Skip to content

FlashAttention と FlashInfer:2 大 GPU アテンションバックエンド

源码版本v0.25.1

役割

NVIDIA GPU で最もよく走るのは 2 つのバックエンドです:FlashAttentionBackend (flash_attn.py:67) は Dao Labs の flash_attn_varlen_func をラップし、FA2 から FA3、FA4 まで自适应します。FlashInferBackend (flashinfer.py:330) は FlashInfer の prefill / decode kernel をラップし、Hopper/Blackwell で GQA、FP8 KV cache、page size 128+ をネイティブにサポートします。両者とも AttentionBackend を継承し、同じ抽象を実装しますが、内部の KV cache 形状、metadata フィールド、cudagraph の経路は異なります。

FlashAttentionMetadataBuilder.build (flash_attn.py:428) は CommonAttentionMetadata を prefill / decode の 2 つのセグメントに切り、それぞれ query_start_locseq_lensblock_table を埋め、ついでに AOT scheduling も計算します。FlashInferMetadataBuilder.build (flashinfer.py:1057) はさらに split_decodes_and_prefills を行い、batch を「先に decode、次に prefill」のレイアウトに並べ替えます。FlashInfer の decode kernel は batch の前段がすべて同長の decode query であることを前提とするためです。

FlashAttentionImpl.forward (flash_attn.py:749) と FlashInferImpl.forward (flashinfer.py:1599) が本当に kernel を呼ぶ場所です。これらは Q/K/V、kv_cacheattn_metadata を受け取り、まず kernel が必要な [num_tokens, num_heads, head_size] に reshape し、対応する flash_attn_varlen_func または FlashInfer の BatchPrefillWithRaggedKVCacheWrapper / BatchDecodeWithPagedKVCacheWrapper を呼びます。forward_includes_kv_cache_update は FA では False (FA kernel 自身が KV cache を書く)、FlashInfer も同様ですが、KV cache の書き込み経路は unified_kv_cache_update op を経由します。

設計動機

  • FA バージョンの自适应:get_flash_attn_version (flash_attn.py:714) が head_size、alibi、プラットフォームの演算力に基づいて FA2/3/4 から選びます。FA3 でのみ sink、FP8 KV cache、per-head quant scales をサポートします。
  • block size は 16 アラインメント:FA の get_supported_kernel_block_sizesMultipleOf(16) を返し (flash_attn.py:76-77)、kernel tile が細切れにならないようにします。
  • NHD / HND の 2 種 layout:get_kv_cache_stride_order (flash_attn.py:134-153) が KVCacheLayoutType に応じて次元の並びを決めます。FA3 は NHD を好み、FlashInfer は kernel ごとに異なる layout を好みます。
  • FlashInfer の大 page サポート:get_supported_kernel_block_sizes は Blackwell + GQA のとき [16, 32, 64, 128, 256, 512, 1024] を返し (flashinfer.py:343-362)、大規模モデルがより大きな KV page を使って block table 参照を減らせるようにします。
  • cudagraph の段階分け:FA builder はデフォルトで _cudagraph_support = UNIFORM_SINGLE_TOKEN_DECODE、FlashInfer は GQA + uniform batch で ALWAYS まで行け、spec-decode も cudagraph に乗せられます。
  • cascade attention で prefix 共有:FA は use_cascade_attention (flash_attn.py:670-671) で複数リクエストが同じ KV prefix を共有します。FlashInfer は独自の cascade 経路を使い、どちらもバッチ内 prefix を抽出して 1 回だけ計算します。

主要ファイル

データフロー

2 つのバックエンドの forward が受け取るテンサルの形状は同じで、Q は [num_tokens, num_heads, head_size]、KV cache はバックエンド自身の get_kv_cache_shape で決まります。以下は FA の KV cache 形状と layout の選択です:

