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) はそのうちの 1 つの実装で、純 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 モードでは inline scale を収めるため head_dim を pad します。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 はアテンションの前に K/V と一緒に計算されるのではなく、q_pe (RoPE に通す) と q_nope (KV latent と直接投影) に分かれます。KV cache に保存するのは [kv_c (kv_lora_rank), k_pe (qk_rope_head_dim)] を結合した単一ベクトルで、decode 時に MQA として 1 回取得し、その後 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 に応じて動的に計算し、長いシーケンスほど複数の split に切って複数の 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) はこの 2 つだけを返し、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 未初期化時のフォールバック:
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 は展開された K/V ではなく圧縮 latent を保存します。どちらも /attention/backend の抽象の下に掛けられ、get_attn_backend が初期化時に選定します。Triton と MLA の KV cache 形状は最終的に /kv-cache/kv-cache-manager で統一管理され、サンプリングとの接続は /sampling/sampler を参照してください。