"""Un multigrille DEUX NIVEAUX monté sur l'opérateur RÉEL d'une base apprise. La question posée ----------------- Les transferts rapides du dépôt supposent une base **nodale** : ``CellwiseTransfer`` prolonge par interpolation aux nœuds de la maille fille, ce qui ne coïncide avec la projection L2 que dans ce cas -- il refuse d'ailleurs explicitement une base modale plutôt que de l'approcher. Une base **apprise** n'est ni nodale ni modale connue : aucun transfert rapide ne s'y applique, et c'est pour cela que le préconditionneur de ``advection_diffusion_2d_learned_basis_mg.py`` reste monté sur une hiérarchie en base classique -- donc sur un opérateur qui n'est pas celui qu'on résout. Ce fichier vérifie que l'autre voie est praticable : :class:`~...transfer.modal.CellwiseL2Transfer`, la projection L2 **locale** (bloc-diagonale, un ``M_K^-1`` par maille), qui ne suppose rien de la base. Si elle tient, un multigrille peut être monté sur l'opérateur réel. Trois vérifications, dans cet ordre ------------------------------------ 1. **Exactitude sur base polynomiale.** Quand ``V_H`` est inclus dans ``V_h``, la projection doit reproduire la fonction grossière au chiffre près. C'est le test qui distingue « le transfert est correct » de « le transfert a l'air de marcher », et il se fait sans multigrille. 2. **Propriété variationnelle.** ``restrict_dual`` doit être la transposée EXACTE de ``prolong`` -- c'est ce qui garantit ``A_H = P^T A_h P``, donc que la correction grossière est une projection A-orthogonale. Elle est dérivée par ``jax.linear_transpose`` et non écrite à la main, donc on mesure ici que la dérivation fait ce qu'on croit. 3. **Le cycle contracte**, sur base APPRISE et sur base classique, en mesurant le facteur de réduction de l'erreur par cycle sur ``A e = 0``. ⚠ Ce que ce fichier ne fait PAS : entraîner. Le réseau est tiré au hasard et laissé tel quel -- ce qui suffit, et est même plus honnête : une base non entraînée est une base quelconque, donc le cas le plus défavorable pour un transfert qui ne doit rien supposer. ⚠ **Les matrices du transfert dépendent de la base**, donc d'un paramètre appris. Elles sont construites au point où la base se trouve : une époque plus tard, le transfert est périmé et l'objet doit être reconstruit. C'est le pendant, côté transferts, de ``MG.relinearise``. ⚠ **Et le préconditionneur est coupé du gradient.** Il est bâti sur ``jax.lax.stop_gradient`` de l'espace : un préconditionneur est un accélérateur, pas une partie du modèle, et laisser l'optimiseur voir à travers lui reviendrait à optimiser la façon de résoudre plutôt que ce qu'on résout. Lancer : python mg_learned_basis_2levels_1d.py """ from __future__ import annotations import jax import jax.numpy as jnp import numpy as np 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.nonlinear_approximation.networks.mlp import MLP 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, ) DIM = 1 EPS = 0.05 ORDER = 2 NB_BASIS = ORDER + 1 N_FINE, N_COARSE = 16, 8 N_CYCLES = 12 SEED = 0 _MAPPING = Mapping(mappings=[InvertibleFunction(lambda x: x, lambda y: y)]) _FLUX = SIPGFlux(sigma=10.0 * ORDER * (ORDER + 1), h=None) def make_mesh(n_cells: 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, ) def _taylor(y, i, mesh): """La base de Taylor, un objet stable (pas une lambda par appel).""" return local_taylor_basis(y, i, mesh, order=ORDER, out_dim=1) def classical_space(n_cells: int) -> VariablesDG: """L'espace DG polynomial.""" return VariablesDG( basis=AnalyticBasis( nb_basis=NB_BASIS, out_dim=1, mesh=make_mesh(n_cells), basis_type="scalar", local_basis=_taylor, ), 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 _enriched(u, y, i, mesh): """``m(x) T_k(x)`` avec ``m = exp(2 tanh(N))``.""" return local_taylor_basis(y, i, mesh, order=ORDER, out_dim=1) * jnp.exp( 2.0 * jnp.tanh(u[0]) ) def learned_space(network, n_cells: int) -> VariablesDG: """Le MÊME champ appris, posé sur un maillage donné. ⚠ C'est ce qui rend une hiérarchie possible sur une base apprise : ``m(x)`` est un champ sur le DOMAINE, pas sur un maillage. Les deux niveaux portent donc le même réseau -- un seul jeu de poids, deux discrétisations -- et rien n'a besoin d'être « transféré » du réseau fin au réseau grossier. """ return VariablesDG( basis=PatchwiseParametricBasis( nb_basis=NB_BASIS, out_dim=1, mesh=make_mesh(n_cells), patchwise_parametric_function=network, local_basis=_enriched, basis_type="scalar", use_local_coords=True, ), nb_variables=1, ) def make_scheme(space: VariablesDG) -> EllipticDGscheme: """Le schéma DG d'un niveau : advection-diffusion, Dirichlet homogène.""" form = EllipticWeakForm( dim=DIM, A=lambda x: EPS * jnp.eye(DIM), b=lambda x: jnp.array([1.0]), c=lambda x: jnp.zeros(()), f=lambda x: jnp.ones(()), ) model = AbstractPhysicalWeakModel.from_weak_form( form, dirichlet=lambda x: jnp.zeros(1) ) return EllipticDGscheme(model, space, _FLUX) # ── 1. Le transfert est-il exact quand il doit l'être ? ────────────────────── def check_exactness(coarse: VariablesDG, fine: VariablesDG, transfer, label: str): """Prolonger puis évaluer doit rendre la MÊME fonction, si ``V_H`` est dans ``V_h``. Args: coarse: L'espace grossier. fine: L'espace fin. transfer: Le transfert à vérifier. label: Le nom du cas, pour l'affichage. Returns: L'écart maximal entre les deux évaluations. """ key = jax.random.PRNGKey(SEED + 1) dofs_coarse = jax.random.normal(key, coarse.dofsl.shape) dofs_fine = transfer.prolong(dofs_coarse) # ⚠ Loin des interfaces : l'expansion DG est DISCONTINUE aux faces, et les # deux niveaux n'y ont pas les mêmes discontinuités. Comparer là mesurerait # le saut, pas le transfert. points = jnp.linspace(0.012, 0.988, 401)[:, None] read_coarse = jax.vmap( lambda p: VariablesDG._classical_local_evaluate_pure(coarse, dofs_coarse, p) )(points) read_fine = jax.vmap( lambda p: VariablesDG._classical_local_evaluate_pure(fine, dofs_fine, p) )(points) gap = float(jnp.max(jnp.abs(read_fine - read_coarse))) print(f" {label:34s} écart max {gap:.2e}") return gap def check_transpose(transfer, coarse: VariablesDG, fine: VariablesDG) -> float: """``
== vs ':34s} écart relatif {gap:.2e}")
return gap
# ── 3. Le cycle contracte-t-il ? ─────────────────────────────────────────────
def contraction_rate(space_coarse, space_fine, label: str) -> float:
"""Le facteur de réduction de l'erreur par cycle, sur ``A e = 0``.
⚠ Second membre NUL et itéré de départ aléatoire : la solution est alors
exactement zéro, donc l'itéré EST l'erreur et le taux se lit sans référence.
C'est la façon standard de mesurer un cycle, et la seule qui ne confonde pas
la contraction avec la qualité du second membre.
Args:
space_coarse: L'espace du niveau grossier.
space_fine: L'espace du niveau fin.
label: Le nom du cas.
Returns:
Le facteur de contraction asymptotique.
"""
scheme_coarse = make_scheme(space_coarse)
scheme_fine = make_scheme(space_fine)
transfer = CellwiseL2Transfer(
space_coarse, space_fine, structured_children_nd([N_COARSE])
)
hierarchy = Hierarchy(
[level_from_scheme(scheme_coarse), level_from_scheme(scheme_fine)],
[transfer],
)
mg = MG(hierarchy, BlockJacobiSmoother(omega=0.8), nu_pre=2, nu_post=2)
key = jax.random.PRNGKey(SEED + 3)
error = jax.random.normal(key, (hierarchy.finest.n_dofs,))
zero = jnp.zeros_like(error)
norms = [float(jnp.linalg.norm(error))]
for _ in range(N_CYCLES):
error = mg.cycle(error, zero)
norms.append(float(jnp.linalg.norm(error)))
rates = [norms[k + 1] / norms[k] for k in range(len(norms) - 1) if norms[k] > 0]
asymptotic = float(np.exp(np.mean(np.log(np.asarray(rates[-5:])))))
print(
f" {label:34s} taux {asymptotic:.3f} "
f"({norms[0]:.1e} -> {norms[-1]:.1e} en {N_CYCLES} cycles)"
)
return asymptotic
def main() -> None:
"""Les trois vérifications, sur base classique puis sur base apprise."""
children = structured_children_nd([N_COARSE])
key = jax.random.PRNGKey(SEED)
network = BasisNN(key)
coarse_classical = classical_space(N_COARSE)
fine_classical = classical_space(N_FINE)
# ⚠ Le préconditionneur est bâti sur une COPIE COUPÉE DU GRADIENT : un
# cycle est un accélérateur, pas une partie du modèle. Sans cela,
# l'optimiseur verrait à travers le transfert et le lisseur, et
# optimiserait la façon de résoudre au lieu de ce qu'on résout.
frozen_network = jax.lax.stop_gradient(network)
coarse_learned = learned_space(frozen_network, N_COARSE)
fine_learned = learned_space(frozen_network, N_FINE)
print("1. Exactitude du transfert (V_H inclus dans V_h ?)")
transfer_classical = CellwiseL2Transfer(coarse_classical, fine_classical, children)
gap_classical = check_exactness(
coarse_classical, fine_classical, transfer_classical, "base polynomiale"
)
transfer_learned = CellwiseL2Transfer(coarse_learned, fine_learned, children)
gap_learned = check_exactness(
coarse_learned, fine_learned, transfer_learned, "base apprise"
)
print("\n2. Propriété variationnelle (restrict_dual = transpose(prolong))")
check_transpose(transfer_classical, coarse_classical, fine_classical)
check_transpose(transfer_learned, coarse_learned, fine_learned)
print(f"\n3. Contraction d'un cycle V(2,2), {N_COARSE} -> {N_FINE} mailles")
rate_classical = contraction_rate(
coarse_classical, fine_classical, "base polynomiale"
)
rate_learned = contraction_rate(coarse_learned, fine_learned, "base apprise")
print("\n" + "=" * 64)
verdict = "OUI" if max(rate_classical, rate_learned) < 1.0 else "NON"
print(f"un MG sur l'opérateur réel d'une base apprise : {verdict}")
print(
f" exactitude polynomiale {gap_classical:.1e} "
f"(la base apprise ne l'attend pas : {gap_learned:.1e})"
)
print(
f" contraction : {rate_classical:.3f} en polynomial, "
f"{rate_learned:.3f} en appris"
)
if __name__ == "__main__":
main()