Pathwise Derivatives Beyond the Reparameterization Trick

Martin Jankowiak, Fritz Obermeyer

Introduction

Maximizing objective functions via gradient methods is ubiquitous in machine learning. When the objective function L\mathcal{L} is defined as an expectation of a (differentiable) test function fθ(z)f_{\bm{\theta}}({\bm{z}}) w.r.t. a probability distribution qθ(z)q_{{\bm{\theta}}}({\bm{z}}),

computing exact gradients w.r.t. the parameters θ{\bm{\theta}} is often unfeasible so that optimization methods must instead make due with stochastic gradient estimates. If the gradient estimator is unbiased, then stochastic gradient descent with an appropriately chosen sequence of step sizes can be shown to have nice convergence properties (Robbins & Monro, 1951). If, however, the gradient estimator exhibits large variance, stochastic optimization algorithms may be impractically slow. Thus it is of general interest to develop gradient estimators with reduced variance.

We revisit the class of gradient estimators popularized in (Kingma & Welling, 2013; Rezende et al., 2014; Titsias & Lázaro-Gredilla, 2014), which go under the name of the pathwise derivative or the reparameterization trick. While this class of gradient estimators is not applicable to all choices of probability distribution qθ(z)q_{{\bm{\theta}}}({\bm{z}}), empirically it has been shown to yield suitably low variance in many cases of practical interest and thus has seen wide use. We show that the pathwise derivative in the literature is in fact a particular instance of a continuous family of gradient estimators. Drawing a connection to tangent fields in the field of optimal transport,See (Villani, 2003; Ambrosio et al., 2008) for a review. we show that one can define a unique pathwise gradient that is optimal in the sense of optimal transport. For the purposes of this paper, we will refer to these optimal gradients as OMT (optimal mass transport) gradients.

The resulting geometric picture is particularly intriguing in the case of multivariate distributions, where each choice of gradient estimator specifies a velocity field on the sample space. To make this picture more concrete, in Figure 1 we show the velocity fields that correspond to two different gradient estimators for the off-diagonal element of the Cholesky factor parameterizing a bivariate Normal distribution. We note that the velocity field that corresponds to the reparameterization trick has a large rotational component that makes it suboptimal in the sense of optimal transport. In Sec. 7 we show that this suboptimality can result in reduced performance when fitting a Gaussian Process to data.

The rest of this paper is organized as follows. In Sec. 2 we provide a brief overview of stochastic gradient variational inference. In Sec. 3 we show how to compute pathwise gradients for univariate distributions. In Sec. 4 we expand our discussion of pathwise gradients to the case of multivariate distributions, introduce the connection to the transport equation, and provide an analytic formula for the OMT gradient in the case of the multivariate Normal. In Sec. 5 we discuss how we can compute high precision approximate pathwise gradients for the Gamma, Beta, and Dirichlet distributions. In Sec. 6 we place our work in the context of related research. In Sec. 7 we demonstrate the performance of our gradient estimators with a variety of synthetic experiments and experiments on real world datasets. Finally, in Sec. 8 we conclude with a discussion of directions for future work.

Stochastic Gradient Variational Inference

One area where stochastic gradient estimators play a particularly central role is stochastic variational inference (Hoffman et al., 2013). This is especially the case for black-box methods (Wingate & Weber, 2013; Ranganath et al., 2014), where conjugacy and other simplifying structural assumptions are unavailable, with the consequence that Monte Carlo estimators become necessary. For concreteness, we will refer to this class of methods as Stochastic Gradient Variational Inference (SGVI). In this section we give a brief overview of this line of research, as it serves as the motivating use case for our work. Furthermore, in Sec. 7 SGVI will serve as the main testbed for our proposed methods.

Let p(x,z)p({\bm{x}},{\bm{z}}) define a joint probability distribution over observed data x{\bm{x}} and latent random variables z{\bm{z}}. One of the main tasks in Bayesian inference is to compute the posterior distribution p(z∣x)=p(x,z)p(x)p({\bm{z}}|{\bm{x}})=\frac{p({\bm{x}},{\bm{z}})}{p({\bm{x}})}. For many models of interest, this is an intractably hard problem and so approximate methods become necessary. Variational inference recasts Bayesian inference as an optimization problem. Specifically we define a variational family of distributions qθ(z)q_{\bm{\theta}}({\bm{z}}) parameterized by θ{\bm{\theta}} and seek to find a value of θ{\bm{\theta}} that minimizes the KL divergence between qθ(z)q_{\bm{\theta}}({\bm{z}}) and the (unknown) posterior p(z∣x)p({\bm{z}}|{\bm{x}}). This is equivalent to maximizing the ELBO (Jordan et al., 1999), defined as

For general choices of p(x,z)p({\bm{x}},{\bm{z}}) and qθ(z)q_{\bm{\theta}}({\bm{z}}), this expectation—much less its gradients—cannot be computed analytically. In these circumstances a natural approach is to build a Monte Carlo estimator of the ELBO and its gradient w.r.t. θ{\bm{\theta}}. The properties of the chosen gradient estimator—especially its bias and variance—play a critical rule in determining the viability of the resulting stochastic optimization. Next, we review two commonly used gradient estimators; we leave a brief discussion of more elaborate variants to Sec. 6.

The score function estimator, also referred to as the log-derivative trick or reinforce (Glynn, 1990; Williams, 1992), provides a simple and broadly applicable recipe for estimating ELBO gradients (Paisley et al., 2012). The score function estimator expresses the gradient as an expectation with respect to qθ(z)q_{\bm{\theta}}({\bm{z}}), with the simplest variant given by

