Reparameterizing the Birkhoff Polytope for Variational Permutation Inference

Scott W. Linderman, Gonzalo E. Mena, Hal Cooper, Liam Paninski, John P. Cunningham

Introduction

Permutation inference is central to many modern machine learning problems. Identity management (Guibas, 2008) and multiple-object tracking (Shin et al., 2005; Kondor et al., 2007) are fundamentally concerned with finding a permutation that maps an observed set of items to a set of canonical labels. Ranking problems, critical to search and recommender systems, require inference over the space of item orderings (Meilă et al., 2007; Lebanon and Mao, 2008; Adams and Zemel, 2011). Furthermore, many probabilistic models, like preferential attachment network models (Bloem-Reddy and Orbanz, 2016) and repulsive point process models (Rao et al., 2016), incorporate a latent permutation into their generative processes; inference over model parameters requires integrating over the set of permutations that could have given rise to the observed data. In neuroscience, experimentalists now measure whole-brain recordings in C. Elegans (Kato et al., 2015; Nguyen et al., 2016), a model organism with a known synaptic network (White et al., 1986); a current challenge is matching the observed neurons to corresponding nodes in the reference network. In Section 5, we address this problem from a Bayesian perspective in which permutation inference is a central component of a larger inference problem involving unknown model parameters and hierarchical structure.

The task of computing optimal point estimates of permutations under various loss functions has been well studied in the combinatorial optimization literature (Kuhn, 1955; Munkres, 1957; Lawler, 1963). However, many probabilistic tasks, like the aforementioned neural identity inference problem, require reasoning about the posterior distribution over permutation matrices. A variety of Bayesian permutation inference algorithms have been proposed, leveraging sampling methods (Diaconis, 1988; Miller et al., 2013; Harrison and Miller, 2013), Fourier representations (Kondor et al., 2007; Huang et al., 2009), as well as convex (Lim and Wright, 2014) and continuous (Plis et al., 2011) relaxations for approximating the posterior distribution. Here, we address this problem from an alternative direction, leveraging stochastic variational inference (Hoffman et al., 2013) and reparameterization gradients (Rezende et al., 2014; Kingma and Welling, 2014) to derive a scalable and efficient permutation inference algorithm.

Section 2 lays the necessary groundwork, introducing definitions, prior work on permutation inference, variational inference, and continuous relaxations. Section 3 presents our primary contribution: a pair of transformations that enable variational inference over doubly-stochastic matrices, and, in the zero-temperature limit, permutations, via stochastic variational inference. In the process, we show how these transformations connect to recent work on discrete variational inference (Maddison et al., 2017; Jang et al., 2017; Balog et al., 2017). Sections 4 and 5 present a variety of experiments that illustrate the benefits of the proposed variational approach. Further details are in the supplement.

Background

A permutation is a bijective mapping of a set onto itself. When this set is finite, the mapping is conveniently represented as a binary matrix X∈{0,1}N×N{X\in\{0,1\}^{N\times N}} where Xm,n=1{X_{m,n}=1} implies that element mm is mapped to element nn. Since permutations are bijections, both the rows and columns of XX must sum to one. From a geometric perspective, the Birkhoff-von Neumann theorem states that the convex hull of the set of permutation matrices is the set of doubly-stochastic matrices; i.e. non-negative square matrices whose rows and columns sum to one. The set of doubly-stochastic matrices is known as the Birkhoff polytope, and it is defined by,

2 Related Work

A number of previous works have considered approximate methods of posterior inference over the space of permutations. When a point estimate will not suffice, sampling methods like Markov chain Monte Carlo (MCMC) algorithms may yield a reasonable approximate posterior for simple problems (Diaconis, 1988). Harrison and Miller (2013) developed an importance sampling algorithm that fills in count matrices one row at a time, showing promising results for matrices with O(100)O(100) rows and columns. Li et al. (2013) considered using the Hungarian algorithm within a Perturb-and-MAP algorithm for approximate sampling. Another line of work considers inference in the spectral domain, approximating distributions over permutations with the low frequency Fourier components (Kondor et al., 2007; Huang et al., 2009). Perhaps most relevant to this work, Plis et al. (2011) propose a continuous relaxation from permutation matrices to points on a hypersphere, and then use the von Mises-Fisher (vMF) distribution to model distributions on the sphere’s surface. We will relax permutations to points in the Birkhoff polytope and derive temperature-controlled densities such that as the temperature goes to zero, the distribution converges to an atomic density on permutation matrices. This will enable efficient variational inference with the reparameterization trick, which we describe next.

