本章前置:你已读过第 01–06 章。也就是说你知道:对数与概率的基本运算、softmax 把 logits 变成概率分布、Transformer 在每个位置输出一个长度为 (词表大小)的 logits 向量、以及数据被打包成形状 的 token 批次。
你将学到:模型到底根据什么信号学习。我们从”语言模型就是一串条件概率的连乘”出发,经由最大似然估计(MLE)→ 负对数似然(NLL)→ 交叉熵(cross-entropy)一步步推导出训练损失;手推交叉熵对 logits 的梯度 ,搞懂”为什么交叉熵这么好训”;弄清输入和标签错位一位(shift)的具体下标;推导困惑度(perplexity)并理解它的直觉;最后看 SFT 里如何用掩码损失只对 assistant 的回答算损失(为第 12 章埋伏笔)。动手部分用一个 3 类小例子手算交叉熵,并用
torch.nn.functional.cross_entropy对照。👈 上一章:注意力机制 · 完整推导 | 返回总览
7.1 语言模型就是一串条件概率
我们想让模型学会”语言”。“会语言”用数学说,就是模型能对任意一句话出现的概率给出一个合理的估计——常见的话概率高,胡言乱语概率低。
设一句话(一个序列)由 个 token 组成:。整句话的概率,可以用概率的链式法则一字不漏地拆成连乘:
逐符号解释: 表示”由参数为 (模型所有权重)的模型给出的概率”; 是整句话; 是连乘; 表示” 之前的所有 token”(即 )。这个式子的大白话是:一句话的概率 = 第 1 个 token 出现的概率 × (在第 1 个已知的前提下第 2 个的概率) × (在前 2 个已知的前提下第 3 个的概率) × ……。
为什么能这样拆?这就是链式法则,永远成立( 的多变量推广),不需要任何假设。它之所以重要,是因为它把”建模一整句话”这个大问题,化简成了一个反复出现的小问题:给定前文,预测下一个 token。而”给定前文预测下一个 token”——正是我们第 06 章那个带因果掩码的 Transformer 干的事!每个位置 输出的 logits,经过 softmax,就是 这个分布。一切都对上了。
7.2 从最大似然到负对数似然
7.2.1 最大似然估计(MLE):让真实数据”最不意外”
我们手上有一大堆真实文本(训练语料)。训练的目标朴素地讲就是:调整参数 ,让模型觉得”这些真实文本”出现的概率尽可能大。模型如果给真实文本打了高概率,说明它”觉得真实的句子很正常、不意外”,这正是我们想要的。这个原则叫最大似然估计(Maximum Likelihood Estimation,MLE)。
对单个序列,我们要最大化的就是 7.1 的 :
逐符号解释: 的意思是”找到让后面那个式子取最大值的 “。也就是说,我们要找一组权重,使模型赋予真实序列的概率最大。
7.2.2 取对数:把连乘变成连加
直接对一长串概率的连乘做优化非常糟糕:每个概率都是 0 到 1 之间的小数,几百上千个一连乘,结果会小到计算机的浮点数都表示不了(数值下溢)。救星是对数。对数有个黄金性质:,它能把连乘变成连加。而且对数是单调递增函数——最大化 和最大化 得到的 完全一样,不改变答案。于是两边取对数:
逐句解释:连乘的对数,等于各项对数的连加。右边这个 叫对数似然(log-likelihood)。它在数值上稳定得多,也方便求导(加法的导数好算)。
7.2.3 取负:把”最大化”变成”最小化”损失
机器学习的惯例是最小化一个损失函数(loss,越小越好),而不是最大化。这只需在对数似然前面加个负号,得到负对数似然(Negative Log-Likelihood,NLL):
逐句解释:加负号后,“最大化对数似然”就等价于”最小化负对数似然”。这就是我们的损失函数雏形。直觉上:如果模型给真实 token 的概率 很接近 1, 接近 0, 也接近 0,损失小(好);如果模型给真实 token 的概率很小(接近 0), 会冲向 ,损失巨大(差)。模型越是”对真实答案感到意外”,罚得越狠,这非常合理。
到这里请记住这条主线:最大化似然 = 最小化 NLL。接下来我们说明,单个位置的 NLL 其实就是大名鼎鼎的”交叉熵”。
7.3 单个位置:交叉熵就是 真实 token 的概率
聚焦某一个位置 。模型在这里输出一个长度为 的 logits 向量 ( 是词表大小,本仓库约 5 万)。softmax 把它变成在整个词表上的概率分布:
逐符号解释: 是 logits 的第 个分量(对应词表里第 个 token 的”原始得分”); 是模型认为”下一个 token 是第 个词”的概率。所有 加起来等于 1。
真实的下一个 token 是某个确定的 id,记为 。我们可以把这个”标准答案”写成一个 one-hot 向量 :它只在第 个位置是 1,其余全是 0(意思是”正确答案 100% 是第 个词”)。交叉熵衡量”真实分布 “和”模型分布 “之间的差距,定义为:
逐符号解释:对词表里每个词,把”真实分布在该词上的值”乘以”模型概率的对数”,求和再取负。但 只有第 个位置是 1、别的都是 0,所以这个和里只有 这一项活下来,其余全乘以 0 消失了:
逐句解释:交叉熵化简成了非常干净的一句话——就是”真实那个 token 的概率取对数,再加负号”。这恰好就是 7.2 里单个位置的负对数似然!所以在语言模型里,“交叉熵损失”和”负对数似然”是同一个东西的两个名字。模型给正确 token 的概率越高,这一项损失越小。
7.3.1 关键推导:交叉熵对 logits 的梯度是
为什么交叉熵配 softmax”特别好训”?秘密在它的梯度异常干净。我们要算损失 对每个 logit 的偏导。这一步是本章的数学高潮,我们一步不跳地推。
先把交叉熵用 logits 完整写开。因为 ,取对数:
逐句解释: 把分式变成”分子的对数减分母的对数”,即 ,再整体取负就得到上式。记右边那个 为 (log-sum-exp)。
现在对某个 求偏导,分两部分。
第一部分, 对 求导:只有当 时它是 ,否则是 。用 one-hot 记号写,就是 。
第二部分, 对 求导。用链式法则(外层 的导数是 ,内层 对 的导数只有 那一项 留下):
逐句解释:这正好又是 softmax 的定义! 对第 个 logit 的导数,恰好等于模型给第 个词的概率 。
把两部分合起来:
逐句解释这为什么是”好训”的关键:梯度就是**“模型预测的概率分布”减去”真实的 one-hot 分布”**,干净得不能再干净。它的含义极其直观——
- 对真实那个 token():梯度是 ,是个负数。梯度下降会沿梯度反方向走,于是它抬高 ,让模型更倾向选真实 token。模型当前越没把握( 离 1 越远),这股推力越大。
- 对其他错误 token():梯度是 ,是个正数,于是它压低 。模型当前给某个错词的概率越高,压它的力越大。
- 当模型完全预测对了(),梯度处处为 0,不再调整。
没有讨厌的饱和、没有梯度消失、推力大小自动正比于”错得多严重”。这就是交叉熵 + softmax 成为分类/语言建模标配的根本原因。
7.4 整段序列与一个 batch 的平均损失
7.3 处理的是一个位置。一条序列有 个位置、一个 batch 有 条序列。我们把所有位置的交叉熵取平均作为最终损失(取平均而非求和,是为了让损失值不随 batch 大小、序列长度变化,便于比较):
逐符号解释:外层两个求和遍历 batch 里每条序列 和每个位置 ;括号里是该位置的交叉熵( 是这条序列在位置 的真实下一个 token);前面的 求平均。一句话:整批数据上,每个位置交叉熵的平均值。
7.4.1 输入与标签错位一位(shift)的具体下标
上面式子里反复出现”位置 的真实下一个 token “。这”下一个”在工程上就是一次错位一位(shift by one)。具体地,给模型的输入和它要预测的标签是同一串 token,只是错开一格:
逐句解释下标对应:模型在看到 (输入第 0 位)时,要预测的标签是 (标签第 0 位);看到前缀 时要预测 ;……一句话:标签就是输入整体往左挪一位。位置 的输入是 ,它要预测的目标是 ,这正是”用前文预测下一个 token”。
在本仓库里,这个错位有两种写法,理解上等价。预训练主路径 src/models/transformer.py 的 forward 里,数据加载时已经把 targets 准备成”比 idx 超前一位”的张量了,所以 forward 直接在所有位置上一把算交叉熵,不必在函数里再手动错位:
x = self.forward_hidden(idx)
logits = self.lm_head(x)
loss = None
if targets is not None:
B, T, C = logits.shape
flat_logits = logits.reshape(B * T, C)
targets = targets.reshape(B * T).long()
loss = F.cross_entropy(flat_logits, targets)
return logits, loss
逐行解释:logits 形状是 (这里变量名 C 实为词表大小 );reshape(B * T, C) 把前两维 拍平成一长条,变成 ;targets.reshape(B * T) 把标签也拍平成 的整数向量;然后 F.cross_entropy 一次性对这 个位置算交叉熵并求平均——这正是 。(顺带一提:注释里特意用 reshape 而不是 view,是因为 targets 来自一个非连续的张量切片,在 CPU 上 view 会报错,reshape 更稳妥。)
而 SFT 路径 src/post_training/sft.py 则把错位显式地写出来(因为它还要配一个掩码,见 7.6 节):
# Predict token t+1 from position t (same shift the base model uses).
logits = logits[:, :-1, :]
targets = tokens[:, 1:]
逐行解释:logits[:, :-1, :] 取除最后一个位置外的所有位置(因为最后一个位置没有”下一个 token”可预测,丢掉);tokens[:, 1:] 取从第二个 token 开始的序列作为标签。这两行合起来,就把”位置 的输出”和”第 个真实 token”严丝合缝地对齐了——和上面那张错位下标表完全一致。
7.5 困惑度:平均每步在多少个选项里纠结
损失值 (平均交叉熵)是个抽象的数,不太好”体感”。困惑度(perplexity,PPL)给它换了个更直观的尺度。定义就是把平均交叉熵放到指数上:
逐符号解释: 是以 为底的指数函数,正好”抵消”交叉熵里的自然对数 。
为什么这样定义有意义?我们从直觉推一遍。先看单个位置:若模型给真实 token 的概率是 ,那一步的交叉熵是 ,它的困惑度是 。也就是说,单步困惑度 = 真实 token 概率的倒数。如果模型很确定(),困惑度是 1(一点都不困惑);如果模型在 个候选里完全均匀地猜(每个概率 ),困惑度就是 。所以困惑度的直觉是:
平均每预测一个 token,模型大约在多少个”等可能的选项”之间纠结。
困惑度越低越好:1 是完美(每步都笃定),越大说明模型越”懵”。再看一个有用的标尺——一个完全没训练、瞎猜的模型,会给词表里每个 token 大致均等的概率 ,此时平均交叉熵约为 ,困惑度约为
也就是”在整个词表 个词里均匀乱猜”。本仓库词表 ,对应初始交叉熵约 。所以你训练时如果看到 loss 从约 10.8 开始往下掉、困惑度从约 5 万往下掉,就说明模型正在从”瞎猜”逐步学会”把概率集中到合理的下一个词上”。这也是第 11 章你会亲眼盯着看的曲线。
7.6 掩码损失:SFT 只对”回答”算账(第 12 章伏笔)
预训练时,语料是一整条 token 流,每个位置都要预测下一个,没有”该不该算”的区分。但到了 SFT(指令微调)阶段,一条训练样本长这样:
[系统/用户的 prompt 部分] [assistant 的回答部分]
我们希望模型学会的是怎么回答,而不是去背诵、复现用户的提问。如果对 prompt 部分也算损失,模型会把一部分”学习力气”浪费在预测用户会问什么上——那不是我们要的。
解决办法是给每个位置配一个 0/1 的掩码 :assistant 回答的 token 处 (要算损失),prompt 的 token 处 (不算)。于是 SFT 损失变成”只在掩码为 1 的位置上求平均的交叉熵”:
逐符号解释: 是位置 的交叉熵;分子把每个位置的交叉熵乘上它的掩码再求和——掩码为 0 的位置(prompt)直接被乘没了,贡献为 0;分母 是”被算账的位置总数”,用它做平均,保证结果是”每个被监督的 token 的平均损失”,不受 prompt 长短影响。
对照 src/post_training/sft.py 的实现:
mask = loss_mask[:, 1:].to(logits.dtype)
V = logits.size(-1)
ce = F.cross_entropy(logits.reshape(-1, V).float(), targets.reshape(-1).long(), reduction="none")
ce = ce.view(targets.shape) * mask
return ce.sum() / mask.sum().clamp(min=1.0)
逐行解释:
mask = loss_mask[:, 1:]:掩码同样要和标签一起错位一位([:, 1:]),才能和targets对齐——这呼应 7.4.1 的 shift。reduction="none":让cross_entropy不要自动求平均,而是返回每个位置各自的交叉熵(一个向量),因为我们要先乘掩码再自己求平均。这正是上式里的 。ce = ce.view(targets.shape) * mask:把逐位置交叉熵恢复成 形状,逐元素乘掩码——对应分子里的 ,prompt 位置被清零。ce.sum() / mask.sum().clamp(min=1.0):分子求和、除以掩码之和(即被监督 token 数)——正是上式的分式。clamp(min=1.0)是个保险:万一某批一个回答 token 都没有(分母为 0),避免除以零。
联系 \text{ignore_index} 的常见做法:很多代码库不用乘掩码,而是把不算损失的标签位置设成一个特殊值(常用
-100),再传给cross_entropy(..., ignore_index=-100),效果一样——那些位置被直接跳过。本仓库选择了”乘 0/1 掩码”这种更显式、更易读的写法。两种思路你都会在第 12 章细讲到。
7.7 动手:3 类小例子手算交叉熵,再用 torch 对照
我们用一个最小的 3 分类例子(把”3 类”想成”词表只有 3 个词”)亲手算一遍交叉熵,再让 PyTorch 验证。假设模型在某个位置输出 logits
真实答案是第 0 类()。第一步,算 softmax。先指数化:;它们的和约为 。于是
第二步,交叉熵就是 真实类(第 0 类)的概率:
第三步,验证一下梯度公式 :这里 ,所以梯度约为 ——第 0 类是负的(会被抬高),另两类是正的(会被压低),和 7.3.1 的结论一致。
现在用 torch.nn.functional.cross_entropy 对照(注意:PyTorch 的 cross_entropy 直接吃 logits,内部自带 softmax+log,你不要自己先 softmax 再传进去,否则等于做了两次):
import torch
import torch.nn.functional as F
logits = torch.tensor([[2.0, 1.0, 0.1]]) # 形状 (1, 3): 1 个样本, 3 个类
target = torch.tensor([0]) # 真实类别是第 0 类
# 1) 手动: softmax -> 取真实类概率 -> -log
probs = F.softmax(logits, dim=-1)
print("softmax 概率:", probs) # ≈ [[0.659, 0.242, 0.099]]
manual_ce = -torch.log(probs[0, 0])
print("手算交叉熵:", manual_ce.item()) # ≈ 0.417
# 2) 直接用 cross_entropy(吃 logits, 不要自己先 softmax)
ce = F.cross_entropy(logits, target)
print("F.cross_entropy:", ce.item()) # ≈ 0.417, 与手算一致
# 3) 顺便验证梯度 = p - onehot(y)
logits2 = logits.clone().requires_grad_(True)
F.cross_entropy(logits2, target).backward()
print("logits 的梯度:", logits2.grad) # ≈ [[-0.341, 0.242, 0.099]]
运行后你会看到三件事互相印证:F.cross_entropy 的输出 ≈ 0.417,和你手算的 一致;logits2.grad ≈ [-0.341, 0.242, 0.099],和 一致。把 target 改成 torch.tensor([2])(假装真实答案是模型最不看好的第 2 类),你会看到交叉熵一下子涨到约 ——模型对真实答案越意外,罚得越重,这就是 7.2.3 说的那股劲。
再算个困惑度感受一下:当 CE ≈ 0.417 时,这一步的困惑度是 ,意思是”模型在大约 1.5 个等可能选项之间纠结”,已经相当笃定;而 CE ≈ 2.31 时困惑度是 ——在词表只有 3 个词的情况下,这已经比”瞎猜的 3”还差,说明模型把概率押错了地方。
小结
- 语言模型 = 一串条件概率连乘:,把”建模整句话”化简成”给定前文预测下一个 token”。
- 训练目标由 MLE → 取对数(连乘变连加)→ 取负 推出负对数似然 NLL:“最大化似然 = 最小化 NLL”。
- 单个位置的 NLL 就是交叉熵 ;它对 logits 的梯度是 ——干净、无饱和、推力正比于错误程度,这是交叉熵”好训”的根本。
- 整批损失是所有位置交叉熵的平均 ;输入与标签错位一位(输入 、标签 ,代码里
logits[:, :-1]配tokens[:, 1:])。 - 困惑度 :平均每步在多少个等可能选项里纠结;瞎猜模型 PPL ≈ (本仓库约 5 万,对应初始 loss ≈ 10.83)。
- 掩码损失:SFT 只对 assistant 回答的 token 算交叉熵(掩码为 1),prompt 部分掩码为 0 不计,损失对被监督 token 数求平均——为第 12 章铺路。
自测题
- 为什么训练时要对似然取对数?它解决了什么数值问题、又带来了什么计算上的方便?
- 写出交叉熵对 logits 的梯度公式,并用一句话解释:对”真实 token”对应的 logit,梯度是正还是负?它会被抬高还是压低?
- 输入序列是
[t0, t1, t2, t3],那么标签序列是什么?代码里用哪两个切片实现这个错位? - 某模型在某条 100 个 token 的句子上平均交叉熵是 。它的困惑度是多少?这个数直觉上代表什么?
- SFT 损失里,如果不乘掩码、对所有位置(含 prompt)都算交叉熵,会带来什么问题?
深入参考
- 本仓库精炼版参考:目标、损失与困惑度
- 源码:
src/models/transformer.py(forward里的交叉熵)、src/post_training/sft.py(sft_loss掩码损失)
下一章我们有了损失这个”方向盘”,接下来要解决”怎么沿着它稳稳地走”——优化器与训练系统:从梯度下降到 Adam/AdamW、学习率调度、梯度累积与混合精度。
下一章 👉 优化与训练系统