Skip to content

Solvers

Both solvers expose the same solve(...) entry point. Importing modax enables JAX float64.

Rodas5P

modax.rodas5P.solve

solve(
    ode_fn,
    y0,
    t_span,
    params,
    *,
    rtol=1e-08,
    atol=1e-10,
    first_step=None,
    max_steps=100000,
    return_stats=False,
    error_weights=None,
    pcoeff=0.0,
    icoeff=1.0,
    dcoeff=0.0,
    lu_precision: str = "fp32",
    trajectories_per_block=None,
    sens_error_control=True,
    sens_param_columns=None,
    sparsity=None,
    ordering="amd",
    tf_index=None,
    max_registers=None,
    array_rhs=None,
    save_hook=None,
    hook_size: int = 0,
    save_history: bool = True,
)

JAX-callable Rodas5 custom-kernel solve.

Only the right-hand side is supplied. The Jacobian df/dy, and the partial time derivative df/dt that a non-autonomous system needs to keep fifth-order accuracy, are both differentiated out of ode_fn with Enzyme, so a non-autonomous problem needs nothing extra from the caller.

lu_precision ("fp32" or "fp64") selects the precision of the per-step LU factorisation and triangular solves. The state, right-hand side, Jacobian and error estimate are always float64; because the Rosenbrock--Wanner order conditions retain full order under an approximate Jacobian, the "fp32" default does not reduce the method's order while halving the LU shared-memory footprint and exploiting FP32 throughput. "fp64" is available for ill-conditioned systems where the FP32 factorisation degrades step-size control.

first_step pins the initial step size; the default None (like any non-positive value) lets the kernel start from 1e-6 of the integration window.

error_weights is an optional per-component weight array, shape (n_vars,) or (N, n_vars), applied in the weighted RMS step-size error norm; a weight of 0 excludes that component from step-size control.

pcoeff/icoeff/dcoeff are the PID step-controller gains; the default (0, 1, 0) is the classic I-controller.

The solve is an XLA custom call into the numba-cuda kernel, so it carries a jax.custom_jvp rule rather than being differentiated by XLA: asking for a derivative integrates the continuous forward-sensitivity system alongside the state (see modax/_sensitivity.py). jax.jvp, jax.jacfwd, jax.grad, jax.jacrev and jax.value_and_grad all work with respect to y0 and params; t_span is not differentiable. An undifferentiated call runs the plain kernel and pays nothing.

The joint system has n_vars * (1 + n_sens) components, but its iteration matrix is not factorised whole: it is block lower triangular with the same M0 = I/(h*gamma) - J_y on every diagonal block, so one n_vars factorisation serves the state and every sensitivity column and the coupling is a forward substitution. What grows with n_sens is the thread's own working set -- ten stage vectors of the augmented state, in local memory -- rather than the block's shared budget, which carries only the n_vars matrix. This is still a solver for problems with few parameters relative to the state dimension.

sens_param_columns restricts the parameter sensitivities to the columns named, rather than carrying one block per parameter. The joint system is n_vars * (1 + n_sens) components and each direction costs a second-order Enzyme sweep per stage, so this is the difference between paying for the parameters you want and paying for the whole params row -- which matters most when that row also carries things that are not parameters at all, such as integration bounds or an integer selecting a table. The default None carries every column, as before.

sens_error_control decides whether the sensitivity components take part in the step-size error norm. The default True controls them to the same rtol/atol as the state, so the gradient is as accurate as the value. False drops them from the norm, which makes the joint solve take exactly the step sequence the plain solve takes -- the value then matches a plain call bit for bit -- at the cost of nothing tying the sensitivities' accuracy to rtol.

The joint Jacobian's lower-left block d(J_y S + J_p)/dy is a second derivative of ode_fn, and it is formed rather than dropped: Enzyme's forward-over-forward sweep applies it to the state increment without ever materialising the matrix. Dropping it would be legitimate under the W property -- order 5 survives an approximate Jacobian -- but not cheap: the error constant it costs was measured at 200x the steps on a right-hand side bilinear in state and parameters, which is most reaction networks.

sparsity is where a structured problem pays off, and it is the only thing a caller has to supply to get one. Pass an (n_vars, n_vars) mask, a scipy sparse matrix, or an (nnz, 2) array of indices, and two things follow. The Jacobian costs one Enzyme sweep per colour of the pattern's column intersection graph rather than one per column, since columns sharing no row can be seeded together and the pattern says which output component belongs to which (modax._sparsity). And the iteration matrix is ordered, factorised symbolically and given an in-kernel sparse LU and sparse triangular solves compiled for that exact structure (modax._sparse_direct). The default -- no pattern -- colours every column apart and factorises densely, which is the same mechanism at its uninformative end rather than a second path.

