La inferencia de modelos de lenguaje grandes (LLMs) con billones de parámetros presenta un desafío fundamental en la computación distribuida: cómo gestionar la escala masiva de los modelos y el contexto de entrada dentro de las limitaciones de memoria y ancho de banda de las GPUs, manteniendo al mismo tiempo una latencia y un throughput aceptables. Tradicionalmente, NVIDIA ha dominado este espacio debido a su ecosistema de software maduro y hardware optimizado. Sin embargo, con el crecimiento exponencial del tamaño de los modelos, la capacidad de memoria HBM por GPU se ha convertido en un factor crítico, abriendo una ventana de oportunidad para arquitecturas alternativas.

Este artículo explora cómo las GPUs AMD MI355X, con su alta capacidad de VRAM, están desafiando el status quo en el despliegue de LLMs de escala "hyperscaler". La tesis central es que, a pesar de las brechas históricas en el soporte de software, la combinación de una mayor densidad de memoria HBM y un costo por GPU significativamente menor puede hacer que las plataformas AMD sean competitivas, e incluso superiores en rendimiento por dólar, para la inferencia de LLMs masivos, especialmente aquellos que exceden la capacidad de memoria de una sola GPU de la competencia.

El problema se agrava con modelos como Kimi K3, que requieren más de 1.5TB de VRAM solo para los pesos, sin contar el KV cache para contextos largos. Esto empuja los límites de las configuraciones de una sola GPU o incluso de un solo nodo, forzando la paralelización de modelos (TP) a través de múltiples nodos, lo que introduce sobrecargas de comunicación. La optimización del software para explotar eficientemente el hardware subyacente se vuelve crucial para mitigar estas sobrecargas y maximizar la utilización de los recursos.

Arquitectura del Sistema

La arquitectura de inferencia se basa en un despliegue distribuido del modelo Kimi K3, que con 2.8 billones de parámetros requiere una estrategia de paralelización de modelos (TP) debido a su tamaño. En el caso de las GPUs AMD MI355X, se utiliza una configuración TP8 (8 GPUs por nodo) que permite alojar el modelo y un KV cache de 1M tokens en un solo nodo. Para las GPUs NVIDIA B200, la menor capacidad de VRAM por GPU (192GB) obliga a una configuración TP16 que abarca dos nodos, introduciendo comunicación cross-node a través de RoCE v2.

El sistema de inferencia emplea un framework como sglang, que gestiona la carga de trabajo y la ejecución de kernels. La optimización clave se centra en dos áreas: la decodificación especulativa y la optimización del prefill. La decodificación especulativa se implementa mediante un "block-diffusion draft" externo (RadixArk’s Kimi-K3-DSpark), que genera tokens borradores para ser verificados por el modelo principal. El verifier de sglang utiliza una lógica de muestreo que, en su "dense path", requiere una operación top_k_renorm_prob para construir la distribución objetivo. La implementación inicial en ROCm carecía de un kernel específico para esta operación, lo que se resolvió con una función PyTorch que realiza un sort, masked_fill y divide.

La optimización del prefill, crítica para el "time-to-first-token" (TTFT), se abordó mejorando la ejecución del kernel de atención. Kimi K3 en ROCm inicialmente recurría a un kernel de atención Triton genérico y lento. La causa raíz fue un desajuste de forma: el kernel AITER MLA prefill, más rápido, esperaba un número de cabezas de atención (attention heads) que fuera 4, 8 o un múltiplo de 16, mientras que Kimi K3 en TP8 presentaba 12 cabezas por rank. La solución fue un simple "zero-padding" del conteo de cabezas de 12 a 16 para permitir la ejecución del kernel AITER MLA, y luego extraer las 12 cabezas reales de la salida. Estas optimizaciones se realizaron sin necesidad de desarrollar kernels personalizados, sino corrigiendo definiciones y desajustes de configuración.

