ARTICLE DETAIL

资讯详情

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

TRL DistillationTrainer 实战指南:基于广义 JSD 的 On-Policy 知识蒸馏

TRL DistillationTrainer 实战指南:基于广义 JSD 的 On-Policy 知识蒸馏 TRL DistillationTrainer 实战指南基于广义 JSD 的 On-Policy 知识蒸馏【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trlTRL 的DistillationTrainer实现了《On-Policy Distillation of Language Models: Learning from Self-Generated Mistakes》论文提出的 Generalized Knowledge DistillationGKD方法让较小的 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, splittrain) trainer DistillationTrainer( modelQwen/Qwen2.5-0.5B-Instruct, teacher_modelQwen/Qwen2.5-1.5B-Instruct, train_datasetdataset, ) 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)-混合。端点退化为纯散度beta0.0前向 KL (\mathbb{D}_{\mathrm{KL}}[p_T | p_S])beta1.0反向 KL (\mathbb{D}_{\mathrm{KL}}[p_S | p_T])默认值让 student 分布收窄去贴合 teacher 的高概率区是实际中最常用的选择beta0.5标准 JSD。注意这里的beta与 GRPO 的beta含义完全不同。GRPO 的beta是相对参考模型的 KL 惩罚系数而这里它直接选择散度本身没有参考模型 KL 惩罚见 distillation_config.py 中beta字段的注释。内存优化实现实际计算时vocab 投影与散度按 chunk 分块进行——有效位置被argsort打包到前部每 chunk 只保留[chunk_size, vocab_size]大小的 logitschunk 大小为 256见_CHUNKED_LM_HEAD_CHUNK_SIZE并配合梯度检查点torch.utils.checkpoint因此峰值激活内存不再随全 vocab × 序列长度的 logits 张量增长。详见 reducing_memory_usage.md。源码中散度计算对logit_scaleCohere 系与final_logit_softcappingGemma 系做了兼容处理且 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的 kwargsrevision也会用于加载 processing classteacher_model_name_or_pathNoneteacher 模型名或路径本地加载时使用teacher_model_revisionNoneteacher 的 revision分支、tag 或 commit hashteacher_model_init_kwargsNone实例化 teacher 时传给from_pretrained的 kwargstrust_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.2repetition_penalty1.0重复惩罚1.0鼓励新 token1.0鼓励重复generation_kwargsNone额外传给GenerationConfig/SamplingParams的 kwargs如suppress_tokens、num_beams与上述参数冲突时以它为准chat_template_kwargsNone传给apply_chat_template的额外 kwargscache_implementationNone非 vLLM 生成时的 cache 实现vLLM 加速参数默认值说明use_vllmFalse是否用 vLLM 生成 on-policy completionsvllm_modecolocatecolocate同进程共享 GPU或server独立进程/GPUHTTP 通信vllm_model_implvllmvLLM 后端vllm或transformersvllm_enable_sleep_modeFalse优化器步骤期间 offload student 权重的 sleep 模式vllm_server_base_urlNone若提供忽略 host/portvllm_server_host/vllm_server_port0.0.0.0/8000server 模式的 host/portvllm_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前向 KL1.0反向 KL0.5JSD范围校验[0, 1]越界直接抛ValueErrormax_tool_calling_iterationsNoneAgent 训练时工具调用轮数上限None表示无限制模型生成无工具调用的回复轮即停止log_completionsFalse每logging_steps记录一批 (prompt, completion) 样本可用 rich 打印、wandb/trackio 记录并保存 parquetnum_completions_to_printNonerich 打印的 completion 数量None表示全部log_unique_promptsFalse日志中是否只保留唯一 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使用工具时只统计非工具 tokenstep_time每个训练步平均耗时秒含生成completions/mean_length、completions/min_length、completions/max_length生成 completion 的平均/最小/最大长度工具场景只统计非工具 tokencompletions/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 1Colocate 模式默认。vLLM 在 trainer 进程内运行与训练模型共享 GPU 显存无需启动独立服务可提升 GPU 利用率但可能与训练竞争显存。from trl import DistillationConfig training_args DistillationConfig( ..., use_vllmTrue, # vllm_modecolocate by default )Option 2Server 模式。vLLM 在独立进程及独立 GPU中运行通过 HTTP 与 trainer 通信适合有专用推理 GPU 的场景。启动 vLLM serverVLLM_SERVER_DEV_MODE1 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_vllmTrue, vllm_modeserver, )⚠️ 警告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而不是训练整个 studentfrom datasets import load_dataset from trl import DistillationTrainer from peft import LoraConfig dataset load_dataset(trl-lib/ultrafeedback-prompt, splittrain) trainer DistillationTrainer( modelQwen/Qwen2.5-0.5B-Instruct, teacher_modelQwen/Qwen2.5-1.5B-Instruct, train_datasetdataset, peft_configLoraConfig(), ) trainer.train()⚠️ 警告蒸馏损失直接读取lm_head.weight并通过 backbone 前向_get_last_hidden_state绕过PeftModel.forward()。因此在lm_head上挂 adaptertarget_modules含lm_head会被拒绝——head 上的可训练 adapter 位于损失永远看不到的独立子模块中会静默得不到梯度源码中显式抛出ValueErrorPrompt 学习类方法PromptTuning、PrefixTuning、P-Tuning同样会被拒绝因为虚拟 token 通过PeftModel.forward()注入而损失直接调用 backbone 会漏掉它们如需训练 head请改用modules_to_save[lm_head]。另外源码对 PEFT DeepSpeed ZeRO-3 场景做了适配非量化模型下自动传autocast_adapter_dtypeFalse规避混合 dtype 的 TypeErrorZeRO-3 下强制use_reentrantTrue的梯度检查点并为 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-V3TRL 会在启用工具时自动替换为打了补丁的训练模板完整清单见 chat_templates.md。使用DistillationConfig的max_tool_calling_iterations限制工具调用轮数默认无限制student 生成不包含工具调用的回复轮即停止。⚠️ 警告暂不支持异步工具async请传入同步函数——源码中inspect.iscoroutinefunction检查会直接抛ValueError。另外启用工具要求 transformers ≥ 5.0.0且低版本需要jmespathtransformers ≥ 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。八、训练视觉语言模型VLMDistillationTrainer支持在含文本与图片的多模态数据集上蒸馏 VLMstudent 与 teacher 都传 VLM数据集为仅 prompt 格式带image列单图或images列多图列表。数据集结构见 dataset_formats.md。已在以下模型上验证Gemma 3——如google/gemma-3-4b-itLLaVA-NeXT——如llava-hf/llava-v1.6-mistral-7b-hfQwen2-VL——如Qwen/Qwen2-VL-2B-InstructQwen2.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_modeTrue通过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_batchesFalse与dataloader_num_workers0源码会强制覆盖并告警以保证生成批次的分组顺序断点续训安全_buffered_inputsNone时如从 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),仅供参考
返回列表