"""A learned DG basis beats the classical one at the same number of DOFs. -u'' = f on (0, 1), u = 0 at both ends, u(x) = sin(2 pi x), solved by SIPG DG on 10 cells of degree 4 (50 DOFs). The local basis is the Taylor basis MULTIPLIED by a small network shared by every cell (a ``PatchwiseParametricBasis``): phi_k(x) = T_k(x) * (2 + 0.1 N(x)), so the DOF count and the sparsity are the classical ones, and only the space changes. The DG solve sits inside the approximation space: the optimiser never sees the DOFs -- those are what the solve returns -- and adjusts the network so that the DG solution minimises the strong PDE residual (``Projector``, the natural-gradient ENG optimiser by default). The exact solution is used only for the errors. The classical reference is the SAME scheme with the plain Taylor basis on the SAME mesh: 50 DOFs against 50 DOFs. The comparison with other problems (convection-dominated, the disk), other learned bases (cellwise, with a learned mapping) and their recorded results live in the benchmark ``benchmarks/benchmarks_jax/dg_learned_basis``, whose row ``laplacian_1d_sin2pix / patchwise`` is this example. Measured 2026-09-25 (80 epochs of ENG, 193 trained parameters, loss 8.9e-04 -> 3.8e-10, about 7 s of training on CPU), relative L2 error: ====================== ========== classical DG, 50 DOFs 7.54e-07 learned DG, initial 7.53e-07 learned DG, trained 9.31e-10 ====================== ========== a gain of about 800x at equal DOFs. ⚠ Every error is computed after the SAME solve the training sees -- the assembled Jacobian and a direct solve. The former version of this file re-solved the final schemes matrix-free (CG, one Newton step, ``tol=1e-6``) and printed 4.4e-05 for the classical scheme and 9.4e-04 for the trained one: a learned basis WORSE than the classical one after its loss fell by six orders of magnitude. Both numbers were Krylov errors, not discretisation errors: the trained basis is worse conditioned and CG stalls on it (still 1.6e-05 at ``tol=1e-12`` and 20 Newton steps), while the direct solve reads 9.3e-10, the same as the DOFs the space itself returns. Note: the former version of this file announced ``sin(4 pi x)`` in its docstring while its code solved ``sin(2 pi x)``; the code is what is kept. """ import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt 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, PatchwiseParametricBasis, ) 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.networks.mlp import MLP 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 jax.config.update("jax_enable_x64", True) DIM = 1 OUT_DIM = 1 SEED = 1 N_CELLS = 10 DEGREE = 4 QUAD_ORDER = 5 NB_BASIS = DEGREE + 1 SIGMA = (DEGREE + 1) * (DEGREE + DIM) / DIM HIDDEN = [12, 12] N_COLLOC = 6000 N_EPOCHS = 80 # The DG solve inside the space: assembled Jacobian and a direct solve (the # space's default, ``matrix_free=False``). The two final schemes are solved the # same way -- see the module docstring for why that matters. SOLVE_KWARGS = {"max_iter": 1, "tol": 1e-6} FINAL_SOLVE_KWARGS = {"matrix_free": False, **SOLVE_KWARGS} # ── Problem: module-level callables, one object each (a stable compile key) ── def identity(x): """The identity mapping of (0, 1).""" return x def f_rhs(x: jnp.ndarray) -> jnp.ndarray: """``f = -u''`` for ``u = sin(2 pi x)``.""" return (4 * jnp.pi**2.0) * jnp.sin(2.0 * jnp.pi * x[0:1]) def u_exact(x: jnp.ndarray) -> jnp.ndarray: """``u = sin(2 pi x)``.""" return jnp.sin(2.0 * jnp.pi * x[0:1]) def dirichlet_bc(x: jnp.ndarray) -> jnp.ndarray: """Homogeneous Dirichlet datum.""" return jnp.zeros(OUT_DIM) def taylor_basis(coords, i, mesh): """The classical Taylor basis of degree ``DEGREE``.""" return local_taylor_basis(coords, i, mesh, order=DEGREE, out_dim=OUT_DIM) def enriched_basis(u, coords, i, mesh): """The Taylor basis times ``2 + 0.1 N(x)``, ``N`` the shared network.""" return taylor_basis(coords, i, mesh) * (2.0 + 0.1 * u[0]) class BasisNN(MLP): """The network shared by every cell.""" def __init__(self, key): super().__init__(in_size=DIM, out_size=1, hidden_sizes=HIDDEN, key=key) # ── Discretisation ────────────────────────────────────────────────────────── mesh = Mesh( dim=DIM, n_cells=(N_CELLS,), ref_quad=UnitSquareTensorized(dim=DIM, order=QUAD_ORDER), mapping=Mapping(mappings=[InvertibleFunction(identity, identity)]), ) model_weak = AbstractPhysicalWeakModel.from_weak_form( LaplacianWeakForm(dim=DIM, f=f_rhs), dirichlet=dirichlet_bc ) flux = SIPGFlux(sigma=SIGMA, h=mesh.h) # Classical DG, the same mesh and flux: the reference at equal DOFs. classical_scheme = EllipticDGscheme( model_weak, VariablesDG( basis=AnalyticBasis( nb_basis=NB_BASIS, out_dim=OUT_DIM, mesh=mesh, local_basis=taylor_basis, basis_type="scalar", ), nb_variables=OUT_DIM, ), flux, ) started = time.perf_counter() solved_classical = EllipticDGscheme.solve(classical_scheme, **FINAL_SOLVE_KWARGS) solved_classical.variables.dofsl.block_until_ready() t_classical = time.perf_counter() - started # Learned DG: the same scheme on the enriched basis. key = jax.random.PRNGKey(SEED) key, subkey = jax.random.split(key) basis_nn = BasisNN(key=subkey) learned_scheme = EllipticDGscheme( model_weak, VariablesDG( basis=PatchwiseParametricBasis( nb_basis=NB_BASIS, out_dim=OUT_DIM, mesh=mesh, patchwise_parametric_function=basis_nn, local_basis=enriched_basis, basis_type="scalar", ), nb_variables=OUT_DIM, ), flux, ) space = DGEllipticApproximationSpace( dims={"x": DIM, "dofsl": 1}, list_assemblers=[learned_scheme], model_type="x_dofsl", newton_kwargs=SOLVE_KWARGS, ) # ── Training: the strong residual of the DG solution ───────────────────────── domain = Segment1D((0.0, 1.0), is_main_domain=True) model = LaplacianDirichletDG(main_domain=domain, f_rhs=f_rhs, bc="weak") sampler = TensorizedSampler([DomainSampler(domain)], bc=True) key, sample_dict = sampler.sample(key, N_COLLOC) projector = Projector(model, space, sampler) # The acceptance criterion of a learned model: the optimiser moves the declared # network and nothing else -- not the mesh, not the flux, not the DOFs. n_theta = int(projector.optimizer.n_theta) n_declared = sum(leaf.size for leaf in jax.tree_util.tree_leaves(basis_nn)) print(f"trained parameters n_theta = {n_theta}, network weights = {n_declared}") assert n_theta == n_declared, "a frozen array leaked into the trained parameters" print(f"initial loss: {projector.evaluate_loss(space, sample_dict):.6e}") print(f"training ({N_EPOCHS} epochs) ...") started = time.perf_counter() key, projector = projector.project(key, space, N_EPOCHS, N_COLLOC) jax.block_until_ready(jax.tree_util.tree_leaves(projector.best_loss)) t_learned = time.perf_counter() - started trained_space = projector.space print(f"final loss: {projector.best_loss['total']:.6e}") # ── Errors ─────────────────────────────────────────────────────────────────── solved_learned = EllipticDGscheme.solve( trained_space.assemblers[0], **FINAL_SOLVE_KWARGS ) solved_initial = EllipticDGscheme.solve(learned_scheme, **FINAL_SOLVE_KWARGS) l2_classical = float(l2_error(solved_classical, u_exact, relative=True)) l2_initial = float(l2_error(solved_initial, u_exact, relative=True)) l2_learned = float(l2_error(solved_learned, u_exact, relative=True)) n_dofs = N_CELLS * NB_BASIS print() print(f"{'':22s} {'rel. L2 error':>14s} {'time (s)':>10s} {'DOFs':>6s}") print(f"{'classical DG':22s} {l2_classical:>14.4e} {t_classical:>10.3f} {n_dofs:>6d}") print(f"{'learned DG (initial)':22s} {l2_initial:>14.4e} {'':>10s} {n_dofs:>6d}") print( f"{'learned DG (trained)':22s} {l2_learned:>14.4e} {t_learned:>10.2f} {n_dofs:>6d}" ) print(f"gain at equal DOFs: {l2_classical / l2_learned:.1f}x") # ── Plots ──────────────────────────────────────────────────────────────────── x_plot = jnp.linspace(0.0, 1.0, 300)[:, None] u_ref = jax.vmap(u_exact)(x_plot)[:, 0] u_classical = solved_classical.variables.evaluate(x_plot)[:, 0] u_initial = solved_initial.variables.evaluate(x_plot)[:, 0] u_learned = solved_learned.variables.evaluate(x_plot)[:, 0] fig, axs = plt.subplots(1, 3, figsize=(17, 4.8)) axs[0].plot(x_plot[:, 0], u_ref, "k--", linewidth=2, label="exact") axs[0].plot(x_plot[:, 0], u_classical, label="classical DG") axs[0].plot(x_plot[:, 0], u_initial, ":", alpha=0.6, label="learned DG (initial)") axs[0].plot(x_plot[:, 0], u_learned, label="learned DG (trained)") axs[0].set_xlabel("x") axs[0].set_ylabel("u(x)") axs[0].set_title(f"Solutions, {N_CELLS} cells x {NB_BASIS} DOFs") axs[0].legend(fontsize=8) axs[0].grid(True, alpha=0.3) loss_total = jnp.asarray(projector.losses.losses_history["total"]).reshape(-1) axs[1].semilogy(loss_total, "o-", markersize=3) axs[1].set_xlabel("epoch") axs[1].set_ylabel("strong residual (loss)") axs[1].set_title(f"Training, ENG, {n_theta} parameters") axs[1].grid(True, alpha=0.3) axs[2].semilogy( x_plot[:, 0], jnp.abs(u_classical - u_ref), label=f"classical DG (L2 rel. {l2_classical:.2e})", ) axs[2].semilogy( x_plot[:, 0], jnp.abs(u_learned - u_ref), label=f"learned DG (L2 rel. {l2_learned:.2e})", ) axs[2].set_xlabel("x") axs[2].set_ylabel(r"$|u_h - u|$") axs[2].set_title("Pointwise error, equal DOFs") axs[2].legend(fontsize=8) axs[2].grid(True, alpha=0.3) plt.tight_layout() plt.show()