TRL AsyncDistillationTrainer 实战手册:三终端部署、关键参数取舍与训练观测
2026/9/17 23:07:21 网站建设 项目流程

TRL AsyncDistillationTrainer 实战手册:三终端部署、关键参数取舍与训练观测

【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl

这份手册面向正在做大模型后训练的工程师,讲透 TRL 实验模块里的AsyncDistillationTrainer:学生模型在自己的 vLLM 服务器上生成 on-policy 样本(即由当前策略自己生成的完成结果),教师通过 HTTP 打分,生成与梯度更新并发进行。读完你能够独立在 3 块 GPU 上跑通一次异步蒸馏,会判断betateacher_top_kmax_staleness等关键取舍,并看懂rollout/sample/perf/指标去定位生成侧或训练侧的瓶颈。

先跑起来:装依赖、起三终端、写脚本

先把依赖安装顺序搞对

硬性前提是vllm>=0.22.0transformers>=5.2.0。这两者目前存在冲突的依赖约束,必须装 vLLM 之后再用--no-deps强制装 transformers,否则 pip 会互相降级:

pip install 'vllm>=0.22.0' pip install 'transformers>=5.2.0' --no-deps

分布式训练只支持 FSDP2,DeepSpeed ZeRO 不可用。

三终端如何分工

一个教师服务器、一个学生推理服务器、训练进程,三者必须分在三块不同的 GPU 上。两个服务器都是普通的vllm serve,但启动参数完全不同——原因在下一章展开:

# GPU 0:教师(只做打分,权重永不更新) CUDA_VISIBLE_DEVICES=0 vllm serve Qwen/Qwen2.5-1.5B-Instruct \ --port 8001 --logprobs-mode processed_logprobs --max-logprobs -1 # GPU 1:学生 vLLM(生成 + NCCL 权重传输) CUDA_VISIBLE_DEVICES=1 VLLM_SERVER_DEV_MODE=1 vllm serve Qwen/Qwen2.5-0.5B-Instruct \ --port 8000 --weight-transfer-config '{"backend":"nccl"}' # GPU 2:训练 CUDA_VISIBLE_DEVICES=2 accelerate launch train_async_distillation.py

教师那两个 flag 都有具体用途:--logprobs-mode processed_logprobsteacher_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", split="train") trainer = AsyncDistillationTrainer( model="Qwen/Qwen2.5-0.5B-Instruct", train_dataset=dataset, ) trainer.train()

缺省配置下教师指向http://localhost:8001、学生指向http://localhost:8000,学习率默认1e-6(不是TrainingArguments5e-5)。更贴近实战的单教师示例在examples/async_distillation_math/async_distillation_math.py:GSM8K 数据、max_steps=100、trackio 上报,照抄改路径即可。

它是怎么运转的:四个角色各干什么

与同步版DistillationTrainer(教师、学生、优化器在同一进程里顺序执行)不同,异步版把过程拆成四个角色:

  1. rollout workerspawn出的子进程,显式清空CUDA_VISIBLE_DEVICES):调学生 vLLM 的/v1/completions采样出完成结果,再调教师 vLLM 的/v1/completionsprompt_logprobs做 teacher-forced 打分——教师只对学生自己的完成结果逐位置报 logprob,不生成任何新 token。一次 rollout = 一次生成 + 一次打分,恰好产出一个训练样本。蒸馏没有 GRPO 那种组内基线,所以 prompt 不会被重复采样。
  2. mp.Queue 缓冲队列(容量queue_maxsize,默认 1024):worker 把带教师稀疏分布的打分样本推进来。
  3. 训练循环(主进程):每次拉一个样本,staleness 超过max_staleness就丢弃(计入sample/dropped_stale_total);拉够后规划成 rows——每个 DP rank 一行,行内样本拼成单条序列且position_ids逐样本重置,按 Σ Lᵢ² 贪心分桶平衡;gradient_accumulation_steps个 micro-batch 构成一次优化器步,损失是广义 JSD。
  4. 权重同步:每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 分布的广义 JSD,beta在两者之间插值:

beta散度行为实际参与计算的支撑集
0.0(默认)前向 KLmean-seeking完整teacher_top_k宽的教师支撑 + 尾桶
1.0反向 KLmode-seeking仅两个候选:教师 top-1 与完成结果的实际 token
中间值插值走收窄路径,同1.0

为什么支撑集要变?线上协议只能保证教师 logprob 对"教师 top-1"和"实际 token"这两个身份一定可用;beta>0 时若把支撑再放宽,学生可能采到的 token 只能概率性地被覆盖,而不是保证。默认前向 KL 是最稳的入口;MOPD 场景应显式写beta=1.0(见下文)。

实现上是内存友好的:(chunk, vocab)形状的 logits 是唯一随词表规模扩展的张量,学生隐藏状态按 256 个 token 一 chunk 投影过lm_headtorch.utils.checkpoint在前向后丢弃、反向时重算,峰值 logits 内存是chunk × vocab而不是"全部有效 token × 词表"。

teacher_top_k 用 8 还是 16~64

teacher_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_totalsample/staleness_mean再动
max_inflight_tasks-1(自动)自动值为max_staleness × per_device_train_batch_size × gradient_accumulation_steps × num_processes
queue_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 > 1sp_size > 1的序列维并行直接报错(蒸馏在生成之后才构建模型输入,transformers 的 context/Ulysses 并行无法作用于原始生成 batch);max_completion_length默认 2048;logging_steps默认 1、gradient_checkpointing默认True、未设fp16bf16默认True,均与TrainingArguments不同。日志侧,log_completions(默认False)每log_completions_steps(100)个"被打分的样本"记录一批 (prompt, completion) 对——计数口径是 worker 打分的样本数而非优化器步,因为 worker 与 trainer 是不同进程,看不到global_stepnum_completions_to_printNone时全部打印。

什么时候需要多个教师: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。要复刻它的配置就显式写beta=1.0,默认的0.0是前向 KL。

可运行的双教师示例在examples/async_distillation_math/async_distillation_mopd.py:GSM8K 路由给数学教师Qwen/Qwen2.5-1.5B-Instructiamtarun/python_code_instructions_18k_alpaca路由给代码教师Qwen/Qwen2.5-Coder-1.5B-Instruct,学生是Qwen/Qwen2.5-0.5B-Instruct,配置中显式beta=1.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_srollout/backpressure_s永远不会同时大,队列大小 + 哪边高,直接定位瓶颈在哪一侧。

症状二:生成被队列压住(训练受限)。队列接近满、rollout/backpressure_s高,说明生成被节流、产出在队列里老化——盯住sample/time_in_queue_s(单个样本在队列里待了多久)与sample/staleness_mean是否攀升。原因通常是训练侧太慢,结合症状四排查。

症状三:教师变慢。教师调用在每个 rollout 的关键路径上,延迟直接抬升rollout/duration_s(一次 generate+score 往返的墙钟时间)与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% 满),这是量化效应而非 bug,token_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_sperf/optimizer_s再把一步的计算与优化器耗时拆开,权重同步自身拆成_pause_s(等 vLLM)/_barrier_s(rank 偏斜)/_transfer_s(字节传输)三段。顺手扫一眼completions/mean_lengthcompletions/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 worker:trl/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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询