La creciente escala de los modelos de lenguaje grandes (LLM), con modelos como Kimi K3 superando los 2.8 billones de parámetros, presenta un desafío fundamental en la computación distribuida: cómo servir estos modelos de manera eficiente y económica. La memoria de alto ancho de banda (HBM) se ha convertido en un cuello de botella crítico, ya que el tamaño de los pesos del modelo y el caché KV para contextos extensos exceden la capacidad de VRAM de las GPUs convencionales. Esto obliga a la partición del modelo a través de múltiples dispositivos o nodos, introduciendo latencia de comunicación y complejidad.
Históricamente, NVIDIA ha dominado el mercado de GPUs para IA, estableciendo un ecosistema de software maduro (CUDA). Sin embargo, la emergencia de GPUs con mayor capacidad de HBM por parte de competidores como AMD, específicamente la MI355X con 288GB de VRAM, plantea la pregunta de si la "moat" de software de NVIDIA puede ser superada por una ventaja de hardware en capacidad de memoria. Este artículo examina si la capacidad de HBM de AMD puede traducirse en una ventaja de costo-rendimiento para la inferencia de LLM a escala de hyperscaler, incluso con las barreras de software existentes.
Arquitectura del Sistema
La arquitectura de inferencia para LLMs de escala trillonaria se basa en la partición del modelo (Tensor Parallelism, TP) a través de múltiples GPUs, a menudo distribuidas en varios nodos. Para Kimi K3 (2.8T parámetros), se requiere más de 1.5TB de VRAM solo para los pesos, sin contar el KV cache para contextos de 1M tokens. Esto implica configuraciones como TP16 en nodos B200 (8 GPUs de 192GB cada una, requiriendo 2 nodos) o TP8 en nodos MI355X/B300 (8 GPUs de 288GB cada una, cabiendo en un solo nodo). La elección de la GPU y la configuración de TP impacta directamente la latencia de comunicación inter-GPU/inter-nodo, especialmente en operaciones como all-reduce en el path crítico de decodificación.
Para optimizar el rendimiento, se emplean dos técnicas clave: decodificación especulativa y optimizaciones de prefill. La decodificación especulativa acelera la generación de tokens prediciendo secuencias futuras con un modelo borrador más pequeño y verificándolas con el modelo grande. En este caso, se utiliza un borrador externo (RadixArk’s Kimi-K3-DSpark). La implementación de esta técnica en ROCm requirió la corrección de una NameError en el verificador sglang debido a la ausencia de una implementación de top_k_renorm_prob para GPUs AMD, resuelta con una función PyTorch que emula la lógica de top-k (sort, masked_fill, divide).
Las optimizaciones de prefill se centran en reducir el tiempo hasta el primer token (TTFT). Para Kimi K3 en ROCm, se identificó que el kernel de atención AITER MLA, optimizado para prefill, no se cargaba debido a un shape mismatch (12 cabezas de atención por rank en TP8 vs. 4, 8 o múltiplos de 16 esperados por el kernel). La solución fue un simple zero-padding de las cabezas de atención de 12 a 16 para permitir el uso del kernel rápido, mejorando el prefill de 2 a 3 veces. Estas optimizaciones demuestran que, a menudo, los cuellos de botella en el software AMD no son la ausencia de kernels, sino problemas de compatibilidad o definiciones faltantes que pueden resolverse con adaptaciones de código de alto nivel.
Flujo de Decodificación Especulativa (ROCm)
- 1 Generar Borrador Modelo borrador externo (Kimi-K3-DSpark) genera tokens especulativos.
- 2 Verificar con Modelo Grande El modelo Kimi K3 verifica los tokens propuestos.
- 3 SGLang Verifier Componente de sglang para aceptar/rechazar tokens.
- 4 Path Denso (ROCm) Intenta usar `top_k_renorm_prob` para construir distribución objetivo.
- 5 Error `NameError` Falla por definición faltante de `top_k_renorm_prob` en ROCm.
- 6 Aplicar Fix PyTorch Implementación manual de `top-k renorm` (sort, masked_fill, divide).
- 7 Generar Siguiente Token El verificador acepta tokens y el modelo genera el siguiente.
Flujo de Optimización de Prefill (ROCm)
- 1 Solicitud de Prefill Inicio de una solicitud de prefill de contexto largo (ej. 172k tokens).
- 2 Carga de Kernel de Atención Intento de cargar el kernel AITER MLA para prefill rápido.
- 3 Shape Mismatch Falla por 12 cabezas de atención vs. 4/8/16 esperadas por AITER.
- 4 Fallback a Triton Genérico Uso de kernel de atención Triton más lento.
- 5 Aplicar Fix Zero-Padding Zero-padding de cabezas de atención de 12 a 16.
- 6 Carga Exitosa AITER MLA El kernel AITER MLA se carga y ejecuta correctamente.
- 7 Prefill Acelerado Procesamiento de prefill 2-3x más rápido, reduciendo TTFT.
| Capa | Tecnología | Justificación |
|---|---|---|
| compute | AMD MI355X GPU | Hardware de inferencia principal, elegido por su alta capacidad de HBM (288GB) y mejor relación rendimiento/dólar para modelos de gran escala. vs NVIDIA B200 GPU, NVIDIA B300 GPU 8x MI355X en configuración TP8 (Tensor Parallelism de 8) |
| compute | NVIDIA B200 GPU | Hardware de inferencia alternativo, con menor capacidad de HBM (192GB) por GPU, requiriendo configuración TP16 (2 nodos) para Kimi K3. 2x8 B200 en configuración TP16 (Tensor Parallelism de 16) |
| compute | NVIDIA B300 GPU | Hardware de inferencia alternativo, con alta capacidad de HBM (288GB) por GPU, pero con mayor costo por unidad de rendimiento. 8x B300 en configuración TP8+DCP8 |
| orchestration | ROCm | Plataforma de software de AMD para computación GPU, utilizada para ejecutar los modelos en hardware AMD. Requiere adaptaciones para compatibilidad con frameworks de inferencia. vs CUDA |
| data-processing | sglang | Framework de inferencia para LLMs, utilizado para la decodificación especulativa y la gestión de la generación de tokens. Necesitó parches para ROCm. |
| data-processing | PyTorch | Framework de aprendizaje profundo, utilizado para implementar las correcciones de software (ej. `top-k renorm`) y para la ejecución de kernels. |
| data-processing | RadixArk’s Kimi-K3-DSpark | Modelo borrador externo utilizado para la decodificación especulativa de Kimi K3. vs MTP (Medusa/Tree Attention), EAGLE (Efficient and Accurate Group-wise Language Model Evaluation) |
Trade-offs
Ganancias
- ▲ Rendimiento por dólar (MI355X vs B300)
- ▲ Throughput agregado por nodo (MI355X vs B200)
- △ Throughput single-stream (MI355X vs B200)
- ▲ Tiempo hasta el primer token (TTFT) en MI355X
Costes
- ▲ Soporte de software y madurez del ecosistema (ROCm vs CUDA)
- △ Throughput agregado (MI355X vs B300)
- ▲ Esfuerzo de ingeniería para optimización de kernels en ROCm
```python
def top_k_renorm_prob_rocm(prob_vector, k):
# Sort values to find top-k
sorted_probs, indices = torch.sort(prob_vector, descending=True)
# Create a mask for top-k elements
mask = torch.zeros_like(prob_vector, dtype=torch.bool)
mask[indices[:k]] = True
# Zero out non-top-k elements
renormed_prob = prob_vector.masked_fill(~mask, 0.0)
# Rescale to sum to 1
sum_renormed = renormed_prob.sum()
if sum_renormed > 0:
renormed_prob /= sum_renormed
return renormed_prob
# In sglang's ROCm sampling branch:
# Replace call to undefined top_k_renorm_prob with top_k_renorm_prob_rocm
``````python
def pad_attention_heads(input_tensor, original_heads=12, target_heads=16):
# Assuming input_tensor has shape (batch, seq_len, original_heads, head_dim)
batch, seq_len, _, head_dim = input_tensor.shape
if original_heads == target_heads:
return input_tensor
# Create a new tensor with target_heads and copy original data
padded_tensor = torch.zeros(
batch, seq_len, target_heads, head_dim,
device=input_tensor.device, dtype=input_tensor.dtype
)
padded_tensor[:, :, :original_heads, :] = input_tensor
return padded_tensor
# Before calling AITER MLA kernel:
# attention_output = AITER_MLA_kernel(pad_attention_heads(input_attention_heads))
# Extract original heads from output: attention_output[:, :, :original_heads, :]
```Fundamentos Teóricos
El problema de la inferencia eficiente de modelos de gran escala se relaciona con los principios fundamentales de la arquitectura de computadoras y los sistemas distribuidos. La necesidad de particionar modelos a través de múltiples dispositivos y nodos evoca los desafíos de la comunicación inter-proceso y la coherencia de memoria en sistemas distribuidos, donde la latencia de red y el ancho de banda son factores críticos.
La decodificación especulativa se basa en el concepto de predicción y verificación, un patrón común en la optimización de pipelines de CPU (ej. branch prediction) y en algoritmos de búsqueda heurística. La optimización del prefill, por otro lado, aborda el problema del "cold start" o la inicialización de estado, un desafío recurrente en sistemas de bases de datos (ej. warm-up de caches, carga inicial de datos en un LSM-tree) y en el rendimiento interactivo de cualquier sistema. La gestión de la memoria de alto ancho de banda (HBM) es un tema central en la investigación de arquitecturas de memoria, donde la capacidad y el rendimiento son trade-offs constantes, como se discute en trabajos sobre jerarquías de memoria y acceso no uniforme a la memoria (NUMA).