已合并
fix(triton_experimental): adapt launcher to vendored triton core 3.5.0 #43760
fix(triton_experimental): adapt launcher to vendored triton core 3.5.0 #43760
已合并
AllenGuan创建于 8月4日
共 3 个文件变更+99-26
@@ -72,6 +72,7 @@ import contextlib
72 72 
73from torch._inductor import config73from torch._inductor import config
74from .. import device_props74from .. import device_props
75+from ..compat import IS_TRITON_36_PLUS
75import torch76import torch
76 77 
77 78 
@@ -3608,7 +3609,6 @@ class NPUTritonKernel(TritonKernel):
3608 2. Add total_size arg for NPU 40CU group dispatch3609 2. Add total_size arg for NPU 40CU group dispatch
3609 3. Wrap kernel body in group-based loop3610 3. Wrap kernel body in group-based loop
3610 """3611 """
3611- from torch._inductor.utils import triton_version_uses_attrs_dict
3612 from torch._inductor.codegen.triton_utils import (3612 from torch._inductor.codegen.triton_utils import (
3613 config_of, signature_to_meta, equal_1_arg_indices, non_constexpr_signature3613 config_of, signature_to_meta, equal_1_arg_indices, non_constexpr_signature
3614 )3614 )
@@ -3823,9 +3823,9 @@ class NPUTritonKernel(TritonKernel):
3823 3823 
3824 # Add BLOCK constexpr args3824 # Add BLOCK constexpr args
3825 def add_constexpr_arg(arg_name):3825 def add_constexpr_arg(arg_name):
3826- if triton_version_uses_attrs_dict():3826+ if IS_TRITON_36_PLUS:
3827 signature.append(ConstexprArg(arg_name))3827 signature.append(ConstexprArg(arg_name))
3828- argdefs.append(ArgName(arg_name, is_constexpr=True))3828+ argdefs.append(ArgName(arg_name, is_constexpr=True)) # pre-3.6: constexprs stay out of signature
3829 3829 
3830 for tree in self.range_trees:3830 for tree in self.range_trees:
3831 if tree.tensor_dim is None:3831 if tree.tensor_dim is None:
@@ -5087,19 +5087,27 @@ class NPUTritonScheduling(TritonScheduling):
5087 size_hints = {"x": int(x_total_hint), "r0_": total_cores}5087 size_hints = {"x": int(x_total_hint), "r0_": total_cores}
5088 from torch._inductor.codegen.triton_utils import _type_of5088 from torch._inductor.codegen.triton_utils import _type_of
5089 dt_star = _type_of(out_dtype) # e.g. "*fp32"5089 dt_star = _type_of(out_dtype) # e.g. "*fp32"
5090- # Build the AttrsDescriptor config directly (config_of() resolves alignment via5090+ # Build the combine-kernel attrs config directly (config_of() resolves alignment
5091- # scheduler.name_to_buf, but our "in_ptr0"/"out_ptr0" aren't real graph buffers).5091+ # via scheduler.name_to_buf, but our "in_ptr0"/"out_ptr0" aren't real graph buffers).
5092 # Mark pointers (0,1) and static xnumel (2) 16-divisible: workspace+output are5092 # Mark pointers (0,1) and static xnumel (2) 16-divisible: workspace+output are
5093 # fresh aligned allocations and xnumel is a multiple of 16. r0_numel (3) is the5093 # fresh aligned allocations and xnumel is a multiple of 16. r0_numel (3) is the
5094 # constexpr core count.5094 # constexpr core count.
5095- from triton.compiler.compiler import AttrsDescriptor
5096 div16 = [0, 1]5095 div16 = [0, 1]
5097 if int(x_total_hint) % 16 == 0:5096 if int(x_total_hint) % 16 == 0:
5098 div16.append(2)5097 div16.append(2)
5099- combine_attrs = AttrsDescriptor.from_dict({5098+ if IS_TRITON_36_PLUS:
5100- "arg_properties": {"tt.divisibility": tuple(div16), "tt.equal_to": ()},5099+ # triton-ascend >= 3.6 (vendored core >= 3.5.0): AttrsDescriptor is removed;
5101- "cls": "AttrsDescriptor",5100+ # configs entries are plain dicts keyed by (arg_idx,) with
5102- })5101+ # [["tt.divisibility", 16]] payloads, consumed by triton's
5102+ # ASTFunction.deserialize. Mirrors torch AttrsDescriptorWrapper's dict branch.
5103+ combine_attrs = {(i,): [["tt.divisibility", 16]] for i in div16}
5104+ else:
5105+ # triton-ascend 3.2.x: V2 AttrsDescriptor object.
5106+ from triton.compiler.compiler import AttrsDescriptor
5107+ combine_attrs = AttrsDescriptor.from_dict({
5108+ "arg_properties": {"tt.divisibility": tuple(div16), "tt.equal_to": ()},
5109+ "cls": "AttrsDescriptor",
5110+ })
5103 triton_meta = {5111 triton_meta = {
5104 "signature": {5112 "signature": {
5105 "in_ptr0": dt_star,5113 "in_ptr0": dt_star,
@@ -0,0 +1,41 @@
1+# Copyright (c) 2026, Huawei Technologies Co., Ltd
2+#
3+"""Single source of truth for the triton-ascend API-generation switch.
4+ 
5+IS_TRITON_36_PLUS is True iff the installed triton-ascend exposes the >= 3.6
6+API surface (vendored triton core >= 3.5.0): launch hooks moved to
7+triton.knobs.runtime.* AND AttrsDescriptor removed everywhere (constants
8+passed as a plain dict). False covers triton-ascend 3.2.x (vendored core
9+3.2.0): hooks are CompiledKernel class attrs, AttrsDescriptor exists.
10+ 
11+Capability probe (no version-string parsing), mirroring torch's own
12+get_triton_attrs_descriptor_version heuristic:
13+ 3.2.2 -> (no knobs, has AttrsDescriptor) -> False
14+ 3.6 -> (has knobs, no AttrsDescriptor) -> True
15+A hypothetical mixed shape (knobs present, AttrsDescriptor still there)
16+resolves to False -- the safe/legacy direction. Equivalent to torch's
17+triton_version_uses_attrs_dict() on every released triton-ascend.
18+"""
19+ 
20+try:
21+ from triton import knobs # noqa: F401
22+ 
23+ _has_knobs = True
24+except ImportError:
25+ _has_knobs = False
26+ 
27+_has_attrs_descriptor = False
28+try:
29+ import triton.backends.compiler as _backends_compiler
30+ 
31+ _has_attrs_descriptor |= hasattr(_backends_compiler, "AttrsDescriptor")
32+except ImportError:
33+ pass
34+try:
35+ import triton.compiler.compiler as _compiler_compiler
36+ 
37+ _has_attrs_descriptor |= hasattr(_compiler_compiler, "AttrsDescriptor")
38+except ImportError:
39+ pass
40+ 
41+IS_TRITON_36_PLUS = _has_knobs and not _has_attrs_descriptor
@@ -45,6 +45,7 @@ from torch._inductor.runtime.hints import (
45)45)
46from .device_props import get_npu_vector_core_count, get_npu_ub_size_bytes46from .device_props import get_npu_vector_core_count, get_npu_ub_size_bytes
47from . import device_props47from . import device_props
48+from .compat import IS_TRITON_36_PLUS
48from torch._inductor.runtime.triton_heuristics import (49from torch._inductor.runtime.triton_heuristics import (
49 CachingAutotuner,50 CachingAutotuner,
50 TritonCompileResult,51 TritonCompileResult,
@@ -57,6 +58,7 @@ from torch._inductor.runtime.triton_compat import (
57 ASTSource,58 ASTSource,
58 Config,59 Config,
59 GPUTarget,60 GPUTarget,
61+ knobs,
60)62)
61# 2.13.0: upstream triton_compat removed cc_warp_size; warp_size now uses the63# 2.13.0: upstream triton_compat removed cc_warp_size; warp_size now uses the
62# DeviceProperties.warp_size field (None on NPU -> falls back to 32), matching the64# DeviceProperties.warp_size field (None on NPU -> falls back to 32), matching the
@@ -230,7 +232,6 @@ class NPUTritonCompileResult(TritonCompileResult):
230 """232 """
231 233 
232 def make_launcher(self):234 def make_launcher(self):
233- from torch._inductor.utils import triton_version_uses_attrs_dict
234 from torch._inductor.runtime.triton_heuristics import (235 from torch._inductor.runtime.triton_heuristics import (
235 config_to_dict,236 config_to_dict,
236 )237 )
@@ -254,21 +255,33 @@ class NPUTritonCompileResult(TritonCompileResult):
254 255 
255 NPU_CU_COUNT = get_npu_vector_core_count()256 NPU_CU_COUNT = get_npu_vector_core_count()
256 257 
257- if triton_version_uses_attrs_dict():258+ if IS_TRITON_36_PLUS:
258 call_args = list(fn.arg_names)259 call_args = list(fn.arg_names)
259 def_args = list(fn.arg_names)260 def_args = list(fn.arg_names)
260- if (261+ # Config constants (the constexprs in fn.constexprs -- XBLOCK/YBLOCK/
261- "num_warps" in compile_meta["constants"]262+ # ZBLOCK/R0_BLOCK -- plus the implicit num_warps/num_stages) are baked
262- or "num_stages" in compile_meta["constants"]263+ # into the compiled kernel from the chosen Config, NOT passed by the
263- ):264+ # launcher caller: the generated wrapper calls kernel.run(ptrs, xnumel,
265+ # stream=...) with no block sizes. Exclude them from def_args and splice
266+ # their literal value into call_args. Mirrors upstream triton_heuristics
267+ # _get_arg_lists implicit_constants; without it the launcher signature
268+ # requires e.g. XBLOCK that the caller never passes ("launcher() missing
269+ # 1 required positional argument: 'XBLOCK'").
270+ implicit_constants = {"num_warps", "num_stages"} | set(known_constants)
271+ implicit_constants &= set(compile_meta["constants"].keys())
272+ if implicit_constants:
273+ # Both def_args (launcher Python signature) and call_args (passed to
274+ # the ascend C runner and to bin.launch_metadata) must drop the config
275+ # constants: the caller never passes them (def_args), and the ascend C
276+ # launch wrapper takes only runtime args, not constexprs -- including
277+ # them raises "function takes exactly N arguments (N+1 given)". Mirrors
278+ # the non-attrs_dict branch's `i not in fn.constexprs` filter.
264 def_args = [279 def_args = [
265- arg for arg in def_args if arg not in ("num_warps", "num_stages")280+ arg for arg in def_args if arg not in implicit_constants
281+ ]
282+ call_args = [
283+ arg for arg in call_args if arg not in implicit_constants
266 ]284 ]
267- repl = {
268- k: str(compile_meta["constants"].get(k))
269- for k in ("num_warps", "num_stages")
270- }
271- call_args = [repl.get(arg, arg) for arg in call_args]
272 else:285 else:
273 call_args = [286 call_args = [
274 arg287 arg
@@ -290,11 +303,22 @@ class NPUTritonCompileResult(TritonCompileResult):
290 if pm is not None and not isinstance(pm, dict) and hasattr(pm, "_asdict"):303 if pm is not None and not isinstance(pm, dict) and hasattr(pm, "_asdict"):
291 binary.packed_metadata = pm._asdict()304 binary.packed_metadata = pm._asdict()
292 305 
306+ # triton-ascend >= 3.6 (vendored core >= 3.5.0) moved launch hooks off the
307+ # CompiledKernel class attribute onto knobs.runtime.* (HookChain). IS_TRITON_36_PLUS's
308+ # probe includes the knobs import, so knobs is not None in that branch.
309+ if IS_TRITON_36_PLUS:
310+ launch_enter = knobs.runtime.launch_enter_hook
311+ launch_exit = knobs.runtime.launch_exit_hook
312+ else:
313+ # legacy (triton-ascend 3.2.x): hooks are CompiledKernel class attributes
314+ launch_enter = binary.__class__.launch_enter_hook
315+ launch_exit = binary.__class__.launch_exit_hook
316+ 
293 scope = {317 scope = {
294 "grid_meta": cfg.kwargs,318 "grid_meta": cfg.kwargs,
295 "bin": binary,319 "bin": binary,
296- "launch_enter_hook": binary.__class__.launch_enter_hook,320+ "launch_enter_hook": launch_enter,
297- "launch_exit_hook": binary.__class__.launch_exit_hook,321+ "launch_exit_hook": launch_exit,
298 "metadata": (322 "metadata": (
299 binary.packed_metadata323 binary.packed_metadata
300 if hasattr(binary, "packed_metadata")324 if hasattr(binary, "packed_metadata")
@@ -398,7 +422,7 @@ class NPUTritonCompileResult(TritonCompileResult):
398 else:422 else:
399 launch_metadata = (423 launch_metadata = (
400 f"bin.launch_metadata((grid_0, 1, 1), stream, {', '.join(call_args)})"424 f"bin.launch_metadata((grid_0, 1, 1), stream, {', '.join(call_args)})"
401- if binary.__class__.launch_enter_hook else "None"425+ if launch_enter else "None"
402 )426 )
403 runner_args = [427 runner_args = [
404 "grid_0",428 "grid_0",
@@ -563,7 +587,7 @@ class NPUTritonCompileResult(TritonCompileResult):
563 if launcher.store_cubin:587 if launcher.store_cubin:
564 launcher.fn = fn588 launcher.fn = fn
565 launcher.bin = binary589 launcher.bin = binary
566- if triton_version_uses_attrs_dict():590+ if IS_TRITON_36_PLUS:
567 cfg_dict = config_to_dict(cfg)591 cfg_dict = config_to_dict(cfg)
568 def_args = [x for x in def_args if x not in cfg_dict]592 def_args = [x for x in def_args if x not in cfg_dict]
569 call_args = [593 call_args = [