TextClassificationTrainer 類別
定義
重要
部分資訊涉及發行前產品,在發行之前可能會有大幅修改。 Microsoft 對此處提供的資訊,不做任何明確或隱含的瑕疵擔保。
用於 IEstimator<TTransformer> 訓練深度神經網路(DNN)來分類文本。
public class TextClassificationTrainer : Microsoft.ML.TorchSharp.NasBert.NasBertTrainer<uint,long>
type TextClassificationTrainer = class
inherit NasBertTrainer<uint32, int64>
Public Class TextClassificationTrainer
Inherits NasBertTrainer(Of UInteger, Long)
- 繼承
-
TextClassificationTrainer
- 繼承
備註
要建立這個訓練器,請使用 TextClassification。
輸入與輸出欄位
輸入標籤欄位資料必須是 鍵 型別,句子欄位必須是型別 TextDataViewType。
此訓練器會輸出下列欄位:
| 輸出欄位名稱 | 欄位類型 | Description |
|---|---|---|
PredictedLabel |
金鑰 類型 | 預測標籤的索引。 如果其值為 i,則其實際標籤為鍵值輸入的標籤型別中的第 i個類別。 |
Score |
的向量Single | 所有班級的分數。 較高的值表示較高的機率會落入相關聯的類別。 如果 i-th 元素具有最大值,則預測的標籤索引會 i。 注意 是一個 i 以零為基礎的指標。 |
訓練師特性
| 特徵 | Value |
|---|---|
| 機器學習任務 | 多類別分類 |
| 需要正規化嗎? | No |
| 快取是必須的嗎? | No |
| 除了 Microsoft.ML 之外,必須使用 NuGet | Microsoft.ML.TorchSharp 以及 libtorch-cpu 或 libtorch-cuda-11.3,或任何作業系統專用的變體。 |
| 可匯出至 ONNX | No |
訓練演算法細節
利用現有的預訓練 NAS-BERT roBERTa 模型,訓練深度神經網路(DNN)以進行文字分類。
方法
| 名稱 | Description |
|---|---|
| Fit(IDataView) |
用於 IEstimator<TTransformer> 訓練深度神經網路(DNN)來分類文本。 (繼承來源 NasBertTrainer<TLabelCol,TTargetsCol>) |
| GetOutputSchema(SchemaShape) |
用於 IEstimator<TTransformer> 訓練深度神經網路(DNN)來分類文本。 (繼承來源 NasBertTrainer<TLabelCol,TTargetsCol>) |
擴充方法
| 名稱 | Description |
|---|---|
| AppendCacheCheckpoint<TTrans>(IEstimator<TTrans>, IHostEnvironment) |
在估計鏈中附加一個「快取檢查點」。 這將確保下游估計器能針對快取資料進行訓練。 在訓練師接受多次資料通行前設置快取檢查點會很有幫助。 |
| WithOnFitDelegate<TTransformer>(IEstimator<TTransformer>, Action<TTransformer>) |
給定一個估計器,回傳一個包裹物件,該物件會呼叫一次 Fit(IDataView) 代理。 估計器通常回傳擬合的資訊很重要,因此該 Fit(IDataView) 方法回傳一個特定型別的物件,而非一般 ITransformer的 。 然而,同時, IEstimator<TTransformer> 通常會被組成包含許多物件的管線,因此我們可能需要建立一條估計鏈,將 EstimatorChain<TLastTransformer> 我們想要取得變壓器的估計器埋藏在這條鏈的某處。 在這種情況下,我們可以透過此方法附加一個代理,當 fit 被呼叫時會被呼叫。 |