Skip to content

Multiplicative Filter Networks (MFN)

Multiplicative Filter Networks (Fathony, Sahu, Willmott, Kolter — ICLR 2021) replace MLP composition with multiplicative filter chaining. Instead of deeply composing nonlinearities, each layer multiplies the previous activation by a new filter evaluated directly on the original input:

\[ z_1 = g_1(x), \qquad z_{i+1} = g_{i+1}(x) \odot \bigl(W_i z_i + b_i\bigr), \qquad y = W_L z_L + b_L. \]

Two filter families ship in pyrox_nn:

  • FourierNet — \(g_i(x) = \sin(\Omega_i x + \varphi_i)\), frequency-domain filters. Products of sinusoids span exponentially many frequencies with depth \(L\).
  • GaborNet — \(g_i(x) = \sin(\Omega_i x + \varphi_i) \odot \exp(-\tfrac{\gamma_i}{2}\|x - \mu_i\|^2)\), Gabor atoms with learned frequency \(\Omega_i\), phase \(\varphi_i\), location \(\mu_i\), and bandwidth \(\gamma_i\).

Connection to RBFFourierFeatures: A GaborNet with depth=1 and \(\mu = 0\) is a localized variant of random Fourier features. As \(\gamma \to 0\) (very wide envelope) it recovers the plain RBF-RFF feature map. See HSGPFeatures for the related Hilbert-space GP basis.

Quick example

import jax.random as jr
from pyrox_nn import GaborNet
from numpyro import handlers

key = jr.PRNGKey(0)

# Deterministic GaborNet
net = GaborNet.init(in_features=2, hidden_features=64, out_features=1, depth=3, key=key)

import jax.numpy as jnp
x = jnp.ones((100, 2))
y = net(x)          # (100, 1)

# Bayesian GaborNet — sample sites registered for every parameter
from pyrox_nn import BayesianGaborNet
bnet = BayesianGaborNet.init(
    in_features=2, hidden_features=64, out_features=1, depth=3, key=key,
    pyrox_name="gabor",
)
with handlers.seed(rng_seed=1):
    y_sample = bnet(x)  # weights sampled from prior

Filter primitives

FourierFilter

Bases: Module

Single Fourier filter: \(g(x) = \sin(\Omega x + \varphi)\).

One multiplicative filter primitive for use inside a FourierNet.

Init follows Fathony et al. (2021) §4.1: frequencies are drawn as \(\Omega_{ij} \sim \mathcal{N}(0,\,\sigma_f^2/D)\) where \(D\) is in_features and \(\sigma_f\) is freq_scale; phases are drawn as \(\varphi_i \sim \mathrm{Uniform}(-\pi, \pi)\).

Attributes:

Name Type Description
Omega Float[Array, 'out in']

Frequency matrix of shape (out_features, in_features).

phi Float[Array, ' out']

Phase vector of shape (out_features,).

in_features int

Input dimension.

out_features int

Output (filter) dimension.

Source code in .venv/lib/python3.12/site-packages/geonnax/mfn.py
class FourierFilter(eqx.Module):
    r"""Single Fourier filter: $g(x) = \sin(\Omega x + \varphi)$.

    One multiplicative filter primitive for use inside a
    `FourierNet`.

    Init follows Fathony et al. (2021) §4.1: frequencies are drawn as
    $\Omega_{ij} \sim \mathcal{N}(0,\,\sigma_f^2/D)$ where
    $D$ is ``in_features`` and $\sigma_f$ is
    ``freq_scale``; phases are drawn as
    $\varphi_i \sim \mathrm{Uniform}(-\pi, \pi)$.

    Attributes:
        Omega: Frequency matrix of shape ``(out_features, in_features)``.
        phi: Phase vector of shape ``(out_features,)``.
        in_features: Input dimension.
        out_features: Output (filter) dimension.
    """

    Omega: Float[Array, "out in"]
    phi: Float[Array, " out"]
    in_features: int = eqx.field(static=True)
    out_features: int = eqx.field(static=True)

    @classmethod
    def init(
        cls,
        in_features: int,
        out_features: int,
        *,
        key: PRNGKeyArray,
        freq_scale: float = 256.0,
    ) -> FourierFilter:
        """Construct with Fathony-et-al. §4.1 initialization.

        Examples:
            >>> import jax.numpy as jnp, jax.random as jr
            >>> from geonnax.mfn import FourierFilter
            >>> f = FourierFilter.init(3, 8, key=jr.PRNGKey(0))
            >>> f(jnp.ones(3)).shape  # (3,) -> (8,)
            (8,)
        """
        _require_positive(
            in_features=in_features,
            out_features=out_features,
            freq_scale=freq_scale,
        )
        k_omega, k_phi = jax.random.split(key)
        omega_std = freq_scale / math.sqrt(in_features)
        Omega = jax.random.normal(k_omega, (out_features, in_features)) * omega_std
        phi = jax.random.uniform(k_phi, (out_features,), minval=-jnp.pi, maxval=jnp.pi)
        return cls(
            Omega=Omega,
            phi=phi,
            in_features=in_features,
            out_features=out_features,
        )

    def __call__(self, x: Float[Array, " D"]) -> Float[Array, " H"]:
        r"""Evaluate the filter ``g(x) = sin(Ω x + φ)``.

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

        Returns:
            Filter response of shape ``(out_features,)``.

        Examples:
            >>> import jax.numpy as jnp, jax.random as jr
            >>> from geonnax.mfn import FourierFilter
            >>> f = FourierFilter.init(2, 6, key=jr.PRNGKey(0))
            >>> f(jnp.zeros(2)).shape  # (2,) -> (6,)
            (6,)
        """
        # Frequency projection: (D,) · (H, D) -> (H,).
        proj = einx.dot("d, h d -> h", x, self.Omega)
        return jnp.sin(proj + self.phi)  # g(x) = sin(Ω x + φ), shape (H,)

