Library / SDK
pyro-ppl/numpyro avatar
pyro-ppl/numpyro

NumPyro: JAX-backed probabilistic programming for MCMC and SVI

Probabilistic programming with NumPy powered by JAX for autograd and JIT compilation to GPU/TPU/CPU.

2,758 stars316 forksPythonApache-2.0

At a glance

What is it?
NumPyro wraps Pyro's modelling primitives around JAX, so NUTS and ADVI run as compiled XLA kernels on CPU, GPU or TPU. It is a good fit when you already write JAX and want to keep the sampler on the same device as the model.
Who is it for?
Adopt NumPyro if your model is already expressed in JAX or NumPy and you want NUTS and ADVI compiled to the device you train on; if you need discrete latent variables in a large model, read the MixedHMC and TraceEnum_ELBO docs before committing. Skip it if you want a batteries-included modelling language, since NumPyro is deliberately a substrate rather than a DSL.
Can I use it commercially?
Yes. Apache-2.0 is a permissive licence: you can use, modify and sell software built on it, as long as you keep its copyright and licence notices.
Is it still maintained?
Yes. The repository last received commits 7 days ago.
What is it written in?
Mainly Python, according to GitHub's language statistics.

Answers come from the project's GitHub data, last synced on September 24, 2026, and from our analysis. They are not legal advice.

Editorial analysis

What NumPyro is for, and who ends up using it

NumPyro is a probabilistic programming library that provides a NumPy backend for Pyro. The models you write look like Pyro models: ordinary Python plus primitives such as sample and param, with NumPy-style array code instead of PyTorch tensor code. The README is explicit that the design goal is to be lightweight and to act as a flexible substrate for users to build on, which is a different promise from a modelling language that ships a fixed set of declared model types.

The audience follows from that. If your likelihood is already a JAX function, or your data pipeline is NumPy arrays, NumPyro lets you write the generative model in the same idiom and hand the whole thing to an inference algorithm. The repository ships an examples directory with models such as bnn.py, hmm.py, gp.py, horseshoe_regression.py and ode.py, which is a fair picture of the intended range: hierarchical regression, hidden Markov models, Gaussian processes, neural network priors and differential equation models. The README also warns that NumPyro is under active development and that you should expect brittleness, bugs and API changes as the design evolves. Take that sentence at face value rather than as boilerplate.

How JAX compilation changes the sampler, not just the speed

The README states one of the motivations for NumPyro directly: speeding up Hamiltonian Monte Carlo by JIT compiling the Verlet integrator, which contains multiple gradient computations. With JAX, jit and grad compose, so the entire integration step becomes one XLA-optimized kernel. NumPyro also removes Python overhead from NUTS by compiling the whole tree-building stage, which the README attributes to Iterative NUTS.

That is an architectural decision with consequences beyond throughput. Because the sampler is compiled, the model function has to be traceable and its shapes have to be static enough for XLA. Plates exist for exactly this reason: numpyro.plate declares a batch dimension so that the sample sites inside it are vectorized rather than looped over in Python. If you write a model whose control flow depends on sampled values in a way JAX cannot trace, the compilation step is where you find out.

The inference surface is wider than NUTS alone. The README lists MixedHMC for models with discrete latent variables and HMCECS, which computes the likelihood on subsets of the data each iteration. On the variational side there is a basic ADVI implementation with autoguides, plus TraceGraph_ELBO and TraceEnum_ELBO for models containing discrete latent variables. Distributions, constraints and bijective transforms live in numpyro.distributions, and the README says the design largely follows torch.distributions, with a major subset of that API implemented. TensorFlow Probability distributions can be used inside NumPyro models through numpyro.contrib.tfp.distributions.TFPDistribution. Effect handlers in numpyro.handlers give sample and param nonstandard interpretations, which is how custom inference utilities get built.

Installing NumPyro and running the eight schools model

