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(RolloutSample) | prompt + 学生完成结果 + 教师逐位置稀疏分布 | 跨进程边界的唯一数据载体 |
| row | 一个 DP rank 在一个 micro-batch 中前向的内容 | 若干样本拼成的一条序列,position_ids逐样本重置 |
| row-slot | 一个优化器步容纳的行数 | grad_accum × world_size,是校验 batch 指标的基准 |
| staleness | 样本落后当前模型版本多少个权重更新 | 数据的"年龄",超过max_staleness即丢弃 |
| generated / forwarded / trained tokens | 学生生成的 / 前向处理的 / 损失实际计算的 token | trained ⊆ 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=1、prompt_logprobs=teacher_top_k、temperature=teacher_temperature,教师不生成任何新 token)。本地教师前向无法在这个无 CUDA 的子进程里跑,而 HTTP 打分让教师硬件与学生、trainer 完全解耦——这是与DistillationTrainer(教师本地加载,生成、教师前向、更新在同一进程顺序执行)的根本分叉。
rollout_buffer。mp.Queue,maxsize=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外面套两层:默认TokenBudgetBatcher(token_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),中间值线性插值。与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(去重后宽度 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_kwargs | None | from_pretrained关键字参数,revision同时用于加载 tokenizer | 模型需要特殊加载参数时 |
dtype | "float32" | 学生加载精度,model_init_kwargs中的dtype优先 | 学生 vLLM 服务 dtype 不一致导致 mismatch 时 |
trust_remote_code | False | 允许加载 Hub 自定义代码模型 | 使用自定义代码仓库时 |
生成采样
| 参数 | 默认值 | 说明 | 何时需修改 |
|---|---|---|---|
max_completion_length | 2048 | 每完成结果最大生成 token 数 | 长思维链或需要截断时 |
temperature | 1.0 | on-policy 采样温度 | 探索与稳定性权衡 |
top_p | 1.0 | nucleus 采样参数 | 同上 |
top_k | 0 | top-k 采样,0禁用 | 同上 |
min_p | None | 最小 token 概率,按最可能 token 概率缩放,典型0.01–0.2 | 抑制低概率 token |
repetition_penalty | 1.0 | 惩罚 prompt 与已生成文本中已出现 token | 重复退化时 |
chat_template_kwargs | None | 传给apply_chat_template的额外参数 | 模板需要开关(如 think 模式)时 |
vLLM 服务器
| 参数 | 默认值 | 说明 | 何时需修改 |
|---|---|---|---|
vllm_server_base_url | "http://localhost:8000" | 学生服务器,用于生成与权重更新 | 跨机部署时 |
vllm_server_timeout | 240.0 | 等待学生服务器就绪的总超时(秒) | 大模型加载慢时 |
teacher_server_urls | {"default": "http://localhost:8001"} | 教师服务器映射;多条目启用 MOPD,每行teacher_id选打分者 | 多教师或跨机部署 |
request_timeout | 600 | 单个 HTTP 请求超时(秒),对任意服务器 | 长序列打分慢时 |
weight_sync_timeout | 1800 | 权重传输超时(秒),超时 raise 而非挂死 | 大模型传输慢时 |
蒸馏损失
| 参数 | 默认值 | 说明 | 何时需修改 |
|---|---|---|---|
beta | 0.0 | 广义 JSD 插值,0前向 KL、1反向 KL | MOPD 按论文取1.0 |
teacher_temperature | 1.0 | 散度 softmax 温度,作用于教师(服务端)与学生两侧 | 软化/锐化教师分布 |
teacher_top_k | 8 | 每位置请求的教师候选数,完整词表从不传输 | 正式训练提到16–64 |
add_tail_bucket | True | 追加尾部桶,避免小候选集下散度趋零 | 一般不改 |
token_budget | None | 单行最大真实 token 数;None时取学生 vLLM 的max_model_len | 控制峰值内存与行填充率 |
异步流水线
| 参数 | 默认值 | 说明 | 何时需修改 |
|---|---|---|---|
max_inflight_tasks | -1 | 在途生成+打分任务上限;-1自动取max(max_staleness, 1) × samples_per_step | 生成吞吐不足时 |
max_staleness | 4 | 样本可落后当前版本的最大权重更新步数 | on-policy 性要求高时调小 |
queue_maxsize | 1024 | rollout 队列缓冲上限 | 生成快于训练时 |
weight_sync_steps | 1 | 两次权重同步之间的训练步数 | 同步开销占比高时调大 |
heartbeat_stale_after_s | 300.0 | worker 心跳超时秒数,超过判挂起并中止 | 一般不改 |
日志
| 参数 | 默认值 | 说明 | 何时需修改 |
|---|---|---|---|
log_completions | False | 每 N 个已打分样本记录一批 (prompt, completion) | 需要人工抽检时 |
log_completions_steps | 100 | 两次记录之间被打分的样本数;按 worker 打分计数,非优化器步 | 配合上行 |
num_completions_to_print | None | 用 rich 打印的完成结果数,None全部 | 日志刷屏时 |
⚠️ 与
TrainingArguments默认值不同:logging_steps默认1(非500);gradient_checkpointing默认True(非False);bf16在未设置fp16时默认True;learning_rate默认1e-6(非5e-5);ignore_data_skip默认True(非False,skip-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_batches=True、dispatch_batches=True:主进程驱动 dataloader,batch 广播而非各进程独立拉取。- 前向实现硬编码为 FlashAttention(
kernels-community/flash-attn3),padding-free 模式依赖position_ids重置;use_liger_kernel=True抛NotImplementedError。
部署与运行
最小训练脚本(完整可运行示例见 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.0且transformers>=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 上:
- 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。
- 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 启动训练:
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 |
jsd、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_s | trainer 因队列空阻塞了多久 | 持续走高 | 生成侧产速不足 |
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_staleness | off-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_s与rollout/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 占比 | 某教师逼近 0 | teacher_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_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。IterableDataset无len(),恢复时 worker 从 prompt 0 重启。
本模块不做什么
- 不支持本地(进程内 GPU)教师前向:无 CUDA 的子进程跑不了,回主进程的路径未实现。
- 不支持序列维并行(
cp_size/sp_size> 1 抛错)。 - 不支持
use_liger_kernel(抛NotImplementedError)。 - 不做跨教师集成/平均;没有奖励函数与分组基线(区别于 GRPO)。
- 分布式仅 FSDP2,DeepSpeed 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.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/:
ServerDistillationTrainer,beta支撑集收窄逻辑的镜像来源 - trl/experimental/async_grpo/:
AsyncGRPOTrainer,本 trainer 的架构原型,planner/worker 机制逐行移植自它
【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考