GEM 训练实践:Meta 如何将 LLM 规模广告推荐基础模型训练效率翻倍
- 这是 Instagram 和 Facebook 广告推荐背后的基础模型,目前已在数千块最新一代 GPU 上以 LLM 规模进行训练。本文将详细介绍我们是如何通过协同设计内核、精度、并行策略、网络和内存,在 12 个月内将端到端(E2E)训练效率翻倍至 20–25% 的 Model FLOPs Utilization(MFU),同时将训练 FLOPs 扩展 4 倍。
- GEM 的训练面临着推荐系统与 LLM 交叉领域独有的工程挑战,因为该模型结合了混合架构以及不同于典型 LLM 负载的推荐领域数据特性。
- 为 LLM 训练优化的 AI 基础设施(内核、并行策略、低精度方案等)并不能直接迁移,需要大量创新以及软硬件协同设计,才能高效地在 LLM 规模下训练推荐模型。
-
- 底层优化:通过定制化的推荐内核库(Jagged Flash Attention(JFA)、Generalized Dot-Product Attention(GDPA)、BlockAttention 等)以及针对推荐负载混合使用超低精度训练(包括 MXFP8 attention 和 MLP)实现,充分利用最新一代 GPU 的架构特性。
- 系统级并行:拓扑感知的 5D 并行配合无 SM 参与的集合通信——稠密参数采用 2D FSDP + Expert Parallelism,稀疏参数采用 Fully Sharded 2D Model Parallelism——与 Meta 的多层网络层次结构协同设计,以降低通信开销。
- 成果:过去 12 个月内,我们将 GEM 的 E2E 训练效率翻倍至 20-25% MFU,同时将总训练 FLOPs 扩展了 4 倍。
GEM 是 Meta 广告系统背后的核心推荐基础模型,采用混合架构,包含数万亿稀疏嵌入参数和数十亿稠密参数。GEM 基于广告内容和用户互动数据进行训练,数据特征分为两类:序列特征(例如用户活动历史)和非序列特征(例如用户位置、广告创意表示)。模型对每类特征分别应用定制化的注意力机制,同时也支持跨特征学习。

正是这种混合架构与推荐域数据特性之间的相互作用,让 GEM 的训练变得格外困难。
挑战一:实现高 GPU 利用率
如今数据中心的 GPU 及其软件栈主要针对 LLM 工作负载做了优化,而推荐工作负载由于其独特的数据特征以及丰富的用户与广告信号交互模式,与 LLM 截然不同,这使得在训练 GEM 这种规模的推荐基础模型时,很难让 GPU 算力得到充分利用。
- 输入长度参差不齐:训练样本的序列长度差异很大,因为用户行为历史长短不一。如果统一填充到最大长度,算力浪费最多可达 50%。
- 交互模式多样且序列不对称:自注意力作用于极长的序列(行为历史)但注意力窗口较短;交叉注意力用于学习用户与广告的交互,查询很长而键/值较短;池化多头注意力(PMA)则压缩用户行为历史,导致查询短而键/值长。这些不对称的张量形状使得核内流水线难以有效占满计算单元。
- 访存密集型操作:例如 MLP 的嵌入维度较小,以及为保证模型质量和训练稳定性而引入的各种归一化层,都会让计算单元利用率偏低。
- 数值敏感性强:广告优化任务(如 CTR/CVR 预测)对数值变化(例如精度)极为敏感,简单的低精度训练很容易引发模型质量下降。
挑战二:在数千张 GPU 上高效扩展
在数千张 GPU 上训练拥有数万亿稀疏嵌入参数和数十亿稠密参数的 GEM,需要的是高效扩展,而不仅仅是规模扩大。简单堆叠 GPU 数量并不能换来线性的加速比。在分布式训练中,每个训练步骤的端到端延迟由下式决定:
E2E Latency = Max across GPU Rank (Max(Local Compute Time, Communication Time))
接近线性的扩展需要同时满足以下四个条件:
- 总计算时间远大于总通信时间。
- 通信被计算隐藏,且二者不产生资源竞争。
- 因显存压力导致的重复计算尽可能少。
- 各 rank 之间负载均衡良好。
GEM 的工作负载对上述每一条都构成了挑战:
- O(Trillion) 级别的稀疏参数与 O(Billion) 级别的稠密参数带来大量通信,并形成混合计算模式。
- 各层架构差异较大,使得重叠窗口不均匀;通信与计算之间的资源竞争让隐藏通信开销变得不容易。
- 长序列带来巨大的激活值,将显存推到极限,不得不进行激活重计算,从而侵蚀效率。
- 不同样本的序列长度参差不齐,产生数据驱动的负载倾斜,且这种倾斜在不同 rank 之间各不相同。
我们的方法与效率框架
面对上述挑战,我们需要一个框架,把庞大的协同设计工作收敛到少数几个技术抓手。我们用 E2E MFU 来衡量训练效率,它可以分解为两个因子:
E2E MFU = Local MFU(计算效率)× Scaling Ratio(扩展效率)
这两个因子对应着两个相关但不同的优化问题。
Local MFU(计算效率)衡量单卡计算单元被利用的程度——即负载距离硬件 roofline 有多近。它由 kernel 设计、数值精度,以及负载的计算模式(数据维度、序列长度)与 GPU 架构(Tensor Core、存储层次、SM 调度)的契合程度共同决定。
Scaling Ratio(扩展效率)衡量当负载扩展到上千张 GPU 时,单卡性能被保留多少。Scaling Ratio 等于 1.0 表示完美的线性扩展;实际中,通信开销、负载不均、straggler 效应以及显存压力下的激活重计算都会拉低它。
为了单独测算 Local MFU,我们在单卡上逐层运行模型,并在不引入激活重计算和通信的情况下,按权重加权平均得到 MFU。Scaling Ratio 则由 Local MFU 与 E2E MFU 的比值推导得出。
这种分解很关键,它让我们能够把计算效率和扩展效率视为相关但独立的优化问题,分别用各自的技术手段去解决:
- 计算效率是 kernel 层面和数值精度层面的问题,核心抓手是 kernel 设计和超低精度训练——两者都瞄准单卡 roofline。
- 扩展效率是一个分布式系统问题。可调节的杠杆包括并行策略、网络拓扑映射、网络效率、内存管理和负载均衡——全部针对单 GPU 与多 GPU 吞吐量之间的差距。
两者必须同时优化,才能最大化端到端 MFU。
通过推荐内核与超低精度训练优化计算效率
为应对上述推荐系统特有的挑战,并提升 GPU FLOPS 利用率,我们构建了一套定制内核库和超低精度训练方案,专为最新 GPU 硬件上的推荐负载定制和优化。
- JFA——消除因填充不规则输入造成的最高达 50% 的算力浪费。
- BlockAttention——将长用户历史自注意力成本从 O(L²) 降至 O(L),同时保持模型质量和效率
- GDPA——统一并加速 GEM 中多样、非对称的注意力模块,突破了 FlashAttention 稠密长序列假设的局限
- MXFP8 attention + MLP——将低精度 Tensor Core 吞吐转化为真实的端到端加速,且不会损害对精度敏感的 CTR/CVR 目标
深入定制化推荐内核库
不规则序列 Flash Attention
FlashAttention 面向 LLM 中常见的稠密定长序列而设计。在推荐模型中,用户序列天然是不规则的——每个样本的 token 数从数百到数万不等——填充到最大长度会浪费最高达 50% 的算力。
标准 FlashAttention 实现假设序列长度均匀,以高效分块和并行化;面对不规则输入时,朴素方案要么填充(浪费算力),要么在短序列提前完成时让 SM 闲置。我们开发了 JFA,一种直接作用于变长不规则张量的定制 FlashAttention 实现,既消除了填充开销,又支持推荐场景特有的功能,例如自定义注意力偏置、非对称 query/key-value 长度以及高效的反向传播。
JFA 历经四代演进,逐步从慢于带填充的 SDPA(缩放点积注意力),到在最新一代 GPU 上追平 SOTA CUDA/Cutlass 性能:
- 基于减法的不规则掩码:传统的二维不规则边界掩码(用 -inf 标记无效位置)会消耗大量非 Tensor Core 指令(约占执行指令的 28%)。我们改用一种新的减法方案——将 Query/Key 掩码置零(Tensor Memory Accelerator (TMA) 可以免费完成这一操作),然后减去多余的指数——在消除掩码开销的同时获得数值上等价的结果。
- 反向并行化:FlashAttention 的反向传播需要在各个序列 tile 间累加 dQ,通常依赖代价较高的 atomic add。我们尝试了多种方案(序列并行 + atomic、无序列并行、序列并行 + 重计算、拆分 dQ/dKdV),发现对于 batch × heads 较大的推荐负载,不使用序列并行并拆分 dQ 计算的方案能消除 atomic 写入和冗余重计算,从而带来 21–40% 的反向加速。
- Warp 特化与持久化 Kernel:升级到 Triton Low-Level Extensions (TLX) 后,可以显式进行 warp 特化,并配合 TMA 和持久化 Kernel 调度——借助最新的硬件特性,获得了 30–100% 的 TFLOPS 提升。
JFA v4 (TLX) 相比 JFA v2 实现了 40–140% 的 TFLOPS 提升,在生产环境的真实不规则分布(稀疏度 0.5)下均能保持稳定的加速效果,相对本地 MFU 提升 18.5%,QPS 提升 12%。
广义点积注意力(GDPA)
GEM 中存在多种类注意力的交互模式——self-attention、PMA 和 cross-attention——它们共享一个共同结构:两次矩阵乘法中间夹一个逐元素激活函数,但用 GELU 或 SiLU 等激活函数替代了 softmax。我们将这些模块统一到一个 GDPA Kernel 中,专门针对最新一代 GPU 上的生产 RecSys 训练负载进行优化。
现有的 FlashAttention Kernel 是为 LLM 风格的稠密长序列输入设计的,在真实的线上流量下表现很差。我们观察到,真实负载与合成 benchmark 之间存在 2.6 倍的前向性能差距,极端情况下差距高达 4 倍,这主要由短/不对称的 K/V 序列、不规则输入以及破坏流水线 occupancy 假设的大 batch size 造成。

