Skip to content

Geo encoders

Geophysical inputs usually arrive as longitude/latitude in degrees, while downstream neural-field and GP features typically want periodic encodings or unit-sphere coordinates. The geo encoders in pyrox_nn make those preprocessing steps first-class and composable.

The canonical spherical-harmonic pipeline is:

import equinox as eqx

from pyrox_nn import (
    Cartesian3DEncoder,
    Deg2Rad,
    SphericalHarmonicEncoder,
)

encoder = eqx.nn.Sequential(
    [
        Deg2Rad(),
        Cartesian3DEncoder(input_unit="radians"),
        SphericalHarmonicEncoder(l_max=8, input_mode="cartesian"),
    ]
)
features = encoder(lonlat_deg)  # (N, 81)

Cartesian3DEncoder uses the same axis convention expected by pyrox_gp.SphericalHarmonicInducingFeatures, so the NN and GP spherical paths line up. For temporal complements, see fourier_features and seasonal_features.

Stateful encoder layers

Deg2Rad

Bases: Module

Element-wise degrees-to-radians conversion.

Stateless wrapper around geonnax.geo.deg2rad — no learnable parameters.

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.encoders import Deg2Rad
>>> Deg2Rad()(jnp.array([0.0, 90.0, 180.0])).shape
(3,)
Source code in .venv/lib/python3.12/site-packages/geonnax/encoders.py
class Deg2Rad(eqx.Module):
    """Element-wise degrees-to-radians conversion.

    Stateless wrapper around `geonnax.geo.deg2rad` — no learnable
    parameters.

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.encoders import Deg2Rad
        >>> Deg2Rad()(jnp.array([0.0, 90.0, 180.0])).shape
        (3,)
    """

    def __call__(self, x: Float[Array, ...]) -> Float[Array, ...]:
        """Convert ``x`` from degrees to radians, preserving its shape.

        Examples:
            >>> import jax.numpy as jnp
            >>> from geonnax.encoders import Deg2Rad
            >>> Deg2Rad()(jnp.zeros((4,))).shape
            (4,)
        """
        return deg2rad(x)

LonLatScale

Bases: Module

Affine-rescale a single lon/lat pair into [-1, 1].

Values inside the given ranges map into [-1, 1]; out-of-range values are not clipped. The default ranges assume lonlat is in degrees.

Attributes:

Name Type Description
lon_range tuple[float, float]

(min, max) longitude domain (must satisfy min < max).

lat_range tuple[float, float]

(min, max) latitude domain (must satisfy min < max).

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.encoders import LonLatScale
>>> LonLatScale()(jnp.array([180.0, 90.0])).shape
(2,)
Source code in .venv/lib/python3.12/site-packages/geonnax/encoders.py
class LonLatScale(eqx.Module):
    """Affine-rescale a single lon/lat pair into ``[-1, 1]``.

    Values inside the given ranges map into ``[-1, 1]``; out-of-range
    values are *not* clipped. The default ranges assume ``lonlat`` is
    in degrees.

    Attributes:
        lon_range: ``(min, max)`` longitude domain (must satisfy
            ``min < max``).
        lat_range: ``(min, max)`` latitude domain (must satisfy
            ``min < max``).

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.encoders import LonLatScale
        >>> LonLatScale()(jnp.array([180.0, 90.0])).shape
        (2,)
    """

    lon_range: tuple[float, float] = eqx.field(static=True, default=(-180.0, 180.0))
    lat_range: tuple[float, float] = eqx.field(static=True, default=(-90.0, 90.0))

    def __post_init__(self) -> None:
        _validate_range(self.lon_range, name="lon_range")
        _validate_range(self.lat_range, name="lat_range")

    def __call__(self, lonlat: Num[Array, " 2"]) -> Float[Array, " 2"]:
        """Rescale a single ``(2,)`` lon/lat pair into ``[-1, 1]``.

        Examples:
            >>> import jax.numpy as jnp
            >>> from geonnax.encoders import LonLatScale
            >>> # Domain max maps to +1 in each column.
            >>> out = LonLatScale()(jnp.array([180.0, 90.0]))
            >>> bool(jnp.allclose(out, 1.0))
            True
        """
        # (2,) -> (1, 2) so the batched helper applies, then drop the batch dim.
        out = lonlat_scale(
            lonlat[None, :],
            lon_range=self.lon_range,
            lat_range=self.lat_range,
        )
        return out[0]  # (1, 2) -> (2,)

Cartesian3DEncoder

Bases: Module

Lift a single lon/lat coordinate onto the unit sphere \(S^2\).

Stateless wrapper around geonnax.geo.lonlat_to_cartesian3d.

Attributes:

Name Type Description
input_unit Literal['degrees', 'radians']

Whether the input is in "degrees" or "radians".

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.encoders import Cartesian3DEncoder
>>> Cartesian3DEncoder()(jnp.array([0.0, 0.0])).shape
(3,)
Source code in .venv/lib/python3.12/site-packages/geonnax/encoders.py
class Cartesian3DEncoder(eqx.Module):
    """Lift a single lon/lat coordinate onto the unit sphere $S^2$.

    Stateless wrapper around `geonnax.geo.lonlat_to_cartesian3d`.

    Attributes:
        input_unit: Whether the input is in ``"degrees"`` or
            ``"radians"``.

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.encoders import Cartesian3DEncoder
        >>> Cartesian3DEncoder()(jnp.array([0.0, 0.0])).shape
        (3,)
    """

    input_unit: Literal["degrees", "radians"] = eqx.field(
        static=True, default="radians"
    )

    def __post_init__(self) -> None:
        _validate_input_unit(self.input_unit)

    def __call__(self, lonlat: Float[Array, " 2"]) -> Float[Array, " 3"]:
        """Map a single ``(2,)`` lon/lat pair to a ``(3,)`` unit vector.

        Examples:
            >>> import jax.numpy as jnp
            >>> from geonnax.encoders import Cartesian3DEncoder
            >>> # (lon, lat) = (0, 0) → +x on the unit sphere.
            >>> out = Cartesian3DEncoder()(jnp.array([0.0, 0.0]))
            >>> bool(jnp.allclose(out, jnp.array([1.0, 0.0, 0.0]), atol=1e-6))
            True
        """
        # (2,) -> (1, 2) for the batched helper, then drop the batch dim.
        out = lonlat_to_cartesian3d(lonlat[None, :], input_unit=self.input_unit)
        return out[0]  # (1, 3) -> (3,)

