1. Training on data

Given pairs (f_i, F_i) — a source and its known solution — a PhysicNOProjector trains the same way a PINN does, through a data residual instead of a PDE one. Each training example is its own physical model carrying its own f; a batch is a batch of models, not a batch of points.

from scimba_jax.neural_operator.physic_no.abstract_deeponet import AbstractDeepONet
from scimba_jax.nonlinear_approximation.approximation_spaces.physic_no_approximation_spaces import (  # noqa: E501
    PhysicNOApproximationSpace,
)
from scimba_jax.nonlinear_approximation.numerical_solvers.physic_no_projectors import (
    PhysicNOProjector,
)
from scimba_jax.physical_models.ode.anti_derivative import AntiDerivative1DWithData

# A DeepONet learns the MAP f -> F, not one instance of it: every training
# example is its own physical model, so a batch is a batch of PDEs.
# pre_encoder says how a model exposes its own f to the branch network.
class DeepONetAntiDerivative(AbstractDeepONet):
    def pre_encoder_size(self):
        return self.grid_size

    def pre_encoder(self, physical_model):
        return physical_model.data_residuals["interior"].data_vals[0][..., 0]


no = DeepONetAntiDerivative(grid_size=20)
space = PhysicNOApproximationSpace(dims={"x": 1}, list_models=[no], model_type="x")
# Trained against DATA: each pde carries its own (f, F) pair.
# only_data=True turns off every physics term.
projector = PhysicNOProjector(pdes, space, sampler, only_data=True, weights=custom_weights)
key, projector = projector.project(key, space, n_epochs=20, batch_size=4, n_dl=20)
# 4 antiderivative instances, grid of 20, 20 epochs: loss 0.18 -> 0.048.

2. Training on physics — PhysicNO

Swap the data residual for a PDE residual and nothing about the architecture moves — only the physical model does, and only its pre_encoder reads f as a callable instead of a precomputed array. No labelled solution is needed anywhere: the network learns to satisfy the equation at sampled points, the same contract a PINN trains against, but for a whole family of right-hand sides at once instead of one.

from scimba_jax.physical_models.ode.anti_derivative import AntiDerivative1D

# Same architecture -- only pre_encoder changes, reading f AS A FUNCTION,
# and the physical model carries the ODE residual u' - f. No F data anywhere.
pdes = [AntiDerivative1D(main_domain=domain_x, f_rhs=f_i) for f_i in sampled_fs]
projector = PhysicNOProjector(pdes, space, sampler)  # only_data defaults to False
key, projector = projector.project(key, space, n_epochs=20, batch_size=4, n_colloc=50, n_bc_colloc=1)
# Same 4 instances, 20 epochs, collocation instead of data: loss 0.70 -> 0.50.
# Slower to fall -- a PDE residual is a weaker signal than a labelled value --
# but it needs no F, exactly when F is what a solver cannot supply.

The batch axis is the simulation

A PINN batches points of one problem; a PhysicNO batches problems. Every training instance is a physical model of the same class carrying its own f and its own data, and the N of them are stacked once (create_batch: every array leaf of the model gains a leading axis of length N; the domain, the residual definitions and anything shared are stored a single time). A mini-batch is then a slice of that stack by index — never a fresh stack, which would recompile at every step.

At each step the sampler draws one set of collocation points, shared by the whole mini-batch (in_axes=None), while each model brings its own f and its own observations (in_axes=0); the residual is vmap-ed over the models and the loss is its mean over the simulations. Shapes (batch_size, n_colloc) are fixed, so the step compiles once. What is common to all instances — an observation grid, say — lives in the model as a plain numpy array and is not duplicated batch_size times; only the values that differ per simulation carry the batch axis.

from scimba_jax.physical_models.model_sampler import ModelSampler

# N simulations = N physical models of one class, each with its own f.
pdes = [AntiDerivative1D(main_domain=domain_x, f_rhs=f_i) for f_i in sampled_fs]

