From 848ee57067f5e402aaeb7c135749019af96ee588 Mon Sep 17 00:00:00 2001 From: elvischenv <219235043+elvischenv@users.noreply.github.com> Date: Sat, 29 Nov 2025 16:05:37 +0800 Subject: [PATCH] feat: support flashinfer kernel autotune (#12306) Co-authored-by: Qiaolin Yu Co-authored-by: Kangyan-Zhou --- .../attention/hybrid_linear_attn_backend.py | 4 +- .../sglang/srt/layers/quantization/mxfp4.py | 4 +- .../sglang/srt/model_executor/model_runner.py | 305 +++++++++++++++++- python/sglang/srt/server_args.py | 7 + 4 files changed, 316 insertions(+), 4 deletions(-) diff --git a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py index fe436e95b..6e23a4639 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -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 diff --git a/python/sglang/srt/layers/quantization/mxfp4.py b/python/sglang/srt/layers/quantization/mxfp4.py index d44444a3a..4616879ce 100644 --- a/python/sglang/srt/layers/quantization/mxfp4.py +++ b/python/sglang/srt/layers/quantization/mxfp4.py @@ -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) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 92f9e8788..b70fb1fc2 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -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 diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index e2026160f..abe466632 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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(