$npx skillfedfor your agent

blackjax

Flexible and fast sampling in Python

Worth itPyPI Scientific/EngineeringReleased Jul 2026227.1K downloads / moApache License 2.0Pure Python

Decision gist · record as of 2026-08-14

pure-Python wheel — blackjax-1.6.2-py3-none-any.whl
v1.6.2 · released 2026-07-16 · Python >=3.11 · 6 runtime deps: jax, jaxlib, numpy, optax, scipy, typing-extensions

Yes. BlackJAX is actively maintained, has no known vulnerabilities, installs with low friction, and is licensed permissively. It is well-suited if you need MCMC sampling on JAX with GPU support, want composable algorithm building blocks, or are integrating sampling into a larger system. Not necessary if you only need sampling from standard distributions or are already using a full probabilistic programming framework.AI-flagged interpretation of the facts on this page — verify before relying

Before you install

  • Requires Python 3.11 or later.
  • GPU/TPU sampling requires separate JAX installation with hardware support; default JAX install runs on CPU only.
  • Installation is straightforward via pip with low friction.

License · maintenance · safety

Apache License 2.0 (permissive) — Licensed under Apache License 2.0 (permissive), allowing use in commercial and proprietary projects with minimal restrictions—you must include a copy of the license and state significant changes, but can use the code freely.

last release 2026-07-16 (29 days) · last repo commit 2026-08-14 · 1,109 stars

0 known vulnerabilities (OSV.dev, 2026-08-14) · 227,060 downloads/mo, #9,189 on PyPI

Verify before relying

pip install blackjax

import blackjax
import jax
import jax.numpy as jnp

def logdensity_fn(x):
    return -0.5 * jnp.sum(x**2)

step_size = 1e-3
inverse_mass_matrix = jnp.array([1., 1.])
nuts = blackjax.nuts(logdensity_fn, step_size, inverse_mass_matrix)
initial_position = jnp.array([0., 0.])
state = nuts.init(initial_position)

rng_key = jax.random.key(0)
state, info = nuts.step(rng_key, state)
  • Whether the library supports all modern JAX features and remains compatible with future JAX releases.
  • Performance benchmarks comparing BlackJAX samplers to other MCMC libraries on standard problems.
Same gist for agents: .md · .json

What it is and what it does

BlackJAX is a library of sampling algorithms—primarily MCMC methods like NUTS—built on top of JAX and designed to run efficiently on both CPU and GPU hardware. It is not a probabilistic programming language itself, but rather a collection of composable, reusable sampling kernels that integrate well with any tool that can provide a log-probability density function compatible with JAX.

The library targets two audiences: users who have a logpdf and just need a working sampler (via simple high-level APIs like `blackjax.nuts()`), and researchers building custom inference algorithms who want to reuse robust, well-tested building blocks like integrators, proposal generators, and momentum samplers. All kernels follow a stateless functional pattern—taking a random key and state, returning a new state and metadata—making them easy to compose and swap. The package depends on jax, jaxlib, numpy, optax, scipy, and typing-extensions.

Use it for

  • Sample from a posterior distribution when you have a log-probability function but no full probabilistic programming framework.
  • Build custom MCMC samplers by composing elementary kernels (integrators, proposals, momentum generators) without reimplementing from scratch.
  • Run inference on GPU or TPU by leveraging JAX's hardware acceleration for large-scale Bayesian computations.
  • Integrate sampling into an existing probabilistic programming language by providing a modular, decoupled sampler.
  • Learn how sampling algorithms work by studying and experimenting with the library's composable, well-documented building blocks.
  • Accelerate research on new sampling schemes using robust, performant, and reusable algorithm components.

Worth the install?

AI-flagged interpretation of the facts on this page. Verify before relying on it.

Worth it

Yes.

BlackJAX is actively maintained, has no known vulnerabilities, installs with low friction, and is licensed permissively. It is well-suited if you need MCMC sampling on JAX with GPU support, want composable algorithm building blocks, or are integrating sampling into a larger system. Not necessary if you only need sampling from standard distributions or are already using a full probabilistic programming framework.

Install

blackjax on PyPI

Before you install

Installation is straightforward via pip with low friction. The package is actively maintained with a recent release (29 days ago) and receives regular updates. Note that JAX defaults to CPU-only; GPU/TPU support requires separate JAX installation following JAX's hardware-specific instructions.

Requires Python 3.11 or later. GPU/TPU sampling requires separate JAX installation with hardware support; default JAX install runs on CPU only.

License in practice

Licensed under Apache License 2.0 (permissive), allowing use in commercial and proprietary projects with minimal restrictions—you must include a copy of the license and state significant changes, but can use the code freely.

Quickstart

pip install blackjax

import blackjax
import jax
import jax.numpy as jnp

def logdensity_fn(x):
    return -0.5 * jnp.sum(x**2)

step_size = 1e-3
inverse_mass_matrix = jnp.array([1., 1.])
nuts = blackjax.nuts(logdensity_fn, step_size, inverse_mass_matrix)
initial_position = jnp.array([0., 0.])
state = nuts.init(initial_position)

rng_key = jax.random.key(0)
state, info = nuts.step(rng_key, state)

Verify before relying

  • Whether the library supports all modern JAX features and remains compatible with future JAX releases.
  • Performance benchmarks comparing BlackJAX samplers to other MCMC libraries on standard problems.

Package facts

LicenseApache License 2.0 permissive
Python supportSupports the current Python release >=3.11
Install frictionLow. Pure-Python wheel
Runtime dependencies
6 packages
jaxjaxlibnumpyoptaxscipytyping-extensions
MaintenanceActively maintained 29 days since the last release
Last repo commit
First released
Downloads227,060 / month, #9,189 on PyPI 30-day window, as of 2026-08-14
Known vulnerabilitiesNone known OSV.dev, checked 2026-08-14
Classifiers
Development Status :: 5 - Production/StableIntended Audience :: DevelopersIntended Audience :: Information TechnologyIntended Audience :: Science/ResearchLicense :: OSI Approved :: Apache Software LicenseOperating System :: MacOSOperating System :: POSIXProgramming Language :: Python :: 3.11Programming Language :: Python :: 3.12Programming Language :: Python :: 3.13Topic :: EducationTopic :: Scientific/EngineeringTopic :: Scientific/Engineering :: Artificial IntelligenceTopic :: Scientific/Engineering :: Mathematics

Evidence: blackjax-1.6.2-py3-none-any.whl

Tags

Capabilities
mcmc sampling jaxbayesian inference gpunuts sampler pythonprobabilistic sampling libraryjax-based mcmchamiltonian monte carlocomposable sampling algorithms
Topics
mcmcbayesian-inferencegpu-accelerated
PyPI keywords
probabilitymachine learningstatisticsmcmcsampling

Let your AI agent find packages like this

Example. Real query, live index.

You found this page by searching. An agent finds it by wishing: SkillFed indexes 14,416 PyPI packages by what they can do, searchable in plain language.

wish › “mcmc sampling jax”

  • blackjaxBlackJAX provides a collection of MCMC and other sampling algorithms…
  • numpyroNumPyro is a probabilistic programming library that uses JAX for…
  • tensorflow-probabilityTensorFlow Probability provides probabilistic modeling, statistical…

Give your agent the search over MCP, or paste the wish link into any chat.

More Scientific/Engineering packages

numpy Worth it
PyPI · Software Development · released Aug 2026

NumPy provides an N-dimensional array object and a comprehensive suite of mathematical, linear algebra, Fourier transform, and random number functions for scientific computing in Python.

BSD-3-Clause AND 0BSD AND MIT AND Zlib AND CC0-1.0compiled wheel · 3.12+
1.1Bdownloads / mo
pandas Worth it
PyPI · Scientific/Engineering · released Jul 2026

pandas provides fast, flexible data structures (Series and DataFrame) for loading, cleaning, transforming, and analyzing labeled or relational data in Python.

BSD-3-Clausecompiled wheel · 3.11+
769.1Mdownloads / mo
scipy Worth it
PyPI · Libraries · released Jun 2026

scipy provides numerical algorithms for mathematics, science, and engineering—including optimization, integration, linear algebra, Fourier transforms, signal and image processing, and ODE solvers—built on numpy arrays.

BSD-3-Clausecompiled wheel · 3.12+
449.0Mdownloads / mo
scikit-learn Worth it
PyPI · Software Development · released Jun 2026

scikit-learn provides a comprehensive Python library for supervised and unsupervised machine learning, including classification, regression, clustering, dimensionality reduction, and model evaluation tools built on NumPy and SciPy.

Install it if you need to train, evaluate, or deploy supervised or unsupervised learning models.

BSD-3-Clausecompiled wheel · 3.11+
235.5Mdownloads / mo
dill Worth it
PyPI · Software Development · released Jan 2026

dill extends Python's pickle module to serialize and deserialize a much wider range of Python objects, including functions, lambdas, classes, and interpreter sessions, to byte streams for storage or network transmission.

BSD-3-Clausepure Python · 3.9+
208.1Mdownloads / mo
multiprocess Worth it
PyPI · Software Development · released Jan 2026

Multiprocess is an enhanced fork of Python's standard multiprocessing library that uses dill for better serialization, allowing you to spawn processes with a threading-like API and share complex objects between them.

Install it if you use multiprocessing and encounter pickle serialization limits with lambdas or complex objects.

BSD-3-Clausepure Python · 3.9+
202.7Mdownloads / mo

See also numpyro · jax-cuda13-pjrt · jax-cuda13-plugin · jax-cuda12-plugin · nutpie · jax-cuda12-pjrt · jaxlib · emcee · jax · distrax