已合并
test(fx): add testcases for torch.fx.experimental.symbolic_shapes APIs master #36820
xuanzhi-2026创建于 5月26日
test(fx): add testcases for torch.fx.experimental.symbolic_shapes APIs master #36820
已合并
共 1 个文件变更+83-0
| @@ -13,6 +13,11 @@ Add validation cases for torch.fx symbolic_shapes related APIs on NPU: | |||
| 13 | - symbolic_shapes.DimConstraints.solve | 13 | - symbolic_shapes.DimConstraints.solve |
| 14 | - symbolic_shapes.DimConstraints.forced_specializations | 14 | - symbolic_shapes.DimConstraints.forced_specializations |
| 15 | - symbolic_shapes.DimConstraints.prettify_results | 15 | - symbolic_shapes.DimConstraints.prettify_results |
| 16 | + - torch.fx.experimental.symbolic_shapes.ShapeEnv.size_hint | ||
| 17 | + - torch.fx.experimental.symbolic_shapes.ShapeEnv.suppress_guards | ||
| 18 | + - torch.fx.experimental.symbolic_shapes.ShapeEnvSettings | ||
| 19 | + - torch.fx.experimental.symbolic_shapes.StatefulSymbolicContext | ||
| 20 | + - torch.fx.experimental.symbolic_shapes.StatelessSymbolicContext | ||
| 16 | - symbolic_shapes._lru_cache | 21 | - symbolic_shapes._lru_cache |
| 17 | - symbolic_shapes.CallMethodKey | 22 | - symbolic_shapes.CallMethodKey |
| 18 | - symbolic_shapes.CallMethodKey.get | 23 | - symbolic_shapes.CallMethodKey.get |
| @@ -24,6 +29,8 @@ Add validation cases for torch.fx symbolic_shapes related APIs on NPU: | |||
| 24 | - symbolic_shapes.sym_eq | 29 | - symbolic_shapes.sym_eq |
| 25 | """ | 30 | """ |
| 26 | 31 | ||
| 32 | +import dataclasses | ||
| 33 | +import importlib | ||
| 27 | import inspect | 34 | import inspect |
| 28 | 35 | ||
| 29 | import sympy | 36 | import sympy |
| @@ -37,6 +44,8 @@ from torch.utils._sympy.functions import FloorDiv | |||
| 37 | from torch.utils._sympy.value_ranges import ValueRanges | 44 | from torch.utils._sympy.value_ranges import ValueRanges |
| 38 | 45 | ||
| 39 | 46 | ||
| 47 | +importlib.import_module("torch_npu") | ||
| 48 | + | ||
| 40 | device_type = acc.type if (acc := torch.accelerator.current_accelerator()) else "cpu" | 49 | device_type = acc.type if (acc := torch.accelerator.current_accelerator()) else "cpu" |
| 41 | 50 | ||
| 42 | 51 | ||
| @@ -219,6 +228,80 @@ class TestSymbolicShapesAPI(TestCase): | |||
| 219 | self.assertTrue(symbolic_shapes.sym_eq(1, 1)) | 228 | self.assertTrue(symbolic_shapes.sym_eq(1, 1)) |
| 220 | self.assertFalse(symbolic_shapes.sym_eq(1, 2)) | 229 | self.assertFalse(symbolic_shapes.sym_eq(1, 2)) |
| 221 | 230 | ||
| 231 | + def test_shape_env_size_hint(self): | ||
| 232 | + shape_env = symbolic_shapes.ShapeEnv() | ||
| 233 | + self.assertEqual(shape_env.size_hint(sympy.Integer(8)), 8) | ||
| 234 | + | ||
| 235 | + signature = inspect.signature(shape_env.size_hint) | ||
| 236 | + self.assertIn("expr", signature.parameters) | ||
| 237 | + self.assertIn("allow_none", signature.parameters) | ||
| 238 | + self.assertEqual(signature.parameters["allow_none"].default, False) | ||
| 239 | + | ||
| 240 | + def test_shape_env_suppress_guards(self): | ||
| 241 | + shape_env = symbolic_shapes.ShapeEnv() | ||
| 242 | + with shape_env.suppress_guards(): | ||
| 243 | + self.assertEqual(shape_env.size_hint(sympy.Integer(4)), 4) | ||
| 244 | + | ||
| 245 | + def test_shape_env_settings(self): | ||
| 246 | + field_names = { | ||
| 247 | + field.name for field in dataclasses.fields(symbolic_shapes.ShapeEnvSettings) | ||
| 248 | + } | ||
| 249 | + setting_values = { | ||
| 250 | + "allow_scalar_outputs": True, | ||
| 251 | + "allow_dynamic_output_shape_ops": True, | ||
| 252 | + "assume_static_by_default": False, | ||
| 253 | + "specialize_zero_one": True, | ||
| 254 | + "duck_shape": True, | ||
| 255 | + "prefer_deferred_runtime_asserts_over_guards": False, | ||
| 256 | + "allow_complex_guards_as_runtime_asserts": False, | ||
| 257 | + "trace_asserts": False, | ||
| 258 | + } | ||
| 259 | + | ||
| 260 | + settings = symbolic_shapes.ShapeEnvSettings( | ||
| 261 | + **{ | ||
| 262 | + name: value | ||
| 263 | + for name, value in setting_values.items() | ||
| 264 | + if name in field_names | ||
| 265 | + } | ||
| 266 | + ) | ||
| 267 | + | ||
| 268 | + self.assertIn("allow_scalar_outputs", field_names) | ||
| 269 | + self.assertIn("duck_shape", field_names) | ||
| 270 | + for name, value in setting_values.items(): | ||
| 271 | + if name in field_names: | ||
| 272 | + self.assertEqual(getattr(settings, name), value) | ||
| 273 | + | ||
| 274 | + def test_stateless_symbolic_context(self): | ||
| 275 | + context = symbolic_shapes.StatelessSymbolicContext( | ||
| 276 | + dynamic_sizes=[symbolic_shapes.DimDynamic.DUCK, symbolic_shapes.DimDynamic.DUCK], | ||
| 277 | + ) | ||
| 278 | + | ||
| 279 | + self.assertEqual( | ||
| 280 | + context.dynamic_sizes, | ||
| 281 | + [symbolic_shapes.DimDynamic.DUCK, symbolic_shapes.DimDynamic.DUCK], | ||
| 282 | + ) | ||
| 283 | + self.assertEqual( | ||
| 284 | + context.dynamic_strides, | ||
| 285 | + [symbolic_shapes.DimDynamic.INFER_STRIDE, symbolic_shapes.DimDynamic.INFER_STRIDE], | ||
| 286 | + ) | ||
| 287 | + self.assertEqual(context.constraint_sizes, [None, None]) | ||
| 288 | + self.assertEqual(context.constraint_strides, [None, None]) | ||
| 289 | + | ||
| 290 | + def test_stateful_symbolic_context(self): | ||
| 291 | + tensor_source = ConstantSource("x") | ||
| 292 | + context = symbolic_shapes.StatefulSymbolicContext( | ||
| 293 | + dynamic_sizes=[symbolic_shapes.DimDynamic.DUCK, symbolic_shapes.DimDynamic.DUCK], | ||
| 294 | + tensor_source=tensor_source, | ||
| 295 | + ) | ||
| 296 | + | ||
| 297 | + self.assertEqual(context.tensor_source, tensor_source) | ||
| 298 | + self.assertEqual(context.shape_env_to_source_to_symbol_cache, {}) | ||
| 299 | + self.assertEqual( | ||
| 300 | + context.dynamic_sizes, | ||
| 301 | + [symbolic_shapes.DimDynamic.DUCK, symbolic_shapes.DimDynamic.DUCK], | ||
| 302 | + ) | ||
| 303 | + | ||
| 304 | + | ||
| 222 | 305 | ||
| 223 | class TestSymbolicShapesTargetApiNPU(TestCase): | 306 | class TestSymbolicShapesTargetApiNPU(TestCase): |
| 224 | def test_lru_cache_handles_hits_clears_and_maxsize(self): | 307 | def test_lru_cache_handles_hits_clears_and_maxsize(self): |