TRL DistillationTrainer 实战指南:基于广义 JSD 的 On-Policy 知识蒸馏
2026/9/13 10:22:06 网站建设 项目流程

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 类文档)分为两步:

  1. 生成:每个训练步,student 对采样的 prompt 自回归生成一批 completions(可选 vLLM 加速);
  2. 匹配分布:对生成的 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_kwargsNone传给AutoModelForCausalLM.from_pretrained的 kwargs,revision也会用于加载 processing class
teacher_model_name_or_pathNoneteacher 模型名或路径(本地加载时使用)
teacher_model_revisionNoneteacher 的 revision(分支、tag 或 commit hash)
teacher_model_init_kwargsNone实例化 teacher 时传给from_pretrained的 kwargs
trust_remote_codeFalse是否允许加载 Hub 上带自定义代码的模型/分词器(student 与 teacher 都生效)
disable_dropoutFalse训练时是否关闭 student 的 dropout

数据预处理

参数默认值说明
remove_unused_columnsFalse默认不删列,trainer 直接消费原始prompt列并 on-policy 生成
max_completion_length512每个 completion 最多生成的 token 数
ds3_gather_for_generationTrueDeepSpeed ZeRO-3 下是否聚合权重用于生成(提速);关闭可训练超单卡显存的模型但生成变慢,且与 vLLM 不兼容
shuffle_datasetTrue是否打乱训练集
pad_to_multiple_ofNone若设置,prompt/completion ids 填充到该值的倍数

生成控制

参数默认值说明
temperature1.0采样与损失计算共同的温度,越高分布越软
top_p1.0nucleus 采样参数
top_k0top-k 采样,0表示关闭
min_pNone最小 token 概率(按最可能 token 概率缩放),典型值0.01–0.2
repetition_penalty1.0重复惩罚,>1.0鼓励新 token,<1.0鼓励重复
generation_kwargsNone额外传给GenerationConfig/SamplingParams的 kwargs(如suppress_tokensnum_beams);与上述参数冲突时以它为准
chat_template_kwargsNone传给apply_chat_template的额外 kwargs
cache_implementationNone非 vLLM 生成时的 cache 实现

vLLM 加速

参数默认值说明
use_vllmFalse是否用 vLLM 生成 on-policy completions
vllm_mode"colocate""colocate"(同进程共享 GPU)或"server"(独立进程/GPU,HTTP 通信)
vllm_model_impl"vllm"vLLM 后端,"vllm""transformers"
vllm_enable_sleep_modeFalse优化器步骤期间 offload student 权重的 sleep 模式
vllm_server_base_urlNone若提供,忽略 host/port
vllm_server_host/vllm_server_port"0.0.0.0"/8000server 模式的 host/port
vllm_server_timeout240.0连接 server 超时
vllm_group_port51216vLLM 权重更新组(NCCL)端口
vllm_gpu_memory_utilization0.3colocate 模式下 vLLM 引擎的 GPU 显存占用比例
vllm_max_model_lengthNonecolocate 引擎最大序列长度
vllm_tensor_parallel_size1colocate 引擎的张量并行度

训练与日志

参数默认值说明
beta1.0广义 JSD 插值系数,0.0=前向 KL,1.0=反向 KL,0.5=JSD;范围校验[0, 1],越界直接抛ValueError
max_tool_calling_iterationsNoneAgent 训练时工具调用轮数上限;None表示无限制,模型生成无工具调用的回复轮即停止
log_completionsFalselogging_steps记录一批 (prompt, completion) 样本,可用 rich 打印、wandb/trackio 记录并保存 parquet
num_completions_to_printNonerich 打印的 completion 数量,None表示全部
log_unique_promptsFalse日志中是否只保留唯一 prompt

另外,DistillationConfigTrainingArguments的几个默认值做了覆盖:logging_steps默认10gradient_checkpointing默认True、未显式设置 fp16 时bf16默认Truelearning_rate默认1e-6

源码中__post_init__还校验了序列并行不兼容性:蒸馏需要在生成后于 trainer 内部构建模型输入,因此 Transformers 的 context-parallel / Ulysses 序列并行(cp_size > 1sp_size > 1)暂不支持,会直接报错提示设置为 1 或关闭parallelism_config

五、训练日志指标

训练与评估过程中记录的指标如下(由 distillation_trainer.py 的_generatecompute_loss产生):

  • num_tokens:迄今处理的 token 总数(含 prompt 与 completion);使用工具时只统计非工具 token;
  • step_time:每个训练步平均耗时(秒,含生成);
  • completions/mean_lengthcompletions/min_lengthcompletions/max_length:生成 completion 的平均/最小/最大长度(工具场景只统计非工具 token);
  • completions/mean_terminated_lengthcompletions/min_terminated_lengthcompletions/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 的场景。

  1. 启动 vLLM server:
VLLM_SERVER_DEV_MODE=1 vllm serve <model_name> \ --weight-transfer-config '{"backend": "nccl"}' \ --logprobs-mode processed_logprobs \ --max-logprobs -1
  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。

使用DistillationConfigmax_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_frequencytools/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=Falsedataloader_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),仅供参考

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

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

立即咨询