已合并
fix(triton_experimental): adapt launcher to vendored triton core 3.5.0 #43760
AllenGuan创建于 8月4日
fix(triton_experimental): adapt launcher to vendored triton core 3.5.0 #43760
已合并
共 3 个文件变更+99-26
| @@ -72,6 +72,7 @@ import contextlib | |||
| 72 | 72 | ||
| 73 | from torch._inductor import config | 73 | from torch._inductor import config |
| 74 | from .. import device_props | 74 | from .. import device_props |
| 75 | +from ..compat import IS_TRITON_36_PLUS | ||
| 75 | import torch | 76 | import torch |
| 76 | 77 | ||
| 77 | 78 | ||
| @@ -3608,7 +3609,6 @@ class NPUTritonKernel(TritonKernel): | |||
| 3608 | 2. Add total_size arg for NPU 40CU group dispatch | 3609 | 2. Add total_size arg for NPU 40CU group dispatch |
| 3609 | 3. Wrap kernel body in group-based loop | 3610 | 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_signature | 3613 | 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 args | 3824 | # 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_of | 5088 | 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 via | 5090 | + # 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 are | 5092 | # 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 the | 5093 | # 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 | ) |
| 46 | from .device_props import get_npu_vector_core_count, get_npu_ub_size_bytes | 46 | from .device_props import get_npu_vector_core_count, get_npu_ub_size_bytes |
| 47 | from . import device_props | 47 | from . import device_props |
| 48 | +from .compat import IS_TRITON_36_PLUS | ||
| 48 | from torch._inductor.runtime.triton_heuristics import ( | 49 | from 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 the | 63 | # 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 the | 64 | # 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 | arg | 287 | 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_metadata | 323 | 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 = fn | 588 | launcher.fn = fn |
| 565 | launcher.bin = binary | 589 | 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 = [ |