Triton und MLA: portable Attention und Latent-Space-Attention
Verantwortung
TritonAttentionBackend(triton_attn.py:271) ist ein rein in Python geschriebener Attention-Kernel, den vLLM mit Triton implementiert, ohne das flash_attn-Paket von Dao Labs oder FlashInfer zu benötigen. Seine Daseinsberechtigung ist die Fallback-Rolle (Fallback): Wenn die GPU-Rechenfähigkeit nicht ausreicht, head_size außerhalb des FA-Supports liegt oder die Plattform CPU/XPU ist, kann der Selector (selector) trotzdem eine lauffähige Implementierung liefern. Es unterstützt fp16/bf16/fp32, mehrere FP8/INT4-KV-Cache-Quantisierungen sowie Inline-Scale-Layouts wie fp8_per_token_head / int8_per_token_head (per-token-per-head).
MLA (Latent-Space-Attention / Multi-head Latent Attention, DeepSeek-Serie) ist ein anderer Weg. Es tauscht nicht den Kernel aus, sondern die Semantik des KV-Caches (KV cache): K/V werden nicht expandiert, sondern nur das komprimierte Latent kv_c (typisch kv_lora_rank=512) plus ein kurzes RoPE k_pe gespeichert. Deshalb ist die Backend-Basisklasse von MLA nicht AttentionBackend, sondern MLACommonBackend(mla_attention.py:1206); die impl-Basisklasse ist nicht AttentionImpl, sondern MLAAttentionImpl(backend.py:941), mit der Schnittstelle forward_mqa / forward_mha und ohne Implementierung eines normalen forward. TritonMLABackend(triton_mla.py:81) ist eine solche Implementierung: ein reiner Triton-Decode-Kernel (decode_attention_fwd), der hardwareübergreifend läuft; zusätzlich gibt es in vllm/v1/attention/backends/mla/ weitere MLA-Backends wie FlashAttn / FlashInfer / FlashMLA / Cutlass / Tokenspeed / Aiter (ROCm).
Entwurfsmotivation
- Triton-Fallback deckt breit ab:
TritonAttentionBackend.supported_dtypesenthältfloat32(triton_attn.py:271-287);supported_kv_cache_dtypesenthältint4_per_token_head/int8_per_token_head/fp8_per_token_head— Randkombinationen, die FA/FI nicht unterstützen, laufen hier. - Nicht-kausale Attention unterstützt:
supports_non_causalliefert True(triton_attn.py:301-303), sodass Prefix LM / ViT mit bidirektionaler Attention über Triton laufen können. - Cascade nicht aktiviert:
TritonAttentionBackend.use_cascade_attentionliefert explizitFalse(triton_attn.py:368-370), es wird nur der einfache Varlen-Pfad genutzt. - MLA mit einem KV-Kanal:
MLACommonBackend.get_kv_cache_shapeliefert(num_blocks, block_size, head_size)(mla_attention.py:1215-1223), ohne K/V-Slot-Aufteilung, weil das Latent + RoPE als verketteter Vektor gespeichert und beim Decode per MQA gelesen wird. - MLA schränkt head_size ein:
get_supported_head_sizesakzeptiert nur[320, 576](mla_attention.py:1236-1238), passend zukv_lora_rank + qk_rope_head_dimvon DeepSeek-V2/V3; andere head_size wird direkt abgewiesen. - lse muss zurückgegeben werden:
TritonMLAImpl.can_return_lse_for_decode = True(triton_mla.py:134-135), weil DCP beim Merge über Shards hinweg auf LSE angewiesen ist;lse_base_on_eist standardmäßig True, und der Merge-Kernel wähltcp_lse_ag_out_rs. - Wettbewerb mehrerer Backends:
AttentionBackendEnumlistetTRITON_MLA/FLASH_ATTN_MLA/FLASHINFER_MLA/FLASHMLA/CUTLASS_MLA/TOKENSPEED_MLAauf(registry.py:80-82); der Selector wählt basierend auf GPU-Rechenfähigkeit das optimale.
Schlüsseldateien
TritonAttentionBackend:271— Triton-Backend-Klasse; unterstützt fp32 und mehrere per_token_head-Quantisierungs-KV-Caches.Triton get_kv_cache_shape:317-345— Im per_token_head-Modus wird head_dim gepadded, um Inline-Scale unterzubringen.TritonAttentionImpl.forward:560— Ruft denpaged_attention-Triton-Kernel auf und lehntoutput_block_scaleexplizit ab.MLAAttention:339— Einstieg auf Modellebene; hältkv_b_proj,qk_nope_head_dim,qk_rope_head_dim,kv_lora_rank.MLACommonBackend:1206— MLA-Backend-Basisklasse; definiertis_mla() = Trueund 3D-KV-Cache-Form.MLACommonImpl:1988— MLA-impl-Basisklasse; deklariert abstraktforward_mqa/forward_mha.TritonMLABackend:81— Reines Triton-MLA-Backend, läuft auf jeder GPU.TritonMLAImpl.forward_mqa:189— Ruft dendecode_attention_fwd-Triton-Kernel mitis_mla=True-Zweig auf._compute_num_kv_splits:41-47— Wählt KV-Split-Anzahl nachmax_seq_len / 512, gedeckelt durchsm_count * 2.FlashAttnMLABackend:43— MLA auf FA-Basis; nutztflash_attn_varlen_func+ MLA-spezifische Metadaten.
Datenfluss
Besonderheit von MLA: Q wird vor der Attention nicht zusammen mit K/V berechnet, sondern in q_pe (mit RoPE) und q_nope (direkte Projektion mit dem KV-Latent) aufgespalten. Im KV-Cache (KV cache) liegt der aus [kv_c (kv_lora_rank), k_pe (qk_rope_head_dim)] verkettete Einzelvektor; beim Decode wird er einmal per MQA gelesen und danach über die Head-Dimension per o_proj projiziert. Wenn TritonMLAImpl.forward_mqa decode_attention_fwd aufruft, wird kv_c_and_k_pe_cache als «K-Cache mit einem einzigen Head» übergeben:
# 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 ist nicht hartkodiert: Es wird dynamisch anhand des längsten seq_len im aktuellen Batch berechnet; lange Sequenzen werden stärker aufgeteilt, damit mehr SMs parallel arbeiten, und beim Merge wird über lse ein LogSumExp-Merge durchgeführt. Der Workspace stammt vom workspace_manager, um torch.empty auf dem Decode-Hotpath zu vermeiden.
Grenzen und Fehler
- Triton unterstützt keinen fused block_scale-Output:
TritonAttentionImpl.forwardwirft direktNotImplementedError, wennoutput_block_scale is not None(triton_attn.py:584-588). - MLA unterstützt alibi/sliding_window/logits_soft_cap nicht:
TritonMLAImpl.__init__prüft explizit und wirftNotImplementedError(triton_mla.py:166-170). - MLA head_size auf 320 oder 576 beschränkt:
MLACommonBackend.get_supported_head_sizes(mla_attention.py:1236-1238) liefert nur diese beiden; alles außerhalb DeepSeek-V2/V3 kommt nicht durch. - FP8-KV-Cache läuft mit BF16-Query:
TritonMLAImplsetzt bei FP8-KV-Cachesupports_quant_query_input = False(triton_mla.py:183-185), weil der Triton-Kernel intern dequantisiert und die obere Schicht Q nicht ebenfalls als FP8 packen darf. - Fallback bei nicht-initialisiertem workspace_manager: Wenn
is_workspace_manager_initialized()False liefert, fälltforward_mqaauftorch.emptyzurück(triton_mla.py:225-232); das funktioniert in Unit-Tests, in Produktion mussGPUModelRunnerden Workspace vorab allokiert haben. - num_kv_splits beeinflusst LSE-Merge-Korrektheit: Eine Änderung der Split-Anzahl ändert die Shape von
attn_logits; der Stage2-Merge-Kernel muss denselben Split-Wert sehen. Deshalb ist_compute_num_kv_splits(triton_mla.py:41) sowohl beim Reservieren des Workspace im Builder als auch im Forward dieselbe Funktion.
Zusammenfassung
Triton ist vLLMs Fallback-Backend, wenn FA/FI fehlen, und deckt fp32 sowie randständige Quantisierungs-KV-Caches ab; MLA ist eine für die Latent-Space-Attention der DeepSeek-Serie geschaffene Sub-Abstraktion, deren KV-Cache (KV cache) das komprimierte Latent statt expandiertem K/V speichert. Beide hängen unter der /attention/backend-Abstraktion und werden durch get_attn_backend bei der Initialisierung ausgewählt. Die Form des KV-Caches wird letztlich von /kv-cache/kv-cache-manager zentral verwaltet; der Übergang zum Sampling siehe /sampling/sampler.