Skip to content

Public API for custom solvers composing exact sub-flows (e.g. AbstractSplittingSolver) #771

Description

@mooresm

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.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions