"""Learning c(x) in -Delta u + c(x) u = f(x) through a differentiable DG solve. Problem: -Delta u + c(x) u = f(x) on [0, 1], u(0) = u(1) = 2 Exact solution: u(x) = 2 + sin(2 pi x) True coefficient: c(x) = 10 x (1 - x) Static source: f(x) = 4 pi^2 sin(2 pi x) + 10 x (1 - x) (2 + sin(2 pi x)) Approach: - fixed mesh and DG-Q1 Lagrange basis; - c(x) modelled by an ApproximationSpace (an MLP): the parameter to learn; - DGEllipticApproximationSpace solves the DG system at every step; c_theta is a pytree child of the scheme, so the gradient goes through the implicit differentiation of the solve; - loss = data fitting: ||u_h(c_theta, x_i) - u_obs(x_i)||^2 over N_OBS measurements; - optimisation by ``Projector`` with its default optimiser, ENG (natural gradient), the same pattern as ``pinns/inverse_problems/helmholtz_inverse_nu.py``. Measured 2026-09-25: data loss 6.24e-10 after 100 epochs, RMS error 8.5e-04 on u and 5.5e-02 on c (about 9 s in all). The same inverse problem on the 2D disk, through four discretisations (DG / FEM, unstructured / block-structured), is the benchopt benchmark ``benchmarks/benchmarks_jax/inverse_reaction/``. """ 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_lagrange_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 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.approximation_spaces import ( ApproximationSpace, ) 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.mlp import MLP 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.classical_weakform.laplacian_inverse_linear_problem import ( LaplacianReactionWeakFormLearnableCoeff, ) from scimba_jax.physical_models.data_residuals import CollocDataResidual jax.config.update("jax_enable_x64", True) # ── Parameters ──────────────────────────────────────────────────────────────── physical_dim = 1 out_dim = 1 n_cells = 40 poly_deg = 1 nb_basis = (poly_deg + 1) ** physical_dim # Q1: 2 nodes per cell in 1D quad_order = 3 sigma = poly_deg * (poly_deg + 1) seed = 42 N_EPOCHS = 100 N_OBS = 40 # measurement points of u # ── Analytic problem ────────────────────────────────────────────────────────── def u_exact(x): return 2.0 + jnp.sin(2.0 * jnp.pi * x[0]) def c_exact(x): return 10.0 * x[0] * (1.0 - x[0]) def f_static(x): # f = -Delta u + c u with u = 2 + sin(2 pi x), c = 10 x (1 - x) # -Delta u = 4 pi^2 sin(2 pi x), c u = 10 x (1 - x) (2 + sin(2 pi x)) return 4.0 * jnp.pi**2 * jnp.sin(2.0 * jnp.pi * x[0]) + 10.0 * x[0] * ( 1.0 - x[0] ) * (2.0 + jnp.sin(2.0 * jnp.pi * x[0])) def dirichlet_bc(_x): return jnp.full(out_dim, 2.0) # u(0) = u(1) = 2 def identity(x): """The identity mesh mapping (and its inverse).""" return x # ── Mesh and DG-Q1 Lagrange variables (fixed) ──────────────────────────────── mapping_id = InvertibleFunction(identity, identity) mesh = Mesh( dim=physical_dim, n_cells=[n_cells], ref_quad=UnitSquareTensorized(dim=physical_dim, order=quad_order), mapping=Mapping(mappings=[mapping_id]), ) def lagrange_basis_fn(coords, i, msh): return local_lagrange_basis(coords, i, msh, order=poly_deg, out_dim=out_dim) basis = AnalyticBasis( nb_basis=nb_basis, out_dim=out_dim, mesh=mesh, local_basis=lagrange_basis_fn, basis_type="scalar", ) variables = VariablesDG(basis=basis, nb_variables=out_dim) flux = SIPGFlux(sigma=sigma, h=1.0 / n_cells) # ── ApproximationSpace for c_theta(x) ───────────────────────────────────────── key = jax.random.PRNGKey(seed) c_network = MLP(in_size=physical_dim, out_size=out_dim, hidden_sizes=[16, 16], key=key) c_space = ApproximationSpace({"x": physical_dim}, [(c_network, "scalar", None)]) pde = LaplacianReactionWeakFormLearnableCoeff(dim=physical_dim, c=c_space, f=f_static) assembler = EllipticDGscheme( AbstractPhysicalWeakModel.from_weak_form(pde, dirichlet=dirichlet_bc), variables, flux, ) # ── DGEllipticApproximationSpace: wraps the DG solve ────────────────────────── space = DGEllipticApproximationSpace( dims={"x": physical_dim, "dofsl": 1}, list_assemblers=[assembler], model_type="x_dofsl", newton_kwargs={"max_iter": 3}, ) (u_fn,) = space.create_variables() # ── Physical model (empty) + data residual ──────────────────────────────────── dx = Segment1D((0.0, 1.0), is_main_domain=True) class DataFittingModel(AbstractPhysicalModel): """A model with no PDE residual, only a data residual.""" def __init__(self, domain, data): super().__init__(main_domain=domain) self.data_residuals["data"] = CollocDataResidual( size=1, model_type="x_dofsl", data=data ) # ── Observations of u ───────────────────────────────────────────────────────── x_obs = jnp.linspace(0.02, 0.98, N_OBS)[:, None] # (N_OBS, 1) u_obs = jax.vmap(u_exact)(x_obs)[:, None] # (N_OBS, 1), noise-free measurements model = DataFittingModel(dx, data=(x_obs, u_obs)) sampler = TensorizedSampler( [DomainSampler(dx)], # a domain is required by TensorizedSampler bc=False, data_samplers=model.data_residuals, ) # ── Projector (default optimiser: ENG, natural gradient) ───────────────────── pinn = Projector(model, space, sampler) print(f"Training for {N_EPOCHS} epochs (Projector, ENG, N_OBS={N_OBS}) ...") key, pinn = pinn.project(key, space, N_EPOCHS, n_dl_colloc=N_OBS, verbose=True) nspace = pinn.space loss_history = pinn.losses.losses_history print(f"\nFinal loss (data): {pinn.best_loss['total']:.4e}") # ── Plots ───────────────────────────────────────────────────────────────────── x_plot = jnp.linspace(0.0, 1.0, 300)[:, None] u_ref_plot = jax.vmap(u_exact)(x_plot) c_ref_plot = jax.vmap(c_exact)(x_plot) (dofsl_final,) = nspace.get_intermediate_values() u_h_final = jax.vmap(u_fn, in_axes=(None, 0, None))(nspace, x_plot, dofsl_final) # c_theta read where the ASSEMBLY reads it: through the coefficient space's own # variable, never its raw network (which would skip any pre/post-processing). c_learned = nspace.assemblers[0].pde.weak_form.c (c_fn,) = c_learned.create_variables() c_learned_plot = jax.vmap(c_fn, in_axes=(None, 0))(c_learned, x_plot)[:, 0] l2_u = float(jnp.sqrt(jnp.mean((u_h_final.squeeze() - u_ref_plot) ** 2))) l2_c = float(jnp.sqrt(jnp.mean((c_learned_plot - c_ref_plot) ** 2))) print(f"L2 error(u) = {l2_u:.4e} L2 error(c) = {l2_c:.4e}") fig, axs = plt.subplots(1, 3, figsize=(15, 4)) axs[0].plot(x_plot[:, 0], u_ref_plot, "k--", label="u_exact") axs[0].plot(x_plot[:, 0], u_h_final, label=f"u_h DG (L2={l2_u:.2e})") axs[0].scatter( x_obs[:, 0], u_obs[:, 0], s=4, color="grey", alpha=0.5, label="measurements" ) axs[0].set_title("DG solution vs exact") axs[0].legend(fontsize=8) axs[0].grid(True, alpha=0.3) axs[1].plot(x_plot[:, 0], c_ref_plot, "k--", label="c_exact = 10*x(1-x)") axs[1].plot(x_plot[:, 0], c_learned_plot, label=f"learned c_θ (L2={l2_c:.2e})") axs[1].set_title("Learned coefficient c(x)") axs[1].legend(fontsize=8) axs[1].grid(True, alpha=0.3) axs[2].semilogy(jnp.asarray(loss_history["total"]).reshape(-1), linewidth=1) axs[2].set_title("Loss history (data fitting)") axs[2].set_xlabel("Epoch") axs[2].set_ylabel("MSE(u_h, u_obs)") axs[2].grid(True, alpha=0.3) plt.suptitle( "-Δu + c(x)u = f(x) — learning c(x) from measurements of u", fontsize=11, ) plt.tight_layout() plt.show()