复旦与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层使用的是GDNGated DeltaNet门控差分网络这是一种循环机制用固定大小的状态向量来压缩历史信息就像人类用几句话总结一段对话的要点无论对话多长总结出来的关键信息大小始终固定不随文本长度增长。另外16层使用的是全注意力机制这种机制需要保存每一个历史词的完整记录就像把整段对话的录音逐字记录文本越长记录就越多存储空间呈线性增长。因此Qwen模型的关键笔记由两部分组成48个GDN层各自留下一份固定大小的循环状态加上16个全注意力层各自留下的键值页面KV Pages。这些键值页面按照上下文并行CPContext 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内存中占用72MB21个索引层的DSA键页面各占用16MB。全部78层的MLA加上21层的DSA索引每块GPU的CPU端存储约为5.81GB32块GPU合计约186GB的CPU内存用于存放提示词状态。从GLM那笔全连接隐藏缓冲区的大小可以直观感受MoE并行的压力65536个位置乘以8路路由展开后有524288行数据每行宽度6144以BF16格式存储光这一个张量就占用6GB显存。传统的全序列训练图不仅要存这个张量还要存前后各层的所有中间结果叠加下来轻易超过单卡上限。LongStraw通过提示词不建立梯度图加上每次只在回答段做一层重新计算的策略彻底绕开了这个爆显存的死局。**四、从32K到210万——一步步排雷的七个关卡**LongStraw的GLM实现不是一蹴而就的而是经历了一次典型的工程调试旅程从最小可行规模开始一个关卡一个关卡地击穿瓶颈。研究团队最先遭遇的问题是在普通的全序列训练模式下32K长度可以跑通但一旦尝试扩展到210万位置GPU显存就会溢出Out of MemoryOOM。而且溢出的位置还在不断漂移——先是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的全架构验证。第六步突破用一个只有单个回答G1的哨兵运行在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模型越来越依赖超长上下文来完成复杂的智能体任务这个方向的研究意义会随着时间推移变得越来越清晰。---QAQ1LongStraw为什么能用更少的GPU处理更长的上下文训练ALongStraw的核心思路是把读取长提示词和处理每个回答拆开来做。提示词以不记录梯度的方式过一遍只保留后续必需的少量状态数据比如注意力键值页面或循环状态然后一次只处理一个回答并立刻清除中间数据。这样GPU显存里同时存活的最大数据量从提示词加所有回答缩减到关键状态加当前一个回答从根本上绕开了全序列梯度图的显存瓶颈。Q2LongStraw目前的实验结果能证明它是正确的分布式训练方法吗A还不能完全证明。论文明确指出当前的成果是执行收据证明了210万位置的训练执行路径在物理上走得通但存在三个尚未修复的问题Qwen的键值投影梯度没有跨GPU汇总、GLM历史版本的稀疏注意力选择只在各GPU本地进行而非全局选择、两个模型的提示词阶段梯度都被截断。修复这些问题并与传统全序列训练做数值对比才能建立更强的正确性证明。Q3GRPO训练里分组大小对显存影响有多大A根据Qwen模型的实测数据从分组大小2增加到8峰值显存只增加了0.208GB增幅仅0.213%而运行时间增加了约1586秒。这是因为LongStraw对各个回答做序列化处理显存峰值主要由单个最大回答决定而非由回答数量决定。不过总量上存储所有回答的输入标签、奖励和预评分结果仍然随分组大小线性增长只是这部分数据比激活图小得多。

本月热点