Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion csrc/api/sparse_decode.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -80,7 +80,14 @@ class Decode_Sm90_Impl : public DecodeImplBase {
void run_(const SparseAttnDecodeParams &params, const std::vector<FeatureT> &required_features) override {
DISPATCH_MODEL_TYPE(params.model_type, MODEL_TYPE, [&]() {
DISPATCH_NUM_HEADS(params.h_q, NUM_HEADS, [&]() {
sm90::decode::sparse::run_flash_splitkv_mla_fp8_sparse_kernel<MODEL_TYPE, NUM_HEADS>(params);
if constexpr (MODEL_TYPE == ModelType::V32) {
DISPATCH_BOOLEAN_FLAG(params.topk_length != nullptr, HAVE_TOPK_LENGTH, [&]() {
sm90::decode::sparse::run_flash_splitkv_mla_fp8_sparse_kernel<MODEL_TYPE, NUM_HEADS, HAVE_TOPK_LENGTH>(params);
});
} else {
// DeepSeek-V4 always handles topk_length, so it has a single instantiation
sm90::decode::sparse::run_flash_splitkv_mla_fp8_sparse_kernel<MODEL_TYPE, NUM_HEADS, false>(params);
}
});
});
}
Expand Down
5 changes: 4 additions & 1 deletion csrc/kernels/sm90/decode/sparse/config.h
Original file line number Diff line number Diff line change
Expand Up @@ -12,11 +12,14 @@ using namespace cute;

namespace sm90::decode::sparse {

template<ModelType MODEL_TYPE, int NUM_HEADS>
template<ModelType MODEL_TYPE, int NUM_HEADS, bool HAVE_TOPK_LENGTH>
class KernelTemplate {
public:

static_assert(NUM_HEADS == 64 || NUM_HEADS == 128);
// DeepSeek-V4 always bounds the index reads by topk_length. For V3.2 it is
// compiled in only when topk_length is given, so the plain path is unchanged.
static constexpr bool CHECK_TOPK_LENGTH = MODEL_TYPE != ModelType::V32 || HAVE_TOPK_LENGTH;
static constexpr int NUM_M_BLOCKS = NUM_HEADS / 64;
static constexpr int CLUSTER_SIZE = NUM_M_BLOCKS;

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,6 @@

namespace sm90::decode::sparse {

template void run_flash_splitkv_mla_fp8_sparse_kernel<ModelType::V32, 128>(const SparseAttnDecodeParams &params);
template void run_flash_splitkv_mla_fp8_sparse_kernel<ModelType::V32, 128, false>(const SparseAttnDecodeParams &params);

}
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
#include "../splitkv_mla.cuh"

namespace sm90::decode::sparse {

template void run_flash_splitkv_mla_fp8_sparse_kernel<ModelType::V32, 128, true>(const SparseAttnDecodeParams &params);

}
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,6 @@

namespace sm90::decode::sparse {

template void run_flash_splitkv_mla_fp8_sparse_kernel<ModelType::V32, 64>(const SparseAttnDecodeParams &params);
template void run_flash_splitkv_mla_fp8_sparse_kernel<ModelType::V32, 64, false>(const SparseAttnDecodeParams &params);

}
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
#include "../splitkv_mla.cuh"

namespace sm90::decode::sparse {

template void run_flash_splitkv_mla_fp8_sparse_kernel<ModelType::V32, 64, true>(const SparseAttnDecodeParams &params);

}
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,6 @@

namespace sm90::decode::sparse {

template void run_flash_splitkv_mla_fp8_sparse_kernel<ModelType::V4, 128>(const SparseAttnDecodeParams &params);
template void run_flash_splitkv_mla_fp8_sparse_kernel<ModelType::V4, 128, false>(const SparseAttnDecodeParams &params);

}
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@

namespace sm90::decode::sparse {

template void run_flash_splitkv_mla_fp8_sparse_kernel<ModelType::V4, 64>(const SparseAttnDecodeParams &params);
template void run_flash_splitkv_mla_fp8_sparse_kernel<ModelType::V4, 64, false>(const SparseAttnDecodeParams &params);

}

41 changes: 23 additions & 18 deletions csrc/kernels/sm90/decode/sparse/splitkv_mla.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -83,9 +83,9 @@ __forceinline__ __device__ void scale_softmax(
*(float2*)(sScale + 2*(idx_in_warpgroup/4)) = *(float2*)(scale_for_olds);
}

template<ModelType MODEL_TYPE, int NUM_HEADS>
template<ModelType MODEL_TYPE, int NUM_HEADS, bool HAVE_TOPK_LENGTH>
template<typename TMAParams>
__device__ void KernelTemplate<MODEL_TYPE, NUM_HEADS>::devfunc(const SparseAttnDecodeParams &params, const TMAParams &tma_params) {
__device__ void KernelTemplate<MODEL_TYPE, NUM_HEADS, HAVE_TOPK_LENGTH>::devfunc(const SparseAttnDecodeParams &params, const TMAParams &tma_params) {
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ == 900)) || (defined(__CLION_IDE__) || defined(__VSCODE_IDE__))
const int head_block_idx = NUM_M_BLOCKS == 1 ? 0 : blockIdx.x;
const int s_q_idx = blockIdx.y;
Expand Down Expand Up @@ -158,20 +158,25 @@ __device__ void KernelTemplate<MODEL_TYPE, NUM_HEADS>::devfunc(const SparseAttnD
int start_block_idx, end_block_idx;
bool is_no_split;

int topk_length;
// The following fields are only valid for DeepSeek-V4
int topk_length, extra_topk_length, num_orig_kv_blocks;
int extra_topk_length, num_orig_kv_blocks;
};
auto get_cur_req_info = [&](int batch_idx) -> MainloopArgs {
MainloopArgs args;
int total_topk_padded;
int topk_length = params.topk;
if constexpr (CHECK_TOPK_LENGTH) {
if (params.topk_length) topk_length = __ldg(params.topk_length + batch_idx);
}
// Must match the per-request block count in get_decoding_sched_meta (at least one block)
int orig_topk_padded = max(ku::ceil(topk_length, (int)TOPK_BLOCK_SIZE), (int)TOPK_BLOCK_SIZE);
args.topk_length = topk_length;
if constexpr (MODEL_TYPE == ModelType::V32) {
total_topk_padded = params.topk;
total_topk_padded = HAVE_TOPK_LENGTH ? orig_topk_padded : params.topk;
} else {
int topk_length = params.topk_length ? __ldg(params.topk_length + batch_idx) : params.topk;
int orig_topk_padded = max(ku::ceil(topk_length, (int)TOPK_BLOCK_SIZE), (int)TOPK_BLOCK_SIZE);
int extra_topk_length = params.extra_topk_length ? __ldg(params.extra_topk_length + batch_idx) : params.extra_topk;
total_topk_padded = orig_topk_padded + ku::ceil(extra_topk_length, (int)TOPK_BLOCK_SIZE);
args.topk_length = topk_length;
args.extra_topk_length = extra_topk_length;
args.num_orig_kv_blocks = orig_topk_padded / TOPK_BLOCK_SIZE;
}
Expand Down Expand Up @@ -528,8 +533,8 @@ __device__ void KernelTemplate<MODEL_TYPE, NUM_HEADS>::devfunc(const SparseAttnD
nxt_token_indexs[round] = __ldg(gExtraIndices + (block_idx+1-args.num_orig_kv_blocks)*TOPK_BLOCK_SIZE + idx_in_cluster*(TOPK_BLOCK_SIZE/2) + my_token_idx);
}

if constexpr (MODEL_TYPE == ModelType::V4) {
// For DeepSeek-V4, we need to check whether the token_index is within topk_length
if constexpr (CHECK_TOPK_LENGTH) {
// Check whether the token_index is within topk_length
if (rel_block_idx*TOPK_BLOCK_SIZE + idx_in_cluster*(TOPK_BLOCK_SIZE/2) + my_token_idx >= topk_length) {
token_index = -1; // To prevent IMA when we have invalid (e.g. INT_MAX) topk indexes outside topk_length
}
Expand Down Expand Up @@ -628,10 +633,10 @@ __device__ void KernelTemplate<MODEL_TYPE, NUM_HEADS>::devfunc(const SparseAttnD
if (idx_in_warpgroup < 32) {
// We put this after fence_view_async_shared() since this won't be read by async proxy
auto is_index_valid = [&](int index, int offset_within_thread) -> bool {
if constexpr (MODEL_TYPE == ModelType::V32) {
return index != -1;
} else {
if constexpr (CHECK_TOPK_LENGTH) {
return index != -1 && rel_block_idx*TOPK_BLOCK_SIZE + lane_idx*2 + offset_within_thread < topk_length;
} else {
return index != -1;
}
};
int2 indices = __ldg((int2*)(indices_base + lane_idx*2));
Expand Down Expand Up @@ -683,8 +688,8 @@ flash_fwd_splitkv_mla_fp8_sparse_kernel(__grid_constant__ const SparseAttnDecode
Kernel::devfunc(params, tma_params);
}

template<ModelType MODEL_TYPE, int NUM_HEADS>
void KernelTemplate<MODEL_TYPE, NUM_HEADS>::run(const SparseAttnDecodeParams &params) {
template<ModelType MODEL_TYPE, int NUM_HEADS, bool HAVE_TOPK_LENGTH>
void KernelTemplate<MODEL_TYPE, NUM_HEADS, HAVE_TOPK_LENGTH>::run(const SparseAttnDecodeParams &params) {
KU_ASSERT(params.h_kv == 1);
KU_ASSERT(params.topk % TOPK_BLOCK_SIZE == 0);
KU_ASSERT(params.d_qk == HEAD_DIM_K);
Expand All @@ -698,7 +703,7 @@ void KernelTemplate<MODEL_TYPE, NUM_HEADS>::run(const SparseAttnDecodeParams &pa
}
} else {
KU_ASSERT(params.extra_kv == nullptr, "V3.2 does not support extra KV cache");
KU_ASSERT(params.topk_length == nullptr, "V3.2 does not support dynamic topk length");
KU_ASSERT(HAVE_TOPK_LENGTH == (params.topk_length != nullptr));
KU_ASSERT(params.stride_kv_row == 656); // number of bytes per token (512 fp8 + 4 float32 + 64 bfloat16)
}

Expand Down Expand Up @@ -748,7 +753,7 @@ void KernelTemplate<MODEL_TYPE, NUM_HEADS>::run(const SparseAttnDecodeParams &pa
shape_Q, tma_Q,
tensor_map_o
};
auto mla_kernel = &flash_fwd_splitkv_mla_fp8_sparse_kernel<KernelTemplate<MODEL_TYPE, NUM_HEADS>, decltype(tma_params)>;
auto mla_kernel = &flash_fwd_splitkv_mla_fp8_sparse_kernel<KernelTemplate<MODEL_TYPE, NUM_HEADS, HAVE_TOPK_LENGTH>, decltype(tma_params)>;

constexpr size_t smem_size = sizeof(SharedMemoryPlan);
KU_CUDA_CHECK(cudaFuncSetAttribute(mla_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
Expand Down Expand Up @@ -779,9 +784,9 @@ void KernelTemplate<MODEL_TYPE, NUM_HEADS>::run(const SparseAttnDecodeParams &pa
KU_CHECK_KERNEL_LAUNCH();
}

template<ModelType MODEL_TYPE, int NUM_HEADS>
template<ModelType MODEL_TYPE, int NUM_HEADS, bool HAVE_TOPK_LENGTH>
void run_flash_splitkv_mla_fp8_sparse_kernel(const SparseAttnDecodeParams &params) {
KernelTemplate<MODEL_TYPE, NUM_HEADS>::run(params);
KernelTemplate<MODEL_TYPE, NUM_HEADS, HAVE_TOPK_LENGTH>::run(params);
}

}
3 changes: 1 addition & 2 deletions csrc/kernels/sm90/decode/sparse/splitkv_mla.h
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,7 @@

namespace sm90::decode::sparse {

template<ModelType MODEL_TYPE, int NUM_HEADS>
template<ModelType MODEL_TYPE, int NUM_HEADS, bool HAVE_TOPK_LENGTH>
void run_flash_splitkv_mla_fp8_sparse_kernel(const SparseAttnDecodeParams &params);

}

2 changes: 2 additions & 0 deletions setup.py
Original file line number Diff line number Diff line change
Expand Up @@ -95,6 +95,8 @@ def get_nvcc_thread_args():
"csrc/kernels/sm90/decode/sparse/instantiations/v4_persistent_h128.cu",
"csrc/kernels/sm90/decode/sparse/instantiations/v32_persistent_h64.cu",
"csrc/kernels/sm90/decode/sparse/instantiations/v32_persistent_h128.cu",
"csrc/kernels/sm90/decode/sparse/instantiations/v32_persistent_h64_topklen.cu",
"csrc/kernels/sm90/decode/sparse/instantiations/v32_persistent_h128_topklen.cu",

# sm90 sparse prefill
"csrc/kernels/sm90/prefill/sparse/instantiations/phase1_k512.cu",
Expand Down
2 changes: 1 addition & 1 deletion tests/test_flash_mla_sparse_decoding.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@ def gen_testcase() -> List[RawTestParam]:
for d_qk in [576, 512]:
for have_extra_k in ([False, True] if d_qk == 512 else [False]):
for have_extra_topk_len in ([False, True] if have_extra_k else [False]):
for have_topk_len in ([False, True] if d_qk == 512 else [False]):
for have_topk_len in [False, True]:
for h_q in [64, 128]:
cur_correctness_cases = [
RawTestParam(b, h_q, s_q, 1, s_k, is_varlen, topk,
Expand Down