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.

API Reference

Authors
Affiliations
University of Valencia

Core

The base types, model contract, term algebra, forcing, stratification, elliptic caches, and checkpointing that every component builds on (re-exported at the top level as somax and somax.core).

BasisForcing

class

BasisForcing(coeffs: "Float[Array, ' m']", spatial: 'SpatialBasis', temporal: 'TemporalBasis', grid_shape: 'tuple[int, ...]') -> None

Reduced-order forcing: a fixed space-time frame driven by a coefficient vector.

Compose

class

Compose(outer: 'Term', inner: 'Term') -> None

Operator composition: Compose(outer, inner)(state) = outer(inner(state)).

ConstantForcing

class

ConstantForcing(field: 'Array') -> None

Time-independent forcing field.

ConstantInTime

class

ConstantInTime(m: 'int') -> None

Time-independent gate: b(t) = ones(m) (the Phi_t = I case).

Diagnostics

class

Diagnostics() -> None

Base class for on-demand diagnostic quantities.

DirichletHelmholtzCache

class

DirichletHelmholtzCache(solver: 'eqx.Module') -> None

Helmholtz cache for Dirichlet (DST-based) domains.

ForcingProtocol

class

ForcingProtocol() -> None

Base class for forcing terms.

ForcingTerm

class

ForcingTerm(forcing: 'ForcingProtocol', place: 'Callable[[PyTree, Array], PyTree]', grid: 'eqx.Module | None' = None) -> None

Lift a :class:ForcingProtocol field into the term algebra as a tendency.

FourierInTime

class

FourierInTime(freqs: "Float[Array, ' m']", phases: "Float[Array, ' m']") -> None

Cosine temporal gate b_a(t) = cos(omega_a t + phase_a).

GaussianWindowsInTime

class

GaussianWindowsInTime(centers: "Float[Array, ' m']", widths: "Float[Array, ' m']") -> None

Localized temporal gate b_a(t) = exp(-(t - tau_a)^2 / (2 T_a^2)).

HelmholtzCache

class

HelmholtzCache() -> None

Cached Helmholtz solver for repeated (lap - lambda) psi = f solves.

InterpolatedForcing

class

InterpolatedForcing(path: 'dfx.AbstractPath') -> None

Data-driven forcing via diffrax path interpolation.

ModalTransform

class

ModalTransform(Cl2m: "Float[Array, 'nl nl']", Cm2l: "Float[Array, 'nl nl']", eigenvalues: 'Array', rossby_radii: 'Array') -> None

Precomputed layer-to-mode and mode-to-layer transforms.

MultimodalHelmholtzCache

class

MultimodalHelmholtzCache(caches: 'tuple[HelmholtzCache, ...]') -> None

Batched Helmholtz cache for multilayer modal solves.

NeumannHelmholtzCache

class

NeumannHelmholtzCache(solver: 'eqx.Module') -> None

Helmholtz cache for Neumann (DCT-based) domains.

NoForcing

class

NoForcing() -> None

Zero forcing (free evolution).

Params

class

Params() -> None

Base class for differentiable model parameters.

PeriodicHelmholtzCache

class

PeriodicHelmholtzCache(solver: 'eqx.Module', lambda_: 'float', zero_mean: 'bool' = True) -> None

Helmholtz cache for periodic (FFT-based) domains.

PhysConsts

class

PhysConsts() -> None

Base class for frozen physical constants.

Scaled

class

Scaled(term: 'Term', coeff: 'float | Array') -> None

A term multiplied by a scalar coefficient.

ScaledModel

class

ScaledModel(inner: 'SomaxModel', transform: 'StateAffine', time_scale: 'float' = 1.0) -> None

Integrate inner in transformed coordinates.

Scales

class

Scales(L: 'float', U: 'float', H: 'float', f0: 'float', T: 'float', g: 'float' = 9.81, kind: 'ScaleKind' = 'advective') -> None

