已关闭
test(fx): add ShapeEnv API tests without upstream coverage #36447
小辉懂编程创建于 5月22日关闭于 5月22日
test(fx): add ShapeEnv API tests without upstream coverage #36447
已关闭
共 1 个文件变更+25-0
| @@ -1,5 +1,6 @@ | |||
| 1 | import unittest | 1 | import unittest |
| 2 | 2 | ||
| 3 | +import sympy | ||
| 3 | import torch | 4 | import torch |
| 4 | import torch_npu | 5 | import torch_npu |
| 5 | from torch.fx import Graph | 6 | from 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 | ||
| 65 | if __name__ == "__main__": | 90 | if __name__ == "__main__": |
| 66 | run_tests() | 91 | run_tests() |