如何配置 TRL AsyncDistillationTrainer:三终端部署、beta 调参与日志排障一次讲清
2026/9/17 4:34:52 网站建设 项目流程

如何配置 TRL AsyncDistillationTrainer:三终端部署、beta 调参与日志排障一次讲清

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

TRL 的AsyncDistillationTrainer把蒸馏中的生成、教师打分与梯度更新拆成并发执行:教师模型永不在本地加载,只需一个 vLLM 服务器 URL。本文带你部署三终端拓扑、选对betateacher_top_k,并用日志指标定位生成侧或训练侧的瓶颈。

读完本文你会掌握三件事:

  • 如何用三个终端把教师服务器、学生 vLLM 服务器与训练进程跑在独立 GPU 上;
  • 如何按训练阶段选择betateacher_top_kadd_tail_bucket,理解支撑集收窄的原因;
  • 如何用队列类与性能类指标区分生成受限、训练受限与服务退化三类瓶颈。

环境要求:需要vllm>=0.22.0transformers>=5.2.0;分布式训练仅支持 FSDP2,不支持 DeepSpeed ZeRO。由于 vLLM 与 transformers 当前依赖约束冲突,必须先装 vLLM、再强制安装 transformers:先执行pip install 'vllm>=0.22.0',再执行pip install 'transformers>=5.2.0' --no-deps

🏗️ 架构全景:三个角色的并发拓扑

同步蒸馏里,教师是本地模型,生成、教师前向、梯度更新在同一进程内排队执行——学生模型稍大、教师更大时,两块负载往往挤不下一台机器;就算挤得下,GPU 也在"等生成"与"等反向"之间反复空转。AsyncDistillationTrainer的做法是把三者拆到三个角色上,各自占各自的卡:

  • 打分端(教师 vLLM 服务器):纯静态,权重永不更新,因此不需要 dev 模式,也不需要权重传输后端。rollout worker 把学生生成的完整序列发到教师的/v1/completions,用prompt_logprobs做 teacher-forced 打分——教师只回每个位置 top-teacher_top_k的稀疏候选 logprob,不生成任何新 token。完整词表从不经 HTTP 传输。
  • 生成端(学生 vLLM 服务器):只负责按当前权重采样学生的 on-policy 完成结果。它开启了 dev 模式与 NCCL 权重传输,trainer 每weight_sync_steps个训练步把更新后的学生权重推送进来。
  • 训练端(trainer 主进程):内部又分两个环节。后台 rollout worker 是一个清除了 CUDA 设备的 spawn 子进程,跑 asyncio 事件循环,负责"生成 + 打分"后把可训练样本推入进程间队列rollout_buffer;训练循环则不断从队列拉样本、计算广义 JSD 损失、更新权重。

谁通过什么协议调用谁,一句话版本:

  • worker → 学生 vLLM:HTTP/v1/completions,带采样参数;
  • worker → 教师 vLLM:HTTP/v1/completions,带prompt_logprobs=teacher_top_ktemperature=teacher_temperature
  • trainer → 学生 vLLM:NCCL 权重流,每weight_sync_steps步一次;
  • worker ↔ trainer:mp.Queue,容量上限queue_maxsize(默认 1024)。

由于生成始终领先训练,样本可能反映略微过期的策略:每个样本最多落后max_staleness(默认 4)个权重更新,超过即丢弃,丢弃数计入sample/dropped_stale_total

⚡ 三终端部署快速上手

