FlashAttention
Attention compara muchas parejas de fragmentos de texto. Si guardamos todos los resultados intermedios en la gran memoria de la tarjeta de cálculo, transferir los datos puede resultar costoso. Quieres realizar el mismo cálculo reduciendo el traslado y el almacenamiento de una tabla enorme.
FlashAttention organiza el cálculo de attention por bloques en la pequeña memoria rápida del chip. No guarda toda la matriz de comparaciones y pesos en la memoria grande de la GPU, la tarjeta usada para cálculos paralelos.
Cada porción aporta parte del resultado, y el algoritmo actualiza correctamente la normalización y la suma. No basta con calcular medias separadas de las porciones y promediarlas: cada porción puede tener una contribución distinta.
Se conserva la operación matemática de attention exacto, con las diferencias habituales de la aritmética computacional. Para attention denso, el número de comparaciones sigue creciendo de forma cuadrática. La mejora procede de la organización de la memoria y la ejecución, no de eliminar dependencias entre parejas.
Fuente del mecanismo: Dao et al., §2–3.2, algoritmo 1.
Mecanismo y detalles

Para la atención densa, el número de operaciones sigue creciendo cuadráticamente con la longitud de la secuencia. La ganancia procede de la organización del cálculo y de la memoria. Durante el entrenamiento se recalculan algunas cantidades en la pasada hacia atrás en lugar de almacenarlas. «Exacta» significa que no se aproxima la propia operación; un orden diferente de operaciones en coma flotante puede producir pequeñas diferencias numéricas.
No basta con promediar los resultados de los bloques
Softmax normaliza los pesos respecto a todas las claves permitidas. Los bloques pueden tener sumas de pesos sin normalizar muy distintas. Por tanto, hay que asignar a cada uno la contribución correcta al resultado.
Esto puede hacerse gradualmente, manteniendo la puntuación máxima , la suma y el numerador . El resultado es . Cuando un nuevo bloque aumenta el máximo a , hay que reescalar los anteriores y por antes de añadir los nuevos términos. Milakov y Gimelshein, §3, algoritmo 3 explican el mecanismo de normalización estable.
El ejemplo propio utiliza las puntuaciones ya dadas y los valores . Los pesos antes de normalizar tienen proporciones , así que el resultado es . Los resultados de los dos bloques son 0 y 1. Su media normal da el valor incorrecto de 0,5.
La demostración aísla la agregación para una consulta y valores escalares. El kernel completo calcula bloques de puntuaciones a partir de Q y K en la GPU; aquí, los números ya dados permiten comprobar el cálculo a mano. Cambiar el tamaño del bloque no es un benchmark.
La atención de consultas agrupadas reduce el número de cabezas K/V, y FlashAttention cambia cómo se transfieren y procesan los datos. Estas técnicas pueden combinarse. FlashAttention 2 es una evolución posterior de la implementación; el mecanismo descrito aquí procede del primer trabajo.
Utilizo contenido generado por IA como parte de mi proceso de aprendizaje diario.