fuse mla cache

Signed-off-by: Woosuk Kwon <woosuk@inferact.ai>
This commit is contained in:
Woosuk Kwon
2026-03-31 03:38:42 +00:00
parent 34d73a3375
commit b443e6702e
2 changed files with 187 additions and 78 deletions
@@ -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
@@ -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