where log⁡r=log⁡p(x,z)−log⁡qθ(z)\log r=\log p({\bm{x}},{\bm{z}})-\log q_{{\bm{\theta}}}({\bm{z}}). Monte Carlo estimates of Eqn. 3 can be formed by drawing samples from qθ(z)q_{\bm{\theta}}({\bm{z}}) and computing the term in the square brackets. Although the score function estimator is very general (e.g. it applies to discrete random variables) it typically suffers from high variance, although this can be mitigated with the use of variance reduction techniques such as Rao-Blackwellization (Casella & Robert, 1996) and control variates (Ross, 2006).

2 Pathwise Gradient Estimator

The pathwise gradient estimator, a.k.a. the reparameterization trick (RT), is not as broadly applicable as the score function estimator, but it generally exhibits lower variance (Price, 1958; Salimans et al., 2013; Kingma & Welling, 2013; Glasserman, 2013; Rezende et al., 2014; Titsias & Lázaro-Gredilla, 2014). It is applicable to continuous random variables whose probability density qθ(z)q_{{\bm{\theta}}}({\bm{z}}) can be reparameterized such that we can rewrite expectations

where q0(z)q_{0}({\bm{z}}) is a fixed distribution with no dependence on θ{\bm{\theta}} and T(ϵ;θ)\mathcal{T}({\bm{\epsilon}};{\bm{\theta}}) is a differentiable θ{\bm{\theta}}-dependent transformation. Since the expectation w.r.t. q0(ϵ)q_{0}({\bm{\epsilon}}) has no θ{\bm{\theta}} dependence, gradients w.r.t. θ{\bm{\theta}} can be computed by pushing ∇θ\nabla_{\bm{\theta}} through the expectation. This reparameterization can be done for a number of distributions, including for example the Normal distribution. Unfortunately the reparameterization trick is non-trivial to apply to a number of commonly used distributions, e.g. the Gamma and Beta distributions, since the required shape transformations T(ϵ;θ)\mathcal{T}({\bm{\epsilon}};{\bm{\theta}}) inevitably involve special functions.

Univariate Pathwise Gradients

Consider an objective function given as the expectation of a test function fθ(z)f_{\bm{\theta}}(z) with respect to a distribution qθ(z)q_{\bm{\theta}}(z), where zz is a continuous one-dimensional random variable:

Here qθ(z)q_{\bm{\theta}}(z) and fθ(z)f_{\bm{\theta}}(z) are parameterized by θ{\bm{\theta}}, and we would like to compute (stochastic) gradients of L\mathcal{L} w.r.t. θ\theta, where θ\theta is a scalar component of θ{\bm{\theta}}:

Crucially we would like to avoid the log-derivative trick, which yields a gradient estimator that tends to have high variance. Doing so will be easy if we can rewrite the expectation in terms of a fixed distribution that does not depend on θ\theta. A natural choice is to use the standard uniform distribution U\mathcal{U},

where the transformation Fθ−1:u→zF_{\bm{\theta}}^{-1}:u\rightarrow z is the inverse CDF of qθ(z)q_{\bm{\theta}}(z). As desired, all dependence on θ\theta is now inside the expectation. Unfortunately, for many continuous univariate distributions of interest (e.g. the Gamma and Beta distributions) the transformation Fθ−1F_{\bm{\theta}}^{-1} (as well as its derivative w.r.t. θ{\bm{\theta}}) does not admit a simple analytic expression.

Fortunately, by making use of implicit differentiation we can compute the gradient in Eqn. 6 without explicitly introducing Fθ−1F_{\bm{\theta}}^{-1}. To complete the derivation define uu by

and differentiate both sides of Eqn. 8 w.r.t. θ\theta and make use of the fact that u∼Uu\sim\mathcal{U} does not depend on θ\theta to obtain

This then yields our master formula for the univariate case

where the corresponding gradient estimator is given by

While this derivation is elementary, it helps to clarify things: the key ingredient needed to compute pathwise gradients in Eqn. 6 is the ability to compute (or approximate) the derivative of the CDF, i.e. ∂∂θFθ(z)\frac{\partial}{\partial\theta}F_{{\bm{\theta}}}(z). In the supplementary materials we verify that Eqn. 11 results in correct gradients.

It is worth emphasizing how this approach differs from a closely related alternative. Suppose we construct a (differentiable) approximation of the inverse CDF, F^θ−1(u)≈Fθ−1(u)\hat{F}_{{\bm{\theta}}}^{-1}(u)\approx F_{{\bm{\theta}}}^{-1}(u). For example, we might train a neural network nn(u,θ)≈Fθ−1(u){\rm nn}(u,{\bm{\theta}})\approx F_{{\bm{\theta}}}^{-1}(u). We can then push samples u∼Uu\sim\mathcal{U} through nn(u,θ){\rm nn}(u,{\bm{\theta}}) and obtain approximate samples from qθ(z)q_{{\bm{\theta}}}(z) as well as approximate derivatives dzdθ\frac{dz}{d\theta} via the chain rule; in this case, there will be a mismatch between the probability qθ(z)q_{\bm{\theta}}(z) assigned to samples zz and the actual distribution over zz. By contrast, if we use the construction of Eqn. 10, our samples zz will still be exactOr rather their exactness will be determined by the quality of our sampler for qθ(z)q_{\bm{\theta}}(z), which is fully decoupled from how we compute derivatives dzdθ\frac{dz}{d\theta}. and the fidelity of our approximation of (the derivatives of) Fθ(z)F_{{\bm{\theta}}}(z) will only affect the accuracy of our approximation for dzdθ\frac{dz}{d\theta}.

Multivariate Pathwise Gradients

In the previous section we focused on continuous univariate distributions. Pathwise gradients can also be constructed for continuous multivariate distributions, although the analysis is in general expected to be much more complicated than in the univariate case—directly analogous to the difference between ordinary and partial differential equations. Before constructing estimators for particular distributions, we introduce the connection to the transport equation.

