Skip to content

META

Read original ↗

GEM Training: How Meta Doubled the Efficiency of Its LLM-Scale Ads Foundation Model

Summary

Meta's ML Applications team describes how it trains GEM (Generative Ads Recommendation Model) — the foundation model behind ads recommendations across Instagram and Facebook — at LLM scale on several thousand latest-generation GPUs. Over 12 months Meta doubled end-to-end training efficiency to 20–25% MFU while scaling total training FLOPs 4×, through the joint co-design of kernels, numerical precision, parallelism, networking, and memory. The central framing: E2E MFU decomposes into Local MFU (compute efficiency) × Scaling Ratio (scaling efficiency) — two related but distinct optimization problems. Compute efficiency came from a custom recommendation-kernel library (Jagged Flash Attention, GDPA, BlockAttention) plus MXFP8 ultra-low-precision attention + MLP; scaling efficiency came from topology-aware 5D parallelism with SM-free collectives, compiler-based activation checkpointing, and a recsys-specific load-balancing scheme.

Key Takeaways

  1. Recsys training at LLM scale is a distinct problem from LLM training. GPU stacks (kernels, parallelism, low-precision recipes) are tuned for LLMs; GEM's hybrid architecture — trillions of sparse embedding parameters + billions of dense parameters — plus rec-domain data properties (jagged inputs, asymmetric attention shapes, memory-bound ops, numerical sensitivity of CTR/CVR objectives) means LLM infra "does not directly transfer" (Source: sources/2026-08-03-meta-gem-training-how-meta-doubled-the-efficiency-of-its-llm-scal).

  2. The efficiency framework is a product, not a metric. E2E MFU = Local MFU × Scaling Ratio (compute-efficiency-vs-scaling-efficiency) turns a sprawling co-design effort into two levers: Local MFU is a kernel + numerical-precision problem (per-GPU roofline); Scaling Ratio is a distributed-systems problem (communication overlap, load balance, straggler, recompute). Local MFU is isolated by running layers individually on one GPU with no recompute or comm; scaling ratio is derived as E2E MFU / Local MFU.

  3. Jagged inputs are the defining rec-domain hazard. User sequences vary from hundreds to tens of thousands of tokens; padding to max length wastes up to 50% of compute. JFA operates directly on variable-length jagged tensors, and jaggedness reappears as data-driven load skew across ranks that changes every iteration — a recsys-specific straggler problem LLMs sidestep by padding.

  4. Kernel-library wins are large and specific. JFA v4 (TLX) delivers 40–140% TFLOPs improvement over JFA v2 (18.5% relative local-MFU gain, 12% QPS). GDPA unifies self/cross/PMA attention (softmax replaced by GELU/SiLU) and hits 2× forward speedup (1,145 BF16 TFLOPs, ~97% Tensor Core util), up to 3.5× over FlashAttention 4 on short-K/V production shapes; kernels together give >30% E2E throughput. BlockAttention (fixed 64-token blocks + fused RoPE backward) adds +30.6% self-attention-layer MFU over Triton block attention (~+44% over SWA).

  5. MXFP8 turns low-precision FLOPs into real E2E speedups without CTR/CVR regression. FP8 gives 2× peak FLOPS over FP16, FP4 gives 4×. Meta extended the FA4 kernel with end-to-end MXFP8 block-scaled MMA: **>1.3× fwd,

    1.5× bwd** on GEM shapes. Three kernel innovations: TMEM scale-factor overlap, online P-to-MXFP8 conversion in the softmax warp, and transpose-invariant [32,32] block quantization computed once.

  6. Quantization overhead is neutralized by fusing it into the preceding kernel. Weight quantization is done on the FSDP shard before all-gather (pre-all-gather-shard-quantization) and the low-precision payload is communicated (smaller all-gather). Activation quantization is fused into the preceding normalization (PreNorm fusion) or projection (quantization-fused-into-preceding-kernel) — no extra kernel launch, no extra HBM traffic.

  7. Numerical stability is engineered, not assumed. Random Hadamard Transforms spread outliers, stochastic rounding removes deterministic bias, weight-gradient (WGrad) is selectively skipped or kept higher-precision, and later quantization-sensitive layers fall back to BF16 — ultra-low precision applied only where it pays (large GEMMs).

  8. 5D parallelism is matched to a three-tier network hierarchy. GEM's cluster: 8 GPUs/host on NVLink, hosts within an AI zone on RoCE, AI zones on oversubscribed RoCE. Dense: 2D FSDP + Expert Parallelism (3D); sparse: Fully Sharded 2D Model Parallelism. The design principle (topology-aware-parallelism): match each collective's message volume to the bandwidth of its topology tier — heavy EP all-gathers go on intra-node NVLink; cross-zone traffic is only small sharded gradients.

  9. Sparse parallelism evolved V1→V3 to erase memory overhead. 1D model parallelism (poor load balance, very high comm) → 2D model parallelism (good balance, but each replica holds a full O(trillion) copy) → Fully Sharded 2D (near-zero memory overhead; each rank stores a fraction, reconstructed on-demand via extra all-gather/reduce-scatter mapped to NVLink).

  10. SM-free communication reclaims compute. Collectives that run on SMs steal ~24 SMs (up to 15% efficiency loss). NCCLX moves data via the Copy Engine (intra-node NVLink) + RDMA (inter-node), cutting all-gather SM usage from 24 → 1 (~5% E2E QPS at full scale); NVLink SHARP offloads reductions to switch hardware (sm-free-communication).

  11. Memory: big local batches without the full bill. Compiler-based AutoAC replaces a single global recompute budget with per-region budget schedules (memory flows to highest recompute-ROI regions); activation quantization (BF16→FP8/MX4) squeezes checkpointed tensors further — enabling 1K+ sample local batches at modest recompute cost.

  12. Base Batch Shuffling solves the recsys straggler with zero comm. The heaviest rank exceeds average by ~15% every iteration. Global rebalancing (all-to-all per step) negates its own gains; BBS (base-batch-shuffling-load-balancing) has distributed readers emit small 128-sample sub-batches, sort by total sequence length, and interleave heaviest-with-lightest when merging into 1K+ batches — capturing ~90% of optimal balance with zero cross-rank communication: 4% QPS + 4% peak-memory reduction.

Operational Numbers

Metric Value
Training scale several thousand latest-gen GPUs
E2E training efficiency (MFU) 20–25% (doubled in 12 months)
Total training FLOPs growth 4× over 12 months
Sparse parameters O(trillions) embedding params
Dense parameters O(billions)
Padding waste avoided (jagged) up to 50% of compute
JFA v4 (TLX) vs JFA v2 40–140% TFLOPs improvement
JFA local-MFU / QPS gain +18.5% relative local MFU, +12% QPS
GDPA forward 2× (1,145 BF16 TFLOPs, ~97% Tensor Core util)
GDPA vs FlashAttention 4 (short K/V) up to 3.5× forward
Kernel library E2E throughput >30% improvement
BlockAttention self-attn MFU +30.6% over Triton block attn (~+44% over SWA)
SWA long-seq self-attn latency −68% (neutral NE)
MXFP8 fwd / bwd kernel >1.3× / >1.5×
FP8 / FP4 peak FLOPS vs FP16 2× / 4×
SM reclaim (all-gather) 24 SMs → 1 SM (~5% E2E QPS)
SM-contention efficiency cost avoided up to 15%
Local batch size up to 1K+ samples
Heaviest-rank overshoot ~15% per iteration
Base Batch Shuffling gain +4% QPS, −4% peak memory (~90% of optimal balance)
Network tiers NVLink (intra-host) / RoCE (intra-zone) / oversubscribed RoCE (inter-zone)
GPUs per host 8 (NVLink)
FSDP shard group ~128–256 GPUs

Caveats

  • Architecture-overview voice — no absolute QPS, GPU count, vendor/model (H100/GB200/B200 not named; "latest-generation" only), fleet size, or wall-clock training time disclosed.
  • GEM model architecture (layer topology, DHEN expert count, exact sparse/dense param split) not given beyond O(trillions)/O(billions).
  • Thread/expert classification, load-heuristic details for parallelism-dimension sizing, and BBS sub-batch merge algorithm are described qualitatively.
  • MXFP8 kernel innovations reference the FA4 kernel and TMEM/TMA hardware features without naming the specific GPU generation.
  • Companion PyTorch blog posts referenced for GDPA and TLX BlockAttention (external, not ingested).

Source

Last updated · 766 distilled / 2,225 read