From 3a6de13cd82266c61a50fb07c6277a871871c1bf Mon Sep 17 00:00:00 2001 From: YAMY <74099316+YAMY1234@users.noreply.github.com> Date: Wed, 20 May 2026 22:22:41 -0700 Subject: [PATCH] perf(dsv4): add MHC token-count prewarm (#25810) --- python/sglang/srt/layers/mhc.py | 18 +++ .../sglang/srt/model_executor/model_runner.py | 9 +- python/sglang/srt/models/deepseek_v4.py | 110 ++++++++++++++++++ python/sglang/srt/models/deepseek_v4_nextn.py | 5 + 4 files changed, 141 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/layers/mhc.py b/python/sglang/srt/layers/mhc.py index 414cda649..0593039d5 100644 --- a/python/sglang/srt/layers/mhc.py +++ b/python/sglang/srt/layers/mhc.py @@ -450,6 +450,24 @@ def _compute_num_split_for_mhc_pre(num_tokens: int, hc_hidden_size: int) -> int: return max(1, min(n_sms // max(grid_size, 1), num_block_k // 4)) +def get_mhc_pre_token_count_representatives( + max_num_tokens: int, hc_hidden_size: int +) -> Tuple[int, ...]: + """Return one token-count representative for each MHC pre split bucket.""" + if max_num_tokens <= 0: + return tuple() + + representatives_by_split: dict[int, int] = {} + for num_tokens in range(1, max_num_tokens + 1): + n_splits = _compute_num_split_for_mhc_pre(num_tokens, hc_hidden_size) + representatives_by_split[n_splits] = num_tokens + + return tuple( + representatives_by_split[n_splits] + for n_splits in sorted(representatives_by_split) + ) + + @tilelang.jit( pass_configs={ tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True, diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 373c0720f..075c57bdd 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -2292,7 +2292,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): def kernel_warmup(self): """ Warmup and tune kernels before cuda graph capture. - Currently only doing FlashInfer autotune. + Covers framework-level warmups and optional model-specific warmups. """ if self.device != "cuda": return @@ -2300,6 +2300,13 @@ class ModelRunner(ModelRunnerKVCacheMixin): if self._should_run_flashinfer_autotune(): self._flashinfer_autotune() + # Models may need their own warmup for model-specific kernels or JIT paths. + # Register those hooks on the model class so ModelRunner can keep this + # warmup entry point generic. + model_kernel_warmup = getattr(self.model, "kernel_warmup", None) + if model_kernel_warmup is not None: + model_kernel_warmup(self) + def _pre_initialize_flashinfer_allreduce_workspace(self): """Pre-initialize flashinfer allreduce fusion workspaces. diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index ead96a4a6..cb57ea400 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -2,6 +2,7 @@ from __future__ import annotations import concurrent.futures import logging +import time from typing import ( TYPE_CHECKING, Iterable, @@ -696,6 +697,70 @@ class DeepseekV4DecoderLayer(nn.Module): self.rms_norm_eps = config.rms_norm_eps self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() + def prewarm_mhc_token_counts( + self, token_counts: Tuple[int, ...], device: torch.device + ) -> None: + paths = ( + ( + "attn", + self.hc_attn_fn, + self.hc_attn_scale, + self.hc_attn_base, + self.input_layernorm, + ), + ( + "ffn", + self.hc_ffn_fn, + self.hc_ffn_scale, + self.hc_ffn_base, + self.post_attention_layernorm, + ), + ) + + with torch.inference_mode(): + for num_tokens in token_counts: + for path_name, hc_fn, hc_scale, hc_base, norm in paths: + tic = time.perf_counter() + residual = torch.empty( + (num_tokens, self.hc_mult, self.hidden_size), + dtype=torch.bfloat16, + device=device, + ) + y, post, comb, _ = self.hc_pre( + residual, + hc_fn, + hc_scale, + hc_base, + norm=norm, + ) + del residual, y, post, comb + torch.cuda.synchronize() + logger.info( + "DeepSeek V4 MHC prewarm path=%s num_tokens=%s completed in %.3fs", + path_name, + num_tokens, + time.perf_counter() - tic, + ) + + def prewarm_mhc_token_count_buckets( + self, max_num_tokens: int, device: torch.device + ) -> Tuple[int, ...]: + from sglang.srt.layers.mhc import get_mhc_pre_token_count_representatives + + token_counts = get_mhc_pre_token_count_representatives( + max_num_tokens, self.hc_mult * self.hidden_size + ) + if not token_counts: + return token_counts + + logger.info( + "DeepSeek V4 MHC prewarm max_num_tokens=%s representative token counts: %s", + max_num_tokens, + token_counts, + ) + self.prewarm_mhc_token_counts(token_counts, device) + return token_counts + def hc_pre( self, x: torch.Tensor, @@ -983,6 +1048,24 @@ class DeepseekV4Model(nn.Module): if self.dsa_enable_prefill_cp: self.cp_size = get_attention_cp_size() + def prewarm_mhc_token_count_buckets( + self, max_num_tokens: int, device: torch.device + ) -> Tuple[int, ...]: + tic = time.perf_counter() + logger.info( + "Running DeepSeek V4 MHC prewarm for max_num_tokens=%s", + max_num_tokens, + ) + token_counts = self.layers[self.start_layer].prewarm_mhc_token_count_buckets( + max_num_tokens, device + ) + logger.info( + "DeepSeek V4 MHC prewarm finished in %.3fs for representative token counts: %s", + time.perf_counter() - tic, + token_counts, + ) + return token_counts + def hc_head( self, x: torch.Tensor, @@ -1134,6 +1217,33 @@ class DeepseekV4ForCausalLM(nn.Module): self.cp_rank = get_attention_cp_rank() self.cp_size = get_attention_cp_size() + def prewarm_mhc_token_count_buckets( + self, max_num_tokens: int, device: torch.device + ) -> Tuple[int, ...]: + return self.model.prewarm_mhc_token_count_buckets(max_num_tokens, device) + + def kernel_warmup(self, model_runner) -> None: + if not model_runner.is_hybrid_swa: + return + if not envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.get(): + return + if not envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.get(): + return + + max_num_tokens = model_runner.server_args.chunked_prefill_size + if max_num_tokens is None or max_num_tokens <= 0: + max_num_tokens = 8192 + + token_counts = self.prewarm_mhc_token_count_buckets( + max_num_tokens, model_runner.device + ) + model_runner.tp_group.barrier() + + logger.info( + "DeepSeek V4 MHC prewarm completed for representative token-count shapes: %s", + token_counts, + ) + @property def routed_experts_weights_of_layer(self): return self._routed_experts_weights_of_layer.value diff --git a/python/sglang/srt/models/deepseek_v4_nextn.py b/python/sglang/srt/models/deepseek_v4_nextn.py index bd116c29a..f6d9a3f7d 100644 --- a/python/sglang/srt/models/deepseek_v4_nextn.py +++ b/python/sglang/srt/models/deepseek_v4_nextn.py @@ -128,6 +128,11 @@ class DeepseekV4ModelNextN(nn.Module): y = torch.sum(pre.unsqueeze(-1) * x.view(shape), dim=1) return y.to(dtype) + def prewarm_mhc_token_count_buckets( + self, max_num_tokens: int, device: torch.device + ) -> Tuple[int, ...]: + return self.decoder.prewarm_mhc_token_count_buckets(max_num_tokens, device) + def forward( self, input_ids: torch.Tensor,