feat: Add FP8 KV cache support for Triton attention backend (#18882)

This commit is contained in:
Zack Yu
2026-03-02 23:38:34 -08:00
committed by GitHub
parent 62480ebb1b
commit 07b8d763ef
6 changed files with 180 additions and 27 deletions
@@ -7,6 +7,7 @@ import torch
import triton import triton
import triton.language as tl 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.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton
from sglang.srt.layers.dp_attention import get_attention_tp_size 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.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.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens
self.speculative_num_steps = model_runner.server_args.speculative_num_steps 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 = ( self.num_head = (
model_runner.model_config.num_attention_heads // get_attention_tp_size() 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) # Save KV cache first (must do this before unified kernel)
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer( if (
layer, forward_batch.out_cache_loc, k, v 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) 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 kv_indices = self.forward_metadata.kv_indices
window_kv_offsets = None 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( self.extend_attention_fwd(
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
k.contiguous(), k.contiguous(),
@@ -864,6 +888,8 @@ class TritonAttnBackend(AttentionBackend):
causal, causal,
self.forward_metadata.mask_indptr, self.forward_metadata.mask_indptr,
self.forward_metadata.max_extend_len, self.forward_metadata.max_extend_len,
k_descale,
v_descale,
layer.scaling, layer.scaling,
logit_cap=logits_soft_cap, logit_cap=logits_soft_cap,
sliding_window_size=sliding_window_size, sliding_window_size=sliding_window_size,
@@ -970,12 +996,21 @@ class TritonAttnBackend(AttentionBackend):
# Convert prefix_lens to int32 for the kernel # Convert prefix_lens to int32 for the kernel
prefix_lens = prefix_lens.to(torch.int32) 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 # Call unified kernel
self.extend_attention_fwd_unified( self.extend_attention_fwd_unified(
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
o.view(-1, layer.tp_q_head_num, layer.v_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_key_buffer(layer.layer_id),
forward_batch.token_to_kv_pool.get_value_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, self.forward_metadata.qo_indptr,
unified_kv_indptr, unified_kv_indptr,
unified_kv_indices, unified_kv_indices,
@@ -1017,9 +1052,22 @@ class TritonAttnBackend(AttentionBackend):
logits_soft_cap = logit_capping_mod(layer.logit_capping_method, layer.logit_cap) logits_soft_cap = logit_capping_mod(layer.logit_capping_method, layer.logit_cap)
if save_kv_cache: if save_kv_cache:
forward_batch.token_to_kv_pool.set_kv_buffer( if self.use_mla: # Triton MLA currently doesn't support quantized kv cache
layer, forward_batch.out_cache_loc, k, v 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: if layer.sliding_window_size is not None and layer.sliding_window_size > -1:
kv_indptr = self.forward_metadata.window_kv_indptr kv_indptr = self.forward_metadata.window_kv_indptr
@@ -1028,6 +1076,13 @@ class TritonAttnBackend(AttentionBackend):
kv_indptr = self.forward_metadata.kv_indptr kv_indptr = self.forward_metadata.kv_indptr
kv_indices = self.forward_metadata.kv_indices 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( self.decode_attention_fwd(
q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id), 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.forward_metadata.num_kv_splits,
self.max_kv_splits, self.max_kv_splits,
layer.scaling, layer.scaling,
k_descale,
v_descale,
logit_cap=logits_soft_cap, logit_cap=logits_soft_cap,
sinks=sinks, sinks=sinks,
xai_temperature_len=layer.xai_temperature_len, xai_temperature_len=layer.xai_temperature_len,
@@ -46,7 +46,7 @@ def _fwd_kernel_stage1(
Q, Q,
K_Buffer, K_Buffer,
V_Buffer, V_Buffer,
sm_scale, sm_scale_withk,
kv_indptr, kv_indptr,
kv_indices, kv_indices,
Att_Out, Att_Out,
@@ -124,7 +124,7 @@ def _fwd_kernel_stage1(
other=0.0, other=0.0,
) )
qk = tl.sum(q[None, :] * k, 1) qk = tl.sum(q[None, :] * k, 1)
qk *= sm_scale qk *= sm_scale_withk
if logit_cap > 0: if logit_cap > 0:
qk = logit_cap * tanh(qk / logit_cap) qk = logit_cap * tanh(qk / logit_cap)
@@ -189,7 +189,7 @@ def _decode_att_m_fwd(
kv_indices, kv_indices,
num_kv_splits, num_kv_splits,
max_kv_splits, max_kv_splits,
sm_scale, sm_scale_withk,
logit_cap, logit_cap,
xai_temperature_len=-1, xai_temperature_len=-1,
): ):
@@ -220,7 +220,7 @@ def _decode_att_m_fwd(
q, q,
k_buffer, k_buffer,
v_buffer, v_buffer,
sm_scale, sm_scale_withk,
kv_indptr, kv_indptr,
kv_indices, kv_indices,
att_out, att_out,
@@ -254,7 +254,7 @@ def _fwd_grouped_kernel_stage1(
Q, Q,
K_Buffer, K_Buffer,
V_Buffer, V_Buffer,
sm_scale, sm_scale_withk,
kv_indptr, kv_indptr,
kv_indices, kv_indices,
Att_Out, Att_Out,
@@ -365,7 +365,7 @@ def _fwd_grouped_kernel_stage1(
other=0.0, other=0.0,
) )
qk += tl.dot(qpe, kpe.to(qpe.dtype)) qk += tl.dot(qpe, kpe.to(qpe.dtype))
qk *= sm_scale qk *= sm_scale_withk
if logit_cap > 0: if logit_cap > 0:
qk = logit_cap * tanh(qk / logit_cap) qk = logit_cap * tanh(qk / logit_cap)
@@ -433,7 +433,7 @@ def _decode_grouped_att_m_fwd(
kv_indices, kv_indices,
num_kv_splits, num_kv_splits,
max_kv_splits, max_kv_splits,
sm_scale, sm_scale_withk,
logit_cap, logit_cap,
xai_temperature_len=-1, xai_temperature_len=-1,
): ):
@@ -479,7 +479,7 @@ def _decode_grouped_att_m_fwd(
q, q,
k_buffer, k_buffer,
v_buffer, v_buffer,
sm_scale, sm_scale_withk,
kv_indptr, kv_indptr,
kv_indices, kv_indices,
att_out, att_out,
@@ -517,6 +517,7 @@ def _fwd_kernel_stage2(
Mid_O, Mid_O,
Mid_O_1, Mid_O_1,
O, O,
v_scale,
kv_indptr, kv_indptr,
num_kv_splits, num_kv_splits,
sink_ptr, sink_ptr,
@@ -577,7 +578,7 @@ def _fwd_kernel_stage2(
tl.store( tl.store(
O + cur_batch * stride_obs + cur_head * stride_oh + offs_d, O + cur_batch * stride_obs + cur_head * stride_oh + offs_d,
acc / e_sum, acc / e_sum * v_scale,
mask=mask_d, mask=mask_d,
) )
@@ -587,6 +588,7 @@ def _decode_softmax_reducev_fwd(
lse, lse,
q, q,
o, o,
v_scale,
v_buffer, v_buffer,
kv_indptr, kv_indptr,
num_kv_splits, num_kv_splits,
@@ -611,6 +613,7 @@ def _decode_softmax_reducev_fwd(
logits, logits,
lse, lse,
o, o,
v_scale,
kv_indptr, kv_indptr,
num_kv_splits, num_kv_splits,
sinks, sinks,
@@ -641,7 +644,8 @@ def decode_attention_fwd_normal(
attn_lse, attn_lse,
num_kv_splits, num_kv_splits,
max_kv_splits, max_kv_splits,
sm_scale, sm_scale_withk,
v_scale,
logit_cap=0.0, logit_cap=0.0,
sinks=None, sinks=None,
xai_temperature_len=-1, xai_temperature_len=-1,
@@ -656,7 +660,7 @@ def decode_attention_fwd_normal(
kv_indices, kv_indices,
num_kv_splits, num_kv_splits,
max_kv_splits, max_kv_splits,
sm_scale, sm_scale_withk,
logit_cap, logit_cap,
xai_temperature_len, xai_temperature_len,
) )
@@ -665,6 +669,7 @@ def decode_attention_fwd_normal(
attn_lse, attn_lse,
q, q,
o, o,
v_scale,
v_buffer, v_buffer,
kv_indptr, kv_indptr,
num_kv_splits, num_kv_splits,
@@ -684,7 +689,8 @@ def decode_attention_fwd_grouped(
attn_lse, attn_lse,
num_kv_splits, num_kv_splits,
max_kv_splits, max_kv_splits,
sm_scale, sm_scale_withk,
v_scale,
logit_cap=0.0, logit_cap=0.0,
sinks=None, sinks=None,
xai_temperature_len=-1, xai_temperature_len=-1,
@@ -699,7 +705,7 @@ def decode_attention_fwd_grouped(
kv_indices, kv_indices,
num_kv_splits, num_kv_splits,
max_kv_splits, max_kv_splits,
sm_scale, sm_scale_withk,
logit_cap, logit_cap,
xai_temperature_len, xai_temperature_len,
) )
@@ -708,6 +714,7 @@ def decode_attention_fwd_grouped(
attn_lse, attn_lse,
q, q,
o, o,
v_scale,
v_buffer, v_buffer,
kv_indptr, kv_indptr,
num_kv_splits, num_kv_splits,
@@ -728,6 +735,8 @@ def decode_attention_fwd(
num_kv_splits, num_kv_splits,
max_kv_splits, max_kv_splits,
sm_scale, sm_scale,
k_scale,
v_scale,
logit_cap=0.0, logit_cap=0.0,
sinks=None, sinks=None,
xai_temperature_len=-1, xai_temperature_len=-1,
@@ -751,7 +760,8 @@ def decode_attention_fwd(
attn_lse, attn_lse,
num_kv_splits, num_kv_splits,
max_kv_splits, max_kv_splits,
sm_scale, sm_scale * k_scale,
v_scale,
logit_cap=logit_cap, logit_cap=logit_cap,
sinks=sinks, sinks=sinks,
xai_temperature_len=xai_temperature_len, xai_temperature_len=xai_temperature_len,
@@ -769,7 +779,8 @@ def decode_attention_fwd(
attn_lse, attn_lse,
num_kv_splits, num_kv_splits,
max_kv_splits, max_kv_splits,
sm_scale, sm_scale * k_scale,
v_scale,
logit_cap=logit_cap, logit_cap=logit_cap,
sinks=sinks, sinks=sinks,
xai_temperature_len=xai_temperature_len, xai_temperature_len=xai_temperature_len,
@@ -232,6 +232,8 @@ def _fwd_kernel(
sink_ptr, sink_ptr,
window_kv_offset_ptr, window_kv_offset_ptr,
sm_scale, sm_scale,
k_scale,
v_scale,
kv_group_num, kv_group_num,
stride_qbs, stride_qbs,
stride_qh, stride_qh,
@@ -386,7 +388,7 @@ def _fwd_kernel(
other=0.0, other=0.0,
) )
qk += tl.dot(qpe.to(kpe.dtype), kpe) qk += tl.dot(qpe.to(kpe.dtype), kpe)
qk *= sm_scale qk *= sm_scale * k_scale
if logit_cap > 0: if logit_cap > 0:
qk = logit_cap * tanh(qk / logit_cap) qk = logit_cap * tanh(qk / logit_cap)
@@ -415,7 +417,7 @@ def _fwd_kernel(
other=0.0, other=0.0,
) )
p = p.to(v.dtype) 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 e_max = n_e_max
@@ -561,6 +563,8 @@ def extend_attention_fwd(
is_causal, is_causal,
mask_indptr, mask_indptr,
max_len_extend, max_len_extend,
k_scale,
v_scale,
sm_scale=None, sm_scale=None,
logit_cap=0.0, logit_cap=0.0,
skip_prefix_custom_mask=True, skip_prefix_custom_mask=True,
@@ -617,6 +621,8 @@ def extend_attention_fwd(
sinks, sinks,
window_kv_offsets, window_kv_offsets,
sm_scale, sm_scale,
k_scale,
v_scale,
kv_group_num, kv_group_num,
q_extend.stride(0), q_extend.stride(0),
q_extend.stride(1), q_extend.stride(1),
@@ -702,7 +708,8 @@ def _fwd_kernel_unified(
mask_indptr, mask_indptr,
sink_ptr, sink_ptr,
window_start_pos, window_start_pos,
sm_scale, sm_scale_withk,
v_scale,
kv_group_num, kv_group_num,
stride_qbs, stride_qbs,
stride_qh, stride_qh,
@@ -887,7 +894,7 @@ def _fwd_kernel_unified(
) )
qk += tl.dot(qpe.to(kpe.dtype), kpe) qk += tl.dot(qpe.to(kpe.dtype), kpe)
qk *= sm_scale qk *= sm_scale_withk
if logit_cap > 0: if logit_cap > 0:
qk = logit_cap * tanh(qk / logit_cap) qk = logit_cap * tanh(qk / logit_cap)
@@ -935,7 +942,7 @@ def _fwd_kernel_unified(
) )
tl.store( tl.store(
O + offs_o, O + offs_o,
acc / deno[:, None], acc / deno[:, None] * v_scale,
mask=mask_m[:, None] & mask_dv[None, :], mask=mask_m[:, None] & mask_dv[None, :],
) )
@@ -945,6 +952,8 @@ def extend_attention_fwd_unified(
o, o,
k_buffer, k_buffer,
v_buffer, v_buffer,
k_scale,
v_scale,
qo_indptr, qo_indptr,
kv_indptr, kv_indptr,
kv_indices, kv_indices,
@@ -1024,7 +1033,8 @@ def extend_attention_fwd_unified(
mask_indptr, mask_indptr,
sinks, sinks,
window_start_pos, window_start_pos,
sm_scale, sm_scale * k_scale,
v_scale,
kv_group_num, kv_group_num,
q.stride(0), q.stride(0),
q.stride(1), q.stride(1),
@@ -251,6 +251,8 @@ class TestTritonAttention(CustomTestCase):
True, True,
mask_indptr, mask_indptr,
max_len_extend, max_len_extend,
1.0,
1.0,
) )
b_seq_mask_len = b_seq_len_extend * b_seq_len b_seq_mask_len = b_seq_len_extend * b_seq_len
@@ -286,6 +288,8 @@ class TestTritonAttention(CustomTestCase):
True, True,
mask_indptr, mask_indptr,
max_len_extend, max_len_extend,
1.0,
1.0,
) )
redundant_attention( redundant_attention(
@@ -395,6 +399,8 @@ class TestTritonAttention(CustomTestCase):
is_causal=True, is_causal=True,
mask_indptr=None, mask_indptr=None,
max_len_extend=max_len_extend, max_len_extend=max_len_extend,
k_scale=1.0,
v_scale=1.0,
sliding_window_size=WINDOW_SIZE, sliding_window_size=WINDOW_SIZE,
) )
@@ -517,6 +523,8 @@ class TestTritonAttention(CustomTestCase):
num_kv_splits, num_kv_splits,
max_kv_splits, max_kv_splits,
sm_scale, sm_scale,
1.0,
1.0,
) )
# Correctness reference (float32, stable softmax) # Correctness reference (float32, stable softmax)
@@ -591,6 +599,7 @@ class TestTritonAttention(CustomTestCase):
num_kv_splits, num_kv_splits,
max_kv_splits, max_kv_splits,
sm_scale, sm_scale,
1.0,
) )
attn_logits1 = torch.empty( attn_logits1 = torch.empty(
@@ -616,6 +625,7 @@ class TestTritonAttention(CustomTestCase):
num_kv_splits, num_kv_splits,
max_kv_splits, max_kv_splits,
sm_scale, sm_scale,
1.0,
) )
cos_sim = torch.nn.functional.cosine_similarity( cos_sim = torch.nn.functional.cosine_similarity(
@@ -722,6 +732,8 @@ class TestTritonAttention(CustomTestCase):
is_causal=True, is_causal=True,
mask_indptr=None, mask_indptr=None,
max_len_extend=max_len_extend, max_len_extend=max_len_extend,
k_scale=1.0,
v_scale=1.0,
) )
# Build unified KV indices # Build unified KV indices
@@ -750,6 +762,8 @@ class TestTritonAttention(CustomTestCase):
o_unified, o_unified,
k_buffer, k_buffer,
v_buffer, v_buffer,
1.0,
1.0,
qo_indptr, qo_indptr,
unified_kv_indptr, unified_kv_indptr,
unified_kv_indices, unified_kv_indices,
@@ -155,6 +155,8 @@ class TestWaveAttention(unittest.TestCase):
is_causal, is_causal,
mask_indptr, mask_indptr,
max_len_extend, max_len_extend,
1.0,
1.0,
) )
o_wave = torch.empty( o_wave = torch.empty(
@@ -240,6 +242,7 @@ class TestWaveAttention(unittest.TestCase):
num_kv_splits, num_kv_splits,
max_kv_splits, max_kv_splits,
sm_scale, sm_scale,
1.0,
logit_cap, logit_cap,
) )
@@ -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()