torch_harmonics.distributed.compute_split_shapes#
- torch_harmonics.distributed.compute_split_shapes(size, num_chunks)[source]#
Compute balanced chunk sizes for distributing a dimension across ranks.
Divides
sizeelements intonum_chunkspieces that differ by at most one element. The firstsize % num_chunkschunks receive one extra element; the remaining chunks get the base sizesize // num_chunks.This is used internally by every distributed module to determine how latitudes, longitudes, and spectral modes are partitioned across process groups.
- Parameters:
- Returns:
Per-rank chunk sizes, ordered by rank.
- Return type:
List[int]
- Raises:
RuntimeError – If
size < num_chunks(every chunk must be non-empty).
Examples
>>> from torch_harmonics.distributed import compute_split_shapes >>> compute_split_shapes(256, 4) [64, 64, 64, 64] >>> compute_split_shapes(128, 3) [43, 43, 42] >>> compute_split_shapes(10, 4) [3, 3, 2, 2]