关于分布式Muon的想法整合及讨论
1. Introduction近年来MuonMomentUm 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*hEmbedding参数是V*hV为词表大小MLP层的两个参数分别是h*4h和4h*hlayernorm的参数是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发给同节点的其他gpubwd则相反梯度首先在节点内部reduce准确来说是其定义的reduce-to-owner到相应的index gpu然后在外部进行通信这种方法用NVLink替换掉了大量的InfiniBand同时也为overlap提供了可能。从这个角度来看Muon针对Matrix-aware重新设计了通信方式。ovelapDMuon将overlap优化分为iteration内部和iteration之间的overlap。iteration之间的overlap基于参数更新后、inter域发送参数与下一次iteration开始第0层的前向计算之间允许gpu在收到进行第0层的参数后立即开始fwd。对于iteration内部设置prefetch策略一个gpu在开始计算第i层时发起第i1层参数的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必须得到完整的Mdeepspeed使用了一种类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 C0A1 B1 C1A2 B2 C2A3 B3 C3]而每个rank期望收到[A0 A1 A2 A3][B0 B1 B2 B3][C0 C1 C2 C3]这样就需要一次重排恢复源码中所定义的ds_shape。于是M构建完成每个rank计算并更新自己的分片比如rank0只管Arank1只管Brank2只管C。从这里来讲Canzona下一章详细阐述比他走的更远其直接用了A2A和group优化。此后各自的M被更新这需要被同步到所有rank上这又是一次all-gather。同时DeepSpeed对NS算子也进行了优化。这一篇先不论述DMuon也对算子进行了优化6. 与TP的结合如果按照Matrix-awareTP会把一个完整参数切分为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 TensorParallelMuonMegatron-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得到总的。我们令于是有在各自的分片上做这样的计算是正确的因为每经过一轮NSX发生变化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提出一些想法。若有模糊以及错误的地方敬请读者谅解并指正。

相关新闻