Skip to content

vllm.model_executor.warmup.jit_warmup

Shared interfaces and tracing helpers for explicit JIT warmup keys.

Classes:

  • JitWarmupRegistry –

    Collect and compile JIT kernels selected during runner setup.

  • VllmJitKernel –

    Kernel wrapper that owns dispatch, warmup keys, and compilation.

  • WarmupChoices –

    Expand an explicit finite set of values in a traced warmup method.

  • WarmupIntRange –

    Expand integers with range semantics or a custom monotonic progression.

Functions:

  • kernel_launcher –

    Delegate a declarative launch specification to its backend owner.

  • zip_inputs –

    Group row-wise dispatch inputs that should be expanded in lockstep.

JitWarmupRegistry

Collect and compile JIT kernels selected during runner setup.

Methods:

  • activate –

    Collect registrations made in this context.

  • capture –

    Collect warmup registrations made while the decorated callable runs.

  • register –

    Register a kernel with the active registry, if one exists.

  • warmup –

    Expand registrations and compile each wrapper/key pair once.

Source code in vllm/model_executor/warmup/jit_warmup.py
class JitWarmupRegistry:
    """Collect and compile JIT kernels selected during runner setup."""

    _active: ContextVar[JitWarmupRegistry | None] = ContextVar(
        "active_jit_warmup_registry",
        default=None,
    )

    def __init__(self, vllm_config: Any) -> None:
        self.vllm_config = vllm_config
        self._registrations: dict[
            VllmJitKernel[Any],
            list[tuple[tuple[Any, ...], dict[str, Any]]],
        ] = {}

    @classmethod
    def capture(cls, init_fn: Callable[..., None]) -> Callable[..., None]:
        """Collect warmup registrations made while the decorated callable runs."""

        @wraps(init_fn)
        def wrapped(instance: Any, vllm_config: Any, *args: Any, **kwargs: Any) -> None:
            registry = cls(vllm_config)
            instance.jit_warmup_registry = registry
            with registry.activate():
                init_fn(instance, vllm_config, *args, **kwargs)

        return wrapped

    @contextmanager
    def activate(self) -> Iterator[None]:
        """Collect registrations made in this context."""
        token = self._active.set(self)
        try:
            yield
        finally:
            self._active.reset(token)

    @classmethod
    def register(
        cls,
        kernel: VllmJitKernel[Any],
        *args: Any,
        **kwargs: Any,
    ) -> None:
        """Register a kernel with the active registry, if one exists."""
        registry = cls._active.get()
        if registry is not None:
            registry._add(kernel, args, kwargs)

    def _add(
        self,
        kernel: VllmJitKernel[Any],
        args: tuple[Any, ...],
        kwargs: dict[str, Any],
    ) -> None:
        registrations = self._registrations.setdefault(kernel, [])
        # Every layer of a deep model appends an identical (args, kwargs) here
        # (e.g. a 61-layer DSA model registers the same pack (dtype, pad_value)
        # 61 times); each expands to the same compile keys, so tracing
        # get_warmup_keys once per distinct registration is sufficient. Dedup on
        # identity-or-equality -- identity short-circuits shared singletons like
        # vllm_config before any deep __eq__.
        if any(
            _same_registration(registered, (args, kwargs))
            for registered in registrations
        ):
            return
        registrations.append((args, kwargs))

    def __len__(self) -> int:
        return sum(len(registrations) for registrations in self._registrations.values())

    def warmup(self) -> None:
        """Expand registrations and compile each wrapper/key pair once."""
        from tqdm import tqdm

        from vllm.distributed import is_global_first_rank

        kernel_items: list[tuple[VllmJitKernel[Any], dict[Any, None]]] = []
        for kernel, registrations in self._registrations.items():
            compile_keys: dict[Any, None] = {}
            for args, kwargs in registrations:
                if (
                    not args
                    and not kwargs
                    and "vllm_config"
                    in inspect.signature(kernel.get_warmup_keys).parameters
                ):
                    kwargs = {"vllm_config": self.vllm_config}
                for compile_key in kernel.get_warmup_keys(*args, **kwargs):
                    compile_keys[compile_key] = None
            if compile_keys:
                kernel_items.append((kernel, compile_keys))

        if not kernel_items:
            return

        total_keys = sum(len(compile_keys) for _, compile_keys in kernel_items)
        with tqdm(
            kernel_items,
            desc=f"JIT kernel warmup ({total_keys} compile keys)",
            disable=not is_global_first_rank(),
            dynamic_ncols=True,
            unit="kernel",
        ) as progress:
            for kernel, compile_keys in progress:
                progress.set_postfix_str(
                    f"{kernel.__class__.__name__} ({len(compile_keys)} keys)",
                    refresh=False,
                )
                kernel.compile_many(compile_keys)

activate()

Collect registrations made in this context.

Source code in vllm/model_executor/warmup/jit_warmup.py
@contextmanager
def activate(self) -> Iterator[None]:
    """Collect registrations made in this context."""
    token = self._active.set(self)
    try:
        yield
    finally:
        self._active.reset(token)

capture(init_fn) classmethod

Collect warmup registrations made while the decorated callable runs.

Source code in vllm/model_executor/warmup/jit_warmup.py
@classmethod
def capture(cls, init_fn: Callable[..., None]) -> Callable[..., None]:
    """Collect warmup registrations made while the decorated callable runs."""

    @wraps(init_fn)
    def wrapped(instance: Any, vllm_config: Any, *args: Any, **kwargs: Any) -> None:
        registry = cls(vllm_config)
        instance.jit_warmup_registry = registry
        with registry.activate():
            init_fn(instance, vllm_config, *args, **kwargs)

    return wrapped

register(kernel, *args, **kwargs) classmethod

Register a kernel with the active registry, if one exists.

Source code in vllm/model_executor/warmup/jit_warmup.py
@classmethod
def register(
    cls,
    kernel: VllmJitKernel[Any],
    *args: Any,
    **kwargs: Any,
) -> None:
    """Register a kernel with the active registry, if one exists."""
    registry = cls._active.get()
    if registry is not None:
        registry._add(kernel, args, kwargs)

warmup()

Expand registrations and compile each wrapper/key pair once.

Source code in vllm/model_executor/warmup/jit_warmup.py
def warmup(self) -> None:
    """Expand registrations and compile each wrapper/key pair once."""
    from tqdm import tqdm

    from vllm.distributed import is_global_first_rank

    kernel_items: list[tuple[VllmJitKernel[Any], dict[Any, None]]] = []
    for kernel, registrations in self._registrations.items():
        compile_keys: dict[Any, None] = {}
        for args, kwargs in registrations:
            if (
                not args
                and not kwargs
                and "vllm_config"
                in inspect.signature(kernel.get_warmup_keys).parameters
            ):
                kwargs = {"vllm_config": self.vllm_config}
            for compile_key in kernel.get_warmup_keys(*args, **kwargs):
                compile_keys[compile_key] = None
        if compile_keys:
            kernel_items.append((kernel, compile_keys))

    if not kernel_items:
        return

    total_keys = sum(len(compile_keys) for _, compile_keys in kernel_items)
    with tqdm(
        kernel_items,
        desc=f"JIT kernel warmup ({total_keys} compile keys)",
        disable=not is_global_first_rank(),
        dynamic_ncols=True,
        unit="kernel",
    ) as progress:
        for kernel, compile_keys in progress:
            progress.set_postfix_str(
                f"{kernel.__class__.__name__} ({len(compile_keys)} keys)",
                refresh=False,
            )
            kernel.compile_many(compile_keys)

VllmJitKernel

Bases: Generic[CompileKeyT], ABC

Kernel wrapper that owns dispatch, warmup keys, and compilation.

Methods:

  • compile –

    Compile one warmup key.

  • compile_many –

    Compile a batch of warmup keys, allowing backend-specific scheduling.

  • dispatch –

    Build one compile key from one concrete dispatch point.

  • get_warmup_keys –

    Return compile keys that should be warmed for this kernel.

  • register_warmup –

    Register this kernel with the active runner's warmup registry.

  • warmup –

    Compile this kernel's warmup keys.

Source code in vllm/model_executor/warmup/jit_warmup.py
class VllmJitKernel(Generic[CompileKeyT], ABC):
    """Kernel wrapper that owns dispatch, warmup keys, and compilation."""

    CompileKey: type[CompileKeyT]
    bind_launch_inputs = True

    def __init__(self) -> None:
        dispatch = type(self).dispatch
        self._dispatch_trace = (
            None
            if dispatch is VllmJitKernel.dispatch
            else _trace_compile_key_dispatch(self.dispatch)
        )
        self._compiled_cache: dict[Any, Any] = {}

    def compile_key(self, kwargs: Mapping[str, Any]) -> CompileKeyT:
        if self._dispatch_trace is None:
            raise TypeError(f"{type(self).__name__} does not define dispatch()")
        return self._dispatch_trace.compile_key(self.CompileKey, kwargs)

    def _get_or_compile(
        self,
        compile_key: CompileKeyT,
        *,
        runtime_context: Mapping[str, Any] | None = None,
    ) -> Any:
        """Return a cached executor, compiling it on a monitored cache miss."""
        if compile_key not in self._compiled_cache:
            self.compile(compile_key)

        try:
            return self._compiled_cache[compile_key]
        except KeyError as exc:
            details = [f"compile_key={compile_key!r}"]
            if runtime_context:
                details.append(f"runtime_context={dict(runtime_context)!r}")
            raise RuntimeError(
                f"{type(self).__name__}.compile(...) did not cache its JIT "
                f"executor ({', '.join(details)})"
            ) from exc

    def _trace_dispatch(
        self, dispatch: CompileKeyDispatchFn[CompileKeyT]
    ) -> Callable[..., list[CompileKeyT]]:
        compile_key_dispatch_trace = _trace_compile_key_dispatch(dispatch)

        def traced(
            *input_groups: _WarmupInputRows,
            _when: WarmupPredicateFn | None = None,
            **kwargs: WarmupValues,
        ) -> list[CompileKeyT]:
            for group in input_groups:
                if not isinstance(group, _WarmupInputRows):
                    raise TypeError(
                        "_trace_dispatch positional arguments must be "
                        "zip_inputs(...) groups"
                    )
            predicate_trace = (
                _trace_warmup_predicate(_when) if _when is not None else None
            )
            predicate_only_names: frozenset[str] = frozenset()
            if predicate_trace is not None:
                compile_key_fields = frozenset(
                    field.name for field in fields(cast(Any, self.CompileKey))
                )
                predicate_only_names = (
                    predicate_trace.input_names
                    - compile_key_dispatch_trace.input_names
                    - compile_key_fields
                )
            # Unmatched **kwargs fields also belong to the expansion space.
            available_names = set(kwargs).union(
                *(group.rows[0] for group in input_groups)
            )
            input_names = compile_key_dispatch_trace.input_names_for(available_names)
            if predicate_trace is not None:
                input_names = input_names | predicate_trace.input_names
            expanded_input_groups = tuple(
                _expand_warmup_input_rows(group.rows, input_names)
                for group in input_groups
            )
            # Expand independent keyword inputs into cartesian-product axes.
            expanded_kwarg_axes = tuple(
                (name, _expand_warmup_values(value))
                for name, value in kwargs.items()
                if name in input_names
            )
            dispatch_value_axes = (
                *expanded_input_groups,
                *(values for _, values in expanded_kwarg_axes),
            )
            input_group_count = len(expanded_input_groups)
            kwarg_names = tuple(name for name, _ in expanded_kwarg_axes)
            compile_keys: dict[CompileKeyT, None] = {}
            for dispatch_value_set in itertools.product(*dispatch_value_axes):
                dispatch_values = _merge_warmup_kwargs(
                    (
                        *dispatch_value_set[:input_group_count],
                        dict(
                            zip(
                                kwarg_names,
                                dispatch_value_set[input_group_count:],
                            )
                        ),
                    )
                )
                if predicate_trace is not None and not predicate_trace.matches(
                    dispatch_values
                ):
                    continue
                compile_key = compile_key_dispatch_trace.compile_key(
                    self.CompileKey,
                    {
                        name: value
                        for name, value in dispatch_values.items()
                        if name not in predicate_only_names
                    },
                )
                compile_keys[compile_key] = None
            return list(compile_keys)

        return traced

    def _expand_warmup_cases(
        self,
        cases_fn: Callable[..., Any],
        *args: Any,
        _value_expander: Callable[[WarmupValues], tuple[Any, ...]] = (
            _expand_warmup_values
        ),
        **kwargs: Any,
    ) -> Iterator[Mapping[str, Any]]:
        """Expand symbolic domains declared inside a warmup-cases method."""
        function_def = get_function_source_node(cases_fn)
        if isinstance(function_def, ast.Lambda):
            raise _dispatch_expr_error(
                function_def, "Warmup cases must be a function definition"
            )
        signature = inspect.signature(cases_fn)
        bound = signature.bind(*args, **kwargs)
        bound.apply_defaults()
        static_values = dict(bound.arguments)
        bound_self = getattr(cases_fn, "__self__", None)
        if bound_self is not None:
            static_values["self"] = bound_self
        globals_ = getattr(cases_fn, "__func__", cases_fn).__globals__

        domains: list[tuple[str, tuple[Any, ...]]] = []
        local_exprs: list[tuple[str, ast.AST]] = []
        predicates: list[ast.AST] = []
        return_expr: ast.AST | None = None
        for statement in function_def.body:
            if (
                isinstance(statement, ast.Expr)
                and isinstance(statement.value, ast.Constant)
                and isinstance(statement.value.value, str)
            ):
                continue
            assignment = _named_assignment(statement, "Warmup case")
            if assignment is not None:
                name, value_expr = assignment
                call_name = (
                    (get_ast_full_name(value_expr.func) or "").split(".")[-1]
                    if isinstance(value_expr, ast.Call)
                    else ""
                )
                if call_name in {"WarmupIntRange", "WarmupChoices"}:
                    value = _eval_dispatch_expr(value_expr, static_values, globals_)
                    domains.append((name, _value_expander(value)))
                else:
                    local_exprs.append((name, value_expr))
                continue
            if (
                isinstance(statement, ast.Expr)
                and isinstance(statement.value, ast.Call)
                and (get_ast_full_name(statement.value.func) or "").endswith("_when")
            ):
                if len(statement.value.args) != 1 or statement.value.keywords:
                    raise _dispatch_expr_error(
                        statement, "_when requires exactly one positional expression"
                    )
                predicates.append(statement.value.args[0])
                continue
            if isinstance(statement, ast.Return):
                if return_expr is not None or statement.value is None:
                    raise _dispatch_expr_error(
                        statement, "Warmup cases require one final return expression"
                    )
                return_expr = statement.value
                continue
            raise _dispatch_expr_error(
                statement,
                "AST-traced warmup cases support assignments, _when, and return",
            )

        if return_expr is None or not isinstance(return_expr, ast.Call):
            raise ValueError("AST-traced warmup cases must return one case(...) call")
        if any(keyword.arg is None for keyword in return_expr.keywords):
            raise _dispatch_expr_error(
                return_expr, "Warmup case expressions do not support **kwargs"
            )

        class _InlineDomainRewriter(ast.NodeTransformer):
            def visit_Call(self, node: ast.Call) -> ast.AST:
                call_name = (get_ast_full_name(node.func) or "").split(".")[-1]
                if call_name not in {"WarmupIntRange", "WarmupChoices"}:
                    return self.generic_visit(node)
                name = f"__warmup_domain_{len(domains)}"
                domain = _eval_dispatch_expr(node, static_values, globals_)
                domains.append((name, _value_expander(domain)))
                return ast.copy_location(ast.Name(id=name, ctx=ast.Load()), node)

        return_expr = cast(ast.Call, _InlineDomainRewriter().visit(return_expr))

        domain_names = tuple(name for name, _ in domains)
        dynamic_names = set(domain_names)
        dynamic_local_exprs: list[tuple[str, ast.AST]] = []
        for name, expr in local_exprs:
            referenced_names = {
                node.id for node in ast.walk(expr) if isinstance(node, ast.Name)
            }
            if referenced_names & dynamic_names:
                dynamic_local_exprs.append((name, expr))
                dynamic_names.add(name)
            else:
                static_values[name] = _eval_dispatch_expr(expr, static_values, globals_)

        for _, expr in dynamic_local_exprs:
            _validate_dispatch_expr(expr)
        for predicate in predicates:
            _validate_dispatch_expr(predicate)
        _validate_dispatch_expr(return_expr)

        case_body: list[ast.stmt] = [
            ast.Assign(
                targets=[ast.Name(id=name, ctx=ast.Store())],
                value=cast(ast.expr, expr),
            )
            for name, expr in dynamic_local_exprs
        ]
        if predicates:
            predicate = (
                predicates[0]
                if len(predicates) == 1
                else ast.BoolOp(
                    op=ast.And(),
                    values=[cast(ast.expr, value) for value in predicates],
                )
            )
            case_body.append(
                ast.If(
                    test=ast.UnaryOp(op=ast.Not(), operand=cast(ast.expr, predicate)),
                    body=[ast.Return(value=ast.Constant(value=None))],
                    orelse=[],
                )
            )
        case_body.append(ast.Return(value=cast(ast.expr, return_expr)))
        # FunctionDef fields vary across the supported Python versions, so no
        # single constructor overload matches every mypy target.
        case_function = ast.FunctionDef(  # type: ignore[call-overload]
            name="__vllm_warmup_case",
            args=ast.arguments(
                posonlyargs=[],
                args=[ast.arg(arg=name) for name in domain_names],
                kwonlyargs=[],
                kw_defaults=[],
                defaults=[],
            ),
            body=case_body,
            decorator_list=[],
            returns=None,
            type_comment=None,
        )
        case_globals = dict(globals_)
        case_globals.update(static_values)
        exec(
            compile(
                ast.fix_missing_locations(
                    ast.Module(body=[case_function], type_ignores=[])
                ),
                "<jit-warmup>",
                "exec",
            ),
            case_globals,
        )
        evaluate_case = cast(Callable[..., Any], case_globals[case_function.name])

        domain_values = tuple(values for _, values in domains)
        for values in itertools.product(*domain_values):
            case = evaluate_case(*values)
            if case is None:
                continue
            if not isinstance(case, Mapping):
                raise TypeError("AST-traced warmup cases must return a mapping")
            if any(
                isinstance(value, WarmupChoices | WarmupIntRange)
                for value in case.values()
            ):
                raise TypeError("Warmup domains must be direct expressions")
            yield case

    def dispatch(self, **kwargs: Any) -> CompileKeyT:
        """Build one compile key from one concrete dispatch point."""
        raise NotImplementedError

    @abstractmethod
    def get_warmup_keys(self, *args: Any, **kwargs: Any) -> list[CompileKeyT]:
        """Return compile keys that should be warmed for this kernel."""
        raise NotImplementedError

    @abstractmethod
    def compile(self, compile_key: CompileKeyT) -> None:
        """Compile one warmup key."""
        raise NotImplementedError

    def compile_many(self, compile_keys: Iterable[CompileKeyT]) -> None:
        """Compile a batch of warmup keys, allowing backend-specific scheduling."""
        for compile_key in compile_keys:
            self.compile(compile_key)

    def register_warmup(self, *args: Any, **kwargs: Any) -> None:
        """Register this kernel with the active runner's warmup registry."""
        JitWarmupRegistry.register(self, *args, **kwargs)

    def warmup(self, *args: Any, **kwargs: Any) -> None:
        """Compile this kernel's warmup keys."""
        self.compile_many(self.get_warmup_keys(*args, **kwargs))

