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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -942,9 +942,10 @@ static void launchFullCacheKernel(
// ────────────────────────────────────────────────────────────────────────────
// Torch op wrapper
// ────────────────────────────────────────────────────────────────────────────
torch::stable::Tensor fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert(
void fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert_out(
torch::stable::Tensor const& q_in, // [N, num_heads_q, 512] bf16
torch::stable::Tensor const& kv, // [N, 512] bf16 (read-only)
torch::stable::Tensor& q_out, // [N, q_head_padded, 512]
torch::stable::Tensor& k_cache, // [num_blocks, block_bytes] uint8
torch::stable::Tensor const& slot_mapping, // [N] int64
torch::stable::Tensor const& position_ids, // [N] int64
Expand All @@ -970,8 +971,16 @@ torch::stable::Tensor fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert(
STD_TORCH_CHECK(kv.dim() == 2 && kv.size(1) == 512, "kv shape [N, 512]");
STD_TORCH_CHECK(q_in.scalar_type() == kv.scalar_type(),
"q_in and kv dtype must match");
STD_TORCH_CHECK(q_out.device() == q_in.device() && q_out.is_contiguous(),
"q_out must be contiguous and on the same device as q_in");
STD_TORCH_CHECK(q_out.scalar_type() == q_in.scalar_type(),
"q_out dtype must match q_in");
STD_TORCH_CHECK(q_head_padded >= q_in.size(1),
"q_head_padded must be >= q_in.size(1) (num_heads_q)");
STD_TORCH_CHECK(q_out.dim() == 3 && q_out.size(0) == q_in.size(0) &&
q_out.size(1) == q_head_padded &&
q_out.size(2) == q_in.size(2),
"q_out shape [N, q_head_padded, 512]");
STD_TORCH_CHECK(k_cache.scalar_type() == torch::headeronly::ScalarType::Byte,
"k_cache must be uint8");
STD_TORCH_CHECK(cos_sin_cache.dim() == 2 && cos_sin_cache.size(1) == 64,
Expand Down Expand Up @@ -999,11 +1008,6 @@ torch::stable::Tensor fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert(
q_in.get_device_index());
const cudaStream_t stream = get_current_cuda_stream(q_in.get_device_index());

// Allocate the padded q output. The kernel writes every element (live
// region gets RMSNorm+RoPE; pad region gets zeros), so `empty` is safe.
auto q_out = torch::stable::new_empty(
q_in, {q_in.size(0), q_head_padded, q_in.size(2)}, q_in.scalar_type());

VLLM_STABLE_DISPATCH_HALF_TYPES(
q_in.scalar_type(), "fused_deepseek_v4_qnorm_rope_kv_insert", [&] {
using qkv_scalar_t = scalar_t;
Expand All @@ -1020,6 +1024,20 @@ torch::stable::Tensor fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert(
num_heads_q_padded, cache_block_size_i, kv_block_stride,
stream);
});
}

torch::stable::Tensor fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert(
torch::stable::Tensor const& q_in, 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, int64_t q_head_padded,
double eps, int64_t cache_block_size) {
auto q_out = torch::stable::new_empty(
q_in, {q_in.size(0), q_head_padded, q_in.size(2)}, q_in.scalar_type());
fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert_out(
q_in, kv, q_out, k_cache, slot_mapping, position_ids, cos_sin_cache,
q_head_padded, eps, cache_block_size);
return q_out;
}

Expand Down
8 changes: 8 additions & 0 deletions csrc/libtorch_stable/ops.h
Original file line number Diff line number Diff line change
Expand Up @@ -269,6 +269,14 @@ 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_quant_insert_out(
torch::stable::Tensor const& q_in, torch::stable::Tensor const& kv,
torch::stable::Tensor& q_out, torch::stable::Tensor& k_cache,
torch::stable::Tensor const& slot_mapping,
torch::stable::Tensor const& position_ids,
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,
Expand Down
7 changes: 7 additions & 0 deletions csrc/libtorch_stable/torch_bindings.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -433,6 +433,11 @@ STABLE_TORCH_LIBRARY_FRAGMENT(_C, ops) {
"Tensor q_in, Tensor kv, Tensor! k_cache, "
"Tensor slot_mapping, Tensor position_ids, Tensor cos_sin_cache, "
"int q_head_padded, float eps, int cache_block_size) -> Tensor");
ops.def(
"fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert_out("
"Tensor q_in, Tensor kv, Tensor! q_out, Tensor! k_cache, "
"Tensor slot_mapping, Tensor position_ids, Tensor cos_sin_cache, "
"int q_head_padded, float eps, int cache_block_size) -> ()");

// 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.
Expand Down Expand Up @@ -751,6 +756,8 @@ 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_quant_insert_out",
TORCH_BOX(&fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert_out));
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));
Expand Down
18 changes: 18 additions & 0 deletions tests/kernels/test_compressor_kv_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@

from vllm import _custom_ops as ops
from vllm.models.deepseek_v4.common.ops import (
compute_global_topk_indices_and_lens,
dequantize_and_gather_k_cache,
quantize_and_insert_k_cache,
)
Expand All @@ -34,6 +35,23 @@
from .test_fused_indexer_q_rope_quant import quantize_to_mxfp4


def test_compute_global_topk_reuses_output_buffers():
device = "cuda"
topk_indices = torch.tensor(
[[0, 3, -1], [1, 2, -1]], dtype=torch.int32, device=device
)
token_to_req = torch.tensor([0, 1], dtype=torch.int32, device=device)
block_table = torch.tensor([[5, 7], [11, 13]], dtype=torch.int32, device=device)
is_valid = torch.tensor([True, False], device=device)
args = (topk_indices, token_to_req, block_table, 2, is_valid)
expected = compute_global_topk_indices_and_lens(*args)
outputs = tuple(torch.empty_like(tensor) for tensor in expected)
actual = compute_global_topk_indices_and_lens(*args, output_buffers=outputs)
for result, output, reference in zip(actual, outputs, expected):
assert result.data_ptr() == output.data_ptr()
torch.testing.assert_close(result, reference)


def _ue8m0_reference(x: torch.Tensor, block_size: int, fp8_max: float):
"""PyTorch reference for UE8M0 FP8 quantization (per-block, power-of-2 scale).

Expand Down
14 changes: 12 additions & 2 deletions tests/kernels/test_fused_deepseek_v4_qnorm_rope_kv_insert.py
Original file line number Diff line number Diff line change
Expand Up @@ -257,8 +257,18 @@ def test_q_path_matches_reference(num_tokens: int, n_heads: int, padded_heads: i
num_blocks, bs, HEAD_BYTES, dtype=torch.uint8, device=device
).view(num_blocks, -1)
slot_mapping = torch.full((num_tokens,), -1, dtype=torch.int64, device=device)
q_out = _call_fused(
q, padded_heads, kv, k_cache, slot_mapping, positions, cos_sin_cache, eps, bs
q_out = torch.empty(num_tokens, padded_heads, HEAD_DIM, dtype=dtype, device=device)
torch.ops._C.fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert_out(
q,
kv,
q_out,
k_cache,
slot_mapping,
positions,
cos_sin_cache,
padded_heads,
eps,
bs,
)

torch.testing.assert_close(q_out[:, :n_heads], q_ref, rtol=1e-2, atol=1e-2)
Expand Down
26 changes: 26 additions & 0 deletions tests/kernels/test_fused_indexer_q_rope_quant.py
Original file line number Diff line number Diff line change
Expand Up @@ -150,6 +150,23 @@ def test_fused_indexer_q_rope_quant_matches_unfused(
q_quant_ref, weights_ref = _reference(
positions, q, cos_sin_cache, weights, softmax_scale, head_scale, use_fp4
)
output_buffers: tuple[torch.Tensor, ...] | None = None
OUTPUT_BUFFER_TEST_NUM_TOKENS = 7
if num_tokens == OUTPUT_BUFFER_TEST_NUM_TOKENS and cache_dtype == torch.float32:
if use_fp4:
q_ref, q_scale_ref = q_quant_ref
output_buffers = (
torch.empty_like(q_ref),
torch.empty_like(q_scale_ref)
.view(torch.uint8)
.reshape(num_tokens, N_HEAD, -1),
torch.empty_like(weights_ref),
)
else:
output_buffers = (
torch.empty_like(q_quant_ref),
torch.empty_like(weights_ref),
)
# use_cutedsl=False: force the triton path even when cutedsl is installed
# by patching the dispatcher's has_cutedsl() binding to return False.
cutedsl_patch = (
Expand All @@ -169,8 +186,17 @@ def test_fused_indexer_q_rope_quant_matches_unfused(
softmax_scale,
head_scale,
use_fp4,
output_buffers=output_buffers,
)

if output_buffers is not None:
if use_fp4:
assert q_quant_fused[0].data_ptr() == output_buffers[0].data_ptr()
assert q_quant_fused[1].data_ptr() == output_buffers[1].data_ptr()
else:
assert q_quant_fused.data_ptr() == output_buffers[0].data_ptr()
assert weights_fused.data_ptr() == output_buffers[-1].data_ptr()

if use_fp4:
q_quant_ref, q_scale_ref = q_quant_ref
q_quant_fused, q_scale_fused = q_quant_fused
Expand Down
37 changes: 35 additions & 2 deletions vllm/models/deepseek_v4/attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
from vllm.models.deepseek_v4.common.ops.fused_indexer_q import MXFP4_BLOCK_SIZE

if TYPE_CHECKING:
from vllm.models.deepseek_v4.eager_scratch import DeepseekV4EagerScratchPool
from vllm.v1.attention.backends.mla.sparse_swa import (
DeepseekSparseSWAMetadata,
)
Expand Down Expand Up @@ -181,6 +182,7 @@ def __init__(
prefix: str,
topk_indices_buffer: torch.Tensor | None = None,
aux_stream_list: list[torch.cuda.Stream] | None = None,
eager_scratch_pool: "DeepseekV4EagerScratchPool | None" = None,
) -> None:
super().__init__()
config = vllm_config.model_config.hf_config
Expand Down Expand Up @@ -269,6 +271,7 @@ def __init__(
)
self.indexer_rotary_emb = self.rotary_emb
self.topk_indices_buffer = topk_indices_buffer
self.eager_scratch_pool = eager_scratch_pool

self.indexer = None
if self.compress_ratio == 4:
Expand All @@ -290,6 +293,7 @@ def __init__(
compress_ratio=self.compress_ratio,
prefix=f"{prefix}.indexer",
aux_stream=indexer_aux_stream,
eager_scratch_pool=eager_scratch_pool,
)

# Will be None on ROCm for now.
Expand Down Expand Up @@ -340,6 +344,7 @@ def __init__(
rotate=True,
prefix=f"{prefix}.compressor",
k_cache_prefix=self.prefix,
eager_scratch_pool=eager_scratch_pool,
)

def forward(
Expand Down Expand Up @@ -568,10 +573,24 @@ def _fused_qnorm_rope_kv_insert(
if cache_dtype == torch.uint8:
# fp8_ds_mla 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.
# the padding head slots.
# 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)
if self.eager_scratch_pool is not None:
q_out = self.eager_scratch_pool.q_out(q.shape[0])
torch.ops._C.fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert_out(
q,
kv,
q_out,
swa_kv_cache_2d,
swa_metadata.slot_mapping,
positions,
cos_sin_cache,
self.padded_heads,
self.eps,
swa_metadata.block_size,
)
return q_out
return torch.ops._C.fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert(
q,
kv,
Expand Down Expand Up @@ -620,6 +639,13 @@ def _fused_qnorm_rope_kv_insert(
)
return q_fp8

def _global_topk_output_buffers(
self, topk_indices: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor] | None:
if self.compress_ratio != 4 or self.eager_scratch_pool is None:
return None
return self.eager_scratch_pool.global_topk_outputs(topk_indices)

def get_attn_backend(self) -> type[AttentionBackend]:
return self.backend_cls

Expand Down Expand Up @@ -699,6 +725,7 @@ def __init__(
compress_ratio: int = 1,
prefix: str = "",
aux_stream: torch.cuda.Stream | None = None,
eager_scratch_pool: "DeepseekV4EagerScratchPool | None" = None,
):
super().__init__()
self.vllm_config = vllm_config
Expand All @@ -711,6 +738,7 @@ def __init__(
self.rope_dim = config.qk_rope_head_dim # 64
self.q_lora_rank = q_lora_rank # 1536
self.compress_ratio = compress_ratio
self.eager_scratch_pool = eager_scratch_pool
self.use_fp4_kv = self.vllm_config.attention_config.use_fp4_indexer_cache
logger.info_once(
"Using %s indexer cache for Lightning Indexer.",
Expand Down Expand Up @@ -774,6 +802,7 @@ def __init__(
prefix=f"{prefix}.compressor",
k_cache_prefix=self.k_cache.prefix,
use_fp4_cache=self.use_fp4_kv,
eager_scratch_pool=eager_scratch_pool,
)

self.indexer_op = SparseAttnIndexer(
Expand Down Expand Up @@ -834,6 +863,9 @@ def wq_b_and_q_quant():
# ReplicatedLinear returns (output, bias); bias is None.
q, _ = self.wq_b(qr)
q = q.view(-1, self.n_head, self.head_dim)
outputs = None
if self.eager_scratch_pool is not None and self.use_fp4_kv:
outputs = self.eager_scratch_pool.indexer_q_outputs(q.shape[0])
return fused_indexer_q_rope_quant(
positions,
q,
Expand All @@ -842,6 +874,7 @@ def wq_b_and_q_quant():
self.softmax_scale,
self.n_head**-0.5,
use_fp4=self.use_fp4_kv,
output_buffers=outputs,
)

# compressor returns None and writes K to the indexer KV cache; the
Expand Down
12 changes: 10 additions & 2 deletions vllm/models/deepseek_v4/common/ops/cache_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -438,6 +438,7 @@ def compute_global_topk_indices_and_lens(
block_table: torch.Tensor,
block_size: int,
is_valid_token: torch.Tensor,
output_buffers: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Map local topk indices to global KV cache slots and count valid entries.

Expand All @@ -447,8 +448,15 @@ def compute_global_topk_indices_and_lens(
3. Masking padding tokens to length 0
"""
num_tokens = topk_indices.shape[0]
global_topk_indices = torch.empty_like(topk_indices)
topk_lens = torch.empty(num_tokens, dtype=torch.int32, device=topk_indices.device)
if output_buffers is None:
global_topk_indices = torch.empty_like(topk_indices)
topk_lens = torch.empty(
num_tokens, dtype=torch.int32, device=topk_indices.device
)
else:
global_topk_indices, topk_lens = output_buffers
assert global_topk_indices.shape == topk_indices.shape
assert topk_lens.shape == (num_tokens,)
_compute_global_topk_indices_and_lens_kernel[(num_tokens,)](
global_topk_indices,
global_topk_indices.stride(0),
Expand Down
Loading
Loading