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,
Bases:
DiscreteContinuousConvDistributed 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.DiscreteContinuousConvS2Serial 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 thetranspose-back; the K-expanded intermediate is saved for backward.
fused=True— reordered a2a: the weight einsum runs before thecollectives 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:
References
[1]