torch_harmonics.DiscreteContinuousConvTransposeS2#

class torch_harmonics.DiscreteContinuousConvTransposeS2(
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,
)[source]#

Bases: DiscreteContinuousConv

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

This is the transpose (adjoint) of DiscreteContinuousConvS2. It uses the same continuous-filter and quadrature construction but applies the \(\Psi\) tensor in the reverse direction – typically to map a coarser grid to a finer one (upsampling), analogous to a transposed/strided convolution in the planar case. It shares the compact-support filter and sparse, linearly scaling evaluation, and the same approximate \(SO(3)\) equivariance.

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)

References

[1]

forward(x)[source]#

Apply the transpose 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