LongStraw:固定 GPU 预算下,把 RL 后训练上下文推到 2M+ token 的执行栈

  • 关联论文:2607.14952
  • 作者:flyP
  • 更新:2026-07-19

一句话结论

LongStraw 是一个面向百万 token 量级 RL 后训练的架构感知执行栈,用「共享 prompt 无 autograd + 短响应分支逐条回放」的方式,把激活显存压到训练图级别,代价是多花一些重放时间;它在 8 张 H20 上把 Qwen3.6-27B 的 GRPO 训练推到 2.1M positions,并在 32 张 H20 上跑通 GLM-5.2 的端到端 2.1M token 全 78 层通路。

解决的真问题

推理侧的 context window 已经逼近 100 万甚至更长,而 RL 后训练(post-training)侧的 context 长度往往仍卡在 256K 及以下,部署时再靠长度泛化硬扛。这个 gap 对 AI Agent 类任务尤其致命:观察、工具输出、文档、历史决策会在一条 trajectory 里不断累积。问题在于,对一条 prompt 内做 GRPO(Group Relative Policy Optimization)时,同一个 prompt 要被多个 candidate response 共享,如果 prompt 已经很长,response backward 阶段整张计算图会保留所有这些 response 的中间激活,激活显存随 group size 线性爆炸,于是百万 token RL 训练被「显存墙」死死卡住,而不是被算力卡住。LongStraw 想做的就是把这条墙从算法端往系统端推一推。

核心方法

LongStraw 的关键观察是:GRPO 的训练图之所以爆显存,是因为 prompt 段的反向传播和 response 段的反向传播被串在同一张 autograd 图里。共享 prompt 的所有 response 都要 backprop,prompt 部分的中间激活被复制了 group_size 份。LongStraw 把这层显式解耦:

  • 共享 prompt 不走 autograd:对共享 prompt 部分做前向推理,但 detach 出 prompt state,只保留「后面 token 真正需要的、模型特异的中间状态」。具体保留哪些由网络结构决定(attention 层需要 KV,recurrent / compressed-attention 层需要各自的压缩状态),所以叫 architecture-aware。
  • 短 response 分支逐条 replay:每个 response 单独作为一条短分支重放前向 + backward,prompt state 从已 detach 的快照里读回。因为 response 长度相对 prompt 短得多(GRPO 的输出本来就被裁短),单条分支的 live graph 远小于全量回放。
  • 代价是 replay 时间:要重放 group_size 次 response forward,整体吞吐会下降,换来的是峰值显存基本只由 prompt 长度 + 单条 response 决定。

伪代码逻辑大致是:

# Stage 1: shared prompt, no grad
with torch.no_grad():
    prompt_state = model.encode(prompt)   # 保留后续 token 所需的最小 state

# Stage 2: response-by-response replay with autograd
losses = []
for resp in responses:                    # group_size 个候选 response
    logits = model.decode(prompt_state, resp)
    loss = grpo_loss(logits, resp, group_baseline)
    losses.append(loss)
total_loss = stack(losses).sum()
total_loss.backward()                     # 每条 response 单独的 backward

要点是:encoder 阶段省下的显存是 group_size * prompt_activations 这一大块;decoder 阶段每次只承受一条 response 的图。

关键实验与数据

作者选择两条结构有代表性的模型验证执行栈是否真的吃下架构差异:

  • Qwen3.6-27B:hybrid recurrent + full attention。8 张 H20 GPU 上完成 grouped Qwen scoring 和 response backward,在 group size 为 2 和 8 时都能跑到 2.1M positionsgroup size 从 2 增加到 8,峰值分配显存只增加 0.21 GB,基本等于常数,这就是解耦 prompt state 的直接收益。另一项独立 stress test 推到 4.46M positions
  • GLM-5.2:compressed-attention MoE,结构更复杂。32 张 H20 上验证端到端 LongStraw 通路,2.1M token prompt 全 78 层都跑通,覆盖了 MoE + 压缩注意力的组合。

但作者自己也强调:这些实验建立的是 execution capacity(执行容量),而不是完整训练正确性——因为 detach 出来的 prompt state 在严格意义上会损失一些梯度耦合,且部分分布式前向和梯度 composition 路径尚未完整实现。所以「能跑」和「训出来的策略严格等价于常规 GRPO」是两件事,本文主要交付的是前者。

亮点与局限

亮点

  • 问题定位非常诚实:它承认现在百万 token RL 训练是被显存卡住的,而不是被算法或数据卡住。这种系统视角在 RL infra 论文里相对少见。
  • 解法对结构透明:用 detach + replay 把图显式拆开,没有改模型、没改损失函数,理论上可以叠加到任意 GRPO / PPO 类算法。
  • 在 hybrid recurrent、compressed-attention MoE 两条迥异的结构上同时验证,architecture-aware 不是嘴上说说。
  • 实测的「group size 加 8 倍只多 0.21 GB 显存」是非常有冲击力的数字,直接说明之前绝大多数显存都浪费在 prompt 重复反向上了。

