TRL AsyncDistillationTrainer 深度解析:解耦生成与梯度更新、稀疏教师打分与指标诊断
2026/9/17 23:19:25 网站建设 项目流程

TRL AsyncDistillationTrainer 深度解析:解耦生成与梯度更新、稀疏教师打分与指标诊断

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

AsyncDistillationTrainer 是 TRL 实验模块中的异步 on-policy 蒸馏训练器:学生完成结果由后台 rollout worker 生成、由远端教师服务器打分,训练与生成并发。与同步的 DistillationTrainer 相比,核心差异只有一句话:教师永不被本地加载,只需一个 vLLM 服务器 URL,教师可运行在与学生、trainer 完全不同的硬件上。

术语速览

术语一句话定义直觉类比或最小示例
rollout一次 prompt 被学生生成一次、被教师打分一次一趟"生成+打分"往返,恰好产出一个训练样本
sample(RolloutSampleprompt + 学生完成结果 + 教师逐位置稀疏分布跨进程边界的唯一数据载体
row一个 DP rank 在一个 micro-batch 中前向的内容若干样本拼成的一条序列,position_ids逐样本重置
row-slot一个优化器步容纳的行数grad_accum × world_size,是校验 batch 指标的基准
staleness样本落后当前模型版本多少个权重更新数据的"年龄",超过max_staleness即丢弃
generated / forwarded / trained tokens学生生成的 / 前向处理的 / 损失实际计算的 tokentrained ⊆ generated;forwarded = prompt + generated
尾部桶(tail bucket)teacher_top_k之外追加的一个候选,承载剩余概率质量把"其他所有 token"显式变成一个选项
MOPD多教师 on-policy 蒸馏,该 trainer 只实现其融合阶段数据集teacher_id列决定哪个教师打分

架构与数据流

ROLLOUT WORKER(spawn 子进程,CUDA_VISIBLE_DEVICES 已清空) PROMPT(message list + 可选 teacher_id) ├─ generate 学生 vLLM /v1/completions(采样完成结果) └─ score 路由教师 vLLM /v1/completions(prompt_logprobs,teacher-forced) └─ RolloutSample(prompt + 完成结果 + 稀疏教师分布) ═══════ 进程边界:rollout_buffer(mp.Queue,maxsize = queue_maxsize)═══════ 主进程(训练循环,FSDP2/DDP) SAMPLE → staleness 检查(> max_staleness 丢弃) └─ Batcher(TokenBudget 默认 / FixedCount)→ ROW(每 DP rank 一行,Σ Lᵢ² 平衡) └─ MICRO-BATCH → PACKED ROW → FORWARD(_jsd_divergence,bs=1) └─ × grad_accum → OPTIMIZER STEP 每 weight_sync_steps 步:NCCL 权重传输 → 学生 vLLM(model_version + 1)

Rollout worker。一个 spawn 出的子进程,入口_child_main先调用_scrub_child_env清空CUDA_VISIBLE_DEVICES等环境变量——子进程没有资格碰 CUDA,任何惰性探测设备的库都会和父进程的显存分配器竞争。进程内运行 asyncio 事件循环_AsyncRolloutLoop:由于蒸馏没有分组基线,生成与打分合并在单个任务_generate_and_score_one中完成,并发度来自最多max_inflight_tasks个在途任务。它通过 OpenAI 兼容的/v1/completions与两类服务器通信:学生服务器用于采样真实完成结果;教师服务器用于 teacher-forced 打分(max_tokens=1prompt_logprobs=teacher_top_ktemperature=teacher_temperature,教师不生成任何新 token)。本地教师前向无法在这个无 CUDA 的子进程里跑,而 HTTP 打分让教师硬件与学生、trainer 完全解耦——这是与DistillationTrainer(教师本地加载,生成、教师前向、更新在同一进程顺序执行)的根本分叉。

rollout_buffermp.Queuemaxsize=queue_maxsize(默认 1024)。worker 每打出一个 sample 就put,队列满时阻塞(阻塞时长计入rollout/backpressure_s);trainer 侧的RolloutQueueDataset.__iter__以 5 秒轮询间隔get,队列为空则轮询并执行check_health_fn心跳检查(heartbeat_stale_after_s默认 300 秒,超时判定 worker 挂起并中止)。

Batcher(规划器)RolloutQueueDataset外面套两层:默认TokenBudgetBatchertoken_budget未设时取学生 vLLM 服务器的max_model_len,训练开始时刻一次查询),把样本按 Σ Lᵢ² 贪心装箱、每行不超预算;token_budget <= 0时换FixedCountBatcher,每 micro-batch 固定打包per_device_train_batch_size × num_processes个样本。装箱是贪心分桶(_balance_by_squared_length:按长度降序,放入当前 Σ Lᵢ² 最小的一行),目的是不让某个 rank 拖尾梯度 all-reduce——注意力成本是 O(L²),所以平衡的是平方和而非 token 数。

训练循环DataCollatorForRollout把每行拼成单条序列、稀疏教师候选按位置 padding 到teacher_top_k + 1宽(不足处以 id-1/ logprob-inf填充),行之间 padding 成矩形供 accelerate 分发,compute_loss前向前剥掉行间 padding。每weight_sync_steps步执行_sync_weight:pause vLLM → 全 rank barrier → 经 NCCL 流式传输可训练参数(FSDP2 下逐参数full_tensor()all-gather,避免整模型物化)→ resume,rank 0 的model_version自增并经共享mp.Value推给 worker。

不解耦会怎样:同步版本里 GPU 在生成阶段完全闲置,教师前向、生成、更新三段串行。解耦的代价是staleness——worker 领先训练最多一个队列深度,样本反映的是旧策略;max_staleness(默认 4)控制一个样本最多落后多少个权重更新,超过即丢弃(计入sample/dropped_stale_total)。另引入两项固定开销:队列内存(1024 个 sample 的缓冲)与每步的 NCCL 权重传输。

核心机制

优化什么

compute_loss最小化学生与教师在逐位置 token 分布上的广义 Jensen-Shannon 散度(generalized JSD)。选它而非策略梯度类目标,是因为蒸馏的信号是分布匹配而非标量奖励:教师在每个完成位置给出完整(稀疏化后的)分布,学生有梯度可用。beta是插值系数:0.0为前向 KL(mean-seeking,默认),1.0为反向 KL(mode-seeking),中间值线性插值。与DistillationTrainerServerDistillationTrainer使用同一目标函数,三者行为一致。

设某位置教师分布为P_T(稀疏,仅在候选集 C 上有定义,外加尾部桶)、学生分布为P_S(学生本地 logits 精确 softmax,非近似):

β=0: L = Σ_{c∈C} P_T(c) · (log P_T(c) − log P_S(c)) # 前向 KL β=1: L = Σ_{c∈C} P_S(c) · (log P_S(c) − log P_T(c)) # 反向 KL 0<β<1: M = (1−β)·P_S + β·P_T L = β·KL(P_T‖M) + (1−β)·KL(P_S‖M) loss = 对行内所有 valid trained token 的 L 求和(再按 token 数归一)

beta 与支撑集的行为对照

beta取值支撑集 C行为
0.0(默认)教师完整teacher_top_k宽支撑 + 尾部桶前向 KL;前向 KL 的权重恰好就是该支撑提供的分布,无信息损失
0 < beta < 1收窄为 2 个候选:教师 top-1 + 完成结果实际 token(去重后宽度 2)JSD 插值;混合分布的前向项需要教师 top-1,反向项只需要实际 token
1.0仅实际 token(宽度 1)纯反向 KL 是纯学生加权期望,教师 top-1 贡献为零,直接丢弃

收窄逻辑(_narrow_top1_actual_support)的动机:线上协议在不传输更宽(或完整)词表的前提下,保证教师 logprob 可用的只有教师 top-1 与实际 token 两个身份(vLLM 的prompt_logprobs总会报告实际 token,即便它落在 top-k 之外);任何更宽的支撑都只是概率性地覆盖学生可能采样的 token,而非保证。beta越界在AsyncDistillationConfig.__post_init__直接抛ValueError,合法域为[0.0, 1.0]

