深入理解 SFT 的 Loss Masking:如何做到“只对回答部分回传梯度“
深入理解 SFT 的 Loss Masking:如何做到"只对回答部分回传梯度"
一、先建立一个核心认知:模型永远在"预测下一个词"
在讲 loss masking 之前,我们必须先明确一件事:无论是预训练还是 SFT,语言模型做的都是同一件事——给定前面的所有 token,预测下一个 token。
模型的输出是一个概率分布:对于序列中的每一个位置,模型都会输出"下一个 token 应该是什么"的预测。训练时,我们把模型的预测和真实的下一个 token 做对比,用**交叉熵损失(Cross Entropy Loss)**来衡量预测得好不好,然后回传梯度、更新参数。
这里有个关键点常被忽略:模型对序列里的每一个位置都会产生一个预测和一个损失。也就是说,一条长度为 N 的序列,天然会产生 N 个位置的损失。默认情况下,这 N 个损失会被加起来(或平均),一起回传梯度。
Loss masking 要做的事,就是"挑出"其中一部分位置的损失参与梯度回传,把另一部分位置的损失"扔掉"。
二、用一个具体例子拆解整个过程
假设我们有这样一条 SFT 训练样本:
用户:法国的首都是哪里?
助手:巴黎。
第 1 步:拼接并加上特殊标记
实际训练时,这条数据会被拼成一整个序列,通常带有角色标记(这里简化表示):
<|user|> 法国 的 首都 是 哪里 ? <|assistant|> 巴黎 。 <|end|>
第 2 步:Tokenize,得到 token 序列
假设分词后得到(为了直观,用词而非真实子词):
| 位置 | 0 | 1 | 2 | 3 | 4 | 5 | 6 | 7 | 8 | 9 | 10 |
|---|---|---|---|---|---|---|---|---|---|---|---|
| token | <|user|> |
法国 | 的 | 首都 | 是 | 哪里 | ? | <|assistant|> |
巴黎 | 。 | <|end|> |
这里,位置 0~7 属于问题部分(prompt),位置 8~10 属于回答部分(response)。
第 3 步:构造 input 和 labels
训练时我们需要两个东西:
- input_ids:模型的输入,就是上面完整的 token 序列。
- labels:训练的"标准答案",用来和模型预测做对比、算损失。
关键就在 labels 上。我们把不想计算损失的位置,标记成一个特殊值 -100:
| 位置 | 0 | 1 | 2 | 3 | 4 | 5 | 6 | 7 | 8 | 9 | 10 |
|---|---|---|---|---|---|---|---|---|---|---|---|
| input_ids | <|user|> |
法国 | 的 | 首都 | 是 | 哪里 | ? | <|assistant|> |
巴黎 | 。 | <|end|> |
| labels | -100 | -100 | -100 | -100 | -100 | -100 | -100 | -100 | 巴黎 | 。 | <|end|> |
问题部分全部被替换成 -100,回答部分保留真实的 token id。
第 4 步:-100 是怎么起作用的?
这是整个机制的技术核心。在 PyTorch 里,交叉熵损失函数 nn.CrossEntropyLoss 有一个默认参数:
ignore_index = -100
它的含义是:凡是 label 等于 -100 的位置,直接跳过,不计入损失,也就不产生梯度。
所以当我们把问题部分的 label 全设为 -100 后,损失函数在计算时会自动忽略这些位置。最终只有位置 8、9、10(也就是"巴黎 。 <|end|>")参与了损失计算和梯度回传。
import torch
import torch.nn as nn
loss_fn = nn.CrossEntropyLoss(ignore_index=-100)
# logits: 模型在每个位置的预测分布, shape = [seq_len, vocab_size]
# labels: 其中问题部分为 -100
loss = loss_fn(logits, labels) # -100 的位置被自动跳过
这样,"只对回答部分回传梯度"就实现了——不是靠什么特殊的训练算法,而仅仅是靠把不想学的位置的 label 设成 -100 这样一个朴素的技巧。
三、一个容易搞混的细节:错位(shift)
前面为了讲清楚"哪些位置被 mask",我把 labels 直接和 input_ids 对齐了。但要注意,因为模型是"用当前及之前的 token 预测下一个 token",所以在真正算损失时,input 和 label 之间要错开一位:
- 用位置 0~9 的输入,去预测位置 1~10 的 token。
- 即:
模型在位置 i 的输出对应的答案是位置 i+1 的真实 token。
在主流框架(如 Hugging Face Transformers)里,这个 shift 操作是在模型内部自动完成的,我们只需要提供对齐好的 input_ids 和 labels,把不想学的部分填成 -100 即可。所以实践中,你构造 labels 时不需要自己手动错位,只要保证"问题部分填 -100、回答部分填真实 id"就行。
四、为什么一定要 mask 掉问题部分?
回到动机层面,理解了机制后,动机就非常清楚了:
-
我们不想让模型学"怎么提问"。 如果不 mask,模型会在"法国的首都是哪里?"这些问题 token 上也计算损失,等于在训练模型去生成用户的问题。但用户的问题是用户输入的,模型的职责是回答,不是造问题。
-
让梯度聚焦在真正重要的信号上。 SFT 的目标是让模型学会"给定这个问题,应该如何组织出理想的回答"。真正承载"回答协议"的,是回答部分的 token。只对回答回传梯度,训练信号更纯粹、更高效。
-
避免问题分布干扰模型。 SFT 数据里的问题往往由人工构造,其风格、分布和真实用户不完全一致。如果强迫模型去拟合这些问题,反而可能带来偏差。
五、总结:一张图看懂 loss masking
完整序列: [ 问题部分 tokens ] [ 回答部分 tokens ]
│ │
input_ids: [ 真实 id ] [ 真实 id ] ← 模型都要"看到"
│ │
labels: [ -100 ... ] [ 真实 id ] ← 只有回答部分算损失
│ │
忽略 ✗ 计算 ✓
(无梯度) (回传梯度)
一句话概括实现原理:
Loss masking 并没有改变模型结构或训练算法,它只是在构造训练标签(labels)时,把问题部分的 token 全部替换成
-100(即ignore_index),从而让交叉熵损失函数自动跳过这些位置——最终只有回答部分的 token 产生损失和梯度。
这就是为什么说:SFT 在工程上与预训练唯一实质性的区别,就是这个 loss masking。任务还是"预测下一个词",损失函数还是交叉熵,唯一变的是——我们只在"回答"这段序列上给模型打分。
梯度是从哪里来的
好问题。这其实要从"损失到底是怎么算出来的"以及"梯度是从哪里来的"讲起。核心结论是:梯度是从损失里"流"回参数的,如果某个位置根本没产生损失,那它自然就没有梯度可以回传。 而 -100 做的事,就是让那些位置"不产生损失"。
下面用具体数字一步步拆开看。
一、先搞清楚:梯度是从损失来的
深度学习训练的链条是这样的:
参数 → 前向计算 → 预测 → 损失(loss) → 反向传播 → 梯度 → 更新参数
注意方向:梯度是从 loss 出发,沿着计算图往回走,一路"分配"到每个参数上的。 换句话说,一个位置能不能产生梯度,取决于它有没有对最终的 loss 做出贡献。
如果某个位置的损失是 0,而且这个 0 是"被硬性排除、压根没参与加总"的,那么它对 loss 的贡献就是"不存在",反向传播时也就没有梯度从它这里流出去。
二、用数字看:损失是怎么在每个位置算出来的
还用上一篇的例子,序列有 11 个位置:
位置: 0 1 2 3 4 5 6 7 8 9 10
token: <user> 法国 的 首都 是 哪里 ? <assistant> 巴黎 。 <end>
模型对每个位置都会输出一个"下一个词的概率分布",然后拿这个分布去和真实的下一个 token 比,得到该位置的损失。单个位置的交叉熵损失公式是:
loss_i = -log( P(正确token) )
假设模型在各个位置预测正确 token 的概率如下(编造的数字,方便演示):
| 位置 i | 该位置要预测的真实token | 模型给它的概率 P | 该位置损失 loss_i = -log§ |
|---|---|---|---|
| 0 | 法国 | 0.01 | 4.61 |
| 1 | 的 | 0.30 | 1.20 |
| 2 | 首都 | 0.05 | 3.00 |
| 3 | 是 | 0.60 | 0.51 |
| … | … | … | … |
| 8 | 巴黎 | 0.40 | 0.92 |
| 9 | 。 | 0.80 | 0.22 |
| 10 | <end> |
0.90 | 0.10 |
三、关键对比:不 mask vs mask
情况 A:不做 mask(问题部分也算损失)
总损失是所有位置的损失加起来再平均:
total_loss = (loss_0 + loss_1 + ... + loss_10) / 11
= (4.61 + 1.20 + 3.00 + ... + 0.92 + 0.22 + 0.10) / 11
这个 total_loss 里包含了位置 0 的 4.61。既然位置 0 对 total_loss 有贡献,反向传播时,梯度就会顺着位置 0 的这条路流回去,模型就会努力去调整参数,让"法国"这个词在位置 0 更容易被预测出来。
结果:模型被训练去"学会生成问题",这不是我们想要的。
情况 B:做 mask(问题部分 label = -100)
损失函数看到某个位置的 label 是 -100,会做两件事:
- 不计算这个位置的损失(连
-log(P)都不算)。 - 不把它计入分母(求平均时也不数它)。
于是总损失变成只对回答部分求平均:
total_loss = (loss_8 + loss_9 + loss_10) / 3
= (0.92 + 0.22 + 0.10) / 3
= 0.41
看这个式子——位置 0 到 7 的 loss 根本没出现在里面。
四、为什么"没出现在 loss 里"就等于"没有梯度"?
这是最核心的一步。反向传播计算某个位置对参数的梯度时,本质是在算:
∂(total_loss) / ∂(该位置相关的参数)
现在 total_loss 的表达式里完全不包含位置 0 的任何项(loss_0 被 -100 挡掉了,从未加进来)。那么对位置 0 相关计算求偏导时:
∂(total_loss) / ∂(位置0的预测) = 0
一个常数(total_loss 里没有位置 0)对位置 0 求导,结果就是 0。梯度为 0,就意味着这个位置不会推动任何参数更新。
打个比方:
你在算一份团队奖金,只把 A、B、C 三个人的业绩加进了总额。现在要按"谁贡献了总额"来分奖金,D 的业绩压根没被加进总额,那 D 分到的当然是 0。不是因为 D 干得差,而是他根本没被算进这笔账。
-100 干的就是"把这个位置从账本里划掉"这件事。划掉了,它对 loss 贡献为 0,对 loss 求导也是 0,梯度自然为 0。
五、区分两个容易混淆的"0"
这里要特别澄清一个常见误解:
loss_i = 0(损失为零):意思是"模型在这个位置预测得完美"(P=1,-log(1)=0)。这个位置仍然参与了 loss 计算,只是碰巧损失值是 0。它是"账本里的一项,只不过金额为 0"。label = -100(被 mask):意思是"这个位置根本不进账本"。它连损失都不算,从计算图里被排除了。
真正让"没有梯度"的是后者。不是"损失值为 0",而是"这一项压根没进入 total_loss 的表达式",所以求导时它对应的偏导恒为 0。
六、一句话总结
label = -100
→ 交叉熵函数跳过该位置,不计算 loss,也不计入平均
→ total_loss 的数学表达式里完全不含这个位置
→ 反向传播对这个位置求偏导时,结果恒为 0
→ 梯度为 0 → 不更新任何与"生成问题"相关的参数
所以"把问题 token 换成 -100 就没有梯度",本质原因是:梯度只从 loss 里流出,而 -100 让这些位置从 loss 的计算中彻底消失了。 没进 loss,就没梯度。
梯度等于0推导
我们用链式法则一步步推导,看清楚 label = -100 时这个偏导为什么恒等于 0。
一、先定义符号
设:
- θ\thetaθ:模型参数(我们要对它求梯度)
- ziz_izi:模型在位置 iii 的输出 logits(一个 vocab 维的向量),它是参数 θ\thetaθ 的函数,即 zi=fi(θ)z_i = f_i(\theta)zi=fi(θ)
- ℓi\ell_iℓi:位置 iii 的交叉熵损失
- LLL:总损失 total_loss
二、总损失是怎么定义的(关键在这里)
带 mask 的总损失,数学上写成只对未被 mask 的位置求和。我们用一个指示变量 mim_imi 表示:
mi={1,labeli≠−100 (参与计算)0,labeli=−100 (被 mask) m_i = \begin{cases} 1, & \text{label}_i \neq -100 \ (\text{参与计算}) \\[4pt] 0, & \text{label}_i = -100 \ (\text{被 mask}) \end{cases} mi={1,0,labeli=−100 (参与计算)labeli=−100 (被 mask)
那么总损失是:
L=∑imi ℓi∑imi L = \frac{\displaystyle\sum_{i} m_i \, \ell_i}{\displaystyle\sum_{i} m_i} L=i∑mii∑miℓi
分子:只把 mi=1m_i = 1mi=1 的位置的损失加起来。
分母:未被 mask 的位置总数(用于求平均)。
注意:当 mi=0m_i = 0mi=0 时,miℓi=0m_i \ell_i = 0miℓi=0,这一项根本没进分子。
三、对参数求偏导(链式法则)
我们要算 ∂L∂θ\dfrac{\partial L}{\partial \theta}∂θ∂L。分母 ∑imi\sum_i m_i∑imi 是个常数(就是个计数,不依赖 θ\thetaθ),记作 NNN,所以:
∂L∂θ=1N∑imi ∂ℓi∂θ \frac{\partial L}{\partial \theta} = \frac{1}{N} \sum_i m_i \, \frac{\partial \ell_i}{\partial \theta} ∂θ∂L=N1i∑mi∂θ∂ℓi
再对每一项用链式法则展开(损失 ℓi\ell_iℓi 通过 logits ziz_izi 依赖参数 θ\thetaθ):
∂L∂θ=1N∑imi⋅∂ℓi∂zi⋅∂zi∂θ \boxed{\ \frac{\partial L}{\partial \theta} = \frac{1}{N} \sum_i m_i \cdot \frac{\partial \ell_i}{\partial z_i} \cdot \frac{\partial z_i}{\partial \theta}\ } ∂θ∂L=N1i∑mi⋅∂zi∂ℓi⋅∂θ∂zi
这就是完整公式。
四、聚焦某个被 mask 的位置 kkk(label = -100)
现在单独看被 mask 的那个位置 kkk,它对总梯度的贡献是求和式里的第 kkk 项:
位置 k 的贡献=1N⋅mk⋅∂ℓk∂zk⋅∂zk∂θ \text{位置 } k \text{ 的贡献} = \frac{1}{N} \cdot m_k \cdot \frac{\partial \ell_k}{\partial z_k} \cdot \frac{\partial z_k}{\partial \theta} 位置 k 的贡献=N1⋅mk⋅∂zk∂ℓk⋅∂θ∂zk
因为 labelk=−100\text{label}_k = -100labelk=−100,所以 mk=0m_k = 0mk=0,代入:
位置 k 的贡献=1N⋅0⏟mk⋅∂ℓk∂zk⋅∂zk∂θ=0 \text{位置 } k \text{ 的贡献} = \frac{1}{N} \cdot \underbrace{0}_{m_k} \cdot \frac{\partial \ell_k}{\partial z_k} \cdot \frac{\partial z_k}{\partial \theta} = 0 位置 k 的贡献=N1⋅mk 0⋅∂zk∂ℓk⋅∂θ∂zk=0
不管 ∂ℓk∂zk\dfrac{\partial \ell_k}{\partial z_k}∂zk∂ℓk 和 ∂zk∂θ\dfrac{\partial z_k}{\partial \theta}∂θ∂zk 是多少,前面乘了个 0,整项就是 0。
五、更本质的视角:ignore_index 是"损失从未定义"
上面用 mim_imi 是一种数学上的"事后乘零"写法,方便看清楚。但要注意,PyTorch 里 ignore_index=-100 的真实实现比乘零更彻底:它压根不计算 ℓk\ell_kℓk,也不把 zkz_kzk 接入 loss 的计算图。
换句话说,从计算图角度看:
L=L(zi1,zi2,… )(只包含未被 mask 的位置) L = L(z_{i_1}, z_{i_2}, \dots)\quad \text{(只包含未被 mask 的位置)} L=L(zi1,zi2,…)(只包含未被 mask 的位置)
zkz_kzk(被 mask 位置的 logits)根本不是 LLL 的自变量。既然 LLL 的表达式里压根没有 zkz_kzk,那么:
∂L∂zk=0(常数对无关变量求导) \frac{\partial L}{\partial z_k} = 0 \quad(\text{常数对无关变量求导}) ∂zk∂L=0(常数对无关变量求导)
再往回一步,梯度是否传到参数 θ\thetaθ,取决于这条链:
∂L∂θ∣经由位置 k=∂L∂zk⏟= 0⋅∂zk∂θ=0 \frac{\partial L}{\partial \theta}\bigg|_{\text{经由位置 }k} = \underbrace{\frac{\partial L}{\partial z_k}}_{=\,0} \cdot \frac{\partial z_k}{\partial \theta} = 0 ∂θ∂L 经由位置 k==0 ∂zk∂L⋅∂θ∂zk=0
第一个因子就是 0,整条路径断了,梯度传不过来。
六、直觉总结
两种理解方式殊途同归:
| 视角 | 说法 | 结论 |
|---|---|---|
| 乘零视角 | mk=0m_k = 0mk=0,那一项被乘成 0 | 该位置贡献 = 0 |
| 计算图视角 | zkz_kzk 不是 LLL 的自变量 | ∂L/∂zk=0\partial L / \partial z_k = 0∂L/∂zk=0,链断掉 |
一句话:
梯度经由某位置回传的通道是 ∂L∂θ∣位置k=∂L∂ℓk⋅∂ℓk∂zk⋅∂zk∂θ\dfrac{\partial L}{\partial \theta}\big|_{\text{位置}k} = \dfrac{\partial L}{\partial \ell_k}\cdot\dfrac{\partial \ell_k}{\partial z_k}\cdot\dfrac{\partial z_k}{\partial \theta}∂θ∂L 位置k=∂ℓk∂L⋅∂zk∂ℓk⋅∂θ∂zk。
label=-100让 ∂L∂ℓk=0\dfrac{\partial L}{\partial \ell_k}=0∂ℓk∂L=0(这一项不在总损失里),于是整条乘积为 0,梯度断在源头,传不到任何参数。
后记
2026年8月12日于上海,在claude opus 4.8辅助下完成。
更多推荐

所有评论(0)