Metric-Free Natural Gradient for Joint-Training of Boltzmann Machines
Guillaume Desjardins, Razvan Pascanu, Aaron Courville, Yoshua Bengio
Introduction
Boltzmann Machines (BM) have become a popular method in Deep Learning for performing feature extraction and probability modeling. The emergence of these models as practical learning algorithms stems from the development of efficient training algorithms, which estimate the negative log-likelihood gradient by either contrastive or stochastic approximations. However, the success of these models has for the most part been limited to the Restricted Boltzmann Machine (RBM) , whose architecture allows for efficient exact inference. Unfortunately, this comes at the cost of the model’s representational capacity, which is limited to a single layer of latent variables. The Deep Boltzmann Machine (DBM) addresses this by defining a joint energy function over multiple disjoint layers of latent variables, where interactions within a layer are prohibited. While this affords the model a rich inference scheme incorporating top-down feedback, it also makes training much more difficult, requiring until recently an initial greedy layer-wise pretraining scheme. Since, Montavon and Muller 2012 have shown that this difficulty stems from an ill-conditioning of the Hessian matrix, which can be addressed by a simple reparameterization of the DBM energy function, a trick called centering (an analogue to centering and skip-connections found in the deterministic neural network literature ). As the barrier to joint-training Joint-training refers to the act of jointly optimizing (the concatenation of all model parameters, across all layers of the DBM) through maximum likelihood. This is in contrast to , where joint-training is preceded by a greedy layer-wise pretraining strategy. is overcoming a challenging optimization problem, it is apparent that second-order gradient methods might prove to be more effective than simple stochastic gradient methods. This should prove especially important as we consider models with increasingly complex posteriors or higher-order interactions between latent variables.
The Natural Gradient
The main insight behind the natural gradient is that the space of all probability distributions forms a Riemannian manifold. Learning, which typically proceeds by iteratively adapting the parameters to fit an empirical distribution , thus traces out a path along this manifold. An immediate consequence is that following the direction of steepest descent in the original Euclidean parameter space does not correspond to the direction of steepest descent along . To do so, one needs to account for the metric describing the local geometry of the manifold, which is given by the Fisher Information matrix , shown in Equation 4. While this metric is typically derived from Information Geometry, a derivation more accessible to a machine learning audience can be obtained as follows.
The natural gradient aims to find the search direction which minimizes a given objective function, such that the Kullback–Leibler divergence remains constant throughout optimization. This constraint ensures that we make constant progress regardless of the curvature of the manifold and enforces an invariance to the parameterization of the model. The natural gradient for maximum likelihood can thus be formalized as:
In order to derive a useful parameter update rule, we will consider the KL divergence under the assumption . We also assume we have a discrete and bounded domain over which we define the probability mass function When clear from context, we will drop the argument of to save space. . Taking the Taylor series expansion of around , and denoting as the column vector of partial derivatives with as the -th entry, and the Hessian matrix with in position , we have:
with the transition stemming from the fact that . Replacing the objective function of Equation 1 by its first-order Taylor expansion and rewriting the constraint as a Lagrangian, we arrive at the following formulation for , the loss function which the natural gradient seeks to minimize.
Setting to zero yields the natural gradient direction :
While its form is reminiscent of the Newton direction, the natural gradient multiplies the estimated gradient by the inverse of the expected Hessian of (Equation 3) or equivalently by the Fisher Information matrix (FIM, Equation 4). The equivalence between both expressions can be shown trivially, with the details appearing in the Appendix. We stress that both of these expectations are computed with respect to the model distribution, and thus computing the metric does not involve the empirical distribution in any way. The FIM for Boltzmann Machines is thus not equal to the uncentered covariance of the maximum likelihood gradients. In the following, we pursue our derivation from the form given in Equation 4.
2 Natural Gradient for Boltzmann Machines
Starting from the expression of found in Equation 3, we can derive the natural gradient metric for Boltzmann Machines.
The natural gradient metric for first-order BMs takes on a surprisingly simple form: it is the expected Hessian of the log-partition function. With a few lines of algebra (whose details are presented in the Appendix), we can rewrite it as follows:
Discussion.
When computing the Taylor expansion of the KL divergence in Equation 2, we glossed over an important detail. Namely, how to handle latent variables in , a topic first discussed in . If , we could just as easily have derived the natural gradient by considering the constraint . Alternatively, since the distinction between visible and hidden units is entirely artificial (since the KL divergence does not involve the empirical distribution), we may simply wish to consider the distribution obtained by analytically integrating out a maximal number of random variables. In a DBM, this would entail marginalizing over all odd or even layers, a strategy employed with great success in the context of AIS . In this work however, we only consider the metric obtained by considering the divergence between the full joint distributions and .
Metric-Free Natural Gradient Implementation
For Boltzmann Machines, the matrix-vector product can be computed in a straightforward manner, without recourse to Pearlmutter’s R-operator . Starting from a sampling approximation to Equation 5, we simply push the dot product inside of the expectation as follows:
Experiments
We performed a proof-of-concept experiment to determine whether our Metric-Free Natural Gradient (MFNG) algorithm is suitable for joint-training of complex Boltzmann Machines. To this end, we compared our method to Stochastic Maximum Likelihood and a diagonal approximation of MFNG on a 3-layer Deep Boltzmann Machine trained on MNIST . All algorithms were run in conjunction with the centering strategy of Montavon and Muller 2012, which proved crucial to successfully joint-train all layers of the DBM (even when using MFNG) The centering coefficients were initialized as in , but were otherwise held fixed during training.. We chose a small 3-layer DBM with 784-400-100 units at the first, second and third layers respectively, to be comparable to . Hyper-parameters were varied as follows. For inference, we ran iterations of either mean-field as implemented in or Gibbs sampling. The learning rate was kept fixed during training and chosen from the set . For MinRes, we set the damping coefficient to and used a fixed tolerance of (used to determine convergence). Finally, we tested all algorithms on minibatch sizes of either , or elements We expect larger minibatch sizes to be preferable, however simulating this number of Markov chains in parallel (on top of all other memory requirements) was sufficient to hit the memory bottlenecks of GPUs.. Finally, since we are comparing optimization algorithms, hyper-parameters were chosen based on the training set likelihood (though we still report the associated test errors). All experiments used the MinRes linear solver, both for its speed and its ability to return pseudo-inverses when faced with ill-conditioning.
Figure 1 (left) shows the likelihood as estimated by Annealed Importance Sampling as a function of the number of epochs While we do not report error margins for AIS likelihood estimates, the numbers proved robust to changes in the number of particles and temperatures being simulated. To obtain such robust estimates, we implemented all the tricks described in Salakhutdinov and Hinton 2009 and : a zero-weight base-rate model whose biases are set by maximum likelihood; interpolating distributions , with the target distribution; and finally analytical integration of all odd-layers.. Under this metric, MFNG achieves the fastest convergence, obtaining a training/test set likelihood of / nats after 94 epochs. In comparison, MFNG-diag obtains / nats and SML / nats in 100 epochs. The picture changes however when plotting likelihood as a function of CPU-time, as shown in Figure 1 (right). Given a wall-time of s for MFNG and SML, and s for MFNG-diag This discrepancy will be resolved in the next revision., SML is able to perform upwards of epochs, resulting in an impressive likelihood score of / . Note that these results were obtained on the binary-version of MNIST (thresholded at 0.5) in order to compare to . These results are therefore not directly comparable to , which binarizes the dataset through sampling (by treating each pixel activation as the probability of a Bernouilli distribution).
Figure 2 shows a breakdown of the algorithm runtime, for various components of the algorithm. These statistics were collected in the early stages of training, but are generally representative of the bigger picture. While the linear solver clearly dominates the runtime, there are a few interesting observations to make. For small models and batch sizes greater than , a single evaluation of appears to be of the same order of magnitude as a gradient evaluation. In all cases, this cost is smaller than that of sampling, which represents a non-negligible part of the total computational budget. This suggests that MFNG could become especially attractive for models which are expensive to sample from. Overall however, restricting the number of CG/MinRes iterations appears key to computational performance, which can be achieved by increasing the damping factor . How this affects convergence in terms of likelihood is left for future work.
Discussion and Future Work
While the wall-clock performance of MFNG is not currently competitive with SML, we believe there are still many avenues to explore to improve computational efficiency. Firstly, we performed almost no optimization of the various MinRes hyper-parameters. In particular, we ran the algorithm to convergence with a fixed tolerance of . While this typically resulted in relatively few iterations (around ), this level of precision might not be required (especially given the stochastic nature of the algorithm). Additionally, it could be worth exploiting the same strategy as HF where the linear solver is initialized by the solution found in the previous iteration. This may prove much more efficient than the current approach of initializing the solver with a zero vector. Pre-conditioning is also a well-known method for accelerating the convergence speed of linear solvers . Our implementation used a simple diagonal regularization of . The Jacobi preconditioner could be implemented easily however by computing the diagonal of in a first-pass.
Finally, while our single experiment offers little evidence in support of either conclusion, it may very well be possible that MFNG is simply not computationally efficient for DBMs, compared to SML with centering. In this case, it would be worth applying the method to either (i) models with known ill-conditioning, such as factored 3-rd order Boltzmann Machines or (ii) models and distributions exhibiting complex posterior distributions. In such scenarios, we may wish to maximize the use of the positive phase statistics (which were obtained at a high computational cost) by performing larger jumps in parameter space. It remains to be seen how this would interact with SML, where the burn-in period of the persistent chains is directly tied to the magnitude of .
Appendix
We include the following derivations for completeness.