From b5235fca2eb7cf5791c94eef079d1387fb498504 Mon Sep 17 00:00:00 2001 From: Yongye Zhu Date: Thu, 4 Jun 2026 10:35:09 -0400 Subject: [PATCH] [DSv4] Adding TRTLLM gen attention kernel (#43827) --- ...deepseek_v4_qnorm_rope_kv_insert_kernel.cu | 444 ++++++++++ csrc/libtorch_stable/ops.h | 17 + csrc/libtorch_stable/torch_bindings.cpp | 20 + csrc/torch_bindings.cpp | 3 +- docs/design/attention_backends.md | 14 + tests/kernels/test_compressor_kv_cache.py | 149 ++++ ..._fused_deepseek_v4_qnorm_rope_kv_insert.py | 249 +++++- .../generate_attention_backend_docs.py | 36 +- vllm/models/deepseek_v4/amd/rocm.py | 2 +- vllm/models/deepseek_v4/attention.py | 183 +++- .../models/deepseek_v4/common/ops/__init__.py | 2 + .../deepseek_v4/common/ops/cache_utils.py | 305 +++++++ vllm/models/deepseek_v4/compressor.py | 53 +- .../deepseek_v4/nvidia/flashinfer_sparse.py | 407 +++++++++ vllm/models/deepseek_v4/nvidia/flashmla.py | 11 +- .../ops/sparse_attn_compress_cutedsl.py | 805 +++++++++++++++++- vllm/utils/flashinfer.py | 17 + vllm/v1/attention/backends/mla/sparse_swa.py | 9 +- vllm/v1/attention/backends/registry.py | 11 + vllm/v1/kv_cache_interface.py | 8 +- 20 files changed, 2659 insertions(+), 86 deletions(-) create mode 100644 vllm/models/deepseek_v4/nvidia/flashinfer_sparse.py diff --git a/csrc/libtorch_stable/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu b/csrc/libtorch_stable/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu index a5f3f03de00..4d34b4b6b50 100644 --- a/csrc/libtorch_stable/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu +++ b/csrc/libtorch_stable/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu @@ -102,6 +102,35 @@ constexpr float NUM_TOKEN_CUTOFF = 1024; constexpr int kNumLanes = 32; constexpr int kElemsPerLane = kHeadDim / kNumLanes; // 16 +// Pack this lane's 16 fp32 elements into per-tensor E4M3 FP8 (one uint4 = 16 +// B), scaling by `scale` (a reciprocal scale) and saturating to ±448. Used by +// the FlashInfer full-cache path for both the Q and KV stores. +__device__ __forceinline__ uint4 packFp8E4M3x16(float const* values, + float const scale) { +#ifndef USE_ROCM + uint4 out; + auto* out2 = reinterpret_cast<__nv_fp8x2_storage_t*>(&out); + #pragma unroll + for (int i = 0; i < kElemsPerLane / 2; i++) { + float2 scaled = + make_float2(values[2 * i] * scale, values[2 * i + 1] * scale); + scaled.x = fminf(fmaxf(scaled.x, -kFp8Max), kFp8Max); + scaled.y = fminf(fmaxf(scaled.y, -kFp8Max), kFp8Max); + out2[i] = __nv_cvt_float2_to_fp8x2(scaled, __NV_SATFINITE, __NV_E4M3); + } + return out; +#else + uint8_t out_bytes[kElemsPerLane]; + #pragma unroll + for (int i = 0; i < kElemsPerLane; i++) { + float scaled = values[i] * scale; + scaled = fminf(fmaxf(scaled, -kFp8Max), kFp8Max); + out_bytes[i] = rocm_cvt_float_to_fp8_e4m3(scaled); + } + return *reinterpret_cast(out_bytes); +#endif +} + // ──────────────────────────────────────────────────────────────────────────── // Small inline helpers // ──────────────────────────────────────────────────────────────────────────── @@ -649,6 +678,257 @@ void launchFusedDeepseekV4QNormRopeKVRopeQuantInsert( #undef DISPATCH } +// ──────────────────────────────────────────────────────────────────────────── +// FlashInfer full-cache kernel +// ──────────────────────────────────────────────────────────────────────────── +// +// Sibling to the FlashMLA kernel above, used by the FlashInfer V4 sparse-MLA +// backend. Differences from the legacy path: +// * No Q head padding — output Q layout matches the input num_heads_q. +// * KV is written as a *contiguous* 512-wide row per token (token-strided), +// not the legacy UE8M0 paged layout with a separate scale tail. +// * Q/KV are stored either as bf16 or as per-tensor E4M3 FP8 (one global +// scale), selected by the STORE_Q_FP8 / STORE_KV_FP8 template flags. +// +// Grid: 1D, gridDim.x = ceil(num_tokens_full * (num_heads_q + 1) / warps). +// Each warp handles one (token, slot): slot < num_heads_q → Q, slot == +// num_heads_q → KV. +template +__global__ void fusedDeepseekV4FullCacheKernel( + scalar_t_in* __restrict__ q_inout, // [N, H, 512], in place (bf16) + uint8_t* __restrict__ q_fp8_out, // [N, H, 512] fp8, optional + int64_t const q_fp8_stride0, // elements (fp8 == bytes) + int64_t const q_fp8_stride1, // elements (fp8 == bytes) + scalar_t_in const* __restrict__ kv_in, // [N, 512] bf16 + uint8_t* __restrict__ k_cache, // contiguous bf16 or fp8 cache + int64_t const* __restrict__ slot_mapping, // [num_tokens_insert] i64 + int64_t const* __restrict__ position_ids, // [N] i64 + float const* __restrict__ cos_sin_cache, // [max_pos, 64] fp32 + float const* __restrict__ fp8_scale_ptr, // scalar, KV fp8 only + float const* __restrict__ q_fp8_scale_inv, // scalar, Q fp8 only + float const eps, + int const num_tokens_full, // = q.size(0) = kv.size(0) + int const num_tokens_insert, // = slot_mapping.size(0) + int const num_heads_q, // H (no padding) + int const cache_block_size, // tokens per cache block + int64_t const kv_block_stride, // bytes per cache block + int64_t const kv_token_stride) { // bytes per cache token +#if (!defined(__CUDA_ARCH__) || __CUDA_ARCH__ < 800) && !defined(USE_ROCM) + if constexpr (std::is_same_v) { + return; + } else { +#endif + using Converter = vllm::_typeConvert; + int const warpsPerBlock = blockDim.x / 32; + int const warpId = threadIdx.x / 32; + int const laneId = threadIdx.x % 32; + int const globalWarpIdx = blockIdx.x * warpsPerBlock + warpId; + + int const slotsPerToken = num_heads_q + 1; + int const tokenIdx = globalWarpIdx / slotsPerToken; + int const slotIdx = globalWarpIdx % slotsPerToken; + if (tokenIdx >= num_tokens_full) return; + bool const isKV = (slotIdx == num_heads_q); + // KV branch: skip DP-padded tokens (no slot reserved for them). + if (isKV && tokenIdx >= num_tokens_insert) return; + +#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900) + cudaGridDependencySynchronize(); +#endif + + int const dim_base = laneId * kElemsPerLane; // in [0, 512) step 16 + scalar_t_in const* src_ptr; + if (isKV) { + src_ptr = kv_in + static_cast(tokenIdx) * kHeadDim + dim_base; + } else { + src_ptr = q_inout + + (static_cast(tokenIdx) * num_heads_q + slotIdx) * + kHeadDim + + dim_base; + } + uint4 const v0 = *reinterpret_cast(src_ptr); + uint4 const v1 = *reinterpret_cast(src_ptr + 8); + + // ── Decode bf16 → 16 fp32 registers ─────────────────────────────────── + float elements[kElemsPerLane]; + { + auto const* p0 = + reinterpret_cast(&v0); + auto const* p1 = + reinterpret_cast(&v1); +#pragma unroll + for (int i = 0; i < 4; i++) { + float2 f2 = Converter::convert(p0[i]); + elements[2 * i] = f2.x; + elements[2 * i + 1] = f2.y; + } +#pragma unroll + for (int i = 0; i < 4; i++) { + float2 f2 = Converter::convert(p1[i]); + elements[8 + 2 * i] = f2.x; + elements[8 + 2 * i + 1] = f2.y; + } + } + + // ── Q branch: RMSNorm (no weight) ───────────────────────────────────── + if (!isKV) { + float sumOfSquares = 0.0f; +#pragma unroll + for (int i = 0; i < kElemsPerLane; i++) { + sumOfSquares += elements[i] * elements[i]; + } + sumOfSquares = warpSum(sumOfSquares); + float const rms_rcp = + rsqrtf(sumOfSquares / static_cast(kHeadDim) + eps); +#pragma unroll + for (int i = 0; i < kElemsPerLane; i++) { + elements[i] = elements[i] * rms_rcp; + } + } + + // ── GPT-J RoPE on dims [NOPE_DIM, HEAD_DIM) ─────────────────────────── + bool const is_rope_lane = dim_base >= kNopeDim; + if (is_rope_lane) { + int64_t const pos = position_ids[tokenIdx]; + constexpr int kHalfRope = kRopeDim / 2; + float const* cos_ptr = cos_sin_cache + pos * kRopeDim; + float const* sin_ptr = cos_ptr + kHalfRope; + int const rope_local_base = dim_base - kNopeDim; + int const half_base = rope_local_base >> 1; + float4 const c0 = *reinterpret_cast(cos_ptr + half_base); + float4 const c1 = *reinterpret_cast(cos_ptr + half_base + 4); + float4 const s0 = *reinterpret_cast(sin_ptr + half_base); + float4 const s1 = *reinterpret_cast(sin_ptr + half_base + 4); + float const cos_arr[8] = {c0.x, c0.y, c0.z, c0.w, c1.x, c1.y, c1.z, c1.w}; + float const sin_arr[8] = {s0.x, s0.y, s0.z, s0.w, s1.x, s1.y, s1.z, s1.w}; +#pragma unroll + for (int p = 0; p < kElemsPerLane / 2; p++) { + float const x_even = elements[2 * p]; + float const x_odd = elements[2 * p + 1]; + elements[2 * p] = x_even * cos_arr[p] - x_odd * sin_arr[p]; + elements[2 * p + 1] = x_even * sin_arr[p] + x_odd * cos_arr[p]; + } + } + + // ── Store ───────────────────────────────────────────────────────────── + if (!isKV) { + if constexpr (STORE_Q_FP8) { + float const scale_inv = VLLM_LDG(q_fp8_scale_inv); + uint4 const out = packFp8E4M3x16(elements, scale_inv); + uint8_t* dst = q_fp8_out + + static_cast(tokenIdx) * q_fp8_stride0 + + static_cast(slotIdx) * q_fp8_stride1 + dim_base; + *reinterpret_cast(dst) = out; + } else { + uint4 out0, out1; + auto* po0 = reinterpret_cast(&out0); + auto* po1 = reinterpret_cast(&out1); +#pragma unroll + for (int i = 0; i < 4; i++) { + po0[i] = Converter::convert( + make_float2(elements[2 * i], elements[2 * i + 1])); + } +#pragma unroll + for (int i = 0; i < 4; i++) { + po1[i] = Converter::convert( + make_float2(elements[8 + 2 * i], elements[8 + 2 * i + 1])); + } + scalar_t_in* dst = + q_inout + + (static_cast(tokenIdx) * num_heads_q + slotIdx) * kHeadDim + + dim_base; + *reinterpret_cast(dst) = out0; + *reinterpret_cast(dst + 8) = out1; + } + } else { + int64_t const slot_id = slot_mapping[tokenIdx]; + if (slot_id >= 0) { + int64_t const block_idx = slot_id / cache_block_size; + int64_t const pos_in_block = slot_id % cache_block_size; + uint8_t* cache_row = + k_cache + block_idx * kv_block_stride + pos_in_block * kv_token_stride; + if constexpr (STORE_KV_FP8) { + float const inv_scale = 1.0f / VLLM_LDG(fp8_scale_ptr); + uint4 const out = packFp8E4M3x16(elements, inv_scale); + *reinterpret_cast(cache_row + dim_base) = out; + } else { + uint4 out0, out1; + auto* po0 = + reinterpret_cast(&out0); + auto* po1 = + reinterpret_cast(&out1); +#pragma unroll + for (int i = 0; i < 4; i++) { + po0[i] = Converter::convert( + make_float2(elements[2 * i], elements[2 * i + 1])); + } +#pragma unroll + for (int i = 0; i < 4; i++) { + po1[i] = Converter::convert( + make_float2(elements[8 + 2 * i], elements[8 + 2 * i + 1])); + } + scalar_t_in* dst = reinterpret_cast(cache_row) + dim_base; + *reinterpret_cast(dst) = out0; + *reinterpret_cast(dst + 8) = out1; + } + } + } + +#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900) + cudaTriggerProgrammaticLaunchCompletion(); +#endif +#if (!defined(__CUDA_ARCH__) || __CUDA_ARCH__ < 800) && !defined(USE_ROCM) + } +#endif +} + +// Configure + launch helper shared by the bf16 and fp8 full-cache launchers. +template +static void launchFullCacheKernel( + scalar_t_in* q_inout, uint8_t* q_fp8_out, int64_t q_fp8_stride0, + int64_t q_fp8_stride1, scalar_t_in const* kv_in, uint8_t* k_cache, + int64_t const* slot_mapping, int64_t const* position_ids, + float const* cos_sin_cache, float const* fp8_scale, + float const* q_fp8_scale_inv, float const eps, int const num_tokens_full, + int const num_tokens_insert, int const num_heads_q, + int const cache_block_size, int64_t const kv_block_stride, + int64_t const kv_token_stride, char const* op_name, cudaStream_t stream) { + constexpr int kBlockSize = 256; + constexpr int kWarpsPerBlock = kBlockSize / 32; + int64_t const total_warps = + static_cast(num_tokens_full) * (num_heads_q + 1); + int const grid = + static_cast((total_warps + kWarpsPerBlock - 1) / kWarpsPerBlock); + auto* kernel = + fusedDeepseekV4FullCacheKernel; +#ifndef USE_ROCM + static int const sm_version = getSMVersion(); + STD_TORCH_CHECK(sm_version >= 80, op_name, + " requires sm_80+ (Ampere or newer); got sm_", sm_version); + cudaLaunchConfig_t config; + config.gridDim = dim3(grid); + config.blockDim = dim3(kBlockSize); + config.dynamicSmemBytes = 0; + config.stream = stream; + cudaLaunchAttribute attrs[1]; + attrs[0].id = cudaLaunchAttributeProgrammaticStreamSerialization; + attrs[0].val.programmaticStreamSerializationAllowed = 1; + config.attrs = attrs; + config.numAttrs = (sm_version >= 90) ? 1 : 0; + cudaLaunchKernelEx(&config, kernel, q_inout, q_fp8_out, q_fp8_stride0, + q_fp8_stride1, kv_in, k_cache, slot_mapping, position_ids, + cos_sin_cache, fp8_scale, q_fp8_scale_inv, eps, + num_tokens_full, num_tokens_insert, num_heads_q, + cache_block_size, kv_block_stride, kv_token_stride); +#else + kernel<<>>( + q_inout, q_fp8_out, q_fp8_stride0, q_fp8_stride1, kv_in, k_cache, + slot_mapping, position_ids, cos_sin_cache, fp8_scale, q_fp8_scale_inv, + eps, num_tokens_full, num_tokens_insert, num_heads_q, cache_block_size, + kv_block_stride, kv_token_stride); +#endif +} + } // namespace deepseek_v4_fused_ops } // namespace vllm @@ -735,3 +1015,167 @@ torch::stable::Tensor fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert( }); return q_out; } + +// ──────────────────────────────────────────────────────────────────────────── +// FlashInfer full-cache torch ops +// ──────────────────────────────────────────────────────────────────────────── +void fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_bf16_insert( + torch::stable::Tensor& q, // [N, H, 512] bf16, in place + torch::stable::Tensor const& kv, // [N, 512] bf16, read-only + torch::stable::Tensor& k_cache, // [num_blocks, bs, 512] bf16 + torch::stable::Tensor const& slot_mapping, // [num_tokens_insert] int64 + torch::stable::Tensor const& position_ids, // [N] int64 + torch::stable::Tensor const& cos_sin_cache, // [max_pos, 64] float32 + double eps, int64_t cache_block_size) { + using torch::headeronly::ScalarType; + STD_TORCH_CHECK(q.device().is_cuda() && q.is_contiguous(), + "q must be contiguous CUDA"); + STD_TORCH_CHECK(kv.device().is_cuda() && kv.is_contiguous(), + "kv must be contiguous CUDA"); + STD_TORCH_CHECK(k_cache.device().is_cuda(), "k_cache must be CUDA"); + STD_TORCH_CHECK(slot_mapping.device().is_cuda() && + slot_mapping.scalar_type() == ScalarType::Long, + "slot_mapping must be int64 CUDA"); + STD_TORCH_CHECK(position_ids.device().is_cuda() && + position_ids.scalar_type() == ScalarType::Long, + "position_ids must be int64 CUDA"); + STD_TORCH_CHECK(cos_sin_cache.device().is_cuda() && + cos_sin_cache.scalar_type() == ScalarType::Float && + cos_sin_cache.dim() == 2 && cos_sin_cache.size(1) == 64, + "cos_sin_cache shape [max_pos, 64] float32"); + STD_TORCH_CHECK(q.dim() == 3 && q.size(2) == 512, "q shape [N, H, 512]"); + STD_TORCH_CHECK(kv.dim() == 2 && kv.size(1) == 512, "kv shape [N, 512]"); + STD_TORCH_CHECK(q.scalar_type() == ScalarType::BFloat16 && + kv.scalar_type() == ScalarType::BFloat16, + "q and kv must be bfloat16"); + STD_TORCH_CHECK(k_cache.dim() == 3 && k_cache.size(1) == cache_block_size && + k_cache.size(2) == 512 && k_cache.stride(2) == 1, + "k_cache shape [num_blocks, cache_block_size, 512] contiguous"); + STD_TORCH_CHECK(k_cache.scalar_type() == ScalarType::BFloat16, + "k_cache must be bfloat16"); + + int const num_tokens_full = static_cast(q.size(0)); + int const num_tokens_insert = static_cast(slot_mapping.size(0)); + STD_TORCH_CHECK(static_cast(kv.size(0)) == num_tokens_full && + static_cast(position_ids.size(0)) == num_tokens_full, + "q/kv/position_ids row counts must match"); + STD_TORCH_CHECK(num_tokens_insert <= num_tokens_full, + "slot_mapping must not exceed q row count"); + int const num_heads_q = static_cast(q.size(1)); + + const torch::stable::accelerator::DeviceGuard device_guard( + q.get_device_index()); + const cudaStream_t stream = get_current_cuda_stream(q.get_device_index()); + + // bf16 cache: 2 bytes/element -> byte strides for the uint8-addressed kernel. + int64_t const kv_block_stride = k_cache.stride(0) * 2; + int64_t const kv_token_stride = k_cache.stride(1) * 2; + + VLLM_STABLE_DISPATCH_HALF_TYPES( + q.scalar_type(), + "fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_bf16_insert", [&] { + vllm::deepseek_v4_fused_ops::launchFullCacheKernel( + reinterpret_cast(q.mutable_data_ptr()), nullptr, 0, 0, + reinterpret_cast(kv.const_data_ptr()), + reinterpret_cast(k_cache.mutable_data_ptr()), + slot_mapping.const_data_ptr(), + position_ids.const_data_ptr(), + cos_sin_cache.const_data_ptr(), nullptr, nullptr, + static_cast(eps), num_tokens_full, num_tokens_insert, + num_heads_q, static_cast(cache_block_size), kv_block_stride, + kv_token_stride, + "fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_bf16_insert", + stream); + }); +} + +void fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_fp8_insert( + torch::stable::Tensor const& q, // [N, H, 512] bf16, read-only + torch::stable::Tensor const& kv, // [N, 512] bf16, read-only + torch::stable::Tensor& q_fp8, // [N, H, 512] fp8 e4m3 + torch::stable::Tensor& k_cache, // [num_blocks, bs, 512] fp8 + torch::stable::Tensor const& slot_mapping, // [num_tokens_insert] int64 + torch::stable::Tensor const& position_ids, // [N] int64 + torch::stable::Tensor const& cos_sin_cache, // [max_pos, 64] float32 + torch::stable::Tensor const& fp8_scale, // scalar float32 (KV scale) + torch::stable::Tensor const& q_fp8_scale_inv, // scalar float32 (1 / Q scale) + double eps, int64_t cache_block_size) { + using torch::headeronly::ScalarType; + STD_TORCH_CHECK(q.device().is_cuda() && q.is_contiguous(), + "q must be contiguous CUDA"); + STD_TORCH_CHECK(kv.device().is_cuda() && kv.is_contiguous(), + "kv must be contiguous CUDA"); + STD_TORCH_CHECK(q_fp8.device().is_cuda() && q_fp8.is_contiguous() && + q_fp8.scalar_type() == ScalarType::Float8_e4m3fn && + q_fp8.dim() == 3 && q_fp8.size(0) == q.size(0) && + q_fp8.size(1) == q.size(1) && q_fp8.size(2) == q.size(2), + "q_fp8 must be a contiguous float8_e4m3fn tensor matching q"); + STD_TORCH_CHECK(k_cache.device().is_cuda(), "k_cache must be CUDA"); + STD_TORCH_CHECK(slot_mapping.device().is_cuda() && + slot_mapping.scalar_type() == ScalarType::Long, + "slot_mapping must be int64 CUDA"); + STD_TORCH_CHECK(position_ids.device().is_cuda() && + position_ids.scalar_type() == ScalarType::Long, + "position_ids must be int64 CUDA"); + STD_TORCH_CHECK(cos_sin_cache.device().is_cuda() && + cos_sin_cache.scalar_type() == ScalarType::Float && + cos_sin_cache.dim() == 2 && cos_sin_cache.size(1) == 64, + "cos_sin_cache shape [max_pos, 64] float32"); + STD_TORCH_CHECK(fp8_scale.device().is_cuda() && + fp8_scale.scalar_type() == ScalarType::Float && + fp8_scale.size(0) == 1, + "fp8_scale must be a scalar float32 CUDA tensor"); + STD_TORCH_CHECK(q_fp8_scale_inv.device().is_cuda() && + q_fp8_scale_inv.scalar_type() == ScalarType::Float && + q_fp8_scale_inv.size(0) == 1, + "q_fp8_scale_inv must be a scalar float32 CUDA tensor"); + STD_TORCH_CHECK(q.dim() == 3 && q.size(2) == 512, "q shape [N, H, 512]"); + STD_TORCH_CHECK(kv.dim() == 2 && kv.size(1) == 512, "kv shape [N, 512]"); + STD_TORCH_CHECK(q.scalar_type() == kv.scalar_type(), + "q and kv dtype must match"); + STD_TORCH_CHECK(k_cache.dim() == 3 && k_cache.size(1) == cache_block_size && + k_cache.size(2) == 512 && k_cache.stride(2) == 1, + "k_cache shape [num_blocks, cache_block_size, 512] contiguous"); + STD_TORCH_CHECK(k_cache.scalar_type() == ScalarType::Float8_e4m3fn, + "k_cache must be float8_e4m3fn"); + + int const num_tokens_full = static_cast(q.size(0)); + int const num_tokens_insert = static_cast(slot_mapping.size(0)); + STD_TORCH_CHECK(static_cast(kv.size(0)) == num_tokens_full && + static_cast(position_ids.size(0)) == num_tokens_full, + "q/kv/position_ids row counts must match"); + STD_TORCH_CHECK(num_tokens_insert <= num_tokens_full, + "slot_mapping must not exceed q row count"); + int const num_heads_q = static_cast(q.size(1)); + + const torch::stable::accelerator::DeviceGuard device_guard( + q.get_device_index()); + const cudaStream_t stream = get_current_cuda_stream(q.get_device_index()); + + VLLM_STABLE_DISPATCH_HALF_TYPES( + q.scalar_type(), + "fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_fp8_insert", [&] { + vllm::deepseek_v4_fused_ops::launchFullCacheKernel( + // q is read-only in the fp8 path (the kernel writes q_fp8); the + // launcher signature is non-const, so cast away const on the ptr. + reinterpret_cast( + const_cast(q.const_data_ptr())), + reinterpret_cast(q_fp8.mutable_data_ptr()), + q_fp8.stride(0), q_fp8.stride(1), + reinterpret_cast(kv.const_data_ptr()), + reinterpret_cast(k_cache.mutable_data_ptr()), + slot_mapping.const_data_ptr(), + position_ids.const_data_ptr(), + cos_sin_cache.const_data_ptr(), + fp8_scale.const_data_ptr(), + q_fp8_scale_inv.const_data_ptr(), static_cast(eps), + num_tokens_full, num_tokens_insert, num_heads_q, + static_cast(cache_block_size), + // fp8 cache: 1 byte/element -> stride already in bytes. + k_cache.stride(0), k_cache.stride(1), + "fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_fp8_insert", + stream); + }); +} diff --git a/csrc/libtorch_stable/ops.h b/csrc/libtorch_stable/ops.h index dd27a6968d0..0a991de76ff 100644 --- a/csrc/libtorch_stable/ops.h +++ b/csrc/libtorch_stable/ops.h @@ -238,6 +238,23 @@ torch::stable::Tensor fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert( torch::stable::Tensor const& cos_sin_cache, int64_t q_head_padded, double eps, int64_t cache_block_size); +void fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_bf16_insert( + torch::stable::Tensor& q, torch::stable::Tensor const& kv, + torch::stable::Tensor& k_cache, torch::stable::Tensor const& slot_mapping, + torch::stable::Tensor const& position_ids, + torch::stable::Tensor const& cos_sin_cache, double eps, + int64_t cache_block_size); + +void fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_fp8_insert( + torch::stable::Tensor const& q, torch::stable::Tensor const& kv, + torch::stable::Tensor& q_fp8, torch::stable::Tensor& k_cache, + torch::stable::Tensor const& slot_mapping, + torch::stable::Tensor const& position_ids, + torch::stable::Tensor const& cos_sin_cache, + torch::stable::Tensor const& fp8_scale, + torch::stable::Tensor const& q_fp8_scale_inv, double eps, + int64_t cache_block_size); + #ifndef USE_ROCM torch::stable::Tensor minimax_allreduce_rms( torch::stable::Tensor const& input, diff --git a/csrc/libtorch_stable/torch_bindings.cpp b/csrc/libtorch_stable/torch_bindings.cpp index e9a62a8666c..511a788eeae 100644 --- a/csrc/libtorch_stable/torch_bindings.cpp +++ b/csrc/libtorch_stable/torch_bindings.cpp @@ -343,6 +343,20 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_C, ops) { "Tensor slot_mapping, Tensor position_ids, Tensor cos_sin_cache, " "int q_head_padded, float eps, int cache_block_size) -> Tensor"); + // FlashInfer V4 full-cache variants: write Q in place (bf16) or to a separate + // FP8 tensor, and KV into a contiguous 512-wide token-strided cache. + ops.def( + "fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_bf16_insert(" + "Tensor! q, Tensor kv, Tensor! k_cache, Tensor slot_mapping, " + "Tensor position_ids, Tensor cos_sin_cache, float eps, " + "int cache_block_size) -> ()"); + ops.def( + "fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_fp8_insert(" + "Tensor q, Tensor kv, Tensor! q_fp8, Tensor! k_cache, " + "Tensor slot_mapping, Tensor position_ids, Tensor cos_sin_cache, " + "Tensor fp8_scale, Tensor q_fp8_scale_inv, float eps, " + "int cache_block_size) -> ()"); + #ifndef USE_ROCM ops.def( "minimax_allreduce_rms(" @@ -591,6 +605,12 @@ STABLE_TORCH_LIBRARY_IMPL(_C, CUDA, ops) { ops.impl("fused_qk_norm_rope", TORCH_BOX(&fused_qk_norm_rope)); ops.impl("fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert", TORCH_BOX(&fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert)); + ops.impl( + "fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_bf16_insert", + TORCH_BOX(&fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_bf16_insert)); + ops.impl( + "fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_fp8_insert", + TORCH_BOX(&fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_fp8_insert)); #ifndef USE_ROCM ops.impl("minimax_allreduce_rms", TORCH_BOX(&minimax_allreduce_rms)); ops.impl("minimax_allreduce_rms_qk", TORCH_BOX(&minimax_allreduce_rms_qk)); diff --git a/csrc/torch_bindings.cpp b/csrc/torch_bindings.cpp index c078222bca0..3351638f574 100644 --- a/csrc/torch_bindings.cpp +++ b/csrc/torch_bindings.cpp @@ -55,7 +55,8 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) { // Horizontally-fused DeepseekV4-MLA: per-head RMSNorm + GPT-J RoPE for Q, and // GPT-J RoPE + UE8M0 FP8 quant + paged cache insert for KV, all in one - // kernel launch. Registered in _C_stable_libtorch. + // kernel launch. Registered in _C_stable_libtorch (incl. the FlashInfer V4 + // full-cache bf16/fp8 variants). // Quantization ops #ifndef USE_ROCM diff --git a/docs/design/attention_backends.md b/docs/design/attention_backends.md index 329a4aacfb6..bd3eed7d55b 100644 --- a/docs/design/attention_backends.md +++ b/docs/design/attention_backends.md @@ -228,3 +228,17 @@ MLA decode backends are selected using the standard | `TOKENSPEED_MLA` | fp16, bf16 | `fp8`, `fp8_e4m3` | 32, 64 | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | 10.x | | `TRITON_MLA` | fp16, bf16 | `auto`, `float16`, `bfloat16`, `fp8`, `fp8_e4m3` | %16 | Any | ❌ | ❌ | ❌ | ❌ | ✅ | Decoder | Any | | `XPU_MLA_SPARSE` | fp16, bf16 | `auto`, `float16`, `bfloat16` | Any | 576 | ❌ | ❌ | ✅ | ❌ | ❌ | Decoder | Any | + +### DeepSeek V4 Decode Backends + +DeepSeek V4 sparse MLA uses its own decode backends, selected via +`--attention-backend=` (e.g., `FLASHMLA_SPARSE_DSV4`, +`FLASHINFER_MLA_SPARSE_DSV4`). They share the V4 sparse-index +pipeline (compressor + SWA + indexer, 256-token blocks, head 512); +default on NVIDIA is `FLASHMLA_SPARSE_DSV4`. + +| Backend | Dtypes | KV Dtypes | Block Sizes | Head Sizes | Sink | Non-Causal | Sparse | MM Prefix | DCP | Attention Types | Compute Cap. | +| ------- | ------ | --------- | ----------- | ---------- | ---- | ---------- | ------ | --------- | --- | --------------- | ------------ | +| `FLASHINFER_MLA_SPARSE_DSV4` | fp16, bf16 | `auto` | Any | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | Any | +| `FLASHMLA_SPARSE_DSV4` | fp16, bf16 | `auto` | 256 | 512 | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | Any | +| `ROCM_FLASHMLA_SPARSE_DSV4` | fp16, bf16 | `auto` | Any | Any | ❌ | ❌ | ❌ | ❌ | ❌ | Decoder | N/A | diff --git a/tests/kernels/test_compressor_kv_cache.py b/tests/kernels/test_compressor_kv_cache.py index c6daab2d86b..74dc01472a8 100644 --- a/tests/kernels/test_compressor_kv_cache.py +++ b/tests/kernels/test_compressor_kv_cache.py @@ -468,6 +468,7 @@ def _reference_kv_compress_norm_rope( use_fp4: bool = False, rms_eps: float = 1e-6, fp8_max: float = 448.0, + return_full_cache: bool = False, ): """Compress → RMSNorm → GPT-J RoPE → quantize. @@ -521,6 +522,12 @@ def _reference_kv_compress_norm_rope( results.append(torch.cat([nope, rope]).to(state_cache.dtype)) result = torch.stack(results) + if return_full_cache: + # Contiguous 512-wide bf16 row (nope unrotated + rope rotated), matching + # the FlashInfer full-cache layout before any per-tensor fp8 quant. The + # kernel rounds the fp32 result to bf16 once at the store. + return result.to(torch.bfloat16) + if use_fp4: return quantize_to_mxfp4(result) else: @@ -667,3 +674,145 @@ def test_fused_kv_insert_indexer(num_tokens: int, kv_block_size: int, use_fp4: b assert torch.equal(actual_scale, scale[i : i + 1]), ( f"token {i}: scale {actual_scale.item()} != {scale[i].item()}" ) + + +@pytest.mark.parametrize("compress_ratio", [4, 128]) +@pytest.mark.parametrize("store_fp8", [False, True]) +def test_cutedsl_full_cache_store(compress_ratio: int, store_fp8: bool): + """CuTeDSL compressor full-cache (FlashInfer) store parity for head=512. + + Exercises the contiguous bf16 / per-tensor fp8 store branch of both the C4 + fused kernel and the C128 split kernel against the PyTorch reference. + """ + cutedsl = pytest.importorskip("cutlass") # noqa: F841 + from vllm.models.deepseek_v4.nvidia.ops.sparse_attn_compress_cutedsl import ( + fused_kv_compress_norm_rope_insert_sparse_attn_cutedsl, + split_kv_compress_norm_rope_insert_sparse_attn_cutedsl, + ) + + HEAD_DIM = 512 + ROPE_DIM = 64 + RMS_EPS = 1e-6 + FP8_MAX = 448.0 + # C128 compress (Block8 kernel) requires state-cache block_size=8; C4 uses 16. + BLOCK_SIZE = 8 if compress_ratio == 128 else 16 + KV_BLOCK_SIZE = 64 + device = "cuda" + torch.manual_seed(7) + + overlap = 1 if compress_ratio == 4 else 0 + coff = 1 + overlap + num_tokens = 8 + + num_pages = (compress_ratio * num_tokens - 1) // BLOCK_SIZE + 2 + # The production CompressorStateCache is fp32. + state_cache = torch.randn( + num_pages, BLOCK_SIZE, 2 * coff * HEAD_DIM, dtype=torch.float32, device=device + ) + block_table = torch.arange(num_pages, dtype=torch.int32, device=device).unsqueeze(0) + token_to_req = torch.zeros(num_tokens, dtype=torch.int32, device=device) + slot_mapping = torch.arange(num_tokens, dtype=torch.int64, device=device) + positions = torch.arange( + compress_ratio - 1, + compress_ratio * num_tokens, + compress_ratio, + dtype=torch.int64, + device=device, + ) + rms_weight = torch.randn(HEAD_DIM, dtype=torch.bfloat16, device=device) + cos_sin_cache = torch.randn( + compress_ratio * num_tokens, ROPE_DIM, dtype=torch.float32, device=device + ) + + dtype = torch.float8_e4m3fn if store_fp8 else torch.bfloat16 + kv_n_blocks = (num_tokens + KV_BLOCK_SIZE - 1) // KV_BLOCK_SIZE + 1 + k_cache = torch.zeros( + kv_n_blocks, KV_BLOCK_SIZE, HEAD_DIM, dtype=dtype, device=device + ) + fp8_scale = torch.tensor( + [0.5 if store_fp8 else 1.0], dtype=torch.float32, device=device + ) + + if compress_ratio == 4: + fused_kv_compress_norm_rope_insert_sparse_attn_cutedsl( + state_cache, + token_to_req, + positions, + slot_mapping, + block_table, + BLOCK_SIZE, + rms_weight, + RMS_EPS, + cos_sin_cache, + k_cache, + slot_mapping, + KV_BLOCK_SIZE, + k_cache.stride(0), + head_size=HEAD_DIM, + state_width=coff * HEAD_DIM, + rope_head_dim=ROPE_DIM, + fp8_max=FP8_MAX, + quant_block=64, + token_stride=576, + scale_dim=8, + compress_ratio=compress_ratio, + overlap=True, + store_full_kv=True, + store_full_fp8=store_fp8, + fp8_scale=fp8_scale, + ) + else: + compressed_kv = torch.empty( + (num_tokens, HEAD_DIM), dtype=torch.float32, device=device + ) + split_kv_compress_norm_rope_insert_sparse_attn_cutedsl( + state_cache, + token_to_req, + positions, + slot_mapping, + block_table, + BLOCK_SIZE, + compressed_kv, + rms_weight, + RMS_EPS, + cos_sin_cache, + k_cache, + slot_mapping, + KV_BLOCK_SIZE, + k_cache.stride(0), + head_size=HEAD_DIM, + state_width=coff * HEAD_DIM, + rope_head_dim=ROPE_DIM, + fp8_max=FP8_MAX, + quant_block=64, + token_stride=576, + scale_dim=8, + compress_ratio=compress_ratio, + overlap=bool(overlap), + store_full_kv=True, + store_full_fp8=store_fp8, + fp8_scale=fp8_scale, + ) + + ref = _reference_kv_compress_norm_rope( + state_cache, + block_table, + positions, + rms_weight, + cos_sin_cache, + compress_ratio, + overlap, + rms_eps=RMS_EPS, + return_full_cache=True, + ) # [num_tokens, HEAD_DIM] bf16 + + actual = torch.stack( + [k_cache[i // KV_BLOCK_SIZE, i % KV_BLOCK_SIZE] for i in range(num_tokens)] + ) + if store_fp8: + ref_fp8 = torch.clamp(ref.float() / fp8_scale, -FP8_MAX, FP8_MAX).to( + torch.float8_e4m3fn + ) + torch.testing.assert_close(actual.float(), ref_fp8.float(), rtol=0.0, atol=0.3) + else: + torch.testing.assert_close(actual.float(), ref.float(), rtol=3e-2, atol=3e-2) diff --git a/tests/kernels/test_fused_deepseek_v4_qnorm_rope_kv_insert.py b/tests/kernels/test_fused_deepseek_v4_qnorm_rope_kv_insert.py index a49ea498e5e..e568ce57638 100644 --- a/tests/kernels/test_fused_deepseek_v4_qnorm_rope_kv_insert.py +++ b/tests/kernels/test_fused_deepseek_v4_qnorm_rope_kv_insert.py @@ -67,7 +67,7 @@ def apply_rope_gptj_last_k( head_dim = x.shape[-1] nope_dim = head_dim - rope_dim - cs = cos_sin_cache[positions].to(torch.float32) + cs = cos_sin_cache[positions.long()].to(torch.float32) cos = cs[..., :half] sin = cs[..., half:] @@ -114,6 +114,18 @@ def _op_available() -> bool: return hasattr(torch.ops._C, "fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert") +def _full_cache_fp8_op_available() -> bool: + return hasattr( + torch.ops._C, "fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_fp8_insert" + ) + + +def _full_cache_bf16_op_available() -> bool: + return hasattr( + torch.ops._C, "fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_bf16_insert" + ) + + pytestmark = pytest.mark.skipif( not torch.cuda.is_available() or not _op_available(), reason="CUDA not available or fused DeepseekV4 op not built in", @@ -415,3 +427,238 @@ def test_combined_q_and_kv( "padded head slots must be exact zero" ) torch.testing.assert_close(k_cache_fused, k_cache_ref, rtol=0, atol=0) + + +# ── Full-cache (FlashInfer) path parity ────────────────────────────────────── + + +def _call_full_cache_fp8_fused( + q, + kv, + q_fp8, + k_cache, + slot_mapping, + positions, + cos_sin_cache, + fp8_scale, + q_fp8_scale_inv, + eps, + bs, +): + torch.ops._C.fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_fp8_insert( + q, + kv, + q_fp8, + k_cache, + slot_mapping, + positions.long(), + cos_sin_cache, + fp8_scale, + q_fp8_scale_inv, + eps, + bs, + ) + + +def _call_full_cache_bf16_fused( + q, + kv, + k_cache, + slot_mapping, + positions, + cos_sin_cache, + eps, + bs, +): + torch.ops._C.fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_bf16_insert( + q, + kv, + k_cache, + slot_mapping, + positions.long(), + cos_sin_cache, + eps, + bs, + ) + + +def _fp8_full_cache_reference( + q, + kv, + k_cache, + q_fp8, + slot_mapping, + positions, + cos_sin_cache, + eps, + block_size, + fp8_scale, + q_fp8_scale_inv, +): + q_ref = rmsnorm_no_weight(q, eps) + q_ref = apply_rope_gptj_last_k(q_ref, positions, cos_sin_cache) + q_fp8.copy_( + torch.clamp(q_ref.float() * q_fp8_scale_inv, -FP8_MAX, FP8_MAX).to( + torch.float8_e4m3fn + ) + ) + + kv_ref = apply_rope_gptj_last_k(kv, positions, cos_sin_cache) + valid = slot_mapping >= 0 + slots = slot_mapping[valid] + block_idx = slots // block_size + pos_in_block = slots % block_size + k_cache[block_idx, pos_in_block] = torch.clamp( + kv_ref[valid].float() / fp8_scale, -FP8_MAX, FP8_MAX + ).to(torch.float8_e4m3fn) + + +def _bf16_full_cache_reference( + q, + kv, + k_cache, + slot_mapping, + positions, + cos_sin_cache, + eps, + block_size, +): + q_ref = rmsnorm_no_weight(q, eps) + # Kernel keeps RMSNorm+RoPE in fp32 and rounds to bf16 once at the store. + q_ref = apply_rope_gptj_last_k(q_ref, positions, cos_sin_cache).to(q.dtype) + + kv_ref = apply_rope_gptj_last_k(kv, positions, cos_sin_cache) + valid = slot_mapping >= 0 + slots = slot_mapping[valid] + block_idx = slots // block_size + pos_in_block = slots % block_size + k_cache[block_idx, pos_in_block] = kv_ref[valid] + return q_ref + + +@pytest.mark.skipif( + not _full_cache_fp8_op_available(), + reason="full-cache per-tensor FP8 DeepseekV4 op not built in", +) +@pytest.mark.parametrize("num_tokens", [4, 17]) +@pytest.mark.parametrize("n_heads", [8, 17]) +@pytest.mark.parametrize("positions_dtype", [torch.int32, torch.int64]) +def test_full_cache_per_tensor_fp8_matches_reference( + num_tokens: int, + n_heads: int, + positions_dtype: torch.dtype, +): + torch.manual_seed(4) + device = "cuda" + dtype = torch.bfloat16 + eps = 1e-6 + block_size = 16 + max_pos = 4096 + + q = torch.randn(num_tokens, n_heads, HEAD_DIM, dtype=dtype, device=device) + kv = torch.randn(num_tokens, HEAD_DIM, dtype=dtype, device=device) + positions = torch.arange(num_tokens, dtype=positions_dtype, device=device) + cos_sin_cache = make_cos_sin_cache(max_pos, ROPE_DIM, torch.float32, device) + + num_blocks = (num_tokens + block_size - 1) // block_size + 1 + slot_mapping = torch.arange(num_tokens, dtype=torch.int64, device=device) + fp8_scale = torch.tensor([1.0], dtype=torch.float32, device=device) + q_fp8_scale_inv = torch.tensor([1.0], dtype=torch.float32, device=device) + + q_fp8_ref = torch.empty_like(q, dtype=torch.float8_e4m3fn) + q_fp8_fused = torch.empty_like(q, dtype=torch.float8_e4m3fn) + k_cache_ref = torch.zeros( + num_blocks, block_size, HEAD_DIM, dtype=torch.float8_e4m3fn, device=device + ) + k_cache_fused = torch.zeros_like(k_cache_ref) + + _fp8_full_cache_reference( + q, + kv, + k_cache_ref, + q_fp8_ref, + slot_mapping, + positions, + cos_sin_cache, + eps, + block_size, + fp8_scale, + q_fp8_scale_inv, + ) + _call_full_cache_fp8_fused( + q.clone(), + kv, + q_fp8_fused, + k_cache_fused, + slot_mapping, + positions, + cos_sin_cache, + fp8_scale, + q_fp8_scale_inv, + eps, + block_size, + ) + + torch.testing.assert_close( + q_fp8_fused.float(), q_fp8_ref.float(), rtol=0, atol=0.25 + ) + torch.testing.assert_close( + k_cache_fused.float(), k_cache_ref.float(), rtol=0, atol=0.25 + ) + + +@pytest.mark.skipif( + not _full_cache_bf16_op_available(), + reason="full-cache BF16 DeepseekV4 op not built in", +) +@pytest.mark.parametrize("num_tokens", [4, 17]) +@pytest.mark.parametrize("n_heads", [8, 17]) +@pytest.mark.parametrize("positions_dtype", [torch.int32, torch.int64]) +def test_full_cache_bf16_matches_reference( + num_tokens: int, + n_heads: int, + positions_dtype: torch.dtype, +): + torch.manual_seed(5) + device = "cuda" + dtype = torch.bfloat16 + eps = 1e-6 + block_size = 16 + max_pos = 4096 + + q = torch.randn(num_tokens, n_heads, HEAD_DIM, dtype=dtype, device=device) + kv = torch.randn(num_tokens, HEAD_DIM, dtype=dtype, device=device) + positions = torch.arange(num_tokens, dtype=positions_dtype, device=device) + cos_sin_cache = make_cos_sin_cache(max_pos, ROPE_DIM, torch.float32, device) + + num_blocks = (num_tokens + block_size - 1) // block_size + 1 + slot_mapping = torch.arange(num_tokens, dtype=torch.int64, device=device) + + q_fused = q.clone() + k_cache_ref = torch.zeros( + num_blocks, block_size, HEAD_DIM, dtype=torch.bfloat16, device=device + ) + k_cache_fused = torch.zeros_like(k_cache_ref) + q_ref = _bf16_full_cache_reference( + q, + kv, + k_cache_ref, + slot_mapping, + positions, + cos_sin_cache, + eps, + block_size, + ) + _call_full_cache_bf16_fused( + q_fused, + kv, + k_cache_fused, + slot_mapping, + positions, + cos_sin_cache, + eps, + block_size, + ) + + torch.testing.assert_close(q_fused, q_ref, rtol=1e-2, atol=1e-2) + torch.testing.assert_close(k_cache_fused, k_cache_ref, rtol=0, atol=0) diff --git a/tools/pre_commit/generate_attention_backend_docs.py b/tools/pre_commit/generate_attention_backend_docs.py index 1a93068537b..1f7150ce6a7 100644 --- a/tools/pre_commit/generate_attention_backend_docs.py +++ b/tools/pre_commit/generate_attention_backend_docs.py @@ -1562,7 +1562,9 @@ def generate_legend() -> str: def generate_mla_section( - prefill_backends: list[dict[str, Any]], decode_backends: list[dict[str, Any]] + prefill_backends: list[dict[str, Any]], + decode_backends: list[dict[str, Any]], + v4_decode_backends: list[dict[str, Any]] | None = None, ) -> str: """Generate the complete MLA section with prefill and decode tables.""" lines = [ @@ -1611,6 +1613,22 @@ def generate_mla_section( columns = _build_columns(is_mla=True, has_versions=False) lines.extend(_render_table(columns, decode_backends)) + if v4_decode_backends: + lines.extend( + [ + "", + "### DeepSeek V4 Decode Backends", + "", + "DeepSeek V4 sparse MLA uses its own decode backends, selected via", + "`--attention-backend=` (e.g., `FLASHMLA_SPARSE_DSV4`,", + "`FLASHINFER_MLA_SPARSE_DSV4`). They share the V4 sparse-index", + "pipeline (compressor + SWA + indexer, 256-token blocks, head 512);", + "default on NVIDIA is `FLASHMLA_SPARSE_DSV4`.", + "", + ] + ) + lines.extend(_render_table(columns, v4_decode_backends)) + lines.append("") return "\n".join(lines) @@ -1651,9 +1669,15 @@ def generate_docs() -> str: if fi_features: all_backends = _expand_flashinfer_variants(all_backends, fi_features) - # Split into MLA and non-MLA - mla_backends = [b for b in all_backends if b["is_mla"]] - non_mla_backends = [b for b in all_backends if not b["is_mla"]] + # DeepSeek V4 (*_DSV4) decode backends get their own subsection rather than + # mixing into the main MLA / standard tables (the ROCm V4 backend isn't + # flagged is_mla by the AST heuristic, so filter purely on the name). + def _is_v4(b: dict[str, Any]) -> bool: + return b["name"].endswith("_DSV4") + + v4_decode_backends = [b for b in all_backends if _is_v4(b)] + mla_backends = [b for b in all_backends if b["is_mla"] and not _is_v4(b)] + non_mla_backends = [b for b in all_backends if not b["is_mla"] and not _is_v4(b)] # Generate documentation script_path = "tools/pre_commit/generate_attention_backend_docs.py" @@ -1703,7 +1727,9 @@ def generate_docs() -> str: doc_lines.append("\n>\n".join(footnotes) + "\n") # Add MLA section with prefill and decode backends - doc_lines.append(generate_mla_section(mla_prefill_backends, mla_backends)) + doc_lines.append( + generate_mla_section(mla_prefill_backends, mla_backends, v4_decode_backends) + ) return "\n".join(doc_lines) diff --git a/vllm/models/deepseek_v4/amd/rocm.py b/vllm/models/deepseek_v4/amd/rocm.py index 2af93fba31e..7298f18365d 100644 --- a/vllm/models/deepseek_v4/amd/rocm.py +++ b/vllm/models/deepseek_v4/amd/rocm.py @@ -576,7 +576,7 @@ class DeepseekV4ROCMAiterSparseSWAMetadataBuilder(DeepseekSparseSWAMetadataBuild class DeepseekV4ROCMAiterMLASparseBackend(DeepseekV4FlashMLASparseBackend): @staticmethod def get_name() -> str: - return "ROCM_V4_FLASHMLA_SPARSE" + return "ROCM_FLASHMLA_SPARSE_DSV4" @staticmethod def get_builder_cls() -> type["DeepseekV4ROCMAiterMLASparseMetadataBuilder"]: diff --git a/vllm/models/deepseek_v4/attention.py b/vllm/models/deepseek_v4/attention.py index 55cb3d94ba6..5f13d1bd8d0 100644 --- a/vllm/models/deepseek_v4/attention.py +++ b/vllm/models/deepseek_v4/attention.py @@ -55,9 +55,6 @@ from vllm.utils.multi_stream_utils import ( maybe_execute_in_parallel, ) from vllm.v1.attention.backend import AttentionBackend, AttentionMetadata -from vllm.v1.attention.backends.mla.flashmla_sparse import ( - FlashMLASparseBackend, -) from vllm.v1.attention.backends.mla.indexer import ( DeepseekV4IndexerBackend, get_max_prefill_buffer_size, @@ -73,21 +70,82 @@ if TYPE_CHECKING: logger = init_logger(__name__) -def _select_v4_sparse_impl() -> "type[DeepseekV4SparseMLAAttentionImpl]": - """Pick the platform-specific V4 sparse MLA impl class. Sole platform check.""" +def _resolve_dsv4_backend(vllm_config: VllmConfig | None): + """Return the explicitly-requested DSv4 sparse backend enum, or None.""" + if vllm_config is None: + return None + attn_config = getattr(vllm_config, "attention_config", None) + return getattr(attn_config, "backend", None) if attn_config is not None else None + + +def _select_v4_sparse_impl( + vllm_config: VllmConfig | None = None, +) -> "type[DeepseekV4SparseMLAAttentionImpl]": + """Pick the V4 sparse MLA impl class. + + An explicit ``--attention-backend FLASHINFER_MLA_SPARSE_DSV4`` selects the + FlashInfer TRTLLM-gen path; otherwise the platform default (FlashMLA on + NVIDIA, ROCm Aiter on AMD) is used. + """ + from vllm.v1.attention.backends.registry import AttentionBackendEnum + + backend = _resolve_dsv4_backend(vllm_config) + if backend == AttentionBackendEnum.FLASHINFER_MLA_SPARSE_DSV4: + from vllm.models.deepseek_v4.nvidia.flashinfer_sparse import ( + DeepseekV4FlashInferMLASparseImpl, + ) + + logger.info_once("Using FLASHINFER_MLA_SPARSE_DSV4 backend.") + return DeepseekV4FlashInferMLASparseImpl if current_platform.is_rocm(): from vllm.models.deepseek_v4.amd.rocm import ( DeepseekV4ROCMAiterMLASparseImpl, ) + logger.info_once("Using ROCM_FLASHMLA_SPARSE_DSV4 backend.") return DeepseekV4ROCMAiterMLASparseImpl from vllm.models.deepseek_v4.nvidia.flashmla import ( DeepseekV4FlashMLASparseImpl, ) + logger.info_once("Using FLASHMLA_SPARSE_DSV4 backend.") return DeepseekV4FlashMLASparseImpl +def _resolve_dsv4_kv_cache_dtype( + backend, + kv_cache_dtype: str, + cache_config: CacheConfig | None, +) -> tuple[str, torch.dtype]: + """Map ``(backend, --kv-cache-dtype)`` to ``(cache_dtype_str, torch_dtype)``. + + FlashInfer V4 reads a contiguous 512-wide KV row (bf16 or per-tensor FP8 + E4M3); FlashMLA V4 reads the legacy UE8M0 paged layout (uint8 / + ``fp8_ds_mla``). For FlashMLA the canonical ``fp8_ds_mla`` string is + written back onto ``cache_config`` so the page-size specs pick the 576B + layout. + """ + from vllm.v1.attention.backends.registry import AttentionBackendEnum + + if backend == AttentionBackendEnum.FLASHINFER_MLA_SPARSE_DSV4: + if kv_cache_dtype.startswith("fp8"): + return kv_cache_dtype, torch.float8_e4m3fn + # auto / bfloat16 -> contiguous BF16 cache. + return kv_cache_dtype, torch.bfloat16 + + # FlashMLA (and ROCm Aiter): legacy UE8M0 paged uint8 cache. + assert kv_cache_dtype.startswith("fp8"), ( + f"DeepseekV4 FlashMLA sparse backend only supports fp8 kv-cache, " + f"got {kv_cache_dtype}" + ) + if kv_cache_dtype != "fp8_ds_mla": + if cache_config is not None: + cache_config.cache_dtype = "fp8_ds_mla" + kv_cache_dtype = "fp8_ds_mla" + logger.info_once("Using DeepSeek's fp8_ds_mla KV cache format.") + return kv_cache_dtype, torch.uint8 + + class DeepseekV4MLA(nn.Module): def __init__( self, @@ -194,10 +252,17 @@ class DeepseekV4MLA(nn.Module): self.ln_events = [torch.cuda.Event() for _ in range(4)] assert cache_config is not None, "DeepseekV4 attention requires cache_config" + # Resolve the SWA cache tensor dtype from the selected backend: FlashMLA + # uses the legacy UE8M0 paged uint8 layout; FlashInfer uses a contiguous + # bf16 / per-tensor fp8 row. + backend = _resolve_dsv4_backend(vllm_config) + _, swa_cache_torch_dtype = _resolve_dsv4_kv_cache_dtype( + backend, cache_config.cache_dtype, cache_config + ) self.swa_cache_layer = DeepseekV4SWACache( head_dim=self.head_dim, window_size=self.window_size, - dtype=torch.uint8, + dtype=swa_cache_torch_dtype, prefix=f"{prefix}.swa_cache", cache_config=cache_config, ) @@ -478,25 +543,66 @@ class DeepseekV4MLA(nn.Module): assert swa_metadata is not None swa_kv_cache = self.swa_cache_layer.kv_cache - swa_kv_cache_2d = swa_kv_cache.view(swa_kv_cache.shape[0], -1) + # The fused insert ops require int64 position_ids; the runner's positions + # buffer is already int64, so no cast is needed. + assert positions.dtype == torch.int64 + cos_sin_cache = self.rotary_emb.cos_sin_cache + cache_dtype = swa_kv_cache.dtype - # Horizontally fused: - # Q side: q_head_norm (per-head RMSNorm, no weight) + GPT-J RoPE, - # with zero-fill for the padding head slots. The kernel - # allocates and returns the padded q tensor. - # KV side: GPT-J RoPE + UE8M0 FP8 quant + paged cache insert # kv is unchanged; mla_attn reads kv solely via swa_kv_cache. - return torch.ops._C.fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert( + if cache_dtype == torch.uint8: + # Legacy FlashMLA UE8M0 paged path. Horizontally fused: + # Q side: per-head RMSNorm (no weight) + GPT-J RoPE, zero-filling + # the padding head slots; the kernel allocates and returns + # the padded q tensor. + # KV side: GPT-J RoPE + UE8M0 FP8 quant + paged cache insert. + swa_kv_cache_2d = swa_kv_cache.view(swa_kv_cache.shape[0], -1) + return torch.ops._C.fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert( + q, + kv, + swa_kv_cache_2d, + swa_metadata.slot_mapping, + positions, + cos_sin_cache, + self.padded_heads, + self.eps, + swa_metadata.block_size, + ) + + # FlashInfer full-cache path: contiguous [num_blocks, block_size, 512] + # cache (no Q padding). bf16 rewrites q in place; per-tensor fp8 writes a + # separately-allocated fp8 q and quantizes the KV row. + block_size = swa_metadata.block_size + swa_kv_cache_3d = swa_kv_cache.view(-1, block_size, self.head_dim) + if cache_dtype == torch.bfloat16: + torch.ops._C.fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_bf16_insert( + q, + kv, + swa_kv_cache_3d, + swa_metadata.slot_mapping, + positions, + cos_sin_cache, + self.eps, + block_size, + ) + return q + + # per-tensor fp8 (torch.float8_e4m3fn) + q_fp8 = torch.empty_like(q, dtype=torch.float8_e4m3fn) + torch.ops._C.fused_deepseek_v4_qnorm_rope_kv_rope_full_cache_fp8_insert( q, kv, - swa_kv_cache_2d, + q_fp8, + swa_kv_cache_3d, swa_metadata.slot_mapping, - positions.to(torch.int64), - self.rotary_emb.cos_sin_cache, - self.padded_heads, + positions, + cos_sin_cache, + self.mla_attn._flashinfer_fp8_kv_scale, + self.mla_attn._flashinfer_fp8_q_scale_inv, self.eps, - swa_metadata.block_size, + block_size, ) + return q_fp8 class DeepseekV4MLAAttention(nn.Module, AttentionLayerBase): @@ -524,7 +630,8 @@ class DeepseekV4MLAAttention(nn.Module, AttentionLayerBase): **extra_impl_args, ) -> None: super().__init__() - self.impl_cls = _select_v4_sparse_impl() + vllm_config = get_current_vllm_config() + self.impl_cls = _select_v4_sparse_impl(vllm_config) self.backend_cls = self.impl_cls.backend_cls self.num_heads = num_heads self.num_kv_heads = 1 @@ -556,34 +663,23 @@ class DeepseekV4MLAAttention(nn.Module, AttentionLayerBase): self.swa_cache_layer: DeepseekV4SWACache = swa_cache_layer # Get vllm config for cache setup - vllm_config = get_current_vllm_config() self.max_num_batched_tokens = ( vllm_config.scheduler_config.max_num_batched_tokens ) self.max_model_len = vllm_config.model_config.max_model_len - # DeepseekV4 only supports fp8 kv-cache format for now. + + # Resolve the kv-cache dtype from the selected backend. FlashMLA uses + # the legacy UE8M0 paged uint8 (fp8_ds_mla) layout; FlashInfer uses a + # contiguous bf16 / per-tensor fp8 row. + backend = _resolve_dsv4_backend(vllm_config) kv_cache_dtype = cache_config.cache_dtype if cache_config is not None else "fp8" - - assert kv_cache_dtype.startswith("fp8"), ( - f"DeepseekV4 only supports fp8 kv-cache format for now, " - f"got {kv_cache_dtype}" + self.kv_cache_dtype, self.kv_cache_torch_dtype = _resolve_dsv4_kv_cache_dtype( + backend, kv_cache_dtype, cache_config ) - assert issubclass(self.get_attn_backend(), FlashMLASparseBackend), ( - "Only FlashMLA Sparse Attention backend is supported for DeepseekV4 for now" - ) - # FlashMLA Sparse Attention fp8 backend uses "fp8_ds_mla" kv-cache format - # Automatically convert fp8 kv-cache format to "fp8_ds_mla" - if ( - issubclass(self.get_attn_backend(), FlashMLASparseBackend) - and kv_cache_dtype.startswith("fp8") - and kv_cache_dtype != "fp8_ds_mla" - ): - assert cache_config is not None - cache_config.cache_dtype = "fp8_ds_mla" - kv_cache_dtype = "fp8_ds_mla" - logger.info_once("Using DeepSeek's fp8_ds_mla KV cache format.") - self.kv_cache_dtype = kv_cache_dtype + # Per-impl layer buffers (e.g. FlashInfer FP8 scale buffers). No-op for + # the FlashMLA / ROCm impls. + self.impl_cls.init_layer_buffers(self) # Register with compilation context for metadata lookup compilation_config = vllm_config.compilation_config @@ -602,14 +698,17 @@ class DeepseekV4MLAAttention(nn.Module, AttentionLayerBase): self.compress_ratio <= 1 ): # SWA part. Allocated separately as DeepseekV4SWACache. return None + # FlashMLA uses the UE8M0 paged uint8 layout (576B aligned); FlashInfer + # uses a contiguous bf16 / per-tensor fp8 cache with no extra alignment. + is_flashmla = self.kv_cache_dtype == "fp8_ds_mla" return MLAAttentionSpec( block_size=vllm_config.cache_config.block_size, num_kv_heads=1, head_size=self.head_dim, - dtype=torch.uint8, + dtype=torch.uint8 if is_flashmla else self.kv_cache_torch_dtype, compress_ratio=self.compress_ratio, cache_dtype_str=self.kv_cache_dtype, - alignment=576, # NOTE: FlashMLA requires 576B alignment + alignment=576 if is_flashmla else None, # FlashMLA needs 576B model_version="deepseek_v4", ) diff --git a/vllm/models/deepseek_v4/common/ops/__init__.py b/vllm/models/deepseek_v4/common/ops/__init__.py index dc6f3c608d9..ff6ee22996d 100644 --- a/vllm/models/deepseek_v4/common/ops/__init__.py +++ b/vllm/models/deepseek_v4/common/ops/__init__.py @@ -2,6 +2,7 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from .cache_utils import ( + build_flashinfer_mixed_sparse_indices, combine_topk_swa_indices, compute_global_topk_indices_and_lens, dequantize_and_gather_k_cache, @@ -15,6 +16,7 @@ from .save_partial_states import save_partial_states __all__ = [ "MXFP4_BLOCK_SIZE", + "build_flashinfer_mixed_sparse_indices", "combine_topk_swa_indices", "compute_global_topk_indices_and_lens", "dequantize_and_gather_k_cache", diff --git a/vllm/models/deepseek_v4/common/ops/cache_utils.py b/vllm/models/deepseek_v4/common/ops/cache_utils.py index ac66751e311..8adf219dbbe 100644 --- a/vllm/models/deepseek_v4/common/ops/cache_utils.py +++ b/vllm/models/deepseek_v4/common/ops/cache_utils.py @@ -592,3 +592,308 @@ def _combine_topk_swa_indices_kernel( combined_len = topk_len + swa_len tl.store(combined_lens_ptr + token_idx, combined_len) + + +def build_flashinfer_mixed_sparse_indices( + decode_swa_indices: torch.Tensor, + decode_compressed_indices: torch.Tensor | None, + decode_compressed_topk_lens: torch.Tensor | None, + prefill_topk_indices: torch.Tensor, + query_start_loc: torch.Tensor, + seq_lens: torch.Tensor, + token_to_req_indices: torch.Tensor, + swa_block_table: torch.Tensor, + swa_block_size: int, + compressed_block_table: torch.Tensor | None, + compressed_block_size: int, + window_size: int, + compress_ratio: int, + topk: int, + decode_compressed_indices_are_local: bool = False, + decode_is_valid_token: torch.Tensor | None = None, +) -> tuple[torch.Tensor, torch.Tensor]: + """Build the FlashInfer DSV4 sparse-index matrix for decode-first batches. + + Produces ``sparse_indices`` of shape ``[num_tokens, window_size + + padded_topk]`` (the first ``window_size`` columns are SWA slot ids, the rest + are compressed/top-k slot ids) and ``sparse_topk_lens`` (active length per + token). Decode tokens read precomputed SWA/compressed indices; prefill tokens + derive their SWA window from the position and translate local compressed + indices to global slots via the block tables. + """ + assert decode_swa_indices.dtype == torch.int32 + assert decode_swa_indices.dim() == 2 + assert decode_swa_indices.shape[-1] == window_size + if decode_compressed_topk_lens is not None: + assert decode_compressed_topk_lens.dtype == torch.int32 + assert prefill_topk_indices.dtype == torch.int32 + assert prefill_topk_indices.dim() == 2 + assert query_start_loc.dtype == torch.int32 + assert seq_lens.dtype == torch.int32 + assert token_to_req_indices.dtype == torch.int32 + assert swa_block_table.dtype == torch.int32 + + num_decode_tokens = decode_swa_indices.shape[0] + num_prefill_tokens = prefill_topk_indices.shape[0] + num_tokens = num_decode_tokens + num_prefill_tokens + assert token_to_req_indices.shape[0] >= num_tokens + if decode_compressed_topk_lens is not None: + assert decode_compressed_topk_lens.shape[0] >= num_decode_tokens + + decode_compressed_topk = 0 + if decode_compressed_indices is None: + decode_compressed_indices = prefill_topk_indices + else: + assert decode_compressed_indices.dtype == torch.int32 + assert decode_compressed_indices.dim() == 2 + assert decode_compressed_indices.shape[0] == num_decode_tokens + decode_compressed_topk = decode_compressed_indices.shape[-1] + if decode_compressed_topk > 0 and decode_compressed_indices_are_local: + assert decode_is_valid_token is not None + assert decode_is_valid_token.dtype == torch.bool + assert decode_is_valid_token.shape[0] >= num_decode_tokens + else: + decode_is_valid_token = token_to_req_indices + + if compressed_block_table is None: + compressed_block_table = swa_block_table + assert compressed_block_table.dtype == torch.int32 + has_decode_compressed_lens = decode_compressed_topk_lens is not None + if decode_compressed_topk_lens is None: + decode_compressed_topk_lens = token_to_req_indices + + # The FlashInfer TRTLLM-gen sparse-MLA kernels require every per-token topk + # index row to start on a 16-byte boundary: the kernel loads the compressed + # indices with 128-bit (16-byte) vectorized loads, so a misaligned row would + # fault or read across rows. 16 bytes = 4 int32 indices, so round the topk + # width (and hence the row stride, since the SWA columns are fixed-width) up + # to a multiple of 4. The extra columns are filled with -1 (invalid) and bounded + # by ``sparse_topk_lens``, so padding never changes the attention result. + padded_topk = max(topk, decode_compressed_topk) + padded_topk = (padded_topk + 3) // 4 * 4 + sparse_indices = torch.empty( + (num_tokens, window_size + padded_topk), + dtype=torch.int32, + device=decode_swa_indices.device, + ) + sparse_topk_lens = torch.empty( + num_tokens, dtype=torch.int32, device=decode_swa_indices.device + ) + if num_tokens == 0: + return sparse_indices, sparse_topk_lens + + window_block_size = triton.next_power_of_2(max(window_size, 1)) + topk_block_size = triton.next_power_of_2(max(padded_topk, 1)) + max_block_size = max(window_block_size, topk_block_size) + num_warps = 4 if max_block_size >= 256 else 1 + + _build_flashinfer_mixed_sparse_indices_kernel[(num_tokens,)]( + sparse_indices, + sparse_indices.stride(0), + sparse_topk_lens, + decode_swa_indices, + decode_swa_indices.stride(0), + decode_compressed_indices, + decode_compressed_indices.stride(0), + decode_compressed_topk_lens, + decode_is_valid_token, + prefill_topk_indices, + prefill_topk_indices.stride(0), + query_start_loc, + seq_lens, + token_to_req_indices, + swa_block_table, + swa_block_table.stride(0), + swa_block_size, + compressed_block_table, + compressed_block_table.stride(0), + compressed_block_size, + NUM_DECODE_TOKENS=num_decode_tokens, + WINDOW_SIZE=window_size, + COMPRESS_RATIO=compress_ratio, + TOP_K=topk, + PADDED_TOP_K=padded_topk, + PREFILL_TOPK_STRIDE=prefill_topk_indices.shape[-1], + DECODE_COMPRESSED_TOPK=decode_compressed_topk, + DECODE_COMPRESSED_INDICES_ARE_LOCAL=decode_compressed_indices_are_local, + HAS_DECODE_COMPRESSED_LENS=has_decode_compressed_lens, + WINDOW_BLOCK_SIZE=window_block_size, + TOPK_BLOCK_SIZE=topk_block_size, + num_warps=num_warps, + ) + return sparse_indices, sparse_topk_lens + + +@triton.jit( + do_not_specialize=[ + "sparse_indices_stride", + "decode_swa_stride", + "decode_compressed_stride", + "prefill_topk_stride", + "swa_block_table_stride", + "swa_block_size", + "compressed_block_table_stride", + "compressed_block_size", + "NUM_DECODE_TOKENS", + "PREFILL_TOPK_STRIDE", + ] +) +def _build_flashinfer_mixed_sparse_indices_kernel( + sparse_indices_ptr, + sparse_indices_stride, + sparse_topk_lens_ptr, + decode_swa_indices_ptr, + decode_swa_stride, + decode_compressed_indices_ptr, + decode_compressed_stride, + decode_compressed_topk_lens_ptr, + decode_is_valid_token_ptr, + prefill_topk_indices_ptr, + prefill_topk_stride, + query_start_loc_ptr, + seq_lens_ptr, + token_to_req_indices_ptr, + swa_block_table_ptr, + swa_block_table_stride, + swa_block_size, + compressed_block_table_ptr, + compressed_block_table_stride, + compressed_block_size, + NUM_DECODE_TOKENS, + WINDOW_SIZE: tl.constexpr, + COMPRESS_RATIO: tl.constexpr, + TOP_K: tl.constexpr, + PADDED_TOP_K: tl.constexpr, + PREFILL_TOPK_STRIDE, + DECODE_COMPRESSED_TOPK: tl.constexpr, + DECODE_COMPRESSED_INDICES_ARE_LOCAL: tl.constexpr, + HAS_DECODE_COMPRESSED_LENS: tl.constexpr, + WINDOW_BLOCK_SIZE: tl.constexpr, + TOPK_BLOCK_SIZE: tl.constexpr, +): + token_idx = tl.program_id(0) + + if token_idx < NUM_DECODE_TOKENS: + for i in range(0, WINDOW_SIZE, WINDOW_BLOCK_SIZE): + offset = i + tl.arange(0, WINDOW_BLOCK_SIZE) + mask = offset < WINDOW_SIZE + values = tl.load( + decode_swa_indices_ptr + token_idx * decode_swa_stride + offset, + mask=mask, + other=-1, + ) + tl.store( + sparse_indices_ptr + token_idx * sparse_indices_stride + offset, + values, + mask=mask, + ) + + compressed_len = tl.zeros((), dtype=tl.int32) + for i in range(0, PADDED_TOP_K, TOPK_BLOCK_SIZE): + offset = i + tl.arange(0, TOPK_BLOCK_SIZE) + mask = offset < PADDED_TOP_K + values = tl.load( + decode_compressed_indices_ptr + + token_idx * decode_compressed_stride + + offset, + mask=offset < DECODE_COMPRESSED_TOPK, + other=-1, + ) + if DECODE_COMPRESSED_INDICES_ARE_LOCAL: + token_valid = tl.load(decode_is_valid_token_ptr + token_idx) + is_valid = values >= 0 + req_idx = tl.load(token_to_req_indices_ptr + token_idx) + block_indices = values // compressed_block_size + block_numbers = tl.load( + compressed_block_table_ptr + + req_idx * compressed_block_table_stride + + block_indices, + mask=mask & is_valid, + other=-1, + ) + block_offsets = values % compressed_block_size + values = block_numbers * compressed_block_size + block_offsets + values = tl.where(is_valid, values, -1) + compressed_len += tl.sum((is_valid & token_valid).to(tl.int32), axis=0) + tl.store( + sparse_indices_ptr + + token_idx * sparse_indices_stride + + WINDOW_SIZE + + offset, + values, + mask=mask, + ) + + if DECODE_COMPRESSED_TOPK == 0: + compressed_len = tl.zeros((), dtype=tl.int32) + elif not DECODE_COMPRESSED_INDICES_ARE_LOCAL: + if HAS_DECODE_COMPRESSED_LENS: + compressed_len = tl.load(decode_compressed_topk_lens_ptr + token_idx) + else: + compressed_len = tl.full((), DECODE_COMPRESSED_TOPK, dtype=tl.int32) + + tl.store(sparse_topk_lens_ptr + token_idx, WINDOW_SIZE + compressed_len) + return + + prefill_idx = token_idx - NUM_DECODE_TOKENS + req_idx = tl.load(token_to_req_indices_ptr + token_idx) + query_start = tl.load(query_start_loc_ptr + req_idx) + query_end = tl.load(query_start_loc_ptr + req_idx + 1) + query_len = query_end - query_start + seq_len = tl.load(seq_lens_ptr + req_idx) + start_pos = seq_len - query_len + token_idx_in_query = token_idx - query_start + pos = start_pos + token_idx_in_query + swa_len = tl.minimum(pos + 1, WINDOW_SIZE) + swa_start_pos = pos - swa_len + 1 + topk_len = tl.minimum((pos + 1) // COMPRESS_RATIO, TOP_K) + + for i in range(0, WINDOW_SIZE, WINDOW_BLOCK_SIZE): + offset = i + tl.arange(0, WINDOW_BLOCK_SIZE) + mask = offset < WINDOW_SIZE + pos_offset = swa_start_pos + offset + block_indices = pos_offset // swa_block_size + block_numbers = tl.load( + swa_block_table_ptr + req_idx * swa_block_table_stride + block_indices, + mask=mask & (offset < swa_len), + other=-1, + ) + block_offsets = pos_offset % swa_block_size + slot_ids = block_numbers * swa_block_size + block_offsets + slot_ids = tl.where(offset < swa_len, slot_ids, -1) + tl.store( + sparse_indices_ptr + token_idx * sparse_indices_stride + offset, + slot_ids, + mask=mask, + ) + + for i in range(0, PADDED_TOP_K, TOPK_BLOCK_SIZE): + offset = i + tl.arange(0, TOPK_BLOCK_SIZE) + mask = offset < PADDED_TOP_K + local_idx = tl.load( + prefill_topk_indices_ptr + prefill_idx * prefill_topk_stride + offset, + mask=(offset < PREFILL_TOPK_STRIDE) & (offset < topk_len), + other=-1, + ) + is_valid = local_idx >= 0 + block_indices = local_idx // compressed_block_size + block_numbers = tl.load( + compressed_block_table_ptr + + req_idx * compressed_block_table_stride + + block_indices, + mask=mask & is_valid, + other=-1, + ) + block_offsets = local_idx % compressed_block_size + slot_ids = block_numbers * compressed_block_size + block_offsets + slot_ids = tl.where((offset < topk_len) & is_valid, slot_ids, -1) + tl.store( + sparse_indices_ptr + + token_idx * sparse_indices_stride + + WINDOW_SIZE + + offset, + slot_ids, + mask=mask, + ) + + tl.store(sparse_topk_lens_ptr + token_idx, WINDOW_SIZE + topk_len) diff --git a/vllm/models/deepseek_v4/compressor.py b/vllm/models/deepseek_v4/compressor.py index f36dc8f1762..2ceec0789a1 100644 --- a/vllm/models/deepseek_v4/compressor.py +++ b/vllm/models/deepseek_v4/compressor.py @@ -155,13 +155,17 @@ class CompressorStateCache(torch.nn.Module, AttentionLayerBase): raise ValueError(f"Invalid compress ratio: {compress_ratio}") def get_kv_cache_spec(self, vllm_config: VllmConfig) -> KVCacheSpec: + # FlashMLA's UE8M0 paged layout needs 576B alignment; the FlashInfer + # full-cache path shares state pages with contiguous KV pages, so + # padding would break page matching. + is_flashmla = vllm_config.cache_config.cache_dtype == "fp8_ds_mla" return SlidingWindowMLASpec( # only has one vector instead of K + V block_size=self.block_size, num_kv_heads=1, head_size=self.state_dim, dtype=self.dtype, sliding_window=self.sliding_window, - alignment=576, # NOTE: FlashMLA requires 576B alignment + alignment=576 if is_flashmla else None, ) def forward(self): ... @@ -333,26 +337,40 @@ class DeepseekCompressor(nn.Module): # - position used: (positions // compress_ratio) * compress_ratio cos_sin_cache = rotary_emb.cos_sin_cache k_cache_metadata = cast(Any, attn_metadata[self.k_cache_prefix]) - kv_cache = self._static_forward_context[self.k_cache_prefix].kv_cache + k_cache_layer = self._static_forward_context[self.k_cache_prefix] + kv_cache = k_cache_layer.kv_cache - if current_platform.is_cuda(): - # NVIDIA GPUs. - if self.head_dim == 512: - from .nvidia.ops.sparse_attn_compress_cutedsl import ( - compress_norm_rope_store_cutedsl, - ) + # FlashInfer V4 reads a contiguous bf16 / per-tensor fp8 cache row; the + # legacy FlashMLA path uses the UE8M0 paged uint8 layout. + store_full_kv = self.head_dim == 512 and kv_cache.dtype != torch.uint8 + store_full_fp8 = kv_cache.dtype == torch.float8_e4m3fn + fp8_scale = ( + getattr(k_cache_layer, "_flashinfer_fp8_kv_scale", None) + if store_full_fp8 + else None + ) - # Main compressor path. - # Use a cutedsl kernel for better performance. - compress_norm_rope_store_fn = compress_norm_rope_store_cutedsl - else: - # Indexer path (head_dim == 128). - # Use a triton kernel. - compress_norm_rope_store_fn = compress_norm_rope_store_triton + # cutedsl (head=512) accepts the full-cache flags; triton (indexer/AMD) + # does not, so the two callables have different signatures. + compress_norm_rope_store_fn: Any + if current_platform.is_cuda() and self.head_dim == 512: + from .nvidia.ops.sparse_attn_compress_cutedsl import ( + compress_norm_rope_store_cutedsl, + ) + + # head=512 on CUDA always uses cutedsl, for both the legacy UE8M0 + # layout and the FlashInfer full-cache layout. The full-cache flags + # are consumed only here. + compress_norm_rope_store_fn = compress_norm_rope_store_cutedsl + extra_kwargs: dict[str, Any] = dict( + store_full_kv=store_full_kv, + store_full_fp8=store_full_fp8, + fp8_scale=fp8_scale, + ) else: - # AMD GPUs. - # Always use a triton kernel. + # Indexer path (head_dim == 128) or AMD: triton, legacy UE8M0 only. compress_norm_rope_store_fn = compress_norm_rope_store_triton + extra_kwargs = {} compress_norm_rope_store_fn( state_cache=state_cache, @@ -377,4 +395,5 @@ class DeepseekCompressor(nn.Module): quant_block=self._quant_block, token_stride=self._token_stride, scale_dim=self._scale_dim, + **extra_kwargs, ) diff --git a/vllm/models/deepseek_v4/nvidia/flashinfer_sparse.py b/vllm/models/deepseek_v4/nvidia/flashinfer_sparse.py new file mode 100644 index 00000000000..71ea4fe506e --- /dev/null +++ b/vllm/models/deepseek_v4/nvidia/flashinfer_sparse.py @@ -0,0 +1,407 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""DeepSeek V4 FlashInfer TRTLLM-gen sparse MLA backend. + +Uses FlashInfer's public ``trtllm_batch_decode_sparse_mla_dsv4`` launcher with a +contiguous bf16 / per-tensor FP8 KV cache. Shares the V4 sparse-index pipeline +(SWA cache + compressor + indexer, 256-token blocks, head_size 512) with the +FlashMLA V4 backend; only the attention forward differs. +""" + +from typing import TYPE_CHECKING, cast + +import torch + +from vllm.forward_context import get_forward_context +from vllm.models.deepseek_v4.common.ops import ( + build_flashinfer_mixed_sparse_indices, +) +from vllm.models.deepseek_v4.nvidia.flashmla import ( + DeepseekV4FlashMLASparseBackend, + DeepseekV4SparseMLAAttentionImpl, +) +from vllm.utils.flashinfer import flashinfer_trtllm_batch_decode_sparse_mla_dsv4 +from vllm.v1.attention.backends.mla.flashmla_sparse import FlashMLASparseMetadata + +if TYPE_CHECKING: + from vllm.models.deepseek_v4.attention import DeepseekV4MLAAttention + from vllm.v1.attention.backends.mla.sparse_swa import DeepseekSparseSWAMetadata + +# 128 MB TRTLLM-gen workspace, allocated once per device and zero-initialized +# (required for first use). Reused across all FlashInfer V4 layers. +_FLASHINFER_DSV4_WORKSPACE_BUFFER_SIZE = 128 * 1024 * 1024 +_flashinfer_dsv4_workspace_by_device: dict[torch.device, torch.Tensor] = {} + + +def _get_flashinfer_dsv4_workspace(device: torch.device) -> torch.Tensor: + workspace = _flashinfer_dsv4_workspace_by_device.get(device) + if workspace is None: + workspace = torch.zeros( + _FLASHINFER_DSV4_WORKSPACE_BUFFER_SIZE, + dtype=torch.uint8, + device=device, + ) + _flashinfer_dsv4_workspace_by_device[device] = workspace + return workspace + + +class DeepseekV4FlashInferMLASparseBackend(DeepseekV4FlashMLASparseBackend): + """Shares the FlashMLA V4 metadata/cache pipeline; swaps the attention impl. + + Inheriting from the FlashMLA V4 backend reuses its ``FlashMLASparseMetadata`` + builder (which the V4 sparse-index pipeline needs — the V3.2 FlashInfer + builder lacks the ``c128a_*`` fields), 256-token blocks, head_size 512, and + the contiguous (num_blocks, block_size, 512) cache shape for non-``fp8_ds_mla`` + dtypes. + """ + + @staticmethod + def get_name() -> str: + return "FLASHINFER_MLA_SPARSE_DSV4" + + @staticmethod + def get_impl_cls() -> type["DeepseekV4FlashInferMLASparseImpl"]: + return DeepseekV4FlashInferMLASparseImpl + + +class DeepseekV4FlashInferMLASparseImpl(DeepseekV4SparseMLAAttentionImpl): + """FlashInfer TRTLLM-gen sparse MLA implementation for DeepSeek V4.""" + + backend_cls = DeepseekV4FlashInferMLASparseBackend + + @classmethod + def get_padded_num_q_heads(cls, num_heads: int) -> int: + # FP8 decode kernel only supports h_q = 64 or 128. + if num_heads > 128: + raise ValueError( + f"DeepseekV4 Flashinfer MLA Sparse does not support {num_heads} heads " + "(FP8 decode kernel requires h_q in {64, 128})." + ) + return 64 if num_heads <= 64 else 128 + + @classmethod + def init_layer_buffers(cls, layer: "DeepseekV4MLAAttention") -> None: + # Per-tensor FP8 scale buffers + precomputed scalar BMM scales. Only the + # per-tensor FP8 cache path consumes these; bf16 reads ``layer.scale``. + if layer.kv_cache_torch_dtype != torch.float8_e4m3fn: + return + # TODO: load real per-tensor Q/KV scales from the checkpoint; unit + # scales until the scale tensor names are wired. + fp8_q_scale = 1.0 + fp8_kv_scale = 1.0 + layer.register_buffer( + "_flashinfer_fp8_q_scale", + torch.tensor([fp8_q_scale], dtype=torch.float32), + persistent=False, + ) + layer.register_buffer( + "_flashinfer_fp8_q_scale_inv", + torch.tensor([1.0 / fp8_q_scale], dtype=torch.float32), + persistent=False, + ) + layer.register_buffer( + "_flashinfer_fp8_kv_scale", + torch.tensor([fp8_kv_scale], dtype=torch.float32), + persistent=False, + ) + # TRTLLM-gen takes scalar scale args on a distinct (correct) C++ path + # vs 1-elem tensors, so these are Python floats. bmm1 folds the softmax + # scale and the Q/KV per-tensor scales; bmm2 is the KV scale. + layer._flashinfer_fp8_bmm1_scale = layer.scale * fp8_q_scale * fp8_kv_scale + layer._flashinfer_fp8_bmm2_scale = fp8_kv_scale + + @classmethod + def forward_mqa( # type: ignore[override] + cls, + layer: "DeepseekV4MLAAttention", + q: torch.Tensor, + kv: torch.Tensor, + positions: torch.Tensor, + output: torch.Tensor, + ) -> None: + # The TRTLLM-gen kernel requires h_q in {64, 128}, so the output buffer + # is allocated at the padded head count while q arrives at the local + # head count; _forward pads q to match before the launcher. + assert output.shape[0] == q.shape[0] and output.shape[-1] == q.shape[-1], ( + f"output buffer shape {output.shape} incompatible with q shape {q.shape}" + ) + assert output.shape[1] >= q.shape[1], ( + f"output heads {output.shape[1]} must be >= q heads {q.shape[1]}" + ) + # Per-tensor FP8 q produces a bf16 attention output. + expected_output_dtype = ( + torch.bfloat16 if q.dtype == torch.float8_e4m3fn else q.dtype + ) + assert output.dtype == expected_output_dtype, ( + f"output dtype {output.dtype} must match expected {expected_output_dtype} " + f"for q dtype {q.dtype}" + ) + + forward_context = get_forward_context() + attn_metadata = forward_context.attn_metadata + if attn_metadata is None: + # Warmup dummy run: FlashInfer reads the cache directly and lazily + # allocates its workspace, so nothing to reserve here. + output.zero_() + return + + assert isinstance(attn_metadata, dict) + flashmla_metadata = cast( + FlashMLASparseMetadata | None, attn_metadata.get(layer.prefix) + ) + swa_metadata = cast( + "DeepseekSparseSWAMetadata | None", + attn_metadata.get(layer.swa_cache_layer.prefix), + ) + assert swa_metadata is not None + + swa_only = layer.compress_ratio <= 1 + # SWA-only layers don't allocate their own compressed KV cache. + self_kv_cache = layer.kv_cache if not swa_only else None + swa_kv_cache = layer.swa_cache_layer.kv_cache + + cls._forward( + layer=layer, + q=q, + kv_cache=self_kv_cache, + swa_k_cache=swa_kv_cache, + swa_metadata=swa_metadata, + attn_metadata=flashmla_metadata, + swa_only=swa_only, + output=output, + ) + + @classmethod + def _build_sparse_index_metadata( + cls, + layer: "DeepseekV4MLAAttention", + kv_cache: torch.Tensor | None, + swa_k_cache: torch.Tensor, + swa_metadata: "DeepseekSparseSWAMetadata", + attn_metadata: FlashMLASparseMetadata | None, + swa_only: bool, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Build the combined sparse-index tensors for the mixed batch. + + Returns ``(compressed_kv_cache, seq_lens, sparse_indices, + sparse_topk_lens)``. + """ + num_decodes = swa_metadata.num_decodes + num_prefills = swa_metadata.num_prefills + num_decode_tokens = swa_metadata.num_decode_tokens + num_prefill_tokens = swa_metadata.num_prefill_tokens + num_reqs = num_decodes + num_prefills + num_tokens = num_decode_tokens + num_prefill_tokens + + assert swa_metadata.seq_lens is not None + assert swa_metadata.query_start_loc is not None + assert swa_metadata.token_to_req_indices is not None + assert swa_metadata.decode_swa_indices is not None + assert swa_metadata.block_table is not None + + decode_swa_indices = swa_metadata.decode_swa_indices.reshape( + num_decode_tokens, layer.window_size + ) + decode_compressed_topk_lens = None + decode_compressed_indices_are_local = False + decode_is_valid_token = None + + if swa_only: + assert layer.topk_indices_buffer is not None + compressed_kv_cache = swa_k_cache + decode_compressed_indices = None + prefill_topk_indices = layer.topk_indices_buffer[ + num_decode_tokens:num_tokens, :0 + ] + compressed_block_table = None + compressed_block_size = swa_metadata.block_size + top_k = 0 + else: + assert kv_cache is not None + assert attn_metadata is not None + compressed_kv_cache = kv_cache + compressed_block_table = attn_metadata.block_table[:num_reqs] + compressed_block_size = attn_metadata.block_size // layer.compress_ratio + + if layer.compress_ratio == 4: + assert layer.topk_indices_buffer is not None + if num_prefill_tokens > 0: + prefill_topk_indices = layer.topk_indices_buffer[ + num_decode_tokens:num_tokens + ] + top_k = prefill_topk_indices.shape[-1] + else: + prefill_topk_indices = layer.topk_indices_buffer[:0, :0] + top_k = 0 + + decode_compressed_indices_are_local = True + assert swa_metadata.is_valid_token is not None + decode_is_valid_token = swa_metadata.is_valid_token[:num_decode_tokens] + if num_decode_tokens > 0: + decode_compressed_indices = layer.topk_indices_buffer[ + :num_decode_tokens + ] + else: + # Keep the logical width aligned with the mixed-batch case so + # pure-prefill steps reuse the same Triton specialization. + decode_compressed_indices = prefill_topk_indices[:0] + else: + if num_prefill_tokens > 0: + assert attn_metadata.c128a_prefill_topk_indices is not None + prefill_topk_indices = attn_metadata.c128a_prefill_topk_indices + top_k = prefill_topk_indices.shape[-1] + else: + prefill_topk_indices = decode_swa_indices[:0, :0] + top_k = 0 + + if num_decode_tokens > 0: + assert attn_metadata.c128a_global_decode_topk_indices is not None + assert attn_metadata.c128a_decode_topk_lens is not None + decode_compressed_indices = ( + attn_metadata.c128a_global_decode_topk_indices.view( + num_decode_tokens, -1 + ) + ) + decode_compressed_topk_lens = attn_metadata.c128a_decode_topk_lens + if num_prefill_tokens == 0: + prefill_topk_indices = decode_compressed_indices[:0, :0] + else: + decode_compressed_indices = prefill_topk_indices[:0] + decode_compressed_topk_lens = swa_metadata.seq_lens[:0] + + query_start_loc = swa_metadata.query_start_loc[: num_reqs + 1] + seq_lens = swa_metadata.seq_lens[:num_reqs] + assert seq_lens.dtype == torch.int32 + sparse_indices, sparse_topk_lens = build_flashinfer_mixed_sparse_indices( + decode_swa_indices, + decode_compressed_indices, + decode_compressed_topk_lens, + prefill_topk_indices[:num_prefill_tokens], + query_start_loc, + seq_lens, + swa_metadata.token_to_req_indices[:num_tokens], + swa_metadata.block_table[:num_reqs], + swa_metadata.block_size, + compressed_block_table, + compressed_block_size, + layer.window_size, + layer.compress_ratio, + top_k, + decode_compressed_indices_are_local=decode_compressed_indices_are_local, + decode_is_valid_token=decode_is_valid_token, + ) + return compressed_kv_cache, seq_lens, sparse_indices, sparse_topk_lens + + @classmethod + def _forward( + cls, + layer: "DeepseekV4MLAAttention", + q: torch.Tensor, + kv_cache: torch.Tensor | None, + swa_k_cache: torch.Tensor, + swa_metadata: "DeepseekSparseSWAMetadata", + attn_metadata: FlashMLASparseMetadata | None, + swa_only: bool, + output: torch.Tensor, + ) -> None: + assert layer.kv_cache_torch_dtype in (torch.bfloat16, torch.float8_e4m3fn) + num_decodes = swa_metadata.num_decodes + num_prefills = swa_metadata.num_prefills + num_decode_tokens = swa_metadata.num_decode_tokens + num_prefill_tokens = swa_metadata.num_prefill_tokens + num_reqs = num_decodes + num_prefills + num_tokens = num_decode_tokens + num_prefill_tokens + if num_tokens == 0: + return + + ( + compressed_kv_cache, + seq_lens, + sparse_indices, + sparse_topk_lens, + ) = cls._build_sparse_index_metadata( + layer=layer, + kv_cache=kv_cache, + swa_k_cache=swa_k_cache, + swa_metadata=swa_metadata, + attn_metadata=attn_metadata, + swa_only=swa_only, + ) + + # CUDA graph execution can pad q/output past the scheduled token count; + # restrict to the real tokens (the launcher validates sparse indices). + query = q[:num_tokens] + output = output[:num_tokens] + bmm1_scale: float | torch.Tensor = layer.scale + bmm2_scale: float | torch.Tensor = 1.0 + if layer.kv_cache_torch_dtype == torch.float8_e4m3fn: + assert query.dtype == torch.float8_e4m3fn + bmm1_scale = layer._flashinfer_fp8_bmm1_scale + bmm2_scale = layer._flashinfer_fp8_bmm2_scale + else: + assert query.dtype == torch.bfloat16 + query = query.contiguous() + + # The TRTLLM-gen sparse-MLA kernel requires h_q in {64, 128}; zero-pad + # the query heads to the allocated output head count. Padded heads attend + # to the shared KV and are sliced off downstream (output is padded too). + padded_heads = output.shape[1] + if query.shape[1] < padded_heads: + padded_query = query.new_zeros( + (query.shape[0], padded_heads, query.shape[2]) + ) + padded_query[:, : query.shape[1], :] = query + query = padded_query + + workspace = _get_flashinfer_dsv4_workspace(q.device) + query_start_loc = swa_metadata.query_start_loc + query_start_loc_cpu = swa_metadata.query_start_loc_cpu + assert query_start_loc is not None and query_start_loc_cpu is not None + + # Keep Perkz's two-call decode/prefill split: the TRTLLM-gen launcher is + # tuned for uniform-q batches, and collapsing the mixed batch into a + # single call is the suspected source of the prior IMA. + if num_decode_tokens > 0: + decode_cu = query_start_loc[: num_decodes + 1] + decode_cu_cpu = query_start_loc_cpu[: num_decodes + 1] + decode_lens_cpu = decode_cu_cpu[1:] - decode_cu_cpu[:-1] + flashinfer_trtllm_batch_decode_sparse_mla_dsv4( + query=query[:num_decode_tokens], + swa_kv_cache=swa_k_cache, + workspace_buffer=workspace, + sparse_indices=sparse_indices[:num_decode_tokens], + compressed_kv_cache=compressed_kv_cache, + sparse_topk_lens=sparse_topk_lens[:num_decode_tokens], + seq_lens=seq_lens[:num_decodes], + out=output[:num_decode_tokens], + bmm1_scale=bmm1_scale, + bmm2_scale=bmm2_scale, + sinks=layer.attn_sink, + cum_seq_lens_q=decode_cu, + max_q_len=int(decode_lens_cpu.max().item()), + ) + + if num_prefill_tokens > 0: + # The prefill query view re-anchors at offset 0, so rebase the + # cumulative query offsets to start at 0. + prefill_cu = ( + query_start_loc[num_decodes : num_reqs + 1] + - query_start_loc[num_decodes] + ) + prefill_cu_cpu = query_start_loc_cpu[num_decodes : num_reqs + 1] + prefill_lens_cpu = prefill_cu_cpu[1:] - prefill_cu_cpu[:-1] + flashinfer_trtllm_batch_decode_sparse_mla_dsv4( + query=query[num_decode_tokens:num_tokens], + swa_kv_cache=swa_k_cache, + workspace_buffer=workspace, + sparse_indices=sparse_indices[num_decode_tokens:num_tokens], + compressed_kv_cache=compressed_kv_cache, + sparse_topk_lens=sparse_topk_lens[num_decode_tokens:num_tokens], + seq_lens=seq_lens[num_decodes:num_reqs], + out=output[num_decode_tokens:num_tokens], + bmm1_scale=bmm1_scale, + bmm2_scale=bmm2_scale, + sinks=layer.attn_sink, + cum_seq_lens_q=prefill_cu, + max_q_len=int(prefill_lens_cpu.max().item()), + ) diff --git a/vllm/models/deepseek_v4/nvidia/flashmla.py b/vllm/models/deepseek_v4/nvidia/flashmla.py index 5c8b08d4c12..e9b9c678306 100644 --- a/vllm/models/deepseek_v4/nvidia/flashmla.py +++ b/vllm/models/deepseek_v4/nvidia/flashmla.py @@ -75,6 +75,15 @@ class DeepseekV4SparseMLAAttentionImpl(SparseMLAAttentionImpl[FlashMLASparseMeta """ raise NotImplementedError + @classmethod + def init_layer_buffers(cls, layer: "DeepseekV4MLAAttention") -> None: + """Register impl-specific buffers on the layer at construction. + + No-op by default; FlashInfer overrides this to register its per-tensor + FP8 scale buffers. + """ + return None + class DeepseekV4FlashMLASparseBackend(FlashMLASparseBackend): @staticmethod @@ -83,7 +92,7 @@ class DeepseekV4FlashMLASparseBackend(FlashMLASparseBackend): @staticmethod def get_name() -> str: - return "V4_FLASHMLA_SPARSE" + return "FLASHMLA_SPARSE_DSV4" @staticmethod def get_impl_cls() -> type["DeepseekV4SparseMLAAttentionImpl"]: diff --git a/vllm/models/deepseek_v4/nvidia/ops/sparse_attn_compress_cutedsl.py b/vllm/models/deepseek_v4/nvidia/ops/sparse_attn_compress_cutedsl.py index 18e596029ff..4ff4b232d10 100644 --- a/vllm/models/deepseek_v4/nvidia/ops/sparse_attn_compress_cutedsl.py +++ b/vllm/models/deepseek_v4/nvidia/ops/sparse_attn_compress_cutedsl.py @@ -508,6 +508,426 @@ class SparseAttnCompressNormRopeStoreC4Kernel: ) +class SparseAttnCompressNormRopeStoreFullC4Kernel( + SparseAttnCompressNormRopeStoreC4Kernel +): + def __init__( + self, + head_size: int, + state_width: int, + rope_head_dim: int, + fp8_max: float, + quant_block: int, + token_stride: int, + scale_dim: int, + compress_ratio: int, + overlap: bool, + store_full_fp8: bool = False, + ): + super().__init__( + head_size, + state_width, + rope_head_dim, + fp8_max, + quant_block, + token_stride, + scale_dim, + compress_ratio, + overlap, + ) + self.store_full_fp8 = store_full_fp8 + + @cute.jit + def __call__( + self, + state_cache: cute.Tensor, + token_to_req_indices: cute.Tensor, + positions: cute.Tensor, + slot_mapping: cute.Tensor, + block_table: cute.Tensor, + block_size: Int64, + rms_norm_weight: cute.Tensor, + rms_norm_eps: Float32, + cos_sin_cache: cute.Tensor, + k_cache: cute.Tensor, + kv_slot_mapping: cute.Tensor, + kv_cache_block_size: Int64, + fp8_scale: cute.Tensor, + stream: CUstream, + ): + grid = (slot_mapping.shape[0], 1, 1) + self.kernel( + state_cache, + token_to_req_indices, + positions, + slot_mapping, + block_table, + block_size, + rms_norm_weight, + rms_norm_eps, + cos_sin_cache, + k_cache, + kv_slot_mapping, + kv_cache_block_size, + fp8_scale, + ).launch(grid=grid, block=(self.tb_size, 1, 1), stream=stream) + + @cute.kernel + def kernel( + self, + state_cache: cute.Tensor, + token_to_req_indices: cute.Tensor, + positions: cute.Tensor, + slot_mapping: cute.Tensor, + block_table: cute.Tensor, + block_size: Int64, + rms_norm_weight: cute.Tensor, + rms_norm_eps: Float32, + cos_sin_cache: cute.Tensor, + k_cache: cute.Tensor, + kv_slot_mapping: cute.Tensor, + kv_cache_block_size: Int64, + fp8_scale: cute.Tensor, + ): + token_idx, _, _ = cute.arch.block_idx() + tid, _, _ = cute.arch.thread_idx() + warp_id = cute.arch.make_warp_uniform(tid // 32) + lane_id = tid % 32 + group_lane = lane_id % self.lanes_per_group + group_idx = warp_id * self.groups_per_warp + lane_id // self.lanes_per_group + elem_base = group_idx * self.quant_block + group_lane * self.elems_per_lane + + slot_id = slot_mapping[token_idx] + has_position = token_idx < positions.shape[0] + position = Int64(0) + if has_position: + position = positions[token_idx] + boundary = has_position and ( + (position + Int64(1)) % Int64(self.compress_ratio) == Int64(0) + ) + has_req_idx = token_idx < token_to_req_indices.shape[0] + has_kv_slot_idx = token_idx < kv_slot_mapping.shape[0] + kv_slot_idx = Int64(-1) + if has_kv_slot_idx: + kv_slot_idx = kv_slot_mapping[token_idx] + active = ( + slot_id >= Int64(0) and has_req_idx and boundary and kv_slot_idx >= Int64(0) + ) + + if active: + req_idx = token_to_req_indices[token_idx] + start = position - Int64(self.window - 1) + + smem = cutlass.utils.SmemAllocator() + s_block_numbers = smem.allocate_tensor( + Int32, cute.make_layout((self.window,)), byte_alignment=4 + ) + partial_sums = smem.allocate_tensor( + Float32, cute.make_layout((self.num_warps,)), byte_alignment=4 + ) + rrms_shared = smem.allocate_tensor( + Float32, cute.make_layout((1,)), byte_alignment=4 + ) + + for row in cutlass.range_constexpr(self.window): + pos = start + Int64(row) + if tid == row: + block_number_i32 = Int32(0) + if pos >= Int64(0): + block_index = pos // block_size + block_number_i32 = block_table[req_idx, block_index] + s_block_numbers[row] = block_number_i32 + cute.arch.sync_threads() + + local_max = cute.make_rmem_tensor((self.elems_per_lane,), Float32) + local_sum = cute.make_rmem_tensor((self.elems_per_lane,), Float32) + local_product = cute.make_rmem_tensor((self.elems_per_lane,), Float32) + + for e in cutlass.range_constexpr(self.elems_per_lane): + local_max[e] = -Float32.inf + local_sum[e] = Float32(0.0) + local_product[e] = Float32(0.0) + + cp_f32x4 = cute.make_copy_atom( + cute.nvgpu.CopyUniversalOp(), Float32, num_bits_per_copy=128 + ) + copy_layout = cute.make_layout( + (self.copy_chunks, self.copy_elems), + stride=(self.copy_elems, 1), + ) + kv_vals = cute.make_rmem_tensor(copy_layout, Float32) + score_vals = cute.make_rmem_tensor(copy_layout, Float32) + + for row in cutlass.range_constexpr(self.window): + pos = start + Int64(row) + if pos >= Int64(0): + block_index = pos // block_size + block_offset = pos - block_index * block_size + block_number = s_block_numbers[row].to(Int64) + head_offset = Int64((row // self.compress_ratio) * self.head_dim) + row_tensor = state_cache[block_number, block_offset, None] + for chunk in cutlass.range_constexpr(self.copy_chunks): + copy_elem = const_expr(chunk * self.copy_elems) + col_tile = ( + head_offset + (elem_base + Int32(copy_elem)).to(Int64) + ) // Int64(self.copy_elems) + kv_src = cute.local_tile( + row_tensor, + tiler=(self.copy_elems,), + coord=(col_tile,), + ) + score_src = cute.local_tile( + row_tensor, + tiler=(self.copy_elems,), + coord=( + col_tile + Int64(self.state_width // self.copy_elems), + ), + ) + cute.copy(cp_f32x4, kv_src, kv_vals[chunk, None]) + cute.copy(cp_f32x4, score_src, score_vals[chunk, None]) + + for e in cutlass.range_constexpr(self.elems_per_lane): + chunk = const_expr(e // self.copy_elems) + copy_elem = const_expr(e % self.copy_elems) + score = score_vals[chunk, copy_elem] + kv = kv_vals[chunk, copy_elem] + new_max = cute.arch.fmax(local_max[e], score) + old_scale = cute.math.exp2( + (local_max[e] - new_max) * Float32(self.rcp_ln2), + fastmath=True, + ) + new_scale = cute.math.exp2( + (score - new_max) * Float32(self.rcp_ln2), + fastmath=True, + ) + local_sum[e] = local_sum[e] * old_scale + new_scale + local_product[e] = local_product[e] * old_scale + kv * new_scale + local_max[e] = new_max + + x = cute.make_rmem_tensor((self.elems_per_lane,), Float32) + local_sumsq = Float32(0.0) + for e in cutlass.range_constexpr(self.elems_per_lane): + x[e] = local_product[e] / local_sum[e] + local_sumsq += x[e] * x[e] + + warp_sum = local_sumsq + for step in cutlass.range_constexpr(5): + offset = const_expr(16 >> step) + warp_sum += cute.arch.shuffle_sync_bfly(warp_sum, offset) + + if lane_id == 0: + partial_sums[warp_id] = warp_sum + cute.arch.sync_threads() + if tid == 0: + total = Float32(0.0) + for i in cutlass.range_constexpr(self.num_warps): + total += partial_sums[i] + rrms_shared[0] = cute.math.rsqrt( + total / Float32(self.head_dim) + rms_norm_eps, fastmath=True + ) + cute.arch.sync_threads() + + rrms = rrms_shared[0] + for e in cutlass.range_constexpr(self.elems_per_lane): + elem = elem_base + e + x[e] = x[e] * rrms * rms_norm_weight[elem].to(Float32) + + page = kv_slot_idx // kv_cache_block_size + kv_offset = kv_slot_idx - page * kv_cache_block_size + value_base = page * k_cache.stride[0] + kv_offset * k_cache.stride[1] + + if const_expr(self.store_full_fp8): + k_cache_u16 = cute.recast_tensor(k_cache, Uint16) + inv_fp8 = Float32(1.0) / fp8_scale[0] + if group_idx == self.nope_blocks: + compressed_pos = (position // Int64(self.compress_ratio)) * Int64( + self.compress_ratio + ) + for pair in cutlass.range_constexpr(self.elems_per_lane // 2): + elem = const_expr(pair * 2) + pair_idx = (elem_base - self.nope_dim) // 2 + Int32(pair) + cos_v = cos_sin_cache[compressed_pos, pair_idx] + sin_v = cos_sin_cache[ + compressed_pos, pair_idx + Int32(self.rope_dim // 2) + ] + real = x[elem] * cos_v - x[elem + 1] * sin_v + imag = x[elem] * sin_v + x[elem + 1] * cos_v + packed_bf16 = _fp32x2_to_bf16x2(real, imag) + b0, b1 = _bf16x2_to_fp32(packed_bf16) + y0 = cutlass.min( + cutlass.max(b0 * inv_fp8, Float32(-self.fp8_max)), + Float32(self.fp8_max), + ) + y1 = cutlass.min( + cutlass.max(b1 * inv_fp8, Float32(-self.fp8_max)), + Float32(self.fp8_max), + ) + packed_fp8 = _fp32x2_to_fp8e4m3x2(y0, y1) + out_base = value_base + (elem_base + Int32(elem)).to(Int64) + k_cache_u16.iterator[out_base // Int64(2)] = packed_fp8 + else: + for pair in cutlass.range_constexpr(self.elems_per_lane // 2): + elem = const_expr(pair * 2) + packed_bf16 = _fp32x2_to_bf16x2(x[elem], x[elem + 1]) + b0, b1 = _bf16x2_to_fp32(packed_bf16) + y0 = cutlass.min( + cutlass.max(b0 * inv_fp8, Float32(-self.fp8_max)), + Float32(self.fp8_max), + ) + y1 = cutlass.min( + cutlass.max(b1 * inv_fp8, Float32(-self.fp8_max)), + Float32(self.fp8_max), + ) + packed_fp8 = _fp32x2_to_fp8e4m3x2(y0, y1) + out_base = value_base + (elem_base + Int32(elem)).to(Int64) + k_cache_u16.iterator[out_base // Int64(2)] = packed_fp8 + else: + k_cache_u32 = cute.recast_tensor(k_cache, Uint32) + if group_idx == self.nope_blocks: + compressed_pos = (position // Int64(self.compress_ratio)) * Int64( + self.compress_ratio + ) + for pair in cutlass.range_constexpr(self.elems_per_lane // 2): + elem = const_expr(pair * 2) + pair_idx = (elem_base - self.nope_dim) // 2 + Int32(pair) + cos_v = cos_sin_cache[compressed_pos, pair_idx] + sin_v = cos_sin_cache[ + compressed_pos, pair_idx + Int32(self.rope_dim // 2) + ] + real = x[elem] * cos_v - x[elem + 1] * sin_v + imag = x[elem] * sin_v + x[elem + 1] * cos_v + packed_bf16 = _fp32x2_to_bf16x2(real, imag) + out_base = value_base + ((elem_base + Int32(elem)) * 2).to( + Int64 + ) + k_cache_u32.iterator[out_base // Int64(4)] = packed_bf16 + else: + for pair in cutlass.range_constexpr(self.elems_per_lane // 2): + elem = const_expr(pair * 2) + packed_bf16 = _fp32x2_to_bf16x2(x[elem], x[elem + 1]) + out_base = value_base + ((elem_base + Int32(elem)) * 2).to( + Int64 + ) + k_cache_u32.iterator[out_base // Int64(4)] = packed_bf16 + + @cache + @staticmethod + def compile( + head_size: int = 512, + state_width: int = 1024, + rope_head_dim: int = 64, + fp8_max: float = 448.0, + quant_block: int = 64, + token_stride: int = 576, + scale_dim: int = 8, + kv_block_stride: int = 74752, + compress_ratio: int = 4, + overlap: bool = True, + store_full_fp8: bool = False, + norm_weight_dtype: type[cutlass.Numeric] = Float32, + ): + if compress_ratio != 4 or not overlap: + raise ValueError("CuTe DSL C4 fused sparse-attn requires C4 overlap.") + if head_size != 512: + raise ValueError( + "CuTe DSL C4 fused sparse-attn currently requires head_size=512." + ) + if state_width != 2 * head_size: + raise ValueError( + "CuTe DSL C4 fused sparse-attn requires state_width=2*head_size." + ) + if quant_block != 64: + raise ValueError( + "CuTe DSL C4 fused sparse-attn currently requires quant_block=64." + ) + if rope_head_dim != 64: + raise ValueError( + "CuTe DSL C4 fused sparse-attn currently requires rope_head_dim=64." + ) + num_positions = cute.sym_int() + num_slots = cute.sym_int() + num_req_indices = cute.sym_int() + num_kv_slots = cute.sym_int() + num_state_blocks = cute.sym_int() + num_kv_blocks = cute.sym_int() + state_cache_block_size = cute.sym_int() + block_table_width = cute.sym_int() + max_pos = cute.sym_int() + state_cache_width = state_width * 2 + + state_cache = cute.runtime.make_fake_tensor( + Float32, + (num_state_blocks, state_cache_block_size, state_cache_width), + stride=( + cute.sym_int64(divisibility=16), + cute.sym_int64(divisibility=16), + 1, + ), + assumed_align=16, + ) + token_to_req_indices = make_fake_tensor( + Int32, (num_req_indices,), divisibility=4 + ) + positions = make_fake_tensor(Int64, (num_positions,), divisibility=8) + slot_mapping = make_fake_tensor(Int64, (num_slots,), divisibility=8) + block_table = make_fake_tensor( + Int32, (cute.sym_int(), block_table_width), divisibility=1 + ) + rms_norm_weight = make_fake_tensor( + norm_weight_dtype, (head_size,), divisibility=4 + ) + cos_sin_cache = cute.runtime.make_fake_tensor( + Float32, + (max_pos, rope_head_dim), + stride=(cute.sym_int64(divisibility=4), 1), + assumed_align=4, + ) + k_cache = cute.runtime.make_fake_tensor( + Uint8, + (num_kv_blocks, cute.sym_int(), cute.sym_int()), + stride=( + cute.sym_int64(divisibility=16), + cute.sym_int64(divisibility=8), + 1, + ), + assumed_align=16, + ) + kv_slot_mapping = make_fake_tensor(Int64, (num_kv_slots,), divisibility=8) + fp8_scale = make_fake_tensor(Float32, (1,), divisibility=1) + + kernel = SparseAttnCompressNormRopeStoreFullC4Kernel( + head_size, + state_width, + rope_head_dim, + fp8_max, + quant_block, + token_stride, + scale_dim, + compress_ratio, + overlap, + store_full_fp8, + ) + stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True) + return cute.compile( + kernel, + state_cache, + token_to_req_indices, + positions, + slot_mapping, + block_table, + Int64(0), + rms_norm_weight, + Float32(0.0), + cos_sin_cache, + k_cache, + kv_slot_mapping, + Int64(0), + fp8_scale, + stream, + options="--enable-tvm-ffi", + ) + + class SparseAttnCompressC128Block8Kernel: head_tile = 64 rows_per_warp = 16 @@ -1109,6 +1529,273 @@ class SparseAttnNormRopeStoreKernel: ) +class SparseAttnNormRopeStoreFullKernel: + def __init__( + self, + head_size: int, + rope_head_dim: int, + fp8_max: float, + quant_block: int, + token_stride: int, + scale_dim: int, + compress_ratio: int, + store_full_fp8: bool = False, + ): + # Standalone (not inheriting the #44230-restructured legacy kernel): + # set attrs directly so the full-cache C128 path is decoupled. + self.head_dim = head_size + self.rope_dim = rope_head_dim + self.nope_dim = head_size - rope_head_dim + self.fp8_max = fp8_max + self.quant_block = quant_block + self.token_stride = token_stride + self.scale_dim = scale_dim + self.num_warps = head_size // quant_block + self.nope_blocks = self.nope_dim // quant_block + self.tb_size = head_size // 2 + self.compress_ratio = compress_ratio + self.store_full_fp8 = store_full_fp8 + + @cute.jit + def __call__( + self, + compressed_kv: cute.Tensor, + positions: cute.Tensor, + slot_mapping: cute.Tensor, + rms_norm_weight: cute.Tensor, + rms_norm_eps: Float32, + cos_sin_cache: cute.Tensor, + k_cache: cute.Tensor, + kv_slot_mapping: cute.Tensor, + kv_cache_block_size: Int64, + fp8_scale: cute.Tensor, + stream: CUstream, + ): + grid = (slot_mapping.shape[0], 1, 1) + self.kernel( + compressed_kv, + positions, + slot_mapping, + rms_norm_weight, + rms_norm_eps, + cos_sin_cache, + k_cache, + kv_slot_mapping, + kv_cache_block_size, + fp8_scale, + ).launch(grid=grid, block=(self.tb_size, 1, 1), stream=stream) + + @cute.kernel + def kernel( + self, + compressed_kv: cute.Tensor, + positions: cute.Tensor, + slot_mapping: cute.Tensor, + rms_norm_weight: cute.Tensor, + rms_norm_eps: Float32, + cos_sin_cache: cute.Tensor, + k_cache: cute.Tensor, + kv_slot_mapping: cute.Tensor, + kv_cache_block_size: Int64, + fp8_scale: cute.Tensor, + ): + token_idx, _, _ = cute.arch.block_idx() + tid, _, _ = cute.arch.thread_idx() + warp_id = cute.arch.make_warp_uniform(tid // 32) + lane_id = tid % 32 + elem0 = tid * 2 + + slot_id = slot_mapping[token_idx] + has_position = token_idx < positions.shape[0] + position = Int64(0) + if has_position: + position = positions[token_idx] + boundary = has_position and ( + (position + Int64(1)) % Int64(self.compress_ratio) == Int64(0) + ) + has_kv_slot_idx = token_idx < kv_slot_mapping.shape[0] + kv_slot_idx = Int64(-1) + if has_kv_slot_idx: + kv_slot_idx = kv_slot_mapping[token_idx] + active = slot_id >= Int64(0) and boundary and kv_slot_idx >= Int64(0) + + if active: + base = token_idx.to(Int64) * compressed_kv.stride[0] + elem0.to(Int64) + x0 = compressed_kv.iterator[base] + x1 = compressed_kv.iterator[base + Int64(1)] + + local_sumsq = x0 * x0 + x1 * x1 + warp_sum = local_sumsq + for step in cutlass.range_constexpr(5): + offset = const_expr(16 >> step) + warp_sum += cute.arch.shuffle_sync_bfly(warp_sum, offset) + + smem = cutlass.utils.SmemAllocator() + partial_sums = smem.allocate_tensor( + Float32, cute.make_layout((self.num_warps,)), byte_alignment=4 + ) + rrms_shared = smem.allocate_tensor( + Float32, cute.make_layout((1,)), byte_alignment=4 + ) + + if lane_id == 0: + partial_sums[warp_id] = warp_sum + cute.arch.sync_threads() + if tid == 0: + total = Float32(0.0) + for i in cutlass.range_constexpr(self.num_warps): + total += partial_sums[i] + rrms_shared[0] = cute.math.rsqrt( + total / Float32(self.head_dim) + rms_norm_eps, fastmath=True + ) + cute.arch.sync_threads() + + rrms = rrms_shared[0] + x0 = x0 * rrms * rms_norm_weight[elem0].to(Float32) + x1 = x1 * rrms * rms_norm_weight[elem0 + 1].to(Float32) + + page = kv_slot_idx // kv_cache_block_size + kv_offset = kv_slot_idx - page * kv_cache_block_size + value_base = page * k_cache.stride[0] + kv_offset * k_cache.stride[1] + + if const_expr(self.store_full_fp8): + k_cache_u16 = cute.recast_tensor(k_cache, Uint16) + inv_fp8 = Float32(1.0) / fp8_scale[0] + fp8_v0 = x0 + fp8_v1 = x1 + if warp_id == self.nope_blocks: + compressed_pos = (position // Int64(self.compress_ratio)) * Int64( + self.compress_ratio + ) + pair_idx = lane_id + cs_base = compressed_pos * cos_sin_cache.stride[0] + pair_idx.to( + Int64 + ) + cos_v = cos_sin_cache.iterator[cs_base] + sin_v = cos_sin_cache.iterator[cs_base + Int64(self.rope_dim // 2)] + fp8_v0 = x0 * cos_v - x1 * sin_v + fp8_v1 = x0 * sin_v + x1 * cos_v + fp8_packed_bf16 = _fp32x2_to_bf16x2(fp8_v0, fp8_v1) + b0, b1 = _bf16x2_to_fp32(fp8_packed_bf16) + y0 = cutlass.min( + cutlass.max(b0 * inv_fp8, Float32(-self.fp8_max)), + Float32(self.fp8_max), + ) + y1 = cutlass.min( + cutlass.max(b1 * inv_fp8, Float32(-self.fp8_max)), + Float32(self.fp8_max), + ) + packed_fp8 = _fp32x2_to_fp8e4m3x2(y0, y1) + out_base = value_base + elem0.to(Int64) + k_cache_u16.iterator[out_base // Int64(2)] = packed_fp8 + else: + k_cache_u32 = cute.recast_tensor(k_cache, Uint32) + bf16_v0 = x0 + bf16_v1 = x1 + if warp_id == self.nope_blocks: + compressed_pos = (position // Int64(self.compress_ratio)) * Int64( + self.compress_ratio + ) + pair_idx = lane_id + cs_base = compressed_pos * cos_sin_cache.stride[0] + pair_idx.to( + Int64 + ) + cos_v = cos_sin_cache.iterator[cs_base] + sin_v = cos_sin_cache.iterator[cs_base + Int64(self.rope_dim // 2)] + bf16_v0 = x0 * cos_v - x1 * sin_v + bf16_v1 = x0 * sin_v + x1 * cos_v + bf16_packed = _fp32x2_to_bf16x2(bf16_v0, bf16_v1) + out_base = value_base + (elem0 * 2).to(Int64) + k_cache_u32.iterator[out_base // Int64(4)] = bf16_packed + + @cache + @staticmethod + def compile( + head_size: int = 512, + rope_head_dim: int = 64, + fp8_max: float = 448.0, + quant_block: int = 64, + token_stride: int = 576, + scale_dim: int = 8, + kv_block_stride: int = 74752, + compress_ratio: int = 128, + store_full_fp8: bool = False, + norm_weight_dtype: type[cutlass.Numeric] = Float32, + ): + if quant_block != 64: + raise ValueError( + "CuTe DSL sparse-attn store currently requires quant_block=64." + ) + if rope_head_dim != 64: + raise ValueError( + "CuTe DSL sparse-attn store currently requires rope_head_dim=64." + ) + if head_size % quant_block != 0: + raise ValueError("head_size must be divisible by quant_block.") + num_positions = cute.sym_int() + num_slots = cute.sym_int() + num_kv_slots = cute.sym_int() + max_pos = cute.sym_int() + num_blocks = cute.sym_int() + + compressed_kv = cute.runtime.make_fake_tensor( + Float32, + (num_slots, head_size), + stride=(cute.sym_int64(divisibility=4), 1), + assumed_align=4, + ) + positions = make_fake_tensor(Int64, (num_positions,), divisibility=8) + slot_mapping = make_fake_tensor(Int64, (num_slots,), divisibility=8) + rms_norm_weight = make_fake_tensor( + norm_weight_dtype, (head_size,), divisibility=4 + ) + cos_sin_cache = cute.runtime.make_fake_tensor( + Float32, + (max_pos, rope_head_dim), + stride=(cute.sym_int64(divisibility=4), 1), + assumed_align=4, + ) + k_cache = cute.runtime.make_fake_tensor( + Uint8, + (num_blocks, cute.sym_int(), cute.sym_int()), + stride=( + cute.sym_int64(divisibility=16), + cute.sym_int64(divisibility=8), + 1, + ), + assumed_align=16, + ) + kv_slot_mapping = make_fake_tensor(Int64, (num_kv_slots,), divisibility=8) + fp8_scale = make_fake_tensor(Float32, (1,), divisibility=1) + + kernel = SparseAttnNormRopeStoreFullKernel( + head_size, + rope_head_dim, + fp8_max, + quant_block, + token_stride, + scale_dim, + compress_ratio, + store_full_fp8, + ) + stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True) + return cute.compile( + kernel, + compressed_kv, + positions, + slot_mapping, + rms_norm_weight, + Float32(0.0), + cos_sin_cache, + k_cache, + kv_slot_mapping, + Int64(0), + fp8_scale, + stream, + options="--enable-tvm-ffi", + ) + + def compile_split_sparse_attn_cutedsl( head_size: int, state_width: int, @@ -1123,6 +1810,8 @@ def compile_split_sparse_attn_cutedsl( compress_ratio: int, overlap: bool, rms_norm_weight_dtype: torch.dtype, + store_full_kv: bool = False, + store_full_fp8: bool = False, ): if not ( head_size == 512 @@ -1141,18 +1830,33 @@ def compile_split_sparse_attn_cutedsl( state_width=state_width, ) norm_weight_dtype = _TORCH_TO_CUTE[rms_norm_weight_dtype] - store = SparseAttnNormRopeStoreKernel.compile( - head_size, - rope_head_dim, - fp8_max, - quant_block, - token_stride, - scale_dim, - kv_block_stride, - compress_ratio, - norm_weight_dtype, - kv_cache_block_size, - ) + if store_full_kv: + # FlashInfer contiguous bf16/fp8 cache: standalone full-cache store. + store = SparseAttnNormRopeStoreFullKernel.compile( + head_size=head_size, + rope_head_dim=rope_head_dim, + fp8_max=fp8_max, + quant_block=quant_block, + token_stride=token_stride, + scale_dim=scale_dim, + kv_block_stride=kv_block_stride, + compress_ratio=compress_ratio, + store_full_fp8=store_full_fp8, + norm_weight_dtype=norm_weight_dtype, + ) + else: + store = SparseAttnNormRopeStoreKernel.compile( + head_size, + rope_head_dim, + fp8_max, + quant_block, + token_stride, + scale_dim, + kv_block_stride, + compress_ratio, + norm_weight_dtype, + kv_cache_block_size, + ) return compress, store @@ -1180,13 +1884,16 @@ def split_kv_compress_norm_rope_insert_sparse_attn_cutedsl( scale_dim: int = 8, compress_ratio: int = 128, overlap: bool = False, + store_full_kv: bool = False, + store_full_fp8: bool = False, + fp8_scale: torch.Tensor | None = None, ) -> None: if k_cache.ndim != 3: raise ValueError( "CuTe DSL sparse-attn store expects the real DeepSeek V4 " f"3D k_cache layout [num_blocks, block_size, 584], got ndim={k_cache.ndim}." ) - if kv_cache_block_size != k_cache.shape[1]: + if not store_full_kv and kv_cache_block_size != k_cache.shape[1]: raise ValueError( "CuTe DSL split sparse-attn wrapper expected kv_cache_block_size " f"to match k_cache.shape[1], got {kv_cache_block_size} and " @@ -1199,6 +1906,8 @@ def split_kv_compress_norm_rope_insert_sparse_attn_cutedsl( "CuTe DSL sparse-attn store supports rms_norm_weight dtype " f"bf16/fp32, got {rms_norm_weight.dtype}." ) + if store_full_fp8 and not store_full_kv: + raise ValueError("store_full_fp8 requires store_full_kv.") compress, store = compile_split_sparse_attn_cutedsl( head_size, state_width, @@ -1213,6 +1922,8 @@ def split_kv_compress_norm_rope_insert_sparse_attn_cutedsl( compress_ratio, overlap, rms_norm_weight.dtype, + store_full_kv=store_full_kv, + store_full_fp8=store_full_fp8, ) compress( state_cache, @@ -1223,6 +1934,25 @@ def split_kv_compress_norm_rope_insert_sparse_attn_cutedsl( compressed_kv, ) + if store_full_kv: + # Byte-addressed contiguous cache; block size + per-tensor scale are + # passed at call time (not baked into compile). + if fp8_scale is None: + fp8_scale = torch.ones(1, dtype=torch.float32, device=k_cache.device) + store( + compressed_kv, + positions, + slot_mapping, + rms_norm_weight, + rms_norm_eps, + cos_sin_cache, + k_cache.view(torch.uint8), + kv_slot_mapping, + kv_cache_block_size, + fp8_scale, + ) + return + store( compressed_kv, positions, @@ -1258,6 +1988,9 @@ def fused_kv_compress_norm_rope_insert_sparse_attn_cutedsl( scale_dim: int = 8, compress_ratio: int = 4, overlap: bool = True, + store_full_kv: bool = False, + store_full_fp8: bool = False, + fp8_scale: torch.Tensor | None = None, ) -> None: if positions.numel() == 0: return @@ -1272,6 +2005,43 @@ def fused_kv_compress_norm_rope_insert_sparse_attn_cutedsl( "CuTe DSL sparse-attn fused store expects the real DeepSeek V4 " f"3D k_cache layout [num_blocks, block_size, 584], got ndim={k_cache.ndim}." ) + if store_full_fp8 and not store_full_kv: + raise ValueError("store_full_fp8 requires store_full_kv.") + if store_full_kv: + # FlashInfer contiguous bf16/fp8 cache: byte-addressed full-cache C4 store. + if fp8_scale is None: + fp8_scale = torch.ones(1, dtype=torch.float32, device=k_cache.device) + compiled = SparseAttnCompressNormRopeStoreFullC4Kernel.compile( + head_size=head_size, + state_width=state_width, + rope_head_dim=rope_head_dim, + fp8_max=fp8_max, + quant_block=quant_block, + token_stride=token_stride, + scale_dim=scale_dim, + kv_block_stride=kv_block_stride, + compress_ratio=compress_ratio, + overlap=overlap, + store_full_fp8=store_full_fp8, + norm_weight_dtype=norm_weight_dtype, + ) + compiled( + state_cache, + token_to_req_indices, + positions, + slot_mapping, + block_table, + block_size, + rms_norm_weight, + rms_norm_eps, + cos_sin_cache, + k_cache.view(torch.uint8), + kv_slot_mapping, + kv_cache_block_size, + fp8_scale, + ) + return + compiled = SparseAttnCompressNormRopeStoreC4Kernel.compile( head_size=head_size, state_width=state_width, @@ -1324,6 +2094,9 @@ def compress_norm_rope_store_cutedsl( quant_block: int, token_stride: int, scale_dim: int, + store_full_kv: bool = False, + store_full_fp8: bool = False, + fp8_scale: torch.Tensor | None = None, ) -> None: if compress_ratio == 4: # For C4A, the single fused kernel is faster than the two-kernel version. @@ -1350,6 +2123,9 @@ def compress_norm_rope_store_cutedsl( scale_dim=scale_dim, compress_ratio=compress_ratio, overlap=overlap, + store_full_kv=store_full_kv, + store_full_fp8=store_full_fp8, + fp8_scale=fp8_scale, ) else: # For C128, the two-kernel version is faster than the single fused kernel. @@ -1382,4 +2158,7 @@ def compress_norm_rope_store_cutedsl( scale_dim=scale_dim, compress_ratio=compress_ratio, overlap=overlap, + store_full_kv=store_full_kv, + store_full_fp8=store_full_fp8, + fp8_scale=fp8_scale, ) diff --git a/vllm/utils/flashinfer.py b/vllm/utils/flashinfer.py index f7ed180a730..95f8b4b7ec0 100644 --- a/vllm/utils/flashinfer.py +++ b/vllm/utils/flashinfer.py @@ -72,6 +72,14 @@ def _missing(*_: Any, **__: Any) -> NoReturn: ) +def _missing_dsv4_sparse_mla(*_: Any, **__: Any) -> NoReturn: + raise RuntimeError( + "flashinfer.mla.trtllm_batch_decode_sparse_mla_dsv4 is not available. " + "Install a FlashInfer build that includes DeepSeek V4 sparse MLA " + "TRTLLM-GEN support." + ) + + def _get_submodule(module_name: str) -> Any | None: """Safely import a submodule and return it, or None if not available.""" try: @@ -141,6 +149,14 @@ flashinfer_b12x_fused_moe = _lazy_import_wrapper( trtllm_fp4_block_scale_moe = _lazy_import_wrapper( "flashinfer", "trtllm_fp4_block_scale_moe" ) +# DeepSeek V4 sparse MLA TRTLLM-GEN decode launcher (public wrapper). Handles +# the SWA + compressed KV pools, the concatenated sparse-index matrix, and +# per-tensor FP8 / BF16 inputs with BF16 output. +flashinfer_trtllm_batch_decode_sparse_mla_dsv4 = _lazy_import_wrapper( + "flashinfer.mla", + "trtllm_batch_decode_sparse_mla_dsv4", + fallback_fn=_missing_dsv4_sparse_mla, +) # Special case for autotune since it returns a context manager autotune = _lazy_import_wrapper( "flashinfer.autotuner", @@ -965,6 +981,7 @@ __all__ = [ "flashinfer_b12x_fused_moe", "flashinfer_convert_sf_to_mma_layout", "trtllm_fp4_block_scale_moe", + "flashinfer_trtllm_batch_decode_sparse_mla_dsv4", "autotune", "has_flashinfer_moe", "has_flashinfer_comm", diff --git a/vllm/v1/attention/backends/mla/sparse_swa.py b/vllm/v1/attention/backends/mla/sparse_swa.py index f0e444e493c..59698442f98 100644 --- a/vllm/v1/attention/backends/mla/sparse_swa.py +++ b/vllm/v1/attention/backends/mla/sparse_swa.py @@ -73,9 +73,14 @@ class DeepseekV4SWACache(torch.nn.Module, AttentionLayerBase): # determines the SWA block size of 64 tokens per block. # TODO(yifan): make SWA block size automatically determined and configurable. self.block_size = 64 - assert self.dtype == torch.uint8 + # uint8: legacy FlashMLA UE8M0 paged layout. bfloat16 / float8_e4m3fn: + # FlashInfer contiguous full-cache layout. + assert self.dtype in (torch.uint8, torch.bfloat16, torch.float8_e4m3fn) def get_kv_cache_spec(self, vllm_config: VllmConfig) -> KVCacheSpec: + # FlashMLA's UE8M0 paged layout needs 576B alignment; FlashInfer's + # contiguous bf16/fp8 cache uses the natural element-size page. + is_flashmla = self.cache_config.cache_dtype == "fp8_ds_mla" return SlidingWindowMLASpec( block_size=self.block_size, num_kv_heads=1, @@ -83,7 +88,7 @@ class DeepseekV4SWACache(torch.nn.Module, AttentionLayerBase): dtype=self.dtype, sliding_window=self.window_size, cache_dtype_str=self.cache_config.cache_dtype, - alignment=576, # NOTE: FlashMLA requires 576B alignment + alignment=576 if is_flashmla else None, model_version="deepseek_v4", ) diff --git a/vllm/v1/attention/backends/registry.py b/vllm/v1/attention/backends/registry.py index 87abb688431..0f91dac5aeb 100644 --- a/vllm/v1/attention/backends/registry.py +++ b/vllm/v1/attention/backends/registry.py @@ -76,6 +76,17 @@ class AttentionBackendEnum(Enum, metaclass=_AttentionBackendEnumMeta): FLASHMLA_SPARSE = ( "vllm.v1.attention.backends.mla.flashmla_sparse.FlashMLASparseBackend" ) + # DeepSeek V4 sparse MLA backends (model-driven; selected via the V4 layer). + FLASHMLA_SPARSE_DSV4 = ( + "vllm.models.deepseek_v4.nvidia.flashmla.DeepseekV4FlashMLASparseBackend" + ) + FLASHINFER_MLA_SPARSE_DSV4 = ( + "vllm.models.deepseek_v4.nvidia.flashinfer_sparse." + "DeepseekV4FlashInferMLASparseBackend" + ) + ROCM_FLASHMLA_SPARSE_DSV4 = ( + "vllm.models.deepseek_v4.amd.rocm.DeepseekV4ROCMAiterMLASparseBackend" + ) FLASH_ATTN_MLA = "vllm.v1.attention.backends.mla.flashattn_mla.FlashAttnMLABackend" NO_ATTENTION = "vllm.v1.attention.backends.no_attention.NoAttentionBackend" FLEX_ATTENTION = "vllm.v1.attention.backends.flex_attention.FlexAttentionBackend" diff --git a/vllm/v1/kv_cache_interface.py b/vllm/v1/kv_cache_interface.py index 3bbfba1a0fe..2f8048c7966 100644 --- a/vllm/v1/kv_cache_interface.py +++ b/vllm/v1/kv_cache_interface.py @@ -547,10 +547,12 @@ class SlidingWindowMLASpec(SlidingWindowSpec): @property def real_page_size_bytes(self) -> int: - if self.model_version == "deepseek_v4": - # DeepseekV4: 448B NoPE + 128B RoPE + 8B fp8 scale = 584B per token. + if self.model_version == "deepseek_v4" and self.cache_dtype_str == "fp8_ds_mla": + # DeepseekV4 FlashMLA: 448B NoPE + 128B RoPE + 8B fp8 scale = 584B + # per token. FlashInfer's contiguous bf16/fp8 cache falls through to + # the element-size formula below. return self.storage_block_size * 584 - assert self.model_version is None, ( + assert self.model_version in (None, "deepseek_v4"), ( f"Unsupported model version: {self.model_version}" ) return (