1. An unknown coefficient, in place of a fixed one

\(-\Delta u + c(x)\, u = f\)

On a disk, with u observed at a handful of points and c unknown. Replacing the fixed coefficient with a small network wrapped in an ApproximationSpace is the entire change — the forward FEM solve, the mesh, the basis, are all untouched.

import jax
from scimba_jax.linear_approximation.solvers import LinearSolve
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.networks.mlp import MLP
from scimba_jax.physical_models.classical_weakform.laplacian_inverse_linear_problem import (  # noqa: E501
    LaplacianReactionWeakFormLearnableCoeff,
)

# u OBSERVED at a handful of points, c UNKNOWN (mesh, model, variables as on
# the FEM page). c(x) = softplus(MLP(x)) is itself an ApproximationSpace --
# softplus keeps it positive, which the equation needs -- and plugs straight
# into the weak form like any other coefficient.
def positive(*args):
    return jax.nn.softplus(args[0])


coefficient = ApproximationSpace(
    {"x": 2}, [(MLP(in_size=2, out_size=1, hidden_sizes=[16, 16], key=key), "scalar", None)],
    post_processing=positive,
)
model.add_weak_form(
    "main", LaplacianReactionWeakFormLearnableCoeff(dim=2, c=coefficient, f=f_source)
)
# The FORWARD solve is unchanged -- one linear system, same as the FEM page.
scheme = EllipticFEscheme(model, variables)
space = DGEllipticApproximationSpace(
    dims={"x": 2, "dofsl": 1}, list_assemblers=[scheme], model_type="x_dofsl",
    newton_kwargs={"solver": LinearSolve(tol=1e-7)},
)
# 89 DOFs, 40 observations, 20 ENG epochs: loss 5.7e-07, c recovered to
# 8.0e-02 relative L2 -- against 0.36 for the best CONSTANT c.

2. A data loss, and the natural gradient

The Projector here has no PDE residual to minimize — its only loss compares the solve's output to what was observed. The forward solve sits inside that loss, so training it differentiates straight through the linear solve. The map from c to u is smoothing, which makes plain gradient descent a poor fit; the natural-gradient optimizer (ENG, see PINNs) rescales by the parametrization's own metric and is what makes this class of problem converge at all.

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_model import AbstractPhysicalModel
from scimba_jax.physical_models.data_residuals import CollocDataResidual

# The loss is DATA, not a PDE residual: compares u_h at the observation
# points to what was measured there.
class DataFittingModel(AbstractPhysicalModel):
    def __init__(self, domain, data):
        super().__init__(main_domain=domain)
        self.data_residuals["data"] = CollocDataResidual(
            size=1, model_type="x_dofsl", data=data
        )


model = DataFittingModel(domain, data=(points, observations))
sampler = TensorizedSampler([DomainSampler(domain)], bc=False, data_samplers=model.data_residuals)
# ENG is what makes this converge AT ALL: c -> u is smoothing, so plain
# gradient descent brings the loss down while c stays far off.
projector = Projector(model, space, sampler, optimizer="ENG", matrix_regularization=1e-6)
key, projector = projector.project(key, space, n_epochs=20, n_dl=len(points))
# `matrix_regularization` is not a tuning knob: the Gram matrix is rank
# deficient (directions the data cannot see), and its Cholesky returns nan
# without the damping.

The same idea works with any discretization here — DG, finite volumes, or a collocation basis — and with more than one coefficient at once, as long as each is wrapped the same way and stays identifiable from the data actually available.