Conversation
…tensor_to_devices
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Problem
gather_tensor_shapesilently corrupts shapes for tensors with zero-sized dimensions (e.g. empty sequence slices or empty batch tensors):Passing
t = torch.zeros(3, 0, 4)intogather_tensor_shape(t)returns shape[3, 4]instead of(3, 0, 4)— the zero dimension is completely stripped.copy_tensor_to_devicescrashes withTypeErroron worker ranks:Root cause
gather_tensor_shapeusedbase_tensor[base_tensor.nonzero()]to filter non-zero elements. Any legitimate0dimension (such as(0, 10)or(3, 0, 4)) gets stripped out along with padding..nonzero()indices returns a 2D tensor of shape(N, 1). When worker ranks (tensor is None) callcopy_tensor_to_devices,torch.zeros(shape, ...)fails becauseshapeis a 2D tensor rather than a tuple of ints.torch.emptywas previously allocated forbase_tensor, which risks reading uninitialized dirty memory from the PyTorch caching allocator on ranks wheretensor is None.Fix
num_dims + 1at index 0 ofbase_tensorand write/read shape dimensions directly by slicing without relying onnonzero().torch.zeros(sized to 66 elements, since standard PyTorch supports up to 64 dimensions).tupleof ints, ensuringtorch.zeros(shape, ...)works properly across all worker ranks.Verification
test_gather_tensor_shapeandtest_copy_tensor_to_devicesintests/test_utils.pytesting scalars(), multi-dim tensors, and tensors with zero-sized dimensions ((3, 0, 4)and(0, 10)).tests/test_utils.pypass (50 passed, 4 skipped).ruff format --check.