From 01b003255a0bdaf31f866959cd02cba9fb871914 Mon Sep 17 00:00:00 2001 From: zhengyao Date: Thu, 16 Jul 2026 16:05:04 +0800 Subject: [PATCH] [AMD]Reuse fused FP8 KV cache write on standard aiter prefill/decode (#26852) --- .../srt/layers/attention/aiter_backend.py | 48 +- .../attention/test_fused_fp8_kv_write.py | 437 ++++++++++++++++++ 2 files changed, 479 insertions(+), 6 deletions(-) create mode 100644 test/registered/attention/test_fused_fp8_kv_write.py diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 60ef93923..513461823 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -1873,6 +1873,17 @@ class AiterAttnBackend(AttentionBackend): # base class NotImplementedError. pass + def _use_fused_fp8_kv_write(self, layer: RadixAttention) -> bool: + # Fused write reuses K's num_heads/head_dim for V, so it needs FP8, + # non-MLA, non-SWA, and matching K/V head count + head_dim. + return ( + self.kv_cache_dtype == fp8_dtype + and not self.use_mla + and not self.use_sliding_window_kv_pool + and layer.tp_k_head_num == layer.tp_v_head_num + and layer.qk_head_dim == layer.v_head_dim + ) + def forward_extend( self, q: torch.Tensor, @@ -1895,7 +1906,7 @@ class AiterAttnBackend(AttentionBackend): v_descale = None if self.kv_cache_dtype == fp8_dtype: k_descale = layer.k_scale if layer.k_scale is not None else self.k_scale - v_descale = layer.v_scale if layer.v_scale is not None else self.k_scale + v_descale = layer.v_scale if layer.v_scale is not None else self.v_scale if k is not None: assert v is not None @@ -1947,6 +1958,24 @@ class AiterAttnBackend(AttentionBackend): ) elif self.use_mla: self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v) + elif self._use_fused_fp8_kv_write(layer): + # FP8: fuse bf16->fp8 cast + paged write in one kernel. + k_cache, v_cache = self.token_to_kv_pool.get_kv_buffer( + layer.layer_id + ) + launch_reshape_and_cache_flash( + k.view(-1, layer.tp_k_head_num, layer.qk_head_dim), + v.view(-1, layer.tp_v_head_num, layer.v_head_dim), + k_cache.view( + -1, self.page_size, layer.tp_k_head_num, layer.qk_head_dim + ), + v_cache.view( + -1, self.page_size, layer.tp_v_head_num, layer.v_head_dim + ), + cache_loc, + k_scale=k_descale, + v_scale=v_descale, + ) else: self.token_to_kv_pool.set_kv_buffer( layer, @@ -2374,7 +2403,7 @@ class AiterAttnBackend(AttentionBackend): v_descale = None if self.kv_cache_dtype == fp8_dtype: k_descale = layer.k_scale if layer.k_scale is not None else self.k_scale - v_descale = layer.v_scale if layer.v_scale is not None else self.k_scale + v_descale = layer.v_scale if layer.v_scale is not None else self.v_scale if save_kv_cache: # SHUFFLE 5D pool path — see forward_extend for rationale. @@ -2414,10 +2443,13 @@ class AiterAttnBackend(AttentionBackend): k_scale=k_descale, v_scale=v_descale, ) - elif self.use_triton_unified_attention and self.kv_cache_dtype == fp8_dtype: - # [PATCH] FP8 non-SWA: use launch_reshape_and_cache_flash to - # fuse bf16→fp8 cast + paged write in one Triton kernel, - # eliminating separate float8_copy + store_kvcache overhead. + elif self.use_mla: + # MLA pool has its own set_kv_buffer (no scale args). + self.token_to_kv_pool.set_kv_buffer( + layer, forward_batch.out_cache_loc, k, v + ) + elif self._use_fused_fp8_kv_write(layer): + # FP8: fuse bf16->fp8 cast + paged write in one kernel. token_to_kv_pool = self.token_to_kv_pool k_cache, v_cache = token_to_kv_pool.get_kv_buffer(layer.layer_id) launch_reshape_and_cache_flash( @@ -2430,6 +2462,8 @@ class AiterAttnBackend(AttentionBackend): -1, self.page_size, layer.tp_v_head_num, layer.v_head_dim ), forward_batch.out_cache_loc, + k_scale=k_descale, + v_scale=v_descale, ) else: self.token_to_kv_pool.set_kv_buffer( @@ -2440,6 +2474,8 @@ class AiterAttnBackend(AttentionBackend): ), k, v, + k_descale, + v_descale, ) if self.use_mla: diff --git a/test/registered/attention/test_fused_fp8_kv_write.py b/test/registered/attention/test_fused_fp8_kv_write.py new file mode 100644 index 000000000..33d52094a --- /dev/null +++ b/test/registered/attention/test_fused_fp8_kv_write.py @@ -0,0 +1,437 @@ +"""Accuracy and dispatch tests for the fused bf16->fp8 KV-cache write on the +standard aiter FP8 path. Skipped on CPU (Triton requires a GPU).""" + +import unittest + +import torch + +from sglang.test.ci.ci_register import register_amd_ci + +register_amd_ci(est_time=30, suite="stage-b-test-1-gpu-small-amd") + +_HAS_CUDA = torch.cuda.is_available() + + +def _naive_bf16_write(k, v, k_cache, v_cache, loc, k_scale, v_scale, fp8_dtype): + """Naive reference: bf16 in-place divide -> fp8 cast -> scatter.""" + ck = k.clone() + cv = v.clone() + ck.div_(k_scale) + cv.div_(v_scale) + k_cache[loc, 0] = ck.to(fp8_dtype) + v_cache[loc, 0] = cv.to(fp8_dtype) + + +def _fp32_divide_write(k, v, k_cache, v_cache, loc, k_scale, v_scale, fp8_dtype): + """Fused-kernel reference: fp32 divide -> fp8 cast -> scatter.""" + k_cache[loc, 0] = (k.float() / k_scale.float()).to(fp8_dtype) + v_cache[loc, 0] = (v.float() / v_scale.float()).to(fp8_dtype) + + +@unittest.skipUnless(_HAS_CUDA, "Triton kernels require a GPU") +class TestFusedFp8KvWrite(unittest.TestCase): + def _run(self, num_tokens, num_heads, head_dim, total_slots=None, seed=0xC0FFEE): + from sglang.srt.layers.attention.utils import ( + launch_reshape_and_cache_flash, + ) + from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype + + torch.manual_seed(seed) + dev = "cuda" + page_size = 1 + if total_slots is None: + total_slots = num_tokens + assert num_tokens <= total_slots + + k = torch.randn( + (num_tokens, num_heads, head_dim), dtype=torch.bfloat16, device=dev + ) + v = torch.randn( + (num_tokens, num_heads, head_dim), dtype=torch.bfloat16, device=dev + ) + + # Per-tensor descale = amax / fp8_max (matches the production path). + fp8_max = torch.finfo(fp8_dtype).max + k_scale = (k.abs().amax().float() / fp8_max).clamp_min(1e-12).view(1) + v_scale = (v.abs().amax().float() / fp8_max).clamp_min(1e-12).view(1) + + loc = torch.randperm(total_slots, device=dev)[:num_tokens].to(torch.int64) + + # Fused path (kernel under test). + k_cache_f = torch.zeros( + (total_slots, page_size, num_heads, head_dim), dtype=fp8_dtype, device=dev + ) + v_cache_f = torch.zeros_like(k_cache_f) + launch_reshape_and_cache_flash( + k.view(-1, num_heads, head_dim), + v.view(-1, num_heads, head_dim), + k_cache_f.view(-1, page_size, num_heads, head_dim), + v_cache_f.view(-1, page_size, num_heads, head_dim), + loc, + k_scale=k_scale, + v_scale=v_scale, + ) + + return { + "k": k, + "v": v, + "k_scale": k_scale, + "v_scale": v_scale, + "loc": loc, + "fp8_dtype": fp8_dtype, + "k_cache_f": k_cache_f, + "v_cache_f": v_cache_f, + } + + def test_matches_fp32_divide_reference(self): + """Fused write matches the fp32-divide reference (within fp8 rounding).""" + r = self._run(num_tokens=32, num_heads=1, head_dim=256) + loc, fp8_dtype = r["loc"], r["fp8_dtype"] + + k_ref = torch.zeros_like(r["k_cache_f"]) + v_ref = torch.zeros_like(r["v_cache_f"]) + _fp32_divide_write( + r["k"], r["v"], k_ref, v_ref, loc, r["k_scale"], r["v_scale"], fp8_dtype + ) + + k_mismatch = (r["k_cache_f"] != k_ref).float().mean().item() + v_mismatch = (r["v_cache_f"] != v_ref).float().mean().item() + self.assertLess(k_mismatch, 1e-3, f"K vs fp32-divide ref: {k_mismatch:.4%}") + self.assertLess(v_mismatch, 1e-3, f"V vs fp32-divide ref: {v_mismatch:.4%}") + + def test_no_precision_regression_vs_naive(self): + """Fused write is at least as close to the fp32 reference as naive.""" + r = self._run(num_tokens=8000, num_heads=1, head_dim=256) + loc, fp8_dtype = r["loc"], r["fp8_dtype"] + + k_naive = torch.zeros_like(r["k_cache_f"]) + v_naive = torch.zeros_like(r["v_cache_f"]) + _naive_bf16_write( + r["k"], r["v"], k_naive, v_naive, loc, r["k_scale"], r["v_scale"], fp8_dtype + ) + + ref_k = r["k"].float() / r["k_scale"].float() + ref_v = r["v"].float() / r["v_scale"].float() + + err_fused = (r["k_cache_f"][loc, 0].float() - ref_k).abs().mean().item() + ( + r["v_cache_f"][loc, 0].float() - ref_v + ).abs().mean().item() + err_naive = (k_naive[loc, 0].float() - ref_k).abs().mean().item() + ( + v_naive[loc, 0].float() - ref_v + ).abs().mean().item() + + mismatch = ( + (r["k_cache_f"][loc, 0] != k_naive[loc, 0]).float().mean().item() + + (r["v_cache_f"][loc, 0] != v_naive[loc, 0]).float().mean().item() + ) / 2.0 + print( + f"[fused_fp8_kv_write] prefill fused_vs_naive_mismatch={mismatch:.4%} " + f"err_fused={err_fused:.5f} err_naive={err_naive:.5f}" + ) + + self.assertLessEqual( + err_fused, + err_naive * 1.02, + f"fused err {err_fused:.5f} regressed vs naive {err_naive:.5f}", + ) + + def test_writes_only_target_slots(self): + """Fused write touches only the target slots; others stay zero.""" + r = self._run(num_tokens=16, num_heads=4, head_dim=128, total_slots=64) + loc, fp8_dtype = r["loc"], r["fp8_dtype"] + + k_ref = torch.zeros_like(r["k_cache_f"]) + v_ref = torch.zeros_like(r["v_cache_f"]) + _fp32_divide_write( + r["k"], r["v"], k_ref, v_ref, loc, r["k_scale"], r["v_scale"], fp8_dtype + ) + + # Target slots match the fp32-divide reference. + self.assertTrue( + torch.equal(r["k_cache_f"][loc], k_ref[loc]), "target-slot K mismatch" + ) + self.assertTrue( + torch.equal(r["v_cache_f"][loc], v_ref[loc]), "target-slot V mismatch" + ) + + # Non-target slots must remain zero (no stray writes). + untouched = torch.ones( + r["k_cache_f"].shape[0], dtype=torch.bool, device=loc.device + ) + untouched[loc] = False + self.assertEqual( + r["k_cache_f"][untouched].float().abs().sum().item(), + 0.0, + "fused wrote into non-target K slots", + ) + self.assertEqual( + r["v_cache_f"][untouched].float().abs().sum().item(), + 0.0, + "fused wrote into non-target V slots", + ) + + +class _StopForward(Exception): + """Short-circuit forward_* after the KV write to skip the attention math.""" + + +@unittest.skipUnless(_HAS_CUDA, "Triton kernels require a GPU") +class TestAiterFp8KvDispatch(unittest.TestCase): + """Backend dispatch: scale passing (incl. v_scale fallback) and the + head-dim-mismatch fallback.""" + + def _make_backend(self, fp8_dtype, self_k_scale, self_v_scale): + from sglang.srt.layers.attention.aiter_backend import AiterAttnBackend + + be = AiterAttnBackend.__new__(AiterAttnBackend) + be.kv_cache_dtype = fp8_dtype + be.k_scale = self_k_scale + be.v_scale = self_v_scale + be.kv_cache_is_vectorized_5d = False + be.use_triton_unified_attention = False + be.use_sliding_window_kv_pool = False + be.use_mla = False + be.page_size = 1 + + class _Meta: + swa_out_cache_loc = None + + be.forward_metadata = _Meta() + return be + + def _make_layer(self, heads, qk_dim, v_dim): + class _Layer: + k_scale = None + v_scale = None + layer_id = 0 + tp_q_head_num = heads + tp_k_head_num = heads + tp_v_head_num = heads + qk_head_dim = qk_dim + v_head_dim = v_dim + is_cross_attention = False + + return _Layer() + + def _make_pool(self, num_slots, heads, qk_dim, v_dim, fp8_dtype, on_set): + kc = torch.zeros((num_slots, 1, heads, qk_dim), dtype=fp8_dtype, device="cuda") + vc = torch.zeros((num_slots, 1, heads, v_dim), dtype=fp8_dtype, device="cuda") + + class _Pool: + def get_kv_buffer(self, layer_id): + return kc, vc + + def set_kv_buffer(self, *args, **kwargs): + on_set(args, kwargs) + raise _StopForward + + return _Pool() + + def test_decode_uses_fused_and_v_scale_falls_back_to_self(self): + """Standard FP8 decode takes the fused path; with layer.v_scale=None, + v_scale falls back to self.v_scale (not self.k_scale).""" + from unittest import mock + + from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype + + dev = "cuda" + heads, dim, n = 2, 64, 4 + k_scale = torch.tensor([0.111], device=dev) + v_scale = torch.tensor([0.222], device=dev) + + be = self._make_backend(fp8_dtype, k_scale, v_scale) + be.token_to_kv_pool = self._make_pool( + n, heads, dim, dim, fp8_dtype, on_set=lambda a, k: None + ) + layer = self._make_layer(heads, dim, dim) # k_scale/v_scale None -> fallback + + q = torch.randn((n, heads, dim), dtype=torch.bfloat16, device=dev) + k = torch.randn((n, heads, dim), dtype=torch.bfloat16, device=dev) + v = torch.randn((n, heads, dim), dtype=torch.bfloat16, device=dev) + + class _FB: + out_cache_loc = torch.arange(n, device=dev, dtype=torch.int64) + + captured = {} + + def fake_launch(*args, **kwargs): + captured["kwargs"] = kwargs + raise _StopForward + + with mock.patch( + "sglang.srt.layers.attention.aiter_backend.launch_reshape_and_cache_flash", + fake_launch, + ): + try: + be.forward_decode(q, k, v, layer, _FB(), save_kv_cache=True) + except _StopForward: + pass + + self.assertIn("kwargs", captured, "fused path was not taken for standard FP8") + self.assertIs( + captured["kwargs"]["v_scale"], + v_scale, + "v_scale must fall back to self.v_scale, not self.k_scale", + ) + self.assertIs(captured["kwargs"]["k_scale"], k_scale) + + def test_decode_head_dim_mismatch_falls_back_to_set_kv_buffer(self): + """qk_head_dim != v_head_dim falls back to set_kv_buffer (kernel reuses + K's head_dim for V).""" + from unittest import mock + + from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype + + dev = "cuda" + heads, qk_dim, v_dim, n = 2, 64, 32, 4 # qk_head_dim != v_head_dim + k_scale = torch.tensor([0.111], device=dev) + v_scale = torch.tensor([0.222], device=dev) + be = self._make_backend(fp8_dtype, k_scale, v_scale) + set_calls = [] + be.token_to_kv_pool = self._make_pool( + n, + heads, + qk_dim, + v_dim, + fp8_dtype, + on_set=lambda a, k: set_calls.append((a, k)), + ) + layer = self._make_layer(heads, qk_dim, v_dim) + + q = torch.randn((n, heads, qk_dim), dtype=torch.bfloat16, device=dev) + k = torch.randn((n, heads, qk_dim), dtype=torch.bfloat16, device=dev) + v = torch.randn((n, heads, v_dim), dtype=torch.bfloat16, device=dev) + + class _FB: + out_cache_loc = torch.arange(n, device=dev, dtype=torch.int64) + + launch_calls = {"n": 0} + + def fake_launch(*args, **kwargs): + launch_calls["n"] += 1 + raise _StopForward + + with mock.patch( + "sglang.srt.layers.attention.aiter_backend.launch_reshape_and_cache_flash", + fake_launch, + ): + try: + be.forward_decode(q, k, v, layer, _FB(), save_kv_cache=True) + except _StopForward: + pass + + self.assertEqual( + launch_calls["n"], 0, "fused kernel must be skipped on head-dim mismatch" + ) + self.assertEqual( + len(set_calls), 1, "must fall back to set_kv_buffer on head-dim mismatch" + ) + # Fallback must still forward the descales, else the fp8 cache is wrong. + args, kwargs = set_calls[0] + passed = list(args) + list(kwargs.values()) + self.assertTrue(any(a is k_scale for a in passed), "fallback dropped k_scale") + self.assertTrue(any(a is v_scale for a in passed), "fallback dropped v_scale") + + def test_decode_mla_uses_set_kv_buffer_without_scales(self): + """MLA decode must call the MLA pool's set_kv_buffer(layer, loc, k, v) + without scale args (its signature takes no scales).""" + from unittest import mock + + from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype + + dev = "cuda" + heads, dim, n = 2, 64, 4 + be = self._make_backend( + fp8_dtype, + torch.tensor([0.111], device=dev), + torch.tensor([0.222], device=dev), + ) + be.use_mla = True + set_calls = [] + be.token_to_kv_pool = self._make_pool( + n, heads, dim, dim, fp8_dtype, on_set=lambda a, k: set_calls.append((a, k)) + ) + layer = self._make_layer(heads, dim, dim) + + q = torch.randn((n, heads, dim), dtype=torch.bfloat16, device=dev) + k = torch.randn((n, heads, dim), dtype=torch.bfloat16, device=dev) + v = torch.randn((n, heads, dim), dtype=torch.bfloat16, device=dev) + + class _FB: + out_cache_loc = torch.arange(n, device=dev, dtype=torch.int64) + + launch_calls = {"n": 0} + + def fake_launch(*args, **kwargs): + launch_calls["n"] += 1 + raise _StopForward + + with mock.patch( + "sglang.srt.layers.attention.aiter_backend.launch_reshape_and_cache_flash", + fake_launch, + ): + try: + be.forward_decode(q, k, v, layer, _FB(), save_kv_cache=True) + except _StopForward: + pass + + self.assertEqual(launch_calls["n"], 0, "MLA must not use the fused kernel") + self.assertEqual(len(set_calls), 1, "MLA must use set_kv_buffer") + args, kwargs = set_calls[0] + # MLA pool signature is (layer, loc, k, v) — exactly 4 args, no scales. + self.assertEqual( + len(args), 4, f"MLA set_kv_buffer got {len(args)} args (scales leaked?)" + ) + self.assertEqual(kwargs, {}) + + +class TestUseFusedFp8KvWritePredicate(unittest.TestCase): + """Unit test for the shared _use_fused_fp8_kv_write predicate.""" + + def _backend(self, **overrides): + from sglang.srt.layers.attention.aiter_backend import AiterAttnBackend + from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype + + be = AiterAttnBackend.__new__(AiterAttnBackend) + be.kv_cache_dtype = overrides.get("kv_cache_dtype", fp8_dtype) + be.use_mla = overrides.get("use_mla", False) + be.use_sliding_window_kv_pool = overrides.get( + "use_sliding_window_kv_pool", False + ) + return be + + def _layer(self, tp_k=2, tp_v=2, qk_dim=64, v_dim=64): + class _Layer: + tp_k_head_num = tp_k + tp_v_head_num = tp_v + qk_head_dim = qk_dim + v_head_dim = v_dim + + return _Layer() + + def test_predicate(self): + # Eligible: FP8 + non-MLA + non-SWA + matching K/V heads. + self.assertTrue(self._backend()._use_fused_fp8_kv_write(self._layer())) + self.assertFalse( + self._backend(kv_cache_dtype=torch.bfloat16)._use_fused_fp8_kv_write( + self._layer() + ) + ) + self.assertFalse( + self._backend(use_mla=True)._use_fused_fp8_kv_write(self._layer()) + ) + self.assertFalse( + self._backend(use_sliding_window_kv_pool=True)._use_fused_fp8_kv_write( + self._layer() + ) + ) + self.assertFalse( + self._backend()._use_fused_fp8_kv_write(self._layer(qk_dim=64, v_dim=32)) + ) + self.assertFalse( + self._backend()._use_fused_fp8_kv_write(self._layer(tp_k=2, tp_v=1)) + ) + + +if __name__ == "__main__": + unittest.main()