TRL AsyncDistillationTrainer 异步蒸馏实战指南:三终端部署与训练指标分诊
【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl
AsyncDistillationTrainer是 TRL 实验模块中的异步 on-policy 蒸馏训练器:学生模型在 vLLM 服务器上自己生成样本,教师经 vLLM 服务打分,生成与梯度更新并发进行。本文覆盖选型、部署、beta/teacher_top_k设置与按问题排查指标。
何时选它:先看清与同步蒸馏的差异
同步版DistillationTrainer把教师加载在本地,生成、教师前向、梯度更新在同一进程内顺序执行。异步版把教师完全搬到线上:只需一个教师 vLLM 服务器的 URL,教师可以跑在独立硬件上,规模也可以大到与学生并排装不下。两者对比如下。
| 维度 | 同步DistillationTrainer | 异步AsyncDistillationTrainer |
|---|---|---|
| 教师加载方式 | 本地加载完整模型 | 永不本地加载,仅一个 HTTP URL |
| 硬件摆放 | 师生与训练同进程同机 | 教师可独立部署,甚至跨机 |
| 教师权重是否更新 | 不参与训练,静态 | 静态:trainer 从不向任何教师传权重 |
| 生成与训练关系 | 交替执行 | 后台 worker 生成打分,主进程并行训练 |
| on-policy 程度 | 学生自生成 | 学生自生成(每个训练样本都由学生产生) |
| 分布式支持 | 常规后端的通用支持 | 仅 FSDP2,不支持 DeepSpeed ZeRO |
| 额外依赖 | 无 | vllm>=0.22.0+ 学生侧 NCCL 权重传输 |
选型结论:教师与学生必须同机同进程、且追求最简流程时用同步版;教师更大、需要独立硬件、或希望训练不等生成时用该 trainer。
一个样本的完整旅程
先固定五个术语,全文沿用:
- rollout:一个 prompt 被学生生成一次、被教师打分一次;
- sample:prompt + 学生完成结果 + 教师逐位置的稀疏 top-k 候选,是跨进程传输的单元;
- row:一个 DP rank 在一次 micro-batch 中前向的内容,若干 sample 拼接成单条序列、
position_ids在样本边界重置; - micro-batch:
world_size个 row 的集合,每个 rank 前向自己的那行; - optimizer step:
gradient_accumulation_steps个 micro-batch 累积而成的一次优化器步。所有指标中的 "step" 恒指它,而非 micro-batch。
数据集行 = PROMPT(消息列表 + 可选 `teacher_id` 列) └─ ROLLOUT 1 个 prompt -> 1 次学生生成 -> 1 次教师打分 ├─ generate 学生 vLLM /v1/completions,按 temperature/top_p 采样 └─ score 路由到的教师 /v1/completions,带 `prompt_logprobs`,teacher-forced └─ SAMPLE prompt + completion + 教师稀疏分布 ════════════ 进程边界:`rollout_buffer`(mp.Queue)════════════ └─ SAMPLE 逐个拉取;staleness 超过 `max_staleness` 即丢弃 └─ ROW 规划器分入 dp 个 row 之一,按 Σ Lᵢ² 平衡 └─ MICRO-BATCH dp 个 row,每 rank 一个 └─ PACKED ROW 拼接成 1 条序列,position_ids 逐样本重置 └─ FORWARD `compute_loss`,bs=1,rank 间填充先被剥掉 └─ OPTIMIZER STEP 由 `grad_accum` 个 micro-batch 构成蒸馏没有 GRPO 那种需要跨生成求的 advantage 基线,所以 prompt 不重复、一次 rollout 恰好产出一个训练样本。三个环节的解耦体现在:生成由后台 spawn 子进程负责(CUDA_VISIBLE_DEVICES被清空,worker 内部跑 asyncio 循环,最多max_inflight_tasks个任务在途),教师经 HTTP 打分而非本地前向,主进程只做拉样本、算损失、更新权重。每weight_sync_steps步,更新后的学生权重经 NCCL 推给学生 vLLM 服务器;生成始终领先于训练,样本可能反映过期策略,落后超过max_staleness个权重更新的样本被丢弃并计入sample/dropped_stale_total。
三终端如何点亮
三个角色必须分卡运行。教师和学生都是普通的vllm serve,但启动参数不同;trainer 脚本保持最小即可。
最小训练脚本(train_async_distillation.py):
# 加载数据集并启动异步蒸馏,默认配置即可跑通 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()终端 1:教师服务器(GPU 0)
# 教师只被打分、永不更新,需要服务端温度与解除 logprob 上限 CUDA_VISIBLE_DEVICES=0 vllm serve Qwen/Qwen2.5-1.5B-Instruct \ --port 8001 \ --logprobs-mode processed_logprobs \ --max-logprobs -1终端 2:学生 vLLM 服务器(GPU 1)
# 学生只被生成,且必须开启 dev 模式与 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"}'终端 3:训练(GPU 2)
# trainer 与两个服务器分卡,accelerate 单卡启动即可 CUDA_VISIBLE_DEVICES=2 accelerate launch train_async_distillation.py两个 flag 各自对应一个机制:
--logprobs-mode processed_logprobs:teacher_temperature由 vLLM 在服务端作用于返回的 logprobs,而非客户端重缩放。缺失该 flag 时教师静默返回原始 logprobs,teacher_temperature只作用于学生侧。--max-logprobs -1:解除 vLLM 默认每 token 最多 20 个 logprob 的上限,teacher_top_k超过 20 时必须有它。
学生的 dev 模式与 NCCL 配置则解决另一件事:trainer 每weight_sync_steps步要把新权重推进生成服务器,没有这个通道生成会一直停留在初始策略。
约束:先装 vLLM 再强制装 transformers:
pip install 'vllm>=0.22.0'之后pip install 'transformers>=5.2.0' --no-deps。原因:当前两包的依赖约束互相冲突,直接联装会降级 transformers。
损失与两个旋钮:beta和teacher_top_k
目标是最小化学生与教师在逐位置 token 分布上的广义 Jensen-Shannon 散度。它决定了两个行为。
beta决定加权方式,也决定在哪些候选上算散度。beta=0.0是前向 KL(mean-seeking),beta=1.0是反向 KL(mode-seeking),中间值在两者间插值,取值必须在[0.0, 1.0]内。支撑集选择随beta变化:
beta=0.0:用教师上报的完整teacher_top_k宽支撑(外加尾桶)。前向 KL 的期望恰好按教师概率加权,教师上报的这组候选正好够用。beta != 0.0:支撑收窄到最多两个候选——教师自己的 top-1 与完成结果的实际 token。线上协议在不传输更宽词表的前提下只保证这两个 token 有教师 logprob,更宽的支撑只能算概率性覆盖学生可能采样的 token,不是保证。特例是beta=1.0:纯反向 KL 是学生自身加权的期望,教师 top-1 对目标无贡献,支撑进一步收窄为仅实际 token(宽度 1,去重后)。
后果是实际的:调teacher_top_k在beta != 0.0时几乎不改变训练信号,它只影响beta=0.0的精度与teacher_entropy的下界。
teacher_top_k与尾桶决定教师分布的近似质量。线上只传输每个位置teacher_top_k个候选,外加 vLLM 总会报告的 realized token(即使它在 top-k 之外)和add_tail_bucket的尾桶(log(1 - Σexp(top_k_logps)),候选少时防止散度平凡地接近零)。学生侧的散度是精确的,完整 logits 本地可得;被近似的只有教师侧。默认8是冒烟测试级,邻近 RL 框架的 on-policy 蒸馏生产配置在16(miles 默认)到64(EasyOPD)之间,正式训练上调到16–64合理。
内存代价由分块实现控制:投影学生隐藏状态过lm_head时按 256 个 token 一块进行,(chunk_size, vocab_size)的 logits 是唯一随词表扩展的张量,checkpoint 机制在反向时按需重算。峰值 logits 内存是256 × vocab_size,而不是全部有效 token 乘词表。
参数速查
以下每项回答"改它改变什么、何时该改"。
| 参数 | 默认值 | 改它改变什么 / 何时该改 |
|---|---|---|
beta | 0.0 | 0=前向 KL,1=反向 KL,中间插值;同时切换支撑集(见上节)。MOPD 场景论文用反向 KL,需显式设1.0 |
teacher_top_k | 8 | beta=0.0下决定教师分布近似宽度。冒烟用 8,正式训练提至 16–64;超过 20 需教师--max-logprobs -1 |
teacher_temperature | 1.0 | 散度两侧的 softmax 温度,服务端计算教师 logprobs 且作用于学生 logits。与采样temperature无关。想平滑教师分布再改 |
add_tail_bucket | True | 是否追加尾桶吸收 top-k 之外的概率质量。teacher_top_k很小时关掉会使散度偏小,不建议关 |
max_staleness | 4 | 样本允许落后多少权重更新,超过即丢。trainer-bound 时调大换吞吐、代价是 off-policy;生成侧慢时它是丢弃率的主要来源 |
max_inflight_tasks | -1 | 两个服务器的在途任务上限。自动值 =max(max_staleness, 1) × per_device_train_batch_size × gradient_accumulation_steps × num_processes(下限 1 防止max_staleness=0清零调度导致空队列挂死)。想加大并发改它 |
queue_maxsize | 1024 | rollout 队列缓冲上限,决定样本最多能"存"多少份。显存/内存紧张或想减小 staleness 上界时调小 |
weight_sync_steps | 1 | 两次 NCCL 权重同步之间的训练步数。调大省同步开销,但生成侧使用的策略更旧 |
token_budget | None | 单 row 打包的最大真实 token 数。None时取学生服务器max_model_len(启动时查询);<=0切换为固定样本数打包 |
dtype | "float32" | 学生加载精度。默认 float32 因为 training-inference mismatch 度量对 trainer 精度敏感;端到端弥合还要求vllm serve --dtype一致。要省显存改bfloat16并同步服务器 |
request_timeout | 600 | 对任一 vLLM 服务器的单请求超时(秒)。大 batch 打分慢时调大 |
weight_sync_timeout | 1800 | 一次权重传输的超时,超时 raise 而非挂起。模型大或网络慢时调大 |
heartbeat_stale_after_s | 300.0 | worker 心跳超过该秒数视为挂起并中止。长请求阻塞 worker 时调大 |
log_completions | False | 是否周期性记录 (prompt, completion) 对。调试生成质量时开 |
log_completions_steps | 100 | 按 worker 已打分样本数计数(worker 是独立进程,看不到global_step),不是优化器步 |
num_completions_to_print | None | 每次打印的完成结果数,None为全部 |
另有五项默认值不同于 transformersTrainingArguments,容易按旧习惯误判:
| 参数 | 该 trainer 默认 | TrainingArguments默认 | 说明 |
|---|---|---|---|
learning_rate | 1e-6 | 5e-5 | 蒸馏更新幅度小,学习率相应低两个量级 |
bf16 | 未设fp16时为True | False | 默认走 bf16 |
gradient_checkpointing | True | False | 默认开启以省显存 |
logging_steps | 1 | 500 | 每步都记,指标密度高 |
ignore_data_skip | True | False | 基础 Trainer 的 skip-and-replay 不适用于实时 rollout 队列,见 MOPD 一节的恢复语义 |
指标分诊:按四个问题查
口径先讲清:(numerator, denominator)对按 Σnum/Σden 聚合为比率;名字含total的是计数器求和;含max/min的取窗口极值;其余 gauge 取窗口均值。token 分三种口径:generated(学生实际生成)、forwarded(前向处理的全部 token,prompt+生成)、trained(completion_mask == 1覆盖、损失真正计算的子集)。注意 trained ≠ generated:教师未对某完成位置报告任何候选时,该位置在散度中被掩掉但仍参与前向。
问题一:瓶颈在生成侧还是训练侧?
rollout 队列相当于生成与训练之间的车间缓冲区,四个指标描述它,其中两个等待指标互为镜像:perf/rollout_wait_s统计队列为空、rollout/backpressure_s统计队列已满,同一时刻只可能成立其一,因此两者永远不会同时很大。结合队列水位读:
| 指标 | 口径 | 含义 |
|---|---|---|
sample/rollout_queue_size | gauge,窗口均值 | 当前排队等待的样本数 |
sample/time_in_queue_s | gauge,窗口均值 | 单个样本入队到被训练的等待秒数 |
perf/rollout_wait_s | gauge,窗口均值 | 训练因队列为空阻塞的秒数 |
rollout/backpressure_s | gauge,窗口均值 | 生成因队列已满被节流的秒数 |
判读:队列近空且perf/rollout_wait_s高 → 生成受限,再看rollout/generated_tok_s、rollout/inflight、rollout/score_s;队列近满且rollout/backpressure_s高 → 训练受限,注意sample/staleness_mean随之攀升;两者都近零 → 平衡。
慢教师直接体现在每次往返上,教师调用在每个 rollout 的关键路径里(不同于 GRPO 的独立打分循环):
| 指标 | 口径 | 含义 |
|---|---|---|
rollout/duration_s | gauge,窗口均值 | 一次 rollout 墙钟时间,生成加教师调用 |
rollout/score_s | gauge,窗口均值 | 其中教师/v1/completions耗时 |
teacher_score_s/<id> | gauge,仅 MOPD | 按教师拆分的打分耗时,慢专家只拖慢路由给它的 rollout |
rollout/generated_tok_s | 比率(窗口) | 生成吞吐,停顿会显现 |
rollout/inflight | gauge | 在途 rollout 数(正在生成或打分) |
rollout/vllm_retry_total | total计数 | 重试过的 vLLM 请求(学生或教师服务器);退化服务器否则只是"莫名的变慢" |
生成侧本身的健康度看完成结果:
| 指标 | 口径 | 含义 |
|---|---|---|
completions/mean_length | gauge,窗口均值 | 每个 rollout 的生成 token 数 |
completions/min_length/completions/max_length | 极值 | 窗口内最短/最长完成结果 |
completions/clipped_ratio | 比率 | 未以 EOS 结束、被max_completion_length截断的占比 |
问题二:样本在变陈旧吗?
| 指标 | 口径 | 含义 |
|---|---|---|
sample/forwarded_tokens_mean | gauge | 单样本 token 数(prompt+生成) |
sample/forwarded_tokens_max | 极值 | 单样本最大 token 数 |
sample/trained_tokens_mean | gauge | 其中损失覆盖的 token 数 |
sample/staleness_mean/sample/staleness_max | gauge / 极值 | 数据落后当前策略多少个版本;jsd展示 off-policy 对损失的效果,这里展示原因 |
sample/dropped_stale_total | total计数 | 因超过max_staleness被丢弃的样本数 |
问题三:学生在收敛还是在坍缩?
这些是无前缀指标,窗口内 trained token 上的均值。
| 指标 | 口径 | 含义 |
|---|---|---|
jsd | 窗口均值 | 损失最小化的广义 JSD(按配置的beta)。下降=学生在向教师分布收敛 |
entropy | 窗口均值 | 学生自身预测熵。jsd下降的同时这里崩塌=学生在收窄而非学习 |
teacher_entropy | 窗口均值 | 教师在所报候选上的熵;只有teacher_top_k个候选过线,它从下方界定真实值 |
teacher_jsd/<id> | 窗口均值,仅 MOPD | 限定在该教师打分 token 上的jsd,混合jsd会把不同领域教师的不同速率混为一谈 |
teacher_entropy/<id> | 窗口均值,仅 MOPD | teacher_entropy的同样拆分 |
teacher_token_frac/<id> | 比率,仅 MOPD | 该教师打分 token 的占比。路由偏斜否则不可见:被饿死的教师仍会报告健康的teacher_jsd/<id> |
没有 per-teacher 的entropy:学生熵是其自身策略的属性,与哪个教师打分无关,混合指标已覆盖。
问题四:算力真正花在训练上的比例?
吞吐与 MFU 各报两次,基于同一次优化器步,仅除数不同:_fwd_bwd除以perf/fwd_bwd_s(纯计算),回答"有数据时 trainer 跑得多高效",低则问题在 trainer;_wall_clock除以perf/step_s(完整一步,含等 rollout),回答"分配到的算力有多少变成了训练",远低于前者则瓶颈大概率在生成。两者之差是perf/rollout_wait_s加优化器与权重同步时间。
| 指标 | 口径 | 含义 |
|---|---|---|
perf/step_s | gauge | 优化器步之间的墙钟时间,含计算、优化器、权重同步、队列等待 |
perf/fwd_bwd_s | gauge | 前向+反向,对一步的 micro-batch 求和;所有_fwd_bwd指标的分母 |
perf/fwd_s | gauge | 其中前向部分;fwd_s / fwd_bwd_s接近 1/3 是常见划分,更高意味着反向便宜或重计算落在前向 |
perf/optimizer_s | gauge | optimizer.step()耗时 |
perf/rollout_wait_s | gauge | 见问题一 |
perf/weight_sync_s | gauge | 一次完整同步,另有_pause_s(等 vLLM)、_barrier_s(rank 偏斜)、_transfer_s(字节传输) |
perf/forwarded_tok_s_fwd_bwd/perf/forwarded_tok_s_wall_clock | 比率 | 两种口径的每秒前向 token 数 |
perf/trained_tok_s_wall_clock | 比率 | 同口径,只统计损失见过的 token |
perf/mfu_fwd_bwd/perf/mfu_wall_clock | 比率 | 两种口径的模型 FLOPs 利用率 |
batch 级指标在问题四之外还要会读打包健康度。一步固定含gradient_accumulation_steps × world_size个 row-slot;_per_step指标跨全部 rank 求和,row_*是行均值,两者按samples_per_step ≈ row-slots × samples_per_row校验,允许百分之零点几的偏差(per-step 是求和、per-row 是均值)。
| 指标 | 口径 | 含义 |
|---|---|---|
batch/samples_per_step | 求和 | 每优化器步的训练样本数 |
batch/forwarded_tokens_per_step/batch/trained_tokens_per_step | 求和 | 一步内前向/训练的 token 数 |
batch/microbatches_per_step | 实测计数 | 实际值,非配置读出 |
batch/masked_token_frac | 比率 | completion_mask == 0的前向 token 占比,前向中不产生梯度的部分 |
batch/samples_per_row | gauge | 规划器每行打包的样本数 |
batch/row_tokens_mean/batch/row_tokens_max | gauge / 极值 | 行内 token 规模 |
batch/row_fill_frac | 比率 | 行 token 数相对token_budget。低=预算没填满。样本很长时这是量化效应而非 bug:1 万 token 样本铺 3.2 万预算(3 个放得下、4 个永远放不下,常只能放 2 个、行约 77% 满),1k 样本则几乎完美铺满;token_budget是调节杠杆,且注意力按序列 O(L²),更满的长行不线性增加内存 |
batch/row_imbalance | gauge | 各行max Σ Lᵢ² / mean Σ Lᵢ²。注意力 O(L²),它预测哪个 rank 拖慢 all-reduce;1.0 为完美 |
batch/pad_frac | 比率 | rank 间填充,只增加广播字节,前向之前剥掉 |
batch/dropped_oversize_total | total计数 | 超过token_budget被丢弃的样本数 |
两种规划器对应token_budget的两个 regime:>0走TokenBudgetBatcher(token 预算,动态样本数,默认);<=0走FixedCountBatcher(每 micro-batch 固定per_device_train_batch_size × num_processes个样本)。两者都做 Σ Lᵢ² 平衡,避免某个 rank 拖尾。
多教师 MOPD:路由规则与恢复语义
MOPD 不是该 trainer 核心目标所在论文(2306.13649,同步单教师场景)的一部分,而是独立方法:通用 SFT、各域独立 RL 专家训练、MOPD 融合三阶段中的第三阶段。该 trainer 只实现融合阶段——各域专家必须已存在(例如分别用GRPOTrainer/RLOOTrainer训练)并经 HTTP 提供服务,再把teacher_server_urls指向它们。注意论文自己的 Stage 3 用反向 KL(beta=1.0),而该 trainer 默认beta=0.0,配置 MOPD 时应显式设置。
- 单个条目:所有样本由该教师打分。
- 多个条目:每行的
teacher_id列选择打分者,如数学 prompt 给数学专家、代码 prompt 给代码专家。每个样本只发给匹配的那一个教师,绝不跨教师平均或集成。teacher_id缺失或未映射直接 raiseValueError,不静默回退。
约束:每个教师必须与学生共享同一个 tokenizer。原因:完成结果以原始 token id 传输,教师回报的候选 id 直接索引学生词表;词表不同的教师会把学生训练到错误的 token 上,除非其词表大于学生,否则这种错误是静默的。
检查点保存的是已训练位置而非生成器位置:首个未被训练的 prompt 索引随每个检查点写入rollout_state.json,恢复时 worker 直接快进到该位置。worker 领先训练一个队列深度,缓冲中未训练的样本在运行结束时丢失——若从生成器位置恢复,会跳过那些"已生成但未训练"的 prompt。流式数据集(IterableDataset)无法重新定位,其 worker 恢复时从 prompt 0 重启。
红线与常见坑
- 依赖装反:现象——transformers 被 vLLM 拉到不满足要求的版本。原因——两包依赖约束冲突。对策——先装
vllm>=0.22.0,再pip install 'transformers>=5.2.0' --no-deps。 - 用 DeepSpeed ZeRO 跑分布式:现象——配置不生效或行为异常。原因——该 trainer 分布式只支持 FSDP2。对策——换 FSDP2 配置。
teacher_id未映射:现象——启动即报ValueError。原因——拒绝静默路由到错误教师。对策——补齐数据集teacher_id列或对应teacher_server_urls条目。teacher_top_k超过 20 教师报错/截断:现象——候选数不足。原因——vLLM 默认每 token logprob 上限 20。对策——教师加--max-logprobs -1。- 超过
token_budget的样本消失:现象——batch/dropped_oversize_total增长、日志有警告。原因——该样本放不进任何行。对策——调大token_budget或缩短max_completion_length。 - 序列维并行直接抛错:现象——
cp_size>1或sp_size>1时__post_init__raise。原因——蒸馏在生成之后才于 trainer 内部构建模型输入,context-parallel/Ulysses 的输入分片无法应用于原始生成 batch。对策——cp_size=sp_size=1。 teacher_temperature不生效:现象——只有学生侧变化。原因——教师未带--logprobs-mode processed_logprobs,静默返回原始 logprobs。对策——按三终端章节的 flag 重启教师。max_staleness=0空队列挂死:现象——队列始终为空。原因——max_inflight_tasks自动公式含max_staleness因子,下限max(max_staleness, 1)是防呆,显式设 0 会清零调度。对策——不要显式压 0,或显式给max_inflight_tasks设正值。- 检查点恢复后数据"重复/跳段":现象——流式数据集恢复后从 prompt 0 重来。原因——
IterableDataset无法重新定位。对策——MOPD/流式场景预期行为;可重定位数据集按rollout_state.json快进。
边界与延伸阅读
该 trainer 刻意保持最小化,不打算成长为通用方案;需要不支持的功能时官方建议直接克隆仓库改造(仓库地址:git clone https://gitcode.com/GitHub_Trending/tr/trl)。有两个官方注入口,测试即通过注入 no-op 实现脱离真实 vLLM 运行:
RolloutWorkerProtocol:可替换的 rollout worker(生成+打分循环);WeightTransferProtocol:可替换的权重同步后端。
关键源码与示例(均为仓库相对路径):
trl/experimental/async_distillation/async_distillation_config.py(参数默认值与启动校验)trl/experimental/async_distillation/async_distillation_trainer.py(损失、指标规约、打包与检查点)trl/experimental/async_distillation/async_rollout_worker.py(spawn 子进程、生成打分循环、RolloutSample)trl/experimental/async_distillation/weight_transfer.py(NCCL 权重传输客户端)trl/experimental/async_distillation/vllm_client.py(vLLM HTTP 客户端)examples/async_distillation_math/async_distillation_math.py(单教师可运行示例:GSM8K、max_steps=100、learning_rate=1e-6)examples/async_distillation_math/async_distillation_mopd.py(双教师 MOPD 示例:数学/代码教师路由,显式beta=1.0)
【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考