Skip to content

Triton y MLA: atención portable y atención en espacio latente

源码版本v0.25.1

Responsabilidades

TritonAttentionBackend(triton_attn.py:271) es un kernel de atención en Python puro escrito con Triton por vLLM, sin depender del paquete flash_attn de Dao Labs ni de FlashInfer. Su razón de ser es servir de respaldo (fallback): cuando la GPU no tiene suficiente capacidad, head_size cae fuera del rango soportado por FA, o la plataforma es CPU/XPU, el selector aún puede devolver una implementación que funcione. Soporta fp16/bf16/fp32, varias cuantizaciones de KV cache en FP8/INT4, y layouts con escala inline por token por head del tipo fp8_per_token_head / int8_per_token_head.

MLA (atención en espacio latente / Multi-head Latent Attention, series DeepSeek) es otra vía. No se trata de «cambiar kernel», sino de «cambiar la semánttica del KV cache»: K/V no se expanden, solo se almacena el latent comprimido kv_c (típicamente kv_lora_rank=512) junto con un pequeño tramo de RoPE k_pe. Por eso la clase base del backend MLA no es AttentionBackend, sino MLACommonBackend(mla_attention.py:1206); la clase base del impl no es AttentionImpl, sino MLAAttentionImpl(backend.py:941), con interfaz forward_mqa / forward_mha y sin implementar el forward ordinario. TritonMLABackend(triton_mla.py:81) es una de sus implementaciones: escribe el kernel de decode en Triton puro (decode_attention_fwd) y corre en cualquier hardware; además existen otros backends MLA —FlashAttn / FlashInfer / FlashMLA / Cutlass / Tokenspeed / Aiter (ROCm)—, todos bajo vllm/v1/attention/backends/mla/.

Motivación de diseño

  • Triton como respaldo de amplio alcance: TritonAttentionBackend.supported_dtypes incluye float32(triton_attn.py:271-287), y supported_kv_cache_dtypes incluye int4_per_token_head / int8_per_token_head / fp8_per_token_head; las combinaciones periféricas que FA/FI no soportan aquí corren.
  • Soporta atención no causal: supports_non_causal devuelve True(triton_attn.py:301-303), de modo que Prefix LM / ViT, que requieren atención bidireccional, pueden ir por Triton.
  • No activa cascade: TritonAttentionBackend.use_cascade_attention devuelve explícitamente False(triton_attn.py:368-370), solo va por la ruta varlen ingenua.
  • MLA con un único canal KV: MLACommonBackend.get_kv_cache_shape devuelve (num_blocks, block_size, head_size)(mla_attention.py:1215-1223), sin slot separado para K/V, porque lo que se almacena es el vector conjugado latent + RoPE, recuperado por MQA en decode.
  • MLA limita el head_size: get_supported_head_sizes solo admite [320, 576](mla_attention.py:1236-1238), que se corresponden con las combinaciones kv_lora_rank + qk_rope_head_dim de DeepSeek-V2/V3; cualquier otro head_size se rechaza directamente.
  • lse obligatorio: TritonMLAImpl.can_return_lse_for_decode = True(triton_mla.py:134-135), porque la combinación entre shards de DCP necesita el LSE; lse_base_on_e es True por defecto, y el kernel de combinación entre shards va por cp_lse_ag_out_rs.
  • Competencia entre backends: AttentionBackendEnum enumera TRITON_MLA / FLASH_ATTN_MLA / FLASHINFER_MLA / FLASHMLA / CUTLASS_MLA / TOKENSPEED_MLA(registry.py:80-82), y el selector elige el mejor según la capacidad de la GPU.

Archivos clave

Flujo de datos

La particularidad de MLA: Q no se computa junto con K/V antes de la atención, sino que se divide en q_pe (que pasa por RoPE) y q_nope (que se proyecta directamente contra el latent de KV). El KV cache almacena un único vector [kv_c (kv_lora_rank), k_pe (qk_rope_head_dim)]; en decode se recupera por MQA y luego se hace o_proj en la dimensión de head. TritonMLAImpl.forward_mqa pasa kv_c_and_k_pe_cache a decode_attention_fwd como si fuera «un K cache de un solo head»:

