Skip to content

Internals

The modules behind solve. They are private — the package's surface is the two solve functions — but the design argument in the guide refers to them, so they are documented here.

Sparsity and colouring

modax._sparsity

Turning a Jacobian sparsity pattern into the fewest forward sweeps.

Forward-mode AD does not hand back a Jacobian; it hands back J v for a direction v. Seeding the unit vectors one at a time costs n_vars sweeps. But two columns that share no row are structurally orthogonal: their contributions to J v never land on the same component, so seeding both at once -- v = e_i + e_j -- returns both columns uncorrupted in a single sweep, and the sparsity pattern says which component belongs to which column.

Partitioning the columns into as few such groups as possible is exactly vertex colouring of the column intersection graph S^T S, which NetworkX's greedy colouring does well enough here: the bound that matters is the largest set of mutually overlapping columns, and greedy hits it on the patterns these solvers see.

The pattern a caller supplies must be a superset of the true nonzeros -- colouring a superset is conservative, colouring a subset silently corrupts entries where two columns in a group turn out to overlap after all. It need not cover the factorisation's fill-in: that has slots of its own, laid out by modax._sparse_direct, and no column of J writes them.

CompressedJacobian dataclass

CompressedJacobian(
    n_vars: int,
    n_colours: int,
    colour: tuple[int, ...],
    packed: tuple[int, ...] | None = None,
    packed_diagonal: tuple[int, ...] | None = None,
    n_slots: int | None = None,
)

A column-compressed Jacobian layout, and the seeds that fill it.

Entry (r, c) of the Jacobian lives at slot r * n_colours + colour[c]. The layout is dense in the rows and compressed in the columns, so a slot exists for every row of every colour; slots whose colour group has no nonzero in that row are written but never read, which costs nothing and keeps the write a straight run over a column block.

With no pattern every column conflicts with every other, colouring gives n_colours == n_vars, and this degenerates exactly to the dense row-major matrix -- so the dense path is not a special case in the kernel, just the uninformative end of the same mechanism.

modax._sparse_direct overrides the grid with packed: the slot of each entry in its own CSR image of L + U, so the sweeps deposit -J straight into the buffer the factorisation will work in.

size property

size: int

Elements in one trajectory's matrix.

diagonal property

diagonal: tuple[int, ...]

Slot of each (i, i) entry, for the 1/(h*gamma) term.

slot

slot(row: int, col: int) -> int

Where entry (row, col) lives.

Source code in modax/_sparsity.py
def slot(self, row: int, col: int) -> int:
    """Where entry ``(row, col)`` lives."""
    grid = row * self.n_colours + self.colour[col]
    if self.packed is None:
        return grid
    slot = self.packed[grid]
    if slot < 0 and row == col:
        assert self.packed_diagonal is not None
        return self.packed_diagonal[row]
    if slot < 0:
        raise ValueError(
            f"entry ({row}, {col}) is not in the pattern, so a packed "
            "layout has no slot for it; declare it in the sparsity pattern"
        )
    return slot

store_slots

store_slots() -> ndarray | None

Where each (row, colour) sweep value goes, or None for a run.

On the grid a colour group's values are a contiguous block and the kernel writes them straight down; packed, each one is scattered to its own slot and the positions no entry claims are simply not written.

Source code in modax/_sparsity.py
def store_slots(self) -> np.ndarray | None:
    """Where each ``(row, colour)`` sweep value goes, or ``None`` for a run.

    On the grid a colour group's values are a contiguous block and the
    kernel writes them straight down; packed, each one is scattered to its
    own slot and the positions no entry claims are simply not written.
    """
    if self.packed is None:
        return None
    return np.asarray(self.packed, dtype=np.int32)

seed_table

seed_table(min_width: int = 0) -> ndarray

(n_colours + 1, max(n_vars, min_width)) of tangent directions.

Row g is the indicator of colour group g, so one sweep seeded with it yields that whole group. The extra final row is all zeros, for the sweeps that vary t or a parameter rather than the state -- which is why the rows are widened to min_width: the zero row doubles as the null parameter direction, and there may be more parameters than state variables.

Source code in modax/_sparsity.py
def seed_table(self, min_width: int = 0) -> np.ndarray:
    """``(n_colours + 1, max(n_vars, min_width))`` of tangent directions.

    Row ``g`` is the indicator of colour group ``g``, so one sweep seeded
    with it yields that whole group. The extra final row is all zeros, for
    the sweeps that vary ``t`` or a parameter rather than the state -- which
    is why the rows are widened to ``min_width``: the zero row doubles as
    the null *parameter* direction, and there may be more parameters than
    state variables.
    """
    width = max(self.n_vars, min_width)
    seeds = np.zeros((self.n_colours + 1, width), dtype=np.float64)
    for c, g in enumerate(self.colour):
        seeds[g, c] = 1.0
    return seeds

dense_jacobian

dense_jacobian(n_vars: int) -> CompressedJacobian

Every column its own colour: the row-major dense matrix.

Source code in modax/_sparsity.py
def dense_jacobian(n_vars: int) -> CompressedJacobian:
    """Every column its own colour: the row-major dense matrix."""
    return CompressedJacobian(
        n_vars=n_vars, n_colours=n_vars, colour=tuple(range(n_vars))
    )

normalize_sparsity

normalize_sparsity(sparsity, n_vars: int) -> tuple[tuple[int, ...], ...]

Accept a dense mask, a scipy sparse matrix, or (row, col) pairs.

Returned as a hashable tuple of row tuples so the colouring can be cached on it: a kernel is rebuilt per pattern, not per call.

Source code in modax/_sparsity.py
def normalize_sparsity(sparsity, n_vars: int) -> tuple[tuple[int, ...], ...]:
    """Accept a dense mask, a scipy sparse matrix, or ``(row, col)`` pairs.

    Returned as a hashable tuple of row tuples so the colouring can be cached
    on it: a kernel is rebuilt per pattern, not per call.
    """
    if hasattr(sparsity, "tocoo"):  # scipy sparse
        coo = sparsity.tocoo()
        pairs = zip(coo.row.tolist(), coo.col.tolist())
    else:
        arr = np.asarray(sparsity)
        if arr.ndim == 2 and arr.shape == (n_vars, n_vars):
            rows, cols = np.nonzero(arr)
            pairs = zip(rows.tolist(), cols.tolist())
        elif arr.ndim == 2 and arr.shape[1] == 2:
            pairs = ((int(r), int(c)) for r, c in arr.tolist())
        else:
            raise ValueError(
                "sparsity must be an (n_vars, n_vars) mask, a scipy sparse "
                f"matrix, or an (nnz, 2) array of (row, col); got shape "
                f"{arr.shape} for n_vars={n_vars}"
            )
    by_row: list[set[int]] = [set() for _ in range(n_vars)]
    for r, c in pairs:
        if not (0 <= r < n_vars and 0 <= c < n_vars):
            raise ValueError(
                f"sparsity entry ({r}, {c}) lies outside an {n_vars}x{n_vars} Jacobian"
            )
        by_row[r].add(c)
    return tuple(tuple(sorted(cs)) for cs in by_row)

colour_sparsity cached

colour_sparsity(pattern: tuple[tuple[int, ...], ...]) -> CompressedJacobian

Colour a pattern's column intersection graph, fewest colours wins.

Source code in modax/_sparsity.py
@functools.cache
def colour_sparsity(pattern: tuple[tuple[int, ...], ...]) -> CompressedJacobian:
    """Colour a pattern's column intersection graph, fewest colours wins."""
    import networkx as nx

    n_vars = len(pattern)
    # Columns conflict when some row holds both. Building the adjacency from the
    # rows directly is cheaper than forming S^T S for the patterns seen here,
    # where a row has a handful of entries out of hundreds of columns.
    graph = nx.Graph()
    graph.add_nodes_from(range(n_vars))
    for cols in pattern:
        for a in range(len(cols)):
            for b in range(a + 1, len(cols)):
                graph.add_edge(cols[a], cols[b])

    best: dict[int, int] | None = None
    for strategy in _STRATEGIES:
        colouring = nx.coloring.greedy_color(graph, strategy=strategy)
        if best is None or max(colouring.values()) < max(best.values()):
            best = colouring
    assert best is not None

    colour = tuple(int(best.get(c, 0)) for c in range(n_vars))
    compressed = CompressedJacobian(
        n_vars=n_vars, n_colours=max(colour) + 1, colour=colour
    )
    _check_orthogonal(pattern, compressed)
    return compressed

The compiled sparse direct solver

modax._sparse_direct

A compiled sparse direct solver for the Rodas5P iteration matrix.

