datawhale--llm-algo-leetcode9️⃣ SFT Training Loop

链接:笔记

09. SFT Training Loop | 监督微调训练循环

本章目标:沿着一条固定数据链路,理解 decoder-only 模型怎样从 prompt + response 得到 SFT loss,并把原理落到 notebook 的两个 TODO。

数据构造 → 标签屏蔽 → Padding → Next-token 对齐 → Cross Entropy → 训练更新

本章不重新讲 RMSNorm、RoPE、GQA、SwiGLU;它们已经在第 01~05 章解释。本章只回答:模型输出以后,训练目标是怎样定义的?


模型架构总览

SFT 不会替换 Decoder Layer。它把监督数据接到第 05 章的 LLaMA-style decoder 和第 08 章的 LM Head 后面:
#mermaid-svg-Ba0Oz5eUnvmeZtlt{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-Ba0Oz5eUnvmeZtlt .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .error-icon{fill:#552222;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .marker{fill:#333333;stroke:#333333;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .marker.cross{stroke:#333333;}#mermaid-svg-Ba0Oz5eUnvmeZtlt svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-Ba0Oz5eUnvmeZtlt p{margin:0;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .cluster-label text{fill:#333;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .cluster-label span{color:#333;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .cluster-label span p{background-color:transparent;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .label text,#mermaid-svg-Ba0Oz5eUnvmeZtlt span{fill:#333;color:#333;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .node rect,#mermaid-svg-Ba0Oz5eUnvmeZtlt .node circle,#mermaid-svg-Ba0Oz5eUnvmeZtlt .node ellipse,#mermaid-svg-Ba0Oz5eUnvmeZtlt .node polygon,#mermaid-svg-Ba0Oz5eUnvmeZtlt .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .rough-node .label text,#mermaid-svg-Ba0Oz5eUnvmeZtlt .node .label text,#mermaid-svg-Ba0Oz5eUnvmeZtlt .image-shape .label,#mermaid-svg-Ba0Oz5eUnvmeZtlt .icon-shape .label{text-anchor:middle;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .rough-node .label,#mermaid-svg-Ba0Oz5eUnvmeZtlt .node .label,#mermaid-svg-Ba0Oz5eUnvmeZtlt .image-shape .label,#mermaid-svg-Ba0Oz5eUnvmeZtlt .icon-shape .label{text-align:center;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .node.clickable{cursor:pointer;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .arrowheadPath{fill:#333333;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-Ba0Oz5eUnvmeZtlt .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-Ba0Oz5eUnvmeZtlt .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-Ba0Oz5eUnvmeZtlt .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .cluster text{fill:#333;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .cluster span{color:#333;}#mermaid-svg-Ba0Oz5eUnvmeZtlt div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-Ba0Oz5eUnvmeZtlt rect.text{fill:none;stroke-width:0;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .icon-shape,#mermaid-svg-Ba0Oz5eUnvmeZtlt .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .icon-shape p,#mermaid-svg-Ba0Oz5eUnvmeZtlt .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .icon-shape .label rect,#mermaid-svg-Ba0Oz5eUnvmeZtlt .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-Ba0Oz5eUnvmeZtlt .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-Ba0Oz5eUnvmeZtlt .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-Ba0Oz5eUnvmeZtlt :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} Prompt + Response

原始样本
数据构造

input_ids / labels
Padding

attention_mask
Embedding

B,S,d

Decoder Layer × N

RMSNorm / Attention / RoPE / SwiGLU
LM Head

