Skip to content

Accelerating Sharded Data Parallelism at Scale with Federated Learning ​

本文由 paper-daily 使用 DeepSeek 自动生成,仅供快速了解论文;关键结论请以原文为准。

论文原文 · PDF · 源文件

【一句话总结】受联邦学习启发,将大规模分片数据并行解耦为松耦合联邦组,通过组间低频 FedAvg 聚合大幅降低通信开销,实现 8.04 倍训练加速并提升模型质量。

基本信息 ​

属性内容
作者Gianluca Mittone, Marco Aldinucci
来源arXiv:2609.20359
发布日期2026-09-17
抓取领域分布式/并行训练系统
学科方向分布式系统 · 人工智能 · 性能优化
arXiv 分类cs.DC, cs.AI, cs.PF
适用层次进阶
标签【标签】分片数据并行, 联邦学习, 大规模训练, 通信优化, 基础模型
PDF在线阅读
代码仓库暂无

问题的初衷(Why - 为什么要做这个研究) ​

【问题的初衷】随着基础模型(Foundation Models, FMs)规模的指数级增长,其训练过程需要在数千块高端 GPU 上持续数月之久,这对高性能计算系统提出了严峻挑战。分片数据并行(Sharded Data Parallelism, DP)作为当前主流的加速策略,通过将数据和模型切分到多块 GPU 上来提升吞吐量。然而,当部署规模扩大时,分片数据并行会产生难以承受的通信开销(Communication Overhead),尤其是在具有异构性能的多层互连(Multi-tier Interconnects)架构上,跨节点、跨机架的梯度同步与参数聚合成为瓶颈。现有方法如完全分片数据并行(Fully Sharded Data Parallelism, FSDP)和混合分片数据并行(Hybrid Sharded Data Parallelism, HSDP)虽然在一定程度上缓解了显存压力,但其全局通信模式在大规模集群中仍然导致通信时间随 GPU 数量近似线性甚至超线性增长,严重制约了训练效率。与此同时,联邦学习(Federated Learning, FL)在分布式边缘设备上展现出高效的通信特性,其 FedAvg 式聚合机制通过减少同步频率和通信量实现了松耦合的分布式训练。受此启发,本文试图将联邦学习的高效通信原则引入大规模分片数据并行训练中,以解决多层级互连下的通信瓶颈问题,从而在保持模型质量的同时显著提升训练速度。


问题的解决(What - 提出了什么方案) ​

【问题的解决】本文提出了两种混合算法——FL+FSDP 和 FL+HSDP,其核心思想是将大规模数据并行部署解耦为多个较小的、松耦合的联邦组(Federation Groups),在组内采用分片数据并行进行高效计算,而在组间采用 FedAvg 风格的聚合方式进行周期性同步。具体而言,FL+FSDP 在完全分片数据并行的基础上,将全局进程组划分为多个联邦组,每个组内部独立执行 FSDP 的前向与反向传播,组间仅通过低频率的 FedAvg 聚合交换模型参数或梯度。FL+HSDP 则进一步结合混合分片策略,在组内同时使用分片与复制,以平衡显存与通信。与现有方法相比,本文方法的本质区别在于:第一,通过解耦全局同步为组内同步和组间松耦合聚合,大幅减少了跨组通信量,尤其适合异构多层级互连;第二,全局批大小(Global Batch Size)不再随 GPU 总数线性增长,而是受限于联邦组的大小,从而避免了超大批量带来的优化困难;第三,形式化分析了通信成本,证明了该方法在扩展性上的优势。实验表明,在 512 块 A100 GPU 上预训练 Llama3.1 8B 模型时,FL+FSDP 和 FL+HSDP 在相同超参数下实现了高达 8.04 倍的数据处理速度提升和 4.48 倍的评估困惑度(Perplexity)降低,兼顾了计算效率与模型质量。


技术方法详解(How - 怎么实现的) ​

