Page 31 · SimLabs LLM Visual

交叉熵与反向传播:模型如何知道自己错了

上一页讲了自动微分,它像一台发动机,能把“误差对每个参数的敏感度”算出来。但发动机还需要燃料——这个燃料就是损失函数。在语言模型里,这个损失函数就是交叉熵。这一页会把模型从“说出 logits”到“知道自己错在哪里、每个方向该改多少”这条完整链路,拆成你可以亲手操作的实验。

看懂 logits 与概率 理解交叉熵为什么这样设计 掌握 p − y 梯度公式 把损失接到参数更新

为什么需要交叉熵

模型做完一次前向传播后,最后会输出一组logits——每个候选 token 一个分数,分数越高代表模型越看好它。但分数本身不是概率,也不能直接告诉我们“模型错得有多离谱”。我们需要一个标量指标,把“预测分布”和“正确答案”之间的差距量化出来,这样才能反向传播、更新参数。

隐藏状态 h → 线性层 W → logits → softmax → 概率 p → 交叉熵损失 L

① 为什么不用 MSE

如果对 softmax 输出用均方误差,梯度会包含 p(1−p) 因子。当预测非常错误(p 接近 0 或 1)时,这个因子会趋近于 0,导致梯度消失,模型几乎学不动。

② 交叉熵的直觉

交叉熵只关心“正确类别的概率有多高”。正确类概率越接近 1,损失越接近 0;越接近 0,损失越大,惩罚越猛烈。

③ 梯度极其干净

softmax + 交叉熵组合后,反向传播的梯度化简为 p − y。没有额外因子,数值稳定,实现简单,这是它成为分类任务标配的重要原因。

先抓住一句话: 交叉熵是语言模型的“痛苦计量器”。它不看模型说了什么,只看模型给正确 token 分配了多少概率——分配得越少,痛苦越大,梯度越强,参数就被推得越狠。

实验台:从 logits 到梯度

下面这个实验台模拟一个语言模型在某个位置的输出。词汇表有 4 个候选 token,模型先给它们打分(logits),然后经过 softmax 变成概率,再用交叉熵和正确答案比较,最后算出每个 logit 的梯度。你可以调整 logits,观察每一环怎么变。

调模型输出的 logits

1Softmax:把 logits 变成概率分布

softmax 用指数函数把任意实数变成正数,再除以总和。分数越高的 token,概率越大,所有概率加起来等于 1。

// 数值稳定版 softmax
m = max(logits) // 减去最大值防止溢出
pi = exp(logiti − m) / Σj exp(logitj − m)

2交叉熵:量化预测与目标的差距

交叉熵只取正确类别的概率,取负对数。正确类概率越接近 1,损失越小;越接近 0,损失越大。

当前损失 L
L = −log(p好)
—
−log(p) 曲线:正确类概率越低,惩罚越陡

3反向传播:每个 logit 该往哪边调

交叉熵对 logits 的梯度有一个非常优美的结果:概率减去 one-hot 目标。这就是 softmax + 交叉熵组合之所以流行的核心原因。

∂L/∂logiti = pi − yi
// y 是 one-hot:正确类为 1,其他为 0
Token logit 概率 p 目标 y 梯度 p − y 方向
怎么读这张表: 正确 token 的梯度是 p − 1 < 0,所以更新时会抬高它的 logit;错误 token 的梯度是 p − 0 > 0,所以更新时会压低它的 logit。概率越高的错误 token,被压得越狠。

4训练演示:损失如何一步步下降

每次“梯度下降一步”,就是让每个 logit 沿着梯度的反方向走一小步。重复这个过程,正确 token 的 logit 会升高,错误 token 的 logit 会降低,损失随之下降。

点击「自动训练 30 步」查看 loss 曲线
训练步数 → —

为什么梯度恰好是 p − y

这个结果看起来像魔法,但其实只需要几步链式法则。下面把它完整推一遍,你会看到它为什么如此干净。

1

写出交叉熵的展开形式

交叉熵损失 L = −Σj yj · log(pj)。因为 y 是 one-hot,只有正确类那一项不为零,所以 L = −log(ptarget)。

2

把 softmax 代入 log 里

log(pk) = logitk − logsumexp(logits)。所以 L = −logittarget + logsumexp(logits)。这一步把 softmax 和 log 合在一起,数值上更稳定。

3

对每个 logit 求偏导

对 logitk 求导:
· 如果是正确类:∂L/∂logittarget = −1 + ptarget = p − 1
· 如果是其他类:∂L/∂logitk = 0 + pk = p
合并起来就是 ∂L/∂logits = p − y。

一句话理解: 交叉熵 + softmax 的梯度之所以这么干净,是因为 log 和 exp 互为逆运算,求导后正好抵消,只剩下“预测概率 − 真实标签”。这也是为什么几乎所有的语言模型、分类模型都用这套组合。

在语言模型里,这套机制怎么用

真实语言模型不是只预测一个位置,而是每个位置都预测下一个 token。假设输入是“今天天气很”,模型要预测下一个 token。如果有 5 个位置,就有 5 个交叉熵损失,把它们平均一下,就是这次前向传播的总损失。

位置 1:今 → 预测「天」 → L₁ | 位置 2:天 → 预测「气」 → L₂ | … | 总损失 L = mean(L₁, L₂, …)

Teacher Forcing

