训练BertForSequenceClassification时损失与准确率无变化的问题排查
核心错误(最可能的原因)
你的验证集编码完全用错了数据:
val_encodings = tokenizer(val_winners, truncation=True, padding=True)
这里应该传入案件事实文本val_facts,而不是标签数据val_winners。模型训练时用事实文本做输入,验证时却喂了0/1的标签值,相当于让模型从无意义的数字文本里预测标签,自然无法学到有效特征,导致准确率完全不动,损失也没有合理下降。
其他可能的原因
数据集类别不平衡
准确率一直固定在0.665,大概率是数据集里某一类的占比刚好是这个数值,模型一直在无脑预测多数类,所以准确率保持不变。你可以统计train_winners和val_winners中0、1的数量占比,验证这个猜测。自定义Trainer冗余且可能有隐藏问题
你重写的CustomTrainer的compute_loss逻辑和Trainer默认实现几乎完全一致,BERTForSequenceClassification本身就用交叉熵损失做二分类,没必要自定义,反而可能引入不必要的bug。学习率设置不合理
TrainingArguments没有指定学习率,默认是5e-5,但BERT微调通常更适合2e-5或1e-5的学习率,过大的学习率可能导致模型无法收敛,损失波动且不下降。指标单一且无参考性
只使用准确率指标,在类别不平衡场景下完全无法反映模型真实性能,你之前注释掉的precision、recall、F1才是更有效的评估指标。
修复建议
优先修正验证集编码错误
把验证集编码代码改成:val_encodings = tokenizer(val_facts, truncation=True, padding=True)检查并处理类别不平衡
- 统计训练集和验证集的标签分布:
import collections print("训练集标签分布:", collections.Counter(train_winners)) print("验证集标签分布:", collections.Counter(val_winners)) - 如果不平衡,可选择:
- 在
CrossEntropyLoss中设置weight参数,给少数类更高权重 - 对少数类进行过采样,或对多数类进行欠采样
- 改用F1-score作为主要评估指标
- 在
- 统计训练集和验证集的标签分布:
简化Trainer,使用默认实现
删除CustomTrainer类,直接用原生Trainer:trainer = Trainer( model=model, args=training_args, train_dataset=train_dataset, eval_dataset=val_dataset, compute_metrics=compute_metrics, )调整训练参数
- 设置合适的学习率与训练策略:
training_args = TrainingArguments( output_dir="test_trainer", logging_dir='logs', evaluation_strategy="epoch", per_device_train_batch_size=16, per_device_eval_batch_size=32, num_train_epochs=8, # 先减少epoch数,观察收敛情况 logging_steps=50, learning_rate=2e-5, # BERT微调常用学习率 save_strategy="epoch", load_best_model_at_end=True, # 保存最优模型,避免过拟合 )
- 设置合适的学习率与训练策略:
添加多维度评估指标
恢复多指标计算:import evaluate metric = evaluate.load("accuracy", "precision", "recall", "f1") def compute_metrics(eval_pred): logits, labels = eval_pred predictions = np.argmax(logits, axis=1) return { "accuracy": metric.compute(predictions=predictions, references=labels)["accuracy"], "precision": metric.compute(predictions=predictions, references=labels)["precision"], "recall": metric.compute(predictions=predictions, references=labels)["recall"], "f1": metric.compute(predictions=predictions, references=labels)["f1"] }验证数据一致性
确保train_facts、val_facts与对应的标签train_winners、val_winners是一一对应的,没有出现文本和标签错位的情况。
内容的提问来源于stack exchange,提问作者camdenmcgath

