torch_harmonics.distributed.distributed_transpose_azimuth#

torch_harmonics.distributed.distributed_transpose_azimuth(input_, dims_, shapes_)[source]#

All-to-all transpose across the azimuth process group.

Redistributes input_ so that data sharded along dims_[0] becomes sharded along dims_[1]. This is the core communication pattern used when switching between spatial and spectral partitioning of the longitude axis.

Parameters:
  • input (torch.Tensor) – Input tensor, partitioned along dims_[0].

  • dims (tuple[int, int]) – (source_dim, target_dim) — the dimension to scatter from and the dimension to gather into.

  • shapes (list[int]) – Per-rank sizes along dims_[1] (i.e. the expected receive sizes).

Returns:

Transposed tensor, now partitioned along dims_[1].

Return type:

torch.Tensor