Skip to main content

tinydiffeq

CI Docs PyPI Python versions License: MIT Ruff

Tiny differentiable ODE/SDE/DAE/SDAE solvers for JAX: fixed-step Euler/RK4, adaptive Tsit5, linearly implicit Rodas5P for stiff ODEs and index-1 DAEs, and Euler–Maruyama for Itô SDEs and semi-explicit index-1 SDAEs. The package also includes primal, vmap-friendly finite-state DTMC and CTMC simulators with sequential and associative parallel-prefix execution. Deterministic probability forecasts are differentiable in the initial mass and include DTMC matrix powers, dense CTMC exponentials, and matrix-free Arnoldi/Krylov actions over probability pytrees. The same dense and matrix-free backends are available directly through solve_linear_ode for any fixed homogeneous linear array or pytree operator; jvp_linear_ode and vjp_linear_ode apply the exact initial-state tangent and adjoint exponential actions without differentiating Arnoldi orthogonalization. Fixed stepping and the default adaptive path use bounded lax.scan loops with exactly max_steps attempt slots. Shapes stay static as tolerances or curvature change, and these solves support forward mode, reverse mode, and reverse-over-forward. Adaptive ODE and DAE solves may instead select adaptive_loop="forward": a dynamic lax.while_loop that executes only actual attempts and supports JVP and nested forward AD, but not reverse mode. A vmapped forward loop runs until its slowest lane finishes.

This is a deliberately small, jvp/vjp-friendly package. Rodas5P is a JAX adaptation of Steinebach's method and follows SciML's OrdinaryDiffEqRosenbrock implementation. Use diffrax or SciML if you need general mass matrices, fully implicit or higher-index DAEs, events, continuous solution objects, sparse/Krylov linear solvers for ODE/DAE stages, or specialized adjoints. Initial DAE consistency and explicit DAE stages use nlls-gram; the same nlls solve supplies the square root's implicit derivative, whose default is a direct nonsymmetric LU() solve. LMRootSolver requires residual-only stopping (gtol=xtol=0) and accepts only CONVERGED roots whose Euclidean residual norm is below the root atol. Its max_steps_is_success field remains for source compatibility but does not make MAX_STEPS a valid DAE root.

The linear exponential-action API follows SciML ExponentialUtilities.expv. It includes fixed and residual-controlled adaptive matrix-free time slicing; the latter keeps the Krylov dimension static for predictable JAX compilation. SciML's ExponentialIntegrators.jl is the reference for the broader nonlinear exponential-integrator family.

2.4.0 migration note

  • SaveAt(ts=..., exact=True) now gathers realized knots for explicit fixed-step ODEs. Every query must align with a knot; adaptive methods, Rodas5P, DAEs, SDEs, and SDAEs continue to reject exact mode.
  • Solution.num_steps and DAESolution.num_steps count logical attempts, including rejections. DAE results additionally expose num_root_solves and num_root_steps; num_accepted retains its existing meaning.
  • Adaptive ODE and DAE solves may opt into adaptive_loop="forward" for an actual-work loop. It supports primal, JVP, and nested forward AD but not reverse mode; adaptive_loop="bounded" remains the reverse-mode-capable default. Under vmap, the forward loop runs to the slowest lane.
  • LMRootSolver(predictor="secant") is an opt-in continuation warm start for locally unique algebraic branches; predictor="previous" remains the default.

DAE root acceptance is stricter in 2.4.0. nlls-gram owns both the primal root solve and implicit derivative; square implicit AD defaults to direct LU(). Only CONVERGED roots whose residual norm is below atol are accepted, so gtol and xtol must both be zero. max_steps_is_success remains for source compatibility, now defaults to False, and never makes MAX_STEPS a valid root. Upgrading configurations should remove nonzero gtol/xtol; if they relied on budget exhaustion, increase the root budget or adjust the residual tolerance instead.

Install

uv add tinydiffeq

For GPU use, install the JAX accelerator build that matches your hardware, for example:

uv add tinydiffeq "jax[cuda13]"

Minimal example