The pattern must be a superset of the true nonzeros. Colouring a superset only costs sweeps; colouring a subset silently corrupts the entries where two columns of a group do overlap after all. It need not include the factorisation's fill-in, which the symbolic pass works out and gives slots of its own.

ordering picks the fill-reducing permutation: "amd" by default, SuiteSparse's approximate minimum degree out of cvxopt, or "natural" to skip the ordering. It is ignored without a pattern.

Forward sensitivities work with either solver: the joint iteration matrix is block lower triangular with the same M0 on every diagonal block, so only the n_vars block is ever factorised and the coupling between blocks is a forward substitution the kernel does itself.

tf_index names a column of params holding each trajectory's own end time, for ensembles whose members finish at different times; save times past a trajectory's end hold its final state. The default None ends every trajectory at t_span[-1].

trajectories_per_block is one thread's worth of work each, defaulting to a warp; nothing on chip bounds it, since every per-trajectory buffer is thread-local. max_registers caps the kernel's per-thread register count, trading spills against occupancy.

save_hook is a device function hook(save_idx, y, t, p_row, acc) the kernel calls at every save time -- the initial state, each dense-output save, and the frozen saves past a trajectory's own end time -- with the state at that time, so a consumer of the history can be evaluated inside the launch instead of after it. acc is the trajectory's row of the (n, hook_size) output, zeroed at the start and persistent across saves, so the hook can accumulate (a line-of-sight integral, say) or store derived quantities per save. With save_history=False the history output shrinks to the final state, shape (n, 1, n_vars). A solve with a hook returns (hist, hook_out) (plus the stats when asked) and supports neither jax.vmap nor differentiation.