我们重新设计了内核流水线、调度策略与计算逻辑,以此缩小实际流量负载与硬件理论峰值之间的性能差距。
- 面向非 softmax 激活的流水线重构:移除 softmax 修正阶段,释放四个 warp 及其占用的寄存器。对于较短的 K/V 序列,外层软件流水线可弥补内层流水线仅运行 1–2 次迭代所带来的约 10% 性能损失。
- 面向不规则张量的软件级分片调度:在 CPU 端预先计算有效分片、完全跳过空分片,并在 SM 之间采用 zigzag 分配,将负载不均衡从 6 倍降至接近均衡。
- 纯 ALU 激活近似:用 6 阶泰勒展开(纯 ALU)替代 GELU 中依赖 SFU 的 tanh,在 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^2)$ 降到 $O(L \times \text{window})$,从而使更长的序列变得切实可行。滑动窗口注意力(SWA)内核在 JFA 中跳过了窗口外的 tile,将长序列自注意力的延迟降低了最多 68%,同时对 NE(normalized entropy,归一化熵,一种模型质量指标)没有任何影响。
随后,我们利用块对齐注意力(block-aligned attention)进一步优化了结构。鉴于 GEM 可以安全地使用固定的 64 token 块,每个 Q 块只需关注对应的 K/V 块,将注意力计算转化为独立的 64×64 运算。这样做消除了 SWA 中仍然存在的部分窗口掩码和多 tile 迭代,并让一个专用的 TLX 内核省去了 FlashAttention 中的一些开销,比如在线 softmax 修正、logsumexp 的 HBM 读写,以及单独的 Di 预处理。
将 RoPE 反向计算融合到注意力的尾部(epilogue)中,又消除了一个受内存带宽限制的内核,同时让梯度保持在 FP32 寄存器中。这两处优化叠加起来,TLX 块注意力 + 融合的 rotary 机制相比 Triton 块注意力,将自注意力层的 MFU 提升了 30.6%,相比 SWA 基线则提升约 44%。

