TRL DistillationTrainer 实战指南:基于广义 JSD 的 On-Policy 知识蒸馏
【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl
TRL 的DistillationTrainer实现了《On-Policy Distillation of Language Models: Learning from Self-Generated Mistakes》论文提出的 Generalized Knowledge Distillation(GKD)方法:让较小的 student 模型在自己生成的回复(on-policy)上匹配 teacher 模型的完整 next-token 分布,从而克服传统蒸馏中训练与推理分布不一致的问题。读完本文,你将掌握如何用几行代码完成一次蒸馏训练、理解广义 JSD 损失的数学含义与内存优化实现、并用 vLLM 加速生成、接入 PEFT/LoRA、训练 Agent 与多模态 VLM。
一、核心原理:什么是 On-Policy 知识蒸馏
传统知识蒸馏(KD)用固定的 teacher 输出序列训练 student,但训练时看到的序列分布与 student 推理时自己生成的序列分布存在分布偏移(distribution mismatch)。GKD 的解决思路是:让 student 在自生成的输出序列上学习,由 teacher 对这些序列给出反馈(next-token 分布),学生据此修正自己的错误。
DistillationTrainer的具体实现(见 distillation_trainer.py 类文档)分为两步:
- 生成:每个训练步,student 对采样的 prompt 自回归生成一批 completions(可选 vLLM 加速);
- 匹配分布:对生成的 completion token,用 teacher 的完整 next-token 分布与 student 的分布计算**分块 Jensen-Shannon 散度(JSD)**损失——teacher 的稠密 logits 从不整体物化到显存,因此显存开销可控。
该 trainer 的贡献者之一是 Carlos Miguel Patiño。若需要将生成与训练解耦、并且 teacher 通过 HTTP 服务打分而非本地前向,可参考其异步版本AsyncDistillationTrainer。
二、快速开始:一行脚本完成蒸馏
以下示例把 Qwen2.5-0.5B-Instruct 蒸馏自Qwen/Qwen2.5-1.5B-Instructteacher,训练数据使用trl-lib/ultrafeedback-prompt数据集的 prompt(纯 prompt 数据集,仅需 prompt 列)。
# train_distillation.py from datasets import load_dataset from trl import DistillationTrainer dataset = load_dataset("trl-lib/ultrafeedback-prompt", split="train") trainer = DistillationTrainer( model="Qwen/Qwen2.5-0.5B-Instruct", teacher_model="Qwen/Qwen2.5-1.5B-Instruct", train_dataset=dataset, ) trainer.train()执行:
accelerate launch train_distillation.py训练完成后模型会保存在配置的output_dir中,并自动附带模型卡片(_save_checkpoint会调用create_model_card)。
三、深入理解蒸馏方法
3.1 生成 completions
在每个训练步,student 为采样的 prompts 生成一批 completions。源码中对数据加载做了专门优化:get_train_dataloader返回的 batch 大小是per_device_train_batch_size × gradient_accumulation_steps,即一次加载一个"生成批次"(generation batch),每个累积窗口内只生成一次 completions 并复用,避免在每个微步重复生成(见 distillation_trainer.py 的get_train_dataloader与_prepare_inputs)。
3.2 计算损失:广义 JSD
损失是 student 分布 (p_S) 与 teacher 分布 (p_T) 之间的广义 Jensen-Shannon 散度,由beta插值:
[ \mathcal{L}\beta = \beta , \mathbb{D}{\mathrm{KL}}!\left[ p_T | p_M \right] + (1 - \beta) , \mathbb{D}_{\mathrm{KL}}!\left[ p_S | p_M \right], \qquad p_M = (1 - \beta) , p_S + \beta , p_T ]
其中 (p_M) 是两个分布的 (\beta)-混合。端点退化为纯散度:
beta=0.0:前向 KL (\mathbb{D}_{\mathrm{KL}}[p_T | p_S]);beta=1.0:反向 KL (\mathbb{D}_{\mathrm{KL}}[p_S | p_T])(默认值,让 student 分布"收窄"去贴合 teacher 的高概率区,是实际中最常用的选择);beta=0.5:标准 JSD。
注意:这里的
beta与 GRPO 的beta含义完全不同。GRPO 的beta是相对参考模型的 KL 惩罚系数;而这里它直接选择散度本身,没有参考模型 KL 惩罚(见 distillation_config.py 中beta字段的注释)。
内存优化实现:实际计算时,vocab 投影与散度按 chunk 分块进行——有效位置被argsort打包到前部,每 chunk 只保留[chunk_size, vocab_size]大小的 logits(chunk 大小为 256,见_CHUNKED_LM_HEAD_CHUNK_SIZE),并配合梯度检查点(torch.utils.checkpoint),因此峰值激活内存不再随"全 vocab × 序列长度"的 logits 张量增长。详见 reducing_memory_usage.md。
源码中散度计算对logit_scale(Cohere 系)与final_logit_softcapping(Gemma 系)做了兼容处理,且 teacher 侧在torch.no_grad()下投影,不产生任何梯度图。从 test_distillation_trainer.py 可以看到完整的单元测试覆盖:beta ∈ {0.0, 0.5, 1.0}与朴素全 vocab 实现的数值一致性、teacher/student 隐藏宽度不同、logit scale/softcap、温度缩放、lm_head bias 等场景。
3.3 预期数据集格式
数据集应为 对话式(conversational) 且 仅 prompt 格式,因为 student 会自己 on-policy 生成回复,数据只需 prompt:
{"prompt": [{"role": "user", "content": "What color is the sky?"}]}也可以使用纯文本(standard)格式。如果使用IterableDataset(流式数据集),必须在训练参数中设置max_steps,否则无法推断数据集长度来配置学习率调度器与训练循环。
四、DistillationConfig 关键参数速查
DistillationConfig继承自 transformers 的TrainingArguments,额外声明了以下参数(默认值以 distillation_config.py 为准)。这些参数均可通过HfArgumentParser转为命令行参数,这也是trl distillationCLI 能够工作的基础。
模型与 teacher 相关
| 参数 | 默认值 | 说明 |
|---|---|---|
model_init_kwargs | None | 传给AutoModelForCausalLM.from_pretrained的 kwargs,revision也会用于加载 processing class |
teacher_model_name_or_path | None | teacher 模型名或路径(本地加载时使用) |
teacher_model_revision | None | teacher 的 revision(分支、tag 或 commit hash) |
teacher_model_init_kwargs | None | 实例化 teacher 时传给from_pretrained的 kwargs |
trust_remote_code | False | 是否允许加载 Hub 上带自定义代码的模型/分词器(student 与 teacher 都生效) |
disable_dropout | False | 训练时是否关闭 student 的 dropout |
数据预处理
| 参数 | 默认值 | 说明 |
|---|---|---|
remove_unused_columns | False | 默认不删列,trainer 直接消费原始prompt列并 on-policy 生成 |
max_completion_length | 512 | 每个 completion 最多生成的 token 数 |
ds3_gather_for_generation | True | DeepSpeed ZeRO-3 下是否聚合权重用于生成(提速);关闭可训练超单卡显存的模型但生成变慢,且与 vLLM 不兼容 |
shuffle_dataset | True | 是否打乱训练集 |
pad_to_multiple_of | None | 若设置,prompt/completion ids 填充到该值的倍数 |
生成控制
| 参数 | 默认值 | 说明 |
|---|---|---|
temperature | 1.0 | 采样与损失计算共同的温度,越高分布越软 |
top_p | 1.0 | nucleus 采样参数 |
top_k | 0 | top-k 采样,0表示关闭 |
min_p | None | 最小 token 概率(按最可能 token 概率缩放),典型值0.01–0.2 |
repetition_penalty | 1.0 | 重复惩罚,>1.0鼓励新 token,<1.0鼓励重复 |
generation_kwargs | None | 额外传给GenerationConfig/SamplingParams的 kwargs(如suppress_tokens、num_beams);与上述参数冲突时以它为准 |
chat_template_kwargs | None | 传给apply_chat_template的额外 kwargs |
cache_implementation | None | 非 vLLM 生成时的 cache 实现 |
vLLM 加速
| 参数 | 默认值 | 说明 |
|---|---|---|
use_vllm | False | 是否用 vLLM 生成 on-policy completions |
vllm_mode | "colocate" | "colocate"(同进程共享 GPU)或"server"(独立进程/GPU,HTTP 通信) |
vllm_model_impl | "vllm" | vLLM 后端,"vllm"或"transformers" |
vllm_enable_sleep_mode | False | 优化器步骤期间 offload student 权重的 sleep 模式 |
vllm_server_base_url | None | 若提供,忽略 host/port |
vllm_server_host/vllm_server_port | "0.0.0.0"/8000 | server 模式的 host/port |
vllm_server_timeout | 240.0 | 连接 server 超时 |
vllm_group_port | 51216 | vLLM 权重更新组(NCCL)端口 |
vllm_gpu_memory_utilization | 0.3 | colocate 模式下 vLLM 引擎的 GPU 显存占用比例 |
vllm_max_model_length | None | colocate 引擎最大序列长度 |
vllm_tensor_parallel_size | 1 | colocate 引擎的张量并行度 |
训练与日志
| 参数 | 默认值 | 说明 |
|---|---|---|
beta | 1.0 | 广义 JSD 插值系数,0.0=前向 KL,1.0=反向 KL,0.5=JSD;范围校验[0, 1],越界直接抛ValueError |
max_tool_calling_iterations | None | Agent 训练时工具调用轮数上限;None表示无限制,模型生成无工具调用的回复轮即停止 |
log_completions | False | 每logging_steps记录一批 (prompt, completion) 样本,可用 rich 打印、wandb/trackio 记录并保存 parquet |
num_completions_to_print | None | rich 打印的 completion 数量,None表示全部 |
log_unique_prompts | False | 日志中是否只保留唯一 prompt |
另外,DistillationConfig对TrainingArguments的几个默认值做了覆盖:logging_steps默认10、gradient_checkpointing默认True、未显式设置 fp16 时bf16默认True、learning_rate默认1e-6。
源码中__post_init__还校验了序列并行不兼容性:蒸馏需要在生成后于 trainer 内部构建模型输入,因此 Transformers 的 context-parallel / Ulysses 序列并行(cp_size > 1或sp_size > 1)暂不支持,会直接报错提示设置为 1 或关闭parallelism_config。
五、训练日志指标
训练与评估过程中记录的指标如下(由 distillation_trainer.py 的_generate与compute_loss产生):
num_tokens:迄今处理的 token 总数(含 prompt 与 completion);使用工具时只统计非工具 token;step_time:每个训练步平均耗时(秒,含生成);completions/mean_length、completions/min_length、completions/max_length:生成 completion 的平均/最小/最大长度(工具场景只统计非工具 token);completions/mean_terminated_length、completions/min_terminated_length、completions/max_terminated_length:以 EOS 正常终止的 completion 的长度统计;completions/clipped_ratio:被截断(clip)的 completion 占比;tools/call_frequency:生成批次中每条 completion 平均工具调用次数(仅当提供tools时记录);tools/failure_frequency:工具调用失败比例(工具未找到、抛异常或调用类型不支持),无调用时为0.0(仅当提供tools时记录);entropy:生成 completions 上 token 预测的平均熵(单位 nats)。
评估模式下这些指标会自动加上eval_前缀。
六、定制与加速
6.1 用 vLLM 加速生成
On-policy 方法的生成常常是训练瓶颈。vLLM 是高吞吐、低延迟的推理引擎,先安装:
pip install trl[vllm]支持两种模式:
Option 1:Colocate 模式(默认)。vLLM 在 trainer 进程内运行,与训练模型共享 GPU 显存,无需启动独立服务,可提升 GPU 利用率,但可能与训练竞争显存。
from trl import DistillationConfig training_args = DistillationConfig( ..., use_vllm=True, # vllm_mode="colocate" by default )Option 2:Server 模式。vLLM 在独立进程(及独立 GPU)中运行,通过 HTTP 与 trainer 通信,适合有专用推理 GPU 的场景。
- 启动 vLLM server:
VLLM_SERVER_DEV_MODE=1 vllm serve <model_name> \ --weight-transfer-config '{"backend": "nccl"}' \ --logprobs-mode processed_logprobs \ --max-logprobs -1- 训练脚本开启 server 模式:
from trl import DistillationConfig training_args = DistillationConfig( ..., use_vllm=True, vllm_mode="server", )⚠️ 警告:server 必须使用与 trainer 不同的 GPU,否则可能触发 NCCL 错误;可用
CUDA_VISIBLE_DEVICES环境变量指定 GPU。💡 提示:根据模型规模与训练显存需求,可能需要调整
vllm_gpu_memory_utilization,避免显存利用不足或 OOM。
使用 vLLM 时,trainer 会在每个global_step变化后同步 student 权重到 vLLM 引擎(源码中_generate_single_turn内的vllm_generation.sync_weights())。更多细节见 speeding_up_training.md。
6.2 用 PEFT 训练适配器
支持与 🤗 PEFT 深度集成,只训练 LoRA 适配器并分享到 Hub,而不是训练整个 student:
from datasets import load_dataset from trl import DistillationTrainer from peft import LoraConfig dataset = load_dataset("trl-lib/ultrafeedback-prompt", split="train") trainer = DistillationTrainer( model="Qwen/Qwen2.5-0.5B-Instruct", teacher_model="Qwen/Qwen2.5-1.5B-Instruct", train_dataset=dataset, peft_config=LoraConfig(), ) trainer.train()⚠️ 警告:蒸馏损失直接读取
lm_head.weight,并通过 backbone 前向(_get_last_hidden_state)绕过PeftModel.forward()。因此:
- 在
lm_head上挂 adapter(target_modules含"lm_head")会被拒绝——head 上的可训练 adapter 位于损失永远看不到的独立子模块中,会静默得不到梯度,源码中显式抛出ValueError;- Prompt 学习类方法(PromptTuning、PrefixTuning、P-Tuning)同样会被拒绝,因为虚拟 token 通过
PeftModel.forward()注入,而损失直接调用 backbone 会漏掉它们;- 如需训练 head,请改用
modules_to_save=["lm_head"]。
另外,源码对 PEFT + DeepSpeed ZeRO-3 场景做了适配:非量化模型下自动传autocast_adapter_dtype=False规避混合 dtype 的 TypeError;ZeRO-3 下强制use_reentrant=True的梯度检查点,并为 PEFT 开启enable_input_require_grads()。QLoRA(量化模型)时 adapter 权重会转为 bf16。
七、Agent 训练:工具调用与多模态工具响应
DistillationTrainer支持 Agent 训练:student 在生成过程中调用工具,并在整个轨迹上进行蒸馏。工具结果 token 会被 mask 出损失,student 只在自己生成的 token 上被训练。
7.1 定义工具
tools参数接收一组 Python 函数。每个工具必须是带类型注解的参数与返回值、并配有Google 风格 docstring(说明用途、参数与返回值)的标准函数:
from trl import DistillationTrainer def multiply(a: int, b: int) -> int: """ Multiplies two integers. Args: a: The first integer. b: The second integer. Returns: The product of the two integers. """ return a * b trainer = DistillationTrainer( tools=[multiply], ..., )💡 提示:工具调用循环要求 chat template 是prefix-preserving的(追加工具消息不得改变先前消息的渲染)。对已知模型族(如 Qwen3、DeepSeek-V3),TRL 会在启用工具时自动替换为打了补丁的训练模板,完整清单见 chat_templates.md。
使用DistillationConfig的max_tool_calling_iterations限制工具调用轮数;默认无限制,student 生成不包含工具调用的回复轮即停止。
⚠️ 警告:暂不支持异步工具(async),请传入同步函数——源码中
inspect.iscoroutinefunction检查会直接抛ValueError。另外启用工具要求 transformers ≥ 5.0.0,且低版本需要jmespath(transformers ≥ 5.13 不再需要)。
7.2 多模态工具响应
工具可以返回"图片 + 文本"的内容块列表,适用于 VLM Agent 训练(截图、图表、摄像头画面等视觉反馈):
from PIL import Image def take_screenshot() -> list: """ Takes a screenshot of the current screen. Returns: The screenshot image with a description. """ img = Image.open("screenshot.png") return [{"type": "image", "image": img}, {"type": "text", "text": "Here is the screenshot."}]返回的图片会自动注入对话,并在后续生成轮次中传给 VLM。工具循环中的工具结果、失败统计分别由tools/call_frequency与tools/failure_frequency指标反映;损失通过completion_mask × tool_mask精确屏蔽工具结果 token。
八、训练视觉语言模型(VLM)
DistillationTrainer支持在含文本与图片的多模态数据集上蒸馏 VLM:student 与 teacher 都传 VLM,数据集为仅 prompt 格式,带image列(单图)或images列(多图列表)。数据集结构见 dataset_formats.md。
已在以下模型上验证:
- Gemma 3——如
google/gemma-3-4b-it; - LLaVA-NeXT——如
llava-hf/llava-v1.6-mistral-7b-hf; - Qwen2-VL——如
Qwen/Qwen2-VL-2B-Instruct; - Qwen2.5-VL——如
Qwen/Qwen2.5-VL-3B-Instruct。
💡 提示:不保证兼容所有 VLM。如果你认为某个模型应当被支持,可以提交 issue 或直接提交 PR。
源码为 VLM 做了大量细节处理:_tokenize_prompts从对话消息中提取图片并调用apply_chat_template;前向时通过base_model(多模态包装器)注入视觉 token;处理了 Qwen 的image_grid_thw、Gemma/SmolVLM2/LLaVa-Next 的pixel_values、LLaVa-Next 的image_sizes、LFM2-VL 的spatial_shapes等各模型字段;工具图片混入后还会重建mm_token_type_ids/token_type_ids。
九、命令行接口
用trl distillationCLI 可从命令行直接启动蒸馏训练,支持完整训练与 LoRA,复用标准ModelConfig参数。相关命令的注册位于 cli/commands/init.py,脚本入口见 scripts/distillation.py。
# 完整训练: trl distillation \ --model_name_or_path Qwen/Qwen2.5-0.5B-Instruct \ --teacher_model_name_or_path Qwen/Qwen2.5-1.5B-Instruct \ --dataset_name trl-lib/ultrafeedback-prompt \ --learning_rate 2e-5 \ --per_device_train_batch_size 4 \ --gradient_accumulation_steps 8 \ --output_dir distilled-model \ --num_train_epochs 1# LoRA: trl distillation \ --model_name_or_path Qwen/Qwen2.5-0.5B-Instruct \ --teacher_model_name_or_path Qwen/Qwen2.5-1.5B-Instruct \ --dataset_name trl-lib/ultrafeedback-prompt \ --learning_rate 2e-4 \ --per_device_train_batch_size 4 \ --gradient_accumulation_steps 8 \ --output_dir distilled-model \ --num_train_epochs 1 \ --use_peft \ --lora_r 64 \ --lora_alpha 16在 scripts/distillation.py 中还提供了等价的python trl/scripts/distillation.py ...用法。脚本内部会校验--teacher_model_name_or_path必须提供;student 的量化配置走 trainer 的quantization_config参数,teacher 的量化配置则放在teacher_model_init_kwargs中(两者不能同时出现在model_init_kwargs,否则 trainer 拒绝)。
十、实现细节与边界条件
从源码可以提炼出几个值得注意的实现事实(对应 distillation_trainer.py):
- 词表必须一致:
__init__中会校验 student 与 teacher 的vocab_size必须相等,因为损失比较的是两者在全词表上的完整 next-token 分布;跨分词器蒸馏需要改用 GOLD 方法; - teacher 不参与梯度:teacher 以
evaluation_mode=True通过accelerator.prepare_model准备(DeepSpeed 下走prepare_deepspeed),前向包在torch.no_grad()中,teacher 参数不会累积梯度; - teacher 与 student 隐藏宽度可以不同:每个模型按自己的隐藏宽度扁平化后分别通过各自的
lm_head投影,只有词表必须一致(分块损失函数_chunked_divergence_loss明确支持); - 生成批次的复用机制:生成只发生在每个梯度累积窗口的开头(
_prepare_inputs中_step % gradient_accumulation_steps == 0时),通过RepeatSampler与_buffered_inputs将一次生成的结果切成多个微批,显著节省生成开销;生成批次大小 =per_device_train_batch_size × num_processes × gradient_accumulation_steps; - 流式数据集约束:
IterableDataset要求dispatch_batches=False与dataloader_num_workers=0(源码会强制覆盖并告警),以保证生成批次的分组顺序; - 断点续训安全:
_buffered_inputs=None时(如从 checkpoint 恢复)会在首个微步重新生成,保证正确性; - 数值正确性有测试保障:
tests/test_distillation_trainer.py中对分块损失与朴素全 vocab 实现做了逐 beta 的数值对比,并覆盖 bf16 hidden + fp32 weight、不同隐藏宽度、logit scale/softcap、温度、bias 等边界。
至此,你已掌握DistillationTrainer从原理到实战的完整路径:理解 GKD 的 on-policy 思想与广义 JSD 的beta语义、按参数表调优生成与显存、用 vLLM/PEFT 加速与轻量化、扩展 Agent 工具调用与 VLM 多模态蒸馏,并可通过 CLI 一键启动完整训练或 LoRA 训练。
【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考