feat: support flashinfer kernel autotune (#12306)
Co-authored-by: Qiaolin Yu <liin1211@outlook.com> Co-authored-by: Kangyan-Zhou <zky314343421@gmail.com>
This commit is contained in:
@@ -96,7 +96,9 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
if forward_batch.spec_info.topk > 1:
|
||||
retrieve_next_token = forward_batch.spec_info.retrive_next_token
|
||||
retrieve_next_sibling = forward_batch.spec_info.retrive_next_sibling
|
||||
retrieve_parent_token = torch.empty_like(retrieve_next_token)
|
||||
# retrieve_next_token is None during dummy run so skip tensor creation
|
||||
if retrieve_next_token is not None:
|
||||
retrieve_parent_token = torch.empty_like(retrieve_next_token)
|
||||
else:
|
||||
query_start_loc = torch.empty(
|
||||
(bs + 1,), dtype=torch.int32, device=self.device
|
||||
|
||||
@@ -47,6 +47,7 @@ from sglang.srt.utils import (
|
||||
is_triton_kernels_available,
|
||||
log_info_on_rank0,
|
||||
mxfp_supported,
|
||||
next_power_of_2,
|
||||
round_up,
|
||||
set_weight_attrs,
|
||||
)
|
||||
@@ -634,7 +635,7 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
|
||||
)
|
||||
elif self.flashinfer_mxfp4_moe_precision == "default":
|
||||
x_quant, x_scale = mxfp8_quantize(x, False, alignment=self.hidden_size)
|
||||
x_scale = x_scale.view(torch.float8_e4m3fn).reshape(-1)
|
||||
x_scale = x_scale.view(torch.float8_e4m3fn).reshape(*x.shape[:-1], -1)
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
@@ -684,6 +685,7 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
|
||||
None, # tile_tokens_dim
|
||||
1, # routing_method_type, renormalize
|
||||
True, # do finalize
|
||||
tune_max_num_tokens=next_power_of_2(x_quant.shape[0]),
|
||||
output=symm_output,
|
||||
)[0]
|
||||
return StandardCombineInput(hidden_states=trtllm_gen_output)
|
||||
|
||||
@@ -84,9 +84,12 @@ from sglang.srt.layers.attention.attention_registry import (
|
||||
)
|
||||
from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
DpPaddingMode,
|
||||
get_attention_tp_group,
|
||||
get_attention_tp_size,
|
||||
initialize_dp_attention,
|
||||
set_dp_buffer_len,
|
||||
set_is_extend_in_batch,
|
||||
)
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||
from sglang.srt.layers.pooler import EmbeddingPoolerOutput
|
||||
@@ -121,9 +124,18 @@ from sglang.srt.mem_cache.memory_pool import (
|
||||
SWAKVPool,
|
||||
)
|
||||
from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner
|
||||
from sglang.srt.model_executor.cuda_graph_runner import CudaGraphRunner
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||
from sglang.srt.model_executor.cuda_graph_runner import (
|
||||
CudaGraphRunner,
|
||||
set_torch_compile_config,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import (
|
||||
CaptureHiddenMode,
|
||||
ForwardBatch,
|
||||
ForwardMode,
|
||||
PPProxyTensors,
|
||||
)
|
||||
from sglang.srt.model_executor.hook_manager import register_forward_hooks
|
||||
from sglang.srt.model_executor.input_buffers import GraphInputBuffers
|
||||
from sglang.srt.model_executor.npu_graph_runner import NPUGraphRunner
|
||||
from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
|
||||
PiecewiseCudaGraphRunner,
|
||||
@@ -157,6 +169,9 @@ from sglang.srt.utils import (
|
||||
is_npu,
|
||||
log_info_on_rank0,
|
||||
monkey_patch_p2p_access_check,
|
||||
require_attn_tp_gather,
|
||||
require_gathered_buffer,
|
||||
require_mlp_tp_gather,
|
||||
reserve_rope_cache_for_long_sequences,
|
||||
set_cuda_arch,
|
||||
slow_rank_detector,
|
||||
@@ -531,6 +546,7 @@ class ModelRunner:
|
||||
if self.device == "cuda":
|
||||
self.init_cublas()
|
||||
self.init_attention_backend()
|
||||
self.kernel_warmup()
|
||||
self.init_device_graphs()
|
||||
elif self.device in ["npu", "cpu"]:
|
||||
self.init_attention_backend()
|
||||
@@ -2135,6 +2151,291 @@ class ModelRunner:
|
||||
.cuda()
|
||||
)
|
||||
|
||||
def kernel_warmup(self):
|
||||
"""
|
||||
Warmup and tune kernels before cuda graph capture.
|
||||
Currently only doing FlashInfer autotune.
|
||||
"""
|
||||
if self.device != "cuda":
|
||||
return
|
||||
|
||||
if self._should_run_flashinfer_autotune():
|
||||
self._flashinfer_autotune()
|
||||
|
||||
def _should_run_flashinfer_autotune(self) -> bool:
|
||||
"""Check if flashinfer autotune should be run."""
|
||||
if not self.server_args.enable_flashinfer_autotune:
|
||||
return False
|
||||
|
||||
backend_str = self.server_args.attention_backend
|
||||
if backend_str not in ["flashinfer", "trtllm_mla", "trtllm_mha"]:
|
||||
return False
|
||||
|
||||
major, _ = torch.cuda.get_device_capability()
|
||||
if major < 9:
|
||||
return False
|
||||
|
||||
if (
|
||||
self.spec_algorithm.is_eagle()
|
||||
or self.spec_algorithm.is_standalone()
|
||||
or self.spec_algorithm.is_ngram()
|
||||
):
|
||||
return not self.is_draft_worker
|
||||
|
||||
return True
|
||||
|
||||
def _flashinfer_autotune(self):
|
||||
"""Run flashinfer autotune."""
|
||||
from flashinfer.autotuner import autotune
|
||||
|
||||
logger.info("Running FlashInfer autotune...")
|
||||
|
||||
with torch.inference_mode(), autotune():
|
||||
self._dummy_run(batch_size=self.req_to_token_pool.size)
|
||||
|
||||
logger.info("FlashInfer autotune completed.")
|
||||
|
||||
def _dummy_run(self, batch_size: int):
|
||||
"""Run a dummy forward pass for warmup/profiling."""
|
||||
if self.is_generation:
|
||||
capture_forward_mode = ForwardMode.DECODE
|
||||
else:
|
||||
capture_forward_mode = ForwardMode.EXTEND
|
||||
capture_hidden_mode = CaptureHiddenMode.NULL
|
||||
num_tokens_per_bs = 1
|
||||
if (
|
||||
self.spec_algorithm.is_eagle()
|
||||
or self.spec_algorithm.is_standalone()
|
||||
or self.spec_algorithm.is_ngram()
|
||||
):
|
||||
if self.is_draft_worker:
|
||||
raise RuntimeError("This should not happen")
|
||||
else:
|
||||
capture_forward_mode = ForwardMode.TARGET_VERIFY
|
||||
num_tokens_per_bs = self.server_args.speculative_num_draft_tokens
|
||||
|
||||
if self.server_args.enable_return_hidden_states:
|
||||
capture_hidden_mode = CaptureHiddenMode.FULL
|
||||
|
||||
num_tokens = batch_size * num_tokens_per_bs
|
||||
|
||||
seq_len_fill_value = self.attn_backend.get_cuda_graph_seq_len_fill_value()
|
||||
|
||||
if self.server_args.enable_torch_compile:
|
||||
set_torch_compile_config()
|
||||
|
||||
if self.spec_algorithm.is_eagle3():
|
||||
self.model.set_eagle3_layers_to_capture()
|
||||
|
||||
require_mlp_tp_gather_ = require_mlp_tp_gather(self.server_args)
|
||||
if require_gathered_buffer(self.server_args):
|
||||
assert require_mlp_tp_gather_ or require_attn_tp_gather(self.server_args)
|
||||
|
||||
buffers: GraphInputBuffers = GraphInputBuffers.create(
|
||||
device=self.device,
|
||||
max_bs=batch_size,
|
||||
max_num_token=num_tokens,
|
||||
hidden_size=self.model_config.hidden_size,
|
||||
vocab_size=self.model_config.vocab_size,
|
||||
dtype=self.model_config.dtype,
|
||||
dp_size=self.server_args.dp_size,
|
||||
pp_size=self.server_args.pp_size,
|
||||
is_encoder_decoder=self.model_config.is_encoder_decoder,
|
||||
require_mlp_tp_gather=require_mlp_tp_gather_,
|
||||
seq_len_fill_value=seq_len_fill_value,
|
||||
encoder_len_fill_value=0,
|
||||
num_tokens_per_bs=num_tokens_per_bs,
|
||||
cache_loc_dtype=torch.int64,
|
||||
)
|
||||
buffers.num_token_non_padded[...] = num_tokens
|
||||
|
||||
# For extend mode
|
||||
if not self.is_generation:
|
||||
extend_prefix_lens_cpu = [0] * batch_size
|
||||
extend_seq_lens_cpu = [seq_len_fill_value] * batch_size
|
||||
extend_num_tokens = num_tokens
|
||||
extend_seq_lens = torch.full(
|
||||
(batch_size,), seq_len_fill_value, dtype=torch.int32, device=self.device
|
||||
)
|
||||
extend_prefix_lens = torch.zeros(
|
||||
(batch_size,), dtype=torch.int32, device=self.device
|
||||
)
|
||||
extend_start_loc = torch.arange(
|
||||
0, num_tokens, num_tokens_per_bs, dtype=torch.int32, device=self.device
|
||||
)
|
||||
else:
|
||||
extend_prefix_lens_cpu = None
|
||||
extend_seq_lens_cpu = None
|
||||
extend_num_tokens = None
|
||||
extend_seq_lens = None
|
||||
extend_prefix_lens = None
|
||||
extend_start_loc = None
|
||||
|
||||
if self.server_args.pp_size > 1:
|
||||
pp_proxy_tensors = PPProxyTensors(
|
||||
{k: v[:num_tokens] for k, v in buffers.pp_proxy_tensors.items()}
|
||||
)
|
||||
|
||||
if require_mlp_tp_gather_:
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens] * self.server_args.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=self.device,
|
||||
)
|
||||
)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens] * self.server_args.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=self.device,
|
||||
)
|
||||
)
|
||||
global_dp_buffer_len = num_tokens * self.server_args.dp_size
|
||||
elif require_attn_tp_gather(self.server_args):
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens],
|
||||
dtype=torch.int32,
|
||||
device=self.device,
|
||||
)
|
||||
)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens],
|
||||
dtype=torch.int32,
|
||||
device=self.device,
|
||||
)
|
||||
)
|
||||
global_dp_buffer_len = num_tokens
|
||||
else:
|
||||
global_dp_buffer_len = None
|
||||
|
||||
def get_spec_info():
|
||||
spec_info = None
|
||||
if self.spec_algorithm.is_eagle() or self.spec_algorithm.is_standalone():
|
||||
from sglang.srt.speculative.eagle_info import EagleVerifyInput
|
||||
|
||||
if self.is_draft_worker:
|
||||
raise RuntimeError("This should not happen.")
|
||||
else:
|
||||
spec_info = EagleVerifyInput(
|
||||
draft_token=None,
|
||||
custom_mask=buffers.custom_mask,
|
||||
positions=None,
|
||||
retrive_index=None,
|
||||
retrive_next_token=None,
|
||||
retrive_next_sibling=None,
|
||||
retrive_cum_len=None,
|
||||
spec_steps=self.server_args.speculative_num_steps,
|
||||
topk=self.server_args.speculative_eagle_topk,
|
||||
draft_token_num=self.server_args.speculative_num_draft_tokens,
|
||||
capture_hidden_mode=CaptureHiddenMode.FULL,
|
||||
seq_lens_sum=None,
|
||||
seq_lens_cpu=None,
|
||||
)
|
||||
|
||||
elif self.spec_algorithm.is_ngram():
|
||||
from sglang.srt.speculative.ngram_info import NgramVerifyInput
|
||||
|
||||
spec_info = NgramVerifyInput(
|
||||
draft_token=None,
|
||||
tree_mask=buffers.custom_mask,
|
||||
positions=None,
|
||||
retrive_index=None,
|
||||
retrive_next_token=None,
|
||||
retrive_next_sibling=None,
|
||||
draft_token_num=num_tokens_per_bs,
|
||||
)
|
||||
spec_info.capture_hidden_mode = CaptureHiddenMode.NULL
|
||||
|
||||
return spec_info
|
||||
|
||||
spec_info = get_spec_info()
|
||||
if capture_hidden_mode != CaptureHiddenMode.FULL:
|
||||
capture_hidden_mode = (
|
||||
spec_info.capture_hidden_mode if spec_info else CaptureHiddenMode.NULL
|
||||
)
|
||||
|
||||
if self.server_args.enable_lora:
|
||||
lora_ids = [None] * batch_size
|
||||
else:
|
||||
lora_ids = None
|
||||
|
||||
forward_batch = ForwardBatch(
|
||||
forward_mode=capture_forward_mode,
|
||||
batch_size=batch_size,
|
||||
input_ids=buffers.input_ids,
|
||||
req_pool_indices=buffers.req_pool_indices,
|
||||
seq_lens=buffers.seq_lens,
|
||||
seq_lens_cpu=buffers.seq_lens_cpu,
|
||||
next_token_logits_buffer=buffers.next_token_logits_buffer,
|
||||
orig_seq_lens=buffers.seq_lens,
|
||||
req_to_token_pool=self.req_to_token_pool,
|
||||
token_to_kv_pool=self.token_to_kv_pool,
|
||||
attn_backend=self.attn_backend,
|
||||
out_cache_loc=buffers.out_cache_loc,
|
||||
seq_lens_sum=buffers.seq_lens.sum().item(),
|
||||
encoder_lens=buffers.encoder_lens,
|
||||
return_logprob=False,
|
||||
positions=buffers.positions,
|
||||
extend_num_tokens=extend_num_tokens,
|
||||
extend_seq_lens=extend_seq_lens,
|
||||
extend_prefix_lens=extend_prefix_lens,
|
||||
extend_start_loc=extend_start_loc,
|
||||
extend_prefix_lens_cpu=extend_prefix_lens_cpu,
|
||||
extend_seq_lens_cpu=extend_seq_lens_cpu,
|
||||
global_num_tokens_gpu=buffers.global_num_tokens_gpu,
|
||||
global_num_tokens_for_logprob_gpu=buffers.global_num_tokens_for_logprob_gpu,
|
||||
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
|
||||
global_dp_buffer_len=global_dp_buffer_len,
|
||||
mrope_positions=buffers.mrope_positions,
|
||||
spec_algorithm=self.spec_algorithm,
|
||||
spec_info=spec_info,
|
||||
capture_hidden_mode=capture_hidden_mode,
|
||||
num_token_non_padded=buffers.num_token_non_padded,
|
||||
global_forward_mode=capture_forward_mode,
|
||||
lora_ids=lora_ids,
|
||||
)
|
||||
|
||||
if lora_ids is not None:
|
||||
self.lora_manager.prepare_lora_batch(forward_batch)
|
||||
|
||||
self.attn_backend.init_forward_metadata(forward_batch)
|
||||
|
||||
def run_once():
|
||||
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
|
||||
set_dp_buffer_len(
|
||||
global_dp_buffer_len,
|
||||
num_tokens,
|
||||
forward_batch.dp_padding_mode.is_max_len(),
|
||||
)
|
||||
set_is_extend_in_batch(False)
|
||||
|
||||
kwargs = {}
|
||||
if (
|
||||
self.server_args.pp_size > 1
|
||||
and "pp_proxy_tensors"
|
||||
in inspect.signature(self.model.forward).parameters
|
||||
):
|
||||
kwargs["pp_proxy_tensors"] = PPProxyTensors(
|
||||
{k: v.clone() for k, v in pp_proxy_tensors.tensors.items()}
|
||||
)
|
||||
if not self.is_generation:
|
||||
kwargs["get_embedding"] = True
|
||||
|
||||
logits_output_or_pp_proxy_tensors = self.model.forward(
|
||||
buffers.input_ids,
|
||||
forward_batch.positions,
|
||||
forward_batch,
|
||||
**kwargs,
|
||||
)
|
||||
return logits_output_or_pp_proxy_tensors
|
||||
|
||||
torch.get_device_module(self.device).synchronize()
|
||||
self.tp_group.barrier()
|
||||
run_once()
|
||||
|
||||
def init_device_graphs(self):
|
||||
"""Capture device graphs."""
|
||||
self.graph_runner = None
|
||||
|
||||
@@ -378,6 +378,7 @@ class ServerArgs:
|
||||
mm_attention_backend: Optional[str] = None
|
||||
nsa_prefill_backend: str = "flashmla_sparse"
|
||||
nsa_decode_backend: str = "fa3"
|
||||
enable_flashinfer_autotune: bool = False
|
||||
|
||||
# Speculative decoding
|
||||
speculative_algorithm: Optional[str] = None
|
||||
@@ -2882,6 +2883,12 @@ class ServerArgs:
|
||||
type=str,
|
||||
choices=NSA_CHOICES,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--enable-flashinfer-autotune",
|
||||
default=ServerArgs.enable_flashinfer_autotune,
|
||||
action="store_true",
|
||||
help="Enable FlashInfer autotuning for optimal kernel selection.",
|
||||
)
|
||||
|
||||
# Speculative decoding
|
||||
parser.add_argument(
|
||||
|
||||
Reference in New Issue
Block a user