Skip to content
Merged
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
12 changes: 12 additions & 0 deletions csrc/api/common.h
Original file line number Diff line number Diff line change
Expand Up @@ -111,6 +111,18 @@ inline int int64_stride_to_int(int64_t orig_stride) {
} \
} ();

// Like DISPATCH_MODEL_TYPE, but also covers the NVFP4 format (SM100-only kernel).
// Kept separate so that SM90 kernel templates are never instantiated for NVFP4.
#define DISPATCH_MODEL_TYPE_SM100(MODEL_TYPE, CONSTEXPR_NAME, ...) \
[&] () { \
if (MODEL_TYPE == ModelType::V32_NVFP4_FP8ROPE) { \
static constexpr ModelType CONSTEXPR_NAME = ModelType::V32_NVFP4_FP8ROPE; \
return __VA_ARGS__(); \
} else { \
return DISPATCH_MODEL_TYPE(MODEL_TYPE, CONSTEXPR_NAME, __VA_ARGS__); \
} \
} ();

// The following code is adapted from https://ykiko.me/en/articles/680412313/, which converts enum values to string names.
template<auto value>
constexpr auto get_static_enum_name(){
Expand Down
52 changes: 33 additions & 19 deletions csrc/api/sparse_decode.h
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,10 @@ enum class DecodeFeatures : int {
ATTN_SINK,
TOPK_LENGTH,
EXTRA_KVCACHE,
EXTRA_TOPK_LENGTH
EXTRA_TOPK_LENGTH,

// NVFP4 (e2m1 + per-16 e4m3 scales) NoPE with e4m3 RoPE. SM100-only.
NVFP4_FP8ROPE_KVCACHE_FORMAT
};

struct DecodeImplMeta {
Expand Down Expand Up @@ -82,6 +85,7 @@ class Decode_Sm100_Head64_Impl : public DecodeImplBase {
DecodeFeatures::HEAD_DIM_576,
DecodeFeatures::V32_KVCACHE_FORMAT,
DecodeFeatures::MODEL1_KVCACHE_FORMAT,
DecodeFeatures::NVFP4_FP8ROPE_KVCACHE_FORMAT,
DecodeFeatures::ATTN_SINK,
DecodeFeatures::TOPK_LENGTH,
DecodeFeatures::EXTRA_KVCACHE,
Expand All @@ -100,7 +104,7 @@ class Decode_Sm100_Head64_Impl : public DecodeImplBase {

protected:
void run_(const SparseAttnDecodeParams &params, const std::vector<FeatureT> &required_features) override {
DISPATCH_MODEL_TYPE(params.model_type, MODEL_TYPE, [&]() {
DISPATCH_MODEL_TYPE_SM100(params.model_type, MODEL_TYPE, [&]() {
sm100::decode::head64::run_flash_splitkv_mla_fp8_sparse_kernel<MODEL_TYPE>(params);
});
}
Expand All @@ -116,6 +120,7 @@ class Decode_Sm100_Head64x2_Impl : public DecodeImplBase {
DecodeFeatures::HEAD_DIM_576,
DecodeFeatures::V32_KVCACHE_FORMAT,
DecodeFeatures::MODEL1_KVCACHE_FORMAT,
DecodeFeatures::NVFP4_FP8ROPE_KVCACHE_FORMAT,
DecodeFeatures::ATTN_SINK,
DecodeFeatures::TOPK_LENGTH,
DecodeFeatures::EXTRA_KVCACHE,
Expand All @@ -134,7 +139,7 @@ class Decode_Sm100_Head64x2_Impl : public DecodeImplBase {

protected:
void run_(const SparseAttnDecodeParams &params, const std::vector<FeatureT> &required_features) override {
DISPATCH_MODEL_TYPE(params.model_type, MODEL_TYPE, [&]() {
DISPATCH_MODEL_TYPE_SM100(params.model_type, MODEL_TYPE, [&]() {
for (int start_head_idx = 0; start_head_idx < 128; start_head_idx += 64) {
SparseAttnDecodeParams cur_params = params;
cur_params.q += start_head_idx * params.stride_q_h_q;
Expand Down Expand Up @@ -286,18 +291,34 @@ sparse_attn_decode_interface(

// Check shape
KU_CHECK_SHAPE(q, b, s_q, h_q, d_qk);
ModelType model_type;
{
int bytes_per_token;
// Infer the quantized KV cache format from the KV cache's bytes-per-token
// (i.e. its last dim)
const int bytes_per_token = static_cast<int>(kv.size(3));
if (d_qk == 576 && d_v == 512) {
// V3.2 style
bytes_per_token = 512 + 64*2 + (512/128)*4;
if (bytes_per_token == 512 + 64*2 + (512/128)*4) {
// V3.2 style, 656B/token: 512B e4m3 NoPE | 128B bf16 RoPE | 16B fp32 NoPE SF
model_type = ModelType::V32;
} else if (bytes_per_token == 352) {
// NVFP4 NoPE + FP8 RoPE, 352B/token: 256B e2m1 NoPE | 64B e4m3 RoPE (unscaled) | 32B e4m3 NoPE SF
// Must match KernelTemplate<V32_NVFP4_FP8ROPE>::BYTES_PER_TOKEN
model_type = ModelType::V32_NVFP4_FP8ROPE;
} else {
STD_TORCH_CHECK(false, "Cannot infer the sparse KV cache format: with d_qk == ", d_qk, " and d_v == ", d_v,
", kv.size(-1) (bytes per token) must be 656 (fp8) or 352 (nvfp4 nope + fp8 rope), but got ", bytes_per_token);
}
} else if (d_qk == 512 && d_v == 512) {
// MODEL1 style
bytes_per_token = 448 + 64*2 + (448/64)*1 + 1;
if (bytes_per_token == 448 + 64*2 + (448/64)*1 + 1) {
// MODEL1 style, 584B/token
model_type = ModelType::MODEL1;
} else {
STD_TORCH_CHECK(false, "Cannot infer the sparse KV cache format: with d_qk == ", d_qk, " and d_v == ", d_v,
", kv.size(-1) (bytes per token) must be 584 (fp8), but got ", bytes_per_token);
}
} else {
STD_TORCH_CHECK(false, "Unsupported head sizes for is_fp8_kvcache == True");
STD_TORCH_CHECK(false, "Unsupported head sizes for sparse decoding: d_qk == ", d_qk, ", d_v == ", d_v);
}
KU_CHECK_SHAPE(kv, num_blocks, page_block_size, h_kv, bytes_per_token);
KU_CHECK_SHAPE(extra_kv, extra_num_blocks, extra_page_block_size, h_kv, bytes_per_token);
STD_TORCH_CHECK(kv.stride(1) == bytes_per_token, "The whole block must be contiguous when is_fp8_cache is True for kv cache");
if (extra_kv.has_value()) {
Expand All @@ -324,15 +345,6 @@ sparse_attn_decode_interface(
}
Tensor lse = torch::stable::new_empty(q, {b, s_q, h_q}, ScalarType::Float);

ModelType model_type;
if (d_qk == 576) {
model_type = ModelType::V32;
} else if (d_qk == 512) {
model_type = ModelType::MODEL1;
} else {
STD_TORCH_CHECK(false, "Unsupported d_qk: ", d_qk);
}

std::vector<DecodeFeatures> features;
if (h_q == 64) {
features.push_back(DecodeFeatures::HEAD_64);
Expand All @@ -352,6 +364,8 @@ sparse_attn_decode_interface(
features.push_back(DecodeFeatures::V32_KVCACHE_FORMAT);
} else if (model_type == ModelType::MODEL1) {
features.push_back(DecodeFeatures::MODEL1_KVCACHE_FORMAT);
} else if (model_type == ModelType::V32_NVFP4_FP8ROPE) {
features.push_back(DecodeFeatures::NVFP4_FP8ROPE_KVCACHE_FORMAT);
} else {
STD_TORCH_CHECK(false, "Unsupported model type: ", (int)model_type);
}
Expand Down
5 changes: 4 additions & 1 deletion csrc/params.h
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,10 @@

enum class ModelType {
V32,
MODEL1
MODEL1,
// V3.2 geometry (d_qk=576) with NVFP4 (e2m1, per-16 e4m3 scales) NoPE and
// e4m3 RoPE. SM100-only. See csrc/sm100/decode/head64/config.h for the layout.
V32_NVFP4_FP8ROPE
};

struct __align__(4*8) DecodingSchedMeta {
Expand Down
80 changes: 69 additions & 11 deletions csrc/sm100/decode/head64/config.h
Original file line number Diff line number Diff line change
Expand Up @@ -30,17 +30,67 @@ enum NamedBarriers : uint32_t {
template<ModelType MODEL_TYPE>
struct KernelTemplate {

static constexpr int D_Q = MODEL_TYPE == ModelType::V32 ? 576 : 512;
// NVFP4 format: V3.2 geometry, NoPE stored as e2m1 (2 elems/byte) with per-16-element
// e4m3 scale factors; RoPE stored as plain e4m3 with no scale factors at all — e4m3 has
// 4 exponent bits, so it spans the RoPE magnitude range unaided and a block scale would
// only cost bytes.
static constexpr bool IS_NVFP4 = MODEL_TYPE == ModelType::V32_NVFP4_FP8ROPE;
// "V32 geometry": d_qk = 576 = 512 NoPE + 64 RoPE, V = NoPE only
static constexpr bool IS_V32_GEOM = MODEL_TYPE == ModelType::V32 || IS_NVFP4;

static constexpr int D_Q = IS_V32_GEOM ? 576 : 512;
static constexpr int D_K = D_Q;
static constexpr int D_V = 512;
static constexpr int D_NOPE = MODEL_TYPE == ModelType::V32 ? 512 : 448;
static constexpr int D_NOPE = IS_V32_GEOM ? 512 : 448;
static constexpr int D_ROPE = 64;
static constexpr int QUANT_TILE_SIZE = MODEL_TYPE == ModelType::V32 ? 128 : 64;
static constexpr bool V_HAVE_ROPE = MODEL_TYPE == ModelType::V32 ? false : true;
static constexpr int NUM_SCALES_EACH_TOKEN = MODEL_TYPE == ModelType::V32 ? 4 : 8; // Padding is included
static constexpr int TMA_K_STRIDE = MODEL_TYPE == ModelType::V32 ? D_NOPE+2*D_ROPE+4*(D_NOPE/QUANT_TILE_SIZE) : D_NOPE+2*D_ROPE; // Stride of K's tensormap. This stride must 1) be a factor of the actual stride between tokens 2) large enough to cover the entire KV cache. Since TMA copy's coordinate can only be 32bit signed integers, this number must >= 128, perferrably >= 256. So we set this to 656 for V32 and 576 for MODEL1. Extra padding may be necessary for KV blocks.
static constexpr int QUANT_TILE_SIZE = IS_NVFP4 ? 16 : (MODEL_TYPE == ModelType::V32 ? 128 : 64);
static constexpr bool V_HAVE_ROPE = IS_V32_GEOM ? false : true;
static constexpr int NUM_SCALES_EACH_TOKEN = MODEL_TYPE == ModelType::V32 ? 4 : 8; // Padding is included. Unused for NVFP4 (scales live in the raw tail buffer).

// Raw (quantized) byte counts per token
static constexpr int NOPE_RAW_BYTES = IS_NVFP4 ? D_NOPE/2 : D_NOPE; // e2m1 packs 2/byte
static constexpr int ROPE_RAW_BYTES = D_ROPE; // e4m3, 1 byte per element
// NVFP4 tail region = [64B e4m3 rope | 32B e4m3 nope SF], already a multiple of 16B.
// The tail is TMA-gathered as one box, so scales arrive together with the data they scale.
static constexpr int NVFP4_NUM_NOPE_SCALES = D_NOPE/16; // 32
static constexpr int NVFP4_SF_NOPE_OFFSET = ROPE_RAW_BYTES; // within tail
// The 32 SF bytes are NOT stored in element-block order. A dequant thread owns the scale
// groups {4c + q : c = 0..7} for its own q = (idx_in_group/2) in [0, 4), which in element
// order are 8 bytes with stride 4. They are permuted at quantization time so that the scale
// for element block s (covering NoPE dims [16s, 16s+16)) lives at
// NVFP4_SF_NOPE_OFFSET + NVFP4_SF_BYTE(s), NVFP4_SF_BYTE(s) = 8*(s&3) + (s>>2)
// which maps thread q's eight scales onto the contiguous bytes [8q, 8q+8) -> one LDS.64.
// Keep this in lockstep with the quantizer (tests/quant.py and the vLLM cache writer).
static constexpr int nvfp4_sf_byte(int s) { return 8*(s & 3) + (s >> 2); }
// Thread q's scale groups {4c+q} must land on the contiguous bytes 8q+c. That also
// makes it a permutation of [0, 32), since both sides cover that range exactly once.
static constexpr bool nvfp4_sf_perm_ok() {
for (int q = 0; q < 4; ++q)
for (int c = 0; c < 8; ++c)
if (nvfp4_sf_byte(4*c + q) != 8*q + c) return false;
return true;
}
static_assert(nvfp4_sf_perm_ok());
static constexpr int TAIL_BYTES = IS_NVFP4 ? ROPE_RAW_BYTES + NVFP4_NUM_NOPE_SCALES : 16; // 96; dummy 16 for non-NVFP4
// SMEM staging stride for one tma_gather4 call (4 tokens' tails, packed): TMA requires the
// unswizzled SMEM destination to be 128B-aligned, so pad between 4-token groups.
static constexpr int TAIL_GROUP_STRIDE = ku::ceil(4*TAIL_BYTES, 128); // 384
static constexpr int BYTES_PER_TOKEN =
MODEL_TYPE == ModelType::V32 ? D_NOPE + 2*D_ROPE + 4*(D_NOPE/128) : // 656
MODEL_TYPE == ModelType::MODEL1 ? D_NOPE + 2*D_ROPE + 8 : // 584 (per-block scale suffix layout)
NOPE_RAW_BYTES + TAIL_BYTES; // NVFP4: 352
static constexpr int TMA_K_STRIDE = MODEL_TYPE == ModelType::V32 ? D_NOPE+2*D_ROPE+4*(D_NOPE/QUANT_TILE_SIZE) :
MODEL_TYPE == ModelType::MODEL1 ? D_NOPE+2*D_ROPE :
BYTES_PER_TOKEN; // Stride of K's tensormap. This stride must 1) be a factor of the actual stride between tokens 2) large enough to cover the entire KV cache. Since TMA copy's coordinate can only be 32bit signed integers, this number must >= 128, perferrably >= 256. So we set this to 656 for V32, 576 for MODEL1, and BYTES_PER_TOKEN for NVFP4. Extra padding may be necessary for KV blocks.
static_assert(D_NOPE + D_ROPE == D_Q);
static_assert(V_HAVE_ROPE ? (D_NOPE + D_ROPE == D_V) : (D_NOPE == D_V));
static_assert(!IS_NVFP4 || (TMA_K_STRIDE % 16 == 0 && TMA_K_STRIDE >= 256));
static_assert(!IS_NVFP4 || BYTES_PER_TOKEN == 352); // Keep in sync with the bytes_per_token literal in csrc/api/sparse_decode.h
// The permuted SF layout is read 8 bytes at a time from the 128B-aligned raw_tail staging
// buffer, at (t/4)*TAIL_GROUP_STRIDE + (t%4)*TAIL_BYTES + NVFP4_SF_NOPE_OFFSET + 8q, so
// every term of that offset must be a multiple of 8.
static_assert(!IS_NVFP4 || (TAIL_GROUP_STRIDE % 8 == 0 && TAIL_BYTES % 8 == 0 &&
NVFP4_SF_NOPE_OFFSET % 8 == 0 && NVFP4_NUM_NOPE_SCALES == 32));

static constexpr int B_H = 64;
static constexpr int B_TOPK = 64;
Expand All @@ -50,9 +100,9 @@ static constexpr int NUM_THREADS = 128*3; // 128 exp + 1/32 utcmma + 1/32 raw K
static constexpr float MAX_INIT_VAL = -1e30f; // To avoid (-inf) - (-inf) = NaN

static constexpr int D_Q_SW128 = 512;
static constexpr int D_Q_SW64 = MODEL_TYPE == ModelType::V32 ? 64 : 0;
static constexpr int D_Q_SW64 = IS_V32_GEOM ? 64 : 0;
static_assert(D_Q_SW128 + D_Q_SW64 == D_Q);
static constexpr int K_ROPE_SW = MODEL_TYPE == ModelType::V32 ? 64 : 128; // RoPE part stored in SW64 (for V32) or SW128 (for MODEL1), in bytes
static constexpr int K_ROPE_SW = IS_V32_GEOM ? 64 : 128; // RoPE part stored in SW64 (for V32 geometry) or SW128 (for MODEL1), in bytes

template<
typename Shape_Q_SW128, typename TMA_Q_SW128,
Expand Down Expand Up @@ -172,7 +222,13 @@ struct SharedMemoryPlan {
array_aligned<bf16, B_H*D_ROPE> rope; // RoPE part, dequantized. SW64 in v32 mode, SW128 in MODEL1 mode
} dequant[NUM_BUFS];
static_assert(sizeof(dequant) >= sizeof(bf16) * (B_H*D_Q)); // So that Q does not covers raw_nope
array_aligned<e4m3, B_H*D_NOPE> raw_nope[NUM_BUFS]; // Raw (quantized) NoPE part
array_aligned<uint8_t, B_TOPK*NOPE_RAW_BYTES> raw_nope[NUM_BUFS]; // Raw (quantized) NoPE part
// NVFP4 only: raw tail region per token = [rope raw | nope SFs],
// TMA-gathered as one box per 4 tokens; groups of 4 packed tails are padded to
// TAIL_GROUP_STRIDE for TMA alignment. Token t lives at
// (t/4)*TAIL_GROUP_STRIDE + (t%4)*TAIL_BYTES.
// Dummy-sized (16B total, unused) for other formats.
array_aligned<uint8_t, IS_NVFP4 ? (B_TOPK/4)*TAIL_GROUP_STRIDE : 16, IS_NVFP4 ? 128 : 16> raw_tail[IS_NVFP4 ? NUM_BUFS : 1];
} kv;
} u;
union {
Expand All @@ -182,13 +238,15 @@ struct SharedMemoryPlan {
CUTE_ALIGNAS(16) float rowwise_max_buf[128];
char is_token_valid[NUM_INDEX_BUFS][B_TOPK/8];
int tma_coord[NUM_INDEX_BUFS][B_TOPK];
e8m0 scales[NUM_INDEX_BUFS][B_TOPK][NUM_SCALES_EACH_TOKEN];
e8m0 scales[NUM_INDEX_BUFS][B_TOPK][IS_NVFP4 ? 1 : NUM_SCALES_EACH_TOKEN]; // Unused for NVFP4 (scales live in raw_tail)
array_aligned<uint32_t, 1> tmem_start_addr;
transac_bar_t bar_last_store_done;
transac_bar_t bar_q_tma, bar_q_utccp;
transac_bar_t bar_rope_ready[NUM_BUFS];
transac_bar_t bar_rope_ready[NUM_BUFS]; // Non-NVFP4: rope TMA done (init 1, expect_tx). NVFP4: rope dequant done (init 128, arrived by the dequant warpgroup)
transac_bar_t bar_nope_ready[NUM_BUFS];
transac_bar_t bar_raw_ready[NUM_BUFS], bar_raw_free[NUM_BUFS];
// NVFP4 only: raw tail TMA done (init 1, expect_tx) / raw tail consumed by dequant WG (init 128)
transac_bar_t bar_rawtail_ready[NUM_BUFS], bar_rawtail_free[NUM_BUFS];
transac_bar_t bar_valid_coord_scale_ready[NUM_INDEX_BUFS], bar_valid_coord_scale_free[NUM_INDEX_BUFS];
transac_bar_t bar_qk_done[NUM_BUFS], bar_so_ready[NUM_BUFS], bar_sv_done[NUM_BUFS];
};
Expand Down
8 changes: 8 additions & 0 deletions csrc/sm100/decode/head64/instantiations/v32_nvfp4_fp8rope.cu
Original file line number Diff line number Diff line change
@@ -0,0 +1,8 @@
#include "../kernel.cuh"

namespace sm100::decode::head64 {

template
void run_flash_splitkv_mla_fp8_sparse_kernel<ModelType::V32_NVFP4_FP8ROPE>(const SparseAttnDecodeParams &params);

}
Loading