Skip to main content

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)

Source distribution for numba4jax 0.0.14
File Size Uploaded
numba4jax-0.0.14.tar.gz 12.0 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for numba4jax 0.0.14
File Interpreter ABI Platform
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

Release history Release notifications | RSS feed

This release

0.0.14 This release

2 release files

0.0.9

2 release files

0.0.8

2 release files

0.0.7

2 release files

0.0.6

2 release files

0.0.5

2 release files

0.0.4

2 release files

0.0.3

2 release files

0.0.2

2 release files

0.0.1

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