B,S,d\] → \[B,S,V

Shift + CrossEntropyLoss

ignore_index=-100
backward()

optimizer.step()

SFT 和预训练的差别

在标准 causal-LM SFT 中,模型前向结构与预训练相同。主要差异是监督标签:

text 复制代码
预训练:通常对所有有效 next-token target 计算 loss
SFT:只对 response target 计算 loss,prompt/padding target 设为 -100

这不代表 prompt 不经过模型。response 的预测必须依赖 prompt 的 hidden states。

本章固定数据流

#mermaid-svg-KOBKKHgszRz6p8i7{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-KOBKKHgszRz6p8i7 .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-KOBKKHgszRz6p8i7 .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-KOBKKHgszRz6p8i7 .error-icon{fill:#552222;}#mermaid-svg-KOBKKHgszRz6p8i7 .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-KOBKKHgszRz6p8i7 .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-KOBKKHgszRz6p8i7 .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-KOBKKHgszRz6p8i7 .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-KOBKKHgszRz6p8i7 .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-KOBKKHgszRz6p8i7 .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-KOBKKHgszRz6p8i7 .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-KOBKKHgszRz6p8i7 .marker{fill:#333333;stroke:#333333;}#mermaid-svg-KOBKKHgszRz6p8i7 .marker.cross{stroke:#333333;}#mermaid-svg-KOBKKHgszRz6p8i7 svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-KOBKKHgszRz6p8i7 p{margin:0;}#mermaid-svg-KOBKKHgszRz6p8i7 .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-KOBKKHgszRz6p8i7 .cluster-label text{fill:#333;}#mermaid-svg-KOBKKHgszRz6p8i7 .cluster-label span{color:#333;}#mermaid-svg-KOBKKHgszRz6p8i7 .cluster-label span p{background-color:transparent;}#mermaid-svg-KOBKKHgszRz6p8i7 .label text,#mermaid-svg-KOBKKHgszRz6p8i7 span{fill:#333;color:#333;}#mermaid-svg-KOBKKHgszRz6p8i7 .node rect,#mermaid-svg-KOBKKHgszRz6p8i7 .node circle,#mermaid-svg-KOBKKHgszRz6p8i7 .node ellipse,#mermaid-svg-KOBKKHgszRz6p8i7 .node polygon,#mermaid-svg-KOBKKHgszRz6p8i7 .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-KOBKKHgszRz6p8i7 .rough-node .label text,#mermaid-svg-KOBKKHgszRz6p8i7 .node .label text,#mermaid-svg-KOBKKHgszRz6p8i7 .image-shape .label,#mermaid-svg-KOBKKHgszRz6p8i7 .icon-shape .label{text-anchor:middle;}#mermaid-svg-KOBKKHgszRz6p8i7 .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-KOBKKHgszRz6p8i7 .rough-node .label,#mermaid-svg-KOBKKHgszRz6p8i7 .node .label,#mermaid-svg-KOBKKHgszRz6p8i7 .image-shape .label,#mermaid-svg-KOBKKHgszRz6p8i7 .icon-shape .label{text-align:center;}#mermaid-svg-KOBKKHgszRz6p8i7 .node.clickable{cursor:pointer;}#mermaid-svg-KOBKKHgszRz6p8i7 .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-KOBKKHgszRz6p8i7 .arrowheadPath{fill:#333333;}#mermaid-svg-KOBKKHgszRz6p8i7 .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-KOBKKHgszRz6p8i7 .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-KOBKKHgszRz6p8i7 .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-KOBKKHgszRz6p8i7 .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-KOBKKHgszRz6p8i7 .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-KOBKKHgszRz6p8i7 .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-KOBKKHgszRz6p8i7 .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-KOBKKHgszRz6p8i7 .cluster text{fill:#333;}#mermaid-svg-KOBKKHgszRz6p8i7 .cluster span{color:#333;}#mermaid-svg-KOBKKHgszRz6p8i7 div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-KOBKKHgszRz6p8i7 .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-KOBKKHgszRz6p8i7 rect.text{fill:none;stroke-width:0;}#mermaid-svg-KOBKKHgszRz6p8i7 .icon-shape,#mermaid-svg-KOBKKHgszRz6p8i7 .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-KOBKKHgszRz6p8i7 .icon-shape p,#mermaid-svg-KOBKKHgszRz6p8i7 .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-KOBKKHgszRz6p8i7 .icon-shape .label rect,#mermaid-svg-KOBKKHgszRz6p8i7 .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-KOBKKHgszRz6p8i7 .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-KOBKKHgszRz6p8i7 .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-KOBKKHgszRz6p8i7 :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} prompt_ids

10,20,30

response_ids

40,50,60,70

拼接

10,20,30,40,50,60,70

labels

-100,-100,-100,40,50,60,70

padding 到 8

input_ids 末尾加 0

labels 末尾加 -100
模型前向

logits 1,8,V
shift

logits 前 7 个

labels 后 7 个
loss

只保留 40,50,60,70

主线模型与 toy 示例

主线模型继续使用:

text 复制代码
B = 2       batch size
S = 16      sequence length
d = 4096    hidden size
V           vocabulary size

为了手算,本章所有数字演示固定使用下面这一条 toy 样本:

text 复制代码
prompt   = [10, 20, 30]
response = [40, 50, 60, 70]
pad_id   = 0
max_len  = 8

1070 只是 token ID,不代表具体文字;0 只是本 toy 示例的 padding ID。

本节涉及的模块

模块 shape 本章作用
input_ids [B,S] 模型真正读取的 token ID
attention_mask [B,S] 标记 padding 是否是有效输入
labels [B,S] loss 使用的目标 token;-100 表示忽略
Decoder hidden states [B,S,d] 第 05 章已经学过的模型内部表示
logits [B,S,V] LM Head 对词表中每个 token 的原始分数
shift_logits [B,S-1,V] 去掉最后一个没有下一个目标的位置
shift_labels [B,S-1] 去掉第一个没有前置位置的目标
loss scalar 有效 response target 的平均交叉熵

0. 学习本章前先认识这些名词

这一节只定义名词,不进入代码实现。所有例子都使用固定数据 [10,20,30] + [40,50,60,70]

名词 英文 本章中的准确含义
SFT Supervised Fine-Tuning 在带有标准回答的监督数据上继续训练预训练模型
Prompt Prompt 给模型的输入问题或条件;本例为 [10,20,30]
Response Response 希望模型生成的回答;本例为 [40,50,60,70]
Token Token tokenizer 切分出的文本片段;模型实际处理的是它的整数 ID
input_ids Input IDs 送进 Embedding 的 token ID 序列
Label / Target Label / Target loss 用来检查模型预测是否正确的目标 token
Logits Logits 模型对词表中每个 token 给出的未归一化分数
Loss Loss 预测分布与 target 之间的误差
Loss Masking Loss Masking 选择哪些 target 进入 loss
ignore_index Ignore Index loss 中需要忽略的特殊 label 值,本章为 -100
Padding Padding 为了把不同长度样本堆成同一个 Tensor 而追加的占位位置
pad_id Padding ID padding 占位 token 的真实 ID;本例为 0
attention_mask Attention Mask 标记输入位置是否有效;padding 位置通常为 0
Next-token prediction Next-token prediction 用当前位置及之前的 token 预测下一个 token
Shift Shift logits/labels 让位置 t 的输出对应位置 t+1 的 target
Causal mask Causal attention mask Attention 中禁止位置 t 读取未来位置 >t

0.1 三个 mask 的区别

text 复制代码
causal mask:限制 Attention 的时间方向,不能偷看未来
attention_mask:标记 padding 是否是有效输入
labels == -100:标记某个 target 是否计入 loss

它们发生在不同阶段,不能互相替代。

0.2 input_idslabels 的区别

固定数据 padding 后:

text 复制代码
input_ids = [10,20,30,40,50,60,70,0]
labels    = [-100,-100,-100,40,50,60,70,-100]
text 复制代码
input_ids:模型读什么
labels:    loss 检查什么

labels 不是第二份输入,也不是随便复制的序列。


1. 核心差异与机制

1.1 为什么要把 Prompt 和 Response 拼接?

decoder-only 模型是自回归模型,目标是学习:

P ( response ∣ prompt ) P(\text{response}\mid\text{prompt}) P(response∣prompt)

因此输入必须包含完整上下文:

text 复制代码
input_ids = [10,20,30,40,50,60,70]

模型可以根据:

text 复制代码
[10,20,30] → 预测 40
[10,20,30,40] → 预测 50
[10,20,30,40,50] → 预测 60
[10,20,30,40,50,60] → 预测 70

如果只输入 [40,50,60,70],模型就看不到 prompt,无法学习"给定 prompt 后怎样回答"。

1.2 为什么 labels 只保留 Response?

SFT 通常希望优化 response 的条件概率,所以 labels 对齐写成:

text 复制代码
input_ids = [10,20,30,40,50,60,70]
labels    = [-100,-100,-100,40,50,60,70]

解释:

text 复制代码
prompt 位置:作为上下文,不加入本例的监督目标
response 位置:保留真实 token ID,加入 loss

-100CrossEntropyLoss(ignore_index=-100) 约定的忽略标记。它不是 token ID,不会送入 Embedding。

精确地说,-100 位置不产生自己的 loss 项;不能笼统说 prompt 相关的所有参数都没有梯度,因为 response loss 仍可通过 Attention、残差和 Embedding 的计算图反传。

1.3 为什么需要 padding?

拼接得到的真实序列长度为 7:

text 复制代码
[10,20,30,40,50,60,70]

当 batch 需要固定长度 max_len=8 时,在末尾追加 pad_id=0

text 复制代码
input_ids = [10,20,30,40,50,60,70,0]

padding 不是 response 的一部分,只是为了得到规则 Tensor shape [B,8]

Padding 不是模型的必需步骤

如果只训练这一条样本,完全可以令:

text 复制代码
max_len = len(input_ids) = 7
input_ids = [10,20,30,40,50,60,70]

此时不需要 pad_id,模型也可以直接处理 shape [1,7] 的输入。max_len=8 是本 notebook toy 函数为了演示"固定长度输出"而指定的练习参数,不是 decoder-only 模型的硬性要求。

真正需要 padding 的情况是:一个 batch 中不同样本的长度不同,而普通 Tensor 必须是矩形 shape [B,S]。常见选择是把每个 batch 补到该 batch 的最长样本,而不是永远补到全局最大长度;也可以按长度分桶,或使用更复杂的 packed/ragged 方案。

对应的 attention mask 是:

text 复制代码
attention_mask = [1,1,1,1,1,1,1,0]

对应的 labels 是:

text 复制代码
labels = [-100,-100,-100,40,50,60,70,-100]

两者作用不同:

text 复制代码
attention_mask=0:padding 不作为有效输入
labels=-100:padding 不作为 loss target

1.4 为什么要 Shift?

模型位置 t 的输出预测下一个位置 t+1 的 token:

text 复制代码
位置       0    1    2    3    4    5    6    7
输入      10   20   30   40   50   60   70    0
输出       z0   z1   z2   z3   z4   z5   z6   z7

对应预测目标:

text 复制代码
z0 → 20
z1 → 30
z2 → 40
z3 → 50
z4 → 60
z5 → 70
z6 → 0(padding,忽略)
z7 → 没有下一个位置,丢弃

因此代码是:

python 复制代码
shift_logits = logits[..., :-1, :]
shift_labels = labels[..., 1:]

本例中:

text 复制代码
shift_labels = [-100,-100,40,50,60,70,-100]

真正有效的配对是:

text 复制代码
z2 → 40
z3 → 50
z4 → 60
z5 → 70

1.5 Causal mask 与 Shift 的区别

text 复制代码
causal mask:作用在 Attention score [B,H,S,S]
              防止模型读取未来 token

shift:       作用在 loss 输入 [B,S,V] 和 [B,S]
              让输出与下一个 target 对齐

即使模型已经使用 causal mask,计算 loss 时仍然需要 shift。


2. 数学公式与具体数值

2.1 Response-only SFT loss

T 是 response target 的位置集合,则:

L S F T = − 1 ∣ T ∣ ∑ t ∈ T log ⁡ p θ ( y t ∣ x < t ) \mathcal{L}{\mathrm{SFT}} =-\frac{1}{|T|}\sum{t\in T} \log p_\theta(y_t\mid x_{<t}) LSFT=−∣T∣1t∈T∑logpθ(yt∣x<t)

本例有效 target 是:

text 复制代码
40、50、60、70

因此:

text 复制代码
|T| = 4

假设模型对四个正确 token 的概率分别是:

text 复制代码
0.5、0.25、0.8、0.1

则:

L = − log ⁡ 0.5 + log ⁡ 0.25 + log ⁡ 0.8 + log ⁡ 0.1 4 ≈ 1.04 \mathcal{L} =-\frac{\log0.5+\log0.25+\log0.8+\log0.1}{4} \approx1.04 L=−4log0.5+log0.25+log0.8+log0.1≈1.04

三个 prompt 位置和一个 padding 位置不进入这个平均值。

2.2 单个 token 的交叉熵

如果某个有效 target 的正确概率为 0.8

ℓ = − log ⁡ ( 0.8 ) ≈ 0.223 \ell=-\log(0.8)\approx0.223 ℓ=−log(0.8)≈0.223

这个概率对应本例中的某一个 response token,例如 40;token ID 是固定数据的一部分,0.8 是模型 softmax 后给出的概率。

2.3 Shape 链路

主线模型:

text 复制代码
input_ids [B,S]
→ Embedding [B,S,d]
→ Decoder × N [B,S,d]
→ LM Head [B,S,V]
→ shift_logits [B,S-1,V]
→ flatten [B*(S-1),V]

labels:

text 复制代码
labels [B,S]
→ shift_labels [B,S-1]
→ flatten [B*(S-1)]

view(-1,V) 只合并 batch 维和时间维,词表维 V 保留为类别维。


3. 代码实现框架

3.1 build_sft_data

python 复制代码
def build_sft_data(prompt_ids: list[int], response_ids: list[int],
                   pad_id: int = 0, max_len: int = 16):
    # 1. 拼接 prompt 和 response
    input_ids = prompt_ids + response_ids

    # 2. prompt 不监督,response 保留真实 token id
    labels = [-100] * len(prompt_ids) + response_ids

    # 3. input_ids 和 labels 必须同步截断
    if len(input_ids) > max_len:
        input_ids = input_ids[:max_len]
        labels = labels[:max_len]
    else:
        # 4. 输入补 pad_id,labels 补 ignore_index
        pad_len = max_len - len(input_ids)
        input_ids = input_ids + [pad_id] * pad_len
        labels = labels + [-100] * pad_len

    return torch.tensor(input_ids, dtype=torch.long), \
           torch.tensor(labels, dtype=torch.long)

逐行说明:

  1. prompt_ids + response_ids 是 Python list 拼接,长度从 S_pS_r 变为 S_p+S_r
  2. [-100] * len(prompt_ids) + response_ids 让 prompt 位置被忽略、response 位置保留真实目标。
  3. 截断必须同时作用于两个列表,否则时间维无法对齐。
  4. padding 的 input_ids 使用 pad_id,padding 的 labels 使用 -100
  5. 返回的单条样本 shape 是 [max_len];DataLoader 堆叠后才是 [B,max_len]

本 toy 数据 max_len=8 的结果:

text 复制代码
input_ids = [10,20,30,40,50,60,70,0]
labels    = [-100,-100,-100,40,50,60,70,-100]

3.2 compute_sft_loss

python 复制代码
def compute_sft_loss(logits: torch.Tensor, labels: torch.Tensor):
    # 1. 对齐 next-token prediction
    shift_logits = logits[..., :-1, :].contiguous()
    shift_labels = labels[..., 1:].contiguous()

    # 2. 展平 batch 和时间维
    loss_fct = nn.CrossEntropyLoss(ignore_index=-100)
    shift_logits = shift_logits.view(-1, shift_logits.size(-1))
    shift_labels = shift_labels.view(-1)

    # 3. 只对非 -100 target 计算平均交叉熵
    return loss_fct(shift_logits, shift_labels)

逐行说明:

  1. logits[..., :-1, :] 的 shape 从 [B,S,V] 变为 [B,S-1,V]
  2. labels[..., 1:] 的 shape 从 [B,S] 变为 [B,S-1]
  3. contiguous() 确保切片后的内存布局适合后面的 view();它不改变数值。
  4. view(-1,V)[B,S-1,V] 变为 [B*(S-1),V]
  5. labels 同步变为 [B*(S-1)],每个整数是一个词表类别 ID。
  6. ignore_index=-100 排除 prompt 和 padding target;默认 reduction="mean" 只除以有效 target 数。

4. 动手实战

下面的 notebook 是独立 toy 示例,使用固定数据 [10,20,30][40,50,60,70],不改变主线模型的 B=2,S=16,d=4096

4.1 TODO 1:构造 labels

补全:

python 复制代码
labels = [-100] * len(prompt_ids) + response_ids

检查长度:

text 复制代码
len(labels) == len(prompt_ids) + len(response_ids)

4.2 TODO 2:截断与 padding

题目要求的 toy 策略是:

text 复制代码
超长:从末尾截断
不足:input_ids 补 pad_id,labels 补 -100

真实项目需要额外检查:截断后不能把整个 response 都删掉,否则一个样本可能没有有效 target。

4.3 TODO 3:Shift

python 复制代码
shift_logits = logits[..., :-1, :].contiguous()
shift_labels = labels[..., 1:].contiguous()

4.4 TODO 4:交叉熵

python 复制代码
loss_fct = nn.CrossEntropyLoss(ignore_index=-100)
shift_logits = shift_logits.view(-1, shift_logits.size(-1))
shift_labels = shift_labels.view(-1)
loss = loss_fct(shift_logits, shift_labels)

5. 测试与真实训练框架

5.1 测试分别验证什么?

python 复制代码
assert labels.tolist() == [-100, -100, -100, 40, 50, 60, 70, -100]

验证 labels 内容和位置正确。

python 复制代码
assert torch.isfinite(loss).item()
assert loss.item() < 0.01

在 notebook 的测试中,只有 z2..z5 被设置为正确 response 的高分,因此低 loss 可以检查 shift 和 ignore_index 是否符合预期。

python 复制代码
loss.backward()
assert logits.grad is not None

只证明 loss 与 logits 的计算图连通,不能单独证明真实模型一定收敛。

5.2 全 -100 的边界

python 复制代码
empty_labels = torch.full((1, 4), -100, dtype=torch.long)

如果一个 batch 的所有 target 都是 -100reduction="mean" 没有有效分母,通常得到 NaN。生产数据管道应过滤或显式拒绝这种样本。

5.3 完整训练循环

python 复制代码
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-5)