边界与隐含假设

  • teacher_top_k默认 8,是冒烟测试量级;超过 20 必须教师服务器以--max-logprobs -1启动。
  • add_tail_bucket=True(默认)时,_add_tail_bucket追加第 K+1 个元素log(1 − Σ exp(top_k_logps))logsumexp被 clamp 到 −1e−7 以下保证尾质量为正),避免候选集较小时散度平凡地趋零。
  • 教师未对某完成位置报告任何候选时(has_teacher_signal为假),该位置被token_mask_1d从损失中排除——否则经尾部桶会退化成"双方 100% 尾部"的伪造近零散度。该位置仍参与前向,所以 trained ≠ forwarded。
  • 归一化按全局 trained token 数:DDP/FSDP 对梯度求均值,compute_loss乘以world_size / global_n_tokens,再除以gradient_accumulation_steps;教师信号缺失的位置不参与计数,导致该窗口略微欠归一,被接受为教师侧数据缺口的代价。

实现注记:分块 lm_head 投影

DistillationTrainer相同,(chunk_size, vocab_size)的 logits 是唯一随词表规模扩展的张量。_chunked_jsd_loss把 backbone 输出的有效位置按_CHUNKED_LM_HEAD_CHUNK_SIZE = 256切块,每块在torch.utils.checkpoint下投影过lm_head,前向完成后丢弃、反向时重算——峰值 logits 内存是256 × vocab_size而非total_valid_tokens × vocab_size。与同步版的两点差异:只有一个模型需要投影(教师的稀疏候选已在线下算好),目标 ids 是稀疏候选集而非完整词表。FSDP2 下lm_head.weight是 DTensor,在分块前一次性full_tensor(),all-gather 只发生一次。

配置速查

模型加载

参数默认值说明何时需修改
model_init_kwargsNonefrom_pretrained关键字参数,revision同时用于加载 tokenizer模型需要特殊加载参数时
dtype"float32"学生加载精度,model_init_kwargs中的dtype优先学生 vLLM 服务 dtype 不一致导致 mismatch 时
trust_remote_codeFalse允许加载 Hub 自定义代码模型使用自定义代码仓库时

生成采样

参数默认值说明何时需修改
max_completion_length2048每完成结果最大生成 token 数长思维链或需要截断时
temperature1.0on-policy 采样温度探索与稳定性权衡
top_p1.0nucleus 采样参数同上
top_k0top-k 采样,0禁用同上
min_pNone最小 token 概率,按最可能 token 概率缩放,典型0.010.2抑制低概率 token
repetition_penalty1.0惩罚 prompt 与已生成文本中已出现 token重复退化时
chat_template_kwargsNone传给apply_chat_template的额外参数模板需要开关(如 think 模式)时

vLLM 服务器

参数默认值说明何时需修改
vllm_server_base_url"http://localhost:8000"学生服务器,用于生成与权重更新跨机部署时
vllm_server_timeout240.0等待学生服务器就绪的总超时(秒)大模型加载慢时
teacher_server_urls{"default": "http://localhost:8001"}教师服务器映射;多条目启用 MOPD,每行teacher_id选打分者多教师或跨机部署
request_timeout600单个 HTTP 请求超时(秒),对任意服务器长序列打分慢时
weight_sync_timeout1800权重传输超时(秒),超时 raise 而非挂死大模型传输慢时

蒸馏损失

参数默认值说明何时需修改
beta0.0广义 JSD 插值,0前向 KL、1反向 KLMOPD 按论文取1.0
teacher_temperature1.0散度 softmax 温度,作用于教师(服务端)与学生两侧软化/锐化教师分布
teacher_top_k8每位置请求的教师候选数,完整词表从不传输正式训练提到1664
add_tail_bucketTrue追加尾部桶,避免小候选集下散度趋零一般不改
token_budgetNone单行最大真实 token 数;None时取学生 vLLM 的max_model_len控制峰值内存与行填充率

异步流水线

