已合并
test(fx): add NPU coverage for PropagateUnbackedSymInts and rebind_unbacked #37266
lihaofei-2026创建于 5月31日
test(fx): add NPU coverage for PropagateUnbackedSymInts and rebind_unbacked #37266
已合并
lihaofei-2026创建于 5月31日
已删除 :symbolic-shapes-propagate-unbacked-symints-v2.10.0合入到Ascend/pytorchv2.10.0
1 个文件变更+132-0
@@ -0,0 +1,132 @@
1+"""
2+Add validation cases for torch.fx.experimental.symbolic_shapes.PropagateUnbackedSymInts on NPU.
3+ 
4+1. PyTorch community lacks sufficient and direct API validations for PropagateUnbackedSymInts,
5+ so this file is added.
6+2. This file validates PropagateUnbackedSymInts.run, PropagateUnbackedSymInts.run_node,
7+ PropagateUnbackedSymInts.placeholder, PropagateUnbackedSymInts.output, and rebind_unbacked.
8+"""
9+ 
10+import torch
11+from torch.testing._internal.common_utils import TestCase, run_tests
12+from torch._dynamo.utils import detect_fake_mode
13+from torch.fx.experimental.symbolic_shapes import PropagateUnbackedSymInts, rebind_unbacked
14+ 
15+device_type = acc.type if (acc := torch.accelerator.current_accelerator()) else "cpu"
16+torch.zeros(3, 4).to(device_type)
17+ 
18+ 
19+class TestPropagateUnbackedSymInts(TestCase):
20+ 
21+ def test_propagate_unbacked_symints_run(self):
22+ """Test PropagateUnbackedSymInts.run with NPU tensor."""
23+ 
24+ class M(torch.nn.Module):
25+ def forward(self, x: torch.Tensor):
26+ return torch.nonzero(x)
27+ 
28+ inp = (torch.tensor([1, 0, 1, 0]).to(device_type),)
29+ gm = torch.export.export(M(), inp, strict=True).module()
30+ fake_inputs = [
31+ node.meta.get("val") for node in gm.graph.nodes if node.op == "placeholder"
32+ ]
33+ fake_mode = detect_fake_mode(fake_inputs)
34+ with fake_mode:
35+ result = PropagateUnbackedSymInts(gm).run(*fake_inputs)
36+ self.assertIsNotNone(result)
37+ 
38+ def test_propagate_unbacked_symints_run_node(self):
39+ """Test PropagateUnbackedSymInts.run_node with NPU tensor."""
40+ 
41+ class RunNodeCapturingInterpreter(PropagateUnbackedSymInts):
42+ def __init__(self, *args, **kwargs):
43+ super().__init__(*args, **kwargs)
44+ self.captured_results = {}
45+ 
46+ def run_node(self, n):
47+ result = super().run_node(n)
48+ self.captured_results[n] = result
49+ return result
50+ 
51+ class M(torch.nn.Module):
52+ def forward(self, x: torch.Tensor):
53+ return torch.nonzero(x)
54+ 
55+ inp = (torch.tensor([1, 0, 1, 0]).to(device_type),)
56+ gm = torch.export.export(M(), inp, strict=True).module()
57+ fake_inputs = [
58+ node.meta.get("val") for node in gm.graph.nodes if node.op == "placeholder"
59+ ]
60+ fake_mode = detect_fake_mode(fake_inputs)
61+ with fake_mode:
62+ interpreter = RunNodeCapturingInterpreter(gm)
63+ result = interpreter.run(*fake_inputs)
64+ self.assertIsNotNone(result)
65+ for node in gm.graph.nodes:
66+ if node.op == "call_function":
67+ self.assertIn(node, interpreter.captured_results)
68+ 
69+ def test_propagate_unbacked_symints_placeholder(self):
70+ """Test PropagateUnbackedSymInts.placeholder with NPU tensor."""
71+ 
72+ class M(torch.nn.Module):
73+ def forward(self, x: torch.Tensor):
74+ return torch.nonzero(x)
75+ 
76+ inp = (torch.tensor([1, 0, 1, 0]).to(device_type),)
77+ gm = torch.export.export(M(), inp, strict=True).module()
78+ fake_inputs = [
79+ node.meta.get("val") for node in gm.graph.nodes if node.op == "placeholder"
80+ ]
81+ fake_mode = detect_fake_mode(fake_inputs)
82+ with fake_mode:
83+ interpreter = PropagateUnbackedSymInts(gm)
84+ interpreter.args_iter = iter(fake_inputs)
85+ for node in gm.graph.nodes:
86+ if node.op == "placeholder":
87+ result = interpreter.placeholder(node.target, node.args, node.kwargs)
88+ self.assertIsNotNone(result)
89+ 
90+ def test_propagate_unbacked_symints_output(self):
91+ """Test PropagateUnbackedSymInts.output with NPU tensor."""
92+ 
93+ class M(torch.nn.Module):
94+ def forward(self, x: torch.Tensor):
95+ return torch.nonzero(x)
96+ 
97+ inp = (torch.tensor([1, 0, 1, 0]).to(device_type),)
98+ gm = torch.export.export(M(), inp, strict=True).module()
99+ fake_inputs = [
100+ node.meta.get("val") for node in gm.graph.nodes if node.op == "placeholder"
101+ ]
102+ fake_mode = detect_fake_mode(fake_inputs)
103+ with fake_mode:
104+ interpreter = PropagateUnbackedSymInts(gm)
105+ interpreter.run(*fake_inputs)
106+ for node in gm.graph.nodes:
107+ if node.op == "output":
108+ result = interpreter.output(node.target, node.args, node.kwargs)
109+ self.assertIsNotNone(result)
110+ 
111+ def test_rebind_unbacked(self):
112+ """Test rebind_unbacked with NPU tensor."""
113+ 
114+ class M(torch.nn.Module):
115+ def forward(self, x: torch.Tensor):
116+ return torch.nonzero(x)
117+ 
118+ inp = (torch.tensor([1, 0, 1, 0]).to(device_type),)
119+ gm = torch.export.export(M(), inp, strict=True).module()
120+ fake_inputs = [
121+ node.meta.get("val") for node in gm.graph.nodes if node.op == "placeholder"
122+ ]
123+ fake_mode = detect_fake_mode(fake_inputs)
124+ shape_prop_gm = torch.fx.passes.shape_prop.ShapeProp(
125+ gm=gm, fake_mode=fake_mode
126+ )
127+ shape_prop_gm.propagate(*fake_inputs)
128+ self.assertEqual(len(fake_mode.shape_env.pending_fresh_unbacked_symbols), 0)
129+ 
130+ 
131+if __name__ == "__main__":
132+ run_tests()