Here the velocity field vθ\bm{v}^{\theta} is a vector field defined on the sample space that displaces samples (i.e. particles) z\bm{z} as we vary θ\theta infinitesimally. Note that there is a velocity field vθ\bm{v}^{\theta} for each component θ\theta of θ{\bm{\theta}}. This equation is readily interpreted in the language of fluid dynamics. In order for the the total probability to be conserved, the term ∂∂θqθ(z)\frac{\partial}{\partial\theta}q_{\bm{\theta}}({\bm{z}})—which is the rate of change of the number of particles in the infinitesimal volume element at z\bm{z}—has to be counterbalanced by the in/out-flow of particles—as given by the divergence term.

2 Gradient Estimator

Given a solution to Eqn. 12, we can form the gradient estimator

which generalizes Eqn. 11 to the multivariate case. That this is an unbiased gradient estimator follows directly from the divergence theorem (see the supplementary materials).

3 Tangent Fields

In general Eqn. 12 admits an infinite dimensional space of solutions. In the context of our derivation of Eqn. 10, we might loosely say that different solutions of Eqn. 12 correspond to different ways of specifying quantiles of qθ(z)q_{{\bm{\theta}}}(\bm{z}). To determine a uniqueWe refer the reader to Ch. 8 of (Ambrosio et al., 2008) for details. solution—the tangent field from the theory of optimal transport—we require that

In this case it can be shown that vOMT\bm{v}^{\rm OMT} minimizes the total kinetic energy, which is given byNote that the univariate solution, Eqn. 10, is automatically the OMT solution.

4 Gradient variance

5 The Multivariate Normal

Note that Eqn. 17 is just a particular instance of the solution to the transport equation that is implicitly provided by the reparameterization trick, namely

In the supplementary materials we verify that Eqn. 17 satisfies the transport equation Eqn. 12. However, it is evidently not optimal in the sense of optimal transport, since ∂viRT∂zj=δiaLbj−1\frac{\partial v_{i}^{\rm RT}}{\partial z_{j}}=\delta_{ia}L^{-1}_{bj} is not symmetric in ii and jj. In fact the tangent field takes the form

where SabS^{ab} is a symmetric matrix whose precise form we give in the supplementary materials. We note that computing gradients with Eqn. 19 is O(D3)\mathcal{O}(D^{3}), since it involves a singular value decomposition of the covariance matrix. In Sec. 7 we show that the resulting gradient estimator can lead to reduced variance.

Numerical Recipes

In this section we show how Eqn. 10 can be used to obtain pathwise gradients in practice. In many cases of interest we will need to derive approximations to ∂∂θF(z)\frac{\partial}{\partial\theta}F(z) that balance the need for high accuracy (thus yielding gradient estimates with negligible bias) with the need for computational efficiency. In particular we will derive accurate approximations to Eqn. 10 for the Gamma, Beta, and Dirichlet distributions. These approximations will involve three basic components:

The Lugannani-Rice saddlepoint expansion (Lugannani & Rice, 1980; Butler, 2007)

Rational polynomial approximations in regions of (z,θ)(z,\theta) that are analytically intractable

The CDF of the Gamma distribution involves the (lower) incomplete gamma function γ(⋅)\gamma(\cdot): Fα,β(z)=γ(α,βz)Γ(α)F_{\alpha,\beta}(z)=\frac{\gamma(\alpha,\beta z)}{\Gamma(\alpha)}. Unfortunately γ(⋅)\gamma(\cdot) does not admit simple analytic expressions for derivatives w.r.t. its first argument, and so we must resort to numerical approximations. Since z∼Gamma(α,β=1)⇔z/β∼Gamma(α,β)z\sim\rm{Gamma}(\alpha,\beta=1)\Leftrightarrow z/\beta\sim\rm{Gamma}(\alpha,\beta) it is sufficient to consider dzdα\frac{dz}{d\alpha} for the standard Gamma distribution with β=1\beta=1.

To give a flavor for the kinds of approximations we use, consider how we can approximate ∂∂αγ(α,z)\frac{\partial}{\partial\alpha}\gamma(\alpha,z) in the limit z≪1z\ll 1. We simply do a Taylor series in powers of zz:

In practice we use 6 terms in this expansion, which is accurate for z<0.8z<0.8. Details for the remaining approximations can be found in the supplementary materials.

2 Beta

The CDF of the Beta distribution, FBetaF_{\rm Beta}, is the (regularized) incomplete beta function; just like in the case of the Gamma distribution, its derivatives do not admit simple analytic expressions. We describe the numerical approximations we used in the supplementary materials.

3 Dirichlet

Let z∼Dir(α)\bm{z}\sim\rm{Dir}(\bm{\alpha}) be Dirichlet distributed with nn components. Noting that the ziz_{i} are constrained to lie within the unit (n−1)(n-1)-simplex, we proceed by representing z\bm{z} in terms of n−1n-1 mutually independent Beta variates (Wilks, 1962):

Note that Eqn. 20 implies that ddα∑izi=0\frac{d}{d\bm{\alpha}}\sum_{i}z_{i}=0, as it must because of the simplex constraint. Since we have already developed an approximation for ∂FBeta∂θ\frac{\partial F_{\rm{Beta}}}{\partial\theta}, Eqn. 20 provides a complete recipe for pathwise Dirichlet gradients. Note that although we have used a stick-breaking construction to derive Eqn. 20, this in no way dictates the sampling scheme we use when generating z∼Dir(α)\bm{z}\sim\rm{Dir}(\bm{\alpha}). In the supplementary materials we verify that Eqn. 20 satisfies the transport equation.

4 Implementation