init(in_features: int, out_features: int, *, key: PRNGKeyArray, freq_scale: float = 256.0) -> FourierFilter classmethod

Construct with Fathony-et-al. §4.1 initialization.

Examples:

>>> import jax.numpy as jnp, jax.random as jr
>>> from geonnax.mfn import FourierFilter
>>> f = FourierFilter.init(3, 8, key=jr.PRNGKey(0))
>>> f(jnp.ones(3)).shape  # (3,) -> (8,)
(8,)
Source code in .venv/lib/python3.12/site-packages/geonnax/mfn.py
@classmethod
def init(
    cls,
    in_features: int,
    out_features: int,
    *,
    key: PRNGKeyArray,
    freq_scale: float = 256.0,
) -> FourierFilter:
    """Construct with Fathony-et-al. §4.1 initialization.

    Examples:
        >>> import jax.numpy as jnp, jax.random as jr
        >>> from geonnax.mfn import FourierFilter
        >>> f = FourierFilter.init(3, 8, key=jr.PRNGKey(0))
        >>> f(jnp.ones(3)).shape  # (3,) -> (8,)
        (8,)
    """
    _require_positive(
        in_features=in_features,
        out_features=out_features,
        freq_scale=freq_scale,
    )
    k_omega, k_phi = jax.random.split(key)
    omega_std = freq_scale / math.sqrt(in_features)
    Omega = jax.random.normal(k_omega, (out_features, in_features)) * omega_std
    phi = jax.random.uniform(k_phi, (out_features,), minval=-jnp.pi, maxval=jnp.pi)
    return cls(
        Omega=Omega,
        phi=phi,
        in_features=in_features,
        out_features=out_features,
    )

GaborFilter

Bases: Module

Single Gabor filter: \(g(x) = \sin(\Omega x + \varphi) \odot \exp(-\tfrac{\gamma}{2}\|x - \mu\|^2)\).

Init follows Fathony et al. (2021) §4.2: per-filter \(\gamma_i \sim \mathrm{Gamma}(\alpha, \beta)\), \(\mu_i \sim \mathrm{Uniform}(\text{domain})\), \(\Omega_{i,:} \sim \mathcal{N}(0, \gamma_i\,I_D)\) (the load-bearing tied initialization).

\(\gamma\) is stored in log space so positivity is preserved without optimizer constraints.

Attributes:

Name Type Description
Omega Float[Array, 'out in']

Frequency matrix (out_features, in_features).

phi Float[Array, ' out']

Phase vector (out_features,).

mu Float[Array, 'out in']

Envelope centres (out_features, in_features).

log_gamma Float[Array, ' out']

Log-bandwidth (out_features,).

in_features int

Input dimension.

out_features int

Output (filter) dimension.

domain tuple[float, float]

(low, high) used for \(\mu\) initialization (static).

Source code in .venv/lib/python3.12/site-packages/geonnax/mfn.py
class GaborFilter(eqx.Module):
    r"""Single Gabor filter:
    $g(x) = \sin(\Omega x + \varphi) \odot \exp(-\tfrac{\gamma}{2}\|x - \mu\|^2)$.

    Init follows Fathony et al. (2021) §4.2: per-filter
    $\gamma_i \sim \mathrm{Gamma}(\alpha, \beta)$,
    $\mu_i \sim \mathrm{Uniform}(\text{domain})$,
    $\Omega_{i,:} \sim \mathcal{N}(0, \gamma_i\,I_D)$ (the
    load-bearing tied initialization).

    $\gamma$ is stored in log space so positivity is preserved
    without optimizer constraints.

    Attributes:
        Omega: Frequency matrix ``(out_features, in_features)``.
        phi: Phase vector ``(out_features,)``.
        mu: Envelope centres ``(out_features, in_features)``.
        log_gamma: Log-bandwidth ``(out_features,)``.
        in_features: Input dimension.
        out_features: Output (filter) dimension.
        domain: ``(low, high)`` used for $\mu$ initialization (static).
    """

    Omega: Float[Array, "out in"]
    phi: Float[Array, " out"]
    mu: Float[Array, "out in"]
    log_gamma: Float[Array, " out"]
    in_features: int = eqx.field(static=True)
    out_features: int = eqx.field(static=True)
    domain: tuple[float, float] = eqx.field(static=True)

    @classmethod
    def init(
        cls,
        in_features: int,
        out_features: int,
        *,
        key: PRNGKeyArray,
        domain: tuple[float, float] = (-1.0, 1.0),
        gamma_alpha: float = 6.0,
        gamma_beta: float = 1.0,
    ) -> GaborFilter:
        """Construct with Fathony-et-al. §4.2 initialization.

        Examples:
            >>> import jax.numpy as jnp, jax.random as jr
            >>> from geonnax.mfn import GaborFilter
            >>> g = GaborFilter.init(2, 6, key=jr.PRNGKey(0))
            >>> g(jnp.zeros(2)).shape  # (2,) -> (6,)
            (6,)
        """
        _require_positive(
            in_features=in_features,
            out_features=out_features,
            gamma_alpha=gamma_alpha,
            gamma_beta=gamma_beta,
        )
        if domain[0] >= domain[1]:
            raise ValueError(f"domain must satisfy low < high; got domain={domain}.")
        k_gamma, k_mu, k_omega, k_phi = jax.random.split(key, 4)
        # jax.random.gamma samples Gamma(alpha, 1); divide by beta for rate=beta.
        gamma = jax.random.gamma(k_gamma, gamma_alpha, (out_features,)) / gamma_beta
        log_gamma = jnp.log(gamma)
        mu = jax.random.uniform(
            k_mu, (out_features, in_features), minval=domain[0], maxval=domain[1]
        )
        # Tied init: Omega_i ~ N(0, gamma_i * I_D).
        Omega = jax.random.normal(k_omega, (out_features, in_features)) * jnp.sqrt(
            gamma[:, None]
        )
        phi = jax.random.uniform(k_phi, (out_features,), minval=-jnp.pi, maxval=jnp.pi)
        return cls(
            Omega=Omega,
            phi=phi,
            mu=mu,
            log_gamma=log_gamma,
            in_features=in_features,
            out_features=out_features,
            domain=domain,
        )

    def __call__(self, x: Float[Array, " D"]) -> Float[Array, " H"]:
        r"""Evaluate ``g(x) = sin(Ω x + φ) ⊙ exp(-γ/2 · ‖x - μ‖²)``.

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

        Returns:
            Filter response of shape ``(out_features,)``.

        Examples:
            >>> import jax.numpy as jnp, jax.random as jr
            >>> from geonnax.mfn import GaborFilter
            >>> g = GaborFilter.init(3, 4, key=jr.PRNGKey(1))
            >>> g(jnp.ones(3)).shape  # (3,) -> (4,)
            (4,)
        """
        gamma = jnp.exp(self.log_gamma)  # positivity via log-param, (H,)
        # ‖x - μ‖² = ‖x‖² + ‖μ‖² - 2 x·μ, expanded to reuse the einx dot.
        x_norm_sq = jnp.sum(x**2)  # scalar
        mu_norm_sq = jnp.sum(self.mu**2, axis=-1)  # (H,)
        cross = einx.dot("d, h d -> h", x, self.mu)  # (D,)·(H,D) -> (H,)
        sq_dist = jnp.maximum(x_norm_sq + mu_norm_sq - 2.0 * cross, 0.0)  # (H,)
        envelope = jnp.exp(-0.5 * gamma * sq_dist)  # Gaussian window, (H,)
        # Oscillation sin(Ω x + φ): (D,)·(H,D) -> (H,).
        sinusoidal = jnp.sin(einx.dot("d, h d -> h", x, self.Omega) + self.phi)
        return sinusoidal * envelope  # elementwise modulation, (H,)

