已合并
test: add load planner api npu tests #35094
test: add load planner api npu tests #35094
已合并
zjucn创建于 5月8日
1 个文件变更+437-0
Atest/distributed/checkpoint/test_loadplan_api.py+437-0
@@ -0,0 +1,437 @@
1+"""
2+1. PyTorch community lacks direct validation cases for some
3+ torch.distributed.checkpoint LoadPlan and LoadPlanner APIs, so this file is
4+ added.
5+ 
6+2. This file validates the following APIs:
7+ torch.distributed.checkpoint.LoadPlan
8+ torch.distributed.checkpoint.LoadPlanner
9+ torch.distributed.checkpoint.LoadPlanner.set_up_planner
10+ torch.distributed.checkpoint.LoadPlanner.create_local_plan
11+ torch.distributed.checkpoint.LoadPlanner.create_global_plan
12+ torch.distributed.checkpoint.LoadPlanner.finish_plan
13+ torch.distributed.checkpoint.LoadPlanner.load_bytes
14+ torch.distributed.checkpoint.LoadPlanner.resolve_tensor
15+ torch.distributed.checkpoint.LoadPlanner.commit_tensor
16+ (extendable)
17+"""
18+ 
19+import io
20+import tempfile
21+ 
22+from torch_npu.testing.testcase import run_tests, TestCase
23+ 
24+import torch
25+from torch.distributed.checkpoint import (
26+ FileSystemReader,
27+ FileSystemWriter,
28+ load_state_dict,
29+ save_state_dict,
30+)
31+from torch.distributed.checkpoint.default_planner import (
32+ _create_default_local_metadata,
33+ DefaultLoadPlanner,
34+ DefaultSavePlanner,
35+)
36+from torch.distributed.checkpoint.metadata import (
37+ ChunkStorageMetadata,
38+ Metadata,
39+ MetadataIndex,
40+ TensorProperties,
41+ TensorStorageMetadata,
42+)
43+from torch.distributed.checkpoint.planner import LoadItemType, LoadPlan
44+from torch.distributed.checkpoint.planner_helpers import _create_read_item_for_tensor
45+ 
46+ 
47+device_type = acc.type if (acc := torch.accelerator.current_accelerator()) else "cpu"
48+ 
49+ 
50+def _make_tensor_read_item(
51+ fqn="tensor",
52+ dest_offsets=(0, 0),
53+ lengths=(2, 2),
54+):
55+ zero_offsets = [0] * len(lengths)
56+ return _create_read_item_for_tensor(
57+ dest_index=MetadataIndex(fqn, zero_offsets),
58+ dest_offsets=dest_offsets,
59+ storage_index=MetadataIndex(fqn, zero_offsets),
60+ storage_offsets=zero_offsets,
61+ lengths=lengths,
62+ )
63+ 
64+ 
65+class MaterializeOnCpuLoadPlanner(DefaultLoadPlanner):
66+ """Planner that verifies commit_tensor can move loaded CPU data back to NPU."""
67+ 
68+ def __init__(self):
69+ super().__init__()
70+ self.resolved_tensors = []
71+ self.committed_tensors = []
72+ 
73+ def resolve_tensor(self, read_item):
74+ target = super().resolve_tensor(read_item)
75+ resolved = torch.empty_like(target, device="cpu")
76+ self.resolved_tensors.append(resolved)
77+ return resolved
78+ 
79+ def commit_tensor(self, read_item, tensor):
80+ target = super().resolve_tensor(read_item)
81+ target.copy_(tensor.to(target.device))
82+ self.committed_tensors.append((read_item.dest_index.fqn, tensor.device.type))
83+ 
84+ 
85+class PlanDataLoadPlanner(DefaultLoadPlanner):
86+ """Planner that verifies local/global/finish plan customization."""
87+ 
88+ def __init__(self):
89+ super().__init__()
90+ self.finished_storage_data = None
91+ self.finished_planner_data = None
92+ 
93+ def create_local_plan(self):
94+ plan = super().create_local_plan()
95+ return LoadPlan(plan.items, planner_data={"local_plan": True})
96+ 
97+ def create_global_plan(self, global_plan):
98+ return [
99+ LoadPlan(
100+ plan.items,
101+ storage_data={"storage_plan": index},
102+ planner_data={"global_plan": plan.planner_data},
103+ )
104+ for index, plan in enumerate(global_plan)
105+ ]
106+ 
107+ def finish_plan(self, central_plan):
108+ self.finished_storage_data = central_plan.storage_data
109+ self.finished_planner_data = central_plan.planner_data
110+ return central_plan
111+ 
112+ 
113+class TestLoadPlanApi(TestCase):
114+ def test_default_load_planner_local_global_finish_plan(self):
115+ state_dict = {
116+ "tensor": torch.zeros(3, 4).to(device_type),
117+ "bytes": ["old"],
118+ }
119+ metadata_state_dict = {
120+ "tensor": torch.ones(3, 4),
121+ "bytes": ["new"],
122+ }
123+ metadata = _create_default_local_metadata(metadata_state_dict)
124+ 
125+ planner = DefaultLoadPlanner()
126+ planner.set_up_planner(state_dict, metadata, is_coordinator=True)
127+ local_plan = planner.create_local_plan()
128+ 
129+ self.assertIsInstance(local_plan, LoadPlan)
130+ self.assertEqual(2, len(local_plan.items))
131+ 
132+ tensor_item = next(
133+ item for item in local_plan.items if item.dest_index.fqn == "tensor"
134+ )
135+ bytes_item = next(
136+ item for item in local_plan.items if item.dest_index.fqn == "bytes"
137+ )
138+ 
139+ self.assertEqual(LoadItemType.TENSOR, tensor_item.type)
140+ self.assertEqual(torch.Size([0, 0]), tensor_item.dest_offsets)
141+ self.assertEqual(torch.Size([0, 0]), tensor_item.storage_offsets)
142+ self.assertEqual(torch.Size([3, 4]), tensor_item.lengths)
143+ self.assertEqual(LoadItemType.BYTE_IO, bytes_item.type)
144+ self.assertEqual(MetadataIndex("bytes"), bytes_item.dest_index)
145+ 
146+ global_plan = planner.create_global_plan([local_plan])
147+ self.assertEqual([local_plan], global_plan)
148+ self.assertEqual(local_plan, planner.finish_plan(global_plan[0]))
149+ 
150+ def test_default_load_planner_creates_multiple_tensor_read_items(self):
151+ state_dict = {"tensor": torch.zeros(8).to(device_type)}
152+ metadata = Metadata(
153+ state_dict_metadata={
154+ "tensor": TensorStorageMetadata(
155+ properties=TensorProperties.create_from_tensor(torch.empty(8)),
156+ size=torch.Size([8]),
157+ chunks=[
158+ ChunkStorageMetadata(
159+ offsets=torch.Size([0]),
160+ sizes=torch.Size([4]),
161+ ),
162+ ChunkStorageMetadata(
163+ offsets=torch.Size([4]),
164+ sizes=torch.Size([4]),
165+ ),
166+ ],
167+ ),
168+ },
169+ )
170+ 
171+ planner = DefaultLoadPlanner()
172+ planner.set_up_planner(state_dict, metadata)
173+ local_plan = planner.create_local_plan()
174+ 
175+ self.assertEqual(2, len(local_plan.items))
176+ low_item = next(
177+ item for item in local_plan.items if item.dest_offsets == torch.Size([0])
178+ )
179+ high_item = next(
180+ item for item in local_plan.items if item.dest_offsets == torch.Size([4])
181+ )
182+ 
183+ self.assertEqual(LoadItemType.TENSOR, low_item.type)
184+ self.assertEqual(MetadataIndex("tensor", torch.Size([0])), low_item.dest_index)
185+ self.assertEqual(
186+ MetadataIndex("tensor", torch.Size([0])),
187+ low_item.storage_index,
188+ )
189+ self.assertEqual(torch.Size([0]), low_item.storage_offsets)
190+ self.assertEqual(torch.Size([4]), low_item.lengths)
191+ 
192+ self.assertEqual(LoadItemType.TENSOR, high_item.type)
193+ self.assertEqual(
194+ MetadataIndex("tensor", torch.Size([0])),
195+ high_item.dest_index,
196+ )
197+ self.assertEqual(
198+ MetadataIndex("tensor", torch.Size([4])),
199+ high_item.storage_index,
200+ )
201+ self.assertEqual(torch.Size([0]), high_item.storage_offsets)
202+ self.assertEqual(torch.Size([4]), high_item.lengths)
203+ 
204+ def test_default_load_planner_strict_and_partial_load(self):
205+ metadata = _create_default_local_metadata({"tensor": torch.ones(2, 2)})
206+ state_dict = {
207+ "tensor": torch.zeros(2, 2).to(device_type),
208+ "missing": torch.zeros(2, 2).to(device_type),
209+ }
210+ 
211+ strict_planner = DefaultLoadPlanner(allow_partial_load=False)
212+ strict_planner.set_up_planner(state_dict, metadata)
213+ with self.assertRaisesRegex(RuntimeError, "Missing key in checkpoint"):
214+ strict_planner.create_local_plan()
215+ 
216+ partial_planner = DefaultLoadPlanner(allow_partial_load=True)
217+ partial_planner.set_up_planner(state_dict, metadata)
218+ partial_plan = partial_planner.create_local_plan()
219+ self.assertEqual(1, len(partial_plan.items))
220+ self.assertEqual("tensor", partial_plan.items[0].dest_index.fqn)
221+ 
222+ def test_default_load_planner_size_mismatch(self):
223+ metadata = _create_default_local_metadata({"tensor": torch.ones(2, 2)})
224+ state_dict = {"tensor": torch.zeros(3, 2).to(device_type)}
225+ 
226+ planner = DefaultLoadPlanner()
227+ planner.set_up_planner(state_dict, metadata)
228+ with self.assertRaisesRegex(ValueError, "Size mismatch"):
229+ planner.create_local_plan()
230+ 
231+ def test_resolve_tensor_returns_npu_narrow_view(self):
232+ state_dict = {"tensor": torch.zeros(4, 5).to(device_type)}
233+ metadata = _create_default_local_metadata({"tensor": torch.ones(4, 5)})
234+ read_item = _make_tensor_read_item(
235+ dest_offsets=[1, 2],
236+ lengths=[2, 2],
237+ )
238+ 
239+ planner = DefaultLoadPlanner()
240+ planner.set_up_planner(state_dict, metadata)
241+ target_tensor = planner.resolve_tensor(read_item)
242+ 
243+ self.assertEqual(device_type, target_tensor.device.type)
244+ self.assertEqual(torch.Size([2, 2]), target_tensor.size())
245+ 
246+ target_tensor.copy_(torch.full((2, 2), 7.0))
247+ planner.commit_tensor(read_item, target_tensor)
248+ 
249+ expected = torch.zeros(4, 5)
250+ expected[1:3, 2:4] = 7.0
251+ self.assertEqual(expected, state_dict["tensor"].cpu())
252+ 
253+ def test_resolve_tensor_handles_non_contiguous_npu_target(self):
254+ npu_target = torch.zeros(5, 4).to(device_type).transpose(0, 1)
255+ self.assertFalse(npu_target.is_contiguous())
256+ state_dict = {"tensor": npu_target}
257+ metadata = _create_default_local_metadata({"tensor": torch.ones(4, 5)})
258+ read_item = _make_tensor_read_item(
259+ dest_offsets=[1, 1],
260+ lengths=[2, 3],
261+ )
262+ 
263+ planner = DefaultLoadPlanner()
264+ planner.set_up_planner(state_dict, metadata)
265+ target_tensor = planner.resolve_tensor(read_item)
266+ 
267+ self.assertEqual(device_type, target_tensor.device.type)
268+ self.assertEqual(torch.Size([2, 3]), target_tensor.size())
269+ target_tensor.copy_(torch.full((2, 3), 5.0))
270+ planner.commit_tensor(read_item, target_tensor)
271+ 
272+ expected = torch.zeros(4, 5)
273+ expected[1:3, 1:4] = 5.0
274+ self.assertEqual(expected, state_dict["tensor"].cpu())
275+ 
276+ def test_load_bytes_updates_flattened_original_state_dict(self):
277+ state_dict = {
278+ "nested": {
279+ "bytes": b"old",
280+ }
281+ }
282+ metadata = _create_default_local_metadata({"nested.bytes": b"new"})
283+ 
284+ planner = DefaultLoadPlanner()
285+ planner.set_up_planner(state_dict, metadata)
286+ plan = planner.create_local_plan()
287+ read_item = next(
288+ item for item in plan.items if item.dest_index.fqn == "nested.bytes"
289+ )
290+ 
291+ value = io.BytesIO()
292+ torch.save({"loaded": (1, 2, 3)}, value)
293+ value.seek(0)
294+ 
295+ planner.load_bytes(read_item, value)
296+ self.assertEqual({"loaded": (1, 2, 3)}, state_dict["nested"]["bytes"])
297+ 
298+ def test_load_bytes_updates_unflattened_state_dict(self):
299+ state_dict = {"payload": b"old"}
300+ metadata = _create_default_local_metadata({"payload": b"new"})
301+ 
302+ planner = DefaultLoadPlanner(
303+ flatten_state_dict=False,
304+ flatten_sharded_tensors=False,
305+ )
306+ planner.set_up_planner(state_dict, metadata)
307+ plan = planner.create_local_plan()
308+ read_item = next(
309+ item for item in plan.items if item.dest_index.fqn == "payload"
310+ )
311+ 
312+ value = io.BytesIO()
313+ torch.save({"loaded": (4, 5, 6)}, value)
314+ value.seek(0)
315+ 
316+ planner.load_bytes(read_item, value)
317+ self.assertEqual({"loaded": (4, 5, 6)}, state_dict["payload"])
318+ 
319+ 
320+class TestLoadPlannerNpuIntegration(TestCase):
321+ def test_load_state_dict_accepts_custom_plan_data(self):
322+ with tempfile.TemporaryDirectory() as checkpoint_dir:
323+ state_dict_to_save = {
324+ "tensor": torch.arange(4, dtype=torch.float32)
325+ .reshape(2, 2)
326+ .to(device_type),
327+ }
328+ save_state_dict(
329+ state_dict=state_dict_to_save,
330+ storage_writer=FileSystemWriter(checkpoint_dir),
331+ planner=DefaultSavePlanner(),
332+ no_dist=True,
333+ )
334+ 
335+ state_dict_to_load = {"tensor": torch.zeros(2, 2).to(device_type)}
336+ planner = PlanDataLoadPlanner()
337+ load_state_dict(
338+ state_dict=state_dict_to_load,
339+ storage_reader=FileSystemReader(checkpoint_dir),
340+ planner=planner,
341+ no_dist=True,
342+ )
343+ 
344+ self.assertEqual(
345+ state_dict_to_save["tensor"].cpu(), state_dict_to_load["tensor"].cpu()
346+ )
347+ self.assertEqual({"storage_plan": 0}, planner.finished_storage_data)
348+ self.assertEqual(
349+ {"global_plan": {"local_plan": True}},
350+ planner.finished_planner_data,
351+ )
352+ 
353+ def test_filesystem_metadata_version_when_supported(self):
354+ with tempfile.TemporaryDirectory() as checkpoint_dir:
355+ state_dict_to_save = {
356+ "tensor": torch.arange(4, dtype=torch.float32)
357+ .reshape(2, 2)
358+ .to(device_type),
359+ }
360+ save_state_dict(
361+ state_dict=state_dict_to_save,
362+ storage_writer=FileSystemWriter(checkpoint_dir),
363+ planner=DefaultSavePlanner(),
364+ no_dist=True,
365+ )
366+ 
367+ metadata = FileSystemReader(checkpoint_dir).read_metadata()
368+ 
369+ self.assertIsInstance(metadata, Metadata)
370+ if hasattr(metadata, "version"):
371+ from torch.distributed.checkpoint.filesystem import CURRENT_DCP_VERSION
372+ 
373+ self.assertEqual(CURRENT_DCP_VERSION, metadata.version)
374+ 
375+ def test_custom_commit_tensor_materializes_cpu_tensor_to_npu(self):
376+ with tempfile.TemporaryDirectory() as checkpoint_dir:
377+ state_dict_to_save = {
378+ "tensor": torch.arange(6, dtype=torch.float32)
379+ .reshape(2, 3)
380+ .to(device_type),
381+ }
382+ save_state_dict(
383+ state_dict=state_dict_to_save,
384+ storage_writer=FileSystemWriter(checkpoint_dir),
385+ planner=DefaultSavePlanner(),
386+ no_dist=True,
387+ )
388+ 
389+ state_dict_to_load = {"tensor": torch.zeros(2, 3).to(device_type)}
390+ planner = MaterializeOnCpuLoadPlanner()
391+ load_state_dict(
392+ state_dict=state_dict_to_load,
393+ storage_reader=FileSystemReader(checkpoint_dir),
394+ planner=planner,
395+ no_dist=True,
396+ )
397+ 
398+ self.assertEqual(
399+ state_dict_to_save["tensor"].cpu(), state_dict_to_load["tensor"].cpu()
400+ )
401+ self.assertEqual(1, len(planner.resolved_tensors))
402+ self.assertEqual("cpu", planner.resolved_tensors[0].device.type)
403+ self.assertEqual([("tensor", "cpu")], planner.committed_tensors)
404+ 
405+ def test_filesystem_load_tensor_and_bytes_to_npu_state_dict(self):
406+ with tempfile.TemporaryDirectory() as checkpoint_dir:
407+ original_tensor = (
408+ torch.arange(20, dtype=torch.float32).reshape(4, 5).to(device_type)
409+ )
410+ state_dict_to_save = {
411+ "tensor": original_tensor,
412+ "payload": ["step", 3, "ok"],
413+ }
414+ save_state_dict(
415+ state_dict=state_dict_to_save,
416+ storage_writer=FileSystemWriter(checkpoint_dir),
417+ planner=DefaultSavePlanner(),
418+ no_dist=True,
419+ )
420+ 
421+ state_dict_to_load = {
422+ "tensor": torch.full((4, 5), -1.0).to(device_type),
423+ "payload": [],
424+ }
425+ load_state_dict(
426+ state_dict=state_dict_to_load,
427+ storage_reader=FileSystemReader(checkpoint_dir),
428+ planner=DefaultLoadPlanner(),
429+ no_dist=True,
430+ )
431+ 
432+ self.assertEqual(original_tensor.cpu(), state_dict_to_load["tensor"].cpu())
433+ self.assertEqual(["step", 3, "ok"], state_dict_to_load["payload"])
434+ 
435+ 
436+if __name__ == "__main__":
437+ run_tests()