Engineering at Meta

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

8.5内容质量
GEM Training: How Meta Doubled the Efficiency of Its LLM-Scale Ads Foundation Model

TL;DR · AI 摘要

Meta通过自定义内核库和混合精度训练,将GEM广告推荐模型训练效率提升至20-25% MFU,同时实现4倍FLOPs扩展。

核心要点

  • Jagged Flash Attention(JFA)等自定义内核库使GPU利用率提升100%
  • 混合超低精度训练(MXFP8)优化推荐工作负载,减少内存占用30%
  • 5D并行性架构降低通信开销,实现2倍训练吞吐量提升

结构提纲

按章节快速跳转。

  1. 介绍Meta GEM模型在LLM规模训练中的效率突破及工程挑战。

  2. 解析混合架构模型在推荐系统数据特性下的训练复杂性。

  3. 通过JFA/GDPA等自定义内核实现GPU利用率翻倍。

  4. 5D并行性架构结合网络分层设计降低通信开销。

  5. 12个月内实现4倍FLOPs扩展且保持20-25% MFU。

思维导图

用一张图看清主题之间的关系。

查看大纲文本(无障碍 / 无 JS 友好)
  • GEM训练优化方法
    • 计算效率提升
      • JFA/GDPA自定义内核
      • MXFP8混合精度训练
    • 扩展效率优化
      • 5D并行性架构
      • 网络分层通信优化
    • 成果指标
      • 20-25% MFU利用率
      • 4倍FLOPs扩展

金句 / Highlights

值得收藏与分享的关键句。

#推荐系统#LLM训练#GPU优化#Meta#广告技术
打开原文

GEM训练:Meta如何将LLM级广告基础模型的训练效率提升一倍 - Meta工程实践

POSTED ON

AUGUST 3, 2026

TO

AI Research

,

Data Infrastructure

ML Applications

GEM训练:Meta如何将LLM级广告基础模型的训练效率提升一倍

.entry-meta

.entry-header

By

Darren Liu

Huayu Li

Raghav Boinepalli

Yuzhen Huang

Jackie (Jiaqi) Xu

Richard Qiu

Chunzhi Yang

Rich Zhu

Dev (Devashish) Shankar

Huaqing Xiong

Lei Tian

Ruilin Chen

Xiaoyi (Leo) Liu

Yasmine Badr

Wentao Duan

  • Meta的生成式广告推荐模型(GEM)作为Instagram和Facebook广告推荐的基础模型,现已在数千块最新一代GPU上实现LLM规模训练。本文将详细说明我们如何实现:在12个月内将端到端(E2E)训练效率提升至20-25%模型FLOPs利用率(MFU),同时将训练FLOPs规模扩大4倍,通过协同设计内核、精度、并行性、网络和内存等技术。
  • GEM训练在推荐系统与LLM的交叉领域面临独特的工程挑战,该模型结合了混合架构和推荐领域特有的数据属性,这些属性与典型LLM工作负载存在显著差异。
  • 针对LLM训练优化的AI基础设施(内核、并行性、低精度方案等)无法直接迁移,需要通过大量创新和软硬件协同设计,才能高效实现推荐模型的LLM级训练。
  • 我们通过计算效率创新和扩展效率创新的协同方案应对这些挑战:计算效率:通过定制化推荐内核库(包括锯齿状快速注意力JFA、广义点积注意力GDPA、BlockAttention等)和针对推荐工作负载优化的混合超低精度训练(包含MXFP8注意力和MLP),充分利用最新一代GPU架构。扩展效率:基于拓扑感知的5D并行性(采用无流式多处理器SM的集合操作)——密集参数使用2D FSDP+专家并行,稀疏参数采用全分片2D模型并行——与Meta的多层级网络架构协同设计,降低通信开销。
  • 成果:我们在保持训练FLOPs规模扩大4倍的同时,将GEM的端到端训练效率提升至20-25% MFU。

GEM架构及其独特的训练挑战