The vector field may take (x), (x, t), (x, t, args), or (x, t, args, p) — always in that order. args is pass-through data (not an AD target by convention); p holds differentiable parameters (any pytree). The state may also be any JAX pytree. It must contain at least one leaf, and every leaf must be a nonempty real floating array with the same dtype; vector fields and project preserve that structure. Output keeps the structure and adds the saved-time axis to each leaf.

import jax
import jax.numpy as jnp
from tinydiffeq import solve_ode, Tsit5, IController, SaveAt

jax.config.update("jax_enable_x64", True)  # your call — the library never sets it


def f(x, t, args, p):
    return -p * x


sol = solve_ode(
    f, Tsit5(), 0.0, 2.0, jnp.asarray(1.0),
    p=jnp.asarray(1.3),
    dt_0=0.1,
    controller=IController(rtol=1e-8, atol=1e-10),
    max_steps=512,
    save_at=SaveAt(ts=jnp.linspace(0.0, 2.0, 21)),  # fixed output shape,
)                                                  # however many steps adapt
print(sol.xs)   # states on the grid
print(sol.ok)   # reached t_1 with every requested output valid?

IController() and PIController() choose tolerances from x_0.dtype: rtol=1e-4, atol=1e-6 for float32 and rtol=1e-7, atol=1e-9 for float64. Pass explicit values when tolerances are part of your model's scientific specification. The default dt_min is 10 * finfo(dtype).eps * max(1, abs(t_1)).

max_steps is the total internal attempt budget: accepted steps plus rejections. It is not normally the number of returned times. Endpoint mode returns one time/state, SaveAt(ts=...) returns the requested grid, and SaveAt(steps=True) returns the initial state and accepted internal steps as a contiguous prefix of max_steps + 1 rows. The remaining rows repeat the last accepted state by default; sol.accepted distinguishes data from padding. Rejected attempts never appear in the returned trajectory. sol.num_steps reports the number of attempts actually made, while sol.num_accepted excludes rejections.

SaveAt(ts=...) also accepts a Python sequence. These are observation times: the adaptive controller still chooses its own internal mesh. Explicit methods use cubic Hermite interpolation; Rodas5P uses its published stiff-aware fourth-order continuous extension. For an explicit fixed-step ODE, SaveAt(ts=..., exact=True) instead requires every requested time to be an internal knot and gathers the stored state directly. Exact mode does not apply to adaptive ODEs, Rodas5P, DAEs, SDEs, or SDAEs.

Semi-explicit DAEs

For a square index-1 system dy/dt = f(y, z, t, args, p) and 0 = g(y, z, t, args, p):

from tinydiffeq import IController, Rodas5P, Tsit5, solve_semi_explicit_dae


def dae_f(y, z, t, args, p):
    dy = p * z
    return dy, {"flow": dy}


def dae_g(y, z, t, args, p):
    return z - y


dae_sol = solve_semi_explicit_dae(
    dae_f, dae_g, Tsit5(), 0.0, 1.0,
    jnp.asarray(1.0), jnp.asarray(0.5),
    p=jnp.asarray(2.0), dt_0=0.1,
    controller=IController(), max_steps=128,
)
print(dae_sol.ys, dae_sol.zs, dae_sol.aux["flow"])

# One initial nonlinear consistency solve, then linear Rodas5P stages.
stiff_dae_sol = solve_semi_explicit_dae(
    dae_f, dae_g, Rodas5P(), 0.0, 1.0,
    jnp.asarray(1.0), jnp.asarray(0.5),
    p=jnp.asarray(2.0), dt_0=0.1,
    controller=IController(), max_steps=128,
)

z_0 is a guess and is made consistent automatically. RK4 and Tsit5 restore the algebraic root at every stage. Rodas5P performs no nonlinear solves after initialization: it advances the corresponding block mass-matrix system using one reused LU factorization per attempt. Differential fields may return a floating saved-aux pytree stored at accepted nodes and interpolated on requested deterministic grids. Algebraic equations may separately return internal context passed to the dynamics. On the default bounded path, JVP, VJP, and reverse-over-forward propagate through both implicit initialization and the time integrator. See the DAE documentation for root controls, SaveAt, and scope limits.

