优化器 (3)(Gram Newton-Schulz)

📄 原文:Zifeng Mai · 阅读原文 · 代码

引言

在上一篇文章中,我们深入介绍了 Muon 优化器的数学原理。Muon 的核心创新在于使用 Newton-Schulz 迭代来近似矩阵的 msign 函数,也称为极分解(polar decomposition),从而实现对参数矩阵的正交化更新。

然而,标准的 Newton-Schulz 算法存在一个显著的计算瓶颈:它需要在每次迭代中对原始的矩形参数矩阵进行多次矩阵乘法运算。对于一个 的权重矩阵,每次 Newton-Schulz 迭代需要 的时间复杂度,这在大规模训练中成为了不可忽视的开销。

近期,DAO Lab 提出了 Gram Newton-Schulz 算法,通过以下创新显著加速了 Muon 优化器:

  1. 数学等价性:在 Gram 矩阵 上迭代而非原矩阵 ,输出与标准 Newton-Schulz 数学等价
  2. 数值稳定性:通过重启策略和 float16 算术,实现与标准方法相当的稳定性
  3. 硬件感知:对称 GEMM 核充分利用 Hopper/Blackwell 架构特性

在 Kimi K2 这种万亿参数量级别的模型上,Gram Newton-Schulz 可以将优化器时间减少高达 50%。

一、Muon 与 Newton-Schulz 回顾

1.1. Muon 优化器更新规则

在第 步训练时,令 为权重矩阵, 是 的梯度。Muon 的更新规则为:

(1)

其中 是动量系数, 是学习率。 是动量矩阵,。

Definition 1(极分解):若 是矩阵 的 SVD 分解,则 。

由于精确计算 需要进行完整的 SVD 分解,计算开销较大,因此 Muon 使用 Newton-Schulz 迭代来近似它。

1.2. 标准 Newton-Schulz 迭代

Newton-Schulz 是一种基于矩阵多项式的迭代方法。从初始矩阵 开始,每次迭代按照以下规则更新近似值 :

(2)

其中 是预先设定的多项式系数。

下面分析标准 Newton-Schulz 迭代的计算复杂度。

令 表示迭代次数(在 Muon 中 ),并假设 ,定义纵横比 。

每次迭代包含三次矩阵乘法:

  • : FLOPs
  • (其中 ): FLOPs
  • (其中 ): FLOPs

总计算量为:

(3)

当 时,总计算量为 ,分布在 15 次 GEMM 操作上。

因此,标准的 Newton-Schulz 迭代存在以下问题:

  1. 未利用对称性: 和 都是对称矩阵,但标准实现未利用这一结构。
  2. 强依赖 :当 时,计算量主要由昂贵的 GEMM 操作主导。

二、Gram Newton-Schulz 的数学原理

2.1. 核心思想

Gram Newton-Schulz 的核心洞察在于下面的公式 (4)(证明见上一篇博客)。

(4)

基于公式 (4),Gram Newton-Schulz 的策略是:

  1. 计算 Gram 矩阵
  2. 使用迭代方法近似
  3. 计算

Gram Newton-Schulz 的关键优势在于,第 2 步(占据绝大部分计算量)仅在小的 对称矩阵上操作,仅需两次矩形矩阵乘法(初始的 和最终的 )。

2.2. 从 Newton-Schulz 到 Gram Newton-Schulz

Proposition 1(奇异向量不变性):设 是 的 SVD 分解。若 按照方程 (2) 迭代,则对任意 ,存在奇次多项式 使得:

(5)

其中 , 为 的秩。

根据 Proposition 1,如果多项式序列 对所有奇异值成立,则 。

那么,如何将迭代多项式方法 转换为近似 的迭代方法?

Lemma 1(奇次多项式的性质) 每个奇次多项式 可以重写为 的形式,其中 是低一次的多项式。

因此,如果 ,则 ,因此 Newton-Schulz 多项式隐式地提供了近似逆平方根的方法。下面的定理展示了二者的对应关系。

Theorem 1(迭代求逆平方根):设任取 都有

(6)

则有

(7)

其中 由以下迭代定义:

(8)

迭代的起始条件为 ,。

注意到当 时,有 。因此,我们就可以通过迭代来近似 。

2.3. Naive Gram Newton-Schulz 算法

将定理 1 的迭代提升到矩阵形式,作者提出了 Naive Gram Newton-Schulz 算法:

算法 1(Naive Gram Newton-Schulz)

输入:(),系数

  1. (将奇异值归一化到 ,)

  2. 对于 按照下面的公式迭代

    (9)
  3. 返回

2.4. 计算复杂度分析

下面分析 Naive Gram Newton-Schulz 的 FLOPs 数量级。每次迭代包含四次矩阵乘法(使用对称 GEMM 核):

  • : FLOPs
  • : FLOPs
  • : FLOPs

