r"""Physics-informed FNO on the 2D Laplacian: data + a STRONG residual. Same family as ``laplacian_2d_deep_ritz_unet.py``, same sensors, same manufactured solutions -- and the loss that file could not write:: -Delta u = f pointwise, at sampled collocation points ⚠ **That is the whole reason this operator exists.** The U-Net decodes onto a Q1 expansion, which is C0: its second derivatives vanish inside every cell, so ``-Delta u`` would be identically zero and the loss would see nothing. Deep Ritz works around it by asking only for FIRST derivatives. A truncated spectral series has no such limit -- it is a finite sum of smooth modes, so autodiff differentiates it to any order at any point. Two things come for free here, and neither is available to the U-Net: * **the strong residual**, above; * **the boundary condition, exactly.** With a SINE basis every mode vanishes on the boundary, so the expansion does too. The catch is the projection ``Q``: it acts after the series and its bias makes ``Q(0) != 0``. So the example passes an unbiased linear projection, and then ``u = 0`` on the boundary to machine precision -- measured below. No boundary residual, no boundary weight to tune, no penalty fighting the interior. ⚠ **One honesty note about this family.** The manufactured solutions are sums of ``sin(k pi x) sin(l pi y)``, so they lie EXACTLY in the sine basis. That is a genuine advantage of matching the basis to the problem, but it flatters this particular case: on a family that is not built from sine modes, the basis would still give the boundary condition for free and would no longer contain the solution. The Fourier and cosine bases are one argument away if you want to see the difference. Run: python laplacian_2d_strong_residual_fno.py """ 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.neural_operator.data_for_no.grid_data import GridData from scimba_jax.neural_operator.physic_no.grid_based.field_encoding import ( interior_rhs_callables, ) from scimba_jax.neural_operator.physic_no.grid_based.physic_informed_fno import ( PhysicInformedFNO, ) 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.numerical_solvers.physic_no_projectors import ( PhysicNOProjector, ) from scimba_jax.nonlinear_approximation.optimizers.losses import ScimbaMSE from scimba_jax.physical_models.abstract_physical_model import AbstractPhysicalModel from scimba_jax.physical_models.abstract_residuals import InteriorResidual from scimba_jax.physical_models.data_residuals import CollocDataResidual from scimba_jax.utils.functional_fields import make_functional_field_class from scimba_jax.utils.scimba_pytree import ScimbaPytree, trainable # ── Parameters, kept at the U-Net example's values where they mean the same ── # ⚠ ALIGNED on `discrete_no/fno_source_to_solution.py`, deliberately: that file # learns the same family from DATA ALONE, so it is the witness that says whether # a shortfall here comes from the architecture or from the regime. Any axis left # different makes the comparison meaningless -- and four of them were. N_GRID = 32 # points per direction; Nyquist caps the modes at N_GRID // 2 # ⚠ Eight, not twenty. Twenty to thirty is the usual range on real problems, but # the witness keeps "twice the modes of the signal" -- the family has three, and # 20 modes cost 102 894 parameters against roughly 16k here. N_MODES_FNO = 8 CHANNELS = 8 N_LAYERS = 4 # 3 full blocks plus the split one, as the witness has 3 blocks BASIS = "sine" # "sine" | "cosine" | "fourier" -- see the module docstring N_MODES = 3 # modes of the manufactured family # ⚠ 256 PDEs and the WHOLE grid supervised, as a FNO is normally trained -- not # 64 PDEs seen at 200 scattered sensors, which was the Deep Ritz setup and gave # 20x less supervision than the witness. N_TRAIN, N_TEST = 256, 64 N_EPOCHS_DATA, N_EPOCHS_PHYSICS, BATCH_SIZE = 1200, 600, 32 LEARNING_RATE = 3.0e-3 # ⚠ Collocation points for a STRONG residual, not for an energy: a pointwise # second-order residual samples a much rougher function than an integral of # first derivatives, and needs more of them. N_COLLOC = 4000 # ⚠ The residual weight is NOT comparable to the Ritz one. Measured on this # family: the raw residual sits around 2.6e4 against 0.17 for the data, so 1e-2 # does not balance them -- the run reports the two terms separately. WEIGHT_DATA, WEIGHT_RESIDUAL = 1.0, 1.0e-2 domain = HypercubeND([(0.0, 1.0), (0.0, 1.0)], is_main_domain=True) grid_data = GridData(2, HypercubeND([(0.0, 1.0), (0.0, 1.0)]), (N_GRID, N_GRID)) source_class = make_functional_field_class("strong_laplacian_source") class StrongLaplacianResidual(InteriorResidual): r"""``-Delta u = f``, pointwise. ⚠ The residual the Q1-decoded U-Net cannot carry: inside a cell its second derivatives are zero, so this would be ``0 = f`` everywhere and the optimiser would chase an unreachable constant. A spectral series has the derivatives. Args: domain: the domain, f_rhs: the source, model_type: the variables; defaults to ``"x"``. """ def __init__(self, domain, f_rhs=None, model_type="x"): super().__init__(domain=domain, size=1, model_type=model_type, f_rhs=f_rhs) def construct_residual(self, *variables): """The left-hand side. Args: *variables: the approximation space's variables. Returns: ``-Delta u`` as a parametric function. """ return -variables[0].laplacian("x") class LaplacianStrongWithData(AbstractPhysicalModel): """2D Laplacian: a strong residual and sensor observations. ⚠ **No boundary residual.** The sine basis makes the expansion vanish on the boundary exactly, so a Dirichlet penalty would only add a term that is already zero -- and a weight to tune for nothing. Args: main_domain: the domain, f_rhs: the source, data: the pair (sensor points, observed values), model_type: the variables; defaults to ``"x"``. """ def __init__(self, main_domain, f_rhs=None, data=(), model_type="x"): super().__init__(main_domain=main_domain) self.physical_residuals = { self.main_domain.get_label(): StrongLaplacianResidual( domain=main_domain, f_rhs=f_rhs, model_type=model_type ) } self.add_data_residual( "data", CollocDataResidual( size=1, model_type=model_type, data=data, batchable_args=False ), ) class UnbiasedProjection(ScimbaPytree): """``Q`` without a bias, so that ``Q(0) = 0``. ⚠ The one piece the boundary condition hangs on. A sine basis gives a series that vanishes on the boundary; a projection WITH a bias then maps that zero to ``b`` and the condition is gone -- silently, since nothing raises and the field merely stops being zero at the edge. Measured at +0.089 on the boundary with the default two-layer projection. Args: in_channels: the latent width, out_channels: the solution's size, key: a random generator state. """ weight = trainable(True) def __init__(self, in_channels: int, out_channels: int, 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 # ── The manufactured family, and the sources DEDUCED from it ───────────────── modes = jnp.arange(1, N_MODES + 1) eigenvalues = jnp.pi**2 * (modes[:, None] ** 2 + modes[None, :] ** 2) def solution_and_source(coefficients): """The solution and its source, both exact. Args: coefficients: the series coefficients, ``(k, l)``. Returns: the pair of callables. """ def basis(x): return ( jnp.sin(jnp.pi * modes * x[0])[:, None] * jnp.sin(jnp.pi * modes * x[1])[None, :] ) def solution(x): return jnp.sum(coefficients * basis(x))[None] def source(x): return jnp.sum(coefficients * eigenvalues * basis(x))[None] return solution, source def make_batch(key, n_models, sensors): """Draw a family of PDEs and the matching observations. Args: key: the generator state, n_models: how many models, sensors: the observation points. Returns: (new key, list of models, exact values at the sensors, list of exact solutions as callables). """ key, subkey = jax.random.split(key) coefficients = jax.random.normal(subkey, (n_models, N_MODES, N_MODES)) coefficients = coefficients / (modes[None, :, None] * modes[None, None, :]) models, exact, solutions = [], [], [] for index in range(n_models): solution, source = solution_and_source(coefficients[index]) values = jax.vmap(solution)(sensors) exact.append(values) solutions.append(solution) models.append( LaplacianStrongWithData( main_domain=domain, f_rhs=source_class(source), data=(sensors, values), ) ) return key, models, jnp.stack(exact), solutions key = jax.random.PRNGKey(0) # ⚠ The WHOLE grid, not scattered sensors. A FNO is supervised on its grid; # sampling 200 points instead was a Deep Ritz habit and starved the operator. sensors = grid_data.grid.reshape(-1, 2) N_SENSORS = sensors.shape[0] key, train_models, train_exact, _ = make_batch(key, N_TRAIN, sensors) key, test_models, test_exact, test_solutions = make_batch(key, N_TEST, sensors) # ── The operator, the space, the projector ─────────────────────────────────── n_fields = PhysicInformedFNO.count_encoded_fields(train_models[0], grid_data) # ⚠ ONE constant, read off the training set, exactly as # `discrete_no/fno_source_to_solution.py` does. The source carries the # Laplacian's EIGENVALUES, so it is about a hundred times larger than the # solution it must produce; handing that to a GELU network unscaled is a known # way to learn very little. Measured below so the number is not a guess. _grid_points = grid_data.grid.reshape(-1, 2) INPUT_SCALE = float( jnp.std( jnp.stack( [ jax.vmap(field)(_grid_points) for model in train_models for field in interior_rhs_callables(model) ] ) ) ) print(f"input_scale (std of the training sources) = {INPUT_SCALE:.3f}") key, key_net, key_projection = jax.random.split(key, 3) operator = PhysicInformedFNO( grid_data, n_fields, 1, key_net, channels=CHANNELS, n_modes=N_MODES_FNO, n_layers=N_LAYERS, basis=BASIS, input_scale=INPUT_SCALE, # ⚠ The three arguments that keep the boundary condition exact: an # unbiased projection, and nothing between the series and it. projection=UnbiasedProjection(CHANNELS, 1, key_projection), final_channel_mlp=False, final_activation=False, ) space = PhysicNOApproximationSpace( dims={"x": 2}, list_models=[operator], model_type="x" ) # ⚠ `bc=False`: there is no boundary residual to sample for. sampler = TensorizedSampler([DomainSampler(domain)], model_type="x", bc=False) interior_label = domain.get_label() losses = { interior_label: (ScimbaMSE(WEIGHT_RESIDUAL),), "data": (ScimbaMSE(WEIGHT_DATA),), } #: Stage one: ``only_data=True`` drops the physical residuals entirely, so the #: Hessians are not even computed -- cheaper than weighting them to zero. data_projector = PhysicNOProjector( train_models, space, sampler, optimizer="Adam", learning_rate=LEARNING_RATE, only_data=True, weights={"data": [WEIGHT_DATA]}, ) def relative_errors(trained, models, exact, points) -> np.ndarray: """Relative L2 error per model, at the given points. Args: trained: the trained projector, models: the physical models, exact: the exact values, ``(n_models, n_points, 1)``, points: the evaluation points. Returns: ``(n_models,)``. """ errors = [] for index, model in enumerate(models): predicted = trained.evaluate(model, points) reference = exact[index] errors.append( float(jnp.linalg.norm(predicted - reference) / jnp.linalg.norm(reference)) ) return np.asarray(errors) def boundary_deviation(trained, model) -> float: """How far the prediction is from zero ON the boundary. ⚠ The measurement this file exists to report. With a sine basis and an unbiased projection it should be at machine precision -- not small, ZERO -- because no term of the expansion is non-zero there. Anything else means something was inserted after the series. Args: trained: the trained projector, model: any model of the family. Returns: the largest absolute value found on the boundary. """ edge = jnp.linspace(0.0, 1.0, 101) zeros, ones = jnp.zeros_like(edge), jnp.ones_like(edge) points = jnp.concatenate( [ jnp.stack([edge, zeros], axis=-1), jnp.stack([edge, ones], axis=-1), jnp.stack([zeros, edge], axis=-1), jnp.stack([ones, edge], axis=-1), ] ) return float(jnp.abs(trained.evaluate(model, points)).max()) def report(trained, label: str) -> np.ndarray: """Print the train and test errors, and the boundary deviation. Args: trained: the trained projector, label: what to call this run. Returns: the test errors. """ train_errors = relative_errors(trained, train_models, train_exact, sensors) test_errors = relative_errors(trained, test_models, test_exact, sensors) print(f"\n### {label}") print(f" train relative L2 : {train_errors.mean():.3e}") print(f" test relative L2 : {test_errors.mean():.3e}") print( f" |u| on the boundary : {boundary_deviation(trained, test_models[0]):.3e}" f" (basis {BASIS!r}; zero means the basis imposed it, not the loss)" ) # ⚠ Per TERM, and it is the number that says what to tune. A total loss # says nothing about whether the physics or the data is driving the fit: # with weights four orders apart, one of the two can be doing all the work # while the other is already converged -- or ignored. history = trained.losses.losses_history # ⚠ NOT weighted -- verified: `1e-2 * interior + 1.0 * data` reproduces the # total exactly on two separate runs (2.780 and 0.063). Reading these as # weighted made the total look inconsistent with its own terms. print(" loss by term (last epoch, UNWEIGHTED):") for label, values in history.items(): if label == "total": continue print(f" {label:24s} {float(jnp.asarray(values)[-1].sum()):.3e}") print(f" {'total':24s} {float(jnp.asarray(history['total'])[-1]):.3e}") return test_errors if __name__ == "__main__": # ── Stage one: data only ───────────────────────────────────────────────── key, data_projector = data_projector.project( key, space, N_EPOCHS_DATA, BATCH_SIZE, n_colloc=0, n_bc_colloc=0, n_dl_colloc=N_SENSORS, verbose=True, ) data_errors = report(data_projector, "stage 1 -- data only") # ── Stage two: the physics, from where stage one left off ──────────────── projector = PhysicNOProjector( train_models, data_projector.space, sampler, optimizer="Adam", learning_rate=LEARNING_RATE, losses=losses, ) key, projector = projector.project( key, data_projector.space, N_EPOCHS_PHYSICS, BATCH_SIZE, n_colloc=N_COLLOC, n_bc_colloc=0, n_dl_colloc=N_SENSORS, verbose=True, ) test_errors = report(projector, "stage 2 -- data + strong residual") print("\n### What the physics term bought") print(f" test error, data only : {data_errors.mean():.3e}") print(f" test error, data + physics : {test_errors.mean():.3e}") print( f" ratio : {data_errors.mean() / test_errors.mean():.2f}x" ) # ⚠ **No second-order refinement here**, and it is not an omission: the # U-Net example finishes with SS-Broyden, which keeps a DENSE inverse # Hessian. This operator has enough parameters that the matrix would not # fit in memory -- the process is killed by the OS, silently and without a # traceback. Measured; the arithmetic is printed below. A limited-memory # method (L-BFGS) is the option if a second-order pass is wanted. print( f"\n n_theta = {operator.ndof()} -> a dense inverse Hessian would " f"be {operator.ndof() ** 2 * 8 / 1e9:.0f} GB" ) # ── Figures ────────────────────────────────────────────────────────────── side = jnp.linspace(0.0, 1.0, 80) mesh_x, mesh_y = jnp.meshgrid(side, side, indexing="ij") plot_points = jnp.stack([mesh_x.ravel(), mesh_y.ravel()], axis=-1) figure, axes = plt.subplots(3, 3, figsize=(12, 11), constrained_layout=True) for column in range(3): predicted = projector.evaluate(test_models[column], plot_points)[:, 0] reference = jax.vmap(test_solutions[column])(plot_points)[:, 0] for row, (values, title) in enumerate( ( (reference, "exact"), (predicted, "FNO, strong residual"), (jnp.abs(predicted - reference), "|error|"), ) ): image = axes[row, column].pcolormesh( np.asarray(mesh_x), np.asarray(mesh_y), np.asarray(values).reshape(80, 80), shading="auto", ) figure.colorbar(image, ax=axes[row, column]) axes[row, column].set_title(f"{title} -- test {column}", fontsize=9) output = __file__.replace(".py", ".png") figure.savefig(output, dpi=110) print(f"\nfigure: {output}") plt.show()