参数默认值说明何时需修改
max_inflight_tasks-1在途生成+打分任务上限;-1自动取max(max_staleness, 1) × samples_per_step生成吞吐不足时
max_staleness4样本可落后当前版本的最大权重更新步数on-policy 性要求高时调小
queue_maxsize1024rollout 队列缓冲上限生成快于训练时
weight_sync_steps1两次权重同步之间的训练步数同步开销占比高时调大
heartbeat_stale_after_s300.0worker 心跳超时秒数,超过判挂起并中止一般不改

日志

参数默认值说明何时需修改
log_completionsFalse每 N 个已打分样本记录一批 (prompt, completion)需要人工抽检时
log_completions_steps100两次记录之间被打分的样本数;按 worker 打分计数,非优化器步配合上行
num_completions_to_printNone用 rich 打印的完成结果数,None全部日志刷屏时

⚠️ 与TrainingArguments默认值不同:logging_steps默认1(非500);gradient_checkpointing默认True(非False);bf16在未设置fp16时默认Truelearning_rate默认1e-6(非5e-5);ignore_data_skip默认True(非False,skip-and-replay 循环不适用于实时 rollout 队列,trainer 会强制置True)。

约束关系(__post_init____init__强制):

  1. beta必须在[0.0, 1.0],否则ValueError
  2. 序列维并行不支持:parallelism_configcp_size > 1sp_size > 1直接抛错——蒸馏在生成之后才于 trainer 内部构建模型输入,transformers 的 context/Ulysses 输入分片无法作用于原始生成 batch。
  3. teacher_server_urls至少一个条目(None时回填{"default": "http://localhost:8001"})。
  4. accelerator_config被强制为split_batches=Truedispatch_batches=True:主进程驱动 dataloader,batch 广播而非各进程独立拉取。
  5. 前向实现硬编码为 FlashAttention(kernels-community/flash-attn3),padding-free 模式依赖position_ids重置;use_liger_kernel=TrueNotImplementedError

部署与运行

