Metadata-Version: 2.4
Name: rng-jax
Version: 0.0.3
Summary: JAX random number generation as a NumPy generator
Project-URL: Repository, https://github.com/glass-dev/rng-jax
Project-URL: Issues, https://github.com/glass-dev/rng-jax
Author-email: Nicolas Tessore <n.tessore@ucl.ac.uk>
License-Expression: MIT
Classifier: License :: OSI Approved :: MIT License
Classifier: Operating System :: OS Independent
Classifier: Programming Language :: Python :: 3
Requires-Python: >=3.9
Requires-Dist: jax
Provides-Extra: doc
Requires-Dist: furo; extra == 'doc'
Requires-Dist: sphinx; extra == 'doc'
Provides-Extra: test
Requires-Dist: pytest; extra == 'test'
Requires-Dist: pytest-cov; extra == 'test'
Description-Content-Type: text/markdown

# `rng-jax` — JAX random number generation as a NumPy generator

**This is a proof of concept only.**

Wraps JAX's stateless random number generation in a class implementing the
[`numpy.random.Generator`](generator) interface.

## Example

```py
>>> 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](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 `.key()`
method to obtain a new random key and advance the internal state of the
generator:

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

The right way to compile array-agnostic code is usually to compile the "main"
function at the highest level of the code. Using the `rng_jax.Generator` class
fully _within_ a compiled function works without issue.

[array-api]: https://data-apis.org/array-api/latest/
[generator]: https://numpy.org/doc/stable/reference/random/generator.html
