基于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参数时会自动计算对应交叉熵损失,这个问题不影响运行结果。
修复方案
- 训练过程中新增最优权重保存逻辑,在
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')
- 评估阶段加载训练好的权重,不要重新初始化随机模型:
best_model = BERT().to(device) best_model.load_state_dict(torch.load('best_bert_cls.pt')) evaluate(best_model, test_iter)
- 修改evaluate函数中的argmax维度:
将y_pred.extend(torch.argmax(output, 2).tolist())改为
y_pred.extend(torch.argmax(output, 1).tolist())
- 补全粘贴时截断的代码(train函数参数、print语句等部分内容显示不全),确保语法正确。
内容的提问来源于stack exchange,提问作者coding zombie
相关产品推荐
相关产品推荐

