r"""Solves a 2D magnetostatic problem using PINNs. The domain is the unit square :math:`\Omega = (0, 1)^2`, containing a rectangular magnet sub-domain :math:`\Omega_m = (0.3, 0.5) \times (0.3, 0.6)`. """ 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.data_sampler import DataSampler 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, DataResidual, ) 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", ): 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): """2D magnetostatic PDE: Laplace on vacuum and magnet. ``create_variables`` returns (vacuum_full, magnet, vacuum_smooth). The vacuum Laplacian residual acts on the SMOOTH part (index 2): the singular basis (index 0 = full) is harmonic (Δ=0), hence legitimately excluded from the residual — and this avoids the blow-up of the autodiff of the Laplacian at the corner. """ def __init__( self, main_domain: VolumetricDomain, model_type: str = "x", ): super().__init__(main_domain=main_domain) self.physical_residuals: PHYSICAL_RESIDUALS_TYPE = { "vacuum": IndexedLaplacianResidual( domain=main_domain, var_index=2, # vacuum_smooth (without the singular basis) f_rhs=lambda x: x * 0.0, model_type=model_type, ), "magnet": IndexedLaplacianResidual( domain=main_domain, var_index=1, f_rhs=lambda x: x * 0.0, model_type=model_type, ), } def phi_dir(x, centers, r0): """RBF basis (Gaussian) centered on the outer boundary points.""" diffs = x[None, :] - centers dists2 = jnp.sum(diffs**2, axis=-1) return jnp.exp(-dists2 / r0**2) def phi_mid(x, centers, r0): """RBF basis (Gaussian) centered on the interface points.""" diffs = x[None, :] - centers dists2 = jnp.sum(diffs**2, axis=-1) return jnp.exp(-dists2 / r0**2) def phi_min(x, centers, normals, r0): """RBF basis vanishing on the interface but with a non-zero normal derivative. Used to correct the jump of the normal derivative across the interface without perturbing the continuity correction (phi_mid). """ 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) def get_correction_vector(x, collocation, r0_b, r0_i): """Returns the row of RBF coefficients for a given point x (size Nb + 2*Ni).""" v_dir = phi_dir(x, collocation["xy_bound"], r0_b) v_mid = phi_mid(x, collocation["xy_int"], r0_i) v_min = phi_min(x, collocation["xy_int"], collocation["n_int"], r0_i) return jnp.concatenate([v_dir, v_mid, v_min], axis=0) def get_correction_normal_deriv(x, n_vec, collocation, r0_b, r0_i): """Normal derivative of the correction basis with respect to n_vec.""" def f(x_col): return get_correction_vector(x_col, collocation, r0_b, r0_i) _, deriv = jax.jvp(f, (x,), (n_vec,)) return deriv 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) # ── Corner singular enrichment (OPTION) ─────────────────────────────────────── # Corners of the magnet rectangle and the "towards the inside of the magnet" # direction at each corner (to orient the branch cut of the log OUTSIDE the # vacuum). MAGNET_CORNERS = jnp.array([[0.3, 0.3], [0.5, 0.3], [0.5, 0.6], [0.3, 0.6]]) CORNER_BETA = jnp.array( # bisector towards the magnet: π/4, 3π/4, 5π/4, 7π/4 [jnp.pi / 4, 3 * jnp.pi / 4, 5 * jnp.pi / 4, 7 * jnp.pi / 4] ) # Source points shifted SLIGHTLY into the magnet (MFS idea): interface/boundary # points sometimes fall exactly on a corner → r=0 → derivative of arctan2 = nan. # Placing the source at SING_OFFSET inside the magnet keeps every evaluation # point at a distance ≥ SING_OFFSET → bounded derivatives, log cut inside the # magnet. The singularity seen from the vacuum side stays peaked at the corner. SING_OFFSET = 0.015 SING_SOURCES = MAGNET_CORNERS + SING_OFFSET * jnp.stack( [jnp.cos(CORNER_BETA), jnp.sin(CORNER_BETA)], axis=-1 ) def corner_singular_vec(x, corners, beta): """Flattened corner singular basis: [Re(w·logw), Im(w·logw)]_k → (2K,). w = x − corner (complex). Re/Im(w·log w) are HARMONIC (Δ=0) and their gradient diverges as log(1/r) at the corner — the field singularity of a polygonal magnet. ``beta`` places the log cut INSIDE the magnet (continuous on the vacuum side). Used ONLY through its value and first derivative (columns of the RBF solve of ``compute_coeff`` and field reconstruction), NEVER in the Laplacian residual → no second derivative, no blow-up. """ d = x[None, :] - corners # (K, 2) zx, zy = d[:, 0], d[:, 1] r = jnp.sqrt(zx**2 + zy**2 + 1e-18) ang = jnp.arctan2(zy, zx) theta = jnp.mod(ang - beta, 2 * jnp.pi) - jnp.pi # jump at ang = beta (magnet) logr = jnp.log(r) re = zx * logr - zy * theta # Re(w log w) im = zy * logr + zx * theta # Im(w log w) return jnp.stack([re, im], axis=-1).reshape(-1) # (2K,) def sing_normal_deriv(x, n_vec, corners, beta): """Normal derivative (along n_vec) of the singular basis: (2K,).""" def f(xx): return corner_singular_vec(xx, corners, beta) _, deriv = jax.jvp(f, (x,), (n_vec,)) return deriv class RBFHardConstrainsApproximation(AbstractApproxSpace): """A RBF approximation with hard constraints.""" rbf_scale: float = 0.1 x_bnd: NDARRAY_TYPE n_bnd: NDARRAY_TYPE x_interface: NDARRAY_TYPE n_interface: NDARRAY_TYPE physical_params: dict # OPTION: corner singular enrichment (basis added to the RBF solve) sing_corners: NDARRAY_TYPE sing_beta: NDARRAY_TYPE use_singular: bool = False 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_interface: NDARRAY_TYPE = None, n_interface: NDARRAY_TYPE = None, physical_params: dict = None, list_models: list = [], sing_corners: NDARRAY_TYPE = None, sing_beta: NDARRAY_TYPE = None, use_singular: bool = False, 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_interface = x_interface self.n_interface = n_interface self.physical_params = physical_params self.sing_corners = sing_corners self.sing_beta = sing_beta self.use_singular = use_singular 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) if not (len(pre_processings) == len(models)): raise ValueError( f"Number of pre_processings ({len(pre_processings)}) must match number of models ({len(models)})" ) if not (len(post_processings) == len(models)): raise ValueError( f"Number of post_processings ({len(post_processings)}) must match number of models ({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 compute_ndof(self) -> int: return sum(m.ndof() for m in self.models) def create_variables(self) -> tuple[ParamFunction, ...]: collocation = { "xy_bound": self.x_bnd, "xy_int": self.x_interface, "n_int": self.n_interface, } r0_b = self.physical_params["r0_b"] r0_i = self.physical_params["r0_i"] n_rbf = self.x_bnd.shape[0] + 2 * self.x_interface.shape[0] corners, beta, use_sing = self.sing_corners, self.sing_beta, self.use_singular def eval_vacuum_smooth(space, *args): # SMOOTH part: MLP + RBF correction (Gaussians). This is THE field # the Laplacian residual sees → bounded derivatives, no blow-up. x, c = args[0], args[-1] c_rbf = c[:n_rbf] corr = jnp.dot(get_correction_vector(x, collocation, r0_b, r0_i), c_rbf) return space.models[0](x)[0] + corr def eval_vacuum_full(space, *args): # FULL field = smooth + Σ aₖ χₖ (harmonic singular basis). Used for # the hard BCs, the data and the plots — never the Laplacian residual. x, c = args[0], args[-1] smooth = eval_vacuum_smooth(space, *args) if not use_sing: return smooth a_sing = c[n_rbf:] sing = jnp.dot(corner_singular_vec(x, corners, beta), a_sing) return smooth + sing def eval_magnet(space, *args): return space.models[1](args[0])[0] var_full = ParamScalarFunction( self.dims, eval_vacuum_full, f_type=self.model_type ) var_magnet = ParamScalarFunction(self.dims, eval_magnet, f_type=self.model_type) var_smooth = ParamScalarFunction( self.dims, eval_vacuum_smooth, f_type=self.model_type ) # order: (vacuum_full, magnet, vacuum_smooth) # → data/plots use index 0; the vacuum Laplacian residual uses index 2. return (var_full, var_magnet, var_smooth) def compute_coeff(self) -> NDARRAY_TYPE: """Solves the linear system enforcing the hard constraints and returns the RBF coefficients (size Nb + 2*Ni). Uses ``self.models`` directly, so gradients with respect to the network parameters of ``model_m`` and ``model_v`` flow through ``c``. """ model_v, model_m = self.models collocation = { "xy_bound": self.x_bnd, "xy_int": self.x_interface, "n_int": self.n_interface, } r0_b = self.physical_params["r0_b"] r0_i = self.physical_params["r0_i"] mu_m = self.physical_params["mu_m"] Bc = self.physical_params["Bc"] # --- step 1: build the interpolation matrix A --- A_block1 = jax.vmap( lambda x: get_correction_vector(x, collocation, r0_b, r0_i) )(collocation["xy_bound"]) A_block2 = jax.vmap( lambda x: get_correction_vector(x, collocation, r0_b, r0_i) )(collocation["xy_int"]) A_block3 = -mu_m * jax.vmap( lambda x, n: get_correction_normal_deriv(x, n, collocation, r0_b, r0_i) )(collocation["xy_int"], collocation["n_int"]) A = jnp.concatenate([A_block1, A_block2, A_block3], axis=0) # --- OPTION: corner singular columns (same 3 conditions) --- # The χₖ are harmonic → their amplitudes aₖ are solved for WITHIN this # least-squares problem (like c). Block 3 (RHS ∝ B_c·n_y, discontinuous # at the corner) activates them: Gaussians cannot fit a corner jump. if self.use_singular: corners, beta = self.sing_corners, self.sing_beta S1 = jax.vmap(lambda x: corner_singular_vec(x, corners, beta))( collocation["xy_bound"] ) S2 = jax.vmap(lambda x: corner_singular_vec(x, corners, beta))( collocation["xy_int"] ) S3 = -mu_m * jax.vmap(lambda x, n: sing_normal_deriv(x, n, corners, beta))( collocation["xy_int"], collocation["n_int"] ) A_sing = jnp.concatenate([S1, S2, S3], axis=0) A = jnp.concatenate([A, A_sing], axis=1) # extra columns # --- step 2: build the right-hand side b (errors of the NNs) --- b1 = -jax.vmap(lambda x: model_v(x)[0])(collocation["xy_bound"]) b2 = jax.vmap(lambda x: model_m(x)[0])(collocation["xy_int"]) - jax.vmap( lambda x: model_v(x)[0] )(collocation["xy_int"]) dn_NN_m = jax.vmap(lambda x, n: get_nn_normal_deriv(model_m, x, n))( collocation["xy_int"], collocation["n_int"] ) dn_NN_v = jax.vmap(lambda x, n: get_nn_normal_deriv(model_v, x, n))( collocation["xy_int"], collocation["n_int"] ) b3 = Bc * collocation["n_int"][:, 1] - dn_NN_m + mu_m * dn_NN_v b = jnp.concatenate([b1, b2, b3], axis=0) # --- step 3: solve for [c ; a] (RBF coeffs + singular amplitudes) --- rcond = self.physical_params.get("rcond", 1e-6) c = jnp.linalg.pinv(A, rcond=rcond) @ b return c 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_i = self.x_interface.shape[0] n = n_b + 2 * n_i if self.use_singular: n += 2 * self.sing_corners.shape[0] # singular amplitudes return ((n,),) # main domain: unit square (0, 1)^2 domain_x = [(0.0, 1.0), (0.0, 1.0)] # magnet sub-domain: (0.3, 0.5) x (0.3, 0.6) domain_ix = [(0.3, 0.5), (0.3, 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) # boundary groups: outer boundary of the square vs. magnet interface vacuum.set_boundaries_dict( { "boundary": ["bc south", "bc east", "bc north", "bc west"], "interface": ["bc magnet"], } ) sampler_init = TensorizedSampler([DomainSampler(vacuum)], bc=True) key = jax.random.PRNGKey(0) key, bc_samples = sampler_init.bc_sample(key, {"boundary": 40, "interface": 50}) x_bnd, n_bnd = bc_samples["boundary"] x_interface, n_interface = bc_samples["interface"] # equidistant interface: points spread proportionally to the length of each side def make_interface_equidist(domain_ix, n_total): x0, x1 = domain_ix[0] y0, y1 = domain_ix[1] w, h = x1 - x0, y1 - y0 perim = 2 * w + 2 * h n_b = max(1, round(n_total * w / perim)) n_t = max(1, round(n_total * w / perim)) n_l = max(1, round(n_total * h / perim)) n_r = n_total - n_b - n_t - n_l sides = [ ( jnp.linspace(x0, x1, n_b, endpoint=False), jnp.full(n_b, y0), jnp.array([0.0, -1.0]), ), ( jnp.linspace(x0, x1, n_t, endpoint=False), jnp.full(n_t, y1), jnp.array([0.0, +1.0]), ), ( jnp.full(n_l, x0), jnp.linspace(y0, y1, n_l, endpoint=False), jnp.array([-1.0, 0.0]), ), ( jnp.full(n_r, x1), jnp.linspace(y0, y1, n_r, endpoint=False), jnp.array([+1.0, 0.0]), ), ] pts = jnp.concatenate([jnp.stack([xs, ys], axis=-1) for xs, ys, _ in sides]) nrm = jnp.concatenate([jnp.tile(n, (len(xs), 1)) for xs, _, n in sides]) return pts, nrm x_interface, n_interface = make_interface_equidist(domain_ix, n_total=50) print(x_bnd.shape) # (40, 2) print(x_interface.shape) # (32, 2) sampler = TensorizedSampler([DomainSampler(vacuum)], bc=False) pde = MagnetostaticNoBc(main_domain=vacuum, model_type="x_dofsl") # OPTION: enrich the vacuum field with the corner singular basis. The functions # Re/Im(w log w) are HARMONIC → extra columns of the RBF solve of compute_coeff # (amplitudes fixed by the B_c interface condition at the corner, like the c # coefficients), and EXCLUDED from the Laplacian residual (Δ=0). It therefore # works with physics alone; at False, behaviour is identical to the original. USE_CORNER_SINGULARITY = False key, subkey = jax.random.split(key) network1 = MLP( in_size=2, out_size=1, hidden_sizes=[16] * 5, key=key, activation_type="silu" ) network2 = MLP( in_size=2, out_size=1, hidden_sizes=[16] * 5, key=subkey, activation_type="silu" ) # adaptive r0: mean spacing × alpha # alpha > 1 → wider RBFs → correction Laplacian of order 1/(alpha*r0)² # reduces the ill-conditioning of the loss without degrading BC enforcement too much RBF_ALPHA = 1.0 boundary_perimeter = 2 * (domain_x[0][1] - domain_x[0][0]) + 2 * ( domain_x[1][1] - domain_x[1][0] ) interface_perimeter = 2 * (domain_ix[0][1] - domain_ix[0][0]) + 2 * ( domain_ix[1][1] - domain_ix[1][0] ) r0_b_val = RBF_ALPHA * boundary_perimeter / (1 + x_bnd.shape[0]) r0_i_val = RBF_ALPHA * interface_perimeter / (1 + x_interface.shape[0]) print(f"r0_b={r0_b_val:.4f} r0_i={r0_i_val:.4f}") base_physical_params = {"r0_b": r0_b_val, "r0_i": r0_i_val, "mu_m": 1.01, "Bc": 1.0} space = RBFHardConstrainsApproximation( dims={"x": 2}, model_type="x", rbf_scale=0.1, x_bnd=x_bnd, n_bnd=n_bnd, x_interface=x_interface, n_interface=n_interface, physical_params={**base_physical_params, "rcond": 1e-2}, list_models=[(network1, "scalar", None), (network2, "scalar", None)], sing_corners=SING_SOURCES, sing_beta=CORNER_BETA, use_singular=USE_CORNER_SINGULARITY, ) weights_dict = {"vacuum": [1.0], "magnet": [1.0]} # ── rcond annealing: the constraints are tightened progressively ───────────── # Large rcond at the start → damped gradient through c → stable training # Small rcond at the end → tight constraints → accurate solution rcond_schedule = [3e-5] # [1e-4, 3.0e-5, 1e-5, 7e-6, 5e-6] epochs_schedule = [2000] # [1000, 500, 500, 500, 500] N_COLLOC = 5000 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, # optimizer="ENG", matrix_regularization=1e-5, ) 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") space_full = pinn.space # "full PINN" result (physics only) # ══════════════════════════════════════════════════════════════════════════════ # SECOND TRAINING — hybrid: physics loss + data loss (500 FEM points) # ══════════════════════════════════════════════════════════════════════════════ # # We load the reference FEM solution (Magneto_bi_validation.csv), draw 500 # points from it, and add a data loss comparing the PINN field to A_FEM at those # points — ON TOP OF the physics loss (vacuum + magnet Laplacian). A point # inside the magnet constrains vars[1] (magnet); outside it, vars[0] (vacuum). CSV_PATH = Path(__file__).parent / "Magneto_bi_validation.csv" N_DATA_OUT = 350 # data points outside the magnet N_DATA_IN = 150 # data points INSIDE the magnet (forced, to cover it) DATA_WEIGHT = 5.0 # weight of the data loss (physics = 1.0) # The whole hybrid run below needs the reference FEM solution. The CSV is not # shipped with the repository, so it is guarded: without it the script still # produces the physics-only PINN and its plots. if CSV_PATH.exists(): _raw = np.genfromtxt(CSV_PATH, delimiter=";", skip_header=1) xy_fem = jnp.asarray(_raw[:, :2]) A_fem = jnp.asarray(_raw[:, 2]) def _in_magnet(xy): return ( (xy[:, 0] >= domain_ix[0][0]) & (xy[:, 0] <= domain_ix[0][1]) & (xy[:, 1] >= domain_ix[1][0]) & (xy[:, 1] <= domain_ix[1][1]) ) # FEM data points: drawn separately INSIDE and OUTSIDE the magnet, to guarantee # coverage of the magnet (otherwise very few points land in it). mask_fem_in = np.array(_in_magnet(xy_fem)) idx_in_all = np.where(mask_fem_in)[0] idx_out_all = np.where(~mask_fem_in)[0] key, k_in, k_out = jax.random.split(key, 3) sel_in = np.array(jax.random.choice(k_in, idx_in_all, (N_DATA_IN,), replace=False)) sel_out = np.array( jax.random.choice(k_out, idx_out_all, (N_DATA_OUT,), replace=False) ) xy_in, A_in = xy_fem[sel_in], A_fem[sel_in] xy_out, A_out = xy_fem[sel_out], A_fem[sel_out] print(f"\nFEM data: {N_DATA_OUT} outside the magnet + {N_DATA_IN} inside") class IndexedDataResidual(DataResidual): """Data loss on vars[var_index]: u_var(x) = A_FEM(x).""" var_index: int = 0 def __init__( self, var_index: int = 0, size: int = 1, model_type: str = "x_dofsl" ): super().__init__(size=size, 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] # Hybrid model: same physics residuals + two data residuals. pde_hybrid = MagnetostaticNoBc(main_domain=vacuum, model_type="x_dofsl") pde_hybrid.add_data_residual("data_vacuum", IndexedDataResidual(var_index=0)) pde_hybrid.add_data_residual("data_magnet", IndexedDataResidual(var_index=1)) sampler_hybrid = TensorizedSampler( [DomainSampler(vacuum)], bc=False, data_samplers={ "data_vacuum": DataSampler((xy_out, A_out[:, None])), "data_magnet": DataSampler((xy_in, A_in[:, None])), }, ) weights_hybrid = { "vacuum": [1.0], "magnet": [1.0], "data_vacuum": [DATA_WEIGHT], "data_magnet": [DATA_WEIGHT], } # NEW network, trained FROM SCRATCH (same architecture/schedule as the full # PINN), for a fair "physics only" vs "physics + data" comparison. key, k1h, k2h = jax.random.split(key, 3) net1_h = MLP( in_size=2, out_size=1, hidden_sizes=[20] * 2, key=k1h, activation_type="silu" ) net2_h = MLP( in_size=2, out_size=1, hidden_sizes=[20] * 2, key=k2h, activation_type="silu" ) space_hybrid = RBFHardConstrainsApproximation( dims={"x": 2}, model_type="x", rbf_scale=0.1, x_bnd=x_bnd, n_bnd=n_bnd, x_interface=x_interface, n_interface=n_interface, physical_params={**base_physical_params, "rcond": 1e-2}, list_models=[(net1_h, "scalar", None), (net2_h, "scalar", None)], sing_corners=SING_SOURCES, sing_beta=CORNER_BETA, use_singular=USE_CORNER_SINGULARITY, ) # SS-Broyden optimizer (self-scaled quasi-Newton with line search) for the data # case: converges in far fewer iterations than ENG. matrix_regularization is # ENG-specific → it is not passed here. N_EPOCHS_HYBRID = 400 space_hybrid.physical_params = {**base_physical_params, "rcond": rcond_schedule[-1]} start_h = timeit.default_timer() pinn_hybrid = Projector( pde_hybrid, space_hybrid, sampler_hybrid, weights=weights_hybrid, optimizer="SS-Broyden", ) key, pinn_hybrid = pinn_hybrid.project(key, space_hybrid, N_EPOCHS_HYBRID, N_COLLOC) space_hybrid = pinn_hybrid.space print(f"hybrid best loss: {pinn_hybrid.best_loss}") print(f"hybrid time: {timeit.default_timer() - start_h:.1f}s") # ── Error vs FEM (Magneto_bi_validation.csv) for full PINN AND hybrid ──────── def physical_field(space, xy): """PINN physical field: vars[1] inside the magnet, vars[0] elsewhere.""" dofsl_s = space.get_intermediate_values() v_vac, v_mag, _ = space.create_variables() # (full, magnet, smooth) uv = jax.vmap(v_vac, in_axes=(None, 0, None))(space, xy, dofsl_s[0]) um = jax.vmap(v_mag, in_axes=(None, 0, None))(space, xy, dofsl_s[0]) return jnp.where(_in_magnet(xy), um, uv) A_full = physical_field(space_full, xy_fem) A_hyb = physical_field(space_hybrid, xy_fem) abserr_full = np.abs(np.array(A_full - A_fem)) # pointwise absolute error abserr_hyb = np.abs(np.array(A_hyb - A_fem)) l2_ref = float(jnp.sqrt(jnp.mean(A_fem**2))) l2_full = float(jnp.sqrt(jnp.mean((A_full - A_fem) ** 2))) / l2_ref l2_hyb = float(jnp.sqrt(jnp.mean((A_hyb - A_fem) ** 2))) / l2_ref print( f"\nRelative L2 error vs FEM: full PINN = {l2_full:.3e} hybrid = {l2_hyb:.3e}" ) figE, axE = plt.subplots(1, 2, figsize=(13, 5.2)) figE.suptitle( f"Pointwise error |$u_{{PINN}}-A_{{FEM}}$| — " f"full PINN: L2={l2_full:.2e} | hybrid (+{N_DATA_OUT + N_DATA_IN} data): " f"L2={l2_hyb:.2e}", fontsize=11, ) xy_fem_np = np.array(xy_fem) # Each panel on ITS own scale (99th percentile to ignore the few spikes), # otherwise the full PINN (small errors) is crushed. for ax, err, ti in zip( axE, [abserr_full, abserr_hyb], [f"full PINN — $L_2$={l2_full:.2e}", f"hybrid +data — $L_2$={l2_hyb:.2e}"], ): vmax = float(np.percentile(err, 99)) sc = ax.scatter( xy_fem_np[:, 0], xy_fem_np[:, 1], c=err, s=5, cmap="turbo", vmin=0.0, vmax=vmax, ) ax.add_patch( 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="w", linewidth=1.2, ) ) ax.set_title(ti) ax.set_aspect("equal") figE.colorbar(sc, ax=ax, label=r"$|u_{PINN}-A_{FEM}|$") plt.tight_layout() else: print(f"\nValidation CSV not found: {CSV_PATH}") print("Skipping the hybrid (physics + data) training and the FEM comparison.") # ── Plots ───────────────────────────────────────────────────────────────────── dofsl = pinn.space.get_intermediate_values() var_vacuum, var_magnet, _ = pinn.space.create_variables() # (full, magnet, smooth) 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) fig, axs = plt.subplots(1, 2, figsize=(11, 5)) for ax, u, title in zip( axs, (u_vacuum, u_magnet), ("vacuum field (network 1)", "magnet field (network 2 + RBF correction)"), ): c = ax.contourf(X, Y, u, levels=40, cmap="jet") ax.contour(X, Y, u, levels=40, colors="k", linewidths=0.5) fig.colorbar(c, 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(title) ax.set_xlabel("x") ax.set_ylabel("y") ax.set_aspect("equal") plt.tight_layout() # physical field: vacuum in the air, magnet inside the magnet in_magnet = ( (X >= domain_ix[0][0]) & (X <= domain_ix[0][1]) & (Y >= domain_ix[1][0]) & (Y <= domain_ix[1][1]) ) u_physical = jnp.where(in_magnet, u_magnet, u_vacuum) fig2, ax2 = plt.subplots(figsize=(6, 5)) cf = ax2.contourf(X, Y, u_physical, levels=20, cmap="jet") ax2.contour(X, Y, u_physical, levels=20, colors="k", linewidths=0.5) fig2.colorbar(cf, ax=ax2) rect2 = 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, ) ax2.add_patch(rect2) ax2.set_title("physical field (vacuum outside the magnet, magnet inside)") ax2.set_xlabel("x") ax2.set_ylabel("y") ax2.set_aspect("equal") plt.tight_layout() # ── Check that the gradient flows through c ────────────────────────────────── # Note: JIT-compiled functions are avoided (the cache fixes the trace, so the # patch would not be taken into account). We differentiate get_intermediate_values # directly. print("\n=== Checking ∂c/∂θ ===") # 1) Direct Jacobian through jac_intermediate_value_theta (≡ jacrev on compute_coeff) jac_c_theta = pinn.space.jac_intermediate_value_theta(0) print(f"∂c/∂θ shape : {jac_c_theta.shape}") print(f"max |∂c/∂θ| : {jnp.max(jnp.abs(jac_c_theta)):.4e}") print(f"mean |∂c/∂θ| : {jnp.mean(jnp.abs(jac_c_theta)):.4e}") # 2) Gradient of ‖c‖² with respect to the weights (no JIT, no patch) def _sum_c2(s): c = s.get_intermediate_values()[0] return jnp.sum(c**2) g_c_tree = jax.grad(_sum_c2)(pinn.space) g_c_max = jnp.max( jnp.abs( jnp.concatenate([leaf.ravel() for leaf in jax.tree_util.tree_leaves(g_c_tree)]) ) ) print( f"max |∂‖c‖²/∂θ| : {g_c_max:.4e} " f"({'OK, the gradient flows through c' if g_c_max > 1e-10 else 'PROBLEM'})" ) # ── Loss history over all runs ──────────────────────────────────────────────── 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 history — all runs") ax_l.legend(fontsize=8) ax_l.grid(True, alpha=0.3) plt.tight_layout() plt.show()