Base networks
MLP, ResNet and ICNN — the building blocks. Each is a
ScimbaPytree: jit/vmap/grad-compatible, with
ndof() and layer-level access.
Input embeddings
Fourier features, periodic encodings, Neumann boundary maps — composable pre-lifts that run before the first linear layer.
Approximation spaces
ApproximationSpace wraps one or more networks with
pre/post-processing and exposes them as
ParamFunction objects via
create_variables().
Activations
14+ choices: tanh (adaptive), sin/cos, SiLU, GELU, Swish, ReLUk, hat, rational, isotropic and anisotropic RBF. All plugged in by name string.
ScalarParam
A single learnable scalar — used in inverse problems to identify an unknown coefficient (viscosity, diffusion, …) jointly with the PDE solve.
Pre / post-processing
Coordinate changes (polar, cylindrical, …) via
pre_processing; strong boundary conditions and output
constraints via post_processing or
ParamFunction algebra.
1. MLP and ResNet
MLP is the standard building block: a sequence of
Linear → activation layers. Key constructor
arguments:
in_size,out_size— input / output dimensionhidden_sizes— list of hidden widths, e.g.[64, 64]activation— activation name string (see §4)embedding— optional embedding type applied before the first layer (see §3)final_bias— disable the output bias to enforce u(0)=0 patternskey— JAX random key for weight initialisation
ResNet adds skip connections: it projects input to
a shared hidden width, applies n_blocks residual
block_class instances (each must preserve that
width), then projects to the output. Any callable
ScimbaPytree can be used as the block.
import jax
from scimba_jax.nonlinear_approximation.networks.mlp import MLP
from scimba_jax.nonlinear_approximation.networks.resnet_mlp import ResNet
key = jax.random.PRNGKey(0)
# Standard MLP: 2D input -> 1D output, two hidden layers of width 64, tanh
net = MLP(in_size=2, out_size=1, hidden_sizes=[64, 64], activation="tanh", key=key)
# MLP without bias on the last layer
net_no_bias = MLP(in_size=2, out_size=1, hidden_sizes=[64, 64],
activation="tanh", final_bias=False, key=key)
# ResNet: input/output projections + 3 residual MLP blocks of width 64
net_res = ResNet(
in_size=2, out_size=1,
block_class=MLP,
block_kwargs={"in_size": 64, "out_size": 64, "hidden_sizes": [64], "activation": "tanh"},
n_blocks=3,
key=key,
)
print(net.ndof()) # total learnable parameters
2. ICNN and ScalarParam
ICNN (Input Convex Neural Network, Amos et al. 2017)
guarantees that the output is convex in x for
any set of learned weights, by using non-negative inner-layer
weights (enforced via softplus) and skip connections
from the input at every layer. It is the natural choice for
Monge–Ampère problems and optimal transport.
ScalarParam is not a function of its input — it
returns a single trainable scalar regardless of x.
Use it to turn an unknown PDE coefficient into a learnable
parameter that is optimised jointly with the network weights.
import jax
from scimba_jax.nonlinear_approximation.networks.icnn import ICNN
from scimba_jax.nonlinear_approximation.networks.scalar_param import ScalarParam
key = jax.random.PRNGKey(0)
# ICNN: output is convex in x (Amos et al., 2017)
# -- required activation must be convex + non-decreasing (softplus, relu, ...)
icnn = ICNN(in_size=2, out_size=1, hidden_sizes=[32, 32],
activation="softplus", key=key)
# ScalarParam: a single trainable scalar -- for unknown PDE coefficients
nu = ScalarParam(init=1.0) # e.g. unknown viscosity in an inverse problem
print(nu.ndof()) # 1
3. Input embeddings
Embeddings are pre-lifts applied to the raw input before the
first linear layer. They are ScimbaPytrees and can
be chained with ComposedEmbedding.
FourierEmbedding — draws n_features
frequency vectors from a Gaussian of standard deviation
std and appends sin/cos projections. Frequencies
are learnable; the output dimension grows by
2 × n_features. Pass axes to apply the
lifting to a subset of input dimensions only.
PeriodicEmbedding — encodes known spatial
periodicity exactly via cos(2πkx/T),
sin(2πkx/T) for harmonics
k = 1 … n_periodic_features.
Frequencies are fixed (not learned).
NeumannEmbeddingOnSquare — maps each coordinate
through φ(x) = 3x²−2x³ (normalised to the unit
interval), so that both the value and its derivative vanish at
the boundary. Useful for enforcing zero-flux (Neumann) boundary
conditions strongly.
The same embedding type can be passed as the embedding
keyword of MLP, which constructs it automatically and
adjusts in_size.
import jax
import jax.numpy as jnp
from scimba_jax.nonlinear_approximation.networks.embeddings import (
FourierEmbedding,
PeriodicEmbedding,
NeumannEmbeddingOnSquare,
ComposedEmbedding,
)
from scimba_jax.nonlinear_approximation.networks.mlp import MLP
key = jax.random.PRNGKey(0)
# Fourier features: learnable Gaussian frequencies -> sin/cos pairs
# (2,) input -> (2 + 2*32,) = (66,) output, then passed to MLP
fourier = FourierEmbedding(key=key, in_size=2, n_features=32, std=1.0)
net = MLP(in_size=2, out_size=1, hidden_sizes=[64, 64], activation="tanh",
embedding="fourier", n_features=32, std=1.0, key=key)
# Periodic embedding: fixed harmonics k=1..3, no learning -- exact periodicity
periodic = PeriodicEmbedding(
key=key, in_size=2,
periods=jnp.array([1.0, 1.0]),
axes=(0, 1),
n_periodic_features=3,
)
# Neumann boundary embedding: squashes x through 3x^2-2x^3 so the output
# (and its gradient) vanishes at the domain boundary -- strong zero-flux BC
neumann = NeumannEmbeddingOnSquare(
key=key, in_size=2,
domain_bounds=((0.0, 1.0), (0.0, 1.0)),
)
# Chain embeddings: apply Neumann map then lift with Fourier features
composed = ComposedEmbedding(neumann, fourier)
4. ApproximationSpace, pre/post-processing and create_variables()
ApproximationSpace is the glue between raw networks
and the rest of Scimba. It holds:
dims— dict of named input dimensions, e.g.{"x": 2, "mu": 1}list_models— list of(network, type, size)triples;typeis one of"scalar","field", or"vec"model_type— call signature:"x_mu","x_t_mu","x_t", …pre_processing— maps raw inputs to the array the network actually receives; the default concatenates all args into one arraypost_processing— transforms the raw network output before it is wrapped; one callable per model, or a single callable broadcast to all
Coordinate changes via pre_processing.
Any smooth bijection can be applied here: polar coordinates
(r, θ), cylindrical, log-radial, or a learned
domain mapping. The Jacobian of the coordinate change is
handled automatically by JAX autodiff through the
ParamFunction operators
(.gradient_x(), .jacobian_x(), …).
Strong boundary conditions via post_processing
or ParamFunction algebra.
To enforce a homogeneous Dirichlet condition exactly, multiply
the raw network output by a distance-to-boundary function
φ(x) that vanishes on ∂Ω. Because the result is
still a ParamFunction, all differential operators
remain available. Alternatively, a post_processing
hook can enforce positivity, zero mean, or any pointwise
constraint that depends only on the output value.
Calling create_variables() wraps each network
in a ParamFunction —
ParamScalarFunction, ParamVecFunction,
or ParamFieldFunction — that supports the full
symbolic algebra: arithmetic, composition f << g,
and automatic differentiation operators.
These objects are the inputs to AbstractResidual
and to all Scimba physical models.
import jax
import jax.numpy as jnp
from scimba_jax.nonlinear_approximation.networks.mlp import MLP
from scimba_jax.nonlinear_approximation.approximation_spaces.approximation_spaces import (
ApproximationSpace,
)
from scimba_jax.nonlinear_approximation.model_class.funcparam_vectorial import (
ParamScalarFunction,
)
key = jax.random.PRNGKey(0)
dims = {"x": 2}
net_u = MLP(in_size=2, out_size=1, hidden_sizes=[64, 64], key=key)
net_q = MLP(in_size=2, out_size=2, hidden_sizes=[64, 64], key=key)
# ── Pre-processing: polar coordinate change ──────────────────────────────
# The network receives (r, theta) instead of (x, y).
def polar_pre(x):
r = jnp.linalg.norm(x, axis=-1, keepdims=True)
theta = jnp.arctan2(x[..., 1:2], x[..., 0:1])
return (jnp.concatenate([r, theta], axis=-1),)
space = ApproximationSpace(
dims=dims,
list_models=[
(net_u, "scalar", None), # -> ParamScalarFunction
(net_q, "vec", 2), # -> ParamVecFunction (size 2)
],
model_type="x",
pre_processing=polar_pre,
)
u_raw, q = space.create_variables()
# ── Post-processing: strong Dirichlet BC via ParamFunction algebra ────────
# phi(x) = 1 - |x|^2 vanishes on the unit circle => u = phi * net satisfies u=0
phi = ParamScalarFunction(
dims,
fn=lambda model, x: 1.0 - jnp.sum(x ** 2, keepdims=True),
f_type="x",
)
u = phi * u_raw # u = 0 on |x| = 1 for *any* network weights
# ── Or: post_processing hook in the space (output-only transform) ─────────
def enforce_positivity(out): return jnp.abs(out)
space2 = ApproximationSpace(
dims=dims,
list_models=[(net_u, "scalar", None)],
model_type="x",
post_processing=enforce_positivity,
)
(u_pos,) = space2.create_variables()
# All returned objects support the full ParamFunction algebra:
grad_u = u.gradient_x() # ParamFieldFunction
lap_u = u.laplacian_x() # ParamScalarFunction
div_q = q.divergence_x() # ParamScalarFunction
For worked examples see the Basics of scimba_jax tutorial or the Scimba Jax API reference.