跳到主要内容

数据截至 (上游 commit 92d63d4e8bb4)

03 · 后训练链:SFT 与 RL

这一章讲什么: base 模型怎么变成 chat 模型。两步:SFT(scripts/chat_sft.py)教它对话格式与能力,RL(scripts/chat_rl.py)在数学题上按「答对奖励」再推一把。


1. 它要解决的小问题

预训练只教会模型「续写文本」。要变成 ChatGPT 还缺三样:

  • 格式:听懂 <|user_start|> 之后的提问,用 <|assistant_end|> 收尾,会写工具调用。
  • 能力:多选题按格式作答、数学题列式计算、写简单代码。
  • 拔高:SFT 只能模仿示例,RL 可以按「对错」直接优化——示例里没有的策略也能被发现。

2. SFT:任务混合 + mask + packing

数据:Task 抽象与混合

  • 所有数据集实现统一接口 Tasktasks/common.py:85-126):get_example(i) 返回 conversation,evaluate(conversation, completion) 打分,eval_type 标记 generative / categorical。
  • 五个任务分工:
任务教什么eval_type
SmolTalk(46 万行通用对话)聊天(仅训练用)
MMLU / ARC多选题格式categorical
GSM8K数学 + 计算器工具调用generative
HumanEval简单 Python(evaluate 直接 execute_code 跑测试,tasks/humaneval.py:79-96generative
  • TaskMixturetasks/common.py:129-161):把多个 Task 拼成一个逻辑数据集,种子 42 确定性打乱;想过采样某个任务,就在列表里重复传几遍——SFT 里 MMLU ×3、GSM8K ×4 就是这么来的(scripts/chat_sft.py:162-168)。
  • 数据集下载不依赖 HF datasets 库:load_hub_dataset 直接调 hub 的 parquet 导出 API,FileLock 保证多 rank 只下载一次(tasks/common.py:45-82)。

装载:bestfit-pad(对比预训练的 bestfit-crop)

sft_data_generator_bos_bestfitscripts/chat_sft.py:180-298)和预训练装载是同一个 best-fit 骨架,但有一条关键差异:

预训练SFT
放不下时裁剪最短文档填满padding BOS 填满
为什么文档海量,裁得起对话有限,一条都不能丢
代价~35% token 被裁padding 位不算 loss

mask 落地方式:渲染给的 mask 平移一位对齐 targets,mask=0 的位置写成 -1(PyTorch cross_entropy 的 ignore_index),padding 位同样 -1(scripts/chat_sft.py:288-296)。

训练:从预训练「继承一切」

  • 超参继承max_seq_len / device_batch_size / total_batch_size / 三个 lr 默认从 base checkpoint 的 meta 里继承(scripts/chat_sft.py:96-115)。
  • 优化器热启动:加载 base 的 optimizer state(动量 buffer 有价值),但 load_state_dict 会把 group 元数据(lr 等)也覆盖成预训练末期的 ≈0 值——所以先存新 lr、加载后再写回(scripts/chat_sft.py:134-147)。这是容易踩的坑,注释写得很直白。
  • **lr 调度改用 progress(0→1)**而不是绝对步数:SFT 是数据驱动停止,总步数事先不知道(scripts/chat_sft.py:304-314);last_stepall_reduce(MAX) 跨 rank 同步,防分布式挂起(scripts/chat_sft.py:333-337)。
  • 评估:val bpb + ChatCORE——五个任务各自减去随机基线再平均(0 = 随机,1 = 满分),定义在 scripts/chat_sft.py:376-381scripts/chat_eval.py:228-238

3. RL:GSM8K 上的极简策略梯度

直觉

作者把算法叫「"GRPO"(带引号)」,文件头自己交代了四刀(scripts/chat_rl.py:1-10):

  1. 删 trust region——没有 reference model、没有 KL 正则。
  2. on-policy——采样和训练同分布,不需要 PPO ratio + clip。
  3. DAPO 式 token 级归一化,不按序列归一。
  4. 优势不做 z-score,只用 r − mean(r)

砍完剩下的就是 REINFORCE:同一道题采一组答案,组内奖励减均值当优势,按优势加权每个 token 的 logp

原理演示

# 示意,非源码
for question in gsm8k_train: # 各 rank 分片
prompt = render_for_completion(question) # 请 assistant 作答
samples = engine.generate_batch(prompt, n=16) # 一次 prefill,克隆 KV,采 16 答
rewards = [score(s) for s in samples] # 答对 1,答错 0
advantages = rewards - mean(rewards) # 组内相对好坏
logp = -model(inputs, targets, loss_reduction='none') # 每 token 的 log p
loss = -(logp * advantages).sum() / num_valid_tokens # token 级归一
loss.backward()

重点看:没有 reference model、没有重要性采样比、没有 clip——因为采样的模型就是要更新的模型(on-policy),三步合成一步。

真实实现

  • 采样批get_batchscripts/chat_rl.py:86-146)。engine.generate_batch 的 num_samples 路共享一次 prefill(第四章);返回的 mask 里 prompt 和工具 forced token 都是 0,随后写成 -1 不进 loss(scripts/chat_rl.py:128-140,注释明确「correctly 不训练 prompt 和工具输出」)。
  • 奖励GSM8K.reward 直接复用 evaluate——extract_answer#### (\-?[0-9\.\,]+) 后的数字与标准答案比对(tasks/gsm8k.py:21-33109-116)。
  • 优势advantages = rewards - rewards.mean()scripts/chat_rl.py:142-144)。
  • 损失logp = -model(inputs, targets, loss_reduction='none')pg_obj = (logp * advantages).sum();除以 num_valid × num_passes × examples_per_rank(token 级归一化);loss = -pg_objscripts/chat_rl.py:263-272)。
  • lrinit_lr_frac = 0.05(比 SFT 的 0.8 小一个数量级),再线性降到 0(scripts/chat_rl.py:55209-212)。

评测:pass@k

run_gsm8k_evalscripts/chat_rl.py:150-191)在 test 集上每题采 k 个答案,pass@k = 至少一个答对的比例;分布式各 rank 分片评测再 all_reduce(scripts/chat_rl.py:225-243)。


4. 关键细节 / 坑

  • SFT 的「停止」是数据驱动的:不吃完一遍数据不停(除非显式 --num-iterations),进度条是 consumed / dataset_size 的近似(scripts/chat_sft.py:268-277)。
  • RL 采样要探索:temperature=1.0 + top_k=50;评测时 temperature=0(scripts/chat_rl.py:47-49150-156)。
  • RL 不存 optimizer state:checkpoint 只存模型(scripts/chat_rl.py:314-322)。
  • 多选格式的小心思render_mc 把字母放在选项之后、且 = 与字母之间无空格——因为 tokenizer 对 " A""A" 是不同 token,而 assistant 要回答的是裸 "A"tasks/common.py:187-206 的 docstring)。小模型对这种细节敏感。
  • RL 暂不支持 fp16(README dtype 节原话:「SFT supports this too but RL currently does not」);SFT 走 GradScaler 路径(scripts/chat_sft.py:151-154)。
  • SFT 混合比例是手拍的:SmolTalk + MMLU×3 + GSM8K×4,val 集也按近似比例抽样对齐(scripts/chat_sft.py:169-173);换配方要自己改这段。

下一章:04 · Engine 推理