Skip to content

Triton et MLA : attention portable et attention à espace latent

源码版本v0.25.1

Responsabilités

TritonAttentionBackend(triton_attn.py:271) est un kernel d'attention en pur Python écrit avec Triton, qui ne dépend ni du package flash_attn de Dao Labs ni de FlashInfer. Son rôle est le « fallback » : quand la capacité GPU est insuffisante, que head_size n'est pas dans la plage supportée par FA, ou que la plateforme est CPU/XPU, le selector peut toujours fournir une implémentation fonctionnelle. Il supporte fp16/bf16/fp32, plusieurs quantifications de cache KV FP8/INT4, ainsi qu'un layout d'inline scale per-token-per-head comme fp8_per_token_head / int8_per_token_head.

MLA (attention à espace latent / Multi-head Latent Attention, séries DeepSeek) est une autre voie. Il ne s'agit pas de « changer de kernel » mais de « changer la sémantique de la cache KV » : K/V ne sont pas développés, on ne met en cache que le latent compressé kv_c (typiquement kv_lora_rank=512) plus un court segment RoPE k_pe. De ce fait, la classe de base du backend MLA n'est pas AttentionBackend mais MLACommonBackend(mla_attention.py:1206) ; la classe de base de l'impl n'est pas AttentionImpl mais MLAAttentionImpl(backend.py:941), avec pour interface forward_mqa / forward_mha, sans forward standard. TritonMLABackend(triton_mla.py:81) est une implémentation parmi d'autres : kernel decode en pur Triton (decode_attention_fwd), qui tourne sur tout matériel ; d'autres backends MLA existent (FlashAttn / FlashInfer / FlashMLA / Cutlass / Tokenspeed / Aiter pour ROCm), tous sous vllm/v1/attention/backends/mla/.

Motivation de conception

  • Couverture large du fallback Triton : TritonAttentionBackend.supported_dtypes inclut float32(triton_attn.py:271-287) ; supported_kv_cache_dtypes inclut int4_per_token_head / int8_per_token_head / fp8_per_token_head — les combinaisons marginales que FA/FI ne gèrent pas passent ici.
  • Attention non causale supportée : supports_non_causal renvoie True(triton_attn.py:301-303), ce qui permet à Prefix LM / ViT d'utiliser Triton pour l'attention bidirectionnelle.
  • Pas de cascade : TritonAttentionBackend.use_cascade_attention renvoie explicitement False(triton_attn.py:368-370) ; seul le chemin varlen naïf est emprunté.
  • MLA : canal KV unique : MLACommonBackend.get_kv_cache_shape renvoie (num_blocks, block_size, head_size)(mla_attention.py:1215-1223) ; pas de séparation K/V, car ce qui est caché est le vecteur latent + RoPE concaténé, lu en MQA au decode.
  • head_size restreint pour MLA : get_supported_head_sizes ne reconnaît que [320, 576](mla_attention.py:1236-1238), qui correspondent aux combinaisons kv_lora_rank + qk_rope_head_dim de DeepSeek-V2/V3 ; tout autre head_size est rejeté.
  • LSE obligatoire : TritonMLAImpl.can_return_lse_for_decode = True(triton_mla.py:134-135) car la fusion inter-shards DCP s'appuie sur la LSE ; lse_base_on_e vaut True par défaut, et la fusion inter-shards passe par cp_lse_ag_out_rs.
  • Concurrence de plusieurs backends : AttentionBackendEnum énumère TRITON_MLA / FLASH_ATTN_MLA / FLASHINFER_MLA / FLASHMLA / CUTLASS_MLA / TOKENSPEED_MLA(registry.py:80-82) ; le selector choisit le plus performant selon la capacité GPU.

Fichiers clés

Flux de données

La particularité de MLA : Q n'est pas calculé avec K/V avant l'attention, mais séparé en q_pe (passe par RoPE) et q_nope (projette directement avec le latent KV). La cache KV contient le vecteur unique [kv_c (kv_lora_rank), k_pe (qk_rope_head_dim)] concaténé, lu une fois en MQA au decode, puis une o_proj projette en dimension de head. TritonMLAImpl.forward_mqa passe kv_c_and_k_pe_cache à decode_attention_fwd comme « K cache mono-head » :

python
# vllm/v1/attention/backends/mla/triton_mla.py L189-L256
def forward_mqa(
    self,
    q: torch.Tensor | tuple[torch.Tensor, torch.Tensor],
    kv_c_and_k_pe_cache: torch.Tensor,
    attn_metadata: MLACommonMetadata,
    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 type(q) is tuple:
        q = torch.cat(q, dim=-1)

    assert isinstance(q, torch.Tensor)
    B = q.shape[0]
    q_num_heads = q.shape[1]
    o = torch.zeros(
        B, q_num_heads, self.kv_lora_rank, dtype=q.dtype, device=q.device
    )
    lse = torch.zeros(B, q_num_heads, dtype=q.dtype, device=q.device)

    # For batch invariance, use only 1 split to ensure deterministic reduction
    if envs.VLLM_BATCH_INVARIANT:
        num_kv_splits = 1
    else:
        num_kv_splits = _compute_num_kv_splits(
            attn_metadata.max_seq_len, self._sm_count
        )

    # NOTE: the +1 stores the LogSumExp (LSE) that the stage2 kernel uses
    # to merge partial attention outputs across splits.
    logits_shape = (B, q_num_heads, num_kv_splits, self.kv_lora_rank + 1)
    if is_workspace_manager_initialized():
        (attn_logits,) = current_workspace_manager().get_simultaneous(
            (logits_shape, torch.float32),
        )
    else:
        attn_logits = torch.empty(
            logits_shape, dtype=torch.float32, device=q.device
        )

    # Add a head dim of 1
    kv_c_and_k_pe_cache = kv_c_and_k_pe_cache.unsqueeze(2)
    kv_c_cache = kv_c_and_k_pe_cache[..., : self.kv_lora_rank]
    PAGE_SIZE = kv_c_and_k_pe_cache.size(1)

    # Run MQA — always pass layer scales. When KV cache is
    # BF16 the kernel's `if dtype.is_fp8()` check is a no-op.
    decode_attention_fwd(
        q,
        kv_c_and_k_pe_cache,
        kv_c_cache,
        o,
        lse,
        attn_metadata.decode.block_table,
        attn_metadata.decode.seq_lens,
        attn_logits,
        num_kv_splits,
        self.scale,
        PAGE_SIZE,
        k_scale=layer._k_scale,
        v_scale=layer._k_scale,
        is_mla=True,
    )

    return o, lse

num_kv_splits n'est pas codé en dur : il est calculé dynamiquement à partir du seq_len max du batch, pour découper les longues séquences en plusieurs splits et faire travailler les SM en parallèle ; la phase de fusion s'appuie sur lse pour le LogSumExp. L'espace de travail est récupéré via workspace_manager afin d'éviter un torch.empty sur le hot path decode.

Limites et échecs

  • Triton ne supporte pas la sortie fused block_scale : TritonAttentionImpl.forward lève NotImplementedError si output_block_scale is not None(triton_attn.py:584-588).
  • MLA ne supporte pas alibi/sliding_window/logits_soft_cap : TritonMLAImpl.__init__ vérifie et lève NotImplementedError(triton_mla.py:166-170).
  • head_size MLA restreint à 320 ou 576 : MLACommonBackend.get_supported_head_sizes(mla_attention.py:1236-1238) ne renvoie que ces deux valeurs ; tout ce qui n'est pas DeepSeek-V2/V3 est essentiellement exclu.
  • Cache KV FP8 implique query BF16 : TritonMLAImpl met supports_quant_query_input = False en cache KV FP8(triton_mla.py:183-185), car le kernel Triton fait lui-même le dequant et la couche supérieure ne peut pas aussi compresser Q en FP8.
  • Fallback quand workspace_manager n'est pas initialisé : forward_mqa retombe sur torch.empty si is_workspace_manager_initialized() renvoie False(triton_mla.py:225-232), suffisant pour les tests, mais en production c'est GPUModelRunner qui doit préallouer le workspace.
  • num_kv_splits impacte la correction de la fusion LSE : un nombre de splits différent change le shape de attn_logits, et le kernel stage2 de merge doit utiliser la même valeur — c'est pourquoi _compute_num_kv_splits(triton_mla.py:41) est la même fonction utilisée à la réservation du workspace côté builder et à l'exécution côté forward.

Résumé

Triton est le backend fallback de vLLM quand ni FA ni FI ne sont disponibles, couvrant fp32 et les caches KV quantifiées marginales ; MLA est une sous-abstraction dédiée à l'attention à espace latent des séries DeepSeek, où la cache KV stocke le latent compressé plutôt que K/V développé. Les deux se rattachent à la couche d'abstraction /attention/backend et sont sélectionnés à l'initialisation par get_attn_backend. La forme de la cache KV côté Triton et MLA est in fine gérée de façon unifiée par /kv-cache/kv-cache-manager ; la connexion au sampling est sur /sampling/sampler.

Voir la documentation officielle : Documentation vLLM · README