Characteristic scales of a run. All fields static.

SeasonalWindForcing

class

SeasonalWindForcing(tau0: 'Array', omega: 'float', phase: 'float' = 0.0) -> None

Sinusoidal wind forcing with learnable amplitude.

SimulationCheckpointer

class

SimulationCheckpointer(checkpoint_dir: 'str', checkpoint_interval: 'int') -> None

Manages periodic checkpointing during long simulations.

SomaxModel

class

SomaxModel() -> None

Abstract base class defining the somax model contract.

SpatialBasis

class

SpatialBasis(Phi: "Float[Array, ' Ngrid m']", std: "Float[Array, ' m']") -> None

A precomputed spatial dictionary plus the per-mode prior std.

State

class

State() -> None

Base class for model state vectors.

StateAffine

class

StateAffine(loc: 'PyTree', scale: 'PyTree') -> None

Per-leaf affine map on a State pytree: y = (x - loc) / scale.

StratificationProfile

class

StratificationProfile(H: 'Array', g_prime: 'Array', rho: 'Array | None' = None) -> None

Discrete vertical stratification for a layered ocean model.

Sum

class

Sum(terms: 'tuple[Term, ...]') -> None

Additive composition of terms: sum_i terms[i](t, state, args).

TemporalBasis

class

TemporalBasis() -> None

Maps a scalar time to per-atom temporal weights b(t).

Term

class

Term() -> None

A single additive contribution to a model right-hand side.

TermFn

class

TermFn(fn: 'Callable[[float, PyTree, PyTree | None], PyTree]', _kind: 'Kind' = 'explicit') -> None

Adapt a plain callable into a :class:Term.

TermModel

class

TermModel(terms: 'Term') -> None

A :class:SomaxModel whose RHS is an assembled :class:Term tree.

TransformedForcing

class

TransformedForcing(base: 'ForcingProtocol', inverse: 'Callable[[Array], Array]') -> None

Apply a pointwise transform to a base forcing (e.g. log-space synthesis).

VectorBasisForcing

class

VectorBasisForcing(coeffs: "Float[Array, ' m']", spatial: 'VectorSpatialBasis', temporal: 'TemporalBasis', grid_shape: 'tuple[int, ...]') -> None

Reduced-order vector forcing: a fixed vector frame driven by coeffs.

VectorSpatialBasis

class

VectorSpatialBasis(Phi: "Float[Array, ' Ngrid m ncomp']", std: "Float[Array, ' m']") -> None

A precomputed vector dictionary plus the per-mode prior std.

add_to

function

add_to(component: 'str', layer: 'int | None' = None) -> 'Callable[[PyTree, Array], PyTree]'

Build a placement that adds a field onto one named state component.

add_vector_to

function

add_vector_to(components: 'tuple[str, ...]', layer: 'int | None' = None) -> 'Callable[[PyTree, Array], PyTree]'

Build a placement that adds a component-major field onto named components.

as_parameter

function

as_parameter(value: 'ArrayLike | Parameterize | NonTrainable') -> 'Any'

Coerce a factory argument into a Params leaf.

build_diffrax_terms

function

build_diffrax_terms(term: 'Term', *, state_fn: 'Callable[[PyTree], PyTree] | None' = None) -> 'dfx.AbstractTerm'

Build the diffrax term object for term, IMEX-aware.

control_filter

function

control_filter(forcing: 'BasisForcing') -> 'BasisForcing'

Boolean filter selecting only coeffs for gradient updates.

explicit

function

explicit(term: 'Term') -> 'Term'

Tag term for the explicit stage of an IMEX integrator.

frozen

function

frozen(value: 'ArrayLike') -> 'NonTrainable'

Hide a parameter from gradients while keeping it a runtime value.

geostrophic_currents

function

