$npx skillfedfor your agent

distrax

Distrax: Probability distributions in JAX.

Worth itPyPI Python ModulesReleased Jun 2026188.5K downloads / mopermissive licensePure Python

Decision gist · record as of 2026-08-14

pure-Python wheel — distrax-0.1.9-py3-none-any.whl
v0.1.9 · released 2026-06-12 · Python >=3.11 · 6 runtime deps: absl-py, chex, jax, jaxlib, numpy, tfp-nightly

Yes. Distrax is actively maintained, has no known vulnerabilities, and offers a well-designed JAX-native alternative to TensorFlow Probability for teams already using JAX. The permissive Apache 2.0 license and cross-compatibility with TFP make it low-risk to adopt. Install it if you need probability distributions in JAX and value extensibility and mathematical clarity over a comprehensive feature set.AI-flagged interpretation of the facts on this page — verify before relying

Before you install

  • Requires Python >= 3.11 and a working JAX installation (including jaxlib, which may require a C++ compiler or prebuilt binaries for your platform).
  • Installation is straightforward with low friction; the package is actively maintained with a recent release and no known vulnerabilities.
  • It depends on JAX, jaxlib, and TensorFlow Probability nightly, which are substantial but standard dependencies in the JAX ecosystem.

License · maintenance · safety

permissive license (permissive) — Licensed under Apache 2.0 (permissive), allowing free use, modification, and distribution in commercial and private projects with minimal restrictions.

last release 2026-06-12 (63 days) · last repo commit 2026-07-30 · 649 stars

0 known vulnerabilities (OSV.dev, 2026-08-14) · 188,525 downloads/mo, #9,941 on PyPI

Verify before relying

pip install distrax

import distrax
import jax
import jax.numpy as jnp

key = jax.random.PRNGKey(1234)
mu = jnp.array([-1., 0., 1.])
sigma = jnp.array([0.1, 0.2, 0.3])
dist = distrax.MultivariateNormalDiag(mu, sigma)
samples = dist.sample(seed=key)
log_prob = dist.log_prob(samples)
  • Whether the library's API stability is guaranteed or if breaking changes are expected before a 1.0 release.
  • Performance characteristics compared to TensorFlow Probability for common distributions and sampling operations.
  • Coverage of which TFP distributions and bijectors are actually implemented in Distrax.
Same gist for agents: .md · .json

What it is and what it does

Distrax is a JAX-native probability library that reimplements a subset of TensorFlow Probability with emphasis on readability and extensibility. It provides distributions (like MultivariateNormalDiag) and bijectors (invertible transformations with Jacobian tracking) that work seamlessly with JAX's functional and JIT-compilation paradigms. The library is designed to be cross-compatible with TensorFlow Probability—you can mix Distrax and TFP distributions in the same computation, or wrap one for use in the other's meta-distributions.

The package targets researchers and practitioners building probabilistic models, particularly in reinforcement learning where custom policy distributions are common. It emphasizes mathematical clarity in implementations and makes it simple to define custom distributions or bijectors. While not intended to replace TensorFlow Probability entirely, it fills the gap for teams already committed to JAX and wanting a lighter-weight, more extensible alternative for probability operations.

Use it for

  • Build probabilistic agent policies in reinforcement learning with custom distribution definitions.
  • Sample from and compute log-probabilities of multivariate distributions in JAX-based machine learning pipelines.
  • Compose complex distributions using bijectors (e.g., transformed distributions via Tanh or other invertible functions).
  • Migrate TensorFlow Probability code to JAX by using Distrax distributions with cross-compatible APIs.
  • Implement variational inference or other probabilistic inference methods using JAX's autodiff and JIT compilation.

Worth the install?

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

Worth it

Yes.

Distrax is actively maintained, has no known vulnerabilities, and offers a well-designed JAX-native alternative to TensorFlow Probability for teams already using JAX. The permissive Apache 2.0 license and cross-compatibility with TFP make it low-risk to adopt. Install it if you need probability distributions in JAX and value extensibility and mathematical clarity over a comprehensive feature set.

Install

distrax on PyPI

Before you install

Installation is straightforward with low friction; the package is actively maintained with a recent release and no known vulnerabilities. It depends on JAX, jaxlib, and TensorFlow Probability nightly, which are substantial but standard dependencies in the JAX ecosystem.

Requires Python >= 3.11 and a working JAX installation (including jaxlib, which may require a C++ compiler or prebuilt binaries for your platform).

