Skip to article frontmatterSkip to article content
Site not loading correctly?

This may be due to an incorrect BASE_URL configuration. See the MyST Documentation for reference.

Lorenz '63 — Your First somax Simulation

Authors
Affiliations
University of Valencia

The Lorenz '63 system is the canonical low-dimensional chaotic attractor. It is the simplest model in somax and a good starting point for understanding the API.

What you’ll learn:

  1. How to create a model, state, and run a forward simulation

  2. How the SomaxModel contract works (vector_field, integrate, diagnose)

  3. How to differentiate through a simulation with jax.grad

  4. How to run an ensemble with jax.vmap

Background

The system is defined by three coupled ODEs with parameters σ\sigma (Prandtl number), ρ\rho (Rayleigh number), and β\beta (geometric factor):

dxdt=σ(y−x),dydt=x(ρ−z)−y,dzdt=xy−βz\frac{dx}{dt} = \sigma(y - x), \qquad \frac{dy}{dt} = x(\rho - z) - y, \qquad \frac{dz}{dt} = xy - \beta z

At the standard parameters (σ,ρ,β)=(10,28,8/3)(\sigma, \rho, \beta) = (10, 28, 8/3) the system exhibits deterministic chaos: nearby trajectories diverge exponentially, tracing the famous butterfly-shaped strange attractor.

1. Create the model

Lorenz63.create() builds a model with the standard parameter values. The model is an eqx.Module — an immutable pytree whose fields (including params) are visible to jax.grad and jax.jit.

Lorenz63(params=L63Params(sigma=weak_f32[], rho=weak_f32[], beta=weak_f32[]))

2. Forward simulation

model.integrate() wraps diffrax.diffeqsolve with automatic boundary-condition enforcement. We save the trajectory at dense output times to visualize the attractor.

Trajectory shape: x=(4000,), y=(4000,), z=(4000,)

3. Visualize the attractor

The classic 3D butterfly and the three time series.

<Figure size 1400x500 with 2 Axes>

4. Diagnostics

model.diagnose() computes on-demand quantities from the state. For L63 it returns the kinetic energy E=12(x2+y2+z2)E = \tfrac{1}{2}(x^2 + y^2 + z^2).

<Figure size 1000x300 with 1 Axes>

5. Adjoint methods for differentiation

Differentiating through an ODE solve requires an adjoint method that trades off memory, accuracy, and speed. diffrax provides three:

AdjointMemoryGradientsUse case
RecursiveCheckpointAdjointO(N)O(\sqrt{N})ExactDefault
DirectAdjointO(N)O(N)ExactShort windows
BacksolveAdjointO(1)O(1)ApproxLong windows
ImplicitAdjointO(1)O(1)ExactSteady states
  • RecursiveCheckpointAdjoint (the default) uses Griewank--Walther optimal checkpointing: it re-computes forward steps from saved checkpoints during the backward pass, giving exact gradients with sub-linear memory.

  • DirectAdjoint stores the entire forward trajectory — exact but O(N)O(N) memory. Also supports forward-mode AD.

  • BacksolveAdjoint solves the continuous adjoint ODE backwards in time with O(1)O(1) memory. Gradients are approximate (optimize-then-discretize ≠\neq discretize-then-optimize).

  • ImplicitAdjoint differentiates through the implicit function theorem: if the solver finds a fixed point u∗u^* such that g(u∗,θ)=0g(u^*, \theta) = 0, the gradient is du∗/dθ=−(dg/du)−1 dg/dθdu^*/d\theta = -(dg/du)^{-1}\, dg/d\theta. Memory is O(1)O(1) and gradients are exact, but only applies to steady-state / fixed-point problems (not time-stepping).

5a. Gradient with respect to parameters

Use case: parameter estimation. Given observations, find the parameters θ=(σ,ρ,β)\theta = (\sigma, \rho, \beta) that minimize a loss.

We compute ∇θL\nabla_\theta \mathcal{L} where L=∑t∥u(t)∥2\mathcal{L} = \sum_t \| \mathbf{u}(t) \|^2 and the state u(t)\mathbf{u}(t) depends on θ\theta through the ODE:

∂L∂θ=∫0Tλ(t)⊤∂f∂θ dt\frac{\partial \mathcal{L}}{\partial \theta} = \int_0^T \lambda(t)^\top \frac{\partial f}{\partial \theta}\, dt

where λ(t)\lambda(t) is the adjoint state satisfying the backward ODE λ˙=−(∂f/∂u)⊤λ\dot{\lambda} = -(\partial f / \partial u)^\top \lambda.

--- Gradient w.r.t. parameters ---
  dL/d(sigma) = -2.7039
  dL/d(rho)   = 15.1506
  dL/d(beta)  = 331.9404

5b. Gradient with respect to the initial state

Use case: state estimation / data assimilation. Given a model with known parameters, find the initial condition u0\mathbf{u}_0 that best fits observations (the 4D-Var problem).

We compute ∇u0L\nabla_{\mathbf{u}_0} \mathcal{L}:

∂L∂u0=λ(0)\frac{\partial \mathcal{L}}{\partial \mathbf{u}_0} = \lambda(0)

where λ(t)\lambda(t) is the same adjoint state, but now evaluated at t=0t = 0. This is the gradient that 4D-Var minimizes.

--- Gradient w.r.t. initial state ---
  dL/d(x0) = -21.0361
  dL/d(y0) = -18.2364
  dL/d(z0) = 39.6907

5c. Joint gradient — parameters and state

Use case: bi-level optimization. Simultaneously estimate the initial state and model parameters (e.g. weak-constraint 4D-Var).

We compute (∇u0L(\nabla_{\mathbf{u}_0} \mathcal{L},∇θL)\nabla_\theta \mathcal{L}) in a single backward pass using eqx.partition to separate the differentiable leaves.

--- Joint gradient ---
  dL/d(x0)    = -21.0361
  dL/d(y0)    = -18.2364
  dL/d(z0)    = 39.6907
  dL/d(sigma) = -2.7039
  dL/d(rho)   = 15.1506
  dL/d(beta)  = 331.9404

5d. Comparing adjoint methods

We time RecursiveCheckpointAdjoint (default) vs DirectAdjoint on the same problem. BacksolveAdjoint requires passing the model as an explicit argument to diffeqsolve (not via closure), so it needs a different API pattern — see the diffrax docs for details.

--- Adjoint method comparison ---
  RecursiveCheckpoint                  dL/d(sigma)=-2.7039  (5.1276s)
  Direct                               dL/d(sigma)=-2.7036  (3.1110s)

6. Ensemble simulation with jax.vmap

Chaotic systems are sensitive to initial conditions. We can explore this by running an ensemble of trajectories with small perturbations and watching them diverge.

<Figure size 1500x400 with 3 Axes>

Summary

Conceptsomax API
Create a modelLorenz63.create(sigma=10, rho=28, beta=8/3)
Initial conditionL63State(x=..., y=..., z=...)
Forward simulationmodel.integrate(state0, t0, t1, dt, saveat=...)
Diagnosticsmodel.diagnose(state)
Grad w.r.t. paramseqx.filter_grad(loss)(model)
Grad w.r.t. statejax.grad(loss)(state0)
Joint gradjax.grad(loss, argnums=(0, 1))(state0, model)
Ensembleeqx.filter_vmap(integrate_one)(batch_states)

Next steps: