FlashAttention und FlashInfer: die beiden großen GPU-Attention-Backends
Verantwortung
Auf NVIDIA-GPUs laufen am häufigsten zwei Backends: FlashAttentionBackend(flash_attn.py:67) kapselt flash_attn_varlen_func von Dao Labs und passt sich von FA2 über FA3 bis FA4 automatisch an; FlashInferBackend(flashinfer.py:330) kapselt FlashInfers Prefill-/Decode-Kernel und bietet auf Hopper/Blackwell native Unterstützung für GQA, FP8-KV-Cache (KV cache) und page size 128+. Beide erben von AttentionBackend und implementieren dieselbe Abstraktion, aber ihre internen KV-Cache-Formen, Metadata-Felder und cudagraph-Pfade unterscheiden sich.
FlashAttentionMetadataBuilder.build(flash_attn.py:428) zerlegt CommonAttentionMetadata in einen Prefill- und einen Decode-Abschnitt, füllt jeweils query_start_loc, seq_lens, block_table und berechnet nebenbei AOT-Scheduling. FlashInferMetadataBuilder.build(flashinfer.py:1057) muss zusätzlich split_decodes_and_prefills ausführen und den Batch in ein «Decode zuerst, dann Prefill»-Layout umsortieren, weil FlashInfers Decode-Kernel erwartet, dass der vordere Teil des Batches ausschließlich gleichlange Decode-Queries enthält.
FlashAttentionImpl.forward(flash_attn.py:749) und FlashInferImpl.forward(flashinfer.py:1599) sind die Stellen, die den Kernel tatsächlich aufrufen: Sie nehmen Q/K/V, kv_cache und attn_metadata, reshappen auf die vom Kernel geforderte Form [num_tokens, num_heads, head_size] und rufen dann flash_attn_varlen_func bzw. FlashInfers BatchPrefillWithRaggedKVCacheWrapper / BatchDecodeWithPagedKVCacheWrapper auf. forward_includes_kv_cache_update ist bei FA False (der FA-Kernel schreibt den KV-Cache (KV cache) selbst), bei FlashInfer analog, aber der Schreibpfad für den KV-Cache läuft über den unified_kv_cache_update-Op.
Entwurfsmotivation
- FA-Versions-Autoanpassung:
get_flash_attn_version(flash_attn.py:714) wählt basierend auf head_size, alibi und Plattform-Rechenfähigkeit zwischen FA2/3/4 aus; erst FA3 unterstützt sink, FP8-KV-Cache und per-head Quant-Scales. - Blockgröße 16-ausgerichtet: FAs
get_supported_kernel_block_sizesliefertMultipleOf(16)(flash_attn.py:76-77), damit der Kernel-Tile nicht zerschnitten wird. - NHD / HND zwei Layouts:
get_kv_cache_stride_order(flash_attn.py:134-153) entscheidet über die Dimensionsreihenfolge anhand vonKVCacheLayoutType; FA3 bevorzugt NHD, FlashInfer je nach Kernel unterschiedliches Layout. - FlashInfer-Large-Page-Support:
get_supported_kernel_block_sizesliefert auf Blackwell + GQA[16, 32, 64, 128, 256, 512, 1024](flashinfer.py:343-362), damit große Modelle größere KV-Seiten nutzen und Block-Table-Abfragen seltener werden. - cudagraph abgestuft: FA-Builder nutzt standardmäßig
_cudagraph_support = UNIFORM_SINGLE_TOKEN_DECODE, FlashInfer kann bei GQA + uniform Batch auf ALWAYS kommen, sodass auch Spec-Decode über cudagraph läuft. - Cascade-Attention teilt Prefix: FA nutzt
use_cascade_attention(flash_attn.py:670-671), damit mehrere Anfragen denselben KV-Prefix (prefix) teilen; FlashInfer geht seinen eigenen Cascade-Pfad; beide extrahieren den Prefix im Batch und berechnen ihn nur einmal.
Schlüsseldateien
FlashAttentionBackend:67— Backend-Klasse; deklariert unterstützte 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), die 2 ist die K/V-Slot-Aufteilung.FlashAttentionMetadataBuilder.build:428— Prefill- + Decode-Segmentierung, AOT-Schedule-Entscheidung.FlashAttentionImpl:674— impl-Klasse,forward_includes_kv_cache_update = False.FlashAttentionImpl.forward:749— Haupteinstieg fürflash_attn_varlen_func, Q/K/V-Shape[num_tokens, num_heads, head_size].FlashInferBackend:330— FlashInfer-Backend-Klasse;supported_kv_cache_dtypesenthältfp8/nvfp4.FlashInfer get_kv_cache_shape:380-389— Liefert je nach dtype unterschiedliche Form; nvfp4 nutzt packed-Layout.FlashInferMetadataBuilder:600— HältFIPrefill/FIDecodeals Sub-Builder, verzweigt je nach Phase.FlashInferMetadataBuilder.build:1057—split_decodes_and_prefillssortiert den Batch um, danach werden Metadaten gefüllt.FlashInferImpl.forward:1599— Wählt Prefill- oder Decode-Wrapper für den FlashInfer-Kernel.
Datenfluss
Beide Backends erhalten im forward dieselbe Tensor-Form: Q ist [num_tokens, num_heads, head_size], der KV-Cache (KV cache) wird vom backend-eigenen get_kv_cache_shape bestimmt. Im Folgenden FAs KV-Cache-Form und Layout-Auswahl:
# 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 exponiert auf Blackwell + GQA größere Page-Größen, sodass die Anzahl der KV-Blöcke sinkt und Block-Table-Abfragen billiger werden:
# 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 sortiert den Batch nach dem Empfang von CommonAttentionMetadata zuerst um und schiebt den Decode-Abschnitt nach vorn, sodass der Decode-Kernel einen durchgehenden Block gleich langer Queries erhält und keine Varlen-Indizierung braucht. FA sortiert nicht um, weil es direkt flash_attn_varlen_func aufruft und Varlen ohnehin beliebige Längen unterstützt.
Grenzen und Fehler
- Nur FA4 unterstützt 512 head_size:
supports_head_size(flash_attn.py:155-163) liefert beihead_size > 256nur dann True, wennis_fa_version_supported(4)ist. - FP8-KV-Cache-Einschränkungen: FA akzeptiert
fp8/fp8_e4m3nur mit FA3 + Hopper (is_device_capability_family(90))(flash_attn.py:166-176), andere Kombinationen fallen auf andere Backends zurück. - block_size muss Vielfaches von 16 sein:
get_supported_kernel_block_sizesvon FA und FI verlangenMultipleOf(16); block_size, das keine 16er-Vielfache ist, führt dazu, dass der Selector (selector) das Backend überspringt. - FlashInfer kann keine nicht-kausale Query-Query-Attention: Wenn in
buildcausal=False ist, wird direkt der native Prefill-Pfad gewählt(flashinfer.py:1074-1080), weil FlashInfers Decode-/TRTLLM-Pfade bidirektionale Attention nicht ausdrücken. - FlashInfer liefert keine post-top-k Logits:
TopKTopPSamplernutzt bei logprobs_modeprocessed_logits/processed_logprobsnichtforward_cuda(topk_topp_sampler.py:86-95), weil der FlashInfer-Sampler (sampler) die gefilterten Logits nicht exponiert. - DCP-Combine-Zweig:
FlashAttentionImpl.__init__wählt basierend aufdcp_comm_backendentwedercp_lse_ag_out_rsoderdcp_a2a_lse_reduce(flash_attn.py:738-743); eine falsche Wahl macht den schardübergreifenden Softmax-Merge falsch.
Zusammenfassung
FlashAttention ist vLLMs Standard-Backend auf NVIDIA-GPUs und deckt die breiteste Kombination aus dtype / head_size ab; FlashInfer bietet auf Hopper/Blackwell nativeren Support für GQA, FP8-KV-Cache (KV cache) und große Pages und hält Decode-Batches stabiler. Beide werden über die AttentionBackend-Abstraktion nach oben hin freigegeben, und der Selector (selector) wählt basierend auf Modell und Hardware eines aus. Andere Backends siehe /attention/backend (Abstraktionsschicht) und /attention/triton-mla (Triton / MLA). Übergang zur Sampling-Schicht siehe /sampling/sampler.