# Stacked ONCE (create_batch): every array leaf gets a leading axis of length
# N; the domain and the residual definitions are stored a single time.
sampler = ModelSampler(pdes)
# A mini-batch is a SLICE of the stack -- one pytree holding 4 models. Never
# a fresh stack: that would be a new treedef, hence a compilation, per step.
key, mini_batch, idx = sampler.sample(key, batch_size=4)

# Inside a training step, the projector draws one set of collocation points
# and vmaps the residual over the models -- the points are shared, f is not:
#     vmap(residual, in_axes=(None, None, 0))(space, points, mini_batch)
# The loss is the mean over the 4 simulations; (batch_size, n_colloc) is
# fixed, so the step compiles once. PhysicNOProjector does all of this from
# the plain list:
projector = PhysicNOProjector(pdes, space, sampler_x)
key, projector = projector.project(key, space, n_epochs=20, batch_size=4, n_colloc=50)
Pick a shape for the map

3. Architectures

What differs between them is the geometry the operator is willing to see and how it moves information across the domain — never the training loop above, which is shared by all of them.

FNO

Fourier Neural Operator, on a Cartesian grid: mixes channels in Fourier space, where a global convolution is a pointwise product — cheap, and exact for periodic problems.

Geo-FNO

An FNO that first warps an irregular domain onto a regular grid through a learned change of coordinates, so the same spectral machinery still applies off a square or a box.

φ-FEM-FNO

The domain itself is an input, not a fixed shape: a level set marks which grid cells are inside, so one trained network answers on a whole family of variable geometries.

U-Net

Downsamples while doubling channels, then upsamples back, with a skip connection at every level joining encoder and decoder before each convolution — the structure that keeps fine detail alive.

GINO

Geometry-Informed Neural Operator: reads an irregular point cloud onto a latent regular grid, applies an FNO there, and reads back out — an FNO's cost with a mesh's geometry.

Pointwise operators

The same map applied independently at every point — no mixing across the domain at all. The smallest member of the family, and the one the shared API is validated against first.

GNN / GNO on a mesh

Message passing on a mesh graph, where the neighbourhood decides whether the result is a network on this graph or an operator on the domain — the mesh's edges, or a physical ball summed with the mesh's quadrature weights. See below.

DeepONet

A branch network reads the source function at a fixed set of sensor points, a trunk network reads the query point, and their combination is the operator's output — architecture-agnostic encoders and decoders, unlike the grid- or cloud-bound ones above.

One more architecture lives in scimba_jax but not on this page: a classical FEM or DG solve, read as encode/propagate/decode in its own right. It belongs with mesh-based solvers rather than here, since nothing about it is learned unless a coefficient is — see inverse problems for that case.

Meshes as graphs

4. Graph layers — neural network, or neural operator?

One layer, written once, and everything the literature names is a choice of two functions in it. What decides whether the result depends on the mesh or on the domain is neither of them.

One layer: message, aggregation, update

