"""Solve the 2D Laplacian DG scheme via Newton-Raphson. Exact solution: u(x,y) = sin(π(x-a1)/(b1-a1)) sin(π(y-a2)/(b2-a2)) on [a1,b1]×[a2,b2] PDE: -Δu = π²(1/(b1-a1)² + 1/(b2-a2)²) u, u = 0 on ∂Ω """ import time from pathlib import Path import jax import jax.numpy as jnp import matplotlib.patches as mpatches import matplotlib.pyplot as plt import numpy as np from scimba_jax.domains.meshless_domains.domains_2d import Square2D from scimba_jax.linear_approximation.basis.analytic_bases import local_taylor_basis from scimba_jax.linear_approximation.basis.general_bases import AnalyticBasis from scimba_jax.linear_approximation.galerkin.dg.elliptic_dg_scheme import ( EllipticDGscheme, ) from scimba_jax.linear_approximation.galerkin.dg.flux import MagnetoStaticSIPGFlux 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.physical_models.abstract_physical_weak_model import ( AbstractPhysicalWeakModel, ) from scimba_jax.physical_models.classical_weakform.magneto_static_weak_form import ( MagnetoStaticWeakForm, ) # ── Plotting the magnet domain ───────────────────────────────────────────────────── def plot_domain_and_mesh(mesh, square_domain, magnet_domain): """Plot the domain with the magnet inclusion and mesh.""" nx, ny = mesh.n_cells x0_v, x1_v = float(square_domain.bounds[0, 0]), float(square_domain.bounds[0, 1]) y0_v, y1_v = float(square_domain.bounds[1, 0]), float(square_domain.bounds[1, 1]) x0_m, x1_m = float(magnet_domain.bounds[0, 0]), float(magnet_domain.bounds[0, 1]) y0_m, y1_m = float(magnet_domain.bounds[1, 0]), float(magnet_domain.bounds[1, 1]) x_nodes = np.linspace(x0_v, x1_v, nx + 1) y_nodes = np.linspace(y0_v, y1_v, ny + 1) fig, ax = plt.subplots(figsize=(6, 6)) for x in x_nodes: ax.axvline(x, color="gray", linewidth=0.5, zorder=1) for y in y_nodes: ax.axhline(y, color="gray", linewidth=0.5, zorder=1) ax.add_patch( mpatches.Rectangle( (x0_v, y0_v), x1_v - x0_v, y1_v - y0_v, linewidth=2, edgecolor="black", facecolor="lightblue", alpha=0.4, zorder=2, label=r"$\Omega_v$", ) ) ax.add_patch( mpatches.Rectangle( (x0_m, y0_m), x1_m - x0_m, y1_m - y0_m, linewidth=2, edgecolor="red", facecolor="salmon", alpha=0.6, zorder=3, label=r"$\Omega_m$", ) ) ax.text(x0_v + 0.05, y1_v - 0.08, r"$\Omega_v$", fontsize=14, color="steelblue") ax.text( (x0_m + x1_m) / 2, (y0_m + y1_m) / 2, r"$\Omega_m$", fontsize=14, color="darkred", ha="center", va="center", ) margin = 0.02 ax.set_xlim(x0_v - margin, x1_v + margin) ax.set_ylim(y0_v - margin, y1_v + margin) ax.set_aspect("equal") ax.set_xlabel("x") ax.set_ylabel("y") ax.set_title(r"Domaine $\Omega_v$ avec inclusion $\Omega_m$") ax.legend(loc="upper right", fontsize=12) plt.tight_layout() plt.savefig("domain_plot.png", dpi=150) plt.show() # ── Plot source term ───────────────────────────────────────────────────────── def plot_source_term(pde, square_domain, magnet_domain, n_plot=51): """Plot the source term and magnet domain.""" x0_v, x1_v = float(square_domain.bounds[0, 0]), float(square_domain.bounds[0, 1]) y0_v, y1_v = float(square_domain.bounds[1, 0]), float(square_domain.bounds[1, 1]) x0_m, x1_m = float(magnet_domain.bounds[0, 0]), float(magnet_domain.bounds[0, 1]) y0_m, y1_m = float(magnet_domain.bounds[1, 0]), float(magnet_domain.bounds[1, 1]) xs = np.linspace(x0_v, x1_v, n_plot) ys = np.linspace(y0_v, y1_v, n_plot) XX, YY = np.meshgrid(xs, ys) pts = jnp.array(np.stack([XX.ravel(), YY.ravel()], axis=-1)) f_vals = np.array(jax.vmap(lambda p: pde.source_term(p).ravel()[0])(pts)) F = f_vals.reshape(n_plot, n_plot) print("max(f_vals) =", f_vals.max(), "min(f_vals) =", f_vals.min()) rect_x = [x0_m, x1_m, x1_m, x0_m, x0_m] rect_y = [y0_m, y0_m, y1_m, y1_m, y0_m] fig, ax = plt.subplots(figsize=(6, 5)) fig.suptitle("Source term $f$") im = ax.pcolormesh(XX, YY, F, shading="auto", cmap="RdBu_r") ax.plot(rect_x, rect_y, "k-", lw=1.5, label="magnet boundary") ax.set_aspect("equal") ax.legend() plt.colorbar(im, ax=ax) plt.tight_layout() plt.show() # ── Solution plots ────────────────────────────────────────────────────── def plot_solution_2d(assembler, square_domain, magnet_domain, n_plot=51): """Plot u_h on a regular grid — style identique à magnetostatic.py.""" mesh = assembler.variables.mesh x0_v, x1_v = float(square_domain.bounds[0, 0]), float(square_domain.bounds[0, 1]) y0_v, y1_v = float(square_domain.bounds[1, 0]), float(square_domain.bounds[1, 1]) x0_m, x1_m = float(magnet_domain.bounds[0, 0]), float(magnet_domain.bounds[0, 1]) y0_m, y1_m = float(magnet_domain.bounds[1, 0]), float(magnet_domain.bounds[1, 1]) xs = np.linspace(x0_v, x1_v, n_plot) ys = np.linspace(y0_v, y1_v, n_plot) XX, YY = np.meshgrid(xs, ys) pts = jnp.array(np.stack([XX.ravel(), YY.ravel()], axis=-1)) UH = np.array(assembler.variables.evaluate(pts))[:, 0].reshape(n_plot, n_plot) n_cells_val = mesh.n_cells[0] LEVELS = 40 fig, ax = plt.subplots(figsize=(6, 5)) cf = ax.contourf(XX, YY, UH, levels=LEVELS, cmap="jet") ax.contour(XX, YY, UH, levels=LEVELS, colors="k", linewidths=0.3) fig.colorbar(cf, ax=ax) rect = plt.Rectangle( (x0_m, y0_m), x1_m - x0_m, y1_m - y0_m, fill=False, edgecolor="white", linewidth=1.5, ) ax.add_patch(rect) ax.set_title(f"DG SIPG — {n_cells_val}×{n_cells_val} mailles — $u_h$") ax.set_xlabel("x") ax.set_ylabel("y") ax.set_aspect("equal") plt.tight_layout() plt.show() # ── Problem setup ───────────────────────────────────────────────────────────── physical_dim = 2 out_dim = 1 order = 1 # polynomial degree, Q1 elements nb_basis = (order + 1) ** physical_dim quad_order = 4 n_cells = 50 h_phys = 1.0 / n_cells mapping_id = InvertibleFunction(lambda x: x, lambda y: y) square_domain = Square2D(bounds=[(0.0, 1.0), (0.0, 1.0)], is_main_domain=True) magnet_domain = Square2D(bounds=[(0.3, 0.5), (0.3, 0.6)], is_main_domain=False) def dirichlet_bc(x): return jnp.zeros(out_dim) # ── PDE setup ─────────────────────────────────────────────────────────────── pde = MagnetoStaticWeakForm( magnet_domain=magnet_domain, ) # print("\nPlotting source term...") # plot_source_term(pde, square_domain, magnet_domain) # ── Mesh setup ───────────────────────────────────────────────────────────── m = Mesh( dim=physical_dim, n_cells=(n_cells, n_cells), ref_quad=UnitSquareTensorized(dim=physical_dim, order=quad_order), mapping=Mapping(mappings=[mapping_id]), is_identity_mapping=False, ) # print("\nPlotting domain and mesh...") # plot_domain_and_mesh(m, square_domain, magnet_domain) # ── Solver setup ───────────────────────────────────────────────────────── taylor_basis = AnalyticBasis( nb_basis=nb_basis, out_dim=out_dim, mesh=m, local_basis=lambda coords, i, mesh: local_taylor_basis( coords, i, mesh, order=order, out_dim=out_dim ), basis_type="scalar", ) flux = MagnetoStaticSIPGFlux(sigma=order * (order + 1), h=h_phys) variables = VariablesDG(basis=taylor_basis, nb_variables=out_dim) assembler = EllipticDGscheme( AbstractPhysicalWeakModel.from_weak_form(pde, dirichlet=dirichlet_bc), variables, flux, ) assembler = EllipticDGscheme.solve(assembler, max_iter=1, matrix_free=True, tol=1e-6) print("\nPlotting solution..") plot_solution_2d(assembler, square_domain, magnet_domain, n_plot=201) # ── Validation vs CSV ───────────────────────────────────────────────────────── csv_path = ( Path(__file__).parent.parent.parent.parent / "pinns/stationary_pdes/elliptic_pdes/Magneto_bi_validation.csv" ) if csv_path.exists(): data_val = np.genfromtxt(csv_path, delimiter=";", skip_header=1) x_val = jnp.array(data_val[:, 0]) y_val = jnp.array(data_val[:, 1]) A_ref = jnp.array(data_val[:, 2]) xy_val = jnp.stack([x_val, y_val], axis=-1) u_val = np.array(assembler.variables.evaluate(xy_val))[:, 0] err = u_val - np.array(A_ref) err_abs = np.abs(err) l2_err = float(np.sqrt(np.mean(err**2))) l2_ref = float(np.sqrt(np.mean(np.array(A_ref) ** 2))) l2_rel = l2_err / l2_ref linf_err = float(np.max(err_abs)) linf_rel = linf_err / float(np.max(np.abs(np.array(A_ref)))) x0_m, x1_m = float(magnet_domain.bounds[0, 0]), float(magnet_domain.bounds[0, 1]) y0_m, y1_m = float(magnet_domain.bounds[1, 0]), float(magnet_domain.bounds[1, 1]) mask_m = ( (np.array(x_val) >= x0_m) & (np.array(x_val) <= x1_m) & (np.array(y_val) >= y0_m) & (np.array(y_val) <= y1_m) ) l2_magnet = ( float(np.sqrt(np.mean(err[mask_m] ** 2))) if mask_m.any() else float("nan") ) l2_vaccum = float(np.sqrt(np.mean(err[~mask_m] ** 2))) print( f"\nValidation CSV → L2={l2_err:.3e} L2_rel={l2_rel:.3e} " f"Linf={linf_err:.3e} Linf_rel={linf_rel:.3e}" ) print(f" L2 vaccum={l2_vaccum:.3e} L2 magnet={l2_magnet:.3e}") sort_idx = np.argsort(np.array(x_val)) x_s = np.array(x_val)[sort_idx] y_s = np.array(y_val)[sort_idx] A_s = np.array(A_ref)[sort_idx] u_s = u_val[sort_idx] err_s = err_abs[sort_idx] fig3, axs3 = plt.subplots(1, 3, figsize=(18, 5)) sc0 = axs3[0].scatter(x_s, y_s, c=A_s, cmap="jet", s=1) fig3.colorbar(sc0, ax=axs3[0]) axs3[0].set_title("Référence CSV A(x,y)") axs3[0].set_aspect("equal") sc1 = axs3[1].scatter( x_s, y_s, c=u_s, cmap="jet", s=1, vmin=A_s.min(), vmax=A_s.max() ) fig3.colorbar(sc1, ax=axs3[1]) axs3[1].set_title("DG SIPG $u_h(x,y)$") axs3[1].set_aspect("equal") vmax_err = float(np.percentile(err_abs, 99)) sc2 = axs3[2].scatter(x_s, y_s, c=err_s, cmap="hot_r", s=1, vmin=0, vmax=vmax_err) fig3.colorbar(sc2, ax=axs3[2]) axs3[2].set_title(f"|DG − ref| L2={l2_err:.2e} L2_rel={l2_rel:.2e}") axs3[2].set_aspect("equal") for ax in axs3: ax.set_xlabel("x") ax.set_ylabel("y") rect = plt.Rectangle( (x0_m, y0_m), x1_m - x0_m, y1_m - y0_m, fill=False, edgecolor="white", linewidth=1.5, ) ax.add_patch(rect) plt.suptitle(f"DG SIPG {n_cells}×{n_cells} vs CSV — magnétostatique", fontsize=12) plt.tight_layout() plt.show() else: print(f"\nCSV de validation non trouvé : {csv_path}") # ── Benchmark : boucle Python vs vmap sur 200 valeurs de Bc ────────────────── # # Problème avec jax.vmap(EllipticDGscheme.solve) : # → effet de bord : scheme_pytree.variables.dofsl = dofsl_sol (mutation Python) # → vmap interdit tout accès en écriture à un état global pendant la trace # # Solution générale : solve_pure — encapsule les fonctions internes pures # (assembly_fn, jacobian_fn) + jnp.linalg.solve, sans aucune mutation. # Fonctionne pour tout paramètre (Bc, μ, géométrie…) dès que : # - le weak form utilise des ops JAX traçables (Bc * array, pas jnp.array([0., Bc])) # - le paramètre est passé comme scalaire JAX à solve_pure N_BC = 100 bc_values = jnp.linspace(0.4, 1.04, N_BC) _assembly_fn = EllipticDGscheme._make_assembly_fn() _jacobian_fn = EllipticDGscheme._make_jacobian_fn() def solve_pure(bc): """Solve pur : construit assembler(bc), assemble K et b, résout K·u = b.""" pde_bc = MagnetoStaticWeakForm(magnet_domain=magnet_domain, Bc=bc) asm_bc = EllipticDGscheme( AbstractPhysicalWeakModel.from_weak_form(pde_bc, dirichlet=dirichlet_bc), variables, flux, ) d0 = jnp.zeros_like(asm_bc.variables.dofsl) b = -_assembly_fn(asm_bc, d0) # b varie avec bc K = _jacobian_fn(asm_bc, d0) # K varie avec μ (ici μ fixe) return jnp.linalg.solve(K, b) # ── vmap ────────────────────────────────────────────────────────────────────── # solve_pure est une fonction pure JAX → compatible vmap print("Warmup vmap (JIT)...") solve_vmap = jax.jit(jax.vmap(solve_pure)) jax.block_until_ready(solve_vmap(bc_values[:2])) print("vmap run...") t0 = time.perf_counter() dofsl_vmap = solve_vmap(bc_values) jax.block_until_ready(dofsl_vmap) t_vmap = time.perf_counter() - t0 print(f" vmap : {t_vmap:.3f} s ({N_BC} solves)") # ── Évaluateur DG vectorisé sur une coupe 1D ────────────────────────────────── # _classical_local_evaluate_pure prend (variables_pytree, dofsl, point) — # dofsl est un argument explicite → compatible jax.vmap. n_cut = 201 x_cut = jnp.linspace(0.0, 1.0, n_cut) pts_cut = jnp.stack([x_cut, jnp.full(n_cut, 0.6)], axis=-1) dofsl_shape = assembler.variables.dofsl.shape # (n_cells_total, n_basis, out_dim) def evaluate_cut(dofsl_flat): """Évalue u(x, y=0.6) pour un dofsl flat donné — pur, sans effet de bord.""" dofsl_shaped = dofsl_flat.reshape(dofsl_shape) def eval_pt(pt): return VariablesDG._classical_local_evaluate_pure( assembler.variables, dofsl_shaped, pt ) return jax.vmap(eval_pt)(pts_cut)[:, 0] # shape (n_cut,) print("\nÉvaluation vectorisée sur la coupe y=0.6...") evaluate_cut_batch = jax.jit(jax.vmap(evaluate_cut)) jax.block_until_ready(evaluate_cut_batch(dofsl_vmap[:2])) # warmup JIT u_cut_all = np.array(evaluate_cut_batch(dofsl_vmap)) # (N_BC, n_cut) u_mean = u_cut_all.mean(axis=0) u_std = u_cut_all.std(axis=0) # ── Plot coupe 1D : 10 courbes + moyenne ± écart-type ───────────────────────── x_cut_np = np.array(x_cut) idx_10 = np.linspace(0, N_BC - 1, 10, dtype=int) bc_10 = np.array(bc_values)[idx_10] colors_10 = plt.cm.viridis(np.linspace(0, 1, 10)) fig, ax = plt.subplots(figsize=(9, 4)) for i, (idx, bc_val) in enumerate(zip(idx_10, bc_10)): ax.plot( x_cut_np, u_cut_all[idx], color=colors_10[i], lw=0.9, alpha=0.7, label=f"Bc={bc_val:.2f}", ) ax.plot(x_cut_np, u_mean, "k-", lw=2.0, label="moyenne") ax.fill_between( x_cut_np, u_mean - u_std, u_mean + u_std, alpha=0.25, color="gray", label="±σ" ) ax.axvline(0.3, color="r", linestyle="--", lw=0.8) ax.axvline(0.5, color="r", linestyle="--", lw=0.8, label="bords aimant") ax.set_xlabel("x") ax.set_ylabel(r"$u(x,\,y=0.6)$") ax.set_title( f"Coupe à $y=0.6$ (interface haute) — {N_BC} valeurs de $B_c \\in [0.4, 1.04]$" ) ax.legend(fontsize=7, ncol=4, loc="upper right") ax.grid(True, alpha=0.3) plt.tight_layout() plt.show()