for batch in dataloader:
    optimizer.zero_grad()

    logits = model(
        input_ids=batch["input_ids"],
        attention_mask=batch.get("attention_mask"),
    ).logits

    loss = compute_sft_loss(logits, batch["labels"])
    loss.backward()
    optimizer.step()

这与第 00 章的训练闭环一致:

text 复制代码
zero_grad → forward → loss → backward → step

本章新增的是 labels 构造、shift 和 response-only loss。

5.4 Hugging Face causal LM

也可以把 labels 直接交给模型:

python 复制代码
outputs = model(
    input_ids=batch["input_ids"],
    attention_mask=batch["attention_mask"],
    labels=batch["labels"],
)
loss = outputs.loss

模型内部通常会完成同样的 shift 和交叉熵,因此最重要的检查仍然是 labels 是否正确。


6. 工程要点

6.1 ignore_index 不会跳过模型前向

ignore_index=-100 只排除对应的 loss 项:

text 复制代码
仍然计算 prompt 的 Embedding、Attention、Decoder 和 logits
只是 prompt target 不进入交叉熵平均

所以不能把 loss masking 理解成节省了 prompt 的全部显存和前向时间。

6.2 Padding 与截断

  1. pad_id 应使用 tokenizer 的真实 pad_token_id;toy 示例的 0 只是演示值。
  2. padding 的 labels 必须是 -100,否则模型会被训练去预测 padding。
  3. 真实 batch 必须传 attention_mask,否则 padding 可能成为无效上下文。
  4. 截断时应尽量保留 response,并统计被截断样本的比例。
  5. 若使用 EOS,需明确 EOS 是否属于 response 的监督目标。
  6. 如果 prompt 为空,通常需要 BOS 或明确处理第一个 response token,因为它没有前置位置可用于 shift。

6.3 有效 token 数

默认 reduction="mean" 是对有效 target 求平均。不同样本的 response 长度差异很大时,应监控每个 batch 的有效 token 数;做梯度累积时,也要考虑按有效 token 数进行一致归一化。


7. 与 LoRA 的关系

SFT loss 和 LoRA 解决的是不同层面的问题:

text 复制代码
SFT labels:决定哪些 target 计入 loss
LoRA:      决定哪些参数被更新

因此可以组合使用:

text 复制代码
prompt/response → labels mask → SFT loss
                                  ↓
                           只更新 LoRA 参数

LoRA 不会替代 labels 构造,也不会改变 shift 的位置关系。


8. 本章最终数据流

text 复制代码
prompt_ids [B,S_p]
response_ids [B,S_r]
        ↓ 拼接
input_ids [B,S]
        ↓ padding
input_ids [B,max_len]
attention_mask [B,max_len]
labels [B,max_len]
        ↓ 第 05 章的模型前向
hidden_states [B,max_len,d]
        ↓ 第 08 章的 LM Head
logits [B,max_len,V]
        ↓ shift
shift_logits [B,max_len-1,V]
shift_labels [B,max_len-1]
        ↓ ignore_index=-100
loss scalar
        ↓ backward → optimizer.step
更新模型参数