geostrophic_currents(domain: 'Domain', *, num_basis_per_dim: 'int | tuple[int, ...]' = 8, length_scale: 'float' = 1.0, nu: 'float' = 1.5, variance: 'float' = 1.0) -> 'VectorBasisForcing'

Vector preset: incompressible (u, v) current-error forcing.

implicit

function

implicit(term: 'Term') -> 'Term'

Tag term for the implicit stage of an IMEX integrator.

interval

function

interval(value: 'ArrayLike', lower: 'float', upper: 'float') -> 'Parameterize'

Constrain a parameter to the open interval (lower, upper).

matern_spectral_density

function

matern_spectral_density(sqrt_lambda: "Float[Array, ' m']", *, variance: 'float' = 1.0, length_scale: 'float' = 1.0, nu: 'float' = 1.5, ndim: 'int' = 2) -> "Float[Array, ' m']"

Matérn power spectral density S(omega) at omega = sqrt_lambda.

partition

function

partition(term: 'Term') -> 'tuple[Term | None, Term | None]'

Split term into (explicit, implicit) sub-terms by kind.

positive

function

positive(value: 'ArrayLike') -> 'Parameterize'

Constrain a parameter to be strictly positive.

spatial_from_divfree

function

spatial_from_divfree(domain: 'Domain', *, num_basis_per_dim: 'int | tuple[int, ...]', length_scale: 'float' = 1.0, nu: 'float' = 1.5, variance: 'float' = 1.0) -> 'VectorSpatialBasis'

Build a :class:VectorSpatialBasis of divergence-free velocity atoms.

spatial_from_eof

function

spatial_from_eof(data: "Float[Array, 'T N']", n_modes: 'int', *, center: 'bool' = True) -> 'SpatialBasis'

Build a :class:SpatialBasis from empirical orthogonal functions (PCA).

spatial_from_fourier

function

spatial_from_fourier(domain: 'Domain', *, num_basis_per_dim: 'int | tuple[int, ...]', length_scale: 'float' = 1.0, nu: 'float' = 1.5, variance: 'float' = 1.0) -> 'SpatialBasis'

Build a :class:SpatialBasis from the box-Laplacian (HSGP) eigenbasis.

spatial_from_gabor

function

spatial_from_gabor(domain: 'Domain', *, n_scales: 'int', base_scale: 'float', slope: 'float' = 4.0, amplitude: 'float' = 1.0, oversample: 'float' = 1.0) -> 'SpatialBasis'

Build a :class:SpatialBasis from a geonnax dyadic radial-Gabor frame.

spatial_from_graph_laplacian

function

spatial_from_graph_laplacian(adjacency: "Float[Array, 'V V']", n_modes: 'int', *, normalized: 'bool' = True, regularization: 'float' = 0.001, smoothness: 'float' = 2.0) -> 'SpatialBasis'

Build a :class:SpatialBasis from graph-Laplacian eigenvectors.

spatial_from_rbf

function

spatial_from_rbf(domain: 'Domain', centers: "Float[Array, 'm ndim']", widths: "Float[Array, ' m']", *, kernel: 'str' = 'gaussian', std: "Float[Array, ' m'] | float" = 1.0) -> 'SpatialBasis'

Build a :class:SpatialBasis from a geonnax placeable radial basis.

spatial_from_spherical_rbf

function

spatial_from_spherical_rbf(domain: 'Domain', centers_lonlat: "Float[Array, 'm 2']", widths: "Float[Array, ' m']", *, kernel: 'str' = 'gaussian', std: "Float[Array, ' m'] | float" = 1.0, degrees: 'bool' = True) -> 'SpatialBasis'

Build a :class:SpatialBasis from a geodesic (on-sphere) radial basis.

spatial_from_wavelet

function

spatial_from_wavelet(domain: 'Domain', *, wavelet: 'str' = 'haar', levels: 'int | None' = None, std: "Float[Array, ' m'] | float" = 1.0) -> 'SpatialBasis'