python
# vllm/v1/attention/backends/mla/triton_mla.py L189-L256
def forward_mqa(
    self,
    q: torch.Tensor | tuple[torch.Tensor, torch.Tensor],
    kv_c_and_k_pe_cache: torch.Tensor,
    attn_metadata: MLACommonMetadata,
    layer: AttentionLayer,
) -> tuple[torch.Tensor, torch.Tensor | None]:
    assert kv_c_and_k_pe_cache.numel() > 0
    assert attn_metadata.decode is not None

    if type(q) is tuple:
        q = torch.cat(q, dim=-1)

    assert isinstance(q, torch.Tensor)
    B = q.shape[0]
    q_num_heads = q.shape[1]
    o = torch.zeros(
        B, q_num_heads, self.kv_lora_rank, dtype=q.dtype, device=q.device
    )
    lse = torch.zeros(B, q_num_heads, dtype=q.dtype, device=q.device)

    # For batch invariance, use only 1 split to ensure deterministic reduction
    if envs.VLLM_BATCH_INVARIANT:
        num_kv_splits = 1
    else:
        num_kv_splits = _compute_num_kv_splits(
            attn_metadata.max_seq_len, self._sm_count
        )

    # NOTE: the +1 stores the LogSumExp (LSE) that the stage2 kernel uses
    # to merge partial attention outputs across splits.
    logits_shape = (B, q_num_heads, num_kv_splits, self.kv_lora_rank + 1)
    if is_workspace_manager_initialized():
        (attn_logits,) = current_workspace_manager().get_simultaneous(
            (logits_shape, torch.float32),
        )
    else:
        attn_logits = torch.empty(
            logits_shape, dtype=torch.float32, device=q.device
        )

    # Add a head dim of 1
    kv_c_and_k_pe_cache = kv_c_and_k_pe_cache.unsqueeze(2)
    kv_c_cache = kv_c_and_k_pe_cache[..., : self.kv_lora_rank]
    PAGE_SIZE = kv_c_and_k_pe_cache.size(1)

    # Run MQA — always pass layer scales. When KV cache is
    # BF16 the kernel's `if dtype.is_fp8()` check is a no-op.
    decode_attention_fwd(
        q,
        kv_c_and_k_pe_cache,
        kv_c_cache,
        o,
        lse,
        attn_metadata.decode.block_table,
        attn_metadata.decode.seq_lens,
        attn_logits,
        num_kv_splits,
        self.scale,
        PAGE_SIZE,
        k_scale=layer._k_scale,
        v_scale=layer._k_scale,
        is_mla=True,
    )

    return o, lse

num_kv_splits no está fijo: se calcula dinámicamente según el seq_len máximo del batch actual; las secuencias largas se parten en más trozos para que los SM paralelicen, y la fase de combinación usa lse para hacer el merge por LogSumExp. La zona de trabajo se toma del workspace_manager para evitar torch.empty en el hot path de decode.

Límites y fallos

  • Triton no soporta salida block_scale fusionada: TritonAttentionImpl.forward lanza directamente NotImplementedError cuando output_block_scale is not None(triton_attn.py:584-588).
  • MLA no soporta alibi/sliding_window/logits_soft_cap: TritonMLAImpl.__init__ lo comprueba explícitamente y lanza NotImplementedError(triton_mla.py:166-170).
  • head_size de MLA limitado a 320 o 576: MLACommonBackend.get_supported_head_sizes(mla_attention.py:1236-1238) solo devuelve esos dos; cualquier modelo ajeno a DeepSeek-V2/V3 queda fuera.
  • FP8 KV cache va con query BF16: TritonMLAImpl pone supports_quant_query_input = False con KV cache FP8(triton_mla.py:183-185), porque el kernel de Triton hace el dequant internamente y no se puede dejar que la capa superior comprima también Q a FP8.
  • Fallback cuando workspace_manager no está inicializado: forward_mqa recurre a torch.empty cuando is_workspace_manager_initialized() devuelve False(triton_mla.py:225-232); basta para pruebas unitarias, pero en producción es GPUModelRunner quien debe preasignar el workspace.
  • num_kv_splits afecta a la corrección del merge LSE: el cambio del número de splits cambia la shape de attn_logits, y el kernel de merge del stage2 debe usar el mismo valor, por lo que _compute_num_kv_splits(triton_mla.py:41) es la misma función en el momento en que el builder reserva el workspace y en el del forward.

Resumen

Triton es el backend de respaldo de vLLM cuando no hay FA/FI, y cubre fp32 y KV cache cuantizado en combinaciones periféricas; MLA es la subabstracción específica para la atención en espacio latente de la serie DeepSeek, donde el KV cache guarda el latent comprimido en vez de K/V expandidos. Ambos cuelgan de la misma capa de abstracción, /attention/backend, y son seleccionados por get_attn_backend en la inicialización. La forma final del KV cache de Triton y MLA la gestiona de forma unificada /kv-cache/kv-cache-manager; la conexión con el muestreo está en /sampling/sampler.

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