Stein Variational Gradient Descent: A General Purpose Bayesian Inference Algorithm

Qiang Liu, Dilin Wang

Introduction

Bayesian inference provides a powerful tool for modeling complex data and reasoning under uncertainty, but casts a long standing challenge on computing intractable posterior distributions. Markov chain Monte Carlo (MCMC) has been widely used to draw approximate posterior samples, but is often slow and has difficulty accessing the convergence. Variational inference instead frames the Bayesian inference problem into a deterministic optimization that approximates the target distribution with a simpler distribution by minimizing their KL divergence. This makes variational methods efficiently solvable by using off-the-shelf optimization techniques, and easily applicable to large datasets (i.e., "big data") using the stochastic gradient descent trick [e.g., 1]. In contrast, it is much more challenging to scale up MCMC to big data settings [see e.g., 2, 3].

Meanwhile, both the accuracy and computational cost of variational inference critically depend on the set of distributions in which the approximation is defined. Simple approximation sets, such as these used in the traditional mean field methods, are too restrictive to resemble the true posterior distributions, while more advanced choices cast more difficulties on the subsequent optimization tasks. For this reason, efficient variational methods often need to be derived on a model-by-model basis, causing is a major barrier for developing general purpose, user-friendly variational tools applicable for different kinds of models, and accessible to non-ML experts in application domains.

This case is in contrast with the maximum a posteriori (MAP) optimization tasks for finding the posterior mode (sometimes known as the poor man’s Bayesian estimator, in contrast with the full Bayesian inference for approximating the full posterior distribution), for which variants of (stochastic) gradient descent serve as a simple, generic, yet extremely powerful toolbox. There has been a recent growth of interest in creating user-friendly variational inference tools [e.g., 4, 5, 6, 7], but more efforts are still needed to develop more efficient general purpose algorithms.

In this work, we propose a new general purpose variational inference algorithm which can be treated as a natural counterpart of gradient descent for full Bayesian inference (see Algorithm 1). Our algorithm uses a set of particles for approximation, on which a form of (functional) gradient descent is performed to minimize the KL divergence and drive the particles to fit the true posterior distribution. Our algorithm has a simple form, and can be applied whenever gradient descent can be applied. In fact, it reduces to gradient descent for MAP when using only a single particle, while automatically turns into a full Bayesian approach with more particles.

Underlying our algorithm is a new theoretical result that connects the derivative of KL divergence w.r.t. smooth variable transforms and a recently introduced kernelized Stein discrepancy , which allows us to derive a closed form solution for the optimal smooth perturbation direction that gives the steepest descent on the KL divergence within the unit ball of a reproducing kernel Hilbert space (RKHS). This new result is of independent interest, and can find wide application in machine learning and statistics beyond variational inference.

This paper is organized as follows. Section 2 introduces backgrounds on kernelized Stein discrepancy (KSD). Our main results are presented in Section 3 in which we clarify the connection between KSD and KL divergence, and leverage it to develop our novel variational inference method. Section 4 discusses related works, and Section 5 presents numerical results. The paper is concluded in Section 6.

Background

Stein’s Identity and Kernelized Stein Discrepancy

Here the choice of this function set F{\mathcal{F}} is critical, and decides the discriminative power and computational tractability of Stein discrepancy. Traditionally, F{\mathcal{F}} is taken to be sets of functions with bounded Lipschitz norms, which unfortunately casts a challenging functional optimization problem that is computationally intractable or requires special considerations (see Gorham and Mackey and reference therein).

Kernelized Stein discrepancy bypasses this difficulty by maximizing ϕ{\boldsymbol{\phi}} in the unit ball of a reproducing kernel Hilbert space (RKHS) for which the optimization has a closed form solution. Following Liu et al. , KSD is defined as

where we assume the kernel k(x,x′)k(x,x^{\prime}) of RKHS H\mathcal{H} is in the Stein class of pp as a function of xx for any fixed x′∈Xx^{\prime}\in\mathcal{X}. The optimal solution of (2) has been shown to be ϕ(x)=ϕq,p∗(x)/∣∣ϕq,p∗∣∣Hd{\boldsymbol{\phi}}(x)={\boldsymbol{\phi}}^{*}_{q,p}(x)/||{\boldsymbol{\phi}}^{*}_{q,p}||_{\mathcal{H}^{d}}, where

