torch_harmonics.distributed.DistributedSpectralConvS2#
- class torch_harmonics.distributed.DistributedSpectralConvS2(
- in_shape,
- out_shape,
- in_channels,
- out_channels,
- num_groups=1,
- grid_in='equiangular',
- grid_out='equiangular',
- bias=False,
Bases:
ModuleDistributed spectral convolution layer on \(S^2\) implemented with distributed real SHT (Driscoll-Healy formulation, see https://api.semanticscholar.org/CorpusID:122817218). Computation is split across polar and azimuth communicator groups.
Distribution scheme. The forward and inverse SHTs are performed by
DistributedRealSHTandDistributedInverseRealSHTrespectively — see their docstrings for the all-to-all transpose sequence. After the forward SHT, the spectral coefficients are split so that degreeslare distributed across polar ranks and ordersmacross azimuth ranks. The learnable spectral weightK[groups, c_in, c_out, l]is stored with itsldimension sharded across polar ranks (each rank holds only its locallmax_localslice). The spectral contraction is therefore fully local — no communication is needed for the channel mixing.Note
When saving a checkpoint to a single file, the weight tensor must be gathered across polar ranks along the
ldimension to recover the full(groups, c_in, c_out, lmax)shape. Likewise, when loading a serial checkpoint into the distributed module, theldimension must be split according tocompute_split_shapes().Note
The spectral weight
K[g, c_in, c_out, l]has nomdimension — it is broadcast over the spectral orders during the contraction. Because orders are split across azimuth ranks, each rank computes only a partial sum of the weight gradient (over its localmmodes). For correct gradients the user must all-reduce (sum) the weight gradients across azimuth ranks. This can be implemented viatorch.nn.parallel.DistributedDataParallelcommunication hooks ortorch.Tensor.register_post_accumulate_grad_hook().See also
torch_harmonics.SpectralConvS2Serial counterpart with full mathematical description and parameter documentation.
- Parameters:
in_shape (Tuple[int]) – Spatial input grid shape
(nlat, nlon).out_shape (Tuple[int]) – Spatial output grid shape
(nlat, nlon).in_channels (int) – Number of input channels.
out_channels (int) – Number of output channels.
num_groups (int, optional) – Number of channel groups for grouped spectral weights, by default 1.
grid_in (str, optional) – Grid used for the forward distributed SHT (
"equiangular","legendre-gauss","lobatto","equiangular-trapezoidal"), by default"equiangular".grid_out (str, optional) – Grid used for the inverse distributed SHT, same options as
grid_in.bias (bool, optional) – If
True, adds a learnable spectral bias computed from the spatial integral (replicated across process groups as needed), by defaultFalse.
- Raises:
AssertionError – If
in_channelsorout_channelsis not divisible bynum_groups.- Returns:
Tensor of shape
(..., out_channels, out_shape[0], out_shape[1]).- Return type:
Notes
The layer truncates
lmax/mmaxto the distributed SHT limits, and uses locallmax/mmaxslices when constructing spectral weights. The grouped contraction is performed with_contract_lwise.