tinydiffeq
Tiny differentiable ODE/SDE/DAE/SDAE solvers for JAX: fixed-step Euler/RK4,
adaptive Tsit5 with integral or proportional-integral step-size control,
Euler–Maruyama for Itô SDEs, and nonstiff semi-explicit index-1 deterministic
and stochastic DAEs.
One bounded lax.scan of exactly max_steps iterations serves fixed and
adaptive stepping, so shapes are static, nothing recompiles as tolerances or
curvature change, and every solve is differentiable in both forward and
reverse mode — including reverse-over-forward, the pattern a
Levenberg–Marquardt optimizer with geodesic acceleration needs when it
differentiates through a rollout. After a solve reaches its horizon, a
lax.cond skips solver and controller work during the padded scan tail.
This is a deliberately small, jvp/vjp-friendly subset of
diffrax. Use diffrax if you need stiff
or fully implicit solvers, higher-index DAEs, full
derivative-term PID control, events, continuous interpolation objects, or
checkpointed/backsolve adjoints. The DAE algebraic solve uses nlls-gram.
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 within the max_steps budget?
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.
SaveAt(ts=...) also accepts a Python sequence. These are observation times:
the adaptive controller still chooses its own internal mesh, and cubic
Hermite interpolation evaluates the solution at every requested point.
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, Tsit5, solve_semi_explicit_dae
def dae_f(y, z, t, args, p):
return p * z
def dae_g(y, z, t, args, p):
return z - y, {"flow": p * z}
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, has_aux=True,
)
print(dae_sol.ys, dae_sol.zs, dae_sol.aux["flow"])
z_0 is a guess and is made consistent automatically. Algebraic equations
may return a floating aux pytree stored at every accepted node and
Hermite-interpolated with z on requested deterministic grids. Algebraic
roots use an implicitly differentiated square LM solve, so JVP, VJP, and
reverse-over-forward propagate with respect to y, t, and p. See the
DAE documentation
for root controls, SaveAt, and scope limits.
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); the states differentiate fully through the RK stages. 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
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file tinydiffeq-1.0.0.tar.gz.
File metadata
- Download URL: tinydiffeq-1.0.0.tar.gz
- Upload date:
- Size: 127.8 kB
- Tags: Source
- Uploaded using Trusted Publishing? Yes
- Uploaded via:
twine/6.1.0 CPython/3.13.12
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
7a13a56a6509e03b1bf4e514d895b843b8c5d32f74aa36bab99637aca80b4af2
|
|
| MD5 |
e18c47d9ef5c239814ff22223af0831e
|
|
| BLAKE2b-256 |
fabb63def7643842d7ae1c1689e856eafa2818938b8ffcdcc208222615f90562
|
Provenance
The following attestation bundles were made for tinydiffeq-1.0.0.tar.gz:
Publisher:
publish.yml on HighDimensionalEconLab/tinydiffeq
-
Statement:
-
Statement type:
https://in-toto.io/Statement/v1 -
Predicate type:
https://docs.pypi.org/attestations/publish/v1 -
Subject name:
tinydiffeq-1.0.0.tar.gz -
Subject digest:
7a13a56a6509e03b1bf4e514d895b843b8c5d32f74aa36bab99637aca80b4af2 - Sigstore transparency entry: 2138845194
- Sigstore integration time:
-
Permalink:
HighDimensionalEconLab/tinydiffeq@23fa371742eef93055269807e10d74d403a6b0d9 -
Branch / Tag:
refs/tags/v1.0.0 - Owner: https://github.com/HighDimensionalEconLab
-
Access:
public
-
Token Issuer:
https://token.actions.githubusercontent.com -
Runner Environment:
github-hosted -
Publication workflow:
publish.yml@23fa371742eef93055269807e10d74d403a6b0d9 -
Trigger Event:
release
-
Statement type:
File details
Details for the file tinydiffeq-1.0.0-py3-none-any.whl.
File metadata
- Download URL: tinydiffeq-1.0.0-py3-none-any.whl
- Upload date:
- Size: 31.1 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? Yes
- Uploaded via:
twine/6.1.0 CPython/3.13.12
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
2b86a439c0f140173afdc6c083d444c067c962e2483503d4adb4a18e5914a42d
|
|
| MD5 |
ee800b713eda29f20482b16e627e404b
|
|
| BLAKE2b-256 |
ac17c55578c32c2148b4f53d7670a618b3641d9a6bcd5429cd91f620703d76dc
|
Provenance
The following attestation bundles were made for tinydiffeq-1.0.0-py3-none-any.whl:
Publisher:
publish.yml on HighDimensionalEconLab/tinydiffeq
-
Statement:
-
Statement type:
https://in-toto.io/Statement/v1 -
Predicate type:
https://docs.pypi.org/attestations/publish/v1 -
Subject name:
tinydiffeq-1.0.0-py3-none-any.whl -
Subject digest:
2b86a439c0f140173afdc6c083d444c067c962e2483503d4adb4a18e5914a42d - Sigstore transparency entry: 2138845201
- Sigstore integration time:
-
Permalink:
HighDimensionalEconLab/tinydiffeq@23fa371742eef93055269807e10d74d403a6b0d9 -
Branch / Tag:
refs/tags/v1.0.0 - Owner: https://github.com/HighDimensionalEconLab
-
Access:
public
-
Token Issuer:
https://token.actions.githubusercontent.com -
Runner Environment:
github-hosted -
Publication workflow:
publish.yml@23fa371742eef93055269807e10d74d403a6b0d9 -
Trigger Event:
release
-
Statement type: