· 机器学习 ·阅读时长约 4 分钟

交叉熵损失(Cross-Entropy Loss)

交叉熵源自信息论,衡量两个概率分布之间的差异。给定真实分布 $p$ 和预测分布 $q$,交叉熵定义为:

交叉熵损失函数信息论

所属专题:机器学习基础 (ml-basics·07)

交叉熵损失(Cross-Entropy Loss)

一、原理

交叉熵源自信息论,衡量两个概率分布之间的差异。给定真实分布 pp 和预测分布 qq,交叉熵定义为:

H(p,q)=ip(i)logq(i)H(p, q) = -\sum_{i} p(i) \log q(i)

核心思想:用预测分布 qq 去编码真实分布 pp 的样本时,所需的平均比特数。qq 越接近 pp,交叉熵越小;两者相等时,交叉熵等于 pp 的熵 H(p)H(p),达到下界。

与 KL 散度的关系:

H(p,q)=H(p)+DKL(pq)H(p, q) = H(p) + D_{KL}(p \| q)

由于 H(p)H(p) 在训练中是常数,最小化交叉熵等价于最小化 KL 散度,即让预测分布逼近真实分布。

为什么用它做损失函数?

  • 对错误分类的惩罚是指数级的(logq-\log qq0q \to 0 时趋于无穷),梯度大,学得快
  • 与 softmax 配合时,梯度形式简洁:L/zi=qipi\partial L / \partial z_i = q_i - p_i,不会出现 sigmoid + MSE 那种梯度消失
  • 概率解释清晰:等价于最大似然估计(MLE)

二、计算

1. 二分类交叉熵(BCE, Binary Cross-Entropy)

真实标签 y{0,1}y \in \{0, 1\},预测概率 y^(0,1)\hat{y} \in (0, 1)(通常来自 sigmoid):

L=[ylogy^+(1y)log(1y^)]L = -\big[y \log \hat{y} + (1 - y) \log (1 - \hat{y})\big]

对一批 NN 个样本取平均。

2. 多分类交叉熵(Categorical CE)

CC 个类别,真实标签为 one-hot 向量 yy,预测分布 y^=softmax(z)\hat{y} = \text{softmax}(z)

L=c=1Cyclogy^cL = -\sum_{c=1}^{C} y_c \log \hat{y}_c

若真实类别为 kk(硬标签),则简化为 L=logy^kL = -\log \hat{y}_k

3. 数值稳定的实现

直接计算 log(softmax(z))\log(\text{softmax}(z)) 会溢出。工程上用 log-sum-exp 技巧 合并 softmax 和 log:

logy^k=zklogiezi=zk(m+logiezim),m=maxizi\log \hat{y}_k = z_k - \log \sum_{i} e^{z_i} = z_k - \Big(m + \log \sum_i e^{z_i - m}\Big), \quad m = \max_i z_i

因此框架(PyTorch/TF)中:

  • PyTorchnn.CrossEntropyLoss 直接接受 logits(未过 softmax),内部合并 log-softmax + NLL
  • PyTorchnn.BCEWithLogitsLoss 同理接受 logits,比 Sigmoid + BCELoss 数值更稳
  • 不要手动 softmax 再传入 CrossEntropyLoss,会做两次

4. 常见扩展

变体用途
加权交叉熵类别不平衡:L=cwcyclogy^cL = -\sum_c w_c y_c \log \hat{y}_c
Focal Loss极端不平衡(如目标检测):L=(1y^k)γlogy^kL = -(1-\hat{y}_k)^\gamma \log \hat{y}_k,降低易分样本权重
Label Smoothing防过拟合:把 one-hot 的 1 换成 1ϵ1-\epsilon,其余分 ϵ/(C1)\epsilon/(C-1)
KL Divergence Loss软标签场景(如知识蒸馏),目标本身是分布而非 one-hot

三、使用场景

分类任务(最典型)

  • 图像分类:ResNet、ViT 等的最后一层 softmax + 交叉熵
  • 文本分类:情感分析、意图识别
  • 语义分割:逐像素的多分类,等价于每个像素做交叉熵

语言模型

  • 自回归 LM(GPT 系):预测下一个 token,词表上的多分类交叉熵,即负对数似然,也是 perplexity 的定义基础:PPL=eL\text{PPL} = e^{L}
  • BERT MLM:被 mask 位置上的交叉熵

检索与对比学习

  • InfoNCE / CLIP 损失:把正样本从一批负样本中”分类”出来,本质是交叉熵
  • Softmax 检索:候选集上的多分类

知识蒸馏

  • 学生模型拟合教师的软分布,用带温度的交叉熵(或等价的 KL 散度)

强化学习

  • 策略梯度中的 log-likelihood 项 logπ(as)\log \pi(a|s) 与交叉熵同源
  • PPO / SFT 阶段直接用交叉熵监督

四、什么时候 用交叉熵

  • 回归任务:目标是连续值,用 MSE / MAE / Huber
  • 排序任务:更适合 pairwise(RankNet)或 listwise(ListNet)损失
  • 度量学习:Triplet / Contrastive Loss 更直接
  • 输出无法归一化为概率:如生成模型的像素输出(用 L1/L2 或对抗损失)

五、常见坑

  1. logits 双 softmax:手动 softmax 后再传给 CrossEntropyLoss,梯度会不对
  2. 标签格式:PyTorch CrossEntropyLoss 要类别索引(LongTensor),不是 one-hot
  3. 忽略 padding:序列任务里 padding 位置要用 ignore_index 排除,否则稀释损失
  4. 类别极度不平衡:直接用会被多数类主导,需加权或换 Focal Loss
  5. 概率为 0log0=\log 0 = -\infty,务必用 log-sum-exp 稳定实现,或加 ϵ\epsilon 裁剪

评论