Skip to main content

numba-enzyme-cuda

Differentiate Numba-compiled Python functions via Enzyme.

This is the CUDA-enabled fork of numba-enzyme, published under its own name because the upstream release carries no CUDA backend. It provides the same numba_enzyme import package, so it is a drop-in replacement -- and the two must not be installed side by side.

Install

pip install numba-enzyme-cuda

Supported platforms

OS / architecture Linux x86_64 only
glibc ≥ 2.39 (e.g. Ubuntu 24.04+, Debian 13+, Fedora 39+)
Python CPython 3.10, 3.11, 3.12, 3.13

Usage

Write an ordinary Python function and pass it to grad, jvp, jacfwd, or any of the other transforms below. No type annotations are needed: like a Numba njit function, each derivative callable infers the argument types from the values it is called with and compiles one specialization per distinct set of types, reusing it afterwards. Arguments must be floating-point scalars. Alternatively, annotate every parameter and the return value with numba_enzyme.types to fix the types up front, in which case the function is compiled as soon as it is transformed. You can also decorate your function f with @differentiable to expose each operation as an attribute.

import math  # import numpy as np
from numba_enzyme import (
    differentiable,
    grad,
    jacfwd,
    jacrev,
    jvp,
    vjp,
)
from numba_enzyme.types import Float64

def f(x, y):
    return x * y + math.cos(x * y)
    # alternatively, x * y + np.cos(x * y)

grad(f)(1.0, 2.0)               # -> (df/dx, df/dy), compiled for float64 args
jvp(f)((1.0, 2.0), (1.0, 0.0))  # -> directional derivative along (1.0, 0.0)

def f_vec(x, y):
    return x * y, x * x + y

jacfwd(f_vec)(1.0, 2.0)            # -> ((2.0, 1.0), (2.0, 1.0))
jvp(f_vec)((1.0, 2.0), (1.0, 0.0)) # -> column 0, (2.0, 2.0)
jacrev(f_vec)(1.0, 2.0)            # -> ((2.0, 1.0), (2.0, 1.0))
vjp(f_vec)((1.0, 2.0), (0.0, 1.0)) # -> row 1, (2.0, 1.0)

@differentiable
def g(x: Float64, y: Float64) -> Float64:  # annotations fix the types up front
    return x * y + math.cos(x * y)
    # alternatively, x * y + np.cos(x * y)

g(1.0, 2.0)           # calls the original Python function directly
g.grad(1.0, 2.0)      # reverse-mode gradient, built lazily on first access
g.jvp((1.0, 2.0), (1.0, 0.0))  # forward-mode JVP

grad requires its target function to return exactly one scalar. It does not accept tuple-valued targets. jvp and jacfwd handle both scalar and vector outputs, as do their reverse-mode counterparts vjp and jacrev. CPU vector outputs use fixed-size homogeneous tuples. Full Jacobians are returned as output-by-input tuple matrices.

Choosing a differentiation operation

API Mode Result and typical use
grad Reverse Gradient of a single scalar output in one reverse sweep. Prefer this for scalar losses, particularly with many inputs.
jacfwd Forward Complete Jacobian, using one sweep per input. Prefer it when there are relatively few inputs.
jacrev Reverse Complete Jacobian, using one sweep per output. Prefer it when there are relatively few outputs.
jvp Forward Jacobian-vector product J @ tangent without constructing the Jacobian. Use it for a known input direction, and for a single column, with a unit direction.
vjp Reverse Vector-Jacobian product cotangent @ J without constructing the Jacobian. Use it for backpropagation from a known output cotangent, and for a single row, with a unit cotangent.

For a scalar-output function, grad, jacfwd, and jacrev have the same values and tuple shape. Their computational paths differ: grad and jacrev use reverse mode, while jacfwd uses one forward sweep per input.

CUDA device functions

Install the optional Numba-CUDA-MLIR integration and a CUDA toolkit:

pip install 'numba-enzyme-cuda[cuda]'

An @cuda.jit(device=True) function can be differentiated on the host without adding Python type annotations. The result is another device callable that Numba-CUDA-MLIR automatically specializes and links into every kernel that uses it:

from numba_cuda_mlir import cuda, types
from numba_enzyme import grad, jvp

@cuda.jit(device=True)
def f(x, y):
    return x * y + x * x

df = grad(f)
jf = jvp(f)

