Can We Remove the Square-Root in Adaptive Gradient Methods? A Second-Order Perspective

Wu Lin, Felix Dangel, Runa Eschenhagen, Juhan Bae, Richard E. Turner, Alireza Makhzani

Introduction

Adaptive gradient-based methods like Adam (Kingma & Ba, 2015) play a significant role in training modern deep learning models such as transformers. A better understanding of these methods allows us to address their shortcomings and develop new adaptive methods to reduce training time and improve performance. This is essential as deep learning models become increasingly large and complex to train.

Despite their success on architectures like transformers, adaptive methods tend to generalize worse than stochastic gradient descent (SGD) on convolutional architectures (Wilson et al., 2017). Our understanding of this discrepancy is limited. Balles & Hennig (2018) dissect Adam into two concepts, sign descent and adaptive step sizes, and hypothesize that the connection to sign descent could cause the generalization gap with SGD on CNNs. Similarly, Kunstner et al. (2023); Chen et al. (2023) argue that adaptive methods outperform SGD on transformers due to their connection to sign descent.

It is challenging to isolate the sign descent component of adaptive methods like Adam or RMSProp (Tieleman & Hinton, 2012), as they typically introduce a square root to the preconditioner, which conflates the sign and adaptivity aspects. The root is motivated by reports of improved performance (Tieleman & Hinton, 2012) and to stabilize convergence near the optimum (Kingma & Ba, 2015; Kunstner et al., 2019; Martens, 2020). However, it conflicts with the motivation of adaptive methods as approximate second-order methods based on the empirical Fisher, which is also commonly mentioned in works introducing adaptive methods (e.g. Kingma & Ba, 2015).

Here, we investigate how the behavior of adaptive methods changes when we remove the root. Our idea is to strengthen their often-mentioned link to second-order methods that is weakened by the root. Conceptually, this cleanly disentangles the aforementioned adaptivity aspect from the sign aspect. Practically, it provides an opportunity to revisit the root’s role in the context of modern training strategies, such as using non-constant learning rate schedules (Loshchilov & Hutter, 2016) and hyperparameter tuning schemes (Choi et al., 2019), that differ from the original pipelines in which the square root was introduced. Computationally, removing the root is beneficial for lowering memory consumption and per-iteration cost of non-diagonal matrix adaptive methods, which require costly matrix decompositions that need to be carried out in high precision to avoid numerical instabilities (Gupta et al., 2018; Anil et al., 2020).

There are some challenges to establishing a rigorous second-order perspective on adaptive methods; one cannot just remove the root. We overcome those challenges and make the following contributions:

We establish a rigorous second-order view of adaptive methods: we remove the square root (Section 2), show how to interpret the gradient outer product as a new empirical Fisher variant (Section 3), and adjust the preconditioner initialization (Section 4).

Empirically, we show that—surprisingly—removing the root not only closes the generalization gap between adaptive methods and SGD on convolutional neural networks, but maintains the performance of root-based methods on vision transformers (Section 2).

We demonstrate the second-order perspective’s conceptual and computational merits to develop matrix adaptive methods that, thanks to recent approaches on inverse-free second-order methods (Lin et al., 2023), require neither matrix inverses nor matrix square roots and can therefore run in low precision—a crucial ingredient for modern training (Section 4).

First-order View of Adaptive Methods

For many deep learning tasks, training a neural network (NN) means solving an unconstrained optimization problem. For simplicity, consider a supervised learning task with a set of NN data points \{y_{i},\mbox{\mbox{x\mathbf{x}}}_{i}\}_{i=1}^{N}, where yiy_{i} and \mbox{\mbox{x\mathbf{x}}}_{i} represent a label and a feature vector. Given a NN f(\cdot;\mbox{\mbox{μ\boldsymbol{\mu}}}) with learnable weights μ\boldsymbol{\mu}, the optimization problem is

where \hat{y}_{i}:=f(\mbox{\mbox{x\mathbf{x}}}_{i};\mbox{\mbox{μ\boldsymbol{\mu}}}) is a predicted label given a feature vector \mbox{\mbox{x\mathbf{x}}}_{i} as an input to the NN and c(yi,y^i)c(y_{i},\hat{y}_{i}) is a loss function that measures a discrepancy between the true (yiy_{i}) and predicted (y^i\hat{y}_{i}) label.

To solve this optimization problem, we can use adaptive gradient methods, which use the following preconditioned gradient update

We have to use a symmetric square-root decomposition \hat{\mbox{\mbox{S\mathbf{S}}}}^{\frac{1}{2}} to compute the preconditioner.