CyclicEncoder

Bases: Module

Encode a single periodic input as concatenated cos/sin features.

Stateless wrapper around geonnax.geo.cyclic_encode.

Examples:

>>> import jax.numpy as jnp
>>> import jax
>>> from geonnax.encoders import CyclicEncoder
>>> jax.vmap(CyclicEncoder())(jnp.array([0.0, jnp.pi])).shape
(2, 2)
Source code in .venv/lib/python3.12/site-packages/geonnax/encoders.py
class CyclicEncoder(eqx.Module):
    """Encode a single periodic input as concatenated cos/sin features.

    Stateless wrapper around `geonnax.geo.cyclic_encode`.

    Examples:
        >>> import jax.numpy as jnp
        >>> import jax
        >>> from geonnax.encoders import CyclicEncoder
        >>> jax.vmap(CyclicEncoder())(jnp.array([0.0, jnp.pi])).shape
        (2, 2)
    """

    def __call__(
        self,
        angles: Float[Array, ""] | Float[Array, " D"],
    ) -> Float[Array, " F"]:
        """Encode a scalar or ``(D,)`` angle vector into ``[cos, sin]`` features.

        Examples:
            >>> import jax.numpy as jnp
            >>> from geonnax.encoders import CyclicEncoder
            >>> # Scalar angle → 2 features (cos, sin).
            >>> CyclicEncoder()(jnp.array(0.0)).shape
            (2,)
            >>> # (D,) angle vector → 2·D features.
            >>> CyclicEncoder()(jnp.zeros((3,))).shape
            (6,)
        """
        promoted = jnp.atleast_1d(angles)  # () or (D,) -> (D,)
        out = cyclic_encode(promoted[None, :])  # (D,) -> (1, D) -> (1, 2·D)
        return out[0]  # (1, 2·D) -> (2·D,)

SphericalHarmonicEncoder

Bases: Module

Real spherical-harmonic features on the unit 2-sphere for a single point.

Stateless wrapper that evaluates geonnax._basis.real_spherical_harmonics on either an already- cartesian input (input_mode='cartesian') or a lon/lat pair (input_mode='lonlat', assumed in radians).

Attributes:

Name Type Description
l_max int

Maximum harmonic degree (must be >= 0). The output has (l_max + 1) ** 2 features.

input_mode Literal['cartesian', 'lonlat']

"cartesian" for a (3,) unit-sphere input or "lonlat" for a (2,) lon/lat pair in radians.

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.encoders import SphericalHarmonicEncoder
>>> # l_max=3 → (l_max + 1)^2 = 16 features for a (3,) cartesian input.
>>> SphericalHarmonicEncoder(l_max=3)(jnp.array([1.0, 0.0, 0.0])).shape
(16,)
Source code in .venv/lib/python3.12/site-packages/geonnax/encoders.py
class SphericalHarmonicEncoder(eqx.Module):
    """Real spherical-harmonic features on the unit 2-sphere for a single point.

    Stateless wrapper that evaluates
    `geonnax._basis.real_spherical_harmonics` on either an already-
    cartesian input (``input_mode='cartesian'``) or a lon/lat pair
    (``input_mode='lonlat'``, assumed in radians).

    Attributes:
        l_max: Maximum harmonic degree (must be ``>= 0``). The output
            has ``(l_max + 1) ** 2`` features.
        input_mode: ``"cartesian"`` for a ``(3,)`` unit-sphere input
            or ``"lonlat"`` for a ``(2,)`` lon/lat pair in radians.

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.encoders import SphericalHarmonicEncoder
        >>> # l_max=3 → (l_max + 1)^2 = 16 features for a (3,) cartesian input.
        >>> SphericalHarmonicEncoder(l_max=3)(jnp.array([1.0, 0.0, 0.0])).shape
        (16,)
    """

    l_max: int = eqx.field(static=True)
    input_mode: Literal["cartesian", "lonlat"] = eqx.field(
        static=True, default="cartesian"
    )

    def __post_init__(self) -> None:
        if self.l_max < 0:
            raise ValueError(f"l_max must be >= 0; got {self.l_max}.")
        if self.input_mode not in {"cartesian", "lonlat"}:
            raise ValueError(
                f"input_mode must be 'cartesian' or 'lonlat'; got {self.input_mode!r}."
            )

    @property
    def num_features(self) -> int:
        """Number of harmonic features, ``(l_max + 1) ** 2``.

        Examples:
            >>> from geonnax.encoders import SphericalHarmonicEncoder
            >>> SphericalHarmonicEncoder(l_max=2).num_features
            9
        """
        return (self.l_max + 1) ** 2  # degrees 0..l_max give Σ(2l+1) = (l_max+1)^2

    def __call__(
        self,
        x: Float[Array, " 3"] | Float[Array, " 2"],
    ) -> Float[Array, " M"]:
        """Evaluate real spherical harmonics at a single point.

        Examples:
            >>> import jax.numpy as jnp
            >>> from geonnax.encoders import SphericalHarmonicEncoder
            >>> # lonlat mode: (2,) radian pair → (l_max + 1)^2 features.
            >>> enc = SphericalHarmonicEncoder(l_max=2, input_mode="lonlat")
            >>> enc(jnp.array([0.0, 0.0])).shape
            (9,)
        """
        if self.input_mode == "cartesian":
            if x.ndim != 1 or x.shape[-1] != 3:
                raise ValueError(
                    f"x must be (3,) when input_mode='cartesian'; got shape {x.shape}."
                )
            unit_xyz = x[None, :]  # (3,) -> (1, 3)
        else:
            if x.ndim != 1 or x.shape[-1] != 2:
                raise ValueError(
                    f"x must be (2,) when input_mode='lonlat'; got shape {x.shape}."
                )
            # (2,) lon/lat radians -> (1, 3) on the unit sphere.
            unit_xyz = lonlat_to_cartesian3d(x[None, :], input_unit="radians")
        # (1, 3) -> (1, (l_max+1)^2) -> drop batch dim -> ((l_max+1)^2,)
        return real_spherical_harmonics(unit_xyz, l_max=self.l_max)[0]