Both Stein operator and KSD depend on pp only through the score function ∇xlog⁡p(x)\nabla_{x}\log p(x), which can be calculated without knowing the normalization constant of pp, because we have ∇xlog⁡p(x)=∇xlog⁡pˉ(x)\nabla_{x}\log p(x)=\nabla_{x}\log\bar{p}(x) when p(x)=pˉ(x)/Zp(x)=\bar{p}(x)/Z. This property makes Stein’s identity a powerful tool for handling unnormalized distributions that appear widely in machine learning and statistics.

Variational Inference Using Smooth Transforms

Variational inference approximates the target distribution p(x)p(x) using a simpler distribution q∗(x)q^{*}(x) found in a predefined set Q={q(x)}\mathcal{Q}=\{q(x)\} of distributions by minimizing the KL divergence, that is,

where we do not need to calculate the constant log⁡Z\log Z for solving the optimization. The choice of set Q\mathcal{Q} is critical and defines different types of variational inference methods. The best set Q\mathcal{Q} should strike a balance between i) accuracy, broad enough to closely approximate a large class of target distributions, ii) tractability, consisting of simple distributions that are easy for inference, and iii) solvability so that the subsequent KL minimization problem can be efficiently solved.

In this work, we focus on the sets Q{\mathcal{Q}} consisting of distributions obtained by smooth transforms from a tractable reference distribution, that is, we take Q{\mathcal{Q}} to be the set of distributions of random variables of form z=T(x)z={\boldsymbol{T}}(x) where T ⁣:X→X{\boldsymbol{T}}\colon\mathcal{X}\to\mathcal{X} is a smooth one-to-one transform, and xx is drawn from a tractable reference distribution q0(x)q_{0}(x). By the change of variables formula, the density of zz is

where T−1{\boldsymbol{T}}^{-1} denotes the inverse map of T{\boldsymbol{T}} and ∇zT−1\nabla_{z}{\boldsymbol{T}}^{-1} the Jacobian matrix of T−1{\boldsymbol{T}}^{-1}. Such distributions are computationally tractable, in the sense that the expectation under q[T]q_{[{\boldsymbol{T}}]} can be easily evaluated by averaging {zi}\{z_{i}\} when zi=T(xi)z_{i}={\boldsymbol{T}}(x_{i}) and xi∼q0.x_{i}\sim q_{0}. Such Q{\mathcal{Q}} can also in principle closely approximate almost arbitrary distributions: it can be shown that there always exists a measurable transform T{\boldsymbol{T}} between any two distributions without atoms (i.e. no single point carries a positive mass); in addition, for Lipschitz continuous densities pp and qq, there always exist transforms between them that are least as smooth as both pp and qq. We refer the readers to Villani for in-depth discussion on this topic.

In practice, however, we need to restrict the set of transforms T{\boldsymbol{T}} properly to make the corresponding variational optimization in (4) practically solvable. One approach is to consider T{\boldsymbol{T}} with certain parametric form and optimize the corresponding parameters [e.g., 13, 14]. However, this introduces a difficult problem on selecting the proper parametric family to balance the accuracy, tractability and solvability, especially considering that T{\boldsymbol{T}} has to be an one-to-one map and has to have an efficiently computable Jacobian matrix.

Instead, we propose a new algorithm that iteratively constructs incremental transforms that effectively perform steepest descent on T{\boldsymbol{T}} in RKHS. Our algorithm does not require to explicitly specify parametric forms, nor to calculate the Jacobian matrix, and has a particularly simple form that mimics the typical gradient descent algorithm, making it easily implementable even for non-experts in variational inference.

To explain how we minimize the KL divergence in (4), we consider an incremental transform formed by a small perturbation of the identity map: T(x)=x+ϵϕ(x){\boldsymbol{T}}(x)=x+\epsilon{\boldsymbol{\phi}}(x), where ϕ(x){\boldsymbol{\phi}}(x) is a smooth function that characterizes the perturbation direction and the scalar ϵ\epsilon represents the perturbation magnitude. When ∣ϵ∣|\epsilon| is sufficiently small, the Jacobian of T{\boldsymbol{T}} is full rank (close to the identity matrix), and hence T{\boldsymbol{T}} is guaranteed to be an one-to-one map by the inverse function theorem.

The following result, which forms the foundation of our method, draws an insightful connection between Stein operator and the derivative of KL divergence w.r.t. the perturbation magnitude ϵ\epsilon.

Let T(x)=x+ϵϕ(x){\boldsymbol{T}}(x)=x+\epsilon{\boldsymbol{\phi}}(x) and q[T](z)q_{[{\boldsymbol{T}}]}(z) the density of z=T(x)z={\boldsymbol{T}}(x) when x∼q(x)x\sim q(x), we have

where Apϕ(x)=∇xlog⁡p(x)ϕ(x)⊤+∇xϕ(x){\mathcal{A}}_{p}{\boldsymbol{\phi}}(x)=\nabla_{x}\log p(x){\boldsymbol{\phi}}(x)^{\top}+\nabla_{x}{\boldsymbol{\phi}}(x) is the Stein operator.

Relating this to the definition of KSD in (2), we can identify the ϕq,p∗{\boldsymbol{\phi}}^{*}_{q,p} in (3) as the optimal perturbation direction that gives the steepest descent on the KL divergence in zero-centered balls of Hd\mathcal{H}^{d}.

Let T(x)=x+f(x){\boldsymbol{T}}(x)=x+{\boldsymbol{f}}(x), where f∈Hd{\boldsymbol{f}}\in\mathcal{H}^{d}, and q[T]q_{[{\boldsymbol{T}}]} the density of z=T(x)z={\boldsymbol{T}}(x) when x∼qx\sim q,

This suggests that T∗(x)=x+ϵ⋅ϕq,p∗(x){\boldsymbol{T}}^{*}(x)=x+\epsilon\cdot{\boldsymbol{\phi}}^{*}_{q,p}(x) is equivalent to a step of functional gradient descent in RKHS. However, what is critical in the iterative procedure (7) is that we also iteratively apply the variable transform so that every time we would only need to evaluate the functional gradient descent at zero perturbation f=0\boldsymbol{f}=0 on the identity map T(x)=x{\boldsymbol{T}}(x)=x. This brings a critical advantage since the gradient at f≠0\boldsymbol{f}\neq 0 is more complex and would require to calculate the inverse Jacobian matrix [∇xT(x)]−1[\nabla_{x}{\boldsymbol{T}}(x)]^{-1} that casts computational or implementation hurdles.

2 Stein Variational Gradient Descent

Algorithm 1 mimics a gradient dynamics at the particle level, where the two terms in ϕ^∗(x){\boldsymbol{\hat{\phi}}}{}^{*}(x) in (8) play different roles: the first term drives the particles towards the high probability areas of p(x)p(x) by following a smoothed gradient direction, which is the weighted sum of the gradients of all the points weighted by the kernel function. The second term acts as a repulsive force that prevents all the points to collapse together into local modes of p(x)p(x); to see this, consider the RBF kernel k(x,x′)=exp⁡(−1h∣∣x−x′∣∣2)k(x,x^{\prime})=\exp(-\frac{1}{h}||x-x^{\prime}||^{2}), the second term reduces to ∑j2h(x−xj)k(xj,x)\sum_{j}\frac{2}{h}(x-x_{j})k(x_{j},x), which drives xx away from its neighboring points xjx_{j} that have large k(xj,x)k(x_{j},x). If we let bandwidth h→0h\to 0, the repulsive term vanishes, and update (8) reduces to a set of independent chains of typical gradient ascent for maximizing log⁡p(x)\log p(x) (i.e., MAP) and all the particles would collapse into the local modes.

Another interesting case is when we use only a single particle (n=1n=1), in which case Algorithm 1 reduces to a single chain of typical gradient ascent for MAP for any kernel that satisfies ∇xk(x,x)=0\nabla_{x}k(x,x)=0 (for which RBF holds). This suggests that our algorithm can generalize well for supervised learning tasks even with a very small number nn of particles, since gradient ascent for MAP (n=1n=1) has been shown to be very successful in practice. This property distinguishes our particle method with the typical Monte Carlo methods that requires to average over many points. The key difference here is that we use a deterministic repulsive force, other than Monte Carlo randomness, to get diverse points for distributional approximation.

