Skip to main content

rng-jax — NumPy random number generator API for JAX

This is a proof of concept only.

Wraps JAX's stateless random number generation in a class implementing the numpy.random.Generator interface.

Example

>>> import rng_jax
>>> rng = rng_jax.Generator(42)  # same arguments as jax.random.key()
>>> rng.standard_normal(3)
Array([-0.5675502 ,  0.28439185, -0.9320608 ], dtype=float32)
>>> rng.standard_normal(3)
Array([ 0.67903334, -1.220606  ,  0.94670606], dtype=float32)

Rationale

The Array API makes it possible to write array-agnostic Python libraries. The rng-jax package makes it easy to extend this to random number generation in NumPy and JAX. End users only need to provide a rng object, as usual, which can either be a NumPy one or a rng_jax.Generator instance wrapping JAX's stateless random number generation.

How it works

The rng_jax.Generator class works in the obvious way: it keeps track of the JAX key and calls jax.random.split() before every random operation.

JIT and native JAX code

The problem with a stateful RNG is that it cannot be passed into a compiled JAX function. In practice, this is not usually an issue, since the goal of this package is to work in tandem with the Array API: array-agnostic code is not usually compiled at low level. Conversely, native JAX code usually expects a key, anyway, not a rng_jax.Generator instance.

To interface with a native JAX function expecting a key, use the .split() method to obtain a new random key and advance the internal state of the generator:

>>> import jax
>>> rng = rng_jax.Generator(42)
>>> key = rng.split()
>>> jax.random.normal(key, 3)
Array([-0.5675502 ,  0.28439185, -0.9320608 ], dtype=float32)
>>> key = rng.split()
>>> jax.random.normal(key, 3)
Array([ 0.67903334, -1.220606  ,  0.94670606], dtype=float32)

Using the rng_jax.Generator class fully within a compiled JAX function works without issue.

Metadata

Release files for rng-jax 0.0.4

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

Source distribution (sdist)

Source distribution for rng-jax 0.0.4
File Size Uploaded
rng_jax-0.0.4.tar.gz 5.5 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for rng-jax 0.0.4
File Interpreter ABI Platform
rng_jax-0.0.4-py3-none-any.whl Python 3 none any Details

Total release size: 9.9 kB

Release files / rng_jax-0.0.4.tar.gz

Download URL rng_jax-0.0.4.tar.gz
Size 5.5 kB
Tags Source
SHA-256 checksum
How to use checksums
6c92631774d72a5e0bf4c46df2b49a0ec0cc6d9e1ebefc2ea05a7e3902545354
BLAKE2b-256 checksum
How to use checksums
1196ade8e7fdfafa0f7bc1a836085aa33a561a50940d0c5c081a3a52e1e2531d
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/6.0.1 CPython/3.12.8

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 Dec 29, 2024.

Transparency log

Release files / rng_jax-0.0.4-py3-none-any.whl

Download URL rng_jax-0.0.4-py3-none-any.whl
Size 4.4 kB
Tags Python 3
SHA-256 checksum
How to use checksums
746ec5cf57a8e57b60289a69d241e8c02ba7c521250c59d0bc7a02171d16a26e
BLAKE2b-256 checksum
How to use checksums
1c6e156a7fc2b27839ecc2ba3adc514ff78c30926f3e062bfe1982160251889e
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/6.0.1 CPython/3.12.8

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 Dec 29, 2024.

Transparency log

Release history Release notifications | RSS feed

This release

0.0.4 This release

2 release files

0.0.3

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