本章掌握检查

学完后应能用固定数据回答:

  1. 为什么 input_ids 同时包含 [10,20,30][40,50,60,70]
  2. 为什么 prompt 仍在 input_ids 中,但它的 labels 是 -100
  3. pad_id=0 在本 toy 数据中解决什么问题?
  4. attention_mask=[1,1,1,1,1,1,1,0] 和 labels 中的 -100 分别控制什么?
  5. 为什么 z2 对应 target 40
  6. 为什么 shift 后时间长度从 8 变成 7?
  7. 为什么全 -100 labels 可能得到 NaN

9. 参考代码与解析

9.1 build_sft_data 参考实现

python 复制代码
def build_sft_data(prompt_ids: list[int], response_ids: list[int],
                   pad_id: int = 0, max_len: int = 16):
    input_ids = prompt_ids + response_ids
    labels = [-100] * len(prompt_ids) + response_ids

    if len(input_ids) > max_len:
        input_ids = input_ids[:max_len]
        labels = labels[:max_len]
    else:
        pad_len = max_len - len(input_ids)
        input_ids = input_ids + [pad_id] * pad_len
        labels = labels + [-100] * pad_len

    return torch.tensor(input_ids, dtype=torch.long), \
           torch.tensor(labels, dtype=torch.long)

