diff --git a/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py b/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py index f0cc1a44b..f4ae6663c 100644 --- a/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py +++ b/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py @@ -101,6 +101,7 @@ class DeepGemmKernelType(IntEnum): GROUPED_GEMM_NT_BF16_CONTIG = auto() GEMM_NT_F8F8BF16 = auto() GEMM_NT_BF16BF16F32 = auto() + TF32_HC_PRENORM_GEMM = auto() _INITIALIZATION_DICT: Dict[Tuple[DeepGemmKernelType, int, int, int], bool] = dict() @@ -209,7 +210,7 @@ def _compile_deep_gemm_one_type_all( deep_gemm.set_compile_mode(1) # TODO can use multi thread - for m in tqdm(m_list, desc=f"DeepGEMM warmup"): + for m in tqdm(m_list, desc="DeepGEMM warmup"): executor.execute(m=m) if has_compile_mode_api: deep_gemm.set_compile_mode(old_compile_mode) @@ -233,6 +234,7 @@ class _BaseWarmupExecutor: DeepGemmKernelType.GEMM_NT_BF16BF16F32: _BF16F32WarmupExecutor, DeepGemmKernelType.GROUPED_GEMM_NT_BF16_CONTIG: _BF16GroupedContWarmupExecutor, DeepGemmKernelType.GROUPED_GEMM_NT_BF16_MASKED: _BF16GroupedMaskedWarmupExecutor, + DeepGemmKernelType.TF32_HC_PRENORM_GEMM: _TF32HcPrenormWarmupExecutor, }[kernel_type](**kwargs) @staticmethod @@ -266,6 +268,11 @@ class _BaseWarmupExecutor: + num_groups * 4 + num_groups * max_m * n * 2 ) / _GB + elif kernel_type == DeepGemmKernelType.TF32_HC_PRENORM_GEMM: + # The generic hook's fourth dimension is num_splits for MHC. + # A value of 0 represents DeepGEMM's unsplit num_splits=None path. + num_splits = num_groups if num_groups > 0 else 1 + return (max_m * k * 2 + n * k * 4 + num_splits * max_m * (n + 1) * 4) / _GB else: raise ValueError(f"Invalid kernel type: {kernel_type}") @@ -396,6 +403,37 @@ class _BF16GroupedMaskedWarmupExecutor(_BaseWarmupExecutor): ) +class _TF32HcPrenormWarmupExecutor(_BaseWarmupExecutor): + def __init__(self, max_m: int, n: int, k: int, num_groups: int): + self.x = torch.empty((max_m, k), device="cuda", dtype=torch.bfloat16) + self.fn = torch.empty((n, k), device="cuda", dtype=torch.float32) + self.n = n + # The generic warmup executor's num_groups argument is num_splits here. + # A value of 0 represents DeepGEMM's unsplit num_splits=None path. + self.num_splits = num_groups if num_groups > 0 else None + + def execute(self, m): + if self.num_splits is None: + out = torch.empty((m, self.n), device="cuda", dtype=torch.float32) + sqrsum = torch.empty((m,), device="cuda", dtype=torch.float32) + else: + # Slicing the middle dimension of a preallocated + # (num_splits, max_m, n) output would create a strided view. + out = torch.empty( + (self.num_splits, m, self.n), device="cuda", dtype=torch.float32 + ) + sqrsum = torch.empty( + (self.num_splits, m), device="cuda", dtype=torch.float32 + ) + deep_gemm.tf32_hc_prenorm_gemm( + self.x[:m], + self.fn, + out, + sqrsum, + num_splits=self.num_splits, + ) + + def deep_gemm_execution_hook( m: int, n: int, k: int, num_groups: int, kernel_type: DeepGemmKernelType ): diff --git a/python/sglang/srt/layers/deep_gemm_wrapper/entrypoint.py b/python/sglang/srt/layers/deep_gemm_wrapper/entrypoint.py index 764b5345b..08fa2159e 100644 --- a/python/sglang/srt/layers/deep_gemm_wrapper/entrypoint.py +++ b/python/sglang/srt/layers/deep_gemm_wrapper/entrypoint.py @@ -185,6 +185,25 @@ def gemm_nt_bf16bf16f32( deep_gemm.bf16_gemm_nt(lhs, rhs, out) +def tf32_hc_prenorm_gemm( + x: torch.Tensor, + fn: torch.Tensor, + out: torch.Tensor, + sqrsum: torch.Tensor, + num_splits: Optional[int], +): + m, k = x.shape + n, _ = fn.shape + num_splits_key = num_splits if num_splits is not None else 0 + kernel_type = compile_utils.DeepGemmKernelType.TF32_HC_PRENORM_GEMM + + if m == 0: + return + + with compile_utils.deep_gemm_execution_hook(m, n, k, num_splits_key, kernel_type): + deep_gemm.tf32_hc_prenorm_gemm(x, fn, out, sqrsum, num_splits=num_splits) + + def update_deep_gemm_config(gpu_id: int, server_args: ServerArgs): compile_utils.update_deep_gemm_config(gpu_id, server_args) diff --git a/python/sglang/srt/layers/mhc.py b/python/sglang/srt/layers/mhc.py index 0593039d5..d7d0d3c7b 100644 --- a/python/sglang/srt/layers/mhc.py +++ b/python/sglang/srt/layers/mhc.py @@ -450,24 +450,6 @@ 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, @@ -690,8 +672,6 @@ def mhc_pre( ) if envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.get(): - import deep_gemm - n_splits = _compute_num_split_for_mhc_pre(num_tokens, hc_hidden_size) gemm_out_mul = torch.empty( @@ -701,12 +681,14 @@ def mhc_pre( n_splits, num_tokens, dtype=torch.float32, device=residual.device ) - deep_gemm.tf32_hc_prenorm_gemm( + from sglang.srt.layers.deep_gemm_wrapper.entrypoint import tf32_hc_prenorm_gemm + + tf32_hc_prenorm_gemm( residual_flat.view(num_tokens, hc_hidden_size), fn_flat, gemm_out_mul, gemm_out_sqrsum, - num_splits=n_splits, + n_splits, ) gemm_last_dim = hc_mult3 big_fuse_n_splits = n_splits diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index fad8a8868..afe4be2cd 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -2305,7 +2305,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): def kernel_warmup(self): """ Warmup and tune kernels before cuda graph capture. - Covers framework-level warmups and optional model-specific warmups. + Currently only doing FlashInfer autotune. """ if self.device != "cuda": return @@ -2313,13 +2313,6 @@ 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 fc82de69f..78e1495c0 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -2,7 +2,6 @@ from __future__ import annotations import concurrent.futures import logging -import time from contextlib import nullcontext from typing import ( TYPE_CHECKING, @@ -978,70 +977,6 @@ 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, @@ -1111,7 +1046,9 @@ class DeepseekV4DecoderLayer(nn.Module): return y, post.squeeze(-1), comb, False if envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.get(): - import deep_gemm + from sglang.srt.layers.deep_gemm_wrapper.entrypoint import ( + tf32_hc_prenorm_gemm, + ) x_flat = x.flatten(1).bfloat16() @@ -1119,7 +1056,7 @@ class DeepseekV4DecoderLayer(nn.Module): mix_hc = hc_fn.size(0) d_out = torch.empty((m, mix_hc), dtype=torch.float, device=x.device) s_out = torch.empty((m,), dtype=torch.float, device=x.device) - deep_gemm.tf32_hc_prenorm_gemm( + tf32_hc_prenorm_gemm( x_flat, hc_fn.float().contiguous(), d_out, s_out, num_splits=None ) rsqrt = torch.rsqrt(s_out / k + self.rms_norm_eps) @@ -1368,24 +1305,6 @@ 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, @@ -1547,33 +1466,6 @@ 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 069a4ec8c..ba71a0f2d 100644 --- a/python/sglang/srt/models/deepseek_v4_nextn.py +++ b/python/sglang/srt/models/deepseek_v4_nextn.py @@ -129,11 +129,6 @@ 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,