Diseño de kernels de GPU de alto rendimiento con TileLang: GEMM con Tensor Core, Softmax fusionado, FlashAttention y ajuste automático
Puntos clave
- •TileLang permite a los desarrolladores crear kernels de GPU optimizados a nivel de tile en Python mientras el compilador gestiona automáticamente el mapeo de hilos, las distribuciones de memoria, la sincronización y la generación de instrucciones CUDA.
- •El tutorial implementa y valida progresivamente kernels para suma de vectores, multiplicación de matrices con tensor core, epílogos de GEMM fusionados con sesgo y GELU, softmax por filas y FlashAttention frente a líneas base de PyTorch.
- •El decorador de ajuste automático de TileLang busca entre tamaños de tile, profundidades de pipeline y recuentos de hilos para identificar automáticamente configuraciones óptimas dependientes de la arquitectura, abordando el desafío de que los parámetros ideales varían entre generaciones de GPU como Ampere y Hopper.
- •Las técnicas de fusión de kernels demostradas en el tutorial reducen el tráfico de memoria global intermedio al completar operaciones como suma de sesgo y activación GELU dentro de acumuladores residentes en registros antes de escribir las salidas finales.
- •La implementación de FlashAttention procesa los tiles de query, key y value sin materializar la matriz completa de puntuaciones de atención, reduciendo la huella de memoria de un espacio cuadrático a uno lineal mediante actualizaciones de softmax en línea.
- •La inspección del código fuente generado, la impresión del lado del dispositivo y un generador de perfiles integrado proporcionan a los desarrolladores visibilidad sobre las operaciones de tensor core emitidas por el compilador, las copias asíncronas y las barreras de sincronización para depuración y optimización.

