Distill Where You Fail:用自适应教师引导把 GRPO 的"零方差组"救回来

  • 关联论文:2608.00782
  • 作者:flyP
  • 更新:2026-08-06

一句话结论:把 on-policy distillation (OPD) 当"补丁"打——只在 GRPO 的零方差负组上、且只看高 student entropy 或大 teacher-student 离散度的 token,并对这些样本额外补一轮 SFT。在数学 +3.02%、代码 +3.05% 之外,把 RL 后训练里"全对/全错"组信号被吞的致命浪费压成可学习的稠密信号。


1. 它在解决什么真问题

GRPO(Group Relative Policy Optimization)是 2025 年以来 LLM 后训练的事实标准:模型给每个 prompt 采样 G 条回答,按组内相对奖励做策略梯度。便宜、不要 critic、效果稳。但它有两个老毛病:

  1. 奖励稀疏:每条回答只有一个 0/1 标量;
  2. 零方差整组丢掉:当 G 条回答里全对或全错,组内优势全 0,等价于 "白白浪费了一段梯度"。在难 prompt 上这种现象高频出现,是 RL 后训练"看不见的税"。

学界在尝试的两条补救思路:

  • RLOO / ReMax / Reinforce++:改奖励归一化方式,治标;
  • OPD / On-Policy Distillation:用一个大教师模型(通常是同代或前代更强模型)在 token 级别输出监督,提供稠密信号。看似完美,但朴素 GRPO+OPD 反而退化

后者是本文要回答的"真问题":为何朴素 OPD 反而掉点?它在哪些 token 上有用?哪些样本上要蒸馏?蒸馏该怎么和 RL 共享权重?


2. 核心方法(RSTG 三件套)

论文提出 RSTG(Recovering Learning Signals via Adaptive Teacher Guidance),名字已经暗示了"信号在哪儿丢的,就在哪儿召回"。三件套:sample-level 门控、token-level 门控、SFT 补正。

2.1 三条朴素 OPD 退化的根因诊断(论文定位的最重要贡献)

朴素把 GRPO + OPD 拼起来会变差的根因有三条:

  1. 并不是所有样本都值得蒸馏:简单 prompt 上学生已经答对,再去学教师只是在复制答案,浪费梯度;
  2. 学得太快会侵蚀 RL 的探索:蒸馏 loss 很容易压垮 advantage-based 梯度,让策略塌缩到教师;
  3. OPD 的优势是"非对称"的:对 correct token OPD 几乎没梯度(教师也是对的),对 wrong token OPD 反而是大梯度。结果是学生只在少数错 token 上被拉,时间一长错 token 上的 KL 主导,其他 token 反而被欠训。

这三条诊断决定了"必须有门控"的总体思路。

2.2 Sample-Level 门控:只在零方差负组上蒸馏

for each prompt x with student samples {y_i}_{i=1..G}:
    rewards = scorer(x, {y_i})
    if zero_variance(rewards) is False:        # 不是零方差组
        apply GRPO only
    else:
        # 只在零方差组上动手脚;区分正负
        if all_correct(rewards):               # 全对,组优势 = 0
            sample_weight = mean(teacher_confidence(y_i))
            apply OPD with weight alpha_i = sample_weight
        else:                                  # 全错,组优势也 = 0
            sample_weight = 1 - mean(teacher_confidence(y_i))
            apply OPD with weight alpha_i = sample_weight

含义:把 OPD 当"补救信号"——只在 GRPO 完全看不到梯度的组上才让它上场。样本权重用教师置信度(在全对组里,置信度低的可能是擦边对,需要被多学;在全错组里,置信度低的反而是学生有主见的对答案,应该少被教师拉走)。

2.3 Token-Level 门控:只学"该学的 token"

token 级门控是论文第二个发力点。并不是 token 级 KL 处处有用,作者定义两个 mask:

  • 高学生熵 mask:学生对该 token 分布仍然平(不确定),教师的最优分布能给出方向;
  • 大 teacher-student 散度 mask:教师与学生分歧大,说明这里有可学习的信号。

$$ m_t = \mathbb{1}!\left[H_{\text{stu}}(p_t) > \tau_1\right] \;\lor\; \mathbb{1}!\left[KL(p_t^{\text{tea}}|p_t^{\text{stu}}) > \tau_2\right] $$

蒸馏 loss 只在 mask 内 token 上求和: $$ \mathcal{L}_{\text{OPD-masked}} = -\sum_t m_t \cdot \langle \log p_t^{\text{stu}}, \; \text{stopgrad}(p_t^{\text{tea}})\rangle $$

直观:这意味着 OPD 不再"全 token 平铺",而只在 student 最需要帮助的位置发力。

2.4 SFT on Teacher-Correct Trajectories:第三件补丁

