El entrenamiento de modelos de recomendación a escala de hyperscaler, como el Generative Ads Recommendation Model (GEM) de Meta, presenta desafíos fundamentales que no pueden resolverse con las optimizaciones estándar para Large Language Models (LLMs). La tesis central es que la eficiencia en este dominio híbrido (LLMs + RecSys) requiere un codesign profundo y holístico de todo el stack tecnológico, desde los kernels de GPU hasta las estrategias de paralelismo distribuido y la gestión de memoria. La naturaleza de los datos de recomendación, con secuencias de longitud variable (jagged inputs) y patrones de interacción asimétricos, exige soluciones personalizadas que difieren significativamente de las arquitecturas de LLM tradicionales, donde las secuencias suelen ser densas y de longitud fija. Este enfoque integrado es crucial para superar las limitaciones de rendimiento y escalabilidad en entornos de entrenamiento con miles de GPUs.
Históricamente, los sistemas distribuidos han buscado la escalabilidad lineal mediante la minimización de la sobrecarga de comunicación y la maximización del solapamiento entre cómputo y comunicación. Sin embargo, en el contexto de modelos de recomendación masivos, la heterogeneidad de los datos y la complejidad arquitectónica introducen nuevas fricciones. La necesidad de mantener una alta utilización de la GPU (Local MFU) y una escalabilidad eficiente (Scaling Ratio) simultáneamente, obliga a repensar cómo se diseñan los algoritmos de atención, cómo se maneja la precisión numérica y cómo se orquestra el paralelismo a través de una jerarquía de red multinivel. Este artículo demuestra que la optimización de la eficiencia de entrenamiento en este nuevo paradigma es un problema de ingeniería de sistemas distribuidos que exige una atención meticulosa a cada capa del stack.
Arquitectura del Sistema
La arquitectura de entrenamiento de GEM se basa en un enfoque de codesign que aborda la eficiencia de cómputo y la eficiencia de escalado de forma independiente pero coordinada. Para la eficiencia de cómputo (Local MFU), se desarrolló una biblioteca de kernels personalizados que incluye Jagged Flash Attention (JFA), Generalized Dot-Product Attention (GDPA) y BlockAttention. JFA optimiza el manejo de 'jagged inputs' eliminando el padding y utilizando un esquema de sustracción para el enmascaramiento, así como paralelización de backward pass y especialización de warps con Triton Low-Level Extensions (TLX). GDPA unifica y acelera diversos patrones de atención asimétricos, rediseñando el pipeline del kernel y utilizando aproximaciones ALU-only para activaciones como GELU. BlockAttention reduce la complejidad de la auto-atención de O(L²) a O(L) para secuencias largas mediante atención alineada por bloques y fusión de operaciones de backward pass. Además, se implementó entrenamiento de precisión ultra-baja (MXFP8) para atención y MLP, con innovaciones en la colocación de factores de escala en TMEM, conversión online de P a MXFP8 y cuantificación por bloques, mitigando la sobrecarga de cuantificación mediante fusión de kernels y abordando la estabilidad numérica con Random Hadamard Transforms, redondeo estocástico y uso selectivo de mayor precisión.
Para la eficiencia de escalado (Scaling Ratio), se emplea un paralelismo 5D adaptado a la topología de red de Meta. Esto incluye 2D FSDP (Fully Sharded Data Parallel) con Expert Parallelism (EP) para parámetros densos, y Fully Sharded 2D Model Parallelism para parámetros sparse. El paralelismo 3D denso (EP + 2D FSDP) mapea las comunicaciones a la jerarquía de red (NVLink intra-nodo, RoCE inter-nodo dentro de zona, RoCE sobre-suscrito entre zonas) para optimizar el ancho de banda. El paralelismo sparse evoluciona a Fully Sharded 2D para eliminar la sobrecarga de memoria de O(T) al sharding de parámetros. La eficiencia de red se mejora con comunicación SM-free usando NCCLX para offload de movimiento de datos a Copy Engines y RDMA, y NVLink SHARP para reducción in-network. La eficiencia de memoria se logra con Automatic Activation Checkpointing (AutoAC) basado en compilador con presupuestos por región y cuantificación de activaciones (FP8/MX4) en los tensores checkpointed. Finalmente, el balanceo de carga para 'jagged sequences' se aborda con Base Batch Shuffling (BBS), que ordena y entrelaza sub-batches para reducir el sesgo de carga sin comunicación cross-rank.
Flujo de Entrenamiento de GEM (Simplificado)
- 1 Carga de Datos Lectura de sub-batches, ordenamiento por longitud de secuencia y entrelazado ...
- 2 Pre-procesamiento (GPU) Cuantificación de pesos en shards FSDP, fusión de cuantificación en kernels d...
- 3 Forward Pass (Dense) FSDP all-gather de parámetros de expertos (inter-nodo), EP all-gather de acti...
- 4 Forward Pass (Sparse) All-gather de shards de tabla, all-to-all de distribución de features, embedd...
- 5 Cómputo de Atención Uso de kernels JFA, GDPA, BlockAttention con MXFP8 para atención y MLP.
- 6 Checkpointing de Activaciones AutoAC con presupuestos por región y cuantificación de activaciones (FP8/MX4).
- 7 Backward Pass EP all-gather de gradientes de salida (intra-nodo), cómputo de gradientes de ...
- 8 Actualización de Pesos Reduce-scatter de parámetros sparse, aplicación de gradientes.
| Capa | Tecnología | Justificación |
|---|---|---|
| compute | NVIDIA GPUs (última generación) | Unidades de procesamiento principal para el entrenamiento del modelo, optimizadas con Tensor Cores para operaciones de baja precisión. |
| compute | Triton Low-Level Extensions (TLX) | Framework para el desarrollo de kernels de GPU personalizados de alto rendimiento, permitiendo especialización de warps y scheduling persistente. vs CUDA, Cutlass |
| networking | NVLink | Interconexión de alta velocidad intra-nodo para comunicación entre GPUs, utilizada para colectivos de baja latencia. |
| networking | RoCE (RDMA over Converged Ethernet) | Protocolo de red inter-nodo para comunicación entre hosts y zonas de IA, con diferentes niveles de ancho de banda. |
| networking | NCCLX (Meta's NCCL extension) | Biblioteca de colectivos optimizada para comunicación SM-free, utilizando Copy Engines y RDMA para offload de movimiento de datos. vs NCCL estándar |
| networking | NVLink SHARP | Tecnología de reducción in-network para offload de cómputo de reducción de los SMs a los switches de red. |
| orchestration | PyTorch (con compilador) | Framework de aprendizaje automático utilizado para el entrenamiento del modelo, con extensiones para Automatic Activation Checkpointing. |
| storage | HBM (High Bandwidth Memory) | Memoria de alta velocidad en la GPU, gestionada cuidadosamente para evitar recomputación y permitir grandes tamaños de batch. |
Trade-offs
Ganancias
- ▲ Eficiencia de cómputo (Local MFU)
- ▲ Eficiencia de escalado (Scaling Ratio)
- ▲ Utilización de GPU
- ▲ Reducción de latencia de entrenamiento
- ▲ Reducción de uso de memoria
Costes
- ▲ Complejidad de ingeniería (codesign hardware/software)
- ▲ Esfuerzo de desarrollo de kernels personalizados
- △ Sensibilidad numérica con baja precisión
Fundamentos Teóricos
El problema de la eficiencia en el entrenamiento de modelos de gran escala se remonta a los desafíos de la computación paralela y distribuida, donde la Ley de Amdahl establece los límites teóricos de la aceleración. La gestión de la comunicación y el cómputo, así como el balanceo de carga, son temas centrales en la literatura de sistemas distribuidos. El uso de técnicas como el sharding de datos y modelos, y la agregación de gradientes, se basa en principios de paralelismo de datos y modelos bien establecidos en la investigación de redes neuronales distribuidas. La atención a la topología de red y la asignación de colectivos a diferentes niveles de ancho de banda resuena con los principios de diseño de redes de interconexión de alto rendimiento, como los discutidos en trabajos sobre redes de fat-tree o toroidales.
La optimización de kernels y el uso de precisión mixta se conecta con la investigación en aritmética de punto flotante y la estabilidad numérica de algoritmos, un campo estudiado por Wilkinson y otros en la década de 1960. La introducción de FlashAttention por Dao et al. (2022) revolucionó la eficiencia de la atención al reducir la complejidad de memoria, y este trabajo extiende esos principios para manejar las complejidades de las secuencias irregulares y los patrones de atención asimétricos, un problema que no fue el foco original de FlashAttention. La gestión de la memoria a través de técnicas como el checkpointing de activaciones se inspira en trabajos sobre la optimización del uso de memoria en el entrenamiento de redes profundas, buscando un equilibrio entre el re-cómputo y el almacenamiento, un trade-off clásico en la optimización de recursos computacionales.