【技术方法详解】

  • 联邦组划分与松耦合聚合:将全局数据并行进程组划分为 G 个联邦组,每组包含 K 个 GPU,组内采用分片数据并行(FSDP 或 HSDP),组间每 T 步执行一次 FedAvg 聚合。聚合时仅交换模型参数或梯度,通信量从 O(N) 降至 O(G),其中 N 为总 GPU 数。
  • 组内分片策略:在 FL+FSDP 中,组内每个 GPU 持有模型参数的 1/K 分片,前向和反向传播时通过全收集(All-Gather)和减少散射(Reduce-Scatter)操作同步分片参数;在 FL+HSDP 中,组内进一步划分为复制组和分片组,以降低显存占用。
  • 全局批大小控制:传统数据并行的全局批大小随 N 线性增长,而本文方法中全局批大小 Bglobal=BlocalimesK,仅与组大小 K 相关,避免了超大批量导致的收敛问题。
  • 通信成本形式化分析:论文推导了 FL+FSDP 和 FL+HSDP 的通信成本模型,考虑了多层级互连的带宽差异。设组内通信成本为 Cintra,组间通信成本为 Cinter,则总通信成本为 Ctotal=Cintra+Cinter/T,其中 T 为聚合周期。通过增大 T,可进一步降低组间通信开销。
  • 与现有框架的兼容性:该方法可无缝集成到 PyTorch FSDP 和 HSDP 中,仅需修改进程组划分和聚合逻辑,无需改动模型定义或优化器。
  • 超参数调优:关键超参数包括联邦组数量 G、组大小 K、聚合周期 T 和学习率调整策略。论文通过实验给出了推荐配置,例如在 512 GPU 上设置 G=8、K=64、T=4。

系统架构图 ​

方法流程图 ​

核心公式与算法 ​

【核心公式】

Ctotal=Cintra+CinterT

其中 Ctotal 为总通信成本,Cintra 为组内通信成本,Cinter 为组间通信成本,T 为聚合周期。该公式表明,通过增大 T 可降低组间通信开销。

Bglobal=Blocal×K

其中 Bglobal 为全局批大小,Blocal 为单卡本地批大小,K 为联邦组大小。该公式说明全局批大小仅与组大小相关,而非总 GPU 数。

θglobal(t+1)=1G∑g=1Gθg(t)

其中 θglobal(t+1) 为第 t+1 轮聚合后的全局模型参数,θg(t) 为第 g 个联邦组在第 t 轮的本地模型参数,G 为联邦组数量。该公式描述了 FedAvg 风格的组间聚合。


应用场景(Where - 在哪落地) ​

【应用场景】

  • 大规模基础模型预训练:在拥有数千块 GPU 的超级计算集群中,使用 FL+FSDP 或 FL+HSDP 可以显著减少跨机架通信,加速 Llama、GPT 等模型的预训练过程。例如,在 512 块 A100 上训练 Llama3.1 8B 时,数据处理速度提升 8.04 倍,意味着原本需要数月的训练可缩短至数周。
  • 异构集群上的分布式训练:在由不同代际 GPU 和多种互连技术(如 NVLink、InfiniBand、以太网)组成的异构集群中,本文方法通过联邦组划分将通信限制在性能相近的组内,组间仅进行低频聚合,从而有效利用异构资源,避免慢节点拖累整体训练。
  • 边缘-云协同训练:在边缘设备与云端服务器协同训练的场景中,边缘设备可视为联邦组,云端作为聚合中心,利用 FedAvg 式聚合减少边缘与云之间的通信,同时组内使用分片数据并行加速本地训练。预期效果是在保证模型精度的同时,降低通信成本并提升训练效率。

具体技术细节示例(How in Action - 算法如何执行) ​

【具体技术细节示例】假设我们在 8 块 GPU 上训练一个简化模型,总 GPU 数 N=8,划分为 G=2 个联邦组,每组 K=4 块 GPU,聚合周期 T=2。模型参数为 W∈R4×4,初始化为全零矩阵。本地批大小 Blocal=2,全局批大小 Bglobal=2×4=8。

步骤 1:初始化:全局模型参数 Wglobal 广播到两个联邦组,组 1 和组 2 的初始参数均为 Wglobal。

步骤 2:组内 FSDP 前向与反向:组 1 内 4 块 GPU 各持有 W 的 1/4 分片(即每块 GPU 持有 1×4 的行分片)。对于本地批次数据,每块 GPU 计算其分片对应的前向输出,通过全收集(All-Gather)获取完整参数,计算损失后反向传播,再通过减少散射(Reduce-Scatter)将梯度分片到各 GPU。组 2 同理。

步骤 3:本地更新:组 1 和组 2 分别更新其本地模型参数,得到 W1(1) 和 W2(1)。假设经过一次本地迭代后,W1(1)=[1111111111111111],W2(1)=[2222222222222222]。

步骤 4:判断聚合周期:当前迭代步数 t=1,T=2,未达到聚合周期,因此继续本地迭代。

步骤 5:第二次本地迭代:组 1 和组 2 再次执行前向与反向,更新本地参数。假设得到 W1(2)=[3333333333333333],W2(2)=[4444444444444444]。

步骤 6:组间 FedAvg 聚合:达到聚合周期 T=2,执行 FedAvg 聚合:Wglobal(2)=12(W1(2)+W2(2))=[3.53.53.53.53.53.53.53.53.53.53.53.53.53.53.53.5]。

