El problema
La atencion multi-cabeza es O(n²) en memoria respecto a la longitud de secuencia. Para n=128K tokens, la matriz de atencion completa ocupa 128K² × 2 bytes = 32 GB — no cabe ni en la VRAM de una H100 (80 GB) si tienes mas de un ejemplo. La solucion de FlashAttention 1 y 2 fue evitar materializar esa matriz usando tiling sobre SRAM. FA3 lleva esto al extremo para la arquitectura Hopper (H100/H800), introduciendo tres mecanismos nuevos: WGMMA, TMA y solapamiento asimetrico producer-consumer.
Arquitectura y metodo
FlashAttention 1 y 2 ya hacian tiling: dividian Q en bloques de Br filas y K,V en bloques de Bc filas, los cargaban a SRAM, computaban la atencion parcial sin escribir la matriz completa a HBM, y mantenian el online softmax (estadisticos m y l) para corregir la normalizacion al final. El resultado: O(n) memoria en lugar de O(n²).
FA3 introduce tres mejoras especificas para Hopper:
1. WGMMA (Warpgroup Matrix Multiply Accumulate): La H100 tiene una instruccion de tensor core que opera sobre warp groups (4 warps = 128 threads) en lugar de warps individuales. WGMMA puede multiplicar matrices de 64×16 × 16×256 directamente desde SRAM a registros, con latencia asincrona. FA3 reestructura el kernel para explotar WGMMA con wgmma.mma_async.
2. TMA (Tensor Memory Accelerator): Hardware dedicado en H100 para copias asincronas HBM→SRAM. TMA permite al SM delegar las cargas de tiles y seguir computando — un thread lanza el cp.async.bulk y el resto del warp no espera.
3. Solapamiento producer-consumer con ping-pong: FA3 divide los warps en dos grupos. El producer (1 warp) gestiona las cargas TMA. El consumer (3 warps) ejecuta WGMMA sobre el buffer que ya esta listo mientras el producer carga el siguiente. Dos buffers SRAM en alternancia (ping-pong) eliminan la sincronizacion entre carga y computo.
Para FP8, FA3 implementa atencion en FP8 con cuantizacion por bloque en linea, aprovechando que los tensor cores H100 hacen FP8 a doble velocidad que BF16.
Contribuciones clave
- Primer kernel de atencion con solapamiento producer-consumer real en GPU — anterior a FA3 todo era sincronico.
- Soporte FP8 nativo con cuantizacion en linea (sin dequantizar a BF16 en cada paso).
- Atencion causal sin padding en el kernel — los bloques triangulares inferiores se saltan directamente.
Metricas y resultados
En H100 SXM5 con secuencias de 8K tokens, d_head=128:
| Metodo | TFLOP/s (BF16) | % MFU |
|---|---|---|
| Atencion estandar PyTorch | 120 | 13% |
| FlashAttention-2 | 580 | 61% |
| FlashAttention-3 | 740 | 78% |
| FA3 FP8 | 1140 | 75% FP8 MFU |
Speedup sobre FA2: 1.5-2.0x segun longitud de secuencia. Para secuencias cortas (512) el beneficio es menor porque el kernel overhead domina. Para 32K+ tokens, FA3 FP8 es la unica opcion practica.
Retos de implementacion
- Solo H100/H800: WGMMA y TMA son instrucciones exclusivas de Hopper. No hay fallback automatico a A100. El codigo fuente usa
#ifdef __CUDA_ARCH__ >= 900. - CUDA 12.3+ obligatorio: Las intrinsicas de TMA requieren el toolkit nuevo.
- Instalacion:
pip install flash-attn==2.6.0instala FA2. Para FA3 hay que compilar desde el branchhopperdel repo o usar la version empaquetada enflash-attn>=2.7.0. - Cabezas no multiplos de 64: El kernel WGMMA opera en tiles de 64. Si
d_headno es multiplo (ej: 96 en algunos modelos antiguos), hay padding con overhead. - Debugging imposible en SRAM: Los errores en el ping-pong son silenciosos — si el producer no sincroniza bien, el consumer lee datos corruptos sin excepcion. Usar
cuda-memcheckocompute-sanitizer --tool initcheck.
Como lo integraria en Zeropithos o Dibro
Dibro usa Qwen2.5-Coder-32B via llama-server. llama.cpp ya integra FA2 con el flag --flash-attn. Si el servidor corriera en H100 (chuck tiene RTX 3090, arquitectura Ampere — sin WGMMA), FA3 no aplica directamente. Pero si en algun momento se mueve a un nodo con H100 alquilado para inferencia larga, el salto de FA2 a FA3 daria un 60-80% mas de throughput en contextos de 32K+ tokens — exactamente el rango donde Dibro hace analisis de binarios grandes con el reverse tool. La integracion seria trivial: actualizar llama.cpp a una version con soporte FA3 y pasar --flash-attn --flash-attn-v3.
Conclusion
FA3 no es un paper de ideas nuevas — es ingenieria de bajo nivel impecable. Su contribucion real es demostrar que el cuello de botella en LLM inference no es el algoritmo sino la jerarquia de memoria GPU, y que explotar hardware especifico (WGMMA, TMA) da saltos de 2x que ningun cambio de arquitectura puede igualar. Para quien trabaje con contextos largos en H100, FA3 es obligatorio.