@cuda.jit
def use_derivatives(xs, ys, gradients, directional_derivatives):
    i = cuda.grid(1)
    if i < xs.size:
        dx, dy = df(xs[i], ys[i])
        gradients[i, 0] = dx
        gradients[i, 1] = dy
        directional_derivatives[i] = jf(
            (xs[i], ys[i]), (1.0, 0.0)
        )

For CUDA, every transform specializes lazily from the concrete argument types at each call site, so none needs a signature. grad and jvp differentiate an ordinary scalar-returning device function. A primal with several outputs returns a homogeneous tuple, which Numba-CUDA-MLIR lowers to an LLVM struct returned by value: forward modes take Enzyme's tangent struct directly, and reverse modes differentiate an internal scalarisation sum_k w_k * f_k(x) with the weights inactive, since Enzyme does not accept an aggregate differential return. Nothing is staged through an output array on either side of the call. One forward sweep gives a whole Jacobian column:

import numba
from numba_enzyme import jacfwd, jvp

@cuda.jit(device=True)
def f_tuple(x, y):
    return (x * y, x * x + y)

whole = jacfwd(f_tuple)
sweep = jvp(f_tuple)

@cuda.jit
def use_jacfwd(xs, ys):
    i = cuda.grid(1)
    if i < xs.size:
        column = cuda.local.array(2, numba.float64)
        jac = cuda.local.array((2, 2), numba.float64)

        whole(jac, xs[i], ys[i])              # the whole matrix
        sweep(column, xs[i], ys[i], 0.0, 1.0) # just column 1

A jvp call takes the primal's arguments and then a direction set mirroring them, one set per sweep. jacfwd is that same loop with the identity supplied internally, and a Jacobian column is the special case of a unit direction — written out at the call site as above, it costs no more than the column, because the derivative links as LTO IR and nvJitLink inlines it before constant propagation: the zero components kill their tangent arithmetic and the sweep collapses to the one column that survives.

Tuple arguments. A primal may also take homogeneous tuples. Every derivative call then mirrors the primal's own argument list, with each tuple argument supplied as a contiguous array of its elements — so a wide primal is called with a fixed handful of arguments rather than one per scalar. The entry point loads the elements itself, before the Enzyme marker, so the function being differentiated is the same flat-scalar one either way:

from numba_cuda_mlir import types

@cuda.jit(device=True)
def rhs(ys, t, ps):                        # 48 states, 1 parameter
    return (ys[0] * ps[0], ys[1] + t, ...)

signature = types.UniTuple(types.float64, 48)(
    types.UniTuple(types.float64, 48), types.float64, types.UniTuple(types.float64, 1)
)
sweep = jvp(rhs, signature=signature)

# inside CUDA-compiled code, six arguments whatever the tuple lengths are:
sweep(tangent, ys[i], t, ps[i], dys, dt, dps)

An array cannot say how long the tuple it stands for is, so a primal with tuple arguments needs an explicit signature; the call only has to get array-ness right. The direction set mirrors the primal's arguments, so seeding dt = 1.0 with dys and dps zero yields d rhs / d t — no separate index into the flattened arguments, and nothing to materialise in a shape the caller does not already have.

vjp and jacrev accept either a scalar-returning or a tuple-returning primal, and tell them apart by how many arguments the call passes. They use the same output-by-input Jacobian layout:

from numba_enzyme import jacrev, vjp

reverse = jacrev(f_tuple)
product = vjp(f_tuple)

@cuda.jit
def use_reverse(x, y, cotangent, jacobian, input_cotangent):
    # The complete Jacobian.
    reverse(jacobian, x, y)

    # cotangent @ J; a unit cotangent picks out one output's row.
    product(cotangent, input_cotangent, x, y)

jacobian has shape (n_out, n_args). vjp accepts an n_out cotangent and fills input_cotangent. All arrays must be contiguous and use the primal's floating-point dtype.

Prefer one sweep at a time on a GPU: jacfwd's matrix costs n_out * n_args per thread, which stops being viable well before the dimensions jvp handles comfortably. Because the primal shape differs, the tuple modes cannot share a build with the scalar-return modes, and none of them is implied by the default.

Derivatives are emitted as NVVM LTO IR rather than PTX. Numba-CUDA-MLIR compiles the calling kernel to LTO IR too whenever a link item is LTO IR, so nvJitLink inlines the derivative into its caller instead of leaving an opaque call carrying a parameter per primal argument.

