HyperParallel-FSDP:面向昇腾超级Pod的拓扑感知全分片训练与布局驱动Muon优化器
HyperParallel-FSDP: Topology-Aware Fully Sharded Training with Layout-Driven Muon on Ascend SuperPods
- Zhejiang University(浙江大学)
- Huawei Technologies Co., Ltd(华为技术有限公司)
机构由 AI 辅助整理,请以论文原文为准。
AI总结:
针对超级Pod拓扑,提出HyperParallel-FSDP,通过双模式DTensor、拓扑感知FSDP和布局驱动Muon,实现505B MoE模型高效训练,步骤时间较FSDP2降低29.7%。
AI中文摘要:
声明式SPMD编程使用张量分片描述来驱动分布式执行,将并行化与模型代码分离。然而,所评估的PyTorch DTensor栈在autograd之下分派每个算子,导致重复的分派和元数据开销,且缺乏廉价的端到端验证路径。现有的FSDP和分布式Muon实现也与两级超级节点拓扑不匹配:FSDP依赖显式的参数打包和解包,而Muon的整矩阵正交化与参数分片冲突。我们观察到,分布式张量只需在autograd之上的张量API边界表达分片语义,从而允许微分和内核在普通张量上操作。基于这一见解,我们提出了HyperParallel-FSDP,其特点包括:(1)双模式DTensor执行,使用单一分片计划同时支持生产模式(一次性布局解析,无稳态分派开销)和验证模式(端到端元数据传播、快速失败检查和梯度等价性测试);(2)拓扑感知FSDP,具有超级节点内零拷贝集合通信、超级节点间融合归约以及跨层反向流水线,避免在慢速链路上等待;(3)布局驱动的分布式Muon,具有从分片派生的通信组、去重正交化和形状融合的Newton-Schulz迭代。在Atlas 900 A3 SuperPoD上,HyperParallel-FSDP从16个die扩展到384张卡(768个rank),为505B参数的MoE模型维持421k tokens/s的吞吐量,而FSDP通信仅占步骤时间的2.9%。与PyTorch FSDP2相比,平均步骤时间减少29.7%,与Megatron DDP相比减少25.5%,在1000步内Pearson相关系数高于0.999997。分布式Muon在profiler步骤时间上比竞争系统提升5.4-16.0%。源代码可在该https URL获取。
英文摘要:
Declarative SPMD programming uses tensor sharding descriptions to drive distributed execution, separating parallelization from model code. However, the evaluated PyTorch DTensor stack dispatches every operator below autograd, incurring repeated dispatch and metadata costs, while lacking an inexpensive end-to-end validation path. Existing FSDP and distributed Muon implementations also mismatch two-tier supernode topologies: FSDP relies on explicit parameter packing and unpacking, and Muon's whole-matrix orthogonalization conflicts with parameter sharding. We observe that distributed tensors need only express sharding semantics at the tensor API boundary above autograd, allowing differentiation and kernels to operate on plain tensors. Based on this insight, we present HyperParallel-FSDP, featuring: (1) dual-mode DTensor execution, using one sharding plan for both a production mode with one-time layout resolution and no steady-state dispatch overhead, and a validation mode with end-to-end metadata propagation, fail-fast checks, and gradient-equivalence testing; (2) topology-aware FSDP, with zero-copy intra-supernode collectives, fused inter-supernode reduction, and a cross-layer backward pipeline that avoids waits on slow links; and (3) layout-driven distributed Muon, with sharding-derived communication groups, deduplicated orthogonalization, and shape-fused Newton-Schulz iterations. On Atlas 900 A3 SuperPoD, HyperParallel-FSDP scales from 16 dies to 384 cards (768 ranks), sustaining 421k tokens/s for a 505B-parameter MoE while FSDP communication uses 2.9% of step time. It reduces mean step time by 29.7% versus PyTorch FSDP2 and 25.5% versus Megatron DDP, with Pearson correlation above 0.999997 over 1,000 steps. Distributed Muon improves profiler step time by 5.4-16.0% over competing systems. Source code is available at https://atomgit.com/mindspore/hyper-parallel.