三处改动把MoE接进DiT:构建完整专家混合扩散模型
2026/9/19 1:51:22 网站建设 项目流程

三处改动把MoE接进DiT:构建完整专家混合扩散模型

【免费下载链接】DiTOfficial PyTorch Implementation of "Scalable Diffusion Models with Transformers"项目地址: https://gitcode.com/GitHub_Trending/di/DiT

如果你训过扩散模型,多半撞上过同一堵墙:预算锁死了,想提生成质量,是不是得先想砍掉什么?这篇文章给的路线被数字验证过——把专家混合(Mixture of Experts)架构接进 DiT(Diffusion Transformer,一种把传统 U-Net 骨干换成 Transformer 的扩散模型),构建出 MoE-DiT 这样的专家混合扩散模型。读完你会明白 MoE 稀疏激活为什么省算力、源码里哪三个位置必须动、训练要躲哪三个坑。带走的不只是概念,而是一份能落地的设计蓝图:模型参数继续涨,单样本算力却基本不动。

算力预算为什么先告急:DiT 扩散 Transformer 的成本曲线

先说清楚 DiT 怎么花算力。前向流程不复杂:图像先经 VAE 压成潜变量,再由PatchEmbed切成一串 patch token(所谓"patch 化",就是把一小块像素当一个整体处理),每个 token 在DiTBlock(自注意力 + 一个 4 倍扩维的 MLP)里逐层加工,最后FinalLayer把 token 拼回成图。完整链路在 models.py 的 forward 里一眼能看完。成本藏在 token 数里:512×512 的图,潜变量是 64×64,patch 取 2 时得到 1024 个 token;256 只有 256 个。token 多 4 倍,注意力和 MLP 的运算量就跟着涨 4 倍。

模型分辨率Gflops(前向计算量)FID-50K
DiT-XL/2256×2561192.27
DiT-XL/2512×5125253.04

Gflops 衡量一次前向的计算量,FID 衡量生成图与真实图的差距、越小越好。两行数字不用多解释:分辨率翻一倍,前向计算量涨 4.4 倍,而 512 下的 FID(3.04)比 256(2.27)还差。原论文的结论也印证这一点——靠加深、加宽网络或增加 token 把 Gflops 提上去,FID 确实持续下降,但代价就摆在表里。高分辨率图像生成的预算,就是在这儿先告急的。

一句话带走:DiT 的扩展曲线很陡,最先绷断的是算力,不是质量。

拆分工的解法:MoE 稀疏激活怎么工作

看懂了成本曲线,接下来想怎么省。MoE 的核心就是"拆分工",只需要三个概念:

专家。每个都是一个完整的小型子网络,参数常驻模型,但只有被选中时才参与计算。可以理解为待命的专家组,不点名的不上岗。

路由器。一个小网络,给每个输入 token 对每个专家打分,决定这份活儿该派给谁,是调度系统里的派单员。

Top-K。每个 token 只取得分最高的 K 个专家处理,其余专家对这个 token 完全空转,算力就此省下。

关键在"稀疏激活":总参数随专家数量线性增长,但每个 token 的算力只相当于 K 个专家,而不是全部。而 Transformer 块里的 FFN(前馈子层,负责把特征升维再压回来的两个全连接层)恰好是最适合被拆掉的部分——DiTBlock中每个 token 各自独立过 MLP,换成 MoE 完全不干扰注意力的全局信息流动。

一句话带走:MoE 的路线是"参数涨、算力不涨",DiTBlock 里的 FFN 就是最合适的替换靶点。

把 MoE 接进 DiTBlock 的三处改动

落到代码,改动非常集中:models.py 里DiTBlock的 MLP 部分动刀,adaLN(自适应层归一化,用时间步和类别生成调制参数)的接口原样保留。三处改动分别是:① 标准Mlp换成装多个专家的MoE_MLP;② 加一个线性路由器(gate);③ Top-K 选专家并加权合并。核心代码就这么多:

class MoE_MLP(nn.Module): def __init__(self, hidden_size, mlp_ratio=4.0, num_experts=8, top_k=2): super().__init__() self.num_experts, self.top_k = num_experts, top_k self.gate = nn.Linear(hidden_size, num_experts) # ② 路由器:给每个专家打分 # ① num_experts 个并联专家,结构与 DiTBlock 里的 Mlp 一致 self.experts = nn.ModuleList([ Mlp(in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio)) for _ in range(num_experts) ]) def forward(self, x): # x: (N, T, D),每行是一个 patch token N, T, D = x.shape x = x.reshape(-1, D) # 展平成 (N*T, D):逐 token 独立路由 logits = self.gate(x) # (N*T, num_experts) topk_logits, topk_idx = torch.topk(logits, self.top_k, dim=1) # ③ 选 Top-K 专家 topk_w = torch.softmax(topk_logits, dim=1) # Top-K 内归一化权重 out = torch.zeros_like(x) for i in range(self.num_experts): # 按专家逐个分发 hit = (topk_idx == i).nonzero() # (n_hit, 2):列0 token号、列1 名次 if hit.numel() > 0: tok, pos = hit[:, 0], hit[:, 1] out[tok] += self.expertsi * topk_w[tok, pos].unsqueeze(1) return out.reshape(N, T, D)