The README links its Installation section to the documentation rather than spelling out a pip command in the body, so the package name to use is the project name numpyro. What the repository does specify precisely is the dependency floor and the optional extras. The pyproject.toml sets requires-python to >=3.11 and depends on jax>=0.7.0, jaxlib>=0.7.0, multipledispatch, numpy and tqdm. The extras are named cpu, cuda12, cuda13 and tpu, each pinning the matching JAX build at version 0.7.0 or newer. Picking the wrong extra is the most common first mistake, because it decides which JAX wheel you get.

bash
pip install -e '.[dev,doc,test,examples]'

That command is the editable development install from the repository Makefile, which is what contributors use rather than what an application should depend on. The extras named in pyproject.toml are what an application selects from.

The README's worked example is the eight schools model from Gelman et al. The data are eight treatment effects y and their standard errors sigma. The model samples a population mean mu from a Normal(0, 5), a between-school scale tau from a HalfCauchy(5), then inside a plate over the eight schools samples theta from Normal(mu, tau) and registers the observation with obs=y.

python
import numpyro
import numpyro.distributions as dist

def eight_schools(J, sigma, y=None):
    mu = numpyro.sample('mu', dist.Normal(0, 5))
    tau = numpyro.sample('tau', dist.HalfCauchy(5))
    with numpyro.plate('J', J):
        theta = numpyro.sample('theta', dist.Normal(mu, tau))
        numpyro.sample('obs', dist.Normal(theta, sigma), obs=y)

Inference is then run with MCMC using the No-U-Turn Sampler. The README highlights the extra_fields argument to MCMC.run: by default the run collects samples from the posterior only, and extra_fields is how you ask for additional quantities such as potential energy or the acceptance probability. The README truncates before listing every available field, so check the MCMC.run documentation for the accepted names rather than guessing.

Where NumPyro is the wrong choice

The same compilation that makes NUTS fast makes some models awkward. A model whose structure is decided at runtime by Python control flow over sampled values will not trace cleanly, and the fix is usually to rewrite the model in terms of plates and vectorized operations. That rewriting is real work, and for a one-off fit on a few hundred rows it buys nothing.

The README's own warning about brittleness matters for production. It says the API changes as the design evolves, and the release cadence supports reading that seriously: 0.20.1 in March 2026, 0.21.0 in May 2026 and 0.22.0 in September 2026. If you pin NumPyro, pin JAX with it, because the jax>=0.7.0 floor means an unrelated JAX upgrade can move your sampler's behaviour. There is no compatibility shim described for that pairing.

Discrete latent variables are a second boundary. MixedHMC and the enumeration-based ELBOs exist, and the README presents them as the answer for models with discrete latents, but they are separate algorithms with their own constraints rather than a transparent fallback. If your model is mostly discrete, check those pages before assuming the NUTS path applies.

Finally, NumPyro is a substrate, not a modelling language. It gives you primitives, distributions, handlers and algorithms. It does not give you a declarative syntax that a non-programmer can edit, and it does not hide the array shapes from you.

NumPyro against Pyro, PyMC and Stan

The closest comparison is Pyro, since NumPyro is described as a NumPy backend for it. The model code is meant to look very similar, with differences coming from the NumPy versus PyTorch API split, and the distributions module deliberately mirrors torch.distributions so that Pyro users can carry over their batching intuition. The practical difference is the execution stack: Pyro runs on PyTorch, NumPyro on JAX, which is what puts GPU and TPU compilation and the fused HMC integration step within reach. If your team already ships PyTorch, moving to NumPyro means adopting a second array framework, and that cost is often larger than the sampler speedup.

PyMC and Stan take the other approach: a modelling language with a compiler or a graph builder in front, and an ecosystem of diagnostics and reporting around it. NumPyro asks you to write Python and manage shapes yourself. The trade is expressiveness for convenience. A model that is a few lines of Stan may be more lines in NumPyro, but a model that needs a custom likelihood calling into JAX code that already exists is natural in NumPyro and awkward elsewhere.

