Adam scales every parameter by its own running gradient magnitude. That works well when the loss curvature lines up with the coordinate axes, and poorly when it does not: if two weights are strongly coupled, the useful direction is a mix of both, and a per-coordinate scale cannot represent it. Shampoo fixes this with Kronecker-factored preconditioners that capture correlations between rows and between columns of a weight matrix, but it costs matrix roots and extra hyperparameters.

SOAP, short for ShampoO with Adam in the Preconditioner's eigenbasis, was introduced by Vyas and colleagues in 2024 (arXiv 2409.11321). Its idea fits in one sentence: rotate each gradient matrix into the eigenbasis of Shampoo's preconditioner factors, run ordinary Adam there, and rotate the update back. This article derives why that works, walks through the reference implementation step by step, gives numpy code and a toy comparison, and sizes memory and compute for real layer shapes.

Why a diagonal optimizer needs the right basis

Write a layer's weight as an m by n matrix W with gradient G. Adam keeps two moving averages per entry, the mean M and the mean square V, and steps by M divided by the square root of V. The division is elementwise, so Adam is a diagonal preconditioner: it can stretch the coordinate axes, but it cannot rotate them.

Picture a two-weight valley whose long axis runs diagonally. Both coordinates see similar gradient magnitudes, so Adam barely changes the direction and the step zig-zags across the valley. Rotate the coordinates by 45 degrees so that one axis lies along the valley, and the same diagonal rescaling now does exactly the right thing: big steps along the valley, small ones across it. A diagonal method is only as good as the basis it works in.

Shampoo estimates such a basis cheaply. For a matrix parameter it accumulates two factors, L = sum of G G^T (m by m, correlations between rows) and R = sum of G^T G (n by n, correlations between columns). The Kronecker product of these factors approximates the full mn by mn gradient covariance, which nobody can store for a real layer. The eigenvectors of L and R, QL and QR, define a rotated coordinate system for the matrix: the rotated gradient is QLT G QR.

Shampoo is Adafactor in a rotated basis

The paper's key observation is that Shampoo, when used with exponent 1/2 rather than the original 1/4, is already a diagonal method in disguise. Diagonalize the factors: L = QL diag(a) QLT and R = QR diag(b) QRT. The Shampoo update L-1/2 G R-1/2 becomes

