Qué cambia FlashAttention en la inferencia de transformers

Traducción automática

Este artículo se tradujo automáticamente a partir de la versión original en inglés.

FlashAttention calcula la atención exacta sin almacenar las matrices completas de puntuaciones y probabilidades de atención en la memoria principal de la GPU. Procesa bloques más pequeños y combina sus resultados mediante un softmax en línea. Esto reduce el tráfico de memoria y el almacenamiento intermedio, pero mantiene el coste aritmético cuadrático de la atención densa respecto a la longitud de la secuencia.

Un menor uso de memoria no hace que el cálculo de la atención densa crezca de forma lineal.

Separar el almacenamiento del cálculo

La atención materializada estándar calcula puntuaciones para cada par visible de consulta y clave, normaliza esas puntuaciones y las multiplica por los valores. Una matriz de atención cuadrada sin máscara contiene N² entradas para N tokens.

El artículo original de FlashAttention reorganiza ese cálculo en bloques que caben en el chip. Mantiene estadísticas de normalización y acumuladores de salida que actualiza durante el cálculo, en lugar de escribir todas las probabilidades en HBM. Con dimensiones de cabeza fijas, el almacenamiento necesario para los resultados intermedios de atención crece linealmente con la longitud de la secuencia. Sigue evaluando las interacciones densas entre consultas y claves que exige la máscara de atención.

El algoritmo calcula atención exacta; distintos órdenes de ejecución de las operaciones en coma flotante pueden causar pequeñas diferencias numéricas. No es un método aproximado de atención dispersa que descarta determinados pares de tokens.

La implementación y la carga de trabajo siguen importando

PreguntaPor qué comprobarlo
¿El backend admite esta GPU y esta precisión?Las implementaciones de los kernels están dirigidas a hardware concreto
¿Se admiten la dimensión de las cabezas, la máscara y la disposición de la caché?Las formas no compatibles pueden hacer que se seleccione otro backend
¿La atención supone una parte importante del tiempo de la petición?Una gran mejora en el kernel puede tener un efecto pequeño en el conjunto
¿La carga de trabajo es prefill o decodificación de un solo token?El trabajo que pueden ejecutar en paralelo es distinto

FlashAttention-2 también paraleliza bloques de la secuencia de consultas, además de lotes y cabezas, y reduce otros costes de ejecución. FlashAttention-3 se centra en funciones de ejecución específicas de Hopper. Los nombres de versión describen algoritmos e implementaciones; no demuestran qué backend ha seleccionado realmente el servidor.

Comprueba el backend seleccionado en el runtime de serving con la versión fijada. Mide la latencia del prompt, el pico de memoria y el rendimiento de extremo a extremo con longitudes y tamaños de lote representativos. Compara máscaras y precisión idénticas. Un benchmark que solo informa de los TFLOPS del kernel no demuestra una mejora en la latencia por token que percibe el usuario.

Guía de ingeniería: FlashAttention trata las sucesivas versiones y sus mediciones dentro de sus respectivos ámbitos.