From b443e6702e1d260c455c57c7f708bf52e6aa8a43 Mon Sep 17 00:00:00 2001 From: Woosuk Kwon Date: Tue, 31 Mar 2026 03:38:42 +0000 Subject: [PATCH] fuse mla cache Signed-off-by: Woosuk Kwon --- .../deepseek_v3_2_monolithic/decoder_layer.py | 101 ++++++++--- .../models/deepseek_v3_2_monolithic/ops.py | 164 ++++++++++++------ 2 files changed, 187 insertions(+), 78 deletions(-) diff --git a/vllm/model_executor/models/deepseek_v3_2_monolithic/decoder_layer.py b/vllm/model_executor/models/deepseek_v3_2_monolithic/decoder_layer.py index a2176c2ff93..30d15d4ef4b 100644 --- a/vllm/model_executor/models/deepseek_v3_2_monolithic/decoder_layer.py +++ b/vllm/model_executor/models/deepseek_v3_2_monolithic/decoder_layer.py @@ -165,23 +165,33 @@ class MonolithicDecoderLayer(nn.Module): index_k, _ = self.attn.indexer_wk(hidden_states) index_weights, _ = self.attn.indexer_weights_proj(hidden_states) - # Step 2. Q RMS norm - # + KV RMS norm + KV RoPE - # + Index K layer norm + RoPE + FP8 quant + cache write - # + Init topk indices - # # Fetch slot_mapping early so fused_norm_rope can write FP8 data # directly into the indexer KV cache (saves a separate kernel). from vllm.forward_context import get_forward_context - attn_metadata = get_forward_context().attn_metadata + fwd_ctx = get_forward_context() + attn_metadata = fwd_ctx.attn_metadata if isinstance(attn_metadata, dict): idx_meta = attn_metadata[self.attn.indexer_k_cache.prefix] + # Indexer and MLA caches share the same block_size and track + # the same requests, so their slot_mappings are identical. slot_mapping = idx_meta.slot_mapping else: slot_mapping = None + if slot_mapping is not None: + indexer_k_cache = self.attn.indexer_k_cache.kv_cache + mla_kv_cache = self.attn.mla_attn.kv_cache + mla_k_scale = self.attn.mla_attn._k_scale + else: + indexer_k_cache = None + mla_kv_cache = None + mla_k_scale = None - q_c, kv_c = fused_norm_rope( + # Step 2. Q RMS norm + # + KV RMS norm + KV RoPE + MLA cache write + # + Index K layer norm + RoPE + FP8 quant + cache write + # + Init topk indices + q_c = fused_norm_rope( positions, # Q RMS norm q_c, @@ -202,11 +212,12 @@ class MonolithicDecoderLayer(nn.Module): self.attn.indexer_rope_emb.cos_sin_cache, # Top k indices self.attn.topk_indices_buffer, - # Fused FP8 quant + cache write + # Fused cache writes (single slot_mapping for both caches) slot_mapping=slot_mapping, - indexer_k_cache=self.attn.indexer_k_cache.kv_cache - if slot_mapping is not None - else None, + indexer_k_cache=indexer_k_cache, + mla_kv_cache=mla_kv_cache, + mla_kv_cache_dtype=self.attn.mla_attn.kv_cache_dtype, + mla_k_scale=mla_k_scale, ) # Step 3. q_c -> q @@ -248,16 +259,9 @@ class MonolithicDecoderLayer(nn.Module): self.attn.topk_indices_buffer, ) - # Step 6. MLA attention. - attn_out = self.attn.mla_attn( - q, - kv_c, - k_pe, - output_shape=( - hidden_states.shape[0], - self.attn.num_local_heads * self.attn.v_head_dim, - ), - ) + # Step 6. MLA sparse decode attention (inlined). + # The KV cache update was already done in fused_norm_rope (step 2). + attn_out = self._mla_sparse_decode(q, slot_mapping, hidden_states.shape[0]) # Step 7. Output projection (AllReduce disabled when fused). hidden_states, _ = self.attn.o_proj(attn_out) @@ -278,3 +282,58 @@ class MonolithicDecoderLayer(nn.Module): hidden_states = self.mlp(hidden_states) return hidden_states, residual + + def _mla_sparse_decode( + self, + q: torch.Tensor, + slot_mapping: torch.Tensor | None, + num_padded_tokens: int, + ) -> torch.Tensor: + mla = self.attn.mla_attn + output_shape = (num_padded_tokens, mla.num_heads * mla.v_head_dim) + + from vllm.forward_context import get_forward_context + + fwd_ctx = get_forward_context() + attn_metadata = fwd_ctx.attn_metadata + if isinstance(attn_metadata, dict): + attn_metadata = attn_metadata[mla.layer_name] + if attn_metadata is None or slot_mapping is None: + return torch.zeros(output_shape, dtype=q.dtype, device=q.device) + + num_actual_toks = attn_metadata.num_actual_tokens + q = q[:num_actual_toks] + kv_cache = mla.kv_cache + + fp8_attention = mla.kv_cache_dtype.startswith("fp8") + if fp8_attention and mla.kv_cache_dtype != "fp8_ds_mla": + kv_cache = kv_cache.view(torch.float8_e4m3fn) + + impl = mla.impl + + # 1. Q absorption: q_nope @ W_UK^T → ql_nope + q_nope, q_pe = q.split([mla.qk_nope_head_dim, mla.qk_rope_head_dim], dim=-1) + q_nope = q_nope.transpose(0, 1) # (B, N, P) → (N, B, P) + ql_nope = q_nope.new_empty( + q_nope.shape[0], q_nope.shape[1], mla.W_UK_T.shape[2] + ) + torch.bmm(q_nope, mla.W_UK_T, out=ql_nope) + ql_nope = ql_nope.transpose(0, 1) # (N, B, L) → (B, N, L) + + # 2. FP8 query quantization (if needed) + if fp8_attention and impl.supports_quant_query_input: + mqa_q = mla._decode_concat_quant_fp8_op(ql_nope, q_pe, mla._q_scale) + else: + mqa_q = (ql_nope, q_pe) + + # 3. Forward MQA (topk conversion + FlashInfer kernel) + attn_out, _ = impl.forward_mqa(mqa_q, kv_cache, attn_metadata, mla) + + # 4. V up-projection: attn_out @ W_UV → output + # (N, B, L) x (N, L, V) → (N, B, V) → (B, N*V) + output = torch.empty(output_shape, dtype=q.dtype, device=q.device) + x = attn_out.view(-1, mla.num_heads, mla.kv_lora_rank).transpose(0, 1) + out = output[:num_actual_toks].view(-1, mla.num_heads, mla.v_head_dim) + out = out.transpose(0, 1) + torch.bmm(x, mla.W_UV, out=out) + return output diff --git a/vllm/model_executor/models/deepseek_v3_2_monolithic/ops.py b/vllm/model_executor/models/deepseek_v3_2_monolithic/ops.py index ef01e75c1c7..ef67f5c11e2 100644 --- a/vllm/model_executor/models/deepseek_v3_2_monolithic/ops.py +++ b/vllm/model_executor/models/deepseek_v3_2_monolithic/ops.py @@ -449,8 +449,6 @@ def _fused_norm_rope_kernel( kv_stride, kv_rms_norm_w_ptr, kv_rms_eps, - kv_c_out_ptr, - kv_c_out_stride, KV_DIM: tl.constexpr, # KV RoPE kpe_ptr, @@ -472,12 +470,19 @@ def _fused_norm_rope_kernel( INDEX_K_HALF_ROT_DIM: tl.constexpr, # Index K fp32 scratch buffer for layernorm → RoPE handoff index_k_normed_ptr, - # Index K FP8 quant + cache write + # Cache params (shared by indexer K and MLA) slot_mapping_ptr, - kv_cache_ptr, - kv_cache_scale_ptr, - cache_block_size, - cache_stride, + # Index K FP8 cache + indexer_cache_ptr, + indexer_cache_scale_ptr, + indexer_cache_block_size, + indexer_cache_stride, + # MLA KV cache (concat kv_c_normed + k_pe_roped, uses slot_mapping_ptr) + mla_cache_ptr, + mla_cache_block_stride, + mla_cache_entry_stride, + MLA_CACHE_FP8: tl.constexpr, + mla_cache_scale_ptr, # Top k indices topk_indices_ptr, topk_indices_stride, @@ -486,7 +491,7 @@ def _fused_norm_rope_kernel( ): pid = tl.program_id(0) tok_idx = tl.program_id(1) - if pid == 4: + if pid == 3: # Fill top k indices buffer with -1 for i in range(0, TOPK, TOPK_BLOCK_SIZE): offset = i + tl.arange(0, TOPK_BLOCK_SIZE) @@ -506,7 +511,7 @@ def _fused_norm_rope_kernel( # Padding return - if pid == 1: + if pid == 2: # Q RMS norm q_block = tl.arange(0, Q_BLOCK_SIZE) q_mask = q_block < Q_DIM @@ -514,15 +519,20 @@ def _fused_norm_rope_kernel( q_c_rms_w = tl.load(q_rms_norm_w_ptr + q_block, mask=q_mask) q_c = _rms_norm(q_c, q_c_rms_w, q_rms_eps, Q_DIM) tl.store(q_c_out_ptr + tok_idx * q_c_out_stride + q_block, q_c, mask=q_mask) - elif pid == 3: - # KV RMS Norm + elif pid == 1: + # KV RMS Norm + KV RoPE + MLA concat_and_cache. + # Merged so the normed kv_c and RoPE'd k_pe can be written + # to the MLA KV cache directly without a separate kernel. + + # KV RMS Norm (result stays in registers for MLA cache write) kv_block = tl.arange(0, KV_DIM) kv_c = tl.load(kv_ptr + tok_idx * kv_stride + kv_block) kv_c_rms_w = tl.load(kv_rms_norm_w_ptr + kv_block) kv_c = _rms_norm(kv_c, kv_c_rms_w, kv_rms_eps, KV_DIM) - tl.store(kv_c_out_ptr + tok_idx * kv_c_out_stride + kv_block, kv_c) - elif pid == 2: - # KV RoPE + + # KV RoPE (interleaved) on k_pe — in registers only. + # k_pe is not needed after the cache write (MLA decode reads + # from kv_cache), so we skip writing back to kpe_ptr. pos = tl.load(pos_ptr + tok_idx) cos, sin = _cos_sin_cache_kernel( kpe_rope_cos_sin_cache_ptr, @@ -530,16 +540,39 @@ def _fused_norm_rope_kernel( pos, KPE_HALF_ROT_DIM, ) - _rope_kernel( - kpe_ptr + tok_idx * kpe_stride, - 0, - cos, - sin, - 1, - KPE_HALF_ROT_DIM, - 0, - True, + dim_off = tl.arange(0, KPE_HALF_ROT_DIM) + kpe_base = kpe_ptr + tok_idx * kpe_stride + x1 = tl.load(kpe_base + dim_off * 2).to(tl.float32) + x2 = tl.load(kpe_base + dim_off * 2 + 1).to(tl.float32) + r1 = x1 * cos - x2 * sin + r2 = x2 * cos + x1 * sin + + # MLA concat_and_cache: write [kv_c_normed, k_pe_roped] to cache. + if mla_cache_entry_stride == 0: + return + + mla_block_size = mla_cache_block_stride // mla_cache_entry_stride + mla_block_idx = slot_idx // mla_block_size + mla_block_off = slot_idx % mla_block_size + dst = ( + mla_cache_ptr + + mla_block_idx * mla_cache_block_stride + + mla_block_off * mla_cache_entry_stride ) + # kv_c_normed (KV_DIM elements) + if MLA_CACHE_FP8: + scale = tl.load(mla_cache_scale_ptr) + kv_c_fp8 = (kv_c.to(tl.float32) / scale).to(tl.float8e4nv) + tl.store(dst + kv_block, kv_c_fp8) + else: + tl.store(dst + kv_block, kv_c) + # k_pe_roped (from registers, interleaved layout) + if MLA_CACHE_FP8: + tl.store(dst + KV_DIM + dim_off * 2, (r1 / scale).to(tl.float8e4nv)) + tl.store(dst + KV_DIM + dim_off * 2 + 1, (r2 / scale).to(tl.float8e4nv)) + else: + tl.store(dst + KV_DIM + dim_off * 2, r1) + tl.store(dst + KV_DIM + dim_off * 2 + 1, r2) elif pid == 0: # Fused: Index K LayerNorm + RoPE + FP8 quant + cache write. # Eliminates the separate indexer_k_quant_and_cache kernel launch. @@ -610,10 +643,10 @@ def _fused_norm_rope_kernel( result, index_k_mask, slot_idx, - kv_cache_ptr, - kv_cache_scale_ptr, - cache_block_size, - cache_stride, + indexer_cache_ptr, + indexer_cache_scale_ptr, + indexer_cache_block_size, + indexer_cache_stride, index_k_block, INDEX_K_DIM, ) @@ -635,10 +668,13 @@ def fused_norm_rope( index_k_layer_norm_eps: float, index_k_rope_cos_sin_cache: torch.Tensor, topk_indices_buffer: torch.Tensor, - # Cache params for fused index-k FP8 quant + write + # Cache params for fused writes (single slot_mapping for both caches) slot_mapping: torch.Tensor | None = None, indexer_k_cache: torch.Tensor | None = None, -) -> tuple[torch.Tensor, torch.Tensor]: + mla_kv_cache: torch.Tensor | None = None, + mla_kv_cache_dtype: str = "auto", + mla_k_scale: torch.Tensor | None = None, +) -> torch.Tensor: assert positions.ndim == 1 assert q_c.ndim == 2 assert kv_c.ndim == 2 @@ -651,37 +687,47 @@ def fused_norm_rope( kv_dim = kv_c.shape[-1] index_k_dim = index_k.shape[-1] topk = topk_indices_buffer.shape[-1] + device = positions.device - # When indexer_k_cache is provided, program 0 writes FP8 data + scale - # directly into the cache, eliminating a separate - # indexer_k_quant_and_cache call. + # --- Indexer K cache setup --- if indexer_k_cache is not None: assert slot_mapping is not None - cache_scale_view = indexer_k_cache.view(torch.uint8).view(torch.float32) - cache_block_size = indexer_k_cache.shape[1] - cache_stride = indexer_k_cache.shape[2] - # Ensure the pointer is fp8-typed so tl.store accepts fp8 values. + idx_cache_scale_view = indexer_k_cache.view(torch.uint8).view(torch.float32) + idx_cache_block_size = indexer_k_cache.shape[1] + idx_cache_stride = indexer_k_cache.shape[2] if indexer_k_cache.dtype == torch.uint8: indexer_k_cache = indexer_k_cache.view(torch.float8_e4m3fn) else: - # Dummy values — program 0 will still do LayerNorm + RoPE but - # skip the FP8 cache write (slot_idx will be < 0 for all tokens). - cache_scale_view = torch.empty(0, dtype=torch.float32, device=positions.device) - indexer_k_cache = torch.empty( - 0, dtype=torch.float8_e4m3fn, device=positions.device - ) - slot_mapping = torch.full( - (num_tokens,), -1, dtype=torch.int64, device=positions.device - ) - cache_block_size = 1 - cache_stride = 1 + idx_cache_scale_view = torch.empty(0, dtype=torch.float32, device=device) + indexer_k_cache = torch.empty(0, dtype=torch.float8_e4m3fn, device=device) + slot_mapping = torch.full((num_tokens,), -1, dtype=torch.int64, device=device) + idx_cache_block_size = 1 + idx_cache_stride = 1 + + # --- MLA KV cache setup --- + mla_cache_fp8 = mla_kv_cache_dtype != "auto" + if mla_kv_cache is not None: + mla_block_stride = mla_kv_cache.stride(0) + mla_entry_stride = mla_kv_cache.stride(1) + if mla_cache_fp8 and mla_kv_cache.dtype == torch.uint8: + mla_kv_cache = mla_kv_cache.view(torch.float8_e4m3fn) + if mla_k_scale is None: + mla_k_scale = torch.ones(1, dtype=torch.float32, device=device) + else: + # Dummy values — pid 2 will skip the MLA cache write because + # slot_mapping is all -1. + mla_kv_cache = torch.empty(0, dtype=torch.bfloat16, device=device) + mla_block_stride = 0 + mla_entry_stride = 0 + mla_k_scale = torch.ones(1, dtype=torch.float32, device=device) # fp32 scratch buffer for layernorm output → RoPE handoff. - index_k_normed = torch.empty_like(index_k, dtype=torch.float32) + index_k_normed = torch.empty( + num_tokens, index_k_dim, dtype=torch.float32, device=device + ) q_c_out = torch.empty_like(q_c) - kv_c_out = torch.empty_like(kv_c) - _fused_norm_rope_kernel[(5, num_tokens)]( + _fused_norm_rope_kernel[(4, num_tokens)]( positions, # Q RMS norm q_c, @@ -697,8 +743,6 @@ def fused_norm_rope( kv_c.stride(0), kv_rms_norm_w, kv_rms_eps, - kv_c_out, - kv_c_out.stride(0), kv_dim, # KV RoPE k_pe, @@ -718,19 +762,25 @@ def fused_norm_rope( index_k_rope_cos_sin_cache.stride(0), index_k_rope_cos_sin_cache.shape[-1] // 2, index_k_normed, - # FP8 cache write + # Cache params slot_mapping, indexer_k_cache, - cache_scale_view, - cache_block_size, - cache_stride, + idx_cache_scale_view, + idx_cache_block_size, + idx_cache_stride, + # MLA KV cache (uses same slot_mapping) + mla_kv_cache, + mla_block_stride, + mla_entry_stride, + mla_cache_fp8, + mla_k_scale, # Top k indices buffer topk_indices_buffer, topk_indices_buffer.stride(0), topk, TOPK_BLOCK_SIZE=1024, ) - return q_c_out, kv_c_out + return q_c_out @triton.jit