diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index e98136c1f..0599fcd9a 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -8,7 +8,6 @@ from typing import ( TYPE_CHECKING, Iterable, List, - Literal, Optional, Set, Tuple, @@ -158,6 +157,11 @@ logger = logging.getLogger(__name__) _FP8_WO_A_GEMM = envs.SGLANG_OPT_FP8_WO_A_GEMM.get() _MHC_POST_MULT_VALUE = 2.0 +DEEPSEEK_V4_STACKED_PARAMS_MAPPING: List[Tuple[str, str, int]] = [ + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), +] + def _is_fused_mhc_post_pre_enabled() -> bool: # The fused path directly reuses TileLang mhc_post/mhc_pre kernels and their @@ -199,6 +203,52 @@ def _fused_rmsnorm_fp8_quant(hidden_states, weight, eps): return x_quant, x_bf16 +def make_hc_mixing_params( + hc_mult: int, hidden_size: int +) -> Tuple[ + nn.Parameter, nn.Parameter, nn.Parameter, nn.Parameter, nn.Parameter, nn.Parameter +]: + mix_hc = (2 + hc_mult) * hc_mult + hc_dim = hc_mult * hidden_size + return ( + nn.Parameter(torch.empty(mix_hc, hc_dim, dtype=torch.float32)), + nn.Parameter(torch.empty(mix_hc, hc_dim, dtype=torch.float32)), + nn.Parameter(torch.empty(mix_hc, dtype=torch.float32)), + nn.Parameter(torch.empty(mix_hc, dtype=torch.float32)), + nn.Parameter(torch.empty(3, dtype=torch.float32)), + nn.Parameter(torch.empty(3, dtype=torch.float32)), + ) + + +def make_hc_head_params( + hc_mult: int, hidden_size: int +) -> Tuple[nn.Parameter, nn.Parameter, nn.Parameter]: + hc_dim = hc_mult * hidden_size + return ( + nn.Parameter(torch.empty(hc_mult, hc_dim, dtype=torch.float32)), + nn.Parameter(torch.empty(hc_mult, dtype=torch.float32)), + nn.Parameter(torch.empty(1, dtype=torch.float32)), + ) + + +def hc_head_torch( + x: torch.Tensor, + hc_fn: torch.Tensor, + hc_scale: torch.Tensor, + hc_base: torch.Tensor, + *, + norm_eps: float, + hc_eps: float, +) -> torch.Tensor: + shape, dtype = x.size(), x.dtype + x = x.flatten(-2).float() + rsqrt = torch.rsqrt(x.square().mean(-1, keepdim=True) + norm_eps) + mixes = F.linear(x, hc_fn) * rsqrt + pre = torch.sigmoid(mixes * hc_scale + hc_base) + hc_eps + y = torch.sum(pre.unsqueeze(-1) * x.view(shape), dim=-2) + return y.to(dtype) + + _FREQS_CIS_TO_COS_SIN: dict[ Tuple[int, torch.dtype, torch.device], Tuple[torch.Tensor, torch.Tensor] ] = {} @@ -284,137 +334,90 @@ bcg_deepseek_v4_attention_with_output = eager_on_graph(True)( ) -class MQALayer(nn.Module): +class MqaAttentionBase(nn.Module): + def __init__( self, config: DeepSeekV4Config, layer_id: int, - quant_config: Optional[QuantizationConfig] = None, - prefix: str = "", - alt_streams: Optional[List[torch.cuda.Stream]] = None, - compress_ratio_override: Optional[int] = None, + quant_config: Optional[QuantizationConfig], + prefix: str, + *, + attn_tp_rank: Optional[int] = None, + attn_tp_size: Optional[int] = None, + compress_ratio: Optional[int] = None, + fuse_wqa_wkv: Optional[bool] = None, + wo_a_fp8: Optional[bool] = None, + wo_a_keeps_quant_config: Optional[bool] = None, + wo_b_reduce_results: Optional[bool] = None, + rope_original_seq_len: Optional[int] = None, ) -> None: super().__init__() - self.tp_rank = attn_tp_rank = get_parallel().attn_tp_rank - self.tp_size = attn_tp_size = get_parallel().attn_tp_size self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() - if self.dsa_enable_prefill_cp: - self.cp_size = get_parallel().attn_cp_size - self.tp_rank = attn_tp_rank = 0 - self.tp_size = attn_tp_size = 1 + if attn_tp_rank is None or attn_tp_size is None: + attn_tp_rank = get_parallel().attn_tp_rank + attn_tp_size = get_parallel().attn_tp_size + if self.dsa_enable_prefill_cp: + self.cp_size = get_parallel().attn_cp_size + attn_tp_rank, attn_tp_size = 0, 1 + self.attn_tp_rank: int = attn_tp_rank + self.attn_tp_size: int = attn_tp_size + self.layer_id = layer_id self.dim = config.hidden_size + self.hidden_size = config.hidden_size self.qk_rope_head_dim = config.qk_rope_head_dim self.qk_nope_head_dim = config.head_dim - config.qk_rope_head_dim self.head_dim = self.qk_rope_head_dim + self.qk_nope_head_dim - self.n_heads = config.num_attention_heads - self.n_local_heads = self.n_heads // attn_tp_size - self.n_groups = config.o_groups - self.n_local_groups = self.n_groups // attn_tp_size self.rope_head_dim = config.qk_rope_head_dim - self.softmax_scale = self.head_dim**-0.5 - self.hidden_size = config.hidden_size + self.n_heads = config.num_attention_heads + self.n_local_heads = self.n_heads // self.attn_tp_size + self.n_groups = config.o_groups + self.n_local_groups = self.n_groups // self.attn_tp_size self.q_lora_rank = config.q_lora_rank self.o_lora_rank = config.o_lora_rank self.eps = config.rms_norm_eps - compress_ratio = ( - compress_ratio_override - if compress_ratio_override is not None + self.softmax_scale = self.head_dim**-0.5 + + self.compress_ratio: int = ( + compress_ratio + if compress_ratio is not None else config.compress_ratios[layer_id] ) - - assert compress_ratio in ( + assert self.compress_ratio in ( 0, 4, 128, - ), f"V4 compress_ratio: expected one of (0, 4, 128), got {compress_ratio}" - self.compress_ratio: Literal[0, 4, 128] = compress_ratio + ), f"V4 compress_ratio: expected one of (0, 4, 128), got {self.compress_ratio}" assert self.head_dim == config.head_dim assert config.num_key_value_heads == 1 - rope_theta, rope_scaling = get_rope_config(config) - if rope_scaling: - rope_scaling["rope_type"] = "deepseek_yarn" - - rope_base = config.compress_rope_theta if self.compress_ratio else rope_theta - - self.rotary_emb = get_rope_wrapper( - head_size=self.rope_head_dim, - rotary_dim=self.rope_head_dim, - max_position=config.max_position_embeddings, - base=rope_base, - rope_scaling=rope_scaling, - is_neox_style=False, - device=get_server_args().device, + fuse: bool = ( + envs.SGLANG_OPT_FUSE_WQA_WKV.get() if fuse_wqa_wkv is None else fuse_wqa_wkv ) - - from sglang.srt.layers.deepseek_v4_rope import precompute_freqs_cis - - # YARN correction applies to ALL layers (dense and compressed share the same - # YARN-corrected inv_freq); only the rope base differs (rope_theta vs compress_rope_theta). - original_seq_len = rope_scaling["original_max_position_embeddings"] - - freqs_cis = precompute_freqs_cis( - dim=self.qk_rope_head_dim, - seqlen=config.max_position_embeddings, - original_seq_len=original_seq_len, - base=rope_base, - factor=rope_scaling["factor"], - beta_fast=rope_scaling["beta_fast"], - beta_slow=rope_scaling["beta_slow"], + fp8: bool = _FP8_WO_A_GEMM if wo_a_fp8 is None else wo_a_fp8 + reduce_results: bool = ( + (self.attn_tp_size == get_parallel().tp_size and self.attn_tp_size > 1) + if wo_b_reduce_results is None + else wo_b_reduce_results ) - self.register_buffer("freqs_cis", freqs_cis, persistent=False) - self.freqs_cis: torch.Tensor - - if _is_hip: - cos_cache = freqs_cis.real.to(torch.bfloat16).unsqueeze(-2).unsqueeze(-2) - sin_cache = freqs_cis.imag.to(torch.bfloat16).unsqueeze(-2).unsqueeze(-2) - self.register_buffer("cos_cache", cos_cache, persistent=False) - self.register_buffer("sin_cache", sin_cache, persistent=False) - - if envs.SGLANG_OPT_USE_MULTI_STREAM_OVERLAP.get() and alt_streams is not None: - self.alt_streams = alt_streams[:3] - self.alt_streams_indexer = alt_streams[-2:] - else: - self.alt_streams = None - self.alt_streams_indexer = None - - from sglang.srt.utils import is_blackwell_supported - - self._multi_stream_bs_limit = 128 if is_blackwell_supported() else 64 - - self.compressor = None - self.indexer = None - if self.compress_ratio in (4, 128): - self.compressor = Compressor( - config, - layer_id=self.layer_id, - is_in_indexer=False, - freqs_cis=freqs_cis, - compress_ratio=self.compress_ratio, - head_dim=self.head_dim, - rotate=False, - prefix=add_prefix("compressor", prefix), - rotary_emb=getattr(self, "rotary_emb", None), + if wo_a_keeps_quant_config is None: + wo_a_quant_config: Optional[QuantizationConfig] = ( + quant_config if fp8 else None ) - if self.compress_ratio == 4: - self.indexer = C4Indexer( - config, - freqs_cis=freqs_cis, - layer_id=layer_id, - quant_config=quant_config, - prefix=add_prefix("indexer", prefix), - alt_streams=self.alt_streams_indexer, - rotary_emb=getattr(self, "rotary_emb", None), - ) + elif wo_a_keeps_quant_config: + wo_a_quant_config = quant_config + else: + wo_a_quant_config = None + + self.fuse_wqa_wkv = fuse self.attn_sink = nn.Parameter(torch.empty(self.n_heads, dtype=torch.float32)) self._attn_sink_local: Optional[torch.Tensor] = ( - self.attn_sink if attn_tp_size == 1 else None + self.attn_sink if self.attn_tp_size == 1 else None ) - self.fuse_wqa_wkv = envs.SGLANG_OPT_FUSE_WQA_WKV.get() - if self.fuse_wqa_wkv: + if fuse: self.wqkv_a = ReplicatedLinear( self.hidden_size, self.q_lora_rank + self.head_dim, @@ -444,21 +447,21 @@ class MQALayer(nn.Module): bias=False, quant_config=quant_config, prefix=add_prefix("wq_b", prefix), - tp_rank=attn_tp_rank, - tp_size=attn_tp_size, + tp_rank=self.attn_tp_rank, + tp_size=self.attn_tp_size, ) self.kv_norm = RMSNorm(self.head_dim, eps=self.eps) self.wo_a = ColumnParallelLinear( self.n_heads * self.head_dim // self.n_groups, self.n_groups * self.o_lora_rank, bias=False, - quant_config=quant_config if _FP8_WO_A_GEMM else None, + quant_config=wo_a_quant_config, prefix=add_prefix("wo_a", prefix), - tp_rank=attn_tp_rank, - tp_size=attn_tp_size, - **({} if _FP8_WO_A_GEMM else {"params_dtype": torch.bfloat16}), + tp_rank=self.attn_tp_rank, + tp_size=self.attn_tp_size, + **({} if fp8 else {"params_dtype": torch.bfloat16}), ) - if _FP8_WO_A_GEMM: + if fp8: assert hasattr( self.wo_a, "weight_scale_inv" ), "FP8 quant_config must create weight_scale_inv" @@ -468,12 +471,116 @@ class MQALayer(nn.Module): self.hidden_size, bias=False, quant_config=quant_config, - reduce_results=attn_tp_size == get_parallel().tp_size and attn_tp_size > 1, + reduce_results=reduce_results, prefix=add_prefix("wo_b", prefix), - tp_rank=attn_tp_rank, - tp_size=attn_tp_size, + tp_rank=self.attn_tp_rank, + tp_size=self.attn_tp_size, ) + from sglang.srt.layers.deepseek_v4_rope import precompute_freqs_cis + + rope_theta, rope_scaling = get_rope_config(config) + self.rope_scaling = rope_scaling + scaling = rope_scaling or {} + self.rope_base = ( + config.compress_rope_theta if self.compress_ratio else rope_theta + ) + original_seq_len: int = ( + rope_original_seq_len + if rope_original_seq_len is not None + else scaling["original_max_position_embeddings"] + ) + freqs_cis = precompute_freqs_cis( + dim=self.qk_rope_head_dim, + seqlen=config.max_position_embeddings, + original_seq_len=original_seq_len, + base=self.rope_base, + factor=scaling.get("factor", 1.0), + beta_fast=scaling.get("beta_fast", 32), + beta_slow=scaling.get("beta_slow", 1), + ) + self.register_buffer("freqs_cis", freqs_cis, persistent=False) + self.freqs_cis: torch.Tensor + + +class MQALayer(MqaAttentionBase): + def __init__( + self, + config: DeepSeekV4Config, + layer_id: int, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + alt_streams: Optional[List[torch.cuda.Stream]] = None, + compress_ratio_override: Optional[int] = None, + ) -> None: + super().__init__( + config, + layer_id, + quant_config, + prefix, + compress_ratio=compress_ratio_override, + ) + self.tp_rank = self.attn_tp_rank + self.tp_size = self.attn_tp_size + + if self.rope_scaling: + self.rope_scaling["rope_type"] = "deepseek_yarn" + self.rotary_emb = get_rope_wrapper( + head_size=self.rope_head_dim, + rotary_dim=self.rope_head_dim, + max_position=config.max_position_embeddings, + base=self.rope_base, + rope_scaling=self.rope_scaling, + is_neox_style=False, + device=get_server_args().device, + ) + + if _is_hip: + cos_cache = ( + self.freqs_cis.real.to(torch.bfloat16).unsqueeze(-2).unsqueeze(-2) + ) + sin_cache = ( + self.freqs_cis.imag.to(torch.bfloat16).unsqueeze(-2).unsqueeze(-2) + ) + self.register_buffer("cos_cache", cos_cache, persistent=False) + self.register_buffer("sin_cache", sin_cache, persistent=False) + + if envs.SGLANG_OPT_USE_MULTI_STREAM_OVERLAP.get() and alt_streams is not None: + self.alt_streams = alt_streams[:3] + self.alt_streams_indexer = alt_streams[-2:] + else: + self.alt_streams = None + self.alt_streams_indexer = None + + from sglang.srt.utils import is_blackwell_supported + + self._multi_stream_bs_limit = 128 if is_blackwell_supported() else 64 + + self.compressor = None + self.indexer = None + if self.compress_ratio in (4, 128): + self.compressor = Compressor( + config, + layer_id=self.layer_id, + is_in_indexer=False, + freqs_cis=self.freqs_cis, + compress_ratio=self.compress_ratio, + head_dim=self.head_dim, + rotate=False, + prefix=add_prefix("compressor", prefix), + rotary_emb=getattr(self, "rotary_emb", None), + ) + if self.compress_ratio == 4: + self.indexer = C4Indexer( + config, + freqs_cis=self.freqs_cis, + layer_id=layer_id, + quant_config=quant_config, + prefix=add_prefix("indexer", prefix), + alt_streams=self.alt_streams_indexer, + rotary_emb=getattr(self, "rotary_emb", None), + ) + self.attn_mqa = RadixAttention( self.n_local_heads, self.head_dim, @@ -1139,7 +1246,7 @@ class DeepseekV4DecoderLayer(nn.Module): self.config = config self.hidden_size = config.hidden_size self.layer_id = layer_id - self.self_attn = MQALayer( + self.self_attn = self._build_self_attn( config=config, layer_id=layer_id, quant_config=quant_config, @@ -1173,20 +1280,39 @@ class DeepseekV4DecoderLayer(nn.Module): self.hc_mult = hc_mult = config.hc_mult self.hc_sinkhorn_iters = config.hc_sinkhorn_iters self.hc_eps = config.hc_eps - mix_hc = (2 + hc_mult) * hc_mult - hc_dim = hc_mult * config.hidden_size - self.hc_attn_fn = nn.Parameter(torch.empty(mix_hc, hc_dim, dtype=torch.float32)) - self.hc_ffn_fn = nn.Parameter(torch.empty(mix_hc, hc_dim, dtype=torch.float32)) - self.hc_attn_base = nn.Parameter(torch.empty(mix_hc, dtype=torch.float32)) - self.hc_ffn_base = nn.Parameter(torch.empty(mix_hc, dtype=torch.float32)) - self.hc_attn_scale = nn.Parameter(torch.empty(3, dtype=torch.float32)) - self.hc_ffn_scale = nn.Parameter(torch.empty(3, dtype=torch.float32)) + ( + self.hc_attn_fn, + self.hc_ffn_fn, + self.hc_attn_base, + self.hc_ffn_base, + self.hc_attn_scale, + self.hc_ffn_scale, + ) = make_hc_mixing_params(hc_mult, config.hidden_size) 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 _build_self_attn( + self, + *, + config: DeepSeekV4Config, + layer_id: int, + quant_config: Optional[QuantizationConfig], + prefix: str, + alt_streams: Optional[List[torch.cuda.Stream]], + compress_ratio_override: Optional[int], + ) -> nn.Module: + return MQALayer( + config=config, + layer_id=layer_id, + quant_config=quant_config, + prefix=prefix, + alt_streams=alt_streams, + compress_ratio_override=compress_ratio_override, + ) + 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 = ( @@ -1473,6 +1599,29 @@ class DeepseekV4DecoderLayer(nn.Module): if not norm_fused: hidden_states = self.post_attention_layernorm(hidden_states) + hidden_states = self._run_moe_ffn_dp_sync( + hidden_states, + forward_batch, + input_ids=input_ids, + input_ids_global=input_ids_global, + ) + + if not use_fused: + hidden_states = self.hc_post(hidden_states, residual, post, comb) + return hidden_states, None, None, None + + # 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 + + def _run_moe_ffn_dp_sync( + self, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + *, + input_ids: torch.Tensor, + input_ids_global: torch.Tensor, + ) -> torch.Tensor: _use_cp = self.dsa_enable_prefill_cp and dsa_use_prefill_cp(forward_batch) _use_tp_moe_gather = ( not _use_cp @@ -1605,14 +1754,7 @@ class DeepseekV4DecoderLayer(nn.Module): gathered = [torch.empty_like(t) for t in _a2a_scatter_chunks] attn_tp_all_gather(gathered, hidden_states.contiguous()) hidden_states = torch.cat(gathered) - - if not use_fused: - hidden_states = self.hc_post(hidden_states, residual, post, comb) - return hidden_states, None, None, None - - # 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 + return hidden_states # ------------------------------------------------------------------ # TBO op decomposition (prefill two-batch-overlap, EP / mori path) @@ -1882,12 +2024,11 @@ class DeepseekV4Model(nn.Module): self.hc_mult = hc_mult = config.hc_mult self.norm_eps = config.rms_norm_eps if self.pp_group.is_last_rank: - hc_dim = hc_mult * config.hidden_size - self.hc_head_fn = nn.Parameter( - torch.empty(hc_mult, hc_dim, dtype=torch.float32) - ) - self.hc_head_base = nn.Parameter(torch.empty(hc_mult, dtype=torch.float32)) - self.hc_head_scale = nn.Parameter(torch.empty(1, dtype=torch.float32)) + ( + self.hc_head_fn, + self.hc_head_base, + self.hc_head_scale, + ) = make_hc_head_params(hc_mult, config.hidden_size) self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() self.use_fused_mhc_post_pre = _is_fused_mhc_post_pre_enabled() @@ -1912,13 +2053,14 @@ class DeepseekV4Model(nn.Module): norm_eps=self.norm_eps, hc_eps=self.hc_eps, ) - shape, dtype = x.size(), x.dtype - x = x.flatten(1).float() - rsqrt = torch.rsqrt(x.square().mean(-1, keepdim=True) + self.norm_eps) - mixes = F.linear(x, hc_fn) * rsqrt - pre = torch.sigmoid(mixes * hc_scale + hc_base) + self.hc_eps - y = torch.sum(pre.unsqueeze(-1) * x.view(shape), dim=1) - return y.to(dtype) + return hc_head_torch( + x, + hc_fn, + hc_scale, + hc_base, + norm_eps=self.norm_eps, + hc_eps=self.hc_eps, + ) def _can_run_tbo(self, forward_batch: ForwardBatch) -> bool: """DSV4 prefill-only two-batch-overlap gate. @@ -2475,10 +2617,7 @@ class DeepseekV4ForCausalLM(nn.Module): else: logger.info("Skip dequant fp8 wo_a") - stacked_params_mapping = [ - ("gate_up_proj", "gate_proj", 0), - ("gate_up_proj", "up_proj", 1), - ] + stacked_params_mapping = DEEPSEEK_V4_STACKED_PARAMS_MAPPING expert_params_mapping = FusedMoE.make_expert_params_mapping( ckpt_gate_proj_name="gate_proj",