Skip to content

AttentionBackend: capa de abstracción de backends de atención

源码版本v0.25.1

Responsabilidades

vLLM soporta más implementaciones de atención (attention) de las que se cuentan con una mano: FlashAttention 2/3/4, FlashInfer, Triton, MLA, Mamba, ROCm AITER, XPU… Cada una difiere en la disposición del KV cache, en el block size del kernel y en el soporte de cudagraph. La capa AttentionBackend existe para absorber esas diferencias, de modo que la capa superior GPUModelRunner solo invoque un único forward unificado, sin preocuparse por qué kernel se está ejecutando.

La abstracción se divide en dos mitades: AttentionBackend(backend.py:56) es una fábrica de atributos de clase más métodos estáticos que responde a «cómo me llamo, qué clase impl uso, qué metadata builder uso y qué forma debe tener el KV cache»; AttentionMetadataBuilder(backend.py:600) se encarga, en cada step de planificación, de traducir CommonAttentionMetadata (metadatos por batch compartidos entre backends, como query_start_loc, seq_lens, block_table_tensor) al AttentionMetadata que el kernel del backend actual realmente consume. Más abajo, AttentionImpl(backend.py:858) es el objeto que efectivamente expone el método forward y al que la capa de modelo llama una vez por capa.

Los backends son intercambiables: AttentionBackendEnum(registry.py:34) enumera las rutas de clase de todos los backends, y en tiempo de ejecución get_attn_backend(selector.py:54) elige uno en función de head_size, dtype, kv_cache_dtype, si se usa MLA, la capacidad de la plataforma, etc., y carga el módulo correspondiente de forma perezosa. La capa de modelo solo necesita obtener type[AttentionBackend] en la inicialización; el resto lo maneja esta abstracción.

Motivación de diseño

  • Unificación entre hardwares: el mismo código de modelo corre en NVIDIA / AMD / Intel XPU / CPU; la diferencia está solo en qué backend se elige, y el selector centraliza esa decisión.
  • Desacoplamiento de la forma del KV cache: el get_kv_cache_shape de cada backend decide el layout físico del pool de bloques (FlashAttention usa [num_blocks, 2, block_size, num_kv_heads, head_size], MLA usa [num_blocks, block_size, head_size]); KVCacheManager solo mira el spec y no se preocupa por el backend concreto.
  • Soporte de cudagraph por niveles: AttentionCGSupport(backend.py:583) define cuatro grados —ALWAYS / UNIFORM_BATCH / UNIFORM_SINGLE_TOKEN_DECODE / NEVER— para que GPUModelRunner sepa si el backend actual puede entrar en cudagraph.
  • Clasificación de invariantes de argmax: AttentionImplBase(backend.py:769) distingue atributos como «¿puede devolver softmax lse?», «¿lse es ln o log2?», «¿soporta Prefill Context Parallelism?»; el kernel de combinación entre shards de DCP elige la rama según estos flags.
  • MLA va por una vía separada: MLAAttentionImpl(backend.py:941) no implementa forward, sino forward_mqa / forward_mha, porque el KV cache de MLA guarda el latent comprimido, no un K/V convencional.

Archivos clave

  • AttentionBackend ABC:56 — clase abstracta base; define los métodos estáticos get_name / get_impl_cls / get_builder_cls / get_kv_cache_shape.
  • cuatro abstractmethod núcleo:74-97get_name, get_impl_cls, get_builder_cls, get_kv_cache_shape, que las subclases deben implementar.
  • CommonAttentionMetadata:395 — metadatos por batch compartidos entre backends; query_start_loc, seq_lens, block_table_tensor, etc., viven aquí.
  • AttentionMetadataBuilder ABC:600 — el método build traduce CommonAttentionMetadata al metadata propio del backend.
  • AttentionImplBase:769 — clase base del verdadero forward; expone los flags de capacidad can_return_lse_for_decode, lse_base_on_e, supports_pcp.
  • AttentionImpl:858 — ABC para la implementación estándar de atención; define forward(layer, query, key, value, kv_cache, attn_metadata, output).
  • MLAAttentionImpl:941 — ABC específica de MLA; la interfaz es forward_mqa / forward_mha, no el forward estándar.
  • get_attn_backend:54 — selecciona el backend según head_size/dtype/kv_cache_dtype/use_mla, con @cache.
  • AttentionBackendEnum:34 — enumera las rutas de clase de todos los backends; los valores por defecto apuntan a vllm.v1.attention.backends.*.
  • register_backend:233 — sobrescribe en tiempo de ejecución la clase a la que apunta un miembro del enum; se usa para plugins.