初始化和输出步骤:

  • : FLOPs
  • : FLOPs(非对称)

下面的计算量可以优化:

  • 不需要计算:节省 FLOPs
  • 不需要计算:节省 FLOPs

因此 Naive Gram Newton-Schulz 的总 FLOPs 为:

(10)

当 时,总计算量为 ,分布在 19 次 GEMM 操作上。

使用对称 GEMM 的标准 Newton-Schulz 需要 FLOPs。当 (非方阵)时,Gram 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)

根据定义, 应该是半正定矩阵(因为 )。然而,在 bfloat16 下, 会引入微小的负特征值(如 ),这是由于浮点误差导致的。下面的命题指出,只要出现了负特征值,该特征值的幅度就会随着迭代而指数增长。

Proposition 2(负特征值指数增长):设 为 Gram Newton-Schulz 迭代的初始 Gram 矩阵。若 存在负特征值 ,则该特征值的幅度随迭代步数指数增长。

Proof:回顾更新规则 。

代入 ,当 时:

(11)

因此,若 ,伪特征值的幅度指数增长,引发链式反应:当 时,,导致 和 。

原因 2:特征向量漂移(Eigenvector Drift)

在精确算术下,所有中间矩阵的特征向量与 的左奇异向量 匹配,但在有限精度下会发生漂移。这导致 和 的特征值偏离理论值。

3.2. 重启策略(Restarting)

为了缓解 Naive Gram Newton-Schulz 在低精度下的数值稳定性问题,作者提出了重启策略。重启的核心思想在于不直接计算 ,而是运行少量迭代(如 5 步)得到 ,然后在 上再次应用 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 的范围(约 到 )足够使用,且精度更高。在某些测试矩阵上,float16 给出更准确的 近似。

引入安全因子

大多数 Newton-Schulz 多项式设计为在 时收敛。由于数值误差可能导致奇异值略大于 1,因此作者建议引入安全因子:

(12)

这确保即使奇异值达到 1.05 也能收敛。

避免显式添加单位矩阵

计算矩阵二次型 时,显式添加 可能降低稳定性。更准确的方式是:

(13)

这样所有涉及 的算术都在 float32 中进行,精度更高。

3.4. Stabilized Gram Newton-Schulz

综合以上分析,作者提出了完整的 Stabilized Gram Newton-Schulz 算法:

算法 2(Stabilized Gram Newton-Schulz)

输入:(),系数

  1. (转换为半精度)
  2. 第一次迭代:
    • 对于 :
    • (重启)
  3. 第二次迭代:
    • 对于 :
  4. 返回

下面来计算 Stabilized Gram Newton-Schulz 的计算复杂度(带一次重启)。

一次重启需要额外两次矩阵乘法:

  • : FLOPs
  • : FLOPs

同时可以省略三次乘法:

  • :节省 FLOPs
  • :节省 FLOPs

因此,带一次重启的 Stabilized Gram Newton-Schulz 的计算复杂度为:

(14)

对于 :

  • 相比使用对称 GEMM 的标准 Newton-Schulz,减少 42% FLOPs
  • 相比典型实现(无对称 GEMM),减少 58% FLOPs

此外,如果使用 次重启,Gram Newton-Schulz 就完全退化为标准 Newton-Schulz。因此,添加重启可以视为在运行时间和稳定性之间权衡。

四、对称 GEMM 核优化

为了充分利用 Gram Newton-Schulz 带来的对称矩阵乘法机会,作者实现了专门的对称 GEMM 核。

  1. 三角调度器 (Triangular Scheduler):作者仅将下三角部分的工作块分配给 warp group,上三角部分的值通过对称性获得()。
  2. 融合二次型核:对称 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(奇异向量不变性):设 是 的 SVD 分解。若 按照方程 (2) 迭代,则对任意 ,存在奇次多项式 使得:

(16)

其中 , 为 的秩。

Proof:我们使用数学归纳法。

当 时,结论显然成立。

假设当 时,,其中 是某个奇次多项式。

则当 时:

(17)

其中 。

因此结论对 也成立。由数学归纳法,命题得证。

Apd.2. Proof of Theorem 1

Theorem 1(迭代求逆平方根):设任取 都有

(18)

则有

(19)

其中 由以下迭代定义:

(20)

迭代的起始条件为 ,。

Proof:任取对 ,定义 且 。我们将通过归纳法证明 且 对所有 成立。

当 时,由定义 ,,结论成立。

假设结论对 成立,即 且 。

由假设 ,我们有:

(21)

注意到 ,因此 。

两边平方得:

(22)

两边除以 得:

(23)

因此结论对 也成立。由数学归纳法,定理得证。