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 |
|---|---|---|
|
scalar |
arithmetic, |
|
vector |
|
|
vector field |
|
|
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 asf(model, x);"x_mu": parametric spatial quantity, called asf(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 |
|
Basis evaluation and cell/face quadrature |
|
DG consistency, symmetry and penalty |
selected |
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.