Skip to content

vllm.model_executor.warmup.jit_warmup_triton_helper

Classes:

Functions:

TritonCompileKey dataclass

Triton-derived compile key with one non-comparing replay input.

Source code in vllm/model_executor/warmup/jit_warmup_triton_helper.py
@dataclass(frozen=True)
class TritonCompileKey:
    """Triton-derived compile key with one non-comparing replay input."""

    jit_keys: frozenset[TritonJitKey]
    inputs: tuple[tuple[str, Any], ...] = field(compare=False, hash=False, repr=False)

TritonJitKey dataclass

Process-local identity of one Triton JIT specialization.

Source code in vllm/model_executor/warmup/jit_warmup_triton_helper.py
@dataclass(frozen=True)
class TritonJitKey:
    """Process-local identity of one Triton JIT specialization."""

    jit_function_id: int
    jit_function_key: Hashable
    device: Hashable
    cache_key: Hashable

TritonKernelDispatcher

Bases: Protocol[P]

Callable Triton launcher with compile-time warmup support.

Source code in vllm/model_executor/warmup/jit_warmup_triton_helper.py
class TritonKernelDispatcher(Protocol[P]):
    """Callable Triton launcher with compile-time warmup support."""

    def __call__(self, *args: P.args, **kwargs: P.kwargs) -> Any: ...

    def register_warmup(self, *args: Any, **kwargs: Any) -> None: ...

TritonWarmupTensor dataclass

Compile-only tensor metadata used by Triton warmup.

strides=None represents compact row-major storage. Pass explicit strides whenever the runtime tensor can be padded, transposed, or otherwise strided.

Source code in vllm/model_executor/warmup/jit_warmup_triton_helper.py
@dataclass(frozen=True)
class TritonWarmupTensor:
    """Compile-only tensor metadata used by Triton warmup.

    ``strides=None`` represents compact row-major storage. Pass explicit strides
    whenever the runtime tensor can be padded, transposed, or otherwise strided.
    """

    dtype: Any
    aligned: bool = True
    shape: tuple[int, ...] = (1,)
    strides: tuple[int, ...] | None = None
    init: Any = 0

    def data_ptr(self) -> int:
        return 0 if self.aligned else 1

    def ptr_range(self) -> int:
        return 0

    @property
    def device(self) -> Any:
        import torch

        return torch.device(current_platform.device_type)

    def stride(self, dim: int | None = None) -> int | tuple[int, ...]:
        if self.strides is None:
            strides: list[int] = []
            stride = 1
            for size in reversed(self.shape):
                strides.append(stride)
                stride *= size
            result = tuple(reversed(strides))
        else:
            result = self.strides
        return result if dim is None else result[dim]

VllmTritonJitKernel

Bases: VllmJitKernel[CompileKeyT], Generic[CompileKeyT]

Triton owner whose runtime launch specification is reused for warmup.

Methods:

  • compile_many –

    Compile CUDA startup warmup variants in parallel, then wait.

  • warmup_inputs –

    Return runtime-shaped inputs that reproduce one compile key.

