Qué es FlashAttention y por qué importa a tu startup
FlashAttention es un algoritmo de atención exacto y IO-aware que reorganiza cómo se mueve la información entre la HBM (memoria de alto ancho de banda, fuera del chip) y la SRAM (memoria rápida dentro del chip) de la GPU. La idea, en tres palabras según el handbook que motiva esta nota, es tiling + online softmax + recomputation: dividir la atención en bloques pequeños que caben en SRAM y procesarlos en streaming, sin necesidad de materializar la matriz N×N de scores en memoria.
Para un founder hispanohablante esto no es curiosidad académica: el costo de entrenamiento e inferencia de un LLM depende más de cuántas veces la GPU va a HBM que de cuántas multiplicaciones ejecuta. La optimización original de Tri Dao y colaboradores (2022) reportó mejoras medibles frente a baselines públicos: 15% de speedup end-to-end en BERT-large (seq. length 512) comparado con el récord de velocidad de entrenamiento de MLPerf 1.1, 3× de speedup en GPT-2 (seq. length 1K) y 2.4× en long-range arena (seq. length 1K-4K), según el paper original publicado en arXiv (mayo de 2022).
El problema que FlashAttention resuelve
La implementación clásica de atención densa tiene un patrón costoso: calcula el score matrix S, lo escribe en HBM, lo lee de vuelta, calcula las probabilidades P, las escribe, las lee de nuevo y las multiplica por V. Cada valor intermedio se mueve varias veces entre HBM y SRAM, y el tamaño de S y P crece de forma cuadrática con la longitud de la secuencia.
🤖 La IA no es solo para leer sobre ella
En la comunidad la aplicamos: automatización, agentes IA y herramientas reales para emprender, no solo para informarte.
👥 Aplicarla en la comunidadEl handbook pone el tamaño en perspectiva con un ejemplo: para un solo tensor de atención con batch=1, heads=32, seq_len=8192 y dtype FP16, la matriz completa es de 4 GiB ([1, 32, 8192, 8192]). A longitudes largas, ese materializado domina el ancho de banda y bloquea la GPU, aunque las multiplicaciones en sí sean baratas.
La idea clave del paper original es que la complejidad aritmética de la atención densa sigue siendo O(N²·d) — FlashAttention no cambia esa fórmula. Lo que cambia es la cantidad de lecturas y escrituras a HBM: el algoritmo analiza formalmente cuántas transferencias necesita y demuestra que, bajo un modelo de memoria de dos niveles (HBM + SRAM), es óptimo para un rango de tamaños de SRAM.
Tiling, online softmax y recomputation: el truco en tres piezas
Tiling. En lugar de procesar toda la matriz de atención a la vez, FlashAttention divide Q, K y V en bloques (tiles). Un tile de queries Qi permanece en SRAM mientras se hace streaming de múltiples tiles de Kj y Vj que entran y salen. Esto maximiza la reutilización de datos en chip: un bloque de Q puede contribuir a muchos productos QKᵀ antes de volver a escribirse a HBM.
Online softmax. El reto es que softmax no se paraleliza trivialmente: cada fila necesita su máximo global y su denominador estable para evitar overflow. La solución es una recurrencia streaming: mantener estado por fila con m (máximo visto hasta ahora), ℓ (suma exponencial estable bajo esa referencia) y a (acumulador de valores ponderados). Cuando llega un nuevo bloque y aparece un máximo mayor, se rescala todo lo anterior multiplicando por e^(mviejo − mnuevo). La recurrencia no es una aproximación: es la identidad matemática e^(xj − mviejo) · e^(mviejo − mnuevo) = e^(xj − mnuevo).
El handbook lo demuestra con un ejemplo paso a paso: para scores [2, 1, 4, 3] y valores asociados, el bloque [2, 1] se procesa con m₁=2, acumulador [1, e⁻¹]; cuando llega [4, 3] con m_b=4, se aplica α = e⁻², se reescalan las estadísticas previas y se obtiene el mismo vector final que calcularía softmax en un solo paso sobre toda la fila.
Recomputation. En el backward pass, FlashAttention vuelve a generar los tiles de scores en lugar de almacenarlos, ahorrando memoria a cambio de algo de cómputo adicional. En hardware moderno, donde el throughput de matmuls ha crecido más rápido que el ancho de banda, ese trade es favorable.
Lo que FlashAttention cambia y lo que no cambia
El paper original insiste en un punto que muchos pasan por alto: FlashAttention no aproxima la atención. La función matemática evaluada es exactamente la atención densa de producto punto escalado. Lo que cambia es:
- El schedule de ejecución: cómo se particionan y secuencian los bloques.
- El tráfico de memoria: menos idas y vueltas a HBM.
- Los intermedios almacenados: la matriz N×N nunca se materializa en HBM.
Lo que no cambia:
- La complejidad aritmética: sigue siendo O(N²·d) para atención densa.
- La función matemática: misma softmax, mismas interacciones query-key.
- La compatibilidad funcional: se puede combinar con masking causal, RoPE, MQA/GQA y otras variantes, porque es una familia de kernels, no un reemplazo del mecanismo.
El paper también introdujo una extensión block-sparse, pero ese caso es distinto: omitir bloques sí cambia qué interacciones se calculan, y por tanto es matemáticamente diferente a la atención densa exacta.
Por qué el IO importa más que los FLOPs
Las GPUs modernas (A100, H100 y siguientes) tienen jerarquías de memoria desbalanceadas: SRAM es ~20× más rápida que HBM, pero cabe órdenes de magnitud menos. La métrica relevante para entrenamiento a long context no es cuántos FLOPs ejecutas, sino cuántas veces tu kernel toca HBM. Por eso la afirmación central del paper — “IO-aware” — es más una consigna de diseño que un detalle de implementación.
Esto tiene una consecuencia práctica para founders que entrenan o fine-tunean modelos: las ganancias de FlashAttention son especialmente grandes durante el backward pass, donde naive autograd querría mantener grandes intermedios para retropropagar. En el paper original, esa optimización permitió entrenar Transformers con secuencias de 16K y 64K por primera vez con resultados mejores que el azar (Path-X 61.4% y Path-256 63.1%, según el abstract de arXiv).
¿Qué significa esto para tu startup?
Si estás construyendo sobre LLMs o entrenando tu propio modelo, FlashAttention debería ser default, no optimización.
Acciones concretas:
- Verifica qué kernel usa tu framework. En PyTorch y HuggingFace Transformers, FlashAttention se activa típicamente con flags como
attn_implementation="flash_attention_2"o"flash_attention_3"en la config del modelo. Si no lo estás activando explícitamente, probablemente estás corriendo una variante menos eficiente. - Dimensiona tu contexto con cabeza. Como la atención densa sigue siendo O(N²·d), doblar el contexto cuatriplica aproximadamente las FLOPs de atención y el tráfico de HBM. FlashAttention te compra margen — más contexto con la misma GPU — pero no convierte contexto largo en “gratis”. Para reducir costo de verdad a contextos muy largos necesitas sparsity, sliding window o arquitecturas alternativas como Mamba/State Space Models.
- Mira la generación del kernel, no solo el nombre. FA-2, FA-3 y FA-4 explotan diferentes generaciones de hardware (Hopper, Ada/Blackwell) con tiling, async copies y tensor cores específicos. Si entrenas en hardware reciente, asegúrate de usar la versión compatible con tu GPU; el speedup reportado en el paper de 2022 es sobre A100, y las versiones posteriores aprovechan instrucciones nuevas.
Limitaciones y cuándo no usarlo
FlashAttention no es la respuesta universal:
- Sparsity real: si necesitas reducir el número de pares query-key (no solo reorganizar cómo se mueven), necesitas sparse attention, local attention o un patrón específico de atención lineal.
- Hardware sin suficientes recursos de SRAM: en GPUs muy viejas o con SRAM limitado, el tiling deja de pagar.
- Casos donde el modelo ya está dominado por otras partes (FFN, embeddings, decoders): si la atención no es el cuello de botella, la ganancia marginal se diluye.
Fuentes
- Understanding FlashAttention Pt 1: Personal Notes (handbook original)
- FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness (paper original, arXiv 2205.14135)
🤖 La IA no es solo para leer sobre ella
En la comunidad la aplicamos: automatización, agentes IA y herramientas reales para emprender, no solo para informarte.
👥 Aplicarla en la comunidad