It is worth emphasizing that pathwise gradient estimators of the form in Eqn. 13 have the advantage of being ‘plug-and-play.’ We simply plug an approximate or exact velocity field into our favorite automatic differentiation engineOur approximations for pathwise gradients for the Gamma, Beta, and Dirichlet distributions are available in the 0.4 release of PyTorch (Paszke et al., 2017). so that samples z\bm{z} and fθ(z)f_{\bm{\theta}}(\bm{z}) are differentiable w.r.t. θ{\bm{\theta}}. There is no need to construct a surrogate objective function to form the gradient estimator.

Related Work

A number of lines of research bears upon our work. There is a large body of work on constructing gradient estimators with reduced variance, much of which can be understood in terms of control variates (Ross, 2006): for example, (Mnih & Gregor, 2014) construct neural baselines for score-function gradients; (Schulman et al., 2015) discuss gradient estimators for stochastic computation graphs and their Rao-Blackwellization; and (Tucker et al., 2017; Grathwohl et al., 2017) construct adaptive control variates for discrete random variables. Another example of this line of work is reference (Miller et al., 2017), where the authors construct control variates that are applicable when qθ(z)q_{\bm{\theta}}({\bm{z}}) is a diagonal Normal distribution. While our OMT gradient for the multivariate Normal distribution, Eqn. 19, can also be understood in the language of control variates,See Sec. 8 and the supplementary materials for a brief discussion. (Miller et al., 2017) relies on Taylor expansions of the test function fθ(z)f_{\bm{\theta}}({\bm{z}}).In addition, note that in their approach variance reduction for gradients w.r.t. the scale parameter σ\bm{\sigma} necessitates a multi-sample estimator (at least for high-dimensional models where computing the diagonal of the Hessian is prohibitively expensive).

In (Graves, 2016), the author derives formula Eqn. 10 and uses it to construct gradient estimators for mixture distributions. Unfortunately, the resulting gradient estimator is expensive, relying on a recursive computation that scales with the dimension of the sample space.

Another line of work constructs partially reparameterized gradient estimators for cases where the reparameterization trick is difficult to apply. The generalized reparameterization gradient (G-Rep) (Ruiz et al., 2016) uses standardization via sufficient statistics to obtain a transformation T(ϵ;θ)\mathcal{T}({\bm{\epsilon}};{\bm{\theta}}) that minimizes the dependence of q(ϵ)q({\bm{\epsilon}}) on θ{\bm{\theta}}. This results in a partially reparameterized gradient estimator that also includes a score function-like term.That is a term in the gradient estimator that is proportional to the test function fθ(z)f_{\bm{\theta}}({\bm{z}}). In rsvi (Naesseth et al., 2017) the authors consider gradient estimators in the case that qθ(z)q_{\bm{\theta}}({\bm{z}}) can be sampled from efficiently via rejection sampling. This results in a gradient estimator with the same generic structure as G-Rep, although in the case of rsvi the score function-like term can often be dropped in practice at the cost of small bias (with the benefit of reduced variance). Besides the fact that this gradient estimator is not fully pathwise, one key difference with our approach is that for many distributions of interest (e.g. the Beta and Dirichlet distributions), rejection sampling introduces auxiliary random variables, which results in additional stochasticity and thus higher variance (cf. Figure 2). In contrast our pathwise gradients for the Beta and Dirichlet distributions are deterministic for a given z{\bm{z}} and θ{\bm{\theta}}. Finally, (Knowles, 2015) uses (somewhat imprecise) approximations to the inverse CDF to derive gradient estimators for Gamma random variables.

As the final version of this manuscript was being prepared, we became aware of (Figurnov et al., 2018), which has some overlap with this work. In particular, (Figurnov et al., 2018) derives Eqn. 10 and an interesting generalization to the multivariate case. This allows the authors to construct pathwise derivatives for the Gamma, Beta, and Dirichlet distributions. For the latter two distributions, however, the derivatives include additional stochasticity that our pathwise derivatives avoid. Also, the authors do not draw the connection to the transport equation and optimal transport or consider the multivariate Normal distribution in any detail.

Experiments

All experiments in this section use single-sample gradient estimators.

In this section we validate our pathwise gradients for the Beta, Dirichlet, and multivariate Normal distributions. Where appropriate we compare to the RT gradient, the score function gradient, or rsvi.

In Fig. 3 we compare the performance of our OMT gradient for Beta random variables to the rsvi gradient estimator. We use a test function f(z)=z3f(z)=z^{3} for which we can compute the gradient exactly. We see that the OMT gradient performs favorably over the entire range of parameter α\alpha that defines the distribution Beta(α,α){\rm Beta}(\alpha,\alpha) used to compute L\mathcal{L}. For smaller α\alpha, where L\mathcal{L} exhibits larger curvature, the variance of the estimator is noticeably reduced. Notice that one reason for the reduced variance of the OMT estimator as compared to the rsvi estimator is the presence of an auxiliary random variable in the latter case (cf. Figure 2).

1.2 Dirichlet Distribution

In Fig. 4 we compare the variance of our pathwise gradient for the Dirichlet distribution to the rsvi gradient estimator. We compute stochastic gradients of the ELBO for a Multinomial-Dirichlet model initialized at the exact posterior (where the exact gradient is zero). The Dirichlet distribution has 1995 components, and the single data point is a bag of words from a natural language document. We see that the pathwise gradient performs favorably over the entire range of the model hyperparameter α0\alpha_{0} considered. Note that as we crank up the shape augmentation setting BB, the rsvi variance approaches that of the pathwise gradient.As discussed in Sec. 6, the variance of the rsvi gradient estimator can also be reduced by dropping the score function-like term (at the cost of some bias).

1.3 Multivariate Normal

In Fig. 5 we use synthetic test functions to illustrate the amount of variance reduction that can be achieved with the OMT gradient estimator for the multivariate Normal distribution. The dimension is D=50D=50; the results are qualitatively similar for different dimensions.