dense_lu_solver factorises M = I/(h*gamma) - J as though it were dense, which is n_vars ** 3 / 3 operations however few nonzeros M has. A hand-written solver such as DISCO-EB's Schur/Einstein-Boltzmann block-LU does far better, but only for the one block structure it was written against. This module takes the sparsity pattern the kernel is already given and compiles a direct solver for it -- any pattern, no structure assumed -- so the saving a hand-written solver bought is available to any caller who can say where its nonzeros are.

The analysis happens once, on the host, when the kernel is built:

  1. Order. The pattern is symmetrised and handed to SuiteSparse's AMD, which returns a fill-reducing permutation (see fill_reducing_order for why AMD and not COLAMD).
  2. Factor symbolically. The exact pattern of L + U for the permuted matrix, fill-in included, falls out of a pure-pattern Gaussian elimination (fill_pattern) -- no numbers, no device, no sample matrix.
  3. Lay out. That pattern becomes one CSR image of L + U, and the (row, col) -> slot map it defines becomes the CompressedJacobian the AD sweeps write into. The buffer is exactly nnz(L + U) elements: the symbolic pass is the memory footprint.
  4. Compile. The factorisation and the two triangular solves become cuda.jit(device=True) functions, in the same factorize_local(lu, ipiv) / solve_local(lu, ipiv, rhs) shape modax.rodas5P.dense_lu_solver has, so the kernel calls one or the other and has no branch. Where the structure is small enough they are emitted as straight-line code with every slot a literal; above that they fall back to loops over index tables in constant memory.

One trajectory per thread, as everywhere else in this kernel. Every thread runs the same pattern, so there is no divergence to pay for whichever form is emitted.

Why the straight-line form matters. Table-driven, a sparse routine spends a broadcast load on the index of every value it is about to read, and the load of the value cannot issue until that index arrives. With a big enough ensemble the other trajectories cover that latency; DISCO-EB's single-cosmology case is 128 trajectories, four warps on a 46-SM device, and there is nothing to cover it with. Spelling the indices out removes the dependency entirely, and it costs nothing anywhere else: neither the matrix nor the right-hand side can leave local memory, because the kernel indexes both with loop variables of its own, so this trades index loads for instruction count and not for registers. Measured on DISCO-EB at N128: 528 ms table-driven, 419 ms with the solves unrolled, 398 ms with the factorisation unrolled too, against 509 ms for the hand-written Schur solver it replaced.

No pivoting. The pattern has to be fixed at compile time and the same in every thread, so rows cannot be swapped on the numbers. Two things make that sound here. The permutation is symmetric, so M's diagonal stays on the diagonal and I/(h*gamma) guarantees every pivot is structurally there and grows without bound as the step shrinks. And Rodas5P is a Rosenbrock-W method: order 5 survives an approximate factorisation, so a badly conditioned pivot costs step-size control rather than correctness -- and the controller is what notices. A pivot that reaches exactly zero leaves an infinity in the factors, the error norm goes to NaN, the step is rejected and the next attempt has a larger 1/(h*gamma) on that diagonal. This is the same stance dense_lu_solver takes on a singular column, for the same reason.

ORDERINGS module-attribute

ORDERINGS = ('amd', 'natural')

MAX_UNROLLED_SUBSTITUTIONS module-attribute

MAX_UNROLLED_SUBSTITUTIONS = 2048

MAX_UNROLLED_UPDATES module-attribute

MAX_UNROLLED_UPDATES = 8192

SparseLULayout dataclass

SparseLULayout(
    n_vars: int,
    order: tuple[int, ...],
    row_ptr: ndarray,
    col_ind: ndarray,
    diag_ptr: ndarray,
)

One CSR image of L + U, in elimination order.

CSR rather than CSC because every one of the three routines that reads this reads it by rows: the up-looking factorisation takes row i and subtracts multiples of the rows above it, the forward substitution is a dot product of row i of L with the solution so far, and the back substitution is the same over row i of U. One row-major image serves all three; CSC would have to be transposed for two of them, and a column-oriented factorisation would still leave the solves wanting rows.

L and U share the image -- L strictly left of the diagonal, U from it rightwards -- because the factorisation is in place and a unit diagonal needs no storage. row_ptr/diag_ptr bracket the two halves of each row.

Indices here are elimination indices: row i of this structure is variable order[i] of the caller's system. :attr:row_origin and :attr:col_origin carry the translation, so the device code never applies a permutation to a vector -- it visits the rows in elimination order and reads and writes the right-hand side where the caller left it.

nnz property

nnz: int

Elements in one trajectory's L + U, i.e. the whole footprint.

row_origin property

row_origin: ndarray

Caller's index of each elimination row.

col_origin property

col_origin: ndarray

Caller's index of the column each slot belongs to.

slot_table

slot_table() -> ndarray

(n_vars, n_vars) of slots, -1 where the structure has nothing.

Source code in modax/_sparse_direct.py
def slot_table(self) -> np.ndarray:
    """``(n_vars, n_vars)`` of slots, ``-1`` where the structure has nothing."""
    table = np.full((self.n_vars, self.n_vars), -1, dtype=np.int64)
    origin = np.asarray(self.order, dtype=np.int64)
    for i in range(self.n_vars):
        start, stop = int(self.row_ptr[i]), int(self.row_ptr[i + 1])
        table[origin[i], origin[self.col_ind[start:stop]]] = np.arange(start, stop)
    return table

factorization_ops

factorization_ops() -> tuple[ndarray, ndarray, ndarray]

The rank-one updates, addressed off the multiplier that drives them.

Up-looking LU visits row i and, for each entry (i, k) left of the diagonal in increasing k, forms the multiplier L[i, k] and subtracts L[i, k] * U[k, k:] from row i. With the pattern fixed, which slot each of those subtractions reads and writes is fixed too, so the whole inner merge -- the part a runtime sparse solver spends its time on, searching row i for the column it has to update -- is resolved here.

The sources need no table at all: they are row k's upper entries, which CSR already holds contiguously, so the inner loop can simply run over pivot_of[p] + 1 .. upper_end[p] and read lu[s] straight. That leaves one index load per multiply-add rather than three, which is the whole cost of the inner loop besides the arithmetic.

Returned as (dst_start, destination, upper_end), all indexed by the multiplier's own slot p. The destinations for one multiplier are contiguous, and in the same order as the sources, so the inner loop walks destination from dst_start[p] alongside the sources.

Source code in modax/_sparse_direct.py
def factorization_ops(self) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
    """The rank-one updates, addressed off the multiplier that drives them.

    Up-looking LU visits row ``i`` and, for each entry ``(i, k)`` left of the
    diagonal in increasing ``k``, forms the multiplier ``L[i, k]`` and
    subtracts ``L[i, k] * U[k, k:]`` from row ``i``. With the pattern fixed,
    *which* slot each of those subtractions reads and writes is fixed too, so
    the whole inner merge -- the part a runtime sparse solver spends its time
    on, searching row ``i`` for the column it has to update -- is resolved
    here.

    The *sources* need no table at all: they are row ``k``'s upper entries,
    which CSR already holds contiguously, so the inner loop can simply run
    over ``pivot_of[p] + 1 .. upper_end[p]`` and read ``lu[s]`` straight.
    That leaves one index load per multiply-add rather than three, which is
    the whole cost of the inner loop besides the arithmetic.

    Returned as ``(dst_start, destination, upper_end)``, all indexed by the
    multiplier's own slot ``p``. The destinations for one multiplier are
    contiguous, and in the same order as the sources, so the inner loop walks
    ``destination`` from ``dst_start[p]`` alongside the sources.
    """
    n = self.n_vars
    dst_start = np.zeros(self.nnz, dtype=np.int64)
    upper_end = np.asarray(self.row_ptr[1:], dtype=np.int64)[
        np.asarray(self.col_ind, dtype=np.int64)
    ]
    dst: list[int] = []
    table = np.full((n, n), -1, dtype=np.int64)
    for i in range(n):
        start, stop = int(self.row_ptr[i]), int(self.row_ptr[i + 1])
        table[i, self.col_ind[start:stop]] = np.arange(start, stop)
    for i in range(n):
        for p in range(int(self.row_ptr[i]), int(self.diag_ptr[i])):
            k = int(self.col_ind[p])
            dst_start[p] = len(dst)
            for q in range(int(self.diag_ptr[k]) + 1, int(self.row_ptr[k + 1])):
                dst.append(int(table[i, int(self.col_ind[q])]))
    return (
        dst_start.astype(np.int32),
        np.asarray(dst, dtype=np.int32),
        upper_end.astype(np.int32),
    )

flop_counts

flop_counts() -> tuple[int, int, int]

(multipliers, updates, substitutions).

The first two are one factorisation; the third is one solve, both substitutions together. These decide whether the device code is emitted as straight-line or as loops, so they are counted from the structure rather than by building what they are counting.

