Optimizing Neural Networks with Kronecker-factored Approximate Curvature

James MartensRoger Grosse

article2015ICML1,459 citations

Introduces Kronecker-Factored Approximate Curvature (K-FAC), a scalable second-order optimization algorithm that makes natural gradient descent practical for deep neural networks by tractably approximating the Fisher information matrix using layer-wise Kronecker products.

Listen

Training deep neural networks is central to modern machine learning, yet standard optimization methods present significant trade-offs. Standard first-order methods like stochastic gradient descent (SGD) are computationally cheap per step but require many thousands of iterations to converge on complex tasks. Conversely, traditional second-order methods incorporate curvature information to make substantial progress per iteration, but their updates are computationally expensive and degrade in noisy, mini-batch training regimes.

The article develops and evaluates Kronecker-factored Approximate Curvature (K-FAC), a second-order optimization method designed to approximate natural gradient descent efficiently in deep feed-forward neural networks. The objective is to provide powerful curvature-corrected parameter updates at a computational cost per iteration that remains close to standard stochastic gradient methods.

To achieve this, the approach approximates the network's Fisher information matrix by structuring it into layer-wise blocks, factoring each block into the Kronecker product of two smaller matrices, and approximating the inverse matrix as either block-diagonal or block-tridiagonal. Curvature statistics are accumulated online across mini-batches using exponentially decaying averages, decoupling storage and inversion costs from the sample size. The algorithm incorporates adaptive damping techniques and a parameter-free momentum scheme to stabilize step sizes and improve local quadratic optimization. The authors evaluated K-FAC against a well-tuned SGD baseline with Nesterov momentum across three standard deep autoencoder benchmark problems (MNIST, CURVES, and FACES) on a single GPU.

The evaluation yielded several key findings. First, K-FAC achieved per-iteration progress rates that were orders of magnitude higher than SGD with momentum, reducing the total iterations required for convergence from tens of thousands to a few hundred. Second, when using an exponentially increasing mini-batch schedule, K-FAC delivered substantially faster overall wall-clock training times than the baseline. Third, the block-tridiagonal formulation improved per-iteration progress by approximately 25% to 40% compared to the simpler block-diagonal version. Fourth, the method exhibited strong mathematical invariance properties, ensuring optimization behavior remains consistent under affine reparameterizations and changes in network activation functions.

These findings demonstrate that second-order curvature information can be applied practically and efficiently to neural network training without relying on expensive iterative sub-solvers. For organizations training deep networks, K-FAC can dramatically compress training timelines and improve computational efficiency. Because K-FAC achieves massive per-step progress and requires far fewer parameter updates, it is particularly well-suited for distributed computing architectures where network communication and synchronization across nodes represent the primary performance bottleneck.

Organizations evaluating advanced optimization frameworks should consider implementing K-FAC for compute-heavy neural network workloads, prioritizing the block-diagonal variant for general use due to its simpler implementation and competitive per-second throughput. Moving forward, engineering efforts should explore parallelized implementations that compute matrix operations asynchronously across layers. Further research is recommended to expand K-FAC approximations to convolutional and recurrent network architectures and to develop adaptive mini-batch sizing strategies.

Confidence in these results is supported by rigorous mathematical derivations and consistent empirical gains across multiple recognized benchmark problems. However, decision-makers should note that the empirical evaluations focused specifically on deep autoencoder architectures on a single computer system. Application to broader model types, such as modern vision or language models, requires tailored factorizations, and the performance advantages rely on maintaining appropriate damping parameters and sufficiently large mini-batch sizes.

arXiv: 1503.05671
Cover for Optimizing Neural Networks with Kronecker-factored Approximate Curvature

Abstract

We propose an efficient method for approximating natural gradient descent in neural networks which we call Kronecker-Factored Approximate Curvature (K-FAC). K-FAC is based on an efficiently invertible approximation of a neural network's Fisher information matrix which is neither diagonal nor low-rank, and in some cases is completely non-sparse. It is derived by approximating various large blocks of the Fisher (corresponding to entire layers) as being the Kronecker product of two much smaller matrices. While only several times more expensive to compute than the plain stochastic gradient, the updates produced by K-FAC make much more progress optimizing the objective, which results in an algorithm that can be much faster than stochastic gradient descent with momentum in practice. And unlike some previously proposed approximate natural-gradient/Newton methods which use high-quality non-diagonal curvature matrices (such as Hessian-free optimization), K-FAC works very well in highly stochastic optimization regimes. This is because the cost of storing and inverting K-FAC's approximation to the curvature matrix does not depend on the amount of data used to estimate it, which is a feature typically associated only with diagonal or low-rank approximations to the curvature matrix.

Table of Contents

  • 1 Introduction
  • 2 Background and notation
  • 2.1 Neural Networks
  • 2.2 The Natural Gradient
  • 3 A block-wise Kronecker-factored Fisher approximation
  • 3.1 Interpretations of this approximation
  • 4 Additional approximations to F~\tilde{F} and inverse computations
  • 4.1 Structured inverses and the connection to linear regression
  • 4.2 Approximating F~−1\tilde{F}^{-1} as block-diagonal
  • 4.3 Approximating F~−1\tilde{F}^{-1} as block-tridiagonal
  • 4.4 Examining the approximation quality
  • 5 Estimating the required statistics
  • 6 Update damping
  • 6.1 Background and motivation
  • 6.2 A highly effective damping scheme for K-FAC
  • 6.3 A factored Tikhonov regularization technique
  • 6.4 Re-scaling according to the exact FF
  • 6.5 Adapting λ\lambda
  • 6.6 Maintaining a separate damping strength for the approximate Fisher
  • 7 Momentum
  • 8 Computational Costs and Efficiency Improvements
  • 9 Pseudocode for K-FAC
  • 10 Invariance Properties and the Relationship to Whitening and Centering
  • 11 Related Work
  • 12 Heskes’ interpretation of the block-diagonal approximation
  • 13 Experiments
  • 14 Conclusions and future directions
  • References
  • A Derivation of the expression for the approximation from Section
  • B Efficient techniques for inverting A⊗B±C⊗DA\otimes B\pm C\otimes D
  • C Computing v⊤​F​vv^{\top}Fv and u⊤​F​vu^{\top}Fv more efficiently
  • D Proofs for Section

