Triton y MLA: atención portable y atención en espacio latente
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_dtypesincluyefloat32(triton_attn.py:271-287), ysupported_kv_cache_dtypesincluyeint4_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_causaldevuelve 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_attentiondevuelve explícitamenteFalse(triton_attn.py:368-370), solo va por la ruta varlen ingenua. - MLA con un único canal KV:
MLACommonBackend.get_kv_cache_shapedevuelve(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_sizessolo admite[320, 576](mla_attention.py:1236-1238), que se corresponden con las combinacioneskv_lora_rank + qk_rope_head_dimde 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_ees True por defecto, y el kernel de combinación entre shards va porcp_lse_ag_out_rs. - Competencia entre backends:
AttentionBackendEnumenumeraTRITON_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
TritonAttentionBackend:271— clase del backend Triton; soporta fp32 y KV cache cuantizado en varios formatos per_token_head.Triton get_kv_cache_shape:317-345— en modo per_token_head rellena head_dim para alojar la escala inline.TritonAttentionImpl.forward:560— invoca el kernel Tritonpaged_attentiony rechaza explícitamenteoutput_block_scale.MLAAttention:339— entrada de la capa de modelo; contienekv_b_proj,qk_nope_head_dim,qk_rope_head_dim,kv_lora_rank.MLACommonBackend:1206— clase base del backend MLA; defineis_mla() = Truey la forma 3D del KV cache.MLACommonImpl:1988— clase base del impl MLA; declara los métodos abstractosforward_mqa/forward_mha.TritonMLABackend:81— backend MLA en Triton puro; corre en cualquier GPU.TritonMLAImpl.forward_mqa:189— invoca el kernel Tritondecode_attention_fwdcon la ramais_mla=True._compute_num_kv_splits:41-47— elige el número de splits KV segúnmax_seq_len / 512, con topesm_count * 2.FlashAttnMLABackend:43— MLA por la ruta de FA, conflash_attn_varlen_func+ metadata específica de MLA.
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»:
# 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, lsenum_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.forwardlanza directamenteNotImplementedErrorcuandooutput_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 lanzaNotImplementedError(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:
TritonMLAImplponesupports_quant_query_input = Falsecon 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_mqarecurre atorch.emptycuandois_workspace_manager_initialized()devuelve False(triton_mla.py:225-232); basta para pruebas unitarias, pero en producción esGPUModelRunnerquien 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.