Support Triton MLA FP8 KV cache (#20479)

Co-authored-by: b8zhong <b8zhong@users.noreply.github.com>
This commit is contained in:
Brayden Zhong
2026-05-06 18:32:39 -07:00
committed by GitHub
co-authored by b8zhong
parent 2e642ea187
commit 3fe8bc987e
3 changed files with 87 additions and 30 deletions
@@ -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 = (