r"""phi-FEM-FNO: a neural operator that solves on a DIFFERENT domain each time. After Duprez, Lleras, Lozinski, Vigon and Vuillemot, *Phi-FEM-FNO: a new approach to train a Neural Operator as a fast PDE solver for variable geometries* (arXiv 2502.10033). A neural operator normally learns one geometry. Here the geometry is an INPUT: the domain is described by its level set, which is fed to the network as a channel alongside the source. The network then answers on shapes it has never seen, without remeshing and without retraining:: (f, phi) --FNO--> w and the solution is u = phi * w ⚠ **The multiplication by ``phi`` is the whole point, and it is free.** The network never predicts ``u`` directly: it predicts ``w``, and ``u = phi * w`` vanishes wherever ``phi`` does -- that is, exactly on the boundary of the domain, whatever the network has learned. The Dirichlet condition is a property of the ARCHITECTURE, not something the loss has to enforce. The paper measures this too: its variant predicting ``u`` directly does slightly worse. ⚠ **Same idea as the sine basis** of ``laplacian_2d_strong_residual_fno.py``, one level up: there the basis vanished on the boundary of a fixed square, here ``phi`` vanishes on a boundary that MOVES with the input. Where the data comes from ------------------------- :mod:`~....linear_approximation.phi_fem.masked_poisson` solves the whole family in one compiled program -- which is possible because it weights cells instead of selecting them, the list of active cells being of geometry-dependent length. Measured: 32 geometries in 224 ms. That solver returns ``w`` already, and on a Q1 space over a Cartesian mesh its DOFs ARE the grid nodes, so nothing has to be interpolated between the solver and the network. ⚠ **Two channels here, three in the paper.** The third is ``g``, the boundary data, with ``u = phi w + g``. The solver ported so far imposes ``u = 0``, so a ``g`` channel would be identically zero and teach nothing. Adding it means adding the lifting to the solver first. ⚠ **Volume terms only** in the phiFEM solver: the ghost penalty that stabilises small cut cells is not ported, so the reference itself carries about 2% error. What is measured below is how well the operator reproduces THAT reference, not the exact solution. Run: python phifem_fno_variable_geometry.py """ import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.domains.meshless_domains.domains_nd import HypercubeND from scimba_jax.linear_approximation.basis.analytic_bases import local_lagrange_basis from scimba_jax.linear_approximation.basis.general_bases import AnalyticBasis from scimba_jax.linear_approximation.meshes.mesh import Mesh from scimba_jax.linear_approximation.phi_fem import masked_poisson as phifem from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.linear_approximation.variables.variables_fe import VariablesFE from scimba_jax.mapping.mapping import InvertibleFunction, Mapping from scimba_jax.neural_operator.data_for_no.grid_data import GridData from scimba_jax.neural_operator.discrete_no.grid_based.fno import FNO from scimba_jax.neural_operator.discrete_no.grid_based.phifem_fno import PhiFEMFNO from scimba_jax.nonlinear_approximation.numerical_solvers.no_projectors import ( NOProjector, ) from scimba_jax.utils.scimba_pytree import ScimbaPytree N_CELLS = 32 # background mesh; the DOF grid is (N_CELLS + 1)^2 ORDER = 1 N_TRAIN, N_TEST = 384, 96 N_MODES, HIDDEN_CHANNELS, N_BLOCKS = 12, 16, 3 N_EPOCHS, BATCH_SIZE, LEARNING_RATE = 800, 32, 3.0e-3 SEED = 0 class Blob(ScimbaPytree): """A smooth star-shaped domain: ``sum(((x-c)/ab)^2) - 1 + pinch``. ⚠ Written WITHOUT a square root and with a degree-4 correction, both learnt the hard way. A ``sqrt`` has an infinite derivative at the centre -- which usually falls on a mesh vertex -- and a degree-5 correction outgrows ``r^2``, leaving the zero set unbounded whatever its coefficient: measured, it covered 52% of the square and ran into its edges. Args: mu: ``(5,)`` -- centre x, centre y, half-axis x, half-axis y, pinch. """ def __init__(self, mu): self.mu = jnp.asarray(mu, dtype=float) def __call__(self, x): """The level set at one point. Args: x: a point, ``(2,)``. Returns: a scalar, negative inside. """ offset = x - self.mu[:2] base = jnp.sum((offset / self.mu[2:4]) ** 2) - 1.0 return base + self.mu[4] * (offset[0] * offset[1]) ** 2 class Source(ScimbaPytree): """``f(x) = a + b sin(pi x) sin(pi y)``, carried by a leaf. Args: mu: ``(2,)`` -- the constant and the oscillating amplitude. """ def __init__(self, mu): self.mu = jnp.asarray(mu, dtype=float) def __call__(self, x): """The source at one point. Args: x: a point, ``(2,)``. Returns: a scalar. """ return self.mu[0] + self.mu[1] * jnp.sin(jnp.pi * x[0]) * jnp.sin(jnp.pi * x[1]) def make_space(): """The background mesh and its Q1 space. ⚠ Q1 on a Cartesian mesh: the DOFs sit exactly on the ``(N+1)^2`` grid nodes, so the solver's output is already the network's output grid. No interpolation anywhere between the two. Returns: the variables. """ mesh = Mesh( dim=2, n_cells=(N_CELLS, N_CELLS), ref_quad=UnitSquareTensorized(dim=2, order=2 * ORDER + 2), mapping=Mapping(mappings=[InvertibleFunction(lambda x: x, lambda y: y)]), ) basis = AnalyticBasis( nb_basis=(ORDER + 1) ** 2, out_dim=1, mesh=mesh, basis_type="scalar", local_basis=lambda y, i, m: local_lagrange_basis( y, i, m, order=ORDER, out_dim=1 ), ) return VariablesFE(basis=basis, nb_variables=1) def node_grid(): """The Q1 node positions, as a grid. Returns: ``(n_side, n_side, 2)``. """ side = jnp.linspace(0.0, 1.0, N_CELLS + 1) return jnp.stack(jnp.meshgrid(side, side, indexing="ij"), axis=-1) def draw(key, count): """Geometries and sources for ``count`` problems. Returns: ``(shape parameters, source parameters)``. """ keys = jax.random.split(key, 7) shapes = jnp.stack( [ jax.random.uniform(keys[0], (count,), minval=0.44, maxval=0.56), jax.random.uniform(keys[1], (count,), minval=0.44, maxval=0.56), jax.random.uniform(keys[2], (count,), minval=0.20, maxval=0.32), jax.random.uniform(keys[3], (count,), minval=0.20, maxval=0.32), jax.random.uniform(keys[4], (count,), minval=0.0, maxval=18.0), ], axis=-1, ) sources = jnp.stack( [ jax.random.uniform(keys[5], (count,), minval=0.5, maxval=1.5), jax.random.uniform(keys[6], (count,), minval=-1.0, maxval=1.0), ], axis=-1, ) return shapes, sources def build_dataset(variables, shapes, sources): """Solve the family with phiFEM and lay it out for the network. Args: variables: the FE space, shapes: ``(n, 5)``, sources: ``(n, 2)``. Returns: ``(inputs, targets)`` of shapes ``(n, side, side, 2)`` and ``(n, side, side, 1)``. """ side = N_CELLS + 1 nodes = node_grid().reshape(-1, 2) def one(shape, source): level_set, forcing = Blob(shape), Source(source) w = phifem.solve_phifem(variables, level_set, forcing) phi_values = jax.vmap(level_set)(nodes) f_values = jax.vmap(forcing)(nodes) channels = jnp.stack([f_values, phi_values], axis=-1) return channels.reshape(side, side, 2), w.reshape(side, side, 1) return jax.jit(jax.vmap(one))(shapes, sources) def relative_error(prediction, target, inside=None): """Mean relative L2 error over a batch, optionally restricted to the domain. ⚠ **The restriction is not cosmetic.** Outside the domain ``phi > 0`` and ``u = phi w`` is an extrapolation that means nothing: the paper says as much -- values are "extrapolated by 0 outside Omega_h with no impact on loss computation". Measured here, scoring ``u`` over the whole grid gave 3.24e-01 against 1.43e-01 for ``w``, which is backwards, since ``phi`` weights errors DOWN near the boundary. The extra was noise from outside. Args: prediction: ``(n, ...)``, target: ``(n, ...)``, inside: ``(n, ...)`` boolean, or None to score everywhere. Returns: a float. """ if inside is not None: prediction = jnp.where(inside, prediction, 0.0) target = jnp.where(inside, target, 0.0) numerator = jnp.linalg.norm((prediction - target).reshape(len(target), -1), axis=1) denominator = jnp.linalg.norm(target.reshape(len(target), -1), axis=1) return float(jnp.mean(numerator / denominator)) def main(): """Train the operator, then measure it on geometries never seen.""" key = jax.random.PRNGKey(SEED) variables = make_space() key, sub = jax.random.split(key) shapes, sources = draw(sub, N_TRAIN + N_TEST) start = time.perf_counter() inputs, targets = jax.block_until_ready(build_dataset(variables, shapes, sources)) print( f"phiFEM reference: {N_TRAIN + N_TEST} geometries in " f"{time.perf_counter() - start:.1f} s inputs {inputs.shape}" ) # ⚠ One constant per channel, read off the training set. The level set and # the source have no reason to share a scale, and a network fed values of # wildly different sizes learns the larger one first -- measured elsewhere # today at a factor 11 on the final error. scales = jnp.std(inputs[:N_TRAIN], axis=(0, 1, 2)) inputs = inputs / scales target_scale = float(jnp.std(targets[:N_TRAIN])) targets = targets / target_scale print(f" input scales {np.asarray(scales).round(4)}, target {target_scale:.4f}") grid = GridData(2, HypercubeND([(0.0, 1.0)] * 2), (N_CELLS + 1,) * 2) key, key_net, key_fit = jax.random.split(key, 3) # ⚠ The FNO is the BACKBONE; `PhiFEMFNO` is what knows that channel 1 is the # level set and that the answer is `phi * w`. Keeping the two apart is what # lets the backbone be swapped without touching the reconstruction. operator = PhiFEMFNO( FNO( grid, 2, 1, key_net, n_modes=N_MODES, hidden_channels=HIDDEN_CHANNELS, n_blocks=N_BLOCKS, ), phi_channel=1, ) print(f" phi-FEM-FNO: {operator.ndof()} parameters") projector = NOProjector( operator, (inputs[:N_TRAIN], targets[:N_TRAIN]), learning_rate=LEARNING_RATE, ) start = time.perf_counter() _, projector = projector.project( key_fit, operator, N_EPOCHS, batch_size=BATCH_SIZE, tqdm_desc="phi-FEM-FNO" ) trained = getattr(projector, "operator", projector) print(f" trained in {time.perf_counter() - start:.0f} s") predicted = jax.vmap(trained)(inputs) train = slice(0, N_TRAIN) test = slice(N_TRAIN, N_TRAIN + N_TEST) print("\n### Relative L2 error against the phiFEM reference") print(f" on w, train : {relative_error(predicted[train], targets[train]):.3e}") print(f" on w, test : {relative_error(predicted[test], targets[test]):.3e}") # ⚠ The solution itself, rebuilt as phi * w. The errors differ from those on # w because phi weights them: what happens near the boundary, where phi is # small, counts for less in u than it does in w. u_predicted = trained.reconstruct(predicted, inputs) * target_scale u_reference = trained.reconstruct(targets, inputs) * target_scale # ⚠ INSIDE the domain only -- see `relative_error`. phi_grid = inputs[..., 1] * scales[1] inside = (phi_grid <= 0.0)[..., None] print( f" on u, test : {relative_error(u_predicted[test], u_reference[test], inside[test]):.3e}" " (inside the domain)" ) print( f" on u, test : {relative_error(u_predicted[test], u_reference[test]):.3e}" " (whole grid -- outside is extrapolation, not a solution)" ) # ⚠ The measurement this architecture exists for. residual = jax.vmap(trained.boundary_residual)(inputs) * target_scale print( f"\n |u| where phi = 0 (max) : {float(jnp.abs(residual).max()):.2e}" " <- the Dirichlet condition, by construction" ) # ── Figure ─────────────────────────────────────────────────────────────── figure, axes = plt.subplots(3, 4, figsize=(15, 11), constrained_layout=True) for column in range(4): index = N_TRAIN + column * (N_TEST // 4) phi = np.asarray(phi_grid[index]) inside = phi <= 0 for row, (field, title) in enumerate( ( (np.asarray(u_reference[index, ..., 0]), "phiFEM reference"), (np.asarray(u_predicted[index, ..., 0]), "phi-FEM-FNO"), ( np.abs( np.asarray(u_predicted[index, ..., 0]) - np.asarray(u_reference[index, ..., 0]) ), "|error|", ), ) ): image = axes[row, column].imshow( np.where(inside, field, np.nan).T, origin="lower" ) figure.colorbar(image, ax=axes[row, column]) axes[row, column].set_title(f"{title} -- test {column}", fontsize=9) axes[row, column].set_xticks([]) axes[row, column].set_yticks([]) output = __file__.replace(".py", ".png") figure.savefig(output, dpi=110) print(f"\nfigure: {output}") plt.show() if __name__ == "__main__": main()