FlashAttention et FlashInfer : les deux principaux backends d'attention sur GPU
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_sizesde FA renvoieMultipleOf(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 selonKVCacheLayoutType; 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_sizesrenvoie[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
FlashAttentionBackend:67— classe de backend, déclare les dtype / kv_cache_dtype / block size / head_size supportés.FlashAttention get_kv_cache_shape:122-132—(num_blocks, 2, block_size, num_kv_heads, head_size), le 2 sépare K/V.FlashAttentionMetadataBuilder.build:428— découpage prefill + decode, décision AOT schedule.FlashAttentionImpl:674— classe impl,forward_includes_kv_cache_update = False.FlashAttentionImpl.forward:749— point d'entrée appelantflash_attn_varlen_func, shape Q/K/V[num_tokens, num_heads, head_size].FlashInferBackend:330— classe backend FlashInfer,supported_kv_cache_dtypesinclutfp8/nvfp4.FlashInfer get_kv_cache_shape:380-389— shape différente selon dtype, nvfp4 utilise un layout packed.FlashInferMetadataBuilder:600— détient des sous-buildersFIPrefill/FIDecode, dispatch par phase.FlashInferMetadataBuilder.build:1057—split_decodes_and_prefillsréarrange le batch puis remplit le metadata.FlashInferImpl.forward:1599— choisit prefill ou decode wrapper pour appeler le kernel FlashInfer.
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 :
# 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 :
# 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 pourhead_size > 256que siis_fa_version_supported(4). - Restrictions FP8 KV cache : FA n'accepte
fp8/fp8_e4m3qu'avec FA3 + Hopper (is_device_capability_family(90))(flash_attn.py:166-176) ; les autres combinaisons retombent sur un autre backend. block_sizedoit être multiple de 16 :get_supported_kernel_block_sizesde FA et FI exigentMultipleOf(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 :
TopKTopPSampleren logprobs_modeprocessed_logits/processed_logprobsn'emprunte pasforward_cuda(topk_topp_sampler.py:86-95), car le sampler FlashInfer n'expose pas les logits filtrés. - Branche de combine DCP :
FlashAttentionImpl.__init__choisitcp_lse_ag_out_rsoudcp_a2a_lse_reduceselondcp_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