From 7c5708cba734f8dddb504a0d330d6360aa8005c6 Mon Sep 17 00:00:00 2001 From: Qichao Li Date: Sat, 30 May 2026 17:04:51 +0800 Subject: [PATCH] [DeepSeek-V4] Add mhc_fused_post_pre kernel (#25976) Co-authored-by: Qichao Li --- python/sglang/srt/environ.py | 1 + python/sglang/srt/layers/mhc.py | 497 ++++++++++++++++++ python/sglang/srt/models/deepseek_v4.py | 310 +++++++++-- python/sglang/srt/models/deepseek_v4_nextn.py | 6 +- tests/kernels/test_mhc_kernels.py | 111 ++++ 5 files changed, 876 insertions(+), 49 deletions(-) create mode 100644 tests/kernels/test_mhc_kernels.py diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index dd6d85628..33adc21ee 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -657,6 +657,7 @@ class Envs: SGLANG_OPT_USE_TILELANG_MHC_PRE = EnvBool(True) SGLANG_OPT_USE_TILELANG_MHC_POST = EnvBool(True) SGLANG_OPT_USE_TRITON_FUSED_MHC = EnvBool(True) + SGLANG_OPT_FUSE_MHC_POST_PRE = EnvBool(False) SGLANG_OPT_USE_TILELANG_INDEXER = EnvBool(False) SGLANG_OPT_USE_AITER_INDEXER = EnvBool(False) SGLANG_OPT_USE_JIT_INDEXER_METADATA = EnvBool(True) diff --git a/python/sglang/srt/layers/mhc.py b/python/sglang/srt/layers/mhc.py index d7d0d3c7b..76bf69557 100644 --- a/python/sglang/srt/layers/mhc.py +++ b/python/sglang/srt/layers/mhc.py @@ -896,3 +896,500 @@ def mhc_post( residual.shape[-1], ) return out + + +@tilelang.jit( + pass_configs={ + tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True, + tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True, + tilelang.PassConfigKey.TL_PTXAS_REGISTER_USAGE_LEVEL: 10, + }, +) +def mhc_fused_post_pre_fma_tilelang( + prev_comb_mix, + prev_residual, + prev_post_mix, + hidden_in, + pre_fn, + mixes_partial_out, + sqrsum_partial_out, + cur_residual_out, + hc: int, + hidden_size: int, + num_mix_outputs: int, + n_thr: int = 256, + tile_mix_outputs: int = 1, + split_k: int = 1, +) -> tilelang.JITKernel: + num_tokens = T.dynamic("num_tokens") + split_k = T.dynamic("split_k") + + hidden_per_split = (hidden_size + split_k - 1) // split_k + num_mix_output_tiles = (num_mix_outputs + tile_mix_outputs - 1) // tile_mix_outputs + + prev_comb_mix: T.Tensor((num_tokens, hc, hc), T.float32) + prev_residual: T.Tensor((num_tokens, hc, hidden_size), T.bfloat16) + prev_post_mix: T.Tensor((num_tokens, hc), T.float32) + hidden_in: T.Tensor((num_tokens, hidden_size), T.bfloat16) + pre_fn: T.Tensor((num_mix_outputs, hc, hidden_size), T.float32) + + mixes_partial_out: T.Tensor((split_k, num_tokens, num_mix_outputs), T.float32) + sqrsum_partial_out: T.Tensor((split_k, num_tokens), T.float32) + cur_residual_out: T.Tensor((num_tokens, hc, hidden_size), T.bfloat16) + + hidden_iters_per_thread = (hidden_per_split + n_thr - 1) // n_thr + num_warps = n_thr // 32 + + ENABLE_PDL = is_arch_support_pdl() + + # CTA assignment: + # token_idx : this CTA handles one token. + # mix_output_tile_idx : this CTA handles a small tile of mix output columns. + # For HC=4, num_mix_outputs = 24: + # [0:4] -> pre logits + # [4:8] -> post logits + # [8:24] -> comb logits + # hidden_split_idx : this CTA handles one split of the hidden dimension. + # + # Thread assignment inside one CTA: + # Each thread owns several hidden positions in this hidden split: + # hidden_idx = hidden_split_start + hidden_iter * n_thr + thread_idx + # + # For each owned hidden_idx, the thread computes: + # 1. post result: cur_residual[token, :, hidden_idx] + # 2. sqrsum partial for pre RMS + # 3. GEMM partial for several mix output columns + with T.Kernel( + num_tokens, + num_mix_output_tiles, + split_k, + threads=n_thr, + ) as (token_idx, mix_output_tile_idx, hidden_split_idx): + thread_idx = T.get_thread_binding() + warp_idx = T.get_warp_idx() + lane_idx = T.get_lane_idx() + + warp_partials = T.alloc_shared((num_warps, tile_mix_outputs + 1), T.float32) + post_mix_smem = T.alloc_shared((hc,), T.float32) + comb_mix_smem = T.alloc_shared((hc, hc), T.float32) + + post_mix_for_token = T.alloc_local((hc,), T.float32) + comb_mix_for_token = T.alloc_local((hc, hc), T.float32) + + mix_acc = T.alloc_local((tile_mix_outputs,), T.float32) + sqrsum_acc = T.alloc_local((1,), T.float32) + cur_residual_values = T.alloc_local((hc,), T.float32) + + T.clear(mix_acc) + T.clear(sqrsum_acc) + + hidden_split_start = hidden_split_idx * hidden_per_split + + if ENABLE_PDL: + T.pdl_sync() + + # Load post/comb coefficients for this token. + # + # PyTorch equivalent: + # post = prev_post_mix[token_idx] # [HC] + # comb = prev_comb_mix[token_idx] # [HC, HC] + T.copy(prev_post_mix[token_idx, 0], post_mix_smem) + T.copy(prev_comb_mix[token_idx, 0, 0], comb_mix_smem) + + for route_idx in T.unroll(hc): + post_mix_for_token[route_idx] = post_mix_smem[route_idx] + + for old_route_idx in T.unroll(hc): + for new_route_idx in T.unroll(hc): + comb_mix_for_token[old_route_idx, new_route_idx] = comb_mix_smem[ + old_route_idx, new_route_idx + ] + + for hidden_iter in T.serial(hidden_iters_per_thread): + hidden_idx = hidden_split_start + hidden_iter * n_thr + thread_idx + + if hidden_idx < hidden_size: + # Step A: fused post. + # + # PyTorch equivalent: + # cur_residual = + # post.unsqueeze(-1) * hidden_in.unsqueeze(1) + # + ( + # comb.unsqueeze(-1) + # * prev_residual.unsqueeze(2) + # ).sum(dim=1) + # + # Scalar form for this token and this hidden position: + # cur_residual[j, h] + # = post[j] * hidden_in[h] + # + sum_k comb[k, j] * prev_residual[k, h] + for new_route_idx in T.unroll(hc): + cur_residual_values[new_route_idx] = ( + post_mix_for_token[new_route_idx] + * hidden_in[token_idx, hidden_idx] + ) + + for old_route_idx in T.unroll(hc): + cur_residual_values[new_route_idx] += ( + comb_mix_for_token[old_route_idx, new_route_idx] + * prev_residual[token_idx, old_route_idx, hidden_idx] + ) + + # Match the unfused path: + # mhc_post writes bf16 residual, + # then mhc_pre reads bf16 residual. + for route_idx in T.unroll(hc): + cur_residual_values[route_idx] = T.bfloat16( + cur_residual_values[route_idx] + ) + + # Step B1: pre sqrsum partial. + # + # PyTorch equivalent: + # x_flat = cur_residual.reshape(T, HC * H).float() + # sqrsum = (x_flat * x_flat).sum(dim=-1) + # + # Only mix_output_tile_idx == 0 writes cur_residual and sqrsum, + # otherwise different output-column CTAs would duplicate this work. + if mix_output_tile_idx == 0: + for route_idx in T.unroll(hc): + cur_residual_out[token_idx, route_idx, hidden_idx] = ( + cur_residual_values[route_idx] + ) + sqrsum_acc[0] += ( + cur_residual_values[route_idx] + * cur_residual_values[route_idx] + ) + + # Step B2: pre GEMM partial. + # + # PyTorch equivalent: + # mixes = F.linear(x_flat, fn) + # + # Scalar form: + # mixes[token, o] += + # pre_fn[o, route, hidden] * cur_residual[route, hidden] + # + # This CTA computes only tile_mix_outputs columns of mixes. + for tile_col_idx in T.unroll(tile_mix_outputs): + mix_output_idx = ( + mix_output_tile_idx * tile_mix_outputs + tile_col_idx + ) + + if mix_output_idx < num_mix_outputs: + for route_idx in T.unroll(hc): + mix_acc[tile_col_idx] += ( + pre_fn[mix_output_idx, route_idx, hidden_idx] + * cur_residual_values[route_idx] + ) + + # Reduce thread partials inside each warp. + for tile_col_idx in T.unroll(tile_mix_outputs): + mix_acc[tile_col_idx] = T.warp_reduce_sum(mix_acc[tile_col_idx]) + + if mix_output_tile_idx == 0: + sqrsum_acc[0] = T.warp_reduce_sum(sqrsum_acc[0]) + + # One lane per warp writes warp-level partials to shared memory. + if lane_idx == 0: + for tile_col_idx in T.unroll(tile_mix_outputs): + warp_partials[warp_idx, tile_col_idx] = mix_acc[tile_col_idx] + + if mix_output_tile_idx == 0: + warp_partials[warp_idx, tile_mix_outputs] = sqrsum_acc[0] + + T.sync_threads() + + # Reduce across warps and write split partials. + # + # The full PyTorch result would be: + # mixes = F.linear(cur_residual.reshape(T, HC * H), fn) + # sqrsum = (cur_residual.float() ** 2).sum(dim=(1, 2)) + # + # This kernel is split along hidden, so each CTA writes only: + # mixes_partial_out[hidden_split_idx, token, o] + # sqrsum_partial_out[hidden_split_idx, token] + # + # Later mhc_pre_big_fuse does: + # mixes = mixes_partial_out.sum(dim=0) + # sqrsum = sqrsum_partial_out.sum(dim=0) + # rms = rsqrt(sqrsum / (HC * H) + eps) + # mixes *= rms + # mixes -> pre/post/comb + # layer_input = sum_j pre[j] * cur_residual[j] + if warp_idx == 0: + for tile_col_idx in T.unroll(tile_mix_outputs): + mix_output_idx = mix_output_tile_idx * tile_mix_outputs + tile_col_idx + + if mix_output_idx < num_mix_outputs and lane_idx == tile_col_idx: + mix_output_partial = T.alloc_var(T.float32, init=0.0) + + for reduce_warp_idx in T.unroll(num_warps): + mix_output_partial += warp_partials[ + reduce_warp_idx, tile_col_idx + ] + + mixes_partial_out[hidden_split_idx, token_idx, mix_output_idx] = ( + mix_output_partial + ) + + if mix_output_tile_idx == 0 and lane_idx == 0: + sqrsum_partial = T.alloc_var(T.float32, init=0.0) + + for reduce_warp_idx in T.unroll(num_warps): + sqrsum_partial += warp_partials[reduce_warp_idx, tile_mix_outputs] + + sqrsum_partial_out[hidden_split_idx, token_idx] = sqrsum_partial + + if ENABLE_PDL: + T.pdl_trigger() + + +def mhc_fused_post_pre( + x: torch.Tensor, + residual: torch.Tensor, + post_layer_mix: torch.Tensor, + comb_res_mix: torch.Tensor, + fn: torch.Tensor, + hc_scale: torch.Tensor, + hc_base: torch.Tensor, + rms_eps: float, + hc_pre_eps: float, + hc_sinkhorn_eps: float, + hc_post_mult_value: float, + sinkhorn_repeat: int, + n_splits: int = 1, + tile_n: int = 1, + *, + norm_weight: torch.Tensor | None = None, + norm_eps: float | None = None, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Fuse the boundary between one mHC post step and the next mHC pre step. + + The unfused sequence is ``mhc_post -> pre-norm GEMM -> mhc_pre big_fuse``. + This wrapper keeps the numerically sensitive ``mhc_pre_big_fuse`` stage, + including optional RMSNorm, but removes the separate post/pre boundary. + Small token batches use the FMA kernel above to combine ``mhc_post`` and the + pre-norm GEMM in one launch; larger batches keep DeepGEMM for throughput and + only fuse the Python/model-level scheduling boundary. + + Returns: + residual_cur: post-mapped residual, shape (..., hc_mult, hidden_size) + post_mix_cur: shape (..., hc_mult, 1) + comb_mix_cur: shape (..., hc_mult, hc_mult) + layer_input_cur: shape (..., hidden_size) + """ + + assert residual.dtype == torch.bfloat16 + assert x.dtype == torch.bfloat16 + assert post_layer_mix.dtype == torch.float32 + assert comb_res_mix.dtype == torch.float32 + assert fn.dtype == torch.float32 + assert hc_scale.dtype == torch.float32 + assert hc_base.dtype == torch.float32 + + hc_mult = residual.shape[-2] + hidden_size = residual.shape[-1] + hc_mult2 = hc_mult * hc_mult + hc_mult3 = hc_mult * 2 + hc_mult2 + hc_hidden_size = hc_mult * hidden_size + outer_shape = residual.shape[:-2] + + assert x.shape == (*outer_shape, hidden_size) + assert post_layer_mix.shape in ( + (*outer_shape, hc_mult, 1), + (*outer_shape, hc_mult), + ) + assert comb_res_mix.shape == (*outer_shape, hc_mult, hc_mult) + assert fn.shape == (hc_mult3, hc_hidden_size) + assert hc_scale.shape == (3,) + assert hc_base.shape == (hc_mult3,) + + residual_flat = residual.view(-1, hc_mult, hidden_size) + num_tokens = residual_flat.shape[0] + if num_tokens == 0: + # Some DP/EP ranks can receive no tokens; return correctly typed empty + # tensors so later fused layers keep the same contracts as mhc_pre/hc_post. + return ( + torch.empty_like(residual), + torch.empty( + (*outer_shape, hc_mult, 1), dtype=torch.float32, device=residual.device + ), + torch.empty( + (*outer_shape, hc_mult, hc_mult), + dtype=torch.float32, + device=residual.device, + ), + torch.empty( + (*outer_shape, hidden_size), + dtype=torch.bfloat16, + device=residual.device, + ), + ) + x_flat = x.view(num_tokens, hidden_size) + + # The scalar-FMA kernel wins only for small batches where launch + # overhead dominates; beyond the threshold DeepGEMM's tensor-core path wins. + fma_token_threshold = 32 + if num_tokens <= fma_token_threshold: + tile_n = 2 if num_tokens < 8 else 3 + n_splits = 8 if (num_tokens < 8 and hidden_size <= 4096) else 4 + else: + n_splits = _compute_num_split_for_mhc_pre(num_tokens, hc_hidden_size) + + gemm_out_mul = torch.empty( + n_splits, + num_tokens, + hc_mult3, + dtype=torch.float32, + device=residual.device, + ) + gemm_out_sqrsum = torch.empty( + n_splits, + num_tokens, + dtype=torch.float32, + device=residual.device, + ) + residual_cur = torch.empty_like(residual_flat) + + if num_tokens <= fma_token_threshold: + # Small-batch path: one TileLang launch computes hc_post, the bf16 + # residual write, GEMM partials, and the RMS square-sum partials. + mhc_fused_post_pre_fma_tilelang( + comb_res_mix.view(num_tokens, hc_mult, hc_mult), + residual_flat, + post_layer_mix.view(num_tokens, hc_mult), + x_flat, + fn.view(hc_mult3, hc_mult, hidden_size), + gemm_out_mul, + gemm_out_sqrsum, + residual_cur, + hc_mult, + hidden_size, + hc_mult3, + tile_mix_outputs=tile_n, + split_k=n_splits, + ) + else: + # Large-batch path: keep the existing high-throughput TileLang hc_post + + # DeepGEMM pre-norm GEMM decomposition instead of replacing tensor cores. + mhc_post_tilelang( + comb_res_mix.view(num_tokens, hc_mult, hc_mult), + residual_flat, + post_layer_mix.view(num_tokens, hc_mult), + x_flat, + residual_cur, + hc_mult, + hidden_size, + ) + + if envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.get(): + import deep_gemm + + deep_gemm.tf32_hc_prenorm_gemm( + residual_cur.view(num_tokens, hc_hidden_size), + fn, + gemm_out_mul, + gemm_out_sqrsum, + num_splits=n_splits, + ) + else: + # Fallback mirrors mhc_pre when DeepGEMM prenorm is disabled. + n_splits = 1 + gemm_out_mul_2d = torch.empty( + num_tokens, hc_mult3, dtype=torch.float32, device=residual.device + ) + gemm_out_sqrsum_1d = torch.empty( + num_tokens, dtype=torch.float32, device=residual.device + ) + mhc_pre_gemm_sqrsum_tilelang( + residual_cur.view(num_tokens, hc_hidden_size), + fn, + gemm_out_mul_2d, + gemm_out_sqrsum_1d, + hc_mult3, + hc_hidden_size, + ) + gemm_out_mul = gemm_out_mul_2d.unsqueeze(0) + gemm_out_sqrsum = gemm_out_sqrsum_1d.unsqueeze(0) + + post_mix_cur = torch.empty( + num_tokens, + hc_mult, + dtype=torch.float32, + device=residual.device, + ) + comb_mix_cur = torch.empty( + num_tokens, + hc_mult2, + dtype=torch.float32, + device=residual.device, + ) + layer_input_cur = torch.empty( + num_tokens, + hidden_size, + dtype=torch.bfloat16, + device=residual.device, + ) + + if norm_weight is not None: + # Final mhc_pre stage: convert GEMM partials into post/comb/layer_input + # and fuse the following RMSNorm when the model passed a norm weight. + assert norm_eps is not None + assert norm_weight.shape == (hidden_size,) + norm_weight_bf = ( + norm_weight.bfloat16() + if norm_weight.dtype != torch.bfloat16 + else norm_weight + ) + if not norm_weight_bf.is_contiguous(): + norm_weight_bf = norm_weight_bf.contiguous() + mhc_pre_big_fuse_with_norm_tilelang( + gemm_out_mul, + gemm_out_sqrsum, + hc_scale, + hc_base, + residual_cur, + post_mix_cur, + comb_mix_cur, + layer_input_cur, + norm_weight_bf, + hidden_size, + rms_eps, + hc_pre_eps, + hc_sinkhorn_eps, + hc_post_mult_value, + sinkhorn_repeat, + norm_eps, + n_splits, + hc_mult, + hc_mult3, + ) + else: + # Same mhc_pre finalization without the model-layer RMSNorm. + mhc_pre_big_fuse_tilelang( + gemm_out_mul, + gemm_out_sqrsum, + hc_scale, + hc_base, + residual_cur, + post_mix_cur, + comb_mix_cur, + layer_input_cur, + hidden_size, + rms_eps, + hc_pre_eps, + hc_sinkhorn_eps, + hc_post_mult_value, + sinkhorn_repeat, + n_splits, + hc_mult, + hc_mult3, + ) + + return ( + residual_cur.view(*outer_shape, hc_mult, hidden_size), + post_mix_cur.view(*outer_shape, hc_mult, 1), + comb_mix_cur.view(*outer_shape, hc_mult, hc_mult), + layer_input_cur.view(*outer_shape, hidden_size), + ) diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 78e1495c0..e9a477614 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 contextlib import nullcontext from typing import ( TYPE_CHECKING, @@ -61,6 +62,7 @@ from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear from sglang.srt.layers.logits_processor import LogitsProcessor +from sglang.srt.layers.mhc import mhc_fused_post_pre from sglang.srt.layers.moe import get_moe_a2a_backend from sglang.srt.layers.moe.fused_moe_triton import FusedMoE from sglang.srt.layers.quantization.fp8_kernel import sglang_per_token_group_quant_fp8 @@ -110,6 +112,18 @@ from sglang.srt.utils.hf_transformers_utils import get_rope_config logger = logging.getLogger(__name__) _FP8_WO_A_GEMM = envs.SGLANG_OPT_FP8_WO_A_GEMM.get() +_MHC_POST_MULT_VALUE = 2.0 + + +def _is_fused_mhc_post_pre_enabled() -> bool: + # The fused path directly reuses TileLang mhc_post/mhc_pre kernels and their + # tensor layout assumptions, so keep it disabled when either dependency is off. + return ( + envs.SGLANG_OPT_FUSE_MHC_POST_PRE.get() + and envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.get() + and envs.SGLANG_OPT_USE_TILELANG_MHC_POST.get() + ) + _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip _is_gfx95_supported = is_gfx95_supported() @@ -976,6 +990,133 @@ class DeepseekV4DecoderLayer(nn.Module): self.hc_ffn_scale = nn.Parameter(torch.empty(3, dtype=torch.float32)) self.rms_norm_eps = config.rms_norm_eps self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() + self.use_fused_mhc_post_pre = _is_fused_mhc_post_pre_enabled() + self._input_layernorm_weight_bf16 = None + self._post_attention_layernorm_weight_bf16 = None + + def refresh_mhc_norm_weight_cache(self): + # Cache bf16 norm weights so the fused path does not allocate/cast per forward. + self._input_layernorm_weight_bf16 = ( + self.input_layernorm.weight.data.bfloat16().contiguous() + ) + self._post_attention_layernorm_weight_bf16 = ( + self.post_attention_layernorm.weight.data.bfloat16().contiguous() + ) + + 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, + ) + + if self.use_fused_mhc_post_pre: + for num_tokens in token_counts: + for path_name, hc_fn, hc_scale, hc_base, norm in paths: + tic = time.perf_counter() + # Dummy inputs matching the fused kernel's expected shapes. + x = torch.empty( + (num_tokens, self.hidden_size), + dtype=torch.bfloat16, + device=device, + ) + residual = torch.empty( + (num_tokens, self.hc_mult, self.hidden_size), + dtype=torch.bfloat16, + device=device, + ) + post_mix = torch.empty( + (num_tokens, self.hc_mult, 1), + dtype=torch.float32, + device=device, + ) + comb_mix = torch.empty( + (num_tokens, self.hc_mult, self.hc_mult), + dtype=torch.float32, + device=device, + ) + norm_weight = norm.weight.data.bfloat16().contiguous() + mhc_fused_post_pre( + x, + residual, + post_mix, + comb_mix, + hc_fn, + hc_scale, + hc_base, + self.rms_norm_eps, + self.hc_eps, + self.hc_eps, + _MHC_POST_MULT_VALUE, + self.hc_sinkhorn_iters, + norm_weight=norm_weight, + norm_eps=norm.variance_epsilon, + ) + del x, residual, post_mix, comb_mix, norm_weight + torch.cuda.synchronize() + logger.info( + "DeepSeek V4 MHC fused 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, @@ -1001,9 +1142,9 @@ class DeepseekV4DecoderLayer(nn.Module): if x.shape[0] == 0: y = torch.empty((0, shape[-1]), dtype=dtype, device=x.device) - post = torch.empty((0, self.hc_mult), dtype=dtype, device=x.device) + post = torch.empty((0, self.hc_mult), dtype=torch.float32, device=x.device) comb = torch.empty( - (0, self.hc_mult, self.hc_mult), dtype=dtype, device=x.device + (0, self.hc_mult, self.hc_mult), dtype=torch.float32, device=x.device ) return y, post, comb, False @@ -1023,7 +1164,7 @@ class DeepseekV4DecoderLayer(nn.Module): rms_eps=self.rms_norm_eps, hc_pre_eps=self.hc_eps, hc_sinkhorn_eps=self.hc_eps, - hc_post_mult_value=2.0, + hc_post_mult_value=_MHC_POST_MULT_VALUE, sinkhorn_repeat=self.hc_sinkhorn_iters, **norm_kwargs, ) @@ -1040,7 +1181,7 @@ class DeepseekV4DecoderLayer(nn.Module): rms_eps=self.rms_norm_eps, hc_pre_eps=self.hc_eps, hc_sinkhorn_eps=self.hc_eps, - hc_post_mult_value=2.0, + hc_post_mult_value=_MHC_POST_MULT_VALUE, sinkhorn_repeat=self.hc_sinkhorn_iters, ) return y, post.squeeze(-1), comb, False @@ -1122,27 +1263,60 @@ class DeepseekV4DecoderLayer(nn.Module): input_ids: torch.Tensor, forward_batch: ForwardBatch, input_ids_global: torch.Tensor, - ) -> torch.Tensor: - residual = hidden_states - hidden_states, post, comb, norm_fused = self.hc_pre( - hidden_states, - self.hc_attn_fn, - self.hc_attn_scale, - self.hc_attn_base, - norm=self.input_layernorm, - ) - if not norm_fused: - if _use_aiter and _is_gfx95_supported: - x_quant, hidden_states = _fused_rmsnorm_fp8_quant( - hidden_states, - self.input_layernorm.weight, - self.rms_norm_eps, - ) - else: - hidden_states = self.input_layernorm(hidden_states) - x_quant = None - else: + prev_residual: Optional[torch.Tensor] = None, + prev_post: Optional[torch.Tensor] = None, + prev_comb: Optional[torch.Tensor] = None, + ) -> Tuple[ + torch.Tensor, + Optional[torch.Tensor], + Optional[torch.Tensor], + Optional[torch.Tensor], + ]: + use_fused = self.use_fused_mhc_post_pre + + if prev_residual is not None and use_fused: + residual, post, comb, hidden_states = mhc_fused_post_pre( + hidden_states, + prev_residual, + prev_post, + prev_comb, + self.hc_attn_fn, + self.hc_attn_scale, + self.hc_attn_base, + self.rms_norm_eps, + self.hc_eps, + self.hc_eps, + _MHC_POST_MULT_VALUE, + self.hc_sinkhorn_iters, + norm_weight=( + self._input_layernorm_weight_bf16 + if self._input_layernorm_weight_bf16 is not None + else self.input_layernorm.weight.data + ), + norm_eps=self.input_layernorm.variance_epsilon, + ) x_quant = None + else: + residual = hidden_states + hidden_states, post, comb, norm_fused = self.hc_pre( + hidden_states, + self.hc_attn_fn, + self.hc_attn_scale, + self.hc_attn_base, + norm=self.input_layernorm, + ) + if not norm_fused: + if _use_aiter and _is_gfx95_supported: + x_quant, hidden_states = _fused_rmsnorm_fp8_quant( + hidden_states, + self.input_layernorm.weight, + self.rms_norm_eps, + ) + else: + hidden_states = self.input_layernorm(hidden_states) + x_quant = None + else: + x_quant = None hidden_states = self.self_attn( x=hidden_states, @@ -1151,35 +1325,58 @@ class DeepseekV4DecoderLayer(nn.Module): x_quant=x_quant, ) - fused_mhc = try_fused_hc_post_pre( - hidden_states, - residual, - post, - comb, - self.hc_ffn_fn.T, - self.hc_ffn_scale, - self.hc_ffn_base, - self.hc_mult, - self.rms_norm_eps, - self.hc_eps, - 2.0, - self.hc_sinkhorn_iters, - _is_gfx95_supported, - ) - if fused_mhc is not None: - residual, hidden_states, post, comb, norm_fused = fused_mhc + if use_fused: + fused_mhc = try_fused_hc_post_pre( + hidden_states, + residual, + post, + comb, + self.hc_ffn_fn.T, + self.hc_ffn_scale, + self.hc_ffn_base, + self.hc_mult, + self.rms_norm_eps, + self.hc_eps, + _MHC_POST_MULT_VALUE, + self.hc_sinkhorn_iters, + _is_gfx95_supported, + ) + if fused_mhc is not None: + residual, hidden_states, post, comb, norm_fused = fused_mhc + else: + residual, post, comb, hidden_states = mhc_fused_post_pre( + hidden_states, + residual, + post.unsqueeze(-1) if post.ndim == 2 else post, + comb, + self.hc_ffn_fn, + self.hc_ffn_scale, + self.hc_ffn_base, + self.rms_norm_eps, + self.hc_eps, + self.hc_eps, + _MHC_POST_MULT_VALUE, + self.hc_sinkhorn_iters, + norm_weight=( + self._post_attention_layernorm_weight_bf16 + if self._post_attention_layernorm_weight_bf16 is not None + else self.post_attention_layernorm.weight.data + ), + norm_eps=self.post_attention_layernorm.variance_epsilon, + ) + norm_fused = True else: hidden_states = self.hc_post(hidden_states, residual, post, comb) - residual = hidden_states # [n, hc, d] + residual = hidden_states hidden_states, post, comb, norm_fused = self.hc_pre( hidden_states, self.hc_ffn_fn, self.hc_ffn_scale, self.hc_ffn_base, norm=self.post_attention_layernorm, - ) # -> [n, d] - if not norm_fused: - hidden_states = self.post_attention_layernorm(hidden_states) + ) + if not norm_fused: + hidden_states = self.post_attention_layernorm(hidden_states) _use_cp = self.dsa_enable_prefill_cp and dsa_use_prefill_cp(forward_batch) _use_tp_moe_gather = ( @@ -1233,9 +1430,13 @@ class DeepseekV4DecoderLayer(nn.Module): attn_tp_all_gather(gathered, hidden_states.contiguous()) hidden_states = torch.cat(gathered) - hidden_states = self.hc_post(hidden_states, residual, post, comb) + if not use_fused: + hidden_states = self.hc_post(hidden_states, residual, post, comb) + return hidden_states, None, None, None - return hidden_states + # Return the deferred FFN hc_post state; the next layer consumes it with + # cross-layer fusion, and the final layer is completed in DeepseekV4Model. + return hidden_states, residual, post, comb class DeepseekV4Model(nn.Module): @@ -1302,6 +1503,7 @@ class DeepseekV4Model(nn.Module): self.hc_head_scale = nn.Parameter(torch.empty(1, dtype=torch.float32)) self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() + self.use_fused_mhc_post_pre = _is_fused_mhc_post_pre_enabled() if self.dsa_enable_prefill_cp: self.cp_size = get_attention_cp_size() @@ -1375,21 +1577,32 @@ class DeepseekV4Model(nn.Module): # forks alt-streams; later per-layer calls become no-ops. get_attn_backend()._maybe_upgrade_forward_metadata() + use_fused = self.use_fused_mhc_post_pre + prev_residual, prev_post, prev_comb = None, None, None + last_layer = None for i in range(self.start_layer, self.end_layer): layer = self.layers[i] + last_layer = layer ctx = ( nullcontext() if not get_global_server_args().disable_piecewise_cuda_graph else get_global_expert_distribution_recorder().with_current_layer(i) ) with ctx: - hidden_states = layer( + hidden_states, prev_residual, prev_post, prev_comb = layer( positions=positions, hidden_states=hidden_states, forward_batch=forward_batch, input_ids=input_ids, input_ids_global=input_ids_global, + prev_residual=prev_residual, + prev_post=prev_post, + prev_comb=prev_comb, ) + if use_fused and last_layer is not None: + hidden_states = last_layer.hc_post( + hidden_states, prev_residual, prev_post, prev_comb + ) # CP all-gather only on the last PP rank; PP IPC carries CP-split tensors. if self.pp_group.is_last_rank and dsa_use_prefill_cp(forward_batch): @@ -1589,6 +1802,7 @@ class DeepseekV4ForCausalLM(nn.Module): and not self_attn.indexer.compressor.ape_converted ): self_attn.indexer.compressor.apply_ape_hotfix() + layer.refresh_mhc_norm_weight_cache() @staticmethod def remap_weight_name_to_dpsk_hf_format( diff --git a/python/sglang/srt/models/deepseek_v4_nextn.py b/python/sglang/srt/models/deepseek_v4_nextn.py index ba71a0f2d..f116b8446 100644 --- a/python/sglang/srt/models/deepseek_v4_nextn.py +++ b/python/sglang/srt/models/deepseek_v4_nextn.py @@ -170,13 +170,17 @@ class DeepseekV4ModelNextN(nn.Module): hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states) positions = cp_split_and_rebuild_position(forward_batch, positions) - hidden_states = self.decoder( + hidden_states, residual, post, comb = self.decoder( positions=positions, hidden_states=hidden_states, forward_batch=forward_batch, input_ids=input_ids, input_ids_global=input_ids_global, ) + if residual is not None: + # NextN has a single decoder layer, so no later layer can consume a + # deferred fused hc_post state. + hidden_states = self.decoder.hc_post(hidden_states, residual, post, comb) if dsa_use_prefill_cp(forward_batch): hidden_states = cp_all_gather_rerange_output( diff --git a/tests/kernels/test_mhc_kernels.py b/tests/kernels/test_mhc_kernels.py new file mode 100644 index 000000000..989c61ac1 --- /dev/null +++ b/tests/kernels/test_mhc_kernels.py @@ -0,0 +1,111 @@ +import pytest +import torch + +import sglang.srt.layers.mhc as mhc +from sglang.srt.layers.mhc import mhc_fused_post_pre, mhc_post, mhc_pre + + +@pytest.mark.parametrize("hidden_size", [4096, 7168]) +@pytest.mark.parametrize("num_tokens", [0, 1, 8, 17, 32, 64]) +@pytest.mark.parametrize("use_norm", [False, True]) +def test_mhc_fused_post_pre_matches_unfused( + monkeypatch, hidden_size, num_tokens, use_norm +): + if not torch.cuda.is_available(): + pytest.skip("CUDA is required for TileLang mHC kernels") + + monkeypatch.setattr(mhc, "is_dsa_prefill_cp_round_robin_split", lambda: False) + torch.manual_seed(0) + device = torch.device("cuda") + hc_mult = 4 + hc_mult3 = hc_mult * 2 + hc_mult * hc_mult + hc_hidden_size = hc_mult * hidden_size + + x = torch.randn(num_tokens, hidden_size, device=device, dtype=torch.bfloat16) * 0.1 + residual = ( + torch.randn( + num_tokens, hc_mult, hidden_size, device=device, dtype=torch.bfloat16 + ) + * 0.1 + ) + post_prev = torch.rand(num_tokens, hc_mult, 1, device=device, dtype=torch.float32) + comb_prev = ( + torch.rand(num_tokens, hc_mult, hc_mult, device=device, dtype=torch.float32) + * 0.25 + ) + fn = ( + torch.randn(hc_mult3, hc_hidden_size, device=device, dtype=torch.float32) * 0.01 + ) + hc_scale = torch.tensor([0.5, 0.25, 0.25], device=device, dtype=torch.float32) + hc_base = torch.zeros(hc_mult3, device=device, dtype=torch.float32) + norm_weight = ( + torch.ones(hidden_size, device=device, dtype=torch.bfloat16) + if use_norm + else None + ) + norm_eps = 1e-6 if use_norm else None + + rms_eps = 1e-6 + hc_eps = 1e-6 + sinkhorn_repeat = 2 + + residual_ref = post_ref = comb_ref = layer_ref = None + if num_tokens > 0: + residual_ref = mhc_post(x, residual, post_prev, comb_prev) + post_ref, comb_ref, layer_ref = mhc_pre( + residual_ref, + fn, + hc_scale, + hc_base, + rms_eps, + hc_eps, + hc_eps, + 2.0, + sinkhorn_repeat, + norm_weight=norm_weight, + norm_eps=norm_eps, + ) + residual_out, post_out, comb_out, layer_out = mhc_fused_post_pre( + x, + residual, + post_prev, + comb_prev, + fn, + hc_scale, + hc_base, + rms_eps, + hc_eps, + hc_eps, + 2.0, + sinkhorn_repeat, + norm_weight=norm_weight, + norm_eps=norm_eps, + ) + + torch.cuda.synchronize() + if num_tokens == 0: + assert residual_out.shape == residual.shape + assert post_out.shape == (0, hc_mult, 1) + assert comb_out.shape == (0, hc_mult, hc_mult) + assert layer_out.shape == (0, hidden_size) + assert residual_out.dtype == torch.bfloat16 + assert post_out.dtype == torch.float32 + assert comb_out.dtype == torch.float32 + assert layer_out.dtype == torch.bfloat16 + return + + assert residual_ref is not None + assert post_ref is not None + assert comb_ref is not None + assert layer_ref is not None + assert residual_out.shape == residual_ref.shape + assert post_out.shape == post_ref.shape + assert comb_out.shape == comb_ref.shape + assert layer_out.shape == layer_ref.shape + + torch.testing.assert_close(residual_out, residual_ref, atol=0, rtol=0) + torch.testing.assert_close(post_out, post_ref, atol=1e-3, rtol=1e-3) + torch.testing.assert_close(comb_out, comb_ref, atol=1e-3, rtol=1e-3) + layer_atol = 2e-2 if use_norm else 2e-3 + layer_rtol = 2e-2 if use_norm else 2e-3 + torch.testing.assert_close(layer_out, layer_ref, atol=layer_atol, rtol=layer_rtol)