
TRL AsyncDistillationTrainer 实战手册三终端部署、关键参数取舍与训练观测【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl这份手册面向正在做大模型后训练的工程师讲透 TRL 实验模块里的AsyncDistillationTrainer学生模型在自己的 vLLM 服务器上生成 on-policy 样本即由当前策略自己生成的完成结果教师通过 HTTP 打分生成与梯度更新并发进行。读完你能够独立在 3 块 GPU 上跑通一次异步蒸馏会判断beta、teacher_top_k、max_staleness等关键取舍并看懂rollout/、sample/、perf/指标去定位生成侧或训练侧的瓶颈。先跑起来装依赖、起三终端、写脚本先把依赖安装顺序搞对硬性前提是vllm0.22.0与transformers5.2.0。这两者目前存在冲突的依赖约束必须装 vLLM 之后再用--no-deps强制装 transformers否则 pip 会互相降级pip install vllm0.22.0 pip install transformers5.2.0 --no-deps分布式训练只支持 FSDP2DeepSpeed ZeRO 不可用。三终端如何分工一个教师服务器、一个学生推理服务器、训练进程三者必须分在三块不同的 GPU 上。两个服务器都是普通的vllm serve但启动参数完全不同——原因在下一章展开# GPU 0教师只做打分权重永不更新 CUDA_VISIBLE_DEVICES0 vllm serve Qwen/Qwen2.5-1.5B-Instruct \ --port 8001 --logprobs-mode processed_logprobs --max-logprobs -1 # GPU 1学生 vLLM生成 NCCL 权重传输 CUDA_VISIBLE_DEVICES1 VLLM_SERVER_DEV_MODE1 vllm serve Qwen/Qwen2.5-0.5B-Instruct \ --port 8000 --weight-transfer-config {backend:nccl} # GPU 2训练 CUDA_VISIBLE_DEVICES2 accelerate launch train_async_distillation.py教师那两个 flag 都有具体用途--logprobs-mode processed_logprobs让teacher_temperature真正作用于教师返回的 logprobs否则教师静默报告原始 logprobs你的温度设置只影响学生侧--max-logprobs -1解除 vLLM 每 token 默认 20 个 logprob 的上限teacher_top_k超过 20 时必须启动它。学生侧则需要 dev 模式加 NCCL 权重传输后端trainer 才能把更新后的权重推进去。最小训练脚本from datasets import load_dataset from trl.experimental.async_distillation import AsyncDistillationTrainer dataset load_dataset(trl-lib/DeepMath-103K, splittrain) trainer AsyncDistillationTrainer( modelQwen/Qwen2.5-0.5B-Instruct, train_datasetdataset, ) trainer.train()缺省配置下教师指向http://localhost:8001、学生指向http://localhost:8000学习率默认1e-6不是TrainingArguments的5e-5。更贴近实战的单教师示例在examples/async_distillation_math/async_distillation_math.pyGSM8K 数据、max_steps100、trackio 上报照抄改路径即可。它是怎么运转的四个角色各干什么与同步版DistillationTrainer教师、学生、优化器在同一进程里顺序执行不同异步版把过程拆成四个角色rollout workerspawn出的子进程显式清空CUDA_VISIBLE_DEVICES调学生 vLLM 的/v1/completions采样出完成结果再调教师 vLLM 的/v1/completions带prompt_logprobs做 teacher-forced 打分——教师只对学生自己的完成结果逐位置报 logprob不生成任何新 token。一次 rollout 一次生成 一次打分恰好产出一个训练样本。蒸馏没有 GRPO 那种组内基线所以 prompt 不会被重复采样。mp.Queue 缓冲队列容量queue_maxsize默认 1024worker 把带教师稀疏分布的打分样本推进来。训练循环主进程每次拉一个样本staleness 超过max_staleness就丢弃计入sample/dropped_stale_total拉够后规划成 rows——每个 DP rank 一行行内样本拼成单条序列且position_ids逐样本重置按 Σ Lᵢ² 贪心分桶平衡gradient_accumulation_steps个 micro-batch 构成一次优化器步损失是广义 JSD。权重同步每weight_sync_steps步默认 1把更新后的学生权重经 NCCL 推给学生 vLLM 服务器。三个容易踩坑的点教师永远不在训练机上加载。既然 worker 是禁 GPU 的子进程打分只能走 HTTP这也是教师能放在完全不同硬件上、甚至是用更大学员模型的原因。线上只传教师分布的稀疏 top-k 切片外加 vLLM 总会报告的 realized token 和尾桶完整词表从不过网学生侧 logits 是本地全量、精确计算的。生成永远领先于训练。从队列拉出的样本可能由落后若干个权重版本的策略产生max_staleness默认 4是容忍上限超过即丢。检查点有个特殊机制基类 Trainer 的 skip-and-replay 循环不适用于实时队列ignore_data_skip固定为True。每个检查点额外写rollout_state.json记录的是第一个还没被训练的 prompt的位置而不是生成器位置——worker 最多领先一个队列深度从生成位置恢复会跳过已生成未训练的 prompt流式 IterableDataset 无法重新定位恢复时 worker 从 prompt 0 重新开始。关键参数怎么定beta、teacher_top_k 与异步预算这一章只讲真正需要做决策的点其余参数保持默认即可完整清单见trl/experimental/async_distillation/async_distillation_config.py里的AsyncDistillationConfig。beta 取 0 还是 1损失是学生与教师逐位置 token 分布的广义 JSDbeta在两者之间插值beta散度行为实际参与计算的支撑集0.0默认前向 KLmean-seeking完整teacher_top_k宽的教师支撑 尾桶1.0反向 KLmode-seeking仅两个候选教师 top-1 与完成结果的实际 token中间值插值—走收窄路径同1.0为什么支撑集要变线上协议只能保证教师 logprob 对教师 top-1和实际 token这两个身份一定可用beta0 时若把支撑再放宽学生可能采到的 token 只能概率性地被覆盖而不是保证。默认前向 KL 是最稳的入口MOPD 场景应显式写beta1.0见下文。实现上是内存友好的(chunk, vocab)形状的 logits 是唯一随词表规模扩展的张量学生隐藏状态按 256 个 token 一 chunk 投影过lm_headtorch.utils.checkpoint在前向后丢弃、反向时重算峰值 logits 内存是chunk × vocab而不是全部有效 token × 词表。teacher_top_k 用 8 还是 16~64teacher_top_k默认 8是经prompt_logprobs向教师请求的每位置候选 token 数8 是冒烟测试级的轻量默认。邻近 RL 框架的 on-policy 蒸馏生产配置在同一量级miles 默认 16、EasyOPD 64真正开训时提到 16~64 是合理的——它直接决定教师分布近似的精度学生侧不受影响全量 logits 在本地。候选不足teacher_top_k 1宽时collator 用 id-1/ logprob-inf补齐并在损失里掩掉。add_tail_bucket保持默认的True它在 top-k 之外追加一个尾桶项log(1 - sum(exp(top_k_logps)))防止 top-k 较小时散度平凡地趋近于零。两个 temperature 别混用参数默认作用对象temperature1.0学生 on-policy 完成结果的采样teacher_temperature1.0散度本身的 softmax 温度两侧都作用teacher_temperature发给教师让 vLLM 在服务端按该温度算 logprobs精确处理不是客户端重缩放同时在compute_loss里作用于学生 logits。这里最容易错的是教师服务器没带--logprobs-mode processed_logprobs时它静默返回原始 logprobs这个参数就只剩学生侧生效且不会有任何报错。max_staleness、队列与同步节奏参数默认怎么定max_staleness4样本最多落后几个权重版本。调小则样本更新鲜但丢弃更多调大则队列更稳但样本更 off-policy。观察sample/dropped_stale_total与sample/staleness_mean再动max_inflight_tasks-1自动自动值为max_staleness × per_device_train_batch_size × gradient_accumulation_steps × num_processesqueue_maxsize1024队列容量与背压指标一起看观测章weight_sync_steps1学生推理服务器跟随训练的节奏调大省同步开销代价是推理侧更旧heartbeat_stale_after_s300.0worker 心跳停更超过该秒数即视为挂起并中止token_budget 与行打包token_budget默认None是一行一个 DP rank 的前向允许打包的最大真实 token 数。None时在训练启动时取学生 vLLM 服务器的max_model_len保证任何 rollout 样本都不会超预算超出预算的样本进不了任何行被警告丢弃并计入batch/dropped_oversize_total。设0则关闭 token 预算改为每 micro-batch 固定打包per_device_train_batch_size × num_processes个样本行间仍做 Σ Lᵢ² 平衡。几个一句话带过的项dtype默认float32异步 trainer 所针对的 training-inference mismatch 度量对 trainer 自身精度敏感要端到端弥合差距学生 vLLM 服务器也要以相同 dtype 服务cp_size 1或sp_size 1的序列维并行直接报错蒸馏在生成之后才构建模型输入transformers 的 context/Ulysses 并行无法作用于原始生成 batchmax_completion_length默认 2048logging_steps默认 1、gradient_checkpointing默认True、未设fp16时bf16默认True均与TrainingArguments不同。日志侧log_completions默认False每log_completions_steps100个被打分的样本记录一批 (prompt, completion) 对——计数口径是 worker 打分的样本数而非优化器步因为 worker 与 trainer 是不同进程看不到global_stepnum_completions_to_print为None时全部打印。什么时候需要多个教师MOPD 路由MOPD多教师 on-policy 蒸馏不是该 trainer 核心目标所基于论文的一部分而是独立方法先做通用 SFT再对各领域独立做基于 RL 的专家训练最后用 MOPD 把冻结的专家融合进单个学生。这个 trainer 只实现第三阶段——各领域专家教师必须已经存在例如用 GRPO/RLOO 单独训好并以 HTTP 服务形式提供把teacher_server_urls指向它们即可。启用方式teacher_server_urls写多个条目数据集每行加一列teacher_id指定打分者。比如数学 prompt 路由给数学专家、代码 prompt 路由给代码专家。每个样本只分发给它匹配的那一个教师绝不跨教师平均或集成teacher_id缺失或未映射时直接 raiseValueError不会静默回退到错误教师。两个大坑⚠️ 每个教师必须与学生共享同一个 tokenizer。完成结果以原始 token id 发给教师教师报回来的候选 id 会直接拿去索引学生自己的词表。词表不同的教师会把学生训到错误的 token 上且只要它的词表不大于学生的这个错误完全静默。同家族的专家如 Qwen2.5 学生由 Qwen2.5 与 Qwen2.5-Coder 两个专家融合满足要求。MOPD 论文自己的第三阶段用的是反向 KL。要复刻它的配置就显式写beta1.0默认的0.0是前向 KL。可运行的双教师示例在examples/async_distillation_math/async_distillation_mopd.pyGSM8K 路由给数学教师Qwen/Qwen2.5-1.5B-Instructiamtarun/python_code_instructions_18k_alpaca路由给代码教师Qwen/Qwen2.5-Coder-1.5B-Instruct学生是Qwen/Qwen2.5-0.5B-Instruct配置中显式beta1.0四个服务器各占一块 GPU。队列满了先查什么按症状读指标指标按键名后缀决定聚合方式(分子, 分母)对按 Σnum/Σden 聚合成比率键名含total的是计数器求和含max/min的取极值其余是 gauge 取窗口均值。下面按症状组织——step 在全部指标里恒指一次完整优化器步per-step 指标是跨所有 rank 的求和per-row 指标是均值。症状一训练在等队列生成受限。队列为空时 trainer 阻塞的时间是perf/rollout_wait_s它高而队列接近空说明训练在挨饿。接着看rollout/generated_tok_s窗口生成吞吐掉线即学生服务器出问题、rollout/inflight在途生成打分任务数填不满上限说明调度不足、rollout/score_s单次教师调用耗时。记住镜像关系perf/rollout_wait_s与rollout/backpressure_s永远不会同时大队列大小 哪边高直接定位瓶颈在哪一侧。症状二生成被队列压住训练受限。队列接近满、rollout/backpressure_s高说明生成被节流、产出在队列里老化——盯住sample/time_in_queue_s单个样本在队列里待了多久与sample/staleness_mean是否攀升。原因通常是训练侧太慢结合症状四排查。症状三教师变慢。教师调用在每个 rollout 的关键路径上延迟直接抬升rollout/duration_s一次 generatescore 往返的墙钟时间与rollout/score_s。MOPD 下看拆分指标teacher_score_s/id慢专家只拖慢路由给它的那些 rollout混合均值会掩盖这一点。服务器间歇性故障看rollout/vllm_retry_total重试过的 vLLM 请求计数放在 rollout 命名空间下统计的是对服务器的请求而非生成文本——退化的服务器否则看起来像莫名的变慢。症状四行负载不均。batch/row_imbalance各行 Σ Lᵢ² 的最大/均值比1.0 为完美明显大于 1 时某个 rank 的行偏长注意力 O(L²) 意味着它会拖慢整组梯度 all-reduce。配batch/row_fill_frac行 token 数相对token_budget的占比一起看长样本难以铺满预算1 万 token 的样本放进 3.2 万预算3 个放得下、4 个永远放不下打包器常只能放 2 个行约 77% 满这是量化效应而非 bugtoken_budget是调节杠杆。batch 指标之间存在自洽关系batch/samples_per_step约等于grad_accum × rank 数 × batch/samples_per_row两侧差百分之零点几属正常。症状五学习信号不对。jsd损失最小化的广义 JSD下降即学生分布向教师收敛与entropy学生自己的预测熵对照着看jsd 下降的同时 entropy 崩塌说明学生在收窄而不是在学习。teacher_entropy是教师在其报告候选上的熵因只有 top-k 个候选过线它从下方界定真实值。MOPD 下加看teacher_jsd/id与teacher_token_frac/id路由偏斜时某个教师会被饿死但它自己的teacher_jsd/id依旧健康没有 token 占比这个指标偏斜不可见。症状六吞吐对不上。吞吐与 MFU 各报两次后缀标注分母_fwd_bwd除以perf/fwd_bwd_s纯计算时间回答有数据时 trainer 跑得多高效低则问题在 trainer_wall_clock除以perf/step_s含队列等待的完整一步回答分配的算力有多少真变成了训练远低于前者则瓶颈在生成或打分。两者之差约等于perf/rollout_wait_s加上优化器与权重同步时间perf/fwd_s与perf/optimizer_s再把一步的计算与优化器耗时拆开权重同步自身拆成_pause_s等 vLLM/_barrier_srank 偏斜/_transfer_s字节传输三段。顺手扫一眼completions/mean_length与completions/clipped_ratio未以 EOS 结束、被max_completion_length截断的完成结果占比batch/masked_token_frac告诉你前向 token 里有多大比例不产生梯度——教师未对某完成位置打分任何候选时该位置在散度中被掩掉但仍参与前向。收尾它刻意不做什么以及代码在哪这个 trainer 刻意保持最小化官方态度是不打算让它长成通用方案新特性只在有显著社区需求时才考虑不支持的功能由使用者在自己的副本上按需扩展。代码里为此留了两个注入点——RolloutWorkerProtocol自定义 rollout worker与WeightTransferProtocol自定义权重同步后端构造 trainer 时传入替代实现即可测试正是靠注入 no-op 实现不依赖真实 vLLM 服务器就能跑。三处定位实现的关键位置训练器与损失trl/experimental/async_distillation/async_distillation_trainer.py分块 JSD、行规划器、DataCollatorForRollout、检查点恢复rollout workertrl/experimental/async_distillation/async_rollout_worker.py异步生成-打分循环、RolloutSample可运行示例examples/async_distillation_math/单教师async_distillation_math.py与 MOPD 双教师async_distillation_mopd.py【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考