Delayed Sampling and Automatic Rao-Blackwellization of Probabilistic Programs

Lawrence M. Murray, Daniel Lundén, Jan Kudlicka, David Broman, Thomas B. Schön

INTRODUCTION

Probabilistic programs extend graphical models with support for stochastic branches, in the form of conditionals, loops, and recursion. Because they are highly expressive, they pose a challenge in the design of appropriate inference algorithms. This work focuses on Sequential Monte Carlo (SMC) inference algorithms , extending an arc of research that includes probabilistic programming languages (PPLs) such as Venture , Anglican , Probabilistic C , WebPPL , Figaro , and Turing , as well as similarly-motivated software such as LibBi and BiiPS .

The simplest SMC method, the bootstrap particle filter , requires only simulation—not pointwise evaluation—of the prior distribution. While widely applicable, it may be suboptimal with respect to Monte Carlo variance in situations where, in fact, pointwise evaluation is possible, so that other options are viable. One way of reducing Monte Carlo variance is to exploit analytical relationships between random variables, such as conjugate priors and affine transformations. Within SMC, this translates to improvements such as the locally-optimal proposal, variable elimination, and Rao–Blackwellization (see for an overview). The present work seeks to automate such improvements for the user of a PPL.

Typically, a probabilistic program must be run in order to discover the relationships between random variables. Because of stochastic branches, different runs may discover different relationships, or even different random variables. While an equivalent graphical model might be constructed for any single run, it would constitute only partial observation. It may take many runs to observe the full model, if this is possible in finite time at all. We therefore seek a runtime mechanism for the solution of analytically-tractable substructure, rather than a compile-time mechanism of static analysis.

A general-purpose programming language can be augmented with some additional constructs, called checkpoints, to produce a PPL (see e.g. ). Two checkpoints are usual, denoted \funcsample\func{sample} and \funcobserve\func{observe}. The first suggests that a value for a random variable needs to be sampled, the second that a value for a random variable is given and needs to be conditioned upon. At these checkpoints, random behavior may occur in the otherwise-deterministic execution of the program, and intervention may be required by an inference algorithm to produce a correct result.

The simplest inference algorithm instantiates a random variable when first encountered at a \funcsample\func{sample} checkpoint, and updates a weight with the likelihood of a given value at an \funcobserve\func{observe} checkpoint. This produces samples from the prior distribution, weighted by their likelihood under the observations. It corresponds to importance sampling with the posterior as the target and the prior as the proposal. A more sophisticated inference algorithm runs multiple instances of the program simultaneously, pausing after each \funcobserve\func{observe} checkpoint to resample amongst executions. This corresponds to the bootstrap particle filter (see e.g. ).

These are forward methods, in the sense that checkpoints are executed in the order encountered, and sampling is myopic of future observations. The present work introduces a mechanism to change the order in which checkpoints are executed so that sampling can be informed by future observations, exploiting analytical relationships between random variables. This facilitates more sophisticated forward-backward methods, in the sense that information from future observations can be propagated backward through the program.

We refer to this new mechanism as delayed sampling. When a \funcsample\func{sample} checkpoint is reached, its execution is delayed. Instead, a new node representing the random variable is inserted into a graph that is maintained alongside the running program. This graph resembles a directed graphical model of those random variables encountered so far that are involved in analytically-tractable relationships. Each node of the graph is marginalized and conditioned by analytical means for as long as possible until, eventually, it must be instantiated for the program to continue execution. This occurs when the random variable is passed as an argument to a function for which no analytical overload is provided. It is at this last possible moment that sampling is executed and the random variable instantiated.

Operations on the graph are forward-backward. The forward pass is a filter, marginalizing each latent variable over its parents and conditioning on observations, in all cases analytically. The backward pass produces a joint sample. This has some similarity to belief propagation , but the backward passes differ: belief propagation typically obtains the marginal posterior distribution of each variable, not a joint sample. Furthermore, in delayed sampling the graph evolves dynamically as the program executes, and at any time represents only a fraction of the full model. This means that some heuristic decisions must be made without complete knowledge of the model structure.

