Skip to content

Steady State

steady_state

Steady states of autonomous ODE right-hand sides.

:func:steady_state finds y* with rhs(t, y*) = 0 by pseudo-transient continuation (PTC). Each iteration solves the implicit-Euler step

(I / dt - J) delta = rhs(t, y)

and grows dt by the ratio of successive residuals, so early iterations follow the dynamics and later ones become Newton steps. Three features of epidemic models need handling (issue #192):

  • Absorbing states. Cumulative counters (deaths, incidence) never reach a zero derivative. Pass them as fixed: they keep their initial value and are excluded from the convergence norms. :func:sink_states finds states that no derivative depends on.
  • Conserved totals. When w @ rhs(t, y) = 0 for every y, the total w @ y is conserved, the steady states form a family indexed by it, and J is singular. As dt grows, I / dt - J becomes ill conditioned and the total drifts. Pass the conserved directions as invariants: the step solves a bordered system that holds W @ y at its initial value exactly and stays nonsingular as dt grows without bound. :func:conserved_quantities gives the structural invariants of a stoichiometry matrix; :func:linear_invariants also finds totals that are conserved because rates balance (births mu * N), which the stoichiometry alone misses.
  • Slow modes. A small residual can hide a large state error along a slowly decaying mode (error ~ residual / rate). Convergence therefore also requires a small Newton correction -J^{-1} rhs, which measures that error, and the final correction is applied to the returned state.

The iteration count is static and every diagnostic stays an array, so the solver runs under jax.jit and jax.vmap. Pass loop=jax.lax.fori_loop there, so the iteration is compiled once rather than unrolled.

SteadyStateConfig(dt0=1.0, dt_max=1000000000000.0, min_growth=2.0, max_growth=10.0, max_iterations=100, residual_tol=1e-09, step_tol=1e-09, atol=1e-08, early_exit=True) dataclass

Controls for :func:steady_state.

Attributes:

Name Type Description
dt0 float

Initial pseudo-time step, in the RHS's time units.

dt_max float

Largest pseudo-time step.

min_growth float

Smallest factor by which dt grows when the unscaled residual did not increase. Without it, a slowly decaying mode barely shrinks the residual and dt stalls. dt is held while the residual grows by less than this factor per step, as it does while an epidemic takes off, and shrinks beyond that.

max_growth float

Largest factor by which dt grows, or shrinks, in one iteration. Between the two bounds dt follows the ratio of successive unscaled residual norms (switched evolution relaxation).

max_iterations int

Iteration budget; the loop is unrolled this many times when early_exit is false.

residual_tol float

Bound on the RMS of rhs / (atol + |y|) over the unknowns, in inverse time units.

step_tol float

Bound on the RMS of the Newton correction (I / dt_max - J)^{-1} rhs / (atol + |y|), a relative state-error estimate.

atol float

Absolute floor of the per-state scale atol + |y|.

early_exit bool

With the default Python loop, stop as soon as the solve converges. This converts the convergence flag to a host boolean, so under jax.jit or jax.vmap either set it to False or, better, pass loop=jax.lax.fori_loop.

__post_init__()

Validate the controls.

Raises:

Type Description
ValueError

If a control is outside its valid range.

Source code in src/op_engine/steady_state.py
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
def __post_init__(self) -> None:
    """Validate the controls.

    Raises:
        ValueError: If a control is outside its valid range.
    """
    for name in ("dt0", "dt_max", "residual_tol", "step_tol", "atol"):
        value = getattr(self, name)
        if (
            not isinstance(value, Real)
            or isinstance(value, bool)
            or not math.isfinite(value)
            or value <= 0.0
        ):
            msg = f"{name} must be finite and positive"
            raise ValueError(msg)
    if self.dt_max < self.dt0:
        msg = "dt_max must be at least dt0"
        raise ValueError(msg)
    for name in ("min_growth", "max_growth"):
        value = getattr(self, name)
        if (
            not isinstance(value, Real)
            or isinstance(value, bool)
            or not math.isfinite(value)
            or value < 1.0
        ):
            msg = f"{name} must be finite and at least one"
            raise ValueError(msg)
    if self.max_growth < self.min_growth:
        msg = "max_growth must be at least min_growth"
        raise ValueError(msg)
    if (
        not isinstance(self.max_iterations, Integral)
        or isinstance(self.max_iterations, bool)
        or self.max_iterations < 1
    ):
        msg = "max_iterations must be a positive integer"
        raise ValueError(msg)

SteadyStateConvergenceError(result)

Bases: RuntimeError

Raised at an eager boundary when a steady-state solve did not converge.

Store the failed result.

Source code in src/op_engine/steady_state.py
173
174
175
176
177
178
179
180
181
def __init__(self, result: SteadyStateResult) -> None:
    """Store the failed result."""
    self.result = result
    super().__init__(
        "Steady-state solve did not converge after "
        f"{int(result.iterations.item())} iterations; scaled residual "
        f"{float(result.residual_norm.item())!r}, Newton correction "
        f"{float(result.step_norm.item())!r}"
    )

SteadyStateResult(state, residual, converged, iterations, residual_norm, step_norm, dt, invariant_drift) dataclass

Candidate steady state and array-valued diagnostics.

Every diagnostic is a zero-dimensional array in the state's namespace, so a traced solve returns them without a host conversion; check them afterwards, or call :func:require_steady_state at an eager boundary.

Attributes:

Name Type Description
state Array

The candidate steady state, including fixed states.

residual Array

rhs(t, state), zero at fixed states.

converged Array

Whether both the residual and the Newton correction met their tolerances.

iterations Array

PTC iterations performed before convergence.

residual_norm Array

Scaled residual RMS at state.

step_norm Array

Scaled RMS of the last Newton correction.

dt Array

Final pseudo-time step.

invariant_drift Array

Largest |W @ state - W @ y0| relative to 1 + |W @ y0|; zero without invariants.

conserved_quantities(stoichiometry, *, tol=1e-10)

Return the structural conservation laws of a reaction network.

Rows w satisfy w @ stoichiometry = 0, so w @ y is unchanged by every reaction, whatever the rates. Totals conserved only because rates balance (births mu * N against deaths) are not structural; use :func:linear_invariants for those.

Parameters:

Name Type Description Default
stoichiometry object

(n_state, n_reactions) net change matrix, such as CompiledReactionNetwork.stoichiometry.

required
tol float

Relative singular-value threshold.

1e-10

Returns:

Type Description
ndarray

(k, n_state) orthonormal rows, possibly empty.

Source code in src/op_engine/steady_state.py
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
def conserved_quantities(stoichiometry: object, *, tol: float = 1e-10) -> np.ndarray:
    """Return the structural conservation laws of a reaction network.

    Rows ``w`` satisfy ``w @ stoichiometry = 0``, so ``w @ y`` is unchanged by
    every reaction, whatever the rates. Totals conserved only because rates
    balance (births ``mu * N`` against deaths) are not structural; use
    :func:`linear_invariants` for those.

    Args:
        stoichiometry: ``(n_state, n_reactions)`` net change matrix, such as
            ``CompiledReactionNetwork.stoichiometry``.
        tol: Relative singular-value threshold.

    Returns:
        ``(k, n_state)`` orthonormal rows, possibly empty.
    """
    return _left_null_space(np.asarray(stoichiometry, dtype=np.float64), tol=tol)

linear_invariants(rhs, states, *, t=0.0, jacobian=None, fixed=None, tol=1e-06)

Find totals w @ y that the dynamics conserve at sampled states.

A row w is returned when w @ J(y) = 0 and w @ rhs(t, y) = 0 at every sampled y, which includes rate-balanced totals that :func:conserved_quantities misses. Sample a few varied states away from special points such as the disease-free state.

The test is numerical. A mode decaying more slowly than tol times the fastest rate looks conserved, and each row's accuracy is limited by the Jacobian's error relative to the slowest real rate: a slow mode tilts the rows toward itself. The default Jacobian uses central differences for that reason; an exact one (jax.jacfwd, for example) is better still. When you know an invariant, such as a row of ones over the living compartments, pass that exact row to :func:steady_state and use this function to confirm that nothing else is conserved.

Parameters:

Name Type Description Default
rhs SteadyStateRhs

rhs(t, y) -> dy/dt.

required
states Sequence[Array]

Sample states, each a one-dimensional array.

required
t float

Evaluation time.

0.0
jacobian SteadyStateJacobian | None

Dense Jacobian callback; central differences by default.

None
fixed Sequence[int] | Array | None

Indices, or a boolean mask, of states that :func:steady_state will hold fixed. They are left out, so the rows have zeros there.

None
tol float

Singular-value threshold relative to the fastest rate.

1e-06

Returns:

Type Description
ndarray

(k, n_state) orthonormal rows, possibly empty.

Raises:

Type Description
ValueError

If states is empty.

Source code in src/op_engine/steady_state.py
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
def linear_invariants(  # noqa: PLR0913
    rhs: SteadyStateRhs,
    states: Sequence[Array],
    *,
    t: float = 0.0,
    jacobian: SteadyStateJacobian | None = None,
    fixed: Sequence[int] | Array | None = None,
    tol: float = 1e-6,
) -> np.ndarray:
    """Find totals ``w @ y`` that the dynamics conserve at sampled states.

    A row ``w`` is returned when ``w @ J(y) = 0`` and ``w @ rhs(t, y) = 0`` at
    every sampled ``y``, which includes rate-balanced totals that
    :func:`conserved_quantities` misses. Sample a few varied states away from
    special points such as the disease-free state.

    The test is numerical. A mode decaying more slowly than ``tol`` times the
    fastest rate looks conserved, and each row's accuracy is limited by the
    Jacobian's error relative to the slowest real rate: a slow mode tilts the
    rows toward itself. The default Jacobian uses central differences for
    that reason; an exact one (``jax.jacfwd``, for example) is better still.
    When you know an invariant, such as a row of ones over the living
    compartments, pass that exact row to :func:`steady_state` and use this
    function to confirm that nothing else is conserved.

    Args:
        rhs: ``rhs(t, y) -> dy/dt``.
        states: Sample states, each a one-dimensional array.
        t: Evaluation time.
        jacobian: Dense Jacobian callback; central differences by default.
        fixed: Indices, or a boolean mask, of states that
            :func:`steady_state` will hold fixed. They are left out, so the
            rows have zeros there.
        tol: Singular-value threshold relative to the fastest rate.

    Returns:
        ``(k, n_state)`` orthonormal rows, possibly empty.

    Raises:
        ValueError: If ``states`` is empty.
    """
    jacobian = jacobian or _central_difference_jacobian(rhs)
    blocks: list[np.ndarray] = []
    free: np.ndarray | None = None
    for y in states:
        values = np.asarray(y, dtype=np.float64)
        if free is None:
            free = _free_mask(fixed, values.shape[0])
        jac = np.asarray(jacobian(t, y), dtype=np.float64)[np.ix_(free, free)]
        # Divide by the state's size so both blocks are rates, comparable
        # with one threshold.
        rate = np.asarray(rhs(t, y), dtype=np.float64)[free] / max(
            float(np.abs(values).max()), 1.0
        )
        blocks.extend((jac, rate.reshape(-1, 1)))
    if free is None:
        msg = "linear_invariants needs at least one sample state"
        raise ValueError(msg)
    rows = _left_null_space(np.concatenate(blocks, axis=1), tol=tol)
    embedded = np.zeros((rows.shape[0], free.shape[0]))
    embedded[:, free] = rows
    return embedded

require_steady_state(result)

Return a converged result or raise.

This converts the convergence flag to a Python boolean; do not call it inside jax.jit.

Parameters:

Name Type Description Default
result SteadyStateResult

Result to check.

required

Returns:

Type Description
SteadyStateResult

result unchanged.

Raises:

Type Description
SteadyStateConvergenceError

If the solve did not converge.

Source code in src/op_engine/steady_state.py
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
def require_steady_state(result: SteadyStateResult) -> SteadyStateResult:
    """Return a converged result or raise.

    This converts the convergence flag to a Python boolean; do not call it
    inside ``jax.jit``.

    Args:
        result: Result to check.

    Returns:
        ``result`` unchanged.

    Raises:
        SteadyStateConvergenceError: If the solve did not converge.
    """
    if not bool(result.converged.item()):
        raise SteadyStateConvergenceError(result)
    return result

sink_states(rhs, states, *, t=0.0, jacobian=None, tol=1e-09)

Return a mask of states no derivative depends on.

Cumulative counters and absorbing compartments that feed nothing back have a zero Jacobian column. They never reach a zero derivative while their inflow is positive, so hold them fixed in :func:steady_state.

Parameters:

Name Type Description Default
rhs SteadyStateRhs

rhs(t, y) -> dy/dt.

required
states Sequence[Array]

Sample states.

required
t float

Evaluation time.

0.0
jacobian SteadyStateJacobian | None

Dense Jacobian callback; central differences by default.

None
tol float

Largest column entry, relative to the largest Jacobian entry, still treated as zero.

1e-09

Returns:

Type Description
ndarray

Boolean (n_state,) mask.

Source code in src/op_engine/steady_state.py
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
def sink_states(
    rhs: SteadyStateRhs,
    states: Sequence[Array],
    *,
    t: float = 0.0,
    jacobian: SteadyStateJacobian | None = None,
    tol: float = 1e-9,
) -> np.ndarray:
    """Return a mask of states no derivative depends on.

    Cumulative counters and absorbing compartments that feed nothing back have
    a zero Jacobian column. They never reach a zero derivative while their
    inflow is positive, so hold them fixed in :func:`steady_state`.

    Args:
        rhs: ``rhs(t, y) -> dy/dt``.
        states: Sample states.
        t: Evaluation time.
        jacobian: Dense Jacobian callback; central differences by default.
        tol: Largest column entry, relative to the largest Jacobian entry,
            still treated as zero.

    Returns:
        Boolean ``(n_state,)`` mask.
    """
    jacobian = jacobian or _central_difference_jacobian(rhs)
    columns = np.zeros(0)
    for y in states:
        magnitude = np.abs(np.asarray(jacobian(t, y), dtype=np.float64)).max(axis=0)
        columns = magnitude if columns.size == 0 else np.maximum(columns, magnitude)
    return columns <= tol * max(float(columns.max(initial=0.0)), 1e-300)

steady_state(rhs, y0, *, t=0.0, jacobian=None, fixed=None, invariants=None, config=None, loop=None)

Find a steady state of rhs near y0 by pseudo-transient continuation.

Parameters:

Name Type Description Default
rhs SteadyStateRhs

rhs(t, y) -> dy/dt for a one-dimensional state y, such as lambda t, y: compiled.eval_fn(t, y, **params).

required
y0 Array

Starting state. It also fixes the absorbing states' values and the invariant totals.

required
t float

Time at which the autonomous RHS is evaluated.

0.0
jacobian SteadyStateJacobian | None

jacobian(t, y) -> (n, n) dense array. Defaults to forward differences (n RHS evaluations per iteration); pass lambda t, y: jax.jacfwd(lambda z: rhs(t, z))(y) under JAX.

None
fixed Sequence[int] | Array | None

Indices, or a boolean mask, of states held at y0.

None
invariants Array | Sequence[Sequence[float]] | None

(k, n) linearly independent rows W whose totals W @ y the dynamics conserve; held at W @ y0. They are setup data, checked on the host, so they must be concrete values (not traced inside jax.jit).

None
config SteadyStateConfig | None

Iteration controls.

None
loop SteadyStateLoop | None

Optional loop(lower, upper, body, carry) driver for the iterations, such as jax.lax.fori_loop. Under jax.jit it compiles the iteration once instead of unrolling max_iterations copies; it runs every iteration, ignoring early_exit. None uses a Python loop.

None

Returns:

Type Description
SteadyStateResult

The candidate steady state and diagnostics.

Raises:

Type Description
ValueError

If shapes are inconsistent.

TypeError

If y0 is not a floating-point vector.

Source code in src/op_engine/steady_state.py
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
def steady_state(  # noqa: PLR0913, PLR0914, PLR0915
    rhs: SteadyStateRhs,
    y0: Array,
    *,
    t: float = 0.0,
    jacobian: SteadyStateJacobian | None = None,
    fixed: Sequence[int] | Array | None = None,
    invariants: Array | Sequence[Sequence[float]] | None = None,
    config: SteadyStateConfig | None = None,
    loop: SteadyStateLoop | None = None,
) -> SteadyStateResult:
    """Find a steady state of ``rhs`` near ``y0`` by pseudo-transient continuation.

    Args:
        rhs: ``rhs(t, y) -> dy/dt`` for a one-dimensional state ``y``, such as
            ``lambda t, y: compiled.eval_fn(t, y, **params)``.
        y0: Starting state. It also fixes the absorbing states' values and the
            invariant totals.
        t: Time at which the autonomous RHS is evaluated.
        jacobian: ``jacobian(t, y) -> (n, n)`` dense array. Defaults to forward
            differences (``n`` RHS evaluations per iteration); pass
            ``lambda t, y: jax.jacfwd(lambda z: rhs(t, z))(y)`` under JAX.
        fixed: Indices, or a boolean mask, of states held at ``y0``.
        invariants: ``(k, n)`` linearly independent rows ``W`` whose totals
            ``W @ y`` the dynamics conserve; held at ``W @ y0``. They are
            setup data, checked on the host, so they must be concrete values
            (not traced inside ``jax.jit``).
        config: Iteration controls.
        loop: Optional ``loop(lower, upper, body, carry)`` driver for the
            iterations, such as ``jax.lax.fori_loop``. Under ``jax.jit`` it
            compiles the iteration once instead of unrolling
            ``max_iterations`` copies; it runs every iteration, ignoring
            ``early_exit``. ``None`` uses a Python loop.

    Returns:
        The candidate steady state and diagnostics.

    Raises:
        ValueError: If shapes are inconsistent.
        TypeError: If ``y0`` is not a floating-point vector.
    """
    config = config or SteadyStateConfig()
    xp = _namespace_of(y0)
    if len(y0.shape) != 1 or y0.shape[0] < 1:
        msg = f"y0 must be a non-empty vector; got shape {y0.shape}"
        raise ValueError(msg)
    if not xp.isdtype(y0.dtype, "real floating"):
        msg = "y0 must have a real floating-point dtype"
        raise TypeError(msg)
    n = int(y0.shape[0])
    dtype = y0.dtype
    jacobian = jacobian or _finite_difference_jacobian(rhs)
    free_host = _free_mask(fixed, n)
    n_free = int(free_host.sum())
    free = xp.asarray(free_host)
    w = _invariant_matrix(invariants, n, xp, dtype)
    k = 0 if w is None else int(w.shape[0])
    totals = None if w is None else xp.matmul(w, y0)

    eye = xp.eye(n, dtype=dtype)
    free_pair = xp.logical_and(xp.reshape(free, (n, 1)), xp.reshape(free, (1, n)))
    inv_dt_max = xp.asarray(1.0 / config.dt_max, dtype=dtype)

    def masked_rhs(y: Any) -> Any:  # noqa: ANN401
        return xp.where(free, rhs(t, y), xp.zeros_like(y))

    def step(y: Any, value: Any, jac: Any, inv_dt: Any) -> Any:  # noqa: ANN401
        """Solve ``(I/dt - J) delta = rhs`` with fixed rows and invariants.

        Returns:
            The step ``delta`` for the state.
        """
        matrix = xp.where(free_pair, xp.subtract(xp.multiply(eye, inv_dt), jac), eye)
        if w is None:
            return xp.linalg.solve(matrix, value)
        bordered = xp.concat(
            (
                xp.concat((matrix, xp.matrix_transpose(w)), axis=1),
                xp.concat((w, xp.zeros((k, k), dtype=dtype)), axis=1),
            ),
            axis=0,
        )
        target = xp.concat((value, xp.subtract(totals, xp.matmul(w, y))))
        return xp.linalg.solve(bordered, target)[:n]

    value0 = masked_rhs(y0)
    scale0 = xp.add(config.atol, xp.abs(y0))
    initial: _Carry = (
        y0,
        value0,
        scale0,
        _scaled_rms(value0, scale0, free, n_free, xp),
        xp.sqrt(xp.mean(xp.multiply(value0, value0))),
        xp.asarray(np.inf, dtype=dtype),
        xp.asarray(config.dt0, dtype=dtype),
        xp.zeros((), dtype=xp.bool),
        xp.zeros((), dtype=xp.int32),
    )

    def iterate(_index: int, carry: _Carry) -> _Carry:  # noqa: PLR0914
        """Advance one PTC iteration; a converged carry passes through.

        Returns:
            The next carry.
        """
        (y, value, scale, residual_norm, raw_norm, _step_norm, dt, converged,
         iterations) = carry  # fmt: skip
        jac = jacobian(t, y)
        # The Newton correction is taken at dt_max rather than an infinite
        # step: identical for every mode faster than 1/dt_max, and finite
        # (but huge, so unconverged) along a conserved direction that was
        # not declared, where pure Newton would be singular.
        newton = step(y, value, jac, inv_dt_max)
        ptc = step(y, value, jac, xp.divide(1.0, dt))
        newton_norm = _scaled_rms(newton, scale, free, n_free, xp)
        done_now = xp.logical_and(
            xp.logical_and(
                xp.less_equal(residual_norm, config.residual_tol),
                xp.less_equal(newton_norm, config.step_tol),
            ),
            xp.all(xp.isfinite(newton)),
        )
        # A converged iterate takes its Newton correction as a final polish.
        candidate = xp.add(y, xp.where(done_now, newton, ptc))
        candidate_value = masked_rhs(candidate)
        candidate_scale = xp.add(config.atol, xp.abs(candidate))
        candidate_raw = xp.sqrt(xp.mean(xp.multiply(candidate_value, candidate_value)))
        positive = xp.greater(candidate_raw, 0)
        ratio = xp.where(
            positive,
            # Both branches are evaluated, so keep the denominator nonzero.
            xp.divide(raw_norm, xp.where(positive, candidate_raw, 1.0)),
            xp.asarray(config.max_growth, dtype=dtype),
        )
        # Grow while the residual falls; hold through moderate rises, which
        # ordinary dynamics produce (an epidemic taking off); shrink only when
        # the residual jumps by more than min_growth in one step.
        factor = xp.where(
            xp.greater_equal(ratio, 1.0),
            xp.clip(ratio, config.min_growth, config.max_growth),
            xp.where(
                xp.greater_equal(ratio, 1.0 / config.min_growth),
                xp.ones_like(ratio),
                xp.maximum(ratio, 1.0 / config.max_growth),
            ),
        )
        updated: _Carry = (
            candidate,
            candidate_value,
            candidate_scale,
            _scaled_rms(candidate_value, candidate_scale, free, n_free, xp),
            candidate_raw,
            newton_norm,
            xp.minimum(xp.multiply(dt, factor), config.dt_max),
            xp.logical_or(converged, done_now),
            xp.add(iterations, xp.ones_like(iterations)),
        )
        held = (
            xp.where(converged, old, new)
            for old, new in zip(carry[:-2], updated[:-2], strict=True)
        )
        return cast(
            "_Carry",
            (*held, updated[-2], xp.where(converged, iterations, updated[-1])),
        )

    carry = initial
    if loop is not None:
        carry = loop(0, config.max_iterations, iterate, carry)
    else:
        for index in range(config.max_iterations):
            carry = iterate(index, carry)
            if config.early_exit and bool(carry[7]):
                break
    (y, value, _, residual_norm, _, step_norm, dt, converged, iterations) = carry

    drift = (
        xp.asarray(0.0, dtype=dtype)
        if w is None
        else xp.max(
            xp.divide(
                xp.abs(xp.subtract(xp.matmul(w, y), totals)),
                xp.add(1.0, xp.abs(totals)),
            )
        )
    )
    return SteadyStateResult(
        state=y,
        residual=cast("Array", value),
        converged=cast("Array", converged),
        iterations=cast("Array", iterations),
        residual_norm=cast("Array", residual_norm),
        step_norm=cast("Array", step_norm),
        dt=cast("Array", dt),
        invariant_drift=cast("Array", drift),
    )