r"""Magnetostatic PINN with explicit lifting — zero interface loss. We set u = N_θ + w, where the lifting w is chosen so as to satisfy the flux jump μ∂_n u_1 - ∂_n u_2 = -Bc·n_y exactly: w(x, y) = -Bc/μ · (y - y_top) if y > y_top (vacuum above the magnet) w(x, y) = Bc/μ · (y_bot - y) if y < y_bot (vacuum below) w(x, y) = 0 otherwise Verification. The jump condition reads μ ∂_n u_vac - ∂_n u_mag = -Bc · n_y, where n is the outward normal OF THE MAGNET (leaving the magnet, entering the vacuum); that convention has to be fixed before checking anything, as the sign of ∂_n w depends on it. Top side: n_mag = (0, +1), n_y = +1 ∂_n w_vac = ∂_y w · (+1) = (-Bc/μ) · (+1) = -Bc/μ μ ∂_n w_vac - 0 = μ · (-Bc/μ) = -Bc = -Bc · n_y ✓ Bottom side: n_mag = (0, -1), n_y = -1 ∂_n w_vac = ∂_y w · (-1) = (-Bc/μ) · (-1) = Bc/μ μ ∂_n w_vac - 0 = Bc = -Bc · (-1) = -Bc · n_y ✓ The network N_θ therefore learns v = u - w, which satisfies: Δv = 0 everywhere (since Δw = 0 for a linear w) [v] = 0 (automatic with a single network) [μ ∂_n v] = 0 (residual ≈ 0 since μ_m = 1.01 ≈ 1) v = -w on ∂Ω (Dirichlet BC, enforced through the RBF correction) """ import timeit from pathlib import Path from typing import Callable import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.domains.meshless_domains.base import VolumetricDomain from scimba_jax.domains.meshless_domains.domains_2d import Square2D from scimba_jax.nonlinear_approximation.approximation_spaces.approximation_spaces import ( AbstractApproxSpace, default_post_processing, ) from scimba_jax.nonlinear_approximation.approximation_spaces.dg_approximation_spaces import ( dg_pre_processing, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.model_class.funcparam import ParamFunction from scimba_jax.nonlinear_approximation.model_class.funcparam_vectorial import ( ParamScalarFunction, ) from scimba_jax.nonlinear_approximation.networks.mlp import MLP from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector from scimba_jax.physical_models.abstract_physical_model import ( PHYSICAL_RESIDUALS_TYPE, AbstractPhysicalModel, ) from scimba_jax.physical_models.abstract_residuals import PARAM_FUNC_TYPE from scimba_jax.physical_models.elliptic_pde.laplacians import LaplacianResidual from scimba_jax.utils.typing_protocols import NDARRAY_TYPE # ── PDE ─────────────────────────────────────────────────────────────────────── class IndexedLaplacianResidual(LaplacianResidual): """Laplacian residual constraining vars[var_index].""" var_index: int = 0 def __init__( self, domain: VolumetricDomain, var_index: int = 0, f_rhs=None, model_type: str = "x_dofsl", ): super().__init__(domain=domain, f_rhs=f_rhs, model_type=model_type) self.var_index = var_index def construct_residual(self, *vars: PARAM_FUNC_TYPE) -> PARAM_FUNC_TYPE: return -vars[self.var_index].laplacian("x") class LiftingMagnetostaticPDE(AbstractPhysicalModel): """Δu = 0 everywhere — single variable vars[0] = N_θ + w (lifting included).""" def __init__(self, main_domain: VolumetricDomain, model_type: str = "x_dofsl"): super().__init__(main_domain=main_domain) self.physical_residuals: PHYSICAL_RESIDUALS_TYPE = { "vacuum": IndexedLaplacianResidual( domain=main_domain, var_index=0, f_rhs=lambda x: x[0:1] * 0.0, model_type=model_type, ), "magnet": IndexedLaplacianResidual( domain=main_domain, var_index=0, f_rhs=lambda x: x[0:1] * 0.0, model_type=model_type, ), } # ── Analytic lifting ────────────────────────────────────────────────────────── def lifting_w(x, y_bot: float, y_top: float, Bc: float, mu_m: float) -> jnp.ndarray: """w(x,y) satisfying the flux jump μ∂_n u_v - ∂_n u_m = -Bc·n_y.""" y = x[1] C = Bc / mu_m w_above = -C * (y - y_top) w_below = C * (y_bot - y) return jnp.where(y > y_top, w_above, jnp.where(y < y_bot, w_below, 0.0)) # ── RBF (boundary only) ─────────────────────────────────────────────────────── def phi_dir(x, centers, r0): diffs = x[None, :] - centers dists2 = jnp.sum(diffs**2, axis=-1) return jnp.exp(-dists2 / r0**2) def get_correction_vector_bc(x, xy_bound, r0_b): """RBF correction vector — outer boundary only.""" return phi_dir(x, xy_bound, r0_b) # ── Approximation space ─────────────────────────────────────────────────────── class LiftingRBFApprox(AbstractApproxSpace): """Single network + lifting w + RBF correction on the boundary only. No interface points, no flux-jump loss. The Dirichlet BC u = 0 on ∂Ω is enforced through the RBF correction: N_θ + corr = -w. """ x_bnd: NDARRAY_TYPE n_bnd: NDARRAY_TYPE physical_params: dict def __init__( self, dims: dict[str, int], model_type: str = "x_dofsl", x_bnd: NDARRAY_TYPE = None, n_bnd: NDARRAY_TYPE = None, physical_params: dict = None, list_models: list = [], pre_processing: Callable | list[Callable] = [dg_pre_processing], post_processing: Callable | list[Callable] = [default_post_processing], ): super().__init__(dims=dims, model_type=model_type) self.x_bnd = x_bnd self.n_bnd = n_bnd self.physical_params = physical_params models = [m for m, _, _ in list_models] types_models = [t for _, t, _ in list_models] size_models = [1 if s is None else s for _, _, s in list_models] pre_processings = ( [pre_processing] if callable(pre_processing) else pre_processing ) post_processings = ( [post_processing] if callable(post_processing) else post_processing ) if len(pre_processings) == 1: pre_processings *= len(models) if len(post_processings) == 1: post_processings *= len(models) model_type = model_type + "_dofsl" if "dofsl" not in model_type else model_type super().__init__( dims, model_type, models, types_models, size_models, pre_processings, post_processings, ) def _w(self, x): p = self.physical_params return lifting_w(x, p["y_bot"], p["y_top"], p["Bc"], p["mu_m"]) def compute_ndof(self) -> int: return sum(m.ndof() for m in self.models) def compute_coeff(self) -> NDARRAY_TYPE: """Boundary-only linear system: A·c = -(N_θ(x_b) + w(x_b)). The matrix A is square (n_b × n_b) — either pinv or solve can be used. """ model = self.models[0] xy_b = self.x_bnd r0_b = self.physical_params["r0_b"] rcond = self.physical_params.get("rcond", 1e-6) A = jax.vmap(lambda x: get_correction_vector_bc(x, xy_b, r0_b))(xy_b) b = -jax.vmap(lambda x: model(x)[0] + self._w(x))(xy_b) return jnp.linalg.pinv(A, rcond=rcond) @ b def get_intermediate_values(self) -> tuple[NDARRAY_TYPE, ...]: return (self.compute_coeff(),) def get_intermediate_values_shapes(self) -> tuple[tuple[int, ...]]: return ((self.x_bnd.shape[0],),) def create_variables(self) -> tuple[ParamFunction, ...]: xy_b = self.x_bnd def eval_u(space, *args): x, c = args[0], args[-1] r0_b = space.physical_params["r0_b"] corr = jnp.dot(get_correction_vector_bc(x, xy_b, r0_b), c) return space.models[0](x)[0] + corr + space._w(x) return (ParamScalarFunction(self.dims, eval_u, f_type=self.model_type),) # ── Domains ─────────────────────────────────────────────────────────────────── domain_x = [(0.0, 1.0), (0.0, 1.0)] domain_ix = [(0.3, 0.5), (0.3, 0.6)] y_bot = domain_ix[1][0] # 0.3 y_top = domain_ix[1][1] # 0.6 vacuum = Square2D(domain_x, is_main_domain=True, label_str="vacuum") magnet = Square2D(domain_ix, is_main_domain=False, label_str="magnet") vacuum.add_subdomain(magnet) vacuum.set_boundaries_dict({"boundary": ["bc south", "bc east", "bc north", "bc west"]}) key = jax.random.PRNGKey(0) sampler_init = TensorizedSampler([DomainSampler(vacuum)], bc=True) key, bc_samples = sampler_init.bc_sample(key, {"boundary": 40}) x_bnd, n_bnd = bc_samples["boundary"] print(f"x_bnd : {x_bnd.shape}") # ── Adaptive r0 ─────────────────────────────────────────────────────────────── RBF_ALPHA = 1.0 boundary_perimeter = 2 * (domain_x[0][1] - domain_x[0][0]) + 2 * ( domain_x[1][1] - domain_x[1][0] ) r0_b_val = RBF_ALPHA * boundary_perimeter / (1 + x_bnd.shape[0]) base_physical_params = { "r0_b": r0_b_val, "mu_m": 1.01, "Bc": 1.0, "y_bot": y_bot, "y_top": y_top, } # ── Network and space ───────────────────────────────────────────────────────── key, subkey = jax.random.split(key) network = MLP( in_size=2, out_size=1, hidden_sizes=[20] * 4, key=subkey, activation="silu" ) pde = LiftingMagnetostaticPDE(main_domain=vacuum, model_type="x_dofsl") space = LiftingRBFApprox( dims={"x": 2}, model_type="x", x_bnd=x_bnd, n_bnd=n_bnd, physical_params={**base_physical_params, "rcond": 1e-2}, list_models=[(network, "scalar", None)], ) sampler = TensorizedSampler([DomainSampler(vacuum)], bc=False) # ── Training ────────────────────────────────────────────────────────────────── RCOND_VAL = 1e-4 REG_ENG = 5.0e-4 N_EPOCHS = 2000 N_COLLOC = 5000 start = timeit.default_timer() space.physical_params = {**base_physical_params, "rcond": RCOND_VAL} pinn = Projector( pde, space, sampler, weights={"vacuum": [1.0], "magnet": [1.0]}, matrix_regularization=REG_ENG, linesearch="armijo", alpha=0.01, beta=0.5, learning_rate=0.0001, nb_max_steps=20, ) key, pinn = pinn.project(key, space, N_EPOCHS, N_COLLOC) end = timeit.default_timer() print(f"best loss: {pinn.best_loss['total']:.3e} | time: {end - start:.1f}s") # ── Evaluation ──────────────────────────────────────────────────────────────── dofsl = pinn.space.get_intermediate_values() (var_u,) = pinn.space.create_variables() n_plot = 100 x_lin = jnp.linspace(0.0, 1.0, n_plot) y_lin = jnp.linspace(0.0, 1.0, n_plot) X, Y = jnp.meshgrid(x_lin, y_lin) xy_plot = jnp.stack([X.ravel(), Y.ravel()], axis=-1) u_pred = jax.vmap(var_u, in_axes=(None, 0, None))( pinn.space, xy_plot, dofsl[0] ).reshape(n_plot, n_plot) in_magnet = ( (X >= domain_ix[0][0]) & (X <= domain_ix[0][1]) & (Y >= domain_ix[1][0]) & (Y <= domain_ix[1][1]) ) fig, ax = plt.subplots(figsize=(6, 5)) cf = ax.contourf(X, Y, u_pred, levels=40, cmap="jet") ax.contour(X, Y, u_pred, levels=40, colors="k", linewidths=0.3) fig.colorbar(cf, ax=ax) rect = plt.Rectangle( (domain_ix[0][0], domain_ix[1][0]), domain_ix[0][1] - domain_ix[0][0], domain_ix[1][1] - domain_ix[1][0], fill=False, edgecolor="white", linewidth=1.5, ) ax.add_patch(rect) ax.set_title("Lifting PINN — physical field u = N_θ + w") ax.set_xlabel("x") ax.set_ylabel("y") ax.set_aspect("equal") plt.tight_layout() # ── PDE residual ────────────────────────────────────────────────────────────── def lap_u(x): return jnp.trace(jax.hessian(lambda x_: var_u(pinn.space, x_, dofsl[0]))(x)) res = jax.vmap(lap_u)(xy_plot).reshape(n_plot, n_plot) print(f"L2 residual of Δu : {float(jnp.sqrt(jnp.mean(res**2))):.3e}") # ── Validation against the CSV ──────────────────────────────────────────────── csv_path = Path(__file__).parent / "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 = jax.vmap(var_u, in_axes=(None, 0, None))(pinn.space, xy_val, dofsl[0]) err = u_val - A_ref l2_err = float(jnp.sqrt(jnp.mean(err**2))) l2_ref = float(jnp.sqrt(jnp.mean(A_ref**2))) l2_rel = l2_err / l2_ref linf_err = float(jnp.max(jnp.abs(err))) linf_rel = linf_err / float(jnp.max(jnp.abs(A_ref))) print( f"CSV validation → L2={l2_err:.3e} L2_rel={l2_rel:.3e} " f"Linf={linf_err:.3e} Linf_rel={linf_rel:.3e}" ) sort_idx = jnp.argsort(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 = np.array(u_val[sort_idx]) err_s = np.array(jnp.abs(err)[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("CSV reference") sc1 = axs3[1].scatter(x_s, y_s, c=u_s, cmap="jet", s=1) fig3.colorbar(sc1, ax=axs3[1]) axs3[1].set_title("Lifting PINN") vmax_err = float(jnp.percentile(jnp.abs(err), 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"|PINN − ref| L2_rel={l2_rel:.2e}") for ax in axs3: ax.set_xlabel("x") ax.set_ylabel("y") ax.set_aspect("equal") plt.suptitle("Lifting PINN vs CSV — two-material magnetostatics", fontsize=12) plt.tight_layout() else: print(f"CSV not found : {csv_path}") plt.show()