AutoIRT/BERT/bert_train.py

160 lines
5.2 KiB
Python

import pandas as pd
import numpy as np
from sklearn.model_selection import train_test_split
from transformers import BertTokenizer, BertForSequenceClassification, Trainer, TrainingArguments
import torch
from sklearn.metrics import f1_score, accuracy_score, precision_score, recall_score
import logging
import os
import data_split
train_df, val_df, test_df = data_split.split_dataset("dataset/new_autoirt.csv")
# 打印数据集大小
print(len(train_df), len(test_df), len(val_df))
# 加载tokenizer
tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')
# 编码数据
train_encodings = tokenizer(list(train_df['element']), truncation=True, padding=True)
val_encodings = tokenizer(list(val_df['element']), truncation=True, padding=True)
test_encodings = tokenizer(list(test_df['element']), truncation=True, padding=True)
# 定义数据集类
class DefineDataset(torch.utils.data.Dataset):
def __init__(self, encodings, labels):
self.encodings = encodings
self.labels = labels
def __getitem__(self, idx):
item = {key: torch.tensor(val[idx]) for key, val in self.encodings.items()}
item['labels'] = torch.tensor(int(self.labels[idx]))
return item
def __len__(self):
return len(self.labels)
# 创建数据集实例
train_dataset = DefineDataset(train_encodings, list(train_df['label']))
val_dataset = DefineDataset(val_encodings, list(val_df['label']))
test_dataset = DefineDataset(test_encodings, list(test_df['label']))
# 加载模型
model = BertForSequenceClassification.from_pretrained('bert-base-uncased', num_labels=6)
# 定义训练参数
training_args = TrainingArguments(
output_dir='./BERT_training_output',
evaluation_strategy="epoch",
learning_rate=5e-5,
num_train_epochs=8,
warmup_steps=2000,
per_device_train_batch_size=8,
per_device_eval_batch_size=32,
seed=120,
optim="adamw_torch",
save_strategy="epoch",
load_best_model_at_end=True,
metric_for_best_model='accuracy',
)
def preprocess_logits_for_metrics(logits, labels):
if isinstance(logits, tuple):
# Depending on the model and config, logits may contain extra tensors,
# like past_key_values, but logits always come first
logits = logits[0]
return logits
# 定义文件名
output_dir = "./metrics/"
os.makedirs(output_dir, exist_ok=True) # 创建目录
output_file = os.path.join(output_dir, "val_metrics.txt")
file_labels = {
'describe': 0,
'expected': 1,
'reproduce': 2,
'actual': 3,
'environment': 4,
'additional': 5
}
def compute_metrics(eval_pred):
predictions, labels = eval_pred
y_pred = np.argmax(predictions, axis=1)
y_true = labels
# 计算每个类别的指标
f1_scores = f1_score(y_true, y_pred, average=None)
accuracy = accuracy_score(y_true, y_pred)
precision_scores = precision_score(y_true, y_pred, average=None)
recall_scores = recall_score(y_true, y_pred, average=None)
# 计算总体指标
overall_f1 = f1_score(y_true, y_pred, average='weighted')
overall_accuracy = accuracy_score(y_true, y_pred)
overall_precision = precision_score(y_true, y_pred, average='weighted')
overall_recall = recall_score(y_true, y_pred, average='weighted')
# 将数字标签转换为对应的文字标签
label_texts = {v: k for k, v in file_labels.items()}
# 写入文件
with open(output_file, "a") as f:
f.write("Class Metrics:\n")
for i in range(len(f1_scores)):
label_name = label_texts[i] # 使用文字标签
f.write(f"Class: {label_name}\n")
f.write(f"F1 Score: {f1_scores[i]}\n")
f.write(f"Accuracy: {accuracy_score(y_true[y_true == i], y_pred[y_true == i])}\n")
f.write(f"Precision: {precision_scores[i]}\n")
f.write(f"Recall: {recall_scores[i]}\n")
f.write("\n")
# 写入总体指标
f.write("Overall Metrics:\n")
f.write(f"Overall F1 Score: {overall_f1}\n")
f.write(f"Overall Accuracy: {overall_accuracy}\n")
f.write(f"Overall Precision: {overall_precision}\n")
f.write(f"Overall Recall: {overall_recall}\n")
f.write("--------------------------------------------------\n")
# 保存每个类别的指标
metrics = {}
for i in range(len(f1_scores)):
label_name = label_texts[i] # 使用文字标签
metrics[label_name] = {
'f1_score': f1_scores[i],
'accuracy': accuracy_score(y_true[y_true == i], y_pred[y_true == i]),
'precision': precision_scores[i],
'recall': recall_scores[i]
}
# 保存总体指标
metrics['overall'] = {
'f1_score': overall_f1,
'accuracy': overall_accuracy,
'precision': overall_precision,
'recall': overall_recall
}
return {
'accuracy': overall_accuracy, # 返回准确度作为评估指标
}
# 创建Trainer实例
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=val_dataset,
compute_metrics=compute_metrics,
preprocess_logits_for_metrics=preprocess_logits_for_metrics, # 预处理评估指标
)
# 训练模型
trainer.train()