ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

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

如何配置 TRL AsyncDistillationTrainer:三终端部署、beta 调参与日志排障一次讲清 如何配置 TRL AsyncDistillationTrainer三终端部署、beta 调参与日志排障一次讲清【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trlTRL 的AsyncDistillationTrainer把蒸馏中的生成、教师打分与梯度更新拆成并发执行教师模型永不在本地加载只需一个 vLLM 服务器 URL。本文带你部署三终端拓扑、选对beta与teacher_top_k并用日志指标定位生成侧或训练侧的瓶颈。读完本文你会掌握三件事如何用三个终端把教师服务器、学生 vLLM 服务器与训练进程跑在独立 GPU 上如何按训练阶段选择beta、teacher_top_k与add_tail_bucket理解支撑集收窄的原因如何用队列类与性能类指标区分生成受限、训练受限与服务退化三类瓶颈。环境要求需要vllm0.22.0与transformers5.2.0分布式训练仅支持 FSDP2不支持 DeepSpeed ZeRO。由于 vLLM 与 transformers 当前依赖约束冲突必须先装 vLLM、再强制安装 transformers先执行pip install vllm0.22.0再执行pip install transformers5.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 → 学生 vLLMHTTP/v1/completions带采样参数worker → 教师 vLLMHTTP/v1/completions带prompt_logprobsteacher_top_k与temperatureteacher_temperaturetrainer → 学生 vLLMNCCL 权重流每weight_sync_steps步一次worker ↔ trainermp.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.pyGSM8K、max_steps100、learning_rate1e-6、report_totrackio。# train_async_distillation.py 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()教师服务器GPU 0。两个参数都不可省--logprobs-mode processed_logprobs让teacher_temperature在服务端作用于返回的 logprobs否则教师静默返回原始 logprobs温度只影响学生侧--max-logprobs -1解除 vLLM 每 token 20 个 logprob 的上限teacher_top_k才能超过 20# 终端 1GPU 0 —— 教师静态永不更新无需 dev 模式 CUDA_VISIBLE_DEVICES0 vllm serve Qwen/Qwen2.5-1.5B-Instruct \ --port 8001 \ --logprobs-mode processed_logprobs \ --max-logprobs -1学生 vLLM 服务器GPU 1。VLLM_SERVER_DEV_MODE1与 NCCL 权重传输缺一不可否则 trainer 无法把新权重推进去# 终端 2GPU 1 —— 学生 vLLMdev 模式 NCCL 权重传输 CUDA_VISIBLE_DEVICES1 VLLM_SERVER_DEV_MODE1 vllm serve Qwen/Qwen2.5-0.5B-Instruct \ --port 8000 \ --weight-transfer-config {backend:nccl}训练进程GPU 2# 终端 3GPU 2 —— trainer 主进程 CUDA_VISIBLE_DEVICES2 accelerate launch train_async_distillation.py注意该 trainer 的默认值刻意区别于 transformers 的TrainingArgumentslearning_rate为1e-6非5e-5、bf16在未设fp16时默认True、gradient_checkpointing默认True、logging_steps默认1、ignore_data_skip默认True。️ 损失与调参beta、teacher_top_k 怎么选损失最小化的是学生与教师在逐位置 token 分布上的广义 Jensen-Shannon 散度与DistillationTrainer、ServerDistillationTrainer使用同一目标。beta必须落在[0.0, 1.0]越界直接ValueError是插值系数beta0.0为前向 KLmean-seeking权重落在教师概率高的区域是默认值beta1.0为反向 KLmode-seeking权重落在学生自己采样的区域中间值平滑过渡。调参时最关键的一点是支撑集计算散度用的候选集合随 beta 变化。原因不在数学而在传输协议——线上传的只有教师 top-k 切片不是完整词表beta0.0用教师报告的完整teacher_top_k宽支撑外加尾桶。前向 KL 恰好只需要教师分布给出的权重这份支撑正好够用beta ! 0.0支撑收窄为两个候选——教师自己的 top-1 token 和完成结果的实际 token。因为这是协议在不传更宽词表的前提下唯一能保证拿到教师 logprob 的两个身份任何更宽的覆盖都只是概率性近似不是保证。其中beta1.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适合冒烟测试真正训练时提到16–64合理邻近 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 乘词表。 多教师路由MOPDMOPD 是独立方法不属于本 trainer 核心目标所依据的单教师论文。它的完整流程三阶段通用 SFT → 每个领域独立做基于 RL 的专家训练 → 用 MOPD 把冻结的专家融合进一个学生。AsyncDistillationTrainer只实现第三阶段融合阶段各领域的专家必须已经存在例如用GRPOTrainer/RLOOTrainer单独训好、已经以 HTTP 提供服务你再把teacher_server_urls指向它们。MOPD 论文自己的 Stage 3 用的是反向 KL配置时应显式写beta1.0而不是沿用默认的beta0.0。路由规则与约束teacher_server_urls单个条目所有样本都由该教师打分多个条目每行的teacher_id列决定谁打分例如数学 prompt 走数学专家、代码 prompt 走代码专家各自独立服务每个样本只分发给它匹配的那一个教师绝不跨教师求平均或集成teacher_id缺失或映射不到任何条目的样本会直接抛ValueError不存在静默回退到默认教师。可运行的双教师示例在examples/async_distillation_math/async_distillation_mopd.pyGSM8K 路由给Qwen/Qwen2.5-1.5B-InstructmathPython 代码指令数据路由给Qwen/Qwen2.5-Coder-1.5B-Instructcode学生是Qwen/Qwen2.5-0.5B-Instruct显式beta1.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/inflight、rollout/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 占比行打包不紧、长样本多量化效应而非 bug1 万 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_k1宽的位置以 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_rowper-step 是求和、per-row 是均值预期有零点几百分比的偏差而非精确相等。所有指标里 step 都指完整优化器步而非 micro-batchmicro-batch 级量会明说batch/microbatches_per_step或以 per-row 形式给出batch/row_*、batch/samples_per_row。token 口径分三种generated是学生实际生成的 tokencompletions/*forwarded是前向处理过的全部 tokenprompt 生成trained是completion_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 classdtypefloat32学生加载精度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.0–1.0典型0.01–0.2repetition_penalty1.01.0鼓励新 token1.0鼓励重复chat_template_kwargsNone传给apply_chat_template的额外 kwargsvLLM 服务器参数默认值说明vllm_server_base_urlhttp://localhost:8000学生服务器基础 URL用于生成与权重流vllm_server_timeout240.0等学生服务器就绪的总超时秒teacher_server_urls{default: http://localhost:8001}教师服务器。每个都需--logprobs-mode processed_logprobs --max-logprobs -1静态从不向其传权重也不需要 dev 模式。多条目启用 MOPDrequest_timeout600对任一 vLLM 服务器的单请求超时秒weight_sync_timeout1800权重传输超时秒超时 raise 而非挂起蒸馏损失参数默认值说明beta0.0广义 JSD 插值系数0.0前向 KL1.0反向 KL越界抛ValueErrorteacher_temperature1.0散度两侧共用的 softmax 温度与服务端 logprobs 计算绑定teacher_top_k8每位置请求的候选数超过20需教师带--max-logprobs -1add_tail_bucketTrue是否追加尾桶token_budgetNone单行最大真实 token 数。None时取学生服务器max_model_len训练开始时查询保证任何 rollout 样本不超预算超长样本丢弃并计入batch/dropped_oversize_total0禁用预算改为每 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_processesmax_staleness4样本最多落后多少个权重更新超过丢弃queue_maxsize1024rollout 队列容量上限weight_sync_steps1两次权重同步间隔的训练步数heartbeat_stale_after_s300.0worker 心跳停滞超过该秒数即视为挂起并中止日志参数默认值说明log_completionsFalse是否周期性记录 (prompt, completion) 对log_completions_steps100记录间隔按 worker 打分的样本数计而非优化器步worker 与 trainer 是不同进程看不到global_stepnum_completions_to_printNone用rich打印的完成结果数None全部记录与TrainingArguments默认值不一致的参数已在上文标出learning_rate1e-6对5e-5、logging_steps1对500、bf16未设fp16时True对False、gradient_checkpointingTrue对False、ignore_data_skipTrue对False。另外__post_init__强制三条约束不支持序列维并行cp_size 1或sp_size 1直接抛错因为蒸馏在 trainer 内部于生成之后才构建模型输入context-parallel / Ulysses 输入分片无法作用于原始生成 batchteacher_server_urls至少一个条目accelerator_config强制split_batchesTrue、dispatch_batchesTrue主进程驱动 dataloader、batch 广播给其他进程是异步 IterableDataset 正确工作的前提。设计边界与延伸阅读该 trainer 刻意保持最小化不打算长成通用方案官方建议需要缺失功能时直接克隆仓库https://gitcode.com/GitHub_Trending/tr/trl后按自身需求改造新功能只会在有显著社区需求时考虑。源码中的RolloutWorkerProtocol与WeightTransferProtocol两个 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_buckettrl/experimental/async_distillation/async_distillation_trainer.pyrollout worker_AsyncRolloutLoop、RolloutSample、_generate_and_score_onetrl/experimental/async_distillation/async_rollout_worker.py权重传输与 vLLM 客户端trl/experimental/async_distillation/weight_transfer.py、trl/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),仅供参考
返回列表