Source code in modax/rodas5P.py
def solve(
    ode_fn,
    y0,
    t_span,
    params,
    *,
    rtol=1e-8,
    atol=1e-10,
    first_step=None,
    max_steps=100000,
    return_stats=False,
    error_weights=None,
    pcoeff=0.0,
    icoeff=1.0,
    dcoeff=0.0,
    lu_precision: str = "fp32",
    trajectories_per_block=None,
    sens_error_control=True,
    sens_param_columns=None,
    sparsity=None,
    ordering="amd",
    tf_index=None,
    max_registers=None,
    array_rhs=None,
    save_hook=None,
    hook_size: int = 0,
    save_history: bool = True,
):
    """JAX-callable Rodas5 custom-kernel solve.

    Only the right-hand side is supplied. The Jacobian ``df/dy``, and the
    partial time derivative ``df/dt`` that a non-autonomous system needs to
    keep fifth-order accuracy, are both differentiated out of ``ode_fn`` with
    Enzyme, so a non-autonomous problem needs nothing extra from the caller.

    ``lu_precision`` (``"fp32"`` or ``"fp64"``) selects the precision of the
    per-step LU factorisation and triangular solves. The state, right-hand
    side, Jacobian and error estimate are always float64; because the
    Rosenbrock--Wanner order conditions retain full order under an approximate
    Jacobian, the ``"fp32"`` default does not reduce the method's order while
    halving the LU shared-memory footprint and exploiting FP32 throughput.
    ``"fp64"`` is available for ill-conditioned systems where the FP32
    factorisation degrades step-size control.

    ``first_step`` pins the initial step size; the default ``None`` (like any
    non-positive value) lets the kernel start from 1e-6 of the integration
    window.

    ``error_weights`` is an optional per-component weight array, shape
    ``(n_vars,)`` or ``(N, n_vars)``, applied in the weighted RMS step-size
    error norm; a weight of 0 excludes that component from step-size control.

    ``pcoeff``/``icoeff``/``dcoeff`` are the PID step-controller gains; the
    default ``(0, 1, 0)`` is the classic I-controller.


    The solve is an XLA custom call into the numba-cuda kernel, so it carries a
    ``jax.custom_jvp`` rule rather than being differentiated by XLA: asking for
    a derivative integrates the continuous forward-sensitivity system alongside
    the state (see ``modax/_sensitivity.py``). ``jax.jvp``, ``jax.jacfwd``,
    ``jax.grad``, ``jax.jacrev`` and ``jax.value_and_grad`` all work with
    respect to ``y0`` and ``params``; ``t_span`` is not differentiable. An
    undifferentiated call runs the plain kernel and pays nothing.

    The joint system has ``n_vars * (1 + n_sens)`` components, but its iteration
    matrix is *not* factorised whole: it is block lower triangular with the same
    ``M0 = I/(h*gamma) - J_y`` on every diagonal block, so one ``n_vars``
    factorisation serves the state and every sensitivity column and the coupling
    is a forward substitution. What grows with ``n_sens`` is the thread's own
    working set -- ten stage vectors of the augmented state, in local memory --
    rather than the block's shared budget, which carries only the ``n_vars``
    matrix. This is still a solver for problems with few parameters relative to
    the state dimension.

    ``sens_param_columns`` restricts the parameter sensitivities to the columns
    named, rather than carrying one block per parameter. The joint system is
    ``n_vars * (1 + n_sens)`` components and each direction costs a
    second-order Enzyme sweep per stage, so this is the difference between
    paying for the parameters you want and paying for the whole ``params`` row
    -- which matters most when that row also carries things that are not
    parameters at all, such as integration bounds or an integer selecting a
    table. The default ``None`` carries every column, as before.

    ``sens_error_control`` decides whether the sensitivity components take part
    in the step-size error norm. The default ``True`` controls them to the same
    ``rtol``/``atol`` as the state, so the gradient is as accurate as the value.
    ``False`` drops them from the norm, which makes the joint solve take exactly
    the step sequence the plain solve takes -- the value then matches a plain
    call bit for bit -- at the cost of nothing tying the sensitivities' accuracy
    to ``rtol``.

    The joint Jacobian's lower-left block ``d(J_y S + J_p)/dy`` is a second
    derivative of ``ode_fn``, and it is formed rather than dropped: Enzyme's
    forward-over-forward sweep applies it to the state increment without ever
    materialising the matrix. Dropping it would be legitimate under the W
    property -- order 5 survives an approximate Jacobian -- but not cheap: the
    error constant it costs was measured at 200x the steps on a right-hand side
    bilinear in state and parameters, which is most reaction networks.

    ``sparsity`` is where a structured problem pays off, and it is the only
    thing a caller has to supply to get one. Pass an ``(n_vars, n_vars)`` mask,
    a scipy sparse matrix, or an ``(nnz, 2)`` array of indices, and two things
    follow. The Jacobian costs one Enzyme sweep per *colour* of the pattern's
    column intersection graph rather than one per column, since columns sharing
    no row can be seeded together and the pattern says which output component
    belongs to which ([`modax._sparsity`][]). And the iteration matrix is
    ordered, factorised symbolically and given an in-kernel sparse LU and sparse
    triangular solves compiled for that exact structure
    ([`modax._sparse_direct`][]). The default -- no pattern -- colours every
    column apart and factorises densely, which is the same mechanism at its
    uninformative end rather than a second path.

    The pattern must be a **superset** of the true nonzeros. Colouring a
    superset only costs sweeps; colouring a subset silently corrupts the entries
    where two columns of a group do overlap after all. It need *not* include the
    factorisation's fill-in, which the symbolic pass works out and gives slots
    of its own.

    ``ordering`` picks the fill-reducing permutation: ``"amd"`` by default,
    SuiteSparse's approximate minimum degree out of ``cvxopt``, or
    ``"natural"`` to skip the ordering. It is ignored without a pattern.

    Forward sensitivities work with either solver: the joint iteration matrix is
    block lower triangular with the same ``M0`` on every diagonal block, so only
    the ``n_vars`` block is ever factorised and the coupling between blocks is a
    forward substitution the kernel does itself.

    ``tf_index`` names a column of ``params`` holding each trajectory's own end
    time, for ensembles whose members finish at different times; save times
    past a trajectory's end hold its final state. The default ``None`` ends
    every trajectory at ``t_span[-1]``.

    ``trajectories_per_block`` is one thread's worth of work each, defaulting to
    a warp; nothing on chip bounds it, since every per-trajectory buffer is
    thread-local. ``max_registers`` caps the
    kernel's per-thread register count, trading spills against occupancy.

    ``save_hook`` is a device function ``hook(save_idx, y, t, p_row, acc)`` the
    kernel calls at every save time -- the initial state, each dense-output
    save, and the frozen saves past a trajectory's own end time -- with the
    state at that time, so a consumer of the history can be evaluated inside
    the launch instead of after it. ``acc`` is the trajectory's row of the
    ``(n, hook_size)`` output, zeroed at the start and persistent across saves,
    so the hook can accumulate (a line-of-sight integral, say) or store derived
    quantities per save. With ``save_history=False`` the history output shrinks
    to the final state, shape ``(n, 1, n_vars)``. A solve with a hook returns
    ``(hist, hook_out)`` (plus the stats when asked) and supports neither
    ``jax.vmap`` nor differentiation.
    """
    n_vars = jnp.shape(y0)[-1]
    # Normalised here, not in the kernel builder, because that is cached on its
    # arguments and a pattern has to arrive as the same hashable value twice.
    if sparsity is not None:
        sparsity = normalize_sparsity(sparsity, n_vars)
    options = KernelOptions(
        pcoeff=pcoeff,
        icoeff=icoeff,
        dcoeff=dcoeff,
        lu_precision=lu_precision,
        trajectories_per_block=trajectories_per_block_or_default(
            trajectories_per_block
        ),
        sparsity=sparsity,
        ordering=ordering,
        tf_index=-1 if tf_index is None else int(tf_index),
        max_registers=max_registers,
        array_rhs=array_rhs,
        save_hook=save_hook,
        hook_size=int(hook_size),
        save_history=bool(save_history),
    )
    settings = dict(
        rtol=rtol,
        atol=atol,
        first_step=first_step,
        max_steps=max_steps,
        return_stats=return_stats,
        error_weights=error_weights,
    )

    if save_hook is not None:
        if hook_size <= 0:
            raise ValueError("a save_hook needs a positive hook_size")
        # The hook's accumulator is a second output the vmap and JVP rules do
        # not know, so a hooked solve is the plain ensemble launch.
        return _solve_impl(ode_fn, y0, t_span, params, options=options, **settings)
    if not save_history:
        raise ValueError("save_history=False needs a save_hook to consume the saves")

    # The JVP rule wraps the vmap-aware solvers rather than the other way
    # round: custom_vmap's own JVP path instantiates symbolic zeros, which is
    # what tells the rule which sensitivity blocks it has to integrate.
    def solver_for(options):
        return make_custom_vmap_solver(
            functools.partial(_solve_impl, ode_fn, options=options, **settings),
            return_stats=return_stats,
        )

    primal_solver = solver_for(options)

    def joint_solver_for(spec):
        return solver_for(replace(options, spec=spec))

    return make_sensitivity_solver(
        primal_solver,
        joint_solver_for,
        jnp.shape(y0)[-1],
        jnp.shape(params)[-1],
        return_stats,
        sens_error_control,
        None
        if sens_param_columns is None
        else tuple(int(c) for c in sens_param_columns),
    )(y0, t_span, params)

