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,
Bases:
DiscreteContinuousConvDiscrete-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: