微信扫一扫,关注公众号

  • 科技行者

  • 算力行者

见证连接与计算的「力量」

首页 把每一块显存都榨干:训练超长上下文MoE模型的四个隐藏陷阱

把每一块显存都榨干:训练超长上下文MoE模型的四个隐藏陷阱

2026-10-02 16:19
分享至:
----..---.-...-/--...-.-......./-...-....-..--../-............-.- ----..---.-...-/--...-.-......./-...-....-..--../-............-.- ----..---.-...-/--...-.-......./-...-....-..--../-............-.- ----..---.-...-/--...-.-......./-...-....-..--../-............-.-
2026-10-02 16:19 • 科技行者

你有没有想过,训练一个几百亿参数的AI模型时,最先崩溃的往往不是算力,而是显存?

这事说起来挺反直觉的。大家平时聊起大模型训练,总觉得瓶颈是算力不够、GPU太少。但实际情况往往是,一台训练任务跑得好好的,突然就因为某个环节内存爆了,整个训练直接中断。更麻烦的是,这个"某个环节"每次可能都不一样。有时候是路由矩阵撑爆了显存,有时候是词表投影层吃掉了所有空间,有时候又是优化器状态占满了内存。

Salesforce AI Research 的研究者们在这篇论文里指出了一个很朴素但常被忽视的道理:训练能不能成功,不取决于平均显存占用,而取决于每一个组件的峰值占用有没有超过设备容量。哪怕你把三个瓶颈都压下去了,只要第四个还在疯长,训练照样会在某个临界点崩掉。

这就好比一个装修队给一栋楼做防水,四面墙分别请了四个不同的师傅施工,前三面墙都做得天衣无缝,第四面墙却随便糊了糊。下雨的时候,水照样会从第四面墙渗进来,前面三面墙做得多好都没用。楼漏水这件事只需要一个薄弱点,不需要四个。如果这四个师傅不能同时把关,楼迟早要出问题。

MoE模型训练里的四个"薄弱墙面"

混合专家模型

MoE:Mixture-of-Experts,混合专家模型,一种让不同"专家"网络分别处理不同输入的架构,可以在不显著增加计算量的前提下大幅扩展模型参数量

之所以特别容易在长上下文或大批量训练时出问题,是因为它天生带着四个会随着配置变化而疯长的"计算包袱"。

第一个包袱是专家调度。在MoE架构里,每个GPU只负责一部分专家,但输入的token(可以理解成文本被切分后的最小处理单元)要先被路由到它该去的专家那里。如果路由不均衡,某几个专家突然被大量token选中,负责这些专家的GPU瞬间就要处理海量数据,内存占用直接飙升。

第二个包袱是词表投影层。模型输出的时候,需要把每个token的表示映射到整个词表上,算出每个词的概率。这一步产生的张量大小是"token数乘以词表大小",词表通常有几万甚至几十万个词,长上下文一来,这个张量能轻松撑爆显存。

第三个包袱是梯度检查点。

梯度检查点:训练时为了省显存,不保留每一层的中间计算结果,反向传播时需要用到就重新算一遍,用计算时间换内存空间

这个技术本身是为了省内存的,但每个检查点边界处仍然要保留一个"活着"的输入张量,直到反向传播真正用上它。层数越深、序列越长,这些悬而未决的张量加起来也是一笔不小的开销。

第四个包袱是优化器状态。像AdamW这类优化器,需要给每个参数额外存一份动量、一份二阶矩估计,参数量一大,这部分状态占的内存能轻松超过模型本身。

论文里说得很直白:这四个包袱谁先"爆表"完全取决于具体配置。词表大、上下文长,词表投影先出问题;路由本身不均衡,专家调度先崩;网络很深,检查点边界先撐不住;参数量巨大但设备数量有限,优化器状态先超标。

这也是为什么之前很多方案总是"治标不治本"。你解决了一个问题,换个场景另一个问题立刻冒出来,像打地鼠游戏一样,永远打不完。

四个针对性的解决方案

论文提出了四个操作,分别对准上面四个包袱,每一个都只改变计算的顺序和粒度,不碰模型本身、不碰精度、不碰损失函数的计算方式。换句话说,训练出来的结果和标准全参数BF16训练完全一致,不是那种"用精度换内存"的取巧做法。

