AI Agent 想"长记性",先撞穿显存墙——一篇论文说:别再只换模型,把训练执行图拆开,百万 token 的 RL 后训练在 8 张 H20 上就跑起来了

  • 关联论文:2607.14952

如果你是 AI Agent 平台的负责人,或者你在做长上下文训练栈,过去一年大概率被同一个问题反复折磨过:

为什么推理侧已经能跑 100 万 token,后训练侧仍然卡在 256K?

  • 你给 agent 配更长的 trajectory,prompt 越攒越长;
  • 你换更大的集群,显存还是先爆;
  • 你调 group size,reward 反而波动;
  • 你换算法、换数据,瓶颈纹丝不动。

你大概率以为:"百万 token 的 RL 后训练,只能等下一代 GPU。"

但 2026 年 7 月这篇叫 LongStraw(arXiv 2607.14952)的论文,正面打脸这个假设:

百万 token RL 训练,根本不是被算力卡住,是被显存卡住;而显存是被"共享 prompt 的反向传播图"白白放大的。把这张图显式拆开,峰值显存就能从 group_size × prompt 退化到 prompt + 单条 response,8 张 H20 就能跑 2.1M token 端到端。

最狠的是——代码已经开源(github.com/MindLab-Research/longstraw),你今天就能在自己 cluster 上跑一遍。

今天这篇科普用 5 分钟把它讲透:为什么"拆图"比"换 GPU"更值钱?Agent 平台为什么必须重做训练栈?

一、Agent 时代最痛的"显存墙"长什么样

在说论文之前,先把痛点摆清楚。

今天的 Agent trajectory 已经长到让人无语:一条对话里塞下工具输出、文档节选、历史决策、observation 累积,prompt 长度 50 万、100 万 token 是家常便饭。推理侧(inference)早就突破这个长度——GPT 系、Claude 系、Gemini 系纷纷支持 1M+ context。

RL 后训练(post-training)侧始终跟不上。瓶颈不在算法、不在数据,而在一个非常具体的系统事实:

GRPO(Group Relative Policy Optimization)训练里,同一个 prompt 要被 group_size 个候选 response 共享。一次 backward 跑下来,prompt 段的中间激活会被复制 group_size 份,激活显存随 group_size 线性爆炸。

直觉上你以为是"模型太大",其实显存里 80% 都是 prompt 段重复的反向传播中间结果——这些状态对每条 response 来说完全相同,却在被无脑地反复保留。

也就是说:百万 token RL 训练,是被一张"低效的图"卡住的,不是被算法卡住的。你想真把 Agent 训到 100 万 token,必须从执行图下手。

二、LongStraw 的反直觉结论:拆图,而不是换模型

论文的核心结论非常直接:

把"共享 prompt 的反向传播"从 autograd 图里显式拆开:prompt 段走 no_grad + 架构感知的 state 快照,response 段逐条 replay forward+backward,峰值显存基本只由 prompt 长度 + 单条 response 决定,与 group_size 几乎无关。

实测数据极有冲击力:

  • Qwen3.6-27B(hybrid recurrent + full attention):在 8 张 H20 上跑通 2.1M positions 的 grouped Qwen scoring + response backward;
  • GLM-5.2(compressed-attention MoE):在 32 张 H20 上验证 端到端 2.1M token 全 78 层通路;
  • 关键数字:group_size 从 2 增加到 8,峰值分配显存只增加 0.21 GB——基本是常数。这条数字直接说明,之前 95% 的显存都浪费在 prompt 段重复反向上了;
  • stress test 把 Qwen3.6-27B 推到 4.46M positions

这件事的颠覆性在于:它不是"再训一个更大的模型",而是"用同一个模型、同一个算法,把现有集群的有效上下文拉到下一个量级"——对于买不起 1024 卡 H100 集群的中小团队,这是一条绕开硬件焦虑的工程路径

三、LongStraw 怎么 work:三步拆图

