Cache FlashInfer autotune configs (#24156)

This commit is contained in:
Sam Shleifer
2026-05-05 02:00:40 +08:00
committed by GitHub
parent c545a5b1c0
commit e6f252e9b8
@@ -18,6 +18,7 @@ from __future__ import annotations
import contextlib import contextlib
import datetime import datetime
import gc import gc
import hashlib
import inspect import inspect
import logging import logging
import os import os
@@ -27,6 +28,7 @@ import time
import uuid import uuid
from collections import defaultdict from collections import defaultdict
from dataclasses import dataclass from dataclasses import dataclass
from pathlib import Path
from typing import Callable, List, Optional, Tuple, Union from typing import Callable, List, Optional, Tuple, Union
import torch import torch
@@ -339,6 +341,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.tp_size = tp_size self.tp_size = tp_size
self.moe_ep_rank = moe_ep_rank self.moe_ep_rank = moe_ep_rank
self.moe_ep_size = moe_ep_size self.moe_ep_size = moe_ep_size
self.dp_rank = dp_rank
self.dp_size = server_args.dp_size if server_args.enable_dp_attention else 1 self.dp_size = server_args.dp_size if server_args.enable_dp_attention else 1
self.pp_rank = pp_rank self.pp_rank = pp_rank
self.pp_size = pp_size self.pp_size = pp_size
@@ -2424,19 +2427,54 @@ class ModelRunner(ModelRunnerKVCacheMixin):
"""Run flashinfer autotune.""" """Run flashinfer autotune."""
from flashinfer.autotuner import autotune from flashinfer.autotuner import autotune
logger.info("Running FlashInfer autotune...") cache_path = self._flashinfer_autotune_cache_path()
logger.info("Running FlashInfer autotune with cache: %s", cache_path)
# Run warmup on the non-default stream to avoid NCCL 2.29+ cudaMemcpyBatchAsync # Run warmup on the non-default stream to avoid NCCL 2.29+ cudaMemcpyBatchAsync
# calls on default stream (unsupported by CUDA) when --enable-symm-mem is used. # calls on default stream (unsupported by CUDA) when --enable-symm-mem is used.
self.forward_stream.wait_stream(torch.cuda.current_stream()) self.forward_stream.wait_stream(torch.cuda.current_stream())
with torch.get_device_module(self.device).stream(self.forward_stream): with torch.get_device_module(self.device).stream(self.forward_stream):
with torch.inference_mode(), autotune(): with torch.inference_mode(), autotune(True, cache=str(cache_path)):
self._dummy_run( self._dummy_run(batch_size=self.req_to_token_pool.size)
batch_size=self.req_to_token_pool.size, run_ctx=autotune()
)
torch.cuda.current_stream().wait_stream(self.forward_stream) torch.cuda.current_stream().wait_stream(self.forward_stream)
logger.info("FlashInfer autotune completed.") logger.info("FlashInfer autotune completed.")
def _flashinfer_autotune_cache_path(self) -> Path:
import flashinfer
major, minor = torch.cuda.get_device_capability(self.device)
arch = f"sm{major}{minor}"
flashinfer_version = getattr(flashinfer, "__version__", "unknown")
server_args = self.server_args
model_key = "|".join(
[
str(server_args.model_path),
str(self.dtype),
str(server_args.quantization),
str(server_args.moe_runner_backend),
str(self.tp_size),
str(self.pp_size),
str(self.dp_size),
str(self.moe_ep_size),
str(self.model_config.hf_config.__class__.__name__),
]
)
cache_key = hashlib.sha256(model_key.encode()).hexdigest()[:16]
cache_dir = (
Path(envs.SGLANG_CACHE_DIR.get())
/ "flashinfer"
/ "autotune"
/ flashinfer_version
/ arch
/ cache_key
)
cache_dir.mkdir(parents=True, exist_ok=True)
return (
cache_dir
/ f"rank_tp{self.tp_rank}_pp{self.pp_rank}_dp{self.dp_rank or 0}.json"
)
def _dummy_run(self, batch_size: int, run_ctx=None): def _dummy_run(self, batch_size: int, run_ctx=None):
"""Run a dummy forward pass for warmup/profiling.""" """Run a dummy forward pass for warmup/profiling."""
if self.is_generation: if self.is_generation: