Steady States¶
op_engine.steady_state finds an equilibrium y* with rhs(t, y*) = 0, such
as the endemic state of an epidemic model, without a long burn-in
integration. It uses pseudo-transient continuation: each iteration takes an
implicit-Euler step (I/dt - J) delta = rhs(y) and grows dt as the residual
falls, so early iterations follow the dynamics and later ones become Newton
steps.
A first solve¶
import numpy as np
from op_engine import require_steady_state, steady_state
a = np.array([[-2.0, 1.0], [0.5, -1.0]])
b = np.array([1.0, 2.0])
result = require_steady_state(steady_state(lambda t, y: a @ y + b, np.zeros(2)))
print(result.state, int(result.iterations))
steady_state returns a SteadyStateResult whose diagnostics are arrays:
converged, iterations, residual_norm, step_norm, dt, and
invariant_drift. require_steady_state raises
SteadyStateConvergenceError when the solve did not converge.
Epidemic models: sinks and conserved totals¶
Two features of epidemic models need declaring.
- Absorbing states. A cumulative counter (deaths, incidence) never reaches
a zero derivative while its inflow is positive. Pass it as
fixed, by index or mask: it keeps its starting value and is excluded from the convergence test.sink_states(rhs, samples)finds the states no derivative depends on. - Conserved totals. When births balance deaths (
birth = mu * N), the total population is conserved, so the steady states form a family indexed by it and the Jacobian is singular. Pass the conserved directions asinvariants: each rowwholdsw @ yat its starting value exactly, and the solve stays well posed asdtgrows. conserved_quantities(stoichiometry)returns the structural laws of a reaction network, such asS + I + Rfor a closed SIR.linear_invariants(rhs, samples)also finds totals conserved because rates balance, which births make invisible to the stoichiometry.
from op_engine import linear_invariants, sink_states
beta, gamma, mu, omega = 0.4, 0.1, 1 / (70 * 365), 1 / 365
def sirs(t, y):
s, i, r, deaths = y
n = s + i + r
infection = beta * s * i / n
return np.array([
mu * n - infection - mu * s + omega * r,
infection - (gamma + mu) * i,
gamma * i - (mu + omega) * r,
mu * i,
])
rng = np.random.default_rng(0)
samples = [rng.uniform(10.0, 3000.0, 4) for _ in range(3)]
sinks = sink_states(sirs, samples) # [False, False, False, True]
invariants = linear_invariants(sirs, samples, fixed=sinks) # S + I + R
y0 = np.array([4990.0, 10.0, 0.0, 0.0])
endemic = require_steady_state(
steady_state(sirs, y0, fixed=sinks, invariants=invariants)
)
print(endemic.state[:3], float(endemic.invariant_drift))
linear_invariants is numerical: a mode decaying more slowly than tol
(default 1e-6) times the fastest rate looks conserved, and slow modes limit
the accuracy of the rows it returns. When you know an invariant, pass the
exact row (here [[1, 1, 1, 0]]) and use linear_invariants to confirm that
nothing else is conserved.
For an op_system model, use the compiled RHS. The reaction network gives the structural laws:
from op_system import compile_spec
from op_engine import conserved_quantities, from_compiled_rhs
compiled = compile_spec({
"kind": "transitions",
"state": ["S", "I", "R"],
"transitions": [
{"name": "infect", "from": "S", "to": "I", "rate": "b * I / (S + I + R)"},
{"name": "recover", "from": "I", "to": "R", "rate": "g"},
{"name": "wane", "from": "R", "to": "S", "rate": "w"},
],
})
params = {"b": 0.4, "g": 0.1, "w": 0.01}
network = from_compiled_rhs(compiled, params)
closed = conserved_quantities(network.stoichiometry) # one row: S + I + R
def rhs(t, y):
return np.asarray(compiled.eval_fn(t, y, **params))
sirs_state = require_steady_state(
steady_state(rhs, np.array([990.0, 10.0, 0.0]), invariants=closed)
).state
Convergence¶
A solve converges when two conditions hold, with the scale atol + |y| per
state:
- the scaled residual RMS is at most
residual_tol; - the Newton correction, scaled the same way, is at most
step_tol.
The second condition matters for slowly decaying modes. Near a steady state
the error is about residual / rate, so a residual of 1e-9 along a mode
decaying at 1e-7 per day still leaves an error of about 1e-2. The Newton
correction measures that error directly, and the final correction is applied
to the returned state.
Pseudo-transient continuation converges to a steady state: started next to
an unstable one, such as the disease-free state, it can stop there. The dt
schedule holds the step while the residual grows moderately, as it does while
an epidemic takes off, so a small dt0 (the default is 1 time unit) lets the
iterates follow the epidemic to the endemic state. Warm starts from a nearby
solution converge fastest; raise dt0 for those.
SteadyStateConfig collects the controls: dt0, dt_max, min_growth,
max_growth, max_iterations, residual_tol, step_tol, and atol. The
default tolerances assume double precision.
JAX: jit and vmap¶
Pass an automatic-differentiation Jacobian and loop=jax.lax.fori_loop. The
loop compiles the iteration once, where the default Python loop would unroll
max_iterations copies of it. Invariants and the fixed mask are setup data,
so build them outside the traced function.
import jax
import jax.numpy as jnp
from op_engine import SteadyStateConfig
jax.config.update("jax_enable_x64", True)
total = np.ones((1, 3))
def endemic_state(beta, y0):
def rhs(t, y):
s, i, r = y[0], y[1], y[2]
n = s + i + r
return jnp.stack([
mu * n - beta * s * i / n - mu * s + omega * r,
beta * s * i / n - (gamma + mu) * i,
gamma * i - (mu + omega) * r,
])
result = steady_state(
rhs,
y0,
jacobian=lambda t, y: jax.jacfwd(lambda z: rhs(t, z))(y),
invariants=total,
config=SteadyStateConfig(max_iterations=60),
loop=jax.lax.fori_loop,
)
return result.state, result.converged
states, converged = jax.jit(jax.vmap(endemic_state, in_axes=(0, None)))(
jnp.array([0.3, 0.4, 0.5]), jnp.array([4990.0, 10.0, 0.0])
)
Without a jacobian, steady_state uses forward differences, which cost one
RHS evaluation per state per iteration.