class SparseMLACommonImpl(MLACommonBaseImpl[T], Generic[T]):
"""Sparse MLA base with dense and masked-MHA prefill paths."""
is_sparse = True
def __init__(
self,
num_heads: int,
head_size: int,
scale: float,
num_kv_heads: int,
alibi_slopes: list[float] | None,
sliding_window: int | None,
kv_cache_dtype: str,
logits_soft_cap: float | None,
attn_type: str,
kv_sharing_target_layer_name: str | None,
q_lora_rank: int | None,
kv_lora_rank: int,
qk_nope_head_dim: int,
qk_rope_head_dim: int,
qk_head_dim: int,
v_head_dim: int,
kv_b_proj: "ColumnParallelLinear",
indexer: object | None = None,
topk_indices_buffer: torch.Tensor | None = None,
q_pad_num_heads: int | None = None,
) -> None:
super().__init__(
num_heads,
head_size,
scale,
num_kv_heads,
kv_cache_dtype,
kv_lora_rank,
qk_nope_head_dim,
qk_rope_head_dim,
qk_head_dim,
v_head_dim,
kv_b_proj,
)
# The indexer carries the shared buffer for normal layers and tests;
# the explicitly-passed buffer covers backbone skip layers, whose
# indexer is not constructed (see deepseek_v2.py).
self.topk_indices_buffer: torch.Tensor | None = (
indexer.topk_indices_buffer # type: ignore[attr-defined]
if indexer is not None
else topk_indices_buffer
)
self._use_flashinfer_concat_mla_k = (
has_flashinfer()
and which("ninja") is not None
and (self.num_heads == 128)
and (self.qk_nope_head_dim == 128)
and (self.qk_rope_head_dim == 64)
)
self.masked_mha_available = _is_masked_mha_available(
num_heads_total=num_heads * get_tensor_model_parallel_world_size(),
kv_lora_rank=kv_lora_rank,
qk_nope_head_dim=qk_nope_head_dim,
qk_rope_head_dim=qk_rope_head_dim,
v_head_dim=v_head_dim,
kv_cache_dtype=kv_cache_dtype,
)
@staticmethod
def masked_mha_workspace_fits(prefill: MLACommonPrefillMetadata) -> bool:
"""Whether this prefill batch's top-k masks fit the workspace."""
workspace = prefill.topk_mask_workspace
if workspace is None or prefill.query_lens_cpu is None:
return False
max_context_chunk_seq_len = 0
if prefill.chunked_context is not None:
max_context_chunk_seq_len = max(
chunk.max_seq_len for chunk in prefill.chunked_context.chunks
)
fits = _masked_mha_workspace_fits(
batch_size=len(prefill.query_lens_cpu),
max_query_len=prefill.max_query_len,
max_context_chunk_seq_len=max_context_chunk_seq_len,
workspace_numel=workspace.numel(),
)
if not fits:
logger.warning_once(
"Sparse MLA top-k mask workspace (%d MiB) is too small for some "
"prefill batches; those fall back to slower sparse MQA.",
workspace.numel() * torch.int32.itemsize // (1024 * 1024),
)
return fits
@staticmethod
def _slice_topk_per_req(
topk_all: torch.Tensor,
q_lens: list[int],
) -> list[torch.Tensor]:
topk_per_req = []
offset = 0
for q_len in q_lens:
topk_per_req.append(topk_all[offset : offset + q_len])
offset += q_len
return topk_per_req
@staticmethod
def _remap_topk_to_ranges(
topk_per_req: list[torch.Tensor],
range_starts: list[int] | torch.Tensor,
range_lens: list[int],
) -> list[torch.Tensor]:
remapped = []
for topk, start, length in zip(topk_per_req, range_starts, range_lens):
valid = (topk >= start) & (topk < start + length)
remapped.append(torch.where(valid, topk - start, -1))
return remapped
def _project_kv(
self, kv_c_normed: torch.Tensor, k_pe: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor]:
kv_nope = self.kv_b_proj(kv_c_normed)[0].view(
-1,
self.num_heads,
self.qk_nope_head_dim + self.v_head_dim,
)
k_nope, v = kv_nope.split([self.qk_nope_head_dim, self.v_head_dim], dim=-1)
return self._concat_k_nope_k_pe(k_nope, k_pe), v
@staticmethod
def _try_build_global_mask(
topk_per_req: list[torch.Tensor],
q_lens: list[int],
max_query_len: int,
max_seq_len: int,
topk_mask_workspace: torch.Tensor,
) -> torch.Tensor | None:
"""Build a full-sequence top-k mask if it fits within the budget.
When the mask fits, it is reused across the suffix and all context
chunks, avoiding per-chunk mask rebuilds. Returns None when the
mask is too large, signalling the caller to fall back to per-chunk
index remapping.
"""
batch_size, padded_q_len, num_words_padded = _topk_mask_shape(
len(q_lens),
max_query_len,
max_seq_len,
reserve_key_starts_word=True,
)
needed = batch_size * padded_q_len * num_words_padded
if needed > topk_mask_workspace.numel():
return None
mask = topk_mask_workspace[:needed].view(
batch_size, padded_q_len, num_words_padded
)
_build_topk_mask(
topk_per_req,
q_lens,
padded_q_len,
max_seq_len,
mask,
)
return mask
def _run_masked_mha(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
cu_seqlens_q: torch.Tensor,
cu_seqlens_k: torch.Tensor,
max_seqlen_q: int,
max_seqlen_k: int,
topk_per_req: list[torch.Tensor],
q_lens: list[int],
causal: bool,
return_softmax_lse: bool = False,
dense_mask: torch.Tensor | None = None,
key_starts: torch.Tensor | None = None,
topk_mask_workspace: torch.Tensor | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
from vllm.model_executor.layers.attention.sparse_mla_mask import (
dense_mask_mod,
offset_dense_mask_mod,
)
from vllm.vllm_flash_attn import flash_attn_varlen_func
if dense_mask is None:
assert topk_mask_workspace is not None
batch_size, padded_q_len, num_words = _topk_mask_shape(
len(q_lens), max_seqlen_q, max_seqlen_k
)
words_needed = batch_size * padded_q_len * num_words
if words_needed > topk_mask_workspace.numel():
raise ValueError(
f"Sparse MLA top-k mask needs {words_needed} int32 words (batch="
f"{len(q_lens)}, q={max_seqlen_q}, k={max_seqlen_k}) but the "
f"workspace holds {topk_mask_workspace.numel()}."
)
workspace_3d = topk_mask_workspace[:words_needed].view(
batch_size, padded_q_len, num_words
)
dense_mask = _build_topk_mask(
topk_per_req,
q_lens,
padded_q_len,
max_seqlen_k,
workspace_3d,
)
if key_starts is not None:
dense_mask[:, 0, -1].copy_(key_starts)
kwargs = {
"q": q,
"k": k,
"v": v,
"cu_seqlens_q": cu_seqlens_q,
"cu_seqlens_k": cu_seqlens_k,
"max_seqlen_q": max_seqlen_q,
"max_seqlen_k": max_seqlen_k,
"softmax_scale": self.scale,
"return_softmax_lse": return_softmax_lse,
"fa_version": 4,
"mask_mod": dense_mask_mod if key_starts is None else offset_dense_mask_mod,
"aux_tensors": [dense_mask],
"aux_tensor_leading_dims": [2],
"causal": causal,
}
return flash_attn_varlen_func(**kwargs)
def _compute_context_mha(
self,
q: torch.Tensor,
kv_c_and_k_pe_cache: torch.Tensor,
prefill_metadata: MLACommonPrefillMetadata,
k_scale: torch.Tensor,
q_lens: list[int],
topk_per_req: list[torch.Tensor],
dense_mask: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
if self.dcp_world_size > 1:
raise NotImplementedError(
"Masked MHA with context does not yet support decode context "
"parallelism"
)
chunked_context = prefill_metadata.chunked_context
assert chunked_context is not None
output: torch.Tensor | None = None
output_lse: torch.Tensor | None = None
workspace = chunked_context.workspace
for chunk in chunked_context.chunks:
toks = chunk.num_context_tokens
requests = chunk.request_slice
ops.gather_and_maybe_dequant_cache(
src_cache=kv_c_and_k_pe_cache,
dst=workspace,
block_table=prefill_metadata.block_table[requests],
cu_seq_lens=chunk.cu_seq_lens,
token_to_seq=chunk.token_to_seq,
num_tokens=toks,
kv_cache_dtype=self.kv_cache_dtype,
scale=k_scale,
seq_starts=chunk.starts,
)
chunk_kv_c = workspace[:toks, : self.kv_lora_rank]
chunk_k_pe = workspace[:toks, self.kv_lora_rank :].unsqueeze(1)
k, v = self._project_kv(chunk_kv_c, chunk_k_pe)
if dense_mask is not None:
chunk_mask: torch.Tensor | None = dense_mask[requests]
chunk_topk = topk_per_req[requests]
key_starts: torch.Tensor | None = chunk.starts
else:
chunk_mask = None
chunk_topk = self._remap_topk_to_ranges(
topk_per_req[requests],
chunk.starts,
chunk.seq_lens.tolist(),
)
key_starts = None
attn_out, lse = self._run_masked_mha(
q=q[chunk.token_slice],
k=k,
v=v,
cu_seqlens_q=chunk.query_start_loc,
cu_seqlens_k=chunk.cu_seq_lens,
max_seqlen_q=chunk.max_query_len,
max_seqlen_k=chunk.max_seq_len,
topk_per_req=chunk_topk,
q_lens=q_lens[requests],
causal=False,
return_softmax_lse=True,
dense_mask=chunk_mask,
key_starts=key_starts,
topk_mask_workspace=prefill_metadata.topk_mask_workspace,
)
if output is None:
if (
len(chunked_context.chunks) == 1
and not chunked_context.empty_token_slices
):
return attn_out, lse
output, output_lse = init_mla_context_partial(
chunked_context,
attn_out,
lse,
num_tokens=q.shape[0],
)
accumulate_mla_context_chunk(chunk, attn_out, lse, output, output_lse)
assert output is not None and output_lse is not None
return output, output_lse
def forward_mha( # type: ignore[override]
self,
q: torch.Tensor,
kv_c_normed: torch.Tensor,
k_pe: torch.Tensor,
kv_c_and_k_pe_cache: torch.Tensor,
attn_metadata: T,
k_scale: torch.Tensor,
output: torch.Tensor,
output_scale: torch.Tensor | None = None,
) -> None:
prefill_max_seq_len = attn_metadata.prefill_max_seq_len # type: ignore[attr-defined]
topk_tokens = attn_metadata.topk_tokens # type: ignore[attr-defined]
force_dense = getattr(self, "_sparse_mla_force_dense_mha", False)
force_masked = getattr(self, "_sparse_mla_force_masked_mha", False)
if force_dense or (prefill_max_seq_len <= topk_tokens and not force_masked):
return super().forward_mha(
q,
kv_c_normed,
k_pe,
kv_c_and_k_pe_cache,
cast(MLACommonMetadata, attn_metadata),
k_scale,
output,
output_scale,
)
assert output_scale is None
assert self.masked_mha_available
prefill_metadata = attn_metadata.prefill # type: ignore[attr-defined]
assert prefill_metadata is not None
assert prefill_metadata.query_lens_cpu is not None
assert self.topk_indices_buffer is not None
q_lens = prefill_metadata.query_lens_cpu.tolist()
num_decode_tokens = attn_metadata.num_decode_tokens # type: ignore[attr-defined]
topk_all = self.topk_indices_buffer[
num_decode_tokens : num_decode_tokens + q.shape[0]
]
topk_per_req = self._slice_topk_per_req(topk_all, q_lens)
k, v = self._project_kv(kv_c_normed, k_pe)
chunked_context = prefill_metadata.chunked_context
if chunked_context is None:
attn_out = self._run_masked_mha(
q=q,
k=k,
v=v,
cu_seqlens_q=prefill_metadata.query_start_loc,
cu_seqlens_k=prefill_metadata.query_start_loc,
max_seqlen_q=prefill_metadata.max_query_len,
max_seqlen_k=prefill_metadata.max_query_len,
topk_per_req=topk_per_req,
q_lens=q_lens,
causal=True,
topk_mask_workspace=prefill_metadata.topk_mask_workspace,
)
assert isinstance(attn_out, torch.Tensor)
output.copy_(attn_out[..., : self.v_head_dim].flatten(start_dim=-2))
return
context_lens = chunked_context.context_lens_list
dense_mask = self._try_build_global_mask(
topk_per_req,
q_lens,
prefill_metadata.max_query_len,
prefill_max_seq_len,
prefill_metadata.topk_mask_workspace,
)
if dense_mask is not None:
suffix_topk = topk_per_req
else:
suffix_topk = self._remap_topk_to_ranges(topk_per_req, context_lens, q_lens)
suffix_output, suffix_lse = self._run_masked_mha(
q=q,
k=k,
v=v,
cu_seqlens_q=prefill_metadata.query_start_loc,
cu_seqlens_k=prefill_metadata.query_start_loc,
max_seqlen_q=prefill_metadata.max_query_len,
max_seqlen_k=prefill_metadata.max_query_len,
topk_per_req=suffix_topk,
q_lens=q_lens,
causal=True,
return_softmax_lse=True,
dense_mask=dense_mask,
key_starts=(
chunked_context.context_lens if dense_mask is not None else None
),
topk_mask_workspace=prefill_metadata.topk_mask_workspace,
)
context_output, context_lse = self._compute_context_mha(
q=q,
kv_c_and_k_pe_cache=kv_c_and_k_pe_cache,
prefill_metadata=prefill_metadata,
k_scale=k_scale,
q_lens=q_lens,
topk_per_req=topk_per_req,
dense_mask=dense_mask,
)
merge_attn_states(
output=output.view(-1, self.num_heads, self.v_head_dim),
prefix_output=context_output[..., : self.v_head_dim],
prefix_lse=context_lse,
suffix_output=suffix_output[..., : self.v_head_dim],
suffix_lse=suffix_lse,
)