GEM是Meta广告系统的核心推荐基础模型。它采用混合架构,包含数万亿稀疏嵌入参数和数十亿密集参数。GEM在广告内容和用户互动数据上进行训练,数据包含两类特征:序列特征(如用户活动历史)和非序列特征(如用户位置、广告创意表示)。针对每类特征组分别应用定制化注意力机制,同时实现跨特征学习。

这种混合架构与推荐领域数据属性的相互作用,使GEM训练面临独特挑战。

挑战1:实现高GPU利用率

今天的数据中心GPU及其软件堆栈主要针对大型语言模型(LLM)工作负载进行优化,而推荐系统工作负载由于独特的数据特性和丰富的用户与广告信号交互模式,呈现出完全不同的特征,使得训练GEM规模的基础推荐模型时难以实现高GPU计算利用率。

  • 不规则输入:训练样本的序列长度高度可变,因为用户活动历史差异极大。将序列填充到最大长度会导致高达50%的计算资源浪费。
  • 多样的交互模式与非对称序列:自注意力机制处理极长序列(活动历史)但使用短注意力窗口;交叉注意力学习用户与广告的交互时使用长查询但短键/值;池化多头注意力(PMA)压缩用户活动历史,导致短查询但长键/值。这些非对称结构使内核级流水线难以有效利用计算单元。
  • 内存受限操作:例如,MLP的小型嵌入维度和为模型质量及训练稳定性设计的各种归一化操作,会导致计算单元利用率低下。
  • 数值敏感性:广告优化任务(如CTR/CVR预测)对数值变化(如精度)高度敏感,使得简单的低精度训练容易导致质量退化。

挑战2:在数千个GPU上高效扩展

在数千个GPU上训练GEM时,需要处理万亿级稀疏嵌入参数和数十亿级密集参数,这要求实现高效扩展而非单纯扩大规模。简单增加GPU数量并不能带来成比例的加速。在分布式训练中,每个训练步骤的端到端(E2E)延迟由以下公式决定:

E2E 延迟 = 各GPU Rank中的最大值(max(本地计算时间, 通信时间))

近线性扩展需要满足四个条件:

  • 总计算时间远大于总通信时间。
  • 通信操作隐藏在计算过程中且无资源竞争。
  • 内存压力导致的重复计算最小化。
  • 各Rank之间负载均衡良好。

GEM的工作负载对这些条件构成威胁:

  • 万亿级稀疏参数和数十亿级密集参数导致通信量巨大,且计算模式复杂混合。
  • 各层架构多样性导致重叠窗口不均衡;通信与计算之间的资源竞争使隐藏通信变得非 trivial。
  • 长序列和大激活值将内存使用推向极限,迫使重复计算激活值,降低效率。
  • 样本间的不规则序列导致数据驱动的负载偏斜,在不同Rank间差异显著。

我们的方案与效率框架

鉴于上述挑战,我们需要一个框架,将复杂的协同设计工作转化为少数关键技术杠杆。我们通过端到端(E2E)MFU(计算利用率)衡量训练效率,其分解为两个因素:

E2E MFU = 本地MFU(计算效率) × 扩展比例(扩展效率)

这两个因素描述了两个相关但独立的优化问题。

本地MFU(计算效率)衡量单个GPU的计算单元利用率——工作负载与硬件性能边界(roofline)的接近程度。它由内核设计、数值精度以及工作负载的计算模式(数据维度、序列长度)与GPU架构(Tensor核心、内存层次、流式多处理器调度)的匹配程度决定。

扩展比例(scaling efficiency)衡量在将任务分布到数千个GPU时,单个GPU性能的保留程度。扩展比例为1.0表示完美的线性扩展;但实际上,通信开销、负载不平衡、拖后腿效应以及由内存压力引起的激活重新计算都会削弱这一比例。

为了隔离本地MFU,我们单独在单个GPU上运行模型层,并在不进行激活重新计算或通信暴露的情况下计算加权平均MFU。扩展比例是本地MFU与端到端(E2E)MFU的比率。

这种分解很重要,因为它使我们能够将计算效率和扩展效率视为相关但不同的优化问题,每个问题都有其专用的技术集:

  • 计算效率是内核级和数值精度问题。关键手段包括内核设计和超低精度训练——两者都针对单GPU性能上限(roofline)进行优化。
  • 扩展效率是分布式系统问题。关键手段包括并行策略、网络拓扑映射、网络效率、内存管理和负载均衡——所有手段都针对单GPU与多GPU吞吐量之间的差距进行优化。

