Training System#
A single TrainingConfig dataclass centralises every hyperparameter with runtime
validation, and is passed to any solver’s train() method — no kwargs scattered across
multiple calls.
TrainingConfig — full field reference#
Field |
Type |
Default |
Description |
|---|---|---|---|
|
|
|
Total training epochs |
|
|
|
Base learning rate |
|
optax schedule |
|
Overrides |
|
|
|
Collocation mini-batch size |
|
|
|
Initial-condition mini-batch size |
|
|
|
Boundary-condition mini-batch size |
|
|
|
Print interval (used by |
|
|
|
PRNG seed |
|
|
|
List of |
|
|
|
Fuse N steps into one XLA kernel ( |
|
|
|
RAR-D resampling every N outer steps ( |
|
|
|
Candidate pool size ( |
|
|
|
Exponent in |
|
|
|
Output directory; enables auto-restart when non-empty |
|
|
|
Snapshot interval in epochs ( |
Network architectures#
Architecture |
Description |
|---|---|
MLP |
Standard multi-layer perceptron with tanh activations. Configurable depth and
width via a simple layer list, e.g. |
GatedMLP |
Modified MLP (Wang et al., 2022): two input encoders U/V are gate-blended into
every hidden layer. Cures pathological gradient flow on stiff PDEs. Select with
|
FourierMLP |
Trainable random Fourier feature embeddings prepended to a standard MLP. Essential for oscillatory solutions — Helmholtz, wave, high-Re flows — where plain MLPs exhibit spectral bias. |
FBPINN + SimpleGate |
Overlapping subdomain decomposition with sigmoid partition-of-unity windows.
|
See also
Neural operators (FNO, DeepONet, CViT) live in a separate family — see Neural Operators (PINO / DeepONet / CViT).
lax.scan XLA fusion#
Instead of a Python for loop that calls back into Python every epoch, lax.scan
unrolls N gradient steps into a single compiled XLA program. The interpreter only
touches the computation once per n_scan_steps iterations, dramatically reducing
dispatch overhead.
config = TrainingConfig(
epochs = 5000,
lr = 1e-3,
n_scan_steps = 100, # 50 outer Python calls instead of 5000
callbacks = [ConsoleLogger(log_every=500)],
)
solver.train(*data, config=config)
|
Python calls / 5 000 epochs |
Callback granularity |
Use case |
|---|---|---|---|
1 (default) |
5 000 |
every epoch |
Development / debugging |
100 |
50 |
every 100 epochs |
GPU training, medium runs |
500 |
10 |
every 500 epochs |
Long GPU runs, production |
RAR-D adaptive collocation resampling#
At every resample_period outer steps, the solver:
Evaluates the PDE residual
r(x)at a pool ofresample_candidatescandidate pointsComputes sampling probabilities
p(x) ∝ |r(x)|^kReplaces the lowest-residual collocation points with new draws from this distribution
This concentrates compute on high-error regions without changing the total batch size or requiring any geometry change (Lu et al., 2021).
training:
n_scan_steps : 100
resample_period : 5 # every 5 outer steps = every 500 epochs
resample_k : 1.0 # linear in |residual|
Note
For shock-dominated problems (ramp, sod_shock, toro3) underPINN instead uses
RAR/RAD shock-focused resampling (rad_resample, Wu et al. 2023),
p ∝ r^k / E[r^k] + c, tuned via rar_period, rar_candidates, rar_k, rar_c.
QR-DEIM-R adaptive collocation resampling#
underPINN.utils.sampling.qr_deim_resample is an alternative selection rule for the
same resampling slot, addressing a specific weakness of magnitude-proportional
sampling: because RAR-D/RAD draw randomly from p(x) ∝ |r(x)|^k, nothing stops
dozens of draws landing on top of each other on one narrow shock spike while other
high-residual regions go unsampled.
QR-DEIM instead selects points deterministically well-spread across the residual field, using the column-pivoted-QR selection rule from the Discrete Empirical Interpolation Method (Chaturantabut & Sorensen, 2010; Drmač & Gugercin, 2016) to keep the chosen subset as mutually independent as possible.
from underPINN.utils.sampling import qr_deim_resample
new_points = qr_deim_resample(
pde, params, sampler,
n_keep=20000, # collocation batch size
n_candidates=200000, # candidate pool to select from
augment_coords=True, # add residual-scaled coordinates to the feature basis
)
Plain QR-DEIM is capped at one point per basis column, which cannot fill a collocation
batch of thousands. The R (randomized) is a leverage-score-weighted fill (Drineas,
Mahoney & Muthukrishnan, 2006) that reaches the remaining n_keep points while
respecting the same small basis.
Note
This is not a transcription of a specific published “QR-DEIM-R” algorithm — it is a
construction inspired by that selection philosophy. Cost is O(n_candidates × r0)
throughout (r0 = a handful of basis columns); no dense
(n_candidates, n_keep) matrix is ever formed, which at real batch sizes
(e.g. 40,000 × 200,000) would allocate tens of gigabytes.
Shock capturing — artificial viscosity#
The compressible Euler cases add global Laplacian dissipation −ε∇²U on the conserved
variables. ε can be:
Fixed — set directly via
art_viscLearned — jointly optimised as
ε = softplus(log_av)alongside the network weights (trainable_visc: true)
When learned, the parameters become a {"net", "log_av"} pytree optimized by a
single optax chain, so log_av shares the network’s optimizer, learning rate and
cosine schedule rather than having its own.
Note
The shipped shock configs (examples/toro3, examples/ramp, examples/sod_shock)
all set trainable_visc: false and run with a fixed art_visc. If you want the
learned coefficient, you must opt in explicitly.
Time-marching transfer learning#
Long-horizon unsteady flows (e.g. the 3-D pulsatile pipe) are split into time windows. Each window warm-starts from the end-state of the previous window, with per-window checkpoints and window-level restart. See Transfer Learning.
RBA — residual-based adaptive weighting#
Residual-based adaptivity assigns per-point loss weights so boundary and collocation
losses are automatically balanced during training, especially effective for stiff
boundary conditions. Enable with loss.rba: true.
Gauss-Newton / natural-gradient training#
An optional second-order alternative to Adam, for cases where the PDE-residual loss surface is too ill-conditioned for a first-order method to make progress.
Every underPINN loss is a sum of squared residuals, $L(\theta) = \tfrac{1}{2}\lVert r(\theta) \rVert^2$ (PDE residual plus weighted IC/BC residuals, concatenated). Gauss-Newton approximates the Hessian by $J^\mathsf{T}J$ — curvature taken from the residual’s own Jacobian rather than from gradient statistics — and solves the Levenberg-Marquardt-damped normal equations
$$(J^\mathsf{T}J + \lambda I),\Delta\theta = J^\mathsf{T} r$$
with the damping $\lambda$ adapted step-to-step: each step is accepted only if it lowers the loss, damping shrinks after an accepted step and grows after a rejected one (a trust-region scheme). Loss is therefore monotonically non-increasing.
from underPINN.training import train_gauss_newton
# residual_fn: params pytree -> 1-D residual vector
def residual_fn(params):
return pde.residual(params, x_r).reshape(-1)
final_params, loss_hist, damping_hist = train_gauss_newton(
residual_fn, params0,
epochs=200,
damping0=1e-3, # initial LM damping
damping_up=3.0, # grow after a rejected step
damping_down=0.5, # shrink after an accepted step
log_every=20,
)
final_params keeps the same pytree structure as params0; loss_hist and
damping_hist have length epochs + 1 (the initial values are included).
Warning
Scoped to small networks. The step performs an explicit dense solve costing
$O(n_\text{params}^3)$, so it is practical only for networks in the low hundreds to
a few thousand parameters. It is not a replacement for the Adam-based solvers on
the large 3-D or compressible-flow benchmarks — for those, optax.lbfgs
limited-memory quasi-Newton refinement after Adam is the scalable second-order
option.
Note
Two numerical precautions are load-bearing and were each verified necessary:
the step solves an augmented least-squares system rather than forming
$J^\mathsf{T}J$ explicitly (forming it squares the Jacobian’s condition number),
and the whole traced body runs under forced jax.default_matmul_precision("float32")
— JAX’s default GPU matmul precision corrupts jacfwd’s internal matmuls enough to
make the Jacobian wrong in exactly the low-order digits Newton’s method depends on.
See Performance for the precision naming caveat.
Callbacks#
ConsoleLogger
ConsoleLogger(log_every=500)
# Prints: [epoch / total] loss=X.XXe-04 pde=X.XXe-04 ic=X.XXe-03 ...
EarlyStopping
EarlyStopping(patience=400, monitor="loss", min_delta=1e-8)
Monitors a metric (default: total loss) and halts training after patience epochs
without improvement. Correctly fires at the outer-step boundary even inside lax.scan
loops.
ModelCheckpoint
from underPINN.callbacks.checkpoint import ModelCheckpoint
ModelCheckpoint(
out_dir="outputs/burgers/",
monitor="loss", # metric key from the loss aux dict
mode="min", # "min" or "max"
save_best_only=True, # skip non-improving epochs
metadata={
"problem": "burgers",
"network": {"type": "mlp", "layers": [2, 64, 64, 64, 1]},
},
)
Writes params.msgpack + params_meta.json whenever a new best is reached. Full
reference: Model Checkpointing & Inference.
Performance tips#
Problem class |
Recommended |
|---|---|
Fast ODEs |
200 |
Medium PDEs (Burgers, wave, Helmholtz) |
400 – 800 |
Complex PDEs (LDC, airfoil, 3-D) |
1 000 – 2 000 |
Combine early stopping with cosine LR decay (optax.cosine_decay_schedule) for
runs longer than 2 000 epochs — it delivers free accuracy improvement at no extra cost.
See Performance for a full GPU throughput tuning guide.