"""Système couplé advection-diffusion avec deux espaces CG-FEM indépendants. Même PDE que le cas DG multi-space (examples_jax/dg/.../solve_2d_system_diffusion_advection_multi_space.py) : -ε Δu + b₁·∇u + c₁·∇v = f sur Ω, u|∂Ω = 0 -ε Δv + b₂·∇v + c₂·∇u = f sur Ω, v|∂Ω = 0 u et v vivent chacun sur un maillage indépendant (tailles différentes). Le couplage croisé (c₁·∇v dans l'eq de u, etc.) est traité via évaluation cross-space : find_cell_index localise la maille de l'autre mesh contenant chaque point de quadrature. En CG-FEM, pas de flux : uniquement les termes de volume + Dirichlet fort aux nœuds de bord de chaque espace. """ import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.linear_approximation.basis.analytic_bases import local_lagrange_basis from scimba_jax.linear_approximation.basis.general_bases import AnalyticBasis from scimba_jax.linear_approximation.galerkin.fem.elliptic_fe_scheme_multi_space import ( EllipticFEschemeMultipleSpaces, ) 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_fe import VariablesFE from scimba_jax.mapping.mapping import InvertibleFunction, Mapping from scimba_jax.nonlinear_approximation.model_class.funcparam_vectorial import ( ParamVecFunction, ) from scimba_jax.physical_models.abstract_physical_weak_model import ( AbstractPhysicalWeakModel, ) from scimba_jax.physical_models.abstract_weak_form import AbstractWeakForm from scimba_jax.physical_models.weak_boundary_conditions import Dirichlet # ── Paramètres (identiques au cas DG multi-space) ───────────────────────────── physical_dim = 2 quad_order = 3 poly_order = 1 n_cells_u = 40 n_cells_v = 30 sigma_gauss = 0.07 x0 = jnp.array([0.5, 0.5]) def gaussian_source(x): return jnp.exp(-jnp.sum((x - x0) ** 2) / (2.0 * sigma_gauss**2)) / ( 2.0 * jnp.pi * sigma_gauss**2 ) # ── Forme faible (identique au cas DG) ─────────────────────────────────────── class SystemDiffAdvWeakForm(AbstractWeakForm): def __init__(self, dim, eps, b1, b2, c1, c2, f): super().__init__(dim=dim) self.A = lambda x: eps * jnp.eye(dim) self.b1 = b1 self.b2 = b2 self.c1 = c1 self.c2 = c2 self.f = f def bilinear_form( self, u: ParamVecFunction, v: ParamVecFunction ) -> ParamVecFunction: u_1, u_2 = u # un argument par espace v_1, v_2 = v grad_u1 = u_1.gradient("x") grad_u2 = u_2.gradient("x") grad_v1 = v_1.gradient("x") grad_v2 = v_2.gradient("x") fields = self.get_fields() A = fields["A"] b1 = fields["b1"] b2 = fields["b2"] c1 = fields["c1"] c2 = fields["c2"] diff_u = grad_v1.dot(A @ grad_u1) adv_u = grad_u1.dot(b1) * v_1 cross_adv_u = grad_u2.dot(c1) * v_1 diff_v = grad_v2.dot(A @ grad_u2) adv_v = grad_u2.dot(b2) * v_2 cross_adv_v = grad_u1.dot(c2) * v_2 return ParamVecFunction.cat( [diff_u + adv_u + cross_adv_u, diff_v + adv_v + cross_adv_v] ) def linear_form(self, v: ParamVecFunction) -> ParamVecFunction: v_1, v_2 = v f = self.get_fields()["f"] return ParamVecFunction.cat([f * v_1, f * v_2]) # ── Construction du problème avec deux espaces ─────────────────────────────── b1_vec = jnp.array([1.0, 0.0]) b2_vec = jnp.array([-1.0, 0.0]) c1_vec = jnp.array([0.0, 0.3]) c2_vec = jnp.array([0.0, -0.3]) mapping_id = InvertibleFunction(lambda x: x, lambda y: y) eps = 0.06 pde = SystemDiffAdvWeakForm( dim=physical_dim, eps=eps, b1=lambda x: b1_vec, b2=lambda x: b2_vec, c1=lambda x: c1_vec, c2=lambda x: c2_vec, f=gaussian_source, ) def make_variables(n_cells): mesh = Mesh( dim=physical_dim, n_cells=(n_cells, n_cells), ref_quad=UnitSquareTensorized(dim=physical_dim, order=quad_order), mapping=Mapping(mappings=[mapping_id]), ) basis = AnalyticBasis( nb_basis=(poly_order + 1) ** physical_dim, out_dim=1, mesh=mesh, local_basis=lambda coords, i, m: local_lagrange_basis( coords, i, m, order=poly_order, out_dim=1 ), basis_type="scalar", ) return VariablesFE(basis=basis, nb_variables=1) variables_u = make_variables(n_cells_u) # espace 0 variables_v = make_variables(n_cells_v) # espace 1 def dirichlet_bc_u(x): return jnp.zeros(1) def dirichlet_bc_v(x): return jnp.zeros(1) # Une CL par espace, dans l'ordre des espaces. model = AbstractPhysicalWeakModel.from_weak_form(pde) model.add_boundary_condition("0/boundary", Dirichlet(dirichlet_bc_u)) model.add_boundary_condition("1/boundary", Dirichlet(dirichlet_bc_v)) assembler = EllipticFEschemeMultipleSpaces( pde=model, variables_list=[variables_u, variables_v], equation_spaces=[0, 1], ) # ── Résolution ──────────────────────────────────────────────────────────────── ndof_u = variables_u.ndof_linear ndof_v = variables_v.ndof_linear print(f"Espace u : {n_cells_u}²×Q{poly_order}, DOFs = {ndof_u}") print(f"Espace v : {n_cells_v}²×Q{poly_order}, DOFs = {ndof_v}") print(f"Total DOFs : {ndof_u + ndof_v}") print("Résolution (Newton matrix-free, BiCGStab)…") # Advection + couplage croisé → Jacobien non symétrique → BiCGStab (pas CG). assembler = EllipticFEschemeMultipleSpaces.solve( assembler, max_iter=1, matrix_free=True, cg_solver="bicgstab", tol=1e-6, verbose=True, ) print("Résolution terminée.") # ── Évaluation sur une grille commune ──────────────────────────────────────── n_plot = 60 xp = np.linspace(0.0, 1.0, n_plot) XC, YC = np.meshgrid(xp, xp) pts = jnp.array(np.stack([XC.ravel(), YC.ravel()], axis=-1)) U = np.array(assembler.variables_list[0].evaluate(pts)).reshape(n_plot, n_plot) V = np.array(assembler.variables_list[1].evaluate(pts)).reshape(n_plot, n_plot) def grad_squared(var, pts): """|∇u_h|² par autodiff de l'évaluation FE (find_cell_index + base). jax.grad traverse local_evaluate_pure : l'indice de maille est entier (dérivée nulle), le gradient vient de la base locale — en Q1 il est constant par maille. """ def u_scalar(x): return var.local_evaluate_pure(var, var.dofsl, x)[0] g = jax.vmap(jax.grad(u_scalar))(pts) # (N, dim) return np.array(jnp.sum(g**2, axis=-1)) GU = grad_squared(assembler.variables_list[0], pts).reshape(n_plot, n_plot) GV = grad_squared(assembler.variables_list[1], pts).reshape(n_plot, n_plot) # ── Visualisation ───────────────────────────────────────────────────────────── def draw_mesh(ax, n_cells, **kwargs): """Trace les lignes de maillage (grille Cartésienne uniforme sur [0,1]²).""" style = {"color": "w", "linewidth": 0.4, "alpha": 0.5} | kwargs for k in range(n_cells + 1): ax.axhline(k / n_cells, **style) ax.axvline(k / n_cells, **style) fig, axes = plt.subplots(1, 2, figsize=(10, 4)) fig.suptitle( f"Multi-space CG-FEM — u:{n_cells_u}² v:{n_cells_v}² ×Q{poly_order}\n" f"Cross-convection coupling (maillage de chaque espace superposé)", fontsize=11, ) w_h = np.concatenate([U.ravel()[:, None], V.ravel()[:, None]], axis=-1) vmin, vmax = float(w_h.min()), float(w_h.max()) for ax, Z, n_cells_s, title in zip( axes, [U, V], [n_cells_u, n_cells_v], [ f"$u$ — advection $(1,0)$ — mesh {n_cells_u}²", f"$v$ — advection $(-1,0)$ — mesh {n_cells_v}²", ], ): im = ax.pcolormesh(XC, YC, Z, shading="auto", cmap="turbo", vmin=vmin, vmax=vmax) ax.contour(XC, YC, Z, levels=10, colors="k", linewidths=0.5, alpha=0.4) draw_mesh(ax, n_cells_s) ax.set_title(title) ax.set_aspect("equal") plt.colorbar(im, ax=ax) plt.tight_layout() # ── |∇u|² et |∇v|² (autodiff de la solution discrète) ───────────────────────── fig_g, axes_g = plt.subplots(1, 2, figsize=(10, 4)) fig_g.suptitle( "Gradient carré de la solution — $|\\nabla u_h|^2$, $|\\nabla v_h|^2$ " "(constant par maille en Q1)", fontsize=11, ) for ax, Z, n_cells_s, title in zip( axes_g, [GU, GV], [n_cells_u, n_cells_v], [ f"$|\\nabla u_h|^2$ — mesh {n_cells_u}²", f"$|\\nabla v_h|^2$ — mesh {n_cells_v}²", ], ): im = ax.pcolormesh(XC, YC, Z, shading="auto", cmap="turbo") draw_mesh(ax, n_cells_s, alpha=0.25) ax.set_title(title) ax.set_aspect("equal") plt.colorbar(im, ax=ax) plt.tight_layout() # ── Les deux maillages seuls, superposés ────────────────────────────────────── fig2, ax2 = plt.subplots(figsize=(4.5, 4.5)) draw_mesh(ax2, n_cells_u, color="tab:blue", alpha=0.8, linewidth=0.6) draw_mesh(ax2, n_cells_v, color="tab:red", alpha=0.8, linewidth=0.9) ax2.plot([], [], color="tab:blue", label=f"mesh u : {n_cells_u}²") ax2.plot([], [], color="tab:red", label=f"mesh v : {n_cells_v}²") ax2.set_xlim(0, 1) ax2.set_ylim(0, 1) ax2.set_aspect("equal") ax2.set_title("Maillages indépendants des deux espaces") ax2.legend(loc="upper right", framealpha=0.9) plt.tight_layout() plt.show()