Cache FlashInfer autotune configs (#24156)
This commit is contained in:
@@ -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:
|
||||||
|
|||||||
Reference in New Issue
Block a user