"""The basis is told nothing but the FIELD, sampled on a grid. Three operators, one family, and a chain of two questions --------------------------------------------------------- ``advection_diffusion_1d_learned_basis.py`` learns ONE basis for the whole family. ``..._basis_conditioned_by_mu.py`` hands the basis the parameters. This file hands it only the transport field ``b``, evaluated at fixed points, and asks the network to recover from those samples whatever the basis needs:: shared phi_k(x) = m(N(x)) T_k(x) by mu phi_k(x; mu) = m(N(x, mu)) T_k(x) by the field phi_k(x; z) = m(N(x, z)) T_k(x), z = MLP(first modes of b on a grid) The three are trained here TOGETHER, on the same family, from the same keys, so that the two questions can be read off one table: * does conditioning help at all? ``by mu`` against ``shared``; * can it be done without the parameters? ``by the field`` against ``by mu``. The second is the honest target: ``by mu`` is an ORACLE -- it is given exact information for free -- so the field version succeeding means it has RECOVERED that information, and failing means the encoder is what needs work, not the idea. Why this family, and how it was chosen --------------------------------------- ⚠ The tight family of the two earlier files (``b = mu_0 + mu_1 x``, ``mu`` in ``[0.8, 1.2]^2``) cannot answer either question, for two independent reasons: * **``b`` stays positive and barely varies**, so every member is wrong in the same place, in the same way. Measured below with the ``1 - corr`` diagnostic: **0.0036**. One basis enriched there serves everybody, the shared basis is never put in difficulty, and conditioning has nothing to win; * **``b`` is affine**, so its spectrum is ``(mu_0, mu_1)`` up to a bijection. So the transport becomes a smooth SUM OF SINES, which fixes both:: b(x; mu) = B0 + sum_k mu_k sin(k pi x), k = 1..4, mu_k in [-0.5, 0.5] Measured, against the tight family's 0.0036 (``1 - corr`` / relative L2 error of the classical scheme):: family 4c p1 10c p1 max|u_ref| tight, b affine 0.0036/0.48 -- -- B0=2.5, K=4, a<0.5, eps=0.05 0.0549/0.82 0.0109/0.31 4.3e-01 B0=2.5, K=6, a<0.35, eps=0.02 -- 0.0151/0.31 4.4e-01 B0=2, K=4, a<1.5, eps=0.05 0.1241/1.57 -- 1.2e+01 B0=0, K=4, a<3, changes sign 0.0953/3.36 -- 5.2e+07 ⚠ **``B0`` must exceed ``sum_k |mu_k|``.** The last row is the warning: as soon as ``b`` can vanish and change sign, the continuous solution grows like ``exp(int b / eps)`` and ``u_ref`` reaches 1e7 -- the family stops testing anything. With ``B0 = 2.5`` and six amplitudes below 0.35, ``b`` stays above 0.4. ⚠ **The discrimination falls under refinement**, and that is not a defect of the family: 0.055 on 4 cells against 0.011 on 8, on the same problem. The richer the space, the less the best basis depends on the PDE -- the thesis of the shared-basis file, seen from the other side. So the sweep sits at 10 cells, where the classical scheme is at 31% (p1) and 12% (p2): poor enough for a basis to matter, far from the 87% of a 4-cell space where nothing solves at all. ⚠ **The spectrum of a sum of sines IS its coefficients.** That is deliberate: it means the field encoder CAN reach its oracle exactly, so a gap between ``by field`` and ``by mu`` is the encoder's doing and nothing else. It also means this file does not yet prove anything about non-parametric fields -- that is what a tabulated field would be for. ⚠ **The family therefore differs from the two earlier files, and their numbers do not transfer.** That is why all three operators are retrained here. ⚠ The diagnostic runs FIRST and costs nothing. Run: python advection_diffusion_1d_basis_conditioned_by_field.py """ from __future__ import annotations import functools 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_1d import Segment1D from scimba_jax.linear_approximation.basis.analytic_bases import local_taylor_basis from scimba_jax.linear_approximation.basis.general_bases import ( AnalyticBasis, PatchwiseParametricBasis, ) from scimba_jax.linear_approximation.galerkin.dg.flux import SIPGFlux 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_dg import VariablesDG from scimba_jax.mapping.mapping import InvertibleFunction, Mapping from scimba_jax.neural_operator.physic_no.solver_based import ( ConditionedDGOperator, DGSolverOperator, FieldConditionedDGOperator, FourierLatentNet, ) 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.physical_models.abstract_physical_model import AbstractPhysicalModel from scimba_jax.physical_models.abstract_physical_weak_model import ( AbstractPhysicalWeakModel, ) 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.utils.scimba_pytree import ScimbaPytree # ── The family: a turning point whose position varies ──────────────────────── DIM = 1 EPS = 0.02 #: ``b = B0 + sum_k mu_k sin(k pi x)``. ⚠ ``B0`` must exceed ``sum_k |mu_k|``, #: or ``b`` changes sign and the continuous solution grows exponentially -- #: measured at ``max|u_ref| = 5e7``, which tests nothing. B0 = 2.5 N_MU = 6 MU_AMPLITUDE = 0.35 #: What the field encoder reads, and how much of it survives. N_POINTS = 64 N_MODES = 20 LATENT = 4 #: ⚠ Ten cells, not four. Measured on this family, the classical scheme is at #: 31% (p1) and 12% (p2) error there -- poor enough for a basis to be worth #: learning, far enough from the 87% of a 4-cell space to be a real solve. CASES = ((10, 1), (10, 2)) N_CELLS_REF, ORDER_REF = 128, 4 B_TRAIN, B_TEST = 32, 32 BATCH_SIZE = 16 GRID_SIZE = 101 N_EPOCHS = 400 OPTIMIZER = "SS-BFGS" SEED = 0 _MAPPING = Mapping(mappings=[InvertibleFunction(lambda x: x, lambda y: y)]) def source(x: jnp.ndarray) -> jnp.ndarray: """``f = 1``. Args: x: One physical point. Returns: The source there. """ return jnp.ones(()) class SinTransport(ScimbaPytree): """``b(x; mu) = B0 + sum_k mu_k sin(k pi x)``, carried by a LEAF. ⚠ ``mu`` is a ``jnp`` array, hence a pytree child without any marker: the spelling is what places it, and it is what makes the family a stackable batch. Nothing declares it trainable, so it is frozen -- it is data. ⚠ ``B0`` is a module constant, not a parameter: it is the margin that keeps ``b`` positive. Letting it vary down to where ``b`` vanishes turns a test family into an exponentially growing one (see the module docstring). Args: mu: ``(N_MU,)``, the sine amplitudes. """ def __init__(self, mu): self.mu = jnp.asarray(mu, dtype=float) def __call__(self, x: jnp.ndarray) -> jnp.ndarray: """``b(x)``, of shape ``(1,)``. Args: x: One physical point. Returns: The transport there. """ modes = jnp.arange(1, N_MU + 1, dtype=float) return jnp.array([B0 + jnp.sum(self.mu * jnp.sin(modes * jnp.pi * x[0]))]) class AdvectionDiffusion1D(AbstractPhysicalModel): """What the projector samples: ``mu``, and the reference data. Args: main_domain: The segment. mu: ``(N_MU,)``, this PDE's transport. data: ``(x, u_ref(x))``, the reference values. """ def __init__(self, main_domain, mu, data): super().__init__(main_domain=main_domain) self.mu = jnp.asarray(mu, dtype=float) self.data_residuals["data"] = CollocDataResidual( size=1, model_type="x", data=data, # ⚠ False: the grid is the SAME for every PDE, so it stays in # aux_data rather than being duplicated B times. Only the VALUES # are batched. batchable_args=False, ) def mu_of(pde: AdvectionDiffusion1D) -> jnp.ndarray: """The ORACLE's code: the parameters, handed over. Args: pde: The sampled model. Returns: ``mu``. """ return pde.mu def field_of(pde: AdvectionDiffusion1D) -> SinTransport: """The transport field of one PDE, as a callable of ``x``. ⚠ What the field-conditioned operator is given, and ALL it is given. It then samples it; the parameters are never read downstream. Args: pde: The sampled model. Returns: ``b``. """ return SinTransport(pde.mu) def weak_model_of(pde: AdvectionDiffusion1D) -> AbstractPhysicalWeakModel: """The bridge: the projector's model becomes a weak form. Args: pde: The sampled model. Returns: The variational model. """ form = EllipticWeakForm( dim=DIM, A=lambda x: EPS * jnp.eye(DIM), b=SinTransport(pde.mu), c=lambda x: jnp.zeros(()), f=source, ) return AbstractPhysicalWeakModel.from_weak_form( form, dirichlet=lambda x: jnp.zeros(1) ) # ── The spaces and the networks ────────────────────────────────────────────── def make_mesh(n_cells: int, order: int) -> Mesh: """The 1D Cartesian mesh, quadrature exact for the wanted order. Args: n_cells: Number of cells. order: Local degree. Returns: The mesh. """ return Mesh( dim=DIM, n_cells=[n_cells], ref_quad=UnitSquareTensorized(dim=DIM, order=2 * order + 2), mapping=_MAPPING, ) def multiplier(u: jnp.ndarray) -> jnp.ndarray: """``exp(2 tanh(N))``: positive, worth 1 at ``N = 0``, and wide. Args: u: The network's output. Returns: The multiplier. """ return jnp.exp(2.0 * jnp.tanh(u[0])) @functools.lru_cache(maxsize=None) def _taylor_of(order: int): """The local Taylor basis of one order, built ONCE. ⚠ Memoised: a callable lives in ``aux_data``, so two equivalent lambdas built separately give two treedefs and two compilations. Args: order: Local degree. Returns: ``(y, i, mesh) -> values``. """ def local(y, i, mesh): return local_taylor_basis(y, i, mesh, order=order, out_dim=1) return local @functools.lru_cache(maxsize=None) def _enriched_of(order: int): """``m(x) T_k(x)`` for one order, built ONCE (see above). Args: order: Local degree. Returns: ``(u, y, i, mesh) -> values``. """ def enriched(u, y, i, mesh): return local_taylor_basis(y, i, mesh, order=order, out_dim=1) * multiplier(u) return enriched class BasisNN(MLP): """The multiplier before bounding. ONE network for the whole domain. Args: key: Random generator state. in_size: Input width -- ``DIM``, or ``DIM + latent`` when conditioned. """ def __init__(self, key, in_size: int = DIM): super().__init__(in_size=in_size, out_size=1, hidden_sizes=[16, 16], key=key) class ConditionedBasisNN(ScimbaPytree): """``N(x, z)``: the network sees the point AND this PDE's code. ⚠ ``z`` is a plain ``jnp`` field: it lands in ``children`` (so a batch can carry one per PDE) and, being undeclared, it is frozen -- the pair a code that is DATA wants. Its shape at construction is the FINAL shape. Args: key: Random generator state. latent_size: Length of ``z``. """ def __init__(self, key, latent_size: int): self.net = BasisNN(key, in_size=DIM + latent_size) self.z = jnp.zeros((latent_size,)) def __call__(self, x: jnp.ndarray) -> jnp.ndarray: """The network at one point, conditioned by ``z``. Args: x: One point, ``(dim,)``. Returns: ``(1,)``. """ return self.net(jnp.concatenate([x, self.z])) def make_latent_net(key) -> FourierLatentNet: """``b`` on a grid -> its first modes -> an MLP -> ``z``. ⚠ 20 complex modes of a 64-point sample, so 40 real features, compressed to ``LATENT`` numbers. Since ``b`` is a sum of ``N_MU`` sines, its spectrum CONTAINS ``mu`` exactly -- deliberately, so that this operator can reach its oracle and any gap between the two is the encoder's doing. It also means the file proves nothing yet about fields that have no parameters. ⚠ The count matters as much as the architecture: the table at the end reports ``n_theta`` for the three operators, and a field encoder that dwarfs the basis network turns the comparison into one about capacity. Args: key: Random generator state. Returns: The latent network. """ return FourierLatentNet( MLP( in_size=FourierLatentNet.feature_size(N_MODES), out_size=LATENT, # ⚠ Same width as the basis network. Wider, the encoder alone # outweighs everything else -- measured 2 500 weights against the # basis network's 385 with [32, 32] -- and the comparison stops # being about conditioning and becomes about capacity. hidden_sizes=[16, 16], key=key, ), n_points=N_POINTS, n_modes=N_MODES, latent_size=LATENT, ) def classical_space(n_cells: int, order: int) -> VariablesDG: """The usual DG space: Taylor per cell, nothing learned. Args: n_cells: Number of cells. order: Local degree. Returns: The space. """ mesh = make_mesh(n_cells, order) basis = AnalyticBasis( nb_basis=order + 1, out_dim=1, mesh=mesh, basis_type="scalar", local_basis=_taylor_of(order), ) return VariablesDG(basis=basis, nb_variables=1) def learned_space(network, n_cells: int, order: int) -> VariablesDG: """The same space, whose local basis carries ``network``. ⚠ ``use_local_coords=True``: the network sees the coordinate normalised within the cell. Without it, it can learn to zero out whole regions of the domain, which makes the DG matrix singular during optimisation. Args: network: The parametric function -- plain or conditioned. n_cells: Number of cells. order: Local degree. Returns: The space. """ mesh = make_mesh(n_cells, order) basis = PatchwiseParametricBasis( nb_basis=order + 1, out_dim=1, mesh=mesh, patchwise_parametric_function=network, local_basis=_enriched_of(order), basis_type="scalar", use_local_coords=True, ) return VariablesDG(basis=basis, nb_variables=1) def make_flux(order: int) -> SIPGFlux: """The SIPG flux for this order. Args: order: Local degree. Returns: The flux. """ return SIPGFlux(sigma=10.0 * order * (order + 1), h=None) def field_grid() -> jnp.ndarray: """Where the transport field is sampled, ``(N_POINTS, 1)``. Returns: The grid. """ return jnp.linspace(0.0, 1.0, N_POINTS)[:, None] def build_operators(key, n_cells: int, order: int): """The three operators of one mesh, from ONE key. ⚠ One key for all three, so that what differs between them is the conditioning and not the draw. They cannot have identical weights -- the input widths differ -- but they start from the same seed. Args: key: Random generator state. n_cells: Number of cells. order: Local degree. Returns: ``{name: operator}``, in reading order. """ flux = make_flux(order) shared = DGSolverOperator( learned_space(BasisNN(key), n_cells, order), flux, weak_model_fn=weak_model_of, ) by_mu = ConditionedDGOperator( learned_space(ConditionedBasisNN(key, N_MU), n_cells, order), flux, mu_of, latent_size=N_MU, weak_model_fn=weak_model_of, ) by_field = FieldConditionedDGOperator( learned_space(ConditionedBasisNN(key, LATENT), n_cells, order), flux, field_of, field_grid(), make_latent_net(key), weak_model_fn=weak_model_of, ) return {"shared": shared, "by mu": by_mu, "by field": by_field} # ── Measurements ───────────────────────────────────────────────────────────── def relative_l2(prediction, reference) -> float: """Relative L2 error averaged over the batch. Args: prediction: ``(B, n, 1)``. reference: ``(B, n, 1)``. Returns: A float. """ num = jnp.linalg.norm(prediction - reference, axis=(1, 2)) den = jnp.linalg.norm(reference, axis=(1, 2)) return float(jnp.mean(jnp.where(den > 0, num / den, num))) def shape_spread(prediction, reference) -> tuple[float, np.ndarray]: """How differently the members of the family are WRONG. Each pointwise error profile is normalised to unit norm, so amplitude drops out and only the SHAPE remains; the number returned is one minus the mean correlation to the family's mean profile. Zero means every member is wrong in the same way. ⚠ It replaced a barycentre of the error, which was the wrong instrument: a position only sees the difficulty MOVE, and a boundary layer whose thickness varies by a factor of ten would block a shared basis just as well without moving an inch. Measured on the tight family, the barycentre said 0.0% spread where this says 0.0036 -- both small, but only one of them for the right reason. Args: prediction: ``(B, n, 1)``, the classical DG solutions. reference: ``(B, n, 1)``. Returns: ``(spread, profiles)``. """ error = np.abs(np.asarray(prediction - reference))[..., 0] profiles = error / np.maximum(np.linalg.norm(error, axis=1, keepdims=True), 1e-30) mean = profiles.mean(axis=0) mean = mean / max(float(np.linalg.norm(mean)), 1e-30) return 1.0 - float((profiles @ mean).mean()), profiles #: ⚠ A PROXY, and a weakly calibrated one -- it warns, it does not predict. #: The only value it is anchored on is the tight family of the earlier files, #: which measures 0.0036 and showed little gain; no threshold above that has #: ever been checked against an actual training. So it fires only at that order #: of magnitude, and a value above it is not a promise of anything. DISCRIMINATION_FLOOR = 0.005 def can_the_family_discriminate(classical, data, label: str) -> float: """Report how differently the family is solved, and warn if it is not. Args: classical: ``(B, n, 1)``, the classical DG solutions. data: What :func:`prepare_family` produced. label: The space this was measured on. Returns: The spread. """ spread, _ = shape_spread(classical, data["test_ref"]) verdict = ( "the family cannot discriminate" if spread < DISCRIMINATION_FLOOR else "conditioning has room" ) print(f" {label:14s} 1 - corr = {spread:.4f} {verdict}") return spread def prepare_family(key, domain, x_eval): """The ``mu``, the reference solutions, and the models carrying them. Args: key: Random generator state. domain: The segment. x_eval: ``(n, 1)``, the common grid. Returns: ``(key, data)``. """ key, sub = jax.random.split(key) mus = jax.random.uniform( sub, (B_TRAIN + B_TEST, N_MU), minval=-MU_AMPLITUDE, maxval=MU_AMPLITUDE, ) print(f"reference: DG {N_CELLS_REF} cells, order {ORDER_REF} ...") start = time.perf_counter() reference_op = DGSolverOperator( classical_space(N_CELLS_REF, ORDER_REF), make_flux(ORDER_REF), weak_model_fn=weak_model_of, ) blank = [ AdvectionDiffusion1D(domain, mu, (x_eval, jnp.zeros((x_eval.shape[0], 1)))) for mu in mus ] u_ref = jax.block_until_ready( reference_op.evaluate_batch(type(blank[0]).create_batch(blank), x_eval) ) print(f" {time.perf_counter() - start:.1f} s for {len(mus)} PDEs") pdes = [ AdvectionDiffusion1D(domain, mu, (x_eval, u_ref[i])) for i, mu in enumerate(mus) ] return key, { "mus": mus, "train_pdes": pdes[:B_TRAIN], "train_batch": type(pdes[0]).create_batch(pdes[:B_TRAIN]), "test_batch": type(pdes[0]).create_batch(pdes[B_TRAIN:]), "train_ref": u_ref[:B_TRAIN], "test_ref": u_ref[B_TRAIN:], } def train(key, operator, data, domain, x_eval): """Train one operator. ⚠ The key is passed IN rather than split inside, so that the three operators see the same sampled batches in the same order. Args: key: Random generator state. operator: The operator to train. data: What :func:`prepare_family` produced. domain: The segment. x_eval: The common grid. Returns: ``(trained operator, history, seconds, n_theta)``. """ space = PhysicNOApproximationSpace( dims={"x": DIM}, list_models=[operator], model_type="x" ) # ⚠ `bc=False`: the Dirichlet condition is imposed by the DG FLUX, not by a # loss term. sampler = TensorizedSampler([DomainSampler(domain)], bc=False, model_type="x") projector = PhysicNOProjector( data["train_pdes"], space, sampler, only_data=True, weights={"data": [1.0]}, optimizer=OPTIMIZER, ) start = time.perf_counter() _, projector = projector.project( key, space, N_EPOCHS, BATCH_SIZE, n_colloc=0, n_bc_colloc=0, n_dl_colloc=x_eval.shape[0], verbose=True, ) seconds = time.perf_counter() - start history = np.asarray(projector.losses.losses_history["total"]).reshape(-1) return projector.space.models[0], history, seconds, int(space.ndof) def evaluate(operator, data, x_eval) -> tuple[float, float]: """The relative errors of one operator, on train and on test. Args: operator: The operator. data: What :func:`prepare_family` produced. x_eval: The common grid. Returns: ``(train error, test error)``. """ return ( relative_l2( operator.evaluate_batch(data["train_batch"], x_eval), data["train_ref"] ), relative_l2( operator.evaluate_batch(data["test_batch"], x_eval), data["test_ref"] ), ) def run_case(key, n_cells: int, order: int, data, domain, x_eval): """One mesh: the classical scheme, then the three learned operators. Args: key: Random generator state. n_cells: Cells of the coarse grid. order: Local degree. data: What :func:`prepare_family` produced. domain: The segment. x_eval: The common grid. Returns: ``(key, results)``. """ coarse = DGSolverOperator( classical_space(n_cells, order), make_flux(order), weak_model_fn=weak_model_of ) result = { "n_cells": n_cells, "order": order, "n_dofs": n_cells * (order + 1), "label": f"{n_cells} cells, p{order}", "classical": evaluate(coarse, data, x_eval), "classical_curve": coarse.evaluate_batch(data["test_batch"], x_eval), "initial": {}, "trained": {}, "curves": {}, "history": {}, "n_theta": {}, "seconds": {}, } key, weights_key, train_key = jax.random.split(key, 3) for name, operator in build_operators(weights_key, n_cells, order).items(): result["initial"][name] = evaluate(operator, data, x_eval) operator, history, seconds, n_theta = train( train_key, operator, data, domain, x_eval ) result["trained"][name] = evaluate(operator, data, x_eval) result["curves"][name] = operator.evaluate_batch(data["test_batch"], x_eval) result["history"][name] = history result["n_theta"][name] = n_theta result["seconds"][name] = seconds return key, result def geometric_mean(values) -> float: """The GEOMETRIC mean of a list of ratios. Args: values: The ratios. Returns: Their geometric mean. """ return float(np.exp(np.mean(np.log(np.asarray(values))))) NAMES = ("shared", "by mu", "by field") def report(results) -> None: """The table, and the two questions it answers. Args: results: What :func:`run_case` returned for each mesh. """ print("\n### SOLVE error (test set)") print( " `classical` has no network. `init` is the learned basis BEFORE\n" " training -- at epoch 0 the network is random, so the starting basis\n" " is already an enriched one, and it is the honest baseline." ) header = f"\n{'space':13s} {'dof':>4s} {'classical':>10s}" for name in NAMES: header += f" {name + ' init':>12s} {name:>10s}" print(header) print("-" * len(header)) for r in results: line = f"{r['label']:13s} {r['n_dofs']:4d} {r['classical'][1]:10.3e}" for name in NAMES: line += f" {r['initial'][name][1]:12.3e} {r['trained'][name][1]:10.3e}" print(line) print("\n### The two questions") print( f"{'space':13s} {'by mu / shared':>16s} {'by field / shared':>19s} " f"{'by field / by mu':>18s}" ) print("-" * 70) for r in results: shared = r["trained"]["shared"][1] by_mu = r["trained"]["by mu"][1] by_field = r["trained"]["by field"][1] print( f"{r['label']:13s} {shared / by_mu:15.2f}x {shared / by_field:18.2f}x " f"{by_mu / by_field:17.2f}x" ) print("\n### Parameter count") print(f"{'space':13s} " + " ".join(f"{name:>10s}" for name in NAMES)) print("-" * 48) for r in results: print( f"{r['label']:13s} " + " ".join(f"{r['n_theta'][name]:10d}" for name in NAMES) ) oracle = geometric_mean( [r["trained"]["shared"][1] / r["trained"]["by mu"][1] for r in results] ) field = geometric_mean( [r["trained"]["shared"][1] / r["trained"]["by field"][1] for r in results] ) recovered = geometric_mean( [r["trained"]["by mu"][1] / r["trained"]["by field"][1] for r in results] ) print(f"\n conditioning helps (oracle vs shared): {oracle:.2f}x") print(f" the field version helps (vs shared): {field:.2f}x") print(f" the field version vs its ORACLE: {recovered:.2f}x") # ⚠ The order matters, and an earlier version got it wrong: it read the # geometric mean of the oracle column first and blamed the FAMILY, on a run # where the field version was winning everywhere. A mean of 0.57 and 1.37 # is 0.88, which says nothing about either. So the pathological case is # tested FIRST, per space, and the family is only blamed when nothing wins. beaten = [ r["label"] for r in results if r["trained"]["by mu"][1] > r["trained"]["shared"][1] ] if beaten: print( "\n ⚠ On " + ", ".join(beaten) + " the ORACLE is worse than the shared basis, which it\n" " cannot be on capacity -- it has strictly more information and\n" " more weights. So this is OPTIMISATION, not expressiveness: a\n" " local minimum, too few epochs, or badly scaled inputs (mu is\n" " fed raw, and its range may be small next to the normalised x\n" " the network sees). Read the oracle column as a floor, not as a\n" " measure of what conditioning can do." ) elif oracle < 1.15 and field < 1.15: print( "\n ⚠ Neither version helps: the family, not the encoder, is what\n" " limits this. Widen it before touching the encoder." ) elif recovered < 0.7: print( "\n ⚠ The oracle helps but the field version does not follow: the\n" " ENCODER is what to work on -- more modes, a wider latent, a\n" " deeper MLP -- not the family and not the basis." ) def plot(results, data, x_eval) -> None: """One row per mesh: the solution, its error, and the training losses. Args: results: What :func:`run_case` returned for each mesh. data: What :func:`prepare_family` produced. x_eval: The common grid. """ test_ref = data["test_ref"] x = np.asarray(x_eval).ravel() worst = int( np.argmax( np.asarray( jnp.linalg.norm(results[0]["classical_curve"] - test_ref, axis=(1, 2)) ) ) ) mu = np.asarray(data["mus"][B_TRAIN + worst]) figure, axes = plt.subplots( len(results), 3, figsize=(16, 4.2 * len(results)), squeeze=False, constrained_layout=True, ) for row, r in zip(axes, results): curves = [("classical DG", r["classical_curve"], "--")] + [ (name, r["curves"][name], "-") for name in NAMES ] row[0].plot(x, np.asarray(test_ref[worst, :, 0]), "k-", lw=2, label="reference") for label, values, style in curves: row[0].plot(x, np.asarray(values[worst, :, 0]), style, label=label) row[0].set_title(f"{r['label']} -- mu = " + ", ".join(f"{v:.2f}" for v in mu)) row[0].set_xlabel("x") row[0].legend(fontsize=8) for label, values, style in curves: row[1].semilogy( x, np.abs(np.asarray(values[worst, :, 0] - test_ref[worst, :, 0])), style, label=label, ) for edge in np.linspace(0.0, 1.0, r["n_cells"] + 1): row[1].axvline(edge, color="0.85", lw=0.6, zorder=0) row[1].set_title(f"pointwise error -- {r['label']} (cells in grey)") row[1].set_xlabel("x") row[1].legend(fontsize=8) for name in NAMES: row[2].semilogy(r["history"][name], label=name) row[2].set_title(f"training loss, {r['label']} ({OPTIMIZER})") row[2].set_xlabel("epoch") row[2].legend(fontsize=8) output = __file__.replace(".py", ".png") figure.savefig(output, dpi=110) print(f"\nfigure: {output}") plt.show() def main() -> None: """Can a basis be conditioned by the field alone, and does it pay?""" key = jax.random.PRNGKey(SEED) domain = Segment1D((0.0, 1.0), is_main_domain=True) x_eval = jnp.linspace(0.0, 1.0, GRID_SIZE)[:, None] key, data = prepare_family(key, domain, x_eval) # ⚠ The diagnostic FIRST, on the coarsest mesh and without any training: # if the family cannot discriminate, nothing that follows means anything. print("\n### Can this family discriminate? (no training involved)") print( " How differently the members are WRONG, shape only. Near zero, one\n" " basis serves everybody and nothing below is an effect of the\n" " conditioning. The tight family of the earlier files measures 0.0036." ) spreads = [ can_the_family_discriminate( DGSolverOperator( classical_space(n_cells, order), make_flux(order), weak_model_fn=weak_model_of, ).evaluate_batch(data["test_batch"], x_eval), data, f"{n_cells} cells, p{order}", ) for n_cells, order in CASES ] if max(spreads) < DISCRIMINATION_FLOOR: print("\n ⚠ Stop here: widen the family before reading any gain below.") results = [] for n_cells, order in CASES: key, result = run_case(key, n_cells, order, data, domain, x_eval) results.append(result) print( f" {n_cells} cells p{order} ({result['n_dofs']:2d} dof): " + " -> ".join(f"{name} {result['trained'][name][1]:.3e}" for name in NAMES) ) report(results) plot(results, data, x_eval) if __name__ == "__main__": main()