Skip to article frontmatterSkip to article content
Site not loading correctly?

This may be due to an incorrect BASE_URL configuration. See the MyST Documentation for reference.

Learned priors — composing a flow onto a physical scaling

Authors
Affiliations
University of Valencia

A data-assimilation prior has to say which ocean states are plausible. A Gaussian says “close to this mean, with this covariance”; a normalising flow can say something much richer. But a flow trained directly on a dimensional ocean state has to spend its capacity learning that thickness is measured in hundreds of metres and velocity in tenths of a metre per second — facts we already know.

somax.flows lets you hand the flow that knowledge for free. A StateAffine carries the physical scaling; the flow is composed on top and learns only the departure from it.

Requires the optional extra: pip install somax[flows], or uv sync --extra flows.

What you’ll see:

  1. Building a StateAffine from characteristic Scales

  2. Adapting it to a flowjax bijection with to_flowjax

  3. Composing it with a flow into a prior over physical states

  4. Why the log-density picks up a constant from the scaling

1. A state and its physical scales

A shallow-water state whose fields sit four orders of magnitude apart: thickness around 1000 m, velocities around 0.1 m/s.

h  ~ 1000.0 m
u  ~ 0.100 m/s
geostrophic height scale = 1.019 m

2. The transform, as a flowjax bijection

StateAffine already exposes transform_and_log_det; to_flowjax wraps it in the flat-array interface flowjax bijections use. The flat layout is the same one somax.da uses, so a single vector feeds both a filter and a flow.

flat state size      : 192
standardised |h| max : 2.50
log|det J|           : 293.50

The standardised state is O(1) in every field. The log-determinant is a constant — the Jacobian of a fixed affine map does not depend on where you evaluate it — and it is exactly what a density has to pick up when it is pushed through the scaling.

3. A prior over physical states

learned_prior composes base → flow → out of standardised coordinates, so the resulting distribution is over physical states while the flow itself only ever sees O(1) numbers. Any flowjax bijection works as the flow; a trained masked_autoregressive_flow is the realistic choice, and Identity here keeps the arithmetic checkable by hand.

sampled h mean   : 1000.0 m   (loc 1000.0)
sampled u spread : 0.100 m/s (scale 0.100)

Samples come out in physical units with the right magnitudes, from a base distribution that is standard normal. That is the scaling doing its job.

4. The density, term by term

For ϕ\phi the standardising map,

log⁡pX(x)=log⁡pY(ϕ(x))+log⁡∣det⁡Jϕ(x)∣.\log p_X(x) = \log p_Y(\phi(x)) + \log|\det J_\phi(x)|.

With an identity flow both terms are computable directly, so the composition can be checked rather than taken on trust.

prior.log_prob(x)              = 18.2926
log p_base(phi(x)) + log|det J| = 18.2926

Why flowjax is optional

flowjax is not a core somax dependency. Its own Affine forces the scale through a trainable softplus parameterisation, which is the opposite of what a fixed physical scale wants, and it brings the whole flow stack with it. The log-determinant only matters when transforming densities; everywhere else in somax — the DA flatten bridge, ScaledModel, the nondimensional factories — the affine map is used directly and no flow library is involved.