refactor(dsv4): route MHC prenorm through DeepGEMM wrapper (#26238)
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user