3 Variational inference and the reparameterization trick

Given an intractable model with data yy, likelihood p(y ∣ x)p(y\,|\,x), and prior p(x)p(x), variational Bayesian inference algorithms aim to approximate the posterior distribution p(x ∣ y)p(x\,|\,y) with a more tractable distribution q(x;θ)q(x;\theta), where “tractable” means that, at a minimum, we can sample qq and evaluate it pointwise (including its normalization constant) (Blei et al., 2017). We find this approximate distribution by searching for the parameters θ\theta that minimize the Kullback-Leibler (KL) divergence between qq and the true posterior, or equivalently, maximize the evidence lower bound (ELBO),

Perhaps the simplest method of optimizing the ELBO is stochastic gradient ascent. However, computing ∇θL(θ)\nabla_{\theta}\mathcal{L}(\theta) requires some care since the ELBO contains an expectation with respect to a distribution that depends on these parameters.

When xx is a continuous random variable, we can sometimes leverage the reparameterization trick (Salimans and Knowles, 2013; Kingma and Welling, 2014). Specifically, in some cases we can simulate from qq via the following equivalence,

where rr is a distribution on the “noise” zz and where g(z;θ)g(z;\theta) is a deterministic and differentiable function. The reparameterization trick effectively “factors out” the randomness of qq. With this transformation, we can bring the gradient inside the expectation as follows,

This gradient can be estimated with Monte Carlo, and, in practice, this leads to lower variance estimates of the gradient than, for example, the score function estimator (Williams, 1992; Glynn, 1990).

Critically, the gradients in (1) can only be computed if xx is continuous. Recently, Maddison et al. (2017) and Jang et al. (2017) proposed the “Gumbel-softmax” method for discrete variational inference. It is based on the following observation: discrete probability mass functions q(x;θ)q(x;\theta) can be seen as densities with atoms on the vertices of the simplex; i.e. on the set of one-hot vectors {en}n=1N{\{e_{n}\}_{n=1}^{N}}, where en=(0,0,…,1,…,0)T{e_{n}=(0,0,\ldots,1,\ldots,0)^{\mathsf{T}}} is a length-NN binary vector with a single 1 in the nn-th position. This motivates a natural relaxation: let q(x;θ)q(x;\theta) be a density on the interior of the simplex instead, and anneal this density such that it converges to an atomic density on the vertices. Fig. 1a illustrates this idea. Gumbel random variates, are mapped through a temperature-controlled softmax function, {g_{\tau}(\psi)=\big{[}e^{\psi_{1}/\tau}/Z,\ldots,e^{\psi_{N}/\tau}/Z\big{]}}, where Z=∑n=1Neψn/τ{Z=\sum_{n=1}^{N}e^{\psi_{n}/\tau}}, to obtain points in the simplex. As τ\tau goes to zero, the density concentrates on one-hot vectors. We build on these ideas for variational permutation inference.

Variational permutation inference via reparameterization

The Gumbel-softmax method scales linearly with the support of the discrete distribution, rendering it prohibitively expensive for direct use on the set of N!N! permutations. Instead, we develop two transformations to map O(N2)O(N^{2})-dimensional random variates to points in or near the Birkhoff polytope.While Gumbel-softmax does not immediately extend to permutation inference, the methods presented herein easily extend to categorical inference. We explored this direction experimentally and show results in the supplement. Like the Gumbel-softmax method, these transformations will be controlled by a temperature that concentrates the resulting density near permutation matrices. The first method is a novel “stick-breaking” construction; the second rounds points toward permutations with the Hungarian algorithm. We present these in turn and then discuss their relative merits. We provide further implementation details for both methods in the supplement.

Stick-breaking is well-known as a construction for the Dirichlet process (Sethuraman, 1994); here we show how the same intuition can be extended to more complex discrete objects. Let BB be a matrix in (N−1)×(N−1){^{(N-1)\times(N-1)}}; we will transform it into a doubly-stochastic matrix X∈N×N{X\in^{N\times N}} by filling in entry by entry, starting in the top left and raster scanning left to right then top to bottom. Denote the (m,n)(m,n)-th entries of BB and XX by βmn\beta_{mn} and xmn{x}_{mn}, respectively.

