Amari's Natural Gradient
Natural gradient descent measures distance between probability distributions rather than arbitrary parameter coordinates, following steepest descent on a Riemannian manifold.
Why Does This Exist?
Standard gradient descent defines "steepest descent" using the standard Euclidean distance metric in coordinate space: . When parameters represent physical coordinates on a flat plane, this metric makes total sense.
In modern probabilistic machine learning — including neural networks with softmax classification, policy gradients in reinforcement learning, and variational autoencoders — parameters do not represent physical locations. They define a probability distribution .
Euclidean parameter distance has virtually no relationship with the actual change in the underlying probability distribution. Consider a Gaussian distribution :
- Shifting mean by when barely nudges the distribution; the Kullback-Leibler (KL) divergence is negligible.
- Shifting by that same when produces two non-overlapping spikes; the KL divergence explodes.
Standard gradient descent takes identical step sizes in both scenarios because it is blinded by coordinate Euclidean distance. Even worse, if you change units (e.g., scaling weights or swapping to polar coordinates), the standard gradient changes direction entirely.
Pioneered by mathematician Shun-ichi Amari, Information Geometry models parametric probability distributions as points on a curved Riemannian manifold. Amari's Natural Gradient replaces the arbitrary Euclidean metric with the intrinsic Fisher Information Matrix (FIM), ensuring that optimization steps are invariant to parameterization and strictly bounded by true distributional divergence.
For background on directional derivatives and coordinate slopes, see our guide on gradients.
Think of It Like This
Navigating by flat paper map inches vs. true spherical Earth distances
Imagine piloting an aircraft across Greenland using a flat Mercator projection wall map.
On the flat paper map, one inch measured near the North Pole corresponds to only a few miles of real physical terrain. Near the equator, however, that same one-inch line on the map covers hundreds of miles of ocean.
If your autopilot blindly commands the aircraft to "fly 2 inches on the paper map every hour," your physical airspeed will fluctuate wildly depending on where you are. Near the poles you will crawl at a walking pace; near the equator you will break the sound barrier. The paper grid coordinates distort the underlying physical reality.
Standard gradient descent is that naive autopilot: it measures steps in arbitrary paper inches (parameter coordinates).
Information geometry recognizes that the physical surface of the Earth is a curved sphere with its own intrinsic geometry. Amari's natural gradient calculates compass headings and distances directly on the spherical globe. It ensures that an update step moves the probability distribution by a calibrated physical distance (measured in KL divergence), completely ignoring the distortions of whatever arbitrary coordinate projection you chose.
The analogy stops because probability manifolds generally do not have constant spherical curvature; their curvature varies at every point according to the Fisher information metric.
How It Actually Works
The Fisher Information Metric and Invariant Riemannian Descent
Let be a parametric statistical model. Information geometry treats as a smooth -dimensional Riemannian manifold .
To define distance on this manifold, we measure the dissimilarity between two infinitesimally close distributions and using Kullback-Leibler (KL) divergence:
Taking the second-order Taylor expansion of the KL divergence around yields:
where is the Fisher Information Matrix (FIM), which acts as the Riemannian metric tensor :
Deriving the Natural Gradient
Standard gradient descent defines steepest descent under a Euclidean step constraint: subject to .
Amari defined the Natural Gradient by constraining the step in distribution space via the Fisher metric:
Linearizing the objective and setting up the Lagrangian:
Setting the derivative with respect to to zero:
We define the Natural Gradient vector as:
The natural gradient descent update rule is:
Fundamental Properties
- Reparameterization Invariance: Suppose parameters are mapped bijectively to . The standard gradient directions and do not correspond under the Jacobian transformation. In contrast, the natural gradient transforms contravariantly: the updated distribution sequence is mathematically identical regardless of parameterization!
- Equivalence to Gauss-Newton: When the loss function is the negative log-likelihood , the expected Hessian of the loss is precisely the Fisher Information Matrix: . Thus, natural gradient descent is an optimal, guaranteed positive semi-definite Gauss-Newton step on the statistical manifold.
Worked Example
Consider a 1D Gaussian distribution with unknown mean and fixed known variance :
-
Calculate the score function (gradient of log-likelihood):
-
Calculate the scalar Fisher Information :
The metric tensor is , so its inverse is .
-
Compare two training scenarios with an observed loss gradient :
-
Scenario A (High Certainty, ): The standard gradient step is . The natural gradient step is:
Because the distribution is extremely narrow, the natural gradient shrinks the parameter displacement by a factor of , preventing the distribution from jumping out of overlap!
-
Scenario B (High Uncertainty, ): The standard gradient step is . The natural gradient step is:
Because the distribution is wide and dispersed, shifting the mean by a small amount barely registers in KL divergence. The natural gradient amplifies the step by .
-
In both scenarios, the resulting KL divergence of the step is identical:
The step size is calibrated directly to distributional shift rather than coordinate distance.
Code
import numpy as np
def natural_gradient_gaussian_1d( mu: float, variance: float, loss_grad_mu: float, learning_rate: float = 0.1,) -> tuple[float, float, float]: """Compute standard Euclidean update vs Amari Natural Gradient update.
Returns: (delta_mu_standard, delta_mu_natural, step_kl_divergence) """ # Fisher Information for Gaussian mean is 1 / variance fisher_info = 1.0 / variance inv_fisher = variance
# Standard gradient update: -lr * grad delta_standard = -learning_rate * loss_grad_mu
# Natural gradient update: -lr * (F^-1 * grad) delta_natural = -learning_rate * (inv_fisher * loss_grad_mu)
# KL divergence of natural step: 0.5 * delta^2 * F step_kl = 0.5 * (delta_natural**2) * fisher_info
return delta_standard, delta_natural, step_kl
# Gradient of loss with respect to meang = 2.0lr = 0.1
# Case A: Sharp distribution (variance = 0.04)d_std_a, d_nat_a, kl_a = natural_gradient_gaussian_1d( mu=0.0, variance=0.04, loss_grad_mu=g, learning_rate=lr)print(f"Sharp Gaussian (var=0.04): standard={d_std_a:.3f}, natural={d_nat_a:.4f}, KL={kl_a:.6f}")# -> Sharp Gaussian (var=0.04): standard=-0.200, natural=-0.0080, KL=0.000800
# Case B: Wide distribution (variance = 4.0)d_std_b, d_nat_b, kl_b = natural_gradient_gaussian_1d( mu=0.0, variance=4.0, loss_grad_mu=g, learning_rate=lr)print(f"Wide Gaussian (var=4.0): standard={d_std_b:.3f}, natural={d_nat_b:.4f}, KL={kl_b:.6f}")# -> Wide Gaussian (var=4.0): standard=-0.200, natural=-0.8000, KL=0.080000Watch Out For
Attempting direct inversion of high-dimensional Fisher Information Matrices
For a neural network with weights, the Fisher Information Matrix is of size , containing entries. Storing this matrix in single-precision floating point requires 400 terabytes of memory, and inverting it via exact Cholesky decomposition costs FLOPs.
Writing naive natural gradient code np.linalg.inv(F) @ grad instantly triggers out-of-memory crashes on standard machine learning hardware.
Fix: Use structured approximations of the Fisher matrix. In deep feedforward and convolutional networks, use K-FAC (Kronecker-Factored Approximate Curvature), which factors the layer-wise Fisher matrix into Kronecker products of two small matrices. In reinforcement learning, use TRPO (Trust Region Policy Optimization) or Natural Actor-Critic with conjugate gradient iterations that compute Fisher-vector products through automatic differentiation without ever materializing .
The Quick Version
- Standard gradient descent measures steps using Euclidean coordinate distance, which distorts probability distributions and changes with reparameterization.
- Information geometry defines parametric distributions as points on a Riemannian manifold whose metric tensor is the Fisher Information Matrix .
- Amari's natural gradient update follows the true steepest descent direction measured by Kullback-Leibler divergence.
- When minimizing negative log-likelihood, the natural gradient is exactly equivalent to the Fisher-Gauss-Newton optimization method.
- Direct matrix inversion is intractable for deep networks; practical implementations employ Kronecker-factored curvature (K-FAC) or Hessian-free conjugate gradient methods (TRPO).