Composable Effects for Flexible and Accelerated Probabilistic Programming in NumPyro
Du Phan, Neeraj Pradhan, Martin Jankowiak
Introduction
Many probabilistic programming languages (PPLs) are embedded as a DSL within a host language. The advantage of embedding is availability of host language infrastructure along with frameworks for automatic differentiation and hardware acceleration. Within the Python community, some examples of embedded PPLs include Pyro and ProbTorch based on PyTorch , TensorFlow Probability and Edward2 based on TensorFlow, and PyMC3 based on Theano. NumPyro is a package for probabilistic programming built atop JAX , which is a high-level tracing library for program transformations (e.g. automatic differentiation, vectorization and JIT compilation) of Python and NumPy functions. Thus NumPyro enables users to write probabilistic programs using familiar NumPy arrays and operations.
NumPyro is built around the same effect handling abstraction as Pyro. Effect handlers provide a way to inject effectful computation into primitive statements in a probabilistic program, e.g. recording the random choices made in an execution trace. In NumPyro these effects can be easily composed with the JAX tracer that operates at the level of NumPy operations and its own set of primitives for control flow. This allows us to expose a modeling language that is the same as in Pyro. Under the hood, inference algorithms can use effect handlers to inspect and modify program behavior and freely compose with JAX transformations to speed up critical subroutines via parallelization and JIT compilation. As an example (see Sec. 3.1), we implement an iterative version of the No-U-Turn Sampler that can leverage JAX’s jit transformation for end-to-end compilation and optimization by the XLA compiler .
Support for Pyro’s Modeling Interface
NumPyro retains the same language primitives and modeling and inference interface as in Pyro.A generic API for modeling and inference for dispatch to different Pyro backends can be found at https://github.com/pyro-ppl/pyro-api In particular, NumPyro supports sample and param statements that allow users to designate random variables and learnable parameters, respectively. It also has effect handlers like trace, replay and condition to provide nonstandard interpretations to these statements. Table 1 lists some commonly used effect handlers.
Effect handlers have emerged as a composable abstraction for program transformations in PPLs . We have found the effect handling abstraction to be particularly useful in designing a common interface to probabilistic programming, despite fundamental differences in the underlying backend. While PyTorch tries to accommodate much of Python’s dynamism and facilities for object-oriented programming with mutable objects (e.g. PyTorch optimizers update parameters in-place), JAX encourages a functional style of programming as required by the tracer. As an example, unlike PyTorch, JAX uses a functional pseudo-random number generator , which mandates passing an explicit random number generator key (PRNGKey) to distribution samplers. In practice, NumPyro inference algorithms take a single PRNGKey from the user that is split to generate new keys when passed to downstream functions. This does not result, however, in any change to probabilistic programs formulated in NumPyro, as this splitting mechanism is abstracted into a seed handler that operates on sample statements (see Table 1).
Leveraging JAX Transformations in Inference Subroutines
JAX is a Python library that provides a high-level tracer for implementing transformations of programs in Python and NumPy. Currently, three main transformations are available: i) automatic differentiation (grad); ii) JIT compilation (jit) to multiple backends using XLA; and iii) automatic vectorization (vmap). Frostig et al. note that ML workloads are composed of many pure-and-statically-composed (PSC) subroutines that are good candidates for acceleration. This is also true of the inference subroutines that lie at the core of NumPyro. While NumPyro’s frontend—i.e. its modeling and inference API—is close to Pyro, we have taken care to ensure that the core inference algorithms and utilities are purely functional so that we can make extensive use of JAX transformations like jit and vmap. This allows us to implement highly optimized and parallelizable subroutines. We provide two such examples: i) composing jit and grad to implement an iterative version of the NUTS sampler that can be end-to-end JIT compiled by JAX; and ii) composing vmap with effect handlers like trace or condition to implement vectorized subroutines.
The No-U-Turn Sampler (NUTS) is an extension of the Hamiltonian Monte Carlo (HMC) algorithm , which can efficiently sample from high-dimensional continuous probability distributions. NUTS adaptively sets the trajectory length parameter in HMC, which along with the adaptation of the step size and mass matrix parameters, ensures that HMC runs efficiently on a variety of models without extensive hand-tuning. This provides a highly attractive black-box inference algorithm for PPLs to have in their toolkit.
To JIT compile NUTS sampling, we need to jit transform a key component of the algorithm, namely the BuildTree subroutine that recursively builds an implicit balanced binary tree by running the LeapFrog integrator (Appendix A). While this can be written as a PSC subroutine, tracing it is hard for two reasons. First, the form of the LeapFrog integrator requires us to JIT through a gradient computation. JAX can handle this, since transformations like jit and grad are composable.Note, however, that for example PyTorch’s tracing JIT does not allow for this. Second—and more problematically—the complex control flow of the recursive formulation cannot be traced for JIT compilation in JAX.This obstacle was also noted in Tran et al. .
An alternative would be to JIT compile a single LeapFrog step or the potential energy function.This is the approach adopted by the NUTS implementation in Pyro. However, drawing a single sample involves many LeapFrog steps, and the overhead in terms of Python function dispatch calls is significant. In addition, this approach significantly reduces opportunities for operator fusion in XLA compilation. To overcome these limitations, we propose an iterative version of the NUTS algorithm that can be fully JIT compiled. In particular, this involves converting the BuildTree procedure into an iterative procedure, paving the way for a NUTS implementation that can take full advantage of XLA acceleration. As we demonstrate in benchmarking experiments in Sec. 4, the result is an algorithm that is much faster than existing implementations. More details on the algorithm are available in Appendix A.
2 Vectorizing Subroutines with vmap
Many utilities and subroutines for inference, e.g. model prediction, Monte Carlo estimation, or running MCMC chains, can be batched to make use of SIMD vectorization. In many frameworks this kind of batching requires laborious manual threading and/or significant cognitive overhead in managing explicit batch dimensions. JAX provides a vectorizing map (vmap) transformation that makes it easy to represent batched computations as mapping over function arguments along an outermost axis. This requires no changes to the underlying code but maintains the efficiency of manual batching.
Since JAX transformations are fully composable with Pyro’s effect handlers like seed, trace, and condition, and since the latter are implemented within the Python runtime and thus traceable, vmap becomes very powerful. As an example, Fig. 1 shows how we can use vmap to batch three common computations: i) sampling from the prior; ii) sampling from the posterior predictive distribution; iii) and computing log-likelihoods. Note that without vmap we would need to explicitly handle an additional batch dimension within the logistic_regression model and the utility functions in Fig. 1(b), which is particularly cumbersome for more involved models. As a final example, in Stochastic Variational Inference (SVI) , we optimize a loss function that is a Monte Carlo estimate of the Evidence Lower Bound (ELBO). This requires running the model as well as the inference network multiple times, all of which can be elegantly parallelized using vmap (see Appendix D).
Experiments
We compare the performance of NumPyro’s NUTS implementation with that of other frameworks (Stan and Pyro) in both the small and large data regimes. Recall that NumPyro’s NUTS implementation is end-to-end JIT compiled, while in Pyro only the potential energy computation is compiled. We use three benchmark models: i) a Hidden Markov Model (HMM) on a small synthetic dataset; ii) logistic regression on the Forest CoverType dataset ; and iii) a sparse kernel interaction model (SKIM) on synthetic datasets with varying dimensionalities. Refer to Appendix C for details on the benchmarking experiments.
Since we use a small dataset for this experiment, we expect poor performance on the GPU; consequently we limit ourselves to a CPU-only comparison. Note that, although the dataset is small, the potential energy computation involves a loop that can be expensive to differentiate through. From Table 2(a), we see that for the HMM, NumPyro is around X faster than Pyro and X faster than Stan. The iterative procedure in Algorithm 2 introduces insignificant overhead, and the end-to-end compilation allows XLA to output highly optimized code.
Logistic Regression
For this dataset, which contains more than half a million datapoints, GPU acceleration significantly outperforms the CPU, as expected. The time spent in computing gradients in the LeapFrog integrator exceeds the time spent building the tree or computing the terminating condition. Since the bottleneck primarily lies in large tensor operations, we expect the difference between the various GPU implementations to be narrower. Nevertheless, on this problem NumPyro is about 2X faster than Pyro.
Sparse Kernel Interaction Model
SKIM is Bayesian model for sparse regression that can be used to discover pairwise interactions in high dimensional data. Since the sparsity-inducing prior introduces a latent variable for each of the input dimensions, this represents a difficult inference problem when is large. Fig. 2(b) shows how the time per effective sample using NUTS scales with the dimensionality of the dataset for Stan and NumPyro. We observe that NumPyro has consistently lower overhead as compared to Stan. NumPyro offers the flexibility to run inference in single or double precision as well as on different backends such as CPU, GPU, or TPU. While inference with double precision yields a higher effective sample size on average, it is not enough to compensate for the higher time taken to run inference, and hence the time per effective sample is lower for single precision. This model also particularly benefits from GPU acceleration, resulting in a major improvement in execution time when compared to the CPU backend.
Summary
We describe NumPyro, a package for probabilistic programming using Python and NumPy that uses JAX transformations under the hood for hardware acceleration, automatic differentiation, and vectorization. NumPyro has a functional core, where inference subroutines are pure-and-statically composed functions that can be traced by JAX for parallelization and JIT compilation. These subroutines also make use of effect handlers to inspect and transform probabilistic programs. Effect handlers operate on core language primitives within the Python runtime, are transparent to the JAX tracer, and are, therefore, fully composable with JAX’s transformations. This composability allows us to offer the same modeling language as Pyro, and at the same time leverage JAX tranformations to parallelize and JIT compile inference subroutines for significant speed ups. In particular we show that the judicious application of these program transformations allows us to implement an iterative version of the NUTS algorithm that offers strong performance on both the CPU (for small models) and the GPU (for larger models).
Acknowledgments
We would like to thank Noah Goodman for feedback, and the JAX development team—in particular Matthew Johnson and Peter Hawkins—for their invaluable help with many JAX issues and feature requests.
References
Appendix A Iterative NUTS - Algorithm Details
Computing a trajectory in NUTS involves a doubling procedure where at each iteration we run the LeapFrog integrator for twice the number of steps taken in the previous iteration, with the direction (forward or reverse) chosen randomly. This has the effect of building an implicit balanced binary tree. The doubling process is terminated when a subtrajectory from the leftmost to the rightmost node of any balanced subtree begins to double back on itself.
Existing NUTS implementations use a recursive tree building formulation, the BuildTree subroutine, to double the trajectory length (see Hoffman and Gelman [14, Algorithm 6]). A simplified version of this subroutine is presented in Algorithm 1. To build a tree at depth , it builds two subtrees at depth and combines them. This recursive procedure also takes care to ensure that memory usage scales as (where ) rather than by only storing data per subtree. This is important because storing all momentum-position pairs might be prohibitive for large models. The key to JIT compiling NUTS sampling is the ability to convert the recursive BuildTree subroutine into an iterative implementation.
The iterative version of BuildTree takes an initial node (position-momentum pair) and a tree depth argument, and runs the LeapFrog integrator for steps. We will be using -based indexing in the following discussion.
For node , we need to check the U-Turn condition with respect to the leftmost nodes of any binary subtree for which is the rightmost node. Let us denote the indices for these candidate nodes by , and the binary representation of by . Indices in have the same binary representation as except that trailing contiguous s in are progressively masked by ; e.g. for , , and the set of candidate nodes for checking the U-Turn condition are indexed by . Note that this implies that we only need to check the U-Turn condition at odd-numbered nodes against a subset of previous even-numbered nodes.
This allows us to iteratively build the binary tree by running the LeapFrog integrator for steps, and terminating early if the U-Turn condition stated above is satisfied. However, in a naive implementation we would still need memory, since we would need to store the position-momentum pairs at each step of the integrator, which would be an unacceptable regression from the memory requirement of the recursive algorithm.
Memory Efficiency
At an odd step , the data for the candidate nodes indexed by must be present in because the masking procedure ensures that these candidates are the largest even nodes less than for their corresponding bit counts. Figure 3 illustrates the iterative procedure at step , and full details of the algorithm are provided in Algorithm 2.
Appendix B Code for Vectorized Sampling - Logistic Regression
Appendix C Experimental Details
All experiments are conducted on a system using an AMD Ryzen Threadripper 1920X processor and an NVIDIA GeForce RTX 2080 Ti graphics card. Framework versions: PyStan , Pyro (with PyTorch ), NumPyro (with JAX and jaxlib ). For each experiment, we conduct runs with different random seeds and report the average across runs. Code used to benchmark all experiments can be found on the benchmarks branch of the NumPyro GitHub repository.https://github.com/pyro-ppl/numpyro/tree/benchmarks-20191222/benchmarks
To test performance in the small dataset regime, we use an HMM. Following [22, Section 2.6], we construct a semi-supervised HMM model with 3-dimensional latent states and 10-dimensional observations. Using fixed transition and emission matrices, we sample data points and treat the first latent states as observed. For benchmarking we take warmup steps and draw NUTS samples for both Stan and NumPyro. Because Pyro’s NUTS implementation is extremely slow on this problem, we fix the step size to and only draw samples for each run.
Logistic Regression
We consider a logistic regression model on the Forest CoverType dataset, which has datapoints and features. The prior on the weights is a unit normal distribution. Following , we normalize all features and transform the multi-class problem into a binary class problem by merging all the classes except for the most frequent one. For benchmarking we fix the step size in all frameworks to and draw samples. That step size value was obtained by running warmup adaptation steps in NumPyro.
Sparse Kernel Interaction Model (SKIM)
For each value of dimensionality , we produce an artificial dataset with data points that contains 3 randomly selected pairwise interaction terms amongst the covariates. For each of the frameworks, we adapt the step size and mass matrix using 1000 warmup adaptation steps and then compute the time per effective sample based on the next 1000 drawn samples averaged over 5 random runs.