Each row and column has an associated unit-length “stick” that we allot to its entries. The first entry in the matrix is given by x11=β11x_{11}=\beta_{11}. As we work left to right in the first row, the remaining stick length decreases as we add new entries. This reflects the row normalization constraints. The first row follows the standard stick-breaking construction,

This is illustrated in Fig. 1b, where points in the unit square map to points in the simplex. Here, the blue dots are two-dimensional N(0,4I)\mathcal{N}(0,4I) variates mapped through a coordinate-wise logistic function.

Subsequent rows are more interesting, requiring a novel advance on the typical uses of stick breaking. Here we need to conform to row and column sums (which introduce upper bounds), and a lower bound induced by stick remainders that must allow completion of subsequent sum constraints. Specifically, the remaining rows must now conform to both row- and column-constraints. That is,

Moreover, there is also a lower bound on xmnx_{mn}. This entry must claim enough of the stick such that what is leftover fits within the confines imposed by subsequent column sums. That is, each column sum places an upper bound on the amount that may be attributed to any subsequent entry. If the remaining stick exceeds the sum of these upper bounds, the matrix will not be doubly-stochastic. Thus,

where θ={μmn,νmn2}m,n=1N{\theta=\{\mu_{mn},\nu^{2}_{mn}\}_{m,n=1}^{N}} are the mean and variance parameters of the intermediate Gaussian matrix Ψ\Psi, σ(u)=(1+e−u)−1{\sigma(u)=(1+e^{-u})^{-1}} is the logistic function, and τ\tau is a temperature parameter. As τ→0\tau\to 0, the values of βmn\beta_{mn} are pushed to either zero or one, depending on whether the input to the logistic function is negative or positive, respectively. As a result, the doubly-stochastic output matrix XX is pushed toward the extreme points of the Birkhoff polytope, the permutation matrices. This map is illustrated in Fig. 1c for permutations of N=3{N=3} elements. Here, the blue dots are samples of BB with μmn=0\mu_{mn}=0, νmn=2\nu_{mn}=2, and τ=1\tau=1.

We compute gradients of this transformation with automatic differentiation. Since this transformation is “feed-forward,” its Jacobian is lower triangular. The determinant of the Jacobian, necessary for evaluating the density qτ(X;θ)q_{\tau}(X;\theta), is a simple function of the upper and lower bounds and is derived in Appendix B. While this map is peculiar in its reliance on an ordering of the elements, as discussed in Section 3.3, it is a novel transformation to the Birkhoff polytope that supports gradient-based variational permutation inference.

2 Rounding toward permutation matrices

While relaxing permutations to the Birkhoff polytope is intuitively appealing, it is not strictly required. For example, consider the following procedure for sampling a point near the Birkhoff polytope:

Map M→M~M\to\widetilde{M}, a point in the Birkhoff polytope, using the Sinkhorn-Knopp algorithm;

Set Ψ=M~+V⊙Z{\Psi=\widetilde{M}+V\odot Z} where ⊙\odot denotes elementwise multiplication;

Find round(Ψ)\mathsf{round}(\Psi), the nearest permutation matrix to Ψ\Psi, using the Hungarian algorithm;

Output X=τΨ+(1−τ)round(Ψ){X=\tau\Psi+(1-\tau)\mathsf{round}(\Psi)}.

This procedure defines a mapping X=gτ(Z;θ){X=g_{\tau}(Z;\theta)} with θ={M,V}{\theta=\{M,V\}}. When the elements of ZZ are independently sampled from a standard normal distribution, it implicitly defines a distribution over matrices XX parameterized by θ{\theta}. Furthermore, as τ\tau goes to zero, the density concentrates on permutation matrices. A simple example is shown in Fig. 1d, where M=1N11T{M=\tfrac{1}{N}\boldsymbol{1}\boldsymbol{1}^{\mathsf{T}}} with 1\boldsymbol{1} a vector of all ones, V=0.4211T{V=0.4^{2}\boldsymbol{1}\boldsymbol{1}^{\mathsf{T}}}, and τ=0.5{\tau=0.5}. We use this procedure to define a variational distribution with density qτ(X;θ)q_{\tau}(X;\theta).

To compute the ELBO and its gradient (1), we need to evaluate qτ(X;θ)q_{\tau}(X;\theta). By construction, steps (i) and (ii) involve differentiable transformations of parameter MM to set the mean close to the Birkhoff polytope, but since these do not influence the distribution of ZZ, the non-invertibility of the Sinkhorn-Knopp algorithm poses no problems. Had we applied this algorithm directly to ZZ, this would not be true. The challenge in computing the density stems from the rounding in steps (iv) and (v).

