Skip to content

AttentionBackend : couche d'abstraction des backends d'attention

源码版本v0.25.1

Responsabilités

Les implémentations d'attention supportées par vLLM se comptent sur plus d'une main : FlashAttention 2/3/4, FlashInfer, Triton, MLA, Mamba, ROCm AITER, XPU… Chacune diffère sur la disposition de la cache KV, la taille de block kernel, le support cudagraph. La couche AttentionBackend existe pour absorber ces différences afin que GPUModelRunner n'appelle qu'un forward unifié, sans se soucier du kernel en cours d'exécution.

L'abstraction se scinde en deux. AttentionBackend(backend.py:56) est la fabrique (attributs de classe + méthodes statiques) qui répond aux questions « quel est mon nom, quelle classe d'impl, quel metadata builder, quelle forme pour la cache KV ». AttentionMetadataBuilder(backend.py:600) traduit, à chaque step d'ordonnancement, le CommonAttentionMetadata (métadonnées per-batch partagées entre backends, comme query_start_loc, seq_lens, block_table_tensor) en AttentionMetadata réellement consommé par le kernel du backend courant. Plus bas, AttentionImpl(backend.py:858) est l'objet qui détient réellement la méthode forward et que la couche modèle appelle à chaque couche.

Les backends sont pluggables : AttentionBackendEnum(registry.py:34) énumère les « chemins de classe » de tous les backends, et au runtime get_attn_backend(selector.py:54) choisit le bon selon head_size, dtype, kv_cache_dtype, utilisation de MLA, capacité de la plateforme, etc., puis lazy-load le module correspondant. La couche modèle n'a qu'à récupérer un type[AttentionBackend] à l'initialisation ; la suite est gérée par cette abstraction.

Motivation de conception

  • Unification multi-material : un même code modèle tourne sur NVIDIA / AMD / Intel XPU / CPU, la seule différence étant le backend choisi ; le selector concentre la décision en un point.
  • Découplage de la forme de la cache KV : la get_kv_cache_shape de chaque backend décide de la disposition physique du pool de blocks (FlashAttention : [num_blocks, 2, block_size, num_kv_heads, head_size], MLA : [num_blocks, block_size, head_size]) ; KVCacheManager ne regarde que le spec, sans se soucier du backend concret.
  • Support cudagraph gradué : AttentionCGSupport(backend.py:583) distingue quatre niveaux ALWAYS / UNIFORM_BATCH / UNIFORM_SINGLE_TOKEN_DECODE / NEVER, pour que GPUModelRunner sache si le backend peut entrer dans un cudagraph.
  • Classification de l'invariance argmax : AttentionImplBase(backend.py:769) distingue les capacités « peut renvoyer le softmax lse », « lse en ln ou en log2 », « supporte le Prefill Context Parallelism », etc. Le kernel de fusion inter-shards DCP s'appuie sur ces drapeaux pour choisir sa branche.
  • MLA passe par une voie dédiée : MLAAttentionImpl(backend.py:941) n'implémente pas forward mais forward_mqa / forward_mha, car la cache KV de MLA stocke le latent compressé et non un K/V classique.

Fichiers clés

  • AttentionBackend ABC:56 — classe abstraite, définit get_name / get_impl_cls / get_builder_cls / get_kv_cache_shape etc. en statiques.
  • quatre abstractmethod centrales:74-97get_name, get_impl_cls, get_builder_cls, get_kv_cache_shape, à implémenter par les sous-classes.
  • CommonAttentionMetadata:395 — métadonnées per-batch partagées entre backends : query_start_loc, seq_lens, block_table_tensor, etc.
  • AttentionMetadataBuilder ABC:600 — la méthode build traduit CommonAttentionMetadata dans le metadata du backend.
  • AttentionImplBase:769 — classe de base du forward réel, porte les drapeaux de capacité can_return_lse_for_decode, lse_base_on_e, supports_pcp, etc.
  • AttentionImpl:858 — ABC standard de l'impl d'attention, définit forward(layer, query, key, value, kv_cache, attn_metadata, output).
  • MLAAttentionImpl:941 — ABC dédiée à MLA, interface forward_mqa / forward_mha, sans forward standard.
  • get_attn_backend:54 — choisit un backend selon head_size/dtype/kv_cache_dtype/use_mla, avec @cache.
  • AttentionBackendEnum:34 — énumère les chemins de classe de tous les backends, valeurs par défaut vers vllm.v1.attention.backends.*.
  • register_backend:233 — surcharge runtime du pointeur de classe d'un membre de l'enum, pour les plugins.

Flux de données

À chaque step, GPUModelRunner prépare Q/K/V et CommonAttentionMetadata pour tout le batch, puis appelle AttentionImpl.forward couche par couche. Avant chaque couche, le AttentionMetadataBuilder.build du backend a déjà transformé les métadonnées communes en la forme attendue par le kernel. La signature de AttentionMetadataBuilder.build, commune à tous les backends :

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.

Le choix du backend est décidé par get_attn_backend, qui à partir du vllm_config courant + dtype/head_size du modèle récupère le chemin de classe depuis AttentionBackendEnum et lazy-load :

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,
)

À l'initialisation de chaque Attention / MLAAttention, la couche modèle récupère cette classe de backend, puis construit l'impl via backend.get_impl_cls()(...) et le metadata builder via backend.get_builder_cls()(...). Ces deux objets sont réutilisés tout au long de l'inférence.

Limites et échecs

  • Head size non supporté : AttentionBackend.supports_head_size(backend.py:159) est interrogé en premier ; FlashAttention exige head_size multiple de 8 et ≤ 256 (FA4 pousse à 512), sinon le backend est écarté.
  • Dtype / KV cache dtype non supportés : supports_dtype / supports_kv_cache_dtype(backend.py:163-173) vérifient séparément ; FlashInfer supporte fp8/nvfp4, FlashAttention n'accepte fp8 que sur Hopper + FA3.
  • Block size doit être aligné sur 16 : get_supported_kernel_block_sizes de FlashAttention / Triton / MLA renvoient tous MultipleOf(16), un block_size non aligné lève directement ValueError(flash_attn.py:130-131).
  • Le backend MLA ne réutilise pas forward : MLAAttentionImpl(backend.py:941) ne déclare que forward_mqa / forward_mha ; le MLAAttention côté modèle appelle l'un ou l'autre selon le contexte, et un forward mal utilisé lèverait AttributeError.
  • La base LSE pour DCP ne peut pas être fausse : AttentionImplBase.lse_base_on_e(backend.py:795) distingue ln vs log2 ; une erreur ici corrompt silencieusement le dénominateur softmax inter-shards.
  • Le cache du selector suppose une config immuable : _cached_get_attn_backend(selector.py:114) utilise @cache ; modifier dtype/head_size en cours d'exécution ne re-sélectionne pas de backend, il faut redémarrer le processus.

Résumé

La couche AttentionBackend est la clé de la coexistence multi-material et multi-kernel dans vLLM : la couche modèle ne connaît que l'interface AttentionImpl.forward, et la façon dont le backend dispose la cache KV, génère son metadata, ou choisit entre FlashAttention et Triton est fixée à l'initialisation par le couple selector + enum. Pour le détail de chaque backend, voir /attention/flash (FlashAttention / FlashInfer) et /attention/triton-mla (Triton / MLA). L'endroit où la couche supérieure appelle ces impls est sur /worker/gpu-model-runner.

Voir la documentation officielle : Documentation vLLM · README