Build a :class:SpatialBasis from the orthonormal 2D wavelet basis.

ssh_geostrophic

function

ssh_geostrophic(domain: 'Domain', *, n_scales: 'int' = 6, base_scale: 'float' = 20000.0, slope: 'float' = 4.0, amplitude: 'float' = 2e-06, oversample: 'float' = 1.0, windows: "tuple[Float[Array, ' m_t'], Float[Array, ' m_t']] | None" = None) -> 'BasisForcing'

SSH geostrophic preset: a radial-Gabor frame with a wavenumber-law prior.

sss_coastal

function

sss_coastal(domain: 'Domain', centers: "Float[Array, 'm ndim']", widths: "Float[Array, ' m']", *, kernel: 'str' = 'wendland_c2', std: "Float[Array, ' m'] | float" = 1.0, windows: "tuple[Float[Array, ' m_t'], Float[Array, ' m_t']] | None" = None) -> 'BasisForcing'

Coastal SSS preset: placeable radial atoms with a prescribed prior.

sst_frontal

function

sst_frontal(domain: 'Domain', *, num_basis_per_dim: 'int | tuple[int, ...]' = 12, length_scale: 'float' = 1.0, nu: 'float' = 1.5, variance: 'float' = 1.0, windows: "tuple[Float[Array, ' m_t'], Float[Array, ' m_t']] | None" = None) -> 'BasisForcing'

SST preset: a smooth Matérn (HSGP) field over the box-Laplacian eigenbasis.

tile_in_time

function

tile_in_time(spatial: 'SpatialBasis', centers: "Float[Array, ' m_t']", widths: "Float[Array, ' m_t']") -> 'tuple[SpatialBasis, GaussianWindowsInTime]'

Lift a spatial dictionary into a separable space-time frame.

trainable_mask

function

trainable_mask(tree: 'PyTree') -> 'PyTree'

Boolean pytree marking which leaves an optimiser may update.

Models

Dynamical-system and ocean model classes, with their state, parameter, and diagnostic companions.

BaroclinicQG

class

BaroclinicQG(params: 'BaroclinicQGParams', consts: 'BaroclinicQGPhysConsts', grid: 'CartesianGrid2D', diff: 'Difference2D', interp: 'Interpolation2D', mask: 'Mask2D | None', modal: 'ModalTransform', strat: 'StratificationProfile', beta_y: "Float[Array, 'Ny Nx']", wind_forcing: "Float[Array, 'Ny Nx']", helmholtz_lambdas: 'Array', poisson_bc: 'str' = 'dst') -> None

Multilayer quasi-geostrophic model on an Arakawa C-grid.

BarotropicQG

class

BarotropicQG(params: 'BarotropicQGParams', consts: 'BarotropicQGPhysConsts', grid: 'CartesianGrid2D', diff: 'Difference2D', interp: 'Interpolation2D', mask: 'Mask2D | None', beta_y: "Float[Array, 'Ny Nx']", wind_forcing: "Float[Array, 'Ny Nx']", poisson_bc: 'str' = 'dst') -> None

Barotropic quasi-geostrophic model on an Arakawa C-grid.

Burgers1D

class

Burgers1D(params: 'Burgers1DParams', grid: 'CartesianGrid1D', diff: 'Difference1D', advection: 'Advection1D', mask: 'Mask1D | None', periodic: 'bool' = True, method: 'str' = 'upwind1') -> None

1D Burgers equation on an Arakawa C-grid.

Burgers2D

class

Burgers2D(params: 'Burgers2DParams', grid: 'CartesianGrid2D', diff: 'Difference2D', advection: 'FVXAdvection2D', interp: 'Interpolation2D', mask: 'Mask2D | None', method: 'str' = 'upwind1') -> None

2D Burgers equation on an Arakawa C-grid.

Diffusion1D

class

Diffusion1D(params: 'Diffusion1DParams', grid: 'CartesianGrid1D', diff: 'Difference1D', mask: 'Mask1D | None', periodic: 'bool' = True) -> None