num_features: int property

Number of harmonic features, (l_max + 1) ** 2.

Examples:

>>> from geonnax.encoders import SphericalHarmonicEncoder
>>> SphericalHarmonicEncoder(l_max=2).num_features
9

Slepian encoders

Region-localized spherical encoders built on the Slepian concentration problem. The deterministic SlepianEncoder and HybridSphericalSlepianEncoder are re-exported from geonnax; BayesianSlepianEncoder adds NumPyro sites over the cap radius and centre.

SlepianEncoder

Bases: Module

Deterministic Slepian-cap positional encoder on \(S^2\).

Slepian functions are the band-limited modes maximally concentrated inside a spherical cap; each is a weighted sum of spherical harmonics up to l_max with concentration eigenvalue λ ∈ [0, 1].

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.slepian import SlepianEncoder
>>> enc = SlepianEncoder.from_cap(
...     l_max=3,
...     cap_radius_deg=20.0,
...     cap_centre_lonlat_deg=(0.0, 0.0),
...     eig_threshold=0.0,
...     n_modes=4,
... )
>>> enc(jnp.array([0.0, 0.0])).shape == (enc.num_features,)
True
Source code in .venv/lib/python3.12/site-packages/geonnax/slepian.py
class SlepianEncoder(eqx.Module):
    """Deterministic Slepian-cap positional encoder on $S^2$.

    Slepian functions are the band-limited modes maximally concentrated
    inside a spherical cap; each is a weighted sum of spherical harmonics
    up to ``l_max`` with concentration eigenvalue ``λ ∈ [0, 1]``.

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.slepian import SlepianEncoder
        >>> enc = SlepianEncoder.from_cap(
        ...     l_max=3,
        ...     cap_radius_deg=20.0,
        ...     cap_centre_lonlat_deg=(0.0, 0.0),
        ...     eig_threshold=0.0,
        ...     n_modes=4,
        ... )
        >>> enc(jnp.array([0.0, 0.0])).shape == (enc.num_features,)
        True
    """

    basis: SlepianCapBasis
    input_mode: Literal["cartesian", "lonlat"] = eqx.field(
        static=True, default="lonlat"
    )
    weight_by_eigenvalue: bool = eqx.field(static=True, default=True)

    def __post_init__(self) -> None:
        if self.input_mode not in {"cartesian", "lonlat"}:
            raise ValueError(
                f"input_mode must be 'cartesian' or 'lonlat'; got {self.input_mode!r}."
            )

    @classmethod
    def from_cap(
        cls,
        *,
        l_max: int,
        cap_radius_deg: float,
        cap_centre_lonlat_deg: tuple[float, float],
        eig_threshold: float = 0.05,
        n_modes: int | None = None,
        input_mode: Literal["cartesian", "lonlat"] = "lonlat",
        weight_by_eigenvalue: bool = True,
    ) -> SlepianEncoder:
        """Construct an encoder from cap geometry specified in degrees.

        Examples:
            >>> from geonnax.slepian import SlepianEncoder
            >>> enc = SlepianEncoder.from_cap(
            ...     l_max=3,
            ...     cap_radius_deg=20.0,
            ...     cap_centre_lonlat_deg=(0.0, 0.0),
            ...     eig_threshold=0.0,
            ...     n_modes=4,
            ... )
            >>> enc.num_features
            4
        """
        # Degrees in the public API → radians for the basis solver.
        basis = slepian_cap_basis(
            l_max=l_max,
            cap_radius=math.radians(cap_radius_deg),
            n_modes=n_modes,
            eig_threshold=eig_threshold,
            lonlat_centre=jnp.radians(jnp.asarray(cap_centre_lonlat_deg)),
        )
        return cls(
            basis=basis,
            input_mode=input_mode,
            weight_by_eigenvalue=weight_by_eigenvalue,
        )

    @property
    def num_features(self) -> int:
        """Number of output features (retained Slepian modes).

        Examples:
            >>> from geonnax.slepian import SlepianEncoder
            >>> enc = SlepianEncoder.from_cap(
            ...     l_max=3,
            ...     cap_radius_deg=20.0,
            ...     cap_centre_lonlat_deg=(0.0, 0.0),
            ...     eig_threshold=0.0,
            ...     n_modes=4,
            ... )
            >>> enc.num_features
            4
        """
        return self.basis.num_modes

    def __call__(
        self,
        x: Float[Array, " 3"] | Float[Array, " 2"],
    ) -> Float[Array, " K"]:
        """Evaluate the ``K`` retained Slepian modes at a single point.

        Examples:
            >>> import jax.numpy as jnp
            >>> from geonnax.slepian import SlepianEncoder
            >>> enc = SlepianEncoder.from_cap(
            ...     l_max=3,
            ...     cap_radius_deg=20.0,
            ...     cap_centre_lonlat_deg=(0.0, 0.0),
            ...     eig_threshold=0.0,
            ...     n_modes=4,
            ... )
            >>> enc(jnp.array([0.0, 0.0])).shape
            (4,)
        """
        unit_xyz = _unit_xyz(x, self.input_mode)  # (2,)/(3,) -> (3,)
        # (3,) -> (1, 3) -> evaluate K modes -> (1, K) -> drop batch -> (K,)
        features = self.basis.evaluate(unit_xyz[None, :])[0]
        if self.weight_by_eigenvalue:
            # Scale mode k by √λ_k so well-concentrated modes dominate.
            features = features * jnp.sqrt(self.basis.eigenvalues)
        return features

