"""Le MÊME cas 1D, mais le propagateur est un MULTIGRILLE. Test de pipeline. Ce que ce fichier vérifie -------------------------- ``advection_diffusion_1d_learned_basis.py`` résout chaque EDP du lot par une factorisation dense -- imbattable à quelques dizaines de degrés de liberté. Ici on remplace **le solveur, et rien d'autre** : même famille, même base apprise, même entraînement, même mesure. La question n'est donc pas « le multigrille va-t-il plus vite » (à 24 ddl il ne peut pas), mais : **le NO hybride donne-t-il le même résultat quand le propagateur change ?** C'est un test de PIPELINE. S'il passe, le propagateur est bien un paramètre du montage et pas une propriété de l'architecture -- ce que la classe prétend -- et la même chose tiendra en 2-D ou 3-D, là où le multigrille est la seule option. Le multigrille est monté sur l'opérateur RÉEL ----------------------------------------------- C'est la différence avec un préconditionneur de commodité. Les deux niveaux portent la **base apprise** -- le même réseau, deux maillages, puisque ``m(x)`` est un champ sur le domaine et non sur une grille -- et le transfert est :class:`~...transfer.modal.CellwiseL2Transfer`, la projection L2 locale. ⚠ **Aucun transfert rapide ne conviendrait.** ``CellwiseTransfer`` prolonge par interpolation aux nœuds de la maille fille, ce qui ne coïncide avec la projection L2 que pour une base NODALE ; une base apprise n'est ni nodale ni modale connue. La projection L2 locale, elle, ne suppose rien : ``M_c^-1 B_c`` par maille. Vérifié dans ``mg_learned_basis_2levels_1d.py`` : exactitude ``1.3e-15`` sur base polynomiale, contraction ``0.165`` sur base apprise contre ``0.143`` en polynomial. ⚠ **Le cycle est monté une fois, sur la base INITIALE, et coupé du gradient.** Les matrices du transfert dépendent de la base, donc d'un paramètre appris : une époque plus tard, le préconditionneur est monté sur un opérateur légèrement périmé. C'est licite -- un préconditionneur approché coûte des itérations, il ne ment pas -- et c'est ce que ``MG.relinearise`` corrigerait, côté opérateur, si le projecteur offrait un point d'accroche par époque. Le ``stop_gradient`` dit l'autre moitié : un cycle est un accélérateur, pas une partie du modèle. ⚠ **BiCGSTAB et non CG.** Le transport rend la jacobienne non symétrique, et le CG n'y converge pas lentement -- il stagne ou diverge en rendant un itéré qui a l'air convergé. Lancer : python advection_diffusion_1d_learned_basis_mg.py """ from __future__ import annotations import functools import time import jax import jax.numpy as jnp 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.elliptic_dg_scheme import ( EllipticDGscheme, ) 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.solvers.multigrid import MG from scimba_jax.linear_approximation.solvers.smoothers import BlockJacobiSmoother from scimba_jax.linear_approximation.transfer.hierarchy import ( Hierarchy, level_from_scheme, ) from scimba_jax.linear_approximation.transfer.modal import ( CellwiseL2Transfer, structured_children_nd, ) 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 : ceux du cas de référence, à l'identique ───────────────────── DIM = 1 EPS = 0.05 MU_LOW, MU_HIGH = 0.8, 1.2 #: ⚠ Les finesses du cas de référence, chacune avec SON niveau grossier -- il #: faut deux mailles au minimum en bas, donc on divise par deux et pas plus. #: Les finesses balayées, chacune avec SON niveau grossier -- il faut deux #: mailles au minimum en bas, donc on divise par deux et pas plus. CASES = ((4, 1, 2), (4, 2, 2), (8, 1, 4), (8, 2, 4)) # (mailles, ordre, grossier) 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`` -- la source du cas où le gain est le plus net.""" return jnp.ones(()) # ⚠ **Des fonctions de MODULE, pas des lambdas.** Un callable vit dans # l'``aux_data`` et s'y compare par IDENTITÉ : une lambda reconstruite à chaque # appel donne un treedef neuf, donc aucun exécutable partagé -- et, ici, un # préconditionneur qu'on ne peut plus rafraîchir sans tout recompiler. Mesuré : # c'est cette seule cause qui faisait diverger les treedefs de deux ``MG`` # construits avec les mêmes arguments, une fois l'égalité des lisseurs réparée. def _diffusion(x: jnp.ndarray) -> jnp.ndarray: """``A = eps I``.""" return EPS * jnp.eye(DIM) def _reaction(x: jnp.ndarray) -> jnp.ndarray: """``c = 0``.""" return jnp.zeros(()) def _dirichlet(x: jnp.ndarray) -> jnp.ndarray: """``g = 0`` sur tout le bord.""" return jnp.zeros(1) # ── La famille d'EDP ───────────────────────────────────────────────────────── class LinearTransport(ScimbaPytree): """``b(x; mu) = mu_0 + mu_1 x``, porté par une FEUILLE (donc batchable).""" 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 échantillonné : ``mu``, et 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, batchable_args=False ) def weak_model_from(pde: AdvectionDiffusion1D) -> AbstractPhysicalWeakModel: """Le pont : le modèle du projecteur devient une forme faible.""" form = EllipticWeakForm( dim=DIM, A=_diffusion, b=LinearTransport(pde.mu), c=_reaction, f=source, ) return AbstractPhysicalWeakModel.from_weak_form(form, dirichlet=_dirichlet) # ── Les espaces ────────────────────────────────────────────────────────────── def make_mesh(n_cells: int, order: int) -> Mesh: """Le maillage cartésien 1D.""" 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 de Taylor d'un ordre, construite UNE fois (treedef stable).""" 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.""" def enriched(u, y, i, mesh): return local_taylor_basis(y, i, mesh, order=order, out_dim=1) * jnp.exp( 2.0 * jnp.tanh(u[0]) ) return enriched def classical_space(n_cells: int, order: int) -> VariablesDG: """L'espace DG habituel.""" return VariablesDG( basis=AnalyticBasis( nb_basis=order + 1, out_dim=1, mesh=make_mesh(n_cells, order), basis_type="scalar", local_basis=_taylor_of(order), ), nb_variables=1, ) class BasisNN(MLP): """Le multiplicateur avant bornage, partagé par toutes les mailles.""" def __init__(self, key): super().__init__(in_size=DIM, out_size=1, hidden_sizes=[16, 16], key=key) def learned_space(network, n_cells: int, order: int) -> VariablesDG: """Le MÊME champ appris, posé sur un maillage donné. ⚠ Prend un RÉSEAU et non une clé : c'est ce qui permet aux deux niveaux du cycle de partager un seul jeu de poids. ``m(x)`` est un champ sur le domaine, pas sur une grille -- rien n'a donc à être « transféré » du réseau fin au réseau grossier. """ return VariablesDG( basis=PatchwiseParametricBasis( nb_basis=order + 1, out_dim=1, mesh=make_mesh(n_cells, order), patchwise_parametric_function=network, local_basis=_enriched_of(order), basis_type="scalar", use_local_coords=True, ), nb_variables=1, ) def make_operator(space: VariablesDG, order: int, solver=None) -> DGSolverOperator: """L'opérateur : l'espace, le flux, la traduction, et le solveur.""" return DGSolverOperator( space, SIPGFlux(sigma=10.0 * order * (order + 1), h=None), solver=solver, weak_model_fn=weak_model_from, ) def make_scheme(space: VariablesDG, order: int) -> EllipticDGscheme: """Le schéma d'un NIVEAU du cycle, sur une EDP de référence de la famille. ⚠ Le ``mu`` central : le cycle est un préconditionneur, il est monté sur un membre et sert à tous. Un préconditionneur approché coûte des itérations, il ne ment pas -- contrairement à une factorisation, qui doit être exacte. """ middle = jnp.array([0.5 * (MU_LOW + MU_HIGH)] * 2) pde = weak_model_from( AdvectionDiffusion1D( Segment1D((0.0, 1.0), is_main_domain=True), middle, (jnp.zeros((1, DIM)), jnp.zeros((1, 1))), ) ) return EllipticDGscheme( pde, space, SIPGFlux(sigma=10.0 * order * (order + 1), h=None) ) def build_mg(network, n_cells: int, order: int, n_coarse: int) -> MG: """Le cycle deux niveaux, monté sur la base APPRISE elle-même. Args: network: Le réseau de la base, déjà coupé du gradient. n_cells: Mailles du niveau fin. order: Degré local. n_coarse: Mailles du niveau grossier. Returns: Le cycle. """ space_fine = learned_space(network, n_cells, order) space_coarse = learned_space(network, n_coarse, order) transfer = CellwiseL2Transfer( space_coarse, space_fine, structured_children_nd([n_coarse]) ) hierarchy = Hierarchy( [ level_from_scheme(make_scheme(space_coarse, order)), level_from_scheme(make_scheme(space_fine, order)), ], [transfer], ) return MG(hierarchy, BlockJacobiSmoother(omega=0.8), nu_pre=2, nu_post=2) # ── Les mesures ────────────────────────────────────────────────────────────── def relative_l2(prediction, reference) -> float: """Erreur L2 relative moyennée sur le lot.""" 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 prepare_family(key, domain, x_eval): """Les ``mu``, les références résolues EN LOT, et les modèles.""" 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) 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, { "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:], } @functools.lru_cache(maxsize=None) def _scheme_of_level_for(order: int): """``space -> schéma``, un objet STABLE (voir la note sur les lambdas).""" def scheme_of_level(space): return make_scheme(space, order) return scheme_of_level def run_case( key, n_cells, order, n_coarse, data, domain, x_eval, use_mg: bool, ): """Entraîne la base avec l'un des deux propagateurs, et mesure. Args: key: L'état du générateur. n_cells: Mailles du niveau fin. order: Degré local. n_coarse: Mailles du niveau grossier (ignoré si ``use_mg`` est faux). data: Ce que :func:`prepare_family` a produit. domain: Le segment. x_eval: La grille. use_mg: Multigrille si vrai, factorisation dense sinon. Returns: ``(key, résultats)``. Raises: SystemExit: Si les actifs ne sont pas exactement les poids du réseau. """ key, sub = jax.random.split(key) network = BasisNN(sub) scheme_of_level = _scheme_of_level_for(order) operator = make_operator(learned_space(network, n_cells, order), order) # ⚠ L'espace grossier est construit UNE fois ; les rafraîchissements ne font # qu'y substituer le réseau, sans jamais retoucher au maillage. coarse_space = learned_space(network, n_coarse, order) if use_mg: # ⚠ **Une ligne, et c'est tout le propos de l'API.** `with_multigrid` # monte le cycle sur la base COURANTE : transfert par projection L2 # locale (le seul valide pour une base apprise), niveaux coupés du # gradient, BiCGSTAB parce que l'advection rend la jacobienne non # symétrique, plafond de Krylov borné parce qu'un `while_loop` sous # `vmap` ne s'arrête pas par membre. L'utilisateur ne fournit que ce # que lui seul sait : les espaces grossiers -- portant LE MÊME réseau, # puisqu'un champ appris vit sur le domaine et pas sur une grille -- et # la façon de bâtir le schéma d'un niveau. operator = operator.with_multigrid( [coarse_space], scheme_for_level=scheme_of_level ) 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(network)) 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." ) 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] return key, { "seconds": seconds, "train": relative_l2( trained.evaluate_batch(data["train_batch"], x_eval), data["train_ref"] ), "test": relative_l2( trained.evaluate_batch(data["test_batch"], x_eval), data["test_ref"] ), } def main() -> None: """Les mêmes cas 1D, résolus deux fois : LU dense, puis multigrille.""" 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) rows = [] for n_cells, order, n_coarse in CASES: label = f"{n_cells} mailles, p{order}" print(f"\n=== {label} ({n_cells * (order + 1)} ddl), grossier {n_coarse} ===") # ⚠ La MÊME clé pour les deux propagateurs : la seule différence entre # les deux colonnes doit être le solveur, pas l'initialisation. key_case, _ = jax.random.split(key) key, lu = run_case( key_case, n_cells, order, n_coarse, data, domain, x_eval, use_mg=False ) print(f" LU dense : {lu['test']:.4e} en {lu['seconds']:6.1f} s") _, mg = run_case( key_case, n_cells, order, n_coarse, data, domain, x_eval, use_mg=True ) print(f" MG : {mg['test']:.4e} en {mg['seconds']:6.1f} s") rows.append((label, n_cells * (order + 1), lu, mg)) print( f"\n{'espace':16s} {'ddl':>4s} {'LU (test)':>12s} {'MG (test)':>12s} " f"{'écart':>9s} {'LU (s)':>8s} {'MG (s)':>8s}" ) print("-" * 76) for label, n_dofs, lu, mg in rows: gap = abs(mg["test"] - lu["test"]) / lu["test"] print( f"{label:16s} {n_dofs:4d} {lu['test']:12.4e} {mg['test']:12.4e} " f"{gap:8.1e} {lu['seconds']:8.1f} {mg['seconds']:8.1f}" ) worst = max(abs(mg["test"] - lu["test"]) / lu["test"] for _, _, lu, mg in rows) print(f"\nécart relatif maximal entre les deux propagateurs : {worst:.1e}") print("pipeline validée" if worst < 5e-2 else "⚠ les deux propagateurs DIVERGENT") if __name__ == "__main__": main()