Más allá de Transformers: qué viene después de la atención
El coste cuadrático de la atención en Transformers es un cuello de botella real. Atención lineal, residuales e híbridos marcan el futuro.

La arquitectura Transformer lleva ocho años impulsando la IA. Todos los modelos de lenguaje importantes, la mayoría de los sistemas de generación de imágenes y un número creciente de modelos de audio y vídeo se construyen sobre el mecanismo de self-attention presentado en el artículo 'Attention Is All You Need'. Pero la self-attention tiene un problema de fondo: su coste de cómputo y memoria crece cuadráticamente con la longitud de la secuencia. Duplicas la entrada y el coste se cuadruplica.
Para secuencias cortas no importa. Pero para las ventanas de contexto de 128K tokens hacia las que avanzamos, y para los millones de tokens que la gente quiere, es un cuello de botella serio. Una oleada de investigación explora alternativas: residuales de atención que reutilizan cálculos entre capas, variantes de atención lineal que eliminan el coste cuadrático y arquitecturas híbridas que mezclan atención con mecanismos más baratos. El Transformer no desaparece, pero se está rediseñando.
Por qué la self-attention es cara
Para entender las alternativas hay que entender qué calcula realmente la self-attention. Dada una secuencia de N tokens, la self-attention calcula una puntuación de relevancia entre cada par de tokens. Token 1 frente al token 2, token 1 frente al token 3, ..., token 1 frente al token N, luego el token 2 frente a todos los demás, y así sucesivamente. Eso son N² pares.
import torch
import torch.nn.functional as F
def self_attention(Q, K, V):
"""
Standard self-attention.
Q, K, V: (batch, seq_len, d_model)
The attention matrix is seq_len × seq_len.
For seq_len = 1024: ~1M entries (manageable)
For seq_len = 32768: ~1B entries (expensive)
For seq_len = 131072: ~17B entries (very expensive)
"""
d_k = Q.size(-1)
# This matmul creates the N×N attention matrix
scores = torch.matmul(Q, K.transpose(-2, -1)) / (d_k ** 0.5)
weights = F.softmax(scores, dim=-1)
return torch.matmul(weights, V)
Con 4K tokens, la matriz de atención tiene 16 millones de entradas, nada que no puedan manejar las GPU modernas. Con 128K tokens, tiene 16 mil millones de entradas. Con 1 millón de tokens, supera el billón. Incluso con Flash Attention (que no reduce el cómputo, pero mejora mucho los patrones de acceso a memoria), el escalado cuadrático acaba imponiéndose.
Por eso los primeros Transformers se limitaban a 512 o 1024 tokens. Cada generación de hardware y optimización ha subido ese techo, pero seguimos chocando con un muro matemático. El escalado lineal (O(N)) sería fundamentalmente mejor que el cuadrático (O(N²)), y eso es lo que persiguen la mayoría de las arquitecturas alternativas.
Residuales de atención: reutilizar lo que ya has calculado
Uno de los enfoques más pragmáticos para reducir el coste de la atención no la reemplaza: hace que cada capa de atención sea más barata reutilizando el cálculo de las capas anteriores.
La observación: en un Transformer profundo (digamos, de 32 capas), los patrones de atención de capas adyacentes suelen parecerse muchísimo. La capa 15 y la 16 tienden a atender a posiciones similares, con pequeños ajustes. Calcular la matriz de atención N² completa desde cero en cada capa es redundante: buena parte del trabajo ya se hizo una capa antes.
Las residuales de atención aprovechan esto calculando un patrón de atención 'residual': la diferencia entre a qué quiere atender esta capa y lo que calculó la capa anterior. Si la diferencia es pequeña (que suele serlo en las capas intermedias), el cálculo sale más barato. El patrón de atención completo es la suma del patrón de la capa anterior más el residual de la capa actual.
Es análogo a cómo funciona la compresión de vídeo: en lugar de guardar cada fotograma por separado, guardas un fotograma clave y después una serie de diferencias (residuales) respecto a ese fotograma. Las diferencias suelen ser mucho más pequeñas que el fotograma completo, así que la compresión mejora muchísimo.
En la práctica, las residuales de atención reducen el coste de cómputo de la atención entre un 30 y un 50 % en las capas intermedias de modelos profundos, con un impacto mínimo en la calidad. Las primeras y las últimas capas siguen necesitando atención completa (sus patrones son más distintivos), pero las capas intermedias, que son la mayoría, consiguen aceleraciones importantes.
Atención lineal: eliminar el coste cuadrático
Las variantes de atención lineal intentan reformular la atención para que escale como O(N) en lugar de O(N²). La idea general: en vez de calcular explícitamente la matriz de atención N×N, encontrar la forma de obtener la misma salida (o aproximadamente la misma) usando operaciones lineales.
El truco matemático se basa en la descomposición en kernel de softmax. La atención estándar calcula softmax(QK^T)V. Si sustituyes softmax por una función kernel distinta que pueda descomponerse como φ(Q) · φ(K)^T, puedes reordenar el cálculo: en lugar de (φ(Q) · φ(K)^T) · V (que tiene un intermedio N×N), calculas φ(Q) · (φ(K)^T · V) (que tiene un intermedio d×d, donde d es la dimensión del modelo). Como d << N en secuencias largas, esto es muchísimo más barato.
def linear_attention(Q, K, V, feature_map=None):
"""
Linear attention via kernel feature maps.
Cost: O(N * d^2) instead of O(N^2 * d)
"""
if feature_map is None:
# ELU+1 is a common choice (from Katharopoulos et al.)
feature_map = lambda x: F.elu(x) + 1
Q = feature_map(Q) # (batch, seq_len, d)
K = feature_map(K) # (batch, seq_len, d)
# Key insight: compute K^T @ V first (d × d matrix)
# instead of Q @ K^T first (N × N matrix)
KV = torch.einsum('bnd,bnm->bdm', K, V) # (batch, d, d)
# Then multiply by Q
output = torch.einsum('bnd,bdm->bnm', Q, KV) # (batch, N, d)
# Normalize
Z = torch.einsum('bnd,bd->bn', Q, K.sum(dim=1)) # normalization
output = output / Z.unsqueeze(-1)
return output
El problema: sustituir softmax por otra función kernel cambia la distribución de la atención, y los modelos entrenados con atención softmax no siempre se adaptan bien a la atención lineal. La brecha de calidad se ha reducido notablemente (las variantes lineales recientes alcanzan entre el 95 y el 98 % de la calidad de la atención softmax), pero persiste, sobre todo en tareas que requieren recuperar información precisa a larga distancia.
Modelos de espacio de estados: otro paradigma
Los modelos de espacio de estados (SSM) como Mamba adoptan un enfoque radicalmente distinto. En lugar de calcular relaciones por pares entre tokens, procesan la secuencia mediante una recurrencia, manteniendo un estado oculto de tamaño fijo que se actualiza en cada token. Esto es O(N) por naturaleza: procesar el doble de tokens tarda el doble, no el cuádruple.
La innovación de los SSM modernos es hacer que los parámetros de la recurrencia dependan de la entrada (espacios de estado selectivos). Así, el modelo tiene una forma de atención basada en el contenido: puede 'elegir' qué información recordar y cuál olvidar, sin el coste cuadrático. Los modelos tipo Mamba igualan la calidad de los Transformers en muchos benchmarks y son notablemente más rápidos con secuencias largas.
La contrapartida: los SSM procesan los tokens de forma secuencial, lo que dificulta su paralelización durante el entrenamiento frente a los Transformers (que pueden procesar todos los tokens a la vez). La eficiencia en el entrenamiento importa: un modelo que es dos veces más rápido en inferencia pero tres veces más lento de entrenar no es necesariamente una ganancia, ya que la mayor parte del cómputo total va al entrenamiento.
Arquitecturas híbridas: el camino pragmático
La tendencia actual en los modelos de producción son las arquitecturas híbridas que combinan distintos mecanismos de atención. El razonamiento es sencillo: diferentes partes de un modelo se benefician de distintos tipos de cómputo.
- Atención completa para el razonamiento global. Algunas capas necesitan atender a toda la secuencia, encontrando contexto relevante a miles de tokens de distancia. Estas capas usan self-attention estándar (posiblemente optimizada con Flash Attention).
- Atención local para el contexto cercano. Muchas capas atienden principalmente a tokens cercanos (atención con ventana deslizante). Usar una ventana fija de 256-1024 tokens reduce el coste a O(N·W), donde W es el tamaño de la ventana.
- Atención lineal para el contexto amplio. Algunas capas necesitan agregar información de toda la secuencia, pero no requieren pesos de atención precisos. La atención lineal lo ofrece con un coste O(N).
- Capas SSM para el procesamiento secuencial. Las capas tipo Mamba pueden procesar dependencias secuenciales de forma eficiente, sin ningún cálculo de atención.
Modelos como Jamba (AI21) y diversas arquitecturas de investigación alternan estos mecanismos según el papel de cada capa. Las capas iniciales usan atención local (procesan sintaxis y patrones locales). Las intermedias usan atención lineal o SSM (construyen representaciones más amplias). Unas pocas capas estratégicas usan atención completa (razonamiento global y recuperación). Así se consigue un escalado casi lineal en conjunto sin renunciar a la calidad que requiere algo de atención completa.
Lo que deberían vigilar los desarrolladores
Si construyes aplicaciones sobre modelos de lenguaje, los cambios arquitectónicos que ocurren por debajo afectan a tu trabajo de formas concretas.
- Las ventanas de contexto seguirán creciendo. A medida que bajan los costes de atención, las ventanas se amplían. Esto cambia la arquitectura de las aplicaciones: en vez de construir pipelines RAG complejos para meter el contexto relevante en una ventana de 4K, quizá simplemente metas todo en un prompt de 1M tokens. La simplicidad es atractiva, pero la latencia y el coste varían según la arquitectura.
- Los perfiles de latencia cambian. Los Transformers tienen una latencia relativamente plana hasta cierto punto, y luego crece de forma cuadrática. Los modelos de atención lineal y SSM muestran aumentos de latencia más graduales y lineales. En aplicaciones donde el tiempo de respuesta importa, entender el comportamiento de escalado de tu modelo es clave.
- Las diferencias de calidad dependen de la tarea. Los modelos de atención lineal pueden quedarse ligeramente por detrás en tareas que requieren recuperar información precisa de posiciones concretas en contextos largos ('¿cuál era el tercer elemento de la lista de la página 47?'). En tareas de comprensión general rinden igual. Conoce tu caso de uso.
- La optimización de la inferencia importa más. A medida que los modelos se vuelven más complejos en su arquitectura (mezclando distintos tipos de atención), los motores de inferencia necesitan manejar cómputo heterogéneo de forma eficiente. vLLM, TensorRT-LLM y frameworks similares se están adaptando, pero las arquitecturas personalizadas pueden no tener soporte inmediato.
El Transformer no está siendo reemplazado, está evolucionando. La self-attention sigue siendo el mecanismo más expresivo que tenemos para modelar relaciones entre tokens. Pero no hace falta usarla en todas partes, en cada capa, con el coste N² completo. Los modelos de los próximos años usarán la atención de forma quirúrgica: máxima precisión donde más importa y alternativas más baratas en el resto. El resultado serán modelos más rápidos, que manejan contextos más largos y cuestan menos de ejecutar, igualando o superando la calidad actual. Eso merece atención.