先说专家调度这块。之前已经有一个叫LLEP的方法(Least-Loaded Expert Parallelism,最闲专家并行),通过把热门专家的溢出任务挪给闲着的GPU,来解决负载不均衡的问题。但LLEP有个问题:它虽然把任务挪走了,却还是一次性把整批数据都塞进内存,该来的那批数据有多大,内存占用就有多大。

论文提出的PipelinedLLEP做了一个关键改动:把整批数据切成一个个小块(chunk),规定每个GPU往每个小块里最多能塞多少token,然后用流水线的方式,让传输和计算重叠进行。这样一来,无论路由多不均衡,单个小块的大小是被严格限制住的,不会因为某个专家突然爆红就让内存跟着炸开。

这就像是一个食堂打饭,如果让所有想吃红烧肉的人一次性冲上来打饭,窗口肯定被挤爆。但如果规定每一轮最多放20个人进来打饭,不管这一轮里有多少人想吃红烧肉,这一轮的压力都是可控的。要是不设这个限流规则,某一天突然全食堂的人都想吃红烧肉,窗口直接被冲垮,没人能吃上饭。

不过这里有个细节值得说说:如果单纯把批次切块,并行执行计算和通信,内存峰值其实降不下来,因为一个循环调用多次的时候,每次调用产生的计算图(autograd图)都会一直挂在内存里,直到反向传播算完才释放。论文用了一个嵌套的梯度检查点技巧解决了这个问题,把每个小块的专家计算单独包一层可重入检查点,让前一个块的中间结果能及时释放,后一个块才分配自己的空间。这个细节听起来技术性很强,但本质就是"用一次就扔",而不是"攒着最后一起处理"。

实测数据很直接:在65536个token、128个专家、top-8路由的配置下,标准专家并行一旦碰上严重的路由不均衡直接内存溢出崩溃,而PipelinedLLEP相比LLEP能省下56.9%到59.3%的峰值内存,速度还基本不受影响,甚至有的配置下反而更快。

词表投影:让巨大的logit张量"隐形"

第二个方案叫Ring-DTP,针对的是词表投影这个环节。

要理解这个方案,得先明白为什么词表投影这么费内存。计算交叉熵损失(cross-entropy loss,衡量模型预测和真实答案差距的常用损失函数)的时候,理论上每个token只需要三个数:一个最大值、一个指数和、以及目标词对应的那个logit值。但传统做法是先把整个"token数乘词表大小"的logit矩阵算出来,再从里面提取需要的那几个数。这就好比为了知道一群人里谁最高,先把每个人的详细体检报告全部打印出来堆在桌上,而实际上你只需要一个数字。

之前已经有一些单机上的融合交叉熵内核(比如Cut Cross-Entropy和Liger Kernel),通过在线维护log-sum-exp的方式避免生成完整的logit矩阵。但这些方法有个前提:整批数据要在同一块GPU上。而Megatron的做法虽然把权重分片了,却要求每个GPU持有相同的批次数据,这直接砍掉了有效批量大小。

Ring-DTP的做法是让每个GPU保留自己独有的一批数据,同时也持有词表的一部分权重分片,然后让数据或权重像接力赛一样在GPU之间循环传递,直到每一批数据都和每一片权重"见过面"。每次相遇只计算一小条logit,提取出需要的三个统计量后立刻释放这条logit,绝不让完整的大矩阵成型。

这有点像一个大型联谊活动,如果要求所有参与者同时挤进一个大厅认识每一个人,场地压力巨大。但如果换成轮转制,每一轮只有一小批人和一小批人碰头,聊完记录一下关键信息就换下一轮,场地压力就小多了,而且最终每个人还是能认识到所有该认识的人。如果不这么设计,非要一次性凑齐所有人,场地根本装不下这么多人。

实验数据显示,在8路分片、16384个token的配置下,标准方法的峰值内存是42.5GB,而Ring-DTP只用了7.3GB,省了82.8%,时间只多花了5.1%。当token数翻倍到32768时,标准方法的内存几乎翻倍涨到79.5GB,Ring-DTP却只涨到10.6GB,说明这个方法在长上下文场景下的优势会越来越明显。这正是让词表投影层能在百万token上下文里跑起来的关键。