1D diffusion equation on an Arakawa C-grid.

Diffusion2D

class

Diffusion2D(params: 'Diffusion2DParams', grid: 'CartesianGrid2D', diff: 'Difference2D', mask: 'Mask2D | None') -> None

2D diffusion equation on an Arakawa C-grid.

HelmholtzSolver2D

class

HelmholtzSolver2D(grid: 'CartesianGrid2D', bc_type: 'str' = 'dirichlet', lambda_: 'float' = 1.0) -> None

Solve the 2D Helmholtz equation: :math:(\nabla^2 - \lambda) \phi = f.

IncompressibleNS2D

class

IncompressibleNS2D(params: 'NSParams', grid: 'CartesianGrid2D', diff: 'Difference2D', interp: 'Interpolation2D', advection: 'FVXAdvection2D', mask: 'Mask2D | None', problem: 'str' = 'cavity', poisson_bc: 'str' = 'dst', u_lid: 'float' = 1.0, body_force: 'float' = 0.0, method: 'str' = 'upwind1') -> None

2D incompressible Navier-Stokes (vorticity-streamfunction).

LaplaceSolver2D

class

LaplaceSolver2D(grid: 'CartesianGrid2D', bc_type: 'str' = 'dirichlet') -> None

Solve the 2D Laplace equation: :math:\nabla^2 \phi = 0.

LinearConvection1D

class

LinearConvection1D(params: 'LinearConvection1DParams', grid: 'CartesianGrid1D', diff: 'Difference1D', interp: 'Interpolation1D', mask: 'Mask1D | None', periodic: 'bool' = True) -> None

1D linear convection equation on an Arakawa C-grid.

LinearConvection2D

class

LinearConvection2D(params: 'LinearConvection2DParams', grid: 'CartesianGrid2D', diff: 'Difference2D', interp: 'Interpolation2D', mask: 'Mask2D | None') -> None

2D linear convection on an Arakawa C-grid.

LinearShallowWater1D

class

LinearShallowWater1D(params: 'LinearSW1DParams', consts: 'LinearSW1DPhysConsts', grid: 'CartesianGrid1D', diff: 'Difference1D', interp: 'Interpolation1D', mask: 'Mask1D | None') -> None

1D linear shallow water model on an Arakawa C-grid.

LinearShallowWater2D

class

LinearShallowWater2D(params: 'LinearSW2DParams', consts: 'LinearSW2DPhysConsts', grid: 'CartesianGrid2D', diff: 'Difference2D', interp: 'Interpolation2D', coriolis: 'Coriolis2D', mask: 'Mask2D | None', f_field: "Float[Array, 'Ny Nx']", bc_type: 'str' = 'periodic') -> None

2D linear shallow water model on an Arakawa C-grid.

Lorenz63

class

Lorenz63(params: 'L63Params') -> None

Lorenz '63 three-variable chaotic system.

Lorenz96

class

Lorenz96(params: 'L96Params', advection: 'bool' = True) -> None

Lorenz '96 periodic 1D chaotic system.

Lorenz96t

class

Lorenz96t(params: 'L96TParams', advection: 'bool' = True) -> None

Lorenz '96 two-tier (slow-fast) coupled system.

MultilayerShallowWater2D

class

MultilayerShallowWater2D(params: 'MultilayerSW2DParams', consts: 'MultilayerSW2DPhysConsts', grid: 'CartesianGrid2D', diff: 'Difference2D', interp: 'Interpolation2D', coriolis: 'Coriolis2D', vorticity: 'Vorticity2D', advection: 'FVXAdvection2D', diffusion: 'FVXDiffusion2D', mask: 'Mask2D | None', strat: 'StratificationProfile', modal: 'ModalTransform', f_field: "Float[Array, 'Ny Nx']", f_field_ml: "Float[Array, 'nl Ny Nx']", wind_stress_x: "Float[Array, 'Ny Nx']", wind_stress_y: "Float[Array, 'Ny Nx']", bc_type: 'str' = 'periodic', method: 'str' = 'upwind1') -> None

