案例

On-Policy Self-Distillation

分类:On-Policy Self-Distillation;训练 ¥41.81;评测 ¥83.61

代码出处与消耗

  • 完整代码来自 KMnO4-zx/llm-agent-rl-lab/04-opsd
  • 100-step 正式训练的 PyTRIO 实测花销为 ¥41.81;Base Model 与 Step 25、50、75、100 共 5 次完整 AIME25 评测实际采样 18.05M tokens,花销 ¥83.61;合计 ¥125.42
  • 100 个训练 step 的累计 step 时间约为 2 小时 6 分钟,不包含数据准备和完整评测。训练曲线见 SwanLab 实验记录
  • 以上都是该次参考运行的实测值,不是固定报价。由于该案例包含数据、同步/异步训练和评测等多个文件,本文展示关键代码与完整逻辑;可运行的全部文件请以上述源仓库为准。

OPSD 100-step 正式训练的 PyTRIO 消耗

OPSD Base Model 与四个 checkpoint 的 AIME25 评测消耗

介绍

OPSD 的全称是 On-Policy Self-Distillation。它让同一个初始模型同时扮演 Student 和 Teacher:

  • Student 只看题目,按当前策略生成自己的推理轨迹;
  • Teacher 额外看到参考解答,但不生成另一条答案;
  • Teacher 沿着 Student 已经生成的同一条 completion,计算每个 token 的 logprob;
  • Student 根据两者的逐 token 差异更新。

可以把它理解成:

Student:闭卷解题
Teacher:拿着参考解答,沿着 Student 的原始答案逐 token 判卷

这里最重要的边界是:

Teacher 不会重新采样一条“标准答案”。整条训练轨迹只由 Student 生成一次,Teacher 只对这条 Student 轨迹执行 logprob forward。

因此训练仍发生在当前 Student 真正访问到的状态上,保持 on-policy。

OPSD 的 Student、privileged Teacher 与逐 token 学习目标

OPSD 与其他方法的区别

方法训练轨迹监督信号Teacher在 Student 自己的状态上学习
SFT固定专家轨迹token-level CE不需要在线 Teacher
GRPO当前策略 rolloutsequence-level rewardReward / Verifier
普通 OPD当前 Student rolloutTeacher token logprob独立 Teacher 模型
OPSD当前 Student rolloutprivileged Teacher token logprob同一个初始模型,不同 prompt

OPSD 同时保留了 on-policy 轨迹、逐 token 稠密反馈和 self-distillation,不需要部署一个更大的 Teacher。

本案例的目标函数

OPSD 论文讨论了 full-vocabulary logit distillation 和 sampled-token distillation。这个 PyTRIO 案例实现的是第二种:只比较 Student 实际采样 token 在 Student 与 Teacher 下的 logprob。

对于 Student 在位置 t 采样的 token \hat{y}_t

reverse_klt=logpS(y^tx,y^<t)logpT(y^tx,y,y^<t)\operatorname{reverse\_kl}_t = \log p_S(\hat{y}_t \mid x, \hat{y}_{<t}) - \log p_T(\hat{y}_t \mid x, y^*, \hat{y}_{<t})

At=βreverse_klt=β(logpTlogpS)A_t = -\beta \cdot \operatorname{reverse\_kl}_t = \beta(\log p_T - \log p_S)

如果 Teacher 比 Student 更认可这个 token,advantage 为正;如果 Teacher 更不认可,它会被压低。

这是一种 sampled-token reverse-KL 训练实现,不等同于论文主实验的 full-vocabulary JSD。阅读结果时应保留这条实现边界。

实验配置

项目本案例配置
Base ModelQwen/Qwen3.5-4B
StudentLoRA rank 64,训练 attention + MLP
Teacher固定的 step-0 Qwen/Qwen3.5-4B
训练数据siyanzhao/Openthoughts_math_30k_opsd
数据量29,434 对 problem + solution
训练区间100 steps
每个 step32 道题
每道题1 条 Student completion
最大 completion1,024 tokens
最大远程并发32
Student / Teacher thinking均关闭
Samplingtemperature 1.1 / top-p 0.95 / top-k 20
KL coefficient1.0
Learning rate5e-6
Sampler refresh每 step 刷新
Lossimportance_sampling
Checkpoint每 25 step 保存 state + sampler weights

代码结构

OPSD 由数据准备、训练、评测和分析文件共同组成:

04-opsd/
├── 00-datasets.py       # 下载并校验 OPSD / AIME25 数据
├── 00-eval-aime25.py    # Base Model 与 checkpoint 统一评测
├── 01-opsd-async.py     # 推荐的异步 OPSD 训练实现
├── 01-opsd-sync.py      # 便于逐步阅读的同步实现
└── analysis.py          # 汇总 AIME25 checkpoint 结果

完整实现请查看 源代码目录。下面以异步版为主,拆解一次真实训练 step。

环境与数据

本地只需要可联网的 CPU 环境。Student sampling、Teacher logprob、LoRA 前向反向和参数更新由 PyTRIO 远端执行。

git clone https://github.com/KMnO4-zx/llm-agent-rl-lab.git
cd llm-agent-rl-lab

uv sync
trio login
swanlab login

下载固定 revision 的 siyanzhao/Openthoughts_math_30k_opsdyentinglin/aime_2025

uv run python 04-opsd/00-datasets.py

脚本会保存并校验:

04-opsd/datasets/
├── openthoughts_math_30k_opsd/   # 29,434 条训练数据
└── aime_2025/                     # 30 道评测题

每条训练数据至少包含:

problem   # Student 和 Teacher 都能看到
solution  # 只有 Teacher 能看到

solution 不是 SFT label。它只作为 privileged information 改变 Teacher 的条件分布;Student 的训练 target 仍然来自 Student 自己的 completion。

核心逻辑

一次训练 step 可以概括为:

for step in range(total_steps):
    student_sampler = refresh_latest_student_weights()

    rollouts = await asyncio.gather(
        *[
            student_sample_then_teacher_score(problem)
            for problem in batch
        ]
    )

    datums = build_importance_sampling_datums(rollouts)
    await forward_backward(datums)
    await optim_step()

下面展开其中最关键的对齐关系。

1. 同一个模型使用两份不同 Prompt

Student 只看到问题:

def build_student_prompt_ids(tokenizer, problem, enable_thinking):
    user_message = (
        f"Problem: {problem.strip()}\n\n"
        "Please reason step by step, and put your final answer within \\boxed{}."
    )
    return render_chat_prompt(tokenizer, user_message, enable_thinking)

Teacher 额外看到参考解答:

def build_teacher_prompt_ids(tokenizer, problem, solution, enable_thinking):
    user_message = (
        f"Problem: {problem.strip()}\n\n"
        "Here is a reference solution to this problem:\n"
        "=== Reference Solution Begin ===\n"
        f"{solution.strip()}\n"
        "=== Reference Solution End ===\n\n\n"
        f"{TEACHER_TRANSITION}\n\n"
        f"{STUDENT_INSTRUCTION}"
    )
    return render_chat_prompt(tokenizer, user_message, enable_thinking)

两者使用相同 tokenizer,也来自同一个 base model。Teacher 的额外能力来自参考解答上下文,而不是更多参数。

同一个模型使用 Student Prompt 与 privileged Teacher Prompt

2. 创建可训练 Student 和固定 Teacher

01-opsd-async.py 创建一个 LoRA TrainingClient 和一个不带 model_path 的固定 SamplingClient:

service_client = trio.ServiceClient()

training_client = await service_client.create_lora_training_client_async(
    base_model=args.base_model,
    rank=args.lora_rank,
    seed=args.seed,
    train_attn=True,
    train_mlp=True,
    train_unembed=args.train_unembed,
)

teacher_client = await service_client.create_sampling_client_async(
    base_model=args.base_model,
)

每个 step 只更新 Student LoRA。Teacher 保持 step-0 base policy,没有 optimizer,也不会随 Student 一起变化。

3. Student 先生成当前策略轨迹

每个 step 开始时刷新 Student sampler:

student_sampler = (
    await training_client.save_weights_and_get_sampling_client_async()
)

然后 Student 只看 problem 采样:

sample_result = await student_sampler.sample_async(
    prompt=trio.ModelInput.from_ints(student_prompt_ids),
    num_samples=args.group_size,
    sampling_params=sampling_params,
    return_text=False,
)

参考配置中 group_size=1sampler_refresh_steps=1,因此每道题采样一条 completion,并且下一 step 总是使用上一次 optimizer update 后的最新 Student 权重。

4. Teacher 只计算同一条 completion 的 Logprob

Student 完成采样后,Teacher 收到的是:

teacher_prompt_ids + student_completion_ids

Teacher 调用 compute_logprobs_async(),而不是 sample_async()

all_ids = teacher_prompt_ids + completion_ids
all_logprobs = await teacher_client.compute_logprobs_async(
    trio.ModelInput.from_ints(all_ids)
)
teacher_logprobs = all_logprobs[len(teacher_prompt_ids):]

以下三者必须严格等长:

Student completion tokens
Student rollout logprobs
Teacher completion logprobs

源码会检查长度和 None,不会悄悄截断。只要 token 对不上,逐 token reverse KL 就没有意义。

5. 把 Logprob 差变成逐 Token Advantage

对每条有效 Student completion:

student_lps = [float(value) for value in sequence.logprobs]
reverse_kl = np.asarray(student_lps) - np.asarray(teacher_lps)
advantages = -args.kl_penalty_coef * reverse_kl

GRPO 的 advantage 来自同题多条完整轨迹的最终 reward 差;OPSD 的 advantage 来自 Teacher 与 Student 对每个 Student token 的 logprob 差。因此即使最终答案错误,只要两者对中间 token 的偏好不同,这条轨迹仍可能提供训练信号。

6. Prompt 只作上下文,只训练 Completion

build_opd_datum() 对所有字段做相同的自回归右移,并把 prompt 区间填零:

prompt_loss_len = len(student_prompt_ids) - 1
input_ids = student_prompt_ids + completion_ids[:-1]