检查点边界:把"睡着的"张量挪到CPU去

第三个方案叫SCO(Selective Checkpoint Offload,选择性检查点卸载),解决的是梯度检查点技术留下的一个尾巴问题。

梯度检查点省内存的方式是"不存中间结果,需要时重新算",但每一层的输入张量必须一直留在显存里,从前向传播开始一直等到反向传播真正重新计算这一层为止。这个等待期可能很长,尤其是层数很多的时候,相当于很多个"半成品"张量同时占着显存位置,谁也不动,就是干占着。

SCO的思路很简单:把这些暂时用不上但又必须留着的张量,先挪到CPU内存里存着,快用到的时候再提前一层取回来。这样GPU显存里同时最多只需要留两个正在恢复的张量,而不是所有层的输入张量全部囤在那里。

这就像是搬家公司打包东西,如果所有箱子都堆在客厅里等着装车,客厅会被堵得进出不了。但如果把还没轮到装车的箱子先挪到走廊或者阳台上,等快轮到了再搬回客厅门口,客厅始终只有一两个箱子占地,整个流程一样能顺利完成,只是多了一次搬运的功夫。要是不做这个中转,客厅堆满了箱子,连走路的空间都没有,搬家反而更慢。

论文用gpt-oss-20b模型做了测试,给CPU内存不同的预算,GPU的显存峰值确实随着预算的增加单调下降,吞吐量的变化不到2%,但省下的显存换来的是能跑更大的批次,最大批次提升了17.7%。这说明SCO这种"分批卸载"的策略,在几乎不影响速度的前提下,实实在在地扩大了训练的可承受空间。

优化器更新:别让GPU干等着

第四个方案叫OffloadStreamAdamW,针对的是优化器状态卸载这件事本身带来的新问题。

前面提到,把优化器状态(比如AdamW需要的动量和二阶矩)放到CPU内存里能省显存,这是ZeRO-Offload等已有方法的做法。但问题是,CPU算这个更新的速度远不如GPU,更新期间GPU完全闲着,而且此时GPU的显存本来就因为激活值已经释放而处于空闲状态,这段时间等于是资源双重浪费。

OffloadStreamAdamW的解法是反过来利用这个"空闲期":把CPU上存的参数状态分成一个个小批次(bucket),轮流传到GPU上,让GPU来做真正的计算,算完再传回CPU保存。整个过程用三条并行的流水线(传输、计算、写回)交替进行,让GPU不再是干等待着,而是变成真正干活的角色。

这就好比一个仓库管理员,原本的做法是把所有货物一次性搬到隔壁小屋去逐个称重登记,称重的人手速很慢,搬运工在旁边干站着没事干。改进后的做法是让搬运工按批次不停地搬运货物,称重的人也不停地称重,两边同时忙碌,谁都不闲着。如果继续用老办法一次性全搬过去再干等,搬运工的体力(GPU算力)就白白浪费了。

实测效果是,相比CPU上的AVX向量化AdamW实现,OffloadStreamAdamW把优化器更新这一步从3.95秒降到了1.93秒,快了2.05倍。而且论文还发现,增加缓冲槽位(staging slots)数量并不能进一步提速,说明这个流程本身已经被传输带宽卡住了,这是理论上能达到的速度上限附近。

四个方案合体:一百万token上下文成真

把这四个方案组合起来放进一个叫MoP(Mixture-of-Parallelisms,混合并行)的整体架构里,论文在120B、241B、667B三个不同规模的MoE模型上做了端到端测试。

结果相当亮眼:对比精心调优过的FSDP2(一种常见的分布式训练框架配置)基线,组合方案能训练一百万token的上下文长度,是基线能达到的上下文长度的8到32倍。而在两者都能跑的最长上下文长度上做比较,新方案的吞吐量还能达到基线的7.6到10.4倍。

这个数字差距挺夸张的。FSDP2基线在241B模型这个规模上,超过32K token就直接内存溢出,而新方案的最短测试配置就是128K,是基线极限的四倍起步。批量大小方面同样有优势,最大能跑的全局批次是基线的3到12倍。

