r"""Graph networks on JET: learn ``u_0 |-> u(T)`` for an advected, diffused Gaussian. The operator is exact and cheap to write down, which is what makes it a good first test of the graph layers -- nothing else can be blamed. A Gaussian of amplitude ``A``, width ``sigma_0`` and centre ``x_0``, advected at velocity ``v`` and diffused with coefficient ``D`` for a time ``T``, is still a Gaussian:: u(x, T) = A sigma_0^2 / (sigma_0^2 + 2 D T) exp( -|x - x_0 - v T|^2 / (2 (sigma_0^2 + 2 D T)) ) as long as it never reaches the wall, which the draw of ``x_0`` guarantees. The map ``u_0 |-> u(T)`` is a convolution with a Gaussian shifted by ``v T``: a kernel integral over a ball of radius ``|v T|`` plus a few widths. Two modes, the same network ----------------------------- :class:`~scimba_jax.neural_operator.layers.graph_based.GraphNetwork` stacks residual slide layers -- a message ``phi`` aggregated over a neighbourhood, then a pointwise update ``gamma``. The MODE is the neighbourhood: * ``"gnn"`` -- the mesh's edges. Trained and scored on the coarse level 0 of a nested JET hierarchy; then read on the four times finer level 1, where it has no reason to work: its neighbourhood is two edges, whose length has halved; * ``"gno"`` -- the ball ``B(x_i, r)`` with the mesh's quadrature weights. Trained on level 0, scored there, then read on level 1: the same physical radius, integrated with the finer level's own weights -- the operator test. Every message of the slides is run in both modes with the same residual MLP update, the same width and the same training. The table is what the two modes are for. Run: python gnn_jet_advection_diffusion.py """ import itertools 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.neural_operator.layers.graph_based import ( # noqa: E402 MESSAGE_KINDS, MODES, GraphNetwork, GraphOperator, ) from scimba_jax.nonlinear_approximation.numerical_solvers.no_projectors import ( # noqa: E402 NOProjector, ) # ── The physics ───────────────────────────────────────────────────────────── VELOCITY = jnp.array([0.15, 0.10]) DIFFUSION = 0.005 T_FINAL = 1.0 SIGMA_RANGE = (0.10, 0.18) AMPLITUDE_RANGE = (0.5, 1.5) # ⚠ The ball is sized to hold about as many neighbours as the mesh's edges # do (4.3 per cell against 3.8 on level 0), so the two modes cost the same # and differ ONLY in what a neighbourhood is: two edges, or a physical radius # with quadrature weights -- which is 18.6 neighbours on the finer level. Four # layers compose to a reach of 0.52, past the shift |v T| = 0.18 and the # diffusion width sqrt(2 D T) = 0.1 of the operator. RADIUS = 0.13 # The pushed operator: 9 cells per ball on level 0, twice the epochs. RADIUS_WIDE, N_EPOCHS_WIDE, MESSAGE_WIDE = 0.2, 2000, "kernel" # A centre stays this far from the wall: no mass leaves the domain by T. WALL_MARGIN = 3 * SIGMA_RANGE[1] + float(jnp.linalg.norm(VELOCITY)) * T_FINAL # ── The training ──────────────────────────────────────────────────────────── N_TRAIN, N_TEST = 128, 32 N_EPOCHS, BATCH_SIZE, LEARNING_RATE = 1000, 16, 3e-3 HIDDEN, N_BLOCKS, UPDATE = 16, 4, "mlp" MESH_SIZE = 0.13 def gaussian(points, center, amplitude, sigma, t): """The exact solution at time ``t``, read at ``points``. Args: points: ``(n, 2)``, center: ``(2,)``, amplitude: the initial amplitude, sigma: the initial width, t: the time. Returns: ``(n, 1)``. """ variance = sigma**2 + 2.0 * DIFFUSION * t shift = points - center - VELOCITY * t scale = amplitude * sigma**2 / variance return (scale * jnp.exp(-jnp.sum(shift * shift, axis=-1) / (2.0 * variance)))[ :, None ] def draw_centers(key, n, candidates, wall_tree): """Centres drawn among cell centres far enough from the wall. Args: key: a random state, n: how many, candidates: the cell centres, ``(m, 2)``, wall_tree: a KD-tree on the wall's points. Returns: ``(n, 2)``. """ distance, _ = wall_tree.query(np.asarray(candidates)) allowed = jnp.asarray(np.nonzero(distance > WALL_MARGIN)[0]) picked = jax.random.choice(key, allowed, (n,), replace=True) return candidates[picked] def make_dataset(key, n, level, wall_tree): """``(u_0, u_T)`` on the cell centres of one level. Args: key: a random state, n: how many samples, level: the graph level, wall_tree: a KD-tree on the wall's points. Returns: ``(inputs, targets)``, each ``(n, n_cells, 1)``, and the draws. """ keys = jax.random.split(key, 3) centers = draw_centers(keys[0], n, level.point_cloud, wall_tree) sigmas = jax.random.uniform( keys[1], (n,), minval=SIGMA_RANGE[0], maxval=SIGMA_RANGE[1] ) amplitudes = jax.random.uniform( keys[2], (n,), minval=AMPLITUDE_RANGE[0], maxval=AMPLITUDE_RANGE[1] ) at = jax.vmap(gaussian, in_axes=(None, 0, 0, 0, None)) return ( at(level.point_cloud, centers, amplitudes, sigmas, 0.0), at(level.point_cloud, centers, amplitudes, sigmas, T_FINAL), (centers, amplitudes, sigmas), ) @jax.jit def predict(operator, inputs): """The operator on a batch -- the operator as an argument of the jit, so the mesh is an operand and not a constant folded into the executable. 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 ) def train(mode, message_kind, seed, radius=RADIUS, n_epochs=N_EPOCHS): """One network, trained on the coarse level, scored on both levels. Args: mode: ``"gnn"`` or ``"gno"``, message_kind: the message, seed: the network's seed, radius: the ball's radius (gno), n_epochs: how long. Returns: ``(trained operator, ndof, seconds, error on level 0, error on level 1)``. """ network = GraphNetwork( 2, 1, 1, HIDDEN, N_BLOCKS, mode, radius, message_kind, UPDATE, message_kwargs={"hidden_sizes": [32]} if message_kind in ("kernel", "mp") else {}, key=jax.random.fold_in(key, seed), ) operator = GraphOperator(network, coarse) projector = NOProjector(operator, (u0_train, uT_train), learning_rate=LEARNING_RATE) started = time.perf_counter() _, projector = projector.project( jax.random.fold_in(key, 3), operator, n_epochs, batch_size=BATCH_SIZE, tqdm_desc=f" {mode} {message_kind:6s} r={radius}", ) seconds = time.perf_counter() - started trained = projector.operator on_coarse = float(jnp.mean(relative_error(predict(trained, u0_test), uT_test))) on_fine = GraphOperator(trained.network, fine) on_fine_error = float(jnp.mean(relative_error(predict(on_fine, u0_fine), uT_fine))) return trained, network.ndof(), seconds, on_coarse, on_fine_error # ══════════════════════════════════════════════════════════════════════════════ # 1. The hierarchy: train on level 0, test the operator on level 1 # ══════════════════════════════════════════════════════════════════════════════ hierarchy = jet_hierarchy( n_levels=2, mesh_size=MESH_SIZE, ball_radii=(RADIUS, RADIUS_WIDE) ) coarse, fine = hierarchy.levels wall_tree = cKDTree(jet_wall()) print( f"JET: level 0 {coarse.n_nodes} cells ({coarse.ball(RADIUS)[0].shape[0]} ball pairs), " f"level 1 {fine.n_nodes} cells ({fine.ball(RADIUS)[0].shape[0]} ball pairs)" ) key = jax.random.key(0) u0_train, uT_train, _ = make_dataset( jax.random.fold_in(key, 1), N_TRAIN, coarse, wall_tree ) u0_test, uT_test, draws = make_dataset( jax.random.fold_in(key, 2), N_TEST, coarse, wall_tree ) # The SAME test draws read on the fine level: the operator test. at = jax.vmap(gaussian, in_axes=(None, 0, 0, 0, None)) u0_fine, uT_fine = ( at(fine.point_cloud, *draws, 0.0), at(fine.point_cloud, *draws, T_FINAL), ) baseline = float(jnp.mean(relative_error(u0_test, uT_test))) print(f"identity baseline (predict u_0): relative L2 {baseline:.3e}\n") # ══════════════════════════════════════════════════════════════════════════════ # 2. Every message, in both modes # ══════════════════════════════════════════════════════════════════════════════ rows = [] print( f"{'mode':5s} {'message':8s} {'ndof':>7s} {'train':>7s} " f"{'L2 level 0':>11s} {'L2 level 1':>11s}" ) for index, (mode, message_kind) in enumerate(itertools.product(MODES, MESSAGE_KINDS)): trained, ndof, seconds, on_coarse, on_fine = train(mode, message_kind, 10 + index) rows.append((mode, message_kind, ndof, seconds, on_coarse, on_fine)) print( f"{mode:5s} {message_kind:8s} {ndof:7d} {seconds:6.0f}s " f"{on_coarse:11.3e} {on_fine:11.3e}" ) print(f"{'identity':14s} {'':>7s} {'':>7s} {baseline:11.3e}") # ══════════════════════════════════════════════════════════════════════════════ # 3. The operator side, pushed: a wider ball, longer -- then seen on both levels # ══════════════════════════════════════════════════════════════════════════════ print(f"\ngno {MESSAGE_WIDE} at r={RADIUS_WIDE}, {N_EPOCHS_WIDE} epochs") trained, ndof, seconds, on_coarse, on_fine = train( "gno", MESSAGE_WIDE, 20, radius=RADIUS_WIDE, n_epochs=N_EPOCHS_WIDE ) print( f"{'gno':5s} {MESSAGE_WIDE:8s} {ndof:7d} {seconds:6.0f}s " f"{on_coarse:11.3e} {on_fine:11.3e}" ) sample = 0 # One colour scale for every panel -- an amplitude off by 20 % is invisible # otherwise -- and the exact solution on level 1 next to the prediction. u_max = float(jnp.max(uT_test[sample])) panels = [ ("u_0, level 0", coarse, u0_test[sample], 10), ("u_T exact, level 0", coarse, uT_test[sample], 10), ("u_T predicted, level 0", coarse, trained(u0_test[sample]), 10), ("u_T exact, level 1", fine, uT_fine[sample], 4), ( "u_T predicted, level 1 (never seen)", fine, GraphOperator(trained.network, fine)(u0_fine[sample]), 4, ), ] figure, axes = plt.subplots(1, 5, figsize=(22, 4.8)) wall = jet_wall() for axis, (title, level, values, size) in zip(axes, panels): points = np.asarray(level.point_cloud) picture = axis.scatter( points[:, 0], points[:, 1], c=np.asarray(values)[:, 0], s=size, cmap="viridis", vmin=0.0, vmax=u_max, ) 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"advected-diffused Gaussian on JET: gno {MESSAGE_WIDE} + {UPDATE}, r={RADIUS_WIDE} " f"-- rel. L2 {on_coarse:.3f} on level 0, {on_fine:.3f} on level 1" ) figure.tight_layout() figure.savefig(Path(__file__).with_suffix(".png"), dpi=110) plt.show()