Skip to content

vllm.v1.attention.backends.mla.triton_mla

Classes:

TritonMLAMetadataBuilder

Bases: MLACommonMetadataBuilder[MLACommonMetadata]

Source code in vllm/v1/attention/backends/mla/triton_mla.py
class TritonMLAMetadataBuilder(MLACommonMetadataBuilder[MLACommonMetadata]):
    # forward_mqa flattens a uniform multi-token block to one decode row per
    # query token, so causal and non-causal blocks both take the decode path.
    _cudagraph_support: ClassVar[AttentionCGSupport] = AttentionCGSupport.UNIFORM_BATCH
    query_len_support: ClassVar[QueryLenSupport] = QueryLenSupport.UNIFORM
    supports_non_causal_multi_token_decode: ClassVar[bool] = True

    def __init__(self, kv_cache_spec, layer_names, vllm_config, device):
        super().__init__(kv_cache_spec, layer_names, vllm_config, device)
        # DCP local sequence lengths are not advanced between draft steps.
        self.supports_draft_decode_metadata_update = self.dcp_world_size == 1
        self._reserve_attn_logits_workspace()

    def update_draft_decode_metadata(self, _metadata: MLACommonMetadata) -> None:
        pass

    def _reserve_attn_logits_workspace(self) -> None:
        """Pre-size the shared workspace for the decode split-KV attn logits.

        Reserving at the worst case (max_model_len -> max num_kv_splits,
        max_num_seqs decode tokens) before warmup/cudagraph capture means the
        per-call ``get_simultaneous`` in ``forward_mqa`` never has to grow the
        buffer at runtime (which would raise once the workspace is locked).
        """
        if not is_workspace_manager_initialized():
            return
        # forward_mqa flattens each request's block to query_len decode rows,
        # and query_len is bounded by the reorder threshold.
        B = (
            self.vllm_config.scheduler_config.max_num_seqs
            * self.reorder_batch_threshold
        )
        # DCP all-gathers the query heads before forward_mqa.
        q_num_heads = self.num_heads * self.dcp_world_size
        max_splits = _compute_num_kv_splits(
            self.model_config.max_model_len,
            current_platform.num_compute_units(),
        )
        lse_dim = self.mla_dims.kv_lora_rank + 1
        current_workspace_manager().get_simultaneous(
            ((B, q_num_heads, max_splits, lse_dim), torch.float32),
        )

_reserve_attn_logits_workspace()

Pre-size the shared workspace for the decode split-KV attn logits.

Reserving at the worst case (max_model_len -> max num_kv_splits, max_num_seqs decode tokens) before warmup/cudagraph capture means the per-call get_simultaneous in forward_mqa never has to grow the buffer at runtime (which would raise once the workspace is locked).

Source code in vllm/v1/attention/backends/mla/triton_mla.py
def _reserve_attn_logits_workspace(self) -> None:
    """Pre-size the shared workspace for the decode split-KV attn logits.

    Reserving at the worst case (max_model_len -> max num_kv_splits,
    max_num_seqs decode tokens) before warmup/cudagraph capture means the
    per-call ``get_simultaneous`` in ``forward_mqa`` never has to grow the
    buffer at runtime (which would raise once the workspace is locked).
    """
    if not is_workspace_manager_initialized():
        return
    # forward_mqa flattens each request's block to query_len decode rows,
    # and query_len is bounded by the reorder threshold.
    B = (
        self.vllm_config.scheduler_config.max_num_seqs
        * self.reorder_batch_threshold
    )
    # DCP all-gathers the query heads before forward_mqa.
    q_num_heads = self.num_heads * self.dcp_world_size
    max_splits = _compute_num_kv_splits(
        self.model_config.max_model_len,
        current_platform.num_compute_units(),
    )
    lse_dim = self.mla_dims.kv_lora_rank + 1
    current_workspace_manager().get_simultaneous(
        ((B, q_num_heads, max_splits, lse_dim), torch.float32),
    )