LongStraw:固定 GPU 预算下超越 200 万 token 的长上下文强化学习训练 LongStraw: Long-Context RL Beyond 2M Tokens under a Fixed GPU Budget
在固定 GPU 预算下让 GRPO 训练跑到 2M+ token 上下文
前置知识
GRPO(Group Relative Policy Optimization)
DeepSeekMath 提出的 RL 后训练目标,是 PPO 去 critic 版本。把共享同一 prompt 的若干 response 组成 group,用组内归一化奖励作优势 $A_i$,再优化裁剪重要性比 $\rho$ 的 clipped surrogate加 KL 项。
本文全部系统设计都是为了让这个 GRPO 更新图在 2M token 上下文下跑通——一个 prompt 同时条件化 $G$ 条 response 的打分与反向传播。
Context Parallelism(CP,上下文并行)
把序列长度切分到多卡,每 rank 只持一段 KV 历史。Qwen 用 CP8 按 $\mathrm{owner}(p)=p\bmod 8$ 分 page;GLM 用 zigzag CP32 每rank 持 65536 token。响应查询须跨所有 CP rank 做全局 softmax 合成。
CP 是把 prompt 长 KV 状态塞进固定显存的核心手段,但前向合并与反向梯度同步是分布式正确性的关键,也是本文最痛的缺口所在。
Expert Parallelism(EP,专家并行)
把 MoE 的 expert 参数分到多卡。GLM 共 256 个 routed expert,EP32 名义每 rank 8 个。token 经 top-8 router 展开成 8 个 expert 分配,通过 all-to-all 派发到 owner、计算后再 inverse-combine。
CP 分 prompt 位置、EP 分 expert 参数,二者不可互相替代——这是 GLM 路径复杂度的根源。
QLoRA / LoRA
LoRA 在冻结 base 权重上加 $W+BA$ 低秩增量;QLoRA 把 base 量化到 NF4。本文 Qwen 用 NF4 QLoRA(rank 16、lr $2\times 10^{-4}$,$1.17\times 10^8$ 可训练参数),GLM 用 rank-8 LoRA。
PEFT 只减少持久参数存储,不减少 KV 长度或 response 激活寿命——这是单靠 QLoRA 不能让 2M 跑通的根本原因。
MLA + DSA + IndexShare
MLA 把每 token KV 压成 latent(GLM 吸收后 key 宽 576);DSA 在 MLA 上加 top-k indexer(每查询选 top-2048 位置);IndexShare 让 21 层把 top-k 发布给 57 个消费者层复用。
GLM 全部复杂度源于此:要同时保存 MLA latent + DSA index-key page,且分布式下「本地 top-2048」≠「全局 top-2048」。
Activation Checkpointing(激活检查点)
用算力换显存:前向只存层输入和元数据,反向时重算层内中间量。本文关键决策是把边界设为完整 decoder layer 而非仅 attention——否则 MoE router 输出、dispatch 排列、expert 输入仍留在图里。
这是把 GLM response 激活压到一层范围内的必要条件,也是「检查点只用在短 response 而非 2M prompt」论断的核心。
GDN(Gated DeltaNet)
一种线性 recurrent 模块,用 gated delta rule 维护固定形状的递归状态矩阵。关键性质是跨 prompt 边界的状态大小与 prompt 长度 $P$ 无关。Qwen3.6-27B 共 48 层 GDN + 16 层 full attention。
这是 Qwen 比 GLM「干净」的根本原因——48/64 的层把 prompt 状态压缩成常数大小,存储压力集中在少数 full-attention 层。
研究动机
AI agent 正从一次性回答转向使用工具、查阅文档、在长轨迹上持续行动(ReAct 范式)。对这些 agent,context 承载证据、环境观测、工具输出与历史决策,长度随轨迹线性累积。然而推理系统与 RL 后训练之间出现鸿沟:推理端已逼近百万 token,RL 后训练仍停留在 256K 以下,部署时只能寄望于长度外推。训练与推理用显存方式根本不同——推理服务器可以 prefill 完 prompt 就丢弃前向图(如 vLLM),但 GRPO 必须对同一 prompt 条件下的多条 response 同时打分并反向传播。二次复杂度的 attention 与长期存活的反向状态让 GPU 显存成为扩展上下文的硬瓶颈。作者指出:FlashAttention、memory-efficient attention、LoRA、QLoRA 都只减少「一部分」开销,prompt 图、response 图、缓存状态和分布式通信仍要挤在同一份显存里,单独用任何一种都不能让固定 GPU 的 GRPO 跑得下。
本文的目标是本文要回答一个与「scale-out」互补的问题:当 GPU 数量保持固定、不允许随上下文增长而加卡时,一条 GRPO 执行路径能走多远?具体目标是:8 张 H20 上让 Qwen3.6-27B 的 GRPO 跑到 2.1M 位置(并扩展到 4.25M),32 张 H20 上让 GLM-5.2(MoE+MLA/DSA)跑完 2.1M token prompt 的端到端执行。作者反复强调不是要刷「最长上下文」纪录,而是降低长上下文训练的硬件门槛,让算力有限的小团队也能探索这个方向。最终目标是把这套执行栈做成 MinT 训练框架的 opt-in 长上下文扩展。
与已有工作不同的是,已有长上下文训练几乎全是 scale-out 路线:Ring Attention 在 32 张 A100 上跑 4.096M,DeepSpeed-Ulysses 在 256 张 A100 上跑 1M,ByteScale 在 1024 张 GPU 上跑 2M LLaMA-7B——都是用「加卡」换序列长度。本文独特切入是「张量生命周期与物理所有权」:作者观察到,长 prompt 的中间量(attention scratch、FFN 中间量、MoE 路由和 expert-token 排列)只要在前向后立即释放就根本不需要参与反向图;真正要保留到 response 阶段的,只是后续 token 依赖的架构特定条件状态。于是可把「live autograd 图」从 $P+R$ 压到只剩 $R$。这个观察在已有工作中几乎没有被作为「固定预算下的核心调度变量」系统化讨论过。
核心方法
LongStraw 一句话概括:GRPO 反向需要的只是「条件于某个 prompt 状态的 response 梯度」,不需要把整个 prompt 放进可微图。技术路线是四阶段事务(per update $k$):Phase 1 用 $\theta_k$ 关掉 autograd 跑一遍 $x_{1:P}$,每层只保留架构特定 prompt 状态、立即释放 hidden、attention scratch、FFN 中间量和 MoE 路由缓冲;Phase 2 把 prompt 状态当只读,在 policy 反向前先把每条 response 的 old/reference log-prob 全算出来(参数全程不变);Phase 3 逐个 member 在 autograd 下重建短 response 图、复用同一份只读 prompt 状态、反向后立即释放;Phase 4 累加完 $G$ 个本地梯度后每 worker 调一次 optimizer。主导激活规模因此从 $P+R$ 变成 $R$,但 prompt 前向算力仍要付。
核心创新是把 group size $G$ 从「显存维度」改造成「时间维度」。作者给出主导资源账本 $M_{live}\approx M_{fixed}+M_{prompt}(P)+M_{grad}+\max_i M_{branch}(R_i)+M_{score}\sum_i R_i$ 与对应时间公式 $T_{update}=T_{prompt}(P)+\sum_i T_{replay}(R_i)$。活跃 policy 图被「最大单个 member」约束而非被 member 数量约束——串行 replay 只增时间不增峰值显存。这与 Ring Attention/Ulysses「把计算摊到更多卡」的本质差异是:LongStraw 不加卡,而是把 prompt 状态从可微图剥离、串行重放。第二个差异是把 prompt 状态 stop-grad $\bar{z}_P=\mathrm{stopgrad}(z_P(\theta))$,明确放弃 $\partial\ell/\partial z_P\cdot\partial z_P/\partial\theta$ 项换执行可行性。
方法步骤详情
六步事务:(1)Prompt 捕获:关 autograd 跑完 2.1M prompt,每层只保留架构特定状态。Qwen 保留 48 层 GDN 状态 + 16 层 KV page;GLM 保留 MLA latent page($[1,64,1,576]$)和 21 层 DSA indexer-key page。(2)物理所有权:page 拷成 right-sized tensor 而非大 chunk 的 view——否则释放 view 时父 chunk 不释放。(3)并行:Qwen CP8;GLM TP1/CP32/EP32。(4)前向:Qwen 用 stable LSE 全局 softmax(MAX+两次 SUM);GLM 把 page 放 CPU,每层 replay 只暂存该层(72–88 MiB)。(5)Response replay:Qwen 切四个 2048-token block 反向回放;GLM 整层 checkpointing+RNG 保留逐层($77\to 0$)重算。(6)所有 member 反向完每 worker 调一次 optimizer。
技术新颖性
技术新颖性在四点。第一,把 prompt 状态、response 状态、分布式所有权放进「同一张量生命周期契约」,明确指出「物理所有权也是算法的一部分」——逻辑分片若仍是父分配的 view 则 allocator 不真释放,必须显式拷贝。第二,把 checkpoint 边界从 attention-only 提升到「完整 decoder layer」:实测只检查 attention 时 MoE router 输出、dispatch 排列、expert 输入仍存活,必须整层才能约束 response 图。第三,把 GLM 路径拆成 7 个 Stage(full-graph 定位→prefix 捕获→layer-0 可微 replay→all-layer 闭合→并行所有权→2M 单 member canary→fresh grouped),每 stage 只隔离一类依赖。第四,首创「四级证据阶梯」(执行容量/算子保真/更新一致/全梯度等价)并坦白当前 receipt 只达第 1–2 级,这种「诚实的系统 receipt 范式」与传统「报 SOTA 数字」差异最大。
实验结果
**Qwen 2.1M(8 H20,CP8)**:$G=2$ 跑 5198.780s、峰值 97.503 GB;$G=8$ 跑 6785.225s、峰值 97.711 GB。组数 2→8(4 倍)峰值显存只增 0.208 GB(+0.213%),时间增 1586s——验证「group size 是时间维度而非显存维度」。**Qwen 4.25M(仍 8 H20)**:$4{,}456{,}448$ 位置完成 $G=8$ replay,峰值 82.960 GB/rank;prefix-frozen 下连续 8 步 optimizer 峰值 83.894 GB;train-block 到 4,542,464 才 OOM。**GLM 2.1M(32 H20)**:完成两次 78 层 policy 反向 + DistOpt,wall 2975.138s,CPU 每 rank 5.8125 GiB 状态。capture 峰值 112.571–145.148 GB/rank(32.577 GB spread),是 replay 前采的所以只是「2M 可行点」非上限。数字单 run 无方差。
查看结构化数据
| 任务 | 指标 | 本文 | 基线 | 提升 |
|---|---|---|---|---|
| Qwen3.6-27B GRPO 在 2.1M 上下文(固定 8 H20) | 峰值 max_memory_allocated(decimal GB/rank) | G=2 时 97.503 GB;G=8 时 97.711 GB | Qwen 原生 max_position 仅 262,144;本文 receipt 恰为 8× | 上下文 8× 于原生设置仍塞进同一份 8-H20 预算;G 从 2→8 仅 +0.208 GB(+0.213%) |
| Qwen3.6-27B GRPO 在 4.25M 上下文(固定 8 H20) | 可执行上下文长度 / 峰值显存 | 4,456,448 位置、G=8 replay、峰值 82.960 GB;prefix-frozen 下 8 步训练峰值 83.894 GB | 同一 8-H20 预算的 2.1M receipt | 上下文再翻倍仍可执行;train-block 代理到 4,542,464 才 OOM |
| GLM-5.2 GRPO 在 2.1M prompt 上下文(固定 32 H20) | 可执行 prompt 长度 / CPU prompt 状态 / rank | 2,097,152 prompt token,每 rank CPU 5.8125 GiB;层暂存 72–88 MiB;32/32 rank 终止 | GLM-5.2 公开配置 max_position 1,048,576;本文恰 2× | 上下文 2× 于公开配置;完成两次 78 层 policy 反向 + DistOpt 调用 |
| Group size 对显存敏感度(Qwen) | G=2 → G=8 峰值显存增量 | +0.208 GB(+0.213%) | 传统 full-sequence autograd:激活随 G 线性增长 | group 增 4 倍显存基本不变;时间 +1586s(证实 G 是时间维度) |
局限与改进
作者极坦诚罗列限制。**第一,prompt 状态 detach**:两条路径都存 $\bar{z}_P=\mathrm{stopgrad}(z_P)$,只算 $\partial\ell/\partial\theta$,丢 $\partial\ell/\partial z_P\cdot\partial z_P/\partial\theta$——本文只达第 1–2 级证据。**第二,Qwen 反向不完整**:只 all-reduce $dQ$,$dK^r$、$dV^r$ 留本地,但 K/V LoRA 是 replicated,需 $\nabla W_K=\sum_r\nabla W_K^{(r)}$;8 个 AdamW 各自 step 可能 divergence。**第三,GLM CP-local DSA**:每 rank 只在 65536-token 选 top-2048;跳过 finalize。**第四,工作负载合成**:reward 确定,$\beta=0$ 或首步 ratio=1 使 clip/KL 不激活;无在线 rollout。**第五**,时间单 run 无方差。
独立分析的弱点
我识别四个弱点。**弱点 1(头号):分布式梯度不一致**。Qwen 缺 $dK/dV$ 跨 rank 归约,GLM 跳过 `finalize_model_grads`,即使 receipt 跑完更新方向也可能不是真实 GRPO 梯度。改进:Qwen 加 selective reducer;GLM 接回 `finish_grad_sync` 并做 post-step adapter hash 测试。**弱点 2:GLM DSA 是 CP-local 近似**,长程任务会有偏差。改进:实现 candidate merge + 全局 top-2048 + selected-value 交换。**弱点 3:detach 让 response 共享 stale 状态**,重复训练必须每步重捕。改进:32K–64K 上和全序列执行做逐 shard 梯度 parity。**弱点 4:串行 replay 慢、rank 负载不均**——Qwen $G=8$ 要 1.9 小时。改进:partial-overlapped 并行、MoE expert skew 均衡。
未来方向
作者 Section 12.6 给出依赖序路线:**(1)** 恢复 Qwen K/V reductions 并在 32K 验证 GLM gradient-finalization;**(2)** 验证 global cross-CP DSA 的 candidate selection 和 output 合成;**(3)** 在 32K–64K 上把每个 LoRA gradient shard 与 optimizer delta 和传统全序列参考逐参数对比;**(4)** 跑真实 rollout group(采样、reward、update、checkpoint、reload);**(5)** 从 pinned stack 重跑 2M,再做 >2M 的 GLM sweep。基于成果可延伸:把 stale-state approximation 形式化并分析偏差界;探索 hybrid GPU/CPU 放置;把 prompt-state 契约推广到 Mamba、Hyena 等架构;与 AReaL 异步 rollout-training 系统结合。
复现评估
**开源**:代码在 `github.com/MindLab-Research/longstraw`,是 MinT 的 opt-in 扩展。**pinned 栈**:Table 6 给出 MinT Runtime、Megatron-LM、Megatron-Bridge(GLM-5.2 patch)、verl、GLM-5.2 snapshot 精确 commit,但都 prospective——历史 2M receipt 早于 pinned 栈不能反向绑定,当前代码树只有 32K canary 覆盖修复后路径。**算力**:8 或 32 张 H20,每卡约 141 GiB——比 Ring Attention 的 32 A100 或 ByteScale 的 1024 GPU 友好但仍是可观投入。**复现难度**:高,需精通 Megatron、MoE EP、MLA/DSA、QLoRA、IndexShare。**缺失**:没绑学习率、token checksum、LoRA seed、NCCL 版本;时间单 run 无方差。开源+pinned 栈是加分,receipt 未到「按按钮复现」程度。
论文图表
三部分。(a) CP32 zigzag 所有权:32,768 个 page 分成 64 个 chunk(每 chunk 512 page),rank $r$ 持 chunk $r$ 与镜像 $63-r$,共 1024 page=65536 token;本地 tensor 顺序保留 Megatron 双 chunk 拼接,全局 page ID 仍可用于因果 mask。(b) 接受 receipt 中 DSA 是 local:每 rank 只对自己 65536 key 算分、选本地 top-2048、做本地 sparse attention——缺 cross-CP candidate merge、全局 top-2048、selected K/V 交换、output 合成四步。(c) 更新缺口:backward 写本地 DDP grad buffer 后直接调 DistOpt,跳过 `finalize_model_grads`,CP-replicated 非 expert adapter 梯度未归约;parameter all-gather 可让 replica 看起来相等但无法补回缺失的 CP 梯度求和。
这张图同时展示 GLM 路径的「前向算子近似」和「更新不一致」两个独立缺口,是判断 GLM receipt 实际语义强度(只达执行容量级、未达算子保真级)的关键证据。