Source code in vllm/model_executor/warmup/jit_warmup_triton_helper.py
class VllmTritonJitKernel(VllmJitKernel[CompileKeyT], Generic[CompileKeyT]):
    """Triton owner whose runtime launch specification is reused for warmup."""

    kernel: Any
    _warming = False
    _warming_compile_key: CompileKeyT | None = None
    _run_autotune = False

    @abstractmethod
    def warmup_inputs(self, compile_key: CompileKeyT) -> dict[str, Any]:
        """Return runtime-shaped inputs that reproduce one compile key."""
        raise NotImplementedError

    def compile(self, compile_key: CompileKeyT) -> None:
        inputs = self.warmup_inputs(compile_key)
        self._warming = True
        self._warming_compile_key = compile_key
        try:
            cast(Callable[..., None], self)(**inputs)
        finally:
            self._warming = False
            self._warming_compile_key = None

    def compile_many(self, compile_keys: Iterable[CompileKeyT]) -> None:
        """Compile CUDA startup warmup variants in parallel, then wait."""
        keys = list(compile_keys)
        num_threads = min(envs.VLLM_TRITON_JIT_WARMUP_NUM_THREADS, len(keys))
        async_compile = getattr(triton, "AsyncCompileMode", None)
        # AMD LLVM code generation can abort under concurrent compilation.
        if (
            not current_platform.is_cuda()
            or self._run_autotune
            or async_compile is None
            or num_threads <= 1
        ):
            return super().compile_many(keys)

        def compile_all() -> None:
            with (
                ThreadPoolExecutor(max_workers=num_threads) as executor,
                async_compile(executor),
            ):
                # Dispatch stays on this thread: compile() mutates owner state.
                # Only Triton's underlying compiler work runs in the pool.
                super(VllmTritonJitKernel, self).compile_many(keys)

        # Older Triton versions leave async mode active on compilation failure.
        # Isolate that state so it cannot affect later runtime JIT compilation.
        copy_context().run(compile_all)

    @cached_property
    def _kernel_arg_names(self) -> tuple[str, ...]:
        arg_names = getattr(self.kernel, "arg_names", None)
        if arg_names is not None:
            return tuple(arg_names)
        wrapped = getattr(self.kernel, "func", None)
        if wrapped is not None:
            return tuple(inspect.signature(wrapped).parameters)
        raise TypeError(
            f"Cannot inspect kernel parameters for {type(self.kernel).__name__}"
        )

    def _prepare_launch_kwargs(
        self,
        input_names: tuple[str, ...],
        input_values: tuple[Any, ...],
        launch_kwargs: dict[str, Any],
    ) -> dict[str, Any]:
        plan = _launch_binding_plan(
            self._kernel_arg_names, input_names, tuple(launch_kwargs)
        )
        for index, target in plan.input_targets:
            launch_kwargs[target] = input_values[index]
        for target, source, dim in plan.stride_targets:
            launch_kwargs[target] = launch_kwargs[source].stride(dim)
        return launch_kwargs

    def launch(
        self,
        launch_spec: LaunchSpec,
        inputs: Mapping[str, Any],
    ) -> Any:
        if len(launch_spec) == 3:
            grid, launch_kwargs, outputs = launch_spec
        else:
            grid, launch_kwargs = launch_spec
            outputs = None
        kwargs = self._prepare_launch_kwargs(
            tuple(inputs), tuple(inputs.values()), dict(launch_kwargs)
        )
        kernel = kwargs.pop("kernel", None) or self.kernel
        runtime_launcher = kwargs.pop("_runtime_launcher", None)
        runtime_launcher_arg_count = kwargs.pop("_runtime_launcher_arg_count", 0)
        if self._warming:
            if self._run_autotune:
                assert grid is not None
                kernel[grid](**kwargs)
                return outputs
            kwargs = {
                name: _triton_metadata_arg(value) for name, value in kwargs.items()
            }
            if (
                self._warming_compile_key is not None
                and "launch_pdl" in kwargs
                and hasattr(self._warming_compile_key, "launch_pdl")
            ):
                kwargs["launch_pdl"] = self._warming_compile_key.launch_pdl
            warmup = getattr(kernel, "warmup", None)
            assert warmup is not None
            warmup(grid=(1,), **kwargs)
            return outputs
        if grid is None:
            return outputs
        if runtime_launcher is not None:
            regular_args = [
                kwargs.pop(name)
                for name in self._kernel_arg_names[:runtime_launcher_arg_count]
            ]
            runtime_launcher(kernel, grid, *regular_args, **kwargs)
        else:
            kernel[grid](**kwargs)
        return outputs

compile_many(compile_keys)

Compile CUDA startup warmup variants in parallel, then wait.

Source code in vllm/model_executor/warmup/jit_warmup_triton_helper.py
def compile_many(self, compile_keys: Iterable[CompileKeyT]) -> None:
    """Compile CUDA startup warmup variants in parallel, then wait."""
    keys = list(compile_keys)
    num_threads = min(envs.VLLM_TRITON_JIT_WARMUP_NUM_THREADS, len(keys))
    async_compile = getattr(triton, "AsyncCompileMode", None)
    # AMD LLVM code generation can abort under concurrent compilation.
    if (
        not current_platform.is_cuda()
        or self._run_autotune
        or async_compile is None
        or num_threads <= 1
    ):
        return super().compile_many(keys)

    def compile_all() -> None:
        with (
            ThreadPoolExecutor(max_workers=num_threads) as executor,
            async_compile(executor),
        ):
            # Dispatch stays on this thread: compile() mutates owner state.
            # Only Triton's underlying compiler work runs in the pool.
            super(VllmTritonJitKernel, self).compile_many(keys)

    # Older Triton versions leave async mode active on compilation failure.
    # Isolate that state so it cannot affect later runtime JIT compilation.
    copy_context().run(compile_all)

warmup_inputs(compile_key) abstractmethod

Return runtime-shaped inputs that reproduce one compile key.

Source code in vllm/model_executor/warmup/jit_warmup_triton_helper.py
@abstractmethod
def warmup_inputs(self, compile_key: CompileKeyT) -> dict[str, Any]:
    """Return runtime-shaped inputs that reproduce one compile key."""
    raise NotImplementedError

