论文

LongStraw: 在固定GPU预算下超越200万Token的长上下文强化学习

LongStraw: 在固定GPU预算下超越200万Token的长上下文强化学习

推理系统的上下文长度正接近百万token,但后训练强化学习(RL)的工作量通常仍停留在256K token以下,依赖部署时的长度泛化。这种差距对AI智能体尤为重要,因为其观测、工具输出、文档和先前的决策会随长轨迹累计。 LongStraw 是一个架构感知的执行栈,在固定GPU预算下实现百万token级别的RL后训练,并基于GRPO实例化。它评估共享提示时不进行自动求导(autograd),仅保留后续token所需的模型特定状态,并逐步重放短响应分支,以牺牲额外重放时间为代价缩减实时训练图。 在8块H20 GPU上,LongStraw完成了对Qwen3.6-27B(混合循环与全注意力)的2.1M位置分组评分与响应反向传播(backward),组大小为2和8;增大组规模仅增加0.21 GB峰值内存,压力测试更达到4.46M位置。在32块H20上,针对GLM-5.2(压缩注意力混合专家模型)的78层,验证了2.1M token提示的端到端执行路径。 这些实验证明了执行容量,但尚未验证完整的训练正确性,因为捕获的提示状态已分离,且部分分布式前向与梯度组合路径尚未完善。

论文精读

TL;DR LongStraw 在固定 GPU 预算下,通过将共享 prompt 评估与梯度计算分离、重放短响应分支,将 GRPO 后训练的上下文扩展至 200 万 token 以上,并在 Qwen3.6-27B 和 GLM-5.2 混合模型上验证了 210 万 token 的执行能力。

问题

长上下文强化学习后训练 是构建能处理超长轨迹 AI 代理(agent)的关键环节。随着推理系统上下文窗口突破百万 token,后训练工作负载仍普遍停留在 256K 以下,主要依赖部署时的长度泛化。

现有方法局限

当前 RL 后训练框架(如 OpenRLHF、veRL)为通用训练设计,缺乏对长上下文特化:

  • 内存爆炸:全注意力计算图显存随序列长度平方增长,无法在固定 GPU 预算下扩展到百万 token。
  • 重复计算:分组评分(group scoring)时,多条响应分支会反复重算共享提示(prompt)的激活,浪费算力。
  • 架构假设僵硬:主流框架假定完整前向与反向图常驻内存,难以利用混合架构(如循环 + 注意力、压缩注意力)的稀疏性。

为什么这个问题难/重要

挑战在于 解耦提示与响应的计算图 并保持训练正确性。提示部分往往极长但被多条响应共享,必须在不破坏梯度流的前提下重放。此外,现代大模型混合了全注意力、滑动窗口、MoE 等机制,要求执行栈深度感知层间依赖。业界关注度极高:从 AI 研究员到产品经理都清楚,若无法在训练中直接优化长上下文表现,依靠推理时的长度泛化是不可靠的,会导致 agent 在长任务中逐步遗忘或决策失真。

行业类比

类似视频大模型训练中,长序列将自回归成本从 O(n) 推到 O(n²),必须用记忆重放或稀疏注意力才能支撑起端到端训练。

核心洞察

  • **训练与推理上下文长度之间的鸿沟**:长上下文推理系统已迈向百万 token 级,但 RL 后训练仍停留在 256K 以下,依赖部署时的长度泛化。LongStraw 首次在固定 GPU 预算下,将 GRPO 后训练的执行容量推到 2.1M token,揭示架构感知执行栈是弥合该鸿沟的关键路径,而非仅靠模型架构改进。
  • **分离共享提示与短响应重放的内存策略**:LongStraw 将共享提示的前向计算剥离出 autograd 图,仅保留后续 token 所需的模型特定状态,再逐条重放短响应分支进行训练。这与全图保留或序列并行方案相比,在 8×H20 GPU 上仅增加 0.21 GB 峰值内存即可支持组大小从 2 增至 8,体现出针对混合循环与全注意力模型进行状态裁剪的有效性,为 agent 长轨迹训练提供实用基线。

方法

