① 为什么不用 MSE
如果对 softmax 输出用均方误差,梯度会包含 p(1−p) 因子。当预测非常错误(p 接近 0 或 1)时,这个因子会趋近于 0,导致梯度消失,模型几乎学不动。
上一页讲了自动微分,它像一台发动机,能把“误差对每个参数的敏感度”算出来。但发动机还需要燃料——这个燃料就是损失函数。在语言模型里,这个损失函数就是交叉熵。这一页会把模型从“说出 logits”到“知道自己错在哪里、每个方向该改多少”这条完整链路,拆成你可以亲手操作的实验。
模型做完一次前向传播后,最后会输出一组logits——每个候选 token 一个分数,分数越高代表模型越看好它。但分数本身不是概率,也不能直接告诉我们“模型错得有多离谱”。我们需要一个标量指标,把“预测分布”和“正确答案”之间的差距量化出来,这样才能反向传播、更新参数。
如果对 softmax 输出用均方误差,梯度会包含 p(1−p) 因子。当预测非常错误(p 接近 0 或 1)时,这个因子会趋近于 0,导致梯度消失,模型几乎学不动。
交叉熵只关心“正确类别的概率有多高”。正确类概率越接近 1,损失越接近 0;越接近 0,损失越大,惩罚越猛烈。
softmax + 交叉熵组合后,反向传播的梯度化简为 p − y。没有额外因子,数值稳定,实现简单,这是它成为分类任务标配的重要原因。
下面这个实验台模拟一个语言模型在某个位置的输出。词汇表有 4 个候选 token,模型先给它们打分(logits),然后经过 softmax 变成概率,再用交叉熵和正确答案比较,最后算出每个 logit 的梯度。你可以调整 logits,观察每一环怎么变。
softmax 用指数函数把任意实数变成正数,再除以总和。分数越高的 token,概率越大,所有概率加起来等于 1。
交叉熵只取正确类别的概率,取负对数。正确类概率越接近 1,损失越小;越接近 0,损失越大。
交叉熵对 logits 的梯度有一个非常优美的结果:概率减去 one-hot 目标。这就是 softmax + 交叉熵组合之所以流行的核心原因。
| Token | logit | 概率 p | 目标 y | 梯度 p − y | 方向 |
|---|
每次“梯度下降一步”,就是让每个 logit 沿着梯度的反方向走一小步。重复这个过程,正确 token 的 logit 会升高,错误 token 的 logit 会降低,损失随之下降。
这个结果看起来像魔法,但其实只需要几步链式法则。下面把它完整推一遍,你会看到它为什么如此干净。
交叉熵损失 L = −Σj yj · log(pj)。因为 y 是 one-hot,只有正确类那一项不为零,所以 L = −log(ptarget)。
log(pk) = logitk − logsumexp(logits)。所以 L = −logittarget + logsumexp(logits)。这一步把 softmax 和 log 合在一起,数值上更稳定。
对 logitk 求导:
· 如果是正确类:∂L/∂logittarget = −1 + ptarget = p − 1
· 如果是其他类:∂L/∂logitk = 0 + pk = p
合并起来就是 ∂L/∂logits = p − y。
真实语言模型不是只预测一个位置,而是每个位置都预测下一个 token。假设输入是“今天天气很”,模型要预测下一个 token。如果有 5 个位置,就有 5 个交叉熵损失,把它们平均一下,就是这次前向传播的总损失。
训练时,模型在位置 t 的输入是真实的前缀,而不是自己上一步的预测。这样每个位置都能并行计算损失,训练更稳定高效。
像 padding token、prompt 部分有时不计入损失。用一个 mask 把不需要的位置置零,只对真正需要预测的位置求平均。
算出 ∂L/∂logits 后,再往前一层传:∂L/∂W = (∂L/∂logits) · hᵀ,∂L/∂h = Wᵀ · (∂L/∂logits)。梯度就这样一层层传回整个 Transformer。
下面用六张卡片,把这一页涉及的每一个概念完整讲清楚。你可以按顺序读,也可以挑自己最关心的部分看。
交叉熵(Cross-Entropy)是衡量两个概率分布差异的指标。设真实分布为 y,预测分布为 p,则 H(y, p) = −Σ yi log(pi)。在分类任务中,y 通常是 one-hot,所以交叉熵退化为 −log(ptarget)。
反向传播(Backpropagation)是利用链式法则,从损失出发,逐层计算每个参数梯度的算法。当损失是交叉熵、最后一层是 softmax 时,对 logits 的梯度化简为 p − y,这个结果既简洁又数值稳定。
想象你在玩猜词游戏。你给每个候选词打一个“我觉得它有多对”的分数。交叉熵就是一位严格的裁判:它只看你给正确答案打了多少分。你给正确答案的分数越高,它越满意;你越自信地猜错,它越严厉地惩罚你。
反向传播则是裁判给你的反馈:它会告诉你,每个候选词的分数应该调高还是调低。而且反馈非常直接——你给某个词的概率比它“应得”的高多少,就往下压多少;比它“应得”的低多少,就往上抬多少。这就是 p − y 的直觉。
它解决了“如何衡量模型预测分布与真实标签的差距,并高效回传梯度”这两个问题。
交叉熵 + 反向传播是几乎所有现代神经网络训练的核心引擎。语言模型的预训练、微调、RLHF,本质上都在最小化某种形式的交叉熵。
它让模型可以:
没有交叉熵,模型就没有明确的优化目标;没有反向传播,梯度就无法从损失传到参数。两者缺一不可。
logits 是原始分数,softmax 把它们压成合法概率分布。指数放大差距,归一化保证总和为 1。
因为标签是 one-hot,损失就是 −log(ptarget)。正确类概率越低,惩罚越猛烈。
正确类梯度为负(抬高 logit),错误类梯度为正(压低 logit)。概率越大的错误类,被压得越狠。