尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

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

TRL AsyncDistillationTrainer 深度解析:解耦生成与梯度更新、稀疏教师打分与指标诊断 TRL AsyncDistillationTrainer 深度解析解耦生成与梯度更新、稀疏教师打分与指标诊断【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trlAsyncDistillationTrainer 是 TRL 实验模块中的异步 on-policy 蒸馏训练器学生完成结果由后台 rollout worker 生成、由远端教师服务器打分训练与生成并发。与同步的 DistillationTrainer 相比核心差异只有一句话教师永不被本地加载只需一个 vLLM 服务器 URL教师可运行在与学生、trainer 完全不同的硬件上。术语速览术语一句话定义直觉类比或最小示例rollout一次 prompt 被学生生成一次、被教师打分一次一趟生成打分往返恰好产出一个训练样本sampleRolloutSampleprompt 学生完成结果 教师逐位置稀疏分布跨进程边界的唯一数据载体row一个 DP rank 在一个 micro-batch 中前向的内容若干样本拼成的一条序列position_ids逐样本重置row-slot一个优化器步容纳的行数grad_accum × world_size是校验 batch 指标的基准staleness样本落后当前模型版本多少个权重更新数据的年龄超过max_staleness即丢弃generated / forwarded / trained tokens学生生成的 / 前向处理的 / 损失实际计算的 tokentrained ⊆ generatedforwarded prompt generated尾部桶tail bucketteacher_top_k之外追加的一个候选承载剩余概率质量把其他所有 token显式变成一个选项MOPD多教师 on-policy 蒸馏该 trainer 只实现其融合阶段数据集teacher_id列决定哪个教师打分架构与数据流ROLLOUT WORKERspawn 子进程CUDA_VISIBLE_DEVICES 已清空 PROMPTmessage list 可选 teacher_id ├─ generate 学生 vLLM /v1/completions采样完成结果 └─ score 路由教师 vLLM /v1/completionsprompt_logprobsteacher-forced └─ RolloutSampleprompt 完成结果 稀疏教师分布 ═══════ 进程边界rollout_buffermp.Queuemaxsize queue_maxsize═══════ 主进程训练循环FSDP2/DDP SAMPLE → staleness 检查 max_staleness 丢弃 └─ BatcherTokenBudget 默认 / FixedCount→ ROW每 DP rank 一行Σ Lᵢ² 平衡 └─ MICRO-BATCH → PACKED ROW → FORWARD_jsd_divergencebs1 └─ × grad_accum → OPTIMIZER STEP 每 weight_sync_steps 步NCCL 权重传输 → 学生 vLLMmodel_version 1Rollout 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_tokens1、prompt_logprobsteacher_top_k、temperatureteacher_temperature教师不生成任何新 token。本地教师前向无法在这个无 CUDA 的子进程里跑而 HTTP 打分让教师硬件与学生、trainer 完全解耦——这是与DistillationTrainer教师本地加载生成、教师前向、更新在同一进程顺序执行的根本分叉。rollout_buffer。mp.Queuemaxsizequeue_maxsize默认 1024。worker 每打出一个 sample 就put队列满时阻塞阻塞时长计入rollout/backpressure_strainer 侧的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_weightpause vLLM → 全 rank barrier → 经 NCCL 流式传输可训练参数FSDP2 下逐参数full_tensor()all-gather避免整模型物化→ resumerank 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为前向 KLmean-seeking默认1.0为反向 KLmode-seeking中间值线性插值。与DistillationTrainer和ServerDistillationTrainer使用同一目标函数三者行为一致。设某位置教师分布为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去重后宽度 2JSD 插值混合分布的前向项需要教师 top-1反向项只需要实际 token1.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_bucketTrue默认时_add_tail_bucket追加第 K1 个元素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模型需要特殊加载参数时dtypefloat32学生加载精度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.01–0.2抑制低概率 tokenrepetition_penalty1.0惩罚 prompt 与已生成文本中已出现 token重复退化时chat_template_kwargsNone传给apply_chat_template的额外参数模板需要开关如 think 模式时vLLM 服务器参数默认值说明何时需修改vllm_server_base_urlhttp://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.0teacher_temperature1.0散度 softmax 温度作用于教师服务端与学生两侧软化/锐化教师分布teacher_top_k8每位置请求的教师候选数完整词表从不传输正式训练提到16–64add_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非500gradient_checkpointing默认True非Falsebf16在未设置fp16时默认Truelearning_rate默认1e-6非5e-5ignore_data_skip默认True非Falseskip-and-replay 循环不适用于实时 rollout 队列trainer 会强制置True。约束关系__post_init__与__init__强制beta必须在[0.0, 1.0]否则ValueError。序列维并行不支持parallelism_config中cp_size 1或sp_size 1直接抛错——蒸馏在生成之后才于 trainer 内部构建模型输入transformers 的 context/Ulysses 输入分片无法作用于原始生成 batch。teacher_server_urls至少一个条目None时回填{default: http://localhost:8001}。accelerator_config被强制为split_batchesTrue、dispatch_batchesTrue主进程驱动 dataloaderbatch 广播而非各进程独立拉取。前向实现硬编码为 FlashAttentionkernels-community/flash-attn3padding-free 模式依赖position_ids重置use_liger_kernelTrue抛NotImplementedError。部署与运行最小训练脚本完整可运行示例见 examples/async_distillation_math/async_distillation_math.pyfrom 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()⚠️ 环境要求vllm0.22.0且transformers5.2.0分布式训练仅支持 FSDP2不支持 DeepSpeed ZeRO。两者当前存在冲突的依赖约束先装 vLLM 再强制装 transformerspip install vllm0.22.0 pip install transformers5.2.0 --no-deps三终端部署教师、学生 vLLM、trainer 必须在不同 GPU 上GPU 0 启动教师服务器静态永不更新无需 dev 模式CUDA_VISIBLE_DEVICES0 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。GPU 1 启动学生 vLLM 服务器dev 模式 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 启动训练CUDA_VISIBLE_DEVICES2 accelerate launch train_async_distillation.py常见启动失败排查现象最可能原因修复动作卡在等待学生服务器vllm_server_timeout超时学生服务器未启动或 model id 不一致核对vllm serve启动参数与model字符串权重传输超时异常weight_sync_timeout学生服务器缺VLLM_SERVER_DEV_MODE1或 NCCL 后端未启用补环境变量与--weight-transfer-configteacher_top_k设到 20 以上报错教师未带--max-logprobs -1重启教师服务器并加该参数sample/dropped_stale_total持续增长生成慢于训练且max_staleness过紧调大max_inflight_tasks/queue_maxsize或降weight_sync_stepsjsd、entropy整窗口为 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_sMOPD 看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-reducebatch/row_fill_frac行 token 数相对token_budget长期偏低长度量化效应1 万 token 样本铺不满 3.2 万预算调token_budgetperf/rollout_wait_s与rollout/backpressure_s是镜像不会同时很大——二者读数直接告诉你瓶颈在哪一侧。目标函数是否真的收敛你会看到jsd下降但entropy同步崩塌学生在收窄而非学习。指标回答的子问题异常方向指向的根因jsd广义 JSD 是否在按beta收敛平台期或回升学习停滞或 off-policy 干扰对照 stalenessentropy学生自身预测熵随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_s高rank 偏斜perf/fwd_s / perf/fwd_bwd_s前向占比高于 1/3反向便宜或重计算发生在前向30 秒体检只看四个数——jsd学没学到、sample/rollout_queue_sizeperf/rollout_wait_svsrollout/backpressure_s瓶颈在哪侧、perf/mfu_wall_clock算力利用率、sample/staleness_meanoff-policy 积累。四个数健康则系统平衡再下钻到对应小节。高级用法与边界MOPD多教师路由蒸馏前提各领域的专家教师必须已存在例如分别用GRPOTrainer/RLOOTrainer训练并经 HTTP 服务每个教师必须与学生共享 tokenizer——完成结果以原始 token id 传输教师报告的候选 id 直接索引学生词表词表不同的教师会把学生训练到错误的 token 上且这种错误是静默的除非教师词表比学生大。关键 diffteacher_server_urls多条目如{math: ..., code: ...}数据集每行携带teacher_id列缺失或未映射直接报错而非回退每个样本只分发给其匹配的一个教师绝不跨教师平均或集成。论文自身 Stage 3 使用反向 KL需显式beta1.0trainer 默认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。IterableDataset无len()恢复时 worker 从 prompt 0 重启。本模块不做什么不支持本地进程内 GPU教师前向无 CUDA 的子进程跑不了回主进程的路径未实现。不支持序列维并行cp_size/sp_size 1 抛错。不支持use_liger_kernel抛NotImplementedError。不做跨教师集成/平均没有奖励函数与分组基线区别于 GRPO。分布式仅 FSDP2DeepSpeed ZeRO 不支持。扩展点RolloutWorkerProtocol需暴露rollout_buffer、metrics_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.pyNCCL 权重传输客户端trl/experimental/async_distillation/vllm_client.pyvLLM HTTP 客户端就绪等待、max_model_len查询可运行示例examples/async_distillation_math/async_distillation_math.py单教师 GSM8Kexamples/async_distillation_math/async_distillation_mopd.py双教师 MOPD论文核心目标同步单教师 on-policy 蒸馏arXiv:2306.13649MOPD多教师能力融合仅其 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),仅供参考
返回列表