For SMC, delayed sampling yields locally-optimal proposals, variable elimination, and Rao–Blackwellization, with some limitations, to be detailed later. At worst, it provides no benefit. There is little intrusion of the inference algorithm into modeling code, and possibly no intrusion with appropriate language support. This is important, as we consider the user experience and ergonomics of a PPL to be of primary importance.

Related work has considered analytical solutions to probabilistic programs. Where a full analytical solution is possible, it can be achieved via symbolic manipulations in Hakaru . Where not, partial solutions using compile-time program transformations are considered in to improve the acceptance rate of Metropolis–Hastings algorithms. This compile-time approach requires careful treatment of stochastic branches, and even then it may not be possible to propagate analytical solutions through them. Delayed sampling instead operates dynamically, at runtime. It handles stochastic branches without problems, but may introduce some additional execution overhead.

The paper is organized as follows. Section 2 introduces the delayed sampling mechanism. Section 3 provides a set of pedagogical examples and two empirical case studies. Section 4 discusses some limitations and future work. Supplementary material includes further details of the case studies and implementations.

METHODS

As a probabilistic program runs, its memory state evolves dynamically and stochastically over time, and can be considered a stochastic process. Let t=1,2,…t=1,2,\ldots index a sequence of checkpoints. These checkpoints may differ across program runs (this is one of the challenges of inference for probabilistic programs, see e.g. ). In contrast to the two-checkpoint \funcsample\func{sample}-\funcobserve\func{observe} formulation, we define three checkpoint types:

\funcassume(X,p(⋅))\func{assume}(X,p(\cdot)) to initialize a random variable XX with prior distribution p(⋅)p(\cdot),

\funcobserve(x,p(⋅))\func{observe}(x,p(\cdot)) to condition on a random variable XX with likelihood p(⋅)p(\cdot) having some value xx,

\funcvalue(X)\func{value}(X) to realize a value for a random variable XX previously encountered at an \funcassume\func{assume} checkpoint.

We use the statistics convention that an uppercase character (e.g. XX) denotes a random variable, while the corresponding lowercase character (e.g. xx) denotes an instantiation of it.

An \funcassume\func{assume} checkpoint does not result in a random variable being sampled: its sampling is delayed until later. A \funcvalue\func{value} checkpoint occurs the first time that a random variable, previously encountered by an \funcassume\func{assume}, is used in such a way that its value is required. At this point it cannot be delayed any longer, and is sampled.

The program is a sequence of functions ftf_{t} that each maps a starting state Xt−1=xt−1X_{t-1}=x_{t-1} and random input Ut=utU_{t}=u_{t} to an end state Xt=xtX_{t}=x_{t}, so that xt=ft(xt−1,ut)x_{t}=f_{t}(x_{t-1},u_{t}). Note that ftf_{t} is a deterministic function given its arguments. It is not permitted that ftf_{t} has any intrinsic randomness, only the extrinsic randomness provided by UtU_{t}.

We are motivated by variance reduction in Monte Carlo estimators. Consider some functional φ(X)\varphi(X) of interest. We wish to compute expectations of the form:

A classic aim is to reduce mean squared error:

The Rao–Blackwellized estimator does not instantiate XMX_{M}, but rather marginalizes it out:

2 Delayed sampling

Delayed sampling uses analytical relationships to reorder the execution of checkpoints and reduce variance. Each \funcobserve\func{observe} is executed as early as possible, and the sampling associated with \funcassume\func{assume} is delayed for as long as possible, to be informed by observations in between.

I⊆VI\subseteq V be the set of nodes in an initialized state,

M⊆VM\subseteq V be the set of nodes in a marginalized state,

R⊆VR\subseteq V be the set of nodes in a realized state.

At some checkpoint, the program would usually have instantiated all variables in VV with a simulated or observed value, whereas under delayed sampling only those in RR are instantiated, while those in I∪MI\cup M are delayed.

We will restrict the graph GG to be a forest of zero or more disjoint trees, such that each node has at most one parent. This condition is easily ensured by construction: the implementation makes anything else impossible, i.e. only relationships between pairs of random variables are coded. There are some interesting relationships that cannot be represented as trees, such as a normal distribution with conjugate prior over both mean and variance, or multivariate normal distributions. We deal with these as special cases, collecting multiple nodes into single supernodes and implementing relationships between pairs of supernodes, much like the structure achieved by the junction tree algorithm .

The following invariants are preserved at all times:

These imply that the nodes of MM form marginalized paths: one in each of the disjoint trees of GG, from the root node to a node (possibly itself) in the same tree. We will refer to the unique such path in each tree as its MM-path. The node at the start of the MM-path is a root node, while the node at the end is referred to as a terminal node. Terminal nodes have a special place in the algorithms below, and are denoted by the set TT.

where qvq_{v} equals the prior for nodes in II, some updated distribution for nodes in MM, and all nodes in RR are instantiated. The distribution suggests why terminals (in the set TT) are important: they are the nodes informed by all instantiated random variables up to the current point in the program, and can be immediately instantiated themselves. Other nodes in MM await information to be propagated backward from their forward set before they, too, can be instantiated.

When the program reaches a checkpoint, it triggers operations on the graph (details follow):