num_features: int property

Number of output features (retained Slepian modes).

Examples:

>>> from geonnax.slepian import SlepianEncoder
>>> enc = SlepianEncoder.from_cap(
...     l_max=3,
...     cap_radius_deg=20.0,
...     cap_centre_lonlat_deg=(0.0, 0.0),
...     eig_threshold=0.0,
...     n_modes=4,
... )
>>> enc.num_features
4

from_cap(*, l_max: int, cap_radius_deg: float, cap_centre_lonlat_deg: tuple[float, float], eig_threshold: float = 0.05, n_modes: int | None = None, input_mode: Literal['cartesian', 'lonlat'] = 'lonlat', weight_by_eigenvalue: bool = True) -> SlepianEncoder classmethod

Construct an encoder from cap geometry specified in degrees.

Examples:

>>> from geonnax.slepian import SlepianEncoder
>>> enc = SlepianEncoder.from_cap(
...     l_max=3,
...     cap_radius_deg=20.0,
...     cap_centre_lonlat_deg=(0.0, 0.0),
...     eig_threshold=0.0,
...     n_modes=4,
... )
>>> enc.num_features
4
Source code in .venv/lib/python3.12/site-packages/geonnax/slepian.py
@classmethod
def from_cap(
    cls,
    *,
    l_max: int,
    cap_radius_deg: float,
    cap_centre_lonlat_deg: tuple[float, float],
    eig_threshold: float = 0.05,
    n_modes: int | None = None,
    input_mode: Literal["cartesian", "lonlat"] = "lonlat",
    weight_by_eigenvalue: bool = True,
) -> SlepianEncoder:
    """Construct an encoder from cap geometry specified in degrees.

    Examples:
        >>> from geonnax.slepian import SlepianEncoder
        >>> enc = SlepianEncoder.from_cap(
        ...     l_max=3,
        ...     cap_radius_deg=20.0,
        ...     cap_centre_lonlat_deg=(0.0, 0.0),
        ...     eig_threshold=0.0,
        ...     n_modes=4,
        ... )
        >>> enc.num_features
        4
    """
    # Degrees in the public API → radians for the basis solver.
    basis = slepian_cap_basis(
        l_max=l_max,
        cap_radius=math.radians(cap_radius_deg),
        n_modes=n_modes,
        eig_threshold=eig_threshold,
        lonlat_centre=jnp.radians(jnp.asarray(cap_centre_lonlat_deg)),
    )
    return cls(
        basis=basis,
        input_mode=input_mode,
        weight_by_eigenvalue=weight_by_eigenvalue,
    )

HybridSphericalSlepianEncoder

Bases: Module

Concatenate low-bandwidth global SHs with local Slepian cap modes.

