Optimization2015advanced10 min read

Optimizing Neural Networks with Kronecker-Factored Approximate Curvature

أمثَلة الشبكات العصبية بتقريب الانحناء عبر تحليل كرونيكر

Martens, J. · Grosse, R. — ICML

The problem

First-order optimizers like treat the loss landscape as if it curves equally in every direction. In reality, neural network loss surfaces have wildly different curvatures along different directions — some very steep, some nearly flat. This mismatch forces small learning rates (to avoid exploding along steep directions) which means painfully slow progress along flat directions. Second-order methods like fix this by using the to rescale gradients, but the Fisher is an N×N matrix where N is the number of parameters — storing and inverting it is impossible for modern networks with millions of parameters.

The contribution

K-FAC approximates the Fisher information matrix as a matrix (one block per layer), then factors each block into the of two much smaller matrices: the activation covariance A and the gradient covariance G. Because the inverse of a Kronecker product equals the Kronecker product of the inverses, the full approximate can be computed by inverting only small per-layer matrices. The result is only several times more expensive than plain SGD per step, but each step makes dramatically more progress, yielding faster overall .

The impact

K-FAC made second-order practical for deep networks. It directly inspired Shampoo (which uses Kronecker-structured preconditioners without the Fisher interpretation), distributed K-FAC for large-scale , and extensions to CNNs, RNNs, and Transformers. The Kronecker factorization idea became a standard tool in the designer's toolkit, bridging the gap between cheap first-order methods and principled but expensive second-order ones.

Imagine you're hiking down a mountain in thick fog. SGD gives you a compass pointing downhill, but treats every direction as equally easy to walk — so you take the same cautious step size everywhere, even on a wide gentle meadow.

Natural gives you a topographic map of the terrain around you, so you stride confidently across flat meadows and tiptoe along narrow ridges. But the map covers the entire mountain and takes hours to draw.

K-FAC is a shortcut: instead of mapping every square meter, you survey each section of the trail independently, then combine the section maps with a clever folding trick (a Kronecker product). You get 90% of the topographic insight in 1% of the surveying time.

The problem: gradient descent ignores the shape of the landscape

Standard gradient descent computes the steepest-descent direction in Euclidean space — it asks "which small change to the weights reduces the loss the most?" But this question depends on how you measure "small." In Euclidean space, a step of 0.01 in any direction looks the same. In the loss landscape, it doesn't: moving 0.01 along a steep valley wall might launch you out of the valley entirely, while moving 0.01 along the valley floor barely changes the loss.

The of the loss — how fast the gradient itself changes — varies enormously across different parameter directions. First-order methods are blind to this curvature, which forces practitioners to use small learning rates that accommodate the steepest direction, wasting potential progress in all other directions.

Open in Lab
SGD takes identical steps in every direction. The natural gradient adapts to the curvature. Toggle to compare.
The demo wakes as you arrive…

The fix: natural gradient descent

Amari (1998) proposed a better question: instead of "which parameter change of Euclidean length ε reduces the loss the most?", ask "which parameter change that moves the model's output distribution by at most ε reduces the loss most?" This is the natural gradient.

The natural gradient replaces the identity matrix in the gradient update with the Fisher information matrix FF, which measures how sensitive the model's output distribution is to each parameter direction. The update becomes:

θt+1=θt−ηF−1∇L(θt)\theta_{t+1} = \theta_t - \eta F^{-1} \nabla L(\theta_t)

By pre-multiplying the gradient by F−1F^{-1}, the natural gradient stretches the step along flat directions (where the Fisher eigenvalues are small) and compresses it along steep directions (where they are large). The result is reparameterization invariant — the optimization behaves the same regardless of how you parameterize the network.

θt+1=θt−η F−1∇L(θt)\theta_{t+1} = \theta_t - \eta \, F^{-1} \nabla L(\theta_t)
Natural gradient update rule — F = Fisher information matrix · F⁻¹ rescales the gradient to account for curvature · η = learning rate · result: uniform progress across all parameter directions

Think of the Fisher as a ruler that measures distances in "distribution space" rather than "parameter space." Two parameter vectors that are far apart in Euclidean terms but produce nearly identical output distributions should be considered close; two that are nearby in parameters but produce very different distributions should be considered far apart. The Fisher encodes exactly this geometry.

The K-FAC insight: two approximations that make it tractable

K-FAC makes two structural approximations to the Fisher:

Approximation 1 — Block diagonal. Assume parameters in different layers are statistically independent. This turns the giant N×NN \times N Fisher into a block-diagonal matrix where each block covers one layer. Cross-layer interactions are ignored — a reasonable trade-off because most curvature structure is within layers.

