"""NO hybride : le propagateur est un SOLVEUR DG exact, et on apprend la BASE. Le problème, une famille resserrée ---------------------------------- Advection-diffusion 1D stationnaire, Dirichlet homogène :: -eps u'' + b(x; mu) u' = 1 sur (0, 1), u(0) = u(1) = 0 avec ``eps = 0.05`` (régime modéré : Péclet de maille ~2.5 sur la grille grossière) et un transport qui change d'une EDP à l'autre :: b(x; mu) = mu_0 + mu_1 x, mu ~ U([0.8, 1.2] x [0.8, 1.2]) Famille **resserrée** au sens de ``uq_batched_transport_1d.py`` : des convections proches, pas un catalogue de formes sans rapport. C'est ce qui donne un sens à la question « quelle est la meilleure base pour CETTE famille ? » -- sur une famille très écartée, la meilleure base est une base riche, et un raffinement ferait mieux qu'un apprentissage. L'opérateur, et ce qui le distingue d'un DeepONet -------------------------------------------------- :class:`~scimba_jax.neural_operator.physic_no.solver_based.DGSolverOperator` a les trois temps habituels, mais le milieu est un vrai solveur : * ``encoder`` : le modèle physique devient un **schéma DG** posé sur l'espace ; * ``propagator`` : le schéma est **résolu** -- ici par une factorisation exacte (24 ddl, un LU dense est imbattable) ; le gradient traverse par le théorème des fonctions implicites, pas en dépliant Newton ; * ``decoder`` : l'expansion est **évaluée** en un point. L'état latent ``beta`` n'est donc pas un vecteur appris : ce sont les **degrés de liberté** de la solution discrète. L'opérateur non entraîné résout déjà -- c'est le DG classique -- et ce qu'on apprend est la **base**. What is learned --------------- A :class:`~...basis.general_bases.PatchwiseParametricBasis` whose local basis is ``phi_k(x) = m(x) T_k(x)``, ``m(x) = exp(2 tanh(N(x_local)))`` with ``N`` a small MLP **shared by every cell**. The multiplier carries the three requirements the learned-basis disk (``benchmarks/benchmarks_jax/ dg_learned_basis``, row ``disk_2d``) paid for one by one: *positive* (a multiplier that vanishes makes the local system singular), *equal to 1 at initialisation* (``tanh(0) = 0``, so the first epochs improve something that already solves), and *wide enough* (``exp`` of a bounded argument spans a factor of 55, where ``1 + tanh/2`` spans only 3). L'entraînement : celui d'un DeepONet ------------------------------------- Rien de spécifique ici. L'opérateur entre dans un :class:`~...approximation_spaces.physic_no_approximation_spaces.PhysicNOApproximationSpace` et c'est un :class:`~...numerical_solvers.physic_no_projectors.PhysicNOProjector` qui l'entraîne, exactement comme ``anti_derivative_data_informed.py`` entraîne son DeepONet : ``only_data=True``, la loss est l'écart aux solutions de référence. ⚠ **Le découpage encoder/propagator/decoder n'est pas cosmétique dans ce cadre.** L'espace appelle ``propagator(encoder(pde))`` UNE fois par EDP pour obtenir ``beta``, puis ``decoder(beta, x)`` en chaque point de collocation. Un solve par EDP, pas un solve par point : c'est précisément ce que la séparation des trois temps achète. Le lot ------ ⚠ **Le transport est une FEUILLE, pas un callable.** ``LinearTransport`` porte ``mu`` en enfant du pytree ; l'empilement voit un ``(B, 2)`` et tout le monde partage un exécutable. Un coefficient laissé en ``lambda`` resterait dans l'``aux_data``, et ``create_batch`` garderait celui du PREMIER modèle en jetant les autres, **sans rien lever**. C'est la sortie « famille paramétrique » que ``uq_batched_transport_1d.py`` recommande : une seule forme de transport, donc pas de registre de fonctions ni de ``select_n``. ⚠ **Le lot porte sur les modèles, jamais sur l'espace.** La base apprise est partagée par toute la famille -- c'est l'objet même de l'exercice, et c'est ce qui garde le lot possible : une géométrie, un jeu de poids, B opérateurs. Ce que l'exemple mesure ----------------------- Trois nombres, sur la MÊME grille grossière et à ddl égaux : le DG classique, le DG à base apprise sur les EDP d'entraînement, et le même sur des ``mu`` jamais vus. La référence est un DG fin (128 mailles, ordre 4) produit par le **même** opérateur, pour qu'un écart de convention ne se lise pas comme un gain. Lancer : python advection_diffusion_1d_learned_basis.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 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 # ── Paramètres ─────────────────────────────────────────────────────────────── DIM = 1 EPS = 0.05 F_FREQ = 3.0 # fréquence de la source oscillante MU_LOW, MU_HIGH = 0.8, 1.2 #: Les finesses balayées, du plus PAUVRE au plus riche. ⚠ C'est le paramètre #: central de l'exercice : optimiser une base n'a de sens que là où l'espace est #: trop pauvre pour la solution. Sur une grille déjà fine, le schéma classique #: est excellent partout sauf dans la couche limite, et l'apprentissage ne fait #: plus que redistribuer l'erreur -- il dégrade là où elle était négligeable #: pour gagner là où elle domine la norme. CASES = ((4, 1), (4, 2), (8, 1), (8, 2)) #: La référence : le même schéma, assez fin pour tenir lieu de solution exacte. N_CELLS_REF, ORDER_REF = 128, 4 B_TRAIN, B_TEST = 32, 32 # EDP d'entraînement / de test BATCH_SIZE = 16 # EDP tirées à chaque époque GRID_SIZE = 101 # points de donnée par EDP N_EPOCHS = 400 OPTIMIZER = "SS-BFGS" # PhysicNOProjector n'expose pas ENG (voir la note en tête) SEED = 0 _MAPPING = Mapping(mappings=[InvertibleFunction(lambda x: x, lambda y: y)]) def _constant_source(x: jnp.ndarray) -> jnp.ndarray: """``f = 1`` : la solution est une rampe douce plus une couche limite.""" return jnp.ones(()) def _oscillating_source(x: jnp.ndarray) -> jnp.ndarray: """``f = sin(3 pi x)`` : la difficulté est répartie sur TOUT le domaine.""" return jnp.sin(F_FREQ * jnp.pi * x[0]) #: Les deux sources, et elles ne posent pas le même problème à une base apprise. #: Avec ``f = 1`` l'erreur du schéma est concentrée dans la couche limite, qu'un #: multiplicateur ``m(x)`` bien placé corrige beaucoup ; avec un sinus, elle est #: partout et vient d'un manque de degrés POLYNOMIAUX, que moduler l'amplitude #: ne remplace pas. SOURCES = (("f = 1", _constant_source), ("f = sin(3 pi x)", _oscillating_source)) # ── La famille d'EDP ───────────────────────────────────────────────────────── class LinearTransport(ScimbaPytree): """``b(x; mu) = mu_0 + mu_1 x``, porté par une FEUILLE. ⚠ ``mu`` est un tableau ``jnp``, donc un enfant du pytree sans qu'aucun marqueur ne le déclare -- c'est l'écriture qui place, et c'est ce qui fait d'une famille d'EDP un lot empilable. En ``lambda``, le transport partirait dans l'``aux_data`` et l'empilement garderait celui du premier modèle. ⚠ Et il n'est pas entraînable : rien ne l'a déclaré, donc il est gelé. C'est la donnée de l'EDP, pas un paramètre du modèle. Args: mu: ``(2,)``, les deux coefficients du transport. """ def __init__(self, mu): self.mu = jnp.asarray(mu, dtype=float) def __call__(self, x: jnp.ndarray) -> jnp.ndarray: """``b(x)``, de forme ``(1,)``.""" return jnp.array([self.mu[0] + self.mu[1] * x[0]]) class AdvectionDiffusion1D(AbstractPhysicalModel): """Le modèle que le projecteur échantillonne : ``mu``, et les données. ⚠ Il ne porte AUCUN résidu physique. La loss est entièrement supervisée (``only_data=True``), et le résidu fort n'aurait de toute façon pas grand sens sur une expansion DG discontinue : il ne voit pas les sauts, qui sont précisément là où le schéma impose sa physique. Args: main_domain: Le segment. mu: ``(2,)``, le transport de CETTE EDP. data: ``(x, u_ref(x))``, les valeurs de référence. """ 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 : la grille est la MÊME pour toutes les EDP, donc elle # reste dans l'aux_data et n'est pas dupliquée B fois. Ce que # l'aux_data achète ici n'est pas « pas de gradient » mais « pas de # duplication » -- seules les VALEURS sont batchées. batchable_args=False, ) @functools.lru_cache(maxsize=None) def weak_model_fn_for(source): """La traduction modèle -> forme faible, pour une source donnée. ⚠ Mémoïsée sur la source : le ``weak_model_fn`` vit dans l'``aux_data`` de l'opérateur, donc une fermeture refabriquée à chaque appel donnerait un treedef différent et une compilation par construction. Args: source: ``x -> f(x)``. Returns: ``AdvectionDiffusion1D -> AbstractPhysicalWeakModel``. """ def weak_model_from(pde): return _weak_model(pde, source) return weak_model_from def _weak_model(pde: AdvectionDiffusion1D, source) -> AbstractPhysicalWeakModel: """Le pont : le modèle du projecteur devient une forme faible. C'est le seul endroit où la physique est écrite deux fois, et c'est volontaire : un projecteur raisonne en résidus, un schéma Galerkin en formes faibles, et la traduction est de la physique -- donc à l'utilisateur. Args: pde: Le modèle échantillonné (il porte ``mu``). Returns: Le modèle variationnel correspondant. """ 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) ) # ── Les deux espaces ───────────────────────────────────────────────────────── def make_mesh(n_cells: int, order: int) -> Mesh: """Le maillage cartésien 1D, quadrature exacte pour l'ordre voulu.""" 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): """La base locale de Taylor d'un ordre, construite UNE fois. ⚠ Mémoïsée, et ce n'est pas une optimisation : un callable vit dans l'``aux_data``, donc deux lambdas équivalentes construites séparément donnent deux treedefs différents et deux compilations. """ 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)`` pour un ordre donné, construite UNE fois (voir ci-dessus).""" def enriched(u, y, i, mesh): return local_taylor_basis(y, i, mesh, order=order, out_dim=1) * multiplier(u) return enriched def classical_space(n_cells: int, order: int) -> VariablesDG: """L'espace DG habituel : Taylor par maille, rien d'appris.""" 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) class BasisNN(MLP): """Le multiplicateur avant bornage. UN seul réseau pour tout le domaine.""" def __init__(self, key): super().__init__(in_size=DIM, out_size=1, hidden_sizes=[16, 16], key=key) def multiplier(u: jnp.ndarray) -> jnp.ndarray: """``exp(2 tanh(N))`` : positif, valant 1 en ``N = 0``, et large.""" return jnp.exp(2.0 * jnp.tanh(u[0])) def learned_space(key, n_cells: int, order: int) -> VariablesDG: """Le MÊME espace, dont la base locale porte un réseau. ⚠ ``use_local_coords=True`` : le réseau voit la coordonnée normalisée dans la maille, pas le point physique. Sans cela il peut apprendre à annuler des régions entières du domaine, ce qui rend la matrice DG singulière en cours d'optimisation. """ mesh = make_mesh(n_cells, order) basis = PatchwiseParametricBasis( nb_basis=order + 1, out_dim=1, mesh=mesh, patchwise_parametric_function=BasisNN(key), local_basis=_enriched_of(order), basis_type="scalar", use_local_coords=True, ) return VariablesDG(basis=basis, nb_variables=1) def make_operator(space: VariablesDG, order: int, source, solver=None): """L'opérateur : l'espace, le flux, et la traduction du modèle. Args: space: L'espace d'approximation. order: Le degré local (fixe la pénalité SIPG). source: La source de la famille. solver: La stratégie non linéaire ; le défaut de l'opérateur sinon. Returns: L'opérateur. """ return DGSolverOperator( space, SIPGFlux(sigma=10.0 * order * (order + 1), h=None), solver=solver, weak_model_fn=weak_model_fn_for(source), ) # ── Les données ────────────────────────────────────────────────────────────── def relative_l2(prediction, reference) -> float: """Erreur L2 relative moyennée sur le lot. Args: prediction: ``(B, n, 1)``. reference: ``(B, n, 1)``. Returns: Un flottant. """ 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, reference, points) -> float: """L'erreur de la MEILLEURE approximation de la famille dans l'espace. Le pendant discret de ce qu'une POD mesure : on oublie le schéma et on demande à l'espace ce qu'il pourrait faire de mieux, en projetant chaque solution de référence dessus au sens des moindres carrés sur la grille. ⚠ La matrice de la base s'obtient en passant la base canonique des ddl au DÉCODEUR : ``phi_k`` évalué aux points est exactement ``decoder(e_k, x)``. Aucune API de projection n'est nécessaire, et l'on est sûr de mesurer la base que le solve utilise, celle du réseau courant comprise. Args: operator: L'opérateur, dont on ne lit que l'espace. reference: ``(B, n, 1)``, les solutions de référence. points: ``(n, dim)``. Returns: L'erreur L2 relative moyenne de la meilleure approximation. """ shape = operator.variables.dofsl.shape n_dofs = int(np.prod(shape)) basis_matrix = jax.vmap( lambda e: operator.evaluate_at(e.reshape(shape), points)[:, 0] )(jnp.eye(n_dofs)).T # (n_points, n_dofs) def project_one(values): coefficients, *_ = jnp.linalg.lstsq(basis_matrix, values[:, 0], rcond=None) return basis_matrix @ coefficients fitted = jax.vmap(project_one)(reference)[..., None] return relative_l2(fitted, reference) def prepare_family(key, domain, x_eval, source): """Les ``mu``, les solutions de référence, et les modèles qui les portent. Calculé UNE fois pour tout le balayage : la famille et sa vérité de terrain ne dépendent pas de la finesse à laquelle on la résout. Args: key: L'état du générateur. domain: Le segment. x_eval: ``(n, 1)``, la grille commune. source: La source de la famille. Returns: ``(key, données)`` où les données réunissent lots, références et ``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"référence : DG {N_CELLS_REF} mailles, ordre {ORDER_REF} …") start = time.perf_counter() reference_op = make_operator( classical_space(N_CELLS_REF, ORDER_REF), ORDER_REF, source ) 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 pour {len(mus)} EDP") 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 run_case(key, n_cells: int, order: int, data, domain, x_eval, source): """Un cas du balayage : classique, base initiale, base apprise, meilleure approx. Args: key: L'état du générateur. n_cells: Mailles de la grille grossière. order: Degré local. data: Ce que :func:`prepare_family` a produit. domain: Le segment. x_eval: La grille commune. source: La source de la famille. Returns: ``(key, résultats)``. Raises: SystemExit: Si les paramètres actifs ne sont pas exactement les poids du réseau -- voir le critère d'acceptation du dépôt. """ train_batch, test_batch = data["train_batch"], data["test_batch"] train_ref, test_ref = data["train_ref"], data["test_ref"] coarse_op = make_operator(classical_space(n_cells, order), order, source) train_classical = coarse_op.evaluate_batch(train_batch, x_eval) test_classical = coarse_op.evaluate_batch(test_batch, x_eval) key, sub = jax.random.split(key) operator = make_operator(learned_space(sub, n_cells, order), order, source) space = PhysicNOApproximationSpace( dims={"x": DIM}, list_models=[operator], model_type="x" ) n_weights = sum(int(leaf.size) for leaf in jax.tree_util.tree_leaves(BasisNN(sub))) if int(space.ndof) != n_weights: raise SystemExit( f"{space.ndof} paramètres pour {n_weights} poids : quelque chose de " "gelé s'est glissé dans les actifs (maillage ? quadrature ?)." ) # ⚠ La ligne de CONTRÔLE : le réseau est initialisé au hasard, donc # `m = exp(2 tanh(N))` ne vaut pas exactement 1 à l'époque 0 -- la base de # départ est déjà une base ENRICHIE. Sans cette mesure on ne sait pas si le # gain vient de l'apprentissage ou de la seule perturbation de la base. train_initial = operator.evaluate_batch(train_batch, x_eval) test_initial = operator.evaluate_batch(test_batch, x_eval) # ⚠ `bc=False` : aucun résidu de bord -- le Dirichlet est imposé par le FLUX # du schéma DG, pas par un terme de loss. 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() key, 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 trained = projector.space.models[0] train_learned = trained.evaluate_batch(train_batch, x_eval) test_learned = trained.evaluate_batch(test_batch, x_eval) return key, { "n_dofs": n_cells * (order + 1), "label": f"{n_cells} mailles, p{order}", "seconds": seconds, "classical": ( relative_l2(train_classical, train_ref), relative_l2(test_classical, test_ref), ), "initial": ( relative_l2(train_initial, train_ref), relative_l2(test_initial, test_ref), ), "learned": ( relative_l2(train_learned, train_ref), relative_l2(test_learned, test_ref), ), "best_classical": ( best_approximation(coarse_op, train_ref, x_eval), best_approximation(coarse_op, test_ref, x_eval), ), "best_learned": ( best_approximation(trained, train_ref, x_eval), best_approximation(trained, test_ref, x_eval), ), "curves": (test_classical, test_learned), "history": np.asarray(projector.losses.losses_history["total"]).reshape(-1), } def geometric_mean(values) -> float: """La moyenne GÉOMÉTRIQUE d'une liste de rapports. ⚠ Géométrique et non arithmétique : ce sont des gains, donc des rapports. La moyenne arithmétique de ``1/4`` et ``4`` vaut ``2,1`` alors que les deux se compensent exactement ; la géométrique rend ``1``. Args: values: Les rapports. Returns: Leur moyenne géométrique. """ return float(np.exp(np.mean(np.log(np.asarray(values))))) def report(source_name: str, results) -> tuple[float, float]: """Les deux tableaux d'une source, et ses deux gains moyens. Args: source_name: Le nom de la source. results: Ce que :func:`run_case` a rendu pour chaque finesse. Returns: ``(gain moyen sur le solve, gain moyen sur la base)``, en test. """ print(f"\n### {source_name} — erreur du SOLVE (jeu de test)") print( f"{'espace':16s} {'ddl':>4s} {'classique':>11s} {'init.':>11s} " f"{'appris':>11s} {'gain':>7s}" ) print("-" * 66) for r in results: print( f"{r['label']:16s} {r['n_dofs']:4d} {r['classical'][1]:11.3e} " f"{r['initial'][1]:11.3e} {r['learned'][1]:11.3e} " f"{r['classical'][1] / r['learned'][1]:6.2f}x" ) print(f"\n### {source_name} — MEILLEURE APPROXIMATION dans l'espace (test)") print(f"{'espace':16s} {'ddl':>4s} {'classique':>11s} {'appris':>11s} {'gain':>7s}") print("-" * 54) for r in results: print( f"{r['label']:16s} {r['n_dofs']:4d} {r['best_classical'][1]:11.3e} " f"{r['best_learned'][1]:11.3e} " f"{r['best_classical'][1] / r['best_learned'][1]:6.2f}x" ) gain_solve = geometric_mean([r["classical"][1] / r["learned"][1] for r in results]) gain_basis = geometric_mean( [r["best_classical"][1] / r["best_learned"][1] for r in results] ) print( f"\n gain moyen (géométrique) — solve {gain_solve:.2f}x, " f"base {gain_basis:.2f}x" ) return gain_solve, gain_basis def plot_source(axes, source_name, data, results, x_eval): """Une ligne de figure pour une source : le cas le plus PAUVRE, et les loss.""" poorest = min(results, key=lambda r: r["n_dofs"]) test_classical, test_learned = poorest["curves"] test_ref = data["test_ref"] x = np.asarray(x_eval).ravel() worst = int( np.argmax(np.asarray(jnp.linalg.norm(test_classical - test_ref, axis=(1, 2)))) ) n_cells = CASES[0][0] axes[0].plot(x, np.asarray(test_ref[worst, :, 0]), "k-", lw=2, label="référence") axes[0].plot(x, np.asarray(test_classical[worst, :, 0]), "--", label="DG classique") axes[0].plot(x, np.asarray(test_learned[worst, :, 0]), "-", label="base apprise") mu = np.asarray(data["mus"][B_TRAIN + worst]) axes[0].set_title( f"{source_name} — {poorest['label']}, mu = ({mu[0]:.2f}, {mu[1]:.2f})" ) axes[0].set_xlabel("x") axes[0].legend(fontsize=8) axes[1].semilogy( x, np.abs(np.asarray(test_classical[worst, :, 0] - test_ref[worst, :, 0])), "--", label="DG classique", ) axes[1].semilogy( x, np.abs(np.asarray(test_learned[worst, :, 0] - test_ref[worst, :, 0])), "-", label="base apprise", ) for edge in np.linspace(0.0, 1.0, n_cells + 1): axes[1].axvline(edge, color="0.85", lw=0.6, zorder=0) axes[1].set_title("erreur ponctuelle (mailles en gris)") axes[1].set_xlabel("x") axes[1].legend(fontsize=8) for r in results: axes[2].semilogy(r["history"], label=r["label"]) axes[2].set_title(f"loss d'entraînement ({OPTIMIZER})") axes[2].set_xlabel("époque") axes[2].legend(fontsize=8) def main() -> None: """Deux sources, quatre finesses : où optimiser une base a un sens.""" 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] everything = [] for source_name, source in SOURCES: print(f"\n{'=' * 66}\n{source_name}\n{'=' * 66}") key, data = prepare_family(key, domain, x_eval, source) results = [] for n_cells, order in CASES: key, result = run_case(key, n_cells, order, data, domain, x_eval, source) results.append(result) print( f" {n_cells} mailles p{order} ({result['n_dofs']:2d} ddl) : " f"{N_EPOCHS} époques en {result['seconds']:5.1f} s " f"classique {result['classical'][1]:.3e} -> " f"appris {result['learned'][1]:.3e}" ) everything.append((source_name, data, results)) gains = [] for source_name, _, results in everything: gains.append(report(source_name, results)) print(f"\n{'=' * 66}") print(f"{'source':20s} {'gain solve':>12s} {'gain base':>12s}") print("-" * 46) for (source_name, _, _), (g_solve, g_basis) in zip(everything, gains): print(f"{source_name:20s} {g_solve:11.2f}x {g_basis:11.2f}x") print( f"{'toutes sources':20s} " f"{geometric_mean([g[0] for g in gains]):11.2f}x " f"{geometric_mean([g[1] for g in gains]):11.2f}x" ) fig, axes = plt.subplots( len(everything), 3, figsize=(16, 4.2 * len(everything)), squeeze=False ) for row, (source_name, data, results) in enumerate(everything): plot_source(axes[row], source_name, data, results, x_eval) plt.tight_layout() plt.savefig("advection_diffusion_1d_learned_basis.png", dpi=120) print("\nfigure : advection_diffusion_1d_learned_basis.png") if __name__ == "__main__": main()