lu1tr0nMoonshot AI y Alibaba reemplazaron la atención softmax por Kimi Delta Attention (KDA), una variante lineal derivada de DeltaNet que usa un estado de t
Kimi K2 y los modelos recientes de la familia Qwen3 ya no comparan cada token con todos los anteriores en cada paso de atención: usan Kimi Delta Attention (KDA), un mecanismo que guarda todo el historial relevante en un estado de tamaño fijo y lo actualiza token a token. Es la pieza de arquitectura que le permite a estos modelos procesar contextos largos sin que el costo de cómputo crezca al cuadrado con la longitud de la secuencia.
Un artículo técnico publicado en el blog de Doubleword, 'You Could Have Come Up With Kimi Delta Attention', reconstruye la derivación completa: parte de la atención softmax clásica, quita la normalización, llega a la atención lineal, después a DeltaNet, después a Gated DeltaNet y finalmente a KDA. El resultado es una cadena de decisiones de diseño que, vistas una por una, dejan de parecer magia.
El blog técnico de Doubleword publicó una derivación didáctica de Kimi Delta Attention, el mecanismo de atención que Moonshot AI usa en su familia Kimi y que, según el mismo artículo, también adoptaron modelos recientes de Qwen. En vez de presentar las ecuaciones de KDA como un bloque cerrado, el autor las reconstruye desde cero, mostrando qué problema resuelve cada término.
Esto importa porque las variantes modernas de atención lineal (RetNet, RWKV, Mamba2, DeltaNet, Gated DeltaNet, KDA) suelen presentarse con notación densa que oculta la idea central. Sin ese contexto, un lector técnico ve una ecuación con matrices diagonales y productos externos y no entiende por qué existe cada pieza.
La atención softmax calcula, para cada token, un puntaje de similitud contra todos los tokens anteriores, normaliza esos puntajes con una exponencial y usa el resultado para promediar los vectores de valor. Es preciso, pero el costo crece con el cuadrado de la longitud de la secuencia: con T tokens hay T² pares clave-consulta.
La atención lineal nace de una observación simple: si se quita la normalización softmax, el producto interno escalar entre clave y consulta se puede mover de lugar dentro de la suma. Eso permite agrupar todo lo que depende del pasado en una sola matriz de tamaño fijo, el estado S, en vez de guardar cada clave y cada valor por separado.
La identidad que hace posible el truco es esta: un producto externo |v⟩⟨k| es una matriz, y aplicarlo sobre una consulta |q⟩ da como resultado ⟨k|q⟩|v⟩, exactamente lo mismo que calcular primero el producto interno ⟨k|q⟩ y después escalar el vector v. Sumar esos productos externos token a token construye, de forma incremental, la misma cuenta que antes requería mirar todo el historial de nuevo en cada paso.
El estado de la atención lineal reemplaza el historial completo por una matriz de tamaño fijo.
DeltaNet corrige una limitación de la atención lineal pura: si solo se suman productos externos, el estado nunca 'olvida' ni corrige información vieja, y con secuencias largas puede saturarse. DeltaNet introduce una regla delta: en vez de sumar el valor nuevo sin más, primero predice qué valor 'recordaría' el estado actual para esa clave, calcula el error entre el valor real y esa predicción, y solo escribe ese error. Gated DeltaNet suma, además, un decaimiento (gate) que atenúa el estado anterior antes de aplicar la corrección delta, dándole al modelo una forma explícita de olvidar información irrelevante con el tiempo.
Kimi Delta Attention lleva esta idea un paso más allá: en vez de un decaimiento único (escalar) para todo el estado, aplica un decaimiento distinto por canal mediante una matriz diagonal Diag(α_t). Esto le da al modelo control fino sobre qué dimensiones del estado se atenúan más rápido y cuáles se conservan por más tiempo, algo que un decaimiento escalar no puede expresar.
La secuencia completa de operaciones que ejecuta KDA en cada token t es la siguiente: primero decae el estado anterior canal por canal, después usa ese estado decaído para predecir qué valor recuerda para la clave actual, calcula el error entre el valor real y esa predicción escalado por un factor β_t, corrige el estado con ese error y finalmente lee la salida proyectando el estado corregido sobre la consulta actual (escalada por la raíz inversa de la dimensión de clave, igual que en softmax).
Traducido a código, la versión más simple de atención lineal (sin regla delta, solo para entender la mecánica de lectura y escritura del estado) se ve así:
import torch
def linear_attention_step(S, k_t, v_t, q_t):
# S: estado de forma (d_v, d_k), acumula el historial
S = S + torch.outer(v_t, k_t) # escribir: sumar v_t x k_t
o_t = S @ q_t # leer: proyectar la query sobre el estado
return S, o_t
Esta versión solo acumula, nunca corrige. El paso equivalente para Kimi Delta Attention, con decaimiento por canal y regla delta, agrega tres operaciones más:
import torch
def kda_step(S, k_t, v_t, q_t, alpha_t, beta_t, d_k):
# alpha_t: vector de decaimiento por canal (Diag(alpha_t))
S_tilde = S * alpha_t.unsqueeze(0) # decae el estado anterior, canal a canal
v_hat_t = S_tilde @ k_t # que valor 'recuerda' el estado para esta key
e_t = beta_t * (v_t - v_hat_t) # error entre el valor real y el recordado
S = S_tilde + torch.outer(e_t, k_t) # corrige el estado con la regla delta
o_t = S @ (q_t / d_k ** 0.5) # lee el estado con la query escalada
return S, o_t
Cada llamada a kda_step procesa un token y devuelve el estado actualizado junto con la salida de esa posición. En producción esto no corre como un bucle en Python token por token (sería demasiado lento): se ejecuta con kernels de Triton que procesan la secuencia en bloques (chunkwise), pero la lógica matemática es exactamente la de estas cinco líneas.
El siguiente diagrama resume el ciclo de lectura y escritura que ejecuta KDA en cada paso:
flowchart TD
A["Token t: k_t, v_t, q_t"] --> B["Decae el estado: Stilde = Sprev x Diag(alpha_t)"]
B --> C["Predice: v_hat = Stilde x k_t"]
C --> D["Calcula error: e_t = beta_t x (v_t - v_hat)"]
D --> E["Corrige: S_t = Stilde + e_t x k_t"]
E --> F["Lee salida: o_t = S_t x q_t"]
F --> G["S_t pasa al token t+1"]
💭 Clave: la identidad |v⟩⟨k|q⟩ = ⟨k|q⟩|v⟩ es el único paso matemático que convierte una suma de T² comparaciones en una actualización de estado de tamaño fijo por token. Todo lo demás (DeltaNet, Gated DeltaNet, KDA) son formas cada vez más finas de decidir qué se escribe en ese estado.
Cada variante de esta familia agrega un mecanismo sobre la anterior. La siguiente tabla resume qué aporta cada una y en qué costo por token queda, comparadas contra la atención softmax original:
VarianteQué agregaCosto por tokenDónde apareceAtención softmaxNormaliza y compara cada token contra todo el historialO(T) por token, O(T²) totalTransformer clásicoAtención lineal (sin softmax)Colapsa el historial en un estado de tamaño fijoO(1) por tokenBase teórica de variantes como RetNetDeltaNetRegla delta: corrige el estado en vez de solo acumularO(1) por tokenDeltaNetGated DeltaNetDecaimiento (gate) escalar antes de la corrección deltaO(1) por tokenArquitecturas híbridas recientesKimi Delta Attention (KDA)Decaimiento por canal Diag(α_t) + regla delta con β_tO(1) por tokenKimi (Moonshot AI), Qwen3DeltaNet introduce la corrección delta; KDA la combina con decaimiento por canal.
La implementación de referencia para DeltaNet, Gated DeltaNet y otras variantes de atención lineal vive en el repositorio flash-linear-attention en GitHub, con kernels escritos en Triton. Requiere GPU NVIDIA para correr a velocidad real; en CPU o Apple Silicon funciona en modo eager (más lento, útil solo para entender la lógica).
Instalación en Linux:
python3 -m venv fla-env
source fla-env/bin/activate
pip install --upgrade pip
pip install flash-linear-attention transformers torch
Instalación en macOS (Apple Silicon, sin kernels Triton, corre en modo eager):
python3 -m venv fla-env
source fla-env/bin/activate
pip install --upgrade pip
pip install flash-linear-attention transformers torch
Instalación en Windows (los kernels Triton necesitan WSL2 con GPU NVIDIA):
wsl --install -d Ubuntu
wsl
# dentro de la WSL2, repetir los pasos de instalación de Linux
Una vez instalado, una capa DeltaNet se puede instanciar y probar con un tensor de ejemplo:
from fla.layers import DeltaNet
import torch
layer = DeltaNet(hidden_size=1024, num_heads=8)
x = torch.randn(2, 128, 1024) # (batch, longitud, hidden)
out, *_ = layer(x)
print(out.shape) # torch.Size([2, 128, 1024])
Para confirmar que el estado se mantiene de tamaño fijo (y no crece con la longitud de la secuencia como en softmax), la forma más directa es medir memoria pico con distintas longitudes de entrada y comparar:
torch.cuda.reset_peak_memory_stats()
out, *_ = layer(x)
print(torch.cuda.max_memory_allocated() / 1e6, "MB")
Si se repite esta medición duplicando la longitud de secuencia y la memoria pico se mantiene prácticamente constante, esa es la evidencia directa de que la capa está operando en modo recurrente de estado fijo y no recalculando atención completa.
💡 Tip: antes de escribir kernels propios en Triton para experimentar con variantes de atención lineal, conviene partir del código de flash-linear-attention: ya resuelve el modo chunkwise (procesar la secuencia en bloques) que hace viable entrenar estos modelos a escala real.
La razón por la que esta familia de mecanismos importa en producción es el costo de servir contexto largo. En atención softmax estándar, el caché de claves y valores (KV cache) crece de forma proporcional a la longitud de la conversación, y cada token nuevo tiene que comparar contra ese caché completo. En un mecanismo de estado fijo como KDA, el 'caché' es la matriz de estado S, cuyo tamaño no depende de cuántos tokens ya se procesaron.
Que Moonshot AI y modelos de la familia Qwen3 hayan adoptado variantes de esta línea de investigación, según documenta el artículo de Doubleword, es una señal de que el problema del costo cuadrático de softmax en contextos largos ya no se resuelve solo con más memoria de GPU: se resuelve también cambiando la arquitectura de atención.
⚠️ Ojo: la derivación de DeltaNet asume que las claves llegan normalizadas. Si esa normalización no se aplica correctamente en la implementación, la regla delta puede volverse inestable numéricamente y el estado diverge en secuencias largas.
El compromiso que se paga por este ahorro es la pérdida de la selectividad exacta que da softmax: la normalización exponencial permite que un modelo 'ignore' casi por completo tokens irrelevantes de forma muy precisa, mientras que un estado de tamaño fijo, por más gating fino que tenga, comprime información y puede perder detalle en secuencias extremadamente largas. Por eso varias arquitecturas híbridas combinan capas de atención lineal con capas de atención completa intercaladas, en vez de reemplazar softmax en toda la red.
El propio artículo de Doubleword plantea la derivación como una base para entender variantes futuras: la familia DeltaNet, Gated DeltaNet, KDA no es un punto final sino una progresión, y cada nueva generación de modelos que necesite contextos más largos con menor costo de inferencia es candidata a introducir otra forma de gating o de corrección sobre el mismo esqueleto (decaer, predecir, corregir, leer).
Para un equipo que evalúa arquitecturas de atención lineal hoy, el punto de partida práctico sigue siendo el mismo: leer la implementación en flash-linear-attention, correr los tests del repositorio y comparar memoria y estabilidad numérica contra una capa de atención softmax estándar en el propio caso de uso, antes de decidir si el compromiso vale la pena.
📖 Resumen en Telegram: Ver resumen
Probalo vos: cloná flash-linear-attention, instalá las dependencias con el bloque de arriba y corré el snippet de DeltaNet para ver el estado de tamaño fijo en acción en minutos.
Significa que el trabajo por token no crece con la cantidad de tokens ya procesados. En softmax, el token 10.000 compara contra los 9.999 anteriores; en atención lineal, compara contra un estado de tamaño fijo que ya resume ese historial.
La regla delta calcula un error entre el valor real y el valor que el estado 'predice' para una clave dada. Si las claves no están normalizadas, esa predicción pierde la escala correcta y la corrección puede desestabilizar el estado en secuencias largas.
El artículo de Doubleword describe la derivación matemática de KDA como mecanismo de atención lineal; muchas arquitecturas que adoptan este tipo de mecanismos lo combinan con capas de atención completa intercaladas en vez de eliminar softmax de toda la red.
Gated DeltaNet aplica un decaimiento escalar (un solo número) a todo el estado antes de la corrección delta. KDA aplica un decaimiento por canal mediante la matriz diagonal Diag(α_t), dando control más fino sobre qué dimensiones del estado se olvidan más rápido.
El repositorio flash-linear-attention expone capas como DeltaNet listas para integrarse en una arquitectura de transformer estándar, reemplazando el bloque de atención; entrenar un modelo completo requiere además el resto del pipeline de entrenamiento (datos, tokenizador, loop de optimización).
El código de las capas de atención lineal está en flash-linear-attention en GitHub. Los pesos de los modelos Kimi de Moonshot AI se publican en Hugging Face bajo la organización moonshotai.
📱 ¿Te gusta este contenido? Únete a nuestro canal de Telegram @programacion donde publicamos a diario lo más relevante de tecnología, IA y desarrollo. Resúmenes rápidos, contenido fresco todos los días.