要最大化端到端MFU,必须同时解决这两个问题。

通过推荐内核和超低精度训练优化计算效率

为了解决上述推荐系统特有的挑战并提升GPU FLOPS利用率,我们构建了定制内核库和超低精度训练方案,这些方案专为最新GPU硬件上的推荐工作负载进行了定制和优化。

  • JFA — 消除因填充不规则输入而产生的高达50%的计算浪费。
  • BlockAttention — 在保持模型质量和效率的同时,将长用户历史自注意力成本从O(L²)降至O(L)。
  • GDPA — 统一并加速GEM的多样化、非对称注意力模块,这些模块在FlashAttention的密集长序列假设失效时表现不佳。
  • MXFP8 attention + MLP — 在不退化对精度敏感的CTR/CVR目标的情况下,将低精度Tensor Core吞吐量转化为实际的端到端加速。

推荐定制内核库内部实现

#### 不规则序列Flash Attention

FlashAttention专为LLM中常见的密集固定长度序列设计。在推荐模型中,用户序列本质上是不规则的——每个样本的token数量从数百到数万不等——填充到最大长度可能导致高达50%的计算浪费。

标准FlashAttention实现假设统一序列长度以实现高效分块和并行化;面对不规则输入时,简单方法要么填充(浪费计算资源),要么在短序列提前完成时让SM空闲。我们开发了JFA,这是一种专为不规则张量设计的定制FlashAttention实现,直接操作可变长度不规则张量,在消除填充开销的同时支持推荐系统特有的功能,如自定义注意力偏置、非对称查询/键值长度以及高效的反向传播。

我们通过四代迭代改进JFA,逐步缩小与填充SDPA(缩放点积注意力)的性能差距,最终在最新一代GPU上达到SOTA CUDA/Cutlass性能水平:

  • 通过减法方案实现锯齿掩码:传统的2D锯齿边界掩码(用-inf标记无效位置)消耗大量非Tensor Core指令(约占执行指令的28%)。我们采用了一种新颖的减法方案——通过零值对Query/Key进行掩码(张量内存加速器(TMA)可免费实现),并减去额外的指数,在不增加掩码开销的情况下生成数值等效的结果。
  • 反向传播并行化:FlashAttention的反向传播需要跨序列块累积dQ,通常通过代价高昂的原子加法实现。我们探索了多种方案(带原子操作的序列并行、无序列并行、带重计算的序列并行、dQ/dKdV拆分),发现对于高batch x heads的rec工作负载,采用非序列并行方案并拆分dQ计算可消除原子写入和冗余重计算,实现21-40%的反向传播加速。
  • 线程束专业化和持久内核:升级到Triton低级扩展(TLX)后,实现了显式的线程束专业化、TMA使用和持久内核调度,通过利用最新硬件特性,实现了30-100%的TFLOPS性能提升。

JFA v4(TLX)相比JFA v2实现了40-140%的TFLOPS性能提升,在生产环境锯齿分布(稀疏度0.5)下保持稳定收益,贡献了18.5%的相对本地MFU提升和12%的QPS提升。

#### 广义点积注意力(GDPA)

GEM使用多种类似注意力的交互模式——自注意力、PMA和交叉注意力——这些模式具有共同结构:两个矩阵乘法之间夹着逐元素激活函数,但将softmax替换为GELU或SiLU等激活函数。我们将其统一到一个针对最新一代GPU生产推荐系统训练工作负载优化的GDPA内核中。

现有FlashAttention内核设计用于LLM风格的密集长序列输入,在真实生产流量下表现不佳。我们观察到真实工作负载与基于短/不对称K/V序列、锯齿输入和大批次的合成基准之间存在2.6倍的前向性能差距,最坏情况下差距可达4倍,这打破了流水线占用率的假设。

我们重新设计了内核流水线、调度和数学计算,以弥合真实流量与硬件性能上限之间的差距。

  • 非softmax激活的流水线重设计:移除softmax校正阶段可释放四个线程束及其寄存器。对于短K/V序列,外层循环软件流水线在内层循环仅运行1-2次时,可恢复约10%因内层流水线导致的性能损失。
  • 锯齿张量的软件级分块调度:在CPU上预计算有效分块,完全跳过空分块,并在SM间采用之字形分配——将工作负载偏斜从6倍降低到接近平衡。
  • 仅ALU的激活近似:用GELU的SFU绑定tanh替换为6阶泰勒展开(仅使用ALU),在QK-norm(查询/键归一化)强制的输入范围下保持精度。在前向和反向传播中均消除了SFU争用。

通过这些优化,优化后的GDPA内核实现了2倍的前向加速(1,145 BF16 TFLOPs,约97% Tensor Core利用率)和1.6倍的反向加速。在短K/V生产场景下,其前向速度比Flash Attention 4(FA4)快达3.5倍。在全模型应用中,这些内核可带来超过30%的端到端训练吞吐量提升。

#### BlockAttention

对于GEM自注意力机制,核心效率挑战在于在不支付完整注意力二次成本的情况下扩展长用户序列。我们首先将层从完整自注意力转移到滑动窗口注意力,限制每个token仅关注附近事件,将复杂度从O(L²)降低到O(L * window)。这使得处理更长的序列成为可能。滑动窗口注意力(SWA)内核在JFA中跳过了窗口外的tile,通过中性NE(归一化熵,一种模型质量指标)实现了长序列自注意力延迟降低高达68%。

随后我们通过块对齐注意力进一步优化结构。由于GEM可以安全使用固定的64-token块,每个Q块仅关注对应的K/V块,将注意力转化为独立的64×64问题。这消除了SWA中仍然存在的部分窗口掩码和多tile迭代,使专用的TLX内核能够消除FlashAttention的开销,如在线softmax修正、logsumexp HBM流量和单独的Di预处理。

将RoPE反向计算融合到注意力尾部消除了另一个内存受限的内核,使梯度保留在FP32寄存器中。综合来看,TLX块注意力+融合旋转机制使自注意力层MFU相比Triton块注意力提升了+30.6%,或相比SWA基线提升了约+44%。

混合超低精度训练

在GPU上,更低的精度直接转化为更高的张量核心吞吐量。对于最新一代GPU,FP8相比FP16可提供2倍峰值FLOPS,FP4可提供4倍。我们预计未来几代GPU的低精度峰值FLOPS增长速度将快于FP16。这使得随着硬件厂商加速提升低精度FLOPS,低精度训练变得越来越有吸引力。

然而,在不造成质量退化的情况下实现低精度训练——解决数值稳定性和量化开销问题——仍然是行业性挑战。我们开发了具有数值稳定性增强的MXFP8注意力和MLP,解决了训练稳定性和量化开销问题。

#### 低精度Flash Attention

我们通过端到端MXFP8块缩放MMA扩展了FA4内核,利用最新一代GPU对低精度的原生支持,实现了前向和反向传播。主要挑战在于低精度注意力不仅仅是数据类型转换。每个GEMM(通用矩阵乘法)的K维度必须生成缩放因子,尽管FA4已完全占用TMEM(张量内存),仍需通过共享内存(SMEM)/张量内存(TMEM)分阶段传输,并在线计算softmax P和反向dS等中间结果的缩放因子。

为使张量核心加速在模块层面持续生效,量化被融合到上游归一化和投影内核中,直接生成FP8激活值和张量核心友好的缩放布局,同时避免额外的BF16全局内存流量。对于GEM的锯齿状推荐工作负载,FP8数据保留在未填充位置,仅对TMA(张量内存访问)进行紧凑缩放因子的分散/填充。这使MXFP8块缩放MMA支持转化为实际端到端注意力加速,而不会引入模型质量退化。

