Triton 與 MLA:可移植注意力與潛空間注意力
職責
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,不實現普通 forward。TritonMLABackend(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_dtypes含float32(triton_attn.py:271-287),supported_kv_cache_dtypes含int4_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。 - 多後端競爭:
AttentionBackendEnum把TRITON_MLA/FLASH_ATTN_MLA/FLASHINFER_MLA/FLASHMLA/CUTLASS_MLA/TOKENSPEED_MLA全列出來(registry.py:80-82),selector 按 GPU 算力選最優。
關鍵檔案
TritonAttentionBackend:271— Triton 後端類,支持 fp32 和多種 per_token_head 量化 KV cache。Triton get_kv_cache_shape:317-345— per_token_head 模式下會 pad head_dim 容納 inline scale。TritonAttentionImpl.forward:560— 調paged_attentionTriton kernel,顯式拒絕output_block_scale。MLAAttention:339— 模型層入口,持kv_b_proj、qk_nope_head_dim、qk_rope_head_dim、kv_lora_rank。MLACommonBackend:1206— MLA 後端基類,定義is_mla() = True和 3D KV cache shape。MLACommonImpl:1988— MLA impl 基類,聲明forward_mqa/forward_mha抽象方法。TritonMLABackend:81— 純 Triton MLA 後端,所有 GPU 都能跑。TritonMLAImpl.forward_mqa:189— 調decode_attention_fwdTriton kernel,帶is_mla=True分支。_compute_num_kv_splits:41-47— 按max_seq_len / 512選 KV 分片數,封頂sm_count * 2。FlashAttnMLABackend:43— FA 路線 MLA,用flash_attn_varlen_func+ MLA 專用 metadata。
資料流
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_proj。TritonMLAImpl.forward_mqa 調 decode_attention_fwd 時把 kv_c_and_k_pe_cache 當成「單 head 的 K cache」傳進去:
# 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 不是寫死的:它按當前 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_mqa在is_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。