Knowls

  1. Knowl 1 — Kronecker-Factored Fisher Approximation

    model/method

    In a feed-forward neural network with ℓ\ell layers, let layer i∈{1,…,ℓ}i \in \{1, \dots, \ell\} compute pre-activations si=Wiaˉi−1s_i = W_i \bar{a}_{i-1} and activations ai=ϕi(si)a_i = \phi_i(s_i), where Wi∈Rdi×(di−1+1)W_i \in \mathbb{R}^{d_i \times (d_{i-1}+1)} is the weight matrix (including bias in its final column) and aˉi−1=[ai−1⊤,1]⊤∈Rdi−1+1\bar{a}_{i-1} = [a_{i-1}^\top, 1]^\top \in \mathbb{R}^{d_{i-1}+1} is the activity vector augmented with a homogeneous coordinate. The unit-gradient vector is gi=Dsi=−dlog⁡p(y∣x,θ)dsi∈Rdig_i = \mathcal{D}s_i = -\frac{d \log p(y \mid x, \theta)}{d s_i} \in \mathbb{R}^{d_i}.

    The derivative of the loss with respect to WiW_i is DWi=giaˉi−1⊤DW_i = g_i \bar{a}_{i-1}^\top, so in vectorized form vec(DWi)=aˉi−1⊗gi\text{vec}(DW_i) = \bar{a}_{i-1} \otimes g_i, where ⊗\otimes denotes the Kronecker product. The Fisher information matrix F=E[vec(Dθ)vec(Dθ)⊤]F = \mathbb{E}[\text{vec}(D\theta) \text{vec}(D\theta)^\top] partitions into ℓ×ℓ\ell \times \ell blocks, where the (i,j)(i, j)-th block is: Fi,j=E[vec(DWi)vec(DWj)⊤]=E[aˉi−1aˉj−1⊤⊗gigj⊤]F_{i,j} = \mathbb{E}\left[\text{vec}(DW_i) \text{vec}(DW_j)^\top\right] = \mathbb{E}\left[\bar{a}_{i-1}\bar{a}_{j-1}^\top \otimes g_i g_j^\top\right]

    K-FAC defines an initial approximation F~\tilde{F} of the Fisher by approximating the expectation of each Kronecker product as the Kronecker product of expectations: F~i,j=Aˉi−1,j−1⊗Gi,j\tilde{F}_{i,j} = \bar{A}_{i-1,j-1} \otimes G_{i,j} where Aˉi−1,j−1=E[aˉi−1aˉj−1⊤]\bar{A}_{i-1,j-1} = \mathbb{E}\left[\bar{a}_{i-1}\bar{a}_{j-1}^\top\right] and Gi,j=E[gigj⊤]G_{i,j} = \mathbb{E}\left[g_i g_j^\top\right]. This approximation is equivalent to assuming statistical independence between products of unit activities aˉi−1aˉj−1⊤\bar{a}_{i-1}\bar{a}_{j-1}^\top and products of unit derivatives gigj⊤g_i g_j^\top.

  2. Knowl 2 — Block-Diagonal Kronecker-Factored Fisher Inversion

    model/method

    The block-diagonal approximation Fˇ\check{F} of the Kronecker-factored Fisher F~\tilde{F} retains only the diagonal blocks corresponding to individual layers: Fˇ=diag(F~1,1,F~2,2,…,F~ℓ,ℓ)=diag(Aˉ0,0⊗G1,1,Aˉ1,1⊗G2,2,…,Aˉℓ−1,ℓ−1⊗Gℓ,ℓ)\check{F} = \text{diag}(\tilde{F}_{1,1}, \tilde{F}_{2,2}, \dots, \tilde{F}_{\ell,\ell}) = \text{diag}(\bar{A}_{0,0} \otimes G_{1,1}, \bar{A}_{1,1} \otimes G_{2,2}, \dots, \bar{A}_{\ell-1,\ell-1} \otimes G_{\ell,\ell}) where Aˉi−1,i−1=E[aˉi−1aˉi−1⊤]∈R(di−1+1)×(di−1+1)\bar{A}_{i-1,i-1} = \mathbb{E}\left[\bar{a}_{i-1}\bar{a}_{i-1}^\top\right] \in \mathbb{R}^{(d_{i-1}+1) \times (d_{i-1}+1)} and Gi,i=E[gigi⊤]∈Rdi×diG_{i,i} = \mathbb{E}\left[g_i g_i^\top\right] \in \mathbb{R}^{d_i \times d_i}.

    Using the identity (A⊗B)−1=A−1⊗B−1(A \otimes B)^{-1} = A^{-1} \otimes B^{-1}, the inverse is block-diagonal: Fˇ−1=diag(Aˉ0,0−1⊗G1,1−1,Aˉ1,1−1⊗G2,2−1,…,Aˉℓ−1,ℓ−1−1⊗Gℓ,ℓ−1)\check{F}^{-1} = \text{diag}\left(\bar{A}_{0,0}^{-1} \otimes G_{1,1}^{-1}, \bar{A}_{1,1}^{-1} \otimes G_{2,2}^{-1}, \dots, \bar{A}_{\ell-1,\ell-1}^{-1} \otimes G_{\ell,\ell}^{-1}\right)

    Applying Fˇ−1\check{F}^{-1} to a gradient vector vv, whose layer blocks correspond to gradient matrices Vi=∇Wih(θ)V_i = \nabla_{W_i} h(\theta), yields update blocks UiU_i (such that vec(Ui)=(Aˉi−1,i−1−1⊗Gi,i−1)vec(Vi)\text{vec}(U_i) = (\bar{A}_{i-1,i-1}^{-1} \otimes G_{i,i}^{-1})\text{vec}(V_i)) via the matrix equation: Ui=Gi,i−1ViAˉi−1,i−1−1U_i = G_{i,i}^{-1} V_i \bar{A}_{i-1,i-1}^{-1} This eliminates large matrix operations by requiring the inversion of only 2ℓ2\ell small factor matrices.

  3. Knowl 3 — Block-Tridiagonal Inverse Fisher Approximation via Directed Gaussian Graphical Models

    model/method

    Approximating the inverse Fisher F~−1\tilde{F}^{-1} as block-tridiagonal corresponds to modeling the joint distribution of layer gradient vectors vec(DWi)\text{vec}(DW_i) as a Markov chain across layers. The resulting approximation F^−1\hat{F}^{-1} has the generalized Cholesky decomposition: F^−1=Ξ⊤ΛΞ\hat{F}^{-1} = \Xi^\top \Lambda \Xi where Λ=diag(Σ1∣2−1,Σ2∣3−1,…,Σℓ−1∣ℓ−1,Σℓ−1)\Lambda = \text{diag}\left(\Sigma_{1 \mid 2}^{-1}, \Sigma_{2 \mid 3}^{-1}, \dots, \Sigma_{\ell-1 \mid \ell}^{-1}, \Sigma_\ell^{-1}\right) and Ξ\Xi is an upper block-bidiagonal matrix: Ξ=[I−Ψ1,20⋯00I−Ψ2,3⋯0⋮⋮⋱⋱⋮00⋯I−Ψℓ−1,ℓ00⋯0I]\Xi = \begin{bmatrix} I & -\Psi_{1,2} & 0 & \cdots & 0 \\ 0 & I & -\Psi_{2,3} & \cdots & 0 \\ \vdots & \vdots & \ddots & \ddots & \vdots \\ 0 & 0 & \cdots & I & -\Psi_{\ell-1,\ell} \\ 0 & 0 & \cdots & 0 & I \end{bmatrix}

    The transition matrices factorize as Ψi,i+1=ΨAˉ,i−1,i⊗ΨG,i,i+1\Psi_{i,i+1} = \Psi_{\bar{A}, i-1, i} \otimes \Psi_{G, i, i+1}, where ΨAˉ,i−1,i=Aˉi−1,iAˉi,i−1\Psi_{\bar{A}, i-1, i} = \bar{A}_{i-1,i}\bar{A}_{i,i}^{-1} and ΨG,i,i+1=Gi,i+1Gi+1,i+1−1\Psi_{G, i, i+1} = G_{i, i+1}G_{i+1,i+1}^{-1}.

    The conditional covariances are Σℓ=Aˉℓ−1,ℓ−1⊗Gℓ,ℓ\Sigma_\ell = \bar{A}_{\ell-1,\ell-1} \otimes G_{\ell,\ell} and for i≤ℓ−1i \le \ell - 1: Σi∣i+1=Aˉi−1,i−1⊗Gi,i−(ΨAˉ,i−1,iAˉi,iΨAˉ,i−1,i⊤)⊗(ΨG,i,i+1Gi+1,i+1ΨG,i,i+1⊤)\Sigma_{i \mid i+1} = \bar{A}_{i-1,i-1} \otimes G_{i,i} - \left(\Psi_{\bar{A}, i-1, i} \bar{A}_{i,i} \Psi_{\bar{A}, i-1, i}^\top\right) \otimes \left(\Psi_{G, i, i+1} G_{i+1,i+1} \Psi_{G, i, i+1}^\top\right)

    Matrix-vector multiplication u=F^−1vu = \hat{F}^{-1} v with gradient blocks ViV_i proceeds in three steps:

    1. u(1)=Ξvu^{(1)} = \Xi v, with Ui(1)=Vi−ΨG,i,i+1Vi+1ΨAˉ,i−1,i⊤U_i^{(1)} = V_i - \Psi_{G, i, i+1} V_{i+1} \Psi_{\bar{A}, i-1, i}^\top for i<ℓi < \ell, and Uℓ(1)=VℓU_\ell^{(1)} = V_\ell.
    2. u(2)=Λu(1)u^{(2)} = \Lambda u^{(1)}, with vec(Ui(2))=Σi∣i+1−1vec(Ui(1))\text{vec}(U_i^{(2)}) = \Sigma_{i \mid i+1}^{-1} \text{vec}(U_i^{(1)}).
    3. u=Ξ⊤u(2)u = \Xi^\top u^{(2)}, with Ui=Ui(2)−ΨG,i−1,i⊤Ui−1(2)ΨAˉ,i−2,i−1U_i = U_i^{(2)} - \Psi_{G, i-1, i}^\top U_{i-1}^{(2)} \Psi_{\bar{A}, i-2, i-1} for i>1i > 1, and U1=U1(2)U_1 = U_1^{(2)}.
  4. Knowl 4 — Factored Tikhonov Regularization for Kronecker Factors

    model/method

    Standard Tikhonov damping on a diagonal block adds (λ+η)I(\lambda + \eta)I to Aˉi−1,i−1⊗Gi,i\bar{A}_{i-1,i-1} \otimes G_{i,i}, which breaks Kronecker product inversion. Factored Tikhonov regularization replaces this by damping the individual Kronecker factors: (Aˉi−1,i−1+πiγI)⊗(Gi,i+1πiγI)\left(\bar{A}_{i-1,i-1} + \pi_i \gamma I\right) \otimes \left(G_{i,i} + \frac{1}{\pi_i} \gamma I\right) where γ=λ+η\gamma = \sqrt{\lambda + \eta}, λ\lambda is the Tikhonov damping parameter, and η\eta is the L2L_2 weight-decay coefficient. Expanding this product yields Aˉi−1,i−1⊗Gi,i+(λ+η)I⊗I\bar{A}_{i-1,i-1} \otimes G_{i,i} + (\lambda + \eta) I \otimes I plus a residual error term πiγI⊗Gi,i+πi−1γAˉi−1,i−1⊗I\pi_i \gamma I \otimes G_{i,i} + \pi_i^{-1} \gamma \bar{A}_{i-1,i-1} \otimes I.

    To balance the norms and minimize the trace-norm upper bound of the residual error, the scalar constant πi\pi_i is chosen as: πi=tr(Aˉi−1,i−1)/(di−1+1)tr(Gi,i)/di\pi_i = \sqrt{\frac{\text{tr}(\bar{A}_{i-1,i-1}) / (d_{i-1} + 1)}{\text{tr}(G_{i,i}) / d_i}} where did_i is the number of units in layer ii and di−1+1d_{i-1}+1 is the input dimension including bias.

  5. Knowl 5 — Exact-Fisher Quadratic Model Re-scaling and Levenberg-Marquardt Damping Adaptation

    model/method

    Given a proposed update direction Δ\Delta computed using the approximate inverse Fisher (with factored Tikhonov damping), K-FAC re-scales Δ\Delta to produce the step δ=α∗Δ\delta = \alpha^* \Delta by minimizing the quadratic Taylor model M(δ)=12δ⊤(F+(λ+η)I)δ+∇h(θ)⊤δ+h(θ)M(\delta) = \frac{1}{2} \delta^\top (F + (\lambda + \eta)I) \delta + \nabla h(\theta)^\top \delta + h(\theta) evaluated on the current mini-batch using the exact Fisher FF: α∗=−∇h(θ)⊤ΔΔ⊤FΔ+(λ+η)∥Δ∥22\alpha^* = \frac{-\nabla h(\theta)^\top \Delta}{\Delta^\top F \Delta + (\lambda + \eta) \|\Delta\|_2^2}

    The damping parameter λ\lambda is adapted every T1T_1 iterations (with T1=5T_1 = 5, λ0=150\lambda_0 = 150) based on the reduction ratio ρ=h(θ+δ)−h(θ)M(δ)−M(0)=h(θ+δ)−h(θ)12∇h(θ)⊤δ\rho = \frac{h(\theta + \delta) - h(\theta)}{M(\delta) - M(0)} = \frac{h(\theta + \delta) - h(\theta)}{\frac{1}{2}\nabla h(\theta)^\top \delta}:

    • If ρ>3/4\rho > 3/4, λ←ω1λ\lambda \leftarrow \omega_1 \lambda.
    • If ρ<1/4\rho < 1/4, λ←λ/ω1\lambda \leftarrow \lambda / \omega_1, where ω1=(19/20)T1\omega_1 = (19/20)^{T_1}.

    The factor damping parameter γ\gamma in the factored Tikhonov regularization is maintained separately from λ+η\sqrt{\lambda + \eta} and updated greedily every T2T_2 iterations (T2=20T_2 = 20) by comparing candidates {γ0,ω2γ0,γ0/ω2}\{\gamma_0, \omega_2 \gamma_0, \gamma_0 / \omega_2\} (where ω2=(19/20)T2\omega_2 = (\sqrt{19/20})^{T_2}) and picking the one that minimizes M(α∗Δ)M(\alpha^* \Delta).

  6. Knowl 6 — Parameter-Free Exact-Curvature Momentum

    model/method

    K-FAC incorporates momentum by formulating the update as a linear combination of the current approximate natural gradient proposal Δ\Delta and the previous parameter update δ0\delta_0: δ=αΔ+μδ0\delta = \alpha \Delta + \mu \delta_0 where the learning rate α\alpha and momentum decay μ\mu are chosen jointly to minimize the quadratic model M(δ)=12δ⊤(F+(λ+η)I)δ+∇h(θ)⊤δ+h(θ)M(\delta) = \frac{1}{2} \delta^\top (F + (\lambda + \eta)I)\delta + \nabla h(\theta)^\top \delta + h(\theta) using the exact Fisher FF.

    The optimal parameters [α∗,μ∗]⊤[\alpha^*, \mu^*]^\top are given in closed form by: [α∗μ∗]=−[Δ⊤FΔ+(λ+η)∥Δ∥22Δ⊤Fδ0+(λ+η)Δ⊤δ0Δ⊤Fδ0+(λ+η)Δ⊤δ0δ0⊤Fδ0+(λ+η)∥δ0∥22]−1[∇h(θ)⊤Δ∇h(θ)⊤δ0]\begin{bmatrix} \alpha^* \\ \mu^* \end{bmatrix} = - \begin{bmatrix} \Delta^\top F \Delta + (\lambda + \eta)\|\Delta\|_2^2 & \Delta^\top F \delta_0 + (\lambda + \eta)\Delta^\top \delta_0 \\ \Delta^\top F \delta_0 + (\lambda + \eta)\Delta^\top \delta_0 & \delta_0^\top F \delta_0 + (\lambda + \eta)\|\delta_0\|_2^2 \end{bmatrix}^{-1} \begin{bmatrix} \nabla h(\theta)^\top \Delta \\ \nabla h(\theta)^\top \delta_0 \end{bmatrix}

    This automatically tunes the momentum decay parameter μ\mu at each iteration without manual decay schedules and is mathematically equivalent to preconditioned linear Conjugate Gradient when minimizing a fixed quadratic objective.

  7. Knowl 7 — Online Fisher Statistics Estimation via Predictive Distribution Sampling

    model/method

    The Kronecker factor matrices are defined by expectations over data and model distributions: Aˉi,j=Ex∼Q^x[aˉiaˉj⊤]\bar{A}_{i,j} = \mathbb{E}_{x \sim \hat{Q}_x}[\bar{a}_i \bar{a}_j^\top] and Gi,j=Ex∼Q^x,y∼Py∣x(θ)[gigj⊤]G_{i,j} = \mathbb{E}_{x \sim \hat{Q}_x, y \sim P_{y \mid x}(\theta)}[g_i g_j^\top].

    1. Activations aˉi\bar{a}_i are collected from the standard forward pass on mini-batch inputs xx.
    2. Unit gradients gi=−dlog⁡p(y∣x,θ)dsig_i = -\frac{d \log p(y \mid x, \theta)}{ds_i} require taking the expectation with respect to the network's predictive model distribution Py∣x(θ)P_{y \mid x}(\theta), rather than the empirical data targets y^data\hat{y}_{\text{data}}. K-FAC computes an unbiased Monte-Carlo estimate by sampling pseudo-targets y^∼Py∣x(θ)\hat{y} \sim P_{y \mid x}(\theta) and running an additional backward pass using y^\hat{y} as the target.
    3. Running estimates are updated across iterations kk via an exponentially decaying moving average: Aˉi,j(k)=ϵAˉi,j(k−1)+(1−ϵ)Aˉi,j(batch),Gi,j(k)=ϵGi,j(k−1)+(1−ϵ)Gi,j(batch)\bar{A}_{i,j}^{(k)} = \epsilon \bar{A}_{i,j}^{(k-1)} + (1 - \epsilon) \bar{A}_{i,j}^{(\text{batch})}, \quad G_{i,j}^{(k)} = \epsilon G_{i,j}^{(k-1)} + (1 - \epsilon) G_{i,j}^{(\text{batch})} where ϵ=min⁡{1−1/k,0.95}\epsilon = \min\{1 - 1/k, 0.95\}.
  8. Knowl 8 — Network Transformation Invariance and Equivalence to Gradient Whitening

    theoretical result

    Let a neural network undergo invertible affine transformations of pre-activations and activations across layers: si†=Wi†aˉi−1†,aˉi†=Ωiϕˉi(Φisi†)s_i^\dagger = W_i^\dagger \bar{a}_{i-1}^\dagger, \quad \bar{a}_i^\dagger = \Omega_i \bar{\phi}_i(\Phi_i s_i^\dagger) where ϕˉi\bar{\phi}_i computes ϕi\phi_i and appends a homogeneous coordinate 1, Ωi\Omega_i and Φi\Phi_i are arbitrary invertible matrices with Ωℓ=I\Omega_\ell = I, and aˉ0†=Ω0aˉ0\bar{a}_0^\dagger = \Omega_0 \bar{a}_0.

    Theorem: There exists an invertible linear reparameterization θ=ζ(θ†)\theta = \zeta(\theta^\dagger) such that f†(x,θ†)=f(x,θ)f^\dagger(x, \theta^\dagger) = f(x, \theta). Updating θ\theta by δ=−αFˇ−1∇h\delta = -\alpha \check{F}^{-1} \nabla h (block-diagonal) or δ=−αF^−1∇h\delta = -\alpha \hat{F}^{-1} \nabla h (block-tridiagonal) in the original network is exactly equivalent to updating θ†\theta^\dagger by δ†=−α(Fˇ†)−1∇h†\delta^\dagger = -\alpha (\check{F}^\dagger)^{-1} \nabla h^\dagger or δ†=−α(F^†)−1∇h†\delta^\dagger = -\alpha (\hat{F}^\dagger)^{-1} \nabla h^\dagger in the transformed network, in the sense that ζ(θ†+δ†)=θ+δ\zeta(\theta^\dagger + \delta^\dagger) = \theta + \delta.

    Corollary: An undamped block-diagonal K-FAC update step −αFˇ−1∇h-\alpha \check{F}^{-1} \nabla h is algebraically equivalent to standard gradient descent −α∇h†-\alpha \nabla h^\dagger on a reparameterized network whose unit activations ai†a_i^\dagger and unit gradients gi†g_i^\dagger are centered and whitened w.r.t. the model distribution (i.e. Gi,i†=IG_{i,i}^\dagger = I and Aˉi,i†=I\bar{A}_{i,i}^\dagger = I).

  9. Knowl 9 — Fast Exact Fisher Quadratic Products and Low-Rank Matrix Inverse Multiplication

    model/method

    To minimize computational overhead, K-FAC uses two algebraic shortcuts:

    1. Fast Exact Fisher Quadratic Forms: The Fisher matrix is F=EQ^x[J⊤FRJ]F = \mathbb{E}_{\hat{Q}_x}[J^\top F_R J], where J=∂f(x,θ)∂θJ = \frac{\partial f(x, \theta)}{\partial \theta} is the network output Jacobian and FRF_R is the metric tensor of the output predictive distribution. Computing scalar quadratic forms Δ⊤FΔ\Delta^\top F \Delta, Δ⊤Fδ0\Delta^\top F \delta_0, and δ0⊤Fδ0\delta_0^\top F \delta_0 does not require full matrix-vector products FΔ=J⊤FRJΔF \Delta = J^\top F_R J \Delta. Instead, evaluating only the forward linearized map JΔJ \Delta allows computing the scalar directly via (JΔ)⊤FR(JΔ)(J \Delta)^\top F_R (J \Delta), cutting the computational cost by half.

    2. Low-Rank Inversion for Small Mini-Batches: When the mini-batch size mm is smaller than layer width dd, the empirical gradient ∇Wih=1mGiAˉi−1⊤\nabla_{W_i} h = \frac{1}{m} \mathcal{G}_i \bar{\mathcal{A}}_{i-1}^\top is of rank mm (where Gi∈Rdi×m\mathcal{G}_i \in \mathbb{R}^{d_i \times m} and Aˉi−1∈R(di−1+1)×m\bar{\mathcal{A}}_{i-1} \in \mathbb{R}^{(d_{i-1}+1) \times m}). The block-diagonal inverse update Ui=Gi,i−1(∇Wih)Aˉi−1,i−1−1U_i = G_{i,i}^{-1} (\nabla_{W_i} h) \bar{A}_{i-1,i-1}^{-1} is evaluated as: Ui=1m(Gi,i−1Gi)(Aˉi−1⊤Aˉi−1,i−1−1)U_i = \frac{1}{m} \left(G_{i,i}^{-1} \mathcal{G}_i\right) \left(\bar{\mathcal{A}}_{i-1}^\top \bar{A}_{i-1,i-1}^{-1}\right) reducing the per-iteration matrix multiplication cost from O(d3)O(d^3) to O(d2m)O(d^2 m).

  10. Knowl 10 — The K-FAC Optimization Algorithm

    algorithm

    K-FAC performs approximate natural gradient updates by maintaining running estimates of Kronecker factors Aˉi,j\bar{A}_{i,j} and Gi,jG_{i,j}, computing their damped inverses, evaluating the candidate natural gradient step Δ\Delta, and solving for the exact-Fisher damped scaling and momentum coefficients.

    Input: Training set S, learning objective h(theta), initial parameters theta_1, decay parameters omega_1, omega_2, periods T_1, T_2, T_3, subset ratios tau_1, tau_2
    Initialize lambda = 150, gamma = sqrt(lambda + eta), k = 1
    while theta_k is not satisfactory do
        Select random mini-batch S_prime subset S of size m
        Select random subsets S_1 subset S_prime (size tau_1 m) and S_2 subset S_prime (size tau_2 m)
        Compute stochastic gradient nabla h(theta_k) on S_prime via forward/backward pass
        Sample pseudo-targets y_hat from P_{y|x}(theta_k) on S_1
        Perform backward pass with y_hat on S_1 to obtain predictive derivatives g_i
        Update running estimates of A_bar_{i,j} and G_{i,j} using S_1 data with decay epsilon = min(1 - 1/k, 0.95)
        if k mod T_2 == 0 then
            Set candidate set Gamma = {gamma, omega_2 * gamma, gamma / omega_2}
        else
            Set candidate set Gamma = {gamma}
        end if
        for gamma_cand in Gamma do
            if k mod T_3 == 0 or k <= 3 then
                Compute approximate Fisher inverse using factored Tikhonov damping with gamma_cand
            end if
            Compute update proposal Delta = - F_approx^{-1} nabla h(theta_k)
            Compute final update delta = alpha * Delta + mu * delta_prev by minimizing exact quadratic model M(delta) on S_2
        end for
        Select delta and gamma corresponding to the lowest value of M(delta)
        if k mod T_1 == 0 then
            Compute reduction ratio rho = (h(theta_k + delta) - h(theta_k)) / (0.5 * nabla h(theta_k)^T delta)
            if rho > 3/4 then
                lambda = omega_1 * lambda
            else if rho < 1/4 then
                lambda = lambda / omega_1
            end if
        end if
        theta_{k+1} = theta_k + delta
        delta_prev = delta
        k = k + 1
    end while
  11. Knowl 11 — Optimization Acceleration and Batch Size Scaling on Deep Autoencoders

    empirical result

    K-FAC was evaluated on three deep autoencoder benchmarks: CURVES (784-400-200-100-50-25-6-25-50-100-200-400-784), MNIST (784-1000-500-250-30-250-500-1000-784), and FACES (625-2000-1000-500-30-500-1000-2000-625) using L2L_2 weight decay η=10−5\eta = 10^{-5} and compared against Nesterov-accelerated stochastic gradient descent (SGD).

    Key empirical findings:

    1. Mini-Batch Size Scaling: The per-iteration progress of K-FAC with momentum scales superlinearly with mini-batch size mm, whereas SGD per-iteration progress scales sublinearly. Using an exponentially increasing mini-batch schedule mk=min⁡(m1exp⁡((k−1)/b),∣S∣)m_k = \min(m_1 \exp((k-1)/b), |S|) (growing from m1=1000m_1 = 1000 to full batch size ∣S∣|S| by iteration 500) yields optimal per-second wall-clock progress for K-FAC.
    2. Convergence Speed: In wall-clock time and iteration count, K-FAC reaches lower reconstruction errors orders of magnitude faster than SGD with momentum across all three benchmarks.
    3. Block-Tridiagonal vs. Block-Diagonal: The block-tridiagonal inverse Fisher approximation F^−1\hat{F}^{-1} yields 25%25\% to 40%40\% greater progress per iteration than the block-diagonal approximation Fˇ−1\check{F}^{-1}, producing moderately faster wall-clock convergence on smaller network architectures.

