ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

AReaL On-Policy 知识蒸馏与 KDRL 联合框架:原理、源码实现与配置实战

AReaL On-Policy 知识蒸馏与 KDRL 联合框架:原理、源码实现与配置实战 AReaL On-Policy 知识蒸馏与 KDRL 联合框架原理、源码实现与配置实战【免费下载链接】AReaLThe RL Bridge for LLM-based Agent Applications. Made Simple Flexible.项目地址: https://gitcode.com/GitHub_Trending/are/AReaL本指南深入解析 AReaL 中新增的On-Policy 知识蒸馏On-Policy Distillation与KDRLKD RL 联合训练框架先阐明 Forward KL 与 Reverse KL 两种蒸馏目标的差异及其与 exposure bias 的关系再结合areal/trainer/ppo/actor.py的源码逐步还原纯 KD 与联合损失的实际计算逻辑最后给出完整的teacher配置示例与可复现的命令行。读完本文你将掌握如何在 AReaL 中让一个小规模学生模型在自身采样轨迹上同时向教师学习 通过 RL 探索并能够自行调优rl_loss_weight/distill_loss_weight等关键超参。AReaL 此前主要支持 RL 后训练GRPO/PPO 等本文所介绍的实现为其新增了on-policy 知识蒸馏与KDRL 联合框架使学生可以在同一批 on-policy 轨迹上同时进行教师模仿与强化学习探索从而提升训练效率与稳定性。核心思想蒸馏目标的散度选择知识蒸馏的目标是让学生策略 $\pi_\theta$ 拟合一个更强的教师策略 $\pi_T$。蒸馏目标中采用的散度形式与采样分布会显著影响学生的最终表现与 exposure bias。监督微调Forward KL一种简单有效的方法是在教师生成的数据上最大化对数似然即 SFT。这等价于最小化 $\pi_T$ 与 $\pi_\theta$ 之间的 Forward KL$$\arg \min_{\theta} D_{KL}(\pi_T \parallel \pi_\theta) \arg \max_{\theta} \mathbb{E}_{q \sim Q,\ o \sim \pi_T(\cdot|q)} \left[ \log \pi_\theta(o|q) \right]$$SFT 高效且实现简单但存在一个固有缺陷在 off-policy 数据上训练会产生 exposure bias——训练时前缀来自教师推理时前缀由学生自回归生成。对长链路推理模型而言前缀分布偏移会被逐步放大问题尤为明显。On-Policy 蒸馏Reverse KL为缓解 exposure bias可以在学生自采样轨迹上训练这等价于最小化 Reverse KLRKL$$\arg \min_{\theta} D_{KL}(\pi_\theta \parallel \pi_T) \arg \max_{\theta} \mathbb{E}{q \sim Q,\ o \sim \pi\theta(\cdot|q)} \left[ \log \frac{\pi_T(o|q)}{\pi_\theta(o|q)} \right]$$两者差异的本质在于Forward KL 在教师分布下采样教师前缀、教师轨迹目标是让学生覆盖教师所有高概率区域Reverse KL 在学生分布下采样学生前缀、学生轨迹目标是让学生模式对齐到教师的高概率模式上与自回归推理时的实际分布一致。为什么 RKL 可以看作 REINFORCE最小化 RKL 可视为一种 REINFORCE奖励是教师与学生概率的对数比。采用 GRPO 框架时优化目标为$$J_{RKL}(\theta) \mathbb{E}{q,\ {o_i} \sim \pi{\theta_{old}}} \left[ \frac{1}{G} \sum_{i1}^{G} \frac{1}{|o_i|} \sum_{t1}^{|o_i|} \frac{\pi_\theta(o_{i,t})}{\pi_{\theta_{old}}(o_{i,t})} R_{i,t} \right]$$其中奖励为 $R_{i,t} \log \pi_T(o_{i,t}) - \log \pi_\theta(o_{i,t})$。这一设计会提升教师偏好 token 的概率并抑制教师认为不合理的 token。源码实现纯 KD 分支在 areal/trainer/ppo/actor.py 的grpo_loss_fn中纯 KD 场景rl_loss_weight 0的代码与上述公式一一对应if rl_loss_weight 0: rkl_reward teacher_logp - logprobs.detach() # R_{i,t} importance_weight torch.exp(logprobs - old_logp) # π_θ / π_θ_old rkl_weighted_term importance_weight * rkl_reward * loss_mask loss ( -distill_loss_weight * rkl_weighted_term.sum() / loss_mask.sum().clamp(min1) )实现要点以teacher_logp - logprobs作为奖励 $R_{i,t}$logprobs经detach()处理奖励不参与对当前策略的梯度通过torch.exp(logprobs - old_logp)计算重要性采样权重估计 RKL 梯度与文档中J_{RKL}的 GRPO 形式一致损失前冠以负号系数-distill_loss_weight将最大化对数比奖励的目标转换为最小化形式最终按loss_mask.sum()做 token 级归一化纯 KD 时无需 GAE/优势计算。KDRLGRPO 与 KD 联合联合目标Joint Loss在 KDRL 联合场景下AReaL 采用在 GRPO 目标上增加辅助 KL 项的方案。为保持与 GRPO 的 on-policy 特性一致辅助项使用 Reverse KL$$J_{KDRL}(\theta) J_{GRPO}(\theta) - \beta \cdot D_{KL}(\pi_\theta \parallel \pi_T)$$其中 $\beta$ 对应配置中的distill_loss_weight。由 RKL 目标本身是一个以 $\pi_\theta$ 为采样分布的无偏估计$\nabla_\theta J_{KDRL}(\theta)$ 是 $\nabla_\theta J_{GRPO}(\theta) \beta \cdot \nabla_\theta J_{RKL}(\theta)$ 的无偏估计。源码实现联合损失分支在联合损失场景rl_loss_weight 0中RKL 作为直接正则项出现。最小化logprobs - teacher_logp在学生分布 $\pi_\theta$ 采样下与最小化 $D_{KL}(\pi_\theta \parallel \pi_T)$ 等价。对应代码else: rkl_penalty_per_token (logprobs - teacher_logp) * loss_mask rkl_penalty rkl_penalty_per_token.sum() / loss_mask.sum().clamp(min1) loss rl_loss_weight * loss distill_loss_weight * rkl_penalty即loss rl_loss_weight * loss distill_loss_weight * rkl_penalty。注意此处logprobs是当前训练前向学生策略的 log 概率teacher_logp已在前置逻辑中被detach()冻结因此(logprobs - teacher_logp)的梯度只作用于学生构成一个不破坏 GRPO 主目标的可微正则项。两种模式速查场景rl_loss_weightdistill_loss_weight损失形式含义纯 KD0 0如0.005重要性采样加权 RKL 奖励学生只向教师学习不做 RL 探索KD RL 联合 0如1.0 0如0.005rl_loss_weight * loss distill_loss_weight * rkl_penalty同批轨迹上同时模仿教师与 RL 优化两个分支的默认权重在 areal/api/cli_args.py 的TeacherConfig中定义rl_loss_weight默认1.0distill_loss_weight默认0.005。实战如何配置教师模型teacher 配置项详解TeacherConfig位于 areal/api/cli_args.py支持两种教师引擎类型字段类型 / 默认值说明engine_typestr默认rolloutrollout使用推理引擎vLLM/SGLang对轨迹打教师 logptrain使用传统训练引擎教师路径rolloutInferenceEngineConfig推理引擎教师配置engine_typerollout时必填trainPPOActorConfig传统训练引擎教师配置engine_typetrain时必填本文示例即该路径pathstr教师模型路径设置后将覆盖共享 rollout 后端模型路径offloadbool默认false是否在训练步间卸载教师模型teacher.offload见 areal/trainer/rl_trainer.py 中的_should_offload_teacherrl_loss_weightfloat默认1.0RL 损失权重distill_loss_weightfloat默认0.005蒸馏损失权重需要注意的是teacher与多教师蒸馏配置mopd不可同时配置二者互斥cli_args.py的__post_init__会校验并抛错。完整示例配置训练引擎教师原文档给出的 teacher 配置片段如下完整可运行版本见 examples/distillation/gsm8k_grpo_distill_mode_trainEngine.yamlteacher: backend: fsdp:d1p1t4 rl_loss_weight: 1.0 distill_loss_weight: 0.005 experiment_name: ${experiment_name} trial_name: ${trial_name} path: Qwen/Qwen3-32B init_from_scratch: false disable_dropout: true dtype: ${actor.dtype} mb_spec: max_tokens_per_mb: 10240 optimizer: null scheduling_spec: ${actor.scheduling_spec}字段解读backend: fsdp:d1p1t4教师以 FSDP 并行1 数据并行 × 1 流水并行 × 4 张量并行加载。示例配置中教师为Qwen/Qwen2.5-14B-Instruct14B占用 4 卡学生为Qwen/Qwen3-0.6BFSDPd1p1t1单卡合共 5 卡 rollout 推理引擎rl_loss_weight: 1.0与distill_loss_weight: 0.005即上面联合损失中的rl_loss_weight与 $\beta$disable_dropout: true关闭 dropout保证教师打 logp 时行为确定、可复现dtype: ${actor.dtype}教师精度与学生保持一致示例中为bfloat16optimizer: null教师模型只做前向打分不构造优化器、不更新参数scheduling_spec: ${actor.scheduling_spec}复用 actor 的调度规格每个 worker 1 卡、32GB 内存、运行areal.infra.rpc.rpc_server。配套地学生actor侧的关键配置actor: backend: fsdp:d1p1t1 path: Qwen/Qwen3-0.6B dtype: bfloat16 gradient_checkpointing: true optimizer: type: adam lr: 1.70e-5 weight_decay: 0.017 lr_scheduler_type: constant gradient_clipping: 1.0 eps_clip: 0.4 reward_scaling: 10.0 reward_bias: -0.5 kl_ctl: 0.0 ppo_n_minibatches: 1 use_decoupled_loss: true其中kl_ctl: 0.0关闭了传统 GRPO 中对参考模型的 KL 惩罚蒸馏本身已提供了足够的正则约束reward_scaling: 10.0/reward_bias: -0.5用于缩放 GSM8K 的 0/1 任务奖励use_decoupled_loss: true启用解耦 PPO 损失。推理引擎教师模式rollout除训练引擎教师外AReaL 还支持将教师作为独立推理引擎vLLM运行见 examples/distillation/gsm8k_grpo_distill_mode_rolloutEngine.yamlteacher: engine_type: rollout path: Qwen/Qwen2.5-14B-Instruct rollout: backend: vllm:d1p1t2 scheduling_spec: - task_type: worker port_count: 2 gpu: 1 mem: 32 cmd: python3 -m areal.infra.rpc.rpc_server env_vars: {} rl_loss_weight: 1.0 distill_loss_weight: 5e-3该模式下教师由 vLLM 推理引擎承载vllm:d1p1t21 数据并行 × 2 张量并行负责在学生采样的同一批轨迹上计算teacher_logp。这种部署方式将教师打分与训练前向解耦显存占用更低、扩缩容更灵活。数据流教师 logp 如何进入训练从源码角度看教师 logp 的注入发生在 areal/trainer/rl_trainer.py 的训练循环中teacher_logps self.teacher.compute_logp(rollout_batch) for traj, logp in zip(rollout_batch, teacher_logps): traj[teacher_logp] logp traj[rl_loss_weight] self.config.teacher.rl_loss_weight traj[distill_loss_weight] self.config.teacher.distill_loss_weight即学生先通过 rollout 引擎采样轨迹然后教师引擎在同一批轨迹上compute_logp将teacher_logp与两个权重字段写入每条轨迹最终在grpo_loss_fn中按上文公式参与损失计算。整个过程共用同一份 on-policy 轨迹这正是向教师学习 通过 RL 探索可以在同一批数据上同时进行的原因。运行示例启动命令使用本地调度器启动训练命令来自原文档配置路径已核对存在python3 examples/math/gsm8k_rl.py --config examples/distillation/gsm8k_grpo_distill_mode_trainEngine.yaml scheduler.typelocal experiment_namegsm8k-grpo-distillation trial_nametrial0命令行参数说明--config examples/distillation/gsm8k_grpo_distill_mode_trainEngine.yaml指定训练配置也可改用..._rolloutEngine.yaml走推理引擎教师模式scheduler.typelocal使用本地调度器local调度器实现在 areal/infra/scheduler配置文件中scheduler.type: null时由 CLI 覆盖experiment_namegsm8k-grpo-distillation/trial_nametrial0实验与试次命名用于日志、checkpoint 与统计落盘cluster.fileroot默认/tmp/areal/experiments。其他关键训练配置两份示例配置中其他值得关注的设置rolloutvllm:d1p1t1推理后端max_concurrent_rollouts: 256、max_head_offpolicyness: 2、dump_to_file: truegconfign_samples: 4每组 4 条采样、max_new_tokens: 2048、temperature: 1.0greedy: false数据集train_dataset.path: openai/gsm8kbatch_size: 256max_length: 1024type: rl基础设施单机 8 卡n_nodes: 1、n_gpus_per_node: 8、NFS 名字解析、total_train_epochs: 10权重同步actor.weight_update_mode: xccl支持训练后向 rollout 引擎同步学生权重。结果下图为 Qwen2.5-14B-Instruct教师与 Qwen3-0.6B学生在 FSDP vLLM 条件下的 on-policy KD RL 奖励曲线原图见 docs/zh/algorithms/reward_curve.png。从源码结构与示例配置可以推断该实验采用了本文所述的联合损失rl_loss_weight: 1.0、distill_loss_weight: 5e-3即学生模型在 GSM8K 上一边以 0.6B 的规模模仿 14B 教师的推理行为一边通过 GRPO 探索正确答案最终在奖励曲线上呈现稳定上升。参考原文档docs/zh/algorithms/distillation.md实现源码areal/trainer/ppo/actor.pygrpo_loss_fn中 RKL/联合损失分支、areal/trainer/rl_trainer.py教师 logp 注入、areal/api/cli_args.pyTeacherConfig参数定义示例配置examples/distillation/gsm8k_grpo_distill_mode_trainEngine.yaml、examples/distillation/gsm8k_grpo_distill_mode_rolloutEngine.yaml论文Xu H, Zhu Q, Deng H, et al.KDRL: Post-training Reasoning LLMs via Unified Knowledge Distillation and Reinforcement Learning文中 RKL 与联合损失公式的理论出处【免费下载链接】AReaLThe RL Bridge for LLM-based Agent Applications. Made Simple Flexible.项目地址: https://gitcode.com/GitHub_Trending/are/AReaL创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表