Support Triton MLA FP8 KV cache (#20479)
Co-authored-by: b8zhong <b8zhong@users.noreply.github.com>
This commit is contained in:
@@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, List, Optional
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
from sgl_kernel.utils import is_arch_support_pdl
|
||||
|
||||
from sglang.srt.configs.model_config import AttentionArch
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
@@ -28,6 +29,19 @@ if TYPE_CHECKING:
|
||||
from sglang.srt.speculative.spec_info import SpecInput
|
||||
|
||||
|
||||
_MLA_DECODE_MIN_BLOCK_KV = 32
|
||||
|
||||
|
||||
def _mla_decode_kv_splits_cap(
|
||||
base_max_kv_splits: int, sm_count: int, max_context_len: int
|
||||
) -> int:
|
||||
if sm_count <= 0:
|
||||
return base_max_kv_splits
|
||||
sm_cap = next_power_of_2(sm_count)
|
||||
ctx_cap = next_power_of_2(triton.cdiv(max_context_len, _MLA_DECODE_MIN_BLOCK_KV))
|
||||
return max(base_max_kv_splits, min(sm_cap, ctx_cap))
|
||||
|
||||
|
||||
def logit_capping_mod(logit_capping_method, logit_cap):
|
||||
# positive logit_cap -> tanh cap
|
||||
if logit_capping_method == "tanh":
|
||||
@@ -128,6 +142,13 @@ class TritonAttnBackend(AttentionBackend):
|
||||
"SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS", "false"
|
||||
)
|
||||
self.max_kv_splits = model_runner.server_args.triton_attention_num_kv_splits
|
||||
if self.use_mla:
|
||||
self.max_kv_splits = _mla_decode_kv_splits_cap(
|
||||
self.max_kv_splits,
|
||||
self.device_core_count,
|
||||
self.max_context_len,
|
||||
)
|
||||
self.use_pdl = is_arch_support_pdl()
|
||||
|
||||
self.allow_bidirectional_attention_in_extend = (
|
||||
model_runner.server_args.disable_cuda_graph
|
||||
@@ -887,15 +908,24 @@ class TritonAttnBackend(AttentionBackend):
|
||||
else:
|
||||
# Save KV cache first (must do this before unified kernel)
|
||||
if save_kv_cache:
|
||||
if (
|
||||
self.use_mla or layer.k_scale is None
|
||||
): # Triton MLA currently doesn't support quantized kv cache
|
||||
if layer.k_scale is None:
|
||||
forward_batch.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
forward_batch.out_cache_loc,
|
||||
k,
|
||||
v,
|
||||
)
|
||||
elif self.use_mla:
|
||||
# For MLA, scale K manually before storing since MLATokenToKVPool
|
||||
# doesn't accept scale parameters. Clone to protect k from mutation
|
||||
# since it's used later in the attention kernel.
|
||||
k_scaled = k.clone().div_(layer.k_scale)
|
||||
forward_batch.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
forward_batch.out_cache_loc,
|
||||
k_scaled,
|
||||
v,
|
||||
)
|
||||
else:
|
||||
forward_batch.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
@@ -1139,7 +1169,11 @@ class TritonAttnBackend(AttentionBackend):
|
||||
logits_soft_cap = logit_capping_mod(layer.logit_capping_method, layer.logit_cap)
|
||||
|
||||
if save_kv_cache:
|
||||
if self.use_mla: # Triton MLA currently doesn't support quantized kv cache
|
||||
if self.use_mla:
|
||||
if layer.k_scale is not None:
|
||||
# MLATokenToKVPool doesn't accept scale parameters; k is unused
|
||||
# after this point in decode, so scale in place.
|
||||
k.div_(layer.k_scale)
|
||||
forward_batch.token_to_kv_pool.set_kv_buffer(
|
||||
layer,
|
||||
forward_batch.out_cache_loc,
|
||||
@@ -1197,6 +1231,8 @@ class TritonAttnBackend(AttentionBackend):
|
||||
logit_cap=logits_soft_cap,
|
||||
sinks=sinks,
|
||||
xai_temperature_len=layer.xai_temperature_len,
|
||||
has_mla=self.use_mla,
|
||||
use_pdl=self.use_pdl,
|
||||
)
|
||||
return o
|
||||
|
||||
|
||||
@@ -281,6 +281,8 @@ def _fwd_grouped_kernel_stage1(
|
||||
xai_temperature_len: tl.constexpr,
|
||||
Lk: tl.constexpr,
|
||||
Lv: tl.constexpr,
|
||||
HAS_MLA: tl.constexpr = False,
|
||||
USE_PDL: tl.constexpr = False,
|
||||
):
|
||||
cur_batch = tl.program_id(0)
|
||||
cur_head_id = tl.program_id(1)
|
||||
@@ -329,36 +331,36 @@ def _fwd_grouped_kernel_stage1(
|
||||
e_sum = tl.zeros([BLOCK_H], dtype=tl.float32)
|
||||
acc = tl.zeros([BLOCK_H, BLOCK_DV], dtype=tl.float32)
|
||||
|
||||
# Hoist loop-invariant base offsets
|
||||
base_offs_k = cur_kv_head * stride_buf_kh + offs_d[:, None]
|
||||
if BLOCK_DPE > 0:
|
||||
base_offs_kpe = cur_kv_head * stride_buf_kh + offs_dpe[:, None]
|
||||
if not HAS_MLA:
|
||||
base_offs_v = cur_kv_head * stride_buf_vh + offs_dv[None, :]
|
||||
|
||||
if split_kv_end > split_kv_start:
|
||||
q = tl.load(Q + offs_q, mask=(mask_h[:, None]) & (mask_d[None, :]), other=0.0)
|
||||
q_k = q.to(K_Buffer.dtype.element_ty)
|
||||
if BLOCK_DPE > 0:
|
||||
qpe = tl.load(
|
||||
Q + off_qpe, mask=(mask_h[:, None]) & (mask_dpe[None, :]), other=0.0
|
||||
)
|
||||
for start_n in range(split_kv_start, split_kv_end, BLOCK_N):
|
||||
for start_n in tl.range(split_kv_start, split_kv_end, BLOCK_N):
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N)
|
||||
kv_loc = tl.load(
|
||||
kv_indices + cur_batch_kv_start_idx + offs_n,
|
||||
mask=offs_n < split_kv_end,
|
||||
other=0,
|
||||
)
|
||||
offs_buf_k = (
|
||||
kv_loc[None, :] * stride_buf_kbs
|
||||
+ cur_kv_head * stride_buf_kh
|
||||
+ offs_d[:, None]
|
||||
)
|
||||
offs_buf_k = kv_loc[None, :] * stride_buf_kbs + base_offs_k
|
||||
k = tl.load(
|
||||
K_Buffer + offs_buf_k,
|
||||
mask=(offs_n[None, :] < split_kv_end) & (mask_d[:, None]),
|
||||
other=0.0,
|
||||
)
|
||||
qk = tl.dot(q, k.to(q.dtype))
|
||||
qk = tl.dot(q_k, k)
|
||||
if BLOCK_DPE > 0:
|
||||
offs_buf_kpe = (
|
||||
kv_loc[None, :] * stride_buf_kbs
|
||||
+ cur_kv_head * stride_buf_kh
|
||||
+ offs_dpe[:, None]
|
||||
)
|
||||
offs_buf_kpe = kv_loc[None, :] * stride_buf_kbs + base_offs_kpe
|
||||
kpe = tl.load(
|
||||
K_Buffer + offs_buf_kpe,
|
||||
mask=(offs_n[None, :] < split_kv_end) & (mask_dpe[:, None]),
|
||||
@@ -376,17 +378,15 @@ def _fwd_grouped_kernel_stage1(
|
||||
qk = tl.where(
|
||||
mask_h[:, None] & (offs_n[None, :] < split_kv_end), qk, float("-inf")
|
||||
)
|
||||
|
||||
offs_buf_v = (
|
||||
kv_loc[:, None] * stride_buf_vbs
|
||||
+ cur_kv_head * stride_buf_vh
|
||||
+ offs_dv[None, :]
|
||||
)
|
||||
v = tl.load(
|
||||
V_Buffer + offs_buf_v,
|
||||
mask=(offs_n[:, None] < split_kv_end) & (mask_dv[None, :]),
|
||||
other=0.0,
|
||||
)
|
||||
if HAS_MLA:
|
||||
v = tl.trans(k)
|
||||
else:
|
||||
offs_buf_v = kv_loc[:, None] * stride_buf_vbs + base_offs_v
|
||||
v = tl.load(
|
||||
V_Buffer + offs_buf_v,
|
||||
mask=(offs_n[:, None] < split_kv_end) & (mask_dv[None, :]),
|
||||
other=0.0,
|
||||
)
|
||||
|
||||
n_e_max = tl.maximum(tl.max(qk, 1), e_max)
|
||||
re_scale = tl.exp(e_max - n_e_max)
|
||||
@@ -422,6 +422,9 @@ def _fwd_grouped_kernel_stage1(
|
||||
mask=mask_h,
|
||||
)
|
||||
|
||||
if USE_PDL:
|
||||
tl.extra.cuda.gdc_launch_dependents()
|
||||
|
||||
|
||||
def _decode_grouped_att_m_fwd(
|
||||
q,
|
||||
@@ -436,6 +439,8 @@ def _decode_grouped_att_m_fwd(
|
||||
sm_scale_withk,
|
||||
logit_cap,
|
||||
xai_temperature_len=-1,
|
||||
has_mla=False,
|
||||
use_pdl=False,
|
||||
):
|
||||
BLOCK = 32
|
||||
Lk = k_buffer.shape[-1]
|
||||
@@ -508,6 +513,8 @@ def _decode_grouped_att_m_fwd(
|
||||
num_stages=num_stages,
|
||||
Lk=Lk,
|
||||
Lv=Lv,
|
||||
HAS_MLA=has_mla,
|
||||
USE_PDL=use_pdl,
|
||||
**extra_kargs,
|
||||
)
|
||||
|
||||
@@ -531,10 +538,14 @@ def _fwd_kernel_stage2(
|
||||
BLOCK_DV: tl.constexpr,
|
||||
Lv: tl.constexpr,
|
||||
HAS_SINK: tl.constexpr,
|
||||
USE_PDL: tl.constexpr = False,
|
||||
):
|
||||
cur_batch = tl.program_id(0)
|
||||
cur_head = tl.program_id(1)
|
||||
|
||||
if USE_PDL:
|
||||
tl.extra.cuda.gdc_wait()
|
||||
|
||||
cur_batch_seq_len = tl.load(kv_indptr + cur_batch + 1) - tl.load(
|
||||
kv_indptr + cur_batch
|
||||
)
|
||||
@@ -553,7 +564,7 @@ def _fwd_kernel_stage2(
|
||||
tl.cdiv(tl.cdiv(cur_batch_seq_len, kv_splits), MIN_BLOCK_KV) * MIN_BLOCK_KV
|
||||
)
|
||||
|
||||
for split_kv_id in range(0, MAX_KV_SPLITS):
|
||||
for split_kv_id in tl.range(0, MAX_KV_SPLITS, num_stages=2):
|
||||
split_kv_start = kv_len_per_split * split_kv_id
|
||||
split_kv_end = tl.minimum(split_kv_start + kv_len_per_split, cur_batch_seq_len)
|
||||
|
||||
@@ -594,6 +605,7 @@ def _decode_softmax_reducev_fwd(
|
||||
num_kv_splits,
|
||||
max_kv_splits,
|
||||
sinks=None,
|
||||
use_pdl=False,
|
||||
):
|
||||
batch, head_num = q.shape[0], q.shape[1]
|
||||
Lv = v_buffer.shape[-1]
|
||||
@@ -627,8 +639,10 @@ def _decode_softmax_reducev_fwd(
|
||||
BLOCK_DV=BLOCK_DV,
|
||||
Lv=Lv,
|
||||
HAS_SINK=HAS_SINK,
|
||||
USE_PDL=use_pdl,
|
||||
num_warps=4,
|
||||
num_stages=2,
|
||||
**({"launch_pdl": True} if use_pdl else {}),
|
||||
**extra_kargs,
|
||||
)
|
||||
|
||||
@@ -694,6 +708,8 @@ def decode_attention_fwd_grouped(
|
||||
logit_cap=0.0,
|
||||
sinks=None,
|
||||
xai_temperature_len=-1,
|
||||
has_mla=False,
|
||||
use_pdl=False,
|
||||
):
|
||||
_decode_grouped_att_m_fwd(
|
||||
q,
|
||||
@@ -708,6 +724,8 @@ def decode_attention_fwd_grouped(
|
||||
sm_scale_withk,
|
||||
logit_cap,
|
||||
xai_temperature_len,
|
||||
has_mla=has_mla,
|
||||
use_pdl=use_pdl,
|
||||
)
|
||||
_decode_softmax_reducev_fwd(
|
||||
attn_logits,
|
||||
@@ -720,6 +738,7 @@ def decode_attention_fwd_grouped(
|
||||
num_kv_splits,
|
||||
max_kv_splits,
|
||||
sinks,
|
||||
use_pdl=use_pdl,
|
||||
)
|
||||
|
||||
|
||||
@@ -740,6 +759,8 @@ def decode_attention_fwd(
|
||||
logit_cap=0.0,
|
||||
sinks=None,
|
||||
xai_temperature_len=-1,
|
||||
has_mla=False,
|
||||
use_pdl=False,
|
||||
):
|
||||
assert max_kv_splits == attn_logits.shape[2]
|
||||
assert q.shape[0] <= kv_indptr.shape[0] - 1
|
||||
@@ -784,4 +805,6 @@ def decode_attention_fwd(
|
||||
logit_cap=logit_cap,
|
||||
sinks=sinks,
|
||||
xai_temperature_len=xai_temperature_len,
|
||||
has_mla=has_mla,
|
||||
use_pdl=use_pdl,
|
||||
)
|
||||
|
||||
@@ -382,7 +382,6 @@ def _fwd_kernel(
|
||||
mask=(mask_n[None, :]) & (mask_d[:, None]),
|
||||
other=0.0,
|
||||
)
|
||||
|
||||
qk = tl.dot(q.to(k.dtype), k)
|
||||
if BLOCK_DPE > 0:
|
||||
offs_kpe = (
|
||||
@@ -887,7 +886,6 @@ def _fwd_kernel_unified(
|
||||
other=0.0,
|
||||
)
|
||||
|
||||
# Compute QK
|
||||
qk = tl.dot(q.to(k.dtype), k)
|
||||
if BLOCK_DPE > 0:
|
||||
offs_kpe = (
|
||||
|
||||
Reference in New Issue
Block a user