r"""GINO on the JET tokamak: an FNO run on a shape no grid fits. -Delta u = f in the JET poloidal cross-section, u = 0 on the wall The domain is the real JET wall read from its EQDSK file -- non-convex, with a divertor notch -- and ``f`` is a random mixture of three Gaussians. The data is a batched Q1 finite-element solve; the operator to learn is ``f |-> u``. ⚠ **Why this domain is the point.** An FNO needs a periodic Cartesian grid, and JET is neither. GINO's answer is not to deform the FNO but to put a kernel integral on each side of it:: cloud --GNO--> latent grid --FNO--> latent grid --GNO--> cloud The latent grid here is a plain 32x32 box over the bounding rectangle. It does not fit the domain, and it does not have to: the signed distance function, fed to the FNO as a channel, is what tells it where JET is. **The encoder is an integral over the physical domain**, ``v_0(x) = sum_i kappa(x, y_i) f(y_i) mu_i``, so the point cloud is a QUADRATURE of JET: points AND weights. The finite-element solution can be evaluated anywhere, so any quadrature will do, and the two cases use the two honest ones: * **case 1 -- a fixed cloud** (:class:`~....pointcloud_based.gino.GINO`): the mesh's own quadrature -- one point per cell, weighted by the cell's measure (the midpoint rule, from :class:`~...data_for_no.mesh_data.MeshData`). The geometry is known at construction and its stencils are built once. The call is then ``(u, mu)``, the discrete family's own contract, and it goes into ``NOProjector`` unchanged. * **case 2 -- a cloud per sample** (:class:`~....pointcloud_based.gino.VaryingCloudGINO`): a fresh uniform draw of points for every sample, weighted ``|JET| / n`` -- Monte-Carlo quadrature. The operator never sees the same discretisation twice. More expensive -- the stencils are rebuilt and differentiated through at every step -- and it buys something case 1 cannot have, which the last section measures: the trained operator answers on a cloud of a size it was never trained on, with that cloud's own weights. ⚠ Bare positions with no weights would be summed with ``1 / n``: a quadrature of nothing in particular, off by the domain's area, which the kernel then has to absorb. The weights are the whole reason the encoder is an operator and not a lookup. ⚠ **Coordinates are channels**, which is how case 2 reaches ``NOProjector`` at all. The projector hands the operator one array per sample; the cloud travels inside it, as the last ``dim`` channels, and :class:`_CloudInChannels` splits them again. That is the rule the discrete family already states, applied to a support rather than to a field. Run: python gino_jet_laplacian.py """ import re import time from pathlib import Path import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from matplotlib.path import Path as Polygon from scimba_jax.domains.meshless_domains.domains_nd import HypercubeND from scimba_jax.domains.tokamak import read_eqdsk_wall from scimba_jax.linear_approximation.basis.analytic_bases import ( local_lagrange_basis, local_lagrange_basis_by_logical, ) from scimba_jax.linear_approximation.basis.dof_map import UnstructuredLagrangeDofMap 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.unstructured_mesh import UnstructuredMesh from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.linear_approximation.variables.variables_fe import VariablesFE from scimba_jax.mapping.macro_mesh import macro_mesh_from_points from scimba_jax.neural_operator.data_for_no.grid_data import GridData from scimba_jax.neural_operator.data_for_no.mesh_data import MeshData from scimba_jax.neural_operator.data_for_no.point_cloud_data import PointCloudData from scimba_jax.neural_operator.discrete_no.abstract_neural_operator import ( AbstractNeuralOperator, ) from scimba_jax.neural_operator.discrete_no.pointcloud_based.gino import ( GINO, VaryingCloudGINO, ) from scimba_jax.nonlinear_approximation.numerical_solvers.no_projectors import ( NOProjector, ) from scimba_jax.physical_models.abstract_physical_weak_model import ( AbstractPhysicalWeakModel, ) from scimba_jax.physical_models.classical_weakform.laplacian_weak_form import ( LaplacianWeakForm, ) from scimba_jax.utils.scimba_pytree import ScimbaPytree _HERE = Path(__file__).resolve().parent _DATA = _HERE.parents[3] / "src/scimba_jax/domains/tokamak/data" _TOKAMAK_HELPERS = _HERE.parents[1] / "mesh/meshes_gmesh/example_macro_mesh_tokamak.py" MESH_SIZE = 0.13 N_GAUSSIANS = 3 N_TRAIN, N_TEST = 256, 64 N_CLOUD = 400 GRID_SIDE = 32 RADIUS = 0.25 N_EPOCHS, BATCH_SIZE, LEARNING_RATE = 1000, 32, 2e-3 ARCHITECTURE = dict(hidden_channels=32, gno_channels=8, n_modes=12, n_blocks=4) # ══════════════════════════════════════════════════════════════════════════════ # 1. The domain and the finite-element space # ══════════════════════════════════════════════════════════════════════════════ def cached_jet_mesh(order=1): """The JET wall, meshed once and cached. ⚠ Order 1 and not the 3 of ``solve_laplacian_unstructured_2d.py``: the DOF map is isoparametric, so a Q1 unknown needs a P1 geometry. GMSH does not repeat itself either, hence the cache -- see that example for the measurement. Args: order: the geometric degree. Returns: ``(nodes, cells)``. """ cache = _HERE / f"mesh_jet_p{order}_h{MESH_SIZE}.npz" if cache.exists(): stored = np.load(cache) return stored["nodes"], stored["cells"] source = _TOKAMAK_HELPERS.read_text() namespace = {"np": np} for name in ("_dedupe_polygon", "_densify_polygon"): match = re.search(rf"def {name}\(.*?(?=\ndef |\n# %%)", source, re.S) exec(match.group(0), namespace) # noqa: S102 radial, vertical = read_eqdsk_wall(str(_DATA / "eqdsk_jet_compare.dat")) wall = namespace["_densify_polygon"]( namespace["_dedupe_polygon"](np.stack([radial[:-1], vertical[:-1]], axis=1)), max_seg=0.03, ) macro = macro_mesh_from_points(wall, order=order, mesh_size=MESH_SIZE, smooth=True) np.savez_compressed(cache, nodes=macro.nodes, cells=macro.cells) return macro.nodes, macro.cells class GaussianMixture(ScimbaPytree): """``f(x) = sum_i A_i exp(-|x - c_i|^2 / 2 sigma_i^2)``. ⚠ A ``ScimbaPytree`` and not a closure, and that is what makes the batch work: the parameters are ``jnp`` LEAVES, so one draw per batch element is exactly what ``jax.vmap`` needs. A plain callable has no leaf, lands in the ``aux_data``, and a stacked model would silently keep the FIRST draw. 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))) nodes, cells = cached_jet_mesh() mesh = UnstructuredMesh( nodes=nodes, cells=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), ) variables = VariablesFE(basis=basis, nb_variables=1, dof_map=UnstructuredLagrangeDofMap) def build_scheme(source): """One Poisson-Dirichlet scheme for one source. Args: source: the right-hand side. Returns: The scheme. """ model = AbstractPhysicalWeakModel.from_weak_form( LaplacianWeakForm(dim=2, f=source), dirichlet=lambda x: jnp.zeros(1) ) return EllipticFEscheme(model, variables) # ⚠ The operator does not depend on the source, so the matrix is factorised ONCE # outside the vmap and only the right-hand side and the back substitution are # batched -- the pattern of # `solve_2d_helmholtz_robin_unstructured_disk_batched_sources.py`. _reference = build_scheme( GaussianMixture( jnp.zeros((N_GAUSSIANS, 2)), jnp.zeros(N_GAUSSIANS), jnp.ones(N_GAUSSIANS) ) ) _dofs_init = _reference._initial_dofs() _factorisation = EllipticFEscheme.factorise(_reference, _dofs_init) _back_solve = EllipticFEscheme._make_back_solve_fn() @jax.jit def solve_batch(centers, amplitudes, sigmas): """Every draw solved in ONE compiled program. Args: centers: ``(batch, n_gaussians, 2)``, amplitudes: ``(batch, n_gaussians)``, sigmas: ``(batch, n_gaussians)``. Returns: the DOFs, ``(batch, n_dofs, 1)``. """ def one(c, a, s): scheme = build_scheme(GaussianMixture(c, a, s)) return _back_solve(scheme, _dofs_init, _factorisation.lu, _factorisation.pivots) return jax.vmap(one)(centers, amplitudes, sigmas) @jax.jit def evaluate_batch(dofs, points): """The finite-element solutions read at arbitrary points. ⚠ This is what frees the cloud from the mesh: nothing forces the query points to be nodes, so the same data serves a fixed cloud and a fresh cloud per sample. Args: dofs: ``(batch, n_dofs, 1)``, points: ``(batch, n_points, 2)`` -- one cloud per sample. Returns: ``(batch, n_points, 1)``. """ read = jax.vmap( lambda d, p: VariablesFE._classical_local_evaluate_pure(variables, d, p), in_axes=(None, 0), ) return jax.vmap(read)(dofs, points) # ══════════════════════════════════════════════════════════════════════════════ # 2. Sampling inside the wall, and the signed distance to it # ══════════════════════════════════════════════════════════════════════════════ _radial, _vertical = read_eqdsk_wall(str(_DATA / "eqdsk_jet_compare.dat")) WALL = np.stack([_radial[:-1], _vertical[:-1]], axis=1) _polygon = Polygon(WALL) LOW, HIGH = WALL.min(axis=0), WALL.max(axis=0) def sample_inside(key, n_points, margin=-0.05): """Rejection sampling in the wall. Args: key: a PRNG state, n_points: how many points, margin: how far inside the wall to stay; negative shrinks the domain. Returns: ``(key, points)``. """ kept = np.empty((0, 2)) while len(kept) < n_points: key, subkey = jax.random.split(key) draw = np.asarray(jax.random.uniform(subkey, (4 * n_points, 2))) draw = draw * (HIGH - LOW) + LOW kept = np.concatenate( [kept, draw[_polygon.contains_points(draw, radius=margin)]] ) return key, jnp.asarray(kept[:n_points]) def signed_distance(points): """Distance to the wall, negative inside. Args: points: ``(n, 2)``. Returns: ``(n,)``. """ distance = np.full(len(points), np.inf) for start, end in zip(WALL[:-1], WALL[1:]): edge = end - start along = np.clip(((points - start) @ edge) / (edge @ edge), 0.0, 1.0) foot = start + along[:, None] * edge distance = np.minimum(distance, np.linalg.norm(points - foot, axis=1)) return distance * np.where(_polygon.contains_points(points), -1.0, 1.0) # ══════════════════════════════════════════════════════════════════════════════ # 3. The dataset # ══════════════════════════════════════════════════════════════════════════════ print("JET, -Delta u = f with u = 0 on the wall") print(f" Q1 space: {variables.n_nodes_total} DOFs on {mesh.n_cells_total} cells") key = jax.random.key(0) n_samples = N_TRAIN + N_TEST key, centers = sample_inside(key, n_samples * N_GAUSSIANS, margin=-0.25) centers = centers.reshape(n_samples, N_GAUSSIANS, 2) key, subkey = jax.random.split(key) amplitudes = jax.random.uniform( subkey, (n_samples, N_GAUSSIANS), minval=0.5, maxval=2.0 ) key, subkey = jax.random.split(key) sigmas = jax.random.uniform(subkey, (n_samples, N_GAUSSIANS), minval=0.20, maxval=0.40) start = time.perf_counter() dofs = jax.block_until_ready(solve_batch(centers, amplitudes, sigmas)) print(f" {n_samples} finite-element solves in {time.perf_counter() - start:.2f} s") # The two quadratures: the mesh's own, shared by everybody; a Monte-Carlo # draw per sample. mesh_data = MeshData(2, mesh) fixed_cloud, fixed_weights = mesh_data.point_cloud, mesh_data.integration_weights() N_FIXED = int(fixed_cloud.shape[0]) AREA = float(jnp.sum(fixed_weights)) print(f" mesh quadrature: {N_FIXED} cell centres, |JET| = {AREA:.4f}") key, subkey = jax.random.split(key) varying_clouds = [] for index in range(n_samples): subkey, cloud_points = sample_inside(subkey, N_CLOUD) varying_clouds.append(cloud_points) varying_clouds = jnp.stack(varying_clouds) sources = jax.vmap(lambda c, a, s, p: jax.vmap(GaussianMixture(c, a, s))(p)) f_fixed = sources( centers, amplitudes, sigmas, jnp.broadcast_to(fixed_cloud, (n_samples, N_FIXED, 2)) )[..., None] f_varying = sources(centers, amplitudes, sigmas, varying_clouds)[..., None] u_fixed = evaluate_batch(dofs, jnp.broadcast_to(fixed_cloud, (n_samples, N_FIXED, 2))) u_varying = evaluate_batch(dofs, varying_clouds) # ⚠ Normalised, and not for tidiness: an input whose scale is far from one is # the single factor that dominated the physics-informed FNO here (measured x11). F_SCALE, U_SCALE = float(jnp.std(f_fixed)), float(jnp.std(u_fixed)) print(f" scales: f {F_SCALE:.3f}, u {U_SCALE:.4f}") grid = GridData( 2, HypercubeND( [(float(LOW[0]), float(HIGH[0])), (float(LOW[1]), float(HIGH[1]))], is_main_domain=True, ), (GRID_SIDE, GRID_SIDE), ) sdf = signed_distance(np.asarray(grid.grid).reshape(-1, 2)) SDF = jnp.asarray(sdf.reshape(GRID_SIDE, GRID_SIDE, 1) / sdf.std()) @jax.jit def predict(operator, inputs): """The operator on a batch -- the operator as an argument of the jit, so its stencils and cloud are operands, not constants folded into the HLO. Args: operator: a neural operator with the ``(u, mu)`` call, inputs: ``(batch, n_points, channels)``. Returns: ``(batch, n_points, out_channels)``. """ return jax.vmap(operator)(inputs) def relative_error(predicted, truth): """Per-sample relative L2 error. Args: predicted: ``(batch, n_points, 1)``, truth: same shape. Returns: ``(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 ) # ══════════════════════════════════════════════════════════════════════════════ # 4. Case 1: a fixed cloud # ══════════════════════════════════════════════════════════════════════════════ cloud = PointCloudData.from_points( fixed_cloud, [0.3], integration_weights=fixed_weights ) fixed_operator = GINO( cloud, grid, RADIUS, cloud_channels=1, out_channels=1, grid_input=SDF, key=jax.random.key(1), kernel_kwargs=dict(hidden_sizes=[32, 32]), **ARCHITECTURE, ) print( f"\ncase 1 -- fixed cloud: {fixed_operator.encoder.n_neighbors} stencil slots, " f"{fixed_operator.ndof()} parameters" ) projector = NOProjector( fixed_operator, (f_fixed[:N_TRAIN] / F_SCALE, u_fixed[:N_TRAIN] / U_SCALE), learning_rate=LEARNING_RATE, ) start = time.perf_counter() _, projector = projector.project( jax.random.key(2), fixed_operator, N_EPOCHS, batch_size=BATCH_SIZE, tqdm_desc=" fixed cloud", ) fixed_trained = projector.operator fixed_seconds = time.perf_counter() - start fixed_prediction = predict(fixed_trained, f_fixed[N_TRAIN:] / F_SCALE) * U_SCALE fixed_error = relative_error(fixed_prediction, u_fixed[N_TRAIN:]) baseline = relative_error( jnp.broadcast_to(jnp.mean(u_fixed[:N_TRAIN], axis=0), u_fixed[N_TRAIN:].shape), u_fixed[N_TRAIN:], ) print(f" trained in {fixed_seconds:.0f} s") print(f" relative L2 on test : {float(jnp.mean(fixed_error)):.3e}") print(f" mean-predictor : {float(jnp.mean(baseline)):.3e}") # ══════════════════════════════════════════════════════════════════════════════ # 5. Case 2: a fresh cloud for every sample # ══════════════════════════════════════════════════════════════════════════════ class _CloudInChannels(AbstractNeuralOperator): """Carry the query points inside the input channels. ⚠ Not a trick: "coordinates are channels" is the discrete family's own rule, and this applies it to the SUPPORT. ``NOProjector`` hands an operator one array per sample, so a cloud that changes per sample has nowhere else to travel. The array is ``[fields | x | y]``; this splits it and calls the operator with a real cloud -- and with the Monte-Carlo weights ``|JET| / n`` a uniform draw of ``n`` points deserves. The wrapper holds the operator rather than copying it, so training through it trains the operator it holds. Args: operator: the varying-cloud GINO, n_fields: how many leading channels are fields rather than coordinates. """ def __init__(self, operator, n_fields=1): self.operator = operator self.n_fields = int(n_fields) def __call__(self, u, mu=None): """Split ``[fields | coordinates]`` and run the operator. Args: u: ``(n_points, n_fields + dim)``, mu: the PDE's parameters, passed through. Returns: ``(n_points, out_channels)``. """ n_points = u.shape[0] return self.operator( u[:, : self.n_fields], mu, cloud=u[:, self.n_fields :], grid_input=SDF, cloud_weights=jnp.full((n_points,), AREA / n_points), ) varying_operator = VaryingCloudGINO( grid, RADIUS, cloud_channels=1, out_channels=1, grid_input=SDF, key=jax.random.key(3), kernel_kwargs=dict(hidden_sizes=[32, 32]), **ARCHITECTURE, ) wrapped = _CloudInChannels(varying_operator) print(f"\ncase 2 -- a cloud per sample: {wrapped.ndof()} parameters (the same network)") with_coordinates = jnp.concatenate([f_varying / F_SCALE, varying_clouds], axis=-1) projector = NOProjector( wrapped, (with_coordinates[:N_TRAIN], u_varying[:N_TRAIN] / U_SCALE), learning_rate=LEARNING_RATE, ) start = time.perf_counter() _, projector = projector.project( jax.random.key(4), wrapped, N_EPOCHS, batch_size=BATCH_SIZE, tqdm_desc=" varying cloud", ) varying_trained = projector.operator varying_seconds = time.perf_counter() - start varying_prediction = predict(varying_trained, with_coordinates[N_TRAIN:]) * U_SCALE varying_error = relative_error(varying_prediction, u_varying[N_TRAIN:]) print( f" trained in {varying_seconds:.0f} s ({varying_seconds / fixed_seconds:.1f}x case 1)" ) print(f" relative L2 on test : {float(jnp.mean(varying_error)):.3e}") # ── What case 2 buys: a cloud size it never saw ────────────────────────────── # ⚠ The real neural-operator claim, and the reason the expensive case exists. # Nothing here was retrained: the same weights answer on a different number of # query points, because the kernel is a function of position, not of an index # -- and the quadrature weights follow the cloud, ``|JET| / 3n``. key, other_cloud = sample_inside(key, 3 * N_CLOUD) other_points = jnp.broadcast_to(other_cloud, (N_TEST, 3 * N_CLOUD, 2)) other_source = sources( centers[N_TRAIN:], amplitudes[N_TRAIN:], sigmas[N_TRAIN:], other_points )[..., None] other_truth = evaluate_batch(dofs[N_TRAIN:], other_points) other_prediction = ( jax.jit( lambda op, u, x, w: jax.vmap( lambda u_, x_: op(u_, None, cloud=x_, cloud_weights=w) )(u, x) )( varying_trained.operator, other_source / F_SCALE, other_points, jnp.full((3 * N_CLOUD,), AREA / (3 * N_CLOUD)), ) * U_SCALE ) print( f" zero-shot on {3 * N_CLOUD} points (trained on {N_CLOUD}): " f"{float(jnp.mean(relative_error(other_prediction, other_truth))):.3e}" ) # ══════════════════════════════════════════════════════════════════════════════ # 6. Figure # ══════════════════════════════════════════════════════════════════════════════ # ⚠ The top row is the point of the method, not decoration: the latent grid is # a box that does NOT fit JET, and the signed distance is the only thing that # tells the FNO where the domain is. The bottom row is the two cases on the same # test sample, so their errors are directly comparable. SAMPLE = 0 figure, axes = plt.subplots(2, 4, figsize=(17, 9), constrained_layout=True) fixed_points = np.asarray(fixed_cloud) grid_points = np.asarray(grid.grid).reshape(-1, 2) def draw_wall(axis): """Outline JET and square the axes. Args: axis: the matplotlib axes. """ axis.plot(WALL[:, 0], WALL[:, 1], color="k", linewidth=1.0) axis.set_aspect("equal") axis.set_xticks([]) axis.set_yticks([]) def on_cloud(axis, points, values, title, cmap="viridis", limit=None): """Scatter a field carried by a point cloud. Args: axis: the matplotlib axes, points: ``(n, 2)``, values: ``(n,)``, title: the panel title, cmap: the colour map, limit: symmetric colour limit, or None to span the data. Returns: The scatter artist. """ draw_wall(axis) picture = axis.scatter( points[:, 0], points[:, 1], c=values, s=13, cmap=cmap, vmin=None if limit is None else -limit, vmax=None if limit is None else limit, ) figure.colorbar(picture, ax=axis) axis.set_title(title, fontsize=9) return picture # ── (a) the supports: a box grid over a shape it does not fit ──────────────── draw_wall(axes[0, 0]) axes[0, 0].plot( grid_points[:, 0], grid_points[:, 1], ".", color="0.75", markersize=2.5, label=f"latent grid {GRID_SIDE}x{GRID_SIDE}", ) axes[0, 0].plot( fixed_points[:, 0], fixed_points[:, 1], ".", color="tab:blue", markersize=3.5, label=f"mesh quadrature ({N_FIXED} cells)", ) circle = plt.Circle( tuple(fixed_points[0]), RADIUS, fill=False, color="tab:red", linewidth=1.4 ) axes[0, 0].add_patch(circle) axes[0, 0].legend(fontsize=7, loc="upper right") axes[0, 0].set_title( f"supports: the grid is a BOX\nred = kernel ball, r = {RADIUS}", fontsize=9 ) # ── (b) the signed distance, the FNO's only clue about the domain ──────────── draw_wall(axes[0, 1]) picture = axes[0, 1].pcolormesh( np.asarray(grid.grid[..., 0]), np.asarray(grid.grid[..., 1]), np.asarray(SDF[..., 0]), shading="auto", cmap="RdBu_r", ) figure.colorbar(picture, ax=axes[0, 1]) axes[0, 1].set_title( "signed distance on the latent grid\n(negative inside)", fontsize=9 ) # ── (c, d) one test sample: the source, and the reference solution ─────────── on_cloud( axes[0, 2], fixed_points, np.asarray(f_fixed[N_TRAIN + SAMPLE, :, 0]), f"source f (sample {SAMPLE})", cmap="magma", ) truth = np.asarray(u_fixed[N_TRAIN + SAMPLE, :, 0]) on_cloud(axes[0, 3], fixed_points, truth, "reference u (Q1 finite elements)") # ── (e, f) case 1 ──────────────────────────────────────────────────────────── guess = np.asarray(fixed_prediction[SAMPLE, :, 0]) on_cloud(axes[1, 0], fixed_points, guess, "case 1: GINO, fixed cloud") scale = np.abs(truth).max() on_cloud( axes[1, 1], fixed_points, guess - truth, f"error rel. L2 = {float(fixed_error[SAMPLE]):.1e}", cmap="coolwarm", limit=0.15 * scale, ) # ── (g, h) case 2, on its own cloud ────────────────────────────────────────── varying_points = np.asarray(varying_clouds[N_TRAIN + SAMPLE]) varying_truth = np.asarray(u_varying[N_TRAIN + SAMPLE, :, 0]) varying_guess = np.asarray(varying_prediction[SAMPLE, :, 0]) on_cloud(axes[1, 2], varying_points, varying_guess, "case 2: GINO, a cloud per sample") on_cloud( axes[1, 3], varying_points, varying_guess - varying_truth, f"error rel. L2 = {float(varying_error[SAMPLE]):.1e}", cmap="coolwarm", limit=0.15 * scale, ) figure.suptitle( "GINO on JET: -Delta u = f, u = 0 on the wall " f"test error {float(jnp.mean(fixed_error)):.2e} (fixed cloud) / " f"{float(jnp.mean(varying_error)):.2e} (cloud per sample)", fontsize=11, ) output = __file__.replace(".py", ".png") figure.savefig(output, dpi=110) print(f"\nfigure: {output}") # ── A second figure: the error distributions and the zero-shot check ───────── second, panels = plt.subplots(1, 2, figsize=(12, 4.2), constrained_layout=True) panels[0].hist(np.asarray(fixed_error), bins=20, alpha=0.65, label="fixed cloud") panels[0].hist(np.asarray(varying_error), bins=20, alpha=0.65, label="cloud per sample") panels[0].axvline( float(jnp.mean(baseline)), color="k", linestyle="--", label=f"mean predictor ({float(jnp.mean(baseline)):.2f})", ) panels[0].set_xlabel("relative L2 error") panels[0].set_ylabel("test samples") panels[0].set_title("both cases against the trivial baseline", fontsize=10) panels[0].legend(fontsize=8) zero_shot = relative_error(other_prediction, other_truth) panels[1].plot(np.asarray(varying_error), np.asarray(zero_shot), ".", markersize=6) diagonal = [0.0, float(max(varying_error.max(), zero_shot.max())) * 1.05] panels[1].plot(diagonal, diagonal, color="k", linewidth=0.8) panels[1].set_xlabel(f"error on {N_CLOUD} points (as trained)") panels[1].set_ylabel(f"error on {3 * N_CLOUD} points (never trained)") panels[1].set_title( "case 2 zero-shot: same weights, a cloud size it never saw", fontsize=10 ) panels[1].set_aspect("equal") second_output = __file__.replace(".py", "_errors.png") second.savefig(second_output, dpi=110) print(f"figure: {second_output}") plt.show()