Skip to content

vllm.model_executor.warmup.kernel_warmup

Warmup kernels used during model execution. This is useful specifically for JIT'ed kernels as we don't want JIT'ing to happen during model execution.

Functions:

_all_ranks_have_matching_cache(path, group)

True iff every rank in group has its own, mutually consistent file.

Source code in vllm/model_executor/warmup/kernel_warmup.py
def _all_ranks_have_matching_cache(path: Path, group) -> bool:
    """True iff every rank in ``group`` has its own, mutually consistent file."""
    fingerprint = _autotune_cache_fingerprint(path)
    if group.world_size == 1:
        return fingerprint is not None
    gathered: list[tuple[str, int] | None] = [None] * group.world_size
    torch.distributed.all_gather_object(gathered, fingerprint, group=group.cpu_group)
    return fingerprint is not None and all(f == fingerprint for f in gathered)

_autotune_cache_fingerprint(path)

Identify a saved autotune file by its FlashInfer metadata and size.

Source code in vllm/model_executor/warmup/kernel_warmup.py
def _autotune_cache_fingerprint(path: Path) -> tuple[str, int] | None:
    """Identify a saved autotune file by its FlashInfer metadata and size."""
    try:
        configs = json.loads(path.read_text())
        metadata = configs.pop("_metadata", None)
    except (OSError, ValueError, AttributeError):
        return None
    return json.dumps(metadata, sort_keys=True), len(configs)

_flashinfer_deferred_moe_token_counts(runner)

Return bounded token counts that exercise deferred MoE finalization.

Source code in vllm/model_executor/warmup/kernel_warmup.py
def _flashinfer_deferred_moe_token_counts(
    runner: "GPUModelRunner",
) -> tuple[int, ...]:
    """Return bounded token counts that exercise deferred MoE finalization."""
    from vllm.model_executor.layers.fused_moe import MoERunner

    max_tokens = runner.scheduler_config.max_num_batched_tokens
    token_counts: list[int] = []
    for module in runner.get_model().modules():
        if not isinstance(module, MoERunner):
            continue

        moe_config = module.moe_config
        max_deferred_tokens = moe_config.defer_moe_finalize_max_num_tokens
        if moe_config.use_deferred_moe_finalize and max_deferred_tokens > 0:
            token_counts.append(min(max_tokens, max_deferred_tokens))

    return tuple(dict.fromkeys(token_counts))

flashinfer_autotune(runner)

Autotune FlashInfer operations. FlashInfer have many implementations for the same operation, autotuning runs benchmarks for each implementation and stores the results. The results are cached transparently and future calls to FlashInfer will use the best implementation. Without autotuning, FlashInfer will rely on heuristics, which may be significantly slower.

With PP > 1, stages run different layers and may profile different ops, so each stage's TP group tunes separately with its own cache file; otherwise the world group tunes together. Per-tactic timings are averaged over the tuning group so all its ranks select the same tactic.

Results are persisted per rank: FlashInfer keys MoE entries by tp/ep rank (MoERunner.get_cache_key_extras), so one rank's file only hits on that rank. A rank with a cache hit skips the per-tactic reduce the others block in, so ranks keep loaded configs only if every rank in the tuning group has a matching file and successfully loads it.

Source code in vllm/model_executor/warmup/kernel_warmup.py
def flashinfer_autotune(runner: "GPUModelRunner") -> None:
    """Autotune FlashInfer operations.
    FlashInfer have many implementations for the same operation,
    autotuning runs benchmarks for each implementation and stores
    the results. The results are cached transparently and
    future calls to FlashInfer will use the best implementation.
    Without autotuning, FlashInfer will rely on heuristics, which may
    be significantly slower.

    With PP > 1, stages run different layers and may profile different ops,
    so each stage's TP group tunes separately with its own cache file;
    otherwise the world group tunes together. Per-tactic timings are
    averaged over the tuning group so all its ranks select the same tactic.

    Results are persisted per rank: FlashInfer keys MoE entries by tp/ep rank
    (``MoERunner.get_cache_key_extras``), so one rank's file only hits on that
    rank. A rank with a cache hit skips the per-tactic reduce the others block
    in, so ranks keep loaded configs only if every rank in the tuning group
    has a matching file and successfully loads it.
    """
    from flashinfer.autotuner import AutoTuner, set_autotune_process_group

    import vllm.utils.flashinfer as fi_utils
    from vllm.distributed.parallel_state import (
        get_pp_group,
        get_tp_group,
        get_world_group,
    )

    world = get_world_group()
    pp_size = get_pp_group().world_size
    tune_group = get_tp_group() if pp_size > 1 else world
    is_leader = tune_group.rank_in_group == 0
    tuner = AutoTuner.get()

    autotune_kwargs: dict = {}
    skip_ops = _flashinfer_autotune_skip_ops(runner)
    if skip_ops:
        logger.info_once(
            "Skipping FlashInfer autotuning for ops %s",
            tuple(sorted(skip_ops)),
        )
        autotune_kwargs["skip_ops"] = skip_ops

    cache_path = resolve_flashinfer_autotune_file(runner)
    # The world group only spans DP ranks when vLLM folds DP into it.
    dp_rank = runner.vllm_config.parallel_config.data_parallel_rank
    cache_path = cache_path.with_name(
        f"{cache_path.stem}_dp{dp_rank}_rank{world.rank_in_group}{cache_path.suffix}"
    )
    if is_leader:
        logger.info_once("Using FlashInfer autotune cache file: %s", cache_path)

    # We skip EPLB here since we don't want to record dummy metrics.
    # Randomize inputs to avoid every token pick the same experts,
    # which lead to some EP ranks receiving no tokens and skipping their
    # MoE kernel entirely, and cause hang due to all-reduce collective
    # during synchronized autotuning.
    if _all_ranks_have_matching_cache(cache_path, tune_group):
        loaded = tuner.load_configs(str(cache_path))
        if tune_group.world_size > 1:
            loaded_by_rank: list[bool | None] = [None] * tune_group.world_size
            torch.distributed.all_gather_object(
                loaded_by_rank, loaded, group=tune_group.cpu_group
            )
            loaded = all(loaded_by_rank)
        if not loaded:
            tuner.clear_cache()

    group = tune_group.cpu_group if tune_group.world_size > 1 else None
    set_autotune_process_group(group)
    try:
        with (
            torch.inference_mode(),
            fi_utils.autotune(tune_mode=True, **autotune_kwargs),
        ):
            hisparse_enabled = (
                runner.vllm_config.attention_config.hisparse_config is not None
            )
            if hisparse_enabled:
                # HiSparse hot-buffer attention is bounded by decode batch
                # size, not the prefill-sized batch used for the full model.
                autotune_hisparse_flashinfer_attention(runner)
            _run_flashinfer_autotune_dummy_runs(runner, skip_attn=hisparse_enabled)
            replayssm_autotune_warmup(runner)
            _autotune_kimi_k3_kda_qkvg(runner.get_model())
    finally:
        set_autotune_process_group(None)

    if world.world_size > 1:
        world.barrier()
    # Skip the rewrite when nothing was tuned this start (every entry came from
    # the file). FlashInfer gates its own autotune(cache=...) save the same way.
    if tuner._dirty:
        tuner.save_configs(str(cache_path))