r"""Learn a 1D mesh mapping that refines where the boundary layer is. -eps u'' + u' = 1 on (0, 1), u(0) = u(1) = 0 The reduced problem (eps = 0) is u' = 1, so u = x, which satisfies the left condition and misses the right one: the layer sits at **x = 1** and is O(eps) wide. On a uniform mesh a coarse DG space cannot see it. The mapping's job is to move the cells there, without being told where "there" is. Nothing supervises the mapping. The loss is the **physical residual** of the PDE, exactly as in ordinary PINN training; the only difference is what carries the parameters -- here the mesh mapping rather than a basis or a network. The mapping is what the optimizer moves, and the layer is what the residual complains about, so refinement is a consequence rather than an instruction. The mapping, ``phi : [0,1] -> [0,1]``, is parameterized by hand rather than by a network, so what it learned can be read off. It has ONE parameter ``s``: phi(xi) = (e^{s xi} - 1) / (e^s - 1) - phi(0) = 0 and phi(1) = 1 for every ``s``, so **the boundary points do not move**; - ``phi' > 0`` for every ``s``, so phi is **monotone whatever the parameter**: the optimizer cannot fold the mesh, and there is no positivity penalty; - ``s = 0`` is the identity, ``s < 0`` packs cells at x = 1; - the inverse is a **Newton** solve, which converges because phi' > 0. Reading the result: cells are packed where ``phi'`` is *small*. With the layer at x = 1 we expect phi' small on the right and large on the left -- the mesh stretched over the smooth part and compressed into the layer. What was measured (2026-09-25), to compare with ``dgelliptic_mapping_2d.py``: - The loss, scanned over ``s`` (3 collocation draws averaged), decreases monotonically down to ``s = -5`` (2.07 at s = 0, 5.6e-3 at s = -3.75, 3.0e-3 at s = -5). Training settles near ``s = -3.85`` (it oscillates within [-4.05, -3.6]); the printed "final loss" 2.19e-3 is the smallest of 300 noisy per-epoch values, not the loss at the final ``s``. - Why it stops before the scan's minimum: the loss is a pointwise residual of a DISCONTINUOUS (DG) function, evaluated at fixed collocation points. When the mapping moves an interface across a point, that point's residual jumps; the autodiff gradient sees only the smooth part (it matches finite differences at ds = 1e-4), not those jumps. Over ds = -0.05 at s = -3.1, points that stay in their cell carry -3.6e-3 of the change and the 26 points that change cell -2.3e-3: here the smooth part dominates and has the right sign, so training works -- until about s = -4, where the smooth part changes sign. - Compile: ``s`` used to be built from a Python float (weakly typed), so the training step compiled twice; with an explicit dtype the run drops from 27.7 s to 14.4 s with the same final loss (2.193997e-03). What remains is one compilation of the training step (~8 s: trace 3.9 s, XLA 3.3 s) and one of the initial-loss evaluation (~4 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_1d import Segment1D 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 ( DGEllipticApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.networks.structure_preserving_nets.invertible_nn import ( 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.utils.scimba_pytree import trainable EPS = 0.02 # layer width; smaller makes the mapping matter more N_CELLS = 12 # deliberately coarse: uniform cells cannot resolve the layer POLY_ORDER = 3 # the PINN residual needs u''; Q1 cannot represent it QUAD_ORDER = 6 # chosen EVEN when an odd Gauss rule (a point on the cell # centre) gave a nan mapping gradient through the Taylor basis. No longer the # case since the basis uses repeated multiplications: measured 2026-09-25 with # QUAD_ORDER = 5, the gradient is finite and matches finite differences. DIM, OUT_DIM = 1, 1 S_INIT = 0.3 # mapping starts near the identity (s = 0 is exactly it) N_EPOCHS, N_COLLOC = 300, 800 SEED = 0 def u_exact(x): """Exact solution ``x - (e^{x/eps} - 1) / (e^{1/eps} - 1)``.""" return jnp.array([x[0] - (jnp.exp(x[0] / EPS) - 1.0) / (math.exp(1.0 / EPS) - 1.0)]) # %% The mapping: one parameter, positive derivative for free. class RefinementMap1D(InvertibleNet): r"""Exponential stretching ``phi(xi) = (e^{s xi} - 1) / (e^s - 1)`` on [0, 1]. One parameter, and the properties come for free rather than by constraint: - ``phi(0) = 0`` and ``phi(1) = 1`` identically, so **the boundary points never move**, whatever ``s``; - ``phi'(xi) = s e^{s xi} / (e^s - 1) > 0`` for **every** ``s``, so the Jacobian can neither vanish nor change sign -- nothing to clamp, no penalty term, no way for the optimizer to fold the mesh. A sum of sine modes would need its amplitudes squashed to stay monotone; this does not; - ``s -> 0`` is the identity, so training starts from the uniform mesh. It is also the family one would pick by hand for a boundary layer: ``s < 0`` packs cells near ``x = 1``, ``s > 0`` near ``x = 0``. That makes the learned value readable -- the sign alone says which end got refined. The inverse is analytic here, but is taken by **Newton** anyway: that is the mechanism a mapping without a closed-form inverse would need, and ``phi' > 0`` is exactly what makes it converge. """ #: 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) def __init__(self, s_init: float = S_INIT): super().__init__(size=1, 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 is 0/0. Returns: ``s`` if it is not tiny, otherwise a harmless stand-in; the caller selects the linear branch in that case. """ return jnp.where(jnp.abs(self.s) < 1e-6, 1.0, self.s) def __call__(self, xi: jnp.ndarray) -> jnp.ndarray: """Logical -> physical. Args: xi: Logical coordinate, shape ``(1,)``. Returns: The physical coordinate, shape ``(1,)``. """ xi = jnp.atleast_1d(xi) s = self._s_safe() stretched = jnp.expm1(s * xi) / jnp.expm1(s) return jnp.where(jnp.abs(self.s) < 1e-6, xi, stretched) def derivative(self, xi: jnp.ndarray) -> jnp.ndarray: """``phi'(xi)``, the local cell-size factor. Args: xi: Logical coordinate (scalar). Returns: ``phi'(xi)``; small means cells are packed there. """ s = self._s_safe() d = s * jnp.exp(s * xi) / jnp.expm1(s) return jnp.where(jnp.abs(self.s) < 1e-6, jnp.ones_like(d), d) def backward(self, x: jnp.ndarray) -> jnp.ndarray: """Physical -> logical, by Newton. Args: x: Physical coordinate, shape ``(1,)``. Returns: The logical coordinate, shape ``(1,)``. """ x = jnp.atleast_1d(x) def newton_step(xi, _): return xi - (self(xi) - x) / self.derivative(xi[0]), None xi, _ = jax.lax.scan(newton_step, x, None, length=8) return xi # %% The physics, as a PINN residual: this is what is minimized. class AdvectionDiffusionResidual(InteriorResidual): """``-eps u'' + 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(jnp.ones(1)) class AdvectionDiffusionDG(AbstractPhysicalModel): """Interior residual plus a weak Dirichlet residual on each end.""" def __init__(self, main_domain): super().__init__(main_domain=main_domain) self.physical_residuals = { self.main_domain.get_label(): AdvectionDiffusionResidual(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,), 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, 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: jnp.ones(DIM), c=lambda x: jnp.zeros(()), f=lambda x: jnp.ones(OUT_DIM), ) model = AbstractPhysicalWeakModel(dim=DIM) model.weak_forms = {"interior": weak_form} model.boundary_conditions = {"boundary": Dirichlet(lambda x: jnp.zeros(OUT_DIM))} scheme = EllipticDGscheme( model, VariablesDG(basis=basis, nb_variables=OUT_DIM), # h=None: each face uses its own size, which is the point once a # mapping makes the cells differ by a factor of tens. UpwindFlux # stabilizes the advective term, which SIPG does not touch. 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-12}, ) domain = Segment1D([0.0, 1.0], is_main_domain=True) space = make_space(RefinementMap1D()) model = AdvectionDiffusionDG(main_domain=domain) sampler = TensorizedSampler([DomainSampler(domain)], bc=True) key = jax.random.PRNGKey(SEED) key, sample_dict = sampler.sample(key, N_COLLOC) pinn = Projector(model, space, sampler) print(f"1D DG, eps = {EPS}, {N_CELLS} cells, Q{POLY_ORDER}, 1 mapping parameter") print(f" initial physical loss: {pinn.evaluate_loss(space, sample_dict):.6e}") # %% Train: the residual is the only signal, the mapping the only unknown moved. def node_positions(sp): """Physical positions of the cell interfaces of a space's mesh.""" phi = sp.assemblers[0].variables.mesh.mapping.mappings[0] xi = jnp.linspace(0.0, 1.0, N_CELLS + 1) return np.asarray(jax.vmap(lambda t: phi(t[None])[0])(xi)) nodes_before = node_positions(space) key, pinn = pinn.project(key, space, N_EPOCHS, N_COLLOC) space_opt = pinn.space nodes_after = node_positions(space_opt) print(f" final physical loss : {pinn.best_loss['total']:.6e}") smallest = np.diff(nodes_after).argmin() print( f" smallest cell is #{smallest} of {N_CELLS}, centred at " f"x = {0.5 * (nodes_after[smallest] + nodes_after[smallest + 1]):.3f} " f"(the layer sits at x = 1)" ) # %% Read the result off the mesh. fig, ax = plt.subplots(1, 3, figsize=(15, 4)) fig.suptitle( f"Learned 1D mapping, $-{EPS}\\,u'' + u' = 1$ — physical residual only, " f"{N_CELLS} cells" ) xs = np.linspace(0.0, 1.0, 800) u_ref = np.asarray(jax.vmap(u_exact)(xs[:, None]))[:, 0] # The cell interfaces themselves, drawn through the solution: where they crowd # is where the mapping decided to spend resolution, and it should be the layer. for k, xn in enumerate(nodes_after): ax[0].axvline( xn, color="crimson", lw=1.0, alpha=0.8, label="learned cell interfaces" if k == 0 else None, ) for k, xn in enumerate(nodes_before): ax[0].axvline( xn, color="gray", lw=0.8, ls=":", alpha=0.7, label=f"initial interfaces (s={S_INIT})" if k == 0 else None, ) ax[0].plot(xs, u_ref, "k-", lw=1.8, label="exact solution") ax[0].plot( nodes_after, np.interp(nodes_after, xs, u_ref), "o", color="crimson", ms=4.5, zorder=5, ) ax[0].set_title("Cell interfaces against the solution") ax[0].set_xlabel("x") ax[0].set_ylabel("u") ax[0].legend(fontsize=7, loc="lower left") ax[1].semilogy(np.diff(nodes_before), "o-", ms=4, label=f"initial (s={S_INIT})") ax[1].semilogy(np.diff(nodes_after), "s-", ms=4, label="learned") ax[1].set_title("Cell width (log): small = refined") ax[1].set_xlabel("cell index (left to right)") ax[1].legend(fontsize=8) ax[1].grid(alpha=0.3) 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()