init(in_features: int, out_features: int, *, key: PRNGKeyArray, domain: tuple[float, float] = (-1.0, 1.0), gamma_alpha: float = 6.0, gamma_beta: float = 1.0) -> GaborFilter classmethod

Construct with Fathony-et-al. §4.2 initialization.

Examples:

>>> import jax.numpy as jnp, jax.random as jr
>>> from geonnax.mfn import GaborFilter
>>> g = GaborFilter.init(2, 6, key=jr.PRNGKey(0))
>>> g(jnp.zeros(2)).shape  # (2,) -> (6,)
(6,)
Source code in .venv/lib/python3.12/site-packages/geonnax/mfn.py
@classmethod
def init(
    cls,
    in_features: int,
    out_features: int,
    *,
    key: PRNGKeyArray,
    domain: tuple[float, float] = (-1.0, 1.0),
    gamma_alpha: float = 6.0,
    gamma_beta: float = 1.0,
) -> GaborFilter:
    """Construct with Fathony-et-al. §4.2 initialization.

    Examples:
        >>> import jax.numpy as jnp, jax.random as jr
        >>> from geonnax.mfn import GaborFilter
        >>> g = GaborFilter.init(2, 6, key=jr.PRNGKey(0))
        >>> g(jnp.zeros(2)).shape  # (2,) -> (6,)
        (6,)
    """
    _require_positive(
        in_features=in_features,
        out_features=out_features,
        gamma_alpha=gamma_alpha,
        gamma_beta=gamma_beta,
    )
    if domain[0] >= domain[1]:
        raise ValueError(f"domain must satisfy low < high; got domain={domain}.")
    k_gamma, k_mu, k_omega, k_phi = jax.random.split(key, 4)
    # jax.random.gamma samples Gamma(alpha, 1); divide by beta for rate=beta.
    gamma = jax.random.gamma(k_gamma, gamma_alpha, (out_features,)) / gamma_beta
    log_gamma = jnp.log(gamma)
    mu = jax.random.uniform(
        k_mu, (out_features, in_features), minval=domain[0], maxval=domain[1]
    )
    # Tied init: Omega_i ~ N(0, gamma_i * I_D).
    Omega = jax.random.normal(k_omega, (out_features, in_features)) * jnp.sqrt(
        gamma[:, None]
    )
    phi = jax.random.uniform(k_phi, (out_features,), minval=-jnp.pi, maxval=jnp.pi)
    return cls(
        Omega=Omega,
        phi=phi,
        mu=mu,
        log_gamma=log_gamma,
        in_features=in_features,
        out_features=out_features,
        domain=domain,
    )

Composite networks

FourierNet

Bases: Module

Multiplicative Fourier Filter Network (Fathony et al., ICLR 2021).

Chains FourierFilter primitives multiplicatively:

\[ z_1 = g_1(x), \quad z_{i+1} = g_{i+1}(x) \odot (W_i z_i + b_i), \quad y = W_L z_L + b_L. \]

Each \(g_i\) is a FourierFilter of width hidden_features; the last linear is the readout projecting to out_features.

Attributes:

Name Type Description
filters list[FourierFilter]

Length-depth list of FourierFilter.

linears list[Linear]

Length-depth list of equinox.nn.Linear.

in_features int

Input dimension.

hidden_features int

Filter / hidden width.

out_features int

Output dimension.

depth int

Number of filter layers \(L\).