modax.rodas5P.KernelOptions dataclass

KernelOptions(
    pcoeff: float = 0.0,
    icoeff: float = 1.0,
    dcoeff: float = 0.0,
    lu_precision: str = "fp32",
    trajectories_per_block: int = _DEFAULT_TRAJECTORIES_PER_BLOCK,
    spec: SensitivitySpec | None = None,
    sparsity: tuple[tuple[int, ...], ...] | None = None,
    ordering: str = "amd",
    tf_index: int = -1,
    max_registers: int | None = None,
    array_rhs: object = None,
    save_hook: object = None,
    hook_size: int = 0,
    save_history: bool = True,
)

Everything that selects a compiled kernel besides ode_fn and the shapes.

Hashable, so it is the kernel cache key: two solves with equal options and the same callback share one compiled kernel. solve builds it once from its keyword arguments and hands it down unchanged; the sensitivity rule substitutes spec and nothing else.

modax.rodas5P.trajectories_per_block_or_default

trajectories_per_block_or_default(requested=None) -> int

How many trajectories a block carries, one per thread.

Nothing on chip bounds this any more: every per-trajectory buffer, the iteration matrix included, is thread-local, so the choice is purely how wide a block should be. A warp is the default because the trajectories in a block step adaptively and diverge, and a warp is the granularity at which that divergence costs nothing.

Source code in modax/rodas5P.py
def trajectories_per_block_or_default(requested=None) -> int:
    """How many trajectories a block carries, one per thread.

    Nothing on chip bounds this any more: every per-trajectory buffer, the
    iteration matrix included, is thread-local, so the choice is purely how wide
    a block should be. A warp is the default because the trajectories in a block
    step adaptively and diverge, and a warp is the granularity at which that
    divergence costs nothing.
    """
    if requested is None:
        return _DEFAULT_TRAJECTORIES_PER_BLOCK
    requested = int(requested)
    if requested < 1:
        raise ValueError(f"trajectories_per_block must be positive, got {requested}")
    return requested