混合超低精度训练
在 GPU 上,更低的精度直接意味着更高的 Tensor Core 吞吐量。以最新一代 GPU 为例,FP8 的峰值 FLOPS 是 FP16 的 2 倍,FP4 则是 4 倍。我们预计下一代 GPU 中低精度峰值 FLOPS 的增速会更快。随着硬件厂商对低精度 FLOPS 的扩展速度超过 FP16,低精度训练将变得越来越有吸引力。
然而,要在保证模型质量不下降的前提下实现低精度训练——既要解决数值稳定性问题,也要解决量化开销问题——这仍然是整个行业面临的挑战。我们开发了带有数值稳定性增强的 MXFP8 Attention 和 MLP,同时解决了训练稳定性和量化开销两个问题。
低精度 Flash Attention
我们为 FA4 内核扩展了端到端的 MXFP8 块缩放矩阵乘累加(覆盖前向和反向),充分利用了最新一代 GPU 对低精度的原生支持。核心难点在于,低精度注意力并非简单的数据类型替换:必须沿各 GEMM(通用矩阵乘法)的 K 维度生成缩放因子,在 FA4 的 TMEM 已被占满的情况下将其经由共享内存(SMEM)/ 张量内存(TMEM)进行中转分阶段传递,并对 softmax 输出 P 以及反向过程中的 dS 等中间结果进行在线计算。
为了使 Tensor Core 的加速效果在模块层面也得以保留,我们将量化操作融合到上游的归一化和投影内核中,直接产出 FP8 激活值以及适配 Tensor Core 的缩放布局,避免额外的 BF16 全局内存读写。针对 GEM 中变长的推荐负载,FP8 数据保持在非填充位置,仅对紧凑的缩放因子进行分散/填充以适配 TMA。这样既打通了 MXFP8 块缩放 MMA 的端到端实现路径,让注意力真正获得加速,又不会带来模型质量上的回退。

为满足这些独特需求,我们在内核层面开发了三项创新:
- TMEM 缩放因子布局:FA4 原本已用满 512 列的 TMEM 来存放累加器,没有给块缩放因子留出空间。我们通过将缩放因子与临时空闲的 TMEM 区域重叠来解决(例如把 S(i) 的缩放因子放到 S(1-i) 的累加器区域),整个过程只需额外增加一道轻量级同步屏障,并可被现有 GEMM 的延迟所掩盖。
- P 到 MXFP8 的在线转换:在 softmax 所在的 warp 内,就地将 softmax 输出 P 量化为 MXFP8,复用行最大值(softmax 归一化中已经算过)以避免重复规约。缩放因子通过优化的 PTX 位操作指令得出,省去了昂贵的 log2/round/clamp 操作。
- 块级量化:采用 [32, 32] 的方形分块量化,每个 32×32 块通过 redux.sync.max.abs.f32 的 warp 级规约计算一个缩放因子,从而做到量化结果与转置无关,每个张量只需量化一次。这对反向过程尤为有用,因为反向需要用到转置后的 Q、K。
在 GEM 的代表性 shape 上、基于 Meta 内部功耗受限的最新一代 GPU 测得:使用 MXFP8 时,前向 kernel 加速超过 1.3 倍;反向 kernel 加速超过 1.5 倍。

