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
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]
|
|
lat_range |
tuple[float, float]
|
|
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
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 |
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
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
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 |
input_mode |
Literal['cartesian', 'lonlat']
|
|
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
num_features: int
property
¶
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
37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 | |
num_features: int
property
¶
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
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
156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 | |
num_features: int
property
¶
Number of concatenated output features (SH modes + Slepian modes).
Examples:
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
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
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
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 |
required |
lon_range
|
tuple[float, float]
|
|
(-180.0, 180.0)
|
lat_range
|
tuple[float, float]
|
|
(-90.0, 90.0)
|
Returns:
| Type | Description |
|---|---|
Float[Array, 'N 2']
|
Rescaled lon/lat array of shape |
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
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
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 |
required |
input_unit
|
Literal['degrees', 'radians']
|
Whether |
'radians'
|
Returns:
| Type | Description |
|---|---|
Float[Array, 'N 3']
|
Unit Cartesian coordinates of shape |
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
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 |
required |
Returns:
| Type | Description |
|---|---|
Float[Array, 'N F']
|
|
Float[Array, 'N F']
|
laid out as |
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
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 |
required |
l_max
|
int
|
Maximum harmonic degree. |
required |
input_unit
|
Literal['degrees', 'radians']
|
Whether |
'radians'
|
Returns:
| Type | Description |
|---|---|
Float[Array, 'N M']
|
Real spherical-harmonic features of shape |
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)