HOMEWORK 05

拒绝采样

生成 N 条、用奖励模型挑一条、拿去做 SFT——最朴素的偏好优化算法。这一组作业真正要教的不是流程,而是「你怎么知道涨的那几个点是奖励模型的功劳」。

对应章节:第 9 章 · 拒绝采样 参考实现:_src/code/rejection_sampling/ 配置:rejection_sampling/configs/*.yaml 状态:⚠ 本机未实跑

0. 任务目标

本组尚未实跑 本组作业在本次会话中没有实际运行。原因见 homework/RESULTS.md:这条流水线要先批量生成、再逐条打分,显存与时间预算都明显大于前几组——策略模型是 Qwen/Qwen3-1.7B,奖励模型是 nvidia/AceMath-7B-RM,而本机实际只有约 6.9 GiB 可用显存。

因此,下面的内容是任务设计与预期观察点,不是实测结果。本页不给出任何具体的实验数值——只讲该盯哪些指标、指标之间该有什么关系、以及怎么设计才能让结论站得住。真正的实测记录只在 homework/RESULTS.md 里,那里目前只有 HW1、HW2、HW4 三组。

拒绝采样(Rejection Sampling, RS)是第 9 章的主题,也是最简单的一种「用奖励改进策略」的方法:让模型对每个 prompt 生成 $N$ 条回复,用奖励模型给每条打分,挑出好的那些,再拿它们做一轮 SFT。没有 RL 循环,没有优势估计,没有 KL 惩罚——它就是「采样 + 筛选 + 模仿」。

正因为简单,它在工业界用得极多(Llama 2/3 的后训练流水线里都有它)。但也正因为简单,它有一个特别容易被忽略的方法论陷阱,而这个陷阱就是本组作业的全部重点:

本组要回答的唯一问题 假设你跑完 top_per_prompt,测试集准确率涨了几个点。这几个点是奖励模型的功劳吗?

不一定。至少还有三个解释同样成立:(1) 你只是在模型自己生成的数据上又做了一轮 SFT——这是自蒸馏,它本身就能提升表现;(2) 你只是多采样了——从 8 个候选里挑一个再训,哪怕随机挑,训练分布也已经变了;(3) 你只是强化了某种格式——被选中的回复恰好都以某种方式结尾,评测脚本的答案抽取正则因此匹配率更高(这在 GSM8K 上特别常见,而它和推理能力毫无关系)。

要把这三个解释排除掉,唯一的办法是同预算的随机对照组。这是本组的灵魂,第 2、3 节整整两节讲它。

所以这一组作业的交付物不是「一个准确率数字」,而是一对准确率数字之差。学会这件事之后,你看任何后训练论文里的消融表格,视角都会变。

1. 流水线与运行命令

参考实现把第 9 章的四个阶段完整实现了一遍,任务是 GSM8K(小学数学应用题):

  1. Stage 1 · 生成:用 Qwen/Qwen3-1.7B 对每个训练 prompt 生成 $N$ 条解题过程(num_completions_per_prompt: 8)。
  2. Stage 2 · 打分:用 nvidia/AceMath-7B-RM 给每一条打分。
  3. Stage 3a · 选择:按某种策略挑出一个子集。四种策略的差别全在这一步,前两个阶段完全一样。
  4. Stage 3b · SFT:在选中的 (prompt, completion) 对上微调同一个 Qwen3-1.7B。
  5. Stage 3c · 评测:在 GSM8K 测试切片上量 exact-match 准确率,记为 test_accuracy。

默认切片是 1000 条训练 / 200 条测试,生成用 temperature: 0.8、top_p: 0.95、max_new_tokens: 512,SFT 用 lr 5e-6、2 个 epoch,评测用 eval_temperature: 0.0(贪心,保证可复现)。

先建缓存,再跑四个配置

cd _src/code/

# (1) 只做一次:Stage 1 + Stage 2,产出一份 JSONL 缓存
python -m rejection_sampling.preprocess \
    --config rejection_sampling/configs/top_per_prompt.yaml

# (2) 四个选择策略各训一次(train.py 首次调用会自动跑 preprocess,之后命中缓存)
python -m rejection_sampling.train --config rejection_sampling/configs/top_per_prompt.yaml
python -m rejection_sampling.train --config rejection_sampling/configs/random_per_prompt.yaml
python -m rejection_sampling.train --config rejection_sampling/configs/top_k_overall.yaml
python -m rejection_sampling.train --config rejection_sampling/configs/random_k_overall.yaml

缓存文件每行是一个 prompt 的全部打分结果:

{"question": "...", "answer": "72",
 "completions": ["...", "..."], "rewards": [0.71, 0.33, ...]}
缓存机制不只是省时间 Stage 1 和 Stage 2 很贵,选择很便宜。实现把生成/打分的全部参数哈希成一个短 key,写到 rejection_sampling/output/rollouts/<hash>.jsonl;四个 YAML 的生成与打分参数刻意设成逐字相同,于是它们哈希到同一个文件。

这不只是工程优化——它是「四个策略面对的候选池完全相同」这个对照前提的机制保证。任何一个生成参数被改动(奖励模型、策略模型、数据切片、采样参数、种子、$N$、max_new_tokens),哈希就变、缓存就失效、对照就断了。想强制重跑就删掉那个哈希文件。

四种选择策略

策略保留什么产出对数角色
top_per_prompt每个 prompt 里奖励最高的那一条$M$经典 RS:每题一条,覆盖全部 prompt
random_per_prompt每个 prompt 里随机一条$M$top_per_prompt 的对照组
top_k_overall整个 $M\times N$ 池里奖励最高的 $K$ 条$K$允许同一题贡献多条,也允许某些题一条都没有
random_k_overall扁平池里随机 $K$ 条$K$top_k_overall 的对照组

两个 *_k_overall 配置的 selection.top_k 必须写成同一个值(默认 1000),否则样本预算不对齐,对照就无效了。

2. 为什么必须有 random 对照组

这一节是本组作业的灵魂,重要到值得占一整节。

问题的形状

拒绝采样的完整操作是「生成 → 打分 → 选择 → SFT」。你想检验的假设是「用奖励做选择比不用好」,但你实际做的动作里包含了至少四个可能独立产生效果的成分:

成分它自己能不能涨分和奖励模型有关吗
在自己生成的数据上再做一轮 SFT(自蒸馏)能——收敛到自身高概率模式、减少采样方差、强化格式一致性无关
从 $N$ 个候选里挑(哪怕随机挑)能——训练分布已经和原始输出分布不同了无关
被选中的回复恰好格式更规整能——评测的答案抽取正则匹配率上升间接、且是混杂因素
奖励模型真的识别出了更好的解法能这才是你想测的

只跑一个 top_per_prompt,然后拿它和「原始模型」比,你得到的是这四项的总和。而论文里、简历里、周报里想写的那句话——「我们的奖励模型带来了 X 个点的提升」——对应的只是最后一项。

常见误区 把「RS 流水线的总收益」当成「奖励模型的收益」。这是后训练实验里最高频的归因错误之一,而且它几乎总是朝着「高估自己方法」的方向偏——因为自蒸馏和多采样的基础收益往往不小。

更糟的是,这个错误无法通过跑更久、跑更大来修正:它是实验设计的问题,不是统计功效的问题。设计错了,跑一万次也还是错的。

对照组把式子改成了减法

设 $A_{\text{base}}$ 是原始模型的准确率,$A_{\text{rand}}$ 是随机选择后训练的准确率,$A_{\text{top}}$ 是奖励选择后训练的准确率。那么:

差值它测的是什么
$A_{\text{rand}} - A_{\text{base}}$「多采样 + 再训一轮」本身值多少分。这个值往往不小,正是没有对照组时最容易被误读成 RM 效果的那部分。
$A_{\text{top}} - A_{\text{rand}}$奖励模型的净贡献。这才是本组作业要报告的数字。
$A_{\text{top}} - A_{\text{base}}$整条流水线的总收益。有工程价值,但不能用来论证奖励模型好。

换句话说:随机对照组把一个绝对值问题变成了差值问题。而差值对很多系统性误差是免疫的——两组共享同一份候选池、同一套训练超参、同一个评测脚本,任何影响两边的因素都会在相减时抵消掉。这正是对照实验的全部威力所在。

上游的一个负结果,值得先知道 上游 README 记录了参考运行的定性结论:在 1k 训练 / 200 测试的 GSM8K 切片上,top_k_overall 胜过它的随机对照,而 top_per_prompt 与 random_per_prompt 基本打平;README 明确提醒把这种小差距当作切片噪声,不要当成稳定的策略排序。(这是上游的观察,不是本机的实测结果。)

这个负结果本身就是最好的教学材料:如果你的 top 跑不赢它的随机基线,正确的结论不是「RS 没用」,而是「在这个切片上,奖励模型或候选池没有提供有用的信号」——然后你应该回去看诊断(第 4 节),而不是继续调训练超参。

3. 对照组必须锁死哪些变量

一个合格的对照组必须只改变你想检验的那一个变量,其余全部保持一致。这里想检验的变量是「选择时用不用奖励」,所以必须锁死下面五项:

必须锁死为什么实现里怎么保证
候选池 如果两组面对的 8 条候选不是同一批,你比的就是两次不同的采样运气 生成/打分参数逐字相同 → 哈希相同 → 共享同一份 rollouts/<hash>.jsonl
训练样本数 样本多的一方天然占优,差值就没法归因了 random_per_prompt 输出 $M$ 条,和 top_per_prompt 一样;random_k_overall 的 top_k 必须和 top_k_overall 写成同一个值
选择的结构 「每题一条」和「扁平池取 K 条」的 prompt 覆盖度完全不同,这本身就会影响结果 per-prompt 版本都是全 prompt 覆盖,overall 版本都从扁平池取——只能同结构互比
训练超参 lr / epoch / batch / 序列长度任何一项不同,差值都会被污染 四个 YAML 的 Stage 3b 段落逐字相同
随机种子 SFT 的数据顺序、dropout、评测采样都受它影响 四个配置都是 seed: 42;随机选择本身也由 cfg.seed 播种,因此可复现

锁死之后,两个配置之间唯一的差别就落在选择那一行代码上:max(rewards) 换成了 rng.randrange(len(completions))。这就是「受控对照」在代码层面的样子。

配对方式不能交叉 只能 top_per_prompt ↔ random_per_prompt、top_k_overall ↔ random_k_overall。 拿 top_k_overall 去和 random_per_prompt 比是无效对照——它们的样本数和结构都不一样,差值里混着「选择结构」这个第二变量。

一般原则:对照组是「一对」,不是「一组基线」。每加一个实验条件,就要问一句「它的匹配对照是谁」。

还有两个容易漏掉的坑

其一,评测脚本本身要固定。GSM8K 的 exact-match 依赖答案抽取逻辑(从生成文本里把最终数字抠出来)。如果你在两次运行之间改了这个正则,得到的差值毫无意义。配置里 eval_temperature: 0.0 就是出于同样的考虑——贪心解码去掉了评测阶段的随机性。

其二,别忘了记录基线。$A_{\text{base}}$(完全不训练、直接评测原始模型)不在四个配置里,但它是第 2 节那张差值表的第一行。建议手动补一次:它能告诉你「多采样 + 再训一轮」这部分收益有多大,而这往往比 RM 的净贡献大得多——这个对比本身就很有教育意义。

4. 该盯什么指标

训练与评测阶段

指标含义怎么读
test_accuracyGSM8K 测试切片上的 exact-match 准确率本组的最终因变量。但永远成对读,单独一个数说明不了任何事
selection/num_pairs选出来的训练对数对照双方必须相等。不相等就说明对照设计破了,别看后面的数了
selection/strategy本次运行用的策略名用来在看板里配对,防止把两次运行搞混
sft/lossStage 3b 的训练损失只用来确认训练正常。它低不代表模型好——在「更容易拟合」的数据上损失自然更低,这恰恰可能是选择偏置的副作用
sft/grad_norm梯度范数看有没有异常尖峰

建议整理成这样一张表(数字栏留给你自己填):

配置训练样本数test_accuracy相对随机对照的 gap
基线(不训练)0——
random_per_prompt$M$—(基准)
top_per_prompt$M$—← RM 的净贡献
random_k_overall$K$—(基准)
top_k_overall$K$—← RM 的净贡献

先跑诊断,再跑训练

这条建议能省掉大量算力。rejection_sampling/diagnostics.py 直接读缓存文件,不需要训练任何模型:

python -m rejection_sampling.diagnostics \
    --cache rejection_sampling/output/rollouts/<hash>.jsonl \
    --out-dir rejection_sampling/output/diagnostics

它给出三个视图加一个汇总数:

  1. 奖励直方图(正确 vs 错误回复)——两个分布分不分得开,就是奖励模型的原始分辨力。如果重叠得几乎完全,后面的训练大概率也不会有 gap,可以直接省下来。
  2. 行内选择命中率——在可判定的 prompt 上(即那些既有正确又有错误回复的行),$\argmax(\text{reward})$ 选中正确回复的比例,对比随机基线的比例。这是在训练之前就能拿到的「RM 净贡献」估计,比跑完整训练便宜几百倍。
  3. best-of-$N$ 扫描——$N$ 从 1 到 8,「按奖励取前 $N$ 条里至少有一条正确」的比例,对比「随机取 $N$ 条」。两条线的间距如果很快饱和,说明再加 $N$ 也没用。
  4. decidable_fraction——可判定 prompt 的占比,也就是「行内选择有可能起作用」的那部分数据有多少。
decidable_fraction 是本组最该理解的一个数 如果一个 prompt 的 8 条回复全对,选谁都一样;全错,选谁也都一样。这两类 prompt 对「top vs random」的差值贡献严格为零——无论奖励模型多准。

所以:行内选择的效果上限,被 decidable_fraction 直接卡死。在 GSM8K 这种对 1.7B 模型已经不算难的数据集上,大量 prompt 是 8 次全对,可判定比例天然很低,top_per_prompt 和 random_per_prompt 打平是完全可以预期的——那是数据侧的天花板,不是奖励模型的错。

留意这和 HW3 的「组内对比度」是同一个概念的两种化身:GRPO 里全对/全错的组优势为 0,RS 里全对/全错的行选择无效。凡是靠「同一 prompt 内部比较」吃饭的方法,都会被这件事卡住。

5. 自选消融

消融一:扫 $N$(每题采几条)

改 num_completions_per_prompt(4 / 8 / 16),其余不动。这是第 9 章那条「收益随 $N$ 增长而递减」曲线的实测版本。

观察目标有两个:(1) 候选池变大,是不是真的让「最好的那条」变好了?(2) 还是只让奖励模型有更多机会被噪声骗到($N$ 越大,越可能出现一条「奖励很高但其实是错的」的回复——这是 best-of-$N$ 版本的过度优化)?

算力预算要先算清楚 $N$ 一改,哈希就变,整份缓存必须重新生成和重新打分。这是本组最贵的一个消融:生成量与 $N$ 成正比,打分量也与 $N$ 成正比,而打分用的是一个 7B 模型。改 $N$ 之前先估一遍时间和显存。

消融二:top_per_prompt 与全局 top-$K$ 的结构差异

这两个策略的差别不是「哪个更好」,而是它们在优化不同的东西:

top_per_prompttop_k_overall
prompt 覆盖全覆盖,每题恰好一条不均匀:有的题贡献多条,有的题一条都没有
依赖奖励的什么性质只依赖同一题内部的排序(跨题尺度可以不一致)依赖跨题可比的绝对尺度
失败模式decidable_fraction 低时无效把训练数据集中到「奖励模型打分偏高的那类题」上,引入分布偏置

建议顺便统计一件事:$K$ 取当前值时,有多少 prompt 完全没有代表?这个数直接量化了上面那条偏置。$K$ 越小,覆盖越偏。

再进一步的变体(值得一试):先做组内归一化再全局取 top-$K$。对每一行做 $\tilde r_{i,j} = (r_{i,j} - \mu_i)/\sigma_i$,然后在归一化后的分数上取全局 top-$K$。如果这个变体明显好过原始的 top_k_overall,你就实证了「跨 prompt 的奖励尺度不可比」确实是个真问题——顺便也把自己推到了 GRPO 的门口(HW3 的组内 z-score 正是同一个思路)。

消融三:把 decidable_fraction 当自变量

前面说过,行内选择的效果上限被可判定比例卡死。那就主动去改变它:

  • 换一个更弱的策略模型(同时调小 max_train_samples 省算力)。弱模型在 GSM8K 上全对的行会变少,decidable_fraction 应该显著升高。如果奖励模型有真信号,那么 top 与 random 的 gap 也应该跟着变大。这是一个非常干净的因果检验,也回答了一个关键问题:RS 是在「拯救糟糕的生成」,还是只是在「一堆本来就还行的答案里挑挑拣拣」?
  • 换一个更难的数据集切片,效果类似。
  • 迭代两轮:用第一轮 RS 训出的模型重新生成、重新打分、再训一次。预期是第二轮增益明显小于第一轮,同时 decidable_fraction 下降(模型变强了,全对的行变多了)。这是第 9 章「迭代式 RS」的最小演示,也能让你亲眼看到为什么奖励模型需要跟着策略一起更新。

这三个消融的共同点是:它们都不是在调训练超参,而是在改变「奖励信号有没有发挥空间」。这正是本组作业希望你养成的直觉。

本组小结

问题答案
本组最重要的一件事?设同预算的随机对照组。没有它,你分不清涨分来自选择策略,还是单纯多采样 / 自蒸馏 / 格式强化。
怎么读结果?成对读:top_per_prompt 对 random_per_prompt,top_k_overall 对 random_k_overall。差值才是奖励模型的净贡献。
对照组要锁死什么?同一份 rollout 缓存、同样的样本数、同样的选择结构、同样的训练超参和种子。只留「用不用奖励」这一个变量。
训练之前能做什么?跑 diagnostics:奖励直方图、行内命中率、best-of-$N$ 扫描、decidable_fraction。比跑完整训练便宜几百倍。
如果 top 跑不赢 random?结论不是「RS 没用」,而是「在这个切片上,奖励模型或候选池没提供有用信号」。回去看诊断,别调训练超参。

交付物清单

  • 四个配置的完整日志与 test_accuracy,以及成对整理的对照表。
  • rollout 缓存的哈希文件名 —— 用来证明四次训练确实共享了同一个候选池。
  • diagnostics 的三张图和 decidable_fraction。
  • 一段结论,明确区分「多采样 + 再训一轮的贡献」和「奖励模型的净贡献」两部分。