如果你直接把原始 GRPO 训练 loop 灌给百万 token prompt,显存墙照样爆。LongStraw 用 三步解耦 把这个墙推平:

第 1 步:共享 prompt 段走 no_grad,只保留"架构感知的最小 state"

# Stage 1 — 共享 prompt 不走 autograd
with torch.no_grad():
    prompt_state = model.encode(prompt)   # 仅保留后续 tokens 真正需要的最小 state

这里的关键词是 architecture-aware:

  • attention 层,需要保留 KV cache;
  • recurrent 层(Mamba/Jamba 类),需要保留 recurrent state;
  • compressed-attention 层,需要保留 压缩后的 attention state

每种架构具体要保留什么、丢什么,由网络结构决定——所以 LongStraw 不是"通用 trick",而是"针对每种架构的工程适配"。

第 2 步:response 段逐条 replay forward + backward

# Stage 2 — response 分支逐条 replay
losses = []
for resp in responses:                    # group_size 个候选 response
    logits = model.decode(prompt_state, resp)   # prompt_state 从快照读回
    loss = grpo_loss(logits, resp, group_baseline)
    losses.append(loss)
total_loss = stack(losses).sum()
total_loss.backward()                     # 每条 response 单独的 backward

要点是:单条 response 的 live graph 远小于全量回放——response 长度本身被 GRPO 裁短过(几百到几千 token),而 prompt 已经 detach,二者不会叠加成"巨图"。

第 3 步:代价是 wall-clock 时间

既然要 replay group_size 次 response forward,端到端训练吞吐会下降——用显存换时间。原文没有给完整的 Pareto 曲线,工程团队需要自己跑 profiling 才能确认 sweet spot。

伪代码完整长这样:

prompt_state = no_grad_encode(prompt)          # 仅一次

losses = []
for resp in responses:                         # replay group_size 次
    logits = decode(prompt_state, resp)
    losses.append(grpo_loss(logits))
sum(losses).backward()                         # 每条独立 backward

四、为什么这件事对 2026 年的 AI Agent 平台至关重要

如果你是下面任一种角色,LongStraw 几乎就是必读:

  • AI Agent 平台 / RL infra 负责人:trajectory 越长、共享 prompt 越长,group size 越不能省。LongStraw 的"group size 基本不增显存"如果成立,意味着可以放心开大 group 拿更稳的 GRPO baseline——这是直接提升训练质量的杠杆;
  • 长上下文训练栈工程师:在 hybrid recurrent、MoE、compressed-attention 模型上做后训练,这是当前为数不多的可借鉴工程模板;
  • 买不起超算的中等团队:它提供一条 绕开 H100 集群焦虑的路径——8 张 H20 就能跑 2.1M token RL,直接压低实验成本;
  • 算法研究者:它不改 GRPO 算法本身,只改执行图,理论上可以插件式接到 GRPO / Dr.GRPO / RLOO / DAPO 等所有变体上;
  • 不太适合纯算法党:本文不主张改算法,改的是图,纯算法视角会觉得"工程味太重"。

最关键的一句话:Agent 时代,训练栈必须从"模型 + 数据 + 损失函数"的三件套,扩展成"模型 + 数据 + 损失函数 + 执行图"的四件套——你不优化执行图,光换模型,百万 token 训练永远跑不起来。

五、三处落地风险别踩

风险 1:detach 切断了 prompt 内部梯度耦合,policy 可能不收敛到标准 GRPO 的等价解

这是 LongStraw 最大的不确定性。prompt 端 detach 等于手动断开 prompt 内部的自反传路径——严格说,损失函数的梯度和标准 GRPO 有偏差。论文自承只建立 execution capacity(执行容量),不保证训练质量

落地前必须做的验证:

  1. 在 7B 量级的小模型上,用标准 GRPO 和 LongStraw GRPO 跑同一份数据;
  2. 比对最终 checkpoint 在同一 test set 上的 reward 分布;
  3. 差异 <5% 才可接受——超过就重新审视 detach 策略,或者考虑 hybrid 方案(关键层保留 autograd,次要层 detach)。

