如果你最近在研究 DiT(Diffusion Transformer),大概率会碰到一个很实际的困惑:论文里那张 FID 对“计算规模”的曲线,横坐标到底是怎么算出来的?同样是 DiT-XL/2,有人算出来一张图一次前向只要两百多 GFLOPs,有人换一种口径却算出几十 TFLOPs,数字能差几十倍。这其实不是模型变了,而是“在哪个空间算”“按哪个口径算”没有先统一。
这篇文章就从一个做项目落地的角度,把 DiT 计算规模这事完整拆开讲一遍。我会把一个 DiT 的前向 FLOPs 公式手推出来,拿 DiT-XL/2 的实际配置算一遍,再补上训练、采样、显存这些工程上真正关心的口径,最后给几个我在实际排查计算量时踩过的坑。不管你是要复现 DiT,还是想估算自己改出来的模型在 A100/H100 上能不能跑,这套方法都够用。
1. 先分清:DiT 计算规模到底在算哪几个量
我最早被 DiT 计算规模搞晕,就是因为把几个不同口径的量混在一起比了。这里先把口径钉死,后面所有公式才不会乱套。
| 口径 | 含义 | 常见数量级 |
|---|---|---|
| 参数量(Params) | 模型里所有可学习权重的个数 | 33M ~ 675M |
| 前向 FLOPs | 一张图做一次去噪推理需要的浮点运算量 | 12G ~ 237G |
| 训练 FLOPs / 样本 | 训练时处理一个样本的一次优化步 | 大约为前向的 3 倍 |
| 采样 FLOPs / 图 | 生成一张图需要跑完整采样链路 | N个采样步 × 前向 FLOPs |
这几个概念里有两条最容易乱的分界线。
第一条是“前向 FLOPs”和“训练 FLOPs”的区别。很多人以为扩散模型训练时要像采样那样跑完整个去噪链,实际上训练时每个样本是随机抽一个时间步 t,只跑一次加噪样本的前向和反向传播。所以训练成本大约等于“前向 × 3”,这里的 3 指的是反向传播大约是前向的 2 倍。这个细节如果搞错,你估出来的总训练算力会直接膨胀几十倍。
第二条是“理论 FLOPs”和“实际耗时”的区别。FLOPs 只能告诉你计算量的数学上限,不能直接等成 GPU 跑出来的时间。尤其是当序列长度 T 比较小的时候,矩阵乘法形状也跟着变小,A100 这种卡根本喂不满,实际吞吐会比理论值低很多。所以我后面会专门说,算完 FLOPs 之后怎么做工程层面的折算。
还有一个前置条件必须确认:DiT 是在 VAE 的 latent 空间里跑,还是在 pixel 空间里跑。DiT 论文默认用的是 latent 空间,256×256 的图经过 VAE 下采样 8 倍,变成 32×32 的 latent 特征图,再按 patch size 切块。这个选择对 T 的影响是数量级的,因为序列长度 T 等于(latent 分辨率 / patch size)²,而 FLOPs 里面大头是 T 的二次项和 d 的二次项。后续所有计算,我都先明确告诉你是基于 latent 空间还是 pixel 空间。
2. DiT 前向 FLOPs 的推导:一个 DiT Block 里都有哪些乘法
想快速估算 DiT 的计算规模,不需要真的拿 profiler 去跑一遍,先把一个 DiT Block 的结构拆开,逐段数乘加次数就行。我按“乘加各算一次,统称 1 FLOP”的常见口径来算。这里全部按前向计算来推,不考虑反向。
先约定几个符号:
L:DiT Block 数量d:模型宽度,也就是 transformer 的 embedding 维度T:输入 token 数,也就是 patch 序列长度p:patch sizeC:输入通道数。latent 空间通常 C=4,pixel 空间 C=3h:注意力头数,多头的设置只影响线性层的分块方式,不改变总 FLOPs 公式
2.1 patch embedding 与输出解码层
DiT 先要把二维的 latent 特征图切成 patch,然后线性投影成 d 维向量。每个 patch 展平后长度是C × p²,从C × p²映射到 d,是一个全连接层:
patch embedding FLOPs = 2 × T × d × C × p²对应的输出解码层也差不多,从 d 再映射回C × p²,成本同样是2 × T × d × C × p²。
这部分通常只占总量的一小部分。拿 DiT-XL/2 来说,T=256, d=1152, C=4, p=2,嵌入层两端的 FLOPs 大约是 940 MFLOPs,相比单 block 的 8.45 GFLOPs 完全可以忽略,但参数别忘了,它占显存。
2.2 注意力模块:QKV、attention map、加权求和
DiT 用的是标准的非因果 multi-head self-attention,没有 mask。这一步是很多估算公式最容易丢东西的地方。
首先是把输入序列同时映射成 Q、K、V。这个操作通常实现成一个d → 3d的线性层:
QKV 投影 FLOPs = 2 × T × d × 3d = 6Td²然后计算 Q 和 K 的注意力分数。Q 的形状是T × d,K 的转置是d × T,矩阵乘法一次是T × d × T,乘加各算一次:
attention score FLOPs = 2 × T × d × T = 2T²dsoftmax 本身是逐元素操作,不做矩阵乘法,通常可以忽略。
接下来用 attention map 去加权 V,Q 的加权注意矩阵T × T乘以 V 矩阵T × d:
attention 加权 FLOPs = 2 × T × T × d = 2T²d最后还有一个注意力输出的线性投影d → d:
output projection FLOPs = 2 × T × d × d = 2Td²所以注意力部分汇总如下:
注意力模块 FLOPs = 6Td² + 2T²d + 2T²d + 2Td² = 8Td² + 4T²d这里的4T²d就是很多人会用错的地方。如果直接用视觉 Transformer 里流传的“每层 12Td²”这种速算公式,会把注意力分数和注意力加权漏掉。对 ViT 这种 patch 数量只有一两百的模型,漏掉问题不大;但 DiT 一旦处理高分辨率 latent,序列长度很容易涨到几千甚至上万,T²d就会成为主导项。
2.3 MLP 与 AdaNorm:别把 4 倍宽度展开漏掉
DiT Block 的 MLP 基本沿用 Transformer 的 MLP 结构,隐藏层宽度通常取4d,激活函数用 GELU。
第一个线性层从 d 升到 4d:
MLP 第一层 FLOPs = 2 × T × d × 4d = 8Td²第二个线性层从 4d 压回 d:
MLP 第二层 FLOPs = 2 × T × 4d × d = 8Td²MLP 合计:
MLP FLOPs = 16Td²这里我要提醒一下,虽然激活函数 GELU 也有计算成本,但在比较模型计算规模时,行业里通常只统计矩阵乘法和卷积这类主导算子。逐元素的 GELU、LayerNorm、AdaNorm 的 scale/shift,量级都是O(Td),跟Td²相比小两三个数量级,日常估算直接忽略。
2.4 合并成速算公式
现在把一个 DiT Block 的前向 FLOPs 汇总:
FLOPs_block = 注意力模块 + MLP = (8Td² + 4T²d) + 16Td² = 24Td² + 4T²d全模型前向:
FLOPs_forward = L × (24Td² + 4T²d) + 2 × T × d × C × p²这个公式是我做计算规模估算时最常用的主力公式。有两个点需要说明:
- 这里用的是“乘加各算 1 FLOP”的口径。如果某些工具返回的是 MACs(乘累加次数),那转换成 FLOPs 需要再乘以 2,否则数字会对不上。
- AdaNorm 和 modulation 的线性投影我没有放进主公示,因为它们只在时间步 embedding 上做线性变换,不跟在 batch、token 维度上展开,总占比通常小于 1%。
3. 拿 DiT-XL/2 手算一遍:从配置到最终数字
先看 DiT 论文里常用的四个模型配置,这些数字在后面所有手算中都会用到。
| 模型 | 层数 L | 宽度 d | 多头数 | 参数量 |
|---|---|---|---|---|
| DiT-S/2 | 12 | 384 | 6 | 33M |
| DiT-B/2 | 12 | 768 | 12 | 131M |
| DiT-L/2 | 24 | 1024 | 16 | 458M |
| DiT-XL/2 | 28 | 1152 | 16 | 675M |
如果你用的是 256×256 的输入图像,并且走 VAE latent 空间,下采样 8 倍后 latent 是 32×32,通道数 C=4。当 patch size=2 时:
T = (32 / 2)² = 256对 DiT-XL/2 来说,d=1152, L=28。先算单个 block:
24Td² = 24 × 256 × 1152 × 1152 = 8.15e9 4T²d = 4 × 256 × 256 × 1152 = 3.02e8 FLOPs_block ≈ 8.45e9全模型:
FLOPs_forward ≈ 28 × 8.45e9 ≈ 2.37e11也就是大约 236 GFLOPs。这样一个数才算出来。对应地,DiT-S/2、DiT-B/2、DiT-L/2 的前向 FLOPs 分别大约是 12 GFLOPs、46 GFLOPs、161 GFLOPs。
现在把口径变化一下,你会看到非常夸张的差异。还是 DiT-XL/2:
| 口径 | 结果 |
|---|---|
| 单样本前向一次 | 236 GFLOPs |
| 训练一个样本优化一步 | 约 708 GFLOPs |
| 50 步采样生成一张图 | 11.8 TFLOPs |
| 50 步采样并且开启 classifier-free guidance | 23.6 TFLOPs |
| 如果把模型搬到 pixel 空间跑同一张 256×256 图 | 49.3 TFLOPs |
注意最后一行:如果直接在像素空间跑,T = (256 / 2)² = 16384,序列长度是 latent 空间的 64 倍,FLOPs 直接从 0.236 T 飙到 49.3 T,这还是在没考虑 VAE 编解码成本的情况下。所以你看任何文章讨论 DiT 计算量,第一反应一定是去确认它说的是 latent 空间还是 pixel 空间,否则数字根本没可比性。
4. 为什么计算规模对“分辨率”会二次爆炸:复杂度敏感性分析
DiT 的计算规模不是线性变化的,它的脾气很鲜明:T和d对计算量的影响完全不同,而且这个差异会随着输入尺寸变化突然反转。
看公式24Td² + 4T²d,这两项一个是线性于 T,一个是平方于 T。什么时候 attention 的平方项开始主导?让两项相等:
4T²d = 24Td² T = 6d也就是说,当序列长度 T 超过大约 6 倍的模型宽度 d 时,注意力矩阵计算的 FLOPs 就会反超 MLP 和 QKV 投影,成为新的主要矛盾。
拿 DiT-XL/2 来说,d=1152,临界点大约是T=6912。在 latent 空间里,256×256 图像只有 256 个 token,离临界点非常远,这时候注意力平方项占比只有大约 3.6%,计算规模主要由 d 决定。但是一旦跑到高分辨率场景,比如 latent 变成 128×128,T=(128/2)²=4096还不算太夸张,再上去到 256×256,T=16384,注意力平方项就会占到七成以上。这也是为什么高分辨率 DiT 一定要做窗口注意力或者改成 masked attention,否则计算量会不可控。
patch size 的影响就更直接了。latent 分辨率不变的情况下,p 每翻一倍,T 变成原来的四分之一,attention 平方项理论上变成十六分之一。我拿 DiT-XL/2 在 32×32 latent 上算过:
| patch size | T | 前向 FLOPs |
|---|---|---|
| 2 | 256 | 236.8 GFLOPs |
| 4 | 64 | 57.7 GFLOPs |
| 8 | 16 | 14.3 GFLOPs |
从 p=2 换到 p=8,计算量直接砍掉 16 倍以上。代价也很明确,patch 太大等于信息在入口就被暴力压缩,图像细节尤其是高频结构会丢失。DiT 论文里实验也验证了,patch size=8 时计算效率很好,但 FID 会变差;patch size=2 的细节保留能力最强,但最贵。实际项目里做这个权衡时,我建议先看你的目标分辨率,如果画面里有大量小物体或文字,就不要为了省算力一上来就 p=4 或者 p=8。
还有一个常见做法是“分辨率分层”:先在低分辨率 latent 上用大步长 patch 走流程,最后用一个小步长 patch 的高分辨率 DiT 做精修。这样总计算量不是简单相加,因为高分辨率分支的T²会被限制在一个小范围内,整体成本可控。
5. 从单次前向到训练和采样总消耗
5.1 训练总算力估算
如果你要估整个训练任务,可以用下面这条公式:
训练总 FLOPs ≈ 总训练样本数 × 迭代轮数 × 3 × FLOPs_forward这里的 3 就是“1 次前向 + 约 2 次反向”的训练开销系数。扩散模型训练时每个样本每次只随机抽一个时间步,所以不需要乘扩散步数,这一点我再强调一遍。
举一个真实的项目估算例子。假设 DiT-XL/2 在 256×256 latent 空间训练,FLOPs_forward=236 GFLOPs,训练集 128 万张图,训练 300 个 epoch:
总样本数 = 1.28M × 300 = 3.84e8 训练总 FLOPs = 3.84e8 × 3 × 236e9 = 2.72e20也就是大约 272 EFLOPs。如果单卡有效吞吐按 150 TFLOPs 算,单卡需要大约 500 个小时。实际跑出来往往要再放宽 1.5 到 2 倍,因为 DiT 在 latent 空间的矩阵形状偏小,GPU 利用率很难跑到 50% 以上。
5.2 采样总算力估算
采样就完全是另一套算法了。生成一张图要迭代 N 个去噪步,每步都要跑一次前向,所以:
采样 FLOPs / 图 = N × FLOPs_forward × guidance_倍数guidance_倍数只有在用 classifier-free guidance 时才需要乘。CFG 的一次采样里,每个 step 都要同时跑条件模型和无条件模型,所以相当于乘 2。50 步 DDIM 就已经很贵了,如果用 Euler、Heun 这类更高阶的 solver,实际 step 数还会叠加。
这也是为什么 DiT 这类模型实际部署时,主流方案都是蒸馏或者步数压缩。一个 DiT-XL/2 在 256×256 上,50 步加 CFG 的生成成本大约是 23.6 TFLOPs,对比 Stable Diffusion 那一类 U-Net 同样是 50 步加 CFG 的 512×512 图像,成本已经不在一个量级了。部署前不做步数压缩,推理费用会非常难看。
5.3 显存估算:公式之外的工程账本
FLOPs 算完还要看显存,否则你会出现“算力够但卡放不下”的窘境。显存这块有两条快速估算线。
第一条线是参数量带来的权重和优化器显存。以 fp16 混合精度训练为例,DiT-XL/2 的 675M 参数:
- fp16 权重:
675M × 2B = 1.35 GB - fp16 梯度:同样 1.35 GB
- fp32 的 master weight:
675M × 4B = 2.7 GB - Adam 的两个 fp32 状态:
675M × 8B = 5.4 GB
单模型一轮训练,光权重和优化器就要吃掉大约 10.8 GB。所以 DiT-XL/2 想在 24 GB 的消费级卡上微调,必须用 LoRA/QLoRA 一类的参数高效微调方案,否则 activation 根本放不下。
第二条线是 activation 显存。训练时要留存部分中间激活做反向传播,DiT 的 activation 大头有两块:一个是每层T × d的特征,一个是某些实现里没做 flash attention 的batch × head × T × T注意力矩阵。后者在 T=256 时不大,但当 T 进入几千级别后,一张草图就能把显存写满。所以我做高分辨率训练时,一定会开 flash attention,并且配合 activation checkpointing,否则就算 FLOPs 再低,显存也会先爆。
6. 公式、工具和三个最容易踩的坑
最终我一般会把公式直接写成一个 Python 函数,方便在调整分辨率、patch size、模型宽度时快速出数。下面这段是我一直沿用的小工具:
def dit_forward_flops( d_model: int, layers: int, token_len: int, patch: int = 2, latent_channels: int = 4, ) -> float: """估算 DiT 单图单次前向的 FLOPs,乘加各算一次。""" block_flops = ( 24 * token_len * d_model * d_model + 4 * token_len * token_len * d_model ) embedding_flops = ( 2 * token_len * d_model * latent_channels * patch * patch ) return layers * block_flops + 2 * embedding_flops if __name__ == "__main__": token_len = (32 // 2) ** 2 # 32x32 latent, patch=2 flops = dit_forward_flops( d_model=1152, layers=28, token_len=token_len, patch=2, ) print(f"{flops / 1e9:.2f} GFLOPs")如果想用工具做交叉验证,常见选择有fvcore.nn.FlopCountAnalysis、thop、pytorch profiler。但我实测下来,这类工具对 Transformer 的 attention 部分统计并不一致。有的工具把torch.matmul Q @ K.T正确算进去了,有的直接当成普通乘加算,有的还会漏掉 softmax 后面的attn @ V。所以工具只能做二次确认,不能当唯一答案,还是要以手推公式为准。
我建议的验证办法很简单:先拿一个很小配置,比如 DiT-S/2,用 profiler 跑一个真实输入,把 profiler 给的总 FLOPs 和我公式算出来的做比较,差异在 10% 以内就说明口径基本一致。如果差异很大,多半是工具把某个重复计数了,或者把非矩阵乘法的操作也算进去了。
最后是我项目里总结的三个高频坑,每个都曾经让我对不上数字:
- 第一个坑是把 MACs 当 FLOPs。很多工具返回的是 MACs,一个乘加算 1,换算 FLOPs 时要乘以 2。
- 第二个坑是拿 ViT 的
12Td²速算公式套 DiT。ViT 输入 patch 少,attention 平方项占比低,但 DiT 的 T 经常远大于临界点6d,这时候套那个公式会严重低估。 - 第三个坑是训练时乘了扩散步数。训练每样本随机抽一个 t,只跑一次模型;真正反复跑模型的是采样阶段。
我个人现在复盘 DiT 计算规模的习惯是:先把“latent 还是 pixel”“前向还是训练”“FLOPs 还是 MACs”这三个问题写在纸最上面,再开始套公式。把口径锁死之后,DiT 的计算规模其实很好算,它比 U-Net 那种结构清晰的卷积网络更容易手推,因为所有成本都集中在24Td² + 4T²d这一个式子里。把这一个式子用熟,以后不管是 Scaled DiT、DiT 变体还是视频 DiT,核心估算都能快速给出靠谱答案。