Co-Diseño de atención en modelos de IA para inferencia rápida

  • El rendimiento de la atención densa está determinado por el tamaño del grupo (cabezales de consulta por cabezal KV), la dimensión del cabezal y la longitud de la secuencia, afectando de forma distinta las fases de prefill (limitada por cómputo) y decode (limitada por memoria); el análisis de intensidad aritmética y la forma GEMM revelan que la eficiencia de decode escala con el tamaño del grupo, mientras que el prefill está dominado por la longitud de la secuencia.
  • Para un rendimiento de inferencia óptimo y utilización de GPU en hardware NVIDIA, establezca un tamaño de grupo ($G$) alto para decode, utilice dimensiones de cabezal de 128 o 256 para coincidir con el tile de la GPU y la alineación de memoria, y minimice el estado KV efectivo mediante compresión de caché, atención dispersa/ventana deslizante o arquitecturas híbridas (por ejemplo, NVIDIA Nemotron 3).
  • Las estrategias de paralelismo deben estar dictadas por el conteo de cabezales KV: el paralelismo de tensores (TP) no debe exceder los cabezales KV (KH) para evitar la duplicación de KV, y los modelos con pocos cabezales KV deben emplear paralelismo de datos de atención (ADP), paralelismo KV (KVP) o enfoques híbridos (Wide EP, Helix Parallelism) según se implementa en TensorRT-LLM.

El contenido generado por IA puede resumir información de manera incompleta. Verifique la información importante. Obtenga más información.

A medida que las cargas de trabajo de agentes y contextos largos se vuelven comunes, las longitudes de contexto aumentan y la atención consume una mayor parte del tiempo de inferencia (Figura 1). Debido a que la atención ahora domina ese costo, cómo se diseña —no solo cómo se implementa— determina cada vez más el rendimiento de inferencia de un modelo. Dar forma a la arquitectura del modelo según cómo las GPUs la ejecutan es la premisa del co-diseño de modelos de IA. Para una discusión sobre cómo las opciones de diseño de modelos impactan tanto el throughput como la interactividad sin sacrificar la precisión, consulte la publicación anterior, AI Model Co-Design: Hardware-Friendly LLM Design.

Esta publicación examina cómo el tamaño del grupo (cabezales de consulta por cabezal KV), la dimensión del cabezal y la longitud de la secuencia dan forma al rendimiento de la atención densa, donde cada consulta atiende a todas las claves y valores a lo largo de la longitud de la secuencia. Destilamos ese análisis, junto con cómo se paraleliza la atención entre GPUs, en cuatro pautas prácticas: una lista de verificación de co-diseño que ayuda a los desarrolladores de modelos a aumentar el throughput de inferencia y la interactividad en GPUs NVIDIA. Manténgase atento a una publicación que cubrirá la atención dispersa.

Figura 1. Desglose del tiempo de prefill de DeepSeek-R1 a longitudes de contexto de 4K, 32K y 128K, donde la participación de la atención aumenta del 18% al 85%
Figura 1. Desglose del tiempo de prefill de DeepSeek-R1 a longitudes de contexto de 4K, 32K y 128K, donde la participación de la atención aumenta del 18% al 85%

Cada análisis se basa en dos fuentes: fórmulas analíticas de la aritmética de forma GEMM y datos medidos de kernels de prefill y decode con FP8 tanto para el cómputo de atención como para la caché KV.

¿Cómo son el prefill y decode dos problemas diferentes?

El prefill procesa el prompt completo en paralelo, produciendo grandes matmuls GEMM-M (= ISL × \(G\)) que están limitados por cómputo. Sin speculative decoding, el decode genera un token a la vez, produciendo pequeños matmuls GEMM-M (= \(G\)) y quedando limitado por memoria debido a las lecturas de caché KV desde la memoria de alto ancho de banda (HBM).

El speculative decoding aumenta GEMM-M y puede desplazar el decode hacia un estado limitado por cómputo. Debido a que las longitudes de consulta, el acceso KV y el cuello de botella difieren para prefill y decode (Tabla 2), cada parámetro se analiza por separado para cada fase.

Nota: Con el almacenamiento en caché de prefijos, común en aplicaciones de agentes y de múltiples turnos, un nuevo turno puede tener un ISL corto mientras atiende a un gran caché de prefijo. Con un ISL corto pero un caché de prefijo largo, el prefill se comporta como el decode.

¿Cómo gobierna la intensidad aritmética el comportamiento limitado por cómputo versus memoria?

El modelo roofline limita el rendimiento de la GPU mediante techos de cómputo y ancho de banda, como se explicó anteriormente. La intensidad aritmética determina qué limita el rendimiento (Ecuación 1):

Intensidad Aritmética = FLOPs Totales / Bytes Totales accedidos

El punto de cresta marca la transición de limitado por memoria a limitado por cómputo. El prefill se encuentra muy por encima y está limitado por cómputo, mientras que el decode se encuentra por debajo y está limitado por memoria (Figura 2). El speculative decoding aumenta la intensidad aritmética del decode y puede moverlo hacia la cresta.

Figura 2. Modelo roofline que muestra el decode en la rampa limitada por memoria y el prefill en la meseta limitada por cómputo
Figura 2. Modelo roofline que muestra el decode en la rampa limitada por memoria y el prefill en la meseta limitada por cómputo

¿Cómo calcula el kernel FlashAttention la atención en la GPU?

FlashAttention calcula la atención sin materializar la matriz de atención completa. Transmite tiles de \(Q\), \(K\) y \(V\) desde HBM a SRAM en el chip y fusiona tres pasos en una sola pasada:

  • Primero, matmul por lotes (BMM1) puntúa consultas contra claves
  • Segundo, softmax en línea normaliza las puntuaciones usando un máximo y una suma en ejecución
  • Tercero, segundo matmul por lotes (BMM2) pondera los valores

Los BMM se ejecutan en Tensor Cores mientras que los exponenciales de softmax se ejecutan en unidades de función especial. Las formas de BMM impulsan el análisis de intensidad aritmética que sigue.

Figura 3. Kernel FlashAttention con BMM1 y BMM2 en Tensor Cores y softmax en línea fusionado en unidades de función especial. Imagen adaptada de FlashAttention: Fast and Memory-Efficient Exact Attention with
Figura 3. Kernel FlashAttention con BMM1 y BMM2 en Tensor Cores y softmax en línea fusionado en unidades de función especial. Imagen adaptada de FlashAttention: Fast and Memory-Efficient Exact Attention with

Formas GEMM

El rendimiento de la atención sigue las formas de sus dos matmuls. La Tabla 3 enumera las dimensiones por fase (Batch, M, N, K) de BMM1 y BMM2.

Para el decode, GEMM-M = \(G\), generalmente 8-16, muy por debajo de un tile-M de GPU de 64 o 128, lo que limita el trabajo paralelo por tile. Un \(G\) más grande carga menos KV por token y amortiza cada carga en más cabezales de consulta, mejorando la utilización. La siguiente sección cuantifica este efecto.

Tamaño del grupo

El tamaño del grupo (\(G\)) es el número de cabezales de consulta que comparten un cabezal KV. MHA tiene \(G\) = 1, GQA tiene \(G\) = 4, 8, 16, …, y MQA tiene \(G\) = QH.

Intensidad aritmética en función de \(G\). En las siguientes fórmulas, “Bytes” se refiere a los bytes de HBM movidos. Para simplificar, suponga 1 byte por elemento (es decir, caché KV FP8).

Prefill

A medida que \(G\) crece, el 1/\(G\) desaparece y la intensidad aritmética se acerca a 2 × ISL. A ISL = 32K, aumentar \(G\) de 8 a 16 mejora la intensidad aritmética en menos del 6%. En otras palabras, el prefill está dominado por ISL, no por \(G\). La Figura 4 confirma esto, al variar \(G\) de 1 (MHA) a 64 (MQA) el tiempo de ejecución del prefill cambia en menos del 1%. Ecuaciones 2, 3 y 4:

FLOPs = 4 × PB × QH × ISL² × Hsz (constante en \(G\)) Bytes = 2 × PB × KH × Hsz × ISL × (\(G\) + 1) Intensidad Aritmética = 2 × \(G\) × ISL / (\(G\) + 1) = 2 × ISL / (1 + 1/\(G\)) → 2 × ISL cuando \(G\) → ∞

Decode (GEMM-M = \(G\))

Duplicar \(G\) duplica la intensidad aritmética del decode. Aumentar \(G\) de 1 a 8 proporciona una ganancia de 8x al reducir el tráfico de memoria y mejorar la utilización del cómputo de la GPU. Es independiente de KVSL: la intensidad aritmética se mantiene cerca de 2 × \(G\), por lo que el decode permanece limitado por memoria a menos que \(G\) sea muy grande. Modelos como NVIDIA Nemotron 3 adoptaron GQA con dos cabezales KV, lo que hace que el decode sea más eficiente. Ecuaciones 5, 6 y 7:

FLOPs = 4 × DB × QH × KVSL × Hsz (constante en \(G\)) Bytes = 2 × DB × KH × Hsz × (\(G\) + KVSL) Intensidad Aritmética = 2 × \(G\) × KVSL / (\(G\) + KVSL) ≈ 2 × \(G\) (cuando KVSL ≫ \(G\))

La Figura 4 muestra que el tiempo de ejecución del decode cae aproximadamente 2x por cada duplicación de \(G\) porque reducir a la mitad los cabezales KV reduce a la mitad los datos cargados por token. Más allá de \(G\) = 16, la curva KVSL = 32K se aplana. Su kernel por paso es lo suficientemente pequeño como para que dominen dos costos: la configuración fija y la sobrecarga de posprocesamiento, y la reducción de flash-decoding al dividir KV entre los SM para mantenerse paralelo con pocos cabezales KV.

Vía NVIDIA Developer.