_expand_warmup_cases(cases_fn, *args, _value_expander=_expand_warmup_values, **kwargs)

Expand symbolic domains declared inside a warmup-cases method.

Source code in vllm/model_executor/warmup/jit_warmup.py
def _expand_warmup_cases(
    self,
    cases_fn: Callable[..., Any],
    *args: Any,
    _value_expander: Callable[[WarmupValues], tuple[Any, ...]] = (
        _expand_warmup_values
    ),
    **kwargs: Any,
) -> Iterator[Mapping[str, Any]]:
    """Expand symbolic domains declared inside a warmup-cases method."""
    function_def = get_function_source_node(cases_fn)
    if isinstance(function_def, ast.Lambda):
        raise _dispatch_expr_error(
            function_def, "Warmup cases must be a function definition"
        )
    signature = inspect.signature(cases_fn)
    bound = signature.bind(*args, **kwargs)
    bound.apply_defaults()
    static_values = dict(bound.arguments)
    bound_self = getattr(cases_fn, "__self__", None)
    if bound_self is not None:
        static_values["self"] = bound_self
    globals_ = getattr(cases_fn, "__func__", cases_fn).__globals__

    domains: list[tuple[str, tuple[Any, ...]]] = []
    local_exprs: list[tuple[str, ast.AST]] = []
    predicates: list[ast.AST] = []
    return_expr: ast.AST | None = None
    for statement in function_def.body:
        if (
            isinstance(statement, ast.Expr)
            and isinstance(statement.value, ast.Constant)
            and isinstance(statement.value.value, str)
        ):
            continue
        assignment = _named_assignment(statement, "Warmup case")
        if assignment is not None:
            name, value_expr = assignment
            call_name = (
                (get_ast_full_name(value_expr.func) or "").split(".")[-1]
                if isinstance(value_expr, ast.Call)
                else ""
            )
            if call_name in {"WarmupIntRange", "WarmupChoices"}:
                value = _eval_dispatch_expr(value_expr, static_values, globals_)
                domains.append((name, _value_expander(value)))
            else:
                local_exprs.append((name, value_expr))
            continue
        if (
            isinstance(statement, ast.Expr)
            and isinstance(statement.value, ast.Call)
            and (get_ast_full_name(statement.value.func) or "").endswith("_when")
        ):
            if len(statement.value.args) != 1 or statement.value.keywords:
                raise _dispatch_expr_error(
                    statement, "_when requires exactly one positional expression"
                )
            predicates.append(statement.value.args[0])
            continue
        if isinstance(statement, ast.Return):
            if return_expr is not None or statement.value is None:
                raise _dispatch_expr_error(
                    statement, "Warmup cases require one final return expression"
                )
            return_expr = statement.value
            continue
        raise _dispatch_expr_error(
            statement,
            "AST-traced warmup cases support assignments, _when, and return",
        )

    if return_expr is None or not isinstance(return_expr, ast.Call):
        raise ValueError("AST-traced warmup cases must return one case(...) call")
    if any(keyword.arg is None for keyword in return_expr.keywords):
        raise _dispatch_expr_error(
            return_expr, "Warmup case expressions do not support **kwargs"
        )

    class _InlineDomainRewriter(ast.NodeTransformer):
        def visit_Call(self, node: ast.Call) -> ast.AST:
            call_name = (get_ast_full_name(node.func) or "").split(".")[-1]
            if call_name not in {"WarmupIntRange", "WarmupChoices"}:
                return self.generic_visit(node)
            name = f"__warmup_domain_{len(domains)}"
            domain = _eval_dispatch_expr(node, static_values, globals_)
            domains.append((name, _value_expander(domain)))
            return ast.copy_location(ast.Name(id=name, ctx=ast.Load()), node)

    return_expr = cast(ast.Call, _InlineDomainRewriter().visit(return_expr))

    domain_names = tuple(name for name, _ in domains)
    dynamic_names = set(domain_names)
    dynamic_local_exprs: list[tuple[str, ast.AST]] = []
    for name, expr in local_exprs:
        referenced_names = {
            node.id for node in ast.walk(expr) if isinstance(node, ast.Name)
        }
        if referenced_names & dynamic_names:
            dynamic_local_exprs.append((name, expr))
            dynamic_names.add(name)
        else:
            static_values[name] = _eval_dispatch_expr(expr, static_values, globals_)

    for _, expr in dynamic_local_exprs:
        _validate_dispatch_expr(expr)
    for predicate in predicates:
        _validate_dispatch_expr(predicate)
    _validate_dispatch_expr(return_expr)

    case_body: list[ast.stmt] = [
        ast.Assign(
            targets=[ast.Name(id=name, ctx=ast.Store())],
            value=cast(ast.expr, expr),
        )
        for name, expr in dynamic_local_exprs
    ]
    if predicates:
        predicate = (
            predicates[0]
            if len(predicates) == 1
            else ast.BoolOp(
                op=ast.And(),
                values=[cast(ast.expr, value) for value in predicates],
            )
        )
        case_body.append(
            ast.If(
                test=ast.UnaryOp(op=ast.Not(), operand=cast(ast.expr, predicate)),
                body=[ast.Return(value=ast.Constant(value=None))],
                orelse=[],
            )
        )
    case_body.append(ast.Return(value=cast(ast.expr, return_expr)))
    # FunctionDef fields vary across the supported Python versions, so no
    # single constructor overload matches every mypy target.
    case_function = ast.FunctionDef(  # type: ignore[call-overload]
        name="__vllm_warmup_case",
        args=ast.arguments(
            posonlyargs=[],
            args=[ast.arg(arg=name) for name in domain_names],
            kwonlyargs=[],
            kw_defaults=[],
            defaults=[],
        ),
        body=case_body,
        decorator_list=[],
        returns=None,
        type_comment=None,
    )
    case_globals = dict(globals_)
    case_globals.update(static_values)
    exec(
        compile(
            ast.fix_missing_locations(
                ast.Module(body=[case_function], type_ignores=[])
            ),
            "<jit-warmup>",
            "exec",
        ),
        case_globals,
    )
    evaluate_case = cast(Callable[..., Any], case_globals[case_function.name])

    domain_values = tuple(values for _, values in domains)
    for values in itertools.product(*domain_values):
        case = evaluate_case(*values)
        if case is None:
            continue
        if not isinstance(case, Mapping):
            raise TypeError("AST-traced warmup cases must return a mapping")
        if any(
            isinstance(value, WarmupChoices | WarmupIntRange)
            for value in case.values()
        ):
            raise TypeError("Warmup domains must be direct expressions")
        yield case

