FlashAttention と FlashInfer:2 大 GPU アテンションバックエンド
役割
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_loc、seq_lens、block_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_cache、attn_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_sizesはMultipleOf(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 回だけ計算します。
主要ファイル
FlashAttentionBackend:67— バックエンドクラス、サポートする dtype / kv_cache_dtype / block size / head_size を宣言。FlashAttention get_kv_cache_shape:122-132—(num_blocks, 2, block_size, num_kv_heads, head_size)、2 は K/V のスロット分け。FlashAttentionMetadataBuilder.build:428— prefill + decode セグメント分割、AOT schedule の意思決定。FlashAttentionImpl:674— impl クラス、forward_includes_kv_cache_update = False。FlashAttentionImpl.forward:749—flash_attn_varlen_funcを呼ぶメインエントリ、Q/K/V shape は[num_tokens, num_heads, head_size]。FlashInferBackend:330— FlashInfer バックエンドクラス、supported_kv_cache_dtypesはfp8/nvfp4を含む。FlashInfer get_kv_cache_shape:380-389— dtype に応じて異なる形状を返し、nvfp4 は packed layout に。FlashInferMetadataBuilder:600—FIPrefill/FIDecodeのサブ builder を持ち、phase ごとに振り分け。FlashInferMetadataBuilder.build:1057—split_decodes_and_prefillsで batch を並べ替えてから metadata を埋めます。FlashInferImpl.forward:1599— prefill または decode wrapper を選んで FlashInfer kernel を呼びます。
データフロー
2 つのバックエンドの forward が受け取るテンサルの形状は同じで、Q は [num_tokens, num_heads, head_size]、KV cache はバックエンド自身の get_kv_cache_shape で決まります。以下は FA の KV cache 形状と layout の選択です:
# 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 参照を節約します:
# 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.build は CommonAttentionMetadata を受け取ったあとにまず 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 を参照してください。