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
|
||||||
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,8 +815,23 @@ 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:
|
||||||
|
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(
|
forward_batch.token_to_kv_pool.set_kv_buffer(
|
||||||
layer, forward_batch.out_cache_loc, k, v
|
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,8 +1052,21 @@ 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:
|
||||||
|
if self.use_mla: # Triton MLA currently doesn't support quantized kv cache
|
||||||
forward_batch.token_to_kv_pool.set_kv_buffer(
|
forward_batch.token_to_kv_pool.set_kv_buffer(
|
||||||
layer, forward_batch.out_cache_loc, k, v
|
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:
|
||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user