The Jacobian is more challenging due to the non-differentiability of round\mathsf{round}. However, since the nearest permutation output only changes at points that are equidistant from two or more permutation matrices, round\mathsf{round} is a piecewise constant function with discontinuities only at a set of points with zero measure. Thus, the change of variables theorem still applies.

With the inverse and its Jacobian, we have

where zmn=[gτ−1(X;θ)]mn{z_{mn}=[g_{\tau}^{-1}(X;\theta)]_{mn}} and νmn\nu_{mn} are the entries of VV. In the zero-temperature limit we recover a discrete distribution on permutation matrices; otherwise the density concentrates near the vertices as τ→0{\tau\to 0}. This transformation leverages computationally efficient algorithms like Sinkhorn-Knopp and the Hungarian algorithm to define a temperature-controlled variational distribution near the Birkhoff polytope, and it enjoys many theoretical and practical benefits.

3 Theoretical considerations

The stick-breaking and rounding transformations introduced above each have their strengths and weaknesses. Here we list some of their conceptual differences. While these considerations aid in understanding the differences between the two transformations, the ultimate test is in their empirical performance, which we study in Section 4.

Rounding uses the O(N3)O(N^{3}) Hungarian algorithm in its sampling process, whereas stick-breaking has O(N2)O(N^{2}) complexity. In practice, the stick-breaking computations are slightly more efficient.

Rounding can easily incorporate constraints. If certain mappings are invalid, i.e. xmn≡0{x_{mn}\equiv 0}, they are given an infinite cost in the Hungarian algorithm. This is hard to do this with stick breaking as it would change the computation of the upper and lower bounds. (In both cases, constraints of the form xmn≡1x_{mn}\equiv 1 simply reduce the dimension of the inference problem.)

Stick-breaking introduces a dependence on ordering. While the mapping is bijective, a desired distribution on the Birkhoff polytope may require a complex distribution for BB. Rounding, by contrast, is more “symmetric” in this regard.

In summary, stick-breaking offers an intuitive advantage—an exact relaxation to the Birkhoff polytope—but it suffers from its sensitivity to ordering and its inability to easily incorporate constraints. As we show next, these concerns ultimately lead us to favor the rounding based methods in practice.

Synthetic Experiments

We are interested in two principal questions: (i) how well can the stick-breaking and rounding re-parameterizations of the Birkhoff polytope approximate the true posterior distribution over permutations in tractable, low-dimensional cases? and (ii) when do our proposed continuous relaxations offer advantages over alternative Bayesian permutation inference algorithms?

To assess the quality of our approximations for distributions over permutations, we considered a toy matching problem in which we are given the locations of NN cluster centers and a corresponding set of NN observations, one for each cluster, corrupted by Gaussian noise. Moreover, the observations are permuted so there is no correspondence between the order of observations and the order of the cluster centers. The goal is to recover the posterior distribution over permutations. For N=6N=6, we can explicitly enumerate the N!=720N!=720 permutations and compute the posterior exactly.

As a baseline, we consider the Mallows distribution Mallows (1957) with density over a permutations ϕ\phi given by pθ,ϕ0(ϕ)∝exp⁡(−θd(ϕ,ϕ0))p_{\theta,\phi_{0}}(\phi)\propto\exp(-\theta d(\phi,\phi_{0})), where ϕ0\phi_{0} is a central permutation, d(ϕ,ϕ0)=∑i=1N∣ϕ(i)−ϕ0(i)∣{d(\phi,\phi_{0})=\sum_{i=1}^{N}|\phi(i)-\phi_{0}(i)|} is a distance between permutations, and θ\theta controls the spread around ϕ0\phi_{0}. This is the most popular exponential family model for permutations, but since it is necessarily unimodal, it can fail to capture complex permutation distributions.

We measured the discrepancy between true posterior and an empirical estimate of the inferred posteriors using using the Battacharya distance (BD). We fit qτ(X;θ)q_{\tau}(X;\theta) with an annealing schedule for both stick-breaking and rounding transformations, sampled the variational posterior, and rounded the samples to the nearest permutation matrix with the Hungarian algorithm. For the Mallows distribution, we set ϕ0\phi_{0} to the MAP estimate, also found with the Hungarian algorithm, and sampled using MCMC.

