"""Le même lot d'EDP, mais résolu par Krylov + multigrille. Même problème que ``uq_batched_transport_1d.py`` -- même advection-diffusion 1D, mêmes huit transports, même espace. **Seul le solveur change**, donc la LU dense de l'autre fichier sert de vérité terrain et la comparaison est exacte. L'idée tient en une phrase : **on ne batche que le solve**. L'espace FEM et les multigrilles sont construits hors du ``vmap`` ; celui-ci n'enveloppe que « assembler le second membre, itérer ». ⚠ **Et la construction ne peut PAS entrer dans le ``vmap``** -- ni celle du multigrille, ni celle des niveaux grossiers. ``TracerArrayConversionError``, les tracers venant de la construction de la quadrature et du maillage (``abstractquad.py``, ``Mesh.__init__``) atteintes en bâtissant la hiérarchie. ⚠⚠ Et le piège est pire que l'échec : **ça passe parfois**. Un multigrille construit sous ``vmap`` puis suivi d'un simple ``cycle`` ne lève rien ; c'est en le mettant dans la boucle de ``bicgstab`` que ça casse, parce que le ``while_loop`` force réellement les valeurs. Une construction sous traçage qui « marche » tant qu'on ne force rien est plus dangereuse qu'une qui échoue franchement -- d'où la règle : **construire dehors, empiler, ne vmaper que le calcul**. ⚠ **Le multigrille est partagé par tout le lot, et c'est légitime.** Un préconditionneur n'a pas à être exact ; seul l'opérateur doit l'être, et lui porte bien le modèle de chaque membre. Mesuré sur ces huit EDP, les diagonales ne s'écartent que de 4,1e-6 -- régime dominé par la diffusion. Ça cesserait d'être vrai en régime advectif. ⚠ **``bicgstab`` et non CG** : l'advection rend l'opérateur non symétrique, et un CG y stagne ou diverge -- ce que ``linear_solve`` avertit. Mesuré (CPU M3, Q2, 512 mailles, 1025 ddl, 8 EDP):: solveur solve du lot ecart / LU LU dense batchee (reference) 85 ms -- bicgstab nu 98 ms 1.1e+02 ne converge pas bicgstab + MG partage 13.9 ms 4.5e-10 Six fois la LU directe, et une réponse juste là où le Krylov nu n'en donne aucune. ⚠ La LU dense reste pourtant le bon outil en 1D : ce rapport est ce que le multigrille promet en 2D/3D, où factoriser cesse d'être une option. """ 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 # ⚠ La physique n'est PAS recopiee : c'est ce qui garantit que les deux fichiers # resolvent le meme probleme et que la comparaison porte sur le solveur seul. from uq_batched_transport_1d import ( # noqa: E402 ORDER, TRANSPORT_FIELDS, TRANSPORTS, compile_batch, make_model, nodal_positions, ) 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.solvers.krylov import linear_solve from scimba_jax.linear_approximation.solvers.multigrid import MG from scimba_jax.linear_approximation.solvers.smoothers import DampedJacobiSmoother from scimba_jax.linear_approximation.transfer.hierarchy import ( build_hierarchy_structured, ) from scimba_jax.linear_approximation.variables.variables_fe import VariablesFE from scimba_jax.mapping.mapping import InvertibleFunction, Mapping N_CELLS, N_LEVELS = 512, 2 TOL, MAX_ITER = 1e-10, 500 def make_scheme(model, n_cells: int) -> EllipticFEscheme: """L'espace FEM à une finesse donnée. ⚠ ``n_cells`` en argument : la hiérarchie multigrille demande le même espace à plusieurs finesses. Args: model: Le modèle variationnel. n_cells: Mailles sur ``[0, 1]``. Returns: Le schéma. """ mesh = Mesh( dim=1, n_cells=(int(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 stacked_multigrids(models): """UN multigrille par membre, bâtis dehors puis EMPILÉS. ⚠ C'est la seule façon d'avoir un préconditionneur propre à chaque EDP : le construire dans le ``vmap`` échoue, mais un ``MG`` EST un pytree (40 feuilles ici), donc on peut en bâtir B et empiler leurs feuilles exactement comme ``create_batch`` empile des modèles. Le ``vmap`` prend alors le multigrille en second argument. ⚠ Ce que ça coûte : les B constructions sont séquentielles et concrètes -- mesuré 3,97 s pour huit, contre 0,86 s pour un seul. Args: models: Les modèles, un par EDP. Returns: Le multigrille empilé, avec un axe de lot sur chaque feuille. """ grids = [ MG( build_hierarchy_structured( lambda counts, m=model: make_scheme(m, counts[0]), (N_CELLS,), N_LEVELS, ), DampedJacobiSmoother(2.0 / 3.0), nu_pre=2, nu_post=2, ) for model in models ] _, treedef = jax.tree_util.tree_flatten(grids[0]) leaves = [jax.tree_util.tree_flatten(g)[0] for g in grids] return jax.tree_util.tree_unflatten( treedef, [jnp.stack(group) for group in zip(*leaves)] ) def batched_krylov(models, precond=None, per_member=None): """Le seul endroit où il y a un ``vmap`` : autour du solve. Args: models: Les modèles, un par EDP. precond: ``r -> M^-1 r`` partagé, ou ``None``. per_member: Un multigrille EMPILÉ, qui devient alors un argument vmapé -- exclusif avec ``precond``. Returns: ``(run, args)`` -- ``run(*args)`` rend ``((B, n_ddl, 1), (B,))``, les ddl et le résidu final de chaque membre. """ reference = make_scheme(models[0], N_CELLS) assembly = EllipticFEscheme._make_assembly_fn() jvp = EllipticFEscheme._make_jvp_base_fn() zero = reference._initial_dofs() shape = zero.shape def solve_one(pde, grid=None): # ⚠ L'OPERATEUR porte toujours le modele du membre ; seul le # preconditionneur est, ou non, partage. scheme = copy.copy(reference) scheme.pde = pde own = precond if grid is None else (lambda r: grid.cycle(jnp.zeros_like(r), r)) rhs = -assembly(scheme, zero).reshape(-1) matvec = lambda v: jvp(scheme, zero, v.reshape(shape)).reshape(-1) # noqa: E731 # ⚠ Le RESIDU est remonte, pas jete. Un lot qui tape le plafond # d'iterations rend une reponse fausse sans rien lever -- mesure a # 4.9e-02 d'ecart quand le preconditionneur partage sort de son regime. # Sans cette valeur, l'exemple ne saurait pas le dire. correction, _, residual = linear_solve( "bicgstab", matvec, rhs, TOL, own, MAX_ITER ) return (zero.reshape(-1) + correction).reshape(shape), residual batched = type(models[0]).create_batch(models) if per_member is None: return jax.jit(jax.vmap(lambda pde: solve_one(pde))), (batched,) return jax.jit(jax.vmap(solve_one)), (batched, per_member) def timed(run, batched, repeats: int = 5): """``(compilation, arithmétique, solution)`` d'une callable de lot. Args: run: La callable jitée, gardée entre les appels. batched: Le tuple d'arguments à lui passer. repeats: Appels à chaud moyennés. Returns: ``(secondes de compilation, secondes par appel, dofs)``. """ start = time.perf_counter() dofs = jax.block_until_ready(run(*batched)) cold = time.perf_counter() - start start = time.perf_counter() for _ in range(repeats): again = run(*batched) jax.block_until_ready(again) warm = (time.perf_counter() - start) / repeats return cold - warm, warm, dofs if __name__ == "__main__": models = [make_model(field) for field in TRANSPORT_FIELDS.values()] print(f"{len(models)} EDP : {', '.join(TRANSPORTS)}") # 1. la reference : la LU dense batchee du fichier voisin scheme, run_direct, batched = compile_batch(models) direct = timed(run_direct, (batched,)) rows = [("LU dense batchee (reference)", direct[0], direct[1], direct[2], None)] truth = direct[2] # 2. LE multigrille : bati une fois, hors du vmap start = time.perf_counter() hierarchy = build_hierarchy_structured( lambda counts: make_scheme(models[0], counts[0]), (N_CELLS,), N_LEVELS ) multigrid = MG(hierarchy, DampedJacobiSmoother(2.0 / 3.0), nu_pre=2, nu_post=2) print(f"setup du multigrille (une fois) : {time.perf_counter() - start:.2f} s") # 3. le lot, avec et sans preconditionneur start = time.perf_counter() per_member = stacked_multigrids(models) print( f"setup des {len(models)} multigrilles par membre : " f"{time.perf_counter() - start:.2f} s" ) for tag, precond, grids in ( ("bicgstab nu", None, None), ( "bicgstab + MG partage", lambda r: multigrid.cycle(jnp.zeros_like(r), r), None, ), ("bicgstab + MG par membre", None, per_member), ): run, args = batched_krylov(models, precond, grids) compile_s, solve_s, (dofs, residual) = timed(run, args) rows.append((tag, compile_s, solve_s, dofs, float(jnp.max(residual)))) print( f"\n{'solveur':30s} {'compile':>9s} {'solve':>10s} {'ecart/LU':>10s}" f" {'residu max':>11s}" ) print("-" * 75) for tag, compile_s, solve_s, dofs, residual in rows: gap = float(jnp.max(jnp.abs(dofs - truth))) shown = " --" if residual is None else f"{residual:11.1e}" print(f"{tag:30s} {compile_s:8.2f}s {solve_s * 1e3:9.1f}ms {gap:10.1e} {shown}") # 4. l'incertitude, lue sur la solution du multigrille values = np.asarray(rows[-1][3][:, :, 0]) mean, std = values.mean(axis=0), values.std(axis=0) x = nodal_positions(scheme) figure, axis = plt.subplots(figsize=(6.5, 4.2), constrained_layout=True) for row in values: axis.plot(x, row, lw=0.9, alpha=0.55) axis.fill_between( x, mean - std, mean + std, alpha=0.3, color="tab:blue", label="moyenne ± écart-type", ) axis.plot(x, mean, lw=2, color="tab:blue", label="moyenne") axis.set(title="Krylov + multigrille partagé, 8 EDP", xlabel="x", ylabel="u") axis.legend() axis.grid(alpha=0.25) output = Path(__file__).with_name("uq_batched_transport_1d_mg.png") figure.savefig(output, dpi=110) print(f"\nfigure : {output}") plt.show()