步骤 7:广播与继续:将 Wglobal(2) 广播到两个联邦组,开始下一轮本地迭代。

通过这个示例可以看出,组间通信仅发生在聚合周期 T=2 时,且通信量为模型参数大小,远小于每步都进行全局同步的传统数据并行。


实验结果(Results - 效果如何) ​

【实验结果】论文在 512 块 NVIDIA A100 GPU 上进行了 Llama3.1 8B 模型的预训练实验,采用多层级互连架构(节点内 NVLink、节点间 InfiniBand)。基准数据集为大规模文本语料,对比方法包括标准的 FSDP 和 HSDP。实验设置中,FL+FSDP 和 FL+HSDP 均使用与基线相同的超参数(学习率、批大小等),联邦组数量 G=8,组大小 K=64,聚合周期 T=4。关键结果如下:在数据处理速度方面,FL+FSDP 和 FL+HSDP 分别比对应的 FSDP 和 HSDP 基线快 8.04 倍和 6.12 倍;在模型质量方面,评估困惑度(Perplexity)分别降低了 4.48 倍和 3.21 倍。此外,通信开销分析表明,组间通信量仅占总通信量的 5% 以下,验证了松耦合聚合的有效性。消融实验进一步显示,聚合周期 T 在 2 到 8 之间时性能最佳,过大或过小均会导致效率下降。

实验结果可视化 ​


优势与不足 ​

【优势与不足】 优势:

  • 显著降低大规模训练中的通信开销,尤其适合多层级异构互连环境,通过解耦全局同步为组内同步和组间松耦合聚合,实现了高达 8.04 倍的速度提升。
  • 全局批大小受限于联邦组大小而非总 GPU 数,避免了超大批量导致的优化困难,同时提升了模型质量(困惑度降低 4.48 倍)。
  • 方法通用性强,可无缝集成到现有 FSDP 和 HSDP 框架中,仅需修改进程组划分和聚合逻辑,无需改动模型定义。 不足:
  • 联邦组划分和聚合周期 T 等超参数需要针对具体集群拓扑和模型规模进行调优,缺乏自适应机制。
  • 组间松耦合可能导致模型参数更新存在滞后,在极端非独立同分布(Non-IID)数据分布下可能影响收敛性,尽管论文未深入探讨此场景。
  • 实验仅在 Llama3.1 8B 模型和 512 GPU 上验证,未展示在更大规模(如千亿参数模型、数千 GPU)下的扩展性。

相关工作 ​

【相关工作】

  • 完全分片数据并行(FSDP):PyTorch 提出的显存优化数据并行方法,通过分片模型参数、梯度和优化器状态来降低单卡显存占用,但通信开销随 GPU 数量增加而显著增长。本文在其基础上引入联邦组划分以降低通信。
  • 混合分片数据并行(HSDP):结合分片与复制策略,在节点内使用 FSDP、节点间使用数据并行,以平衡通信与显存。本文将其扩展为联邦组内 HSDP、组间 FedAvg 聚合。
  • 联邦学习(FL)与 FedAvg:在分布式边缘设备上通过周期性聚合本地模型来减少通信,本文借鉴其松耦合聚合思想,将其应用于大规模 GPU 集群训练。
  • 大规模分布式训练通信优化:如 Ring All-Reduce、梯度压缩、通信重叠等技术,本文从进程组拓扑层面进行优化,与这些技术互补。
  • 基础模型预训练:如 GPT、Llama 系列的大规模训练,本文以 Llama3.1 8B 为例验证了方法的有效性。

未来研究方向 ​

【未来方向】

  • 自适应联邦组划分与聚合周期调优:研究根据集群拓扑、带宽差异和模型规模自动确定最优的联邦组数量 G、组大小 K 和聚合周期 T,以进一步提升扩展性和易用性。
  • 非独立同分布数据下的收敛性分析:当前方法假设各组数据分布相似,未来可探索在 Non-IID 数据分布下联邦组间聚合对模型收敛的影响,并设计相应的加权聚合或正则化策略。
  • 与梯度压缩和通信重叠技术结合:将本文的松耦合聚合与梯度量化、稀疏化、通信计算重叠等技术结合,进一步降低组内和组间通信开销,实现端到端的训练加速。

一句话总结 ​

【一句话总结】受联邦学习启发,将大规模分片数据并行解耦为松耦合联邦组,通过组间低频 FedAvg 聚合大幅降低通信开销,实现 8.04 倍训练加速并提升模型质量。


本解读由 DeepSeek AI 自动生成,仅供参考。

Built with curiosity and a little stardust.