[minimax-m3] fp8 attention GEMMs on SM100 (fp8_e4m3 KV + trtllm_mha) (#30971)
Co-authored-by: qiuyue <qiuyue@minimaxi.com> Co-authored-by: xuebi <xuebi@minimaxi.com> Co-authored-by: Xiaoyu Zhang <1182563586@qq.com>
This commit is contained in:
co-authored by
qiuyue
xuebi
Xiaoyu Zhang
parent
e6a4cefc69
commit
bae8eb8d6c
@@ -9,7 +9,7 @@
|
||||
|
||||
#include <cfloat>
|
||||
#include <cstdint>
|
||||
#if defined(__HIP_PLATFORM_AMD__)
|
||||
#ifdef USE_ROCM
|
||||
static constexpr unsigned long long kWarpSyncMask = 0xFFFFFFFFFFFFFFFFull;
|
||||
#else
|
||||
#include <math_constants.h>
|
||||
@@ -30,7 +30,9 @@ namespace {
|
||||
// The trivial case num_blocks <= topk (every block selected) is handled by the
|
||||
// kernels below, outside the Trait.
|
||||
struct TopKTrait {
|
||||
static constexpr uint32_t kMaxTopK = 32;
|
||||
// Also sizes the kernels' smem staging for the ascending-order emit; the
|
||||
// block-id path's test contract goes up to topk == 64.
|
||||
static constexpr uint32_t kMaxTopK = 64;
|
||||
static constexpr uint32_t kCTASize = 512;
|
||||
static constexpr uint32_t kNumWarps = kCTASize / device::kWarpThreads;
|
||||
static constexpr uint32_t kMaxNumBlocks = 4096; // block topk
|
||||
@@ -67,8 +69,14 @@ struct TopKTrait {
|
||||
};
|
||||
constexpr auto warp_inclusive_sum = [](uint32_t lane_id, uint32_t val) {
|
||||
#pragma unroll
|
||||
for (uint32_t offset = 1; offset < 32; offset *= 2) {
|
||||
uint32_t n = __shfl_up_sync(kWarpSyncMask, val, offset, 32);
|
||||
for (uint32_t offset = 1; offset < device::kWarpThreads; offset *= 2) {
|
||||
// Width-32 up-shuffle. On wave64 HIP the un-suffixed __shfl_up takes the
|
||||
// logical-warp width directly; CUDA needs the active mask.
|
||||
#ifdef USE_ROCM
|
||||
uint32_t n = __shfl_up(val, offset, device::kWarpThreads);
|
||||
#else
|
||||
uint32_t n = __shfl_up_sync(kWarpSyncMask, val, offset, device::kWarpThreads);
|
||||
#endif
|
||||
if (lane_id >= offset) val += n;
|
||||
}
|
||||
return val;
|
||||
@@ -264,8 +272,10 @@ struct TopKTrait {
|
||||
// Trait; otherwise the Trait selects the top-k block ids.
|
||||
// -------------------------------------------------------------------------
|
||||
|
||||
// Block-id output: topk_idx[h, b, 0:k_eff) = selected block ids (front-packed,
|
||||
// unordered), [k_eff:topk) = -1.
|
||||
// Block-id output: topk_idx[h, b, 0:k_eff) = selected block ids (sorted
|
||||
// ascending), [k_eff:topk) = -1. Ascending order is a hard requirement of the
|
||||
// MSA fmha_sm100 consumer (kv_block_indexes must be strictly ascending; its
|
||||
// sorted-order early-exit otherwise mis-masks the partial last block).
|
||||
template <typename SeqLenT, bool kUsePDL>
|
||||
__global__ void minimax_decode_topk_block_kernel(
|
||||
const float* __restrict__ score,
|
||||
@@ -297,7 +307,43 @@ __global__ void minimax_decode_topk_block_kernel(
|
||||
|
||||
const float* __restrict__ row = score + (static_cast<int64_t>(h) * batch + b) * max_seqblock;
|
||||
__shared__ TopKTrait::Smem smem;
|
||||
TopKTrait::forward(row, static_cast<uint32_t>(num_blocks), out, static_cast<uint32_t>(topk), &smem);
|
||||
__shared__ int32_t s_topk[TopKTrait::kMaxTopK];
|
||||
TopKTrait::forward(row, static_cast<uint32_t>(num_blocks), s_topk, static_cast<uint32_t>(topk), &smem);
|
||||
__syncthreads(); // s_topk fully written before the sort reads it
|
||||
|
||||
// Emit ascending: num_blocks > topk here, so all topk slots hold distinct
|
||||
// ids and rank(v) = |{x : x < v}| is a permutation. deepseek_v4
|
||||
// topk_impl.cuh warp sort (32x32 / 64x64 branches; topk <= kMaxTopK = 64):
|
||||
// lanes hold the elements in registers (INT32_MAX sentinel past topk), warp
|
||||
// w ranks targets {w, w + kNumWarps, ...} via ballot+popc, lane 0 emits.
|
||||
static_assert(TopKTrait::kMaxTopK <= 2 * device::kWarpThreads);
|
||||
const auto warp_id = tx / device::kWarpThreads;
|
||||
const auto lane_id = tx % device::kWarpThreads;
|
||||
#ifdef USE_ROCM
|
||||
// wave64: __ballot spans the full 64-lane wave (would count the sibling
|
||||
// 32-lane logical warp too); use the file's width-32 shuffle reduction.
|
||||
const auto count_lt = [](int32_t x, int32_t v) { return device::warp::reduce_sum(static_cast<int>(x < v)); };
|
||||
#else
|
||||
const auto count_lt = [](int32_t x, int32_t v) { return __popc(__ballot_sync(kWarpSyncMask, x < v)); };
|
||||
#endif
|
||||
if (topk <= static_cast<int>(device::kWarpThreads)) { // 32 x 32
|
||||
const int32_t tie = (lane_id < static_cast<uint32_t>(topk)) ? s_topk[lane_id] : INT32_MAX;
|
||||
for (uint32_t t = warp_id; t < static_cast<uint32_t>(topk); t += TopKTrait::kNumWarps) {
|
||||
const int32_t target = s_topk[t];
|
||||
const auto rank = count_lt(tie, target);
|
||||
if (lane_id == 0) out[rank] = target;
|
||||
}
|
||||
} else { // 64 x 64: each lane takes 2 elements
|
||||
const int32_t tie_0 = s_topk[lane_id];
|
||||
const int32_t tie_1 = (lane_id + device::kWarpThreads < static_cast<uint32_t>(topk))
|
||||
? s_topk[lane_id + device::kWarpThreads]
|
||||
: INT32_MAX;
|
||||
for (uint32_t t = warp_id; t < static_cast<uint32_t>(topk); t += TopKTrait::kNumWarps) {
|
||||
const int32_t target = s_topk[t];
|
||||
const auto rank = count_lt(tie_0, target) + count_lt(tie_1, target);
|
||||
if (lane_id == 0) out[rank] = target;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Page-table output: for each (batch b, kv-head h) pseudo-request emit the
|
||||
@@ -436,6 +482,7 @@ void minimax_decode_topk(
|
||||
topk_i,
|
||||
")");
|
||||
RuntimeCheck(block_size > 0, "block_size must be > 0, got ", block_size);
|
||||
RuntimeCheck(topk <= static_cast<int64_t>(TopKTrait::kMaxTopK), "topk exceeds kMaxTopK (ascending-sort smem buffer)");
|
||||
if (batch == 0 || num_heads == 0) return;
|
||||
|
||||
const dim3 grid(static_cast<unsigned>(batch), static_cast<unsigned>(num_heads));
|
||||
|
||||
@@ -3,8 +3,10 @@
|
||||
Drop-in replacement for the 2-stage split-K Triton topk
|
||||
(``_topk_index_partial_kernel`` + ``_topk_index_merge_kernel``): given the
|
||||
decode score tensor ``[num_heads, batch, max_seqblock]`` it produces
|
||||
``topk_idx`` ``[num_heads, batch, topk]`` (0-indexed block ids, front-packed,
|
||||
``-1`` padded), matching the consumer ``_gqa_share_sparse_decode_kernel``.
|
||||
``topk_idx`` ``[num_heads, batch, topk]`` (0-indexed block ids, sorted
|
||||
ascending, ``-1`` padded at the tail). Ascending order is required by the MSA
|
||||
fmha_sm100 consumer; the Triton ``_gqa_share_sparse_decode_kernel`` is
|
||||
order-insensitive.
|
||||
|
||||
``minimax_decode_topk_page_table`` additionally fuses the page-table transform
|
||||
for the dense paged backend (trtllm_mha / fa3) and returns the page table plus
|
||||
@@ -17,7 +19,12 @@ from typing import TYPE_CHECKING, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args
|
||||
from sglang.kernels.jit.utils import (
|
||||
cache_once,
|
||||
is_arch_support_pdl,
|
||||
load_jit,
|
||||
make_cpp_args,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from tvm_ffi.module import Module
|
||||
@@ -25,7 +32,7 @@ if TYPE_CHECKING:
|
||||
|
||||
@cache_once
|
||||
def _jit_module(seq_dtype: torch.dtype) -> Module:
|
||||
args = make_cpp_args(seq_dtype, True) # SeqLenT, kUsePDL
|
||||
args = make_cpp_args(seq_dtype, is_arch_support_pdl()) # SeqLenT, kUsePDL
|
||||
return load_jit(
|
||||
"minimax_decode_topk",
|
||||
*args,
|
||||
|
||||
@@ -10,11 +10,14 @@ import triton.language as tl
|
||||
|
||||
_tma_keep_alive_buf = deque(maxlen=200)
|
||||
|
||||
# Q is always bf16/fp16. The paged main K/V cache may be fp8 (unit-scaled) under
|
||||
# --kv-cache-dtype fp8_*; the kernel widens it to the Q dtype on load (IS_FP8
|
||||
# branch). Accepted on both HIP and CUDA (the bf16->fp8 cache write is unit-scaled,
|
||||
# so the widening cast is the exact inverse dequant). The bf16/fp16-only MSA
|
||||
# (fmha_sm100) kernel is excluded for fp8 KV by the backend use_msa gate.
|
||||
# The paged main K/V cache may be fp8 (unit-scaled) under --kv-cache-dtype
|
||||
# fp8_*; with a bf16/fp16 Q the kernel widens K/V to the Q dtype on load
|
||||
# (IS_FP8 branch — the bf16->fp8 cache write is unit-scaled, so the widening
|
||||
# cast is the exact inverse dequant). Under fp8 attn-GEMM mode Q itself is
|
||||
# fp8_e4m3: the IS_FP8 casts become no-ops and tl.dot runs fp8x8 on tensor
|
||||
# cores (P is quantized to the V dtype for the PV MMA, same contract as the
|
||||
# fmha_sm100 fp8 kernel). Accepted on both HIP and CUDA. MSA (fmha_sm100)
|
||||
# accepts fp8 only in the uniform-e4m3 fp8 attn-GEMM mode (backend gate).
|
||||
SPARSE_KV_FP8_DTYPES = (
|
||||
torch.float8_e4m3fn,
|
||||
torch.float8_e5m2,
|
||||
@@ -25,26 +28,56 @@ SPARSE_KV_FP8_DTYPES = (
|
||||
def check_sparse_kv_fp8(
|
||||
q: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
v_cache: torch.Tensor,
|
||||
v_cache: Optional[torch.Tensor],
|
||||
*,
|
||||
label: str,
|
||||
) -> bool:
|
||||
"""Validate the sparse-attention KV cache dtype contract.
|
||||
"""Validate the sparse-attention Q/KV dtype contract.
|
||||
|
||||
Returns True iff the K cache is fp8 (then widened to Q dtype in the kernel).
|
||||
Raises AssertionError otherwise, mirroring the contract the decode and prefill
|
||||
topk kernels both enforce. fp8 is accepted on both HIP and CUDA.
|
||||
Returns True iff the K cache is fp8 (drives the kernels' IS_FP8 constexpr).
|
||||
Two fp8 modes are allowed:
|
||||
* widening (Q bf16/fp16, K/V any fp8): K/V widened to Q dtype on load;
|
||||
* all-fp8 GEMM (fp8 attn-GEMM mode): Q/K/V all fp8_e4m3fn. e5m2 Q is
|
||||
rejected — fmha_sm100's variant lookup silently mis-dispatches e5m2
|
||||
to the e4m3 kernel, so uniform e4m3 is enforced on the sglang side.
|
||||
"""
|
||||
assert q.dtype in (torch.bfloat16, torch.float16)
|
||||
is_fp8 = k_cache.dtype in SPARSE_KV_FP8_DTYPES
|
||||
assert k_cache.dtype == q.dtype or is_fp8, (
|
||||
f"sparse {label} expects K cache dtype == Q dtype ({q.dtype}) "
|
||||
f"or fp8, got {k_cache.dtype}"
|
||||
)
|
||||
assert v_cache.dtype == k_cache.dtype
|
||||
if q.dtype == torch.float8_e4m3fn:
|
||||
assert k_cache.dtype == torch.float8_e4m3fn, (
|
||||
f"sparse {label} with fp8 Q requires an fp8_e4m3fn K cache, "
|
||||
f"got {k_cache.dtype}"
|
||||
)
|
||||
else:
|
||||
assert q.dtype in (
|
||||
torch.bfloat16,
|
||||
torch.float16,
|
||||
), f"sparse {label} expects Q dtype bf16/fp16/fp8_e4m3fn, got {q.dtype}"
|
||||
assert k_cache.dtype == q.dtype or is_fp8, (
|
||||
f"sparse {label} expects K cache dtype == Q dtype ({q.dtype}) "
|
||||
f"or fp8, got {k_cache.dtype}"
|
||||
)
|
||||
if v_cache is not None:
|
||||
assert v_cache.dtype == k_cache.dtype
|
||||
return is_fp8
|
||||
|
||||
|
||||
def sparse_out_dtype(q: torch.Tensor) -> torch.dtype:
|
||||
"""Attention output dtype: bf16 for fp8 Q (fp8 accumulates to bf16 out,
|
||||
matching fmha_sm100's fp8 variant), else the Q dtype."""
|
||||
return torch.bfloat16 if q.dtype == torch.float8_e4m3fn else q.dtype
|
||||
|
||||
|
||||
def unit_scale(scale: Optional[float]) -> float:
|
||||
"""Normalize an optional per-tensor dequant scale: None means unit scale.
|
||||
|
||||
All sparse-op entry points take ``Optional[float] = None`` scales (matching
|
||||
``k_scale_float`` / ``v_scale_float`` on RadixAttention, which are None
|
||||
unless a checkpoint provides them) and normalize here at the kernel-launch
|
||||
boundary, where a concrete float is needed.
|
||||
"""
|
||||
return 1.0 if scale is None else scale
|
||||
|
||||
|
||||
try:
|
||||
make_tensor_descriptor = tl.make_tensor_descriptor
|
||||
except Exception:
|
||||
@@ -260,3 +293,26 @@ def _bitonic_merge(
|
||||
for i in tl.static_range(stage):
|
||||
x, ids = _compare_and_swap(x, ids, flip, i + (n_dims - stage), n_dims)
|
||||
return x, ids
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _sort_ids_ascending(ids, valid_n, BLOCK_SIZE_T: tl.constexpr):
|
||||
"""Sort a top-k id row ascending, invalids packed as -1 at the tail.
|
||||
|
||||
``ids``: int32 [BLOCK_SIZE_T] block ids (0-indexed, -1 = invalid); entries at
|
||||
positions >= ``valid_n`` are also treated as invalid. Valid ids must be
|
||||
distinct and < 2**30 (block ids always are). The MSA fmha_sm100 consumer
|
||||
requires kv_block_indexes strictly ascending with -1 tail padding — its
|
||||
sorted-order early-exit otherwise mis-masks the partial last block.
|
||||
|
||||
O(T^2) rank sort (T = BLOCK_SIZE_T <= 64): invalid entries get unique keys
|
||||
above every valid id so ranks form a permutation.
|
||||
"""
|
||||
off = tl.arange(0, BLOCK_SIZE_T)
|
||||
invalid = (off >= valid_n) | (ids < 0)
|
||||
key = tl.where(invalid, 0x40000000 + off, ids)
|
||||
rank = tl.sum(tl.where(key[None, :] < key[:, None], 1, 0), axis=1)
|
||||
sorted_key = tl.sum(
|
||||
tl.where(rank[None, :] == off[:, None], key[None, :], 0), axis=1
|
||||
)
|
||||
return tl.where(sorted_key >= 0x40000000, -1, sorted_key)
|
||||
|
||||
@@ -8,7 +8,14 @@ import triton.language as tl
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
|
||||
from ..common.utils import _bitonic_merge, robust_allocator
|
||||
from ..common.utils import (
|
||||
_bitonic_merge,
|
||||
_sort_ids_ascending,
|
||||
check_sparse_kv_fp8,
|
||||
robust_allocator,
|
||||
sparse_out_dtype,
|
||||
unit_scale,
|
||||
)
|
||||
|
||||
|
||||
@triton.heuristics(
|
||||
@@ -53,6 +60,8 @@ def _decode_score_kernel(
|
||||
topk: tl.constexpr,
|
||||
# sm_scale
|
||||
sm_scale,
|
||||
# per-tensor K dequant scale (1.0 when the cache is unit-scaled)
|
||||
k_scale,
|
||||
# init and local blocks
|
||||
init_blocks,
|
||||
local_blocks,
|
||||
@@ -75,6 +84,7 @@ def _decode_score_kernel(
|
||||
NUM_KV_CHUNKS: tl.constexpr,
|
||||
SCORE_TYPE: tl.constexpr,
|
||||
SKIP_TRIVIAL_TOPK_SCORE: tl.constexpr,
|
||||
IS_FP8: tl.constexpr,
|
||||
):
|
||||
tl.static_assert(SCORE_TYPE == "max" or SCORE_TYPE == "lse")
|
||||
sm_scale_log2e = sm_scale * 1.4426950409
|
||||
@@ -160,11 +170,15 @@ def _decode_score_kernel(
|
||||
mask=dim_mask[:, None] & pos_mask[None, :],
|
||||
other=0.0,
|
||||
)
|
||||
if IS_FP8:
|
||||
# fp8 index K cache: widening cast with bf16/fp16 Q, no-op with fp8
|
||||
# Q (fp8 attn-GEMM mode; tl.dot runs fp8x8). Compiled out for bf16.
|
||||
k = k.to(q.dtype)
|
||||
# compute qk
|
||||
qk = tl.zeros((BLOCK_SIZE_H, BLOCK_SIZE_N), dtype=tl.float32)
|
||||
qk += tl.where(off_n[None, :] < chunk_end - i, 0, float("-inf"))
|
||||
# [H, D], [D, N] -> [H, N]
|
||||
qk += tl.dot(q, k) * sm_scale_log2e
|
||||
qk += tl.dot(q, k) * (sm_scale_log2e * k_scale)
|
||||
# save qk to score
|
||||
score = tl.reshape(
|
||||
qk,
|
||||
@@ -237,6 +251,9 @@ def _decode_score_attn_kernel(
|
||||
topk: tl.constexpr,
|
||||
# sm_scale
|
||||
sm_scale,
|
||||
# per-tensor KV dequant scales (1.0 when the cache is unit-scaled)
|
||||
k_scale,
|
||||
v_scale,
|
||||
# init and local blocks
|
||||
init_blocks,
|
||||
local_blocks,
|
||||
@@ -271,6 +288,7 @@ def _decode_score_attn_kernel(
|
||||
HAS_SINK: tl.constexpr,
|
||||
SCORE_TYPE: tl.constexpr,
|
||||
SKIP_TRIVIAL_TOPK_SCORE: tl.constexpr,
|
||||
IS_FP8: tl.constexpr,
|
||||
):
|
||||
tl.static_assert(SCORE_TYPE == "max" or SCORE_TYPE == "lse")
|
||||
sm_scale_log2e = sm_scale * 1.4426950409
|
||||
@@ -367,6 +385,10 @@ def _decode_score_attn_kernel(
|
||||
mask=dim_mask[:, None] & pos_mask[None, :],
|
||||
other=0.0,
|
||||
)
|
||||
if IS_FP8:
|
||||
# fp8 index K cache: widening cast with bf16/fp16 Q, no-op with fp8
|
||||
# Q (fp8 attn-GEMM mode; tl.dot runs fp8x8). Compiled out for bf16.
|
||||
k = k.to(q.dtype)
|
||||
# load V as (BLOCK_SIZE_N, head_dim) via indirect addressing
|
||||
v_off = (
|
||||
slots[:, None] * stride_v_s
|
||||
@@ -378,11 +400,15 @@ def _decode_score_attn_kernel(
|
||||
mask=pos_mask[:, None] & dim_mask[None, :],
|
||||
other=0.0,
|
||||
)
|
||||
if IS_FP8:
|
||||
# Cast V to the compute dtype (widening for bf16/fp16 Q; no-op for
|
||||
# fp8 Q where P is quantized to e4m3 for the fp8 PV MMA).
|
||||
v = v.to(q.dtype)
|
||||
# compute qk
|
||||
qk = tl.zeros((BLOCK_SIZE_H, BLOCK_SIZE_N), dtype=tl.float32)
|
||||
qk += tl.where(off_n[None, :] < chunk_end - i, 0, float("-inf"))
|
||||
# [H, D], [D, N] -> [H, N]
|
||||
qk += tl.dot(q, k) * sm_scale_log2e
|
||||
qk += tl.dot(q, k) * (sm_scale_log2e * k_scale)
|
||||
# save qk to score
|
||||
score = tl.reshape(
|
||||
qk,
|
||||
@@ -426,7 +452,7 @@ def _decode_score_attn_kernel(
|
||||
acc_o_scale = tl.exp2(m_i - m_ij)
|
||||
acc_o = acc_o * acc_o_scale[:, None]
|
||||
# [H, N], [N, D] -> [H, D]
|
||||
acc_o += tl.dot(p.to(v.dtype), v)
|
||||
acc_o += tl.dot(p.to(v.dtype), v) * v_scale
|
||||
m_i = m_ij
|
||||
l_i = l_i * acc_o_scale + l_ij
|
||||
# update ptrs
|
||||
@@ -744,8 +770,14 @@ def _topk_index_merge_kernel(
|
||||
+ pid_b * stride_tif_b
|
||||
+ off_t * stride_tif_t
|
||||
)
|
||||
topk_idx_final = tl.where(off_t < tl.minimum(topk, num_blocks), topk_idx_final, -1)
|
||||
tl.store(tif_ptrs, topk_idx_final.to(ti_final_ptr.dtype.element_ty))
|
||||
# Ascending by block id, -1 tail: the MSA fmha_sm100 consumer requires
|
||||
# sorted kv_block_indexes (the bitonic pass above orders by score).
|
||||
topk_idx_final = _sort_ids_ascending(
|
||||
topk_idx_final, tl.minimum(topk, num_blocks), BLOCK_SIZE_T
|
||||
)
|
||||
tl.store(
|
||||
tif_ptrs, topk_idx_final.to(ti_final_ptr.dtype.element_ty), mask=off_t < topk
|
||||
)
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
@@ -768,18 +800,21 @@ def flash_decode_with_topk_idx(
|
||||
disable_index_value: bool = False,
|
||||
use_dense_main_attn: bool = False, # NOTE: need transform idx in this case
|
||||
page_size: int = 1,
|
||||
q_scale: Optional[float] = None,
|
||||
k_scale: Optional[float] = None,
|
||||
v_scale: Optional[float] = None,
|
||||
) -> torch.Tensor:
|
||||
assert score_type in (
|
||||
"max",
|
||||
"lse",
|
||||
), f"score_type must be 'max' or 'lse', got {score_type!r}"
|
||||
triton.set_allocator(robust_allocator)
|
||||
# dtype check
|
||||
assert (
|
||||
q.dtype == torch.bfloat16
|
||||
or q.dtype == torch.float16
|
||||
and k_cache.dtype == q.dtype
|
||||
# dtype check (v_cache is None under disable_index_value)
|
||||
is_fp8 = check_sparse_kv_fp8(
|
||||
q, k_cache, None if disable_index_value else v_cache, label="decode indexer"
|
||||
)
|
||||
k_scale = unit_scale(k_scale)
|
||||
v_scale = unit_scale(v_scale)
|
||||
if not disable_index_value:
|
||||
assert v_cache is not None
|
||||
# shape
|
||||
@@ -793,6 +828,10 @@ def flash_decode_with_topk_idx(
|
||||
# sm scale
|
||||
if sm_scale is None:
|
||||
sm_scale = head_dim**-0.5
|
||||
# q_scale folds exactly into sm_scale: it multiplies every Q-side logit —
|
||||
# the QK dot AND the sink logit — unlike k_scale, which must not touch the
|
||||
# sink term and therefore stays a separate kernel argument.
|
||||
sm_scale = sm_scale * unit_scale(q_scale)
|
||||
# NUM_KV_CHUNKS controls how many parallel chunks each (batch, kv_head) gets.
|
||||
# Total CTAs = batch_size * NUM_KV_CHUNKS * num_kv_heads.
|
||||
# TARGET_GRID is the desired total CTA count; NUM_KV_CHUNKS is derived by:
|
||||
@@ -848,6 +887,7 @@ def flash_decode_with_topk_idx(
|
||||
block_size,
|
||||
topk,
|
||||
sm_scale,
|
||||
k_scale,
|
||||
init_blocks,
|
||||
local_blocks,
|
||||
q.stride(0),
|
||||
@@ -863,6 +903,7 @@ def flash_decode_with_topk_idx(
|
||||
NUM_KV_CHUNKS=NUM_KV_CHUNKS,
|
||||
SCORE_TYPE=score_type,
|
||||
SKIP_TRIVIAL_TOPK_SCORE=skip_trivial_topk_score,
|
||||
IS_FP8=is_fp8,
|
||||
)
|
||||
else:
|
||||
assert v_cache is not None
|
||||
@@ -871,7 +912,7 @@ def flash_decode_with_topk_idx(
|
||||
batch_size,
|
||||
num_q_heads,
|
||||
head_dim,
|
||||
dtype=q.dtype,
|
||||
dtype=sparse_out_dtype(q),
|
||||
device=q.device,
|
||||
)
|
||||
lse = torch.empty(
|
||||
@@ -895,6 +936,8 @@ def flash_decode_with_topk_idx(
|
||||
block_size,
|
||||
topk,
|
||||
sm_scale,
|
||||
k_scale,
|
||||
v_scale,
|
||||
init_blocks,
|
||||
local_blocks,
|
||||
q.stride(0),
|
||||
@@ -922,6 +965,7 @@ def flash_decode_with_topk_idx(
|
||||
NUM_KV_CHUNKS=NUM_KV_CHUNKS,
|
||||
SCORE_TYPE=score_type,
|
||||
SKIP_TRIVIAL_TOPK_SCORE=skip_trivial_topk_score,
|
||||
IS_FP8=is_fp8,
|
||||
)
|
||||
# Fused top-k + page-table transform: emit the dense backend's page table
|
||||
# directly (page-size-aware) instead of block ids, skipping a separate gather.
|
||||
|
||||
@@ -6,7 +6,12 @@ import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from ..common.utils import check_sparse_kv_fp8, robust_allocator
|
||||
from ..common.utils import (
|
||||
check_sparse_kv_fp8,
|
||||
robust_allocator,
|
||||
sparse_out_dtype,
|
||||
unit_scale,
|
||||
)
|
||||
|
||||
|
||||
@triton.heuristics(
|
||||
@@ -49,6 +54,9 @@ def _gqa_share_sparse_decode_kernel(
|
||||
max_kv_len,
|
||||
# sm_scale
|
||||
sm_scale,
|
||||
# per-tensor KV dequant scales (1.0 when the cache is unit-scaled)
|
||||
k_scale,
|
||||
v_scale,
|
||||
# stride
|
||||
stride_q_b,
|
||||
stride_q_h,
|
||||
@@ -178,10 +186,12 @@ def _gqa_share_sparse_decode_kernel(
|
||||
other=0.0,
|
||||
)
|
||||
if IS_FP8:
|
||||
# fp8 KV cache is unit-scaled (set_kv_buffer casts bf16->fp8 with no
|
||||
# scale), so dequant is just a widening cast to the Q compute dtype
|
||||
# before the tl.dot. Matches the bf16 path bit-for-bit when the cache
|
||||
# is bf16 (IS_FP8 False -> this branch is compiled out).
|
||||
# fp8 KV cache: with bf16/fp16 Q this widens K to the compute dtype
|
||||
# (unit-scaled cache -> exact inverse dequant; k_scale covers
|
||||
# calibrated caches). With fp8 Q (fp8 attn-GEMM mode) the cast is a
|
||||
# no-op and tl.dot below runs fp8x8 on tensor cores. Matches the
|
||||
# bf16 path bit-for-bit when the cache is bf16 (IS_FP8 False ->
|
||||
# this branch is compiled out).
|
||||
k = k.to(q.dtype)
|
||||
# load V as (BLOCK_SIZE_N, head_dim) via indirect addressing
|
||||
v_off = (
|
||||
@@ -195,15 +205,16 @@ def _gqa_share_sparse_decode_kernel(
|
||||
other=0.0,
|
||||
)
|
||||
if IS_FP8:
|
||||
# Widen V before the P@V dot. This also makes the `p.to(v.dtype)`
|
||||
# below cast P to the compute dtype (not to fp8, which would be
|
||||
# catastrophic precision loss on the attention weights).
|
||||
# Cast V to the compute dtype. With bf16/fp16 Q this widens (so the
|
||||
# `p.to(v.dtype)` below keeps P in the compute dtype); with fp8 Q it
|
||||
# is a no-op and P is quantized to e4m3 for the fp8 PV MMA — the
|
||||
# same accuracy contract as fmha_sm100's fp8 kernel.
|
||||
v = v.to(q.dtype)
|
||||
# compute qk
|
||||
qk = tl.zeros((BLOCK_SIZE_H, BLOCK_SIZE_N), dtype=tl.float32)
|
||||
qk += tl.where(off_n[None, :] < seq_len - c, 0, float("-inf"))
|
||||
# [H, D], [D, N] -> [H, N]
|
||||
qk += tl.dot(q, k) * sm_scale
|
||||
qk += tl.dot(q, k) * (sm_scale * k_scale)
|
||||
# compute m_ij and l_ij
|
||||
m_ij = tl.maximum(m_i, tl.max(qk, axis=1))
|
||||
p = tl.exp(qk - m_ij[:, None])
|
||||
@@ -212,9 +223,8 @@ def _gqa_share_sparse_decode_kernel(
|
||||
acc_o_scale = tl.exp(m_i - m_ij)
|
||||
acc_o = acc_o * acc_o_scale[:, None]
|
||||
# load v and update acc_o
|
||||
p = p.to(v.dtype)
|
||||
# [H, N], [N, D] -> [H, D]
|
||||
acc_o += tl.dot(p.to(v.dtype), v)
|
||||
acc_o += tl.dot(p.to(v.dtype), v) * v_scale
|
||||
# update statistics
|
||||
m_i = m_ij
|
||||
lse_i = m_ij + tl.log(tl.exp(lse_i - m_ij) + l_ij)
|
||||
@@ -308,9 +318,14 @@ def flash_decode_with_gqa_share_sparse(
|
||||
topk_idx: torch.Tensor, # [num_kv_heads, batch_size, topk]
|
||||
sm_scale: Optional[float] = None,
|
||||
use_tma: bool = True,
|
||||
q_scale: Optional[float] = None,
|
||||
k_scale: Optional[float] = None,
|
||||
v_scale: Optional[float] = None,
|
||||
) -> torch.Tensor:
|
||||
triton.set_allocator(robust_allocator)
|
||||
is_fp8 = check_sparse_kv_fp8(q, k_cache, v_cache, label="decode")
|
||||
k_scale = unit_scale(k_scale)
|
||||
v_scale = unit_scale(v_scale)
|
||||
# shape
|
||||
batch_size, num_q_heads, head_dim = q.shape
|
||||
max_slots, num_kv_heads, _ = k_cache.shape
|
||||
@@ -328,6 +343,9 @@ def flash_decode_with_gqa_share_sparse(
|
||||
# sm scale
|
||||
if sm_scale is None:
|
||||
sm_scale = head_dim**-0.5
|
||||
# q_scale multiplies every Q-side logit (QK dot and sink), so it folds into
|
||||
# sm_scale; k_scale must not touch the sink term and stays a kernel arg.
|
||||
sm_scale = sm_scale * unit_scale(q_scale)
|
||||
# Pick NUM_TOPK_CHUNKS so total grid ≈ TARGET_GRID. Same constraints as
|
||||
# flash_decode_with_topk_idx: must be power of 2 (Triton arange) and must
|
||||
# only depend on shape constants (so grid is fixed within a cuda graph).
|
||||
@@ -345,7 +363,7 @@ def flash_decode_with_gqa_share_sparse(
|
||||
batch_size,
|
||||
num_q_heads,
|
||||
head_dim,
|
||||
dtype=q.dtype,
|
||||
dtype=sparse_out_dtype(q),
|
||||
device=q.device,
|
||||
)
|
||||
lse_partial = torch.empty(
|
||||
@@ -375,6 +393,8 @@ def flash_decode_with_gqa_share_sparse(
|
||||
max_topk,
|
||||
max_kv_len,
|
||||
sm_scale,
|
||||
k_scale,
|
||||
v_scale,
|
||||
q.stride(0),
|
||||
q.stride(1),
|
||||
q.stride(2),
|
||||
|
||||
@@ -6,7 +6,15 @@ import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from ..common.utils import _bitonic_merge, get_cu_seqblocks, robust_allocator
|
||||
from ..common.utils import (
|
||||
_bitonic_merge,
|
||||
_sort_ids_ascending,
|
||||
check_sparse_kv_fp8,
|
||||
get_cu_seqblocks,
|
||||
robust_allocator,
|
||||
sparse_out_dtype,
|
||||
unit_scale,
|
||||
)
|
||||
|
||||
|
||||
@triton.heuristics(
|
||||
@@ -81,6 +89,9 @@ def _flash_attn_fwd_with_block_score_kernel(
|
||||
block_size: tl.constexpr,
|
||||
# sm_scale
|
||||
sm_scale,
|
||||
# per-tensor KV dequant scales (1.0 when the cache is unit-scaled)
|
||||
k_scale,
|
||||
v_scale,
|
||||
# gumbel topk
|
||||
use_gumbel_topk: tl.constexpr,
|
||||
gumbel_seed,
|
||||
@@ -112,6 +123,7 @@ def _flash_attn_fwd_with_block_score_kernel(
|
||||
HAS_SINK: tl.constexpr,
|
||||
SCORE_TYPE: tl.constexpr,
|
||||
DISABLE_INDEX_VALUE: tl.constexpr,
|
||||
IS_FP8: tl.constexpr,
|
||||
):
|
||||
tl.static_assert(SCORE_TYPE == "max" or SCORE_TYPE == "lse")
|
||||
sm_scale_log2e = sm_scale * 1.4426950409
|
||||
@@ -170,7 +182,10 @@ def _flash_attn_fwd_with_block_score_kernel(
|
||||
if HAS_SINK:
|
||||
m_i = tl.zeros((BLOCK_SIZE_Q,), dtype=tl.float32)
|
||||
lse_i = tl.zeros((BLOCK_SIZE_Q,), dtype=tl.float32)
|
||||
qsink = tl.sum(q * sink[None, :], axis=1) * sm_scale_log2e # (BLOCK_SIZE_Q,)
|
||||
qsink = (
|
||||
tl.sum(q.to(tl.float32) * sink[None, :].to(tl.float32), axis=1)
|
||||
* sm_scale_log2e
|
||||
) # (BLOCK_SIZE_Q,)
|
||||
m_i += qsink
|
||||
lse_i += qsink
|
||||
else:
|
||||
@@ -199,8 +214,12 @@ def _flash_attn_fwd_with_block_score_kernel(
|
||||
mask=kd_mask[:, None] & pos_mask[None, :],
|
||||
other=0.0,
|
||||
)
|
||||
if IS_FP8:
|
||||
# fp8 index K cache: widening cast with bf16/fp16 Q, no-op with fp8
|
||||
# Q (fp8 attn-GEMM mode; tl.dot runs fp8x8). Compiled out for bf16.
|
||||
k = k.to(q.dtype)
|
||||
# compute qk
|
||||
qk = tl.dot(q, k) * sm_scale_log2e
|
||||
qk = tl.dot(q, k) * (sm_scale_log2e * k_scale)
|
||||
if i >= diag_start:
|
||||
qk = tl.where(off_q[:, None] >= (i + off_k)[None, :], qk, float("-inf"))
|
||||
# K boundary mask: positions beyond seq_len contribute -inf
|
||||
@@ -251,8 +270,12 @@ def _flash_attn_fwd_with_block_score_kernel(
|
||||
mask=pos_mask[:, None] & vd_mask[None, :],
|
||||
other=0.0,
|
||||
)
|
||||
if IS_FP8:
|
||||
# Cast V to the compute dtype (widening for bf16/fp16 Q; no-op
|
||||
# for fp8 Q where P is quantized to e4m3 for the fp8 PV MMA).
|
||||
v = v.to(q.dtype)
|
||||
p = p.to(v.dtype)
|
||||
acc_o += tl.dot(p, v)
|
||||
acc_o += tl.dot(p, v) * v_scale
|
||||
# update statistics
|
||||
m_i = m_ij
|
||||
lse_i = m_ij + tl.log2(tl.exp2(lse_i - m_ij) + l_ij)
|
||||
@@ -402,6 +425,9 @@ def _topk_index_kernel(
|
||||
* tl.reshape(topk_idx - 1, [BLOCK_SIZE_K // BLOCK_SIZE_T, BLOCK_SIZE_T]),
|
||||
axis=0,
|
||||
)
|
||||
# Ascending by block id, -1 tail: the MSA fmha_sm100 consumer requires
|
||||
# sorted kv_block_indexes (the bitonic pass above orders by score).
|
||||
topk_idx = _sort_ids_ascending(topk_idx, min(topk, valid_blocks), BLOCK_SIZE_T)
|
||||
# save topk
|
||||
ti_ptrs = (
|
||||
ti_ptr
|
||||
@@ -438,15 +464,21 @@ def flash_prefill_with_topk_index(
|
||||
cu_seqblocks_q: Optional[torch.Tensor] = None,
|
||||
max_seqblock_q: Optional[int] = None,
|
||||
all_seqblock_q: Optional[int] = None,
|
||||
q_scale: Optional[float] = None,
|
||||
k_scale: Optional[float] = None,
|
||||
v_scale: Optional[float] = None,
|
||||
):
|
||||
assert score_type in (
|
||||
"max",
|
||||
"lse",
|
||||
), f"score_type must be 'max' or 'lse', got {score_type!r}"
|
||||
triton.set_allocator(robust_allocator)
|
||||
# dtype check
|
||||
assert q.dtype == torch.bfloat16 or q.dtype == torch.float16
|
||||
assert k_cache.dtype == q.dtype
|
||||
# dtype check (v_cache is None under disable_index_value)
|
||||
is_fp8 = check_sparse_kv_fp8(
|
||||
q, k_cache, None if disable_index_value else v_cache, label="prefill indexer"
|
||||
)
|
||||
k_scale = unit_scale(k_scale)
|
||||
v_scale = unit_scale(v_scale)
|
||||
assert cu_seqlens.dtype == torch.int32
|
||||
# shape
|
||||
total_q, num_heads, qk_head_dim = q.shape
|
||||
@@ -455,7 +487,7 @@ def flash_prefill_with_topk_index(
|
||||
# placeholder for BLOCK_SIZE_VD; V is never loaded
|
||||
v_head_dim = qk_head_dim
|
||||
else:
|
||||
assert v_cache is not None and v_cache.dtype == q.dtype
|
||||
assert v_cache is not None
|
||||
assert v_cache.shape[1] == k_cache.shape[1]
|
||||
v_head_dim = v_cache.shape[-1]
|
||||
gqa_group_size = num_heads // num_kv_heads
|
||||
@@ -468,6 +500,9 @@ def flash_prefill_with_topk_index(
|
||||
), "init_blocks + local_blocks must be less than topk"
|
||||
if sm_scale is None:
|
||||
sm_scale = qk_head_dim**-0.5
|
||||
# q_scale multiplies every Q-side logit (QK dot and sink), so it folds into
|
||||
# sm_scale; k_scale must not touch the sink term and stays a kernel arg.
|
||||
sm_scale = sm_scale * unit_scale(q_scale)
|
||||
if cu_seqblocks_q is None or max_seqblock_q is None or all_seqblock_q is None:
|
||||
cu_seqblocks_q, max_seqblock_q, all_seqblock_q, _, _, _ = get_cu_seqblocks(
|
||||
cu_seqlens, max_seqlen_q, block_size_q, block_size_k
|
||||
@@ -476,7 +511,9 @@ def flash_prefill_with_topk_index(
|
||||
if disable_index_value:
|
||||
o = None
|
||||
else:
|
||||
o = torch.empty(total_q, num_heads, v_head_dim, dtype=q.dtype, device=q.device)
|
||||
o = torch.empty(
|
||||
total_q, num_heads, v_head_dim, dtype=sparse_out_dtype(q), device=q.device
|
||||
)
|
||||
score = torch.full(
|
||||
(num_heads, total_q, max_seqblock_k),
|
||||
float("-inf"),
|
||||
@@ -507,6 +544,8 @@ def flash_prefill_with_topk_index(
|
||||
v_head_dim,
|
||||
block_size_k,
|
||||
sm_scale,
|
||||
k_scale,
|
||||
v_scale,
|
||||
False,
|
||||
1,
|
||||
q.stride(0),
|
||||
@@ -529,6 +568,7 @@ def flash_prefill_with_topk_index(
|
||||
req_to_token.stride(0),
|
||||
SCORE_TYPE=score_type,
|
||||
DISABLE_INDEX_VALUE=disable_index_value,
|
||||
IS_FP8=is_fp8,
|
||||
)
|
||||
|
||||
# topk extraction kernel
|
||||
|
||||
@@ -6,7 +6,13 @@ import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from ..common.utils import check_sparse_kv_fp8, get_cu_seqblocks, robust_allocator
|
||||
from ..common.utils import (
|
||||
check_sparse_kv_fp8,
|
||||
get_cu_seqblocks,
|
||||
robust_allocator,
|
||||
sparse_out_dtype,
|
||||
unit_scale,
|
||||
)
|
||||
|
||||
|
||||
@triton.heuristics(
|
||||
@@ -66,6 +72,9 @@ def _gqa_share_sparse_fwd_kernel(
|
||||
num_q_loop,
|
||||
# sm_scale
|
||||
sm_scale,
|
||||
# per-tensor KV dequant scales (1.0 when the cache is unit-scaled)
|
||||
k_scale,
|
||||
v_scale,
|
||||
# stride
|
||||
stride_qn,
|
||||
stride_qh,
|
||||
@@ -201,8 +210,10 @@ def _gqa_share_sparse_fwd_kernel(
|
||||
other=0.0,
|
||||
)
|
||||
if IS_FP8:
|
||||
# fp8 main K cache is unit-scaled; widen to the Q compute dtype
|
||||
# before the tl.dot (compiled out when the cache is bf16).
|
||||
# fp8 main K cache: widening cast with bf16/fp16 Q (unit-scaled
|
||||
# cache -> exact inverse dequant; k_scale covers calibrated
|
||||
# caches), no-op with fp8 Q (fp8 attn-GEMM mode) so tl.dot runs
|
||||
# fp8x8. Compiled out when the cache is bf16.
|
||||
k = k.to(q.dtype)
|
||||
# compute qk
|
||||
qk = tl.zeros((BLOCK_SIZE_Q, BLOCK_SIZE_H, BLOCK_SIZE_K), dtype=tl.float32)
|
||||
@@ -211,7 +222,7 @@ def _gqa_share_sparse_fwd_kernel(
|
||||
qk = tl.reshape(qk, BLOCK_SIZE_QH, BLOCK_SIZE_K)
|
||||
# [BLOCK_SIZE_QH, qk_head_dim] @ [qk_head_dim, BLOCK_SIZE_K]
|
||||
# -> [BLOCK_SIZE_QH, BLOCK_SIZE_K]
|
||||
qk += tl.dot(q, k) * sm_scale_log2e
|
||||
qk += tl.dot(q, k) * (sm_scale_log2e * k_scale)
|
||||
# K boundary mask: positions beyond seq_len contribute -inf
|
||||
qk += tl.where(pos_mask[None, :], 0, float("-inf"))
|
||||
# compute m_ij and l_ij
|
||||
@@ -231,11 +242,13 @@ def _gqa_share_sparse_fwd_kernel(
|
||||
other=0.0,
|
||||
)
|
||||
if IS_FP8:
|
||||
# Widen V so `p.to(v.dtype)` casts P to the compute dtype rather
|
||||
# than to fp8 (which would wreck attention-weight precision).
|
||||
# Cast V to the compute dtype: widening with bf16/fp16 Q (so
|
||||
# `p.to(v.dtype)` keeps P in the compute dtype), no-op with fp8
|
||||
# Q where P is quantized to e4m3 for the fp8 PV MMA — the same
|
||||
# accuracy contract as fmha_sm100's fp8 kernel.
|
||||
v = v.to(q.dtype)
|
||||
p = p.to(v.dtype)
|
||||
acc_o += tl.dot(p, v)
|
||||
acc_o += tl.dot(p, v) * v_scale
|
||||
# update statistics
|
||||
m_i = m_ij
|
||||
lse_i = m_ij + tl.log2(tl.exp2(lse_i - m_ij) + l_ij)
|
||||
@@ -273,9 +286,14 @@ def flash_prefill_with_gqa_share_sparse(
|
||||
use_tma: bool = True,
|
||||
cu_seqblocks_q: Optional[torch.Tensor] = None,
|
||||
max_seqblock_q: Optional[int] = None,
|
||||
q_scale: Optional[float] = None,
|
||||
k_scale: Optional[float] = None,
|
||||
v_scale: Optional[float] = None,
|
||||
) -> torch.Tensor:
|
||||
triton.set_allocator(robust_allocator)
|
||||
is_fp8 = check_sparse_kv_fp8(q, k_cache, v_cache, label="prefill")
|
||||
k_scale = unit_scale(k_scale)
|
||||
v_scale = unit_scale(v_scale)
|
||||
assert block_size_q in {1, 2, 4, 8, 16, 32, 64}
|
||||
assert block_size_k in {16, 32, 64, 128}
|
||||
# shape
|
||||
@@ -292,12 +310,17 @@ def flash_prefill_with_gqa_share_sparse(
|
||||
assert gqa_group_size * block_size_q <= 128
|
||||
if sm_scale is None:
|
||||
sm_scale = qk_head_dim**-0.5
|
||||
# q_scale multiplies every Q-side logit (QK dot and sink), so it folds into
|
||||
# sm_scale; k_scale must not touch the sink term and stays a kernel arg.
|
||||
sm_scale = sm_scale * unit_scale(q_scale)
|
||||
if cu_seqblocks_q is None or max_seqblock_q is None:
|
||||
cu_seqblocks_q, max_seqblock_q, _, _, _, _ = get_cu_seqblocks(
|
||||
cu_seqlens, max_seqlen_q, block_size_q, block_size_k
|
||||
)
|
||||
# output tensor
|
||||
o = torch.empty(total_q, num_q_heads, v_head_dim, device=q.device, dtype=q.dtype)
|
||||
o = torch.empty(
|
||||
total_q, num_q_heads, v_head_dim, device=q.device, dtype=sparse_out_dtype(q)
|
||||
)
|
||||
# launch kernel
|
||||
num_q_loop = (
|
||||
max_seqblock_q // 131072 + 1
|
||||
@@ -330,6 +353,8 @@ def flash_prefill_with_gqa_share_sparse(
|
||||
topk,
|
||||
num_q_loop,
|
||||
sm_scale,
|
||||
k_scale,
|
||||
v_scale,
|
||||
q.stride(0),
|
||||
q.stride(1),
|
||||
q.stride(2),
|
||||
|
||||
@@ -512,13 +512,26 @@ def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
overrides["disable_custom_all_reduce"] = True
|
||||
elif is_sm100_supported():
|
||||
if server_args.is_attention_backend_not_set():
|
||||
overrides["attention_backend"] = "fa4"
|
||||
if (
|
||||
server_args.kv_cache_dtype == "fp8_e4m3"
|
||||
and not envs.SGLANG_DISABLE_M3_FP8_ATTN_GEMM.get()
|
||||
):
|
||||
# fp8 attention GEMMs activate whenever possible
|
||||
# (m3_fp8_attn_gemm_enabled); only trtllm_mha serves the dense
|
||||
# fp8-q path, so prefer it over fa4 for fp8 KV. The
|
||||
# SGLANG_DISABLE_M3_FP8_ATTN_GEMM kill switch keeps the fa4
|
||||
# default (pre-fp8 behavior).
|
||||
overrides["attention_backend"] = "trtllm_mha"
|
||||
else:
|
||||
overrides["attention_backend"] = "fa4"
|
||||
backend_resolved = overrides.get(
|
||||
"attention_backend", server_args.attention_backend
|
||||
)
|
||||
page_resolved = server_args.page_size
|
||||
if (
|
||||
page_resolved is None
|
||||
and overrides.get("attention_backend", server_args.attention_backend)
|
||||
== "fa4"
|
||||
):
|
||||
# fa4 (fmha_sm100) and trtllm_mha both allow the page_size == 128
|
||||
# sparse block MSA needs (trtllm_mha via trtllm-gen's dynamic
|
||||
# tokens-per-page kernels).
|
||||
if page_resolved is None and backend_resolved in ("fa4", "trtllm_mha"):
|
||||
overrides["page_size"] = 128
|
||||
page_resolved = 128
|
||||
if server_args.moe_runner_backend == "auto" and quant_resolved == "mxfp8":
|
||||
@@ -550,6 +563,41 @@ def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
"(MSA is SM100-only; sparse attention runs on the Triton path)."
|
||||
)
|
||||
|
||||
# fp8 attention GEMMs have no opt-in flag: m3_fp8_attn_gemm_enabled
|
||||
# (server_args.py) derives the mode from kv_cache_dtype (fp8_e4m3) +
|
||||
# attention_backend (trtllm_mha) + SM100 at runtime. Surface the
|
||||
# resolution here: warn on fp8_e5m2 (fmha_sm100's variant lookup would
|
||||
# silently dispatch the e4m3 kernel, so e5m2 stays on the widening Triton
|
||||
# path), log when the fp8 GEMM mode is active, and log when the
|
||||
# SGLANG_DISABLE_M3_FP8_ATTN_GEMM kill switch suppresses it.
|
||||
if server_args.kv_cache_dtype == "fp8_e5m2":
|
||||
logger.warning(
|
||||
"MiniMax-M3 with kv_cache_dtype fp8_e5m2: fp8 attention GEMMs stay "
|
||||
"DISABLED (fmha_sm100's variant lookup would silently dispatch the "
|
||||
"e4m3 kernel for e5m2); sparse attention runs on the widening "
|
||||
"Triton path. Use --kv-cache-dtype fp8_e4m3 for fp8 attention GEMMs."
|
||||
)
|
||||
elif (
|
||||
server_args.kv_cache_dtype == "fp8_e4m3"
|
||||
and overrides.get("attention_backend", server_args.attention_backend)
|
||||
== "trtllm_mha"
|
||||
and is_sm100_supported()
|
||||
):
|
||||
if envs.SGLANG_DISABLE_M3_FP8_ATTN_GEMM.get():
|
||||
logger.info(
|
||||
"MiniMax-M3 fp8 attention GEMMs DISABLED by "
|
||||
"SGLANG_DISABLE_M3_FP8_ATTN_GEMM: bf16 indexer + widening "
|
||||
"Triton sparse path, bf16 q; dense layers keep trtllm_mha's "
|
||||
"fp8 KV cache."
|
||||
)
|
||||
else:
|
||||
logger.info(
|
||||
"MiniMax-M3 fp8 attention GEMMs active (kv_cache_dtype fp8_e4m3 + "
|
||||
"trtllm_mha on SM100): fp8 main/index KV, fp8-cast q, fp8 "
|
||||
"sparse/MSA kernels. Set SGLANG_DISABLE_M3_FP8_ATTN_GEMM=1 to "
|
||||
"force the pre-fp8 numerics."
|
||||
)
|
||||
|
||||
moe_runner_resolved = overrides.get(
|
||||
"moe_runner_backend", server_args.moe_runner_backend
|
||||
)
|
||||
@@ -1794,9 +1842,12 @@ def _mla_backend_page_constraints(view: Any) -> dict:
|
||||
or view.decode_attention_backend == "trtllm_mha"
|
||||
or view.prefill_attention_backend == "trtllm_mha"
|
||||
):
|
||||
if page_size not in [16, 32, 64]:
|
||||
# 128 runs on trtllm-gen's dynamic tokens-per-page kernels (flashinfer
|
||||
# >= 0.6.12), which require GQA and equal QK/V head dims — validated at
|
||||
# TRTLLMHAAttnBackend init where the model config is known.
|
||||
if page_size not in [16, 32, 64, 128]:
|
||||
logger.warning(
|
||||
f"TensorRT-LLM MHA only supports page_size of 16, 32 or 64, changing page_size from {page_size} to 64."
|
||||
f"TensorRT-LLM MHA only supports page_size of 16, 32, 64 or 128, changing page_size from {page_size} to 64."
|
||||
)
|
||||
page_size = 64
|
||||
if (
|
||||
|
||||
@@ -1153,6 +1153,10 @@ class Envs:
|
||||
SGLANG_OPT_USE_MINIMAX_DENSE_SPARSE_DECODE = EnvBool(False)
|
||||
SGLANG_DISABLE_MSA = EnvBool(False)
|
||||
SGLANG_OPT_USE_MSA_DECODE_UNDER_GRAPH = EnvBool(False)
|
||||
# Kill switch for the derived fp8 attention-GEMM mode (m3_fp8_attn_gemm_enabled):
|
||||
# forces the pre-fp8 behavior (bf16 indexer + widening sparse path, bf16 q)
|
||||
# even when kv_cache_dtype fp8_e4m3 + trtllm_mha + SM100 would activate it.
|
||||
SGLANG_DISABLE_M3_FP8_ATTN_GEMM = EnvBool(False)
|
||||
|
||||
# MiniMax-M3 sparse decode indexer: single JIT radix-select kernel replaces the 2-stage split-K Triton topk.
|
||||
SGLANG_OPT_USE_MINIMAX_DECODE_TOPK_RADIX = EnvBool(True)
|
||||
|
||||
@@ -18,6 +18,7 @@ from sglang.srt.layers.attention.minimax_sparse_ops.minimax_sparse import (
|
||||
)
|
||||
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.server_args import m3_fp8_attn_gemm_enabled
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
@@ -25,12 +26,26 @@ if TYPE_CHECKING:
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _quant_q_fp8(q: torch.Tensor, q_scale: Optional[float]) -> torch.Tensor:
|
||||
# Same convention as the KV pools: the fp8 tensor stores value/scale and
|
||||
# the attention kernels multiply the logits back by the scale (None = unit).
|
||||
if q_scale is not None:
|
||||
q = q / q_scale
|
||||
return q.to(torch.float8_e4m3fn)
|
||||
|
||||
|
||||
class MiniMaxSparseAttnBackend(AttentionBackend):
|
||||
def __init__(self, runner: ModelRunner):
|
||||
assert isinstance(runner.token_to_kv_pool, MiniMaxSparseKVPool)
|
||||
self.kv_pool = runner.token_to_kv_pool
|
||||
self.req_to_token = runner.req_to_token_pool.req_to_token
|
||||
self.max_context_len = int(runner.model_config.context_len)
|
||||
self.fp8_attn_gemm = m3_fp8_attn_gemm_enabled(runner.server_args)
|
||||
if self.fp8_attn_gemm:
|
||||
assert self.kv_pool.main_pool.dtype == torch.float8_e4m3fn, (
|
||||
"fp8 attn-GEMM mode requires an fp8_e4m3fn main KV pool, got "
|
||||
f"{self.kv_pool.main_pool.dtype}"
|
||||
)
|
||||
|
||||
hf_config = runner.model_config.hf_config
|
||||
sparse_cfg = get_minimax_sparse_attention_config(hf_config)
|
||||
@@ -72,19 +87,26 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
||||
msa_available,
|
||||
)
|
||||
|
||||
# MSA (fmha_sm100) is bf16/fp16-only; an fp8 main KV cache must stay on the
|
||||
# Triton sparse path (it dequants fp8 on load).
|
||||
# MSA (fmha_sm100) runs bf16, or uniform fp8_e4m3 under fp8 attn-GEMM mode
|
||||
# (which also casts q to fp8). An fp8 main KV cache WITHOUT the flag
|
||||
# would pair a bf16 q with fp8 K/V — unsupported by fmha_sm100's
|
||||
# uniform-dtype kernels — so it stays on the Triton sparse path (which
|
||||
# dequants fp8 on load). e5m2 is never allowed into MSA (fmha_sm100's
|
||||
# variant lookup would silently dispatch the e4m3 kernel).
|
||||
_main_kv_is_fp8 = self.kv_pool.main_pool.dtype in (
|
||||
torch.float8_e4m3fn,
|
||||
torch.float8_e5m2,
|
||||
)
|
||||
_msa_fp8_ok = (
|
||||
self.fp8_attn_gemm and self.kv_pool.main_pool.dtype == torch.float8_e4m3fn
|
||||
)
|
||||
self.use_msa = (
|
||||
not envs.SGLANG_DISABLE_MSA.get()
|
||||
and msa_available()
|
||||
and self.block_size_k == 128
|
||||
and self.kv_pool.page_size == self.block_size_k
|
||||
and self.topk_blocks in (4, 8, 16, 32)
|
||||
and not _main_kv_is_fp8
|
||||
and (not _main_kv_is_fp8 or _msa_fp8_ok)
|
||||
)
|
||||
if (
|
||||
not self.use_msa
|
||||
@@ -96,8 +118,8 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
||||
logger.warning(
|
||||
"MiniMax-M3 MSA decode disabled: page_size=%d != sparse block size "
|
||||
"%d. Pass --page-size 128 (with an attention backend that allows it, "
|
||||
"e.g. fa4) to enable the faster MSA kernel; falling back to the "
|
||||
"Triton sparse path.",
|
||||
"e.g. fa4 or trtllm_mha) to enable the faster MSA kernel; falling "
|
||||
"back to the Triton sparse path.",
|
||||
self.kv_pool.page_size,
|
||||
self.block_size_k,
|
||||
)
|
||||
@@ -118,11 +140,10 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
||||
self.use_dense_sparse_decode = (
|
||||
envs.SGLANG_OPT_USE_MINIMAX_DENSE_SPARSE_DECODE.get()
|
||||
and self.block_size_k % self.page_size == 0
|
||||
# _dense_sparse_main_decode calls trtllm decode with a bf16 q and
|
||||
# unit bmm scales — no fp8 handling yet (follow-up).
|
||||
and not self.fp8_attn_gemm
|
||||
)
|
||||
# MSA fmha_sm100 decode is NOT cuda-graph-safe: captured/replayed it returns
|
||||
# wrong results (~14% GSM8K loss on B200). Gate capture via cuda_graph_config,
|
||||
# not legacy disable_* flags — they disagree under config-native flags and would
|
||||
# capture the unsafe MSA decode kernel.
|
||||
from sglang.srt.model_executor.cuda_graph_config import (
|
||||
Backend,
|
||||
Phase,
|
||||
@@ -158,8 +179,18 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
||||
f"[MiniMaxSparse] Backend initialized "
|
||||
f"(score_type={self.score_type!r}, "
|
||||
f"main_attn={'MSA' if self.use_msa else 'triton'}, "
|
||||
f"msa_decode={self._use_msa_decode}, "
|
||||
f"msa_owns_decode={self._msa_owns_decode}, "
|
||||
f"decode_cuda_graph={_decode_cuda_graph}, "
|
||||
f"fp8_attn_gemm={self.fp8_attn_gemm}, "
|
||||
f"disable_value_layers={sorted(self.disable_value_layer_ids)})"
|
||||
)
|
||||
if self.fp8_attn_gemm and self.use_msa:
|
||||
logger.info(
|
||||
"[MiniMaxSparse] fp8 MSA active: the first forward may "
|
||||
"JIT-compile fmha_sm100 fp8 kernel variants (cold cache can "
|
||||
"take minutes; compiles serialize across TP ranks)."
|
||||
)
|
||||
|
||||
def init_forward_metadata_out_graph(
|
||||
self, forward_batch: ForwardBatch, in_capture: bool = False
|
||||
@@ -201,6 +232,7 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
||||
self.topk_blocks,
|
||||
bs,
|
||||
device=device,
|
||||
is_fp8=self.fp8_attn_gemm,
|
||||
)
|
||||
kv_indices_buf = torch.zeros(
|
||||
bs * self._msa_nb_max, dtype=torch.int32, device=device
|
||||
@@ -288,6 +320,10 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
||||
v,
|
||||
idx_k,
|
||||
None if disable_value else idx_v,
|
||||
layer.k_scale_float,
|
||||
layer.v_scale_float,
|
||||
layer.idx_k_scale_float,
|
||||
layer.idx_v_scale_float,
|
||||
)
|
||||
k_cache, v_cache = self.kv_pool.get_kv_buffer(layer.layer_id)
|
||||
if disable_value:
|
||||
@@ -321,6 +357,12 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
||||
q = q[:actual_num_tokens]
|
||||
idx_q = idx_q[:actual_num_tokens]
|
||||
|
||||
# fp8 attention GEMMs: quantize q/idx_q AFTER the KV store (which reads
|
||||
# the bf16 k/v) and the DP trim.
|
||||
if self.fp8_attn_gemm:
|
||||
q = _quant_q_fp8(q, layer.q_scale_float)
|
||||
idx_q = _quant_q_fp8(idx_q, layer.idx_q_scale_float)
|
||||
|
||||
idx_o, o = minimax_sparse_prefill(
|
||||
q,
|
||||
k_cache,
|
||||
@@ -346,6 +388,12 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
||||
disable_index_value=disable_value,
|
||||
use_msa=self.use_msa,
|
||||
seqlens_cpu=forward_batch.extend_seq_lens_cpu,
|
||||
q_scale=layer.q_scale_float,
|
||||
k_scale=layer.k_scale_float,
|
||||
v_scale=layer.v_scale_float,
|
||||
idx_q_scale=layer.idx_q_scale_float,
|
||||
idx_k_scale=layer.idx_k_scale_float,
|
||||
idx_v_scale=layer.idx_v_scale_float,
|
||||
)
|
||||
|
||||
if actual_num_tokens < original_num_tokens:
|
||||
@@ -424,6 +472,10 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
||||
v,
|
||||
idx_k,
|
||||
None if disable_value else idx_v,
|
||||
layer.k_scale_float,
|
||||
layer.v_scale_float,
|
||||
layer.idx_k_scale_float,
|
||||
layer.idx_v_scale_float,
|
||||
)
|
||||
k_cache, v_cache = self.kv_pool.get_kv_buffer(layer.layer_id)
|
||||
if disable_value:
|
||||
@@ -458,6 +510,12 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
||||
"did not prepare the plan for this forward (gate mismatch)."
|
||||
)
|
||||
|
||||
# fp8 attention GEMMs: quantize q/idx_q AFTER the KV store (which reads
|
||||
# the bf16 k/v).
|
||||
if self.fp8_attn_gemm:
|
||||
q = _quant_q_fp8(q, layer.q_scale_float)
|
||||
idx_q = _quant_q_fp8(idx_q, layer.idx_q_scale_float)
|
||||
|
||||
idx_o, o = minimax_sparse_decode(
|
||||
q,
|
||||
None,
|
||||
@@ -483,6 +541,12 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
||||
use_msa=self._use_msa_decode,
|
||||
msa_kv_indices=msa_kv_indices,
|
||||
msa_plan=msa_plan,
|
||||
q_scale=layer.q_scale_float,
|
||||
k_scale=layer.k_scale_float,
|
||||
v_scale=layer.v_scale_float,
|
||||
idx_q_scale=layer.idx_q_scale_float,
|
||||
idx_k_scale=layer.idx_k_scale_float,
|
||||
idx_v_scale=layer.idx_v_scale_float,
|
||||
)
|
||||
return (
|
||||
None if idx_o is None else idx_o.reshape(q.shape[0], -1).contiguous(),
|
||||
|
||||
@@ -67,6 +67,12 @@ def minimax_sparse_prefill(
|
||||
max_seqblock_q: Optional[int] = None,
|
||||
all_seqblock_q: Optional[int] = None,
|
||||
seqlens_cpu: Optional[List[int]] = None,
|
||||
q_scale: Optional[float] = None,
|
||||
k_scale: Optional[float] = None,
|
||||
v_scale: Optional[float] = None,
|
||||
idx_q_scale: Optional[float] = None,
|
||||
idx_k_scale: Optional[float] = None,
|
||||
idx_v_scale: Optional[float] = None,
|
||||
):
|
||||
"""Run MiniMax-M3 sparse prefill.
|
||||
|
||||
@@ -106,6 +112,9 @@ def minimax_sparse_prefill(
|
||||
cu_seqblocks_q=cu_seqblocks_q,
|
||||
max_seqblock_q=max_seqblock_q,
|
||||
all_seqblock_q=all_seqblock_q,
|
||||
q_scale=idx_q_scale,
|
||||
k_scale=idx_k_scale,
|
||||
v_scale=idx_v_scale,
|
||||
)
|
||||
# Step 2: Reduce topk idx if num_idx_heads > num_kv_heads
|
||||
num_idx_heads = idx_q.shape[1]
|
||||
@@ -134,6 +143,9 @@ def minimax_sparse_prefill(
|
||||
prefix_lens=prefix_lens,
|
||||
block_size_k=block_size_k,
|
||||
sm_scale=sm_scale,
|
||||
q_scale=q_scale,
|
||||
k_scale=k_scale,
|
||||
v_scale=v_scale,
|
||||
)
|
||||
except MSAUnavailableError as err:
|
||||
_warn_msa_fallback(err)
|
||||
@@ -154,6 +166,9 @@ def minimax_sparse_prefill(
|
||||
sm_scale=sm_scale,
|
||||
cu_seqblocks_q=cu_seqblocks_q,
|
||||
max_seqblock_q=max_seqblock_q,
|
||||
q_scale=q_scale,
|
||||
k_scale=k_scale,
|
||||
v_scale=v_scale,
|
||||
)
|
||||
else:
|
||||
o = flash_prefill_with_gqa_share_sparse(
|
||||
@@ -173,6 +188,9 @@ def minimax_sparse_prefill(
|
||||
sm_scale=sm_scale,
|
||||
cu_seqblocks_q=cu_seqblocks_q,
|
||||
max_seqblock_q=max_seqblock_q,
|
||||
q_scale=q_scale,
|
||||
k_scale=k_scale,
|
||||
v_scale=v_scale,
|
||||
)
|
||||
return idx_o, o
|
||||
|
||||
@@ -208,6 +226,12 @@ def minimax_sparse_decode(
|
||||
torch.Tensor
|
||||
] = None, # per-forward MSA page table (cached)
|
||||
msa_plan=None, # per-forward MSA fmha_sm100 plan (cached)
|
||||
q_scale: Optional[float] = None,
|
||||
k_scale: Optional[float] = None,
|
||||
v_scale: Optional[float] = None,
|
||||
idx_q_scale: Optional[float] = None,
|
||||
idx_k_scale: Optional[float] = None,
|
||||
idx_v_scale: Optional[float] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# Step 1: Flash decode with topk index (using index head). When the dense main
|
||||
# attention is used, the indexer emits the page table directly (fused
|
||||
@@ -230,6 +254,9 @@ def minimax_sparse_decode(
|
||||
disable_index_value=disable_index_value,
|
||||
use_dense_main_attn=dense_main_attn_fn is not None,
|
||||
page_size=page_size,
|
||||
q_scale=idx_q_scale,
|
||||
k_scale=idx_k_scale,
|
||||
v_scale=idx_v_scale,
|
||||
)
|
||||
num_idx_heads = idx_q.shape[1]
|
||||
num_kv_heads = k_cache.shape[1]
|
||||
@@ -262,6 +289,9 @@ def minimax_sparse_decode(
|
||||
sm_scale=sm_scale,
|
||||
kv_indices=msa_kv_indices,
|
||||
plan=msa_plan,
|
||||
q_scale=q_scale,
|
||||
k_scale=k_scale,
|
||||
v_scale=v_scale,
|
||||
)
|
||||
except MSAUnavailableError as err:
|
||||
_warn_msa_fallback(err)
|
||||
@@ -276,6 +306,9 @@ def minimax_sparse_decode(
|
||||
block_size=block_size_k,
|
||||
topk_idx=topk_idx,
|
||||
sm_scale=sm_scale,
|
||||
q_scale=q_scale,
|
||||
k_scale=k_scale,
|
||||
v_scale=v_scale,
|
||||
)
|
||||
else:
|
||||
o = flash_decode_with_gqa_share_sparse(
|
||||
@@ -289,5 +322,8 @@ def minimax_sparse_decode(
|
||||
block_size=block_size_k,
|
||||
topk_idx=topk_idx,
|
||||
sm_scale=sm_scale,
|
||||
q_scale=q_scale,
|
||||
k_scale=k_scale,
|
||||
v_scale=v_scale,
|
||||
)
|
||||
return idx_o, o
|
||||
|
||||
@@ -3,6 +3,14 @@
|
||||
# Replaces only step 3 of MiniMax sparse prefill/decode. The lightning indexer
|
||||
# (steps 1-2) is unchanged and still produces `topk_idx`.
|
||||
# NVIDIA Blackwell (SM100/sm_103) only; callers gate on `msa_available()`.
|
||||
#
|
||||
# Dtypes: bf16 end-to-end, or uniform fp8_e4m3fn Q/K/V under fp8 attn-GEMM mode
|
||||
# (output bf16). fmha_sm100 selects its kernel variant from q.dtype alone and
|
||||
# casts k/v pointers to the same element type, so mixed bf16-q/fp8-KV is NOT
|
||||
# possible on the cutlass path and e5m2 would silently dispatch the e4m3
|
||||
# kernel — `_check_msa_dtypes` enforces uniformity here. The fp8 kernel
|
||||
# quantizes the unnormalized softmax P to e4m3 before the PV MMA (same
|
||||
# contract as the Triton fp8 path).
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -11,11 +19,32 @@ from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.attention.minimax_sparse.common.utils import unit_scale
|
||||
|
||||
|
||||
class MSAUnavailableError(RuntimeError):
|
||||
"""Raised when fmha_sm100 cannot serve the MiniMax MSA path."""
|
||||
|
||||
|
||||
def _check_msa_dtypes(q: torch.Tensor, k_cache: torch.Tensor, v_cache: torch.Tensor):
|
||||
# Uniform dtype required in BOTH modes: fmha_sm100 keys its kernel variant
|
||||
# on q.dtype alone and casts the k/v pointers to the same element type, so
|
||||
# a mismatched cache would be silently reinterpreted.
|
||||
if q.dtype == torch.bfloat16:
|
||||
assert (
|
||||
k_cache.dtype == torch.bfloat16
|
||||
), f"MSA bf16 requires a bf16 K cache, got {k_cache.dtype}"
|
||||
elif q.dtype == torch.float8_e4m3fn:
|
||||
# e5m2 is rejected here too: fmha_sm100's variant lookup falls back to
|
||||
# the e4m3 kernel for unknown dtype codes.
|
||||
assert (
|
||||
k_cache.dtype == torch.float8_e4m3fn
|
||||
), f"MSA fp8 requires an fp8_e4m3fn K cache, got {k_cache.dtype}"
|
||||
else:
|
||||
raise AssertionError(f"MSA supports bf16 or fp8_e4m3fn Q, got {q.dtype}")
|
||||
assert v_cache.dtype == k_cache.dtype
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def _load_fmha_sm100():
|
||||
try:
|
||||
@@ -101,12 +130,23 @@ def msa_sparse_prefill_main(
|
||||
prefix_lens: torch.Tensor, # [batch]
|
||||
block_size_k: int, # == page_size == 128 for M3
|
||||
sm_scale: Optional[float] = None,
|
||||
q_scale: Optional[float] = None,
|
||||
k_scale: Optional[float] = None,
|
||||
v_scale: Optional[float] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Drop-in for flash_prefill_with_gqa_share_sparse using MSA fmha_sm100.
|
||||
|
||||
Returns o [total_q, num_q_heads, head_dim].
|
||||
Returns o [total_q, num_q_heads, head_dim] (bf16 for fp8 inputs).
|
||||
|
||||
Scale semantics (per-tensor, None = unit): attention runs on Q*q_scale,
|
||||
K*k_scale, V*v_scale. The long-q cute path honors only sm_scale, so
|
||||
q_scale*k_scale is folded into sm_scale (exact for softmax) and v_scale is
|
||||
applied on the output; the short-q cutlass path gets the same folded values.
|
||||
"""
|
||||
fmha_sm100, _ = _load_fmha_sm100()
|
||||
_check_msa_dtypes(q, k_cache, v_cache)
|
||||
is_fp8 = q.dtype == torch.float8_e4m3fn
|
||||
v_scale = unit_scale(v_scale)
|
||||
|
||||
max_slots, num_kv_heads, head_dim = k_cache.shape
|
||||
num_q_heads = q.shape[1]
|
||||
@@ -116,6 +156,7 @@ def msa_sparse_prefill_main(
|
||||
raise ValueError(f"max_slots={max_slots} not divisible by page_size={P}")
|
||||
if sm_scale is None:
|
||||
sm_scale = head_dim**-0.5
|
||||
sm_scale = sm_scale * unit_scale(q_scale) * unit_scale(k_scale)
|
||||
|
||||
# Whole pool as MSA paged KV: [num_phys_pages, num_kv_heads, P, head_dim].
|
||||
n_phys_pages = max_slots // P
|
||||
@@ -138,6 +179,7 @@ def msa_sparse_prefill_main(
|
||||
kv_block_num=topk,
|
||||
causal=True,
|
||||
qo_offset=prefix_lens.to(torch.int32),
|
||||
use_fp8_kvcache=is_fp8,
|
||||
)
|
||||
o, _ = fmha_sm100(
|
||||
q,
|
||||
@@ -148,6 +190,10 @@ def msa_sparse_prefill_main(
|
||||
kv_indices=kv_indices,
|
||||
kv_block_indexes=kv_block_indexes,
|
||||
)
|
||||
# The cute (long-q) sparse prefill backend honors sm_scale only; apply the
|
||||
# V dequant scale on the output (exact: softmax normalization excludes V).
|
||||
if v_scale != 1.0:
|
||||
o = o * v_scale
|
||||
return o
|
||||
|
||||
|
||||
@@ -159,6 +205,7 @@ def build_msa_decode_meta(
|
||||
num_q_heads: int,
|
||||
block_size_k: int,
|
||||
topk: int,
|
||||
is_fp8: bool = False,
|
||||
):
|
||||
"""Per-forward MSA decode metadata (page table + fmha plan), shared across layers.
|
||||
|
||||
@@ -187,25 +234,28 @@ def build_msa_decode_meta(
|
||||
kv_block_num=topk,
|
||||
causal=False,
|
||||
qo_offset=seq_lens_i32 - 1, # decode query sits at the last cached position
|
||||
use_fp8_kvcache=is_fp8,
|
||||
)
|
||||
return kv_indices, plan
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Eager-only MSA decode plan (NOT used under CUDA graph)
|
||||
# MSA decode plan (persistent per batch size; used by eager decode AND under
|
||||
# CUDA graph)
|
||||
#
|
||||
# WARNING: the fmha_sm100 sparse decode kernel is NOT cuda-graph-safe — captured
|
||||
# and replayed it returns silently wrong results that compound across replays
|
||||
# (~14% GSM8K loss on B200). The backend routes decode to the cuda-graph-safe
|
||||
# Triton sparse path whenever decode runs under a CUDA graph (see
|
||||
# MiniMaxSparseAttnBackend._use_msa_decode); this plan is reachable ONLY in eager
|
||||
# decode (no decode CUDA graph), where there is no capture/replay. Do NOT wire it
|
||||
# back into a captured graph — that reintroduces the ~14% regression.
|
||||
# History: MSA decode under CUDA graph was disabled after silently wrong
|
||||
# results (~14% GSM8K on B200). Root cause (2026-07): the topk producers
|
||||
# emitted block ids in score order, violating fmha_sm100's strictly-ascending
|
||||
# kv_block_indexes contract — its sorted-order early-exit then mis-masked the
|
||||
# partial last block for any row with seq_len > topk*block_size. The producers
|
||||
# now sort ascending (minimax_decode_topk.cuh, _topk_index_merge_kernel,
|
||||
# prefill _topk_index_kernel), and capture/replay of the full pipeline is
|
||||
# bit-exact vs eager (see tests/repro_msa_decode_degenerate.py).
|
||||
#
|
||||
# The build-once / replay-update structure below (refreshing the four length
|
||||
# tensors ``{kv_segment_lens, kv_segment_offsets, kv_page_indptr, qo_offset}`` and
|
||||
# the page table in place) is a leftover from the abandoned capture-once attempt;
|
||||
# it is kept only because eager decode reuses one per-forward plan across layers.
|
||||
# The build-once / update-in-place structure below refreshes the four length
|
||||
# tensors ``{kv_segment_lens, kv_segment_offsets, kv_page_indptr, qo_offset}``
|
||||
# and the page table each forward; the captured graph reads the same tensor
|
||||
# addresses on replay.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_MSA_CG_LEN_KEYS = (
|
||||
@@ -242,6 +292,7 @@ def build_msa_decode_cg_plan(
|
||||
topk: int,
|
||||
batch_size: int,
|
||||
device: Optional[torch.device] = None,
|
||||
is_fp8: bool = False,
|
||||
):
|
||||
"""Persistent fmha_sm100 decode plan for one batch size (CUDA-graph stable).
|
||||
|
||||
@@ -264,6 +315,7 @@ def build_msa_decode_cg_plan(
|
||||
causal=False,
|
||||
qo_offset=kv - 1,
|
||||
device=device,
|
||||
use_fp8_kvcache=is_fp8,
|
||||
)
|
||||
_check_cg_plan_layout(plan)
|
||||
return plan
|
||||
@@ -280,37 +332,53 @@ def update_msa_decode_cg_meta(
|
||||
num_q_heads: int,
|
||||
num_kv_heads: int,
|
||||
):
|
||||
"""Refresh the persistent decode plan's length-dependent tensors + page table IN PLACE.
|
||||
"""Refresh the persistent decode plan's length-dependent tensors + page table
|
||||
IN PLACE, entirely with device-side ops (no host<->device sync).
|
||||
|
||||
Host-side (calls fmha_sm100_plan and one ``.item()``); MUST run outside CUDA-graph
|
||||
capture — i.e. only from ``init_forward_metadata_out_graph``. The captured graph then
|
||||
reads the same plan-tensor and ``kv_indices_buf`` addresses on replay.
|
||||
Runs once per decode forward from ``init_forward_metadata_out_graph``; the
|
||||
captured graph then reads the same plan-tensor and ``kv_indices_buf``
|
||||
addresses on replay. A device sync here stalls the overlap scheduler, so the
|
||||
previous implementation — a throwaway ``fmha_sm100_plan`` build per step
|
||||
(``.tolist()``/``.item()`` syncs + a plan-kernel launch) just to copy four
|
||||
length tensors — is replaced by computing their contents directly, matching
|
||||
``_fmha_sm100_plan``'s sparse-decode (qo_len==1, causal=False) layout:
|
||||
|
||||
kv_segment_lens = seq_lens
|
||||
kv_segment_offsets = [0, cumsum(seq_lens)]
|
||||
kv_page_indptr = [0, cumsum(ceil(seq_lens / P))]
|
||||
qo_offset = broadcast max(seq_lens) (causal=False planner quirk)
|
||||
|
||||
The worklist tensors stay untouched: the plan schedules from the constant
|
||||
``topk * P`` per request, never the real lengths (see build_msa_decode_cg_plan).
|
||||
"""
|
||||
P = block_size_k
|
||||
B = seq_lens.shape[0]
|
||||
if B == 0: # idle batch: serving guards this, but keep the helper total
|
||||
return
|
||||
pd = plan[3]
|
||||
seq_lens_i32 = seq_lens.to(torch.int32)
|
||||
# Fresh plan for the real lengths; copy only its four length-dependent tensors into the
|
||||
# persistent plan (same shapes — they depend on batch size, not length). The fresh
|
||||
# worklist is identical to the persistent one (topk*P based) and is discarded.
|
||||
# qo_offset is clamped: graph replay pads the batch with seq_len==0 slots
|
||||
# (masked via kv_segment_lens==0, but seq_len-1 would be -1).
|
||||
fresh = _run_fmha_sm100_plan(
|
||||
torch.ones(B, dtype=torch.int32),
|
||||
seq_lens_i32,
|
||||
num_q_heads,
|
||||
num_kv_heads=num_kv_heads,
|
||||
page_size=P,
|
||||
kv_block_num=topk,
|
||||
causal=False,
|
||||
qo_offset=(seq_lens_i32 - 1).clamp_min(0),
|
||||
device=seq_lens.device,
|
||||
)
|
||||
_check_cg_plan_layout(fresh)
|
||||
pd, fd = plan[3], fresh[3]
|
||||
for k in _MSA_CG_LEN_KEYS:
|
||||
pd[k].copy_(fd[k])
|
||||
table = _build_page_table(req_to_token, slot_ids, seq_lens, P)
|
||||
kv_indices_buf[: table.numel()].copy_(table)
|
||||
pd["kv_segment_lens"].copy_(seq_lens_i32)
|
||||
kv_off = pd["kv_segment_offsets"]
|
||||
kv_off[0].zero_()
|
||||
torch.cumsum(seq_lens_i32, 0, out=kv_off[1:])
|
||||
n_pages = torch.div(seq_lens_i32 + (P - 1), P, rounding_mode="floor")
|
||||
indptr = pd["kv_page_indptr"]
|
||||
indptr[0].zero_()
|
||||
torch.cumsum(n_pages, 0, out=indptr[1:])
|
||||
pd["qo_offset"].copy_(seq_lens_i32.max().expand(B))
|
||||
|
||||
# Page table, sync-free: fill the WHOLE persistent buffer (fixed size, no
|
||||
# host-side total-page count). Packed slot -> (request, logical page) via
|
||||
# searchsorted; slots beyond the live page count land on clamped reads and
|
||||
# are never dereferenced (the kernel walks kv_page_indptr ranges only).
|
||||
n = kv_indices_buf.numel()
|
||||
idx = torch.arange(n, device=kv_indices_buf.device)
|
||||
ends = torch.cumsum(n_pages.to(torch.int64), 0)
|
||||
req = torch.searchsorted(ends, idx, right=True).clamp_max_(B - 1)
|
||||
starts = ends - n_pages
|
||||
logical_first = ((idx - starts[req]) * P).clamp_(0, req_to_token.shape[1] - 1)
|
||||
rows = slot_ids[req].to(torch.int64)
|
||||
kv_indices_buf.copy_((req_to_token[rows, logical_first] // P).to(torch.int32))
|
||||
|
||||
|
||||
def msa_sparse_decode_main(
|
||||
@@ -327,18 +395,27 @@ def msa_sparse_decode_main(
|
||||
torch.Tensor
|
||||
] = None, # precomputed page table (per-forward cache)
|
||||
plan=None, # precomputed fmha_sm100 plan (per-forward cache)
|
||||
q_scale: Optional[float] = None,
|
||||
k_scale: Optional[float] = None,
|
||||
v_scale: Optional[float] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Drop-in for flash_decode_with_gqa_share_sparse using MSA fmha_sm100.
|
||||
|
||||
Each request is one decode query at absolute position seq_len-1 attending to its
|
||||
cached KV through the topk selected 128-blocks. Returns o [batch, num_q_heads, head_dim].
|
||||
cached KV through the topk selected 128-blocks. Returns o [batch, num_q_heads,
|
||||
head_dim] (bf16 for fp8 inputs).
|
||||
|
||||
``kv_indices`` / ``plan`` are shared across all sparse layers of a forward; the serving
|
||||
backend builds them once via ``build_msa_decode_cg_plan`` + ``update_msa_decode_cg_meta``
|
||||
(eager decode only) and passes them in. When omitted (only the standalone parity
|
||||
harnesses) they are built here via ``build_msa_decode_meta``.
|
||||
|
||||
Scales (None = unit) are passed natively: the cutlass decode path folds
|
||||
q_scale*k_scale into the softmax scale and applies v_scale on the output
|
||||
in-kernel.
|
||||
"""
|
||||
fmha_sm100, _ = _load_fmha_sm100()
|
||||
_check_msa_dtypes(q, k_cache, v_cache)
|
||||
|
||||
max_slots, num_kv_heads, head_dim = k_cache.shape
|
||||
H = q.shape[1]
|
||||
@@ -355,7 +432,14 @@ def msa_sparse_decode_main(
|
||||
|
||||
if kv_indices is None or plan is None:
|
||||
kv_indices, plan = build_msa_decode_meta(
|
||||
k_cache, req_to_token, slot_ids, seq_lens, H, P, topk
|
||||
k_cache,
|
||||
req_to_token,
|
||||
slot_ids,
|
||||
seq_lens,
|
||||
H,
|
||||
P,
|
||||
topk,
|
||||
is_fp8=q.dtype == torch.float8_e4m3fn,
|
||||
)
|
||||
kv_block_indexes = topk_idx.permute(1, 0, 2).contiguous().to(torch.int32)
|
||||
|
||||
@@ -365,6 +449,9 @@ def msa_sparse_decode_main(
|
||||
v_paged,
|
||||
plan,
|
||||
sm_scale=sm_scale,
|
||||
q_scale=unit_scale(q_scale),
|
||||
k_scale=unit_scale(k_scale),
|
||||
v_scale=unit_scale(v_scale),
|
||||
kv_indices=kv_indices,
|
||||
kv_block_indexes=kv_block_indexes,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,455 @@
|
||||
"""Unit tests for fp8 (fp8 attn-GEMM mode) support in the M3 sparse Triton kernels.
|
||||
|
||||
Strategy: quantize random bf16 tensors to fp8_e4m3fn, then compare the fp8
|
||||
kernel run against the SAME kernel run in bf16 on the *dequantized* tensors.
|
||||
Both runs see numerically identical Q/K values, so the QK GEMMs match closely
|
||||
and top-k selection is stable; the only intended divergence is the fp8 PV MMA
|
||||
(P quantized to e4m3), which the tolerances cover. This isolates kernel
|
||||
arithmetic from quantization error.
|
||||
|
||||
Covers: step-3 decode/prefill (gqa-share sparse), step-1 decode/prefill
|
||||
indexer, non-unit k_scale/v_scale semantics, and bf16-path regression.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.attention.minimax_sparse_ops.decode.flash_with_topk_idx import (
|
||||
flash_decode_with_topk_idx,
|
||||
)
|
||||
from sglang.srt.layers.attention.minimax_sparse_ops.decode.topk_sparse import (
|
||||
flash_decode_with_gqa_share_sparse,
|
||||
)
|
||||
from sglang.srt.layers.attention.minimax_sparse_ops.prefill.flash_with_topk_idx import (
|
||||
flash_prefill_with_topk_index,
|
||||
)
|
||||
from sglang.srt.layers.attention.minimax_sparse_ops.prefill.topk_sparse import (
|
||||
flash_prefill_with_gqa_share_sparse,
|
||||
)
|
||||
|
||||
DEVICE = "cuda"
|
||||
FP8 = torch.float8_e4m3fn
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not torch.cuda.is_available(), reason="requires CUDA (Triton fp8 kernels)"
|
||||
)
|
||||
# fp8 PV (P quantized to e4m3, ~3-bit mantissa on [0,1] weights) dominates the
|
||||
# fp8-vs-dequantized-ref error; QK matches to fp32-accumulation noise.
|
||||
FP8_ATOL = 6e-2
|
||||
FP8_RTOL = 6e-2
|
||||
# widening mode (bf16 Q, fp8 KV) computes on exactly the dequantized values.
|
||||
WIDEN_ATOL = 1e-3
|
||||
WIDEN_RTOL = 1e-3
|
||||
|
||||
|
||||
def qdq(x: torch.Tensor):
|
||||
"""Quantize to e4m3 and return (fp8, dequantized-bf16) views of it."""
|
||||
x8 = x.to(FP8)
|
||||
return x8, x8.to(torch.bfloat16)
|
||||
|
||||
|
||||
def build_decode_inputs(
|
||||
batch_size=4,
|
||||
num_q_heads=8,
|
||||
num_kv_heads=1,
|
||||
head_dim=128,
|
||||
block_size=128,
|
||||
topk=8,
|
||||
seq_lens_list=(513, 1023, 257, 769),
|
||||
):
|
||||
seq_lens_list = list(seq_lens_list)[:batch_size]
|
||||
max_kv_len = max(seq_lens_list)
|
||||
max_slots = batch_size * max_kv_len
|
||||
q = torch.randn(
|
||||
batch_size, num_q_heads, head_dim, dtype=torch.bfloat16, device=DEVICE
|
||||
)
|
||||
k = torch.randn(
|
||||
max_slots, num_kv_heads, head_dim, dtype=torch.bfloat16, device=DEVICE
|
||||
)
|
||||
v = torch.randn(
|
||||
max_slots, num_kv_heads, head_dim, dtype=torch.bfloat16, device=DEVICE
|
||||
)
|
||||
req_to_token = torch.zeros(batch_size, max_kv_len, dtype=torch.int32, device=DEVICE)
|
||||
slot_ids = torch.arange(batch_size, dtype=torch.int64, device=DEVICE)
|
||||
seq_lens = torch.tensor(seq_lens_list, dtype=torch.int32, device=DEVICE)
|
||||
for i in range(batch_size):
|
||||
base = i * max_kv_len
|
||||
req_to_token[i] = (torch.randperm(max_kv_len, device=DEVICE) + base).to(
|
||||
torch.int32
|
||||
)
|
||||
topk_idx = torch.full(
|
||||
(num_kv_heads, batch_size, topk), -1, dtype=torch.int32, device=DEVICE
|
||||
)
|
||||
for kh in range(num_kv_heads):
|
||||
for b in range(batch_size):
|
||||
nb = (seq_lens_list[b] + block_size - 1) // block_size
|
||||
ak = min(topk, nb)
|
||||
# sorted ascending, matching the production topk contract
|
||||
sel = torch.randperm(nb, device=DEVICE)[:ak].sort().values
|
||||
topk_idx[kh, b, :ak] = sel.to(torch.int32)
|
||||
return q, k, v, req_to_token, seq_lens, slot_ids, topk_idx
|
||||
|
||||
|
||||
def build_prefill_inputs(
|
||||
batch_size=2,
|
||||
num_q_heads=8,
|
||||
num_kv_heads=1,
|
||||
head_dim=128,
|
||||
seq_lens_list=(513, 769),
|
||||
prefix_lens_list=(0, 257),
|
||||
):
|
||||
seq_lens_list = list(seq_lens_list)[:batch_size]
|
||||
prefix_lens_list = list(prefix_lens_list)[:batch_size]
|
||||
q_lens = [s - p for s, p in zip(seq_lens_list, prefix_lens_list)]
|
||||
total_q = sum(q_lens)
|
||||
max_kv_len = max(seq_lens_list)
|
||||
max_slots = batch_size * max_kv_len
|
||||
q = torch.randn(total_q, num_q_heads, head_dim, dtype=torch.bfloat16, device=DEVICE)
|
||||
k = torch.randn(
|
||||
max_slots, num_kv_heads, head_dim, dtype=torch.bfloat16, device=DEVICE
|
||||
)
|
||||
v = torch.randn(
|
||||
max_slots, num_kv_heads, head_dim, dtype=torch.bfloat16, device=DEVICE
|
||||
)
|
||||
req_to_token = torch.zeros(batch_size, max_kv_len, dtype=torch.int32, device=DEVICE)
|
||||
slot_ids = torch.arange(batch_size, dtype=torch.int64, device=DEVICE)
|
||||
for i in range(batch_size):
|
||||
base = i * max_kv_len
|
||||
req_to_token[i] = (torch.randperm(max_kv_len, device=DEVICE) + base).to(
|
||||
torch.int32
|
||||
)
|
||||
cu_seqlens = torch.zeros(batch_size + 1, dtype=torch.int32, device=DEVICE)
|
||||
cu_seqlens[1:] = torch.tensor(q_lens, device=DEVICE).cumsum(0)
|
||||
seq_lens = torch.tensor(seq_lens_list, dtype=torch.int32, device=DEVICE)
|
||||
prefix_lens = torch.tensor(prefix_lens_list, dtype=torch.int32, device=DEVICE)
|
||||
return (
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
req_to_token,
|
||||
slot_ids,
|
||||
cu_seqlens,
|
||||
seq_lens,
|
||||
prefix_lens,
|
||||
max(q_lens),
|
||||
max(seq_lens_list),
|
||||
)
|
||||
|
||||
|
||||
def run_step3_decode(q, k, v, req_to_token, seq_lens, slot_ids, topk_idx, **kw):
|
||||
return flash_decode_with_gqa_share_sparse(
|
||||
q, None, k, v, req_to_token, seq_lens, slot_ids, 128, topk_idx, **kw
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# step-3 decode (gqa-share sparse)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_step3_decode_all_fp8_vs_dequant_ref():
|
||||
torch.manual_seed(0)
|
||||
q, k, v, r2t, seq_lens, sids, tidx = build_decode_inputs()
|
||||
q8, qr = qdq(q)
|
||||
k8, kr = qdq(k)
|
||||
v8, vr = qdq(v)
|
||||
out8 = run_step3_decode(q8, k8, v8, r2t, seq_lens, sids, tidx)
|
||||
ref = run_step3_decode(qr, kr, vr, r2t, seq_lens, sids, tidx)
|
||||
assert out8.dtype == torch.bfloat16
|
||||
torch.testing.assert_close(out8.float(), ref.float(), atol=FP8_ATOL, rtol=FP8_RTOL)
|
||||
|
||||
|
||||
def test_step3_decode_widening_mode():
|
||||
torch.manual_seed(1)
|
||||
q, k, v, r2t, seq_lens, sids, tidx = build_decode_inputs()
|
||||
k8, kr = qdq(k)
|
||||
v8, vr = qdq(v)
|
||||
out = run_step3_decode(q, k8, v8, r2t, seq_lens, sids, tidx)
|
||||
ref = run_step3_decode(q, kr, vr, r2t, seq_lens, sids, tidx)
|
||||
torch.testing.assert_close(
|
||||
out.float(), ref.float(), atol=WIDEN_ATOL, rtol=WIDEN_RTOL
|
||||
)
|
||||
|
||||
|
||||
def test_step3_decode_scales():
|
||||
torch.manual_seed(2)
|
||||
q, k, v, r2t, seq_lens, sids, tidx = build_decode_inputs()
|
||||
q8, qr = qdq(q)
|
||||
k8, kr = qdq(k)
|
||||
v8, vr = qdq(v)
|
||||
q_scale, k_scale, v_scale = 1.5, 0.5, 2.0
|
||||
out8 = run_step3_decode(
|
||||
q8,
|
||||
k8,
|
||||
v8,
|
||||
r2t,
|
||||
seq_lens,
|
||||
sids,
|
||||
tidx,
|
||||
q_scale=q_scale,
|
||||
k_scale=k_scale,
|
||||
v_scale=v_scale,
|
||||
)
|
||||
# reference: bf16 kernel on pre-scaled dequantized Q/K/V (scale semantics:
|
||||
# the tensor stores value/scale; attention runs on value = stored * scale)
|
||||
ref = run_step3_decode(
|
||||
(qr * q_scale).to(torch.bfloat16),
|
||||
(kr * k_scale).to(torch.bfloat16),
|
||||
(vr * v_scale).to(torch.bfloat16),
|
||||
r2t,
|
||||
seq_lens,
|
||||
sids,
|
||||
tidx,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
out8.float(), ref.float(), atol=FP8_ATOL * v_scale, rtol=FP8_RTOL * v_scale
|
||||
)
|
||||
|
||||
|
||||
def test_step3_decode_bf16_regression():
|
||||
torch.manual_seed(3)
|
||||
q, k, v, r2t, seq_lens, sids, tidx = build_decode_inputs()
|
||||
out = run_step3_decode(q, k, v, r2t, seq_lens, sids, tidx)
|
||||
out_scaled = run_step3_decode(
|
||||
q, k, v, r2t, seq_lens, sids, tidx, q_scale=1.0, k_scale=1.0, v_scale=1.0
|
||||
)
|
||||
assert out.dtype == q.dtype
|
||||
torch.testing.assert_close(out, out_scaled, atol=0, rtol=0)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# step-3 prefill (gqa-share sparse)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def run_step3_prefill(q, k, v, r2t, sids, tidx, cu, seq_lens, prefix, max_q, **kw):
|
||||
return flash_prefill_with_gqa_share_sparse(
|
||||
q=q,
|
||||
k_cache=k,
|
||||
v_cache=v,
|
||||
sink=None,
|
||||
req_to_token=r2t,
|
||||
slot_ids=sids,
|
||||
topk_idx=tidx,
|
||||
block_size_q=1,
|
||||
block_size_k=128,
|
||||
cu_seqlens=cu,
|
||||
seq_lens=seq_lens,
|
||||
prefix_lens=prefix,
|
||||
max_seqlen_q=max_q,
|
||||
**kw,
|
||||
)
|
||||
|
||||
|
||||
def _prefill_topk_idx(cu_seqlens, seq_lens, prefix_lens, num_kv_heads, topk, block):
|
||||
# per-token (block_size_q=1) causal topk: for query at absolute position p,
|
||||
# any blocks with start <= p, sorted ascending, -1 padded.
|
||||
total_q = int(cu_seqlens[-1].item())
|
||||
tidx = torch.full(
|
||||
(num_kv_heads, total_q, topk), -1, dtype=torch.int32, device=DEVICE
|
||||
)
|
||||
row = 0
|
||||
for b in range(len(seq_lens)):
|
||||
q_len = int(cu_seqlens[b + 1] - cu_seqlens[b])
|
||||
prefix = int(prefix_lens[b])
|
||||
for j in range(q_len):
|
||||
nb = (prefix + j) // block + 1 # blocks fully/partially before pos
|
||||
ak = min(topk, nb)
|
||||
sel = torch.randperm(nb, device=DEVICE)[:ak].sort().values
|
||||
for kh in range(num_kv_heads):
|
||||
tidx[kh, row, :ak] = sel.to(torch.int32)
|
||||
row += 1
|
||||
return tidx
|
||||
|
||||
|
||||
def test_step3_prefill_all_fp8_vs_dequant_ref():
|
||||
torch.manual_seed(4)
|
||||
q, k, v, r2t, sids, cu, seq_lens, prefix, max_q, _ = build_prefill_inputs()
|
||||
tidx = _prefill_topk_idx(cu.cpu(), seq_lens.cpu(), prefix.cpu(), 1, 8, 128)
|
||||
q8, qr = qdq(q)
|
||||
k8, kr = qdq(k)
|
||||
v8, vr = qdq(v)
|
||||
out8 = run_step3_prefill(q8, k8, v8, r2t, sids, tidx, cu, seq_lens, prefix, max_q)
|
||||
ref = run_step3_prefill(qr, kr, vr, r2t, sids, tidx, cu, seq_lens, prefix, max_q)
|
||||
assert out8.dtype == torch.bfloat16
|
||||
torch.testing.assert_close(out8.float(), ref.float(), atol=FP8_ATOL, rtol=FP8_RTOL)
|
||||
|
||||
|
||||
def test_step3_prefill_scales():
|
||||
torch.manual_seed(5)
|
||||
q, k, v, r2t, sids, cu, seq_lens, prefix, max_q, _ = build_prefill_inputs()
|
||||
tidx = _prefill_topk_idx(cu.cpu(), seq_lens.cpu(), prefix.cpu(), 1, 8, 128)
|
||||
q8, qr = qdq(q)
|
||||
k8, kr = qdq(k)
|
||||
v8, vr = qdq(v)
|
||||
q_scale, k_scale, v_scale = 1.5, 0.5, 2.0
|
||||
out8 = run_step3_prefill(
|
||||
q8,
|
||||
k8,
|
||||
v8,
|
||||
r2t,
|
||||
sids,
|
||||
tidx,
|
||||
cu,
|
||||
seq_lens,
|
||||
prefix,
|
||||
max_q,
|
||||
q_scale=q_scale,
|
||||
k_scale=k_scale,
|
||||
v_scale=v_scale,
|
||||
)
|
||||
ref = run_step3_prefill(
|
||||
(qr * q_scale).to(torch.bfloat16),
|
||||
(kr * k_scale).to(torch.bfloat16),
|
||||
(vr * v_scale).to(torch.bfloat16),
|
||||
r2t,
|
||||
sids,
|
||||
tidx,
|
||||
cu,
|
||||
seq_lens,
|
||||
prefix,
|
||||
max_q,
|
||||
)
|
||||
# v_scale multiplies the fp8 kernel's quantized-P PV output, amplifying
|
||||
# the P-quantization error by the same factor vs the pre-scaled bf16 ref.
|
||||
torch.testing.assert_close(
|
||||
out8.float(), ref.float(), atol=FP8_ATOL * v_scale, rtol=FP8_RTOL * v_scale
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# step-1 decode indexer
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def run_indexer_decode(q, k, v, r2t, seq_lens, sids, **kw):
|
||||
return flash_decode_with_topk_idx(
|
||||
q=q,
|
||||
sink=None,
|
||||
k_cache=k,
|
||||
v_cache=v,
|
||||
req_to_token=r2t,
|
||||
seq_lens=seq_lens,
|
||||
slot_ids=sids,
|
||||
max_seqlen=int(seq_lens.max().item()),
|
||||
block_size=128,
|
||||
topk=4,
|
||||
init_blocks=1,
|
||||
local_blocks=1,
|
||||
**kw,
|
||||
)
|
||||
|
||||
|
||||
def _topk_overlap(a: torch.Tensor, b: torch.Tensor) -> float:
|
||||
"""Mean per-row overlap of the selected (non-negative) block-id sets."""
|
||||
total, hit = 0, 0
|
||||
af, bf = a.reshape(-1, a.shape[-1]), b.reshape(-1, b.shape[-1])
|
||||
for i in range(af.shape[0]):
|
||||
sa = set(af[i][af[i] >= 0].tolist())
|
||||
sb = set(bf[i][bf[i] >= 0].tolist())
|
||||
if not sa and not sb:
|
||||
continue
|
||||
total += max(len(sa), len(sb))
|
||||
hit += len(sa & sb)
|
||||
return hit / max(total, 1)
|
||||
|
||||
|
||||
def test_indexer_decode_all_fp8_vs_dequant_ref():
|
||||
torch.manual_seed(6)
|
||||
q, k, v, r2t, seq_lens, sids, _ = build_decode_inputs(
|
||||
num_q_heads=1, num_kv_heads=1, seq_lens_list=(1023, 769, 513, 257)
|
||||
)
|
||||
q8, qr = qdq(q)
|
||||
k8, kr = qdq(k)
|
||||
v8, vr = qdq(v)
|
||||
o8, tidx8, _ = run_indexer_decode(q8, k8, v8, r2t, seq_lens, sids)
|
||||
oref, tidxref, _ = run_indexer_decode(qr, kr, vr, r2t, seq_lens, sids)
|
||||
assert o8.dtype == torch.bfloat16
|
||||
# o is the full (non-sparse) indexer attention output — independent of the
|
||||
# topk side-channel — so it must track the dequantized reference.
|
||||
torch.testing.assert_close(o8.float(), oref.float(), atol=FP8_ATOL, rtol=FP8_RTOL)
|
||||
# topk selection runs on QK scores that only differ by fp32-accumulation
|
||||
# noise; require near-total agreement (ties may flip an occasional block).
|
||||
assert _topk_overlap(tidx8, tidxref) >= 0.9
|
||||
|
||||
|
||||
def test_indexer_decode_score_only_fp8():
|
||||
torch.manual_seed(7)
|
||||
q, k, _, r2t, seq_lens, sids, _ = build_decode_inputs(
|
||||
batch_size=2, num_q_heads=1, num_kv_heads=1, seq_lens_list=(1023, 769)
|
||||
)
|
||||
q8, qr = qdq(q)
|
||||
k8, kr = qdq(k)
|
||||
o8, tidx8, _ = run_indexer_decode(
|
||||
q8, k8, None, r2t, seq_lens, sids, disable_index_value=True
|
||||
)
|
||||
oref, tidxref, _ = run_indexer_decode(
|
||||
qr, kr, None, r2t, seq_lens, sids, disable_index_value=True
|
||||
)
|
||||
assert o8 is None and oref is None
|
||||
assert _topk_overlap(tidx8, tidxref) >= 0.9
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# step-1 prefill indexer
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def run_indexer_prefill(q, k, v, r2t, sids, cu, seq_lens, prefix, max_q, max_k, **kw):
|
||||
return flash_prefill_with_topk_index(
|
||||
q=q,
|
||||
k_cache=k,
|
||||
v_cache=v,
|
||||
sink=None,
|
||||
req_to_token=r2t,
|
||||
slot_ids=sids,
|
||||
cu_seqlens=cu,
|
||||
seq_lens=seq_lens,
|
||||
prefix_lens=prefix,
|
||||
max_seqlen_q=max_q,
|
||||
max_seqlen_k=max_k,
|
||||
block_size_q=1,
|
||||
block_size_k=128,
|
||||
topk=4,
|
||||
init_blocks=1,
|
||||
local_blocks=1,
|
||||
**kw,
|
||||
)
|
||||
|
||||
|
||||
def test_indexer_prefill_all_fp8_vs_dequant_ref():
|
||||
torch.manual_seed(8)
|
||||
q, k, v, r2t, sids, cu, seq_lens, prefix, max_q, max_k = build_prefill_inputs(
|
||||
num_q_heads=1, num_kv_heads=1
|
||||
)
|
||||
q8, qr = qdq(q)
|
||||
k8, kr = qdq(k)
|
||||
v8, vr = qdq(v)
|
||||
o8, tidx8 = run_indexer_prefill(
|
||||
q8, k8, v8, r2t, sids, cu, seq_lens, prefix, max_q, max_k
|
||||
)
|
||||
oref, tidxref = run_indexer_prefill(
|
||||
qr, kr, vr, r2t, sids, cu, seq_lens, prefix, max_q, max_k
|
||||
)
|
||||
assert o8.dtype == torch.bfloat16
|
||||
torch.testing.assert_close(o8.float(), oref.float(), atol=FP8_ATOL, rtol=FP8_RTOL)
|
||||
assert _topk_overlap(tidx8, tidxref) >= 0.9
|
||||
|
||||
|
||||
def test_dtype_contract_rejects_e5m2_q():
|
||||
q = torch.randn(2, 8, 128, dtype=torch.bfloat16, device=DEVICE).to(
|
||||
torch.float8_e5m2
|
||||
)
|
||||
k = torch.randn(256, 1, 128, dtype=torch.bfloat16, device=DEVICE).to(FP8)
|
||||
from sglang.kernels.ops.attention.minimax_sparse.common.utils import (
|
||||
check_sparse_kv_fp8,
|
||||
)
|
||||
|
||||
with pytest.raises(AssertionError):
|
||||
check_sparse_kv_fp8(q, k, None, label="test")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import sys
|
||||
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
@@ -0,0 +1,286 @@
|
||||
"""Parity tests for MSA (fmha_sm100) all-fp8 sparse attention (fp8 attn-GEMM mode).
|
||||
|
||||
No upstream fp8 test exists for fmha_sm100's cutlass sparse-decode path, so
|
||||
this is the reference check: MSA fp8 vs the Triton fp8 sparse kernels on the
|
||||
same quantized tensors, plus fp8-vs-bf16 error bounds and CUDA-graph
|
||||
capture/replay bit-exactness of the fp8 decode.
|
||||
|
||||
Both fp8 paths quantize the unnormalized softmax P to e4m3 before the PV MMA,
|
||||
but their QK/accumulation orders differ, so MSA-fp8 vs Triton-fp8 tolerances
|
||||
cover two independent P-quantization errors (~1e-1 worst-case elementwise).
|
||||
|
||||
Requires SM100 + fmha_sm100 (first run JIT-compiles the fp8 variants).
|
||||
Run: pytest python/sglang/srt/layers/attention/minimax_sparse_ops/tests/test_msa_fp8_parity.py -v
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.attention.minimax_sparse_ops.decode.topk_sparse import (
|
||||
flash_decode_with_gqa_share_sparse,
|
||||
)
|
||||
from sglang.srt.layers.attention.minimax_sparse_ops.msa import (
|
||||
build_msa_decode_cg_plan,
|
||||
msa_available,
|
||||
msa_sparse_decode_main,
|
||||
msa_sparse_prefill_main,
|
||||
update_msa_decode_cg_meta,
|
||||
)
|
||||
from sglang.srt.layers.attention.minimax_sparse_ops.prefill.topk_sparse import (
|
||||
flash_prefill_with_gqa_share_sparse,
|
||||
)
|
||||
|
||||
DEVICE = "cuda"
|
||||
FP8 = torch.float8_e4m3fn
|
||||
P = 128 # sparse block == page size
|
||||
# two independent e4m3-P quantizations (MSA + Triton)
|
||||
X_ATOL = 1e-1
|
||||
X_RTOL = 1e-1
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not torch.cuda.is_available() or not msa_available(),
|
||||
reason="requires SM100 + fmha_sm100",
|
||||
)
|
||||
|
||||
|
||||
def qdq(x):
|
||||
x8 = x.to(FP8)
|
||||
return x8, x8.to(torch.bfloat16)
|
||||
|
||||
|
||||
def build_paged_inputs(seq_lens_list, num_q_heads=16, num_kv_heads=1, head_dim=128):
|
||||
"""Page-aligned pool: each logical 128-token page maps to one physical page
|
||||
(contiguous 128 slots), as MSA's page-table builder requires."""
|
||||
batch = len(seq_lens_list)
|
||||
pages_per_req = [(s + P - 1) // P for s in seq_lens_list]
|
||||
max_pages = max(pages_per_req)
|
||||
total_pages = batch * max_pages
|
||||
max_slots = total_pages * P
|
||||
page_perm = torch.randperm(total_pages, device=DEVICE)
|
||||
req_to_token = torch.zeros(batch, max_pages * P, dtype=torch.int32, device=DEVICE)
|
||||
for b in range(batch):
|
||||
for p in range(max_pages):
|
||||
phys = page_perm[b * max_pages + p]
|
||||
req_to_token[b, p * P : (p + 1) * P] = phys * P + torch.arange(
|
||||
P, device=DEVICE
|
||||
)
|
||||
k = torch.randn(
|
||||
max_slots, num_kv_heads, head_dim, dtype=torch.bfloat16, device=DEVICE
|
||||
)
|
||||
v = torch.randn(
|
||||
max_slots, num_kv_heads, head_dim, dtype=torch.bfloat16, device=DEVICE
|
||||
)
|
||||
slot_ids = torch.arange(batch, dtype=torch.int64, device=DEVICE)
|
||||
seq_lens = torch.tensor(seq_lens_list, dtype=torch.int32, device=DEVICE)
|
||||
return k, v, req_to_token, slot_ids, seq_lens
|
||||
|
||||
|
||||
def make_topk_idx(seq_lens_list, num_kv_heads, rows_per_req, topk):
|
||||
"""Sorted-ascending causal-safe topk over each request's block count.
|
||||
rows_per_req[b] = number of query rows for request b (1 for decode)."""
|
||||
total_rows = sum(rows_per_req)
|
||||
tidx = torch.full(
|
||||
(num_kv_heads, total_rows, topk), -1, dtype=torch.int32, device=DEVICE
|
||||
)
|
||||
row = 0
|
||||
for b, s in enumerate(seq_lens_list):
|
||||
nb = (s + P - 1) // P
|
||||
for _ in range(rows_per_req[b]):
|
||||
ak = min(topk, nb)
|
||||
sel = torch.randperm(nb, device=DEVICE)[:ak].sort().values
|
||||
for kh in range(num_kv_heads):
|
||||
tidx[kh, row, :ak] = sel.to(torch.int32)
|
||||
row += 1
|
||||
return tidx
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# decode
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _decode_setup(seq_lens_list=(1023, 769, 513, 130), topk=4):
|
||||
"""topk*P = 512 < 1023/769: exercises the degenerate partial-last-block
|
||||
masking (the historical unsorted-topk failure shape)."""
|
||||
torch.manual_seed(0)
|
||||
batch = len(seq_lens_list)
|
||||
k, v, r2t, sids, seq_lens = build_paged_inputs(seq_lens_list)
|
||||
q = torch.randn(batch, 16, 128, dtype=torch.bfloat16, device=DEVICE)
|
||||
tidx = make_topk_idx(seq_lens_list, 1, [1] * batch, topk)
|
||||
return q, k, v, r2t, sids, seq_lens, tidx
|
||||
|
||||
|
||||
def test_msa_fp8_decode_vs_triton_fp8():
|
||||
q, k, v, r2t, sids, seq_lens, tidx = _decode_setup()
|
||||
q8, qr = qdq(q)
|
||||
k8, kr = qdq(k)
|
||||
v8, vr = qdq(v)
|
||||
o_msa = msa_sparse_decode_main(
|
||||
q8, k8, v8, tidx, r2t, sids, seq_lens, block_size_k=P
|
||||
)
|
||||
o_triton = flash_decode_with_gqa_share_sparse(
|
||||
q8, None, k8, v8, r2t, seq_lens, sids, P, tidx
|
||||
)
|
||||
assert o_msa.dtype == torch.bfloat16
|
||||
torch.testing.assert_close(
|
||||
o_msa.float(), o_triton.float(), atol=X_ATOL, rtol=X_RTOL
|
||||
)
|
||||
# bf16 MSA on the dequantized tensors bounds the pure fp8-kernel error
|
||||
o_bf16 = msa_sparse_decode_main(
|
||||
qr, kr, vr, tidx, r2t, sids, seq_lens, block_size_k=P
|
||||
)
|
||||
err = (o_msa.float() - o_bf16.float()).abs().mean()
|
||||
ref = o_bf16.float().abs().mean()
|
||||
assert err / ref < 0.06, f"mean rel err {err/ref:.4f} too high vs bf16 MSA"
|
||||
|
||||
|
||||
def test_msa_fp8_decode_scales():
|
||||
q, k, v, r2t, sids, seq_lens, tidx = _decode_setup()
|
||||
q8, qr = qdq(q)
|
||||
k8, kr = qdq(k)
|
||||
v8, vr = qdq(v)
|
||||
q_scale, k_scale, v_scale = 1.5, 0.5, 2.0
|
||||
o = msa_sparse_decode_main(
|
||||
q8,
|
||||
k8,
|
||||
v8,
|
||||
tidx,
|
||||
r2t,
|
||||
sids,
|
||||
seq_lens,
|
||||
block_size_k=P,
|
||||
q_scale=q_scale,
|
||||
k_scale=k_scale,
|
||||
v_scale=v_scale,
|
||||
)
|
||||
# reference: bf16 MSA on pre-scaled dequantized Q/K/V
|
||||
ref = msa_sparse_decode_main(
|
||||
(qr * q_scale).to(torch.bfloat16),
|
||||
(kr * k_scale).to(torch.bfloat16),
|
||||
(vr * v_scale).to(torch.bfloat16),
|
||||
tidx,
|
||||
r2t,
|
||||
sids,
|
||||
seq_lens,
|
||||
block_size_k=P,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
o.float(), ref.float(), atol=X_ATOL * v_scale, rtol=X_RTOL * v_scale
|
||||
)
|
||||
|
||||
|
||||
def test_msa_fp8_decode_capture_replay_bitexact():
|
||||
q, k, v, r2t, sids, seq_lens, tidx = _decode_setup()
|
||||
q8, _ = qdq(q)
|
||||
k8, _ = qdq(k)
|
||||
v8, _ = qdq(v)
|
||||
bs = q8.shape[0]
|
||||
nb_max = r2t.shape[1] // P
|
||||
plan = build_msa_decode_cg_plan(
|
||||
16, 1, P, tidx.shape[-1], bs, device=q8.device, is_fp8=True
|
||||
)
|
||||
kv_indices = torch.zeros(bs * nb_max, dtype=torch.int32, device=DEVICE)
|
||||
update_msa_decode_cg_meta(
|
||||
plan, kv_indices, r2t, sids, seq_lens, P, tidx.shape[-1], 16, 1
|
||||
)
|
||||
|
||||
def run():
|
||||
return msa_sparse_decode_main(
|
||||
q8,
|
||||
k8,
|
||||
v8,
|
||||
tidx,
|
||||
r2t,
|
||||
sids,
|
||||
seq_lens,
|
||||
block_size_k=P,
|
||||
kv_indices=kv_indices,
|
||||
plan=plan,
|
||||
)
|
||||
|
||||
# eager warmups (also pays the fp8 JIT) on a side stream, as capture does
|
||||
s = torch.cuda.Stream()
|
||||
s.wait_stream(torch.cuda.current_stream())
|
||||
with torch.cuda.stream(s):
|
||||
for _ in range(2):
|
||||
o_eager = run()
|
||||
torch.cuda.current_stream().wait_stream(s)
|
||||
|
||||
g = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(g):
|
||||
o_graph = run()
|
||||
g.replay()
|
||||
torch.cuda.synchronize()
|
||||
assert torch.equal(o_graph, o_eager), "fp8 MSA decode capture/replay not bit-exact"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# prefill (cutlass short-q and cute long-q branches)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _prefill_setup(seq_lens_list, prefix_lens_list, topk=4):
|
||||
torch.manual_seed(1)
|
||||
k, v, r2t, sids, seq_lens = build_paged_inputs(seq_lens_list)
|
||||
q_lens = [s - p for s, p in zip(seq_lens_list, prefix_lens_list)]
|
||||
total_q = sum(q_lens)
|
||||
q = torch.randn(total_q, 16, 128, dtype=torch.bfloat16, device=DEVICE)
|
||||
cu = torch.zeros(len(q_lens) + 1, dtype=torch.int32, device=DEVICE)
|
||||
cu[1:] = torch.tensor(q_lens, device=DEVICE).cumsum(0)
|
||||
prefix = torch.tensor(prefix_lens_list, dtype=torch.int32, device=DEVICE)
|
||||
# causal-valid per-token topk (block_size_q == 1)
|
||||
tidx = torch.full((1, total_q, topk), -1, dtype=torch.int32, device=DEVICE)
|
||||
row = 0
|
||||
for b, (s, p) in enumerate(zip(seq_lens_list, prefix_lens_list)):
|
||||
for j in range(s - p):
|
||||
nb = (p + j) // P + 1
|
||||
ak = min(topk, nb)
|
||||
sel = torch.randperm(nb, device=DEVICE)[:ak].sort().values
|
||||
tidx[0, row, :ak] = sel.to(torch.int32)
|
||||
row += 1
|
||||
return q, k, v, r2t, sids, cu, seq_lens, prefix, tidx, max(q_lens)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"seq_lens,prefix_lens,branch",
|
||||
[
|
||||
((530, 700), (500, 680), "cutlass_short_q"), # qo <= 32
|
||||
((513, 769), (0, 257), "cute_long_q"), # qo > 32
|
||||
],
|
||||
ids=["cutlass_short_q", "cute_long_q"],
|
||||
)
|
||||
def test_msa_fp8_prefill_vs_triton_fp8(seq_lens, prefix_lens, branch):
|
||||
q, k, v, r2t, sids, cu, seq_lens_t, prefix, tidx, max_q = _prefill_setup(
|
||||
list(seq_lens), list(prefix_lens)
|
||||
)
|
||||
q8, _ = qdq(q)
|
||||
k8, _ = qdq(k)
|
||||
v8, _ = qdq(v)
|
||||
o_msa = msa_sparse_prefill_main(
|
||||
q8, k8, v8, tidx, r2t, sids, cu, seq_lens_t, prefix, block_size_k=P
|
||||
)
|
||||
o_triton = flash_prefill_with_gqa_share_sparse(
|
||||
q=q8,
|
||||
k_cache=k8,
|
||||
v_cache=v8,
|
||||
sink=None,
|
||||
req_to_token=r2t,
|
||||
slot_ids=sids,
|
||||
topk_idx=tidx,
|
||||
block_size_q=1,
|
||||
block_size_k=P,
|
||||
cu_seqlens=cu,
|
||||
seq_lens=seq_lens_t,
|
||||
prefix_lens=prefix,
|
||||
max_seqlen_q=max_q,
|
||||
)
|
||||
assert o_msa.dtype == torch.bfloat16
|
||||
torch.testing.assert_close(
|
||||
o_msa.float(), o_triton.float(), atol=X_ATOL, rtol=X_RTOL
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import sys
|
||||
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
@@ -204,6 +204,32 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
||||
# KV fp8: q_type = fp8, out_type=model_runner.dtype
|
||||
self.is_xqa_impl = is_sm90_supported() or is_sm120_supported()
|
||||
|
||||
# trtllm-gen serves page_size >= 128 only through its dynamic
|
||||
# tokens-per-page kernels, which exist solely for GQA with equal QK/V
|
||||
# head dims (power-of-2 pages). Mirror that precondition here so an
|
||||
# unsupported combo fails at construction instead of as a
|
||||
# "Missing TRTLLM-GEN kernel" error during CUDA-graph capture.
|
||||
# XQA (SM90/SM120 decode) has native page-128 kernels; no check needed.
|
||||
if self.page_size >= 128 and not self.is_xqa_impl:
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
attn_tp_size = get_parallel().attn_tp_size
|
||||
num_q_heads = config.num_attention_heads // attn_tp_size
|
||||
num_kv_heads = config.get_num_kv_heads(attn_tp_size)
|
||||
if (
|
||||
num_q_heads // num_kv_heads <= 1
|
||||
or config.head_dim != config.v_head_dim
|
||||
or self.page_size & (self.page_size - 1) != 0
|
||||
):
|
||||
raise ValueError(
|
||||
f"trtllm_mha with page_size={self.page_size} requires "
|
||||
f"trtllm-gen's dynamic tokens-per-page kernels, which only "
|
||||
f"support GQA (q heads per kv head > 1, got "
|
||||
f"{num_q_heads}/{num_kv_heads}) with equal QK/V head dims "
|
||||
f"(got {config.head_dim}/{config.v_head_dim}) and a "
|
||||
f"power-of-2 page size. Use --page-size 64 instead."
|
||||
)
|
||||
|
||||
def _check_decode_kv_access(self) -> None:
|
||||
supported_kinds = {
|
||||
KVCacheAttentionAccessKind.PLAIN,
|
||||
|
||||
@@ -128,6 +128,13 @@ class RadixAttention(nn.Module):
|
||||
self.v_scale = None
|
||||
self.k_scale_float = None
|
||||
self.v_scale_float = None
|
||||
# MiniMax-M3 fp8 attention-GEMM scales (fp8 attn-GEMM mode): main q and
|
||||
# lightning-indexer q/k/v. No checkpoint loader populates them yet;
|
||||
# None means unit scale.
|
||||
self.q_scale_float = None
|
||||
self.idx_q_scale_float = None
|
||||
self.idx_k_scale_float = None
|
||||
self.idx_v_scale_float = None
|
||||
self.quant_method = None
|
||||
|
||||
if quant_config is not None:
|
||||
|
||||
@@ -397,11 +397,18 @@ class HybridCacheController(BaseHiCacheController):
|
||||
if not self.write_queue:
|
||||
return
|
||||
op = CacheOperation.merge_ops(self.write_queue)
|
||||
# Page-first write-back JIT kernels can keep destination host indices on CPU.
|
||||
# Page-first staged write-back kernels need CPU destination host indices.
|
||||
# A HostPoolGroup may mix staged and non-staged child pools, so let it
|
||||
# normalize indices per child instead of moving the whole operation here.
|
||||
if (
|
||||
self.io_backend == "kernel"
|
||||
and self.mem_pool_host.layout == "page_first"
|
||||
and getattr(self.mem_pool_host, "can_use_write_back_jit", False)
|
||||
and (
|
||||
getattr(self.mem_pool_host, "can_use_write_back_jit", False)
|
||||
or getattr(
|
||||
self.mem_pool_host, "supports_per_pool_backup_indices", False
|
||||
)
|
||||
)
|
||||
):
|
||||
host_indices = op.host_indices
|
||||
device_indices = op.device_indices
|
||||
|
||||
@@ -1271,6 +1271,8 @@ class KVCacheConfigurator:
|
||||
return token_to_kv_pool
|
||||
|
||||
def _build_minimax_sparse_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
|
||||
from sglang.srt.server_args import m3_fp8_attn_gemm_enabled
|
||||
|
||||
_hf_config = self.model_config.hf_config
|
||||
sparse_cfg = get_minimax_sparse_attention_config(_hf_config)
|
||||
dense_layer_ids, sparse_layer_ids = get_minimax_sparse_layer_ids(sparse_cfg)
|
||||
@@ -1281,7 +1283,15 @@ class KVCacheConfigurator:
|
||||
size=max_total_num_tokens,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
index_dtype=self.model_dtype,
|
||||
# fp8 attn-GEMM mode opts the lightning-indexer cache into
|
||||
# fp8 too (fp8 indexer GEMMs); fp8 KV without the mode
|
||||
# (e5m2 or non-trtllm_mha backend) keeps the indexer bf16
|
||||
# with the widening-dequant contract.
|
||||
index_dtype=(
|
||||
self.kv_cache_dtype
|
||||
if m3_fp8_attn_gemm_enabled(self.server_args)
|
||||
else self.model_dtype
|
||||
),
|
||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||
head_dim=self.model_config.head_dim,
|
||||
idx_head_dim=sparse_cfg["sparse_index_dim"],
|
||||
|
||||
@@ -4810,6 +4810,11 @@ class MiniMaxSparseKVPool(KVCache):
|
||||
) -> None:
|
||||
self.layer_transfer_counter = layer_transfer_counter
|
||||
|
||||
def get_kv_cache_quant_method(self) -> Any:
|
||||
# The base unwrap chain only knows full_kv_pool/swa_kv_pool; the dense
|
||||
# KV (what attention backends quantize against) lives in main_pool here.
|
||||
return self.main_pool.get_kv_cache_quant_method()
|
||||
|
||||
def _wait_for_layer(self, layer_id: int) -> None:
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||
@@ -4858,10 +4863,14 @@ class MiniMaxSparseKVPool(KVCache):
|
||||
loc: torch.Tensor,
|
||||
cache_k: torch.Tensor,
|
||||
cache_v: torch.Tensor,
|
||||
k_scale: float = 1.0,
|
||||
v_scale: float = 1.0,
|
||||
k_scale: Optional[float] = None,
|
||||
v_scale: Optional[float] = None,
|
||||
) -> None:
|
||||
"""Write main K/V at `loc`. Works for any layer (dense or sparse)."""
|
||||
"""Write main K/V at `loc`. Works for any layer (dense or sparse).
|
||||
|
||||
Scale semantics follow MHATokenToKVPool: None means unit scale;
|
||||
a non-None scale is applied with an in-place div_ before the fp8 cast.
|
||||
"""
|
||||
self.main_pool.set_kv_buffer(
|
||||
layer,
|
||||
loc,
|
||||
@@ -4877,8 +4886,8 @@ class MiniMaxSparseKVPool(KVCache):
|
||||
loc: torch.Tensor,
|
||||
cache_idx_k: torch.Tensor,
|
||||
cache_idx_v: torch.Tensor,
|
||||
k_scale: float = 1.0,
|
||||
v_scale: float = 1.0,
|
||||
k_scale: Optional[float] = None,
|
||||
v_scale: Optional[float] = None,
|
||||
) -> None:
|
||||
mapped_id = self.index_kv_layer_id_mapping.get(layer.layer_id)
|
||||
if mapped_id is None:
|
||||
@@ -4902,6 +4911,7 @@ class MiniMaxSparseKVPool(KVCache):
|
||||
layer: RadixAttention,
|
||||
loc: torch.Tensor,
|
||||
cache_idx_k: torch.Tensor,
|
||||
k_scale: Optional[float] = None,
|
||||
) -> None:
|
||||
mapped_id = self.index_k_layer_id_mapping.get(layer.layer_id)
|
||||
if mapped_id is None:
|
||||
@@ -4912,6 +4922,8 @@ class MiniMaxSparseKVPool(KVCache):
|
||||
)
|
||||
sub_pool = self.index_k_pool
|
||||
if cache_idx_k.dtype != sub_pool.dtype:
|
||||
if k_scale is not None:
|
||||
cache_idx_k = cache_idx_k / k_scale
|
||||
cache_idx_k = cache_idx_k.to(sub_pool.dtype)
|
||||
if sub_pool.store_dtype != sub_pool.dtype:
|
||||
cache_idx_k = cache_idx_k.view(sub_pool.store_dtype)
|
||||
@@ -4950,6 +4962,10 @@ class MiniMaxSparseKVPool(KVCache):
|
||||
cache_v: torch.Tensor,
|
||||
cache_idx_k: torch.Tensor,
|
||||
cache_idx_v: Optional[torch.Tensor],
|
||||
k_scale: Optional[float] = None,
|
||||
v_scale: Optional[float] = None,
|
||||
idx_k_scale: Optional[float] = None,
|
||||
idx_v_scale: Optional[float] = None,
|
||||
) -> None:
|
||||
"""Store main K/V + index K (+ optional index V) for a sparse layer in
|
||||
one fused JIT launch, falling back to separate stores when not applicable."""
|
||||
@@ -4984,12 +5000,24 @@ class MiniMaxSparseKVPool(KVCache):
|
||||
)
|
||||
return
|
||||
|
||||
# Fallback: separate stores (identical semantics).
|
||||
self.set_kv_buffer(layer, loc, cache_k, cache_v)
|
||||
# Fallback: separate stores (identical semantics; quantizes for fp8
|
||||
# pools — the fused raw-byte path is disqualified there by
|
||||
# _can_fuse_kv_index_store's dtype-equality checks). Scales use the
|
||||
# None-means-unit convention throughout: MHATokenToKVPool.set_kv_buffer
|
||||
# applies any non-None scale with an IN-PLACE div_ (extra kernel +
|
||||
# caller-tensor mutation), which must not fire for unit scale.
|
||||
self.set_kv_buffer(layer, loc, cache_k, cache_v, k_scale, v_scale)
|
||||
if disable_value:
|
||||
self.set_index_k_buffer(layer, loc, cache_idx_k)
|
||||
self.set_index_k_buffer(layer, loc, cache_idx_k, idx_k_scale)
|
||||
else:
|
||||
self.set_index_kv_buffer(layer, loc, cache_idx_k, cache_idx_v)
|
||||
self.set_index_kv_buffer(
|
||||
layer,
|
||||
loc,
|
||||
cache_idx_k,
|
||||
cache_idx_v,
|
||||
idx_k_scale,
|
||||
idx_v_scale,
|
||||
)
|
||||
|
||||
def get_kv_size_bytes(self):
|
||||
sub_pools = [self.main_pool, self.index_kv_pool, self.index_k_pool]
|
||||
|
||||
@@ -1528,10 +1528,12 @@ class HostPoolGroup:
|
||||
self.page_size = self.anchor_entry.host_pool.page_size
|
||||
self.device = self.anchor_entry.host_pool.device
|
||||
self.size = self.anchor_entry.host_pool.size
|
||||
self.can_use_write_back_jit = all(
|
||||
child_write_back_jit = [
|
||||
getattr(entry.host_pool, "can_use_write_back_jit", False)
|
||||
for entry in entries
|
||||
)
|
||||
]
|
||||
self.can_use_write_back_jit = all(child_write_back_jit)
|
||||
self.supports_per_pool_backup_indices = any(child_write_back_jit)
|
||||
|
||||
@property
|
||||
def kv_buffer(self):
|
||||
@@ -1632,6 +1634,39 @@ class HostPoolGroup:
|
||||
io_backend,
|
||||
)
|
||||
|
||||
def _backup_uses_cpu_host_indices(self, host_pool, io_backend) -> bool:
|
||||
return (
|
||||
io_backend == "kernel"
|
||||
and getattr(host_pool, "layout", None) == "page_first"
|
||||
and getattr(host_pool, "can_use_write_back_jit", False)
|
||||
)
|
||||
|
||||
def _kernel_index_device(self, entry, device_indices):
|
||||
if device_indices is not None and device_indices.is_cuda:
|
||||
return device_indices.device
|
||||
return getattr(entry.device_pool, "device", None)
|
||||
|
||||
def _normalize_backup_indices(
|
||||
self, entry, host_indices, device_indices, io_backend
|
||||
):
|
||||
if io_backend != "kernel":
|
||||
return host_indices, device_indices
|
||||
|
||||
if self._backup_uses_cpu_host_indices(entry.host_pool, io_backend):
|
||||
if host_indices.is_cuda:
|
||||
host_indices = host_indices.cpu()
|
||||
return host_indices, device_indices
|
||||
|
||||
if not host_indices.is_cuda:
|
||||
target_device = self._kernel_index_device(entry, device_indices)
|
||||
if target_device is not None:
|
||||
host_indices = host_indices.to(target_device, non_blocking=True)
|
||||
if host_indices.is_cuda:
|
||||
host_indices.record_stream(
|
||||
torch.cuda.current_stream(host_indices.device)
|
||||
)
|
||||
return host_indices, device_indices
|
||||
|
||||
def backup_from_device_all_layer(
|
||||
self,
|
||||
device_pool,
|
||||
@@ -1641,10 +1676,13 @@ class HostPoolGroup:
|
||||
pool_transfers: Optional[list] = None,
|
||||
) -> None:
|
||||
# 1. Anchor (KV) backup
|
||||
anchor_host_indices, anchor_device_indices = self._normalize_backup_indices(
|
||||
self.anchor_entry, host_indices, device_indices, io_backend
|
||||
)
|
||||
self.anchor_entry.host_pool.backup_from_device_all_layer(
|
||||
self.anchor_entry.device_pool,
|
||||
host_indices,
|
||||
device_indices,
|
||||
anchor_host_indices,
|
||||
anchor_device_indices,
|
||||
io_backend,
|
||||
)
|
||||
# 2. Extra pool backup
|
||||
@@ -1652,10 +1690,18 @@ class HostPoolGroup:
|
||||
entry = self.entry_map.get(transfer.name)
|
||||
if entry is None or transfer.host_indices is None:
|
||||
continue
|
||||
transfer_host_indices, transfer_device_indices = (
|
||||
self._normalize_backup_indices(
|
||||
entry,
|
||||
transfer.host_indices,
|
||||
transfer.device_indices,
|
||||
io_backend,
|
||||
)
|
||||
)
|
||||
entry.host_pool.backup_from_device_all_layer(
|
||||
entry.device_pool,
|
||||
transfer.host_indices,
|
||||
transfer.device_indices,
|
||||
transfer_host_indices,
|
||||
transfer_device_indices,
|
||||
io_backend,
|
||||
)
|
||||
|
||||
|
||||
@@ -9025,6 +9025,28 @@ class ServerArgs:
|
||||
}
|
||||
|
||||
|
||||
def m3_fp8_attn_gemm_enabled(args) -> bool:
|
||||
"""Whether MiniMax-M3 attention GEMMs run in fp8 (no opt-in flag; active
|
||||
whenever possible): fp8_e4m3 main + index KV caches, fp8-cast q, fp8
|
||||
sparse/MSA kernels, with dense layers on trtllm_mha's fp8-q path. Needs
|
||||
kv_cache_dtype fp8_e4m3 (e5m2 would silently mis-dispatch fmha_sm100's
|
||||
e4m3 kernel), the trtllm_mha backend (the only dense backend with fp8-q
|
||||
GEMMs), and SM100 (MSA fp8 variants and trtllm-gen fp8 dense kernels are
|
||||
sm100-only). SGLANG_DISABLE_M3_FP8_ATTN_GEMM=1 is the kill switch:
|
||||
it forces the pre-fp8 numerics (bf16 indexer + widening sparse path,
|
||||
bf16 q) without having to move off trtllm_mha.
|
||||
"""
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.utils.common import is_sm100_supported
|
||||
|
||||
return (
|
||||
args.kv_cache_dtype == "fp8_e4m3"
|
||||
and args.attention_backend == "trtllm_mha"
|
||||
and is_sm100_supported()
|
||||
and not envs.SGLANG_DISABLE_M3_FP8_ATTN_GEMM.get()
|
||||
)
|
||||
|
||||
|
||||
# NOTE: The process-wide ServerArgs is owned by the runtime context
|
||||
# (sglang.srt.runtime_context). The two functions below are LEGACY shims kept
|
||||
# for the existing call-sites; they publish/read the same live object by
|
||||
|
||||
Reference in New Issue
Block a user