"""Un LOT d'EDP résolu d'un coup, et l'incertitude qu'on en tire (advection-diffusion 1D). -d/dx (A du/dx) + b(x) du/dx = f sur [0, 1], u(0) = u(1) = 0 La diffusion ``A`` et la source ``f`` sont les mêmes pour tout le monde ; c'est le champ de transport ``b`` qui change d'une EDP à l'autre -- et ce sont de vraies FONCTIONS différentes (``1+x``, ``1+x^2``, ``1+sin(pi x/2)``, ...), pas un même profil à coefficient variable. On les résout toutes dans **un seul programme compilé**, puis on lit la moyenne et l'écart-type des solutions : c'est la quantification d'incertitude la plus simple qui soit, et elle sert surtout de contrôle -- si la mécanique de batch est fausse quelque part, la variance le montre avant n'importe quel test. Comment un lot d'EDP se construit --------------------------------- Deux briques, qui viennent des opérateurs neuronaux et marchent telles quelles côté forme faible : * :class:`~scimba_jax.utils.functional_fields.AbstractFunctionalField` fait d'une FONCTION une feuille de pytree. Elle range la fonction dans un registre de classe et n'en garde que l'indice, un entier ; l'appel devient un ``lax.switch`` sur cet indice. C'est ce qui permet à ``b`` de varier d'une EDP à l'autre, ce qu'un callable nu ne permet pas ; * :meth:`~scimba_jax.physical_models.abstract_physical_weak_model.AbstractPhysicalWeakModel.create_batch` empile les feuilles d'une liste de modèles en UN modèle batché. ⚠ **Ce qui varie doit être une FEUILLE.** Un coefficient laissé en callable nu reste dans l'``aux_data``, et l'empilement garde alors celui du PREMIER modèle en jetant les autres, **sans rien lever**. Le champ fonctionnel est exactement ce qui évite ce piège. ⚠ **Pourquoi pas ``jax.vmap(EllipticFEscheme.solve)``.** ``Galerkin.solve`` finit par ``self.variables.dofsl = ...``, une mutation Python, incompatible avec le traçage. On passe donc par les briques PURES -- ``Galerkin.factorise`` et ``Galerkin._make_back_solve_fn`` -- comme ``solve_2d_helmholtz_robin_unstructured_disk_batched_sources.py``. ⚠ **Mais ici l'OPÉRATEUR change**, pas seulement le second membre. Dans l'exemple Helmholtz seule la source variait, donc la factorisation était partagée et sortie du ``vmap`` ; ici ``b`` est dans la forme bilinéaire, donc chaque EDP a sa propre matrice et la factorisation entre DANS le ``vmap``. En 1D c'est le bon outil -- quelques centaines de ddl, une LU dense par membre est gratuite. En 2D/3D ce sera un Krylov préconditionné, et c'est la seule ligne à changer. Ce que ça coûte (mesuré, CPU M3, Q2, 8 EDP) ------------------------------------------ ⚠ Il faut séparer trois choses que le premier jet de ce fichier confondait : construire l'espace et TRACER, COMPILER, et résoudre. Sans ça on lit des centaines de millisecondes là où l'arithmétique en vaut une, et on croit que le lot est lent alors qu'on mesure un traçage :: mailles ddl trace+compile SOLVE 8 EDP boucle Python 64 129 0.83 s 1.0 ms 4.08 s 128 257 0.65 s 3.3 ms 4.14 s 256 513 0.67 s 16.8 ms 4.13 s 512 1025 0.72 s 86.5 ms 4.30 s Le solve suit la factorisation dense, donc grimpe vite avec les ddl ; le traçage, lui, ne dépend pas de la taille -- c'est le même graphe. Et les deux solutions coïncident à 8,8e-15. ⚠ **La colonne « boucle Python » est PLATE**, et c'est ce qu'elle a de plus instructif : 4,1 s quelle que soit la finesse. Elle ne mesure pas des solves, elle mesure huit compilations. Le gain du lot n'est donc pas « du calcul plus rapide » mais « une compilation au lieu de B » -- et il s'évapore le jour où on résout assez gros pour que l'arithmétique domine. ⚠ **Le registre des champs fonctionnels est global au processus et ne décroît jamais**, et c'est la seule chose à savoir avant d'en faire un usage sérieux. ``function_registry`` est un attribut de CLASSE : chaque champ construit y ajoute une entrée pour de bon. Or ``__call__`` fait ``lax.switch`` sur TOUTES les branches enregistrées -- pas seulement celles du lot -- et ``lax.switch`` sous ``vmap`` dégénère en ``select_n``, donc les exécute toutes. Le coût par évaluation suit donc le nombre de fonctions JAMAIS enregistrées. Mesuré, lot figé à B=4, en ne faisant grossir que le registre avec des fonctions que le lot n'utilise pas : 566 ms à 8 entrées, 627 à 20, 765 à 68, 800 à 132. Trois sorties, par ordre de propreté croissante : * une **sous-classe jetable par lot** (registre neuf, largeur = B) -- vérifié, le registre reste à 4 quand celui de l'exemple monte à 16. En échange, une classe neuve est un type de pytree neuf, donc aucun exécutable n'est partagé d'un lot à l'autre : bon pour UN lot, mauvais pour une boucle de lots ; * si la famille est **paramétrique** (une seule forme, un coefficient qui varie), ne pas passer par le registre du tout : une feuille ``jnp`` donne une branche unique et le plein partage d'exécutable. Ce n'est pas le cas ici, où les formes diffèrent vraiment ; * **tabuler** : une forme faible ne lit jamais ``b`` qu'aux points de quadrature, et ceux-là sont FIXÉS par le maillage. On peut donc évaluer le ``b`` de chaque membre une fois, hors du lot, et ne transporter qu'un tableau ``(B, n_points)`` -- plus de registre, plus de ``switch``, « fonctions différentes » devient « données différentes ». ⚠ C'est précisément ce qui sépare ce cas de celui des opérateurs neuronaux : un PINN retire ses points de collocation à chaque époque, donc la fonction doit rester appelable ; une quadrature FEM, elle, ne bouge pas. Non implémenté ici. """ from __future__ import annotations import copy import time from pathlib import Path import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.linear_approximation.basis.analytic_bases import local_lagrange_basis from scimba_jax.linear_approximation.basis.general_bases import AnalyticBasis from scimba_jax.linear_approximation.galerkin.fem.elliptic_fe_scheme import ( EllipticFEscheme, ) from scimba_jax.linear_approximation.meshes.mesh import Mesh from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.linear_approximation.variables.variables_fe import VariablesFE from scimba_jax.mapping.mapping import InvertibleFunction, Mapping 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.weak_boundary_conditions import Dirichlet from scimba_jax.utils.functional_fields import AbstractFunctionalField ORDER = 2 N_CELLS = 512 DIFFUSION = 0.05 class TransportField(AbstractFunctionalField): """``b(x)``, rangé dans un registre pour devenir une feuille de pytree. Le motif de l'API : une sous-classe qui redéclare ``function_registry`` et ``branches``. Ce n'est pas une redondance -- ce sont des attributs de CLASSE, et les redéclarer est ce qui empêche deux champs de natures différentes de partager un registre et de se renvoyer les fonctions l'un de l'autre. La fonction est rangée dans le registre, le champ n'en garde que l'indice (un entier, donc une feuille), et l'appel devient un ``lax.switch`` sur cet indice. C'est ce qui permet à ``b`` de varier d'une EDP à l'autre, ce qu'un callable nu ne permet pas : il resterait dans l'``aux_data``. """ label: str = "transport" function_registry: list = [] branches: list = [] # ⚠ Une famille RESSERRÉE, et c'est ce qui fait que l'UQ dit quelque chose. # Les huit transports valent tous 1 en x=0 et 2 en x=1 : même comportement # global, formes intérieures différentes. Avec des transports très écartés # (x contre x^2 contre 1+x), la bande d'écart-type couvre presque toute la # figure et ne mesure plus que l'étendue du catalogue -- ce qui est un choix de # l'utilisateur, pas une incertitude du problème. Resserrés, l'écart-type dit # ce qu'on veut lui faire dire : de combien la solution bouge quand on ne # connaît la forme du transport qu'approximativement. TRANSPORTS = { "1 + x": lambda x: jnp.array([1.0 + x[0]]), "1 + x^2": lambda x: jnp.array([1.0 + x[0] ** 2]), "1 + x^3": lambda x: jnp.array([1.0 + x[0] ** 3]), "1 + x(3-x)/2": lambda x: jnp.array([1.0 + x[0] * (3.0 - x[0]) / 2.0]), "1 + sin(pi x/2)": lambda x: jnp.array([1.0 + jnp.sin(0.5 * jnp.pi * x[0])]), "1 + (1-cos(pi x))/2": lambda x: jnp.array( [1.0 + 0.5 * (1.0 - jnp.cos(jnp.pi * x[0]))] ), "1 + tanh(2x)/tanh(2)": lambda x: jnp.array( [1.0 + jnp.tanh(2.0 * x[0]) / jnp.tanh(2.0)] ), "1 + (e^x-1)/(e-1)": lambda x: jnp.array( [1.0 + (jnp.exp(x[0]) - 1.0) / (jnp.e - 1.0)] ), } # ⚠ Les champs sont construits UNE FOIS, ici, et réutilisés. Les construire # dans `make_model` -- ce que faisait ce fichier -- ré-enregistre les MÊMES # fonctions à chaque appel : mesuré, le registre passe de 8 à 16, 24, 32 sur # quatre constructions du même lot, alors qu'il n'y a jamais eu que 8 fonctions. # Et ce n'est pas cosmétique : chaque entrée du registre ajoute **2 `select_n` # au graphe tracé, définitivement, qu'elle serve ou non** -- mesuré 137 pour un # registre de 16, 377 pour 136, parfaitement linéaire. Réutiliser les objets # fige le registre (+0 sur quatre passes) sans rien changer au mécanisme. TRANSPORT_FIELDS = {name: TransportField(fn) for name, fn in TRANSPORTS.items()} def make_model(transport: TransportField) -> AbstractPhysicalWeakModel: """Le modèle variationnel d'UNE des EDP du lot. Args: transport: Le champ de transport, pris dans :data:`TRANSPORT_FIELDS`. ⚠ Un CHAMP déjà construit, pas une fonction : en construire un ici ré-enregistrerait la même fonction à chaque appel. Returns: Le modèle, prêt à être empilé avec ses semblables. """ form = EllipticWeakForm( dim=1, A=lambda _x: jnp.eye(1) * DIFFUSION, # ⚠ Le SEUL coefficient emballé : c'est le seul qui varie. Emballer # aussi `A` ou `f` marcherait, mais ajouterait des branches au switch # pour des fonctions identiques d'un membre à l'autre. b=transport, c=lambda _x: jnp.array(0.0), f=lambda _x: jnp.ones(()), ) model = AbstractPhysicalWeakModel(dim=1) model.add_weak_form("main", form) for side in ("west", "east"): model.add_boundary_condition(side, Dirichlet(lambda _x: jnp.zeros(1))) return model def make_scheme(model: AbstractPhysicalWeakModel) -> EllipticFEscheme: """L'espace FEM, construit UNE fois et partagé par tout le lot. Le maillage, la base et la numérotation ne dépendent pas de ``b`` : les rebâtir par EDP serait payer B fois une géométrie identique, et empêcherait le lot de tenir dans un seul programme compilé. Args: model: Un modèle du lot, n'importe lequel. Returns: Le schéma de référence. """ mesh = Mesh( dim=1, n_cells=(N_CELLS,), ref_quad=UnitSquareTensorized(dim=1, order=2 * ORDER + 2), mapping=Mapping(mappings=[InvertibleFunction(lambda x: x, lambda y: y)]), ) basis = AnalyticBasis( nb_basis=ORDER + 1, out_dim=1, mesh=mesh, basis_type="scalar", local_basis=lambda y, i, m: local_lagrange_basis( y, i, m, order=ORDER, out_dim=1 ), ) return EllipticFEscheme(model, VariablesFE(basis=basis, nb_variables=1)) def compile_batch(models: list[AbstractPhysicalWeakModel]): """Prépare le lot : l'espace FEM, le modèle empilé, et la callable jitée. ⚠ La callable est RENDUE plutôt qu'appelée, pour que l'appelant puisse la garder. Une fonction jitée neuve a un cache vide : la rebâtir à chaque résolution retrace tout le résidu, ce qui fait lire des centaines de millisecondes là où l'arithmétique en vaut une. Args: models: Les modèles, un par EDP. Returns: ``(scheme, run, batched)`` -- ``run(batched)`` rend ``(B, n_ddl, 1)``. """ reference = make_scheme(models[0]) batched = type(models[0]).create_batch(models) back_solve = EllipticFEscheme._make_back_solve_fn() start = reference._initial_dofs() def solve_one(pde): # ⚠ `copy.copy` et non un schéma neuf : on veut LE MÊME espace, avec un # modèle différent -- exactement ce que `_with_pde` fait dans les # conteneurs multi-patchs. scheme = copy.copy(reference) scheme.pde = pde factorisation = EllipticFEscheme.factorise(scheme, start) return back_solve(scheme, start, factorisation.lu, factorisation.pivots) return reference, jax.jit(jax.vmap(solve_one)), batched def nodal_positions(scheme: EllipticFEscheme) -> np.ndarray: """Abscisses des ddl, pour tracer sans réévaluer la base. Args: scheme: Le schéma. Returns: ``(n_ddl,)``. """ return np.linspace(0.0, 1.0, scheme.variables.n_nodes_total) if __name__ == "__main__": build_start = time.perf_counter() models = [make_model(field) for field in TRANSPORT_FIELDS.values()] build_seconds = time.perf_counter() - build_start print(f"{len(models)} EDP : {', '.join(TRANSPORTS)}") print( f"registre : {len(TransportField.function_registry)} fonctions, " f"donc {2 * len(TransportField.function_registry)} select_n dans le graphe" ) # ── Le lot, decompose ──────────────────────────────────────────────── # ⚠ La callable jitee est GARDEE entre les deux appels. Rappeler # `solve_batch` mesurerait un nouveau tracage a chaque fois, et non # l'arithmetique -- erreur commise en ecrivant ce fichier, et qui faisait # lire 0,6 s la ou le solve vaut 1 ms. start = time.perf_counter() scheme, run, batched = compile_batch(models) trace_seconds = time.perf_counter() - start start = time.perf_counter() dofs = jax.block_until_ready(run(batched)) cold_seconds = time.perf_counter() - start start = time.perf_counter() for _ in range(20): warm = run(batched) jax.block_until_ready(warm) warm_seconds = (time.perf_counter() - start) / 20 # ── La reference : une boucle Python sur les memes EDP ──────────────── reference = make_scheme(models[0]) back_solve = EllipticFEscheme._make_back_solve_fn() zero = reference._initial_dofs() start = time.perf_counter() loop = [] for model in models: one = copy.copy(reference) one.pde = model factorisation = EllipticFEscheme.factorise(one, zero) loop.append(back_solve(one, zero, factorisation.lu, factorisation.pivots)) loop = jax.block_until_ready(jnp.stack(loop)) loop_seconds = time.perf_counter() - start print(f"\n{'etape':34s} {'secondes':>10s}") print("-" * 45) print(f"{'construction des modeles':34s} {build_seconds:10.2f}") print(f"{'espace FEM + trace du lot':34s} {trace_seconds:10.2f}") print(f"{'compilation (1er appel)':34s} {cold_seconds:10.2f}") print(f"{'SOLVE des 8 EDP (arithmetique)':34s} {warm_seconds:10.4f}") print(f"{'boucle Python, meme resultat':34s} {loop_seconds:10.2f}") print(f"{'ecart lot / boucle':34s} {float(jnp.max(jnp.abs(dofs - loop))):10.2e}") values = np.asarray(dofs[:, :, 0]) mean, std = values.mean(axis=0), values.std(axis=0) # ⚠ Le controle qui compte. Si l'empilement avait garde le premier modele # en jetant les autres -- le piege de l'aux_data -- toutes les colonnes # seraient egales et l'ecart-type serait nul, ce qui RESSEMBLE a un # resultat au lieu de lever. print(f"{'ecart-type max (non degenere)':28s} {std.max():9.2e}") x = nodal_positions(scheme) figure, axes = plt.subplots(1, 2, figsize=(11, 4.2), constrained_layout=True) for (label, _), row in zip(TRANSPORTS.items(), values): axes[0].plot(x, row, lw=1.0, alpha=0.8, label=label) axes[0].set(title=f"{len(models)} EDP, un seul vmap", xlabel="x", ylabel="u") axes[0].legend(fontsize=7, ncol=2) axes[0].grid(alpha=0.25) axes[1].fill_between( x, mean - std, mean + std, alpha=0.3, label="moyenne ± écart-type" ) axes[1].plot(x, mean, lw=2, label="moyenne") axes[1].set(title="incertitude sur la forme du transport", xlabel="x", ylabel="u") axes[1].legend() axes[1].grid(alpha=0.25) output = Path(__file__).with_name("uq_batched_transport_1d.png") figure.savefig(output, dpi=110) print(f"\nfigure : {output}") plt.show()