TRL 异步蒸馏部署指南:3 个终端跑通 AsyncDistillationTrainer 并看懂关键指标
TRL 异步蒸馏部署指南3 个终端跑通 AsyncDistillationTrainer 并看懂关键指标【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trlAsyncDistillationTrainer是 TRL 中的异步 on-policy 蒸馏训练器学生模型自己生成样本、教师通过 vLLM 服务打分生成与梯度更新并行跑教师永不需要本地加载。它直接解决同步蒸馏的两个老问题——生成和更新轮流占卡、教师和学生抢同一块 GPU。三个角色并发流水线怎么搭起来的同步版DistillationTrainer里生成、教师前向、梯度更新在同一进程内串行执行。异步版把这件事拆成两个长期并发的角色外加一条权重回传通道后台 rollout worker一个 spawn 出来的子进程启动时清空CUDA_VISIBLE_DEVICES内部跑 asyncio 事件循环。它从数据集取 prompt先调学生 vLLM 的/v1/completions采样出完成结果再调路由到的教师服务器做 teacher-forced 打分——请求带prompt_logprobs、max_tokens1教师不生成任何新 token只报每个位置的 logprob。一次 rollout 就是「1 个 prompt → 1 条学生完成 → 1 次教师打分」恰好产出一个训练样本蒸馏没有可跨生成计算的基线所以 prompt 不会被重复采样。生成与打分在同一个任务_generate_and_score_one内完成并发度来自最多max_inflight_tasks个在途任务。主进程训练循环从队列里逐个拉取样本计算广义 JSD 损失并更新学生权重。权重回传每weight_sync_steps默认 1个优化器步更新后的学生权重经 NCCL 推回学生的 vLLM 服务器保证生成侧的策略跟着走。打分协议只在线上传稀疏切片完成序列 教师逐位置 top-teacher_top_k候选 logprob外加 vLLM 总会报告的 realized token 和尾桶完整词表从不经 HTTP 传输。实现见 async_rollout_worker.py生成/打分循环与RolloutSample和 async_distillation_trainer.py损失与训练循环。因为生成始终领先于训练样本可能反映略微过期的策略。max_staleness默认4限定一个样本最多可以落后多少个权重更新超了直接丢弃并计入sample/dropped_stale_total——这是控制 off-policy 程度的阀门。三终端部署GPU 分配与启动命令三个进程各占一块 GPU。教师是静态服务器只被打分、从不更新所以不需要 dev 模式学生的 vLLM 服务需要开 dev 模式和 NCCL 权重传输trainer 才能把新权重推进去。下面这条命令把三个终端一次性列全# 终端 1教师GPU 0静态 CUDA_VISIBLE_DEVICES0 vllm serve Qwen/Qwen2.5-1.5B-Instruct \ --port 8001 --logprobs-mode processed_logprobs --max-logprobs -1 # 终端 2学生推理服务GPU 1接收权重更新 CUDA_VISIBLE_DEVICES1 VLLM_SERVER_DEV_MODE1 vllm serve Qwen/Qwen2.5-0.5B-Instruct \ --port 8000 --weight-transfer-config {backend:nccl} # 终端 3训练GPU 2 CUDA_VISIBLE_DEVICES2 accelerate launch train_async_distillation.py关键参数就四个教师侧--logprobs-mode processed_logprobs让teacher_temperature在服务端真正作用于返回的 logprobs--max-logprobs -1解除 vLLM 默认 20 的每位置 logprob 上限学生侧VLLM_SERVER_DEV_MODE1和--weight-transfer-config {backend:nccl}是权重推入 vLLM 的前提。依赖方面该 trainer 要求vllm0.22.0和transformers5.2.0两者当前存在冲突的依赖约束必须先装 vLLM、再用--no-deps强制安装 transformers顺序反了装不上。分布式训练只支持 FSDP2不支持 DeepSpeed ZeRO。另外学生模型的dtype默认是float32异步 trainer 度量的 training-inference mismatch 对 trainer 自身精度敏感要端到端弥合精度差还需学生 vLLM 服务用相同 dtypevllm serve --dtype教师服务器不受影响。最小训练脚本如下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()更贴近实战的完整示例是 async_distillation_math.pyGSM8K、max_steps100、report_totrackio。注意几个与TrainingArguments不同的默认值learning_rate默认1e-6不是5e-5、bf16在未设fp16时默认True、gradient_checkpointing默认True、logging_steps默认1。⚙️ beta 与 teacher_top_k 怎么选beta是广义 JSD 的插值系数0.0为前向 KLmean-seeking默认值1.0为反向 KLmode-seeking中间值在两者间插值超出[0.0, 1.0]会在__post_init__抛 ValueError。beta还会改变散度计算的支撑集supportbeta0.0用教师报告的完整teacher_top_k宽度支撑加尾桶因为前向 KL 的加权恰好就是该支撑能提供的beta ! 0.0支撑收窄到两个候选——教师 top-1 和完成结果的实际 token。原因在打分协议本身不传更宽或完整词表时线上协议保证教师 logprob 可用的只有这两个身份再宽的支撑都只是概率性覆盖学生可能采样的 tokenbeta1.0宽度进一步降到 1纯反向 KL 是纯学生加权期望教师 top-1 不贡献。teacher_top_k是每位置向教师请求的候选数默认8只是适合冒烟测试的轻量档位生产上提到16–64是合理的邻近 RL 框架的 on-policy 蒸馏用稀疏教师支撑miles 默认16、EasyOPD64。超过20就要求教师以--max-logprobs -1启动。学生侧散度是精确的本地算、不近似。add_tail_bucket默认True在候选之外补一个尾桶值为log(1 - sum(exp(top_k_logps)))兜住 top-k 之外的剩余概率质量避免teacher_top_k较小时散度平凡地偏小。 指标解读按排障口径组织的四个命名空间规约规则按键名后缀走(numerator, denominator)对按 Σnum/Σden 聚成比率名字含total的计数器求和含max/min的取极值其余 gauge 取窗口均值。所有_per_step指标里的 step 恒指一次完整的优化器步gradient_accumulation_steps个 micro-batch一步共覆盖gradient_accumulation_steps × world_size个 row 槽位。先判瓶颈rollout 队列的四只仪表worker 推样本进队列trainer 从队列拉四个指标描述这一个缓冲区指标回答的问题sample/rollout_queue_size现在有几个样本在排队sample/time_in_queue_s单个样本在队列里待了多久它 off-policy 程度的时间部分perf/rollout_wait_s训练因队列空阻塞了多久rollout/backpressure_s生成因队列满被卡了多久perf/rollout_wait_s与rollout/backpressure_s互为镜像不会同时大。结合队列大小判断瓶颈在哪一侧队列接近空、perf/rollout_wait_s高 →生成受限训练在挨饿接着看rollout/generated_tok_s、rollout/inflight、rollout/score_s队列接近满、rollout/backpressure_s高 →训练受限生成被节流产出在队列里老化盯住sample/staleness_mean是否攀升两者都接近零 → 两侧平衡。两种吞吐口径别把生成的慢算到训练头上吞吐与 MFU 各报两次基于同一步只差分母_fwd_bwd除以perf/fwd_bwd_s纯前向反向回答「trainer 有数据时跑得多高效」低就说明问题在训练侧_wall_clock除以perf/step_s完整一步含等 rollout 的时间回答「分配到的算力有多少真变成了训练」它远低于前者说明瓶颈大概率在生成。两者之差就是perf/rollout_wait_s加上优化器和权重同步耗时。只看_fwd_bwd会掩盖生成与打分的 GPU 时数只报_wall_clock又可能把教师延迟归咎于 trainer。对应指标perf/forwarded_tok_s_fwd_bwd/perf/forwarded_tok_s_wall_clock、perf/trained_tok_s_wall_clock、perf/mfu_fwd_bwd/perf/mfu_wall_clock。其余命名空间速查rollout/一次生成打分往返rollout/duration_s从派发到打分落库的墙钟时间生成加教师调用rollout/score_s其中教师 HTTP 调用占比——教师调用在关键路径上慢教师直接抬高 durationrollout/generated_tok_s窗口化生成吞吐rollout/inflight在途 rollout 数rollout/vllm_retry_total重试过的 vLLM 请求数退化的服务器否则看起来只是「莫名变慢」。completions/学生为一个 prompt 生成了什么completions/mean_length、completions/min_length/completions/max_length、completions/clipped_ratio未以 EOS 结束、被max_completion_length截断的占比。sample/样本进入训练时sample/forwarded_tokens_mean/_maxprompt生成样本级sample/trained_tokens_mean损失真正覆盖的 token 数——教师没打分任何候选的位置会被掩掉但仍参与前向所以 trained ≠ generatedsample/staleness_mean/_max数据落后当前策略多少个版本jsd展示 off-policy 的效果staleness 展示原因sample/dropped_stale_total。batch/per-step 量是对整步跨所有 rank 的求和——batch/samples_per_step、batch/forwarded_tokens_per_step、batch/trained_tokens_per_step、batch/microbatches_per_step实际计数per-row 量是均值——batch/samples_per_row、batch/row_tokens_mean/_max。batch/samples_per_step ≈ 每步 row 槽位数 × batch/samples_per_row两侧只在步内 micro-batch 方差范围内有零点几的偏差求和 vs 均值。batch/masked_token_frac是不产生梯度的前向 token 占比batch/row_fill_frac是行相对token_budget的填充率长样本下偏低是量化效应token_budget是调节杠杆batch/row_imbalance为max Σ Lᵢ² / mean Σ Lᵢ²注意力 O(L²)它预测哪个 rank 会拖慢梯度 all-reduce1.0 为完美batch/pad_frac是 rank 间填充只增加广播字节、前向前会被剥掉batch/dropped_oversize_total是因超过token_budget被丢弃的样本数。散度无前缀学习信号本身jsd按配置的beta计算的广义 JSD下降意味着学生分布向教师收敛entropy是学生自身预测熵若它随jsd下降而崩塌说明学生在收窄而不是学习teacher_entropy是教师在所报候选上的熵因只有teacher_top_k个候选过线它是真实值的下界。MOPD多教师按领域路由打分MOPDMulti-Teacher On-Policy Distillation是独立于本 trainer 核心目标论文的另一个方法流程三段——先通用 SFT再对每个领域做基于 RL 的专家训练最后用 MOPD 把冻结的领域专家融合进单个学生。AsyncDistillationTrainer只实现第三段融合各领域专家教师必须已单独训练好例如用GRPOTrainer/RLOOTrainer并以 HTTP 提供推理服务。论文自己的 Stage 3 用反向 KL所以 MOPD 场景应显式beta1.0别用默认的前向 KL。teacher_server_urls给多个条目后每行数据的teacher_id列决定由哪个教师打分——数学 prompt 路由给数学专家、代码 prompt 路由给代码专家各自独立服务。样本只发给它匹配的那一个教师绝不跨教师求平均或集成teacher_id缺失或未映射会直接抛 ValueError而不是静默回退到某个教师。可运行的双教师示例见 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。每个教师必须与学生共享同一个 tokenizer。完成结果以原始 token id 发给教师教师报回的候选 id 直接索引学生自己的词表。词表不同的教师会把学生训练到错误的 token 上——除非它的词表严格大于学生否则这个错误是静默的。MOPD 专属指标按教师 id 拆分teacher_jsd/id限定在该教师打分 token 上的 jsd不同领域教师可以以非常不同的速率发散混合jsd会把它们混为一谈teacher_entropy/id同理拆分teacher_token_frac/id是该教师打分的 token 占比路由偏斜靠它才可见——一个被饿死的教师照样报告健康的teacher_jsd/idteacher_score_s/id把教师耗时按 id 拆开一个慢专家只拖慢路由给它的那些 rollout混合均值会把这个事实盖住。没有 per-teacher 的entropy学生熵是其自身策略的属性与哪个教师打分无关。常见坑与已知边界检查点与恢复ignore_data_skip默认True基础 Trainer 的 skip-and-replay 循环对实时 rollout 队列不适用。取而代之每个检查点会写一个rollout_state.json记录第一个尚未被训练的 prompt 索引恢复时 worker 直接快进到该位置、无需重放。注意保存的是已训练位置而非生成器位置worker 领先训练最多一个队列深度那些已缓冲但未训练的样本在运行结束时丢掉若按生成器位置恢复就会跳过「已生成但未训练」的 prompt。流式数据集IterableDataset无法重新定位恢复时 worker 从 prompt 0 重启。典型坑依赖冲突vLLM 与 transformers 约束打架必须先装 vLLM、再--no-deps装 transformers且分布式只认 FSDP2。教师漏掉--logprobs-mode processed_logprobsteacher_temperature会静默地只作用于学生侧 logits教师照报原始 logprobs。序列维并行不支持cp_size 1或sp_size 1直接抛错——蒸馏在 trainer 内部于生成之后构建模型输入transformers 的 context-parallel / Ulysses 输入分片套不到原始生成 batch 上。数据加载被强制为split_batchesTrue、dispatch_batchesTrue主进程驱动 dataloader、batch 广播给其他进程这是异步 IterableDataset 数据加载器正确工作的前提不要试图覆盖。设计定位这个 trainer 刻意保持最小化官方态度是不打算把它养成通用方案。缺少某个功能时建议直接克隆仓库改造源码里RolloutWorkerProtocol与WeightTransferProtocol两个 Protocol 定义了可注入的自定义 rollout worker 与权重同步后端测试就是靠注入 no-op 实现脱离真实 vLLM 服务器跑的。源码与示例索引配置async_distillation_config.py —— 全部默认值、beta范围校验与__post_init__约束训练器async_distillation_trainer.py —— 广义 JSD、分块lm_head投影chunk 256checkpoint 重算控峰值显存、行规划与指标规约Rollout workerasync_rollout_worker.py —— 生成/打分循环、RolloutSample、教师路由与teacher_id校验权重传输weight_transfer.pyvLLM 客户端vllm_client.py单教师示例async_distillation_math.py双教师 MOPD 示例async_distillation_mopd.py同步版蒸馏distillation_trainer.py同样用 JSD 目标的 server_distillation【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考