为满足我们的特殊需求,我们在内核层面开发了三个创新技术:

  • TMEM缩放因子放置:原始FA4完全利用了512列TMEM作为累加器,没有为块级缩放因子留下空间。我们通过将缩放因子与暂时未使用的TMEM区域重叠(例如将S(i)缩放因子放置在S(1-i)累加器区域)来解决这个问题,仅需要一个额外的轻量级屏障,该屏障隐藏在现有GEMM延迟之后。
  • 在线P到MXFP8转换:softmax输出(P)在softmax warp内原地量化为MXFP8,复用已计算的softmax归一化行最大值以避免冗余缩减。缩放因子通过优化的PTX位操作序列生成,而不是使用代价高昂的log2/round/clamp操作。
  • 分块量化:我们使用[32, 32]方形量化方案,通过warp级的redux.sync.max.abs.f32缩减操作,为每个32×32块计算一个缩放因子——使量化过程与转置无关,从而确保每个张量仅量化一次。这在反向传播中特别有用,因为需要转置后的Q,K值。

在GEM代表形状上,基于Meta内部最新一代GPU(功耗受限)的实测数据,使用MXFP8时前向核实现了>1.3倍加速。对于反向核,使用MXFP8实现了>1.5倍加速。

#### 处理量化开销

量化开销主要来自两个方面:模型参数(权重)和中间张量(激活值)。如果处理不当,额外的类型转换、缩放和数据移动可能会抵消低精度Tensor核心带来的计算加速。

  • 权重——全分片数据并行(FSDP)分片上的量化
  • 预全部收集分片量化:在FSDP全部收集之前对每个rank的本地分片进行量化,将量化成本分摊到各个rank上,避免在每个rank上对完全收集后的权重重新量化。量化FSDP通信:通过传输低精度负载(而非BF16)来减少全部收集体积并降低全部收集延迟,从而进一步抵消量化开销。
  • 激活值——内核融合
  • 线性模块:我们通过将激活值量化融合到前导归一化操作(PreNorm融合)中,避免了单独量化步骤带来的额外内核启动和HBM流量。注意模块:除了PreNorm融合外,我们还通过将量化操作融合到前导投影中,使注意内核可以直接使用低精度激活值,无需额外量化步骤。

#### 应对数值稳定性

量化误差、异常值和舍入偏差会使低精度训练在数值上变得脆弱,尤其是在梯度计算中。我们通过以下方式应对这些挑战:

  • 异常值缓解:我们应用随机哈达玛变换来分散异常值并平滑量化前的分布。
  • 配方调优(细粒度控制):我们使用随机舍入消除确定性舍入偏差。权重梯度(WGrad)跳过/更高精度:我们观察到激活值和梯度可能表现出更严重的异常值行为;选择性跳过WGrad或使用更高精度可以显著提升模型质量。
  • 混合精度:我们在能带来最大收益的场景(如大GEMM)使用超低精度,在超低精度不足以满足模型质量目标时(如模型后期层对量化误差更敏感)回退到BF16。

如上所述,对于大规模分布式训练:

近线性扩展需要满足四个条件:总计算时间大于通信时间、计算与通信无竞争地重叠、最小化重复计算,以及良好的负载均衡。我们的优化针对每个条件,提升GEM的扩展效率。

| 条件 | GEM的挑战 | 优化 | |------|-----------|------| | 总计算时间 > 总通信时间 | O(万亿)稀疏参数和O(十亿)密集参数在混合计算模式下导致大量通信 | 拓扑感知的5D并行 | | 通信隐藏在计算中且无竞争 | 通信与计算之间的资源竞争 | SM Free通信 | | 最小化因内存压力导致的重复计算 | 长序列与大激活值将内存使用推向极限,迫使激活值重复计算 | 带量化自动激活检查点 | | 跨rank的良好负载均衡 | 样本间锯齿状序列导致数据驱动的负载偏斜在不同rank间变化 | 序列长度感知的负载均衡 |

5D并行,基于Meta网络拓扑的优化

GEM的混合架构要求为密集参数和稀疏参数分别采用不同的并行策略,因为它们具有不同的计算和通信模式。我们使用5D并行在数千个GPU上高效扩展GEM训练:密集参数使用2D FSDP与专家并行(EP),稀疏参数使用全分片2D模型并行。

