"""Visualisation de maillages 2D avec différents mappings géométriques.""" import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np 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 physical_dim = 2 quad_order = 2 n_cells = 12 def make_mesh(mapping_fn, mapping_inv): return Mesh( dim=physical_dim, n_cells=(n_cells, n_cells), ref_quad=UnitSquareTensorized(dim=physical_dim, order=quad_order), mapping=Mapping(mappings=[InvertibleFunction(mapping_fn, mapping_inv)]), ) # ── Mappings ────────────────────────────────────────────────────────────────── # 1. Carré de référence def id_fwd(x): return x def id_inv(y): return y # 2. Losange — rotation 45° : coins à (±1,0) et (0,±1) # X = x+y-1, Y = y-x | inverse : x=(X-Y+1)/2, y=(X+Y+1)/2 def diamond_fwd(x): return jnp.array([x[0] + x[1] - 1.0, x[1] - x[0]]) def diamond_inv(y): return jnp.array([(y[0] - y[1] + 1.0) / 2.0, (y[0] + y[1] + 1.0) / 2.0]) # 3. Trapèze — largeur 1 en bas, 0.5 centrée en haut def trapeze_fwd(x): return jnp.array([0.25 * x[1] + x[0] * (1.0 - 0.5 * x[1]), x[1]]) def trapeze_inv(y): return jnp.array([(y[0] - 0.25 * y[1]) / (1.0 - 0.5 * y[1]), y[1]]) # 4. Chebyshev — clustering fort vers les 4 bords def cheby_fwd(x): return jnp.array( [0.5 * (1.0 - jnp.cos(jnp.pi * x[0])), 0.5 * (1.0 - jnp.cos(jnp.pi * x[1]))] ) def cheby_inv(y): return jnp.array( [jnp.arccos(1.0 - 2.0 * y[0]) / jnp.pi, jnp.arccos(1.0 - 2.0 * y[1]) / jnp.pi] ) # 5. Couche limite — clustering exponentiel vers y=0 K_BL = 3.0 def bl_fwd(x): return jnp.array([x[0], (jnp.exp(K_BL * x[1]) - 1.0) / (jnp.exp(K_BL) - 1.0)]) def bl_inv(y): return jnp.array([y[0], jnp.log(1.0 + y[1] * (jnp.exp(K_BL) - 1.0)) / K_BL]) # ── Mappings polaires (tous sans couture ni singularité) ────────────────────── def _sector_fwd(x, r_min, r_max, t_min, t_max): r = r_min + x[0] * (r_max - r_min) theta = t_min + x[1] * (t_max - t_min) return jnp.array([r * jnp.cos(theta), r * jnp.sin(theta)]) def _sector_inv(y, r_min, r_max, t_min, t_max): r = jnp.sqrt(y[0] ** 2 + y[1] ** 2) theta = jnp.arctan2(y[1], y[0]) return jnp.array([(r - r_min) / (r_max - r_min), (theta - t_min) / (t_max - t_min)]) # 6. Quart d'anneau — r∈[0.25,0.8], θ∈[0,π/2] def quart_anneau_fwd(x): return _sector_fwd(x, 0.25, 0.8, 0.0, jnp.pi / 2.0) def quart_anneau_inv(y): return _sector_inv(y, 0.25, 0.8, 0.0, jnp.pi / 2.0) # 7. Demi-anneau — r∈[0.2,0.85], θ∈[0,π] → forme en D def demi_anneau_fwd(x): return _sector_fwd(x, 0.2, 0.85, 0.0, jnp.pi) def demi_anneau_inv(y): return _sector_inv(y, 0.2, 0.85, 0.0, jnp.pi) # 8. Secteur large symétrique — r∈[0.05,1.0], θ∈[-π/3,π/3] → pointe à droite def wedge_fwd(x): return _sector_fwd(x, 0.05, 1.0, -jnp.pi / 3.0, jnp.pi / 3.0) def wedge_inv(y): return _sector_inv(y, 0.05, 1.0, -jnp.pi / 3.0, jnp.pi / 3.0) # 9. Arc 240° symétrique — r∈[0.35,0.7], θ∈[-2π/3, 2π/3] → forme en C épaisse def arc240_fwd(x): return _sector_fwd(x, 0.35, 0.7, -2.0 * jnp.pi / 3.0, 2.0 * jnp.pi / 3.0) def arc240_inv(y): return _sector_inv(y, 0.35, 0.7, -2.0 * jnp.pi / 3.0, 2.0 * jnp.pi / 3.0) MAPPINGS = [ ("Carré (ref)", id_fwd, id_inv), ("Losange", diamond_fwd, diamond_inv), ("Trapèze", trapeze_fwd, trapeze_inv), ("Chebyshev", cheby_fwd, cheby_inv), ("Couche limite", bl_fwd, bl_inv), ("Quart d'anneau", quart_anneau_fwd, quart_anneau_inv), ("Demi-anneau (D)", demi_anneau_fwd, demi_anneau_inv), ("Secteur 60°", wedge_fwd, wedge_inv), ("Arc 240°", arc240_fwd, arc240_inv), ] # ── Construction des maillages ──────────────────────────────────────────────── meshes = [(name, make_mesh(fwd, inv)) for name, fwd, inv in MAPPINGS] # ── Plot ────────────────────────────────────────────────────────────────────── def _cell_polygons(mesh, n_edge=12): t = jnp.linspace(0.0, 1.0, n_edge, endpoint=False) edges_unit = jnp.concatenate( [ jnp.stack([t, jnp.zeros_like(t)], axis=1), jnp.stack([jnp.ones_like(t), t], axis=1), jnp.stack([1.0 - t, jnp.ones_like(t)], axis=1), jnp.stack([jnp.zeros_like(t), 1.0 - t], axis=1), ], axis=0, ) def cell_boundary(cell_idx): pts_ref = mesh._unit_hypercube_to_cell(cell_idx, edges_unit) return mesh.mapping.local_mapping(pts_ref) return np.array(jax.vmap(cell_boundary)(mesh.cells_idx)) def plot_all_meshes(meshes, with_quad=True): n = len(meshes) ncols = 3 nrows = (n + ncols - 1) // ncols fig, axes = plt.subplots(nrows, ncols, figsize=(5 * ncols, 4 * nrows)) axes = axes.ravel() for ax, (name, mesh) in zip(axes, meshes): for boundary in _cell_polygons(mesh): poly = np.vstack([boundary, boundary[0]]) ax.plot(poly[:, 0], poly[:, 1], "b-", lw=0.5) if with_quad: _, x_all = mesh.evaluate_mesh_weights_points() x_flat = np.array(x_all.reshape(-1, 2)) ax.scatter(x_flat[:, 0], x_flat[:, 1], s=3, c="r", zorder=3) ax.set_aspect("equal") ax.set_title(name, fontsize=10) ax.axis("off") for ax in axes[n:]: ax.axis("off") fig.suptitle(f"Maillages {n_cells}×{n_cells}", fontsize=13) plt.tight_layout() plt.show() plot_all_meshes(meshes, with_quad=True)