已合并
feat(_inductor): add ascendc backend support for aclgraph capture/replay #40264
dingdairong创建于 7月7日
feat(_inductor): add ascendc backend support for aclgraph capture/replay #40264
已合并
共 3 个文件变更+125-5
| @@ -0,0 +1,53 @@ | |||
| 1 | +"""AscendC backend: basic compilation and numeric correctness.""" | ||
| 2 | +import unittest | ||
| 3 | +import torch | ||
| 4 | +import torch_npu | ||
| 5 | +from torch.testing._internal.common_utils import ( | ||
| 6 | + instantiate_parametrized_tests, | ||
| 7 | + parametrize, | ||
| 8 | + run_tests, | ||
| 9 | + TestCase, | ||
| 10 | +) | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | +class TestAscendcBasic(TestCase): | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + def setUpClass(cls): | ||
| 18 | + super().setUpClass() | ||
| 19 | + cls._ascendc_ok = False | ||
| 20 | + try: | ||
| 21 | + x = torch.randn(4, 4, device="npu") | ||
| 22 | + torch.compile(lambda t: t + 1, backend="inductor", | ||
| 23 | + options={"npu_backend": "ascendc"})(x) | ||
| 24 | + cls._ascendc_ok = True | ||
| 25 | + except Exception: | ||
| 26 | + pass | ||
| 27 | + | ||
| 28 | + def setUp(self): | ||
| 29 | + super().setUp() | ||
| 30 | + if not self._ascendc_ok: | ||
| 31 | + self.skipTest("ascendc backend not available") | ||
| 32 | + | ||
| 33 | + | ||
| 34 | + def test_mul_sub(self, dtype): | ||
| 35 | + """Pointwise pattern: compiled output matches eager.""" | ||
| 36 | + def fn(x, y): | ||
| 37 | + return (x * y - x) | ||
| 38 | + | ||
| 39 | + x = torch.randn(64, 64, dtype=dtype, device="npu") | ||
| 40 | + y = torch.randn(64, 64, dtype=dtype, device="npu") | ||
| 41 | + | ||
| 42 | + eager_out = fn(x, y) | ||
| 43 | + compiled_fn = torch.compile(fn, backend="inductor", | ||
| 44 | + options={"npu_backend": "ascendc"}) | ||
| 45 | + compiled_out = compiled_fn(x, y) | ||
| 46 | + | ||
| 47 | + torch.testing.assert_close(compiled_out, eager_out, rtol=1e-3, atol=1e-3) | ||
| 48 | + | ||
| 49 | + | ||
| 50 | +instantiate_parametrized_tests(TestAscendcBasic) | ||
| 51 | + | ||
| 52 | +if __name__ == "__main__": | ||
| 53 | + run_tests() | ||
| @@ -0,0 +1,52 @@ | |||
| 1 | +"""AscendC backend: reduce-overhead mode (aclgraph capture/replay).""" | ||
| 2 | +import unittest | ||
| 3 | +import torch | ||
| 4 | +import torch_npu | ||
| 5 | +from torch.testing._internal.common_utils import run_tests, TestCase | ||
| 6 | + | ||
| 7 | + | ||
| 8 | + | ||
| 9 | +class TestAscendcReduceOverhead(TestCase): | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + def setUpClass(cls): | ||
| 13 | + super().setUpClass() | ||
| 14 | + cls._ascendc_ok = False | ||
| 15 | + try: | ||
| 16 | + x = torch.randn(4, 4, device="npu") | ||
| 17 | + torch.compile(lambda t: t + 1, backend="inductor", | ||
| 18 | + options={"npu_backend": "ascendc"})(x) | ||
| 19 | + cls._ascendc_ok = True | ||
| 20 | + except Exception: | ||
| 21 | + pass | ||
| 22 | + | ||
| 23 | + def setUp(self): | ||
| 24 | + super().setUp() | ||
| 25 | + if not self._ascendc_ok: | ||
| 26 | + self.skipTest("ascendc backend not available") | ||
| 27 | + | ||
| 28 | + def test_capture_replay(self): | ||
| 29 | + """Verify reduce-overhead triggers graph capture and results are correct.""" | ||
| 30 | + def fn(x, y): | ||
| 31 | + return x * y - x | ||
| 32 | + | ||
| 33 | + x = torch.randn(64, 64, device="npu") | ||
| 34 | + y = torch.randn(64, 64, device="npu") | ||
| 35 | + | ||
| 36 | + compiled_fn = torch.compile( | ||
| 37 | + fn, backend="inductor", | ||
| 38 | + options={"npu_backend": "ascendc", "triton.cudagraphs": True}, | ||
| 39 | + ) | ||
| 40 | + | ||
| 41 | + # Warmup (triggers compilation + first capture) | ||
| 42 | + with torch.no_grad(): | ||
| 43 | + for _ in range(3): | ||
| 44 | + out = compiled_fn(x, y) | ||
| 45 | + | ||
| 46 | + # Verify correctness | ||
| 47 | + eager_out = fn(x, y) | ||
| 48 | + torch.testing.assert_close(out, eager_out, rtol=1e-3, atol=1e-3) | ||
| 49 | + | ||
| 50 | + | ||
| 51 | +if __name__ == "__main__": | ||
| 52 | + run_tests() | ||
| @@ -18,19 +18,31 @@ from .codegen.common import register_device_op_overrides_npu, patch_cache_base_g | |||
| 18 | from .shape_handling import NPUShapeHandling, patch_shape_handling | 18 | from .shape_handling import NPUShapeHandling, patch_shape_handling |
| 19 | from ._npu_meta_registration import npu_patch_meta | 19 | from ._npu_meta_registration import npu_patch_meta |
| 20 | 20 | ||
| 21 | +# 顶层 patch:所有 inductor backend(triton / mlir / dvm / ascendc)都需要的 NPU 设备级patch, | ||
| 22 | +# 与 codegen 后端选择无关,在任何 backend loader 之前无条件执行 | ||
| 21 | npu_patch_meta() | 23 | npu_patch_meta() |
| 22 | register_device_op_overrides_npu() | 24 | register_device_op_overrides_npu() |
| 23 | -patch_has_triton() | 25 | + |
| 24 | -patch_is_gpu() | 26 | + |
| 25 | -patch_device_supports_tma() | 27 | +def _apply_common_patches(): |
| 26 | -patch_codegen_with_cpp_wrapper() | 28 | + # triton / mlir 后端共用的 patch |
| 27 | -patch_cache_base_get_system() | 29 | + patch_has_triton() |
| 30 | + patch_is_gpu() | ||
| 31 | + patch_device_supports_tma() | ||
| 32 | + patch_codegen_with_cpp_wrapper() | ||
| 33 | + patch_cache_base_get_system() | ||
| 34 | + | ||
| 28 | 35 | ||
| 29 | def _get_backend() -> str: | 36 | def _get_backend() -> str: |
| 30 | return os.getenv("TORCHINDUCTOR_NPU_BACKEND", "default") | 37 | return os.getenv("TORCHINDUCTOR_NPU_BACKEND", "default") |
| 31 | 38 | ||
| 32 | 39 | ||
| 40 | +def _load_ascendc_backend(): | ||
| 41 | + from . import ascendc | ||
| 42 | + | ||
| 43 | + | ||
| 33 | def _load_mlir_backend(): | 44 | def _load_mlir_backend(): |
| 45 | + _apply_common_patches() | ||
| 34 | import torch | 46 | import torch |
| 35 | # Prevent RecursionError when formatting LoweringException for huge output tuples (e.g. many permute nodes). | 47 | # Prevent RecursionError when formatting LoweringException for huge output tuples (e.g. many permute nodes). |
| 36 | from .mfusion.safe_inductor_exc import apply_safe_operator_str_patch_if_enabled | 48 | from .mfusion.safe_inductor_exc import apply_safe_operator_str_patch_if_enabled |
| @@ -51,6 +63,7 @@ def _load_mlir_backend(): | |||
| 51 | 63 | ||
| 52 | 64 | ||
| 53 | def _load_dvm_backend(): | 65 | def _load_dvm_backend(): |
| 66 | + _apply_common_patches() | ||
| 54 | import torch | 67 | import torch |
| 55 | from .ascend_npu_ir.ascend_npu_ir.npu import npu_inductor_plugin | 68 | from .ascend_npu_ir.ascend_npu_ir.npu import npu_inductor_plugin |
| 56 | from .lowering_patch import apply_mlir_inductor_patch | 69 | from .lowering_patch import apply_mlir_inductor_patch |
| @@ -70,6 +83,7 @@ def _load_dvm_backend(): | |||
| 70 | 83 | ||
| 71 | 84 | ||
| 72 | def _load_triton_backend(): | 85 | def _load_triton_backend(): |
| 86 | + _apply_common_patches() | ||
| 73 | import os | 87 | import os |
| 74 | import torch | 88 | import torch |
| 75 | has_triton = torch.utils._triton.has_triton() | 89 | has_triton = torch.utils._triton.has_triton() |
| @@ -318,6 +332,7 @@ def _load_triton_backend(): | |||
| 318 | _BACKEND_LOADERS = { | 332 | _BACKEND_LOADERS = { |
| 319 | "mlir": _load_mlir_backend, | 333 | "mlir": _load_mlir_backend, |
| 320 | "dvm": _load_dvm_backend, | 334 | "dvm": _load_dvm_backend, |
| 335 | + "ascendc": _load_ascendc_backend, | ||
| 321 | "default": _load_triton_backend, | 336 | "default": _load_triton_backend, |
| 322 | } | 337 | } |
| 323 | 338 | ||
🟠 High Priority
变更在
_load_ascendc_backend()函数(第 46 行)中添加了from . import ascendc,但仓库中不存在torch_npu/_inductor/ascendc模块(目录或.py文件均不存在)。当用户设置TORCHINDUCTOR_NPU_BACKEND=ascendc并触发_load_backend()时,该 import 会抛出ModuleNotFoundError,导致整个后端加载失败。触发条件:
TORCHINDUCTOR_NPU_BACKEND=ascendc+ 调用_load_backend()。建议:确保
torch_npu/_inductor/ascendc模块存在并与本 PR 一同合入;或在 import 外包裹 try/except 给出明确错误提示,避免静默崩溃。