优化器 (3)(Gram Newton-Schulz)
引言
在上一篇文章中,我们深入介绍了 Muon 优化器的数学原理。Muon 的核心创新在于使用 Newton-Schulz 迭代来近似矩阵的 msign 函数,也称为极分解(polar decomposition),从而实现对参数矩阵的正交化更新。
然而,标准的 Newton-Schulz 算法存在一个显著的计算瓶颈:它需要在每次迭代中对原始的矩形参数矩阵进行多次矩阵乘法运算。对于一个
近期,DAO Lab 提出了 Gram Newton-Schulz 算法,通过以下创新显著加速了 Muon 优化器:
- 数学等价性:在 Gram 矩阵
上迭代而非原矩阵 ,输出与标准 Newton-Schulz 数学等价 - 数值稳定性:通过重启策略和 float16 算术,实现与标准方法相当的稳定性
- 硬件感知:对称 GEMM 核充分利用 Hopper/Blackwell 架构特性
在 Kimi K2 这种万亿参数量级别的模型上,Gram Newton-Schulz 可以将优化器时间减少高达 50%。
一、Muon 与 Newton-Schulz 回顾
1.1. Muon 优化器更新规则
在第
其中
Definition 1(极分解):若
由于精确计算
1.2. 标准 Newton-Schulz 迭代
Newton-Schulz 是一种基于矩阵多项式的迭代方法。从初始矩阵
其中
下面分析标准 Newton-Schulz 迭代的计算复杂度。
令
每次迭代包含三次矩阵乘法:
: FLOPs (其中 ): FLOPs (其中 ): FLOPs
总计算量为:
当
因此,标准的 Newton-Schulz 迭代存在以下问题:
- 未利用对称性:
和 都是对称矩阵,但标准实现未利用这一结构。 - 强依赖
:当 时,计算量主要由昂贵的 GEMM 操作主导。
二、Gram Newton-Schulz 的数学原理
2.1. 核心思想
Gram Newton-Schulz 的核心洞察在于下面的公式 (4)(证明见上一篇博客)。
基于公式 (4),Gram Newton-Schulz 的策略是:
- 计算
Gram 矩阵 - 使用迭代方法近似
- 计算
Gram Newton-Schulz 的关键优势在于,第 2 步(占据绝大部分计算量)仅在小的
2.2. 从 Newton-Schulz 到 Gram Newton-Schulz
Proposition 1(奇异向量不变性):设
其中
根据 Proposition 1,如果多项式序列
那么,如何将迭代多项式方法
Lemma 1(奇次多项式的性质) 每个奇次多项式
因此,如果
Theorem 1(迭代求逆平方根):设任取
则有
其中
迭代的起始条件为
注意到当
2.3. Naive Gram Newton-Schulz 算法
将定理 1 的迭代提升到矩阵形式,作者提出了 Naive Gram Newton-Schulz 算法:
算法 1(Naive Gram Newton-Schulz)
输入:
-
(将奇异值归一化到 , ) -
-
-
对于
按照下面的公式迭代 (9) -
返回
2.4. 计算复杂度分析
下面分析 Naive Gram Newton-Schulz 的 FLOPs 数量级。每次迭代包含四次矩阵乘法(使用对称 GEMM 核):
: FLOPs : FLOPs : FLOPs
初始化和输出步骤:
: FLOPs : FLOPs(非对称)
下面的计算量可以优化:
不需要计算:节省 FLOPs 不需要计算:节省 FLOPs
因此 Naive Gram Newton-Schulz 的总 FLOPs 为:
当
使用对称 GEMM 的标准 Newton-Schulz 需要
在 Muon 的场景下(
- 相比使用对称 GEMM 的标准 Newton-Schulz 节省 55% FLOPs
- 相比典型实现(无对称 GEMM)节省 68% FLOPs
三、Stabilized Gram Newton-Schulz
3.1. Naive Gram Newton-Schulz 的数值稳定性问题s
尽管 Naive Gram Newton-Schulz 在精确算术下与标准 Newton-Schulz 等价,但在有限精度(尤其是半精度)下表现较差。实验发现,使用 bfloat16 训练 Transformer 时损失函数曲线会频繁出现尖峰,最终输出充满 Inf。
作者通过研究特征值和奇异值的演化,给出了矩阵发散的几个原因。
原因 1:伪负特征值(Spurious Negative Eigenvalues)
根据定义,
Proposition 2(负特征值指数增长):设
Proof:回顾更新规则
代入
因此,若
原因 2:特征向量漂移(Eigenvector Drift)
在精确算术下,所有中间矩阵的特征向量与
3.2. 重启策略(Restarting)
为了缓解 Naive Gram Newton-Schulz 在低精度下的数值稳定性问题,作者提出了重启策略。重启的核心思想在于不直接计算
每次重启时,由于重新计算了 Gram 矩阵
作者提出,重启的最佳时机取决于所使用的多项式系数序列。对于 Muon 来说,作者使用 Polar Express 的五次多项式:
| 1 | 8.123737 | -22.232240 | 16.373715 |
| 2 | 4.026529 | -2.776323 | 0.514551 |
| 3 | 3.870284 | -2.739120 | 0.520999 |
| 4 | 3.253351 | -2.343223 | 0.481420 |
| 5 | 2.300652 | -1.668904 | 0.418807 |
通过数值模拟发现,在第 2 次迭代后重启可以最好地平衡稳定性和速度:
的最小特征值保持在 以上 的条件数保持在 以下
3.3. 进一步的安全措施
使用 float16 而非 bfloat16
float16 的范围(约
引入安全因子
大多数 Newton-Schulz 多项式设计为在
这确保即使奇异值达到 1.05 也能收敛。
避免显式添加单位矩阵
计算矩阵二次型
这样所有涉及
3.4. Stabilized Gram Newton-Schulz
综合以上分析,作者提出了完整的 Stabilized Gram Newton-Schulz 算法:
算法 2(Stabilized Gram Newton-Schulz)
输入:
(转换为半精度)- 第一次迭代:
- 对于
: (重启)
- 第二次迭代:
- 对于
:
- 返回
下面来计算 Stabilized Gram Newton-Schulz 的计算复杂度(带一次重启)。
一次重启需要额外两次矩阵乘法:
: FLOPs : FLOPs
同时可以省略三次乘法:
:节省 FLOPs :节省 FLOPs
因此,带一次重启的 Stabilized Gram Newton-Schulz 的计算复杂度为:
对于
- 相比使用对称 GEMM 的标准 Newton-Schulz,减少 42% FLOPs
- 相比典型实现(无对称 GEMM),减少 58% FLOPs
此外,如果使用
四、对称 GEMM 核优化
为了充分利用 Gram Newton-Schulz 带来的对称矩阵乘法机会,作者实现了专门的对称 GEMM 核。
- 三角调度器 (Triangular Scheduler):作者仅将下三角部分的工作块分配给 warp group,上三角部分的值通过对称性获得(
)。 - 融合二次型核:对称 GEMM 核可以融合矩阵二次型的计算,通过在寄存器级别添加
到所有对角线元素,完全避免了加载单位矩阵的 I/O 操作。 (15)
五、实验结果
5.1. 数值稳定性验证
在 Llama-430M 模型上的实验表明:
- Naive Gram Newton-Schulz:出现损失尖峰,最终输出充满 Inf
- Stabilized Gram Newton-Schulz:训练稳定,验证困惑度差异在 0.01 以内
5.2. 加速效果
在 Kimi K2 上:
- Newton-Schulz 正交化步骤运行时间减少 40-50%
- 端到端训练时间减少 2-17%(取决于训练配置)
加速比随矩阵纵横比
(方阵):无加速(退化为标准 Newton-Schulz) (典型 MLP):42% FLOP 减少 (细粒度 MoE):55% FLOP 减少
5.3. 对称 GEMM 核性能
在 Hopper (H100) 和 Blackwell (B200) GPU 架构上 benchmark:
- 相比 cuBLAS,对称 GEMM 核实现 1.5-2x 加速
- 在 Gram Newton-Schulz 中贡献了额外的 10-15% 加速
- 对于
以上矩阵,核效率超过 80%
Appendix
Apd.1. Proof of Proposition 1
Proposition 1(奇异向量不变性):设
其中
Proof:我们使用数学归纳法。
当
假设当
则当
其中
因此结论对
Apd.2. Proof of Theorem 1
Theorem 1(迭代求逆平方根):设任取
则有
其中
迭代的起始条件为
Proof:任取对
当
假设结论对
由假设
注意到
两边平方得:
两边除以
因此结论对