""" Example : Utilisation de Mesh1D avec mappings simples et apprenables Ce script démontre les capacités du maillage 1D : 1. Création d'un maillage avec des mappings fixes (analytiques) 2. Création d'un maillage avec des mappings apprenables (réseaux de neurones) 3. Apprentissage des paramètres du mapping pour approcher un mapping cible 4. Visualisation de l'évolution du maillage au cours de l'apprentissage """ import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt # from scimba_jax.linear_approximation.meshes.mesh_1d import Mesh1D from scimba_jax.linear_approximation.meshes.mesh import Mesh as Mesh1D from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.mapping.mapping import InvertibleFunction, Mapping from scimba_jax.utils.scimba_pytree import ScimbaPytree, trainable # ============================================================================= # PARTIE 1 : Définition d'un modèle apprenable # ============================================================================= class LearnableQuadratic(ScimbaPytree): """Mapping quadratique f(x) = alpha * x^2 avec alpha apprenable. ``alpha`` est declare ``trainable(True)`` : le defaut est GELE, donc un champ n'est optimise que parce que quelqu'un l'a dit. Ici la descente se fait par un ``jax.grad`` direct sur le modele, qui deriverait de toute facon par rapport a tout enfant du pytree ; la declaration dit l'intention a l'endroit ou on la lit, et c'est elle qui compte des que ce mapping passe par un projecteur. """ alpha = trainable(True) def __init__(self, key=None, init_value=1.0): self.alpha = jnp.array(init_value) def __call__(self, x): return self.alpha * x**2 def inverse(self, y): return jnp.sqrt(y / self.alpha) # ============================================================================= # PARTIE 2 : Test du maillage avec mappings analytiques fixes # ============================================================================= # Créer deux mappings simples : # - identity : f(x) = x, inverse f^{-1}(y) = y (pas de déformation) # - x^2 : f(x) = x^2, inverse f^{-1}(y) = sqrt(y) (compression vers zéro) mapping_id = InvertibleFunction(lambda x: x, lambda y: y) mapping_x2 = InvertibleFunction(lambda x: x**2, lambda y: jnp.sqrt(y)) print("\n\n@@@@@@@@@@@@@@@@ Test 1 : Maillage avec mapping x^2 fixe @@@@@@@@@@@@@@@") # Créer un maillage avec 20 cellules et 4 points de quadrature par cellule (order=4) # Le mapping compose identity (sur x) et x^2 (sur la deformation) m = Mesh1D( dim=1, n_cells=(20,), ref_quad=UnitSquareTensorized(dim=1, order=4), mapping=Mapping(mappings=[mapping_id, mapping_x2]), ) # Compiler (JIT) et évaluer le maillage pour les performances key = jax.random.PRNGKey(0) evaluate_jit = jax.jit(m.evaluate_mesh_points) # Première exécution (compilation) start = time.perf_counter() y = evaluate_jit() y.block_until_ready() end = time.perf_counter() print(f"Première évaluation (compilation) : {end - start:.6f} s, shape: {y.shape}") # Deuxième exécution (déjà compilée, plus rapide) start = time.perf_counter() y = evaluate_jit() y.block_until_ready() end = time.perf_counter() print(f"Deuxième évaluation (compilée) : {end - start:.6f} s, shape: {y.shape}") w_all, x_all = m.evaluate_mesh_weights_points() print("w_all shape:", w_all.shape) print("x_all shape:", x_all.shape) print("\n\n@@@@@@@@@@@@@@@@ Test 2 : Maillage avec mapping apprenable @@@@@@@@@@@@@@@") # Créer un modèle apprenable avec alpha initial = 2.0 model = LearnableQuadratic(init_value=2.0) # Créer un maillage utilisant ce modèle apprenable m2 = Mesh1D( dim=1, n_cells=(200,), ref_quad=UnitSquareTensorized(dim=1, order=4), mapping=Mapping(mappings=[mapping_id, model]), ) # Compiler et évaluer key = jax.random.PRNGKey(0) evaluate_jit = jax.jit(m2.evaluate_mesh_points) start = time.perf_counter() y = evaluate_jit() y.block_until_ready() end = time.perf_counter() print(f"Première évaluation (200 cellules) : {end - start:.6f} s, shape: {y.shape}") start = time.perf_counter() y = evaluate_jit() y.block_until_ready() end = time.perf_counter() print(f"Deuxième évaluation (compilée) : {end - start:.6f} s, shape: {y.shape}") # ============================================================================= # PARTIE 3 : Apprentissage d'un mapping apprenable # ============================================================================= print("\n\n@@@@@@@@@@@@@@@@ Test 3 : Apprentissage du paramètre alpha @@@@@@@@@@@@@@@") print( "Objectif : apprendre un mapping alpha*x^2 qui approche le mapping cible (alpha=3.0)" ) # Modèle initial avec alpha = 2.0 (celui qu'on va apprendre) model_test = LearnableQuadratic(init_value=2.0) print(f"Alpha initial: {model_test.alpha}") # Modèle cible avec alpha = 3.0 (celui qu'on veut approcher) model_target = LearnableQuadratic(init_value=3.0) # Créer les maillages correspondants mapping_learnable = Mapping(mappings=[mapping_id, model_test]) mapping_target = Mapping(mappings=[mapping_id, model_target]) mesh_learnable = Mesh1D( dim=1, n_cells=(20,), ref_quad=UnitSquareTensorized(dim=1, order=4), mapping=mapping_learnable, ) mesh_target = Mesh1D( dim=1, n_cells=(20,), ref_quad=UnitSquareTensorized(dim=1, order=4), mapping=mapping_target, ) # Évaluer les points du maillage target (données cibles à approcher) y_target_mesh = mesh_target.evaluate_mesh_points() print(f"Points du maillage target: {y_target_mesh.shape}") # Fonction de loss : calculer l'erreur quadratique moyenne entre points approchés et cibles def loss_fn_mesh(model_param): """Loss basée sur les points du maillage. On crée un maillage temporaire avec le modèle actuel, on évalue ses points, et on compare avec les points cibles. """ mapping_current = Mapping(mappings=[mapping_id, model_param]) mesh_current = Mesh1D( dim=1, n_cells=(20,), ref_quad=UnitSquareTensorized(dim=1, order=4), mapping=mapping_current, ) y_pred = mesh_current.evaluate_mesh_points() return jnp.mean((y_pred - y_target_mesh) ** 2) # Créer la fonction de gradient grad_fn = jax.grad(loss_fn_mesh) learning_rate = 0.5 # Stocker l'évolution du modèle, des paramètres et des losses models = [model_test] alphas = [float(model_test.alpha)] losses = [float(loss_fn_mesh(model_test))] print(f"Loss initiale : {losses[-1]:.6f}") # Boucle d'apprentissage : 4 étapes de descente de gradient current_model = model_test for step in range(4): # Calculer le gradient de la loss par rapport aux paramètres du modèle grads = grad_fn(current_model) # Mettre à jour le paramètre alpha : alpha_new = alpha_old - lr * grad_alpha current_model = jax.tree_util.tree_map( lambda p, g: p - learning_rate * g, current_model, grads ) # Stocker le modèle et évaluer models.append(current_model) alphas.append(float(current_model.alpha)) losses.append(float(loss_fn_mesh(current_model))) print( f"Step {step + 1}: α = {current_model.alpha:.4f}, loss = {losses[-1]:.6f}, grad = {grads.alpha:.6f}" ) print(f"\nÉvolution de alpha: {[f'{a:.4f}' for a in alphas]}") print("Alpha cible: 3.0000") # ============================================================================= # PARTIE 4 : Visualisation de l'apprentissage # ============================================================================= print("\n\n@@@@@@@@@@@@@@@@ Visualisation de l'apprentissage @@@@@@@@@@@@@@@") # Créer une figure avec deux sous-graphiques : maillage et convergence fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(15, 6)) # ==== Graphique 1: Points du maillage au cours de l'apprentissage ==== # Évaluer les points de maillage pour chaque étape d'apprentissage mesh_points = [] for model_i in models: mapping_i = Mapping(mappings=[mapping_id, model_i]) mesh_i = Mesh1D( dim=1, n_cells=(20,), ref_quad=UnitSquareTensorized(dim=1, order=4), mapping=mapping_i, ) # Applatir les points pour une visualisation plus facile points_i = mesh_i.evaluate_mesh_points().flatten() mesh_points.append(points_i) # Données cibles y_target_flat = y_target_mesh.flatten() # Trier les points par valeur pour une meilleure visualisation (plutôt que par index cellule) sort_idx_target = jnp.argsort(mesh_target.evaluate_mesh_points().flatten()) x_ref = jnp.linspace(0, 1, len(sort_idx_target)) # Tracer la solution cible ax1.plot( x_ref, y_target_flat[sort_idx_target], "g-", linewidth=3, label="Target (α=3.0)", marker="o", markersize=3, ) # Tracer l'évolution du maillage à chaque step colors = ["b", "orange", "purple", "brown", "r"] for i, (points_i, alpha_i) in enumerate(zip(mesh_points, alphas)): sort_idx_i = jnp.argsort(points_i) # Pointillé pour les étapes intermédiaires, trait plein pour la dernière linestyle = "--" if i < len(models) - 1 else "-" linewidth = 1.5 if i < len(models) - 1 else 2.5 label = f"Step {i}: α={alpha_i:.4f}" if i > 0 else f"Initial: α={alpha_i:.4f}" ax1.plot( x_ref, points_i[sort_idx_i], linestyle=linestyle, color=colors[i], linewidth=linewidth, label=label, marker="x", markersize=4, ) ax1.set_xlabel("Index du point (normalisé)", fontsize=12) ax1.set_ylabel("Valeur après mapping", fontsize=12) ax1.set_title("Points du maillage après mapping: Apprentissage de α", fontsize=14) ax1.legend(fontsize=9) ax1.grid(True, alpha=0.3) # ==== Graphique 2: Évolution de alpha et de la loss ==== # Axe gauche : α (paramètre du modèle) ax2.plot(range(len(alphas)), alphas, "bo-", linewidth=2, markersize=8, label="α") # Ligne cible pour alpha = 3.0 ax2.axhline(y=3.0, color="g", linestyle="--", linewidth=2, label="Target α=3.0") ax2.set_xlabel("Step", fontsize=12) ax2.set_ylabel("α", fontsize=12, color="b") ax2.tick_params(axis="y", labelcolor="b") ax2.set_title("Évolution de α et loss", fontsize=14) ax2.grid(True, alpha=0.3) # Axe droit : loss (erreur quadratique moyenne) ax2_twin = ax2.twinx() ax2_twin.plot( range(len(losses)), losses, "rs-", linewidth=2, markersize=8, label="Loss" ) ax2_twin.set_ylabel("Loss (MSE)", fontsize=12, color="r") ax2_twin.tick_params(axis="y", labelcolor="r") # Combiner les légendes des deux axes lines1, labels1 = ax2.get_legend_handles_labels() lines2, labels2 = ax2_twin.get_legend_handles_labels() ax2.legend(lines1 + lines2, labels1 + labels2, fontsize=9, loc="upper right") plt.tight_layout() plt.show()