设计原则是将通信量与拓扑层级的可用带宽匹配。当某层级的集合通信成为瓶颈时,我们引入新的并行维度以减少该层级的消息量或组规模。

GEM使用的Meta训练集群具有三层网络拓扑:每台主机的8个GPU通过NVLink连接,AI区域内的主机通过RoCE连接,AI区域之间通过带宽缩减的RoCE连接。

#### 密集参数并行演化:从1D到3D并行

GEM的O(十亿)密集参数使用FSDP进行分片。参数在GPU间分布,计算前通过all-gather重建,梯度通过reduce-scatter同步。我们在FSDP基础上增加两个维度——副本(DDP)维度(形成2D FSDP)和EP,总共形成三个密集参数并行维度(3D密集参数并行)。

| 并行维度 | 集合通信 | 拓扑层级 | 带宽 | |----------|----------|----------|------| | EP(专家并行) | all-gather / reduce-scatter | 节点内NVLink | 高 | | FSDP(组内) | inter-node(AI区域内部) | 中等 | | DDP(跨组) | all-reduce | inter-node(可能跨区域) | 低(过载) |

这种拓扑感知的分布式训练使3D密集参数并行高效——每个维度的通信成本与其拓扑层级的可用带宽匹配。

为何使用2D FSDP:通过减小组规模提升带宽

在数千GPU规模下,标准FSDP需要跨全部rank进行集合通信,而有效带宽会随着组规模扩大而下降——尤其是在跨多个AI区域时。2D FSDP通过将通信拆分为两个拓扑感知层级解决此问题:

  • FSDP分片组:参数在更小的组(如128-256 GPU)内通过all-gather / reduce-scatter进行分片和重建。减小组规模可实现更高的有效带宽。
  • DDP 副本组:梯度通过副本组内的 all-reduce 进行同步。由于参数已经通过 FSDP 进行分片,每个 rank 仅发送参数的一部分——消息大小足够小,即使在跨区域带宽较低的情况下也能容忍。

我们积极预取参数的 all-gather 操作,将每个模块的通信与前一个模块的计算进行流水线处理,以最大程度实现重叠。这种方法对大多数模块效果良好——然而,像 DHEN(深度分层集成网络)这样的大模块专家,其参数规模仍导致通信时间超过相邻计算时间,暴露出瓶颈并降低端到端效率。

添加专家并行:将重型通信转移到最快的链路

为了解决大型密集专家模块的通信暴露问题,我们在 2D FSDP 基础上叠加 EP(专家并行)。通过 EP,每个 rank 仅持有单个专家,将 FSDP 的 all-gather 缩小到单个专家的参数——减少组规模和消息规模。

额外的 EP 通信通过节点内高带宽的 NVLink 进行,使其易于隐藏。前向和反向传播协调 FSDP 和 EP 的集合操作:

  • 前向:FSDP 对专家参数进行 all-gather(16 路,跨节点)→ EP 对激活值进行 all-gather(2 路,节点内 NVLink)→ 在完整批次上计算本地专家 → EP 对输出进行 reduce-scatter(2 路,节点内 NVLink)。
  • 反向:FSDP 对专家参数进行 all-gather(16 路,跨节点)→ EP 对输出梯度进行 all-gather(2 路,节点内 NVLink)→ 计算专家梯度 → EP 对输入梯度进行 reduce-scatter(2 路,节点内 NVLink)→ FSDP 对参数梯度进行 reduce-scatter(16 路,跨节点)。

#### 稀疏并行性演进:从 1D 到 2D 无内存开销的并行性

GEM 的稀疏参数(万亿级嵌入表)带来了与密集参数不同的独特扩展挑战。嵌入表需要通过 all-to-all 通信进行模型并行分片以实现特征分布,其庞大的规模使内存开销成为主要限制因素。我们通过三代稀疏并行性演进解决了这些挑战。

负载不均衡

内存开销

通信成本

V1: 1D 模型并行

非常高 – 全 rank

V2: 2D 模型并行

良好

高 – 每个副本组维护完整的稀疏参数副本(万亿级)

中等 – 减小组规模

V3: 全分片 2D 模型并行

