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

PyTorch二分类模型全预测为多数类问题排查求助

二分类不平衡问题排查与解决

1. CrossEntropyLoss的常见误用点

  • 未设置类别权重:默认CrossEntropyLoss对所有类别权重均等,在类别不平衡场景下,多数类的损失占比会碾压少数类,模型只需全预测多数类就能获得低损失,自然会偏向这种偷懒的预测方式。正确做法是给少数类分配更高权重:
    import torch.nn as nn
    # 假设0为多数类(占比90%),1为少数类(占比10%)
    class_weights = torch.tensor([0.1, 0.9], dtype=torch.float)
    # 确保权重与模型在同一设备(GPU/CPU)
    criterion = nn.CrossEntropyLoss(weight=class_weights.to(device))
    
  • 输出层与标签格式不匹配:二分类任务中,CrossEntropyLoss要求模型输出为(batch_size, 2)的二维张量,标签为一维的类别索引(0或1)。如果错误地给输出层加sigmoid输出单值,或用one-hot编码作为标签,都会导致损失计算逻辑混乱。正确的输出层写法:
    # 二分类输出层无需额外加sigmoid,CrossEntropyLoss内置softmax计算
    self.fc = nn.Linear(hidden_dim, 2)
    
    若标签是one-hot格式,需转换为类别索引:
    # 将(batch_size, 2)的one-hot标签转为(batch_size,)的索引
    labels = torch.argmax(labels, dim=1)
    

2. 类别不平衡的补充处理策略

  • 数据层面:对少数类进行过采样(复制样本)、对多数类进行欠采样(随机丢弃部分样本),或用SMOTE算法合成少数类的新样本。
  • 评估指标:放弃仅用准确率判断模型效果,改用精确率、召回率、F1分数或AUC-ROC,这些指标能更准确反映模型对少数类的识别能力。
  • 训练技巧:
    • 调小学习率,避免模型快速收敛到全预测多数类的局部最优;
    • 加入早停机制,当验证集上的召回率不再提升时停止训练;
    • 在模型中加入Dropout层,降低模型对多数类样本的过拟合程度。

3. 代码排查核心要点

如果上述调整无效,检查以下细节:

  • 确认标签加载逻辑:少数类的标签是否被正确标记,有没有出现类别颠倒的情况;
  • 打印训练过程中的预测分布:每个batch的预测结果里,少数类的预测占比是否始终为0;
  • 验证权重设备一致性:确保类别权重和模型、数据在同一设备,否则权重会失效;
  • 检查梯度更新:打印模型参数的梯度值,确认是否存在梯度消失或爆炸导致模型无法更新。

内容的提问来源于stack exchange,提问作者АЛлександр Рудинский

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 14:34:53