训练时,模型在位置 t 的输入是真实的前缀,而不是自己上一步的预测。这样每个位置都能并行计算损失,训练更稳定高效。

Loss Masking

像 padding token、prompt 部分有时不计入损失。用一个 mask 把不需要的位置置零,只对真正需要预测的位置求平均。

连接回参数

算出 ∂L/∂logits 后,再往前一层传:∂L/∂W = (∂L/∂logits) · hᵀ,∂L/∂h = Wᵀ · (∂L/∂logits)。梯度就这样一层层传回整个 Transformer。

在真实 LLM 训练里: 一个 batch 可能有几百万个 token 位置,每个位置都算一次 softmax + 交叉熵,再一起反向传播。所谓“大模型训练”,本质上就是在海量 token 上反复做这件“预测下一个词、算交叉熵、回传梯度、更新参数”的事。

知识点完整总结

下面用六张卡片,把这一页涉及的每一个概念完整讲清楚。你可以按顺序读,也可以挑自己最关心的部分看。

📖 标准介绍

交叉熵(Cross-Entropy)是衡量两个概率分布差异的指标。设真实分布为 y,预测分布为 p,则 H(y, p) = −Σ yi log(pi)。在分类任务中,y 通常是 one-hot,所以交叉熵退化为 −log(ptarget)。

反向传播(Backpropagation)是利用链式法则,从损失出发,逐层计算每个参数梯度的算法。当损失是交叉熵、最后一层是 softmax 时,对 logits 的梯度化简为 p − y,这个结果既简洁又数值稳定。

💡 通俗介绍

想象你在玩猜词游戏。你给每个候选词打一个“我觉得它有多对”的分数。交叉熵就是一位严格的裁判:它只看你给正确答案打了多少分。你给正确答案的分数越高,它越满意;你越自信地猜错,它越严厉地惩罚你。

反向传播则是裁判给你的反馈:它会告诉你,每个候选词的分数应该调高还是调低。而且反馈非常直接——你给某个词的概率比它“应得”的高多少,就往下压多少;比它“应得”的低多少,就往上抬多少。这就是 p − y 的直觉。

🔧 解决了什么问题

它解决了“如何衡量模型预测分布与真实标签的差距,并高效回传梯度”这两个问题。

  • 衡量差距:交叉熵给模型一个标量损失,让优化目标明确。
  • 避免梯度消失:相比 MSE + softmax,交叉熵不会因为预测太错而梯度消失。
  • 梯度简洁:反向传播只需要 p − y,计算量小,实现简单。
  • 数值稳定:配合 log-sum-exp 技巧,即使 logits 很大也不会溢出。

⭐ 为什么重要

交叉熵 + 反向传播是几乎所有现代神经网络训练的核心引擎。语言模型的预训练、微调、RLHF,本质上都在最小化某种形式的交叉熵。

它让模型可以:

  • 在每一个位置上知道自己错在哪里;
  • 把错误信号公平地分配到每个参数;
  • 在数十亿参数上高效地做梯度下降。

没有交叉熵,模型就没有明确的优化目标;没有反向传播,梯度就无法从损失传到参数。两者缺一不可。

🚀 应用场景

  • 语言模型预训练:GPT、LLaMA、Qwen 等,每个 token 位置都用交叉熵。
  • 文本分类:情感分析、意图识别、垃圾邮件过滤。
  • 图像分类:ResNet、ViT 等,最后一层 softmax + 交叉熵。
  • 语音识别:CTC 损失、注意力解码器都用到交叉熵。
  • 机器翻译:序列到序列模型的标准训练目标。
  • 知识蒸馏:用教师模型的软标签做交叉熵,让学生模型模仿。
  • 对比学习:InfoNCE 损失本质上是交叉熵的变体。

⚠️ 缺陷与局限

  • 对错误标签敏感:如果训练数据里有噪声标签,交叉熵会强行让模型拟合错误答案,导致过拟合。需要标签平滑(Label Smoothing)或鲁棒损失来缓解。
  • 过度自信问题:交叉熵鼓励模型把正确类概率推到 1,容易导致模型过度自信,影响校准(Calibration)。
  • 类别不平衡:在类别极不平衡时,交叉熵可能被多数类主导。常用 Focal Loss、类别权重来修正。
  • 不适用于连续输出:交叉熵假设输出是概率分布。对于回归任务,应使用 MSE、Huber 等损失。
  • 数值稳定性要求高:直接实现 exp 和 log 容易溢出,必须使用 log-sum-exp 技巧或框架内置的稳定实现。
  • 反向传播内存开销:需要保存前向中间结果,大模型训练需要梯度检查点等技术来降低显存。
一句话理解本页: 交叉熵负责告诉模型“你错得有多离谱”,反向传播负责告诉每个参数“你该怎么改”。而 p − y 这个优雅的梯度公式,让这件本来复杂的事变得异常简洁高效。它和上一页的自动微分一起,构成了所有大模型能够“学会”的数学基础。

学完这一页,最好记住三件事

Softmax 把分数变概率

logits 是原始分数,softmax 把它们压成合法概率分布。指数放大差距,归一化保证总和为 1。

交叉熵只看正确类

因为标签是 one-hot,损失就是 −log(ptarget)。正确类概率越低,惩罚越猛烈。

梯度是 p − y

正确类梯度为负(抬高 logit),错误类梯度为正(压低 logit)。概率越大的错误类,被压得越狠。