Backend and solve-strategy boundaries¶
op_engine selects the numerical namespace from the state array with
array_api_compat.array_namespace(). This supports arrays with the standard
__array_namespace__ protocol and native arrays such as torch.Tensor
through compatibility namespaces. Explicit,
dense linearly implicit, and IMEX methods use that namespace's Array-API
operations, including linalg.solve. The numerical method does not change when
the array namespace changes.
Namespace portability and automatic differentiation are related but distinct.
The Array API does not define grad, jit, or traced control flow. A namespace
such as JAX can nevertheless differentiate the ordinary op_engine methods
because their fixed-step numerical operations remain in that namespace.
The optional torch dependency group qualifies representative eager
fixed-step kernels with PyTorch autograd. Compiler integration and complete
method qualification remain provider/backend capabilities rather than
consequences of namespace selection.
Portable fixed-step methods¶
With JAX state, fixed-step Euler, Heun, RK4, Dormand--Prince, dense IMEX
(including paired ARK3), and dense linearly implicit methods can be used inside
jax.jit and differentiated with jax.grad. Dynamic values can include the
initial state, RHS parameters, dense operator values, and Jacobian values. The
time grid, method selection, state shape, and Python solver configuration are
structural inputs and must remain static during a trace.
Sparse acceleration is backend-specific. SciPy sparse solves are not a JAX differentiation path; use dense JAX operators when gradients through a solve are required.
The built-in adaptive controller is an eager Python controller. It preserves a
JAX array namespace, and jax.grad can differentiate the numerical operations
on the branch accepted by the nominal solve. Step acceptance itself extracts
scalar error values and uses Python loops and branches, however, so a live
adaptive=True solve is not a JAX-jit-compatible execution path. This
limitation belongs to the controller, not to the explicit Runge--Kutta, IMEX,
or linearly implicit formulas.
Adaptive-controller boundary¶
After a successful adaptive run, CoreSolver.last_adaptive_schedule contains
the accepted internal step sizes. A new solver can pass that value to
replay_adaptive_schedule. Replay skips error norms and acceptance decisions,
but invokes the same high-order Array-API step kernels. The static schedule can
therefore be used inside jax.jit(jax.value_and_grad(...)) while parameters,
initial state, dense operators, and Jacobians remain dynamic.
The resulting derivative is conditional on the recorded mesh. This is also the branchwise meaning of an eager live-adaptive gradient: neither mode differentiates the discrete accept/reject decision. Refresh the schedule as optimization parameters move, and always refresh it after a material model or tolerance change. A schedule is validated against the output grid, while matching the method and other configuration is the caller's responsibility.
Core schedule replay retains the eager Python loop by default. Select
RunConfig(replay_loop="auto") to use a namespace's registered scan driver,
or replay_loop="scan" to require one. JAX supplies lax.scan; NumPy,
PyTorch, and CuPy retain their original eager paths under "auto". Set
replay_checkpoint=True to rematerialize the scan body during reverse mode.
The default replay_loop="unroll" preserves existing behavior for every
namespace. These options affect frozen-mesh replay, not live adaptive runs
or fixed-step run.
The rolled driver invokes the same accepted-step kernels, including the two
half steps used by adaptive Euler and RK4. It preserves Dormand--Prince FSAL
reuse and selects stored output states from the internal mesh. Explicit
methods, dense IMEX methods (including ARK3), implicit Euler, trapezoidal, and
ROS2 support this path. BDF2 still requires fixed stepping. SDIRK2's nonlinear
diagnostics retain eager replay under "auto"; forced "scan" raises an
error. Dense operator factories and Jacobian callbacks must handle traced
scalar times and step sizes using array operations.
For example, record a nominal mesh eagerly, then freeze it during differentiation:
from dataclasses import replace
import jax
import jax.numpy as jnp
import numpy as np
from op_engine import CoreSolver, ModelCore
from op_engine.core_solver import RunConfig
from op_engine.model_core import ModelCoreOptions
output_times = np.asarray([0.0, 1.0])
nominal_solver = CoreSolver(ModelCore(1, 1, output_times))
nominal_solver.core.set_initial_state(np.ones((1, 1)))
adaptive_config = RunConfig(method="dopri5", adaptive=True)
nominal_solver.run(lambda t, y: -0.3 * y, config=adaptive_config)
schedule = nominal_solver.last_adaptive_schedule
replay_config = replace(
adaptive_config, replay_loop="scan", replay_checkpoint=True
)
def loss(rate):
core = ModelCore(
1, 1, output_times, options=ModelCoreOptions(dtype=np.float32)
)
core.set_initial_state(jnp.ones((1, 1), dtype=jnp.float32))
CoreSolver(core).replay_adaptive_schedule(
lambda t, y: rate * y, schedule, config=replay_config
)
return jnp.sum(core.get_current_state() ** 2)
hessian = jax.jit(jax.hessian(loss))(jnp.asarray(-0.3))
op_engine.loop_ops.LoopAdapter separates iteration from numerical stages.
register_loop_adapter(namespace_name, adapter) accepts an optional backend's
scan(body, initial, inputs) and checkpoint operations. The namespace name is
the module's __name__, as returned by op_engine.array_namespace(state).
JAX is imported lazily only for its own namespace; namespaces without an
adapter fall back to eager replay under "auto". A forced scan or checkpoint
request fails clearly when the selected adapter cannot provide it.
The optional flepimop2 provider also consumes the functional kernels and defaults explicit JAX replay to its compact scan driver.
Run uv run python scripts/benchmark_replay.py to measure gradients,
Hessians, and reverse-over-reverse at 8, 64, and 160 steps for Dormand--Prince
and ROS2. Each case runs in a fresh CPU process and reports tracing,
compilation, execution, process peak RSS, and XLA buffer requirements as JSON.
RSS includes backend startup; XLA temporary buffers still scale with the
saved step states, even though the traced graph stays fixed in size. Add
--loops scan unroll --steps 8 16 for a bounded comparison with eager replay;
--no-checkpoint compares storage strategies and --timeout limits each case.
Public functional steps and external loops¶
CoreSolver.fixed_explicit_step(rhs_func, *, method, t, dt, y,
first_stage=None) is the supported public boundary for an external loop
driver. It returns (y_next, fsal) for euler, heun, rk4, or dopri5,
using the same numerical kernels as fixed-step CoreSolver.run. It reads the
solver's configured state shape for validation but does not update its
ModelCore state or history. Supply arrays with that shape and an RHS that
preserves their shape and namespace. t and dt can be traced backend
scalars; the method and state shape remain static.
For example, a non-provider JAX caller can compile a 160-step solve and differentiate it twice using only public APIs:
import jax
import jax.numpy as jnp
import numpy as np
from op_engine import CoreSolver, ModelCore
solver = CoreSolver(ModelCore(1, 1, np.asarray([0.0, 1.0])))
times = jnp.linspace(0.0, 1.0, 161)
def loss(rate):
def rhs(t, y):
return rate * y
def advance(y, step):
t, dt = step
y_next, fsal = solver.fixed_explicit_step(
rhs, method="dopri5", t=t, dt=dt, y=y
)
return y_next, None
final, _ = jax.lax.scan(
jax.checkpoint(advance),
jnp.ones((1, 1)),
(times[:-1], jnp.diff(times)),
)
return 0.5 * jnp.sum(final**2)
hessian = jax.jit(jax.hessian(loss))(jnp.asarray(-0.3))
reverse_over_reverse = jax.jit(
jax.grad(lambda rate: jnp.sum(jax.grad(loss)(rate) ** 2))
)(jnp.asarray(-0.3))
lax.scan traces the step body once. jax.checkpoint lets reverse mode
recompute intermediate stage values rather than retaining all of them.
The example discards the optional FSAL derivative for a simple carry.
For Dormand--Prince, reuse that derivative as first_stage on the next step
to save one RHS evaluation. Take the first step outside the scan to obtain
an array-valued derivative, then scan the remaining steps with
(y_next, fsal) as the carry. A scan carry must retain its shape and structure;
starting with None and returning an array changes that structure. Other
explicit methods return None. Reuse is valid only while the next RHS
evaluation starts at the same time and state with the same model parameters.
An external driver can also flatten a recorded AdaptiveStepSchedule into
step start times and sizes. Heun and Dormand--Prince's fixed steps reproduce
their accepted adaptive updates. Euler and RK4's adaptive updates use two
half steps; an external fixed-step replay must split each recorded step in
two to reproduce those updates. In all cases, freeze the mesh during
differentiation and refresh it when the nominal solve changes.
Projects that require a compiled adaptive controller or other solver-specific
capabilities can provide those at an external plugin or provider boundary. A
specialized integration may return a complete trajectory and adopt it through
ModelCore.apply_trajectory; it does not need to add a backend-specific method
to CoreSolver.
This keeps optional packages such as Diffrax out of the core dependency and method surfaces. The portable fixed-step methods and their differentiation contract remain identical across conforming array namespaces.
Stochastic sampling boundary¶
DirectSSASolver, AdaptiveTauLeapingSolver, and TauLeapingSolver keep
reaction-channel arithmetic in the state array's namespace. Direct SSA injects
an SSASampler, fixed tau injects a PoissonSampler, and adaptive tau uses
both for exact critical events and noncritical firing counts. This is necessary
because the Array API does not define random-number generation and because
NumPy's stateful generator and JAX's explicit keys have intentionally different
semantics.
NumpySSASampler and NumpyPoissonSampler are provided as conveniences.
JAX users can construct fresh keys from the stable direct-SSA draw_index or
tau-leaping step_index.
Direct SSA's optional forcing schedule is static Python configuration. It
preserves the state namespace and uses fresh draw indices when a forcing
boundary invalidates a pending event; observation times retain that event.
The current safety checks and stochastic event loops are eager, so these
stochastic solvers are not JIT-compatible paths. Adaptive selection also reads
drift/variance reductions, critical classifications, and accept/reject results
back to Python. Moreover, categorical reaction events and integer Poisson
samples do not have an ordinary pathwise derivative. JAX remains useful for
eager array execution and for differentiating a deterministic version of the
same model, but jax.grad through a sampled trajectory is not part of this API
contract. Score-function, reparameterized, or other stochastic gradient
estimators should be implemented explicitly by an inference/provider layer
rather than implied by the array namespace.