[Kernel] Add SM90 Q8KV8 FP8 Sparse MLA Prefill JIT Kernel with Tests and Benchmark (#25751)

This commit is contained in:
Ho-Ren (Jack) Chuang
2026-06-30 09:00:11 +08:00
committed by GitHub
parent a6bc432fd8
commit bc8b3ab1f5
11 changed files with 2500 additions and 0 deletions
@@ -0,0 +1,55 @@
/* Copyright 2025 SGLang Team. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#pragma once
#include <cute/tensor.hpp>
#include <cutlass/arch/arch.h>
#include <cutlass/arch/reg_reconfig.h>
#include <cutlass/cluster_launch.hpp>
#include "defines.h"
#include "params.h"
#include <cstdio>
#include <cstdlib>
#include <math_constants.h>
#define KU_ASSERT(cond) \
do { \
if (!(cond)) { \
fprintf(stderr, "KU_ASSERT failed (%s:%d): %s\n", __FILE__, __LINE__, #cond); \
exit(1); \
} \
} while (0)
#define KU_CUDA_CHECK(call) \
do { \
cudaError_t err = (call); \
if (err != cudaSuccess) { \
fprintf(stderr, "CUDA error (%s:%d): %s\n", __FILE__, __LINE__, cudaGetErrorString(err)); \
exit(1); \
} \
} while (0)
#define KU_CHECK_KERNEL_LAUNCH() KU_CUDA_CHECK(cudaGetLastError())
namespace ku {
template <typename T>
__host__ __device__ __forceinline__ T ceil_div(T a, T b) {
return (a + b - 1) / b;
}
} // namespace ku
@@ -0,0 +1,26 @@
/* Copyright 2025 SGLang Team. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#pragma once
#include <cutlass/arch/barrier.h>
#include <cutlass/bfloat16.h>
using bf16 = cutlass::bfloat16_t;
using fp8 = cutlass::float_e4m3_t;
using transac_bar_t = cutlass::arch::ClusterTransactionBarrier;
using cutlass::arch::fence_barrier_init;
using cutlass::arch::fence_view_async_shared;
using cutlass::arch::NamedBarrier;
@@ -0,0 +1,100 @@
/*
* Taken from FlashMLA PR https://github.com/deepseek-ai/FlashMLA/pull/54
* originally authored by @endurehero
*/
/**
* ref to Fa3's SmemTranspose64x64:
* https://github.com/Dao-AILab/flash-attention/blob/0823cf7b5d96499c1c79a4f64b1e256a035ba4b4/hopper/mainloop_fwd_sm90_tma_gmma_ws.hpp#L26
*/
#pragma once
template <int kBlockN, int kHeadDim>
struct SmemTransposeFp8_64x64 {
static_assert((kBlockN % 64 == 0) && (kHeadDim % 64 == 0));
using Element = cutlass::float_e4m3_t;
using TransposeShapeAtomV = Shape<_64, _64>;
using SmemLayoutAtomV = decltype(tile_to_shape(GMMA::Layout_K_SW64_Atom<Element>{}, TransposeShapeAtomV{}));
using SmemLayoutV = decltype(tile_to_shape(SmemLayoutAtomV{}, Shape<Int<kBlockN>, Int<kHeadDim>>{}));
// for fp8 in-kernel transpose -- src layout
using SmemLayoutDivideV = decltype(tiled_divide(SmemLayoutV{}, TransposeShapeAtomV{}));
using SmemShapeLDSM = Shape<Shape<_8, _8>, Shape<_16, _4>>;
using FactoringShapeV =
decltype(make_shape(SmemShapeLDSM{}, shape<1>(SmemLayoutDivideV{}), shape<2>(SmemLayoutDivideV{})));
using SmemLayoutTransposeV = decltype(composition(SmemLayoutDivideV{}, make_layout(FactoringShapeV{})));
// For fp8, this is the memory transpose.
using SmemLayoutAtomVt = decltype(tile_to_shape(GMMA::Layout_K_SW64_Atom<Element>{}, TransposeShapeAtomV{}));
using SmemLayoutVt = decltype(tile_to_shape(SmemLayoutAtomVt{}, Shape<Int<kHeadDim>, Int<kBlockN>>{}));
// for fp8 in-kernel transpose -- dst layout
using SmemLayoutVtTrans =
decltype(composition(SmemLayoutVt{}, make_ordered_layout(product_each(shape(SmemLayoutV{})), Step<_2, _1>{})));
using SmemLayoutDivideVt = decltype(tiled_divide(SmemLayoutVtTrans{}, TransposeShapeAtomV{}));
using SmemShapeSTSM = Shape<Shape<_16, _4>, Shape<_16, _4>>;
using FactoringShapeVt =
decltype(make_shape(SmemShapeSTSM{}, shape<1>(SmemLayoutDivideVt{}), shape<2>(SmemLayoutDivideVt{})));
using SmemLayoutTransposeVt = decltype(composition(SmemLayoutDivideVt{}, make_layout(FactoringShapeVt{})));
using ldsm_thread_shape = Shape<_4, _1, _8, _4>;
using ldsm_value_shape = Shape<_2, _8, _2, _1>;
using ldsm_value_stride = Stride<_2, _4, _1, _0>;
using TiledCopyLDSM = decltype(make_tiled_copy(
Copy_Atom<SM75_U16x8_LDSM_T, Element>{},
Layout<ldsm_thread_shape>{},
Layout<ldsm_value_shape, ldsm_value_stride>{}));
TiledCopyLDSM tiled_copy_ldsm;
using stsm_thread_shape = Shape<_4, _1, _8, _4>;
using stsm_value_shape = Shape<_4, _4, _2, _1>;
using stsm_value_stride = Stride<_1, _8, _4, _0>;
using TiledCopySTSM = decltype(make_tiled_copy(
Copy_Atom<SM90_U32x4_STSM_N, Element>{},
Layout<stsm_thread_shape>{},
Layout<stsm_value_shape, stsm_value_stride>{}));
TiledCopySTSM tiled_copy_stsm;
template <class SmemTensor, class SmemTensorOut>
CUTLASS_DEVICE void
transpose_pair(SmemTensor&& s_in0, SmemTensorOut&& s_out0, SmemTensor&& s_in1, SmemTensorOut&& s_out1) {
using namespace cute;
auto tid = threadIdx.x % cutlass::NumThreadsPerWarpGroup;
auto thr_copy_ldsm = tiled_copy_ldsm.get_thread_slice(tid);
auto thr_copy_stsm = tiled_copy_stsm.get_thread_slice(tid);
auto tXsX0 = thr_copy_ldsm.partition_S(s_in0);
auto tXrX0 = make_tensor<Element>(shape(tXsX0));
auto tXsX_out0 = thr_copy_stsm.partition_D(s_out0);
auto tXsX1 = thr_copy_ldsm.partition_S(s_in1);
auto tXrX1 = make_tensor<Element>(shape(tXsX1));
auto tXsX_out1 = thr_copy_stsm.partition_D(s_out1);
auto data0 = tXrX0.data();
auto data1 = tXrX1.data();
cute::copy(tiled_copy_ldsm, tXsX0, tXrX0);
cute::copy(tiled_copy_ldsm, tXsX1, tXrX1);
CUTLASS_PRAGMA_UNROLL
for (int n = 0; n < size(tXrX0); n += 8) {
uint32_t* d0 = reinterpret_cast<uint32_t*>(&data0[n]);
uint32_t* d1 = reinterpret_cast<uint32_t*>(&data1[n]);
auto upper0 = d0[0];
auto lower0 = d0[1];
auto upper1 = d1[0];
auto lower1 = d1[1];
d0[0] = __byte_perm(upper0, lower0, 0x6420);
d0[1] = __byte_perm(upper0, lower0, 0x7531);
d1[0] = __byte_perm(upper1, lower1, 0x6420);
d1[1] = __byte_perm(upper1, lower1, 0x7531);
}
cute::copy(tiled_copy_stsm, tXrX0, tXsX_out0);
cute::copy(tiled_copy_stsm, tXrX1, tXsX_out1);
}
};
@@ -0,0 +1,107 @@
/*
* Taken from FlashMLA PR https://github.com/deepseek-ai/FlashMLA/pull/54
* originally authored by @endurehero
*/
// Adapted from https://github.com/Dao-AILab/flash-attention/blob/main/hopper/utils.h
#pragma once
#include <cute/tensor.hpp>
#include <cutlass/array.h>
#include <cutlass/cutlass.h>
#include <cutlass/numeric_conversion.h>
#include <cutlass/numeric_types.h>
namespace flash {
////////////////////////////////////////////////////////////////////////////////////////////////////
// For SM80, convert acc_layout from (MMA=4, MMA_M, MMA_N) to ((4, 2), MMA_M, MMA_N / 2)
// if using m16n8k16, or to (4, MMA_M, MMA_N) if using m16n8k8.
// For SM90, FP16/BF16, convert acc_layout from ((2, 2, N / 8), MMA_M, MMA_N) to ((2, 2, 2), MMA_M, (N / 16, MMA_N))
// For SM90, FP8, convert acc_layout from ((2, 2, N / 8), MMA_M, MMA_N) to ((4, 2, 2), MMA_M, (N / 32, MMA_N))
template <typename MMA_Traits, typename Layout0>
__forceinline__ __device__ auto convert_layout_acc_Aregs(Layout0 acc_layout) {
using X = Underscore;
if constexpr (decltype(rank<0>(acc_layout))::value == 3) { // SM90
static_assert(decltype(size<0, 0>(acc_layout))::value == 2);
static_assert(decltype(size<0, 1>(acc_layout))::value == 2);
static_assert(decltype(rank(acc_layout))::value == 3);
static_assert(decltype(rank(get<0>(acc_layout)))::value == 3);
if constexpr (sizeof(typename MMA_Traits::ValTypeA) == 2) {
auto l = logical_divide(get<0, 2>(acc_layout), Tile<_2>{}); // ((2, N / 16))
return make_layout(
make_layout(get<0, 0>(acc_layout), get<0, 1>(acc_layout), get<0, 0>(l)),
get<1>(acc_layout),
coalesce(make_layout(get<0, 1>(l), get<2>(acc_layout))));
} else {
static_assert(sizeof(typename MMA_Traits::ValTypeA) == 1);
static_assert(decltype(stride<0, 0>(acc_layout))::value == 1);
static_assert(decltype(stride<0, 1>(acc_layout))::value == 2);
auto l = logical_divide(get<0, 2>(acc_layout), Tile<Layout<Shape<_2, _2>>>{}); // (((2, 2), N / 32))
// This combines the first two modes (<0, 0> and <0, 1>) into one mode.
// Will require register shuffling later to be correct.
return make_layout(
make_layout(Layout<_4>{}, get<0, 0, 0>(l), get<0, 0, 1>(l)),
get<1>(acc_layout),
coalesce(make_layout(get<0, 1>(l), get<2>(acc_layout)))); // ((4, 2, 2), MMA_M, N / 32 * MMA_N)
// This combination is right but doesn't work with register shuffling.
// return make_layout(make_layout(coalesce(make_layout(get<0, 0>(acc_layout), get<0, 0, 0>(l))), get<0,
// 1>(acc_layout), get<0, 0, 1>(l)),
// get<1>(acc_layout),
// coalesce(make_layout(get<0, 1>(l), get<2>(acc_layout))));
}
} else { // SM80
static_assert(decltype(size<0>(acc_layout))::value == 4);
static_assert(decltype(rank(acc_layout))::value == 3);
constexpr int mma_shape_K = get<2>(typename MMA_Traits::Shape_MNK{});
static_assert(mma_shape_K == 8 || mma_shape_K == 16);
if constexpr (mma_shape_K == 8) {
return acc_layout;
} else {
auto l = logical_divide(acc_layout, Shape<X, X, _2>{}); // (4, MMA_M, (2, MMA_N / 2)))
return make_layout(make_layout(get<0>(l), get<2, 0>(l)), get<1>(l), get<2, 1>(l));
}
}
};
template <typename Fragment>
CUTLASS_DEVICE void permute_Cregs_fp8(Fragment& frag) {
// frag has shape ((2, 2, N / 8), MMA_M, MMA_N), each element is 32 bits
static_assert(decltype(size<0, 0>(frag))::value == 2);
static_assert(decltype(size<0, 1>(frag))::value == 2);
static_assert(decltype(size<0, 2>(frag))::value % 2 == 0);
static_assert(decltype(stride<0, 0>(frag))::value == 1);
static_assert(sizeof(typename Fragment::value_type) == 4);
Tensor frag_64b = group_modes<1, 3>(recast<uint2>(frag)); // ((1, 2, N / 8), (MMA_M, MMA_N))
#pragma unroll
for (int mi = 0; mi < size<1>(frag_64b); ++mi) {
#pragma unroll
for (int i = 0; i < size<0, 2>(frag_64b) / 2; ++i) {
cutlass::swap(frag_64b(make_coord(_0{}, _1{}, 2 * i), mi), frag_64b(make_coord(_0{}, _0{}, 2 * i + 1), mi));
}
}
}
////////////////////////////////////////////////////////////////////////////////////////////////////
template <typename Engine, typename Layout, typename EngineOut>
CUTLASS_DEVICE void convert_type_out(Tensor<Engine, Layout> const& tensor, Tensor<EngineOut, Layout>& out) {
// Somehow if we allocate out inside this function and return it, e2e is slower and the output can be wrong.
using From_type = typename Engine::value_type;
using To_type = typename EngineOut::value_type;
static constexpr int FragmentSize =
std::max(sizeof(From_type) / sizeof(To_type), sizeof(To_type) / sizeof(From_type));
static_assert(CUTE_STATIC_V(size(tensor)) % FragmentSize == 0, "Fragment size does not vectorize properly");
Tensor frag = recast<cutlass::Array<From_type, FragmentSize> const>(tensor);
Tensor out_frg = recast<cutlass::Array<To_type, FragmentSize>>(out);
static_assert(size(frag) == size(out_frg));
cutlass::NumericArrayConverter<To_type, From_type, FragmentSize> convert_op;
#pragma unroll
for (int i = 0; i < size(frag); ++i) {
out_frg[i] = convert_op(frag[i]);
}
}
} // namespace flash
@@ -0,0 +1,202 @@
/* Copyright 2025 SGLang Team. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
// JIT dispatch entry for the SM90 Q8KV8 sparse MLA prefill kernel.
#pragma once
#include <dlpack/dlpack.h>
#include <tvm/ffi/container/tensor.h>
#include "kernel.cuh"
#include <cmath>
#include <cstdint>
#include <cuda_runtime.h>
namespace {
static inline void
_set_device_and_stream(SparseMlaQ8Kv8PrefillParams& params, tvm::ffi::TensorView q, int64_t cuda_stream) {
DLDevice dev = q.device();
cudaSetDevice(dev.device_id);
params.stream = reinterpret_cast<cudaStream_t>(cuda_stream);
}
template <int D_QK>
static inline void
_run_q8kv8_for_head_dim(SparseMlaQ8Kv8PrefillParams& params, bool have_topk_length, bool have_attn_sink) {
if (have_topk_length) {
if (have_attn_sink) {
sm90::fwd::run_sparse_mla_q8kv8_prefill_kernel<D_QK, true, true>(params);
} else {
sm90::fwd::run_sparse_mla_q8kv8_prefill_kernel<D_QK, true, false>(params);
}
} else {
if (have_attn_sink) {
sm90::fwd::run_sparse_mla_q8kv8_prefill_kernel<D_QK, false, true>(params);
} else {
sm90::fwd::run_sparse_mla_q8kv8_prefill_kernel<D_QK, false, false>(params);
}
}
}
static inline void _run_q8kv8(SparseMlaQ8Kv8PrefillParams& params, bool have_topk_length, bool have_attn_sink) {
switch (params.d_qk) {
case 512:
_run_q8kv8_for_head_dim<512>(params, have_topk_length, have_attn_sink);
return;
case 576:
_run_q8kv8_for_head_dim<576>(params, have_topk_length, have_attn_sink);
return;
default:
fprintf(stderr, "sparse_prefill_q8kv8: unsupported d_qk=%d (must be 512 or 576)\n", params.d_qk);
exit(1);
}
}
static inline SparseMlaQ8Kv8PrefillParams _make_common_params(
tvm::ffi::TensorView q,
tvm::ffi::TensorView kv,
tvm::ffi::TensorView indices,
tvm::ffi::TensorView q_scale,
tvm::ffi::TensorView kv_scale,
tvm::ffi::TensorView out,
tvm::ffi::TensorView max_logits,
tvm::ffi::TensorView lse,
int64_t s_q_val,
int64_t s_kv_val,
int64_t h_q_val,
int64_t h_kv_val,
int64_t d_qk_val,
int64_t d_v_val,
int64_t topk_val,
double sm_scale_val,
int64_t cuda_stream) {
SparseMlaQ8Kv8PrefillParams params;
params.s_q = (int)s_q_val;
params.s_kv = (int)s_kv_val;
params.h_q = (int)h_q_val;
params.h_kv = (int)h_kv_val;
params.d_qk = (int)d_qk_val;
params.d_v = (int)d_v_val;
params.topk = (int)topk_val;
params.sm_scale_div_log2 = (float)sm_scale_val * (float)M_LOG2E;
params.q = reinterpret_cast<const uint8_t*>(q.data_ptr());
params.kv = reinterpret_cast<const uint8_t*>(kv.data_ptr());
params.indices = static_cast<int*>(indices.data_ptr());
params.attn_sink = nullptr;
params.topk_length = nullptr;
params.q_scale_ptr = static_cast<const float*>(q_scale.data_ptr());
params.kv_scale_ptr = static_cast<const float*>(kv_scale.data_ptr());
params.stride_q_s_q = (int)h_q_val * (int)d_qk_val;
params.stride_q_h_q = (int)d_qk_val;
params.stride_kv_s_kv = (int64_t)h_kv_val * (int64_t)d_qk_val;
params.stride_kv_h_kv = (int)d_qk_val;
params.stride_indices_s_q = (int)h_kv_val * (int)topk_val;
params.stride_indices_h_kv = (int)topk_val;
params.out = reinterpret_cast<cutlass::bfloat16_t*>(out.data_ptr());
params.max_logits = static_cast<float*>(max_logits.data_ptr());
params.lse = static_cast<float*>(lse.data_ptr());
_set_device_and_stream(params, q, cuda_stream);
return params;
}
void sparse_prefill_q8kv8_dispatch(
tvm::ffi::TensorView q,
tvm::ffi::TensorView kv,
tvm::ffi::TensorView indices,
tvm::ffi::TensorView q_scale,
tvm::ffi::TensorView kv_scale,
tvm::ffi::TensorView out,
tvm::ffi::TensorView max_logits,
tvm::ffi::TensorView lse,
int64_t s_q_val,
int64_t s_kv_val,
int64_t h_q_val,
int64_t h_kv_val,
int64_t d_qk_val,
int64_t d_v_val,
int64_t topk_val,
double sm_scale_val,
int64_t cuda_stream) {
SparseMlaQ8Kv8PrefillParams params = _make_common_params(
q,
kv,
indices,
q_scale,
kv_scale,
out,
max_logits,
lse,
s_q_val,
s_kv_val,
h_q_val,
h_kv_val,
d_qk_val,
d_v_val,
topk_val,
sm_scale_val,
cuda_stream);
_run_q8kv8(params, false, false);
}
void sparse_prefill_q8kv8_dispatch_full(
tvm::ffi::TensorView q,
tvm::ffi::TensorView kv,
tvm::ffi::TensorView indices,
tvm::ffi::TensorView q_scale,
tvm::ffi::TensorView kv_scale,
tvm::ffi::TensorView attn_sink,
tvm::ffi::TensorView topk_length,
tvm::ffi::TensorView out,
tvm::ffi::TensorView max_logits,
tvm::ffi::TensorView lse,
int64_t s_q_val,
int64_t s_kv_val,
int64_t h_q_val,
int64_t h_kv_val,
int64_t d_qk_val,
int64_t d_v_val,
int64_t topk_val,
double sm_scale_val,
int64_t cuda_stream) {
SparseMlaQ8Kv8PrefillParams params = _make_common_params(
q,
kv,
indices,
q_scale,
kv_scale,
out,
max_logits,
lse,
s_q_val,
s_kv_val,
h_q_val,
h_kv_val,
d_qk_val,
d_v_val,
topk_val,
sm_scale_val,
cuda_stream);
params.attn_sink = static_cast<float*>(attn_sink.data_ptr());
params.topk_length = static_cast<int*>(topk_length.data_ptr());
_run_q8kv8(params, true, true);
}
} // namespace
@@ -0,0 +1,100 @@
/* Copyright 2025 SGLang Team. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#pragma once
#include <cute/tensor.hpp>
#include <cutlass/arch/barrier.h>
namespace sm90 {
__forceinline__ __device__ void
cp_async_cacheglobal_l2_prefetch_256B(const void* src, void* dst, bool pred, int64_t cache_policy) {
uint32_t dst_addr = cute::cast_smem_ptr_to_uint(dst);
asm volatile(
"cp.async.cg.shared.global.L2::cache_hint.L2::256B [%0], [%1], 16, %2, %3;\n" ::"r"(dst_addr),
"l"(src),
"r"(pred ? 16 : 0),
"l"(cache_policy));
}
__forceinline__ __device__ int64_t createpolicy_evict_last() {
int64_t res;
asm volatile("createpolicy.fractional.L2::evict_last.b64 %0, 1.0; \n\t" : "=l"(res) :);
return res;
}
__forceinline__ __device__ int64_t createpolicy_evict_first() {
int64_t res;
asm volatile("createpolicy.fractional.L2::evict_first.b64 %0, 1.0; \n\t" : "=l"(res) :);
return res;
}
__forceinline__ __device__ int get_AorC_row_idx(int local_row_idx, int idx_in_warpgroup) {
// In the layout of fragment A and fragment C during WGMMA, the data each thread holds resides in two particular rows.
// This function converts the local_row_idx (0~2) to the actual row_idx You may refer to this link for the detailed
// layout: https://docs.nvidia.com/cuda/parallel-thread-execution/#wgmma-64n16-a
int row_idx = (idx_in_warpgroup / 32) * 16 + local_row_idx * 8 + (idx_in_warpgroup % 32 / 4);
return row_idx;
}
// A simpler version of gemm
template <typename Tensor0, typename Tensor1, typename Tensor2, typename TiledMma>
__forceinline__ __device__ void gemm_ss(
bool clear_accum,
TiledMma tiled_mma,
Tensor0 const& sA,
Tensor1 const& sB,
Tensor2& rC_frag,
int idx_in_warpgroup) {
using namespace cute;
ThrMMA thr_mma = tiled_mma.get_slice(idx_in_warpgroup);
Tensor sA_frag = thr_mma.partition_fragment_A(sA);
Tensor sB_frag = thr_mma.partition_fragment_B(sB);
static_assert(size<2>(sA_frag) == size<2>(sB_frag));
warpgroup_fence_operand(rC_frag);
warpgroup_arrive();
tiled_mma.accumulate_ = clear_accum ? GMMA::ScaleOut::Zero : GMMA::ScaleOut::One;
CUTLASS_PRAGMA_UNROLL
for (int k = 0; k < size<2>(sA_frag); ++k) {
cute::gemm(tiled_mma, sA_frag(_, _, k), sB_frag(_, _, k), rC_frag);
tiled_mma.accumulate_ = GMMA::ScaleOut::One;
}
warpgroup_fence_operand(rC_frag);
}
template <typename Tensor0, typename Tensor1, typename Tensor2, typename TiledMma>
__forceinline__ __device__ void gemm_rs(
bool clear_accum, TiledMma tiled_mma, Tensor0 rA_frag, Tensor1 const& sB, Tensor2& rC_frag, int idx_in_warpgroup) {
using namespace cute;
ThrMMA thr_mma = tiled_mma.get_slice(idx_in_warpgroup);
Tensor sB_frag = thr_mma.partition_fragment_B(sB);
static_assert(size<2>(rA_frag) == size<2>(sB_frag));
warpgroup_fence_operand(const_cast<Tensor0&>(rA_frag));
warpgroup_fence_operand(rC_frag);
warpgroup_arrive();
tiled_mma.accumulate_ = clear_accum ? GMMA::ScaleOut::Zero : GMMA::ScaleOut::One;
CUTLASS_PRAGMA_UNROLL
for (int k = 0; k < size<2>(rA_frag); ++k) {
cute::gemm(tiled_mma, rA_frag(_, _, k), sB_frag(_, _, k), rC_frag);
tiled_mma.accumulate_ = GMMA::ScaleOut::One;
}
warpgroup_fence_operand(rC_frag);
warpgroup_fence_operand(const_cast<Tensor0&>(rA_frag));
}
} // namespace sm90
@@ -0,0 +1,968 @@
/* Copyright 2025 SGLang Team. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
// SM90 FP8 native sparse MLA prefill kernel.
//
// Algorithm inspired by DeepSeek FlashMLA
// (https://github.com/deepseek-ai/FlashMLA); the kernel itself is a
// clean-room re-implementation targeting the Q8KV8 sparse prefill path.
//
// Design: Native fp8 GMMA path
// QK GEMM: fp8 SS (E4M3 x E4M3 -> F32, k=32, 2x throughput vs bf16)
// PV GEMM: fp8 RS/SS (E4M3 x E4M3 -> F32, V physically transposed in smem)
// Producer: loads fp8 KV from gmem via cp.async.cg direct to smem, then transposes V
// Q: consumer WG0 loads fp8 Q from gmem directly to fp8 smem
#pragma once
#include "config.h"
#include "helpers.h"
#include <cuda_fp8.h>
// using namespace cute must be at global scope BEFORE including dense_fp8 headers
// (they use bare Tensor, make_tensor etc. from cute namespace)
using namespace cute;
// Include the fp8 transpose utility
#include "dense_fp8_transpose_v.h"
// Include the dense_fp8 utils for permute_Cregs_fp8, convert_layout_acc_Aregs, convert_type_out
#include "dense_fp8_utils.h"
namespace sm90 {
namespace fwd {
template <typename Kernel, typename TMAParamsT>
__global__ void sparse_mla_q8kv8_prefill_kernel(
__grid_constant__ const SparseMlaQ8Kv8PrefillParams params, __grid_constant__ const TMAParamsT tma_params);
template <int D_QK, bool HAVE_TOPK_LENGTH, bool HAVE_ATTN_SINK>
struct SparseMlaQ8Kv8PrefillKernel {
static constexpr int D_Q = D_QK;
static constexpr int D_K = D_QK;
static constexpr int D_V = 512;
static constexpr int B_H = 64;
static constexpr int B_TOPK = 64;
static constexpr int NUM_THREADS = 128 * 3;
static constexpr float MAX_INIT_VAL = -1e30f;
using fp8_t = cutlass::float_e4m3_t;
enum NamedBarriers : uint32_t {
wg0_bunch_0_ready = 0, // WG0 publishes max logits and local P buffer.
wg1_bunch_0_ready = 1, // WG1 publishes max logits and local P buffer.
vt0_left_ready = 2, // V[0] left half done (producer + WG0 arrivals).
vt0_right_ready = 3, // V[0] right half done (producer + WG1 arrivals).
sL_ready = 4, // post-loop only
warpgroup0_sync = 5, // post-loop only; reused in-loop as vt1_for_wg0
warpgroup1_sync = 6, // post-loop only; reused in-loop as vt1_for_wg1
epilogue_sync = 7, // never used as call; alias for q_load_done
// SM90: max 8 user NamedBarrier IDs (PTX 8-15).
// All 8 IDs used: 0-3 in-loop only, 4-7 temporally reused between in-loop and post-loop.
};
// Barrier ID aliases -- temporally disjoint reuse:
static constexpr uint32_t q_load_done = epilogue_sync; // pre-loop (256 arrivals)
static constexpr uint32_t vt1_for_wg0 = warpgroup0_sync; // in-loop (256 = prod+WG0)
static constexpr uint32_t vt1_for_wg1 = warpgroup1_sync; // in-loop (256 = prod+WG1)
static constexpr uint32_t s_consumed_ready = sL_ready; // in-loop (256 = WG0+WG1)
// ========================================================================
// FP8 Smem Layouts -- native fp8 in smem
// ========================================================================
// Q: fp8, K-major for QK SS GMMA A-operand
// SW64 because D_QK=576, 576/64=9 (int), 576/128=4.5 (not int)
template <int NUM_TILES>
using SmemLayoutQTiles_FP8 = decltype(coalesce(
tile_to_shape(GMMA::Layout_K_SW64_Atom<fp8_t>{}, Shape<Int<B_H>, Int<64 * NUM_TILES>>{}, Step<_1, _2>{}),
Shape<_1, _1>{}));
// K: fp8, K-major for QK SS GMMA B-operand
template <int NUM_TILES>
using SmemLayoutKTiles_FP8 = decltype(coalesce(
tile_to_shape(GMMA::Layout_K_SW64_Atom<fp8_t>{}, Shape<Int<B_TOPK>, Int<64 * NUM_TILES>>{}, Step<_1, _2>{}),
Shape<_1, _1>{}));
// Vt (transposed V): fp8, K-major for PV GMMA B-operand
// Shape: (D_V, B_TOPK) = (512, 64)
template <int NUM_TILES>
using SmemLayoutVtTiles_FP8 = decltype(coalesce(
tile_to_shape(GMMA::Layout_K_SW64_Atom<fp8_t>{}, Shape<Int<64 * NUM_TILES>, Int<B_TOPK>>{}, Step<_1, _2>{}),
Shape<_1, _1>{}));
// O: bf16 output (unchanged from q16)
template <int NUM_TILES>
using SmemLayoutOTiles = decltype(coalesce(
tile_to_shape(GMMA::Layout_K_SW128_Atom<bf16>{}, Shape<Int<B_H>, Int<64 * NUM_TILES>>{}, Step<_1, _2>{}),
Shape<_1, _1>{}));
using SmemLayoutQ = SmemLayoutQTiles_FP8<D_Q / 64>;
using SmemLayoutK = SmemLayoutKTiles_FP8<D_Q / 64>;
using SmemLayoutVt = SmemLayoutVtTiles_FP8<D_V / 64>; // (512, 64) fp8
using SmemLayoutHalfVt = SmemLayoutVtTiles_FP8<D_V / 64 / 2>; // (256, 64) fp8
using SmemLayoutO = SmemLayoutOTiles<D_V / 64>;
using SmemTransposeV = SmemTransposeFp8_64x64<B_TOPK, D_V>;
// ========================================================================
// FP8 GMMA atoms -- native E4M3, k=32
// ========================================================================
// QK: SS, both K-major, 64x64x32 -> 2x throughput vs bf16 k=16
using TiledMMA_QK = decltype(make_tiled_mma(GMMA::MMA_64x64x32_F32E4M3E4M3_SS_TN<>{}, Layout<Shape<_1, _1, _1>>{}));
// PV local: RS, fp8 P in regs x fp8 Vt in smem (K-major)
using TiledMMA_PV_LocalP =
decltype(make_tiled_mma(GMMA::MMA_64x256x32_F32E4M3E4M3_RS_TN<>{}, Layout<Shape<_1, _1, _1>>{}));
// PV remote: SS, fp8 P from sS x fp8 Vt in smem (K-major)
using TiledMMA_PV_RemoteP =
decltype(make_tiled_mma(GMMA::MMA_64x256x32_F32E4M3E4M3_SS_TN<>{}, Layout<Shape<_1, _1, _1>>{}));
// ========================================================================
// Shared Memory Plan -- native fp8
// ========================================================================
struct SharedMemoryPlan {
union {
array_aligned<fp8_t, cosize_v<SmemLayoutQ>> q; // B_H * D_Q fp8
array_aligned<bf16, cosize_v<SmemLayoutO>> o; // B_H * D_V/2 bf16
} q_o;
array_aligned<fp8_t, cosize_v<SmemLayoutK>> k[2]; // 2x K double-buffer, fp8
array_aligned<fp8_t, cosize_v<SmemLayoutVt>> vt[2]; // 2x Vt transposed buffer, fp8
array_aligned<fp8_t, 128 * 36> s[2]; // 2x S buffer, padded to a 36B row stride to avoid bank conflicts.
bool is_kv_valid[2][B_TOPK];
float2 sM[32];
float2 sL[64];
float final_max_logits[64], final_lse[64];
transac_bar_t bar_q, bar_k0_ready[2], bar_k1_ready[2], bar_is_kv_valid_ready;
transac_bar_t bar_k0_free, bar_k1_free;
// Consumers arrive after PV drains; the producer waits before reusing the Vt buffer.
// These barriers are separate from K-free so K buffers can be released earlier.
transac_bar_t bar_vt_free[2]; // bar_vt_free[0] protects Vt[0], bar_vt_free[1] protects Vt[1]
};
struct TmaParams_t {
CUtensorMap tensor_map_O;
};
// ========================================================================
// devfunc -- main kernel logic, native fp8 GMMA
// ========================================================================
template <typename TMAParamType>
static __device__ __forceinline__ void
devfunc(const SparseMlaQ8Kv8PrefillParams& params, const TMAParamType& tma_params) {
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 900)) || (defined(__CLION_IDE__) || defined(__VSCODE_IDE__))
const int q_h_idx = blockIdx.x % (params.h_q / B_H);
const int s_q_idx = blockIdx.x / (params.h_q / B_H);
const int warpgroup_idx = cutlass::canonical_warp_group_idx();
const int warp_idx = cutlass::canonical_warp_idx_sync();
const int idx_in_warpgroup = threadIdx.x % 128;
extern __shared__ char wksp_buf[];
SharedMemoryPlan& plan = *reinterpret_cast<SharedMemoryPlan*>(wksp_buf);
const float q_scale = __ldg(params.q_scale_ptr);
const float kv_scale = __ldg(params.kv_scale_ptr);
const float qk_combined_scale_div_log2 = q_scale * kv_scale * params.sm_scale_div_log2;
if (warp_idx == 0 && elect_one_sync()) {
cute::prefetch_tma_descriptor(&tma_params.tensor_map_O);
plan.bar_q.init(1);
plan.bar_k0_free.init(128);
plan.bar_k1_free.init(128);
CUTE_UNROLL
for (int i = 0; i < 2; ++i) {
plan.bar_k0_ready[i].init(128);
plan.bar_k1_ready[i].init(128);
}
plan.bar_is_kv_valid_ready.init(16);
CUTE_UNROLL
for (int i = 0; i < 2; ++i) {
// Transaction barriers for Vt buffer safety: 128 arrivals from each consumer WG.
plan.bar_vt_free[i].init(256);
}
fence_barrier_init();
}
__syncthreads();
const int topk_length = HAVE_TOPK_LENGTH ? __ldg(params.topk_length + s_q_idx) : params.topk;
const int num_topk_blocks = HAVE_TOPK_LENGTH ? ku::ceil_div(topk_length, (int)B_TOPK)
: (int)((unsigned int)params.topk / (unsigned int)B_TOPK);
// ================================================================
// Consumer WG0/WG1
// ================================================================
if (warpgroup_idx == 0 || warpgroup_idx == 1) {
cutlass::arch::warpgroup_reg_alloc<216>();
// --------------------------------------------------------
// Load Q from global fp8 -> fp8 smem (thread-based writes)
// Only WG0 loads Q, then sync with WG1 via NamedBarrier
// --------------------------------------------------------
if (warpgroup_idx == 0) {
const fp8_t* gQ = reinterpret_cast<const fp8_t*>(params.q) + s_q_idx * (int64_t)params.stride_q_s_q +
q_h_idx * B_H * (int64_t)params.stride_q_h_q;
// Vectorized Q loading via cp.async.cg (16 bytes per op)
constexpr int Q_GROUP_SIZE = 8;
constexpr int Q_NUM_GROUPS = 128 / Q_GROUP_SIZE;
constexpr int Q_ROWS_PER_GROUP = B_H / Q_NUM_GROUPS;
int q_ig = idx_in_warpgroup % Q_GROUP_SIZE;
int q_gg = idx_in_warpgroup / Q_GROUP_SIZE;
fp8_t* sQ_base = &(make_tensor(make_smem_ptr(plan.q_o.q.data()), SmemLayoutQTiles_FP8<1>{})(q_gg, q_ig * 16));
constexpr int NUM_Q_TILES = D_Q / 64;
int64_t q_cache_policy = createpolicy_evict_first();
CUTE_UNROLL
for (int lr = 0; lr < Q_ROWS_PER_GROUP; ++lr) {
CUTE_UNROLL
for (int ti = 0; ti < NUM_Q_TILES; ++ti) {
// Guard against OOB: last tile may be partial when D_Q%64!=0
bool q_pred = (ti * 64 + q_ig * 16 + 16) <= D_Q;
cp_async_cacheglobal_l2_prefetch_256B(
gQ + (q_gg + lr * Q_NUM_GROUPS) * (int64_t)params.stride_q_h_q + ti * 64 + q_ig * 16,
sQ_base + ti * (B_H * 64) + lr * Q_NUM_GROUPS * 64,
q_pred,
q_cache_policy);
}
}
asm volatile("cp.async.commit_group;\n" ::);
asm volatile("cp.async.wait_group 0;\n" ::);
}
fence_view_async_shared();
NamedBarrier::arrive_and_wait(256, q_load_done);
// --------------------------------------------------------
// Register fragments
// --------------------------------------------------------
float rM[2] = {MAX_INIT_VAL, MAX_INIT_VAL};
float rL[2] = {0.0f, 0.0f};
Tensor rO = partition_fragment_C(TiledMMA_PV_LocalP{}, Shape<Int<B_H>, Int<D_V / 2>>{});
Tensor rP = partition_fragment_C(TiledMMA_QK{}, Shape<Int<B_H>, Int<B_TOPK>>{});
cute::fill(rO, 0.0f);
// fp8 P register for local PV RS GMMA
// Use the same layout that convert_layout_acc_Aregs will produce
using rP_fp8_layout_t = decltype(flash::convert_layout_acc_Aregs<TiledMMA_PV_LocalP>(
partition_fragment_C(TiledMMA_QK{}, Shape<Int<B_H>, Int<B_TOPK>>{}).layout()));
Tensor rP_fp8_local = make_tensor<fp8_t>(rP_fp8_layout_t{});
bool cur_bar_wait_phase = 0;
struct Warpgroup0 {};
struct Warpgroup1 {};
// fp8 QK GEMM: 64-wide tiles, k=32, so 64/32=2 k-steps per tile
auto qkt_gemm_one_tile = [&](auto wg_tag, int tile_idx, bool clear_accum) {
constexpr bool IS_WG1 = std::is_same_v<decltype(wg_tag), Warpgroup1>;
TiledMMA_QK tiled_mma_QK;
Tensor sQ_tile = make_tensor(make_smem_ptr(plan.q_o.q.data() + tile_idx * B_H * 64), SmemLayoutQTiles_FP8<1>{});
Tensor sK_tile =
make_tensor(make_smem_ptr(plan.k[(int)IS_WG1].data() + tile_idx * B_TOPK * 64), SmemLayoutKTiles_FP8<1>{});
gemm_ss(clear_accum, tiled_mma_QK, sQ_tile, sK_tile, rP, idx_in_warpgroup);
};
auto mask_rP = [&](auto wg_tag) {
constexpr bool IS_WG1 = std::is_same_v<decltype(wg_tag), Warpgroup1>;
plan.bar_is_kv_valid_ready.wait(cur_bar_wait_phase);
CUTE_UNROLL
for (int row_idx = 0; row_idx < 2; ++row_idx) {
CUTE_UNROLL
for (int i = row_idx * 2; i < size(rP); i += 4) {
int col = 8 * (i / 4) + (idx_in_warpgroup % 4) * 2;
if (!plan.is_kv_valid[IS_WG1][col]) rP(i) = -INFINITY;
if (!plan.is_kv_valid[IS_WG1][col + 1]) rP(i + 1) = -INFINITY;
}
}
};
// online_softmax: compute softmax on rP (f32), then convert to fp8
auto online_softmax_and_rescale_o = [&](auto wg_tag) {
// mask_rP already waits for the validity mask.
constexpr bool IS_WG1 = std::is_same_v<decltype(wg_tag), Warpgroup1>;
const float scale = qk_combined_scale_div_log2;
float r_sM[2];
if constexpr (IS_WG1) {
*(float2*)r_sM = plan.sM[idx_in_warpgroup / 4];
}
float new_maxs[2];
CUTE_UNROLL
for (int row_idx = 0; row_idx < 2; ++row_idx) {
float cur_max = -INFINITY;
CUTE_UNROLL
for (int i = row_idx * 2; i < size(rP); i += 4) {
cur_max = max(cur_max, max(rP(i), rP(i + 1)));
}
cur_max = max(cur_max, __shfl_xor_sync(0xffffffff, cur_max, 1));
cur_max = max(cur_max, __shfl_xor_sync(0xffffffff, cur_max, 2));
cur_max *= scale;
new_maxs[row_idx] = max(IS_WG1 ? r_sM[row_idx] : rM[row_idx], cur_max);
float scale_for_o = exp2f(rM[row_idx] - new_maxs[row_idx]);
CUTE_UNROLL
for (int i = row_idx * 2; i < size(rO); i += 4) {
rO(i) *= scale_for_o;
rO(i + 1) *= scale_for_o;
}
float cur_sum = 0;
CUTE_UNROLL
for (int i = row_idx * 2; i < size(rP); i += 4) {
float p0 = exp2f(rP(i) * scale - new_maxs[row_idx]);
float p1 = exp2f(rP(i + 1) * scale - new_maxs[row_idx]);
rP(i) = p0;
rP(i + 1) = p1;
cur_sum += p0 + p1;
}
rL[row_idx] = rL[row_idx] * scale_for_o + cur_sum;
}
__syncwarp();
if (idx_in_warpgroup % 4 == 0) {
plan.sM[idx_in_warpgroup / 4] = *(float2*)new_maxs;
}
rM[0] = new_maxs[0];
rM[1] = new_maxs[1];
// Convert rP f32 (GMMA C layout) -> fp8 (RS A-operand layout)
// permute_Cregs_fp8 reorders C regs for fp8 A-operand layout
flash::permute_Cregs_fp8(rP);
// Reinterpret layout: C -> A-operand
Tensor rP_acc = make_tensor(rP.data(), flash::convert_layout_acc_Aregs<TiledMMA_PV_LocalP>(rP.layout()));
// f32 -> fp8
flash::convert_type_out(rP_acc, rP_fp8_local);
};
auto reduce_L = [&]() {
rL[0] += __shfl_xor_sync(0xffffffff, rL[0], 1);
rL[0] += __shfl_xor_sync(0xffffffff, rL[0], 2);
rL[1] += __shfl_xor_sync(0xffffffff, rL[1], 1);
rL[1] += __shfl_xor_sync(0xffffffff, rL[1], 2);
if (idx_in_warpgroup % 4 == 0) plan.sL[threadIdx.x / 4] = *(float2*)(rL);
NamedBarrier::arrive_and_wait(256, NamedBarriers::sL_ready);
float2 peer_L = plan.sL[(threadIdx.x / 4) ^ 32];
rL[0] += peer_L.x;
rL[1] += peer_L.y;
};
auto store_O = [&]() {
float scale_factors[2];
CUTE_UNROLL
for (int i = 0; i < 2; ++i) {
if constexpr (HAVE_ATTN_SINK) {
int attn_sink_idx = q_h_idx * B_H + get_AorC_row_idx(i, idx_in_warpgroup);
float attn_sink = __ldg(params.attn_sink + attn_sink_idx) * CUDART_L2E_F;
scale_factors[i] = kv_scale / (rL[i] + exp2f(attn_sink - rM[i]));
} else {
scale_factors[i] = kv_scale / rL[i];
}
if (rL[i] == 0.0f) scale_factors[i] = 0.0f;
}
Tensor sO_tile =
make_tensor(make_smem_ptr(plan.q_o.o.data() + warpgroup_idx * B_H * (D_V / 2)), SmemLayoutOTiles<4>{});
bf16* stsm_addrs[4];
int stsm_row = (idx_in_warpgroup / 32) * 16 + (idx_in_warpgroup % 16);
CUTE_UNROLL
for (int i = 0; i < 64 / 16; ++i) {
stsm_addrs[i] = &sO_tile(stsm_row, (idx_in_warpgroup % 32 / 16 * 8) + 16 * i);
}
bool s2g_pred = idx_in_warpgroup == 0;
warpgroup_wait<0>();
warpgroup_fence_operand(rO);
CUTE_UNROLL
for (int tile_idx = 0; tile_idx < (D_V / 2) / 64; tile_idx += 1) {
constexpr int NUM_ELEMS_EACH_TILE = B_H * 64 / 128;
bf16 cur_rOb[NUM_ELEMS_EACH_TILE];
CUTE_UNROLL
for (int i = 0; i < NUM_ELEMS_EACH_TILE; ++i) {
float out_value = rO(tile_idx * NUM_ELEMS_EACH_TILE + i) * scale_factors[i % 4 >= 2];
cur_rOb[i] = (bf16)out_value;
}
CUTE_UNROLL
for (int i = 0; i < 64 / 16; ++i) {
SM90_U32x4_STSM_N::copy(
*reinterpret_cast<uint32_t*>(cur_rOb + i * 8 + 0),
*reinterpret_cast<uint32_t*>(cur_rOb + i * 8 + 2),
*reinterpret_cast<uint32_t*>(cur_rOb + i * 8 + 4),
*reinterpret_cast<uint32_t*>(cur_rOb + i * 8 + 6),
*reinterpret_cast<uint128_t*>(stsm_addrs[i] + tile_idx * (B_H * 64)));
}
// Make the STSM writes visible to the subsequent TMA store proxy.
cute::tma_store_fence();
NamedBarrier::arrive_and_wait(
128, warpgroup_idx ? NamedBarriers::warpgroup1_sync : NamedBarriers::warpgroup0_sync);
if (s2g_pred) {
int g_tile_idx = warpgroup_idx * 4 + tile_idx;
SM90_TMA_STORE_3D::copy(
&tma_params.tensor_map_O,
plan.q_o.o.data() + g_tile_idx * (B_H * 64),
g_tile_idx * 64,
q_h_idx * B_H,
s_q_idx);
}
}
cute::tma_store_arrive();
};
// Save/load P regs to/from smem using a flat thread-indexed layout.
// Each thread writes/reads its 32 fp8 values at a unique offset.
// This preserves the RS A-reg ordering exactly, so the reader can
// load back and use RS GMMA directly (no SS GMMA layout issues).
constexpr int kP_per_thread = 32; // ((4,2,2),1,2) = 32 fp8 per thread
// Pad stride from 32 to 36 bytes to avoid shared-memory bank conflicts.
// Stride 32B = only 4 banks (8-way conflict). Stride 36B = 9 banks
// (gcd(9,32)=1 -> zero conflicts: every warp thread hits a unique bank).
constexpr int kP_stride = 36;
auto save_rP_fp8_to_sS = [&](fp8_t* sS_data) {
uint32_t* dst = reinterpret_cast<uint32_t*>(sS_data + idx_in_warpgroup * kP_stride);
uint32_t* src = reinterpret_cast<uint32_t*>(&rP_fp8_local(0));
CUTE_UNROLL
for (int i = 0; i < kP_per_thread / 4; i++) {
dst[i] = src[i];
}
};
auto load_sS_to_rP = [&](fp8_t* sS_data) {
uint32_t* src = reinterpret_cast<uint32_t*>(sS_data + idx_in_warpgroup * kP_stride);
uint32_t* dst = reinterpret_cast<uint32_t*>(&rP_fp8_local(0));
CUTE_UNROLL
for (int i = 0; i < kP_per_thread / 4; i++) {
dst[i] = src[i];
}
};
auto undo_v_transpose_col_permutation = [&]() {
// Undo the column permutation from the fp8 V transpose before writing O.
// CLayout_64x256: col bit0 = t1_bit0 (thread), col bit3 = v1 (register).
// V transpose introduces bit0<->bit3 swap. Fix by cross-thread exchange:
// thread with t1_bit0=0, v1=1 <-> thread with t1_bit0=1, v1=0
// Within each 4-element group (same v2=row): idx%4 in {0,1} are v1=0, {2,3} are v1=1.
int t1_bit0 = (threadIdx.x >> 2) & 1;
#pragma unroll
for (int g = 0; g < 32; g++) {
float a = rO(4 * g + 0);
float b = rO(4 * g + 1);
float c = rO(4 * g + 2);
float d = rO(4 * g + 3);
float send0 = t1_bit0 ? a : c;
float send1 = t1_bit0 ? b : d;
float recv0 = __shfl_xor_sync(0xFFFFFFFF, send0, 4);
float recv1 = __shfl_xor_sync(0xFFFFFFFF, send1, 4);
if (t1_bit0 == 0) {
rO(4 * g + 2) = recv0;
rO(4 * g + 3) = recv1;
} else {
rO(4 * g + 0) = recv0;
rO(4 * g + 1) = recv1;
}
}
};
// ============================================================
// WG0 Pipeline -- native fp8
// ============================================================
if (warpgroup_idx == 0) {
auto pipelined_wait_and_qkt_gemm_l = [&]() __attribute__((always_inline)) {
plan.bar_k0_ready[0].wait(cur_bar_wait_phase);
qkt_gemm_one_tile(Warpgroup0{}, 0, true);
qkt_gemm_one_tile(Warpgroup0{}, 1, false);
qkt_gemm_one_tile(Warpgroup0{}, 2, false);
qkt_gemm_one_tile(Warpgroup0{}, 3, false);
warpgroup_commit_batch();
};
auto pipelined_wait_and_qkt_gemm_r = [&]() __attribute__((always_inline)) {
plan.bar_k0_ready[1].wait(cur_bar_wait_phase);
qkt_gemm_one_tile(Warpgroup0{}, 4, false);
qkt_gemm_one_tile(Warpgroup0{}, 5, false);
qkt_gemm_one_tile(Warpgroup0{}, 6, false);
qkt_gemm_one_tile(Warpgroup0{}, 7, false);
if constexpr (D_QK == 576) {
qkt_gemm_one_tile(Warpgroup0{}, 8, false);
}
warpgroup_commit_batch();
};
auto rescale_rO = [&](float scales[2]) {
CUTE_UNROLL
for (int row = 0; row < 2; ++row) {
CUTE_UNROLL
for (int i = row * 2; i < size(rO); i += 4) {
rO(i) *= scales[row];
rO(i + 1) *= scales[row];
}
rL[row] *= scales[row];
}
};
CUTE_NO_UNROLL
for (int block_idx = 0; block_idx < num_topk_blocks; block_idx += 2) {
// Vt[0] left half: (256, 64) fp8 -- only half we transpose & use
Tensor sVt0l = make_tensor(make_smem_ptr(plan.vt[0].data()), SmemLayoutHalfVt{});
if (block_idx == 0) {
pipelined_wait_and_qkt_gemm_l();
pipelined_wait_and_qkt_gemm_r();
warpgroup_wait<0>();
warpgroup_fence_operand(rP);
plan.bar_k0_free.arrive();
}
mask_rP(Warpgroup0{});
online_softmax_and_rescale_o(Warpgroup0{});
save_rP_fp8_to_sS(plan.s[0].data());
NamedBarrier::arrive(256, NamedBarriers::wg0_bunch_0_ready);
// Wait for Vt[0] left half only (producer + WG0 arrivals).
// V[0]-RIGHT may still be transposing; WG0 doesn't need it.
NamedBarrier::arrive_and_wait(256, vt0_left_ready);
// Local PV: rP_fp8 x Vt0 left half -> RS fp8 GMMA
gemm_rs(false, TiledMMA_PV_LocalP{}, rP_fp8_local, sVt0l, rO, idx_in_warpgroup);
warpgroup_commit_batch();
// Overlap PV-local GMMA drain with barrier waits, sM read, and peer P load.
NamedBarrier::arrive_and_wait(256, NamedBarriers::wg1_bunch_0_ready);
float new_rM[2], scale_factors_arr[2];
*(float2*)new_rM = plan.sM[idx_in_warpgroup / 4];
CUTE_UNROLL
for (int i = 0; i < 2; ++i) {
scale_factors_arr[i] = exp2f(rM[i] - new_rM[i]);
rM[i] = new_rM[i];
}
warpgroup_wait<0>();
warpgroup_fence_operand(rO);
warpgroup_fence_operand(rP_fp8_local);
plan.bar_vt_free[0].arrive();
load_sS_to_rP(plan.s[1].data());
NamedBarrier::arrive_and_wait(256, s_consumed_ready);
// Wait for Vt[1] transpose (prod+WG0 barrier)
NamedBarrier::arrive_and_wait(256, vt1_for_wg0);
// Rescale rO: must be after wait<0> since rO is PV-local accumulator
rescale_rO(scale_factors_arr);
Tensor sVt1l = make_tensor(make_smem_ptr(plan.vt[1].data()), SmemLayoutHalfVt{});
gemm_rs(false, TiledMMA_PV_LocalP{}, rP_fp8_local, sVt1l, rO, idx_in_warpgroup);
warpgroup_commit_batch();
cur_bar_wait_phase ^= 1;
if (block_idx + 2 < num_topk_blocks) {
pipelined_wait_and_qkt_gemm_l();
warpgroup_wait<1>();
warpgroup_fence_operand(rO);
warpgroup_fence_operand(rP_fp8_local);
plan.bar_vt_free[1].arrive();
pipelined_wait_and_qkt_gemm_r();
warpgroup_wait<0>();
warpgroup_fence_operand(rP);
plan.bar_k0_free.arrive();
} else {
warpgroup_wait<0>();
warpgroup_fence_operand(rO);
plan.bar_vt_free[1].arrive();
}
}
undo_v_transpose_col_permutation();
reduce_L();
store_O();
} else {
// ============================================================
// WG1 Pipeline -- native fp8
// ============================================================
// Split QK into R/L halves for loop-end overlap (mirrors WG0 pattern)
auto pipelined_wait_and_qkt_gemm_r_wg1 = [&]() __attribute__((always_inline)) {
// Right half first: K[1]-right arrives earlier from producer
plan.bar_k1_ready[1].wait(cur_bar_wait_phase);
qkt_gemm_one_tile(Warpgroup1{}, 4, true);
qkt_gemm_one_tile(Warpgroup1{}, 5, false);
qkt_gemm_one_tile(Warpgroup1{}, 6, false);
qkt_gemm_one_tile(Warpgroup1{}, 7, false);
if constexpr (D_QK == 576) {
qkt_gemm_one_tile(Warpgroup1{}, 8, false);
}
warpgroup_commit_batch();
};
auto pipelined_wait_and_qkt_gemm_l_wg1 = [&]() __attribute__((always_inline)) {
plan.bar_k1_ready[0].wait(cur_bar_wait_phase);
qkt_gemm_one_tile(Warpgroup1{}, 0, false);
qkt_gemm_one_tile(Warpgroup1{}, 1, false);
qkt_gemm_one_tile(Warpgroup1{}, 2, false);
qkt_gemm_one_tile(Warpgroup1{}, 3, false);
warpgroup_commit_batch();
};
CUTE_NO_UNROLL
for (int block_idx = 0; block_idx < num_topk_blocks; block_idx += 2) {
// Vt[1] right half: (256, 64) fp8 -- only half we transpose & use
Tensor sVt1r = make_tensor(make_smem_ptr(plan.vt[1].data() + 256 * B_TOPK), SmemLayoutHalfVt{});
if (block_idx == 0) {
pipelined_wait_and_qkt_gemm_r_wg1();
pipelined_wait_and_qkt_gemm_l_wg1();
warpgroup_wait<0>();
warpgroup_fence_operand(rP);
plan.bar_k1_free.arrive();
}
mask_rP(Warpgroup1{});
NamedBarrier::arrive_and_wait(256, NamedBarriers::wg0_bunch_0_ready);
online_softmax_and_rescale_o(Warpgroup1{});
save_rP_fp8_to_sS(plan.s[1].data());
NamedBarrier::arrive(256, NamedBarriers::wg1_bunch_0_ready);
// Wait for Vt[1] transpose (prod+WG1 barrier)
NamedBarrier::arrive_and_wait(256, vt1_for_wg1);
// Local PV: rP_fp8 x Vt1 right half -> RS
gemm_rs(false, TiledMMA_PV_LocalP{}, rP_fp8_local, sVt1r, rO, idx_in_warpgroup);
warpgroup_commit_batch();
warpgroup_wait<0>();
warpgroup_fence_operand(rO);
warpgroup_fence_operand(rP_fp8_local);
plan.bar_vt_free[1].arrive();
load_sS_to_rP(plan.s[0].data());
NamedBarrier::arrive_and_wait(256, s_consumed_ready);
// Wait for Vt[0] right half only (producer + WG1 arrivals).
// V[0]-LEFT was signaled earlier; WG1 doesn't need it.
NamedBarrier::arrive_and_wait(256, vt0_right_ready);
Tensor sVt0r = make_tensor(make_smem_ptr(plan.vt[0].data() + 256 * B_TOPK), SmemLayoutHalfVt{});
gemm_rs(false, TiledMMA_PV_LocalP{}, rP_fp8_local, sVt0r, rO, idx_in_warpgroup);
warpgroup_commit_batch();
if (block_idx + 2 < num_topk_blocks) {
cur_bar_wait_phase ^= 1;
// Overlap: start next-iteration QK-right while PV drains
pipelined_wait_and_qkt_gemm_r_wg1();
warpgroup_wait<1>();
warpgroup_fence_operand(rO);
warpgroup_fence_operand(rP_fp8_local);
plan.bar_vt_free[0].arrive();
pipelined_wait_and_qkt_gemm_l_wg1();
warpgroup_wait<0>();
warpgroup_fence_operand(rP);
plan.bar_k1_free.arrive();
} else {
warpgroup_wait<0>();
warpgroup_fence_operand(rO);
plan.bar_vt_free[0].arrive();
}
}
undo_v_transpose_col_permutation();
reduce_L();
store_O();
if (idx_in_warpgroup % 4 == 0) {
for (int row = 0; row < 2; ++row) {
int real_row = get_AorC_row_idx(row, idx_in_warpgroup);
bool is_no_valid_tokens = rL[row] == 0.0f;
plan.final_max_logits[real_row] = is_no_valid_tokens ? -INFINITY : rM[row] * CUDART_LN2_F;
plan.final_lse[real_row] = is_no_valid_tokens ? +INFINITY : logf(rL[row]) + rM[row] * CUDART_LN2_F;
}
// Regular stores are not async-proxy operations; the barrier provides ordering.
asm volatile("" ::: "memory");
}
NamedBarrier::arrive_and_wait(128, NamedBarriers::warpgroup1_sync);
if (idx_in_warpgroup == 0) {
int g_offset = s_q_idx * params.h_q + q_h_idx * B_H;
SM90_BULK_COPY_S2G::copy(plan.final_max_logits, params.max_logits + g_offset, B_H * sizeof(float));
SM90_BULK_COPY_S2G::copy(plan.final_lse, params.lse + g_offset, B_H * sizeof(float));
cute::tma_store_arrive();
}
}
} else {
// ================================================================
// Producer WG2: load fp8 KV via cp.async, then transpose V in smem
// ================================================================
cutlass::arch::warpgroup_reg_dealloc<72>();
constexpr int GROUP_SIZE = 8, NUM_GROUPS = 128 / GROUP_SIZE;
constexpr int NUM_ROWS_PER_GROUP = B_TOPK / NUM_GROUPS;
int idx_in_group = idx_in_warpgroup % GROUP_SIZE;
int group_idx = idx_in_warpgroup / GROUP_SIZE;
int* gIndices = params.indices + s_q_idx * params.stride_indices_s_q;
int tile_shift = idx_in_group / 4;
int col_in_tile = (idx_in_group % 4) * 16;
fp8_t* my_sK_base =
&(make_tensor(make_smem_ptr(plan.k[0].data()), SmemLayoutKTiles_FP8<1>{})(group_idx, col_in_tile)) +
tile_shift * (B_TOPK * 64);
const fp8_t* my_gKV_base = reinterpret_cast<const fp8_t*>(params.kv) + idx_in_group * 16;
int64_t token_indices[2][NUM_ROWS_PER_GROUP];
bool is_token_valid[2][NUM_ROWS_PER_GROUP];
auto load_token_indices = [&](int block_idx) {
CUTE_UNROLL
for (int buf_idx = 0; buf_idx < 2; ++buf_idx) {
CUTE_UNROLL
for (int local_row = 0; local_row < NUM_ROWS_PER_GROUP; ++local_row) {
int offs = (block_idx + buf_idx) * B_TOPK + local_row * NUM_GROUPS + group_idx;
int t = __ldg(gIndices + offs);
bool is_cur_token_valid = t >= 0 && t < params.s_kv;
if constexpr (HAVE_TOPK_LENGTH) {
is_cur_token_valid &= offs < topk_length;
}
token_indices[buf_idx][local_row] = (int64_t)t * (int64_t)params.stride_kv_s_kv;
is_token_valid[buf_idx][local_row] = is_cur_token_valid;
}
}
};
int64_t cache_policy = createpolicy_evict_last();
auto copy_tiles = [&](int buf_idx, int smem_buf, int tile_start, int tile_end) {
CUTE_UNROLL
for (int local_row = 0; local_row < NUM_ROWS_PER_GROUP; ++local_row) {
int64_t token_index = token_indices[buf_idx][local_row];
CUTE_UNROLL
for (int tile_idx = tile_start; tile_idx < tile_end; tile_idx += 2) {
int phys_tile = tile_idx + tile_shift;
if constexpr ((D_K % 128) != 0) {
if (phys_tile >= (D_K / 64)) continue;
}
bool kv_pred = is_token_valid[buf_idx][local_row] && phys_tile < (D_K / 64);
cp_async_cacheglobal_l2_prefetch_256B(
my_gKV_base + token_index + tile_idx * 64,
my_sK_base +
(smem_buf * cosize_v<SmemLayoutK> + tile_idx * (B_TOPK * 64) + local_row * NUM_GROUPS * 64),
kv_pred,
cache_policy);
}
}
};
auto commit_to_mbar = [&](transac_bar_t& bar) { cutlass::arch::cpasync_barrier_arrive_noinc((uint64_t*)(&bar)); };
// V transpose helper instance
SmemTransposeV smem_transpose_v;
using SmemLayoutTransposeV_t = typename SmemTransposeV::SmemLayoutTransposeV;
using SmemLayoutTransposeVt_t = typename SmemTransposeV::SmemLayoutTransposeVt;
// Use the FA3-style STSM thread layout for the fp8 V transpose.
// but same composition-based framework as before.
auto transpose_v_half = [&](int smem_k_buf, int vt_buf, int tile_start, int tile_end) {
Tensor sV_src = as_position_independent_swizzle_tensor(
make_tensor(make_smem_ptr(plan.k[smem_k_buf].data()), SmemLayoutTransposeV_t{}));
Tensor sVt_dst = as_position_independent_swizzle_tensor(
make_tensor(make_smem_ptr(plan.vt[vt_buf].data()), SmemLayoutTransposeVt_t{}));
static_assert((D_V / 64 / 2) % 2 == 0, "half tile count must be even for pair transpose");
CUTE_UNROLL
for (int j = tile_start; j < tile_end; j += 2) {
smem_transpose_v.transpose_pair(
flatten(sV_src(_, 0, j)),
flatten(sVt_dst(_, 0, j)),
flatten(sV_src(_, 0, j + 1)),
flatten(sVt_dst(_, 0, j + 1)));
}
asm volatile("" ::: "memory");
};
int cur_bar_wait_phase_prod = 1;
// Prologue: prefetch the first iteration's indices before the loop.
// Subsequent iterations' indices are prefetched during V transpose
// of the prior iteration, hiding __ldg latency behind compute.
load_token_indices(0);
CUTE_NO_UNROLL
for (int block_idx = 0; block_idx < num_topk_blocks; block_idx += 2) {
// Indices are already loaded by the prologue or the previous iteration's prefetch.
plan.bar_k0_free.wait(cur_bar_wait_phase_prod);
plan.bar_k1_free.wait(cur_bar_wait_phase_prod);
// is_kv_valid write: AFTER bar_k_free waits to avoid race condition.
// Consumers may still be reading prior iteration's is_kv_valid during
// mask_rP until they signal k_free. Writing before waits could overwrite
// values consumers are still reading.
if (idx_in_group == 0) {
CUTE_UNROLL
for (int buf_idx = 0; buf_idx < 2; ++buf_idx)
CUTE_UNROLL
for (int local_row = 0; local_row < NUM_ROWS_PER_GROUP; ++local_row)
plan.is_kv_valid[buf_idx][local_row * NUM_GROUPS + group_idx] = is_token_valid[buf_idx][local_row];
plan.bar_is_kv_valid_ready.arrive();
}
copy_tiles(0, 0, 0, 4);
commit_to_mbar(plan.bar_k0_ready[0]);
asm volatile("cp.async.commit_group;\n" ::);
constexpr int kv_tile_end = D_K / 64;
copy_tiles(1, 1, 4, kv_tile_end);
commit_to_mbar(plan.bar_k1_ready[1]);
copy_tiles(0, 0, 4, kv_tile_end);
commit_to_mbar(plan.bar_k0_ready[1]);
copy_tiles(1, 1, 0, 4);
commit_to_mbar(plan.bar_k1_ready[0]);
asm volatile("cp.async.commit_group;\n" ::);
// Wait for K[0]-left (group-0)
asm volatile("cp.async.wait_group 1;\n" ::);
// fence.proxy.async: make cp.async data visible through generic proxy
// (required for LDSM reads in V transpose; cp.async uses async proxy)
fence_view_async_shared();
asm volatile("bar.sync 7, 128;\n" ::: "memory");
if (block_idx > 0) {
plan.bar_vt_free[0].wait(cur_bar_wait_phase_prod);
}
// Prefetch next iteration's indices before V[0]-LEFT transpose so the
// __ldg latency is hidden behind the full V transpose window.
if (block_idx + 2 < num_topk_blocks) {
load_token_indices(block_idx + 2);
}
transpose_v_half(0, 0, 0, 4);
NamedBarrier::arrive(256, vt0_left_ready);
// Transpose V[1] left before V[0] right to match the consumer handoff order.
// WG0 is on the critical path (feeds WG1 via sM/wg0_bunch).
// WG0 waits for vt1_for_wg0 (V[1]-LEFT) for PV-remote.
// Moving V[1]-LEFT earlier (2nd instead of 4th) reduces WG0
// critical-path stall by ~768 cycles per iteration.
//
// v52 CRASH FIX: K[1]-left tiles 0-3 are in cp.async group-1,
// NOT group-0. wait_group 1 only waits for group-0. Under high
// CTA counts (512+), memory bandwidth saturation delays group-1
// completion past the V[0]-LEFT transpose timing margin, causing
// the V[1]-LEFT transpose to read stale/partial smem data.
// Fix: wait_group 0 before V[1]-LEFT ensures group-1 has completed.
// V[0]-LEFT transpose still overlaps with group-1 async copies.
// Wait for all groups before V[1]-LEFT transpose
asm volatile("cp.async.wait_group 0;\n" ::);
// fence.proxy.async: make cp.async group-1 data visible through
// generic proxy for LDSM reads in V transpose
fence_view_async_shared();
asm volatile("bar.sync 7, 128;\n" ::: "memory");
// V[1]-LEFT: tiles 0-3 from K[1] -- WG0 needs this for PV-remote
if (block_idx > 0) {
plan.bar_vt_free[1].wait(cur_bar_wait_phase_prod);
}
transpose_v_half(1, 1, 0, 4);
NamedBarrier::arrive(256, vt1_for_wg0);
// V[0]-RIGHT: tiles 4-7 from K[0]
transpose_v_half(0, 0, 4, 8);
NamedBarrier::arrive(256, vt0_right_ready);
// V[1]-RIGHT: tiles 4-7 from K[1]
transpose_v_half(1, 1, 4, 8);
NamedBarrier::arrive(256, vt1_for_wg1);
asm volatile("bar.sync 7, 128;\n" ::: "memory");
cur_bar_wait_phase_prod ^= 1;
}
}
cute::tma_store_wait<0>();
#else
if (cute::thread0()) {
CUTE_INVALID_CONTROL_PATH("This kernel only supports sm90");
}
#endif
}
// ========================================================================
// run() -- host-side launch
// ========================================================================
static void run(const SparseMlaQ8Kv8PrefillParams& params) {
KU_ASSERT(params.h_kv == 1);
KU_ASSERT(params.topk % (2 * B_TOPK) == 0);
KU_ASSERT(params.topk > 0);
KU_ASSERT(params.h_q % B_H == 0);
CUtensorMap tensor_map_O;
{
uint64_t size[3] = {(uint64_t)D_V, (uint64_t)params.h_q, (uint64_t)params.s_q};
uint64_t stride[2] = {D_V * sizeof(bf16), D_V * params.h_q * sizeof(bf16)};
uint32_t box_size[3] = {64, B_H, 1};
uint32_t elem_stride[3] = {1, 1, 1};
CUresult res = CUTLASS_CUDA_DRIVER_WRAPPER_CALL(cuTensorMapEncodeTiled)(
&tensor_map_O,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_BFLOAT16,
3,
params.out,
size,
stride,
box_size,
elem_stride,
CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B,
CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,
CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
KU_ASSERT(res == CUresult::CUDA_SUCCESS);
}
TmaParams_t tma_p = {tensor_map_O};
auto kernel = &sparse_mla_q8kv8_prefill_kernel<
SparseMlaQ8Kv8PrefillKernel<D_QK, HAVE_TOPK_LENGTH, HAVE_ATTN_SINK>,
TmaParams_t>;
constexpr size_t smem_size = sizeof(SharedMemoryPlan);
KU_CUDA_CHECK(cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
cutlass::ClusterLaunchParams launch_params = {
dim3((params.h_q / B_H) * params.s_q, 1, 1), dim3(NUM_THREADS, 1, 1), dim3(1, 1, 1), smem_size, params.stream};
cutlass::launch_kernel_on_cluster(launch_params, (void*)kernel, params, tma_p);
KU_CHECK_KERNEL_LAUNCH();
}
};
// ============================================================================
// Global kernel entry point
// ============================================================================
template <typename Kernel, typename TMAParamsT>
__global__ void __launch_bounds__(Kernel::NUM_THREADS, 1, 1) sparse_mla_q8kv8_prefill_kernel(
__grid_constant__ const SparseMlaQ8Kv8PrefillParams params, __grid_constant__ const TMAParamsT tma_params) {
Kernel::devfunc(params, tma_params);
}
// ============================================================================
// External dispatch function
// ============================================================================
template <int D_QK, bool HAVE_TOPK_LENGTH, bool HAVE_ATTN_SINK>
void run_sparse_mla_q8kv8_prefill_kernel(const SparseMlaQ8Kv8PrefillParams& params) {
SparseMlaQ8Kv8PrefillKernel<D_QK, HAVE_TOPK_LENGTH, HAVE_ATTN_SINK>::run(params);
}
} // namespace fwd
} // namespace sm90
@@ -0,0 +1,47 @@
/* Copyright 2025 SGLang Team. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#pragma once
#include "cutlass/bfloat16.h"
#include <cstdint>
#include <cuda_runtime.h>
struct SparseMlaQ8Kv8PrefillParams {
int s_q, s_kv, h_q, h_kv, d_qk, d_v, topk;
float sm_scale_div_log2;
const uint8_t* __restrict__ q;
const uint8_t* __restrict__ kv;
int* __restrict__ indices;
float* __restrict__ attn_sink;
int* __restrict__ topk_length;
const float* __restrict__ q_scale_ptr;
const float* __restrict__ kv_scale_ptr;
int stride_q_s_q;
int stride_q_h_q;
int64_t stride_kv_s_kv;
int stride_kv_h_kv;
int stride_indices_s_q;
int stride_indices_h_kv;
cutlass::bfloat16_t* __restrict__ out;
float* __restrict__ max_logits;
float* __restrict__ lse;
cudaStream_t stream;
};
@@ -0,0 +1,329 @@
"""JIT-compiled Q8KV8 sparse prefill attention kernel for SM90 (Hopper/H200).
Uses native FP8 GMMA instructions via CUTLASS/CUTE for MLA attention
with FP8 quantized Q and KV tensors.
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Optional
import torch
from sglang.jit_kernel.utils import cache_once, load_jit, override_jit_cuda_arch
from sglang.kernel_api_logging import debug_kernel_api
from sglang.srt.utils.custom_op import register_custom_op
if TYPE_CHECKING:
from tvm_ffi.module import Module
# ---------------------------------------------------------------------------
# Build flags
# ---------------------------------------------------------------------------
def _q8kv8_cuda_flags() -> list[str]:
# Minimal flag set, verified by per-flag ablation on SM90/H200 (CUDA 12.9).
# The original list was lifted from DeepSeek FlashMLA's AOT setup.py; under
# this tvm_ffi JIT build only --use_fast_math has any measurable effect, so
# the rest are dropped.
#
# --use_fast_math maps the softmax exp2f to the ex2.approx.f32 MUFU op. Cost
# of removing it: ~+4.3% at short-context / large-topk (s_kv=8192,
# topk=2048), ~+1-2% mid, ~0% at long context -- with no accuracy change
# (its ~2^-22 relative error is far below the fp8-e4m3 quantization noise).
#
# Dropped, all confirmed to leave perf and accuracy bit-identical here:
# * -U__CUDA_NO_HALF*/__CUDA_NO_BFLOAT16_CONVERSIONS__: these only matter
# when the toolchain pre-defines the matching -D__CUDA_NO_* macros, as
# torch.utils.cpp_extension's AOT path does (COMMON_NVCC_FLAGS). The JIT
# toolchain never defines them, so undefining is a no-op.
# * --expt-relaxed-constexpr and -O3: already supplied by the JIT default
# target flags (see utils._get_default_target_flags).
# * --expt-extended-lambda, -lineinfo, -D_USE_MATH_DEFINES: not required
# by this single-translation-unit kernel.
return [
"-O3",
"-DNDEBUG",
"-DCUTE_USE_PACKED_TUPLE=1",
"-DCUTLASS_ENABLE_TENSOR_CORE_MMA=1",
"--use_fast_math",
]
# ---------------------------------------------------------------------------
# Module loader
# ---------------------------------------------------------------------------
@cache_once
def _jit_sparse_mla_q8kv8_prefill_module() -> Module:
with override_jit_cuda_arch(9, 0, "a"):
return load_jit(
"sparse_mla_q8kv8_prefill_sm90",
cuda_files=[
"sparse_mla_q8kv8_prefill_sm90/entry.cuh",
],
cuda_wrappers=[
("dispatch", "sparse_prefill_q8kv8_dispatch"),
("dispatch_full", "sparse_prefill_q8kv8_dispatch_full"),
],
extra_cuda_cflags=_q8kv8_cuda_flags(),
extra_dependencies=["cutlass"],
)
# Pre-resolve entry-point callables on first use to avoid per-call module
# dictionary lookups.
_resolved_entries: Optional[tuple] = None
def _get_entries() -> tuple:
global _resolved_entries
if _resolved_entries is None:
m = _jit_sparse_mla_q8kv8_prefill_module()
_resolved_entries = (
m["dispatch"],
m["dispatch_full"],
)
return _resolved_entries
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
# torch._C._cuda_getCurrentRawStream returns the cudaStream_t pointer expected
# by the JIT wrapper. torch._C._cuda_getCurrentStream returns a packed stream
# id and must not be used here.
_get_current_stream_raw = torch._C._cuda_getCurrentRawStream
# Module-level cache for kernel-write-only output tensors. The active s_q rows
# are overwritten every call; buffers grow monotonically by device/head shape.
def _check_out_buffer(
t: torch.Tensor,
name: str,
shape: tuple,
dtype: torch.dtype,
device: torch.device,
) -> None:
if tuple(t.shape) != tuple(shape):
raise ValueError(f"{name} must have shape {tuple(shape)}, got {tuple(t.shape)}")
if t.dtype != dtype:
raise ValueError(f"{name} must have dtype {dtype}, got {t.dtype}")
if t.device != device:
raise ValueError(f"{name} must be on device {device}, got {t.device}")
if not t.is_contiguous():
raise ValueError(f"{name} must be contiguous")
# Internal custom-op wrappers so the JIT kernel calls participate in
# torch.library / torch.compile tracing and kernel-API debug logging.
# The dispatch_full variant carries the optional attn_sink / topk_length
# tensors as required args; the public API chooses which op to call.
@register_custom_op(
op_name="sparse_mla_q8kv8_prefill",
mutates_args=["out", "max_logits", "lse"],
)
def _sparse_mla_q8kv8_prefill_op(
q: torch.Tensor,
kv: torch.Tensor,
indices: torch.Tensor,
q_scale: torch.Tensor,
kv_scale: torch.Tensor,
out: torch.Tensor,
max_logits: torch.Tensor,
lse: torch.Tensor,
s_q: int,
s_kv: int,
h_q: int,
h_kv: int,
d_qk: int,
d_v: int,
topk: int,
sm_scale: float,
cuda_stream: int,
) -> None:
dispatch_fn, _ = _get_entries()
dispatch_fn(
q,
kv,
indices,
q_scale,
kv_scale,
out,
max_logits,
lse,
s_q,
s_kv,
h_q,
h_kv,
d_qk,
d_v,
topk,
sm_scale,
cuda_stream,
)
@register_custom_op(
op_name="sparse_mla_q8kv8_prefill_full",
mutates_args=["out", "max_logits", "lse"],
)
def _sparse_mla_q8kv8_prefill_full_op(
q: torch.Tensor,
kv: torch.Tensor,
indices: torch.Tensor,
q_scale: torch.Tensor,
kv_scale: torch.Tensor,
attn_sink: torch.Tensor,
topk_length: torch.Tensor,
out: torch.Tensor,
max_logits: torch.Tensor,
lse: torch.Tensor,
s_q: int,
s_kv: int,
h_q: int,
h_kv: int,
d_qk: int,
d_v: int,
topk: int,
sm_scale: float,
cuda_stream: int,
) -> None:
_, dispatch_full_fn = _get_entries()
dispatch_full_fn(
q,
kv,
indices,
q_scale,
kv_scale,
attn_sink,
topk_length,
out,
max_logits,
lse,
s_q,
s_kv,
h_q,
h_kv,
d_qk,
d_v,
topk,
sm_scale,
cuda_stream,
)
@debug_kernel_api
def sparse_mla_q8kv8_prefill_fwd(
q: torch.Tensor, # [s_q, h_q, d_qk], float8_e4m3fn
kv: torch.Tensor, # [s_kv, h_kv, d_qk], float8_e4m3fn
indices: torch.Tensor, # [s_q, h_kv, topk], int32
sm_scale: float,
q_scale: torch.Tensor, # scalar tensor on GPU, float32
kv_scale: torch.Tensor, # scalar tensor on GPU, float32
d_v: int = 512,
attn_sink: Optional[torch.Tensor] = None, # [h_q], float32
topk_length: Optional[torch.Tensor] = None, # [s_q], int32
*,
out: Optional[torch.Tensor] = None, # [s_q, h_q, d_v], bfloat16
max_logits: Optional[torch.Tensor] = None, # [s_q, h_q], float32
lse: Optional[torch.Tensor] = None, # [s_q, h_q], float32
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Run Q8KV8 (FP8) sparse prefill attention on SM90.
The kernel writes into three output tensors. By default fresh tensors
are allocated and returned; callers that want to reuse buffers (e.g.
for CUDA graph capture) may pass pre-allocated ``out`` / ``max_logits``
/ ``lse`` tensors of the expected shape/dtype/device. The three output
tensors must not alias each other.
Returns:
out: [s_q, h_q, d_v], bfloat16
max_logits: [s_q, h_q], float32
lse: [s_q, h_q], float32
"""
s_q, h_q, d_qk = q.shape
s_kv = kv.shape[0]
h_kv = kv.shape[1]
topk = indices.shape[2]
if d_v != 512:
raise ValueError(
f"sparse_mla_q8kv8_prefill_fwd only supports d_v=512, got {d_v}"
)
if (attn_sink is None) != (topk_length is None):
raise ValueError("attn_sink and topk_length must be provided together")
device = q.device
if out is None:
out = torch.empty(s_q, h_q, d_v, dtype=torch.bfloat16, device=device)
else:
_check_out_buffer(out, "out", (s_q, h_q, d_v), torch.bfloat16, device)
if max_logits is None:
max_logits = torch.empty(s_q, h_q, dtype=torch.float32, device=device)
else:
_check_out_buffer(max_logits, "max_logits", (s_q, h_q), torch.float32, device)
if lse is None:
lse = torch.empty(s_q, h_q, dtype=torch.float32, device=device)
else:
_check_out_buffer(lse, "lse", (s_q, h_q), torch.float32, device)
# The three output tensors are written independently by the kernel; any
# aliasing among them would corrupt results, so reject it explicitly.
out_ptr = out.data_ptr()
ml_ptr = max_logits.data_ptr()
lse_ptr = lse.data_ptr()
if out_ptr == ml_ptr or out_ptr == lse_ptr or ml_ptr == lse_ptr:
raise ValueError("out, max_logits and lse must not alias each other")
cuda_stream = _get_current_stream_raw(q.device.index)
if attn_sink is not None and topk_length is not None:
_sparse_mla_q8kv8_prefill_full_op(
q,
kv,
indices,
q_scale,
kv_scale,
attn_sink,
topk_length,
out,
max_logits,
lse,
s_q,
s_kv,
h_q,
h_kv,
d_qk,
d_v,
topk,
sm_scale,
cuda_stream,
)
else:
_sparse_mla_q8kv8_prefill_op(
q,
kv,
indices,
q_scale,
kv_scale,
out,
max_logits,
lse,
s_q,
s_kv,
h_q,
h_kv,
d_qk,
d_v,
topk,
sm_scale,
cuda_stream,
)
return out, max_logits, lse