Approximation 2 — Kronecker factorization. For a fully-connected layer with W∈Rm×nW \in \mathbb{R}^{m \times n}, the Fisher block is an mn×mnmn \times mn matrix. K-FAC approximates it as the Kronecker product of two much smaller matrices: F^ℓ≈Aℓ−1⊗Gℓ\hat{F}_\ell \approx A_{\ell-1} \otimes G_\ell where Aℓ−1=E[aℓ−1aℓ−1⊤]A_{\ell-1} = \mathbb{E}[a_{\ell-1} a_{\ell-1}^\top] is the n×nn \times n covariance of layer inputs (activations), and Gℓ=E[gℓgℓ⊤]G_\ell = \mathbb{E}[g_\ell g_\ell^\top] is the m×mm \times m covariance of backpropagated gradients.

Open in Lab
The full Fisher (left) vs K-FAC's block-diagonal Kronecker approximation (right). Click a layer block to see its Kronecker factorization.
The demo wakes as you arrive…

The Kronecker factorization is equivalent to assuming that the activations aa and the backpropagated gradients gg are statistically independent. While not exactly true, this assumption captures the dominant structure: AA encodes which input features co-activate and GG encodes which output gradients co-vary. Their Kronecker product naturally models how combinations of input correlations and gradient correlations shape the curvature.

Why Kronecker products are the key: cheap inversion

The Kronecker product A⊗GA \otimes G of an n×nn \times n matrix and an m×mm \times m matrix produces an mn×mnmn \times mn matrix — but you never need to form it explicitly. The critical identity is:

(A⊗G)−1=A−1⊗G−1(A \otimes G)^{-1} = A^{-1} \otimes G^{-1}

Instead of inverting one mn×mnmn \times mn matrix (cost O(m3n3)O(m^3 n^3)), you invert an n×nn \times n matrix and an m×mm \times m matrix separately (cost O(n3+m3)O(n^3 + m^3)). For a layer with 512 inputs and 256 outputs, this reduces the inversion from a 131,072×131,072131{,}072 \times 131{,}072 problem to a 512×512512 \times 512 plus a 256×256256 \times 256 problem — a factor of roughly one billion savings.

(A⊗G)−1=A−1⊗G−1(A \otimes G)^{-1} = A^{-1} \otimes G^{-1}
The Kronecker inversion identity — K-FAC's engine — Inverting the Kronecker product = inverting each factor separately · this is what makes K-FAC's per-step cost manageable
Open in Lab
See how two small matrices combine into one large Kronecker product. Drag the slider to change matrix size and compare inversion costs.
The demo wakes as you arrive…

Computing the K-FAC update step by step

For each layer ℓ\ell with weight matrix WℓW_\ell, the K-FAC update has three stages:

Stage 1 — Accumulate statistics. During training, maintain running averages of the activation covariance Aℓ−1=E[aℓ−1aℓ−1⊤]A_{\ell-1} = \mathbb{E}[a_{\ell-1} a_{\ell-1}^\top] and the gradient covariance Gℓ=E[gℓgℓ⊤]G_\ell = \mathbb{E}[g_\ell g_\ell^\top]. These are computed from the same forward and backward passes you already do for SGD — almost no extra work.

Stage 2 — Invert the factors. Periodically (not every step), invert Aℓ−1A_{\ell-1} and GℓG_\ell separately. Each is a small matrix, so this is cheap and can run asynchronously on CPU while the GPU does the next .

Stage 3 — Apply the preconditioned gradient. The natural gradient update for the vectorized weight is F−1vec(∇L)F^{-1} \text{vec}(\nabla L). Using the Kronecker structure, this simplifies to a sandwich: ΔWℓ=Gℓ−1 (∇WℓL) Aℓ−1−1\Delta W_\ell = G_\ell^{-1} \, (\nabla_{W_\ell} L) \, A_{\ell-1}^{-1} This is just two matrix multiplications — the same computational pattern as a forward pass.

ΔWℓ=Gℓ−1 (∇WℓL) Aℓ−1−1\Delta W_\ell = G_\ell^{-1} \, (\nabla_{W_\ell} L) \, A_{\ell-1}^{-1}
K-FAC update — the preconditioned gradient as a matrix sandwich — G⁻¹ rescales along the output dimension · A⁻¹ rescales along the input dimension · the gradient is "squeezed" between two curvature-aware transforms
Open in Lab
Step through the three stages of a K-FAC update for one layer.
The demo wakes as you arrive…

Damping: keeping the approximation stable

