交叉熵损失(Cross-Entropy Loss)
交叉熵源自信息论,衡量两个概率分布之间的差异。给定真实分布 $p$ 和预测分布 $q$,交叉熵定义为:
所属专题:机器学习基础 (ml-basics·07)
交叉熵损失(Cross-Entropy Loss)
一、原理
交叉熵源自信息论,衡量两个概率分布之间的差异。给定真实分布 和预测分布 ,交叉熵定义为:
核心思想:用预测分布 去编码真实分布 的样本时,所需的平均比特数。 越接近 ,交叉熵越小;两者相等时,交叉熵等于 的熵 ,达到下界。
与 KL 散度的关系:
由于 在训练中是常数,最小化交叉熵等价于最小化 KL 散度,即让预测分布逼近真实分布。
为什么用它做损失函数?
- 对错误分类的惩罚是指数级的( 在 时趋于无穷),梯度大,学得快
- 与 softmax 配合时,梯度形式简洁:,不会出现 sigmoid + MSE 那种梯度消失
- 概率解释清晰:等价于最大似然估计(MLE)
二、计算
1. 二分类交叉熵(BCE, Binary Cross-Entropy)
真实标签 ,预测概率 (通常来自 sigmoid):
对一批 个样本取平均。
2. 多分类交叉熵(Categorical CE)
个类别,真实标签为 one-hot 向量 ,预测分布 :
若真实类别为 (硬标签),则简化为 。
3. 数值稳定的实现
直接计算 会溢出。工程上用 log-sum-exp 技巧 合并 softmax 和 log:
因此框架(PyTorch/TF)中:
- PyTorch:
nn.CrossEntropyLoss直接接受 logits(未过 softmax),内部合并 log-softmax + NLL - PyTorch:
nn.BCEWithLogitsLoss同理接受 logits,比Sigmoid + BCELoss数值更稳 - 不要手动 softmax 再传入
CrossEntropyLoss,会做两次
4. 常见扩展
| 变体 | 用途 |
|---|---|
| 加权交叉熵 | 类别不平衡: |
| Focal Loss | 极端不平衡(如目标检测):,降低易分样本权重 |
| Label Smoothing | 防过拟合:把 one-hot 的 1 换成 ,其余分 |
| KL Divergence Loss | 软标签场景(如知识蒸馏),目标本身是分布而非 one-hot |
三、使用场景
分类任务(最典型)
- 图像分类:ResNet、ViT 等的最后一层 softmax + 交叉熵
- 文本分类:情感分析、意图识别
- 语义分割:逐像素的多分类,等价于每个像素做交叉熵
语言模型
- 自回归 LM(GPT 系):预测下一个 token,词表上的多分类交叉熵,即负对数似然,也是 perplexity 的定义基础:
- BERT MLM:被 mask 位置上的交叉熵
检索与对比学习
- InfoNCE / CLIP 损失:把正样本从一批负样本中”分类”出来,本质是交叉熵
- Softmax 检索:候选集上的多分类
知识蒸馏
- 学生模型拟合教师的软分布,用带温度的交叉熵(或等价的 KL 散度)
强化学习
- 策略梯度中的 log-likelihood 项 与交叉熵同源
- PPO / SFT 阶段直接用交叉熵监督
四、什么时候 不 用交叉熵
- 回归任务:目标是连续值,用 MSE / MAE / Huber
- 排序任务:更适合 pairwise(RankNet)或 listwise(ListNet)损失
- 度量学习:Triplet / Contrastive Loss 更直接
- 输出无法归一化为概率:如生成模型的像素输出(用 L1/L2 或对抗损失)
五、常见坑
- logits 双 softmax:手动 softmax 后再传给
CrossEntropyLoss,梯度会不对 - 标签格式:PyTorch
CrossEntropyLoss要类别索引(LongTensor),不是 one-hot - 忽略 padding:序列任务里 padding 位置要用
ignore_index排除,否则稀释损失 - 类别极度不平衡:直接用会被多数类主导,需加权或换 Focal Loss
- 概率为 0:,务必用 log-sum-exp 稳定实现,或加 裁剪