_get_or_compile(compile_key, *, runtime_context=None)

Return a cached executor, compiling it on a monitored cache miss.

Source code in vllm/model_executor/warmup/jit_warmup.py
def _get_or_compile(
    self,
    compile_key: CompileKeyT,
    *,
    runtime_context: Mapping[str, Any] | None = None,
) -> Any:
    """Return a cached executor, compiling it on a monitored cache miss."""
    if compile_key not in self._compiled_cache:
        self.compile(compile_key)

    try:
        return self._compiled_cache[compile_key]
    except KeyError as exc:
        details = [f"compile_key={compile_key!r}"]
        if runtime_context:
            details.append(f"runtime_context={dict(runtime_context)!r}")
        raise RuntimeError(
            f"{type(self).__name__}.compile(...) did not cache its JIT "
            f"executor ({', '.join(details)})"
        ) from exc

compile(compile_key) abstractmethod

Compile one warmup key.

Source code in vllm/model_executor/warmup/jit_warmup.py
@abstractmethod
def compile(self, compile_key: CompileKeyT) -> None:
    """Compile one warmup key."""
    raise NotImplementedError

compile_many(compile_keys)

Compile a batch of warmup keys, allowing backend-specific scheduling.

Source code in vllm/model_executor/warmup/jit_warmup.py
def compile_many(self, compile_keys: Iterable[CompileKeyT]) -> None:
    """Compile a batch of warmup keys, allowing backend-specific scheduling."""
    for compile_key in compile_keys:
        self.compile(compile_key)

