class FlashInferMLAImpl(MLACommonImpl[FlashInferMLAMetadata]):
can_return_lse_for_decode: bool = True
supports_dcp: bool = True
# DCP is the only path that consumes LSE. It uses monolithic CuTeDSL,
# whose public LSE contract is natural-log.
lse_base_on_e: bool = 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,
# MLA Specific Arguments
**mla_args,
) -> None:
super().__init__(
num_heads,
head_size,
scale,
num_kv_heads,
alibi_slopes,
sliding_window,
kv_cache_dtype,
logits_soft_cap,
attn_type,
kv_sharing_target_layer_name,
**mla_args,
)
unsupported_features = [alibi_slopes, sliding_window, logits_soft_cap]
if any(unsupported_features):
raise NotImplementedError(
"FlashInferMLAImpl does not support one of the following: "
"alibi_slopes, sliding_window, logits_soft_cap"
)
if attn_type != AttentionType.DECODER:
raise NotImplementedError(
"Encoder self-attention and "
"encoder/decoder cross-attention "
"are not implemented for "
"FlashInferMLAImpl"
)
self.bmm1_scale: float | None = None
self.bmm2_scale: float | None = None
# Worst-case decode batch for the persistent trtllm-gen multi-CTA-KV
# counter buffer (see _get_multi_ctas_kv_counter_buffer). Captured here
# (config is in scope during construction) so the byte size can be
# resolved once on the first decode and never grows after CUDA-graph
# capture. The impl has no _vllm_config, hence get_current_vllm_config().
_sched = get_current_vllm_config().scheduler_config
self._mla_counter_max_batch: int = (
_sched.max_num_batched_tokens or _sched.max_num_seqs
)
self._mla_counter_bytes: int | None = None
def forward_mqa(
self,
q: torch.Tensor | tuple[torch.Tensor, torch.Tensor],
kv_c_and_k_pe_cache: torch.Tensor,
attn_metadata: FlashInferMLAMetadata,
layer: AttentionLayer,
) -> tuple[torch.Tensor, torch.Tensor | None]:
assert kv_c_and_k_pe_cache.numel() > 0
assert attn_metadata.decode is not None
if isinstance(q, tuple):
q_nope, q_pe = q
q = torch.cat([q_nope, q_pe], dim=-1)
block_table = attn_metadata.decode.block_table
seq_lens = attn_metadata.decode.seq_lens
# Led by the promised bound (baked into the graph), not
# num_decode_tokens // num_decodes -- that average is a per-request length
# only for uniform batches.
multi_token_decode = (
attn_metadata.decode.max_query_len > 1
or attn_metadata.num_decode_tokens > attn_metadata.num_decodes
)
cum_seq_lens_q: torch.Tensor | None = None
max_q_len: int | None = None
row_req: torch.Tensor | None = None
if not attn_metadata.causal:
# FlashInfer decode has no causal flag. Flatten each non-causal
# query block into independent single-token rows.
q = q.unsqueeze(1)
if multi_token_decode:
block_table, seq_lens, row_req = self._flattened_decode_metadata(
attn_metadata, q.shape[0]
)
elif attn_metadata.decode.max_query_len > 1 and self.dcp_world_size == 1:
# Causal spec decode: keep q compact and let the kernel tile each
# request's length (uniform 1+k and adaptive ragged). Mirrors vllm #52157.
# DCP keeps the uniform reshape below: it needs LSE, which the kernels
# do not return on the ragged path (flashinfer #3238).
cum_seq_lens_q = attn_metadata.decode.query_start_loc
max_q_len = attn_metadata.decode.max_query_len
# trtllm API requires extra dimension q_len_per_request for MTP
elif attn_metadata.num_decode_tokens % attn_metadata.num_decodes != 0:
logger.warning_once(
"""FlashInferMLAImpl got a query of uneven length.
This usually indicates an issue in batch reordering
or incorrect setup in dummy_run."""
)
q = q.unsqueeze(1)
else:
q = q.view(attn_metadata.num_decodes, -1, q.shape[-2], q.shape[-1])
if self.bmm1_scale is None:
self.bmm1_scale = self.scale
if is_quantized_kv_cache(self.kv_cache_dtype):
self.bmm1_scale *= layer._q_scale_float * layer._k_scale_float
if self.bmm2_scale is None:
self.bmm2_scale = 1.0
if is_quantized_kv_cache(self.kv_cache_dtype):
self.bmm2_scale *= layer._k_scale_float
return_lse = self.need_to_return_lse_for_decode
workspace_buffer = _get_workspace_buffer(return_lse)
# Parallel gathers can change the runtime Q heads from TP-local num_heads.
runtime_num_heads = q.shape[-2]
extra_kwargs: dict[str, Any] = {}
decode_backend: str | None
if self.dcp_world_size > 1:
causal_seqlens_kv_global = attn_metadata.decode.dcp_tot_seq_lens
assert causal_seqlens_kv_global is not None
if row_req is not None:
causal_seqlens_kv_global = causal_seqlens_kv_global[row_req]
extra_kwargs.update(
enable_dcp=True,
cp_world=self.dcp_world_size,
cp_rank=self.dcp_rank,
causal_seqlens_kv_global=causal_seqlens_kv_global,
)
decode_backend = "cute-dsl"
else:
# trtllm-gen rejects MLA head counts it can't tile (e.g. 96);
# fall back to cute-dsl for those.
decode_backend = _select_mla_decode_backend(runtime_num_heads)
if cum_seq_lens_q is not None:
# Neither decode backend returns LSE on the ragged path
# (flashinfer #3238); DCP, the only LSE consumer, took the uniform
# branch above, so this only guards a future caller wiring the two.
assert not return_lse, (
"FlashInferMLA ragged decode cannot return LSE; DCP and adaptive "
"variable-length decode are mutually exclusive."
)
extra_kwargs["cum_seq_lens_q"] = cum_seq_lens_q
extra_kwargs["max_q_len"] = max_q_len
if decode_backend:
extra_kwargs["backend"] = decode_backend
elif kv_c_and_k_pe_cache.shape[-2] in (32, 64):
# The auto path can dispatch to trtllm-gen, whose multi-CTA-KV decode
# kernel self-resets its semaphore counter after each launch (so it
# only needs zeroing once). Pass a persistent counter buffer to skip
# the per-step re-allocate + re-zero the public entry point would
# otherwise do. Guarded to configs where a trtllm-gen runner is
# eligible (page/block size in {32, 64}); the arg is rejected when
# only a cute-dsl runner can run.
if self._mla_counter_bytes is None:
self._mla_counter_bytes = get_trtllm_gen_multi_ctas_kv_counter_bytes(
self._mla_counter_max_batch,
runtime_num_heads,
get_device_sm_count(q.device),
)
extra_kwargs["multi_ctas_kv_counter_buffer"] = (
_get_multi_ctas_kv_counter_buffer(self._mla_counter_bytes, q.device)
)
kernel_out = trtllm_batch_decode_with_kv_cache_mla(
query=q,
kv_cache=kv_c_and_k_pe_cache.unsqueeze(1),
workspace_buffer=workspace_buffer,
qk_nope_head_dim=self.qk_nope_head_dim,
kv_lora_rank=self.kv_lora_rank,
qk_rope_head_dim=self.qk_rope_head_dim,
block_tables=block_table,
seq_lens=seq_lens,
max_seq_len=attn_metadata.max_seq_len,
bmm1_scale=self.bmm1_scale,
bmm2_scale=self.bmm2_scale,
return_lse=return_lse,
**extra_kwargs,
)
if return_lse:
o, lse = kernel_out
lse = lse.view(-1, lse.shape[-1])
else:
o, lse = kernel_out, None
# Flatten the output for consistent shape
o = o.view(-1, o.shape[-2], o.shape[-1])
return o, lse
def _flattened_decode_metadata(
self,
attn_metadata: FlashInferMLAMetadata,
num_rows: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Expand per-request decode tensors to one row per query token.
Cached across the layers of a group, which all see the same batch.
"""
decode = attn_metadata.decode
assert decode is not None
if decode.query_len != num_rows:
cu = decode.query_start_loc
assert cu is not None
# searchsorted on the device offsets, not repeat_interleave(uniform_len):
# the latter only lines up for uniform batches and hands ragged rows
# another request's KV (silent garbage).
rows = torch.arange(num_rows, device=cu.device, dtype=cu.dtype)
row_req = torch.searchsorted(cu[1:], rows, right=True).clamp_(
max=decode.block_table.shape[0] - 1
)
decode.flattened_row_req = row_req
decode.flattened_block_table = decode.block_table[row_req]
decode.flattened_seq_lens = decode.seq_lens[row_req]
decode.query_len = num_rows
assert decode.flattened_block_table is not None
assert decode.flattened_seq_lens is not None
assert decode.flattened_row_req is not None
return (
decode.flattened_block_table,
decode.flattened_seq_lens,
decode.flattened_row_req,
)