The major computation bottleneck in (8) lies on calculating the gradient ∇xlog⁡p(x)\nabla_{x}\log p(x) for all the points {xi}i=1n\{x_{i}\}_{i=1}^{n}; this is especially the case in big data settings when p(x)∝p0(x)∏k=1Np(Dk∣x)p(x)\propto p_{0}(x)\prod{}^{N}_{k=1}p(D_{k}|x) with a very large NN. We can conveniently address this problem by approximating ∇xlog⁡p(x)\nabla_{x}\log p(x) with subsampled mini-batches Ω⊂{1,…,N}\Omega\subset\{1,\ldots,N\} of the data

Additional speedup can be obtained by parallelizing the gradient evaluation of the nn particles.

The update (8) also requires to compute the kernel matrix {k(xi,xj)}\{k(x_{i},x_{j})\} which costs O(n2)\mathcal{O}{\left(n^{2}\right)}; in practice, this cost can be relatively small compared with the cost of gradient evaluation, since it can be sufficient to use a relatively small nn (e.g., several hundreds) in practice. If there is a need for very large nn, one can approximate the summation ∑i=1n\sum_{i=1}^{n} in (8) by subsampling the particles, or using a random feature expansion of the kernel k(x,x′)k(x,x^{\prime}) .

Related Works

Our algorithm maintains and updates a set of particles, and is of similar style with the Gaussian mixture variation inference methods whose mean parameters can be treated as a set of particles. . Optimizing such mixture KL objectives often requires certain approximation, and this was done most recently in Gershman et al. by approximating the entropy using Jensen’s inequality and the expectation term using Taylor approximation. There is also a large set of particle-based Monte Carlo methods, including variants of sequential Monte Carlo [e.g., 27, 28], as well as a recent particle mirror descent for optimizing the variational objective function ; compared with these methods, our method does not have the weight degeneration problem, and is much more “particle-efficient” in that we reduce to MAP with only one single particle.

Experiments

We test our algorithm on both toy and real world examples, on which we find our method tends to outperform a variety of baseline methods. Our code is available at https://github.com/DartML/Stein-Variational-Gradient-Descent.

We set our target distribution to be p(x)=1/3N(x; −2,1)+2/3N(x; 2,1)p(x)=1/3\mathcal{N}(x;~{}-2,1)+2/3\mathcal{N}(x;~{}2,1), and initialize the particles using q0(x)=N(x;−10,1)q_{0}(x)=\mathcal{N}(x;-10,1). This creates a challenging situation since the probability mass of p(x)p(x) and q0(x)q_{0}(x) are far away each other (with almost zero overlap). Figure 1 shows how the distribution of the particles (n=1)(n=1) of our method evolve at different iterations. We see that despite the small overlap between q0(x)q_{0}(x) and p(x)p(x), our method can push the particles towards the target distribution, and even recover the mode that is further away from the initial point. We found that other particle based algorithms, such as Dai et al. , tend to experience weight degeneracy on this toy example due to the ill choice of q0(x)q_{0}(x).

Bayesian Logistic Regression

We consider Bayesian logistic regression for binary classification using the same setting as Gershman et al. , which assigns the regression weights ww with a Gaussian prior p0(w∣α)=N(w,α−1)p_{0}(w|\alpha)=\mathcal{N}(w,\alpha^{-1}) and p0(α)=Gamma(α,1,0.01)p_{0}(\alpha)=Gamma(\alpha,1,0.01). The inference is applied on posterior p(x∣D)p(x|D) with x=[w,log⁡α]x=[w,\log\alpha]. We compared our algorithm with the no-U-turn sampler (NUTS)code: http://www.cs.princeton.edu/ mdhoffma/ and non-parametric variational inference (NPV)code: http://gershmanlab.webfactional.com/pubs/npv.v1.zip on the 8 datasets (N>500N>500) used in Gershman et al. , and find they tend to give very similar results on these (relatively simple) datasets; see Appendix for more details.

Bayesian Neural Network

We find our algorithm consistently improves over PBP both in terms of the accuracy and speed (except on Yacht); this is encouraging since PBP were specifically designed for Bayesian neural network. We also find that our results are comparable with the more recent results reported on the same datasets [e.g., 32, 33, 34] which leverage some advanced techniques that we can also benefit from.

Conclusion

We propose a simple general purpose variational inference algorithm for fast and scalable Bayesian inference. Future directions include more theoretical understanding on our method, more practical applications in deep learning models, and other potential applications of our basic Theorem in Section 3.1.

