SandAI-org/MagiAttention
A Distributed Attention Towards Linear Scalability for Ultra-Long Context, Heterogeneous Data Training
解决的问题
MagiAttention 解决了在分布式训练中,超长上下文和异构掩码模式下注意力机制的扩展性挑战。旨在提供线性可扩展性和高性能,降低通常与上下文并行(CP)训练设置相关的计算和通信开销。
工作原理
通过以下关键技术实现分布式注意力机制:
- Flex-Flash-Attention (FFA):一种通用的注意力掩码公式化(
AttnSlice)和定制内核,支持多种掩码类型的紧凑表达和高效的分布式分区。 - 计算负载均衡:采用细粒度的块级分片策略和调度求解器,确保 CP 排名之间的负载均衡。
- 零冗余通信:用新型的
GroupCast和GroupReduce原语替代传统的环形 P2P 通信,消除前向和反向传播中的冗余通信量。 - 自适应多阶段重叠:调度计算与通信,隐藏延迟并最大化 GPU 利用率。
适用人群
专为训练超长上下文大模型的研究人员和开发者设计,例如自回归视频生成模型(如 Magi-1),以及需要与 Megatron-LM、PyTorch FSDP 和 HuggingFace Transformers 等框架集成的场景。
核心亮点
- 线性可扩展性:在 H100 和 B200 GPU 上的分布式训练环境中实现近乎线性的可扩展性。
- 广泛硬件支持:在 Hopper GPU 上性能媲美 Flash-Attention 3,并通过 Flash-Attention 4 早期支持 Blackwell GPU。
- 框架集成:可轻松集成至 Megatron-LM、PyTorch FSDP 和 HuggingFace Transformers。
- 高级掩码支持:原生支持多种重叠的注意力掩码模式。
- 注意力池:包含对分布式可学习注意力池机制的扩展。
相关
- 项目
- 项目
- 项目
- Dispatch
- 项目