We found our method outperforms the simple Mallows distribution and reasonably approximates non-trivial distributions over permutations. Fig 2 illustrates our findings, showing (a) sample experiment configurations; (b) examples of inferred, discrete, posteriors for stick breaking, rounding, and Mallows at various levels of noise; and (c) histogram of Battacharya distance. The latter are summarized in Table 1.

Inferring neuron identities in C. elegans

Finally, we consider an application motivated by the study of the neural dynamics in C. elegans. This worm is a model organism in neuroscience as its neural network is stereotyped from animal to animal and its complete neural wiring diagram is known (Varshney et al., 2011). We represent this network, or connectome, as a binary adjacency matrix A∈{0,1}N×N{A\in\{0,1\}^{N\times N}}, shown in Fig. 3a. The hermaphrodite has N=278{N=278} somatic neurons, and (undirected) synaptic connections between neurons mm and nn are denoted by Amn=1A_{mn}=1.

Modern recording technology enables simultaneous measurements of hundreds of these neurons simultaneously (Kato et al., 2015; Nguyen et al., 2016). However, matching the observed neurons to nodes in the reference connectome is still a manual task. Experimenters consider the location of the neuron along with its pattern of activity to perform this matching, but the process is laborious and the results prone to error. We prototype an alternative solution, leveraging the location of neurons and their activity in a probabilistic model. We resolve neural identity by integrating different sources of information from the connectome, some covariates (e.g. position) and neural dynamics. Moreover, we combine information from many individuals to facilitate identity resolution. The hierarchical nature of this problem and the plethora of prior constraints and observations motivates our Bayesian approach.