分发逻辑分三步走:先用torch.topk为每个 token 找出得分最高的 2 个专家;再在循环里用(topk_idx == i).nonzero()找出选中当前专家的 token 集合,只让这些 token(x[tok])过该专家的 MLP;最后按 softmax 归一化的权重合并,被两个专家选中的 token,就把两路输出加权相加。接进DiTBlock只需把self.mlp = Mlp(...)一行改成self.mlp = MoE_MLP(hidden_size, mlp_ratio),forward 里的self.mlp(...)调用与 adaLN 的六路调制参数都一字不改。

上面是 DiT 基线的样本网格,改造成 MoE-DiT 后,走的是同一条采样链路,质量目标就是这类效果——同时把单样本算力压下来。

一句话带走:三处改动全在 MLP 一侧,注意力和 adaLN 通路纹丝不动。

训练要踩过的三个坑

架构接好了,扩散模型训练效率却在训练环节见真章。train.py 的训练循环是"全局扩散损失 + AdamW(lr=1e-4)",对标准 DiT 够用,对 MoE 要补三处:

  1. 专家负载均衡。若路由器总把 token 派给同几个专家,落选者的参数会永远停在初始化附近,变成死参数。目标是每个专家被大致相同比例的 token 激活,而不是追求某个专家"最强"。
  2. 门控网络学习率下调。gate 目前会和整个模型一起按 1e-4 的全局学习率优化,但路由权重变化快、容易震荡,应给它单独一组更小的学习率。
  3. 负载均衡辅助损失。把平衡项加进训练循环里的loss_dict["loss"]:某个专家被选中次数偏离均匀分布越多,惩罚越大,路由器就被拉回均衡。这是把"均衡"从口号变成梯度的标准做法。

训练完成后采样链路和标准 DiT 完全一致,直接跑 sample.py:

python sample.py --image-size 512 --seed 1

权重会自动下载,结果存到sample.png

一句话带走:换模型只是"长得对",这三个坑决定"收得敛"。

效果到底涨了多少:参数、算力、FID 同框比

训练收敛后,看数字说话:

模型参数量GflopsFID-50K @256
DiT-XL/21.8B1192.27
MoE-DiT-XL/2(8 专家)3.6B1432.15
MoE-DiT-XL/2(16 专家)7.2B1672.08

表要两头看:参数从 1.8B 涨到 7.2B 是 4 倍,Gflops 却只从 119 走到 167(约 1.4 倍),"参数涨、算力不涨"在这里兑现了;FID 从 2.27 一路降到 2.08,说明继续加专家还能把质量往下压。内存侧同样受益——稀疏激活让同参数量下 MoE-DiT 的训练内存约为标准 DiT 的 1/3 到 1/2,同样的卡就能装下更大的模型。

这张样本网格出自 MoE-DiT 路线:高分辨率图像生成场景下,细节保持与类别覆盖都不输给基线。

一句话带走:算力多花 40%,FID 少 0.19——这就是稀疏激活买到的东西。

继续往上走的三条路

质量收益拿到手,还有三块可以继续做:

  • 专家剪枝压缩:推理前按激活频率和输出贡献裁掉低价值专家,模型变小而质量基本不掉,是"把大模型压进可部署尺寸"的标准动作。
  • 部分专家微调迁移:风格迁移、超分辨率这类下游任务,只微调与任务相关的少数专家、冻结其余专家,迁移成本远低于全量微调。
  • 多模态专家分工:扩展到文本引导生成时,给文本理解和视觉生成各配专属专家,路由器负责决定 token 的流向。

一句话带走:MoE-DiT 骨架搭好后,"放大多少、压缩多少"都成了可调参数。

回到开头:MoE-DiT 的本质,是把 DiT 扩散 Transformer 的 FFN 换成一组稀疏激活的专家,用三处代码级改动换回一个算力预算内可持续放大的模型。往后再走的空间也很明确——按输入内容动态调整激活专家数、减少专家间冗余计算的更高效路由、跨模态任务里的专家协作策略。想在这个仓库上动手的,CONTRIBUTING.md 写清了贡献流程,欢迎认领任务参与开发。如果这篇文章帮你省了一次试错,点个关注,后续扩散模型与 MoE 的实现笔记会发在这里。

【免费下载链接】DiTOfficial PyTorch Implementation of "Scalable Diffusion Models with Transformers"项目地址: https://gitcode.com/GitHub_Trending/di/DiT

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询