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.
slot ¶
Where entry (row, col) lives.
Source code in modax/_sparsity.py
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
seed_table ¶
(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
dense_jacobian ¶
dense_jacobian(n_vars: int) -> CompressedJacobian
Every column its own colour: the row-major dense matrix.
normalize_sparsity ¶
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
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
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:
- Order. The pattern is symmetrised and handed to SuiteSparse's AMD, which
returns a fill-reducing permutation (see
fill_reducing_orderfor why AMD and not COLAMD). - Factor symbolically. The exact pattern of
L + Ufor the permuted matrix, fill-in included, falls out of a pure-pattern Gaussian elimination (fill_pattern) -- no numbers, no device, no sample matrix. - Lay out. That pattern becomes one CSR image of
L + U, and the(row, col) -> slotmap it defines becomes theCompressedJacobianthe AD sweeps write into. The buffer is exactlynnz(L + U)elements: the symbolic pass is the memory footprint. - Compile. The factorisation and the two triangular solves become
cuda.jit(device=True)functions, in the samefactorize_local(lu, ipiv)/solve_local(lu, ipiv, rhs)shapemodax.rodas5P.dense_lu_solverhas, 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.
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.
slot_table ¶
slot_table() -> ndarray
(n_vars, n_vars) of slots, -1 where the structure has nothing.
Source code in modax/_sparse_direct.py
factorization_ops ¶
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
flop_counts ¶
(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
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
fill_reducing_order ¶
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
fill_pattern ¶
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
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
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
sparse_direct_solver ¶
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
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
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
¶
Which parameter each direction seeds, in direction order.
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
¶
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
make_second_tangent
cached
¶
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
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
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
clear_caches ¶
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
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
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
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
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
331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 | |
Host-side kernel support¶
modax._numba_common ¶
Shared host-side helpers for numba-cuda custom-kernel solvers.
HOOK_OUT_ARGTYPE
module-attribute
¶
Per-trajectory rows a save hook accumulates into ((n, hook_size)).
build_error_weights ¶
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
initial_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
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
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
The JAX glue¶
modax._jax_common ¶
Shared scaffolding for exposing the numba-cuda ensemble solvers to JAX.
normalize_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
make_custom_vmap_solver ¶
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
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 the generic CUDA launcher with JAX once per process.
Source code in modax/_jax_numba_custom_call.py
compile_kernel ¶
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
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
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
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 ¶
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.