已合并
test(distributed): add test for WriteItem.tensor_storage_size #35505
test(distributed): add test for WriteItem.tensor_storage_size #35505
已合并
Flipped创建于 5月13日
1 个文件变更+77-0
Atest/distributed/checkpoint/test_planner_api.py+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()