The full-matrix ideal, and why nobody can afford it

Start from the thing Shampoo is approximating. Flatten one layer’s weights into a vector g ∈ R^N with N = m*n, and run full-matrix AdaGrad: accumulate the outer products H_t = Σ_{s≤t} g_s g_s^T and step with x_{t+1} = x_t - η (H_t + εI)^(-1/2) g_t. This is the ideal because it whitens the gradient: directions that have already accumulated a lot of gradient energy get damped, and the update becomes invariant to any linear reparameterization of the layer.

It is also unaffordable by a margin that is hard to overstate. H has N^2 entries and its inverse square root costs O(N^3). For a modest 1024 × 1024 layer, N ≈ 1.05e6, so H holds 1.1e12 numbers — 4.4 TB in fp32 — and one root costs ~1.2e18 FLOPs. Per layer, per update. Every practical adaptive optimizer is therefore a structured approximation of H. Adam takes the diagonal and discards all coupling; Shampoo takes a different slice.

Advertisement

Kronecker factorization: two small factors, not one huge one

Shampoo’s move is to stop flattening. Keep the gradient in its natural matrix shape and accumulate statistics on each side separately:

G_t : [m, n]                     gradient of one weight matrix
L_t = L_{t-1} + G_t G_t^T  : [m, m]   row-space (output-unit) statistics
R_t = R_{t-1} + G_t^T G_t  : [n, n]   column-space (input-feature) statistics

W_{t+1} = W_t - η * L_t^(-1/4) G_t R_t^(-1/4)

L sees how output units co-vary; R sees how input features co-vary. Together they store m^2 + n^2 numbers where the full matrix needed (m*n)^2. Back to the 1024 × 1024 layer: 2 × 1024^2 ≈ 2.1e6 floats, about 8.4 MB against 4.4 TB, and the two roots cost ~2.1e9 FLOPs against 1.2e18. The implicit preconditioner is still a full mn × mn object — the Kronecker product of the two factors — but it has only m^2 + n^2 degrees of freedom.

Advertisement

Where the fourth root comes from

The exponent looks arbitrary until you write the Kronecker algebra down. Using the column-major convention (A ⊗ B) vec(X) = vec(B X A^T), Gupta, Koren and Singer prove the key bound: the true AdaGrad matrix is dominated by the square roots of the two factors,

Σ_s vec(G_s) vec(G_s)^T  ⪯  R_t^(1/2) ⊗ L_t^(1/2).

AdaGrad wants H^(-1/2), so apply that exponent to the bound: (R^(1/2) ⊗ L^(1/2))^(-1/2) = R^(-1/4) ⊗ L^(-1/4), and by the vec identity that acts on a gradient as L^(-1/4) G R^(-1/4). Half of the -1/2 lands on each side.

A degree check confirms it independently. L is quadratic in G, so L^(-1/4) scales like G^(-1/2), and the product L^(-1/4) G R^(-1/4) is homogeneous of degree -1/2 + 1 - 1/2 = 0 — exactly like H^(-1/2) g. Doubling every gradient leaves the step unchanged. The general rule for an order-k tensor is k factors each raised to -1/(2k); two factors give -1/4.

Rotation, not just rescaling

Why should this help? Eigendecompose L = U Λ U^T; then L^(-1/4) = U Λ^(-1/4) U^T. Multiplying the gradient on the left by that matrix rotates it into the basis of output-unit correlations, damps whichever directions have accumulated the most gradient energy, and rotates back. R does the same on the input side. Adam, restricted to a diagonal, can only rescale entries in the fixed coordinate basis; it has no way to express ‘these two output units have been moving together, so treat their shared direction as one stiff axis.’ Shampoo can, along m row directions and n column directions.

Two honest caveats. The factorization assumes curvature separates into a row part and a column part; correlations that do not factor that way remain invisible. And despite the name, this is not a Newton method: the statistic is accumulated gradient second moments, the same raw material AdaGrad uses, not the Hessian.

Computing the inverse fourth root

Two routes, with different hardware profiles. The direct one is a symmetric eigendecomposition: L is symmetric PSD, so eigh gives U, Λ and L^(-1/4) = U (Λ + εI)^(-1/4) U^T, at a cost of c*m^3. Do it in float64: fp32 eigensolvers on a near-singular accumulator return garbage or fail outright. The consoling fact is that a fourth root is gentle — a condition number κ = 1e12 becomes κ^(1/4) = 1e3 in the root, where an inverse square root would leave 1e6.

The alternative is a coupled Newton iteration for the inverse p-th root, which uses only matrix multiplies. Normalize so ||A|| ≤ 1,start from X = I, M = A, and repeatedly form the factor ((p+1)I - M)/p, updating X ← X * factor and M ← factor^p * M. Convergence is quadratic once M is near the identity, which is precisely what the normalization buys. On hardware where large matmuls vastly outrun an eigensolver, this wins.

The cost dial: update frequency, memory, blocking

Shampoo has two separable costs. Per step: L += G G^T is 2m^2 n, R += G^T G is 2m n^2, and applying both roots is another 2m^2 n + 2m n^2 — 4mn(m+n) total, against a layer forward-plus-backward of roughly 6*B*m*n for B tokens. The overhead ratio is 2(m+n)/(3B): for m = n = 4096 at B ≈ 1e6 that is 0.5%, but at m = n = 2048, B = 4096 it is ~67%. That term, not the eigendecomposition, is what usually kills Shampoo on small batches.

Periodically: accumulate every step, recompute roots only every T steps (50–1000 is typical). Amortized as c(m^3+n^3)/T, that is a fraction of a percent for large layers. The accumulators take m^2 + n^2, so a square layer costs 2m^2 — matching Adam’s two moments, though caching both roots between recomputes doubles that to roughly 4 floats per parameter. A 4096 × 11008 MLP is worse still: 138M accumulator floats against Adam’s 90M. Blocking into b × b tiles fixes both: accumulators return to exactly 2 floats per parameter, and the per-step overhead drops to 4b/(3B).