2 Real World Datasets

In this section we investigate the performance of our gradient estimators for the Gamma, Beta, and multivariate Normal distributions in two variational inference tasks on real world datasets. Note that we include an additional experiment for the multivariate Normal distribution in the supplementary materials, see Sec. 9.11. All the experiments in this section were implemented in the Pyrohttp://pyro.ai probabilistic programming language.

2.2 Gaussian Process Regression

In this section we investigate the performance of our OMT gradient for the multivariate Normal distribution, Eqn. 19, in the context of a Gaussian Process regression task. We model the Mauna Loa CO2{\rm CO}_{2} data from (Keeling & Whorf, 2004) considered in (Rasmussen, 2004). We use a structured kernel that accommodates a long term linear trend as well as a periodic component. We fit the GP using a single-sample Monte Carlo ELBO gradient estimator and all D=468D=468 data points. The variational family is a multivariate Normal distribution with a Cholesky parameterization for the covariance matrix. Progress on the ELBO during the course of training is depicted in Fig. 7. We can see that the OMT gradient estimator has superior sample efficiency due to its lower variance. By iteration 270 the OMT gradient estimator has attained the same ELBO that the RT estimator attains at iteration 500. Since each iteration of the OMT estimator is ∼ ⁣1.9\sim\!1.9x slower than the corresponding RT iteration, the superior sample efficiency of the OMT estimator is largely canceled when judged by wall clock time. Nevertheless, the lower variance of the OMT estimator results in a higher ELBO than that obtained by the RT estimator.

Discussion and Future Work

We have seen that optimal transport offers a fruitful perspective on pathwise gradients. On the one hand it has helped us formulate pathwise gradients in situations where this was assumed to be impractical. On the other hand it has focused our attention on a particular notion of optimality, which led us to develop a new gradient estimator for the multivariate Normal distribution. A better understanding of this notion of optimality and, more broadly, a better understanding of when pathwise gradients are preferable over score function gradients (or vice versa) would be useful in guiding the practical application of these methods.

Since each solution of the transport equation Eqn. 12 yields an unbiased gradient estimator, the difference between any two such estimators can be thought of as a control variate. In the case of the multivariate Normal distribution, where computing the OMT gradient has a cost O(D3)\mathcal{O}(D^{3}), an attractive alternative to using vOMT\bm{v}^{\rm OMT} is to adaptively choose v\bm{v} during the course of optimization in direct analogy to adaptive control variate techniques. In future work we will explore this approach in detail, which promises lower variance than the OMT estimator at reduced computational cost.

The geometric picture from optimal transport—and thus the potential for non-trivial derivative applications—is especially rich for multivariate distributions. Here we have explored the multivariate Normal and Dirichlet distributions in some detail, but this just scratches the surface of multivariate distributions. It would be of general interest to develop pathwise gradients for a broader class of multivariate distributions, including for example mixture distributions. Rich distributions with low variance gradient estimators are of special interest in the context of SGVI, where the need to approximate complex posteriors demands rich families of distributions that lend themselves to stochastic optimization. In future work we intend to explore this connection further.

Acknowledgements

We thank Peter Dayan and Zoubin Ghahramani for feedback on a draft manuscript and other colleagues at Uber AI Labs—especially Noah Goodman and Theofanis Karaletsos— for stimulating conversations during the course of this work. We also thank Christian Naesseth for clarifying details of the experimental setup for the deep exponential family experiment in (Naesseth et al., 2017).

References

Supplementary Materials

For completeness we show explicitly that the formula

yields the correct gradient. Without loss of generality we assume that f(z)f(z) has no explicit dependence on θ\theta. Substituting Eqn. 21 for dzdθ\frac{dz}{d\theta} we have

In the second line we changed the order of integration and in the third we appealed to the fundamental theorem of calculus, assuming that f(z)f(z) is sufficiently regular that we can drop the boundary term at infinity.

Note that Eqn. 21 is the unique solution v=dzdθv=\frac{dz}{d\theta} to the one-dimensional version of the transport equation that satisfies the boundary condition lim⁡z→∞qθv=0\lim_{z\to\infty}q_{\bm{\theta}}v=0:

We consider an illustrative case where Eqn. 21 can be computed in closed form. For simplicity we consider the unit Normal distribution truncatedAs one would expect, Eqn. 21 yields the standard reparameterized gradient in the case of an non-truncated Normal distribution. Also note that the truncated unit normal is amenable to the reparameterization trick provided that one can compute the inverse error function erf−1{\rm erf}^{-1}. to the interval [0,κ][0,\kappa] with κ\kappa as the only free parameter. A simple computation yields

First, notice that for z=κz=\kappa we have dzdκ=1\frac{dz}{d\kappa}=1, which is what we would expect, since u=1u=1 is mapped to the rightmost edge of the interval at z=κz=\kappa, i.e. Fκ−1(1)=κF_{\kappa}^{-1}(1)=\kappa. Similarly we have dzdκ=0\frac{dz}{d\kappa}=0 for z=0z=0. For z∈(0,κ)z\in(0,\kappa) the derivative dzdκ\frac{dz}{d\kappa} interpolates smoothly between 0 and 1. This makes sense, since for a fixed value of uu as we get further into the tails of the distribution, nudging κ\kappa to the right has a correspondingly larger effect on z=Fκ−1(u)z=F_{\kappa}^{-1}(u), while it has a correspondingly smaller effect for uu in the bulk of the distribution.

1.2 Example: Univariate Mixture Distributions

Consider a mixture of univariate distributions:

If we have analytic control over the individual CDFs (or know how to approximate them and their derivatives w.r.t. the parameters) then we can immediately appeal to Eqn. 21. Concretely for derivatives w.r.t. the parameters of each component distribution we have:

for a mixture of univariate Normal distributions.