Multilayer 2D nonlinear shallow water model (vector-invariant form).

NonlinearConvection1D

class

NonlinearConvection1D(grid: 'CartesianGrid1D', advection: 'Advection1D', mask: 'Mask1D | None', periodic: 'bool' = True, method: 'str' = 'upwind1') -> None

1D nonlinear convection (inviscid Burgers) on an Arakawa C-grid.

NonlinearConvection2D

class

NonlinearConvection2D(grid: 'CartesianGrid2D', advection: 'FVXAdvection2D', interp: 'Interpolation2D', mask: 'Mask2D | None', method: 'str' = 'upwind1') -> None

2D nonlinear convection on an Arakawa C-grid.

NonlinearShallowWater1D

class

NonlinearShallowWater1D(params: 'NonlinearSW1DParams', consts: 'NonlinearSW1DPhysConsts', grid: 'CartesianGrid1D', diff: 'Difference1D', interp: 'Interpolation1D', advection: 'Advection1D', mask: 'Mask1D | None', method: 'str' = 'upwind1') -> None

1D nonlinear shallow water model on an Arakawa C-grid.

NonlinearShallowWater2D

class

NonlinearShallowWater2D(params: 'NonlinearSW2DParams', consts: 'NonlinearSW2DPhysConsts', grid: 'CartesianGrid2D', diff: 'Difference2D', interp: 'Interpolation2D', coriolis: 'Coriolis2D', vorticity: 'Vorticity2D', advection: 'FVXAdvection2D', diffusion: 'FVXDiffusion2D', mask: 'Mask2D | None', f_field: "Float[Array, 'Ny Nx']", wind_stress_x: "Float[Array, 'Ny Nx']", wind_stress_y: "Float[Array, 'Ny Nx']", bc_type: 'str' = 'periodic', method: 'str' = 'upwind1') -> None

2D nonlinear shallow water model (vector-invariant form).

PoissonSolver2D

class

PoissonSolver2D(grid: 'CartesianGrid2D', bc_type: 'str' = 'dirichlet') -> None

Solve the 2D Poisson equation: :math:\nabla^2 \phi = f.

ReparameterizedQG

class

ReparameterizedQG(swm: 'MultilayerShallowWater2D', helmholtz_lambdas: 'Array', poisson_bc: 'str' = 'dst') -> None

Reparameterized QG model: multilayer SWM + geostrophic projection.

SphericalQG

class

SphericalQG(params: 'SphericalQGParams', consts: 'SphericalQGPhysConsts', grid: 'SphericalGrid2D', diff: 'SphericalDifference2D', interp: 'Interpolation2D', laplacian: 'SphericalLaplacian2D', advection: 'SphericalAdvection2D', diffusion: 'SphericalDiffusion2D', mask: 'Mask2D | None', f_field: "Float[Array, 'Ny Nx']", wind_forcing: "Float[Array, 'Ny Nx']", method: 'str' = 'upwind1', cg_tol: 'float' = 1e-06, cg_max_steps: 'int' = 500) -> None

Barotropic quasi-geostrophic flow on a sphere.

SphericalSWM

class

SphericalSWM(params: 'SphericalSWMParams', consts: 'SphericalSWMPhysConsts', grid: 'SphericalGrid2D', diff: 'SphericalDifference2D', interp: 'Interpolation2D', vorticity: 'SphericalVorticity2D', advection: 'SphericalAdvection2D', diffusion: 'SphericalDiffusion2D', mask: 'Mask2D | None', f_field: "Float[Array, 'Ny Nx']", wind_stress_x: "Float[Array, 'Ny Nx']", wind_stress_y: "Float[Array, 'Ny Nx']", method: 'str' = 'upwind1') -> None