风险 2:replay 把显存压力换成 wall-clock 时间,单 token 训练成本上升

replay 次数 = group_size,大 group(16、32)下,单步时间可能是标准 GRPO 的 16x / 32x。8 张 H20 上,端到端训练时间可能从 1 天变成 2-3 天。建议:

  • 先用小 group(2、4)验证流程跑通;
  • 再跑一次 throughput vs group_size 的 profiling,找自有集群的 Pareto 最优;
  • 不要直接在大规模生产训练上试。

风险 3:架构适配有门槛,纯 dense Transformer 收益有限

LongStraw 的 capture_prompt_state 必须针对具体架构实现。纯 dense Transformer(LLaMA、Mistral)理论上更容易实现(只需保留 KV cache),但 没有 recurrent / compressed state 那种更强的状态压缩,实际收益可能打折扣。跨模型迁移时,要为每类架构重新实现适配层。

风险 4:ZeRO / FSDP / TP 叠加效果未验证

生产训练通常是 ZeRO-3(分参数)+ FSDP(分 optimizer state)+ TP(分 tensor)的组合。LongStraw 的 detach + replay 在 ZeRO-3 下,prompt_state snapshot 的跨 GPU 广播FSDP 下每条 replay 分支的梯度 reduce——原文实验都没覆盖,需要自己趟坑。

风险 5:代码审计必须做

开源仓库 MindLab-Research/longstraw 是入口,但 README 不一定覆盖完整分布式路径。落地前务必审计:detached prompt_state 的 gradient composition 是否完整、分布式 forward 路径是否闭合、与自家 training framework(veRL / OpenRLHF / TRL)的接口是否匹配。

六、写在最后

LongStraw 最有价值的,不是 "8 张 H20 跑 2.1M token" 这个标题数字,也不是开源仓库本身,而是它给 Agent 时代训练栈的工程师 一个被默认忽视的"换执行图"杠杆:

百万 token RL 训练,不是被算力卡住的,是被显存卡住的;显存不是被模型卡住的,是被"共享 prompt 的反向图"卡住的。把图拆开,峰值显存就能从 group_size × prompt 退化到 prompt + response。

下次再有人跟你说"百万 token RL 训练只能等下一代 GPU",你可以问三个问题:

「你的 prompt 段走了 autograd 吗?」 「你的 response 段是逐条 replay 还是一次全跑?」 「你的 architecture-aware state 适配了什么层?」

——三个问题就能判断对方是真的踩过"显存墙",还是只抱怨过"模型不够大"。


延伸阅读 - 论文:arXiv 2607.14952(LongStraw:固定 GPU 预算下,把 RL 后训练上下文推到 2M+ token 的执行栈) - 开源仓库:github.com/MindLab-Research/longstraw - 同方向工作:Megatron / DeepSpeed 的 sequence parallel(前向能跑,未优化 backward 复制放大问题)、GRPO 变体(Dr.GRPO / RLOO / DAPO,本文与之正交)、hybrid recurrent 路线(Jamba / RecurrentGemma / Transfusion 系,模型架构层面) - 工程模板:Qwen3.6-27B / GLM-5.2 是验证架构,标准 dense Transformer 迁移路径未完整验证


三个标题变体

  1. 百万 token 的 RL 后训练别再等下一代 GPU——这篇论文说"显存墙"是被低效的反向传播图卡住的,8 张 H20 就能跑 2.1M token
  2. AI Agent 想"长记性",先撞穿显存墙——LongStraw 把训练执行图拆开,group size 翻 4 倍显存只多 0.21 GB
  3. Agent 时代训练栈必须升级成四件套——LongStraw 用"detach prompt + replay response"把百万 token RL 训练成本打下来

小红书风格卡片文案(可直接发布)

🤖 AI Agent 想"长记性",先撞穿显存墙

百万 token 的 RL 后训练 真的只能等下一代 GPU 吗? 💸