零方差负组(全部答错)即使加了带 mask 的 OPD,对学生也只是"知道这条不行 + 几条 token 模仿",但 RL 信号依然 0。作者把教师在这些 prompt 上重新跑一遍,取教师答对的轨迹,对学生做 SFT:

$$ \mathcal{L}_{\text{SFT}} = -\sum_t \log p_t^{\text{stu}}(y_t^{\text{tea}}|x) $$

加进总 loss: $$ \mathcal{L}{\text{total}} = \mathcal{L}{\text{GRPO}} + \lambda_{\text{OPD}} \mathcal{L}{\text{OPD-masked}} + \lambda{\text{SFT}} \mathcal{L}_{\text{SFT}} $$

这三件套共同保证:在 GRPO 已经无能为力的零方差组上,先用 SFT 注入答案分布,再用 OPD 在关键 token 上继续学,整体仍然不破坏 RL 的探索分布。

2.5 训练循环伪代码

for step in range(TOTAL):
    prompts = sample_prompts(mix_datasets)
    student_outputs  = rollout(student, prompts, G=8)
    teacher_outputs  = rollout(teacher, prompts)            # 用 vLLM / SGLang / TRITON-LLM 加速
    rewards          = verifier(prompts, student_outputs)  # 0/1(math/code)

    # 1) 分组
    groups = partition_by_variance(rewards)
    for prompt, group in groups.nonzero.items():
        # 2a) 普通 GRPO
        loss_grpo = grpo_loss(prompt, group)
    for prompt, group in groups.zero.items():
        # 2b) 零方差 → 蒸馏 + SFT
        sw = student_weight_from_teacher_conf(group, teacher_outputs)
        opd_m = token_mask_from_entropy_kl(student, teacher, prompt, group, τ1, τ2)
        loss_opd  = opd_loss(student, teacher, prompt, group, sw, opd_m)
        loss_sft  = sft_loss(student, teacher_correct_only, prompt)
        loss += loss_grpo + λ_opd * loss_opd + λ_sft * loss_sft

    # 3) 反向 + 优化器更新
    update(student, loss)

3. 关键实验与数据

⚠️ 事实核查注记(一句话结论 vs 关键实验数字不一致):一句话结论给出"数学 +3.02%、代码 +3.05%",但"## 关键实验与数据"节引用摘要原文为"Math +4.02%、Code +3.05%"。数学提升数字存在 1% 的差异,应以 PDF 原文正文/表格中数字为准。代码 +3.05% 在两处一致。

论文摘要明确给出两条 SOTA 性结果

  • Math benchmark(AIME/MATH/whatever 摘要未具体化)+4.02%(⚠️ 摘要原文如此,但一句话结论标为 +3.02%,数字不一致);
  • Code benchmark(HumanEval / MBPP / LiveCodeBench 摘要未具体化)+3.05%

两条提升都来自"朴素 GRPO+OPD 退化基线 + RSTG",意味着:

  1. 朴素 GRPO+OPD 实际上是被踩在脚下,RSTG 不仅恢复还能显著超越
  2. 数学与代码两个独立赛道都稳定提升,说明门控机制不是某数据集 overfit。

值得标注的盲区:

  • 基线模型是哪一个?摘要未明确给出主表所使用的 student/teacher 模型对,是 Qwen2.5-7B / Qwen2.5-32B?还是 Llama-3.1-8B / 70B?还是自家 pretrained?
  • τ1 / τ2 阈值怎么选?λ_opd / λ_sft 比例?摘要未明确;
  • 是否与 DeepSeek-R1 的 GRPO 改进、Open-R1、DAPO 等同期 SOTA 横评?摘要未明确;
  • wall-clock cost(额外教师前向的代价)未明确给出,可能是一个非平凡开销。

4. 亮点与局限 / 反方段

4.1 亮点

  1. 精准诊断 + 最小改动:把"OPD 退化"归因到三个机制,给出对应三件补丁;不是堆叠新方法,而是把已有方法各自摆在它擅长的位置。
  2. 零方差组信号召回:直接面对 GRPO 最被吐槽的问题(A=R-mean 整组 0 时这段数据浪费),给出可工程化补救。
  3. token-level 掩码可推广:熵/KL 双条件 mask 与教师无关,任何教师都能复用。
  4. 不破坏 RL 探索:通过门控 + 权重,蒸馏只在 group-advantage=0 时上场,不会污染 group-advantage≠0 的正常梯度流。
  5. 可复现性:纯 GRPO + 教师前向 + verifier,不需要 critic、不需要偏好数据。