For \funcassume(Xv,p(⋅)\func{assume}(X_{v},p(\cdot)), call \procInitialize(v,p(⋅))\proc{Initialize}(v,p(\cdot)), which inserts a new node vv into the graph.

For \funcobserve(xv,p(⋅))\func{observe}(x_{v},p(\cdot)), call \procInitialize(v,p(⋅))\proc{Initialize}(v,p(\cdot)), then \procGraft(v)\proc{Graft}(v), which turns vv into a terminal node, then \procObserve(v)\proc{Observe}(v), which assigns the observed value to vv and updates its parent by conditioning.

For \funcvalue(Xv)\func{value}(X_{v}), call \procGraft(v)\proc{Graft}(v), then \procSample(v)\proc{Sample}(v), which samples a value for vv.

Figure 1 provides pseudocode for all operations; Figure 2 illustrates their combination. Operations are of two types: local and recursive. Local operations modify a single node and possibly its parent:

\procMarginalize(v)\proc{Marginalize}(v), where vv is the child of a terminal node, moves vv from II to MM and updates its distribution by marginalizing over its parent.

\procSample(v)\proc{Sample}(v) or \procObserve(v)\proc{Observe}(v), where vv is a terminal node, assigns a value to the associated random variable by either sampling or observing, moves vv from MM to RR, and updates the distribution of its parent node by conditioning. Both \procSample(v)\proc{Sample}(v) and \procObserve(v)\proc{Observe}(v) use an auxiliary function \procRealize(v)\proc{Realize}(v) for their common operations.

As shown in the pseudocode, these local operations have strict preconditions that limit their use to only a subset of the nodes of the graph, e.g. only terminal nodes may be sampled or observed. As long as these preconditions are satisfied, the invariants (1) and (2) are maintained, and the graph GG encodes the representation (3). This is straightforward to check.

The recursive operations realign the MM-path to establish the preconditions for any given node, so that local operations may be applied to it. These have side effects, in that other nodes may be modified to achieve the realignment. The key recursive operation is \procGraft\proc{Graft}, which combines local operations to extend the MM-path to a given node, making it a terminal node. Internally, \procGraft\proc{Graft} may call another recursive operation, \procPrune\proc{Prune}, to shorten the existing MM-path by realizing one or more variables.

EXAMPLES

We have implemented delayed sampling in Anglican (see also ) and a new PPL called Birch. Details are given in Appendices C and D.

Table 1 provides pedagogical examples using a Birch-like syntax, showing the sequence of checkpoints and graph operations triggered as some simple programs execute. They show how delayed sampling behaves through programming structures such as conditionals and loops, including stochastic branches.

In addition, we provide two case studies where delayed sampling improves inference, firstly a linear-nonlinear state-space model with simulated data, secondly a vector-borne disease model with real data from an outbreak of dengue virus in Micronesia. We use a simple random-weight or pseudo-marginal-type importance sampling algorithm for both of these examples:

Run SMC on the probabilistic program with delayed sampling enabled, producing NN number of samples x1,…,xNx^{1},\ldots,x^{N} with associated weights w1,…,wNw^{1},\ldots,w^{N} and a marginal likelihood estimate Z^\hat{Z}.

Draw a∈{1,…,N}a\in\{1,\ldots,N\} from the categorical distribution defined by P(a)=wa/∑n=1NwnP(a)=w^{a}/\sum_{n=1}^{N}w^{n}.

This produces one sample with associated weight, but may be repeated as many times as necessary—in parallel, even—to produce an importance sample as large as desired. The success of the approach depends on the variance of Z^\hat{Z}. This variance can be reduced by marginalizing out one or more variables (recall Section 2.1). This is what delayed sampling achieves, and so we compare the variance of Z^\hat{Z} with delayed sampling enabled and disabled. When disabled, the SMC algorithm is simply a bootstrap particle filter. When enabled, it yields a Rao–Blackwellized particle filter. Where parameters are involved (as in the second case study), the diversity of parameter values depletes through the resampling step of SMC. This has motivated more sophisticated methods for parameter estimation such as particle Markov chain Monte Carlo methods , also applied to probabilistic programs . Particle Gibbs is an obvious candidate here. We find, however, that the reduction in variance afforded by marginalizing out one or more variables with delayed sampling is sufficient to enable the above importance sampling algorithm for the two case studies here.

The first example is that of a mixed linear-nonlinear state-space model. For this model, delayed sampling yields a particle filter with locally-optimal proposal and Rao–Blackwellization.

The model is given by and repeated in Appendix A. It consists of both nonlinear and linear-Gaussian state variables, as well as nonlinear and linear-Gaussian observations. Parameters are fixed. Ideally, the linear-Gaussian substructure is solved analytically (e.g. using a Kalman filter), leaving only the nonlinear substructure to sample (e.g. using a particle filter). The Rao–Blackwellized particle filter, also known as the marginalized particle filter, was designed to achieve precisely this .

Delayed sampling automatically yields this method for this model, as long as analytical relationships between multivariate Gaussian distributions are encoded. In Birch these are implemented as supernodes: single nodes in the graph that contain multiple random variables. While the relationships between individual variables in a multivariate Gaussian have, in general, directed acyclic graph structure, their implementation as supernodes maintains the required tree structure.

The model is run for 100 time steps to simulate data. It is run again with SMC, conditioning on this data. For various numbers of particles, it is run 100 times to estimate Z^\hat{Z}, with delayed sampling enabled and disabled. Figure 3 (left) plots the distribution of these estimates. Clearly, with delayed sampling enabled, fewer particles are needed to achieve comparable variance in the log-likelihood estimate.

2 Vector-borne disease model

The second example is an epidemiological case study of an outbreak of dengue virus: a mosquito-borne tropical disease with an estimated 50-100 million cases and 10000 deaths worldwide each year . It is based on the study in , which jointly models two outbreaks of dengue virus and one of Zika virus in two separate locations (and populations) in Micronesia. Presented here is a simpler study limited to one of those outbreaks, specifically that of dengue on the Yap Main Islands in 2011. The data used consists of 172 observations of reported cases, on a daily basis during the main outbreak, and on a weekly basis before and after.

The model consists of two components, representing the human and mosquito populations, coupled via cross-infection. Each population is further divided into subpopulations of susceptible, exposed, infectious and recovered individuals. At each time step a binomial transfer occurs between subpopulations, parameterized with conjugate beta priors. Details are in Appendix B.

The task is both parameter and state estimation. For this model, delayed sampling produces a Rao–Blackwellized particle filter where parameters, rather than state variables, are marginalized out. While the state variables are sampled immediately, the parameters are maintained in a marginalized state, conditioned on the samples of these state variables. This is a consequence of conjugacy between the beta priors on parameters and the binomial likelihoods of the state variables (as pseudo-observations).

For various numbers of particles, SMC is run 100 times to estimate Z^\hat{Z}, with delayed sampling enabled and disabled. Figure 3 (right) plots the distribution of these estimates. Clearly, with delayed sampling enabled, fewer particles are needed to achieve comparable variance in the log-likelihood estimate. Some posterior results are given in Appendix B.

DISCUSSION AND CONCLUSION

Table 1 demonstrates how delayed sampling operates through typical program structures such as conditionals and loops, including stochastic branches as encountered in probabilistic programs. Figure 3 demonstrates the potential gains. These are particularly encouraging given that the mechanism is mostly automatic.

Some limitations are worth noting. The graph of analytically-tractable relationships must be a forest of disjoint trees. It is unclear whether this is a significant limitation in practice, but support for more general structures may be desirable. It is worth emphasizing that this relates to the structure of analytically-tractable relationships and the ability of the mechanism to utilize them, not to the structure of the model as a whole. At present, for more general structures, some opportunities for variance reduction are missed. One remedy is to encode supernodes, as for the multivariate Gaussian distributions in Section 3.1.

While delayed sampling may reduce the number of samples required for comparable variance, it does require additional computation per sample. For univariate relationships (e.g. beta-binomial, gamma-Poisson), this overhead is constant and—we conjecture—likely worthwhile for any fixed computational budget. For multivariate relationships the overhead is more complex and may not be worthwhile (e.g. multivariate Gaussian conjugacies require matrix inversions that are O(N3)\mathcal{O}(N^{3}) in the number of dimensions). A thorough empirical comparison is beyond the scope of this article.

Finally, while the focus of this work is SMC, delayed sampling may be useful in other contexts. With undirected graphical models, for example, delayed sampling may produce a collapsed Gibbs sampler. This is left to future work.

Acknowledgements

This research was financially supported by the Swedish Foundation for Strategic Research (SSF) via the project ASSEMBLE. Jan Kudlicka was supported by the Swedish Research Council grant 2013-4853.

Supplementary material

Appendix A details the linear-nonlinear state-space model, and Appendix B the vector-borne disease model. Appendix C details the Anglican implementation, and Appendix D the Birch implementation. Code is included for the pedagogical examples in both Anglican and Birch, and for the empirical case studies, along with data sets, in Birch only.

References

Appendix A Details of the linear-nonlinear state-space model

The full model is described in . The state model contains both nonlinear (XtnX_{t}^{n}) and linear-Gaussian (XtlX_{t}^{l}) state variables, and is given by:

The observation model contains both nonlinear (YtnY_{t}^{n}) and linear-Gaussian (YtlY_{t}^{l}) observations, and is given by:

Appendix B Details of the vector-borne disease model

The process model is a discrete-time and discrete-state stochastic model based on the continuous-time and continuous-state deterministic mean-field approximation used in . It consists of two SEIR (susceptible, exposed, infectious, recovered) compartmental models, one for the human population, the other for the mosquito population, coupled via cross-infection terms. Each component consists of state variables giving population counts in each of the four compartments: ss (susceptible), ee (exposed), ii (infectious), and rr (recovered), along with a total population nn that maintains the identity n=s+e+i+rn=s+e+i+r, and parameters ν\nu (birth probability), μ\mu (death probability), λ\lambda (transmission probability), δ\delta (infectious probability), and γ\gamma (recovery probability). A susceptible human may become infected when bitten by an infectious mosquito, while a susceptible mosquito may become infected when biting an infectious human.

For the setting of Yap Main Islands in 2011, the following initial conditions are prescribed:

B.2 Transition model

The model transitions in two steps. The first step is an exchange between compartments that preserves total population. Denoting with primes the intermediate state after this first step, we have:

with the newly exposed, infectious, and recovered populations distributed as:

for parameters λh\lambda^{h}, δh\delta^{h}, γh\gamma^{h}, λm\lambda^{m}, δm\delta^{m}, γm\gamma^{m}. The τth\tau_{t}^{h} gives the number of susceptible humans bitten by at least one infectious mosquito, and τtm\tau_{t}^{m} the number of susceptible mosquitos that bite at least one infectious human:

The second step accounts for births and deaths:

with parameters νh\nu^{h} and νm\nu^{m}, and deaths as

B.3 Observation model

Observations are of the number of new infectious cases reported at health centers, aggregated over the time since the last such observation (this is daily during the peak time of the outbreak and weekly either side). For times t∈{1,…,T}t\in\left\{1,\ldots,T\right\} where observations are available, the observation model is given by

where ltl_{t} (lag) indicates the number of days since the last observation. Significant under-reporting of cases is expected, reflected in the parameter ρ\rho.

B.4 Parameter model

The following fixed values and priors are assigned to parameters, translating prior knowledge on rates in to prior knowledge on probabilities here:

Birth and death in the human population are assumed to be of minimal impact over the course of the outbreak, and so their rates are fixed to zero. The expected lifespan of a mosquito is one week, with birth and death rates fixed accordingly. Mosquitos do not recover before death.

Finally, the prior over the reporting probability is

B.5 Inference results

Inference is performed by drawing 10000 weighted samples, each time running SMC with 8192 particles. The effective sample size of these 10000 weighted samples is computed to be 2260. Some results are shown in Figure 4.

Appendix C Anglican implementation

Anglican is a functional probabilistic programming language integrated with Clojure. Clojure, in turn, is a Lisp dialect which compiles to Java virtual machine bytecode, enabling reuse of the Java infrastructure. The Anglican compiler is built with Clojure macros, and compiles Anglican programs into continuation-passing-style Clojure code. This transformation enables inference algorithms to affect the control flow and record information at checkpoints. Manipulations are performed both on the continuations themselves and on the state, which is passed along as an argument in each continuation call.

For simplicity, delayed sampling is implemented entirely on top of the existing Anglican language, leaving the original language constructs and functionality untouched. A set of new keywords and functions are added for usage of delayed sampling: ds-, ds-value, and ds-observe. The ds-value and ds-observe functions loosely correspond to the \procSample\proc{Sample} and \procObserve\proc{Observe} operations in Section 2.2, but ds-value also includes functionality for retrieving values for already-sampled nodes. The set of ds- functions correspond to the \procInitialize\proc{Initialize} operations in Section 2.2, for various probability distributions, e.g. ds-normal. The delayed sampling graph is conveniently encoded in the already existing Anglican state.

As an example, consider the following line of code:

This binds x to a graph node which is normally distributed with mean mean and standard deviation sd. To subsequently introduce another normally distributed graph node with the node x as mean, one can write

passing the previous graph node x as a parameter. This will initialize a conjugate prior relationship between them. If y is then observed, x will be conditioned on the observed value of y.

Appendix D Birch implementation

Birch is a compiled, imperative, object-oriented, generic, and probabilistic programming language. The latter is its primary research concern. The Birch compiler uses C++ as a target language.

Delayed sampling has been implemented using the Birch type system. Special types are used when declaring variables to make them eligible for delayed sampling. For example, a variable that might ordinarily be declared to be of type Real may be declared to be of type Random to make it eligible for delayed sampling. The generic class Random implements the behavior required for delayed sampling, and is specialized into classes that encode distributions (e.g. Gaussian), then further into classes that encode distributions with analytical relationships to others (e.g. GaussianWithGaussianMean). The graph required for delayed sampling is formed implicitly through objects of these classes and their member attributes.

Birch supports implicit type conversion, compiling directly to the same feature in C++. These implicit conversions are used to automatically trigger the \funcvalue\func{value} checkpoint, and are resolved at compile time. For example, a Random object may be passed to a function that requires a Real argument. An implicit conversion is used to trigger a \funcvalue\func{value} checkpoint, realizing a value of type Real from the object of type Random. In this way, the programmer need not explicitly indicate \funcvalue\func{value} checkpoints.