r"""Learn the post-processing operator of an elliptic DG solve, by physical loss. -U'' = -2 on (0, 1), U = (1+x)**2, U = Op(u_h) = u_h + N(u_h) Same problem as ``tests/.../test_post_processing_solve.py``, where ``Op`` was the square, written by hand, and Q1 was exact. Here ``N`` is a small network and nothing says it should square: it is learned, and the only thing driving it is the **physical loss** -- no reference solution anywhere. What makes that possible is worth stating, because it is not obvious. The DG solve inside the loss drives the *discrete* residual to zero, so measuring that again would measure nothing: it is zero for every ``N``. The loss is the **continuous** residual of the PDE, sampled at random points, plus the boundary condition -- and that one is not zero. It sees the difference between "this solves the discrete equations" and "this solves the PDE", which is exactly the gap a Q1 space leaves unless the operator makes the solution representable. The gradient reaches ``N`` through the Newton solve by implicit differentiation, which the scheme already provides. So the learning problem has a genuine minimum: with ``N`` the square, ``u_h`` only has to hold the line ``1+x``, the Q1 solution is exact, and the continuous residual vanishes. Any other ``N`` leaves a residual. What is *not* claimed is that the network recovers the square: it is determined only up to a rescaling of its input (``N(t) = (t/lam)**2`` with ``u_h = lam (1+x)`` does just as well), and only on the range of values ``u_h`` visits. The check is the error, not the resemblance. """ # %% 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.approximation_spaces import ( # noqa: E501 ApproximationSpace, ) 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.abstract_postprocessing_op import ( AbstractPostProcessingOp, ) from scimba_jax.physical_models.classical_weakform.laplacian_weak_form import ( LaplacianWeakForm, ) from scimba_jax.physical_models.elliptic_pde.laplacians import LaplacianDirichletDG DIM, OUT_DIM = 1, 1 N_CELLS, POLY_ORDER = 4, 1 QUAD_ORDER = 6 N_EPOCHS, N_COLLOC = 200, 600 NEWTON_ITER = 20 # the solve is nonlinear now: one iteration is not enough SEED = 0 def u_exact_fn(x): """``U = (1+x)**2``, the exact solution; used for the data and the error.""" return jnp.array([(1.0 + x[0]) ** 2]) def f_rhs(x): """``f = -U'' = -2``.""" return jnp.full((OUT_DIM,), -2.0) def g_bc(x, n): """The Dirichlet data of the *loss*; boundary residuals also receive ``n``.""" return u_exact_fn(x) # %% The operator whose nonlinearity is learned. class ComposedNonlinearityOp(AbstractPostProcessingOp): """``Op(u) = u + N(u)``: the identity, plus a learned correction. The field is picked up by ``get_fields`` like any field of a weak form, and the composition is the algebra's own ``<<``. Being an ``ApproximationSpace`` makes it dynamic with nothing to declare -- the pytree carries it, and the implicit differentiation of the Newton solve carries the gradient back. Why ``u +`` and not ``N(u)`` alone: an untrained network is nearly flat (here ``|N| < 0.3`` and ``N' < 0.23`` while the solution runs from 1 to 4), so the discrete system asks the DOFs to run off to infinity to compensate, and the Newton solve diverges before training has begun -- measured: ``max|dofs|`` reaches ``8e16``, then NaN. Starting from the identity makes the initial state the *plain DG solve*, which is well-posed, and learning deforms it from there. The global Newton has no damping, so a starting point that solves is not a nicety. """ def __init__(self, net, dim: int = DIM): super().__init__(dim=dim, linear=False, kills_constants=False) self.net = net def formula(self, u): return u + (self.get_fields()["net"] << u) # %% The DG space whose post-processing is the only dynamic thing. def make_scheme(op, order=POLY_ORDER): """An elliptic DG scheme for ``-U'' = -2`` with ``U = Op(u_h)``.""" 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=order + 1, out_dim=OUT_DIM, mesh=mesh, local_basis=lambda c, i, m: local_taylor_basis( c, i, m, order=order, out_dim=OUT_DIM ), basis_type="scalar", ) model = AbstractPhysicalWeakModel.from_weak_form( LaplacianWeakForm(dim=DIM, f=f_rhs), dirichlet=u_exact_fn ) return EllipticDGscheme( model, VariablesDG(basis=basis, nb_variables=OUT_DIM, post_processing=op), SIPGFlux(sigma=order * (order + 1), h=1.0 / N_CELLS), ) def make_space(op, order=POLY_ORDER): """The approximation space that solves the DG system inside the loss.""" return DGEllipticApproximationSpace( dims={"x": DIM, "dofsl": 1}, list_assemblers=[make_scheme(op, order)], model_type="x_dofsl", newton_kwargs={"max_iter": NEWTON_ITER, "tol": 1e-11}, ) def solved_error(space): """Solve with the space's current operator and measure the L2 error on U.""" scheme = EllipticDGscheme.solve( space.assemblers[0], max_iter=NEWTON_ITER, tol=1e-11 ) return scheme, float(l2_error(scheme, u_exact_fn)) # %% Train: the network inside the operator is the only parameter. net = MLP( in_size=1, out_size=OUT_DIM, hidden_sizes=[16, 16], key=jax.random.PRNGKey(SEED) ) op_init = ComposedNonlinearityOp(ApproximationSpace({"x": 1}, [(net, "scalar", None)])) domain = Segment1D([0.0, 1.0], is_main_domain=True) space = make_space(op_init) pinn_model = LaplacianDirichletDG( main_domain=domain, f_rhs=f_rhs, bc="weak", f_bc_rhs=g_bc ) sampler = TensorizedSampler([DomainSampler(domain)], bc=True) pinn = Projector(pinn_model, space, sampler) key = jax.random.PRNGKey(SEED) key, sample_dict = sampler.sample(key, N_COLLOC) print( f"-U'' = -2 on (0,1), U = Op(u_h) = u_h + N(u_h), {N_CELLS} cells Q{POLY_ORDER}" ) print(f" initial physical loss = {float(pinn.evaluate_loss(space, sample_dict)):.3e}") print(f" initial L2 error on U = {solved_error(space)[1]:.3e}") key, pinn = pinn.project(key, space, N_EPOCHS, N_COLLOC) space_opt = pinn.space scheme_opt, err = solved_error(space_opt) op_opt = space_opt.assemblers[0].variables.post_processing print(f" final physical loss = {float(pinn.best_loss['total']):.3e}") print(f" final L2 error on U = {err:.3e}") # What the same space does without any operator, and with the square written by # hand: the two ends of the range the learned operator is judged against. plain = make_space(None) print(f" the same Q1 space, no operator at all: {solved_error(plain)[1]:.3e}") # %% What was learned, and what it buys. var = scheme_opt.variables _, x_all = var.mesh.evaluate_mesh_weights_points() xs = np.asarray(x_all).reshape(-1) u_post = np.asarray(var.evaluate_quad(x_all)).reshape(-1) expansion = np.asarray( jax.vmap( lambda i, p: jnp.einsum("iv,qiv->qv", var.dofsl[i], var.trial_basis(i, p)) )(var.mesh.cells_idx, x_all) ).reshape(-1) fig, ax = plt.subplots(1, 3, figsize=(15, 4)) fig.suptitle( f"Elliptic DG with a learned post-processing — {N_CELLS} cells, Q{POLY_ORDER}" ) ts = np.linspace(expansion.min(), expansion.max(), 200) net_vals = np.asarray( jax.vmap(lambda s: jnp.ravel(op_opt.net.models[0](jnp.atleast_1d(s)))[0])( jnp.asarray(ts) ) ) scale = net_vals[len(ts) // 2] / ts[len(ts) // 2] ** 2 ax[0].plot(ts, net_vals, "r-", lw=1.6, label="learned N") ax[0].plot(ts, scale * ts**2, "k--", lw=1.4, label=f"{scale:.3f} t^2, for reference") ax[0].set_title("what N settled on (not required to be a square)") ax[0].set_xlabel("t (the expansion's value)") ax[0].legend(fontsize=8) ax[0].grid(alpha=0.3) ax[1].plot(xs, (1.0 + xs) ** 2, "k-", lw=2, label="exact U = (1+x)^2") ax[1].plot(xs, u_post, "r--", lw=1.4, label="N(u_h), solved") ax[1].plot(xs, expansion, "b-", lw=1.2, label="the Q1 expansion u_h") ax[1].set_title("a quadratic solution out of a Q1 space") ax[1].set_xlabel("x") ax[1].legend(fontsize=8) ax[1].grid(alpha=0.3) curve = np.asarray(jnp.asarray(pinn.losses.losses_history["total"]).reshape(-1)) ax[2].semilogy(curve, lw=1.3) ax[2].set_title("physical loss") ax[2].set_xlabel("epoch") ax[2].grid(alpha=0.3) plt.tight_layout() plt.show()