已关闭
test(fx): add ShapeEnv API tests without upstream coverage #36447
小辉懂编程创建于 5月22日关闭于 5月22日
test(fx): add ShapeEnv API tests without upstream coverage #36447
已关闭
小辉懂编程创建于 5月22日关闭于 5月22日
共 1 个文件变更+25-0
@@ -1,5 +1,6 @@
1import unittest1import unittest
2 2 
3+import sympy
3import torch4import torch
4import torch_npu5import torch_npu
5from torch.fx import Graph6from torch.fx import Graph
@@ -61,6 +62,30 @@ class TestSymbolicShapes(TestCase):
61 self.assertFalse(is_symbolic(True))62 self.assertFalse(is_symbolic(True))
62 self.assertTrue(is_symbolic(sym_bool))63 self.assertTrue(is_symbolic(sym_bool))
63 64 
65+ def test_shape_env_get_pruned_guards(self):
66+ shape_env = ShapeEnv()
67+ sym_int = shape_env.create_unbacked_symint()
68+ 
69+ pruned_guards = shape_env.get_pruned_guards([sym_int])
70+ 
71+ self.assertIsInstance(pruned_guards, list)
72+ 
73+ def test_shape_env_is_unbacked_symint(self):
74+ shape_env = ShapeEnv()
75+ unbacked_symint = shape_env.create_unbacked_symint()
76+ unbacked_symbol = unbacked_symint.node.expr
77+ regular_symbol = sympy.Symbol("s0", integer=True)
78+ 
79+ self.assertTrue(shape_env.is_unbacked_symint(unbacked_symbol))
80+ self.assertFalse(shape_env.is_unbacked_symint(regular_symbol))
81+ 
82+ def test_shape_env_ignore_fresh_unbacked_symbols(self):
83+ shape_env = ShapeEnv()
84+
85+ with shape_env.ignore_fresh_unbacked_symbols():
86+ sym_int = shape_env.create_unbacked_symint()
87+ self.assertIsNotNone(sym_int)
88+ 
64 89 
65if __name__ == "__main__":90if __name__ == "__main__":
66 run_tests()91 run_tests()