r"""Geo-FNO on a DISK: the reference is FEM, and the deformation is the question. An FFT wants a uniform grid; a disk is not one. Geo-FNO's answer (Li, Huang, Liu, Anandkumar, 2022) is to learn a map that straightens the domain, do everything spectral in the latent square, and carry the geometry in the map. This file asks how much of that map has to be LEARNED, by running the same operator three times on the same data with three deformations:: analytic the elliptical grid map, known in closed form -- nothing to learn, and the ceiling the other two are measured against learned an InvertibleNet, free to find its own straightening pre-trained the same network, first regressed onto the analytic map, then FROZEN -- the middle case, and the one a practitioner has when a geometry is fixed and a map was fitted once ⚠ **Same data, same seed, same budget.** The three differ by the deformation and by nothing else, or the comparison measures something other than what it claims to. Where the data comes from ------------------------- Not from an analytic solution: on a disk with a moving Gaussian source there isn't one. The reference is scimba's own FEM -- Q2 on a structured mesh pushed through the elliptical map -- and the whole family is solved in ONE compiled program, the way ``uq_batched_transport_1d.py`` does it. Measured: 8 PDEs in 4 ms once compiled. ⚠ **The FEM model and the operator's model are not the same object**, and that is not duplication: a Galerkin scheme wants a weak form, a neural operator wants residuals. The translation is physics, so it belongs here rather than in the library. What the geometry buys, and what it costs ----------------------------------------- ⚠ **The Dirichlet condition is exact on the CURVED boundary**, measured at 2.5e-16. The sine basis vanishes on the edge of the latent square, and the analytic map sends the unit circle exactly onto that edge (checked: distance 0.0 at eight points), so the composition vanishes on the circle. No boundary residual, no weight to tune. ⚠ **A learned deformation carries no such guarantee.** Nothing constrains an InvertibleNet to send the circle onto the square's edge, so the exact boundary condition is a property of the analytic map -- and the run reports the boundary value for each variant precisely so that this stays visible rather than assumed. ⚠ **The projection must have no bias**, or ``Q(0) != 0`` and the whole construction is undone silently. Run: python laplacian_2d_disk_geofno.py """ import copy 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_2d import Disk2D 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.galerkin.fem.elliptic_fe_scheme import ( EllipticFEscheme, ) from scimba_jax.linear_approximation.meshes.mesh import Mesh 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.physic_no.grid_based.physic_informed_fno import ( CONTINUOUS_SYNTHESES, GeoFNO, ResidualDeformation, ) from scimba_jax.nonlinear_approximation.approximation_spaces.physic_no_approximation_spaces import ( # noqa: E501 PhysicNOApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.networks.mlp import MLP from scimba_jax.nonlinear_approximation.numerical_solvers.physic_no_projectors import ( PhysicNOProjector, ) from scimba_jax.nonlinear_approximation.optimizers.native_optimizers import ( ScimbaAdam, optimize_pytree, ) from scimba_jax.physical_models.abstract_physical_model import AbstractPhysicalModel from scimba_jax.physical_models.abstract_physical_weak_model import ( AbstractPhysicalWeakModel, ) from scimba_jax.physical_models.abstract_residuals import InteriorResidual from scimba_jax.physical_models.classical_weakform.diffusion_advection_reaction_weak_form import ( # noqa: E501 EllipticWeakForm, ) from scimba_jax.physical_models.data_residuals import CollocDataResidual from scimba_jax.physical_models.weak_boundary_conditions import Dirichlet from scimba_jax.utils.functional_fields import make_functional_field_class from scimba_jax.utils.scimba_pytree import ScimbaPytree # ── Parameters ─────────────────────────────────────────────────────────────── DIM, RADIUS = 2, 1.0 FEM_CELLS, FEM_ORDER = 12, 2 # the reference solver N_GRID = 32 # the LATENT grid the FNO works on # ⚠ 16 modes, not 8. On the basis nodes the source reconstructs at 3.1e-02 # with 8 and 1.1e-02 with 16 -- and the encoder cannot pass on what it could # not represent. N_LAYERS = 4 is THREE full FNO blocks plus the split final # one, matching the discrete FNO example's three. N_MODES, CHANNELS, N_LAYERS = 16, 8, 4 BASIS = "sine" N_TRAIN, N_TEST = 256, 64 N_EPOCHS, BATCH_SIZE, LEARNING_RATE = 1200, 32, 3.0e-3 PRETRAIN_STEPS, PRETRAIN_POINTS = 2000, 4096 #: How much of the residual network is added at the start. Small keeps the map #: near the affine placement, which is what "initialised at the identity" buys. RESIDUAL_SCALE = 0.1 SEED = 0 #: The source family: a Gaussian whose CENTRE moves. ⚠ The centre is what makes #: the family discriminating -- a family whose difficulty always sits in the #: same place cannot tell a good basis from a bad one, which is the lesson of #: `advection_diffusion_1d_basis_conditioned_by_mu.py`. MU_LOW = jnp.array([-0.45, -0.45, jnp.log(0.15)]) MU_HIGH = jnp.array([0.45, 0.45, jnp.log(0.35)]) N_MU = 3 # ── The geometry ───────────────────────────────────────────────────────────── def square_to_disk(y): """``[0,1]^2 -> disk``, through the elliptical grid map. ⚠ Not the polar map: that one sends a whole side of the square onto the centre and its Jacobian vanishes there. This one is smooth, bijective, and sends the four sides onto the circle -- which is what a single structured mesh needs to cover a disk without a hole or a singular centre. Args: y: reference points, ``(..., 2)``. Returns: physical points, ``(..., 2)``. """ u, v = 2.0 * y[..., 0] - 1.0, 2.0 * y[..., 1] - 1.0 return RADIUS * jnp.stack( [u * jnp.sqrt(1.0 - v**2 / 2.0), v * jnp.sqrt(1.0 - u**2 / 2.0)], axis=-1 ) def disk_to_square(p): """The closed-form inverse, so no geometric Newton is ever run. Args: p: physical points, ``(..., 2)``. Returns: reference points, ``(..., 2)``. """ x, y = p[..., 0] / RADIUS, p[..., 1] / RADIUS root = jnp.sqrt(2.0) def half(a, b, s): return 0.5 * jnp.sqrt(jnp.maximum(a + s, 0.0)) - 0.5 * jnp.sqrt( jnp.maximum(b - s, 0.0) ) u = half(2.0 + x**2 - y**2, 2.0 + x**2 - y**2, 2.0 * root * x) v = half(2.0 - x**2 + y**2, 2.0 - x**2 + y**2, 2.0 * root * y) return jnp.stack([0.5 * (u + 1.0), 0.5 * (v + 1.0)], axis=-1) # ⚠ ONE class, at module level, shared by every PDE -- as # `laplacian_2d_strong_residual_fno.py` and `uq_batched_transport_1d.py` both # do. Built per residual instead, each source would land at index 0 of its own # registry and the batch would silently collapse onto the first one. SOURCE_FIELD = make_functional_field_class("geofno_disk_source") # ── The family ─────────────────────────────────────────────────────────────── class GaussianSource(ScimbaPytree): """``f(x) = exp(-|x - c|^2 / 2w^2)``, carried by a LEAF. ⚠ ``mu`` is written as ``jnp``, which is what makes the family a stackable batch; and nothing declares it trainable, so it stays the PDE's data. Args: mu: ``(3,)`` -- centre x, centre y, log width. """ def __init__(self, mu): self.mu = jnp.asarray(mu, dtype=float) def __call__(self, x): """``f(x)``, of shape ``(1,)``. Args: x: one physical point. Returns: the source there. """ centre, width = self.mu[:2], jnp.exp(self.mu[2]) return jnp.exp(-jnp.sum((x - centre) ** 2) / (2.0 * width**2))[None] def weak_model(mu): """The FEM's view of one PDE: ``-Delta u = f``, ``u = 0`` on the boundary. Args: mu: the source parameters. Returns: the variational model. """ form = EllipticWeakForm( dim=DIM, A=lambda _x: jnp.eye(DIM), b=lambda _x: jnp.zeros(DIM), c=lambda _x: jnp.array(0.0), f=GaussianSource(mu), ) model = AbstractPhysicalWeakModel(dim=DIM) model.add_weak_form("main", form) for side in ("west", "east", "south", "north"): model.add_boundary_condition(side, Dirichlet(lambda _x: jnp.zeros(1))) return model class StrongLaplacian(InteriorResidual): """``-Delta u = f``, the operator's view of the same PDE. Args: domain: the disk, f_rhs: the source. """ def __init__(self, domain, f_rhs): super().__init__(domain=domain, size=1, model_type="x", f_rhs=f_rhs) def construct_residual(self, *variables): """The left-hand side. Args: *variables: the space's variables. Returns: ``-Delta u``. """ return -variables[0].laplacian("x") class LaplacianOnDisk(AbstractPhysicalModel): """What the projector samples: the source, and the observed values. Args: domain: the disk, mu: the source parameters, data: ``(points, values)``. """ def __init__(self, domain, mu, data): super().__init__(main_domain=domain) self.mu = jnp.asarray(mu, dtype=float) # ⚠ WRAPPED in a SHARED functional-field class, and this is the whole # correctness of the batch. Passing `GaussianSource(mu)` directly sends # `_construct_rhs` down its `else` branch, which calls # `make_functional_field_class("auto")` -- a NEW class per residual, # whose registry is a CLASS attribute. Every model then holds a # one-entry registry and the index 0; `create_batch` takes the FIRST # model's treedef, stacks four zeros and keeps its registry. Measured: # func_id = [0 0 0 0], and every PDE of the batch reads the SAME source. # Nothing raises, and the operator dutifully learns the family mean. self.physical_residuals = { "interior": StrongLaplacian(domain, SOURCE_FIELD(GaussianSource(mu))) } self.add_data_residual( "data", # ⚠ False: the sensor points are the SAME for every PDE, so only # the values are batched. CollocDataResidual(size=1, model_type="x", data=data, batchable_args=False), ) # ── The reference: FEM, solved as one batch ────────────────────────────────── def fem_scheme(model): """The Q2 space on the mapped disk, shared by the whole batch. Args: model: any model of the family. Returns: the scheme. """ mesh = Mesh( dim=DIM, n_cells=(FEM_CELLS, FEM_CELLS), ref_quad=UnitSquareTensorized(dim=DIM, order=2 * FEM_ORDER + 2), mapping=Mapping(mappings=[InvertibleFunction(square_to_disk, disk_to_square)]), is_identity_mapping=False, ) basis = AnalyticBasis( nb_basis=(FEM_ORDER + 1) ** DIM, out_dim=1, mesh=mesh, basis_type="scalar", local_basis=lambda y, i, m: local_lagrange_basis( y, i, m, order=FEM_ORDER, out_dim=1 ), ) return EllipticFEscheme(model, VariablesFE(basis=basis, nb_variables=1)) def reference_values(mus, points): """Solve the whole family with FEM and read it at ``points``. ⚠ One compiled program for the batch, as ``uq_batched_transport_1d.py`` does: the space, the mesh and the numbering do not depend on the source, so rebuilding them per PDE would pay for an identical geometry N times. Args: mus: ``(n, 3)``, the source parameters, points: ``(m, 2)``, where the solution is read. Returns: ``(n, m, 1)``. """ models = [weak_model(mu) for mu in mus] reference = fem_scheme(models[0]) batched = type(models[0]).create_batch(models) back_solve = EllipticFEscheme._make_back_solve_fn() zero = reference._initial_dofs() variables = reference.variables def solve_one(pde): scheme = copy.copy(reference) scheme.pde = pde factorisation = EllipticFEscheme.factorise(scheme, zero) dofs = back_solve(scheme, zero, factorisation.lu, factorisation.pivots) return jax.vmap( lambda p: type(variables)._classical_local_evaluate_pure(variables, dofs, p) )(points) return jax.jit(jax.vmap(solve_one))(batched) # ── The three deformations ─────────────────────────────────────────────────── def new_deformation(key, hidden=32, n_layers=3): """The paper's map: an affine placement plus a small residual network. ⚠ ``f`` is an ordinary MLP, not an invertible network. With the continuous analysis the inverse is never called, so invertibility is a property of the MATHEMATICS (the map should stay a diffeomorphism) rather than a requirement of the code. The paper uses a three-layer feedforward network of width 32 on sinusoidal features; this keeps the shape and the width. ⚠ ``factor`` and ``offset`` place the disk of radius ``RADIUS`` inside the unit cube BEFORE the residual acts, so the map starts somewhere sensible instead of at a random place -- which is the whole point of the residual form. Args: key: a random generator state, hidden: width of the residual network, n_layers: how many hidden layers. Returns: the deformation. """ inner = MLP( in_size=DIM, out_size=DIM, hidden_sizes=[hidden] * n_layers, activation="tanh", key=key, ) return ResidualDeformation( inner, factor=0.5 / RADIUS, offset=0.5, scale=RESIDUAL_SCALE ) def pretrain_deformation(key, network): """Regress the network onto the analytic map, then hand it back. ⚠ The middle variant exists because it is what one actually has: a fixed geometry, a map fitted once, and no wish to keep paying for it. Freezing is the caller's job (``learn_deformation=False``), not this function's. Args: key: a random generator state, network: the invertible network to fit. Returns: ``(network, final loss)``. """ key, sub = jax.random.split(key) latent = jax.random.uniform(sub, (PRETRAIN_POINTS, DIM)) physical = square_to_disk(latent) target = disk_to_square(physical) del key def loss(net): return jnp.mean((jax.vmap(net)(physical) - target) ** 2) optimiser = optimize_pytree(ScimbaAdam, network, loss, learning_rate=1.0e-3) @jax.jit def step(net, opt): value = loss(net) _, net, opt = opt.update(net, {}) return net, opt, value value = jnp.inf for _ in range(PRETRAIN_STEPS): network, optimiser, value = step(network, optimiser) return network, float(value) class UnbiasedProjection(ScimbaPytree): """``Q`` without a bias, so that ``Q(0) = 0`` and the boundary stays exact. Args: in_channels: the latent width, out_channels: the solution's size, key: a random generator state. """ from scimba_jax.utils.scimba_pytree import trainable as _trainable weight = _trainable(True) def __init__(self, in_channels, out_channels, key): self.weight = jax.random.normal(key, (in_channels, out_channels)) / jnp.sqrt( in_channels ) def __call__(self, value): """Mix the channels linearly. Args: value: ``(in_channels,)``. Returns: ``(out_channels,)``. """ return value @ self.weight def make_operator(key, deformation, learn, input_scale, sample_points, sample_area): """One GeoFNO, differing from its siblings only by the deformation. Args: key: a random generator state, deformation: anything answering ``__call__`` on a point, learn: leave the deformation as it is (True) or freeze it (False), input_scale: what the encoder divides the sampled source by, sample_points: the physical points where the source is read, sample_area: the physical area each of them stands for. Returns: the operator. """ key_net, key_projection = jax.random.split(key) latent_grid = GridData(DIM, HypercubeND([(0.0, 1.0)] * DIM), (N_GRID,) * DIM) return GeoFNO( latent_grid, 1, 1, key_net, deformation=deformation, sample_points=sample_points, sample_area=sample_area, learn_deformation=learn, channels=CHANNELS, n_modes=N_MODES, n_layers=N_LAYERS, basis=BASIS, input_scale=input_scale, projection=UnbiasedProjection(CHANNELS, 1, key_projection), final_channel_mlp=False, final_activation=False, ) # ── Training and reporting ─────────────────────────────────────────────────── def relative_errors(projector, models, exact, points): """Relative L2 error per model, at the sensor points. Args: projector: the trained projector, models: the physical models, exact: ``(n, m, 1)``, the FEM reference, points: ``(m, 2)``. Returns: ``(n,)``. """ return np.asarray( [ float( jnp.linalg.norm(projector.evaluate(model, points) - exact[i]) / jnp.linalg.norm(exact[i]) ) for i, model in enumerate(models) ] ) def boundary_value(projector, model): """The largest ``|u|`` found on the CIRCLE. ⚠ The measurement that separates the variants: with the analytic map the sine basis makes this machine zero, because the map sends the circle onto the square's edge. A learned map is under no such constraint, and this is where that shows. Args: projector: the trained projector, model: any model of the family. Returns: the largest absolute value on the circle. """ theta = jnp.linspace(0.0, 2.0 * jnp.pi, 201) circle = RADIUS * jnp.stack([jnp.cos(theta), jnp.sin(theta)], axis=-1) return float(jnp.abs(projector.evaluate(model, circle)).max()) def train(key, operator, models, sampler, n_sensors): """Train one operator on data alone. ⚠ ``only_data=True``: the physical residuals are not even assembled, which is what "full data first" means -- and it is also the diagnostic. An operator that cannot interpolate its own observations will not be rescued by a physics term. Args: key: a random generator state, operator: the operator, models: the training models, sampler: the domain sampler, n_sensors: how many observation points there are. Returns: ``(projector, seconds, n_theta)``. """ space = PhysicNOApproximationSpace( dims={"x": DIM}, list_models=[operator], model_type="x" ) projector = PhysicNOProjector( models, space, sampler, optimizer="Adam", learning_rate=LEARNING_RATE, only_data=True, weights={"data": [1.0]}, ) start = time.perf_counter() _, projector = projector.project( key, space, N_EPOCHS, BATCH_SIZE, n_colloc=0, n_bc_colloc=0, n_dl_colloc=n_sensors, verbose=True, ) return projector, time.perf_counter() - start, int(space.ndof) def main(): """FEM for the data, then the same operator with three deformations.""" key = jax.random.PRNGKey(SEED) disk = Disk2D((0.0, 0.0), RADIUS, is_main_domain=True) # ⚠ The sensors are the latent grid pushed through the ANALYTIC map, and # they are the same for the three variants. A learned deformation samples # the source elsewhere -- that is its business -- but what the three are # scored on must not depend on the thing being compared. # ⚠ The NODES OF THE BASIS -- the most expensive mistake of this file's # history. A sine basis is orthogonal on the INTERIOR points (j+1)/(n+1), # not on cell centres (j+0.5)/n. Measured, reconstructing the source: # # modes cell centres sine nodes # 8 8.04e-02 3.05e-02 # 16 6.96e-02 1.09e-02 # 32 6.00e-02 3.76e-15 <- exact # # On the wrong nodes the reconstruction FLOORS at 6e-2 whatever the mode # count -- which reads as a limit of the basis and is not one. With that # floor inside the encoder, predicting the family MEAN is a reasonable # answer, and it is exactly what every variant did: 0.51 against 0.52 for # the constant predictor. # # ⚠ They are also strictly interior, so they avoid the circle where # `disk_to_square` has an infinite derivative and the Jacobian weight # returns NaN. One choice, two problems. latent_nodes = ( CONTINUOUS_SYNTHESES[BASIS](DIM, (N_GRID,) * DIM, [(0.0, 1.0)] * DIM) .nodes() .reshape(-1, DIM) ) sensors = square_to_disk(latent_nodes) # ⚠ The PHYSICAL area each sensor stands for. The points come from an even # latent grid pushed through the map, so their physical spacing follows # |det J| of that map -- and it is exactly what cancels the Jacobian the # encoder puts back. Getting it wrong is silent: with 1/n alone the # coefficients came out 1024x too small and the operator learned the family # MEAN, scoring the same 0.5 whatever its size. sensor_area = ( jnp.abs(jnp.linalg.det(jax.vmap(jax.jacfwd(square_to_disk))(latent_nodes))) / latent_nodes.shape[0] ) key, sub = jax.random.split(key) mus = MU_LOW + jax.random.uniform(sub, (N_TRAIN + N_TEST, N_MU)) * ( MU_HIGH - MU_LOW ) print(f"FEM Q{FEM_ORDER}, {FEM_CELLS}x{FEM_CELLS} cells on the disk ...") start = time.perf_counter() exact = jax.block_until_ready(reference_values(mus, sensors)) print( f" {time.perf_counter() - start:.1f} s for {len(mus)} PDEs, " f"{sensors.shape[0]} points each" ) models = [ LaplacianOnDisk(disk, mu, (sensors, exact[i])) for i, mu in enumerate(mus) ] train_models, test_models = models[:N_TRAIN], models[N_TRAIN:] train_exact, test_exact = exact[:N_TRAIN], exact[N_TRAIN:] # ⚠ One constant read off the training sources, as every FNO example here # does. Without it the network is fed values it cannot use. input_scale = float( jnp.std(jnp.stack([jax.vmap(GaussianSource(mu))(sensors) for mu in mus[:32]])) ) print(f"input_scale = {input_scale:.4f}") key, key_net, key_ops = jax.random.split(key, 3) print( f"pre-training the invertible map on the analytic one " f"({PRETRAIN_STEPS} steps) ..." ) start = time.perf_counter() pretrained, pretrain_loss = pretrain_deformation(key_net, new_deformation(key_net)) print(f" {time.perf_counter() - start:.1f} s, final MSE {pretrain_loss:.3e}") variants = { "analytic": (InvertibleFunction(disk_to_square, square_to_disk), True), "learned": (new_deformation(key_net), True), "pre-trained": (pretrained, False), } sampler = TensorizedSampler([DomainSampler(disk)], model_type="x", bc=False) results = {} for name, (deformation, learn) in variants.items(): print(f"\n=== {name} ===") operator = make_operator( key_ops, deformation, learn, input_scale, sensors, sensor_area ) projector, seconds, n_theta = train( key_ops, operator, train_models, sampler, sensors.shape[0] ) results[name] = { "train": relative_errors(projector, train_models, train_exact, sensors), "test": relative_errors(projector, test_models, test_exact, sensors), "boundary": boundary_value(projector, test_models[0]), "n_theta": n_theta, "seconds": seconds, "projector": projector, } print("\n### Relative L2 error against the FEM reference") print( f"{'deformation':14s} {'n_theta':>8s} {'train':>10s} {'test':>10s} " f"{'|u| on circle':>14s} {'seconds':>8s}" ) print("-" * 70) for name, r in results.items(): print( f"{name:14s} {r['n_theta']:8d} {r['train'].mean():10.3e} " f"{r['test'].mean():10.3e} {r['boundary']:14.2e} {r['seconds']:8.0f}" ) print( "\n ⚠ `|u| on circle` is not a quality metric, it is a STRUCTURAL one:\n" " machine zero means the boundary condition came from the basis and\n" " the map, not from the loss. Only a map that sends the circle onto\n" " the latent square's edge can give it." ) # ── Figure ─────────────────────────────────────────────────────────────── side = jnp.linspace(-RADIUS, RADIUS, 120) gx, gy = jnp.meshgrid(side, side, indexing="ij") plot_points = jnp.stack([gx.ravel(), gy.ravel()], axis=-1) inside = (gx.ravel() ** 2 + gy.ravel() ** 2) <= RADIUS**2 figure, axes = plt.subplots( 1, len(results) + 1, figsize=(4.2 * (len(results) + 1), 4.0), constrained_layout=True, ) fem = reference_values(mus[N_TRAIN : N_TRAIN + 1], plot_points)[0, :, 0] fem = jnp.where(inside, fem, jnp.nan) image = axes[0].pcolormesh( np.asarray(gx), np.asarray(gy), np.asarray(fem).reshape(120, 120), shading="auto", ) figure.colorbar(image, ax=axes[0]) axes[0].set_title("FEM reference") for axis, (name, r) in zip(axes[1:], results.items()): predicted = r["projector"].evaluate(test_models[0], plot_points)[:, 0] error = jnp.where(inside, jnp.abs(predicted - fem), jnp.nan) image = axis.pcolormesh( np.asarray(gx), np.asarray(gy), np.asarray(error).reshape(120, 120), shading="auto", ) figure.colorbar(image, ax=axis) axis.set_title(f"|error| -- {name}") for axis in axes: axis.set_aspect("equal") output = __file__.replace(".py", ".png") figure.savefig(output, dpi=110) print(f"\nfigure: {output}") plt.show() if __name__ == "__main__": main()