torch_harmonics.DiscreteContinuousConvS2#

class torch_harmonics.DiscreteContinuousConvS2(
in_channels,
out_channels,
in_shape,
out_shape,
kernel_shape,
basis_type='piecewise linear',
basis_norm_mode='nodal',
groups=1,
grid_in='equiangular',
grid_out='equiangular',
bias=True,
theta_cutoff=None,
optimized_kernel=True,
fused=False,
)[source]#

Bases: DiscreteContinuousConv

Discrete-continuous (DISCO) convolution on the 2-sphere, as described in [1].

The layer evaluates a spherical convolution with a compactly supported filter of angular radius theta_cutoff. The filter is parameterised as a learnable linear combination of fixed basis functions \(\{\phi_k\}\), and the integral is computed by sparse quadrature over the input grid, giving \(O(N)\) cost in the number of grid points. The forward pass is

\[g^{c_o}(\theta'_j, \lambda'_q) = \sum_{c_i} \sum_k w_k^{c_o,c_i} \sum_{i,\,p} \Psi_{k,\,j,\,(i,p)}\; f^{c_i}(\theta_i, \lambda'_q + \lambda_p)\]

where \(\Psi\) is a precomputed sparse convolution tensor that encodes the basis function values at rotated input grid positions, weighted by the quadrature weights. Because the grid is equispaced in longitude, \(\Psi\) is independent of the output longitude (p-shift symmetry).

See also

DISCO convolutions on the sphere

User guide with the full mathematical derivation, filter basis visualisations, and worked examples.

Parameters:
  • in_channels (int) – Number of input channels

  • out_channels (int) – Number of output channels

  • in_shape (Tuple[int]) – Input shape of the convolution tensor

  • out_shape (Tuple[int]) – Output shape of the convolution tensor

  • kernel_shape (Union[int, Tuple[int], Tuple[int, int]]) – Shape of the kernel

  • basis_type (Optional[str]) – Type of the basis functions

  • basis_norm_mode (Optional[str]) – Mode for basis normalization

  • groups (Optional[int]) – Number of groups

  • grid_in (Optional[str]) – Input grid type

  • grid_out (Optional[str]) – Output grid type

  • bias (Optional[bool]) – Whether to use bias

  • theta_cutoff (Optional[float]) – Theta cutoff for the filter basis functions

  • optimized_kernel (Optional[bool]) – Whether to use the optimized kernel (if available)

  • fused (Optional[bool]) – When True, fuses the sparse contraction and weight multiplication into a single autograd region to avoid storing the K-expanded intermediate in the graph. Trades one extra contraction recompute in backward for K× memory savings. Only effective when optimized_kernel is True.

References

[1]

forward(x)[source]#

Apply the discrete-continuous convolution.

Parameters:

x (torch.Tensor) – Input signal of shape (batch, in_channels, nlat_in, nlon_in).

Returns:

Convolved signal of shape (batch, out_channels, nlat_out, nlon_out).

Return type:

torch.Tensor