It is unclear how exactly the square root emerged in adaptive methods. Here, we hypothesize how it got introduced, critically assess why it may be desirable to remove, and highlight connections to our experimental results. We find that few works study square-root-free adaptive methods, and the benefits of removing the root are often overlooked in the literature. Our work fills in this gap, which we think is crucial to better understand these methods and our empirical results underline the great potential of square-root-free adaptive methods.

The most important motivation in favor of the square root is the strong empirical performance of square-root-based methods that has been demonstrated in various works and rightfully established them as state-of-the-art for training transformers. However, the context in which the root was introduced significantly differed from the training schemes that are used nowadays. Original works such as Tieleman & Hinton (2012) show that including the square root improves the performance of adaptive methods when only tuning a learning rate that is fixed throughout training (Bottou et al., 2018). Beyond this original setting, square-root-based methods have demonstrated great capabilities to train a wide range of NNs, e.g. when using a non-constant learning rate (Loshchilov & Hutter, 2016) and random search (Bergstra & Bengio, 2012; Choi et al., 2019) to tune all available hyper-parameters. We wonder whether the square root, while necessary to achieve good performance in outdated training schemes, might not be required in contemporary schemes. To investigate this hypothesis, we conducted an experiment that compares square-root-based and square-root-free methods using the original training scheme in which the square root was introduced (Figure 2). Indeed, we find that the square root is beneficial in this context.

Interpretability & generalization

Square-root-based methods are state-of-the-art for training attention-based models, but exhibit a generalization gap with SGD on CNNs (Wilson et al., 2017). The square root complicates the understanding of this phenomenon. Recent studies attribute the superior performance of adaptive methods over SGD on attention models to their sign-based update rules (Kunstner et al., 2023; Chen et al., 2023). Balles & Hennig (2018) hypothesize that sign descent may cause poor generalization of Adam on CNNs. However, they neither consider the direct effect of the square root nor remove it from an existing method. Therefore, the role of adaptivity as another main factor behind the performance gap is unclear since adding the root conflates sign descent and adaptivity. Removing the square root could clarify this issue, as a square-root-free method no longer performs sign descent. To investigate the role of adaptivity, we experiment with square-root-free RMSProp (Figure 4) on attention and convolutional models (Figure 1). For transformer-based architectures, we find that removing the square root does not negatively affect the performance of RMSProp. This suggests that not only sign descent, but also adaptivity, might be a key factor behind the performance gap between adaptive methods and SGD on transformers. For convolutional architectures, removing the square root closes the generalization gap between RMSProp and SGD. This suggests that square-root-free adaptive methods can generalize well, and raises novel questions on the understanding of the role of adaptivity.

Computational cost

As noted by Duchi et al. (2011), the square root poses numerical and computational challenges on matrix adaptive methods like Equation 4 as it requires matrix square roots which must be carried out in high precision, to avoid numerical instabilities. This increases the run time, memory footprint, and complicates the implementation (Anil et al., 2020; Shi et al., 2023). Using low-precision data types (Micikevicius et al., 2017) is a key technique to boost training speed and reduce memory. The use of the square root makes these matrix methods undesirable for modern low-precision training schemes. By removing the square root, and using the latest advances on second-order methods (Lin et al., 2023), we can develop inverse-free matrix adaptive methods (Section 4.3) that are suitable for mixed-precision training. We present empirical results for the low-precision setting in Figure 3. Thanks to removing the square root, we can consistently train in BFP-16, whereas square-root-based adaptive matrix methods like Shampoo need single—in some cases, even double—precision. On modern vision transformers, we find that adaptive matrix methods outperform diagonal adaptive methods and perform similarly to Shampoo while requiring less time due to mixed-precision training. This shows that removing the square root allows us to overcome the challenges of existing matrix adaptive methods and expand their applicability to modern training pipelines.

Theoretical considerations

Invariances. Adding the square root makes a descent step invariant to the scale of the loss – the square root adjusts the scale of the “squared” gradient to be consistent with the gradient when the loss function is scaled. This is useful as there is no need for users to pay attention to whether the loss function is averaged or summed over data points. While adding the square root fixes the scaling issue, it breaks the affine reparameterization invariance (Nesterov & Nemirovskii, 1994). We can fix the scaling issue and preserve the affine invariance without the square root, as will be shown in Section 3; see Appendix A for an example of the affine invariance of square-root-free methods.

