已合并
test(distributed): add test for WriteItem.tensor_storage_size #35505
Flipped创建于 5月13日
test(distributed): add test for WriteItem.tensor_storage_size #35505
已合并
共 1 个文件变更+77-0
| @@ -0,0 +1,77 @@ | |||
| 1 | +""" | ||
| 2 | +Add validation cases for torch.distributed.checkpoint.planner APIs on NPU: | ||
| 3 | +1. PyTorch community tests lack sufficient validation for WriteItem.tensor_storage_size, so this file is added. | ||
| 4 | +2. This file validates torch.distributed.checkpoint.planner.WriteItem.tensor_storage_size. | ||
| 5 | +""" | ||
| 6 | + | ||
| 7 | +import torch | ||
| 8 | +from torch.distributed.checkpoint.metadata import ( | ||
| 9 | + ChunkStorageMetadata, | ||
| 10 | + MetadataIndex, | ||
| 11 | + TensorProperties, | ||
| 12 | +) | ||
| 13 | +from torch.distributed.checkpoint.planner import ( | ||
| 14 | + TensorWriteData, | ||
| 15 | + WriteItem, | ||
| 16 | + WriteItemType, | ||
| 17 | +) | ||
| 18 | +from torch.testing._internal.common_utils import TestCase, run_tests | ||
| 19 | + | ||
| 20 | + | ||
| 21 | +device_type = ( | ||
| 22 | + torch.accelerator.current_accelerator().type | ||
| 23 | + if torch.accelerator.is_available() | ||
| 24 | + else "cpu" | ||
| 25 | +) | ||
| 26 | + | ||
| 27 | + | ||
| 28 | +class TestPlannerAPI(TestCase): | ||
| 29 | + | ||
| 30 | + def _make_tensor_write_item(self, tensor, write_item_type): | ||
| 31 | + tensor_data = TensorWriteData( | ||
| 32 | + chunk=ChunkStorageMetadata( | ||
| 33 | + offsets=torch.Size([0] * tensor.dim()), | ||
| 34 | + sizes=tensor.size(), | ||
| 35 | + ), | ||
| 36 | + properties=TensorProperties.create_from_tensor(tensor), | ||
| 37 | + size=tensor.size(), | ||
| 38 | + ) | ||
| 39 | + | ||
| 40 | + return WriteItem( | ||
| 41 | + index=MetadataIndex("tensor"), | ||
| 42 | + type=write_item_type, | ||
| 43 | + tensor_data=tensor_data, | ||
| 44 | + ) | ||
| 45 | + | ||
| 46 | + def test_write_item_tensor_storage_size_for_tensor(self): | ||
| 47 | + for dtype in (torch.float32, torch.float16, torch.int8): | ||
| 48 | + tensor = torch.empty((2, 3), dtype=dtype).to(device_type) | ||
| 49 | + write_item = self._make_tensor_write_item( | ||
| 50 | + tensor, | ||
| 51 | + WriteItemType.TENSOR, | ||
| 52 | + ) | ||
| 53 | + | ||
| 54 | + expected_size = tensor.numel() * tensor.element_size() | ||
| 55 | + self.assertEqual(write_item.tensor_storage_size(), expected_size) | ||
| 56 | + | ||
| 57 | + def test_write_item_tensor_storage_size_for_shard(self): | ||
| 58 | + tensor = torch.empty((2, 3), dtype=torch.float32).to(device_type) | ||
| 59 | + write_item = self._make_tensor_write_item( | ||
| 60 | + tensor, | ||
| 61 | + WriteItemType.SHARD, | ||
| 62 | + ) | ||
| 63 | + | ||
| 64 | + expected_size = tensor.numel() * tensor.element_size() | ||
| 65 | + self.assertEqual(write_item.tensor_storage_size(), expected_size) | ||
| 66 | + | ||
| 67 | + def test_write_item_tensor_storage_size_for_non_tensor(self): | ||
| 68 | + write_item = WriteItem( | ||
| 69 | + index=MetadataIndex("bytes"), | ||
| 70 | + type=WriteItemType.BYTE_IO, | ||
| 71 | + ) | ||
| 72 | + | ||
| 73 | + self.assertIsNone(write_item.tensor_storage_size()) | ||
| 74 | + | ||
| 75 | + | ||
| 76 | +if __name__ == "__main__": | ||
| 77 | + run_tests() | ||