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。