Source code in modax/_sparse_direct.py
def flop_counts(self) -> tuple[int, int, int]:
    """``(multipliers, updates, substitutions)``.

    The first two are one factorisation; the third is one solve, both
    substitutions together. These decide whether the device code is emitted
    as straight-line or as loops, so they are counted from the structure
    rather than by building what they are counting.
    """
    lower = self.diag_ptr - self.row_ptr[:-1]
    upper = self.row_ptr[1:] - self.diag_ptr - 1
    multipliers = int(lower.sum())
    updates = int(upper[self.col_ind[_lower_mask(self)]].sum())
    return multipliers, updates, multipliers + int(upper.sum())

SparseDirectSolver

SparseDirectSolver(layout: SparseLULayout, compressed: CompressedJacobian)

A Rodas5P linear solver compiled for one sparsity pattern.

rodas5P.solve builds this itself from the sparsity it is given, so a caller never holds one; :attr:compressed is the layout it laid out and the layout the kernel then fills, which is what keeps the Enzyme sweeps and the factorisation from disagreeing about where an entry lives.

It pivots nothing, so :attr:ipiv_size is 1 -- the kernel still allocates the array and the device functions still take it, because the dense solver needs it and the two have one shape.

Source code in modax/_sparse_direct.py
def __init__(self, layout: SparseLULayout, compressed: CompressedJacobian):
    self.layout = layout
    self.compressed = compressed
    # Resolved once and handed to both, since it is the expensive part of
    # the analysis and the factorisation is the only thing that needs it.
    ops = layout.factorization_ops()
    self.factorize_local = _make_factorize(layout, ops)
    self.solve_local = _make_solve(layout)

fill_reducing_order

fill_reducing_order(
    pattern: tuple[tuple[int, ...], ...], ordering: str = "amd"
)

A permutation of the variables that keeps L + U small.

AMD, on the symmetrised pattern S + S.T.

Why AMD rather than COLAMD, which is the other obvious candidate: COLAMD orders the columns so that fill stays bounded whatever row permutation partial pivoting later chooses. That is the right objective exactly when there will be pivoting -- and there will not be here, because the pattern is compiled into the kernel and cannot depend on the numbers. COLAMD's permutation is also one-sided, so it moves the diagonal off the diagonal, and this factorisation needs the diagonal precisely where I/(h*gamma) puts it. AMD instead minimises (approximately) the fill of the Cholesky factor of S + S.T, which is the standard bound on the fill of an unpivoted LU of S, and it does so with a symmetric permutation P S P.T that leaves every diagonal entry on the diagonal. It is what UMFPACK and SuperLU use in their "symmetric mode" for the same reasons, and an iteration matrix I/(h*gamma) - J is about as close to structurally symmetric as an unsymmetric matrix gets.

On DISCO-EB's Einstein-Boltzmann Jacobian this returns a perfect elimination order -- nnz(L + U) == nnz(J), not one entry of fill -- and the order it finds is the hand-written Schur solver's: peel each free-streaming multipole hierarchy from its truncated end inwards, then eliminate the densely coupled core last.

SuiteSparse's AMD reaches this through cvxopt, which ships it in a manylinux wheel. It used to come through scikit-sparse's CHOLMOD bindings, which have no wheels and compile against SuiteSparse's headers -- so the default ordering worked only where someone had already run an apt install, and pip install modax failed at that build. The two give the same fill on every pattern measured; where they differ it is in how they break ties between orders that are equally good.

ordering is "amd" or "natural"; see :data:ORDERINGS for what became of CHOLMOD's others.

Source code in modax/_sparse_direct.py
def fill_reducing_order(pattern: tuple[tuple[int, ...], ...], ordering: str = "amd"):
    """A permutation of the variables that keeps ``L + U`` small.

    AMD, on the symmetrised pattern ``S + S.T``.

    Why AMD rather than COLAMD, which is the other obvious candidate: COLAMD
    orders the *columns* so that fill stays bounded whatever row permutation
    partial pivoting later chooses. That is the right objective exactly when
    there will be pivoting -- and there will not be here, because the pattern is
    compiled into the kernel and cannot depend on the numbers. COLAMD's
    permutation is also one-sided, so it moves the diagonal off the diagonal,
    and this factorisation needs the diagonal precisely where ``I/(h*gamma)``
    puts it. AMD instead minimises (approximately) the fill of the Cholesky
    factor of ``S + S.T``, which is the standard bound on the fill of an
    unpivoted ``LU`` of ``S``, and it does so with a *symmetric* permutation
    ``P S P.T`` that leaves every diagonal entry on the diagonal. It is what
    UMFPACK and SuperLU use in their "symmetric mode" for the same reasons, and
    an iteration matrix ``I/(h*gamma) - J`` is about as close to structurally
    symmetric as an unsymmetric matrix gets.

    On DISCO-EB's Einstein-Boltzmann Jacobian this returns a *perfect*
    elimination order -- ``nnz(L + U) == nnz(J)``, not one entry of fill -- and
    the order it finds is the hand-written Schur solver's: peel each
    free-streaming multipole hierarchy from its truncated end inwards, then
    eliminate the densely coupled core last.

    SuiteSparse's AMD reaches this through cvxopt, which ships it in a
    manylinux wheel. It used to come through scikit-sparse's CHOLMOD bindings,
    which have no wheels and compile against SuiteSparse's headers -- so the
    default ordering worked only where someone had already run an `apt install`,
    and `pip install modax` failed at that build. The two give the same fill
    on every pattern measured; where they differ it is in how they break ties
    between orders that are equally good.

    ``ordering`` is ``"amd"`` or ``"natural"``; see :data:`ORDERINGS` for what
    became of CHOLMOD's others.
    """
    if ordering not in ORDERINGS:
        raise ValueError(f"unknown ordering {ordering!r}; expected one of {ORDERINGS}")
    n = len(pattern)
    if ordering == "natural":
        return tuple(range(n))
    try:
        from cvxopt import amd, spmatrix  # ty: ignore[unresolved-import]
    except ImportError as exc:  # pragma: no cover - depends on the environment
        raise ImportError(
            "the sparse direct solver orders its variables with SuiteSparse's "
            "AMD, which reaches it through cvxopt. Install it, or pass "
            "ordering='natural' to skip the ordering entirely."
        ) from exc

    # ``amd.order`` reads the *lower triangle* of a symmetric matrix and wants
    # the pattern alone, so the symmetrised pattern goes in with a unit on every
    # entry. Nothing has to be made positive definite for it, which CHOLMOD's
    # route did need: that reached the permutation through an actual
    # factorisation, and the numbers had to survive it.
    entries = {(r, c) for r, cs in enumerate(pattern) for c in cs}
    entries |= {(c, r) for r, c in entries}
    entries |= {(i, i) for i in range(n)}
    lower = sorted((r, c) for r, c in entries if r >= c)
    matrix = spmatrix(1.0, [r for r, _ in lower], [c for _, c in lower], (n, n))
    return tuple(int(i) for i in amd.order(matrix))

fill_pattern

fill_pattern(pattern: tuple[tuple[int, ...], ...]) -> ndarray

The exact L + U pattern of an unpivoted LU, as a boolean matrix.

Gaussian elimination on the pattern alone: eliminating column k makes every row below it inherit row k's entries to the right of the diagonal. The result is exact -- it is what the numeric factorisation will touch, no more and no less -- which is why the footprint is taken this way rather than by factorising a sample matrix and counting. A sample cannot be exact: a coefficient that happens to vanish for those numbers, or an exact cancellation, drops an entry that another right-hand side needs, and the buffer is then one slot short in a kernel that has no way to say so. It is also cheaper, since it needs neither a plausible matrix nor a device.

Bit-per-entry over the whole matrix, so the analysis is O(n ** 3 / 64) time and O(n ** 2) bits. These are solvers for systems of tens to a few hundred variables (a 200-variable pattern analyses in single-digit milliseconds), so a sparse symbolic factorisation would buy nothing but code.

Source code in modax/_sparse_direct.py
def fill_pattern(pattern: tuple[tuple[int, ...], ...]) -> np.ndarray:
    """The exact ``L + U`` pattern of an unpivoted LU, as a boolean matrix.

    Gaussian elimination on the pattern alone: eliminating column ``k`` makes
    every row below it inherit row ``k``'s entries to the right of the diagonal.
    The result is exact -- it is what the numeric factorisation will touch, no
    more and no less -- which is why the footprint is taken this way rather than
    by factorising a sample matrix and counting. A sample cannot be exact: a
    coefficient that happens to vanish for those numbers, or an exact
    cancellation, drops an entry that another right-hand side needs, and the
    buffer is then one slot short in a kernel that has no way to say so. It is
    also cheaper, since it needs neither a plausible matrix nor a device.

    Bit-per-entry over the whole matrix, so the analysis is ``O(n ** 3 / 64)``
    time and ``O(n ** 2)`` bits. These are solvers for systems of tens to a few
    hundred variables (a 200-variable pattern analyses in single-digit
    milliseconds), so a sparse symbolic factorisation would buy nothing but code.
    """
    n = len(pattern)
    filled = np.zeros((n, n), dtype=bool)
    for row, cols in enumerate(pattern):
        filled[row, list(cols)] = True
    # I/(h*gamma) lands on every diagonal whether or not J has anything there,
    # so the diagonal is part of the pattern being factorised.
    np.fill_diagonal(filled, True)
    for k in range(n):
        below = np.flatnonzero(filled[k + 1 :, k]) + k + 1
        if below.size:
            filled[below, k + 1 :] |= filled[k, k + 1 :]
    return filled

