FlashAttention y FlashInfer: los dos grandes backends de atención en GPU
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_sizesde FA devuelveMultipleOf(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únKVCacheLayoutType; FA3 prefiere NHD y FlashInfer prefiere layouts distintos según el kernel. - Soporte de páginas grandes en FlashInfer:
get_supported_kernel_block_sizesdevuelve[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
FlashAttentionBackend:67— clase del backend; declara los dtype / kv_cache_dtype / block size / head_size soportados.FlashAttention get_kv_cache_shape:122-132—(num_blocks, 2, block_size, num_kv_heads, head_size); el 2 es el slot de K/V.FlashAttentionMetadataBuilder.build:428— partición prefill + decode, decisión de schedule AOT.FlashAttentionImpl:674— clase impl, conforward_includes_kv_cache_update = False.FlashAttentionImpl.forward:749— entrada principal que invocaflash_attn_varlen_func; Q/K/V con shape[num_tokens, num_heads, head_size].FlashInferBackend:330— clase del backend FlashInfer;supported_kv_cache_dtypesincluyefp8/nvfp4.FlashInfer get_kv_cache_shape:380-389— devuelve formas distintas según el dtype; nvfp4 va por un layout packed.FlashInferMetadataBuilder:600— mantiene los sub-buildersFIPrefill/FIDecodey despacha por fase.FlashInferMetadataBuilder.build:1057— trassplit_decodes_and_prefillsreordena el batch y rellena el metadata.FlashInferImpl.forward:1599— elige prefill o decode wrapper para invocar el kernel de FlashInfer.
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:
# 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:
# 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 parahead_size > 256cuandois_fa_version_supported(4). - Restricciones de FP8 KV cache: FA solo acepta
fp8/fp8_e4m3con 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_sizesde FA y FI exigenMultipleOf(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:
buildcon 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,TopKTopPSamplerno va porforward_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 entrecp_lse_ag_out_rsydcp_a2a_lse_reducesegúndcp_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.