modax.rodas5P.dense_lu_solver cached

dense_lu_solver(n_vars: int)

Dense LU with partial pivoting, one system per thread.

This is what the kernel uses when it is given no sparsity pattern: right-looking LU over the thread's own row-major buffer, then the two triangular solves in place. Returned as the same (factorize_local, solve_local) pair sparse_direct_solver builds from a pattern, so the kernel calls one or the other and has no branch.

It replaced nvmath's LUPivotSolver, whose block-collective API was the only reason the kernel ever put a matrix in shared memory. That cost a barrier around every factorisation and every stage solve, and shared memory that an ensemble's occupancy could not spare. A thread owning its whole trajectory needs neither, and these matrices -- tens of variables, not thousands -- are far too small for cooperation to pay for itself.

Source code in modax/rodas5P.py
@functools.cache
def dense_lu_solver(n_vars: int):
    """Dense LU with partial pivoting, one system per thread.

    This is what the kernel uses when it is given no sparsity pattern:
    right-looking LU over the thread's own row-major buffer, then the two
    triangular solves in place. Returned as the same
    ``(factorize_local, solve_local)`` pair
    [`sparse_direct_solver`][modax._sparse_direct.sparse_direct_solver]
    builds from a pattern, so the kernel calls one or the other and has no branch.

    It replaced nvmath's ``LUPivotSolver``, whose block-collective API was the
    only reason the kernel ever put a matrix in shared memory. That cost a
    barrier around every factorisation and every stage solve, and shared memory
    that an ensemble's occupancy could not spare. A thread owning its whole
    trajectory needs neither, and these matrices -- tens of variables, not
    thousands -- are far too small for cooperation to pay for itself.
    """
    n = n_vars

    @cuda.jit(device=True)
    def factorize_local(lu, ipiv):
        for i in range(n):
            # Partial pivoting: the largest remaining entry in this column.
            max_val = abs(lu[i * n + i])
            pivot = i
            for k in range(i + 1, n):
                val = abs(lu[k * n + i])
                if val > max_val:
                    max_val = val
                    pivot = k
            ipiv[i] = pivot

            if pivot != i:
                for j in range(n):
                    tmp = lu[i * n + j]
                    lu[i * n + j] = lu[pivot * n + j]
                    lu[pivot * n + j] = tmp

            # A singular column is left alone rather than guarded against: the
            # Rosenbrock-W property tolerates an approximate factorisation, and
            # the step controller rejects whatever comes out of one that is not.
            piv = lu[i * n + i]
            if piv != 0.0:
                inv_piv = 1.0 / piv
                for k in range(i + 1, n):
                    factor = lu[k * n + i] * inv_piv
                    lu[k * n + i] = factor
                    for j in range(i + 1, n):
                        lu[k * n + j] -= factor * lu[i * n + j]

    @cuda.jit(device=True)
    def solve_local(lu, ipiv, rhs):
        # Forward substitution through L, applying the pivots as they come.
        for i in range(n):
            pivot = ipiv[i]
            if pivot != i:
                tmp = rhs[i]
                rhs[i] = rhs[pivot]
                rhs[pivot] = tmp
            acc = rhs[i]
            for j in range(i):
                acc -= lu[i * n + j] * rhs[j]
            rhs[i] = acc
        # Back substitution through U.
        for i in range(n - 1, -1, -1):
            acc = rhs[i]
            for j in range(i + 1, n):
                acc -= lu[i * n + j] * rhs[j]
            rhs[i] = acc / lu[i * n + i]

    return factorize_local, solve_local

Tsit5

modax.tsit5.solve

solve(
    ode_fn,
    y0,
    t_span,
    params,
    *,
    rtol=1e-08,
    atol=1e-10,
    first_step=None,
    max_steps=100000,
    return_stats=False,
    error_weights=None,
    pcoeff=0.0,
    icoeff=1.0,
    dcoeff=0.0,
    backend="auto",
    sens_error_control=True,
    sens_param_columns=None,
)

JAX-callable Tsit5 custom-kernel solve.

backend chooses where the kernel keeps the state and its stage vectors: "shared" in per-block shared memory, "local" in the thread's own local memory. The two are bit-identical; shared is faster where the device is under-occupied (small ensembles, low dimension) and local where it is saturated, and "auto" picks by the ensemble's size and the system's, taking shared whenever the system fits and the ensemble is small enough.

The solve is an XLA custom call into the numba-cuda kernel, so it carries a jax.custom_jvp rule rather than being differentiated by XLA: asking for a derivative integrates the continuous forward-sensitivity system alongside the state (see modax/_sensitivity.py). jax.jvp, jax.jacfwd, jax.grad, jax.jacrev and jax.value_and_grad all work with respect to y0 and params; t_span is not differentiable. An undifferentiated call runs the plain kernel and pays nothing.

sens_error_control decides whether the sensitivity components take part in the step-size error norm. The default True controls them to the same rtol/atol as the state, so the gradient is as accurate as the value. False drops them from the norm, which makes the joint solve take exactly the step sequence the plain solve takes -- the value then matches a plain call bit for bit -- at the cost of nothing tying the sensitivities' accuracy to rtol.

Source code in modax/tsit5.py
def solve(
    ode_fn,
    y0,
    t_span,
    params,
    *,
    rtol=1e-8,
    atol=1e-10,
    first_step=None,
    max_steps=100000,
    return_stats=False,
    error_weights=None,
    pcoeff=0.0,
    icoeff=1.0,
    dcoeff=0.0,
    backend="auto",
    sens_error_control=True,
    sens_param_columns=None,
):
    """JAX-callable Tsit5 custom-kernel solve.

    ``backend`` chooses where the kernel keeps the state and its stage vectors:
    ``"shared"`` in per-block shared memory, ``"local"`` in the thread's own
    local memory. The two are bit-identical; shared is faster where the device
    is under-occupied (small ensembles, low dimension) and local where it is
    saturated, and ``"auto"`` picks by the ensemble's size and the system's,
    taking shared whenever the system fits and the ensemble is small enough.

    The solve is an XLA custom call into the numba-cuda kernel, so it carries a
    ``jax.custom_jvp`` rule rather than being differentiated by XLA: asking for
    a derivative integrates the continuous forward-sensitivity system alongside
    the state (see ``modax/_sensitivity.py``). ``jax.jvp``, ``jax.jacfwd``,
    ``jax.grad``, ``jax.jacrev`` and ``jax.value_and_grad`` all work with
    respect to ``y0`` and ``params``; ``t_span`` is not differentiable. An
    undifferentiated call runs the plain kernel and pays nothing.

    ``sens_error_control`` decides whether the sensitivity components take part
    in the step-size error norm. The default ``True`` controls them to the same
    ``rtol``/``atol`` as the state, so the gradient is as accurate as the value.
    ``False`` drops them from the norm, which makes the joint solve take exactly
    the step sequence the plain solve takes -- the value then matches a plain
    call bit for bit -- at the cost of nothing tying the sensitivities' accuracy
    to ``rtol``.
    """

    settings = dict(
        rtol=rtol,
        atol=atol,
        first_step=first_step,
        max_steps=max_steps,
        return_stats=return_stats,
        error_weights=error_weights,
        pcoeff=pcoeff,
        icoeff=icoeff,
        dcoeff=dcoeff,
        backend=backend,
    )
    # The JVP rule wraps the vmap-aware solvers rather than the other way
    # round: custom_vmap's own JVP path instantiates symbolic zeros, which is
    # what tells the rule which sensitivity blocks it has to integrate.
    primal_solver = make_custom_vmap_solver(
        functools.partial(_solve_impl, ode_fn, **settings),
        return_stats=return_stats,
    )

    def joint_solver_for(spec):
        return make_custom_vmap_solver(
            functools.partial(
                _solve_impl,
                ode_fn,
                spec=spec,
                **settings,
            ),
            return_stats=return_stats,
        )

    return make_sensitivity_solver(
        primal_solver,
        joint_solver_for,
        jnp.shape(y0)[-1],
        jnp.shape(params)[-1],
        return_stats,
        sens_error_control,
        None
        if sens_param_columns is None
        else tuple(int(c) for c in sens_param_columns),
    )(y0, t_span, params)

modax.tsit5.clear_caches

clear_caches() -> None

Drop the compiled kernels.

Useful when sweeping problem sizes in a single process: each unique n_vars compiles a separate kernel, and nothing releases it because the module-level caches hold it.

Source code in modax/tsit5.py
def clear_caches() -> None:
    """Drop the compiled kernels.

    Useful when sweeping problem sizes in a single process: each unique
    ``n_vars`` compiles a separate kernel, and nothing releases it because the
    module-level caches hold it.
    """
    _make_body.cache_clear()
    _make_kernel.cache_clear()
    _make_jax_launch.cache_clear()
    clear_sensitivity_caches()
    gc.collect()