dispatch(**kwargs)

Build one compile key from one concrete dispatch point.

Source code in vllm/model_executor/warmup/jit_warmup.py
def dispatch(self, **kwargs: Any) -> CompileKeyT:
    """Build one compile key from one concrete dispatch point."""
    raise NotImplementedError

get_warmup_keys(*args, **kwargs) abstractmethod

Return compile keys that should be warmed for this kernel.

Source code in vllm/model_executor/warmup/jit_warmup.py
@abstractmethod
def get_warmup_keys(self, *args: Any, **kwargs: Any) -> list[CompileKeyT]:
    """Return compile keys that should be warmed for this kernel."""
    raise NotImplementedError

register_warmup(*args, **kwargs)

Register this kernel with the active runner's warmup registry.

Source code in vllm/model_executor/warmup/jit_warmup.py
def register_warmup(self, *args: Any, **kwargs: Any) -> None:
    """Register this kernel with the active runner's warmup registry."""
    JitWarmupRegistry.register(self, *args, **kwargs)

warmup(*args, **kwargs)

Compile this kernel's warmup keys.

Source code in vllm/model_executor/warmup/jit_warmup.py
def warmup(self, *args: Any, **kwargs: Any) -> None:
    """Compile this kernel's warmup keys."""
    self.compile_many(self.get_warmup_keys(*args, **kwargs))

