torch_harmonics.distributed.init#

torch_harmonics.distributed.init(polar_process_group, azimuth_process_group)[source]#

Initialize the torch-harmonics distributed backend.

This must be called before any distributed SHT, convolution, or other distributed module is used. It registers two orthogonal process groups that define a 2-D process grid over the sphere: one group for the polar (latitudinal) dimension and one for the azimuthal (longitudinal) dimension.

The two groups are typically created from a single torch.distributed.new_subgroups_by_enumeration() (or equivalent) call so that every global rank belongs to exactly one polar group and one azimuth group. See the distributed guide for a complete example of how to build the orthogonal communicator grid.

Parameters:
  • polar_process_group (torch.distributed.ProcessGroup) – Process group whose members share the same azimuthal index and collectively own all latitude chunks.

  • azimuth_process_group (torch.distributed.ProcessGroup) – Process group whose members share the same polar index and collectively own all longitude chunks.

Examples

Setting up a 2 x 4 process grid on 8 GPUs (2 polar ranks, 4 azimuth ranks):

import torch.distributed as dist
import torch_harmonics.distributed as thd

dist.init_process_group(backend="nccl")
world_rank = dist.get_rank()
world_size = dist.get_world_size()   # 8

num_polar, num_azimuth = 2, 4

# ranks in the same row share a polar index -> azimuth group
azimuth_group = dist.new_group(
    ranks=[r for r in range(world_size)
           if r // num_azimuth == world_rank // num_azimuth]
)
# ranks in the same column share an azimuth index -> polar group
polar_group = dist.new_group(
    ranks=[r for r in range(world_size)
           if r % num_azimuth == world_rank % num_azimuth]
)

thd.init(polar_group, azimuth_group)