torch_harmonics.RealSHT#

class torch_harmonics.RealSHT(
nlat,
nlon,
lmax=None,
mmax=None,
grid='equiangular',
norm='ortho',
csphase=True,
)[source]#

Bases: Module

Defines a module for computing the forward (real-valued) SHT. Precomputes Legendre Gauss nodes, weights and associated Legendre polynomials on these nodes. The SHT is applied to the last two dimensions of the input.

Given a real-valued signal \(f(\theta, \lambda)\) sampled on the sphere, the forward scalar SHT computes the spherical harmonic coefficients via a longitudinal FFT followed by Legendre quadrature:

\[\hat{f}_l^m = 2\pi \sum_{k=0}^{N_\theta - 1} \tilde{f}_m(\theta_k)\, P_l^m(\cos\theta_k)\, q_k\]

where \(\tilde{f}_m\) are the Fourier modes and \(q_k\) are the quadrature weights.

See also

Spherical harmonic transforms

User guide with the full mathematical derivation, normalization conventions, grid types, and worked examples.

Parameters:
  • nlat (int) – Number of latitude points

  • nlon (int) – Number of longitude points

  • lmax (int) – Maximum spherical harmonic degree

  • mmax (int) – Maximum spherical harmonic order

  • grid (str) – Grid type ("equiangular", "legendre-gauss", "lobatto", "equiangular-trapezoidal"), by default "equiangular"

  • norm (str) – Normalization convention ("ortho", "schmidt", "unnorm"), by default "ortho".

  • csphase (bool) – Whether to include the Condon–Shortley phase factor \((-1)^m\), by default True.

Examples

>>> import torch
>>> import torch_harmonics as th
>>> nlat, nlon = 128, 256
>>> sht = th.RealSHT(nlat, nlon).cuda()
>>> signal = torch.randn(1, nlat, nlon, device="cuda")
>>> coeffs = sht(signal)   # shape (1, lmax, mmax), complex
>>> coeffs.shape
torch.Size([1, 128, 129])

Note

This module uses cuFFT (via torch.fft.rfft()) to compute the longitudinal Fourier transform efficiently. When running in float16 or bfloat16 precision, cuFFT requires the transformed dimension (nlon) to be a power of two. If your grid does not satisfy this constraint and the module is called inside a torch.autocast context, guard it with torch.autocast(device_type="cuda", enabled=False):

with torch.autocast(device_type="cuda", dtype=torch.float16):
    # ... other half-precision work ...
    with torch.autocast(device_type="cuda", enabled=False):
        coeffs = sht(signal.float())

References

[2], [3]

forward(x)[source]#

Compute the forward (real) spherical harmonic transform.

Parameters:

x (torch.Tensor) – Real-valued signal on the sphere of shape (..., nlat, nlon).

Returns:

Complex spherical harmonic coefficients of shape (..., lmax, mmax).

Return type:

torch.Tensor