处理量化开销
量化开销主要来自两部分:模型参数(权重)和中间张量(激活)。如果处理方式过于粗糙,额外的类型转换、缩放以及数据搬运会抵消低精度 Tensor Core 带来的计算加速。
- 权重——基于 FSDP(Fully Sharded Data Parallel)分片进行量化
-
- All-gather 前的分片量化:在 FSDP all-gather 之前先对各 rank 的本地分片做量化,使量化开销在各 rank 之间摊薄,避免每个 rank 对完整 gather 后的权重重复量化。
- FSDP 量化通信:传输低精度 payload(相比 BF16),降低 all-gather 的数据量与时延,进一步抵消量化开销。
- 激活——kernel fusion
-
- 线性模块:避免单独再做一次量化(会带来额外 kernel launch 和 HBM 读写),将激活量化融合到它前面的 normalization 中(PreNorm fusion),以此消除开销。
- 注意力模块:除 PreNorm fusion 外,还将量化融合到它前面的 projection 中,使 attention kernel 直接消费低精度激活,无需再额外做一步量化。
解决数值稳定性问题
量化误差、离群值和舍入偏差会让低精度训练在数值上变得脆弱,梯度计算尤其如此。我们通过以下方法应对这些挑战:
- 离群值缓解:
- 在低精度量化前施加随机 Hadamard 变换,将离群值摊薄、平滑数据分布。
- 训练配方调优(细粒度控制):
- 采用 stochastic rounding,消除确定性舍入带来的偏差。
- 跳过 / 使用更高精度的权重梯度(WGrad):我们观察到激活值和梯度可能出现更严重的离群行为,有选择地跳过 WGrad 或使用更高精度能够显著提升模型质量。
- 混合精度:
- 在最受益的场景(如大型 GEMM)中使用超低精度,在超低精度无法满足模型质量要求的场景(例如模型的后几层对量化误差更敏感)则回退到 BF16。
扩展效率:5D 并行、网络、内存与负载均衡
如上所述,对于大规模分布式训练:
端到端延迟 = 各 GPU Rank 取最大值(本地计算时间、通信时间 中的较大者)
要实现近线性扩展,需要满足四个条件:总计算时间大于通信时间、计算与通信可无争用地重叠、最小化重计算、以及良好的负载均衡。我们的优化针对每个条件发力,持续提升 GEM 的扩展效率。
| 条件 | GEM 面临的挑战 | 优化手段 |
|---|---|---|
| 总计算时间 > 总通信时间 | 万亿级稀疏参数与十亿级稠密参数带来大量通信,并伴随混合计算模式。 | 拓扑感知的 5D 并行 |
| 通信完全隐藏在计算背后且无资源争用 | 通信与计算之间存在资源争用 | SM 空闲通信(SM Free Communication) |
| 因内存压力导致的重计算最小化 | 长序列配合大激活值将内存使用推向极限,迫使对激活进行重计算 | 带量化的自动激活检查点 |
| 各 Rank 之间良好的负载均衡 | 样本间序列长度参差不齐,导致各 Rank 上的数据驱动负载倾斜 | 感知序列长度的负载均衡 |
5D 并行:结合 Meta 网络拓扑进行优化
GEM 的混合架构要求各组成部分采用不同的并行策略,因为稠密参数与稀疏参数具有不同的计算和通信模式。我们使用 5D 并行将 GEM 的训练高效扩展到上千块 GPU:稠密参数采用带 Expert Parallelism(EP)的 2D FSDP,稀疏参数采用全分片 2D 模型并行(Fully Sharded 2D Model Parallelism)。
设计原则是让通信量与拓扑各层级的可用带宽相匹配。当某个集合通信在某一层成为瓶颈时,我们就引入新的并行维度来降低该层的消息量或组大小。
GEM 使用的 Meta 训练集群采用三层网络拓扑:每台主机内的 8 块 GPU 通过 NVLink 互连,同一 AI zone 内的主机通过 RoCE 连接,不同 AI zone 之间则通过过订阅的 RoCE 相连,带宽进一步降低。