三个角色必须跑在不同的 GPU上。下面的最小脚本只依赖默认配置(teacher_server_urls缺省指向http://localhost:8001,学生服务器缺省指向http://localhost:8000),完整参数版可对照仓库示例examples/async_distillation_math/async_distillation_math.py(GSM8K、max_steps=100learning_rate=1e-6report_to="trackio")。

# 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()

教师服务器(GPU 0)。两个参数都不可省:--logprobs-mode processed_logprobsteacher_temperature在服务端作用于返回的 logprobs(否则教师静默返回原始 logprobs,温度只影响学生侧);--max-logprobs -1解除 vLLM 每 token 20 个 logprob 的上限,teacher_top_k才能超过 20:

# 终端 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

学生 vLLM 服务器(GPU 1)。VLLM_SERVER_DEV_MODE=1与 NCCL 权重传输缺一不可,否则 trainer 无法把新权重推进去:

# 终端 2: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"}'

训练进程(GPU 2):

# 终端 3:GPU 2 —— trainer 主进程 CUDA_VISIBLE_DEVICES=2 accelerate launch train_async_distillation.py

注意该 trainer 的默认值刻意区别于 transformers 的TrainingArgumentslearning_rate1e-6(非5e-5)、bf16在未设fp16时默认Truegradient_checkpointing默认Truelogging_steps默认1ignore_data_skip默认True

🎛️ 损失与调参:beta、teacher_top_k 怎么选

损失最小化的是学生与教师在逐位置 token 分布上的广义 Jensen-Shannon 散度,与DistillationTrainerServerDistillationTrainer使用同一目标。beta(必须落在[0.0, 1.0],越界直接ValueError)是插值系数:beta=0.0为前向 KL,mean-seeking,权重落在教师概率高的区域,是默认值;beta=1.0为反向 KL,mode-seeking,权重落在学生自己采样的区域;中间值平滑过渡。

调参时最关键的一点是:支撑集(计算散度用的候选集合)随 beta 变化。原因不在数学,而在传输协议——线上传的只有教师 top-k 切片,不是完整词表:

  • beta=0.0:用教师报告的完整teacher_top_k宽支撑(外加尾桶)。前向 KL 恰好只需要教师分布给出的权重,这份支撑正好够用;
  • beta != 0.0:支撑收窄为两个候选——教师自己的 top-1 token 和完成结果的实际 token。因为这是协议在不传更宽词表的前提下唯一能保证拿到教师 logprob 的两个身份;任何更宽的覆盖都只是概率性近似,不是保证。其中beta=1.0时支撑宽度进一步缩到 1(纯反向 KL 是纯学生加权期望,教师 top-1 不贡献)。

teacher_top_k(默认8)控制每位置从教师索取的候选数,vLLM 总会额外报告 realized token(即便在 top-k 之外)。尾桶add_tail_bucket(默认True)控制:在 top-k 候选之外追加一个表示剩余概率质量的元素,取值为log(1 - sum(exp(top_k_logps))),避免候选集过小时散度平凡地趋近于零。选型建议:8适合冒烟测试;真正训练时提到1664合理(邻近 RL 框架的参考默认是 miles 用16、EasyOPD 用64)。

另外两个容易混淆的点:

  • teacher_temperature(默认1.0)作用在散度两侧——随请求发给教师让 vLLM 在服务端计算,同时在compute_loss中作用于学生 logits。它与采样用的temperature完全无关;
  • 实现上是分块的:只有(chunk_size, vocab_size)的 logits 张量随词表扩展,chunk 前向后用torch.utils.checkpoint丢弃并在反向重算,峰值 logits 内存约等于单个 chunk(常量256)乘词表大小,而非全部有效 token 乘词表。

🧭 多教师路由(MOPD)

MOPD 是独立方法,不属于本 trainer 核心目标所依据的单教师论文。它的完整流程三阶段:通用 SFT → 每个领域独立做基于 RL 的专家训练 → 用 MOPD 把冻结的专家融合进一个学生。AsyncDistillationTrainer只实现第三阶段(融合阶段):各领域的专家必须已经存在(例如用GRPOTrainer/RLOOTrainer单独训好)、已经以 HTTP 提供服务,你再把teacher_server_urls指向它们。MOPD 论文自己的 Stage 3 用的是反向 KL,配置时应显式写beta=1.0,而不是沿用默认的beta=0.0

路由规则与约束:

  • teacher_server_urls单个条目:所有样本都由该教师打分;多个条目:每行的teacher_id列决定谁打分,例如数学 prompt 走数学专家、代码 prompt 走代码专家,各自独立服务;
  • 每个样本只分发给它匹配的那一个教师,绝不跨教师求平均或集成
  • teacher_id缺失或映射不到任何条目的样本会直接抛ValueError,不存在"静默回退到默认教师"。

可运行的双教师示例在examples/async_distillation_math/async_distillation_mopd.py:GSM8K 路由给Qwen/Qwen2.5-1.5B-Instruct(math),Python 代码指令数据路由给Qwen/Qwen2.5-Coder-1.5B-Instruct(code),学生是Qwen/Qwen2.5-0.5B-Instruct,显式beta=1.0

⚠️每个教师必须与学生共享 tokenizer。完成结果以原始 token id 发给教师,教师回报的候选 id 在compute_loss里直接索引学生自己的词表。同模型家族的教师(如示例中的 Qwen2.5 学生由 Qwen2.5 与 Qwen2.5-Coder 专家融合)满足条件;词表不同的教师会把学生训到错误的 token 上,而且只要它的词表不比学生大,这个错误是静默的。

🩺 用日志定位瓶颈

排障从队列状态入手。四个指标描述同一个队列,其中两个是镜像、永远不会同时偏大:sample/rollout_queue_size告诉你现在有多少样本在等;sample/time_in_queue_s单个样本入队后等了多久(off-policy 性的"秒数"部分);perf/rollout_wait_s是训练端因队列为空阻塞的时长;rollout/backpressure_s是生成端因队列已满阻塞的时长。

诊断流程按"现象 → 疑似原因 → 看哪个指标"走:

现象疑似原因关键指标
队列接近空,perf/rollout_wait_s生成受限(generation-bound),训练在挨饿rollout/generated_tok_s(窗口吞吐,停顿会显现)、rollout/inflightrollout/score_s
队列接近满,rollout/backpressure_s训练受限(trainer-bound),产出在队列中老化sample/staleness_mean是否持续攀升、sample/dropped_stale_total
两者都接近零平衡,无需动作可转看batch/row_fill_frac
吞吐莫名下降某台 vLLM 服务器退化rollout/vllm_retry_total(学生或教师被重试的请求数;它统计的是对服务器的请求而非生成文本,所以归在rollout/下)
教师慢只拖累部分 rolloutMOPD 下某个专家过慢,混合均值会掩盖它teacher_score_s/<id>teacher_jsd/<id>
某个教师"看起来健康"却不见进展路由偏斜,该教师被饿死,却仍报告健康的散度teacher_token_frac/<id>(该教师分到的打分 token 占比)
行打包不紧、长样本多量化效应而非 bug:1 万 token 的样本难铺满 3.2 万 token 的预算,3 个放得下、4 个放不下,打包器常只能放 2 个batch/row_fill_frac,调节token_budget

吞吐与 MFU 各有两份口径,基于同一个优化器步,差别只在除数:*_fwd_bwd除以perf/fwd_bwd_s(纯计算,回答"有数据时 trainer 跑得多高效",偏低说明问题在 trainer);*_wall_clock除以perf/step_s(完整一步含排队等待,回答"分配到的算力有多少真正变成训练")。两者之差约等于perf/rollout_wait_s加优化器与权重同步时间(perf/weight_sync_s还细分_pause_s_barrier_s_transfer_s)。只看前者会掩盖花在生成与打分上的 GPU 时数,只引用后者则可能把生成器或教师的延迟算到 trainer 头上——两个都要看。

数据流水线与检查点恢复

一条数据从 prompt 变成一次梯度,链路是:rollout → sample → row → micro-batch → 优化器步

  • rollout:一个 prompt 生成一次、打分一次。蒸馏没有可跨生成计算的 advantage 基线,所以没有 group、prompt 不重复,一次 rollout 恰好产出一个训练样本;
  • sample:完成结果加上教师对每个完成位置给出的 top-teacher_top_k候选——这是唯一跨进程边界(rollout_buffer)传输的内容,拉取时按max_staleness丢弃过期样本;
  • row:规划器把样本分给 DP rank(按 Σ Lᵢ² 贪心分桶,避免某个 rank 拖尾),一行是若干样本拼成的单条序列,position_ids在每个样本边界重置;候选不足teacher_top_k+1宽的位置以 id-1/ logprob-inf填充,损失中掩掉;
  • micro-batch:每个 DP rank 一行,共world_size行;gradient_accumulation_steps个 micro-batch 累积成一次优化器步。一步覆盖的 row 槽位数恒等于gradient_accumulation_steps × world_size,因此batch/samples_per_step ≈ row 槽位数 × batch/samples_per_row(per-step 是求和、per-row 是均值,预期有零点几百分比的偏差而非精确相等)。

所有指标里 "step" 都指完整优化器步而非 micro-batch;micro-batch 级量会明说(batch/microbatches_per_step)或以 per-row 形式给出(batch/row_*batch/samples_per_row)。token 口径分三种:generated是学生实际生成的 token(completions/*);forwarded是前向处理过的全部 token(prompt + 生成);trainedcompletion_mask == 1覆盖、损失真正计算的子集。trained ≠ generated:教师没给某位置打分任何候选时,该位置在散度中被掩掉但仍参与前向,这部分的占比可以看batch/masked_token_frac

检查点恢复走的是另一套逻辑:ignore_data_skip默认True,基础 Trainer 的 skip-and-replay 不适用于实时 rollout 队列。每个检查点会往rollout_state.json写入第一个尚未被训练的 prompt 索引{"prompt_index": ...}),恢复时 worker 直接快进到该位置,无需重放。存的是已训练位置而非生成器位置——worker 领先训练最多一个队列深度,缓冲里已生成但未训练的样本在运行结束时即丢失;若从生成器位置恢复,就会跳过这批 prompt。流式数据集(IterableDataset)无法重新定位,其 worker 恢复时从 prompt 0 重启。

关键参数速查

以下参数定义于trl/experimental/async_distillation/async_distillation_config.py,该配置只含异步蒸馏特有项,其余训练参数沿用 transformersTrainingArguments

模型

参数默认值说明
model_init_kwargsNone传给AutoModelForCausalLM.from_pretrained的 kwargs;其中revision也用于加载 processing class
dtype"float32"学生加载精度("auto"/"bfloat16"/"float16"/"float32")。默认 float32 因异步 trainer 针对的 training-inference mismatch 度量对 trainer 自身精度敏感;端到端弥合还需学生 vLLM 以相同 dtype 服务。model_init_kwargs中的dtype优先;教师服务器不受影响
trust_remote_codeFalse允许加载 Hub 上带自定义代码的模型/分词器

生成

参数默认值说明
max_completion_length2048每完成结果最多生成的 token 数
temperature1.0采样学生 on-policy 完成结果的温度
top_p1.0nucleus 采样
top_k0top-k 采样;0禁用
min_pNone最小 token 概率(按最可能 token 概率缩放),须落在0.01.0,典型0.010.2
repetition_penalty1.0>1.0鼓励新 token,<1.0鼓励重复
chat_template_kwargsNone传给apply_chat_template的额外 kwargs

vLLM 服务器

参数默认值说明
vllm_server_base_url"http://localhost:8000"学生服务器基础 URL,用于生成与权重流
vllm_server_timeout240.0等学生服务器就绪的总超时(秒)
teacher_server_urls{"default": "http://localhost:8001"}教师服务器。每个都需--logprobs-mode processed_logprobs --max-logprobs -1;静态,从不向其传权重,也不需要 dev 模式。多条目启用 MOPD
request_timeout600对任一 vLLM 服务器的单请求超时(秒)
weight_sync_timeout1800权重传输超时(秒);超时 raise 而非挂起

蒸馏损失

参数默认值说明
beta0.0广义 JSD 插值系数;0.0前向 KL,1.0反向 KL;越界抛ValueError
teacher_temperature1.0散度两侧共用的 softmax 温度;与服务端 logprobs 计算绑定
teacher_top_k8每位置请求的候选数;超过20需教师带--max-logprobs -1
add_tail_bucketTrue是否追加尾桶
token_budgetNone单行最大真实 token 数。None时取学生服务器max_model_len(训练开始时查询),保证任何 rollout 样本不超预算;超长样本丢弃并计入batch/dropped_oversize_total<=0禁用预算,改为每 micro-batch 固定per_device_train_batch_size × num_processes个样本

异步流水线

参数默认值说明
max_inflight_tasks-1在途生成+打分任务上限;-1自动设为max_staleness × per_device_train_batch_size × gradient_accumulation_steps × num_processes
max_staleness4样本最多落后多少个权重更新,超过丢弃
queue_maxsize1024rollout 队列容量上限
weight_sync_steps1两次权重同步间隔的训练步数
heartbeat_stale_after_s300.0worker 心跳停滞超过该秒数即视为挂起并中止

日志

参数默认值说明
log_completionsFalse是否周期性记录 (prompt, completion) 对
log_completions_steps100记录间隔,按 worker 打分的样本数计而非优化器步(worker 与 trainer 是不同进程,看不到global_step
num_completions_to_printNonerich打印的完成结果数;None全部记录

TrainingArguments默认值不一致的参数已在上文标出:learning_rate1e-65e-5)、logging_steps1500)、bf16(未设fp16TrueFalse)、gradient_checkpointingTrueFalse)、ignore_data_skipTrueFalse)。另外__post_init__强制三条约束:不支持序列维并行(cp_size > 1sp_size > 1直接抛错,因为蒸馏在 trainer 内部于生成之后才构建模型输入,context-parallel / Ulysses 输入分片无法作用于原始生成 batch);teacher_server_urls至少一个条目;accelerator_config强制split_batches=Truedispatch_batches=True(主进程驱动 dataloader、batch 广播给其他进程,是异步 IterableDataset 正确工作的前提)。

设计边界与延伸阅读

该 trainer 刻意保持最小化,不打算长成通用方案:官方建议需要缺失功能时直接克隆仓库(https://gitcode.com/GitHub_Trending/tr/trl)后按自身需求改造;新功能只会在有显著社区需求时考虑。源码中的RolloutWorkerProtocolWeightTransferProtocol两个 Protocol 定义了可注入的自定义 rollout worker 与权重同步后端——测试正是靠注入 no-op 实现来脱离真实 vLLM 服务器运行的,这也给你自定义留了明确的接缝。

延伸阅读的仓库内路径:

  • 配置类:trl/experimental/async_distillation/async_distillation_config.py
  • 训练器与损失实现(_jsd_divergence_jsd_loss_chunk_narrow_top1_actual_support_add_tail_bucket):trl/experimental/async_distillation/async_distillation_trainer.py
  • rollout worker(_AsyncRolloutLoopRolloutSample_generate_and_score_one):trl/experimental/async_distillation/async_rollout_worker.py
  • 权重传输与 vLLM 客户端:trl/experimental/async_distillation/weight_transfer.pytrl/experimental/async_distillation/vllm_client.py
  • 单教师示例:examples/async_distillation_math/async_distillation_math.py
  • 双教师 MOPD 示例:examples/async_distillation_math/async_distillation_mopd.py
  • 文档:docs/source/async_distillation_trainer.md

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

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

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

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

立即咨询