TileLang es un lenguaje de dominio específico (DSL) de Python de alto nivel para diseñar y compilar kernels de GPU orientados al rendimiento a través de TVM. A medida que los grandes modelos de lenguaje y otras arquitecturas transformer aumentan las demandas computacionales, la capacidad de crear kernels de GPU personalizados que aprovechen al máximo los tensor cores —los aceleradores dedicados de multiplicación matricial de NVIDIA— se ha vuelto cada vez más crítica para la eficiencia en entrenamiento e inferencia. Sin embargo, escribir dichos kernels tradicionalmente requiere un conocimiento profundo de CUDA C++, programación a nivel de warp y jerarquías de memoria específicas del hardware. TileLang aborda esta brecha al permitir que los desarrolladores expresen los cálculos a nivel de tile mientras el compilador gestiona los detalles de bajo nivel. Este tutorial ofrece un recorrido integral de las capacidades de TileLang, comenzando con la validación del entorno y el establecimiento de utilidades reutilizables de benchmarking y verificación numérica. A partir de ahí, implementa progresivamente suma de vectores, multiplicación de matrices en tiles con tensor core, exploración de esquemas, epílogos de GEMM fusionados, softmax por filas y FlashAttention.
A lo largo del tutorial, los desarrolladores trabajan directamente con los tiles de memoria compartida, fragmentos de registro, bucles en pipeline, primitivas de iteración paralela, reducciones y operadores de GEMM con tensor core de TileLang. El compilador gestiona automáticamente el mapeo de hilos, las distribuciones de memoria, la sincronización, la vectorización y la generación de instrucciones CUDA de bajo nivel. Los kernels se comparan con líneas base de PyTorch y cuBLAS, inspeccionando el código fuente CUDA generado, evaluando el rendimiento de memoria y cómputo, y utilizando ajuste automático para identificar configuraciones de kernel dependientes de la arquitectura. TileLang está disponible en GitHub.
El tutorial comienza configurando el entorno CUDA de Google Colab, instalando TileLang con una opción de respaldo nightly, e importando los módulos necesarios de PyTorch y TileLang. Se definen utilidades reutilizables de benchmarking, validación y reporte para medir la latencia del kernel y comparar las salidas numéricas utilizando el error relativo. Luego, se implementa un kernel de suma de vectores en TileLang, se ejecuta en la GPU, se compara su ancho de banda con el de PyTorch y se inspecciona el código fuente CUDA generado por el compilador.
A continuación, se implementa un kernel de multiplicación de matrices en tiles con tensor core que mueve las entradas en tiles a través de memoria global, memoria compartida y fragmentos de registro. Las dimensiones de los tiles, las etapas de pipeline, el número de hilos y el swizzling de L2 se controlan manualmente, mientras que TileLang genera las instrucciones de tensor core, la sincronización y la lógica de transferencia de memoria. Se evalúan varias configuraciones de esquema, se verifica su precisión numérica y se identifica la configuración de kernel dependiente de la arquitectura con mejor rendimiento. El código completo del tutorial está disponible aquí.
El kernel de multiplicación de matrices se amplía fusionando la suma de sesgo y la activación GELU directamente en el acumulador residente en registros. Este enfoque reduce el tráfico de memoria global intermedio al completar el epílogo antes de escribir el tensor de salida final. La fusión de kernels es una técnica bien establecida para reducir los cuellos de botella en el ancho de banda de memoria, que a menudo dominan la latencia en cargas de trabajo de redes neuronales a gran escala. La implementación fusionada se compara con la ejecución eager de PyTorch. Además, se implementa un kernel de softmax por filas utilizando reducciones de máximo y suma a nivel de fragmento, manteniendo el proceso de normalización en gran medida dentro de los registros.
Se implementa un kernel forward de FlashAttention fusionado que procesa los tiles de query, key y value sin materializar la matriz completa de puntuaciones de atención en la memoria global. Se aplican actualizaciones de softmax en línea utilizando máximos acumulados, sumas de normalización, factores de reescalado y multiplicaciones matriciales en tiles con tensor core. FlashAttention, originalmente introducido por investigadores de Stanford, se ha convertido en una técnica ampliamente adoptada para reducir la huella de memoria del self-attention de un espacio cuadrático a uno lineal, y ahora está integrada en los principales frameworks, incluida la API de scaled dot-product attention de PyTorch. Se validan tanto la atención causal como la no causal frente a la scaled dot-product attention de PyTorch, comparando latencia y rendimiento computacional.
Se define un espacio de búsqueda de ajuste automático que abarca tamaños de tile matricial, dimensiones de bloque K, profundidades de pipeline y recuentos de hilos, filtrando las configuraciones que exceden el presupuesto de memoria compartida. Se utiliza el decorador de ajuste automático de TileLang para compilar, evaluar, validar y almacenar en caché múltiples esquemas de kernel para la misma carga de trabajo de multiplicación de matrices. El kernel seleccionado se ejecuta, se verifica frente a PyTorch y se evalúa en términos de latencia alcanzada y rendimiento del tensor core. Esta búsqueda automatizada aborda un desafío práctico en el desarrollo de kernels de GPU: los tamaños de tile y las profundidades de pipeline óptimos varían entre arquitecturas de GPU como Ampere y Hopper, lo que hace que el ajuste manual sea frágil cuando se apunta a múltiples generaciones de hardware.
Se presenta el flujo de trabajo de depuración e introspección de TileLang a través de impresión del lado del dispositivo, inspección del código CUDA generado y el generador de perfiles de kernel integrado. Se examinan los puntos de referencia emitidos por el compilador, como operaciones de tensor core, copias asíncronas, barreras de sincronización e instrucciones de carga matricial. Todas las secciones del tutorial se organizan en un ejecutor tolerante a fallos que registra el estado de ejecución, reporta información de temporización e imprime una referencia compacta de programación en TileLang.
El tutorial demuestra cómo TileLang traduce programas de Python a nivel de tile en kernels de GPU optimizados sin requerir la gestión manual de índices de hilos, distribuciones de datos a nivel de warp, instrucciones de tensor core ni barreras de memoria asíncrona. Los kernels implementados y validados cubren operaciones elementales limitadas por ancho de banda, cargas de trabajo de GEMM intensivas en cómputo, epílogos de redes neuronales fusionados, reducciones residentes en registros y atención con softmax en línea. La exploración también muestra cómo las dimensiones de bloque, el consumo de memoria compartida, la profundidad de pipeline, el recuento de hilos, la forma de los tiles y el swizzling de L2 influyen en el rendimiento a través de diferentes arquitecturas de GPU. La inspección del código fuente generado, la depuración del lado del dispositivo, la generación de perfiles y la búsqueda automatizada de esquemas establecen conjuntamente un flujo de trabajo completo para desarrollar, verificar, evaluar y refinar kernels personalizados de TileLang.