Optimizing Neural Networks with Kronecker-factored Approximate Curvature
James Martens, Roger Grosse
Introduction
The problem of training neural networks is one of the most important and highly investigated ones in machine learning. Despite work on layer-wise pretraining schemes, and various sophisticated optimization methods which try to approximate Newton-Raphson updates or natural gradient updates, stochastic gradient descent (SGD), possibly augmented with momentum, remains the method of choice for large-scale neural network training (Sutskever et al., 2013).
From the work on Hessian-free optimization (HF) (Martens, 2010) and related methods (e.g. Vinyals and Povey, 2012) we know that updates computed using local curvature information can make much more progress per iteration than the scaled gradient. The reason that HF sees fewer practical applications than SGD are twofold. Firstly, its updates are much more expensive to compute, as they involve running linear conjugate gradient (CG) for potentially hundreds of iterations, each of which requires a matrix-vector product with the curvature matrix (which are as expensive to compute as the stochastic gradient on the current mini-batch). Secondly, HF’s estimate of the curvature matrix must remain fixed while CG iterates, and thus the method is able to go through much less data than SGD can in a comparable amount of time, making it less well suited to stochastic optimizations.
As discussed in Martens and Sutskever (2012) and Sutskever et al. (2013), CG has the potential to be much faster at local optimization than gradient descent, when applied to quadratic objective functions. Thus, insofar as the objective can be locally approximated by a quadratic, each step of CG could potentially be doing a lot more work than each iteration of SGD, which would result in HF being much faster overall than SGD. However, there are examples of quadratic functions (e.g. Li, 2005), characterized by curvature matrices with highly spread-out eigenvalue distributions, where CG will have no distinct advantage over well-tuned gradient descent with momentum. Thus, insofar as the quadratic functions being optimized by CG within HF are of this character, HF shouldn’t in principle be faster than well-tuned SGD with momentum. The extent to which neural network objective functions give rise to such quadratics is unclear, although Sutskever et al. (2013) provides some preliminary evidence that they do.
CG falls victim to this worst-case analysis because it is a first-order method. This motivates us to consider methods which don’t rely on first-order methods like CG as their primary engines of optimization. One such class of methods which have been widely studied are those which work by directly inverting a diagonal, block-diagonal, or low-rank approximation to the curvature matrix (e.g. Becker and LeCun, 1989; Schaul et al., 2013; Zeiler, 2013; Le Roux et al., 2008; Ollivier, 2013). In fact, a diagonal approximation of the Fisher information matrix is used within HF as a preconditioner for CG. However, these methods provide only a limited performance improvement in practice, especially compared to SGD with momentum (see for example Schraudolph et al., 2007; Zeiler, 2013), and many practitioners tend to forgo them in favor of SGD or SGD with momentum.
We know that the curvature associated with neural network objective functions is highly non-diagonal, and that updates which properly respect and account for this non-diagonal curvature, such as those generated by HF, can make much more progress minimizing the objective than the plain gradient or updates computed from diagonal approximations of the curvature (usually HF updates are required to adequately minimize most objectives, compared to the required by methods that use diagonal approximations). Thus, if we had an efficient and direct way to compute the inverse of a high-quality non-diagonal approximation to the curvature matrix (i.e. without relying on first-order methods like CG) this could potentially yield an optimization method whose updates would be large and powerful like HF’s, while being (almost) as cheap to compute as the stochastic gradient.
In this work we develop such a method, which we call Kronecker-factored Approximate Curvature (K-FAC). We show that our method can be much faster in practice than even highly tuned implementations of SGD with momentum on certain standard neural network optimization benchmarks.
The main ingredient in K-FAC is a sophisticated approximation to the Fisher information matrix, which despite being neither diagonal nor low-rank, nor even block-diagonal with small blocks, can be inverted very efficiently, and can be estimated in an online fashion using arbitrarily large subsets of the training data (without increasing the cost of inversion).
This approximation is built in two stages. In the first, the rows and columns of the Fisher are divided into groups, each of which corresponds to all the weights in a given layer, and this gives rise to a block-partitioning of the matrix (where the blocks are much larger than those used by Le Roux et al. (2008) or Ollivier (2013)). These blocks are then approximated as Kronecker products between much smaller matrices, which we show is equivalent to making certain approximating assumptions regarding the statistics of the network’s gradients.
In the second stage, this matrix is further approximated as having an inverse which is either block-diagonal or block-tridiagonal. We justify this approximation through a careful examination of the relationships between inverse covariances, tree-structured graphical models, and linear regression. Notably, this justification doesn’t apply to the Fisher itself, and our experiments confirm that while the inverse Fisher does indeed possess this structure (approximately), the Fisher itself does not.
The rest of this paper is organized as follows. Section 2 gives basic background and notation for neural networks and the natural gradient. Section 3 describes our initial Kronecker product approximation to the Fisher. Section 4 describes our further block-diagonal and block-tridiagonal approximations of the inverse Fisher, and how these can be used to derive an efficient inversion algorithm. Section 5 describes how we compute online estimates of the quantities required by our inverse Fisher approximation over a large “window” of previously processed mini-batches (which makes K-FAC very different from methods like HF or KSD, which base their estimates of the curvature on a single mini-batch). Section 6 describes how we use our approximate Fisher to obtain a practical and robust optimization algorithm which requires very little manual tuning, through the careful application of various theoretically well-founded “damping” techniques that are standard in the optimization literature. Note that damping techniques compensate both for the local quadratic approximation being implicitly made to the objective, and for our further approximation of the Fisher, and are non-optional for essentially any 2nd-order method like K-FAC to work properly, as is well established by both theory and practice within the optimization literature (Nocedal and Wright, 2006). Section 7 describes a simple and effective way of adding a type of “momentum” to K-FAC, which we have found works very well in practice. Section 8 describes the computational costs associated with K-FAC, and various ways to reduce them to the point where each update is at most only several times more expensive to compute than the stochastic gradient. Section 9 gives complete high-level pseudocode for K-FAC. Section 10 characterizes a broad class of network transformations and reparameterizations to which K-FAC is essentially invariant. Section 11 considers some related prior methods for neural network optimization. Proofs of formal results are located in the appendix.
Background and notation
In this section we will define the basic notation for feed-forward neural networks which we will use throughout this paper. Note that this presentation closely follows the one from Martens (2014).
We let denote the loss function which measures the disagreement between a prediction made by the network, and a target . The training objective function is the average (or expectation) of losses with respect to a training distribution over input-target pairs . is a proxy for the objective which we actually care about but don’t have access to, which is the expectation of the loss taken with respect to the true data distribution .
We will assume that the loss is given by the negative log probability associated with a simple predictive distribution for parameterized by , i.e. that we have
where is ’s density function. This is the case for both the standard least-squares and cross-entropy objective functions, where the predictive distributions are multivariate normal and multinomial, respectively.
We will let denote the conditional distribution defined by the neural network, as parameterized by , and its density function. Note that minimizing the objective function can be seen as maximum likelihood learning of the model .
For convenience we will define the following additional notation:
Algorithm 1 shows how to compute the gradient of the loss function of a neural network using standard backpropagation.
2 The Natural Gradient
Because our network defines a conditional model , it has an associated Fisher information matrix (which we will simply call “the Fisher”) which is given by
Here, the expectation is taken with respect to the data distribution over inputs , and the model’s predictive distribution over . Since we usually don’t have access to , and the above expectation would likely be intractable even if we did, we will instead compute using the training distribution over inputs .
The well-known natural gradient (Amari, 1998) is defined as . Motivated from the perspective of information geometry (Amari and Nagaoka, 2000), the natural gradient defines the direction in parameter space which gives the largest change in the objective per unit of change in the model, as measured by the KL-divergence. This is to be contrasted with the standard gradient, which can be defined as the direction in parameter space which gives the largest change in the objective per unit of change in the parameters, as measured by the standard Euclidean metric.
The GGN has served as the curvature matrix of choice in HF and related methods, and so in light of its equivalence to the Fisher, these 2nd-order methods can be seen as approximate natural gradient methods. And perhaps more importantly from a practical perspective, natural gradient-based optimization methods can conversely be viewed as 2nd-order optimization methods, which as pointed out by Martens (2014)), brings to bare the vast wisdom that has accumulated about how to make such methods work well in both theory and practice (e.g Nocedal and Wright, 2006). In Section 6 we productively make use of these connections in order to design a robust and highly effective optimization method using our approximation to the natural gradient/Fisher (which is developed in Sections 3 and 4).
For some good recent discussion and analysis of the natural gradient, see Arnold et al. (2011); Martens (2014); Pascanu and Bengio (2014).
A block-wise Kronecker-factored Fisher approximation
The main computational challenge associated with using the natural gradient is computing (or its product with ). For large networks, with potentially millions of parameters, computing this inverse naively is computationally impractical. In this section we develop an initial approximation of which will be a key ingredient in deriving our efficiently computable approximation to and the natural gradient.
Noting that and that we have , and thus we can rewrite as
Note that the Kronecker product satisfies many convenient properties that we will make use of in this paper, especially the identity . See Van Loan (2000) for a good discussion of the Kronecker product.
where and .
which has the form of what is known as a Khatri-Rao product in multivariate statistics.
The expectation of a Kronecker product is, in general, not equal to the Kronecker product of expectations, and so this is indeed a major approximation to make, and one which likely won’t become exact under any realistic set of assumptions, or as a limiting case in some kind of asymptotic analysis. Nevertheless, it seems to be fairly accurate in practice, and is able to successfully capture the “coarse structure” of the Fisher, as demonstrated in Figure 2 for an example network.
As we will see in later sections, this approximation leads to significant computational savings in terms of storage and inversion, which we will be able to leverage in order to design an efficient algorithm for computing an approximation to the natural gradient.
Consider an arbitrary pair of weights and from the network, where denotes the value of the -th entry. We have that the corresponding derivatives of these weights are given by and , where we denote for convenience , , , and .
The approximation given by eqn. 1 is equivalent to making the following approximation for each pair of weights:
And thus one way to interpret the approximation in eqn. 1 is that we are assuming statistical independence between products of unit activities and products of unit input derivatives.
Another more detailed interpretation of the approximation emerges by considering the following expression for the approximation error (which is derived in the appendix):
Here denotes the cumulant of its arguments. Cumulants are a natural generalization of the concept of mean and variance to higher orders, and indeed 1st-order cumulants are means and 2nd-order cumulants are covariances. Intuitively, cumulants of order measure the degree to which the interaction between variables is intrinsically of order , as opposed to arising from many lower-order interactions.
A basic upper bound for the approximation error is
which will be small if all of the higher-order cumulants are small (i.e. those of order 3 or higher). Note that in principle this upper bound may be loose due to possible cancellations between the terms in eqn. 3.
Because higher-order cumulants are zero for variables jointly distributed according to a multivariate Gaussian, it follows that this upper bound on the approximation error will be small insofar as the joint distribution over , , , and is well approximated by a multivariate Gaussian. And while we are not aware of an argument for why this should be the case in practice, it does seem to be the case that for the example network from Figure 2, the size of the error is well predicted by the size of the higher-order cumulants. In particular, the total approximation error, summed over all pairs of weights in the middle 4 layers, is , and is of roughly the same size as the corresponding upper bound (), whose size is tied to that of the higher order cumulants (due to the impossibility of cancellations in eqn. 4).
Suppose we are given a multivariate distribution whose associated covariance matrix is .
Define the matrix so that for , is the coefficient on the -th variable in the optimal linear predictor of the -th variable from all the other variables, and for , . Then define the matrix to be the diagonal matrix where is the variance of the error associated with such a predictor of the -th variable.
Pourahmadi (2011) showed that and can be obtained from the inverse covariance by the formulas
from which it follows that the inverse covariance matrix can be expressed as
Intuitively, this result says that each row of the inverse covariance is given by the coefficients of the optimal linear predictor of the -th variable from the others, up to a scaling factor. So if the -th variable is much less “useful” than the other variables for predicting the -th variable, we can expect that the -th entry of the inverse covariance will be relatively small.
Note that “usefulness” is a subtle property as we have informally defined it. In particular, it is not equivalent to the degree of correlation between the -th and -th variables, or any such simple measure. As a simple example, consider the case where the -th variable is equal to the -th variable plus independent Gaussian noise. Since any linear predictor can achieve a lower variance simply by shifting weight from the -th variable to the -th variable, we have that the -th variable is not useful (and its coefficient will thus be zero) in the task of predicting the -th variable for any setting of other than or .
Now while in reality the ’s are generated using information from adjacent layers according to a process that is neither linear nor Gaussian, it nonetheless stands to reason that their joint statistics might be reasonably approximated by such a model. In fact, the idea of approximating the distribution over loss gradients with a directed graphical model forms the basis of the recent FANG method of Grosse and Salakhutdinov (2015).
Figure 3 examines the extent to which the inverse Fisher is well approximated as block-diagonal or block-tridiagonal for an example network.
Using the identity we can easily compute the inverse of as
Then to compute , we can make use of the well-known identity to get
Note that block-diagonal approximations to the Fisher information have been proposed before in TONGA (Le Roux et al., 2008), where each block corresponds to the weights associated with a particular unit. In our block-diagonal approximation, the blocks correspond to all the parameters in a given layer, and are thus much larger. In fact, they are so large that they would be impractical to invert as general matrices.
To establish that such a matrix is well defined and can be inverted efficiently, we first observe that assuming that is block-tridiagonal is equivalent to assuming that it is the precision matrix of an undirected Gaussian graphical model (UGGM) over (as depicted in Figure 4), whose density function is proportional to . As this graphical model has a tree structure, there is an equivalent directed graphical model with the same distribution and the same (undirected) graphical structure (e.g. Bishop, 2006), where the directionality of the edges is given by a directed acyclic graph (DAG). Moreover, this equivalent directed model will also be linear/Gaussian, and hence a directed Gaussian Graphical model (DGGM).
Next we will show how the parameters of such a DGGM corresponding to can be efficiently recovered from the tridiagonal blocks of , so that is uniquely determined by these blocks (and hence well-defined). We will assume here that the direction of the edges is from the higher layers to the lower ones. Note that a different choice for these directions would yield a superficially different algorithm for computing the inverse of that would nonetheless yield the same output.
For each , we will denote the conditional covariance matrix of on by and the linear coefficients from to by the matrix , so that the conditional distributions defining the model are
The conditional covariance is thus given by
Following the work of Grosse and Salakhutdinov (2015), we use the block generalization of well-known “Cholesky” decomposition of the precision matrix of DGGMs (Pourahmadi, 1999), which gives
Thus, matrix-vector multiplication with amounts to performing matrix-vector multiplication by , followed by , and then by .
As in the block-diagonal case considered in the previous subsection, matrix-vector products with (and ) can be efficiently computed using the well-known identity . In particular, can be computed as
and similarly can be computed as
where the ’s and ’s are defined in terms of and as in the previous subsection.
Multiplying a vector by amounts to multiplying each by the corresponding . This is slightly tricky because is the difference of Kronecker products, so we cannot use the straightforward identity . Fortunately, there are efficient techniques for inverting such matrices which we discuss in detail in Appendix B.
4 Examining the approximation quality
Estimating the required statistics
Recall that and . Both approximate Fisher inverses discussed in Section 4 require some subset of these. In particular, the block-diagonal approximation requires them for , while the block-tridiagonal approximation requires them for (noting that and ).
Since the ’s don’t depend on , we can take the expectation with respect to just the training distribution over the inputs . On the other hand, the ’s do depend on , and so the expectationIt is important to note this expectation should not be taken with respect to the training/data distribution over (i.e. or ). Using the training/data distribution for would perhaps give an approximation to a quantity known as the “empirical Fisher information matrix”, which lacks the previously discussed equivalence to the Generalized Gauss-Newton matrix, and would not be compatible with the theoretical analysis performed in Section 3.1 (in particular, Lemma 4 would break down). Moreover, such a choice would not give rise to what is usually thought of as the natural gradient, and based on the findings of Martens (2010), would likely perform worse in practice as part of an optimization algorithm. See Martens (2014) for a more detailed discussion of the empirical Fisher and reasons why it may be a poor choice for a curvature matrix compared to the standard Fisher. must be taken with respect to both and the network’s predictive distribution .
While computing matrix-vector products with the could be done exactly and efficiently for a given input (or small mini-batch of ’s) by adapting the methods of Schraudolph (2002), there doesn’t seem to be a sufficiently efficient method for computing the entire matrix itself. Indeed, the hardness results of Martens et al. (2012) suggest that this would require, for each example in the mini-batch, work that is asymptotically equivalent to matrix-matrix multiplication involving matrices the same size as . While a small constant number of such multiplications is arguably an acceptable cost (see Section 8), a number which grows with the size of the mini-batch would not be.
Instead, we will approximate the expectation over by a standard Monte-Carlo estimate obtained by sampling ’s from the network’s predictive distribution and then rerunning the backwards phase of backpropagation (see Algorithm 1) as if these were the training targets.
Note that computing/estimating the required /’s involves computing averages over outer products of various ’s from network’s usual forward pass, and ’s from the modified backwards pass (with targets sampled as above). Thus we can compute/estimate these quantities on the same input data used to compute the gradient , at the cost of one or more additional backwards passes, and a few additional outer-product averages. Fortunately, this turns out to be quite inexpensive, as we have found that just one modified backwards pass is sufficient to obtain a good quality estimate in practice, and the required outer-product averages are similar to those already used to compute the gradient in the usual backpropagation algorithm.
In the case of online/stochastic optimization we have found that the best strategy is to maintain running estimates of the required ’s and ’s using a simple exponentially decaying averaging scheme. In particular, we take the new running estimate to be the old one weighted by , plus the estimate on the new mini-batch weighted by , for some . In our experiments we used , where is the iteration number.
Note that the more naive averaging scheme where the estimates from each iteration are given equal weight would be inappropriate here. This is because the ’s and ’s depend on the network’s parameters , and these will slowly change over time as optimization proceeds, so that estimates computed many iterations ago will become stale.
This kind of exponentially decaying averaging scheme is commonly used in methods involving diagonal or block-diagonal approximations (with much smaller blocks than ours) to the curvature matrix (e.g. LeCun et al., 1998; Park et al., 2000; Schaul et al., 2013). Such schemes have the desirable property that they allow the curvature estimate to depend on much more data than can be reasonably processed in a single mini-batch.
Notably, for methods like HF which deal with the exact Fisher indirectly via matrix-vector products, such a scheme would be impossible to implement efficiently, as the exact Fisher matrix (or GGN) seemingly cannot be summarized using a compact data structure whose size is independent of the amount of data used to estimate it. Indeed, it seems that the only representation of the exact Fisher which would be independent of the amount of data used to estimate it would be an explicit matrix (which is far too big to be practical). Because of this, HF and related methods must base their curvature estimates only on subsets of data that can be reasonably processed all at once, which limits their effectiveness in the stochastic optimization regime.
Update damping
The idealized natural gradient approach is to follow the smooth path in the Riemannian manifold (implied by the Fisher information matrix viewed as a metric tensor) that is generated by taking a series of infinitesimally small steps (in the original parameter space) in the direction of the natural gradient (which gets recomputed at each point). While this is clearly impractical as a real optimization method, one can take larger steps and still follow these paths approximately. But in our experience, to obtain an update which satisfies the minimal requirement of not worsening the objective function value, it is often the case that one must make the step size so small that the resulting optimization algorithm performs poorly in practice.
The reason that the natural gradient can only be reliably followed a short distance is that it is defined merely as an optimal direction (which trades off improvement in the objective versus change in the predictive distribution), and not a discrete update.
Fortunately, as observed by Martens (2014), the natural gradient can be understood using a more traditional optimization-theoretic perspective which implies how it can be used to generate updates that will be useful over larger distances. In particular, when is an exponential family model with as its natural parameters (as it will be in our experiments), Martens (2014) showed that the Fisher becomes equivalent to the Generalized Gauss-Newton matrix (GGN), which is a positive semi-definite approximation of the Hessian of . Additionally, there is the well-known fact that when is the negative log-likelihood function associated with a given pair (as we are assuming in this work), the Hessian of and the Fisher are closely related in the sense is the expected Hessian of under the training distribution , while is the expected Hessian of under the model’s distribution (defined by the density ).
From the interpretation of the natural gradient as the minimizer of , we can see that it fails to be useful as a local update only insofar as fails to be a good local approximation to . And so as argued by Martens (2014), it is natural to make use of the various “damping” techniques that have been developed in the optimization literature for dealing with the breakdowns in local quadratic approximations that inevitably occur during optimization. Notably, this breakdown usually won’t occur in the final “local convergence” stage of optimization where the function becomes well approximated as a convex quadratic within a sufficiently large neighborhood of the local optimum. This is the phase traditionally analyzed in most theoretical results, and while it is important that an optimizer be able to converge well in this final phase, it is arguably much more important from a practical standpoint that it behaves sensibly before this phase.
This initial “exploration phase” (Darken and Moody, 1990) is where damping techniques help in ways that are not apparent from the asymptotic convergence theorems alone, which is not to say there are not strong mathematical arguments that support their use (see Nocedal and Wright, 2006). In particular, in the exploration phase it will often still be true that is accurately approximated by a convex quadratic locally within some region around , and that therefor optimization can be most efficiently performed by minimizing a sequence of such convex quadratic approximations within adaptively sized local regions.
Note that well designed damping techniques, such as the ones we will employ, automatically adapt to the local properties of the function, and effectively “turn themselves off” when the quadratic model becomes a sufficiently accurate local approximation of , allowing the optimizer to achieve the desired asymptotic convergence behavior (Moré, 1978).
Successful and theoretically well-founded damping techniques include Tikhonov damping (aka Tikhonov regularization, which is closely connected to the trust-region method) with Levenberg-Marquardt style adaptation (Moré, 1978), line-searches, and trust regions, truncation, etc., all of which tend to be much more effective in practice than merely applying a learning rate to the update, or adding a fixed multiple of the identity to the curvature matrix. Indeed, a subset of these techniques was exploited in the work of Martens (2010), and primitive versions of them have appeared implicitly in older works such as Becker and LeCun (1989), and also in many recent diagonal methods like that of Zeiler (2013), although often without a good understanding of what they are doing and why they help.
Crucially, more powerful 2nd-order optimizers like HF and K-FAC, which have the capability of taking much larger steps than 1st-order methods (or methods which use diagonal curvature matrices), require more sophisticated damping solutions to work well, and will usually completely fail without them, which is consistent with predictions made in various theoretical analyses (e.g. Nocedal and Wright, 2006). As an analogy one can think of such powerful 2nd-order optimizers as extremely fast racing cars that need more sophisticated control systems than standard cars to prevent them from flying off the road. Arguably one of the reasons why high-powered 2nd-order optimization methods have historically tended to under-perform in machine learning applications, and in neural network training in particular, is that their designers did not understand or take seriously the issue of quadratic model approximation quality, and did not employ the more sophisticated and effective damping techniques that are available to deal with this issue.
For a detailed review and discussion of various damping techniques and their crucial role in practical 2nd-order optimization methods, we refer the reader to Martens and Sutskever (2012).
2 A highly effective damping scheme for K-FAC
Methods like HF which use the exact Fisher seem to work reasonably well with an adaptive Tikhonov regularization technique where is added to , and where is adapted according to Levenberg-Marquardt style adjustment rule. This common and well-studied method can be shown to be equivalent to imposing an adaptive spherical region (known as a “trust region”) which constrains the optimization of the quadratic model (e.g Nocedal and Wright, 2006). However, we found that this simple technique is insufficient when used with our approximate natural gradient update proposals. In particular, we have found that there never seems to be a “good” choice for that gives rise to updates which are of a quality comparable to those produced by methods that use the exact Fisher, such as HF.
One possible explanation for this finding is that, unlike quadratic models based on the exact Fisher (or equivalently, the GGN), the one underlying K-FAC has no guarantee of being accurate up to 2nd-order. Thus, must remain large in order to compensate for this intrinsic 2nd-order inaccuracy of the model, which has the side effect of “washing out” the small eigenvalues (which represent important low-curvature directions).
Fortunately, through trial and error, we were able to find a relatively simple and highly effective damping scheme, which combines several different techniques, and which works well within K-FAC. Our scheme works by computing an initial update proposal using a version of the above described adaptive Tikhonov damping/regularization method, and then re-scaling this according to quadratic model computed using the exact Fisher. This second step is made practical by the fact that it only requires a single matrix-vector product with the exact Fisher, and this can be computed efficiently using standard methods. We discuss the details of this scheme in the following subsections.
3 A factored Tikhonov regularization technique
In the first stage of our damping scheme we generate a candidate update proposal by applying a slightly modified form of Tikhononv damping to our approximate Fisher, before multiplying by its inverse.
Because this is the sum of two Kronecker products we cannot use the simple identity anymore. Fortunately however, there are efficient techniques for inverting such matrices, which we discuss in detail in Appendix B.
As this is a single Kronecker product, all of the computations described in Sections 4.2 and 4.3 can still be used here too, simply by replacing each and with their modified versions and .
To see why the expression in eqn. 7 is a reasonable approximation to eqn. 6, note that expanding it gives
which differs from eqn. 6 by the residual error expression
While the choice of is simple and can sometimes work well in practice, a slightly more principled choice can be found by minimizing the obvious upper bound (following from the triangle inequality) on the matrix norm of this residual expression, for some matrix norm . This gives
Evaluating this expression can be done efficiently for various common choices of the matrix norm . For example, for a general we have where is the height/dimension of , and also .
In our experience, one of the best and must robust choices for the norm is the trace-norm, which for PSD matrices is given by the trace. With this choice, the formula for has the following simple form:
where is the dimension (number of units) in layer . Intuitively, the inner fraction is just the average eigenvalue of divided by the average eigenvalue of .
Interestingly, we have found that this factored approximate Tikhonov approach, which was originally motivated by computational concerns, often works better than the exact version (eqn. 6) in practice. The reasons for this are still somewhat mysterious to us, but it may have to do with the fact that the inverse of the product of two quantities is often most robustly estimated as the inverse of the product of their individually regularized estimates.
4 Re-scaling according to the exact F𝐹F
Given an update proposal produced by multiplying the negative gradient by our approximate Fisher inverse (subject to the Tikhonov technique described in the previous subsection), the second stage of our proposed damping scheme re-scales according to the quadratic model as computed with the exact , to produce a final update .
More precisely, we optimize according to the value of the quadratic model
To evaluate this formula we use the current stochastic gradient (i.e. the same one used to produce ), and compute matrix-vector products with using the input data from the same mini-batch. While using a mini-batch to compute gets away from the idea of basing our estimate of the curvature on a long history of data (as we do with our approximate Fisher), it is made slightly less objectionable by the fact that we are only using it to estimate a single scalar quantity (). This is to be contrasted with methods like HF which perform a long and careful optimization of using such an estimate of .
Because the matrix-vector products with are only used to compute scalar quantities in K-FAC, we can reduce their computational cost by roughly one half (versus standard matrix-vector products with ) using a simple trick which is discussed in Appendix C.
Intuitively, this second stage of our damping scheme effectively compensates for the intrinsic inaccuracy of the approximate quadratic model (based on our approximate Fisher) used to generate the initial update proposal , by essentially falling back on a more accurate quadratic model based on the exact Fisher.
Interestingly, by re-scaling according to , K-FAC can be viewed as a version of HF which uses our approximate Fisher as a preconditioning matrix (instead of the traditional diagonal preconditioner), and runs CG for only 1 step, initializing it from 0. This observation suggests running CG for longer, thus obtaining an algorithm which is even closer to HF (although using a much better preconditioner for CG). Indeed, this approach works reasonably well in our experience, but suffers from some of the same problems that HF has in the stochastic setting, due its much stronger use of the mini-batch–estimated exact .
Figure 7 demonstrates the effectiveness of this re-scaling technique versus the simpler method of just using the raw as an update proposal. We can see that , without being re-scaled, is a very poor update to , and won’t even give any improvement in the objective function unless the strength of the factored Tikhonov damping terms is made very large. On the other hand, when the update is re-scaled, we can afford to compute using a much smaller strength for the factored Tikhonov damping terms, and overall this yields a much larger and more effective update to .
5 Adapting λ𝜆\lambda
Tikhonov damping can be interpreted as implementing a trust-region constraint on the update , so that in particular the constraint is imposed for some , where depends on and the curvature matrix (e.g. Nocedal and Wright, 2006). While some approaches adjust and then seek to find the matching , it is often simpler just to adjust directly, as the precise relationship between and is complicated, and the curvature matrix is constantly evolving as optimization takes place.
The theoretically well-founded Levenberg-Marquardt style rule used by HF for doing this, which we will adopt for K-FAC, is given by
if then
if then
where is the “reduction ratio” and is some decay constant, and all quantities are computed on the current mini-batch (and uses the exact ).
Intuitively, this rule tries to make as small as possible (and hence the implicit trust-region as large as possible) while maintaining the property that the quadratic model remains a good local approximation to (in the sense that it accurately predicts the value of for the which gets chosen at each iteration). It has the desirable property that as the optimization enters the final convergence stage where becomes an almost exact approximation in a sufficiently large neighborhood of the local minimum, the value of will go rapidly enough towards that it doesn’t interfere with the asymptotic local convergence theory enjoyed by 2nd-order methods (Moré, 1978).
In our experiments we applied this rule every iterations of K-FAC, with and , from a starting value of . Note that the optimal value of and the starting value of may be application dependent, and setting them inappropriately could significantly slow down K-FAC in practice.
Computing can be done quite efficiently. Note that for the optimal , , and is available from the usual forward pass. The only remaining quantity which is needed to evaluate is thus , which will require an additional forward pass. But fortunately, we only need to perform this once every iterations.
6 Maintaining a separate damping strength for the approximate Fisher
While the scheme described in the previous sections works reasonably well in most situations, we have found that in order to avoid certain failure cases and to be truly robust in a large variety of situations, the Tikhonov damping strength parameter for the factored Tikhonov technique described in Section 6.3 should be maintained and adjusted independently of . To this end we replace the expression in Section 6.3 with a separate constant , which we initialize to but which is then adjusted using a different rule, which is described at the end of this section.
The reasoning behind this modification is as follows. The role of , according to the Levenberg Marquardt theory (Moré, 1978), is to be as small as possible while maintaining the property that the quadratic model remains a trust-worthy approximation of the true objective. Meanwhile, ’s role is to ensure that the initial update proposal is as good an approximation as possible to the true optimum of (as computed using a mini-batch estimate of the exact ), so that in particular the re-scaling performed in Section 6.4 is as benign as possible. While one might hope that adding the same multiple of the identity to our approximate Fisher as we do to the exact (as it appears in ) would produce the best in this regard, this isn’t obviously the case. In particular, using a larger multiple may help compensate for the approximation we are making to the Fisher when computing , and thus help produce a more “conservative” but ultimately more useful initial update proposal , which is what we observe happens in practice.
A simple measure of the quality of our choice of is the (negative) value of the quadratic model for the optimally chosen . To adjust based on this measure (or others like it) we use a simple greedy adjustment rule. In particular, every iterations during the optimization we try 3 different values of (, , and , where is the current value) and choose the new to be the best of these, as measured by our quality metric. In our experiments we used (which must be a multiple of the constant as defined in Section 8), and .
We have found that works well in practice as a measure of the quality of , and has the added bonus that it can be computed at essentially no additional cost from the incidental quantities already computed when solving for the optimal . In our initial experiments we found that using it gave similar results to those obtained by using other obvious measures for the quality of , such as .
Momentum
Sutskever et al. (2013) found that momentum (Polyak, 1964; Plaut et al., 1986) was very helpful in the context of stochastic gradient descent optimization of deep neural networks. A version of momentum is also present in the original HF method, and it plays an arguably even more important role in more “stochastic” versions of HF (Martens and Sutskever, 2012; Kiros, 2013).
A natural way of adding momentum to K-FAC, and one which we have found works well in practice, is to take the update to be , where is the final update computed at the previous iteration, and where and are chosen to minimize . This allows K-FAC to effectively build up a better solution to the local quadratic optimization problem (where uses the exact ) over many iterations, somewhat similarly to how Matrix Momentum (Scarpetta et al., 1999) and HF do this (see Sutskever et al., 2013).
The optimal solution for and can be computed as
The main cost in evaluating this formula is computing the two matrix-vector products and . Fortunately, the technique discussed in Appendix C can be applied here to compute the 4 required scalars at the cost of only two forwards passes (equivalent to the cost of only one matrix-vector product with ).
Empirically we have found that this type of momentum provides substantial acceleration in regimes where the gradient signal has a low noise to signal ratio, which is usually the case in the early to mid stages of stochastic optimization, but can also be the case in later stages if the mini-batch size is made sufficiently large. These findings are consistent with predictions made by convex optimization theory, and with older empirical work done on neural network optimization (LeCun et al., 1998).
Notably, because the implicit “momentum decay constant” in our method is being computed on the fly, one doesn’t have to worry about setting schedules for it, or adjusting it via heuristics, as one often does in the context of SGD.
Interestingly, if is a quadratic function (so the definition of remains fixed at each iteration) and all quantities are computed deterministically (i.e. without noise), then using this type of momentum makes K-FAC equivalent to performing preconditioned linear CG on , with the preconditioner given by our approximate Fisher. This follows from the fact that linear CG can be interpreted as a momentum method where the learning rate and momentum decay coefficient are chosen to jointly minimize at the current iteration.
Computational Costs and Efficiency Improvements
Let be the typical number of units in each layer and the mini-batch size. The significant computational tasks required to compute a single update/iteration of K-FAC, and rough estimates of their associated computational costs, are as follows:
Here the are various constants that account for implementation details, and we are assuming the use of the naive cubic matrix-matrix multiplication and inversion algorithms when producing the cost estimates. Note that it it is hard to assign precise values to the constants, as they very much depend on how these various tasks are implemented.
Note that most of the computations required for these tasks will be sped up greatly by performing them in parallel across units, layers, training cases, or all of these. The above cost estimates however measure sequential operations, and thus may not accurately reflect the true computation times enjoyed by a parallel implementation. In our experiments we used a vectorized implementation that performed the computations in parallel over units and training cases, although not over layers (which is possible for computations that don’t involve a sequential forwards or backwards “pass” over the layers).
Tasks 1 and 2 represent the standard stochastic gradient computation.
The cost of task 8 can be made relatively insignificant by making the adjustment period for large enough. We used in our experiments.
The costs of tasks 5 and 6 are hard to compare directly with the costs associated with computing the gradient, as their relative sizes will depend on factors such as the architecture of the neural network being trained, as well as the particulars of the implementation. However, one quick observation we can make is that both tasks 5 and 6 involve computations that be performed in parallel across the different layers, which is to be contrasted with many of the other tasks which require sequential passes over the layers of the network.
Clearly, if , then the cost of tasks 5 and 6 becomes negligible in comparison to the others. However, it is more often the case that is comparable or perhaps smaller than . Moreover, while algorithms for inverses and SVDs tend to have the same asymptotic cost as matrix-matrix multiplication, they are at least several times more expensive in practice, in addition to being harder to parallelize on modern GPU architectures (indeed, CPU implementations are often faster in our experience). Thus, and will typically be (much) larger than and , and so in a basic/naive implementation of K-FAC, task 5 can dominate the overall per-iteration cost.
Fortunately, there are several possible ways to mitigate the cost of task 5. As mentioned above, one way is to perform the computations for each layer in parallel, and even simultaneously with the gradient computation and other tasks. In the case of our block-tridiagonal approximation to the inverse, one can avoid computing any SVDs or matrix square roots by using an iterative Stein-equation solver (see Appendix B). And there are also ways of reducing matrix-inversion (and even matrix square-root) to a short sequence of matrix-matrix multiplications using iterative methods (Pan and Schreiber, 1991). Furthermore, because the matrices in question only change slowly over time, one can consider hot-starting these iterative inversion methods from previous solutions. In the extreme case where is very large, one can also consider using low-rank + diagonal approximations of the and matrices maintained online (e.g. using a similar strategy as Le Roux et al. (2008)) from which inverses and/or SVDs can be more easily computed. Although based on our experience such approximations can, in some cases, lead to a substantial degradation in the quality of the updates.
While these ideas work reasonably well in practice, perhaps the simplest method, and the one we ended up settling on for our experiments, is to simply recompute the approximate Fisher inverse only every iterations (we used in our experiments). As it turns out, the curvature of the objective stays relatively stable during optimization, especially in the later stages, and so in our experience this strategy results in only a modest decrease in the quality of the updates.
Let and be matrices whose columns are the ’s and ’s (resp.) associated with the current mini-batch. Let denote the gradient of with respect to , shaped as a matrix (instead of a vector). The estimate of over the mini-batch is given by , which is of rank-. From Section 4.2, computing the amounts to computing . Substituting in our mini-batch estimate of gives
Note that the adjustment technique for described in Section 6.6 requires that, at every iterations, we compute 3 different versions of the update for each of 3 candidate values of . In an ideal implementation these could be computed in parallel with each other, although in the summary analysis below we will assume they are computed serially.
Summarizing, we have that with all of the various efficiency improvements discussed in this section, the average per-iteration computational cost of K-FAC, in terms of serial arithmetic operations, is
where are flag variables indicating whether momentum and the block-tridiagonal inverse approximation (resp.) are used.
Plugging in the values of these various constants that we used in our experiments, for the block-diagonal inverse approximation () this becomes
and for the block-tridiagonal inverse approximation ()
which is to be compared to the per-iteration cost of SGD, as given by
Pseudocode for K-FAC
Algorithm 2 gives high-level pseudocode for the K-FAC method, with the details of how to perform the computations required for each major step left to their respective sections.
Invariance Properties and the Relationship to Whitening and Centering
When computed with the exact Fisher, the natural gradient specifies a direction in the space of predictive distributions which is invariant to the specific way that the model is parameterized. This invariance means that the smooth path through distribution space produced by following the natural gradient with infinitesimally small steps will be similarly invariant.
For a practical natural gradient based optimization method which takes large discrete steps in the direction of the natural gradient, this invariance of the optimization path will only hold approximately. As shown by Martens (2014), the approximation error will go to zero as the effects of damping diminish and the reparameterizing function tends to a locally linear function. Note that the latter will happen as becomes smoother, or the local region containing the update shrinks to zero.
Because K-FAC uses an approximation of the natural gradient, these invariance results are not applicable in our case. Fortunately, as was shown by Martens (2014), one can establish invariance of an update direction with respect to a given reparameterization of the model by verifying certain simple properties of the curvature matrix used to compute the update. We will use this result to show that, under the assumption that damping is absent (or negligible in its affect), K-FAC is invariant to a broad and natural class of transformations of the network.
This class of transformations is given by the following modified network definition (c.f. the definition in Section 2.1):
The following theorem describes the main technical result of this section.
There exists an invertible linear function so that , and thus the transformed network can be viewed as a reparameterization of the original network by . Moreover, additively updating by or in the original network is equivalent to additively updating by or (resp.) in the transformed network, in the sense that .
This immediately implies the following corollary which characterizes the invariance of a basic version of K-FAC to the given class of network transformations.
The optimization path taken by K-FAC (using either of our Fisher approximations or ) through the space of predictive distributions is the same for the default network as it is for the transformed network (where the ’s and ’s remain fixed). This assumes the use of an equivalent initialization (), and a basic version of K-FAC where damping is absent or negligible in effect, momentum is not used, and where the learning rates are chosen in a way that is independent of the network’s parameterization.
While this corollary assumes that the ’s and ’s are fixed, if we relax this assumption so that they are allowed to vary smoothly with , then will be a smooth function of , and so as discussed in Martens (2014), invariance of the optimization path will hold approximately in a way that depends on the smoothness of (which measures how quickly the ’s and ’s change) and the size of the update. Moreover, invariance will hold exactly in the limit as the learning rate goes to 0.
Note that the network transformations can be interpreted as replacing the network’s nonlinearity at each layer with a “transformed” version . So since the well-known logistic sigmoid and tanh functions are related to each other by such a transformation, an immediate consequence of Corollary 2 is that K-FAC is invariant to the choice of logistic sigmoid vs. tanh activation functions (provided that equivalent initializations are used and that the effect of damping is negligible, etc.).
Also note that because the network inputs are also transformed by , K-FAC is thus invariant to arbitrary affine transformations of the input, which includes many popular training data preprocessing techniques.
Many other natural network transformations, such as ones which “center” and normalize unit activities so that they have mean 0 and variance 1 can be described using diagonal choices for the ’s and ’s which vary smoothly with . In addition to being approximately invariant to such transformations (or exactly, in the limit as the step size goes to 0), K-FAC is similarly invariant to a more general class of such transformations, such as those which transform the units so that they have a mean of 0, so they are “centered”, and a covariance matrix of , so they are “whitened”, which is a much stronger condition than the variances of the individual units each being 1.
In the case where we use the block-diagonal approximation and compute updates without damping, Theorem 1 affords us an additional elegant interpretation of what K-FAC is doing. In particular, the updates produced by K-FAC end up being equivalent to those produced by standard gradient descent using a network which is transformed so that the unit activities and the unit-gradients are both centered and whitened (with respect to the model’s distribution). This is stated formally in the following corollary.
Additively updating by in the original network is equivalent to additively updating by the gradient descent update (where as in Theorem 1) in a transformed version of the network where the unit activities and the unit-gradients are both centered and whitened with respect to the model’s distribution.
Related Work
The Hessian-free optimization method of Martens (2010) uses linear conjugate gradient (CG) to optimize local quadratic models of the form of eqn. 5 (subject to an adaptive Tikhonov damping technique) in lieu of directly solving it using matrix inverses. As discussed in the introduction, the main advantages of K-FAC over HF are twofold. Firstly, K-FAC uses an efficiently computable direct solution for the inverse of the curvature matrix and thus avoids the costly matrix-vector products associated with running CG within HF. Secondly, it can estimate the curvature matrix from a lot of data by using an online exponentially-decayed average, as opposed to relatively small-sized fixed mini-batches used by HF. The cost of doing this is of course the use of an inexact approximation to the curvature matrix.
Le Roux et al. (2008) proposed a neural network optimization method known as TONGA based on a block-diagonal approximation of the empirical Fisher where each block corresponds to the weights associated with a particular unit. By contrast, K-FAC uses much larger blocks, each of which corresponds to all the weights within a particular layer. The matrices which are inverted in K-FAC are roughly the same size as those which are inverted in TONGA, but rather than there being one per unit as in TONGA, there are only two per layer. Therefore, K-FAC is significantly less computationally intensive than TONGA, despite using what is arguably a much more accurate approximation to the Fisher. Note that to help mitigate the cost of the many matrix inversions it requires, TONGA approximates the blocks as being low-rank plus a diagonal term, although this introduces further approximation error.
Centering methods work by either modifying the gradient (Schraudolph, 1998) or dynamically reparameterizing the network itself (Raiko et al., 2012; Vatanen et al., 2013; Wiesler et al., 2014), so that various unit-wise scalar quantities like the activities (the ’s) and local derivatives (the ’s) are 0 on average (i.e. “centered”), as they appear in the formula for the gradient. Typically, these methods require the introduction of additional “skip” connections (which bypass the nonlinearities of a given layer) in order to preserve the expressive power/efficiency of the network after these transformations are applied.
It is argued by Raiko et al. (2012) that the application of the centering transformation makes the Fisher of the resulting network closer to a diagonal matrix, and thus makes its gradient more closely resemble its natural gradient. However, this argument uses the strong approximating assumption that the correlations between various network-dependent quantities, such as the activities of different units within a given layer, are zero. In our notation, this would be like assuming that the ’s are diagonal, and that the ’s are rank-1 plus a diagonal term. Indeed, using such an approximation within the block-diagonal version of K-FAC would yield an algorithm similar to standard centering, although without the need for skip connections (and hence similar to the version of centering proposed by Wiesler et al. (2014)).
As shown in Corollary 3, K-FAC can also be interpreted as using the gradient of a transformed network as its update direction, although one in which the ’s and ’s are both centered and whitened (with respect to the model’s distribution). Intuitively, it is this whitening which accounts for the correlations between activities (or back-propagated gradients) within a given layer.
Ollivier (2013) proposed a neural network optimization method which uses a block-diagonal approximation of the Fisher, with the blocks corresponding to the incoming weights (and bias) of each unit. This method is similar to TONGA, except that it approximates the Fisher instead of the empirical Fisher (see Martens (2014) for a discussion of the difference between these). Because computing blocks of the Fisher is expensive (it requires backpropagations, where is the number of output units), this method uses a biased deterministic approximation which can be computed more efficiently, and is similar in spirit to the deterministic approximation used by LeCun et al. (1998). Note that while such an approximation could hypothetically be used within K-FAC to compute the ’s, we have found that our basic unbiased stochastic approximation works nearly as well as the exact values in practice.
The work most closely related to ours is that of Heskes (2000), who proposed an approximation of the Fisher of feed-forward neural networks similar to our Kronecker-factored block-diagonal approximation from Section 4.2, and used it to derive an efficient approximate natural-gradient based optimization method by exploiting the identity . K-FAC differs from Heskes’ method in several important ways which turn out to be crucial to it working well in practice.
In Heskes’ method, update damping is accomplished using a basic factored Tikhonov technique where is added to each and for a fixed parameter which is set by hand. By contrast, K-FAC uses a factored Tikhonov technique where adapted dynamically as described in Section 6.6, combined with a re-scaling technique based on a local quadratic model computed using the exact Fisher (see Section 6.4). Note that the adaptation of is important since what constitutes a good or even merely acceptable value of will change significantly over the course of optimization. And the use of our re-scaling technique, or something similar to it, is also crucial as we have observed empirically that basic Tikhonov damping is incapable of producing high quality updates by itself, even when is chosen optimally at each iteration (see Figure 7 of Section 6.4).
Also, while Heskes’ method computes the ’s exactly, K-FAC uses a stochastic approximation which scales efficiently to neural networks with much higher-dimensional outputs (see Section 5).
Other advances we have introduced include the more accurate block-tridiagonal approximation to the inverse Fisher, a parameter-free type of momentum (see Section 7), online estimation of the and matrices, and various improvements in computational efficiency (see Section 8). We have found that each of these additional elements is important in order for K-FAC to work as well as it does in various settings.
Concurrently with this work Povey et al. (2015) has developed a neural network optimization method which uses a block-diagonal Kronecker-factored approximation similar to the one from Heskes (2000). This approach differs from K-FAC in numerous ways, including its use of the empirical Fisher (which doesn’t work as well as the standard Fisher in our experience – see Section 5), and its use of only a basic factored Tikhonov damping technique without adaptive re-scaling or any form of momentum. One interesting idea introduced by Povey et al. (2015) is a particular method for maintaining an online low-rank plus diagonal approximation of the factor matrices for each block, which allows their inverses to be computed more efficiently (although subject to an approximation). While our experiments with similar kinds of methods for maintaining such online estimates found that they performed poorly in practice compared to the solution of refreshing the inverses only occasionally (see Section 8), the particular one developed by Povey et al. (2015) could potentially still work well, and may be especially useful for networks with very wide layers.
Heskes’ interpretation of the block-diagonal approximation
Heskes (2000) discussed an alternative interpretation of the block-diagonal approximation which yields some useful insight to complement our own theoretical analysis. In particular, he observed that the block-diagonal Fisher approximation is the curvature matrix corresponding to the following quadratic function which measures the difference between the new parameter value and the current value :
Here, , and the ’s and ’s are determined by and are independent of (which determines the ’s).
can be interpreted as a reweighted sum of squared changes of each of the ’s. The reweighing matrix is given by
where is the network’s predictive distribution as parameterized by , and is its Fisher information matrix, and where the expectation is taken w.r.t. the distribution on (as induced by the distribution on the network’s input ). Thus, the effect of reweighing by the ’s is to (approximately) translate changes in into changes in the predictive distribution over , although using the expected/average Fisher instead of the more specific Fisher .
Interestingly, if one used instead of in the expression for , then would correspond to a basic layer-wise block-diagonal approximation of where the blocks are computed exactly (i.e. without the Kronecker-factorizing approximation introduced in Section 3). Such an approximate Fisher would have the interpretation of being the Hessian w.r.t. of either of the measures
It is not clear whether , with its Kronecker-factorizing structure can similarly be interpreted as the Hessian of such a self-evidently intrinsic measure. If it could be, then this would considerably simplify the proof of our Theorem 1 (e.g. using the techniques of Arnold et al. (2011)). Note that itself doesn’t work, as it isn’t obviously intrinsic. Despite this, as shown in Section 10, both and our more advanced approximation produce updates which have strong invariance properties.
Experiments
As our baseline we used the version of SGD with momentum based on Nesterov’s Accelerated Gradient (Nesterov, 1983) described in Sutskever et al. (2013), which was calibrated to work well on these particular deep autoencoder problems. For each problem we followed the prescription given by Sutskever et al. (2013) for determining the learning rate, and the increasing schedule for the decay parameter . We did not compare to methods based on diagonal approximations of the curvature matrix, as in our experience such methods tend not perform as well on these kinds of optimization problems as the baseline does (an observation which is consistent with the findings of Schraudolph (2002); Zeiler (2013)).
Our implementation of K-FAC used most of the efficiency improvements described in Section 8, except that all “tasks” were computed serially (and thus with better engineering and more hardware, a faster implementation could likely be obtained). Because the mini-batch size tended to be comparable to or larger than the typical/average layer size , we did not use the technique described at the end of Section 8 for accelerating the computation of the approximate inverse, as this only improves efficiency in the case where , and will otherwise decrease efficiency.
Both K-FAC and the baseline were implemented using vectorized MATLAB code accelerated with the GPU package Jacket. The code for K-FAC is available for downloadhttp://www.cs.toronto.edu/~jmartens/docs/KFAC3-MATLAB.zip. All tests were performed on a single computer with a 4.4 Ghz 6 core Intel CPU and an NVidia GTX 580 GPU with 3GB of memory. Each method used the same initial parameter setting, which was generated using the “sparse initialization” technique from Martens (2010) (which was also used by Sutskever et al. (2013)).
To help mitigate the detrimental effect that the noise in the stochastic gradient has on the convergence of the baseline (and to a lesser extent K-FAC as well) we used a exponentially decayed iterate averaging approach based loosely on Polyak averaging (e.g. Swersky et al., 2010). In particular, at each iteration we took the “averaged” parameter estimate to be the previous such estimate, multiplied by , plus the new iterate produced by the optimizer, multiplied by , for . Since the training error associated with the optimizer’s current iterate may sometimes be lower than the training error associated with the averaged estimate (which will often be the case when the mini-batch size is very large), we report the minimum of these two quantities.
To be consistent with the numbers given in previous papers we report the reconstruction error instead of the actual objective function value (although these are almost perfectly correlated in our experience). And we report the error on the training set as opposed to the test set, as we are chiefly interested in optimization speed and not the generalization capabilities of the networks themselves.
In our first experiment we examined the relationship between the mini-batch size and the per-iteration rate of progress made by K-FAC and the baseline on the MNIST problem. The results from this experiment are plotted in Figure 9. They strongly suggest that the per-iteration rate of progress of K-FAC tends to a superlinear function of (which can be most clearly seen by examining the plots of training error vs training cases processed), which is to be contrasted with the baseline, where increasing has a much smaller effect on the per-iteration rate of progress, and with K-FAC without momentum, where the per-iteration rate of progress seems to be a linear or slightly sublinear function of . It thus appears that the main limiting factor in the convergence of K-FAC (with momentum applied) is the noise in the gradient, at least in later stages of optimization, and that this is not true of the baseline to nearly the same extent. This would seem to suggest that K-FAC, much more than SGD, would benefit from a massively parallel distributed implementation which makes use of more computational resources than a single GPU.
But even in the single CPU/GPU setting, the fact that the per-iteration rate of progress tends to a superlinear function of , while the per-iteration computational cost of K-FAC is a roughly linear function of , suggests that in order to obtain the best per-second rate of progress with K-FAC, we should use a rapidly increasing schedule for . To this end we designed an exponentially increasing schedule for , given by , where is the current iteration, , and where is chosen so that . The approach of increasing the mini-batch size in this way is analyzed by Friedlander and Schmidt (2012). Note that for other neural network optimization problems, such as ones involving larger training datasets than these autoencoder problems, a more slowly increasing schedule, or one that stops increasing well before reaches , may be more appropriate. One may also consider using an approach similar to that of Byrd et al. (2012) for adaptively determining a suitable mini-batch size.
In our second experiment we evaluated the performance of our implementation of K-FAC versus the baseline on all 3 deep autoencoder problems, where we used the above described exponentially increasing schedule for for K-FAC, and a fixed setting of for the baseline and momentum-less K-FAC (which was chosen from a small range of candidates to give the best overall per-second rate of progress). The relatively high values of chosen for the baseline ( for CURVES, and for MNIST and FACES, compared to the which was used by Sutskever et al. (2013)) reflect the fact that our implementation of the baseline uses a high-performance GPU and a highly optimized linear algebra package, which allows for many training cases to be efficiently processed in parallel. Indeed, after a certain point, making much smaller didn’t result in a significant reduction in the baseline’s per-iteration computation time.
Note that in order to process the very large mini-batches required for the exponentially increasing schedule without overwhelming the memory of the GPU, we partitioned the mini-batches into smaller “chunks” and performed all computations involving the mini-batches, or subsets thereof, one chunk at a time.
The results from this second experiment are plotted in Figures 10 and 11. For each problem K-FAC had a per-iteration rate of progress which was orders of magnitude higher than that of the baseline’s (Figure 11), provided that momentum was used, which translated into an overall much higher per-second rate of progress (Figure 10), despite the higher cost of K-FAC’s iterations (due mostly to the much larger mini-batch sizes used). Note that Polyak averaging didn’t produce a significant increase in convergence rate of K-FAC in this second experiment (actually, it hurt a bit) as the increasing schedule for provided a much more effective (although expensive) solution to the problem of noise in the gradient.
The importance of using some form of momentum on these problems is emphasized in these experiments by the fact that without the momentum technique developed in Section 7, K-FAC wasn’t significantly faster than the baseline (which itself used a strong form of momentum). These results echo those of Sutskever et al. (2013), who found that without momentum, SGD was orders of magnitude slower on these particular problems. Indeed, if we had included results for the baseline without momentum they wouldn’t even have appeared in the axes boundaries of the plots in Figure 10.
Recall that the type of momentum used by K-FAC compensates for the inexactness of our approximation to the Fisher by allowing K-FAC to build up a better solution to the exact quadratic model minimization problem (defined using the exact Fisher) across many iterations. Thus, if we were to use a much stronger approximation to the Fisher when computing our update proposals , the benefit of using this type of momentum would have likely been much smaller than what we observed. One might hypothesize that it is the particular type of momentum used by K-FAC that is mostly responsible for its advantages over the SGD baseline. However in our testing we found that for SGD the more conventional type of momentum used by Sutskever et al. (2013) performs significantly better.
From Figure 11 we can see that the block-tridiagonal version of K-FAC has a per-iteration rate of progress which is typically 25% to 40% larger than the simpler block-diagonal version. This observation provides empirical support for the idea that the block-tridiagonal approximate inverse Fisher is a more accurate approximation of than the block-diagonal approximation . However, due to the higher cost of the iterations in the block-tridiagonal version, its overall per-second rate of progress seems to be only moderately higher than the block-diagonal version’s, depending on the problem.
Note that while matrix-matrix multiplication, matrix inverse, and SVD computation all have the same computational complexity, in practice their costs differ significantly (in increasing order as listed). Computation of the approximate Fisher inverse, which is performed in our experiments once every 20 iterations (and for the first 3 iterations), requires matrix inverses for the block-diagonal version, and SVDs for the block-tridiagonal version. For the FACES problem, where the layers can have as many as 2000 units, this accounted for a significant portion of the difference in the average per-iteration computational cost of the two versions (as these operations must be performed on sized matrices).
While our results suggest that the block-diagonal version is probably the better option overall due to its greater simplicity (and comparable per-second progress rate), the situation may be different given a more efficient implementation of K-FAC where the more expensive SVDs required by the tri-diagonal version are computed approximately and/or in parallel with the other tasks, or perhaps even while the network is being optimized.
Our results also suggest that K-FAC may be much better suited than the SGD baseline for a massively distributed implementation, since it would require far fewer synchronization steps (by virtue of the fact that it requires far fewer iterations).
Conclusions and future directions
In this paper we developed K-FAC, an approximate natural gradient-based optimization method. We started by developing an efficiently invertible approximation to a neural network’s Fisher information matrix, which we justified via a theoretical and empirical examination of the statistics of the gradient of a neural network. Then, by exploiting the interpretation of the Fisher as an approximation of the Hessian, we designed a developed a complete optimization algorithm using quadratic model-based damping/regularization techniques, which yielded a highly effective and robust method virtually free from the need for hyper-parameter tuning. We showed the K-FAC preserves many of natural gradient descent’s appealing theoretical properties, such as invariance to certain reparameterizations of the network. Finally, we showed that K-FAC, when combined with a form of momentum and an increasing schedule for the mini-batch size , far surpasses the performance of a well-tuned version of SGD with momentum on difficult deep auto-encoder optimization benchmarks (in the setting of a single GPU machine). Moreover, our results demonstrated that K-FAC requires orders of magnitude fewer total updates/iterations than SGD with momentum, making it ideally suited for a massively distributed implementation where synchronization is the main bottleneck.
Some potential directions for future development of K-FAC include:
a better/more-principled handling of the issue of gradient stochasticity than a pre-determined increasing schedule for
extensions of K-FAC to recurrent or convolutional architectures, which may require specialized approximations of their associated Fisher matrices
an implementation that better exploits opportunities for parallelism described in Section 8
exploitation of massively distributed computation in order to compute high-quality estimates of the gradient
Acknowledgments
We gratefully acknowledge support from Google, NSERC, and the University of Toronto. We would like to thank Ilya Sutskever for his constructive comments on an early draft of this paper.
References
Appendix A Derivation of the expression for the approximation from Section 3.1
The only specific property of the distribution over , , , and which we will require to do this is captured by the following lemma.
Suppose is a scalar variable which is independent of when conditioned on the network’s output , and is some intermediate quantity computed during the evaluation of (such as the activities of the units in some layer). Then we have
Our proof of this lemma (which is at the end of this section) makes use of the fact that the expectations are taken with respect to the network’s predictive distribution as opposed to the training distribution .
Intuitively, this lemma says that the intermediate quantities computed in the forward pass of Algorithm 1 (or various functions of these) are statistically uncorrelated with various derivative quantities computed in the backwards pass, provided that the targets are sampled according to the network’s predictive distribution (instead of coming from the training set). Valid choices for include , for , and products of these. Examples of invalid choices for include expressions involving , since these will depend on the derivative of the loss, which is not independent of given .
According to a well-known general formula relating moments to cumulants we may write as a sum of 15 terms, each of which is a product of various cumulants corresponding to one of the 15 possible ways to partition the elements of into non-overlapping sets. For example, the term corresponding to the partition is .
Observing that 1st-order cumulants correspond to means and 2nd-order cumulants correspond to covariances, for Lemma 4 gives
where , and (so that ). And similarly for it gives
Using these identities we can eliminate 10 of the terms.
The remaining expression for is thus
That the inner expectation above is follows from the fact that the expected score of a distribution, when taken with respect to that distribution, is .
It is well known that , and that matrix-vector products with this matrix can thus be computed as , where is the matrix representation of (so that ).
Somewhat less well known is that there are also formulas for which can be efficiently computed and likewise give rise to efficient methods for computing matrix-vector products.
First, note that is equivalent to , which is equivalent to the linear matrix equation , where and . This is known as a generalized Stein equation, and different examples of it have been studied in the control theory literature, where they have numerous applications. For a recent survey of this topic, see Simoncini (2014).
One well-known class of methods called Smith-type iterations (Smith, 1968) involve rewriting this matrix equation as a fixed point iteration and then carrying out this iteration to convergence. Interestingly, through the use of a special squaring trick, one can simulate of these iterations with only matrix-matrix multiplications.
Another class of methods for solving Stein equations involves the use of matrix decompositions (e.g. Chu, 1987; Gardiner et al., 1992). Here we will present such a method particularly well suited for our application, as it produces a formula for , which after a fixed overhead cost (involving the computation of SVDs and matrix square roots), can be repeatedly evaluated for different choices of using only a few matrix-matrix multiplications.
We will assume that , , , and are symmetric positive semi-definite, as they always are in our applications. We have
Inverting both sides of the above equation gives
Using the symmetric eigen/SVD-decomposition, we can write and , where for the are diagonal matrices and the are unitary matrices.
Note that both and are diagonal matrices, and thus the middle matrix is just the inverse of a diagonal matrix, and so can be computed efficiently.
where and .
And so matrix-vector products with can be computed as
where denotes element-wise division of by , , and is the vector of ones (sized as appropriate). Note that if we wish to compute multiple matrix-vector products with (as we will in our application), the quantities , , and only need to be computed the first time, thus reducing the cost of any future such matrix-vector products, and in particular avoiding any additional SVD computations.
In the considerably simpler case where and are both scalar multiples of the identity, and is the product of these multiples, we have
where and are the symmetric eigen/SVD-decompositions of and , respectively. And so matrix-vector products with can be computed as
where is the Jacobian of and is the Fisher information matrix of the network’s predictive distribution , evaluated at (where we treat as the “parameters”).
To compute the matrix-vector product as estimated from a mini-batch we simply compute for each in the mini-batch, and average the results. This latter operation can be computed in 3 stages (e.g. Martens, 2014), which correspond to multiplication of the vector first by , then by , and then by .
Multiplication by can be performed by a forward pass which is like a linearized version of the standard forward pass of Algorithm 1. As is usually diagonal or diagonal plus rank-1, matrix-vector multiplications with it are cheap and easy. Finally, multiplication by can be performed by a backwards pass which is essentially the same as that of Algorithm 1. See Schraudolph (2002); Martens (2014) for further details.
The naive way of computing is to compute as above, and then compute the inner product of with . Additionally computing and would require another such matrix-vector product .
However, if we instead just compute the matrix-vector products (which requires only half the work of computing ), then computing as is essentially free. And with computed, we can similarly obtain as and as .
This trick thus reduces the computational cost associated with computing these various scalars by roughly half.
Appendix D Proofs for Section 10
First we will show that the given network transformation can be viewed as reparameterization of the network according to an invertible linear function .
If the transformed network uses in place of then we have
which we can prove by a simple induction. First note that by definition. Then, assuming by induction that , we have
The following lemma is adapted from Martens (2014) (see the section titled “A critical analysis of parameterization invariance”).
Let be some invertible affine function mapping to , which reparameterizes the objective as . Suppose that and are invertible matrices satisfying
Then, additively updating by is equivalent to additively updating by , in the sense that .
Using it is straightforward to verify that
Because and the fact that the networks compute the same outputs (so the loss derivatives are identical), we have by the chain rule that, , and therefore
We now turn our attention to the (see Section 4.3 for the relevant notation).
Inverting both sides gives .
Combining and we have
Inverting both sides gives as required.
First note that a network which is transformed so that and will satisfy the required properties. To see this, note that means that is whitened with respect to the model’s distribution by definition (since the expectation is taken with respect to the model’s distribution), and furthermore we have that by default (e.g. using Lemma 4), so is centered. And since is the square submatrix of which leaves out the last row and column, we also have that and so is whitened. Finally, observe that is given by the final column (or row) of , excluding the last entry, and is thus equal to , and so is centered.
Next, we note that if and then
and so is indeed a standard gradient descent update.
Finally, we observe that there are choices of and which will make the transformed model satisfy and . In particular, from the proof of Theorem 1 we have that and , and so taking and works.