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.
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 , which measures how sensitive the model's output distribution is to each parameter direction. The update becomes:
By pre-multiplying the gradient by , 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.
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 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 , the Fisher block is an matrix. K-FAC approximates it as the Kronecker product of two much smaller matrices: where is the covariance of layer inputs (activations), and is the covariance of backpropagated gradients.
The Kronecker factorization is equivalent to assuming that the activations and the backpropagated gradients are statistically independent. While not exactly true, this assumption captures the dominant structure: encodes which input features co-activate and 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 of an matrix and an matrix produces an matrix — but you never need to form it explicitly. The critical identity is:
Instead of inverting one matrix (cost ), you invert an matrix and an matrix separately (cost ). For a layer with 512 inputs and 256 outputs, this reduces the inversion from a problem to a plus a problem — a factor of roughly one billion savings.
Computing the K-FAC update step by step
For each layer with weight matrix , the K-FAC update has three stages:
Stage 1 — Accumulate statistics. During training, maintain running averages of the activation covariance and the gradient covariance . 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 and 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 . Using the Kronecker structure, this simplifies to a sandwich: This is just two matrix multiplications — the same computational pattern as a forward pass.
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 to the diagonal before inverting:
But naively adding to 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 :
The scalar is chosen to minimize the approximation error. This preserves the Kronecker structure, so the cheap inversion identity still applies.
The idea in code
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.
The full K-FAC algorithm at a glance
Impact and legacy
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.
2010
Hessian-Free Optimization (Martens)
Used conjugate gradients to approximately multiply by the inverse Hessian without forming it. Effective but slow in stochastic settings.
2015
K-FAC (Martens & Grosse)
Kronecker-factored approximation of the Fisher. Made second-order optimization practical for deep networks with stochastic mini-batch training.
2016
K-FAC for Convolutions (Grosse & Martens)
Extended the Kronecker factorization to convolutional layers, broadening K-FAC's applicability to vision models.
2017
Distributed K-FAC (Ba, Grosse, Martens)
Parallelized K-FAC across multiple GPUs, enabling second-order optimization at scale for ImageNet-level training.
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
- Fisher Information Matrixمصفوفة معلومات فيشر
- Natural Gradientالمُتدرِّج الطبيعي
- Kronecker Productجداء كرونيكر
- Curvatureانحناء
- Second Momentالعزم الثاني
- Preconditionerمُهيِّئ التقارب
- Hessianمصفوفة هيسي
- Block-Diagonalقُطرية كتلية
- Dampingتخميد
- Covariance Matrixمصفوفة التباين المشترك
- Gradient Descentالانحدار التدريجي
- Convergenceالتقارب الحسابي