El Cuello de Botella del Que Nadie Habla: HBM, No FLOPs

¡Hola Devs! Todo equipo que escala LLMs choca con la misma pared. Tus GPUs no están limitadas por compute — están limitadas por memoria. Pesos del modelo, gradientes, estados del optimizador, buffers de comunicación y activaciones intermedias compiten por el mismo pool de HBM. En el momento en que subes sequence length o batch size: OOM.

Aquí entra el host offloading. En lugar de recomputar activaciones en el backward pass (rematerialization), haces streaming de ellas a memoria pinned del host en el forward pass y las traes de vuelta cuando se necesiten. El trade-off cambia de "compute extra" a "bandwidth extra".

En clusters commodity, ese trade es malo — PCIe es demasiado lento. Pero en NVIDIA Grace Blackwell, se vuelve un superpoder. La CPU Grace y la GPU Blackwell están conectadas por NVLink-C2C a 900 GB/s bidireccional, y la plataforma Vera Rubin lo duplica a 1.8 TB/s. De repente, la memoria del host es un área de staging legítima, no un castigo.

Si te gusta ver cómo otras infraestructuras repiensan jerarquías de memoria bajo carga real, este análisis de la migración de headless browsers de Cloudflare es un caso paralelo excelente. 🚀

AI developer analyzing GPU memory bottleneck charts for LLM training with JAX host offloading Coding Session Visual

Los Resultados: DeepSeek-V3 671B y Llama 3.1 405B

Todos los benchmarks corrieron en NVIDIA GB200 NVL72 con 128 GPUs, usando MaxText (JAX + XLA) como framework.

DeepSeek-V3 671B (MoE + MLA)

La política hace offload de proyecciones MLA query/key-value e intermedios de up-projection del MoE — activaciones lo suficientemente grandes para decidir si una config de batch cabe o no.

ConfigMicro BatchGlobal BatchTFLOPs/s/deviceGPU Peak (GiB)Host Mem (GiB)
Host offload + LHS + pipelined81024908.2165.2145.1
Sin offload + rematerialization81024578.3151.30.0
Host offload, sin LHS, sin pipeline81024541.6145.6145.1
Sin offload, save on device2256425.3113.30.0
Sin offload, save on device81024OOM0.0

¡Eso es 57% más rápido que rematerialization y 67.7% más rápido que offloading ingenuo sin LHS ni pipelining! Y checa la fila 5: sin offloading, micro batch 8 / global batch 1024 simplemente no cabe. El offloading lo hace viable.

Llama 3.1 405B (Denso)

ConfigLHSPipelinedTFLOPs/s/deviceGPU Peak (GiB)Host Mem (GiB)
Sin offloadONOFF2,669149.60
QKV offloadONOFF2,746149.970.9
QKV offloadOFFOFF2,569139.770.9
QKV offloadONON2,718151.070.9

Un 2.9% de ganancia sobre baseline — menor que DeepSeek, pero revelador. Aquí LHS solo ya esconde casi toda la latencia, así que pipelining no aporta nada. Para modelos densos, offloading es optimización de performance. Para modelos MoE con footprint gigante de activaciones, es desbloquear capacidad.

Las Flags de XLA Que Hacen Esto Funcionar

# Habilita el latency hiding scheduler
--xla_gpu_enable_latency_hiding_scheduler=true

# Habilita pipelined host offloading (crítico para workloads MoE)
--xla_gpu_enable_pipelined_host_offloading=true

# Aumenta el trabajo async en vuelo para que LHS solape copies con colectivas NCCL
--xla_gpu_experimental_parallel_async_compute_limit=8

La tercera flag está subestimada. Le da espacio al scheduler para solapar copies device-to-host con comunicación colectiva. Sin ella, dejas throughput en la mesa.

NVIDIA Blackwell GB200 server rack with NVLink-C2C interconnect for host offloading architecture Development Concept Image

Limitaciones y Cuidados

Host offloading no es bala de plata. Aquí donde falla:

  • Tensores chicos: si el offload no reduce presión de HBM significativamente, el costo de transferencia domina.
  • Workloads con poco overlap: si no hay compute o comunicación independiente suficiente para esconder latencia, pierdes performance. El Llama 3.1 405B lo muestra — LHS ya bastaba.
  • Cuello de botella no es memoria: si tu workload está limitado por colectivas NCCL o overhead de kernel launch, offloading no ayuda.
  • Estimaciones estáticas mienten: la memoria en runtime incluye scratch de NCCL, workspace de attention de cuDNN y buffers del framework. Siempre profilea con runs reales.

Una cosa más: el run con offload usó más memoria de GPU (165.2 vs 145.6 GiB) porque mantiene copy buffers y activaciones prefetchadas residentes. Estás cambiando capacidad de HBM por throughput. Es una decisión deliberada, no almuerzo gratis. 🍽️

Cómo Empezar

  1. Empieza con un run JAX chico y representativo — no un toy, no producción.
  2. Elige activaciones grandes de forward paths caros (proyecciones QKV, up-projections MoE).
  3. Habilita offloading vía jax.remat con memory_kind="pinned_host".
  4. Mide tres cosas: memoria GPU en runtime, memoria del host, y step time end-to-end.
  5. Profilea con Nsight Systems para confirmar que copies D2H y H2D realmente solapan con compute y NCCL.

Para quienes también piensan en cómo las plataformas cloud toman decisiones de escala, vale la pena ver la integración de Google Cloud Workbench con VS Code — es otra capa de la misma pregunta: "¿dónde corre realmente el trabajo?".

Cloud GPU cluster diagram showing HBM to pinned host memory streaming for DeepSeek-V3 training System Abstract Visual

La Conclusión

Host offloading resignifica un trade-off clásico. En lugar de pagar con compute (rematerialization), pagas con bandwidth (streaming). En interconnects de la clase NVLink-C2C, esa bandwidth es lo suficientemente barata para ganar.

El resultado de DeepSeek-V3 es el titular: 908 TFLOPs/s/device con micro batch 8 / global batch 1024, una configuración que literalmente da OOM sin offloading. Eso no es un 5% de tuning — es un nuevo régimen operacional.

Pero no copies ciegamente. Profilea primero. Offloading es una decisión de posicionamiento de memoria que debe validarse con mediciones, no suposiciones. Si tus activaciones son chicas o tu workload ya tiene buen overlap, no ganas nada. Si HBM es la pared entre tú y el siguiente batch size, esta puede ser la palanca que buscabas.

Próximos pasos:

  • Lee el tutorial de host offloading de JAX (jax.remat, checkpoint policies, memory_kind="pinned_host").
  • Prueba los containers de NVIDIA JAX-Toolbox como punto de partida mantenido.
  • Profilea tu workload actual con Nsight Systems para ver si la latencia de transferencia está expuesta.

Fuente: NVIDIA Developer Blog — Reducing HBM Bottlenecks in JAX-Based LLM Training with Host Offloading

Este contenido fue redactado con la asistencia de herramientas de IA, basándose en fuentes confiables, y fue revisado por nuestro equipo editorial antes de su publicación. No reemplaza el asesoramiento de un profesional especializado.