WarmupChoices dataclass

Expand an explicit finite set of values in a traced warmup method.

Source code in vllm/model_executor/warmup/jit_warmup.py
@dataclass(frozen=True)
class WarmupChoices:
    """Expand an explicit finite set of values in a traced warmup method."""

    values: tuple[Any, ...]

    def __init__(self, *values: Any) -> None:
        object.__setattr__(self, "values", values)

WarmupIntRange dataclass

Expand integers with range semantics or a custom monotonic progression.

Source code in vllm/model_executor/warmup/jit_warmup.py
@dataclass(frozen=True)
class WarmupIntRange:
    """Expand integers with range semantics or a custom monotonic progression."""

    start: int
    stop: int
    step: int = 1
    advance: Callable[[int], int] | None = None

_WarmupInputRows dataclass

Warmup dispatch inputs expanded in lockstep.

Source code in vllm/model_executor/warmup/jit_warmup.py
@dataclass(frozen=True)
class _WarmupInputRows:
    """Warmup dispatch inputs expanded in lockstep."""

    rows: tuple[Mapping[str, WarmupValues], ...]

_same_registration(left, right)

True if two (args, kwargs) registrations are equivalent.

Source code in vllm/model_executor/warmup/jit_warmup.py
def _same_registration(
    left: tuple[tuple[Any, ...], dict[str, Any]],
    right: tuple[tuple[Any, ...], dict[str, Any]],
) -> bool:
    """True if two ``(args, kwargs)`` registrations are equivalent."""
    left_args, left_kwargs = left
    right_args, right_kwargs = right
    if len(left_args) != len(right_args) or left_kwargs.keys() != right_kwargs.keys():
        return False
    return all(
        _same_value(x, y) for x, y in zip(left_args, right_args, strict=True)
    ) and all(_same_value(left_kwargs[k], right_kwargs[k]) for k in left_kwargs)

