Rate this Page

Source code for helion.runtime.kernel

from __future__ import annotations

import ast
from collections.abc import Sequence
import contextlib
import dataclasses
import functools
import inspect
import itertools
import logging
import operator
import os
import re
import sys
import textwrap
import threading
import types
from typing import TYPE_CHECKING
from typing import Any
from typing import Callable
from typing import Generic
from typing import Hashable
from typing import Literal
from typing import TypeVar
from typing import cast
from typing import overload
from typing_extensions import Protocol
import weakref

import torch
from torch._dynamo.source import GetItemSource
from torch._dynamo.source import LocalSource
from torch._dynamo.source import TensorProperty
from torch._dynamo.source import TensorPropertySource
from torch._inductor.codecache import PyCodeCache
from torch._inductor.codecache import compiled_fx_graph_hash
from torch._subclasses import FakeTensor
from torch._subclasses.fake_tensor import unset_fake_temporarily
import torch.distributed as dist
from torch.utils._pytree import tree_map_only
from torch.utils.weak import WeakIdKeyDictionary

from .. import exc
from .._compat import target_device_capability
from .._compile_time import measure
from .._compiler.ast_extension import unparse
from .._compiler.autotuner_heuristics import compiler_seed_configs
from .._compiler.compile_environment import CompileEnvironment
from .._compiler.compile_environment import TensorDescriptorLayoutGuard
from .._compiler.compile_environment import _symint_free_symbols
from .._compiler.compile_environment import (
    tensor_descriptor_layout_signature_from_strides,
)
from .._compiler.generate_ast import generate_ast
from .._compiler.inductor_lowering_extra import patch_inductor_lowerings
from .._compiler.kernel_compiler import KernelCompiler
from .._compiler.output_header import assert_no_conflicts
from .._compiler.variable_origin import ArgumentOrigin
from .._dist_utils import _find_process_group_name
from .._dist_utils import check_config_consistancy as dist_check_config_consistancy
from .._dist_utils import kernel_declares_process_group
from .._dist_utils import kernel_uses_symm_mem
from .._logging import LazyString
from .._utils import counters
from ..autotuner.base_search import _AutotunableKernel
from ..language.constexpr import ConstExpr
from .config import Config
from .ref_mode import RefModeContext
from .ref_mode import is_ref_mode_enabled
from .settings import Settings

if TYPE_CHECKING:
    from collections.abc import Generator
    from collections.abc import Hashable

    from torch._guards import Source

    from .._compiler.host_function import HostFunction
    from ..autotuner import ConfigSpec
    from ..autotuner.base_cache import BoundKernelInMemoryCacheKey

    ConfigLike = Config | dict[str, object]

log: logging.Logger = logging.getLogger(__name__)


def _indexing_config_uses_tensor_descriptor(indexing: object, index: int) -> bool:
    if indexing == "tensor_descriptor":
        return True
    if isinstance(indexing, list):
        return index < len(indexing) and indexing[index] == "tensor_descriptor"
    return False


def _td_layout_guard_active_for_config(
    guard: TensorDescriptorLayoutGuard, config: Config
) -> bool:
    return any(
        _indexing_config_uses_tensor_descriptor(config.indexing, index)
        for index in guard.memory_op_indices
    ) or any(
        _indexing_config_uses_tensor_descriptor(config.atomic_indexing, index)
        for index in guard.atomic_op_indices
    )


_R = TypeVar("_R")
CompiledConfig = Callable[..., _R]

# Opt-in: auto-capture Pallas kernels under torch.compile (see
# pallas._tpu_compile_capture).
# Off by default so the eager dispatch path is unchanged.
_TPU_COMPILE_CAPTURE = os.environ.get("HELION_TPU_COMPILE_CAPTURE", "0") == "1"

# Cache for GraphModule hashes
_graph_module_hash_cache: WeakIdKeyDictionary = WeakIdKeyDictionary()

_INT32_INDEX_LIMIT = torch.iinfo(torch.int32).max


def _current_device_index(device_type: str) -> int:
    device_module = getattr(torch, device_type, None)
    current_device = getattr(device_module, "current_device", None)
    if callable(current_device):
        return cast("int", current_device())

    accelerator = getattr(torch, "accelerator", None)
    current_accelerator = getattr(accelerator, "current_accelerator", None)
    current_device_index = getattr(accelerator, "current_device_index", None)
    if callable(current_accelerator) and callable(current_device_index):
        accelerator_device = cast("torch.device", current_accelerator())
        if accelerator_device.type == device_type:
            return cast("int", current_device_index())
    raise exc.InvalidAPIUsage(
        f"autotune_multi requires a current indexed accelerator, got {device_type!r}"
    )


def _canonicalize_multi_shape_device(device: torch.device) -> torch.device:
    if device.type in ("cpu", "meta", "mps"):
        raise exc.InvalidAPIUsage(
            f"autotune_multi requires an indexed accelerator device, got {device}"
        )
    if device.index is not None:
        return device
    return torch.device(device.type, _current_device_index(device.type))


def _has_unspecialized_numeric_value(value: object) -> bool:
    """Return whether normal specialization records only a numeric value's type."""
    if isinstance(value, ConstExpr):
        return False
    if type(value) in (bool, int, float):
        return True
    if dataclasses.is_dataclass(value) and not isinstance(value, type):
        return any(
            _has_unspecialized_numeric_value(getattr(value, field.name))
            for field in dataclasses.fields(value)
        )
    if isinstance(value, dict):
        return any(_has_unspecialized_numeric_value(item) for item in value.values())
    if isinstance(value, (list, tuple)):
        return any(_has_unspecialized_numeric_value(item) for item in value)
    return False


def _resolve_index_dtype(
    settings: Settings,
    args: Sequence[object] | tuple[object, ...],
) -> torch.dtype:
    if (index_dtype := settings.index_dtype) is not None:
        limit = torch.iinfo(index_dtype).max
    else:
        limit = _INT32_INDEX_LIMIT
    over_limit = False

    def _check(tensor: torch.Tensor) -> None:
        nonlocal over_limit
        if over_limit:
            return
        try:
            over_limit = bool(tensor.numel() > limit)
        except RuntimeError:  # unbacked SymInt
            if index_dtype is None:
                over_limit = True

    tree_map_only(torch.Tensor, _check, args)
    # pyrefly: ignore [unbound-name]
    if index_dtype is None:  # Auto-select when not provided
        return torch.int64 if over_limit else torch.int32
    if over_limit:
        # pyrefly: ignore [unbound-name]
        raise exc.InputTensorNumelExceedsIndexType(index_dtype=index_dtype)
    # pyrefly: ignore [unbound-name]
    return index_dtype


def _device_specialization_key(
    args: Sequence[object],
) -> tuple[str | None, tuple[int, int] | None]:
    """Return the recursive argument device key used by bound-kernel caching.

    `_find_device` intentionally searches tensors and bare `torch.device`
    objects inside supported containers to match binding device selection.
    Capability splits mixed-sm cache entries. Device index remains excluded so
    same-capability devices share bound kernels, matching prior behavior.
    """
    try:
        device = _find_device(tuple(args))
    except exc.NoTensorArgs:
        return None, None
    return device.type, target_device_capability(device)


@dataclasses.dataclass
class OutputCodeOptions:
    """Options for :meth:`BoundKernel.to_code`.

    Passing ``options=None`` (the default) keeps ``to_code``'s original behavior.

    Attributes:
        allow_helion_deps: When ``False``, emit a self-contained module that does
            not import ``helion`` at runtime -- the dependency-free launcher is
            inlined (and any in-kernel runtime helpers are embedded) so the only
            deps are ``torch`` + the backend DSL.
        jax_fn: Pallas only. When ``True``, emit a module whose entrypoint operates
            on ``jax.Array`` inputs instead of TorchTPU tensors. Orthogonal to
            ``allow_helion_deps``: combine with ``allow_helion_deps=False`` for a
            pure-JAX module (launch core inlined), or leave ``allow_helion_deps=True``
            to import the launch core from helion.
    """

    allow_helion_deps: bool = True
    jax_fn: bool = False


