[3/n] deepseek_v2.py Refactor: Migrate MLA forward method in deepseek_v2.py (#19122)

This commit is contained in:
Baizhou Zhang
2026-02-27 13:37:29 -08:00
committed by GitHub
parent 7e46aafebb
commit 776709efe8
9 changed files with 906 additions and 818 deletions
@@ -25,7 +25,7 @@ class AttentionBackendRegistry:
def _dispatch_mla_subtype(attn, forward_batch): def _dispatch_mla_subtype(attn, forward_batch):
if _is_hip: if _is_hip:
if attn.rocm_fused_decode_mla and forward_batch.forward_mode.is_decode(): 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: else:
return AttnForwardMethod.MLA return AttnForwardMethod.MLA
else: else:
@@ -1,7 +1,13 @@
from .forward_methods import AttnForwardMethod from .forward_methods import AttnForwardMethod
from .forward_mha import DeepseekMHAForwardMixin 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__ = [ __all__ = [
"AttnForwardMethod", "AttnForwardMethod",
"DeepseekMHAForwardMixin", "DeepseekMHAForwardMixin",
"DeepseekMLACpuForwardMixin",
"DeepseekMLAForwardMixin",
"DeepseekMLARocmForwardMixin",
] ]
@@ -17,7 +17,7 @@ class AttnForwardMethod(IntEnum):
MHA_ONE_SHOT = auto() MHA_ONE_SHOT = auto()
# Use MLA but with fused RoPE # Use MLA but with fused RoPE
MLA_FUSED_ROPE = auto() MLA_FUSED_ROPE_ROCM = auto()
# Use MLA with fused RoPE kernel for CPU # Use MLA with fused RoPE kernel for CPU
MLA_FUSED_ROPE_CPU = auto() MLA_FUSED_ROPE_CPU = auto()
@@ -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
)
@@ -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
@@ -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
+22 -811
View File
@@ -19,7 +19,6 @@
from __future__ import annotations from __future__ import annotations
import logging import logging
import os
from contextlib import nullcontext from contextlib import nullcontext
from typing import Any, Dict, Iterable, List, Optional, Tuple, Union from typing import Any, Dict, Iterable, List, Optional, Tuple, Union
@@ -33,7 +32,6 @@ from sglang.srt.batch_overlap.two_batch_overlap import (
MaybeTboDeepEPDispatcher, MaybeTboDeepEPDispatcher,
model_forward_maybe_tbo, model_forward_maybe_tbo,
) )
from sglang.srt.compilation.piecewise_context_manager import is_in_piecewise_cuda_graph
from sglang.srt.configs.model_config import ( from sglang.srt.configs.model_config import (
get_nsa_index_head_dim, get_nsa_index_head_dim,
get_nsa_index_n_heads, 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.base_config import QuantizationConfig
from sglang.srt.layers.quantization.fp8 import Fp8Config 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.radix_attention import RadixAttention
from sglang.srt.layers.rotary_embedding import get_rope_wrapper from sglang.srt.layers.rotary_embedding import get_rope_wrapper
from sglang.srt.layers.utils import PPMissingLayer 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 ( from sglang.srt.models.deepseek_common.attention_forward_methods import (
AttnForwardMethod, AttnForwardMethod,
DeepseekMHAForwardMixin, DeepseekMHAForwardMixin,
DeepseekMLACpuForwardMixin,
DeepseekMLAForwardMixin,
DeepseekMLARocmForwardMixin,
) )
from sglang.srt.models.deepseek_common.deepseek_weight_loader import ( from sglang.srt.models.deepseek_common.deepseek_weight_loader import (
DeepseekV2WeightLoaderMixin, DeepseekV2WeightLoaderMixin,
) )
from sglang.srt.models.deepseek_common.utils import ( from sglang.srt.models.deepseek_common.utils import (
FORWARD_ABSORB_CORE_ATTENTION_BACKENDS,
_device_sm, _device_sm,
_get_llama_4_scaling, _get_llama_4_scaling,
_is_cpu, _is_cpu,
_is_cpu_amx_available, _is_cpu_amx_available,
_is_cublas_ge_129,
_is_cuda, _is_cuda,
_is_gfx95_supported, _is_gfx95_supported,
_is_hip, _is_hip,
@@ -152,45 +146,23 @@ from sglang.srt.utils import (
BumpAllocator, BumpAllocator,
LazyValue, LazyValue,
add_prefix, add_prefix,
get_bool_env_var,
is_non_idle_and_non_empty, is_non_idle_and_non_empty,
log_info_on_rank0, log_info_on_rank0,
make_layers, make_layers,
use_intel_amx_backend, 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: 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 ( from sglang.srt.layers.rocm_linear_utils import (
aiter_dsv3_router_gemm, aiter_dsv3_router_gemm,
fused_qk_rope_cat_and_cache_mla,
get_dsv3_gemm_output_zero_allocator_size, get_dsv3_gemm_output_zero_allocator_size,
) )
if _use_aiter: if _use_aiter:
from sglang.srt.layers.rocm_linear_utils import fused_qk_rope_cat_and_cache_mla pass
if _is_cuda: if _is_cuda:
from sgl_kernel import bmm_fp8, dsv3_fused_a_gemm, dsv3_router_gemm from sgl_kernel import 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,
)
elif _is_npu: elif _is_npu:
from sglang.srt.hardware_backend.npu.modules.deepseek_v2_attention_mla_npu import ( from sglang.srt.hardware_backend.npu.modules.deepseek_v2_attention_mla_npu import (
forward_dsa_core_npu, forward_dsa_core_npu,
@@ -206,16 +178,6 @@ else:
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
FORWARD_ABSORB_CORE_ATTENTION_BACKENDS = [
"fa3",
"nsa",
"flashinfer",
"cutlass_mla",
"trtllm_mla",
"ascend",
]
class DeepseekV2MLP(nn.Module): class DeepseekV2MLP(nn.Module):
def __init__( def __init__(
self, self,
@@ -1077,7 +1039,13 @@ class DeepseekV2MoE(nn.Module):
state.hidden_states_mlp_output = final_hidden_states state.hidden_states_mlp_output = final_hidden_states
class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): class DeepseekV2AttentionMLA(
nn.Module,
DeepseekMHAForwardMixin,
DeepseekMLAForwardMixin,
DeepseekMLARocmForwardMixin,
DeepseekMLACpuForwardMixin,
):
def __init__( def __init__(
self, self,
@@ -1266,35 +1234,20 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin):
self.w_scale_v = None self.w_scale_v = None
self.use_deep_gemm_bmm = False self.use_deep_gemm_bmm = False
self.flashinfer_mla_disable_ragged = (
get_global_server_args().flashinfer_mla_disable_ragged
)
self.current_attention_backend = ( self.current_attention_backend = (
None # Attention backend used by current forward batch 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 self.has_fused_proj = hasattr(self, "fused_qkv_a_proj_with_mqa")
# which requires self.w_kc and self.w_vc to be packed. self.is_packed_weight = (
# If not, we will use torch.bmm and weight shouldn't be packed in this case self.has_fused_proj
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
and hasattr(self.fused_qkv_a_proj_with_mqa.quant_method, "quant_config") 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() and self.fused_qkv_a_proj_with_mqa.quant_method.quant_config.get_name()
in {"awq", "awq_marlin", "moe_wna16"} in {"awq", "awq_marlin", "moe_wna16"}
) )
self.use_min_latency_fused_a_gemm = ( self.use_min_latency_fused_a_gemm = (
has_fused_proj self.has_fused_proj
and not is_packed_weight 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.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[0] == 2112
and self.fused_qkv_a_proj_with_mqa.weight.shape[1] == 7168 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 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_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( def dispatch_attn_forward_method(
self, forward_batch: ForwardBatch self, forward_batch: ForwardBatch
@@ -1433,7 +1360,7 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin):
inner_state = self.forward_absorb_prepare( inner_state = self.forward_absorb_prepare(
positions, hidden_states, forward_batch, zero_allocator, llama_4_scaling 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( inner_state = self.forward_absorb_fused_mla_rope_prepare(
positions, hidden_states, forward_batch, zero_allocator 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) return self.forward_normal_one_shot_core(*inner_state)
elif attn_forward_method == AttnForwardMethod.MLA: elif attn_forward_method == AttnForwardMethod.MLA:
return self.forward_absorb_core(*inner_state) 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) return self.forward_absorb_fused_mla_rope_core(*inner_state)
elif attn_forward_method == AttnForwardMethod.MLA_FUSED_ROPE_CPU: elif attn_forward_method == AttnForwardMethod.MLA_FUSED_ROPE_CPU:
return self.forward_absorb_fused_mla_rope_cpu_core(*inner_state) 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] qkv_latent = self.fused_qkv_a_proj_with_mqa(hidden_states)[0]
return qkv_latent 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): def rebuild_cp_kv_cache(self, latent_cache, forward_batch, k_nope, k_pe):
# support allgather+rerrange # support allgather+rerrange
latent_cache[..., : self.kv_lora_rank] = k_nope.squeeze(1) 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) k_pe = latent_cache_output[..., self.kv_lora_rank :].unsqueeze(1)
return k_nope, k_pe 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 @staticmethod
def _get_q_b_proj_quant_config(quant_config): def _get_q_b_proj_quant_config(quant_config):
if envs.SGLANG_NVFP4_CKPT_FP8_GEMM_IN_ATTN.get(): if envs.SGLANG_NVFP4_CKPT_FP8_GEMM_IN_ATTN.get():
+3 -3
View File
@@ -8,7 +8,7 @@ from utils import (
precision, 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 from sglang.test.test_utils import CustomTestCase
convert_weight_packed = torch.ops.sgl_kernel.convert_weight_packed 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] key_rot = k_pe[..., :rotary_dim]
cos_sin = cos_sin_cache[pos] cos_sin = cos_sin_cache[pos]
cos, sin = cos_sin.chunk(2, dim=-1) cos, sin = cos_sin.chunk(2, dim=-1)
query_rot = _apply_rotary_emb(query_rot, cos, sin, False) query_rot = apply_rotary_emb(query_rot, cos, sin, False)
key_rot = _apply_rotary_emb(key_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) return query_rot.to(orig_dtype), key_rot.to(orig_dtype)
+2 -2
View File
@@ -3,9 +3,9 @@ import unittest
import torch import torch
from utils import precision 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, DeepseekScalingRotaryEmbedding,
RotaryEmbedding,
) )
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase