拒绝采样
生成 N 条、用奖励模型挑一条、拿去做 SFT——最朴素的偏好优化算法。这一组作业真正要教的不是流程,而是「你怎么知道涨的那几个点是奖励模型的功劳」。
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(小学数学应用题):
- Stage 1 · 生成:用
Qwen/Qwen3-1.7B对每个训练 prompt 生成 $N$ 条解题过程(num_completions_per_prompt: 8)。 - Stage 2 · 打分:用
nvidia/AceMath-7B-RM给每一条打分。 - Stage 3a · 选择:按某种策略挑出一个子集。四种策略的差别全在这一步,前两个阶段完全一样。
- Stage 3b · SFT:在选中的 (prompt, completion) 对上微调同一个
Qwen3-1.7B。 - 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, ...]}
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 个点的提升」——对应的只是最后一项。
更糟的是,这个错误无法通过跑更久、跑更大来修正:它是实验设计的问题,不是统计功效的问题。设计错了,跑一万次也还是错的。
对照组把式子改成了减法
设 $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}}$ | 整条流水线的总收益。有工程价值,但不能用来论证奖励模型好。 |
换句话说:随机对照组把一个绝对值问题变成了差值问题。而差值对很多系统性误差是免疫的——两组共享同一份候选池、同一套训练超参、同一个评测脚本,任何影响两边的因素都会在相减时抵消掉。这正是对照实验的全部威力所在。
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_accuracy | GSM8K 测试切片上的 exact-match 准确率 | 本组的最终因变量。但永远成对读,单独一个数说明不了任何事 |
selection/num_pairs | 选出来的训练对数 | 对照双方必须相等。不相等就说明对照设计破了,别看后面的数了 |
selection/strategy | 本次运行用的策略名 | 用来在看板里配对,防止把两次运行搞混 |
sft/loss | Stage 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
它给出三个视图加一个汇总数:
- 奖励直方图(正确 vs 错误回复)——两个分布分不分得开,就是奖励模型的原始分辨力。如果重叠得几乎完全,后面的训练大概率也不会有 gap,可以直接省下来。
- 行内选择命中率——在可判定的 prompt 上(即那些既有正确又有错误回复的行),$\argmax(\text{reward})$ 选中正确回复的比例,对比随机基线的比例。这是在训练之前就能拿到的「RM 净贡献」估计,比跑完整训练便宜几百倍。
- best-of-$N$ 扫描——$N$ 从 1 到 8,「按奖励取前 $N$ 条里至少有一条正确」的比例,对比「随机取 $N$ 条」。两条线的间距如果很快饱和,说明再加 $N$ 也没用。
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$ 版本的过度优化)?
消融二:top_per_prompt 与全局 top-$K$ 的结构差异
这两个策略的差别不是「哪个更好」,而是它们在优化不同的东西:
top_per_prompt | top_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。- 一段结论,明确区分「多采样 + 再训一轮的贡献」和「奖励模型的净贡献」两部分。