torch_harmonics.distributed.split_tensor_along_dim#

torch_harmonics.distributed.split_tensor_along_dim(tensor, dim, num_chunks)[source]#

Split a tensor along a given dimension into balanced chunks.

Uses compute_split_shapes() to determine chunk sizes, so the split is consistent with the partitioning used by all distributed modules in torch-harmonics. Chunk sizes differ by at most one element.

Parameters:
  • tensor (torch.Tensor) – The tensor to split.

  • dim (int) – The dimension along which to split.

  • num_chunks (int) – Number of chunks (typically the process-group size).

Returns:

A tuple of num_chunks tensor views.

Return type:

tuple[torch.Tensor, …]

Raises:

RuntimeError – If dim is out of range or tensor.shape[dim] < num_chunks.

Examples

>>> import torch
>>> from torch_harmonics.distributed import split_tensor_along_dim
>>> x = torch.arange(10).unsqueeze(0)          # shape (1, 10)
>>> parts = split_tensor_along_dim(x, dim=1, num_chunks=3)
>>> [p.shape[1] for p in parts]
[4, 3, 3]