Skip to content

FlashAttention und FlashInfer: die beiden großen GPU-Attention-Backends

源码版本v0.25.1

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_sizes liefert MultipleOf(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 von KVCacheLayoutType; FA3 bevorzugt NHD, FlashInfer je nach Kernel unterschiedliches Layout.
  • FlashInfer-Large-Page-Support: get_supported_kernel_block_sizes liefert 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

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:

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 exponiert auf Blackwell + GQA größere Page-Größen, sodass die Anzahl der KV-Blöcke sinkt und Block-Table-Abfragen billiger werden:

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 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 bei head_size > 256 nur dann True, wenn is_fa_version_supported(4) ist.
  • FP8-KV-Cache-Einschränkungen: FA akzeptiert fp8/fp8_e4m3 nur 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_sizes von FA und FI verlangen MultipleOf(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 build causal=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: TopKTopPSampler nutzt bei logprobs_mode processed_logits / processed_logprobs nicht forward_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 auf dcp_comm_backend entweder cp_lse_ag_out_rs oder dcp_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.

Siehe offizielle Dokumentation: vLLM 文档 · README.