深入理解 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_idslabels,把不想学的部分填成 -100 即可。所以实践中,你构造 labels 时不需要自己手动错位,只要保证"问题部分填 -100、回答部分填真实 id"就行。


四、为什么一定要 mask 掉问题部分?

回到动机层面,理解了机制后,动机就非常清楚了:

  1. 我们不想让模型学"怎么提问"。 如果不 mask,模型会在"法国的首都是哪里?"这些问题 token 上也计算损失,等于在训练模型去生成用户的问题。但用户的问题是用户输入的,模型的职责是回答,不是造问题

  2. 让梯度聚焦在真正重要的信号上。 SFT 的目标是让模型学会"给定这个问题,应该如何组织出理想的回答"。真正承载"回答协议"的,是回答部分的 token。只对回答回传梯度,训练信号更纯粹、更高效。

  3. 避免问题分布干扰模型。 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,会做两件事:

  1. 不计算这个位置的损失(连 -log(P) 都不算)。
  2. 不把它计入分母(求平均时也不数它)。

于是总损失变成只对回答部分求平均:

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_ii:位置 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=imiimii

分子:只把 mi=1m_i = 1mi=1 的位置的损失加起来。
分母:未被 mask 的位置总数(用于求平均)。

注意:当 mi=0m_i = 0mi=0 时,miℓi=0m_i \ell_i = 0mii=0,这一项根本没进分子


三、对参数求偏导(链式法则)

我们要算 ∂L∂θ\dfrac{\partial L}{\partial \theta}θL。分母 ∑imi\sum_i m_iimi 是个常数(就是个计数,不依赖 θ\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=N1imiθi

再对每一项用链式法则展开(损失 ℓi\ell_ii 通过 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=N1imiziiθ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 的贡献=N1mkzkkθ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 的贡献=N1mk 0zkkθzk=0

不管 ∂ℓk∂zk\dfrac{\partial \ell_k}{\partial z_k}zkk∂zk∂θ\dfrac{\partial z_k}{\partial \theta}θzk 是多少,前面乘了个 0,整项就是 0。


五、更本质的视角:ignore_index 是"损失从未定义"

上面用 mim_imi 是一种数学上的"事后乘零"写法,方便看清楚。但要注意,PyTorch 里 ignore_index=-100 的真实实现比乘零更彻底:它压根不计算 ℓk\ell_kk,也不把 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{常数对无关变量求导}) zkL=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 zkLθzk=0

第一个因子就是 0,整条路径断了,梯度传不过来。


六、直觉总结

两种理解方式殊途同归:

视角 说法 结论
乘零视角 mk=0m_k = 0mk=0,那一项被乘成 0 该位置贡献 = 0
计算图视角 zkz_kzk 不是 LLL 的自变量 ∂L/∂zk=0\partial L / \partial z_k = 0L/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=kLzkkθzklabel=-100∂L∂ℓk=0\dfrac{\partial L}{\partial \ell_k}=0kL=0(这一项不在总损失里),于是整条乘积为 0,梯度断在源头,传不到任何参数。

后记

2026年8月12日于上海,在claude opus 4.8辅助下完成。

Logo

智能硬件社区聚焦AI智能硬件技术生态,汇聚嵌入式AI、物联网硬件开发者,打造交流分享平台,同步全国赛事资讯、开展 OPC 核心人才招募,助力技术落地与开发者成长。

更多推荐