已关闭
test(fx): add ShapeEnv API tests (self-written + upstream patch) for v2.7.1 #36524
小辉懂编程创建于 5月22日关闭于 7月4日
test(fx): add ShapeEnv API tests (self-written + upstream patch) for v2.7.1 #36524
已关闭
小辉懂编程创建于 5月22日关闭于 7月4日
共 2 个文件变更+102-16
@@ -294,5 +294,89 @@ class TestShapeEnvNPU(TestCase):
294 self.assertEqual(env.simplify((a + b) - b), a)294 self.assertEqual(env.simplify((a + b) - b), a)
295 295 
296 296 
297+device_type = (
298+ acc.type if (acc := torch.accelerator.current_accelerator()) else "cpu"
299+)
300+ 
301+ 
302+class TestSymbolicShapes(TestCase):
303+ 
304+ def test_is_accessor_node_with_call_method(self):
305+ """Verify is_accessor_node returns True for call_method nodes with NPU tensor example_value."""
306+ graph = Graph()
307+ x = graph.placeholder("x")
308+ x.meta["example_value"] = torch.randn(2, 3).to(device_type)
309+ 
310+ size_node = graph.call_method("size", args=(x, 0))
311+ self.assertTrue(is_accessor_node(size_node))
312+ 
313+ def test_is_accessor_node_with_call_function(self):
314+ """Verify is_accessor_node correctly identifies sym_size nodes vs regular add nodes."""
315+ graph = Graph()
316+ x = graph.placeholder("x")
317+ 
318+ size_node = graph.call_function(torch.ops.aten.sym_size.int, args=(x, 0))
319+ add_node = graph.call_function(torch.ops.aten.add.Tensor, args=(x, x))
320+ 
321+ self.assertTrue(is_accessor_node(size_node))
322+ self.assertFalse(is_accessor_node(add_node))
323+ 
324+ def test_is_concrete_int_with_literal_and_device_shape(self):
325+ """Verify is_concrete_int and is_symbolic for literals, NPU tensor sizes, and unbacked symints."""
326+ x = torch.randn(2, 3).to(device_type)
327+ sym_int = ShapeEnv().create_unbacked_symint()
328+ 
329+ self.assertTrue(is_concrete_int(3))
330+ self.assertTrue(is_concrete_int(x.size(0)))
331+ self.assertFalse(is_concrete_int(sym_int))
332+ self.assertFalse(is_symbolic(x.size(0)))
333+ self.assertTrue(is_symbolic(sym_int))
334+ 
335+ def test_is_concrete_float_with_literal_and_symbolic_value(self):
336+ """Verify is_concrete_float and is_symbolic for literals and unbacked symfloats."""
337+ sym_float = ShapeEnv().create_unbacked_symfloat()
338+ 
339+ self.assertTrue(is_concrete_float(1.5))
340+ self.assertFalse(is_concrete_float(sym_float))
341+ self.assertFalse(is_symbolic(1.5))
342+ self.assertTrue(is_symbolic(sym_float))
343+ 
344+ def test_is_concrete_bool_with_literal_and_symbolic_value(self):
345+ """Verify is_concrete_bool and is_symbolic for literals and unbacked symbools."""
346+ sym_bool = ShapeEnv().create_unbacked_symbool()
347+ 
348+ self.assertTrue(is_concrete_bool(True))
349+ self.assertFalse(is_concrete_bool(sym_bool))
350+ self.assertFalse(is_symbolic(True))
351+ self.assertTrue(is_symbolic(sym_bool))
352+ 
353+ def test_shape_env_get_pruned_guards(self):
354+ """Verify get_pruned_guards returns a list of guards filtered by given symints."""
355+ shape_env = ShapeEnv()
356+ sym_int = shape_env.create_unbacked_symint()
357+ 
358+ pruned_guards = shape_env.get_pruned_guards([sym_int])
359+ 
360+ self.assertIsInstance(pruned_guards, list)
361+ 
362+ def test_shape_env_ignore_fresh_unbacked_symbols(self):
363+ """Verify ignore_fresh_unbacked_symbols context manager suppresses fresh unbacked symbol registration."""
364+ shape_env = ShapeEnv()
365+ 
366+ with shape_env.ignore_fresh_unbacked_symbols():
367+ sym_int = shape_env.create_unbacked_symint()
368+ self.assertIsNotNone(sym_int)
369+ 
370+ def test_shape_env_is_unbacked_symint(self):
371+ """Verify is_unbacked_symint correctly distinguishes unbacked symbols from regular sympy symbols."""
372+ shape_env = ShapeEnv()
373+ unbacked_symint = shape_env.create_unbacked_symint()
374+ unbacked_symbol = unbacked_symint.node.expr
375+ regular_symbol = sympy.Symbol("s0", integer=True)
376+ 
377+ self.assertTrue(shape_env.is_unbacked_symint(unbacked_symbol))
378+ self.assertFalse(shape_env.is_unbacked_symint(regular_symbol))
379+ 
380+ 
297if __name__ == "__main__":381if __name__ == "__main__":
298 run_tests()382 run_tests()
@@ -1,22 +1,24 @@
1diff --git a/test/test_proxy_tensor.py b/test/test_proxy_tensor.py1diff --git a/test/test_proxy_tensor.py b/test/test_proxy_tensor.py
2-index 26131d5..02704bb 1006442+index 26131d5..e24bb47 100644
3--- a/test/test_proxy_tensor.py3--- a/test/test_proxy_tensor.py
4+++ b/test/test_proxy_tensor.py4+++ b/test/test_proxy_tensor.py
5-@@ -32,7 +32,7 @@ import re5+@@ -1821,7 +1821,7 @@ def forward(self, x_1):
6+ def f(a, b):
7+ assert a.shape[0] == 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+@@ -1889,6 +1889,10 @@ L['a'].size()[1] <= 18""")
6 15
7- import functools16+ # NB: Numbers are carefully chosen to avoid duck shaping from applying
8- import itertools
9--
10-+from torch_npu.contrib import transfer_to_npu
11- aten = torch.ops.aten
12-
13- HAS_CUDA = torch.cuda.is_available()
14-@@ -2148,7 +2148,7 @@ class TestProxyTensorOpInfo(TestCase):
15- _test_make_fx_helper(self, device, dtype, op, "symbolic", out=True)
16-
17-
18--only_for = ("cpu")
19-+only_for = ("cpu",)
20- instantiate_device_type_tests(TestProxyTensorOpInfo, globals(), only_for=only_for)
21 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)
22 24