Source code in .venv/lib/python3.12/site-packages/geonnax/mfn.py
class FourierNet(eqx.Module):
    r"""Multiplicative Fourier Filter Network (Fathony et al., ICLR 2021).

    Chains `FourierFilter` primitives multiplicatively:

    $$
    z_1 = g_1(x), \quad
    z_{i+1} = g_{i+1}(x) \odot (W_i z_i + b_i), \quad
    y = W_L z_L + b_L.
    $$


    Each $g_i$ is a `FourierFilter` of width
    ``hidden_features``; the last linear is the readout projecting to
    ``out_features``.

    Attributes:
        filters: Length-``depth`` list of `FourierFilter`.
        linears: Length-``depth`` list of `equinox.nn.Linear`.
        in_features: Input dimension.
        hidden_features: Filter / hidden width.
        out_features: Output dimension.
        depth: Number of filter layers $L$.
    """

    filters: list[FourierFilter]
    linears: list[eqx.nn.Linear]
    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)

    @classmethod
    def init(
        cls,
        in_features: int,
        hidden_features: int,
        out_features: int,
        *,
        depth: int,
        key: PRNGKeyArray,
        freq_scale: float = 256.0,
    ) -> FourierNet:
        """Construct a ``FourierNet`` with ``depth`` filters and readout linears.

        Examples:
            >>> import jax.numpy as jnp, jax.random as jr
            >>> from geonnax.mfn import FourierNet
            >>> net = FourierNet.init(2, 16, 1, depth=3, key=jr.PRNGKey(0))
            >>> net(jnp.zeros(2)).shape  # (2,) -> (1,)
            (1,)
        """
        if depth < 1:
            raise ValueError(f"depth must be at least 1, got {depth}.")
        keys = jax.random.split(key, 2 * depth)
        filter_keys = keys[:depth]
        linear_keys = keys[depth:]
        filters = [
            FourierFilter.init(
                in_features, hidden_features, key=filter_keys[i], freq_scale=freq_scale
            )
            for i in range(depth)
        ]
        linears = [
            eqx.nn.Linear(
                hidden_features,
                hidden_features if i < depth - 1 else out_features,
                key=linear_keys[i],
            )
            for i in range(depth)
        ]
        return cls(
            filters=filters,
            linears=linears,
            in_features=in_features,
            hidden_features=hidden_features,
            out_features=out_features,
            depth=depth,
        )

    def __call__(self, x: Float[Array, " D"]) -> Float[Array, " O"]:
        """Run the multiplicative Fourier-filter forward pass.

        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.mfn import FourierNet
            >>> net = FourierNet.init(2, 8, 4, depth=2, key=jr.PRNGKey(0))
            >>> net(jnp.ones(2)).shape  # (2,) -> (4,)
            (4,)
        """
        # (D,) -> (O,) via multiplicative filter chaining (see mfn_forward).
        return mfn_forward(x, self.filters, self.linears)

init(in_features: int, hidden_features: int, out_features: int, *, depth: int, key: PRNGKeyArray, freq_scale: float = 256.0) -> FourierNet classmethod

Construct a FourierNet with depth filters and readout linears.

Examples:

>>> import jax.numpy as jnp, jax.random as jr
>>> from geonnax.mfn import FourierNet
>>> net = FourierNet.init(2, 16, 1, depth=3, key=jr.PRNGKey(0))
>>> net(jnp.zeros(2)).shape  # (2,) -> (1,)
(1,)
Source code in .venv/lib/python3.12/site-packages/geonnax/mfn.py
@classmethod
def init(
    cls,
    in_features: int,
    hidden_features: int,
    out_features: int,
    *,
    depth: int,
    key: PRNGKeyArray,
    freq_scale: float = 256.0,
) -> FourierNet:
    """Construct a ``FourierNet`` with ``depth`` filters and readout linears.

    Examples:
        >>> import jax.numpy as jnp, jax.random as jr
        >>> from geonnax.mfn import FourierNet
        >>> net = FourierNet.init(2, 16, 1, depth=3, key=jr.PRNGKey(0))
        >>> net(jnp.zeros(2)).shape  # (2,) -> (1,)
        (1,)
    """
    if depth < 1:
        raise ValueError(f"depth must be at least 1, got {depth}.")
    keys = jax.random.split(key, 2 * depth)
    filter_keys = keys[:depth]
    linear_keys = keys[depth:]
    filters = [
        FourierFilter.init(
            in_features, hidden_features, key=filter_keys[i], freq_scale=freq_scale
        )
        for i in range(depth)
    ]
    linears = [
        eqx.nn.Linear(
            hidden_features,
            hidden_features if i < depth - 1 else out_features,
            key=linear_keys[i],
        )
        for i in range(depth)
    ]
    return cls(
        filters=filters,
        linears=linears,
        in_features=in_features,
        hidden_features=hidden_features,
        out_features=out_features,
        depth=depth,
    )

GaborNet

Bases: Module

Multiplicative Gabor Filter Network (Fathony et al., ICLR 2021).

Same MFN topology as FourierNet but each \(g_i\) is a GaborFilter — a sinusoidal oscillation modulated by a Gaussian envelope:

\[ g_i(x) = \sin(\Omega_i x + \varphi_i) \odot \exp\!\bigl(-\tfrac{\gamma_i}{2}\|x - \mu_i\|^2\bigr). \]

Attributes:

Name Type Description
filters list[GaborFilter]

Length-depth list of GaborFilter.

linears list[Linear]

Length-depth list of equinox.nn.Linear.

in_features int

Input dimension.

hidden_features int

Filter / hidden width.

out_features int

Output dimension.

depth int

Number of filter layers \(L\).

domain tuple[float, float]

