about
BlackJAX: Composable Bayesian Inference in Jax (arxiv.org)
3 points by sebg on Feb 20, 2024 | hide | past | pdf | discuss on HN

In plain words: A toolkit of Bayesian inference building blocks—samplers and approximation methods—written as mix-and-match functions that compile to run fast on ordinary, graphics, or Google chips. Unlike fixed one-piece routines, it lets users snap parts together to build methods, with one-call versions for beginners.

Abstract · BlackJAX: Composable Bayesian inference in JAX

BlackJAX is a library implementing sampling and variational inference algorithms commonly used in Bayesian computation. It is designed for ease of use, speed, and modularity by taking a functional approach to the algorithms' implementation. BlackJAX is written in Python, using JAX to compile and run NumpPy-like samplers and variational methods on CPUs, GPUs, and TPUs. The library integrates well with probabilistic programming languages by working directly with the (un-normalized) target log density function. BlackJAX is intended as a collection of low-level, composable implementations of basic statistical 'atoms' that can be combined to perform well-defined Bayesian inference, but also provides high-level routines for ease of use. It is designed for users who need cutting-edge methods, researchers who want to create complex sampling methods, and people who want to learn how these work.

Alberto Cabezas, Adrien Corenflos, Junpeng Lao, Rémi Louf, Antoine Carnec, Kaustubh Chaudhari, Reuben Cohn-Gordon, Jeremie Coullon, Wei Deng, Sam Duffield, Gerardo Durán-Martín, Marcin Elantkowski, et al.
arXiv:2402.10797 · cs.MS, cs.LG, stat.CO, stat.ML · submitted Feb 16, 2024 · updated Feb 22, 2024
abstract · pdf · html · Companion paper for the library https://github.com/blackjax-devs/blackjax Update: minor changes and updated the list of authors to include technical contributors

add comment on HN