License in practice

Licensed under Apache 2.0 (permissive), allowing free use, modification, and distribution in commercial and private projects with minimal restrictions.

Quickstart

pip install distrax

import distrax
import jax
import jax.numpy as jnp

key = jax.random.PRNGKey(1234)
mu = jnp.array([-1., 0., 1.])
sigma = jnp.array([0.1, 0.2, 0.3])
dist = distrax.MultivariateNormalDiag(mu, sigma)
samples = dist.sample(seed=key)
log_prob = dist.log_prob(samples)

Verify before relying

  • Whether the library's API stability is guaranteed or if breaking changes are expected before a 1.0 release.
  • Performance characteristics compared to TensorFlow Probability for common distributions and sampling operations.
  • Coverage of which TFP distributions and bijectors are actually implemented in Distrax.

Package facts

Licensepermissive license permissive
Python supportSupports the current Python release >=3.11
Install frictionLow. Pure-Python wheel
Runtime dependencies
6 packages
absl-pychexjaxjaxlibnumpytfp-nightly
MaintenanceActively maintained 63 days since the last release
Last repo commit
First released
Downloads188,525 / month, #9,941 on PyPI 30-day window, as of 2026-08-14
Known vulnerabilitiesNone known OSV.dev, checked 2026-08-14
Classifiers
Development Status :: 4 - BetaEnvironment :: ConsoleIntended Audience :: DevelopersIntended Audience :: Science/ResearchLicense :: OSI Approved :: Apache Software LicenseOperating System :: OS IndependentProgramming Language :: PythonProgramming Language :: Python :: 3Topic :: Scientific/Engineering :: Artificial IntelligenceTopic :: Software Development :: Libraries :: Python Modules

Evidence: distrax-0.1.9-py3-none-any.whl

Tags

Capabilities
jax probability distributionsbijectors invertible functionstensorflow probability alternativejax random samplingprobabilistic distributions jaxmachine learning distributionsreinforcement learning policies
Topics
jax-ecosystemprobabilistic-modelingreinforcement-learning
PyPI keywords
jaxprobabilitydistributionpythonmachine learning

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 › “jax probability distributions”

  • distraxDistrax provides JAX-native probability distributions and bijectors…
  • tfp-nightlyTensorFlow Probability provides probabilistic modeling, statistical…
  • tensorflow-probabilityTensorFlow Probability provides probabilistic modeling, statistical…

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

More Python Modules packages

idna Worth it
PyPI · Python Modules · released Jun 2026

Converts domain names between Unicode and ASCII-compatible encoding (Punycode) according to IDNA 2008 and Unicode Technical Standard 46, with security validation and broader script coverage than the standard library.

Install it if you work with internationalized domain names, need to validate domains, or use HTTP clients that depend on it transitively.

BSD-3-Clausepure Python · 3.9+
1.8Bdownloads / mo
setuptools Worth it
PyPI · Python Modules · released Aug 2026

Setuptools is a Python build backend and package management tool that handles building, distributing, and installing Python packages, including support for C/C++ extension modules.

MITpure Python · 3.10+
1.6Bdownloads / mo
PyYAML Worth it
PyPI · Python Modules · released Sep 2025

PyYAML parses and emits YAML 1.1 data format, enabling serialization and deserialization of configuration files and Python objects to and from human-readable YAML text.

MITcompiled wheel · 3.8+
1.2Bdownloads / mo
pydantic Worth it
PyPI · Python Modules · released May 2026

Pydantic validates Python data structures against type hints, coercing and checking input at runtime to ensure it matches a declared schema.

MITpure Python · 3.9+
1.1Bdownloads / mo
annotated-types Worth it
PyPI · Python Modules · released Jul 2026

Provides reusable metadata objects for use with PEP-593 `typing.Annotated` to express common constraints like bounds, collection sizes, and predicates on types.

Install it if you use or build libraries that need to express type constraints in a standardized, inspectable way—or if you want to annotate your own types with…

MITpure Python · 3.10+
871.3Mdownloads / mo
typing-inspection Worth it
PyPI · Python Modules · released Aug 2026

Provides runtime tools to inspect and introspect Python type annotations, enabling programmatic examination of type hints at execution time.

MITpure Python · 3.10+
783.0Mdownloads / mo

See also chex · tensorflow-probability · jax · tfp-nightly · numpyro · jraph · blackjax · diffrax · jax-cuda12-pjrt · optax