\(m_i = \sum_{j \in N(i)} w_{ij}\, \varphi_\theta(h_i, h_j, x_j - x_i), \qquad h_i' = \gamma_\theta(h_i, m_i)\)

message_kind picks \(\varphi\): GraphSAGE's \(W h_j\), MoNet's Gaussian bumps on the local coordinates \((\rho, \text{direction})\), an MLP of the pair (MPNN), or a kernel \(\kappa_\theta(x_j - x_i)\, h_j\) whose matrix an MLP returns (GNO). update_kind picks \(\gamma\): an activation, GIN's \((1+\varepsilon) h + m\), a GRU gate, an MLP of \([h, m]\) — each with a residual. MeshGraphNet is the same layer with an edge state updated alongside. None of this decides what the layer is.

The neighbourhood does. With mode="gnn", \(N(i)\) is the mesh's edges and \(w_{ij} = 1/d_i\): the layer reads "one edge away". Refine the mesh and the edges halve — the same weights now see half the distance, and what was learned is a function of the discretisation. With mode="gno", \(N(i)\) is the ball \(B(x_i, r)\) and \(w_{ij} = |K_j|\), the measure of cell \(j\): a Riemann sum of \(\int_{B(x_i, r)} \varphi(\cdot, y)\, dy\) with the mesh's own quadrature — the cell measures, or the lumped nodal masses, the weights a finite-element assembly uses. Refine the mesh and the sum converges to the integral: the weights define a map between functions.

The ball is geometry, so it is computed upstream — once, in MeshData(ball_radii=…), Euclidean or geodesic (the sum of edge lengths, which does not shortcut through a hole) — and stored; the layer only reads it. That is what keeps the layers jittable, lets the mesh arrive with the call, and gives the ball layers a continuous version readable at any \(x\).

from scimba_jax.neural_operator.data_for_no.mesh_data import MeshData
from scimba_jax.neural_operator.layers.graph_based import GraphNetwork, GraphOperator

# Geometry is computed ONCE, upstream, and stored: the mesh's edges, and the
# ball B(x_i, r) for every radius a layer will ask for. Cell centres are the
# nodes, cell measures the quadrature weights.
graph = MeshData(2, mesh, ball_radii=(0.13,))

# The same layers, two neighbourhoods. "gnn": the edges, averaged.
# "gno": the ball, summed with the cell measures -- a Riemann sum.
gnn = GraphNetwork(2, 1, 1, 16, 4, "gnn", 0.13, "kernel", "mlp", key=key)
gno = GraphNetwork(2, 1, 1, 16, 4, "gno", 0.13, "kernel", "mlp", key=key)

operator = GraphOperator(gno, graph)  # the discrete family's (u, mu) call
projector = NOProjector(operator, (u0_train, uT_train))
_, projector = projector.project(key, operator, 1000, batch_size=16)

# The operator test: the SAME weights, read on the nested mesh 4x finer,
# with that mesh's own ball and measures.
finer = MeshData(2, refined_mesh, ball_radii=(0.13,))
uT_fine = GraphOperator(projector.operator.network, finer)(u0_fine)
# gnn kernel: 1.47e-1 on the training mesh, 4.78e-1 on the finer one.
# gno kernel: 2.00e-1                     -> 1.49e-1; r = 0.2: 8.3e-2 -> 4.7e-2.

What the test shows

A Gaussian advected and diffused on the JET cross-section — the exact operator is a shifted Gaussian convolution, so nothing else can be blamed. Every message is trained in both modes on a mesh of 483 cells, with the radius chosen so that the ball holds as many neighbours as the edges do (4.3 against 3.8): the two modes cost the same and differ only in what a neighbourhood is. Then the same weights are read on the nested mesh four times finer.

The GNN learns the coarse mesh well (MoNet 5.3e-2) and loses it on the fine one (7.3e-1). The GNO is worse where it was trained (kernel 2.0e-1) and better on the finer mesh (1.5e-1): four cells per ball are a crude quadrature, and the finer mesh hands it the integral it was approximating. Nine cells per ball (r = 0.2) give 8.3e-2 → 4.7e-2 — under 10 % on a mesh never seen. The lesson is the operator's, not the layer's: with a quadrature and a physical radius, resolution is something the network is given, not something it was taught.

The same choice runs through the graph U-Net: pooling by parenthood on a nested hierarchy (gnn) or by a ball of coarse cells around each fine centre (gno), the radius doubling per level so the coarse levels carry the long-range part — the same job the FNO does in the middle of a GINO.

mode  message   L2 level 0  L2 level 1   (advected-diffused Gaussian, JET)
gnn   kernel     1.47e-01    4.78e-01
gnn   monet      5.26e-02    7.28e-01
gnn   mp         1.26e-01    4.57e-01
gno   kernel     2.00e-01    1.49e-01    r = 0.13, 4 cells per ball
gno   mp         2.07e-01    1.36e-01
gno   kernel     8.27e-02    4.65e-02    r = 0.20, 9 cells per ball
identity         8.81e-01

Trained on level 0 (483 cells), read on level 1 (1932 cells) with the same weights. Relative L2 on 32 test samples.