已合并
feat(_inductor): add ascendc backend support for aclgraph capture/replay #40264
feat(_inductor): add ascendc backend support for aclgraph capture/replay #40264
已合并
dingdairong创建于 7月7日
共 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+@unittest.skipIf(not torch.npu.is_available(), "requires an NPU device")
14+class TestAscendcBasic(TestCase):
15+ 
16+ @classmethod
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+ @parametrize("dtype", [torch.float32, torch.float16])
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+@unittest.skipIf(not torch.npu.is_available(), "requires an NPU device")
9+class TestAscendcReduceOverhead(TestCase):
10+ 
11+ @classmethod
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
18from .shape_handling import NPUShapeHandling, patch_shape_handling18from .shape_handling import NPUShapeHandling, patch_shape_handling
19from ._npu_meta_registration import npu_patch_meta19from ._npu_meta_registration import npu_patch_meta
20 20 
21+# 顶层 patch:所有 inductor backend(triton / mlir / dvm / ascendc)都需要的 NPU 设备级patch,
22+# 与 codegen 后端选择无关,在任何 backend loader 之前无条件执行
21npu_patch_meta()23npu_patch_meta()
22register_device_op_overrides_npu()24register_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 
29def _get_backend() -> str:36def _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
atomgit-bot
atomgit-botatomgit-bot7月7日

🟠 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 给出明确错误提示,避免静默崩溃。

likedislike
不准确?
dingdairong
7月7日 评论:
42+ 
43+ 
33def _load_mlir_backend():44def _load_mlir_backend():
45+ _apply_common_patches()
34 import torch46 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_enabled48 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 
53def _load_dvm_backend():65def _load_dvm_backend():
66+ _apply_common_patches()
54 import torch67 import torch
55 from .ascend_npu_ir.ascend_npu_ir.npu import npu_inductor_plugin68 from .ascend_npu_ir.ascend_npu_ir.npu import npu_inductor_plugin
56 from .lowering_patch import apply_mlir_inductor_patch69 from .lowering_patch import apply_mlir_inductor_patch
@@ -70,6 +83,7 @@ def _load_dvm_backend():
70 83 
71 84 
72def _load_triton_backend():85def _load_triton_backend():
86+ _apply_common_patches()
73 import os87 import os
74 import torch88 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