SandAI-org/MagiAttention

A Distributed Attention Towards Linear Scalability for Ultra-Long Context, Heterogeneous Data Training

解决的问题

MagiAttention 解决了在分布式训练中,超长上下文和异构掩码模式下注意力机制的扩展性挑战。旨在提供线性可扩展性和高性能,降低通常与上下文并行(CP)训练设置相关的计算和通信开销。

工作原理

通过以下关键技术实现分布式注意力机制:

  • Flex-Flash-Attention (FFA):一种通用的注意力掩码公式化(AttnSlice)和定制内核,支持多种掩码类型的紧凑表达和高效的分布式分区。
  • 计算负载均衡:采用细粒度的块级分片策略和调度求解器,确保 CP 排名之间的负载均衡。
  • 零冗余通信:用新型的 GroupCastGroupReduce 原语替代传统的环形 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
  • 项目