最小训练脚本(完整可运行示例见 examples/async_distillation_math/async_distillation_math.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()

⚠️ 环境要求:vllm>=0.22.0transformers>=5.2.0;分布式训练仅支持 FSDP2(不支持 DeepSpeed ZeRO)。两者当前存在冲突的依赖约束,先装 vLLM 再强制装 transformers:

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

三终端部署,教师、学生 vLLM、trainer 必须在不同 GPU 上:

  1. GPU 0 启动教师服务器(静态,永不更新,无需 dev 模式):
CUDA_VISIBLE_DEVICES=0 vllm serve Qwen/Qwen2.5-1.5B-Instruct \ --port 8001 \ --logprobs-mode processed_logprobs \ --max-logprobs -1

--logprobs-mode processed_logprobs使teacher_temperature作用于返回的 logprobs(否则教师静默报告原始 logprobs,该设置只影响学生侧);--max-logprobs -1解除 vLLM 默认 20 的 per-token logprob 上限,使teacher_top_k可超过 20。

  1. GPU 1 启动学生 vLLM 服务器(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"}'
  1. GPU 2 启动训练:
CUDA_VISIBLE_DEVICES=2 accelerate launch train_async_distillation.py

常见启动失败排查:

现象最可能原因修复动作
卡在等待学生服务器,vllm_server_timeout超时学生服务器未启动或 model id 不一致核对vllm serve启动参数与model字符串
权重传输超时异常(weight_sync_timeout学生服务器缺VLLM_SERVER_DEV_MODE=1或 NCCL 后端未启用补环境变量与--weight-transfer-config
teacher_top_k设到 20 以上报错教师未带--max-logprobs -1重启教师服务器并加该参数
sample/dropped_stale_total持续增长生成慢于训练且max_staleness过紧调大max_inflight_tasks/queue_maxsize,或降weight_sync_steps
jsdentropy整窗口为 NaN窗口内无任何可用教师信号,常见于 tokenizer 不共享确认教师与学生同词表(见下节警告)

观测与调优

指标分两类:吞吐/延迟类回答"快不快",学习信号类回答"学没学到"。perf/双后缀指标基于同一次优化器步,仅分母不同:_fwd_bwd除以perf/fwd_bwd_s(纯计算,衡量 trainer 效率),_wall_clock除以perf/step_s(含队列等待,衡量算力利用率)。只看前者会掩盖生成侧的 GPU 时数,只看后者会把教师延迟算到 trainer 头上。

生成侧是否瓶颈(generation-bound)

你会看到训练在挨饿:队列接近空、perf/rollout_wait_s高。

指标回答的子问题异常方向指向的根因
perf/rollout_wait_strainer 因队列空阻塞了多久持续走高生成侧产速不足
sample/rollout_queue_size当前等待样本数贴近 0同上,与上行互为印证
rollout/generated_tok_s窗口内生成吞吐低或出现平台学生 vLLM 服务器产能不足
rollout/score_s(MOPD 看teacher_score_s/<id>教师调用占 rollout 的时间教师慢;MOPD 下只有被路由的 rollout 受影响
rollout/vllm_retry_total重试过的 vLLM 请求数增长服务器退化,否则表现为莫名的变慢

训练侧是否瓶颈(trainer-bound)

你会看到生成被节流、产物在队列中老化:队列接近满、rollout/backpressure_s高,同时sample/staleness_mean攀升。

指标回答的子问题异常方向指向的根因
rollout/backpressure_s生成因队列满阻塞了多久持续走高训练消费慢
sample/rollout_queue_size当前等待样本数贴近queue_maxsize同上
sample/staleness_mean数据落后当前版本多少步逼近max_stalenessoff-policy 性积累,丢弃风险上升
batch/row_imbalance各行 Σ Lᵢ² 的 max/mean远离 1.0某 rank 拖尾 all-reduce
batch/row_fill_frac行 token 数相对token_budget长期偏低长度量化效应(1 万 token 样本铺不满 3.2 万预算),调token_budget

perf/rollout_wait_srollout/backpressure_s是镜像,不会同时很大——二者读数直接告诉你瓶颈在哪一侧。

目标函数是否真的收敛

你会看到jsd下降但entropy同步崩塌:学生在收窄而非学习。

指标回答的子问题异常方向指向的根因
jsd广义 JSD 是否在按beta收敛平台期或回升学习停滞或 off-policy 干扰(对照 staleness)
entropy学生自身预测熵jsd下降而崩塌模式坍缩而非学习
teacher_entropy教师在所报候选上的熵显著低于直觉值teacher_top_k从下方截断,属正常下界
batch/masked_token_frac前向 token 中不产生梯度的占比prompt 占比大或教师信号缺口多

多教师路由是否偏斜(MOPD 专属)

一个被饿死的教师仍会报告健康的teacher_jsd/<id>,混合jsd会把偏斜藏住。

指标回答的子问题异常方向指向的根因
teacher_token_frac/<id>该教师打分的 token 占比某教师逼近 0teacher_id路由偏斜
teacher_jsd/<id>该教师 token 上的jsd某教师远高对应领域收敛慢
teacher_score_s/<id>该教师打分耗时单教师高只有其路由的 rollout 被拖慢

无 per-teacher 的entropy:学生熵是其自身策略的属性,与谁打分无关。

分配到的算力有多少真正变成了训练

指标回答的子问题异常方向指向的根因
perf/mfu_wall_clockvsperf/mfu_fwd_bwd两口径 MFU 之差差距大时间在队列等待,非计算
perf/weight_sync_s(含_pause_s/_barrier_s/_transfer_s一次完整同步的三阶段耗时_barrier_srank 偏斜
perf/fwd_s / perf/fwd_bwd_s前向占比高于 1/3反向便宜或重计算发生在前向

30 秒体检:只看四个数——jsd(学没学到)、sample/rollout_queue_size+perf/rollout_wait_svsrollout/backpressure_s(瓶颈在哪侧)、perf/mfu_wall_clock(算力利用率)、sample/staleness_mean(off-policy 积累)。四个数健康则系统平衡,再下钻到对应小节。

高级用法与边界

MOPD:多教师路由蒸馏

  • 前提:各领域的专家教师必须已存在(例如分别用GRPOTrainer/RLOOTrainer训练)并经 HTTP 服务;每个教师必须与学生共享 tokenizer——完成结果以原始 token id 传输,教师报告的候选 id 直接索引学生词表,词表不同的教师会把学生训练到错误的 token 上,且这种错误是静默的(除非教师词表比学生大)。
  • 关键 diff:teacher_server_urls多条目(如{"math": ..., "code": ...});数据集每行携带teacher_id列,缺失或未映射直接报错而非回退;每个样本只分发给其匹配的一个教师,绝不跨教师平均或集成。论文自身 Stage 3 使用反向 KL,需显式beta=1.0(trainer 默认0.0)。
  • 可运行示例:examples/async_distillation_math/async_distillation_mopd.py(数学 GSM8K 路由 math 教师、代码路由 code 教师,学生为 Qwen2.5-0.5B-Instruct)。

检查点与断点恢复

每个检查点随写rollout_state.json{"prompt_index": ...}),保存的是已训练位置而非生成器位置——worker 领先队列深度,已缓冲未训练的样本在运行结束即丢失,从生成器位置恢复会跳过"已生成但未训练"的 prompt。IterableDatasetlen(),恢复时 worker 从 prompt 0 重启。

本模块不做什么

  • 不支持本地(进程内 GPU)教师前向:无 CUDA 的子进程跑不了,回主进程的路径未实现。
  • 不支持序列维并行(cp_size/sp_size> 1 抛错)。
  • 不支持use_liger_kernel(抛NotImplementedError)。
  • 不做跨教师集成/平均;没有奖励函数与分组基线(区别于 GRPO)。
  • 分布式仅 FSDP2,DeepSpeed ZeRO 不支持。

扩展点

  • RolloutWorkerProtocol:需暴露rollout_buffermetrics_queue两个队列属性,并实现start/stop/update_model_version/check_health。替换后 trainer 不再自建AsyncRolloutWorker,队列归 worker 所有。
  • WeightTransferProtocol:实现init_weight_transfer/pause/send_weights/resume/destroy。传入 no-op 实现即可禁用 trainer 侧权重同步(测试即如此注入,脱离真实 vLLM 服务器运行)。
  • 官方态度:该 trainer 刻意保持最小化,不打算成长为通用解决方案;需要不支持的功能时,官方建议直接克隆仓库(git clone https://gitcode.com/GitHub_Trending/tr/trl)并按需改造,新功能只在出现显著社区需求时考虑。

延伸阅读

源码(相对仓库根目录):

  • trl/experimental/async_distillation/async_distillation_trainer.py:损失计算(_jsd_divergence_chunked_jsd_loss)、两种 Batcher、DataCollatorForRollout、权重同步与指标聚合
  • trl/experimental/async_distillation/async_distillation_config.py:全部专有参数与__post_init__约束
  • trl/experimental/async_distillation/async_rollout_worker.py:_AsyncRolloutLoop生成+打分循环、RolloutSample、子进程环境清理
  • trl/experimental/async_distillation/weight_transfer.py:NCCL 权重传输客户端
  • trl/experimental/async_distillation/vllm_client.py:vLLM HTTP 客户端(就绪等待、max_model_len查询)

可运行示例:

  • examples/async_distillation_math/async_distillation_math.py:单教师 GSM8K
  • examples/async_distillation_math/async_distillation_mopd.py:双教师 MOPD

论文:

  • 核心目标(同步单教师 on-policy 蒸馏):arXiv:2306.13649
  • MOPD(多教师能力融合,仅其 Stage 3 融合阶段由本 trainer 实现):arXiv:2606.30406

同项目关联模块:

  • trl/experimental/distillation/:DistillationTrainer,教师本地加载的同步版本,同一 JSD 目标
  • trl/experimental/server_distillation/:ServerDistillationTrainerbeta支撑集收窄逻辑的镜像来源
  • trl/experimental/async_grpo/:AsyncGRPOTrainer,本 trainer 的架构原型,planner/worker 机制逐行移植自它

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

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

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

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

立即咨询