analyse

analyse(
    pattern: tuple[tuple[int, ...], ...], ordering: str = "amd"
) -> SparseLULayout

Order, factorise symbolically, and lay out L + U in CSR.

pattern is the normalised form, one tuple of column indices per row, as modax._sparsity.normalize_sparsity returns it.

Source code in modax/_sparse_direct.py
def analyse(
    pattern: tuple[tuple[int, ...], ...], ordering: str = "amd"
) -> SparseLULayout:
    """Order, factorise symbolically, and lay out ``L + U`` in CSR.

    ``pattern`` is the normalised form, one tuple of column indices per row, as
    [`modax._sparsity.normalize_sparsity`][] returns it.
    """
    n_vars = len(pattern)
    order = fill_reducing_order(pattern, ordering)
    # The elimination happens in the permuted matrix, so the pattern is
    # relabelled *before* the symbolic pass rather than after it: a permutation
    # of the fill is not the fill of the permutation, and the whole point of the
    # ordering is that it changes what fills in.
    forward = np.asarray(order, dtype=np.int64)
    inverse = np.empty(n_vars, dtype=np.int64)
    inverse[forward] = np.arange(n_vars)
    filled = fill_pattern(
        tuple(
            tuple(sorted(int(inverse[c]) for c in pattern[int(forward[i])]))
            for i in range(n_vars)
        )
    )
    row_ptr = np.zeros(n_vars + 1, dtype=np.int32)
    row_ptr[1:] = np.cumsum(filled.sum(axis=1))
    col_ind = np.concatenate([np.flatnonzero(filled[i]) for i in range(n_vars)]).astype(
        np.int32
    )
    diag_ptr = np.array(
        [
            int(row_ptr[i])
            + int(np.searchsorted(col_ind[row_ptr[i] : row_ptr[i + 1]], i))
            for i in range(n_vars)
        ],
        dtype=np.int32,
    )
    return SparseLULayout(
        n_vars=n_vars,
        order=tuple(int(i) for i in order),
        row_ptr=row_ptr,
        col_ind=col_ind,
        diag_ptr=diag_ptr,
    )

compressed_jacobian

compressed_jacobian(
    pattern: tuple[tuple[int, ...], ...], layout: SparseLULayout
) -> CompressedJacobian

The colour-compressed layout whose slots are layout's CSR slots.

Colouring and storage are separate questions and this is where they meet. The colouring comes from the caller's pattern as it always does -- the sweeps cost n_colours + 1 whatever happens to the storage -- but the slot each (row, colour) writes to is the factorisation's, so the AD deposits -J directly into the buffer the factorisation will work in. Fill-in slots belong to no column of J and so are written by nobody, which is exactly what needs_clear in the kernel is for.

A consequence worth knowing: the pattern may now be coloured as tightly as it really is. A structured solver that owned its own buffer had to declare its fill-in in the pattern too, since an entry with no slot had nowhere to go, and that cost colours -- DISCO-EB's 12 rather than 11. Here the fill has its own slots by construction, so only the true nonzeros need colouring.

Source code in modax/_sparse_direct.py
def compressed_jacobian(
    pattern: tuple[tuple[int, ...], ...], layout: SparseLULayout
) -> CompressedJacobian:
    """The colour-compressed layout whose slots are ``layout``'s CSR slots.

    Colouring and storage are separate questions and this is where they meet.
    The colouring comes from the caller's pattern as it always does -- the sweeps
    cost ``n_colours + 1`` whatever happens to the storage -- but the slot each
    ``(row, colour)`` writes to is the factorisation's, so the AD deposits ``-J``
    directly into the buffer the factorisation will work in. Fill-in slots
    belong to no column of ``J`` and so are written by nobody, which is exactly
    what ``needs_clear`` in the kernel is for.

    A consequence worth knowing: the pattern may now be coloured as tightly as it
    really is. A structured solver that owned its own buffer had to declare its
    fill-in in the pattern too, since an entry with no slot had nowhere to go,
    and that cost colours -- DISCO-EB's 12 rather than 11. Here the fill has its
    own slots by construction, so only the true nonzeros need colouring.
    """
    compressed = colour_sparsity(pattern)
    table = layout.slot_table()
    packed = np.full(layout.n_vars * compressed.n_colours, -1, dtype=np.int64)
    for row, cols in enumerate(pattern):
        for col in cols:
            slot = int(table[row, col])
            if slot < 0:
                raise ValueError(
                    f"entry ({row}, {col}) is in the pattern but not in the "
                    "factorised structure, which cannot happen unless the two "
                    "were built from different patterns"
                )
            packed[row * compressed.n_colours + compressed.colour[col]] = slot
    diagonal = np.array([int(table[i, i]) for i in range(layout.n_vars)])
    return replace(
        compressed,
        packed=tuple(int(s) for s in packed),
        packed_diagonal=tuple(int(s) for s in diagonal),
        n_slots=layout.nnz,
    )

sparse_direct_solver

sparse_direct_solver(sparsity, n_vars: int, *, ordering: str = 'amd')

Compile a direct sparse solver for sparsity.

sparsity takes the same spellings solve does: an (n_vars, n_vars) mask, a scipy sparse matrix, or an (nnz, 2) index array. It must be a superset of the true nonzeros of J, as for any pattern the kernel is given; it need not include the fill-in, which is what this computes.

Source code in modax/_sparse_direct.py
def sparse_direct_solver(sparsity, n_vars: int, *, ordering: str = "amd"):
    """Compile a direct sparse solver for ``sparsity``.

    ``sparsity`` takes the same spellings ``solve`` does: an ``(n_vars, n_vars)``
    mask, a scipy sparse matrix, or an ``(nnz, 2)`` index array. It must be a
    superset of the true nonzeros of ``J``, as for any pattern the kernel is
    given; it need *not* include the fill-in, which is what this computes.
    """
    return sparse_direct_solver_for(normalize_sparsity(sparsity, n_vars), ordering)

sparse_direct_solver_for cached

sparse_direct_solver_for(
    pattern: tuple[tuple[int, ...], ...], ordering: str = "amd"
) -> SparseDirectSolver

Compile a direct sparse solver for a normalised pattern.

This is what the kernel builder calls, since it has already normalised what the caller passed as sparsity. Cached on the pattern, so asking twice gets the same compiled device functions and the kernel's own cache hits.

Source code in modax/_sparse_direct.py
@functools.cache
def sparse_direct_solver_for(
    pattern: tuple[tuple[int, ...], ...], ordering: str = "amd"
) -> SparseDirectSolver:
    """Compile a direct sparse solver for a *normalised* pattern.

    This is what the kernel builder calls, since it has already normalised what
    the caller passed as ``sparsity``. Cached on the pattern, so asking twice
    gets the same compiled device functions and the kernel's own cache hits.
    """
    layout = analyse(pattern, ordering)
    return SparseDirectSolver(layout, compressed_jacobian(pattern, layout))

Forward sensitivities

modax._sensitivity

Continuous forward sensitivity analysis for the numba-cuda ensemble solvers.

The solvers reach JAX through an XLA FFI custom call, so the solve itself is opaque to autodiff: there is no traced graph for JAX to differentiate and no practical way to run Enzyme over the kernel (see "Why forward sensitivities" in the README). Derivatives are therefore supplied by a jax.custom_jvp rule that integrates the sensitivity system alongside the state.

Writing S(t) = dy(t)/dtheta for a direction block theta, differentiating y' = f(t, y, p) with respect to theta gives the variational equation

S' = J_y(t) S + J_p(t),    J_y = df/dy,  J_p = df/dtheta,

with J_p = 0 and S(t0) = I for the initial-state block, and J_p = df/dp and S(t0) = 0 for the parameter block. Stacking it under the state gives one joint system

d/dt [y, S] = [f(t, y, p), J_y S + J_p]

which the existing kernels integrate as an ordinary ODE of n_aug = n_vars * (1 + n_sens) variables. Because the rule returns the primal and the tangent together, jax.value_and_grad costs one joint solve rather than a solve for the value plus another for the derivative.

Layout

The augmented state is the state followed by one n_vars-long block per sensitivity direction::

