En este artículo, exploraremos uno de los trabajos más críticos en IA moderna: Flash Attention. Aprenderás a implementar Flash Attention utilizando cuTile, con un código completo listo para producción, y cómo optimizarlo a través de la experiencia de «trampa y rescate».
Requisitos del entorno:
- CUDA 13.1 o superior
- Arquitectura GPU: NVIDIA Blackwell (por ejemplo, NVIDIA B200, serie GeForce RTX 50)
- Python: 3.10 o superior
Consulta la documentación de inicio rápido para más información sobre la instalación de cuTile Python.
¿Qué es la atención?
El mecanismo de atención es el corazón computacional de los modelos transformadores. Permite que cada token en una secuencia evalúe cada otro token y decida cuánto ponderar sus contribuciones. Matemáticamente, para las matrices de entrada Query (Q), Key (K) y Value (V), la salida se define como:
(O = text{softmax}left(frac{QK^T}{sqrt{d}}right)V)
Donde:
- (Q text{ tiene forma } (N,d), donde N son tokens de consulta, cada uno con dimensión d.)
- (K text{ tiene forma } (N,d), donde N son tokens clave.)
- (V text{ tiene forma } (N,d), donde N son tokens de valor.)
- (text{La matriz intermedia } QK^{T} text{ tiene forma } (N,N), lo que representa un problema.)
El problema del ancho de banda de memoria
Con una longitud de secuencia de (N = 16,384), común en los LLM modernos, la matriz de atención (QK^{T}) contiene (N^2 = 268) millones de elementos. En FP16, eso significa 512 MB de almacenamiento intermedio por cabeza de atención y por elemento de lote.
Las implementaciones estándar de atención son:
- Calcular la matriz de atención completa (N x N) y escribirla en la memoria global (lento).
- Aplicar softmax fila por fila.
- Leer la matriz de vuelta y multiplicar por (V).
Este enfoque es dependiente de la memoria, ya que la GPU pasa la mayor parte del tiempo esperando que los datos se muevan entre HBM y las unidades de cálculo, en lugar de realizar cálculos.
Cómo Flash Attention resuelve el problema del ancho de banda de memoria
Flash Attention es un algoritmo consciente de I/O que nunca materializa la matriz completa (N x N). En cambio,:
- Divide la computación: Procesa (Q, K, V) en bloques pequeños que caben en la rápida SMEM en chip.
- Utiliza softmax en línea: Calcula softmax de manera incremental sin necesidad de la fila completa.
- Fusiona operaciones: Combina la multiplicación de matrices y softmax en una única pasada de kernel.
El resultado es una aceleración de 2 a 4 veces y ahorros significativos de memoria, lo que permite longitudes de contexto más largas.
Entendiendo softmax en línea
La clave del algoritmo de Flash Attention es el truco de softmax en línea. La softmax segura numéricamente estable requiere conocer el valor máximo a lo largo de toda la fila antes de calcular:
(text{softmax}(x_i) = frac{e^{x_i – max(x)}}{sum_j e^{x_j – max(x)}})
Sin embargo, al procesar mosaicos, no tenemos acceso a la fila completa. La softmax en línea resuelve esto manteniendo estadísticas acumulativas que se pueden actualizar de manera incremental.
Atención causal y atención agrupada por consulta
Antes de sumergirnos en la implementación, entendamos dos variantes importantes de atención utilizadas en los LLM modernos:
Atención causal
En modelos de lenguaje autorregresivos como GPT, LLaMA y Claude, cada token solo puede asistir a tokens anteriores en la secuencia, no a futuros. Esto evita «hacer trampa» durante el entrenamiento, donde el modelo mira hacia adelante para predecir la próxima palabra.
Matemáticamente, aplicamos una máscara triangular a las puntuaciones de atención:
(text{mask}_{ij} = begin{cases} 0 & text{si } i geq j text{ (posición de consulta ≥ posición clave)} -infty & text{si } i < j text{ (tokens futuros)} end{cases})
La atención enmascarada se convierte en:
(O = text{softmax}left(frac{QK^T}{sqrt{d}} + text{mask}right)V)
Agregar (-infinito) a posiciones futuras asegura que se conviertan en cero después de softmax, bloqueando efectivamente el flujo de información de los tokens futuros.

Con el enmascaramiento causal, aproximadamente la mitad de la matriz de atención está enmascarada (el triángulo superior). Podemos omitir el cálculo de estos mosaicos enmascarados por completo, proporcionando una aceleración algorítmica de 2x. Esto es crucial para la optimización de división de bucles K.
Parte 1: El kernel de atención flash en CUDA Tile
Implementemos Flash Attention paso a paso. Nuestra base utiliza mosaicos pequeños de 64×64 y un código sencillo: correcto pero aún no optimizado.
1. Definiendo la interfaz del kernel
En cuTile, el @ct.kernel marca una función de Python como un kernel de GPU. Pasamos constantes de tiempo de compilación usando anotaciones de tipo ct.Constant[T]:
import math import cuda.tile as ct # Alias de tipo para constantes de tiempo de compilación ConstInt = ct.Constant[int] ConstBool = ct.Constant[bool] # Factor de conversión: usamos exp2 en lugar de exp por eficiencia INV_LOG_2 = 1.0 / math.log(2) @ct.kernel() def fmha_kernel( Q, K, V, Out, # Tensores de entrada/salida qk_scale: float, # Factor de escala (1/sqrt(d)) input_pos: int, # Desplazamiento de posición para el enmascaramiento causal TILE_D: ConstInt, # Dimensión de cabeza (por ejemplo, 128) H: ConstInt, # Número de cabezas de atención TILE_M: ConstInt, # Tamaño del mosaico para la dimensión Q (por ejemplo, 64) TILE_N: ConstInt, # Tamaño del mosaico para la dimensión K/V (por ejemplo, 64) QUERY_GROUP_SIZE: ConstInt,# Para Atención Agrupada por Consulta (GQA) CAUSAL: ConstBool, # Si se aplica máscara causal EVEN_K: ConstBool, # Si la longitud K es divisible por TILE_N ):
Resumen: La pila de optimización
| Optimización | Perspectiva clave | Impacto |
|---|---|---|
| Base (64×64) | Correcto pero no optimizado | Base |
| Mosaicos grandes (256×128) | TRAMPA: ¡18-43% más lento! | -18% a -43% |
| + Matemáticas rápidas (FTZ, APPROX) | RESCATE: Los mosaicos grandes ahora son rentables | +34% a +72% desde la trampa |
| + División de bucles K | La mayor optimización única | +16% a +32% |
| + Reasignación de ProgramId | Mejor equilibrio de carga | +1% a +3% |
| + Autotuning | Mosaicos óptimos por secuencia | +10% a +45% |
Aceleración final: 1.60x-1.66x en todas las longitudes de secuencia.
cuTile permite a los desarrolladores expresar estas optimizaciones—mosaicos, controles de matemáticas rápidas, división de bucles, autotune—en un código Python limpio y legible mientras genera PTX altamente optimizado para GPUs NVIDIA.
Encuentra el kernel completamente optimizado en el repositorio de TileGym. Feliz programación.

