From 776709efe89390599302a244f6ab1626f02ba9ff Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 27 Feb 2026 13:37:29 -0800 Subject: [PATCH] [3/n] deepseek_v2.py Refactor: Migrate MLA forward method in deepseek_v2.py (#19122) --- .../attention_backend_handler.py | 2 +- .../attention_forward_methods/__init__.py | 6 + .../forward_methods.py | 2 +- .../attention_forward_methods/forward_mla.py | 492 +++++++++++ .../forward_mla_fused_rope_cpu.py | 152 ++++ .../forward_mla_fused_rope_rocm.py | 227 +++++ python/sglang/srt/models/deepseek_v2.py | 833 +----------------- test/srt/cpu/test_qkv_proj_with_rope.py | 6 +- test/srt/cpu/test_rope.py | 4 +- 9 files changed, 906 insertions(+), 818 deletions(-) create mode 100644 python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py create mode 100644 python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_fused_rope_cpu.py create mode 100644 python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_fused_rope_rocm.py diff --git a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py index cc673a9ca..dfc1d4c97 100644 --- a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py +++ b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py @@ -25,7 +25,7 @@ class AttentionBackendRegistry: def _dispatch_mla_subtype(attn, forward_batch): if _is_hip: if attn.rocm_fused_decode_mla and forward_batch.forward_mode.is_decode(): - return AttnForwardMethod.MLA_FUSED_ROPE + return AttnForwardMethod.MLA_FUSED_ROPE_ROCM else: return AttnForwardMethod.MLA else: diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/__init__.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/__init__.py index e8076e2ad..2a9508fe2 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/__init__.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/__init__.py @@ -1,7 +1,13 @@ from .forward_methods import AttnForwardMethod from .forward_mha import DeepseekMHAForwardMixin +from .forward_mla import DeepseekMLAForwardMixin +from .forward_mla_fused_rope_cpu import DeepseekMLACpuForwardMixin +from .forward_mla_fused_rope_rocm import DeepseekMLARocmForwardMixin __all__ = [ "AttnForwardMethod", "DeepseekMHAForwardMixin", + "DeepseekMLACpuForwardMixin", + "DeepseekMLAForwardMixin", + "DeepseekMLARocmForwardMixin", ] diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_methods.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_methods.py index 839ed3a6f..c234a6b38 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_methods.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_methods.py @@ -17,7 +17,7 @@ class AttnForwardMethod(IntEnum): MHA_ONE_SHOT = auto() # Use MLA but with fused RoPE - MLA_FUSED_ROPE = auto() + MLA_FUSED_ROPE_ROCM = auto() # Use MLA with fused RoPE kernel for CPU MLA_FUSED_ROPE_CPU = auto() diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py new file mode 100644 index 000000000..c365d5c93 --- /dev/null +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py @@ -0,0 +1,492 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.srt.compilation.piecewise_context_manager import is_in_piecewise_cuda_graph +from sglang.srt.layers import deep_gemm_wrapper +from sglang.srt.layers.attention.nsa.utils import nsa_use_prefill_cp +from sglang.srt.layers.communicator import get_attn_tp_context +from sglang.srt.layers.quantization.fp8_kernel import ( + fp8_dtype, + per_tensor_quant_mla_fp8, + per_token_group_quant_mla_deep_gemm_masked_fp8, +) +from sglang.srt.model_executor.forward_batch_info import ForwardBatch +from sglang.srt.models.deepseek_common.utils import ( + FORWARD_ABSORB_CORE_ATTENTION_BACKENDS, + _is_cublas_ge_129, + _is_cuda, + _is_gfx95_supported, + _is_hip, + _use_aiter, + _use_aiter_gfx95, +) +from sglang.srt.server_args import get_global_server_args +from sglang.srt.utils import BumpAllocator + +if TYPE_CHECKING: + from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA + +if _is_cuda: + from sgl_kernel import bmm_fp8 + +if _use_aiter: + from aiter.ops.triton.batched_gemm_a8w8_a_per_token_group_prequant_w_per_batched_tensor_quant import ( + batched_gemm_a8w8_a_per_token_group_prequant_w_per_batched_tensor_quant, + ) +if _use_aiter_gfx95: + from aiter.ops.triton.fused_fp8_quant import ( + fused_flatten_fp8_group_quant, + fused_rms_fp8_group_quant, + ) + + from sglang.srt.layers.quantization.rocm_mxfp4_utils import ( + batched_gemm_afp4wfp4_pre_quant, + fused_flatten_mxfp4_quant, + fused_rms_mxfp4_quant, + ) + from sglang.srt.layers.rocm_linear_utils import fused_qk_rope_cat_and_cache_mla + + +class DeepseekMLAForwardMixin: + + def init_mla_forward(self: DeepseekV2AttentionMLA): + self.flashinfer_mla_disable_ragged = ( + get_global_server_args().flashinfer_mla_disable_ragged + ) + + def forward_absorb_prepare( + self: DeepseekV2AttentionMLA, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + zero_allocator: BumpAllocator, + llama_4_scaling: Optional[torch.Tensor] = None, + ): + from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode + + q_lora = None + topk_indices = None + if self.q_lora_rank is not None: + q, latent_cache = ( + get_attn_tp_context() + .fetch_qkv_latent() + .split( + [self.q_lora_rank, self.kv_lora_rank + self.qk_rope_head_dim], + dim=-1, + ) + ) + k_nope = latent_cache[..., : self.kv_lora_rank] + + # overlap qk norm + if self.alt_stream is not None and get_is_capture_mode(): + current_stream = torch.cuda.current_stream() + self.alt_stream.wait_stream(current_stream) + q = self.q_a_layernorm(q) + with torch.cuda.stream(self.alt_stream): + k_nope = self.kv_a_layernorm(k_nope) + current_stream.wait_stream(self.alt_stream) + else: + if _use_aiter_gfx95 and self.q_b_proj.weight.dtype == torch.uint8: + q, _, k_nope, *_ = fused_rms_mxfp4_quant( + q, + self.q_a_layernorm.weight, + self.q_a_layernorm.variance_epsilon, + k_nope, + self.kv_a_layernorm.weight, + self.kv_a_layernorm.variance_epsilon, + ) + else: + q_lora = None + if ( + _use_aiter_gfx95 + and self.q_b_proj.weight.dtype == torch.float8_e4m3fn + ): + if self.use_nsa: + q_quanted, q_lora, k_nope, _ = fused_rms_fp8_group_quant( + q, + self.q_a_layernorm.weight, + self.q_a_layernorm.variance_epsilon, + k_nope, + self.kv_a_layernorm.weight, + self.kv_a_layernorm.variance_epsilon, + group_size=128, + dtype_quant=torch.float8_e4m3fn, + res1=None, + output_unquantized_inp1=True, + ) + q = q_quanted + else: + q, _, k_nope, _ = fused_rms_fp8_group_quant( + q, + self.q_a_layernorm.weight, + self.q_a_layernorm.variance_epsilon, + k_nope, + self.kv_a_layernorm.weight, + self.kv_a_layernorm.variance_epsilon, + group_size=128, + dtype_quant=torch.float8_e4m3fn, + res1=None, + output_unquantized_inp1=False, + ) + + else: + q = self.q_a_layernorm(q) + k_nope = self.kv_a_layernorm(k_nope) + + # q_lora needed by indexer + if self.use_nsa: + if q_lora is None: + q_lora = q + + # overlap q_b_proj and indexer during decode + if ( + self.alt_stream is not None + and get_is_capture_mode() + and forward_batch.forward_mode.is_decode_or_idle() + and q_lora is not None + ): + current_stream = torch.cuda.current_stream() + self.alt_stream.wait_stream(current_stream) + with torch.cuda.stream(self.alt_stream): + k_nope = k_nope.unsqueeze(1) + q = self.q_b_proj(q)[0].view( + -1, self.num_local_heads, self.qk_head_dim + ) + topk_indices = self.indexer( + x=hidden_states, + q_lora=q_lora, + positions=positions, + forward_batch=forward_batch, + layer_id=self.layer_id, + ) + current_stream.wait_stream(self.alt_stream) + else: + k_nope = k_nope.unsqueeze(1) + q = self.q_b_proj(q)[0].view(-1, self.num_local_heads, self.qk_head_dim) + if q_lora is not None: + topk_indices = self.indexer( + x=hidden_states, + q_lora=q_lora, + positions=positions, + forward_batch=forward_batch, + layer_id=self.layer_id, + ) + else: + q = self.q_proj(hidden_states)[0].view( + -1, self.num_local_heads, self.qk_head_dim + ) + latent_cache = self.kv_a_proj_with_mqa(hidden_states)[0] + k_nope = latent_cache[..., : self.kv_lora_rank] + k_nope = self.kv_a_layernorm(k_nope).unsqueeze(1) + + q_nope, q_pe = q.split([self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) + k_pe = latent_cache[..., self.kv_lora_rank :].unsqueeze(1) + + if self.use_deep_gemm_bmm: + q_nope_val, q_nope_scale, masked_m, expected_m, aligned_m = ( + per_token_group_quant_mla_deep_gemm_masked_fp8(q_nope.transpose(0, 1)) + ) + q_nope_out = q_nope.new_empty( + (self.num_local_heads, aligned_m, self.kv_lora_rank) + ) + deep_gemm_wrapper.grouped_gemm_nt_f8f8bf16_masked( + (q_nope_val, q_nope_scale), + (self.w_kc, self.w_scale_k), + q_nope_out, + masked_m, + expected_m, + ) + q_nope_out = q_nope_out[:, :expected_m, :] + elif _is_hip: + # TODO(haishaw): add bmm_fp8 to ROCm + if _use_aiter_gfx95 and self.w_kc.dtype == torch.uint8: + x = q_nope.transpose(0, 1) + q_nope_out = torch.empty( + x.shape[0], + x.shape[1], + self.w_kc.shape[2], + device=x.device, + dtype=torch.bfloat16, + ) + batched_gemm_afp4wfp4_pre_quant( + x, + self.w_kc.transpose(-2, -1), + self.w_scale_k.transpose(-2, -1), + torch.bfloat16, + q_nope_out, + ) + else: + if (_use_aiter_gfx95 and self.w_kc.dtype == torch.float8_e4m3fn) or ( + get_is_capture_mode() and self.w_kc.dtype == torch.float8_e4m3fnuz + ): + # fp8 Triton kernel: always on gfx950, + # cudagraph-only on gfx942 (hides launch overhead) + q_nope_out = batched_gemm_a8w8_a_per_token_group_prequant_w_per_batched_tensor_quant( + X=q_nope, + WQ=self.w_kc.transpose(-1, -2), + w_scale=self.w_scale, + group_size=128, + YQ=None, # allocate (B, M, N) + transpose_bm=False, # (B, M, N) + transpose_bm_in=True, # (M, B, K) + dtype=torch.bfloat16, + ) + + else: + q_nope_out = torch.bmm( + q_nope.to(torch.bfloat16).transpose(0, 1), + self.w_kc.to(torch.bfloat16) * self.w_scale, + ) + + elif self.w_kc.dtype == torch.float8_e4m3fn: + # fix bmm_fp8 error under cublas12.9 caused by bumpallocator, detail in pr#11612 + q_nope_val, q_nope_scale = per_tensor_quant_mla_fp8( + q_nope.transpose(0, 1), + ( + torch.zeros((1,), dtype=torch.float32, device=q_nope.device) + if _is_cublas_ge_129 + else zero_allocator.allocate(1) + ), + ) + q_nope_out = bmm_fp8( + q_nope_val, self.w_kc, q_nope_scale, self.w_scale, torch.bfloat16 + ) + else: + q_nope_out = torch.bmm(q_nope.transpose(0, 1), self.w_kc) + + q_nope_out = q_nope_out.transpose(0, 1) + + if ( + self.rotary_emb is not None + and (not self._fuse_rope_for_trtllm_mla(forward_batch)) + and (not _use_aiter or not _is_gfx95_supported or self.use_nsa) + ): + q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe) + + if nsa_use_prefill_cp(forward_batch): + # support allgather+rerrange + k_nope, k_pe = self.rebuild_cp_kv_cache( + latent_cache, forward_batch, k_nope, k_pe + ) + + return ( + q_pe, + k_pe, + q_nope_out, + k_nope, + forward_batch, + zero_allocator, + positions, + topk_indices, + llama_4_scaling, + ) + + def forward_absorb_core( + self: DeepseekV2AttentionMLA, + q_pe, + k_pe, + q_nope_out, + k_nope, + forward_batch, + zero_allocator, + positions, + topk_indices, + llama_4_scaling, + ): + save_kv_cache = True + + if self.current_attention_backend in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS: + extra_args = {} + if self._fuse_rope_for_trtllm_mla(forward_batch): + extra_args = { + "cos_sin_cache": self.rotary_emb.cos_sin_cache, + "is_neox": self.rotary_emb.is_neox_style, + "llama_4_scaling": llama_4_scaling, + } + + attn_output = self.attn_mqa( + q_nope_out, + k_nope, + k_nope, + forward_batch, + q_rope=q_pe, + k_rope=k_pe, + **extra_args, + **(dict(topk_indices=topk_indices) if topk_indices is not None else {}), + ) + else: + if _use_aiter_gfx95: + cos = self.rotary_emb.cos_cache + sin = self.rotary_emb.sin_cache + + kv_cache_dtype = ( + fp8_dtype if self.kv_cache_dtype == "fp8_e4m3" else q_nope_out.dtype + ) + + q, _, _, k = fused_qk_rope_cat_and_cache_mla( + q_nope_out, + q_pe, + k_nope, + k_pe, + forward_batch.token_to_kv_pool.get_key_buffer( + self.attn_mqa.layer_id + ), + forward_batch.out_cache_loc, + positions, + cos, + sin, + self.attn_mqa.k_scale, + self.rotary_emb.is_neox_style, + q_out_dtype=kv_cache_dtype, + ) + + save_kv_cache = False + else: + q = torch.cat([q_nope_out, q_pe], dim=-1) + k = torch.cat([k_nope, k_pe], dim=-1) + + # Apply llama 4 scaling if provided + if llama_4_scaling is not None: + q *= llama_4_scaling + + attn_output = self.attn_mqa( + q, + k, + k_nope, + forward_batch, + save_kv_cache=save_kv_cache, + **(dict(topk_indices=topk_indices) if topk_indices is not None else {}), + ) + attn_output = attn_output.view(-1, self.num_local_heads, self.kv_lora_rank) + + if self.use_deep_gemm_bmm: + attn_output_val, attn_output_scale, masked_m, expected_m, aligned_m = ( + per_token_group_quant_mla_deep_gemm_masked_fp8( + attn_output.transpose(0, 1) + ) + ) + attn_bmm_output = attn_output.new_empty( + (self.num_local_heads, aligned_m, self.v_head_dim) + ) + deep_gemm_wrapper.grouped_gemm_nt_f8f8bf16_masked( + (attn_output_val, attn_output_scale), + (self.w_vc, self.w_scale_v), + attn_bmm_output, + masked_m, + expected_m, + ) + attn_bmm_output = ( + attn_bmm_output[:, :expected_m, :].transpose(0, 1).flatten(1, 2) + ) + elif _is_hip: + # TODO(haishaw): add bmm_fp8 to ROCm + if _use_aiter_gfx95 and self.w_vc.dtype == torch.uint8: + x = attn_output.transpose(0, 1) + attn_bmm_output = torch.empty( + x.shape[0], + x.shape[1], + self.w_vc.shape[2], + device=x.device, + dtype=torch.bfloat16, + ) + batched_gemm_afp4wfp4_pre_quant( + x, + self.w_vc.transpose(-2, -1), + self.w_scale_v.transpose(-2, -1), + torch.bfloat16, + attn_bmm_output, + ) + else: + if _use_aiter_gfx95 and self.w_kc.dtype == torch.float8_e4m3fn: + attn_bmm_output = batched_gemm_a8w8_a_per_token_group_prequant_w_per_batched_tensor_quant( + X=attn_output, + WQ=self.w_vc.transpose(-1, -2), + w_scale=self.w_scale, + group_size=128, + YQ=None, + transpose_bm=False, + transpose_bm_in=True, + dtype=torch.bfloat16, + ) + else: + attn_bmm_output = torch.bmm( + attn_output.to(torch.bfloat16).transpose(0, 1), + self.w_vc.to(torch.bfloat16) * self.w_scale, + ) + + if self.o_proj.weight.dtype == torch.uint8: + attn_bmm_output = attn_bmm_output.transpose(0, 1) + attn_bmm_output = fused_flatten_mxfp4_quant(attn_bmm_output) + elif self.o_proj.weight.dtype == torch.float8_e4m3fn: + attn_bmm_output = attn_bmm_output.transpose(0, 1) + attn_bmm_output = fused_flatten_fp8_group_quant( + attn_bmm_output, group_size=128, dtype_quant=torch.float8_e4m3fn + ) + else: + attn_bmm_output = attn_bmm_output.transpose(0, 1).flatten(1, 2) + + elif self.w_vc.dtype == torch.float8_e4m3fn: + attn_output_val, attn_output_scale = per_tensor_quant_mla_fp8( + attn_output.transpose(0, 1), + ( + torch.zeros((1,), dtype=torch.float32, device=attn_output.device) + if _is_cublas_ge_129 + else zero_allocator.allocate(1) + ), + ) + attn_bmm_output = bmm_fp8( + attn_output_val, + self.w_vc, + attn_output_scale, + self.w_scale, + torch.bfloat16, + ) + attn_bmm_output = attn_bmm_output.transpose(0, 1).flatten(1, 2) + else: + if is_in_piecewise_cuda_graph(): + # torch dynamo requires out= op was called where output tensor was non-contiguous + attn_bmm_output = ( + torch.bmm(attn_output.transpose(0, 1), self.w_vc) + .transpose(0, 1) + .flatten(1, 2) + ) + else: + attn_bmm_output = torch.empty( + (attn_output.shape[0], self.num_local_heads * self.v_head_dim), + dtype=attn_output.dtype, + device=attn_output.device, + ) + torch.bmm( + attn_output.transpose(0, 1), + self.w_vc, + out=attn_bmm_output.view( + -1, self.num_local_heads, self.v_head_dim + ).transpose(0, 1), + ) + output, _ = self.o_proj(attn_bmm_output) + + return output + + def _fuse_rope_for_trtllm_mla( + self: DeepseekV2AttentionMLA, forward_batch: ForwardBatch + ) -> bool: + """ + Check if we should skip rope and do fused rope+quantize for TRTLLM MLA decode in fp8_e4m3 path. + """ + if self.current_attention_backend == "nsa": + return ( + get_global_server_args().nsa_decode_backend == "trtllm" + or get_global_server_args().nsa_prefill_backend == "trtllm" + ) and forward_batch.attn_backend.kv_cache_dtype == torch.float8_e4m3fn + + return ( + self.current_attention_backend == "trtllm_mla" + and ( + forward_batch.forward_mode.is_decode_or_idle() + or forward_batch.forward_mode.is_target_verify() + ) + and forward_batch.attn_backend.data_type == torch.float8_e4m3fn + ) diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_fused_rope_cpu.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_fused_rope_cpu.py new file mode 100644 index 000000000..43a5eed37 --- /dev/null +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_fused_rope_cpu.py @@ -0,0 +1,152 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch + +from sglang.srt.layers.amx_utils import PackWeightMethod +from sglang.srt.model_executor.forward_batch_info import ForwardBatch +from sglang.srt.models.deepseek_common.utils import ( + _is_cpu, + _is_cpu_amx_available, +) +from sglang.srt.utils import BumpAllocator, use_intel_amx_backend + +if TYPE_CHECKING: + from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA + + +class DeepseekMLACpuForwardMixin: + + def init_mla_fused_rope_cpu_forward(self: DeepseekV2AttentionMLA): + assert hasattr(self, "has_fused_proj") and hasattr(self, "is_packed_weight") + + # If we have self.fused_qkv_a_proj_with_mqa and we're running on CPU, we will choose the torch.ops.sgl_kernel.qkv_proj_with_rope_fused_weight kernel + # which requires self.w_kc and self.w_vc to be packed. + # If not, we will use torch.bmm and weight shouldn't be packed in this case + if self.has_fused_proj and _is_cpu and _is_cpu_amx_available: + self.quant_method = PackWeightMethod( + weight_names=["w_kc", "w_vc"], transpose_dims=[[1, 2], [1, 2]] + ) + + self.qkv_proj_with_rope_is_int8 = ( + self.has_fused_proj + and not self.is_packed_weight + and self.fused_qkv_a_proj_with_mqa.weight.dtype == torch.int8 + ) + self.qkv_proj_with_rope_is_fp8 = ( + self.has_fused_proj + and not self.is_packed_weight + and self.fused_qkv_a_proj_with_mqa.weight.dtype == torch.float8_e4m3fn + ) + + self.weight_block_size = None + if self.qkv_proj_with_rope_is_fp8 and _is_cpu and _is_cpu_amx_available: + assert getattr( + self.fused_qkv_a_proj_with_mqa.quant_method, "block_quant", False + ) == getattr(self.q_b_proj.quant_method, "block_quant", False) + use_block_quant = getattr( + self.fused_qkv_a_proj_with_mqa.quant_method, "block_quant", False + ) + + if use_block_quant: + assert ( + self.fused_qkv_a_proj_with_mqa.quant_method.quant_config.weight_block_size + == self.q_b_proj.quant_method.quant_config.weight_block_size + ) + self.weight_block_size = ( + self.fused_qkv_a_proj_with_mqa.quant_method.quant_config.weight_block_size + ) + + def forward_absorb_fused_mla_rope_cpu_prepare( + self: DeepseekV2AttentionMLA, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + zero_allocator: BumpAllocator, + ): + assert self.q_lora_rank is not None and use_intel_amx_backend( + self + ), "forward_absorb_fused_mla_rope_cpu_prepare requires q_lora_rank is not None and use_intel_amx_backend" + + q_input, k_input, v_input = ( + torch.ops.sgl_kernel.qkv_proj_with_rope_fused_weight( + hidden_states, + self.fused_qkv_a_proj_with_mqa.weight, + self.q_b_proj.weight, + self.w_kc, + self.q_a_layernorm.weight, + self.kv_a_layernorm.weight, + positions, + self.rotary_emb.cos_sin_cache, + self.kv_a_layernorm.variance_epsilon, + self.qkv_proj_with_rope_is_int8, + self.qkv_proj_with_rope_is_fp8, + ( + self.fused_qkv_a_proj_with_mqa.weight_scale + if self.qkv_proj_with_rope_is_int8 + else ( + self.fused_qkv_a_proj_with_mqa.weight_scale_inv + if self.qkv_proj_with_rope_is_fp8 + else None + ) + ), + ( + self.q_b_proj.weight_scale + if self.qkv_proj_with_rope_is_int8 + else ( + self.q_b_proj.weight_scale_inv + if self.qkv_proj_with_rope_is_fp8 + else None + ) + ), + True, # is_vnni + self.weight_block_size, + self.q_lora_rank, + self.kv_lora_rank, + self.qk_rope_head_dim, + ) + ) + return (q_input, k_input, v_input, forward_batch, zero_allocator) + + def forward_absorb_fused_mla_rope_cpu_core( + self: DeepseekV2AttentionMLA, + q_input, + k_input, + v_input, + forward_batch, + zero_allocator, + ): + assert self.q_lora_rank is not None and use_intel_amx_backend( + self + ), "forward_absorb_fused_mla_rope_cpu_core requires q_lora_rank is not None and use_intel_amx_backend" + + attn_output = self.attn_mqa(q_input, k_input, v_input, forward_batch) + attn_output = attn_output.view(-1, self.num_local_heads, self.kv_lora_rank) + + # [Note] Align shapes of bmm inputs. + # Shapes of inputs: + # q_nope: [M, B, K] + # original self.w_kc: [B, K, N] + # current self.w_kc (which has been converted in PackWeightMethod): [B, N, K] + + # Shapes of inputs to sgl_kernel.cpu.bmm: + # out: [B, M, N] + # mat1: [B, M, K] + # mat2: [B, N, K] + B = self.w_vc.size(0) + N = self.w_vc.size(1) + M = attn_output.size(0) + output = torch.empty([M, int(B * N)], dtype=attn_output.dtype) + attn_bmm_output = output.view([M, B, N]).transpose_(0, 1) + torch.ops.sgl_kernel.bmm_cpu( + attn_bmm_output, + attn_output.transpose(0, 1), + self.w_vc, + True, # is_vnni + None, # scale + ) + attn_output = output + output, _ = self.o_proj(attn_output) + + return output diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_fused_rope_rocm.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_fused_rope_rocm.py new file mode 100644 index 000000000..8868897af --- /dev/null +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_fused_rope_rocm.py @@ -0,0 +1,227 @@ +from __future__ import annotations + +import os +from typing import TYPE_CHECKING + +import torch + +from sglang.srt.layers.quantization.fp8_kernel import per_tensor_quant_mla_fp8 +from sglang.srt.model_executor.forward_batch_info import ForwardBatch +from sglang.srt.models.deepseek_common.utils import ( + _is_cuda, + _is_hip, +) +from sglang.srt.utils import BumpAllocator, get_bool_env_var + +if TYPE_CHECKING: + from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA + +if _is_cuda: + from sgl_kernel import bmm_fp8 + +if _is_hip: + from sglang.srt.layers.attention.triton_ops.rocm_mla_decode_rope import ( + decode_attention_fwd_grouped_rope, + ) + + +class DeepseekMLARocmForwardMixin: + + def init_mla_fused_rope_rocm_forward(self: DeepseekV2AttentionMLA): + self.rocm_fused_decode_mla = get_bool_env_var( + "SGLANG_ROCM_FUSED_DECODE_MLA", "false" + ) + + def forward_absorb_fused_mla_rope_prepare( + self: DeepseekV2AttentionMLA, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + zero_allocator: BumpAllocator, + ): + enable_rope_fusion = ( + os.getenv("SGLANG_FUSED_MLA_ENABLE_ROPE_FUSION", "1") == "1" + ) + # NOTE: hidden_states can be a tuple for some quantization paths. + # For shape/device/dtype, use the first tensor; still pass the original + # hidden_states through linear ops which may accept tuple inputs. + hidden_states_tensor = ( + hidden_states[0] if isinstance(hidden_states, tuple) else hidden_states + ) + + q_len = hidden_states_tensor.shape[0] + q_input = hidden_states_tensor.new_empty( + q_len, self.num_local_heads, self.kv_lora_rank + self.qk_rope_head_dim + ) + if self.q_lora_rank is not None: + q, latent_cache = self.fused_qkv_a_proj_with_mqa(hidden_states)[0].split( + [self.q_lora_rank, self.kv_lora_rank + self.qk_rope_head_dim], dim=-1 + ) + q = self.q_a_layernorm(q) + q = self.q_b_proj(q)[0].view(-1, self.num_local_heads, self.qk_head_dim) + else: + q = self.q_proj(hidden_states)[0].view( + -1, self.num_local_heads, self.qk_head_dim + ) + latent_cache = self.kv_a_proj_with_mqa(hidden_states)[0] + q_nope, q_pe = q.split([self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) + + if _is_hip: + # TODO(haishaw): add bmm_fp8 to ROCm + q_nope_out = torch.bmm( + q_nope.to(torch.bfloat16).transpose(0, 1), + self.w_kc.to(torch.bfloat16) * self.w_scale, + ) + elif self.w_kc.dtype == torch.float8_e4m3fn: + q_nope_val, q_nope_scale = per_tensor_quant_mla_fp8( + q_nope.transpose(0, 1), + zero_allocator.allocate(1), + dtype=torch.float8_e4m3fn, + ) + q_nope_out = bmm_fp8( + q_nope_val, self.w_kc, q_nope_scale, self.w_scale, torch.bfloat16 + ) + else: + q_nope_out = torch.bmm(q_nope.transpose(0, 1), self.w_kc) + q_input[..., : self.kv_lora_rank] = q_nope_out.transpose(0, 1) + v_input = latent_cache[..., : self.kv_lora_rank] + v_input = self.kv_a_layernorm(v_input.contiguous()).unsqueeze(1) + k_input = latent_cache.unsqueeze(1) + k_input[..., : self.kv_lora_rank] = v_input + + if not enable_rope_fusion: + k_pe = k_input[..., self.kv_lora_rank :] + q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe) + q_input[..., self.kv_lora_rank :] = q_pe + k_input[..., self.kv_lora_rank :] = k_pe + k_pe_output = None + else: + k_pe_output = torch.empty_like(k_input[..., self.kv_lora_rank :]) + + q_input[..., self.kv_lora_rank :] = q_pe + + # attn_output = self.attn_mqa(q_input, k_input, v_input, forward_batch) + # Use Fused ROPE with use_rope=OFF. + attn_output = torch.empty( + (q_len, self.num_local_heads, self.kv_lora_rank), + dtype=q.dtype, + device=q.device, + ) + attn_logits, _, kv_indptr, kv_indices, _, _, _ = ( + forward_batch.attn_backend.forward_metadata + ) + cos_sin_cache = self.rotary_emb.cos_sin_cache + num_kv_split = forward_batch.attn_backend.num_kv_splits + sm_scale = self.attn_mqa.scaling + if attn_logits is None: + attn_logits = torch.empty( + ( + forward_batch.batch_size, + self.num_local_heads, + num_kv_split, + self.kv_lora_rank + 1, + ), + dtype=torch.float32, + device=q.device, + ) + + # save current latent cache. + forward_batch.token_to_kv_pool.set_kv_buffer( + self.attn_mqa, forward_batch.out_cache_loc, k_input, None + ) + key_cache_buf = forward_batch.token_to_kv_pool.get_key_buffer( + self.attn_mqa.layer_id + ) + val_cache_buf = key_cache_buf[..., : self.kv_lora_rank] + + return ( + q_input, + key_cache_buf, + val_cache_buf, + attn_output, + kv_indptr, + kv_indices, + k_pe_output, + cos_sin_cache, + positions, + attn_logits, + num_kv_split, + sm_scale, + enable_rope_fusion, + k_input, + forward_batch, + zero_allocator, + ) + + def forward_absorb_fused_mla_rope_core( + self: DeepseekV2AttentionMLA, + q_input, + key_cache_buf, + val_cache_buf, + attn_output, + kv_indptr, + kv_indices, + k_pe_output, + cos_sin_cache, + positions, + attn_logits, + num_kv_split, + sm_scale, + enable_rope_fusion, + k_input, + forward_batch, + zero_allocator, + ): + decode_attention_fwd_grouped_rope( + q_input, + key_cache_buf, + val_cache_buf, + attn_output, + kv_indptr, + kv_indices, + k_pe_output, + self.kv_lora_rank, + self.rotary_emb.rotary_dim, + cos_sin_cache, + positions, + attn_logits, + num_kv_split, + sm_scale, + logit_cap=self.attn_mqa.logit_cap, + use_rope=enable_rope_fusion, + is_neox_style=self.rotary_emb.is_neox_style, + ) + + if enable_rope_fusion: + k_input[..., self.kv_lora_rank :] = k_pe_output + forward_batch.token_to_kv_pool.set_kv_buffer( + self.attn_mqa, forward_batch.out_cache_loc, k_input, None + ) + + attn_output = attn_output.view(-1, self.num_local_heads, self.kv_lora_rank) + + if _is_hip: + # TODO(haishaw): add bmm_fp8 to ROCm + attn_bmm_output = torch.bmm( + attn_output.to(torch.bfloat16).transpose(0, 1), + self.w_vc.to(torch.bfloat16) * self.w_scale, + ) + elif self.w_vc.dtype == torch.float8_e4m3fn: + attn_output_val, attn_output_scale = per_tensor_quant_mla_fp8( + attn_output.transpose(0, 1), + zero_allocator.allocate(1), + dtype=torch.float8_e4m3fn, + ) + attn_bmm_output = bmm_fp8( + attn_output_val, + self.w_vc, + attn_output_scale, + self.w_scale, + torch.bfloat16, + ) + else: + attn_bmm_output = torch.bmm(attn_output.transpose(0, 1), self.w_vc) + attn_output = attn_bmm_output.transpose(0, 1).flatten(1, 2) + output, _ = self.o_proj(attn_output) + + return output diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 5348d6a27..eb36952ca 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -19,7 +19,6 @@ from __future__ import annotations import logging -import os from contextlib import nullcontext from typing import Any, Dict, Iterable, List, Optional, Tuple, Union @@ -33,7 +32,6 @@ from sglang.srt.batch_overlap.two_batch_overlap import ( MaybeTboDeepEPDispatcher, model_forward_maybe_tbo, ) -from sglang.srt.compilation.piecewise_context_manager import is_in_piecewise_cuda_graph from sglang.srt.configs.model_config import ( get_nsa_index_head_dim, get_nsa_index_n_heads, @@ -107,11 +105,6 @@ from sglang.srt.layers.moe.utils import ( ) from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.fp8 import Fp8Config -from sglang.srt.layers.quantization.fp8_kernel import ( - fp8_dtype, - per_tensor_quant_mla_fp8, - per_token_group_quant_mla_deep_gemm_masked_fp8, -) from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.rotary_embedding import get_rope_wrapper from sglang.srt.layers.utils import PPMissingLayer @@ -127,17 +120,18 @@ from sglang.srt.models.deepseek_common.attention_backend_handler import ( from sglang.srt.models.deepseek_common.attention_forward_methods import ( AttnForwardMethod, DeepseekMHAForwardMixin, + DeepseekMLACpuForwardMixin, + DeepseekMLAForwardMixin, + DeepseekMLARocmForwardMixin, ) from sglang.srt.models.deepseek_common.deepseek_weight_loader import ( DeepseekV2WeightLoaderMixin, ) from sglang.srt.models.deepseek_common.utils import ( - FORWARD_ABSORB_CORE_ATTENTION_BACKENDS, _device_sm, _get_llama_4_scaling, _is_cpu, _is_cpu_amx_available, - _is_cublas_ge_129, _is_cuda, _is_gfx95_supported, _is_hip, @@ -152,45 +146,23 @@ from sglang.srt.utils import ( BumpAllocator, LazyValue, add_prefix, - get_bool_env_var, is_non_idle_and_non_empty, log_info_on_rank0, make_layers, use_intel_amx_backend, ) -if _use_aiter: - from aiter.ops.triton.batched_gemm_a8w8_a_per_token_group_prequant_w_per_batched_tensor_quant import ( - batched_gemm_a8w8_a_per_token_group_prequant_w_per_batched_tensor_quant, - ) if _use_aiter_gfx95: - from aiter.ops.triton.fused_fp8_quant import ( - fused_flatten_fp8_group_quant, - fused_rms_fp8_group_quant, - ) - - from sglang.srt.layers.quantization.rocm_mxfp4_utils import ( - batched_gemm_afp4wfp4_pre_quant, - fused_flatten_mxfp4_quant, - fused_rms_mxfp4_quant, - ) from sglang.srt.layers.rocm_linear_utils import ( aiter_dsv3_router_gemm, - fused_qk_rope_cat_and_cache_mla, get_dsv3_gemm_output_zero_allocator_size, ) if _use_aiter: - from sglang.srt.layers.rocm_linear_utils import fused_qk_rope_cat_and_cache_mla + pass if _is_cuda: - from sgl_kernel import bmm_fp8, dsv3_fused_a_gemm, dsv3_router_gemm -elif _is_cpu and _is_cpu_amx_available: - pass -elif _is_hip: - from sglang.srt.layers.attention.triton_ops.rocm_mla_decode_rope import ( - decode_attention_fwd_grouped_rope, - ) + from sgl_kernel import dsv3_fused_a_gemm, dsv3_router_gemm elif _is_npu: from sglang.srt.hardware_backend.npu.modules.deepseek_v2_attention_mla_npu import ( forward_dsa_core_npu, @@ -206,16 +178,6 @@ else: logger = logging.getLogger(__name__) -FORWARD_ABSORB_CORE_ATTENTION_BACKENDS = [ - "fa3", - "nsa", - "flashinfer", - "cutlass_mla", - "trtllm_mla", - "ascend", -] - - class DeepseekV2MLP(nn.Module): def __init__( self, @@ -1077,7 +1039,13 @@ class DeepseekV2MoE(nn.Module): state.hidden_states_mlp_output = final_hidden_states -class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): +class DeepseekV2AttentionMLA( + nn.Module, + DeepseekMHAForwardMixin, + DeepseekMLAForwardMixin, + DeepseekMLARocmForwardMixin, + DeepseekMLACpuForwardMixin, +): def __init__( self, @@ -1266,35 +1234,20 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): self.w_scale_v = None self.use_deep_gemm_bmm = False - self.flashinfer_mla_disable_ragged = ( - get_global_server_args().flashinfer_mla_disable_ragged - ) - self.current_attention_backend = ( None # Attention backend used by current forward batch ) - self.rocm_fused_decode_mla = get_bool_env_var( - "SGLANG_ROCM_FUSED_DECODE_MLA", "false" - ) - # If we have self.fused_qkv_a_proj_with_mqa and we're running on CPU, we will choose the torch.ops.sgl_kernel.qkv_proj_with_rope_fused_weight kernel - # which requires self.w_kc and self.w_vc to be packed. - # If not, we will use torch.bmm and weight shouldn't be packed in this case - has_fused_proj = hasattr(self, "fused_qkv_a_proj_with_mqa") - if has_fused_proj and _is_cpu and _is_cpu_amx_available: - self.quant_method = PackWeightMethod( - weight_names=["w_kc", "w_vc"], transpose_dims=[[1, 2], [1, 2]] - ) - - is_packed_weight = ( - has_fused_proj + self.has_fused_proj = hasattr(self, "fused_qkv_a_proj_with_mqa") + self.is_packed_weight = ( + self.has_fused_proj and hasattr(self.fused_qkv_a_proj_with_mqa.quant_method, "quant_config") and self.fused_qkv_a_proj_with_mqa.quant_method.quant_config.get_name() in {"awq", "awq_marlin", "moe_wna16"} ) self.use_min_latency_fused_a_gemm = ( - has_fused_proj - and not is_packed_weight + self.has_fused_proj + and not self.is_packed_weight and self.fused_qkv_a_proj_with_mqa.weight.dtype == torch.bfloat16 and self.fused_qkv_a_proj_with_mqa.weight.shape[0] == 2112 and self.fused_qkv_a_proj_with_mqa.weight.shape[1] == 7168 @@ -1302,36 +1255,10 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): and 90 <= _device_sm < 120 ) - self.qkv_proj_with_rope_is_int8 = ( - has_fused_proj - and not is_packed_weight - and self.fused_qkv_a_proj_with_mqa.weight.dtype == torch.int8 - ) - self.qkv_proj_with_rope_is_fp8 = ( - has_fused_proj - and not is_packed_weight - and self.fused_qkv_a_proj_with_mqa.weight.dtype == torch.float8_e4m3fn - ) - - self.weight_block_size = None - if self.qkv_proj_with_rope_is_fp8 and _is_cpu and _is_cpu_amx_available: - assert getattr( - self.fused_qkv_a_proj_with_mqa.quant_method, "block_quant", False - ) == getattr(self.q_b_proj.quant_method, "block_quant", False) - use_block_quant = getattr( - self.fused_qkv_a_proj_with_mqa.quant_method, "block_quant", False - ) - - if use_block_quant: - assert ( - self.fused_qkv_a_proj_with_mqa.quant_method.quant_config.weight_block_size - == self.q_b_proj.quant_method.quant_config.weight_block_size - ) - self.weight_block_size = ( - self.fused_qkv_a_proj_with_mqa.quant_method.quant_config.weight_block_size - ) - self.init_mha_forward() + self.init_mla_forward() + self.init_mla_fused_rope_rocm_forward() + self.init_mla_fused_rope_cpu_forward() def dispatch_attn_forward_method( self, forward_batch: ForwardBatch @@ -1433,7 +1360,7 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): inner_state = self.forward_absorb_prepare( positions, hidden_states, forward_batch, zero_allocator, llama_4_scaling ) - elif attn_forward_method == AttnForwardMethod.MLA_FUSED_ROPE: + elif attn_forward_method == AttnForwardMethod.MLA_FUSED_ROPE_ROCM: inner_state = self.forward_absorb_fused_mla_rope_prepare( positions, hidden_states, forward_batch, zero_allocator ) @@ -1472,7 +1399,7 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): return self.forward_normal_one_shot_core(*inner_state) elif attn_forward_method == AttnForwardMethod.MLA: return self.forward_absorb_core(*inner_state) - elif attn_forward_method == AttnForwardMethod.MLA_FUSED_ROPE: + elif attn_forward_method == AttnForwardMethod.MLA_FUSED_ROPE_ROCM: return self.forward_absorb_fused_mla_rope_core(*inner_state) elif attn_forward_method == AttnForwardMethod.MLA_FUSED_ROPE_CPU: return self.forward_absorb_fused_mla_rope_cpu_core(*inner_state) @@ -1502,25 +1429,6 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): qkv_latent = self.fused_qkv_a_proj_with_mqa(hidden_states)[0] return qkv_latent - def _fuse_rope_for_trtllm_mla(self, forward_batch: ForwardBatch) -> bool: - """ - Check if we should skip rope and do fused rope+quantize for TRTLLM MLA decode in fp8_e4m3 path. - """ - if self.current_attention_backend == "nsa": - return ( - get_global_server_args().nsa_decode_backend == "trtllm" - or get_global_server_args().nsa_prefill_backend == "trtllm" - ) and forward_batch.attn_backend.kv_cache_dtype == torch.float8_e4m3fn - - return ( - self.current_attention_backend == "trtllm_mla" - and ( - forward_batch.forward_mode.is_decode_or_idle() - or forward_batch.forward_mode.is_target_verify() - ) - and forward_batch.attn_backend.data_type == torch.float8_e4m3fn - ) - def rebuild_cp_kv_cache(self, latent_cache, forward_batch, k_nope, k_pe): # support allgather+rerrange latent_cache[..., : self.kv_lora_rank] = k_nope.squeeze(1) @@ -1535,703 +1443,6 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): k_pe = latent_cache_output[..., self.kv_lora_rank :].unsqueeze(1) return k_nope, k_pe - def forward_absorb_prepare( - self, - positions: torch.Tensor, - hidden_states: torch.Tensor, - forward_batch: ForwardBatch, - zero_allocator: BumpAllocator, - llama_4_scaling: Optional[torch.Tensor] = None, - ): - q_lora = None - topk_indices = None - if self.q_lora_rank is not None: - q, latent_cache = ( - get_attn_tp_context() - .fetch_qkv_latent() - .split( - [self.q_lora_rank, self.kv_lora_rank + self.qk_rope_head_dim], - dim=-1, - ) - ) - k_nope = latent_cache[..., : self.kv_lora_rank] - - # overlap qk norm - if self.alt_stream is not None and get_is_capture_mode(): - current_stream = torch.cuda.current_stream() - self.alt_stream.wait_stream(current_stream) - q = self.q_a_layernorm(q) - with torch.cuda.stream(self.alt_stream): - k_nope = self.kv_a_layernorm(k_nope) - current_stream.wait_stream(self.alt_stream) - else: - if _use_aiter_gfx95 and self.q_b_proj.weight.dtype == torch.uint8: - q, _, k_nope, *_ = fused_rms_mxfp4_quant( - q, - self.q_a_layernorm.weight, - self.q_a_layernorm.variance_epsilon, - k_nope, - self.kv_a_layernorm.weight, - self.kv_a_layernorm.variance_epsilon, - ) - else: - q_lora = None - if ( - _use_aiter_gfx95 - and self.q_b_proj.weight.dtype == torch.float8_e4m3fn - ): - if self.use_nsa: - q_quanted, q_lora, k_nope, _ = fused_rms_fp8_group_quant( - q, - self.q_a_layernorm.weight, - self.q_a_layernorm.variance_epsilon, - k_nope, - self.kv_a_layernorm.weight, - self.kv_a_layernorm.variance_epsilon, - group_size=128, - dtype_quant=torch.float8_e4m3fn, - res1=None, - output_unquantized_inp1=True, - ) - q = q_quanted - else: - q, _, k_nope, _ = fused_rms_fp8_group_quant( - q, - self.q_a_layernorm.weight, - self.q_a_layernorm.variance_epsilon, - k_nope, - self.kv_a_layernorm.weight, - self.kv_a_layernorm.variance_epsilon, - group_size=128, - dtype_quant=torch.float8_e4m3fn, - res1=None, - output_unquantized_inp1=False, - ) - - else: - q = self.q_a_layernorm(q) - k_nope = self.kv_a_layernorm(k_nope) - - # q_lora needed by indexer - if self.use_nsa: - if q_lora is None: - q_lora = q - - # overlap q_b_proj and indexer during decode - if ( - self.alt_stream is not None - and get_is_capture_mode() - and forward_batch.forward_mode.is_decode_or_idle() - and q_lora is not None - ): - current_stream = torch.cuda.current_stream() - self.alt_stream.wait_stream(current_stream) - with torch.cuda.stream(self.alt_stream): - k_nope = k_nope.unsqueeze(1) - q = self.q_b_proj(q)[0].view( - -1, self.num_local_heads, self.qk_head_dim - ) - topk_indices = self.indexer( - x=hidden_states, - q_lora=q_lora, - positions=positions, - forward_batch=forward_batch, - layer_id=self.layer_id, - ) - current_stream.wait_stream(self.alt_stream) - else: - k_nope = k_nope.unsqueeze(1) - q = self.q_b_proj(q)[0].view(-1, self.num_local_heads, self.qk_head_dim) - if q_lora is not None: - topk_indices = self.indexer( - x=hidden_states, - q_lora=q_lora, - positions=positions, - forward_batch=forward_batch, - layer_id=self.layer_id, - ) - else: - q = self.q_proj(hidden_states)[0].view( - -1, self.num_local_heads, self.qk_head_dim - ) - latent_cache = self.kv_a_proj_with_mqa(hidden_states)[0] - k_nope = latent_cache[..., : self.kv_lora_rank] - k_nope = self.kv_a_layernorm(k_nope).unsqueeze(1) - - q_nope, q_pe = q.split([self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) - k_pe = latent_cache[..., self.kv_lora_rank :].unsqueeze(1) - - if self.use_deep_gemm_bmm: - q_nope_val, q_nope_scale, masked_m, expected_m, aligned_m = ( - per_token_group_quant_mla_deep_gemm_masked_fp8(q_nope.transpose(0, 1)) - ) - q_nope_out = q_nope.new_empty( - (self.num_local_heads, aligned_m, self.kv_lora_rank) - ) - deep_gemm_wrapper.grouped_gemm_nt_f8f8bf16_masked( - (q_nope_val, q_nope_scale), - (self.w_kc, self.w_scale_k), - q_nope_out, - masked_m, - expected_m, - ) - q_nope_out = q_nope_out[:, :expected_m, :] - elif _is_hip: - # TODO(haishaw): add bmm_fp8 to ROCm - if _use_aiter_gfx95 and self.w_kc.dtype == torch.uint8: - x = q_nope.transpose(0, 1) - q_nope_out = torch.empty( - x.shape[0], - x.shape[1], - self.w_kc.shape[2], - device=x.device, - dtype=torch.bfloat16, - ) - batched_gemm_afp4wfp4_pre_quant( - x, - self.w_kc.transpose(-2, -1), - self.w_scale_k.transpose(-2, -1), - torch.bfloat16, - q_nope_out, - ) - else: - if (_use_aiter_gfx95 and self.w_kc.dtype == torch.float8_e4m3fn) or ( - get_is_capture_mode() and self.w_kc.dtype == torch.float8_e4m3fnuz - ): - # fp8 Triton kernel: always on gfx950, - # cudagraph-only on gfx942 (hides launch overhead) - q_nope_out = batched_gemm_a8w8_a_per_token_group_prequant_w_per_batched_tensor_quant( - X=q_nope, - WQ=self.w_kc.transpose(-1, -2), - w_scale=self.w_scale, - group_size=128, - YQ=None, # allocate (B, M, N) - transpose_bm=False, # (B, M, N) - transpose_bm_in=True, # (M, B, K) - dtype=torch.bfloat16, - ) - - else: - q_nope_out = torch.bmm( - q_nope.to(torch.bfloat16).transpose(0, 1), - self.w_kc.to(torch.bfloat16) * self.w_scale, - ) - - elif self.w_kc.dtype == torch.float8_e4m3fn: - # fix bmm_fp8 error under cublas12.9 caused by bumpallocator, detail in pr#11612 - q_nope_val, q_nope_scale = per_tensor_quant_mla_fp8( - q_nope.transpose(0, 1), - ( - torch.zeros((1,), dtype=torch.float32, device=q_nope.device) - if _is_cublas_ge_129 - else zero_allocator.allocate(1) - ), - ) - q_nope_out = bmm_fp8( - q_nope_val, self.w_kc, q_nope_scale, self.w_scale, torch.bfloat16 - ) - else: - q_nope_out = torch.bmm(q_nope.transpose(0, 1), self.w_kc) - - q_nope_out = q_nope_out.transpose(0, 1) - - if ( - self.rotary_emb is not None - and (not self._fuse_rope_for_trtllm_mla(forward_batch)) - and (not _use_aiter or not _is_hip or self.use_nsa) - ): - q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe) - - if nsa_use_prefill_cp(forward_batch): - # support allgather+rerrange - k_nope, k_pe = self.rebuild_cp_kv_cache( - latent_cache, forward_batch, k_nope, k_pe - ) - - return ( - q_pe, - k_pe, - q_nope_out, - k_nope, - forward_batch, - zero_allocator, - positions, - topk_indices, - llama_4_scaling, - ) - - def forward_absorb_core( - self, - q_pe, - k_pe, - q_nope_out, - k_nope, - forward_batch, - zero_allocator, - positions, - topk_indices, - llama_4_scaling, - ): - save_kv_cache = True - - if self.current_attention_backend in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS: - extra_args = {} - if self._fuse_rope_for_trtllm_mla(forward_batch): - extra_args = { - "cos_sin_cache": self.rotary_emb.cos_sin_cache, - "is_neox": self.rotary_emb.is_neox_style, - "llama_4_scaling": llama_4_scaling, - } - - attn_output = self.attn_mqa( - q_nope_out, - k_nope, - k_nope, - forward_batch, - q_rope=q_pe, - k_rope=k_pe, - **extra_args, - **(dict(topk_indices=topk_indices) if topk_indices is not None else {}), - ) - else: - if _use_aiter: - cos = self.rotary_emb.cos_cache - sin = self.rotary_emb.sin_cache - - kv_cache_dtype = ( - fp8_dtype if self.kv_cache_dtype == "fp8_e4m3" else q_nope_out.dtype - ) - - q, _, _, k = fused_qk_rope_cat_and_cache_mla( - q_nope_out, - q_pe, - k_nope, - k_pe, - forward_batch.token_to_kv_pool.get_key_buffer( - self.attn_mqa.layer_id - ), - forward_batch.out_cache_loc, - positions, - cos, - sin, - self.attn_mqa.k_scale, - self.rotary_emb.is_neox_style, - q_out_dtype=kv_cache_dtype, - ) - - save_kv_cache = False - else: - q = torch.cat([q_nope_out, q_pe], dim=-1) - k = torch.cat([k_nope, k_pe], dim=-1) - - # Apply llama 4 scaling if provided - if llama_4_scaling is not None: - q *= llama_4_scaling - - attn_output = self.attn_mqa( - q, - k, - k_nope, - forward_batch, - save_kv_cache=save_kv_cache, - **(dict(topk_indices=topk_indices) if topk_indices is not None else {}), - ) - attn_output = attn_output.view(-1, self.num_local_heads, self.kv_lora_rank) - - if self.use_deep_gemm_bmm: - attn_output_val, attn_output_scale, masked_m, expected_m, aligned_m = ( - per_token_group_quant_mla_deep_gemm_masked_fp8( - attn_output.transpose(0, 1) - ) - ) - attn_bmm_output = attn_output.new_empty( - (self.num_local_heads, aligned_m, self.v_head_dim) - ) - deep_gemm_wrapper.grouped_gemm_nt_f8f8bf16_masked( - (attn_output_val, attn_output_scale), - (self.w_vc, self.w_scale_v), - attn_bmm_output, - masked_m, - expected_m, - ) - attn_bmm_output = ( - attn_bmm_output[:, :expected_m, :].transpose(0, 1).flatten(1, 2) - ) - elif _is_hip: - # TODO(haishaw): add bmm_fp8 to ROCm - if _use_aiter_gfx95 and self.w_vc.dtype == torch.uint8: - x = attn_output.transpose(0, 1) - attn_bmm_output = torch.empty( - x.shape[0], - x.shape[1], - self.w_vc.shape[2], - device=x.device, - dtype=torch.bfloat16, - ) - batched_gemm_afp4wfp4_pre_quant( - x, - self.w_vc.transpose(-2, -1), - self.w_scale_v.transpose(-2, -1), - torch.bfloat16, - attn_bmm_output, - ) - else: - if (_use_aiter_gfx95 and self.w_kc.dtype == torch.float8_e4m3fn) or ( - get_is_capture_mode() and self.w_kc.dtype == torch.float8_e4m3fnuz - ): - # fp8 Triton kernel: always on gfx950, - # cudagraph-only on gfx942 (hides launch overhead) - attn_bmm_output = batched_gemm_a8w8_a_per_token_group_prequant_w_per_batched_tensor_quant( - X=attn_output, - WQ=self.w_vc.transpose(-1, -2), - w_scale=self.w_scale, - group_size=128, - YQ=None, - transpose_bm=False, - transpose_bm_in=True, - dtype=torch.bfloat16, - ) - else: - attn_bmm_output = torch.bmm( - attn_output.to(torch.bfloat16).transpose(0, 1), - self.w_vc.to(torch.bfloat16) * self.w_scale, - ) - - if self.o_proj.weight.dtype == torch.uint8: - attn_bmm_output = attn_bmm_output.transpose(0, 1) - attn_bmm_output = fused_flatten_mxfp4_quant(attn_bmm_output) - elif self.o_proj.weight.dtype == torch.float8_e4m3fn: - attn_bmm_output = attn_bmm_output.transpose(0, 1) - attn_bmm_output = fused_flatten_fp8_group_quant( - attn_bmm_output, group_size=128, dtype_quant=torch.float8_e4m3fn - ) - else: - attn_bmm_output = attn_bmm_output.transpose(0, 1).flatten(1, 2) - - elif self.w_vc.dtype == torch.float8_e4m3fn: - attn_output_val, attn_output_scale = per_tensor_quant_mla_fp8( - attn_output.transpose(0, 1), - ( - torch.zeros((1,), dtype=torch.float32, device=attn_output.device) - if _is_cublas_ge_129 - else zero_allocator.allocate(1) - ), - ) - attn_bmm_output = bmm_fp8( - attn_output_val, - self.w_vc, - attn_output_scale, - self.w_scale, - torch.bfloat16, - ) - attn_bmm_output = attn_bmm_output.transpose(0, 1).flatten(1, 2) - else: - if is_in_piecewise_cuda_graph(): - # torch dynamo requires out= op was called where output tensor was non-contiguous - attn_bmm_output = ( - torch.bmm(attn_output.transpose(0, 1), self.w_vc) - .transpose(0, 1) - .flatten(1, 2) - ) - else: - attn_bmm_output = torch.empty( - (attn_output.shape[0], self.num_local_heads * self.v_head_dim), - dtype=attn_output.dtype, - device=attn_output.device, - ) - torch.bmm( - attn_output.transpose(0, 1), - self.w_vc, - out=attn_bmm_output.view( - -1, self.num_local_heads, self.v_head_dim - ).transpose(0, 1), - ) - output, _ = self.o_proj(attn_bmm_output) - - return output - - def forward_absorb_fused_mla_rope_prepare( - self, - positions: torch.Tensor, - hidden_states: torch.Tensor, - forward_batch: ForwardBatch, - zero_allocator: BumpAllocator, - ): - enable_rope_fusion = ( - os.getenv("SGLANG_FUSED_MLA_ENABLE_ROPE_FUSION", "1") == "1" - ) - # NOTE: hidden_states can be a tuple for some quantization paths. - # For shape/device/dtype, use the first tensor; still pass the original - # hidden_states through linear ops which may accept tuple inputs. - hidden_states_tensor = ( - hidden_states[0] if isinstance(hidden_states, tuple) else hidden_states - ) - - q_len = hidden_states_tensor.shape[0] - q_input = hidden_states_tensor.new_empty( - q_len, self.num_local_heads, self.kv_lora_rank + self.qk_rope_head_dim - ) - if self.q_lora_rank is not None: - q, latent_cache = self.fused_qkv_a_proj_with_mqa(hidden_states)[0].split( - [self.q_lora_rank, self.kv_lora_rank + self.qk_rope_head_dim], dim=-1 - ) - q = self.q_a_layernorm(q) - q = self.q_b_proj(q)[0].view(-1, self.num_local_heads, self.qk_head_dim) - else: - q = self.q_proj(hidden_states)[0].view( - -1, self.num_local_heads, self.qk_head_dim - ) - latent_cache = self.kv_a_proj_with_mqa(hidden_states)[0] - q_nope, q_pe = q.split([self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) - - if _is_hip: - # TODO(haishaw): add bmm_fp8 to ROCm - q_nope_out = torch.bmm( - q_nope.to(torch.bfloat16).transpose(0, 1), - self.w_kc.to(torch.bfloat16) * self.w_scale, - ) - elif self.w_kc.dtype == torch.float8_e4m3fn: - q_nope_val, q_nope_scale = per_tensor_quant_mla_fp8( - q_nope.transpose(0, 1), - zero_allocator.allocate(1), - dtype=torch.float8_e4m3fn, - ) - q_nope_out = bmm_fp8( - q_nope_val, self.w_kc, q_nope_scale, self.w_scale, torch.bfloat16 - ) - else: - q_nope_out = torch.bmm(q_nope.transpose(0, 1), self.w_kc) - q_input[..., : self.kv_lora_rank] = q_nope_out.transpose(0, 1) - v_input = latent_cache[..., : self.kv_lora_rank] - v_input = self.kv_a_layernorm(v_input.contiguous()).unsqueeze(1) - k_input = latent_cache.unsqueeze(1) - k_input[..., : self.kv_lora_rank] = v_input - - if not enable_rope_fusion: - k_pe = k_input[..., self.kv_lora_rank :] - q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe) - q_input[..., self.kv_lora_rank :] = q_pe - k_input[..., self.kv_lora_rank :] = k_pe - k_pe_output = None - else: - k_pe_output = torch.empty_like(k_input[..., self.kv_lora_rank :]) - - q_input[..., self.kv_lora_rank :] = q_pe - - # attn_output = self.attn_mqa(q_input, k_input, v_input, forward_batch) - # Use Fused ROPE with use_rope=OFF. - attn_output = torch.empty( - (q_len, self.num_local_heads, self.kv_lora_rank), - dtype=q.dtype, - device=q.device, - ) - attn_logits, _, kv_indptr, kv_indices, _, _, _ = ( - forward_batch.attn_backend.forward_metadata - ) - cos_sin_cache = self.rotary_emb.cos_sin_cache - num_kv_split = forward_batch.attn_backend.num_kv_splits - sm_scale = self.attn_mqa.scaling - if attn_logits is None: - attn_logits = torch.empty( - ( - forward_batch.batch_size, - self.num_local_heads, - num_kv_split, - self.kv_lora_rank + 1, - ), - dtype=torch.float32, - device=q.device, - ) - - # save current latent cache. - forward_batch.token_to_kv_pool.set_kv_buffer( - self.attn_mqa, forward_batch.out_cache_loc, k_input, None - ) - key_cache_buf = forward_batch.token_to_kv_pool.get_key_buffer( - self.attn_mqa.layer_id - ) - val_cache_buf = key_cache_buf[..., : self.kv_lora_rank] - - return ( - q_input, - key_cache_buf, - val_cache_buf, - attn_output, - kv_indptr, - kv_indices, - k_pe_output, - cos_sin_cache, - positions, - attn_logits, - num_kv_split, - sm_scale, - enable_rope_fusion, - k_input, - forward_batch, - zero_allocator, - ) - - def forward_absorb_fused_mla_rope_cpu_prepare( - self, - positions: torch.Tensor, - hidden_states: torch.Tensor, - forward_batch: ForwardBatch, - zero_allocator: BumpAllocator, - ): - assert self.q_lora_rank is not None and use_intel_amx_backend( - self - ), "forward_absorb_fused_mla_rope_cpu_prepare requires q_lora_rank is not None and use_intel_amx_backend" - - q_input, k_input, v_input = ( - torch.ops.sgl_kernel.qkv_proj_with_rope_fused_weight( - hidden_states, - self.fused_qkv_a_proj_with_mqa.weight, - self.q_b_proj.weight, - self.w_kc, - self.q_a_layernorm.weight, - self.kv_a_layernorm.weight, - positions, - self.rotary_emb.cos_sin_cache, - self.kv_a_layernorm.variance_epsilon, - self.qkv_proj_with_rope_is_int8, - self.qkv_proj_with_rope_is_fp8, - ( - self.fused_qkv_a_proj_with_mqa.weight_scale - if self.qkv_proj_with_rope_is_int8 - else ( - self.fused_qkv_a_proj_with_mqa.weight_scale_inv - if self.qkv_proj_with_rope_is_fp8 - else None - ) - ), - ( - self.q_b_proj.weight_scale - if self.qkv_proj_with_rope_is_int8 - else ( - self.q_b_proj.weight_scale_inv - if self.qkv_proj_with_rope_is_fp8 - else None - ) - ), - True, # is_vnni - self.weight_block_size, - self.q_lora_rank, - self.kv_lora_rank, - self.qk_rope_head_dim, - ) - ) - return (q_input, k_input, v_input, forward_batch, zero_allocator) - - def forward_absorb_fused_mla_rope_core( - self, - q_input, - key_cache_buf, - val_cache_buf, - attn_output, - kv_indptr, - kv_indices, - k_pe_output, - cos_sin_cache, - positions, - attn_logits, - num_kv_split, - sm_scale, - enable_rope_fusion, - k_input, - forward_batch, - zero_allocator, - ): - decode_attention_fwd_grouped_rope( - q_input, - key_cache_buf, - val_cache_buf, - attn_output, - kv_indptr, - kv_indices, - k_pe_output, - self.kv_lora_rank, - self.rotary_emb.rotary_dim, - cos_sin_cache, - positions, - attn_logits, - num_kv_split, - sm_scale, - logit_cap=self.attn_mqa.logit_cap, - use_rope=enable_rope_fusion, - is_neox_style=self.rotary_emb.is_neox_style, - ) - - if enable_rope_fusion: - k_input[..., self.kv_lora_rank :] = k_pe_output - forward_batch.token_to_kv_pool.set_kv_buffer( - self.attn_mqa, forward_batch.out_cache_loc, k_input, None - ) - - attn_output = attn_output.view(-1, self.num_local_heads, self.kv_lora_rank) - - if _is_hip: - # TODO(haishaw): add bmm_fp8 to ROCm - attn_bmm_output = torch.bmm( - attn_output.to(torch.bfloat16).transpose(0, 1), - self.w_vc.to(torch.bfloat16) * self.w_scale, - ) - elif self.w_vc.dtype == torch.float8_e4m3fn: - attn_output_val, attn_output_scale = per_tensor_quant_mla_fp8( - attn_output.transpose(0, 1), - zero_allocator.allocate(1), - dtype=torch.float8_e4m3fn, - ) - attn_bmm_output = bmm_fp8( - attn_output_val, - self.w_vc, - attn_output_scale, - self.w_scale, - torch.bfloat16, - ) - else: - attn_bmm_output = torch.bmm(attn_output.transpose(0, 1), self.w_vc) - attn_output = attn_bmm_output.transpose(0, 1).flatten(1, 2) - output, _ = self.o_proj(attn_output) - - return output - - def forward_absorb_fused_mla_rope_cpu_core( - self, q_input, k_input, v_input, forward_batch, zero_allocator - ): - assert self.q_lora_rank is not None and use_intel_amx_backend( - self - ), "forward_absorb_fused_mla_rope_cpu_core requires q_lora_rank is not None and use_intel_amx_backend" - - attn_output = self.attn_mqa(q_input, k_input, v_input, forward_batch) - attn_output = attn_output.view(-1, self.num_local_heads, self.kv_lora_rank) - - # [Note] Align shapes of bmm inputs. - # Shapes of inputs: - # q_nope: [M, B, K] - # original self.w_kc: [B, K, N] - # current self.w_kc (which has been converted in PackWeightMethod): [B, N, K] - - # Shapes of inputs to sgl_kernel.cpu.bmm: - # out: [B, M, N] - # mat1: [B, M, K] - # mat2: [B, N, K] - B = self.w_vc.size(0) - N = self.w_vc.size(1) - M = attn_output.size(0) - output = torch.empty([M, int(B * N)], dtype=attn_output.dtype) - attn_bmm_output = output.view([M, B, N]).transpose_(0, 1) - torch.ops.sgl_kernel.bmm_cpu( - attn_bmm_output, - attn_output.transpose(0, 1), - self.w_vc, - True, # is_vnni - None, # scale - ) - attn_output = output - output, _ = self.o_proj(attn_output) - - return output - @staticmethod def _get_q_b_proj_quant_config(quant_config): if envs.SGLANG_NVFP4_CKPT_FP8_GEMM_IN_ATTN.get(): diff --git a/test/srt/cpu/test_qkv_proj_with_rope.py b/test/srt/cpu/test_qkv_proj_with_rope.py index b0b22d3bf..366043dc9 100644 --- a/test/srt/cpu/test_qkv_proj_with_rope.py +++ b/test/srt/cpu/test_qkv_proj_with_rope.py @@ -8,7 +8,7 @@ from utils import ( precision, ) -from sglang.srt.layers.rotary_embedding import _apply_rotary_emb +from sglang.srt.layers.rotary_embedding.utils import apply_rotary_emb from sglang.test.test_utils import CustomTestCase convert_weight_packed = torch.ops.sgl_kernel.convert_weight_packed @@ -46,8 +46,8 @@ def rotary_emb(q_pe, k_pe, pos, cos_sin_cache): key_rot = k_pe[..., :rotary_dim] cos_sin = cos_sin_cache[pos] cos, sin = cos_sin.chunk(2, dim=-1) - query_rot = _apply_rotary_emb(query_rot, cos, sin, False) - key_rot = _apply_rotary_emb(key_rot, cos, sin, False) + query_rot = apply_rotary_emb(query_rot, cos, sin, False) + key_rot = apply_rotary_emb(key_rot, cos, sin, False) return query_rot.to(orig_dtype), key_rot.to(orig_dtype) diff --git a/test/srt/cpu/test_rope.py b/test/srt/cpu/test_rope.py index 00aace4ed..5eeac5337 100644 --- a/test/srt/cpu/test_rope.py +++ b/test/srt/cpu/test_rope.py @@ -3,9 +3,9 @@ import unittest import torch from utils import precision -from sglang.srt.layers.rotary_embedding import ( +from sglang.srt.layers.rotary_embedding.base import RotaryEmbedding +from sglang.srt.layers.rotary_embedding.rope_variant import ( DeepseekScalingRotaryEmbedding, - RotaryEmbedding, ) from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler from sglang.test.test_utils import CustomTestCase