FlashAttention (ES)
FlashAttention es un algoritmo revolucionario para calcular el mecanismo de atención (attention), diseñado para acelerar significativamente el entrenamiento y la inferencia de modelos de lenguaje grandes (LLM) manteniendo al mismo tiempo la precisión computacional completa. El algoritmo fue presentado por primera vez en 2022 por un equipo de investigadores de la Universidad de Stanford, liderado por Tri Dao[1].
La idea clave de FlashAttention consiste en reorganizar los cálculos teniendo en cuenta la jerarquía de memoria de la GPU, lo que permite minimizar el número de accesos a la memoria lenta y eliminar el principal cuello de botella del mecanismo de atención estándar.
El problema de la atención estándar
El mecanismo de autoatención estándar en los transformadores se calcula mediante la siguiente fórmula: donde Q, K, V son las matrices de consulta (queries), clave (keys) y valor (values).
El principal problema de este enfoque es su complejidad cuadrática en tiempo y memoria (O(N²)) con respecto a la longitud de la secuencia N[1]. En una implementación ingenua, es necesario calcular y almacenar en la memoria de la GPU la matriz de atención completa S de tamaño N×N, lo que conduce a dos problemas críticos:
- Alto consumo de memoria: Almacenar la matriz N×N se vuelve inviable cuando se trabaja con contextos largos.
- Operaciones de entrada/salida (E/S): El principal cuello de botella no es la cantidad de operaciones aritméticas, sino los constantes accesos a la memoria lenta de la GPU.
Jerarquía de memoria de la GPU
Para comprender el problema, es importante distinguir dos tipos de memoria en una GPU (utilizando como ejemplo la NVIDIA A100):
- SRAM (memoria estática): Memoria rápida en el chip de pequeño tamaño (~20 MB) con un enorme ancho de banda (hasta 19 TB/s).
- HBM (memoria de alto ancho de banda): Memoria más lenta de gran capacidad (40–80 GB) con un ancho de banda mucho menor (alrededor de 1.5 TB/s)[2].
Esta asimetría hace que el algoritmo de atención estándar esté limitado por el ancho de banda de la memoria (memory-bound), ya que lee y escribe constantemente grandes matrices desde la HBM lenta, lo que constituye la principal fuente de latencia.
Innovaciones clave de FlashAttention
FlashAttention es un algoritmo consciente de E/S (IO-aware) que resuelve el problema minimizando los accesos a la HBM. Esto se logra mediante tres técnicas principales.
Tiling y procesamiento por bloques
En lugar de procesar la matriz completa de una vez, FlashAttention divide las matrices de entrada Q, K y V en pequeños bloques (tiles) que caben en la rápida SRAM. El algoritmo carga secuencialmente estos bloques, realiza todos los cálculos de atención para ellos y actualiza el resultado final, sin almacenar la matriz de atención completa en la lenta HBM[1].
Cálculo de Softmax en línea
Un avance técnico clave fue el cálculo "en línea" (online) de Softmax. El Softmax estándar requiere conocer todos los elementos del vector de entrada para la normalización. FlashAttention utiliza un algoritmo modificado que permite calcular Softmax por partes. Mantiene dos valores intermedios (el máximo actual y la suma de las exponenciales) que se actualizan a medida que se procesan nuevos bloques, lo que permite obtener el resultado exacto sin acceder a toda la matriz a la vez[2].
Fusión de operaciones en un único kernel de CUDA
Todas las operaciones de atención (multiplicación de matrices QKᵀ, enmascaramiento, Softmax, multiplicación por V) se combinan en un único kernel de CUDA fusionado (fused kernel). Esto reduce drásticamente el número de operaciones de lectura/escritura en la HBM: en lugar de múltiples pasadas sobre toda la matriz, el algoritmo carga un bloque en la SRAM una vez, realiza todos los cálculos y escribe solo el resultado final.
Eficiencia teórica y práctica
Complejidad y optimalidad
FlashAttention reduce el consumo de memoria de O(N²) a O(N), lo que garantiza una escalabilidad lineal. Se ha demostrado que la complejidad de E/S del algoritmo es teóricamente óptima para el cálculo de la atención en una jerarquía de memoria de dos niveles, lo que significa que es imposible realizar una atención exacta más rápido sin modificar el hardware[3].
Resultados empíricos
La primera versión de FlashAttention demostró mejoras significativas:
- Aceleración:
- BERT-large (longitud 512): 15% de aceleración en el entrenamiento.
- GPT-2 (longitud 1K): aceleración de 3 veces.
- Tareas de Long-Range Arena (1K-4K): aceleración de 2.4 veces[1].
- Ahorro de memoria: Hasta 20 veces de ahorro de memoria en comparación con las implementaciones base exactas.
- Mejora en la calidad de los modelos: Gracias a la capacidad de trabajar con contextos más largos, FlashAttention no solo no pierde calidad, sino que la mejora. Por ejemplo, la perplejidad de GPT-2 mejoró en 0.7 puntos, y la precisión en tareas de clasificación de documentos largos aumentó en 6.4 puntos[1].
Evolución y desarrollos posteriores
El éxito de FlashAttention dio lugar a toda una serie de algoritmos orientados al hardware.
FlashAttention-2 (2023)
La segunda versión se centró en un uso más completo de los recursos de la GPU. En el FlashAttention original, la eficiencia en una NVIDIA A100 era solo del 25–40% del máximo. FlashAttention-2 introdujo mejoras en la paralelización de los cálculos, lo que permitió[4]:
- Lograr una aceleración del doble en comparación con la primera versión.
- Aumentar la utilización de la GPU hasta el 50–73% del máximo teórico.
- Ampliar el soporte a cabezales de atención de tamaño 256, así como a arquitecturas Multi-Query Attention (MQA).
FlashAttention-3 (2024)
La tercera versión fue optimizada específicamente para la arquitectura de GPU NVIDIA Hopper (H100)[5]. Utiliza nuevas capacidades de hardware, como la asincronía de los Tensor Cores y el soporte para FP8, lo que permitió:
- Alcanzar una aceleración adicional de 1.5 a 2 veces en comparación con FlashAttention-2.
- Lograr un rendimiento de hasta 740 TFLOPS en FP16 y cerca de 1.2 PFLOPS en FP8.
Soluciones especializadas
Las ideas de FlashAttention se han desarrollado en otros proyectos:
- FlashInfer (2025): Un motor de atención personalizable, optimizado específicamente para tareas de inferencia de LLMs. Se centra en el trabajo eficiente con la caché KV en modo de generación por streaming[6].
- FlashMLA (2024): Una implementación de atención con compresión de la caché de contexto (latent attention), que permite ahorrar memoria en secuencias muy largas con una pérdida mínima de información[7].
Influencia en la industria y el ecosistema
FlashAttention se convirtió en un avance fundamental y rápidamente se transformó en el estándar de la industria para el entrenamiento e inferencia eficientes de LLMs. Ha sido integrado en bibliotecas clave como PyTorch y Hugging Face, y es utilizado en la mayoría de los modelos de lenguaje grandes (LLaMA, MPT, Falcon, Claude, etc.).
Fueron FlashAttention y sus versiones posteriores los que jugaron un papel decisivo en el aumento de las ventanas de contexto de los modelos de lenguaje: de 2–4 mil tokens (GPT-3) a 128 mil tokens (GPT-4) e incluso a millones de tokens en modelos experimentales[8]. El algoritmo eliminó uno de los principales obstáculos para la escalabilidad de los transformadores, abriendo nuevas posibilidades para aplicaciones de IA, desde el análisis de documentos largos hasta la comprensión multimodal.
Enlaces externos
Bibliografía
- Dao, T. et al. (2022). FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. arXiv:2205.14135.
- Dao, T. (2023). FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning. arXiv:2307.08691.
- Shah, J. et al. (2024). FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-Precision. arXiv:2407.08608.
- Kwon, W. et al. (2023). Efficient Memory Management for Large Language Model Serving with PagedAttention. arXiv:2309.06180.
- Hong, K. et al. (2023). FlashDecoding++: Faster Large Language Model Inference on GPUs. arXiv:2311.01282.
- Ye, Z. et al. (2025). FlashInfer: Efficient and Customizable Attention Engine for LLM Inference Serving. arXiv:2501.01005.
- Dege, P. et al. (2025). FlashMLA-ETAP: Efficient Transpose Attention Pipeline for Accelerating MLA Inference on NVIDIA H20 GPUs. arXiv:2506.01969.
- Wang, G. et al. (2025). FlashMask: Efficient and Rich Mask Extension of FlashAttention. OpenReview wUtXB43Chi.
- Dao, T. et al. (2022). FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness (OpenReview version). OpenReview H4DqfPSibmx.
- Gholami, A. et al. (2024). FlashAttention on a Napkin: A Diagrammatic Approach to Deep Learning IO-Awareness. OpenReview pF2ukh7HxA.
Referencias
- ↑ 1.0 1.1 1.2 1.3 1.4 Dao, Tri, et al. «FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness». arXiv:2205.14135 [cs.LG], 28 de mayo de 2022. [1]
- ↑ 2.0 2.1 Dao, Tri, et al. «FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness». OpenReview. [2]
- ↑ «We're Training AI Twice as Fast This Year as Last». IEEE Spectrum. [3]
- ↑ Dao, Tri. «FlashAttention-2». tridao.me. [4]
- ↑ «FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision». PyTorch Blog. [5]
- ↑ «[2501.01005] FlashInfer: Efficient and Customizable Attention Engine for LLM Inference Serving». arXiv. [6]
- ↑ «GitHub - deepseek-ai/FlashMLA: FlashMLA: Efficient MLA decoding kernels». GitHub. [7]
- ↑ «The Evolution of Flash Attention: Revolutionizing Transformer Efficiency». Medium. [8]