接近零

中等 – 通过高速 NVLink 的额外通信

V1 → V2: 解决负载不均衡和通信瓶颈

在数千 GPU 规模下,1D 模型并行遇到两个根本瓶颈:

  • 负载不均衡:将嵌入表分片分布在数千个 rank 上会导致严重的工作负载偏斜——每个 rank 持有的分片过少,无法实现平衡划分。
  • 通信延迟:all-to-all 集合操作的组规模随总 rank 数量增长。跨节点带宽随着组规模扩大而快速下降,特别是在作业跨越多个 AI 区域且带宽被过度订阅时尤为明显。

2D 模型并行通过将 rank 分成更小的模型并行组(例如 256 GPU),并让多个副本组执行数据并行,解决了这两个问题。每个副本组在更小的范围内独立进行分片和通信,降低 all-to-all 延迟并改善负载均衡——在大规模场景下相比 1D 实现显著的 QPS 提升。

V2 → V3: 消除内存开销

V2的权衡在于内存:每个副本组必须保存其分配的分片参数的完整副本。对于GEM的万亿参数稀疏表而言,这种O(T)的开销会显著消耗HBM资源——阻碍模型规模的进一步扩展。

Fully Sharded 2D通过将每个副本的参数副本进一步在组内分片,消除了这一开销。每个计算节点仅存储分片的一部分,参数按需重建:

  • 前向传播:全收集表分片 → 全对全特征分布 → 嵌入查找 → 全对全嵌入返回
  • 反向传播:全收集表分片 → 全对全梯度交换 → 本地更新 → 归约分散参数

V3新增的全收集和归约分散操作映射到节点内NVLink。我们通过流水线技术将这些集合操作与密集计算重叠执行,并调度全收集操作在内存峰值前释放重建副本。

通过这些优化,我们在GEM的训练规模下实现了几乎无开销的稀疏扩展,通信暴露度极低。

网络效率:将通信从SM中剥离

借助5D并行,GEM通过流水线技术将大部分通信隐藏在计算内核之后。然而,通信集合操作仍可能占用SM,导致SM资源竞争。通信内核会占用本可用于并行计算的SM资源(例如全收集和归约分散各占用约24个SM),造成高达15%的效率损失。更糟糕的是,计算内核性能下降可能超过SM占用率的损失,因为波浪调度可能导致更多资源浪费。

因此,我们主要的网络效率优化方向是实现SM无关通信——将数据移动从SM卸载到专用硬件引擎。

对于纯数据移动集合操作(如全收集),我们使用NCCLX(Meta对NCCL库的扩展)实现无复制、无SM占用的通信。NCCLX利用硬件特性实现无需SM参与的数据传输:复制引擎(CE)处理节点内NVLink传输,RDMA处理节点间传输,使全收集的SM占用从24个降至1个。这为计算任务回收了约23个SM,在全规模训练下带来约5%的端到端QPS提升。

对于需要归约的集合操作(如All-Reduce),我们发现通过网络交换机硬件卸载归约计算的NVLink SHARP技术是可行方案,可有效降低SM占用。

内存效率:实现大本地批量而不支付完整内存代价

每张GPU的内存可分为三类:激活内存、嵌入表和密集参数(包括优化器状态)。在将嵌入表和密集参数跨GPU并行分片后,激活内存成为每张GPU内存的主要组成部分,其规模随模型和批量大小增长。

我们采用两种技术应对这一挑战:

基于编译器的自动激活检查点(AutoAC)

基于编译器的激活检查点技术已通过分析联合前向-反向图中的单个节点,超越了传统全有或全无的重新计算方式——节省昂贵操作,仅重新计算廉价的逐点操作。但其仍对整个模型应用单一内存预算,当不同区域(编译子图之间存在图断点)的重新计算ROI(每GB激活内存节省的延迟)存在差异时,仍会损失性能。我们用定制的区域预算计划取代全局预算,使内存流向回报最高的区域。这使得内存-延迟权衡超越了任何统一预算所能实现的水平。

激活量化