LongStraw 针对 强化学习后训练 (RL post-training) 中上下文长度远小于推理阶段的问题,提出一种在固定 GPU 预算下支持 百万级 token 上下文 的架构感知执行栈。其核心思想是:利用共享提示的一次性无梯度预计算,将长计算图拆解为短响应分支的逐一重播,从而大幅降低显存占用。

输入与预处理

  • 输入为超长提示(可达 2M+ token)与多个短响应分支,例如在 GRPO (Group Relative Policy Optimization) 中,同一提示会生成多个候选响应用于组内比较。
  • 提示预计算:对共享提示执行一次无梯度的前向传播,仅保存后续 token 生成所需的 模型特定状态(如 KV 缓存、递归状态等),不构建计算图。这一步将长上下文的显存占用从 O(N) 压缩为仅保留状态,同时避免了主训练循环中的重复计算。

关键模块:分支重播与梯度计算

  • 每个响应分支独立重播:基于保存的提示状态,依次对每个短响应分支执行带梯度的前向/反向传播。由于响应通常远短于提示,计算图尺寸大幅减小,峰值显存占用几乎不随组大小增长(实验显示组大小从 2 增至 8 仅增加 0.21 GB)。
  • 架构感知优化:针对不同模型特性调整状态保存与重播策略。例如,对于混合递归与全注意力模型 (Qwen3.6-27B) 和压缩注意力 Mixture-of-Experts (GLM-5.2),LongStraw 能识别并利用冗余结构,进一步减少需保存的状态量。

输出与执行能力验证

  • 在 8 块 H20 GPU 上实现 Qwen 模型 2.1M 位置的分组评分与反向传播,压力测试更达到 4.46M 位置;在 32 块 H20 上跑通 GLM-5.2 全 78 层的端到端路径。当前验证了执行能力,但提示状态为 detached,部分分布式前向/梯度组合流程尚未完整,后续需完善训练正确性。

与同类方法的差异

现有长上下文 RL 方案多依赖长度泛化(在短上下文训练、部署时外推),或需大量 GPU 并行保存完整计算图;LongStraw 真正在训练时处理百万级 token,通过计算换空间策略,使单次训练所需硬件门槛大幅降低,更贴近实际可部署的 AI Agent 长轨迹学习场景。

实验

实验设计

LongStraw 的实验旨在验证执行容量而非完整训练正确性,因为捕获的提示状态是脱离梯度的,部分分布式前向和梯度合成路径尚未完成。

  • 模型与硬件:在 Qwen3.6-27B(混合循环注意力和全注意力)和 GLM-5.2(压缩注意力混合专家)上验证,使用固定数量的 NVIDIA H20 GPU(8 块或 32 块)。
  • 任务设置:构造超长提示(包含共享前缀和多个短响应分支),用 GRPO 算法进行评分和反向传播。核心操作包括:评估共享提示时不启用自动微分,仅保留后续令牌需要的模型特定状态,然后逐个重放短响应分支以减少活训练图。

关键发现

  1. 分组评分与反向传播:在 8 块 H20 上,Qwen3.6-27B 对 group=2group=8 均成功完成 2.1M 位置的评分与响应反向传播;组大小从 2 增至 8 时,峰值分配内存仅增加 0.21 GB,显示出近乎恒定的内存开销。
  2. 极限长度测试:单独的压力测试将序列推至 4.46M tokens,远超 2M 目标。
  3. 端到端路径:在 32 块 H20 上,GLM-5.2 所有 78 层处理了含 2.1M token 提示的完整执行路径。

启示与对比

  • 传统 RL 后训练受限于 256K 上下文,LongStraw 将可执行上下文提升近一个数量级,且未增加 GPU 数量,对长轨迹 Agent 应用尤为重要
  • 其架构感知设计(区分混合注意力、压缩注意力等)表明,未来超长上下文 RL 训练必须与模型内部结构深度结合,泛用的内存节省策略可能无法发挥最优效果。
  • 当前工作只证明了执行可行性,距离生产级训练尚有差距(如梯度同步、完整分布式组合),但为后续研究提供了清晰的工程基线。

行业影响

落地场景

