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(执行容量),不保证训练质量。
落地前必须做的验证:
- 在 7B 量级的小模型上,用标准 GRPO 和 LongStraw GRPO 跑同一份数据;
- 比对最终 checkpoint 在同一 test set 上的 reward 分布;
- 差异 <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 迁移路径未完整验证
三个标题变体
- 百万 token 的 RL 后训练别再等下一代 GPU——这篇论文说"显存墙"是被低效的反向传播图卡住的,8 张 H20 就能跑 2.1M token
- AI Agent 想"长记性",先撞穿显存墙——LongStraw 把训练执行图拆开,group size 翻 4 倍显存只多 0.21 GB
- 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 这种"拆图"思路对你有启发吗?🤔