Output is [Y_l^m features | Slepian cap features] — global structure from spherical harmonics plus localized detail from the cap.

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.slepian import HybridSphericalSlepianEncoder
>>> enc = HybridSphericalSlepianEncoder.from_cap(
...     sh_l_max=1,
...     slepian_l_max=3,
...     cap_radius_deg=20.0,
...     cap_centre_lonlat_deg=(0.0, 0.0),
...     eig_threshold=0.0,
...     n_modes=2,
... )
>>> enc(jnp.array([0.0, 0.0])).shape == (enc.num_features,)
True
Source code in .venv/lib/python3.12/site-packages/geonnax/slepian.py
class HybridSphericalSlepianEncoder(eqx.Module):
    """Concatenate low-bandwidth global SHs with local Slepian cap modes.

    Output is ``[Y_l^m features | Slepian cap features]`` — global
    structure from spherical harmonics plus localized detail from the cap.

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.slepian import HybridSphericalSlepianEncoder
        >>> enc = HybridSphericalSlepianEncoder.from_cap(
        ...     sh_l_max=1,
        ...     slepian_l_max=3,
        ...     cap_radius_deg=20.0,
        ...     cap_centre_lonlat_deg=(0.0, 0.0),
        ...     eig_threshold=0.0,
        ...     n_modes=2,
        ... )
        >>> enc(jnp.array([0.0, 0.0])).shape == (enc.num_features,)
        True
    """

    sh_encoder: SphericalHarmonicEncoder
    slepian_encoder: SlepianEncoder

    @classmethod
    def from_cap(
        cls,
        *,
        sh_l_max: int,
        slepian_l_max: int,
        cap_radius_deg: float,
        cap_centre_lonlat_deg: tuple[float, float],
        eig_threshold: float = 0.05,
        n_modes: int | None = None,
        input_mode: Literal["cartesian", "lonlat"] = "lonlat",
    ) -> HybridSphericalSlepianEncoder:
        """Construct the hybrid encoder used by Slepian positional encodings.

        Examples:
            >>> from geonnax.slepian import HybridSphericalSlepianEncoder
            >>> enc = HybridSphericalSlepianEncoder.from_cap(
            ...     sh_l_max=1,
            ...     slepian_l_max=3,
            ...     cap_radius_deg=20.0,
            ...     cap_centre_lonlat_deg=(0.0, 0.0),
            ...     eig_threshold=0.0,
            ...     n_modes=2,
            ... )
            >>> # sh_l_max=1 → (1 + 1)^2 = 4 SH features, plus 2 Slepian modes.
            >>> enc.num_features
            6
        """
        return cls(
            sh_encoder=SphericalHarmonicEncoder(l_max=sh_l_max, input_mode=input_mode),
            slepian_encoder=SlepianEncoder.from_cap(
                l_max=slepian_l_max,
                cap_radius_deg=cap_radius_deg,
                cap_centre_lonlat_deg=cap_centre_lonlat_deg,
                eig_threshold=eig_threshold,
                n_modes=n_modes,
                input_mode=input_mode,
            ),
        )

    @property
    def num_features(self) -> int:
        """Number of concatenated output features (SH modes + Slepian modes).

        Examples:
            >>> from geonnax.slepian import HybridSphericalSlepianEncoder
            >>> enc = HybridSphericalSlepianEncoder.from_cap(
            ...     sh_l_max=1,
            ...     slepian_l_max=3,
            ...     cap_radius_deg=20.0,
            ...     cap_centre_lonlat_deg=(0.0, 0.0),
            ...     eig_threshold=0.0,
            ...     n_modes=2,
            ... )
            >>> enc.num_features
            6
        """
        return self.sh_encoder.num_features + self.slepian_encoder.num_features

    def __call__(
        self,
        x: Float[Array, " 3"] | Float[Array, " 2"],
    ) -> Float[Array, " F"]:
        """Concatenate SH and Slepian features for a single point.

        Examples:
            >>> import jax.numpy as jnp
            >>> from geonnax.slepian import HybridSphericalSlepianEncoder
            >>> enc = HybridSphericalSlepianEncoder.from_cap(
            ...     sh_l_max=1,
            ...     slepian_l_max=3,
            ...     cap_radius_deg=20.0,
            ...     cap_centre_lonlat_deg=(0.0, 0.0),
            ...     eig_threshold=0.0,
            ...     n_modes=2,
            ... )
            >>> enc(jnp.array([0.0, 0.0])).shape
            (6,)
        """
        # (num_sh,) ++ (num_slepian,) -> (num_sh + num_slepian,)
        return jnp.concatenate([self.sh_encoder(x), self.slepian_encoder(x)], axis=-1)

num_features: int property

Number of concatenated output features (SH modes + Slepian modes).

Examples:

>>> from geonnax.slepian import HybridSphericalSlepianEncoder
>>> enc = HybridSphericalSlepianEncoder.from_cap(
...     sh_l_max=1,
...     slepian_l_max=3,
...     cap_radius_deg=20.0,
...     cap_centre_lonlat_deg=(0.0, 0.0),
...     eig_threshold=0.0,
...     n_modes=2,
... )
>>> enc.num_features
6

from_cap(*, sh_l_max: int, slepian_l_max: int, cap_radius_deg: float, cap_centre_lonlat_deg: tuple[float, float], eig_threshold: float = 0.05, n_modes: int | None = None, input_mode: Literal['cartesian', 'lonlat'] = 'lonlat') -> HybridSphericalSlepianEncoder classmethod

Construct the hybrid encoder used by Slepian positional encodings.

Examples:

>>> from geonnax.slepian import HybridSphericalSlepianEncoder
>>> enc = HybridSphericalSlepianEncoder.from_cap(
...     sh_l_max=1,
...     slepian_l_max=3,
...     cap_radius_deg=20.0,
...     cap_centre_lonlat_deg=(0.0, 0.0),
...     eig_threshold=0.0,
...     n_modes=2,
... )
>>> # sh_l_max=1 → (1 + 1)^2 = 4 SH features, plus 2 Slepian modes.
>>> enc.num_features
6
Source code in .venv/lib/python3.12/site-packages/geonnax/slepian.py
@classmethod
def from_cap(
    cls,
    *,
    sh_l_max: int,
    slepian_l_max: int,
    cap_radius_deg: float,
    cap_centre_lonlat_deg: tuple[float, float],
    eig_threshold: float = 0.05,
    n_modes: int | None = None,
    input_mode: Literal["cartesian", "lonlat"] = "lonlat",
) -> HybridSphericalSlepianEncoder:
    """Construct the hybrid encoder used by Slepian positional encodings.

    Examples:
        >>> from geonnax.slepian import HybridSphericalSlepianEncoder
        >>> enc = HybridSphericalSlepianEncoder.from_cap(
        ...     sh_l_max=1,
        ...     slepian_l_max=3,
        ...     cap_radius_deg=20.0,
        ...     cap_centre_lonlat_deg=(0.0, 0.0),
        ...     eig_threshold=0.0,
        ...     n_modes=2,
        ... )
        >>> # sh_l_max=1 → (1 + 1)^2 = 4 SH features, plus 2 Slepian modes.
        >>> enc.num_features
        6
    """
    return cls(
        sh_encoder=SphericalHarmonicEncoder(l_max=sh_l_max, input_mode=input_mode),
        slepian_encoder=SlepianEncoder.from_cap(
            l_max=slepian_l_max,
            cap_radius_deg=cap_radius_deg,
            cap_centre_lonlat_deg=cap_centre_lonlat_deg,
            eig_threshold=eig_threshold,
            n_modes=n_modes,
            input_mode=input_mode,
        ),
    )