(low, high) used for \(\mu\) initialization.

gamma_alpha float

Shape parameter of the \(\gamma\) Gamma prior.

gamma_beta float

Rate parameter of the \(\gamma\) Gamma prior.

Source code in .venv/lib/python3.12/site-packages/geonnax/mfn.py
class GaborNet(eqx.Module):
    r"""Multiplicative Gabor Filter Network (Fathony et al., ICLR 2021).

    Same MFN topology as `FourierNet` but each $g_i$ is a
    `GaborFilter` — a sinusoidal oscillation modulated by a
    Gaussian envelope:

    $$
    g_i(x) = \sin(\Omega_i x + \varphi_i)
    \odot \exp\!\bigl(-\tfrac{\gamma_i}{2}\|x - \mu_i\|^2\bigr).
    $$


    Attributes:
        filters: Length-``depth`` list of `GaborFilter`.
        linears: Length-``depth`` list of `equinox.nn.Linear`.
        in_features: Input dimension.
        hidden_features: Filter / hidden width.
        out_features: Output dimension.
        depth: Number of filter layers $L$.
        domain: ``(low, high)`` used for $\mu$ initialization.
        gamma_alpha: Shape parameter of the $\gamma$ Gamma prior.
        gamma_beta: Rate parameter of the $\gamma$ Gamma prior.
    """

    filters: list[GaborFilter]
    linears: list[eqx.nn.Linear]
    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)
    domain: tuple[float, float] = eqx.field(static=True)
    gamma_alpha: float = eqx.field(static=True)
    gamma_beta: float = eqx.field(static=True)

    @classmethod
    def init(
        cls,
        in_features: int,
        hidden_features: int,
        out_features: int,
        *,
        depth: int,
        key: PRNGKeyArray,
        domain: tuple[float, float] = (-1.0, 1.0),
        gamma_alpha: float = 6.0,
        gamma_beta: float = 1.0,
    ) -> GaborNet:
        """Construct a ``GaborNet`` with ``depth`` Gabor filters and readouts.

        Examples:
            >>> import jax.numpy as jnp, jax.random as jr
            >>> from geonnax.mfn import GaborNet
            >>> net = GaborNet.init(2, 16, 1, depth=3, key=jr.PRNGKey(0))
            >>> net(jnp.zeros(2)).shape  # (2,) -> (1,)
            (1,)
        """
        if depth < 1:
            raise ValueError(f"depth must be at least 1, got {depth}.")
        keys = jax.random.split(key, 2 * depth)
        filter_keys = keys[:depth]
        linear_keys = keys[depth:]
        filters = [
            GaborFilter.init(
                in_features,
                hidden_features,
                key=filter_keys[i],
                domain=domain,
                gamma_alpha=gamma_alpha,
                gamma_beta=gamma_beta,
            )
            for i in range(depth)
        ]
        linears = [
            eqx.nn.Linear(
                hidden_features,
                hidden_features if i < depth - 1 else out_features,
                key=linear_keys[i],
            )
            for i in range(depth)
        ]
        return cls(
            filters=filters,
            linears=linears,
            in_features=in_features,
            hidden_features=hidden_features,
            out_features=out_features,
            depth=depth,
            domain=domain,
            gamma_alpha=gamma_alpha,
            gamma_beta=gamma_beta,
        )

    def __call__(self, x: Float[Array, " D"]) -> Float[Array, " O"]:
        """Run the multiplicative Gabor-filter forward pass.

        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.mfn import GaborNet
            >>> net = GaborNet.init(2, 8, 4, depth=2, key=jr.PRNGKey(0))
            >>> net(jnp.ones(2)).shape  # (2,) -> (4,)
            (4,)
        """
        # (D,) -> (O,) via multiplicative filter chaining (see mfn_forward).
        return mfn_forward(x, self.filters, self.linears)

init(in_features: int, hidden_features: int, out_features: int, *, depth: int, key: PRNGKeyArray, domain: tuple[float, float] = (-1.0, 1.0), gamma_alpha: float = 6.0, gamma_beta: float = 1.0) -> GaborNet classmethod

Construct a GaborNet with depth Gabor filters and readouts.

Examples:

>>> import jax.numpy as jnp, jax.random as jr
>>> from geonnax.mfn import GaborNet
>>> net = GaborNet.init(2, 16, 1, depth=3, key=jr.PRNGKey(0))
>>> net(jnp.zeros(2)).shape  # (2,) -> (1,)
(1,)
Source code in .venv/lib/python3.12/site-packages/geonnax/mfn.py
@classmethod
def init(
    cls,
    in_features: int,
    hidden_features: int,
    out_features: int,
    *,
    depth: int,
    key: PRNGKeyArray,
    domain: tuple[float, float] = (-1.0, 1.0),
    gamma_alpha: float = 6.0,
    gamma_beta: float = 1.0,
) -> GaborNet:
    """Construct a ``GaborNet`` with ``depth`` Gabor filters and readouts.

    Examples:
        >>> import jax.numpy as jnp, jax.random as jr
        >>> from geonnax.mfn import GaborNet
        >>> net = GaborNet.init(2, 16, 1, depth=3, key=jr.PRNGKey(0))
        >>> net(jnp.zeros(2)).shape  # (2,) -> (1,)
        (1,)
    """
    if depth < 1:
        raise ValueError(f"depth must be at least 1, got {depth}.")
    keys = jax.random.split(key, 2 * depth)
    filter_keys = keys[:depth]
    linear_keys = keys[depth:]
    filters = [
        GaborFilter.init(
            in_features,
            hidden_features,
            key=filter_keys[i],
            domain=domain,
            gamma_alpha=gamma_alpha,
            gamma_beta=gamma_beta,
        )
        for i in range(depth)
    ]
    linears = [
        eqx.nn.Linear(
            hidden_features,
            hidden_features if i < depth - 1 else out_features,
            key=linear_keys[i],
        )
        for i in range(depth)
    ]
    return cls(
        filters=filters,
        linears=linears,
        in_features=in_features,
        hidden_features=hidden_features,
        out_features=out_features,
        depth=depth,
        domain=domain,
        gamma_alpha=gamma_alpha,
        gamma_beta=gamma_beta,
    )