Convergence analysis. Theoretical works such as Duchi et al. (2011); Reddi et al. (2019) and many others suggest that adding the square root is useful to prove convergence bounds for convex loss functions. Convergence analysis for square-root-based methods is also extended to non-convex settings when certain assumptions such as gradient Lipschitz and Polyak-Łojasiewicz condition are satisfied. However, compared to the regret bound of AdaGrad (Duchi et al., 2011) studied in convex settings, Hazan et al. (2006) give a better regret bound in strongly-convex cases without introducing the square root. Recent works such as Mukkamala & Hein (2017); Wang et al. (2020) prove similar bounds for square-root-free methods for convex problems. Thus, square-root-free methods are theoretically grounded—at least in convex settings.

Behavior near optimum. The square root is often introduced to avoid oscillation near an optimal solution where the preconditioner can be ill-conditioned. For example, in one dimension when near the optimum, the descent direction s−1gs^{-1}g can be unbounded when we use the outer product as the preconditioner s=g2s=g^{2} because the gradient gg is close to at the optimum and the outer product as a ‘squared’ gradient decreases much faster than the gradient gg. Thus, the update without the square root can lead to oscillation when using a constant learning rate. However, even without the square root, a preconditioner can still be well-conditioned when incorporating outer products from previous iterations (Roux et al., 2007) and using Tikhonov damping (Becker et al., 1988; Martens & Sutskever, 2012). For example, the descent direction s−1gs^{-1}g can be bounded without the square root even when near an optimal solution, where gg is a gradient and s=(1−β2)s+β2g2s=(1-\beta_{2})s+\beta_{2}g^{2} is a preconditioner estimated by an exponentially moving average.

A Second-order Perspective

Here, we describe a second-order perspective on adaptive methods without the square root. This naturally resolves the scaling issue and is coherent with the original motivation to develop these methods.

Without the square root, an adaptive method takes a step

where the outer product \mbox{\mbox{H\mathbf{H}}}=\mbox{\mbox{g\mathbf{g}}}\mbox{\mbox{g\mathbf{g}}}^{T} is used as a curvature approximation. For example, when γ=0\gamma=0 and S\mathbf{S} is a full matrix, we obtain the update of full-matrix AdaGrad without the square root.

The above update resembles a second-order method. However, it is not invariant to the scale of the loss, unlike Newton’s method: Scaling the loss by a constant cc scales the gradient and Hessian by cc, but the gradient outer product by c2c^{2}, which is inconsistent with its role as a Hessian approximation and therefore conflicts with the second-order interpretation. We will resolve this particular conflict and improve the justification of square-root-free adaptive methods through approximate second-order methods.

The key step is to define the gradient outer product as a novel empirical Fisher that differs from the standard empirical Fisher discussed in the deep learning literature (Kingma & Ba, 2015; Kunstner et al., 2019; Martens, 2020). Our empirical Fisher relies on the aggregated mini-batch gradient, rather than per-sample gradients. Since the Fisher is tied to a probability distribution that must be normalized, this provides a gauge to automatically make updates invariant to the scale of the loss.

Now we introduce a new Fisher matrix as a FIM over a joint distribution of the labels.

(2) Our Fisher matrices for the original (unscaled) loss

where the labels \mbox{\mbox{y\mathbf{y}}}=(y_{1},\cdots,y_{N}) are considered jointly as a random vector, p(\mbox{\mbox{y\mathbf{y}}}|\mbox{\mbox{X\mathbf{X}}};\mbox{\mbox{μ\boldsymbol{\mu}}}):=\prod_{i=1}^{N}p(y_{i}|\mbox{\mbox{x\mathbf{x}}}_{i};\mbox{\mbox{μ\boldsymbol{\mu}}}) is its joint distribution, and \mbox{\mbox{X\mathbf{X}}}=(\mbox{\mbox{x\mathbf{x}}}_{1},\cdots,\mbox{\mbox{x\mathbf{x}}}_{N}) is a feature matrix.

Our empirical Fisher matrix is defined by replacing the expectation in \mbox{\mbox{F\mathbf{F}}}_{\text{new}}(\mbox{\mbox{μ\boldsymbol{\mu}}}) with observed labels.

where \nabla_{\mu}\log p(\mbox{\mbox{y\mathbf{y}}}|\mbox{\mbox{X\mathbf{X}}};\mbox{\mbox{μ\boldsymbol{\mu}}})=\sum_{i=1}^{N}\nabla_{\mu}\log p(y_{i}|\mbox{\mbox{x\mathbf{x}}}_{i};\mbox{\mbox{μ\boldsymbol{\mu}}})=-\sum_{i=1}^{N}\mbox{\mbox{g\mathbf{g}}}_{i}=-\mbox{\mbox{g\mathbf{g}}} .

In a mini-batch case, we can define a Fisher matrix for a mini-batch with BB data points as

