""" Example : Utilisation de Mesh 2D avec mappings simples et apprenables Ce script démontre les capacités du maillage 2D : 1. Création d'un maillage 2D avec des mappings fixes (analytiques) 2. Création d'un maillage 2D avec des mappings apprenables (paramètres à apprendre) 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 import Mesh 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 2D # ============================================================================= class LearnableQuadratic2D(ScimbaPytree): """Mapping quadratique 2D f(x, y) = (alpha * x^2, beta * y^2). ``alpha`` et ``beta`` sont declares ``trainable(True)`` : le defaut est GELE, donc un champ n'est optimise que parce que quelqu'un l'a dit (voir ``mesh_1d.py``). """ alpha = trainable(True) beta = trainable(True) def __init__(self, key=None, alpha_init=1.0, beta_init=1.0): self.alpha = jnp.array(alpha_init) self.beta = jnp.array(beta_init) def __call__(self, x): return jnp.array([self.alpha * x[0] ** 2, self.beta * x[1] ** 2]) def inverse(self, y): return jnp.array([jnp.sqrt(y[0] / self.alpha), jnp.sqrt(y[1] / self.beta)]) def backward(self, y): return self.inverse(y) # ============================================================================= # PARTIE 2 : Fonction utilitaire de visualisation du maillage 2D # ============================================================================= def plot_mesh_2d(mesh, ax, title="Maillage 2D", xlim=None, ylim=None): """Affiche les arêtes et les points de quadrature d'un maillage 2D. Args: mesh: Instance de Mesh avec dim=2 ax: Axes matplotlib sur lequel dessiner title: Titre du graphique xlim: Tuple (xmin, xmax) pour zoomer sur une région (optionnel) ylim: Tuple (ymin, ymax) pour zoomer sur une région (optionnel) """ n_cells = [int(n) for n in mesh.n_cells] # Construire les nœuds physiques à partir de la grille de référence [0,1]^2 grids = [jnp.linspace(0, 1, n + 1) for n in n_cells] mesh_grids = jnp.meshgrid(*grids, indexing="ij") nodes_ref = jnp.stack([g.ravel() for g in mesh_grids], axis=-1) nodes_phys = mesh.mapping.local_mapping(nodes_ref) nodes = nodes_phys.reshape(*(n + 1 for n in n_cells), 2) # Arêtes verticales (x=const) for i in range(n_cells[0] + 1): for j in range(n_cells[1]): p1 = nodes[i, j] p2 = nodes[i, j + 1] ax.plot([p1[0], p2[0]], [p1[1], p2[1]], "b-", lw=0.8, alpha=0.6) # Arêtes horizontales (y=const) for i in range(n_cells[0]): for j in range(n_cells[1] + 1): p1 = nodes[i, j] p2 = nodes[i + 1, j] ax.plot([p1[0], p2[0]], [p1[1], p2[1]], "r-", lw=0.8, alpha=0.6) # Points de quadrature dans chaque cellule x_all = mesh.evaluate_mesh_points() # (n_cells_total, n_gauss, 2) pts = x_all.reshape(-1, 2) ax.scatter(pts[:, 0], pts[:, 1], s=6, c="black", zorder=5, alpha=0.7) ax.set_title(title, fontsize=11) ax.set_aspect("equal") ax.set_xlabel("x") ax.set_ylabel("y") ax.grid(True, alpha=0.2) if xlim is not None: ax.set_xlim(xlim) if ylim is not None: ax.set_ylim(ylim) # ============================================================================= # PARTIE 3 : Test du maillage avec mappings analytiques fixes # ============================================================================= # Mappings 2D : # - identity : f(x, y) = (x, y), inverse : (x, y) # - x^2 composante par composante : f(x, y) = (x^2, y^2), inverse : (sqrt(x), sqrt(y)) mapping_id = InvertibleFunction(lambda x: x, lambda y: y) mapping_x2_2d = InvertibleFunction(lambda x: x**2, lambda y: jnp.sqrt(y)) print("\n\n@@@@@@@@@@@@@@@@ Test 1 : Maillage 2D avec mapping x^2 fixe @@@@@@@@@@@@@@@") # Créer un maillage 2D avec (10, 10) cellules et 4 points de quadrature par direction # Le mapping composé est : identity puis x^2 composante par composante m = Mesh( dim=2, n_cells=(10, 10), ref_quad=UnitSquareTensorized(dim=2, order=4), mapping=Mapping(mappings=[mapping_id, mapping_x2_2d]), ) # Compiler (JIT) et évaluer le maillage 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 2D avec mapping apprenable @@@@@@@@@@@@@@@" ) # Créer un modèle apprenable avec alpha=2.0, beta=1.5 model = LearnableQuadratic2D(alpha_init=2.0, beta_init=1.5) m2 = Mesh( dim=2, n_cells=(20, 20), ref_quad=UnitSquareTensorized(dim=2, order=4), mapping=Mapping(mappings=[mapping_id, model]), ) 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 (20x20 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 4 : Apprentissage d'un mapping apprenable 2D # ============================================================================= print( "\n\n@@@@@@@@@@@@@@@@ Test 3 : Apprentissage des paramètres alpha et beta @@@@@@@@@@@@@@@" ) print( "Objectif : apprendre un mapping (alpha*x^2, beta*y^2) " "qui approche le mapping cible (alpha=3.0, beta=2.5)" ) # Modèle initial (à apprendre) model_test = LearnableQuadratic2D(alpha_init=2.0, beta_init=1.5) print(f"Alpha initial: {model_test.alpha}, Beta initial: {model_test.beta}") # Modèle cible (à approcher) model_target = LearnableQuadratic2D(alpha_init=3.0, beta_init=2.5) # Créer les maillages correspondants mapping_learnable = Mapping(mappings=[mapping_id, model_test]) mapping_target = Mapping(mappings=[mapping_id, model_target]) mesh_learnable = Mesh( dim=2, n_cells=(10, 10), ref_quad=UnitSquareTensorized(dim=2, order=4), mapping=mapping_learnable, ) mesh_target = Mesh( dim=2, n_cells=(10, 10), ref_quad=UnitSquareTensorized(dim=2, order=4), mapping=mapping_target, ) # Points du maillage cible (données à approcher) y_target_mesh = mesh_target.evaluate_mesh_points() print(f"Points du maillage target: {y_target_mesh.shape}") def loss_fn_mesh(model_param): """Loss basée sur les points du maillage 2D. On crée un maillage temporaire avec le modèle actuel, on évalue ses points, et on compare avec les points cibles (erreur quadratique moyenne). """ mapping_current = Mapping(mappings=[mapping_id, model_param]) mesh_current = Mesh( dim=2, n_cells=(10, 10), ref_quad=UnitSquareTensorized(dim=2, 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 des paramètres et de la loss models = [model_test] alphas = [float(model_test.alpha)] betas = [float(model_test.beta)] losses = [float(loss_fn_mesh(model_test))] print(f"Loss initiale : {losses[-1]:.6f}") # Boucle d'apprentissage : 6 étapes de descente de gradient current_model = model_test for step in range(30): # Calculer le gradient de la loss par rapport aux paramètres du modèle grads = grad_fn(current_model) # Mettre à jour alpha et beta simultanément current_model = jax.tree_util.tree_map( lambda p, g: p - learning_rate * g, current_model, grads ) # Stocker et évaluer models.append(current_model) alphas.append(float(current_model.alpha)) betas.append(float(current_model.beta)) losses.append(float(loss_fn_mesh(current_model))) print( f"Step {step + 1}: α = {current_model.alpha:.4f}, β = {current_model.beta:.4f}, " f"loss = {losses[-1]:.8f}, " f"grad_α = {grads.alpha:.6f}, grad_β = {grads.beta:.6f}" ) print(f"\nÉvolution de alpha: {[f'{a:.4f}' for a in alphas]}") print(f"Évolution de beta: {[f'{b:.4f}' for b in betas]}") print("Alpha cible: 3.0000, Beta cible: 2.5000") # ============================================================================= # PARTIE 5 : Visualisation # ============================================================================= print("\n\n@@@@@@@@@@@@@@@@ Visualisation @@@@@@@@@@@@@@@") fig = plt.figure(figsize=(18, 10)) # Zone de zoom commune : coin bas-gauche où la compression quadratique est la plus visible zoom_xlim = (0.0, 0.2) zoom_ylim = (0.0, 0.2) # ---- Ligne 1 : Maillages initial, final et cible ---- ax1 = fig.add_subplot(2, 3, 1) mesh_initial = Mesh( dim=2, n_cells=(10, 10), ref_quad=UnitSquareTensorized(dim=2, order=4), mapping=Mapping(mappings=[mapping_id, models[0]]), ) plot_mesh_2d( mesh_initial, ax1, title=f"Initial : α={alphas[0]:.2f}, β={betas[0]:.2f}", xlim=zoom_xlim, ylim=zoom_ylim, ) ax2 = fig.add_subplot(2, 3, 2) mesh_final = Mesh( dim=2, n_cells=(10, 10), ref_quad=UnitSquareTensorized(dim=2, order=4), mapping=Mapping(mappings=[mapping_id, models[-1]]), ) plot_mesh_2d( mesh_final, ax2, title=f"Final : α={alphas[-1]:.4f}, β={betas[-1]:.4f}", xlim=zoom_xlim, ylim=zoom_ylim, ) ax3 = fig.add_subplot(2, 3, 3) plot_mesh_2d( mesh_target, ax3, title="Cible : α=3.00, β=2.50", xlim=zoom_xlim, ylim=zoom_ylim ) # ---- Ligne 2 : Évolution des paramètres et de la loss ---- ax4 = fig.add_subplot(2, 3, 4) steps = range(len(alphas)) ax4.plot(steps, alphas, "bo-", linewidth=2, markersize=8, label="α") ax4.axhline( y=3.0, color="b", linestyle="--", linewidth=1.5, label="Cible α=3.0", alpha=0.5 ) ax4.plot(steps, betas, "rs-", linewidth=2, markersize=8, label="β") ax4.axhline( y=2.5, color="r", linestyle="--", linewidth=1.5, label="Cible β=2.5", alpha=0.5 ) ax4.set_xlabel("Step", fontsize=12) ax4.set_ylabel("Valeur du paramètre", fontsize=12) ax4.set_title("Évolution de α et β", fontsize=13) ax4.legend(fontsize=9) ax4.grid(True, alpha=0.3) ax5 = fig.add_subplot(2, 3, 5) ax5.semilogy(steps, losses, "go-", linewidth=2, markersize=8) ax5.set_xlabel("Step", fontsize=12) ax5.set_ylabel("Loss (MSE)", fontsize=12) ax5.set_title("Convergence de la loss", fontsize=13) ax5.grid(True, alpha=0.3) # ---- Maillage de référence avec mapping x^2 fixe ---- ax6 = fig.add_subplot(2, 3, 6) plot_mesh_2d(m, ax6, title="Mapping x² fixe (10×10 cellules)") plt.tight_layout() plt.show()