已合并
docs:更新 torch.fx.experimental.symbolic_shapes API 文档 #36879
cuiyunhao-2026创建于 5月27日
docs:更新 torch.fx.experimental.symbolic_shapes API 文档 #36879
已合并
共 6 个文件变更+139-4
| @@ -0,0 +1,13 @@ | |||
| 1 | +# torch.fx.experimental.symbolic_shapes | ||
| 2 | + | ||
| 3 | +> [!NOTE] | ||
| 4 | +> 若API"是否支持"为"是","限制与说明"为"-",说明此API和原生API支持度保持一致。 | ||
| 5 | + | ||
| 6 | +|API名称|是否支持|限制与说明| | ||
| 7 | +|--|--|--| | ||
| 8 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv](https://pytorch.org/docs/2.10.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html)|是|-| | ||
| 9 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.format_guards](https://pytorch.org/docs/2.10.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.format_guards)|是|-| | ||
| 10 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze](https://pytorch.org/docs/2.10.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze)|是|-| | ||
| 11 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze_runtime_asserts](https://pytorch.org/docs/2.10.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze_runtime_asserts)|是|-| | ||
| 12 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.get_axioms](https://pytorch.org/docs/2.10.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.get_axioms)|是|-| | ||
| 13 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.get_implications](https://pytorch.org/docs/2.10.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.get_implications)|是|-| | ||
| @@ -0,0 +1,13 @@ | |||
| 1 | +# torch.fx.experimental.symbolic_shapes | ||
| 2 | + | ||
| 3 | +> [!NOTE] | ||
| 4 | +> 若API"是否支持"为"是","限制与说明"为"-",说明此API和原生API支持度保持一致。 | ||
| 5 | + | ||
| 6 | +|API名称|是否支持|限制与说明| | ||
| 7 | +|--|--|--| | ||
| 8 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv](https://pytorch.org/docs/2.11.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html)|是|-| | ||
| 9 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.format_guards](https://pytorch.org/docs/2.11.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.format_guards)|是|-| | ||
| 10 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze](https://pytorch.org/docs/2.11.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze)|是|-| | ||
| 11 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze_runtime_asserts](https://pytorch.org/docs/2.11.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze_runtime_asserts)|是|-| | ||
| 12 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.get_axioms](https://pytorch.org/docs/2.11.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.get_axioms)|是|-| | ||
| 13 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.get_implications](https://pytorch.org/docs/2.11.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.get_implications)|是|-| | ||
| @@ -0,0 +1,13 @@ | |||
| 1 | +# torch.fx.experimental.symbolic_shapes | ||
| 2 | + | ||
| 3 | +> [!NOTE] | ||
| 4 | +> 若API"是否支持"为"是","限制与说明"为"-",说明此API和原生API支持度保持一致。 | ||
| 5 | + | ||
| 6 | +|API名称|是否支持|限制与说明| | ||
| 7 | +|--|--|--| | ||
| 8 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv](https://pytorch.org/docs/2.12.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html)|是|-| | ||
| 9 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.format_guards](https://pytorch.org/docs/2.12.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.format_guards)|是|-| | ||
| 10 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze](https://pytorch.org/docs/2.12.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze)|是|-| | ||
| 11 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze_runtime_asserts](https://pytorch.org/docs/2.12.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze_runtime_asserts)|是|-| | ||
| 12 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.get_axioms](https://pytorch.org/docs/2.12.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.get_axioms)|是|-| | ||
| 13 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.get_implications](https://pytorch.org/docs/2.12.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.get_implications)|是|-| | ||
| @@ -0,0 +1,13 @@ | |||
| 1 | +# torch.fx.experimental.symbolic_shapes | ||
| 2 | + | ||
| 3 | +> [!NOTE] | ||
| 4 | +> 若API"是否支持"为"是","限制与说明"为"-",说明此API和原生API支持度保持一致。 | ||
| 5 | + | ||
| 6 | +|API名称|是否支持|限制与说明| | ||
| 7 | +|--|--|--| | ||
| 8 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv](https://pytorch.org/docs/2.7.1/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html)|是|-| | ||
| 9 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.format_guards](https://pytorch.org/docs/2.7.1/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.format_guards)|是|-| | ||
| 10 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze](https://pytorch.org/docs/2.7.1/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze)|是|-| | ||
| 11 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze_runtime_asserts](https://pytorch.org/docs/2.7.1/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze_runtime_asserts)|是|-| | ||
| 12 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.get_axioms](https://pytorch.org/docs/2.7.1/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.get_axioms)|是|-| | ||
| 13 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.get_implications](https://pytorch.org/docs/2.7.1/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.get_implications)|是|-| | ||
| @@ -0,0 +1,13 @@ | |||
| 1 | +# torch.fx.experimental.symbolic_shapes | ||
| 2 | + | ||
| 3 | +> [!NOTE] | ||
| 4 | +> 若API"是否支持"为"是","限制与说明"为"-",说明此API和原生API支持度保持一致。 | ||
| 5 | + | ||
| 6 | +|API名称|是否支持|限制与说明| | ||
| 7 | +|--|--|--| | ||
| 8 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv](https://pytorch.org/docs/2.9.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html)|是|-| | ||
| 9 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.format_guards](https://pytorch.org/docs/2.9.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.format_guards)|是|-| | ||
| 10 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze](https://pytorch.org/docs/2.9.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze)|是|-| | ||
| 11 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze_runtime_asserts](https://pytorch.org/docs/2.9.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.freeze_runtime_asserts)|是|-| | ||
| 12 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.get_axioms](https://pytorch.org/docs/2.9.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.get_axioms)|是|-| | ||
| 13 | +|[torch.fx.experimental.symbolic_shapes.ShapeEnv.get_implications](https://pytorch.org/docs/2.9.0/generated/torch.fx.experimental.symbolic_shapes.ShapeEnv.html#torch.fx.experimental.symbolic_shapes.ShapeEnv.get_implications)|是|-| | ||
| @@ -1,8 +1,9 @@ | |||
| 1 | """ | 1 | """ |
| 2 | -Add validation cases for torch.fx.experimental.symbolic_shapes APIs on NPU: | 2 | +Add validation cases for torch.fx.experimental.symbolic_shapes.ShapeEnv APIs on NPU: |
| 3 | - | 3 | +1. PyTorch community lacks sufficient and direct API validations for some APIs, so this file is added. |
| 4 | -PyTorch community lacks sufficient and direct API validations for some APIs, so this file is added. | 4 | +2. This file validates ShapeEnv.format_guards, ShapeEnv.freeze, |
| 5 | -This file validates ShapeEnv guard / sympy APIs and PropagateUnbackedSymInts interpreter APIs (extendable). | 5 | + ShapeEnv.freeze_runtime_asserts, ShapeEnv.get_axioms, |
| 6 | + ShapeEnv.get_implications (extendable). | ||
| 6 | """ | 7 | """ |
| 7 | 8 | ||
| 8 | import unittest | 9 | import unittest |
| @@ -11,6 +12,7 @@ import sympy | |||
| 11 | import torch | 12 | import torch |
| 12 | import torch_npu | 13 | import torch_npu |
| 13 | from torch import nn | 14 | from torch import nn |
| 15 | +from torch._guards import ShapeGuard, SLoc | ||
| 14 | from torch._subclasses.fake_tensor import FakeTensorMode | 16 | from torch._subclasses.fake_tensor import FakeTensorMode |
| 15 | from torch.fx import Graph, Interpreter, symbolic_trace | 17 | from torch.fx import Graph, Interpreter, symbolic_trace |
| 16 | from torch.fx.experimental.symbolic_shapes import ( | 18 | from torch.fx.experimental.symbolic_shapes import ( |
| @@ -294,5 +296,73 @@ class TestShapeEnvNPU(TestCase): | |||
| 294 | self.assertEqual(env.simplify((a + b) - b), a) | 296 | self.assertEqual(env.simplify((a + b) - b), a) |
| 295 | 297 | ||
| 296 | 298 | ||
| 299 | + | ||
| 300 | + def test_torch_fx_experimental_symbolic_shapes_ShapeEnv_format_guards(self): | ||
| 301 | + """Verify format_guards returns empty string when no guards, and formatted output with guards.""" | ||
| 302 | + se = ShapeEnv() | ||
| 303 | + self.assertEqual(se.format_guards(), "") | ||
| 304 | + | ||
| 305 | + s0 = sympy.Symbol('s0', integer=True, positive=True) | ||
| 306 | + s1 = sympy.Symbol('s1', integer=True, positive=True) | ||
| 307 | + sloc = SLoc("framework_loc", "user_loc") | ||
| 308 | + se.guards.append(ShapeGuard(sympy.Lt(s0, s1, evaluate=False), sloc, False)) | ||
| 309 | + se.guards.append(ShapeGuard(sympy.Ge(s0, sympy.Integer(1), evaluate=False), sloc, False)) | ||
| 310 | + result = se.format_guards() | ||
| 311 | + self.assertIn("s0 < s1", result) | ||
| 312 | + self.assertIn("s0 >= 1", result) | ||
| 313 | + | ||
| 314 | + result_verbose = se.format_guards(verbose=True) | ||
| 315 | + self.assertIn("user_loc", result_verbose) | ||
| 316 | + | ||
| 317 | + | ||
| 318 | + def test_torch_fx_experimental_symbolic_shapes_ShapeEnv_freeze(self): | ||
| 319 | + """Verify freeze sets frozen flag to True.""" | ||
| 320 | + se = ShapeEnv() | ||
| 321 | + self.assertFalse(se.frozen) | ||
| 322 | + se.freeze() | ||
| 323 | + self.assertTrue(se.frozen) | ||
| 324 | + | ||
| 325 | + | ||
| 326 | + def test_torch_fx_experimental_symbolic_shapes_ShapeEnv_freeze_runtime_asserts(self): | ||
| 327 | + """Verify freeze_runtime_asserts sets runtime_asserts_frozen flag to True.""" | ||
| 328 | + se = ShapeEnv() | ||
| 329 | + self.assertFalse(se.runtime_asserts_frozen) | ||
| 330 | + se.freeze_runtime_asserts() | ||
| 331 | + self.assertTrue(se.runtime_asserts_frozen) | ||
| 332 | + | ||
| 333 | + | ||
| 334 | + def test_torch_fx_experimental_symbolic_shapes_ShapeEnv_get_axioms(self): | ||
| 335 | + """Verify get_axioms returns tuple, and filters by symbols parameter.""" | ||
| 336 | + se = ShapeEnv() | ||
| 337 | + self.assertIsInstance(se.get_axioms(), tuple) | ||
| 338 | + | ||
| 339 | + s0 = sympy.Symbol('s0', integer=True, positive=True) | ||
| 340 | + s1 = sympy.Symbol('s1', integer=True, positive=True) | ||
| 341 | + sloc = SLoc("framework_loc", "user_loc") | ||
| 342 | + se.guards.append(ShapeGuard(sympy.Lt(s0, s1, evaluate=False), sloc, False)) | ||
| 343 | + se.guards.append(ShapeGuard(sympy.Ge(s0, sympy.Integer(1), evaluate=False), sloc, False)) | ||
| 344 | + axioms_with_symbols = se.get_axioms(symbols=(s0,)) | ||
| 345 | + self.assertTrue(len(axioms_with_symbols) > 0) | ||
| 346 | + | ||
| 347 | + | ||
| 348 | + def test_torch_fx_experimental_symbolic_shapes_ShapeEnv_get_implications(self): | ||
| 349 | + """Verify get_implications returns implications for Eq, Lt, Ne, Le expressions.""" | ||
| 350 | + se = ShapeEnv() | ||
| 351 | + s0 = sympy.Symbol('s0', integer=True, positive=True) | ||
| 352 | + s1 = sympy.Symbol('s1', integer=True, positive=True) | ||
| 353 | + | ||
| 354 | + impl_eq = se.get_implications(sympy.Eq(s0, s1, evaluate=False)) | ||
| 355 | + self.assertTrue(len(impl_eq) > 0) | ||
| 356 | + | ||
| 357 | + impl_lt = se.get_implications(sympy.Lt(s0, s1, evaluate=False)) | ||
| 358 | + self.assertTrue(len(impl_lt) > 0) | ||
| 359 | + | ||
| 360 | + impl_ne = se.get_implications(sympy.Ne(s0, s1, evaluate=False)) | ||
| 361 | + self.assertTrue(len(impl_ne) > 0) | ||
| 362 | + | ||
| 363 | + impl_le = se.get_implications(sympy.Le(s0, s1, evaluate=False)) | ||
| 364 | + self.assertTrue(len(impl_le) > 0) | ||
| 365 | + | ||
| 366 | + | ||
| 297 | if __name__ == "__main__": | 367 | if __name__ == "__main__": |
| 298 | run_tests() | 368 | run_tests() |