Coupling loop
Purpose
Advance the plasma state in time while the physics modules disagree with each other: transport coefficients depend on the profiles, and the profiles depend on the transport coefficients. Serves brief invariants 3 (nothing impure inside the compiled region) and 4 (the whole simulation differentiable end to end), and is the place where a guarded surrogate is substituted for a physics module.
Method
One time step solves the fixed point
\[p = G(p), \qquad G(p) = \mathcal{S}\big(p_\text{old};\, \chi(p),\, S(p)\big),\]
where \(p = (T_e, T_i, n_e)\) on cell centres, \(\chi\) are transport coefficients on cell faces, \(S\) are volumetric sources on cell centres, and \(\mathcal S\) is one backward-Euler (\(\theta = 1\)) solve of the transport equations with those coefficients held fixed (see transport solver). Units follow the conventions document: SI except temperature in eV and density in m⁻³; the radial coordinate is normalised toroidal flux throughout.
Algorithm
Algorithm 1 — one step (as in tkit.core.loop.make_step):
eq <- equilibrium(state) # quasi-static, recomputed once per step
p_k <- nonlinear_iteration(state.profiles) # Algorithm 2 or 3, fixed iteration count
chi, S <- transport(eq, p_k), sources(eq, p_k)
p <- implicit_solve(p_old, chi, S, dt) # corrector: coefficients frozen at p_k
p <- sawtooth(eq, p) # smoothed trigger
emit trace(profiles, fluxes, sources, ok_masks, residual)
The outer loop over steps is a scan of fixed length with per-step checkpointing, so reverse-mode memory stays bounded over long discharges. No unbounded loop appears anywhere in the compiled region — a deliberate restriction from the brief, and the reason iteration counts are fixed rather than tolerance-driven.
The corrector solve, deliberately
The accepted state is the exact solution of one final solve with coefficients frozen at the last iterate, rather than the last iterate itself. Consequence: the discrete particle and energy balances hold exactly for the coefficients that are recorded in the trace, however poorly the nonlinear iteration converged. Without it, non-convergence silently corrupts conservation; with it, non-convergence shows up only where it should, in the reported residual.
Nonlinear iteration
Two schemes, selected by loop.solver.
Algorithm 2 — relaxed Picard (the M0 scheme, retained as fallback): iterate \(p \leftarrow p + \alpha\,(G(p) - p)\) a fixed number of times with \(\alpha = 0.7\).
Algorithm 3 — Newton-Raphson (default since OQ-4): solve \(F(p) = p - G(p) = 0\).
- work in field-normalised variables \(x\), each field scaled by its magnitude at the start of the step, so the Jacobian is \(O(1)\);
- \(J = I - \partial g/\partial x\) by forward-mode automatic differentiation (dense);
- \(\delta = -J^{-1}F\);
- line search over a fixed candidate set: step lengths \(1, \tfrac12, \tfrac14, \tfrac18\) and the relaxed Picard update, evaluated in parallel; keep the smallest residual, with non-positive states scored as infinite;
- repeat for a fixed iteration count.
# tkit/core/loop.py @ 91d4e55 (one Newton iteration, trimmed)
p_g, ok_t_new, ok_s_new = fixed_point(p)
x = to_x(p)
f = x - to_x(p_g) # F(x) = x - g(x)
jac = jnp.eye(x.size) - jax.jacfwd(g)(x) # dense, forward mode
dx = jnp.linalg.solve(jac, -f)
cands = jnp.concatenate(
[x[None] + lengths[:, None] * dx[None], (x - cfg.relax * f)[None]]
) # Newton steps 1, 1/2, ... and Picard
best = jnp.argmin(jax.vmap(merit)(cands)) # merit = residual, inf if non-positive
return (to_p(cands[best]), ok_t & ok_t_new, ok_s & ok_s_new), resIncluding the Picard update among the candidates is what makes the guarantee “never worse per iteration than the scheme it replaced” hold, and keeps the fallback the owner asked for inside the same code path rather than as a separate mode.
Why not “iterate until converged”
An unbounded loop cannot be reverse-differentiated without special handling, makes compiled cost unpredictable, and hides non-convergence. Fixed count plus reported residual makes the failure mode visible: stall_tol is compared against the final residual of every step and the run prints a stall on stderr. Tolerances are never loosened to silence it (brief §9).
Measured behaviour
| case | Picard | Newton |
|---|---|---|
| toy, default: final residual | 3 × 10−5C-002 | 10−14C-003 |
| toy, stiff (original M0 settings): final residual | stalls at 2 × 10⁻², reported | 3 × 10−11C-005 |
| cost, 100 steps, CPU | baseline | 200× slowerC-004 |
Quadratic convergence is asserted directly: each residual must be below \(50 r^2\) of the previous one until round-off (tests/regression/test_toy_loop.py::test_newton_converges_quadratically).
Fixed-step driver over TORAX
The same loop shape drives TORAX’s own step function. TORAX normally chooses its step adaptively in an interpreted loop; configured for fixed steps, its step is an ordinary compiled function and can be called from inside our scan, with our checkpointing. This gives the brief’s loop shape together with TORAX’s full physics (current diffusion, bootstrap current, fusion power, radiation, pedestal, sawteeth) instead of re-implementing it.
Verified by: identity against its own driver (10−13C-006), reproduction of the stored reference at the reference’s own step (10−14C-007), and gradients through three coupled steps checked against finite differences (10−6C-008).
# tkit/physics/torax_loop.py @ 26f89fb (trimmed)
def simulate(step_fn, state0, post0, n_steps, checkpoint=True):
def body(carry, _):
state, post = step_fn(*carry) # TORAX's own compiled time step
return (state, post), _record(state, post, guard_ok(step_fn, state))
step = jax.checkpoint(body) if checkpoint else body
return jax.lax.scan(step, (state0, post0), None, length=n_steps)TORAX is switched to fixed steps by three config overrides: numerics.fixed_dt, numerics.adaptive_dt = False, and time_step_calculator.calculator_type = "fixed". Each recorded step carries TORAX’s solver status (0 converged, 1 not converged, 2 converged only to the coarse tolerance), its inner iteration count and whether a sawtooth crash was applied, so a stall is visible in the trace rather than only in a log.
Time convergence of the published references
With fixed steps, the step can be halved and the answer watched. For the ramp-up, halving from the reference’s 2 s step changes the core electron temperature by 18%C-012 at t = 2 s; the change shrinks as the step shrinks, as a first-order scheme should (test_fixed_step_first_order_convergence). The flat-top reference uses TORAX’s adaptive step (median 0.12 s, up to 0.42 s) and sits 34%C-013 from the \(\Delta t \to 0\) answer at t = 2 s. The error is concentrated in the initial heating transient. This is not a defect in TORAX, whose examples are tuned for speed, but it means part of any difference from a published curve is time-step error; how the benchmark accounts for it is OQ-5.
Limits. Disabling adaptive stepping also disables TORAX’s retry-on-failure, which is what exposes the stalls below. Gradients work through both of TORAX’s solvers and in both modes: the Newton solver’s lax.while_loop sits inside jax.lax.custom_root, which differentiates the converged solution implicitly instead of the iterations. (This page said until 1 Oct 2026 that reverse mode needed the linear solver; that was wrong.)
Post-mortem: NaN forward-mode gradients through Newton
Recorded 2026-10-01. Forward-mode derivatives (jax.jvp, jacfwd) through one Newton step of the published benchmark were NaN in every component; reverse mode matched finite differences. TORAX’s root finder alone was fine on a toy problem, and the residual’s own forward derivative at the root, the Jacobian (condition number 2 x 103) and its solve were all finite when evaluated directly. The NaN came from how custom_root’s forward rule evaluates that derivative: it differentiates the residual with respect to every closed-over constant, with explicit zero tangents for the ones that do not vary, where an ordinary forward pass skips them. Evaluating the traced derivative operation by operation found the first NaN in TORAX’s trapped-particle fraction (Sauter), \(1 - \sqrt{a}\,(1-\epsilon_\text{eff})/(1 +
2\sqrt{\epsilon_\text{eff}})\): \(\epsilon = 0\) on the magnetic-axis face, the derivative of \(\sqrt{\epsilon}\) is infinite there, and zero times infinity is NaN. It enters the conductivity on the axis face, and the tangent linear solve spreads it to every unknown. Reverse mode puts the same NaN only into the geometry’s cotangent, which nothing reads.
tkit.physics.torax_patches replaces the function, only if its source hash matches TORAX 1.4.3’s, with one whose square root has derivative 0 at 0 (where(e > 0, sqrt(where(e > 0, e, 1)), 0)). Values are unchanged to the last bit; forward mode now equals reverse mode (2 × 10-15 (NaN before the fix)C-062). The same pattern (a square root of a quantity that is exactly zero on the axis) could recur in other TORAX models; the forward-mode test on the benchmark would catch it.
Post-mortem: the stalls that were not a physics event
Recorded 2026-09-22.
Observed. With fixed steps of 0.5 s or less, TORAX’s Newton solver reported non-convergence on up to six steps in a 2.5 s window late in the ITER hybrid ramp-up. Smaller steps failing more is backwards.
What it wasn’t. Plasma current, minimum safety factor, auxiliary and ohmic power, density and internal inductance are all smooth across the window; transport coefficients are smooth too, with no model switching on or off, and only the three pedestal-region faces pinned at their floor, as they are throughout the run. No discrete event, no clipping transition.
What it was. The failing steps exit after 3–5 of 30 allowed iterations, when the line search shrinks the step below tau_min = 0.01 of the original. Neighbouring steps converge only to the coarse tolerance. The residual is hovering near residual_tol = 1e-5, where the line search cannot make progress. Which steps fail is sensitive to compilation details (eager stepping versus scan), confirming marginality rather than a physical cause.
Effect. Raising the iteration limit changes nothing. Lowering the line-search floor to 1e-3 (residual tolerance unchanged) removes all failures and moves stored energy at 80 s by 2 × 10−5 relativeC-011.
Not changed. Defaults are left alone and the stalls are reported; whether fixed-step runs may lower the floor is an owner decision (OQ-5). Recorded as ceiling K-003.
Interface
Public signatures:
Models(equilibrium, transport, sources, sawtooth)— each returns(output, ok_mask).LoopConfig(dt [s], n_steps, solver ∈ {newton, picard}, n_picard, relax, n_newton, n_line_search, stall_tol).simulate(models, edge, cfg, state0) -> (State, Trace), jit-compiled and differentiable.Trace(t, profiles, fluxes, flux_std, sources, ok_transport, ok_sources, picard_residual).
Parameters and defaults
| parameter | default | sensitive to |
|---|---|---|
dt |
5e-4 s (toy) | stiffness; first-order accuracy in time |
n_steps |
100 (toy) | — |
solver |
newton |
stiff transport; Picard stalls (C-005) |
n_newton |
4 | residual reached per step |
n_line_search |
4 | robustness on hard steps; cost per iteration |
relax |
0.7 | Picard mode, and the Picard candidate in Newton’s line search |
stall_tol |
1e-3 | when a stall is reported |
Correctness
| property | test | claim |
|---|---|---|
| 100 coupled steps run compiled, finite, positive, edge condition respected | test_100_steps_under_jit |
C-001 |
| per-step particle and energy balance reproduce from the trace to 1e-9 | tests/conservation |
— |
| Newton converges quadratically | test_newton_converges_quadratically |
C-003 |
| Newton converges where Picard stalls, and the stall is reported | test_newton_converges_where_picard_stalls |
C-005 |
| both solvers agree once converged | test_solvers_agree_on_converged_default |
— |
| gradients finite and finite-difference-checked through 20 coupled steps | tests/diff/test_gradients.py |
— |
| identity against TORAX’s own driver | test_scan_matches_torax_driver |
C-006 |
Ceilings
Changelog
- 2026-09-30: first published version; code excerpts and time-convergence section added.
- 2026-09-22: fixed-step driver section; stall post-mortem (entries 3, 5).
- 2026-09-21 — Newton-Raphson with line search added, Picard retained (entry 1).