42 分钟
AI 数学精要

softmax 与交叉熵:多分类损失的完整推导

从零推导 softmax 交叉熵对 logits 的梯度为何是 s−y,讲清困惑度与标签平滑

  • 能写数值稳定的 softmax,并手算概率、交叉熵损失与梯度
  • 自己独立把 ∂L/∂z = s−y 推导出来,讲明白为什么错误越大、梯度就越大
  • 解释分类为何用交叉熵而非 MSE,理解饱和与梯度尺度
  • 会算困惑度 PPL,说清标签平滑解决什么问题

多分类网络最后一层几乎固定是「softmax + 交叉熵」,这一节把它的梯度亲手推出来

模型最后输出一串任意实数 logits,softmax 把它们变成和为 1 的类别概率,交叉熵再把概率和 one-hot 标签变成标量损失。为什么是这个组合、它的梯度长什么样、为什么比 MSE 更适合分类?这一节全部用推导和手算回答。因果链:logits 经 softmax 归一为概率 → one-hot 下交叉熵退化为 −log sc → 对 logits 求导得到极简的 s−y → 对比 MSE 说明其优势 → 用困惑度度量语言模型、用标签平滑校准过度自信。

3.1 softmax:把任意实数压成和为 1 的概率

softmax 对 logits z=(z₁…zK) 定义 sk=ez_k/Σⱼez_j指数保证结果非负,除以总和保证 Σsk=1;它是「软」的 argmax——最大的 logit 对应最大概率,但不会把概率全占了,两个 logit 差距越大,结果越接近 one-hot。直接算 ez 有问题,z 很大的时候数值会溢出。数值稳定的做法是先减最大值:sk=ez_k−m/Σez_j−m(m=max z),分子分母同乘 e−m,结果不变,但指数不会再爆炸。

手算 z=[2.0,1.0,0.1]:减不减最大值结果相同,e²≈7.389、e¹≈2.718、e0.1≈1.105,总和≈11.212,于是 s≈[0.659,0.242,0.099],和为 1。注意 softmax 对整体加常数不变(只关心 logits 之差),所以最后一层常不设偏置也能正常分类。

softmax 有两个必须会用的代数性质。第一个是平移不变:给所有 logit 同加常数 c,sk=ez_k+c/Σez_j+c=ec·ez_k/(ec·Σez_j)=原式,所以只有 logits 之差有意义,这也正是「减去最大值做稳定化」合法的根据。第二个是它和 argmax 的关系:带温度 T 时 sk=ez_k/T/Σez_j/T,T→0 收敛为 one-hot(最大分量为 1)、T→∞ 收敛为均匀分布,普通 softmax 是 T=1 的折中;argmax 不可导而 softmax 处处可导,需要「软的最大值」做梯度优化时就用它。

填空题填写空白处的代码
z=[2.0,1.0,0.1],正确类 c=0 # softmax s0=e^2/(e^2+e^1+e^0.1) ≈ (填 0.659) # 交叉熵 L=-ln(s0) ≈ (填 0.417) # 对正确类 logit 的梯度 s0-1 ≈ (填 -0.341)

3.2 one-hot 下交叉熵退化为 −log s_c

分类标签是 one-hot 向量 y,正确类 c 处为 1、其余为 0。交叉熵 L=−Σₖyₖln sₖ 里只有 yc=1 的那一项能留下来,所以 L=−ln sc损失只看模型给正确类的概率:给得越接近 1,损失越接近 0;给得越低,损失越大,而且是按对数急剧往上窜的。这就是 l2 恒等式在 H(P)=0 时的特例,也是最大似然的逻辑——最大化正确类的似然 sc,跟最小化 −ln sc 是一回事。

沿用上例 sc=0.659,L=−ln0.659≈0.417。两个极端帮你建立直觉:要是模型完美,sc=1,那 L=−ln1=0;要是把正确类概率压到 0.01,L=−ln0.01≈4.605,梯度也跟着变大,逼着模型快速纠正。交叉熵对「严重错误」给的是强梯度,这刚好是分类训练想要的。

预测输出
若模型对正确类的 softmax 概率 s_c=1(其余为 0),交叉熵损失等于多少?

3.3 核心推导:交叉熵对 logits 的梯度恰为 s−y

这是必须会独立推的结果。由 L=−ln sc 与 ln sc=zc−lnΣⱼez_j,分两种情况对 zk 求导。当 k=c:∂ln sc/∂zc=1−ez_c/Σ=1−sc,故 ∂L/∂zc=−(1−sc)=sc−1。当 k≠c:∂ln sc/∂zk=−ez_k/Σ=−sk,故 ∂L/∂zk=sk

两种情况用 one-hot yk 合并成一个极简式:∂L/∂zk=sk−yk手算上例(y=[1,0,0])梯度=s−y=[−0.341,0.242,0.099]:正确类得到负梯度(−0.341),把它的 logit 往上推;每个错误类得到与其概率成正比的正梯度(0.242、0.099),把它们往下压。梯度之和为 0,这是 softmax 平移不变性的体现。预测越错(sc 越小),正确类梯度 sc−1 越接近 −1,更新越猛;预测越对,梯度越接近 0,自然停止。

推导

为什么分类用交叉熵而不是 MSE:梯度会不会饱和

softmax 后面接最小平方误差,也就是 L_MSE=½Σ(sk−yk)²,链式求导的时候会多一个雅可比因子,梯度里会带 sk(δ−sj) 这一项。 问题出在这:模型错得离谱的时候,正确类的 sc≈0,softmax 处在饱和区,这个因子会趋近 0,直接导致「越错越学不动」的梯度饱和。 换成交叉熵就不一样,它的梯度是 s−y,sc≈0 的时候,正确类的梯度≈−1,错得越狠,学习信号反而越强。 再加上交叉熵本身是 MLE、概率校准效果更好、优化面近似凸,多分类任务几乎全用 softmax-CE 组合;MSE 配 softmax 只在少数回归型概率任务里才会出现。 记个结论就行:交叉熵把 softmax 的导数「约掉了」,留下的是和误差线性相关的干净梯度。

把「交叉熵如何约掉 softmax 导数」用链式法则再推一遍,你会更信这个结果。单个 softmax 的雅可比为 ∂sk/∂zj=skkj−sj):对角元 sk(1−sk)、非对角元 −s_ksj交叉熵对概率的导数 ∂L/∂sk=−yk/sk。链式合并 ∂L/∂zjk(−yk/sk)·skkj−sj)=−yj+sj·Σk yk,one-hot 时 Σyk=1,于是等于 sj−yj复杂的雅可比与 1/s 恰好抵消,只剩误差本身,这也是框架把 LogSoftmax 与 NLLLoss 融合实现、数值更稳的代数原因。

示例代码(可运行)

3.4 困惑度:语言模型的「等概率候选数」

语言模型逐词预测分布,用平均交叉熵衡量,取指数得到困惑度 PPL=exp(平均 NLL)(nat)或 2平均交叉熵 bit它的解释很直观:模型在每个位置相当于在 PPL 个等概率候选词之间犹豫。PPL 越小越好:PPL=1 表示完全确定、百发百中;PPL=10 相当于平均在 10 个等可能词里选。手算例子:若平均交叉熵(nat)=ln10≈2.303,则 PPL=e2.303=10。

注意PPL 和交叉熵是单调一一对应的。降低 PPL 就是降低交叉熵。大家习惯报 PPL,只是因为「候选词个数」比「奈特数」更符合直觉。不同分词器(tokenizer)会改变词表和每个词的概率。跨模型比 PPL 必须用相同分词和相同测试集,否则数字不可比。

记三个口算锚点就行:平均每词负对数似然是 ln2≈0.693 nat 的时候,PPL=2(对应两个候选);是 ln10≈2.303 的时候,PPL=10;是 0 的时候,PPL=1。字符级模型还常用 bits-per-character=平均交叉熵(bit)/字符数,它和词级 PPL 粒度不一样,不能直接比,得先统一到同一 token 粒度再下结论。

选择题

关于困惑度 PPL,下列说法正确的是?

3.5 标签平滑:别让模型把概率押成 0 和 1

one-hot 硬目标逼模型给正确类打 1、其余类打 0,会导致过度自信(over-confidence)、logits 越拉越大、泛化与校准变差。标签平滑把目标改了:每一类都留 ε/K 的保底概率,公式是 y_smooth=(1−ε)y+ε/K——正确类变成 (1−ε)+ε/K,其余各类变成 ε/K。K=3、ε=0.1 时,目标从 [1,0,0] 变成 [0.933,0.033,0.033](各类概率和仍为 1)。它相当于给分布加了个均匀先验,属于一种正则,约束模型别把错误类的概率压到 0,在图像分类、翻译任务里常能提升校准效果和鲁棒性。代价是训练达不到 0 损失——因为目标本身就不是 one-hot 形式的。

标签平滑为什么能改善校准?硬目标训练到后期,正确类 logit 会趋向无穷大——因为 CE 只在 sc=1 时才取 0,模型对错误答案「过度确定」,置信度远高于真实正确率,期望校准误差会变差。平滑后正确类目标概率有上界 (1−ε)+ε/K<1,最优 logit 被拉回有限值,置信度和真实正确率更一致。代价是训练交叉熵的下界变成 −ln(1−ε+ε/K)>0,看到损失稳定降不到 0 别误判为没收敛——那正是平滑在起作用。

负对数似然,梯度是 s−y,和误差线性相关。误差大的时候梯度也大,不会出现梯度饱和的问题。它的输出概率可以校准,是分类任务的默认损失函数。

配对题把部件对到它的作用
找 Bugsoftmax 必须减最大值防 exp 溢出;分类损失用交叉熵以避免 softmax+MSE 的梯度饱和。
# 手写 softmax-CE 的两个高频 bug import math z=[1000.,1001.,1002.] # 错误一:不减最大值,e^1002 直接 OverflowError s=[math.exp(v) for v in z] # 错误二:训练分类时在 softmax 输出上再套 MSE,严重错误处梯度饱和学不动
🐍把整条推导链装进脑子

MLE→负对数似然→one-hot 下 −ln sc→求导得 s−y,这四步是分类问题的统一叙事,逻辑回归、softmax 回归、Transformer 分类头完全同构,只是 logits 的来源从线性函数换成了网络。面试让你手推 softmax 交叉熵梯度,按 3.3 分 k=c 与 k≠c 两种情况写,再合并 s−y,就是满分答案。

⚠️三个数值陷阱

exp 上溢:务必减 max;log(0):配合 log-softmax 或加 eps,框架里用 CrossEntropyLoss=LogSoftmax+NLL 融合实现,比先 softmax 再 log 更稳;类别极不平衡时纯 CE 会偏向多数类,需加类别权重或 focal loss,别误以为是模型容量问题。

ℹ️softmax 的温度与蒸馏伏笔

带温度 T 的 softmax:sk=ez_k/T/Σez_j/TT=1 是普通 softmax;T→∞ 分布趋于均匀、暴露类间相似性(暗知识);T→0 趋于 argmax。l4 知识蒸馏用高 T 教师软标签,先记住旋钮位置。

💡三步自查分类头梯度

loss 不降时:打印 softmax 概率,看是不是等于 1,有没有出现 NaN(溢出/log0 导致的);看正确类概率是不是在上升,错误类的梯度方向对不对;初始化阶段各类的 s≈1/K,正确类梯度≈−(1−1/K),要是初始梯度就接近 0,多半是损失接错了,或者 softmax 重复做了两次。

选择题

softmax 接交叉熵,损失对第 k 个 logit 的梯度是?

本节小结

一条推导链:softmax sk=ez_k/Σez_j(减 max 防溢)把 logits 归一为概率(手算 [0.659,0.242,0.099])→ one-hot 下交叉熵 L=−ln sc(=0.417,完美时为 0)→ 分两种情况求导合并得 ∂L/∂z=s−y=[−0.341,0.242,0.099],错误越大梯度越大、分量和为 0 → 相比之下 softmax+MSE 会梯度饱和,故分类用 CE → PPL=exp(平均CE) 表等概率候选数、标签平滑用 (1−ε)y+ε/K 防过度自信。

资深工程师加餐

底层原理 · 大厂视角 · 工程经验,点卡片展开

一条样本是一个特征向量,一批样本堆成矩阵,神经网络一层的变换本质就是矩阵乘法加激活。换基/特征值分解相当于找数据的主要方向(PCA 降维),GPU 之所以适合深度学习,正是因为它能大规模并行做矩阵运算。把「向量=对象、矩阵=变换」建立起直觉,后面公式就不再抽象。