BayesianSlepianEncoder

Bases: Parameterized

Slepian encoder with NumPyro sites for cap radius and centre.

Truncation is controlled by n_modes only. eig_threshold is not exposed because filtering by concentration ratio under tracing would require a data-dependent output size (jnp.nonzero needs a static size), which breaks JIT and downstream shape contracts.

Source code in packages/pyrox-nn/src/pyrox_nn/_slepian.py
class BayesianSlepianEncoder(Parameterized):
    """Slepian encoder with NumPyro sites for cap radius and centre.

    Truncation is controlled by ``n_modes`` only. ``eig_threshold`` is not
    exposed because filtering by concentration ratio under tracing would
    require a data-dependent output size (``jnp.nonzero`` needs a static
    ``size``), which breaks JIT and downstream shape contracts.
    """

    l_max: int = eqx.field(static=True)
    init_cap_radius_deg: float = eqx.field(static=True)
    init_cap_centre_lonlat_deg: tuple[float, float] = eqx.field(static=True)
    n_modes: int | None = eqx.field(static=True, default=None)
    input_mode: Literal["cartesian", "lonlat"] = eqx.field(
        static=True, default="lonlat"
    )
    weight_by_eigenvalue: bool = eqx.field(static=True, default=True)
    pyrox_name: str | None = eqx.field(static=True, default=None)

    def setup(self) -> None:
        if self.l_max < 0:
            raise ValueError(f"l_max must be >= 0; got {self.l_max}.")
        if self.input_mode not in {"cartesian", "lonlat"}:
            raise ValueError(
                f"input_mode must be 'cartesian' or 'lonlat'; got {self.input_mode!r}."
            )
        self.register_param(
            "cap_radius",
            jnp.asarray(math.radians(self.init_cap_radius_deg)),
            constraint=dist.constraints.interval(0.0, math.pi),
        )
        self.register_param(
            "cap_centre",
            jnp.radians(jnp.asarray(self.init_cap_centre_lonlat_deg)),
            constraint=dist.constraints.real_vector,
        )

    @pyrox_method
    def __call__(
        self,
        x: Float[Array, "N 3"] | Float[Array, "N 2"],
    ) -> Float[Array, "N K"]:
        radius = self.get_param("cap_radius")
        centre = self.get_param("cap_centre")
        basis = slepian_cap_basis(
            self.l_max,
            radius,
            n_modes=self.n_modes,
            eig_threshold=None,
            lonlat_centre=centre,
        )
        features = basis.evaluate(_unit_xyz(x, self.input_mode))
        if self.weight_by_eigenvalue:
            features = features * jnp.sqrt(basis.eigenvalues)[None, :]
        return features

Pure-JAX helper functions

deg2rad(x: Float[Array, ...]) -> Float[Array, ...]

Convert degrees to radians element-wise.

Computes x · π / 180 and preserves the input shape.

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.geo import deg2rad
>>> deg2rad(jnp.array([0.0, 90.0, 180.0])).shape
(3,)
>>> deg2rad(jnp.array([[45.0, -90.0], [270.0, 360.0]])).shape
(2, 2)
Source code in .venv/lib/python3.12/site-packages/geonnax/geo.py
def deg2rad(x: Float[Array, ...]) -> Float[Array, ...]:
    r"""Convert degrees to radians element-wise.

    Computes ``x · π / 180`` and preserves the input shape.

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.geo import deg2rad
        >>> deg2rad(jnp.array([0.0, 90.0, 180.0])).shape
        (3,)
        >>> deg2rad(jnp.array([[45.0, -90.0], [270.0, 360.0]])).shape
        (2, 2)
    """
    return x * (jnp.pi / 180.0)  # deg → rad: scale by π/180, shape preserved

lonlat_scale(lonlat: Num[Array, 'N 2'], *, lon_range: tuple[float, float] = (-180.0, 180.0), lat_range: tuple[float, float] = (-90.0, 90.0)) -> Float[Array, 'N 2']

Affine-rescale lon/lat columns.

Values inside the given ranges map into [-1, 1]; out-of-range values are not clipped and map outside [-1, 1] linearly. The default ranges assume lonlat is in degrees; pass matching lon_range / lat_range in whatever unit you use.

Integer inputs are promoted to float32 before the affine step so (lonlat - lower) / (upper - lower) is not computed in integer arithmetic (which would silently round the output to -1 / 0 / 1).

Parameters:

Name Type Description Default
lonlat Num[Array, 'N 2']

Longitude/latitude matrix of shape (N, 2).

required
lon_range tuple[float, float]

(min, max) longitude domain (must satisfy min < max).

(-180.0, 180.0)
lat_range tuple[float, float]

(min, max) latitude domain (must satisfy min < max).

(-90.0, 90.0)

Returns:

Type Description
Float[Array, 'N 2']

