中文

SimpleFSDP:利用 torch.compile 实现更简单的全分片数据并行

分布式、并行与集群计算 2024-11-07 v2 人工智能

摘要

大模型的分布式训练消耗巨大的计算资源,并且需要大量的工程工作来组合各种训练技术。本文介绍了 SimpleFSDP,一个基于 PyTorch 原生编译器的全分片数据并行(FSDP)框架,该框架实现简单,易于维护和组合,允许完整的计算-通信图追踪,并通过编译器后端优化带来性能提升。SimpleFSDP 的新颖之处在于其独特的、对 torch.compiletorch.compile 友好的集体通信实现,使用了现有的 PyTorch 原语,即参数化、选择性激活检查点和 DTensor。它还首次在 TorchInductor 后端中实现了中间表示(IR)节点分桶和重排序,以实现有效的计算-通信重叠。因此,用户可以应用上述优化来自动或手动包装模型组件,以最小化通信暴露。在 Llama 3 模型(包括超大的 405B 模型)上使用 TorchTitan 对 SimpleFSDP 进行的广泛评估表明,与其他分布式训练技术组合时,与最广泛采用的 FSDP2 急切执行框架相比,内存占用最多减少 28.54%,吞吐量最多提升 68.67%。

关键词

引用

@article{arxiv.2411.00284,
  title  = {SimpleFSDP: Simpler Fully Sharded Data Parallel with torch.compile},
  author = {Ruisi Zhang and Tianyu Liu and Will Feng and Andrew Gu and Sanket Purandare and Wanchao Liang and Francisco Massa},
  journal= {arXiv preprint arXiv:2411.00284},
  year   = {2024}
}