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¶
-
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).
-
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 asE2E MFU / Local MFU. -
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.
-
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).
-
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. -
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.
-
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).
-
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.
-
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).
-
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).
-
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.
-
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¶
- Original: https://engineering.fb.com/2026/08/03/ml-applications/training-gem-at-llm-scale-meta-ads-recommendation-foundation-model/
- Raw markdown:
raw/meta/2026-08-03-gem-training-how-meta-doubled-the-efficiency-of-its-llm-scal-5ddb723e.md
Related¶
- sources/2026-03-31-meta-adaptive-ranking-model-bending-the-inference-scaling-curve — MARM is the inference/serving side of Meta's LLM-scale ads ranking; GEM is the training side. Shared vocabulary: MFU, selective FP8, hardware-aware architecture, sparse-embedding sharding.
- sources/2026-04-02-meta-kernelevolve-how-metas-ranking-engineer-agent-optimizes-ai-infrastructure — KernelEvolve auto-synthesizes GPU kernels in Triton/TLX; GEM's JFA v4 and BlockAttention are hand-written TLX kernels of exactly the kind KernelEvolve targets, and KernelEvolve's headline win is on the Andromeda ads retrieval model.
- sources/2026-05-26-meta-silvertorch-index-as-model-a-new-retrieval-paradigm-for-recommendation-systems — SilverTorch uses TorchRec for sparse-table sharding across the GPU memory hierarchy; GEM's Fully Sharded 2D sparse parallelism is the training-time analogue.
- sources/2026-07-01-databricks-gpu-reliability — Databricks' NCCL deep-dive on payload-size algorithm switching and IB timeouts; GEM's SM-free-communication work is the same collective-communication layer optimized for efficiency rather than reliability.