2026 年 7 月这篇论文(arXiv 2607.14952) 正面反驳了一个被默认了很久的假设:

百万 token RL 训练,不是被算力卡住的 是被显存卡住的 显存不是被模型卡住的 是被"共享 prompt 的反向图"卡住的 🔥

过去大家都默认:

🔹 推理侧 context window 已逼近 100 万+ 🔹 RL 后训练侧仍卡在 256K 🔹 唯一的解法是"换更大的集群 / 等下一代 H 系列" 🔹 Agent trajectory 越长越没救 ❌

但这篇论文给了一个反直觉洞见 💡:

在 GRPO 训练里 同一个 prompt 要被 group_size 个候选 response 共享 一次 backward 跑下来 prompt 段的中间激活被复制 group_size 份 激活显存随 group_size 线性爆炸 💥

你以为是"模型太大" 其实显存里 80% 都是 prompt 段重复的反向传播中间结果 对每条 response 来说完全相同 却被无脑地反复保留

LongStraw 的解法:三步拆图 🛠️

Step 1 · 共享 prompt 走 no_grad 只保留"架构感知的最小 state" attention 层留 KV cache recurrent 层留 recurrent state compressed-attention 层留压缩状态

Step 2 · response 段逐条 replay 每条 response 单独 forward + backward prompt state 从快照读回 单条 response 的 live graph 远小于全量回放

Step 3 · 代价是 wall-clock 时间 replay 次数 = group_size 用显存换时间 需要工程团队跑 profiling 找 sweet spot

效果有多炸 📈:

🔹 Qwen3.6-27B:8 张 H20 跑通 2.1M positions 🔹 GLM-5.2:32 张 H20 验证 2.1M token 全 78 层 🔹 关键数字:group_size 从 2 增加到 8,峰值显存只多 0.21 GB——基本是常数 🔹 stress test:4.46M positions on Qwen3.6-27B

这件事的颠覆性在于 🎯:

不是"再训一个更大的模型" 而是"用同一个模型、同一个算法 把现有集群的有效上下文拉到下一个量级"

对于买不起 1024 卡 H100 集群的中小团队 这是一条 绕开硬件焦虑的工程路径 🚀

最爽的是——代码已经开源 🎉 GitHub:github.com/MindLab-Research/longstraw 你今天就能在自己 cluster 上跑一遍

工程落地点 🛠️:

1️⃣ Agent 平台放心开大 group —— "group size 基本不增显存"如果成立,可以拿更稳的 GRPO baseline 2️⃣ 小模型先验证 reward 偏差 —— 7B 量级对比 GRPO vs LongStraw GRPO,reward 差异 <5% 才可接受 3️⃣ 跑 throughput vs group_size profiling —— 找自有集群的 Pareto 最优,大 group 时间成本上升 4️⃣ 架构适配必须做 —— 纯 dense Transformer 收益有限,hybrid / MoE / compressed-attention 收益最大 5️⃣ ZeRO / FSDP / TP 叠加需额外工程 —— 原文未覆盖分布式并行场景,prompt_state 广播 + 梯度 reduce 要自己实现

⚠️ 必须警惕的边界: - detach 切断 prompt 内部梯度耦合,policy 可能不收敛到标准 GRPO 等价解(论文自承只验证执行容量) - replay 把显存换时间,大 group(16/32)下单步时间可能是 16x / 32x - 纯 dense Transformer 迁移路径未完整验证 - 生产级分布式并行叠加效果未知,需要工程团队自己趟坑 - 代码审计必须做——开源 README 不一定覆盖完整分布式路径

📎 论文 ID:2607.14952

💬 评论区聊聊:你做 RL 后训练时,显存先爆还是算力先爆?LongStraw 这种"拆图"思路对你有启发吗?🤔

人工智能 #AI科普 #LLM训练 #RL后训练 #强化学习 #GRPO #Agent #长上下文 #显存优化 #工程实践 #论文分享 #技术分享 #开发者 #研究者 #AI前沿