局限

  • 梯度正确性未证:prompt 端 detach 等于切断了 prompt 内部的自反传耦合,严格说损失函数的梯度会和标准 GRPO 有偏差。replay 次数越多,这个偏差的累积效应越值得审视。
  • 吞吐换显存:replay 把显存压力换成了 wall-clock 时间,论文没有给端到端训练吞吐 vs 显存 的 Pareto 曲线,单 token 训练成本会上升。
  • 执行 ≠ 完整训练:分布式 forward 和梯度 composition 还没补齐,意味着离「可以直接拿这个栈去训一个 SOTA agent」还有距离。
  • 只在自家架构(Qwen3.6-27B、GLM-5.2)上验证,跨架构(标准 dense Transformer、长 context Mamba 类)的可移植性未明确。

对工程落地的启发

  • 任何想自建百万 token RL 后训练栈的团队,都应先算一笔账:是不是显存先撞墙、还是算力先撞墙。如果是显存,LongStraw 的 detach + replay 思路是可以低成本借鉴的工程模板。
  • 对 Agent 平台来说,trajectory 越长,共享 prompt 越长,group size 越不能省——LongStraw 的「group size 基本不增显存」的结果如果成立,意味着可以放心开大 group 来拿更稳的 GRPO baseline。
  • 代码已开源(MindLab-Research/longstraw),但落地前需要重点验证:在你的模型结构上,detach 是否会让策略真的收敛到同一份 SOTA 检查点。
  • 对长上下文推理栈的启示是反方向的:推理栈不需要 backward,所以不受这个显存墙约束;不要把推理的 context 上限直接当作训练的 context 上限。

与同方向工作的关系

  • hybrid recurrent / Mamba / 压缩注意力 路线(如 Jamba、RecurrentGemma、Transfusion 系)同向,但本文关注点不是模型架构,而是「在这种架构下 RL 训练要怎么扩 context」。
  • GRPO 变体(Dr. GRPO、RLOO、DAPO 等)正交——本文不改算法,只改执行图。理论上可以插件式接进去。
  • long-context RLHF 框架(如 OpenRLHF、verl、TRL 的长序列支持)相比,LongStraw 的卖点是显式的 architecture-aware detach,更贴近带 recurrent state 的非纯 Transformer。
  • Megatron / DeepSpeed 的 sequence parallel 工作也有交集,但那些工作主要解决「前向能跑」,对 backward 的 prompt 复制放大问题缺乏针对性。

适合谁读

  • 做 LLM infra / RL infra 的工程师,尤其是 Agent 平台、长上下文训练栈的负责人;
  • 在 hybrid recurrent、MoE、compressed-attention 模型上做后训练的研究者;
  • 想理解「百万 token RL 训练到底难在哪」的系统方向学生——本文是相当干净的系统教学样本;
  • 不太适合纯算法研究者:本文不主张改 GRPO,改的是图。

不确定处

  • 端到端训练质量(policy 是否收敛到与标准 GRPO 等价的解)原文未明确给出实验,本文自承只建立 execution capacity;
  • 与 ZeRO / FSDP / TP 并行策略叠加时的具体开销数字原文未明确;
  • 在 dense Transformer(非 hybrid、非 MoE)上的表现原文未明确。

工程落地与核查(Jay)

1. 事实核查小结

  • ✅ 2.1M positions on 8 H20 GPUs with Qwen3.6-27B:原文摘要明确。
  • ✅ group size 从 2 到 8,峰值分配显存只增加 0.21 GB:原文摘要明确数字。
  • ✅ 4.46M positions stress test:原文摘要 stress test 项目。
  • ✅ 32 H20 上 GLM-5.2 端到端 2.1M token 全 78 层跑通:原文摘要明确。
  • ✅ Qwen3.6-27B 是 hybrid recurrent + full attention,GLM-5.2 是 compressed-attention MoE:原文摘要一致。
  • ✅ 本文只建立 execution capacity,不保证训练正确性:原文摘要明确 self-claim,与本文一致。
  • ✅ GitHub: MindLab-Research/longstraw:代码已开源(README 中有)。
  • ⚠️ 重要存疑:全文只交付执行容量,从未声称训练质量收敛到标准 GRPO 等价解。任何基于本文的工程落地,都必须在自有数据上验证策略收敛性,不能假设「能跑就能训出等价的 policy」。
  • ⚠️ 存疑:dense Transformer(非 hybrid)上的可移植性未验证,标准 LLaMA / Mistral 结构不适用。
  • ⚠️ 存疑:与 ZeRO / FSDP / TP 的叠加效果未给出,分布式训练场景需额外工程工作。

2. 实际系统怎么用

LongStraw 的适用场景判断流程

你的 RL 训练撞墙了吗?
  ↓
是显存先爆还是 wall-clock 先爆?
  ↓