Coverage note — Omitted specific algebraic derivations of Generalized Stein equation solvers from Appendix B (such as Smith-type iterations and Chu's matrix decomposition methods), as they are standard existing numerical linear algebra techniques rather than original methodological contributions.

References

  1. 1.S.-I. Amari and H. Nagaoka. Methods of Information Geometry, volume 191 of Translations of Mathematical monographs. Oxford University Press, 2000.
  2. 2.S.-I. Amari. Natural gradient works efficiently in learning. Neural Computation, 10(2):251–276, 1998.
  3. 3.L. Arnold, A. Auger, N. Hansen, and Y. Ollivier. Information-geometric optimization algorithms: A unifying picture via invariance principles. 2011, arXiv:1106.3708.
  4. 4.S. Becker and Y. LeCun. Improving the Convergence of Back-Propagation Learning with Second Order Methods. In Proceedings of the 1988 Connectionist Models Summer School, pages 29–37, 1989.
  5. 5.C. M. Bishop. Pattern Recognition and Machine Learning (Information Science and Statistics). Springer, 2006.
  6. 6.R. H. Byrd, G. M. Chin, J. Nocedal, and Y. Wu. Sample size selection in optimization methods for machine learning. Mathematical programming, 134(1):127–155, 2012.
  7. 7.K.-w. E. Chu. The solution of the matrix equations AXB−CXD=EAXB - CXD = E and (YA−DZ,YC−BZ)=(E,F)(YA - DZ, YC - BZ) = (E, F). Linear Algebra and its Applications, 93(0):93 – 105, 1987.
  8. 8.C. Darken and J. E. Moody. Note on learning rate schedules for stochastic optimization. In Advances in Neural Information Processing Systems, pages 832–838, 1990.
  9. 9.M. P. Friedlander and M. W. Schmidt. Hybrid deterministic-stochastic methods for data fitting. SIAM J. Scientific Computing, 34(3), 2012.
  10. 10.J. D. Gardiner, A. J. Laub, J. J. Amato, and C. B. Moler. Solution of the sylvester matrix equation AXBT+CXDT=EAXB^T + CXD^T = E. ACM Trans. Math. Softw., 18(2):223–231, June 1992.
  11. 11.X. Glorot and Y. Bengio. Understanding the difficulty of training deep feedforward neural net- works. In Proceedings of AISTATS 2010, volume 9, pages 249–256, may 2010.
  12. 12.R. Grosse and R. Salakhutdinov. Scaling up natural gradient by factorizing fisher information. In Proceedings of the 32nd International Conference on Machine Learning (ICML), 2015.
  13. 13.T. Heskes. On “natural” learning and pruning in multilayered perceptrons. Neural Computation, 12(4):881–901, 2000.
  14. 14.G. E. Hinton and R. R. Salakhutdinov. Reducing the dimensionality of data with neural networks. Science, July 2006.
  15. 15.G. E. Hinton, N. Srivastava, A. Krizhevsky, I. Sutskever, and R. Salakhutdinov. Improving neural networks by preventing co-adaptation of feature detectors. CoRR, abs/1207.0580, 2012.
  16. 16.R. Kiros. Training neural networks with stochastic Hessian-free optimization. In International Conference on Learning Representations (ICLR), 2013.
  17. 17.N. Le Roux, P.-a. Manzagol, and Y. Bengio. Topmoumoute online natural gradient algorithm. In Advances in Neural Information Processing Systems 20, pages 849–856. MIT Press, 2008.
  18. 18.Y. LeCun, L. Bottou, G. Orr, and K. Müller. Efficient backprop. Neural networks: Tricks of the trade, pages 546–546, 1998.
  19. 19.R.-C. Li. Sharpness in rates of convergence for CG and symmetric Lanczos methods. Technical Report 05-01, Department of Mathematics, University of Kentucky, 2005.
  20. 20.J. Martens. Deep learning via Hessian-free optimization. In Proceedings of the 27th International Conference on Machine Learning (ICML), 2010.
  21. 21.J. Martens. New insights and perspectives on the natural gradient method. 2014, arXiv:1412.1193.
  22. 22.J. Martens and I. Sutskever. Training deep and recurrent networks with Hessian-free optimization. In Neural Networks: Tricks of the Trade, pages 479–535. Springer, 2012.
  23. 23.J. Martens, I. Sutskever, and K. Swersky. Estimating the Hessian by backpropagating curvature. In Proceedings of the 29th International Conference on Machine Learning (ICML), 2012.
  24. 24.J. Moré. The Levenberg-Marquardt algorithm: implementation and theory. Numerical analysis, pages 105–116, 1978.
  25. 25.Y. Nesterov. A method of solving a convex programming problem with convergence rate O(1/k)O(1/\sqrt{k}). Soviet Mathematics Doklady, 27:372–376, 1983.
  26. 26.J. Nocedal and S. J. Wright. Numerical optimization. Springer, 2. ed. edition, 2006.
  27. 27.Y. Ollivier. Riemannian metrics for neural networks. 2013, arXiv:1303.0818.
  28. 28.V. Pan and R. Schreiber. An improved newton iteration for the generalized inverse of a matrix, with applications. SIAM Journal on Scientific and Statistical Computing, 12(5):1109–1130, 1991.
  29. 29.H. Park, S.-I. Amari, and K. Fukumizu. Adaptive natural gradient learning algorithms for various stochastic models. Neural Networks, 13(7):755–764, September 2000.
  30. 30.R. Pascanu and Y. Bengio. Revisiting natural gradient for deep networks. In International Conference on Learning Representations, 2014.
  31. 31.D. Plaut, S. Nowlan, and G. E. Hinton. Experiments on learning by back propagation. Technical Report CMU-CS-86-126, Department of Computer Science, Carnegie Mellon University, Pittsburgh, PA, 1986.
  32. 32.B. Polyak. Some methods of speeding up the convergence of iteration methods. USSR Computational Mathematics and Mathematical Physics, 4(5):1 – 17, 1964. ISSN 0041-5553.
  33. 33.M. Pourahmadi. Joint mean-covariance models with applications to longitudinal data: unconstrained parameterisation. Biometrika, 86(3):677–690, 1999.
  34. 34.M. Pourahmadi. Covariance Estimation: The GLM and Regularization Perspectives. Statistical Science, 26(3):369–387, August 2011.
  35. 35.D. Povey, X. Zhang, and S. Khudanpur. Parallel training of DNNs with natural gradient and parameter averaging. In International Conference on Learning Representations: Workshop track, 2015.
  36. 36.T. Raiko, H. Valpola, and Y. LeCun. Deep learning made easier by linear transformations in perceptrons. In AISTATS, volume 22 of JMLR Proceedings, pages 924–932, 2012.
  37. 37.S. Scarpetta, M. Rattray, and D. Saad. Matrix momentum for practical natural gradient learning. Journal of Physics A: Mathematical and General, 32(22):4047, 1999.
  38. 38.T. Schaul, S. Zhang, and Y. LeCun. No more pesky learning rates. In Proceedings of the 30th International Conference on Machine Learning (ICML), 2013.
  39. 39.N. N. Schraudolph. Centering neural network gradient factors. In G. B. Orr and K.-R. Müller, editors, Neural Networks: Tricks of the Trade, volume 1524 of Lecture Notes in Computer Science, pages 207–226. Springer Verlag, Berlin, 1998.
  40. 40.N. N. Schraudolph. Fast curvature matrix-vector products for second-order gradient descent. Neural Computation, 14, 2002.
  41. 41.N. N. Schraudolph, J. Yu, and S. Günter. A stochastic quasi-newton method for online convex optimization. In In Proceedings of 11th International Conference on Artificial Intelligence and Statistics, 2007.
  42. 42.V. Simoncini. Computational methods for linear matrix equations. 2014.
  43. 43.R. Smith. Matrix equation XA+BX=CXA + BX = C. SIAM J. Appl. Math., 16(1):198 – 201, 1968.
  44. 44.I. Sutskever, J. Martens, G. Dahl, and G. Hinton. On the importance of initialization and momentum in deep learning. In Proceedings of the 30th International Conference on Machine Learning (ICML), 2013.
  45. 45.K. Swersky, B. Chen, B. Marlin, and N. de Freitas. A tutorial on stochastic approximation algorithms for training restricted boltzmann machines and deep belief nets. In Information Theory and Applications Workshop (ITA), 2010, pages 1–10, Jan 2010.
  46. 46.C. F. Van Loan. The ubiquitous kronecker product. Journal of computational and applied mathematics, 123(1):85–100, 2000.
  47. 47.T. Vatanen, T. Raiko, H. Valpola, and Y. LeCun. Pushing stochastic gradient towards second-order methods – backpropagation learning with transformations in nonlinearities. 2013, arXiv:1301.3476.
  48. 48.O. Vinyals and D. Povey. Krylov subspace descent for deep learning. In International Conference on Artificial Intelligence and Statistics (AISTATS), 2012.
  49. 49.S. Wiesler, A. Richard, R. Schlüter, and H. Ney. Mean-normalized stochastic gradient for large-scale deep learning. In IEEE International Conference on Acoustics, Speech, and Signal Processing, pages 180–184, 2014.
  50. 50.M. D. Zeiler. ADADELTA: An adaptive learning rate method. 2013, arXiv:1212.5701.

Citation

MLA
Martens, J., and R. Grosse. “Optimizing Neural Networks with Kronecker-factored Approximate Curvature”. arXiv, 2015, http://arxiv.org/abs/1503.05671v7.
APA
Martens, J., & Grosse, R. (2015). Optimizing Neural Networks with Kronecker-factored Approximate Curvature. arXiv. http://arxiv.org/abs/1503.05671v7
Chicago
Martens, J., and R. Grosse. 2015. “Optimizing Neural Networks with Kronecker-factored Approximate Curvature”. arXiv. http://arxiv.org/abs/1503.05671v7.
Harvard
Martens, J. and Grosse, R. (2015) “Optimizing Neural Networks with Kronecker-factored Approximate Curvature”, arXiv [Preprint]. Available at: http://arxiv.org/abs/1503.05671v7.
Vancouver
1. Martens J, Grosse R (2015) Optimizing Neural Networks with Kronecker-factored Approximate Curvature. arXiv

BibTeX

@article{martens2015optimizing,
  title = {Optimizing Neural Networks with Kronecker-factored Approximate Curvature},
  author = {Martens, James and Grosse, Roger},
  year = {2015},
  journal = {arXiv},
  url = {http://arxiv.org/abs/1503.05671v7},
  eprint = {1503.05671}
}
Metadata:arXiv

Access the Paper

This paper is available from its original source. Click below to access the PDF.

Open PDF
License: Authors