已合并
test: 补齐torch.fx.experimental.symbolic_shapes NPU 适配验证与统一运行脚本 #36871
cuiyunhao-2026创建于 5月27日
test: 补齐torch.fx.experimental.symbolic_shapes NPU 适配验证与统一运行脚本 #36871
已合并
共 1 个文件变更+84-0
| @@ -11,6 +11,7 @@ import sympy | |||
| 11 | import torch | 11 | import torch |
| 12 | import torch_npu | 12 | import torch_npu |
| 13 | from torch import nn | 13 | from torch import nn |
| 14 | +from torch._guards import ShapeGuard, SLoc | ||
| 14 | from torch._subclasses.fake_tensor import FakeTensorMode | 15 | from torch._subclasses.fake_tensor import FakeTensorMode |
| 15 | from torch.fx import Graph, Interpreter, symbolic_trace | 16 | from torch.fx import Graph, Interpreter, symbolic_trace |
| 16 | from torch.fx.experimental.symbolic_shapes import ( | 17 | from torch.fx.experimental.symbolic_shapes import ( |
| @@ -293,6 +294,89 @@ class TestShapeEnvNPU(TestCase): | |||
| 293 | a, b = sympy.symbols("a b") | 294 | a, b = sympy.symbols("a b") |
| 294 | self.assertEqual(env.simplify((a + b) - b), a) | 295 | self.assertEqual(env.simplify((a + b) - b), a) |
| 295 | 296 | ||
| 297 | + | ||
| 298 | + def test_torch_fx_experimental_symbolic_shapes_ShapeEnv_format_guards(self): | ||
| 299 | + """Verify format_guards returns formatted guard strings.""" | ||
| 300 | + se = ShapeEnv() | ||
| 301 | + self.assertEqual(se.format_guards(), "") | ||
| 302 | + | ||
| 303 | + s0 = sympy.Symbol('s0', integer=True, positive=True) | ||
| 304 | + s1 = sympy.Symbol('s1', integer=True, positive=True) | ||
| 305 | + sloc = SLoc("framework_loc", "user_loc") | ||
| 306 | + se.guards.append(ShapeGuard( | ||
| 307 | + sympy.Lt(s0, s1, evaluate=False), sloc, False)) | ||
| 308 | + se.guards.append(ShapeGuard( | ||
| 309 | + sympy.Ge(s0, sympy.Integer(1), evaluate=False), sloc, False)) | ||
| 310 | + | ||
| 311 | + result = se.format_guards() | ||
| 312 | + self.assertIn("s0 < s1", result) | ||
| 313 | + self.assertIn("s0 >= 1", result) | ||
| 314 | + | ||
| 315 | + result_verbose = se.format_guards(verbose=True) | ||
| 316 | + self.assertIn("s0 < s1", result_verbose) | ||
| 317 | + self.assertIn("user_loc", result_verbose) | ||
| 318 | + | ||
| 319 | + | ||
| 320 | + def test_torch_fx_experimental_symbolic_shapes_ShapeEnv_freeze(self): | ||
| 321 | + """Verify freeze toggles the frozen state of ShapeEnv.""" | ||
| 322 | + se = ShapeEnv() | ||
| 323 | + self.assertFalse(se.frozen) | ||
| 324 | + se.freeze() | ||
| 325 | + self.assertTrue(se.frozen) | ||
| 326 | + | ||
| 327 | + | ||
| 328 | + def test_torch_fx_experimental_symbolic_shapes_ShapeEnv_freeze_runtime_asserts(self): | ||
| 329 | + """Verify freeze_runtime_asserts toggles runtime_asserts_frozen.""" | ||
| 330 | + se = ShapeEnv() | ||
| 331 | + self.assertFalse(se.runtime_asserts_frozen) | ||
| 332 | + se.freeze_runtime_asserts() | ||
| 333 | + self.assertTrue(se.runtime_asserts_frozen) | ||
| 334 | + | ||
| 335 | + | ||
| 336 | + def test_torch_fx_experimental_symbolic_shapes_ShapeEnv_get_axioms(self): | ||
| 337 | + """Verify get_axioms returns axioms tuple, optionally filtered by symbols.""" | ||
| 338 | + se = ShapeEnv() | ||
| 339 | + self.assertIsInstance(se.get_axioms(), tuple) | ||
| 340 | + | ||
| 341 | + s0 = sympy.Symbol('s0', integer=True, positive=True) | ||
| 342 | + s1 = sympy.Symbol('s1', integer=True, positive=True) | ||
| 343 | + sloc = SLoc("framework_loc", "user_loc") | ||
| 344 | + se.guards.append(ShapeGuard( | ||
| 345 | + sympy.Lt(s0, s1, evaluate=False), sloc, False)) | ||
| 346 | + | ||
| 347 | + axioms = se.get_axioms(symbols=(s0,)) | ||
| 348 | + self.assertIn(sympy.Lt(s0, s1, evaluate=False), axioms) | ||
| 349 | + | ||
| 350 | + | ||
| 351 | + def test_torch_fx_experimental_symbolic_shapes_ShapeEnv_get_implications(self): | ||
| 352 | + """Verify get_implications returns implications for Eq/Lt/Ne/Le expressions.""" | ||
| 353 | + se = ShapeEnv() | ||
| 354 | + s0 = sympy.Symbol('s0', integer=True, positive=True) | ||
| 355 | + s1 = sympy.Symbol('s1', integer=True, positive=True) | ||
| 356 | + | ||
| 357 | + # Eq implies equality | ||
| 358 | + impl_eq = dict(se.get_implications( | ||
| 359 | + sympy.Eq(s0, s1, evaluate=False))) | ||
| 360 | + self.assertIn(sympy.Eq(s0, s1, evaluate=False), impl_eq) | ||
| 361 | + | ||
| 362 | + # Lt implies Le and Ne | ||
| 363 | + impl_lt = dict(se.get_implications( | ||
| 364 | + sympy.Lt(s0, s1, evaluate=False))) | ||
| 365 | + self.assertIn(sympy.Lt(s0, s1, evaluate=False), impl_lt) | ||
| 366 | + self.assertIn(sympy.Le(s0, s1, evaluate=False), impl_lt) | ||
| 367 | + self.assertIn(sympy.Ne(s0, s1, evaluate=False), impl_lt) | ||
| 368 | + | ||
| 369 | + # Ne | ||
| 370 | + impl_ne = dict(se.get_implications( | ||
| 371 | + sympy.Ne(s0, s1, evaluate=False))) | ||
| 372 | + self.assertIn(sympy.Ne(s0, s1, evaluate=False), impl_ne) | ||
| 373 | + | ||
| 374 | + # Le implies Lt(a, b+1) | ||
| 375 | + impl_le = dict(se.get_implications( | ||
| 376 | + sympy.Le(s0, s1, evaluate=False))) | ||
| 377 | + self.assertIn(sympy.Le(s0, s1, evaluate=False), impl_le) | ||
| 378 | + self.assertIn(sympy.Lt(s0, s1 + 1, evaluate=False), impl_le) | ||
| 379 | + | ||
| 296 | 380 | ||
| 297 | if __name__ == "__main__": | 381 | if __name__ == "__main__": |
| 298 | run_tests() | 382 | run_tests() |