where \mbox{\mbox{y\mathbf{y}}}_{\text{mini}}:=(y_{1},\cdots,y_{B}) is a label vector, \mbox{\mbox{X\mathbf{X}}}_{\text{mini}}:=(\mbox{\mbox{x\mathbf{x}}}_{1},\cdots,\mbox{\mbox{x\mathbf{x}}}_{B}) is a feature matrix for the mini-batch, and p(\mbox{\mbox{y\mathbf{y}}}_{\text{mini}}|\mbox{\mbox{X\mathbf{X}}}_{\text{mini}};\mbox{\mbox{μ\boldsymbol{\mu}}}):=\prod_{i=1}^{B}p(y_{i}|\mbox{\mbox{x\mathbf{x}}}_{i},\mbox{\mbox{μ\boldsymbol{\mu}}}) is the joint distribution over labels for the mini-batch. This distribution can also be obtained by marginalizing unseen labels in the original joint distribution p(\mbox{\mbox{y\mathbf{y}}}|\mbox{\mbox{X\mathbf{X}}};\mbox{\mbox{μ\boldsymbol{\mu}}}) defined for the full-batch.

This Fisher is an unbiased estimation of our full-batch Fisher as the claim—proof in Appendix B—is stated below.

Our mini-batch Fisher \frac{1}{B}\mbox{\mbox{F\mathbf{F}}}_{\text{mini}}(\mbox{\mbox{μ\boldsymbol{\mu}}}) is an unbiased estimation of the full-batch Fisher \frac{1}{N}\mbox{\mbox{F\mathbf{F}}}_{\text{new}}(\mbox{\mbox{μ\boldsymbol{\mu}}}).

When we replace the expectation with observed labels, we obtain our empirical Fisher for the mini-batch. Note that this empirical Fisher is not an unbiased estimator of the full-batch empirical Fisher. However, we often do not consider the unbiasedness when we incorporate empirical Fisher matrices from previous iterations. For example, consider the update of AdaGrad. Moreover, our interpretation provides the direct link to the Hessian as we can view the outer product as a Hessian approximation (see (7)) when the product corresponds to our Fisher. This interpretation allows us to preserve the scale-invariance to the loss and the affine reparametrization invariance in square-root-free methods.

(3) Our Fisher matrices for a scaled loss

Now, we consider a case when the loss function is scaled. In this case, the outer product H\mathbf{H} no longer coincides with a Fisher matrix. As an example, we consider averaging a loss over N>1N>1 data points, as it is often done in mini-batch settings. The loss is

where the gradient and the outer product are defined as \mbox{\mbox{g\mathbf{g}}}_{\text{scaled}}:=\frac{1}{N}\sum_{i=1}^{N}\nabla_{\mu}c(y_{i},f(\mbox{\mbox{x\mathbf{x}}}_{i};\mbox{\mbox{μ\boldsymbol{\mu}}}))=\frac{1}{N}\sum_{i=1}^{N}\mbox{\mbox{g\mathbf{g}}}_{i} and \mbox{\mbox{H\mathbf{H}}}_{\text{scaled}}:=\mbox{\mbox{g\mathbf{g}}}_{\text{scaled}}\mbox{\mbox{g\mathbf{g}}}_{\text{scaled}}^{T}=\frac{1}{N^{2}}\mbox{\mbox{H\mathbf{H}}} , respectively.

In this case, the empirical Fisher matrix should be defined using the same joint distribution as

Our fix does not break the affine invariance of a square-root-free method. A square-root-free method is affine invariant when using our scaled empirical Fisher matrix as shown in Claim 2 (see Appx. C for a proof).

Our square-root-free update is affine invariant.

By using our empirical Fisher matrices, our square-root-free methods like Newton’s method not only is scale-invariant when we scale a loss function but also preserves the affine reparametrization invariance. Thus, we consider our methods as approximate second-order methods.

2 Difference to the Standard Empirical Fisher

From the above equation, we can see that the standard empirical Fisher on the left is not a rank-one matrix while our empirical Fisher on the right is a rank-one matrix.

3 Disentangling the Outer Product from the Fisher

Matrix Adaptive Methods

Here, we develop a class of matrix adaptive methods without using matrix decompositions. This introduces an additional challenge: a dense matrix-valued preconditioner may be too costly to store. We address this by enforcing the preconditioner to be structured – specifically, Kronecker-factored. Equation 5 seems to suggest that the optimizer’s preconditioner S\mathbf{S} must have the same structure as the curvature approximation H\mathbf{H}, so that the structure is maintained under the update. Thus, an additional structural projection is often needed when using incompatible structures. A common approach to solve the projection sub-problem is to introduce a sequential inner loop. However, this increases the iteration cost due to the use of the loop. A prime example is Newton-CG which requires a computationally expensive Hessian-vector product at each iteration in the loop when the curvature approximation is the Hessian. Instead, we consider an inner-loop-free approach while allowing the curvature approximation H\mathbf{H} has an arbitrary structure. To do so, we decouple these concepts and formulate how to use the chain rule to incorporate arbitrarily structured curvature approximations into a class of structured preconditioners. Inspired by approximate second-order methods like KFAC and Shampoo, we then develop new matrix methods (inverse-free Shampoo) with a Kronecker-factored preconditioner.

We start with the Bayesian learning rule (BLR) (Khan & Lin, 2017; Khan et al., 2018; Zhang et al., 2018; Lin et al., 2020, 2021a, 2023; Tan, 2022; Khan & Rue, 2023). As discussed in Sec. 3, if the product coincides with our empirical Fisher, we can view the outer product as a Hessian approximation. This second-order view allows us to extend the BLR originally developed for Newton’s method. The BLR views the preconditioner S\mathbf{S} in Equation 5 as an inverse covariance of a Gaussian, and the curvature approximation as a partial derivative. This perspective not only allows the preconditioner and the curvature approximation to have their independent structures, but also allows matrix adaptive methods to become inverse-free by reparameterizing the preconditioner as the inverse covariance.

First, we consider a Bayesian problem formulation and solve a variational inference problem with a Gaussian variational approximation. In this setting, we consider NN weights as random variables and use a new symbol w\mathbf{w} to denote these weights as they are no longer learnable. We refer to μ\boldsymbol{\mu} and S\mathbf{S} as the mean and the inverse covariance of the Gaussian q(\mbox{\mbox{w\mathbf{w}}}|\mbox{\mbox{μ\boldsymbol{\mu}}},\mbox{\mbox{S\mathbf{S}}}). This variational inference problem (Barber & Bishop, 1997) (with γ=1\gamma=1) is

The FIM of the Gaussian q(\mbox{\mbox{w\mathbf{w}}}|\mbox{\mbox{θ\boldsymbol{\theta}}}) under parameterization θ\boldsymbol{\theta} is defined as

This Fisher has a closed-form expression and is positive-definite as long as θ\boldsymbol{\theta} is a valid parameterization of the Gaussian. For example, \mbox{\mbox{θ\boldsymbol{\theta}}}:=(\mbox{\mbox{μ\boldsymbol{\mu}}},\mbox{\mbox{S\mathbf{S}}}) is a valid parameterization if and only if the inverse covariance S\mathbf{S} is positive-definite. This FIM should not be confused with other Fisher matrices in Section 3.

We use natural gradient descent (NGD) to solve the problem in Equation 11. An NGD step in parameterization θ\boldsymbol{\theta} is

Khan et al. (2018) show that this update is equivalent to

where we use the following gradient identities which hold when q(\mbox{\mbox{w\mathbf{w}}}|\mbox{\mbox{μ\boldsymbol{\mu}}},\mbox{\mbox{S\mathbf{S}}}) is Gaussian (Opper & Archambeau, 2009),

To recover the square-root-free update in Equation 5, we approximate the update in Equation 13 with

by using a delta approximation at μ\boldsymbol{\mu} to approximate the expectations highlighted in red, replacing the Hessian with the outer product H\mathbf{H} as it coincides with our empirical Fisher in the unscaled case shown in (7), and introducing an extra learning rate β2\beta_{2} to compensate for the error of the curvature approximation. By using the delta approximation, we can use the NGD update derived for the Bayesian problem in Equation 11 to solve the non-Bayesian problem in Equation 1. This Bayesian formulation unveils the Gaussian approximation hidden in square-root-free updates as the NGD update recovers the square-root-free method. This is possible when we view the update from a second-order perspective as we strengthen the perspective in Section 3. If the loss is scaled, we can replace the outer product with our empirical Fisher as a proper Hessian/curvature approximation (c.f. Section 3).

Many square-root-free adaptive gradient methods can be derived from this update rule. For example, Equation 15 becomes the update of full-matrix AdaGrad update without the square root when γ=0\gamma=0:

We obtain a square-root-free RMSProp update when γ=1\gamma=1; similarly, a square-root-free AdaGrad update when γ=0\gamma=0.

Non-zero initialization

For adaptive gradient methods, the preconditioner S\mathbf{S} is often initialised to zero. An immediate consequence of the BLR perspective is that S\mathbf{S} should be initialized to a non-zero value because we view it as an inverse covariance. As shown in Figure 2, this is important for the performance of square-root-free methods in practice. Moreover, the updated S\mathbf{S} in Equation 15 is guaranteed to be positive-definite (PD) when S\mathbf{S} is initialised to a PD matrix since the product \mbox{\mbox{H\mathbf{H}}}=\mbox{\mbox{g\mathbf{g}}}\mbox{\mbox{g\mathbf{g}}}^{T} is semi-positive-definite.