_same_value(left, right)

True if two registration values are the same object or compare equal.

Identity is checked first so shared singletons (e.g. vllm_config, torch dtypes) short-circuit before any potentially deep or non-boolean __eq__.

Source code in vllm/model_executor/warmup/jit_warmup.py
def _same_value(left: Any, right: Any) -> bool:
    """True if two registration values are the same object or compare equal.

    Identity is checked first so shared singletons (e.g. ``vllm_config``, torch
    dtypes) short-circuit before any potentially deep or non-boolean ``__eq__``.
    """
    if left is right:
        return True
    try:
        return bool(left == right)
    except (TypeError, ValueError, RuntimeError):
        return False

_validate_dispatch_expr(node)

Enforce the dispatch AST subset before compiling traced code.

Source code in vllm/model_executor/warmup/jit_warmup.py
def _validate_dispatch_expr(node: ast.AST) -> None:
    """Enforce the dispatch AST subset before compiling traced code."""
    allowed_nodes = (
        ast.Name,
        ast.Constant,
        ast.Lambda,
        ast.arguments,
        ast.arg,
        ast.IfExp,
        ast.Tuple,
        ast.List,
        ast.BoolOp,
        ast.And,
        ast.Or,
        ast.Compare,
        *tuple(_CMP_OPS),
        ast.UnaryOp,
        ast.Not,
        ast.USub,
        ast.BinOp,
        *tuple(_BIN_OPS),
        ast.Call,
        ast.keyword,
        ast.Starred,
        ast.Attribute,
        ast.Subscript,
        ast.Load,
    )
    for child in ast.walk(node):
        if not isinstance(child, allowed_nodes):
            raise _dispatch_expr_error(child, "Unsupported dispatch expression")
        if isinstance(child, ast.Call) and any(
            keyword.arg is None for keyword in child.keywords
        ):
            raise _dispatch_expr_error(
                child, "Dispatch helper calls cannot use **kwargs"
            )

_when(condition)

Filter symbolic cases in an AST-traced warmup-cases method.

Source code in vllm/model_executor/warmup/jit_warmup.py
def _when(condition: bool) -> None:
    """Filter symbolic cases in an AST-traced warmup-cases method."""
    raise RuntimeError("_when() is only valid in an AST-traced warmup method")

kernel_launcher(call_fn)

Delegate a declarative launch specification to its backend owner.

Source code in vllm/model_executor/warmup/jit_warmup.py
def kernel_launcher(call_fn: Callable[..., Any]) -> Callable[..., Any]:
    """Delegate a declarative launch specification to its backend owner."""
    signature = inspect.signature(call_fn)

    @wraps(call_fn)
    def wrapper(self: Any, *args: Any, **kwargs: Any) -> Any:
        launch_spec = call_fn(self, *args, **kwargs)
        if not self.bind_launch_inputs:
            return self.launch(launch_spec, {})

        bound = signature.bind(self, *args, **kwargs)
        bound.apply_defaults()
        inputs = {
            name: value for name, value in bound.arguments.items() if name != "self"
        }
        return self.launch(launch_spec, inputs)

    return wrapper

zip_inputs(*rows)

Group row-wise dispatch inputs that should be expanded in lockstep.

Source code in vllm/model_executor/warmup/jit_warmup.py
def zip_inputs(*rows: Mapping[str, WarmupValues]) -> _WarmupInputRows:
    """Group row-wise dispatch inputs that should be expanded in lockstep."""
    if not rows:
        raise ValueError("zip_inputs requires at least one dispatch input row")
    if not all(isinstance(row, Mapping) for row in rows):
        raise ValueError("zip_inputs rows must be mappings")

    first_names = frozenset(rows[0])
    if not first_names:
        raise ValueError("zip_inputs rows require at least one dispatch input name")
    if not all(isinstance(name, str) for name in first_names):
        raise ValueError("zip_inputs dispatch input names must be strings")

    input_rows: list[Mapping[str, WarmupValues]] = []
    for row in rows:
        names = frozenset(row)
        if names != first_names:
            raise ValueError("zip_inputs rows must use the same dispatch input names")
        input_rows.append(dict(row))

    return _WarmupInputRows(rows=tuple(input_rows))