Triton et MLA : attention portable et attention à espace latent
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_dtypesinclutfloat32(triton_attn.py:271-287) ;supported_kv_cache_dtypesinclutint4_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_causalrenvoie 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_attentionrenvoie explicitementFalse(triton_attn.py:368-370) ; seul le chemin varlen naïf est emprunté. - MLA : canal KV unique :
MLACommonBackend.get_kv_cache_shaperenvoie(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_sizesne reconnaît que[320, 576](mla_attention.py:1236-1238), qui correspondent aux combinaisonskv_lora_rank + qk_rope_head_dimde 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_evaut True par défaut, et la fusion inter-shards passe parcp_lse_ag_out_rs. - Concurrence de plusieurs backends :
AttentionBackendEnuménumèreTRITON_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
TritonAttentionBackend:271— classe du backend Triton, supporte fp32 et plusieurs caches KV quantifiées per_token_head.Triton get_kv_cache_shape:317-345— en mode per_token_head, pad le head_dim pour loger l'inline scale.TritonAttentionImpl.forward:560— appelle le kernel Tritonpaged_attention, rejette explicitementoutput_block_scale.MLAAttention:339— entrée côté modèle, détientkv_b_proj,qk_nope_head_dim,qk_rope_head_dim,kv_lora_rank.MLACommonBackend:1206— classe de base du backend MLA, définitis_mla() = Trueet une cache KV 3D.MLACommonImpl:1988— classe de base de l'impl MLA, déclare les abstract methodsforward_mqa/forward_mha.TritonMLABackend:81— backend MLA en pur Triton, fonctionne sur tout GPU.TritonMLAImpl.forward_mqa:189— appelle le kernel Tritondecode_attention_fwd, avec une brancheis_mla=True._compute_num_kv_splits:41-47— choisit le nombre de splits KV selonmax_seq_len / 512, plafonné àsm_count * 2.FlashAttnMLABackend:43— variante MLA via FA, qui utiliseflash_attn_varlen_func+ metadata dédiée MLA.
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 » :
# 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, lsenum_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.forwardlèveNotImplementedErrorsioutput_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èveNotImplementedError(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 :
TritonMLAImplmetsupports_quant_query_input = Falseen 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_mqaretombe surtorch.emptysiis_workspace_manager_initialized()renvoie False(triton_mla.py:225-232), suffisant pour les tests, mais en production c'estGPUModelRunnerqui doit préallouer le workspace. num_kv_splitsimpacte la correction de la fusion LSE : un nombre de splits différent change le shape deattn_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