已关闭
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
已关闭
共 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 | + | ||
| 309 | if __name__ == "__main__": | 393 | if __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 | + 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 | + 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 | + | ||