In Fig. 9 we demonstrate that the OMT gradient for a mixture of univariate Normal distributions can have much lower variance than the corresponding score function gradient. Here the mixture has two components with μ=(0,1)\bm{\mu}=(0,1) and σ=(1,1)\bm{\sigma}=(1,1). Note that using the reparameterization trick in this setting would be impractical.

2 The Multivariate Case

Suppose we are given a velocity field that satisfies the transport equation:

Then, as discussed in the main text, we can form the gradient estimator

That this gradient estimator is unbiased follows directly from the transport equation and divergence theorem:

and assume that qθfvθq_{\bm{\theta}}f{\bm{v}}^{\theta} is sufficiently well-behaved that we can drop the surface integral. This is just the multivariate generalization of the derivation in the previous section.

3 Multivariate Normal

Note that the transport equation for the multivariate distribution can be written in the form

The homogenous equation (i.e. the transport equation without the source term ∂log⁡q∂Lab\frac{\partial\log q}{\partial L_{ab}}) is then given by

In these coordinates it is evident that infinitesimal rotations, i.e. vector fields of the form

satisfyThese are in fact not the only solutions; in addition there are non-linear solutions. the homogenous equation, since

which is symmetric in ii and jj. This implies that the velocity field can be specified as the gradient of a scalar field (this is generally true for the OMT solution), i.e.

Note, however, that this is not the OMT solution we care about: it minimizes a different kinetic energy functional to the one we care about (namely it minimizes the kinetic energy functional in whitened coordinates and not in natural coordinates).

It is enough to show that the following expectation vanishes:Note that we can thus think of this term as a control variate.

where AijA_{ij} is an antisymmetric matrix. The sum in Eqn. 43 splits up into a sum of paired terms of the form

4 Natural Coordinates

We first show that the velocity field vRT\bm{v}^{\rm RT} that follows from the reparameterization trick satisfies the transport equation in the (given) coordinates z\bm{z}, where we have

Thus, the terms cancel term by term and the transport equation is satisfied.

What about the OMT gradient in the natural (given) coordinates z\bm{z}? To proceed we represent v\bm{v} as a linear vector field with symmetric and antisymmetric parts. Imposing the OMT condition determines the antisymmetric part. Imposing the transport equation determines the symmetric part. We find that

where SabS^{ab} is the unique symmetric matrix that satisfies the equation

To explicitly solve Eqn. 9.4 for SabS^{ab} we use SVD to write

where DD and UU are diagonal and orthogonal matrices, respectively. Then we have that

where ÷\div represents elementwise division and ⊗\otimes is the outer product. Note that a naive implementation of a gradient estimator based on Eqn. 49 would explicitly construct ξijab\xi_{ij}^{ab}, which has size quartic in the dimension. A more efficient implementation will instead make use of ξijab\xi_{ij}^{ab}’s structure as a sum of products and never explicitly constructs ξijab\xi_{ij}^{ab}.Our implementation can be found here: https://github.com/uber/pyro/blob/0.2.1/pyro/distributions/omt_mvn.py

In Fig. 10 we compare the performance of our OMT gradient for a bivariate Normal distribution to the reparameterization trick gradient estimator. We use a test function fθ(z)f_{\bm{\theta}}({\bm{z}}) for which we can compute the gradient exactly. We see that the OMT gradient estimator performs favorably over the entire range of parameters considered.

5 Gradient Variance for Linear Test Functions

We use the following example to give more intuition for when we expect OMT gradients for the multivariate Normal distribution to be lower variance than RT gradients. Let qθ(z)q_{\bm{\theta}}({\bm{z}}) be the unit normal distribution in DD dimensions. Consider the test function

and the derivative w.r.t. the off-diagonal elements of the Cholesky factor LL. A simple computation yields the total variance of the RT estimator:

So if we draw the parameters κi\kappa_{i} from a generic prior we expect the variance of the OMT estimator to be about half of that of the RT estimator. Concretely, if κi∼N(0,1)\kappa_{i}\sim\mathcal{N}(0,1) then the variance of the OMT estimator will be exactly half that of the RT estimator in expectation. While this computation is for a very specific case—a linear test function and a unit normal qθ(z)q_{\bm{\theta}}({\bm{z}})—we find that this magnitude of variance reduction is typical.

6 The Lugannani-Rice Approximation

Saddlepoint approximation methods take advantage of cumulant generating functions (CGFs) to construct (often very accurate) approximations to probability density functions in situations where full analytic control is intractable.We refer the reader to (Butler, 2007) for an overview. These methods are also directly applicable to CDFs, where a particularly useful approximation—often used by statisticians to estimate various tail probabilities—has been developed by Lugannani and Rice (Lugannani & Rice, 1980). This approximation—after additional differentiation w.r.t. the parameters of the distribution qθ(z)q_{\bm{\theta}}(z)—forms the basis of our approximate formulas for pathwise gradients for the Gamma, Beta and Dirichlet distributions in regions of (z,θ)(z,\theta) where the (marginal) density is approximately gaussian. As we will see these approximations attain high accuracy.

For completeness we briefly describe the Lugannani-Rice approximation. It is given by:

7 Gamma Distribution

