TRL 异步蒸馏完整配置指南:AsyncDistillationTrainer 从原理到调参
【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl
🏭 从"教师不能同时在场"说起
想象这样一个场景:你的学生是 0.5B 的小模型,教师是 15B 的大模型,差了十几倍。同步蒸馏的套路是"生成、教师前向、梯度更新"在同一个进程里顺序执行,意味着学生与教师必须挤在同一组显卡上——15B 权重 fp16 下光权重就要 30GB,再叠上训练侧的激活值,OOM 只是时间问题;就算硬塞得下,两套模型的前向也在抢同一块显存的带宽。
TRL 给出的答案就是AsyncDistillationTrainer(实现见 async_distillation_trainer.py)。它做的仍是 on-policy 蒸馏——每一个训练样本都来自学生自己的采样——但教师从此不占用训练卡:它作为一个独立的 vLLM 服务器存在,只对 HTTP 打分请求做出响应。学生生成答案、教师远程复核、trainer 更新权重,三条线各干各的,互不阻塞。
🧭 一条样本的旅程:从 Prompt 到梯度
把整条链路拆成"生成侧"和"训练侧"两个区,中间隔着一条进程边界:
生成侧(rollout worker 子进程,无 GPU) ① 取 prompt ──▶ ② 学生起草 ──▶ ③ 教师复核 │(学生 vLLM 采样生成) │(teacher-forced 打分) ▼ rollout_buffer(跨进程队列) ▼ 训练侧(trainer 主进程) ④ 新鲜度检查 ──▶ ⑤ 分行 ──▶ ⑥ 打包 ──▶ ⑦ 前向 + JSD 损失 ──▶ ⑧ 优化器步 ▲ ⑧ 每 weight_sync_steps 步 ──NCCL 推新权重──▶ 学生 vLLM ─┘(闭环)- ① 取 prompt:worker 子进程从数据集取下一行(一行 = 消息列表,外加可选的
teacher_id)。 - ② 学生起草:调学生 vLLM 服务器的
/v1/completions采样生成一份答案。 - ③ 教师复核:把答案原样发回路由到的教师服务器打分——就像拿着标准答案逐字对照,老师只给每个位置报 logprob,自己不动笔写新字。产物是
RolloutSample:prompt、答案、每个完成位置的 top-teacher_top_k候选 token 及其 logprob。①→③ 是一个在途任务,多个任务按max_inflight_tasks并发。 - ④ 新鲜度检查:trainer 逐个拉取样本,对比样本生成时的模型版本与当前版本,落后超过
max_staleness直接丢弃(sample/dropped_stale_total计数)。 - ⑤ 分行:规划器把样本分给各 DP rank,贪心按 Σ Lᵢ² 平衡,避免某张卡拖慢全体。
- ⑥ 打包:
DataCollatorForRollout把一行的样本拼成一条长序列,样本边界处重置position_ids。 - ⑦ 前向:
compute_loss按 256 个 token 分块投影lm_head(峰值 logits 内存因此只是256 × vocab_size),计算广义 JSD。 - ⑧ 优化器步:
gradient_accumulation_steps个 micro-batch 凑成一步;每weight_sync_steps步,新权重经 NCCL 流进学生 vLLM 服务器,生成侧因此能跟上——这就是回环箭头。
⚖️ 损失与 β 参数选择:支撑集为什么会被收窄
先看约束。教师住在另一台机器上,完整词表根本过不了 HTTP。每个完成位置线上能拿到的只有:教师的 top-teacher_top_k候选 logprob、实际实现 token 的 logprob(vLLM 保证它即使不在 top-k 里也会返回)、以及可选的尾部桶(add_tail_bucket=True,把剩余概率质量收进一个元素,防止候选太少时散度"看起来恒等于零")。学生那一侧是精确的——它就是要训练的模型,完整 logits 本地就有。
在这个约束下,beta决定了散度在哪些候选上计算:
beta=0.0(前向 KL):期望以教师分布加权,只关心"教师把概率放在哪",而教师的 top-k 切片恰好就是这份信息——所以保留全部teacher_top_k宽度支撑。beta≠0.0:混合项里出现了以学生分布加权的分量,重心落在学生自己采到的 token 上。这个 token 未必在教师 top-k 内,而线上协议能保证拿到教师 logprob 的只有两个身份:教师 top-1 与答案实际 token。于是_narrow_top1_actual_support把支撑收窄到 2 个候选(两者相同则去重)。beta=1.0(反向 KL):纯学生加权,教师 top-1 不贡献任何项,支撑进一步缩到 1 个——只剩实际 token。
一句话实操建议:默认beta=0.0适合"全面继承教师分布";若目标是让学生模仿教师的主行为(比如融合 RL 专家,MOPD 论文第三阶段用的就是反向 KL),显式写beta=1.0。注意中间值下teacher_top_k对支撑宽度已经无效,它只影响教师报告的是哪个 top-1。
config = AsyncDistillationConfig( beta=1.0, # 反向 KL,mode-seeking teacher_top_k=16, # 线上每位置的候选数 teacher_server_urls={"math": "http://localhost:8001"}, )🖥️ 三终端部署:教师、学生、训练各占一张卡
踩坑提醒:当前 vLLM 与 transformers 的依赖约束互相冲突,装反了会直接 import 失败。正确顺序是先装 vLLM,再"裸装"transformers(跳过依赖解析):
pip install 'vllm>=0.22.0' pip install 'transformers>=5.2.0' --no-deps另外,分布式训练只支持 FSDP2,DeepSpeed ZeRO 不在支持列表里。
终端 1——教师服务器:它的角色最轻,只接打分请求,不生成新文本、永不更新,所以什么 dev 开关都不需要,只要两个"打分精确性"flag:
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真正作用于返回的 logprob(漏掉它,该参数只影响学生侧);--max-logprobs -1解除 vLLM 默认的 20 上限,teacher_top_k想超过 20 必须开它。
终端 2——学生的 vLLM 服务器:它既负责起草,又要接收 trainer 推来的新权重,所以必须进 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——训练进程:脚本很短,完整可跑版本在 examples/async_distillation_math/async_distillation_math.py:
trainer = AsyncDistillationTrainer( model="Qwen/Qwen2.5-0.5B-Instruct", args=AsyncDistillationConfig( teacher_server_urls={"default": "http://localhost:8001"}, vllm_server_base_url="http://localhost:8000"}, train_dataset=dataset, ) trainer.train()CUDA_VISIBLE_DEVICES=2 accelerate launch examples/async_distillation_math/async_distillation_math.py新手最容易踩的 8 个参数
| 参数 | 默认值 | 为什么容易踩 |
|---|---|---|
teacher_server_urls | {"default": "http://localhost:8001"} | 端口要和终端 1 对上;换成多个条目就进入 MOPD 路由模式 |
beta | 0.0 | 必须在[0, 1],越界直接ValueError;取 0 和取非 0 时支撑宽度完全不同(见上一节) |
teacher_top_k | 8 | 8 只是冒烟测试级别,正式训练建议 16–64;>20 要求教师带--max-logprobs -1 |
teacher_temperature | 1.0 | 同时作用于教师侧(服务端算 logprob)与学生 logits;若教师没开processed_logprobs,它对教师侧静默失效 |
token_budget | None(取学生服务器max_model_len) | 样本超出预算就进不了任何行,被丢弃并计入batch/dropped_oversize_total;显存紧张时它是第一调节杠杆 |
max_staleness | 4 | 样本可落后当前策略的权重版本数;太小队列常空,太大 off-policy 味太重 |
weight_sync_steps | 1 | 每 N 个优化器步推一次权重给学生 vLLM;调大省同步开销,但生成用更旧策略 |
dtype | "float32" | 默认 fp32 而非 bf16(为对齐训练-推理精度度量);要端到端一致,学生 vLLM 也要用相同--dtype起服务 |
其余字段(采样参数、超时、心跳、日志开关)见 async_distillation_config.py。再留个心眼:learning_rate默认1e-6(不是5e-5)、logging_steps默认1、gradient_checkpointing默认True、bf16默认开——带着旧习惯配参数会翻车。
🔀 MOPD 多教师路由:各管一摊,不搞集成
teacher_server_urls里放多个条目后,路由规则只有一条:每个样本由数据里的teacher_id列挑中唯一一位教师打分——数学 prompt 走数学教师,代码 prompt 走代码教师。没有跨教师平均,没有 ensemble,是"分诊"而不是"会诊"。teacher_id缺失或不在映射里时,_resolve_teacher_server_url直接抛ValueError,不会悄悄回退到别的教师。
有一条硬约束:教师必须和学生共用 tokenizer。答案以原始 token id 发给教师,教师回传的候选 id 会在compute_loss里直接索引学生自己的词表——词表对不上就是在给错误的 token 做训练,而且只要教师的词表不比学生小,这个错误是完全静默的。同家族组合(Qwen2.5 学生 + Qwen2.5 教师 + Qwen2.5-Coder 教师)天然满足,参考 async_distillation_mopd.py 的双教师配置。
最后注意范围:MOPD 只接"融合阶段"——各领域专家必须已单独训练好并通过 HTTP 服务,这个 trainer 不负责把它们练出来。
🩺 指标瓶颈定位:哪根柱子在拖后腿
诊断顺序固定为"先问后答"的三步:
第一问:谁在等谁?看镜像指标对perf/rollout_wait_s与rollout/backpressure_s——训练侧因队列空停下的总时长,生成侧因队列满被压回的总时长,两者不会同时很高。
第二问:行装满了吗?看batch/row_fill_frac。1 万 token 的样本塞 3.2 万 token 的预算,3 个放得下、4 个永远放不下,打包器经常只能塞 2 个——填充率偏低时token_budget是第一调节杠杆,顺带查batch/dropped_oversize_total有没有超预算丢弃。
第三问:学生在学习还是收缩?看jsd与entropy的配对走势:jsd降、entropy稳,是正常收敛;jsd降的同时entropy塌方,说明学生只在高置信 token 上收缩分布,没学到东西。
| 指标 | 一句话含义 | 异常时接着看 |
|---|---|---|
sample/rollout_queue_size | 队列里躺着多少份已打分样本 | 与 wait/backpressure 配对判断瓶颈在哪侧 |
sample/time_in_queue_s | 单个样本从打分完到进训练等了多久(off-policy 的"秒数"部分) | sample/staleness_mean、sample/dropped_stale_total |
perf/rollout_wait_s | 训练因队列空而停摆的累计时长 | rollout/score_s(教师慢?)、rollout/generated_tok_s(生成吞吐?)、rollout/inflight |
rollout/backpressure_s | 生成因队列满被压回的累计时长 | sample/staleness_mean(是否在队列里变老)、batch/row_tokens_mean |
rollout/vllm_retry_total | 对两台 vLLM 的 HTTP 重试次数,持续上涨说明有台服务器在退化 | rollout/duration_s |
batch/row_fill_frac | 行的实际 token 数占token_budget的比例,低即预算没吃满 | batch/dropped_oversize_total |
batch/row_imbalance | 各行 Σ Lᵢ² 的最大/均值,1.0 完美,偏高说明有 rank 在拖 all-reduce | batch/row_tokens_max |
jsd/entropy | 损失本体 / 学生自身预测熵,一降一稳=收敛,一降一塌=收缩 | sample/staleness_mean、completions/clipped_ratio |
性能侧只需盯两个:perf/step_s(一步优化器步的墙钟时间,含一切等待)与perf/fwd_bwd_s(其中纯前反向部分)。MFU 的两个口径_fwd_bwd与_wall_clock差只在分母——前者回答"有数据时训练侧多高效",后者回答"分到的算力有多少真变成了训练",后者远小于前者就去找生成侧。MOPD 下还有teacher_jsd/<id>、teacher_token_frac/<id>按教师拆分,路由偏斜在混合jsd里是看不出来的。
下一次跑训练,别盯着 loss 发呆:先翻队列和背压,定位瓶颈在哪一侧;再看填充率,确认预算匹配样本长度;最后用jsd和entropy的配对走势确认学生真在学习。这套诊断路径,全在上面的表里。
【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考