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 であり、通常の 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_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 はアテンションの前に 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_mqadecode_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 に応じて動的に計算し、長いシーケンスほど複数の split に切って複数の SM で並列に走らせ、マージ段階では lse で LogSumExp の統合を行います。ワークスペースは workspace_manager から取得し、decode のホットパスで torch.empty するのを避けます。

境界と失敗

  • Triton は fused block_scale 出力をサポートしない:TritonAttentionImpl.forwardoutput_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_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 は展開された K/V ではなく圧縮 latent を保存します。どちらも /attention/backend の抽象の下に掛けられ、get_attn_backend が初期化時に選定します。Triton と MLA の KV cache 形状は最終的に /kv-cache/kv-cache-manager で統一管理され、サンプリングとの接続は /sampling/sampler を参照してください。

公式資料: vLLM 文档 · README