The Fisher approximation isn't perfect, and the Kronecker factors can have very small eigenvalues, leading to dangerously large update steps. K-FAC uses Tikhonov — adding a small constant λ\lambda to the diagonal before inverting:

(F^ℓ+λI)−1(\hat{F}_\ell + \lambda I)^{-1}

But naively adding λI\lambda I to A⊗GA \otimes G destroys the Kronecker structure and makes inversion expensive again. K-FAC's clever solution is factored Tikhonov damping: distribute the damping across the two factors using a balancing scalar π\pi:

(A+πλ I)−1⊗(G+1πλ I)−1(A + \pi\sqrt{\lambda}\,I)^{-1} \otimes (G + \frac{1}{\pi}\sqrt{\lambda}\,I)^{-1}

The scalar π\pi is chosen to minimize the approximation error. This preserves the Kronecker structure, so the cheap inversion identity still applies.

The idea in code

K-FAC update for one fully-connected layerpython

Simplified to show the idea — not the real implementation.

import numpy as np

def kfac_update(grad_W, A_cov, G_cov, damping=1e-3):
    """
    K-FAC preconditioned gradient for one FC layer.

    grad_W : (m, n) — gradient of loss w.r.t. weight matrix W
    A_cov  : (n, n) — running average of input covariance E[a aᵀ]
    G_cov  : (m, m) — running average of gradient covariance E[g gᵀ]
    damping: scalar λ for Tikhonov regularization
    """
    n = A_cov.shape[0]
    m = G_cov.shape[0]

    # Factored Tikhonov damping: balance λ across factors
    pi = np.sqrt((np.trace(A_cov) / n) / (np.trace(G_cov) / m))
    A_damped = A_cov + pi * np.sqrt(damping) * np.eye(n)
    G_damped = G_cov + (1/pi) * np.sqrt(damping) * np.eye(m)

    # Invert each small factor separately — this is the key trick
    A_inv = np.linalg.inv(A_damped)   # (n, n)
    G_inv = np.linalg.inv(G_damped)   # (m, m)

    # The "sandwich": preconditioned gradient = G⁻¹ · ∇W · A⁻¹
    natural_grad = G_inv @ grad_W @ A_inv
    return natural_grad

# Usage in training loop:
# 1. Forward pass: record activations a for each layer
# 2. Backward pass: record gradients g for each layer
# 3. Update running averages: A ← β·A + (1-β)·a·aᵀ
#                             G ← β·G + (1-β)·g·gᵀ
# 4. Periodically recompute A_inv, G_inv (can be async on CPU)
# 5. W ← W - η · kfac_update(∇W, A, G, λ)

Convergence: fewer steps, bigger strides

The paper's experiments on deep autoencoders and convolutional networks showed K-FAC converging in significantly fewer iterations than SGD with — often reaching the same loss in 3–10× fewer updates. While each K-FAC step costs several times more than an SGD step (due to the covariance estimation and periodic inversions), the net effect is faster wall-clock convergence.

Unlike -free optimization — the previous best second-order method — K-FAC works well in highly stochastic settings with small mini-batches. The running-average covariance estimates naturally smooth out noise, and the factored damping prevents the approximation from becoming unstable.

Open in Lab
Compare convergence curves of SGD, Adam, and K-FAC on a deep autoencoder.
The demo wakes as you arrive…

The full K-FAC algorithm at a glance

Open in Lab
Walk through the complete K-FAC algorithm one phase at a time.
The demo wakes as you arrive…

Impact and legacy

  1. 1998

    Natural Gradient Descent (Amari)

    Introduced the idea of using the Fisher information matrix to rescale gradients according to the geometry of the output distribution space.

  2. 2010

    Hessian-Free Optimization (Martens)

    Used conjugate gradients to approximately multiply by the inverse Hessian without forming it. Effective but slow in stochastic settings.

  3. 2015

    K-FAC (Martens & Grosse)

    Kronecker-factored approximation of the Fisher. Made second-order optimization practical for deep networks with stochastic mini-batch training.

  4. 2016

    K-FAC for Convolutions (Grosse & Martens)

    Extended the Kronecker factorization to convolutional layers, broadening K-FAC's applicability to vision models.

  5. 2017

    Distributed K-FAC (Ba, Grosse, Martens)

    Parallelized K-FAC across multiple GPUs, enabling second-order optimization at scale for ImageNet-level training.

  6. 2018

    Shampoo (Gupta, Koren, Singer)

    Used Kronecker-structured preconditioners with matrix root operations. Simpler than K-FAC (no Fisher interpretation needed), and later scaled to production at Google.

CitationMartens, Grosse. Optimizing Neural Networks with Kronecker-factored Approximate Curvature. ICML, 2015.

Terms in this paper