target_ids = [0] * prompt_loss_len + completion_ids
padded_logprobs = [0.0] * prompt_loss_len + old_logprobs
padded_advantages = [0.0] * prompt_loss_len + advantages.tolist()

datum = trio.Datum(
    model_input=trio.ModelInput.from_ints(input_ids),
    loss_fn_inputs={
        "target_tokens": np.asarray(target_ids, dtype=np.int64),
        "logprobs": np.asarray(padded_logprobs, dtype=np.float32),
        "advantages": np.asarray(padded_advantages, dtype=np.float32),
    },
)

prompt token 只提供上下文,不参与策略优化。真正进入 importance_sampling loss 的只有 Student 自己生成的 completion token。

7. 在题目之间并发,在单题内部保持顺序

一道题的 Student sample 必须先完成,Teacher 才能对同一条 completion 打分;不同题目之间则可以并发:

rollouts = await asyncio.gather(
    *(rollout_and_track(row) for row in batch)
)

异步版使用 asyncio.Semaphore(32) 统一限制 Student sampling 和 Teacher scoring 的远程并发数。当前 batch 的所有 rollout 都来自同一个 Student checkpoint,全部完成后才做 optimizer update,因此并发不会破坏 on-policy 边界。

8. 更新 Student

当前 step 的全部 Datum 展平后,提交一次训练更新:

fwd_bwd_future = await training_client.forward_backward_async(
    datums,
    loss_fn="importance_sampling",
)
optim_future = await training_client.optim_step_async(adam)

fwd_bwd_result = await fwd_bwd_future
await optim_future

PyTRIO 负责远端 forward、backward、LoRA optimizer 和 checkpoint;本地代码控制 Student / Teacher prompt、rollout、advantage、batch 边界和更新时机。

运行训练

执行 100-step 异步 OPSD:

uv run python 04-opsd/01-opsd-async.py \
    --steps 100 \
    --batch-size 32 \
    --group-size 1 \
    --max-tokens 1024 \
    --sample-size 0 \
    --save-every-steps 25 \
    --max-concurrency 32 \
    --swanlab-mode online

每 25 step 会保存:

*-state            # 包含优化器状态,用于断点续训
*-sampler_weights  # 用于采样和 AIME25 评测

参考运行前 100 step 的部分指标为:

trainer/loss_mean: 0.0577 → 0.0425
reverse_kl_mean:   0.0644 → 0.0471
reverse_kl_std:    0.4378 → 0.3687

OPSD 训练过程中的 loss、reverse KL、advantage 与耗时曲线

这些变化说明 Student 正在靠近 privileged Teacher,但不能单独证明数学能力提高,最终仍要看独立 benchmark。

AIME25 评测

评测 Base Model:

uv run python 04-opsd/00-eval-aime25.py \
    --val-n 12 \
    --max-tokens 38912 \
    --temperature 1.0 \
    --enable-thinking false \
    --output 04-opsd/eval-results/aime25-base.jsonl

评测 Step 100 checkpoint:

uv run python 04-opsd/00-eval-aime25.py \
    --val-n 12 \
    --max-tokens 38912 \
    --temperature 1.0 \
    --enable-thinking false \
    --model-path trio://<your_sampler_weights_path> \
    --output 04-opsd/eval-results/aime25-sampler-steps100.jsonl

每个模型状态包含 30 道题 × 每题 12 次采样 = 360 条 completion。max_tokens=38,912 是单条生成上限,不等于实际消耗;参考运行每个模型状态实际使用约 3.40M~3.86M sample tokens。

参考结果:

OPSD Base Model 与 Step 25、50、75、100 的 AIME25 结果

模型 / CheckpointAverage@12Pass@12正确 generations至少答对一次的题目
Qwen3.5-4B Base51.67%80.00%186 / 36024 / 30
Step 2551.11%73.33%184 / 36022 / 30
Step 5050.28%86.67%181 / 36026 / 30
Step 7548.61%76.67%175 / 36023 / 30
Step 10052.78%86.67%190 / 36026 / 30

Step 100 相比 Base Model 的 Average@12 提升 1.11 个百分点,Pass@12 提升 6.67 个百分点。但中间 checkpoint 并不单调,且这里只评测了 30 道题、一次训练、没有多个随机种子,因此更适合说明训练链路跑通并得到小幅正向变化,不能宣称稳定复现了论文收益。

训练时重点观察

指标作用
trainer/loss_meansampled-token 训练目标是否下降
opd/reverse_kl_meanStudent 与 privileged Teacher 的平均差异
opd/reverse_kl_stdtoken-level 差异是否集中
opd/advantage_meanTeacher 信号的平均方向
data/completion_tokens_total当前 step 的 Student completion token 数
time/step_elapsed_time包含采样、Teacher forward 和训练的完整 step 时间

Loss 和 reverse KL 下降只说明 Student 正在靠近 Teacher。Teacher 是否真正利用参考解答给出可靠监督,以及独立 benchmark 是否提升,仍需要单独验证。

完整实验图、费用截图和更多讨论请阅读源仓库的 OPSD README

这篇文档对你有帮助吗?

本页目录