Bayesian variants

BayesianFourierNet

Bases: PyroxModule

FourierNet with Bayesian priors on all filter and linear weights.

A thin subclass of FourierNet that overrides __call__ to register NumPyro sample sites for every parameter:

  • Per filter i: filter_{i}.Omega and filter_{i}.phi.
  • Per linear i: linear_{i}.W and linear_{i}.b.

Total number of sites: \(4L\) where \(L\) is depth.

Priors:

  • \(\Omega_i \sim \mathcal{N}(0, \sigma^2)\) (matrix).
  • \(\varphi_i \sim \mathrm{Uniform}(-\pi, \pi)\).
  • \(W_i \sim \mathcal{N}(0, \sigma^2)\) (matrix).
  • \(b_i \sim \mathcal{N}(0, \sigma^2)\) (vector).

Attributes:

Name Type Description
prior_std float

Prior standard deviation \(\\sigma\) for Gaussian sites (default 1.0). Phase sites always use \(\mathrm{Uniform}(-\pi, \pi)\).

Source code in packages/pyrox-nn/src/pyrox_nn/_mfn.py
class BayesianFourierNet(PyroxModule):
    r"""FourierNet with Bayesian priors on all filter and linear weights.

    A thin subclass of `FourierNet` that overrides ``__call__`` to
    register NumPyro sample sites for every parameter:

    - Per filter *i*: ``filter_{i}.Omega`` and ``filter_{i}.phi``.
    - Per linear *i*: ``linear_{i}.W`` and ``linear_{i}.b``.

    Total number of sites: $4L$ where $L$ is ``depth``.

    Priors:

    - $\Omega_i \sim \mathcal{N}(0, \sigma^2)$ (matrix).
    - $\varphi_i \sim \mathrm{Uniform}(-\pi, \pi)$.
    - $W_i \sim \mathcal{N}(0, \sigma^2)$ (matrix).
    - $b_i \sim \mathcal{N}(0, \sigma^2)$ (vector).

    Attributes:
        prior_std: Prior standard deviation $\\sigma$ for Gaussian
            sites (default 1.0).  Phase sites always use
            $\mathrm{Uniform}(-\pi, \pi)$.
    """

    core: FourierNet
    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,
        key: PRNGKeyArray,
        freq_scale: float = 256.0,
        prior_std: float = 1.0,
        pyrox_name: str | None = None,
    ) -> BayesianFourierNet:
        """Construct a `BayesianFourierNet`.

        Args mirror `FourierNet.init`, plus:

        Args:
            prior_std: Prior standard deviation for Gaussian sites
                (default 1.0).

        Raises:
            ValueError: If ``prior_std`` is non-positive or any
                `FourierNet.init` validation fails.
        """
        _require_positive(prior_std=prior_std)
        core = FourierNet.init(
            in_features,
            hidden_features,
            out_features,
            depth=depth,
            key=key,
            freq_scale=freq_scale,
        )
        return cls(core=core, prior_std=prior_std, pyrox_name=pyrox_name)

    # Convenience accessors so downstream code that reads structural fields
    # off the wrapper (in_features/depth/etc.) keeps working transparently.
    @property
    def filters(self) -> list[FourierFilter]:
        return self.core.filters

    @property
    def linears(self) -> list[eqx.nn.Linear]:
        return self.core.linears

    @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 depth(self) -> int:
        return self.core.depth

    @pyrox_method
    def __call__(self, x: Float[Array, "N D"]) -> Float[Array, "N O"]:
        """Forward pass with sampled parameters.

        Args:
            x: Input of shape ``(N, D)`` or ``(D,)`` (single point).

        Returns:
            Output of shape ``(N, O)`` or ``(O,)`` if input was 1-D.
        """
        squeeze = x.ndim == 1
        x2d = jnp.atleast_2d(x)

        sampled_filters: list[FourierFilter] = []
        for i, f in enumerate(self.core.filters):
            Omega, phi = _sample_filter_omega_phi(self, i, f, self.prior_std)
            sampled_filters.append(
                eqx.tree_at(lambda ff: (ff.Omega, ff.phi), f, (Omega, phi))
            )

        sampled_linears = _sample_normal_linears(
            self, self.core.linears, self.prior_std
        )

        sampled_core = eqx.tree_at(
            lambda c: (c.filters, c.linears),
            self.core,
            (sampled_filters, sampled_linears),
        )
        out = jax.vmap(sampled_core)(x2d)
        return out[0] if squeeze else out

init(in_features: int, hidden_features: int, out_features: int, *, depth: int, key: PRNGKeyArray, freq_scale: float = 256.0, prior_std: float = 1.0, pyrox_name: str | None = None) -> BayesianFourierNet classmethod

Construct a BayesianFourierNet.

Args mirror FourierNet.init, plus:

Parameters:

Name Type Description Default
prior_std float

Prior standard deviation for Gaussian sites (default 1.0).

1.0

Raises:

Type Description
ValueError

If prior_std is non-positive or any FourierNet.init validation fails.

