1. Introduction
近年来,Muon(MomentUm Orthogonalized by Newton–Schulz)由于其相对于AdamW能够使训练更快速收敛的性能,使得它逐渐成为大模型训练中一种值得关注的优化器。与AdamW对参数逐元素进行自适应缩放不同,Muon主要作用于 Transformer 中的二维矩阵参数,并显式利用这些参数的矩阵结构。对于某个权重矩阵,Muon首先维护其梯度的momentum,随后对M进行近似正交化,并使用得到的矩阵作为实际更新方向。
这种计算结构使Muon与传统AdamW在分布式训练中的系统行为明显不同。首先,AdamW的更新是逐元素操作,天然契合任意维度的切片;而Muon要求完整的参数矩阵和梯度,这会制约数据切片的方式。其次,Muon在每一步优化器更新中额外引入了一系列大规模矩阵乘法,这也使得Offload在CPU上的优化器计算可能成为瓶颈。
本文主要基于ZeRO、MatrixFSDP、DMuon等目前的研究成果,对Muon分布式技术以及未来的offload进行一些讨论和展望。我们本次不讨论Dion等对算法进行重构的技术,而是针对普遍的NS迭代Muon。
2. Muon执行过程
Muon的一次step更新公式如下:
可以看到,Muon的计算结构主要为和一系列矩阵乘法,前者也被称为Gram Matrix。重点在于第一个公式对M的归一化,这要求M必须完整存在,从而引发了算法与系统的mismatch。
3. 针对partition的改进
目前对于DP上的数据切分,主要基于DeepSpeed ZeRO-DP,他将其定义的第一类状态(Model States)进行三阶段的切分,每个阶段依次切分优化器状态、梯度、参数。
对于ZeRO-3的原始实现以及FSDP1,数据切分的方法是先将数据进行flatten后纵向切分。这样能够保证分片大小一致,但会出现边界问题:每个DP rank保存的分片可能包含某个完整的矩阵,也可能包含完整矩阵的fragment。这与Muon的要求相悖,如果每个rank的某个梯度是不全的,必须通过all-gather获得完整的梯度,才能进行Muon update。
MatrixFSDP和DMuon选择了一种与Muon相适配的切片方法,前者称之为Matrix-aware。想法大致是,既然Muon需要完整矩阵,而我又必须要进行patition,那不如直接按照矩阵来切,也就是将不同的W看成不可拆分的原子单位。这样在Muon更新时,就完全不需要通信,本地更新即可。但问题随之而来:
第一点,参数的大小不同,self-attention层的QKV投影参数是h*3h,线性投影的参数是h*h,Embedding参数是V*h(V为词表大小),MLP层的两个参数分别是h*4h和4h*h,layernorm的参数是1*h,这要求我们尽可能将参数分成大小均衡的组。MatrixFSDP使用了三种planner来优化这一点,DMuon也尝试进行优化。但这一点是有争议的,Canzona就没有使用owner这种限制,而是把参数flatten之后按照参数边界进行拆分,这样虽然很难实现较为完美的均衡,但可以提升reduce-scatter梯度时的效率。
第二点,对于Muon的计算代价,相同大小的参数参与计算时的计算代价未必相同,这要求参数的分组在计算代价上也要均衡。DMuon进行实际profile来优化这一点。
第三点,在每一层计算之前,每一个相关参数要从其所在GPU发到所有GPU上,如果一层的参数都在同一个节点中的不同gpu中,就会导致通信不均衡问题,也就是fanout,如果不改变通信方式,就要求同一层的参数尽可能分到不同节点上。
第四点,即使使用了比较完美的优化方法,也无法保证负载大小完全相同,这样调用传统all-gather会因为负载大小不均衡而出问题,一种直接的方法是进行padding,但这样会浪费显存;或者直接优化通信原语,比如MatrixFSDP使用其定义的segment communication,其本质是send/recv。
第五点,这种切片方法的改变会导致autograd buffer问题、checkpoint策略也受到影响。
4. DMuon效率优化
通信:我们提到,如果不改变通信方式,fanout问题会成为通信瓶颈。对此,DMuon进行了一些优化,它使用二维结构定位gpu,设置一个二级通信域:节点内(intra)和节点之间(inter)。fwd需要合并参数时,持有相关参数的gpu首先将该参数发给每一个node中相同index的gpu,然后各自通过NVLink发给同节点的其他gpu;bwd则相反,梯度首先在节点内部reduce(准确来说是其定义的reduce-to-owner)到相应的index gpu,然后在外部进行通信,这种方法用NVLink替换掉了大量的InfiniBand,同时也为overlap提供了可能。从这个角度来看,Muon针对Matrix-aware重新设计了通信方式。
ovelap:DMuon将overlap优化分为iteration内部和iteration之间的overlap。iteration之间的overlap,基于参数更新后、inter域发送参数与下一次iteration开始第0层的前向计算之间,允许gpu在收到进行第0层的参数后立即开始fwd。对于iteration内部,设置prefetch策略,一个gpu在开始计算第i层时,发起第i+1层参数的intra域通信hook,同一节点内的gpu步伐可能略有不同,但这种类似于DDP overlap的机制可以天然抑制过快的gpu,使得node内的gpu步调趋于一致。Canzona对于梯度的A2A通信,使用类似DDP bucket的思想,实现通信的均衡分组。
5. DeepSpeed相关工作
DeepSpeed一开始支持ZeRO-1/2的分布式Muon,思路是:不改变原始的flatten切片(祖宗之法不可变),梯度reduce-scatter,然后用自己那份完整分片进行Muon更新。后来发现一个因为分片产生的错误,由于原始的分片方式是可能让某些参数产生碎片的,此时如果直接进行Muon计算,就不是精确的NS语义,而他们没察觉到这点。紧接着就是紧急修复,大致在2026年6月,reduce-scatter被强制为false,那么为了确保得到完整分片,就要走all-reduce,这导致了ZeRO-2的退化。但仔细想一下,如果还是使用reduce-scatter,某个梯度分片中的某些梯度不完整,完全可以再次用某些方法把这些不完整的梯度聚合起来。一个直观的方法就是加一次all-gather,但这样跟整体all-reduce也没什么区别了;其实没必要,我们只需要设定一些metadata,记录每个矩阵的owner,比如一个矩阵被拦腰切断,那么他就有两个owner,等等。有了owner,判断出自己的分片中的哪些矩阵不完整,让这些矩阵复制给所有owner各一份就可以了,deepspeed使用了几种offset元数据来完成这一功能,于是后续reduce-scatter被重新启用。
DeepSpeed对于ZeRO-3的支持值得一说。ZeRO-3明确禁止reduce-scatter,也就是说梯度默认走all-reduce,这样其实一定程度上破坏了梯度的拆分。ZeRO-3在数据上的分片与FSDP2相似,都是把每一个参数切成d份(d为DP degree),均匀分到所有DP rank中,优化器状态和梯度也就都是这个分法。在梯度all-reduce之后,每个rank有完整的梯度,但没有完整的优化器状态,所以只能得到M的分片。接下来,为了进行NS,必须得到完整的M,deepspeed使用了一种类all-gather的方法,假设:
rank0 input = [A0 B0 C0]
rank1 input = [A1 B1 C1]
rank2 input = [A2 B2 C2]
rank3 input = [A3 B3 C3]
传统的all-gather是:
ag(A0,A1,A2,A3)
ag(B0,B1,B2,B3)
ag(C0,C1,C2,C3)
deepspeed则是将每个rank的数据concrete,用大块通信替代了多次小块通信,于是经过这样的all-gather后,每个rank收到:
[A0 B0 C0
A1 B1 C1
A2 B2 C2
A3 B3 C3]
而每个rank期望收到:
[A0 A1 A2 A3]
[B0 B1 B2 B3]
[C0 C1 C2 C3]
这样就需要一次重排,恢复源码中所定义的ds_shape。于是,M构建完成,每个rank计算并更新自己的分片,比如rank0只管A,rank1只管B,rank2只管C。从这里来讲,Canzona(下一章详细阐述)比他走的更远,其直接用了A2A和group优化。此后,各自的M被更新,这需要被同步到所有rank上,这又是一次all-gather。
同时,DeepSpeed对NS算子也进行了优化。(这一篇先不论述,DMuon也对算子进行了优化)
6. 与TP的结合
如果按照Matrix-aware,TP会把一个完整参数切分为t份,放到t个GPU上。与DP不同的是,TP域一般位于节点内部,通信使用NVLink,但存在两个问题:一是为了执行optimizer.step需要all-gather梯度;二是在all-gather之后,每个TP rank需要做完全相同的Muon计算。
Canzona基于ZeRO-1,针对Megatron-LM TP进行优化,其重点在于减少冗余的Muon计算。其第一个小优化在于切片方式,正如3.1说的那样,他更注重参数的连续性,以保持梯度reduce-scatter的效率(这里可能会有点疑问,Canzona明明说基于ZeRO-1,却对梯度使用reduce-scatter而不是all-reduce。其实用了rs就可以拆梯度,从而进化到ZeRO-2了,至于为什么不这么做有待考量,个人猜测可能是懒得改Megatron optimizer)。其结合TP的算法是这样的:为每个参数设定一个owner,只不过这个owner的归属者是TP中的某个rank,而不是DP中的rank;每个rank只进行分配给自己的那部分参数的Muon计算。在bwd后,每个TP rank得到的是梯度的切片(至于按行排列还是按列排列,取决于在哪个块,TP是如何拆分的),为了得到完整的梯度,所有rank把自己的梯度分片发给对应的owner,通信上是all-to-all。每个rank得到了分配给自己的参数对应的完整梯度,就开始Muon计算,然后又是一个all-to-all,分发给相应的参数分片(TP会把参数切成t份),让他们各自进行optimizer.step。同时,为了提升all-to-all的效率,使用类似DDP的思想,将参数进行分组,只不过这里的分组条件并不是DDP那种层的次序,而是实现Muon计算的均衡。
还有一个就是Nvidia本家的Megatron TensorParallelMuon,Megatron-core提供三种方式:duplicated、distributed和blockwise。与上文相同,每个TP rank得到了对于任一参数对应梯度的分片,接下来如何做取决于那三种mode。首先是blockwise,直接对分片做Muon NS运算,这种方法不等价于标准的Muon,只能说是近似实现,这里不评价。重点在于deplicated和distributed,如果是duplicated,会执行一次all-gather,这样所有TP rank都拥有了所有完整梯度。接下来,他们进行重复的Muon计算,然后各取所需更新自己的参数分片。这样只需要一次all-gather,但需要多次重复计算,由于没有像Canzona那样对分片进行改动,所以也只能这样冗余计算。distributed模式比较有Megatron-TP的意思,Gram Matrix在这里被切片:
那么X分片如何获得?我们获得了梯度分片,可以经过一步计算得到M分片,M到X一步归一化,只需all-reduce各自的分母sum结果,最后统一除以该结果即可。
现在归一化的问题解决了,每个TP rank各自计算,然后通过all-reduce,得到总的
。
我们令:,于是有:
,在各自的分片上做这样的计算是正确的,因为:
每经过一轮NS,X发生变化,A也随之变化,所以要增加一次A的all-reduce。所以,总体的通信次数就是NS迭代次数的all-reduce和一次标量all-reduce。
7. Offload思考
目前关于Muon的分布式,这些研究所关心的方向大都是:切片方法、通信优化、算子优化以及与传统并行策略(如TP)的结合。但offload也是未来不可忽视的一个方向,DeepSpeed目前关于Muon的offload也在开展中。Muon计算需要多次大矩阵乘法的特性,就决定其很难放到CPU进行计算,而只能考虑CPU的频繁offload。一个很自然的想法是,使用GPU做offload--这很大胆,因为如果抽出GPU做纯offload,天然就减少了正常fwd/bwd的GPU,但GPU计算能力强、存储能力也不算弱、GPU之间使用NVLink的特性又使得这种想法得以产生。如果真的使用GPU进行offload,有两种可能的方向:第一种是让GPU做纯offload,只要能够承载完整的M和G,这就可以实现,代价是效率会因为工作gpu的减少而降低,假设在一个8gpu的节点中,抽出1个gpu做纯offload,那么效率天然下降12.5%,必须设计overlap(比如GPU-GPU分块传递,及时计算的流水线)来隐藏通信,以抵消效率的下降。在这种情况下,通信将会集中到该offload GPU上,这又跟之前的fanout情况相似。如果考虑节点内硬件拓扑,对于HGX H100/200平台,8张GPU通过4个NVSwitch隔离,或者对于某些硬件拓扑,8张GPU用两个NVSwitch进行隔离,那么每个隔离域选出一个GPU作为offload可能是更好的选择。第二种是GPU既做fwd/bwd,又做offload,这跟DMuon就更像一点,但可以更不均衡一点,比如设定某些GPU更偏向于offload,某些GPU更偏向做compute。更重要的是,抽调GPU作为offload,会影响到TP的性能。不管怎么说,Muon的出现使得二阶动量被消除了,现在更重要的研究方向还是计算上,还没到需要关心存储以及offload策略的时候。本文仅仅整合一些目前的方法,并对未来offload提出一些想法。若有模糊以及错误的地方,敬请读者谅解并指正。