若是显存(group_size * prompt_length 导致的 OOM):LongStraw 思路可能适用
若是算力(throughput 低):LongStraw 不解决问题,换大集群或优化算子

集成 LongStraw 的步骤(基于已开源代码):

# 基于 longstraw 伪代码逻辑的集成示意
from longstraw import ArchitectureAwareStateManager

# Step 1: 替换你的 GRPO dataloader 的 collate 逻辑
manager = ArchitectureAwareStateManager(model)
# 模型需支持:query minimal state(如 KV cache、hidden states 按层导出)

# Step 2: Stage 1 — prompt 前向,不走 autograd
with torch.no_grad():
    prompt_state = manager.capture_prompt_state(prompt_batch)
    # prompt_state 包含后续 tokens 所需的 architecture-specific state

# Step 3: Stage 2 — response 分支逐条 replay
for resp_idx in range(group_size):
    with model.forward(resp_branch[resp_idx]) as branch:
        logits = branch.decode(prompt_state, resp_branch[resp_idx])
        loss = grpo_loss(logits)
        losses.append(loss)

# Step 4: 合并梯度
torch.cat(losses).sum().backward()

关键工程决策:是否复用共享 prompt batching?

LongStraw 的设计允许不同 sample 的 prompt 共享同一份 prompt_state snapshot(如果 prompt 内容相同)。这意味着在 batch 内可以做 prompt 合并,进一步降低显存。实现上需要按 prompt 内容 hash 分组,把内容相同的 sample 合并到同一批做 shared forward。

3. 坑在哪

坑 1:梯度正确性未经验证,收敛可能偏移

这是最大风险。detach 切断了 prompt 端内部的自反传耦合,严格说这一步会引入梯度偏差。如果偏差只在少数层累积可能影响不大,但如果偏差在关键决策层(如 RL 的 advantage estimation)累积,可能导致训出来的 policy 跟标准 GRPO 存在系统性偏移。

落地前必须做的验证

  1. 在小模型(如 7B)上,用标准 GRPO 和 LongStraw GRPO 同时训练相同数据,比对最终 checkpoint 在同一 test set 上的 reward 分布;
  2. 如果 reward 分布显著不同(>5% 差距),说明偏差不可接受,需要重新审视 detach 策略。

坑 2:吞吐换显存的 tradeoff 不透明

原文只给了峰值显存,没给 wall-clock 时间曲线。replay 次数 = group_size,对于大 group(16、32),replay 开销可能是 16x / 32x 的前向延迟。在 8x H20 的小规模集群上,这个开销可能让端到端训练时间从 1 天变成 2-3 天。需要跑一个 throughput vs group_size 的 profiling 再决定生产使用的 group size。

坑 3:分布式并行策略叠加有未知数

生产训练的 RL pipeline 通常用 ZeRO-3(分参数)+ FSDP(分 optimizer state)+ TP(分 tensor)的组合。LongStraw 的 detach + replay 方案在 ZeRO-3 场景下,prompt_state snapshot 的跨 GPU 广播逻辑需要额外实现;FSDP 下每个 replay 分支的梯度需要正确 reduce。这些在原文实验里未覆盖,是工程化路上需要自己趟的坑。

坑 4:架构适配有门槛

LongStraw 明确是 architecture-aware——需要针对具体模型架构实现 capture_prompt_state 的逻辑,把哪些中间状态保留、哪些丢弃。纯 dense Transformer(如 LLaMA、Mistral)理论上更容易实现(只需要保留 KV cache),但也意味着没有 recurrent/compressed attention 那种更强的状态压缩,效果可能打折扣。跨模型迁移需要为每类架构重新实现适配层。

坑 5:与 GRPO 以外的 RL 算法兼容性有限

LongStraw 的设计天然适配 GRPO(group 内共享 prompt),对 PPO(每个 sample 独立 prompt)来说,prompt 重复放大问题本来就不存在,LongStraw 收益有限。如果你的训练 pipeline 用的是 PPO 而不是 GRPO,这套方案收益可能很小。

4. 工程核查 checklist

检查项 建议
梯度正确性验证 先在小模型上做 GRPO vs LongStraw GRPO 的收敛对比实验,reward 差异 <5% 才可接受
Throughput profiling 测不同 group_size(2/4/8/16)下的 wall-clock 时间 vs 显存,找到自有集群的 Pareto 最优
分布式并行叠加 ZeRO-3 + FSDP + TP 组合下,prompt_state 广播和梯度 reduce 需要额外实现
架构适配 确认你的模型有 recurrent/compressed state 才有最大收益;纯 dense transformer 收益有限
RL 算法匹配 GRPO 适配;PPO 场景不需要 LongStraw 改造
开源代码审计 MindLab-Research/longstraw README 中确认具体实现的覆盖范围(是否包含 gradient composition)
仿真先行 不要直接在大规模训练上试,先用 1M token + 小模型跑通验证流程,再 scale up