Resumen: una suma normalizada puede calcularse por bloques
Obtener una fila de atención no exige conservar todos sus pesos. Bastan un máximo, un normalizador y un vector acumulado si cambian juntos de escala al llegar una puntuación mayor. Derivamos la identidad, la verificamos con logits que desbordan la exponencial ingenua y mostramos un fallo: reescalar solo el denominador produce una salida incorrecta aunque softmax parezca estable. Es una derivación comentada con código ejecutado, no un algoritmo nuevo ni un benchmark GPU.
1. Qué objeto queremos calcular
Una consulta q y N claves k_j de dimensión d_k producen puntuaciones adimensionales s_j=qᵀk_j/√d_k, incluyendo cualquier sesgo. Cada clave tiene un valor v_j en R^(d_v). La salida o es una combinación convexa: coeficientes positivos que suman uno. Consideramos una fila, sin dropout y al menos una clave no enmascarada. Las excluidas no contribuyen. Cambiamos el orden de cálculo, no la función matemática.
Restar el mismo máximo multiplica numerador y denominador por exp(−m), sin cambiar el cociente. Con puntuaciones finitas los exponentes son no positivos: ningún término supera uno y alguno vale uno. Así, l está entre 1 y N en aritmética real. Evita el desbordamiento exponencial, no todo error: términos diminutos pueden subdesbordarse y las sumas siguen redondeándose.
2. El invariante y el cambio de escala
Tras t claves mantenemos m_t, máximo del prefijo; l_t, suma exponencial respecto a ese máximo; y u_t, suma ponderada de valores. Llega (s,v). El nuevo máximo es m′=max(m_t,s). Cada contribución antigua exp(s_i−m_t) se multiplica por exp(m_t−m′), convirtiéndose en exp(s_i−m′). El mismo factor multiplica u_t porque usaba los mismos pesos. Después añadimos la nueva clave.
La prueba es un invariante de bucle: sustituir las definiciones de l y u da exactamente las sumas del prefijo ampliado. Una clave inicializa m=s_1, l=1, u=v_1, evitando −∞−(−∞). Un estado vacío debe tratarse explícitamente. Hay que distinguir el numerador u de la salida o: tratar o como u pierde el factor l anterior.
3. Un ejemplo que hace visible el error
Usamos s=(1000,1001,999), v_1=(1,0), v_2=(0,2), v_3=(3,−1). Python desborda exp(1000), pero los pesos estabilizados son proporcionales a (e^−1,1,e^−2). Tras la primera clave l=1 y u=(1,0). La segunda eleva el máximo a 1001: el numerador antiguo pasa a (e^−1,0), no sigue en (1,0). Así l=1+e^−1 y u=(e^−1,2). La tercera mantiene el máximo y añade e^−2·(3,−1).
| Paso | m | l | o₁ | o₂ |
|---|---|---|---|---|
| 1 | 1000.0 | 1.000000000 | 1.000000000 | 0.000000000 |
| 2 | 1001.0 | 1.367879441 | 0.268941421 | 1.462117157 |
| 3 | 1001.0 | 1.503214724 | 0.514820191 | 1.240451338 |
El resultado es o=(0,514820191; 1,240451338). Reescalar l pero no u en el segundo paso da 0,935332675 en el primer componente. El segundo queda correcto por casualidad, pues su contribución anterior era cero: comprobar uno solo podría ocultar el fallo. Ambos resultados son finitos. Ausencia de NaN y desbordamiento no demuestra corrección.

4. Fusionar bloques sin conservar pesos
La misma álgebra fusiona conjuntos disjuntos A y B resumidos por (m,l,u). Se toma m=max(m_A,m_B), se reescalan normalizadores y numeradores y se suman. El resultado representa la unión: en aritmética real es asociativo y conmutativo. En coma flotante el árbol de reducción puede cambiar los últimos bits. Permite bloques paralelos sin prometer identidad bit a bit.
El archivo compara bloques de 1, 7, 32 y 257 con materialización estable para 257 claves y valores de cuatro componentes. El generador Python usa semilla 20260924; puntuaciones uniformes en [−a,a] con a=1,10,1000 y valores en [−2,2]. El máximo error absoluto de doce comparaciones es aproximadamente 6,66×10^−16. Es una comprobación sintética con float Python, no una cota universal ni medida FP16/BF16. Código, semilla y resultados están adjuntos.
5. Memoria, complejidad y mediciones pendientes
Para una fila, el estado persistente necesita d_v+2 escalares, además de consulta, bloques de claves y valores y buffers. Con d_v=64 en FP32 son 264 bytes, no toda la memoria del kernel. Una matriz N×N FP32 ocupa 4N² bytes por cabeza y elemento del batch: 64 MiB con N=4096, 256 MiB con N=8192. Una implementación ingenua puede materializar también pesos normalizados. Evitar matrices no comprime K y V ni elimina la KV cache del decodificador.
Calcular densamente todos los pares consulta-clave sigue siendo cuadrático: O(N²(d_k+d_v)). La recurrencia no convierte atención exacta en atención lineal. Un bucle escalar Python puede ser más lento que multiplicación vectorizada; el beneficio hardware exige bloques, fusión y accesos adecuados. Medirlo requiere GPU, dtype, dimensiones, máscara, batch, versiones, calentamiento y sincronización. No hicimos ese benchmark ni informamos aceleración.
Filas totalmente enmascaradas necesitan una convención: 0/0 no define distribución. El código rechaza un estado sin claves y exige puntuaciones finitas; las claves enmascaradas se filtran antes. Valores de distinto signo pueden cancelarse en el numerador; precisión reducida exige elegir acumulador. Gradientes y dropout requieren lógica adicional: verificar el forward de una fila no certifica un entrenamiento completo.
6. Fuentes y conclusión
Milakov y Gimelshein describen normalización online (2018, arXiv v2); sus experimentos Tesla V100, CUDA 9.1 y distintos batch no trasladan tiempos a nuestro código. FlashAttention de Dao y colegas (2022, arXiv v2) conecta bloques con menor tráfico HBM y recomputación del backward. Leímos algoritmo, análisis y experimentos. Explicamos un fundamento sin reproducir benchmarks ni revisar desarrollos posteriores.
La conclusión verificable es el invariante compartido por numerador y denominador. El máximo no es solo un truco contra desbordamiento: define la escala numérica acumulada. Cuando cambia se convierten todas las contribuciones previas. Eso permite prescindir de la matriz de pesos conservando el resultado en aritmética real.
Milakov M., Gimelshein N. (2018), Online normalizer calculation for softmax, arXiv:1805.02867v2.
from math import exp
scores = [1000., 1001., 999.]
values = [[1., 0.], [0., 2.], [3., -1.]]
m, total, numerator = scores[0], 1., values[0][:]
for s, v in zip(scores[1:], values[1:]):
new_m = max(m, s)
old_scale, new_scale = exp(m-new_m), exp(s-new_m)
total = total*old_scale + new_scale
numerator = [u*old_scale + x*new_scale for u, x in zip(numerator, v)]
m = new_m
print([u/total for u in numerator])
Código, datos e instrucciones · JSON. Cálculos didácticos ejecutados con Python 3.14.0; figuras con Matplotlib 3.11.2. Análisis asistido por IA, sin afirmar revisión por pares ni humana. Portada original ImageGen ilustrativa: no documenta personas, sedes ni instalaciones de EL-AI. Fuentes consultadas el 24 de septiembre de 2026.