Each public transform builds only its own entry points, and each entry point carries its own Enzyme marker call. Scalar-output reverse wrappers return an input-gradient tuple and emit one C-ABI entry point per input argument; tuple-returning jacrev performs one generated reverse sweep per output row at run time.

CUDA derivatives specialize lazily from the concrete types at each compiled call site, just like ordinary Numba-CUDA-MLIR device functions, including wide primals whose calls pass more than 30 arguments. An explicit signature may optionally constrain the accepted specialization. A compute capability may also be selected explicitly; otherwise the current device's compute capability is used:

from numba_cuda_mlir import types

df = grad(
    f,
    signature=types.float64(types.float64, types.float64),
    cc=(8, 0),
)

CUDA support currently covers float32 and float64 device functions with homogeneous scalar inputs, returning one scalar or a homogeneous tuple of them. A primal that writes its outputs through an array argument is not supported; neither are other array inputs, mixed types, activity annotations, general mutation, or kernel (device=False) differentiation. Derivatives must be called from CUDA-compiled code, not the host. The LLVM 15 Enzyme bridge currently supports compute capabilities 7.x through 9.x; Numba-CUDA-MLIR's LLVM 20 path for Blackwell (10.x+) needs a matching newer Enzyme toolchain.

Development

Install the locked Python environment and bootstrap the matching LLVM 15 and Enzyme binaries once:

uv sync
uv run python packaging/bootstrap_dev_toolchain.py

The toolchain is extracted into the ignored .dev-toolchain/ directory and is discovered automatically; no PATH or NUMBA_ENZYME_PLUGIN_PATH changes are needed. The bootstrap is idempotent and accepts --force to refresh the toolchain. The development dependency set includes Numba-CUDA-MLIR's CUDA 12 compiler/runtime components, so a host with a compatible NVIDIA driver can run the complete suite directly:

uv run pytest

Scope

  • It is still not possible to mark the arguments as active or constant.
  • CPU arguments must be floating-point scalars; results may be scalar or fixed-size homogeneous tuples of scalars. General array inputs and outputs are not yet supported.
  • Linear algebra e.g. np.dot, np.linalg.norm etc are not supported.

License

Apache License 2.0 with LLVM Exceptions — see LICENSE and NOTICE.

Metadata

Release files for numba-enzyme-cuda 0.2.2

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

Source distribution (sdist)

Source distribution for numba-enzyme-cuda 0.2.2
File Size Uploaded
numba_enzyme_cuda-0.2.2.tar.gz 144.7 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for numba-enzyme-cuda 0.2.2
File Interpreter ABI Platform
numba_enzyme_cuda-0.2.2-py3-none-manylinux_2_38_x86_64.whl Python 3 none Linux glibc 2.38+ x86-64 Details

Total release size: 73.6 MB

Release files / numba_enzyme_cuda-0.2.2.tar.gz

Download URL numba_enzyme_cuda-0.2.2.tar.gz
Size 144.7 kB
Tags Source
SHA-256 checksum
How to use checksums
d9521f49d960cdaa9b8c3aa6d2a94827a42983c1f3697e60c6325dd325878c0d
BLAKE2b-256 checksum
How to use checksums
0d0e1dc26ce799ef4f344b46985da09aeab6153f62a573364f27e180bbbfd256
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/7.0.0 CPython/3.13.14

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 Oct 5, 2026.

Transparency log

Release files / numba_enzyme_cuda-0.2.2-py3-none-manylinux_2_38_x86_64.whl

Download URL numba_enzyme_cuda-0.2.2-py3-none-manylinux_2_38_x86_64.whl
Size 73.5 MB
Tags Linux glibc 2.38+ x86-64 Python 3
SHA-256 checksum
How to use checksums
e29ae7abf33b110cccebdcba4952ac9573a0c04e26610ebf451d1dcb400cf322
BLAKE2b-256 checksum
How to use checksums
9160d8255b8c2cd198d1fa9eaf308b75c2bbe4da8dd447cf58ff2437c332d2f8
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/7.0.0 CPython/3.13.14

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 Oct 5, 2026.

Transparency log

Release history Release notifications | RSS feed

This release

0.2.2 This release

2 release files

0.2.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