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 dimension
  • hidden_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 patterns
  • key — 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; type is 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 array
  • post_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.