refactor(dsv4): route MHC prenorm through DeepGEMM wrapper (#26238)

This commit is contained in:
YAMY
2026-05-27 17:45:45 -07:00
committed by GitHub
parent 68e5b4fdd6
commit eae03ce3b2
6 changed files with 67 additions and 148 deletions
@@ -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
):
@@ -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)
+4 -22
View File
@@ -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
@@ -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.
+4 -112
View File
@@ -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
@@ -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,