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