已合并
test(fx): add testcases for torch.fx.experimental.symbolic_shapes APIs master #36820
test(fx): add testcases for torch.fx.experimental.symbolic_shapes APIs master #36820
已合并
xuanzhi-2026创建于 5月26日
1 个文件变更+83-0
@@ -13,6 +13,11 @@ Add validation cases for torch.fx symbolic_shapes related APIs on NPU:
13 - symbolic_shapes.DimConstraints.solve13 - symbolic_shapes.DimConstraints.solve
14 - symbolic_shapes.DimConstraints.forced_specializations14 - symbolic_shapes.DimConstraints.forced_specializations
15 - symbolic_shapes.DimConstraints.prettify_results15 - 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_cache21 - symbolic_shapes._lru_cache
17 - symbolic_shapes.CallMethodKey22 - symbolic_shapes.CallMethodKey
18 - symbolic_shapes.CallMethodKey.get23 - symbolic_shapes.CallMethodKey.get
@@ -24,6 +29,8 @@ Add validation cases for torch.fx symbolic_shapes related APIs on NPU:
24 - symbolic_shapes.sym_eq29 - symbolic_shapes.sym_eq
25"""30"""
26 31 
32+import dataclasses
33+import importlib
27import inspect34import inspect
28 35 
29import sympy36import sympy
@@ -37,6 +44,8 @@ from torch.utils._sympy.functions import FloorDiv
37from torch.utils._sympy.value_ranges import ValueRanges44from torch.utils._sympy.value_ranges import ValueRanges
38 45 
39 46 
47+importlib.import_module("torch_npu")
48+ 
40device_type = acc.type if (acc := torch.accelerator.current_accelerator()) else "cpu"49device_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 
223class TestSymbolicShapesTargetApiNPU(TestCase):306class TestSymbolicShapesTargetApiNPU(TestCase):
224 def test_lru_cache_handles_hits_clears_and_maxsize(self):307 def test_lru_cache_handles_hits_clears_and_maxsize(self):