Skip to content

FlashAttention et FlashInfer : les deux principaux backends d'attention sur GPU

源码版本v0.25.1

Responsabilités

Sur GPU NVIDIA, deux backends dominent : FlashAttentionBackend(flash_attn.py:67) enveloppe le flash_attn_varlen_func de Dao Labs, en s'adaptant de FA2 à FA3 et FA4 ; FlashInferBackend(flashinfer.py:330) enveloppe les kernels prefill / decode de FlashInfer, avec un support natif sur Hopper/Blackwell pour GQA, FP8 KV cache et page size 128+. Tous deux héritent de AttentionBackend et implémentent la même abstraction, mais leur disposition de cache KV, leurs champs de metadata et leur chemin cudagraph diffèrent.

FlashAttentionMetadataBuilder.build(flash_attn.py:428) découpe le CommonAttentionMetadata en segments prefill et decode, remplit query_start_loc, seq_lens, block_table, et calcule au passage l'AOT scheduling. FlashInferMetadataBuilder.build(flashinfer.py:1057) fait en plus un split_decodes_and_prefills qui réarrange le batch en « decode d'abord, prefill ensuite », car le kernel decode de FlashInfer suppose que le début du batch est composé de queries decode de même longueur.

FlashAttentionImpl.forward(flash_attn.py:749) et FlashInferImpl.forward(flashinfer.py:1599) sont les points d'appel réels aux kernels : ils prennent Q/K/V, kv_cache, attn_metadata, reshape en [num_tokens, num_heads, head_size] puis appellent flash_attn_varlen_func ou les wrappers BatchPrefillWithRaggedKVCacheWrapper / BatchDecodeWithPagedKVCacheWrapper de FlashInfer. forward_includes_kv_cache_update vaut False pour FA (le kernel FA écrit lui-même la cache KV), et idem pour FlashInfer, mais le chemin d'écriture passe par l'op unified_kv_cache_update.

Motivation de conception

  • Adaptation de version FA : get_flash_attn_version(flash_attn.py:714) choisit entre FA2/3/4 selon head_size, alibi, capacité de la plateforme ; seul FA3 supporte sink, FP8 KV cache, per-head quant scales.
  • Alignement 16 de la block size : get_supported_kernel_block_sizes de FA renvoie MultipleOf(16)(flash_attn.py:76-77), pour ne pas fragmenter les tiles du kernel.
  • Deux layouts NHD / HND : get_kv_cache_stride_order(flash_attn.py:134-153) décide de l'ordre des dimensions selon KVCacheLayoutType ; FA3 préfère NHD, FlashInfer préfère différents layouts selon le kernel.
  • Support des grandes pages dans FlashInfer : get_supported_kernel_block_sizes renvoie [16, 32, 64, 128, 256, 512, 1024] sur Blackwell + GQA(flashinfer.py:343-362), pour permettre aux gros modèles d'utiliser des pages KV plus larges et donc moins de lookups dans la block table.
  • Cudagraph gradué : le builder FA a par défaut _cudagraph_support = UNIFORM_SINGLE_TOKEN_DECODE ; FlashInfer peut monter à ALWAYS en GQA + uniform batch, ce qui permet au spec-decode de profiter du cudagraph.
  • Cascade attention partageant le prefix : FA via use_cascade_attention(flash_attn.py:670-671) permet à plusieurs requêtes de partager le même prefix KV ; FlashInfer a son propre chemin cascade ; les deux extraient le prefix commun du batch pour ne le calculer qu'une fois.

Fichiers clés

Flux de données

Les deux backends reçoivent des tensors de même shape dans forward : Q en [num_tokens, num_heads, head_size], cache KV décidée par leur propre get_kv_cache_shape. La forme et le choix de layout de la cache KV côté FA :

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 expose des page sizes plus larges sur Blackwell + GQA, ce qui réduit le nombre de blocks KV et les lookups dans la 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.build, après avoir reçu le CommonAttentionMetadata, réarrange le batch en plaçant le segment decode en tête, de sorte que le kernel decode reçoive une plage contiguë de queries de même longueur, sans indexation varlen. FA ne réarrange pas, car il passe directement par flash_attn_varlen_func qui gère nativement des segments de longueur arbitraire.

Limites et échecs

  • FA4 seul supporte head_size 512 : supports_head_size(flash_attn.py:155-163) ne renvoie True pour head_size > 256 que si is_fa_version_supported(4).
  • Restrictions FP8 KV cache : FA n'accepte fp8/fp8_e4m3 qu'avec FA3 + Hopper (is_device_capability_family(90))(flash_attn.py:166-176) ; les autres combinaisons retombent sur un autre backend.
  • block_size doit être multiple de 16 : get_supported_kernel_block_sizes de FA et FI exigent MultipleOf(16) ; un block_size non multiple fait sauter le backend par le selector.
  • FlashInfer ne gère pas l'attention query-query non causale : dans build, quand causal=False, on tombe sur un prefill native(flashinfer.py:1074-1080), car les chemins decode/TRTLLM de FlashInfer n'expriment pas l'attention bidirectionnelle.
  • FlashInfer ne renvoie pas les logits post-top-k : TopKTopPSampler en logprobs_mode processed_logits / processed_logprobs n'emprunte pas forward_cuda(topk_topp_sampler.py:86-95), car le sampler FlashInfer n'expose pas les logits filtrés.
  • Branche de combine DCP : FlashAttentionImpl.__init__ choisit cp_lse_ag_out_rs ou dcp_a2a_lse_reduce selon dcp_comm_backend(flash_attn.py:738-743) ; une erreur ici corrompt la fusion softmax inter-shards.

Résumé

FlashAttention est le backend par défaut de vLLM sur GPU NVIDIA et couvre le plus large éventail de combinaisons dtype / head_size ; FlashInfer a un support plus natif sur Hopper/Blackwell pour GQA, FP8 KV cache et grandes pages, et tient mieux en performance sur des batches decode. Tous deux sont exposés via l'abstraction AttentionBackend, et le selector choisit selon le modèle et le matériel. Les autres backends sont sur /attention/backend (abstraction) et /attention/triton-mla (Triton / MLA). La connexion avec la couche de sampling est sur /sampling/sampler.

Voir la documentation officielle : Documentation vLLM · README