流形上的最速下降:6. Muon + 双旋转
TL;DR · AI 摘要
Title: 流形上的最速下降:6. Muon + 双旋转 - 科学空间 Scientific Spaces URL Source: Markdown Content: 我们知道,用Adam、Muon等优化器更新矩阵参数时,奇异值和左右奇异...
核心要点
- 主题聚焦:流形上的最速下降:6. Muon + 双旋转
- 来源:科学空间,建议结合原文判断细节。
- AI 分析暂不可用,本条为保底评分与摘要。
我们知道,用Adam、Muon等优化器更新矩阵参数时,奇异值和左右奇异向量都会随之变化,它们通常都是耦合在一起。也正是因为这种耦合性,我们无法简单地调控矩阵参数的奇异值,因此在奇异值出现异常增长时,我们无法简单有效地阻止它,这可能导致训练失败。
本文在《Pion: A Spectrum-Preserving Optimizer via Orthogonal Equivalence Transformation》(后面简称Pion)的启发下,提出了一种单独更新矩阵左右奇异向量的Muon变体——“旋转Muon(MuonR)”,它能维持矩阵的奇异值分布不变,从而保证训练稳定性。
前文回顾[#](https://spaces.ac.cn/archives/11777#%E5%89%8D%E6%96%87%E5%9B%9E%E9%A1%BE)
由于左右奇异向量组成的矩阵一定是正交矩阵,所以我们先来简单回顾一下正交约束下的Muon。设参数$\mathbf{\mathit{W}} \in \mathbb{R}^{n \times n}$,满足$\mathbf{\mathit{W}}^{\top} \mathbf{\mathit{W}} = \mathbf{\mathit{I}}$,设更新量为$\Delta \mathbf{\mathit{W}} = - \eta \mathbf{\Phi}$,我们希望参数更新之后依然满足正交性,那么对应的谱范数最速下降问题就是
$$ (\text{1}) \underset{\mathbf{\Phi}}{max} \text{tr} \left(\right. \mathbf{\mathit{G}}^{\top} \mathbf{\Phi} \left.\right) \text{s}.\text{t}. \parallel \mathbf{\Phi} \parallel_{2} = 1 , \left(\right. \mathbf{\mathit{W}} - \eta \mathbf{\Phi} \left.\right)^{\top} \left(\right. \mathbf{\mathit{W}} - \eta \mathbf{\Phi} \left.\right) = \mathbf{\mathit{I}} $$
可以解得$\mathbf{\Phi} = \mathbf{\mathit{W}} \mathbf{\mathit{O}}$,其中$\mathbf{\mathit{O}} = \text{msign} \left(\right. \left[\right. \mathbf{\mathit{W}}^{\top} \mathbf{\mathit{G}} \left]\right._{\text{skew}} \left.\right)$,这里$\left[\right. \mathbf{\mathit{X}} \left]\right._{\text{skew}} = \left(\right. \mathbf{\mathit{X}} - \mathbf{\mathit{X}}^{\top} \left.\right) / 2$是反对称化算符。考虑上缩回操作,完整的更新规则是:
$$ (\text{2}) \mathbf{\mathit{W}} \leftarrow \mathbf{\mathit{W}} \left(\right. \mathbf{\mathit{I}} - \eta \mathbf{\mathit{O}} \left.\right) \left(\right. \mathbf{\mathit{I}} - \mathbf{\mathit{O}}^{\top} \mathbf{\mathit{O}} + \frac{\mathbf{\mathit{O}}^{\top} \mathbf{\mathit{O}}}{\sqrt{1 + \eta^{2}}} \left.\right) $$
特别地,如果$\left[\right. \mathbf{\mathit{W}}^{\top} \mathbf{\mathit{G}} \left]\right._{\text{skew}}$是满秩的,那么将简化成
$$ (\text{3}) \mathbf{\mathit{W}} \leftarrow \frac{\mathbf{\mathit{W}} \left(\right. \mathbf{\mathit{I}} - \eta \mathbf{\mathit{O}} \left.\right)}{\sqrt{1 + \eta^{2}}} $$
推导过程在《流形上的最速下降:2. Muon + 正交》可以找到,这里就不详细展开了。不管是式$(\text{2})$还是$(\text{3})$,它们都是完全解析的,仅在Muon的基础上增加了几步矩阵乘法,复杂度没有明显增加,所以这个结果是完全实用的。 现在考虑矩阵$\mathbf{\mathit{W}} \in \mathbb{R}^{n \times m} \left(\right. n \geq m \left.\right)$,如果同时满足$\mathbf{\mathit{W}}^{\top} \mathbf{\mathit{W}} = \mathbf{\mathit{I}}$,那么我们称$\mathbf{\mathit{W}}$位于Stiefel流形上,它是正交矩阵概念的一般化,上述结果理论上可以推广到Stiefel流形上,但非方阵下需要求解一个非线性方程组,比较难实用化,个中细节可以参考《流形上的最速下降:3. Muon + Stiefel》。
瞬时重参[#](https://spaces.ac.cn/archives/11777#%E7%9E%AC%E6%97%B6%E9%87%8D%E5%8F%82)
接着,我们将注意力转移到任意参数矩阵$\mathbf{\mathit{W}} \in \mathbb{R}^{n \times m}$上,目标是使参数更新过程中奇异值保持不变,从而杜绝奇异值异常增长的可能。
为了实现这个目标,我们采用“瞬时重参”的思想:在更新开始前,先将$\mathbf{\mathit{W}}$矩阵重新参数化为$\overset{\sim}{\mathbf{\mathit{W}}} = \mathbf{\mathit{L}} \mathbf{\mathit{W}} \mathbf{\mathit{R}}$,其中$\mathbf{\mathit{L}} \in \mathbb{R}^{n \times n} , \mathbf{\mathit{R}} \in \mathbb{R}^{m \times m}$,并且都初始化为单位阵。这样一来,初始化时满足$\overset{\sim}{\mathbf{\mathit{W}}} = \mathbf{\mathit{W}}$,并且通过记$\mathbf{\mathit{G}} = \nabla_{\mathbf{\mathit{W}}} \mathcal{L}$,我们可以写出
$$ (\text{4}) \nabla_{\mathbf{\mathit{L}}} \mathcal{L} = \mathbf{\mathit{G}} \mathbf{\mathit{W}}^{\top} , \nabla_{\mathbf{\mathit{R}}} \mathcal{L} = \mathbf{\mathit{W}}^{\top} \mathbf{\mathit{G}} $$
随后,我们约定冻结$\mathbf{\mathit{W}}$,只更新$\mathbf{\mathit{L}}$和$\mathbf{\mathit{R}}$,同时在更新过程中保持$\mathbf{\mathit{L}} , \mathbf{\mathit{R}}$的正交性,这样更新后的$\overset{\sim}{\mathbf{\mathit{W}}}$依然具有跟$\mathbf{\mathit{W}}$一样的奇异值。现在从$\mathbf{\mathit{L}} , \mathbf{\mathit{R}}$的视角看,问题再次变成了正交流形上的最速下降,并且这时候$\mathbf{\mathit{L}} , \mathbf{\mathit{R}}$都是方阵,对应的最速下降完全可解析解!根据式$(\text{3})$可以直接写出更新规则
$$ (\text{5}) \mathbf{\mathit{L}} \leftarrow \left(\right. \mathbf{\mathit{I}} - \eta \mathbf{\mathit{O}}_{L} \left.\right) \left(\right. \mathbf{\mathit{I}} - \mathbf{\mathit{O}}_{L}^{\top} \mathbf{\mathit{O}}_{L} + \frac{\mathbf{\mathit{O}}_{L}^{\top} \mathbf{\mathit{O}}_{L}}{\sqrt{1 + \eta^{2}}} \left.\right) \\ (\text{6}) \mathbf{\mathit{R}} \leftarrow \left(\right. \mathbf{\mathit{I}} - \eta \mathbf{\mathit{O}}_{R} \left.\right) \left(\right. \mathbf{\mathit{I}} - \mathbf{\mathit{O}}_{R}^{\top} \mathbf{\mathit{O}}_{R} + \frac{\mathbf{\mathit{O}}_{R}^{\top} \mathbf{\mathit{O}}_{R}}{\sqrt{1 + \eta^{2}}} \left.\right) $$
其中$\mathbf{\mathit{O}}_{L} = \text{msign} \left(\right. \left[\right. \mathbf{\mathit{G}} \mathbf{\mathit{W}}^{\top} \left]\right._{\text{skew}} \left.\right) , \mathbf{\mathit{O}}_{R} = \text{msign} \left(\right. \left[\right. \mathbf{\mathit{W}}^{\top} \mathbf{\mathit{G}} \left]\right._{\text{skew}} \left.\right)$。将新的$\mathbf{\mathit{L}} , \mathbf{\mathit{R}}$与$\mathbf{\mathit{W}}$乘起来,那么就得到完整的更新规则
$$ (\text{7}) \mathbf{\mathit{W}} \leftarrow \mathbf{\mathit{L}} \mathbf{\mathit{W}} \mathbf{\mathit{R}} $$
这便是基于“瞬时重参”思想推导出来“旋转Muon(Muon under Rotation,MuonR)”。实际场景下通常还有动量,我们将它理解为平滑过后的梯度,所以只需要将梯度$\mathbf{\mathit{G}}$换成动量$\mathbf{\mathit{M}}$。
一些细节[#](https://spaces.ac.cn/archives/11777#%E4%B8%80%E4%BA%9B%E7%BB%86%E8%8A%82)
由于$\mathbf{\mathit{L}} , \mathbf{\mathit{R}}$的更新都需要算一次$\text{msign}$,所以即便在最理想($n = m$)的情况下,MuonR计算量也是Muon的两倍。不过对于足够大的模型来说,这个计算量翻倍对端到端的训练时间影响是比较微弱的,通常还可以接受。如果想要减少这部分开销,可以考虑$\mathbf{\mathit{L}} , \mathbf{\mathit{R}}$交替更新的做法,平摊一下计算量。
事实上,MuonR最大的问题在于,它自始至终保持矩阵的全体奇异值不变,这意味着我们必须在初始化时就把参数的全体奇异值确定下来。这并不容易,因为不同位置的矩阵可能需要不同的尺度,强行设置为同一组值大概率是次优的。
一个可行的思路是,在适当的随机初始化基础上,给每个矩阵前或后添加一个Element-wise相乘的向量,用来补偿尺度的自由度。对于紧接在RMSNorm后面的矩阵,RMSNorm自带的gamma参数已经承担了这一角色,所以这部分矩阵则可以省略这一操作。
至于矩阵的初始奇异值如何选取,我们可以考虑常规的随机初始化,也可以考虑按Zipf's law的形式构建。进一步地,我们还可以尝试将奇异值熵调节成《矩阵参数的奇异值熵越高越好吗?》计算出来的最优熵,以期达到更优的效果。
当然,如果我们确实能事先确定矩阵的奇异值——比如本就期望某处参数始终保持正交性——那就不用考虑这些了,直接套用MuonR就行。
中途切换[#](https://spaces.ac.cn/archives/11777#%E4%B8%AD%E9%80%94%E5%88%87%E6%8D%A2)
还有一种可选的做法是“中途切换”,仅将MuonR作为“维稳”的手段。
具体来说,一开始我们还是用常规的Muon,并监测矩阵的谱范数/F范数,一旦矩阵的范数超出我们期望的范围,那就切换到MuonR,由于两种Muon都只依赖于同一个梯度/动量,它们只是计算不同,所以这种切换是允许的。MuonR不改变矩阵的奇异值,那么谱范数/F范数都不再增长,正好可以作为“维稳”手段使用。
但要注意的是,我们要让切换前后的更新幅度要尽可能对齐,避免引入“突变”。为此,我们考虑MuonR的一阶近似
$$ (\text{8}) \mathbf{\mathit{L}} \mathbf{\mathit{W}} \mathbf{\mathit{R}} \approx \left(\right. \mathbf{\mathit{I}} - \eta \mathbf{\mathit{O}}_{L} \left.\right) \mathbf{\mathit{W}} \left(\right. \mathbf{\mathit{I}} - \eta \mathbf{\mathit{O}}_{R} \left.\right) \approx \mathbf{\mathit{W}} - \eta \left(\right. \mathbf{\mathit{O}}_{L} \mathbf{\mathit{W}} + \mathbf{\mathit{W}} \mathbf{\mathit{O}}_{R} \left.\right) $$
由于$\mathbf{\mathit{O}}_{L} , \mathbf{\mathit{O}}_{R}$的奇异值都不超过1(注意我们无法保证$\left[\right. \mathbf{\mathit{G}} \mathbf{\mathit{W}}^{\top} \left]\right._{\text{skew}}$和$\left[\right. \mathbf{\mathit{W}}^{\top} \mathbf{\mathit{G}} \left]\right._{\text{skew}}$都满秩,所以不能直接使用$\mathbf{\mathit{O}}_{L} , \mathbf{\mathit{O}}_{R}$的正交性),,所以有$\parallel \mathbf{\mathit{O}}_{L} \mathbf{\mathit{W}} \parallel_{F} \leq \parallel \mathbf{\mathit{W}} \parallel_{F}$和$\parallel \mathbf{\mathit{W}} \mathbf{\mathit{O}}_{R} \parallel_{F} \leq \parallel \mathbf{\mathit{W}} \parallel_{F}$,于是
$$ (\text{9}) \parallel \mathbf{\mathit{O}}_{L} \mathbf{\mathit{W}} + \mathbf{\mathit{W}} \mathbf{\mathit{O}}_{R} \parallel_{F} \leq \parallel \mathbf{\mathit{O}}_{L} \mathbf{\mathit{W}} \parallel_{F} + \parallel \mathbf{\mathit{W}} \mathbf{\mathit{O}}_{R} \parallel_{F} \leq 2 \parallel \mathbf{\mathit{W}} \parallel_{F} $$
常规Muon是$\mathbf{\mathit{W}} - \eta \text{msign} \left(\right. \mathbf{\mathit{G}} \left.\right)$,$\text{msign} \left(\right. \mathbf{\mathit{G}} \left.\right)$的F范数普遍为$\sqrt{min \left(\right. n , m \left.\right)}$,因此为了对齐更新量F范数,从Muon切换到MuonR,大致上要将学习率乘以$\frac{\sqrt{min \left(\right. n , m \left.\right)}}{2 \parallel \mathbf{\mathit{W}} \parallel_{F}}$。 实践中,上式第一个不等号可能不够紧,$\mathbf{\mathit{O}}_{L} \mathbf{\mathit{W}}$和$\mathbf{\mathit{W}} \mathbf{\mathit{O}}_{R}$更多是近乎正交的关系,根据勾股定理,结果应该约等于$\sqrt{2} \parallel \mathbf{\mathit{W}} \parallel_{F}$,那么这个倍数应当多乘以$\sqrt{2}$。不过考虑到$\sqrt{2}$与$1$差别不是特别大,并且为了保证极端情况下的可用性,建议还是保留上述形式。
异同分析[#](https://spaces.ac.cn/archives/11777#%E5%BC%82%E5%90%8C%E5%88%86%E6%9E%90)
本文在开头就明确表示,MuonR受启发自Pion,我们先来看看它们的联系与区别。
首先,将更新规则限定为左乘和右乘正交矩阵这种双旋转形式的思路主要来源于Pion;在确定这一更新形式后,通过“瞬时重参”获取对应的梯度这一步是比较自然的;接着,Pion和MuonR就开始“分道扬镳”:
1、Pion通过矩阵指数$exp \left(\right. 反对称矩阵 \left.\right)$来实现正交性,实际计算中则通过展开到二阶来近似;
2、Pion走的是Adam路线,并且将$\mathbf{\mathit{L}} , \mathbf{\mathit{R}}$的梯度分别滑动平均,这使得它的缓存变量达到了4组;
3、MuonR走的是Muon路线,跟Muon一样只缓存动量,这使得我们可以随时跟Muon切换;
4、MuonR基于正交流形下最速下降的解析解,只需有限额外步骤即可精确地实现正交性。
总的来说,Pion的正交性设计比较经验化,且四组缓存变量让人有点“望而却步”;MuonR则是Muon、正交流形最速下降等一系列工作下相对自然的产物,笔者认为它整体上要更符合第一性原理一些。
文章小结[#](https://spaces.ac.cn/archives/11777#%E6%96%87%E7%AB%A0%E5%B0%8F%E7%BB%93)
本文提出了MuonR,一种将更新形式约束为左右旋转矩阵的Muon变体,它能够保持矩阵的奇异值分布不变,是一种简洁的维持训练稳定性的训练方案。
_转载到请包括本文地址:[https://spaces.ac.cn/archives/11777](https://spaces.ac.cn/archives/11777 "流形上的最速下降:6. Muon + 双旋转")_
_更详细的转载事宜请参考:_[《科学空间FAQ》](https://spaces.ac.cn/archives/6508#%E6%96%87%E7%AB%A0%E5%A6%82%E4%BD%95%E8%BD%AC%E8%BD%BD/%E5%BC%95%E7%94%A8 "《科学空间FAQ》")