r"""Learn the SIPG penalty constant, kept inside the range that makes it work. -u'' = f on (0, 1), u = 0 at both ends, u_exact = sin(pi x) SIPG needs a penalty ``sigma / h`` on the jump across each face. Too small and the bilinear form loses coercivity, so the scheme is unstable; too large and it over-penalises, degrading both accuracy and conditioning. The usable window is known: ``sigma`` of the order of ``p (p+1)``. So rather than let an optimizer discover that the hard way, the parameterisation itself confines it: sigma(theta) = p(p+1) * (1 + sigmoid(theta)) in ( p(p+1), 2 p(p+1) ) Whatever the optimizer does to ``theta``, ``sigma`` cannot leave that interval. There is no penalty term to balance, no clipping, no run that has to be thrown away because the constant wandered somewhere the scheme is unstable. What is learned is *where inside the window* the constant should sit. **The dynamic field lives in the user's class, not in the library.** The library's :class:`SIPGFlux` declares ``sigma`` static, which is right: it is given data in every ordinary run. Learning it is the inverse-problem case, and scimba's convention is that the user who wants that writes their own class and marks the field dynamic -- here ``theta``, with ``sigma`` derived from it. Overriding ``_penalty`` is the whole of it. """ # %% 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.error_analysis import l2_error from scimba_jax.linear_approximation.galerkin.dg.elliptic_dg_scheme import ( EllipticDGscheme, ) from scimba_jax.linear_approximation.galerkin.dg.flux import SIPGFlux 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 InvertibleFunction, 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.numerical_solvers.projectors import Projector from scimba_jax.physical_models.abstract_physical_weak_model import ( AbstractPhysicalWeakModel, ) from scimba_jax.physical_models.classical_weakform.laplacian_weak_form import ( LaplacianWeakForm, ) from scimba_jax.physical_models.elliptic_pde.laplacians import LaplacianDirichletDG from scimba_jax.physical_models.weak_boundary_conditions import Dirichlet from scimba_jax.utils.scimba_pytree import trainable DIM, OUT_DIM = 1, 1 N_CELLS = 4 # deliberately coarse: with 12 the solution is exact to 1e-10 # and sigma has nothing to trade off. Here the error varies by a factor 400 # across the window. POLY_ORDER = 4 QUAD_ORDER = 6 N_EPOCHS, N_COLLOC = 100, 1200 # The window deliberately straddles the coercivity threshold: its lower end is # well below p(p+1), where SIPG stops being stable, so there is something real # to find. A window entirely inside the safe range leaves sigma almost without # effect -- the error varies by under 2% across it. SIGMA_LO_FACTOR, SIGMA_HI_FACTOR = 0.15, 3.0 THETA_INIT = 0.0 # sigma starts at the middle of the window SEED = 0 def u_exact(x): """Exact solution ``sin(pi x)``.""" return jnp.sin(jnp.pi * x[0]).reshape(OUT_DIM) def f_rhs(x): """Source of ``-u'' = f`` for that solution.""" return (jnp.pi**2 * jnp.sin(jnp.pi * x[0])).reshape(OUT_DIM) # %% The flux whose penalty constant is learned. class LearnableSIPGFlux(SIPGFlux): r"""SIPG whose penalty constant is learned, confined to the usable window. ``sigma = sigma_ref * (1 + sigmoid(theta))``, so it stays strictly between ``sigma_ref`` and ``2 sigma_ref`` for any ``theta`` -- coercivity is a property of the parameterisation rather than something the optimizer has to respect. ``theta`` is the only dynamic field: it is what the optimizer sees, and ``sigma`` is read off from it. Overriding :meth:`_penalty` is enough, since that is the single point where the base class turns its constant into a coefficient. Args: sigma_ref: Lower end of the window, typically ``p (p+1)``. theta_init: Starting value; ``0`` puts sigma in the middle. """ #: Le SEUL parametre du cas. ⚠ Declarer l'EMPLACEMENT (children) ne #: suffit pas, il faut declarer le ROLE : `trainable` est une PORTE, une #: feuille n'est active que si son champ le dit. Sans cette ligne #: l'optimiseur ne voit rien -- n_theta = 0, et l'assemblage de la Gram #: leve "Too few leaves for PyTreeDef; expected 1, got 0". theta: jnp.ndarray = trainable(True) # sigma_lo / sigma_hi ne declarent RIEN, et c'est correct : deux `float` # Python ne sont jamais des feuilles, donc ils partent en aux_data d'eux # memes, et le defaut est gele. Il n'y a qu'un parametre dans ce cas. def __init__(self, sigma_lo: float, sigma_hi: float, theta_init: float = 0.0): # h=None: each face uses its own size, which is what SIPG's sigma/h # asks for on a non-uniform mesh. super().__init__(sigma=0.5 * (sigma_lo + sigma_hi), h=None) self.sigma_lo = float(sigma_lo) self.sigma_hi = float(sigma_hi) self.theta = jnp.array(float(theta_init)) def current_sigma(self) -> jnp.ndarray: """The constant this flux is currently using.""" return self.sigma_lo + (self.sigma_hi - self.sigma_lo) * jax.nn.sigmoid( self.theta ) def _penalty(self, fields): """``sigma(theta) / h_face``. Args: fields: Field values at the face; carries ``"h_face"``. Returns: The penalty coefficient. """ return self.current_sigma() / fields["h_face"] SIGMA_CLASSICAL = POLY_ORDER * (POLY_ORDER + 1) # the textbook value SIGMA_LO = SIGMA_LO_FACTOR * SIGMA_CLASSICAL SIGMA_HI = SIGMA_HI_FACTOR * SIGMA_CLASSICAL # %% The DG space whose flux is the only dynamic thing. def make_space(flux): """A DG approximation space carrying ``flux``.""" mesh = Mesh( dim=DIM, n_cells=(N_CELLS,), ref_quad=UnitSquareTensorized(dim=DIM, order=QUAD_ORDER), mapping=Mapping(mappings=[InvertibleFunction(lambda x: x, lambda y: y)]), ) 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", ) model = AbstractPhysicalWeakModel(dim=DIM) model.weak_forms = {"interior": LaplacianWeakForm(dim=DIM, f=f_rhs)} model.boundary_conditions = {"boundary": Dirichlet(lambda x: jnp.zeros(OUT_DIM))} scheme = EllipticDGscheme( model, VariablesDG(basis=basis, nb_variables=OUT_DIM), flux ) 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(LearnableSIPGFlux(SIGMA_LO, SIGMA_HI, THETA_INIT)) model = LaplacianDirichletDG(main_domain=domain, f_rhs=f_rhs, bc="weak") sampler = TensorizedSampler([DomainSampler(domain)], bc=True) key = jax.random.PRNGKey(SEED) key, sample_dict = sampler.sample(key, N_COLLOC) pinn = Projector(model, space, sampler) flux0 = space.assemblers[0].flux print(f"1D DG Laplacian, {N_CELLS} cells, Q{POLY_ORDER}") print( f" sigma confined to ({SIGMA_LO:.1f}, {SIGMA_HI:.1f}); textbook value {SIGMA_CLASSICAL:.0f}" ) print(f" initial sigma = {float(flux0.current_sigma()):.4f}") print(f" initial physical loss = {pinn.evaluate_loss(space, sample_dict):.6e}") # %% Train: theta is the only parameter the optimizer can move. key, pinn = pinn.project(key, space, N_EPOCHS, N_COLLOC) space_opt = pinn.space flux_opt = space_opt.assemblers[0].flux sigma_learned = float(flux_opt.current_sigma()) print(f" final physical loss = {pinn.best_loss['total']:.6e}") print( f" learned sigma = {sigma_learned:.4f} " f"({(sigma_learned - SIGMA_LO) / (SIGMA_HI - SIGMA_LO):.0%} into the window, " f"{sigma_learned / SIGMA_CLASSICAL:.2f} x the textbook value)" ) # %% What that constant is worth: one solve at the learned value. solved_learned = EllipticDGscheme.solve(space_opt.assemblers[0], max_iter=1) err_learned = float(l2_error(solved_learned, u_exact, relative=True)) print(f" relative L2 error at the learned sigma = {err_learned:.4e}") # %% Read the result. fig, ax = plt.subplots(figsize=(6, 4)) fig.suptitle( f"Learned SIPG penalty — 1D Laplacian, {N_CELLS} cells, Q{POLY_ORDER}, " f"sigma in ({SIGMA_LO:.1f}, {SIGMA_HI:.1f}), textbook {SIGMA_CLASSICAL:.0f}\n" f"learned {sigma_learned:.2f}, relative L2 error {err_learned:.2e}" ) loss_total = jnp.asarray(pinn.losses.losses_history["total"]).reshape(-1) ax.semilogy(np.asarray(loss_total), lw=1.2) ax.set_title("Physical residual during training") ax.set_xlabel("epoch") ax.grid(alpha=0.3) plt.tight_layout() plt.show()