Guard and fallback
Purpose
Invariant 5 of the brief: every surrogate sits behind a guard with an explicit input domain and a Tier A fallback. The guard decides, at every evaluation point (every radial face, every time step), whether the surrogate’s answer may be used. Where it may not, the fallback’s answer is used at that point, and the decision is returned so the loop can record it.
The decision
At each face, with features \(x\) (see core state), ensemble mean \(y\) and spread \(\sigma\):
\[\text{ok} = \underbrace{x \in \mathcal D}_{\text{trained here}} \ \wedge\ \underbrace{\sigma \le \tau}_{\text{members agree}} \ \wedge\ \underbrace{y \in \mathcal Y}_{\text{physically possible}} \ \wedge\ \underbrace{\text{finite}(y, \sigma)}_{\text{no NaN}} ,\]
where \(\mathcal D\) is the verified input domain (box, optionally a hull), \(\tau\) a per-output tolerance, and \(\mathcal Y\) the physical output bounds: non-negative diffusivities below \(\chi_\max\) and \(\lvert v\rvert \le v_\max\). The output is
\[y_\text{out} = \begin{cases} y & \text{ok} \\ y_\text{fallback} & \text{otherwise}\end{cases} \qquad \sigma_\text{out} = \begin{cases} \sigma & \text{ok} \\ 0 & \text{otherwise.}\end{cases}\]
The spread is zeroed where the fallback is used, deliberately: the fallback is a physics model with no ensemble, and carrying the rejected surrogate’s spread would attach an uncertainty to a number that didn’t produce it.
# tkit/core/guard.py (GuardedTransportSurrogate.evaluate)
x = self.features(eq, p)
y0, s0 = jax.lax.stop_gradient(self.ensemble(x))
finite = jnp.all(jnp.isfinite(y0), axis=-1) & jnp.all(jnp.isfinite(s0), axis=-1)
ok = (
self.domain.contains(jax.lax.stop_gradient(x))
& jnp.all(s0 <= self.tol, axis=-1)
& self.output_domain.contains(y0)
& finite
)
y, sigma = _masked_eval(self.ensemble, x, ok)
y_fb, ok_fb = self.fallback(eq, p)
y_out = jnp.where(ok[..., None], y, y_fb.to_array())
sigma_out = jnp.where(ok[..., None], sigma, 0.0)Both branches are always evaluated and the choice is a where, not a Python if: the guard is a pure, differentiable function of its inputs and works inside jit and vmap.
A where alone protects the values but not the derivatives. Where the surrogate is rejected because its output is not finite, its derivative there is usually not finite either, and in reverse mode the zero weight the where gives that face multiplies it: 0 × NaN is NaN, and the sum over faces carries it into every parameter. _masked_eval therefore evaluates the ensemble face by face with each face’s features and parameters passed through where(ok, leaf, stop_gradient(leaf)). The cotangent of a rejected face’s copy is then selected away before the faces are summed, so the derivative at a rejected face is the fallback’s alone. The values are unchanged.
The returned mask is the guard decision: True where the surrogate’s output was used. The fallback’s own validity flag is reported separately (evaluate returns both), so a face the surrogate served is never counted as a fallback face. Construction checks the contract: a Tier A fallback, finite non-negative tolerances, and domains of the right width. Tested by tests/unit/test_guard.py: surrogate used when every check passes; fallback when uncertain, outside the domain (per face), or unphysical; the guard is a differentiable pytree; and with the surrogate rejected at all, none or some faces, reverse- and forward-mode derivatives stay finite even when the surrogate’s own derivative is NaN.
Inside TORAX
Once TORAX computes the physics, the guard has to act inside TORAX’s transport step. TORAX lets users register their own transport models, so the guard is registered as one, model_name = "tkit_guarded", holding two ordinary TORAX transport models: the surrogate (QLKNN) and the fallback (Bohm/gyro-Bohm). It evaluates both and picks per face.
Differences from the core guard:
- No ensemble-spread check (K-005). QLKNN in TORAX is a single network, not an ensemble. Domain box, finiteness and output bounds are active.
- Features are computed from TORAX’s state with the same definitions as the core features, which are not the ones QLKNN itself receives (K-008, see core state). Harmless while the domain box is open or arbitrary; to be settled before verified domains arrive.
- Output bounds default to \(\chi_\max = 100\) m² s⁻¹ and \(v_\max = 50\) m s⁻¹.
- Post-processing is applied once, to the combined result. TORAX transport models apply finishing steps (clipping to limits, inner and outer patches, radial smoothing). Those settings are moved from the wrapped models up to the wrapper, so they act on the guarded result rather than separately on each branch. Sub-channel breakdowns are not mixed between branches.
- The recorded mask is a post-step diagnostic (K-014). TORAX evaluates transport inside its solver, at the start of the step and at each nonlinear iterate, and none of those decisions reach its outputs. The fixed-step loop records the guard decision re-evaluated on the accepted end-of-step state, with TORAX’s own pre-step and pedestal machinery. It says where that state lies relative to the guard’s domain; it is not a record of which model served each face during the step, and is labelled
rejected_fraction_poststep, not a fallback fraction. With the domain always open or always closed (the first two rows below) the two coincide.
Three identity tests pin the behaviour:
| guard setting | expected | result |
|---|---|---|
| domain so wide it never triggers | identical to plain TORAX with QLKNN | bit-identicalC-014 |
| domain empty, always triggers | identical to TORAX with the fallback alone | 10−12C-015 |
| trusted only at low gradients | a partial handover, recorded per face | the trace shows both branches in use |
Post-mortem: a default nobody set
Recorded 2026-09-21. Entry 4.
What was written. The first version of the wrapper split the scenario’s transport settings in two: model-specific settings went to the surrogate, and finishing settings the user had written went to the wrapper.
What was correct. Finishing settings the model’s own validator adds must also go to the wrapper. When TORAX validates a QLKNN configuration that has no smoothing_width, it fills in 0.1. Unwrapped, the QLKNN model is the top-level transport model, so that smoothing is applied. Wrapped, the injected value landed on the inner model, where finishing is never applied, and the wrapper kept its own default of no smoothing. Same physics, different finishing.
What the bug looked like from inside. Nothing looked wrong. Profiles were smooth and plausible, the run converged, and the mask itself was correct: the error was in what happened after the guard’s choice, not in the choice. The difference, 5%C-016 in core electron temperature, is within the range usually attributed to model choice. Had it been found after surrogates arrived, it would have been blamed on the surrogate.
What caught it. The identity test in the first row above: with a guard that never triggers, the wrapped model must reproduce the unwrapped one exactly. It is independent of the bug because it compares against TORAX’s own unwrapped path rather than against anything the wrapper computes, and it has exactly one reason to fail.
| core \(T_e\) difference, open guard vs unwrapped | |
|---|---|
| before | 5%C-016 |
| after | bit-identicalC-014 |
What changed. The wrapper now resolves each wrapped model’s configuration through TORAX’s own validator first and lifts any finishing setting that differs from TORAX’s base default up to the wrapper. The identity test stays in the suite (tests/regression/test_torax_guard.py::test_open_guard_reproduces_plain_qlknn_exactly).
Ceilings
K-005 (no ensemble-spread check yet).
Changelog
- 2026-09-30: first published version.
- 2026-09-21: guard inside TORAX; post-mortem (entry 4).