Operations
Every quantity in jNO — a coordinate, a network output, a residual, a derivative — is a traced
expression, an instance of Placeholder. You build a problem by operating on these expressions.
This page is the complete menu of what you can do to one, in two families:
- Part A — operations available on any traced expression (derivatives, integrals, reductions, units, trackers…).
- Part B — operations that only make sense
when the expression is backed by trainable parameters, i.e. a
Modelreturned byjno.nn(...)(optimizers, LoRA, freezing, Bayesian sampling…).
.scale(...) is overloaded — the receiver decides
The same method name means two different things depending on what it is called on:
- On an expression or a called field —
x.scale(0.1),net(x).scale(U)— it declares a characteristic magnitude for non-dimensionalization (Part A, pairs with.unit). - On a model object —
net.optimizer(optax.adam).scale(lrs.exponential(...))— it sets the learning-rate schedule (Part B).
This is the general rule for the whole page: a handful of controls (.scale, .regularize,
.mask, .freeze, .lora, …) require trainable parameters and belong to Part B; everything in
Part A works on any expression, trainable or not.
Part A — operations on any traced expression
Differentiation
Every expression carries derivative methods; the differentiation scheme rides on the call. The
method forms and the jno.numpy (jnn) free-function forms are equivalent:
import jno.numpy as jnn
u_x = u.d(x) # ∂u/∂x (Jacobian) — same as jnn.grad(u, x)
u_xx = u.d2(x) # ∂²u/∂x² (Hessian) — same as jnn.hessian(u, [x])
u_x = u.d(x, scheme="finite_difference") # the scheme is an argument
lap = u.laplacian(x, y) # ∇²u — same as jnn.laplacian(u, [x, y])
H = u.hessian(x, y) # full Hessian matrix
g = u.grad(x, y) # spatial gradient (VectorView) [∂u/∂x, ∂u/∂y]
Aliases: .diff = .d, .dd = .d2. Higher-order derivatives chain: u.d(x).d(x).
Vector-calculus helpers live on jnn: jnn.jacobian, jnn.divergence, jnn.curl_2d, jnn.curl_3d.
Schemes (scheme= on any derivative). Finite-difference schemes require
compute_mesh_connectivity=True on the domain.
| Scheme string | Grad | Lap/Hess | Notes |
|---|---|---|---|
"automatic_differentiation" (default) |
✅ | ✅ | Exact; any domain |
"automatic_differentiation:forward" / :reverse |
✅ | — | jacfwd / jacrev |
"automatic_differentiation:fwd-over-rev" (default Hessian) / :fwd-over-fwd / :rev-over-rev |
— | ✅ | 2nd-order AD variants |
"finite_difference" |
✅ | ✅ | Area-weighted; general unstructured meshes |
"finite_difference:lsq" / :uniform / :inverse_distance" |
✅ | ✅ | Least-squares / uniform / distance-weighted |
"finite_difference:cotangent" |
— | ✅ | Cotangent Laplacian; 2D only |
Set a project-wide default with jno.setup(__file__, diff_type="forward", hessian_type="fwd-over-fwd").
Spell the Laplacian however reads best. u.xx + u.yy, u.d2(x) + u.d2(y) and
u.laplacian(x, y) describe the same operator, and jno.core compiles all three to the same
single node: a trace pass folds a sum of squared partials over distinct coordinates into one
Laplacian, so the network is evaluated and differentiated once instead of once per coordinate.
On a 2-D+time PINN (513 collocation points, MLP 4×64) that is 308 MFLOP/step for every spelling,
against 390 for the unfused u.xx + u.yy and 470 for u.d2(x) + u.d2(y).
The fold is deliberately conservative — it applies only where the two forms are the same
mathematics. Terms keep their own nodes when they repeat a coordinate (u.xx + u.xx is
2 ∂²u/∂x²), when they are subtracted rather than added, when they sit over different fields or
different AD modes, when the coordinate is temporal (u.tt evaluates through the time path), and
for every finite_difference scheme (:cotangent returns the whole Laplacian for any requested
dimension, so folding would halve it). FEM weak forms are left untouched — the variational route
lowers them by pattern.
FEM weak forms — .bind then attribute derivatives. A finite-element trial/test symbol is bound to
its quadrature coordinates once, after which derivatives read as plain attributes:
ui = u.bind(x=xi, y=yi, t=ti) # bind the symbol to coordinates
ui.x, ui.y, ui.z # spatial derivatives ∂u/∂x, ∂u/∂y, ∂u/∂z
ui.t # time derivative
Integration
.integrate() collapses a field to a scalar by summing over the mesh. The region (volume vs
boundary) is auto-detected from the Variable tags inside the expression — you pass no region
argument. Requires compute_mesh_connectivity=True.
vol = u.integrate() # ∫_Ω u dV (interior tag → volume weights)
bnd = u_b.integrate() # ∫_∂Ω u ds (boundary tag → surface weights) — jnn.integrate(u) is the alias
Flux integrals are written explicitly — request normals and form the dot product yourself:
x_b, y_b, _, nx, ny = dom.variable("boundary", normals=True, split=True)
flux = (u_b.d(x_b) * nx + u_b.d(y_b) * ny).integrate() # ∮ ∂u/∂n ds
An Integral is an ordinary scalar node — differentiable and jax.jit-compatible, so it drops
straight into a loss ((u.integrate() - target).square()) or a tracker.
Reductions, math & comparisons
Reductions are properties returning a squeezed scalar node; the loss helpers are the ones you reach for most:
u.mean u.sum u.min u.max u.std # reductions
u.mse # mean(square(x)) u.mae # mean(abs(x))
u.shape u.T u.real u.imag # structural / complex parts
Symbolic comparisons return trace nodes: a.equal(b), a.not_equal(b), and the operators
>, <, >=, <=. The full elementwise math library (sin, exp, sqrt, where, concat,
stack, dot, matmul, norm, …) lives in jno.numpy / jno.np — see the
jno.numpy Reference for the complete catalog.
Semantic views & binding
Typed views reinterpret an expression without copying, exposing the right accessors for its role:
Every view supports .bind(**named_vars) (alias .partials(...)) to attach the coordinate Variables
a field depends on, so attribute-style derivatives (.x, .t) work even when those coordinates are not
the network's own inputs.
Units & non-dimensionalization
Annotate the dimension and characteristic magnitude of any leaf, and jno.units audits
consistency and extracts the dimensionless groups (Fourier / Péclet numbers) of a residual — then
rewrites it to a well-scaled O(1) form.
x = x.unit("m").scale(L) # dimension + characteristic length
u = net(x, t).unit("K").scale(U) # dimension + characteristic magnitude of the field
res = u.d(t) - alpha * u.d2(x)
jno.units.check(res) # audit dimensional consistency (.warnings is empty if OK)
jno.units.infer(res) # the inferred Unit of an expression
report = jno.units.nondimensionalize(res) # each term's dimensionless group πᵢ = Sᵢ / S_ref
transformed, rescaler = jno.units.rescale(res) # rewrite to O(1) dimensionless form
rdom = rescaler.rescaled_domain(dom) # a copy of the domain with coordinates scaled to O(1)
u_phys = rescaler.to_physical(u_hat) # map a dimensionless solution back: u = U · û
nondimensionalize / rescale operate on the additive terms within a single residual
(πᵢ = Sᵢ / S_ref), not on a ratio between two separate losses. Today only coordinates and the network
output are annotatable through the public API; a bare material coefficient has no public .unit hook
yet. See the
Gradient Conflict tutorial for a worked example.
Custom functions
Wrap an arbitrary JAX function so it joins the symbolic graph and stays differentiable — for nonlinear constitutive laws, lookup tables, or anything cleaner as standalone code:
Trackers, labels & debugging
A tracker is logged every interval steps but does not contribute to the loss:
val_error = jno.np.mean(jno.np.abs(u - u_exact)).tracker(100) # or jno.np.tracker(expr, interval=100)
crux = jno.core([pde.mse, bc.mse, val_error])
Logged values appear in the statistics returned by solve(). Related metadata methods (all return
self, chainable): .name("label") tags an expression for logs / W&B, and .print(what="shape")
emits a runtime shape/stat/value and passes the value through.
Gradient control
.stop_gradient is identity in the forward pass and zero in the backward pass — freeze part of a graph
or turn an expensive quantity into a constant regulariser:
J_sg = u.grad(u_net).stop_gradient # treat the current Jacobian as a constant
ntk_reg = (J_sg @ J_sg.T - target_K).mse
Parameter Jacobian & the Neural Tangent Kernel
u.grad(net) is .grad in its parameter overload: passed a single Model (rather than
coordinates) it returns the Jacobian of the expression w.r.t. the network's trainable parameters,
shape (B, N, P) (or (B, N, D, P) for vector output). It is an ordinary node — usable as a tracker,
a loss, or evaluated after training.
J = u.grad(u_net) # ∂u/∂θ
K = J[0] @ J[0].T # (N, N) Neural Tangent Kernel
# Restrict to a parameter subset (cheaper) via a boolean pytree + net.mask(...)
import equinox as eqx, jax
all_false = jax.tree_util.tree_map(lambda _: False, u_net.module)
output_mask = eqx.tree_at(lambda m: m.output_layer.weight, all_false, True)
J_out = u.grad(u_net.mask(output_mask)) # only the output-layer weights
See the Gradient Conflict tutorial for NTK conditioning and gradient cosine-similarity diagnostics.
Part B — operations that require trainable parameters
These act on a Model — the object returned by jno.nn(module) (any Equinox module: a foundax
model or your own eqx.Module). Calling it, net(x, y), yields a ModelCall. All controls return
self and chain.
import optax, foundax as fx, jno
from jno import LearningRateSchedule as lrs
net = jno.nn(fx.mlp(in_features=2, output_dim=1, hidden_dims=64, num_layers=3), name="u_net")
Trainable scalar parameters
jno.np.parameter creates a trainable array that optimises exactly like a network — the building block
for inverse problems (identify unknown PDE coefficients from residuals):
a = jno.np.parameter((1,), key=k1, name="a")
a.optimizer(optax.adam(1e-2))
residual = a * jno.np.sin(π * x) - target
Optimizer & learning rate
Per-group control via .mask(...) (consumed by the next mutator; a bare optimizer(...) clears all
groups):
net.optimizer(optax.adamw).scale(lrs(1e-3)) # global fallback
net.mask(decoder_mask).optimizer(optax.adam).scale(lrs(5e-4))
jno.optimizers adds custom second-order optimizers not in optax (engd, ssbroyden, ssbfgs,
soap, md) — each an optax GradientTransformation, composable with optax.chain. See
Optimizer & LR and Schedules.
Regularization
.regularize is called on a field (a called network or a FEM nodal parameter) and returns a
pointwise penalty term — FEM-exact where possible, else the autodiff form:
reg = net(x, y).regularize("h1seminorm", x, y) # smooth / tv / nonneg / bounded / l2(FEM)
crux = jno.core([pde.mse, reg])
Parameter selection & freezing
net.freeze() net.unfreeze() # exclude / re-include from training
net.mask(param_mask) # one-shot boolean-pytree scope for the next control
net.constrain(jax.nn.softplus) # reparameterize params before every forward pass
Fine-tuning (LoRA)
Details and per-target configuration: LoRA.
Bayesian & variational inference
Turn a point estimate into a posterior — MCMC (.bayesian) or a variational ELBO fit (.vi, mutually
exclusive):
net.bayesian(blackjax.nuts, warmup=500, keep=1000) # posterior via MCMC
net.vi(blackjax.meanfield_vi, optimizer=optax.adam(1e-3)) # variational approximation
net.posterior_samples net.posterior_diagnostics # draws + per-step diagnostics (or None)
See Bayesian Sampling.
Precision & initialization
net.dtype(jnp.float64) # set params + compute dtype
net.initialize("weights.eqx") # load pretrained weights (path / pytree / initializer)
Diagnostics, sweeps & deployment
net.summary() net.dont_show() # print / suppress the model-control summary
net.tune(optimizer=[optax.adam, optax.sgd], lr=[1e-3, 1e-4]) # declare per-model HP-sweep options
net.reset() # reset training config to defaults
net.to_iree(sample_inputs) # compile to an IREEModel for deployment
Weights are persisted with the free functions jno.save / jno.load (not model methods). See
Hyperparameter Tuning and IREE Deployment.