Skip to content

AttentionBackend:アテンションバックエンド抽象レイヤー

源码版本v0.25.1

役割

vLLM がサポートするアテンション (attention) 実装は片手では足りないほどあります:FlashAttention 2/3/4、FlashInfer、Triton、MLA、Mamba、ROCm AITER、XPU……各者は KV cache のレイアウト、kernel block size、cudagraph サポートの面で異なります。AttentionBackend というレイヤーが存在する意義は、これらの差異を吸収し、上位の GPUModelRunner が統一された 1 つの forward を呼ぶだけで、いまどの kernel で走っているかを意識しなくて済むようにすることです。

抽象は 2 つに分かれます:AttentionBackend (backend.py:56) 自身はクラス属性 + 静的メソッドのファクトリで、「私は何という名前か、どの impl クラスを使うか、どの metadata builder を使うか、KV cache はどんな形状か」に答えます。AttentionMetadataBuilder (backend.py:600) は各スケジュールステップで CommonAttentionMetadata (バックエンド間で共有する per-batch メタデータ、例えば query_start_locseq_lensblock_table_tensor) を現在のバックエンド kernel が本当に必要とする AttentionMetadata に変換します。その下の AttentionImpl (backend.py:858) が実際に forward メソッドを持ち、モデルレイヤーから毎層 1 回呼ばれるオブジェクトです。

バックエンドは交換可能です:AttentionBackendEnum (registry.py:34) が全バックエンドの「クラスパス」を列挙型で並べ、実行時に get_attn_backend (selector.py:54) が head_size、dtype、kv_cache_dtype、MLA かどうか、プラットフォームの演算力などの条件で 1 つを選び、対応するモジュールを遅延ロードします。モデルレイヤーは初期化時に type[AttentionBackend] を受け取れば、後のことはこの抽象に任せます。

設計動機

  • ハードウェア横断の統一:同じモデルコードが NVIDIA / AMD / Intel XPU / CPU で動き、違いはどのバックエンドを選ぶかだけです。selector が意思決定を 1 箇所に集めます。
  • KV cache 形状の切り離し:各バックエンドの get_kv_cache_shape が block プールの物理レイアウトを決めます (FlashAttention は [num_blocks, 2, block_size, num_kv_heads, head_size]、MLA は [num_blocks, block_size, head_size])。KVCacheManager は spec だけを見て具体的なバックエンドを意識しません。
  • cudagraph サポートの段階分け:AttentionCGSupport (backend.py:583) は ALWAYS / UNIFORM_BATCH / UNIFORM_SINGLE_TOKEN_DECODE / NEVER の 4 段階に分け、GPUModelRunner が現在のバックエンドが cudagraph に入れられるかを知ります。
  • argmax 不変性の分類:AttentionImplBase (backend.py:769) は「softmax lse を返せるか」「lse が ln か log2 か」「Prefill Context Parallelism をサポートするか」などの属性を区別し、DCP の断片間マージ kernel はこれらのフラグで分岐を選びます。
  • MLA は専用経路:MLAAttentionImpl (backend.py:941) は forward を実装せず、forward_mqa / forward_mha を使います。MLA の KV cache は圧縮された latent を保存し、通常の K/V ではないためです。

主要ファイル

  • AttentionBackend ABC:56 — 抽象基底クラス、get_name / get_impl_cls / get_builder_cls / get_kv_cache_shape などの静的メソッドを定義。
  • 四个核心 abstractmethod:74-97get_nameget_impl_clsget_builder_clsget_kv_cache_shape、サブクラスは必ず実装します。
  • CommonAttentionMetadata:395 — バックエンド間共有の per-batch メタデータ、query_start_locseq_lensblock_table_tensor などはここにあります。
  • AttentionMetadataBuilder ABC:600build メソッドが CommonAttentionMetadata をバックエンド自身の metadata に変換します。
  • AttentionImplBase:769 — 真の forward の基底クラス、can_return_lse_for_decodelse_base_on_esupports_pcp などの能力フラグを持ちます。
  • AttentionImpl:858 — 標準アテンション実装 ABC、forward(layer, query, key, value, kv_cache, attn_metadata, output) を定義。
  • MLAAttentionImpl:941 — MLA 専用 ABC、インターフェースは forward_mqa / forward_mha、標準 forward は使いません。
  • get_attn_backend:54 — head_size/dtype/kv_cache_dtype/use_mla などの条件でバックエンドを選び、@cache を付けます。
  • AttentionBackendEnum:34 — 全バックエンドのクラスパスを列挙型に並べ、デフォルト値は vllm.v1.attention.backends.* を指します。
  • register_backend:233 — 実行時にある列挙メンバーが指すクラスを上書きします。プラグイン用。

