Skip to content

Fix gather_tensor_shape discarding zero-sized dimensions and crash in copy_tensor_to_devices - #4357

Open
Ryanakml wants to merge 1 commit into
huggingface:mainfrom
Ryanakml:fix-gather-tensor-shape-zero-dims
Open

Ryanakml wants to merge 1 commit into
huggingface:mainfrom
Ryanakml:fix-gather-tensor-shape-zero-dims

Conversation

@Ryanakml

Copy link
Copy Markdown

Problem

  1. gather_tensor_shape silently corrupts shapes for tensors with zero-sized dimensions (e.g. empty sequence slices or empty batch tensors):
    Passing t = torch.zeros(3, 0, 4) into gather_tensor_shape(t) returns shape [3, 4] instead of (3, 0, 4) — the zero dimension is completely stripped.
  2. copy_tensor_to_devices crashes with TypeError on worker ranks:
  File "accelerate/utils/operations.py", line 600, in copy_tensor_to_devices
    tensor = torch.zeros(shape, dtype=TENSOR_INT_TO_DTYPE[dtype]).to(state.device)
TypeError: zeros(): argument 'size' (position 1) must be tuple of ints, not Tensor

Root cause

  1. gather_tensor_shape used base_tensor[base_tensor.nonzero()] to filter non-zero elements. Any legitimate 0 dimension (such as (0, 10) or (3, 0, 4)) gets stripped out along with padding.
  2. In PyTorch, indexing a 1D tensor with its .nonzero() indices returns a 2D tensor of shape (N, 1). When worker ranks (tensor is None) call copy_tensor_to_devices, torch.zeros(shape, ...) fails because shape is a 2D tensor rather than a tuple of ints.
  3. torch.empty was previously allocated for base_tensor, which risks reading uninitialized dirty memory from the PyTorch caching allocator on ranks where tensor is None.

Fix

  • Explicitly encode num_dims + 1 at index 0 of base_tensor and write/read shape dimensions directly by slicing without relying on nonzero().
  • Allocate safely with torch.zeros (sized to 66 elements, since standard PyTorch supports up to 64 dimensions).
  • Return shape as a standard tuple of ints, ensuring torch.zeros(shape, ...) works properly across all worker ranks.

Verification

  • Added test_gather_tensor_shape and test_copy_tensor_to_devices in tests/test_utils.py testing scalars (), multi-dim tensors, and tensors with zero-sized dimensions ((3, 0, 4) and (0, 10)).
  • All 50 tests in tests/test_utils.py pass (50 passed, 4 skipped).
  • Code formatted and clean with ruff format --check.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant