live

FlashAttention-3: Como Exprimir H100 al Limite con Atencion IO-Aware

FlashAttention-3 reescribe la atencion para H100/H800 explotando WGMMA, TMA y solapamiento de computo/memoria, logrando 1.5-2x sobre FA2 y 75% de FP8 MFU.

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.

HBM (80 GB) Q, K, V matrices 2 TB/s bandwidth tile Q [Br x d] tile K [Bc x d] tile V [Bc x d] output O [Br x d] TMA SRAM (228 KB/SM) shared memory 10 TB/s bandwidth softmax stats (m, l) ping-pong buffers WGMMA in-flight async Registers per warp-group ~80 TB/s effective WGMMA accumulators online softmax state FA3: solapamiento producer (TMA load) + consumer (WGMMA compute) con ping-pong

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.0 instala FA2. Para FA3 hay que compilar desde el branch hopper del repo o usar la version empaquetada en flash-attn>=2.7.0.
  • Cabezas no multiplos de 64: El kernel WGMMA opera en tiles de 64. Si d_head no 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-memcheck o compute-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.

aqui cualquier cosa mientras cuadramos el logo