From 23ea7b648135e433995ddd074d90526ee45618f8 Mon Sep 17 00:00:00 2001 From: weireweire Date: Tue, 4 Aug 2026 14:41:56 +0800 Subject: [PATCH] Prewarm DSV4 MHC post kernel at model load (#30741) Co-authored-by: weireweire <20922698+weireweire@users.noreply.github.com> --- python/sglang/srt/models/deepseek_v4.py | 35 ++++++++++++++++++------- 1 file changed, 25 insertions(+), 10 deletions(-) diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 0c59f2c6c..65dd55549 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -2752,8 +2752,8 @@ class DeepseekV4ForCausalLM(nn.Module): return name - def _prewarm_mhc_pre_kernels(self) -> None: - """One-shot mhc_pre() JIT prewarm at load time, synced across ranks. + def _prewarm_mhc_kernels(self) -> None: + """One-shot MHC JIT prewarm at load time, synced across ranks. Runs before any forward so the compile burst stays off the serving path; the barrier keeps ranks from proceeding while a peer is still @@ -2774,16 +2774,17 @@ class DeepseekV4ForCausalLM(nn.Module): if layer is None: return - from sglang.kernels.ops.layernorm.mhc import prewarm_mhc_pre + from sglang.kernels.ops.layernorm.mhc import mhc_post, prewarm_mhc_pre tic = time.perf_counter() + residual = torch.zeros( + (1, layer.hc_mult, layer.hidden_size), + dtype=torch.bfloat16, + device=layer.hc_attn_fn.device, + ) prewarm_mhc_pre( # Template carrying dtype/device; buckets allocate their own sizes. - residual=torch.zeros( - (1, layer.hc_mult, layer.hidden_size), - dtype=torch.bfloat16, - device=layer.hc_attn_fn.device, - ), + residual=residual, fn=layer.hc_attn_fn, hc_scale=layer.hc_attn_scale, hc_base=layer.hc_attn_base, @@ -2797,13 +2798,27 @@ class DeepseekV4ForCausalLM(nn.Module): norm_weight=layer.input_layernorm.weight.data, norm_eps=layer.input_layernorm.variance_epsilon, ) + mhc_post( + x=residual.new_zeros((1, layer.hidden_size)), + residual=residual, + post_layer_mix=torch.zeros( + (1, layer.hc_mult, 1), + dtype=torch.float32, + device=residual.device, + ), + comb_res_mix=torch.zeros( + (1, layer.hc_mult, layer.hc_mult), + dtype=torch.float32, + device=residual.device, + ), + ) torch.cuda.synchronize() compile_secs = time.perf_counter() - tic # Runs before init_memory_pool(); don't let transients skew pool sizing. torch.cuda.empty_cache() get_tp_group().barrier() logger.info( - "DeepSeek V4 MHC prenorm prewarm at load: compile %.1fs, rank sync +%.1fs", + "DeepSeek V4 MHC prewarm at load: compile %.1fs, rank sync +%.1fs", compile_secs, time.perf_counter() - tic - compile_secs, ) @@ -3153,7 +3168,7 @@ class DeepseekV4ForCausalLM(nn.Module): self.post_load_weights(is_nextn=is_nextn, weight_names=weight_names) if not is_nextn: - self._prewarm_mhc_pre_kernels() + self._prewarm_mhc_kernels() def get_embed_and_head(self): return self.model.embed_tokens.weight, self.lm_head.weight