データフロー

1 ステップの中で GPUModelRunner はまずバッチ全体の query/key/value と CommonAttentionMetadata を用意し、層ごとに AttentionImpl.forward を呼びます。各層の前に、バックエンドの AttentionMetadataBuilder.build が共通メタデータを kernel が必要な形状に加工しています。以下は AttentionMetadataBuilder.build のシグネチャで、全バックエンドがこの同じインターフェースを実装します:

python
# vllm/v1/attention/backend.py L666-L678
@abstractmethod
def build(
    self,
    common_prefix_len: int,
    common_attn_metadata: CommonAttentionMetadata,
    fast_build: bool = False,
) -> M:
    """
    Central method that builds attention metadata.
    Some builders (MLA) require reorder_batch to be called prior to build.

    Args:
        common_prefix_len: The length of the common prefix of the batch.
        common_attn_metadata: The common attention metadata.

バックエンドの経路選択は get_attn_backend が決めます。これは現在の vllm_config とモデル dtype/head_size などのパラメータから、AttentionBackendEnum でバックエンドクラスパスを取り出して遅延ロードします:

python
# vllm/v1/attention/selector.py L75-L110
from vllm.config import get_current_vllm_config

vllm_config = get_current_vllm_config()

cache_config = vllm_config.cache_config
if cache_config is not None and cache_config.user_specified_block_size:
    block_size = cache_config.block_size
else:
    block_size = None

kv_transfer_config = vllm_config.kv_transfer_config
use_kv_connector = (
    kv_transfer_config is not None and kv_transfer_config.is_kv_transfer_instance
)

attn_selector_config = AttentionSelectorConfig(
    head_size=head_size,
    dtype=dtype,
    kv_cache_dtype=cast(CacheDtype | None, kv_cache_dtype),
    block_size=block_size,
    use_mla=use_mla,
    ...
)

return _cached_get_attn_backend(
    backend=vllm_config.attention_config.backend,
    attn_selector_config=attn_selector_config,
    num_heads=num_heads,
)

モデルレイヤーは各 Attention / MLAAttention の初期化時にこの backend クラスを受け取り、backend.get_impl_cls()(...) で impl を構築し、backend.get_builder_cls()(...) で metadata builder を構築します。その後は推論の全ラウンドでこの 2 つのオブジェクトを使い回します。

境界と失敗

  • サポートしない head サイズ:AttentionBackend.supports_head_size (backend.py:159) がまずバックエンドに問い合わせます。FlashAttention は head_size が 8 の倍数かつ 256 以下を要求し (FA4 で 512 まで)、条件を満たさなければ別のバックエンドに切り替えます。
  • サポートしないデータ型 / KV cache dtype:supports_dtype / supports_kv_cache_dtype (backend.py:163-173) がそれぞれチェックします。FlashInfer は fp8/nvfp4 をサポートし、FlashAttention は Hopper + FA3 でのみ fp8 を受け付けます。
  • block size は 16 アラインメント必須:FlashAttention / Triton / MLA の get_supported_kernel_block_sizes はいずれも MultipleOf(16) を返し、block_size がアラインされていなければ直接 ValueError を投げます (flash_attn.py:130-131)。
  • MLA バックエンドは forward を使い回さない:MLAAttentionImpl (backend.py:941) は forward_mqa / forward_mha だけを宣言します。モデルレイヤーの MLAAttention はコンテキストに応じて対応するメソッドを呼び、forward を誤って使うと AttributeError になります。
  • DCP の lse 基数を間違えない:AttentionImplBase.lse_base_on_e (backend.py:795) は ln と log2 を区別し、間違えると断片間 softmax の分母が暗黙に狂います。
  • selector のキャッシュは設定不変を仮定:_cached_get_attn_backend (selector.py:114) は @cache を使い、実行中に dtype/head_size を変えてもバックエンドは再選択されず、プロセスの再起動が必要です。

まとめ

AttentionBackend の抽象レイヤーは vLLM がマルチハードウェア・マルチ kernel を共存させる鍵です。モデルレイヤーは AttentionImpl.forward という 1 つのインターフェースだけを認識し、バックエンドがどう KV cache を配置し、どう metadata を生成し、FlashAttention か Triton かはすべて selector + enum が初期化時に決めます。各バックエンドの具体的な実装は /attention/flash (FlashAttention / FlashInfer) と /attention/triton-mla (Triton / MLA) を参照してください。上位がどこでこれらの impl を呼ぶかは /worker/gpu-model-runner を参照してください。

公式資料: vLLM 文档 · README