"""Quatre FORMES d'enrichissement pour la base d'un NO hybride, sur le cas dur. Pourquoi ce fichier existe --------------------------- ``advection_diffusion_1d_learned_basis.py`` mesure une base apprise de la forme ``phi_k(x) = m(x) T_k(x)``, ``m(x) = exp(2 tanh(N(x)))`` sur deux sources. Avec ``f = 1`` elle gagne un facteur 8 ; avec ``f = sin(3 pi x)`` le gain tombe à 2,6 et l'erreur cesse de diminuer près des interfaces de mailles. La raison est *structurelle* et non un défaut de réglage : un multiplicateur scalaire partagé module une AMPLITUDE. Cela représente très bien un changement d'échelle -- une couche limite -- et pas du tout un contenu fréquentiel, que seuls des degrés polynomiaux (ou d'autres atomes) apportent. On compare donc ici quatre façons d'enrichir la même base de Taylor, à **nombre de degrés de liberté identique** (c'est l'espace qu'on change, jamais sa dimension), sur la source oscillante seule -- le cas où la première forme plafonne. Les quatre formes ----------------- ===================== ==================================================== ``mult. partagé`` ``m(x) T_k(x)``, un seul champ pour tous les modes. La forme de référence, celle du premier exemple. ``mult. par mode`` ``m_k(x) T_k(x)``. Le réseau sort un multiplicateur PAR mode, donc il peut étirer les modes hauts sans toucher aux bas -- ce qu'un champ unique ne peut pas. ``additif`` ``T_k(x) + tanh(a_k(x)) / 2``. L'enrichissement n'est plus multiplicatif : il ajoute une forme libre. ⚠ La perturbation est BORNÉE et petite au départ, faute de quoi rien ne garantit que les ``p+1`` fonctions restent indépendantes -- et une base locale dégénérée rend le système singulier sans rien lever. ``Gabor`` ``exp(-((xi-c_k)/s_k)^2/2) cos(omega_k (xi-c_k) + phi_k)`` en coordonnée locale ``xi``, avec **fréquence, position, largeur et phase apprises**, une de chaque par mode. Ce n'est plus un enrichissement d'une base polynomiale : c'est une AUTRE base, et elle porte le contenu fréquentiel que les trois précédentes ne fabriquent pas. ===================== ==================================================== Ce que Gabor change, et ce qu'il coûte --------------------------------------- Quatre nombres par mode contre plusieurs centaines de poids : ``4 (p+1)`` paramètres au total, soit **12** ici, contre ~350 pour un MLP. C'est la différence entre apprendre un champ et apprendre *où regarder* -- une base d'atomes localisés dont on ajuste la place et l'échelle, exactement ce qu'une ondelette fait à la main. Mesuré, il atteint ``9.2e-03`` de meilleure approximation contre ``8.2e-03`` pour le meilleur MLP : 13 % d'écart pour trente fois moins de paramètres. ⚠ **Il lui faut une QUADRATURE PLUS FINE, et c'est la mesure qui l'a dit.** ``2 p + 2`` points intègrent exactement une base polynomiale -- c'est le défaut partout dans le dépôt -- mais pas une gaussienne fois un cosinus dont l'argument atteint ``4 pi`` sur une maille, soit deux périodes pour six points. Mesuré, à tout le reste égal : ``4.3e-02`` de solve à six points (donc *pire* que le schéma classique) contre ``1.7e-02`` à seize, et ``2.2e-02`` contre ``9.2e-03`` de meilleure approximation. ⚠⚠ Et le piège se referme au pire endroit : à l'INITIALISATION la quadrature grossière suffit (``omega_k = k pi``, au plus une période, erreur ``5.60e-01`` contre ``5.59e-01``). C'est l'apprentissage qui pousse les fréquences hors de sa portée -- l'optimiseur se dirige exactement là où l'assemblage cesse d'être juste, et optimise ensuite un problème qui n'est plus le bon. ⚠ **Tous les paramètres sont BORNÉS**, chacun pour sa raison. ``omega`` doit rester intégrable (ci-dessus) ; ``s`` ne doit pas tendre vers zéro, car un atome concentré au centre de la maille ne voit plus les faces -- or le DG *couple par les faces*, donc une enveloppe qui s'y annule détruit le schéma avant même le conditionnement. D'où ``s in (0.2, 0.8)`` : à ``s = 0.5``, un atome vaut encore ``0.61`` au bord de sa maille. ⚠ **Gabor ne vaut PAS le schéma classique à l'initialisation**, contrairement aux trois autres formes (où ``tanh(0) = 0`` laisse la base de Taylor inchangée). La colonne « base initiale » du tableau le montre, et c'est la raison pour laquelle elle y figure. Lancer : python learned_basis_variants_1d.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, trainable # ── Paramètres ─────────────────────────────────────────────────────────────── DIM = 1 EPS = 0.05 F_FREQ = 3.0 MU_LOW, MU_HIGH = 0.8, 1.2 N_CELLS, ORDER = 8, 2 # 24 ddl : la finesse où la forme multiplicative plafonne NB_BASIS = ORDER + 1 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" #: Points de Gauss par direction. ⚠ ``2 p + 2`` intègre EXACTEMENT une base #: polynomiale, et c'est pour cela que c'est le défaut partout dans le dépôt -- #: mais un atome de Gabor n'est pas un polynôme. Son argument peut atteindre #: ``4 pi`` sur une maille, soit deux périodes, que six points échantillonnent à #: trois points par période. D'où la variante à quadrature fine. QUAD_POINTS = 2 * ORDER + 2 QUAD_POINTS_FINE = 16 HIDDEN = [16, 16] SEED = 0 _MAPPING = Mapping(mappings=[InvertibleFunction(lambda x: x, lambda y: y)]) def 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]) # ── La famille d'EDP ───────────────────────────────────────────────────────── class LinearTransport(ScimbaPytree): """``b(x; mu) = mu_0 + mu_1 x``, porté par une FEUILLE (donc batchable). Args: mu: ``(2,)``. """ 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. Args: main_domain: Le segment. mu: ``(2,)``. data: ``(x, u_ref(x))``. """ 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=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 quatre formes d'enrichissement ─────────────────────────────────────── def make_mesh(n_cells: int, order: int, quad_points: int | None = None) -> Mesh: """Le maillage cartésien 1D. Args: n_cells: Nombre de mailles. order: Degré local (fixe la quadrature par défaut). quad_points: Points de Gauss par direction ; ``2 order + 2`` sinon. Returns: Le maillage. """ points = 2 * order + 2 if quad_points is None else quad_points return Mesh( dim=DIM, n_cells=[n_cells], ref_quad=UnitSquareTensorized(dim=DIM, order=points), 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 def classical_space(n_cells: int, order: int) -> VariablesDG: """L'espace DG habituel.""" mesh = make_mesh(n_cells, order) return VariablesDG( basis=AnalyticBasis( nb_basis=order + 1, out_dim=1, mesh=mesh, basis_type="scalar", local_basis=_taylor_of(order), ), nb_variables=1, ) class FieldNN(MLP): """Un champ appris, à ``out`` sorties. ⚠ Ses poids se déclarent eux-mêmes.""" def __init__(self, key, out: int = 1): super().__init__(in_size=DIM, out_size=out, hidden_sizes=HIDDEN, key=key) class GaborParams(ScimbaPytree): """Un PORTEUR de paramètres appris -- pas un champ : il ignore le point. ⚠ C'est le détournement volontaire de ``PatchwiseParametricBasis``, qui attend ``y -> u`` puis ``local_basis(u, y, i, mesh)``. Ici ``u`` ne dépend pas de ``y`` : ce sont les paramètres des atomes, et c'est ``local_basis`` qui s'en sert pour fabriquer la base. On réutilise ainsi toute la mécanique (pytree, dérivation, partition) sans écrire une classe de base de plus. ⚠ ``theta`` DOIT se déclarer ``trainable`` : la règle du dépôt est qu'une feuille est active si, et seulement si, le champ qui la porte l'a dit. Un MLP n'a rien à déclarer parce que ses ``Linear`` le font déjà ; ici, personne d'autre ne peut le faire. Args: n_modes: Nombre d'atomes par maille. """ theta = trainable(True) def __init__(self, n_modes: int): # (fréquence, position, largeur, PHASE) par mode, toutes à leur valeur # neutre : tanh(0) = 0 et sigmoid(0) = 1/2, donc omega_k = k pi, # c_k = 0, s = 0.5, et phi_k = k pi / 2. self.theta = jnp.zeros((4, n_modes)) def __call__(self, y: jnp.ndarray) -> jnp.ndarray: """Les paramètres, quel que soit le point.""" return self.theta def multiplier(u: jnp.ndarray) -> jnp.ndarray: """``exp(2 tanh(N))`` : positif, valant 1 en ``N = 0``, large d'un facteur 55.""" return jnp.exp(2.0 * jnp.tanh(u)) def shared_multiplier_basis(u, y, i, mesh): """``m(x) T_k(x)`` -- un seul champ pour tous les modes.""" return local_taylor_basis(y, i, mesh, order=ORDER, out_dim=1) * multiplier(u[0]) def per_mode_multiplier_basis(u, y, i, mesh): """``m_k(x) T_k(x)`` -- un multiplicateur par mode. ⚠ ``u`` a ``NB_BASIS`` composantes et la base ``(NB_BASIS, 1)`` : le produit se fait mode à mode, d'où le ``[:, None]``. Sans lui, la diffusion NumPy donnerait ``(NB_BASIS, NB_BASIS)`` sans rien lever. """ taylor = local_taylor_basis(y, i, mesh, order=ORDER, out_dim=1) return taylor * multiplier(u)[:, None] def additive_basis(u, y, i, mesh): """``T_k(x) + tanh(a_k(x)) / 2`` -- l'enrichissement n'est plus un facteur. La perturbation est bornée par ``1/2`` et nulle à l'initialisation : la base part de Taylor, et ne peut pas s'en écarter au point de dégénérer. """ taylor = local_taylor_basis(y, i, mesh, order=ORDER, out_dim=1) return taylor + 0.5 * jnp.tanh(u)[:, None] def gabor_basis(u, y, i, mesh): """Des atomes de Gabor à fréquence, position et largeur APPRISES. ``phi_k(xi) = exp(-((xi - c_k) / s_k)^2 / 2) cos(omega_k (xi - c_k))`` où ``xi`` est la coordonnée locale dans la maille, dans ``[-1/2, 1/2]``. ⚠ **La PHASE n'est pas un raffinement, c'est ce qui rend la base viable.** ``cos`` est PAIR ; sans phase et avec des centres initialisés à zéro, les ``p+1`` atomes d'une maille sont tous des fonctions paires de la coordonnée locale, et l'espace local ne contient AUCUNE fonction impaire -- il ne peut donc pas représenter une simple pente. Mesuré sans elle : une erreur de ``9.4e-01`` à l'initialisation (contre ``4e-02`` pour les autres formes) et ``0,88x`` le schéma classique après entraînement, c'est-à-dire pire que lui. ``phi_k = k pi / 2`` au départ fait alterner cosinus et sinus, exactement comme une base de Fourier locale. ⚠ Les autres paramètres sont bornés, chacun pour une raison différente : ``omega`` pour rester intégrable par la quadrature, ``c`` pour rester dans la maille, et ``s`` pour que l'atome ne s'annule pas sur les FACES -- le DG couple par les faces, et une enveloppe qui y meurt casse le schéma bien avant le conditionnement. Args: u: ``(4, NB_BASIS)`` -- les paramètres bruts rendus par le porteur. y: Le point physique. i: L'indice de maille. mesh: Le maillage. Returns: ``(NB_BASIS, 1)``. """ xi = (y - mesh.cell_centroid(i)) * jnp.asarray(mesh.n_cells, dtype=y.dtype) modes = jnp.arange(NB_BASIS, dtype=y.dtype) omega = jnp.pi * (modes + 2.0 * jnp.tanh(u[0])) # k pi a l'initialisation center = 0.5 * jnp.tanh(u[1]) # dans la maille width = 0.2 + 0.6 * jax.nn.sigmoid(u[2]) # 0.5 a l'initialisation phase = 0.5 * jnp.pi * modes + jnp.pi * jnp.tanh(u[3]) # cos, sin, cos, ... z = (xi[0] - center) / width atoms = jnp.exp(-0.5 * z**2) * jnp.cos(omega * (xi[0] - center) + phase) return atoms[:, None] def learned_space(key, kind: str) -> VariablesDG: """L'espace enrichi de la forme demandée. Args: key: L'état du générateur. kind: ``"shared"``, ``"per_mode"``, ``"additive"`` ou ``"gabor"``. Returns: L'espace. Raises: ValueError: Sur une forme inconnue. """ quad = QUAD_POINTS_FINE if kind.endswith("_hq") else QUAD_POINTS mesh = make_mesh(N_CELLS, ORDER, quad) kind = kind.removesuffix("_hq") if kind == "shared": carrier, local, local_coords = FieldNN(key, 1), shared_multiplier_basis, True elif kind == "per_mode": carrier = FieldNN(key, NB_BASIS) local, local_coords = per_mode_multiplier_basis, True elif kind == "additive": carrier = FieldNN(key, NB_BASIS) local, local_coords = additive_basis, True elif kind == "gabor": carrier, local, local_coords = GaborParams(NB_BASIS), gabor_basis, False else: raise ValueError(f"forme inconnue : {kind!r}") return VariablesDG( basis=PatchwiseParametricBasis( nb_basis=NB_BASIS, out_dim=1, mesh=mesh, patchwise_parametric_function=carrier, local_basis=local, basis_type="scalar", use_local_coords=local_coords, ), nb_variables=1, ) def make_operator(space: VariablesDG, order: int) -> DGSolverOperator: """L'opérateur : l'espace, le flux, la traduction du modèle.""" return DGSolverOperator( space, SIPGFlux(sigma=10.0 * order * (order + 1), h=None), weak_model_fn=weak_model_from, ) # ── 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 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. La matrice de la base s'obtient en passant la base canonique des ddl au DÉCODEUR. """ 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 def project_one(values): coefficients, *_ = jnp.linalg.lstsq(basis_matrix, values[:, 0], rcond=None) return basis_matrix @ coefficients return relative_l2(jax.vmap(project_one)(reference)[..., None], reference) 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:], } def run_variant(key, kind: str, data, domain, x_eval): """Entraîne une forme d'enrichissement et la mesure. Args: key: L'état du générateur. kind: La forme. data: Ce que :func:`prepare_family` a produit. domain: Le segment. x_eval: La grille. Returns: ``(key, résultats)``. """ key, sub = jax.random.split(key) operator = make_operator(learned_space(sub, kind), ORDER) space = PhysicNOApproximationSpace( dims={"x": DIM}, list_models=[operator], model_type="x" ) initial = operator.evaluate_batch(data["test_batch"], x_eval) 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] learned = trained.evaluate_batch(data["test_batch"], x_eval) return key, { "kind": kind, "n_params": int(space.ndof), "seconds": seconds, "initial": relative_l2(initial, data["test_ref"]), "solve": relative_l2(learned, data["test_ref"]), "best": best_approximation(trained, data["test_ref"], x_eval), "curve": learned, "history": np.asarray(projector.losses.losses_history["total"]).reshape(-1), } VARIANTS = ("shared", "per_mode", "additive", "gabor", "gabor_hq") LABELS = { "shared": "mult. partagé", "per_mode": "mult. par mode", "additive": "additif", "gabor": f"Gabor ({QUAD_POINTS} pts quad)", "gabor_hq": f"Gabor ({QUAD_POINTS_FINE} pts quad)", } def main() -> None: """Compare les quatre formes sur la source oscillante.""" 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) classical_op = make_operator(classical_space(N_CELLS, ORDER), ORDER) classical = classical_op.evaluate_batch(data["test_batch"], x_eval) e_classical = relative_l2(classical, data["test_ref"]) b_classical = best_approximation(classical_op, data["test_ref"], x_eval) print( f"\nDG classique ({N_CELLS * NB_BASIS} ddl) : solve {e_classical:.3e}, " f"meilleure approx. {b_classical:.3e}" ) results = [] for kind in VARIANTS: print(f"\n=== {LABELS[kind]} ===") key, result = run_variant(key, kind, data, domain, x_eval) results.append(result) print( f" {result['n_params']} paramètres, {N_EPOCHS} époques en " f"{result['seconds']:.1f} s solve {result['solve']:.3e}" ) print( f"\n{'forme':18s} {'params':>7s} {'init.':>11s} {'solve':>11s} " f"{'gain':>7s} {'best approx':>12s} {'gain':>7s}" ) print("-" * 78) print( f"{'DG classique':18s} {'—':>7s} {'—':>11s} {e_classical:11.3e} " f"{'—':>7s} {b_classical:12.3e} {'—':>7s}" ) for r in results: print( f"{LABELS[r['kind']]:18s} {r['n_params']:7d} {r['initial']:11.3e} " f"{r['solve']:11.3e} {e_classical / r['solve']:6.2f}x " f"{r['best']:12.3e} {b_classical / r['best']:6.2f}x" ) # ── La figure ──────────────────────────────────────────────────────── x = np.asarray(x_eval).ravel() test_ref = data["test_ref"] worst = int( np.argmax(np.asarray(jnp.linalg.norm(classical - test_ref, axis=(1, 2)))) ) fig, axes = plt.subplots(1, 3, figsize=(16, 4.4)) axes[0].plot(x, np.asarray(test_ref[worst, :, 0]), "k-", lw=2, label="référence") axes[0].plot( x, np.asarray(classical[worst, :, 0]), "--", color="0.5", label="DG classique" ) for r in results: axes[0].plot(x, np.asarray(r["curve"][worst, :, 0]), label=LABELS[r["kind"]]) axes[0].set_title(f"EDP de test la plus dure ({N_CELLS} mailles, p{ORDER})") axes[0].set_xlabel("x") axes[0].legend(fontsize=8) axes[1].semilogy( x, np.abs(np.asarray(classical[worst, :, 0] - test_ref[worst, :, 0])), "--", color="0.5", label="DG classique", ) for r in results: axes[1].semilogy( x, np.abs(np.asarray(r["curve"][worst, :, 0] - test_ref[worst, :, 0])), label=LABELS[r["kind"]], ) for edge in np.linspace(0.0, 1.0, N_CELLS + 1): axes[1].axvline(edge, color="0.9", 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=LABELS[r["kind"]]) axes[2].set_title(f"loss d'entraînement ({OPTIMIZER})") axes[2].set_xlabel("époque") axes[2].legend(fontsize=8) plt.tight_layout() plt.savefig("learned_basis_variants_1d.png", dpi=120) print("\nfigure : learned_basis_variants_1d.png") if __name__ == "__main__": main()