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:
elvischenv
2025-11-29 16:05:37 +08:00
committed by GitHub
parent ce6b7dfce7
commit 848ee57067
4 changed files with 316 additions and 4 deletions

View File

@@ -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

View File

@@ -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)

View File

@@ -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

View File

@@ -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(