已合并
test: 新增 GraphPickler NPU 适配验证与统一运行脚本 #34859
test: 新增 GraphPickler NPU 适配验证与统一运行脚本 #34859
已合并
yuhongming-2026创建于 4月30日
1 个文件变更+108-0
Atest/fx/test_fx_graph_pickler.py+108-0
@@ -0,0 +1,108 @@
1+#!/usr/bin/env python3
2+# Owner(s): ["module: fx"]
3+ 
4+import inspect
5+ 
6+import torch
7+import torch.fx as fx
8+from torch.fx import _graph_pickler
9+from torch.testing._internal.common_utils import run_tests, TestCase
10+ 
11+ 
12+def _npu_available() -> bool:
13+ return hasattr(torch, "npu") and torch.npu.is_available()
14+ 
15+ 
16+def _build_test_graph(device: str = "cpu") -> tuple[fx.GraphModule, torch.Tensor]:
17+ class SimpleModel(torch.nn.Module):
18+ def __init__(self):
19+ super().__init__()
20+ self.conv = torch.nn.Conv2d(3, 8, 3, padding=1)
21+ self.bn = torch.nn.BatchNorm2d(8)
22+ 
23+ def forward(self, x):
24+ return torch.relu(self.bn(self.conv(x)))
25+ 
26+ model = SimpleModel().eval().to(device)
27+ input_tensor = torch.randn(2, 3, 8, 8, device=device)
28+ traced = fx.symbolic_trace(model)
29+ return traced, input_tensor
30+ 
31+ 
32+def _node_kinds(graph: fx.Graph) -> list[tuple[str, str]]:
33+ return [(node.op, str(node.target)) for node in graph.nodes]
34+ 
35+ 
36+def _extract_graph(loaded_obj: object) -> fx.Graph:
37+ if isinstance(loaded_obj, fx.Graph):
38+ return loaded_obj
39+ return loaded_obj.graph # type: ignore[union-attr]
40+ 
41+ 
42+def _build_options():
43+ options_cls = getattr(_graph_pickler, "Options", None)
44+ if options_cls is None:
45+ return None
46+ 
47+ signature = inspect.signature(options_cls)
48+ has_required_parameter = any(
49+ name != "self" and param.default is inspect._empty
50+ for name, param in signature.parameters.items()
51+ )
52+ if has_required_parameter:
53+ return None
54+ return options_cls()
55+ 
56+ 
57+def _loads_payload(payload: bytes):
58+ loads_signature = inspect.signature(_graph_pickler.GraphPickler.loads)
59+ if "fake_mode" in loads_signature.parameters:
60+ from torch._subclasses.fake_tensor import FakeTensorMode
61+ 
62+ return _graph_pickler.GraphPickler.loads(payload, fake_mode=FakeTensorMode())
63+ return _graph_pickler.GraphPickler.loads(payload)
64+ 
65+ 
66+class TestFxGraphPickler(TestCase):
67+ def test_graphpickler_dumps_and_loads_cpu(self):
68+ self.assertTrue(hasattr(_graph_pickler, "GraphPickler"))
69+ self.assertTrue(hasattr(_graph_pickler.GraphPickler, "dumps"))
70+ self.assertTrue(hasattr(_graph_pickler.GraphPickler, "loads"))
71+ 
72+ traced, _ = _build_test_graph("cpu")
73+ payload = _graph_pickler.GraphPickler.dumps(traced)
74+ loaded_obj = _loads_payload(payload)
75+ loaded_graph = _extract_graph(loaded_obj)
76+ self.assertEqual(_node_kinds(traced.graph), _node_kinds(loaded_graph))
77+ 
78+ def test_graphpickler_dumps_and_loads_npu(self):
79+ if not _npu_available():
80+ self.skipTest("NPU not available")
81+ 
82+ traced, _ = _build_test_graph("npu")
83+ payload = _graph_pickler.GraphPickler.dumps(traced)
84+ loaded_obj = _loads_payload(payload)
85+ loaded_graph = _extract_graph(loaded_obj)
86+ self.assertEqual(_node_kinds(traced.graph), _node_kinds(loaded_graph))
87+ 
88+ def test_graphpickler_options(self):
89+ options = _build_options()
90+ if options is None:
91+ self.skipTest(
92+ "torch.fx._graph_pickler.Options unavailable or requires args"
93+ )
94+ 
95+ traced, _ = _build_test_graph("cpu")
96+ dumps_signature = inspect.signature(_graph_pickler.GraphPickler.dumps)
97+ if "options" in dumps_signature.parameters:
98+ payload = _graph_pickler.GraphPickler.dumps(traced, options=options)
99+ else:
100+ payload = _graph_pickler.GraphPickler.dumps(traced, options)
101+ 
102+ self.assertIsInstance(payload, (bytes, bytearray))
103+ loaded_obj = _loads_payload(payload)
104+ self.assertTrue(isinstance(loaded_obj, (fx.Graph, fx.GraphModule)))
105+ 
106+ 
107+if __name__ == "__main__":
108+ run_tests()