Weak forms and ParamFunction

Weak forms in scimba_jax describe the mathematical integrand. They do not know whether it will be integrated by CG-FEM, DG, a structured mesh or an unstructured one: EllipticFEscheme and EllipticDGscheme own that work. This separation is what lets one PDE definition be reused by different discretisations.

ParamFunction: a typed symbolic function

ParamFunction is a small symbolic layer over JAX callables. It records which arguments a function accepts, then provides algebra and differential operators that retain this information. The most common specialisations are:

Class

Output

Typical operations

ParamScalarFunction

scalar

arithmetic, gradient, laplacian

ParamVecFunction

vector

dot, norm, componentwise algebra

ParamFieldFunction

vector field

divergence, curl, gradient

ParamMatrixFunction

matrix

matrix-vector products and tensor algebra

The dims dictionary describes available variables, for example {"x": 2, "t": 1, "mu": 3}. The f_type string states which of them the underlying callable consumes, in canonical order. Typical forms are:

  • "x": spatial quantity, called as f(model, x);

  • "x_mu": parametric spatial quantity, called as f(model, x, mu);

  • "t_x": time-dependent spatial quantity;

  • "x_n": boundary quantity using an outward normal;

  • "dofsl_x": quantity depending on a current discrete state.

The first argument is always the model/pytree. This makes the same expression usable with JAX transformations and ensures learnable data remain visible to the gradient. Operators compose signatures automatically: for example, u.gradient("x") returns a field with the same physical arguments as u.

from scimba_jax.nonlinear_approximation.model_class.funcparam_scalar import (
    ParamScalarFunction,
)

u = ParamScalarFunction(
    dims={"x": 2, "mu": 1},
    fn=lambda model, x, mu: model(x, mu),
    f_type="x_mu",
)
grad_u = u.gradient("x")

Avoid closing over a learnable field in a Python lambda. Put it on the weak-form object instead, as an approximation space. get_fields() then retrieves it through the live outer pytree during assembly; the optimiser can see its parameters and its gradient is not cut.

Defining a weak form

Subclass AbstractWeakForm, place coefficients on self, and return ParamFunction expressions from bilinear_form and linear_form. A coefficient can be a plain callable such as lambda x: 1.0, or an AbstractApproxSpace when it is learned. Plain callables are inspected and wrapped according to their output shape: scalar, vector field, or matrix.

class DiffusionReaction(AbstractWeakForm):
    def __init__(self, dim, diffusion, source):
        super().__init__(dim=dim)
        self.diffusion = diffusion  # callable or learnable approximation space
        self.source = source

    def bilinear_form(self, u, v):
        fields = self.get_fields()
        return fields["diffusion"] * u.gradient("x").dot(v.gradient("x"))

    def linear_form(self, v):
        return self.get_fields()["source"] * v

EllipticWeakForm is the built-in diffusion–advection–reaction version. It uses the convention (A grad(u)) . grad(v) + (b . grad(u)) v + c u v = f v.

Boundary conditions and fluxes

Boundary conditions are separate objects: Dirichlet, Neumann, Robin, and Nitsche. In CG-FEM, Dirichlet data are normally imposed strongly through the lifted residual. In DG, the boundary value is supplied as a ghost trace to the numerical flux. Neumann and Robin conditions provide face terms in both cases.

The DG scheme builds the ParamFunction values and gradients on both sides of a face, supplies the outward normal n, and then calls its AbstractFlux. Therefore a coefficient which genuinely depends on the state may be declared as coefficient(x, u_value): it is evaluated with the trace from the relevant side of the face.

What belongs where

Concern

Owner

PDE integrand and coefficients

AbstractWeakForm subclass

Basis evaluation and cell/face quadrature

EllipticFEscheme / EllipticDGscheme

DG consistency, symmetry and penalty

selected AbstractFlux

Mesh mapping, normals and physical weights

mesh

Linear/nonlinear solution process

solver and preconditioner

Keeping this boundary is particularly useful for complicated PDEs: introducing a new term should normally change the weak form or flux, not duplicate an assembler for every mesh and solver combination.