复旦与MindLab联手破解AI训练难题:用8块GPU跑通200万上下文秘密

这项由复旦大学与MindLab联合开展的研究,以预印本形式发布于2026年7月,论文编号为arXiv:2607.14952,有兴趣深入了解技术细节的读者可通过该编号检索完整原文。
**一道现实的鸿沟**
现代AI助手变得越来越聪明,但有一个鲜为人知的矛盾正在悄悄加深:AI在正式"上岗"时能处理几百万字的超长文本,可它在"上岗前的培训阶段"却往往只能处理区区几万字,两者之间存在巨大落差。这就好比一名厨师在实际工作中要掌管一张能容纳两百道菜的超长菜单,但他在烹饪学校练习时只接触过十几道菜的简化版本,然后寄希望于正式工作时自己能"举一反三"。
这个问题在AI智能体(Agent)上尤为突出。所谓AI智能体,就是那些能够使用各种工具、查阅资料、一步步完成复杂任务的AI系统。它们在工作时会积累大量上下文信息——用户的需求、工具返回的结果、之前做出的决策,这些全都堆积在记忆里,动辄就是数十万甚至百万量级的文字。
训练这样的AI系统,麻烦比推理(也就是让训练好的AI直接用于工作)要复杂得多。推理时,机器只需要读一遍输入、给出回答,完事之后可以把中间过程全部清掉。但训练时,系统还需要对比AI给出的多个不同回答,评判哪个更好,然后把"反馈信号"从输出一路传回到模型内部——这在技术上叫做"反向传播"。这个过程会在GPU显存里同时堆积大量中间数据,就像一家餐厅不仅要同时做几十道菜,还要把每道菜的每一步操作都拍下来留档,以便事后复盘。显存撑不住,训练就崩溃。
研究团队给出的解法叫做**LongStraw**,核心思路是:把"读完这本长篇小说"和"反思自己写的答案"这两件事彻底分开来做,从而让有限的GPU显存只需要承担其中的一小部分。
**一、为什么训练比推理更吃显存——从一道数学题说起**
以往的AI训练方式,可以用这样一个场景来理解:老师出了一道题,包含一段很长的阅读材料(这就是"提示词",Prompt),然后让学生A和学生B各自写出答案。传统方法要求把阅读材料和两份答案全部堆在桌面上,同时反复对比、修改,桌面面积有限,东西太多就放不下了。
LongStraw的做法是:先把阅读材料仔细看一遍,但不把它铺在桌面上——只抽取出"理解这道题所需要的关键笔记",把阅读材料本身收起来。然后,拿着这份关键笔记,一次只评判一个学生的答案,评完立刻清掉,再去评判下一个学生的答案。这样一来,桌面上最多只需要放"关键笔记"加上"当前正在评判的那一份答案",空间需求大幅下降。
这套方法在技术上的名字叫做**GRPO**(Group Relative Policy Optimization,组相对策略优化)。它的核心逻辑是:AI生成一组回答,通过比较这组回答的相对好坏来计算"谁更优秀",然后以此来调整AI的参数,让它下次表现得更好。LongStraw没有修改这个评判逻辑,它改变的是"如何在有限资源下把这套评判流程跑起来"。
具体来说,LongStraw把一次完整的训练更新分解为四个阶段。第一阶段叫做"提示词捕获":让AI以不追踪梯度(也就是不准备"留档复盘")的方式读完整段长文本,只保留后续需要用到的那份"关键笔记",其余中间过程立即释放。第二阶段叫做"预评分":在参数不做任何修改的前提下,先记录下每个回答在当前AI版本下的得分,冻结这些得分,作为后续比较的基准。第三阶段叫做"策略重演":一次只处理一个回答,开启梯度追踪,让AI重新过一遍这个回答,计算损失,做一次反向传播,然后立刻清掉这个回答的所有中间数据,再处理下一个。第四阶段叫做"优化器更新":等所有回答都处理完毕,把累积下来的梯度一次性应用到参数上,完成本轮训练。
这种把"读长文本"和"处理每个回答"拆开的设计,使得GPU显存里同时存活的最大数据量从"长文本加上所有回答"缩减为"长文本的关键笔记加上当前这一个回答"。
**二、两个截然不同的AI大脑,两套量身定制的"笔记策略"**
LongStraw并非一套万能模板,它需要根据不同模型的内部结构来决定"关键笔记"应该记录什么。研究团队为两个架构差异明显的大模型分别设计了不同的实现方案。
第一个模型是**Qwen3.6-27B**,它有64个解码层,里面混合了两种处理文字的机制。其中48层使用的是"GDN"(Gated DeltaNet,门控差分网络),这是一种循环机制,用固定大小的"状态向量"来压缩历史信息,就像人类用几句话总结一段对话的要点,无论对话多长,总结出来的关键信息大小始终固定,不随文本长度增长。另外16层使用的是"全注意力"机制,这种机制需要保存每一个历史词的完整记录,就像把整段对话的录音逐字记录,文本越长,记录就越多,存储空间呈线性增长。
因此,Qwen模型的"关键笔记"由两部分组成:48个GDN层各自留下一份固定大小的循环状态,加上16个全注意力层各自留下的键值页面(KV Pages)。这些键值页面按照"上下文并行"(CP,Context Parallelism)的方式分散存储在8块GPU上,每块GPU各自保管一部分。等到处理回答时,8块GPU通过一套精确的数学合并操作(基于稳定的对数求和指数公式)把各自管理的那部分结果汇总成正确答案,就像8个人各自保管了一本账簿的不同章节,合账时按章节编号加权汇总。
第二个模型是**GLM-5.2**,它的结构复杂得多。78个解码层全部使用一种叫做**MLA**(Multi-head Latent Attention,多头潜在注意力)的压缩注意力机制,把历史信息压缩成更紧凑的潜在表示来节省存储。更特别的是,它还叠加了一套叫做**DSA**(Dynamic Sparse Attention,动态稀疏注意力)的机制:每次处理一个词时,不去看全部历史词,而是先用一个轻量级的"索引器"对历史词打分,只选出最重要的2048个位置来精读,其余的跳过。
GLM还有另一个独特之处:它的78层中,只有21层会自己计算这个"选哪2048个位置"的索引,其余57层直接复用邻近层算好的索引,从而避免重复计算。这个设计叫做IndexShare(索引共享)。此外,GLM的前3层使用普通的全连接前馈网络,后75层使用**MoE**(Mixture of Experts,专家混合)结构——每层有256个"专家"网络,每个词只激活其中8个,大幅减少每次前向计算的参数量。但这也带来了一个新挑战:这256个专家分散存储在32块GPU上,每次处理数据都需要跨GPU进行数据分发和汇总(EP All-to-All通信)。
GLM的"关键笔记"同样存储在32块GPU对应的CPU内存中(而非GPU显存),包括78层的MLA潜在键值页面和21个索引计算层的DSA索引键页面。处理回答时,每次只把当前层需要的一小份数据从CPU搬到GPU,用完立刻搬回或释放,从根本上控制GPU显存占用的峰值。
**三、每块GPU到底存了多少东西——用具体数字感受一下规模**
这里提供几个具体数字,帮助感受这些设计的实际规模,而不只是停留在概念层面。
对于Qwen模型,研究设定的上下文长度恰好是2,097,152个位置(即2的21次方,约210万)。其中约208.9万个位置是提示词,剩余8192个位置是回答输入。提示词被分成32640个"页面",每页64个位置,8块GPU各自管理其中的4080个页面。仅仅是16个全注意力层的键值数据,每块GPU就需要存储约15.94GB——这还只是键值数据本身,不算模型权重、适配器参数、临时计算缓冲区等其他占用。完整的训练峰值显存被控制在97.5GB左右(8块GPU各自约97GB)。
对于GLM模型,32块GPU按照Megatron框架的"锯齿形"分配方式各自持有1024个页面、对应65536个提示词位置。每层的MLA潜在页面在一块GPU的CPU内存中占用72MB,21个索引层的DSA键页面各占用16MB。全部78层的MLA加上21层的DSA索引,每块GPU的CPU端存储约为5.81GB,32块GPU合计约186GB的CPU内存用于存放提示词状态。
从GLM那笔全连接隐藏缓冲区的大小可以直观感受MoE并行的压力:65536个位置乘以8路路由,展开后有524288行数据,每行宽度6144,以BF16格式存储,光这一个张量就占用6GB显存。传统的全序列训练图不仅要存这个张量,还要存前后各层的所有中间结果,叠加下来轻易超过单卡上限。LongStraw通过"提示词不建立梯度图"加上"每次只在回答段做一层重新计算"的策略,彻底绕开了这个爆显存的死局。
**四、从32K到210万——一步步排雷的七个关卡**
LongStraw的GLM实现不是一蹴而就的,而是经历了一次典型的工程调试旅程,从最小可行规模开始,一个关卡一个关卡地击穿瓶颈。
研究团队最先遭遇的问题是:在普通的全序列训练模式下,32K长度可以跑通,但一旦尝试扩展到210万位置,GPU显存就会溢出(Out of Memory,OOM)。而且溢出的位置还在不断漂移——先是DSA的注意力得分矩阵撑爆了显存,修完之后又轮到专家LoRA(一种参数高效微调方法)的中间计算,再改完又轮到MoE输出拼接操作。这说明问题的根源不是某一个单独的大张量,而是整个全序列自动微分图太重了,优化任何一个局部都只是把瓶颈推到下一个地方。
第一步突破:彻底放弃对提示词建立梯度图,只在提示词结束处保存必要的状态,之后专心处理回答部分。这个决定确立了整个方案的核心架构。
第二步突破:在不带梯度的情况下,把128K、256K、512K、1M、最终到210万位置的提示词全部过一遍,验证MLA和DSA的状态确实可以被正确捕获和存储,证明存储方案本身是可行的。但此时还没有任何训练的能力,只是单纯地读完了一段超长文本。
第三步突破:选取第0层(最靠近输入的那一层),单独做一次带梯度的回答处理和反向传播,验证"读取保存的提示词状态、处理一段短回答、跑一次优化器"这条最小训练路径是通的。在这一步,引入了CPU存储和按层分批传输的方案,让1M和210万规模都能完成这个单层测试。
第四步突破:把所有78层都串联起来,但先在较短的32K和64K规模上验证,专门解决IndexShare的生命周期问题(索引发布层必须在每次前向传播时发布新的索引,消费层必须消费同一次前向传播的索引,不能跨回答或跨参数版本混用)、DSA调用接口在短回答下的兼容性问题,以及激活检查点的正确粒度问题(必须以整个解码层为单位,而非只覆盖注意力部分)。
第五步突破:引入TP1/CP32/EP32的并行拓扑,配合CPU页面存储和单层分批暂存,让全部78层在32块GPU上的显存占用被控制在合理范围,把测试规模推进到32K和64K的全架构验证。
第六步突破:用一个只有单个回答(G=1)的"哨兵"运行,在210万位置规模下完整走过78层的前向、反向和一次优化器调用,确认完整的执行路径在目标规模下是通的。这还不是真正意义上的GRPO训练(因为只有一个回答无法形成有意义的相对评分),但它验证了资源的可行性。
第七步突破:用两个确定性的合成回答(奖励分别为0和1,归一化后优势为-1和1)完整跑通一次分组执行:捕获提示词、评分、两次78层反向传播、一次优化器调用,32块GPU上的全部32个进程全部正常终止。
**五、用数字说话——实验结果的真实面目**
Qwen模型在8块H20 GPU上完成了两个不同分组规模的完整测试。分组大小为2时,整个运行耗时约5199秒,峰值显存97.503GB;分组大小为8时,耗时约6785秒,峰值显存97.711GB。两个规模之间,峰值显存的差距只有0.208GB,增幅仅0.213%,而时间多出约1586秒。这个结果印证了设计的核心思路:序列化地处理各个回答,使得峰值显存主要由最大单个回答决定,而非由回答数量决定。提示词捕获占据了整个运行时间的约89.6%(约4656秒),每个额外的回答大约只需要265秒。把提示词的耗时摊销到所有回答上,每个回答的平均耗时从2599秒降到848秒,摊销效益相当显著。
在更大的规模上,研究团队还在同样的8块H20 GPU上测试了约445万位置(精确值为4,456,448)的场景,并成功完成了8个回答的完整重演和反向传播,峰值显存82.960GB。在"前缀冻结"模式下(即提示词对应的参数不更新),甚至连续完成了8次包含8个回答的完整优化器更新(共64次回答重演),峰值显存83.894GB,为特定训练目标提供了多步训练的执行证明。一个容量探测测试在4,538,368位置通过,在再多4096个位置处溢出,给出了当前配置下的粗略上限。
GLM模型在32块H20 GPU上,用210万位置的提示词和两个极短的合成回答完成了完整的分组执行,提示词捕获加两次78层前向/反向加优化器调用共耗时约2975秒。从CPU存储的角度看,每块GPU持有约5.81GB的提示词状态数据在CPU内存中;从GPU显存的角度看,捕获阶段的峰值分配在112.571GB到145.148GB之间(各个进程的用量有约32.5GB的差距,提示存在负载不均衡问题)。
**六、哪些事做到了,哪些事还没有——诚实的边界划定**
这份研究在技术诚实性方面表现得相当直接,明确区分了"已经证明的"和"尚未完成的"。
在已经证明的层面,研究建立了四件事:完整的执行路径在指定规模下不溢出、不崩溃,每块GPU都能走完所有阶段;Qwen模型通过全局CP8注意力统计合并,实现了正确的全上下文条件前向计算(带有BF16数值精度的轻微误差,非按位精确);分组执行的时序正确,预评分在参数更新前完成,两次反向传播后才执行一次优化器调用;显存峰值数据和每阶段耗时数据有据可查。
在尚未完成的层面,研究坦诚地指出了三个重要问题。第一,**分布式梯度组合不完整**。对于Qwen,前向注意力的全局合并是正确的,但反向传播时,负责存储键/值的各GPU计算了本地的梯度贡献后没有跨GPU汇总,而键/值的投影适配器参数(LoRA权重)是在所有GPU上复制的,它们应该收到来自所有GPU的梯度之和,但实际上每个GPU各自独立更新了自己的副本——这意味着8块GPU上的模型参数会产生分歧。对于GLM,正常的Megatron训练流程在反向传播后会调用一个叫做finalize_model_grads的函数来完成CP维度上的梯度汇总,但历史上的执行版本绕过了这个函数,直接让优化器从未汇总的本地梯度更新参数。第二,**GLM历史执行版本的DSA前向计算是局部的**,每块GPU只在自己持有的65536个提示词位置里选top-2048,而不是在全部210万个位置里全局选top-2048,这在语义上与模型定义的操作不符。第三,**提示词状态的梯度被截断**,两个模型都没有把梯度传回到提示词处理阶段,这意味着当前实现对模型参数的更新只反映了"如何生成更好的回答",而没有反映"提示词理解部分的参数如何改进"。
换句话说,当前的成果是一张"执行收据",证明了这条路是物理上走得通的,但还不是一张"正确训练收据",还需要后续的梯度同步修复工作才能成为真正有意义的分布式训练结果。
**七、这项研究告诉了我们什么更深层的道理**
研究团队从这次工程实践中提炼出几条对整个AI训练系统领域有参考价值的认识。
关于显存容量,核心决定因素是张量的**生命周期**,而不是计算的稀疏程度。DSA减少了注意力计算量,但提示词的索引键依然是长度相关的数据;MoE减少了每个词激活的参数量,但每次路由和分发仍然产生大量临时张量。真正释放显存的,是允许这些临时数据在使用完毕后立即消亡,而非让它们在整个前向/反向图中长期存活。
关于物理所有权,逻辑上分片的数据如果实际存储还是共享一块大内存,释放"自己那份"根本不会降低显存占用。Qwen早期实现中就踩过这个坑——逻辑上每块GPU只保留1/8的页面,但因为这些页面只是大缓冲区的切片视图,父缓冲区没有被释放,显存一点没少。用物理上独立的小缓冲区存储各自的页面,才真正把所有权落实到位。
关于并行维度,CP并行(上下文并行)和EP并行(专家并行)解决的是两个完全不同的问题,虽然可以映射到同一组GPU上,但不能互相替代。前向注意力的全局合并做到了不代表反向梯度的汇总也做到了,这两件事需要分别验证。
**说到底,这项研究想证明什么**
归根结底,这项工作想证明的是:超长上下文的强化学习训练,不一定非得靠堆砌几百上千块GPU才能实现,在合理的架构设计下,用少量GPU也能跑通超过200万位置的训练执行路径。
当然,"跑通"和"训练正确"之间还有一段距离需要弥合——分布式梯度同步、完整的DSA全局选择、与普通全序列训练的数值对比验证,这些都是团队在论文中明确点出的待完成工作。研究者们没有掩盖这些局限,而是把它们清清楚楚地列在了"限制与验证路线图"这一章里,并给出了后续工作应该按照什么顺序推进的具体建议。
这种诚实本身也是有价值的:它让读者能够清楚地区分"执行层面的可行性"和"训练正确性",避免把一个扎实的系统工程探索误读为一个完整的训练方法突破。
对于关注AI基础设施的研究者和工程师而言,这项工作开辟了一个值得深入探索的方向:在固定计算资源的约束下,通过对模型架构特性的深度理解来重新组织训练流程,而不是简单地靠资源堆量来换取更长的上下文能力。当AI模型越来越依赖超长上下文来完成复杂的智能体任务,这个方向的研究意义会随着时间推移变得越来越清晰。
---
Q&A
Q1:LongStraw为什么能用更少的GPU处理更长的上下文训练?
A:LongStraw的核心思路是把"读取长提示词"和"处理每个回答"拆开来做。提示词以不记录梯度的方式过一遍,只保留后续必需的少量状态数据(比如注意力键值页面或循环状态),然后一次只处理一个回答并立刻清除中间数据。这样GPU显存里同时存活的最大数据量从"提示词加所有回答"缩减到"关键状态加当前一个回答",从根本上绕开了全序列梯度图的显存瓶颈。
Q2:LongStraw目前的实验结果能证明它是正确的分布式训练方法吗?
A:还不能完全证明。论文明确指出,当前的成果是"执行收据",证明了210万位置的训练执行路径在物理上走得通,但存在三个尚未修复的问题:Qwen的键值投影梯度没有跨GPU汇总、GLM历史版本的稀疏注意力选择只在各GPU本地进行而非全局选择、两个模型的提示词阶段梯度都被截断。修复这些问题并与传统全序列训练做数值对比,才能建立更强的正确性证明。
Q3:GRPO训练里分组大小对显存影响有多大?
A:根据Qwen模型的实测数据,从分组大小2增加到8,峰值显存只增加了0.208GB,增幅仅0.213%,而运行时间增加了约1586秒。这是因为LongStraw对各个回答做序列化处理,显存峰值主要由单个最大回答决定,而非由回答数量决定。不过总量上,存储所有回答的输入标签、奖励和预评分结果仍然随分组大小线性增长,只是这部分数据比激活图小得多。