You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

训练BertForSequenceClassification时损失与准确率无变化的问题排查

问题分析与解决方案

核心错误(最可能的原因)

你的验证集编码完全用错了数据:

val_encodings = tokenizer(val_winners, truncation=True, padding=True)

这里应该传入案件事实文本val_facts,而不是标签数据val_winners。模型训练时用事实文本做输入,验证时却喂了0/1的标签值,相当于让模型从无意义的数字文本里预测标签,自然无法学到有效特征,导致准确率完全不动,损失也没有合理下降。

其他可能的原因

  1. 数据集类别不平衡
    准确率一直固定在0.665,大概率是数据集里某一类的占比刚好是这个数值,模型一直在无脑预测多数类,所以准确率保持不变。你可以统计train_winners和val_winners中0、1的数量占比,验证这个猜测。

  2. 自定义Trainer冗余且可能有隐藏问题
    你重写的CustomTrainer的compute_loss逻辑和Trainer默认实现几乎完全一致,BERTForSequenceClassification本身就用交叉熵损失做二分类,没必要自定义,反而可能引入不必要的bug。

  3. 学习率设置不合理
    TrainingArguments没有指定学习率,默认是5e-5,但BERT微调通常更适合2e-5或1e-5的学习率,过大的学习率可能导致模型无法收敛,损失波动且不下降。

  4. 指标单一且无参考性
    只使用准确率指标,在类别不平衡场景下完全无法反映模型真实性能,你之前注释掉的precision、recall、F1才是更有效的评估指标。

修复建议

  1. 优先修正验证集编码错误
    把验证集编码代码改成:

    val_encodings = tokenizer(val_facts, truncation=True, padding=True)
    
  2. 检查并处理类别不平衡

    • 统计训练集和验证集的标签分布:
      import collections
      print("训练集标签分布:", collections.Counter(train_winners))
      print("验证集标签分布:", collections.Counter(val_winners))
      
    • 如果不平衡,可选择:
      • 在CrossEntropyLoss中设置weight参数,给少数类更高权重
      • 对少数类进行过采样,或对多数类进行欠采样
      • 改用F1-score作为主要评估指标
  3. 简化Trainer,使用默认实现
    删除CustomTrainer类,直接用原生Trainer:

    trainer = Trainer(
        model=model,
        args=training_args,
        train_dataset=train_dataset,
        eval_dataset=val_dataset,
        compute_metrics=compute_metrics,
    )
    
  4. 调整训练参数

    • 设置合适的学习率与训练策略:
      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,  # 保存最优模型,避免过拟合
      )
      
  5. 添加多维度评估指标
    恢复多指标计算:

    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"]
        }
    
  6. 验证数据一致性
    确保train_facts、val_facts与对应的标签train_winners、val_winners是一一对应的,没有出现文本和标签错位的情况。

内容的提问来源于stack exchange,提问作者camdenmcgath

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.23 03:17:06