feat: Add FP8 KV cache support for Triton attention backend (#18882)
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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),
|
||||
|
||||
Reference in New Issue
Block a user