链接:笔记
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
10 到 70 只是 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_ids 与 labels 的区别
固定数据 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
-100 是 CrossEntropyLoss(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)
逐行说明:
prompt_ids + response_ids是 Python list 拼接,长度从S_p和S_r变为S_p+S_r。[-100] * len(prompt_ids) + response_ids让 prompt 位置被忽略、response 位置保留真实目标。- 截断必须同时作用于两个列表,否则时间维无法对齐。
- padding 的
input_ids使用pad_id,padding 的labels使用-100。 - 返回的单条样本 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)
逐行说明:
logits[..., :-1, :]的 shape 从[B,S,V]变为[B,S-1,V]。labels[..., 1:]的 shape 从[B,S]变为[B,S-1]。contiguous()确保切片后的内存布局适合后面的view();它不改变数值。view(-1,V)把[B,S-1,V]变为[B*(S-1),V]。- labels 同步变为
[B*(S-1)],每个整数是一个词表类别 ID。 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 都是 -100,reduction="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 与截断
pad_id应使用 tokenizer 的真实pad_token_id;toy 示例的0只是演示值。- padding 的 labels 必须是
-100,否则模型会被训练去预测 padding。 - 真实 batch 必须传
attention_mask,否则 padding 可能成为无效上下文。 - 截断时应尽量保留 response,并统计被截断样本的比例。
- 若使用 EOS,需明确 EOS 是否属于 response 的监督目标。
- 如果 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
更新模型参数
本章掌握检查
学完后应能用固定数据回答:
- 为什么
input_ids同时包含[10,20,30]和[40,50,60,70]? - 为什么 prompt 仍在
input_ids中,但它的 labels 是-100? pad_id=0在本 toy 数据中解决什么问题?attention_mask=[1,1,1,1,1,1,1,0]和 labels 中的-100分别控制什么?- 为什么
z2对应 target40? - 为什么 shift 后时间长度从 8 变成 7?
- 为什么全
-100labels 可能得到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.辅助理解