Source code in packages/pyrox-nn/src/pyrox_nn/_mfn.py
@classmethod
def init(
    cls,
    in_features: int,
    hidden_features: int,
    out_features: int,
    *,
    depth: int,
    key: PRNGKeyArray,
    freq_scale: float = 256.0,
    prior_std: float = 1.0,
    pyrox_name: str | None = None,
) -> BayesianFourierNet:
    """Construct a `BayesianFourierNet`.

    Args mirror `FourierNet.init`, plus:

    Args:
        prior_std: Prior standard deviation for Gaussian sites
            (default 1.0).

    Raises:
        ValueError: If ``prior_std`` is non-positive or any
            `FourierNet.init` validation fails.
    """
    _require_positive(prior_std=prior_std)
    core = FourierNet.init(
        in_features,
        hidden_features,
        out_features,
        depth=depth,
        key=key,
        freq_scale=freq_scale,
    )
    return cls(core=core, prior_std=prior_std, pyrox_name=pyrox_name)

BayesianGaborNet

Bases: PyroxModule

GaborNet with Bayesian priors on all filter and linear weights.

A thin subclass of GaborNet that overrides __call__ to register NumPyro sample sites for every parameter:

  • Per filter i: filter_{i}.Omega, filter_{i}.phi, filter_{i}.mu, and filter_{i}.log_gamma.
  • Per linear i: linear_{i}.W and linear_{i}.b.

Total number of sites: \(6L\) where \(L\) is depth.

Priors:

  • \(\Omega_i \sim \mathcal{N}(0, \sigma^2)\) (matrix).
  • \(\varphi_i \sim \mathrm{Uniform}(-\pi, \pi)\).
  • \(\mu_i \sim \mathrm{Uniform}(\texttt{domain\_low},\texttt{domain\_high})\).
  • \(\log\gamma_i \sim \mathcal{N}(0, \sigma^2)\) (log-space).
  • \(W_i \sim \mathcal{N}(0, \sigma^2)\) (matrix).
  • \(b_i \sim \mathcal{N}(0, \sigma^2)\) (vector).

Attributes:

Name Type Description
prior_std float

Prior standard deviation \(\\sigma\) for Gaussian and log-gamma sites (default 1.0).

Source code in packages/pyrox-nn/src/pyrox_nn/_mfn.py
class BayesianGaborNet(PyroxModule):
    r"""GaborNet with Bayesian priors on all filter and linear weights.

    A thin subclass of `GaborNet` that overrides ``__call__`` to
    register NumPyro sample sites for every parameter:

    - Per filter *i*: ``filter_{i}.Omega``, ``filter_{i}.phi``,
      ``filter_{i}.mu``, and ``filter_{i}.log_gamma``.
    - Per linear *i*: ``linear_{i}.W`` and ``linear_{i}.b``.

    Total number of sites: $6L$ where $L$ is ``depth``.

    Priors:

    - $\Omega_i \sim \mathcal{N}(0, \sigma^2)$ (matrix).
    - $\varphi_i \sim \mathrm{Uniform}(-\pi, \pi)$.
    - $\mu_i \sim \mathrm{Uniform}(\texttt{domain\_low},\texttt{domain\_high})$.
    - $\log\gamma_i \sim \mathcal{N}(0, \sigma^2)$ (log-space).
    - $W_i \sim \mathcal{N}(0, \sigma^2)$ (matrix).
    - $b_i \sim \mathcal{N}(0, \sigma^2)$ (vector).

    Attributes:
        prior_std: Prior standard deviation $\\sigma$ for Gaussian
            and log-gamma sites (default 1.0).
    """

    core: GaborNet
    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,
        key: PRNGKeyArray,
        domain: tuple[float, float] = (-1.0, 1.0),
        gamma_alpha: float = 6.0,
        gamma_beta: float = 1.0,
        prior_std: float = 1.0,
        pyrox_name: str | None = None,
    ) -> BayesianGaborNet:
        """Construct a `BayesianGaborNet`.

        Args mirror `GaborNet.init`, plus:

        Args:
            prior_std: Prior standard deviation for Gaussian and
                log-gamma sites (default 1.0).

        Raises:
            ValueError: If ``prior_std`` is non-positive or any
                `GaborNet.init` validation fails.
        """
        _require_positive(prior_std=prior_std)
        core = GaborNet.init(
            in_features,
            hidden_features,
            out_features,
            depth=depth,
            key=key,
            domain=domain,
            gamma_alpha=gamma_alpha,
            gamma_beta=gamma_beta,
        )
        return cls(core=core, prior_std=prior_std, pyrox_name=pyrox_name)

    # Structural accessors mirror the previous direct-inheritance API.
    @property
    def filters(self) -> list[GaborFilter]:
        return self.core.filters

    @property
    def linears(self) -> list[eqx.nn.Linear]:
        return self.core.linears

    @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 depth(self) -> int:
        return self.core.depth

    @property
    def domain(self) -> tuple[float, float]:
        return self.core.domain

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

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

    @pyrox_method
    def __call__(self, x: Float[Array, "N D"]) -> Float[Array, "N O"]:
        """Forward pass with sampled parameters.

        Args:
            x: Input of shape ``(N, D)`` or ``(D,)`` (single point).

        Returns:
            Output of shape ``(N, O)`` or ``(O,)`` if input was 1-D.
        """
        squeeze = x.ndim == 1
        x2d = jnp.atleast_2d(x)
        low, high = self.core.domain

        sampled_filters: list[GaborFilter] = []
        for i, f in enumerate(self.core.filters):
            Omega, phi = _sample_filter_omega_phi(self, i, f, self.prior_std)
            mu = self.pyrox_sample(
                f"filter_{i}.mu",
                dist.Uniform(low, high)
                .expand([f.out_features, f.in_features])
                .to_event(2),
            )
            log_gamma = self.pyrox_sample(
                f"filter_{i}.log_gamma",
                dist.Normal(0.0, self.prior_std).expand([f.out_features]).to_event(1),
            )
            sampled_filters.append(
                eqx.tree_at(
                    lambda ff: (ff.Omega, ff.phi, ff.mu, ff.log_gamma),
                    f,
                    (Omega, phi, mu, log_gamma),
                )
            )

        sampled_linears = _sample_normal_linears(
            self, self.core.linears, self.prior_std
        )

        sampled_core = eqx.tree_at(
            lambda c: (c.filters, c.linears),
            self.core,
            (sampled_filters, sampled_linears),
        )
        out = jax.vmap(sampled_core)(x2d)
        return out[0] if squeeze else out

