Desarrollo de Jagged Flash Attention para el modelo GEM
Meta ha implementado Jagged Flash Attention (JFA) como el núcleo operativo de su Generative Ads Model (GEM), un sistema especializado en la generación de publicidad. El desafío técnico principal reside en que los modelos de anuncios de Meta procesan secuencias de usuarios con longitudes variables, conocidas como secuencias jagged o irregulares. Tradicionalmente, el procesamiento de estas secuencias requería el uso de padding, que consiste en añadir tokens vacíos para ajustar todas las secuencias a una longitud fija, lo que resultaba en un desperdicio de hasta el 50 % de la capacidad de cómputo.
Para solucionar esta ineficiencia, el kernel JFA aplica el algoritmo de FlashAttention directamente sobre tensores Q, K y V empaquetados, utilizando un tensor de offsets para registrar los límites de cada secuencia. Este enfoque elimina la necesidad de materializar tokens de relleno, optimizando el uso de la memoria y el ciclo de procesamiento. El objetivo final de este desarrollo es maximizar el rendimiento en la arquitectura NVIDIA Blackwell B200, donde la atención se identifica como el kernel más lento de todo el sistema GEM.
El rol de TLX en la optimización de kernels
La optimización de JFA se ha llevado a cabo utilizando TLX (Triton Low-level Extensions). TLX es una capa de extensiones que añade un control explícito y consciente del hardware sobre el modelo de programación basado en tiles de Triton. Hasta ahora, alcanzar el rendimiento máximo en la arquitectura Blackwell requería la escritura manual de kernels en CUDA o CuteDSL, lenguajes extremadamente complejos de desarrollar y difíciles de adaptar cuando se introducen nuevas variantes de modelado.
TLX permite cerrar la brecha entre la facilidad de desarrollo de Triton y el rendimiento bruto de CUDA. El kernel de atención desarrollado con TLX consta de aproximadamente 3200 líneas de código, lo que representa una reducción de tres veces en la cantidad de código comparado con los kernels de FlashAttention-4 (FA4) basados en CuteDSL, que alcanzan las 10 000 líneas. Esta simplificación permite que los ingenieros de modelado, y no solo los especialistas en kernels, puedan leer, extender y fusionar el código de manera eficiente.
Mejoras de rendimiento frente a FlashAttention-4
En las pruebas de rendimiento realizadas con el formato bfloat16 sobre hardware B200, el kernel de Meta ha superado la versión de mayo de 2026 de FlashAttention-4 en las formas jagged específicas que requiere el modelo GEM. Los resultados muestran una mejora del 13 % en el forward pass y un incremento del 50 % en el backward pass. Estas cifras demuestran que el control explícito sobre el hardware permite una ejecución más fluida de las operaciones de atención.
El rendimiento en Blackwell depende críticamente de mantener los tensor cores alimentados de forma continua. Mientras que el compilador de Triton estándar gestiona la mayoría de las decisiones de movimiento de datos, TLX expone primitivas de primer nivel. Esto incluye la asignación explícita de memoria compartida (SMEM) y memoria de tensor (TMEM), la especialización de warps mediante tareas asíncronas, el uso de barreras y la implementación de TMA (Tensor Memory Accelerator) y MMA (Matrix Multiply-Accumulate) asíncronos.
Cambios estructurales y gestión de memoria
La reorganización del kernel mediante TLX ha permitido dividir la CTA (Cooperative Thread Array) en tareas asíncronas con roles especializados. Se han asignado warps específicos para las cargas de TMA, otros para las multiplicaciones de los tensor cores, otros para los cálculos de softmax y corrección, y un grupo dedicado al almacenamiento del epílogo. En el proceso de backward, se ha incluido un warp especializado en la reducción de dQ.
Esta especialización evita que los tensor cores se detengan mientras se ejecutan las operaciones de softmax, permitiendo que ambas tareas ocurran concurrentemente. Además, Meta ha implementado una gestión manual de los buffers en chip, definiendo la profundidad de la tubería. Por ejemplo, se utiliza un triple buffering para K y V, permitiendo que el warp de carga se adelante al de procesamiento. Asimismo, se han aliasado los buffers de TMEM con vidas útiles no solapadas, compartiendo una misma asignación para los scores QK, la matriz P y las estadísticas de softmax, mientras que el acumulador PV mantiene su propia asignación.
Optimizaciones específicas para el flujo de trabajo de GEM
Meta ha introducido optimizaciones diseñadas para eliminar cuellos de botella detectados mediante NVIDIA Nsight Compute y TritonBench. Una de las innovaciones más relevantes es la programación de tiles jagged a través de los SM (Streaming Multiprocessors). Debido a que las secuencias varían drásticamente en longitud, una asignación ingenua de tiles provocaría que algunos SM quedaran inactivos mientras otros procesaban secuencias largas.
Para resolver este desequilibrio de carga en el forward pass, Meta implementó un sistema de balanceo de carga en el host. El proceso consiste en ordenar todos los tiles por carga de trabajo de KV de forma descendente y distribuirlos entre los SM siguiendo un patrón zigzag o boustrophedon. En las pasadas pares, los SM se llenan de izquierda a derecha, y en las impares, de derecha a izquierda. Esta técnica ha recuperado aproximadamente un 20 % del rendimiento del kernel de forward pass con un coste insignificante para la CPU.
Optimización del proceso de Backward y dQ
El proceso de backward presenta desafíos adicionales, especialmente cuando se utiliza la configuración broadcast-Q, donde una única Q densa se distribuye entre cada secuencia del lote. Esto convierte el epílogo de dQ en una reducción sumatoria altamente disputada. Para mitigar esto, Meta desarrolló una tubería de staging de SMEM de doble buffer que oculta la latencia de la reducción de dQ, eliminando el principal cuello de botella del backward pass.
Otras mejoras incluyen la liberación temprana de la memoria de tensor dQ, lo que permite que el warp de MMA comience la multiplicación del siguiente tile más rápidamente. También se ha aplicado el loop peeling, que consiste en dividir el bucle KV en una pasada masiva sin ramificaciones y una cola pequeña enmascarada, eliminando la sobrecarga de las máscaras por iteración y evitando el spill de registros. Finalmente, se ha adoptado la técnica de MMA colaborativa entre dos CTAs, proveniente de FA4, para elevar la utilización de los tensor cores en las fases más intensivas de multiplicación de matrices.
Impacto en el desarrollo de modelos generativos
La capacidad de implementar kernels de alto rendimiento sin depender exclusivamente de CUDA o CuteDSL reduce drásticamente el ciclo de iteración para los ingenieros de IA. Dado que la atención es el lugar donde se implementan la mayoría de las innovaciones de modelado, como las ventanas deslizantes o el sparse-block, contar con un kernel extensible en Python mediante TLX acelera la experimentación.
La transición hacia arquitecturas como NVIDIA Blackwell exige una gestión mucho más granular del hardware para evitar que la potencia de cómputo se desperdicie en esperas de memoria. El enfoque de Meta con JFA y TLX establece un precedente sobre cómo optimizar modelos para datos irregulares, asegurando que la eficiencia del hardware se traduzca directamente en una mayor velocidad de entrenamiento e inferencia para modelos de gran escala como GEM.




