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 _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 = len(q_lens)
tile_m = 128 if max_query_len <= 128 else 256
padded_q_len = triton.cdiv(max_query_len, tile_m) * tile_m
num_words_padded = (max_seq_len + 31) // 32 + 1
needed = batch_size * padded_q_len * num_words_padded
if needed * torch.int32.itemsize > GLOBAL_TOPK_MASK_MAX_BYTES:
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
tile_m = 128 if max_seqlen_q <= 128 else 256
padded_q_len = triton.cdiv(max_seqlen_q, tile_m) * tile_m
if dense_mask is None:
batch_size = len(q_lens)
num_words = (max_seqlen_k + 31) // 32
assert topk_mask_workspace is not None
workspace_3d = topk_mask_workspace[
: batch_size * padded_q_len * num_words
].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
use_global_mask = dense_mask is not None
output: torch.Tensor | None = None
output_lse: torch.Tensor | None = None
workspace = chunked_context.workspace
for i, toks in enumerate(chunked_context.seq_tot):
if toks == 0:
continue
ops.gather_and_maybe_dequant_cache(
src_cache=kv_c_and_k_pe_cache,
dst=workspace,
block_table=prefill_metadata.block_table,
cu_seq_lens=chunked_context.cu_seq_lens[i],
token_to_seq=chunked_context.token_to_seq[i],
num_tokens=chunked_context.chunk_total_token[i],
kv_cache_dtype=self.kv_cache_dtype,
scale=k_scale,
seq_starts=chunked_context.starts[i],
)
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)
chunk_lens = chunked_context.seq_lens[i].tolist()
chunk_topk = (
topk_per_req
if use_global_mask
else self._remap_topk_to_ranges(
topk_per_req,
chunked_context.starts[i],
chunk_lens,
)
)
attn_out, lse = self._run_masked_mha(
q=q,
k=k,
v=v,
cu_seqlens_q=prefill_metadata.query_start_loc,
cu_seqlens_k=chunked_context.cu_seq_lens[i],
max_seqlen_q=prefill_metadata.max_query_len,
max_seqlen_k=chunked_context.max_seq_lens[i],
topk_per_req=chunk_topk,
q_lens=q_lens,
causal=False,
return_softmax_lse=True,
dense_mask=dense_mask,
key_starts=(chunked_context.starts[i] if use_global_mask else None),
topk_mask_workspace=prefill_metadata.topk_mask_workspace,
)
if output is None:
output = attn_out
output_lse = lse
else:
assert output_lse is not None
merge_attn_states(
output=output,
output_lse=output_lse,
prefix_output=output,
prefix_lse=output_lse,
suffix_output=attn_out,
suffix_lse=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.seq_lens.sum(dim=0).tolist()
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,
prefill_tokens_with_context=chunked_context.prefill_tokens_with_context,
)