r"""A four-level graph U-Net on JET: ``f |-> u`` for the Laplacian, GNN and GNO. -Delta u = f in the JET poloidal cross-section, u = 0 on the wall The data is the same as ``gino_jet_laplacian.py``: ``f`` a random mixture of three Gaussians, ``u`` a batched P1 finite-element solve, both read at the cell centres of a nested JET hierarchy. The operator is :class:`~scimba_jax.neural_operator.layers.graph_based.GraphUNet` -- residual slide layers at every level, a geometric pooling down, an unpooling back up with a skip -- the graph reading of the grid U-Net. ⚠ **Why a U-Net and not a chain.** ``-Delta`` is local; its inverse is not: ``u`` at a point depends on ``f`` everywhere. A chain of ball layers of radius ``r`` needs a depth of ``diameter / r`` to carry that dependence; a U-Net carries it in one pass, through the coarse levels where the ball -- doubled at every coarser level -- spans a good part of the domain. The two modes, and what each is tested on ------------------------------------------ * ``"gnn"`` -- edges for the layers, PARENTHOOD for the pooling (the L2 average of a cell's children, the exact nesting). Trained on the levels (0, 1, 2, 3) of a five-level hierarchy and scored there; * ``"gno"`` -- balls for the layers, a spatial BALL for the pooling (the L2 average of the fine cells within ``r`` of a coarse centre). Trained on (0, 1, 2, 3), scored there, then the resolution test: the SAME weights pinned to the levels (1, 2, 3, 4), four times finer at every level and never seen. Every ball keeps its physical radius, every pooling its physical reach. The truth there is a finite-element solve ON the finest mesh. The ``"gnn"`` network is read on (1, 2, 3, 4) too, for the record: nothing about it is meant to survive the change of mesh. Measured (kernel + mlp, width 16, 2 blocks per level, 400 epochs with a cosine decay, 128 training sources; relative L2 on 32 test sources):: levels mode train trained tuple one level finer, never seen 4 gnn 93 s 4.15e-01 4.28e-01 4 gno 333 s 7.11e-02 9.95e-02 3 gnn 78 s 2.56e-01 5.45e-01 3 gno 279 s 1.99e-01 2.86e-01 mean predictor 4.4e-01 ⚠ The fourth level is what made the gno an operator worth having: 2.0e-1 to 7.1e-2 for the 868 pairs of a 57-cell bottom, because ``-Delta``'s inverse is global and the bottom is the only level that sees far. The gno is now at 3x the GINO of ``gino_jet_laplacian.py`` (2.3e-2, a global FNO, 2.3M parameters against 143k). The gnn does not benefit: on 57 cells two edges of reach are 0.8 of a 3.5-wide domain, and parenthood carries no geometry -- it sits at the mean predictor. What the example is for is the last column: the gno holds under a change of mesh, the gnn never did. Run: python gnn_unet_jet_laplacian.py """ import math import sys import time from pathlib import Path import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scipy.spatial import cKDTree sys.path.insert(0, str(Path(__file__).resolve().parent)) from jet_hierarchy import jet_hierarchy, jet_wall # noqa: E402 from scimba_jax.linear_approximation.basis.analytic_bases import ( # noqa: E402 local_lagrange_basis, local_lagrange_basis_by_logical, ) from scimba_jax.linear_approximation.basis.dof_map import ( # noqa: E402 UnstructuredLagrangeDofMap, ) from scimba_jax.linear_approximation.basis.general_bases import ( AnalyticBasis, # noqa: E402 ) from scimba_jax.linear_approximation.galerkin.fem.elliptic_fe_scheme import ( # noqa: E402 EllipticFEscheme, ) from scimba_jax.linear_approximation.meshes.unstructured_mesh import ( # noqa: E402 UnstructuredMesh, ) from scimba_jax.linear_approximation.quad.gauss_quad import ( # noqa: E402 UnitSquareTensorized, ) from scimba_jax.linear_approximation.variables.variables_fe import ( VariablesFE, # noqa: E402 ) from scimba_jax.neural_operator.data_for_no.hierarchic_mesh_data import ( # noqa: E402 HierarchicMeshData, ) from scimba_jax.neural_operator.layers.graph_based import ( # noqa: E402 GraphOperator, GraphUNet, ) from scimba_jax.nonlinear_approximation.numerical_solvers.no_projectors import ( # noqa: E402 NOProjector, ) from scimba_jax.nonlinear_approximation.optimizers.schedules import ( # noqa: E402 CosineDecay, ) from scimba_jax.physical_models.abstract_physical_weak_model import ( # noqa: E402 AbstractPhysicalWeakModel, ) from scimba_jax.physical_models.classical_weakform.laplacian_weak_form import ( # noqa: E402 LaplacianWeakForm, ) from scimba_jax.utils.scimba_pytree import ScimbaPytree # noqa: E402 N_GAUSSIANS = 3 N_TRAIN, N_TEST = 128, 32 N_EPOCHS, BATCH_SIZE, LEARNING_RATE = 400, 16, 3e-3 # One optimizer step per mini-batch: the schedule is written in steps. N_STEPS = N_EPOCHS * math.ceil(N_TRAIN / BATCH_SIZE) HIDDEN, BLOCKS_PER_LEVEL = 16, 2 # ⚠ The kernel message, in both modes -- the SAME network, two neighbourhoods. # It is 3x cheaper than the pair MLP on a ball (kappa(x_j - x_i) is evaluated # once for all samples, outside the vmap) and, measured, the better gnn too at # this budget: MeshGraphNets' MLP message sits on a loss plateau for over a # thousand epochs (4.4e-1 after 400, the mean predictor's level), the kernel # gives 2.6e-1 in 78 s. MESSAGE = {"gnn": "kernel", "gno": "kernel"} UPDATE = "mlp" # The ball on the finest level of a triple; it doubles at every coarser # level, so the network's physical receptive field is the same whether it # reads the levels (0, 1, 2) or (1, 2, 3). RADIUS, RADIUS_GROWTH = 0.1, 2.0 MESH_SIZE = 0.4 # Gmsh's length at level 0; level k has cells of MESH_SIZE / 2**k # ⚠ The U-Net's depth is its REACH. -Delta's inverse is global, and the # coarsest level is where the network sees far: with four levels the bottom # is ~50 cells with balls of radius 0.8 -- nearly every pair, for nothing. # Three levels (bottom at 203 cells, r = 0.4) gave 1.99e-1 / 2.86e-1. N_LEVELS = 4 FINEST = N_LEVELS - 1 # the training tuple's finest level TRAIN_LEVELS = tuple(range(N_LEVELS)) TEST_LEVELS = tuple(range(1, N_LEVELS + 1)) class GaussianMixture(ScimbaPytree): """``f(x) = sum_i A_i exp(-|x - c_i|^2 / 2 sigma_i^2)`` -- a pytree, so it batches. Args: centers: ``(n_gaussians, 2)``, amplitudes: ``(n_gaussians,)``, sigmas: ``(n_gaussians,)``. """ def __init__(self, centers, amplitudes, sigmas): self.centers = jnp.asarray(centers) self.amplitudes = jnp.asarray(amplitudes) self.sigmas = jnp.asarray(sigmas) def __call__(self, x): """The source at one point. Args: x: a point, ``(2,)``. Returns: a scalar. """ squared = jnp.sum((x - self.centers) ** 2, axis=1) return jnp.sum(self.amplitudes * jnp.exp(-squared / (2.0 * self.sigmas**2))) class FiniteElementSolver: """A batched P1 Laplacian solve on one level's mesh, read at cell centres. The matrix does not depend on the source, so it is factorised once and only the right-hand side is batched -- the pattern of the GINO example. Args: level: the graph level whose mesh carries the finite elements. """ def __init__(self, level): mesh = UnstructuredMesh( nodes=np.asarray(level.mesh.nodes), cells=np.asarray(level.mesh.cells), ref_quad=UnitSquareTensorized(dim=2, order=4), order=1, ) basis = AnalyticBasis( nb_basis=4, out_dim=1, mesh=mesh, basis_type="scalar", local_basis_by_logical=lambda y, i, m: local_lagrange_basis_by_logical( y, i, m, order=1, out_dim=1 ), local_basis=lambda y, i, m: local_lagrange_basis( y, i, m, order=1, out_dim=1 ), ) self.variables = VariablesFE( basis=basis, nb_variables=1, dof_map=UnstructuredLagrangeDofMap ) self.centres = level.point_cloud reference = self._scheme( GaussianMixture( jnp.zeros((N_GAUSSIANS, 2)), jnp.zeros(N_GAUSSIANS), jnp.ones(N_GAUSSIANS), ) ) self._dofs_init = reference._initial_dofs() self._factorisation = EllipticFEscheme.factorise(reference, self._dofs_init) self._back_solve = EllipticFEscheme._make_back_solve_fn() def _scheme(self, source): model = AbstractPhysicalWeakModel.from_weak_form( LaplacianWeakForm(dim=2, f=source), dirichlet=lambda x: jnp.zeros(1) ) return EllipticFEscheme(model, self.variables) def __call__(self, centers, amplitudes, sigmas): """``(f, u)`` at the cell centres for a batch of sources. Args: centers: ``(batch, n_gaussians, 2)``, amplitudes: ``(batch, n_gaussians)``, sigmas: ``(batch, n_gaussians)``. Returns: ``(f, u)``, each ``(batch, n_cells, 1)``. """ @jax.jit def solve(c, a, s): def one(c1, a1, s1): scheme = self._scheme(GaussianMixture(c1, a1, s1)) dofs = self._back_solve( scheme, self._dofs_init, self._factorisation.lu, self._factorisation.pivots, ) u = jax.vmap( lambda p: VariablesFE._classical_local_evaluate_pure( self.variables, dofs, p ) )(self.centres) f = jax.vmap(GaussianMixture(c1, a1, s1))(self.centres)[:, None] return f, u return jax.vmap(one)(c, a, s) return solve(centers, amplitudes, sigmas) def draw_sources(key, n, candidates, wall_tree, margin=0.25): """Random Gaussian mixtures centred away from the wall. Args: key: a random state, n: how many sources, candidates: points to draw centres among, ``(m, 2)``, wall_tree: a KD-tree on the wall's points, margin: the least distance of a centre to the wall. Returns: ``(centers, amplitudes, sigmas)``. """ keys = jax.random.split(key, 3) distance, _ = wall_tree.query(np.asarray(candidates)) allowed = jnp.asarray(np.nonzero(distance > margin)[0]) picked = jax.random.choice(keys[0], allowed, (n * N_GAUSSIANS,), replace=True) centers = candidates[picked].reshape(n, N_GAUSSIANS, 2) amplitudes = jax.random.uniform(keys[1], (n, N_GAUSSIANS), minval=0.5, maxval=2.0) sigmas = jax.random.uniform(keys[2], (n, N_GAUSSIANS), minval=0.20, maxval=0.40) return centers, amplitudes, sigmas @jax.jit def predict(operator, inputs): """The operator on a batch, jitted with the operator as an ARGUMENT. ⚠ ``jax.jit(jax.vmap(operator))`` would close over the operator, and a closed-over pytree is baked into the HLO as constants -- here the whole hierarchy, a million-pair ball included, which XLA then tries to constant-fold at compile time (measured: minutes on the test hierarchy). Passed as an argument, the hierarchy is an operand and the compile is a second. Args: operator: a ``GraphOperator``, inputs: ``(batch, n_cells, channels)``. Returns: ``(batch, n_cells, out_channels)``. """ return jax.vmap(operator)(inputs) def relative_error(predicted, truth): """Relative L2 error per sample, ``(batch,)``.""" flat = predicted.reshape(predicted.shape[0], -1) - truth.reshape(truth.shape[0], -1) return jnp.linalg.norm(flat, axis=1) / jnp.linalg.norm( truth.reshape(truth.shape[0], -1), axis=1 ) # ══════════════════════════════════════════════════════════════════════════════ # 1. Nested levels: train on the first N_LEVELS, test the operator one level finer # ══════════════════════════════════════════════════════════════════════════════ radii = [RADIUS * RADIUS_GROWTH**k for k in range(N_LEVELS)] # finest ... coarsest # Level k is the finest of one triple and a coarser level of another, so it # carries every radius a triple can ask of it; the cross-level balls are the # finer level's radius of each transition. hierarchy = jet_hierarchy( n_levels=N_LEVELS + 1, mesh_size=MESH_SIZE, # Level k is the finest of one N_LEVELS-tuple and one level coarser in the # next, so it carries the radius of both roles. ball_radii=[ tuple(radii[i] for i in (N_LEVELS - 1 - k, N_LEVELS - k) if 0 <= i < N_LEVELS) for k in range(N_LEVELS + 1) ], cross_radii=tuple(radii[: N_LEVELS - 1]), ) levels = hierarchy.levels train_hierarchy = HierarchicMeshData( levels[:N_LEVELS], hierarchy.parents[: N_LEVELS - 1], hierarchy.cross_radii ) test_hierarchy = HierarchicMeshData( levels[1:], hierarchy.parents[1:], hierarchy.cross_radii ) for k, level in enumerate(levels): balls = ", ".join( f"r={r:.2f}: {level.ball(r)[0].shape[0]}" for r in level.ball_radii ) print(f"level {k}: {level.n_nodes} cells; ball pairs {balls}") wall_tree = cKDTree(jet_wall()) key = jax.random.key(0) draws_train = draw_sources( jax.random.fold_in(key, 1), N_TRAIN, levels[FINEST].point_cloud, wall_tree ) draws_test = draw_sources( jax.random.fold_in(key, 2), N_TEST, levels[FINEST].point_cloud, wall_tree ) started = time.perf_counter() solve_train = FiniteElementSolver(levels[FINEST]) f_train, u_train = solve_train(*draws_train) f_test, u_test = solve_train(*draws_test) solve_fine = FiniteElementSolver(levels[FINEST + 1]) f_fine, u_fine = solve_fine(*draws_test) print(f"finite-element data in {time.perf_counter() - started:.0f} s") F_SCALE, U_SCALE = float(jnp.std(f_train)), float(jnp.std(u_train)) baseline = float( jnp.mean( relative_error( jnp.broadcast_to(jnp.mean(u_train, axis=0), u_test.shape), u_test ) ) ) # ══════════════════════════════════════════════════════════════════════════════ # 2. The U-Net in both modes # ══════════════════════════════════════════════════════════════════════════════ results = {} for mode in ("gnn", "gno"): network = GraphUNet( 2, 1, 1, HIDDEN, n_levels=N_LEVELS, mode=mode, radius=RADIUS, radius_growth=RADIUS_GROWTH, blocks_per_level=BLOCKS_PER_LEVEL, message_kind=MESSAGE[mode], update_kind=UPDATE, message_kwargs={"hidden_sizes": [32]}, key=jax.random.fold_in(key, 3), ) operator = GraphOperator(network, train_hierarchy) print( f"\nGraphUNet mode={mode} ({MESSAGE[mode]} + {UPDATE}, pool {network.pools[0].kind}, " f"unpool {network.unpools[0].kind}): {network.ndof()} parameters" ) # ⚠ At a constant rate the loss sat on a plateau (2.2e-1) from epoch 400 to # 1000 before dropping; the cosine decay is what makes 400 epochs enough. projector = NOProjector( operator, (f_train / F_SCALE, u_train / U_SCALE), learning_rate=LEARNING_RATE, schedule=CosineDecay(N_STEPS, final_ratio=0.05), ) started = time.perf_counter() _, projector = projector.project( jax.random.fold_in(key, 4), operator, N_EPOCHS, batch_size=BATCH_SIZE, tqdm_desc=f" {mode}", ) trained = projector.operator seconds = time.perf_counter() - started prediction = predict(trained, f_test / F_SCALE) * U_SCALE error = float(jnp.mean(relative_error(prediction, u_test))) on_fine = GraphOperator(trained.network, test_hierarchy) started = time.perf_counter() prediction_fine = predict(on_fine, f_fine / F_SCALE) * U_SCALE error_fine = float(jnp.mean(relative_error(prediction_fine, u_fine))) fine_seconds = time.perf_counter() - started results[mode] = (trained, prediction, prediction_fine, error, error_fine) print(f" trained in {seconds:.0f} s") print(f" relative L2 on test, levels {TRAIN_LEVELS}: {error:.3e}") print( f" relative L2 on test, levels {TEST_LEVELS}: {error_fine:.3e} " f"({fine_seconds:.0f} s, {levels[FINEST + 1].n_nodes} cells never seen)" ) print(f"\n mean-predictor: {baseline:.3e}") # ══════════════════════════════════════════════════════════════════════════════ # 3. One source, both resolutions, gno # ══════════════════════════════════════════════════════════════════════════════ _, prediction, prediction_fine, error, error_fine = results["gno"] sample = 0 # ⚠ One colour scale for every u panel, or an amplitude off by 30 % is # invisible -- each panel autoscaled would show only the shape. u_max = float(jnp.max(jnp.abs(u_fine[sample]))) mistake = jnp.abs(prediction_fine[sample] - u_fine[sample]) panels = [ (f"f, level {FINEST}", levels[FINEST], f_test[sample], None), (f"u FEM, level {FINEST}", levels[FINEST], u_test[sample], u_max), (f"u gno U-Net, level {FINEST}", levels[FINEST], prediction[sample], u_max), (f"u FEM, level {FINEST + 1}", levels[FINEST + 1], u_fine[sample], u_max), ( f"u gno U-Net, level {FINEST + 1} (never seen)", levels[FINEST + 1], prediction_fine[sample], u_max, ), ( f"|error|, level {FINEST + 1} (rel. L2 {error_fine:.2f})", levels[FINEST + 1], mistake, None, ), ] figure, axes = plt.subplots(1, 6, figsize=(26, 4.8)) wall = jet_wall() for axis, (title, level, values, limit) in zip(axes, panels): points = np.asarray(level.point_cloud) picture = axis.scatter( points[:, 0], points[:, 1], c=np.asarray(values)[:, 0], s=4, cmap="viridis" if limit is not None else "magma", vmin=0.0 if limit is not None else None, vmax=limit, ) axis.plot(*np.vstack([wall, wall[:1]]).T, "k-", linewidth=0.8) axis.set_title(title) axis.set_aspect("equal") figure.colorbar(picture, ax=axis, shrink=0.8) figure.suptitle( f"-Delta u = f on JET: gno U-Net trained on levels {TRAIN_LEVELS}, read on {TEST_LEVELS} " f"-- rel. L2 {error:.2f} there, {error_fine:.2f} here" ) figure.tight_layout() figure.savefig(Path(__file__).with_suffix(".png"), dpi=110) plt.show()