[3/n] deepseek_v2.py Refactor: Migrate MLA forward method in deepseek_v2.py (#19122)
This commit is contained in:
@@ -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",
|
||||||
]
|
]
|
||||||
|
|||||||
+1
-1
@@ -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
|
||||||
|
)
|
||||||
+152
@@ -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
|
||||||
+227
@@ -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
|
||||||
@@ -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():
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user