2 Decoupling Preconditioner & Curvature

Note that in Equation 13 the preconditioner is the inverse covariance while the curvature approximation H\mathbf{H} appears as a term of the partial derivative ∂μL\partial_{\mu}\mathcal{L}. At first glance, the update of S\mathbf{S} in Equation 13 requires S\mathbf{S} and H\mathbf{H} to have the same structure. However, if a structure in S\mathbf{S} can be obtained via reparameterization, we can perform NGD in this reparameterized space due to the parameterization invariance of natural gradients. By doing so, we allow the curvature approximation to have its own independent structure and use the chain rule to project the curvature approximation onto the reparameterized space of S\mathbf{S}. Importantly, our approach does not introduce significant computational overhead since we do not require an inner loop to solve the projection problem. Thus, the preconditioner S\mathbf{S} and the curvature approximation H\mathbf{H} can have distinct structures since they are decoupled by NGD on the variational problem. We will see an example in the next section. More examples can be found at Lin et al. (2021a, b).

3 Kronecker-factored Adaptive Methods

Shampoo is a square-root-based Kronecker-factored method, where the inner (matrix) square roots are introduced due to the structural approximation and the outer (matrix) square root is inherited from full-matrix AdaGrad (Gupta et al., 2018). The preconditioner S\mathbf{S} in Shampoo is not decoupled from curvature approximation H\mathbf{H}. Thus, the authors approximate \mbox{\mbox{H\mathbf{H}}}\approx{\hat{\mbox{\mbox{S\mathbf{S}}}}}_{C}^{1/2}\otimes{\hat{\mbox{\mbox{S\mathbf{S}}}}}_{K}^{1/2} and use \mbox{\mbox{S\mathbf{S}}}=\left(\hat{\mbox{\mbox{S\mathbf{S}}}}_{C}^{1/2}\otimes\hat{\mbox{\mbox{S\mathbf{S}}}}_{K}^{1/2}\right)^{1/2} as a preconditioner to ensure they have the same structure. The Shampoo update is

In contrast, we consider a structural reparameterization of the inverse covariance \mbox{\mbox{S\mathbf{S}}}=\mbox{\mbox{S\mathbf{S}}}_{C}\otimes\mbox{\mbox{S\mathbf{S}}}_{K} as the preconditioner and perform NGD to update \mbox{\mbox{S\mathbf{S}}}_{C} and \mbox{\mbox{S\mathbf{S}}}_{K}. In our approach, we treat a curvature approximation, such as the gradient outer product H\mathbf{H}, as a partial derivative and use it to update \mbox{\mbox{S\mathbf{S}}}_{C} and \mbox{\mbox{S\mathbf{S}}}_{K} by the chain rule. By doing so, we do not require the approximation H\mathbf{H} to have the same structure as the preconditioner S\mathbf{S}. Thus, we decouple the preconditioner S\mathbf{S} from a curvature approximation H\mathbf{H} without introducing the approximation for H\mathbf{H} in Shampoo. This allows us to obtain a square-root-free update scheme.

where \mbox{\mbox{F\mathbf{F}}}_{\mu\mu}=\mbox{\mbox{S\mathbf{S}}}=\mbox{\mbox{S\mathbf{S}}}_{C}\otimes\mbox{\mbox{S\mathbf{S}}}_{K} , \mbox{\mbox{F\mathbf{F}}}_{CC}=-\frac{d}{2}\frac{\partial\mbox{\mbox{S\mathbf{S}}}_{C}^{-1}}{\partial\mbox{\mbox{S\mathbf{S}}}_{C}} , and \mbox{\mbox{F\mathbf{F}}}_{KK}=-\frac{p}{2}\frac{\partial\mbox{\mbox{S\mathbf{S}}}_{K}^{-1}}{\partial\mbox{\mbox{S\mathbf{S}}}_{K}} . We use the chain rule to compute these derivatives

We can obtain an inverse-free update by reparameterizing \mbox{\mbox{S\mathbf{S}}}^{-1}=\mbox{\mbox{S\mathbf{S}}}_{C}^{-1}\otimes\mbox{\mbox{S\mathbf{S}}}_{K}^{-1}=(\mbox{\mbox{C\mathbf{C}}}\mbox{\mbox{C\mathbf{C}}}^{T})\otimes(\mbox{\mbox{K\mathbf{K}}}\mbox{\mbox{K\mathbf{K}}}^{T}) and directly updating C\mathbf{C} and K\mathbf{K} instead of \mbox{\mbox{S\mathbf{S}}}_{C} and \mbox{\mbox{S\mathbf{S}}}_{K} as suggested by Lin et al. (2023). This update (IF-Shampoo) is inverse-free and square-root-free (see Appx. 5), avoiding numerically unstable matrix inversions and decompositions. (derivation in Appx. D). Moreover, this update is equivalent to directly updating \mbox{\mbox{S\mathbf{S}}}_{C} and \mbox{\mbox{S\mathbf{S}}}_{K} up to a first-order accuracy.