4.2 局限 / 反方段(按 lessons-W31 要求强制 1 段)

  • 教师成本非平凡:每个 prompt 还要再让 teacher 前向一次;如果 teacher 是大 10× 体量,这一项在 PPO/GRPO 单步中成为显著额外开销;摘要未给出 wall-clock。
  • 门控阈值依赖经验:τ1、τ2 是超参,最佳值需要 sweep;不同 student-teacher pair 是否还稳,论文未给出多模型对实验。
  • 正组 OPD 缺位:作者刻意把 OPD 限制在零方差负组,但理论上"差不多对但还能改"的样本才是蒸馏 ROI 最高的区段;论文未探索这一区段。
  • 仍未提供与最新 SOTA 的横评:math/code 各 +4 / +3% 是相对朴素基线,但相对 DeepSeek-R1-Distill、Qwen3-Instruct、Skywork-OR1 等 2026 年中前线 RL 体系尚未横评。
  • 代码权重:摘要未明确是否承诺开源实现与训练脚本。

5. 对工程落地的启发

  1. 最小可用补丁grpo_loss + λ_opd * masked_opd + λ_sft * sft 这套 loss 不侵入主线 RL 工程,接入代价低;
  2. 教师推理基础设施:必须已有 vLLM / SGLang / TensorRT-LLM 类高吞吐教师推理;上线前先做教师推理吞吐压测;
  3. token mask 监控:训中监控 entropy 与 KL 分布,避免 mask 比例塌到 0 或撑到 1;
  4. 零方差组比例作为先验指标:训练前先 estimate 零方差组占比,若超过 ~30% 才考虑上 OPD,节省 wall-clock;
  5. 可移植到 reasoning dataset 拼接:math+code+general reasoning 混合训练,SFT 数据可来自教师在错题上的重生成;
  6. 可与 DPO/SimPO 类偏好蒸馏混合:本文只示范"做错题+教师对答案"路径,偏好对蒸馏仍可叠加。

最小可跑命令(⚠️ 原文未提供开源代码,以下为按论文方法学重建):

# 伪环境示意
python -m train_rstg \
  --student qwen2.5-7b-instruct \
  --teacher qwen2.5-32b-instruct \
  --dataset mix(math,code) \
  --G 8 \
  --tau1 0.6 --tau2 0.3 \
  --lambda_opd 0.5 --lambda_sft 0.2 \
  --use_verifier math_verify,code_exec

需要 8×H100 / 4×H200 一类典型 RL 后训练栈。论文实际硬件与 batch 摘要未明确。


6. 与同方向工作的关系

  • GRPO 系列:DeepSeekMath/DeepSeek-R1 原始 GRPO、Dr. GRPO(去掉长度归一化偏差)、DAPO(动态采样)、GSPO(sequence-level policy)。RSTG 与 Dr. GRPO/DAPO 互补,专治"零方差"。
  • OPD/蒸馏:MiniLLM 的 on-policy KD、DistilWhisper 模式、ULD/ULDormouse。RSTG 把"对哪些样本/哪些 token 蒸馏"做成显式门控,把蒸馏做成 RL 的补丁信号。
  • SFT+RL 混合:Rejection Sampling Fine-Tuning (RFT)、STaR、Self-Rewarding。在这些体系里"RL 之前先做 SFT"几乎已成范式;RSTG 的差异是 SFT 与 RL 同时在线、按 group gate 切换。
  • RLAIF / Self-rewarding:通过自评代替人类反馈;RSTG 是"RLAIF + 显式门控"的实用主义路径,与纯自评路线互补。

7. 适合谁读

  • LLM 后训练 / RL 团队的算法工程师:手上已经在跑 GRPO/DAPO 系列的可直接搬;
  • 训练框架工程团队:要把"teacher 推理 + student 训练"双轨异步流水线常态化的人;
  • 数学/代码垂类模型团队:想给垂类小模型再压 3–4% 准确率的负责人;
  • 学术界研究组:关心 RLAIF/OPD 退化机理的理论研究者。

不适合:教师-学生架构关系不清晰的场景(如纯 self-rewarding without external teacher),或偏好对学习(DPO/RPO)为主路径的团队——门控思路可借鉴但 loss 形式不再直接适用。


原始链接:https://arxiv.org/abs/2608.00782(v1 提交 2026-08-01)· 不确定项:实际 student/teacher 模型对、τ1/τ2、λ 比例、与 2026 年前线 RL SOTA 横评、教师 wall-clock 成本、开源情况均按"原文未明确"标注。

工程落地与核查(Jay)

事实核查摘要

