AREX Feed Article
TRL v1.13 发布:新增 1M+ token 长上下文后训练指南
9 月 10 日,Hugging Face 联合创始人 Thomas Wolf()在 X 上,开源后训练库 发布 v1.13。按帖文说法,这一版把重点放在长上下文训练上:新增一份关于如何后训练 1M+ token 上下文模型的指南,并带来速度与内存占用方面的多项改进。
随版本发布的给出了一个可运行示例和一组实测数字:在单节点 8×H100 上训练 Qwen3-8B,单条序列 1,048,576 token(词元),380 秒/步,每 GPU 峰值 56.2 GB(bf16)。v1.13.0 当天出现在 与 上。文中数字均出自项目自己的发布说明与文档,属官方口径。
四道坎:分块损失、位置重标定、激活卸载、序列切分
指南的标题是「Training Beyond 1M Tokens」。它给出的场景是:agent(智能体)可以连续数小时读文件、跑命令,把会话的每一轮都带着往前走,累积出几十万 token 的上下文,这也是前沿模型开始宣传百万 token 级窗口的原因。要让模型真正擅长这么长的上下文,就得用同样长的序列去训练;麻烦在于,一条百万 token 的序列装不进一张 GPU,八张也不行,除非采取一些手段。
第一道坎是损失函数。序列变长后,最先吃满显存的就是损失计算本身:最后一层与语言建模头相乘会产生一个 logits 矩阵(对词表里每个词的未归一化分数),大小是序列长度 × 词表规模,而词表超过十万条;把它完整算出来并留在显存里,正是显存峰值的来源。TRL 的做法是把 logits 按行切块,每次只计算 256 行的损失再求和,整块矩阵不会完整地占据显存;这已是默认设置(loss_type="chunked_nll"),要回到旧行为需显式改成 "nll"。以 Qwen3-4B、单张 80 GB 卡为例,指南称仅这一项改动就能把显存允许的序列长度从 32k token 推到 160k。
第二道坎是位置。模型只认得训练时见过的位置编号:演示所用的 Qwen3-4B 训练长度是 40,960 token,再往后的位置它没学过任何含义。指南测得模型在 70k 左右开始丢失线索,loss 一路攀升,最终高过整段文本的一元熵(6.45)——比完全无视上下文、只按词频预测还差。解法是用 YaRN 重标定位置:把 100,000 号位置映射回 25,000,让每个 token 都落在模型熟悉的范围内;40,960 的训练长度对应 163,840 的目标长度,缩放因子取 4。重标定后,同一条 160k token 序列的 loss 曲线保持平直,最后 20k token 的平均 loss 为 2.8,不缩放时是 7.3。
第三道坎是激活值。梯度检查点(gradient checkpointing)会丢掉每层的大部分中间结果、在反向传播时重算,只保留每层一个「序列 × 隐层」的张量;但在长序列下,这些保存张量仍是显存占用的大头。它们在正向计算时写入、要到很晚才被读取,因此可以挪到 CPU 内存、等反向传播用到时再取回。对应开关是 gradient_checkpointing_kwargs={"offload": True};指南给出的结果是同一配置下峰值显存从 59.9 GB 降到 48.8 GB,单卡可训长度从 160k 再推到 256k。代价是每个保存张量都要在主机与 GPU 之间往返搬运。
第四道坎是单卡容量本身。到百万 token 级别,仅一层 MLP(多层感知机)在反向重算时同时存在的四个中间张量就要 76 GB(每个按 1,048,576 × 9,728 的规模计算),而 80 GB 的卡还要先分给模型和优化器 30 GB。此时要把序列分到多张卡上:四张卡各拿 262,144 个 token,张量随之缩到 19 GB。注意力要求每个 token 都看到之前所有 token,TRL 支持两种交换方式:上下文并行(CP,按 token 切分)与 Ulysses 序列并行(SP,按注意力头切分)。指南给出的步时数据是:131k token 时,2 张卡 34.6 秒/步,4 张卡 17.8 秒/步,8 张卡 9.5 秒/步;每翻一倍卡数,步时接近减半。
指南的总结是:单卡停在 256k 出头,四张卡才能放下整条百万 token 序列。
可运行的示例:Qwen3-8B 与拼成百万 token 的书籍序列
指南配套的示例放在 TRL 仓库的 examples/sft_qwen3_8b_1m_context/ 目录下,配一份 context_parallel_8gpu.yaml 配置。它把 PG-19 数据集里的书籍首尾相接,拼成每条约 100 万 token 的序列,用来微调 Qwen3-8B,运行方式是一条 accelerate launch 命令。指南展示的第一步日志里,loss 是 4.311、grad_norm 是 29.25;指南同时提醒,长上下文配置稍有偏差的运行,起始 loss 会落在 10 左右,而不是 4.3。这一步只有一条序列:每个 token 都要注意到之前所有 token,而不是一批短序列拼成的 batch(批次)。
指南写明:示例需要一台 8×H100(或更好)的节点。指南还给出一个运行提示:在百万 token 规模上要设置 PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,否则张量又大又要求连续,分配器可能在显存总量足够的情况下仍找不到空间、直接失败。
速度与显存:chunked loss 重回张量核心
速度与显存方面的改进,主要来自对默认损失函数内层循环的一次修复。原来的实现写的是 h.float() @ w.float().t():两个操作数本来就是 bf16,升到 fp32 没有带来任何信息,却把矩阵乘法从张量核心移到了 fp32 SIMT 路径,还会生成整个 lm_head 权重的 fp32 副本——对 248k 词表来说是 2.03 GB,而且每个 chunk(分块)、每次梯度检查点重算都要重建一遍。官方在对 Qwen3.6-35B-A3B 的一次 8×H100 trl sft 剖析中发现,这两个 fp32 GEMM 占全部 GPU kernel(计算内核)时间的 21.6%。
修复后的单 chunk 基准(256 token × 248,320 词表 × 2,048 隐层,1×H100,bf16,前向加反向)是:耗时从 23.37 毫秒降到 3.86 毫秒(约 6 倍),峰值显存从 5.99 GB 降到 3.03 GB。端到端(每步 16,384 token)的 tokens/s/GPU 提升为:gemma-3-270m 全参微调 24,036→31,609(1.32 倍),Qwen3-0.6B 全参微调 21,276→26,641(1.25 倍),Qwen3-8B 全参微调 3,554→6,009(1.69 倍),Qwen3-8B LoRA r16(低秩适配)4,531→7,125(1.57 倍),Qwen3-30B-A3B(MoE,混合专家)4,251→5,101(1.20 倍)。同一修复也用在了蒸馏 trainer(训练器)上,后者每个 chunk 要算两次,学生和老师各一次。发布说明称,在 accelerate 混合精度下数值逐位一致。
v1.13 的其他变化:PPO 移除、融合损失整合、依赖下限上调
破坏性变更方面,PPO(近端策略优化)被整体移除:PPOTrainer、PPOConfig 与 modeling_value_head.py 里的三个包装类一起消失,连带它们的测试、文档页和示例。发布说明交代了背景:PPOTrainer 落在 2020 年 3 月 28 日的提交 dfb6a580,那是这个库的第一个提交,当时包名还叫 lm_ppo;它是 TRL 里最老的一件东西,也是原始代码库的最后残留。移除理由包括:超过一年没有功能开发、唯一没有对齐输入格式的 trainer、记录到的使用量接近于零,以及持续引来自动化 bug 猎手对没人运行的代码提交真实报告。对现有用户,发布说明称 from trl import PPOTrainer 自 v1.10 起就已失效,真正在跑 PPO 的人本来就固定在旧版本上,那些安装会继续按原样解析。
损失函数方面,TRL 对 liger-kernel 唯一的依赖 liger_kernel.chunked_loss(DPO、KTO、GRPO、JSD 的融合线性损失)被收进了 trl.losses:FusedLinearDPOLoss、FusedLinearKTOLoss、FusedLinearGRPOLoss、FusedLinearJSDLoss 四个类,复制自 Liger-Kernel v0.8.2,保留 BSD-2 声明,发布说明称其在随机输入上与已安装的 Liger 逐位一致。用法不变:use_liger_kernel=True 仍选择这套融合损失,也仍然需要安装 liger-kernel。
依赖下限同时上调:peft 至少 0.13.0,deepspeed 至少 0.18.6;新增对 vLLM 0.28.0 的支持,同时移除 vLLM 0.19.0。
适用边界:全注意力模型,CP 与 packing 互斥
指南明确排除了滑动窗口与分块注意力的模型:OpenAI GPT-OSS、Gemma 3 与 Gemma 4、Qwen3.5 及之后版本会被 accelerate 直接拒绝;示例与实验选用的 Qwen3、Qwen3-MoE 是全注意力架构,可以启用上下文并行。上下文并行还要求 causal SDPA(带因果掩码的缩放点积注意力),因此不能与 packing(样本拼接,靠块对角掩码防止文档互相读取)同时使用:TRL 会在同时请求两者时报错;序列长度还必须补齐到 cp_size 的两倍,四张卡时要在 SFTConfig 里设置 pad_to_multiple_of=8。
把这些条件叠在一起,能跑通百万 token 完整示例的配置可以概括为:模型必须是全注意力,序列切分与 packing 要二选一。并行方式还有一条约束:cp_size 对应 FSDP2 后端,sp_size 对应 DeepSpeed,两者目前不能交叉使用。