Our numerical recipe for dzdα\frac{dz}{d\alpha} for the standard Gamma distribution with β=1\beta=1 divides (z,α(z,\alpha) space into three regions. If z<0.8z<0.8 we use the Taylor series expansion given in the main text. If α>8\alpha>8 we use the following set of expressions derived from the Lugannani-Rice approximation. Away from the singularity, for z≷α±δ⋅αz\gtrless\alpha\pm\delta\cdot\alpha, we use:

Near the singularity, i.e. for ∣z−α∣≤δ⋅α|z-\alpha|\leq\delta\cdot\alpha, we use:

Note that Eqn. 59 is derived from Eqn. 58 by a Taylor expansion in powers of (z−α)(z-\alpha). We set δ=0.1\delta=0.1, which is chosen to balance use of Eqn. 58 (which is more accurate) and Eqn. 59 (which is more numerically stable for z≈αz\approx\alpha). Finally, in the remaining region (z>0.8z>0.8 and α<8\alpha<8) we use a bivariate rational polynomial approximation f(z,α)=exp⁡(p(z,α)q(z,α))f(z,\alpha)=\exp\left(\frac{p(z,\alpha)}{q(z,\alpha)}\right) where p,qp,q are polynomials in the coordinates log⁡(z/α)\log(z/\alpha) and log⁡(α)\log(\alpha), with terms up to order 2 in log⁡(z/α)\log(z/\alpha) and order 3 in log⁡(α)\log(\alpha). We fit the rational approximation using least squares on 15696 random (z,α)(z,\alpha) pairs with α\alpha sampled log uniformly between 0.00001 and 10, and zz sampled conditioned on α\alpha. Our complete approximation for dzdα\frac{dz}{d\alpha} is unit tested to have relative accuracy of 0.0005 on a wide range of inputs.

8 Beta Distribution

The CDF of the Beta distribution is given by

where B(z;α,β)B(z;\alpha,\beta) and B(α,β)B(\alpha,\beta) are the incomplete beta function and beta function, respectively. Our numerical recipe for computing dzdα\frac{dz}{d\alpha} and dzdβ\frac{dz}{d\beta} for the Beta distribution divides (z,α,β(z,\alpha,\beta) space into three sets of regions. First suppose that z≪1z\ll 1. Then just like for the Gamma distribution, we can compute a Taylor series of B(z;α,β)B(z;\alpha,\beta) in powers of zz

that can readily be differentiated w.r.t. either α\alpha or β\beta. Combined with the derivatives of the beta function,

this gives a complete recipe for approximating dzdα\frac{dz}{d\alpha} and dzdβ\frac{dz}{d\beta} for small zz.Here ψ(⋅)\psi(\cdot) is the digamma function, which is available in most advanced tensor libraries. By appealing to the symmetry of the Beta distribution

we immediately gain approximations to dzdα\frac{dz}{d\alpha} and dzdβ\frac{dz}{d\beta} for 1−z≪11-z\ll 1. It remains to specify when these various approximations are applicable. Let us define ξ=z(1−z)(α+β)\xi=z(1-z)(\alpha+\beta). Empirically we find that these approximations are accurate for dzdα\frac{dz}{d\alpha} if

with the conditions flipped for dzdβ\frac{dz}{d\beta}. Depending on the precise region, we use 8 to 10 terms in the Taylor series.

Next we describe the set of approximations we derived from the Lugannani-Rice approximation and that we find to be accurate for α>6\alpha>6 and β>6\beta>6. By Eqn. 63 it is sufficient to describe our approximation for dzdα\frac{dz}{d\alpha}. First define σ=αβ(α+β)α+β+1\sigma=\frac{\sqrt{\alpha\beta}}{(\alpha+\beta)\sqrt{\alpha+\beta+1}}, the standard deviation of the Beta distribution. Then away from the singularity, for z≷αα+β±ϵ⋅σz\gtrless\frac{\alpha}{\alpha+\beta}\pm\epsilon\cdot\sigma, we use:

Near the singularity, i.e. for ∣z−αα+β∣≤ϵ⋅σ|z-\frac{\alpha}{\alpha+\beta}|\leq\epsilon\cdot\sigma, we use:

We set ϵ=0.1\epsilon=0.1, which is chosen to balance numerical accuracy and numerical stability (just as in the case of the Gamma distribution).

Finally, in the remaining region we use a rational multivariate polynomial approximation

where p,qp,q are polynomials in the three coordinates log⁡(z)\log(z), log⁡(α/z)\log(\alpha/z), and log⁡((α+β)z/α)\log((\alpha+\beta)z/\alpha) with terms up to order 2, 2, and 3 in the respective coordinates. The rational approximation was minimax fit to 2842 points in the remaining region for 0.01<α,β<10000.01<\alpha,\beta<1000. Test points were randomly sampled using log uniform sampling of α,β\alpha,\beta and stratified sampling of zz conditioned on α,β\alpha,\beta. Minimax fitting achieved about half the maximum error of simple least squares fitting. Our complete approximation for dzdα\frac{dz}{d\alpha} and dzdβ\frac{dz}{d\beta} is unit tested to have relative accuracy of 0.001 on a wide range of inputs.

9 Dirichlet Distribution

For completeness we record the general version of the formula for the pathwise gradient (given implicitly in the main text):

We want to confirm that Eqn. 66 satisfies the transport equation for each choice of j=1,...,nj=1,...,n:

Treating zjz_{j} as a function of z−j=(z1,...,zj−1,zj+1,...,zn){\bf z}_{-j}=(z_{1},...,z_{j-1},z_{j+1},...,z_{n}) everywhere and introducing obvious shorthand for FBeta(⋅)F_{\rm Beta}(\cdot) and Beta(⋅)\rm{Beta}(\cdot) we have:

where (log⁡B)′\left(\log B\right)^{\prime} is differentiated w.r.t. the argument of B(zj)B(z_{j}). We further have that

it becomes clear by comparing the individual terms that everything cancels identically and so Eqn. 67 is in fact satisfied by the velocity field in Eqn. 66.

Finally, we note that Eqn. 66 is not the OMT solution in the coordinates z−j\bm{z}_{-j}. It is the OMT solution in some coordinate system, but it is not readily apparent which coordinate system that might be.

10 Student’s t-Distribution

As another example of how to compute pathwise gradients consider Student’s t-distribution. Although we have not done so ourselves, it should be straightforward to compute an accurate approximation to Eqn. 21. In the absence of such an approximation, however, we can still get a pathwise gradient for the Student’s t-distribution by composing the Normal and Gamma distributions:

Since sampling zz like this introduces an auxiliary random degree of freedom, pathwise gradients dzdν\frac{dz}{d\nu} computed using Eqn. 68 will exhibit a larger variance than a direct computation of Eqn. 21 would yield.Note, however, that this additional variance will decrease as ν\nu increases. The point is that no additional work is needed to obtain this particular form of the pathwise gradient: just use pathwise gradients for the Gamma and Normal distributions and the sampling procedure in Eqn. 68.

11 Baseball Experiment

To gain more insight into when we expect the OMT gradient estimator for the multivariate Normal distribution to outperform the RT gradient estimator, we conduct an additional experiment. We consider a model for repeated binary trial data (baseball players at bat) using the data in (Efron & Morris, 1975) and the modeling setup in (Stan Manual, 2017) with partial pooling. There are 18 baseball players and the data consists of 45 hits/misses for each player. The model has two global latent variables and 18 local latent variables so that the posterior is 20-dimensional. Specifically, the two global latent random variables are ϕ\phi and κ\kappa, with priors Uniform(0,1)\rm{Uniform}(0,1) and Pareto(1,1.5)∝κ−5/2\rm{Pareto}(1,1.5)\propto\kappa^{-5/2}, respectively. The local latent random variables are given by θi\theta_{i} for i=0,...,17i=0,...,17, with p(θi)=Beta(θi∣α=ϕκ,β=(1−ϕ)κ)p(\theta_{i})=\rm{Beta}(\theta_{i}|\alpha=\phi\kappa,\beta=(1-\phi)\kappa). The data likelihood factorizes into 45 Bernoulli observations with mean chance of success θi\theta_{i} for each player ii. The variational approximation is formed in the unconstrained space {logit(ϕ),log⁡(κ−1),logit(θi)}\{\rm{logit}(\phi),\log(\kappa-1),\rm{logit}(\theta_{i})\} and consists of a multivariate Normal distribution with a full-rank Cholesky factor L\bm{L}. We use the Adam optimizer for training with a learning rate of 5×10−35\times 10^{-3} (Kingma & Ba, 2014).

For this particular model mean field SGVI performs reasonably well, since correlations between the latent random variables are not particularly strong. If we initialize L\bm{L} near the identity, we find that the OMT and RT gradient estimators perform nearly identically, with the difference that the former has an increased computational cost of about 25% per iteration. If, however, we initialize L\bm{L} far from the identity—so that the optimizer has to traverse a considerable distance in L\bm{L} space where the covariance matrix exhibits strong correlations—we find that the OMT estimator makes progress more quickly than the RT estimator and converges to a higher ELBO, see Fig. 12. Generalizing from this, we expect the OMT gradient estimator for the multivariate Normal distribution to exhibit better sample efficiency than the RT estimator in problems where the covariance matrix exhibits strong correlations. This is indeed the case for the GP experiment in the main text, where the learned kernel induces strong temporal correlations.

12 Experimental Details

As noted in the main text, we use single-sample gradient estimators in all experiments. Unless noted otherwise, we always include the score function term for rsvi.

In all cases the gradients can be computed analytically, which makes it easier to reliably estimate the variance of the gradient estimators.

12.2 Sparse Gamma def

Following (Naesseth et al., 2017), we use analytic expressions for each entropy term (as opposed to using the sampling estimate). We use the adaptive step sequence ρn\rho^{n} proposed by (Kucukelbir et al., 2016) and also used in (Naesseth et al., 2017), which combines rmsprop (Tieleman & Hinton, 2012) and Adagrad (Duchi et al., 2011):

Here n=1,2,...n=1,2,... is the iteration number and the operations in Eqn. 72 are to be understood element-wise. In our case the gradient g^n\hat{g}^{n} is always a single-sample estimate. We fix δ=10−16\delta=10^{-16} and t=0.1t=0.1. In contrast to (Kucukelbir et al., 2016) but in line with (Naesseth et al., 2017) we initialize s0s_{0} at zero. To choose η\eta we did a grid search for each gradient estimator and each of the two model variants. Specifically, for each η\eta we did 100 training iterations for three trials with different random seeds and then chose the η\eta that yielded the highest mean ELBO after 100 iterations. This procedure led to the selection of η=4.5\eta=4.5 for the first model variant and η=30\eta=30 for the second model variant (note that within each model variant the gradient estimators preferred the same value of η\eta). For the first model variant we included the score function-like term in the rsvi gradient estimator, while we did not include it for the second model variant, as we found that this hurt performance. In both cases we used the shape augmentation setting B=4B=4, which was also used for the results reported in (Naesseth et al., 2017). After fixing η\eta we trained the model for 2000 iterations, initializing with another random number seed. The figure in the main text shows the training curves for that single run. We confirmed that other random number seeds give similar results. A reference implementation can be found here: https://github.com/uber/pyro/blob/0.2.1/examples/sparse_gamma_def.py

12.3 Gaussian Process Regression

We used the Adam optimizer (Kingma & Ba, 2014) to optimize the ELBO with single-sample gradient estimates. We chose the Adam hyperparameters by doing a grid search over the learning rate and β1\beta_{1}. For each combination (lr,β1)({\rm lr},\beta_{1}) we did 20 training iterations for three trials with different random seeds and then chose the combination that yielded the highest mean ELBO after 20 iterations. This procedure led to the selection of a learning rate of 0.0300.030 and β1=0.50\beta_{1}=0.50 for both gradient estimators (OMT and reparameterization trick). We then trained the model for 500 iterations, initializing with another random number seed. The figure in the main text shows the training curves for that single run. We confirmed that other random number seeds give similar results.