torch_harmonics.distributed.DistributedDiscreteContinuousConvS2#

class torch_harmonics.distributed.DistributedDiscreteContinuousConvS2(
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

Distributed version of Discrete-continuous convolutions (DISCO) on the 2-Sphere as described in [1]. We assume the data can be split in polar and azimuthal directions.

See also

torch_harmonics.DiscreteContinuousConvS2

Serial counterpart with full mathematical description and parameter documentation.

The algorithm is all-to-all (azimuth <-> channel swap so the sparse psi contraction runs against the full nlon_in row, polar reduce_scatter completes the H sum, then back to channel-distributed). The fused= flag mirrors the serial conv:

fused=False (default) — standard a2a: einsum after the

transpose-back; the K-expanded intermediate is saved for backward.

fused=True — reordered a2a: the weight einsum runs before the

collectives on the local azimuth channel shard, via the fused contraction+einsum op that recomputes the K-expanded in backward instead of saving it. K× lower activation memory and K× less collective volume, at the cost of one extra contraction in backward. CUDA + optimized kernels only.

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

  • out_channels (int) – Number of output channels

  • in_shape (Tuple[int]) – Shape of the input tensor

  • out_shape (Tuple[int]) – Shape of the output tensor

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

  • basis_type (Optional[str]) – Type of basis to use

  • basis_norm_mode (Optional[str]) – Normalization mode for the filter basis

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

  • grid_in (Optional[str]) – Grid type for the input tensor

  • grid_out (Optional[str]) – Grid type for the output tensor

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

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

  • optimized_kernel (Optional[bool]) – Use the optimized CUDA contraction kernel. Required when fused=True.

  • fused (bool) – Mirrors the serial conv. False (default): standard all-to-all (the K-expanded intermediate is saved for backward). True: reordered all-to-all — the weight einsum runs before the collectives on the local azimuth channel shard and the K-expanded is recomputed in backward instead of saved, for K× lower activation memory and K× less collective volume (CUDA + optimized kernels only).

Returns:

Output tensor

Return type:

torch.Tensor

References

[1]