Shallow water on a sphere, vector-invariant form.

geostrophic_adjustment_2d

function

geostrophic_adjustment_2d(nx: 'int' = 128, ny: 'int' = 128, Lx: 'float' = 1000000.0, Ly: 'float' = 1000000.0, f0: 'float' = 0.0001, H0: 'float' = 100.0, eta_max: 'float' = 1.0) -> 'tuple[LinearShallowWater2D, LinearSW2DState]'

2D geostrophic adjustment: step-function height perturbation.

gravity_wave_1d

function

gravity_wave_1d(nx: 'int' = 400, Lx: 'float' = 1000000.0, g: 'float' = 9.81, H0: 'float' = 100.0, sigma: 'float' = 50000.0) -> 'tuple[LinearShallowWater1D, LinearSW1DState]'

1D gravity wave: Gaussian height perturbation, no rotation.

inertial_oscillation_1d

function

inertial_oscillation_1d(nx: 'int' = 50, Lx: 'float' = 1000000.0, f0: 'float' = 0.0001, u_init: 'float' = 1.0) -> 'tuple[LinearShallowWater1D, LinearSW1DState]'

1D inertial oscillation: uniform initial u, period = 2*pi/f0.

State / parameter / diagnostic companions

Each model carries dataclass companions for its state, differentiable parameters, frozen physical constants, and on-demand diagnostics:

Domain

Spatial and temporal domain descriptors.

Domain

class

Domain(xmin: Union[float, Iterable[float]], xmax: Union[float, Iterable[float]], dx: Union[float, Iterable[float]], Nx: Union[int, Iterable[int]], Lx: Union[float, Iterable[float]])

Domain class for a rectangular domain

TimeDomain

class

TimeDomain(tmin: float, tmax: float, dt: float)

TimeDomain(tmin, tmax, dt)

pipekit Operators

Bridge that exposes somax models as pipekit.Operator stages.

Burgers2DOp

class

Burgers2DOp(nx: 'int' = 64, ny: 'int' = 64, Lx: 'float' = 2.0, Ly: 'float' = 2.0, nu: 'float' = 0.01, method: 'str' = 'upwind1', imex: 'bool' = False, dt: 'float' = 0.001) -> 'None'

Serializable pipekit Operator for the term-based 2D Burgers model.

SomaxModelOp

class

SomaxModelOp(model: 'Any', dt: 'float') -> 'None'

A pipekit Operator wrapping any built somax forward model.

Evaluation Metrics

Reference-free field diagnostics computed on a model’s own grid.

compute_eval_metrics

function

compute_eval_metrics(model: 'Any', state: 'State') -> 'dict[str, float]'

Compute every applicable reference-free metric for (model, state).

geostrophic_imbalance

function

geostrophic_imbalance(model: 'Any', state: 'State', *, interior: 'bool' = True, eps: 'float' = 1e-30) -> "Float[Array, '']"

Dimensionless ageostrophic fraction of a shallow-water-type state.

kinetic_energy

function

kinetic_energy(u: "Float[Array, 'Ny Nx']", v: "Float[Array, 'Ny Nx']", grid: 'Any', *, interior: 'bool' = True) -> "Float[Array, '']"

Domain-integrated kinetic energy 0.5 ∫ (u² + v²) dA.

qg_balance_residual

function

qg_balance_residual(model: 'Any', state: 'State', *, interior: 'bool' = True, eps: 'float' = 1e-30) -> "Float[Array, '']"

Dimensionless PV-inversion residual for a barotropic QG model.

rms_divergence

function

rms_divergence(u: "Float[Array, 'Ny Nx']", v: "Float[Array, 'Ny Nx']", diff: 'Any', *, interior: 'bool' = True) -> "Float[Array, '']"

Root-mean-square horizontal divergence sqrt(<(∇·u)²>).

total_enstrophy

function

total_enstrophy(u: "Float[Array, 'Ny Nx']", v: "Float[Array, 'Ny Nx']", diff: 'Any', grid: 'Any', *, interior: 'bool' = True) -> "Float[Array, '']"

Domain-integrated enstrophy 0.5 ∫ ζ² dA with ζ = ∂v/∂x - ∂u/∂y.

In-JIT Guards

Fail-fast tripwires that halt a run at the offending step.

guard_ceiling

function

guard_ceiling(x: 'Array', *, where: 'str', ceil: 'float') -> 'Array'

Return x unchanged, raising in-JIT if |x| exceeds ceil.

guard_finite

function

guard_finite(x: 'Array', *, where: 'str') -> 'Array'

Return x unchanged, raising in-JIT if it holds any NaN/Inf.

guard_positive

function

guard_positive(x: 'Array', *, where: 'str', floor: 'float' = 0.0) -> 'Array'

Return x unchanged, raising in-JIT if any element <= floor.

Monitors

Chunk-boundary observability for the somax-sim runner.

BaseMonitor

class

BaseMonitor()

Inert base monitor — override only the hooks you care about.

ChunkInfo

class

ChunkInfo(index: 'int', n_chunks: 'int', t0: 'float', t1: 'float', wall_seconds: 'float', is_snapshot: 'bool', stats: 'dict[str, Any]' = <factory>) -> None

Context handed to a monitor at a diagnostic-chunk boundary.

ConservationDriftMonitor

class

ConservationDriftMonitor(rtol_warn: 'float' = 0.01, rtol_fail: 'float | None' = None)

MONITOR (optionally FAIL-HARD): track drift of conserved invariants.

EnergyGrowthMonitor

class

EnergyGrowthMonitor(factor: 'float' = 10.0, hard_factor: 'float | None' = None)

MONITOR (optionally FAIL-HARD): flag large energy jumps between chunks.

Monitor

class

Monitor(*args, **kwargs)

Structural protocol for a simulation monitor.

MonitorVerdict

class

MonitorVerdict(metrics: 'dict[str, float]' = <factory>, messages: 'tuple[str, ...]' = (), terminate: 'bool' = False, reason: 'str | None' = None) -> None

A monitor’s response to a chunk. Inert by default.

NonFiniteMonitor

class

NonFiniteMonitor()

FAIL-HARD: terminate the run if any state field holds NaN/Inf.

SolverHealthMonitor

class

SolverHealthMonitor()

MONITOR: surface diffrax per-chunk solver statistics.

ThroughputMonitor

class

ThroughputMonitor()

MONITOR: report simulated seconds per wallclock second per chunk.

WatchdogMonitor

class

WatchdogMonitor(max_wall_s: 'float')

FAIL-HARD: terminate if cumulative wallclock exceeds a ceiling.

default_monitors

function

default_monitors() -> 'list[BaseMonitor]'

The runner’s default monitor set — preserves legacy behavior.

Solvers

Matrix-free IMEX integration helpers for stiff term models.

imex_solver

function

imex_solver(*, rtol: 'float' = 0.0001, atol: 'float' = 1e-06, gmres_restart: 'int' = 20) -> 'dfx.AbstractSolver'

Build a KenCarp3 IMEX solver with a matrix-free implicit stage.

imex_stepsize_controller

function

imex_stepsize_controller(*, rtol: 'float' = 0.0001, atol: 'float' = 1e-06) -> 'dfx.AbstractStepSizeController'

Build an adaptive PID controller for an IMEX solve.

IO & Persistence

xarray / zarr helpers that round-trip model states and snapshots (requires the sim dependency group).

These symbols live in somax.io and require the optional sim dependency group (uv sync --group sim):

Data Assimilation

Adapters wiring somax models into the filterax / vardax DA stack (requires the optional da dependency group).

These symbols live in somax.da and require the optional da dependency group (uv sync --group da):