z[j]                        = y[j]
z[n_vars + k * n_vars + r]  = S[r, k]

Direction k runs over the initial-state block first (n_vars directions, present only when y0 is differentiated) and then the parameter block (n_params directions, present only when params is differentiated). Only the blocks JAX actually asks for are integrated, so differentiating with respect to parameters alone does not pay for the n_vars initial-state columns.

J_y S_k + J_p_k comes from numba_enzyme.jvp applied to the user's ode_fn, seeded with a direction over the callback's flattened argument list (y_0 ... y_n-1, t, p_0 ... p_m-1): seeding (S_k, 0, e_k) returns that whole right-hand side in one sweep, and seeding the time argument instead returns df/dt (what Rodas5P already uses). So the parameter derivatives the sensitivity system needs come from the same device function the stiff kernel already builds, with nothing extra from the caller and no Jacobian ever formed.

SensitivitySpec dataclass

SensitivitySpec(
    n_vars: int,
    n_params: int,
    wrt_y0: bool,
    wrt_params: bool,
    error_control: bool = True,
    param_columns: tuple[int, ...] | None = None,
)

Which sensitivity blocks a joint solve carries.

Hashable and used as part of the kernel cache keys, so a solve differentiated only with respect to params compiles (and integrates) a smaller augmented system than one differentiated with respect to both.

error_control decides whether the sensitivity components take part in the step-size error norm; it belongs here rather than in the solve because the error-norm denominator it selects is a compile-time constant of the kernel.

param_seed_columns property

param_seed_columns: tuple[int, ...]

Which parameter each direction seeds, in direction order.

n_aug property

n_aug: int

Size of the joint [y, S] system the kernel actually integrates.

n_error property

n_error: int

Component count the weighted-RMS error norm divides by.

With the sensitivities excluded from step-size control their weights are zero, so they contribute exactly 0.0 to the sum; dividing by n_vars rather than n_aug then reproduces the plain solve's error norm bit for bit, and with it its exact step sequence.

make_tangent cached

make_tangent(ode_fn, n_vars: int, n_params: int)

D f(y, t, p)[u] in one forward sweep, for an arbitrary direction.

The sensitivity right-hand side J_y S_k + J_p_k is a directional derivative: seed (S_k, 0, e_k) and it comes out whole. Assembling it from unit columns instead costs n_vars + 1 sweeps where this costs one.

Source code in modax/_sensitivity.py
@functools.cache
def make_tangent(ode_fn, n_vars: int, n_params: int):
    """``D f(y, t, p)[u]`` in one forward sweep, for an arbitrary direction.

    The sensitivity right-hand side ``J_y S_k + J_p_k`` *is* a directional
    derivative: seed ``(S_k, 0, e_k)`` and it comes out whole.  Assembling it
    from unit columns instead costs ``n_vars + 1`` sweeps where this costs one.
    """
    return jvp(as_cuda_device(ode_fn), signature=_primal_signature(n_vars, n_params))

make_second_tangent cached

make_second_tangent(ode_fn, n_vars: int, n_params: int)

D2 f(y, t, p)[u, v] + D f(y, t, p)[w], forward over forward.

The joint system's Jacobian is the derivative of a right-hand side that already contains first derivatives of ode_fn, so its off-diagonal block holds second derivatives of the original problem: d2f/dy2 contracted with S_k, plus the mixed d2f/dy dp_k. Seeding u = (S_k, 0, e_k) and v = (w, 0, 0) applies that block to w without forming it; seeding v = (0, 1, 0) returns the sensitivity block's dF/dt.

jvp composes with itself, so this is literally the forward sweep of the forward sweep -- numba-enzyme records the chain and emits both markers into one module, rather than trying to differentiate an already-compiled derivative. Composing makes the derivative of the whole tangent map, a function of (x, u), so the call carries a fourth direction w for the inner direction's own variation; every caller here passes zero for it, which leaves the plain bilinear form.

Source code in modax/_sensitivity.py
@functools.cache
def make_second_tangent(ode_fn, n_vars: int, n_params: int):
    """``D2 f(y, t, p)[u, v] + D f(y, t, p)[w]``, forward over forward.

    The joint system's Jacobian is the derivative of a right-hand side that
    already contains first derivatives of ``ode_fn``, so its off-diagonal block
    holds second derivatives of the original problem: ``d2f/dy2`` contracted
    with ``S_k``, plus the mixed ``d2f/dy dp_k``.  Seeding ``u = (S_k, 0, e_k)``
    and ``v = (w, 0, 0)`` applies that block to ``w`` without forming it;
    seeding ``v = (0, 1, 0)`` returns the sensitivity block's ``dF/dt``.

    ``jvp`` composes with itself, so this is literally the forward sweep of the
    forward sweep -- numba-enzyme records the chain and emits both markers into
    one module, rather than trying to differentiate an already-compiled
    derivative.  Composing makes the derivative of the *whole* tangent map, a
    function of ``(x, u)``, so the call carries a fourth direction ``w`` for the
    inner direction's own variation; every caller here passes zero for it,
    which leaves the plain bilinear form.
    """
    signature = _primal_signature(n_vars, n_params)
    return jvp(jvp(as_cuda_device(ode_fn), signature=signature), signature=signature)

seed_table cached

seed_table(spec: SensitivitySpec)

Read-only table every unit and zero direction is a window into.

2 * L zeros with a single 1.0 at L, so the window starting at L - k carries its 1.0 at position k, and any window inside [0, L) is all zeros. One table in constant memory replaces a per-thread one-hot buffer, which could not promote to registers because its store index is dynamic.

Source code in modax/_sensitivity.py
@functools.cache
def seed_table(spec: SensitivitySpec):
    """Read-only table every unit and zero direction is a window into.

    ``2 * L`` zeros with a single ``1.0`` at ``L``, so the window starting at
    ``L - k`` carries its ``1.0`` at position ``k``, and any window inside
    ``[0, L)`` is all zeros.  One table in constant memory replaces a per-thread
    one-hot buffer, which could not promote to registers because its store index
    is dynamic.
    """
    length = max(spec.n_vars, spec.n_params)
    table = np.zeros(2 * length, dtype=np.float64)
    table[length] = 1.0
    return table

make_augmented_local_writer cached

make_augmented_local_writer(ode_fn, spec: SensitivitySpec)

Joint [f, J_y S + J_p] writer for a kernel's thread-local state.

Both kernels run one trajectory per thread, so the whole augmented vector -- the state block and every sensitivity direction -- is written by the thread that owns it, out of and into its own local arrays. The directions stay independent of one another (each is its own forward sweep seeded with (S_k, 0, e_k)); nothing is shared, so there is no race and no synchronisation inside a device function whose callers invoke it under a divergent if running.

Source code in modax/_sensitivity.py
@functools.cache
def make_augmented_local_writer(ode_fn, spec: SensitivitySpec):
    """Joint ``[f, J_y S + J_p]`` writer for a kernel's thread-local state.

    Both kernels run one trajectory per thread, so the whole augmented vector --
    the state block and every sensitivity direction -- is written by the thread
    that owns it, out of and into its own local arrays.  The directions stay
    independent of one another (each is its own forward sweep seeded with
    ``(S_k, 0, e_k)``); nothing is shared, so there is no race and no
    synchronisation inside a device function whose callers invoke it under a
    divergent ``if running``.
    """
    n_vars = spec.n_vars
    n_sens = spec.n_sens
    n_y0_dirs = spec.n_y0_dirs
    n_params = spec.n_params
    param_cols = spec.param_seed_columns
    length = max(n_vars, n_params)
    fn_device = as_cuda_device(ode_fn)
    tangent_of = make_tangent(ode_fn, n_vars, n_params)
    seeds = seed_table(spec)

    @cuda.jit(device=True)
    def write_vector(z_row, t, p_row, out):
        seed = cuda.const.array_like(seeds)
        values = fn_device(z_row, t, p_row)
        for j in range(n_vars):
            out[j] = values[j]

        tangent = cuda.local.array(n_vars, types.float64)
        for k in range(n_sens):
            base = n_vars + k * n_vars
            start = 0 if k < n_y0_dirs else length - param_cols[k - n_y0_dirs]
            tangent_of(
                tangent,
                z_row,
                t,
                p_row,
                z_row[base : base + n_vars],
                0.0,
                seed[start : start + n_params],
            )
            for r in range(n_vars):
                out[base + r] = tangent[r]

    return write_vector

clear_caches

clear_caches() -> None

Drop the cached device writers and Enzyme derivatives.

Each (ode_fn, spec) pair compiles its own writer, and nothing releases them: the module-level caches hold them for the process's lifetime.