python
# vllm/v1/attention/backends/flash_attn.py L122-L132
@staticmethod
def get_kv_cache_shape(
    num_blocks: int,
    block_size: int,
    num_kv_heads: int,
    head_size: int,
    cache_dtype_str: str = "auto",
) -> tuple[int, ...]:
    if block_size % 16 != 0:
        raise ValueError("Block size must be a multiple of 16.")
    return (num_blocks, 2, block_size, num_kv_heads, head_size)

FlashInfer は Blackwell + GQA のときにより大きな page size を晒し、KV block 数を減らして block table 参照を節約します:

python
# vllm/v1/attention/backends/flashinfer.py L342-L362
@staticmethod
def get_supported_kernel_block_sizes() -> list[int | MultipleOf]:
    # Page sizes >= 128 only run on the trtllm-gen dynamic kernel (GQA/MQA
    # on Blackwell); advertise them only when usable so selection never
    # picks a large kernel block we cannot serve.
    use_large_pages = False
    vllm_config = get_current_vllm_config_or_none()
    if vllm_config is not None and vllm_config.model_config is not None:
        pc = vllm_config.parallel_config
        mc = vllm_config.model_config
        num_qo_heads = mc.get_num_attention_heads(pc)
        num_kv_heads = mc.get_num_kv_heads(pc)
        use_large_pages = (
            num_kv_heads > 0
            and num_qo_heads // num_kv_heads > 1
            and current_platform.is_device_capability_family(100)
            and can_use_trtllm_attention(num_qo_heads, num_kv_heads)
        )
    if not use_large_pages:
        return [16, 32, 64]
    return [16, 32, 64, 128, 256, 512, 1024]

FlashInferMetadataBuilder.buildCommonAttentionMetadata を受け取ったあとにまず batch の並べ替えを行い、decode セグメントを batch の前に移動します。これにより decode kernel が受け取るのは連続した同長 query になり、varlen のインデックスが不要です。FA は並べ替えを行わず、直接 flash_attn_varlen_func に渡します。varlen はもともと任意の長さのセグメントをサポートします。

境界と失敗

  • FA4 でのみ 512 head_size をサポート:supports_head_size (flash_attn.py:155-163) は head_size > 256 のとき is_fa_version_supported(4) でのみ True を返します。
  • FP8 KV cache の制限:FA は FA3 + Hopper (is_device_capability_family(90)) のときのみ (flash_attn.py:166-176) fp8/fp8_e4m3 を受け付け、それ以外の組合せは別のバックエンドにフォールバックします。
  • block_size は 16 の倍数必須:FA と FI の get_supported_kernel_block_sizes はどちらも MultipleOf(16) を要求し、16 の倍数でない block size は selector にそのバックエンドをスキップさせます。
  • FlashInfer は非因果 query-query アテンションを走らせない:build の中で causal=False のときは直接 native prefill に進みます (flashinfer.py:1074-1080)。FlashInfer decode/TRTLLM 経路は双方向アテンションを表現しないためです。
  • FlashInfer は post-top-k logits を返さない:TopKTopPSampler は logprobs_mode が processed_logits / processed_logprobs のとき forward_cuda を使いません (topk_topp_sampler.py:86-95)。FlashInfer sampler はフィルタ後の logits を晒さないためです。
  • DCP combine の分岐:FlashAttentionImpl.__init__dcp_comm_backend に応じて cp_lse_ag_out_rs または dcp_a2a_lse_reduce を選び (flash_attn.py:738-743)、間違えると断片間 softmax のマージが狂います。

まとめ

FlashAttention は vLLM が NVIDIA GPU で使うデフォルトバックエンドで、dtype / head_size の組合せを最も広くカバーします。FlashInfer は Hopper/Blackwell で GQA、FP8 KV cache、大 page をよりネイティブにサポートし、decode batch の性能が安定します。どちらも AttentionBackend 抽象を通じて上位に晒され、selector がモデルとハードウェアに応じて 1 つを選びます。その他のバックエンドは /attention/backend (抽象層) と /attention/triton-mla (Triton / MLA) を参照してください。サンプリング層との接続は /sampling/sampler を参照してください。

公式資料: vLLM 文档 · README