Rescaled lon/lat array of shape (N, 2).

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.geo import lonlat_scale
>>> lonlat = jnp.array([[-180.0, -90.0], [0.0, 0.0], [180.0, 90.0]])
>>> lonlat_scale(lonlat).shape
(3, 2)
>>> # Midpoint of each range maps exactly to 0.
>>> bool((lonlat_scale(jnp.array([[0.0, 0.0]]))[0] == 0.0).all())
True
>>> # Custom domain (e.g. a regional grid in degrees).
>>> lonlat_scale(
...     jnp.array([[0.0, 50.0]]),
...     lon_range=(-10.0, 10.0),
...     lat_range=(40.0, 60.0),
... ).shape
(1, 2)
Source code in .venv/lib/python3.12/site-packages/geonnax/geo.py
def lonlat_scale(
    lonlat: Num[Array, "N 2"],
    *,
    lon_range: tuple[float, float] = (-180.0, 180.0),
    lat_range: tuple[float, float] = (-90.0, 90.0),
) -> Float[Array, "N 2"]:
    """Affine-rescale lon/lat columns.

    Values inside the given ranges map into ``[-1, 1]``; out-of-range
    values are not clipped and map outside ``[-1, 1]`` linearly. The
    default ranges assume ``lonlat`` is in degrees; pass matching
    ``lon_range`` / ``lat_range`` in whatever unit you use.

    Integer inputs are promoted to ``float32`` before the affine step
    so ``(lonlat - lower) / (upper - lower)`` is not computed in
    integer arithmetic (which would silently round the output to
    ``-1 / 0 / 1``).

    Args:
        lonlat: Longitude/latitude matrix of shape ``(N, 2)``.
        lon_range: ``(min, max)`` longitude domain (must satisfy
            ``min < max``).
        lat_range: ``(min, max)`` latitude domain (must satisfy
            ``min < max``).

    Returns:
        Rescaled lon/lat array of shape ``(N, 2)``.

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.geo import lonlat_scale
        >>> lonlat = jnp.array([[-180.0, -90.0], [0.0, 0.0], [180.0, 90.0]])
        >>> lonlat_scale(lonlat).shape
        (3, 2)
        >>> # Midpoint of each range maps exactly to 0.
        >>> bool((lonlat_scale(jnp.array([[0.0, 0.0]]))[0] == 0.0).all())
        True
        >>> # Custom domain (e.g. a regional grid in degrees).
        >>> lonlat_scale(
        ...     jnp.array([[0.0, 50.0]]),
        ...     lon_range=(-10.0, 10.0),
        ...     lat_range=(40.0, 60.0),
        ... ).shape
        (1, 2)
    """
    _validate_lonlat_shape(lonlat)
    _validate_range(lon_range, name="lon_range")
    _validate_range(lat_range, name="lat_range")

    lonlat = _promote_to_floating(lonlat)
    lower = jnp.asarray([lon_range[0], lat_range[0]], dtype=lonlat.dtype)  # (2,)
    upper = jnp.asarray([lon_range[1], lat_range[1]], dtype=lonlat.dtype)  # (2,)
    # Affine map: 2·(x − lo)/(hi − lo) − 1, so [lo, hi] → [-1, 1]. (N,2) -> (N,2)
    return 2.0 * (lonlat - lower) / (upper - lower) - 1.0

lonlat_to_cartesian3d(lonlat: Float[Array, 'N 2'], *, input_unit: Literal['degrees', 'radians'] = 'radians') -> Float[Array, 'N 3']

Lift lon/lat coordinates onto the unit sphere.

Uses the standard parameterization

\[ x = \cos(\phi)\cos(\lambda), \quad y = \cos(\phi)\sin(\lambda), \quad z = \sin(\phi), \]

where lon = λ and lat = ϕ. This matches the axis convention expected by the analogous GP-side inducing features, so the NN and GP spherical paths line up.

Parameters:

Name Type Description Default
lonlat Float[Array, 'N 2']

Longitude/latitude matrix of shape (N, 2).

required
input_unit Literal['degrees', 'radians']

Whether lonlat is in "degrees" or "radians".

'radians'

Returns:

Type Description
Float[Array, 'N 3']

Unit Cartesian coordinates of shape (N, 3).

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.geo import lonlat_to_cartesian3d
>>> # Prime meridian / equator → +x; output is one row of unit norm.
>>> lonlat_to_cartesian3d(jnp.array([[0.0, 0.0]])).shape
(1, 3)
>>> # North pole (lat = π/2) → +z, so the z-column equals 1.
>>> bool(
...     jnp.allclose(
...         lonlat_to_cartesian3d(jnp.array([[0.0, 0.5 * jnp.pi]]))[:, 2],
...         1.0,
...     )
... )
True
Source code in .venv/lib/python3.12/site-packages/geonnax/geo.py
def lonlat_to_cartesian3d(
    lonlat: Float[Array, "N 2"],
    *,
    input_unit: Literal["degrees", "radians"] = "radians",
) -> Float[Array, "N 3"]:
    r"""Lift lon/lat coordinates onto the unit sphere.

    Uses the standard parameterization

    $$
    x = \cos(\phi)\cos(\lambda), \quad
    y = \cos(\phi)\sin(\lambda), \quad
    z = \sin(\phi),
    $$


    where ``lon = λ`` and ``lat = ϕ``. This matches the axis
    convention expected by
    the analogous GP-side inducing features, so the NN and
    GP spherical paths line up.

    Args:
        lonlat: Longitude/latitude matrix of shape ``(N, 2)``.
        input_unit: Whether ``lonlat`` is in ``"degrees"`` or
            ``"radians"``.

    Returns:
        Unit Cartesian coordinates of shape ``(N, 3)``.

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.geo import lonlat_to_cartesian3d
        >>> # Prime meridian / equator → +x; output is one row of unit norm.
        >>> lonlat_to_cartesian3d(jnp.array([[0.0, 0.0]])).shape
        (1, 3)
        >>> # North pole (lat = π/2) → +z, so the z-column equals 1.
        >>> bool(
        ...     jnp.allclose(
        ...         lonlat_to_cartesian3d(jnp.array([[0.0, 0.5 * jnp.pi]]))[:, 2],
        ...         1.0,
        ...     )
        ... )
        True
    """
    _validate_lonlat_shape(lonlat)
    _validate_input_unit(input_unit)

    # (N,2) in deg/rad → radians; columns are λ (lon) and ϕ (lat).
    angles = deg2rad(lonlat) if input_unit == "degrees" else lonlat
    lon = angles[:, 0]  # λ, shape (N,)
    lat = angles[:, 1]  # ϕ, shape (N,)
    cos_lat = jnp.cos(lat)  # (N,)
    # x = cosϕ·cosλ, y = cosϕ·sinλ, z = sinϕ. Stack to (N,3) on the unit sphere.
    return jnp.stack(
        [
            cos_lat * jnp.cos(lon),
            cos_lat * jnp.sin(lon),
            jnp.sin(lat),
        ],
        axis=-1,
    )

