torch_harmonics.distributed.flatten_and_pad_leading_dims#
- torch_harmonics.distributed.flatten_and_pad_leading_dims(
- tensor,
- min_leading_size,
- num_trailing_dims=2,
Collapse all but the trailing
num_trailing_dimsdims into a single leading dim, padding it to at leastmin_leading_size.The distributed (S)HT redistributes this leading (“channel/batch”) axis across the process grid via all-to-all transposes, which require every rank to receive a non-empty chunk – i.e. the leading dim must be at least the (largest) group size. Uneven splits are fine (e.g. 5 elements over 4 ranks -> [2, 1, 1, 1]), so we only pad when the leading dim is smaller than the group size, never up to a multiple of it. Since the transforms are linear, zero-padding leaves the real entries untouched;
unpad_and_unflatten_leading_dims()restores the original layout afterwards.- Parameters:
tensor (torch.Tensor) – Tensor whose last
num_trailing_dimsdims are the transform dims (everything before them is flattened into the leading axis).min_leading_size (int) – Minimum size the flattened leading dim must reach. Pass
max(comm_size_polar, comm_size_azimuth)– both transpose directions split this same axis, so it must be at least as large as the larger group.num_trailing_dims (int) – Number of trailing dims to keep intact.
2for the scalar SHT (nlat, nlon);3for the vector SHT (2, nlat, nlon), so the component axis is preserved.
- Returns:
tensor (torch.Tensor) – Flattened (and possibly zero-padded) contiguous tensor with shape
(M_pad, *trailing).lead_shape (torch.Size) – The original leading dims, used to restore the shape later.
lead_size (int) – The true (pre-pad) flattened leading size, used to slice off the padding.