On vision transformers, we find that IF-Shampoo performs similarly to Shampoo in terms of per-iteration progress. However, our method can run in BFP-16; in contrast to Shampoo. We observe that one step of IF-Shampoo takes half the time of Shampoo and requires significantly less memory. This is a promising result to make matrix adaptive methods more prominent in modern large-scale training.

Conclusion

We investigated how the behavior of adaptive methods changes when we remove the square root and thereby strengthen their motivation from a second-order perspective. Surprisingly, we found empirically that removing the square root not only closes the generalization gap between adaptive methods and SGD on convolutional NNs, but also maintains the performance of square-root-based methods on vision transformers. Removing the square root eliminates the connection to sign descent which has been hypothesized to cause the gap on convolutional NNs and transformers. However, our findings highlight that adaptivity might be an important concept for the success of such methods that is currently overlooked, which poses new questions regarding the role of adaptivity and the understanding of adaptive methods. Conceptually, we established a rigorous second-order view on square-root-free adaptive methods by viewing their gradient outer product as a novel empirical Fisher that differs from the standard empirical Fisher discussed in the deep learning literature and allows to recover the scale invariance that is inherent to square-root-based methods. This perspective allowed us to develop IF-Shampoo, a matrix adaptive, inverse-free method that stably works in BFP-16 and trains roughly 2x faster than its square-root-based counterpart Shampoo. We provide novel insights for the understanding and development of adaptive gradient methods.

ACKNOWLEDGMENTS

We thank Emtiyaz Khan, Mark Schmidt, Kirill Neklyudov, and Roger Grosse for helpful discussions at the early stage of this work. Resources used in preparing this research were provided, in part, by the Province of Ontario, the Government of Canada through CIFAR, and companies sponsoring the Vector Institute.

References

Appendix A Example: Affine Invariance of Root-Free Methods

We demonstrate the affine invariance by an example and show how adding the root breaks the invariance. Consider a loss function l_{a}(a)=\mbox{\frac{1}{2}}a^{2} with an initial point a0=2a_{0}=2 , a root-based update is anew=a0−sa−1ga=2−1=1a_{\text{new}}=a_{0}-s_{a}^{-1}g_{a}=2-1=1 , where gradient ga=∇ala(a)=a0=2g_{a}=\nabla_{a}l_{a}(a)=a_{0}=2 and preconditioner sa=ga2=∣a0∣=2s_{a}=\sqrt{g_{a}^{2}}=|a_{0}|=2 . Now, consider a reparameterized loss as l_{2}(b)=\mbox{\frac{1}{2}}(2b)^{2} with an initial point b0b_{0}, where a=2ba=2b. Thus, b0=1b_{0}=1 when a0=2a_{0}=2. The update becomes bnew=b0−sb−1gb=1−1=0b_{\text{new}}=b_{0}-s_{b}^{-1}g_{b}=1-1=0 , where gradient gb=∇blb(b)=4b0=4g_{b}=\nabla_{b}l_{b}(b)=4b_{0}=4 and preconditioner sb=gb2=4∣b0∣=4s_{b}=\sqrt{g_{b}^{2}}=4|b_{0}|=4. Unfortunately, the updated anew=1a_{\text{new}}=1 is not equivalent to the updated bnew=0b_{\text{new}}=0 since anew≠2bnewa_{\text{new}}\neq 2b_{\text{new}}.

Now, consider a square-root-free update for the original loss as anew=a0−sa−1ga=2−0.5=1.5a_{\text{new}}=a_{0}-s_{a}^{-1}g_{a}=2-0.5=1.5, where gradient ga=∇ala(a)=a0=2g_{a}=\nabla_{a}l_{a}(a)=a_{0}=2 and preconditioner sa=ga2=a02=4s_{a}=g_{a}^{2}=a_{0}^{2}=4. Similarly, the update for the reparameterized loss is bnew=b0−sb−1gb=1−0.25=0.75b_{\text{new}}=b_{0}-s_{b}^{-1}g_{b}=1-0.25=0.75, where gradient gb=∇blb(b)=4b0=4g_{b}=\nabla_{b}l_{b}(b)=4b_{0}=4 and preconditioner sb=gb2=16b02=16s_{b}=g_{b}^{2}=16b_{0}^{2}=16. Note that the updated anew=1.5a_{\text{new}}=1.5 is equivalent to the updated bnew=0.75b_{\text{new}}=0.75 since anew=2bnewa_{\text{new}}=2b_{\text{new}}. Thus, removing the root preserves the affine invariance.