The comparison to JAX itself is a category error that shows up in search results. JAX is the numerical foundation; NumPyro is the probabilistic layer that provides sample, plate, distributions and the inference algorithms. You do not choose between them.

Maintenance, licensing and the upgrade bill

The repository is not archived. The last push was on 2026-09-23, five days before the date used here, and the most recent release is 0.22.0 from 2026-09-18. That is a live project by any reasonable reading, with three releases in the six months before this article.

NumPyro is licensed under Apache-2.0, and the pyproject.toml declares the license as a file reference to LICENSE and carries the Apache Software License classifier. Apache-2.0 is permissive and includes an explicit patent grant, which matters if you are embedding the library in a commercial product. This is a description of the licence text, not legal advice; if you are redistributing a modified NumPyro, read LICENSE.md and the NOTICE handling yourself.

The upgrade cost is dominated by the JAX coupling rather than by NumPyro's own API. The dependency floor is jax>=0.7.0 and jaxlib>=0.7.0, and the cpu, cuda12, cuda13 and tpu extras each pin their JAX build at the same floor. In practice that means a NumPyro bump can pull a JAX bump, so test the sampler on a representative model rather than only checking that imports succeed. The repository maintains a CHANGELOG linked from pyproject.toml, and that file is the place to look for the API changes the README warns about. The Makefile shows the project's own gates: ruff check, ruff format --check, a header check script, and ty check, followed by pytest. Running those locally is how a contributor verifies a patch, and the doctest target pins JAX_PLATFORM_NAME=cpu so documentation examples do not need a GPU.

Editorial conclusion

Adopt NumPyro if your model is already expressed in JAX or NumPy and you want NUTS and ADVI compiled to the device you train on; if you need discrete latent variables in a large model, read the MixedHMC and TraceEnum_ELBO docs before committing. Skip it if you want a batteries-included modelling language, since NumPyro is deliberately a substrate rather than a DSL. Before writing your own model, run the eight schools example, check the JAX version your environment resolves against the jax>=0.7.0 floor in pyproject.toml, and read the CHANGELOG for the API changes the README warns about.

Frequently asked questions

What is NumPyro?

It is a lightweight probabilistic programming library that provides a NumPy backend for Pyro, using JAX for automatic differentiation and JIT compilation to GPU, CPU or TPU. Models use Pyro primitives such as sample and param, and inference runs through algorithms including NUTS, MixedHMC, HMCECS and ADVI.

How do I install NumPyro?

Install the numpyro package from PyPI, choosing the extra that matches your hardware: cpu, cuda12, cuda13 or tpu. Each extra pins the corresponding JAX build at version 0.7.0 or newer, and the package requires Python 3.11 or later.

How does NumPyro differ from Pyro?

NumPyro is a NumPy backend for Pyro, so the model code is meant to look very similar apart from differences between the NumPy and PyTorch APIs. The distributions module largely follows torch.distributions, and the execution stack is JAX rather than PyTorch.

How does NumPyro relate to JAX?

NumPyro relies on JAX for automatic differentiation and JIT compilation, and its models are JAX-traceable Python. The dependency floor in pyproject.toml is jax>=0.7.0 and jaxlib>=0.7.0, so the two move together.

Official sources

  1. License: Apache-2.0
  2. Project website
  3. pyro-ppl/numpyro on GitHub
  4. README
  5. Releases
Add this badge to your README

If you maintain this project, the badge below links readers to this analysis and shows its maintenance status from the daily GitHub snapshot. Paste the markdown into your README; add ?metric=license or ?metric=stars to the image URL for a different field.

Add this badge to your README

markdown
[![Hysen Labs](https://hysenlabs.com/badge/pyro-ppl-numpyro.svg)](https://hysenlabs.com/projects/pyro-ppl-numpyro)