Source code in modax/_sensitivity.py
def clear_caches() -> None:
    """Drop the cached device writers and Enzyme derivatives.

    Each ``(ode_fn, spec)`` pair compiles its own writer, and nothing releases
    them: the module-level caches hold them for the process's lifetime.
    """
    make_tangent.cache_clear()
    make_second_tangent.cache_clear()
    seed_table.cache_clear()
    make_augmented_local_writer.cache_clear()

augmented_y0

augmented_y0(y0_arr, spec: SensitivitySpec)

(n, n_vars) initial state -> (n, n_aug) joint initial state.

S(t0) is the identity for the initial-state block (dy0/dy0) and zero for the parameter block (the initial state does not depend on p).

Source code in modax/_sensitivity.py
def augmented_y0(y0_arr, spec: SensitivitySpec):
    """``(n, n_vars)`` initial state -> ``(n, n_aug)`` joint initial state.

    ``S(t0)`` is the identity for the initial-state block (``dy0/dy0``) and zero
    for the parameter block (the initial state does not depend on ``p``).
    """
    n = y0_arr.shape[0]
    blocks = [y0_arr]
    if spec.wrt_y0:
        eye = jnp.eye(spec.n_vars, dtype=y0_arr.dtype).reshape(1, -1)
        blocks.append(jnp.broadcast_to(eye, (n, spec.n_vars * spec.n_vars)))
    if spec.wrt_params:
        blocks.append(jnp.zeros((n, spec.n_vars * spec.n_param_dirs), y0_arr.dtype))
    return jnp.concatenate(blocks, axis=1)

augmented_error_weights

augmented_error_weights(weights_arr, spec: SensitivitySpec)

(n, n_vars) error weights -> (n, n_aug).

A zero weight drops a component from the step-size error norm, so spec.error_control=False makes the joint solve take exactly the step sequence the plain solve would -- the value from jax.value_and_grad is then bit-identical to the value from a plain call, and the sensitivities ride along on steps chosen for the state alone. That is not the default: nothing then ties the sensitivities' accuracy to rtol, and a stiff solver taking large steps on an easy state can be badly wrong about them.

Source code in modax/_sensitivity.py
def augmented_error_weights(weights_arr, spec: SensitivitySpec):
    """``(n, n_vars)`` error weights -> ``(n, n_aug)``.

    A zero weight drops a component from the step-size error norm, so
    ``spec.error_control=False`` makes the joint solve take exactly the step
    sequence the plain solve would -- the value from ``jax.value_and_grad`` is
    then bit-identical to the value from a plain call, and the sensitivities
    ride along on steps chosen for the state alone.  That is not the default:
    nothing then ties the sensitivities' accuracy to ``rtol``, and a stiff
    solver taking large steps on an easy state can be badly wrong about them.
    """
    tail = np.ones if spec.error_control else np.zeros
    return np.concatenate(
        [
            weights_arr,
            tail((weights_arr.shape[0], spec.n_aug - spec.n_vars), dtype=np.float64),
        ],
        axis=1,
    )

split_augmented

split_augmented(hist, spec: SensitivitySpec)

(n, n_save, n_aug) -> state (n, n_save, n_vars) and S (n, n_save, n_vars, n_sens).

Source code in modax/_sensitivity.py
def split_augmented(hist, spec: SensitivitySpec):
    """``(n, n_save, n_aug)`` -> state ``(n, n_save, n_vars)`` and ``S`` ``(n, n_save, n_vars, n_sens)``."""
    n, n_save, _ = hist.shape
    state = hist[:, :, : spec.n_vars]
    sens = hist[:, :, spec.n_vars :].reshape(n, n_save, spec.n_sens, spec.n_vars)
    return state, jnp.swapaxes(sens, 2, 3)

make_sensitivity_solver

make_sensitivity_solver(
    primal_solver,
    joint_solver_for,
    n_vars: int,
    n_params: int,
    return_stats: bool,
    sens_error_control: bool,
    param_columns: tuple[int, ...] | None = None,
)

Attach a forward-sensitivity JVP rule to one solver implementation.

primal_solver(y0, t_span, params) runs the plain solve. joint_solver_for(spec) returns the corresponding solver for the joint [y, S] system described by spec; its output has the ordinary solver shape, with n_aug components in place of n_vars. Both are expected to be already wrapped for vmap: the rule has to sit outside custom_vmap, whose JVP path traces to a jaxpr and so instantiates every symbolic zero, which would hide exactly the information the rule uses to decide which sensitivity blocks to integrate.

The rule fires only under a JAX differentiation transform, so an undifferentiated call pays nothing for sensitivities. The tangent is formed by contracting S -- which depends only on the primal inputs -- with the input tangents. That contraction is linear in the tangents, so JAX can transpose it: jax.grad and jax.jacrev work off the same rule as jax.jvp and jax.jacfwd, with no adjoint solve.

Source code in modax/_sensitivity.py
def make_sensitivity_solver(
    primal_solver,
    joint_solver_for,
    n_vars: int,
    n_params: int,
    return_stats: bool,
    sens_error_control: bool,
    param_columns: tuple[int, ...] | None = None,
):
    """Attach a forward-sensitivity JVP rule to one solver implementation.

    ``primal_solver(y0, t_span, params)`` runs the plain solve.
    ``joint_solver_for(spec)`` returns the corresponding solver for the joint
    ``[y, S]`` system described by ``spec``; its output has the ordinary solver
    shape, with ``n_aug`` components in place of ``n_vars``.  Both are expected
    to be already wrapped for ``vmap``: the rule has to sit *outside*
    ``custom_vmap``, whose JVP path traces to a jaxpr and so instantiates every
    symbolic zero, which would hide exactly the information the rule uses to
    decide which sensitivity blocks to integrate.

    The rule fires only under a JAX differentiation transform, so an
    undifferentiated call pays nothing for sensitivities.  The tangent is formed
    by contracting ``S`` -- which depends only on the primal inputs -- with the
    input tangents.  That contraction is linear in the tangents, so JAX can
    transpose it: ``jax.grad`` and ``jax.jacrev`` work off the same rule as
    ``jax.jvp`` and ``jax.jacfwd``, with no adjoint solve.
    """

    @jax.custom_jvp
    def solve_fn(y0, t_span, params):
        return primal_solver(y0, t_span, params)

    @functools.partial(solve_fn.defjvp, symbolic_zeros=True)
    def solve_jvp(primals, tangents):
        y0, t_span, params = primals
        dy0, dt_span, dparams = tangents

        if not isinstance(dt_span, SymbolicZero):
            raise NotImplementedError(
                "differentiating a modax solve with respect to t_span is not "
                "supported; wrap the save times in jax.lax.stop_gradient, or "
                "differentiate with respect to y0 and/or params only"
            )

        wrt_y0 = not isinstance(dy0, SymbolicZero)
        wrt_params = not isinstance(dparams, SymbolicZero)
        if not (wrt_y0 or wrt_params):
            out = primal_solver(y0, t_span, params)
            return out, jax.tree_util.tree_map(_zero_tangent, out)

        spec = SensitivitySpec(
            n_vars, n_params, wrt_y0, wrt_params, sens_error_control, param_columns
        )
        out = joint_solver_for(spec)(y0, t_span, params)
        hist, stats = out if return_stats else (out, None)
        state, sens = split_augmented(hist, spec)

        axis_size = state.shape[0]
        tangent = jnp.zeros_like(state)
        offset = 0
        if wrt_y0:
            tangent += jnp.einsum(
                "nsij,nj->nsi",
                sens[..., : spec.n_y0_dirs],
                _match_shape(dy0, axis_size, "y0"),
            )
            offset = spec.n_y0_dirs
        if wrt_params:
            dp = _match_shape(dparams, axis_size, "params")
            if spec.param_columns is not None:
                # Only the chosen columns were integrated, so contract against
                # just those. The transpose of this gather is a scatter, which
                # is what makes jax.grad return a cotangent of the full width
                # with zeros where no sensitivity was carried.
                dp = dp[:, jnp.asarray(spec.param_columns)]
            tangent += jnp.einsum("nsik,nk->nsi", sens[..., offset:], dp)
        if stats is None:
            return state, tangent
        # Step counters are integers and carry no derivative, but JAX still
        # wants a tangent of the right (float0) type for every output.
        return (state, stats), (tangent, jax.tree_util.tree_map(_zero_tangent, stats))

    return solve_fn

Host-side kernel support

modax._numba_common

Shared host-side helpers for numba-cuda custom-kernel solvers.

HOOK_OUT_ARGTYPE module-attribute

HOOK_OUT_ARGTYPE = _F64_2D

Per-trajectory rows a save hook accumulates into ((n, hook_size)).

build_error_weights

build_error_weights(error_weights, n: int, n_vars: int) -> ndarray

Broadcast a user error_weights argument to a (n, n_vars) array.

None yields all-ones (every component weighted equally); a 1-D array of length n_vars is broadcast across trajectories; a 2-D (n, n_vars) array is used as-is. This array is copied to the device and read per component as the weight argument of the kernel's error-contribution device function.

