Skip to content

NN API

The pyrox_nn subpackage ships uncertainty-aware neural network layers in four families:

  1. Geographic / spherical encoders (re-exported from geonnax) — degree/radian, lon/lat, cyclic, spherical-harmonic, and Slepian preprocessing for geophysical inputs.
  2. Dense / Bayesian-linear layers (pyrox_nn._dense) — reparameterization, Flipout, hierarchical, NCP, DVI, rank-1 ensemble, and variational-dropout variants of Wx + b.
  3. Spectral / GP-flavoured layers — random-feature kernel maps, SNGP and VSSGP heads, deep random-feature expansions.
  4. Ensembles & output heads — BatchEnsemble layers, heteroscedastic Monte-Carlo output heads.
  5. Bayesian Neural Field stack (pyrox_nn._bnf) — five layers that together implement the BNF architecture (Saad et al., Nat. Comms. 2024).
  6. Pure-JAX feature helpers (re-exported from geonnax.basis) — pandas-free building blocks the BNF layers wrap.

See also: Geo encoders for the longitude/latitude and spherical-harmonic API surface.

Dense / Bayesian-linear layers

DenseReparameterization

Bases: PyroxModule

Bayesian dense layer via the reparameterization trick.

Samples weight and bias from learned Gaussian posteriors at every forward pass. Registers NumPyro sample sites so the KL between the variational posterior and the prior is tracked by the ELBO.

\[ W \sim \mathcal{N}(\mu_W, \sigma_W^2), \quad b \sim \mathcal{N}(\mu_b, \sigma_b^2), \quad y = x W + b. \]

Attributes:

Name Type Description
in_features int

Input dimension.

out_features int

Output dimension.

bias bool

Whether to include a bias term.

prior_scale float

