r"""Learn a 2D mesh mapping that refines where the boundary layer is. -eps Laplace(u) + du/dx = 1 on (0,1)^2 u = 0 on the whole boundary (Dirichlet) The 2D twin of ``dgelliptic_mapping_1d.py``: same equation, same exponential mapping, applied to x only (y is the identity), so ONE parameter ``s_x``. Advection points along x: the layer of the 1D problem sits at x = 1. With u = 0 on the whole boundary the solution is NOT the 1D profile extended in y -- it also has thin layers along y = 0 and y = 1, which a stretch of x cannot resolve and which stay in the loss whatever ``s_x``. As in 1D, nothing supervises the mapping: the loss is the physical residual of the PDE, and the mesh mapping is the only dynamic object in the space. What was measured (2026-09-25) -- read this before changing a constant ------------------------------------------------------------------------ The question was "why does the 2D learn slowly?". It is not a matter of epochs. 1. Loss against ``s_x`` (the Projector's own loss, 3 collocation draws averaged, the constants below): global minimum at ``s_x ~ -4.25`` (9.1e-3); training ends at ``s_x = -3.12`` (per-epoch loss 7e-2 to 9e-2). The floor itself is fine: the 1D case at the same 10 cells in x bottoms out at ~6e-3 (s = -4.5). So the 2D does NOT sit at its optimum. 2. The landscape has a POLE at ``s_x ~ -2.45`` (loss 1.6e5; 31 at -2.5 and at -2.4) and a local minimum at ``s_x ~ -1.5`` (0.62) in front of it. It is a loss of coercivity of SIPG: the penalty is ``sigma / h`` with ``h = |K|^(1/d) = sqrt(hx hy)``, too large a length for the thin cells the mapping makes near x = 1. With ``2 sigma`` the pole moves to about -4.25; with ``4 sigma`` it is gone and the landscape is monotone like the 1D one. In 1D ``h`` is the cell width and there is no pole. 3. The main cause: the gradient points the WRONG WAY from ``s_x ~ -3.05`` on. It is exact (it matches finite differences at ds = 1e-4), but for a pointwise residual of a DISCONTINUOUS DG function it sees only the points that stay in their cell. Over ds = -0.05 at s_x = -3.1: the points that stay change the loss by +8.7e-4, the 45 (of 2000) that change cell by -1.6e-2. In 1D it is the other way round (-3.6e-3 against -2.3e-3). So ENG pushes ``s_x`` back towards the pole and the line search, which only shortens the step, keeps it at ``s_x ~ -3.1``: that is the plateau of the loss curve. More cells in y (10), 12 cells in x or an even quadrature do not change the sign; with ``4 sigma`` alone the pole is gone but the wrong gradient then walks ``s_x`` to +4.6 in 80 steps (prototype run), so do not raise sigma alone. 4. What fixes it, measured on a prototype (not in the library): collocation points attached to the mesh -- drawn in the logical square, mapped by the CURRENT mapping and weighted by ``|det D phi|``, an unbiased estimate of the same integral. No point ever changes cell, the gradient becomes consistent (stationary point at ``s_x ~ -4.3``, the scan's minimum), and with ``4 sigma`` the prototype training reaches ``s_x = -4.4`` and loss 9.3e-3 in 4 steps (then drifts along a flat valley, ``s_x ~ -13.7`` and loss ~1e-2 after 30 steps). Without the sigma change it stays at the pole-bounded local minimum (~ -1.6). So the two changes belong to the library (the sampler and the SIPG length scale), not to the constants of this example. Compile and run time, same date: the training step used to compile twice (``s`` weakly typed at first, see ``RefinementMap2D.__init__``), and the post-processing solve used matrix-free GMRES, which ran 79 s WITHOUT converging (20 Newton iterations, residual 3e-5, 90 % away from the solution the loss sees). It now uses the space's own direct solve (3.7 s, residual 7e-15). Whole run: 166.5 s -> 63.9 s, same final loss (3.538604e-02) and ``s_x``. What remains: one compilation of the training step (~20 s: trace 6.4 s, XLA 11.4 s), 80 steps (~31 s), the initial-loss evaluation (~5 s). """ # %% import math import jax import jax.numpy as jnp 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 ( SIPGFlux, SumFlux, UpwindFlux, ) 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 Mapping from scimba_jax.nonlinear_approximation.approximation_spaces.dg_approximation_spaces import ( # noqa: E501 DGEllipticApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.networks.structure_preserving_nets.invertible_nn import ( # noqa: E501 InvertibleNet, ) from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector from scimba_jax.physical_models.abstract_physical_model import AbstractPhysicalModel from scimba_jax.physical_models.abstract_physical_weak_model import ( AbstractPhysicalWeakModel, ) from scimba_jax.physical_models.abstract_residuals import InteriorResidual from scimba_jax.physical_models.boundary_residuals import DirichletResidual from scimba_jax.physical_models.classical_weakform.diffusion_advection_reaction_weak_form import ( # noqa: E501 EllipticWeakForm, ) from scimba_jax.physical_models.weak_boundary_conditions import Dirichlet from scimba_jax.plots.plots_galerkin import plot_solution_2d from scimba_jax.utils.scimba_pytree import trainable jax.config.update("jax_enable_x64", True) EPS = 0.02 N_CELLS_X = 10 # cells per direction, set independently. Neither count N_CELLS_Y = 5 # explains the plateau: see point 3 of the module docstring. POLY_ORDER = 3 # the PINN residual needs second derivatives QUAD_ORDER = 5 DIM, OUT_DIM = 2, 1 S_INIT = -0.8 # x starts near the identity (s = 0 is exactly it) N_EPOCHS, N_COLLOC = 80, 2000 SEED = 0 ADVECTION = jnp.array([1.0, 0.0]) # along x only: nothing to resolve in y def layer_profile(x): """The 1D profile of the same equation, for orientation only. With u = 0 on the whole boundary this is **not** the exact solution -- it does not vanish on south and north. It is drawn only to show where the x-layer sits, since the advection and the x-conditions are the same. """ return jnp.array([x[0] - (jnp.exp(x[0] / EPS) - 1.0) / (math.exp(1.0 / EPS) - 1.0)]) # %% The mapping: one parameter, stretching x only. class RefinementMap2D(InvertibleNet): r"""Exponential stretching of ``x`` alone; ``y`` is left untouched. phi(xi, eta) = ( (e^{s xi} - 1)/(e^s - 1), eta ) The advection points along x and the boundary conditions do not depend on y, so there is nothing in this problem that a y-stretch could resolve. One parameter therefore says everything, and the answer reads off its sign: negative packs cells at x = 1, where the layer is. Giving the map a second parameter only let it compress y, which it did (s_y = -1.065) while leaving x alone -- resolution spent where the solution is constant. Removing the freedom removes the distraction. The properties are the 1D ones, unchanged: each edge maps to itself so the boundary is fixed; ``phi' > 0`` for every ``s`` so the mesh cannot fold; ``s = 0`` is exactly the identity; the inverse is Newton, which converges because ``phi' > 0``. """ #: The ONLY parameter of the case. `trainable` is a DOOR: a leaf is #: active only if its field declares it, so without this line the #: mapping does not train (n_theta = 0). s: jnp.ndarray = trainable(True) # scalar: the x-stretch def __init__(self, s_init: float = S_INIT): super().__init__(size=DIM, conditional_size=0, layers_list=[]) # An explicit dtype, not `jnp.array(s_init)`: from a Python float that # gives a WEAKLY typed array, while the optimizer hands back a strongly # typed one after the first step. The two are different jit keys, so # the whole training step compiled twice (measured 2026-09-25 with # JAX_EXPLAIN_CACHE_MISSES: 1D 27.7 s -> see the module docstring). self.s = jnp.asarray(s_init, dtype=float) def _s_safe(self) -> jnp.ndarray: """``s`` kept away from 0, where the ratio would be 0/0.""" return jnp.where(jnp.abs(self.s) < 1e-6, 1.0, self.s) def _stretch(self, t: jnp.ndarray) -> jnp.ndarray: """The 1D stretch applied to one coordinate.""" s = self._s_safe() return jnp.where(jnp.abs(self.s) < 1e-6, t, jnp.expm1(s * t) / jnp.expm1(s)) def _stretch_deriv(self, t: jnp.ndarray) -> jnp.ndarray: """Its derivative, the local cell-size factor along x.""" s = self._s_safe() d = s * jnp.exp(s * t) / jnp.expm1(s) return jnp.where(jnp.abs(self.s) < 1e-6, jnp.ones_like(d), d) def __call__(self, xi: jnp.ndarray) -> jnp.ndarray: """Logical -> physical: stretch x, keep y. Args: xi: Logical point, shape ``(..., 2)``. Returns: The physical point, same shape. """ xi = jnp.atleast_1d(xi) return jnp.stack([self._stretch(xi[..., 0]), xi[..., 1]], axis=-1) def backward(self, x: jnp.ndarray) -> jnp.ndarray: """Physical -> logical, by Newton on the x component alone. Args: x: Physical point, shape ``(..., 2)``. Returns: The logical point, same shape. """ x = jnp.atleast_1d(x) def newton_step(t, _): return t - (self._stretch(t) - x[..., 0]) / self._stretch_deriv(t), None t, _ = jax.lax.scan(newton_step, x[..., 0], None, length=8) return jnp.stack([t, x[..., 1]], axis=-1) # %% The physics, as a PINN residual: this is what is minimized. class AdvectionDiffusionResidual2D(InteriorResidual): """``-eps Laplace(u) + b . grad u - 1``, written in ParamFunction.""" def __init__(self, domain, model_type="x_dofsl"): super().__init__( domain=domain, size=1, model_type=model_type, f_rhs=lambda x: jnp.ones(1), ) def construct_residual(self, *vars): u = vars[0] return -EPS * u.laplacian("x") + u.gradient("x").dot(ADVECTION) class AdvectionDiffusionDG2D(AbstractPhysicalModel): """Interior residual plus a weak Dirichlet residual on each side.""" def __init__(self, main_domain): super().__init__(main_domain=main_domain) self.physical_residuals = { self.main_domain.get_label(): AdvectionDiffusionResidual2D(main_domain), } for label, boundary in self.boundaries.items(): self.physical_residuals[label] = DirichletResidual( domain=boundary, model_type="x_dofsl" ) # %% The DG space whose mapping is the only dynamic thing. def make_space(mapping_module): """A DG approximation space on a mesh carrying ``mapping_module``.""" mesh = Mesh( dim=DIM, n_cells=(N_CELLS_X, N_CELLS_Y), ref_quad=UnitSquareTensorized(dim=DIM, order=QUAD_ORDER), mapping=Mapping(mappings=[mapping_module]), # Without this the mesh takes its Cartesian fast path and the mapping is # never applied -- the parameters would train against nothing. is_identity_mapping=False, ) basis = AnalyticBasis( nb_basis=(POLY_ORDER + 1) ** DIM, out_dim=OUT_DIM, mesh=mesh, local_basis=lambda c, i, m: local_taylor_basis( c, i, m, order=POLY_ORDER, out_dim=OUT_DIM ), basis_type="scalar", ) weak_form = EllipticWeakForm( dim=DIM, A=lambda x: EPS * jnp.eye(DIM), b=lambda x: ADVECTION, c=lambda x: jnp.zeros(()), f=lambda x: jnp.ones(OUT_DIM), ) model = AbstractPhysicalWeakModel(dim=DIM) model.weak_forms = {"interior": weak_form} # Homogeneous Dirichlet on the whole boundary, exactly matching the PINN # model. The meshless Square2D groups its four sides under a single # "boundary" label, so anything finer here would have the loss ask for a # condition the scheme does not impose and the two would be solving # different problems. model.boundary_conditions = {"boundary": Dirichlet(lambda x: jnp.zeros(OUT_DIM))} scheme = EllipticDGscheme( model, VariablesDG(basis=basis, nb_variables=OUT_DIM), SumFlux( [ SIPGFlux(sigma=POLY_ORDER * (POLY_ORDER + 1) * DIM, h=None), UpwindFlux(), ] ), ) return DGEllipticApproximationSpace( dims={"x": DIM, "dofsl": 1}, list_assemblers=[scheme], model_type="x_dofsl", newton_kwargs={"max_iter": 1, "tol": 1e-7}, ) domain = Square2D([[0.0, 1.0], [0.0, 1.0]], is_main_domain=True) space = make_space(RefinementMap2D()) model = AdvectionDiffusionDG2D(main_domain=domain) sampler = TensorizedSampler([DomainSampler(domain)], bc=True) key = jax.random.PRNGKey(SEED) key, sample_dict = sampler.sample(key, N_COLLOC) # Same optimizer as the 1D example (the default, ENG). ⚠ Measured 2026-09-25: # the landscape is NOT smooth here (a pole at s_x ~ -2.45) and past s_x ~ -3 # the gradient points away from the minimum (module docstring, points 2-3), # which any gradient method -- Adam included -- follows. pinn = Projector(model, space, sampler) print( f"2D DG, eps = {EPS}, {N_CELLS_X}x{N_CELLS_Y}, Q{POLY_ORDER}, " f"1 mapping parameter (x only)" ) print(f" initial physical loss: {pinn.evaluate_loss(space, sample_dict):.6e}") # %% Train. def interfaces(sp): """Physical cell interfaces along each direction.""" phi = sp.assemblers[0].variables.mesh.mapping.mappings[0] # One evaluation per direction rather than the diagonal shortcut: the # counts differ, and the shortcut would also silently assume the map is # separable, which a future non-tensor mapping would not be. tx = jnp.linspace(0.0, 1.0, N_CELLS_X + 1) ty = jnp.linspace(0.0, 1.0, N_CELLS_Y + 1) xs = jax.vmap(lambda a: phi(jnp.stack([a, jnp.zeros(())]))[0])(tx) ys = jax.vmap(lambda b: phi(jnp.stack([jnp.zeros(()), b]))[1])(ty) return np.asarray(xs), np.asarray(ys) x_before, y_before = interfaces(space) key, pinn = pinn.project(key, space, N_EPOCHS, N_COLLOC) space_opt = pinn.space x_after, y_after = interfaces(space_opt) s_learned = float( jnp.ravel(space_opt.assemblers[0].variables.mesh.mapping.mappings[0].s)[0] ) wx, wy = np.diff(x_after), np.diff(y_after) print(f" final physical loss : {pinn.best_loss['total']:.6e}") print(f" learned s_x = {s_learned:+.3f} (y is the identity)") print(f" cell width ratio: x {wx.max() / wx.min():.2f}, y {wy.max() / wy.min():.2f}") print( " expected: s_x clearly negative (layer at x=1); the loss scan puts the " "minimum near s_x = -4.25, training stops near -3.1 (module docstring)" ) # %% Read the result off the mesh. fig, ax = plt.subplots(1, 3, figsize=(15, 4.4)) fig.suptitle( f"Learned 2D mapping — advection along x only, layer at x=1, " f"{N_CELLS_X}x{N_CELLS_Y}, Q{POLY_ORDER}" ) xs = np.linspace(0.0, 1.0, 400) u_ref = np.asarray(jax.vmap(layer_profile)(np.stack([xs, np.zeros_like(xs)], axis=-1)))[ :, 0 ] for k, xn in enumerate(x_after): ax[0].axvline( xn, color="crimson", lw=1.0, alpha=0.85, label="learned x-interfaces" if k == 0 else None, ) for k, xn in enumerate(x_before): ax[0].axvline( xn, color="gray", lw=0.8, ls=":", alpha=0.7, label=f"initial (s={S_INIT})" if k == 0 else None, ) ax[0].plot(xs, u_ref, "k-", lw=1.8, label="1D layer profile (orientation)") ax[0].set_title("x-interfaces against the solution") ax[0].set_xlabel("x") ax[0].set_ylabel("u") ax[0].legend(fontsize=7, loc="lower left") # The learned mesh itself: the horizontal lines should stay evenly spaced. for xn in x_after: ax[1].plot([xn, xn], [0, 1], color="crimson", lw=0.9) for yn in y_after: ax[1].plot([0, 1], [yn, yn], color="steelblue", lw=0.9) ax[1].set_title("Learned mesh: x refined, y left alone") ax[1].set_xlabel("x") ax[1].set_ylabel("y") ax[1].set_aspect("equal") loss_total = jnp.asarray(pinn.losses.losses_history["total"]).reshape(-1) ax[2].semilogy(np.asarray(loss_total), lw=1.2) ax[2].set_title("Physical residual during training") ax[2].set_xlabel("epoch") ax[2].grid(alpha=0.3) plt.tight_layout() plt.show() # %% The solution itself, on the learned mesh: u_h, the exact profile, and the # error -- which should sit in a band at x = 1, exactly where the layer is and # exactly what the mapping was supposed to go and resolve. # The space solves the DG system on the fly inside the loss and does not keep # the result, so it has to be solved once here or the plot shows zeros. # ⚠ Solved with the SPACE'S OWN settings (direct Newton), i.e. exactly the # solution the loss saw. The matrix-free default is a CG, wrong for this # non-symmetric system (Pe = 50), and GMRES matrix-free was no better here: # measured 2026-09-25 at s_x = -3.116, 79 s, not converged, 90 % away from the # direct solution (3.7 s, residual 7e-15). assembler = space_opt.assemblers[0] solved = type(assembler).solve( assembler, matrix_free=space_opt.matrix_free, **space_opt.newton_kwargs ) plot_solution_2d( solved, title=f"DG solution on the learned mesh (s_x = {s_learned:+.3f}) " "-- homogeneous Dirichlet all around", )