flowMC: Normalizing flow enhanced sampling package for probabilistic inference in JAX
Kaze W. K. Wong, Marylou Gabrié, Daniel T. Foreman-Mackey · The Journal of Open Source Software · 2023
Across scientific fields, the Bayesian framework is used to account for uncertainties in inference (Ellison, 2004;Lancaster, 2004;Von Toussaint, 2011).However, for models with more than a few parameters, exact inference is intractable.A frequently used strategy is to approximately sample the posterior distribution on a model's parameters with a Markov chain Monte Carlo (MCMC) method.Yet conventional MCMC methods relying on local updates can take a prohibitive time to converge when posterior distributions have complex geometries (see e.g., Rubinstein & Kroese (2017)).flowMC is a Python library implementing accelerated MCMC leveraging deep generative modelling as proposed by Gabrié et al. (2022), built on top of the machine learning libraries JAX (Bradbury et al., 2018) and Flax (Heek et al., 2020).At its core, flowMC uses a combination of Metropolis-Hastings Markov kernels using local and global proposed moves.While multiple chains are run using local-update Markov kernels to generate approximate samples over the region of interest in the target parameter space, these samples are used to train a normalizing flow (NF) model to approximate the samples' density.The NF is then used in an independent Metropolis-Hastings kernel to propose global jumps across the parameter space.The flowMC sampler can handle non-trivial geometry, such as multimodal distributions and distributions with local correlations.