Skip to content

FlashAttention y FlashInfer: los dos grandes backends de atención en GPU

源码版本v0.25.1

Responsabilidades

Los dos backends más usados en NVIDIA GPU son: FlashAttentionBackend(flash_attn.py:67), que envuelve flash_attn_varlen_func de Dao Labs y se adapta automáticamente de FA2 a FA3 y FA4; y FlashInferBackend(flashinfer.py:330), que envuelve los kernels de prefill / decode de FlashInfer y aporta soporte nativo en Hopper/Blackwell para GQA, FP8 KV cache y page sizes 128+. Ambos heredan de AttentionBackend e implementan la misma abstracción, pero difieren en la forma interna del KV cache, los campos de metadata y las rutas de cudagraph.

FlashAttentionMetadataBuilder.build(flash_attn.py:428) divide CommonAttentionMetadata en dos tramos, prefill y decode, rellena query_start_loc, seq_lens, block_table y, de paso, calcula el scheduling AOT. FlashInferMetadataBuilder.build(flashinfer.py:1057) además debe invocar split_decodes_and_prefills y reordenar el lote a una disposición «decode primero, luego prefill», porque el kernel de decode de FlashInfer asume que el tramo delantero del batch son queries de decode de igual longitud.

FlashAttentionImpl.forward(flash_attn.py:749) y FlashInferImpl.forward(flashinfer.py:1599) son los puntos donde se invocan los kernels: toman Q/K/V, kv_cache y attn_metadata, los reshapean a [num_tokens, num_heads, head_size] y llaman a flash_attn_varlen_func o a los wrappers BatchPrefillWithRaggedKVCacheWrapper / BatchDecodeWithPagedKVCacheWrapper de FlashInfer. forward_includes_kv_cache_update en FA es False (el kernel de FA escribe el KV cache por sí mismo); en FlashInfer ocurre algo similar, pero la escritura del KV cache pasa por la op unified_kv_cache_update.

Motivación de diseño

  • Adaptive de versión FA: get_flash_attn_version(flash_attn.py:714) elige entre FA2/3/4 en función de head_size, alibi y la capacidad de la plataforma; solo FA3 soporta sink, FP8 KV cache y per-head quant scales.
  • Alineación de block size a 16: get_supported_kernel_block_sizes de FA devuelve MultipleOf(16)(flash_attn.py:76-77), para que el tile del kernel no se quede partido.
  • Dos layouts NHD / HND: get_kv_cache_stride_order(flash_attn.py:134-153) decide el orden de las dimensiones según KVCacheLayoutType; FA3 prefiere NHD y FlashInfer prefiere layouts distintos según el kernel.
  • Soporte de páginas grandes en FlashInfer: get_supported_kernel_block_sizes devuelve [16, 32, 64, 128, 256, 512, 1024] en Blackwell + GQA(flashinfer.py:343-362), de modo que los modelos grandes usen páginas KV más grandes y reduzcan las consultas al block table.
  • cudagraph por niveles: el builder de FA usa por defecto _cudagraph_support = UNIFORM_SINGLE_TOKEN_DECODE, mientras que FlashInfer puede llegar a ALWAYS con GQA + uniform batch, de modo que el spec-decode también entre en cudagraph.
  • Cascade attention para prefix compartido: FA expone use_cascade_attention(flash_attn.py:670-671) para que varias peticiones compartan el mismo KV prefix; FlashInfer va por su propia ruta cascade; ambos extraen el prefix del batch y solo lo computan una vez.

Archivos clave

Flujo de datos

Los tensores que reciben los forward de ambos backends tienen la misma forma: Q es [num_tokens, num_heads, head_size], y la forma del KV cache la decide el get_kv_cache_shape de cada backend. Abajo está la forma y la selección de layout del KV cache de 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 expone page sizes mayores en Blackwell + GQA, de modo que haya menos bloques KV y la consulta al block table salga más barata:

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]

Tras recibir CommonAttentionMetadata, FlashInferMetadataBuilder.build reordena primero el lote y desplaza el tramo de decode al frente, de modo que el kernel de decode reciba un segmento continuo de queries de igual longitud sin tener que indexar como varlen. FA no reordena, porque va directamente por flash_attn_varlen_func, que de por sí soporta tramos de cualquier longitud.

Límites y fallos

  • Solo FA4 soporta head_size 512: supports_head_size(flash_attn.py:155-163) solo devuelve True para head_size > 256 cuando is_fa_version_supported(4).
  • Restricciones de FP8 KV cache: FA solo acepta fp8/fp8_e4m3 con FA3 + Hopper (is_device_capability_family(90))(flash_attn.py:166-176); otras combinaciones caen a otro backend.
  • block_size debe ser múltiplo de 16: get_supported_kernel_block_sizes de FA y FI exigen MultipleOf(16); un block_size que no sea múltiplo de 16 hace que el selector descarte el backend.
  • FlashInfer no soporta atención query-query no causal: build con causal=False cae a un prefill nativo(flashinfer.py:1074-1080), porque la ruta decode/TRTLLM de FlashInfer no expresa atención bidireccional.
  • FlashInfer no devuelve logits post-top-k: cuando logprobs_mode es processed_logits / processed_logprobs, TopKTopPSampler no va por forward_cuda(topk_topp_sampler.py:86-95), porque el sampler de FlashInfer no expone los logits ya filtrados.
  • Rama combine de DCP: FlashAttentionImpl.__init__ elige entre cp_lse_ag_out_rs y dcp_a2a_lse_reduce según dcp_comm_backend(flash_attn.py:738-743); elegir mal provoca una combinación errónea del softmax entre shards.

Resumen

FlashAttention es el backend por defecto de vLLM en NVIDIA GPU y cubre la combinación más amplia de dtype / head_size; FlashInfer ofrece soporte más nativo en Hopper/Blackwell para GQA, FP8 KV cache y páginas grandes, y es más estable con lotes de decode. Ambos se exponen a la capa superior a través de la abstracción AttentionBackend, y el selector elige uno en función del modelo y del hardware. Otros backends en /attention/backend (capa de abstracción) y /attention/triton-mla (Triton / MLA). La conexión con la capa de muestreo está en /sampling/sampler.

Véase la documentación oficial: vLLM 文档 · README.