Probabilistic Programming with Programmable Variational Inference
arXiv:2406.15742 · doi:10.1145/3656463
Abstract
Compared to the wide array of advanced Monte Carlo methods supported by modern probabilistic programming languages (PPLs), PPL support for variational inference (VI) is less developed: users are typically limited to a predefined selection of variational objectives and gradient estimators, which are implemented monolithically (and without formal correctness arguments) in PPL backends. In this paper, we propose a more modular approach to supporting variational inference in PPLs, based on compositional program transformation. In our approach, variational objectives are expressed as programs, that may employ first-class constructs for computing densities of and expected values under user-defined models and variational families. We then transform these programs systematically into unbiased gradient estimators for optimizing the objectives they define. Our design enables modular reasoning about many interacting concerns, including automatic differentiation, density accumulation, tracing, and the application of unbiased gradient estimation strategies. Additionally, relative to existing support for VI in PPLs, our design increases expressiveness along three axes: (1) it supports an open-ended set of user-defined variational objectives, rather than a fixed menu of options; (2) it supports a combinatorial space of gradient estimation strategies, many not automated by today's PPLs; and (3) it supports a broader class of models and variational families, because it supports constructs for approximate marginalization and normalization (previously introduced only for Monte Carlo inference). We implement our approach in an extension to the Gen probabilistic programming system (genjax.vi, implemented in JAX), and evaluate on several deep generative modeling tasks, showing minimal performance overhead vs. hand-coded implementations and performance competitive with well-established open-source PPLs.
References in corpus (28)
- Variational Inference: A Review for Statisticians
- Semi-Supervised Learning with Deep Generative Models
- NVAE: A Deep Hierarchical Variational Autoencoder
- Variational Autoencoder for Deep Learning of Images, Labels and Captions
- Markov Chain Monte Carlo and Variational Inference: Bridging the Gap
- Variational Diffusion Models
- Attend, Infer, Repeat: Fast Scene Understanding with Generative Models
- Automatic Differentiation Variational Inference
- Gradient Estimation Using Stochastic Computation Graphs
- Monte Carlo Gradient Estimation in Machine Learning
- Denotational validation of higher-order Bayesian inference
- Simple, Distributed, and Accelerated Probabilistic Programming
- ADEV: Sound Automatic Differentiation of Expected Values of Probabilistic Programs
- Smoothing Methods for Automatic Differentiation Across Conditional Branches
- Automatic Differentiation of Programs with Discrete Randomness
- Functional Tensors for Probabilistic Programming
- DiCE: The Infinitely Differentiable Monte-Carlo Estimator
- Reparameterization Gradient for Non-differentiable Models
- Storchastic: A Framework for General Stochastic Automatic Differentiation
- : Computable Semantics for Differentiable Programming with Higher-Order Functions and Datatypes
- Tensor Variable Elimination for Plated Factor Graphs
- Correctness of Sequential Monte Carlo Inference for Probabilistic Programming Languages
- Learning Proposals for Probabilistic Programs with Inference Combinators
- AIDE: An algorithm for measuring the accuracy of probabilistic inference algorithms
- An Easy to Interpret Diagnostic for Approximate Inference: Symmetric Divergence Over Simulations
- Importance Weighted Hierarchical Variational Inference
- Recursive Monte Carlo and Variational Inference with Auxiliary Variables
- You Only Linearize Once: Tangents Transpose to Gradients