已合并
test(fx): add testcases for torch.fx.experimental.symbolic_shapes APIs v2.12.0 #36818
xuanzhi-2026创建于 5月26日
test(fx): add testcases for torch.fx.experimental.symbolic_shapes APIs v2.12.0 #36818
已合并
共 1 个文件变更+83-0
| @@ -9,6 +9,11 @@ Add validation cases for torch.fx symbolic_shapes related APIs on NPU: | |||
| 9 | 3. Current covered APIs / behaviors include: | 9 | 3. Current covered APIs / behaviors include: |
| 10 | - symbolic_shapes.DimConstraints.add | 10 | - symbolic_shapes.DimConstraints.add |
| 11 | - symbolic_shapes.DimConstraints.add_equality | 11 | - symbolic_shapes.DimConstraints.add_equality |
| 12 | + - torch.fx.experimental.symbolic_shapes.ShapeEnv.size_hint | ||
| 13 | + - torch.fx.experimental.symbolic_shapes.ShapeEnv.suppress_guards | ||
| 14 | + - torch.fx.experimental.symbolic_shapes.ShapeEnvSettings | ||
| 15 | + - torch.fx.experimental.symbolic_shapes.StatefulSymbolicContext | ||
| 16 | + - torch.fx.experimental.symbolic_shapes.StatelessSymbolicContext | ||
| 12 | - symbolic_shapes._lru_cache | 17 | - symbolic_shapes._lru_cache |
| 13 | - symbolic_shapes.CallMethodKey | 18 | - symbolic_shapes.CallMethodKey |
| 14 | - symbolic_shapes.CallMethodKey.get | 19 | - symbolic_shapes.CallMethodKey.get |
| @@ -16,6 +21,9 @@ Add validation cases for torch.fx symbolic_shapes related APIs on NPU: | |||
| 16 | - symbolic_shapes.check_consistent | 21 | - symbolic_shapes.check_consistent |
| 17 | """ | 22 | """ |
| 18 | 23 | ||
| 24 | +import dataclasses | ||
| 25 | +import inspect | ||
| 26 | + | ||
| 19 | import sympy | 27 | import sympy |
| 20 | import torch | 28 | import torch |
| 21 | 29 | ||
| @@ -85,6 +93,81 @@ class TestSymbolicShapesAPI(TestCase): | |||
| 85 | self.assertEqual(constraints._symbolic_equivalences, [(source, symbolic_expr)]) | 93 | self.assertEqual(constraints._symbolic_equivalences, [(source, symbolic_expr)]) |
| 86 | 94 | ||
| 87 | 95 | ||
| 96 | + def test_shape_env_size_hint(self): | ||
| 97 | + shape_env = symbolic_shapes.ShapeEnv() | ||
| 98 | + self.assertEqual(shape_env.size_hint(sympy.Integer(8)), 8) | ||
| 99 | + | ||
| 100 | + signature = inspect.signature(shape_env.size_hint) | ||
| 101 | + self.assertIn("expr", signature.parameters) | ||
| 102 | + self.assertIn("allow_none", signature.parameters) | ||
| 103 | + self.assertEqual(signature.parameters["allow_none"].default, False) | ||
| 104 | + | ||
| 105 | + def test_shape_env_suppress_guards(self): | ||
| 106 | + shape_env = symbolic_shapes.ShapeEnv() | ||
| 107 | + with shape_env.suppress_guards(): | ||
| 108 | + self.assertEqual(shape_env.size_hint(sympy.Integer(4)), 4) | ||
| 109 | + | ||
| 110 | + def test_shape_env_settings(self): | ||
| 111 | + field_names = { | ||
| 112 | + field.name for field in dataclasses.fields(symbolic_shapes.ShapeEnvSettings) | ||
| 113 | + } | ||
| 114 | + setting_values = { | ||
| 115 | + "allow_scalar_outputs": True, | ||
| 116 | + "allow_dynamic_output_shape_ops": True, | ||
| 117 | + "assume_static_by_default": False, | ||
| 118 | + "specialize_zero_one": True, | ||
| 119 | + "duck_shape": True, | ||
| 120 | + "prefer_deferred_runtime_asserts_over_guards": False, | ||
| 121 | + "allow_complex_guards_as_runtime_asserts": False, | ||
| 122 | + "trace_asserts": False, | ||
| 123 | + } | ||
| 124 | + | ||
| 125 | + settings = symbolic_shapes.ShapeEnvSettings( | ||
| 126 | + **{ | ||
| 127 | + name: value | ||
| 128 | + for name, value in setting_values.items() | ||
| 129 | + if name in field_names | ||
| 130 | + } | ||
| 131 | + ) | ||
| 132 | + | ||
| 133 | + self.assertIn("allow_scalar_outputs", field_names) | ||
| 134 | + self.assertIn("duck_shape", field_names) | ||
| 135 | + for name, value in setting_values.items(): | ||
| 136 | + if name in field_names: | ||
| 137 | + self.assertEqual(getattr(settings, name), value) | ||
| 138 | + | ||
| 139 | + def test_stateless_symbolic_context(self): | ||
| 140 | + context = symbolic_shapes.StatelessSymbolicContext( | ||
| 141 | + dynamic_sizes=[symbolic_shapes.DimDynamic.DUCK, symbolic_shapes.DimDynamic.DUCK], | ||
| 142 | + ) | ||
| 143 | + | ||
| 144 | + self.assertEqual( | ||
| 145 | + context.dynamic_sizes, | ||
| 146 | + [symbolic_shapes.DimDynamic.DUCK, symbolic_shapes.DimDynamic.DUCK], | ||
| 147 | + ) | ||
| 148 | + self.assertEqual( | ||
| 149 | + context.dynamic_strides, | ||
| 150 | + [symbolic_shapes.DimDynamic.INFER_STRIDE, symbolic_shapes.DimDynamic.INFER_STRIDE], | ||
| 151 | + ) | ||
| 152 | + self.assertEqual(context.constraint_sizes, [None, None]) | ||
| 153 | + self.assertEqual(context.constraint_strides, [None, None]) | ||
| 154 | + | ||
| 155 | + def test_stateful_symbolic_context(self): | ||
| 156 | + tensor_source = ConstantSource("x") | ||
| 157 | + context = symbolic_shapes.StatefulSymbolicContext( | ||
| 158 | + dynamic_sizes=[symbolic_shapes.DimDynamic.DUCK, symbolic_shapes.DimDynamic.DUCK], | ||
| 159 | + tensor_source=tensor_source, | ||
| 160 | + ) | ||
| 161 | + | ||
| 162 | + self.assertEqual(context.tensor_source, tensor_source) | ||
| 163 | + self.assertEqual(context.shape_env_to_source_to_symbol_cache, {}) | ||
| 164 | + self.assertEqual( | ||
| 165 | + context.dynamic_sizes, | ||
| 166 | + [symbolic_shapes.DimDynamic.DUCK, symbolic_shapes.DimDynamic.DUCK], | ||
| 167 | + ) | ||
| 168 | + | ||
| 169 | + | ||
| 170 | + | ||
| 88 | class TestSymbolicShapesTargetApiNPU(TestCase): | 171 | class TestSymbolicShapesTargetApiNPU(TestCase): |
| 89 | def test_lru_cache_handles_hits_clears_and_maxsize(self): | 172 | def test_lru_cache_handles_hits_clears_and_maxsize(self): |
| 90 | class DummyShapeEnv: | 173 | class DummyShapeEnv: |