Q_L [ G' / sqrt(a_i * b_j) ] Q_R^T        where G' = Q_L^T G Q_R

Rotate, divide entry (i, j) by sqrt(ai bj), rotate back. Now, ai is the i-th diagonal entry of QLT L QL, which is the accumulated sum of squares of row i of the rotated gradient G'. Likewise bj is the sum of squares of column j. Adafactor, the memory-saving cousin of Adam, approximates the second moment of entry (i, j) by row sum times column sum divided by the total sum. So dividing by sqrt(ai bj) is Adafactor's estimate up to a single global constant, which the learning rate absorbs. Shampoo with power 1/2 is Adafactor running in Shampoo's eigenbasis.

The improvement suggests itself. Adafactor's rank-one estimate exists to save memory; if you can afford Adam's full per-entry second moment, keep it, still in the rotated basis. That is SOAP. It inherits Adam's hyperparameters and adds essentially one new one, the preconditioning frequency: how often to recompute the basis.

The SOAP algorithm as implemented

The reference implementation in the authors' repository (nikhilvyas/SOAP, file soap.py) has details that matter when you reimplement or debug it. Per matrix parameter:

  1. On the very first step, accumulate L and R from the gradient, take an exact eigendecomposition (torch.linalg.eigh) to get the first basis, and make no parameter update at all. The code comment says this is so the current gradient is never used in its own projection.
  2. Rotate the gradient: G' = QLT G QR.
  3. Update the moments in the rotated space: M = beta1 M + (1 - beta1) G' and V = beta2 V + (1 - beta2) G'2.
  4. Compute N' = M / (sqrt(V) + eps) with Adam's bias corrections folded into the step size, rotate back with N = QL N' QRT, and subtract lr times N.
  5. Apply decoupled weight decay, as AdamW does: W = W - lr * wd * W.
  6. Update the factors every step as exponential moving averages: L = shampoo_beta L + (1 - shampoo_beta) G GT, and the same for R. When shampoo_beta is left at its default of -1 the code uses beta2.
  7. Every precondition_frequency steps, refresh the basis. Instead of a full eigendecomposition the code warm-starts from the old basis: estimate eigenvalues as the diagonal of QT L Q, sort descending, do one power-iteration step (multiply L by the sorted Q) and orthonormalize with a QR decomposition.

Two state-handling details are easy to get wrong. The first moment M is projected back to the original basis with the old Q and forward again with the new Q around every refresh, so momentum survives a basis change intact. The second moment V is not rotated, because the rotation of a squared quantity is not the square of a rotated one. It is only permuted to follow the eigenvalue sort, and the moving average adapts it over the next steps. That continual re-estimation is why SOAP degrades much less than Shampoo when refreshes are infrequent: Shampoo's scaling freezes between refreshes, SOAP's keeps learning.

A dimension larger than max_precond_dim (default 10,000) is left unrotated, so a 4096 by 50,000 embedding is preconditioned on one side only, and 1-D parameters such as biases and norm gains get plain AdamW unless precondition_1d is set.

Two clocks: per-step work and basis refresh

One SOAP step for a weight matrix W (m x n)Gradient Gm x n, original basisRotateG' = QL^T G QRSecond moment VEMA of G' squaredFirst moment MEMA of G', rotated basisAdam stepM' / (sqrt(V) + eps)Rotate backQL N' QR^TUpdate Wlr step + decoupled wdEvery step: accumulateL += G G^T, R += G^T G (EMA)Every f steps: refreshQL, QR via power step + QROn refreshre-project M, permute Vnew basisAdam runs unchanged inside the rotated coordinates. Only the basis is expensive, and it is refreshed rarely.
The per-step path is cheap: rotate the gradient in, an elementwise Adam update, rotate the update out. The eigenbasis refresh is the expensive part and runs every f steps.

A numpy implementation

Here is a compact numpy version for one matrix parameter. It keeps M in the original basis and rotates it each step, which is mathematically equivalent to the reference code's re-projection on refresh and easier to read. For real training use a maintained PyTorch or JAX implementation.

import numpy as np

def soap(loss_grad, W, steps, lr, b1=0.95, b2=0.95, eps=1e-8, freq=10):
    m, n = W.shape
    L, R = np.zeros((m, m)), np.zeros((n, n))
    M, V = np.zeros_like(W), np.zeros_like(W)
    QL = QR = None
    t = 0
    for _ in range(steps + 1):
        _, G = loss_grad(W)
        if QL is not None:                        # step 0 only builds the basis
            t += 1
            Gr = QL.T @ G @ QR
            M = b1 * M + (1 - b1) * G
            V = b2 * V + (1 - b2) * Gr * Gr
            step = lr * np.sqrt(1 - b2 ** t) / (1 - b1 ** t)
            W = W - step * (QL @ ((QL.T @ M @ QR) / (np.sqrt(V) + eps)) @ QR.T)
        L = b2 * L + (1 - b2) * G @ G.T
        R = b2 * R + (1 - b2) * G.T @ G
        if QL is None:
            QL, QR = np.linalg.eigh(L)[1], np.linalg.eigh(R)[1]
        elif t % freq == 0:
            QL, QR, V = refresh(L, R, QL, QR, V)
    return W

def refresh(L, R, QL, QR, V):
    out = []
    for axis, (P, Q) in enumerate(((L, QL), (R, QR))):
        idx = np.argsort(-np.diag(Q.T @ P @ Q))   # estimated eigenvalues, descending
        Q, V = Q[:, idx], np.take(V, idx, axis=axis)
        Q, _ = np.linalg.qr(P @ Q)                # one power-iteration step
        out.append(Q)
    return out[0], out[1], V

Weight decay is omitted for brevity. Note that there is no matrix inverse or root anywhere, only one eigendecomposition per factor at the start and one QR per factor per refresh.

Worked example: a rotated, ill-conditioned problem

To see the effect, take a 32 by 16 weight and the loss 0.5 ||A (W - W*) B||2, where A and B are randomly rotated matrices whose squared singular values span a factor of 1,000. That is ill-conditioned curvature off the coordinate axes, the case a diagonal method handles badly. Gradients get Gaussian noise with standard deviation 0.1. Both optimizers use betas (0.95, 0.95), SOAP refreshes every 10 steps, and each optimizer gets its best learning rate from the grid 0.003, 0.01, 0.03, 0.1, 0.3, 1.0. The starting loss is about 6.29 million.

StepsAdam: best final loss (lr)SOAP: best final loss (lr)
2001,708 (0.1)370 (0.1)
1,00040.3 (0.01)0.50 (0.01)

These numbers come from the SOAP code above (plus the 0.1 gradient noise) on this problem, against a standard Adam loop with the same betas, one seed. They illustrate the mechanism, not a benchmark: rotated quadratic curvature is the best case for an eigenbasis method. The published evidence is the paper's language-model pre-training runs: at 360M and 660M parameters, in the large-batch regime, SOAP reduced the number of iterations by over 40% and wall-clock time by over 35% compared with AdamW, and improved on Shampoo by roughly 20% on both measures. Keep those qualifiers when you repeat the claim.

Memory and compute for real layer shapes

For an m by n matrix, AdamW stores 2mn numbers. SOAP adds L, R, QL and QR, which is 2(m2 + n2) more. Two shapes from a model with hidden size 4,096 show the range.

LayerAdamW stateExtra for SOAPTotal vs AdamW
Attention projection 4096 x 409633.6M values67.1M (both sides rotated)3.0x
MLP up-projection 4096 x 16384134.2M values33.6M (16,384 exceeds 10,000, so one side only)1.25x

In fp32 the attention projection's optimizer state grows from about 134 MB to 403 MB. Square layers are the expensive case. Sharding optimizer state (ZeRO, FSDP) divides this by the shard count, but each refresh then needs a factor computed by its owner.

For compute, the reference code adds two rotations per step (gradient in, update out; the listing above adds a third for M) and two factor updates, about 6mn(m + n) operations, against the layer's roughly 6mn per token for forward and backward. The overhead ratio is about (m + n) divided by tokens per step: under 1% for a 4096 by 4096 layer at a million tokens per batch. The refresh costs order m3 + n3 and is amortized over f steps; SOAP tolerates large f far better than Shampoo does.

Running it in training

  • Precision. Keep L, R, Q and the moments in fp32 even when training in bf16. The reference code runs eigh in fp32 and retries in fp64 if it fails to converge.
  • Learning rate. SOAP's update has Adam's scale, so start from your tuned AdamW learning rate and warmup, then sweep a factor of about 3 either way. The reference defaults are 3e-3 and betas (0.95, 0.95); the factor EMA inherits beta2 by default.
  • Checkpointing. A custom checkpointer must save L, R and both Q matrices; without them a restart reruns step 0 and rebuilds the basis from one gradient.
  • Monitoring. Log refresh time per layer and how far the basis moves between refreshes; lower f if it moves a lot, raise it if it barely moves.

How it compares

OptimizerWhat it preconditions withState per m x n matrixExtra hyperparameters
AdamWPer-entry second moment, fixed axes2mnnone
AdafactorFactored second moment, fixed axesabout m + nfactoring rules
ShampooL and R roots, refreshed every f stepsmn + 2(m^2 + n^2) or moreexponent, epsilon, grafting, f
SOAPAdam in the eigenbasis of L and R2mn + 2(m^2 + n^2)f
MuonOrthogonalized momentum, no second momentmnNewton-Schulz steps

Muon, the other popular non-diagonal optimizer, keeps only momentum and replaces each update with its nearest orthogonal matrix, so it uses less memory than AdamW but discards per-direction scale information. SOAP keeps that information, in a better basis, and pays for it in memory. Neither dominates at every scale.

Failure modes

  • Rotating V on refresh. A reimplementation that rotates the second moment like the first produces wrong, sometimes negative, variances. Permute it and let the EMA adapt, as the reference code does.
  • Stale bases with large f. Very infrequent refreshes leave Adam running in an outdated basis. SOAP tolerates this better than Shampoo but not without limit; watch loss around refresh boundaries for sawtooth patterns.
  • Unfair baselines. A SOAP win against an untuned AdamW is not a win. Tune both on the same budget.

What to do next

  1. Run the numpy listing on the toy problem above, then change the conditioning and the refresh frequency to see when the gap closes.
  2. Read the reference soap.py alongside this article, checking each of the seven steps in the algorithm section against the code.
  3. Compute SOAP's optimizer-state memory for every distinct layer shape in your model, using max_precond_dim, and confirm it fits after sharding.
  4. Run a short ablation at a small model size: tuned AdamW against SOAP at f = 10 and f = 50, same tokens, same schedule, compared on validation loss against wall-clock time.
  5. If SOAP wins, scale one step at a time and re-check the gap, logging refresh time and basis drift as you go.

Related reading on this site: Shampoo, the second-order optimizer SOAP builds on, Adam and AdamW from first principles, the AdamW math deep dive, Adafactor and its factored second moment and Muon and orthogonalized updates.

Key takeaway: SOAP rotates each gradient matrix into the eigenbasis of the Shampoo factors, runs ordinary Adam there and rotates back, adding one hyperparameter and 2(m^2 + n^2) state per matrix. Budget memory per layer shape and judge it against a tuned AdamW on wall-clock time.