_triton_key_deriver(kernel)

Prepare Triton's key derivation once for one kernel and device.

Source code in vllm/model_executor/warmup/jit_warmup_triton_helper.py
def _triton_key_deriver(
    kernel: Any,
) -> Callable[[Mapping[str, Any]], set[TritonJitKey]]:
    """Prepare Triton's key derivation once for one kernel and device."""
    knobs = triton.knobs
    Autotuner = triton.runtime.autotuner.Autotuner
    Heuristics = triton.runtime.autotuner.Heuristics
    driver = triton.runtime.driver
    JITFunction = triton.runtime.jit.JITFunction
    compute_cache_key = triton.runtime.jit.compute_cache_key

    if isinstance(kernel, Heuristics):
        derive_inner = _triton_key_deriver(kernel.fn)

        def derive_heuristic(kwargs: Mapping[str, Any]) -> set[TritonJitKey]:
            heuristic_kwargs = dict(kwargs)
            for name, heuristic in kernel.values.items():
                heuristic_kwargs[name] = heuristic(heuristic_kwargs)
            return derive_inner(heuristic_kwargs)

        return derive_heuristic

    if isinstance(kernel, Autotuner):
        derive_inner = _triton_key_deriver(kernel.fn)

        def derive_autotuned(kwargs: Mapping[str, Any]) -> set[TritonJitKey]:
            previous_nargs = getattr(kernel, "nargs", None)
            kernel.nargs = {}
            try:
                configs = kernel.prune_configs(dict(kwargs))
            finally:
                kernel.nargs = previous_nargs
            keys: set[TritonJitKey] = set()
            for config in configs:
                conflicts = kwargs.keys() & config.kwargs.keys()
                if conflicts:
                    names = ", ".join(sorted(conflicts))
                    raise ValueError(f"Conflicting autotune parameters: {names}")
                keys.update(derive_inner(dict(kwargs) | config.all_kwargs()))
            return keys

        return derive_autotuned

    if not isinstance(kernel, JITFunction):
        raise TypeError(f"Unsupported Triton kernel wrapper: {type(kernel).__name__}")

    device = cast(Hashable, driver.active.get_current_device())
    _, kernel_key_cache, _, _, binder = kernel.device_caches[device]
    jit_function_id = id(kernel)
    jit_function_key = cast(Hashable, kernel.cache_key)
    debug = kernel.debug or knobs.runtime.debug
    instrumentation_mode = knobs.compilation.instrumentation_mode
    fpsan_homomorphic_casts = getattr(
        knobs.compilation, "fpsan_homomorphic_casts", None
    )
    inspection_hook = knobs.runtime.add_stages_inspection_hook

    def derive_jit(kwargs: Mapping[str, Any]) -> set[TritonJitKey]:
        binder_kwargs = dict(kwargs)
        binder_kwargs["debug"] = binder_kwargs.get("debug", debug) or debug
        binder_kwargs["instrumentation_mode"] = instrumentation_mode
        if fpsan_homomorphic_casts is not None:
            binder_kwargs["fpsan_homomorphic_casts"] = fpsan_homomorphic_casts
        _, specialization, options = binder(**binder_kwargs)
        if inspection_hook is not None:
            _, inspection_hash = inspection_hook()
            specialization.append(f'("custom_pipeline", {inspection_hash})')
        cache_key = cast(
            Hashable,
            compute_cache_key(kernel_key_cache, specialization, options),
        )
        return {
            TritonJitKey(
                jit_function_id,
                jit_function_key,
                device,
                cache_key,
            )
        }

    return derive_jit

triton_kernel_dispatcher_with_warmup(*, warmup_inputs, kernel=None)

triton_kernel_dispatcher_with_warmup(
    *,
    warmup_inputs: Callable[..., WarmupCases],
    kernel: None = None,
) -> Callable[[Any], Any]
triton_kernel_dispatcher_with_warmup(
    *,
    warmup_inputs: Callable[..., WarmupCases],
    kernel: Any,
) -> Callable[
    [Callable[P, DispatchSpec]], TritonKernelDispatcher[P]
]

Decorate a dispatch function or a native Triton kernel for warmup.

Source code in vllm/model_executor/warmup/jit_warmup_triton_helper.py
def triton_kernel_dispatcher_with_warmup(
    *,
    warmup_inputs: Callable[..., WarmupCases],
    kernel: Any | None = None,
) -> Callable[[Any], Any]:
    """Decorate a dispatch function or a native Triton kernel for warmup."""

    def decorate(dispatch: Any) -> Any:
        owner = _DecoratedTritonJitKernel(
            dispatch if kernel is None else kernel,
            warmup_inputs,
            None if kernel is None else dispatch,
        )
        update_wrapper(owner, dispatch, updated=())
        return owner

    return decorate