cyclic_encode(angles: Float[Array, ' N'] | Float[Array, 'N D']) -> Float[Array, 'N F']

Encode periodic inputs as concatenated cos/sin features.

Parameters:

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

Angle vector (N,) or matrix (N, D) in radians.

required

Returns:

Type Description
Float[Array, 'N F']

(N, 2) for vector input or (N, 2 * D) for matrix input,

Float[Array, 'N F']

laid out as [cos_0, ..., cos_{D-1}, sin_0, ..., sin_{D-1}].

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.geo import cyclic_encode
>>> # Vector input (N,) → (N, 2) of [cos, sin].
>>> cyclic_encode(jnp.array([0.0, jnp.pi])).shape
(2, 2)
>>> # Multi-dimensional input: each column is encoded independently.
>>> cyclic_encode(jnp.zeros((3, 2))).shape
(3, 4)
Source code in .venv/lib/python3.12/site-packages/geonnax/geo.py
def cyclic_encode(
    angles: Float[Array, " N"] | Float[Array, "N D"],
) -> Float[Array, "N F"]:
    """Encode periodic inputs as concatenated cos/sin features.

    Args:
        angles: Angle vector ``(N,)`` or matrix ``(N, D)`` in radians.

    Returns:
        ``(N, 2)`` for vector input or ``(N, 2 * D)`` for matrix input,
        laid out as ``[cos_0, ..., cos_{D-1}, sin_0, ..., sin_{D-1}]``.

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.geo import cyclic_encode
        >>> # Vector input (N,) → (N, 2) of [cos, sin].
        >>> cyclic_encode(jnp.array([0.0, jnp.pi])).shape
        (2, 2)
        >>> # Multi-dimensional input: each column is encoded independently.
        >>> cyclic_encode(jnp.zeros((3, 2))).shape
        (3, 4)
    """
    if angles.ndim == 1:
        promoted = einx.id("n -> n 1", angles)  # (N,) -> (N, 1)
    elif angles.ndim == 2:
        promoted = angles  # (N, D)
    else:
        raise ValueError(f"angles must be (N,) or (N, D); got shape {angles.shape}.")
    # [cos θ, sin θ] embeds each angle on the unit circle. (N, D) -> (N, 2·D)
    return jnp.concatenate([jnp.cos(promoted), jnp.sin(promoted)], axis=-1)

spherical_harmonic_encode(lonlat: Float[Array, 'N 2'], l_max: int, *, input_unit: Literal['degrees', 'radians'] = 'radians') -> Float[Array, 'N M']

Lift lon/lat to \(S^2\) and evaluate real spherical harmonics.

Parameters:

Name Type Description Default
lonlat Float[Array, 'N 2']

Longitude/latitude matrix of shape (N, 2).

required
l_max int

Maximum harmonic degree.

required
input_unit Literal['degrees', 'radians']

Whether lonlat is in "degrees" or "radians".

'radians'

Returns:

Type Description
Float[Array, 'N M']

Real spherical-harmonic features of shape (N, (l_max + 1)^2).

Examples:

>>> import jax.numpy as jnp
>>> from geonnax.geo import spherical_harmonic_encode
>>> lonlat = jnp.array([[0.0, 0.0], [1.5707964, 0.0]])
>>> # (N, 2) -> (N, (l_max + 1)^2); here (l_max + 1)^2 = 16.
>>> spherical_harmonic_encode(lonlat, l_max=3).shape
(2, 16)
>>> # l_max=0 keeps only the constant Y_0^0 mode.
>>> spherical_harmonic_encode(lonlat, l_max=0).shape
(2, 1)
Source code in .venv/lib/python3.12/site-packages/geonnax/geo.py
def spherical_harmonic_encode(
    lonlat: Float[Array, "N 2"],
    l_max: int,
    *,
    input_unit: Literal["degrees", "radians"] = "radians",
) -> Float[Array, "N M"]:
    """Lift lon/lat to $S^2$ and evaluate real spherical harmonics.

    Args:
        lonlat: Longitude/latitude matrix of shape ``(N, 2)``.
        l_max: Maximum harmonic degree.
        input_unit: Whether ``lonlat`` is in ``"degrees"`` or
            ``"radians"``.

    Returns:
        Real spherical-harmonic features of shape ``(N, (l_max + 1)^2)``.

    Examples:
        >>> import jax.numpy as jnp
        >>> from geonnax.geo import spherical_harmonic_encode
        >>> lonlat = jnp.array([[0.0, 0.0], [1.5707964, 0.0]])
        >>> # (N, 2) -> (N, (l_max + 1)^2); here (l_max + 1)^2 = 16.
        >>> spherical_harmonic_encode(lonlat, l_max=3).shape
        (2, 16)
        >>> # l_max=0 keeps only the constant Y_0^0 mode.
        >>> spherical_harmonic_encode(lonlat, l_max=0).shape
        (2, 1)
    """
    # (N, 2) -> (N, 3) unit sphere, then evaluate Y_l^m up to l_max.
    unit_xyz = lonlat_to_cartesian3d(lonlat, input_unit=input_unit)
    return real_spherical_harmonics(unit_xyz, l_max=l_max)  # (N, (l_max+1)^2)