论文还专门做了个训练质量的验证实验,用gpt-oss-20b模型在数学题数据集上做微调,对比新方案和FSDP2基线训练出来的模型,在AIME 2025测试集上的准确率分别是59.8%和59.6%,几乎没有差异。这说明省内存这件事没有偷工减料,模型学到的东西是一样的。

这几个方案分别针对四个不同的内存瓶颈,而且论文特意强调每一个都可以单独启用,不依赖其他三个。这个设计思路挺聪明的:不同的训练任务遇到的瓶颈不一样,有的模型词表特别大,有的路由特别不均衡,让用户按需开启相应的方案,而不是被迫承担一整套复杂系统的全部开销。

写在后面

读这篇论文的过程中,最触动我的其实不是某一个具体的技术细节,而是论文标题里那句"flattening every memory peak"背后的思维方式。

大部分工程优化文章讲的是"我们把X降低了多少",但这篇论文一开篇就先给你讲清楚一个残酷的事实:降低一个峰值,如果留着另一个峰值不管,训练照样跑不起来。这种"整体约束"的视角,其实比单点优化更难做,因为你得同时理解四个完全不同的系统组件,还要保证它们之间不互相冲突。

论文里有个细节我觉得值得单独说一说:在讲PipelinedLLEP的时候,他们发现光靠"流水线重叠"这个手段并不足以真正降低内存峰值,还需要配合嵌套的梯度检查点技巧才能把每个块的中间结果及时释放掉。这说明很多看似朴素的"分块处理"思路,实际落地时往往藏着一层容易被忽略的坑,通信重叠解决的是速度问题,内存峰值的问题需要另外的机制去处理,这两件事看起来相关,其实是两个独立的维度。

另外一个有意思的地方是,论文附录里专门讨论了一个"延迟出现双峰分布"的现象:在特定的分块数量下,同样的配置有时候跑得快,有时候莫名其妙慢了一大截,而且这种慢的模式一旦触发就会稳定重现。研究者最后发现根源出在主机端的启动模式上,插入一次同步操作就能让这个慢模式转移到别的配置去。这种"玄学问题"最后被系统性地定位出来,而不是简单归因于"随机波动",这种排查思路本身也挺值得学习的。

这篇论文没有解决的问题其实也挺明显:四个参数(序列并行度、专家并行度、投影并行度、chunk大小、bucket大小)目前还是靠人工调参和实测曲线来选,论文自己也在最后承认这一点。如果未来能有一个自动化的方式根据模型结构和硬件配置直接推荐这套参数,那才是真正把这套方法变成"开箱即用"的工具,而不是需要专家经验才能用好的精密仪器。

Q&A

Q1:论文提出的四个内存优化方法分别解决什么问题?

A:分别对应MoE训练中四个会不受控增长的内存瓶颈:PipelinedLLEP限制专家调度时的token缓冲区大小,Ring-DTP让词表投影不用生成完整的巨大logit矩阵,SCO把梯度检查点的部分张量卸载到CPU内存,OffloadStreamAdamW用GPU流水线加速CPU优化器状态的更新过程。

Q2:这套方法相比传统训练方案能带来多大的提升?

A:在120B到667B参数规模的MoE模型上,组合使用这四个方法后能训练一百万token长度的上下文,是调优后FSDP2基线能达到长度的8到32倍,在相同上下文长度下吞吐量最高能达到基线的10.4倍,最大批量能达到基线的12倍。

Q3:使用这些内存优化方法会不会影响模型训练效果?

A:不会。论文特意做了对比实验,用同样的数据和超参数分别训练模型,新方案和传统FSDP2基线在AIME 2025数学测试集上的准确率分别是59.8%和59.6%,几乎没有差异,说明省内存并没有牺牲训练质量。

分享至
0赞

好文章,需要你的鼓励

推荐文章
----..---.-...-/--...-.-......./-...-....-..--../-............-.- ----..---.-...-/--...-.-......./-...-....-..--../-............-.- ----..---.-...-/--...-.-......./-...-....-..--../-............-.- ----..---.-...-/--...-.-......./-...-....-..--../-............-.-