已合并
test(distributed/checkpoint): Add StorageWriter storage_meta API tests #35076
test(distributed/checkpoint): Add StorageWriter storage_meta API tests #35076
已合并
Jwerr创建于 5月8日
1 个文件变更+278-0
@@ -0,0 +1,278 @@
1+"""
2+PyTorch community lacks some torch.distributed.checkpoint.StorageWriter APIs validation, so this file is added.
3+ 
4+This file validates the following apis:
5+torch.distributed.checkpoint.StorageWriter.storage_meta
6+(extendable)
7+"""
8+ 
9+import os
10+import tempfile
11+import uuid
12+ 
13+import torch
14+import torch.distributed as dist
15+import torch.distributed.checkpoint as dcp
16+from torch.distributed.checkpoint import FileSystemWriter
17+from torch.distributed.checkpoint.default_planner import DefaultSavePlanner
18+from torch.distributed.checkpoint.metadata import Metadata, StorageMeta
19+from torch.distributed.checkpoint.storage import StorageWriter, WriteResult
20+from torch.futures import Future
21+from torch.testing._internal.distributed._shard.sharded_tensor import (
22+ ShardedTensorTestBase,
23+)
24+ 
25+from torch_npu.testing.common_distributed import skipIfUnsupportMultiNPU, with_comms
26+from torch_npu.testing.testcase import run_tests, TestCase
27+ 
28+ 
29+device_type = acc.type if (acc := torch.accelerator.current_accelerator()) else "cpu"
30+ 
31+ 
32+class _MinimalStorageWriter(StorageWriter):
33+ """Minimal subclass implementing only abstractmethods to exercise the base
34+ default storage_meta() which returns None.
35+ """
36+ 
37+ def reset(self, checkpoint_id: str | os.PathLike | None = None) -> None:
38+ return None
39+ 
40+ def set_up_storage_writer(self, is_coordinator: bool) -> None:
41+ return None
42+ 
43+ def prepare_local_plan(self, plan):
44+ return plan
45+ 
46+ def prepare_global_plan(self, plans):
47+ return plans
48+ 
49+ def write_data(self, plan, planner) -> Future[list[WriteResult]]:
50+ fut: Future[list[WriteResult]] = Future()
51+ fut.set_result([])
52+ return fut
53+ 
54+ def finish(self, metadata: Metadata, results: list[list[WriteResult]]) -> None:
55+ return None
56+ 
57+ @classmethod
58+ def validate_checkpoint_id(cls, checkpoint_id: str | os.PathLike) -> bool:
59+ return True
60+ 
61+ 
62+class TestStorageMetaSingleProcess(TestCase):
63+ """Single-process contract validation: base default, FileSystemWriter field
64+ semantics, post-save validity, empty-state-dict boundary, reset lifecycle,
65+ and SavePlanner propagation.
66+ """
67+ 
68+ # --------------------------------------------------------------------- #
69+ # Base-class default #
70+ # --------------------------------------------------------------------- #
71+ def test_base_storage_meta_callable_and_returns_none(self):
72+ self.assertTrue(callable(StorageWriter.storage_meta))
73+ writer = _MinimalStorageWriter()
74+ self.assertIsNone(writer.storage_meta())
75+ 
76+ # --------------------------------------------------------------------- #
77+ # FileSystemWriter field semantics #
78+ # --------------------------------------------------------------------- #
79+ def _assert_storage_meta_fields(self, writer: FileSystemWriter, expected_dir: str):
80+ meta = writer.storage_meta()
81+ self.assertIsInstance(meta, StorageMeta)
82+ self.assertEqual(str(meta.checkpoint_id), expected_dir)
83+ self.assertIsInstance(meta.modules, list)
84+ self.assertEqual(meta.modules, [])
85+ self.assertIsInstance(uuid.UUID(meta.save_id), uuid.UUID)
86+ 
87+ def test_filesystem_writer_storage_meta_fields(self):
88+ with tempfile.TemporaryDirectory() as temp_dir:
89+ writer = FileSystemWriter(temp_dir)
90+ self._assert_storage_meta_fields(writer, temp_dir)
91+ 
92+ def test_storage_meta_after_reset(self):
93+ with tempfile.TemporaryDirectory() as d1, tempfile.TemporaryDirectory() as d2:
94+ writer = FileSystemWriter(d1)
95+ save_id_before = writer.storage_meta().save_id
96+ 
97+ writer.reset(checkpoint_id=d2)
98+ meta_after = writer.storage_meta()
99+ 
100+ self.assertEqual(str(meta_after.checkpoint_id), d2)
101+ self.assertNotEqual(save_id_before, meta_after.save_id)
102+ 
103+ # --------------------------------------------------------------------- #
104+ # Save integration #
105+ # --------------------------------------------------------------------- #
106+ def test_storage_meta_valid_after_save(self):
107+ with tempfile.TemporaryDirectory() as temp_dir:
108+ state_dict = {"weight": torch.randn(4, 4).to(device_type)}
109+ writer = FileSystemWriter(temp_dir)
110+ dcp.save(state_dict, storage_writer=writer, no_dist=True)
111+ self._assert_storage_meta_fields(writer, temp_dir)
112+ 
113+ def test_storage_meta_empty_state_dict(self):
114+ """Boundary: writer.storage_meta() must remain valid even when no tensors
115+ are written (empty state_dict).
116+ """
117+ with tempfile.TemporaryDirectory() as temp_dir:
118+ writer = FileSystemWriter(temp_dir)
119+ dcp.save({}, storage_writer=writer, no_dist=True)
120+ self._assert_storage_meta_fields(writer, temp_dir)
121+ 
122+ def test_storage_meta_reset_between_saves(self):
123+ """save -> reset(new dir) -> save: save_id must change to guarantee
124+ uniqueness across distinct save operations, preventing metadata collisions
125+ in distributed checkpointing scenarios.
126+ """
127+ with tempfile.TemporaryDirectory() as d1, tempfile.TemporaryDirectory() as d2:
128+ writer = FileSystemWriter(d1)
129+ dcp.save(
130+ {"w": torch.randn(4, 4).to(device_type)},
131+ storage_writer=writer,
132+ no_dist=True,
133+ )
134+ save_id_1 = writer.storage_meta().save_id
135+ 
136+ writer.reset(checkpoint_id=d2)
137+ dcp.save(
138+ {"w": torch.randn(4, 4).to(device_type)},
139+ storage_writer=writer,
140+ no_dist=True,
141+ )
142+ save_id_2 = writer.storage_meta().save_id
143+ 
144+ self.assertNotEqual(save_id_1, save_id_2)
145+ 
146+ # --------------------------------------------------------------------- #
147+ # SavePlanner propagation #
148+ # --------------------------------------------------------------------- #
149+ def test_storage_meta_passed_to_planner(self):
150+ captured = {}
151+ 
152+ class _CapturingPlanner(DefaultSavePlanner):
153+ def set_up_planner(
154+ self, state_dict, storage_meta=None, is_coordinator=False
155+ ):
156+ captured["meta"] = storage_meta
157+ super().set_up_planner(state_dict, storage_meta, is_coordinator)
158+ 
159+ with tempfile.TemporaryDirectory() as temp_dir:
160+ state_dict = {"w": torch.randn(4, 4).to(device_type)}
161+ writer = FileSystemWriter(temp_dir)
162+ expected = writer.storage_meta()
163+ 
164+ dcp.save(
165+ state_dict,
166+ storage_writer=writer,
167+ planner=_CapturingPlanner(),
168+ no_dist=True,
169+ )
170+ 
171+ self.assertIsInstance(captured.get("meta"), StorageMeta)
172+ self.assertEqual(
173+ str(captured["meta"].checkpoint_id),
174+ str(expected.checkpoint_id),
175+ )
176+ self.assertEqual(captured["meta"].save_id, expected.save_id)
177+ 
178+ 
179+class TestStorageMetaDistributed(ShardedTensorTestBase):
180+ """Multi-NPU contract: checkpoint_id must be consistent across ranks,
181+ save_id must remain unique per writer instance, and storage_meta must be
182+ obtainable directly from the writer after distributed save.
183+ 
184+ NOTE: We assert writer.storage_meta() rather than Metadata.storage_meta
185+ because PyTorch 2.9+ leaves the latter None on the no_dist=True path
186+ (community Issue #177887). The writer-side API contract is the stable
187+ surface for validation.
188+ """
189+ 
190+ def tearDown(self):
191+ super().tearDown()
192+ self.destroy_pg()
193+ 
194+ def destroy_pg(self) -> None:
195+ """Best-effort cleanup: swallow exceptions to avoid masking the real
196+ test result when the process group is already in a bad state.
197+ """
198+ if dist.is_initialized():
199+ dist.barrier()
200+ dist.destroy_process_group()
201+ 
202+ @property
203+ def world_size(self) -> int:
204+ return 2
205+ 
206+ def _shared_temp_dir(self):
207+ """Yield a single checkpoint directory visible to all ranks.
208+ Rank 0 creates the directory and broadcasts its path; all ranks
209+ synchronize on teardown so rank 0 can safely remove it.
210+ """
211+ 
212+ class _SharedTempDirContext:
213+ def __init__(self):
214+ self._temp_dir_ctx = None
215+ self.temp_dir = ""
216+ 
217+ def __enter__(self):
218+ if dist.get_rank() == 0:
219+ self._temp_dir_ctx = tempfile.TemporaryDirectory()
220+ self.temp_dir = self._temp_dir_ctx.__enter__()
221+ object_list = [self.temp_dir]
222+ dist.broadcast_object_list(object_list, src=0)
223+ self.temp_dir = object_list[0]
224+ return self.temp_dir
225+ 
226+ def __exit__(self, exc_type, exc, tb):
227+ dist.barrier()
228+ if dist.get_rank() == 0 and self._temp_dir_ctx is not None:
229+ self._temp_dir_ctx.__exit__(exc_type, exc, tb)
230+ return False
231+ 
232+ return _SharedTempDirContext()
233+ 
234+ @skipIfUnsupportMultiNPU(2)
235+ @with_comms
236+ def test_storage_meta_checkpoint_id_consistency_across_ranks(self) -> None:
237+ with self._shared_temp_dir() as temp_dir:
238+ writer = FileSystemWriter(temp_dir)
239+ dcp.save({"t": torch.randn(4, 4).to(device_type)}, storage_writer=writer)
240+ 
241+ local_id = str(writer.storage_meta().checkpoint_id)
242+ gathered: list[str | None] = [None] * dist.get_world_size()
243+ dist.all_gather_object(gathered, local_id)
244+ for cid in gathered:
245+ self.assertEqual(cid, gathered[0])
246+ 
247+ @skipIfUnsupportMultiNPU(2)
248+ @with_comms
249+ def test_storage_meta_save_id_changes_after_reset_distributed(self) -> None:
250+ with self._shared_temp_dir() as temp_dir:
251+ writer = FileSystemWriter(temp_dir)
252+ dcp.save({"t": torch.randn(4, 4).to(device_type)}, storage_writer=writer)
253+ save_id_1 = writer.storage_meta().save_id
254+ 
255+ dist.barrier()
256+ writer.reset(checkpoint_id=temp_dir)
257+ dcp.save({"t": torch.randn(4, 4).to(device_type)}, storage_writer=writer)
258+ save_id_2 = writer.storage_meta().save_id
259+ 
260+ self.assertNotEqual(save_id_1, save_id_2)
261+ 
262+ @skipIfUnsupportMultiNPU(2)
263+ @with_comms
264+ def test_storage_meta_returned_in_distributed_save(self) -> None:
265+ with self._shared_temp_dir() as temp_dir:
266+ writer = FileSystemWriter(temp_dir)
267+ dcp.save(
268+ {f"rank_{dist.get_rank()}": torch.randn(4, 4).to(device_type)},
269+ storage_writer=writer,
270+ )
271+ 
272+ meta = writer.storage_meta()
273+ self.assertIsInstance(meta, StorageMeta)
274+ self.assertIsInstance(uuid.UUID(meta.save_id), uuid.UUID)
275+ 
276+ 
277+if __name__ == "__main__":
278+ run_tests()