尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

slime On-Policy Distillation 实战指南:用逆 KL 惩罚在任意 Advantage Estimator 上叠加教师蒸馏信号

slime On-Policy Distillation 实战指南:用逆 KL 惩罚在任意 Advantage Estimator 上叠加教师蒸馏信号 slime On-Policy Distillation 实战指南用逆 KL 惩罚在任意 Advantage Estimator 上叠加教师蒸馏信号【免费下载链接】slimeslime is an LLM post-training framework for RL Scaling.项目地址: https://gitcode.com/GitHub_Trending/slime12/slime导读本文基于 slime一个面向 RL Scaling 的 LLM 后训练框架中 on-policy distillation在策略蒸馏OPD的官方文档与源码系统讲解 OPD 的核心原理、参数配置与两种教师部署模式外部 SGLang 服务器 / Megatron 内嵌。你读完本文后将掌握为什么 OPD 以“学生自己采样的 token”而非教师生成的轨迹为学习对象、逆 KL 惩罚如何以蒙特卡洛形式叠加到 GRPO/PPO/REINFORCE 等任意 advantage estimator 上、两种教师模式各自适用的场景与完整可运行的配置脚本以及如何从 HuggingFace 权重转换教师 checkpoint 并跑通开箱示例。一、什么是 On-Policy Distillation学生沿着自己的轨迹学习传统知识蒸馏off-policy distillation通常用教师模型生成的数据或离线数据集训练学生学生接触的分布与自身当前策略并不一致。On-policy distillationOPD的关键差异在于训练数据来自学生当前策略current policy的在线采样教师并不生成轨迹只在学生访问到的每一个前缀prefix上为学生实际采样的同一个 next token给出评分。数据来源是学生a_t ~ π_θ(·|h_t)其中h_t是学生生成轨迹中 tokena_t之前的历史。教师只打分不生成教师模型固定frozen为每个学生采样出的 token 计算 log-probability。信号是稠密的 token 级信号沿学生自己的整条轨迹每一个位置都得到一个学习信号而不是只有句尾的稀疏奖励。在 slime 中这个信号以**逆 KL 的采样估计sampled reverse-KL penalty**作用于 advantage因此它可以与 GRPO、PPO、REINFORCE、GSPO 等任何 advantage estimator 自由组合当任务奖励为零时同一套机制退化为纯蒸馏pure distillation——学生只模仿教师的 token 级概率分布。从源码结构看该机制在 slime 中被设计为 advantage 之上的“加法项”而非独立的 estimator。slime/utils/arguments.py中--advantage-estimator参数的 help 文本明确写道on-policy distillation (OPD) is now orthogonal to the advantage estimator. Use --opd-kl-coef 0 to enable OPD on top of any estimator.arguments.py与文档表述完全一致。二、核心参数一览OPD 相关的命令行参数在 slime/utils/arguments.py 的add_on_policy_distillation_arguments中注册下表与官方文档保持一致参数说明--use-opd启用在策略蒸馏。使用 OPD 的必需开关actionstore_true默认 False。--opd-typeOPD 类型sglang或megatron二选一choices限定。启用--use-opd时必须显式指定。--opd-kl-coefOPD KL 惩罚系数typefloat默认 1.0。控制蒸馏信号相对于 RL advantage 的权重即公式中的λ_opd。--opd-teacher-load教师模型的 Megatron checkpoint 路径。--opd-typemegatron时必须设置--opd-typesglang时禁止设置。--opd-teacher-ckpt-step可选的教师模型 checkpoint 步数typeint默认 None。参数校验逻辑同样位于 arguments.py--use-opd但未指定--opd-type→ 报错要求从sglang/megatron中选择--opd-typemegatron但未提供--opd-teacher-load→ 报错要求教师 checkpoint 路径--opd-typemegatron且--opd-teacher-load指向的目录不存在、或缺少latest_checkpointed_iteration.txtMegatron checkpoint 的必要标记文件→ 报错--opd-typesglang却同时设置了--opd-teacher-load→ 报冲突错误未启用--use-opd但设置了--opd-teacher-load→ 报错提示需要添加--use-opd。这些校验保证了配置的“强约束性”两种教师模式互斥且路径必填项严格区分避免误配导致的隐性问题。三、原理逆 KL 的蒙特卡洛估计与 advantage 修正3.1 Token 级逆 KL 的定义设学生策略为π_θ教师策略为π_Th_t为学生生成轨迹中采样 tokena_t之前的历史。按 Thinking Machines Lab 给出的定义token 级逆 KL 为D_KL(π_θ(·|h_t) ‖ π_T(·|h_t)) E_{a_t ~ π_θ(·|h_t)} [ log π_θ(a_t|h_t) − log π_T(a_t|h_t) ]这里的顺序很重要KL 的第一个参数被评估的分布是学生分布期望同样对学生分布取值教师只评估学生实际采样的 token教师自身不参与轨迹生成。这与“让学生分布向教师分布靠近”的蒸馏直觉一致——逆 KL 惩罚的是学生在教师高概率区域之外的分配。3.2 蒙特卡洛采样贡献slime 并不遍历完整词表去精确计算该期望词表展开在训练成本上不可行而是对每个采样 token 使用单次蒙特卡洛贡献d̂_t log π_θ(a_t|h_t) − log π_T(a_t|h_t), a_t ~ π_θ(·|h_t)即“学生的采样 log-prob − 教师对同一 token 的 log-prob”。注意尽管 KL 的期望非负单个样本的d̂_t仍可能为负当教师概率高于学生概率时这是单样本估计的正常现象。3.3 Advantage 修正公式slime 将基础 advantage 修改为Â_t A_t − λ_opd · d̂_t其中A_t来自所配置的 estimatorGRPO / PPO / REINFORCE / GSPO 等纯蒸馏时任务奖励为零A_t即为零λ_opd即--opd-kl-coef策略损失使用修正后的Â_t因此OPD 项与 advantage estimator 的选择相互独立。3.4 源码级实现该公式的落地实现位于 slime/backends/megatron_utils/loss.py 的apply_opd_kl_to_advantagesreverse_kl student_log_probs[i] - teacher_log_probs[i] advantages[i] adv - args.opd_kl_coef * reverse_kl该函数在 advantage 计算流程中被调用同一文件第 809 行附近args.use_opd为真时触发并将reverse_kl存入rollout_data[opd_reverse_kl]用于日志上报第 1168-1170 行会将其聚合成opd_reverse_kl指标。apply_opd_kl_to_advantages的 docstring 明确说明其语义Computes reverse KL (student_logp - teacher_logp) and adds weighted penalty to advantages in-place. This is orthogonal to the base advantage estimator.并在 References 中指向 thinking-machines-lab/tinker-cookbook 的同名实现——这是对公式出处Thinking Machines Lab 博客的源码级印证。同时注意teacher_log_probs缺失时会抛出明确异常OPD with opd_type... requires teacher_log_probs这要求训练数据管线中必须携带教师 log-probs这正是两种教师模式要解决的核心问题。四、两种教师模式4.1 SGLang 模式--opd-type sglang教师跑在外部服务器适用场景教师与学生架构不同如 Qwen3-32B 教师 vs Qwen3-8B 学生或教师过大无法与训练模型同时装入显存。由于教师需要为学生的原始 token ID评分学生与教师必须使用兼容的 tokenizer 和词表。工作流程外部 SGLang 服务器加载并运行教师模型。在 rollout 阶段自定义 reward 函数slime.rollout.on_policy_distillation.reward_func将学生采样的 token ID 发送给教师服务器经--rm-url指定的/generate接口获取教师对这些相同 token 的 log-probability。自定义后处理函数slime.rollout.on_policy_distillation.post_process_rewards将教师 log-probs 裁剪到 response 范围并存入sample.teacher_log_probs。训练阶段slime 从基础 advantage 中减去按--opd-kl-coef缩放的采样 log-probability 差值。配置模板与官方文档一致--use-opd --opd-type sglang --opd-kl-coef 1.0 --custom-rm-path slime.rollout.on_policy_distillation.reward_func --custom-reward-post-process-path slime.rollout.on_policy_distillation.post_process_rewards --rm-url http://TEACHER_IP:TEACHER_PORT/generate实现细节reward_func与post_process_rewards均在 slime/rollout/on_policy_distillation.py 中实现reward_func是异步函数构造 payload 时直接传入sample.tokens即学生采样的 token ID并设置return_logprobTrue、logprob_start_len0、max_new_tokens0教师不生成新 token只做评分通过aiohttp异步 POST 到args.rm_url同时支持多模态输入若sample.multimodal_inputs含图片则通过encode_image_for_rollout_engine编码后随请求发送。post_process_rewards从 SGLang 返回的meta_info[input_token_logprobs]中提取教师逐 token log-prob[1:]跳过首个位置再按sample.response_length裁剪到 response 区间写回sample.teacher_log_probs最后对每个样本返回标量奖励0.0——纯蒸馏场景下学习信号完全来自 OPD KL 惩罚如果你有任务奖励可以在此处叠加。4.2 Megatron 模式--opd-type megatron教师内嵌进训练进程适用场景教师与学生/参考模型架构相同且能整体放入 GPU 显存与策略模型同时驻留。工作流程教师模型在初始化阶段作为额外的 Megatron 模型加载经--opd-teacher-load。在训练前向传播阶段教师模型为每个样本计算 log-probs复用 Megatron 的 forward 路径。KL 惩罚内联计算并应用到 advantages无需外部服务器与网络往返。配置模板--use-opd --opd-type megatron --opd-kl-coef 1.0 --opd-teacher-load /path/to/teacher_torch_dist注意教师 checkpoint 必须是 Megatron 格式torch_dist或torch。可以使用 tools/convert_hf_to_torch_dist.py 从 HuggingFace 格式转换。五、运行示例Qwen3-8B 学生 Qwen3-32B 教师官方在 examples/on_policy_distillation/README.md 中提供了完整示例示例脚本位于examples/on_policy_distillation/目录。5.1 SGLang 教师模式第 1 步下载模型与数据hf download Qwen/Qwen3-32B --local-dir /root/Qwen3-32B hf download Qwen/Qwen3-8B --local-dir /root/Qwen3-8B hf download --repo-type dataset zhuzilin/dapo-math-17k --local-dir /root/dapo-math-17k第 2 步转换学生模型为 Megatron 格式cd /root/slime source scripts/models/qwen3-8B.sh PYTHONPATH/root/Megatron-LM python tools/convert_hf_to_torch_dist.py \ ${MODEL_ARGS[]} \ --hf-checkpoint /root/Qwen3-8B \ --save /root/Qwen3-8B_torch_dist其中 scripts/models/qwen3-8B.sh 提供 Qwen3-8B 的模型参数hidden size、层数、注意力头等与环境变量convert_hf_to_torch_dist.py负责将 HuggingFace checkpoint 转为 Megatrontorch_dist格式。第 3 步运行bash examples/on_policy_distillation/run-qwen3-8B-opd.sh该脚本run-qwen3-8B-opd.sh的关键流程如下启动教师服务器在后台CUDA_VISIBLE_DEVICES7 python3 -m sglang.launch_server --model-path /root/Qwen3-32B --host 0.0.0.0 --port $TEACHER_PORT --tp 1 --chunked-prefill-size 4096 --mem-fraction-static 0.6其中TEACHER_IP127.0.0.1、TEACHER_PORT13141健康检查循环curl -sf http://$TEACHER_IP:$TEACHER_PORT/health_generate直到就绪再curl .../get_model_info确认提交 Ray 任务ray job submit将train.py作为任务提交传入以下参数组CKPT_ARGS--hf-checkpoint /root/Qwen3-8B用于初始化、--ref-load /root/Qwen3-8B_torch_dist参考模型、--load/--save /root/Qwen3-8B_slime/、--save-interval 20ROLLOUT_ARGS--prompt-data /root/dapo-math-17k/dapo-math-17k.jsonl、--apply-chat-template、--rollout-shuffle、--num-rollout 300、--rollout-batch-size 16、--n-samples-per-prompt 4、--rollout-max-response-len 16384、--rollout-temperature 1、--global-batch-size 64、--balance-dataRM_ARGS即 4.1 节中的三个参数--custom-rm-path、--custom-reward-post-process-path、--rm-urlGRPO_ARGS--advantage-estimator grpo基础 estimator 任意可选--use-opd --opd-type sglang --opd-kl-coef 1.0并搭配--use-kl-loss --kl-loss-coef 0.00 --kl-loss-type low_var_kl --entropy-coef 0.00关闭常规 KL/熵正则避免与 OPD 信号混叠PERF_ARGS--tensor-model-parallel-size 2 --sequence-parallel --pipeline-model-parallel-size 1、--recompute-granularity full --recompute-method uniform --recompute-num-layers 1、--use-dynamic-batch-size --max-tokens-per-gpu 16384SGLANG_ARGS--rollout-num-gpus-per-engine 1 --sglang-mem-fraction-static 0.4OPTIMIZER_ARGSAdam--lr 1e-6 --lr-decay-style constant --weight-decay 0.1 --adam-beta1 0.9 --adam-beta2 0.98并行规模--actor-num-nodes 1 --actor-num-gpus-per-node 2 --rollout-num-gpus 4收尾清理训练结束后pkill掉 sglang / ray 进程。5.2 Megatron 教师模式# 1. 将学生和教师模型都转换为 Megatron 格式 # 2. 运行 bash examples/on_policy_distillation/run-qwen3-8B-opd-megatron.shrun-qwen3-8B-opd-megatron.sh 与 SGLang 版的主要差异在GRPO_ARGS--advantage-estimator grpo --use-opd --opd-type megatron --opd-kl-coef 1.0 --opd-teacher-load /root/Qwen3-8B_torch_dist # 教师模型路径示例中与参考模型相同仅为演示注意该脚本头部注释明确强调这只是演示配置——示例用原模型当教师自蒸馏实际使用中应换成更强的教师模型如 Qwen3-32B并按任务调整--opd-kl-coef。教师 checkpoint 同样须先经tools/convert_hf_to_torch_dist.py转换README 中的转换示例使用scripts/models/qwen3-8B.sh的MODEL_ARGS可替换为你的教师模型配置。5.3 两种模式的切换与常见错误从 SGLang 模式切换为 Megatron 模式把--opd-type sglang改为--opd-type megatron并补上--opd-teacher-load同时可移除--custom-rm-path/--custom-reward-post-process-path/--rm-url--rm-type math或你的常规 reward 配置接管。参数配错的报错行为源码校验已在第二节说明--use-opd缺--opd-type、megatron缺--opd-teacher-load、sglang却带--opd-teacher-load系统都会抛出清晰异常提示方便快速定位。六、初步实验结果官方文档与 examples/on_policy_distillation/README.md 报告了同一组初步结果使用 Qwen3-8B-Base 模型在 OpenThoughts3-1.2M 数据集的一部分上做 SFT然后在剩余数据上用 Qwen3-32B 教师进行在策略蒸馏Math500 评测结果如下方案Pass1Qwen3-8B-Base SFT76%Qwen3-8B-Base SFT On-Policy Distillation94%这一对比直观展示了 OPD 的价值在不改变学生模型规模、仅以教师 token 级信号做在线蒸馏的情况下Math500 Pass1 从 76% 提升到 94%18 个百分点。需说明这是仓库文档报告的初步实验数据读者在复现时应以自己环境与数据集下的实测结果为准。七、FAQ 与最佳实践要点为什么有两种 OPD 模式sglang模式教师运行在独立 SGLang 服务器上适合教师与学生架构不同、或教师过大无法与策略模型同卡加载的场景megatron模式教师通过与参考模型相同的参数加载机制内嵌进 Megatron 训练进程要求教师与策略模型架构一致。教师 log-probs 从哪里来SGLang 模式rollout 阶段经自定义 reward 函数从外部服务器在线获取Megatron 模式训练前向阶段由内嵌教师模型计算。实践要点总结两种模式都要求学生与教师词表/tokenizer 兼容SGLang 模式按 token ID 评分尤其依赖这一点Megatron 模式教师必须与策略模型架构一致且显存要能同时容纳两者纯蒸馏时任务奖励为零学习信号完全来自--opd-kl-coef缩放的逆 KL 惩罚若需叠加任务奖励可在自定义 reward/后处理函数中返回真实标量奖励--opd-kl-coef的取值需要按任务调试过小蒸馏信号弱过大可能压制任务奖励信号示例中关闭了常规 KL 与熵正则--kl-loss-coef 0.00 --entropy-coef 0.00避免与 OPD 的 KL 惩罚相互干扰这一点在自定义配置时值得留意。相关资源官方文档docs/en/advanced/on-policy-distillation.md、docs/zh/advanced/on-policy-distillation.md示例与脚本examples/on_policy_distillation/、run-qwen3-8B-opd.sh、run-qwen3-8B-opd-megatron.sh核心实现slime/rollout/on_policy_distillation.py、slime/backends/megatron_utils/loss.py、slime/utils/arguments.py模型配置与转换工具scripts/models/qwen3-8B.sh、tools/convert_hf_to_torch_dist.py【免费下载链接】slimeslime is an LLM post-training framework for RL Scaling.项目地址: https://gitcode.com/GitHub_Trending/slime12/slime创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表