In the previous two articles, Rethinking Learning Rate and Batch Size (Part 1): The Current State and Rethinking Learning Rate and Batch Size (Part 2): Mean Field, we mainly proposed the mean-field method to simplify calculations related to the learning rate and batch size. At that time, the optimizers we analyzed were SGD, SignSGD, and SoftSignSGD, and the main purpose was simplification; in essence, there were no new conclusions.
However, in today’s feast of optimizers, how could Muon be left without a seat at the table? So, in this article, we attempt to compute the relevant conclusions for Muon and see whether its relationship between the learning rate and batch size exhibits new patterns.
Basic Notation
As is well known, the main characteristic of Muon is its non-element-wise update rule, so the element-wise computation methods from “When the Batch Size Increases, How Should the Learning Rate Change Accordingly?” and “How Does Adam’s Epsilon Affect the Scaling Law of the Learning Rate?” become completely inapplicable. Fortunately, the mean field introduced in the previous article still works; we only need to adjust a few details.
Let us first introduce some notation. Let the loss function be \(\mathcal{L}(\boldsymbol{W})\), where \(\boldsymbol{W}\in\mathbb{R}^{n\times m}\) is a matrix (assume \(n\geq m\)), \(\boldsymbol{G}\) is its gradient, the gradient of a single sample is denoted \(\tilde{\boldsymbol{G}}\), whose mean is \(\boldsymbol{G}\) and whose variance is \(\sigma^2\); when the batch size is \(B\), the gradient is denoted \(\tilde{\boldsymbol{G}}_B\), whose mean is still \(\boldsymbol{G}\) but whose variance becomes \(\sigma^2/B\). Note that here the variance is just a scalar \(\sigma^2\), unlike before, where we considered the full covariance matrix.
The core reason for this simplification is that the random variable here is itself already a matrix, so its corresponding covariance matrix would actually be a 4th-order tensor, which is rather troublesome to discuss. Would simplifying it to a single scalar severely sacrifice accuracy? Actually, no. Although the previous two articles considered the full covariance matrix \(\boldsymbol{\Sigma}\), a careful look reveals that the final result only depends on \(\mathop{\mathrm{tr}}(\boldsymbol{\Sigma})\), which is equivalent to simplifying it to a scalar from the very beginning.
The Hessian
Similarly, let the update be \(-\eta\tilde{\boldsymbol{\Phi}}_B\), and consider the second-order expansion of the loss function \[\begin{equation}\mathcal{L}(\boldsymbol{W} - \eta\tilde{\boldsymbol{\Phi}}_B) \approx \mathcal{L}(\boldsymbol{W}) - \eta \mathop{\mathrm{tr}}(\tilde{\boldsymbol{\Phi}}{}_B^{\top}\boldsymbol{G}) + \frac{1}{2}\eta^2\mathop{\mathrm{tr}}(\tilde{\boldsymbol{\Phi}}{}_B^{\top}\boldsymbol{H}\tilde{\boldsymbol{\Phi}}_B)\label{eq:loss-2}\end{equation}\] There should be no questions about the first two terms; the third term is the harder one to understand. Similar to the covariance matrix, the Hessian matrix \(\boldsymbol{H}\) here is a 4th-order tensor, which is rather hard to understand.
The simplest way to approach this should be the linear operator perspective, i.e., understanding \(\boldsymbol{H}\) as a linear operator whose input and output are both matrices. We don’t need to know what \(\boldsymbol{H}\) looks like, nor how \(\boldsymbol{H}\) operates on \(\tilde{\boldsymbol{\Phi}}_B\); we only need to know that \(\boldsymbol{H}\tilde{\boldsymbol{\Phi}}_B\) is linear in \(\tilde{\boldsymbol{\Phi}}_B\). In this way, the objects we handle are still matrices, adding no mental burden. Any conforming linear operator can serve as an approximation of the Hessian matrix, without needing to write out the concrete higher-order tensor form.
The protagonist of this article is Muon, and we take \(\tilde{\boldsymbol{\Phi}}_B=\mathop{\mathrm{msign}}(\tilde{\boldsymbol{G}}_B)\) as its approximation for computation. By definition, we write \(\mathop{\mathrm{msign}}(\tilde{\boldsymbol{G}}_B)=\tilde{\boldsymbol{G}}_B(\tilde{\boldsymbol{G}}{}_B^{\top}\tilde{\boldsymbol{G}}_B)^{-1/2}\); from the Newton’s method perspective, this amounts to assuming \(\boldsymbol{H}^{-1}\boldsymbol{X} = \eta_{\max}\boldsymbol{X}(\boldsymbol{G}^{\top}\boldsymbol{G})^{-1/2}\), hence \(\boldsymbol{H}\boldsymbol{X} = \eta_{\max}^{-1}\boldsymbol{X}(\boldsymbol{G}^{\top}\boldsymbol{G})^{1/2}\), which will be used in the computation below.
Computing the Expectation
Taking the expectation of both sides of Eq. \(\eqref{eq:loss-2}\), we get \[\begin{equation}\mathbb{E}[\mathcal{L}(\boldsymbol{W} - \eta\tilde{\boldsymbol{\Phi}}_B)] \approx \mathcal{L}(\boldsymbol{W}) - \eta \mathop{\mathrm{tr}}(\mathbb{E}[\tilde{\boldsymbol{\Phi}}_B]^{\top}\boldsymbol{G}) + \frac{1}{2}\eta^2\mathbb{E}[\mathop{\mathrm{tr}}(\tilde{\boldsymbol{\Phi}}{}_B^{\top}\boldsymbol{H}\tilde{\boldsymbol{\Phi}}_B)]\end{equation}\] First compute \(\mathbb{E}[\tilde{\boldsymbol{\Phi}}_B]\): \[\begin{equation}\mathbb{E}[\tilde{\boldsymbol{\Phi}}_B]=\mathbb{E}[\tilde{\boldsymbol{G}}_B(\tilde{\boldsymbol{G}}{}_B^{\top}\tilde{\boldsymbol{G}}_B)^{-1/2}]\approx\mathbb{E}[\tilde{\boldsymbol{G}}_B](\mathbb{E}[\tilde{\boldsymbol{G}}{}_B^{\top}\tilde{\boldsymbol{G}}_B])^{-1/2} = \boldsymbol{G}(\mathbb{E}[\tilde{\boldsymbol{G}}{}_B^{\top}\tilde{\boldsymbol{G}}_B])^{-1/2}\end{equation}\] We write \(\mathbb{E}[\tilde{\boldsymbol{G}}{}_B^{\top}\tilde{\boldsymbol{G}}_B]\) out component by component, and assume independence between different components; then \[\begin{equation}\mathbb{E}[\tilde{\boldsymbol{G}}{}_B^{\top}\tilde{\boldsymbol{G}}_B]_{i,j} = \mathbb{E}\left[\sum_{k=1}^n (\tilde{G}_B)_{k,i}(\tilde{G}_B)_{k,j}\right] = \left\{\begin{aligned} \mathbb{E}\left[\sum_{k=1}^n (\tilde{G}_B)_{k,i}^2\right] = \left(\sum_{k=1}^n G_{k,i}^2\right) + n\sigma^2/B,\quad (i=j) \\[6pt] \sum_{k=1}^n \mathbb{E}[(\tilde{G}_B)_{k,i}] \mathbb{E}[(\tilde{G}_B)_{k,j}] = \sum_{k=1}^n G_{k,i}G_{k,j},\quad (i\neq j) \end{aligned}\right.\end{equation}\] Combining these gives \(\mathbb{E}[\tilde{\boldsymbol{G}}{}_B^{\top}\tilde{\boldsymbol{G}}_B]=\boldsymbol{G}^{\top}\boldsymbol{G} + (n\sigma^2/B) \boldsymbol{I}\), so \[\begin{equation}\mathbb{E}[\tilde{\boldsymbol{\Phi}}_B]\approx \boldsymbol{G}(\boldsymbol{G}^{\top}\boldsymbol{G} + (n\sigma^2/B) \boldsymbol{I})^{-1/2} = \mathop{\mathrm{msign}}(\boldsymbol{G})(\boldsymbol{I} + (n\sigma^2/B) (\boldsymbol{G}^{\top}\boldsymbol{G})^{-1})^{-1/2}\end{equation}\] To further simplify the dependence on \(B\), we approximate \(\boldsymbol{G}^{\top}\boldsymbol{G}\) by \(\mathop{\mathrm{tr}}(\boldsymbol{G}^{\top}\boldsymbol{G})\boldsymbol{I}/m\), i.e., we keep only the diagonal part of \(\boldsymbol{G}^{\top}\boldsymbol{G}\), and then replace the diagonal entries by their average. In this way, we obtain \[\begin{equation}\mathbb{E}[\tilde{\boldsymbol{\Phi}}_B]\approx \mathop{\mathrm{msign}}(\boldsymbol{G})(1 + \mathcal{B}_{\text{simple}}/B)^{-1/2}\end{equation}\] where \(\mathcal{B}_{\text{simple}} = mn\sigma^2/\mathop{\mathrm{tr}}(\boldsymbol{G}^{\top}\boldsymbol{G})= mn\sigma^2/\Vert\boldsymbol{G}\Vert_F\), which is actually the same as computing the \(\mathcal{B}_{\text{simple}}\) of the previous two articles by treating \(\boldsymbol{G}\) as a vector. The form of the above equation is identical to that of SignSGD, from which we can conjecture that Muon will not have much new in the relationship between the learning rate and the batch size.
The Same Pattern
As for \(\mathbb{E}[\mathop{\mathrm{tr}}(\tilde{\boldsymbol{\Phi}}{}_B^{\top}\boldsymbol{H}\tilde{\boldsymbol{\Phi}}_B)]\), we compute only under the assumption just derived for Muon, i.e., \(\boldsymbol{H}\boldsymbol{X} = \eta_{\max}^{-1}\boldsymbol{X}(\boldsymbol{G}^{\top}\boldsymbol{G})^{1/2}\); then \[\begin{equation}\mathop{\mathrm{tr}}(\tilde{\boldsymbol{\Phi}}{}_B^{\top}\boldsymbol{H}\tilde{\boldsymbol{\Phi}}_B) = \eta_{\max}^{-1}\mathop{\mathrm{tr}}(\tilde{\boldsymbol{\Phi}}{}_B^{\top}\tilde{\boldsymbol{\Phi}}_B(\boldsymbol{G}^{\top}\boldsymbol{G})^{1/2})\end{equation}\] Note that since \(\tilde{\boldsymbol{\Phi}}_B\) is the result of \(\mathop{\mathrm{msign}}\), it is necessarily an orthogonal matrix (full rank), so \(\tilde{\boldsymbol{\Phi}}{}_B^{\top}\tilde{\boldsymbol{\Phi}}_B=\boldsymbol{I}\), i.e., in this case \(\mathop{\mathrm{tr}}(\tilde{\boldsymbol{\Phi}}{}_B^{\top}\boldsymbol{H}\tilde{\boldsymbol{\Phi}}_B)\) is a deterministic constant \(\eta_{\max}^{-1}\mathop{\mathrm{tr}}((\boldsymbol{G}^{\top}\boldsymbol{G})^{1/2})=\eta_{\max}^{-1}\mathop{\mathrm{msign}}(\boldsymbol{G})^{\top}\boldsymbol{G}\), and thus we can obtain \[\begin{equation}\eta^* \approx \frac{\mathop{\mathrm{tr}}(\mathbb{E}[\tilde{\boldsymbol{\Phi}}_B]^{\top}\boldsymbol{G})}{\mathbb{E}[\mathop{\mathrm{tr}}(\tilde{\boldsymbol{\Phi}}{}_B^{\top}\boldsymbol{H}\tilde{\boldsymbol{\Phi}}_B)]}\approx \frac{\eta_{\max}}{\sqrt{1 + \mathcal{B}_{\text{simple}}/B}}\end{equation}\] Sure enough, its form is exactly the same as the result for SignSGD; there is no new pattern.
In fact, thinking carefully, this is only natural: SignSGD directly applies \(\mathop{\mathrm{sign}}\) to the gradient, while Muon’s \(\mathop{\mathrm{msign}}\) applies \(\mathop{\mathrm{sign}}\) to the singular values; intuitively this amounts to applying \(\mathop{\mathrm{sign}}\) in a changed coordinate system. What it brings is a new matrix update rule, while the learning rate \(\eta^*\) and batch size \(B\) are merely scalars; given that the core behind all of them is \(\mathop{\mathrm{sign}}\), the asymptotic relationship between these scalars is very unlikely to change noticeably.
Of course, we have so far only computed one particular \(\boldsymbol{H}\); if we consider a more general \(\boldsymbol{H}\), then, just as with SignSGD, the Surge phenomenon—“as the batch size increases, the learning rate should instead decrease”—could also appear. But as we said in the “Reflection on the Cause” section of the previous article (Rethinking Learning Rate and Batch Size (Part 2): Mean Field), if the Surge phenomenon is really observed, it may be better to switch optimizers rather than to correct the relationship between \(\eta^*\) and \(B\).
Article Summary
In this article, we attempted a simple analysis of Muon using the mean-field approximation. The conclusion is that its relationship between the learning rate and batch size is consistent with that of SignSGD, with no new pattern.
When reprinting, please include the address of this article: https://kexue.fm/archives/11285
For more detailed reprinting matters, please refer to: Scientific Spaces FAQ
Comments