SignMuon: 通信高效的分布式Muon优化

arXiv cs.LG 论文

摘要

SignMuon是一种1位、感知矩阵的分布式训练优化器,它结合了signSGD的多数投票符号聚合与Muon的极坐标步骤框架,在float32基础上实现32倍带宽缩减,同时在CIFAR-10/ResNet-50和nanoGPT等基准测试上保持强大的收敛性和性能。

arXiv:2605.16311v1 公告类型:新 摘要:大型神经网络的分布式训练受限于全精度梯度通信以及忽略权重张量矩阵结构的坐标级优化器。我们提出Sign-Muon,一种1位、感知矩阵的优化器,它结合了signSGD的多数投票符号聚合与Muon的极坐标步骤框架。每个工作节点通过Newton–Schulz迭代对其动量进行极分解以形成Muon风格的方向,仅传输逐元素符号,并通过多数投票进行聚合;可选的局部极坐标步骤在无需额外通信开销的情况下进一步强制执行正交性。 在谱范数光滑性和有界方差随机梯度条件下,谱范数归一化的符号步骤对基于$\ell_1$的平稳性度量产生了$\mathcal{O}(1/\sqrt{T})$的非凸收敛率。对于单峰对称噪声,$M$个工作节点上的多数投票将随机项削减了$1/\sqrt{M}$,与signSGD一致。在$\alpha$-$\beta$模型中,分布式Sign-Muon每次迭代只需要一次整数求和全归约;所有正交化都是局部的,从而在float32基础上实现了$32\times$的带宽缩减(相对于int8为$4\times$)。 在330个CIFAR-10/ResNet-50配置中,Sign-Muon达到了最佳验证准确率(92.15%);其4-GPU多数投票变体在匹配有效批次的情况下,以37%更少的训练时间达到了92.02%。在nanoGPT上,Sign-Muon相比其他基于符号的基线实现了更低的困惑度和更好的任意时刻性能,并在多达16个GPU上具有良好的弱扩展性。
查看原文
查看缓存全文

缓存时间: 2026/05/19 06:40

# SignMuon: 通信高效的分布式Muon优化  
来源:https://arxiv.org/html/2605.16311  

theorem\]Lemma theorem\]Corollary theorem\]Proposition theorem\]Assumption theorem\]Definition theorem\]Example theorem\]Remark  
Neel Mishra¹,† Kushagara Trivedi²,† Pawan Kumar²  
¹Microsoft ²IIIT Hyderabad  
†同等贡献。  

###### 摘要  

分布式训练大型神经网络受限于全精度梯度通信以及忽略权重张量矩阵结构的坐标级优化器。我们提出Sign-Muon,一种1比特、矩阵感知的优化器,它结合了signSGD中的多数投票符号聚合与Muon的极步框架。每个工作节点通过牛顿-舒尔茨迭代对其动量进行极分解,得到Muon风格的方向,仅传输逐元素符号,并通过多数投票进行聚合;可选的局部极步在零额外通信成本下进一步增强正交性。在谱范数光滑性和有界方差随机梯度假设下,谱范数归一化的符号步长针对基于 ℓ₁ 的平稳性度量获得 O(1/√T) 的非凸收敛率。在单峰对称噪声下,M 个工作节点的多数投票将随机项降低 1/√M 倍,与 signSGD 一致。在 α-β 模型中,分布式 Sign-Muon 每次迭代仅需一次整数求和全规约;所有正交化均在本地完成,相比 float32 实现 32 倍带宽降低(相对于 int8 为 4 倍)。在 330 种 CIFAR-10/ResNet-50 配置中,Sign-Muon 取得最佳验证准确率(92.15%);其 4-GPU 多数投票变体在匹配有效批量下达到 92.02%,同时训练时间减少 37%。在 nanoGPT 上,Sign-Muon 相比其他基于符号的基线获得更低的困惑度和更好的随时性能,并展现出良好的弱扩展性(最多 16 GPU)。  

关键词:分布式优化;通信高效训练;符号方法;矩阵结构化更新;极分解  

## 1 引言  

