已合并
test: 补齐torch.fx.experimental.symbolic_shapes NPU 适配验证与统一运行脚本 #36871
cuiyunhao-2026创建于 5月27日
test: 补齐torch.fx.experimental.symbolic_shapes NPU 适配验证与统一运行脚本 #36871
已合并
cuiyunhao-2026创建于 5月27日
共 1 个文件变更+84-0
@@ -11,6 +11,7 @@ import sympy
11import torch11import torch
12import torch_npu12import torch_npu
13from torch import nn13from torch import nn
14+from torch._guards import ShapeGuard, SLoc
14from torch._subclasses.fake_tensor import FakeTensorMode15from torch._subclasses.fake_tensor import FakeTensorMode
15from torch.fx import Graph, Interpreter, symbolic_trace16from torch.fx import Graph, Interpreter, symbolic_trace
16from torch.fx.experimental.symbolic_shapes import (17from 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+ @unittest.skipUnless(torch.npu.is_available(), "requires npu")
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+ @unittest.skipUnless(torch.npu.is_available(), "requires npu")
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+ @unittest.skipUnless(torch.npu.is_available(), "requires npu")
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+ @unittest.skipUnless(torch.npu.is_available(), "requires npu")
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+ @unittest.skipUnless(torch.npu.is_available(), "requires npu")
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 
297if __name__ == "__main__":381if __name__ == "__main__":
298 run_tests()382 run_tests()