Source code in modax/_numba_common.py
def build_error_weights(error_weights, n: int, n_vars: int) -> np.ndarray:
    """Broadcast a user ``error_weights`` argument to a ``(n, n_vars)`` array.

    ``None`` yields all-ones (every component weighted equally); a 1-D array of
    length ``n_vars`` is broadcast across trajectories; a 2-D ``(n, n_vars)``
    array is used as-is. This array is copied to the device and read per
    component as the ``weight`` argument of the kernel's error-contribution
    device function.
    """
    if error_weights is None:
        return np.ones((n, n_vars), dtype=np.float64)
    weights = np.asarray(error_weights, dtype=np.float64)
    if weights.ndim == 1:
        weights = np.broadcast_to(weights, (n, n_vars))
    return np.ascontiguousarray(weights, dtype=np.float64)

initial_step

initial_step(first_step)

The dt0 scalar the kernels take, from a user first_step.

dt0 is a launch-time scalar, so it cannot be derived here from the save times: under jit (and under the custom_vmap rule, which traces unconditionally) those are tracers. A caller that pins no first step passes the non-positive sentinel instead, and each kernel takes 1e-6 of its own integration window off the times array it already reads.

Source code in modax/_numba_common.py
def initial_step(first_step):
    """The ``dt0`` scalar the kernels take, from a user ``first_step``.

    ``dt0`` is a launch-time scalar, so it cannot be derived here from the save
    times: under ``jit`` (and under the ``custom_vmap`` rule, which traces
    unconditionally) those are tracers. A caller that pins no first step passes
    the non-positive sentinel instead, and each kernel takes 1e-6 of its own
    integration window off the ``times`` array it already reads.
    """
    return np.float64(0.0 if first_step is None else first_step)

ensemble_ffi_call

ensemble_ffi_call(
    launch,
    arrays,
    scratch_specs,
    *,
    n: int,
    n_vars: int,
    n_save: int,
    dt0,
    rtol,
    atol,
    max_steps,
    n_save_hist: int | None = None,
    hook_size: int | None = None,
)

Launch a solver kernel from JAX and return its solution outputs.

arrays are the array inputs (y0, times, params, error weights) in kernel order and scratch_specs describes the kernel's scratch arrays, which XLA allocates as extra outputs. Returns (hist, accepted, rejected, loop).

n_save_hist is the history's time extent when it differs from the number of save times (a kernel that keeps only the final state passes 1), and hook_size the width of the save hook's per-trajectory accumulator rows, which the kernel takes as one more output right after the counters; the call then returns (hist, accepted, rejected, loop, hook_out).

Source code in modax/_numba_common.py
def ensemble_ffi_call(
    launch,
    arrays,
    scratch_specs,
    *,
    n: int,
    n_vars: int,
    n_save: int,
    dt0,
    rtol,
    atol,
    max_steps,
    n_save_hist: int | None = None,
    hook_size: int | None = None,
):
    """Launch a solver kernel from JAX and return its solution outputs.

    ``arrays`` are the array inputs (y0, times, params, error weights) in kernel
    order and ``scratch_specs`` describes the kernel's scratch arrays, which XLA
    allocates as extra outputs. Returns ``(hist, accepted, rejected, loop)``.

    ``n_save_hist`` is the history's time extent when it differs from the number
    of save times (a kernel that keeps only the final state passes 1), and
    ``hook_size`` the width of the save hook's per-trajectory accumulator rows,
    which the kernel takes as one more output right after the counters; the
    call then returns ``(hist, accepted, rejected, loop, hook_out)``.
    """
    int_spec = jax.ShapeDtypeStruct((n,), jnp.int32)
    n_hist = n_save if n_save_hist is None else n_save_hist
    output_specs = (
        jax.ShapeDtypeStruct((n, n_hist, n_vars), jnp.float64),
        int_spec,
        int_spec,
        int_spec,
    )
    if hook_size is not None:
        output_specs += (jax.ShapeDtypeStruct((n, hook_size), jnp.float64),)
    output_specs += tuple(scratch_specs)
    result = ffi_abi_call(
        launch,
        arrays,
        output_specs,
        input_kinds=SOLVER_INPUT_KINDS,
        scalar_f64_values=(dt0, rtol, atol),
        scalar_i32_values=(max_steps,),
    )
    return result[:4] if hook_size is None else result[:5]

make_cuda_local_vector_writer

make_cuda_local_vector_writer(fn, n_vars: int)

A vector writer over one trajectory's own thread-local arrays.

Both kernels run one trajectory per thread, so there is no stripe to share: the thread owning the trajectory writes the whole vector, and both y_row and out are its own local memory rather than rows of a global scratch array.

Source code in modax/_numba_common.py
def make_cuda_local_vector_writer(fn, n_vars: int):
    """A vector writer over one trajectory's own thread-local arrays.

    Both kernels run one trajectory per thread, so there is no stripe to share:
    the thread owning the trajectory writes the whole vector, and both
    ``y_row`` and ``out`` are its own local memory rather than rows of a global
    scratch array.
    """
    fn_device = as_cuda_device(fn)

    @cuda.jit(device=True)
    def write_vector(y_row, t, p_row, out):
        values = fn_device(y_row, t, p_row)
        for j in range(n_vars):
            out[j] = values[j]

    return write_vector

The JAX glue

modax._jax_common

Shared scaffolding for exposing the numba-cuda ensemble solvers to JAX.

normalize_y0_params

normalize_y0_params(y0, params)

Broadcast y0 / params to a consistent (N, …) ensemble layout.

Accepts either 1-D ((n_vars,) / (n_params,)) or 2-D ((N, n_vars) / (N, n_params)) inputs and returns 2-D arrays with a common leading axis, so every numba-cuda solver shares one calling convention.

Source code in modax/_jax_common.py
def normalize_y0_params(y0, params):
    """Broadcast ``y0`` / ``params`` to a consistent ``(N, …)`` ensemble layout.

    Accepts either 1-D (``(n_vars,)`` / ``(n_params,)``) or 2-D
    (``(N, n_vars)`` / ``(N, n_params)``) inputs and returns 2-D arrays with a
    common leading axis, so every numba-cuda solver shares one calling
    convention.
    """
    y0_arr = jnp.asarray(y0, dtype=jnp.float64)
    params_arr = jnp.asarray(params, dtype=jnp.float64)

    if y0_arr.ndim not in (1, 2) or params_arr.ndim not in (1, 2):
        raise ValueError(
            "y0 must have shape (n_vars,) or (N, n_vars) and params shape "
            f"(n_params,) or (N, n_params); got y0.shape={y0_arr.shape} and "
            f"params.shape={params_arr.shape}"
        )
    if y0_arr.ndim == 2:
        n = y0_arr.shape[0]
        if params_arr.ndim == 2 and params_arr.shape[0] != n:
            raise ValueError(
                "params must have shape (n_params,) or (N, n_params) when y0 has "
                f"shape (N, n_vars); got y0.shape={y0_arr.shape} and "
                f"params.shape={params_arr.shape}"
            )
    elif params_arr.ndim == 2:
        n = params_arr.shape[0]
    else:
        n = 1

    if y0_arr.ndim == 1:
        y0_arr = jnp.broadcast_to(y0_arr, (n, y0_arr.shape[0]))
    if params_arr.ndim == 1:
        params_arr = jnp.broadcast_to(params_arr, (n, params_arr.shape[0]))
    return y0_arr, params_arr, n, y0_arr.shape[1]

make_custom_vmap_solver

make_custom_vmap_solver(solve_impl: Callable, *, return_stats: bool)

Wrap a solver implementation so outer jax.vmap becomes one ensemble call.

solve_impl must accept (y0, t_span, params) and return the normal public solver result for those arrays. The custom batching rule supports vmapping scalar solves over y0 and/or params and lowers that vmap to a single native ensemble solve with a leading trajectory axis. Every stats field the kernels emit is a per-trajectory counter, so the stats pytree only needs a trailing solve axis added.

