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.
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.
"""
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)
if pp_size > 1:
ranks = "-".join(str(rank) for rank in tune_group.ranks)
cache_path = cache_path.with_name(
f"{cache_path.stem}_tp_{ranks}{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.
# Read cached autotune results and broadcast within the tuning group.
cached_results: bytes | None = None
if is_leader and cache_path.exists():
with open(cache_path, "rb") as f:
cached_results = f.read()
cached_results = tune_group.broadcast_object(cached_results, src=0)
if cached_results is not None:
write_flashinfer_autotune_cache(cache_path, cached_results)
tune_group.barrier()
tuner.load_configs(str(cache_path))
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())
with torch.inference_mode():
_run_flashinfer_bf16_autotune_dummy_run(
runner, skip_ops=skip_ops, skip_attn=hisparse_enabled
)
finally:
set_autotune_process_group(None)
if world.world_size > 1:
world.barrier()
if is_leader:
tuner.save_configs(str(cache_path))
|