[docs] class Kernel(Generic[_R]):
[docs] def __init__( self, fn: Callable[..., _R], *, configs: Sequence[ConfigLike] | None = None, settings: Settings | None, key: Callable[..., Hashable] | None = None, ) -> None: """ Initialize the Kernel object. This is typically called from the `@helion.kernel` decorator. Args: fn: The function to be compiled as a Helion kernel. configs: A list of configurations to use for the kernel. settings: The settings to be used by the Kernel. If None, a new `Settings()` instance is created. key: Optional callable that returns an extra hashable component for specialization. """ super().__init__() assert isinstance(fn, types.FunctionType) assert_no_conflicts(fn) self.name: str = fn.__name__ # pyrefly: ignore [read-only] self.fn: types.FunctionType = fn self.signature: inspect.Signature = inspect.signature(fn) self.settings: Settings = settings or Settings() self._key_fn: Callable[..., Hashable] | None = key # Whether the kernel declares distributed intent via an hl.ProcessGroupName # argument. Computed once so the per-call is_distributed check stays cheap # on the dispatch hot path (avoids re-running inspect.signature). See #3024. self._declares_process_group: bool = kernel_declares_process_group(fn) self.configs: list[Config] = [ # pyrefly: ignore [bad-argument-type] Config(**config) if isinstance(config, dict) else config for config in configs or [] ] self._bind_lock = threading.RLock() self._bound_kernels: dict[BoundKernelInMemoryCacheKey, BoundKernel] = {} # Fast dispatch cache: maps a cheap, fine-grained argument key (exact # dtype/shape/stride/device per tensor) directly to a BoundKernel, # skipping the full specialization-key machinery on repeat calls. self._dispatch_cache: dict[Hashable, BoundKernel] = {} self._specialize_extra: dict[ Hashable, list[Callable[[Sequence[object]], Hashable]] ] = {} self._has_specialization_extras = False self._cute_grouped_static_tail_extra_descriptors: dict[ Hashable, set[Hashable] ] = {} if any( param.kind in ( inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD, inspect.Parameter.KEYWORD_ONLY, ) for param in self.signature.parameters.values() ): raise TypeError( f"Kernel({self.name}) cannot have *args, **kwargs, or keyword-only arguments" ) self._annotations: list[object] = [] for param in self.signature.parameters.values(): ann = param.annotation if isinstance(ann, str) and re.search(r"constexpr", ann, re.IGNORECASE): self._annotations.append(ConstExpr) else: self._annotations.append(ann) # Cache the number of parameters to avoid accessing self.signature.parameters # during torch.compile tracing. self._num_params: int = len(self.signature.parameters) # Expose function attributes for compatibility with torch.library.custom_op # These are set as instance attributes to allow the Kernel to be used # as if it were a regular function for introspection purposes functools.update_wrapper(self, fn) # Manually add function-specific attributes not copied by update_wrapper self.__globals__ = fn.__globals__ self.__code__ = fn.__code__ self.__defaults__ = fn.__defaults__ self.__kwdefaults__ = fn.__kwdefaults__ # Opt-in: register the torch.compile(backend="tpu") capture op now, so an # annotated/functional/benchmark-free Pallas kernel is captured with zero # warm-up (like a hand-written custom_op). None if it must use first-call # registration (unannotated, mutating, or autotuning among configs=[...]). self._capture_op: Callable[..., Any] | None = None if self.settings.backend == "pallas" and _TPU_COMPILE_CAPTURE: from .pallas._tpu_compile_capture import register_decoration_op self._capture_op = register_decoration_op(self)
[docs] @functools.cache # noqa: B019 def kernel_source(self) -> str: """ Return the kernel's source text. This is the stable identifier across processes/runs, suitable for grouping telemetry rows by kernel during analysis. """ return inspect.getsource(self.fn)
def _get_bound_kernel_cache_key( self, args: tuple[object, ...], signature: tuple[Hashable, ...] ) -> BoundKernelInMemoryCacheKey | None: from ..autotuner.base_cache import BoundKernelInMemoryCacheKey extra_fns = self._specialize_extra.get(signature) if extra_fns is not None: extra_results: tuple[Hashable, ...] = tuple([s(args) for s in extra_fns]) return BoundKernelInMemoryCacheKey(signature, extra_results) return None def _create_bound_kernel_cache_key( self, bound_kernel: BoundKernel, args: tuple[object, ...], signature: tuple[Hashable, ...], ) -> BoundKernelInMemoryCacheKey: from ..autotuner.base_cache import BoundKernelInMemoryCacheKey self._specialize_extra[signature] = extra_fns = bound_kernel._specialize_extra() if extra_fns: self._has_specialization_extras = True extra_results: tuple[Hashable, ...] = tuple([s(args) for s in extra_fns]) return BoundKernelInMemoryCacheKey(signature, extra_results) def _extend_bound_kernel_specializations( self, bound_kernel: BoundKernel, signature: tuple[Hashable, ...], extractors: list[Callable[[Sequence[object]], Hashable]], args: Sequence[object], ) -> None: if not extractors: return from ..autotuner.base_cache import BoundKernelInMemoryCacheKey with self._bind_lock: existing_extractors = self._specialize_extra.setdefault(signature, []) existing_extractors.extend(extractors) self._has_specialization_extras = True with unset_fake_temporarily(): current_results = tuple( extractor(args) for extractor in existing_extractors ) stale_keys = [ key for key in self._bound_kernels if key.specialization_key == signature and len(key.extra_results) < len(existing_extractors) ] for stale_key in stale_keys: self._bound_kernels.pop(stale_key) for fast_key, cached_bound in list(self._dispatch_cache.items()): if cached_bound._base_spec_key == signature: self._dispatch_cache.pop(fast_key) self._bound_kernels[ BoundKernelInMemoryCacheKey(signature, current_results) ] = bound_kernel def _compute_is_distributed(self, args: Sequence[object]) -> bool: """Whether this call should compile as a distributed kernel. True when the arguments carry symmetric-memory tensors, or the author declared distributed intent (``settings.distributed`` or an ``hl.ProcessGroupName`` argument) inside an initialized process group. This is folded into the specialization key so a symmetric-memory call and an identically-shaped ordinary call never share a compiled kernel. See GitHub issue #3024. """ return kernel_uses_symm_mem(tuple(args)) or ( dist.is_initialized() and (self.settings.distributed or self._declares_process_group) ) def _fast_dispatch_key(self, args: tuple[object, ...]) -> Hashable | None: """ Build a cheap dispatch key for the fast-path cache in ``__call__``. The key records exact per-argument metadata (dtype/shape/stride/device for tensors, type and value for scalars), which is strictly finer than the full specialization key: any two argument lists that produce different full keys also produce different fast keys. That makes it safe to map a fast key directly to the BoundKernel that a full ``bind()`` resolved for the same arguments. If a base signature has extra specialization extractors, their results are appended to preserve the same no-collision invariant for value-based specializations. Returns None when an argument type is not handled (tensor subclasses, containers, ...), or when there is no tensor argument to pin down the device; callers must then take the regular ``bind()`` path. """ key: list[Hashable] = [] has_tensor = False for a in args: t = type(a) if t is torch.Tensor or t is torch.nn.Parameter: tensor = cast("torch.Tensor", a) has_tensor = True si = getattr(tensor, "_dynamo_static_indices", None) key.append( ( tensor.dtype, tensor.shape, tensor.stride(), tensor.device, None if si is None else frozenset(si), ) ) elif t is int or t is float or t is bool or t is str: key.append((t, a)) elif a is None: key.append(None) elif t is torch.dtype or t is torch.device: key.append(a) else: return None if not has_tensor: return None key.append(self._compute_is_distributed(args)) signature = ( self._base_specialization_key(args) if self._has_specialization_extras else None ) if self._key_fn is not None: key.append(signature[-1] if signature is not None else self._key_fn(*args)) if signature is not None: extra_fns = self._specialize_extra.get(signature) if extra_fns is not None: key.append(tuple(s(args) for s in extra_fns)) return tuple(key)
[docs] def bind(self, args: tuple[object, ...]) -> BoundKernel[_R]: """ Bind the given arguments to the Kernel and return a BoundKernel object. Args: args: The arguments to bind to the Kernel. Returns: BoundKernel: A BoundKernel object with the given arguments bound. """ # Dynamo executes bind while capturing the call but cannot trace an RLock. # Eager binds, including all cache mutations, remain serialized. if torch.compiler.is_compiling(): return self._bind(args) with self._bind_lock: return self._bind(args)
def _bind(self, args: tuple[object, ...]) -> BoundKernel[_R]: with measure("Kernel.bind"): if not isinstance(args, tuple): assert isinstance(args, list), "args must be a tuple or list" args = tuple(args) if len(args) > self._num_params: raise TypeError( f"Too many arguments passed to the kernel, expected: {self._num_params} got: {len(args)}." ) signature = self._base_specialization_key(args) cache_key = self._get_bound_kernel_cache_key(args, signature) bound_kernel = ( None if cache_key is None else self._bound_kernels.get(cache_key, None) ) if bound_kernel is None: normalized_args: tuple[object, ...] = self.normalize_args(*args) if len(normalized_args) != len(args): # we had default args that needed to be applied bound_kernel = self.bind(normalized_args) else: bound_kernel = BoundKernel(self, args) if cache_key is None: cache_key = self._create_bound_kernel_cache_key( bound_kernel, args, signature ) self._bound_kernels[cache_key] = bound_kernel return bound_kernel def _base_specialization_key(self, args: Sequence[object]) -> tuple[Hashable, ...]: """ Generate the base specialization key from input argument metadata only, using the per-type extractor functions defined in `_specialization_extractors`, without any extras discovered during compilation. Used internally for _specialize_extra lookups. """ result: list[Hashable] = [] assert len(args) <= len(self._annotations) for value, annotation in zip(args, self._annotations, strict=False): if isinstance(value, ConstExpr): result.append(value.value) elif annotation is ConstExpr: result.append(value) else: result.append(self._specialization_key(value)) device_type, device_capability = _device_specialization_key(args) is_distributed = self._compute_is_distributed(args) if self._key_fn is not None: return ( *result, device_type, device_capability, is_distributed, self._key_fn(*args), ) return (*result, device_type, device_capability, is_distributed)
[docs] def specialization_key(self, args: Sequence[object]) -> tuple[Hashable, ...]: """ Generate the full specialization key for the given arguments, including any additional specialization constraints discovered during compilation (e.g. from hl.specialize() calls). Before the first compilation, these extras are not yet known and the key may be incomplete. Args: args: The arguments to generate a specialization key for. Returns: Hashable: A hashable key representing the specialization of the arguments. """ base = self._base_specialization_key(args) extra_fns = self._specialize_extra.get(base) if extra_fns is not None: return base + tuple([s(args) for s in extra_fns]) return base
def _specialization_key(self, obj: object) -> Hashable: """ Helper used to generate a specialization key for the given object. This method determines a unique key for the object based on its type and the corresponding extractor function defined in `_specialization_extractors`. Args: obj: The argument to generate a specialization key for. Returns: Hashable: A hashable key representing the specialization of the object. """ extractor = _specialization_extractors.get(type(obj)) if extractor is None: if isinstance(obj, torch.fx.GraphModule): # GraphModule subclasses need special handling extractor = _specialization_extractors[torch.fx.GraphModule] elif isinstance(obj, torch.Tensor): # torch.Tensor subclasses (e.g. the JAX-export adapter) # share the standard tensor specialization key. Use the # SymInt-safe extractor: unlike exact ``torch.Tensor``, # subclasses may carry symbolic sizes/strides. extractor = _specialization_extractors["tensor_subclass"] elif isinstance(obj, tuple) and hasattr(obj, "_fields"): # this is a namedtuple extractor = _specialization_extractors["namedtuple"] elif dataclasses.is_dataclass(obj): extractor = _specialization_extractors["dataclass"] else: raise TypeError(f"unsupported argument type: {type(obj).__name__}") return extractor(self, obj)
[docs] def normalize_args(self, *args: object, **kwargs: object) -> tuple[object, ...]: """ Normalize the given arguments and keyword arguments according to the function signature. Args: args: The positional arguments to normalize. kwargs: The keyword arguments to normalize. Returns: tuple[object, ...]: A tuple of normalized positional arguments. """ bound_args = self.signature.bind(*args, **kwargs) bound_args.apply_defaults() return tuple(bound_args.args)
[docs] def autotune( self, args: Sequence[object], *, force: bool = True, **options: object, ) -> Config: """ Perform autotuning to find the optimal configuration for the kernel. This uses the default setting, you can call helion.autotune.* directly for more customization. If config= or configs= is provided to helion.kernel(), the search will be restricted to the provided configs. Use force=True to ignore the provided configs. Mutates (the bound version of) self so that `__call__` will run the best config found. Args: args: Example arguments used for benchmarking during autotuning. force: If True, force full autotuning even if a config is provided. options: Additional keyword options forwarded to the autotuner. Returns: Config: The best configuration found during autotuning. """ args = self.normalize_args(*args) return self.bind(args).autotune(args, force=force, **options)
[docs] def autotune_multi( self, arg_sets: Sequence[Sequence[object]], *, aggregation: Literal["geomean", "max"] = "geomean", relative_to: Literal["default", "baseline"] | None = None, cache_tag: str | None = None, force: bool = True, **options: object, ) -> Config: """Find one config using an objective measured across several inputs. The first argument set anchors config generation. Each candidate is measured on every set, then reduced with a geometric mean or maximum. ``relative_to`` optionally optimizes per-shape latency relative to each shape's default config or custom baseline. Only the supplied bound specializations are configured. Every argument set must bind normally to the same exact, current accelerator device. Distributed processes are not supported. A non-empty ``cache_tag`` is required for custom callbacks, dynamic-shape tuning, and runtime numeric arguments; callers own tag invalidation in those cases. Args: arg_sets: Non-empty sequence of representative kernel argument sequences. aggregation: Joint objective, either ``"geomean"`` or ``"max"``. relative_to: Optional ``"default"`` or ``"baseline"`` normalization. cache_tag: User-managed cache discriminator for dynamic shapes, runtime numeric arguments, or callbacks. force: If true, ignore pinned configs and cache reads during the search. options: Additional options forwarded to the registered autotuner. Returns: The config selected for all supplied specializations. """ from ..autotuner.benchmark_provider import LocalBenchmarkProvider from ..autotuner.benchmark_provider import _has_valid_multi_shape_measurement from ..autotuner.benchmark_provider import _materialize_multi_shape_config from ..autotuner.benchmark_provider import _MultiShapeAutotuneArgs from .settings import default_autotuner_fn if aggregation not in ("geomean", "max"): raise exc.InvalidAPIUsage( "autotune_multi aggregation must be 'geomean' or 'max'" ) if relative_to not in (None, "default", "baseline"): raise exc.InvalidAPIUsage( "autotune_multi relative_to must be None, 'default', or 'baseline'" ) if cache_tag is not None and (not isinstance(cache_tag, str) or not cache_tag): raise exc.InvalidAPIUsage( "autotune_multi cache_tag must be a non-empty string" ) if ( not isinstance(arg_sets, Sequence) or isinstance(arg_sets, (str, bytes)) or not arg_sets ): raise exc.InvalidAPIUsage( "autotune_multi arg_sets must be a non-empty sequence" ) if dist.is_initialized(): raise exc.InvalidAPIUsage( "autotune_multi does not support an initialized distributed process group" ) if self.settings.backend not in {"triton", "tileir", "cute", "pallas"}: raise exc.InvalidAPIUsage( f"autotune_multi does not support backend {self.settings.backend!r}" ) if self.settings.autotuner_fn is not default_autotuner_fn: raise exc.InvalidAPIUsage( "autotune_multi requires the default registered autotuner" ) if self.settings.autotune_cache == "AOTAutotuneCache": raise exc.InvalidAPIUsage( "autotune_multi does not support AOTAutotuneCache" ) if self.settings.autotune_cache not in { "LocalAutotuneCache", "StrictLocalAutotuneCache", "RemoteAutotuneCache", "StrictRemoteAutotuneCache", }: raise exc.InvalidAPIUsage( "autotune_multi requires a built-in local or remote best-config cache" ) if "benchmark_provider_cls" in options: if options.pop("benchmark_provider_cls") is not LocalBenchmarkProvider: raise exc.InvalidAPIUsage( "autotune_multi does not support a custom benchmark provider" ) if relative_to == "baseline" and self.settings.autotune_baseline_fn is None: raise exc.InvalidAPIUsage( "autotune_multi relative_to='baseline' requires autotune_baseline_fn" ) if self.settings.autotune_benchmark_fn is not None: raise exc.InvalidAPIUsage( "autotune_multi does not support autotune_benchmark_fn" ) custom_callbacks = ( self.settings.autotune_baseline_fn, self.settings.autotune_baseline_accuracy_check_fn, self.settings.autotune_config_filter, ) if cache_tag is None and any(fn is not None for fn in custom_callbacks): raise exc.InvalidAPIUsage( "autotune_multi requires cache_tag when a custom baseline, " "accuracy check, or config filter is configured" ) if cache_tag is None and not self.settings.static_shapes: raise exc.InvalidAPIUsage( "autotune_multi requires cache_tag when static_shapes=False" ) normalized_arg_sets: list[tuple[object, ...]] = [] for case_index, arg_set in enumerate(arg_sets): if not isinstance(arg_set, Sequence) or isinstance(arg_set, (str, bytes)): raise exc.InvalidAPIUsage( f"autotune_multi arg_sets[{case_index}] must be a sequence" ) try: normalized_arg_sets.append(self.normalize_args(*arg_set)) except TypeError as error: raise exc.InvalidAPIUsage( f"autotune_multi arg_sets[{case_index}] does not match the " f"kernel signature: {error}" ) from error if cache_tag is None: for case_index, normalized_args in enumerate(normalized_arg_sets): if any( annotation is not ConstExpr and _has_unspecialized_numeric_value(value) for value, annotation in zip( normalized_args, self._annotations, strict=True ) ): raise exc.InvalidAPIUsage( "autotune_multi requires cache_tag when an argument set " "contains a runtime numeric value; " f"arg_sets[{case_index}] does" ) cases: list[tuple[BoundKernel[_R], tuple[object, ...]]] = [] case_keys: list[BoundKernelInMemoryCacheKey] = [] for case_index, normalized_args in enumerate(normalized_arg_sets): try: bound_kernel = self.bind(normalized_args) except exc.NoTensorArgs as error: raise exc.InvalidAPIUsage( "autotune_multi requires each argument set to have a device " f"discoverable by normal kernel binding; arg_sets[{case_index}] did not" ) from error signature = self._base_specialization_key(normalized_args) case_key = self._get_bound_kernel_cache_key(normalized_args, signature) assert case_key is not None cases.append((bound_kernel, normalized_args)) case_keys.append(case_key) anchor = cases[0][0] canonical_device = _canonicalize_multi_shape_device(anchor.env.device) current_index = _current_device_index(canonical_device.type) if canonical_device.index != current_index: raise exc.InvalidAPIUsage( "autotune_multi requires the indexed accelerator to be current: " f"got {canonical_device}, current index is {current_index}" ) unique_bound_kernels: list[BoundKernel[_R]] = [] seen_bound_kernel_ids: set[int] = set() for bound_kernel, _ in cases: if id(bound_kernel) not in seen_bound_kernel_ids: seen_bound_kernel_ids.add(id(bound_kernel)) unique_bound_kernels.append(bound_kernel) anchor_backend = anchor.env.backend.name anchor_capability = target_device_capability(anchor.env.device) advanced_controls_files = self.settings.autotune_search_acf or None anchor_fingerprint = anchor.config_spec.structural_fingerprint( advanced_controls_files=advanced_controls_files ) for case_index, (bound_kernel, _) in enumerate(cases): if bound_kernel.env.process_group_name is not None: raise exc.InvalidAPIUsage( "autotune_multi does not support a bound kernel with a process group" ) bound_device = _canonicalize_multi_shape_device(bound_kernel.env.device) if bound_device != canonical_device: raise exc.InvalidAPIUsage( "autotune_multi bound a different device for " f"arg_sets[{case_index}]: {bound_device} != {canonical_device}" ) if bound_kernel.env.backend.name != anchor_backend: raise exc.InvalidAPIUsage( "autotune_multi requires every case to use the same backend" ) if target_device_capability(bound_kernel.env.device) != anchor_capability: raise exc.InvalidAPIUsage( "autotune_multi requires every case to have the same device capability" ) fingerprint = bound_kernel.config_spec.structural_fingerprint( advanced_controls_files=advanced_controls_files ) if fingerprint != anchor_fingerprint: raise exc.InvalidAPIUsage( "autotune_multi requires structurally compatible ConfigSpec " f"instances; arg_sets[{case_index}] is incompatible with the anchor" ) multi_args = _MultiShapeAutotuneArgs( cases=tuple(cases), aggregation=aggregation, relative_to=relative_to, cache_tag=cache_tag, workload_key=( "multi_shape:v1", aggregation, relative_to, cache_tag, tuple( (case_key.specialization_key, case_key.extra_results) for case_key in case_keys ), ), reference_latencies=None, ) ephemeral = anchor.env.backend.make_ephemeral_cache() ctx = ephemeral if ephemeral is not None else contextlib.nullcontext() with ctx: config = anchor.env.backend.autotune( anchor, cast("Sequence[object]", multi_args), force=force, **options, ) if multi_args.search_started and not multi_args.found_valid_config: raise exc.NoConfigFound if multi_args.search_started and not _has_valid_multi_shape_measurement( multi_args, anchor.config_spec, config, ): raise exc.NoConfigFound winner = _materialize_multi_shape_config(anchor.config_spec, config) if ephemeral is not None: for bound_kernel in unique_bound_kernels: bound_kernel.env.backend.finalize_ephemeral_cache(bound_kernel, winner) for bound_kernel in unique_bound_kernels: bound_kernel.compile_config(winner) for bound_kernel in unique_bound_kernels: bound_kernel.set_config(winner) return winner
[docs] def __call__(self, *args: object, **kwargs: object) -> _R: """ Call the Kernel with the given arguments and keyword arguments. Args: args: The positional arguments to pass to the Kernel. kwargs: The keyword arguments to pass to the Kernel. Returns: _R: The result of the Kernel function call. """ if kwargs: args = self.normalize_args(*args, **kwargs) elif self._dispatch_cache: # Fast path: repeat call with argument metadata seen before. The # cache is only populated by calls that already took the slow # path below, so hitting it cannot skip autotuning/compilation. # (The cache stays empty under TPU compile capture, so this # cannot bypass auto_capture_call below.) fast_key = self._fast_dispatch_key(args) if fast_key is not None: bound = self._dispatch_cache.get(fast_key) if bound is not None and bound._run is not None: return bound._run(*args) if self.settings.backend == "pallas" and _TPU_COMPILE_CAPTURE: # Local import: _tpu_compile_capture pulls in the dynamo HOP machinery, # not ready when kernel.py first loads during ``import helion``. from .pallas._tpu_compile_capture import RUN_NORMAL from .pallas._tpu_compile_capture import auto_capture_call result = auto_capture_call(self, args) if result is not RUN_NORMAL: return cast("_R", result) return self.bind(args)(*args) bound = self.bind(args) result = bound(*args) if self._has_specialization_extras: # A concurrent first compile can discover a late specialization and # make BoundKernel.__call__ rebind internally. Resolve it again so # the fast-dispatch key never points at the pre-specialization bound. bound = self.bind(args) if bound._run is not None: fast_key = self._fast_dispatch_key(args) if fast_key is not None: self._dispatch_cache[fast_key] = bound return result
[docs] def reset(self) -> None: """ Clears the cache of bound kernels, meaning subsequent calls will recompile and re-autotune. """ self._bound_kernels.clear() self._dispatch_cache.clear()
@property def jax_fn(self) -> Callable[..., Any]: """A pure-JAX callable view of this Helion kernel. Pallas-backend kernels can be called directly with JAX arrays or tracers (i.e. inside ``jax.jit``) by going through this property. The kernel's compile/specialize path runs the first time the callable is invoked; subsequent calls reuse the cached compilation just like ``__call__``. This only supports kernels compiled for the Pallas backend. """ from .pallas.jax_export import make_jax_fn cached = getattr(self, "_jax_fn_callable", None) if cached is not None: return cast("Callable[..., Any]", cached) rv = make_jax_fn(self) # pyrefly: ignore [unsupported-attribute-set] self._jax_fn_callable = rv return rv
class BoundKernel(_AutotunableKernel, Generic[_R]): def __init__( self, kernel: Kernel[_R], args: tuple[object, ...], ) -> None: """ Initialize a BoundKernel object. This constructor sets up the environment, compiles the kernel function, and prepares the arguments for execution. Args: kernel: The Kernel object to bind. args: A tuple of arguments to bind to the kernel. """ super().__init__() self.kernel = kernel # Base specialization key from the REAL args (the same key bind() uses to # distinguish BoundKernels). Reused in compile_config as the PyCodeCache # ``extra`` discriminator for backends that cache shape-specific module state # (Backend.requires_shape_specialized_module). Must be computed on the real # args: fake_args turn scalar int/float/bool into SymInt/SymFloat/SymBool, # which the specialization extractors reject. self._base_spec_key = kernel._base_specialization_key(args) self._run: Callable[..., _R] | None = None self._config: Config | None = None self._compile_cache: dict[Config, CompiledConfig] = {} self._cache_path_map: dict[Config, str | None] = {} # Direct to_code() has no call arguments, so keep this bound kernel's # construction-time tensor values as its stable weak fallback. self._runtime_tensor_refs_by_name = { name: weakref.ref(value) for name, value in zip(self.kernel.signature.parameters, args, strict=False) if isinstance(value, torch.Tensor) } self._first_compile_lock = threading.RLock() self._backward_compiled: ( tuple[Kernel[object], str, BoundKernel[object]] | None ) = None # Distributed detection lives on Kernel so the same value feeds both the # specialization cache key and CompileEnvironment gating; that keeps a # symmetric-memory call from ever aliasing a non-distributed compiled # kernel of the same shape. See GitHub issue #3024. is_distributed = self.kernel._compute_is_distributed(args) self._env = CompileEnvironment( _find_device(args), self.kernel.settings, index_dtype=_resolve_index_dtype(self.kernel.settings, args), is_distributed=is_distributed, ) if is_ref_mode_enabled(self.kernel.settings): self.fake_args = [] # type: ignore[assignment] self.host_function = None # type: ignore[assignment] return with self.env: self._env.process_group_name = _find_process_group_name( kernel.fn, args, is_distributed ) assert len(args) == len(self.kernel.signature.parameters) self.fake_args: list[object] = [] constexpr_args = {} for name, arg, annotation in zip( self.kernel.signature.parameters, args, self.kernel._annotations, strict=False, ): if isinstance(arg, ConstExpr): assert not isinstance(arg.value, torch.Tensor), ( "ConstExpr cannot be a tensor" ) self.fake_args.append(arg.value) constexpr_args[name] = arg.value elif annotation is ConstExpr: assert not isinstance(arg, torch.Tensor), ( "ConstExpr cannot be a tensor" ) self.fake_args.append(arg) constexpr_args[name] = arg else: self.fake_args.append(self.env.to_fake(arg, ArgumentOrigin(name))) self._apply_mark_static(args) with ( _maybe_skip_dtype_check_in_meta_registrations(), patch_inductor_lowerings(), measure("BoundKernel.create_host_function"), ): try: compiler = KernelCompiler(self.env) self.host_function: HostFunction = compiler.compile( self.kernel.fn, self.fake_args, constexpr_args, ) except Exception: config = self.env.config_spec.default_config() self.maybe_log_repro(log.warning, args, config=config) raise self.env.restrict_pid_types_for_persistent(args) self.env.config_spec.configure_epilogue_subtile_autotune(args) self.env.config_spec.compiler_seed_configs = compiler_seed_configs( self.env, self.host_function.device_ir, ) # Post-compile FX-graph scan to detect kernels # whose tcgen05 matmul is followed by an # aux-fused store # (``out[tile] = (acc + residual[tile]).to(...)`` # and variants — see # ``host_function_has_tcgen05_aux_kernel_pattern`` # for the accepted shapes). When detected, the # autotune surface widens to admit # ``tcgen05_strategy=ROLE_LOCAL_WITH_SCHEDULER`` # + ``tcgen05_warp_spec_c_input_warps=1`` so the # productive C-input warp lift is reachable from # the normal autotune path. For pure-matmul # kernels the detector returns False and the # autotune surface keeps the narrow # ``MONOLITHIC + c_input_warps=0`` shape so # autotune cannot sample the strictly-worse # inert C-input warp configuration. # The exact-shape detector is narrower: it gates the # ``tcgen05_aux_load_mode=tma`` seed/search axis. from .._compiler.cute.aux_tensor import ( host_function_has_tcgen05_aux_kernel_pattern, ) from .._compiler.cute.aux_tensor import ( host_function_has_tcgen05_exact_shape_aux_kernel_pattern, ) from .._compiler.cute.aux_tensor import ( host_function_matmul_has_non_tcgen05_operand, ) self.env.config_spec.cute_tcgen05_aux_kernel_detected = ( host_function_has_tcgen05_aux_kernel_pattern(self.host_function) ) self.env.config_spec.cute_tcgen05_exact_shape_aux_kernel_detected = ( host_function_has_tcgen05_exact_shape_aux_kernel_pattern( self.host_function ) ) self.env.config_spec.cute_tcgen05_matmul_has_non_tcgen05_operand = ( host_function_matmul_has_non_tcgen05_operand(self.host_function) ) if not self.env.settings.disable_autotuner_heuristics: for seed_config in self.env.config_spec.autotune_seed_configs(): if ( seed_config not in self.env.config_spec.compiler_seed_configs ): self.env.config_spec.compiler_seed_configs.append( seed_config ) def _apply_mark_static(self, args: tuple[object, ...]) -> None: """ Apply torch._dynamo.mark_static() markings from input tensors. This reads _dynamo_static_indices from each tensor argument and marks the corresponding dimensions as specialized (constant) in the kernel. """ for arg, fake_arg in zip(args, self.fake_args, strict=True): if isinstance(arg, torch.Tensor) and isinstance(fake_arg, torch.Tensor): for dim in getattr(arg, "_dynamo_static_indices", ()): size = fake_arg.size(dim) if isinstance(size, torch.SymInt): self.env.specialized_vars.update(_symint_free_symbols(size)) @property def env(self) -> CompileEnvironment: # pyrefly: ignore[bad-override] return self._env @property def settings(self) -> Settings: """ Retrieve the settings associated with the kernel. Returns: Settings: The settings of the kernel. """ return self.kernel.settings @property def config_spec(self) -> ConfigSpec: """ Retrieve the configuration specification for the kernel. Returns: ConfigSpec: The configuration specification. """ return self.env.config_spec @property def configs(self) -> list[Config]: """Return the kernel's configured configs (alias for `self.kernel.configs`).""" return self.kernel.configs def _normalize_config(self, config: ConfigLike) -> Config: if isinstance(config, Config): return config # pyrefly: ignore [bad-argument-type] return Config(**config) def _normalized_config_copy(self, config: ConfigLike) -> Config: normalized = self._normalize_config(config) normalized = Config(**normalized.config) # pyrefly: ignore[bad-argument-type] self.env.config_spec.normalize(normalized) return normalized def format_kernel_decorator(self, config: Config, settings: Settings) -> str: """Return the @helion.kernel decorator snippet capturing configs and settings that influence Triton code generation.""" parts = [ f"config={config.__repr__()}", f"static_shapes={settings.static_shapes}", ] if settings.index_dtype is not None: parts.append(f"index_dtype={settings.index_dtype}") return f"@helion.kernel({', '.join(parts)})" def to_code( self, config: ConfigLike | None = None, *, options: OutputCodeOptions | None = None, emit_repro_caller: bool = False, output_origin_lines: bool | None = None, ) -> str: """ Generate backend-specific code for the kernel based on the given configuration. Args: config: The configuration to use for code generation. options: Optional :class:`~helion.runtime.precompile.OutputCodeOptions`. With ``allow_helion_deps=False`` the returned module is self-contained (no ``helion`` import at runtime); ``jax_fn=True`` (Pallas only) emits a pure-JAX module operating on ``jax.Array``s. ``None`` keeps the default behavior. emit_repro_caller: Emits a main function to call the kernel with example inputs. Returns: str: The generated code as a string. """ if config is None: config = self._require_implicit_config() with self.env, measure("BoundKernel.to_code"): # Work on a copy so the caller's Config is not mutated with defaults # specific to this BoundKernel's config_spec. config = self._normalized_config_copy(config) with ( self._runtime_arg_values_for_codegen(), measure("BoundKernel.generate_ast"), ): # pyrefly: ignore [bad-argument-type] root = generate_ast(self.host_function, config, emit_repro_caller) self._register_cute_grouped_static_tail_specializations() if output_origin_lines is None: output_origin_lines = self.settings.output_origin_lines import_lines: list[str] = [] body_start = 0 for i, stmt in enumerate(root.body): if isinstance(stmt, (ast.Import, ast.ImportFrom)): if not ( isinstance(stmt, ast.ImportFrom) and stmt.module == "__future__" ): import_lines.append(ast.unparse(stmt)) continue body_start = i break else: body_start = len(root.body) body_root = ast.Module(body=root.body[body_start:], type_ignores=[]) ast.fix_missing_locations(body_root) # One optional AST processing step, then the single unparse. Both rewrites run # after generate_ast and outside the fake-tensor env above: jax_fn's launch # capture runs the compiled kernel on *real* tensors (which specializes # fake_args, so codegen must already be done); dep-free is pure-AST and # unaffected by placement. jax_fn is checked first -- it spans both dep modes. if options is not None and options.jax_fn: from .._compiler.output_code_utils import build_jax_fn_module from .._compiler.output_code_utils import capture_jax_launch_metadata jax_meta = capture_jax_launch_metadata(self, config) body_root = build_jax_fn_module( self, options, import_lines, body_root, jax_meta ) elif options is not None and not options.allow_helion_deps: from .._compiler.output_code_utils import build_dependency_free_code body_root = build_dependency_free_code( self, options, import_lines, body_root ) with measure("BoundKernel.unparse"): body = unparse(body_root, output_origin_lines=output_origin_lines) imports = "\n".join(import_lines) if imports: return f"from __future__ import annotations\n\n{imports}\n\n{body}" return f"from __future__ import annotations\n\n{body}" def to_triton_code( self, config: ConfigLike | None = None, *, emit_repro_caller: bool = False, output_origin_lines: bool | None = None, ) -> str: """Backward-compatible alias for :meth:`to_code`.""" return self.to_code( config, emit_repro_caller=emit_repro_caller, output_origin_lines=output_origin_lines, ) def compile_config( self, config: ConfigLike | None = None, *, allow_print: bool = True ) -> CompiledConfig: """ Compile the kernel for a specific configuration. Args: config: The configuration to compile the kernel with. allow_print: Set to suppress printing the output code when autotuning. Returns: CompiledConfig: A callable object representing the compiled kernel. """ if config is None: config = self._require_implicit_config() requested_config = self._normalize_config(config) config = self._normalized_config_copy(requested_config) dist_check_config_consistancy( config, process_group_name=self._env.process_group_name ) if (rv := self._compile_cache.get(config)) is not None: return rv device_index = ( self._env.device.index if self._env.device.index is not None else 0 ) self.env.backend.setup_compile_cache_dir(device_index) try: triton_code = self.to_triton_code( config, emit_repro_caller=self.settings.print_output_code ) # static_shapes=True keys a distinct BoundKernel per input shape (see # _tensor_key), but PyCodeCache keys compiled modules by SOURCE TEXT, so # two shapes that emit byte-identical source (e.g. compact_worklist, whose # token dim enters only via runtime offsets + a data-dependent loop) would # share one module. For backends whose generated module caches shape- # specific state (Backend.requires_shape_specialized_module -- e.g. Pallas, # whose output-meta descriptor / launcher cache / ds-pad decision / # signature lock are all monomorphic) that shared module returns the first # shape's cached output extent for both. Fold the specialization key into # the PyCodeCache key (via ``extra``, leaving the source untouched) so each # specialization gets its own module. Reusing the same key that # distinguishes BoundKernels (``self._base_spec_key``, computed from the # real args in __init__) guarantees distinct BoundKernel => distinct module # and already covers shape/dtype/stride recursively through nested # container args. cache_extra = "" if ( self.settings.static_shapes and self.env.backend.requires_shape_specialized_module ): cache_extra = repr(self._base_spec_key) with measure("BoundKernel.PyCodeCache.load"): module = PyCodeCache.load(triton_code, extra=cache_extra) self.env.backend.annotate_compiled_module( module, triton_code, self.kernel.name ) except Exception: log.warning( "Helion compiler triton codegen error for %s", self.format_kernel_decorator(requested_config, self.settings), exc_info=True, ) self.maybe_log_repro(log.warning, self.fake_args, config=requested_config) raise if allow_print: log.info("Output code written to: %s", module.__file__) log.debug("Debug string: \n%s", LazyString(lambda: self._debug_str())) # for distributed kernel, print rank1 code since rank0 # code can skip some offset computation. if ( not dist.is_initialized() or dist.get_rank() == 1 ) and self.settings.print_output_code: log.info("Output code: \n%s", triton_code) print(f"# Output code written to: {module.__file__}", file=sys.stderr) print(triton_code, file=sys.stderr) rv = getattr(module, self.kernel.name) self._compile_cache[config] = rv self._cache_path_map[config] = module.__file__ return rv def bench_compile_config( self, config: Config | dict[str, object] | None = None, *, allow_print: bool = True, ) -> Callable[..., object]: return self.compile_config(config, allow_print=allow_print) def extra_cache_key(self) -> str: """Return extra data folded into the disk-cache key. Returns ``""`` by default, leaving the cache key unchanged. """ return "" def supports_subprocess_benchmark(self) -> bool: return True def is_cacheable(self) -> bool: return True def get_cached_path(self, config: ConfigLike | None = None) -> str | None: """ Get the file path of the generated Triton code for a specific configuration. Args: config: The configuration to get the file path for. Returns: str | None: The file path of the generated Triton code, or None if not found. """ if config is None: config = self._require_implicit_config() requested_config = self._normalize_config(config) if requested_config in self._cache_path_map: return self._cache_path_map[requested_config] try: config = self._normalized_config_copy(requested_config) except exc.InvalidConfig: return None return self._cache_path_map.get(config, None) def _debug_str(self) -> str: """ Generate a debug string for the kernel. Returns: str: A string containing debug information about the kernel. """ if self.host_function is None: # In ref mode, host_function is not created return f"<BoundKernel {self.kernel.fn.__name__} in ref mode>" with self.env: return self.host_function.debug_str() def autotune( self, args: Sequence[object], *, force: bool = True, **kwargs: object, ) -> Config: """ Perform autotuning to find the optimal configuration for the kernel. This uses the default setting, you can call helion.autotune.* directly for more customization. If config= or configs= is provided to helion.kernel(), the search will be restricted to the provided configs. Use force=True to ignore the provided configs. Mutates self so that `__call__` will run the best config found. Args: args: Example arguments used for benchmarking during autotuning. force: If True, force full autotuning even if a config is provided. kwargs: Additional keyword options forwarded to the autotuner. Returns: Config: The best configuration found during autotuning. """ ephemeral = self.env.backend.make_ephemeral_cache() ctx = ephemeral if ephemeral is not None else contextlib.nullcontext() with ctx: config = self.env.backend.autotune(self, args, force=force, **kwargs) if ephemeral is not None: self.env.backend.finalize_ephemeral_cache(self, config) self.set_config(config) return config def set_config(self, config: ConfigLike) -> None: """ Set the configuration for the kernel and compile it. Mutates self so that `__call__` will run the provided config. Args: config: The configuration to set. """ config = self._normalize_config(config) self._run = self.compile_config(config) self._config = config counters["best_config_decorator"][ self.format_kernel_decorator(config, self.settings) ] = 1 def _specialize_extra(self) -> list[Callable[[Sequence[object]], Hashable]]: """ Returns a list of functions that will be called to generate extra specialization keys. This is used to specialize on the values hl.specialize()'ed arguments. Returns: list[Callable[[Sequence[object]], Hashable]]: A list of functions that generate extra specialization keys. """ if ( not self.env.specialized_vars and not self.env.specialized_strides and not self.env.tensor_descriptor_layout_guards ): return [] def make_extractor(v: Source) -> Callable[[Sequence[object]], Hashable]: if isinstance(v, TensorPropertySource): index = v.idx assert index is not None inner = make_extractor(v.base) if v.prop == TensorProperty.SIZE: def size_extractor( args: Sequence[object], _inner: Callable[[Sequence[object]], Hashable] = inner, _index: int = index, ) -> Hashable: result = _inner(args) # Handle list of tensors: return tuple of sizes for all tensors if isinstance(result, (list, tuple)): return tuple( cast("torch.Tensor", t).size(_index) for t in result ) return cast("torch.Tensor", result).size(_index) return size_extractor if v.prop == TensorProperty.STRIDE: def stride_extractor( args: Sequence[object], _inner: Callable[[Sequence[object]], Hashable] = inner, _index: int = index, ) -> Hashable: result = _inner(args) # Handle list of tensors: return tuple of strides for all tensors if isinstance(result, (list, tuple)): return tuple( cast("torch.Tensor", t).stride(_index) for t in result ) return cast("torch.Tensor", result).stride(_index) return stride_extractor raise exc.SpecializeArgType(v) if isinstance(v, GetItemSource): if not isinstance(v.index, (int, str)) or v.index_is_slice: raise exc.SpecializeArgType(v) inner = make_extractor(v.base) def getitem_extractor( args: Sequence[object], _inner: Callable[[Sequence[object]], Hashable] = inner, _index: int | str = v.index, ) -> Hashable: result = _inner(args) if isinstance(result, dict): return cast("Hashable", result[_index]) if isinstance(_index, str): return cast("Hashable", getattr(result, _index)) return cast("Sequence[Hashable]", result)[_index] return getitem_extractor if isinstance(v, LocalSource): index = arg_name_to_index[v.local_name] return operator.itemgetter(index) raise exc.SpecializeArgType(v) arg_name_to_index: dict[str, int] = { n: i for i, n in enumerate(self.kernel.signature.parameters.keys()) } extractors: list[Callable[[Sequence[object]], Hashable]] = [] extracted_strides: set[TensorPropertySource] = set() for v in sorted(self.env.specialized_vars, key=lambda v: v.name): source = self.env.shape_env.var_to_sources[v][0] extractors.append(make_extractor(source)) if ( isinstance(source, TensorPropertySource) and source.prop == TensorProperty.STRIDE and source.idx is not None ): extracted_strides.add(source) for source in sorted(self.env.specialized_strides, key=repr): if source in extracted_strides: continue extractors.append(make_extractor(source)) implicit_config = self._fixed_config_for_td_layout_guards() for source, guard in sorted( self.env.tensor_descriptor_layout_guards.items(), key=lambda item: repr(item[0]), ): if implicit_config is not None and not _td_layout_guard_active_for_config( guard, implicit_config ): continue extract_tensor = make_extractor(source) def td_layout_extractor( args: Sequence[object], _extract_tensor: Callable[ [Sequence[object]], Hashable ] = extract_tensor, _ndim: int = guard.ndim, _element_size: int = guard.element_size, ) -> Hashable: tensor = cast("torch.Tensor", _extract_tensor(args)) if tensor.ndim != _ndim: return ("ndim", tensor.ndim) return tensor_descriptor_layout_signature_from_strides( tensor.stride(), _element_size, ) extractors.append(td_layout_extractor) return extractors @contextlib.contextmanager def _runtime_arg_values_for_codegen(self) -> Generator[None, None, None]: values: dict[str, object] = self.env.runtime_arg_values_by_name if not values: values = { name: value for name, ref in self._runtime_tensor_refs_by_name.items() if (value := ref()) is not None } with self.env.use_runtime_arg_values(values): yield def _register_cute_grouped_static_tail_specializations(self) -> None: if self.kernel.settings.backend != "cute": return signature = self._base_spec_key descriptors = _cute_grouped_static_tail_extra_descriptors( self.env.cute_resolved_wrapper_plans ) if not descriptors: return seen = self.kernel._cute_grouped_static_tail_extra_descriptors.setdefault( signature, set(), ) new_descriptors: list[Hashable] = [] new_extractors: list[Callable[[Sequence[object]], Hashable]] = [] for descriptor in descriptors: if descriptor in seen: continue new_descriptors.append(descriptor) new_extractors.append(_make_cute_grouped_static_tail_extractor(descriptor)) runtime_args = tuple( self.env.runtime_arg_values_by_name.get(name) for name in self.kernel.signature.parameters ) self.kernel._extend_bound_kernel_specializations( self, signature, new_extractors, runtime_args, ) seen.update(new_descriptors) def _fixed_config_for_td_layout_guards(self) -> Config | None: """Return the fixed config if TD layout guards can be filtered safely.""" if self._config is not None: return self._config if self.kernel.settings.autotune_effort == "none" and ( len(self.kernel.configs) == 0 or self.settings.force_autotune ): return self.config_spec.default_config() if self.settings.force_autotune: return None if len(self.kernel.configs) == 1: return self.kernel.configs[0] return None def _user_provided_config(self) -> Config | None: """Return a config if the user explicitly provided one, else None. Checks the kernel's config list and settings to determine if a config can be resolved without autotuning. """ configs = self.kernel.configs if self.kernel.settings.autotune_effort == "none" and ( len(configs) == 0 or self.settings.force_autotune ): config = self.config_spec.default_config() if not is_ref_mode_enabled(self.kernel.settings): kernel_decorator = self.format_kernel_decorator(config, self.settings) print( f"Using implicit config: {kernel_decorator}", file=sys.stderr, ) return config if self.settings.force_autotune: return None if len(configs) == 1: return configs[0] return None def _implicit_config(self) -> Config | None: """ Returns a single config that is implicitly used by this kernel, if any. """ if self._config is not None: return self._config return self._user_provided_config() def _require_implicit_config(self) -> Config: """ Returns the implicit config for this kernel, or raises an error if no implicit config is available. """ if (config := self._implicit_config()) is None: raise RuntimeError("no config provided and no implicit config available") return config def ensure_config_exists(self, args: Sequence[object]) -> None: """ Ensure a config is available, triggering autotuning if needed. If an implicit config is available (from configs list or default), it will be used. Otherwise, autotuning will be triggered with the provided args. """ if self._config is not None: return # Already have a config if (config := self._implicit_config()) is not None: with measure("BoundKernel.set_config"): self.set_config(config) else: with measure("BoundKernel.autotune"): self.autotune(args, force=False) # pyrefly: ignore [bad-return] def run_ref(self, *args: object) -> _R: # Unwrap ConstExpr arguments clean_args = [] for arg in args: if isinstance(arg, ConstExpr): clean_args.append(arg.value) else: clean_args.append(arg) # Pass the config to RefModeContext with RefModeContext(self.env, self._config): result = self.kernel.fn(*clean_args) return cast("_R", result) def __call__(self, *args: object) -> _R: """ Execute the kernel with the given arguments. Args: args: The arguments to pass to the kernel. Returns: _R: The result of the kernel execution. """ if self.kernel._has_specialization_extras: rebound = self.kernel.bind(args) if rebound is not self: return rebound(*args) if self._run is None: with self._first_compile_lock: # Another caller may have discovered a late specialization while # this caller waited for the first compile. if self.kernel._has_specialization_extras: rebound = self.kernel.bind(args) if rebound is not self: return rebound(*args) if self._run is None: if is_ref_mode_enabled(self.kernel.settings): if (config := self._implicit_config()) is not None: self._config = config return self.run_ref(*args) runtime_args: dict[str, object] = { name: value for name, value in zip( self.kernel.signature.parameters, args, strict=False ) if isinstance(value, torch.Tensor) } with self.env.use_runtime_arg_values(runtime_args): self.ensure_config_exists(args) assert self._run is not None self.maybe_log_repro(log.warning, args) return self._run(*args) def backend_cache_key(self, config: ConfigLike | None = None) -> str | None: """ Return the backend cache key for the compiled kernel. For the Triton backend, this is the base32 encoding of the SHA-256 hash that Triton uses to cache compiled GPU binaries under ``TRITON_CACHE_DIR/<key>/``. For the CuTe backend, it is the base32 encoding of the SHA-256 hash of the compiled IR module, which names the ``CUTE_DSL_CACHE_DIR/<key>.mlir`` artifact. Args: config: The configuration to look up. Defaults to the implicit config. Returns: str | None: The cache key, or None if the kernel hasn't been JIT-compiled yet or the backend doesn't support cache keys. """ if config is None: config = self._require_implicit_config() config = self._normalized_config_copy(config) compiled_fn = self._compile_cache.get(config) if compiled_fn is None: return None return self.env.backend.compiled_cache_key(self, compiled_fn) def maybe_log_repro( self, log_func: Callable[[str], None], args: Sequence[object], config: Config | None = None, ) -> None: if not self.settings.print_repro: return effective_config = config or self._config assert effective_config is not None # Get kernel source try: raw_source = inspect.getsource(self.kernel.fn) source_lines = textwrap.dedent(raw_source).splitlines() # Skip decorator lines (including multi-line decorators) start_idx = 0 while start_idx < len(source_lines) and not source_lines[ start_idx ].lstrip().startswith("def "): start_idx += 1 kernel_body = "\n".join(source_lines[start_idx:]) except (OSError, TypeError): kernel_body = f"# Source unavailable for {self.kernel.fn.__module__}.{self.kernel.fn.__qualname__}" # Format decorator decorator = self.format_kernel_decorator(effective_config, self.settings) # Build output output_lines = [ "# === HELION KERNEL REPRO ===", "import helion", "import helion.language as hl", "import torch", "from torch._dynamo.testing import rand_strided", "", decorator, kernel_body, ] # Generate caller function if args: def _render_input_arg_assignment(name: str, value: object) -> list[str]: if isinstance(value, torch.Tensor): shape = tuple(int(d) for d in value.shape) stride = tuple(int(s) for s in value.stride()) device = str(value.device) dtype = str(value.dtype) lines = [ f"{name} = rand_strided({shape!r}, {stride!r}, dtype={dtype}, device={device!r})" ] if value.requires_grad: lines.append(f"{name}.requires_grad_(True)") return lines return [f"{name} = {value!r}"] sig_param_names = list(self.kernel.signature.parameters.keys()) assert len(args) == len(sig_param_names) output_lines.extend(["", "def helion_repro_caller():"]) output_lines.append(" torch.manual_seed(0)") arg_names: list[str] = [] for i, value in enumerate(args): var_name = sig_param_names[i] arg_names.append(var_name) # Add assignment lines with indentation for line in _render_input_arg_assignment(var_name, value): output_lines.append(f" {line}") # Add return statement call_args = ", ".join(arg_names) output_lines.append(f" return {self.kernel.name}({call_args})") output_lines.extend(["", "helion_repro_caller()"]) output_lines.append("# === END HELION KERNEL REPRO ===") repro_text = "\n" + "\n".join(output_lines) log_func(repro_text) class _KernelDecorator(Protocol): def __call__( self, fn: Callable[..., _R], ) -> Kernel[_R]: ... @overload def kernel( fn: Callable[..., _R], *, config: ConfigLike | None = None, configs: Sequence[ConfigLike] | None = None, key: Callable[..., Hashable] | None = None, **settings: object, ) -> Kernel[_R]: ... @overload def kernel( fn: None = None, *, config: ConfigLike | None = None, configs: Sequence[ConfigLike] | None = None, key: Callable[..., Hashable] | None = None, **settings: object, ) -> _KernelDecorator: ...
[docs] def kernel( fn: Callable[..., _R] | None = None, *, config: ConfigLike | None = None, configs: Sequence[ConfigLike] | None = None, key: Callable[..., Hashable] | None = None, **settings: object, ) -> Kernel[_R] | _KernelDecorator: """ Decorator to create a Kernel object from a Python function. Args: fn: The function to be wrapped by the Kernel. If None, a decorator is returned. config: A single configuration to use for the kernel. Refer to the ``helion.Config`` class for details. configs: A list of configurations to use for the kernel. Can only specify one of config or configs. Refer to the ``helion.Config`` class for details. key: Optional callable returning a hashable that augments the specialization key. settings: Keyword arguments representing settings for the Kernel. Can also use settings=Settings(...) to pass a Settings object directly. Refer to the ``helion.Settings`` class for available options. Returns: object: A Kernel object or a decorator that returns a Kernel object. See Also: - :class:`~helion.Settings`: Controls compilation behavior and debugging options - :class:`~helion.Config`: Controls GPU execution parameters and optimization strategies """ if config is not None: assert not configs, "Cannot specify both config and configs" configs = [config] elif configs is None: configs = [] if settings_obj := settings.get("settings"): assert len(settings) == 1, "settings must be the only keyword argument" assert isinstance(settings_obj, Settings), "settings must be a Settings object" else: settings_obj = Settings(**settings) if fn is None: return functools.partial( kernel, configs=configs, settings=settings_obj, key=key, ) return Kernel( fn, configs=configs, settings=settings_obj, key=key, )
def _hashable_dim(s: int | torch.SymInt) -> Hashable: if isinstance(s, torch.SymInt): return (id(s.node.shape_env), s.node.expr) return s def _safe_bucket_dim(s: int | torch.SymInt) -> Hashable: if isinstance(s, torch.SymInt): return (id(s.node.shape_env), s.node.expr) # Dynamic-shape kernels should not get separate bound kernels for sizes # 0 or 1. Keep 2 as the canonical "dynamic dimension" bucket that was # already used for all concrete sizes >= 2. return 2 _EMPTY_FROZENSET: frozenset[int] = frozenset() def _bucketed_size(obj: torch.Tensor) -> tuple[Hashable, ...]: sz = obj.size() n = len(sz) if n == 1: return (_safe_bucket_dim(sz[0]),) if n == 2: return (_safe_bucket_dim(sz[0]), _safe_bucket_dim(sz[1])) if n == 3: return ( _safe_bucket_dim(sz[0]), _safe_bucket_dim(sz[1]), _safe_bucket_dim(sz[2]), ) return tuple(_safe_bucket_dim(s) for s in sz) def _hashable_dims(dims: Sequence[int | torch.SymInt]) -> tuple[Hashable, ...]: n = len(dims) if n == 1: return (_hashable_dim(dims[0]),) if n == 2: return (_hashable_dim(dims[0]), _hashable_dim(dims[1])) if n == 3: return (_hashable_dim(dims[0]), _hashable_dim(dims[1]), _hashable_dim(dims[2])) return tuple(_hashable_dim(s) for s in dims) def _concrete_tensor_key(fn: Kernel, obj: torch.Tensor) -> Hashable: # Fast extractor for plain ``torch.Tensor`` / ``torch.nn.Parameter``: # exact-type dispatch guarantees concrete int sizes/strides, so # ``torch.Size`` and the stride tuple can be used directly (both are # tuple subclasses that hash/compare identically to plain int tuples). # The ``_hashable_dims`` wrap in ``_tensor_key`` exists only to # normalize SymInts, which appear on FakeTensors during tracing. si = getattr(obj, "_dynamo_static_indices", None) static_indices = frozenset(si) if si is not None else _EMPTY_FROZENSET if fn.settings.static_shapes: return (obj.dtype, obj.size(), obj.stride(), static_indices) bucketed = _bucketed_size(obj) if fn.settings.index_dtype is None: try: needs_int64 = bool(obj.numel() > _INT32_INDEX_LIMIT) except RuntimeError: needs_int64 = True # unbacked SymInt return ( obj.dtype, bucketed, needs_int64, static_indices, ) return ( obj.dtype, bucketed, static_indices, ) def _tensor_key(fn: Kernel, obj: torch.Tensor) -> Hashable: si = getattr(obj, "_dynamo_static_indices", None) static_indices = frozenset(si) if si is not None else _EMPTY_FROZENSET if fn.settings.static_shapes: return ( obj.dtype, _hashable_dims(obj.size()), _hashable_dims(obj.stride()), static_indices, ) bucketed = _bucketed_size(obj) if fn.settings.index_dtype is None: try: needs_int64 = bool(obj.numel() > _INT32_INDEX_LIMIT) except RuntimeError: needs_int64 = True # unbacked SymInt return ( obj.dtype, bucketed, needs_int64, static_indices, ) return ( obj.dtype, bucketed, static_indices, ) def _sequence_key(fn: Kernel, obj: Sequence) -> Hashable: return type(obj), tuple([fn._specialization_key(item) for item in obj]) def _mapping_key( fn: Kernel, obj: dict[str | int, object], real_type: type[object] ) -> Hashable: return real_type, tuple( sorted((k, fn._specialization_key(v)) for k, v in obj.items()) ) def _number_key(fn: Kernel, n: float | bool) -> object: return type(n) def _function_key(fn: Kernel, obj: types.FunctionType) -> object: if obj.__closure__: closures = [ fn._specialization_key(cell.cell_contents) for cell in obj.__closure__ ] return (obj.__code__, *closures) return obj.__code__ def _cute_grouped_layout_has_m_tail( layout_values: tuple[int, ...], *, bm: int, group_count: int, ) -> bool | None: cursor = 0 has_m_tail = False for expected_group in range(group_count): if expected_group > 0: next_m_boundary = ((cursor + bm - 1) // bm) * bm while ( cursor < len(layout_values) and cursor < next_m_boundary and layout_values[cursor] < 0 ): cursor += 1 if cursor != next_m_boundary or ( cursor < len(layout_values) and layout_values[cursor] < 0 ): return None if cursor >= len(layout_values) or layout_values[cursor] != expected_group: return None start = cursor while cursor < len(layout_values) and layout_values[cursor] == expected_group: cursor += 1 actual_m = cursor - start if start % bm != 0: return None has_m_tail = has_m_tail or actual_m % bm != 0 if cursor != len(layout_values): if all(value < 0 for value in layout_values[cursor:]): cursor = len(layout_values) if cursor != len(layout_values): return None return has_m_tail def _cute_grouped_static_tail_extra_descriptors( plans: Sequence[dict[str, object]], ) -> tuple[Hashable, ...]: descriptors: list[Hashable] = [] for plan in plans: if plan.get("kind") != "tcgen05_grouped_static_persistent" or bool( plan.get("worklist_metadata") ): continue if not ( isinstance(plan.get("grouped_static_has_m_tail"), bool) or isinstance(plan.get("grouped_static_has_n_tail"), bool) ): continue layout_idx = plan.get("layout_bind_idx") group_count = plan.get("group_count") bm = plan.get("bm") if not ( isinstance(layout_idx, int) and isinstance(group_count, int) and isinstance(bm, int) ): continue n_sizes_idx = plan.get("n_sizes_bind_idx") bn = plan.get("bn") descriptors.append( ( "cute_grouped_static_tail", layout_idx, group_count, bm, n_sizes_idx if isinstance(n_sizes_idx, int) else None, bn if isinstance(bn, int) else None, ) ) return tuple(sorted(descriptors, key=repr)) def _cute_int_1d_tensor_values( value: object, cache: WeakIdKeyDictionary, ) -> tuple[int, ...] | None: if not ( isinstance(value, torch.Tensor) and value.ndim == 1 and value.dtype in (torch.int32, torch.int64) ): return None if torch.is_inference(value): return tuple(int(v) for v in value.detach().cpu().tolist()) signature = ( int(value._version), int(value.data_ptr()), tuple(value.shape), tuple(value.stride()), value.dtype, ) try: cached_signature, cached_values = cache[value] except KeyError: pass else: if cached_signature == signature: return cached_values values = tuple(int(v) for v in value.detach().cpu().tolist()) cache[value] = (signature, values) return values def _make_cute_grouped_static_tail_extractor( descriptor: Hashable, ) -> Callable[[Sequence[object]], Hashable]: ( _label, layout_idx, group_count, bm, n_sizes_idx, bn, ) = cast("tuple[object, int, int, int, int | None, int | None]", descriptor) tensor_values_cache: WeakIdKeyDictionary = WeakIdKeyDictionary() def cute_grouped_static_tail_extractor( args: Sequence[object], *, _descriptor: Hashable = descriptor, _layout_idx: int = layout_idx, _group_count: int = group_count, _bm: int = bm, _n_sizes_idx: int | None = n_sizes_idx, _bn: int | None = bn, ) -> Hashable: layout_has_m_tail: bool | None = None if _layout_idx < len(args): layout_values = _cute_int_1d_tensor_values( args[_layout_idx], tensor_values_cache, ) if layout_values is not None: layout_has_m_tail = _cute_grouped_layout_has_m_tail( layout_values, bm=_bm, group_count=_group_count, ) n_sizes_has_n_tail: bool | None = None if _n_sizes_idx is not None and _bn is not None and _n_sizes_idx < len(args): n_sizes_values = _cute_int_1d_tensor_values( args[_n_sizes_idx], tensor_values_cache, ) if n_sizes_values is not None and len(n_sizes_values) == _group_count: n_sizes_has_n_tail = any( group_n % _bn != 0 for group_n in n_sizes_values ) return (_descriptor, layout_has_m_tail, n_sizes_has_n_tail) return cute_grouped_static_tail_extractor def _graph_module_key(fn: Kernel, obj: torch.fx.GraphModule) -> Hashable: """Generate a specialization key for GraphModule arguments.""" # Check if already cached if obj in _graph_module_hash_cache: return _graph_module_hash_cache[obj] # Check for unsupported operations unsupported_ops = { node.op for node in itertools.chain( obj.graph.find_nodes(op="call_module"), obj.graph.find_nodes(op="get_attr"), ) } if unsupported_ops: raise exc.GraphModuleUnsupportedOps(", ".join(sorted(unsupported_ops))) _graph_module_hash_cache[obj] = rv = str(compiled_fx_graph_hash(obj, [], {}, [])) return rv _specialization_extractors: dict[ type[object] | str, Callable[[Kernel, object], Hashable], # pyrefly: ignore [bad-assignment] ] = { # Exact-type dispatch (see ``_specialization_key``): plain tensors and # Parameters always have concrete int sizes/strides and take the fast # extractor. Subclasses (FakeTensor below, or anything hitting the # ``isinstance`` fallback) go through SymInt-safe ``_tensor_key``. torch.Tensor: _concrete_tensor_key, torch.nn.Parameter: _concrete_tensor_key, FakeTensor: _tensor_key, # SymInt-safe extractor for torch.Tensor subclasses reached via the # isinstance fallback in ``_specialization_key`` (string key so the # fallback stays loosely typed, like "namedtuple" / "dataclass"). "tensor_subclass": _tensor_key, torch.dtype: lambda fn, x: x, torch.device: lambda fn, x: x, int: _number_key, float: _number_key, bool: _number_key, str: lambda fn, x: x, list: _sequence_key, tuple: _sequence_key, # pyrefly: ignore [bad-argument-type] dict: lambda fn, x: _mapping_key(fn, x, type(x)), # pyrefly: ignore [missing-attribute] "namedtuple": lambda fn, x: _mapping_key(fn, x._asdict(), type(x)), # pyrefly: ignore [no-matching-overload, bad-argument-type] "dataclass": lambda fn, x: _mapping_key(fn, dataclasses.asdict(x), type(x)), types.FunctionType: _function_key, types.BuiltinFunctionType: lambda fn, x: x, torch.fx.GraphModule: _graph_module_key, # pyrefly: ignore [missing-attribute] ConstExpr: lambda fn, x: x.value, type(None): lambda fn, x: None, } def _find_device(args: tuple[object, ...]) -> torch.device: """ Extract the device from the arguments. Args: args: The arguments to extract the device from. Returns: torch.device: The extracted device """ for arg in args: if isinstance(arg, torch.device): return arg if isinstance(arg, torch.Tensor): return arg.device if isinstance(arg, ConstExpr): continue # a constexpr-wrapped value carries no device; skip it if isinstance(arg, (tuple, list)): for item in arg: try: return _find_device((item,)) except exc.NoTensorArgs: pass elif isinstance(arg, dict): for item in arg.values(): try: return _find_device((item,)) except exc.NoTensorArgs: pass raise exc.NoTensorArgs def _maybe_skip_dtype_check_in_meta_registrations() -> ( contextlib.AbstractContextManager[None, None] ): # pyrefly: ignore [implicit-import] if hasattr(torch.fx.experimental._config, "skip_dtype_check_in_meta_registrations"): # pyrefly: ignore [implicit-import, missing-attribute] return torch.fx.experimental._config.patch( skip_dtype_check_in_meta_registrations=True ) return contextlib.nullcontext()