From 07b8d763ef0b096508d26a4c800738d9cdea852b Mon Sep 17 00:00:00 2001 From: Zack Yu Date: Mon, 2 Mar 2026 23:38:34 -0800 Subject: [PATCH] feat: Add FP8 KV cache support for Triton attention backend (#18882) --- .../srt/layers/attention/triton_backend.py | 69 +++++++++++++++++-- .../attention/triton_ops/decode_attention.py | 41 +++++++---- .../attention/triton_ops/extend_attention.py | 22 ++++-- .../test_triton_attention_kernels.py | 14 ++++ .../attention/test_wave_attention_kernels.py | 3 + test/registered/quant/test_fp8kv_triton.py | 58 ++++++++++++++++ 6 files changed, 180 insertions(+), 27 deletions(-) create mode 100644 test/registered/quant/test_fp8kv_triton.py diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index 657777de1..611680c78 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -7,6 +7,7 @@ import torch import triton import triton.language as tl +from sglang.srt.configs.model_config import AttentionArch from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton from sglang.srt.layers.dp_attention import get_attention_tp_size @@ -86,6 +87,7 @@ class TritonAttnBackend(AttentionBackend): self.token_to_kv_pool_allocator = model_runner.token_to_kv_pool_allocator self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens self.speculative_num_steps = model_runner.server_args.speculative_num_steps + self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA self.num_head = ( model_runner.model_config.num_attention_heads // get_attention_tp_size() ) @@ -813,9 +815,24 @@ class TritonAttnBackend(AttentionBackend): # Save KV cache first (must do this before unified kernel) if save_kv_cache: - forward_batch.token_to_kv_pool.set_kv_buffer( - layer, forward_batch.out_cache_loc, k, v - ) + if ( + self.use_mla or layer.k_scale is None + ): # Triton MLA currently doesn't support quantized kv cache + forward_batch.token_to_kv_pool.set_kv_buffer( + layer, + forward_batch.out_cache_loc, + k, + v, + ) + else: + forward_batch.token_to_kv_pool.set_kv_buffer( + layer, + forward_batch.out_cache_loc, + k.clone(), # cloned to protect k,v from in-place mutation in set_kv_buffer + v.clone(), + layer.k_scale, + layer.v_scale, + ) logits_soft_cap = logit_capping_mod(layer.logit_capping_method, layer.logit_cap) @@ -850,6 +867,13 @@ class TritonAttnBackend(AttentionBackend): kv_indices = self.forward_metadata.kv_indices window_kv_offsets = None + if layer.k_scale is not None and layer.v_scale is not None: + k_descale = layer.k_scale_float + v_descale = layer.v_scale_float + else: + k_descale = 1.0 + v_descale = 1.0 + self.extend_attention_fwd( q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), k.contiguous(), @@ -864,6 +888,8 @@ class TritonAttnBackend(AttentionBackend): causal, self.forward_metadata.mask_indptr, self.forward_metadata.max_extend_len, + k_descale, + v_descale, layer.scaling, logit_cap=logits_soft_cap, sliding_window_size=sliding_window_size, @@ -970,12 +996,21 @@ class TritonAttnBackend(AttentionBackend): # Convert prefix_lens to int32 for the kernel prefix_lens = prefix_lens.to(torch.int32) + if layer.k_scale is not None and layer.v_scale is not None: + k_descale = layer.k_scale_float + v_descale = layer.v_scale_float + else: + k_descale = 1.0 + v_descale = 1.0 + # Call unified kernel self.extend_attention_fwd_unified( q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), o.view(-1, layer.tp_q_head_num, layer.v_head_dim), forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id), forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id), + k_descale, + v_descale, self.forward_metadata.qo_indptr, unified_kv_indptr, unified_kv_indices, @@ -1017,9 +1052,22 @@ class TritonAttnBackend(AttentionBackend): logits_soft_cap = logit_capping_mod(layer.logit_capping_method, layer.logit_cap) if save_kv_cache: - forward_batch.token_to_kv_pool.set_kv_buffer( - layer, forward_batch.out_cache_loc, k, v - ) + if self.use_mla: # Triton MLA currently doesn't support quantized kv cache + forward_batch.token_to_kv_pool.set_kv_buffer( + layer, + forward_batch.out_cache_loc, + k, + v, + ) + else: + forward_batch.token_to_kv_pool.set_kv_buffer( + layer, + forward_batch.out_cache_loc, + k, + v, + layer.k_scale, + layer.v_scale, + ) if layer.sliding_window_size is not None and layer.sliding_window_size > -1: kv_indptr = self.forward_metadata.window_kv_indptr @@ -1028,6 +1076,13 @@ class TritonAttnBackend(AttentionBackend): kv_indptr = self.forward_metadata.kv_indptr kv_indices = self.forward_metadata.kv_indices + if layer.k_scale is not None and layer.v_scale is not None: + k_descale = layer.k_scale_float + v_descale = layer.v_scale_float + else: + k_descale = 1.0 + v_descale = 1.0 + self.decode_attention_fwd( q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id), @@ -1040,6 +1095,8 @@ class TritonAttnBackend(AttentionBackend): self.forward_metadata.num_kv_splits, self.max_kv_splits, layer.scaling, + k_descale, + v_descale, logit_cap=logits_soft_cap, sinks=sinks, xai_temperature_len=layer.xai_temperature_len, diff --git a/python/sglang/srt/layers/attention/triton_ops/decode_attention.py b/python/sglang/srt/layers/attention/triton_ops/decode_attention.py index 1ba5d463d..2b166f3b0 100644 --- a/python/sglang/srt/layers/attention/triton_ops/decode_attention.py +++ b/python/sglang/srt/layers/attention/triton_ops/decode_attention.py @@ -46,7 +46,7 @@ def _fwd_kernel_stage1( Q, K_Buffer, V_Buffer, - sm_scale, + sm_scale_withk, kv_indptr, kv_indices, Att_Out, @@ -124,7 +124,7 @@ def _fwd_kernel_stage1( other=0.0, ) qk = tl.sum(q[None, :] * k, 1) - qk *= sm_scale + qk *= sm_scale_withk if logit_cap > 0: qk = logit_cap * tanh(qk / logit_cap) @@ -189,7 +189,7 @@ def _decode_att_m_fwd( kv_indices, num_kv_splits, max_kv_splits, - sm_scale, + sm_scale_withk, logit_cap, xai_temperature_len=-1, ): @@ -220,7 +220,7 @@ def _decode_att_m_fwd( q, k_buffer, v_buffer, - sm_scale, + sm_scale_withk, kv_indptr, kv_indices, att_out, @@ -254,7 +254,7 @@ def _fwd_grouped_kernel_stage1( Q, K_Buffer, V_Buffer, - sm_scale, + sm_scale_withk, kv_indptr, kv_indices, Att_Out, @@ -365,7 +365,7 @@ def _fwd_grouped_kernel_stage1( other=0.0, ) qk += tl.dot(qpe, kpe.to(qpe.dtype)) - qk *= sm_scale + qk *= sm_scale_withk if logit_cap > 0: qk = logit_cap * tanh(qk / logit_cap) @@ -433,7 +433,7 @@ def _decode_grouped_att_m_fwd( kv_indices, num_kv_splits, max_kv_splits, - sm_scale, + sm_scale_withk, logit_cap, xai_temperature_len=-1, ): @@ -479,7 +479,7 @@ def _decode_grouped_att_m_fwd( q, k_buffer, v_buffer, - sm_scale, + sm_scale_withk, kv_indptr, kv_indices, att_out, @@ -517,6 +517,7 @@ def _fwd_kernel_stage2( Mid_O, Mid_O_1, O, + v_scale, kv_indptr, num_kv_splits, sink_ptr, @@ -577,7 +578,7 @@ def _fwd_kernel_stage2( tl.store( O + cur_batch * stride_obs + cur_head * stride_oh + offs_d, - acc / e_sum, + acc / e_sum * v_scale, mask=mask_d, ) @@ -587,6 +588,7 @@ def _decode_softmax_reducev_fwd( lse, q, o, + v_scale, v_buffer, kv_indptr, num_kv_splits, @@ -611,6 +613,7 @@ def _decode_softmax_reducev_fwd( logits, lse, o, + v_scale, kv_indptr, num_kv_splits, sinks, @@ -641,7 +644,8 @@ def decode_attention_fwd_normal( attn_lse, num_kv_splits, max_kv_splits, - sm_scale, + sm_scale_withk, + v_scale, logit_cap=0.0, sinks=None, xai_temperature_len=-1, @@ -656,7 +660,7 @@ def decode_attention_fwd_normal( kv_indices, num_kv_splits, max_kv_splits, - sm_scale, + sm_scale_withk, logit_cap, xai_temperature_len, ) @@ -665,6 +669,7 @@ def decode_attention_fwd_normal( attn_lse, q, o, + v_scale, v_buffer, kv_indptr, num_kv_splits, @@ -684,7 +689,8 @@ def decode_attention_fwd_grouped( attn_lse, num_kv_splits, max_kv_splits, - sm_scale, + sm_scale_withk, + v_scale, logit_cap=0.0, sinks=None, xai_temperature_len=-1, @@ -699,7 +705,7 @@ def decode_attention_fwd_grouped( kv_indices, num_kv_splits, max_kv_splits, - sm_scale, + sm_scale_withk, logit_cap, xai_temperature_len, ) @@ -708,6 +714,7 @@ def decode_attention_fwd_grouped( attn_lse, q, o, + v_scale, v_buffer, kv_indptr, num_kv_splits, @@ -728,6 +735,8 @@ def decode_attention_fwd( num_kv_splits, max_kv_splits, sm_scale, + k_scale, + v_scale, logit_cap=0.0, sinks=None, xai_temperature_len=-1, @@ -751,7 +760,8 @@ def decode_attention_fwd( attn_lse, num_kv_splits, max_kv_splits, - sm_scale, + sm_scale * k_scale, + v_scale, logit_cap=logit_cap, sinks=sinks, xai_temperature_len=xai_temperature_len, @@ -769,7 +779,8 @@ def decode_attention_fwd( attn_lse, num_kv_splits, max_kv_splits, - sm_scale, + sm_scale * k_scale, + v_scale, logit_cap=logit_cap, sinks=sinks, xai_temperature_len=xai_temperature_len, diff --git a/python/sglang/srt/layers/attention/triton_ops/extend_attention.py b/python/sglang/srt/layers/attention/triton_ops/extend_attention.py index 2fb428fb2..8ce0e35ff 100644 --- a/python/sglang/srt/layers/attention/triton_ops/extend_attention.py +++ b/python/sglang/srt/layers/attention/triton_ops/extend_attention.py @@ -232,6 +232,8 @@ def _fwd_kernel( sink_ptr, window_kv_offset_ptr, sm_scale, + k_scale, + v_scale, kv_group_num, stride_qbs, stride_qh, @@ -386,7 +388,7 @@ def _fwd_kernel( other=0.0, ) qk += tl.dot(qpe.to(kpe.dtype), kpe) - qk *= sm_scale + qk *= sm_scale * k_scale if logit_cap > 0: qk = logit_cap * tanh(qk / logit_cap) @@ -415,7 +417,7 @@ def _fwd_kernel( other=0.0, ) p = p.to(v.dtype) - acc = acc * re_scale[:, None] + tl.dot(p, v) + acc = acc * re_scale[:, None] + tl.dot(p, v) * v_scale e_max = n_e_max @@ -561,6 +563,8 @@ def extend_attention_fwd( is_causal, mask_indptr, max_len_extend, + k_scale, + v_scale, sm_scale=None, logit_cap=0.0, skip_prefix_custom_mask=True, @@ -617,6 +621,8 @@ def extend_attention_fwd( sinks, window_kv_offsets, sm_scale, + k_scale, + v_scale, kv_group_num, q_extend.stride(0), q_extend.stride(1), @@ -702,7 +708,8 @@ def _fwd_kernel_unified( mask_indptr, sink_ptr, window_start_pos, - sm_scale, + sm_scale_withk, + v_scale, kv_group_num, stride_qbs, stride_qh, @@ -887,7 +894,7 @@ def _fwd_kernel_unified( ) qk += tl.dot(qpe.to(kpe.dtype), kpe) - qk *= sm_scale + qk *= sm_scale_withk if logit_cap > 0: qk = logit_cap * tanh(qk / logit_cap) @@ -935,7 +942,7 @@ def _fwd_kernel_unified( ) tl.store( O + offs_o, - acc / deno[:, None], + acc / deno[:, None] * v_scale, mask=mask_m[:, None] & mask_dv[None, :], ) @@ -945,6 +952,8 @@ def extend_attention_fwd_unified( o, k_buffer, v_buffer, + k_scale, + v_scale, qo_indptr, kv_indptr, kv_indices, @@ -1024,7 +1033,8 @@ def extend_attention_fwd_unified( mask_indptr, sinks, window_start_pos, - sm_scale, + sm_scale * k_scale, + v_scale, kv_group_num, q.stride(0), q.stride(1), diff --git a/test/registered/attention/test_triton_attention_kernels.py b/test/registered/attention/test_triton_attention_kernels.py index 80fc86d2e..b38e673b4 100644 --- a/test/registered/attention/test_triton_attention_kernels.py +++ b/test/registered/attention/test_triton_attention_kernels.py @@ -251,6 +251,8 @@ class TestTritonAttention(CustomTestCase): True, mask_indptr, max_len_extend, + 1.0, + 1.0, ) b_seq_mask_len = b_seq_len_extend * b_seq_len @@ -286,6 +288,8 @@ class TestTritonAttention(CustomTestCase): True, mask_indptr, max_len_extend, + 1.0, + 1.0, ) redundant_attention( @@ -395,6 +399,8 @@ class TestTritonAttention(CustomTestCase): is_causal=True, mask_indptr=None, max_len_extend=max_len_extend, + k_scale=1.0, + v_scale=1.0, sliding_window_size=WINDOW_SIZE, ) @@ -517,6 +523,8 @@ class TestTritonAttention(CustomTestCase): num_kv_splits, max_kv_splits, sm_scale, + 1.0, + 1.0, ) # Correctness reference (float32, stable softmax) @@ -591,6 +599,7 @@ class TestTritonAttention(CustomTestCase): num_kv_splits, max_kv_splits, sm_scale, + 1.0, ) attn_logits1 = torch.empty( @@ -616,6 +625,7 @@ class TestTritonAttention(CustomTestCase): num_kv_splits, max_kv_splits, sm_scale, + 1.0, ) cos_sim = torch.nn.functional.cosine_similarity( @@ -722,6 +732,8 @@ class TestTritonAttention(CustomTestCase): is_causal=True, mask_indptr=None, max_len_extend=max_len_extend, + k_scale=1.0, + v_scale=1.0, ) # Build unified KV indices @@ -750,6 +762,8 @@ class TestTritonAttention(CustomTestCase): o_unified, k_buffer, v_buffer, + 1.0, + 1.0, qo_indptr, unified_kv_indptr, unified_kv_indices, diff --git a/test/registered/attention/test_wave_attention_kernels.py b/test/registered/attention/test_wave_attention_kernels.py index fbd347048..93ac52e3c 100644 --- a/test/registered/attention/test_wave_attention_kernels.py +++ b/test/registered/attention/test_wave_attention_kernels.py @@ -155,6 +155,8 @@ class TestWaveAttention(unittest.TestCase): is_causal, mask_indptr, max_len_extend, + 1.0, + 1.0, ) o_wave = torch.empty( @@ -240,6 +242,7 @@ class TestWaveAttention(unittest.TestCase): num_kv_splits, max_kv_splits, sm_scale, + 1.0, logit_cap, ) diff --git a/test/registered/quant/test_fp8kv_triton.py b/test/registered/quant/test_fp8kv_triton.py new file mode 100644 index 000000000..7fb81f3e3 --- /dev/null +++ b/test/registered/quant/test_fp8kv_triton.py @@ -0,0 +1,58 @@ +import unittest +from types import SimpleNamespace +from urllib.parse import urlparse + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.few_shot_gsm8k import run_eval +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +register_cuda_ci(est_time=520, suite="stage-b-test-large-1-gpu") + + +class TestFP8KVCacheTritonBackend(CustomTestCase): + @classmethod + def setUpClass(cls): + cls.model = "neuralmagic/Meta-Llama-3-8B-Instruct-FP8-KV" + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--quantization", + "fp8", + "--kv-cache-dtype", + "fp8_e4m3", + "--attention-backend", + "triton", + ], + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + def test_gsm8k(self): + parsed_url = urlparse(self.base_url) + args = SimpleNamespace( + num_shots=5, + data_path=None, + num_questions=200, + max_new_tokens=512, + parallel=200, + host=f"{parsed_url.scheme}://{parsed_url.hostname}", + port=parsed_url.port, + ) + metrics = run_eval(args) + print(f"{metrics=}") + self.assertGreater(metrics["accuracy"], 0.70) + + +if __name__ == "__main__": + unittest.main()