triton_scalar_specialization_rep(value)

Return an integer with the same default Triton JIT specialization.

For an ordinary integer argument, Triton's cache key contains its inferred type (i32, i64, or u64) and one of three value classes:

  • 1 is specialized as the exact constant 1.
  • Multiples of 16 receive a tt.divisibility = 16 attribute.
  • All other values have no value specialization.

Warmup only needs one concrete value for each cache-key class. This helper returns 1 for the exact-one class and otherwise returns a divisible or generic representative while preserving the inferred integer type.

This applies only to non-constexpr integer arguments using Triton's default specialization. Do not use it for arguments listed in do_not_specialize or do_not_specialize_on_alignment.

Source code in vllm/model_executor/warmup/jit_warmup_triton_helper.py
def triton_scalar_specialization_rep(value: int) -> int:
    """Return an integer with the same default Triton JIT specialization.

    For an ordinary integer argument, Triton's cache key contains its inferred
    type (``i32``, ``i64``, or ``u64``) and one of three value classes:

    * ``1`` is specialized as the exact constant ``1``.
    * Multiples of 16 receive a ``tt.divisibility = 16`` attribute.
    * All other values have no value specialization.

    Warmup only needs one concrete value for each cache-key class. This helper
    returns ``1`` for the exact-one class and otherwise returns a divisible or
    generic representative while preserving the inferred integer type.

    This applies only to non-``constexpr`` integer arguments using Triton's
    default specialization. Do not use it for arguments listed in
    ``do_not_specialize`` or ``do_not_specialize_on_alignment``.
    """
    if value == 1:
        return 1

    if -(1 << 31) <= value < (1 << 31):
        divisible_rep = 16
        generic_rep = 2
    elif -(1 << 63) <= value < (1 << 63):
        divisible_rep = 1 << 31
        generic_rep = (1 << 31) + 1
    elif 0 <= value < (1 << 64):
        divisible_rep = 1 << 63
        generic_rep = (1 << 63) + 1
    else:
        raise OverflowError(f"Integer {value} is outside Triton's scalar range")

    return divisible_rep if value % 16 == 0 else generic_rep

triton_warmup_inputs(kernel, *args, grid, pointer_dtypes=None, **kwargs)

Build launcher inputs from Triton's native positional argument order.

Source code in vllm/model_executor/warmup/jit_warmup_triton_helper.py
def triton_warmup_inputs(
    kernel: Any,
    *args: Any,
    grid: tuple[int, ...],
    pointer_dtypes: (
        Mapping[Any, Iterable[str]] | Iterable[tuple[Any, Iterable[str]]] | None
    ) = None,
    **kwargs: Any,
) -> dict[str, Any]:
    """Build launcher inputs from Triton's native positional argument order."""
    arg_names = tuple(kernel.arg_names)
    if len(args) > len(arg_names):
        raise ValueError(
            f"Received {len(args)} positional inputs for {len(arg_names)} "
            "Triton kernel arguments"
        )
    inputs = dict(zip(arg_names, args))
    duplicate_names = inputs.keys() & kwargs.keys()
    if duplicate_names:
        raise ValueError(
            f"Triton inputs passed twice: {', '.join(sorted(duplicate_names))}"
        )
    inputs.update(kwargs)
    if pointer_dtypes is not None:
        function_def = get_function_source_node(kernel)
        if not isinstance(function_def, ast.FunctionDef):
            raise ValueError("Expected Triton kernel to be defined as a function")
        pointer_names = _pointer_arg_names(function_def, arg_names)
        groups = (
            pointer_dtypes.items()
            if isinstance(pointer_dtypes, Mapping)
            else pointer_dtypes
        )
        for dtype, names in groups:
            for name in names:
                if name not in arg_names:
                    raise ValueError(f"Unknown Triton pointer argument: {name}")
                if name not in pointer_names:
                    raise ValueError(f"Triton argument is not a pointer: {name}")
                if name in inputs:
                    raise ValueError(f"Triton input passed twice: {name}")
                inputs[name] = TritonWarmupTensor(dtype)
        missing_pointers = pointer_names - inputs.keys()
        if missing_pointers:
            names = ", ".join(sorted(missing_pointers))
            raise ValueError(f"Missing Triton pointer inputs: {names}")
    return {"grid": grid, **inputs}