DAE solutions expose num_steps, num_root_solves, and num_root_steps as logical per-trajectory work counters. Explicit methods default to reusing the previous algebraic root as the next stage guess; LMRootSolver(predictor="secant") is an opt-in continuation predictor for locally unique root branches.

Fixed-step semi-explicit Itô SDAEs use the corresponding solve_semi_explicit_sdae interface with EulerMaruyama, a PRNG key, and n_steps; see the SDAE documentation.

Gradients through the solve

def endpoint(p):
    return solve_ode(
        f, Tsit5(), 0.0, 2.0, jnp.asarray(1.0), p=p,
        dt_0=0.1, controller=IController(rtol=1e-10, atol=1e-12),
        max_steps=512,
    ).xs

jax.grad(endpoint)(jnp.asarray(1.3))                      # reverse mode
jax.jvp(endpoint, (jnp.asarray(1.3),), (jnp.asarray(1.0),))  # forward mode
jax.grad(lambda p: jax.jvp(endpoint, (p,), (jnp.asarray(1.0),))[1])(
    jnp.asarray(1.3)
)                                                          # reverse-over-forward

The step-size controller is wrapped in stop_gradient (accept/reject is non-differentiable either way, and the error-ratio power blows up at exactly zero error); states differentiate through the solver stages on the realized, frozen mesh. In particular, adaptive SaveAt(steps=True) does not include mesh motion in its time or state derivatives. See the docs for the design contracts: static shapes and SaveAt, AD through adaptive stepping, SDE key semantics, and the package API.

License

MIT

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

tinydiffeq-2.4.0.tar.gz (614.9 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

tinydiffeq-2.4.0-py3-none-any.whl (67.8 kB view details)

Uploaded Python 3

File details

Details for the file tinydiffeq-2.4.0.tar.gz.

File metadata

  • Download URL: tinydiffeq-2.4.0.tar.gz
  • Upload date:
  • Size: 614.9 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/7.0.0 CPython/3.13.14

File hashes

Hashes for tinydiffeq-2.4.0.tar.gz
Algorithm Hash digest
SHA256 932dc3443fb43b94c3522f6ff481ed41352a5ba613019d718b7d8e54613afd5f
MD5 3b3395fc8ccb2f924d0b3a82d15ef95b
BLAKE2b-256 85ef70d9c158b96f8a9405ad053ae993ab50118331b262334824477cd6ba0002

See more details on using hashes here.

Provenance

The following attestation bundles were made for tinydiffeq-2.4.0.tar.gz:

Publisher: publish.yml on HighDimensionalEconLab/tinydiffeq

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

File details

Details for the file tinydiffeq-2.4.0-py3-none-any.whl.

File metadata

  • Download URL: tinydiffeq-2.4.0-py3-none-any.whl
  • Upload date:
  • Size: 67.8 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/7.0.0 CPython/3.13.14

File hashes

Hashes for tinydiffeq-2.4.0-py3-none-any.whl
Algorithm Hash digest
SHA256 7bca7dc9654079c9f114eaf04cbae4bef476dc25e186585fde3a799c9d937dfb
MD5 1a9ac99666b4c444e267317a26c963e8
BLAKE2b-256 5d3521fda2fb23cfb9c22b831ae4d782bd8759f4e1be018755e497457edba269

See more details on using hashes here.

Provenance

The following attestation bundles were made for tinydiffeq-2.4.0-py3-none-any.whl:

Publisher: publish.yml on HighDimensionalEconLab/tinydiffeq

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

Release history Release notifications | RSS feed

2.6.1

2 files

2.6.0

2 files

2.5.0

2 files

This release

2.4.0 This release

2 files

2.3.0

2 files

2.2.0

2 files

2.1.0

2 files

2.0.0

2 files

1.1.0

2 files

1.0.0

2 files

0.3.0

2 files

0.2.0

2 files

0.1.0

2 files

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page