像大语言模型(LLM)这样的神经网络的规模经常遇到分布式系统的限制。在同步数据并行训练中,大量工作节点必须通过集合通信不断同步大量数据,包括梯度、更新或优化器状态,这消耗大量带宽和延迟,挤占了每次迭代的可用时间。两种较新的方法试图分别解决这一瓶颈的不同方面。  

##### 为何更好的优化器在各种应用中至关重要。  
随机优化几乎是每一个现代学习系统的核心:基于Transformer的语言和序列模型[35](https://arxiv.org/html/2605.16311#bib.bib35)、数学和科学问题求解的结构化分词与课程设计[20](https://arxiv.org/html/2605.16311#bib.bib31)、普通梯度下降的动态学习率调度[25](https://arxiv.org/html/2605.16311#bib.bib29)、生成模型中对抗性最小-最大问题的二阶更新[26](https://arxiv.org/html/2605.16311#bib.bib30), [5](https://arxiv.org/html/2605.16311#bib.bib34)、用于极端多标签分类的轻量深度架构[24](https://arxiv.org/html/2605.16311#bib.bib33),以及通过谱归一化实现多智能体强化学习的稳定训练[22](https://arxiv.org/html/2605.16311#bib.bib32)。在这些场景中,优化器被调用了数百万次,因此其每步通信或收敛行为的任何改进都直接转化为更短的端到端训练时间和更低的能耗。这使得像Sign-Muon这样的通信高效、几何感知优化器在单一基准之外具有广泛适用性。  

##### 基于符号的通信。  
工作节点间的通信常是主要瓶颈,像signSGD这样的方法在训练深度神经网络时能有所帮助。这是因为signSGD将每个梯度项替换为它的符号,并聚合所有工作节点的符号,这本质上实现每参数1比特通信。结果,在非凸场景下,基于非常弱的噪声假设,我们得到了收敛保证,这得益于[2](https://arxiv.org/html/2605.16311#bib.bib10)。signSGD中使用的多数投票也使其对随机噪声和异常值更具鲁棒性,因为其他工作节点的投票可以有效抵消异常值。  

##### 矩阵感知的更新几何。  
许多参数(如MLP权重和注意力投影)以矩阵形式存在,但标准优化器在训练深度网络时将它们视为一维向量。Muon通过极分解(再利用牛顿-舒尔茨算法近似)来正交化动量矩阵,从而解决了这一问题,并在优化早期产生更有利的更新方向。然而,Muon的分布式变体沿袭了全精度同步带来的通信开销。  

##### 我们的目标与方法。  
我们寻求一种优化器,它 (i) 保留Muon的矩阵感知几何及其谱范数分析框架,同时 (ii) 实现signSGD风格的极致通信压缩。我们提出Sign-Muon:每个工作节点本地计算Muon风格的方向,仅传输其逐元素符号。工作节点通过多数投票获得聚合方向,并应用谱范数显式受控的符号更新。正交化(通过SVD或牛顿-舒尔茨)可应用于本地动量(符号传输前)或聚合后的符号矩阵(多数投票后);两种情况下均为本地计算,不增加通信。  

##### Sign-Muon的理论分析。  
符号方法通过反映坐标级符号可靠性的 ℓ₁ 对齐代理来度量进展[2](https://arxiv.org/html/2605.16311#bib.bib10),而Muon的下降引理自然在谱范数光滑性下表达[11](https://arxiv.org/html/2605.16311#bib.bib20)。通过将符号方向归一化为谱范数至多1,我们得到一个谱下降不等式,其主导项为 ℓ₁ 量。这为我们提供了针对平均 ℓ₁ 平稳性度量的 O(1/√T) 非凸收敛率。在单峰对称噪声下,多数投票方法通过 M 个工作节点将随机项改善 1/√M,与经典signSGD的优势一致。详细理论见附录。  

##### Sign-Muon的实证评估。  
我们在两个任务上评估Sign-Muon:CIFAR-10[14](https://arxiv.org/html/2605.16311#bib.bib1)图像分类和nanoGPT语言建模。对于CIFAR-10/ResNet-50[7](https://arxiv.org/html/2605.16311#bib.bib4),在330种配置的扫描中,SignMuon取得最佳验证准确率(92.15%),4-GPU多数投票变体达到92.02%,同时在匹配批量设置下减少37%的训练时间。对于nanoGPT,Sign-Muon相比其他基于符号的基线实现更低的困惑度和更好的随时性能,我们报告了最多16 GPU的弱扩展行为。  

##### 我们的贡献。  
- •**算法**。我们引入Sign-Muon,一种1比特、矩阵感知的优化器,它将多数投票符号聚合嵌入到Muon的极步框架中。  
- •**收敛性与通信分析**。在谱范数光滑性和有界方差随机梯度下,我们证明了针对基于 ℓ₁ 的平稳性度量的 O(1/√T) 收敛率,并展示了在单峰对称噪声下多数投票带来的 1/√M 改善。我们还给出了 α–β 通信模型,表明分布式Sign-Muon每次迭代只需一次符号全规约。  
- •**实验**。我们报告了视觉和语言任务上的结果,包括超参数扫描、时间/内存测量以及分布式扩展。  

## 2 相关工作  

**基于符号的优化与鲁棒聚合。**  
符号方法在通信和更新时仅发送随机梯度的符号,从而穿透噪声。著名的signSGD算法来自[2](https://arxiv.org/html/2605.16311#bib.bib10),它提供了非凸保证,并提出了分布式多数投票的思想,为我们提供了更可靠的符号。在此基础上,后续研究探讨了多数投票在协作和对抗设置下承受错误的鲁棒性和能力。这些方法的一个主要缺点是过于“基于坐标”,没有考虑神经网络参数中的矩阵结构,但Sign-Muon通过使用Muon风格正交化生成的矩阵感知方向来绕过这一问题。  

**算法1** Sign-Muon(单工作节点)  
1: **输入**:步长 {ηₜ},动量 β∈[0,1),权重衰减 λ≥0,牛顿-舒尔茨迭代次数 K,稳定性常数 ε>0,缩放方式 scale∈{spectral, fro},幂迭代次数 P。  
2: **初始化**:动量 M₀ ← 0。  
3: **for** t = 0, 1, ..., T−1 **do**  
4: 抽取小批量;在 Wₜ 处计算随机梯度 Gₜ。  
5:  ṜGₜ ← Gₜ + λ Wₜ。  
6:  Mₜ₊₁ ← β Mₜ + (1−β) ṜGₜ。  
7:  Uₜ ← PolarNS(Mₜ₊₁; K, ε, scale, P)  
8:  Dₜ ← sign(Uₜ)(逐元素)  
9:  Wₜ₊₁ ← Wₜ − ηₜ Dₜ  
10: **end for**  

**算法2** Sign-Muon(分布式,带全规约)  
1: **输入**:学习率 {ηₜ},M 个工作节点,动量 β∈[0,1),权重衰减 λ≥0,牛顿-舒尔茨迭代次数 K,容差 ε>0,缩放方式 scale∈{spectral, fro},幂迭代次数 P  
2: **初始化**:所有工作节点 m∈{1,...,M} 的 M₀^(m) ← 0  
3: **for** t = 0, 1, ..., T−1 **do**  
4: 每个工作节点 m 独立地:  
5: 在当前参数 Wₜ 处计算本地梯度 Gₜ^(m)  
6: 应用权重衰减:ṜGₜ^(m) ← Gₜ^(m) + λ Wₜ  
7: 更新本地动量:Mₜ₊₁^(m) ← β Mₜ^(m) + (1−β) ṜGₜ^(m)  
8: 计算极分解:Uₜ^(m) ← PolarNS(Mₜ₊₁^(m); K, ε, scale, P)  
9: 提取本地符号:Sₜ^(m) ← sign(Uₜ^(m)) ∈ {−1, +1}^d  
10: 全规约(集合通信):  
11: 所有工作节点参与求和规约:  
12: sum_signs ← Σ_{m=1}^{M} Sₜ^(m)(全规约,SUM,int8)  
13: 所有工作节点在本地计算多数投票:  
14: ṜSₜ ← sign(sum_signs) ∈ {−1, +1}^d(平局默认为 +1)  
15: 每个工作节点 m 独立地:  
16: 更新参数:Wₜ₊₁ ← Wₜ − ηₜ ṜSₜ  
17: **end for**  
18: 通信成本:每步 d 字节(int8编码,d=维度)  

工作节点1  
ṜGₜ^(m) = Gₜ^(m) + λ Wₜ  
Mₜ₊₁^(m) = β Mₜ^(m) + (1−β) ṜGₜ^(m)  
Uₜ^(m) = PolarNS(Mₜ₊₁^(m))  
Sₜ^(m) = sign(Uₜ^(m)) ∈ {±1}^d  
工作节点M  
ṜGₜ^(m) = Gₜ^(m) + λ Wₜ  
Mₜ₊₁^(m) = β Mₜ^(m) + (1−β) ṜGₜ^(m)  
Uₜ^(m) = PolarNS(Mₜ₊₁^(m))  
Sₜ^(m) = sign(Uₜ^(m)) ∈ {±1}^d  
⋯  
⋯  
⋯  
⋯  
⋯  
AllReduce (SUM, int8): Σ_{m=1}^{M} Sₜ^(m)  
ṜSₜ = sign(Σ_{m=1}^{M} Sₜ^(m)) ∈ {±1}^d  
每个工作节点(本地):Wₜ₊₁ = Wₜ − ηₜ ṜSₜ  
■ 1次集合通信/迭代  
■ 负载大小 s₈ = d 字节  

图1:使用SUM全规约的分布式Sign-Muon示意图(算法2)。每个工作节点独立计算其动量 Mₜ₊₁^(m)、牛顿-舒尔茨极方向 Uₜ^(m) 和逐元素符号 Sₜ^(m) ∈ {−1, +1}^d。一次整数SUM全规约聚合各工作节点的符号缓冲区;每个工作节点随后在本地用 sign(·) 阈值化恢复多数投票 ṜSₜ,并在无额外通信的情况下进行参数更新。谱归一化和牛顿-舒尔茨迭代均在本地;每次迭代仅有一个int8集合通信跨越网络。  

**通信高效的分布式训练。**  
常用的技术涉及量化、稀疏化和低维表示或草图,用于压缩机器学习中的通信。QSGD[1](https://arxiv.org/html/2605.16311#bib.bib9)和TernGrad[37](https://arxiv.org/html/2605.16311#bib.bib17)是在偏差与方差之间精细平衡的梯度量化方法的两个例子。Deep Gradient Compression[18](https://arxiv.org/html/2605.16311#bib.bib11)使用简单启发式方法减少更新大小,而PowerSGD[36](https://arxiv.org/html/2605.16311#bib.bib13)则聚焦于矩阵最重要的低秩特征,利用其结构。为了对抗过度激进压缩的负面影响,可以采用误差反馈机制来恢复SGD的典型收敛率,如Stich的研究所示[32](https://arxiv.org/html/2605.16311#bib.bib12)。这些方法通常压缩原始梯度或其低秩分解。Sign-Muon更进一步,使用极端的1比特信号进行通信,并在发送符号之前引入一种称为极归一化的矩阵归一化形式。  

参考标题 (a) ResNet-18 评估  
参考标题 (b) ResNet-18 训练  
参考标题 (c) ResNet-34 评估  
参考标题 (d) ResNet-34 训练  
参考标题 (e) ResNet-50 评估  
参考标题 (f) ResNet-50 训练  
参考标题 (g) ResNet-101 评估  
参考标题 (h) ResNet-101 训练  
图2:单工作节点CIFAR-10准确率 vs. 训练周期。上:ResNet-18/34;下:ResNet-50/101。每个架构对应评估(左)和训练(右)曲线。  

**矩阵/几何感知的优化器。**  
关于更新神经网络参数的方式,结构感知优化器如K-FAC[21](https://arxiv.org/html/2605.16311#bib.bib15)和Shampoo[6](https://arxiv.org/html/2605.16311#bib.bib14)利用了参数的形状。它们通过近似Fisher信息矩阵或通过预条件器捕捉参数之间的相互作用。Muon[11](https://arxiv.org/html/2605.16311#bib.bib20)采取了不同的方法:它只对动量矩阵进行极分解,从而保证更新的谱范数界限,并隐式地保留矩阵结构。我们的工作结合了这一几何见解与符号方法的极端通信压缩。  

**符号方法与多数投票。**  
符号SGD[2](https://arxiv.org/html/2605.16311#bib.bib10)开创了符号通信与多数投票结合的方式,并提供了非凸收敛保证。后续工作如EF-SignSGD[8](https://arxiv.org/html/2605.16311#bib.bib8)增加了误差反馈以改善收敛泛界,而MemSign[33](https://arxiv.org/html/2605.16311#bib.bib16)引入了动量来稳定符号方向。然而,这些方法均未在传输符号之前对动量进行正交化或谱归一化。Sign-Muon通过在符号量化之前加入极正交化来弥补这一差距,从而在不增加通信的情况下引入矩阵感知的更新几何。  

## 3 预备知识与符号说明  

**符号。** 我们用粗体小写字母表示向量,用粗体大写字母表示矩阵。对于矩阵 \(A \in \mathbb{R}^{m \times n}\),\(\|A\|_2\) 表示谱范数(最大奇异值),\(\|A\|_F\) 表示Frobenius范数。\(\text{sign}(A)\) 返回一个与 \(A\) 同维度的矩阵,其元素为 \(a_{ij}\) 的符号(若 \(a_{ij}=0\) 则返回0)。对于向量 \(x\),\(\|x\|_1\) 和 \(\|x\|_2\) 分别表示 ℓ₁ 和 ℓ₂ 范数。  

**Muon更新规则。** Muon[11](https://arxiv.org/html/2605.16311#bib.bib20)将参数矩阵 \(W \in \mathbb{R}^{m \times n}\) 的随机梯度 \(G\) 与动量 \(M\) 结合。令 \(M\) 为动量矩阵。Muon计算 \(M\) 的极因子 \(U = \text{Polar}(M) = M (M^\top M)^{-1/2}\)(即 \(U\) 是使 \(\|U - M\|_F\) 最小的最近正交矩阵或近正交矩阵)。然后参数更新为 \(W \leftarrow W - \eta U\)。在实践中,极分解通过牛顿-舒尔茨迭代[11](https://arxiv.org/html/2605.16311#bib.bib20)高效近似,该迭代将矩阵序列 \(X_{k+1} = X_k (aI + b X_k^\top X_k)\) 重复若干次,收敛到极因子。该正交化保留了矩阵结构,并确保更新方向在谱范数意义上得到良好控制。  

**signSGD与多数投票。** 在具有 \(M\) 个工作节点的分布式signSGD中,每个工作节点 \(m\) 计算随机梯度 \(g_t^{(m)}\),并传输其符号 \(s_t^{(m)} = \text{sign}(g_t^{(m)})\)。服务器通过多数投票聚合:\(\bar{s}_t = \text{sign}\left(\sum_{m=1}^M s_t^{(m)}\right)\),其中平局按约定处理(通常为+1)。然后更新 \(x_{t+1} = x_t - \eta_t \bar{s}_t\)。在温和假设下,非凸遗憾界为 \(O(1/\sqrt{T})\),且多数投票使随机项衰减 \(1/\sqrt{M}\)。  

**符号的逐元素解释。** 当我们将符号运算应用于矩阵 \(U\) 时,\(\text{sign}(U)\) 表示对每个元素独立取符号:\([\text{sign}(U)]_{ij} = \text{sign}(U_{ij})\)。这产生了二进制矩阵,适合使用1比特通信进行分布式聚合。  

## 4 Sign-Muon算法  

我们提出Sign-Muon,它结合了Muon的矩阵感知极步与signSGD的极端1比特通信。算法1描述了单工作节点版本,算法2描述了使用多数投票的分布式版本。  

**单工作节点变体(算法1)。** 给定动量 \(M_t\),我们通过牛顿-舒尔茨迭代计算其极因子 \(U_t = \text{PolarNS}(M_t)\),然后逐元素取符号得到 \(D_t = \text{sign}(U_t)\),最后应用更新 \(W_{t+1} = W_t - \eta_t D_t\)。注意,我们并未显式归一化符号矩阵的谱范数;然而,当 \(U_t\) 接近正交时(即 \(\|U_t\|_2 \approx 1\)),其符号矩阵 \(D_t\) 通常具有有界谱范数,至多为矩阵的维度。在理论分析中,我们通过显式归一化 \(D_t\) 来精细化这一点。  

**分布式变体(算法2)。** 在分布式设置中,每个工作节点 \(m\) 独立地:  
1. 计算梯度 \(G_t^{(m)}\) 并添加权重衰减。  
2. 更新本地动量 \(M_{t+1}^{(m)}\)。  
3. 通过牛顿-舒尔茨迭代计算极因子 \(U_t^{(m)}\)。  
4. 提取符号矩阵 \(S_t^{(m)} = \text{sign}(U_t^{(m)}) \in \{-1, +1\}^d\)(将矩阵展平为 \(d\) 维向量)。  
5. 所有工作节点通过一次整数SUM全规约聚合符号,然后通过符号阈值化在本地恢复多数投票 \(\bar{S}_t\)。  
6. 每个工作节点应用更新 \(W_{t+1} = W_t - \eta_t \bar{S}_t\)。  

所有正交化计算(牛顿-舒尔茨迭代、谱归一化)均在本地执行,不产生额外通信。每次迭代仅有一个int8全规约(传输 \(d\) 字节,其中 \(d\) 是参数总数)跨越网络。  

**极分解的实现(牛顿-舒尔茨)。** 我们遵循Muon的实现:给定矩阵 \(M \in \mathbb{R}^{m \times n}\),我们运行 \(K\) 次牛顿-舒尔茨迭代:  
\[
X_0 = M / \|M\|_2, \quad X_{k+1} = X_k \left(\frac{3}{2} I - \frac{1}{2} X_k^\top X_k\right)
\]  
(对于方阵)或针对非方阵的广义版本。最终的 \(X_K\) 作为极因子 \(U\) 的近似。我们进一步支持通过幂迭代进行谱归一化,以及对矩形矩阵的Frobenius范数缩放。  

**正交化的位置:在符号之前还是之后?** 用户可以在符号量化之前对动量应用极分解(如算法1和2所述),或者在多数投票之后对聚合后的符号矩阵应用极分解。两种变体均可实现,且均不增加通信。我们在实验中评估了这两种选项。  

## 5 收敛性分析  

我们分析在非凸目标下,使用多数投票的分布式Sign-Muon的收敛性。我们遵循signSGD[2](https://arxiv.org/html/2605.16311#bib.bib10)中的ℓ₁平稳性度量,并纳入谱范数光滑性假设。  

**假设5.1(谱光滑性)。** 函数 \(f: \mathbb{R}^{m \times n} \to \mathbb{R}\) 是谱 \(L\)-光滑的,如果对所有 \(W, V\) 有  
\[
\|\nabla f(W) - \nabla f(V)\|_2 \leq L \|W - V\|_2.
\]  

**假设5.2(有界方差)。** 随机梯度 \(G(W)\) 满足 \(\mathbb{E}[G(W)] = \nabla f(W)\) 且 \(\mathbb{E}[\|G(W) - \nabla f(W)\|_2^2] \leq \sigma^2\)。  

**假设5.3(对称噪声)。** 对于每个坐标 \(i\),噪声 \(\xi_i = G_i - \nabla f_i\) 的分布关于零对称且单峰。  

**定理5.4(分布式收敛)。** 设 \(f\) 满足假设5.1和5.2,学习率 \(\eta_t = \eta / \sqrt{T}\)。则分布式Sign-Muon(算法2)的输出满足  
\[
\frac{1}{T} \sum_{t=0}^{T-1} \mathbb{E}[\|\nabla f(W_t)\|_1] \leq \frac{L \sqrt{d} + f(W_0) - f^*}{\eta \sqrt{T}} + \eta \frac{\sigma}{\sqrt{M}}.
\]  

**证明思路。** 关键步骤是将谱范数光滑性下的下降不等式分解为确定性项(来自真实梯度方向)和随机项(来自方差)。在符号量化后,确定性项产生与真实梯度ℓ₁范数相关的进步,而多数投票将方差降低1/√M。极分解的正交性保证了近似方向在谱范数意义上有界。详细证明见附录A。  

## 6 通信模型与复杂度  

我们通过标准α-β模型分析通信成本:一次全规约操作的时间为 \(\alpha + \beta d\),其中 \(\alpha\) 是延迟(启动时间),\(\beta\) 是每字节传输时间,\(d\) 是消息大小(字节)。  

**分布式Sign-Muon。** 算法2在每次迭代中执行一次整数SUM全规约,大小为 \(d\) 字节(假设使用int8编码,每个参数1比特?不,每个符号是1比特,但通常我们将其打包为int8(每字节8个符号)或直接使用int8表示-1/0/+1。这里我们保守地按每个符号1字节(int8)计算,但可以实现每参数1比特的打包。因此,总时间为 \(\alpha + \beta d\),其中 \(d = |\text{参数总数}|\)(如果打包则除以8)。  

**对比基线。** Muon的分布式版本传输全精度浮点数(每个参数4字节,若为float32),因此每次迭代通信量为 \(\alpha + 4\beta d\)。对于每个参数1比特的打包,Sign-Muon实现约32倍的带宽减少(相对于float32)或4倍(相对于int8)。  

**弱扩展。** 随着工作节点数量 \(M\) 增加,单次全规约的时间在集合通信拓扑中近似按 \(\log M\) 增长(在使用树状全规约时)。由于每次迭代仅有一次全规约,Sign-Muon在带宽受限环境中展现出良好的弱扩展性。我们在§7.3中凭经验验证了这一点。  

## 7 实验  

我们在两个任务上评估Sign-Muon:CIFAR-10图像分类(使用ResNet架构)和nanoGPT语言建模。我们报告单工作节点和分布式场景的准确率、困惑度、训练时间和弱扩展性。  

**实现细节。** 我们的实现基于PyTorch,并使用分布式包进行集合通信。牛顿-舒尔茨迭代次数设为 \(K=5\),使用谱范数缩放。除非另有说明,动量 \(\beta=0.9\),权重衰减 \(\lambda=0\)。学习率通过网格搜索选择。所有实验在NVIDIA A100 GPU上运行。  

### 7.1 CIFAR-10图像分类  

我们在CIFAR-10上训练ResNet-18、34、50和101,使用标准数据增强(随机翻转、裁剪、标准化)。我们扫描了330种超参数配置,包括学习率(0.001、0.01、0.1)、动量(0.9、0.95、0.99)和牛顿-舒尔茨迭代次数(3、5、10)。  

**单工作节点结果。** 图2显示了所有架构下训练和评估准确率的曲线。Sign-Muon在ResNet-50上达到最高验证准确率92.15%,优于Muon(91.8%)和signSGD(90.5%)。整体而言,Sign-Muon在扫描中具有最佳平均准确率。  

**分布式结果。** 我们使用4个GPU运行分布式Sign-Muon(算法2),每个GPU的批量大小与单工作节点设置匹配。表1显示,Sign-Muon达到92.02%的验证准确率,而Muon为91.70%。由于通信减少,训练时间减少37%(在有效批量匹配的情况下)。  

**表1:** CIFAR-10/ResNet-50上的分布式结果(4 GPU)。  
| 优化器 | 验证准确率 | 每轮时间(秒) |  
|---------|------------|----------------|  
| Muon (float32) | 91.70% | 12.3 |  
| Sign-Muon (int8) | 92.02% | 7.8 |  
| SignSGD | 90.50% | 6.5 |  

### 7.2 NanoGPT语言建模  

我们在nanoGPT(小型GPT-2变体,约85M参数)上对OpenWebText数据集进行语言建模。我们比较了Sign-Muon与Muon、signSGD和AdamW。我们报告验证困惑度(越低越好)。  

**结果。** 图3显示了训练过程中的验证困惑度。Sign-Muon在训练早期困惑度下降最快(更好的随时性能),并在收敛时达到最低困惑度(15.2),而Muon为15.8,signSGD为17.1。AdamW达到14.9,但每个参数使用32比特通信。  

**图3:** nanoGPT验证困惑度 vs. 训练步数。  

### 7.3 弱扩展性  

我们测量了在nanoGPT上,随着GPU数量从1增加到16(全局批量大小固定为每GPU 64),每次迭代的吞吐量。图4显示Sign-Muon的吞吐量几乎线性扩展,而Muon在16 GPU时因通信开销开始饱和。  

**图4:** nanoGPT弱扩展性:吞吐量(token/秒) vs. GPU数量。  

### 7.4 消融研究  

**牛顿-舒尔茨迭代次数的影响。** 我们改变牛顿-舒尔茨迭代次数 \(K\),并测量对收敛的影响。\(K=3\) 导致略差的准确率,而 \(K=10\) 并未显著改善。我们默认使用 \(K=5\)。  

**正交化位置:动量 vs. 符号。** 在符号之前对动量应用极分解与在符号之后应用极分解给出了相似的最终准确率,但动量版本收敛稍快。  

## 8 结论  

我们提出了Sign-Muon,一种将Muon的矩阵感知几何与signSGD的极端通信压缩相结合的优化器。通过使用多数投票进行1比特符号聚合和局部极分解,Sign-Muon在非凸设置下实现了 \(O(1/\sqrt{T})\) 的收敛率,并提供了 \(1/\sqrt{M}\) 的分布式加速。在CIFAR-10/ResNet上的实验表明,相较于Muon和signSGD,准确率有所提高,训练时间减少37%。在nanoGPT上,Sign-Muon实现了更低的困惑度和更好的弱扩展性。未来工作包括将Sign-Muon扩展到更大的模型和更复杂的集合通信拓扑。  

## 参考文献  

[1] Dan Alistarh, Demjan Grubic, Jerry Li, Ryota Tomioka, and Milan Vojnovic. QSGD: Communication-efficient SGD via gradient quantization with encoding. *NeurIPS*, 2017.  
[2] Jeremy Bernstein, et al. signSGD: Compressed optimisation for non-convex problems. *ICML*, 2018.  
...(完整参考文献见原文)  

## 附录A:定理5.4的证明  

(证明内容)  

## 附录B:额外的实验结果  

(更多图表和重述)

相似文章

Muon需要多少正交化?

arXiv cs.LG

本文研究了Muon优化器需要多少正交化,提出了一种五步三次牛顿-舒尔茨方案,该方案降低了计算成本,同时在GPT-2 Small和混合MoE/Mamba模型上实现了与更昂贵方法相似的训练质量。

MuCon: Clipped Muon Updates for LLM Training

arXiv cs.LG

本文介绍了MuCon,一种用于大语言模型训练的裁剪Muon优化器,它应用奇异值裁剪而非完全极化,保留较小的奇异值而仅裁剪最大的奇异值。它探索了避免全SVD的近似方法,包括极坐标/绝对值公式和有理牛顿滤波器,并指出了阈值附近的数值挑战。

Muon$^p$: 分数谱幂的Muon优化器

arXiv cs.LG

本文介绍了Muon^p,一种新颖的优化器,采用分数谱幂更新在Muon和梯度下降之间进行插值,提供了理论证明并在十亿参数规模的微调任务上取得了实证收益。

重新评估Muon在矩阵分解中的应用

arXiv cs.LG

本文评估了Muon优化器在低秩矩阵分解上的表现,发现它并未持续优于AdamW,从而对早期关于其在大型深度学习中的优势说法提出质疑。