9.2 compute_sft_loss 参考实现

python 复制代码
def compute_sft_loss(logits: torch.Tensor, labels: torch.Tensor):
    shift_logits = logits[..., :-1, :].contiguous()
    shift_labels = labels[..., 1:].contiguous()

    loss_fct = nn.CrossEntropyLoss(ignore_index=-100) # 啥意思---看到labels=-100就跳过;创建了一个损失计算器
    shift_logits = shift_logits.view(-1, shift_logits.size(-1)) # 这是在干啥,corssentropy将预测位置摊平;[B,S,V]-->[B*S,V]即[N,V]
    shift_labels = shift_labels.view(-1) # labels也要reshape;[B*S]
    return loss_fct(shift_logits, shift_labels)# 对每一个有效 token 预测,检查模型有没有把真实 token 的概率放高,然后把所有位置的错误平均成一个数字。

9.3 一句话直觉

text 复制代码
input_ids 决定模型看到什么,
labels 决定 loss 检查什么,
shift 决定输出和下一个 token 怎样配对,
ignore_index 决定哪些配对不计入平均。

10.辅助理解

相关推荐
叫我:松哥24 分钟前
基于flask仿小米商城管理系统,使用flask开的一个商场网站
数据库·后端·python·flask
智购科技无人售货机工厂32 分钟前
2026自动售货机防拆机物理安全设计:从安全螺丝到结构互锁的工程实践~YH
android·网络·驱动开发·python·单片机·安全·云原生
2401_8685347838 分钟前
校园网规划与设计
python·pygame
小猴子爱上树43 分钟前
跨境电商AI批量图片翻译工具,视频字幕翻译免费试用
人工智能·python·音视频
zx_741484811 小时前
【Python入门】爬虫实战:Requests + XPath 从基础到实战
开发语言·爬虫·python
高洁011 小时前
Teacher Forcing技术解析
人工智能·python·深度学习·transformer·知识图谱
2601_962065251 小时前
从零创建一个 Django 项目
后端·python·django
溪语流沙1 小时前
【Python项目实战】虚拟环境与依赖管理:venv / pip / requirements.txt实操
开发语言·python·pip
梯度下降者2 小时前
CukeTest 自动化测试工具2023年度回顾白皮书
自动化测试·python·cuketest·qtquick/qml·linuxatk