Matrix-Free Natural Gradient via Stochastic Truncated Metrics
Ke Sun
Abstract
Natural gradient descent is a theoretical pillar of second-order optimization for neural networks. Existing implementations either rely on structural approximations of the curvature matrix or incur substantial computational and memory overhead. Direct application remains difficult at deep-network scale. We build an optimizer based on an unbiased stochastic estimator of the Fisher information matrix. It employs a truncated exponential moving average metric with damping, which we call the Stochastic Truncated Metric (STM), and is otherwise closely aligned with the original natural gradient. This gives a matrix-free parameter update with strictly $O(Kd)$ time and memory per step, where $K$ is a small constant. It is applicable across network architectures without layer-wise assumptions. We bound the error induced by the finite-memory truncation. We show that STM is stable and efficient on vision and language benchmarks.
Chat is not available.
Successful Page Load