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

基于PyTorch实现BERT文本分类时混淆矩阵全为0的问题排查

代码问题诊断
  • 核心问题1:评估时使用了未训练的随机初始化模型
    你所有的训练逻辑都是在model这个实例上执行的,但评估时重新初始化了best_model = BERT().to(device),这个新实例没有加载任何训练得到的权重,输出完全随机,自然预测结果全部错误,混淆矩阵全为0。
  • 核心问题2:argmax维度取值错误
    BertForSequenceClassification返回的分类logits是形状为(batch_size, 类别数)的二维张量,你代码中使用torch.argmax(output, 2)取维度2的最大值,而二维张量不存在维度2,得到的预测值完全不符合预期。
  • 次要问题:训练函数变量未初始化
    代码中best_valid_loss没有提前初始化,直接用if best_valid_loss > average_valid_loss判断会触发未定义变量报错,你当前没报错大概率是代码粘贴时出现了截断。
  • 冗余问题:损失函数定义未使用
    你定义了criterion = nn.BCELoss()但全程没有调用,不过BertForSequenceClassification传入labels参数时会自动计算对应交叉熵损失,这个问题不影响运行结果。
修复方案
  1. 训练过程中新增最优权重保存逻辑,在train函数开头初始化best_valid_loss = float('inf'),当验证损失更优时保存权重:
if best_valid_loss > average_valid_loss:
    best_valid_loss = average_valid_loss
    torch.save(model.state_dict(), 'best_bert_cls.pt')
  1. 评估阶段加载训练好的权重,不要重新初始化随机模型:
best_model = BERT().to(device)
best_model.load_state_dict(torch.load('best_bert_cls.pt'))
evaluate(best_model, test_iter)
  1. 修改evaluate函数中的argmax维度:
    将y_pred.extend(torch.argmax(output, 2).tolist())改为
y_pred.extend(torch.argmax(output, 1).tolist())
  1. 补全粘贴时截断的代码(train函数参数、print语句等部分内容显示不全),确保语法正确。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 05:57:03