资讯动态

TRL 异步蒸馏完整配置指南:AsyncDistillationTrainer 从原理到调参

发布时间:2026/9/17 7:43:59 来源:尧图企业网站定制
TRL 异步蒸馏完整配置指南AsyncDistillationTrainer 从原理到调参【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl 从教师不能同时在场说起想象这样一个场景你的学生是 0.5B 的小模型教师是 15B 的大模型差了十几倍。同步蒸馏的套路是生成、教师前向、梯度更新在同一个进程里顺序执行意味着学生与教师必须挤在同一组显卡上——15B 权重 fp16 下光权重就要 30GB再叠上训练侧的激活值OOM 只是时间问题就算硬塞得下两套模型的前向也在抢同一块显存的带宽。TRL 给出的答案就是AsyncDistillationTrainer实现见 async_distillation_trainer.py。它做的仍是 on-policy 蒸馏——每一个训练样本都来自学生自己的采样——但教师从此不占用训练卡它作为一个独立的 vLLM 服务器存在只对 HTTP 打分请求做出响应。学生生成答案、教师远程复核、trainer 更新权重三条线各干各的互不阻塞。 一条样本的旅程从 Prompt 到梯度把整条链路拆成生成侧和训练侧两个区中间隔着一条进程边界生成侧rollout worker 子进程无 GPU ① 取 prompt ──▶ ② 学生起草 ──▶ ③ 教师复核 │学生 vLLM 采样生成 │teacher-forced 打分 ▼ rollout_buffer跨进程队列 ▼ 训练侧trainer 主进程 ④ 新鲜度检查 ──▶ ⑤ 分行 ──▶ ⑥ 打包 ──▶ ⑦ 前向 JSD 损失 ──▶ ⑧ 优化器步 ▲ ⑧ 每 weight_sync_steps 步 ──NCCL 推新权重──▶ 学生 vLLM ─┘闭环① 取 promptworker 子进程从数据集取下一行一行 消息列表外加可选的teacher_id。② 学生起草调学生 vLLM 服务器的/v1/completions采样生成一份答案。③ 教师复核把答案原样发回路由到的教师服务器打分——就像拿着标准答案逐字对照老师只给每个位置报 logprob自己不动笔写新字。产物是RolloutSampleprompt、答案、每个完成位置的 top-teacher_top_k候选 token 及其 logprob。①→③ 是一个在途任务多个任务按max_inflight_tasks并发。④ 新鲜度检查trainer 逐个拉取样本对比样本生成时的模型版本与当前版本落后超过max_staleness直接丢弃sample/dropped_stale_total计数。⑤ 分行规划器把样本分给各 DP rank贪心按 Σ Lᵢ² 平衡避免某张卡拖慢全体。⑥ 打包DataCollatorForRollout把一行的样本拼成一条长序列样本边界处重置position_ids。⑦ 前向compute_loss按 256 个 token 分块投影lm_head峰值 logits 内存因此只是256 × vocab_size计算广义 JSD。⑧ 优化器步gradient_accumulation_steps个 micro-batch 凑成一步每weight_sync_steps步新权重经 NCCL 流进学生 vLLM 服务器生成侧因此能跟上——这就是回环箭头。⚖️ 损失与 β 参数选择支撑集为什么会被收窄先看约束。教师住在另一台机器上完整词表根本过不了 HTTP。每个完成位置线上能拿到的只有教师的 top-teacher_top_k候选 logprob、实际实现 token 的 logprobvLLM 保证它即使不在 top-k 里也会返回、以及可选的尾部桶add_tail_bucketTrue把剩余概率质量收进一个元素防止候选太少时散度看起来恒等于零。学生那一侧是精确的——它就是要训练的模型完整 logits 本地就有。在这个约束下beta决定了散度在哪些候选上计算beta0.0前向 KL期望以教师分布加权只关心教师把概率放在哪而教师的 top-k 切片恰好就是这份信息——所以保留全部teacher_top_k宽度支撑。beta≠0.0混合项里出现了以学生分布加权的分量重心落在学生自己采到的 token 上。这个 token 未必在教师 top-k 内而线上协议能保证拿到教师 logprob 的只有两个身份教师 top-1 与答案实际 token。于是_narrow_top1_actual_support把支撑收窄到 2 个候选两者相同则去重。beta1.0反向 KL纯学生加权教师 top-1 不贡献任何项支撑进一步缩到 1 个——只剩实际 token。一句话实操建议默认beta0.0适合全面继承教师分布若目标是让学生模仿教师的主行为比如融合 RL 专家MOPD 论文第三阶段用的就是反向 KL显式写beta1.0。注意中间值下teacher_top_k对支撑宽度已经无效它只影响教师报告的是哪个 top-1。config AsyncDistillationConfig( beta1.0, # 反向 KLmode-seeking teacher_top_k16, # 线上每位置的候选数 teacher_server_urls{math: http://localhost:8001}, )️ 三终端部署教师、学生、训练各占一张卡踩坑提醒当前 vLLM 与 transformers 的依赖约束互相冲突装反了会直接 import 失败。正确顺序是先装 vLLM再裸装transformers跳过依赖解析pip install vllm0.22.0 pip install transformers5.2.0 --no-deps另外分布式训练只支持 FSDP2DeepSpeed ZeRO 不在支持列表里。终端 1——教师服务器它的角色最轻只接打分请求不生成新文本、永不更新所以什么 dev 开关都不需要只要两个打分精确性flagCUDA_VISIBLE_DEVICES0 vllm serve Qwen/Qwen2.5-1.5B-Instruct \ --port 8001 \ --logprobs-mode processed_logprobs \ --max-logprobs -1--logprobs-mode processed_logprobs让teacher_temperature真正作用于返回的 logprob漏掉它该参数只影响学生侧--max-logprobs -1解除 vLLM 默认的 20 上限teacher_top_k想超过 20 必须开它。终端 2——学生的 vLLM 服务器它既负责起草又要接收 trainer 推来的新权重所以必须进 dev 模式并打开 NCCL 传输通道CUDA_VISIBLE_DEVICES1 VLLM_SERVER_DEV_MODE1 vllm serve Qwen/Qwen2.5-0.5B-Instruct \ --port 8000 \ --weight-transfer-config {backend:nccl}终端 3——训练进程脚本很短完整可跑版本在 examples/async_distillation_math/async_distillation_math.pytrainer AsyncDistillationTrainer( modelQwen/Qwen2.5-0.5B-Instruct, argsAsyncDistillationConfig( teacher_server_urls{default: http://localhost:8001}, vllm_server_base_urlhttp://localhost:8000}, train_datasetdataset, ) trainer.train()CUDA_VISIBLE_DEVICES2 accelerate launch examples/async_distillation_math/async_distillation_math.py新手最容易踩的 8 个参数参数默认值为什么容易踩teacher_server_urls{default: http://localhost:8001}端口要和终端 1 对上换成多个条目就进入 MOPD 路由模式beta0.0必须在[0, 1]越界直接ValueError取 0 和取非 0 时支撑宽度完全不同见上一节teacher_top_k88 只是冒烟测试级别正式训练建议 16–6420 要求教师带--max-logprobs -1teacher_temperature1.0同时作用于教师侧服务端算 logprob与学生 logits若教师没开processed_logprobs它对教师侧静默失效token_budgetNone取学生服务器max_model_len样本超出预算就进不了任何行被丢弃并计入batch/dropped_oversize_total显存紧张时它是第一调节杠杆max_staleness4样本可落后当前策略的权重版本数太小队列常空太大 off-policy 味太重weight_sync_steps1每 N 个优化器步推一次权重给学生 vLLM调大省同步开销但生成用更旧策略dtypefloat32默认 fp32 而非 bf16为对齐训练-推理精度度量要端到端一致学生 vLLM 也要用相同--dtype起服务其余字段采样参数、超时、心跳、日志开关见 async_distillation_config.py。再留个心眼learning_rate默认1e-6不是5e-5、logging_steps默认1、gradient_checkpointing默认True、bf16默认开——带着旧习惯配参数会翻车。 MOPD 多教师路由各管一摊不搞集成teacher_server_urls里放多个条目后路由规则只有一条每个样本由数据里的teacher_id列挑中唯一一位教师打分——数学 prompt 走数学教师代码 prompt 走代码教师。没有跨教师平均没有 ensemble是分诊而不是会诊。teacher_id缺失或不在映射里时_resolve_teacher_server_url直接抛ValueError不会悄悄回退到别的教师。有一条硬约束教师必须和学生共用 tokenizer。答案以原始 token id 发给教师教师回传的候选 id 会在compute_loss里直接索引学生自己的词表——词表对不上就是在给错误的 token 做训练而且只要教师的词表不比学生小这个错误是完全静默的。同家族组合Qwen2.5 学生 Qwen2.5 教师 Qwen2.5-Coder 教师天然满足参考 async_distillation_mopd.py 的双教师配置。最后注意范围MOPD 只接融合阶段——各领域专家必须已单独训练好并通过 HTTP 服务这个 trainer 不负责把它们练出来。 指标瓶颈定位哪根柱子在拖后腿诊断顺序固定为先问后答的三步第一问谁在等谁看镜像指标对perf/rollout_wait_s与rollout/backpressure_s——训练侧因队列空停下的总时长生成侧因队列满被压回的总时长两者不会同时很高。第二问行装满了吗看batch/row_fill_frac。1 万 token 的样本塞 3.2 万 token 的预算3 个放得下、4 个永远放不下打包器经常只能塞 2 个——填充率偏低时token_budget是第一调节杠杆顺带查batch/dropped_oversize_total有没有超预算丢弃。第三问学生在学习还是收缩看jsd与entropy的配对走势jsd降、entropy稳是正常收敛jsd降的同时entropy塌方说明学生只在高置信 token 上收缩分布没学到东西。指标一句话含义异常时接着看sample/rollout_queue_size队列里躺着多少份已打分样本与 wait/backpressure 配对判断瓶颈在哪侧sample/time_in_queue_s单个样本从打分完到进训练等了多久off-policy 的秒数部分sample/staleness_mean、sample/dropped_stale_totalperf/rollout_wait_s训练因队列空而停摆的累计时长rollout/score_s教师慢、rollout/generated_tok_s生成吞吐、rollout/inflightrollout/backpressure_s生成因队列满被压回的累计时长sample/staleness_mean是否在队列里变老、batch/row_tokens_meanrollout/vllm_retry_total对两台 vLLM 的 HTTP 重试次数持续上涨说明有台服务器在退化rollout/duration_sbatch/row_fill_frac行的实际 token 数占token_budget的比例低即预算没吃满batch/dropped_oversize_totalbatch/row_imbalance各行 Σ Lᵢ² 的最大/均值1.0 完美偏高说明有 rank 在拖 all-reducebatch/row_tokens_maxjsd/entropy损失本体 / 学生自身预测熵一降一稳收敛一降一塌收缩sample/staleness_mean、completions/clipped_ratio性能侧只需盯两个perf/step_s一步优化器步的墙钟时间含一切等待与perf/fwd_bwd_s其中纯前反向部分。MFU 的两个口径_fwd_bwd与_wall_clock差只在分母——前者回答有数据时训练侧多高效后者回答分到的算力有多少真变成了训练后者远小于前者就去找生成侧。MOPD 下还有teacher_jsd/id、teacher_token_frac/id按教师拆分路由偏斜在混合jsd里是看不出来的。下一次跑训练别盯着 loss 发呆先翻队列和背压定位瓶颈在哪一侧再看填充率确认预算匹配样本长度最后用jsd和entropy的配对走势确认学生真在学习。这套诊断路径全在上面的表里。【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价