XLM-RoBERTa文本分类:早停与多轮评估指标实现方案
PyTorch版XLM-RoBERTa二分类训练改造方案
你当前使用的是PyTorch实现的XLMRobertaForSequenceClassification模型,本身不提供Keras风格的compile/fit高层接口,不需要依赖第三方训练框架,直接在你现有手写训练循环基础上修改即可实现全部需求。
1. 配置Binary Cross-Entropy损失
默认情况下num_labels=2时模型内置损失为交叉熵损失,要使用BCE损失需要做两处调整:
- 调整模型初始化配置,将输出改为单神经元,匹配BCE输入要求:
model = XLMRobertaForSequenceClassification.from_pretrained( "xlm-roberta-base", num_labels = 1, # BCE二分类设置为1,输出单logit output_attentions = False, output_hidden_states = False, ) model.to(device)
- 手动替换损失函数,不要依赖模型内置的自动损失计算:
from torch.nn import BCEWithLogitsLoss # 初始化BCE损失(自带sigmoid实现,不需要提前给输出做激活) bce_loss = BCEWithLogitsLoss()
训练和验证阶段的前向传播逻辑修改为:
# 训练/验证阶段前向传播时不再传入labels参数 outputs = model(b_input_ids, attention_mask=b_input_mask) logits = outputs.logits # 标签转为float类型,形状和logits对齐后计算损失 loss = bce_loss(logits.view(-1), b_labels.float().view(-1))
验证阶段的预测逻辑同步修改,不需要argmax,改用sigmoid转概率后按0.5阈值判定类别:
# 所有验证批次预测结果拼接完成后 val_probs = torch.sigmoid(torch.tensor(stacked_val_preds)).view(-1).numpy() y_pred = (val_probs >= 0.5).astype(int) y_true = np.array(targets_list)
2. 实现patience=15的早停机制
不需要引入额外工具,在训练循环外初始化跟踪变量即可:
# 早停配置 early_stopping_patience = 15 # 监控验证集F1分数(值越大效果越好,若监控损失则初始值设为正无穷) best_val_f1 = -float('inf') patience_counter = 0 # 缓存最优模型权重,避免早停时加载到过拟合的权重 best_model_weights = None
注意:早停patience设为15时,建议将总训练轮次epochs设置为大于30的值,否则训练会在早停触发前正常结束
每个epoch验证结束、计算完所有指标后,加入早停判断逻辑:
if val_f1 > best_val_f1: best_val_f1 = val_f1 patience_counter = 0 # 保存当前最优权重 best_model_weights = {k: v.cpu().clone() for k, v in model.state_dict().items()} else: patience_counter += 1 print(f"验证F1未提升,当前等待计数: {patience_counter}/{early_stopping_patience}") if patience_counter >= early_stopping_patience: print(f"连续{early_stopping_patience}个epoch指标无提升,触发早停,终止训练") # 加载最优权重 model.load_state_dict(best_model_weights) model.to(device) break
训练全部结束后再保存模型,此时保存的是泛化性最优的权重,而非最后一个epoch的权重。
3. 输出每个epoch的全量评估指标
你现有代码已经计算了准确率,只需补充导入sklearn的其他指标计算函数,在验证阶段同步计算recall、f1即可:
# 补充导入 from sklearn.metrics import accuracy_score, recall_score, f1_score
验证阶段得到y_true和y_pred后,直接计算所有指标:
val_acc = accuracy_score(y_true, y_pred) # 若数据集类别不平衡,可替换average参数为'weighted',或指定pos_label为正类标签 val_recall = recall_score(y_true, y_pred, average='binary') val_f1 = f1_score(y_true, y_pred, average='binary') # 打印指标 print(f'Val Loss: {total_val_loss:.4f}') print(f'Val Acc: {val_acc:.4f} | Val Recall: {val_recall:.4f} | Val F1: {val_f1:.4f}')
将上述三个指标补充到training_stats的存储字典中,即可保留所有epoch的运行数据。
原有代码bug修复
你现有代码中training_time的计算位置错误,写在了训练批次循环内部,会导致每个批次都覆盖时间值,需要把这行代码移到训练批次循环结束后、打印训练损失之前:
# 训练批次全部迭代完成后再计算耗时 training_time = format_time(time.time() - t0) print("") print('Train loss:' ,total_train_loss) print(" Training epoch took: {:}".format(training_time))
内容的提问来源于stack exchange,提问作者Rachele Franceschini
相关产品推荐
相关产品推荐

