Views
No views yet
pip install pipenv (先安裝pipenv)
pipenv shell (建置虛擬環境)
pipenv install (已有Pipfile檔案則可直接使用以下指令安裝所需套件)pipenv install pandas numpy torch transformers scikit-learn tqdmgit clone https://github.com/elmmaple/AG_news.gitpython classfication_use_BertForSequenceClassification.py(BertForSequenceClassification版本)
or
python classification_use_BertModel.py(BertModel + 分類器版本)程式將會載入 BERT 模型和分詞器
定義 TextClassification 模型(如用BertForSequenceClassification則不需額外寫分類器)
加載訓練和測試數據集
設定優化器和損失函數
進行模型訓練
評估模型在測試集上的表現
訓練和評估流程使用 pandas 載入train.csv和test.csv</div>
pd.read_csv(XXX_FILE_PATH) 加載預訓練的BERT模型。定義 TextClassification 模型,該模型在 BERT 的基礎上添加了全連接層進行分類
tokenizer = BertTokenizerFast.from_pretrained('bert-base-uncased')創建 AGNewsDataset 類別,處理文本數據,進行分詞並準備成模型可接受的格式。
使用 DataLoader 加載訓練數據,定義優化器和損失函數。
進行多個 epoch 的訓練,計算損失並進行反向傳播優化創建測試數據集並使用 DataLoader 加載。在模型評估模式下,對測試數據進行預測,計算精確度、召回率和 F1 分數等指標。根據測試數據集的預測結果,以下是模型的性能指標:
精確度(Precision):{precision}
召回率(Recall):{recall}
測試準確度(Test Accuracy):{accuracy}
F1 分數:{f1}