密集并行的演进:从 1D 到 3D 并行
GEM 中规模达 O(Billion) 的密集参数通过 FSDP 进行分片。参数分布在各 GPU 上,计算前通过 all-gather 重建,梯度通过 reduce-scatter 同步。在 FSDP 之上我们又增加了两个维度——一个副本(DDP)维度(构成 2D FSDP)和 EP——总共形成三个密集并行维度(3D 密集并行)。
| 并行维度 | 集合通信类型 | 拓扑层级 | 带宽 |
|---|---|---|---|
| EP(专家并行) | All-gather / reduce-scatter | 节点内 NVLink | 高 |
| FSDP(组内) | All-gather / reduce-scatter | 节点间(同一 AI zone 内) | 中 |
| DDP(跨组) | All-reduce | 节点间(可能跨 zone) | 低(过订阅) |
这种拓扑感知的分布式训练正是 3D 密集并行高效的原因——每个维度的通信开销与其所在拓扑层级的可用带宽相匹配。
为何选择 2D FSDP:通过缩减组大小获取更高带宽
在数千 GPU 的规模下,标准 FSDP 需要跨全部 rank 执行集合通信,而有效带宽会随组大小下降——尤其在跨越多个 AI zone 时更为明显。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)→ 在完整 batch 上计算本地专家 → 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 的稀疏参数(O(万亿) 级别的 embedding 表)带来了与密集参数截然不同的扩展挑战。Embedding 表需要模型并行分片,并通过 all-to-all 通信进行特征分发,其庞大的规模使得内存开销成为主要约束。我们经历了三代稀疏并行的演进来解决这些挑战。
| 负载不均衡 | 内存开销 | 通信成本 | |
|---|---|---|---|
| V1:1D 模型并行 | 差 | 无 | 非常高——全 rank |
| V2:二维模型并行 | 良好 | 较高——每个副本组都保留一份完整的稀疏参数副本 O(万亿级) | 中等——组规模有所减小 |
| V3:全分片式二维模型并行 | 良好 | 接近零 | 中等——通过高速 NVLink 增加额外通信 |
V1 → V2:解决负载不均与通信瓶颈
在数千卡 GPU 规模下,一维模型并行在效率上会遇到两个根本性瓶颈:
- 负载不均:将嵌入表分片分散到数千个 rank 上会导致严重的工作负载倾斜——每个 rank 持有的分片太少,无法做到均衡划分。
- 通信延迟:all-to-all 集合通信的组规模随总 rank 数线性增长。跨节点带宽随组规模增大而迅速下降,尤其当任务横跨多个 AI 集群、带宽存在超额订阅时更为明显。
二维模型并行通过将 rank 划分为更小的模型并行组(例如 256 张 GPU),同时由多个副本组执行数据并行,从而同时解决这两个问题。每个副本组在更小的范围内独立完成分片与通信,既降低了 all-to-all 延迟,也改善了负载均衡——在大规模场景下相比一维并行带来显著的 QPS 提升。
V2 → V3:消除内存开销
V2 的代价是内存:每个副本组都必须持有其负责分片参数的完整副本。对于 GEM 的万亿级稀疏参数表而言,这种 O(T) 的开销会占用大量 HBM,阻碍模型规模的进一步扩展。
全分片式二维并行通过对每个副本的参数副本在其组内进一步分片,彻底消除了这一开销。每个 rank 只存储分片的一小部分,参数按需重建:
- 前向:all-gather 表分片 → all-to-all 特征分发 → 嵌入查找 → all-to-all 嵌入回收
- 反向:all-gather 表分片 → all-to-all 梯度交换 → 本地更新 → reduce-scatter 参数
V3 额外引入的 all-gather 和 reduce-scatter 被映射到节点内的 NVLink 上。我们通过流水线将这些集合通信与并发的稠密计算重叠执行,并调度 all-gather 在峰值内存使用前释放已重建的副本。
经过这些优化,我们能够在 GEM 的训练规模下几乎无开销地实现稀疏扩展,通信开销也极少。

