DiT-MoE
将MoE稀疏化技术引入扩散变换器,实现16B参数 DiT 的高效训练与推理
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
将MoE稀疏化技术引入扩散变换器,实现16B参数 DiT 的高效训练与推理
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。

图1:DiT-MoE 整体框架图 — 稀疏 MoE 层替代标准 FFN,多专家路由实现高效超大规模扩散模型
2024年,图像生成领域迎来了一场静默的革命。当 Stable Diffusion 3、FLUX 等模型相继将参数量推向数十亿级别时,一个根本性的问题浮出水面:稠密(Dense)扩散变换器的计算成本随参数量线性增长,小团队根本无法承担"把模型做大"的算力代价。
正是这一背景下,北京大学的研究团队提出了 DiT-MoE(Diffusion Transformers with Mixture of Experts),将大语言模型领域已经验证的 Mixture of Experts(MoE)稀疏化技术引入扩散模型训练。该工作发表在 arXiv(arXiv:2407.11633),代码由 Fei Zhengcong 等研究者开源,成为扩散模型 Scaling 研究的里程碑式参考实现。
这一思路的核心逻辑并不复杂:与其让所有参数每次都参与计算,不如在每一层设置多个"专家"网络(前馈神经网络),让每个 token 只激活最相关的 2 个专家。 如此一来,16B 参数的模型实际推理成本可以降至接近 1B 级别,同时保持与稠密模型相当的生成质量。
DiT 最早由 Facebook Research 提出,是一种将标准 ViT 架构与扩散概率模型相结合的图像生成范式。其核心是把图像切分成固定大小的 patch,通过自注意力机制在 latent 空间中进行去噪生成。
DiT-MoE 在此基础上做了以下关键改造:
1. MoE Gate 路由机制
在 DiT 的每个 Transformer Block 中,标准 FFN(前馈网络)被替换为 MoE 层。核心实现在 models.py 中的 MoEGate 类:
class MoEGate(nn.Module):
def __init__(self, embed_dim, num_experts=16, num_experts_per_tok=2, aux_loss_alpha=0.01):
self.top_k = num_experts_per_tok
self.n_routed_experts = num_experts
self.scoring_func = 'softmax'
self.weight = nn.Parameter(torch.empty((self.n_routed_experts, self.gating_dim)))
Gate 网络接收 hidden states,计算每个 expert 的 softmax 得分,然后选择 top-k 个 expert 处理当前 token。同时加入 auxiliary loss(辅助损失),确保路由负载均衡,防止少数 expert 过度集中。
2. 专家并行策略
训练脚本支持 PyTorch DDP(分布式数据并行)和 DeepSpeed ZeRO-2/ZeRO-3 多种分布式策略。ZeRO-3 配置配合 offload 技术,使得在有限显存下训练超大规模模型成为可能。配置示例在 config/zero2.json 和 config/zero3.json。
3. 整流流(Rectified Flow)训练
DiT-MoE 同时支持 DDIM 和整流流两种采样方式。整流流来自 FLUX 和 SD3 的实践,将噪声到数据的回归路径简化为线性插值,收敛更快且生成质量更好。训练命令通过 --rf True 参数启用。
4. 专家专业化分析工具
项目在 analysis/ 目录下提供了一套完整的专家行为分析工具:
expert_data.py:统计不同类别条件下各专家的激活频率heatmap_class.py、heatmap_patch.py、heatmap_step.py:可视化专家选择热力图项目提供了从 DiT-MoE-S 到 DiT-MoE-G 共 5 个规模的预训练权重,全部托管在 HuggingFace(feizhengcong/DiT-MoE):
| 模型 | 专家数 | 每token激活专家数 | 采样方式 | 建议显存 |
|---|---|---|---|---|
| DiT-MoE-S/2-8E2A | 8 | 2 | DDIM | 40GB+ |
| DiT-MoE-S/2-16E2A | 16 | 2 | DDIM | 80GB+ |
| DiT-MoE-B/2-8E2A | 8 | 2 | DDIM | 160GB+ |
| DiT-MoE-XL/2-8E2A | 8 | 2 | Rectified Flow | 640GB+ (多卡) |
| DiT-MoE-G/2-16E2A | 16 | 2 | Rectified Flow | 极大规模集群 |
生成 256×256 图像推理命令(以 XL 模型为例):
python sample.py \
--model DiT-XL/2 \
--ckpt /path/to/model \
--vae-path /path/to/vae \
--image-size 256 \
--cfg-scale 1.5
推理默认使用 torch.float16,对 XL/G 级别模型强制启用 Flash Attention 加速。
普通用户的困境:DiT-MoE 不是一个开箱即用的"AI画图工具",它是一套完整的训练框架。即便是推理使用,也需要:
研究者的工作流:这才是 DiT-MoE 的主战场。如果你研究 MoE 在生成模型中的应用、探索专家专业化规律、或尝试将 MoE 迁移到自己的扩散模型,该项目提供了:
核心代码文件一览:
models.py:DiT-MoE 模型定义(MoE Gate、DiT Block、DiT Model)train.py:PyTorch DDP 单机多卡训练train_deepspeed.py:DeepSpeed 分布式训练(支持 ZeRO-2/3)sample.py:推理脚本diffusion/:扩散/整流流采样器analysis/:专家行为分析工具DiT-MoE 并非没有代价。稀疏化带来的主要挑战是负载均衡:如果某些 expert 长期被选中而另一些闲置,不仅浪费参数,还会影响模型整体容量上限。项目中通过 aux_loss_alpha 参数约束,但这本身是一个需要精细调优的平衡。
另外,推理阶段的多 expert 动态路由引入了额外延迟。尽管激活参数少,但路由逻辑在 onnx/TensorRT 部署时并不友好,工业级部署仍需额外工程优化。
最重要的是,没有 Dockerfile 或一键启动脚本意味着所有环境配置都需要手动完成,对非分布式训练专业的研究者来说门槛较高。
DiT-MoE 的出现,标志着扩散模型正式进入 "参数量可以无限增长,实际算力需求保持可控" 的新阶段。
从更长远的视角看,这一工作与 LLM 领域的 DeepSeek-MoE、Mixtral-8x7B 一脉相承,代表了 "稀疏化激活"作为 Scaling 核心手段 的广泛适用性。当 Stable Diffusion 4、FLUX 2 等商业模型持续扩大规模时,DiT-MoE 的开源实现为学术社区提供了一把理解底层机制的钥匙。
对于 AI 开发者而言,DiT-MoE 的核心参考价值在于:如何设计 Gate、如何处理负载均衡、如何在多卡环境下高效训练。这些经验可以直接迁移到其他模态的 MoE 扩散模型研究中。