"""Does a learned basis get better when it is TOLD which PDE it is solving? The question, and why it is asked this way ------------------------------------------ ``advection_diffusion_1d_learned_basis.py`` learns ONE basis for a whole family of PDEs -- the network is shared, only the model varies. That is the right object when the family is tight. It stops being the right object as soon as the best basis depends on WHICH member is being solved: a single basis cannot place its boundary layer in two places at once. This file changes exactly one thing: the basis network also sees ``mu``, this PDE's transport parameters :: shared phi_k(x) = m(N(x)) T_k(x) conditioned phi_k(x; z) = m(N(x, z)) T_k(x), z = mu ⚠ **``z = mu`` is handed over, not learned.** ``mu`` is a frozen leaf of the model, so this configuration adds no trainable weight of its own -- only the 32 extra first-layer weights a wider input costs (16 x 2, measured and reported below, against a network of about 350). It is therefore the ORACLE of the harder question (read ``z`` off the field ``b`` instead of off ``mu``): it gives the basis the exact information, for free. If it does not help HERE, no encoder will rescue it, and the family is what needs changing. Everything else is kept identical --------------------------------- Same family (``b(x; mu) = mu_0 + mu_1 x``, ``mu ~ U([0.8, 1.2]^2)``), same ``eps``, same source, same mesh sweep, same reference solver, same optimiser, same number of epochs, same seed -- and the two operators are trained from the SAME key, on the SAME sampled batches. Without that, an ordinary run-to-run spread would read as an effect. ⚠ **The family is deliberately TIGHT** in the shared-basis file: that is what made "one basis for everybody" a sensible thing to want. So a small gain here is not a bug, it is an answer -- it says the optimum barely depends on the PDE. The lever to widen it, if that is what the numbers say, is a smaller ``eps`` (a real boundary layer) and a ``b`` that MOVES it. Run: python advection_diffusion_1d_basis_conditioned_by_mu.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, ) 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 # ── Parameters: IDENTICAL to the shared-basis file, so the two compare ─────── DIM = 1 EPS = 0.05 MU_LOW, MU_HIGH = 0.8, 1.2 LATENT = 2 # z = mu, so the latent size is the number of parameters #: The mesh sweep, from POOREST to richest. Optimising a basis only means #: something where the space is too poor for the solution. CASES = ((4, 1), (4, 2), (8, 1), (8, 2)) #: The reference: the same scheme, fine enough to stand in for the exact one. 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``: the solution is a gentle ramp plus a boundary layer. The source where the shared basis gained most -- the error sits in the layer, which a well-placed multiplier corrects -- so it is the one where a conditioned basis has something to beat. Args: x: One physical point. Returns: The source there. """ return jnp.ones(()) # ── The family ─────────────────────────────────────────────────────────────── class LinearTransport(ScimbaPytree): """``b(x; mu) = mu_0 + mu_1 x``, carried by a LEAF. ⚠ ``mu`` is a ``jnp`` array, hence a child of the pytree without any marker declaring it -- the spelling is what places it, and it is what makes a family of PDEs a stackable batch. As a ``lambda`` the transport would land in ``aux_data`` and the stacking would keep the first model's. ⚠ And it is not trainable: nothing declared it, so it is frozen. It is the PDE's data, not a parameter of the model. Args: mu: ``(2,)``, the two transport coefficients. """ 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. """ return jnp.array([self.mu[0] + self.mu[1] * x[0]]) class AdvectionDiffusion1D(AbstractPhysicalModel): """What the projector samples: ``mu``, and the reference data. Args: main_domain: The segment. mu: ``(2,)``, 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 and is not duplicated B times. What aux_data buys here # is not "no gradient" but "no duplication" -- only the VALUES are # batched. batchable_args=False, ) def mu_of(pde: AdvectionDiffusion1D) -> jnp.ndarray: """The latent code: the PDE's parameters, read as they are. ⚠ A module-level function, not a lambda: ``code_fn`` lives in the operator's ``aux_data``, so a closure rebuilt per call would give a fresh treedef and a compilation per construction. Args: pde: The sampled model. Returns: ``mu``, which IS ``z`` here. """ return pde.mu def weak_model_of(pde: AdvectionDiffusion1D) -> AbstractPhysicalWeakModel: """The bridge: the projector's model becomes a weak form. The one place where the physics is written twice, and deliberately so: a projector reasons in residuals, a Galerkin scheme in weak forms, and the translation is physics -- so it belongs to the user. Args: pde: The sampled model (it carries ``mu``). Returns: The corresponding variational model. """ form = EllipticWeakForm( dim=DIM, A=lambda x: EPS * jnp.eye(DIM), b=LinearTransport(pde.mu), c=lambda x: jnp.zeros(()), f=source, ) return AbstractPhysicalWeakModel.from_weak_form( form, dirichlet=lambda x: jnp.zeros(1) ) # ── The three spaces ───────────────────────────────────────────────────────── 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, ) @functools.lru_cache(maxsize=None) def _taylor_of(order: int): """The local Taylor basis of one order, built ONCE. ⚠ Memoised, and not as an optimisation: 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 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])) class BasisNN(MLP): """The multiplier before bounding. ONE network for the whole domain. Args: key: Random generator state. in_size: Input width -- ``DIM`` for the shared basis, ``DIM + LATENT`` for the conditioned one. """ 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. Written that way it lands in ``children`` -- so a batch can carry one per PDE -- and, being undeclared, it is frozen. That pair is exactly what a code that is DATA rather than a parameter wants. Declaring it trainable would put an exact null direction in the Gram. ⚠ Its shape at construction is the FINAL shape: ``beta_shape`` is read off it once, and a ``z`` that changed length would change the treedef on the first call and recompile everything. Args: key: Random generator state. latent_size: Length of ``z``. """ def __init__(self, key, latent_size: int = LATENT): 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 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, not the physical point. 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 plain_operator(space: VariablesDG, order: int, solver=None) -> DGSolverOperator: """The unconditioned operator -- classical or shared-basis space. Args: space: The approximation space. order: Local degree (fixes the SIPG penalty). solver: The nonlinear strategy; the operator's default otherwise. Returns: The operator. """ return DGSolverOperator( space, make_flux(order), solver=solver, weak_model_fn=weak_model_of ) def conditioned_operator(space: VariablesDG, order: int) -> ConditionedDGOperator: """The operator whose basis is told ``mu``. Args: space: The approximation space, carrying a conditioned network. order: Local degree (fixes the SIPG penalty). Returns: The operator. """ return ConditionedDGOperator( space, make_flux(order), mu_of, latent_size=LATENT, weak_model_fn=weak_model_of, ) # ── 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 best_approximation(operator, batch, reference, points) -> float: """The error of the BEST approximation of the family in the space. The discrete counterpart of what a POD measures: forget the scheme and ask the space what it could do at best, by least-squares fitting each reference solution on the grid. ⚠ The basis matrix is obtained by passing the canonical DOF basis to the DECODER: ``phi_k`` at the points IS ``decoder(e_k, x)``. No projection API is needed, and one is sure to measure the basis the solve uses. ⚠ And for a CONDITIONED operator the basis differs per PDE, so the matrix is rebuilt per member -- which is the whole point of the comparison, and the reason this takes a batch where the shared version took none. Args: operator: The operator, of which only the space is read. batch: The stacked models, for their latent codes. reference: ``(B, n, 1)``, the reference solutions. points: ``(n, dim)``. Returns: The mean relative L2 error of the best approximation. """ shape = tuple(operator.variables.dofsl.shape) n_dofs = int(np.prod(shape)) conditioned = isinstance(operator, ConditionedDGOperator) def basis_matrix(z): def column(e): beta = jnp.concatenate([e, z]) if conditioned else e.reshape(shape) return operator.evaluate_at(beta, points)[:, 0] return jax.vmap(column)(jnp.eye(n_dofs)).T # (n_points, n_dofs) def project_one(matrix, values): coefficients, *_ = jnp.linalg.lstsq(matrix, values[:, 0], rcond=None) return matrix @ coefficients if conditioned: codes = jax.vmap(operator.latent)(batch) fitted = jax.vmap(lambda z, v: project_one(basis_matrix(z), v))( codes, reference ) else: matrix = basis_matrix(None) fitted = jax.vmap(lambda v: project_one(matrix, v))(reference) return relative_l2(fitted[..., None], reference) def prepare_family(key, domain, x_eval): """The ``mu``, the reference solutions, and the models carrying them. Computed ONCE for the whole sweep: the family and its ground truth do not depend on the mesh it is solved on. Args: key: Random generator state. domain: The segment. x_eval: ``(n, 1)``, the common grid. Returns: ``(key, data)`` where the data gather batches, references and ``mu``. """ key, sub = jax.random.split(key) mus = jax.random.uniform(sub, (B_TRAIN + B_TEST, 2), minval=MU_LOW, maxval=MU_HIGH) print(f"reference: DG {N_CELLS_REF} cells, order {ORDER_REF} ...") start = time.perf_counter() reference_op = plain_operator(classical_space(N_CELLS_REF, ORDER_REF), ORDER_REF) 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, and return it with its history. ⚠ The key is passed IN rather than split inside, so that the shared and the conditioned operator see the same sampled batches in the same order. An ordinary run-to-run spread would otherwise read as an effect. 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`: no boundary residual -- the Dirichlet condition is imposed # by the DG scheme's 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 point of the sweep: classical, shared basis, conditioned basis. 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 = plain_operator(classical_space(n_cells, order), order) classical = evaluate(coarse, data, x_eval) # ⚠ ONE key for the weights, ONE for the training, both SHARED by the two # learned operators: what differs between them must be the conditioning and # nothing else. key, weights_key, train_key = jax.random.split(key, 3) shared_op = plain_operator( learned_space(BasisNN(weights_key), n_cells, order), order ) shared_initial = evaluate(shared_op, data, x_eval) shared_op, shared_history, shared_seconds, shared_theta = train( train_key, shared_op, data, domain, x_eval ) conditioned_op = conditioned_operator( learned_space(ConditionedBasisNN(weights_key), n_cells, order), order ) conditioned_initial = evaluate(conditioned_op, data, x_eval) conditioned_op, cond_history, cond_seconds, cond_theta = train( train_key, conditioned_op, data, domain, x_eval ) return key, { "n_cells": n_cells, "order": order, "n_dofs": n_cells * (order + 1), "label": f"{n_cells} cells, p{order}", "seconds": (shared_seconds, cond_seconds), "n_theta": (shared_theta, cond_theta), "classical": classical, "shared_initial": shared_initial, "conditioned_initial": conditioned_initial, "shared": evaluate(shared_op, data, x_eval), "conditioned": evaluate(conditioned_op, data, x_eval), "best_classical": best_approximation( coarse, data["test_batch"], data["test_ref"], x_eval ), "best_shared": best_approximation( shared_op, data["test_batch"], data["test_ref"], x_eval ), "best_conditioned": best_approximation( conditioned_op, data["test_batch"], data["test_ref"], x_eval ), "curves": ( coarse.evaluate_batch(data["test_batch"], x_eval), shared_op.evaluate_batch(data["test_batch"], x_eval), conditioned_op.evaluate_batch(data["test_batch"], x_eval), ), "history": (shared_history, cond_history), } def geometric_mean(values) -> float: """The GEOMETRIC mean of a list of ratios. ⚠ Geometric and not arithmetic: these are gains, hence ratios. The arithmetic mean of ``1/4`` and ``4`` is ``2.1`` although the two cancel exactly; the geometric one returns ``1``. Args: values: The ratios. Returns: Their geometric mean. """ return float(np.exp(np.mean(np.log(np.asarray(values))))) 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 #: Below this, the family cannot tell a conditioned basis from a shared one. #: Calibrated, not guessed: the tight family of the earlier files measures #: 0.0036 and shows no gain; a sum-of-sines family measures 0.055 on 4 cells. DISCRIMINATION_FLOOR = 0.02 def can_the_family_discriminate(results, data, x_eval) -> None: """How differently the members of the family are solved, per space. ⚠ Run FIRST and costing nothing: conditioning can only pay where the best basis depends on the PDE. If every member is wrong in the same way, one basis enriched there serves everybody and a gain would measure the extra weights rather than the conditioning. Args: results: What :func:`run_case` returned for each mesh. data: What :func:`prepare_family` produced. x_eval: The common grid (unused; kept for symmetry with the caller). """ print("\n### Can this family discriminate? (no training involved)") print( " How differently the members are WRONG, shape only, on the classical\n" " scheme. Near zero, one basis serves everybody." ) spreads = [] for r in results: spread, _ = shape_spread(r["curves"][0], data["test_ref"]) spreads.append(spread) verdict = ( "cannot discriminate" if spread < DISCRIMINATION_FLOOR else "conditioning has room" ) print(f" {r['label']:14s} 1 - corr = {spread:.4f} {verdict}") if max(spreads) < DISCRIMINATION_FLOOR: print( "\n ⚠ This family cannot show what conditioning does: b stays\n" " positive and barely varies, so the difficulty never moves.\n" " A smooth sum of sines measures 15x higher -- see\n" " advection_diffusion_1d_basis_conditioned_by_field.py." ) def report(results, data, x_eval) -> None: """The two tables, and the verdict this file exists to produce. Args: results: What :func:`run_case` returned for each mesh. data: What :func:`prepare_family` produced. x_eval: The common grid. """ can_the_family_discriminate(results, data, x_eval) print("\n### SOLVE error (test set)") print( " `classical` is the UNLEARNED basis (plain Taylor, no network at all);\n" " `init.` is the same learned basis before training. ⚠ The two differ,\n" " and the second is the honest baseline: at epoch 0 the network is\n" " random, so `m = exp(2 tanh(N))` is not 1 and the starting basis is\n" " ALREADY an enriched one. Reading a gain against `classical` alone\n" " credits training with what a mere perturbation of the basis did." ) print( f"\n{'space':14s} {'dof':>4s} {'classical':>10s} {'sh.init':>10s} " f"{'shared':>10s} {'cd.init':>10s} {'conditioned':>11s} {'cd/sh':>7s}" ) print("-" * 82) for r in results: print( f"{r['label']:14s} {r['n_dofs']:4d} {r['classical'][1]:10.3e} " f"{r['shared_initial'][1]:10.3e} {r['shared'][1]:10.3e} " f"{r['conditioned_initial'][1]:10.3e} {r['conditioned'][1]:11.3e} " f"{r['shared'][1] / r['conditioned'][1]:6.2f}x" ) print("\n### BEST APPROXIMATION in the space (test set)") print( f"{'space':16s} {'dof':>4s} {'classical':>11s} {'shared':>11s} " f"{'conditioned':>12s} {'cond/shared':>12s}" ) print("-" * 72) for r in results: print( f"{r['label']:16s} {r['n_dofs']:4d} {r['best_classical']:11.3e} " f"{r['best_shared']:11.3e} {r['best_conditioned']:12.3e} " f"{r['best_shared'] / r['best_conditioned']:11.2f}x" ) print("\n### Parameter count -- the conditioning is nearly free") print(f"{'space':16s} {'shared':>8s} {'conditioned':>12s} {'extra':>7s}") print("-" * 46) for r in results: shared_theta, cond_theta = r["n_theta"] print( f"{r['label']:16s} {shared_theta:8d} {cond_theta:12d} " f"{cond_theta - shared_theta:+7d}" ) gain_solve = geometric_mean([r["shared"][1] / r["conditioned"][1] for r in results]) gain_basis = geometric_mean( [r["best_shared"] / r["best_conditioned"] for r in results] ) print( f"\n conditioned vs shared (geometric mean) -- solve {gain_solve:.2f}x, " f"basis {gain_basis:.2f}x" ) if gain_basis < 1.15: print( "\n ⚠ The gain is in the noise. That is an ANSWER, not a bug: on\n" " this family the optimum barely depends on the PDE. Widen it\n" " (smaller eps, a b that MOVES the layer) before reading any\n" " further, and before writing the field encoder." ) def plot(results, data, x_eval) -> None: """One row per mesh size: the solution, its error, and the two losses. ⚠ Several rows, and the same PDE on each. A single mesh cannot tell whether the conditioning buys accuracy or merely buys what refining would have given anyway -- and the answer is the whole question, since a basis is worth learning only where the space is too poor for the solution. Args: results: What :func:`run_case` returned for each mesh. data: What :func:`prepare_family` produced. x_eval: The common grid. """ # One case per mesh size, the cheapest of each -- the poorest spaces, where # optimising a basis has something to do. rows = [ min((r for r in results if r["n_cells"] == n), key=lambda r: r["n_dofs"]) for n in sorted({r["n_cells"] for r in results}) ] test_ref = data["test_ref"] x = np.asarray(x_eval).ravel() # ⚠ The SAME member on every row, picked once on the poorest space: what # varies down a column must be the mesh, not the PDE. poorest_classical = rows[0]["curves"][0] worst = int( np.argmax( np.asarray(jnp.linalg.norm(poorest_classical - test_ref, axis=(1, 2))) ) ) mu = np.asarray(data["mus"][B_TRAIN + worst]) figure, axes = plt.subplots( len(rows), 3, figsize=(16, 4.2 * len(rows)), squeeze=False, constrained_layout=True, ) for row, result in zip(axes, rows): classical, shared, conditioned = result["curves"] curves = ( (classical, "--", "classical DG"), (shared, "-", "shared basis"), (conditioned, "-", "conditioned"), ) row[0].plot(x, np.asarray(test_ref[worst, :, 0]), "k-", lw=2, label="reference") for values, style, label in curves: row[0].plot(x, np.asarray(values[worst, :, 0]), style, label=label) row[0].set_title(f"{result['label']} -- mu = ({mu[0]:.2f}, {mu[1]:.2f})") row[0].set_xlabel("x") row[0].legend(fontsize=8) for values, style, label 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, result["n_cells"] + 1): row[1].axvline(edge, color="0.85", lw=0.6, zorder=0) row[1].set_title(f"pointwise error -- {result['label']} (cells in grey)") row[1].set_xlabel("x") row[1].legend(fontsize=8) shared_history, cond_history = result["history"] row[2].semilogy(shared_history, "--", label="shared basis") row[2].semilogy(cond_history, "-", label="conditioned") row[2].set_title(f"training loss, {result['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: """Four meshes: does telling the basis which PDE it solves help?""" 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) 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): " f"classical {result['classical'][1]:.3e} -> " f"shared {result['shared'][1]:.3e} -> " f"conditioned {result['conditioned'][1]:.3e}" ) report(results, data, x_eval) plot(results, data, x_eval) if __name__ == "__main__": main()