Flujo de datos

En un step, GPUModelRunner prepara primero el query/key/value de todo el lote junto con CommonAttentionMetadata, y luego llama a AttentionImpl.forward capa por capa. Antes de cada capa, AttentionMetadataBuilder.build del backend ya ha transformado los metadatos comunes a la forma que el kernel necesita. Abajo está la firma de AttentionMetadataBuilder.build, que todos los backends implementan con esta misma interfaz:

python
# vllm/v1/attention/backend.py L666-L678
@abstractmethod
def build(
    self,
    common_prefix_len: int,
    common_attn_metadata: CommonAttentionMetadata,
    fast_build: bool = False,
) -> M:
    """
    Central method that builds attention metadata.
    Some builders (MLA) require reorder_batch to be called prior to build.

    Args:
        common_prefix_len: The length of the common prefix of the batch.
        common_attn_metadata: The common attention metadata.

El enrutamiento al backend lo decide get_attn_backend, que a partir de la vllm_config actual y de los parámetros dtype/head_size del modelo recupera la ruta de clase desde AttentionBackendEnum y la carga de forma perezosa:

python
# vllm/v1/attention/selector.py L75-L110
from vllm.config import get_current_vllm_config

vllm_config = get_current_vllm_config()

cache_config = vllm_config.cache_config
if cache_config is not None and cache_config.user_specified_block_size:
    block_size = cache_config.block_size
else:
    block_size = None

kv_transfer_config = vllm_config.kv_transfer_config
use_kv_connector = (
    kv_transfer_config is not None and kv_transfer_config.is_kv_transfer_instance
)

attn_selector_config = AttentionSelectorConfig(
    head_size=head_size,
    dtype=dtype,
    kv_cache_dtype=cast(CacheDType | None, kv_cache_dtype),
    block_size=block_size,
    use_mla=use_mla,
    ...
)

return _cached_get_attn_backend(
    backend=vllm_config.attention_config.backend,
    attn_selector_config=attn_selector_config,
    num_heads=num_heads,
)

La capa de modelo obtiene esta clase de backend al inicializar cada Attention / MLAAttention, y luego construye el impl con backend.get_impl_cls()(...) y el metadata builder con backend.get_builder_cls()(...); ambos objetos se reutilizan durante toda la inferencia posterior.

Límites y fallos

  • Head sizes no soportados: AttentionBackend.supports_head_size(backend.py:159) se pregunta primero al backend; FlashAttention exige que head_size sea múltiplo de 8 y no supere 256 (FA4 llega hasta 512); si no se cumple, se cambia de backend.
  • Dtype / KV cache dtype no soportados: supports_dtype / supports_kv_cache_dtype(backend.py:163-173) los comprueban por separado; FlashInfer soporta fp8/nvfp4, mientras que FlashAttention solo acepta fp8 en Hopper + FA3.
  • El block size debe estar alineado a 16: get_supported_kernel_block_sizes de FlashAttention / Triton / MLA devuelve siempre MultipleOf(16); un block_size no alineado lanza directamente ValueError(flash_attn.py:130-131).
  • El backend MLA no reutiliza forward: MLAAttentionImpl(backend.py:941) solo declara forward_mqa / forward_mha; el MLAAttention de la capa de modelo llama al método que corresponde según el contexto, y un uso equivocado de forward lanza AttributeError directamente.
  • La base del lse en DCP no puede equivocarse: AttentionImplBase.lse_base_on_e(backend.py:795) distingue ln de log2; elegir mal provoca un fallo silencioso en el denominador del softmax entre shards.
  • La caché del selector asume configuración invariante: _cached_get_attn_backend(selector.py:114) usa @cache; cambiar dtype/head_size en caliente no re-selecciona el backend, hace falta reiniciar el proceso.

Resumen

La abstracción de AttentionBackend es la clave para que vLLM convivir con múltiples hardwares y kernels: la capa de modelo solo conoce la interfaz AttentionImpl.forward; cómo el backend dispone el KV cache, cómo genera el metadata y si va por FlashAttention o por Triton lo deciden selector + enum en la inicialización. Los detalles de cada backend están en /attention/flash (FlashAttention / FlashInfer) y /attention/triton-mla (Triton / MLA). Dónde invoca la capa superior a estos impl está en /worker/gpu-model-runner.

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