Skip to content

Triton 與 MLA:可移植注意力與潛空間注意力

源码版本v0.25.1

職責

TritonAttentionBackend(triton_attn.py:271)是 vLLM 用 Triton 寫的純 Python 注意力 kernel,不依賴 Dao Labs 的 flash_attn 包也不依賴 FlashInfer。它的存在意義是「兜底」:當 GPU 算力不夠、head_size 不在 FA 支持範圍、或運行平臺是 CPU/XPU 時,selector 仍然能給出一個能跑的實現。它支持 fp16/bf16/fp32、多種 FP8/INT4 KV cache 量化,以及 fp8_per_token_head / int8_per_token_head 這種 per-token-per-head 的 inline scale layout。

MLA(潛空間注意力 / Multi-head Latent Attention,DeepSeek 系列)是另一條路。它不是「換 kernel」而是「換 KV cache 語義」:K/V 不展開,只快取壓縮後的 latent kv_c(典型 kv_lora_rank=512)加一小段 RoPE k_pe。所以 MLA 的後端基類不是 AttentionBackend,而是 MLACommonBackend(mla_attention.py:1206);impl 基類不是 AttentionImpl 而是 MLAAttentionImpl(backend.py:941),介面是 forward_mqa / forward_mha,不實現普通 forwardTritonMLABackend(triton_mla.py:81)是其中一種實現,純 Triton 寫 decode kernel (decode_attention_fwd),跨硬件都能跑;另外還有 FlashAttn / FlashInfer / FlashMLA / Cutlass / Tokenspeed / Aiter(ROCm)等多個 MLA 後端,都在 vllm/v1/attention/backends/mla/ 下。

設計動機

  • Triton 兜底覆蓋廣:TritonAttentionBackend.supported_dtypesfloat32(triton_attn.py:271-287),supported_kv_cache_dtypesint4_per_token_head / int8_per_token_head / fp8_per_token_head,FA/FI 不支持的邊緣組合在這裡能跑。
  • 支持非因果注意力:supports_non_causal 返回 True(triton_attn.py:301-303),Prefix LM / ViT 這種雙向注意力能走 Triton。
  • 不啟用 cascade:TritonAttentionBackend.use_cascade_attention 顯式返回 False(triton_attn.py:368-370),只走樸素 varlen 路徑。
  • MLA 單 KV 通道:MLACommonBackend.get_kv_cache_shape 返回 (num_blocks, block_size, head_size)(mla_attention.py:1215-1223),沒有 K/V 分槽,因為快取的是 latent + RoPE 拼接向量,decode 時按 MQA 取。
  • MLA 限定 head_size:get_supported_head_sizes 只認 [320, 576](mla_attention.py:1236-1238),對應 DeepSeek-V2/V3 的 kv_lora_rank + qk_rope_head_dim 組合,其它 head_size 直接拒絕。
  • lse 必須返回:TritonMLAImpl.can_return_lse_for_decode = True(triton_mla.py:134-135),因為 DCP 跨分片合併要靠 LSE;lse_base_on_e 預設 True,跨分片合併 kernel 走 cp_lse_ag_out_rs
  • 多後端競爭:AttentionBackendEnumTRITON_MLA / FLASH_ATTN_MLA / FLASHINFER_MLA / FLASHMLA / CUTLASS_MLA / TOKENSPEED_MLA 全列出來(registry.py:80-82),selector 按 GPU 算力選最優。

關鍵檔案

資料流

MLA 的特殊之處:Q 在 attention 之前沒和 K/V 一起算,而是分成 q_pe(走 RoPE)和 q_nope(直接和 KV latent 做投影)。KV cache 裡存的是 [kv_c (kv_lora_rank), k_pe (qk_rope_head_dim)] 拼成的單向量,decode 時按 MQA 取一次,然後在 head 維做 o_projTritonMLAImpl.forward_mqa 調 decode_attention_fwd 時把 kv_c_and_k_pe_cache 當成「單 head 的 K cache」傳進去:

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 不是寫死的:它按當前 batch 裡最長的 seq_len 動態算,長序列多切幾片讓多 SM 並行,合併階段靠 lse 做 LogSumExp 歸併。工作區從 workspace_manager 取,避免 decode 熱路徑上 torch.empty

邊界與失敗

  • Triton 不支持 fused block_scale 輸出:TritonAttentionImpl.forward 拿到 output_block_scale is not None 直接拋 NotImplementedError(triton_attn.py:584-588)。
  • MLA 不支持 alibi/sliding_window/logits_soft_cap:TritonMLAImpl.__init__ 顯式檢查並拋 NotImplementedError(triton_mla.py:166-170)。
  • MLA head_size 限定 320 或 576:MLACommonBackend.get_supported_head_sizes(mla_attention.py:1236-1238)只返回這兩個,DeepSeek-V2/V3 之外的基本進不來。
  • FP8 KV cache 走 BF16 query:TritonMLAImpl 在 FP8 KV cache 時把 supports_quant_query_input = False(triton_mla.py:183-185),因為 Triton kernel 內部做 dequant,不能讓上層把 Q 也壓成 FP8。
  • workspace_manager 未初始化時 fallback:forward_mqais_workspace_manager_initialized() 返回 False 時退回 torch.empty(triton_mla.py:225-232),單元測試場景能跑,但生產路徑要靠 GPUModelRunner 把 workspace 預分配好。
  • num_kv_splits 影響 LSE 合併正確性:splits 數變化會讓 attn_logits 的 shape 變,stage2 merge kernel 必須用同一份 splits 值,所以 _compute_num_kv_splits(triton_mla.py:41)在 builder 預留 workspace 時和 forward 時是同一個函數。

小結

Triton 是 vLLM 在沒有 FA/FI 時的兜底後端,覆蓋 fp32 和邊緣量化 KV cache;MLA 是為 DeepSeek 系潛空間注意力專門開闢的子抽象,KV cache 存壓縮 latent 而不是展開的 K/V。兩者都掛在 /attention/backend 這層抽象之下,被 get_attn_backend 在初始化時選定。Triton 與 MLA 的 KV cache 形狀最終由 /kv-cache/kv-cache-manager 統一管理,採樣銜接見 /sampling/sampler

對照官方資料:vLLM 文件 · README