Support NoPE layers in the tokenspeed_mla FP8 prefill hook (#38152)
Co-authored-by: Mohammad Angkad <mohammad.angkad@radixark.ai>
This commit is contained in:
co-authored by
Mohammad Angkad
parent
514b45fd34
commit
09daea94ac
@@ -180,13 +180,13 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
|||||||
enable_ex2_emulation=enable_ex2_emulation,
|
enable_ex2_emulation=enable_ex2_emulation,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
def _fused_rope_fp8_quantize(
|
def _fused_rope_fp8_quantize(
|
||||||
self,
|
|
||||||
q_nope: torch.Tensor,
|
q_nope: torch.Tensor,
|
||||||
q_pe: torch.Tensor,
|
q_pe: torch.Tensor,
|
||||||
k_nope: torch.Tensor,
|
k_nope: torch.Tensor,
|
||||||
k_pe: torch.Tensor,
|
k_pe: torch.Tensor,
|
||||||
cos_sin_cache: torch.Tensor,
|
cos_sin_cache: Optional[torch.Tensor],
|
||||||
positions: torch.Tensor,
|
positions: torch.Tensor,
|
||||||
is_neox: bool,
|
is_neox: bool,
|
||||||
qk_nope_head_dim: int,
|
qk_nope_head_dim: int,
|
||||||
@@ -194,6 +194,9 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
|||||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
"""Fused RoPE + FP8 quantize that also packs nope+pe along the last
|
"""Fused RoPE + FP8 quantize that also packs nope+pe along the last
|
||||||
dim, so FMHA consumes contig FP8 Q/K without an extra concat or cast.
|
dim, so FMHA consumes contig FP8 Q/K without an extra concat or cast.
|
||||||
|
|
||||||
|
``cos_sin_cache`` is None for NoPE layers (``skip_rope``); they keep the
|
||||||
|
FP8 quantize and the packed layout, only the rotation is dropped.
|
||||||
"""
|
"""
|
||||||
num_heads = q_nope.shape[1]
|
num_heads = q_nope.shape[1]
|
||||||
seq_len = q_nope.shape[0]
|
seq_len = q_nope.shape[0]
|
||||||
@@ -218,6 +221,14 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
|||||||
else:
|
else:
|
||||||
k_pe_expanded = k_pe
|
k_pe_expanded = k_pe
|
||||||
|
|
||||||
|
if cos_sin_cache is None:
|
||||||
|
# NoPE layer: quantize straight into the packed buffers.
|
||||||
|
q_fp8[..., :qk_nope_head_dim].copy_(q_nope)
|
||||||
|
q_fp8[..., qk_nope_head_dim:].copy_(q_pe)
|
||||||
|
k_fp8[..., :qk_nope_head_dim].copy_(k_nope)
|
||||||
|
k_fp8[..., qk_nope_head_dim:].copy_(k_pe_expanded)
|
||||||
|
return q_fp8, k_fp8
|
||||||
|
|
||||||
_flashinfer_rope.mla_rope_quantize_fp8(
|
_flashinfer_rope.mla_rope_quantize_fp8(
|
||||||
q_rope=q_pe,
|
q_rope=q_pe,
|
||||||
k_rope=k_pe_expanded,
|
k_rope=k_pe_expanded,
|
||||||
@@ -257,14 +268,15 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
|||||||
v_bf16 = kv[..., layer.qk_nope_head_dim :]
|
v_bf16 = kv[..., layer.qk_nope_head_dim :]
|
||||||
q_nope = q[..., : layer.qk_nope_head_dim]
|
q_nope = q[..., : layer.qk_nope_head_dim]
|
||||||
|
|
||||||
|
rotary_emb = layer.rotary_emb
|
||||||
q_fp8, k_fp8 = self._fused_rope_fp8_quantize(
|
q_fp8, k_fp8 = self._fused_rope_fp8_quantize(
|
||||||
q_nope=q_nope,
|
q_nope=q_nope,
|
||||||
q_pe=q_pe,
|
q_pe=q_pe,
|
||||||
k_nope=k_nope,
|
k_nope=k_nope,
|
||||||
k_pe=k_pe,
|
k_pe=k_pe,
|
||||||
cos_sin_cache=layer.rotary_emb.cos_sin_cache,
|
cos_sin_cache=None if rotary_emb is None else rotary_emb.cos_sin_cache,
|
||||||
positions=positions,
|
positions=positions,
|
||||||
is_neox=getattr(layer.rotary_emb, "is_neox_style", True),
|
is_neox=True if rotary_emb is None else rotary_emb.is_neox_style,
|
||||||
qk_nope_head_dim=layer.qk_nope_head_dim,
|
qk_nope_head_dim=layer.qk_nope_head_dim,
|
||||||
qk_rope_head_dim=layer.qk_rope_head_dim,
|
qk_rope_head_dim=layer.qk_rope_head_dim,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ import unittest
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.layers.attention.tokenspeed_mla_backend import TokenspeedMLABackend
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
from sglang.test.kits.attention_unittest.attention_methods.mla_attention import (
|
from sglang.test.kits.attention_unittest.attention_methods.mla_attention import (
|
||||||
MLAAttentionCase,
|
MLAAttentionCase,
|
||||||
@@ -181,5 +182,84 @@ class TestTokenspeedMLAAttentionBackendCorrectness(CustomTestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestTokenspeedMLANoPEPrefillQuantize(CustomTestCase):
|
||||||
|
"""NoPE MLA layers must still get packed FP8 prefill Q/K.
|
||||||
|
|
||||||
|
Layers built with ``skip_rope`` carry no rotary embedding, so the prefill
|
||||||
|
quantize runs with ``cos_sin_cache=None``. It must return the same packed
|
||||||
|
``[nope | pe]`` FP8 layout as the RoPE path, and the head-0 pe slice it
|
||||||
|
writes to the KV cache must equal the plain FP8 cast of the unroped
|
||||||
|
``k_pe`` that the decode path reads back.
|
||||||
|
"""
|
||||||
|
|
||||||
|
T = 7
|
||||||
|
NUM_HEADS = 4
|
||||||
|
QK_NOPE_HEAD_DIM = 128
|
||||||
|
QK_ROPE_HEAD_DIM = 64
|
||||||
|
KV_LORA_RANK = 512
|
||||||
|
|
||||||
|
def test_nope_prefill_quantize_packs_without_rope(self):
|
||||||
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
head_dim = self.QK_NOPE_HEAD_DIM + self.QK_ROPE_HEAD_DIM
|
||||||
|
|
||||||
|
q_nope = torch.randn(
|
||||||
|
self.T,
|
||||||
|
self.NUM_HEADS,
|
||||||
|
self.QK_NOPE_HEAD_DIM,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
q_pe = torch.randn(
|
||||||
|
self.T,
|
||||||
|
self.NUM_HEADS,
|
||||||
|
self.QK_ROPE_HEAD_DIM,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
k_nope = torch.randn(
|
||||||
|
self.T,
|
||||||
|
self.NUM_HEADS,
|
||||||
|
self.QK_NOPE_HEAD_DIM,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
# k_pe reaches the backend as the strided tail of the latent cache.
|
||||||
|
latent_cache = torch.randn(
|
||||||
|
self.T,
|
||||||
|
1,
|
||||||
|
self.KV_LORA_RANK + self.QK_ROPE_HEAD_DIM,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
k_pe = latent_cache[:, :, self.KV_LORA_RANK :]
|
||||||
|
|
||||||
|
q_fp8, k_fp8 = TokenspeedMLABackend._fused_rope_fp8_quantize(
|
||||||
|
q_nope=q_nope,
|
||||||
|
q_pe=q_pe,
|
||||||
|
k_nope=k_nope,
|
||||||
|
k_pe=k_pe,
|
||||||
|
cos_sin_cache=None,
|
||||||
|
positions=torch.arange(self.T, device=device),
|
||||||
|
is_neox=True,
|
||||||
|
qk_nope_head_dim=self.QK_NOPE_HEAD_DIM,
|
||||||
|
qk_rope_head_dim=self.QK_ROPE_HEAD_DIM,
|
||||||
|
)
|
||||||
|
|
||||||
|
for name, out in (("q", q_fp8), ("k", k_fp8)):
|
||||||
|
with self.subTest(tensor=name):
|
||||||
|
self.assertEqual(out.shape, (self.T, self.NUM_HEADS, head_dim))
|
||||||
|
self.assertEqual(out.dtype, torch.float8_e4m3fn)
|
||||||
|
self.assertTrue(out.is_contiguous())
|
||||||
|
|
||||||
|
fp8 = torch.float8_e4m3fn
|
||||||
|
nope = self.QK_NOPE_HEAD_DIM
|
||||||
|
self.assertTrue(torch.equal(q_fp8[..., :nope], q_nope.to(fp8)))
|
||||||
|
self.assertTrue(torch.equal(q_fp8[..., nope:], q_pe.to(fp8)))
|
||||||
|
self.assertTrue(torch.equal(k_fp8[..., :nope], k_nope.to(fp8)))
|
||||||
|
# The slice prepare_prefill_qkv writes into the KV cache; the decode
|
||||||
|
# NoPE path reads it back as a plain cast of the unroped k_pe.
|
||||||
|
self.assertTrue(torch.equal(k_fp8[:, 0:1, nope:], k_pe.to(fp8)))
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user