Skip to main content

numerax

tests Coverage Status docs DOI

Statistical and numerical computation functions for JAX, focusing on tools not available in the main JAX API.

📖 Documentation

Installation

pip install numerax

# With scientific ML dependencies like equinox
pip install numerax[sciml]

Features

Special Functions

Differentiable special functions missing from JAX:

import jax.numpy as jnp
import numerax

# Inverse functions for statistical distributions
x = numerax.special.gammap_inverse(p, a)  # Gamma quantiles
y = numerax.special.erfcinv(x)  # Inverse complementary error function

# Modified Bessel functions of the first kind, real order
i = numerax.special.ive(v, z)  # exp(-z) I_v(z); stable for large z
i = numerax.special.iv(v, z)   # I_v(z)

# Chi-squared distribution (includes JAX functions + custom ppf)
x = numerax.stats.chi2.ppf(q, df, loc=0, scale=1)

Key features:

  • Inverse functions for statistical distributions missing from JAX
  • Full differentiability and JAX transformation support

Profile Likelihood

Efficient profile likelihood computation for statistical inference with nuisance parameters:

import jax.numpy as jnp
import numerax

# Example: Normal distribution with mean inference, variance profiling
def normal_llh(params, data):
    mu, log_sigma = params
    sigma = jnp.exp(log_sigma)
    return jnp.sum(-0.5 * jnp.log(2 * jnp.pi) - log_sigma 
                   - 0.5 * ((data - mu) / sigma) ** 2)

# Profile over log_sigma, infer mu
is_nuisance = [False, True]  # mu=inference, log_sigma=nuisance

def get_initial_log_sigma(data):
    return jnp.array([jnp.log(jnp.std(data))])

profile_llh = numerax.stats.make_profile_llh(
    normal_llh, is_nuisance, get_initial_log_sigma
)

# Evaluate profile likelihood
data = jnp.array([1.2, 0.8, 1.5, 0.9, 1.1])
llh_val, opt_nuisance, diff, n_iter = profile_llh(jnp.array([1.0]), data)

Key features:

  • Convergence diagnostics and configurable optimization parameters
  • Automatic parameter masking for inference vs. nuisance parameters

Utilities

Utilities for working with PyTree-based models, including parameter counting and model summaries.

from numerax.utils import count_params, tree_summary
import jax.numpy as jnp

# Count parameters in PyTree-based models
model = {"weights": jnp.ones((10, 5)), "bias": jnp.zeros(5)}
num_params = count_params(model)  # 55 parameters

# Pretty-print model structure (similar to Keras model.summary())
model = {
    "encoder": {
        "weights": jnp.ones((10, 20)),
        "bias": jnp.zeros(20),
    },
    "decoder": {
        "weights": jnp.ones((20, 5)),
        "bias": jnp.zeros(5),
    },
}
tree_summary(model)
# ======================================================================
# PyTree Summary
# ======================================================================
# Name                  Shape           Dtype             Params
# ----------------------------------------------------------------------
# root                                                       325
#   encoder                                                  220
#     - weights         [10,20]         float32              200
#     - bias            [20]            float32               20
#   decoder                                                  105
#     - weights         [20,5]          float32              100
#     - bias            [5]             float32                5
# ======================================================================
# Total params: 325
# ======================================================================

Key features:

  • Parameter counting for PyTree-based models including Equinox (requires numerax[sciml])
  • Model structure visualization with shapes, dtypes, and parameter counts
  • Decorators for preserving function metadata when using JAX's advanced features

Acknowledgements

This work is supported by the Department of Energy AI4HEP program.

Citation

If you use numerax in your research, please cite it using the citation information from Zenodo (click the DOI badge at the top of the README) to ensure you get the correct DOI for the version you used.

Metadata

Release files for numerax 1.4.0

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for numerax 1.4.0
File Size Uploaded
numerax-1.4.0.tar.gz 288.0 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for numerax 1.4.0
File Interpreter ABI Platform
numerax-1.4.0-py3-none-any.whl Python 3 none any Details

Total release size: 312.0 kB

Release files / numerax-1.4.0.tar.gz

Download URL numerax-1.4.0.tar.gz
Size 288.0 kB
Tags Source
SHA-256 checksum
How to use checksums
9ef896c5763b8a57bcbc66c9256a68b47a89f6b00dd0f3f9c25ccebc568cdda8
BLAKE2b-256 checksum
How to use checksums
5069ec14436de9745d6e06fe47bbf1e63f009517dfcc6a85cb61f053f5660403
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/6.1.0 CPython/3.13.12

Provenance

Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.

PyPI Publish Attestation

PyPI verified that this artifact, at this checksum, originated from the publisher listed below.

Signed by GitHub Actions, verified by PyPI on Jun 8, 2026.

Transparency log

Release files / numerax-1.4.0-py3-none-any.whl

Download URL numerax-1.4.0-py3-none-any.whl
Size 24.0 kB
Tags Python 3
SHA-256 checksum
How to use checksums
83a7f8f08022a60cb89be9e98584d6ac557bb1338af70fa60e278c7bc1f579d4
BLAKE2b-256 checksum
How to use checksums
2504cd629d8cd4d2c1328b9a2a4cc1f73ae9ebe4f7b1202522a2cd25ffb4930e
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/6.1.0 CPython/3.13.12

Provenance

Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.

PyPI Publish Attestation

PyPI verified that this artifact, at this checksum, originated from the publisher listed below.

Signed by GitHub Actions, verified by PyPI on Jun 8, 2026.

Transparency log

Release history Release notifications | RSS feed

This release

1.4.0 This release

2 release files

1.3.1

2 release files

1.3.0

2 release files

1.2.0

2 release files

1.1.0

2 release files

1.0.2

2 release files

1.0.0

2 release files

0.3.0

2 release files

0.2.0

2 release files

0.1.0

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page