IA 360
Actualidad

Cómo leer una traza de atención en PyTorch antes de optimizarla

El profiler revela kernels, copias y rutas de SDPA que no se ven en el código. La clave es medir el caso propio antes de cambiar backend.

4 min de lectura Generado con IA Read in English
Cómo leer una traza de atención en PyTorch antes de optimizarla

El 10 de julio de 2026, Hugging Face publicó la tercera parte de su tutorial de perfilado en PyTorch, dedicada a leer trazas de atención antes de optimizar. Cuando un modelo Transformer se vuelve lento, la atención suele aparecer pronto en la conversación: combina multiplicaciones de matrices, máscaras, softmax y memoria intermedia; con secuencias largas, todo eso crece rápido. La lección es mirar la traza antes de sustituir unas líneas por una función “optimizada”.

Perfilar no consiste en adivinar qué parte del programa parece cara. Consiste en observar qué operadores y qué kernels se ejecutan realmente, cuánto tiempo consumen y qué movimiento de memoria los acompaña. PyTorch ofrece torch.profiler para registrar actividades de CPU y CUDA, formas de entrada y consumo de memoria. La tabla ayuda a localizar candidatos; la traza explica el orden y las dependencias.

La atención que parece sencilla

Una implementación directa de atención causal calcula los productos entre queries y keys, escala el resultado, aplica una máscara, ejecuta softmax y vuelve a multiplicar por los values. En una traza, esas operaciones no son una abstracción: se convierten en varios kernels de GPU y, a veces, en copias que no estaban en el pseudocódigo.

El tutorial de Hugging Face muestra un ejemplo concreto en una NVIDIA A100. La versión directa lanzaba una copia de memoria junto a los cálculos esperados. Al cambiar masked_fill por su variante in-place, esa copia desaparecía y el forward pasaba de seis a cinco kernels. Es una observación útil, pero no una receta universal: las operaciones in-place pueden sobrescribir valores que autograd necesita para el backward. En el experimento eran seguras porque se ejecutaba inferencia bajo torch.no_grad; ese detalle cambia por completo la decisión en entrenamiento.

Una función corta no siempre implica menos trabajo

PyTorch empaqueta esta operación en torch.nn.functional.scaled_dot_product_attention, o SDPA. Para tensores CUDA puede elegir entre una implementación matemática, FlashAttention, una variante memory-efficient y otras rutas compatibles. La interfaz es cómoda, pero el backend seleccionado depende del tipo de dato, las dimensiones, la máscara y el hardware.

En las pruebas del artículo, fijar el backend matemático de SDPA produjo una sorpresa: lanzó 20 kernels por forward frente a los cinco de la versión ingenua in-place y fue alrededor de 3,7 veces más lento. La razón no es que SDPA sea mala, sino que esa ruta prioriza ser una referencia segura y general. En ese caso concreto materializaba una máscara, usaba un softmax protegido y elevaba operaciones a precisión FP32, con más trabajo y tráfico de memoria.

Las rutas efficient, flash y cuDNN se veían de otra manera: un kernel fusionado por forward. La fusión evita escribir toda la matriz de atención de tamaño secuencia por secuencia en la memoria principal de la GPU y reduce lanzamientos. En la A100 del experimento, Flash fue la ruta CUDA más rápida para esa forma: 146,8 microsegundos de media en GPU, frente a 277,9 para efficient y 186,3 para cuDNN. Pero cuDNN dedicó más tiempo en CPU a preparar su plan, y puede ganar en otras dimensiones. Una traza limpia tampoco demuestra que el trabajo haya desaparecido: a veces sólo se ha desplazado dentro de una biblioteca.

Un método antes que un truco

Para perfilar atención con criterio conviene seguir una secuencia corta. Primero, perfilar un caso representativo tras el calentamiento, separando CPU y CUDA. Segundo, ordenar por tiempo propio y total: el tiempo propio excluye las llamadas hijas, mientras que el total las incluye. Tercero, abrir la traza y buscar kernels repetidos, copias, máscaras reconstruidas o huecos entre CPU y GPU.

Después se pueden comparar backends de SDPA con las mismas entradas, medir latencia y memoria, y verificar que la salida mantiene la precisión necesaria. La documentación de PyTorch permite forzar temporalmente un backend precisamente para hacer esa comparación. Si una opción no está soportada para la forma o el hardware, el aviso también es información útil.

La optimización de atención no se reduce a pronunciar “FlashAttention”. Se trata de aprender a leer qué hace el programa en una GPU concreta y de conservar la corrección al mejorarlo. El perfil no entrega una respuesta automática, pero evita que una mejora elegante en el código empeore el trabajo real.

Fuentes de esta pieza

Esta pieza se apoya en 4 fuente(s) primaria(s), recogidas durante la investigación.

Este artículo se ha elaborado con inteligencia artificial bajo supervisión editorial humana.

Compartir este artículo

Este sitio web utiliza cookies para mejorar la experiencia de navegación. Política de cookies.

↑↓ navegar ↵ abrir esc cerrar