References

Appendix A Proof of Theorem 3.1

Let qq and pp be two smooth densities, and T=Tϵ(x){\boldsymbol{T}}={\boldsymbol{T}}_{\epsilon}(x) an one-to-one transform on X\mathcal{X} indexed by parameter ϵ\epsilon, and T{\boldsymbol{T}} is differentiable w.r.t. both xx and ϵ\epsilon. Define q[T]q_{[{\boldsymbol{T}}]} to be the density of z=Tϵ(x)z={\boldsymbol{T}}_{\epsilon}(x) when x∼qx\sim q, and sp=∇xlog⁡p(x){\boldsymbol{s}}_{p}=\nabla_{x}\log p(x), we have

Denote by p[T−1](z)p_{[{\boldsymbol{T}}^{-1}]}(z) the density of z=T−1(x)z={\boldsymbol{T}}^{-1}(x) when x∼p(x)x\sim p(x), then

We just need to calculate log⁡p[T−1](x)\log p_{[{\boldsymbol{T}}^{-1}]}(x); define sp(x)=∇xlog⁡p(x){\boldsymbol{s}}_{p}(x)=\nabla_{x}\log p(x), we get

When T(x)=x+ϵϕ(x){\boldsymbol{T}}(x)=x+\epsilon{\boldsymbol{\phi}}(x) and ϵ=0\epsilon=0, we have

where II is the identity matrix. Using Lemma A.1 gives the result. ∎

Appendix B Proof of Theorem 3.3

Let Hd=H×⋯×H\mathcal{H}^{d}=\mathcal{H}\times\cdots\times\mathcal{H} be a vector-valued RKHS, and F[f]F[f] be a functional on ff. The gradient ∇fF[f]\nabla_{f}F[f] of F[⋅]F[\cdot] is a function in Hd\mathcal{H}^{d} that satisfies

For the terms in the above equation, we have

Taking f=0f=0 then gives the desirable result. ∎

Appendix C Connection with de Bruijn’s identity and Fisher Divergence

If we take ϕq,p(x)=∇xlog⁡p(x)−∇xlog⁡q(x){\boldsymbol{\phi}}_{q,p}(x)=\nabla_{x}\log p(x)-\nabla_{x}\log q(x) in (5), we can show that (5) reduces to

where F(q, p){\mathcal{F}}(q,~{}p) is the Fisher divergence between pp and qq, defined as

Note that this can be treated as a deterministic version of de Bruijn’s identity , which draws similar connection between KL and Fisher divergence, but uses randomized linear transform T(x)=x+ϵ⋅ξ{\boldsymbol{T}}(x)=x+\sqrt{\epsilon}\cdot\xi, where ξ\xi is a standard Gaussian noise.

Appendix D Additional Experiments

We collect additional experimental results that can not fitted into the main paper.

We consider the Bayesian logistic regression model for binary classification, on which the regression weights ww is assigned with a Gaussian prior p0(w)=N(w,α−1)p_{0}(w)=\mathcal{N}(w,\alpha^{-1}) and p0(α)=Γ(α,a,b)p_{0}(\alpha)=\Gamma(\alpha,a,b), and apply inference on posterior p(x∣D)p(x\mid D), where x=[w,log⁡α]x=[w,\log\alpha]. The hyper-parameter is taken to be a=1a=1 and b=0.01b=0.01. This setting is the same as that in Gershman et al. . We compared our algorithm with the no-U-turn sampler (NUTS)code: http://www.cs.princeton.edu/ mdhoffma/ and non-parametric variational inference (NPV)code: http://gershmanlab.webfactional.com/pubs/npv.v1.zip on the 8 datasets (N>500N>500) as used in Gershman et al. , in which we use 100100 particles, NPV uses 100 mixture components, and NUTS uses 1000 draws with 10001000 burnin period. We find that all these three algorithms almost always performs the same across the 8 datasets (See Figure in Appendix), and this is consistent with Figure 2 of Gershman et al. .

We further experimented on a toy dataset with only two features and visualize the prediction probability of the three algorithms in Figure 5. We again find that all the three algorithms tend to perform similarly. Note, however, that NPV is relatively inconvenient to use since it requires the Hessian matrix, and NUTS tends to be very small when applied on massive datasets.