Implementing a custom diffrax.AbstractSolver that composes exact (closed-form) sub-flows — rather than approximating a vector field with a generic Runge–Kutta-type step — currently requires importing several private (underscore-prefixed) modules. This issue covers (a) which private symbols are needed and why, and (b) a concrete general-purpose solver, AbstractSplittingSolver, that would be a natural public addition built on top of a stable version of that API.
(a) Private types currently needed for a custom AbstractSolver
Subclassing diffrax.AbstractSolver to implement a solver whose step doesn't go through the standard vector-field/Runge–Kutta machinery currently requires:
from diffrax._custom_types import VF, Args, BoolScalarLike, Control, DenseInfo, RealScalarLike, Y
from diffrax._local_interpolation import LocalLinearInterpolation
from diffrax._solver.base import AbstractSolver
None of these are specific to any unusual use case — they're the ordinary building blocks (term_structure, interpolation_cls, the type aliases used in init/step/func signatures) that any custom solver needs, whether it's a novel Runge–Kutta tableau or, as below, a composition of exact sub-flows.
Since these modules are private, there's no compatibility guarantee across releases, so any downstream solver built this way is exposed to silent breakage on a Diffrax version bump that currently wouldn't be flagged as a breaking change.
Minimal repro (Diffrax 0.7.x, JAX 0.4.x)
showing the private imports required just to get a toy custom solver's type signatures to typecheck/run:
import jax.numpy as jnp
import diffrax
from diffrax._custom_types import RealScalarLike, Y, Args, DenseInfo, BoolScalarLike
from diffrax._local_interpolation import LocalLinearInterpolation
from diffrax._solver.base import AbstractSolver
class ToyExactSolver(AbstractSolver):
term_structure = diffrax.ODETerm
interpolation_cls = LocalLinearInterpolation
def order(self, terms):
return 1
def init(self, terms, t0, t1, y0, args):
return None
def step(self, terms, t0, t1, y0, args, solver_state, made_jump):
# toy exact flow: dx/dt = -ax => x(t1) = x(t0) exp(-a(t1-t0))
a = args
y1 = y0 * jnp.exp(-a * (t1 - t0))
dense_info = dict(y0=y0, y1=y1)
return y1, None, dense_info, None, diffrax.RESULTS.successful
sol = diffrax.diffeqsolve(
diffrax.ODETerm(lambda t, y, a: -a * y),
ToyExactSolver(),
t0=0.0, t1=1.0, dt0=0.1, y0=jnp.array(1.0), args=0.5,
)
Everything in step here is public JAX; the only reason the private imports are needed at all is to satisfy AbstractSolver's own type contract.
Question for maintainers:
would you accept re-exporting these specific symbols from the public diffrax namespace (lowest-effort fix), or is a broader public "custom solver" extension API (e.g. documented under "Extending Diffrax") something you'd want designed more deliberately?
Happy to raise a PR for either, depending on preference.
Implementing a custom
diffrax.AbstractSolverthat composes exact (closed-form) sub-flows — rather than approximating a vector field with a generic Runge–Kutta-type step — currently requires importing several private (underscore-prefixed) modules. This issue covers (a) which private symbols are needed and why, and (b) a concrete general-purpose solver,AbstractSplittingSolver, that would be a natural public addition built on top of a stable version of that API.(a) Private types currently needed for a custom AbstractSolver
Subclassing diffrax.AbstractSolver to implement a solver whose step doesn't go through the standard vector-field/Runge–Kutta machinery currently requires:
None of these are specific to any unusual use case — they're the ordinary building blocks (term_structure, interpolation_cls, the type aliases used in init/step/func signatures) that any custom solver needs, whether it's a novel Runge–Kutta tableau or, as below, a composition of exact sub-flows.
Since these modules are private, there's no compatibility guarantee across releases, so any downstream solver built this way is exposed to silent breakage on a Diffrax version bump that currently wouldn't be flagged as a breaking change.
Minimal repro (Diffrax 0.7.x, JAX 0.4.x)
showing the private imports required just to get a toy custom solver's type signatures to typecheck/run:
Everything in
stephere is public JAX; the only reason the private imports are needed at all is to satisfyAbstractSolver's own type contract.Question for maintainers:
would you accept re-exporting these specific symbols from the public diffrax namespace (lowest-effort fix), or is a broader public "custom solver" extension API (e.g. documented under "Extending Diffrax") something you'd want designed more deliberately?
Happy to raise a PR for either, depending on preference.