Title: A Statistical Viewpoint on Modern Matrix-Based Pretraining Methods
Abstract: I will present a new viewpoint for modelling pretraining methods based on whitening. Stochastic gradients are noisy estimates of the full-batch gradient, with some components having high variance and others low variance. We need to take this variance into account; otherwise, we might accidentally take a large step in a high-variance direction, leading the method astray. By modelling the gradient with a normal distribution, I will argue that gradient estimates should be whitened using their covariance.
To make covariance whitening computationally efficient, we assume that the covariance has a Kronecker-factor structure, corresponding to a matrix normal model. If we estimate these factors using empirical moments, we recover the classic Shampoo method. However, I will show that there are better ways of estimating the factors by computing a maximum a posteriori (MAP) estimate. We refer to the resulting method as MN (Matrix Normal). For the non-Bayesian crowd (myself included), I will also present MN from the perspective of a proximal-point algorithm based on the likelihood. The resulting MN methods consistently outperform methods in the same class, including Shampoo, Distributed Shampoo, and KL-Shampoo, when training a small language model.
Time permitting, I will also give a statistical interpretation of Muon: its update can be viewed as the MAP estimate of a matrix Langevin distribution. Using this interpretation, we combine Muon with whitening to obtain a new MN-Muon method that consistently, albeit modestly, improves upon Muon when training small language models. This matters because Muon is one of the innovations driving the success of recent frontier-scale models from Moonshot AI, DeepSeek, and Qwen.