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)
| File | Size | Uploaded | |
|---|---|---|---|
| rng_jax-0.0.4.tar.gz | 5.5 kB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| 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 logRelease 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