ARTICLE DETAIL

资讯详情

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

TRL 降低显存占用实战指南:截断、Packing、Liger、Chunked CE 与激活卸载全解析

TRL 降低显存占用实战指南:截断、Packing、Liger、Chunked CE 与激活卸载全解析 TRL 降低显存占用实战指南截断、Packing、Liger、Chunked CE 与激活卸载全解析【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl本文以 TRL 官方文档 Reducing Memory Usage 为主线系统梳理 TRL 提供的全套显存优化技术从max_length截断与 Packing 数据组织到 Liger Kernel、chunked cross-entropy、padding-free 前向、激活卸载、ZeRO-3 权重 gather 控制与 vLLM sleep mode并结合仓库源码逐一说明每个开关的默认值、底层实现与适用限制帮助你在有限 GPU 资源上跑起更大模型的训练任务。一、截断Truncation控制序列长度是第一道防线训练数据中的序列长度往往差异很大。批量构建时短序列会被填充padding到批内最长序列的长度即使大多数序列都很短显存占用也会很高。因此将序列截断到一个合理长度是降低显存的第一步——TRL 各 Trainer 默认就会对序列做截断但截断长度通常需要你按具体任务调整。各 Trainer 的截断入口都是max_lengthDPOmax_length截断的是 promptcompletion 拼接后的完整序列from trl import DPOConfig training_args DPOConfig(..., max_length...)需要注意旧版的max_prompt_length与max_completion_length参数已被移除。如果数据集存在超长 prompt/completion应在训练前对数据集做过滤或预截断而不是依赖这两个已废弃的参数。SFTmax_length直接作用于输入序列默认值为 1024见 SFTConfig 中max_length: int | None field(default1024, ...)超过max_length的序列按truncation_mode默认keep_start截断设为None则不做截断。from trl import SFTConfig training_args SFTConfig(..., max_length...)如何选择合适的max_length这是一个典型的权衡问题设置过小大量 token 被丢弃无法参与训练损失数据信息设置过大显存占用陡增可能触发 OOM且在没有 Packing 或 padding-free 的情况下大批量 padding token 会导致训练效率低下。TRL 官方文档提供了一个可视化工具数据集序列长度分布分析器见原文档中的 iframe 嵌入来观察你的数据集中序列长度的分布据此选择能覆盖绝大多数样本的max_length。二、Packing把多条短序列拼进同一训练行该技巧仅适用于SFT训练且要求使用FlashAttention或其变体。截断有两个天然缺陷一是信息丢失序列尾部的重要 token 被丢弃二是截断长度两难过短丢数据、过长低效。Packing 通过将多条序列组合到同一条训练行、把每行填满到max_length来缓解这两个问题。从源码看sft_trainer.py 与 L1247启用 packing 时 Trainer 会检查是否使用 FlashAttention未使用时会给出相应提示——这与文档中仅支持 FlashAttention 环境的限制一致。TRL 的 packing 基于Best-Fit DecreasingBFD装箱算法实现核心函数是 data_utils.py 中的pack_dataset它把数据集序列打包进长度为seq_length即max_length的块中序列超长时按不同策略处理溢出 token。共支持三种策略SFTConfig 中choices: [bfd, bfd_split, wrapped]bfd默认Best-Fit Decreasing 装箱。若单条序列超过max_length溢出 token 直接丢弃bfd_split同样是 BFD 装箱但超长序列会被切分为多个≤ max_length的片段再参与装箱保留全部 token对应 Fewer Truncations Improve Language Modeling 论文中的做法。源码中旧的bfd-requeue名称已被重命名为bfd_splitSFTConfig.post_init中会对旧名发出弃用警告并自动改写wrapped把所有 token 连成一条流再按固定长度切块padding 最少但会打断序列边界、把不同样本混在一起会破坏大部分数据的序列连续性据 Qwen3-Coder-Next 技术报告讨论这可能伤害模型表现。注意如果所有序列都短于max_lengthbfd与bfd_split的行为完全一致因为既不需要截断也不需要切分。pack_dataset的 docstring 中还给出了一个直观的示例摘自 trl/data_utils.py from datasets import Dataset from trl import pack_dataset examples { ... input_ids: [[1, 2, 3, 4, 5], [6, 7], [8, 9, 10], [11]], ... attention_mask: [[1, 1, 1, 0, 0], [1, 0], [1, 1, 0], [1]], ... } dataset Dataset.from_dict(examples) # 默认 bfd截断超长序列 packed pack_dataset(dataset, seq_length4, strategybfd) {input_ids: [[1, 2, 3, 4], [8, 9, 10, 11], [6, 7]], ...} # bfd_split保留全部 token packed pack_dataset(dataset, seq_length4, strategybfd_split) {input_ids: [[1, 2, 3, 4], [8, 9, 10, 5], [6, 7, 11]], ...}完整用法from trl import SFTConfig training_args SFTConfig( ..., packingTrue, packing_strategybfd, max_length512, )一个从源码可确认的联动行为当packingTrue且策略为bfd/bfd_split时padding-free 会被自动开启与padding_free参数取值无关sft_trainer.py 中self.padding_free args.padding_free or (args.packing and args.packing_strategy in {bfd, bfd_split})。三、PEFT参数高效微调大幅降低显存PEFTParameter-Efficient Fine-Tuning方法如 LoRA 是降低训练显存最有效的手段之一不训练全部模型参数只训练少量 adapter 参数从而显著降低显存需求让有限硬件也能微调更大的模型。from datasets import load_dataset from peft import LoraConfig from trl import SFTTrainer dataset load_dataset(trl-lib/Capybara, splittrain) peft_config LoraConfig() trainer SFTTrainer( modelQwen/Qwen2.5-0.5B, train_datasetdataset, peft_configpeft_config, )PEFT 还可以与 4-bit / 8-bit 量化叠加获得更大的显存节省。完整的 adapter 方法、量化配置说明见 PEFT Integration。四、Liger Kernel降低峰值显存Liger Kernel 是一组专为 LLM 训练设计的 Triton kernel。根据仓库文档的介绍它能将多卡训练吞吐量提升约 20%、显存占用降低约 60%。在 TRL 中通过统一的use_liger_kernelTrue开关启用支持多个 Trainerfrom trl import SFTConfig training_args SFTConfig(..., use_liger_kernelTrue)from trl import DPOConfig training_args DPOConfig(..., use_liger_kernelTrue)from trl import GRPOConfig training_args GRPOConfig(..., use_liger_kernelTrue)from trl import KTOConfig training_args KTOConfig(..., use_liger_kernelTrue)from trl.experimental.gkd import GKDConfig training_args GKDConfig(..., use_liger_kernelTrue)更细节的集成说明见 Liger Kernel Integration。注意它与下文 chunked cross-entropy 存在互斥关系见第五节。五、Chunked Cross-EntropySFT 默认的 logits 显存优化在大词表模型中LM head 输出的[batch × seq_len × vocab]logits 张量是前向/反向期间最主要的常驻激活之一。SFTConfig的loss_typechunked_nll避免了它的整体物化labels -100的位置在lm_head矩阵乘法之前就被丢弃只对有效 token 做投影交叉熵按 token 分块chunk计算并用梯度检查点使每一块的 logits 只在自己的前向/反向窗口内存活。因此峰值 logits 显存从(batch × seq_len) × vocab_size降为chunk_size × vocab_size。从实现看_chunked_cross_entropy_loss按chunk_size遍历有效 token每个[chunk_size, vocab_size]的 logits 仅在自身的梯度检查点区间内保留。它与标准nll损失数学等价——这是显存优化而非新的损失函数。在 [SFTTrainer] 中它已是默认值SFTConfig.post_init中未显式指定loss_type时默认取chunked_nll若use_liger_kernelTrue则默认取nll二者不兼容。若要退出默认路径from trl import SFTConfig training_args SFTConfig(..., loss_typenll) # opt out of the default chunked path按仓库文档给出的实测数据大词表模型上峰值 VRAM 通常可降约 30%、最高约 50%在Qwen3-1.7B、词表约 151k 上测得单卡约 30%FSDP2 × 4 卡下最高约 50%训练时长通常持平或略快loss_type还支持dftDynamic Fine-Tuning。限制条件源码与配置文档均可确认与use_liger_kernelTrue不兼容sft_trainer.py 中直接raise ValueError与 PEFT 不兼容sft_trainer.py 提示lm_head被 PEFT adapter 包装时不支持与 VLM 不兼容。六、Padding-free批内展平消除 padding 开销Padding-free 批处理是另一种减少显存占用的思路先采样一个 batch再把批内所有序列展平成一条连续序列做前向从而完全避免 padding。与 packing 不同padding-free 保证每条序列完整不被拼接破坏。强烈建议搭配FlashAttention 2 或 3使用否则可能出现 batch 之间注意力相互污染的问题。from trl import SFTConfig training_args SFTConfig(..., padding_freeTrue, model_init_kwargs{attn_implementation: kernels-community/flash-attn2})from trl import DPOConfig training_args DPOConfig(..., padding_freeTrue, model_init_kwargs{attn_implementation: kernels-community/flash-attn2})需要注意 DPO 的一个现状由于 DPO 重构DPOTrainer中padding_freeTrue暂时不可用——设置后只会发出警告并回退到标准 padding未来版本计划恢复。这一点在 dpo_trainer.py 中可以确认检测到padding_freeTrue时打印 temporarily unavailable after a refactor 的警告并把内部标志重置为False。七、激活卸载Activation Offloading用 CPU 内存换 GPU 峰值显存激活卸载通过在前向时把激活张量临时搬到 CPU RAM、仅在反向需要时再取回来降低 GPU 峰值显存代价是训练时间略有增加。启用方式from trl import SFTConfig training_args SFTConfig(..., activation_offloadingTrue)默认值为FalseSFTConfig。其底层实现在 trl/models/activation_offloading.py几个值得了解的细节基于 PyTorch 的saved_tensors_hooks实现OffloadActivations类L83在前向中拦截被 autograd 保存的激活超过min_offload_size默认 1024 字节的张量才会被搬到带 pin memory 的 CPU 内存小张量不值得搬运见 get_act_offloading_ctx_manager 的参数说明默认使用独立的 CUDA 流use_streamsTrue要求 torch ≥ 2.5.0把 CPU-GPU 传输与计算重叠并配合max_fwd_stash_size默认 5在前向中控制缓存深度在传输/计算重叠度与内存占用之间做权衡对lm_head这类输出头会自动豁免卸载通过NoOpManager注册 forward hook见 L736-L739因为它的激活刚卸载就要立即取回、大词表下代价极高对 FSDP v2 的DTensor参数做了特殊处理按 storage 指针去重以避免重复卸载。八、Padding Sequences to a Multiple把长度对齐到硬件友好边界目前支持SFT与Reward两种 Trainer。启用后所有序列会被填充到指定数值的整数倍。这在部分硬件上能通过对齐内存友好的长度边界来提升计算效率from trl import SFTConfig training_args SFTConfig(..., pad_to_multiple_of2048)from trl import RewardConfig training_args RewardConfig(..., pad_to_multiple_of2048)对应字段见 SFTConfig.pad_to_multiple_of默认None即不对齐。九、禁用 ZeRO-3 生成时的权重 gather使用DeepSpeed ZeRO-3时模型权重分片分布在多张 GPU 上。GRPO、RLOO、Online DPO 等在线方法在训练过程中需要模型在线生成 completion默认会把权重临时 gather 到单卡上生成——对超大模型这一步可能直接 OOM对应 TRL issue #2250。遇到该问题可关闭生成时的 gatherfrom trl import GRPOConfig training_args GRPOConfig(..., ds3_gather_for_generationFalse)from trl.experimental.online_dpo import OnlineDPOConfig training_args OnlineDPOConfig(..., ds3_gather_for_generationFalse)from trl import RLOOConfig training_args RLOOConfig(..., ds3_gather_for_generationFalse)该参数在 GRPOConfig 中默认值为True。代价是生成速度会变慢但能规避 gather 导致的 OOM。十、vLLM sleep mode优化步把 vLLM 权重与缓存下沉到 CPU当使用vLLM作为在线训练方法的生成后端时可开启 sleep mode在优化optimizer step阶段把 vLLM 的参数与 KV cache 卸载到 CPU RAM等到权重同步与下一轮生成时再重新载入 GPU 显存。from trl import GRPOConfig training_args GRPOConfig(..., vllm_enable_sleep_modeTrue)from trl import RLOOConfig training_args RLOOConfig(..., vllm_enable_sleep_modeTrue)该参数vllm_enable_sleep_mode默认False见 GRPOConfig。它把 GPU 显存占用压得更低特别适合大模型或资源受限场景代价是唤醒 vLLM 引擎会带来 host-device 传输延迟训练速度会略有影响。十一、梯度检查点Gradient Checking以算力换显存梯度检查点不在前向中保存全部中间激活而在反向时重算是经典的以算力换显存手段from trl import SFTConfig training_args SFTConfig(..., gradient_checkpointingTrue)在 TRL 中所有 Trainer 默认开启gradient checkpointing 以优化显存SFTConfig 的 NOTE 明确列出gradient_checkpointing默认True与transformers.TrainingArguments的默认False不同。如需关闭显式设置gradient_checkpointingFalse。更多通用技巧如与transformers生态配合的调优可参考transformers的性能指南。十二、组合建议与速查官方文档的核心建议是不同技巧之间可以组合建议针对自己的硬件与数据实验不同搭配找到最优配置。结合本文各节的源码事实可整理为如下速查表技术关键参数适用 Trainer默认值主要限制截断max_length各 TrainerSFT 默认 1024需按数据分布选择PackingpackingTrue, packing_strategybfd仅 SFTFalse/bfd需 FlashAttentionPEFT/LoRApeft_config各 Trainer—与 chunked_nll 不兼容Ligeruse_liger_kernelTrueSFT/DPO/GRPO/KTO/GKDFalse与 chunked_nll 互斥Chunked CEloss_typechunked_nllSFTSFT 默认启用不兼容 Liger/PEFT/VLMPadding-freepadding_freeTrueSFTDPO 暂不可用False需 FA2/FA3激活卸载activation_offloadingTrueSFTFalse需要足够 CPU RAM长度对齐pad_to_multiple_of2048SFT、RewardNone—禁 ZeRO-3 gatherds3_gather_for_generationFalseGRPO/RLOO/Online DPOTrue生成变慢vLLM sleepvllm_enable_sleep_modeTrueGRPO/RLOOvLLM 后端False唤醒有传输延迟梯度检查点gradient_checkpointingTrue全部 TrainerTrue重算增加耗时典型组合思路短序列为主的 SFT 用Packingbfd/bfd_split超长上下文优先截断 padding-freeFA2/FA3显存瓶颈在优化步的在线 RLGRPO/RLOO叠加vLLM sleep mode与 ZeRO-3 gather 控制大词表模型在 SFT 上保留默认chunked_nll若改用 Liger 则显式设置loss_typenll避免冲突。所有参数与行为均可在当前仓库中直接查证配置项定义见 SFTConfig、GRPOConfig实现见 sft_trainer.py、dpo_trainer.py、data_utils.py 与 activation_offloading.py对应测试可参考 tests/test_sft_trainer.py、tests/test_dpo_trainer.py 与 tests/test_activation_offloading.py。【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表