Appendix B Proof of Claim 1

We first show that our Fisher matrix coincides with the standard Fisher.

Similarly, we can show our mini-batch Fisher coincides with the standard mini-batch Fisher.

Thus, it is easy to see that \frac{1}{B}\mbox{\mbox{F\mathbf{F}}}_{\text{mini}}(\mbox{\mbox{μ\boldsymbol{\mu}}}) is an unbiased estimation of \frac{1}{N}\mbox{\mbox{F\mathbf{F}}}_{\text{new}}(\mbox{\mbox{μ\boldsymbol{\mu}}}) since the standard mini-batch Fisher is an unbiased estimation of the standard full-batch Fisher.

Now, we show that our Fisher matrix coincides with the standard Fisher. Recall that we define the joint distribution over labels is p(\mathbf{y}|\mbox{\mbox{X\mathbf{X}}};\mbox{\mbox{μ\boldsymbol{\mu}}})=\prod_{i=1}^{N}p(y_{i}|\mbox{\mbox{x\mathbf{x}}}_{i};\mbox{\mbox{μ\boldsymbol{\mu}}})

where the last line is due to the independence of per-sample distributions in our joint distribution as shown below.

Since each a per-sample distribution is independent, we have

where we make use of the following result as the per-sample distribution is normalized.

Appendix C Proof of Claim 2

Recall that the (unscaled) optimization problem in (1) is

Now, consider reparametrizing μ\boldsymbol{\mu} with a known non-singular matrix A\mathbf{A} as \mbox{\mbox{μ\boldsymbol{\mu}}}=\mbox{\mbox{A\mathbf{A}}}\mbox{\mbox{m\mathbf{m}}}. In this case, the optimization problem becomes

We will show that a square-root-free method is affine invariant at each step. In other words, if we use the same square-root-free method to solve these two problems, they are equivalent.

For the first problem, the method takes the following step at iteration tt

For the second problem, we assume \mbox{\mbox{S\mathbf{S}}}_{0}^{\text{rep}} is initialized with \mbox{\mbox{A\mathbf{A}}}^{T}\mbox{\mbox{S\mathbf{S}}}_{0}\mbox{\mbox{A\mathbf{A}}} and \mbox{\mbox{A\mathbf{A}}}^{-1}\mbox{\mbox{μ\boldsymbol{\mu}}}_{0}=\mbox{\mbox{m\mathbf{m}}}_{0} since A\mathbf{A} is known. In this case, the square-root-free update at the first iteration becomes

From abvoe expressions, we can see that both updates are equivalent at the first iteration since \mbox{\mbox{μ\boldsymbol{\mu}}}_{1}=\mbox{\mbox{A\mathbf{A}}}_{1}\mbox{\mbox{m\mathbf{m}}}_{1}. Similarly, we can show that both updates are equivalent at every iteration by induction.

Thus, we can see that full-matrix square-root-free method is affine invariance. For a diagonal square-root-free method, it only preserves a diagonal invariance. Likewise, we can show the update is affine invariance in a scaled case when using our empirical Fisher.

Appendix D Derivation of our matrix inverse-free method

where we can use the chain rule to compute the partial derivative at \mbox{\mbox{η\boldsymbol{\eta}}}_{C}^{\text{cur}}=\mathbf{0}

Thus, the update for block C\mathbf{C} can be re-expressed as

We can also include a damping term \lambda\mbox{\mbox{I\mathbf{I}}}_{dp}=\lambda\mbox{\mbox{I\mathbf{I}}}_{d}\otimes\mbox{\mbox{I\mathbf{I}}}_{p} into the curvature approximation such as \mbox{\mbox{H\mathbf{H}}}=\mbox{\mbox{g\mathbf{g}}}\mbox{\mbox{g\mathbf{g}}}^{T}+\lambda\mbox{\mbox{I\mathbf{I}}}_{d}\otimes\mbox{\mbox{I\mathbf{I}}}_{p}. Recall that we do not assume that the curvature approximation H\mathbf{H} has the same structure as the preconditioner S\mathbf{S}. We can similarly update blocks K\mathbf{K} and μ\boldsymbol{\mu}. The deails of the complete update can be found at Fig. 7.

Appendix E Additional Results in Convex Settings