Source code in modax/_jax_common.py
def make_custom_vmap_solver(solve_impl: Callable, *, return_stats: bool):
    """Wrap a solver implementation so outer ``jax.vmap`` becomes one ensemble call.

    ``solve_impl`` must accept ``(y0, t_span, params)`` and return the normal
    public solver result for those arrays.  The custom batching rule supports
    vmapping scalar solves over ``y0`` and/or ``params`` and lowers that vmap to
    a single native ensemble solve with a leading trajectory axis.  Every stats
    field the kernels emit is a per-trajectory counter, so the stats pytree
    only needs a trailing solve axis added.
    """

    @custom_vmap
    def _solve(y0, t_span, params):
        return solve_impl(y0, t_span, params)

    @_solve.def_vmap
    def _solve_vmap(axis_size, in_batched, y0, t_span, params):
        y0_batched, t_span_batched, params_batched = in_batched
        if t_span_batched:
            t_span_arr = jnp.asarray(t_span)
            if t_span_arr.ndim != 2:
                raise NotImplementedError(
                    "vmap over nested t_span values is not supported; use a shared "
                    "t_span and vmap over y0 and/or params, or call the solver directly."
                )
            # JAX can mark closed-over constant save times as batched inside
            # a larger vmapped function.  Treat that as a shared time grid.
            t_span = t_span_arr[0]

        y0_arr = _broadcast_for_vmap(y0, y0_batched, axis_size, "y0")
        params_arr = _broadcast_for_vmap(params, params_batched, axis_size, "params")
        result = solve_impl(y0_arr, t_span, params_arr)

        if not return_stats:
            return result[:, None, :, :], True

        sol, stats = result
        stats_out = jax.tree_util.tree_map(lambda x: x[:, None], stats)
        stats_batched = jax.tree_util.tree_map(lambda _: True, stats_out)
        return (sol[:, None, :, :], stats_out), (True, stats_batched)

    return _solve

The XLA FFI shim

modax._jax_numba_custom_call

JAX custom-call bridge for launching numba-cuda kernels.

The public surface in this module is intentionally small: compile a CUDA kernel, register a typed XLA FFI launcher for it, and call it from JAX. The C++ FFI shim is built lazily into /tmp so the project can keep using plain uv run python without a package build step.

CudaLaunch dataclass

CudaLaunch(
    function: int,
    grid: tuple[int, int, int],
    block: tuple[int, int, int],
    shared_mem: int = 0,
)

Compiled CUDA kernel launch metadata for an XLA FFI call.

register_target

register_target() -> None

Register the generic CUDA launcher with JAX once per process.

Source code in modax/_jax_numba_custom_call.py
def register_target() -> None:
    """Register the generic CUDA launcher with JAX once per process."""

    global _LOADED_LIB, _REGISTERED
    if _REGISTERED:
        return
    so_path = _build_bridge()
    _LOADED_LIB = ctypes.CDLL(str(so_path))
    symbol = getattr(_LOADED_LIB, _TARGET_NAME)
    address = ctypes.cast(symbol, ctypes.c_void_p).value
    # A symbol resolved out of a loaded library always has one; the
    # `None` is what `c_void_p` carries for a null pointer.
    assert address is not None
    capsule = _pycapsule_new(address)
    jax.ffi.register_ffi_target(_TARGET_NAME, capsule, platform="CUDA", api_version=1)
    _REGISTERED = True

compile_kernel

compile_kernel(kernel: Any, argtypes: Sequence[Any]) -> int

Compile a cuda.jit kernel and return its CUfunction pointer.

numba-cuda-mlir exposes no get_cufunc() -- its MLIRLibrary only offers the textual IR. The linked cubin and the mangled entry name are on the compile result's metadata instead, so load the module through the driver and look the function up by name.

Source code in modax/_jax_numba_custom_call.py
def compile_kernel(kernel: Any, argtypes: Sequence[Any]) -> int:
    """Compile a ``cuda.jit`` kernel and return its ``CUfunction`` pointer.

    numba-cuda-mlir exposes no ``get_cufunc()`` -- its ``MLIRLibrary`` only
    offers the textual IR.  The linked cubin and the mangled entry name are on
    the compile result's metadata instead, so load the module through the driver
    and look the function up by name.
    """

    cres = kernel.compile(tuple(argtypes)).cres
    cubin = cres.metadata["cubin"]
    func_name = cres.metadata["func_name"]

    libcuda = ctypes.CDLL("libcuda.so.1")
    module = ctypes.c_void_p()
    err = libcuda.cuModuleLoadData(ctypes.byref(module), ctypes.c_char_p(cubin))
    if err != 0:
        raise RuntimeError(f"cuModuleLoadData failed with CUDA driver error {err}")
    _LOADED_MODULES.append(module)

    function = ctypes.c_void_p()
    err = libcuda.cuModuleGetFunction(
        ctypes.byref(function), module, func_name.encode()
    )
    if err != 0:
        raise RuntimeError(f"cuModuleGetFunction failed with CUDA driver error {err}")
    if function.value is None:
        raise RuntimeError("cuModuleGetFunction returned a null function pointer")
    return int(function.value)

make_launch

make_launch(
    kernel: Any,
    argtypes: Sequence[Any],
    *,
    grid: int | Sequence[int],
    block: int | Sequence[int],
    shared_mem: int = 0,
) -> CudaLaunch
Source code in modax/_jax_numba_custom_call.py
def make_launch(
    kernel: Any,
    argtypes: Sequence[Any],
    *,
    grid: int | Sequence[int],
    block: int | Sequence[int],
    shared_mem: int = 0,
) -> CudaLaunch:
    return CudaLaunch(
        function=compile_kernel(kernel, argtypes),
        grid=_as_3d(grid),
        block=_as_3d(block),
        shared_mem=int(shared_mem),
    )

ffi_abi_call

ffi_abi_call(
    launch: CudaLaunch,
    inputs: Sequence[Any],
    output_specs: Sequence[ShapeDtypeStruct],
    *,
    input_kinds: Sequence[int],
    scalar_f64_values: Sequence[float] = (),
    scalar_i32_values: Sequence[int] = (),
) -> tuple[Any, ...]

Launch a Numba CUDA kernel using Numba's normal array/scalar ABI.

Outputs are always arrays, so only the input kinds need spelling out.

Source code in modax/_jax_numba_custom_call.py
def ffi_abi_call(
    launch: CudaLaunch,
    inputs: Sequence[Any],
    output_specs: Sequence[jax.ShapeDtypeStruct],
    *,
    input_kinds: Sequence[int],
    scalar_f64_values: Sequence[float] = (),
    scalar_i32_values: Sequence[int] = (),
) -> tuple[Any, ...]:
    """Launch a Numba CUDA kernel using Numba's normal array/scalar ABI.

    Outputs are always arrays, so only the input kinds need spelling out.
    """

    register_target()
    attrs = {
        "function": np.int64(launch.function),
        "grid_x": np.int64(launch.grid[0]),
        "grid_y": np.int64(launch.grid[1]),
        "grid_z": np.int64(launch.grid[2]),
        "block_x": np.int64(launch.block[0]),
        "block_y": np.int64(launch.block[1]),
        "block_z": np.int64(launch.block[2]),
        "shared_mem": np.int64(launch.shared_mem),
        "arg_kinds": np.asarray(
            tuple(input_kinds) + (ABI_ARRAY,) * len(output_specs), dtype=np.int64
        ),
        "scalar_f64_values": np.asarray(tuple(scalar_f64_values), dtype=np.float64),
        "scalar_i32_values": np.asarray(tuple(scalar_i32_values), dtype=np.int32),
    }
    result = jax.ffi.ffi_call(
        _TARGET_NAME,
        tuple(output_specs),
        has_side_effect=False,
        custom_call_api_version=_CUSTOM_CALL_API_VERSION,
    )(*inputs, **attrs)
    if not isinstance(result, tuple):
        return (result,)
    return result

Generated device source

modax._codegen

Compile generated straight-line device source into a cuda.jit function.

Two parts of the Rodas5P kernel are emitted as source rather than written: the sparse factorisation and triangular solves, with every slot a literal (modax._sparse_direct), and the Jacobian writer, with every colour's seed row a literal (modax.rodas5P). Both are load-bearing -- see the measurements in AGENTS.md -- and both need the same three things done to the text they produce, which is what this module does once.

compile_device_source

compile_device_source(
    name: str, lines: list[str], namespace: dict | None = None
)

Exec lines defining name and return it as a device function.

The source is registered with linecache under a filename of its own, so a numba typing error inside it points at the offending line rather than at nothing, and the text stays readable from a debugger as fn._generated_source. namespace supplies whatever the generated body refers to by name.

Source code in modax/_codegen.py
def compile_device_source(name: str, lines: list[str], namespace: dict | None = None):
    """Exec ``lines`` defining ``name`` and return it as a device function.

    The source is registered with ``linecache`` under a filename of its own, so
    a numba typing error inside it points at the offending line rather than at
    nothing, and the text stays readable from a debugger as
    ``fn._generated_source``. ``namespace`` supplies whatever the generated body
    refers to by name.
    """
    source = "\n".join(lines) + "\n"
    filename = f"<modax generated {name} {next(_SOURCES)}>"
    linecache.cache[filename] = (len(source), None, source.splitlines(True), filename)
    scope: dict = {}
    exec(compile(source, filename, "exec"), dict(namespace or {}), scope)  # noqa: S102
    generated = scope[name]
    generated._generated_source = source
    return cuda.jit(device=True)(generated)