diff --git a/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py b/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py index aceb69b19..333832545 100644 --- a/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py +++ b/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py @@ -180,13 +180,13 @@ class TokenspeedMLABackend(TRTLLMMLABackend): enable_ex2_emulation=enable_ex2_emulation, ) + @staticmethod def _fused_rope_fp8_quantize( - self, q_nope: torch.Tensor, q_pe: torch.Tensor, k_nope: torch.Tensor, k_pe: torch.Tensor, - cos_sin_cache: torch.Tensor, + cos_sin_cache: Optional[torch.Tensor], positions: torch.Tensor, is_neox: bool, qk_nope_head_dim: int, @@ -194,6 +194,9 @@ class TokenspeedMLABackend(TRTLLMMLABackend): ) -> tuple[torch.Tensor, torch.Tensor]: """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. + + ``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] seq_len = q_nope.shape[0] @@ -218,6 +221,14 @@ class TokenspeedMLABackend(TRTLLMMLABackend): else: 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( q_rope=q_pe, k_rope=k_pe_expanded, @@ -257,14 +268,15 @@ class TokenspeedMLABackend(TRTLLMMLABackend): v_bf16 = kv[..., 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_nope=q_nope, q_pe=q_pe, k_nope=k_nope, 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, - 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_rope_head_dim=layer.qk_rope_head_dim, ) diff --git a/test/registered/attention/unittests/mla/test_tokenspeed_mla.py b/test/registered/attention/unittests/mla/test_tokenspeed_mla.py index 5aa31f94d..a266d7b1c 100644 --- a/test/registered/attention/unittests/mla/test_tokenspeed_mla.py +++ b/test/registered/attention/unittests/mla/test_tokenspeed_mla.py @@ -3,6 +3,7 @@ import unittest import torch +from sglang.srt.layers.attention.tokenspeed_mla_backend import TokenspeedMLABackend from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.test.kits.attention_unittest.attention_methods.mla_attention import ( 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__": unittest.main()