[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

View File

@@ -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:

View File

@@ -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",
]

View File

@@ -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()

View File

@@ -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
)

View File

@@ -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

View File

@@ -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

View File

@@ -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():