Flujo de Decodificación Especulativa

  1. 1 Modelo Borrador Genera una secuencia de tokens candidatos (draft tokens).
  2. 2 Modelo Principal Evalúa los tokens candidatos en paralelo.
  3. 3 Verifier (sglang) Compara las probabilidades del modelo principal con los tokens candidatos.
  4. 4 top_k_renorm_prob Ajusta la distribución de probabilidad para el muestreo (implementado en PyTo...
  5. 5 Muestreo Selecciona el siguiente token basado en la distribución ajustada.

Flujo de Optimización de Prefill de Atención

  1. 1 Modelo Kimi K3 (TP8) Presenta 12 attention heads por rank.
  2. 2 Zero-Padding Ajusta el conteo de heads de 12 a 16 para compatibilidad con kernel.
  3. 3 Kernel AITER MLA Ejecuta la operación de prefill de atención optimizada.
  4. 4 Extracción de Heads Recupera las 12 attention heads reales de la salida del kernel.
  5. 5 Output Prefill Genera el estado de KV cache para el contexto de entrada.
CapaTecnologíaJustificación
compute AMD MI355X GPU Unidad de procesamiento principal para la inferencia de LLMs, seleccionada por su alta capacidad de VRAM (288GB) y su relación rendimiento/precio. vs NVIDIA B200 GPU, NVIDIA B300 GPU 8 GPUs por nodo (TP8) para Kimi K3.
compute ROCm Plataforma de software para programación de GPUs AMD, incluyendo compiladores, bibliotecas y herramientas de desarrollo. vs CUDA
data-processing sglang Framework de inferencia de LLMs que gestiona la ejecución del modelo, el muestreo y la decodificación especulativa.
networking RoCE v2 Protocolo de red para comunicación de baja latencia entre nodos, utilizado en configuraciones multi-nodo (ej. NVIDIA B200 TP16) para operaciones como all-reduce. ~195 Gb/s

Trade-offs

Ganancias
  • ▲▲ Rendimiento por dólar
  • Capacidad de VRAM por GPU
  • Throughput agregado por nodo (MI355X vs B200)
  • Tiempo a la primera ficha (TTFT) por optimización de prefill
Costes
  • Soporte de software y madurez del ecosistema (ROCm vs CUDA)
  • Throughput agregado absoluto (MI355X vs B300)
  • Esfuerzo de ingeniería inicial para optimización de kernels en ROCm
import torch

def top_k_renorm_prob(prob_vector, k):
    # Obtener los k valores más altos y sus índices
    top_k_values, top_k_indices = torch.topk(prob_vector, k)
    
    # Crear una máscara para los valores no top-k
    mask = torch.ones_like(prob_vector, dtype=torch.bool)
    mask[top_k_indices] = False
    
    # Poner a cero los valores no top-k
    renormed_prob_vector = prob_vector.clone()
    renormed_prob_vector[mask] = 0.0
    
    # Rescalar para que la suma sea 1
    sum_renormed = renormed_prob_vector.sum()
    if sum_renormed > 0:
        renormed_prob_vector /= sum_renormed
    
    return renormed_prob_vector
Implementación de la lógica de `top_k_renorm_prob` para el verifier de decodificación especulativa en ROCm, utilizando operaciones básicas de PyTorch en lugar de un kernel personalizado.

Fundamentos Teóricos

La problemática de la inferencia eficiente de LLMs se conecta directamente con los principios de la arquitectura de computadoras y los algoritmos de procesamiento de datos a gran escala. La necesidad de paralelización de modelos, como se observa en Kimi K3, es una manifestación de la Ley de Amdahl y la Ley de Gustafson, que dictan los límites de la aceleración que se puede lograr mediante la paralelización. La capacidad de memoria HBM es un cuello de botella crítico, un concepto bien estudiado en la jerarquía de memoria y el "memory wall" que ha sido un desafío persistente en el diseño de sistemas de alto rendimiento.

La decodificación especulativa se basa en principios de predicción y corrección de errores, análogos a las técnicas utilizadas en la predicción de ramas en arquitecturas de CPU o en la compresión de datos. Aunque no hay un paper único que la defina para LLMs, se inspira en la idea de generar una hipótesis (tokens borradores) y luego validarla o corregirla con un modelo más potente, un patrón común en la inteligencia artificial y el procesamiento de señales. La optimización de kernels de atención, como el AITER MLA, se relaciona con la investigación en algoritmos de álgebra lineal y optimización de matrices para GPUs, un campo activo desde los primeros trabajos en computación paralela y GPGPU, donde la alineación de datos y el tamaño de los bloques son cruciales para el rendimiento (ej. BLAS, cuBLAS, ROCm-hipBLAS).