Fast yet Simple Natural-Gradient Descent for Variational Inference in Complex Models
Mohammad Emtiyaz Khan, Didrik Nielsen
I Introduction
Modern machine-learning methods, such as deep learning, are capable of producing accurate predictions which has lead to their enormous recent success in fields, e.g., computer vision, speech recognition, and recommendation systems. However, this is not enough for other fields such as robotics and medical diagnostics where we also require an accurate estimate of confidence or uncertainty in the predictions. Bayesian inference provides such uncertainty measures by using the posterior distribution obtained using Bayes’ rule. Unfortunately, this computation requires integrating over all possible values of the model parameters, which is infeasible for large complex models such as Bayesian neural networks.
Sampling methods such as Markov Chain Monte Carlo usually converge slowly when applied to such large problems. In contrast, approximate Bayesian methods such as variational inference (VI) can scale to large problems by obtaining approximations to the posterior distribution by using an optimization method, e.g., stochastic-gradient descent (SGD) methods . These methods could provide reasonable approximations very quickly.
An issue in using SGD is that it ignores the information geometry of the posterior approximation (see Figure 1(a)). Recent approaches address this issue by using stochastic natural-gradient descent methods which exploit the Riemannian geometry of exponential-family approximations to improve the rate of convergence . Unfortunately, these approaches only apply to a restricted class of models known as conditionally-conjugate models, and do not work for nonconjugate models such as Bayesian neural networks.
This paper discusses some recent methods that generalize the use of natural gradients to such large and complex nonconjugate models. We show that, for exponential-family approximations, a duality between their natural and expectation parameter-spaces enables a simple natural-gradient update. The resulting updates are equivalent to a recently proposed method called Conjugate-computation Variational Inference (CVI) . An attractive feature of the method is that it naturally obtains local exponential-family approximations for individual model components. We discuss the application of the CVI method to Bayesian neural networks and show some recent results from a recent work demonstrating faster convergence of natural-gradient VI methods compared to gradient-based VI methods (see Figure 1(b)).
II Problem Formulation
In this section, we discuss the problem of variational inference and show how SGD can be used to optimize it. SGD ignores the geometry of the posterior approximations, and we discuss how natural-gradient methods address this issue. We end the section by mentioning issues with existing natural-gradient methods for variational inference.
We consider models Methods discussed in this paper apply to a more general class of models, e.g., the model class discussed in , but for clarity of presentation we focus on a restricted class. that take the following form:
where is a likelihood function which relates the model parameters to the ’th data-example \mbox{{\cal D}}_{i}, and p(\mbox{\mbox{}}) is the prior distribution which we assume to be an exponential-family distribution ,
where is a vector of sufficient statistics, \mbox{\mbox{}}_{0} is the natural-parameter vector, and A(\mbox{\mbox{}}) is the log-partition function. The model parameter is a random vector here and sometimes is referred to as the latent vector.
Example: Consider Bayesian neural networks (BNN) to model data \mbox{{\cal D}}_{i} that contains input \mbox{\mbox{}}_{i}\in\real^{D} and a scalar output . The vector is the vector of network weights. The likelihood p(\mbox{{\cal D}}_{i}|\mbox{\mbox{}}) could be an exponential-family distribution p(y_{i}|f_{\bf{z}}(\mbox{\mbox{}}_{i})) whose parameter is a neural network parameterized by . We assume an isotropic Gaussian prior p(\mbox{\mbox{}}):=\mbox{{\cal N}}(\mbox{\mbox{}}|0,\mbox{\mbox{}}/\tau) where is a scalar. Its natural parameters are \mbox{\mbox{}}_{0}:=\{0,-\tau\mbox{\mbox{}}/2\}. ∎
For such models, Bayesian approaches can estimate a measure of uncertainty by using the posterior distribution: p(\mbox{\mbox{}}|\mbox{{\cal D}}):=p(\mbox{{\cal D}}|\mbox{\mbox{}})p(\mbox{\mbox{}})/p(\mbox{{\cal D}}). This requires computation of the normalization constant p(\mbox{{\cal D}})=\int p(\mbox{{\cal D}}|\mbox{\mbox{}})p(\mbox{\mbox{}})d\mbox{\mbox{}} which unfortunately is difficult to compute in models such as Bayesian neural networks. One source of difficulty is that the likelihood p(\mbox{{\cal D}}_{i}|\mbox{\mbox{}}) does not take the same form as the prior with respect to , or, in other words, the model is nonconjugate . As a result, the product p(\mbox{{\cal D}}|\mbox{\mbox{}})p(\mbox{\mbox{}}) does not take a form with which p(\mbox{{\cal D}}) can be easily computed.
Variational inference (VI) simplifies the problem by approximating p(\mbox{\mbox{}}|\mbox{{\cal D}}) with a distribution q(\mbox{\mbox{}}) whose normalizing constant is relatively easier to compute. In models (1), a straightforward choice is to choose q(\mbox{\mbox{}}) to be of the same parametric This restriction may not lead to a suboptimal approximation, e.g., in mean-field approximation in conjugate exponential-family models, the optimal form according to the variational objective turns out to be an exponential-family approximation . form as the prior p(\mbox{\mbox{}}) but with a different natural-parameter vector , i.e., q_{\lambda}(\mbox{\mbox{}}):=h(\mbox{\mbox{}})\exp[\mbox{\mbox{}}^{\top}\mbox{\mbox{}}(\mbox{\mbox{}})-A(\mbox{\mbox{}})]. The parameter can be obtained by maximizing the variational objective which is also a lower bound to p(\mbox{{\cal D}}) ,
where is the set of valid variational parameters. Intuitively, the first term favors q_{\lambda}(\mbox{\mbox{}}) which is close to the prior p(\mbox{\mbox{}}) while the second term favors those that obtain high expected log-likelihood values. The variational objective has a very familiar form similar to many other regularized optimization problems in machine learning .
Example: In the BNN example, we can choose q_{\lambda}(\mbox{\mbox{}})=\mbox{{\cal N}}(\mbox{\mbox{}}|\mbox{\mbox{}},\mbox{\mbox{}}) where is the mean and is the covariance. The natural-parameter vector is \mbox{\mbox{}}:=\{\mbox{\mbox{}}^{-1}\mbox{\mbox{}},-\mbox{\frac{1}{2}}\mbox{\mbox{}}^{-1}\}, and our goal in VI is to maximize with respect to these parameters. ∎
II-B VI with Gradient Descent
A straightforward approach to maximize is to use a gradient-based method, e.g., the following stochastic-gradient descent (SGD) algorithm:
where is the iteration number, is a step size, and \widehat{\nabla}_{\lambda}\mathcal{L}(\mbox{\mbox{}}_{t}) is a stochastic estimate of the derivative of at \mbox{\mbox{}}=\mbox{\mbox{}}_{t} (the ‘hat’ here indicates a stochastic estimate). Such stochastic gradients can be easily computed using methods such as REINFORCE and the reparameterization trick . This results in a simple but powerful approach which applies to many models and scales to large data.
Despite this, a direct application of SGD to optimize \mathcal{L}(\mbox{\mbox{}}) is problematic because SGD ignores the information geometry of the distribution q_{\lambda}(\mbox{\mbox{}}). To see this, we can rewrite (4) as,
Equivalence can be established by taking the derivative and setting to 0. The equation (5) implies that SGD moves in the direction of the gradient while remaining close, in terms of the Euclidean distance, to the previous \mbox{\mbox{}}_{t}. However, the Euclidean distance between natural parameters is not appropriate because is the parameter of a distribution and the Euclidean distance is often a poor measure of dissimilarity between distributions. This is illustrated in Figure 1(a). A more informative measure such as a Kullback-Leibler (KL) divergence, which directly measures the distance between distributions, might be more appropriate.
II-C VI with Natural-Gradient Descent
The issue discussed above can be addressed by using natural-gradient methods that exploit the information geometry of . An exponential-family distribution induces a Riemannian manifold with a metric defined by the Fisher Information Matrix (FIM) , e.g. the FIM can be obtained as follows in the natural parameterization,
Natural-gradient descent modifies the SGD step (5) by using the Riemannian metric instead of the Euclidean distance,
where is a scalar step size. This results in an update similar to the SGD update shown in (4),
where the stochastic gradient is scaled by the FIM. The scaled stochastic-gradient is referred to as the stochastic natural gradient defined as follows:
Natural gradients are also naturally suited for VI in certain class of models. A recent work in shows that for conjugate exponential-family models, natural-gradients with respect to the natural-parameterization take a very simple form. For example, consider the first term in (3) which consists of the ratio of two terms that are conjugate to each other. The natural-gradient then is equal to the difference in the natural parameter of the two terms (see Eq. 41 in for more details):
The above natural gradient does not require computation of the FIM, which is surprising. It is natural to ask whether a similar expression is possible when the model contains nonconjugate terms? We show that it is possible to do so if we perform natural-gradient descent in the natural parameter space, but not if we do it in the space of expectation parameters.
III Natural Gradients with Exponential Family
In this section, we show that natural gradient with respect to the natural parameters can be obtained by computing the gradient with respect to the expectation parameter. In the next section, we will show that this enables a simple natural-gradient update which does not require an explicit inversion of FIM.
For an exponential-family in the minimal representation, the natural gradient with respect to is equal to the gradient with respect to , and vice versa, i.e.,
Proof: Using chain rule, we can rewrite the derivative with respect to in terms of :
It is well known that the second derivative of A(\mbox{\mbox{}}) is equal to the FIM for exponential-family distribution, i.e., \mbox{\mbox{}}(\mbox{\mbox{}}):=\nabla_{\lambda\lambda}^{2}A(\mbox{\mbox{}}) . This matrix is invertible when the representation is minimal. Therefore multiplying the above equation with inverse of \mbox{\mbox{}}(\mbox{\mbox{}}) gives us the first equality. Since the FIM with respect to is inverse of the FIM with respect to , the second equality is immediate. ∎
This result is a consequence of a relationship between and . The two vectors are related through the Legendre transform which is the following transformation \mbox{\mbox{}}=\nabla A(\mbox{\mbox{}}). Since A(\mbox{\mbox{}}) is a convex function, the space of and are both Riemannian manifolds which are also duals In information geometry, this is known as the dually-flat Riemannian structure . of each other. An attractive property of this structure is that the FIM in one space is the inverse of the FIM in the other space. This enables us to compute natural gradient in one space using the gradient in the other, as shown in (11). This result is also discussed in an earlier work by Hensman et al. in the context of conjugate models, although they do not explicitly mention the connection to duality.
The natural gradient makes a better choice for conjugate models because assumes a simple form which does not require computation of the FIM. The unfortunately does not have this property. For example, for (10) requires computation of the FIM because it is equal to \mbox{\mbox{}}(\mbox{\mbox{}})(\mbox{\mbox{}}_{0}-\mbox{\mbox{}}). This can be shown by using (11), (9) and (10).
The recent work by propose to use the gradients with respect to to perform natural gradient with respect . They arrive at this conclusion by using the equivalence of mirror descent and natural-gradient descent. Our discussion above complements their work by using the duality of the two spaces.
IV Natural Gradients for Nonconjugate Models
In this section, we show that in some cases the natural gradient of the nonconjugate term can be easily computed by using . We also show that the resulting update takes a simple form.
We start with the expression for . Using (11) and (10), it is straightforward to write this expression:
A stochastic natural-gradient descent update can be obtained by using the gradient of a randomly sampled data example \mbox{{\cal D}}_{i} and multiplying it by , as shown below:
where the gradient is multiplied by to obtain an unbiased stochastic gradient. This update is equivalent to the update obtained in where it is referred to as Conjugate-computation variational inference (CVI). In , this is derived using a mirror-descent formulation, while we use the duality of the exponential family (Theorem 1).
If we approximate the expectations using a single Monte Carlo sample \mbox{\mbox{}}_{t}\sim q_{\lambda}(\mbox{\mbox{}}), we can write the update in (14) as
These updates take a form similar to Newton’s method. The covariance matrix \mbox{\mbox{}}_{t} plays a similar role to the Hessian in Newton’s method and scales the gradient in the update of \mbox{\mbox{}}_{t}. The matrix itself contains a moving average of the past Hessians. It is, however, not common to compute Hessians for deep models, but, as we discuss in Section VI, we can use another approximation to simplify this computation. With such an approximation, these updates can be implemented efficiently within existing deep learning code-bases as discussed in . ∎
Similarly to the above example, it might be possible to employ automatic-gradient methods to compute natural gradients in many models. A recent work explores this possibility. Another stochastic approximation method discussed in is also useful. For simple models, such as generalized linear models, where we can directly derive the distribution of the local variables, we can locally compute the gradients. This is discussed in for generalized linear models, Gaussian processes, and linear dynamical systems with nonlinear likelihoods.
V Local Approximations with Natural Gradients
We now show that natural gradients not only result in simple updates, but they also give rise to local exponential-family approximations of the nonconjugate terms. An attractive feature of these approximation is that the natural gradient of a nonconjugate likelihood is also the natural parameter of its local approximation.
We can then write the update (14) as an approximate Bayesian filter as shown below,
VI Results on Bayesian Neural Networks
In this section, we compare an approximate natural-gradient VI method with a gradient-based VI method. The natural-gradient method employs two approximations to the update (16)-(17). The first approximation is to use a diagonal covariance matrix which enables a fast computation when dimensionality of is large. The second approximation is to use a generalized Gauss-Newton approximation for the Hessian. This avoids the need to compute second-order derivatives making the implementation easier. The resulting method is called Variational Online Gauss-Newton (VOGN) . The updates of this method, as discussed in , is very similar to the Adam optimizer and can be implemented with a few lines of code change. This makes it easy to apply VOGN to large deep-learning problems.
Figure 1(b) compares VOGN with a gradient-based approach called Bayes by Backprop . The latter optimizes using the Adam optimizer. The results are obtained using a neural network with single-hidden layer of 64 hidden units and ReLU activations. A prior precision of , a minibatch size of 128 and 16 Monte-Carlo samples are used for all runs. The two figures show results on the following two datasets: ‘Australian’ ( and ) and ‘Breast Cancer’ ( and ) datasets. We show loss vs epochs, where a lower value indicates a better performance. We clearly see that the natural-gradient method is much faster than the gradient-based method. See for more experimental results.
VII Conclusions
In this paper, we discuss methods for natural-gradient descent in variational inference. Unlike gradient-based approaches, natural-gradient methods exploit the information geometry of the solution and can converge quickly. We review a few recent works and provide new insights using the duality associated with exponential-family approximations. We discuss an attractive property of the natural-gradient to obtain local conjugate approximations for individual model components. Finally, we showed some illustrative examples where these methods have been applied to perform Bayesian deep learning.
Acknowledgment
We would like to thank the following people at RIKEN, AIP for discussions and feedback: Aaron Mishkin, Frederik Kunstner, Voot Tangkaratt, and Wu Lin. We would also like to thank James Hensman and Shun-ichi Amari for discussions.