r"""Solves a 2D magnetostatic problem with 3 materials using PINNs. The domain is the unit square :math:`\Omega = (0, 1)^2`, containing: - a rectangular magnet :math:`\Omega_m = (0.3, 0.5) \times (0.3, 0.6)` - a polar piece :math:`\Omega_p = (0.5, 0.6) \times (0.2, 0.7)` - vacuum everywhere else """ 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 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 MagnetostaticNoBc(AbstractPhysicalModel): """3-material magnetostatic PDE: Laplace on vacuum (0), magnet (1), polar (2).""" 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=1, f_rhs=lambda x: x[0:1] * 0.0, model_type=model_type, ), "polar": IndexedLaplacianResidual( domain=main_domain, var_index=2, f_rhs=lambda x: x[0:1] * 0.0, model_type=model_type, ), } # ── RBF basis functions ──────────────────────────────────────────────────────── def phi_dir(x, centers, r0): """Gaussian RBF centered on boundary/interface points.""" diffs = x[None, :] - centers dists2 = jnp.sum(diffs**2, axis=-1) return jnp.exp(-dists2 / r0**2) phi_mid = phi_dir def phi_min(x, centers, normals, r0): """RBF vanishing on interface but with non-zero normal derivative.""" diffs = x[None, :] - centers dists2 = jnp.sum(diffs**2, axis=-1) proj = jnp.sum(normals * diffs, axis=-1) return proj * jnp.exp(-dists2 / r0**2) # ── Per-material correction vectors ─────────────────────────────────────────── # vacuum: boundary + mag-vac interface + pol-vac interface # magnet: mag-vac interface + mag-pol interface # polar : pol-vac interface + mag-pol interface def get_correction_vector_vac(x, coll, r0_b, r0_i_m, r0_i_p): """RBF row for vacuum field (size = n_b + 2*n_mv + 2*n_pv).""" v_dir = phi_dir(x, coll["xy_bound"], r0_b) v_mid_mv = phi_mid(x, coll["xy_int_mag_vac"], r0_i_m) v_min_mv = phi_min(x, coll["xy_int_mag_vac"], coll["n_int_mag_vac"], r0_i_m) v_mid_pv = phi_mid(x, coll["xy_int_pol_vac"], r0_i_p) v_min_pv = phi_min(x, coll["xy_int_pol_vac"], coll["n_int_pol_vac"], r0_i_p) return jnp.concatenate([v_dir, v_mid_mv, v_min_mv, v_mid_pv, v_min_pv]) def get_correction_vector_mag(x, coll, r0_i_m, r0_i_mp): """RBF row for magnet field (size = 2*n_mv + 2*n_mp).""" v_mid_mv = phi_mid(x, coll["xy_int_mag_vac"], r0_i_m) v_min_mv = phi_min(x, coll["xy_int_mag_vac"], coll["n_int_mag_vac"], r0_i_m) v_mid_mp = phi_mid(x, coll["xy_int_mag_pol"], r0_i_mp) v_min_mp = phi_min(x, coll["xy_int_mag_pol"], coll["n_int_mag_pol"], r0_i_mp) return jnp.concatenate([v_mid_mv, v_min_mv, v_mid_mp, v_min_mp]) def get_correction_vector_pol(x, coll, r0_i_p, r0_i_mp): """RBF row for polar field (size = 2*n_pv + 2*n_mp).""" v_mid_pv = phi_mid(x, coll["xy_int_pol_vac"], r0_i_p) v_min_pv = phi_min(x, coll["xy_int_pol_vac"], coll["n_int_pol_vac"], r0_i_p) v_mid_mp = phi_mid(x, coll["xy_int_mag_pol"], r0_i_mp) v_min_mp = phi_min(x, coll["xy_int_mag_pol"], coll["n_int_mag_pol"], r0_i_mp) return jnp.concatenate([v_mid_pv, v_min_pv, v_mid_mp, v_min_mp]) def _make_normal_deriv(get_cv_fn): def dn(x, n_vec, coll, *r0s): def f(x_): return get_cv_fn(x_, coll, *r0s) _, deriv = jax.jvp(f, (x,), (n_vec,)) return deriv return dn get_correction_normal_deriv_vac = _make_normal_deriv(get_correction_vector_vac) get_correction_normal_deriv_mag = _make_normal_deriv(get_correction_vector_mag) get_correction_normal_deriv_pol = _make_normal_deriv(get_correction_vector_pol) def get_nn_normal_deriv(model, x, n_vec): """Normal derivative of a neural network output.""" grad_u = jax.grad(lambda x_col: model(x_col)[0])(x) return jnp.dot(grad_u, n_vec) # ── Approximation space ──────────────────────────────────────────────────────── class RBFHardConstrainsApproximation(AbstractApproxSpace): """RBF hard-constraint approximation for 3 materials.""" rbf_scale: float = 0.1 x_bnd: NDARRAY_TYPE n_bnd: NDARRAY_TYPE x_int_mag_vac: NDARRAY_TYPE n_int_mag_vac: NDARRAY_TYPE x_int_mag_pol: NDARRAY_TYPE n_int_mag_pol: NDARRAY_TYPE x_int_pol_vac: NDARRAY_TYPE n_int_pol_vac: NDARRAY_TYPE physical_params: dict def __init__( self, dims: dict[str, int], model_type: str = "x_dofsl", rbf_scale: float = 0.1, x_bnd: NDARRAY_TYPE = None, n_bnd: NDARRAY_TYPE = None, x_int_mag_vac: NDARRAY_TYPE = None, n_int_mag_vac: NDARRAY_TYPE = None, x_int_mag_pol: NDARRAY_TYPE = None, n_int_mag_pol: NDARRAY_TYPE = None, x_int_pol_vac: NDARRAY_TYPE = None, n_int_pol_vac: 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.rbf_scale = rbf_scale self.x_bnd = x_bnd self.n_bnd = n_bnd self.x_int_mag_vac = x_int_mag_vac self.n_int_mag_vac = n_int_mag_vac self.x_int_mag_pol = x_int_mag_pol self.n_int_mag_pol = n_int_mag_pol self.x_int_pol_vac = x_int_pol_vac self.n_int_pol_vac = n_int_pol_vac self.physical_params = physical_params models = [m for m, _, _ in list_models] types_models = [t for _, t, _ in list_models] size_models = [1 if size is None else size for _, _, size 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 _r0s(self): p = self.physical_params return p["r0_b"], p["r0_i_m"], p["r0_i_mp"], p["r0_i_p"] def _coll(self): return { "xy_bound": self.x_bnd, "xy_int_mag_vac": self.x_int_mag_vac, "n_int_mag_vac": self.n_int_mag_vac, "xy_int_mag_pol": self.x_int_mag_pol, "n_int_mag_pol": self.n_int_mag_pol, "xy_int_pol_vac": self.x_int_pol_vac, "n_int_pol_vac": self.n_int_pol_vac, } def compute_ndof(self) -> int: return sum(m.ndof() for m in self.models) def compute_coeff(self): """Build and solve the 7-block monolithic linear system. Returns (c_v, c_m, c_p). """ model_v, model_m, model_p = self.models coll = self._coll() r0_b, r0_i_m, r0_i_mp, r0_i_p = self._r0s() mu_m = self.physical_params["mu_m"] mu_p = self.physical_params["mu_p"] Bc = self.physical_params["Bc"] rcond = self.physical_params.get("rcond", 1e-6) delta = self.physical_params.get("delta_flux", 0.0) xy_b = coll["xy_bound"] xy_mv = coll["xy_int_mag_vac"] n_mv = coll["n_int_mag_vac"] xy_mp = coll["xy_int_mag_pol"] n_mp = coll["n_int_mag_pol"] xy_pv = coll["xy_int_pol_vac"] n_pv = coll["n_int_pol_vac"] N_cv = xy_b.shape[0] + 2 * xy_mv.shape[0] + 2 * xy_pv.shape[0] N_cm = 2 * xy_mv.shape[0] + 2 * xy_mp.shape[0] N_cp = 2 * xy_pv.shape[0] + 2 * xy_mp.shape[0] def av(pts): return jax.vmap( lambda x: get_correction_vector_vac(x, coll, r0_b, r0_i_m, r0_i_p) )(pts) # noqa: N802 def am(pts): return jax.vmap( lambda x: get_correction_vector_mag(x, coll, r0_i_m, r0_i_mp) )(pts) # noqa: N802 def ap(pts): return jax.vmap( lambda x: get_correction_vector_pol(x, coll, r0_i_p, r0_i_mp) )(pts) # noqa: N802 def dn_av(pts, ns): return jax.vmap( lambda x, n: get_correction_normal_deriv_vac( x, n, coll, r0_b, r0_i_m, r0_i_p ) )(pts, ns) def dn_am(pts, ns): return jax.vmap( lambda x, n: get_correction_normal_deriv_mag( x, n, coll, r0_i_m, r0_i_mp ) )(pts, ns) def dn_ap(pts, ns): return jax.vmap( lambda x, n: get_correction_normal_deriv_pol( x, n, coll, r0_i_p, r0_i_mp ) )(pts, ns) def zv(n): return jnp.zeros((n, N_cv)) def zm(n): return jnp.zeros((n, N_cm)) def zp(n): return jnp.zeros((n, N_cp)) n_b, n_mv_, n_mp_, n_pv_ = ( xy_b.shape[0], xy_mv.shape[0], xy_mp.shape[0], xy_pv.shape[0], ) # Block 1: outer BC (vacuum) — u_v = 0 on boundary A1 = jnp.concatenate([av(xy_b), zm(n_b), zp(n_b)], axis=1) b1 = -jax.vmap(lambda x: model_v(x)[0])(xy_b) # Block 2: mag-vac continuity — u_m - u_v = 0 A2 = jnp.concatenate([-av(xy_mv), am(xy_mv), zp(n_mv_)], axis=1) b2 = jax.vmap(lambda x: model_v(x)[0])(xy_mv) - jax.vmap( lambda x: model_m(x)[0] )(xy_mv) # Block 3: mag-vac flux — mu_m*dn_u_v - dn_u_m = -Bc/mu_m * ny xy_mv_v = xy_mv + delta * n_mv # point inside the vacuum xy_mv_m = xy_mv - delta * n_mv # point inside the magnet dn_v_mv = jax.vmap(lambda x, n: get_nn_normal_deriv(model_v, x, n))( xy_mv_v, n_mv ) dn_m_mv = jax.vmap(lambda x, n: get_nn_normal_deriv(model_m, x, n))( xy_mv_m, n_mv ) A3 = jnp.concatenate( [mu_m * dn_av(xy_mv, n_mv), -dn_am(xy_mv, n_mv), zp(n_mv_)], axis=1 ) b3 = -Bc / mu_m * n_mv[:, 1] + (1.0 / mu_m) * dn_m_mv - dn_v_mv # Block 4: pol-vac continuity — u_p - u_v = 0 A4 = jnp.concatenate([-av(xy_pv), zm(n_pv_), ap(xy_pv)], axis=1) b4 = jax.vmap(lambda x: model_v(x)[0])(xy_pv) - jax.vmap( lambda x: model_p(x)[0] )(xy_pv) # Block 5: pol-vac flux — (1/mu_p)*dn_u_p - dn_u_v = 0 xy_pv_v = xy_pv + delta * n_pv # point inside the vacuum xy_pv_p = xy_pv - delta * n_pv # point inside the polar piece dn_v_pv = jax.vmap(lambda x, n: get_nn_normal_deriv(model_v, x, n))( xy_pv_v, n_pv ) dn_p_pv = jax.vmap(lambda x, n: get_nn_normal_deriv(model_p, x, n))( xy_pv_p, n_pv ) A5 = jnp.concatenate( [-dn_av(xy_pv, n_pv), zm(n_pv_), (1.0 / mu_p) * dn_ap(xy_pv, n_pv)], axis=1 ) b5 = -dn_v_pv + (1.0 / mu_p) * dn_p_pv # Block 6: mag-pol continuity — u_p - u_m = 0 A6 = jnp.concatenate([zv(n_mp_), -am(xy_mp), ap(xy_mp)], axis=1) b6 = jax.vmap(lambda x: model_m(x)[0])(xy_mp) - jax.vmap( lambda x: model_p(x)[0] )(xy_mp) # Block 7: mag-pol flux — (1/mu_m)*dn_u_m - (1/mu_p)*dn_u_p = 0 xy_mp_m = xy_mp - delta * n_mp # point in the magnet (n_mp points to polar) xy_mp_p = xy_mp + delta * n_mp # point inside the polar piece dn_m_mp = jax.vmap(lambda x, n: get_nn_normal_deriv(model_m, x, n))( xy_mp_m, n_mp ) dn_p_mp = jax.vmap(lambda x, n: get_nn_normal_deriv(model_p, x, n))( xy_mp_p, n_mp ) A7 = jnp.concatenate( [ zv(n_mp_), (1.0 / mu_m) * dn_am(xy_mp, n_mp), -(1.0 / mu_p) * dn_ap(xy_mp, n_mp), ], axis=1, ) b7 = (1.0 / mu_p) * dn_p_mp + (1.0 / mu_m) * dn_m_mp A_global = jnp.concatenate([A1, A2, A3, A4, A5, A6, A7], axis=0) b_global = jnp.concatenate([b1, b2, b3, b4, b5, b6, b7], axis=0) c_global = jnp.linalg.pinv(A_global, rcond=rcond) @ b_global # returns a single concatenated vector → a single dofsl passed to the vmap return c_global def _split_coeff(self, c_all): n_b = self.x_bnd.shape[0] n_mv = self.x_int_mag_vac.shape[0] n_mp = self.x_int_mag_pol.shape[0] n_pv = self.x_int_pol_vac.shape[0] N_cv = n_b + 2 * n_mv + 2 * n_pv N_cm = 2 * n_mv + 2 * n_mp c_v = c_all[:N_cv] c_m = c_all[N_cv : N_cv + N_cm] c_p = c_all[N_cv + N_cm :] return c_v, c_m, c_p def get_intermediate_values(self) -> tuple[NDARRAY_TYPE, ...]: return (self.compute_coeff(),) def get_intermediate_values_shapes(self) -> tuple[tuple[int, ...]]: n_b = self.x_bnd.shape[0] n_mv = self.x_int_mag_vac.shape[0] n_mp = self.x_int_mag_pol.shape[0] n_pv = self.x_int_pol_vac.shape[0] N_cv = n_b + 2 * n_mv + 2 * n_pv N_cm = 2 * n_mv + 2 * n_mp N_cp = 2 * n_pv + 2 * n_mp return ((N_cv + N_cm + N_cp,),) def create_variables(self) -> tuple[ParamFunction, ...]: coll = self._coll() def eval_vacuum(space, *args): x, c_all = args[0], args[-1] c_v, _, _ = space._split_coeff(c_all) r0_b, r0_i_m, _, r0_i_p = space._r0s() corr = jnp.dot( get_correction_vector_vac(x, coll, r0_b, r0_i_m, r0_i_p), c_v ) return space.models[0](x)[0] + corr def eval_magnet(space, *args): x, c_all = args[0], args[-1] _, c_m, _ = space._split_coeff(c_all) _, r0_i_m, r0_i_mp, _ = space._r0s() corr = jnp.dot(get_correction_vector_mag(x, coll, r0_i_m, r0_i_mp), c_m) return space.models[1](x)[0] + corr def eval_polar(space, *args): x, c_all = args[0], args[-1] _, _, c_p = space._split_coeff(c_all) _, _, r0_i_mp, r0_i_p = space._r0s() corr = jnp.dot(get_correction_vector_pol(x, coll, r0_i_p, r0_i_mp), c_p) return space.models[2](x)[0] + corr var1 = ParamScalarFunction(self.dims, eval_vacuum, f_type=self.model_type) var2 = ParamScalarFunction(self.dims, eval_magnet, f_type=self.model_type) var3 = ParamScalarFunction(self.dims, eval_polar, f_type=self.model_type) return (var1, var2, var3) # ── Domains ──────────────────────────────────────────────────────────────────── domain_x = [(0.0, 1.0), (0.0, 1.0)] domain_ix = [(0.3, 0.5), (0.3, 0.6)] # magnet domain_px = [(0.5, 0.6), (0.2, 0.7)] # polar piece vacuum = Square2D(domain_x, is_main_domain=True, label_str="vacuum") magnet = Square2D(domain_ix, is_main_domain=False, label_str="magnet") polar = Square2D(domain_px, is_main_domain=False, label_str="polar") vacuum.add_subdomain(magnet) vacuum.add_subdomain(polar) 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"] # ── Interface sampling ───────────────────────────────────────────────────────── def _uniform_segment(x0, y0, x1, y1, n): """n equidistant points on segment (x0,y0)→(x1,y1) with outward normal.""" xs = jnp.linspace(x0, x1, n) ys = jnp.linspace(y0, y1, n) return jnp.stack([xs, ys], axis=-1) def make_mag_vac_interface(domain_m, n_total): """3 sides of magnet bordering vacuum (left, bottom, top).""" x0, x1 = domain_m[0] y0, y1 = domain_m[1] h, w = y1 - y0, x1 - x0 perim = h + w + w n_l = max(1, round(n_total * h / perim)) n_b = max(1, round(n_total * w / perim)) n_t = max(1, n_total - n_l - n_b) sides = [ ( _uniform_segment(x0, y0, x0, y1, n_l), jnp.tile(jnp.array([-1.0, 0.0]), (n_l, 1)), ), ( _uniform_segment(x0, y0, x1, y0, n_b), jnp.tile(jnp.array([0.0, -1.0]), (n_b, 1)), ), ( _uniform_segment(x0, y1, x1, y1, n_t), jnp.tile(jnp.array([0.0, +1.0]), (n_t, 1)), ), ] pts = jnp.concatenate([s for s, _ in sides]) nrm = jnp.concatenate([n for _, n in sides]) return pts, nrm def make_mag_pol_interface(domain_m, n_total): """Right side of magnet = left side of polar (x=x1_m, y in [y0_m, y1_m]).""" x0, x1 = domain_m[0] y0, y1 = domain_m[1] pts = _uniform_segment(x1, y0, x1, y1, n_total) nrm = jnp.tile(jnp.array([+1.0, 0.0]), (n_total, 1)) return pts, nrm def make_pol_vac_interface(domain_m, domain_p, n_total): """Sides of polar not touching magnet.""" xm0, xm1 = domain_m[0] ym0, ym1 = domain_m[1] xp0, xp1 = domain_p[0] yp0, yp1 = domain_p[1] segs = [ # left-bottom: x=xp0, y in [yp0, ym0] ((xp0, yp0, xp0, ym0), [-1.0, 0.0]), # left-top: x=xp0, y in [ym1, yp1] ((xp0, ym1, xp0, yp1), [-1.0, 0.0]), # bottom: y=yp0, x in [xp0, xp1] ((xp0, yp0, xp1, yp0), [0.0, -1.0]), # right: x=xp1, y in [yp0, yp1] ((xp1, yp0, xp1, yp1), [+1.0, 0.0]), # top: y=yp1, x in [xp0, xp1] ((xp0, yp1, xp1, yp1), [0.0, +1.0]), ] lengths = [ ((x1s - x0s) ** 2 + (y1s - y0s) ** 2) ** 0.5 for (x0s, y0s, x1s, y1s), _ in segs ] total = sum(lengths) counts = [max(1, round(n_total * seg_len / total)) for seg_len in lengths] counts[-1] = max(1, n_total - sum(counts[:-1])) pts_list, nrm_list = [], [] for ((x0s, y0s, x1s, y1s), n), ni in zip(segs, counts): pts_list.append(_uniform_segment(x0s, y0s, x1s, y1s, ni)) nrm_list.append(jnp.tile(jnp.array(n), (ni, 1))) return jnp.concatenate(pts_list), jnp.concatenate(nrm_list) x_int_mag_vac, n_int_mag_vac = make_mag_vac_interface(domain_ix, n_total=30) x_int_mag_pol, n_int_mag_pol = make_mag_pol_interface(domain_ix, n_total=20) x_int_pol_vac, n_int_pol_vac = make_pol_vac_interface(domain_ix, domain_px, n_total=30) print(f"x_bnd : {x_bnd.shape}") print(f"x_int_mag_vac: {x_int_mag_vac.shape}") print(f"x_int_mag_pol: {x_int_mag_pol.shape}") print(f"x_int_pol_vac: {x_int_pol_vac.shape}") # ── adaptive r0 ─────────────────────────────────────────────────────────────── RBF_ALPHA = 1.0 boundary_perimeter = 4.0 # unit square perim_mv = float( 2 * (domain_ix[0][1] - domain_ix[0][0]) + (domain_ix[1][1] - domain_ix[1][0]) ) # 3 sides of magnet perim_mp = float(domain_ix[1][1] - domain_ix[1][0]) # right side of magnet perim_pv = float( 2 * (domain_px[0][1] - domain_px[0][0]) + (domain_px[1][1] - domain_px[1][0]) - (domain_ix[1][1] - domain_ix[1][0]) ) # polar sides minus shared with magnet r0_b_val = RBF_ALPHA * boundary_perimeter / (1 + x_bnd.shape[0]) r0_i_m_val = RBF_ALPHA * perim_mv / (1 + x_int_mag_vac.shape[0]) r0_i_mp_val = RBF_ALPHA * perim_mp / (1 + x_int_mag_pol.shape[0]) r0_i_p_val = RBF_ALPHA * perim_pv / (1 + x_int_pol_vac.shape[0]) print( f"r0_b={r0_b_val:.4f} r0_i_m={r0_i_m_val:.4f} r0_i_mp={r0_i_mp_val:.4f} r0_i_p={r0_i_p_val:.4f}" ) # ── Networks and space ───────────────────────────────────────────────────────── key, k1, k2, k3 = jax.random.split(key, 4) network1 = MLP(in_size=2, out_size=1, hidden_sizes=[14] * 4, key=k1, activation="tanh") network2 = MLP(in_size=2, out_size=1, hidden_sizes=[14] * 4, key=k2, activation="tanh") network3 = MLP(in_size=2, out_size=1, hidden_sizes=[14] * 4, key=k3, activation="tanh") base_physical_params = { "r0_b": r0_b_val, "r0_i_m": r0_i_m_val, "r0_i_mp": r0_i_mp_val, "r0_i_p": r0_i_p_val, "mu_m": 1.01, "mu_p": 2000.0, "Bc": 1.0, "delta_flux": 0.01, # δ > 0 to evaluate gradients slightly off the interface } sampler = TensorizedSampler([DomainSampler(vacuum)], bc=False) pde = MagnetostaticNoBc(main_domain=vacuum, model_type="x_dofsl") space = RBFHardConstrainsApproximation( dims={"x": 2}, model_type="x", x_bnd=x_bnd, n_bnd=n_bnd, x_int_mag_vac=x_int_mag_vac, n_int_mag_vac=n_int_mag_vac, x_int_mag_pol=x_int_mag_pol, n_int_mag_pol=n_int_mag_pol, x_int_pol_vac=x_int_pol_vac, n_int_pol_vac=n_int_pol_vac, physical_params={**base_physical_params, "rcond": 1e-4}, list_models=[ (network1, "scalar", None), (network2, "scalar", None), (network3, "scalar", None), ], ) weights_dict = {"vacuum": [1.0], "magnet": [1.0], "polar": [1.0]} # ── Training ─────────────────────────────────────────────────────────────────── rcond_schedule = [1e-4] epochs_schedule = [4000] N_COLLOC = 4000 start = timeit.default_timer() current_space = space all_losses = [] for rcond_val, n_ep in zip(rcond_schedule, epochs_schedule): current_space.physical_params = {**base_physical_params, "rcond": rcond_val} pinn = Projector( pde, current_space, sampler, weights=weights_dict, matrix_regularization=1.0e-4, linesearch="armijo", alpha=0.01, beta=0.5, learning_rate=0.0001, nb_max_steps=20, ) key, pinn = pinn.project(key, current_space, n_ep, N_COLLOC) current_space = pinn.space all_losses.append( (rcond_val, jnp.asarray(pinn.losses.losses_history["total"]).reshape(-1)) ) print(f" rcond={rcond_val:.0e} best_loss={pinn.best_loss}") end = timeit.default_timer() print("best loss final: ", pinn.best_loss) print(f"time total: {end - start:.1f}s") # ── Evaluation on grid ───────────────────────────────────────────────────────── dofsl = pinn.space.get_intermediate_values() var_vacuum, var_magnet, var_polar = 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_vacuum = jax.vmap(var_vacuum, in_axes=(None, 0, None))( pinn.space, xy_plot, dofsl[0] ).reshape(n_plot, n_plot) u_magnet = jax.vmap(var_magnet, in_axes=(None, 0, None))( pinn.space, xy_plot, dofsl[0] ).reshape(n_plot, n_plot) u_polar = jax.vmap(var_polar, in_axes=(None, 0, None))( pinn.space, xy_plot, dofsl[0] ).reshape(n_plot, n_plot) # physical field: select subdomain-specific network in_magnet = ( (X >= domain_ix[0][0]) & (X <= domain_ix[0][1]) & (Y >= domain_ix[1][0]) & (Y <= domain_ix[1][1]) ) in_polar = ( (X >= domain_px[0][0]) & (X <= domain_px[0][1]) & (Y >= domain_px[1][0]) & (Y <= domain_px[1][1]) ) u_physical = jnp.where(in_magnet, u_magnet, jnp.where(in_polar, u_polar, u_vacuum)) # ── Plots ────────────────────────────────────────────────────────────────────── def add_rects(ax): for domain, color in [(domain_ix, "white"), (domain_px, "cyan")]: rect = plt.Rectangle( (domain[0][0], domain[1][0]), domain[0][1] - domain[0][0], domain[1][1] - domain[1][0], fill=False, edgecolor=color, linewidth=1.5, ) ax.add_patch(rect) fig, axs = plt.subplots(1, 4, figsize=(20, 5)) for ax, u, title in zip( axs, (u_vacuum, u_magnet, u_polar, u_physical), ("vacuum", "magnet", "polar piece", "physical field"), ): c = ax.contourf(X, Y, u, levels=40, cmap="jet") fig.colorbar(c, ax=ax) add_rects(ax) ax.set_title(title) ax.set_xlabel("x") ax.set_ylabel("y") ax.set_aspect("equal") plt.tight_layout() # ── Validation vs CSV ────────────────────────────────────────────────────────── csv_path = Path(__file__).parent / "Magneto_tri_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_v = jax.vmap(var_vacuum, in_axes=(None, 0, None))( pinn.space, xy_val, dofsl[0] ) u_val_m = jax.vmap(var_magnet, in_axes=(None, 0, None))( pinn.space, xy_val, dofsl[0] ) u_val_p = jax.vmap(var_polar, in_axes=(None, 0, None))(pinn.space, xy_val, dofsl[0]) mask_m = ( (x_val >= domain_ix[0][0]) & (x_val <= domain_ix[0][1]) & (y_val >= domain_ix[1][0]) & (y_val <= domain_ix[1][1]) ) mask_p = ( (x_val >= domain_px[0][0]) & (x_val <= domain_px[0][1]) & (y_val >= domain_px[1][0]) & (y_val <= domain_px[1][1]) ) u_val_phys = jnp.where(mask_m, u_val_m, jnp.where(mask_p, u_val_p, u_val_v)) err = u_val_phys - A_ref err_abs = jnp.abs(err) 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(err_abs)) linf_rel = linf_err / float(jnp.max(jnp.abs(A_ref))) l2_m = ( float(jnp.sqrt(jnp.mean(err[mask_m] ** 2))) if jnp.any(mask_m) else float("nan") ) l2_p = ( float(jnp.sqrt(jnp.mean(err[mask_p] ** 2))) if jnp.any(mask_p) else float("nan") ) l2_v = float(jnp.sqrt(jnp.mean(err[~mask_m & ~mask_p] ** 2))) print( f"\nValidation → L2={l2_err:.3e} L2_rel={l2_rel:.3e} Linf={linf_err:.3e} Linf_rel={linf_rel:.3e}" ) print( f" L2 vacuum={l2_v:.3e} L2 magnet={l2_m:.3e} L2 polar={l2_p:.3e}" ) # scatter plots 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_phys[sort_idx]) err_s = np.array(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("CSV reference A(x,y)") 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("PINNS u(x,y)") vmax_err = float(jnp.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"|PINNS−ref| L2={l2_err:.2e} L2_rel={l2_rel:.2e}") for ax in axs3: ax.set_xlabel("x") ax.set_ylabel("y") ax.set_aspect("equal") add_rects(ax) plt.suptitle("PINNs vs CSV validation (3 materials)", fontsize=12) plt.tight_layout() # Bx, By validation csv_path_B = Path(__file__).parent / "Magneto_tri_validation_B.csv" if csv_path_B.exists(): data_B = np.genfromtxt(csv_path_B, delimiter=";", skip_header=1) x_vB = jnp.array(data_B[:, 0]) y_vB = jnp.array(data_B[:, 1]) xy_vB = jnp.stack([x_vB, y_vB], axis=-1) mask_mB = ( (x_vB >= domain_ix[0][0]) & (x_vB <= domain_ix[0][1]) & (y_vB >= domain_ix[1][0]) & (y_vB <= domain_ix[1][1]) ) mask_pB = ( (x_vB >= domain_px[0][0]) & (x_vB <= domain_px[0][1]) & (y_vB >= domain_px[1][0]) & (y_vB <= domain_px[1][1]) ) def a_phys(x_pt): uv = var_vacuum(pinn.space, x_pt, dofsl[0]) um = var_magnet(pinn.space, x_pt, dofsl[0]) up = var_polar(pinn.space, x_pt, dofsl[0]) in_m = ( (x_pt[0] >= domain_ix[0][0]) & (x_pt[0] <= domain_ix[0][1]) & (x_pt[1] >= domain_ix[1][0]) & (x_pt[1] <= domain_ix[1][1]) ) in_p = ( (x_pt[0] >= domain_px[0][0]) & (x_pt[0] <= domain_px[0][1]) & (x_pt[1] >= domain_px[1][0]) & (x_pt[1] <= domain_px[1][1]) ) return jnp.where(in_m, um, jnp.where(in_p, up, uv)) dA = jax.vmap(lambda xi: jax.grad(a_phys)(xi))(xy_vB) Bx_pred = dA[:, 1] By_pred = -dA[:, 0] for key_B, pred, col_idx in [("Bx", Bx_pred, 2), ("By", By_pred, 3)]: B_ref = jnp.array(data_B[:, col_idx]) err_B = jnp.abs(pred - B_ref) l2_B = float(jnp.sqrt(jnp.mean((pred - B_ref) ** 2))) l2_ref_B = float(jnp.sqrt(jnp.mean(B_ref**2))) print( f"{key_B}: L2={l2_B:.3e} L2_rel={l2_B / l2_ref_B:.3e} Linf={float(jnp.max(err_B)):.3e}" ) else: print(f"Validation CSV not found: {csv_path}") # ── Loss history ─────────────────────────────────────────────────────────────── fig_l, ax_l = plt.subplots(figsize=(10, 4)) offset = 0 for rcond_val, curve in all_losses: epochs = jnp.arange(offset, offset + len(curve)) ax_l.semilogy(epochs, curve, label=f"rcond={rcond_val:.0e}") ax_l.axvline(offset, color="gray", linestyle=":", linewidth=0.8) offset += len(curve) ax_l.set_xlabel("Epoch (cumulative)") ax_l.set_ylabel("Loss") ax_l.set_title("Loss — three-material magnetostatics") ax_l.legend(fontsize=8) ax_l.grid(True, alpha=0.3) plt.tight_layout() plt.show()