Our goal is to infer WW and {X(j)}\{X^{(j)}\} given {Y(j)}\{Y^{(j)}\} using variational permutation inference. We place a standard Gaussian prior on WW and a uniform prior on X(j)X^{(j)}, and we use the rounding transformation to approximate the posterior, p(W,{X(j)} ∣ {Y(j)})∝p(W)∏mp(Y(j) ∣ W,X(j)) p(X(j)p(W,\{X^{(j)}\}\,|\,\{Y^{(j)}\})\propto p(W)\prod_{m}p(Y^{(j)}\,|\,W,{X^{(j)}})\,p({X^{(j)}}).

Finally, we use neural position along the worm’s body to constrain the possible neural identities for a given neuron. We use the known positions of each neuron (Lints et al., 2005), approximating the worm as a one-dimensional object with neurons locations distributed as in Fig. 3c. Then, given reported positions of the neurons, we can conceive a binary constraint matrix C(j)C^{(j)} so that Cmn(j)=1C^{(j)}_{mn}=1 if (observed) neuron mm is close enough to (canonical) neuron nn; i.e., if their distance is smaller than a tolerance ν\nu. We enforce this constraint during inference by zeroing corresponding entries in the parameter matrix MM described in 3.2. This modeling choice greatly reduces the number parameters of the model, and facilitates inference.

We find that our method outperforms each baseline. Fig. 4a illustrates convergence to a better solution for a certain parameter configuration. Moreover, Fig. 4b and Fig. 4c show that our method outperforms alternatives when there are many possible candidates and when only a small proportion of neurons are known with certitude. Fig. 4c also shows that these Bayesian methods benefit from combining information across many worms.

Altogether, these results indicate our method enables a more efficient use of information than its alternatives. This is consistent with other results showing faster convergence of variational inference over MCMC (Blei et al., 2017), especially with simple Metropolis-Hastings proposals. We conjecture that MCMC would eventually obtain similar if not better results, but the local proposals—swapping pairs of labels—leads to slow convergence. On the other hand, Fig 4a shows that our method converges much more quickly while still capturing a distribution over permutations, as shown by the overall variance of the samples in Fig 4d and the individual samples in Fig 4e.

Conclusion

Our results provide evidence that variational permutation inference is a valuable tool, especially in complex problems like neural identity inference where information must be aggregated from disparate sources in a hierarchical model. As we apply this to real neural recordings, we must consider more realistic, nonlinear models of neural dynamics. Here, again, we expect variational methods to shine, leveraging automatic gradients of the relaxed ELBO to efficiently explore the space of variational posterior distributions.

We thank Christian Naesseth for many helpful discussions. SWL is supported by the Simons Collaboration on the Global Brain (SCGB) Postdoctoral Fellowship (418011). HC is supported by Graphen, Inc. LP is supported by ARO MURI W91NF-12-1-0594, the SCGB, DARPA SIMPLEX N66001-15-C-4032, IARPA MICRONS D16PC00003, and ONR N00014-16-1-2176. JPC is supported by the Sloan Foundation, McKnight Foundation, and the SCGB.

References

Appendix A Alternative methods of discrete variational inference

We can gain insight and intuition about the stick-breaking and rounding transformations by considering their counterparts for discrete, or categorical, variational inference. Continuous relaxations are an appealing approach for this problem, affording gradient-based inference with the reparameterization trick. First we review the Gumbel-softmax method [Maddison et al., 2017, Jang et al., 2017, Kusner and Hernández-Lobato, 2016]—a recently proposed method for discrete variational inference with the reparameterization trick—then we discuss analogs of our permutation and rounding transformations for the categorical case. These can be considered alternatives to the Gumbel-softmax method, which we compare empirically in Appendix A.5.

Recently there have been a number of proposals for extending the reparameterization trick [Rezende et al., 2014, Kingma and Welling, 2014] to high dimensional discrete problemsDiscrete inference is only problematic in the high dimensional case, since in low dimensional problems we can enumerate the possible values of xx and compute the normalizing constant p(y)=∑xp(y,x)p(y)=\sum_{x}p(y,x). by relaxing them to analogous continuous problems [Maddison et al., 2017, Jang et al., 2017, Kusner and Hernández-Lobato, 2016]. These approaches are based on the following observation: if x∈{0,1}Nx\in\{0,1\}^{N} is a one-hot vector drawn from a categorical distribution, then the support of p(x)p(x) is the set of vertices of the N−1N-1 dimensional simplex. We can represent the distribution of xx as an atomic density on the simplex.

Viewing xx as a vertex of the simplex motivates a natural relaxation: rather than restricting ourselves to atomic measures, consider continuous densities on the simplex. To be concrete, suppose the density of xx is defined by the transformation,

The Gumbel distribution leads to a nicely interpretable model: adding i.i.d. Gumbel noise to log⁡θ{\log\theta} and taking the argmax yields an exact sample from the normalized probability mass function θˉ\bar{\theta}, where θˉn=θn/∑m=1Nθm{\bar{\theta}_{n}=\theta_{n}/\sum_{m=1}^{N}\theta_{m}} [Gumbel, 1954]. The softmax is a natural relaxation. As the temperature τ\tau goes to zero, the softmax converges to the argmax function. Ultimately, however, this is just a continuous relaxation of an atomic density to a continuous density.

Stick-breaking and rounding offer two alternative ways of constructing a relaxed version of a discrete random variable, and both are amenable to reparameterization. However, unlike the Gumbel-Softmax, these relaxations enable extensions to more complex combinatorial objects, notably, permutations.

A.2 Stick-breaking

parameterized by θ=(μn,νn)n=1N−1{\theta=(\mu_{n},\nu_{n})_{n=1}^{N-1}}. Then map this to the unit hypercube in a temperature-controlled manner with the logistic function,

where σ(u)=(1+e−u)−1{\sigma(u)=(1+e^{-u})^{-1}} is the logistic function. Finally, transform the unit hypercube to a point in the simplex:

Here, βn\beta_{n} is the fraction of the remaining “stick” of probability mass assigned to xnx_{n}. This transformation is invertible, the Jacobian is lower-triangular, and the determinant of the Jacobian is easy to compute. Linderman et al. compute the density of xx implied by a Gaussian density on ψ\psi.

The temperature τ\tau controls how concentrated p(x)p(x) is at the vertices of the simplex, and with appropriate choices of parameters, in the limit τ→0{\tau\to 0} we can recover any categorical distribution (we will discuss this in detail in Section A.4. In the other limit, as τ→∞\tau\to\infty, the density concentrates on a point in the interior of the simplex determined by the parameters, and for intermediate values, the density is continuous on the simplex.

A.3 Rounding

Rounding transformations also have a natural analog for discrete variational inference. Let ene_{n} denote a one-hot vector with nn-th entry equal to one. Define the rounding operator,

In the case of a tie, let n∗n^{*} be the smallest index nn such that ψn>ψm\psi_{n}>\psi_{m} for all m<nm<n. Rounding effectively partitions the space into NN disjoint “Voronoi” cells,

By definition, round(ψ)=en∗{\mathsf{round}(\psi)=e_{n^{*}}} for all ψ∈Vn∗{\psi\in V_{n^{*}}}

We define a map that pulls points toward their rounded values,

For τ∈{\tau\in}, the map defined by (2) moves points strictly closer to their rounded values so that round(ψ)=round(x)\mathsf{round}(\psi)=\mathsf{round}(x).

Note that the Voronoi cells are intersections of halfspaces and, as such, are convex sets. Since xx is a convex combination of ψ\psi and en∗e_{n^{*}}, both of which belong to the convex set Vn∗V_{n^{*}}, xx must belong to Vn∗V_{n^{*}} as well. ∎

As long as ψ\psi is in the interior of its Voronoi cell, the round\mathsf{round} function is piecewise constant and the Jacobian is ∂ψ∂x=1τI{\tfrac{\partial\psi}{\partial x}=\tfrac{1}{\tau}I}, and its determinant is τ−N\tau^{-N}. Taken together, we have,

Compare this to the density of the rounded random variables for permutation inference.

A.4 Limit analysis for stick-breaking

We show that stick-breaking for discrete variational inference can converge to any categorical distribution in the zero-temperature limit.

These two facts, combined with the invertibility of the stick-breaking procedure, lead to the following proposition

In the zero-temperature limit, stick-breaking of logistic-normal random variables can realize any categorical distribution on xx.

There is a one-to-one correspondence between π∈ΔN{\pi\in\Delta_{N}} and ρ∈N−1{\rho\in^{N-1}}. Specifically,

Since these are recursively defined, we can substitute the definition of ρm\rho_{m} to obtain an expression for ρn\rho_{n} in terms of π\pi only. Thus, any desired categorical distribution π\pi implies a set of Bernoulli parameters ρ\rho. In the zero temperature limit, any desired ρn\rho_{n} can be obtained with appropriate choice of Gaussian mean μn\mu_{n} and variance νn2\nu_{n}^{2}. Together these imply that stick-breaking can realize any categorical distribution when τ→0{\tau\to 0}. ∎

A.5 Variational Autoencoders (VAE) with categorical latent variables

We considered the density estimation task on MNIST digits, as in Maddison et al. , Jang et al. , where observed digits are reconstructed from a latent discrete code. We used the continuous ELBO for training, and evaluated performance based on the marginal likelihood, estimated with the variational objective of the discretized model. We compared against the methods of Jang et al. , Maddison et al. and obtained the results in Table 2. While stick-breaking and rounding fare slightly worse than the Gumbel-softmax method, they are readily extensible to more complex discrete objects, as shown in the main paper.

Figure 5 shows MNIST reconstructions using Gumbel-Softmax, stick-breaking and rounding reparameterizations. In all the three cases reconstructions are reasonably accurate, and there is diversity in reconstructions.

Appendix B Variational permutation inference details

Here we discuss more of the subtleties of variational permutation inference and present the mathematical derivations in more detail.

Continuous relaxations require re-thinking the objective: the model log-probability is defined with discrete latent variables, but our relaxed posterior is a continuous density. As in Maddison et al. , we instead maximize a relaxed ELBO. We assume the functional form of the likelihood remains unchanged, and simply accepts continuous values instead of discrete. However, we need to specify a new continuous prior p(X)p(X) over the relaxed discrete latent variables, here, over relaxations of permutation matrices. It is important that the prior be sensible: ideally, the prior should penalize values of XX that are far from permutation matrices.

For our categorical experiment on MNIST we use a mixture of Gaussians around each vertex, p(x)=1N∑n=1NN(x ∣ ek,η2){p(x)=\tfrac{1}{N}\sum_{n=1}^{N}\mathcal{N}(x\,|\,e_{k},\eta^{2})}. This can be extended to permutations, where we use a mixture of Gaussians for each coordinate,

Although this prior puts significant mass around invalid points (e.g. (1,1,…,1){(1,1,\ldots,1)}), it penalizes XX that are far from BN\mathcal{B}_{N}.

B.2 Computing the ELBO

Here we show how to evaluate the ELBO. Note that the stick-breaking and rounding transformations are compositions of invertible functions, gτ=hτ∘f{g_{\tau}=h_{\tau}\circ f} with Ψ=f(z;θ){\Psi=f(z;\theta)} and X=hτ(Ψ){X=h_{\tau}(\Psi)}. In both cases, ff takes in a matrix of independent standard Gaussians (z)(z) and transforms it with the means and variances in θ\theta to output a matrix Ψ\Psi with entries ψmn∼N(μmn,νmn2){\psi_{mn}\sim\mathcal{N}(\mu_{mn},\nu^{2}_{mn})}. Stick-breaking and rounding differ in the temperature-controlled transformations hτ(Ψ)h_{\tau}(\Psi) they use to map Ψ\Psi toward the Birkhoff polytope.

To evaluate the ELBO, we must compute the density of qτ(X;θ)q_{\tau}(X;\theta). Let {J_{h_{\tau}}(u)=\frac{\partial h_{\tau}(U)}{\partial U}\big{|}_{U=u}} denote the Jacobian of a function hτh_{\tau} evaluated at value uu. By the change of variables theorem and properties of the determinant,

Now we appeal to the law of the unconscious statistician to compute the entropy of qτ(X;θ)q_{\tau}(X;\theta),

Since Ψ\Psi consists of independent Gaussians with variances νmn2\nu_{mn}^{2}, the entropy is simply,

We estimate the second term of equation (B.2) using Monte-Carlo samples. For both transformations, the Jacobian has a simple form.

As with the standard stick-breaking transformation to the simplex, our transformation to the Birkhoff polytope is feed-forward; i.e. to compute xmnx_{mn} we only need to know the values of β\beta up to and including the (m,n)(m,n)-th entry. Consequently, the Jacobian of the transformation is triangular, and its determinant is simply the product of its diagonal.

We derive an explicit form in two steps. With a slight abuse of notation, note that the Jacobian of hτ(Ψ)h_{\tau}(\Psi) is given by the chain rule,

Since both transformations are bijective, the determinant is,

the product of the individual determinants. The first determinant is,

The second transformation, from Ψ\Psi to BB, is an element-wise, temperature-controlled logistic transformation such that,

It is important to note that the transformation that maps B→XB\rightarrow X is only piecewise continuous: the function is not differentiable at the points where the bounds change; for example, when changing BB causes the active upper bound to switch from the row to the column constraint or vice versa. In practice, we find that our stochastic optimization algorithms still perform reasonably in the face of this discontinuity.

Jacobian of the rounding transformation.

The rounding transformation is given in matrix form in the main text, and we restate it here in coordinate-wise form for convenience,

This transformation is piecewise linear with jumps at the boundaries of the “Voronoi cells;” i.e., the points where round(X)\mathsf{round}(X) changes. The set of discontinuities has Lebesgue measure zero so the change of variables theorem still applies. Within each Voronoi cell, the rounding operation is constant, and the Jacobian is,

For the rounding transformation with given temperature, the Jacobian is constant.

Appendix C Experiment details

We used Tensorflow [Abadi et al., 2016] for the VAE experiments, slightly changing the code made available from Jang et al. . For experiments on synthetic matching and the C. elegans example we used Autograd [Maclaurin et al., 2015], explicitly avoiding propagating gradients through the non-differentiable round\mathsf{round} operation, which requires solving a matching problem.

We used ADAM [Kingma and Ba, 2014] with learning rate 0.1 for optimization. For rounding, the parameter vector VV defined in 3.2 was constrained to lie in the interval [0.1,0.5][0.1,0.5]. Also, for rounding, we used ten iterations of the Sinkhorn-Knopp algorithm, to obtain points in the Birkhoff polytope. For stick-breaking the variances ν\nu defined in 3.1 were constrained between 10−810^{-8} and 11. In either case, the temperature, along with maximum values for the noise variances were calibrated using a grid search.

In the C. elegans example we considered the symmetrized version of the adjacency matrix described in [Varshney et al., 2011]; i.e. we used A′=(A+A⊤)/2A^{\prime}=(A+A^{\top})/2, and the matrix WW was chosen antisymmetric, with entries sampled randomly with the sparsity pattern dictated by A′A^{\prime}. To avoid divergence, the matrix WW was then re-scaled by 1.1 times its spectral radius. This choice, although not essential, induced a reasonably well-behaved linear dynamical system, rich in non-damped oscillations. We used a time window of T=1000T=1000 time samples, and added spherical standard noise at each time. All results in Figure 4 are averages over five experiment simulations with different sampled matrices WW. For results in Figure 4b we considered either one or four worms (squares and circles, respectively), and for the x-axis we used the values ν∈{0.0075,0.01,0.02,0.04,0.05}\nu\in\{0.0075,0.01,0.02,0.04,0.05\}. We fixed the number of known neuron identities to 25 (randomly chosen). For results in Figure 4c we used four worms and considered two values for ν\nu; 0.1 (squares) and 0.05 (circles). Different x-axis values correspond to fixing 110, 83, 55 and 25 neuron identities.