在AutoAC基础上,我们进一步通过激活量化压缩内存使用。该技术作用于检查点张量——AutoAC已确定需要为反向传递保留的中间激活张量集合。启用时,它会在前向和反向图边界处将这些保存的激活节点(例如从BF16量化到FP8/MX4)。

借助这些优化,我们能够使用大本地批量大小(高达1K+样本)并配合适度的激活重新计算成本,高效训练GEM模型。这对扩展至关重要,因为小批量大小和高激活重新计算都会损害MFU。

负载均衡:特定于推荐系统的拖后腿问题

LLM训练可通过将所有序列填充为固定长度来避免负载均衡。但对GEM而言,用户序列本身具有锯齿状特征,填充会浪费50%以上的计算资源。锯齿状内核虽然避免了每级的浪费,却带来了新问题——每次迭代都变化的数据驱动计算偏斜。

最重负载的计算节点每次迭代均超出平均值约15%。

选择合适的再平衡策略

我们考虑了本地和全局再平衡策略以解决工作负载不平衡:

| 方法 | 机制 | 平衡质量 | 开销 | |------|------|----------|------| | 本地(同级内) | 每个节点独立再平衡自己的批次 | 高:达到最优的90% | 无(零跨节点通信) | | 全局(跨节点) | 节点通过all-to-all交换样本 | 近乎完美 | 每次训练步骤引入新的all-to-all集合操作 |

全局方法的开销——每次训练步骤都需要进行集合操作——反而抵消了其旨在实现的效率提升。我们开发了一种新方法,称为基础批次洗牌(BBS),分布式读取器生成小子批次(128样本),在合并为完整训练批次(每节点1k+样本)时按总序列长度排序并交错排列(最重与最轻配对)——在零跨节点通信情况下实现理论最优平衡的大部分效果。

BBS使GEM训练效率提升4%,包含4%的QPS提升和4%的峰值内存减少。启用后,最大负载与平均负载的差距立即显著下降。

向更大规模和更高效率迈进

在大语言模型(LLMs)与推荐系统交叉领域训练基础模型是一个协同设计问题,而非单纯的软件或硬件问题。我们在此描述的2倍效率提升来自于对堆栈每一层的精心优化——内核、精度、并行性、网络和内存必须同步推进。我们预计下一次2倍提升将以类似方式实现,且随着智能体的引入,优化周期的迭代速度将显著加快。随着GEM模型的持续扩展,我们期望继续突破系统边界,在AI基础设施堆栈的不同层级实现更极致的协同设计,进一步提升计算和扩展效率。我们分享这项工作,希望更广泛的社区能在其运行的工作负载中发现类似的机会。

致谢

我们想感谢彭天舒、张嘉生、杨Angel、Shah Rikin、桑柯、Tang Kevin、Kadluczka Pawel、周Jacky、许涵、Palaz Enes、闫浩、Siso Jake、Wu Rupert、Xu Liangbei、胡宇硕、Liu Serena、Yu Hongtao、Su Bor-Yiing、Mohan Santosh、Si Min、Jiang Shali、Chen Laming、Liu Boyang、Zhou Qinghai、Xia Xiaozhen、Rudy Jason、Xu Jiayi、Chanpuriya Dan、Yang Justin、Chadha Mandeep、Au Carmen、Kuang Hairong、Iyengar Subodh、Balasubramanian Balaji、Sullerey Anamaya、Vimawala Viral、Gur Saket、Wang May、Sinha Vibha、Hashimov Rustam、Wang Ernest、Leung Max、Chang Shuo、Sultan Musharaf、Platon Oana、Nie Jade、Falconer Eric、Chen Ping、Reeves Damian、Chen Xian、Wen Ellie、Sun Chonglin、Musumeci GP、Srinivasan Reva、Hansen Brian、Sung Vivienne、Phelps Patrick、Massimi Paolo、Zheng Jie、Madan Anuj、Garg Nikhil、Gan Xiaorui、Bocharov John、Tewari Ritwik、Chen Wenlin、Liu Rocky、Yan Tak、Kolay Santanu、Pandey Sandeep、Steiner Matt,以及所有在大规模高效训练Meta最大广告推荐工作负载背后默默付出的v-team成员。