指令微调:SFT


文档摘要

指令微调:SFT 本节摘要:预训练基础模型能续写序列,但听不懂指令。监督微调(Supervised Fine-Tuning, SFT)是修复它的最小改动:喂模型成对的「指令 + 期望回答」,训练主体预测回答 token。诀窍是你只想让损失算在回答上,不算在指令上。本节构建 Alpaca 式 SFT 循环,用自定义 collate 函数以 掩码指令 token,在 200 对指令-回答上训练,用留出集上的精确匹配(exact-match)评估。技巧全在掩码: 与 边界 token 把序列划成三区,collate 把指令区与填充区都置 ,模型前向看见整序列、注意力能关注指令,但损失只数回答 token——恰是「以指令为条件、预测回答」。

指令微调:SFT

本节摘要:预训练基础模型能续写序列,但听不懂指令。监督微调(Supervised Fine-Tuning, SFT)是修复它的最小改动:喂模型成对的「指令 + 期望回答」,训练主体预测回答 token。诀窍是你只想让损失算在回答上,不算在指令上。本节构建 Alpaca 式 SFT 循环,用自定义 collate 函数以 ignore_index=-100 掩码指令 token,在 200 对指令-回答上训练,用留出集上的精确匹配(exact-match)评估。技巧全在掩码:<INST><RESP> 边界 token 把序列划成三区,collate 把指令区与填充区都置 -100,模型前向看见整序列、注意力能关注指令,但损失只数回答 token——恰是「以指令为条件、预测回答」。

对应原课程:Phase 19 · Lesson 39 · instruction-tuning-sft(原英文 phases/19-capstone-projects/39-instruction-tuning-sft/docs/en.md)。本节属「从零构建 GPT」赛道第十节。

学习目标

阅读完本节,你应当能够:

  1. 把成对的指令-回答格式化成带显式边界 token 的单一因果序列。
  2. 构建 collate 函数,掩码指令 token,使交叉熵只算回答 token。
  3. 在微型 transformer 上用 SFT 目标训练,看着评估指标动起来。
  4. 实现尊重回答起始边界的贪心与温度采样生成。
  5. 在留出集的生成补全上算精确匹配。

一、问题与直觉

用下一 token 预测训练的基础模型不知道什么是指令。给它 What is the capital of France?,它会续写问题或编个新句子——它有语言能力,但没有格式契约。

SFT 契约是一个字符串模板。每个训练例变成单一序列,有三个区:

<INST> What is the capital of France? <RESP> The capital of France is Paris.

边界 token 是训练时保留的特殊 token。模型学到 <RESP> 之后都是回答,而回答是要被打分的。基础模型的下一 token 目标依然适用,只是训练在每例都有这种形状的语料上。

ignore_indextorch.nn.functional.cross_entropy 的特性:任何等于 ignore_index 的目标位贡献零损失、零梯度,PyTorch 约定用 -100。collate 函数每例建两个张量:input_ids(全序列)与 labels(input_ids 的副本,指令位覆写为 -100)。模型前向看见整序列、注意力能关注指令,损失只数回答 token。

边界 token 与填充

分词器是字节级,留三个特殊 token:INST_ID = 256(指令区起点)、RESP_ID = 257(指令/回答分界)、PAD_ID = 258(变长批填充)。序列是 [INST] inst_bytes [RESP] resp_bytes [PAD]*

位移是标准因果技巧:input_idsi 预测位 i+1,故 labels[i] = input_ids[i+1](末位从输入丢、首位从目标丢)。掩码在位移后施加,落在正确位上。三处置 -100:指令区、填充区、RESP_ID 边界位本身(不训练模型预测边界 token,它预测紧随其后)。

二、从零实现

数据是 main.py 里确定性生成的 200 对指令-回答,覆盖六种任务:事实单点(某国首都)、算术、列表抽取、一句话总结、代码(print/sort)、定义。每任务有模板化指令与确定性回答——刻意简单,因精确匹配脆弱,本节用「对答案是一个特定字符串」的夹具。拆 160 训练/40 测试,测试集覆盖全部六类,可报每类精确匹配。

模型够小(隐 96、2 块、最大长 64),能在 CPU 上 2 分钟内训到收敛。循环是标准 PyTorch SFT 循环:Adam、学习率 3e-41e-3、1020 epoch、无调度器。

def collate(batch, tok, max_len): input_ids, labels = [], [] for inst, resp in batch: ids = [INST_ID] + tok.encode(inst) + [RESP_ID] + tok.encode(resp) ids = ids[:max_len] + [PAD_ID] * max(0, max_len - len(ids)) # 因果目标:错位一位;指令/填充/边界位置 -100 lab = ids[1:] + [PAD_ID] lab = [t if (i >= len([INST_ID]+tok.encode(inst)) # 在回答区 and ids[i] != PAD_ID) else -100 for i, t in enumerate(lab)] input_ids.append(ids); labels.append(lab) return torch.tensor(input_ids), torch.tensor(labels)

训练循环与第 34 节同构,只是损失自带 -100 掩码:

for ep in range(epochs): for x, y in train_loader: logits = model(x) # (B, T, V) loss = F.cross_entropy(logits.view(-1, V), y.view(-1), ignore_index=-100) loss.backward(); opt.step(); opt.zero_grad() if ep % 5 == 0: print("exact_match", evaluate_em(model, test_loader))

生成时,模型拿到指令前缀 [INST] inst_bytes [RESP] 并生成 token,直到序列达 max_len 或触发停止启发式(两个连续句尾字节 ./!/?)。精确匹配用贪心解码(温度会让指标随机)。看精确匹配从 epoch 1 的 0.0 升到 epoch 15 的约 0.85,是本节的回报:你能看见模型同时学会格式与答案。

设计要点:精确匹配是最严的文本指标——预测回答串归一化(小写、去空白、折叠双空格)后与同样归一化的参考逐字符比,每例 1 或 0,聚合是均值。它不含糊:说 0.7 就是 70% 的测试指令字符级命中黄金答案。真实 SFT 管线常配 token F1(第 39 节)与判官模型,但精确匹配仍是锚点。

三、框架对比

HuggingFace trlSFTTrainer 把模板、collate、掩码打包:配 chat template、DataCollatorForCompletionOnlyLM 自动在响应区外置 -100。本节手写让你看清掩码怎么落位、边界 token 怎么定义。Alpaca/Vicuna/ShareGPT 各有模板格式,但底层契约相同:指令区不算损失、回答区算。本节用字节级分词器与三个特殊 token,换到 BPE 分词器与 Llama chat template 时,只需改模板字符串与分词器,掩码逻辑不动。Axolotl、LLaMA-Factory 等工具是这套手写循环的产品化封装。

四、可复用产物

main.py:demo 在 CPU 上约 2 分钟训完,每 5 epoch 打印留出集精确匹配,典型从 0.0 升到约 0.85。collate 函数、边界 token 约定、生成停止启发式均可复用——换到任何因果 LM 上,只要模型有「输入 + 错位目标」接口,SFT 循环原样工作。

五、练习

  1. 六类细分:评估时按六类任务分别报精确匹配,找出模型最弱的一类。
  2. 模板变体:换模板格式(加 <system> 角色、换分隔符),重训,观察是否影响收敛。
  3. 停止启发式:把「两个句尾字节」停止改成「生成 <END> token」,重训,对比生成长度控制。
  4. 温度生成:评估时加温度采样 + 多数投票(同一指令采样 5 次取最常见答案),对比贪心的精确匹配。
  5. token F1:除精确匹配外加 token 级 F1(第 39 节),报告两个数字,观察哪个更宽容。

本节要点回顾

  1. 基础模型不懂指令:有语言能力但无格式契约,SFT 是修复它的最小改动。
  2. 契约是模板:每例变 [INST] inst [RESP] resp 三区序列,边界 token 划区。
  3. 掩码是诀窍:ignore_index=-100 让指令/填充/边界位零损失零梯度,只数回答 token。
  4. 因果位移:目标错位一位,掩码在位移后施加落在正确位。
  5. 前向看全序列:注意力能关注指令,损失只算回答——恰是「以指令为条件、预测回答」。
  6. 精确匹配最严:字符级 1/0,真实管线配 F1 与判官,但精确匹配是不含糊的锚点。

下一节,我们做「DPO 从零」——把奖励模型 + PPO 两阶段塌缩成单一监督损失,直接在偏好对上训策略。


发布者: 作者: Rohit Gupta 转发
评论区 (0)
U