网络效率:让通信脱离 SM
借助 5D 并行,GEM 通过流水线把大部分通信隐藏在计算内核之后。但通信集合运算也会占用 SM,从而引发 SM 争用。通信内核会占用约 24 个 SM(例如 all-gather、reduce-scatter),这些 SM 本可被并行的计算内核使用,导致最高约 15% 的效率损失。更糟糕的是,由于 wave 调度可能产生更多浪费,计算内核性能的下降幅度会超过 SM 占用率的损失。
因此,我们网络效率优化工作的核心方向是实现 SM-free 通信——把数据搬运从 SM 卸载到专用硬件引擎。

对于纯数据搬运类的集合运算(例如 all-gather),我们使用 NCCLX——Meta 对 NCCL 库的扩展——来实现免拷贝、SM-free 的通信。NCCLX 利用硬件特性在无 SM 参与的情况下搬运数据:由 Copy Engine (CE) 负责节点内 NVLink 传输,RDMA 负责节点间传输,使 all-gather 的 SM 占用从 24 个降至 1 个。这相当于为计算腾出约 23 个 SM,在完整训练规模下带来约 5% 的端到端 QPS 提升。
对于需要归约操作的集合运算(例如 All-Reduce),我们发现 NVLink SHARP 的网络内归约是一条可行路径,它通过把归约计算从 SM 卸载到网络交换机硬件上来减少 SM 占用。
内存效率:使用大批次而不承担全部内存开销
每块 GPU 的显存可归为三类:激活值、嵌入表和稠密参数(含优化器状态)。经过并行切分后,嵌入表和稠密参数被分摊到各 GPU 上,激活值成为每卡显存的主要占用者,其大小随模型和批次规模线性增长。
我们采用了两种技术来解决这个问题:
基于编译器的自动激活检查点(AutoAC)
PyTorch 基于编译器的激活检查点机制,通过在前向–反向联合图中逐节点分析,已经优于传统的"全量或全不"重计算方式——它会跳过昂贵的算子、重计算廉价的逐元素算子。但它对整个模型仍采用统一的内存预算,当不同区域(图中断点之间的各编译子图)在重计算投入产出比(每 GB 激活所节省的延迟)上存在差异时,这种做法就会损失性能。我们用一套按区域定制的预算调度替代了全局预算,让内存流向收益最高的区域,从而突破了任何统一预算所能达到的内存–延迟权衡上限。 激活量化 在 AutoAC 之上,我们进一步借助激活量化来压缩内存。它作用于已被 AutoAC 判定需要保存以供反向传递使用的检查点张量——即那些已确定的中间激活张量。启用后,在前向图和反向图的交界处,将这些被保存的激活节点从 BF16 量化到 FP8/MX4。 借助这些优化,我们能够使用较大的本地批大小(最高可达 1K+ 样本),同时只付出适度的激活重计算开销来高效训练 GEM 模型。这一点对于扩展性至关重要,因为小批大小和高昂的激活重计算都会拖累 MFU。负载均衡:推荐场景特有的掉队问题
LLM 训练可以通过将所有序列填充到固定长度来规避负载均衡。但在 GEM 场景下,用户序列天然长短不一,填充方式会浪费 50% 以上的算力。不规则(jagged)内核避免了单卡上的浪费,却带来了新问题——每轮迭代都不同的、由数据驱动的算力倾斜。 最重的 rank 每轮迭代都比平均水平高出约 15%。 选择合适的重均衡策略 我们考虑了局部和全局两种重均衡策略来解决负载不均问题: | 方法 | 机制 | 均衡质量 | 开销 | |---|---|---|---| | **局部(卡内)** | 每个 rank 独立重均衡自己的批次。 | 高:达最优值的 90% | 无(零跨 rank 通信)。 | | **全局(跨卡)** | | | |
全局方案的开销——每一步训练都要做集合通信——抵消了它本想带来的效率提升。我们开发了一项新技术,称为基础批次混洗(Base Batch Shuffling,BBS):分布式读取器先生成小型子批次(128 个样本),合并成完整训练批次(每 rank 1000+ 样本)时按总序列长度排序并交错排列(最重的与最轻的配对)——无需跨 rank 通信,即可获得绝大部分理论上的最优负载均衡。