核查项 状态 说明
GRPO 零方差组信号丢失机制 ✅ 合理 与 DeepSeekMath 论文描述一致
零方差组定义(G条全对/全错→组优势=0) ✅ 符合 GRPO 原始公式 A = r - mean(r),全相同则 mean = 自身,优势 = 0
Token-level mask 公式(熵/KL 双条件) ✅ 数学上合理 H_stu > τ1 ∨ KL > τ2,stopgrad 设计防止教师梯度回传
SFT loss on teacher-correct trajectories ✅ 方法合理 标准 SFT loss,教师答案做绿目标
一句话结论数字 vs 关键实验数字不一致 ⚠️ 需核查 一句话:+3.02%;关键实验引摘要:+4.02%;差 1%,应以 PDF 原文表格为准
代码 +3.05% ✅ 两处一致 一句话和关键实验节均引为 +3.05%
教师加速用 vLLM/SGLang/TRITON-LLM ✅ 业界标准 均为 2026 年中可用方案
OPD 退化的三条根因 ✅ 逻辑自洽 诊断合理,支撑三件套设计

实际系统怎么用

RSTG 是一套 RL 后训练层的轻量补丁,不是独立系统。接入路径分为三个阶段:

阶段一:零方差诊断(接入前置评估)

# 先跑一轮不带 OPD 的 GRPO,统计零方差组比例
from grpo import GRPO
model = GRPO(student, verifier)
stats = model.count_zero_variance_groups(val_prompts)
print(f"零方差组比例: {stats.zero_rate:.1%}")
# 若 zero_rate < 10%:OPD 收益有限,优先优化 GRPO 基线
# 若 zero_rate > 30%:RSTG 补丁收益明显,值得接入

阶段二:teacher 推理流水线部署 - 教师模型须与 student 共享相同 tokenizer,以便做 token-level KL 对齐; - 教师推理吞吐是系统瓶颈:建议教师用量化版本(如 AWQ/CGPTQ),student 用 BF16 全精度; - 教师置信度 mean(teacher_confidence(y_i)) 的计算方式取决于模型是否输出 logit——需要检查教师模型是否输出归一化概率分布。

阶段三:混合训练循环(伪代码见 2.5 节)

工程坑位一览

  1. 教师推理是 P99 延迟杀手:每个训练 step 需要跑两路(student + teacher),如果教师是 student 的 10× 体量,单步时间约 ×1.8-2.0 倍。在 GPU 集群上这意味着训练吞吐下降接近一半;建议用 continuous batching 掩盖气泡。

  2. τ1/τ2 阈值需要 sweep:0.6/0.3 只是论文示例值,不同 teacher-student 组合表现差异可能很大。推荐做法:先用默认值跑几个 epoch,观察 mask 比例(mean(m_t)),若 <5% 说明阈值过高,若 >80% 说明阈值过低,逐步调优。

  3. SFT 与 GRPO 同时在线可能导致灾难性遗忘:SFT 直接注入了教师答案分布,如果 λ_sft 过高,student 可能过度依赖 SFT 路径而减少探索。建议从 λ_opd=0.3, λ_sft=0.1 小值开始,逐步放大。

  4. verifier 本身的 0/1 判断必须准确:GRPO 信号全靠 verifier 输出;如果 math verifier 有 5% 的误判率,这些误判样本的梯度会直接毒害训练。需要对 verifier 做人工抽检评估。

  5. 全对组 OPD 的梯度方向有争议:在全对组里,sample_weight = teacher_confidence,对教师置信度低的样本(可能是边缘正确答案)给更高的蒸馏权重——这实际上是鼓励 student 少依赖自己的边缘正确判断,而去模仿教师。方向上合理但效果未经大规模验证。

  6. 开源状态未知:论文截止 2026-08-01 v1 提交时尚未明确开源,工程团队接入前务必确认 arXiv 页面的 GitHub 链接或联系作者。

最小可跑验证路径

# 1. 确认 arXiv 开源代码(截稿时未标注)
open https://arxiv.org/abs/2608.00782

# 2. 搭建教师推理服务(以 SGLang 为例)
pip install sglang
python -m sglang.launch_server \
  --model-path qwen2.5-32b-instruct \
  --port 30000 --dtype half

# 3. 用小数据集跑一轮 GRPO baseline 测零方差率
python -m grpo.train \
  --student qwen2.5-7b-instruct \
  --dataset math subset 500 \
  --G 8 \
  --eval-only  # 仅测零方差率,不更新梯度

# 4. 若 zero_rate > 30%,接入 RSTG
python -m grpo_rstg.train \
  --student qwen2.5-7b-instruct \
  --teacher-url http://localhost:30000 \
  --dataset math subset 500 \
  --G 8 \
  --tau1 0.6 --tau2 0.3 \
  --lambda_opd 0.5 --lambda_sft 0.2 \
  --num-steps 1000

# 5. 监控 mask ratio 是否在合理区间
# tensorboard --logdir ./runs/rstg_experiment/

⚠️ 警告:上述路径依赖论文是否开源代码与预训练脚本。若未开源,需要自行实现 loss 计算逻辑(参考第 2 节公式),工程量约 2-3 人天。