Scale of the isotropic Gaussian prior on weights and bias. The prior mean is zero.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
class DenseReparameterization(PyroxModule):
    r"""Bayesian dense layer via the reparameterization trick.

    Samples weight and bias from learned Gaussian posteriors at every
    forward pass. Registers NumPyro sample sites so the KL between the
    variational posterior and the prior is tracked by the ELBO.

    $$
    W \sim \mathcal{N}(\mu_W, \sigma_W^2), \quad
    b \sim \mathcal{N}(\mu_b, \sigma_b^2), \quad
    y = x W + b.
    $$

    Attributes:
        in_features: Input dimension.
        out_features: Output dimension.
        bias: Whether to include a bias term.
        prior_scale: Scale of the isotropic Gaussian prior on weights
            and bias. The prior mean is zero.
        pyrox_name: Explicit scope name for NumPyro site registration.
    """

    in_features: int = eqx.field(static=True)
    out_features: int = eqx.field(static=True)
    bias: bool = eqx.field(static=True, default=True)
    prior_scale: float = eqx.field(static=True, default=1.0)
    pyrox_name: str | None = eqx.field(static=True, default=None)

    @pyrox_method
    def __call__(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D_out"]:
        prior_w = dist.Normal(
            jnp.zeros((self.in_features, self.out_features)),
            self.prior_scale,
        ).to_event(2)
        W = self.pyrox_sample("weight", prior_w)
        out = einx.dot("... din, din dout -> ... dout", x, W)
        if self.bias:
            prior_b = dist.Normal(
                jnp.zeros(self.out_features), self.prior_scale
            ).to_event(1)
            b = self.pyrox_sample("bias", prior_b)
            out = out + b
        return out

DenseFlipout

Bases: PyroxModule

Bayesian dense layer with Flipout sign-flip structure.

Samples weight from the prior and applies per-example Rademacher sign flips to the weight perturbation (Wen et al., 2018). Under a NumPyro guide that learns the posterior mean, the sign flips decorrelate gradient estimates across minibatch examples.

In model mode (no guide) this is equivalent to DenseReparameterization — the Flipout variance reduction activates when a guide provides a posterior centered at a learned mean.

Attributes:

Name Type Description
in_features int

Input dimension.

out_features int

Output dimension.

bias bool

Whether to include a bias term.

prior_scale float

Scale of the isotropic Gaussian prior.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
class DenseFlipout(PyroxModule):
    r"""Bayesian dense layer with Flipout sign-flip structure.

    Samples weight from the prior and applies per-example Rademacher
    sign flips to the weight perturbation (Wen et al., 2018). Under a
    NumPyro guide that learns the posterior mean, the sign flips
    decorrelate gradient estimates across minibatch examples.

    In model mode (no guide) this is equivalent to
    `DenseReparameterization` — the Flipout variance reduction
    activates when a guide provides a posterior centered at a learned
    mean.

    Attributes:
        in_features: Input dimension.
        out_features: Output dimension.
        bias: Whether to include a bias term.
        prior_scale: Scale of the isotropic Gaussian prior.
        pyrox_name: Explicit scope name for NumPyro site registration.
    """

    in_features: int = eqx.field(static=True)
    out_features: int = eqx.field(static=True)
    bias: bool = eqx.field(static=True, default=True)
    prior_scale: float = eqx.field(static=True, default=1.0)
    pyrox_name: str | None = eqx.field(static=True, default=None)

    @pyrox_method
    def __call__(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D_out"]:
        prior_w = dist.Normal(
            jnp.zeros((self.in_features, self.out_features)),
            self.prior_scale,
        ).to_event(2)
        W = self.pyrox_sample("weight", prior_w)
        out = einx.dot("... din, din dout -> ... dout", x, W)

        if self.bias:
            prior_b = dist.Normal(
                jnp.zeros(self.out_features), self.prior_scale
            ).to_event(1)
            b = self.pyrox_sample("bias", prior_b)
            out = out + b
        return out

DenseVariational

Bases: PyroxModule

Dense layer with a user-supplied prior factory.

Provides flexibility over the weight prior by accepting a callable that builds the prior distribution given the layer shape. The model samples from the prior; the posterior is handled by a NumPyro guide (e.g., AutoNormal).

Attributes:

Name Type Description
in_features int

Input dimension.

out_features int

Output dimension.

make_prior Callable[..., Any]

Callable (in_features, out_features) -> Distribution.

bias bool

Whether to include a bias term.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
class DenseVariational(PyroxModule):
    r"""Dense layer with a user-supplied prior factory.

    Provides flexibility over the weight prior by accepting a callable
    that builds the prior distribution given the layer shape. The
    model samples from the prior; the posterior is handled by a NumPyro
    guide (e.g., ``AutoNormal``).

    Attributes:
        in_features: Input dimension.
        out_features: Output dimension.
        make_prior: Callable ``(in_features, out_features) -> Distribution``.
        bias: Whether to include a bias term.
        pyrox_name: Explicit scope name for NumPyro site registration.
    """

    in_features: int = eqx.field(static=True)
    out_features: int = eqx.field(static=True)
    make_prior: Callable[..., Any] = eqx.field(static=True)
    bias: bool = eqx.field(static=True, default=True)
    pyrox_name: str | None = eqx.field(static=True, default=None)

    @pyrox_method
    def __call__(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D_out"]:
        prior = self.make_prior(self.in_features, self.out_features)
        W = self.pyrox_sample("weight", prior)
        out = einx.dot("... din, din dout -> ... dout", x, W)
        if self.bias:
            b = self.pyrox_sample(
                "bias",
                dist.Normal(jnp.zeros(self.out_features), 1.0).to_event(1),
            )
            out = out + b
        return out

DenseDVI

Bases: PyroxModule

Deterministic Variational Inference dense layer (Wu et al., 2018).

Propagates a Gaussian distribution through the linear layer analytically — there is no Monte Carlo sampling. The input is a diagonal-covariance Gaussian \((\mu_x, \sigma_x^2)\), the output is the (still-diagonal) Gaussian \((\mu_y, \sigma_y^2)\) induced by an independent-Gaussian variational posterior \(q(W) = \mathcal{N}(M, S)\) over the weights and a separate diagonal posterior on the bias.

With weight posterior mean \(M\) (shape \(D_\mathrm{in}\times D_\mathrm{out}\)) and per-element posterior variance \(S\) (same shape):

\[ \mu_y = \mu_x M, \qquad \sigma_y^2 = \sigma_x^2 (M \circ M) + (\mu_x^{\circ 2} + \sigma_x^2)\,S, \]

plus the bias mean / variance if enabled. Compared to MC estimators, DVI gives zero-variance gradients of the ELBO at the cost of propagating second-order statistics layer by layer (so it only really pays off when all dense layers in a block are DVI; a single DVI layer in a sampling stack just adds bookkeeping).

The KL between the diagonal-Gaussian variational posterior and a fixed isotropic Gaussian prior \(p(W) = \mathcal{N}(0, \pi^2)\) is closed-form and is registered with numpyro.factor so SVI's Trace_ELBO picks it up:

\[ \mathrm{KL}\!\bigl[\mathcal{N}(M, S) \,\big\|\, \mathcal{N}(0, \pi^2)\bigr] = \sum_{ij}\Bigl[ \log \pi - \tfrac12\log S_{ij} + \frac{S_{ij} + M_{ij}^2}{2\pi^2} - \tfrac12 \Bigr]. \]
Plate semantics

Same as the rest of the pyrox Bayesian dense family — call this layer outside numpyro.plate("data", ..., subsample_size=...). The KL is a weight-prior term: it sums over the weight and bias matrices, not over the batch, so it's a single scalar per layer. numpyro.factor is still a sample-type site, though, and putting it inside a subsampled plate would broadcast the scalar to the plate dim and apply scale = N/B — the same over-counting trap that affects every per-layer numpyro.factor. Keep this layer at the top of the model (or outside any data plate) and only plate the observation likelihood:

def model(x, y=None):
    mean, var = dvi(x_mean, x_var)        # KL emitted here
    with numpyro.plate("data", x.shape[0]):
        numpyro.sample("obs",
            dist.Normal(mean, jnp.sqrt(var)), obs=y)

Attributes:

Name Type Description
in_features int

Input dimension \(D_\mathrm{in}\).

out_features int

Output dimension \(D_\mathrm{out}\).

bias bool

Whether to include a diagonal-Gaussian bias.

prior_scale float

Std \(\pi\) of the isotropic Gaussian prior.

init_log_var float

Initial value for the log posterior variance (a small negative number keeps initial draws tight).

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Examples:

>>> import jax.numpy as jnp
>>> from numpyro import handlers
>>> dvi = DenseDVI(in_features=3, out_features=2, pyrox_name="dvi")
>>> mean = jnp.ones((4, 3))
>>> var = 0.1 * jnp.ones((4, 3))
>>> with handlers.seed(rng_seed=0):
...     out_mean, out_var = dvi(mean, var)
>>> out_mean.shape, out_var.shape
((4, 2), (4, 2))
References

Wu, A., Nowozin, S., Meeds, E., Turner, R. E., Hernández-Lobato, J. M., & Gaunt, A. L. (2018). Deterministic Variational Inference for Robust Bayesian Neural Networks. ICLR.

Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
class DenseDVI(PyroxModule):
    r"""Deterministic Variational Inference dense layer (Wu et al., 2018).

    Propagates a *Gaussian distribution* through the linear layer
    analytically — there is no Monte Carlo sampling. The input is a
    diagonal-covariance Gaussian $(\mu_x, \sigma_x^2)$, the
    output is the (still-diagonal) Gaussian $(\mu_y, \sigma_y^2)$
    induced by an independent-Gaussian variational posterior
    $q(W) = \mathcal{N}(M, S)$ over the weights and a separate
    diagonal posterior on the bias.

    With weight posterior mean $M$
    (shape $D_\mathrm{in}\times D_\mathrm{out}$) and per-element
    posterior variance $S$ (same shape):

    $$
    \mu_y = \mu_x M, \qquad
    \sigma_y^2 = \sigma_x^2 (M \circ M)
               + (\mu_x^{\circ 2} + \sigma_x^2)\,S,
    $$

    plus the bias mean / variance if enabled. Compared to MC
    estimators, DVI gives zero-variance gradients of the ELBO at the
    cost of propagating second-order statistics layer by layer (so it
    only really pays off when *all* dense layers in a block are DVI;
    a single DVI layer in a sampling stack just adds bookkeeping).

    The KL between the diagonal-Gaussian variational posterior and a
    fixed isotropic Gaussian prior $p(W) = \mathcal{N}(0, \pi^2)$
    is closed-form and is registered with `numpyro.factor` so
    SVI's ``Trace_ELBO`` picks it up:

    $$
    \mathrm{KL}\!\bigl[\mathcal{N}(M, S) \,\big\|\, \mathcal{N}(0, \pi^2)\bigr]
    = \sum_{ij}\Bigl[
        \log \pi - \tfrac12\log S_{ij}
        + \frac{S_{ij} + M_{ij}^2}{2\pi^2} - \tfrac12
      \Bigr].
    $$

    Plate semantics:
        Same as the rest of the pyrox Bayesian dense family — call
        this layer **outside** ``numpyro.plate("data", ..., subsample_size=...)``.
        The KL is a *weight-prior* term: it sums over the weight and
        bias matrices, not over the batch, so it's a single scalar
        per layer. ``numpyro.factor`` is still a sample-type site,
        though, and putting it inside a subsampled plate would broadcast
        the scalar to the plate dim and apply ``scale = N/B`` — the
        same over-counting trap that affects every per-layer
        ``numpyro.factor``. Keep this layer at the top of the model
        (or outside any data plate) and only plate the observation
        likelihood:

            def model(x, y=None):
                mean, var = dvi(x_mean, x_var)        # KL emitted here
                with numpyro.plate("data", x.shape[0]):
                    numpyro.sample("obs",
                        dist.Normal(mean, jnp.sqrt(var)), obs=y)

    Attributes:
        in_features: Input dimension $D_\mathrm{in}$.
        out_features: Output dimension $D_\mathrm{out}$.
        bias: Whether to include a diagonal-Gaussian bias.
        prior_scale: Std $\pi$ of the isotropic Gaussian prior.
        init_log_var: Initial value for the log posterior variance
            (a small negative number keeps initial draws tight).
        pyrox_name: Explicit scope name for NumPyro site registration.

    Examples:
        >>> import jax.numpy as jnp
        >>> from numpyro import handlers
        >>> dvi = DenseDVI(in_features=3, out_features=2, pyrox_name="dvi")
        >>> mean = jnp.ones((4, 3))
        >>> var = 0.1 * jnp.ones((4, 3))
        >>> with handlers.seed(rng_seed=0):
        ...     out_mean, out_var = dvi(mean, var)
        >>> out_mean.shape, out_var.shape
        ((4, 2), (4, 2))

    References:
        Wu, A., Nowozin, S., Meeds, E., Turner, R. E., Hernández-Lobato,
        J. M., & Gaunt, A. L. (2018). *Deterministic Variational
        Inference for Robust Bayesian Neural Networks.* ICLR.
    """

    in_features: int = eqx.field(static=True)
    out_features: int = eqx.field(static=True)
    bias: bool = eqx.field(static=True, default=True)
    prior_scale: float = eqx.field(static=True, default=1.0)
    init_log_var: float = eqx.field(static=True, default=-3.0)
    pyrox_name: str | None = eqx.field(static=True, default=None)

    def __post_init__(self) -> None:
        if self.in_features <= 0 or self.out_features <= 0:
            raise ValueError(
                "in_features and out_features must be > 0; "
                f"got {self.in_features=}, {self.out_features=}."
            )
        if self.prior_scale <= 0:
            raise ValueError(f"prior_scale must be > 0; got {self.prior_scale}.")

    @pyrox_method
    def __call__(
        self,
        mean: Float[Array, "*batch D_in"],
        var: Float[Array, "*batch D_in"],
    ) -> tuple[Float[Array, "*batch D_out"], Float[Array, "*batch D_out"]]:
        if mean.shape != var.shape:
            raise ValueError(f"mean.shape {mean.shape} != var.shape {var.shape}.")
        if mean.shape[-1] != self.in_features:
            raise ValueError(
                f"mean.shape[-1] = {mean.shape[-1]} does not match "
                f"in_features = {self.in_features}."
            )

        W_mean = self.pyrox_param(
            "weight_mean", jnp.zeros((self.in_features, self.out_features))
        )
        W_log_var = self.pyrox_param(
            "weight_log_var",
            jnp.full(
                (self.in_features, self.out_features),
                float(self.init_log_var),
            ),
        )
        W_var = jnp.exp(W_log_var)

        out_mean = einx.dot("... din, din dout -> ... dout", mean, W_mean)
        out_var = einx.dot("... din, din dout -> ... dout", var, W_mean**2) + einx.dot(
            "... din, din dout -> ... dout", mean**2 + var, W_var
        )

        if self.bias:
            b_mean = self.pyrox_param("bias_mean", jnp.zeros(self.out_features))
            b_log_var = self.pyrox_param(
                "bias_log_var",
                jnp.full((self.out_features,), float(self.init_log_var)),
            )
            out_mean = out_mean + b_mean
            out_var = out_var + jnp.exp(b_log_var)

        # Closed-form KL[N(M, S) || N(0, prior_scale^2)] over weights and bias.
        # Use math.log + jnp.asarray cast so prior constants pick up the
        # same dtype as the params (avoid silent float64 promotion under
        # jax_enable_x64 + float32 params).
        log_prior_scale = jnp.asarray(math.log(self.prior_scale), dtype=W_mean.dtype)
        prior_var = jnp.asarray(self.prior_scale**2, dtype=W_mean.dtype)
        kl = jnp.sum(
            _diag_gaussian_kl(
                W_mean,
                W_var,
                W_log_var,
                prior_mean=0.0,
                log_prior_scale=log_prior_scale,
                prior_var=prior_var,
            )
        )
        if self.bias:
            b_var = jnp.exp(b_log_var)
            kl = kl + jnp.sum(
                _diag_gaussian_kl(
                    b_mean,
                    b_var,
                    b_log_var,
                    prior_mean=0.0,
                    log_prior_scale=log_prior_scale,
                    prior_var=prior_var,
                )
            )
        # Add -KL to the model log density. This is a per-layer scalar
        # (sums over the weight / bias matrices, not the batch) — its
        # emission is independent of any data plate the layer lives in.
        numpyro.factor(self._pyrox_fullname("kl"), -kl)
        return out_mean, out_var

DenseHierarchical

Bases: PyroxModule

Hierarchical Bayesian dense layer with multiplicative shrinkage.

Decomposes the effective weight matrix into a deterministic base \(\theta \in \mathbb{R}^{D_\mathrm{in} \times D_\mathrm{out}}\) multiplied row-wise by a per-input-unit local scale \(z^{(\mathrm{loc})} \in \mathbb{R}^{D_\mathrm{in}}\) and an overall global scale \(z^{(\mathrm{glob})} \in \mathbb{R}\),

\[ W_{ij} = \theta_{ij} \cdot z_i^{(\mathrm{loc})} \cdot z^{(\mathrm{glob})}, \]

with isotropic Gaussian priors centred at one,

\[ z_i^{(\mathrm{loc})} \sim \mathcal{N}(1, \sigma_\mathrm{loc}^2), \qquad z^{(\mathrm{glob})} \sim \mathcal{N}(1, \sigma_\mathrm{glob}^2). \]

The local scale prunes individual input units (a column of \(\theta\) whose z_loc posterior concentrates near zero is effectively switched off) while the global scale modulates the overall layer activation — the same hierarchical-shrinkage structure used by horseshoe-style BNNs (Louizos et al., 2017). Both scales are pyrox_sample sites so any standard NumPyro guide (AutoNormal, etc.) drives the variational posterior; the deterministic base \(\theta\) and bias are pyrox_param.

Plate semantics

Same as the rest of pyrox_nn's Bayesian dense layers — call outside numpyro.plate("data", ..., subsample_size=...) and only plate the observation likelihood, otherwise the per-layer prior log-probabilities of z_loc and z_glob get scaled by the subsample ratio.

Attributes:

Name Type Description
in_features int

Input dimension \(D_\mathrm{in}\).

out_features int

Output dimension \(D_\mathrm{out}\).

bias bool

Whether to include a deterministic bias term.

prior_local_scale float

Std \(\sigma_\mathrm{loc}\) of the local scale prior.

prior_global_scale float

Std \(\sigma_\mathrm{glob}\) of the global scale prior.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Examples:

>>> import jax.numpy as jnp
>>> from numpyro import handlers
>>> layer = DenseHierarchical(
...     in_features=4, out_features=2, pyrox_name="hier"
... )
>>> x = jnp.ones((3, 4))
>>> with handlers.seed(rng_seed=0):
...     y = layer(x)
>>> y.shape
(3, 2)
References

Louizos, C., Ullrich, K., & Welling, M. (2017). Bayesian Compression for Deep Learning. NeurIPS.

Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
class DenseHierarchical(PyroxModule):
    r"""Hierarchical Bayesian dense layer with multiplicative shrinkage.

    Decomposes the effective weight matrix into a deterministic base
    $\theta \in \mathbb{R}^{D_\mathrm{in} \times D_\mathrm{out}}$
    multiplied row-wise by a per-input-unit local scale
    $z^{(\mathrm{loc})} \in \mathbb{R}^{D_\mathrm{in}}$ and an
    overall global scale $z^{(\mathrm{glob})} \in \mathbb{R}$,

    $$
    W_{ij} = \theta_{ij} \cdot z_i^{(\mathrm{loc})}
             \cdot z^{(\mathrm{glob})},
    $$

    with isotropic Gaussian priors centred at one,

    $$
    z_i^{(\mathrm{loc})} \sim \mathcal{N}(1, \sigma_\mathrm{loc}^2),
    \qquad
    z^{(\mathrm{glob})} \sim \mathcal{N}(1, \sigma_\mathrm{glob}^2).
    $$

    The local scale prunes individual input units (a column of
    $\theta$ whose ``z_loc`` posterior concentrates near zero is
    effectively switched off) while the global scale modulates the
    overall layer activation — the same hierarchical-shrinkage
    structure used by horseshoe-style BNNs (Louizos et al., 2017).
    Both scales are ``pyrox_sample`` sites so any standard NumPyro
    guide (``AutoNormal``, etc.) drives the variational posterior; the
    deterministic base $\theta$ and bias are ``pyrox_param``.

    Plate semantics:
        Same as the rest of ``pyrox_nn``'s Bayesian dense layers — call
        outside ``numpyro.plate("data", ..., subsample_size=...)`` and
        only plate the observation likelihood, otherwise the
        per-layer prior log-probabilities of ``z_loc`` and ``z_glob``
        get scaled by the subsample ratio.

    Attributes:
        in_features: Input dimension $D_\mathrm{in}$.
        out_features: Output dimension $D_\mathrm{out}$.
        bias: Whether to include a deterministic bias term.
        prior_local_scale: Std $\sigma_\mathrm{loc}$ of the local
            scale prior.
        prior_global_scale: Std $\sigma_\mathrm{glob}$ of the
            global scale prior.
        pyrox_name: Explicit scope name for NumPyro site registration.

    Examples:
        >>> import jax.numpy as jnp
        >>> from numpyro import handlers
        >>> layer = DenseHierarchical(
        ...     in_features=4, out_features=2, pyrox_name="hier"
        ... )
        >>> x = jnp.ones((3, 4))
        >>> with handlers.seed(rng_seed=0):
        ...     y = layer(x)
        >>> y.shape
        (3, 2)

    References:
        Louizos, C., Ullrich, K., & Welling, M. (2017). *Bayesian
        Compression for Deep Learning.* NeurIPS.
    """

    in_features: int = eqx.field(static=True)
    out_features: int = eqx.field(static=True)
    bias: bool = eqx.field(static=True, default=True)
    prior_local_scale: float = 0.1
    prior_global_scale: float = 0.1
    pyrox_name: str | None = eqx.field(static=True, default=None)

    def __post_init__(self) -> None:
        if self.prior_local_scale <= 0:
            raise ValueError(
                f"prior_local_scale must be > 0; got {self.prior_local_scale}."
            )
        if self.prior_global_scale <= 0:
            raise ValueError(
                f"prior_global_scale must be > 0; got {self.prior_global_scale}."
            )

    @pyrox_method
    def __call__(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D_out"]:
        theta = self.pyrox_param(
            "theta", jnp.zeros((self.in_features, self.out_features))
        )
        z_loc = self.pyrox_sample(
            "z_local",
            dist.Normal(jnp.ones(self.in_features), self.prior_local_scale).to_event(1),
        )
        z_glob = self.pyrox_sample(
            "z_global", dist.Normal(1.0, self.prior_global_scale)
        )
        # Effective weight uses the broadcasted multiplicative scaling.
        # Equivalent (and slightly cheaper) to scaling x first then matmul:
        #   y = ((x * z_loc) @ theta) * z_glob.
        out = einx.dot("... din, din dout -> ... dout", x * z_loc, theta) * z_glob
        if self.bias:
            b = self.pyrox_param("b", jnp.zeros(self.out_features))
            out = out + b
        return out

DenseVariationalDropout

Bases: PyroxModule

Sparse variational dropout dense layer.

Implements variational dropout (Kingma et al., 2015) extended by Molchanov et al. (2017) to a log-uniform prior that enables automatic sparsification via per-weight learnable dropout rates. The variational posterior on weights is

\[ q(W_{ij} \mid \theta_{ij}, \alpha_{ij}) = \mathcal{N}\!\bigl(\theta_{ij},\; \alpha_{ij}\,\theta_{ij}^2\bigr). \]

Forward passes use the local reparameterization trick — the pre-activation distribution is closed-form and the noise is sampled once per output unit per batch element rather than once per weight:

\[ \gamma = X\theta, \quad \delta = X^{\circ 2}\,(\alpha \circ \theta^{\circ 2}), \quad Y = \gamma + \sqrt{\delta} \circ \epsilon, \quad \epsilon \sim \mathcal{N}(0, I). \]

The KL between the posterior and the log-uniform prior is approximated analytically (Molchanov et al., 2017) and added to the NumPyro trace via numpyro.factor. SVI then optimizes

\[ \mathcal{L} = \mathbb{E}_q[\log p(y \mid f)] - \mathrm{KL}\bigl[q\,\|\,p\bigr]. \]

Weights with log_alpha > threshold (default 3.0, dropout rate ~0.95) are effectively pruned; inspect the trained pattern via sparsity.

Plate semantics

The KL contribution is registered via numpyro.factor, which is itself a sample-type site and therefore subject to numpyro.plate scaling. To keep the per-layer KL counted once (not scaled by the data-plate's subsample ratio), call the layer outside any plate("data", ..., subsample_size=...) block — the standard pyrox / NumPyro convention for global Bayesian parameters. Plate only the observation likelihood.

Correct (forward outside the data plate):

def model(x, y=None):
    layer = DenseVariationalDropout(in_features=D, out_features=1)
    f = layer(x).squeeze(-1)               # KL emitted here
    with numpyro.plate("data", x.shape[0]):
        numpyro.sample("obs", dist.Normal(f, 0.5), obs=y)

Incorrect (forward inside a subsampled data plate scales KL by N / subsample_size):

def model(x, y=None):
    with numpyro.plate("data", N, subsample_size=B) as idx:
        f = layer(x[idx]).squeeze(-1)      # ⚠ scales KL
        numpyro.sample("obs", ...)

Attributes:

Name Type Description
in_features int

Input dimension.

out_features int

Output dimension.

bias bool

Whether to include a bias term.

log_alpha_init float

Initial value for log_alpha (typically a small negative number, e.g., -5.0).

threshold float

log_alpha threshold for declaring a weight pruned.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Examples:

>>> import jax
>>> import jax.numpy as jnp
>>> from numpyro import handlers
>>> layer = DenseVariationalDropout(
...     in_features=4, out_features=2, pyrox_name="vd"
... )
>>> x = jnp.ones((3, 4))
>>> with handlers.seed(rng_seed=0):
...     y = layer(x)
>>> y.shape
(3, 2)
Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
class DenseVariationalDropout(PyroxModule):
    r"""Sparse variational dropout dense layer.

    Implements variational dropout (Kingma et al., 2015) extended by
    Molchanov et al. (2017) to a log-uniform prior that enables
    automatic sparsification via per-weight learnable dropout rates.
    The variational posterior on weights is

    $$
    q(W_{ij} \mid \theta_{ij}, \alpha_{ij}) =
    \mathcal{N}\!\bigl(\theta_{ij},\; \alpha_{ij}\,\theta_{ij}^2\bigr).
    $$

    Forward passes use the *local reparameterization trick* — the
    pre-activation distribution is closed-form and the noise is sampled
    once per output unit per batch element rather than once per weight:

    $$
    \gamma = X\theta, \quad
    \delta = X^{\circ 2}\,(\alpha \circ \theta^{\circ 2}), \quad
    Y = \gamma + \sqrt{\delta} \circ \epsilon, \quad
    \epsilon \sim \mathcal{N}(0, I).
    $$

    The KL between the posterior and the log-uniform prior is
    approximated analytically (Molchanov et al., 2017) and added to the
    NumPyro trace via `numpyro.factor`. SVI then optimizes

    $$
    \mathcal{L} = \mathbb{E}_q[\log p(y \mid f)] - \mathrm{KL}\bigl[q\,\|\,p\bigr].
    $$

    Weights with ``log_alpha > threshold`` (default 3.0, dropout rate
    ~0.95) are effectively pruned; inspect the trained pattern via
    `sparsity`.

    Plate semantics:
        The KL contribution is registered via `numpyro.factor`,
        which is itself a sample-type site and therefore subject to
        ``numpyro.plate`` scaling. To keep the per-layer KL counted
        once (not scaled by the data-plate's subsample ratio), call
        the layer **outside** any ``plate("data", ..., subsample_size=...)``
        block — the standard pyrox / NumPyro convention for global
        Bayesian parameters. Plate only the observation likelihood.

        Correct (forward outside the data plate):

            def model(x, y=None):
                layer = DenseVariationalDropout(in_features=D, out_features=1)
                f = layer(x).squeeze(-1)               # KL emitted here
                with numpyro.plate("data", x.shape[0]):
                    numpyro.sample("obs", dist.Normal(f, 0.5), obs=y)

        Incorrect (forward inside a subsampled data plate scales KL by
        ``N / subsample_size``):

            def model(x, y=None):
                with numpyro.plate("data", N, subsample_size=B) as idx:
                    f = layer(x[idx]).squeeze(-1)      # ⚠ scales KL
                    numpyro.sample("obs", ...)

    Attributes:
        in_features: Input dimension.
        out_features: Output dimension.
        bias: Whether to include a bias term.
        log_alpha_init: Initial value for ``log_alpha`` (typically a
            small negative number, e.g., ``-5.0``).
        threshold: ``log_alpha`` threshold for declaring a weight pruned.
        pyrox_name: Explicit scope name for NumPyro site registration.

    Examples:
        >>> import jax
        >>> import jax.numpy as jnp
        >>> from numpyro import handlers
        >>> layer = DenseVariationalDropout(
        ...     in_features=4, out_features=2, pyrox_name="vd"
        ... )
        >>> x = jnp.ones((3, 4))
        >>> with handlers.seed(rng_seed=0):
        ...     y = layer(x)
        >>> y.shape
        (3, 2)
    """

    in_features: int = eqx.field(static=True)
    out_features: int = eqx.field(static=True)
    bias: bool = eqx.field(static=True, default=True)
    log_alpha_init: float = eqx.field(static=True, default=-5.0)
    threshold: float = eqx.field(static=True, default=3.0)
    pyrox_name: str | None = eqx.field(static=True, default=None)

    @pyrox_method
    def __call__(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D_out"]:
        theta = self.pyrox_param(
            "theta",
            jnp.zeros((self.in_features, self.out_features)),
        )
        log_alpha = self.pyrox_param(
            "log_alpha",
            jnp.full(
                (self.in_features, self.out_features),
                float(self.log_alpha_init),
            ),
        )
        log_alpha_clamped = jnp.clip(log_alpha, _VD_LOG_ALPHA_MIN, _VD_LOG_ALPHA_MAX)
        alpha = jnp.exp(log_alpha_clamped)

        gamma = einx.dot("... din, din dout -> ... dout", x, theta)
        delta = einx.dot("... din, din dout -> ... dout", x**2, alpha * theta**2)
        # Floor to a tiny positive value: keeps the sqrt gradient finite at
        # delta = 0 without injecting visible noise (sqrt(1e-30) ≈ 1e-15).
        std = jnp.sqrt(jnp.maximum(delta, 1e-30))
        # numpyro.prng_key returns Array | None at the type level, but is
        # always an Array inside a `seed` handler — which is required for
        # the pyrox_param/factor calls above to succeed in any case.
        key = cast(JaxArray, numpyro.prng_key())
        eps = jax.random.normal(key, gamma.shape, dtype=gamma.dtype)
        out = gamma + std * eps

        if self.bias:
            b = self.pyrox_param("bias", jnp.zeros(self.out_features))
            out = out + b

        numpyro.factor(
            self._pyrox_fullname("kl"),
            jnp.sum(_vd_neg_kl(log_alpha)),
        )
        return out

    def sparsity(self, log_alpha: Float[Array, "D_in D_out"]) -> Float[Array, ""]:
        """Fraction of weights with ``log_alpha > threshold``.

        Pass the trained ``log_alpha`` parameter, typically retrieved
        from the SVI param store under ``f"{pyrox_name}.log_alpha"``.
        """
        return jnp.mean((log_alpha > self.threshold).astype(log_alpha.dtype))

sparsity(log_alpha: Float[Array, 'D_in D_out']) -> Float[Array, '']

Fraction of weights with log_alpha > threshold.

Pass the trained log_alpha parameter, typically retrieved from the SVI param store under f"{pyrox_name}.log_alpha".

Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
def sparsity(self, log_alpha: Float[Array, "D_in D_out"]) -> Float[Array, ""]:
    """Fraction of weights with ``log_alpha > threshold``.

    Pass the trained ``log_alpha`` parameter, typically retrieved
    from the SVI param store under ``f"{pyrox_name}.log_alpha"``.
    """
    return jnp.mean((log_alpha > self.threshold).astype(log_alpha.dtype))

DenseNCP

Bases: PyroxModule

Noise Contrastive Prior dense layer (Hafner et al., 2019).

Decomposes a dense layer into a prior-regularized backbone plus a scaled stochastic perturbation:

\[ y = \underbrace{x W_d + b_d}_{\text{backbone}} + \underbrace{\sigma \cdot (x W_s + b_s)}_{\text{perturbation}}, \]

where all weights are pyrox_sample sites with Gaussian priors and \(\sigma\) has a LogNormal prior. The backbone carries the bulk of the signal; the perturbation branch adds calibrated uncertainty that can be trained via a noise contrastive objective.

Attributes:

Name Type Description
in_features int

Input dimension.

out_features int

Output dimension.

init_scale float

Initial value for the perturbation scale \(\sigma\).

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
class DenseNCP(PyroxModule):
    r"""Noise Contrastive Prior dense layer (Hafner et al., 2019).

    Decomposes a dense layer into a prior-regularized backbone plus a
    scaled stochastic perturbation:

    $$
    y = \underbrace{x W_d + b_d}_{\text{backbone}}
      + \underbrace{\sigma \cdot (x W_s + b_s)}_{\text{perturbation}},
    $$

    where all weights are ``pyrox_sample`` sites with Gaussian priors
    and $\sigma$ has a ``LogNormal`` prior. The backbone carries
    the bulk of the signal; the perturbation branch adds calibrated
    uncertainty that can be trained via a noise contrastive objective.

    Attributes:
        in_features: Input dimension.
        out_features: Output dimension.
        init_scale: Initial value for the perturbation scale
            $\sigma$.
        pyrox_name: Explicit scope name for NumPyro site registration.
    """

    in_features: int = eqx.field(static=True)
    out_features: int = eqx.field(static=True)
    init_scale: float = eqx.field(static=True, default=1.0)
    pyrox_name: str | None = eqx.field(static=True, default=None)

    @pyrox_method
    def __call__(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D_out"]:
        W_d = self.pyrox_sample(
            "weight_det",
            dist.Normal(jnp.zeros((self.in_features, self.out_features)), 1.0).to_event(
                2
            ),
        )
        b_d = self.pyrox_sample(
            "bias_det",
            dist.Normal(jnp.zeros(self.out_features), 1.0).to_event(1),
        )
        det = einx.dot("... din, din dout -> ... dout", x, W_d) + b_d

        W_s = self.pyrox_sample(
            "weight_stoch",
            dist.Normal(jnp.zeros((self.in_features, self.out_features)), 1.0).to_event(
                2
            ),
        )
        b_s = self.pyrox_sample(
            "bias_stoch",
            dist.Normal(jnp.zeros(self.out_features), 1.0).to_event(1),
        )
        scale = self.pyrox_sample(
            "scale",
            dist.LogNormal(jnp.log(jnp.maximum(jnp.array(self.init_scale), 1e-6)), 1.0),
        )
        stoch = scale * (einx.dot("... din, din dout -> ... dout", x, W_s) + b_s)

        return det + stoch

NCPContinuousPerturb

Bases: Module

Input perturbation for the Noise Contrastive Prior pattern.

Adds Gaussian noise scaled by a fixed positive scale to the input:

\[ \tilde{x} = x + \sigma \epsilon, \qquad \epsilon \sim \mathcal{N}(0, I). \]

Place before a deterministic network to inject input uncertainty; pair with a Bayesian DenseNCP head for the full NCP architecture (Hafner et al., 2019).

Stochasticity comes from the explicit PRNG key argument.

Attributes:

Name Type Description
scale float | Float[Array, '']

Perturbation scale \(\sigma\).

Examples:

>>> import jax.numpy as jnp
>>> import jax.random as jr
>>> perturb = NCPContinuousPerturb(scale=0.5)
>>> x = jnp.zeros(3)
>>> out = perturb(x, key=jr.PRNGKey(0))  # x̃ = x + σ·ε, ε ~ N(0, I)
>>> out.shape
(3,)
Source code in .venv/lib/python3.12/site-packages/geonnax/ncp.py
class NCPContinuousPerturb(eqx.Module):
    r"""Input perturbation for the Noise Contrastive Prior pattern.

    Adds Gaussian noise scaled by a fixed positive scale to the input:

    $$
    \tilde{x} = x + \sigma \epsilon, \qquad
    \epsilon \sim \mathcal{N}(0, I).
    $$


    Place before a deterministic network to inject input uncertainty;
    pair with a Bayesian ``DenseNCP`` head for the full NCP
    architecture (Hafner et al., 2019).

    Stochasticity comes from the explicit PRNG ``key`` argument.

    Attributes:
        scale: Perturbation scale $\sigma$.

    Examples:
        >>> import jax.numpy as jnp
        >>> import jax.random as jr
        >>> perturb = NCPContinuousPerturb(scale=0.5)
        >>> x = jnp.zeros(3)
        >>> out = perturb(x, key=jr.PRNGKey(0))  # x̃ = x + σ·ε, ε ~ N(0, I)
        >>> out.shape
        (3,)
    """

    scale: float | Float[Array, ""] = 1.0

    def __call__(
        self,
        x: Float[Array, " D"],
        *,
        key: Array,
    ) -> Float[Array, " D"]:
        # x̃ = x + σ·ε with ε ~ N(0, I), same shape/dtype as x → (D,).
        eps = jax.random.normal(key, x.shape, dtype=x.dtype)
        return x + self.scale * eps

NCPNormalOutput

Bases: PyroxModule

Output-side Noise Contrastive Prior layer (Hafner et al., 2018).

Completes the NCP pattern in pyrox_nn: pair with NCPContinuousPerturb at the input and a heteroscedastic network (e.g. an MLP terminating in a mean head and a positive-std head — a softplus or exp of a learned log-scale) so the network produces predictions for both the clean batch and the input-perturbed noisy batch. Given the noisy batch's predictive distribution \(\mathcal{N}(\hat{y}_n, \hat{\sigma}_n^2)\), this layer adds the analytic NCP regulariser

\[ \mathcal{L}_\mathrm{NCP} = \sum_{n} \mathrm{KL}\!\bigl[\mathcal{N}(\hat{y}_n, \hat{\sigma}_n^2) \;\big\|\; \mathcal{N}(\mu_\mathrm{prior}, \sigma_\mathrm{prior}^2)\bigr] \]

to the model log density via numpyro.factor. Pulling the noisy-input predictive distribution toward the fixed prior away from the training distribution gives the network calibrated out-of-distribution uncertainty, which is the central claim of NCP.

The closed-form Gaussian KL used here is

\[ \mathrm{KL}\bigl[\mathcal{N}(\mu, \sigma^2) \,\|\, \mathcal{N}(\mu_p, \sigma_p^2)\bigr] = \log\frac{\sigma_p}{\sigma} + \frac{\sigma^2 + (\mu - \mu_p)^2}{2\sigma_p^2} - \tfrac{1}{2}. \]
Plate semantics

Unlike pyrox's weight-prior KL terms, the NCP KL is data-dependent — every input row contributes its own \(\mathrm{KL}_n\) term. Internally the layer emits the numpyro.factor site as a per-example vector (shape (*batch,)) rather than a pre-summed scalar; that lets NumPyro's plate machinery sum over the batch axis and apply the subsample scaling automatically.

The canonical training pattern is to emit the layer inside numpyro.plate("data", N, subsample_size=B):

def model(x_clean, y_clean, x_noisy):
    clean_mean, _clean_std = network(x_clean)
    noisy_mean, noisy_std = network(x_noisy)
    ncp_out = NCPNormalOutput(prior_std=1.0)
    with numpyro.plate("data", N, subsample_size=B):
        ncp_out(noisy_mean, noisy_std)              # scaled to N
        numpyro.sample("obs",
            dist.Normal(clean_mean, ...), obs=y_clean)

Inside the plate, NumPyro sums the per-example log-densities over the batch dim and multiplies by scale = N / B, producing the standard unbiased estimate of the full-dataset NCP KL Σ_{n=1}^N KL_n. Outside any plate the layer's contribution is just Σ_{n in batch} KL_n (i.e. the raw batch sum), which is the correct full-dataset value when you train on the whole dataset at once.

Attributes:

Name Type Description
prior_mean float

Prior predictive mean \(\mu_\mathrm{prior}\).

prior_std float

Prior predictive std \(\sigma_\mathrm{prior}\) (must be positive).

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Examples:

>>> import jax.numpy as jnp
>>> from numpyro import handlers
>>> ncp = NCPNormalOutput(
...     prior_mean=0.0, prior_std=1.0, pyrox_name="ncp_out"
... )
>>> noisy_mean = jnp.zeros((4, 1))
>>> noisy_std = 0.5 * jnp.ones((4, 1))
>>> with handlers.seed(rng_seed=0):
...     kl = ncp(noisy_mean, noisy_std)
>>> kl.shape
()
References

Hafner, D., Tran, D., Lillicrap, T., Irpan, A., & Davidson, J. (2018). Noise Contrastive Priors for Functional Uncertainty. UAI.

Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
class NCPNormalOutput(PyroxModule):
    r"""Output-side Noise Contrastive Prior layer (Hafner et al., 2018).

    Completes the NCP pattern in ``pyrox_nn``: pair with
    `NCPContinuousPerturb` at the input and a heteroscedastic
    network (e.g. an MLP terminating in a mean head and a positive-std
    head — a softplus or ``exp`` of a learned log-scale) so the
    network produces predictions for both the *clean* batch and the
    input-perturbed *noisy* batch. Given the noisy batch's predictive
    distribution $\mathcal{N}(\hat{y}_n, \hat{\sigma}_n^2)$,
    this layer adds the analytic NCP regulariser

    $$
    \mathcal{L}_\mathrm{NCP} =
    \sum_{n} \mathrm{KL}\!\bigl[\mathcal{N}(\hat{y}_n, \hat{\sigma}_n^2)
        \;\big\|\; \mathcal{N}(\mu_\mathrm{prior}, \sigma_\mathrm{prior}^2)\bigr]
    $$

    to the model log density via `numpyro.factor`. Pulling the
    noisy-input predictive distribution toward the fixed prior away
    from the training distribution gives the network calibrated
    out-of-distribution uncertainty, which is the central claim of NCP.

    The closed-form Gaussian KL used here is

    $$
    \mathrm{KL}\bigl[\mathcal{N}(\mu, \sigma^2)
        \,\|\, \mathcal{N}(\mu_p, \sigma_p^2)\bigr]
    = \log\frac{\sigma_p}{\sigma} +
      \frac{\sigma^2 + (\mu - \mu_p)^2}{2\sigma_p^2} - \tfrac{1}{2}.
    $$

    Plate semantics:
        Unlike pyrox's *weight-prior* KL terms, the NCP KL is
        **data-dependent** — every input row contributes its own
        $\mathrm{KL}_n$ term. Internally the layer emits the
        `numpyro.factor` site as a *per-example* vector
        (shape ``(*batch,)``) rather than a pre-summed scalar; that
        lets NumPyro's plate machinery sum over the batch axis and
        apply the subsample scaling automatically.

        The canonical training pattern is to emit the layer **inside**
        ``numpyro.plate("data", N, subsample_size=B)``:

            def model(x_clean, y_clean, x_noisy):
                clean_mean, _clean_std = network(x_clean)
                noisy_mean, noisy_std = network(x_noisy)
                ncp_out = NCPNormalOutput(prior_std=1.0)
                with numpyro.plate("data", N, subsample_size=B):
                    ncp_out(noisy_mean, noisy_std)              # scaled to N
                    numpyro.sample("obs",
                        dist.Normal(clean_mean, ...), obs=y_clean)

        Inside the plate, NumPyro sums the per-example log-densities
        over the batch dim and multiplies by ``scale = N / B``,
        producing the standard unbiased estimate of the full-dataset
        NCP KL ``Σ_{n=1}^N KL_n``. Outside any plate the layer's
        contribution is just ``Σ_{n in batch} KL_n`` (i.e. the raw
        batch sum), which is the correct full-dataset value when
        you train on the whole dataset at once.

    Attributes:
        prior_mean: Prior predictive mean $\mu_\mathrm{prior}$.
        prior_std: Prior predictive std $\sigma_\mathrm{prior}$
            (must be positive).
        pyrox_name: Explicit scope name for NumPyro site registration.

    Examples:
        >>> import jax.numpy as jnp
        >>> from numpyro import handlers
        >>> ncp = NCPNormalOutput(
        ...     prior_mean=0.0, prior_std=1.0, pyrox_name="ncp_out"
        ... )
        >>> noisy_mean = jnp.zeros((4, 1))
        >>> noisy_std = 0.5 * jnp.ones((4, 1))
        >>> with handlers.seed(rng_seed=0):
        ...     kl = ncp(noisy_mean, noisy_std)
        >>> kl.shape
        ()

    References:
        Hafner, D., Tran, D., Lillicrap, T., Irpan, A., & Davidson, J.
        (2018). *Noise Contrastive Priors for Functional Uncertainty.*
        UAI.
    """

    prior_mean: float = eqx.field(static=True, default=0.0)
    prior_std: float = eqx.field(static=True, default=1.0)
    pyrox_name: str | None = eqx.field(static=True, default=None)

    def __post_init__(self) -> None:
        if self.prior_std <= 0:
            raise ValueError(f"prior_std must be > 0; got {self.prior_std}.")

    @pyrox_method
    def __call__(
        self,
        noisy_mean: Float[Array, "*batch D"],
        noisy_std: Float[Array, "*batch D"],
    ) -> Float[Array, ""]:
        # `noisy_std` is a *standard deviation* — the caller is responsible
        # for ensuring it is non-negative (e.g. via softplus/exp on a
        # learned log-scale head). Negative inputs would silently give a
        # finite-but-wrong KL after squaring; an explicit zero produces
        # `+inf` from `log(0)` which surfaces the bug. We do not floor
        # `noisy_var` because doing so asymmetrically (without a matching
        # floor on `prior_var`) would break the `noisy_std == prior_std`
        # → KL = 0 invariant for tiny `prior_std`.
        if noisy_mean.shape != noisy_std.shape:
            raise ValueError(
                f"noisy_mean shape {noisy_mean.shape} != "
                f"noisy_std shape {noisy_std.shape}."
            )
        # Require an explicit feature axis. With a 1-D ``(B,)`` input,
        # summing along axis=-1 below would collapse the *batch* axis
        # itself, producing a scalar factor that gets broadcast across
        # the data plate — exactly the over-counting bug a per-example
        # factor is designed to avoid. For scalar-regression heads,
        # reshape to ``(B, 1)``.
        if noisy_mean.ndim < 2:
            raise ValueError(
                "noisy_mean / noisy_std must have at least 2 dims "
                "(batch + feature). For a scalar regression head, pass "
                "`noisy_mean[:, None]` and `noisy_std[:, None]`. Got "
                f"shape {noisy_mean.shape}."
            )
        prior_var = jnp.asarray(self.prior_std) ** 2
        noisy_var = noisy_std**2
        kl_per_elem = _diag_gaussian_kl(
            noisy_mean,
            noisy_var,
            jnp.log(noisy_var),
            prior_mean=self.prior_mean,
            log_prior_scale=jnp.log(self.prior_std),
            prior_var=prior_var,
        )
        # Sum only over the trailing feature axis. Keeping the leading
        # batch axis intact is what makes NumPyro's plate machinery do
        # the right thing under `plate("data", N, subsample_size=B)`:
        # the plate handler sums log_probs over the batch dim and then
        # multiplies by `N/B`, giving the unbiased full-dataset estimate
        # `(N/B) * sum_{n in batch} kl_n`. Emitting an already-summed
        # scalar instead would let NumPyro broadcast it across the
        # plate dim and over-count by a factor of B.
        kl_per_example = jnp.sum(kl_per_elem, axis=-1)
        # Add -kl_per_example to the model log density site-by-site.
        # Outside any plate this sums to -total_KL; inside a plate the
        # plate handler scales it correctly.
        numpyro.factor(self._pyrox_fullname("kl"), -kl_per_example)
        return jnp.sum(kl_per_example)

RBFFourierFeatures

Bases: PyroxModule

SSGP-style RFF layer with RBF spectral density.

Both the spectral frequencies \(W\) and the lengthscale \(\ell\) are pyrox_sample sites — \(W\) has a standard normal prior (the RBF spectral density) and \(\ell\) has a LogNormal prior. Under SVI, the guide learns a posterior over both; under a seed handler, they are drawn from the prior.

Attributes:

Name Type Description
in_features int

Input dimension.

n_features int

Number of frequency pairs (output dim 2 * n_features).

init_lengthscale float

Prior location for the lengthscale.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
class RBFFourierFeatures(PyroxModule):
    r"""SSGP-style RFF layer with RBF spectral density.

    Both the spectral frequencies $W$ and the lengthscale
    $\ell$ are ``pyrox_sample`` sites — $W$ has a
    standard normal prior (the RBF spectral density) and $\ell$
    has a ``LogNormal`` prior. Under SVI, the guide learns a posterior
    over both; under a seed handler, they are drawn from the prior.

    Attributes:
        in_features: Input dimension.
        n_features: Number of frequency pairs (output dim
            ``2 * n_features``).
        init_lengthscale: Prior location for the lengthscale.
        pyrox_name: Explicit scope name for NumPyro site registration.
    """

    in_features: int = eqx.field(static=True)
    n_features: int = eqx.field(static=True)
    init_lengthscale: float = 1.0
    pyrox_name: str | None = None

    @classmethod
    def init(
        cls,
        in_features: int,
        n_features: int,
        *,
        lengthscale: float = 1.0,
    ) -> RBFFourierFeatures:
        if lengthscale <= 0:
            raise ValueError(f"lengthscale must be > 0, got {lengthscale}.")
        return cls(
            in_features=in_features,
            n_features=n_features,
            init_lengthscale=lengthscale,
        )

    @pyrox_method
    def __call__(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D_rff"]:
        W = self.pyrox_sample(
            "W",
            dist.Normal(0.0, 1.0)
            .expand([self.in_features, self.n_features])
            .to_event(2),
        )
        ls = self.pyrox_sample(
            "lengthscale",
            dist.LogNormal(jnp.log(jnp.asarray(self.init_lengthscale)), 1.0),
        )
        return _vmap_rff_forward(W, ls, self.n_features, x)

RBFCosineFeatures

Bases: PyroxModule

Cosine-bias variant of random Fourier features for the RBF kernel.

Uses the single-cosine feature map with a bias term:

\[ \phi(x) = \sqrt{2 / D}\,\cos(x W / \ell + b) \]

where \(W \sim \mathcal{N}(0, I)\) and \(b \sim \mathrm{Uniform}(0, 2\pi)\). This variant produces n_features-dimensional output (half the dimension of the [cos, sin] variant in RBFFourierFeatures) and is commonly used in Random Kitchen Sinks implementations.

All parameters (\(W\), \(b\), \(\ell\)) are pyrox_sample sites.

Attributes:

Name Type Description
in_features int

Input dimension.

n_features int

Number of random features (= output dimension).

init_lengthscale float

Prior location for the lengthscale.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
class RBFCosineFeatures(PyroxModule):
    r"""Cosine-bias variant of random Fourier features for the RBF kernel.

    Uses the single-cosine feature map with a bias term:

    $$
    \phi(x) = \sqrt{2 / D}\,\cos(x W / \ell + b)
    $$

    where $W \sim \mathcal{N}(0, I)$ and
    $b \sim \mathrm{Uniform}(0, 2\pi)$. This variant produces
    ``n_features``-dimensional output (half the dimension of the
    ``[cos, sin]`` variant in `RBFFourierFeatures`) and is
    commonly used in Random Kitchen Sinks implementations.

    All parameters ($W$, $b$, $\ell$) are
    ``pyrox_sample`` sites.

    Attributes:
        in_features: Input dimension.
        n_features: Number of random features (= output dimension).
        init_lengthscale: Prior location for the lengthscale.
        pyrox_name: Explicit scope name for NumPyro site registration.
    """

    in_features: int = eqx.field(static=True)
    n_features: int = eqx.field(static=True)
    init_lengthscale: float = 1.0
    pyrox_name: str | None = None

    @classmethod
    def init(
        cls,
        in_features: int,
        n_features: int,
        *,
        lengthscale: float = 1.0,
    ) -> RBFCosineFeatures:
        if lengthscale <= 0:
            raise ValueError(f"lengthscale must be > 0, got {lengthscale}.")
        return cls(
            in_features=in_features,
            n_features=n_features,
            init_lengthscale=lengthscale,
        )

    @pyrox_method
    def __call__(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D_rff"]:
        W = self.pyrox_sample(
            "W",
            dist.Normal(0.0, 1.0)
            .expand([self.in_features, self.n_features])
            .to_event(2),
        )
        b = self.pyrox_sample(
            "b",
            dist.Uniform(0.0, 2.0 * jnp.pi).expand([self.n_features]).to_event(1),
        )
        ls = self.pyrox_sample(
            "lengthscale",
            dist.LogNormal(jnp.log(jnp.asarray(self.init_lengthscale)), 1.0),
        )
        return _vmap_rff_cosine_forward(W, b, ls, self.n_features, x)

MaternFourierFeatures

Bases: PyroxModule

SSGP-style RFF layer with Matern spectral density.

Spectral frequencies \(W\) have a StudentT(df=2\nu) prior (the Matern spectral density). The smoothness \(\nu\) controls the regularity: nu=0.5 (Laplace), nu=1.5 (Matern-3/2), nu=2.5 (Matern-5/2).

Attributes:

Name Type Description
in_features int

Input dimension.

n_features int

Number of frequency pairs.

nu float

Smoothness parameter \(\nu\).

init_lengthscale float

Prior location for the lengthscale.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
class MaternFourierFeatures(PyroxModule):
    r"""SSGP-style RFF layer with Matern spectral density.

    Spectral frequencies $W$ have a ``StudentT(df=2\nu)`` prior
    (the Matern spectral density). The smoothness $\nu$ controls
    the regularity: ``nu=0.5`` (Laplace), ``nu=1.5`` (Matern-3/2),
    ``nu=2.5`` (Matern-5/2).

    Attributes:
        in_features: Input dimension.
        n_features: Number of frequency pairs.
        nu: Smoothness parameter $\nu$.
        init_lengthscale: Prior location for the lengthscale.
        pyrox_name: Explicit scope name for NumPyro site registration.
    """

    in_features: int = eqx.field(static=True)
    n_features: int = eqx.field(static=True)
    nu: float = eqx.field(static=True, default=1.5)
    init_lengthscale: float = 1.0
    pyrox_name: str | None = None

    @classmethod
    def init(
        cls,
        in_features: int,
        n_features: int,
        *,
        nu: float = 1.5,
        lengthscale: float = 1.0,
    ) -> MaternFourierFeatures:
        if lengthscale <= 0:
            raise ValueError(f"lengthscale must be > 0, got {lengthscale}.")
        if nu <= 0:
            raise ValueError(f"nu must be > 0, got {nu}.")
        return cls(
            in_features=in_features,
            n_features=n_features,
            nu=nu,
            init_lengthscale=lengthscale,
        )

    @pyrox_method
    def __call__(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D_rff"]:
        W = self.pyrox_sample(
            "W",
            dist.StudentT(df=2.0 * self.nu, loc=0.0, scale=1.0)
            .expand([self.in_features, self.n_features])
            .to_event(2),
        )
        ls = self.pyrox_sample(
            "lengthscale",
            dist.LogNormal(jnp.log(jnp.asarray(self.init_lengthscale)), 1.0),
        )
        return _vmap_rff_forward(W, ls, self.n_features, x)

MaternCosineFeatures

Bases: PyroxModule

Cosine-bias variant of random Fourier features for the Matern kernel.

Single-cosine analogue of MaternFourierFeatures:

\[ \phi(x) = \sqrt{2 / D}\,\cos(x W / \ell + b) \]

where \(W \sim \mathrm{StudentT}(2\nu)\) (the Matern spectral density) and \(b \sim \mathrm{Uniform}(0, 2\pi)\). Output dim is n_features (vs 2 * n_features for the [cos, sin] variant). Approximates the same kernel as MaternFourierFeatures in expectation but with higher variance per draw — see Sutherland & Schneider (2015).

All parameters (\(W\), \(b\), \(\ell\)) are pyrox_sample sites.

Attributes:

Name Type Description
in_features int

Input dimension.

n_features int

Number of random features (= output dimension).

nu float

Smoothness parameter \(\nu\).

init_lengthscale float

Prior location for the lengthscale.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
class MaternCosineFeatures(PyroxModule):
    r"""Cosine-bias variant of random Fourier features for the Matern kernel.

    Single-cosine analogue of `MaternFourierFeatures`:

    $$
    \phi(x) = \sqrt{2 / D}\,\cos(x W / \ell + b)
    $$

    where $W \sim \mathrm{StudentT}(2\nu)$ (the Matern spectral
    density) and $b \sim \mathrm{Uniform}(0, 2\pi)$. Output dim is
    ``n_features`` (vs ``2 * n_features`` for the ``[cos, sin]``
    variant). Approximates the same kernel as
    `MaternFourierFeatures` in expectation but with higher
    variance per draw — see Sutherland & Schneider (2015).

    All parameters ($W$, $b$, $\ell$) are
    ``pyrox_sample`` sites.

    Attributes:
        in_features: Input dimension.
        n_features: Number of random features (= output dimension).
        nu: Smoothness parameter $\nu$.
        init_lengthscale: Prior location for the lengthscale.
        pyrox_name: Explicit scope name for NumPyro site registration.
    """

    in_features: int = eqx.field(static=True)
    n_features: int = eqx.field(static=True)
    nu: float = eqx.field(static=True, default=1.5)
    init_lengthscale: float = 1.0
    pyrox_name: str | None = None

    @classmethod
    def init(
        cls,
        in_features: int,
        n_features: int,
        *,
        nu: float = 1.5,
        lengthscale: float = 1.0,
    ) -> MaternCosineFeatures:
        if lengthscale <= 0:
            raise ValueError(f"lengthscale must be > 0, got {lengthscale}.")
        if nu <= 0:
            raise ValueError(f"nu must be > 0, got {nu}.")
        return cls(
            in_features=in_features,
            n_features=n_features,
            nu=nu,
            init_lengthscale=lengthscale,
        )

    @pyrox_method
    def __call__(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D_rff"]:
        W = self.pyrox_sample(
            "W",
            dist.StudentT(df=2.0 * self.nu, loc=0.0, scale=1.0)
            .expand([self.in_features, self.n_features])
            .to_event(2),
        )
        b = self.pyrox_sample(
            "b",
            dist.Uniform(0.0, 2.0 * jnp.pi).expand([self.n_features]).to_event(1),
        )
        ls = self.pyrox_sample(
            "lengthscale",
            dist.LogNormal(jnp.log(jnp.asarray(self.init_lengthscale)), 1.0),
        )
        return _vmap_rff_cosine_forward(W, b, ls, self.n_features, x)

LaplaceFourierFeatures

Bases: PyroxModule

SSGP-style RFF layer with Laplace (Matern-1/2) spectral density.

Spectral frequencies \(W\) have a Cauchy prior (Student-t with df = 1).

Attributes:

Name Type Description
in_features int

Input dimension.

n_features int

Number of frequency pairs.

init_lengthscale float

Prior location for the lengthscale.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
class LaplaceFourierFeatures(PyroxModule):
    r"""SSGP-style RFF layer with Laplace (Matern-1/2) spectral density.

    Spectral frequencies $W$ have a ``Cauchy`` prior (Student-t
    with ``df = 1``).

    Attributes:
        in_features: Input dimension.
        n_features: Number of frequency pairs.
        init_lengthscale: Prior location for the lengthscale.
        pyrox_name: Explicit scope name for NumPyro site registration.
    """

    in_features: int = eqx.field(static=True)
    n_features: int = eqx.field(static=True)
    init_lengthscale: float = 1.0
    pyrox_name: str | None = None

    @classmethod
    def init(
        cls,
        in_features: int,
        n_features: int,
        *,
        lengthscale: float = 1.0,
    ) -> LaplaceFourierFeatures:
        if lengthscale <= 0:
            raise ValueError(f"lengthscale must be > 0, got {lengthscale}.")
        return cls(
            in_features=in_features,
            n_features=n_features,
            init_lengthscale=lengthscale,
        )

    @pyrox_method
    def __call__(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D_rff"]:
        W = self.pyrox_sample(
            "W",
            dist.StudentT(df=1.0, loc=0.0, scale=1.0)
            .expand([self.in_features, self.n_features])
            .to_event(2),
        )
        ls = self.pyrox_sample(
            "lengthscale",
            dist.LogNormal(jnp.log(jnp.asarray(self.init_lengthscale)), 1.0),
        )
        return _vmap_rff_forward(W, ls, self.n_features, x)

LaplaceCosineFeatures

Bases: PyroxModule

Cosine-bias variant of random Fourier features for the Laplace kernel.

Single-cosine analogue of LaplaceFourierFeatures (the Matern-1/2 kernel):

\[ \phi(x) = \sqrt{2 / D}\,\cos(x W / \ell + b) \]

where \(W \sim \mathrm{Cauchy}(0, 1)\) (Student-t with df = 1) and \(b \sim \mathrm{Uniform}(0, 2\pi)\). Output dim is n_features.

All parameters (\(W\), \(b\), \(\ell\)) are pyrox_sample sites.

Attributes:

Name Type Description
in_features int

Input dimension.

n_features int

Number of random features (= output dimension).

init_lengthscale float

Prior location for the lengthscale.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
class LaplaceCosineFeatures(PyroxModule):
    r"""Cosine-bias variant of random Fourier features for the Laplace kernel.

    Single-cosine analogue of `LaplaceFourierFeatures` (the
    Matern-1/2 kernel):

    $$
    \phi(x) = \sqrt{2 / D}\,\cos(x W / \ell + b)
    $$

    where $W \sim \mathrm{Cauchy}(0, 1)$ (Student-t with
    ``df = 1``) and $b \sim \mathrm{Uniform}(0, 2\pi)$. Output
    dim is ``n_features``.

    All parameters ($W$, $b$, $\ell$) are
    ``pyrox_sample`` sites.

    Attributes:
        in_features: Input dimension.
        n_features: Number of random features (= output dimension).
        init_lengthscale: Prior location for the lengthscale.
        pyrox_name: Explicit scope name for NumPyro site registration.
    """

    in_features: int = eqx.field(static=True)
    n_features: int = eqx.field(static=True)
    init_lengthscale: float = 1.0
    pyrox_name: str | None = None

    @classmethod
    def init(
        cls,
        in_features: int,
        n_features: int,
        *,
        lengthscale: float = 1.0,
    ) -> LaplaceCosineFeatures:
        if lengthscale <= 0:
            raise ValueError(f"lengthscale must be > 0, got {lengthscale}.")
        return cls(
            in_features=in_features,
            n_features=n_features,
            init_lengthscale=lengthscale,
        )

    @pyrox_method
    def __call__(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D_rff"]:
        W = self.pyrox_sample(
            "W",
            dist.StudentT(df=1.0, loc=0.0, scale=1.0)
            .expand([self.in_features, self.n_features])
            .to_event(2),
        )
        b = self.pyrox_sample(
            "b",
            dist.Uniform(0.0, 2.0 * jnp.pi).expand([self.n_features]).to_event(1),
        )
        ls = self.pyrox_sample(
            "lengthscale",
            dist.LogNormal(jnp.log(jnp.asarray(self.init_lengthscale)), 1.0),
        )
        return _vmap_rff_cosine_forward(W, b, ls, self.n_features, x)

ArcCosineFourierFeatures

Bases: PyroxModule

Random features for the arc-cosine kernel (Cho & Saul, 2009).

The arc-cosine kernel of order \(p\) corresponds to an infinite-width single-layer ReLU network. The random feature map is:

\[ \phi(x) = \sqrt{2 / D}\,\max(0,\, x W / \ell)^p \]

where \(W \sim \mathcal{N}(0, I)\).

order=0 gives the Heaviside (step) feature; order=1 gives the ReLU feature (the most common); order=2 gives the squared ReLU feature.

Attributes:

Name Type Description
in_features int

Input dimension.

n_features int

Number of random features (= output dimension).

order int

Kernel order (0, 1, or 2).

init_lengthscale float

Prior location for the lengthscale.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
class ArcCosineFourierFeatures(PyroxModule):
    r"""Random features for the arc-cosine kernel (Cho & Saul, 2009).

    The arc-cosine kernel of order $p$ corresponds to an
    infinite-width single-layer ReLU network. The random feature map
    is:

    $$
    \phi(x) = \sqrt{2 / D}\,\max(0,\, x W / \ell)^p
    $$

    where $W \sim \mathcal{N}(0, I)$.

    ``order=0`` gives the Heaviside (step) feature; ``order=1`` gives
    the ReLU feature (the most common); ``order=2`` gives the squared
    ReLU feature.

    Attributes:
        in_features: Input dimension.
        n_features: Number of random features (= output dimension).
        order: Kernel order (0, 1, or 2).
        init_lengthscale: Prior location for the lengthscale.
        pyrox_name: Explicit scope name for NumPyro site registration.
    """

    in_features: int = eqx.field(static=True)
    n_features: int = eqx.field(static=True)
    order: int = eqx.field(static=True, default=1)
    init_lengthscale: float = 1.0
    pyrox_name: str | None = None

    @classmethod
    def init(
        cls,
        in_features: int,
        n_features: int,
        *,
        order: int = 1,
        lengthscale: float = 1.0,
    ) -> ArcCosineFourierFeatures:
        return cls(
            in_features=in_features,
            n_features=n_features,
            order=order,
            init_lengthscale=lengthscale,
        )

    @pyrox_method
    def __call__(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D_rff"]:
        W = self.pyrox_sample(
            "W",
            dist.Normal(0.0, 1.0)
            .expand([self.in_features, self.n_features])
            .to_event(2),
        )
        ls = self.pyrox_sample(
            "lengthscale",
            dist.LogNormal(jnp.log(jnp.asarray(self.init_lengthscale)), 1.0),
        )
        z = einx.dot("... din, din f -> ... f", x, W) / ls
        if self.order == 0:
            h = (z > 0.0).astype(x.dtype)
        else:
            h = jnp.maximum(z, 0.0) ** self.order
        return jnp.sqrt(2.0 / self.n_features) * h

RandomKitchenSinks

Bases: PyroxModule

Random Kitchen Sinks: RFF + a learned linear head.

Composes any RFF layer (RBFFourierFeatures, MaternFourierFeatures, LaplaceFourierFeatures) with a trainable linear projection:

\[ y = \phi(x)\, \beta + b \]

The linear head (beta, bias) is registered via pyrox_sample with Normal priors.

Attributes:

Name Type Description
rff RBFFourierFeatures | MaternFourierFeatures | LaplaceFourierFeatures

The underlying RFF feature layer.

init_beta Float[Array, 'D_rff D_out']

Initial linear weights.

init_bias Float[Array, ' D_out']

Initial bias vector.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
class RandomKitchenSinks(PyroxModule):
    r"""Random Kitchen Sinks: RFF + a learned linear head.

    Composes any RFF layer (`RBFFourierFeatures`,
    `MaternFourierFeatures`, `LaplaceFourierFeatures`)
    with a trainable linear projection:

    $$
    y = \phi(x)\, \beta + b
    $$

    The linear head (``beta``, ``bias``) is registered via
    ``pyrox_sample`` with ``Normal`` priors.

    Attributes:
        rff: The underlying RFF feature layer.
        init_beta: Initial linear weights.
        init_bias: Initial bias vector.
        pyrox_name: Explicit scope name for NumPyro site registration.
    """

    rff: RBFFourierFeatures | MaternFourierFeatures | LaplaceFourierFeatures
    init_beta: Float[Array, "D_rff D_out"]
    init_bias: Float[Array, " D_out"]
    pyrox_name: str | None = None

    @classmethod
    def init(
        cls,
        rff: RBFFourierFeatures | MaternFourierFeatures | LaplaceFourierFeatures,
        out_features: int,
    ) -> RandomKitchenSinks:
        """Construct from a pre-built RFF layer with zero-initialized head."""
        beta = jnp.zeros((2 * rff.n_features, out_features))
        bias = jnp.zeros(out_features)
        return cls(rff=rff, init_beta=beta, init_bias=bias)

    @pyrox_method
    def __call__(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D_out"]:
        phi = self.rff(x)
        beta = self.pyrox_sample(
            "beta",
            dist.Normal(self.init_beta, 1.0).to_event(2),
        )
        bias = self.pyrox_sample(
            "bias",
            dist.Normal(self.init_bias, 1.0).to_event(1),
        )
        return einx.dot("... r, r dout -> ... dout", phi, beta) + bias

init(rff: RBFFourierFeatures | MaternFourierFeatures | LaplaceFourierFeatures, out_features: int) -> RandomKitchenSinks classmethod

Construct from a pre-built RFF layer with zero-initialized head.

Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
@classmethod
def init(
    cls,
    rff: RBFFourierFeatures | MaternFourierFeatures | LaplaceFourierFeatures,
    out_features: int,
) -> RandomKitchenSinks:
    """Construct from a pre-built RFF layer with zero-initialized head."""
    beta = jnp.zeros((2 * rff.n_features, out_features))
    bias = jnp.zeros(out_features)
    return cls(rff=rff, init_beta=beta, init_bias=bias)

Wave-4 spectral layers (#41)

VariationalFourierFeatures

Bases: PyroxModule

VSSGP — RFF with a learnable variational posterior over frequencies.

Standard RFF (e.g. RBFFourierFeatures) treats the spectral frequencies \(W\) as a frozen prior draw; VSSGP (Gal & Turner, 2015) treats \(W\) as a latent with a learnable mean-field posterior, recovering spectral uncertainty on top of the feature-space uncertainty.

Prior: \(p(W) = \mathcal{N}(0, I)\) (RBF spectral density in lengthscale-1 units). The lengthscale is itself a sampled site (LogNormal(log init_lengthscale, 1)) so that frequencies are rescaled to the physical kernel.

Under SVI, attach an AutoNormal to learn the posterior on W; under prior-only seeds, behaves identically to RBFFourierFeatures.

Attributes:

Name Type Description
in_features int

Input dimension \(D\).

n_features int

Number of frequency pairs (output dim 2 * n_features).

init_lengthscale float

Prior location for the kernel lengthscale.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
class VariationalFourierFeatures(PyroxModule):
    r"""VSSGP — RFF with a learnable variational posterior over frequencies.

    Standard RFF (e.g. `RBFFourierFeatures`) treats the spectral
    frequencies $W$ as a frozen prior draw; VSSGP (Gal & Turner,
    2015) treats $W$ as a latent with a learnable mean-field
    posterior, recovering spectral *uncertainty* on top of the
    feature-space uncertainty.

    Prior: $p(W) = \mathcal{N}(0, I)$ (RBF spectral density in
    lengthscale-1 units). The lengthscale is itself a sampled site
    (``LogNormal(log init_lengthscale, 1)``) so that frequencies are
    rescaled to the physical kernel.

    Under SVI, attach an `AutoNormal` to
    learn the posterior on ``W``; under prior-only seeds, behaves
    identically to `RBFFourierFeatures`.

    Attributes:
        in_features: Input dimension $D$.
        n_features: Number of frequency pairs (output dim ``2 * n_features``).
        init_lengthscale: Prior location for the kernel lengthscale.
        pyrox_name: Explicit scope name for NumPyro site registration.
    """

    in_features: int = eqx.field(static=True)
    n_features: int = eqx.field(static=True)
    init_lengthscale: float = 1.0
    pyrox_name: str | None = None

    @classmethod
    def init(
        cls,
        in_features: int,
        n_features: int,
        *,
        lengthscale: float = 1.0,
    ) -> VariationalFourierFeatures:
        if lengthscale <= 0:
            raise ValueError(f"lengthscale must be > 0, got {lengthscale}.")
        return cls(
            in_features=in_features,
            n_features=n_features,
            init_lengthscale=lengthscale,
        )

    @pyrox_method
    def __call__(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D_rff"]:
        # Same prior as RBFFourierFeatures — the *posterior* is what differs
        # under SVI: an attached AutoGuide learns q(W) instead of forcing W
        # to its prior draw.
        W = self.pyrox_sample(
            "W",
            dist.Normal(0.0, 1.0)
            .expand([self.in_features, self.n_features])
            .to_event(2),
        )
        ls = self.pyrox_sample(
            "lengthscale",
            dist.LogNormal(jnp.log(jnp.asarray(self.init_lengthscale)), 1.0),
        )
        return _vmap_rff_forward(W, ls, self.n_features, x)

OrthogonalRandomFeatures

Bases: Module

Orthogonal Random Features (Yu et al., 2016) — variance-reduced RFF.

Frequencies are drawn from blocks of Haar-orthogonal matrices scaled by independent chi-distributed magnitudes, giving the same RBF kernel approximation as plain RBFFourierFeatures in expectation but with provably lower variance for finite n_features.

Frozen at construction time — no priors, no SVI on W. The frequency matrix is built once from a key and stored as a static array.

Attributes:

Name Type Description
in_features int

Input dimension \(D\).

n_features int

Number of feature pairs. Must satisfy n_features % in_features == 0 so that ORF blocks tile cleanly.

lengthscale Float[Array, '']

Fixed kernel lengthscale (no prior; pass a value).

W Float[Array, 'D_in D_orf']

Pre-built frequency matrix of shape (in_features, n_features).

The feature map is the shared RFF map \(\phi(x) = \sqrt{1/D}\,[\cos(W^\top x/\ell),\,\sin(W^\top x/\ell)]\), so calling the module on a (in_features,) vector yields a (2 * n_features,) feature vector.

Examples:

>>> import jax.numpy as jnp, jax.random as jr
>>> from geonnax.randfeat import OrthogonalRandomFeatures
>>> orf = OrthogonalRandomFeatures.init(4, 8, key=jr.PRNGKey(0))
>>> orf(jnp.ones(4)).shape
(16,)
Source code in .venv/lib/python3.12/site-packages/geonnax/randfeat.py
class OrthogonalRandomFeatures(eqx.Module):
    r"""Orthogonal Random Features (Yu et al., 2016) — variance-reduced RFF.

    Frequencies are drawn from blocks of Haar-orthogonal matrices scaled by
    independent chi-distributed magnitudes, giving the same RBF kernel
    approximation as plain ``RBFFourierFeatures`` *in expectation* but
    with provably lower variance for finite ``n_features``.

    Frozen at construction time — no priors, no SVI on ``W``. The frequency
    matrix is built once from a ``key`` and stored as a static array.

    Attributes:
        in_features: Input dimension $D$.
        n_features: Number of feature pairs. Must satisfy
            ``n_features % in_features == 0`` so that ORF blocks tile cleanly.
        lengthscale: Fixed kernel lengthscale (no prior; pass a value).
        W: Pre-built frequency matrix of shape ``(in_features, n_features)``.

    The feature map is the shared RFF map
    $\phi(x) = \sqrt{1/D}\,[\cos(W^\top x/\ell),\,\sin(W^\top x/\ell)]$,
    so calling the module on a ``(in_features,)`` vector yields a
    ``(2 * n_features,)`` feature vector.

    Examples:
        >>> import jax.numpy as jnp, jax.random as jr
        >>> from geonnax.randfeat import OrthogonalRandomFeatures
        >>> orf = OrthogonalRandomFeatures.init(4, 8, key=jr.PRNGKey(0))
        >>> orf(jnp.ones(4)).shape
        (16,)
    """

    in_features: int = eqx.field(static=True)
    n_features: int = eqx.field(static=True)
    lengthscale: Float[Array, ""]
    W: Float[Array, "D_in D_orf"]

    @classmethod
    def init(
        cls,
        in_features: int,
        n_features: int,
        *,
        key: jax.Array,
        lengthscale: float = 1.0,
    ) -> OrthogonalRandomFeatures:
        """Build the frozen ORF frequency matrix and wrap it in the module.

        ``n_features`` must be divisible by ``in_features`` so the
        Haar-orthogonal blocks tile cleanly.

        Examples:
            >>> import jax.random as jr
            >>> from geonnax.randfeat import OrthogonalRandomFeatures
            >>> orf = OrthogonalRandomFeatures.init(4, 8, key=jr.PRNGKey(0))
            >>> orf.W.shape
            (4, 8)
        """
        if lengthscale <= 0:
            raise ValueError(f"lengthscale must be > 0, got {lengthscale}.")
        if in_features <= 0 or n_features <= 0:
            raise ValueError(
                "in_features and n_features must be > 0; got "
                f"in_features={in_features}, n_features={n_features}."
            )
        if n_features % in_features != 0:
            raise ValueError(
                f"n_features ({n_features}) must be divisible by in_features "
                f"({in_features}) so ORF blocks tile cleanly."
            )
        n_blocks = n_features // in_features
        W = orthogonal_blocks(in_features, n_blocks, key=key)
        return cls(
            in_features=in_features,
            n_features=n_features,
            lengthscale=jnp.asarray(lengthscale),
            W=W,
        )

    def __call__(self, x: Float[Array, " D_in"]) -> Float[Array, " D_rff"]:
        r"""Map ``x`` to its ORF feature vector ``(2 * n_features,)``.

        Examples:
            >>> import jax.numpy as jnp, jax.random as jr
            >>> from geonnax.randfeat import OrthogonalRandomFeatures
            >>> orf = OrthogonalRandomFeatures.init(2, 4, key=jr.PRNGKey(0))
            >>> orf(jnp.ones(2)).shape
            (8,)
        """
        # (D_in,) -> (2 * n_features,)
        return rff_forward(self.W, self.lengthscale, self.n_features, x)

init(in_features: int, n_features: int, *, key: jax.Array, lengthscale: float = 1.0) -> OrthogonalRandomFeatures classmethod

Build the frozen ORF frequency matrix and wrap it in the module.

n_features must be divisible by in_features so the Haar-orthogonal blocks tile cleanly.

Examples:

>>> import jax.random as jr
>>> from geonnax.randfeat import OrthogonalRandomFeatures
>>> orf = OrthogonalRandomFeatures.init(4, 8, key=jr.PRNGKey(0))
>>> orf.W.shape
(4, 8)
Source code in .venv/lib/python3.12/site-packages/geonnax/randfeat.py
@classmethod
def init(
    cls,
    in_features: int,
    n_features: int,
    *,
    key: jax.Array,
    lengthscale: float = 1.0,
) -> OrthogonalRandomFeatures:
    """Build the frozen ORF frequency matrix and wrap it in the module.

    ``n_features`` must be divisible by ``in_features`` so the
    Haar-orthogonal blocks tile cleanly.

    Examples:
        >>> import jax.random as jr
        >>> from geonnax.randfeat import OrthogonalRandomFeatures
        >>> orf = OrthogonalRandomFeatures.init(4, 8, key=jr.PRNGKey(0))
        >>> orf.W.shape
        (4, 8)
    """
    if lengthscale <= 0:
        raise ValueError(f"lengthscale must be > 0, got {lengthscale}.")
    if in_features <= 0 or n_features <= 0:
        raise ValueError(
            "in_features and n_features must be > 0; got "
            f"in_features={in_features}, n_features={n_features}."
        )
    if n_features % in_features != 0:
        raise ValueError(
            f"n_features ({n_features}) must be divisible by in_features "
            f"({in_features}) so ORF blocks tile cleanly."
        )
    n_blocks = n_features // in_features
    W = orthogonal_blocks(in_features, n_blocks, key=key)
    return cls(
        in_features=in_features,
        n_features=n_features,
        lengthscale=jnp.asarray(lengthscale),
        W=W,
    )

HSGPFeatures

Bases: PyroxModule

Hilbert-Space Gaussian Process feature layer (Riutort-Mayol et al., 2023).

A deterministic Laplacian-eigenfunction basis on the bounded box \([-L, L]^D\) plus learnable per-basis amplitudes with a kernel-spectral-density prior:

\[ \hat{f}(x) = \sum_{j=1}^{M} \alpha_j\,\sqrt{S(\sqrt{\lambda_j})}\,\phi_j(x), \quad \alpha_j \sim \mathcal{N}(0, 1). \]

This is the NN-side dual of pyrox_gp.FourierInducingFeatures — same basis, different prior wiring. As M and L grow, the induced GP converges to the kernel passed in.

Attributes:

Name Type Description
in_features int

Input dimension \(D\).

num_basis_per_dim tuple[int, ...]

Per-axis number of 1D eigenfunctions; total basis count is prod(num_basis_per_dim).

L tuple[float, ...]

Per-axis box half-width.

kernel Kernel

A stationary kernel from pyrox_gp whose spectral density supplies the per-basis prior variance. Currently pyrox_gp.RBF and pyrox_gp.Matern are supported by pyrox_gp._basis.spectral_density.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
class HSGPFeatures(PyroxModule):
    r"""Hilbert-Space Gaussian Process feature layer (Riutort-Mayol et al., 2023).

    A *deterministic* Laplacian-eigenfunction basis on the bounded box
    $[-L, L]^D$ plus learnable per-basis amplitudes with a
    kernel-spectral-density prior:

    $$
    \hat{f}(x) = \sum_{j=1}^{M} \alpha_j\,\sqrt{S(\sqrt{\lambda_j})}\,\phi_j(x),
    \quad \alpha_j \sim \mathcal{N}(0, 1).
    $$

    This is the NN-side dual of `pyrox_gp.FourierInducingFeatures`
    — same basis, different prior wiring. As ``M`` and ``L`` grow, the
    induced GP converges to the kernel passed in.

    Attributes:
        in_features: Input dimension $D$.
        num_basis_per_dim: Per-axis number of 1D eigenfunctions; total
            basis count is ``prod(num_basis_per_dim)``.
        L: Per-axis box half-width.
        kernel: A stationary kernel from `pyrox_gp` whose spectral
            density supplies the per-basis prior variance. Currently
            `pyrox_gp.RBF` and `pyrox_gp.Matern` are
            supported by `pyrox_gp._basis.spectral_density`.
        pyrox_name: Explicit scope name for NumPyro site registration.
    """

    in_features: int = eqx.field(static=True)
    num_basis_per_dim: tuple[int, ...] = eqx.field(static=True)
    L: tuple[float, ...] = eqx.field(static=True)
    kernel: Kernel
    pyrox_name: str | None = None

    @classmethod
    def init(
        cls,
        in_features: int,
        num_basis_per_dim: int | tuple[int, ...],
        L: float | tuple[float, ...],
        *,
        kernel: Kernel,
    ) -> HSGPFeatures:
        if isinstance(num_basis_per_dim, int):
            num_basis_per_dim = (num_basis_per_dim,) * in_features
        if isinstance(L, int | float):
            L = (float(L),) * in_features
        if len(num_basis_per_dim) != in_features:
            raise ValueError(
                f"num_basis_per_dim length ({len(num_basis_per_dim)}) "
                f"must match in_features ({in_features})."
            )
        if len(L) != in_features:
            raise ValueError(
                f"L length ({len(L)}) must match in_features ({in_features})."
            )
        if any(L_d <= 0 for L_d in L):
            raise ValueError(f"L must be all positive; got {L}.")
        if any(M_d < 1 for M_d in num_basis_per_dim):
            raise ValueError(
                f"num_basis_per_dim must be all >= 1; got {num_basis_per_dim}."
            )
        return cls(
            in_features=in_features,
            num_basis_per_dim=tuple(num_basis_per_dim),
            L=tuple(float(L_d) for L_d in L),
            kernel=kernel,
        )

    @property
    def num_basis(self) -> int:
        n = 1
        for m in self.num_basis_per_dim:
            n *= m
        return n

    @pyrox_method
    def __call__(self, x: Float[Array, "N D_in"]) -> Float[Array, " N"]:
        Phi, lam = fourier_basis(x, self.num_basis_per_dim, self.L)  # (N, M), (M,)
        # Spectral density evaluated under the kernel's own context so any
        # priors on (variance, lengthscale) register exactly once.
        with _kernel_context(self.kernel):
            S = spectral_density(self.kernel, lam, D=self.in_features)
        sqrt_S = jnp.sqrt(S)
        alpha = self.pyrox_sample(
            "alpha",
            dist.Normal(0.0, 1.0).expand([self.num_basis]).to_event(1),
        )
        return einx.dot("n m, m -> n", Phi, sqrt_S * alpha)

SIREN — Sinusoidal Representation Networks

SIREN (Sitzmann, Martel, Bergman, Lindell, Wetzstein — NeurIPS 2020) replaces ReLU/GELU with sin and prescribes a three-regime initialisation scheme that keeps pre-activation variance stable across depth.

Three-regime weight initialisation (Theorem 1)

Layer W init Activation
"first" U(-1/d_in, 1/d_in) sin(ω₀ · (W x + b))
"hidden" U(-√(c/d_in)/ω, √(c/d_in)/ω) sin(ω · (W x + b))
"last" U(-√(c/d_in), √(c/d_in)) none (linear) — W x + b

Bias b is initialised U(-1/√d_in, 1/√d_in) for every regime. Typical choice: ω₀ = ω = 30 for image / high-frequency INR tasks.

Usage

import jax.random as jr, jax.numpy as jnp
from pyrox_nn import SirenDense, SIREN, BayesianSIREN

# Single layer
layer = SirenDense.init(3, 64, key=jr.PRNGKey(0), layer_type="first")
y = layer(jnp.ones((5, 3)))  # (5, 64)

# Multi-layer network (depth=5 → first + 3 hidden + last)
net = SIREN.init(2, 64, 1, depth=5, key=jr.PRNGKey(0))
y = net(jnp.zeros((100, 2)))  # (100, 1)

# Bayesian variant (no key needed — weights come from the prior)
from numpyro import handlers
bnet = BayesianSIREN.init(2, 32, 1, depth=3)
with handlers.seed(rng_seed=0):
    y = bnet(jnp.zeros((10, 2)))  # (10, 1)

Alternative INR backbone

SIREN and GaborNet / FourierNet (MFN, #87) are complementary INR backbones: SIREN composes nonlinearities deeply, while MFN uses a product of Gabor filters. Choose based on the signal's smoothness profile.

SirenDense

Bases: Module

Sine-activated dense layer: y = sin(ω · (W x + b)) or y = W x + b.

Single primitive of a SIREN network with three init regimes (Sitzmann et al. 2020, Theorem 1):

+----------+-------------------------------------------+-------------+ | Regime | W init | Activation | +==========+===========================================+=============+ | first | U(-1/d_in, 1/d_in) | sin(ω··)| +----------+-------------------------------------------+-------------+ | hidden | U(-√(c/d_in)/ω, √(c/d_in)/ω) | sin(ω··)| +----------+-------------------------------------------+-------------+ | last | U(-√(c/d_in), √(c/d_in)) | none | +----------+-------------------------------------------+-------------+

Bias b is initialised U(-1/√d_in, 1/√d_in) for every regime.

Attributes:

Name Type Description
W Float[Array, 'in_features out_features']

Weight matrix of shape (in_features, out_features).

b Float[Array, ' out_features']

Bias vector of shape (out_features,).

omega float

Frequency multiplier applied inside the sine.

in_features int

Input dimension.

out_features int

Output dimension.

layer_type SirenLayerType

One of "first", "hidden", "last".

c float

Constant from Theorem 1 (default 6.0).

Source code in .venv/lib/python3.12/site-packages/geonnax/siren.py
class SirenDense(eqx.Module):
    r"""Sine-activated dense layer: ``y = sin(ω · (W x + b))`` or ``y = W x + b``.

    Single primitive of a SIREN network with three init regimes
    (Sitzmann et al. 2020, Theorem 1):

    +----------+-------------------------------------------+-------------+
    | Regime   | ``W`` init                                | Activation  |
    +==========+===========================================+=============+
    | first    | ``U(-1/d_in, 1/d_in)``                    | ``sin(ω··)``|
    +----------+-------------------------------------------+-------------+
    | hidden   | ``U(-√(c/d_in)/ω, √(c/d_in)/ω)``         | ``sin(ω··)``|
    +----------+-------------------------------------------+-------------+
    | last     | ``U(-√(c/d_in), √(c/d_in))``              | none        |
    +----------+-------------------------------------------+-------------+

    Bias ``b`` is initialised ``U(-1/√d_in, 1/√d_in)`` for every regime.

    Attributes:
        W: Weight matrix of shape ``(in_features, out_features)``.
        b: Bias vector of shape ``(out_features,)``.
        omega: Frequency multiplier applied inside the sine.
        in_features: Input dimension.
        out_features: Output dimension.
        layer_type: One of ``"first"``, ``"hidden"``, ``"last"``.
        c: Constant from Theorem 1 (default 6.0).
    """

    W: Float[Array, "in_features out_features"]
    b: Float[Array, " out_features"]
    omega: float = eqx.field(static=True)
    in_features: int = eqx.field(static=True)
    out_features: int = eqx.field(static=True)
    layer_type: SirenLayerType = eqx.field(static=True)
    c: float = eqx.field(static=True, default=6.0)

    @classmethod
    def init(
        cls,
        in_features: int,
        out_features: int,
        *,
        key: Array,
        omega: float = 30.0,
        layer_type: SirenLayerType = "hidden",
        c: float = 6.0,
    ) -> SirenDense:
        """Construct a ``SirenDense`` with Sitzmann-regime weight initialisation.

        Examples:
            >>> import jax.numpy as jnp, jax.random as jr
            >>> from geonnax.siren import SirenDense
            >>> layer = SirenDense.init(
            ...     3, 8, key=jr.PRNGKey(0), layer_type="first"
            ... )
            >>> layer(jnp.ones(3)).shape  # (3,) -> (8,)
            (8,)
        """
        _require_positive(
            in_features=in_features,
            out_features=out_features,
            omega=omega,
            c=c,
        )
        k_w, k_b = jax.random.split(key)
        # siren_W_limit validates layer_type.
        w_limit = siren_W_limit(layer_type, in_features, omega, c)
        W = jax.random.uniform(
            k_w, (in_features, out_features), minval=-w_limit, maxval=w_limit
        )
        b_limit = 1.0 / math.sqrt(in_features)
        b = jax.random.uniform(k_b, (out_features,), minval=-b_limit, maxval=b_limit)
        return cls(
            W=W,
            b=b,
            omega=omega,
            in_features=in_features,
            out_features=out_features,
            layer_type=layer_type,
            c=c,
        )

    def __call__(self, x: Float[Array, " D_in"]) -> Float[Array, " D_out"]:
        r"""Apply the layer: ``sin(ω · (W x + b))``, or ``W x + b`` if ``last``.

        Args:
            x: Input vector of shape ``(in_features,)``.

        Returns:
            Output vector of shape ``(out_features,)``.

        Examples:
            >>> import jax.numpy as jnp, jax.random as jr
            >>> from geonnax.siren import SirenDense
            >>> layer = SirenDense.init(4, 5, key=jr.PRNGKey(0))
            >>> layer(jnp.ones(4)).shape  # (4,) -> (5,)
            (5,)
        """
        # Affine map: (D_in,) · (D_in, D_out) -> (D_out,), then add bias.
        pre = einx.dot("i, i o -> o", x, self.W) + self.b
        if self.layer_type == "last":
            return pre  # readout layer has no activation
        # Sine activation with frequency multiplier ω: y = sin(ω · pre).
        return jnp.sin(self.omega * pre)

init(in_features: int, out_features: int, *, key: Array, omega: float = 30.0, layer_type: SirenLayerType = 'hidden', c: float = 6.0) -> SirenDense classmethod

Construct a SirenDense with Sitzmann-regime weight initialisation.

Examples:

>>> import jax.numpy as jnp, jax.random as jr
>>> from geonnax.siren import SirenDense
>>> layer = SirenDense.init(
...     3, 8, key=jr.PRNGKey(0), layer_type="first"
... )
>>> layer(jnp.ones(3)).shape  # (3,) -> (8,)
(8,)
Source code in .venv/lib/python3.12/site-packages/geonnax/siren.py
@classmethod
def init(
    cls,
    in_features: int,
    out_features: int,
    *,
    key: Array,
    omega: float = 30.0,
    layer_type: SirenLayerType = "hidden",
    c: float = 6.0,
) -> SirenDense:
    """Construct a ``SirenDense`` with Sitzmann-regime weight initialisation.

    Examples:
        >>> import jax.numpy as jnp, jax.random as jr
        >>> from geonnax.siren import SirenDense
        >>> layer = SirenDense.init(
        ...     3, 8, key=jr.PRNGKey(0), layer_type="first"
        ... )
        >>> layer(jnp.ones(3)).shape  # (3,) -> (8,)
        (8,)
    """
    _require_positive(
        in_features=in_features,
        out_features=out_features,
        omega=omega,
        c=c,
    )
    k_w, k_b = jax.random.split(key)
    # siren_W_limit validates layer_type.
    w_limit = siren_W_limit(layer_type, in_features, omega, c)
    W = jax.random.uniform(
        k_w, (in_features, out_features), minval=-w_limit, maxval=w_limit
    )
    b_limit = 1.0 / math.sqrt(in_features)
    b = jax.random.uniform(k_b, (out_features,), minval=-b_limit, maxval=b_limit)
    return cls(
        W=W,
        b=b,
        omega=omega,
        in_features=in_features,
        out_features=out_features,
        layer_type=layer_type,
        c=c,
    )

SIREN

Bases: Module

Multi-layer sinusoidal representation network (Sitzmann et al., NeurIPS 2020).

Topology:

\[ \begin{aligned} z_1 &= \sin(\omega_0 (W_0 x + b_0)), \\ z_{i+1} &= \sin(\omega (W_i z_i + b_i)), \quad i = 1 \ldots L-1, \\ y &= W_L z_L + b_L. \end{aligned} \]

Each layer uses the corresponding Sitzmann Theorem 1 init regime (SirenDense): "first" for layer 0, "hidden" for intermediate layers, and "last" for the readout.

depth counts all layers including the readout; depth=2 gives one first-layer + one last-layer (no hidden layers); depth=5 gives first + 3 hidden + last. Must be ≥ 2.

Source code in .venv/lib/python3.12/site-packages/geonnax/siren.py
class SIREN(eqx.Module):
    r"""Multi-layer sinusoidal representation network (Sitzmann et al., NeurIPS 2020).

    Topology:

    $$
    \begin{aligned}
    z_1 &= \sin(\omega_0 (W_0 x + b_0)), \\
    z_{i+1} &= \sin(\omega (W_i z_i + b_i)), \quad i = 1 \ldots L-1, \\
    y &= W_L z_L + b_L.
    \end{aligned}
    $$


    Each layer uses the corresponding Sitzmann Theorem 1 init regime
    (`SirenDense`):  ``"first"`` for layer 0, ``"hidden"`` for
    intermediate layers, and ``"last"`` for the readout.

    ``depth`` counts *all* layers including the readout; ``depth=2`` gives
    one first-layer + one last-layer (no hidden layers); ``depth=5`` gives
    first + 3 hidden + last.  Must be ≥ 2.
    """

    layers: list[SirenDense]
    in_features: int = eqx.field(static=True)
    hidden_features: int = eqx.field(static=True)
    out_features: int = eqx.field(static=True)
    depth: int = eqx.field(static=True)
    first_omega: float = eqx.field(static=True)
    hidden_omega: float = eqx.field(static=True)

    @classmethod
    def init(
        cls,
        in_features: int,
        hidden_features: int,
        out_features: int,
        *,
        depth: int,
        key: Array,
        first_omega: float = 30.0,
        hidden_omega: float = 30.0,
        c: float = 6.0,
    ) -> SIREN:
        """Construct a SIREN with the correct per-layer init regimes.

        Examples:
            >>> import jax.numpy as jnp, jax.random as jr
            >>> from geonnax.siren import SIREN
            >>> net = SIREN.init(2, 16, 1, depth=4, key=jr.PRNGKey(0))
            >>> net(jnp.zeros(2)).shape  # (2,) -> (1,)
            (1,)
            >>> len(net.layers)  # first + 2 hidden + last
            4
        """
        if depth < 2:
            raise ValueError(f"depth must be >= 2 (first + last); got depth={depth}")
        _require_positive(
            in_features=in_features,
            hidden_features=hidden_features,
            out_features=out_features,
            first_omega=first_omega,
            hidden_omega=hidden_omega,
            c=c,
        )
        specs = build_siren_specs(
            in_features,
            hidden_features,
            out_features,
            depth,
            first_omega,
            hidden_omega,
            c,
        )
        keys = jax.random.split(key, depth)
        layers = [
            SirenDense.init(
                spec.in_features,
                spec.out_features,
                key=k,
                omega=spec.omega,
                layer_type=spec.layer_type,
                c=spec.c,
            )
            for spec, k in zip(specs, keys, strict=True)
        ]
        return cls(
            layers=layers,
            in_features=in_features,
            hidden_features=hidden_features,
            out_features=out_features,
            depth=depth,
            first_omega=first_omega,
            hidden_omega=hidden_omega,
        )

    def __call__(self, x: Float[Array, " D_in"]) -> Float[Array, " D_out"]:
        r"""Run the full forward pass, composing the sine-activated layers.

        Args:
            x: Input vector of shape ``(in_features,)``.

        Returns:
            Output vector of shape ``(out_features,)``.

        Examples:
            >>> import jax.numpy as jnp, jax.random as jr
            >>> from geonnax.siren import SIREN
            >>> net = SIREN.init(3, 32, 2, depth=3, key=jr.PRNGKey(0))
            >>> net(jnp.ones(3)).shape  # (3,) -> (2,)
            (2,)
        """
        # z_1 = sin(ω₀(W₀x+b₀)); z_{i+1}=sin(ω(W_i z_i+b_i)); y=W_L z_L+b_L.
        # Shapes: (D_in,) -> (H,) -> … -> (H,) -> (D_out,).
        z = x
        for layer in self.layers:
            z = layer(z)
        return z

init(in_features: int, hidden_features: int, out_features: int, *, depth: int, key: Array, first_omega: float = 30.0, hidden_omega: float = 30.0, c: float = 6.0) -> SIREN classmethod

Construct a SIREN with the correct per-layer init regimes.

Examples:

>>> import jax.numpy as jnp, jax.random as jr
>>> from geonnax.siren import SIREN
>>> net = SIREN.init(2, 16, 1, depth=4, key=jr.PRNGKey(0))
>>> net(jnp.zeros(2)).shape  # (2,) -> (1,)
(1,)
>>> len(net.layers)  # first + 2 hidden + last
4
Source code in .venv/lib/python3.12/site-packages/geonnax/siren.py
@classmethod
def init(
    cls,
    in_features: int,
    hidden_features: int,
    out_features: int,
    *,
    depth: int,
    key: Array,
    first_omega: float = 30.0,
    hidden_omega: float = 30.0,
    c: float = 6.0,
) -> SIREN:
    """Construct a SIREN with the correct per-layer init regimes.

    Examples:
        >>> import jax.numpy as jnp, jax.random as jr
        >>> from geonnax.siren import SIREN
        >>> net = SIREN.init(2, 16, 1, depth=4, key=jr.PRNGKey(0))
        >>> net(jnp.zeros(2)).shape  # (2,) -> (1,)
        (1,)
        >>> len(net.layers)  # first + 2 hidden + last
        4
    """
    if depth < 2:
        raise ValueError(f"depth must be >= 2 (first + last); got depth={depth}")
    _require_positive(
        in_features=in_features,
        hidden_features=hidden_features,
        out_features=out_features,
        first_omega=first_omega,
        hidden_omega=hidden_omega,
        c=c,
    )
    specs = build_siren_specs(
        in_features,
        hidden_features,
        out_features,
        depth,
        first_omega,
        hidden_omega,
        c,
    )
    keys = jax.random.split(key, depth)
    layers = [
        SirenDense.init(
            spec.in_features,
            spec.out_features,
            key=k,
            omega=spec.omega,
            layer_type=spec.layer_type,
            c=spec.c,
        )
        for spec, k in zip(specs, keys, strict=True)
    ]
    return cls(
        layers=layers,
        in_features=in_features,
        hidden_features=hidden_features,
        out_features=out_features,
        depth=depth,
        first_omega=first_omega,
        hidden_omega=hidden_omega,
    )

BayesianSIREN

Bases: PyroxModule

SIREN with regime-scaled Normal priors on all layer weights.

Replaces the deterministic weight matrices of SIREN with NumPyro sample sites. For layer \(i\) with Sitzmann Theorem 1 half-width \(a_i\) (the uniform bound used by SirenDense):

\[ W_i \sim \mathcal{N}\!\left(0,\, \sigma_0 \cdot \frac{a_i}{\sqrt{3}}\right), \qquad b_i \sim \mathcal{N}\!\left(0,\, \sigma_0 \cdot \frac{1}{\sqrt{3 \, d_i}}\right), \]

where \(\sigma_0\) is prior_std and \(d_i\) is the input dimension of layer \(i\). The \(a_i / \sqrt{3}\) factor makes \(\operatorname{Var}(W_i)\) equal to the variance of Sitzmann's \(\mathcal{U}(-a_i, a_i)\) init exactly, so the Bayesian prior preserves the activation variance prescribed by Theorem 1 — avoiding the saturated-sine pathology that a flat \(\mathcal{N}(0, 1)\) prior would cause.

Registered sites: {scope}.layer_0.W, {scope}.layer_0.b, …, {scope}.layer_{depth-1}.W, {scope}.layer_{depth-1}.b — exactly 2 · depth sites per forward call.

Attributes:

Name Type Description
specs tuple[SirenLayerSpec, ...]

Tuple of per-layer specs (static). Holds each layer's layer_type, in_features, out_features, omega, and c — i.e. everything needed to scale the priors.

in_features int

Input dimension.

hidden_features int

Hidden dimension.

out_features int

Output dimension.

depth int

Total layers including readout. Must be ≥ 2.

first_omega float

Frequency multiplier for the first layer.

hidden_omega float

Frequency multiplier for hidden layers.

prior_std float

Scale factor for the regime-scaled Normal prior (default 1.0).

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Examples:

>>> import jax.random as jr, jax.numpy as jnp
>>> from numpyro import handlers
>>> net = BayesianSIREN.init(2, 32, 1, depth=3)
>>> with handlers.seed(rng_seed=0):
...     y = net(jnp.zeros((4, 2)))
>>> y.shape
(4, 1)
Source code in packages/pyrox-nn/src/pyrox_nn/_siren.py
class BayesianSIREN(PyroxModule):
    r"""SIREN with regime-scaled Normal priors on all layer weights.

    Replaces the deterministic weight matrices of `SIREN` with NumPyro
    sample sites.  For layer $i$ with Sitzmann Theorem 1 half-width
    $a_i$ (the uniform bound used by `SirenDense`):

    $$
    W_i \sim \mathcal{N}\!\left(0,\, \sigma_0 \cdot \frac{a_i}{\sqrt{3}}\right),
    \qquad
    b_i \sim \mathcal{N}\!\left(0,\,
        \sigma_0 \cdot \frac{1}{\sqrt{3 \, d_i}}\right),
    $$

    where $\sigma_0$ is ``prior_std`` and $d_i$ is the input
    dimension of layer $i$.  The $a_i / \sqrt{3}$ factor makes
    $\operatorname{Var}(W_i)$ equal to the variance of Sitzmann's
    $\mathcal{U}(-a_i, a_i)$ init exactly, so the Bayesian prior
    preserves the activation variance prescribed by Theorem 1 — avoiding
    the saturated-sine pathology that a flat $\mathcal{N}(0, 1)$
    prior would cause.

    Registered sites: ``{scope}.layer_0.W``, ``{scope}.layer_0.b``, …,
    ``{scope}.layer_{depth-1}.W``, ``{scope}.layer_{depth-1}.b``
    — exactly ``2 · depth`` sites per forward call.

    Attributes:
        specs: Tuple of per-layer specs (static).  Holds each layer's
            ``layer_type``, ``in_features``, ``out_features``, ``omega``,
            and ``c`` — i.e. everything needed to scale the priors.
        in_features: Input dimension.
        hidden_features: Hidden dimension.
        out_features: Output dimension.
        depth: Total layers including readout.  Must be ≥ 2.
        first_omega: Frequency multiplier for the first layer.
        hidden_omega: Frequency multiplier for hidden layers.
        prior_std: Scale factor for the regime-scaled Normal prior (default 1.0).
        pyrox_name: Explicit scope name for NumPyro site registration.

    Examples:
        >>> import jax.random as jr, jax.numpy as jnp
        >>> from numpyro import handlers
        >>> net = BayesianSIREN.init(2, 32, 1, depth=3)
        >>> with handlers.seed(rng_seed=0):
        ...     y = net(jnp.zeros((4, 2)))
        >>> y.shape
        (4, 1)
    """

    specs: tuple[SirenLayerSpec, ...] = eqx.field(static=True)
    in_features: int = eqx.field(static=True)
    hidden_features: int = eqx.field(static=True)
    out_features: int = eqx.field(static=True)
    depth: int = eqx.field(static=True)
    first_omega: float = eqx.field(static=True)
    hidden_omega: float = eqx.field(static=True)
    prior_std: float = eqx.field(static=True, default=1.0)
    pyrox_name: str | None = eqx.field(static=True, default=None)

    @classmethod
    def init(
        cls,
        in_features: int,
        hidden_features: int,
        out_features: int,
        *,
        depth: int,
        first_omega: float = 30.0,
        hidden_omega: float = 30.0,
        c: float = 6.0,
        prior_std: float = 1.0,
        pyrox_name: str | None = None,
    ) -> BayesianSIREN:
        """Construct a `BayesianSIREN`.

        All weights come from the prior, so no PRNG key is needed at
        construction time — the key enters when sampling inside a
        ``numpyro`` handler (``handlers.seed``, SVI, etc.).

        Args:
            in_features: Input dimension.
            hidden_features: Hidden dimension.
            out_features: Output dimension.
            depth: Total layers including readout.  Must be ≥ 2.
            first_omega: Frequency for the first layer.
            hidden_omega: Frequency for hidden layers.
            c: Theorem-1 constant.
            prior_std: Scale factor for the Normal priors (default 1.0, must be > 0).
            pyrox_name: Optional explicit scope name for NumPyro.

        Returns:
            Initialised `BayesianSIREN`.

        Raises:
            ValueError: If ``depth < 2``, or any of the feature dimensions,
                omegas, ``c``, or ``prior_std`` is non-positive.
        """
        if depth < 2:
            raise ValueError(f"depth must be >= 2 (first + last); got depth={depth}")
        _require_positive(
            in_features=in_features,
            hidden_features=hidden_features,
            out_features=out_features,
            first_omega=first_omega,
            hidden_omega=hidden_omega,
            c=c,
            prior_std=prior_std,
        )
        specs = build_siren_specs(
            in_features,
            hidden_features,
            out_features,
            depth,
            first_omega,
            hidden_omega,
            c,
        )
        return cls(
            specs=specs,
            in_features=in_features,
            hidden_features=hidden_features,
            out_features=out_features,
            depth=depth,
            first_omega=first_omega,
            hidden_omega=hidden_omega,
            prior_std=prior_std,
            pyrox_name=pyrox_name,
        )

    @pyrox_method
    def __call__(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D_out"]:
        """Sample weights from regime-scaled priors and run the forward pass.

        Registers ``layer_{i}.W`` and ``layer_{i}.b`` NumPyro sample sites
        for each layer ``i`` in ``[0, depth)``.

        Args:
            x: Input tensor of shape ``(*batch, in_features)``.

        Returns:
            Output tensor of shape ``(*batch, out_features)``.
        """
        # Normal stddev = (Uniform half-width) / √3 so Var(W) matches
        # Sitzmann's U(-a, a) init exactly.
        inv_sqrt3 = 1.0 / math.sqrt(3.0)
        z = x
        for i, spec in enumerate(self.specs):
            a = siren_W_limit(spec.layer_type, spec.in_features, spec.omega, spec.c)
            w_scale = self.prior_std * a * inv_sqrt3
            b_scale = self.prior_std * inv_sqrt3 / math.sqrt(spec.in_features)
            W = self.pyrox_sample(
                f"layer_{i}.W",
                dist.Normal(0.0, w_scale)
                .expand([spec.in_features, spec.out_features])
                .to_event(2),
            )
            b = self.pyrox_sample(
                f"layer_{i}.b",
                dist.Normal(0.0, b_scale).expand([spec.out_features]).to_event(1),
            )
            pre = einx.dot("... i, i o -> ... o", z, W) + b
            z = pre if spec.layer_type == "last" else jnp.sin(spec.omega * pre)
        return z

init(in_features: int, hidden_features: int, out_features: int, *, depth: int, first_omega: float = 30.0, hidden_omega: float = 30.0, c: float = 6.0, prior_std: float = 1.0, pyrox_name: str | None = None) -> BayesianSIREN classmethod

Construct a BayesianSIREN.

All weights come from the prior, so no PRNG key is needed at construction time — the key enters when sampling inside a numpyro handler (handlers.seed, SVI, etc.).

Parameters:

Name Type Description Default
in_features int

Input dimension.

required
hidden_features int

Hidden dimension.

required
out_features int

Output dimension.

required
depth int

Total layers including readout. Must be ≥ 2.

required
first_omega float

Frequency for the first layer.

30.0
hidden_omega float

Frequency for hidden layers.

30.0
c float

Theorem-1 constant.

6.0
prior_std float

Scale factor for the Normal priors (default 1.0, must be > 0).

1.0
pyrox_name str | None

Optional explicit scope name for NumPyro.

None

Returns:

Type Description
BayesianSIREN

Initialised BayesianSIREN.

Raises:

Type Description
ValueError

If depth < 2, or any of the feature dimensions, omegas, c, or prior_std is non-positive.

Source code in packages/pyrox-nn/src/pyrox_nn/_siren.py
@classmethod
def init(
    cls,
    in_features: int,
    hidden_features: int,
    out_features: int,
    *,
    depth: int,
    first_omega: float = 30.0,
    hidden_omega: float = 30.0,
    c: float = 6.0,
    prior_std: float = 1.0,
    pyrox_name: str | None = None,
) -> BayesianSIREN:
    """Construct a `BayesianSIREN`.

    All weights come from the prior, so no PRNG key is needed at
    construction time — the key enters when sampling inside a
    ``numpyro`` handler (``handlers.seed``, SVI, etc.).

    Args:
        in_features: Input dimension.
        hidden_features: Hidden dimension.
        out_features: Output dimension.
        depth: Total layers including readout.  Must be ≥ 2.
        first_omega: Frequency for the first layer.
        hidden_omega: Frequency for hidden layers.
        c: Theorem-1 constant.
        prior_std: Scale factor for the Normal priors (default 1.0, must be > 0).
        pyrox_name: Optional explicit scope name for NumPyro.

    Returns:
        Initialised `BayesianSIREN`.

    Raises:
        ValueError: If ``depth < 2``, or any of the feature dimensions,
            omegas, ``c``, or ``prior_std`` is non-positive.
    """
    if depth < 2:
        raise ValueError(f"depth must be >= 2 (first + last); got depth={depth}")
    _require_positive(
        in_features=in_features,
        hidden_features=hidden_features,
        out_features=out_features,
        first_omega=first_omega,
        hidden_omega=hidden_omega,
        c=c,
        prior_std=prior_std,
    )
    specs = build_siren_specs(
        in_features,
        hidden_features,
        out_features,
        depth,
        first_omega,
        hidden_omega,
        c,
    )
    return cls(
        specs=specs,
        in_features=in_features,
        hidden_features=hidden_features,
        out_features=out_features,
        depth=depth,
        first_omega=first_omega,
        hidden_omega=hidden_omega,
        prior_std=prior_std,
        pyrox_name=pyrox_name,
    )

SNGP — spectral-normalised GP head

The SNGP output layer (Liu et al., 2020): a random-feature GP last layer whose posterior covariance comes from a Laplace approximation (LaplaceRandomFeatureCovariance, re-exported from geonnax), giving distance-aware uncertainty from a single deterministic forward pass.

RandomFeatureGaussianProcess

Bases: PyroxModule

SNGP output layer (Liu et al., 2020).

A random Fourier feature map followed by a learnable linear head, plus a Laplace-approximation covariance over the linear weights. The forward pass returns the mean prediction and (optionally) a per-input predictive variance summarising distance from the training distribution.

Forward (mean):

\[ \phi(x) = \sqrt{\tfrac{2}{D}}\,\cos\!\bigl(W\, x / \ell + b\bigr), \qquad \mu(x) = \phi(x)\, H + b_H. \]

The frequencies \(W\) and bias \(b\) of the RFF map are frozen (they implicitly define the kernel approximation): they are registered as pyrox_param sites for substitution and checkpointing, then guarded with jax.lax.stop_gradient inside feature_map so SGD-style optimisers leave them untouched. The lengthscale \(\ell\), the linear head \(H, b_H\), and the Laplace precision are the trainable / updated quantities.

Predictive variance — when \(\hat{\Lambda}\) is the current precision matrix:

\[ \sigma^2(x_*) = \phi(x_*)^\top \hat{\Lambda}^{-1}\, \phi(x_*). \]

Training pattern (one minibatch):

  1. mean = layer(x) registers / reuses the trainable params and returns the mean prediction. Compute the loss, take a gradient step on the SVI parameter store as usual.
  2. After the gradient step, call new_layer = layer.update_precision(features) where features is the result of feature_map evaluated on the same minibatch using the updated parameters. This returns a new layer with the LRFC's precision EMA-updated.

At inference, mean, var = layer(x, return_cov=True) produces the mean and the Laplace per-input predictive variance.

Plate semantics

Same as the rest of pyrox_nn's Bayesian / heteroscedastic dense layers — call this layer outside numpyro.plate("data", ..., subsample_size=...) and only plate the observation likelihood.

Attributes:

Name Type Description
in_features int

Input dimension \(D_\mathrm{in}\).

num_features int

Number of random Fourier features \(D\).

out_features int

Output dimension \(D_\mathrm{out}\).

init_lengthscale float

Initial lengthscale \(\ell\). Optimised during training as a positive pyrox_param.

W_init Float[Array, 'D_in D']

Frozen RFF frequencies, shape (D_in, D), drawn from a standard Normal (the RBF spectral density).

bias_init Float[Array, ' D']

Frozen RFF biases, shape (D,), drawn from Uniform(0, 2 pi).

output_linear_init Float[Array, 'D D_out']

Init for the linear head, shape (D, D_out).

covariance LaplaceRandomFeatureCovariance

The LaplaceRandomFeatureCovariance instance.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

References

Liu, J. Z., et al. (2020). Simple and Principled Uncertainty Estimation with Deterministic Deep Learning via Distance Awareness. NeurIPS.

Source code in packages/pyrox-nn/src/pyrox_nn/_sngp.py
class RandomFeatureGaussianProcess(PyroxModule):
    r"""SNGP output layer (Liu et al., 2020).

    A random Fourier feature map followed by a learnable linear head,
    plus a Laplace-approximation covariance over the linear weights.
    The forward pass returns the mean prediction and (optionally) a
    per-input predictive variance summarising distance from the
    training distribution.

    Forward (mean):

    $$
    \phi(x) = \sqrt{\tfrac{2}{D}}\,\cos\!\bigl(W\, x / \ell + b\bigr),
    \qquad \mu(x) = \phi(x)\, H + b_H.
    $$

    The frequencies $W$ and bias $b$ of the RFF map are
    *frozen* (they implicitly define the kernel approximation): they
    are registered as ``pyrox_param`` sites for substitution and
    checkpointing, then guarded with `jax.lax.stop_gradient`
    inside `feature_map` so SGD-style optimisers leave them
    untouched. The lengthscale $\ell$, the linear head
    $H, b_H$, and the Laplace precision are the trainable /
    updated quantities.

    Predictive variance — when $\hat{\Lambda}$ is the current
    precision matrix:

    $$
    \sigma^2(x_*) = \phi(x_*)^\top \hat{\Lambda}^{-1}\, \phi(x_*).
    $$

    Training pattern (one minibatch):

    1. ``mean = layer(x)`` registers / reuses the trainable params and
       returns the mean prediction. Compute the loss, take a gradient
       step on the SVI parameter store as usual.
    2. After the gradient step, call
       ``new_layer = layer.update_precision(features)`` where
       ``features`` is the result of `feature_map` evaluated on
       the same minibatch using the *updated* parameters. This returns
       a new layer with the LRFC's precision EMA-updated.

    At inference, ``mean, var = layer(x, return_cov=True)`` produces
    the mean and the Laplace per-input predictive variance.

    Plate semantics:
        Same as the rest of ``pyrox_nn``'s Bayesian / heteroscedastic
        dense layers — call this layer outside
        ``numpyro.plate("data", ..., subsample_size=...)`` and only
        plate the observation likelihood.

    Attributes:
        in_features: Input dimension $D_\mathrm{in}$.
        num_features: Number of random Fourier features $D$.
        out_features: Output dimension $D_\mathrm{out}$.
        init_lengthscale: Initial lengthscale $\ell$. Optimised
            during training as a positive ``pyrox_param``.
        W_init: Frozen RFF frequencies, shape ``(D_in, D)``, drawn from
            a standard Normal (the RBF spectral density).
        bias_init: Frozen RFF biases, shape ``(D,)``, drawn from
            ``Uniform(0, 2 pi)``.
        output_linear_init: Init for the linear head, shape ``(D, D_out)``.
        covariance: The `LaplaceRandomFeatureCovariance` instance.
        pyrox_name: Explicit scope name for NumPyro site registration.

    References:
        Liu, J. Z., et al. (2020). *Simple and Principled Uncertainty
        Estimation with Deterministic Deep Learning via Distance
        Awareness.* NeurIPS.
    """

    core: geonnax.RandomFeatureGaussianProcess
    pyrox_name: str | None = None

    @classmethod
    def init(
        cls,
        key: PRNGKeyArray,
        in_features: int,
        num_features: int,
        out_features: int,
        *,
        init_lengthscale: float = 1.0,
        momentum: float = 0.999,
        ridge: float = 1.0,
        head_scale: float = 0.01,
        pyrox_name: str | None = None,
    ) -> RandomFeatureGaussianProcess:
        """Construct an SNGP head with frozen RFF freqs and an empty precision."""
        # geonnax validates positive dims and init_lengthscale > 0;
        # momentum / ridge constraints are validated inside the LRFC init.
        core = geonnax.RandomFeatureGaussianProcess.init(
            in_features=in_features,
            num_features=num_features,
            out_features=out_features,
            key=key,
            init_lengthscale=init_lengthscale,
            momentum=momentum,
            ridge=ridge,
            head_scale=head_scale,
        )
        return cls(core=core, pyrox_name=pyrox_name)

    # Read-only property accessors retain the pre-refactor attribute names so
    # external callers (tests, user code reading static dims) keep working.
    @property
    def in_features(self) -> int:
        return self.core.in_features

    @property
    def num_features(self) -> int:
        return self.core.num_features

    @property
    def out_features(self) -> int:
        return self.core.out_features

    @property
    def init_lengthscale(self) -> float:
        return float(self.core.lengthscale)

    @property
    def W_init(self) -> Float[Array, "D_in D"]:
        return self.core.W

    @property
    def bias_init(self) -> Float[Array, " D"]:
        return self.core.bias

    @property
    def output_linear_init(self) -> Float[Array, "D D_out"]:
        return self.core.output_linear

    @property
    def covariance(self) -> geonnax.LaplaceRandomFeatureCovariance:
        return self.core.covariance

    def _swap_feature_core(self) -> geonnax.RandomFeatureGaussianProcess:
        """Register the RFF-map params and swap them into the core.

        Only the frequency/bias/lengthscale arrays are needed for the
        feature map; the linear-head params are registered inside
        ``__call__`` so disabled branches (e.g. pure feature-map usage
        via `feature_map`) don't materialise unused sites.
        """
        W = jax.lax.stop_gradient(self.pyrox_param("W", self.core.W))
        b = jax.lax.stop_gradient(self.pyrox_param("bias", self.core.bias))
        ls = self.pyrox_param(
            "lengthscale",
            self.core.lengthscale,
            constraint=dist.constraints.positive,
        )
        return eqx.tree_at(
            lambda c: (c.W, c.bias, c.lengthscale),
            self.core,
            (W, b, ls),
        )

    @pyrox_method
    def feature_map(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D"]:
        r"""Random Fourier feature map: $\phi(x) = \sqrt{2/D}\,\cos(Wx/\ell + b)$.

        Frequencies and bias are registered as ``pyrox_param`` sites for
        substitution / checkpointing, but `jax.lax.stop_gradient`
        is applied so SVI's gradient-based optimisers leave them
        frozen at their init values. The lengthscale is the active
        bandwidth control and is constrained positive.
        """
        new_core = self._swap_feature_core()
        # geonnax feature_map is single-example `(D_in,) -> (D,)`.
        return vmap_over_flat_batch(new_core.feature_map, x)

    @pyrox_method
    def __call__(
        self,
        x: Float[Array, "*batch D_in"],
        *,
        return_cov: bool = False,
    ) -> (
        Float[Array, "*batch D_out"]
        | tuple[Float[Array, "*batch D_out"], Float[Array, " *batch"]]
    ):
        # Register the RFF + linear-head params and swap them into the core.
        new_core = self._swap_feature_core()
        H = self.pyrox_param("output_linear", self.core.output_linear)
        b_out = self.pyrox_param(
            "output_bias", jnp.zeros(self.out_features, dtype=x.dtype)
        )
        new_core = eqx.tree_at(
            lambda c: (c.output_linear, c.output_bias),
            new_core,
            (H, b_out),
        )

        # geonnax `__call__` is single-example `(D_in,)`. `return_cov=False`
        # keeps the output a plain array; `return_cov=True` returns
        # (mean, var) per example and the helper restores both leaves.
        if return_cov:
            mean, var = vmap_over_flat_batch(
                lambda xi: new_core(xi, return_cov=True), x
            )
            return mean, var

        return vmap_over_flat_batch(new_core, x)

    def update_precision(
        self, features: Float[Array, "*batch D"]
    ) -> RandomFeatureGaussianProcess:
        """Return a new layer with an EMA-updated Laplace precision.

        Pure-functional: ``self`` is unchanged. Pass features computed
        on the current minibatch (e.g. via `feature_map`) — the
        update folds the empirical second moment into the EMA. Call
        this once per training batch *after* the gradient step.
        """
        # Flatten any leading batch dims down to the single batch axis the
        # geonnax `update_precision` expects.
        flat = einx.id("b... d -> (b...) d", features)
        new_core = self.core.update_precision(flat)
        return eqx.tree_at(lambda layer: layer.core, self, new_core)

init(key: PRNGKeyArray, in_features: int, num_features: int, out_features: int, *, init_lengthscale: float = 1.0, momentum: float = 0.999, ridge: float = 1.0, head_scale: float = 0.01, pyrox_name: str | None = None) -> RandomFeatureGaussianProcess classmethod

Construct an SNGP head with frozen RFF freqs and an empty precision.

Source code in packages/pyrox-nn/src/pyrox_nn/_sngp.py
@classmethod
def init(
    cls,
    key: PRNGKeyArray,
    in_features: int,
    num_features: int,
    out_features: int,
    *,
    init_lengthscale: float = 1.0,
    momentum: float = 0.999,
    ridge: float = 1.0,
    head_scale: float = 0.01,
    pyrox_name: str | None = None,
) -> RandomFeatureGaussianProcess:
    """Construct an SNGP head with frozen RFF freqs and an empty precision."""
    # geonnax validates positive dims and init_lengthscale > 0;
    # momentum / ridge constraints are validated inside the LRFC init.
    core = geonnax.RandomFeatureGaussianProcess.init(
        in_features=in_features,
        num_features=num_features,
        out_features=out_features,
        key=key,
        init_lengthscale=init_lengthscale,
        momentum=momentum,
        ridge=ridge,
        head_scale=head_scale,
    )
    return cls(core=core, pyrox_name=pyrox_name)

feature_map(x: Float[Array, '*batch D_in']) -> Float[Array, '*batch D']

Random Fourier feature map: \(\phi(x) = \sqrt{2/D}\,\cos(Wx/\ell + b)\).

Frequencies and bias are registered as pyrox_param sites for substitution / checkpointing, but jax.lax.stop_gradient is applied so SVI's gradient-based optimisers leave them frozen at their init values. The lengthscale is the active bandwidth control and is constrained positive.

Source code in packages/pyrox-nn/src/pyrox_nn/_sngp.py
@pyrox_method
def feature_map(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D"]:
    r"""Random Fourier feature map: $\phi(x) = \sqrt{2/D}\,\cos(Wx/\ell + b)$.

    Frequencies and bias are registered as ``pyrox_param`` sites for
    substitution / checkpointing, but `jax.lax.stop_gradient`
    is applied so SVI's gradient-based optimisers leave them
    frozen at their init values. The lengthscale is the active
    bandwidth control and is constrained positive.
    """
    new_core = self._swap_feature_core()
    # geonnax feature_map is single-example `(D_in,) -> (D,)`.
    return vmap_over_flat_batch(new_core.feature_map, x)

update_precision(features: Float[Array, '*batch D']) -> RandomFeatureGaussianProcess

Return a new layer with an EMA-updated Laplace precision.

Pure-functional: self is unchanged. Pass features computed on the current minibatch (e.g. via feature_map) — the update folds the empirical second moment into the EMA. Call this once per training batch after the gradient step.

Source code in packages/pyrox-nn/src/pyrox_nn/_sngp.py
def update_precision(
    self, features: Float[Array, "*batch D"]
) -> RandomFeatureGaussianProcess:
    """Return a new layer with an EMA-updated Laplace precision.

    Pure-functional: ``self`` is unchanged. Pass features computed
    on the current minibatch (e.g. via `feature_map`) — the
    update folds the empirical second moment into the EMA. Call
    this once per training batch *after* the gradient step.
    """
    # Flatten any leading batch dims down to the single batch axis the
    # geonnax `update_precision` expects.
    flat = einx.id("b... d -> (b...) d", features)
    new_core = self.core.update_precision(flat)
    return eqx.tree_at(lambda layer: layer.core, self, new_core)

LaplaceRandomFeatureCovariance

Bases: Module

Laplace-approximation precision for an SNGP output head.

Stores the precision matrix \(\hat{\Lambda} \in \mathbb{R}^{D \times D}\) over the linear weights of the output layer. Updated as an exponential moving average of feature outer products during training:

\[ \hat{\Lambda}_{t+1} \leftarrow m\,\hat{\Lambda}_t + (1 - m)\,\frac{1}{B} \sum_{b=1}^{B} \phi(x_b)\,\phi(x_b)^\top. \]

At test time the predictive variance for a feature vector \(\phi(x_*)\) is

\[ \sigma^2(x_*) = \phi(x_*)^\top \hat{\Sigma}\, \phi(x_*), \qquad \hat{\Sigma} = \hat{\Lambda}^{-1}, \]

computed stably via a Cholesky solve.

The container is pure-functional: update returns a new instance with an updated precision rather than mutating self, matching how Equinox composes immutable PyTrees with optimisers. A small ridge \(\lambda\) initialises the precision at \(\lambda I\) and is also added at solve-time inside covariance and variance_at so the Cholesky stays numerically well-conditioned even after many EMA steps with low momentum (which would otherwise let the ridge contribution decay geometrically and the precision approach singularity). Equivalently \(\hat\Sigma = (\hat\Lambda + \lambda I)^{-1}\) — the Bayesian-linear-regression interpretation of SNGP, where \(\lambda I\) is a Gaussian prior precision on the head weights.

Attributes:

Name Type Description
precision Float[Array, 'D D']

Current precision matrix \(\hat{\Lambda}\).

momentum float

EMA momentum \(m \in [0, 1]\). Higher values give slower updates; 0.999 works well for most settings.

ridge float

Diagonal ridge \(\lambda\). Used both as the init value of precision and as a solve-time jitter to keep the Cholesky well-defined.

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.sngp import LaplaceRandomFeatureCovariance
>>> cov = LaplaceRandomFeatureCovariance.init(4, ridge=1.0)
>>> cov.precision.shape
(4, 4)
>>> cov.variance_at(jnp.eye(4)).shape
(4,)
Source code in .venv/lib/python3.12/site-packages/geonnax/sngp.py
class LaplaceRandomFeatureCovariance(eqx.Module):
    r"""Laplace-approximation precision for an SNGP output head.

    Stores the precision matrix $\hat{\Lambda} \in \mathbb{R}^{D \times D}$
    over the linear weights of the output layer. Updated as an
    exponential moving average of feature outer products during
    training:

    $$
    \hat{\Lambda}_{t+1} \leftarrow m\,\hat{\Lambda}_t
    + (1 - m)\,\frac{1}{B} \sum_{b=1}^{B} \phi(x_b)\,\phi(x_b)^\top.
    $$


    At test time the predictive variance for a feature vector
    $\phi(x_*)$ is

    $$
    \sigma^2(x_*) = \phi(x_*)^\top \hat{\Sigma}\, \phi(x_*),
    \qquad \hat{\Sigma} = \hat{\Lambda}^{-1},
    $$


    computed stably via a Cholesky solve.

    The container is *pure-functional*: `update` returns a new
    instance with an updated precision rather than mutating ``self``,
    matching how Equinox composes immutable PyTrees with optimisers.
    A small ridge $\lambda$ initialises the precision at
    $\lambda I$ and is *also* added at solve-time inside
    `covariance` and `variance_at` so the Cholesky stays
    numerically well-conditioned even after many EMA steps with low
    momentum (which would otherwise let the ridge contribution decay
    geometrically and the precision approach singularity).
    Equivalently $\hat\Sigma = (\hat\Lambda + \lambda I)^{-1}$ —
    the Bayesian-linear-regression interpretation of SNGP, where
    $\lambda I$ is a Gaussian prior precision on the head weights.

    Attributes:
        precision: Current precision matrix $\hat{\Lambda}$.
        momentum: EMA momentum $m \in [0, 1]$. Higher values give
            slower updates; ``0.999`` works well for most settings.
        ridge: Diagonal ridge $\lambda$. Used both as the init
            value of ``precision`` and as a solve-time jitter to keep
            the Cholesky well-defined.

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.sngp import LaplaceRandomFeatureCovariance
        >>> cov = LaplaceRandomFeatureCovariance.init(4, ridge=1.0)
        >>> cov.precision.shape
        (4, 4)
        >>> cov.variance_at(jnp.eye(4)).shape
        (4,)
    """

    precision: Float[Array, "D D"]
    momentum: float = eqx.field(static=True, default=0.999)
    ridge: float = eqx.field(static=True, default=1.0)

    @classmethod
    def init(
        cls,
        num_features: int,
        *,
        momentum: float = 0.999,
        ridge: float = 1.0,
    ) -> LaplaceRandomFeatureCovariance:
        r"""Construct a fresh covariance container with ``ridge * I`` precision.

        Examples:
            >>> from geonnax.sngp import LaplaceRandomFeatureCovariance
            >>> cov = LaplaceRandomFeatureCovariance.init(4, momentum=0.9)
            >>> cov.precision.shape
            (4, 4)
        """
        if num_features <= 0:
            raise ValueError(f"num_features must be > 0; got {num_features}.")
        if not 0.0 <= momentum <= 1.0:
            raise ValueError(f"momentum must lie in [0, 1]; got {momentum}.")
        if ridge <= 0:
            raise ValueError(f"ridge must be > 0; got {ridge}.")
        return cls(
            precision=ridge * jnp.eye(num_features),
            momentum=momentum,
            ridge=ridge,
        )

    def update(self, features: Float[Array, "B D"]) -> LaplaceRandomFeatureCovariance:
        r"""Return a new container with EMA-updated precision.

        $\hat\Lambda \leftarrow m\,\hat\Lambda + (1-m)\,\Phi^\top\Phi/B$.

        Examples:
            >>> import jax.numpy as jnp
            >>> from geonnax.sngp import LaplaceRandomFeatureCovariance
            >>> cov = LaplaceRandomFeatureCovariance.init(4)
            >>> new = cov.update(jnp.ones((8, 4)))
            >>> new is cov
            False
        """
        B = features.shape[0]
        # Feature Gram ΦᵀΦ / B: contract the batch axis b → (D, D).
        outer = einx.dot("b d, b e -> d e", features, features) / B
        new_precision = self.momentum * self.precision + (1.0 - self.momentum) * outer
        return eqx.tree_at(lambda c: c.precision, self, new_precision)

    def _chol(self) -> Float[Array, "D D"]:
        # Symmetrise to absorb floating-point asymmetry in the EMA, then
        # add ridge jitter so the matrix is guaranteed positive-definite
        # regardless of how the EMA has evolved.
        sym = 0.5 * (self.precision + einx.id("i j -> j i", self.precision))
        D = sym.shape[0]
        return jnp.linalg.cholesky(sym + self.ridge * jnp.eye(D, dtype=sym.dtype))

    def covariance(self) -> Float[Array, "D D"]:
        r"""Inverse of the precision matrix (one-shot Cholesky inversion).

        $\hat\Sigma = (\hat\Lambda + \lambda I)^{-1}$, shape ``(D, D)``.

        Examples:
            >>> from geonnax.sngp import LaplaceRandomFeatureCovariance
            >>> cov = LaplaceRandomFeatureCovariance.init(4)
            >>> cov.covariance().shape
            (4, 4)
        """
        L = self._chol()
        D = self.precision.shape[0]
        return jax.scipy.linalg.cho_solve((L, True), jnp.eye(D))

    def variance_at(self, features: Float[Array, "N D"]) -> Float[Array, " N"]:
        r"""Per-row predictive variance $\phi(x_n)^\top \hat{\Sigma}\,\phi(x_n)$.

        Computed via a triangular solve to avoid materialising the full
        $D \times D$ covariance:

        $$
        y = L^{-1} \phi(x_n)^\top, \qquad
        \sigma^2(x_n) = \lVert y \rVert_2^2
        = \phi(x_n)^\top (L L^\top)^{-1} \phi(x_n).
        $$


        Examples:
            >>> import jax.numpy as jnp
            >>> from geonnax.sngp import LaplaceRandomFeatureCovariance
            >>> cov = LaplaceRandomFeatureCovariance.init(4)
            >>> cov.variance_at(jnp.eye(4)).shape
            (4,)
        """
        L = self._chol()
        # Solve L y = Φᵀ ; (D, D)\(D, N) -> (D, N), one column per row of Φ.
        y = jax.scipy.linalg.solve_triangular(
            L, einx.id("n d -> d n", features), lower=True
        )
        # σ²(xₙ) = ‖yₙ‖² ; sum over D -> (N,)
        return einx.sum("[d] n", y * y)

init(num_features: int, *, momentum: float = 0.999, ridge: float = 1.0) -> LaplaceRandomFeatureCovariance classmethod

Construct a fresh covariance container with ridge * I precision.

Examples:

>>> from geonnax.sngp import LaplaceRandomFeatureCovariance
>>> cov = LaplaceRandomFeatureCovariance.init(4, momentum=0.9)
>>> cov.precision.shape
(4, 4)
Source code in .venv/lib/python3.12/site-packages/geonnax/sngp.py
@classmethod
def init(
    cls,
    num_features: int,
    *,
    momentum: float = 0.999,
    ridge: float = 1.0,
) -> LaplaceRandomFeatureCovariance:
    r"""Construct a fresh covariance container with ``ridge * I`` precision.

    Examples:
        >>> from geonnax.sngp import LaplaceRandomFeatureCovariance
        >>> cov = LaplaceRandomFeatureCovariance.init(4, momentum=0.9)
        >>> cov.precision.shape
        (4, 4)
    """
    if num_features <= 0:
        raise ValueError(f"num_features must be > 0; got {num_features}.")
    if not 0.0 <= momentum <= 1.0:
        raise ValueError(f"momentum must lie in [0, 1]; got {momentum}.")
    if ridge <= 0:
        raise ValueError(f"ridge must be > 0; got {ridge}.")
    return cls(
        precision=ridge * jnp.eye(num_features),
        momentum=momentum,
        ridge=ridge,
    )

update(features: Float[Array, 'B D']) -> LaplaceRandomFeatureCovariance

Return a new container with EMA-updated precision.

\(\hat\Lambda \leftarrow m\,\hat\Lambda + (1-m)\,\Phi^\top\Phi/B\).

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.sngp import LaplaceRandomFeatureCovariance
>>> cov = LaplaceRandomFeatureCovariance.init(4)
>>> new = cov.update(jnp.ones((8, 4)))
>>> new is cov
False
Source code in .venv/lib/python3.12/site-packages/geonnax/sngp.py
def update(self, features: Float[Array, "B D"]) -> LaplaceRandomFeatureCovariance:
    r"""Return a new container with EMA-updated precision.

    $\hat\Lambda \leftarrow m\,\hat\Lambda + (1-m)\,\Phi^\top\Phi/B$.

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.sngp import LaplaceRandomFeatureCovariance
        >>> cov = LaplaceRandomFeatureCovariance.init(4)
        >>> new = cov.update(jnp.ones((8, 4)))
        >>> new is cov
        False
    """
    B = features.shape[0]
    # Feature Gram ΦᵀΦ / B: contract the batch axis b → (D, D).
    outer = einx.dot("b d, b e -> d e", features, features) / B
    new_precision = self.momentum * self.precision + (1.0 - self.momentum) * outer
    return eqx.tree_at(lambda c: c.precision, self, new_precision)

covariance() -> Float[Array, 'D D']

Inverse of the precision matrix (one-shot Cholesky inversion).

\(\hat\Sigma = (\hat\Lambda + \lambda I)^{-1}\), shape (D, D).

Examples:

>>> from geonnax.sngp import LaplaceRandomFeatureCovariance
>>> cov = LaplaceRandomFeatureCovariance.init(4)
>>> cov.covariance().shape
(4, 4)
Source code in .venv/lib/python3.12/site-packages/geonnax/sngp.py
def covariance(self) -> Float[Array, "D D"]:
    r"""Inverse of the precision matrix (one-shot Cholesky inversion).

    $\hat\Sigma = (\hat\Lambda + \lambda I)^{-1}$, shape ``(D, D)``.

    Examples:
        >>> from geonnax.sngp import LaplaceRandomFeatureCovariance
        >>> cov = LaplaceRandomFeatureCovariance.init(4)
        >>> cov.covariance().shape
        (4, 4)
    """
    L = self._chol()
    D = self.precision.shape[0]
    return jax.scipy.linalg.cho_solve((L, True), jnp.eye(D))

variance_at(features: Float[Array, 'N D']) -> Float[Array, ' N']

Per-row predictive variance \(\phi(x_n)^\top \hat{\Sigma}\,\phi(x_n)\).

Computed via a triangular solve to avoid materialising the full \(D \times D\) covariance:

\[ y = L^{-1} \phi(x_n)^\top, \qquad \sigma^2(x_n) = \lVert y \rVert_2^2 = \phi(x_n)^\top (L L^\top)^{-1} \phi(x_n). \]

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.sngp import LaplaceRandomFeatureCovariance
>>> cov = LaplaceRandomFeatureCovariance.init(4)
>>> cov.variance_at(jnp.eye(4)).shape
(4,)
Source code in .venv/lib/python3.12/site-packages/geonnax/sngp.py
def variance_at(self, features: Float[Array, "N D"]) -> Float[Array, " N"]:
    r"""Per-row predictive variance $\phi(x_n)^\top \hat{\Sigma}\,\phi(x_n)$.

    Computed via a triangular solve to avoid materialising the full
    $D \times D$ covariance:

    $$
    y = L^{-1} \phi(x_n)^\top, \qquad
    \sigma^2(x_n) = \lVert y \rVert_2^2
    = \phi(x_n)^\top (L L^\top)^{-1} \phi(x_n).
    $$


    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.sngp import LaplaceRandomFeatureCovariance
        >>> cov = LaplaceRandomFeatureCovariance.init(4)
        >>> cov.variance_at(jnp.eye(4)).shape
        (4,)
    """
    L = self._chol()
    # Solve L y = Φᵀ ; (D, D)\(D, N) -> (D, N), one column per row of Φ.
    y = jax.scipy.linalg.solve_triangular(
        L, einx.id("n d -> d n", features), lower=True
    )
    # σ²(xₙ) = ‖yₙ‖² ; sum over D -> (N,)
    return einx.sum("[d] n", y * y)

Deep spectral GPs

DeepVSSGP

Bases: PyroxModule

Deep Random Feature Expansion for Variational SSGP (Cutajar et al. 2017).

A stack of \(L\) variational SSGP layers, each with random spectral frequencies \(\Omega_l\) and random projection weights \(W_l\):

\[ \begin{aligned} F_0 &= X, \\ F_{l+1} &= \Phi_l(F_l;\, \Omega_l, \ell_l)\, W_l, \quad l = 0, \ldots, L-1, \\ \Phi_l(F;\, \Omega_l, \ell_l) &= \sqrt{1/M}\, [\cos(F\,\Omega_l/\ell_l), \sin(F\,\Omega_l/\ell_l)]. \end{aligned} \]

Each layer registers three sample sites:

  • layer_{l}.W_freq — RFF frequencies, prior \(\mathcal{N}(0, 1)\) (RBF spectral density in lengthscale-1 units).
  • layer_{l}.lengthscale — kernel lengthscale, prior \(\mathrm{LogNormal}(\log \ell_{\mathrm{init}}, 1)\).
  • layer_{l}.W_proj — projection weights, prior \(\mathcal{N}(0, \sigma_W^2)\).

Under SVI an AutoNormal learns mean-field Gaussian posteriors over all \(3L\) sites — one MC sample per forward pass gives the doubly-stochastic reparameterised ELBO of Cutajar et al. (2017).

At depth=1 this reduces to a single VSSGP layer mapping in_features -> out_features via the RFF basis (same model class as VariationalFourierFeatures followed by a DenseReparameterization head). Stacking adds non-stationarity at the cost of a non-Gaussian aggregate likelihood — the layer-wise marginalisation that makes single-layer SSGP closed-form is no longer available, hence the variational treatment.

Attributes:

Name Type Description
in_features int

Input dimension \(D_{\mathrm{in}}\).

hidden_features int

Inter-layer dimension \(D_h\) (constant across hidden layers).

out_features int

Output dimension \(D_{\mathrm{out}}\).

n_features int

Per-layer Fourier-feature pair count \(M\) (so each layer's hidden state is \(2M\)-dim before projection).

depth int

Total number of stacked SSGP layers \(L\). Must be \(\ge 1\).

init_lengthscale float

Prior location for each layer's lengthscale.

prior_std float

Standard deviation of the per-layer projection prior \(\mathcal{N}(0, \sigma_W^2)\).

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Examples:

>>> import jax.random as jr, jax.numpy as jnp
>>> from numpyro import handlers
>>> net = DeepVSSGP.init(in_features=2, hidden_features=4,
...                       out_features=1, depth=3, n_features=16)
>>> with handlers.seed(rng_seed=0):
...     y = net(jnp.zeros((8, 2)))
>>> y.shape
(8, 1)
Source code in packages/pyrox-nn/src/pyrox_nn/_vssgp.py
class DeepVSSGP(PyroxModule):
    r"""Deep Random Feature Expansion for Variational SSGP (Cutajar et al. 2017).

    A stack of $L$ variational SSGP layers, each with random
    spectral frequencies $\Omega_l$ and random projection weights
    $W_l$:

    $$
    \begin{aligned}
    F_0 &= X, \\
    F_{l+1} &= \Phi_l(F_l;\, \Omega_l, \ell_l)\, W_l,
        \quad l = 0, \ldots, L-1, \\
    \Phi_l(F;\, \Omega_l, \ell_l) &=
        \sqrt{1/M}\,
        [\cos(F\,\Omega_l/\ell_l), \sin(F\,\Omega_l/\ell_l)].
    \end{aligned}
    $$

    Each layer registers three sample sites:

    - ``layer_{l}.W_freq`` — RFF frequencies, prior $\mathcal{N}(0, 1)$
      (RBF spectral density in lengthscale-1 units).
    - ``layer_{l}.lengthscale`` — kernel lengthscale, prior
      $\mathrm{LogNormal}(\log \ell_{\mathrm{init}}, 1)$.
    - ``layer_{l}.W_proj`` — projection weights, prior
      $\mathcal{N}(0, \sigma_W^2)$.

    Under SVI an
    `AutoNormal` learns mean-field
    Gaussian posteriors over all $3L$ sites — one MC sample per
    forward pass gives the doubly-stochastic reparameterised ELBO of
    Cutajar et al. (2017).

    At ``depth=1`` this reduces to a single VSSGP layer mapping
    ``in_features -> out_features`` via the RFF basis (same model class
    as `VariationalFourierFeatures` followed by a
    `DenseReparameterization` head). Stacking adds
    non-stationarity at the cost of a non-Gaussian aggregate likelihood
    — the layer-wise marginalisation that makes single-layer SSGP
    closed-form is no longer available, hence the variational
    treatment.

    Attributes:
        in_features: Input dimension $D_{\mathrm{in}}$.
        hidden_features: Inter-layer dimension $D_h$ (constant
            across hidden layers).
        out_features: Output dimension $D_{\mathrm{out}}$.
        n_features: Per-layer Fourier-feature pair count $M$ (so
            each layer's hidden state is $2M$-dim before
            projection).
        depth: Total number of stacked SSGP layers $L$. Must
            be $\ge 1$.
        init_lengthscale: Prior location for each layer's lengthscale.
        prior_std: Standard deviation of the per-layer projection
            prior $\mathcal{N}(0, \sigma_W^2)$.
        pyrox_name: Explicit scope name for NumPyro site registration.

    Examples:
        >>> import jax.random as jr, jax.numpy as jnp
        >>> from numpyro import handlers
        >>> net = DeepVSSGP.init(in_features=2, hidden_features=4,
        ...                       out_features=1, depth=3, n_features=16)
        >>> with handlers.seed(rng_seed=0):
        ...     y = net(jnp.zeros((8, 2)))
        >>> y.shape
        (8, 1)
    """

    core: DeepVSSGPCore
    init_lengthscale: float = 1.0
    prior_std: float = 1.0
    pyrox_name: str | None = None

    # Structural accessors so wrapper exposes the same field surface as before.
    @property
    def in_features(self) -> int:
        return self.core.in_features

    @property
    def hidden_features(self) -> int:
        return self.core.hidden_features

    @property
    def out_features(self) -> int:
        return self.core.out_features

    @property
    def n_features(self) -> int:
        return self.core.n_features

    @property
    def depth(self) -> int:
        return self.core.depth

    @classmethod
    def init(
        cls,
        in_features: int,
        hidden_features: int,
        out_features: int,
        *,
        depth: int,
        n_features: int = 64,
        lengthscale: float = 1.0,
        prior_std: float = 1.0,
        pyrox_name: str | None = None,
    ) -> DeepVSSGP:
        """Construct a `DeepVSSGP`.

        Args:
            in_features: Input dimension. Must be $\\ge 1$.
            hidden_features: Hidden dimension. Must be $\\ge 1$.
            out_features: Output dimension. Must be $\\ge 1$.
            depth: Total stacked SSGP layers (including readout).
                Must be $\\ge 1$.
            n_features: Per-layer Fourier-feature pair count. Must be
                $\\ge 1$.
            lengthscale: Prior location for each layer's lengthscale.
                Must be $> 0$.
            prior_std: Per-layer projection prior standard deviation.
                Must be $> 0$.
            pyrox_name: Optional explicit scope name for NumPyro site
                registration.

        Returns:
            Initialised `DeepVSSGP`.

        Raises:
            ValueError: If ``depth``, any feature dimension, or
                ``n_features`` is $< 1$, or if ``lengthscale`` /
                ``prior_std`` is $\\le 0$.
        """
        if depth < 1:
            raise ValueError(f"depth must be >= 1, got {depth}.")
        if n_features < 1:
            raise ValueError(f"n_features must be >= 1, got {n_features}.")
        if lengthscale <= 0:
            raise ValueError(f"lengthscale must be > 0, got {lengthscale}.")
        if prior_std <= 0:
            raise ValueError(f"prior_std must be > 0, got {prior_std}.")
        for name, dim in (
            ("in_features", in_features),
            ("hidden_features", hidden_features),
            ("out_features", out_features),
        ):
            if dim < 1:
                raise ValueError(f"{name} must be >= 1, got {dim}.")
        # The Bayesian wrapper samples every per-layer array on each
        # forward call, so the core's stored arrays are only used as a
        # structural skeleton.  We build it with a dummy PRNG key — its
        # weight values are irrelevant because eqx.tree_at swaps them
        # for sampled values below.
        core = DeepVSSGPCore.init(
            in_features,
            hidden_features,
            out_features,
            depth=depth,
            key=jax.random.PRNGKey(0),
            n_features=n_features,
            lengthscale=lengthscale,
            prior_std=prior_std,
        )
        return cls(
            core=core,
            init_lengthscale=lengthscale,
            prior_std=prior_std,
            pyrox_name=pyrox_name,
        )

    @pyrox_method
    def __call__(self, x: Float[Array, "*batch D_in"]) -> Float[Array, "*batch D_out"]:
        sampled_W_freqs: list[Float[Array, "d_in n_features"]] = []
        sampled_W_projs: list[Float[Array, "d_rff d_out"]] = []
        sampled_lengthscales: list[Float[Array, ""]] = []
        for layer_idx in range(self.core.depth):
            in_dim = (
                self.core.in_features if layer_idx == 0 else self.core.hidden_features
            )
            out_dim = (
                self.core.out_features
                if layer_idx == self.core.depth - 1
                else self.core.hidden_features
            )
            W_freq = self.pyrox_sample(
                f"layer_{layer_idx}.W_freq",
                dist.Normal(0.0, 1.0)
                .expand([in_dim, self.core.n_features])
                .to_event(2),
            )
            ls = self.pyrox_sample(
                f"layer_{layer_idx}.lengthscale",
                dist.LogNormal(jnp.log(jnp.asarray(self.init_lengthscale)), 1.0),
            )
            W_proj = self.pyrox_sample(
                f"layer_{layer_idx}.W_proj",
                dist.Normal(0.0, self.prior_std)
                .expand([2 * self.core.n_features, out_dim])
                .to_event(2),
            )
            sampled_W_freqs.append(W_freq)
            sampled_W_projs.append(W_proj)
            sampled_lengthscales.append(ls)

        sampled_core = eqx.tree_at(
            lambda c: (c.W_freqs, c.W_projs, c.lengthscales),
            self.core,
            (
                sampled_W_freqs,
                sampled_W_projs,
                jnp.stack(sampled_lengthscales),
            ),
        )
        # Single-example geonnax core over arbitrary leading batch dims.
        return vmap_over_flat_batch(sampled_core, x)

init(in_features: int, hidden_features: int, out_features: int, *, depth: int, n_features: int = 64, lengthscale: float = 1.0, prior_std: float = 1.0, pyrox_name: str | None = None) -> DeepVSSGP classmethod

Construct a DeepVSSGP.

Parameters:

Name Type Description Default
in_features int

Input dimension. Must be \(\ge 1\).

required
hidden_features int

Hidden dimension. Must be \(\ge 1\).

required
out_features int

Output dimension. Must be \(\ge 1\).

required
depth int

Total stacked SSGP layers (including readout). Must be \(\ge 1\).

required
n_features int

Per-layer Fourier-feature pair count. Must be \(\ge 1\).

64
lengthscale float

Prior location for each layer's lengthscale. Must be \(> 0\).

1.0
prior_std float

Per-layer projection prior standard deviation. Must be \(> 0\).

1.0
pyrox_name str | None

Optional explicit scope name for NumPyro site registration.

None

Returns:

Type Description
DeepVSSGP

Initialised DeepVSSGP.

Raises:

Type Description
ValueError

If depth, any feature dimension, or n_features is \(< 1\), or if lengthscale / prior_std is \(\le 0\).

Source code in packages/pyrox-nn/src/pyrox_nn/_vssgp.py
@classmethod
def init(
    cls,
    in_features: int,
    hidden_features: int,
    out_features: int,
    *,
    depth: int,
    n_features: int = 64,
    lengthscale: float = 1.0,
    prior_std: float = 1.0,
    pyrox_name: str | None = None,
) -> DeepVSSGP:
    """Construct a `DeepVSSGP`.

    Args:
        in_features: Input dimension. Must be $\\ge 1$.
        hidden_features: Hidden dimension. Must be $\\ge 1$.
        out_features: Output dimension. Must be $\\ge 1$.
        depth: Total stacked SSGP layers (including readout).
            Must be $\\ge 1$.
        n_features: Per-layer Fourier-feature pair count. Must be
            $\\ge 1$.
        lengthscale: Prior location for each layer's lengthscale.
            Must be $> 0$.
        prior_std: Per-layer projection prior standard deviation.
            Must be $> 0$.
        pyrox_name: Optional explicit scope name for NumPyro site
            registration.

    Returns:
        Initialised `DeepVSSGP`.

    Raises:
        ValueError: If ``depth``, any feature dimension, or
            ``n_features`` is $< 1$, or if ``lengthscale`` /
            ``prior_std`` is $\\le 0$.
    """
    if depth < 1:
        raise ValueError(f"depth must be >= 1, got {depth}.")
    if n_features < 1:
        raise ValueError(f"n_features must be >= 1, got {n_features}.")
    if lengthscale <= 0:
        raise ValueError(f"lengthscale must be > 0, got {lengthscale}.")
    if prior_std <= 0:
        raise ValueError(f"prior_std must be > 0, got {prior_std}.")
    for name, dim in (
        ("in_features", in_features),
        ("hidden_features", hidden_features),
        ("out_features", out_features),
    ):
        if dim < 1:
            raise ValueError(f"{name} must be >= 1, got {dim}.")
    # The Bayesian wrapper samples every per-layer array on each
    # forward call, so the core's stored arrays are only used as a
    # structural skeleton.  We build it with a dummy PRNG key — its
    # weight values are irrelevant because eqx.tree_at swaps them
    # for sampled values below.
    core = DeepVSSGPCore.init(
        in_features,
        hidden_features,
        out_features,
        depth=depth,
        key=jax.random.PRNGKey(0),
        n_features=n_features,
        lengthscale=lengthscale,
        prior_std=prior_std,
    )
    return cls(
        core=core,
        init_lengthscale=lengthscale,
        prior_std=prior_std,
        pyrox_name=pyrox_name,
    )

Ensembles — BatchEnsemble / rank-1

Efficient deep ensembles that share one weight matrix and learn per-member rank-1 perturbations (Wen et al., 2020; Dusenberry et al., 2020).

DenseRank1

Bases: PyroxModule

Rank-1 ensemble dense layer.

Implements the BatchEnsemble (Wen et al., 2020) / rank-1 BNN (Dusenberry et al., 2020) parameterization: a single shared kernel \(W \in \mathbb{R}^{D_\mathrm{in} \times D_\mathrm{out}}\) and per-member rank-1 multiplicative perturbations \(s_i \in \mathbb{R}^{D_\mathrm{in}}\), \(r_i \in \mathbb{R}^{D_\mathrm{out}}\) for \(i = 1, \ldots, M\). The per-member effective weight is

\[ W_i = (s_i \otimes r_i) \circ W, \]

and the efficient forward pass avoids materialising \(W_i\):

\[ y_i = \bigl((x \circ s_i)\, W\bigr) \circ r_i + b_i. \]

Two modes via the bayesian flag:

  • bayesian=False (default) — BatchEnsemble. \(r, s, W, b\) are all deterministic pyrox_param sites and per-member diversity comes purely from the random initialisation of \(r_i, s_i\). Use this for ensemble training under a single shared SGD trajectory.
  • bayesian=True — rank-1 BNN. \(r, s\) are pyrox_sample sites with Normal priors centered at the per-member init values; \(W, b\) remain deterministic. Plug into NumPyro's SVI machinery (an AutoNormal guide on r, s recovers Dusenberry et al., 2020).
Plate semantics

Identical to other pyrox Bayesian dense layers — call this layer outside numpyro.plate("data", ..., subsample_size=...) and only plate the observation likelihood. The model log density picks up \(\log p(r_i)\) and \(\log p(s_i)\) once per layer (not once per example) under the canonical pattern.

Attributes:

Name Type Description
in_features int

Input dimension \(D_\mathrm{in}\).

out_features int

Output dimension \(D_\mathrm{out}\).

ensemble_size int

Number of ensemble members \(M\).

bias bool

Whether to include a per-member bias.

bayesian bool

If True, place Normal priors on \(r, s\).

prior_scale float

Std of the Bayesian priors on \(r, s\). Only used when bayesian=True.

W_init float

Shared kernel init, shape (D_in, D_out).

r_init float

Per-member output-side init, shape (M, D_out).

s_init float

Per-member input-side init, shape (M, D_in).

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Examples:

>>> import jax.random as jr
>>> import jax.numpy as jnp
>>> from numpyro import handlers
>>> layer = DenseRank1.init(
...     jr.PRNGKey(0),
...     in_features=4,
...     out_features=2,
...     ensemble_size=3,
... )
>>> x = jnp.ones((5, 4))
>>> with handlers.seed(rng_seed=0):
...     y = layer(x)
>>> y.shape
(3, 5, 2)
References

Wen, Y., Tran, D., & Ba, J. (2020). BatchEnsemble: An Alternative Approach to Efficient Ensemble and Lifelong Learning. ICLR.

Dusenberry, M. W., et al. (2020). Efficient and Scalable Bayesian Neural Nets with Rank-1 Factors. ICML.

Source code in packages/pyrox-nn/src/pyrox_nn/_ensemble.py
class DenseRank1(PyroxModule):
    r"""Rank-1 ensemble dense layer.

    Implements the BatchEnsemble (Wen et al., 2020) / rank-1 BNN
    (Dusenberry et al., 2020) parameterization: a single shared kernel
    $W \in \mathbb{R}^{D_\mathrm{in} \times D_\mathrm{out}}$ and
    per-member rank-1 multiplicative perturbations
    $s_i \in \mathbb{R}^{D_\mathrm{in}}$,
    $r_i \in \mathbb{R}^{D_\mathrm{out}}$ for
    $i = 1, \ldots, M$. The per-member effective weight is

    $$
    W_i = (s_i \otimes r_i) \circ W,
    $$

    and the efficient forward pass avoids materialising $W_i$:

    $$
    y_i = \bigl((x \circ s_i)\, W\bigr) \circ r_i + b_i.
    $$

    Two modes via the ``bayesian`` flag:

    * ``bayesian=False`` (default) — BatchEnsemble. $r, s, W, b$
      are all deterministic ``pyrox_param`` sites and per-member
      diversity comes purely from the random initialisation of
      $r_i, s_i$. Use this for ensemble training under a single
      shared SGD trajectory.
    * ``bayesian=True`` — rank-1 BNN. $r, s$ are
      ``pyrox_sample`` sites with Normal priors centered at the
      per-member init values; $W, b$ remain deterministic. Plug
      into NumPyro's SVI machinery (an ``AutoNormal`` guide on
      ``r, s`` recovers Dusenberry et al., 2020).

    Plate semantics:
        Identical to other pyrox Bayesian dense layers — call this
        layer **outside** ``numpyro.plate("data", ..., subsample_size=...)``
        and only plate the observation likelihood. The model log
        density picks up $\log p(r_i)$ and $\log p(s_i)$
        once per layer (not once per example) under the canonical
        pattern.

    Attributes:
        in_features: Input dimension $D_\mathrm{in}$.
        out_features: Output dimension $D_\mathrm{out}$.
        ensemble_size: Number of ensemble members $M$.
        bias: Whether to include a per-member bias.
        bayesian: If ``True``, place Normal priors on $r, s$.
        prior_scale: Std of the Bayesian priors on $r, s$. Only
            used when ``bayesian=True``.
        W_init: Shared kernel init, shape ``(D_in, D_out)``.
        r_init: Per-member output-side init, shape ``(M, D_out)``.
        s_init: Per-member input-side init, shape ``(M, D_in)``.
        pyrox_name: Explicit scope name for NumPyro site registration.

    Examples:
        >>> import jax.random as jr
        >>> import jax.numpy as jnp
        >>> from numpyro import handlers
        >>> layer = DenseRank1.init(
        ...     jr.PRNGKey(0),
        ...     in_features=4,
        ...     out_features=2,
        ...     ensemble_size=3,
        ... )
        >>> x = jnp.ones((5, 4))
        >>> with handlers.seed(rng_seed=0):
        ...     y = layer(x)
        >>> y.shape
        (3, 5, 2)

    References:
        Wen, Y., Tran, D., & Ba, J. (2020). *BatchEnsemble: An
        Alternative Approach to Efficient Ensemble and Lifelong
        Learning.* ICLR.

        Dusenberry, M. W., et al. (2020). *Efficient and Scalable
        Bayesian Neural Nets with Rank-1 Factors.* ICML.
    """

    core: geonnax.DenseRank1
    bayesian: bool = eqx.field(static=True, default=False)
    prior_scale: float = 0.5
    pyrox_name: str | None = None

    @classmethod
    def init(
        cls,
        key: PRNGKeyArray,
        in_features: int,
        out_features: int,
        ensemble_size: int,
        *,
        bias: bool = True,
        bayesian: bool = False,
        init_scale: float = 0.5,
        prior_scale: float = 0.5,
        pyrox_name: str | None = None,
    ) -> DenseRank1:
        """Construct a layer with random per-member init vectors."""
        # `prior_scale` is only consulted in Bayesian mode — don't reject
        # configs that pass through a sentinel default in deterministic mode.
        if bayesian and prior_scale <= 0:
            raise ValueError(
                f"prior_scale must be > 0 when bayesian=True; got {prior_scale}."
            )
        # geonnax handles positive-dim / init_scale validation.
        core = geonnax.DenseRank1.init(
            key,
            in_features=in_features,
            out_features=out_features,
            ensemble_size=ensemble_size,
            bias=bias,
            init_scale=init_scale,
        )
        return cls(
            core=core,
            bayesian=bayesian,
            prior_scale=prior_scale,
            pyrox_name=pyrox_name,
        )

    # Convenience read-only attribute access.
    @property
    def in_features(self) -> int:
        return self.core.in_features

    @property
    def out_features(self) -> int:
        return self.core.out_features

    @property
    def ensemble_size(self) -> int:
        return self.core.ensemble_size

    @property
    def bias(self) -> bool:
        return self.core.bias

    @pyrox_method
    def __call__(
        self, x: Float[Array, "*batch D_in"]
    ) -> Float[Array, "M *batch D_out"]:
        W = self.pyrox_param("W", self.core.W)

        if self.bayesian:
            r = self.pyrox_sample(
                "r",
                dist.Normal(self.core.r, self.prior_scale).to_event(2),
            )
            s = self.pyrox_sample(
                "s",
                dist.Normal(self.core.s, self.prior_scale).to_event(2),
            )
        else:
            r = self.pyrox_param("r", self.core.r)
            s = self.pyrox_param("s", self.core.s)

        b = self.pyrox_param("b", self.core.b) if self.bias else self.core.b

        new_core = eqx.tree_at(
            lambda c: (c.W, c.r, c.s, c.b),
            self.core,
            (W, r, s, b),
        )
        return _vmap_collapsed(new_core, x)

init(key: PRNGKeyArray, in_features: int, out_features: int, ensemble_size: int, *, bias: bool = True, bayesian: bool = False, init_scale: float = 0.5, prior_scale: float = 0.5, pyrox_name: str | None = None) -> DenseRank1 classmethod

Construct a layer with random per-member init vectors.

Source code in packages/pyrox-nn/src/pyrox_nn/_ensemble.py
@classmethod
def init(
    cls,
    key: PRNGKeyArray,
    in_features: int,
    out_features: int,
    ensemble_size: int,
    *,
    bias: bool = True,
    bayesian: bool = False,
    init_scale: float = 0.5,
    prior_scale: float = 0.5,
    pyrox_name: str | None = None,
) -> DenseRank1:
    """Construct a layer with random per-member init vectors."""
    # `prior_scale` is only consulted in Bayesian mode — don't reject
    # configs that pass through a sentinel default in deterministic mode.
    if bayesian and prior_scale <= 0:
        raise ValueError(
            f"prior_scale must be > 0 when bayesian=True; got {prior_scale}."
        )
    # geonnax handles positive-dim / init_scale validation.
    core = geonnax.DenseRank1.init(
        key,
        in_features=in_features,
        out_features=out_features,
        ensemble_size=ensemble_size,
        bias=bias,
        init_scale=init_scale,
    )
    return cls(
        core=core,
        bayesian=bayesian,
        prior_scale=prior_scale,
        pyrox_name=pyrox_name,
    )

LayerNormEnsemble

Bases: PyroxModule

Per-ensemble-member LayerNorm.

Drop-in replacement for LayerNorm inside BatchEnsemble / Rank1 architectures. Computes the standard LayerNorm normalisation over the trailing feature dimension and applies a per-member affine transform — each ensemble member \(i \in \{1, \ldots, M\}\) gets its own learnable scale \(\gamma_i \in \mathbb{R}^D\) and bias \(\beta_i \in \mathbb{R}^D\):

\[ \hat{x}_i = \frac{x_i - \mu(x_i)}{\sqrt{\sigma^2(x_i) + \epsilon}}, \qquad y_i = \gamma_i \odot \hat{x}_i + \beta_i, \]

where \(\mu\) and \(\sigma^2\) are the empirical mean and variance over the trailing feature axis (computed independently for each member-batch slice). Without per-member scale/bias, sharing a single LayerNorm across the ensemble would couple all members and erase the diversity introduced by DenseRank1 or any other BatchEnsemble layer upstream.

Input is expected to carry a leading ensemble axis of size ensemble_size and a trailing feature axis of size feature_dim. Any number of intermediate batch / time axes are supported and pass through unchanged.

Attributes:

Name Type Description
ensemble_size int

Number of ensemble members \(M\).

feature_dim int

Trailing feature dimension \(D\) over which the normalisation is computed.

eps float

Small positive constant added to the variance for numerical stability.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Examples:

>>> import jax.numpy as jnp
>>> from numpyro import handlers
>>> ln = LayerNormEnsemble(
...     ensemble_size=3, feature_dim=4, pyrox_name="ln"
... )
>>> x = jnp.ones((3, 5, 4))  # (M, batch, D)
>>> with handlers.seed(rng_seed=0):
...     y = ln(x)
>>> y.shape
(3, 5, 4)
Source code in packages/pyrox-nn/src/pyrox_nn/_ensemble.py
class LayerNormEnsemble(PyroxModule):
    r"""Per-ensemble-member LayerNorm.

    Drop-in replacement for ``LayerNorm`` inside BatchEnsemble / Rank1
    architectures. Computes the standard LayerNorm normalisation over
    the trailing feature dimension and applies a *per-member* affine
    transform — each ensemble member $i \in \{1, \ldots, M\}$
    gets its own learnable scale $\gamma_i \in \mathbb{R}^D$
    and bias $\beta_i \in \mathbb{R}^D$:

    $$
    \hat{x}_i = \frac{x_i - \mu(x_i)}{\sqrt{\sigma^2(x_i) + \epsilon}},
    \qquad
    y_i = \gamma_i \odot \hat{x}_i + \beta_i,
    $$

    where $\mu$ and $\sigma^2$ are the empirical mean and
    variance over the trailing feature axis (computed independently
    for each member-batch slice). Without per-member scale/bias,
    sharing a single LayerNorm across the ensemble would couple all
    members and erase the diversity introduced by `DenseRank1`
    or any other BatchEnsemble layer upstream.

    Input is expected to carry a leading ensemble axis of size
    ``ensemble_size`` and a trailing feature axis of size
    ``feature_dim``. Any number of intermediate batch / time axes
    are supported and pass through unchanged.

    Attributes:
        ensemble_size: Number of ensemble members $M$.
        feature_dim: Trailing feature dimension $D$ over which
            the normalisation is computed.
        eps: Small positive constant added to the variance for
            numerical stability.
        pyrox_name: Explicit scope name for NumPyro site registration.

    Examples:
        >>> import jax.numpy as jnp
        >>> from numpyro import handlers
        >>> ln = LayerNormEnsemble(
        ...     ensemble_size=3, feature_dim=4, pyrox_name="ln"
        ... )
        >>> x = jnp.ones((3, 5, 4))  # (M, batch, D)
        >>> with handlers.seed(rng_seed=0):
        ...     y = ln(x)
        >>> y.shape
        (3, 5, 4)
    """

    core: geonnax.LayerNormEnsemble
    pyrox_name: str | None = eqx.field(static=True, default=None)

    def __init__(
        self,
        ensemble_size: int,
        feature_dim: int,
        eps: float = 1e-5,
        pyrox_name: str | None = None,
        *,
        core: geonnax.LayerNormEnsemble | None = None,
    ) -> None:
        # Keep the prior call signature (positional dims, kw eps / pyrox_name)
        # so existing callers don't need to know about the geonnax core
        # construction.
        if core is None:
            # geonnax validates positive ensemble_size / feature_dim / eps.
            core = geonnax.LayerNormEnsemble.init(
                ensemble_size=ensemble_size,
                feature_dim=feature_dim,
                eps=eps,
            )
        self.core = core
        self.pyrox_name = pyrox_name

    @property
    def ensemble_size(self) -> int:
        return self.core.ensemble_size

    @property
    def feature_dim(self) -> int:
        return self.core.feature_dim

    @property
    def eps(self) -> float:
        return self.core.eps

    @pyrox_method
    def __call__(self, x: Float[Array, "M *batch D"]) -> Float[Array, "M *batch D"]:
        if x.ndim < 2:
            raise ValueError(
                f"x must have at least 2 dims (M and D); got shape {x.shape}."
            )
        if x.shape[0] != self.ensemble_size:
            raise ValueError(
                f"x.shape[0] = {x.shape[0]} does not match "
                f"ensemble_size = {self.ensemble_size}."
            )
        if x.shape[-1] != self.feature_dim:
            raise ValueError(
                f"x.shape[-1] = {x.shape[-1]} does not match "
                f"feature_dim = {self.feature_dim}."
            )

        scales = self.pyrox_param("scales", self.core.scales)
        biases = self.pyrox_param("biases", self.core.biases)

        new_core = eqx.tree_at(
            lambda c: (c.scales, c.biases),
            self.core,
            (scales, biases),
        )

        # Core takes (M, D) → (M, D); collapse arbitrary intermediate batch
        # dims to a single leading axis, vmap over it, then restore. The
        # ensemble axis stays at position 0 because the core is per-member.
        batch_shape = x.shape[1:-1]
        flat = einx.id("m b... d -> m (b...) d", x)  # (M, B, D)
        # vmap over the B axis (position 1 on both input and output).
        out = jax.vmap(new_core, in_axes=1, out_axes=1)(flat)  # (M, B, D)
        return einx.id("m (b...) d -> m b... d", out, b=batch_shape)

MultiHeadAttentionBE

Bases: PyroxModule

Multi-head attention with BatchEnsemble rank-1 projections.

Standard scaled-dot-product multi-head attention where each of the four linear projections — query, key, value, and output — uses a BatchEnsemble parameterisation: a shared full-rank kernel plus per-ensemble-member rank-1 multiplicative perturbations. So for member \(i \in \{1, \ldots, M\}\) and projection \(P \in \{Q, K, V, O\}\),

\[ W_i^{(P)} = (s_i^{(P)} \otimes r_i^{(P)}) \circ W^{(P)}, \]

and the attention itself is the usual

\[ \mathrm{Attn}(Q, K, V) = \mathrm{softmax}\! \Bigl(\frac{Q K^\top}{\sqrt{d_k}}\Bigr) V. \]

The forward consumes un-ensembled inputs (query, key, value of shape (T, D) / (S, D)), adds the ensemble axis when projecting to Q, K, V, runs per-member attention in parallel, and returns the per-member output of shape (M, T, D). Equivalent to running M independent attention heads with rank-1 weight perturbations and stacking their outputs.

Plate semantics

Same convention as DenseRank1 and the rest of the pyrox_nn ensemble / Bayesian dense family — call this layer outside numpyro.plate("data", ..., subsample_size=...) and only plate the observation likelihood. All four projections register their parameters as pyrox_param sites; nothing about the layer is data-dependent so plate-scaling does not come into play unless the user puts the call inside a subsampled plate.

Attributes:

Name Type Description
embed_dim int

Total feature dimension \(D\) of query / key / value (must be divisible by num_heads).

num_heads int

Number of attention heads \(H\). Each head sees embed_dim // num_heads features.

ensemble_size int

Number of ensemble members \(M\).

bias bool

Whether each of the four projections includes a per-member bias. When False, no bias param sites are registered for any of Q / K / V / O.

q_init / k_init / v_init / o_init

Per-projection BatchEnsemble init arrays. Build via init.

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Examples:

>>> import jax.random as jr
>>> import jax.numpy as jnp
>>> from numpyro import handlers
>>> mha = MultiHeadAttentionBE.init(
...     jr.PRNGKey(0),
...     embed_dim=8, num_heads=2, ensemble_size=3,
... )
>>> x = jnp.ones((5, 8))
>>> with handlers.seed(rng_seed=0):
...     y = mha(x, x, x)         # self-attention
>>> y.shape
(3, 5, 8)
Source code in packages/pyrox-nn/src/pyrox_nn/_ensemble.py
class MultiHeadAttentionBE(PyroxModule):
    r"""Multi-head attention with BatchEnsemble rank-1 projections.

    Standard scaled-dot-product multi-head attention where each of the
    four linear projections — query, key, value, and output — uses a
    BatchEnsemble parameterisation: a shared full-rank kernel plus
    per-ensemble-member rank-1 multiplicative perturbations. So for
    member $i \in \{1, \ldots, M\}$ and projection
    $P \in \{Q, K, V, O\}$,

    $$
    W_i^{(P)} = (s_i^{(P)} \otimes r_i^{(P)}) \circ W^{(P)},
    $$

    and the attention itself is the usual

    $$
    \mathrm{Attn}(Q, K, V) = \mathrm{softmax}\!
        \Bigl(\frac{Q K^\top}{\sqrt{d_k}}\Bigr) V.
    $$

    The forward consumes un-ensembled inputs (``query``, ``key``,
    ``value`` of shape ``(T, D)`` / ``(S, D)``), adds the ensemble
    axis when projecting to ``Q``, ``K``, ``V``, runs per-member
    attention in parallel, and returns the per-member output of
    shape ``(M, T, D)``. Equivalent to running ``M`` independent
    attention heads with rank-1 weight perturbations and stacking
    their outputs.

    Plate semantics:
        Same convention as `DenseRank1` and the rest of the
        ``pyrox_nn`` ensemble / Bayesian dense family — call this
        layer **outside** ``numpyro.plate("data", ..., subsample_size=...)``
        and only plate the observation likelihood. All four projections
        register their parameters as ``pyrox_param`` sites; nothing
        about the layer is data-dependent so plate-scaling does not
        come into play unless the user puts the call inside a
        subsampled plate.

    Attributes:
        embed_dim: Total feature dimension $D$ of query / key /
            value (must be divisible by ``num_heads``).
        num_heads: Number of attention heads $H$. Each head sees
            ``embed_dim // num_heads`` features.
        ensemble_size: Number of ensemble members $M$.
        bias: Whether each of the four projections includes a
            per-member bias. When ``False``, no bias param sites are
            registered for any of Q / K / V / O.
        q_init / k_init / v_init / o_init: Per-projection
            BatchEnsemble init arrays. Build via `init`.
        pyrox_name: Explicit scope name for NumPyro site registration.

    Examples:
        >>> import jax.random as jr
        >>> import jax.numpy as jnp
        >>> from numpyro import handlers
        >>> mha = MultiHeadAttentionBE.init(
        ...     jr.PRNGKey(0),
        ...     embed_dim=8, num_heads=2, ensemble_size=3,
        ... )
        >>> x = jnp.ones((5, 8))
        >>> with handlers.seed(rng_seed=0):
        ...     y = mha(x, x, x)         # self-attention
        >>> y.shape
        (3, 5, 8)
    """

    core: geonnax.MultiHeadAttentionBE
    pyrox_name: str | None = eqx.field(static=True, default=None)

    @classmethod
    def init(
        cls,
        key: PRNGKeyArray,
        embed_dim: int,
        num_heads: int,
        ensemble_size: int,
        *,
        bias: bool = True,
        init_scale: float = 0.5,
        pyrox_name: str | None = None,
    ) -> MultiHeadAttentionBE:
        """Construct an MHA-BE layer with random Q/K/V/O projection inits."""
        # geonnax validates positive dims, divisibility, init_scale.
        core = geonnax.MultiHeadAttentionBE.init(
            key,
            embed_dim=embed_dim,
            num_heads=num_heads,
            ensemble_size=ensemble_size,
            bias=bias,
            init_scale=init_scale,
        )
        return cls(core=core, pyrox_name=pyrox_name)

    @property
    def embed_dim(self) -> int:
        return self.core.embed_dim

    @property
    def num_heads(self) -> int:
        return self.core.num_heads

    @property
    def ensemble_size(self) -> int:
        return self.core.ensemble_size

    @property
    def bias(self) -> bool:
        return self.core.bias

    @property
    def head_dim(self) -> int:
        return self.core.head_dim

    # Maintain the pre-refactor attribute names for any external accessors.
    @property
    def q_init(self) -> Rank1ProjInit:
        return self.core.q_proj

    @property
    def k_init(self) -> Rank1ProjInit:
        return self.core.k_proj

    @property
    def v_init(self) -> Rank1ProjInit:
        return self.core.v_proj

    @property
    def o_init(self) -> Rank1ProjInit:
        return self.core.o_proj

    def _register_proj(self, name: str, init: Rank1ProjInit) -> Rank1ProjInit:
        """Register a projection's arrays as ``pyrox_param`` sites.

        Skips the ``b`` site when ``self.bias`` is ``False`` so disabled
        biases don't leak unused params into the SVI parameter store.
        """
        b = (
            self.pyrox_param(f"{name}_b", init.b)
            if self.bias
            else jnp.zeros_like(init.b)
        )
        return Rank1ProjInit(
            W=self.pyrox_param(f"{name}_W", init.W),
            r=self.pyrox_param(f"{name}_r", init.r),
            s=self.pyrox_param(f"{name}_s", init.s),
            b=b,
        )

    @pyrox_method
    def __call__(
        self,
        query: Float[Array, "T D"],
        key: Float[Array, "S D"],
        value: Float[Array, "S D"],
    ) -> Float[Array, "M T D"]:
        # Register the four projections.
        q_proj = self._register_proj("q", self.core.q_proj)
        k_proj = self._register_proj("k", self.core.k_proj)
        v_proj = self._register_proj("v", self.core.v_proj)
        o_proj = self._register_proj("o", self.core.o_proj)

        # Swap all four projections into the core; geonnax's __call__
        # validates query/key/value shapes and runs the per-member
        # attention. The core forward matches the pyrox call signature
        # (un-ensembled `(T, D)` / `(S, D)` in → `(M, T, D)` out), so no
        # vmap is needed here — the ensemble axis is intrinsic to the
        # core, not a data batch.
        new_core = eqx.tree_at(
            lambda c: (c.q_proj, c.k_proj, c.v_proj, c.o_proj),
            self.core,
            (q_proj, k_proj, v_proj, o_proj),
        )
        return new_core(query, key, value)

init(key: PRNGKeyArray, embed_dim: int, num_heads: int, ensemble_size: int, *, bias: bool = True, init_scale: float = 0.5, pyrox_name: str | None = None) -> MultiHeadAttentionBE classmethod

Construct an MHA-BE layer with random Q/K/V/O projection inits.

Source code in packages/pyrox-nn/src/pyrox_nn/_ensemble.py
@classmethod
def init(
    cls,
    key: PRNGKeyArray,
    embed_dim: int,
    num_heads: int,
    ensemble_size: int,
    *,
    bias: bool = True,
    init_scale: float = 0.5,
    pyrox_name: str | None = None,
) -> MultiHeadAttentionBE:
    """Construct an MHA-BE layer with random Q/K/V/O projection inits."""
    # geonnax validates positive dims, divisibility, init_scale.
    core = geonnax.MultiHeadAttentionBE.init(
        key,
        embed_dim=embed_dim,
        num_heads=num_heads,
        ensemble_size=ensemble_size,
        bias=bias,
        init_scale=init_scale,
    )
    return cls(core=core, pyrox_name=pyrox_name)

Heteroscedastic output heads

Monte-Carlo sigmoid / softmax output layers with factor-analysis noise (Collier et al., 2021) for input-dependent label noise.

MCSigmoidDenseFA

Bases: _HeteroscedasticBase

Heteroscedastic multi-label output layer (FA noise + sigmoid).

Identical low-rank-plus-diagonal logit-noise model as MCSoftmaxDenseFA, but the per-class outputs are independent Bernoullis — final probabilities are the MC average of element-wise sigmoids, not a softmax. Use this for multi-label classification or independent binary heads.

\[ \hat{p}(y_k = 1 \mid x) \approx \frac{1}{S}\sum_{s=1}^{S} \sigma\!\bigl(\eta(x) + \epsilon_s\bigr)_k. \]

See MCSoftmaxDenseFA for the noise model, plate semantics, init API, and references.

Source code in packages/pyrox-nn/src/pyrox_nn/_heteroscedastic.py
class MCSigmoidDenseFA(_HeteroscedasticBase):
    r"""Heteroscedastic multi-label output layer (FA noise + sigmoid).

    Identical low-rank-plus-diagonal logit-noise model as
    `MCSoftmaxDenseFA`, but the per-class outputs are
    independent Bernoullis — final probabilities are the MC average of
    *element-wise* sigmoids, not a softmax. Use this for multi-label
    classification or independent binary heads.

    $$
    \hat{p}(y_k = 1 \mid x) \approx
    \frac{1}{S}\sum_{s=1}^{S}
    \sigma\!\bigl(\eta(x) + \epsilon_s\bigr)_k.
    $$

    See `MCSoftmaxDenseFA` for the noise model, plate semantics,
    init API, and references.
    """

    @classmethod
    def _core_cls(cls) -> type[geonnax.HeteroscedasticHead]:
        return geonnax.MCSigmoidDenseFA

    @pyrox_method
    def __call__(self, x: Float[Array, "N D_in"]) -> Float[Array, "N C"]:
        new_core = self._swap_core()
        logits = _batched_logits(new_core, x)
        return jnp.mean(jax.nn.sigmoid(logits), axis=0)

MCSoftmaxDenseFA

Bases: _HeteroscedasticBase

Heteroscedastic multi-class output layer (FA noise + softmax).

Implements Collier et al. (2021): the logit covariance is input-dependent low-rank-plus-diagonal,

\[ \eta(x) = W_\mu x + b_\mu + \epsilon, \qquad \Sigma(x) = V(x) V(x)^\top + \operatorname{diag}\!\bigl(\sigma^2(x)\bigr), \;\; \epsilon \sim \mathcal{N}(0, \Sigma(x)), \]

where \(V(x) = \mathrm{reshape}(W_V x + b_V, [C, r])\) and \(\sigma(x) = \exp(W_\sigma x + b_\sigma)\). Output is the Monte Carlo average of softmaxed perturbed logits

\[ \hat{p}(y = k \mid x) \approx \frac{1}{S}\sum_{s=1}^{S} \mathrm{softmax}_k\!\bigl(\eta(x) + \epsilon_s\bigr). \]

All linear factors are deterministic pyrox_param sites — the layer is heteroscedastic but not Bayesian over its weights. Use it as a drop-in head for classification when label noise is known to be input-dependent (label disagreement, fine-grained categories).

Plate semantics

Same as other pyrox_nn Bayesian dense layers — call outside numpyro.plate("data", ..., subsample_size=...) so the parameter sites are unscaled. The MC noise is drawn from numpyro.prng_key().

Attributes:

Name Type Description
in_features int

Input dimension \(D_\mathrm{in}\).

num_classes int

Number of classes \(C\).

rank int

Rank \(r\) of the low-rank factor \(V(x)\).

num_mc_samples int

Number of MC softmax samples \(S\) per forward call.

diag_init_bias float

Initial value for the diagonal-scale bias b_diag (a small negative number keeps initial noise small).

pyrox_name str | None

Explicit scope name for NumPyro site registration.

Examples:

>>> import jax.random as jr
>>> import jax.numpy as jnp
>>> from numpyro import handlers
>>> layer = MCSoftmaxDenseFA.init(
...     jr.PRNGKey(0), in_features=4, num_classes=3, rank=2,
... )
>>> x = jnp.ones((5, 4))
>>> with handlers.seed(rng_seed=0):
...     probs = layer(x)
>>> probs.shape
(5, 3)
>>> bool(jnp.allclose(probs.sum(axis=-1), 1.0))
True
References

Collier, M., Mustafa, B., Kokiopoulou, E., Jenatton, R., & Berent, J. (2021). Correlated Input-Dependent Label Noise in Large-Scale Image Classification. CVPR.

Source code in packages/pyrox-nn/src/pyrox_nn/_heteroscedastic.py
class MCSoftmaxDenseFA(_HeteroscedasticBase):
    r"""Heteroscedastic multi-class output layer (FA noise + softmax).

    Implements Collier et al. (2021): the logit covariance is
    input-dependent low-rank-plus-diagonal,

    $$
    \eta(x) = W_\mu x + b_\mu + \epsilon, \qquad
    \Sigma(x) = V(x) V(x)^\top + \operatorname{diag}\!\bigl(\sigma^2(x)\bigr),
    \;\; \epsilon \sim \mathcal{N}(0, \Sigma(x)),
    $$

    where $V(x) = \mathrm{reshape}(W_V x + b_V, [C, r])$ and
    $\sigma(x) = \exp(W_\sigma x + b_\sigma)$. Output is the
    Monte Carlo average of softmaxed perturbed logits

    $$
    \hat{p}(y = k \mid x) \approx
    \frac{1}{S}\sum_{s=1}^{S}
    \mathrm{softmax}_k\!\bigl(\eta(x) + \epsilon_s\bigr).
    $$

    All linear factors are deterministic ``pyrox_param`` sites — the
    layer is heteroscedastic but not Bayesian over its weights. Use it
    as a drop-in head for classification when label noise is known to
    be input-dependent (label disagreement, fine-grained categories).

    Plate semantics:
        Same as other ``pyrox_nn`` Bayesian dense layers — call
        outside ``numpyro.plate("data", ..., subsample_size=...)`` so
        the parameter sites are unscaled. The MC noise is drawn from
        ``numpyro.prng_key()``.

    Attributes:
        in_features: Input dimension $D_\mathrm{in}$.
        num_classes: Number of classes $C$.
        rank: Rank $r$ of the low-rank factor $V(x)$.
        num_mc_samples: Number of MC softmax samples $S$ per
            forward call.
        diag_init_bias: Initial value for the diagonal-scale bias
            ``b_diag`` (a small negative number keeps initial noise
            small).
        pyrox_name: Explicit scope name for NumPyro site registration.

    Examples:
        >>> import jax.random as jr
        >>> import jax.numpy as jnp
        >>> from numpyro import handlers
        >>> layer = MCSoftmaxDenseFA.init(
        ...     jr.PRNGKey(0), in_features=4, num_classes=3, rank=2,
        ... )
        >>> x = jnp.ones((5, 4))
        >>> with handlers.seed(rng_seed=0):
        ...     probs = layer(x)
        >>> probs.shape
        (5, 3)
        >>> bool(jnp.allclose(probs.sum(axis=-1), 1.0))
        True

    References:
        Collier, M., Mustafa, B., Kokiopoulou, E., Jenatton, R., &
        Berent, J. (2021). *Correlated Input-Dependent Label Noise in
        Large-Scale Image Classification.* CVPR.
    """

    @classmethod
    def _core_cls(cls) -> type[geonnax.HeteroscedasticHead]:
        return geonnax.MCSoftmaxDenseFA

    @pyrox_method
    def __call__(self, x: Float[Array, "N D_in"]) -> Float[Array, "N C"]:
        new_core = self._swap_core()
        logits = _batched_logits(new_core, x)
        return jnp.mean(jax.nn.softmax(logits, axis=-1), axis=0)

Bayesian Neural Field stack

Standardization

Bases: PyroxModule

Apply a fixed-coefficient affine standardization.

\[ \tilde x \;=\; \frac{x - \mu}{\sigma}. \]

Both mu and std are static (fit-time) constants, not learned. Use pyrox_nn.preprocessing.fit_standardization to construct from a pandas DataFrame.

Attributes:

Name Type Description
mu Float[Array, ' D']

Per-feature mean, shape (D,).

std Float[Array, ' D']

Per-feature standard deviation, shape (D,). Must be strictly positive — guard upstream.

pyrox_name str | None

Optional override for the per-instance scope name.

Source code in packages/pyrox-nn/src/pyrox_nn/_bnf.py
class Standardization(PyroxModule):
    r"""Apply a fixed-coefficient affine standardization.

    $$
    \tilde x \;=\; \frac{x - \mu}{\sigma}.
    $$

    Both ``mu`` and ``std`` are static (fit-time) constants, not
    learned. Use `pyrox_nn.preprocessing.fit_standardization` to
    construct from a pandas DataFrame.

    Attributes:
        mu: Per-feature mean, shape ``(D,)``.
        std: Per-feature standard deviation, shape ``(D,)``. Must be
            strictly positive — guard upstream.
        pyrox_name: Optional override for the per-instance scope name.
    """

    mu: Float[Array, " D"]
    std: Float[Array, " D"]
    pyrox_name: str | None = None

    @pyrox_method
    def __call__(self, x: Float[Array, "N D"]) -> Float[Array, "N D"]:
        return (x - self.mu) / self.std

FourierFeatures

Bases: PyroxModule

Per-input dyadic-frequency Fourier basis.

For each input column, evaluates 2 * degree Fourier features at frequencies \(2\pi \cdot 2^d\) for \(d \in \{0, \dots, \text{degree} - 1\}\). Concatenated across all columns.

Wraps pyrox_nn._features.fourier_features per input dimension.

Attributes:

Name Type Description
degrees tuple[int, ...]

Number of dyadic frequencies per input column, as a Python tuple[int, ...]. A column with degree = 0 contributes no features. Marked static so the loop over columns unrolls at trace time.

rescale bool

If True, divide each (cos_d, sin_d) pair by d + 1 to bias the prior toward lower frequencies.

pyrox_name str | None

Optional scope-name override.

Source code in packages/pyrox-nn/src/pyrox_nn/_bnf.py
class FourierFeatures(PyroxModule):
    r"""Per-input dyadic-frequency Fourier basis.

    For each input column, evaluates ``2 * degree`` Fourier features at
    frequencies $2\pi \cdot 2^d$ for $d \in \{0, \dots,
    \text{degree} - 1\}$. Concatenated across all columns.

    Wraps `pyrox_nn._features.fourier_features` per input
    dimension.

    Attributes:
        degrees: Number of dyadic frequencies per input column, as a
            Python ``tuple[int, ...]``. A column with ``degree = 0``
            contributes no features. Marked ``static`` so the loop
            over columns unrolls at trace time.
        rescale: If ``True``, divide each ``(cos_d, sin_d)`` pair by
            ``d + 1`` to bias the prior toward lower frequencies.
        pyrox_name: Optional scope-name override.
    """

    degrees: tuple[int, ...] = eqx.field(static=True)
    rescale: bool = eqx.field(static=True, default=False)
    pyrox_name: str | None = None

    @pyrox_method
    def __call__(self, x: Float[Array, "N D"]) -> Float[Array, "N F"]:
        feats = []
        for col_idx, d in enumerate(self.degrees):
            if d <= 0:
                continue
            feats.append(fourier_features(x[:, col_idx], d, rescale=self.rescale))
        if not feats:
            return jnp.zeros((x.shape[0], 0), dtype=x.dtype)
        return jnp.concatenate(feats, axis=-1)

SeasonalFeatures

Bases: PyroxModule

Period-and-harmonic cos/sin basis on a scalar time axis.

For each period \(\tau_p\) with \(H_p\) harmonics, emits 2 * H_p cos/sin columns. Total output width is \(2 \sum_p H_p\).

Wraps pyrox_nn._features.seasonal_features. Periods and harmonics are kept as Python tuples (static) so the inner shape structure is known at trace time.

Attributes:

Name Type Description
periods tuple[float, ...]

Period values, tuple[float, ...].

harmonics tuple[int, ...]

Harmonics per period, tuple[int, ...].

rescale bool

If True, divide each (cos, sin) pair by its within-period harmonic index.

pyrox_name str | None

Optional scope-name override.

Source code in packages/pyrox-nn/src/pyrox_nn/_bnf.py
class SeasonalFeatures(PyroxModule):
    r"""Period-and-harmonic cos/sin basis on a scalar time axis.

    For each period $\tau_p$ with $H_p$ harmonics, emits
    ``2 * H_p`` cos/sin columns. Total output width is $2 \sum_p
    H_p$.

    Wraps `pyrox_nn._features.seasonal_features`. Periods and
    harmonics are kept as Python tuples (static) so the inner shape
    structure is known at trace time.

    Attributes:
        periods: Period values, ``tuple[float, ...]``.
        harmonics: Harmonics per period, ``tuple[int, ...]``.
        rescale: If ``True``, divide each ``(cos, sin)`` pair by its
            within-period harmonic index.
        pyrox_name: Optional scope-name override.
    """

    periods: tuple[float, ...] = eqx.field(static=True)
    harmonics: tuple[int, ...] = eqx.field(static=True)
    rescale: bool = eqx.field(static=True, default=False)
    pyrox_name: str | None = None

    @pyrox_method
    def __call__(self, t: Float[Array, " N"]) -> Float[Array, "N F"]:
        return seasonal_features(t, self.periods, self.harmonics, rescale=self.rescale)

InteractionFeatures

Bases: PyroxModule

Element-wise products on selected pairs of input columns.

Wraps pyrox_nn._features.interaction_features.

Attributes:

Name Type Description
pairs tuple[tuple[int, int], ...]

Index pairs, tuple[tuple[int, int], ...]. Empty tuple produces an (N, 0) output. Static so the count K is known at trace time.

pyrox_name str | None

Optional scope-name override.

Source code in packages/pyrox-nn/src/pyrox_nn/_bnf.py
class InteractionFeatures(PyroxModule):
    r"""Element-wise products on selected pairs of input columns.

    Wraps `pyrox_nn._features.interaction_features`.

    Attributes:
        pairs: Index pairs, ``tuple[tuple[int, int], ...]``. Empty
            tuple produces an ``(N, 0)`` output. Static so the count
            ``K`` is known at trace time.
        pyrox_name: Optional scope-name override.
    """

    pairs: tuple[tuple[int, int], ...] = eqx.field(static=True)
    pyrox_name: str | None = None

    @pyrox_method
    def __call__(self, x: Float[Array, "N D"]) -> Float[Array, "N K"]:
        if not self.pairs:
            return jnp.zeros((x.shape[0], 0), dtype=x.dtype)
        return interaction_features(x, jnp.asarray(self.pairs, dtype=jnp.int32))

BayesianNeuralField

Bases: PyroxModule

The full Bayesian Neural Field architecture.

A spatiotemporal MLP with:

  1. A learned per-input log-scale adjustment (Logistic(0, 1) prior).
  2. Four feature blocks concatenated into h_0: rescaled inputs, Fourier features, seasonal features, interaction products.
  3. Per-block softplus(feature_gain) modulation.
  4. A depth-L MLP whose layers are \(h_{\ell+1} = \sigma_\alpha\bigl(g_\ell \cdot W_\ell\, h_\ell / \sqrt{\lvert h_\ell \rvert}\bigr)\), where \(\sigma_\alpha = \mathrm{sig}(\beta) \cdot \mathrm{elu} + (1 - \mathrm{sig}(\beta)) \cdot \mathrm{tanh}\) is a learned mixed activation.
  5. A final linear layer scaled by softplus(output_gain).

All weights, biases, gains, scales, and the activation logit carry independent \(\mathrm{Logistic}(0, 1)\) priors registered via PyroxModule.pyrox_sample.

The \(1/\sqrt{\text{fan-in}}\) pre-normalization is the standard NTK-scaling trick — it makes the layer-wise prior predictive a fan-in-independent Gaussian process in the infinite-width limit (Lee et al., 2018).

Attributes:

Name Type Description
input_scales tuple[float, ...]

Per-input fixed scale (typically training-data inter-quartile range). Static tuple[float, ...].

fourier_degrees tuple[int, ...]

Per-input number of dyadic Fourier frequencies. Static tuple[int, ...]; use 0 to skip a column.

interactions tuple[tuple[int, int], ...]

Pair-index list for interaction features. Static tuple[tuple[int, int], ...]; empty for none.

seasonality_periods tuple[float, ...]

Periods for seasonal features. Static tuple[float, ...]. The time variable is taken from input column time_col.

num_seasonal_harmonics tuple[int, ...]

Harmonics per period. Static tuple[int, ...].

width int

Hidden layer width.

depth int

Number of hidden MLP layers.

time_col int

Index of the time column inside x used for seasonal features (default 0).

pyrox_name str | None

Optional scope-name override.

Source code in packages/pyrox-nn/src/pyrox_nn/_bnf.py
class BayesianNeuralField(PyroxModule):
    r"""The full Bayesian Neural Field architecture.

    A spatiotemporal MLP with:

    1. A learned per-input log-scale adjustment (Logistic(0, 1) prior).
    2. Four feature blocks concatenated into ``h_0``: rescaled inputs,
       Fourier features, seasonal features, interaction products.
    3. Per-block ``softplus(feature_gain)`` modulation.
    4. A depth-``L`` MLP whose layers are
       $h_{\ell+1} = \sigma_\alpha\bigl(g_\ell \cdot W_\ell\, h_\ell
       / \sqrt{\lvert h_\ell \rvert}\bigr)$, where $\sigma_\alpha
       = \mathrm{sig}(\beta) \cdot \mathrm{elu} + (1 - \mathrm{sig}(\beta))
       \cdot \mathrm{tanh}$ is a learned mixed activation.
    5. A final linear layer scaled by ``softplus(output_gain)``.

    All weights, biases, gains, scales, and the activation logit carry
    independent $\mathrm{Logistic}(0, 1)$ priors registered via
    `PyroxModule.pyrox_sample`.

    The $1/\sqrt{\text{fan-in}}$ pre-normalization is the
    standard NTK-scaling trick — it makes the layer-wise prior
    predictive a fan-in-independent Gaussian process in the
    infinite-width limit (Lee et al., 2018).

    Attributes:
        input_scales: Per-input fixed scale (typically training-data
            inter-quartile range). Static ``tuple[float, ...]``.
        fourier_degrees: Per-input number of dyadic Fourier
            frequencies. Static ``tuple[int, ...]``; use ``0`` to skip
            a column.
        interactions: Pair-index list for interaction features. Static
            ``tuple[tuple[int, int], ...]``; empty for none.
        seasonality_periods: Periods for seasonal features. Static
            ``tuple[float, ...]``. The time variable is taken from
            input column ``time_col``.
        num_seasonal_harmonics: Harmonics per period. Static
            ``tuple[int, ...]``.
        width: Hidden layer width.
        depth: Number of hidden MLP layers.
        time_col: Index of the time column inside ``x`` used for
            seasonal features (default 0).
        pyrox_name: Optional scope-name override.
    """

    input_scales: tuple[float, ...] = eqx.field(static=True)
    fourier_degrees: tuple[int, ...] = eqx.field(static=True)
    interactions: tuple[tuple[int, int], ...] = eqx.field(static=True)
    seasonality_periods: tuple[float, ...] = eqx.field(static=True)
    num_seasonal_harmonics: tuple[int, ...] = eqx.field(static=True)
    width: int = eqx.field(static=True)
    depth: int = eqx.field(static=True)
    time_col: int = eqx.field(static=True, default=0)
    pyrox_name: str | None = None

    @staticmethod
    def _logistic_prior(shape: tuple[int, ...]) -> dist.Distribution:
        """Independent Logistic(0, 1) prior over an array of given shape."""
        if not shape:
            return dist.Logistic(0.0, 1.0)
        return dist.Logistic(jnp.zeros(shape), 1.0).to_event(len(shape))

    @pyrox_method
    def __call__(self, x: Float[Array, "N D"]) -> Float[Array, " N"]:
        d_in = len(self.input_scales)
        input_scales = jnp.asarray(self.input_scales, dtype=jnp.float32)

        # 1. Input rescaling: x / (input_scales * exp(log_scale_adjustment)).
        log_scale_adjustment = self.pyrox_sample(
            "log_scale_adjustment",
            self._logistic_prior((d_in,)),
        )
        scaled_x = x / (input_scales * jnp.exp(log_scale_adjustment))

        # 2. Build the four feature blocks.
        feature_blocks: list[Float[Array, "N F_block"]] = [scaled_x]

        # Fourier per input dim (only for degrees > 0).
        for col_idx, d in enumerate(self.fourier_degrees):
            if d > 0:
                feature_blocks.append(
                    fourier_features(scaled_x[:, col_idx], d, rescale=True)
                )

        # Seasonal on the time column.
        if self.seasonality_periods and any(self.num_seasonal_harmonics):
            feature_blocks.append(
                seasonal_features(
                    x[:, self.time_col],
                    self.seasonality_periods,
                    self.num_seasonal_harmonics,
                    rescale=True,
                )
            )

        # Interaction products.
        if self.interactions:
            feature_blocks.append(
                interaction_features(
                    scaled_x, jnp.asarray(self.interactions, dtype=jnp.int32)
                )
            )

        # 3. Per-block softplus(feature_gain) modulation.
        gated_blocks: list[Float[Array, "N F_block"]] = []
        for b_idx, block in enumerate(feature_blocks):
            if block.shape[-1] == 0:
                continue
            gain = self.pyrox_sample(
                f"feature_gain_{b_idx}",
                self._logistic_prior(()),
            )
            gated_blocks.append(block * jax.nn.softplus(gain))
        h = jnp.concatenate(gated_blocks, axis=-1)

        # 4. Mixed elu/tanh activation, learned mix weight.
        logit_activation_weight = self.pyrox_sample(
            "logit_activation_weight",
            self._logistic_prior(()),
        )
        alpha = jax.nn.sigmoid(logit_activation_weight)

        def activation(z: Float[Array, ...]) -> Float[Array, ...]:
            return alpha * jax.nn.elu(z) + (1.0 - alpha) * jnp.tanh(z)

        # 5. Hidden MLP layers.
        for layer_idx in range(self.depth):
            fan_in = h.shape[-1]
            W = self.pyrox_sample(
                f"layer_{layer_idx}_W",
                self._logistic_prior((fan_in, self.width)),
            )
            b = self.pyrox_sample(
                f"layer_{layer_idx}_b",
                self._logistic_prior((self.width,)),
            )
            layer_gain = self.pyrox_sample(
                f"layer_{layer_idx}_gain",
                self._logistic_prior(()),
            )
            h = h / jnp.sqrt(fan_in)
            h = activation(
                jax.nn.softplus(layer_gain)
                * (einx.dot("... i, i o -> ... o", h, W) + b)
            )

        # 6. Output linear, scaled by softplus(output_gain).
        fan_in = h.shape[-1]
        W_out = self.pyrox_sample(
            "output_W",
            self._logistic_prior((fan_in, 1)),
        )
        b_out = self.pyrox_sample(
            "output_b",
            self._logistic_prior((1,)),
        )
        output_gain = self.pyrox_sample(
            "output_gain",
            self._logistic_prior(()),
        )
        h = h / jnp.sqrt(fan_in)
        out = einx.dot("... i, i o -> ... o", h, W_out) + b_out
        return (jax.nn.softplus(output_gain) * out).squeeze(-1)

Pure-JAX feature helpers

fourier_features(x: Float[Array, ' N'], max_degree: int, *, rescale: bool = False) -> Float[Array, 'N two_max_degree']

Cos/sin Fourier basis at dyadic frequencies.

For each input element and each degree \(d \in \{0, \dots, D-1\}\), evaluates

\[ \phi_{d, \cos}(x) = \cos(2\pi \cdot 2^d \cdot x), \qquad \phi_{d, \sin}(x) = \sin(2\pi \cdot 2^d \cdot x). \]

Returns the columns concatenated as [cos_0, ..., cos_{D-1}, sin_0, ..., sin_{D-1}], matching Google's bayesnf layout.

Parameters:

Name Type Description Default
x Float[Array, ' N']

Length-N input vector.

required
max_degree int

Number of dyadic frequencies D. Output has 2 * max_degree columns.

required
rescale bool

If True, divide each (cos_d, sin_d) pair by d + 1 to bias the prior toward lower-frequency basis functions.

False

Returns:

Type Description
Float[Array, 'N two_max_degree']

Array of shape (N, 2 * max_degree).

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.basis import fourier_features
>>> fourier_features(jnp.linspace(0.0, 1.0, 5), max_degree=3).shape
(5, 6)
Source code in .venv/lib/python3.12/site-packages/geonnax/basis.py
def fourier_features(
    x: Float[Array, " N"],
    max_degree: int,
    *,
    rescale: bool = False,
) -> Float[Array, "N two_max_degree"]:
    r"""Cos/sin Fourier basis at dyadic frequencies.

    For each input element and each degree $d \in \{0, \dots,
    D-1\}$, evaluates

    $$
    \phi_{d, \cos}(x) = \cos(2\pi \cdot 2^d \cdot x), \qquad
    \phi_{d, \sin}(x) = \sin(2\pi \cdot 2^d \cdot x).
    $$


    Returns the columns concatenated as ``[cos_0, ..., cos_{D-1},
    sin_0, ..., sin_{D-1}]``, matching Google's bayesnf layout.

    Args:
        x: Length-``N`` input vector.
        max_degree: Number of dyadic frequencies ``D``. Output has
            ``2 * max_degree`` columns.
        rescale: If ``True``, divide each ``(cos_d, sin_d)`` pair by
            ``d + 1`` to bias the prior toward lower-frequency basis
            functions.

    Returns:
        Array of shape ``(N, 2 * max_degree)``.

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.basis import fourier_features
        >>> fourier_features(jnp.linspace(0.0, 1.0, 5), max_degree=3).shape
        (5, 6)
    """
    degrees = jnp.arange(max_degree)
    # Broadcast x to (N, D) frequencies without an explicit reshape.
    z = einx.id("n -> n d", x, d=max_degree) * (2.0 * jnp.pi * 2.0**degrees)
    feats = jnp.concatenate([jnp.cos(z), jnp.sin(z)], axis=-1)
    if rescale:
        denom = jnp.concatenate([degrees + 1, degrees + 1])
        feats = feats / denom
    return feats

seasonal_features(x: Float[Array, ' N'], periods: Sequence[float], harmonics: Sequence[int], *, rescale: bool = False) -> Float[Array, 'N two_F']

Cos/sin features at multiples of \(2\pi / \tau_p\).

For each period \(\tau_p\) with \(H_p\) harmonics, evaluates

\[ \phi_{p, h, \cos}(x) = \cos(2\pi h x / \tau_p), \qquad \phi_{p, h, \sin}(x) = \sin(2\pi h x / \tau_p), \]

for \(h = 1, \dots, H_p\). Returns the cos columns concatenated with the sin columns, length \(F = \sum_p H_p\) each.

periods and harmonics are Python sequences (tuples, lists, or 0-d JAX arrays wrapped at the call site). Keeping them as Python values lets the function run cleanly under jax.jit and lax.scan without triggering a concretization error.

Parameters:

Name Type Description Default
x Float[Array, ' N']

Time/index input, shape (N,).

required
periods Sequence[float]

Period values.

required
harmonics Sequence[int]

Harmonics per period.

required
rescale bool

If True, divide each (cos, sin) pair by its within-period harmonic index, biasing the prior toward longer-wavelength modes within each period.

False

Returns:

Type Description
Float[Array, 'N two_F']

Array of shape (N, 2 * F).

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.basis import seasonal_features
>>> x = jnp.linspace(0.0, 10.0, 4)
>>> # F = 1 + 2 = 3 frequencies -> 2 * F = 6 columns
>>> seasonal_features(x, [7.0, 365.0], [1, 2]).shape
(4, 6)
Source code in .venv/lib/python3.12/site-packages/geonnax/basis.py
def seasonal_features(
    x: Float[Array, " N"],
    periods: Sequence[float],
    harmonics: Sequence[int],
    *,
    rescale: bool = False,
) -> Float[Array, "N two_F"]:
    r"""Cos/sin features at multiples of $2\pi / \tau_p$.

    For each period $\tau_p$ with $H_p$ harmonics, evaluates

    $$
    \phi_{p, h, \cos}(x) = \cos(2\pi h x / \tau_p), \qquad
    \phi_{p, h, \sin}(x) = \sin(2\pi h x / \tau_p),
    $$


    for $h = 1, \dots, H_p$. Returns the cos columns concatenated
    with the sin columns, length $F = \sum_p H_p$ each.

    ``periods`` and ``harmonics`` are **Python sequences** (tuples,
    lists, or 0-d JAX arrays wrapped at the call site). Keeping them as
    Python values lets the function run cleanly under ``jax.jit`` and
    ``lax.scan`` without triggering a concretization error.

    Args:
        x: Time/index input, shape ``(N,)``.
        periods: Period values.
        harmonics: Harmonics per period.
        rescale: If ``True``, divide each ``(cos, sin)`` pair by its
            within-period harmonic index, biasing the prior toward
            longer-wavelength modes within each period.

    Returns:
        Array of shape ``(N, 2 * F)``.

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.basis import seasonal_features
        >>> x = jnp.linspace(0.0, 10.0, 4)
        >>> # F = 1 + 2 = 3 frequencies -> 2 * F = 6 columns
        >>> seasonal_features(x, [7.0, 365.0], [1, 2]).shape
        (4, 6)
    """
    _, freq_list = seasonal_frequencies(periods, harmonics)
    if not freq_list:
        return jnp.zeros((x.shape[0], 0), dtype=x.dtype)
    freqs = jnp.asarray(freq_list, dtype=jnp.float32)
    z = einx.id("n -> n f", x, f=freqs.shape[0]) * (2.0 * jnp.pi * freqs)
    feats = jnp.concatenate([jnp.cos(z), jnp.sin(z)], axis=-1)
    if rescale:
        # Rescale by within-period harmonic index (1, 2, ..., H_p).
        h_within_list: list[float] = []
        for n_h in harmonics:
            h_within_list.extend(range(1, int(n_h) + 1))
        h_within = jnp.asarray(h_within_list, dtype=jnp.float32)
        denom = jnp.concatenate([h_within, h_within])
        feats = feats / denom
    return feats

seasonal_frequencies(periods: Sequence[float], harmonics: Sequence[int]) -> tuple[list[int], list[float]]

Flatten (period, harmonic_count) pairs into Python lists.

For each period \(\tau_p\) with \(H_p\) harmonics, emits frequencies \(f_{p, h} = h / \tau_p\) for \(h = 1, \dots, H_p\). The total length is \(F = \sum_p H_p\).

Inputs are Python sequences, not JAX arrays, so this helper runs at trace time and never triggers a concretization error under jax.jit. Most callers won't use it directly; it's exposed for symmetry with seasonal_features.

Parameters:

Name Type Description Default
periods Sequence[float]

Period values.

required
harmonics Sequence[int]

Number of harmonics per period.

required

Returns:

Type Description
list[int]

(period_index, frequency): two Python lists of length

list[float]

\(F = \sum_p H_p\).

Examples:

>>> from geonnax.basis import seasonal_frequencies
>>> # periods (7, 365) with (1, 2) harmonics -> F = 1 + 2 = 3 freqs
>>> idx, freqs = seasonal_frequencies([7.0, 365.0], [1, 2])
>>> idx
[0, 1, 1]
>>> len(freqs)
3
Source code in .venv/lib/python3.12/site-packages/geonnax/basis.py
def seasonal_frequencies(
    periods: Sequence[float],
    harmonics: Sequence[int],
) -> tuple[list[int], list[float]]:
    r"""Flatten ``(period, harmonic_count)`` pairs into Python lists.

    For each period $\tau_p$ with $H_p$ harmonics, emits
    frequencies $f_{p, h} = h / \tau_p$ for $h = 1, \dots,
    H_p$. The total length is $F = \sum_p H_p$.

    Inputs are **Python sequences**, not JAX arrays, so this helper
    runs at trace time and never triggers a concretization error under
    ``jax.jit``. Most callers won't use it directly; it's exposed for
    symmetry with `seasonal_features`.

    Args:
        periods: Period values.
        harmonics: Number of harmonics per period.

    Returns:
        ``(period_index, frequency)``: two Python lists of length
        $F = \sum_p H_p$.

    Examples:
        >>> from geonnax.basis import seasonal_frequencies
        >>> # periods (7, 365) with (1, 2) harmonics -> F = 1 + 2 = 3 freqs
        >>> idx, freqs = seasonal_frequencies([7.0, 365.0], [1, 2])
        >>> idx
        [0, 1, 1]
        >>> len(freqs)
        3
    """
    period_index: list[int] = []
    freqs: list[float] = []
    for p_idx, (period, n_h) in enumerate(zip(periods, harmonics, strict=True)):
        for h in range(1, int(n_h) + 1):
            period_index.append(p_idx)
            freqs.append(float(h) / float(period))
    return period_index, freqs

interaction_features(x: Float[Array, 'N D'], pairs: Int[Array, 'K 2']) -> Float[Array, 'N K']

Element-wise products on selected pairs of input columns.

For each pair \((i, j)\) and each row \(n\), computes \(x_{n, i} \cdot x_{n, j}\).

Parameters:

Name Type Description Default
x Float[Array, 'N D']

Input matrix, shape (N, D).

required
pairs Int[Array, 'K 2']

Index pairs, shape (K, 2). Empty pairs yield an (N, 0) output.

required

Returns:

Type Description
Float[Array, 'N K']

Array of shape (N, K) of pairwise products.

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.basis import interaction_features
>>> x = jnp.arange(6.0).reshape(2, 3)  # (N=2, D=3)
>>> pairs = jnp.array([[0, 1], [0, 2]])  # (K=2, 2)
>>> interaction_features(x, pairs).shape
(2, 2)
Source code in .venv/lib/python3.12/site-packages/geonnax/basis.py
def interaction_features(
    x: Float[Array, "N D"],
    pairs: Int[Array, "K 2"],
) -> Float[Array, "N K"]:
    r"""Element-wise products on selected pairs of input columns.

    For each pair $(i, j)$ and each row $n$, computes
    $x_{n, i} \cdot x_{n, j}$.

    Args:
        x: Input matrix, shape ``(N, D)``.
        pairs: Index pairs, shape ``(K, 2)``. Empty pairs yield an
            ``(N, 0)`` output.

    Returns:
        Array of shape ``(N, K)`` of pairwise products.

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.basis import interaction_features
        >>> x = jnp.arange(6.0).reshape(2, 3)  # (N=2, D=3)
        >>> pairs = jnp.array([[0, 1], [0, 2]])  # (K=2, 2)
        >>> interaction_features(x, pairs).shape
        (2, 2)
    """
    if pairs.shape[0] == 0:
        return jnp.zeros((x.shape[0], 0), dtype=x.dtype)
    # x[:, pairs] has shape (N, K, 2); reduce the paired axis with prod.
    selected = x[:, pairs]
    return einx.prod("n k [two]", selected)

standardize(x: Float[Array, '*shape'], mu: Float[Array, '*shape'], std: Float[Array, '*shape']) -> Float[Array, '*shape']

Affine standardize: (x - mu) / std.

Broadcasts mu and std against x per the JAX broadcasting rules. std is not clamped; pass a positive value or guard upstream.

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.basis import standardize
>>> x = jnp.array([1.0, 3.0, 5.0])
>>> bool(jnp.allclose(standardize(x, 3.0, 2.0), jnp.array([-1, 0, 1])))
True
Source code in .venv/lib/python3.12/site-packages/geonnax/basis.py
def standardize(
    x: Float[Array, "*shape"],
    mu: Float[Array, "*shape"],
    std: Float[Array, "*shape"],
) -> Float[Array, "*shape"]:
    """Affine standardize: ``(x - mu) / std``.

    Broadcasts ``mu`` and ``std`` against ``x`` per the JAX broadcasting
    rules. ``std`` is *not* clamped; pass a positive value or guard
    upstream.

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.basis import standardize
        >>> x = jnp.array([1.0, 3.0, 5.0])
        >>> bool(jnp.allclose(standardize(x, 3.0, 2.0), jnp.array([-1, 0, 1])))
        True
    """
    return (x - mu) / std

unstandardize(z: Float[Array, '*shape'], mu: Float[Array, '*shape'], std: Float[Array, '*shape']) -> Float[Array, '*shape']

Inverse of standardize: z * std + mu.

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.basis import standardize, unstandardize
>>> x = jnp.array([1.0, 3.0, 5.0])
>>> # unstandardize undoes standardize for the same (mu, std).
>>> z = standardize(x, 3.0, 2.0)
>>> bool(jnp.allclose(unstandardize(z, 3.0, 2.0), x))
True
Source code in .venv/lib/python3.12/site-packages/geonnax/basis.py
def unstandardize(
    z: Float[Array, "*shape"],
    mu: Float[Array, "*shape"],
    std: Float[Array, "*shape"],
) -> Float[Array, "*shape"]:
    """Inverse of `standardize`: ``z * std + mu``.

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.basis import standardize, unstandardize
        >>> x = jnp.array([1.0, 3.0, 5.0])
        >>> # unstandardize undoes standardize for the same (mu, std).
        >>> z = standardize(x, 3.0, 2.0)
        >>> bool(jnp.allclose(unstandardize(z, 3.0, 2.0), x))
        True
    """
    return z * std + mu