Skip to content

Triton und MLA: portable Attention und Latent-Space-Attention

源码版本v0.25.1

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_dtypes enthält float32(triton_attn.py:271-287); supported_kv_cache_dtypes enthält int4_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_causal liefert True(triton_attn.py:301-303), sodass Prefix LM / ViT mit bidirektionaler Attention über Triton laufen können.
  • Cascade nicht aktiviert: TritonAttentionBackend.use_cascade_attention liefert explizit False(triton_attn.py:368-370), es wird nur der einfache Varlen-Pfad genutzt.
  • MLA mit einem KV-Kanal: MLACommonBackend.get_kv_cache_shape liefert (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_sizes akzeptiert nur [320, 576](mla_attention.py:1236-1238), passend zu kv_lora_rank + qk_rope_head_dim von 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_e ist standardmäßig True, und der Merge-Kernel wählt cp_lse_ag_out_rs.
  • Wettbewerb mehrerer Backends: AttentionBackendEnum listet TRITON_MLA / FLASH_ATTN_MLA / FLASHINFER_MLA / FLASHMLA / CUTLASS_MLA / TOKENSPEED_MLA auf(registry.py:80-82); der Selector wählt basierend auf GPU-Rechenfähigkeit das optimale.

Schlüsseldateien

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:

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 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.forward wirft direkt NotImplementedError, wenn output_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 wirft NotImplementedError(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: TritonMLAImpl setzt bei FP8-KV-Cache supports_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ällt forward_mqa auf torch.empty zurück(triton_mla.py:225-232); das funktioniert in Unit-Tests, in Produktion muss GPUModelRunner den 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.

Siehe offizielle Dokumentation: vLLM 文档 · README.