已关闭
test(fx): add ShapeEnv API tests (self-written + upstream patch) for v2.12.0 #36523
小辉懂编程创建于 5月22日关闭于 7月4日
test(fx): add ShapeEnv API tests (self-written + upstream patch) for v2.12.0 #36523
已关闭
小辉懂编程创建于 5月22日关闭于 7月4日
共 2 个文件变更+108-0
@@ -306,5 +306,89 @@ class TestShapeEnvNPU(TestCase):
306 self.assertEqual(env.simplify((a + b) - b), a)306 self.assertEqual(env.simplify((a + b) - b), a)
307 307 
308 308 
309+device_type = (
310+ acc.type if (acc := torch.accelerator.current_accelerator()) else "cpu"
311+)
312+ 
313+ 
314+class TestSymbolicShapes(TestCase):
315+ 
316+ def test_is_accessor_node_with_call_method(self):
317+ """Verify is_accessor_node returns True for call_method nodes with NPU tensor example_value."""
318+ graph = Graph()
319+ x = graph.placeholder("x")
320+ x.meta["example_value"] = torch.randn(2, 3).to(device_type)
321+ 
322+ size_node = graph.call_method("size", args=(x, 0))
323+ self.assertTrue(is_accessor_node(size_node))
324+ 
325+ def test_is_accessor_node_with_call_function(self):
326+ """Verify is_accessor_node correctly identifies sym_size nodes vs regular add nodes."""
327+ graph = Graph()
328+ x = graph.placeholder("x")
329+ 
330+ size_node = graph.call_function(torch.ops.aten.sym_size.int, args=(x, 0))
331+ add_node = graph.call_function(torch.ops.aten.add.Tensor, args=(x, x))
332+ 
333+ self.assertTrue(is_accessor_node(size_node))
334+ self.assertFalse(is_accessor_node(add_node))
335+ 
336+ def test_is_concrete_int_with_literal_and_device_shape(self):
337+ """Verify is_concrete_int and is_symbolic for literals, NPU tensor sizes, and unbacked symints."""
338+ x = torch.randn(2, 3).to(device_type)
339+ sym_int = ShapeEnv().create_unbacked_symint()
340+ 
341+ self.assertTrue(is_concrete_int(3))
342+ self.assertTrue(is_concrete_int(x.size(0)))
343+ self.assertFalse(is_concrete_int(sym_int))
344+ self.assertFalse(is_symbolic(x.size(0)))
345+ self.assertTrue(is_symbolic(sym_int))
346+ 
347+ def test_is_concrete_float_with_literal_and_symbolic_value(self):
348+ """Verify is_concrete_float and is_symbolic for literals and unbacked symfloats."""
349+ sym_float = ShapeEnv().create_unbacked_symfloat()
350+ 
351+ self.assertTrue(is_concrete_float(1.5))
352+ self.assertFalse(is_concrete_float(sym_float))
353+ self.assertFalse(is_symbolic(1.5))
354+ self.assertTrue(is_symbolic(sym_float))
355+ 
356+ def test_is_concrete_bool_with_literal_and_symbolic_value(self):
357+ """Verify is_concrete_bool and is_symbolic for literals and unbacked symbools."""
358+ sym_bool = ShapeEnv().create_unbacked_symbool()
359+ 
360+ self.assertTrue(is_concrete_bool(True))
361+ self.assertFalse(is_concrete_bool(sym_bool))
362+ self.assertFalse(is_symbolic(True))
363+ self.assertTrue(is_symbolic(sym_bool))
364+ 
365+ def test_shape_env_get_pruned_guards(self):
366+ """Verify get_pruned_guards returns a list of guards filtered by given symints."""
367+ shape_env = ShapeEnv()
368+ sym_int = shape_env.create_unbacked_symint()
369+ 
370+ pruned_guards = shape_env.get_pruned_guards([sym_int])
371+ 
372+ self.assertIsInstance(pruned_guards, list)
373+ 
374+ def test_shape_env_ignore_fresh_unbacked_symbols(self):
375+ """Verify ignore_fresh_unbacked_symbols context manager suppresses fresh unbacked symbol registration."""
376+ shape_env = ShapeEnv()
377+ 
378+ with shape_env.ignore_fresh_unbacked_symbols():
379+ sym_int = shape_env.create_unbacked_symint()
380+ self.assertIsNotNone(sym_int)
381+ 
382+ def test_shape_env_is_unbacked_symint(self):
383+ """Verify is_unbacked_symint correctly distinguishes unbacked symbols from regular sympy symbols."""
384+ shape_env = ShapeEnv()
385+ unbacked_symint = shape_env.create_unbacked_symint()
386+ unbacked_symbol = unbacked_symint.node.expr
387+ regular_symbol = sympy.Symbol("s0", integer=True)
388+ 
389+ self.assertTrue(shape_env.is_unbacked_symint(unbacked_symbol))
390+ self.assertFalse(shape_env.is_unbacked_symint(regular_symbol))
391+ 
392+ 
309if __name__ == "__main__":393if __name__ == "__main__":
310 run_tests()394 run_tests()
@@ -0,0 +1,24 @@
1+diff --git a/test/test_proxy_tensor.py b/test/test_proxy_tensor.py
2+index 171c13b..71073ec 100644
3+--- a/test/test_proxy_tensor.py
4++++ b/test/test_proxy_tensor.py
5+@@ -1867,7 +1867,7 @@ def forward(self, x_1):
6+ if a.shape[0] != b.shape[0] * 2:
7+ raise AssertionError("a.shape[0] should equal b.shape[0] * 2")
8+ return a.cos()
9+- fx_g = make_fx(f, tracing_mode="symbolic")(torch.randn(16), torch.randn(8))
10++ fx_g = make_fx(f, tracing_mode="symbolic")(torch.randn(16).npu(), torch.randn(8).npu())
11+ from torch._dynamo.source import LocalSource
12+ self.assertExpectedInline(
13+ str(fx_g.shape_env.produce_guards(fx_placeholder_vals(fx_g), [LocalSource("a"), LocalSource("b")], ignore_static=False))
14+@@ -1943,6 +1943,10 @@ L['a'].size()[1] <= 18""")
15+
16+ # NB: Numbers are carefully chosen to avoid duck shaping from applying
17+
18++ def _trace(f, *args):
19++ inps = [torch.randn(arg).npu() for arg in args]
20++ return make_fx(f, tracing_mode="symbolic")(*inps)
21++
22+ fx_g = _trace(f, (5, 6), (5, 6))
23+ self._assert_no_guards(fx_g, 2)
24+