Skip to content

Fix shape metadata when copying tensors across devices - #4245

Open
MrCapricornLiu wants to merge 1 commit into
huggingface:mainfrom
MrCapricornLiu:MrCapricornLiu/fix/inference-tensor-shape-metadata
Open

MrCapricornLiu wants to merge 1 commit into
huggingface:mainfrom
MrCapricornLiu:MrCapricornLiu/fix/inference-tensor-shape-metadata

Conversation

@MrCapricornLiu

Copy link
Copy Markdown

copy_tensor_to_devices, used by pipeline inference with gather_output=True, reduces a shape buffer whose unused entries and non-source ranks are uninitialized. Those values can become dimensions or an invalid dtype. Filtering metadata with nonzero() also loses zero-sized axes, and the receiver passes a tensor where torch.zeros expects a size tuple.

Initialize the metadata buffer to zero, prefix the dimensions with their count and dtype, and construct receiver tensors from the decoded size tuple. The existing collective and dtype mapping are retained.

Tests cover deterministic uninitialized-memory filling, scalars, empty dimensions, three dtypes, and either the first or last rank as source. All 12 metadata subcases fail on the original code and pass with the fix; the utilities module reports 46 passed and 5 skipped. The distributed operations script passes on eight H800 GPUs, and an eight-stage prepare_pippy(gather_output=True) model matches the unpartitioned reference on every rank. Ruff lint/format and pre-commit pass. XLA and multi-node inference were not tested.

Signed-off-by: Chenghao Liu <chliu@stu.pku.edu.cn>

@SunMarc SunMarc left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for fixing ! just a nit

Comment thread tests/test_utils.py
Comment on lines +80 to +98
@unittest.skipUnless(hasattr(torch.utils, "deterministic"), "requires deterministic memory filling")
def test_gather_tensor_shape_preserves_metadata(self):
enabled = torch.are_deterministic_algorithms_enabled()
warn_only = torch.is_deterministic_algorithms_warn_only_enabled()
fill = torch.utils.deterministic.fill_uninitialized_memory
try:
torch.use_deterministic_algorithms(True)
torch.utils.deterministic.fill_uninitialized_memory = True
for shape in [(2, 3), (), (0, 3), (2, 0, 3)]:
for dtype in [torch.float32, torch.bfloat16, torch.int64]:
with self.subTest(shape=shape, dtype=dtype):
tensor = torch.ones(shape, dtype=dtype, device=PartialState().device)
actual_shape, actual_dtype = gather_tensor_shape(tensor)
self.assertEqual(actual_shape.numel(), len(shape))
self.assertEqual(actual_shape.flatten().tolist(), list(shape))
self.assertEqual(TENSOR_INT_TO_DTYPE[actual_dtype], dtype)
finally:
torch.utils.deterministic.fill_uninitialized_memory = fill
torch.use_deterministic_algorithms(enabled, warn_only=warn_only)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

don't need to add a new test, remove it

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.

2 participants