LongStraw 使超长上下文(>2M tokens)的 RL 后训练在固定 GPU 预算下可行,直接解锁需要全序列推理的复杂 AI 智能体(agent)场景,例如:

  • 代码库级智能代码生成:阅读整个大型仓库(数万行代码、文档、提交历史)后修改或新增功能,需要完整上下文才能保持一致性。
  • 多轮对话式企业数据分析:知识库包含数万份合同、邮件、报表,LLM 需同时引用历史查询与全量文档才能给出无遗漏的对比总结。
  • 金融审计与合规:处理长达数年的交易日志,每笔交易可能依赖上下文中的监管规则、此前判断,模型必须在完整轨迹上学习推理。

商业价值

商业价值体现在 降低长上下文 RL 训练的硬件门槛扩大长文场景的服务范围

  • 降本:使用 8×H20(消费级/低端数据中心 GPU)即可运行超 2M tokens 的 GRPO 训练,使中小团队也能以可控成本自训长上下文 agent,避免租用昂贵的大规模集群。
  • 增收:托管式长文 agent 服务可覆盖更高价值场景,如法律全文审查、保险理赔全文档分析,定价可高于常规 token 生成服务。
  • 体验提升:长上下文 RL 训练的模型可比单纯依赖位置外推的模型更稳定地遵循长轨迹中的复杂指令,减少 agent 在执行长期任务时的偏离与重复,提升用户信任与留存。

与现有产品/工作流的接口

LongStraw 的核心是 一次前向缓存 + 逐条响应回放,可以无缝嵌入现有 RLHF/GRPO 管线:

  • 训练框架侧:可直接嫁接在 OpenRLHF、trl、VERL 等主流框架的 policy loss 计算之前,替换原有的全批量前向 pass,仅需将 prompt 的前向结果缓存为 detached state,响应分支再回放。
  • 模型侧:支持混合架构(如 Hybrid-RNN/Full-Attention 的 Qwen3.6-27B)与 MoE 压缩注意力(GLM-5.2),表明方法具有模型通用性,可适配自研模型。
  • 部署工作流:训练得到的模型仍为标准 checkpoint,无需额外推理定制,可直接部署至 vLLM/SGLang 等推理框架,服务长上下文 agent 应用。

具体落地 use case

  1. 企业级智能 IT 运维 agent:在混合云环境中,agent 需收集过去 48 小时内数万条监控日志、工单、配置变更记录,定位故障根因并提出修复步骤。LongStraw 可让模型在 RL 阶段直接从真实的长上下文轨迹中学习,而非仅靠短序列模拟,显著提高诊断准确率。
  2. 电商平台全流程购物助手:用户对话可能持续数百轮,跨越商品搜索、多商家比价、订单追踪、售后维权,助手需记住所有历史偏好和平台政策。通过 LongStraw 训练,模型可以在完整对话历史上被优化,避免遗忘早期需求,提升成交转化率和售后满意度。

局限

  • 训练正确性未验证:论文明确声明当前仅展示了执行能力(execution capacity),而非完整的训练正确性,因为捕获的提示状态是分离的(detached),部分分布式前向和梯度组合路径尚未完成。这意味着无法保证该方法在真实RL训练中能够有效收敛或产出与标准训练一致的模型质量,后续需完善分布式梯度组合与完整反向传播路径。
  • 架构依赖与泛化受限:LongStraw 是架构感知的执行栈,目前仅针对混合循环与全注意力(Qwen3.6-27B)和压缩注意力 MoE(GLM-5.2)两种特定架构进行了实现与验证。对于更广泛使用的纯 Transformer 或其它非循环模型,状态捕获和重放机制可能需要大幅调整甚至重新设计,其通用性尚待检验。
  • 额外重放时间可能抵消显存收益:方法通过重放短响应分支来减小训练图显存占用,但引入了额外重放时间(additional replay time)。论文未报告端到端训练吞吐量,仅展示了峰值显存节省。在实际大模型 RL 训练中,多次重放可能显著增加 wall-clock 时间,使得总训练效率不一定优于其它重计算或序列并行方案。
论文Changhai Zhou2026-07-16原文

相关内容