numba4jax
A small experimental python package allowing you to use numba-jitted functions from within jax with no overhead.
This package uses the CFFI of Numba to expose the C Function pointer of your compiled function to XLA. It works both for CPU and GPU functions.
This package exports a single decorator @njit4jax, which takes an argument, a function
or Tuple describing the output shape of the function itself.
See the brief example below.
import jax
import jax.numpy as jnp
from numba4jax import ShapedArray, njit4jax
def compute_type(*x):
return x[0]
@njit4jax(compute_type)
def test(args):
y, x, x2 = args
y[:] = x[:] + 1
z = jnp.ones((1, 2), dtype=float)
jax.make_jaxpr(test)(z, z)
print("output: ", test(z, z))
print("output: ", jax.jit(test)(z, z))
z = jnp.ones((2, 3), dtype=float)
print("output: ", jax.jit(test)(z, z))
z = jnp.ones((1, 3, 1), dtype=float)
print("output: ", jax.jit(test)(z, z))
Backend support
This package supports both the CPU and GPU backends of jax.
The GPU backend is only supported on linux, and is highly experimental.
It requires CUDA to be installed in a standard path.
CUDA is found through numba.cuda, so you should first check that numba.cuda
works.
Metadata
Release files for numba4jax 0.0.14
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Source distribution (sdist)
| File | Size | Uploaded | |
|---|---|---|---|
| numba4jax-0.0.14.tar.gz | 12.0 kB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| numba4jax-0.0.14-py3-none-any.whl | Python 3 | none | any | Details |
Total release size: 28.2 kB
Release files / numba4jax-0.0.14.tar.gz
| Download URL | numba4jax-0.0.14.tar.gz |
|---|---|
| Size | 12.0 kB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
a16911c3d3d1ac72cd6d9fdd003c285b4b86fe365ca072b8187c228c5011630f
|
|
BLAKE2b-256 checksum How to use checksums |
e2a4f97a263f88bcd6aed229214ee6508d44ae5dbb22bcd748d071fda9c3a54c
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
Yes |
| Uploaded via |
twine/5.1.0 CPython/3.12.4
|
Release files / numba4jax-0.0.14-py3-none-any.whl
| Download URL | numba4jax-0.0.14-py3-none-any.whl |
|---|---|
| Size | 16.2 kB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
cd4a23b5e25a3a4fc5e9adb21ca06cb1cccaf07a31e0dfa979619bf0447d33c2
|
|
BLAKE2b-256 checksum How to use checksums |
49ded1a5b8df5efaeed0cafd79a9b32f893291274034761ef523de3452e2b123
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
Yes |
| Uploaded via |
twine/5.1.0 CPython/3.12.4
|