BBS 为 GEM 训练带来了 4% 的效率提升,其中 QPS 提升 4%,峰值内存降低 4%。该方案上线后,最大负载与平均负载之间的差距立即缩小。
迈向下一阶段的规模与效率
在 LLM 与推荐系统的交叉领域训练基础模型,是一个协同设计问题,而非单纯的软件问题或硬件问题。我们所获得的 2 倍效率提升,来自对堆栈每一层的精细优化——内核、精度、并行策略、网络与内存必须协同演进。我们预计下一个 2 倍效率提升也将以类似的方式实现,并且随着智能体开始自动化部分优化闭环,迭代速度会进一步加快。在持续扩展 GEM 模型的过程中,我们期望持续突破系统边界,深化 AI 基础设施各层之间的极致协同设计,从而进一步提升算力与规模效率。我们分享这项工作,是希望更广泛的社区也能在自己的工作负载中看到类似的机会。
致谢
感谢 Tianshu Peng、 Jiasheng Zhang、 Angel Yang、 Rikin Shah、 Ke Sang、 Kevin Tang、 Pawel Kadluczka、 Jacky Zhou、 Han Xu、 Enes Palaz、 Hao Yan、 Jake Siso、 Rupert Wu、 Liangbei Xu、 Yusuo Hu、 Serena Liu、 Hongtao Yu、 Bor-Yiing Su、 Santosh Mohan、 Min Si、 Shali Jiang、 Laming Chen、 Boyang Liu、 Qinghai Zhou、 Xiaozhen Xia、 Jason Rudy、 Jiayi Xu、 Dan Chanpuriya、 Justin Yang、 Mandeep Chadha、 Carmen Au、 Hairong Kuang、 Subodh Iyengar、 Balaji Balasubramanian、 Anamaya Sullerey、 Viral Vimawala、 Saket Gur、 May Wang、 Vibha Sinha、 Rustam Hashimov、 Ernest Wang、 Max Leung、 Shuo Chang、 Musharaf Sultan、 Oana Platon、 Jade Nie、 Eric Falconer、 Ping Chen、 Damian Reeves、 Xian Chen、 Ellie Wen、 Chonglin Sun、 GP Musumeci、 Reva Srinivasan、 Brian Hansen、 Vivienne Sung、 Patrick Phelps、 Paolo Massimi、 Jie Zheng、 Anuj Madan、 Nikhil Garg、 Xiaorui Gan、 John Bocharov、 Ritwik Tewari、 Wenlin Chen、 Rocky Liu、 Tak Yan、 Santanu Kolay、 Sandeep Pandey、 Matt Steiner, 以及整个 v-team 团队,正是在他们的努力下,Meta 最大规模的广告推荐工作负载得以高效完成训练。
分享:
- 分享到 Facebook(在新窗口中打开) Facebook
- 分享到 Threads(在新窗口中打开) Threads
- 分享到 WhatsApp(在新窗口中打开) WhatsApp
- 分享到 LinkedIn(新窗口打开) LinkedIn
- 分享到 Reddit(新窗口打开) Reddit
- 分享到 X(新窗口打开) X
- 分享到 Bluesky(新窗口打开) Bluesky
- 分享到 Mastodon(新窗口打开) Mastodon
- 分享到 Hacker News(新窗口打开) Hacker News
- 通过邮件发送给好友(新窗口打开) 邮件


