共计 1807 个字符,预计需要花费 5 分钟才能阅读完成。
背景痛点:为什么分类任务不用 MSE?
刚入门深度学习时,很多人会疑惑:既然均方误差(MSE)在回归任务中表现良好,为什么分类任务非要使用交叉熵(CE)损失?这里隐藏着两个关键问题:

-
梯度消失问题:当使用 Sigmoid 激活 +MSE 时,输出概率接近 0 或 1 时梯度会趋近于零,导致参数更新停滞。例如二分类中预测概率 p =0.99 时,MSE 的梯度仅为(1-0.99)0.99(1-0.99)≈0.0001
-
概率解释偏差:MSE 默认预测值与标签是欧式空间关系,而分类任务需要度量的是概率分布差异。比如预测猫的概率 0.6 vs 0.8,与 0.1 vs 0.3 的差值相同,但前者分布更接近真实标签[1,0]
数学本质:从信息论到 KL 散度
交叉熵的本质是衡量两个概率分布的差异,其推导路径如下:
-
信息熵:描述事件的不确定性,$H(p)=-\sum p(x)\log p(x)$。例如公平硬币抛掷的熵是 1bit
-
KL 散度:衡量两个分布的差异,$D_{KL}(p||q)=\sum p(x)\log\frac{p(x)}{q(x)}$
-
交叉熵分解:$H(p,q)=H(p)+D_{KL}(p||q)$。由于真实分布 p 是 one-hot 向量(固定熵值为 0),因此最小化 CE 等价最小化 KL 散度
实际计算时需要注意:
- 网络输出通常是 logits(未归一化的分数),需要通过 Softmax 转换:$p_i=e^{z_i}/\sum e^{z_j}$
- 为避免数值溢出,PyTorch 内部使用 log_softmax 技巧:$\log p_i = z_i – \text{logsumexp}(z)$
PyTorch 三大实现方式
方式 1:手动分解实现(理解原理)
def manual_ce(logits: torch.Tensor, labels: torch.Tensor):
# logits.shape = [B, C], labels.shape = [B]
log_probs = torch.log_softmax(logits, dim=1) # 数值稳定
return torch.nn.functional.nll_loss(log_probs, labels) # 负对数似然
方式 2:直接调用内置函数(生产推荐)
# 自动处理 softmax + negative log likelihood
loss = torch.nn.CrossEntropyLoss()(logits, labels)
方式 3:标签平滑实现(正则化)
def label_smoothing_ce(logits, labels, alpha=0.1):
# labels 从 [N] 变为 [N, C] 的平滑分布
smoothed = torch.full_like(logits, alpha/(logits.size(1)-1))
smoothed.scatter_(1, labels.unsqueeze(1), 1-alpha)
return -(smoothed * torch.log_softmax(logits, dim=1)).sum(dim=1).mean()
CIFAR-10 实战观察
用 PyTorch Lightning 搭建实验:
-
学习率影响:当 lr=0.1 时,CE 损失快速下降但出现震荡;lr=0.01 时收敛稳定但速度较慢
-
损失曲线解读:初期损失下降快(调整决策边界),后期缓慢(微调分类面)
-
典型数值范围:CIFAR-10 初始 CE≈2.3(-ln(0.1)),收敛时约 0.5
生产环境避坑指南
- 数值稳定技巧:
- 永远不要在 softmax 前单独计算 exp,使用 log_softmax 或 CrossEntropyLoss
-
混合精度训练时,对 CE 损失保持 FP32 计算
-
类别不平衡处理:
# 根据类别频率设置权重 weights = torch.tensor([0.1, 0.9]) loss = nn.CrossEntropyLoss(weight=weights.to(device)) -
多标签分类方案:
- 改用 sigmoid 激活 + 二元交叉熵(BCE)
-
每个类别独立判断,输出层节点数 = 类别数
-
分布式训练同步:
- 确保所有进程的损失计算使用相同的权重
- 使用 DistributedSampler 保证数据划分一致
延伸思考
当验证集 loss 持续下降但准确率停滞时,可能的原因包括:
– 模型在过度优化容易样本(主要影响 loss)
– 存在标签噪声干扰了评估
– 需要调整分类阈值或引入 F1 等指标
建议尝试:
– 可视化混淆矩阵
– 检查预测结果的置信度分布
– 在验证集上计算 ECE(预期校准误差)
