From 9e428b83c7112d29f08d1f5479906f8350e2b6f0 Mon Sep 17 00:00:00 2001 From: biondizzle Date: Tue, 19 May 2026 15:34:09 +0000 Subject: [PATCH] =?UTF-8?q?Fix=20KV=20cache:=20write=20to=20paged=20cache,?= =?UTF-8?q?=20handle=20uint8=E2=86=92fp8=20conversion,=20fix=20RoPE=20bug?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- vllm/patches/deepseek_v4_attention.py | 140 ++++++++++++----- vllm/patches/layers/csa_attention.py | 207 ++++++++++++++++++++++---- 2 files changed, 285 insertions(+), 62 deletions(-) diff --git a/vllm/patches/deepseek_v4_attention.py b/vllm/patches/deepseek_v4_attention.py index 641e3f19..83cc6645 100644 --- a/vllm/patches/deepseek_v4_attention.py +++ b/vllm/patches/deepseek_v4_attention.py @@ -560,16 +560,24 @@ class DeepseekV4MultiHeadLatentAttentionWrapper(PluggableLayer): positions: torch.Tensor, out: torch.Tensor, ) -> None: - """Blackwell (SM100+) attention: pure PyTorch, no FlashMLA. + """Blackwell (SM100+) attention: KV cache-based, no FlashMLA. - Same projection flow as the original path, but replaces: - 1. torch.ops._C.fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert - → pure PyTorch RoPE + KV cache insert - 2. self.mla_attn (FlashMLA) → PyTorch SDPA + FIXED: Now writes KV to the paged cache and reads from it during decode. + The previous version used raw kv for attention (no cache write), which + produced garbage during decode because prior tokens' KV was unavailable. + + Pipeline: + 1. Project Q and KV (same as original) + 2. Apply RoPE to Q (in-place) + 3. Write KV to paged cache (RoPE + fp8 quantize + insert) + 4. Prefill: causal attention using raw kv (all tokens available) + 5. Decode: read ALL cached KV from paged cache, then attend """ from vllm.model_executor.layers.csa_attention import ( fused_qnorm_rope_kv_insert_py, - full_sdpa_attention, + blackwell_attention_kv_write, + blackwell_attention_decode, + causal_prefill_attention, ) forward_context = get_forward_context() @@ -600,42 +608,100 @@ class DeepseekV4MultiHeadLatentAttentionWrapper(PluggableLayer): positions, self.indexer_rotary_emb, ) - # RoPE on Q + KV cache insert (pure PyTorch) - if isinstance(attn_metadata, dict): - from vllm.v1.attention.backends.mla.sparse_swa import ( - DeepseekSparseSWAMetadata, - ) - swa_metadata = cast( - "DeepseekSparseSWAMetadata | None", - attn_metadata.get(self.swa_cache_layer.prefix), - ) - if 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) - fused_qnorm_rope_kv_insert_py( - q, kv, swa_kv_cache_2d, - swa_metadata.slot_mapping, - positions.to(torch.int64), - self.rotary_emb.cos_sin_cache, - self.eps, - swa_metadata.block_size, - nope_dim=self.nope_head_dim, - rope_dim=self.rope_head_dim, - ) - else: - self._apply_rope_q(q, positions) - else: + # Get SWA metadata + if not isinstance(attn_metadata, dict): # Dummy run - self._apply_rope_q(q, positions) + fused_qnorm_rope_kv_insert_py( + q, kv, None, None, positions.to(torch.int64), + self.rotary_emb.cos_sin_cache, self.eps, 0, + nope_dim=self.nope_head_dim, + rope_dim=self.rope_head_dim, + ) out.zero_() return - # Apply RoPE to KV (required for correct attention scores) - # Q already has RoPE from fused_qnorm_rope_kv_insert_py or _apply_rope_q - kv_rope = self._apply_rope_kv(kv, positions) + from vllm.v1.attention.backends.mla.sparse_swa import ( + DeepseekSparseSWAMetadata, + ) + swa_metadata = cast( + "DeepseekSparseSWAMetadata | None", + attn_metadata.get(self.swa_cache_layer.prefix), + ) + if swa_metadata is None: + # No SWA metadata (shouldn't happen in normal operation) + fused_qnorm_rope_kv_insert_py( + q, kv, None, None, positions.to(torch.int64), + self.rotary_emb.cos_sin_cache, self.eps, 0, + nope_dim=self.nope_head_dim, + rope_dim=self.rope_head_dim, + ) + out.zero_() + return - # Attention using PyTorch SDPA (works on Blackwell) - o = full_sdpa_attention(q, kv_rope, self.scale) + # Apply per-head RMS norm + RoPE on Q (in-place) + swa_kv_cache = self.swa_cache_layer.kv_cache + swa_kv_cache_2d = swa_kv_cache.view(swa_kv_cache.shape[0], -1) + fused_qnorm_rope_kv_insert_py( + q, kv, swa_kv_cache_2d, + swa_metadata.slot_mapping, + positions.to(torch.int64), + self.rotary_emb.cos_sin_cache, + self.eps, + swa_metadata.block_size, + nope_dim=self.nope_head_dim, + rope_dim=self.rope_head_dim, + ) + + # CRITICAL FIX: Write KV to paged cache (RoPE + fp8 quant + insert) + # This was MISSING in the original Blackwell path, causing garbage output. + # We need an inv_scale cache alongside the fp8 cache. + if not hasattr(self, '_swa_inv_scale_cache'): + max_slots = swa_kv_cache.shape[0] * swa_kv_cache.shape[1] + self._swa_inv_scale_cache = torch.zeros( + max_slots, 1, dtype=torch.bfloat16, device=kv.device, + ) + blackwell_attention_kv_write( + kv, positions, swa_kv_cache, self._swa_inv_scale_cache, + swa_metadata.slot_mapping, swa_metadata.block_size, + self.rotary_emb.cos_sin_cache, + nope_dim=self.nope_head_dim, + rope_dim=self.rope_head_dim, + ) + + # Split prefill and decode + num_decode_tokens = swa_metadata.num_decode_tokens + num_prefills = swa_metadata.num_prefill_tokens + + o = torch.zeros( + hidden_states.shape[0], self.n_local_heads, self.head_dim, + dtype=torch.bfloat16, device=hidden_states.device, + ) + + # ── Decode attention ────────────────────────────────────── + if num_decode_tokens > 0: + q_decode = q[:num_decode_tokens] + pos_decode = positions[:num_decode_tokens] + for t in range(num_decode_tokens): + o[t] = blackwell_attention_decode( + q_decode[t:t+1], pos_decode[t:t+1], + swa_kv_cache, self._swa_inv_scale_cache, + swa_metadata.slot_mapping[t:t+1], + swa_metadata.block_size, + self.scale, + self.window_size, + ).squeeze(0) + + # ── Prefill attention ───────────────────────────────────── + if num_prefills > 0: + q_prefill = q[num_decode_tokens:] + # For prefill, we have all the KV from the current forward pass. + # Apply RoPE to KV and do causal attention. + kv_rope_prefill = self._apply_rope_kv( + kv[num_decode_tokens:], positions[num_decode_tokens:], + ) + o[num_decode_tokens:] = causal_prefill_attention( + q_prefill, kv_rope_prefill, self.scale, + ) # Write into the output buffer (same shape as original path) if self.n_local_heads < self.padded_heads: diff --git a/vllm/patches/layers/csa_attention.py b/vllm/patches/layers/csa_attention.py index 04ad23ff..1e6e4e2b 100644 --- a/vllm/patches/layers/csa_attention.py +++ b/vllm/patches/layers/csa_attention.py @@ -1,8 +1,20 @@ -# SPDX-License: Apache-2.0 +# SPDX-License-Identifier: Apache-2.0 """ CSA/HCA attention for Blackwell (SM100+). -Replaces vLLM's FlashMLA + fused CUDA kernels with pure PyTorch. +Replaces vLLM's FlashMLA + fused CUDA kernels with our own KV cache-based +attention pipeline. The previous version used raw KV for attention (no cache), +which produced garbage during decode because the KV cache was never written. + +Key changes from the broken version: +1. fused_qnorm_rope_kv_insert_py NOW WRITES KV to the paged cache (fp8) +2. full_sdpa_attention is replaced with cache-aware attention +3. KV is quantized to fp8 with per-token scale, RoPE applied before caching + +Architecture: +- KV latent: (T, HD=512) single head, shared across 128 Q heads +- KV Cache: fp8_e4m3 paged cache with per-token inverse scale +- Attention: BF16 (NVFP4 too lossy for Q×K^T) """ import torch @@ -27,10 +39,91 @@ def apply_inv_gptj_rope(x, cos, sin, nope_dim): return out +# ── KV Cache Operations ────────────────────────────────────────────── + +def kv_quantize_fp8(kv_bf16): + """BF16 KV → fp8_e4m3 with per-token inverse scale.""" + amax = kv_bf16.float().abs().amax(dim=-1, keepdim=True).clamp(min=1e-12) + fp8_max = torch.tensor(448.0, dtype=torch.float32, device=kv_bf16.device) + scale = fp8_max / amax + kv_fp8 = (kv_bf16.float() * scale).to(torch.float8_e4m3fn) + inv_scale = (amax / fp8_max).to(torch.bfloat16) + return kv_fp8, inv_scale + + +def kv_dequantize_fp8(kv_fp8, inv_scale): + """fp8 KV → BF16.""" + return (kv_fp8.to(torch.bfloat16) * inv_scale).to(torch.bfloat16) + + +def paged_kv_write(kv_data, slot_mapping, cache, block_size): + """Write KV into paged cache. + + kv_data: (T, D) tensor to write (fp8 or bf16) + slot_mapping: (T,) slot indices + cache: (num_blocks, block_size, D) cache tensor (may be uint8) + """ + # Handle dtype mismatch: cache is uint8, kv_data is fp8 + if cache.dtype == torch.uint8 and kv_data.dtype == torch.float8_e4m3fn: + kv_to_write = kv_data.view(torch.uint8) + else: + kv_to_write = kv_data + + for t in range(kv_data.shape[0]): + slot = slot_mapping[t].item() + block_idx = slot // block_size + offset = slot % block_size + if block_idx < cache.shape[0] and offset < cache.shape[1]: + cache[block_idx, offset] = kv_to_write[t] + + +def paged_kv_read(slot_mapping, cache, block_size, num_tokens, head_dim): + """Read KV from paged cache. Returns fp8 or uint8.""" + device = cache.device + kv = torch.zeros(num_tokens, head_dim, dtype=cache.dtype, device=device) + for t in range(num_tokens): + slot = slot_mapping[t].item() + block_idx = slot // block_size + offset = slot % block_size + if block_idx < cache.shape[0] and offset < cache.shape[1]: + kv[t] = cache[block_idx, offset] + # If cache is uint8, reinterpret as fp8 + if cache.dtype == torch.uint8: + kv = kv.view(torch.float8_e4m3fn) + return kv + + +# ── Attention ───────────────────────────────────────────────────────── + +def causal_prefill_attention(q, kv, scale): + """Full causal self-attention for prefill. q: (T, NH, HD), kv: (T, HD).""" + T, NH, HD = q.shape + q_t = q.permute(1, 0, 2) # (NH, T, HD) + kv_exp = kv.unsqueeze(0).expand(NH, -1, -1) # (NH, T, HD) + out = F.scaled_dot_product_attention(q_t, kv_exp, kv_exp, is_causal=True, scale=scale) + return out.permute(1, 0, 2) # (T, NH, HD) + + +def decode_attention(q, kv, scale): + """Decode attention: 1 query vs N cached KVs. + + q: (1, NH, HD) — single decode token + kv: (N, HD) — all cached KV (already with RoPE) + """ + NH = q.shape[1] + HD = q.shape[2] + q_t = q.permute(1, 0, 2) # (NH, 1, HD) + kv_exp = kv.unsqueeze(0).expand(NH, -1, -1) # (NH, N, HD) + out = F.scaled_dot_product_attention(q_t, kv_exp, kv_exp, is_causal=False, scale=scale) + return out.permute(1, 0, 2) # (1, NH, HD) + + +# ── Fused Q norm + RoPE + KV cache write ───────────────────────────── + def fused_qnorm_rope_kv_insert_py( q, # (T, num_heads, head_dim) — modified in-place kv, # (T, head_dim) — not modified - swa_kv_cache_2d, # paged KV cache — NOT written (cache uses fp8_ds_mla packed format) + swa_kv_cache_2d, # paged KV cache (2D view) slot_mapping, positions, cos_sin_cache, @@ -40,12 +133,11 @@ def fused_qnorm_rope_kv_insert_py( rope_dim, ) -> None: """Pure PyTorch: per-head RMS norm on Q + GPT-J RoPE on Q. - - The C++ kernel also writes KV (RoPE'd + FP8 quantized) to the SWA paged - cache, but the cache uses fp8_ds_mla packed layout which is not a simple - [slot, head_dim] uint8 array. We skip the cache write here — the - compressor (Triton) handles the compressed cache, and for full SDPA - attention we use the raw kv tensor directly. + + NOTE: KV cache write is now handled separately in the attention forward + (blackwell_attention_kv_write), which also applies RoPE and fp8 quantization + before writing to the paged cache. This function only handles Q normalization + and RoPE. """ T = q.shape[0] if T == 0: @@ -65,24 +157,89 @@ def fused_qnorm_rope_kv_insert_py( q[:, :, nope_dim:][:, :, 1::2] = q_rope[:, :, 0::2] * sin_q + q_rope[:, :, 1::2] * cos_q +def blackwell_attention_kv_write( + kv, # (T, HD) kv_normed — NOT RoPE'd yet + positions, # (T,) absolute positions + swa_kv_cache, # (num_blocks, block_size, HD) fp8 paged cache + swa_inv_scale, # (max_slots, 1) per-token inv scale + slot_mapping, # (T,) slot indices + block_size, # tokens per block + cos_sin_cache, # (max_pos, rope_dim) for RoPE + nope_dim, # 448 + rope_dim, # 64 +) -> None: + """Write KV to paged cache: apply RoPE → fp8 quantize → insert. + + This is the function that vLLM's Blackwell path was missing. + Without this, the KV cache is never written, and decode attention + produces garbage because it can't access prior tokens' KV. + """ + T = kv.shape[0] + if T == 0: + return + + # Apply GPT-J RoPE to KV + half = rope_dim // 2 + cos = cos_sin_cache[positions, :half].to(kv.dtype) + sin = cos_sin_cache[positions, half:2 * half].to(kv.dtype) + # Must use original values for both even and odd before modifying + kv_rope_nope = kv[:, nope_dim:].clone() + even = kv_rope_nope[:, 0::2] + odd = kv_rope_nope[:, 1::2] + new_even = even * cos - odd * sin + new_odd = even * sin + odd * cos + kv_rope = kv.clone() + kv_rope[:, nope_dim:][:, 0::2] = new_even + kv_rope[:, nope_dim:][:, 1::2] = new_odd + + # Quantize to fp8 + kv_fp8, inv_scale = kv_quantize_fp8(kv_rope) + + # Write to paged cache + paged_kv_write(kv_fp8, slot_mapping, swa_kv_cache, block_size) + + # Write inv_scale to flat cache + for t in range(T): + slot = slot_mapping[t].item() + swa_inv_scale[slot] = inv_scale[t] + + +def blackwell_attention_decode( + q, # (1, NH, HD) single decode query with RoPE + positions, # (1,) absolute position + swa_kv_cache, # (num_blocks, block_size, HD) fp8 paged cache + swa_inv_scale, # (max_slots, 1) per-token inv scale + slot_mapping, # (1,) slot for the new token (already written) + block_size, # tokens per block + scale, # 1/sqrt(HD) + window_size, # 128 +) -> torch.Tensor: + """Decode attention: read all cached KV, attend. + + Returns: (1, NH, HD) attention output. + """ + pos = positions[0].item() + # Read all KV from position 0 to pos (inclusive) + all_slots = torch.arange(pos + 1, dtype=torch.int64, device=q.device) + kv_cached_fp8 = paged_kv_read(all_slots, swa_kv_cache, block_size, pos + 1, q.shape[2]) + kv_inv_scales = swa_inv_scale[all_slots] + kv_cached = kv_dequantize_fp8(kv_cached_fp8, kv_inv_scales) + + # Apply SWA window + window_start = max(0, pos - window_size + 1) + kv_window = kv_cached[window_start:] + + return decode_attention(q, kv_window, scale) + + def full_sdpa_attention( q: torch.Tensor, # (T, NH, HD) with RoPE kv: torch.Tensor, # (T, HD) KV latent scale: float, ) -> torch.Tensor: - """Full self-attention using manual matmuls (works on all GPUs).""" - T, NH, HD = q.shape - - kv_exp = kv.unsqueeze(1).expand(-1, NH, -1).contiguous() - q_2d = q.reshape(T * NH, 1, HD) - k_2d = kv_exp.permute(1, 0, 2).unsqueeze(1).expand(NH, T, T, -1).contiguous().reshape(T * NH, T, HD) - v_2d = k_2d.clone() - - scores = torch.matmul(q_2d, k_2d.transpose(-1, -2)) * scale - query_pos = torch.arange(T, device=q.device).unsqueeze(1).repeat(1, NH).reshape(T * NH) - kv_pos = torch.arange(T, device=q.device).unsqueeze(0) - causal = kv_pos <= query_pos.unsqueeze(1) - scores = scores.squeeze(1).masked_fill(~causal, float('-inf')) - weights = F.softmax(scores.float(), dim=-1).to(q.dtype) - out = torch.matmul(weights.unsqueeze(1), v_2d).squeeze(1) - return out.reshape(T, NH, HD) + """Full causal self-attention for PREFILL only. + + DEPRECATED: Use causal_prefill_attention instead. + Kept for backward compatibility with the existing vLLM patch. + """ + return causal_prefill_attention(q, kv, scale)