[AMD] Optimize Kimi-K3 Triton MLA prefill on gfx950 (#35770)

Co-authored-by: clintg6 <7388379+clintg6@users.noreply.github.com>
Co-authored-by: Thomas Wang <thomawan@amd.com>
Co-authored-by: HAI <hixiao@gmail.com>
Co-authored-by: kk <43161300+kkHuang-amd@users.noreply.github.com>
This commit is contained in:
Clint
2026-09-04 17:44:30 -07:00
committed by GitHub
co-authored by clintg6 Thomas Wang HAI kk
parent 3b64169f9d
commit 3c2724c48d
8 changed files with 1847 additions and 135 deletions
@@ -16,14 +16,15 @@ Memory-efficient attention for prefill.
It supports page size = 1 and prefill with KV cache (i.e. extend).
"""
import math
from typing import Optional
import torch
import triton
import triton.language as tl
from sglang.kernels.ops.attention.decode_attention import _extract_kv_strides
from sglang.kernels.ops.attention.prefill_attention import (
context_attention_fwd,
)
from sglang.kernels.ops.attention.prefill_attention import context_attention_fwd
from sglang.kernels.ops.attention.score_mod import unpack_aux_tensors
from sglang.srt.environ import envs
from sglang.srt.utils import (
@@ -86,11 +87,12 @@ def _get_block_sizes_for_extend_attention(Lq: int, Lv: int):
# late-prefill kernel from ~12.57 ms to ~5.24 ms.
BLOCK_M, BLOCK_N = (64, 32)
num_warps = 4
elif _is_gfx95 and Lq <= 256:
# gfx950 (CDNA4), head_dim <= 256: every workgroup streams the whole
# prefix, so a larger query tile halves the KV bytes read per call;
# BLOCK_M / num_warps = 16 rows per warp is exactly one MFMA tile at
# matrix_instr_nonkdim=16. Measured on MI350X at head_dim 64, 128, 256.
elif _is_gfx95 and 128 < Lq <= 256:
# gfx950 (CDNA4), 128 < head_dim <= 256: a larger query tile halves KV bytes
# streamed per call (each workgroup reads the whole prefix); 8 warps
# hide the loads. Measured on MI350X head_dim 256: -36% kernel time,
# 28% -> 44% MFU, numerically equivalent (BLOCK_N reduction order
# unchanged). Other AMD archs / head dims keep the default below.
BLOCK_M, BLOCK_N = (128, 64)
num_warps = 8
else:
@@ -147,6 +149,21 @@ def _get_block_sizes_for_extend_attention(Lq: int, Lv: int):
return BLOCK_DMODEL, BLOCK_DPE, BLOCK_DV, BLOCK_M, BLOCK_N, num_warps
def _get_num_stages_for_extend_attention(
Lq: int, Lv: int, block_n: int | None = None
) -> int:
if _is_gfx95 and Lq == 192 and Lv == 128:
return 2
if (
_is_gfx95
and Lq == 576
and Lv == 512
and (block_n == 32 or (block_n is None and _is_triton_ge_37))
):
return 2
return 1
def _compact_extend_q_tiles_per_head(
*,
batch_size: int,
@@ -356,6 +373,7 @@ def _fwd_kernel(
BLOCK_DV: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_N_PREFIX: tl.constexpr,
USE_CUSTOM_MASK: tl.constexpr,
IS_CAUSAL: tl.constexpr,
SKIP_PREFIX_CUSTOM_MASK: tl.constexpr,
@@ -364,9 +382,14 @@ def _fwd_kernel(
SKIP_EXTEND: tl.constexpr,
STORE_TRANSPOSE: tl.constexpr,
HAS_SINK: tl.constexpr,
USE_COMPACT_TILE_GRID: tl.constexpr,
USE_EXP2: tl.constexpr,
USE_FP8_PREFIX: tl.constexpr,
USE_FP8_EXTEND: tl.constexpr,
FP8_MAX: tl.constexpr,
IS_GFX1250: tl.constexpr = False,
USE_COMPACT_TILE_GRID: tl.constexpr = False,
PAGE_SIZE: tl.constexpr = 1,
IDENTITY_KV_INDICES: tl.constexpr = False,
SCORE_MOD: tl.constexpr = None,
Aux0=None,
aux0_stride_t=0,
@@ -401,6 +424,8 @@ def _fwd_kernel(
cur_head = tl.program_id(1)
cur_block_m = tl.program_id(2)
cur_kv_head = cur_head // kv_group_num
LOG2E: tl.constexpr = 1.4426950408889634
LN2: tl.constexpr = 0.6931471805599453
cur_seq_extend_start_idx = tl.load(qo_indptr + cur_seq)
cur_seq_len_extend = tl.load(qo_indptr + cur_seq + 1) - cur_seq_extend_start_idx
@@ -459,15 +484,17 @@ def _fwd_kernel(
# stage 1: compute scores with prefix
offs_n = tl.arange(0, BLOCK_N)
# The FP8 prefix sweep can use a wider tile than the current-token sweep.
offs_n_prefix = tl.arange(0, BLOCK_N_PREFIX)
acc = tl.zeros([BLOCK_M, BLOCK_DV], dtype=tl.float32)
deno = tl.zeros([BLOCK_M], dtype=tl.float32)
e_max = tl.zeros([BLOCK_M], dtype=tl.float32) - float("inf")
prefix_end = 0 if SKIP_PREFIX else cur_seq_len_prefix
for start_n in range(0, prefix_end, BLOCK_N):
start_n = tl.multiple_of(start_n, BLOCK_N)
mask_n = (start_n + offs_n) < cur_seq_len_prefix
for start_n in range(0, prefix_end, BLOCK_N_PREFIX):
start_n = tl.multiple_of(start_n, BLOCK_N_PREFIX)
mask_n = (start_n + offs_n_prefix) < cur_seq_len_prefix
final_mask = mask_m[:, None] & mask_n[None, :]
if USE_CUSTOM_MASK and not SKIP_PREFIX_CUSTOM_MASK:
@@ -478,7 +505,7 @@ def _fwd_kernel(
* (cur_seq_len + window_kv_offset)
+ window_kv_offset
+ start_n
+ offs_n[None, :],
+ offs_n_prefix[None, :],
mask=(mask_m[:, None] & mask_n[None, :]),
other=0,
)
@@ -488,7 +515,7 @@ def _fwd_kernel(
# q_id = prefix_len + cur_m, kv_id = cur_n
window_mask = (
cur_seq_len_prefix + cur_block_m * BLOCK_M + offs_m[:, None]
) <= (start_n + offs_n[None, :] + SLIDING_WINDOW_SIZE)
) <= (start_n + offs_n_prefix[None, :] + SLIDING_WINDOW_SIZE)
final_mask &= window_mask
SKIP_TILE = False
@@ -496,11 +523,14 @@ def _fwd_kernel(
SKIP_TILE = tl.max(tl.max(final_mask.to(tl.int32), axis=1), axis=0) == 0
if not SKIP_TILE:
offs_kv_loc = tl.load(
kv_indices + cur_seq_kv_start_idx + start_n + offs_n,
mask=mask_n,
other=0,
)
if IDENTITY_KV_INDICES:
offs_kv_loc = cur_seq_kv_start_idx + start_n + offs_n_prefix
else:
offs_kv_loc = tl.load(
kv_indices + cur_seq_kv_start_idx + start_n + offs_n_prefix,
mask=mask_n,
other=0,
)
# Page-aware KV address math. At PAGE_SIZE==1
# (legacy / non-shared / shared-at-ps=1), Triton specializes
@@ -560,7 +590,10 @@ def _fwd_kernel(
qk += tl.dot(qpe, kpe.to(qpe.dtype))
else:
qk += tl.dot(qpe.to(kpe.dtype), kpe)
qk *= sm_scale * k_scale
if USE_EXP2:
qk *= sm_scale * k_scale * LOG2E
else:
qk *= sm_scale * k_scale
if logit_cap > 0:
qk = logit_cap * tanh(qk / logit_cap)
@@ -572,7 +605,7 @@ def _fwd_kernel(
qk = SCORE_MOD(
qk,
(cur_seq_len_prefix + cur_block_m * BLOCK_M + offs_m)[:, None],
start_n + offs_n[None, :],
start_n + offs_n_prefix[None, :],
(cur_seq_extend_start_idx + cur_block_m * BLOCK_M + offs_m)[
:, None
],
@@ -590,8 +623,12 @@ def _fwd_kernel(
row_max_fixed = tl.where(row_max == float("-inf"), -1e20, row_max)
n_e_max = tl.maximum(row_max_fixed, e_max)
re_scale = tl.exp(e_max - n_e_max)
p = tl.exp(qk - n_e_max[:, None])
if USE_EXP2:
re_scale = tl.exp2(e_max - n_e_max)
p = tl.exp2(qk - n_e_max[:, None])
else:
re_scale = tl.exp(e_max - n_e_max)
p = tl.exp(qk - n_e_max[:, None])
deno = deno * re_scale + tl.sum(p, 1)
if PAGE_SIZE == 1:
@@ -612,14 +649,18 @@ def _fwd_kernel(
mask=mask_n[:, None] & mask_dv[None, :],
other=0.0,
)
# keep softmax weights p in fp32 for the P·V dot (do not downcast to bf16)
# on gfx1250; on other platforms restore the original p.to(v.dtype) cast.
# TODO: remove this branch once the gfx1250 bf16 P·V issue is resolved.
if IS_GFX1250:
dot = tl.dot(p, v.to(tl.float32), out_dtype=tl.float32)
if USE_FP8_PREFIX:
p_dot = (p * FP8_MAX).to(v.dtype)
acc = acc * re_scale[:, None] + tl.dot(p_dot, v) * (v_scale / FP8_MAX)
else:
dot = tl.dot(p.to(v.dtype), v)
acc = acc * re_scale[:, None] + dot * v_scale
# keep softmax weights p in fp32 for the P·V dot (do not downcast to bf16)
# on gfx1250; on other platforms restore the original p.to(v.dtype) cast.
# TODO: remove this branch once the gfx1250 bf16 P·V issue is resolved.
if IS_GFX1250:
dot = tl.dot(p, v.to(tl.float32), out_dtype=tl.float32)
else:
dot = tl.dot(p.to(v.dtype), v)
acc = acc * re_scale[:, None] + dot * v_scale
e_max = n_e_max
@@ -631,16 +672,7 @@ def _fwd_kernel(
else tl.minimum(cur_seq_len_extend, (cur_block_m + 1) * BLOCK_M)
)
extend_end = 0 if SKIP_EXTEND else cur_block_m_end
# The mask below keeps (q, kv) iff q <= kv + SLIDING_WINDOW_SIZE, so no tile
# under this floor can hold an unmasked element -- tight for any BLOCK_M/BLOCK_N.
# SKIP_TILE already made those tiles no-ops, so bounding the loop is
# bit-identical and drops their cross-wave tl.max reduction.
extend_start = 0
if SLIDING_WINDOW_SIZE > 0:
extend_start = (
tl.maximum(cur_block_m * BLOCK_M - SLIDING_WINDOW_SIZE, 0) // BLOCK_N
) * BLOCK_N
for start_n in range(extend_start, extend_end, BLOCK_N):
for start_n in range(0, extend_end, BLOCK_N):
start_n = tl.multiple_of(start_n, BLOCK_N)
mask_n = (start_n + offs_n) < cur_block_m_end
@@ -706,7 +738,10 @@ def _fwd_kernel(
)
qk += tl.dot(qpe, kpe)
qk *= sm_scale
if USE_EXP2:
qk *= sm_scale * LOG2E
else:
qk *= sm_scale
if logit_cap > 0:
qk = logit_cap * tanh(qk / logit_cap)
@@ -736,8 +771,12 @@ def _fwd_kernel(
row_max_fixed = tl.where(row_max == float("-inf"), -1e20, row_max)
n_e_max = tl.maximum(row_max_fixed, e_max)
re_scale = tl.exp(e_max - n_e_max)
p = tl.exp(qk - n_e_max[:, None])
if USE_EXP2:
re_scale = tl.exp2(e_max - n_e_max)
p = tl.exp2(qk - n_e_max[:, None])
else:
re_scale = tl.exp(e_max - n_e_max)
p = tl.exp(qk - n_e_max[:, None])
deno = deno * re_scale + tl.sum(p, 1)
offs_v = (
@@ -748,26 +787,41 @@ def _fwd_kernel(
v = tl.load(
V_Extend + offs_v, mask=mask_n[:, None] & mask_dv[None, :], other=0.0
)
# keep softmax weights p in fp32 for the P·V dot (do not downcast to bf16)
# on gfx1250; on other platforms restore the original p.to(v.dtype) cast.
# TODO: remove this branch once the gfx1250 bf16 P·V issue is resolved.
if IS_GFX1250:
dot = tl.dot(p, v.to(tl.float32), out_dtype=tl.float32)
if USE_FP8_EXTEND:
p_dot = (p * FP8_MAX).to(v.dtype)
acc = acc * re_scale[:, None] + tl.dot(p_dot, v) * (1.0 / FP8_MAX)
else:
dot = tl.dot(p.to(v.dtype), v)
acc = acc * re_scale[:, None] + dot
# keep softmax weights p in fp32 for the P·V dot (do not downcast to bf16)
# on gfx1250; on other platforms restore the original p.to(v.dtype) cast.
# TODO: remove this branch once the gfx1250 bf16 P·V issue is resolved.
if IS_GFX1250:
dot = tl.dot(p, v.to(tl.float32), out_dtype=tl.float32)
else:
dot = tl.dot(p.to(v.dtype), v)
acc = acc * re_scale[:, None] + dot
e_max = n_e_max
if HAS_SINK:
cur_sink = tl.load(sink_ptr + cur_head)
deno += tl.exp(cur_sink - e_max)
if USE_EXP2:
deno += tl.exp2(cur_sink * LOG2E - e_max)
else:
deno += tl.exp(cur_sink - e_max)
# A ragged prefix chunk can be empty for some requests. Represent an empty
# partial as output=0 and LSE=-inf so merge_state ignores it exactly.
no_kv = deno == 0.0
if STORE_LSE:
offs_lse = (
cur_seq_extend_start_idx + cur_block_m * BLOCK_M + offs_m
) * stride_lse_bs + cur_head * stride_lse_h
lse = tl.log(deno) + e_max
if USE_EXP2:
lse = tl.log(deno) + e_max * LN2
else:
lse = tl.log(deno) + e_max
lse = tl.where(no_kv, float("-inf"), lse)
tl.store(LSE_Extend + offs_lse, lse, mask=mask_m)
offs_o = (
@@ -776,16 +830,17 @@ def _fwd_kernel(
+ cur_head * stride_oh
+ offs_dv[None, :]
)
deno_safe = tl.where(no_kv, 1.0, deno)
if STORE_TRANSPOSE:
tl.store(
O_Extend + offs_o.T,
(acc / deno[:, None]).T,
(acc / deno_safe[:, None]).T,
mask=(mask_m[:, None] & mask_dv[None, :]).T,
)
else:
tl.store(
O_Extend + offs_o,
acc / deno[:, None],
acc / deno_safe[:, None],
mask=mask_m[:, None] & mask_dv[None, :],
)
@@ -820,6 +875,7 @@ def extend_attention_fwd(
score_mod=None,
aux_tensors=None,
extend_seq_lens_cpu=None,
identity_kv_indices: bool = False,
):
"""
q_extend, k_extend, v_extend, o_extend: contiguous tensors
@@ -832,6 +888,8 @@ def extend_attention_fwd(
respectively so DCP can compute those two parts separately.
``score_mod`` / ``aux_tensors`` add a custom term to the attention logits;
see triton_ops/score_mod.py for the contract.
``identity_kv_indices`` promises that the prefix buffer is densely packed,
allowing direct addressing instead of loading an index for every token.
"""
Lq, Lk, Lv = (
q_extend.shape[-1],
@@ -839,20 +897,90 @@ def extend_attention_fwd(
v_extend.shape[-1],
)
# Get block sizes and configuration
BLOCK_DMODEL, BLOCK_DPE, BLOCK_DV, BLOCK_M, BLOCK_N, num_warps = (
_get_block_sizes_for_extend_attention(Lq, Lv)
)
sm_scale = sm_scale or 1.0 / (Lq**0.5)
batch_size, head_num = qo_indptr.shape[0] - 1, q_extend.shape[1]
kv_group_num = q_extend.shape[1] // k_extend.shape[1]
zero_prefix_shape = (
head_num == 12
and k_extend.shape[1] == 12
and Lq == 192
and Lk == 192
and Lv == 128
)
absorbed_shape = (
head_num == 12
and k_extend.shape[1] == 1
and Lq == 576
and Lk == 576
and Lv == 512
)
kimi_k3_shape = zero_prefix_shape or absorbed_shape
# Match Aiter's opt-in behavior: cast Q, K, and V separately before the
# native-FP8 zero-prefix kernel at every sequence length.
use_fp8_zero_prefix = (
_is_gfx95
and envs.SGLANG_TRITON_FP8_PREFILL_ATTN.get()
and zero_prefix_shape
and q_extend.dtype == torch.bfloat16
and k_extend.dtype == torch.bfloat16
and v_extend.dtype == torch.bfloat16
and k_buffer.dtype == torch.float8_e4m3fn
and v_buffer.dtype == torch.float8_e4m3fn
and custom_mask is None
and is_causal
and sliding_window_size <= 0
and logit_cap <= 0
and xai_temperature_len <= 0
and sinks is None
and score_mod is None
and aux_tensors is None
)
if use_fp8_zero_prefix:
q_extend = q_extend.to(torch.float8_e4m3fn)
k_extend = k_extend.to(torch.float8_e4m3fn)
v_extend = v_extend.to(torch.float8_e4m3fn)
# Get block sizes and configuration for the generic fallback.
BLOCK_DMODEL, BLOCK_DPE, BLOCK_DV, BLOCK_M, BLOCK_N, num_warps = (
_get_block_sizes_for_extend_attention(Lq, Lv)
)
USE_CUSTOM_MASK = custom_mask is not None
# Skip custom mask for prefix part
SKIP_PREFIX_CUSTOM_MASK = skip_prefix_custom_mask
HAS_SINK = sinks is not None
USE_FP8_PREFIX = (
_is_gfx95
and kimi_k3_shape
and k_buffer.dtype == torch.float8_e4m3fn
and v_buffer.dtype == torch.float8_e4m3fn
)
USE_FP8_EXTEND = (
_is_gfx95
and zero_prefix_shape
and k_extend.dtype == torch.float8_e4m3fn
and v_extend.dtype == torch.float8_e4m3fn
)
FP8_MAX = (
torch.finfo(torch.float8_e4m3fn).max
if USE_FP8_PREFIX or USE_FP8_EXTEND
else 1.0
)
# At head_dim 192, FP8 operands allow a 128-column tile while BF16 does
# not fit in LDS. Widen prefix and extend sweeps independently.
BLOCK_N_ARCH = BLOCK_N
FP8_BLOCK_N = 128 if BLOCK_N_ARCH < 128 and Lq <= 192 else BLOCK_N_ARCH
BLOCK_N = FP8_BLOCK_N if USE_FP8_EXTEND else BLOCK_N_ARCH
BLOCK_N_PREFIX = FP8_BLOCK_N if USE_FP8_PREFIX else BLOCK_N_ARCH
USE_EXP2 = (
_is_gfx95
and kimi_k3_shape
and logit_cap <= 0
and xai_temperature_len <= 0
and score_mod is None
)
STORE_LSE = lse_extend is not None
stride_lse_bs = lse_extend.stride(0) if STORE_LSE else 0
stride_lse_h = lse_extend.stride(1) if STORE_LSE else 0
@@ -877,7 +1005,9 @@ def extend_attention_fwd(
grid = (compact_q_tiles, head_num)
else:
grid = (batch_size, head_num, triton.cdiv(max_len_extend, BLOCK_M))
num_stages = 1
num_stages = (
_get_num_stages_for_extend_attention(Lq, Lv, BLOCK_N) if kimi_k3_shape else 1
)
extra_kargs = {}
if _is_hip:
@@ -940,6 +1070,7 @@ def extend_attention_fwd(
BLOCK_DV=BLOCK_DV,
BLOCK_M=BLOCK_M,
BLOCK_N=BLOCK_N,
BLOCK_N_PREFIX=BLOCK_N_PREFIX,
Lq=Lq,
Lv=Lv,
USE_CUSTOM_MASK=USE_CUSTOM_MASK,
@@ -952,7 +1083,12 @@ def extend_attention_fwd(
IS_GFX1250=_is_gfx1250,
STORE_TRANSPOSE=_is_hip,
USE_COMPACT_TILE_GRID=use_compact_tile_grid,
USE_EXP2=USE_EXP2,
USE_FP8_PREFIX=USE_FP8_PREFIX,
USE_FP8_EXTEND=USE_FP8_EXTEND,
FP8_MAX=FP8_MAX,
PAGE_SIZE=page_size,
IDENTITY_KV_INDICES=identity_kv_indices,
SCORE_MOD=score_mod,
Aux0=aux0,
aux0_stride_t=aux0_stride_t,
@@ -1470,3 +1606,536 @@ def extend_attention_fwd_unified(
num_stages=num_stages,
**extra_kargs,
)
@triton.jit
def _dense_prefill_inner(
acc,
deno,
e_max,
q,
qpe,
K,
V,
cur_seq_kv_start,
cur_kv_head,
stride_kbs,
stride_kh,
stride_vbs,
stride_vh,
offs_d,
offs_dpe,
offs_dv,
q_pos,
mask_m,
mask_d,
mask_dv,
start_lo,
start_hi,
kv_end,
qk_scale,
logit_cap: tl.constexpr,
BLOCK_DPE: tl.constexpr,
BLOCK_N: tl.constexpr,
MASKED: tl.constexpr,
IS_CAUSAL: tl.constexpr,
EVEN_D: tl.constexpr,
USE_FP8: tl.constexpr,
LOG2_FP8_MAX: tl.constexpr,
):
"""One online-softmax sweep over ``[start_lo, start_hi)`` of the KV axis.
Instantiated twice per kernel. ``MASKED=False`` covers the interior, where
every key is visible to every query row of the block: the predicates are
gone, so the K/V loads stay contiguous over the head dim and widen to
dwordx4. ``MASKED=True`` covers only the causal diagonal and the ragged
tail. Splitting the sweep is what keeps that per-element ``tl.where`` and
the narrowed loads off the long prefix, which is the bulk of the work.
"""
offs_n = tl.arange(0, BLOCK_N)
for start_n in range(start_lo, start_hi, BLOCK_N):
start_n = tl.multiple_of(start_n, BLOCK_N)
offs_kn = start_n + offs_n
offs_k = (
(cur_seq_kv_start + offs_kn[None, :]) * stride_kbs
+ cur_kv_head * stride_kh
+ offs_d[:, None]
)
if MASKED:
mask_n = offs_kn < kv_end
k = tl.load(K + offs_k, mask=mask_n[None, :] & mask_d[:, None], other=0.0)
elif EVEN_D:
k = tl.load(K + offs_k)
else:
k = tl.load(K + offs_k, mask=mask_d[:, None], other=0.0)
qk = tl.dot(q, k, out_dtype=tl.float32)
if BLOCK_DPE > 0:
offs_kpe = (
(cur_seq_kv_start + offs_kn[None, :]) * stride_kbs
+ cur_kv_head * stride_kh
+ offs_dpe[:, None]
)
if MASKED:
kpe = tl.load(K + offs_kpe, mask=mask_n[None, :], other=0.0)
else:
kpe = tl.load(K + offs_kpe)
qk += tl.dot(qpe, kpe, out_dtype=tl.float32)
if logit_cap > 0:
qk *= qk_scale
qk = logit_cap * tanh(qk / logit_cap)
qk *= 1.4426950408889634
else:
qk *= qk_scale
if MASKED:
final_mask = mask_m[:, None] & mask_n[None, :]
if IS_CAUSAL:
final_mask &= q_pos[:, None] >= offs_kn[None, :]
qk = tl.where(final_mask, qk, float("-inf"))
row_max = tl.max(qk, 1)
# A fully masked row would poison e_max with -inf; -1e20 keeps the
# rescale finite and still contributes nothing to deno.
row_max = tl.where(row_max == float("-inf"), -1e20, row_max)
else:
row_max = tl.max(qk, 1)
n_e_max = tl.maximum(row_max, e_max)
re_scale = tl.exp2(e_max - n_e_max)
if USE_FP8:
# Bias the exponent instead of multiplying P by FP8_MAX after the
# fact: exp2(qk - m + log2(FP8_MAX)) == exp2(qk - m) * FP8_MAX, so
# the lift costs a BLOCK_M-wide subtract rather than a
# BLOCK_M x BLOCK_N one, and skips a rounding step. deno picks up
# the same constant factor, which cancels against acc at the final
# divide -- so out_scale drops its 1/FP8_MAX and LSE subtracts
# log2(FP8_MAX) from e_max.
p = tl.exp2(qk - (n_e_max - LOG2_FP8_MAX)[:, None])
else:
p = tl.exp2(qk - n_e_max[:, None])
deno = deno * re_scale + tl.sum(p, 1)
offs_v = (
(cur_seq_kv_start + offs_kn[:, None]) * stride_vbs
+ cur_kv_head * stride_vh
+ offs_dv[None, :]
)
if MASKED:
v = tl.load(V + offs_v, mask=mask_n[:, None] & mask_dv[None, :], other=0.0)
elif EVEN_D:
v = tl.load(V + offs_v)
else:
v = tl.load(V + offs_v, mask=mask_dv[None, :], other=0.0)
# P is already lifted off the e4m3 denormal floor by the exponent bias
# above when USE_FP8; the cast is all that is left.
#
# Guarding this rescale on `tl.min(re_scale) < 1.0` (it is exactly 1.0
# once the running max settles, which is most of a long prefix) was
# measured at 21% SLOWER: the branch splits the loop body and the
# pipeliner stops prefetching K/V across iterations. Keep it
# unconditional.
acc = acc * re_scale[:, None] + tl.dot(p.to(v.dtype), v)
e_max = n_e_max
return acc, deno, e_max
@triton.jit
def _fwd_kernel_dense_prefill(
Q,
K,
V,
O,
Lse,
qo_indptr,
kv_indptr,
sm_scale,
k_scale,
v_scale,
stride_qbs,
stride_qh,
stride_kbs,
stride_kh,
stride_vbs,
stride_vh,
stride_obs,
stride_oh,
stride_lse_bs,
stride_lse_h,
kv_group_num: tl.constexpr,
logit_cap: tl.constexpr,
Lq: tl.constexpr,
Lv: tl.constexpr,
BLOCK_DMODEL: tl.constexpr,
BLOCK_DPE: tl.constexpr,
BLOCK_DV: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
IS_CAUSAL: tl.constexpr,
STORE_LSE: tl.constexpr,
USE_FP8: tl.constexpr,
LOG2_FP8_MAX: tl.constexpr,
EVEN_D: tl.constexpr,
NUM_BLOCKS_M: tl.constexpr,
):
"""Single-loop dense prefill: K/V hold prefix + current chunk contiguously.
``_fwd_kernel`` needs two stages because its prefix lives in the paged
latent KV cache while its suffix is contiguous. Once the prefix has been
up-projected into dense per-head K/V (``AttnForwardMethod.MHA_ONE_SHOT``),
both halves share one base pointer, one dtype and one scale, so the split
buys nothing and only costs pipelining and registers.
Causal masking is bottom-right aligned: query ``m`` of a sequence sits at
absolute position ``prefix_len + m``, where ``prefix_len = kv_len - q_len``.
"""
cur_seq = tl.program_id(0)
cur_head = tl.program_id(1)
# Causal cost grows with block index: block m sweeps prefix_len + m*BLOCK_M
# keys. Program IDs dispatch roughly in order, so issuing the heavy blocks
# first lets the cheap ones backfill the tail instead of trailing it.
cur_block_m = NUM_BLOCKS_M - 1 - tl.program_id(2)
cur_kv_head = cur_head // kv_group_num
LOG2E: tl.constexpr = 1.4426950408889634
LN2: tl.constexpr = 0.6931471805599453
cur_seq_q_start = tl.load(qo_indptr + cur_seq)
cur_seq_q_len = tl.load(qo_indptr + cur_seq + 1) - cur_seq_q_start
cur_seq_kv_start = tl.load(kv_indptr + cur_seq)
cur_seq_kv_len = tl.load(kv_indptr + cur_seq + 1) - cur_seq_kv_start
cur_seq_prefix_len = cur_seq_kv_len - cur_seq_q_len
# Grid axis 2 spans the batch-max query length; short sequences bail early.
if cur_block_m * BLOCK_M >= cur_seq_q_len:
return
offs_m = tl.arange(0, BLOCK_M)
offs_n = tl.arange(0, BLOCK_N)
offs_d = tl.arange(0, BLOCK_DMODEL)
offs_dv = tl.arange(0, BLOCK_DV)
mask_m = (cur_block_m * BLOCK_M + offs_m) < cur_seq_q_len
mask_d = offs_d < Lq
mask_dv = offs_dv < Lv
offs_q = (
(cur_seq_q_start + cur_block_m * BLOCK_M + offs_m[:, None]) * stride_qbs
+ cur_head * stride_qh
+ offs_d[None, :]
)
q = tl.load(Q + offs_q, mask=mask_m[:, None] & mask_d[None, :], other=0.0)
if BLOCK_DPE > 0:
offs_dpe = BLOCK_DMODEL + tl.arange(0, BLOCK_DPE)
offs_qpe = (
(cur_seq_q_start + cur_block_m * BLOCK_M + offs_m[:, None]) * stride_qbs
+ cur_head * stride_qh
+ offs_dpe[None, :]
)
qpe = tl.load(Q + offs_qpe, mask=mask_m[:, None], other=0.0)
else:
# Never read: BLOCK_DPE is constexpr, so the whole rope branch is
# folded away. These only keep the inner helper's signature uniform.
offs_dpe = offs_d
qpe = q
# Absolute position of each query row inside its sequence.
q_pos = cur_seq_prefix_len + cur_block_m * BLOCK_M + offs_m
acc = tl.zeros([BLOCK_M, BLOCK_DV], dtype=tl.float32)
deno = tl.zeros([BLOCK_M], dtype=tl.float32)
e_max = tl.zeros([BLOCK_M], dtype=tl.float32) - float("inf")
if IS_CAUSAL:
kv_end = tl.minimum(
cur_seq_kv_len, cur_seq_prefix_len + (cur_block_m + 1) * BLOCK_M
)
# The lowest query row of this block sees keys 0..prefix_len+m*BLOCK_M,
# so every whole BLOCK_N below that bound is unmasked for all rows.
n_full = ((cur_seq_prefix_len + cur_block_m * BLOCK_M + 1) // BLOCK_N) * BLOCK_N
n_full = tl.minimum(n_full, kv_end)
else:
kv_end = cur_seq_kv_len
n_full = (kv_end // BLOCK_N) * BLOCK_N
qk_scale = sm_scale * k_scale
if logit_cap <= 0:
qk_scale *= LOG2E
acc, deno, e_max = _dense_prefill_inner(
acc,
deno,
e_max,
q,
qpe,
K,
V,
cur_seq_kv_start,
cur_kv_head,
stride_kbs,
stride_kh,
stride_vbs,
stride_vh,
offs_d,
offs_dpe,
offs_dv,
q_pos,
mask_m,
mask_d,
mask_dv,
0,
n_full,
kv_end,
qk_scale,
logit_cap,
BLOCK_DPE=BLOCK_DPE,
BLOCK_N=BLOCK_N,
MASKED=False,
IS_CAUSAL=IS_CAUSAL,
EVEN_D=EVEN_D,
USE_FP8=USE_FP8,
LOG2_FP8_MAX=LOG2_FP8_MAX,
)
acc, deno, e_max = _dense_prefill_inner(
acc,
deno,
e_max,
q,
qpe,
K,
V,
cur_seq_kv_start,
cur_kv_head,
stride_kbs,
stride_kh,
stride_vbs,
stride_vh,
offs_d,
offs_dpe,
offs_dv,
q_pos,
mask_m,
mask_d,
mask_dv,
n_full,
kv_end,
kv_end,
qk_scale,
logit_cap,
BLOCK_DPE=BLOCK_DPE,
BLOCK_N=BLOCK_N,
MASKED=True,
IS_CAUSAL=IS_CAUSAL,
EVEN_D=EVEN_D,
USE_FP8=USE_FP8,
LOG2_FP8_MAX=LOG2_FP8_MAX,
)
no_kv = deno == 0.0
if STORE_LSE:
offs_lse = (
cur_seq_q_start + cur_block_m * BLOCK_M + offs_m
) * stride_lse_bs + cur_head * stride_lse_h
# e_max is in log2 units because qk carries the folded LOG2E. Under
# FP8 deno also carries the exponent-bias lift; undoing it in log space
# is one more constant on the same term.
if USE_FP8:
lse = tl.log(deno) + (e_max - LOG2_FP8_MAX) * LN2
else:
lse = tl.log(deno) + e_max * LN2
lse = tl.where(no_kv, float("-inf"), lse)
tl.store(Lse + offs_lse, lse, mask=mask_m)
offs_o = (
(cur_seq_q_start + cur_block_m * BLOCK_M + offs_m[:, None]) * stride_obs
+ cur_head * stride_oh
+ offs_dv[None, :]
)
# The FP8 lift applied to P divides out between acc and deno, so v_scale
# is the only surviving factor.
out_scale = v_scale
deno_safe = tl.where(no_kv, 1.0, deno)
tl.store(
O + offs_o,
acc * (out_scale / deno_safe[:, None]),
mask=mask_m[:, None] & mask_dv[None, :],
)
def can_use_dense_prefill_fp8(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
*,
is_causal: bool,
logit_cap: float,
) -> bool:
"""Whether q/k/v may be cast to FP8 ahead of ``dense_prefill_attention_fwd``.
Deliberately as narrow as the zero-prefix gate in ``extend_attention_fwd``:
gfx950, BF16 dense inputs, plain causal softmax. Callers are responsible
for having already rejected custom masks, sinks, SWA and score mods.
"""
return (
_is_gfx95
and envs.SGLANG_TRITON_FP8_PREFILL_ATTN.get()
and q.dtype == torch.bfloat16
and k.dtype == torch.bfloat16
and v.dtype == torch.bfloat16
and is_causal
and logit_cap <= 0
)
def dense_prefill_attention_fwd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
o: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
max_len_q: int,
sm_scale: Optional[float] = None,
k_scale: float = 1.0,
v_scale: float = 1.0,
logit_cap: float = 0.0,
is_causal: bool = True,
lse: Optional[torch.Tensor] = None,
) -> None:
"""Dense varlen prefill over a fully materialized K/V.
q/o are addressed by ``qo_indptr``; k/v by ``kv_indptr``, whose
per-sequence length must be >= the query length. The excess leading rows
are the cached prefix, which every query of that sequence attends. Writes
``o``, and ``lse`` (natural log) when provided.
Shapes::
q : [sum(q_len), H, Lq] o : [sum(q_len), H, Lv]
k : [sum(kv_len), H_kv, Lq] v : [sum(kv_len), H_kv, Lv]
"""
Lq, Lk, Lv = q.shape[-1], k.shape[-1], v.shape[-1]
assert Lq == Lk, f"q/k head dims must match, got {Lq} vs {Lk}"
assert k.shape[0] == v.shape[0], (
f"k/v token counts must match, got {k.shape[0]} vs {v.shape[0]}"
)
sm_scale = sm_scale or 1.0 / (Lq**0.5)
batch_size = qo_indptr.shape[0] - 1
head_num = q.shape[1]
kv_group_num = q.shape[1] // k.shape[1]
# Shares the extend-attention tables: same head dims, same archs, and the
# 192/128 K3 entry is the shape this path exists to serve.
BLOCK_DMODEL, BLOCK_DPE, BLOCK_DV, BLOCK_M, BLOCK_N, num_warps = (
_get_block_sizes_for_extend_attention(Lq, Lv)
)
assert BLOCK_DMODEL + BLOCK_DPE >= Lq, (
f"tile {BLOCK_DMODEL}+{BLOCK_DPE} cannot cover head dim {Lq}"
)
# Both sides of every tl.dot must share a dtype, so a mixed bf16-q/fp8-k
# pair does not compile. That pair is reachable: on gfx95 with MXFP4
# kv_b_proj weights, forward_mha_rocm fuses the up-projection with the FP8
# cast (fused_gemm_afp4wfp4_split_cat) and hands back k/v already in e4m3
# while q is still bf16. Follow the cheap direction -- promote q rather
# than upcast the far larger k/v -- and keep the P scaling consistent with
# it. p is cast to v.dtype, so v decides the underflow lift too.
fp8 = torch.float8_e4m3fn
if fp8 in (q.dtype, k.dtype, v.dtype):
q, k, v = q.to(fp8), k.to(fp8), v.to(fp8)
else:
assert q.dtype == k.dtype == v.dtype, (
f"q/k/v dtypes must match, got {q.dtype}/{k.dtype}/{v.dtype}"
)
use_fp8 = q.dtype == fp8
fp8_max = torch.finfo(fp8).max if use_fp8 else 1.0
store_lse = lse is not None
stride_lse_bs = lse.stride(0) if store_lse else 0
stride_lse_h = lse.stride(1) if store_lse else 0
extra_kargs = {}
if _is_hip:
# No kpack: gfx950 overwrites it to 1 and warns on every launch.
# matrix_instr_nonkdim=32 was measured and is ~1.4x slower here.
extra_kargs = {"waves_per_eu": 1, "matrix_instr_nonkdim": 16}
num_stages = _get_num_stages_for_extend_attention(Lq, Lv, BLOCK_N)
if _is_gfx95 and Lq == 192 and Lv == 128 and use_fp8:
# The shared table is tuned for the two-stage absorbed kernel; this
# sweep is mask-free over its interior and wants a wider KV tile.
# FP8 only: at BLOCK_N=128 the BF16 K/V tiles need 168 KB of LDS
# against a 160 KB limit, so BF16 stays on the narrow tile below.
# Swept BLOCK_M x BLOCK_N x warps x stages over the three shapes this
# path actually sees (FP8, ms):
#
# q16384/p65536 q8192/p32768 q2464/p65536
# 128/64 w4 st3 7.89 4.00 1.28
# 128/128 w4 st2 7.12 3.56 1.18
# 256/128 w4 st2 6.91 3.50 2.08
#
# BLOCK_M=256 edges ahead on full chunks but halves the M-block count,
# which starves the 256 CUs on the short trailing chunk of a request --
# 1.6x slower there. 128/128 wins everywhere. At BLOCK_N=128 the loop
# body already covers the load latency, so the deeper pipeline that
# helped at BLOCK_N=64 no longer pays (7.12 at 2 stages vs 7.24 at 3).
BLOCK_N = 128
num_warps = 4
num_stages = 2
elif _is_gfx95 and Lq == 192 and Lv == 128:
# BF16, stuck on BLOCK_N=64 by LDS: there the mask-free interior does
# pipeline deeper than the table's two stages (10.90ms vs 11.10ms at
# 16K queries over a 64K prefix). 4 stages does not compile (Triton
# asserts in its pipeliner).
num_stages = 3
num_blocks_m = triton.cdiv(max_len_q, BLOCK_M)
grid = (batch_size, head_num, num_blocks_m)
_fwd_kernel_dense_prefill[grid](
q,
k,
v,
o,
lse,
qo_indptr,
kv_indptr,
sm_scale,
k_scale,
v_scale,
q.stride(0),
q.stride(1),
k.stride(0),
k.stride(1),
v.stride(0),
v.stride(1),
o.stride(0),
o.stride(1),
stride_lse_bs,
stride_lse_h,
kv_group_num=kv_group_num,
logit_cap=logit_cap,
Lq=Lq,
Lv=Lv,
BLOCK_DMODEL=BLOCK_DMODEL,
BLOCK_DPE=BLOCK_DPE,
BLOCK_DV=BLOCK_DV,
BLOCK_M=BLOCK_M,
BLOCK_N=BLOCK_N,
IS_CAUSAL=is_causal,
STORE_LSE=store_lse,
USE_FP8=use_fp8,
LOG2_FP8_MAX=math.log2(fp8_max),
EVEN_D=(BLOCK_DMODEL + BLOCK_DPE == Lq and BLOCK_DV == Lv),
NUM_BLOCKS_M=num_blocks_m,
num_warps=num_warps,
num_stages=num_stages,
**extra_kargs,
)
+9
View File
@@ -1037,6 +1037,15 @@ class Envs:
# gfx950 MLA decode stage-1: pick the launch geometry and split count per batch.
# Reorders the fp32 accumulation, so off by default.
SGLANG_MLA_DECODE_TUNE = EnvBool(False)
# Native FP8 prefill for exact gfx950 Kimi-K3 zero-prefix and absorbed
# cached-prefix shapes. Validated at 98% GSM8K accuracy.
SGLANG_TRITON_FP8_PREFILL_ATTN = EnvBool(True)
# Route Triton MLA prefill that carries a cached prefix through dense
# (non-absorbed) one-shot MHA: up-project the prefix out of the latent KV
# cache and run a single dense FP8 kernel instead of the absorbed 576/512
# prefill. Materializes K/V for the whole batch, so it only engages when
# the batch fits the chunk budget.
SGLANG_TRITON_DENSE_PREFILL_ATTN = EnvBool(True)
SGLANG_ENABLE_TORCH_COMPILE = EnvBool(False)
SGLANG_TRITON_PREFILL_TRUNCATION_ALIGN_SIZE = EnvInt(4096)
SGLANG_TRITON_DECODE_SPLIT_TILE_SIZE = EnvInt(256)
@@ -7,6 +7,9 @@ import torch
import triton
from sglang.kernels.ops.attention.metadata import get_num_kv_splits_triton
from sglang.kernels.ops.attention.mla_kv_pack_quantize_fp8 import (
mla_kv_pack_quantize_fp8,
)
from sglang.srt.configs.hybrid_arch import mambaish_config
from sglang.srt.configs.model_config import (
AttentionArch,
@@ -36,12 +39,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
cuda_graph_fully_disabled,
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import (
get_exec,
get_parallel,
get_schedule,
get_spec,
)
from sglang.srt.runtime_context import get_exec, get_parallel, get_schedule, get_spec
from sglang.srt.speculative.spec_utils import (
draft_kv_indices_buffer_width,
draft_kv_indices_used_len,
@@ -90,8 +88,6 @@ def _should_use_verify_shared_kv(model_config, topk, use_mla, use_verify_splitkv
if use_mla:
return is_kimi_k3(model_config.hf_config)
if is_dspark_draft(model_config.hf_config):
# Added for the K3 DSpark draft model, which is qwen3 type attention,
# and using bidirectional (non-causal) mode.
return use_verify_splitkv
return (
use_verify_splitkv
@@ -166,15 +162,13 @@ class TritonAttnBackend(AttentionBackend):
)
from sglang.kernels.ops.attention.extend_attention import (
build_unified_kv_indices,
can_use_dense_prefill_fp8,
dense_prefill_attention_fwd,
extend_attention_fwd,
extend_attention_fwd_unified,
)
from sglang.kernels.ops.attention.verify_mla import (
verify_shared_kv_fwd,
)
from sglang.kernels.ops.attention.verify_splitkv import (
verify_splitkv_fwd,
)
from sglang.kernels.ops.attention.verify_mla import verify_shared_kv_fwd
from sglang.kernels.ops.attention.verify_splitkv import verify_splitkv_fwd
super().__init__()
@@ -189,6 +183,15 @@ class TritonAttnBackend(AttentionBackend):
extend_attention_fwd_unified
)
self.build_unified_kv_indices = torch.compiler.disable(build_unified_kv_indices)
# Dense (non-absorbed) MLA prefill over a materialized prefix; see
# handle_attention_triton for when the dispatcher selects it.
self.dense_prefill_attention_fwd = torch.compiler.disable(
dense_prefill_attention_fwd
)
self.can_use_dense_prefill_fp8 = can_use_dense_prefill_fp8
# Cumulative full sequence lengths addressing the one-shot K/V; built
# on first use per forward and reset by init_forward_metadata.
self._dense_one_shot_kv_indptr = None
# Split-KV EAGLE-verify kernel; enabled below once topk is known (valid only at topk == 1).
self.verify_splitkv_fwd = torch.compiler.disable(verify_splitkv_fwd)
# Grouped-head split-KV verify kernel for MLA or one shared local KV head.
@@ -235,6 +238,24 @@ class TritonAttnBackend(AttentionBackend):
self.num_kv_head = model_runner.model_config.get_num_kv_heads(
get_parallel().attn_tp_size, get_parallel().attn_dcp_size
)
mla_config = model_runner.model_config
self.use_dense_fp8_chunked_prefill = (
self.use_mla
and is_gfx95_supported()
and envs.SGLANG_TRITON_DENSE_PREFILL_ATTN.get()
and envs.SGLANG_TRITON_FP8_PREFILL_ATTN.get()
and model_runner.kv_cache_dtype == torch.float8_e4m3fn
and self.num_head == 12
and mla_config.qk_nope_head_dim + mla_config.qk_rope_head_dim == 192
and mla_config.v_head_dim == 128
and mla_config.kv_lora_rank == 512
)
# forward_mha discovers these hooks dynamically. Hiding them when the
# exact Kimi-K3 FP8 configuration is absent keeps all other models on
# their existing code paths.
if not self.use_dense_fp8_chunked_prefill:
self.prepare_chunked_prefill_qkv = None
self.pack_prefix_chunk_kv = None
# The decode kernel's "// Lv" stride trick requires attn_logits.shape[-1]
# to exactly match the layer's v_head_dim, so hybrid SWA models with
# differing SWA/full v_head_dim need a second buffer for SWA layers.
@@ -745,6 +766,7 @@ class TritonAttnBackend(AttentionBackend):
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init auxiliary variables for triton attention backend."""
self._dense_one_shot_kv_indptr = None
bs = forward_batch.batch_size
window_kv_indptr = self.window_kv_indptr
window_kv_indices = None
@@ -1285,6 +1307,186 @@ class TritonAttnBackend(AttentionBackend):
):
pass
@property
def pack_all_prefix_chunks(self) -> bool:
"""Pack every prefix chunk into one FP8 buffer when capacity allows."""
return self.use_dense_fp8_chunked_prefill
@property
def fuse_prefix_into_extend(self) -> bool:
"""Attend the packed prefix and current chunk in one launch."""
return self.use_dense_fp8_chunked_prefill
def prepare_chunked_prefill_qkv(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
forward_batch: ForwardBatch,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Convert the current dense MHA chunk once and reuse Q for prefix passes."""
fp8_dtype = torch.float8_e4m3fn
output_dtype = q.dtype
if output_dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
output_dtype = torch.bfloat16
forward_batch._triton_dense_fp8_output_dtype = output_dtype
if q.dtype != fp8_dtype:
q = q.to(fp8_dtype)
if k.dtype != fp8_dtype:
k = k.to(fp8_dtype)
if v.dtype != fp8_dtype:
v = v.to(fp8_dtype)
return q.contiguous(), k.contiguous(), v.contiguous()
def pack_prefix_chunk_kv(
self,
k_nope: torch.Tensor,
k_pe: torch.Tensor,
v: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Pack a materialized dense prefix directly into unit-scale FP8 K/V."""
return mla_kv_pack_quantize_fp8(
k_nope,
k_pe,
v,
fp8_dtype=torch.float8_e4m3fn,
enable_pdl=False,
)
def _can_run_dense_fp8_chunked_mha(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
layer: RadixAttention,
forward_batch: ForwardBatch,
) -> bool:
return (
self.use_dense_fp8_chunked_prefill
and forward_batch.attn_attend_prefix_cache is not None
and self.forward_metadata.custom_mask is None
and q.dtype == torch.float8_e4m3fn
and k.dtype == torch.float8_e4m3fn
and v.dtype == torch.float8_e4m3fn
and layer.tp_q_head_num == 12
and layer.tp_k_head_num == 12
and layer.qk_head_dim == 192
and layer.v_head_dim == 128
and (layer.sliding_window_size is None or layer.sliding_window_size <= -1)
and layer.logit_cap <= 0
)
def _forward_dense_fp8_chunked_mha(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
layer: RadixAttention,
forward_batch: ForwardBatch,
):
"""Run current or cached dense FP8 K/V through Triton extend attention."""
output_dtype = getattr(
forward_batch, "_triton_dense_fp8_output_dtype", torch.bfloat16
)
output = torch.empty(
(q.shape[0], layer.tp_q_head_num, layer.v_head_dim),
dtype=output_dtype,
device=q.device,
)
prefix_k = getattr(forward_batch, "fused_prefix_k", None)
if prefix_k is not None:
# Prefix is non-causal while the current chunk is causal. The
# extend kernel already implements precisely that two-stage mask.
prefix_v = forward_batch.fused_prefix_v
self.extend_attention_fwd(
q,
k,
v,
output,
prefix_k,
prefix_v,
self.forward_metadata.qo_indptr,
forward_batch.prefix_chunk_cu_seq_lens[0],
forward_batch.prefix_dense_kv_indices[: prefix_k.shape[0]],
None,
True,
None,
self.forward_metadata.max_extend_len,
1.0,
1.0,
sm_scale=layer.scaling,
page_size=1,
extend_seq_lens_cpu=forward_batch.extend_seq_lens_cpu,
identity_kv_indices=True,
)
return output
lse = torch.empty(
(q.shape[0], layer.tp_q_head_num),
dtype=torch.float32,
device=q.device,
)
if forward_batch.attn_attend_prefix_cache:
chunk_idx = forward_batch.prefix_chunk_idx
assert chunk_idx is not None and chunk_idx >= 0
kv_indptr = forward_batch.prefix_chunk_cu_seq_lens[chunk_idx]
kv_indices = forward_batch.prefix_dense_kv_indices[: k.shape[0]]
self.extend_attention_fwd(
q,
k[:0],
v[:0],
output,
k,
v,
self.forward_metadata.qo_indptr,
kv_indptr,
kv_indices,
None,
False,
None,
self.forward_metadata.max_extend_len,
1.0,
1.0,
sm_scale=layer.scaling,
lse_extend=lse,
skip_extend=True,
page_size=1,
extend_seq_lens_cpu=forward_batch.extend_seq_lens_cpu,
identity_kv_indices=True,
)
# Empty ragged rows are returned as output=0, LSE=-inf, so the
# portable merge_state operation ignores them exactly.
else:
self.extend_attention_fwd(
q,
k,
v,
output,
k[:0],
v[:0],
self.forward_metadata.qo_indptr,
forward_batch.mha_empty_kv_indptr,
self.forward_metadata.kv_indices[:0],
None,
True,
None,
self.forward_metadata.max_extend_len,
1.0,
1.0,
sm_scale=layer.scaling,
lse_extend=lse,
skip_prefix=True,
page_size=1,
extend_seq_lens_cpu=forward_batch.extend_seq_lens_cpu,
)
if forward_batch.mha_return_lse:
return output, lse
return output
def _set_kv_buffer(
self,
forward_batch: ForwardBatch,
@@ -1330,6 +1532,16 @@ class TritonAttnBackend(AttentionBackend):
score_mod=None,
aux_tensors=None,
):
if (
k is not None
and v is not None
and sinks is None
and score_mod is None
and aux_tensors is None
and self._can_run_dense_fp8_chunked_mha(q, k, v, layer, forward_batch)
):
return self._forward_dense_fp8_chunked_mha(q, k, v, layer, forward_batch)
# TODO: reuse the buffer across layers
attn_out = getattr(forward_batch, "_attn_output", None)
if attn_out is not None:
@@ -1398,6 +1610,31 @@ class TritonAttnBackend(AttentionBackend):
):
causal = False
# Dense one-shot MLA prefill (AttnForwardMethod.MHA_ONE_SHOT): k/v were
# up-projected out of the latent cache and span prefix + current chunk,
# so they no longer line up row-for-row with q the way
# extend_attention_fwd requires. Route to the single-loop dense kernel.
# A prefix-chunk phase (attn_attend_prefix_cache) also carries a longer
# k/v, but the dispatcher never hands Triton MHA_CHUNKED_KV.
if (
forward_batch.mha_one_shot
and not forward_batch.attn_attend_prefix_cache
and k is not None
and k.shape[0] != q.shape[0]
):
return self._forward_extend_dense_one_shot(
q,
k,
v,
o,
layer,
forward_batch,
causal,
logits_soft_cap,
sinks=sinks,
score_mod=score_mod,
)
if self.dcp_size > 1:
if score_mod is not None:
raise NotImplementedError(
@@ -1517,6 +1754,81 @@ class TritonAttnBackend(AttentionBackend):
)
return o
def _dense_one_shot_kv_indptr_for(self, forward_batch: ForwardBatch):
"""Cumulative full sequence lengths addressing the one-shot K/V rows.
The MHA one-shot K/V is gathered with fetch_mha_one_shot_kv_indices(),
which lays sequences out back to back at their full seq_len -- so the
row offsets are cumsum(seq_lens), not the prefix-only kv_indptr that
forward_metadata carries for the paged extend path.
"""
if self._dense_one_shot_kv_indptr is None:
bs = forward_batch.batch_size
kv_indptr = torch.zeros(bs + 1, dtype=torch.int32, device=self.device)
kv_indptr[1:] = torch.cumsum(forward_batch.seq_lens[:bs], dim=0)
self._dense_one_shot_kv_indptr = kv_indptr
return self._dense_one_shot_kv_indptr
def _forward_extend_dense_one_shot(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
o: torch.Tensor,
layer: RadixAttention,
forward_batch: ForwardBatch,
causal: bool,
logits_soft_cap: float,
sinks: Optional[torch.Tensor] = None,
score_mod=None,
):
# Guarded rather than silently fallen back on: dropping to
# extend_attention_fwd with a longer-than-q k/v would read the wrong
# rows and quietly return wrong numbers.
if sinks is not None or score_mod is not None:
raise NotImplementedError(
"Triton dense one-shot prefill does not support sinks/score_mod"
)
if layer.sliding_window_size is not None and layer.sliding_window_size > -1:
raise NotImplementedError(
"Triton dense one-shot prefill does not support sliding windows"
)
if layer.k_scale is not None or layer.v_scale is not None:
raise NotImplementedError(
"Triton dense one-shot prefill does not support KV descales"
)
if layer.xai_temperature_len is not None and layer.xai_temperature_len > 0:
raise NotImplementedError(
"Triton dense one-shot prefill does not support xai temperature"
)
q = q.view(-1, layer.tp_q_head_num, layer.qk_head_dim)
k = k.view(-1, layer.tp_k_head_num, layer.qk_head_dim)
v = v.view(-1, layer.tp_k_head_num, layer.v_head_dim)
if self.can_use_dense_prefill_fp8(
q, k, v, is_causal=causal, logit_cap=logits_soft_cap
):
# Cast Q, K and V separately, matching the zero-prefix FP8 gate in
# extend_attention_fwd (and Aiter's opt-in behavior).
q = q.to(torch.float8_e4m3fn)
k = k.to(torch.float8_e4m3fn)
v = v.to(torch.float8_e4m3fn)
self.dense_prefill_attention_fwd(
q,
k.contiguous(),
v.contiguous(),
o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
self.forward_metadata.qo_indptr,
self._dense_one_shot_kv_indptr_for(forward_batch),
self.forward_metadata.max_extend_len,
sm_scale=layer.scaling,
logit_cap=logits_soft_cap,
is_causal=causal,
)
return o
def _forward_extend_dcp(
self,
q: torch.Tensor,
@@ -47,9 +47,20 @@ class ForwardBatchDeepSeekMHAMixin:
prefix_chunk_num_tokens: Optional[List[int]] = None
# KV Indices for each chunk
prefix_chunk_kv_indices: Optional[List[torch.Tensor]] = None
# All chunk KV indices in chunk-major order. The packed dense-prefix path
# gathers this once, then slices materialized K/V by chunk token count.
prefix_all_kv_indices: Optional[torch.Tensor] = None
# Reusable identity indices into a temporary dense prefix K/V buffer.
prefix_dense_kv_indices: Optional[torch.Tensor] = None
# Reusable empty prefix indptr for current-chunk-only attention.
mha_empty_kv_indptr: Optional[torch.Tensor] = None
# For MLA chunked prefix cache used in chunked prefill
# Tell attention backend whether lse needs to be returned
mha_return_lse: Optional[bool] = None
# Packed dense prefix K/V for one fused prefix+extend attention launch.
# These are set only for the duration of an attention call.
fused_prefix_k: Optional[torch.Tensor] = None
fused_prefix_v: Optional[torch.Tensor] = None
# Whether to apply MHA_ONE_SHOT forward method
mha_one_shot: Optional[bool] = None
# KV Indices for MHA_ONE_SHOT forward method
@@ -64,7 +75,12 @@ class ForwardBatchDeepSeekMHAMixin:
def set_attn_attend_prefix_cache(self, attn_attend_prefix_cache: bool):
self.attn_attend_prefix_cache = attn_attend_prefix_cache
def prepare_chunked_kv_indices(self, device: torch.device):
def prepare_chunked_kv_indices(
self,
device: torch.device,
pack_all_prefix_chunks: bool = False,
dense_metadata: bool = False,
):
self.prefix_chunk_kv_indices = []
req_to_token = get_req_to_token_pool().req_to_token
for idx in range(self.num_prefix_chunks):
@@ -95,6 +111,24 @@ class ForwardBatchDeepSeekMHAMixin:
chunk_kv_indices = translator.translate_dcp_read_ids(chunk_kv_indices)
self.prefix_chunk_kv_indices.append(chunk_kv_indices)
if not (pack_all_prefix_chunks or dense_metadata):
return
# Only the packed path gathers the whole prefix in one operation. The
# identity indices and empty indptr describe a dense per-chunk buffer
# and are required for both packed and bounded chunked paths.
if pack_all_prefix_chunks:
self.prefix_all_kv_indices = torch.cat(self.prefix_chunk_kv_indices)
max_dense_chunk_tokens = max(
chunk_indices.numel() for chunk_indices in self.prefix_chunk_kv_indices
)
self.prefix_dense_kv_indices = torch.arange(
max_dense_chunk_tokens, dtype=torch.int32, device=device
)
self.mha_empty_kv_indptr = torch.zeros(
self.batch_size + 1, dtype=torch.int32, device=device
)
# Here we suppose the length of each chunk is equal
# For example, if we have 4 sequences with prefix length [256, 512, 768, 1024], prefix_chunk_len = 256
# num_prefix_chunks = cdiv(1024, 256) = 4
@@ -125,25 +159,23 @@ class ForwardBatchDeepSeekMHAMixin:
# Called before each attention module if using chunked kv cache for prefill
# Some of the codes are adapted from https://github.com/vllm-project/vllm/blob/main/vllm/v1/attention/backends/mla/common.py
def prepare_chunked_prefix_cache_info(self, device: torch.device):
def prepare_chunked_prefix_cache_info(
self,
device: torch.device,
pack_all_prefix_chunks: bool = False,
single_chunk: bool = False,
dense_metadata: bool = False,
):
from sglang.srt.mem_cache.memory_pool import (
HybridLinearKVPool,
MLATokenToKVPool,
)
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
token_to_kv_pool = get_token_to_kv_pool()
assert (
isinstance(token_to_kv_pool, MLATokenToKVPool)
or (
isinstance(token_to_kv_pool, HybridLinearKVPool)
and isinstance(token_to_kv_pool.full_kv_pool, MLATokenToKVPool)
)
or (
isinstance(token_to_kv_pool, SWAKVPool)
and isinstance(token_to_kv_pool.full_kv_pool, MLATokenToKVPool)
)
assert isinstance(token_to_kv_pool, MLATokenToKVPool) or (
isinstance(token_to_kv_pool, HybridLinearKVPool)
and isinstance(token_to_kv_pool.full_kv_pool, MLATokenToKVPool)
), "Currently chunked prefix cache can only be used by Deepseek models"
if not any(self.extend_prefix_lens_cpu):
@@ -158,7 +190,13 @@ class ForwardBatchDeepSeekMHAMixin:
# chunk_capacity is the maximum number of tokens in each chunk
chunk_capacity = self.get_max_chunk_capacity()
self.prefix_chunk_len = chunk_capacity // self.batch_size
if single_chunk:
# The caller has already checked that the full packed prefix fits
# chunk_capacity. One request-major chunk gives the fused kernel the
# same per-request layout described by prefix_chunk_cu_seq_lens[0].
self.prefix_chunk_len = max(self.extend_prefix_lens_cpu)
else:
self.prefix_chunk_len = chunk_capacity // self.batch_size
self.num_prefix_chunks = (
max(self.extend_prefix_lens_cpu) + self.prefix_chunk_len - 1
@@ -210,7 +248,7 @@ class ForwardBatchDeepSeekMHAMixin:
]
# Precompute the kv indices for each chunk
self.prepare_chunked_kv_indices(device)
self.prepare_chunked_kv_indices(device, pack_all_prefix_chunks, dense_metadata)
def fetch_mha_one_shot_kv_indices(self):
if self.mha_one_shot_kv_indices is not None:
@@ -1,3 +1,4 @@
from sglang.srt.environ import envs
from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
from sglang.srt.layers.utils.cp_utils import mla_use_prefill_cp
from sglang.srt.model_executor.forward_context import get_attn_backend
@@ -13,17 +14,21 @@ from sglang.srt.models.deepseek_common.attention_forward_methods.forward_methods
from sglang.srt.models.deepseek_common.utils import _is_hip
from sglang.srt.runtime_context import (
get_exec,
get_parallel,
get_platform,
)
from sglang.srt.utils import use_intel_amx_backend
from sglang.srt.utils import (
is_gfx95_supported,
use_intel_amx_backend,
)
MHA_ONE_SHOT_SUPPORTED_BACKENDS = ["fa3", "flashinfer", "flashmla"]
# ROCm runs dedicated MHA/MLA implementations (forward_mha_rocm.py /
# forward_mla_rocm.py) so the shared CUDA paths carry no AMD branches. Backend
# handlers keep returning the generic method; the platform swap happens here.
# MHA_CHUNKED_KV has no ROCm entry because its accumulation step needs the
# CUDA-only merge_state_v2 kernel.
# MHA_CHUNKED_KV deliberately stays generic on ROCm. Its shared implementation
# selects the ROCm prepare/fetch helpers and the portable merge_state wrapper.
_ROCM_FORWARD_METHODS = {
AttnForwardMethod.MHA: AttnForwardMethod.MHA_ROCM,
AttnForwardMethod.MHA_ONE_SHOT: AttnForwardMethod.MHA_ONE_SHOT_ROCM,
@@ -212,6 +217,26 @@ def handle_attention_dsa(attn, forward_batch):
return AttnForwardMethod.MLA
def _can_use_triton_dense_fp8_prefill(attn, forward_batch) -> bool:
prefix_lens = forward_batch.extend_prefix_lens_cpu
return (
_is_hip
and is_gfx95_supported()
and envs.SGLANG_TRITON_FP8_PREFILL_ATTN.get()
and attn.kv_cache_dtype == "fp8_e4m3"
and attn.num_local_heads == 12
and attn.qk_nope_head_dim == 128
and attn.qk_rope_head_dim == 64
and attn.v_head_dim == 128
and attn.kv_lora_rank == 512
and not get_parallel().dcp_enabled
and not mla_use_prefill_cp(forward_batch)
and forward_batch.forward_mode.is_extend_without_speculative()
and prefix_lens is not None
and any(prefix_lens)
)
def handle_attention_triton(attn, forward_batch):
if is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph():
return AttnForwardMethod.MLA
@@ -220,13 +245,18 @@ def handle_attention_triton(attn, forward_batch):
if get_exec().deterministic.enable_deterministic_inference:
return _dispatch_mla_subtype(attn, forward_batch)
# Kimi-K3 with an FP8 latent cache uses dense 192/128 K/V for cached
# prefixes. Always select chunked-KV here: its fast path packs the prefix
# once and fuses it with the current chunk in the normal extend kernel.
if _can_use_triton_dense_fp8_prefill(attn, forward_batch):
return AttnForwardMethod.MHA_CHUNKED_KV
if (
forward_batch.forward_mode.is_extend_without_speculative()
and sum(forward_batch.extend_prefix_lens_cpu) == 0
):
return AttnForwardMethod.MHA
else:
return _dispatch_mla_subtype(attn, forward_batch)
return _dispatch_mla_subtype(attn, forward_batch)
def handle_attention_intel_xpu(attn, forward_batch):
@@ -21,23 +21,18 @@ from sglang.srt.model_executor.forward_context import (
)
from sglang.srt.models.deepseek_common.utils import (
_is_cuda,
_is_hip,
_is_musa,
_is_npu,
_use_aiter_gfx95,
)
from sglang.srt.runtime_context import (
get_exec,
get_parallel,
get_schedule,
)
from sglang.srt.runtime_context import get_exec, get_parallel, get_schedule
from sglang.srt.utils import BumpAllocator, next_power_of_2
if TYPE_CHECKING:
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
if _is_cuda:
from sgl_kernel import merge_state_v2
from sglang.kernels.ops.attention.concat_mla import concat_mla_k
elif _is_musa:
from sgl_kernel import concat_mla_k
@@ -45,9 +40,52 @@ elif _is_musa:
def resolve_attn_backend(forward_batch: ForwardBatch):
backend = get_attn_backend()
if isinstance(backend, TboAttnBackend):
backend = backend.primary
return backend
while True:
if isinstance(backend, TboAttnBackend):
backend = backend.primary
continue
# Hybrid KDA/MLA models route full-attention calls through an outer
# HybridLinearAttnBackend. Model-side MHA preparation hooks belong to
# its full-attention child.
if hasattr(backend, "full_attn_backend"):
backend = backend.full_attn_backend
continue
# A split prefill/decode HybridAttnBackend may itself be the full-attn
# child. Resolve the backend serving this forward mode as well.
if (
hasattr(backend, "prefill_backend")
and hasattr(backend, "decode_backend")
and hasattr(backend, "_select_backend")
):
backend = backend._select_backend(forward_batch.forward_mode)
continue
return backend
def use_dense_prefix_kv(backend) -> bool:
"""Whether the backend substitutes dense K/V for prefix chunks."""
return (
getattr(backend, "pack_prefix_chunk_kv", None) is not None
and getattr(backend, "pack_all_prefix_chunks", False)
and not get_parallel().dcp_enabled
)
def use_packed_prefix_chunks(backend, forward_batch: ForwardBatch) -> bool:
"""Whether the complete prefix can be packed into one bounded buffer."""
prefix_lens = forward_batch.extend_prefix_lens_cpu
return (
use_dense_prefix_kv(backend)
and prefix_lens is not None
and sum(prefix_lens) <= forward_batch.get_max_chunk_capacity()
)
def use_fused_prefix_extend(backend, forward_batch: ForwardBatch) -> bool:
"""Whether the packed prefix and current chunk can share one launch."""
return getattr(
backend, "fuse_prefix_into_extend", False
) and use_packed_prefix_chunks(backend, forward_batch)
def forward_dsa_indexer_for_mha(
@@ -281,7 +319,14 @@ class DeepseekMHAForwardMixin:
# The top comments in https://github.com/vllm-project/vllm/blob/main/vllm/v1/attention/backends/mla/common.py
# will be helpful for understanding the purpose of this function.
# First do normal mha forward to get output for extended part
# Preserve the ROCm fused RMS/quantized projection path for the current
# chunk; the shared preparation path is primarily the CUDA version.
if _is_hip:
return self.forward_normal_rocm_prepare(
positions, hidden_states, forward_batch, zero_allocator
)
# First do normal mha forward to get output for extended part.
return self.forward_normal_prepare(
positions, hidden_states, forward_batch, zero_allocator
)
@@ -297,12 +342,34 @@ class DeepseekMHAForwardMixin:
has_extend_prefix = forward_batch.extend_prefix_lens_cpu is not None and any(
forward_batch.extend_prefix_lens_cpu
)
backend = resolve_attn_backend(forward_batch)
prepare_qkv_fn = getattr(backend, "prepare_chunked_prefill_qkv", None)
if has_extend_prefix and prepare_qkv_fn is not None:
q, k, v = prepare_qkv_fn(q, k, v, forward_batch)
fused_prefix = has_extend_prefix and use_fused_prefix_extend(
backend, forward_batch
)
# Only initialize the info once
if has_extend_prefix and forward_batch.num_prefix_chunks is None:
forward_batch.prepare_chunked_prefix_cache_info(q.device)
forward_batch.prepare_chunked_prefix_cache_info(
q.device,
pack_all_prefix_chunks=use_packed_prefix_chunks(backend, forward_batch),
single_chunk=fused_prefix,
dense_metadata=use_dense_prefix_kv(backend),
)
if hasattr(get_attn_backend(), "init_mha_chunk_metadata"):
get_attn_backend().init_mha_chunk_metadata(forward_batch)
if fused_prefix:
attn_output = self._fused_prefix_extend_attn_mha(q, k, v, forward_batch)
attn_output = attn_output.reshape(
-1, self.num_local_heads * self.v_head_dim
)
output, _ = self.o_proj(attn_output)
return output
forward_batch.mha_return_lse = has_extend_prefix
# Do mha for extended part without prefix
forward_batch.set_attn_attend_prefix_cache(False)
@@ -356,6 +423,48 @@ class DeepseekMHAForwardMixin:
forward_batch.set_attn_attend_prefix_cache(False)
return self.forward_normal_core(q, k, v, forward_batch, gate)
def _fused_prefix_extend_attn_mha(
self: DeepseekV2AttentionMLA,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
forward_batch: ForwardBatch,
) -> torch.Tensor:
"""Attend the complete prefix and current chunk in one kernel launch."""
assert forward_batch.num_prefix_chunks == 1, (
"fused prefix+extend needs the single-chunk layout, got "
f"{forward_batch.num_prefix_chunks} chunks"
)
backend = resolve_attn_backend(forward_batch)
get_mla_kv_buffer = (
self._get_mla_kv_buffer_rocm if _is_hip else self._get_mla_kv_buffer
)
# With one chunk, these indices are request-major and are described by
# prefix_chunk_cu_seq_lens[0]. Project and pack the prefix only once.
kv_a_normed, k_pe = get_mla_kv_buffer(
forward_batch.prefix_all_kv_indices, torch.bfloat16, forward_batch
)
kv = self.kv_b_proj(kv_a_normed)[0]
kv = kv.view(-1, self.num_local_heads, self.qk_nope_head_dim + self.v_head_dim)
prefix_k, prefix_v = backend.pack_prefix_chunk_kv(
kv[..., : self.qk_nope_head_dim],
k_pe,
kv[..., self.qk_nope_head_dim :],
)
del kv_a_normed, k_pe, kv
forward_batch.mha_return_lse = False
forward_batch.set_attn_attend_prefix_cache(False)
forward_batch.set_prefix_chunk_idx(0)
forward_batch.fused_prefix_k = prefix_k
forward_batch.fused_prefix_v = prefix_v
try:
return self.attn_mha(q, k, v, forward_batch, save_kv_cache=False)
finally:
forward_batch.fused_prefix_k = None
forward_batch.fused_prefix_v = None
def _chunked_prefix_attn_mha(
self: DeepseekV2AttentionMLA,
q: torch.Tensor,
@@ -364,46 +473,82 @@ class DeepseekMHAForwardMixin:
forward_batch: ForwardBatch,
) -> torch.Tensor:
# kv_b_proj needs BF16 input, but legacy q.dtype was BF16 by accident.
from sglang.srt.layers.attention.merge_state import merge_state
backend = resolve_attn_backend(forward_batch)
pack_fn = getattr(backend, "pack_prefix_chunk_kv", None)
kv_a_dtype = torch.bfloat16 if pack_fn is not None else q.dtype
get_mla_kv_buffer = (
self._get_mla_kv_buffer_rocm if _is_hip else self._get_mla_kv_buffer
)
assert forward_batch.num_prefix_chunks is not None
for i in range(forward_batch.num_prefix_chunks):
forward_batch.set_prefix_chunk_idx(i)
kv_indices = forward_batch.prefix_chunk_kv_indices[i]
# Fetch latent cache from memory pool with precomputed chunked kv indices
kv_a_normed, k_pe = self._get_mla_kv_buffer(
kv_indices, kv_a_dtype, forward_batch
)
kv_a_normed, k_pe = all_gather_kv_cache_for_mha_chunk_extend(
kv_a_normed,
k_pe,
forward_batch.prefix_chunk_seq_lens_cpu[i],
forward_batch.prefix_chunk_starts_cpu[i],
# If the complete prefix fits the capacity, gather/project/pack once.
# Larger prefixes retain bounded per-chunk materialization.
packed_prefix_k = None
packed_prefix_v = None
pack_all_prefix_chunks = use_packed_prefix_chunks(backend, forward_batch)
if pack_all_prefix_chunks:
kv_a_normed, k_pe = get_mla_kv_buffer(
forward_batch.prefix_all_kv_indices,
kv_a_dtype,
forward_batch,
)
kv = self.kv_b_proj(kv_a_normed)[0]
kv = kv.view(
-1, self.num_local_heads, self.qk_nope_head_dim + self.v_head_dim
)
v = kv[..., self.qk_nope_head_dim :]
k_nope = kv[..., : self.qk_nope_head_dim]
v_dense = kv[..., self.qk_nope_head_dim :]
packed_prefix_k, packed_prefix_v = pack_fn(k_nope, k_pe, v_dense)
del kv_a_normed, k_pe, kv, k_nope, v_dense
if pack_fn is not None:
k, v = pack_fn(k_nope, k_pe, v)
assert forward_batch.num_prefix_chunks is not None
packed_offset = 0
for i in range(forward_batch.num_prefix_chunks):
forward_batch.set_prefix_chunk_idx(i)
if pack_all_prefix_chunks:
chunk_num_tokens = forward_batch.prefix_chunk_num_tokens[i]
packed_end = packed_offset + chunk_num_tokens
k = packed_prefix_k[packed_offset:packed_end]
v = packed_prefix_v[packed_offset:packed_end]
packed_offset = packed_end
else:
k = torch.empty(
(
k_nope.shape[0],
self.num_local_heads,
self.qk_nope_head_dim + self.qk_rope_head_dim,
),
dtype=v.dtype,
device=v.device,
kv_indices = forward_batch.prefix_chunk_kv_indices[i]
kv_a_normed, k_pe = get_mla_kv_buffer(
kv_indices, kv_a_dtype, forward_batch
)
k[..., : self.qk_nope_head_dim] = k_nope
k[..., self.qk_nope_head_dim :] = k_pe
kv_a_normed, k_pe = all_gather_kv_cache_for_mha_chunk_extend(
kv_a_normed,
k_pe,
forward_batch.prefix_chunk_seq_lens_cpu[i],
forward_batch.prefix_chunk_starts_cpu[i],
)
kv = self.kv_b_proj(kv_a_normed)[0]
kv = kv.view(
-1,
self.num_local_heads,
self.qk_nope_head_dim + self.v_head_dim,
)
v_dense = kv[..., self.qk_nope_head_dim :]
k_nope = kv[..., : self.qk_nope_head_dim]
if pack_fn is not None:
k, v = pack_fn(k_nope, k_pe, v_dense)
else:
v = v_dense
k = torch.empty(
(
k_nope.shape[0],
self.num_local_heads,
self.qk_nope_head_dim + self.qk_rope_head_dim,
),
dtype=v.dtype,
device=v.device,
)
k[..., : self.qk_nope_head_dim] = k_nope
k[..., self.qk_nope_head_dim :] = k_pe
del kv_a_normed, k_pe, kv, k_nope, v_dense
output, lse = self.attn_mha(
q,
@@ -418,9 +563,20 @@ class DeepseekMHAForwardMixin:
)
tmp_output = torch.empty_like(accum_output)
tmp_lse = torch.empty_like(accum_lse)
merge_state_v2(output, lse, accum_output, accum_lse, tmp_output, tmp_lse)
merge_state(
output,
lse,
accum_output,
accum_lse,
tmp_output,
tmp_lse,
)
accum_output, accum_lse = tmp_output, tmp_lse
del kv, k, v, output, lse, tmp_output, tmp_lse
del k, v, output, lse, tmp_output, tmp_lse
if pack_all_prefix_chunks:
assert packed_offset == packed_prefix_k.shape[0]
del packed_prefix_k, packed_prefix_v
return accum_output
@@ -0,0 +1,188 @@
"""Dense (non-absorbed) Triton prefill over a materialized prefix + chunk.
Covers ``AttnForwardMethod.MHA_ONE_SHOT`` for Kimi-K3 on the triton backend,
where the cached prefix is up-projected to the 192/128 MHA shape and attended
in one pass instead of running the 576/512 absorbed kernel.
"""
import unittest
import torch
from sglang.kernels.ops.attention.extend_attention import (
can_use_dense_prefill_fp8,
dense_prefill_attention_fwd,
)
from sglang.srt.environ import envs
from sglang.srt.utils import get_device, is_gfx95_supported, is_hip
from sglang.test.ci.ci_register import register_amd_ci
register_amd_ci(est_time=10, suite="stage-b-test-1-gpu-small-amd-mi35x")
H_Q, D_QK, D_V = 12, 192, 128
FP8 = torch.float8_e4m3fn
def _reference(q, k, v, qo_indptr, kv_indptr, scale, is_causal):
"""Bottom-right aligned causal attention, one sequence at a time, in fp32."""
out = torch.empty(q.shape[0], H_Q, D_V, dtype=torch.float32, device=q.device)
lse = torch.empty(q.shape[0], H_Q, dtype=torch.float32, device=q.device)
for i in range(len(qo_indptr) - 1):
q_lo, q_hi = int(qo_indptr[i]), int(qo_indptr[i + 1])
k_lo, k_hi = int(kv_indptr[i]), int(kv_indptr[i + 1])
q_len, kv_len = q_hi - q_lo, k_hi - k_lo
scores = (
torch.matmul(
q[q_lo:q_hi].float().transpose(0, 1),
k[k_lo:k_hi].float().transpose(0, 1).transpose(1, 2),
)
* scale
)
if is_causal:
# Query m sits at absolute position (kv_len - q_len) + m.
q_pos = torch.arange(q_len, device=q.device)[:, None] + (kv_len - q_len)
k_pos = torch.arange(kv_len, device=q.device)[None, :]
scores = scores.masked_fill(q_pos < k_pos, float("-inf"))
probs = torch.softmax(scores, dim=-1)
out[q_lo:q_hi] = torch.matmul(
probs, v[k_lo:k_hi].float().transpose(0, 1)
).transpose(0, 1)
lse[q_lo:q_hi] = torch.logsumexp(scores, dim=-1).transpose(0, 1)
return out, lse
@unittest.skipUnless(
is_hip() and is_gfx95_supported(), "Kimi-K3 dense Triton prefill requires gfx950"
)
class TestKimiK3TritonDensePrefill(unittest.TestCase):
def setUp(self):
torch.manual_seed(42)
torch.cuda.manual_seed_all(42)
self.device = get_device()
self.scale = D_QK**-0.5
def _run(self, q_lens, prefix_lens, *, mode="bf16", is_causal=True, want_lse=False):
device = self.device
kv_lens = [q + p for q, p in zip(q_lens, prefix_lens)]
qo_indptr = torch.zeros(len(q_lens) + 1, dtype=torch.int32, device=device)
kv_indptr = torch.zeros(len(q_lens) + 1, dtype=torch.int32, device=device)
qo_indptr[1:] = torch.tensor(q_lens, device=device).cumsum(0)
kv_indptr[1:] = torch.tensor(kv_lens, device=device).cumsum(0)
total_q, total_kv = int(qo_indptr[-1]), int(kv_indptr[-1])
q = torch.randn(total_q, H_Q, D_QK, dtype=torch.bfloat16, device=device) * 0.25
k = torch.randn(total_kv, H_Q, D_QK, dtype=torch.bfloat16, device=device) * 0.25
v = torch.randn(total_kv, H_Q, D_V, dtype=torch.bfloat16, device=device) * 0.25
# Quantize before taking the reference so the comparison isolates the
# kernel from the cast: an FP8 run should match FP8 inputs exactly.
if mode == "fp8":
q, k, v = q.to(FP8), k.to(FP8), v.to(FP8)
elif mode == "mixed":
# What forward_mha_rocm actually hands over on gfx95 with MXFP4
# kv_b_proj weights: k/v already e4m3 from the fused up-projection,
# q still bf16.
k, v = k.to(FP8), v.to(FP8)
ref_out, ref_lse = _reference(
q.float() if mode == "bf16" else q.to(FP8).float(),
k.float() if mode == "bf16" else k.float(),
v.float() if mode == "bf16" else v.float(),
qo_indptr,
kv_indptr,
self.scale,
is_causal,
)
out = torch.empty(total_q, H_Q, D_V, dtype=torch.bfloat16, device=device)
lse = (
torch.empty(total_q, H_Q, dtype=torch.float32, device=device)
if want_lse
else None
)
dense_prefill_attention_fwd(
q,
k,
v,
out,
qo_indptr,
kv_indptr,
max(q_lens),
sm_scale=self.scale,
is_causal=is_causal,
lse=lse,
)
return out, lse, ref_out, ref_lse
def test_causal_shapes(self):
# Prefix lengths deliberately straddle the BLOCK_N=64 boundary: the
# kernel splits its KV sweep into an unmasked interior and a masked
# tail at a BLOCK_N multiple, so an off-by-one there is only visible
# when the prefix is not a clean multiple.
cases = [
([128], [0]), # no prefix: every block is diagonal
([7], [0]), # q shorter than BLOCK_M
([1], [1000]), # single query, long prefix
([128], [63]),
([128], [64]),
([128], [65]),
([256], [1025]),
([100, 37, 256, 51], [0, 500, 300, 77]), # ragged, mixed prefixes
]
for q_lens, prefix_lens in cases:
with self.subTest(q=q_lens, prefix=prefix_lens):
out, _, ref, _ = self._run(q_lens, prefix_lens)
torch.testing.assert_close(out.float(), ref, rtol=2e-2, atol=2e-2)
def test_non_causal(self):
out, _, ref, _ = self._run([128, 64], [256, 130], is_causal=False)
torch.testing.assert_close(out.float(), ref, rtol=2e-2, atol=2e-2)
def test_lse_matches_reference(self):
out, lse, ref, ref_lse = self._run([192, 64], [300, 129], want_lse=True)
torch.testing.assert_close(out.float(), ref, rtol=2e-2, atol=2e-2)
# LSE is what the chunked-KV merge would consume, so it has to be a
# natural log in absolute terms, not just proportional.
torch.testing.assert_close(lse, ref_lse, rtol=1e-3, atol=1e-3)
def test_fp8_matches_quantized_reference(self):
out, _, ref, _ = self._run([128, 64], [512, 77], mode="fp8")
torch.testing.assert_close(out.float(), ref, rtol=5e-2, atol=5e-2)
def test_mixed_bf16_query_fp8_kv(self):
# Both operands of a tl.dot must share a dtype, so the wrapper has to
# promote q rather than compile a bf16 x fp8 pair. Regression guard:
# this shape reaches the kernel straight from forward_mha_rocm.
out, _, ref, _ = self._run([128, 64], [512, 77], mode="mixed")
torch.testing.assert_close(out.float(), ref, rtol=5e-2, atol=5e-2)
def test_fp8_gate(self):
device = self.device
q = torch.empty(1, H_Q, D_QK, dtype=torch.bfloat16, device=device)
k = torch.empty(1, H_Q, D_QK, dtype=torch.bfloat16, device=device)
v = torch.empty(1, H_Q, D_V, dtype=torch.bfloat16, device=device)
with envs.SGLANG_TRITON_FP8_PREFILL_ATTN.override(True):
self.assertTrue(
can_use_dense_prefill_fp8(q, k, v, is_causal=True, logit_cap=0.0)
)
# Non-causal and logit-capped softmax are outside the gate, and an
# already-quantized input must not be cast a second time.
self.assertFalse(
can_use_dense_prefill_fp8(q, k, v, is_causal=False, logit_cap=0.0)
)
self.assertFalse(
can_use_dense_prefill_fp8(q, k, v, is_causal=True, logit_cap=1.0)
)
self.assertFalse(
can_use_dense_prefill_fp8(
q.to(FP8), k, v, is_causal=True, logit_cap=0.0
)
)
with envs.SGLANG_TRITON_FP8_PREFILL_ATTN.override(False):
self.assertFalse(
can_use_dense_prefill_fp8(q, k, v, is_causal=True, logit_cap=0.0)
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,310 @@
import unittest
import torch
from sglang.kernels.ops.attention.extend_attention import extend_attention_fwd
from sglang.kernels.ops.attention.extend_attention_split_dim import (
can_use_split_dim_absorbed_extend,
)
from sglang.srt.environ import envs
from sglang.srt.utils import get_device, is_gfx95_supported, is_hip
from sglang.test.ci.ci_register import register_amd_ci
register_amd_ci(est_time=15, suite="stage-b-test-1-gpu-small-amd-mi35x")
@unittest.skipUnless(
is_hip() and is_gfx95_supported(), "Kimi-K3 Triton prefill requires gfx950"
)
class TestKimiK3TritonPrefill(unittest.TestCase):
def setUp(self):
torch.manual_seed(42)
torch.cuda.manual_seed_all(42)
def test_split_dim_ragged_bf16(self):
device = get_device()
h_q, h_kv, d_qk, d_v = 12, 1, 576, 512
extend_lens = (97, 128)
prefix_lens = (53, 91)
scale = 192**-0.5
total_extend, total_prefix = sum(extend_lens), sum(prefix_lens)
q = torch.randn(total_extend, h_q, d_qk, dtype=torch.bfloat16, device=device)
k = torch.randn(total_extend, h_kv, d_qk, dtype=torch.bfloat16, device=device)
v = torch.randn(total_extend, h_kv, d_v, dtype=torch.bfloat16, device=device)
k_buffer = torch.randn(
total_prefix, h_kv, d_qk, dtype=torch.bfloat16, device=device
)
v_buffer = torch.randn(
total_prefix, h_kv, d_v, dtype=torch.bfloat16, device=device
)
qo_indptr = torch.tensor(
[0, extend_lens[0], total_extend], dtype=torch.int32, device=device
)
kv_indptr = torch.tensor(
[0, prefix_lens[0], total_prefix], dtype=torch.int32, device=device
)
kv_indices = torch.arange(total_prefix, dtype=torch.int64, device=device)
output = torch.empty(
total_extend, h_q, d_v, dtype=torch.bfloat16, device=device
)
extend_attention_fwd(
q,
k,
v,
output,
k_buffer,
v_buffer,
qo_indptr,
kv_indptr,
kv_indices,
custom_mask=None,
is_causal=True,
mask_indptr=None,
max_len_extend=max(extend_lens),
k_scale=1.0,
v_scale=1.0,
sm_scale=scale,
)
reference = torch.empty_like(output, dtype=torch.float32)
for batch, (extend_len, prefix_len) in enumerate(zip(extend_lens, prefix_lens)):
q_start = int(qo_indptr[batch])
prefix_start = int(kv_indptr[batch])
q_batch = q[q_start : q_start + extend_len].float()
k_prefix = k_buffer[prefix_start : prefix_start + prefix_len, 0].float()
v_prefix = v_buffer[prefix_start : prefix_start + prefix_len, 0].float()
k_current = k[q_start : q_start + extend_len, 0].float()
v_current = v[q_start : q_start + extend_len, 0].float()
causal = torch.triu(
torch.ones(
extend_len,
extend_len,
dtype=torch.bool,
device=device,
),
diagonal=1,
)
for head in range(h_q):
prefix_scores = q_batch[:, head] @ k_prefix.T * scale
current_scores = q_batch[:, head] @ k_current.T * scale
current_scores.masked_fill_(causal, float("-inf"))
scores = torch.cat([prefix_scores, current_scores], dim=1)
values = torch.cat([v_prefix, v_current], dim=0)
reference[q_start : q_start + extend_len, head] = (
torch.softmax(scores, dim=-1) @ values
)
torch.testing.assert_close(output.float(), reference, rtol=1e-2, atol=1e-2)
def test_split_dim_dispatch_gates(self):
device = get_device()
q = torch.empty(1, 12, 576, dtype=torch.bfloat16, device=device)
k = torch.empty(1, 1, 576, dtype=torch.bfloat16, device=device)
v = torch.empty(1, 1, 512, dtype=torch.bfloat16, device=device)
o = torch.empty(1, 12, 512, dtype=torch.bfloat16, device=device)
k_buffer = torch.empty(1, 1, 576, dtype=torch.bfloat16, device=device)
v_buffer = torch.empty(1, 1, 512, dtype=torch.bfloat16, device=device)
kwargs = dict(
lse=None,
sinks=None,
k_scale=1.0,
v_scale=1.0,
custom_mask=None,
is_causal=True,
sliding_window_size=-1,
logit_cap=0.0,
xai_temperature_len=-1,
skip_prefix=False,
skip_extend=False,
page_size=1,
score_mod=None,
aux_tensors=None,
)
self.assertTrue(
can_use_split_dim_absorbed_extend(q, k, v, o, k_buffer, v_buffer, **kwargs)
)
fp8_k_buffer = k_buffer.to(torch.float8_e4m3fn)
fp8_v_buffer = v_buffer.to(torch.float8_e4m3fn)
with envs.SGLANG_TRITON_FP8_PREFILL_ATTN.override(False):
self.assertFalse(
can_use_split_dim_absorbed_extend(
q, k, v, o, fp8_k_buffer, fp8_v_buffer, **kwargs
)
)
with envs.SGLANG_TRITON_FP8_PREFILL_ATTN.override(True):
self.assertTrue(
can_use_split_dim_absorbed_extend(
q, k, v, o, fp8_k_buffer, fp8_v_buffer, **kwargs
)
)
self.assertTrue(
can_use_split_dim_absorbed_extend(
q,
k,
v,
o,
fp8_k_buffer,
fp8_v_buffer,
**{**kwargs, "k_scale": 0.5, "v_scale": 0.25},
)
)
for override in (
{"page_size": 2},
{"logit_cap": 1.0},
{"sliding_window_size": 128},
{"skip_prefix": True},
{"is_causal": False},
{"lse": torch.empty(1, 12, dtype=torch.float32, device=device)},
{"sinks": torch.empty(12, dtype=torch.float32, device=device)},
{"k_scale": 0.5},
):
self.assertFalse(
can_use_split_dim_absorbed_extend(
q,
k,
v,
o,
k_buffer,
v_buffer,
**{**kwargs, **override},
)
)
def test_zero_prefix_fp8_flag(self):
device = get_device()
tokens, heads, d_qk, d_v = 128, 12, 192, 128
q = torch.randn(tokens, heads, d_qk, dtype=torch.bfloat16, device=device) * 0.25
k = torch.randn(tokens, heads, d_qk, dtype=torch.bfloat16, device=device) * 0.25
v = torch.randn(tokens, heads, d_v, dtype=torch.bfloat16, device=device) * 0.25
k_buffer = torch.empty(1, heads, d_qk, dtype=torch.float8_e4m3fn, device=device)
v_buffer = torch.empty(1, heads, d_v, dtype=torch.float8_e4m3fn, device=device)
qo_indptr = torch.tensor([0, tokens], dtype=torch.int32, device=device)
kv_indptr = torch.tensor([0, 0], dtype=torch.int32, device=device)
kv_indices = torch.empty(0, dtype=torch.int64, device=device)
bf16_output = torch.empty(
tokens, heads, d_v, dtype=torch.bfloat16, device=device
)
fp8_output = torch.empty_like(bf16_output)
scale = d_qk**-0.5
def run(output):
extend_attention_fwd(
q,
k,
v,
output,
k_buffer,
v_buffer,
qo_indptr,
kv_indptr,
kv_indices,
custom_mask=None,
is_causal=True,
mask_indptr=None,
max_len_extend=tokens,
k_scale=1.0,
v_scale=1.0,
sm_scale=scale,
)
with envs.SGLANG_TRITON_FP8_PREFILL_ATTN.override(False):
run(bf16_output)
with envs.SGLANG_TRITON_FP8_PREFILL_ATTN.override(True):
run(fp8_output)
causal = (
torch.arange(tokens, device=device)[None, :]
<= torch.arange(tokens, device=device)[:, None]
)
def reference(q_ref, k_ref, v_ref):
scores = torch.einsum("qhd,khd->qhk", q_ref, k_ref) * scale
scores.masked_fill_(~causal[:, None, :], float("-inf"))
return torch.einsum("qhk,khd->qhd", torch.softmax(scores, dim=-1), v_ref)
bf16_reference = reference(q.float(), k.float(), v.float())
fp8_reference = reference(
q.to(torch.float8_e4m3fn).float(),
k.to(torch.float8_e4m3fn).float(),
v.to(torch.float8_e4m3fn).float(),
)
torch.testing.assert_close(
bf16_output.float(), bf16_reference, rtol=1e-2, atol=1e-2
)
torch.testing.assert_close(
fp8_output.float(), fp8_reference, rtol=2e-2, atol=2e-2
)
def test_absorbed_fp8_prefix(self):
device = get_device()
tokens, prefix, heads, d_qk, d_v = 64, 73, 12, 576, 512
q = torch.randn(tokens, heads, d_qk, dtype=torch.bfloat16, device=device) * 0.25
k = torch.randn(tokens, 1, d_qk, dtype=torch.bfloat16, device=device) * 0.25
v = torch.randn(tokens, 1, d_v, dtype=torch.bfloat16, device=device) * 0.25
k_buffer = (
torch.randn(prefix, 1, d_qk, dtype=torch.bfloat16, device=device) * 0.25
).to(torch.float8_e4m3fn)
v_buffer = (
torch.randn(prefix, 1, d_v, dtype=torch.bfloat16, device=device) * 0.25
).to(torch.float8_e4m3fn)
qo_indptr = torch.tensor([0, tokens], dtype=torch.int32, device=device)
kv_indptr = torch.tensor([0, prefix], dtype=torch.int32, device=device)
kv_indices = torch.arange(prefix, dtype=torch.int64, device=device)
output = torch.empty(tokens, heads, d_v, dtype=torch.bfloat16, device=device)
generic_output = torch.empty_like(output)
scale, k_scale, v_scale = 192**-0.5, 0.5, 0.25
def run(candidate):
extend_attention_fwd(
q,
k,
v,
candidate,
k_buffer,
v_buffer,
qo_indptr,
kv_indptr,
kv_indices,
custom_mask=None,
is_causal=True,
mask_indptr=None,
max_len_extend=tokens,
k_scale=k_scale,
v_scale=v_scale,
sm_scale=scale,
)
with envs.SGLANG_TRITON_FP8_PREFILL_ATTN.override(False):
run(generic_output)
with envs.SGLANG_TRITON_FP8_PREFILL_ATTN.override(True):
run(output)
q_fp8 = q.to(torch.float8_e4m3fn).float()
prefix_scores = (
torch.einsum("qhd,kd->qhk", q_fp8, k_buffer[:, 0].float()) * scale * k_scale
)
current_scores = torch.einsum("qhd,kd->qhk", q.float(), k[:, 0].float())
current_scores *= scale
causal = (
torch.arange(tokens, device=device)[None, :]
<= torch.arange(tokens, device=device)[:, None]
)
current_scores.masked_fill_(~causal[:, None, :], float("-inf"))
scores = torch.cat([prefix_scores, current_scores], dim=-1)
values = torch.cat([v_buffer[:, 0].float() * v_scale, v[:, 0].float()], dim=0)
reference = torch.einsum("qhk,kv->qhv", torch.softmax(scores, dim=-1), values)
torch.testing.assert_close(output.float(), reference, rtol=3e-2, atol=3e-2)
torch.testing.assert_close(
generic_output.float(), reference, rtol=3e-2, atol=3e-2
)
torch.testing.assert_close(
output.float(), generic_output.float(), rtol=1e-2, atol=1e-2
)
if __name__ == "__main__":
unittest.main()