已关闭
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
已关闭
共 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 | + | ||
| 297 | if __name__ == "__main__": | 381 | if __name__ == "__main__": |
| 298 | run_tests() | 382 | run_tests() |
| @@ -1,22 +1,24 @@ | |||
| 1 | diff --git a/test/test_proxy_tensor.py b/test/test_proxy_tensor.py | 1 | diff --git a/test/test_proxy_tensor.py b/test/test_proxy_tensor.py |
| 2 | -index 26131d5..02704bb 100644 | 2 | +index 26131d5..e24bb47 100644 |
| 3 | --- a/test/test_proxy_tensor.py | 3 | --- a/test/test_proxy_tensor.py |
| 4 | +++ b/test/test_proxy_tensor.py | 4 | +++ b/test/test_proxy_tensor.py |
| 5 | -@@ -32,7 +32,7 @@ import re | 5 | +@@ -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 | + L['a'].size()[1] <= 18""") | ||
| 6 | 15 | ||
| 7 | - import functools | 16 | + # 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 | - 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 | ||