init(in_features: int, hidden_features: int, out_features: int, *, depth: int, key: PRNGKeyArray, domain: tuple[float, float] = (-1.0, 1.0), gamma_alpha: float = 6.0, gamma_beta: float = 1.0, prior_std: float = 1.0, pyrox_name: str | None = None) -> BayesianGaborNet classmethod

Construct a BayesianGaborNet.

Args mirror GaborNet.init, plus:

Parameters:

Name Type Description Default
prior_std float

Prior standard deviation for Gaussian and log-gamma sites (default 1.0).

1.0

Raises:

Type Description
ValueError

If prior_std is non-positive or any GaborNet.init validation fails.

Source code in packages/pyrox-nn/src/pyrox_nn/_mfn.py
@classmethod
def init(
    cls,
    in_features: int,
    hidden_features: int,
    out_features: int,
    *,
    depth: int,
    key: PRNGKeyArray,
    domain: tuple[float, float] = (-1.0, 1.0),
    gamma_alpha: float = 6.0,
    gamma_beta: float = 1.0,
    prior_std: float = 1.0,
    pyrox_name: str | None = None,
) -> BayesianGaborNet:
    """Construct a `BayesianGaborNet`.

    Args mirror `GaborNet.init`, plus:

    Args:
        prior_std: Prior standard deviation for Gaussian and
            log-gamma sites (default 1.0).

    Raises:
        ValueError: If ``prior_std`` is non-positive or any
            `GaborNet.init` validation fails.
    """
    _require_positive(prior_std=prior_std)
    core = GaborNet.init(
        in_features,
        hidden_features,
        out_features,
        depth=depth,
        key=key,
        domain=domain,
        gamma_alpha=gamma_alpha,
        gamma_beta=gamma_beta,
    )
    return cls(core=core, prior_std=prior_std, pyrox_name=pyrox_name)

Pure-JAX helper

mfn_forward(x: Float[Array, ' D'], filters: Sequence[Callable[[JaxArray], JaxArray]], linears: Sequence[Callable[[JaxArray], JaxArray]]) -> Float[Array, ' O']

Pure-JAX MFN forward pass given user-supplied filter and linear callables.

Implements the Fathony et al. (2021) multiplicative chaining:

\[ z_1 = g_1(x), \quad z_{i+1} = g_{i+1}(x) \odot (W_i z_i + b_i), \quad y = W_L z_L + b_L. \]

Exists as an escape hatch so users can plug custom filter families into the MFN topology without subclassing FourierNet or GaborNet.

filters and linears must have the same length \(L\). x is a single example of shape (in_features,); use jax.vmap for batched application.

Examples:

>>> import jax.numpy as jnp, jax.random as jr
>>> from geonnax.mfn import FourierNet, mfn_forward
>>> net = FourierNet.init(2, 4, 3, depth=2, key=jr.PRNGKey(0))
>>> mfn_forward(jnp.zeros(2), net.filters, net.linears).shape
(3,)
Source code in .venv/lib/python3.12/site-packages/geonnax/mfn.py
def mfn_forward(
    x: Float[Array, " D"],
    filters: Sequence[Callable[[JaxArray], JaxArray]],
    linears: Sequence[Callable[[JaxArray], JaxArray]],
) -> Float[Array, " O"]:
    """Pure-JAX MFN forward pass given user-supplied filter and linear callables.

    Implements the Fathony et al. (2021) multiplicative chaining:

    $$
    z_1 = g_1(x), \\quad
    z_{i+1} = g_{i+1}(x) \\odot (W_i z_i + b_i), \\quad
    y = W_L z_L + b_L.
    $$


    Exists as an escape hatch so users can plug custom filter families
    into the MFN topology without subclassing `FourierNet` or
    `GaborNet`.

    ``filters`` and ``linears`` must have the same length $L$.
    ``x`` is a single example of shape ``(in_features,)``; use
    `jax.vmap` for batched application.

    Examples:
        >>> import jax.numpy as jnp, jax.random as jr
        >>> from geonnax.mfn import FourierNet, mfn_forward
        >>> net = FourierNet.init(2, 4, 3, depth=2, key=jr.PRNGKey(0))
        >>> mfn_forward(jnp.zeros(2), net.filters, net.linears).shape
        (3,)
    """
    if len(filters) == 0 or len(linears) == 0:
        raise ValueError(
            f"filters and linears must be non-empty; got lengths "
            f"{len(filters)} and {len(linears)}."
        )
    if len(filters) != len(linears):
        raise ValueError(
            f"filters and linears must have equal length; got "
            f"{len(filters)} and {len(linears)}."
        )
    z = filters[0](x)  # z_1 = g_1(x): (D,) -> (H,)
    for f, lin in zip(filters[1:], linears[:-1], strict=True):
        # z_{i+1} = g_{i+1}(x) ⊙ (W_i z_i + b_i): (H,) -> (H,).
        z = f(x) * lin(z)
    return linears[-1](z)  # readout y = W_L z_L + b_L: (H,) -> (O,)