已合并
test(distributed/checkpoint): Add StorageWriter storage_meta API tests #35076
Jwerr创建于 5月8日
test(distributed/checkpoint): Add StorageWriter storage_meta API tests #35076
已合并
共 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 235 | + | ||
| 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 | + | ||
| 248 | + | ||
| 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 | + | ||
| 263 | + | ||
| 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() | ||