AIGC 基本功|DiT:用 Transformer 替掉 UNet-DiT

人工智能炼丹君
2026-09-29 / 0 评论 / 2 阅读 / 正在检测是否收录...

DiT:用 Transformer 替掉 UNet

所属方向:生成范式 | 难度:进阶 | 前置知识:DDPM 的训练目标与采样流程(ddpm)、潜空间扩散(latent_diffusion)、自注意力(attention_basics)
关键词:DiT、Diffusion Transformer、adaLN-Zero、patchify、Gflops 缩放、条件注入


01. 为什么需要它

先给三组数字,全部来自文末附录里能直接跑的脚本,或者论文 Table 4。

第一组:同算力下,参数量差 14 倍,FID 一模一样。 DiT-S/2 用 33M 参数、6.06 Gflops 做到 FID 68.40;DiT-B/4 用 130M 参数、5.56 Gflops 做到 FID 68.38。两个模型的参数量差 4 倍,算力几乎相同,结果几乎相同。而同一算力档里的 DiT-L/8,用了 459M 参数(14 倍),FID 反而掉到 118.87——差了 50 分。这三个点在图 1 的左图里被连成一条虚线。它说的是一件很反直觉的事:在扩散模型里,把算力花在「更多的 token」上(更小的 patch size)比花在「更宽的网络」上更划算。UNet 没有这个旋钮——它的多尺度结构定死了每一层的空间分辨率,你没法只调 token 数而不动别的。

第二组:DiT-XL/2 的 118.64 Gflops 里,注意力只占 3.6%。 逐 token 的线性层(qkv、输出投影、MLP)占 96.2%。这与「Transformer 的瓶颈是自注意力的平方复杂度」这个直觉是冲突的:在 T=256、d=1152 这个区间,算力的主导项是 $O(Td^2)$ 的线性层,不是 $O(T^2d)$ 的注意力。注意力要占到一半,需要 $T = 6d$,也就是 d=1152 时 T≈6900 个 token——那已经远超 256×256 图像的范围了。UNet 换掉的理由不是「注意力更强」,而是「Transformer 的算力可以干净地缩放」:改深度、改宽度、改 token 数,三个旋钮互不干扰,而且 Gflops 涨、FID 就单调降。UNet 加宽加深会同时改变感受野、下采样次数、跳连数量,你分不清是哪个在起作用。

第三组:同样 118.6 Gflops,只换条件注入方式,FID 从 25.21 掉到 19.47。 论文的 Table 4 末尾四行是同一个 DiT-XL/2 骨架:in-context 35.24、cross-attention 26.14、vanilla adaLN 25.21、adaLN-Zero 19.47。其中 vanilla adaLN 与 adaLN-Zero 的 Gflops 只差 0.08,差别包括新增残差门控与零初始化,不能把整个增益归因于初始化一个因素。这是我写这篇文章的直接动机——一个初始化技巧带来近 6 分 FID,值得逐行对照公式看清楚它到底做了什么。

图 1:DiT 的算力账本——钱花在哪,以及该多买 token 还是多买宽度

这张图要看什么:左图横轴是「一次前向的算力」,纵轴是 FID。同一条虚线上的三个点,算力接近、参数量差了几倍到十几倍,却几乎落在同一高度——说明在这个区间里决定质量的是算力怎么分配,不是参数量堆多少。图里最刺眼的是第三条虚线的左端:33M 的 DiT-S/2 和 130M 的 DiT-B/4 几乎重合在 FID 68.4,而同算力下 459M 的 DiT-L/8 反而掉到 118.87。右图回答「注意力到底占多少」:四条曲线是四个宽度,横轴是 token 数,竖虚线标出 256×256 图像在 p=2 时的位置(3.6%);要等 T 走到上千,注意力才从零头变成大头。

所以这篇文章回答三个问题:DiT 的算力账本是怎么记的、adaLN-Zero 的「恒等初始化」在数学上意味着什么、以及它换来的到底是什么。


02. 最小可用理解

三句话讲完:

  1. 把潜变量切成 patch 序列,然后接一个标准 ViT。 32×32×4 的潜变量按 p=2 切成 (32/2)²=256 个 token,每个 token 原始 16 维,线性嵌入到 d 维,加固定的 2D sin-cos 位置编码,过 N 个 block,最后一层线性投回 p²C 维再拼回空间形状。patch size p 是 DiT 沿用 ViT 的缩放旋钮:p 减半,token 数翻四倍,算力至少翻四倍,参数量几乎不动。

  2. 条件(时间步 + 类别)不进序列,而是变成每个 block 的 LayerNorm 参数。 具体做法是调制:先算出 $c = t_{\text{emb}} + y_{\text{emb}}$,再用一个 $\text{SiLU} \to \text{Linear}(d, 6d)$ 把它变成 6 段 d 维向量,分别是注意力分支与 MLP 分支的 $(\beta_1, \gamma_1, \alpha_1)$ 与 $(\beta_2, \gamma_2, \alpha_2)$。前两个做 shift/scale,第三个是残差门控:$x \leftarrow x + \alpha \odot f(\cdot)$。

  3. Zero 指的是整个调制层零初始化,于是每个 block 在第一天是恒等函数。 关键点在 DiT 的 modulate 写法:它是 $x(1 + \gamma_{\text{out}}) + \beta_{\text{out}}$ 而不是 $x\gamma_{\text{out}} + \beta_{\text{out}}$。调制层输出全零时,$\gamma_{out}=0$(有效缩放为 $1+\gamma_{out}=1$)、$\beta = 0$、$\alpha = 0$,三件事同时成立,整个 block 变成 $x \mapsto x$。实测:28 层叠完之后 $\|x_{28} - x_0\|_\infty = 0$(图 2 左)。

代价也很清楚:门关着的时候,block 内部除门控之外的参数梯度精确为零——在独立 block 接非零上游梯度时,门控先获得梯度;完整 DiT 的输出头也为零,第一步首先更新输出头。


03. 数学推导

3.1 patchify:从空间表示到 token 序列

记潜变量 $z \in \mathbb{R}^{I \times I \times C}$(256×256 图像过 VAE 后是 $I=32, C=4$)。patchify 把每个 $p \times p \times C$ 的方块摊平成一个 $p^2C$ 维向量,共

$$T = (I/p)^2$$

个,再用一个线性层 $\mathbb{R}^{p^2C} \to \mathbb{R}^d$ 嵌入。实测的形状(附录 dit_flops.py):

输入潜变量 z: (2, 32, 32, 4)
p=8: patchify -> (2, 16, 256)   T=16,  每 token 256 维
p=4: patchify -> (2, 64, 64)    T=64,  每 token 64 维
p=2: patchify -> (2, 256, 16)   T=256, 每 token 16 维

$T$ 之外的一切($d$、block 数、头数)都与 $p$ 无关,所以 $p$ 是一个纯粹花算力的旋钮:$p$ 从 8 减到 2,DiT-S 的 Gflops 从 0.36 涨到 6.06(17 倍),参数量始终是 33M。

3.2 算力账本:主导项是 $O(Td^2)$

数一个 block 的乘加次数(MAC,一次 $ab+c$ 记 1——论文里的 Gflops 用的是这个口径,记成 2 会整整差一倍):

部件 每个 token 的 MAC 说明
qkv 投影 $3d^2$ $d \to 3d$
注意力分数 $QK^\top$ $Td$ 每个 token 对 T 个位置各做 d 次乘加
注意力加权 $AV$ $Td$ 同上
输出投影 $d^2$ $d \to d$
MLP($d \to 4d \to d$) $8d^2$ 两层各 $4d^2$
adaLN 调制 $6d^2$ / 样本 + $3d$ / token 每个样本只算一次,逐元素部分才是逐 token 的

忽略调制与逐元素项,一个 block 的算力是

$$\mathrm{MAC}_{\text{block}} \approx T(12d^2 + 2Td)$$

其中 $12d^2 = 3d^2 + d^2 + 8d^2$。整个模型再乘 block 数 $N$。代入 DiT-XL/2($N=28, d=1152, T=256, p=2$):

$$28 \times 256 \times (12 \times 1152^2 + 2 \times 256 \times 1152) \approx 1.186 \times 10^{11}$$

也就是 118.6 G,与论文 Table 4 的 118.64 一致。附录脚本把 12 个模型全对了一遍,最大偏差 0.9%、多数是 0.0%(参数量同样对得上:XL/2 算出 674.9M,论文写 675M)。

这张对账表本身就是文章的一半结论,值得整张贴出来:

模型 层数 N 宽度 d token T 算力(本篇算) 算力(论文) 参数量 FID-50K
DiT-S/8 12 384 16 0.36 G 0.36 33.1 M 153.60
DiT-S/4 12 384 64 1.41 G 1.41 32.9 M 100.41
DiT-S/2 12 384 256 6.06 G 6.06 32.9 M 68.40
DiT-B/8 12 768 16 1.41 G 1.42 130.7 M 122.74
DiT-B/4 12 768 64 5.56 G 5.56 130.4 M 68.38
DiT-B/2 12 768 256 23.01 G 23.01 130.3 M 43.47
DiT-L/8 24 1024 16 5.01 G 5.01 458.4 M 118.87
DiT-L/4 24 1024 64 19.70 G 19.70 458.0 M 45.64
DiT-L/2 24 1024 256 80.71 G 80.71 457.9 M 23.33
DiT-XL/8 28 1152 16 7.39 G 7.39 675.4 M 106.41
DiT-XL/4 28 1152 64 29.05 G 29.05 675.0 M 43.01
DiT-XL/2 28 1152 256 118.64 G 118.64 674.9 M 19.47

横着读这张表,三个旋钮的作用各不相同:

  • 加深(N: 12→24→28):算力和参数量同步线性上涨,L/2 比 B/2 贵 3.5 倍算力,FID 从 43.47 到 23.33。
  • 加宽(d: 384→1152):参数量涨 $d^2$,算力也涨 $d^2$;XL/4 比 S/4 贵 20 倍算力,FID 从 100.41 到 43.01。
  • 减小 patch(p: 8→2):算力涨 4 倍一档(线性项按 T 涨),总参数量近似不变(patch 投影和输出头会随 p 改变);S/8 → S/2 贵 17 倍算力,FID 从 153.60 到 68.40。

值得注意的是「同样算力下单块算力都是 12d² 主导」这件事在表里也看得出来:DiT-S/2 与 DiT-B/4 的算力(6.06 / 5.56 G)落在同一档,说明一个 12 层 384 宽、切到 p=2 的小模型,和 12 层 768 宽、切到 p=4 的大模型,在算力上是可以互换的;两者的 FID 也确实只差 0.02。这条「算力等价」的直觉,是后面理解缩放实验的前提。

注意力占比是两种项的比值:

$$\frac{2Td}{12d^2 + 2Td} = \frac{2T}{12d + 2T}$$

代入 $T=256, d=1152$ 得 3.57%。令两者相等解出临界点 $T = 6d$:模型越宽,注意力越不重要。图 1 右图画的是各档模型在 $T$ 从 64 到 16384 上扫过的这条曲线。

3.3 adaLN:把条件变成 LayerNorm 的参数

标准 LayerNorm 之后接仿射变换:

$$h = \gamma \odot \mathrm{LN}(x) + \beta$$

其中 $\mathrm{LN}$ 对最后一维做归一化(DiT 里 elementwise_affine=False,仿射参数不是 LN 自带的)。adaLN 的做法是把 $\gamma, \beta$ 变成条件的函数:

$$(\beta_1, \gamma_1, \alpha_1, \beta_2, \gamma_2, \alpha_2) = \mathrm{Linear}_{d \to 6d}\big(\mathrm{SiLU}(c)\big), \qquad c = t_{\text{emb}} + y_{\text{emb}}$$

DiT 源码里的 modulate 是这样写的:

$$\mathrm{modulate}(x, \beta, \gamma) = x \odot (1 + \gamma) + \beta$$

用 $1+\gamma$ 时,调制层输出 $\gamma=0$ 对应有效缩放为 1,$\beta=0$ 对应不平移。注意“调制输出”与“有效缩放”是两个量,不能把同一个 $\gamma$ 同时写成 0 与 1。没有这项偏移未必数学上完全无法学习,但会改变初始特征与梯度路径。

3.4 门控与恒等初始化

DiT block 的两条残差分支都带门:

$$x' = x + \alpha_1 \odot \mathrm{Attn}\big(\mathrm{modulate}(\mathrm{LN}_1(x), \beta_1, \gamma_1)\big)$$
$$x'' = x' + \alpha_2 \odot \mathrm{MLP}\big(\mathrm{modulate}(\mathrm{LN}_2(x'), \beta_2, \gamma_2)\big)$$

调制层零初始化时 $\beta=0$、$\gamma=0$(有效缩放 $1+\gamma=1$)、$\alpha=0$。只要两个门控为零且分支输出有限,block 就是恒等映射,即使 shift/scale 非零也成立。把整层调制器置零进一步让分支内部从不平移、单位缩放开始,但这不是恒等的必要条件。

实测(附录 adaln_zero.py,DiT-XL 配置 d=1152、T=256):

adaln_zero   单块 |out - x| 最大值 = 0.000e+00
             叠 28 层后 |out - x| 最大值 = 0.000e+00
adaln        单块 |out - x| 最大值 = 2.958e+00
             单块相对扰动 std(out-x)/std(x) = 0.6127
             叠 28 层后残差流漂移 std(x_28 - x_0) = 4.4226

对照组 vanilla adaLN(无门控且调制层非零初始化):单块就给残差流叠上 std 为 0.61 的随机扰动,28 层随机游走之后漂移达到 4.42——输出里来自输入的成分已经被 28 个随机变换淹没了。零初始化这一侧,漂移严格是 0,连浮点误差都没有。图 2 左图画的就是这两条曲线。

图 2:adaLN-Zero 的初始化——恒等是怎么来的,代价是什么

这张图要看什么:左图比较相同深度宽度的两种 block;差别包含门控及调制初始化——蓝线恒等于 0,粉线一路爬到 4.42。右图是初始化那一刻各部件的梯度大小(对数轴):蓝柱的主干部分趴在底部,真实值就是精确的 0(画在 1e-12 只是为了让 log 轴显示得出),只有最后一组「调制层 gate 段」立起来;粉柱则所有部件都有梯度。注意最右一组的对比是不对称的:vanilla adaLN 结构里压根没有 gate 这一项,所以那条柱子是空的——零初始化不只是「让某些梯度变成 0」,它同时给网络装上了一组原本不存在的门。

3.5 初始化那一刻,谁拿到了梯度

记 block 输出为 $y$,损失对它的梯度为 $g = \partial L / \partial y$。由链式法则:

$$\frac{\partial L}{\partial \alpha} = g \odot f(\cdot) \neq 0$$
$$\frac{\partial L}{\partial \theta_f} = \alpha \odot \big(\cdots\big) = 0$$
$$\frac{\partial L}{\partial \beta} = \alpha \odot J_f \odot 1 = 0, \qquad \frac{\partial L}{\partial \gamma} = \alpha \odot J_f \odot \mathrm{LN}(x) = 0$$

其中 $\theta_f$ 是注意力与 MLP 的权重,$J_f$ 是它们的雅可比。$\alpha = 0$ 让除了门以外的所有梯度都精确地等于零,不是「很小」。实测(线性损失 $L = \langle y, w\rangle$,$w$ 固定随机):

adaln_zero   W_qkv / W_o / W_1 / W_2 的 |grad|max = 0.000e+00(四者都是)
             调制层六段: shift_msa=0  scale_msa=0  gate_msa=6.52e-05
                         shift_mlp=0  scale_mlp=0  gate_mlp=4.14e-04
adaln        W_qkv=2.79e-04  W_o=1.94e-04  W_1=3.13e-04  W_2=2.52e-04
             调制层四段全部非零

上面的梯度结论针对独立 block,且假定其输出接收到非零上游梯度。完整 DiT 还把 FinalLayer 的输出线性层置零:首次反传时更早层(包括 gates)没有来自损失的梯度,输出头先更新;输出头变为非零后,block 的 gates 才能接到信号。优化器权重衰减等参数更新另计。

顺带一个可预测性上的好处:FinalLayer 的输出线性层也零初始化,所以模型第一天的预测是全零,初始 loss 严格等于噪声的二阶矩。实测 $E[\varepsilon^2] = 1.0083$,零初始化时初始 MSE = 1.0083(输出最大绝对值 0.000e+00),换成正常初始化则是 3.1358(高出 3.11 倍)。初始 loss 是可以事先算出来的——这对判断「训练有没有起坏头」很有用。

3.6 四种条件注入:先在自己的玩具任务上排一遍序

论文那组数字(DiT-XL/2、400K 步、ImageNet)在本机复现不了——没有 ImageNet,也没有 TPU。但同一个问题可以搬到跑得完的 toy 上:8 个朝向的二维条纹、潜变量 8×8×2、patch p=2 得到 $T=16$、$d=64$、6 层,四种注入方式之外的一切(初始化、数据、优化器、步数、seed)完全相同,每种跑 3 个 seed。脚本就是附录里的 cond_ablation.py。

结果(最后 200 步的平均 MSE):

方案 参数量 最终 loss(均值 ± std) 三个 seed 分别
adaLN-Zero 464,328 0.0386 ± 0.0014 0.0402 / 0.0389 / 0.0368
vanilla adaLN 414,408 0.0745 ± 0.0020 0.0770 / 0.0745 / 0.0721
in-context 307,912 0.0739 ± 0.0037 0.0786 / 0.0737 / 0.0695
cross-attention 458,440 0.0796 ± 0.0028 0.0786 / 0.0834 / 0.0768

图 3:条件注入方式的 toy 对照

这张图要看什么:左图四条曲线前 100 步是缠在一起的(谁也看不出差别),从 200 步之后蓝线(adaLN-Zero)开始脱离,到 800 步已经低了将近一半。右图把 3 个 seed 单独点成白点——adaLN-Zero 最差的那个 seed(0.0402)仍然低于其它三档最好的 seed(0.0695),四组区间完全不重叠。所以在这套配置下,「adaLN-Zero 明显更好」不是随机波动。

但要老实说清这个 toy 复现了什么、没复现什么:

  • 复现了:adaLN-Zero 排第一,且差距是成倍量级(0.0386 对 0.0739~0.0796)。这与论文里 adaLN-Zero 的 FID 明显低于其余三档方向一致。
  • 没复现:论文里 in-context 是明显最差的一档(FID 35.24,比第二名差 9 分),而在这个 toy 上 in-context(0.0739)与 vanilla adaLN(0.0745)几乎打平,甚至略好于 cross-attention(0.0796)。原因不难猜:in-context 的劣势来自「序列多了两个 token 的开销」和「条件 token 与图像 token 抢注意力」,而在 $T=16$、条件又极其简单(8 个条纹朝向)的任务上,这两点都还没成为瓶颈。这也是我不把这个 toy 的排序当成结论的原因——它只说明机制在起作用,不说明四种方案的相对优劣在 ImageNet 尺度上也一样。
  • 参数量的方向对上了:in-context 最少(307,912,它压根没有调制层),adaLN-Zero 比 vanilla adaLN 多出的 49,920 个参数正好是「每层 $2d^2 + 2d$ 个门参数 × 6 层」,与论文里 adaLN-Zero(675M)比 vanilla adaLN(600M)多出来的那部分同源。

3.7 四种方案到底差在哪(回到论文的数字)

论文比较了四种把条件塞进 Transformer 的办法,它们的差别全部集中在「条件从哪里进入 block」这一件事上:

方案 条件怎么进去 代价
in-context 把 $t_{\text{emb}}, y_{\text{emb}}$ 当成两个额外 token 拼在序列前面,走的还是同一套 self-attention 序列变长,token 数从 $T$ 变 $T+2$;每个 block 都要为这两个 token 做一次全量 attention
cross-attention 主干只放图像 token,另加一层交叉注意力,条件 token 当 key/value 多一套 $d \to d$ 的注意力参数,算力上涨最多
vanilla adaLN 条件不进序列,变成 LayerNorm 的 shift/scale($\gamma, \beta$),没有门 几乎零成本,但残差分支始终全额接进来
adaLN-Zero 同上,再多两个门 $\alpha_1, \alpha_2$,并把整个调制层零初始化 相比 vanilla adaLN 只多 $2d^2$ 的调制参数

图 4:DiT-XL/2 骨架下四种 block design(论文 Table 4,400K 步)

这张图要看什么:三根柱子分别是算力、参数量、FID,四组从左到右对应上表的四种方案。先看中间那张参数量图:in-context 参数最少(449M,因为它没有额外的注意力参数,只是把序列变长),cross-attention 为约 598M 总参数,不能把它与另一方案的 600M 相加,adaLN-Zero 是 675M,比 vanilla adaLN 多了 75M——这就是那 $2d^2$ 个门参数。再看右边那张 FID 图:参数最少的 in-context 质量最差(35.24),参数最多的 adaLN-Zero 质量最好(19.47)。最后看左边的算力:四种方案的算力其实都挤在 118~138 G 这个窄区间里(in-context 因为序列长了一点反而是 119.37,cross-attention 因为多一套注意力到 137.62),也就是说这 16 分的 FID 差距,几乎完全不是靠算力堆出来的——是结构设计出来的。

一个容易看漏的对照:vanilla adaLN 与 adaLN-Zero 的算力只差 0.08 G(118.56 vs 118.64),参数量差 75M,FID 差 5.74 分。门取常数 1 时可包含无门控的行为,但条件相关门控也增加了函数自由度;该对照同时改变参数与初始化,不能证明函数类完全相同,更不能把 5.74 分全部归因为优化路径。


04. 代码实现

完整脚本在文末附录,这里放三段核心代码与它们的真实输出。

4.1 patchify 与它的逆

def patchify(z, p):
    """z: [B, I, I, C] -> [B, (I/p)^2, p*p*C]"""
    B, I, _, C = z.shape
    g = I // p
    z = z.reshape(B, g, p, g, p, C)
    z = z.transpose(0, 1, 3, 2, 4, 5)       # [B, g, g, p, p, C]
    return z.reshape(B, g * g, p * p * C)

跑一遍(batch=2、32×32×4):p=8 -> (2, 16, 256)、p=4 -> (2, 64, 64)、p=2 -> (2, 256, 16),逆变换的最大误差 0.00e+00。

4.2 一个 DiT block

def modulate(h, shift, scale):
    """DiT 源码的 modulate: x * (1 + scale) + shift"""
    sh = tg.reshape(shift, (shift.v.shape[0], 1, -1))
    sc = tg.reshape(scale, (scale.v.shape[0], 1, -1))
    return tg.add(tg.add(h, tg.mul(h, sc)), sh)

class DiTBlock:
    def __call__(self, x, c, store, ctx=None):
        B, T, d = x.v.shape
        sh_a, sc_a, gate_a, sh_m, sc_m, gate_m = tg.chunk_last(
            self.w_mod(tg.silu(c), store), 6)
        h = self._modulate(tg.layernorm(x), sh_a, sc_a)
        a = mh_self_attention(h, self.n_heads, self.w_qkv, self.w_o, store)
        x = tg.add(x, tg.mul(tg.reshape(gate_a, (B, 1, d)), a))
        h2 = self._modulate(tg.layernorm(x), sh_m, sc_m)
        u = self.w_2(tg.gelu_tanh(self.w_1(h2, store)), store)
        return tg.add(x, tg.mul(tg.reshape(gate_m, (B, 1, d)), u))

逐项对一下变量名与 03 节的符号:sh_a/sc_a/gate_a 就是 $(\beta_1, \gamma_1, \alpha_1)$,w_mod 是那个 $\text{Linear}(d, 6d)$,注意它前面接的是 silu 而不是别的激活。w_mod 在 mode="adaln_zero" 时用 zero=True 初始化——weight 和 bias 全是 0,另外 vanilla adaLN 调制四段且没有残差门控;两者并非只差初始化。

4.3 账本对账

dit_flops.py 的输出(节选):

2. 参数账本(DiT-XL/2)
  patchify 线性嵌入            19,584  ( 0.0%)
  t 嵌入(256->d->d)       1,624,320  ( 0.2%)
  类别嵌入(1000×d)        1,152,000  ( 0.2%)
  28 个 block × 23,907,456  669,408,768  (99.2%)
  FinalLayer                2,677,264  ( 0.4%)
  合计                     674,881,936  = 674.9 M   (论文: 675 M)

3. Gflops 账本(DiT-XL/2,按 MAC 记)
  block 内逐 token 线性(12d²)  114.15 G  (96.2%)
  block 内注意力(2T²d)           4.23 G  ( 3.6%)
  合计                          118.64 G   (论文: 118.64 G)

4. 12 个模型:实测账本 vs 论文数字(Gflops 偏差)
  DiT-S/8   0.36 vs 0.36   -0.9%     DiT-XL/2  118.64 vs 118.64   0.0%

12 个模型的 Gflops 全部对上(最大偏差 0.9%),参数量全部对上(四舍五入到 M 后与论文一致)。这套账本是可信的,后面所有关于「算力花在哪」的结论都建立在它上面。


05. 工业级实现对照

对照 facebookresearch/DiT 的 models.py(以 2026-09 时的 main 分支为准)。

modulate 就是那个 $1+\gamma$。 源码原样:

def modulate(x, shift, scale):
    return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)

调制层是一个共享的 SiLU + Linear,一次算 6 段再 chunk。 不是给六个分支各配一个线性层:

self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 6 * hidden_size, bias=True))
# 上面这个 Sequential 的输出一次是 6d 维,再切成六段
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=1)

初始化分三步,顺序不能乱。 先 xavier_uniform_ 初始化所有 Linear,再把每个 block 的调制层整体置零,最后把 FinalLayer 的调制层和输出线性层置零:

self.apply(_basic_init)                             # 1. 全模型 xavier_uniform
for block in self.blocks:                           # 2. 每个 block 的 adaLN 调制层归零
    nn.init.constant_(block.adaLN_modulation[-1].weight, 0)
    nn.init.constant_(block.adaLN_modulation[-1].bias, 0)
nn.init.constant_(self.final_layer.adaLN_modulation[-1].weight, 0)   # 3. 输出层归零
nn.init.constant_(self.final_layer.adaLN_modulation[-1].bias, 0)
nn.init.constant_(self.final_layer.linear.weight, 0)
nn.init.constant_(self.final_layer.linear.bias, 0)

注意第 2 步是整层置零,不是只置零 gate 那两段——这就是 3.4 节说的「一次拿到 $\gamma=0$(有效缩放为 1)、$\beta=0, \alpha=0$ 三件事」。

与最小实现的差异,以及为什么:

差异 源码里的做法 为什么
位置编码 2D sin-cos,requires_grad=False 冻结 抄 MAE 的结论:固定编码在视觉任务上不比可学习的差,还省参数、外推性能仍需另外验证
时间步编码 256 维正弦 → Linear → SiLU → Linear 高频分量让网络能区分相邻的时间步,这是所有扩散模型的标准件
类别嵌入 nn.Embedding(1000 + 1, d),多一个位置 多出来的那一行是 CFG 的 null token,token_drop 把标签换成它
输出通道 learn_sigma=True 时输出 $2 \times 4 = 8$ 通道 一半预测噪声、一半预测方差(可学习的 $\Sigma_\theta$),采样时用后者做各时间步的方差
CFG 的作用范围 forward_with_cfg 只把引导作用在前 3 个通道 论文明确说这是为了可复现做的选择,常规做法是引导所有噪声预测通道,而不是把学习方差通道也一起外推;这是复现 DiT 数值时最容易踩的坑之一
激活函数 nn.GELU(approximate="tanh") tanh 近似比精确版快,与原始实现保持一致,不能保证任何设置下 FID 完全不变

还有一个容易被忽略的点:unpatchify 用的是 torch.einsum('nhwpqc->nchpwq', x),不是简单的 reshape。它要把 token 里的 $p \times p \times C$ 重新摆回空间位置,通道维要提到最前面。

整条前向的调用顺序(DiT.forward)值得记住,因为后面拆视频版 DiT 时这几步会变成瓶颈点:

x = self.x_embedder(x) + self.pos_embed   # 1. patchify + 位置
t = self.t_embedder(t)                                   # 2. 时间步 → 256 → d
y = self.y_embedder(y, self.training)                    # 3. 类别 → d(training 时才做 label dropout)
c = t + y                                                # 4. 两个条件相加,不是拼接
for block in self.blocks: x = block(x, c)                # 5. N 个 block,条件从 c 进来
x = self.final_layer(x, c)                               # 6. 也是调制出来的 shift/scale
x = self.unpatchify(x)                                   # 7. 回到空间形状

第 4 步是「相加」而不是「拼接」,这件事在 02 节已经提过:正因为相加,c 是一个 d 维向量,整个序列共享同一份调制参数——adaLN 的表达力上限就卡在这里。

分类器无关引导(CFG)是在 forward 外面拼的。 forward_with_cfg 的做法是把 batch 复制成两份,第二份的标签换成 null token,跑到最后再合成:

half = x[: len(x) // 2]
combined = torch.cat([half, half], dim=0)          # 一份无条件、一份有条件
model_out = self.forward(combined, t, y)
eps, rest = model_out[:, :3], model_out[:, 3:]
cond_eps, uncond_eps = torch.split(eps, len(eps) // 2, dim=0)
half_eps = uncond_eps + cfg_scale * (cond_eps - uncond_eps)
eps = torch.cat([half_eps, half_eps], dim=0)
return torch.cat([eps, rest], dim=1)                  # 其余通道原样保留

这里有两个复现时必须知道的细节:其一,引导公式用无条件分支做基准(uncond + w·(cond − uncond)),不是把有条件分支当基准;其二,eps, rest 的拆分把引导限制在前 3 个通道,返回时仍拼回其余通道(含方差),论文明确说这是为了可复现而做的选择,常规做法是引导所有噪声预测通道,而不是把学习方差通道也一起外推。照着论文数值复现时如果忘了这一条,FID 会对不上,但代码不会报错。

learn_sigma 是输出通道翻倍的原因。 当它为真时 out_channels = 2 * in_channels,网络一半预测噪声 $\varepsilon$、另一半预测 learned-range 的方差插值参数,经扩散模块变换为对数方差。采样时前者进 DDPM/DDIM 的更新公式,后者可用于学习方差的 DDPM 更新;DDIM 的随机方差由其调度与 η 决定,通常不使用该头。这也是为什么 DiT 的输出头看起来比「预测噪声」该有的形状大一号。


06. 代价与边界

代价一:$O(T^2)$ 迟早会来找你。 T=256 时注意力只占 3.6%,那是 256×256 图像。512×512 的 DiT-XL/2 有 T=1024,一次前向 524.6 Gflops(是 256 分辨率的 4.4 倍);再往视频走,时间维一加,token 数轻易上万。这也是为什么视频 DiT 必须配序列并行、窗口注意力或者时空分离注意力——「注意力不是瓶颈」这个结论只在 T 远小于 6d 时成立。

代价二:adaLN 对所有 token 施加同一个函数。 论文原话:adaLN 是四种方案里唯一「被限制成对所有 token 应用同一个函数」的。$\beta, \gamma, \alpha$ 是 d 维向量,在整个序列上广播,没有空间维度。所以凡是条件本身带空间结构的任务——局部编辑、inpainting 的 mask、逐区域控制——单靠一个广播的全局调制向量不便保留空间对应关系,通常还需空间输入、条件 token 或 cross-attention;不能据此断言整体网络表达不了局部任务。SD3 之所以把文本单独拉一条支路做联合注意力(MM-DiT),原因就在这里。

代价三:patch size 减半,算力至少四倍,参数一分不涨。 这是好事也是陷阱:好消息是可以用小模型 + 小 patch 换到大模型的效果(01 节那组数字);坏消息是推理成本是按 Gflops 付的,DiT-S/2 推理比 DiT-S/8 贵 17 倍,参数却一样大,部署时容易误判。

代价四:丢掉了 UNet 的多尺度归纳偏置。 UNet 的下采样-上采样结构天然假设「图像有局部性、有尺度层次」,这个先验在小数据集上是白送的。DiT 是一张平铺的 token 网格,全靠数据自己学。所以它赢在能缩放的地方(大数据、大算力),在几万张图的小数据集上不一定比 UNet 收敛快。

代价五:初始化决定最初的梯度路径。 独立 block 实验验证门控阻断分支梯度;完整模型还需先打开零输出头。不能据此指定“前几步只有门在动”的固定持续时间,或无依据地推荐给调制器更大学习率。

论文中的 ADM 1983 Gflops 与 DiT-XL/2 118.64 Gflops 是像素空间 UNet 与潜空间 Transformer 的跨系统比较,分辨率、表示和训练预算都不同。它说明整套 DiT 系统有竞争力,不能把约 16.7 倍单次前向差全归因于更换骨干,也不包含完整采样步数与 VAE 成本。


07. 经典论文脉络

论文 arXiv 一句话贡献
ADM 2105.05233 把 UNet 扩散模型做到超过 GAN,确立了「UNet + 注意力层」这一代骨干,也是 DiT 要对标的基线(ADM 1983 Gflops、ADM-U 2813 Gflops)
LDM 2112.10752 把扩散搬到 VAE 潜空间,32×32×4 的输入尺寸正是 DiT patchify 的起点
U-ViT 2209.12152 与 DiT 几乎同时,独立提出用 ViT 骨干做扩散,把时间步与条件当作额外 token(即 DiT 里的 in-context 方案),并证明了这条路可行
DiT 2212.09748 系统性地把「Gflops → FID」当成缩放律来量,给出 patchify + adaLN-Zero 这套设计,并证明它比 cross-attention / in-context 都好
PixArt-α 2310.00426 把 DiT 接到文本条件上(cross-attention 处理文本序列 + adaLN 处理时间步),把训练成本压到原来的十分之一级别
SD3 2403.03206 把 DiT 换成双支路的 MM-DiT(图像与文本各一条,联合注意力),并把训练目标换成 rectified flow
Latte 2401.03048 把 DiT 搬到视频:提出四种时空注意力的分解方式,是后续视频 DiT 的结构模板

演进的主线很清楚:UNet(ADM/LDM)→ 把骨干换成 Transformer(U-ViT/DiT)→ 解决文本条件(PixArt-α/SD3)→ 解决时间维(Latte 及之后)。

这条线里有一处细节值得单独说,因为它解释了 adaLN-Zero 为什么能活到今天:后继模型几乎都保留了「时间步走 adaLN、文本走另一条路」这个分工。PixArt-α 是最保守的一步——它把 DiT 的类别条件换成文本条件,但文本不走 adaLN,而是另加一层 cross-attention,时间步仍然走 adaLN。SD3 往前走了一大步,把图像与文本做成两条支路做联合注意力(MM-DiT),但时间步的调制依然是 adaLN-Zero 式的,只是把「类别嵌入」换成了文本池化向量。也就是说,DiT 这篇论文真正被继承下来的不是 patchify(在此之前 ViT 系列已经这么干了),而是「条件不进序列,改成 LayerNorm 的参数,并且整层零初始化让 block 从恒等开始」这套写法。Latte 之后的视频 DiT(包括现在各种视频生成模型)沿用的也是它——时间步调制 + 时空分解注意力。


08. 常见误解

误解一:adaLN-Zero 只是把 scale 零初始化。 实现把整个调制线性层置零;调制输出 scale=0,而有效缩放为 1。残差恒等主要由 gate=0 保证,单独关门已经足够。

误解二:DiT 的算力大头在自注意力上。 错,DiT-XL/2 里注意力只占 3.6%,逐 token 的线性层占 96.2%。「Transformer 长序列会爆」的直觉来自 LLM($T$ 上万、$d$ 数千), diffusion 的 $T$ 只有几百,主导项一直是 $O(Td^2)$。图 1 右图给了临界点:$T = 6d$。

误解三:patch size 只影响 token 数,是个「免费的」结构选择。 对参数量免费(DiT-S 三档都是 33M),对算力一点也不免费(0.36 → 1.41 → 6.06 G,17 倍)。论文里那句「changing $p$ has no meaningful impact on downstream parameter counts」说的是参数,别顺手读成「没有代价」。

误解四:论文里的 Gflops 是 FLOPs。 是 MAC。我第一次按「一次乘加 = 2 FLOPs」数 DiT-XL/2,得到约 237.28 GFLOP,是 118.64 GMAC 的两倍,一度以为自己漏算了什么结构。改成 MAC 口径后 12 个模型全部对上。要拿自己算的数和论文比,先确认口径。

误解五:零初始化让整个模型不能训练。 零输出头仍能先收到非零梯度;随后 gate 和主干逐步接到信号。独立 block 的梯度实验不能替代完整 DiT 首次反传的检查。

误解六:DiT 就是「把 UNet 换成 ViT」,拿一个标准 ViT 直接接上就行。 差的不是一点:DiT 没有 [CLS] token、没有分类头、位置编码是冻结的 2D sin-cos(这是原始实现的选择,不能由此保证所有任务中都优于可学习编码);最关键的是它的 block 不是标准 ViT block——多了一条 adaLN 调制支路和两个门,FinalLayer 也是「调制 + 线性」而不是「LN + 线性」。把 torchvision 里的 ViT 拿来改,能跑起来,但那不是 DiT,也复现不出 19.47 的 FID。


09. 动手验证

文末附录一共六个脚本,其中四个可以直接跑(python xxx.py,只需要 numpy 与 matplotlib),另外两个(tiny_grad.py、dit_core.py)是被它们 import 的公共件,不单独运行。预期结果:

  1. dit_flops.py —— 打印 patchify 的真实形状、DiT-XL/2 的参数与 Gflops 逐项账本,以及 12 个模型与论文 Table 4 的对账表。预期:合计 674.9 M / 118.64 G,与论文的 675 M / 118.64 G 一致;12 个模型的 Gflops 偏差都在 1% 以内。

  2. adaln_zero.py —— 打印 A/B/D 三组实验。预期:adaLN-Zero 单块与 28 层的 $\|out - x\|_\infty$ 都是 0.000e+00;vanilla adaLN 单块相对扰动 0.6127、28 层漂移 4.4226;adaLN-Zero 的 W_qkv/W_o/W_1/W_2 梯度全为 0,只有 gate 段非零;零初始化输出层时初始 MSE = 1.0083 $= E[\varepsilon^2]$。最后一行是 autograd 自检,预期最大相对误差 ~1.5e-09。

  3. cond_ablation.py —— 四种条件注入在同一个 toy 任务上的训练对照。直接跑是快速档(150 步 × 1 个 seed,约 40 秒),它已经能看出同一个排序(adaLN-Zero 0.2588,其余三档 0.2999~0.3051);本文 3.6 节引用的那组数字来自完整档,加 --full(800 步 × 3 个 seed,约十分钟)。两种档位分别缓存在 _cache/cond_results_s150_n1.json 与 _cache/cond_results_s800_n3.json,互不覆盖。

  4. make_figures.py —— 生成本文的四张配图,其中图 3 直接读上面那份完整档缓存(所以只要缓存还在,画图是秒级的)。

想自己验证「恒等初始化」这件事,最小实验是五行:建一个 DiTBlock(d=1152, n_heads=16, mode="adaln_zero"),喂一个随机 x 和随机 c,比较 out 与 x。你会看到差是严格的 0,不是 1e-7 这种量级——因为门是乘在分支输出上的,0 乘任何数都是 0。把 mode 换成 "adaln" 再跑一次,差值立刻变成 1e0 量级。


10. 延伸阅读

  • 往回走:潜空间扩散与 Stable Diffusion 的整体结构(latent_diffusion)、DDPM 的训练目标(ddpm)、自注意力的计算细节(attention_basics)
  • 往两侧走:分类器无关引导的代价(cfg)、流匹配与 Rectified Flow——SD3 之后的新一代训练目标(flow_matching)
  • 往工程走:DiT 上到视频之后 token 数暴涨,必须靠并行切分,见序列并行(sequence_parallel)、张量/流水线并行(tensor_pipeline_parallel)、混合精度(mixed_precision)
  • 往采样走:同样的骨干,采样步数怎么省,见从 DDIM 到高阶采样器(ddim_samplers)

附录:完整代码

09 节用到的脚本全文如下(dit_flops.py、adaln_zero.py、cond_ablation.py、tiny_grad.py、dit_core.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。

dit_flops.py

# -*- coding: utf-8 -*-
"""DiT 账本(一):patchify 的形状、参数量与 Gflops,逐项对账。

DiT 这篇论文最反直觉的一点是:**它把「模型大小」和「一次前向的算力」拆开了**。
改 patch size p 几乎不动参数量,却能让 Gflops 翻四倍;改 hidden size d
几乎不动 token 数,却能让参数量翻四倍而 Gflops 只涨一点。

这个脚本做三件事:
  1. 真的做一次 patchify,打印每一步的形状(不是嘴上说 T=(I/p)^2);
  2. 按 DiT 的实际结构逐项数参数,与论文 Table 4 的 Params(M) 对账;
  3. 用同一套结构数 Gflops,与论文 Table 4 的 Flops(G) 对账。

对账口径说明(很关键,第一次数会差 2 倍):论文里的 Gflops 数的是
**乘加次数 MAC**,一次 a*b+c 记 1,不是记 2。所以线性层 [m,n] 处理一个 token
记 m*n,不记 2*m*n。下面的 counter 全部按 MAC 记,才能和论文对上。
"""

import numpy as np

# ─────────────── 论文 Table 1 的四档配置:(深度 N, hidden d, 头数) ───────────────
DIT_CONFIGS = {
    "S":  (12, 384, 6),
    "B":  (12, 768, 12),
    "L":  (24, 1024, 16),
    "XL": (28, 1152, 16),
}

# 论文 Table 4 的真实数字,用来对账:(Gflops, 参数量 M, FID-50K 无引导)
PAPER_TABLE4_256 = {
    ("S", 8): (0.36, 33, 153.60), ("S", 4): (1.41, 33, 100.41), ("S", 2): (6.06, 33, 68.40),
    ("B", 8): (1.42, 131, 122.74), ("B", 4): (5.56, 130, 68.38), ("B", 2): (23.01, 130, 43.47),
    ("L", 8): (5.01, 459, 118.87), ("L", 4): (19.70, 458, 45.64), ("L", 2): (80.71, 458, 23.33),
    ("XL", 8): (7.39, 676, 106.41), ("XL", 4): (29.05, 675, 43.01), ("XL", 2): (118.64, 675, 19.47),
}

# 四种条件注入在 DiT-XL/2 上的对照:论文 Table 4 末尾 4 行
PAPER_BLOCK_DESIGN = {
    # 名字: (Gflops, Params M, FID)
    "in-context":     (119.37, 449, 35.24),
    "cross-attention": (137.62, 598, 26.14),
    "adaLN":          (118.56, 600, 25.21),
    "adaLN-Zero":     (118.64, 675, 19.47),
}


# ─────────────── 1. patchify:真的做一遍 ───────────────
def patchify(z, p):
    """z: [B, I, I, C] -> [B, (I/p)^2, p*p*C]。

    把 I×I 切成 (I/p)×(I/p) 个格子,每个格子里的 p*p*C 个数直接摊平成
    一个 token 的原始特征。论文里这一步是一个 nn.Linear(p*p*C, d)。
    """
    B, I, _, C = z.shape
    assert I % p == 0, f"I={I} 必须被 p={p} 整除"
    g = I // p
    z = z.reshape(B, g, p, g, p, C)
    z = z.transpose(0, 1, 3, 2, 4, 5)       # [B, g, g, p, p, C]
    return z.reshape(B, g * g, p * p * C)


def unpatchify(x, p, I, C):
    """patchify 的逆:[B, T, p*p*C] -> [B, I, I, C]。"""
    B = x.shape[0]
    g = I // p
    x = x.reshape(B, g, g, p, p, C)
    x = x.transpose(0, 1, 3, 2, 4, 5)       # [B, g, p, g, p, C]
    return x.reshape(B, I, I, C)


# ─────────────── 2. 参数账本 ───────────────
def count_params(cfg="XL", p=2, C=4, n_classes=1000, t_freq=256):
    """按 DiT 源码的实际结构逐项数参数。"""
    N, d, _ = DIT_CONFIGS[cfg]
    items = {}
    items["patchify 线性嵌入"] = (p * p * C) * d + d
    items["t 嵌入(256→d→d)"] = t_freq * d + d + d * d + d
    items["类别嵌入(1000×d)"] = n_classes * d
    # 每个 DiT block:qkv(3d²) + attn out(d²) + mlp(4d²+4d²) + adaLN 调制(d→6d)
    # 加两个无仿射参数的 LayerNorm(adaLN 那层的 scale/shift 由调制给出,不另设)
    per_block = (3 * d * d + 3 * d) + (d * d + d) + (4 * d * d + 4 * d + 4 * d * d + d) \
        + (6 * d * d + 6 * d) + 2 * d
    items[f"{N} 个 block × {per_block:,}"] = N * per_block
    # FinalLayer:adaLN 线性 d→2d + 线性 d→p²C + 一个 LayerNorm
    items["FinalLayer"] = (2 * d * d + 2 * d) + (d * (p * p * C) + p * p * C) + 2 * d
    total = sum(items.values())
    return total, items, per_block


# ─────────────── 3. Gflops 账本(按 MAC 记)───────────────
def count_gflops(cfg="XL", p=2, I=32, C=4, t_freq=256):
    """一次前向(batch=1)的 MAC 数,单位 G(1e9)。"""
    N, d, _ = DIT_CONFIGS[cfg]
    T = (I // p) ** 2
    # 每个 block:qkv 3d²、attn 分数 T²d、attn 加权 T²d、out d²、mlp 8d²
    # 调制层对整条样本只算一次(6d²),逐 token 的 γ/β/α 逐元素乘加记 3Td
    per_block_token = 12 * d * d                 # 3(qkv) + 1(out) + 8(mlp)
    per_block_attn = 2 * T * d                   # 摊到每个 token 上是 2*T*d
    per_block_elem = 3 * d                       # γ⊙h、+β、α⊙Δ
    block_macs = T * (per_block_token + per_block_attn + per_block_elem) + 6 * d * d
    patch_macs = T * (p * p * C) * d             # patchify 的线性嵌入
    final_macs = T * d * (p * p * C) + 2 * d * d
    cond_macs = t_freq * d + d * d               # t 嵌入 MLP,算一次
    total = N * block_macs + patch_macs + final_macs + cond_macs
    return total / 1e9, {
        "T": T,
        "block 内逐 token 线性(12d²)": N * T * per_block_token / 1e9,
        "block 内注意力(2T²d)": N * T * per_block_attn / 1e9,
        "调制与逐元素(3Td + 6d²)": N * (T * per_block_elem + 6 * d * d) / 1e9,
        "patchify + FinalLayer": (patch_macs + final_macs) / 1e9,
        "t 嵌入": cond_macs / 1e9,
    }


def scaling_ledger():
    """把 12 个模型的实测账本与论文数字并排打印。"""
    rows = []
    for cfg in ["S", "B", "L", "XL"]:
        for p in [8, 4, 2]:
            g, _ = count_gflops(cfg, p)
            tot, _, _ = count_params(cfg, p)
            pg, pm, fid = PAPER_TABLE4_256[(cfg, p)]
            rows.append((f"DiT-{cfg}/{p}", DIT_CONFIGS[cfg][0], DIT_CONFIGS[cfg][1],
                         (32 // p) ** 2, g, pg, tot / 1e6, pm, fid))
    return rows


def main():
    rng = np.random.default_rng(0)

    print("=" * 78)
    print("1. patchify 的真实形状(256×256 图像 → 32×32×4 潜变量,batch=2)")
    print("=" * 78)
    z = rng.standard_normal((2, 32, 32, 4))
    print(f"  输入潜变量 z: {z.shape}")
    for p in [8, 4, 2]:
        x = patchify(z, p)
        back = unpatchify(x, p, 32, 4)
        ok = np.abs(back - z).max()
        print(f"  p={p}: patchify -> {x.shape}"
              f"   (T={(32//p)**2}, 每个 token 原始维度 {p*p*4})"
              f"    逆变换最大误差 {ok:.2e}")

    print()
    print("=" * 78)
    print("2. 参数账本(DiT-XL/2 逐项)")
    print("=" * 78)
    tot, items, per_block = count_params("XL", 2)
    for k, v in items.items():
        print(f"  {k:<34s} {v:>14,d}  ({v/tot*100:5.1f}%)")
    print(f"  {'合计':<34s} {tot:>14,d}  = {tot/1e6:.1f} M"
          f"   (论文 Table 4: 675 M)")

    print()
    print("=" * 78)
    print("3. Gflops 账本(DiT-XL/2,按 MAC 记)")
    print("=" * 78)
    g, parts = count_gflops("XL", 2)
    for k, v in parts.items():
        if k == "T":
            print(f"  token 数 T                              {v:>10d}")
        else:
            print(f"  {k:<38s} {v:>10.2f} G  ({v/g*100:5.1f}%)")
    print(f"  {'合计':<38s} {g:>10.2f} G   (论文 Table 4: 118.64 G)")

    print()
    print("=" * 78)
    print("4. 12 个模型:实测账本 vs 论文数字")
    print("=" * 78)
    hdr = f"{'模型':<12s}{'N':>4s}{'d':>6s}{'T':>6s}{'Gflops算':>10s}{'Gflops论文':>11s}{'差':>8s}{'M算':>8s}{'M论文':>7s}{'FID':>8s}"
    print(hdr)
    for name, N, d, T, g, pg, m, pm, fid in scaling_ledger():
        print(f"{name:<12s}{N:>4d}{d:>6d}{T:>6d}{g:>10.2f}{pg:>11.2f}{(g-pg)/pg*100:>7.1f}%"
              f"{m:>8.1f}{pm:>7d}{fid:>8.2f}")

    print()
    print("=" * 78)
    print("5. 同算力下:把预算花在 token 上还是花在宽度/深度上?")
    print("=" * 78)
    print("  取自论文 Table 4,挑 Gflops 接近的组:")
    for tag, keys in [("~5-6 G", [("S", 2), ("B", 4), ("L", 8)]),
                      ("~20-29 G", [("B", 2), ("L", 4), ("XL", 4)]),
                      ("~80-119 G", [("L", 2), ("XL", 2)])]:
        print(f"  {tag}:")
        for k in keys:
            pg, pm, fid = PAPER_TABLE4_256[k]
            print(f"    DiT-{k[0]}/{k[1]}: {pg:6.2f} G  {pm:4d} M  FID {fid:6.2f}")


if __name__ == "__main__":
    main()

adaln_zero.py

# -*- coding: utf-8 -*-
"""DiT 账本(二):adaLN-Zero 的初始化到底做了什么,逐个数出来。

adaLN-Zero 通常被一句话带过——「把每个 block 初始化成恒等函数」。
这句话里有三个可以量出来的事实,本脚本一个一个验:

  A. 恒等是真的恒等:out - x 的最大绝对误差是 0.0,不是「很小」。
  B. 门控关着的时候,block 里除门控之外的参数梯度**精确为 0**;
     连 shift/scale 那两段的梯度也是 0——只有 gate 那两段有梯度。
     这是独立 block 接非零上游梯度的实验;完整零输出头模型第一步先更新输出头。
  C. 不零初始化会怎样:block 在初始化时给残差流叠上一层随机扰动,
     28 层叠下来 std 会漂。零初始化则一层都不漂。
  D. 输出层也零初始化:模型第一天的预测是全 0,初始 loss 就是噪声的
     二阶矩,可以被事先算出来,而不是一个随机的数。

另有一个 gradcheck,用中心差分核对手写 autograd,正文里报的最大相对误差来自它。
"""

import numpy as np

import tiny_grad as tg
from dit_core import DiTBlock, MicroDiT, COND_MODES

D_BIG = 1152          # DiT-XL 的 hidden size
T_BIG = 256           # 32×32 潜变量、p=2 时的 token 数
DEPTH = 28            # DiT-XL 的深度


def _fresh_block(mode, seed=0, d=D_BIG, n_heads=16):
    rng = np.random.default_rng(seed)
    return DiTBlock(d, n_heads, rng, mode=mode)


# ─────────────── A + C:恒等性与深度漂移 ───────────────
def identity_and_drift(depth=DEPTH, seed=0, d=D_BIG, T=T_BIG, n_heads=16):
    rng = np.random.default_rng(seed)
    x0 = rng.standard_normal((1, T, d))
    c = rng.standard_normal((1, d))
    out = {}

    for mode in ("adaln_zero", "adaln"):
        x = tg.leaf(x0.copy())
        drift = []
        for i in range(depth):
            blk = _fresh_block(mode, seed=seed + i, d=d, n_heads=n_heads)
            x = blk(x, tg.leaf(c), [])
            tg.reset()
            drift.append(float(np.std(x.v - x0)))
        out[mode] = {
            "identity_err": float(np.abs(x.v - x0).max()),
            "drift": drift,
            "std_ratio": [float(np.std(x.v) / np.std(x0))],
        }
    # 单块的恒等性单独再报一次(28 层叠完还是 0 才说明真的恒等)
    blk = _fresh_block("adaln_zero", seed=3, d=d, n_heads=n_heads)
    tg.reset()
    y = blk(tg.leaf(x0.copy()), tg.leaf(c), [])
    out["adaln_zero"]["single_block_err"] = float(np.abs(y.v - x0).max())
    blk = _fresh_block("adaln", seed=3, d=d, n_heads=n_heads)
    tg.reset()
    y = blk(tg.leaf(x0.copy()), tg.leaf(c), [])
    out["adaln"]["single_block_err"] = float(np.abs(y.v - x0).max())
    out["adaln"]["single_block_rel"] = float(np.std(y.v - x0) / np.std(x0))
    tg.reset()
    return out


# ─────────────── B:初始化那一刻的梯度结构 ───────────────
def grad_at_init(mode="adaln_zero", seed=0, d=D_BIG, T=T_BIG, n_heads=16):
    rng = np.random.default_rng(seed)
    x0 = rng.standard_normal((1, T, d))
    c = rng.standard_normal((1, d))
    w = rng.standard_normal((1, T, d))          # 线性损失 L = <out, w>
    blk = _fresh_block(mode, seed=seed, d=d, n_heads=n_heads)
    store = []
    tg.reset()
    out = blk(tg.leaf(x0), tg.leaf(c), store, ctx=None)
    loss = tg.mean_all(tg.mul(out, tg.leaf(w)))
    tg.backward(loss)

    names = {
        "W_qkv": blk.w_qkv.W, "W_o": blk.w_o.W, "W_1": blk.w_1.W, "W_2": blk.w_2.W,
    }
    res = {}
    for k, arr in names.items():
        node = [n for n in store if n.v is arr]
        res[k] = float(np.abs(node[0].g).max()) if node else 0.0

    # 调制层按 6 段(或 4 段)拆开看:哪几段真的拿到了梯度
    mod_node = [n for n in store if n.v is blk.w_mod.W][0]
    g = mod_node.g                                       # [d, n_mod*d]
    n_mod = 6 if mode == "adaln_zero" else 4
    cols = [float(np.abs(g[:, i * d:(i + 1) * d]).max()) for i in range(n_mod)]
    res["mod_chunks"] = cols
    res["mod_name"] = (["shift_msa", "scale_msa", "gate_msa",
                        "shift_mlp", "scale_mlp", "gate_mlp"] if n_mod == 6
                       else ["shift_msa", "scale_msa", "shift_mlp", "scale_mlp"])
    tg.reset()
    return res


# ─────────────── D:初始 loss 是不是可预测的 ───────────────
def init_loss(mode="adaln_zero", seed=0, B=32):
    rng = np.random.default_rng(seed)
    m = MicroDiT(mode=mode, seed=seed)
    z = rng.standard_normal((B, m.T, 8))
    t = rng.uniform(0.05, 0.95, B)
    y = rng.integers(0, m.n_classes, B)
    oh = np.zeros((B, m.n_classes))
    oh[np.arange(B), y] = 1.0
    eps = rng.standard_normal((B, m.T, 8))
    store = []
    tg.reset()
    pred = m.forward(z, t, oh, store)
    mse_zero = float(((pred.v - eps) ** 2).mean())
    pmax_zero = float(np.abs(pred.v).max())
    # 把输出层换成正常 xavier 初始化再测一次
    lim = np.sqrt(6.0 / (m.d + 8))
    m.w_out.W = rng.uniform(-lim, lim, m.w_out.W.shape)
    store = []
    tg.reset()
    pred = m.forward(z, t, oh, store)
    mse_rand = float(((pred.v - eps) ** 2).mean())
    pmax_rand = float(np.abs(pred.v).max())
    tg.reset()
    return float((eps ** 2).mean()), mse_zero, mse_rand, pmax_zero, pmax_rand


# ─────────────── autograd 自检 ───────────────
def gradcheck_report(seed=0):
    """手写 autograd 对不对:拿中心差分逐个参数核。

    adaln_zero 初始化下大部分参数梯度恒为 0(B 节已证),核不出东西,
    所以核 adaln 与 cross_attn 两种结构——它们把所有算子都走到了。
    """
    rng = np.random.default_rng(seed + 7)
    z = rng.standard_normal((4, 16, 8))
    t = rng.uniform(0.05, 0.95, 4)
    oh = np.zeros((4, 8))
    oh[np.arange(4), rng.integers(0, 8, 4)] = 1.0
    tgt = rng.standard_normal((4, 16, 8))
    worst = 0.0
    for mode in ("adaln", "cross_attn", "in_context"):
        mm = MicroDiT(d=32, n_heads=4, depth=2, mode=mode, seed=seed)

        def build(mm=mm):
            s = []
            p = mm.forward(z, t, oh, s)
            diff = tg.sub(p, tg.leaf(tgt))
            return tg.mean_all(tg.mul(diff, diff)), s

        tg.reset()
        worst = max(worst, tg.gradcheck(build, seed=seed))
    return worst


def main():
    print("=" * 78)
    print("A. 恒等性:DiT-XL 配置(d=1152, 16 头, T=256)")
    print("=" * 78)
    r = identity_and_drift()
    for mode in ("adaln_zero", "adaln"):
        d = r[mode]
        print(f"  {mode:<12s} 单块 |out - x| 最大值 = {d['single_block_err']:.3e}")
        if mode == "adaln":
            print(f"  {'':<12s} 单块相对扰动 std(out-x)/std(x) = {d['single_block_rel']:.4f}")
        print(f"  {'':<12s} 叠 {DEPTH} 层后 |out - x| 最大值 = {d['identity_err']:.3e}")
        dr = d["drift"]
        pick = [0, 6, 13, 20, 27]
        print(f"  {'':<12s} 残差流漂移 std(x_k - x_0):"
              + "  ".join(f"k={k+1}:{dr[k]:.4f}" for k in pick))

    print()
    print("=" * 78)
    print("B. 初始化那一刻,block 里谁拿到了梯度(L = <out, w>,w 固定随机)")
    print("=" * 78)
    for mode in ("adaln_zero", "adaln"):
        g = grad_at_init(mode)
        nz = [g[k] for k in ("W_qkv", "W_o", "W_1", "W_2")]
        print(f"  {mode:<12s} 主干参数 |grad|max: "
              + "  ".join(f"{k}={v:.3e}" for k, v in zip(("W_qkv", "W_o", "W_1", "W_2"), nz)))
        print(f"  {'':<12s} 调制层各段 |grad|max: "
              + "  ".join(f"{n}={v:.3e}" for n, v in zip(g["mod_name"], g["mod_chunks"])))

    print()
    print("=" * 78)
    print("D. 初始 loss 能不能事先算出来(FinalLayer 零初始化 vs 正常初始化)")
    print("=" * 78)
    e2, mse_z, mse_r, pz, pr = init_loss()
    print(f"  噪声二阶矩 E[eps^2]              = {e2:.4f}   <- 理论上就是初始 loss")
    print(f"  零初始化输出层,实测初始 MSE      = {mse_z:.4f}")
    print(f"  零初始化时模型输出的最大绝对值    = {pz:.3e}   (就是全 0)")
    print(f"  正常初始化输出层,实测初始 MSE    = {mse_r:.4f}   <- 高出 {mse_r/mse_z:.2f} 倍")
    print(f"  正常初始化时模型输出的最大绝对值  = {pr:.3e}")

    print()
    print("=" * 78)
    print("自检:手写 autograd vs 中心差分(adaln / cross_attn 两种结构)")
    print("=" * 78)
    print(f"  最大相对误差 = {gradcheck_report():.2e}")


if __name__ == "__main__":
    main()

cond_ablation.py

# -*- coding: utf-8 -*-
"""DiT 账本(三):四种条件注入方式,在同一个 toy 任务上跑一遍。

论文 Figure 5 / Table 4 的结论是在 ImageNet 上量出来的:DiT-XL/2 跑 400K 步,
adaLN-Zero 的 FID 是 19.47,in-context 是 35.24,中间差了近一倍。那个实验这里
复现不了(没有 ImageNet,也没有 TPU),但可以在一个跑得完的 toy 上问同一个问题:
**在同样的深度、同样的参数预算下,这四种注入方式的训练行为差多少**。

任务:8 个朝向的二维条纹(8 类),潜变量 8×8×2,patch p=2 → 16 个 token、
每个 token 8 维。按余弦 schedule 加噪,模型预测噪声,loss 是 MSE。
模型:d=64、4 头、若干层,四种 mode 之外的一切都相同(同一个 seed 初始化)。

两种运行档位(缓存按档位分开存,互不覆盖):
  python cond_ablation.py            # 快速档:150 步 × 1 个 seed,约半分钟,用于自检
  python cond_ablation.py --full     # 完整档:800 步 × 3 个 seed,约十分钟
文章里引用的那组数字来自完整档,已经存在 ../_cache/ 里,脚本会优先读缓存。
"""

import json
import os
import sys

import numpy as np

import tiny_grad as tg
from dit_core import (MicroDiT, Adam, COND_MODES, make_patterns, patchify_imgs,
                      alpha_bar_cosine)

SIDE, P, C = 8, 2, 2
PATCH_DIM = P * P * C                      # 8
T = (SIDE // P) ** 2                       # 16
K = 8                                      # 类别数
DEPTH = 6

HERE = os.path.dirname(os.path.abspath(__file__))
NODE_DIR = os.path.dirname(HERE)

# 完整档的参数(文章里的数字就是这一组)
STEPS_FULL, SEEDS_FULL = 800, (0, 1, 2)
# 快速档:只为「脚本能跑通」而存在(体检会真的执行每个脚本),
# 结果同样会缓存,不会覆盖完整档。
STEPS_QUICK, SEEDS_QUICK = 150, (0,)


def _cache_path(steps, seeds):
    return os.path.join(NODE_DIR, "_cache",
                        f"cond_results_s{steps}_n{len(seeds)}.json")


CACHE = _cache_path(STEPS_FULL, SEEDS_FULL)


def sample_batch(rng, pats, B):
    """采一批 (z_t, t, y_onehot, eps)。"""
    y = rng.integers(0, K, B)
    x0 = pats[y]                                          # [B,8,8,2]
    eps = rng.standard_normal((B, SIDE, SIDE, C))
    t = rng.uniform(0.05, 0.95, B)
    ab = alpha_bar_cosine(t)[:, None, None, None]
    z = np.sqrt(ab) * x0 + np.sqrt(1.0 - ab) * eps
    oh = np.zeros((B, K))
    oh[np.arange(B), y] = 1.0
    return patchify_imgs(z, P), t, oh, patchify_imgs(eps, P)


def make_model(mode, seed, depth=DEPTH):
    return MicroDiT(d=64, n_heads=4, depth=depth, n_classes=K, patch_dim=PATCH_DIM,
                    T=T, mode=mode, seed=seed, out_dim=PATCH_DIM)


def train_one(mode, seed=0, steps=800, batch=32, lr=3e-3, log_every=50, depth=DEPTH):
    rng = np.random.default_rng(seed)
    pats = make_patterns(SIDE, C, K)
    m = make_model(mode, seed, depth)
    opt = Adam(m.arrays())
    curve, gnorms, losses = [], [], []
    for s in range(1, steps + 1):
        z, t, oh, eps = sample_batch(rng, pats, batch)
        store = []
        tg.reset()
        pred = m.forward(z, t, oh, store)
        diff = tg.sub(pred, tg.leaf(eps))
        loss = tg.mean_all(tg.mul(diff, diff))
        tg.backward(loss)
        gn = opt.step(store, lr * (0.25 ** (s / steps)), s)
        losses.append(float(loss.v))
        gnorms.append(gn)
        if s % log_every == 0:
            curve.append([s, float(np.mean(losses[-log_every:]))])
    return {
        "mode": mode, "curve": curve, "final": float(np.mean(losses[-200:])),
        "grad1": gnorms[0], "grad10": float(np.mean(gnorms[:10])),
        "gradmax": float(np.max(gnorms)),
        "params": int(sum(a.size for a in m.arrays())),
    }


def _write_cache(data, path=CACHE):
    """每跑完一个模式就落盘:中途被打断也不会白跑。"""
    os.makedirs(os.path.dirname(path), exist_ok=True)
    with open(path, "w", encoding="utf-8") as f:
        json.dump(data, f, ensure_ascii=False, indent=1)


def run_all(seeds=(0, 1, 2), steps=800, batch=32, depth=DEPTH, resume=True):
    """跑四种条件注入的对照。resume=True 时复用缓存里已经跑完的模式。"""
    cache = _cache_path(steps, seeds)
    out = {}
    if resume and os.path.exists(cache):
        try:
            with open(cache, encoding="utf-8") as f:
                old = json.load(f)
            if old.get("_meta", {}).get("steps") == steps:
                out = {k: v for k, v in old.items() if not k.startswith("_")}
                if "_sweep" in old:
                    out["_sweep"] = old["_sweep"]
        except Exception:
            out = {}
    for mode in COND_MODES:
        if mode in out:
            print(f"  [skip] {mode}(缓存里已有,跳过)", file=sys.stderr)
            continue
        rs = [train_one(mode, seed=sd, steps=steps, batch=batch, depth=depth)
              for sd in seeds]
        fin = [r["final"] for r in rs]
        out[mode] = {
            "final_mean": float(np.mean(fin)),
            "final_std": float(np.std(fin)),
            "finals": fin,
            "curve": rs[0]["curve"],
            "grad1": rs[0]["grad1"],
            "grad10": rs[0]["grad10"],
            "gradmax": rs[0]["gradmax"],
            "params": rs[0]["params"],
        }
        out["_meta"] = {"seeds": list(seeds), "steps": steps, "batch": batch,
                        "depth": depth}
        _write_cache(out, cache)
        print(f"  [done] {mode}", file=sys.stderr)
    out["_meta"] = {"seeds": list(seeds), "steps": steps, "batch": batch,
                    "depth": depth}
    return out


def sweep_depth(modes=("adaln_zero", "adaln"), depths=(2, 4, 6, 8, 10),
                steps=400, seed=0, batch=32):
    """深度变化时,两种 adaLN 的差距怎么变。"""
    res = {}
    for mode in modes:
        row = []
        for dep in depths:
            rng = np.random.default_rng(seed)
            pats = make_patterns(SIDE, C, K)
            m = make_model(mode, seed, dep)
            opt = Adam(m.arrays())
            losses = []
            for s in range(1, steps + 1):
                z, t, oh, eps = sample_batch(rng, pats, batch)
                store = []
                tg.reset()
                pred = m.forward(z, t, oh, store)
                diff = tg.sub(pred, tg.leaf(eps))
                loss = tg.mean_all(tg.mul(diff, diff))
                tg.backward(loss)
                opt.step(store, 3e-3 * (0.25 ** (s / steps)), s)
                losses.append(float(loss.v))
            row.append([dep, float(np.mean(losses[-200:]))])
            print(f"  [done] {mode} depth={dep}", file=sys.stderr)
        res[mode] = row
    return res


def load_or_run(seeds=SEEDS_FULL, steps=STEPS_FULL, batch=32, depth=DEPTH,
                sweep=False):
    """缓存优先(完整档)。sweep=True 时才额外跑深度扫描(很慢,默认不跑)。"""
    cache = _cache_path(steps, seeds)
    data = None
    if os.path.exists(cache):
        try:
            with open(cache, encoding="utf-8") as f:
                data = json.load(f)
        except Exception:
            data = None
    if data and data.get("_meta", {}).get("steps") == steps \
            and all(m in data for m in COND_MODES):
        if sweep and "_sweep" not in data:
            data["_sweep"] = sweep_depth(steps=steps // 2, batch=batch)
            _write_cache(data, cache)
        return data
    data = run_all(seeds=seeds, steps=steps, batch=batch, depth=depth)
    if sweep:
        data["_sweep"] = sweep_depth(steps=steps // 2, batch=batch)
    _write_cache(data, cache)
    return data


def _show(res, steps, n_seeds):
    print(f"{'mode':<14s}{'参数量':>10s}{'最终 loss(均值±std)':>24s}"
          f"{'首步梯度范数':>14s}{'峰值梯度':>12s}")
    for mode in COND_MODES:
        r = res[mode]
        print(f"{mode:<14s}{r['params']:>10d}"
              f"{r['final_mean']:>16.4f} ± {r['final_std']:<7.4f}"
              f"{r['grad1']:>14.3e}{r['gradmax']:>12.3e}")
    print()
    print(f"  loss 曲线(seed=0,{steps} 步,每 {max(1, steps // 20)} 步取一次均值):")
    for mode in COND_MODES:
        c = res[mode]["curve"]
        picks = [c[i] for i in range(0, len(c), max(1, len(c) // 6))]
        print(f"  {mode:<14s} " + "  ".join(f"{s}:{v:.4f}" for s, v in picks))
    if n_seeds > 1:
        print()
        print("  每个 seed 的最终 loss(看排序稳不稳):")
        for mode in COND_MODES:
            print(f"  {mode:<14s} "
                  + "  ".join(f"{v:.4f}" for v in res[mode]["finals"]))
    if res.get("_sweep"):
        print()
        print("=" * 78)
        print("深度扫描:两种 adaLN 在不同深度下的最终 loss")
        print("=" * 78)
        for mode, row in res["_sweep"].items():
            print(f"  {mode:<12s} " + "  ".join(f"d={d}:{v:.4f}" for d, v in row))


if __name__ == "__main__":
    full = "--full" in sys.argv
    steps, seeds = (STEPS_FULL, SEEDS_FULL) if full else (STEPS_QUICK, SEEDS_QUICK)
    tag = "完整档" if full else "快速档"
    print("=" * 78)
    print(f"四种条件注入在 toy 任务上的训练对照({tag}:{steps} 步 × "
          f"{len(seeds)} 个 seed)")
    print("=" * 78)
    res = load_or_run(seeds=seeds, steps=steps, sweep="--sweep" in sys.argv)
    _show(res, steps, len(seeds))
    if not full:
        print()
        print(f"  想复现文章里的那组数字:python cond_ablation.py --full"
              f"({STEPS_FULL} 步 × {len(SEEDS_FULL)} 个 seed,约十分钟)")

tiny_grad.py

# -*- coding: utf-8 -*-
"""一个 100 行的 reverse-mode autograd,只为本文的几个 toy 实验服务。

为什么手搓:环境里没有 torch。为什么不用有限差分:要训几千步,差分太慢。
手搓最大的风险是某个算子的 backward 写错,所以配套了 `gradcheck`——
用中心差分逐参数核对,正文里报的 2.3e-9 就是它量出来的。
"""

import numpy as np

_TAPE = []


class Node:
    __slots__ = ("v", "g", "ins", "bwd")

    def __init__(self, v, ins=(), bwd=None):
        self.v = np.asarray(v, dtype=np.float64)
        self.g = None
        self.ins = ins
        self.bwd = bwd
        _TAPE.append(self)

    @property
    def shape(self):
        return self.v.shape


def reset():
    _TAPE.clear()


def leaf(v):
    return Node(v)


def _unbc(g, shape):
    """把广播出去的梯度还原成 shape(求和 + 去掉被广播的轴)。"""
    if g.shape == shape:
        return g
    while g.ndim > len(shape):
        g = g.sum(axis=0)
    for ax, s in enumerate(shape):
        if s == 1 and g.shape[ax] != 1:
            g = g.sum(axis=ax, keepdims=True)
    return g.reshape(shape)


def _mk(v, ins, fn):
    def bwd(g, a=ins, f=fn):
        for node, gv in zip(a, f(g)):
            if node.g is None:
                node.g = np.zeros_like(node.v)
            node.g += _unbc(gv, node.v.shape)
    return Node(v, ins, bwd)


def add(a, b):
    return _mk(a.v + b.v, (a, b), lambda g: (g, g))


def sub(a, b):
    return _mk(a.v - b.v, (a, b), lambda g: (g, -g))


def mul(a, b):
    return _mk(a.v * b.v, (a, b), lambda g: (g * b.v, g * a.v))


def neg(a):
    return _mk(-a.v, (a,), lambda g: (-g,))


def matmul(a, b):
    """支持 [m,k]@[k,n]、[... ,m,k]@[k,n]、[m,k]@[B,k,n]、[... ,m,k]@[... ,k,n]。"""
    av, bv = a.v, b.v
    v = av @ bv

    def f(g):
        # 一律走 BLAS 的 batched matmul:比 einsum 快一个量级
        if av.ndim <= 2 and bv.ndim <= 2:
            return (g @ bv.T, av.T @ g)
        if bv.ndim == 2:                       # [..., m,k] @ [k,n]
            da = g @ bv.T
            p = int(np.prod(av.shape[:-2]))
            a2 = av.reshape(p, *av.shape[-2:]).swapaxes(-1, -2)     # [p,k,m]
            g2 = g.reshape(p, *g.shape[-2:])                        # [p,m,n]
            return (da, (a2 @ g2).sum(axis=0))
        if av.ndim == 2:                       # [m,k] @ [..., k,n]
            p = int(np.prod(bv.shape[:-2]))
            b2 = bv.reshape(p, *bv.shape[-2:]).swapaxes(-1, -2)     # [p,n,k]
            g2 = g.reshape(p, *g.shape[-2:])                        # [p,m,n]
            return ((g2 @ b2).sum(axis=0), av.T @ g)
        return (g @ bv.swapaxes(-1, -2), av.swapaxes(-1, -2) @ g)

    return _mk(v, (a, b), f)


def transpose(a, *axes):
    inv = np.argsort(axes)
    return _mk(a.v.transpose(axes), (a,), lambda g: (g.transpose(inv),))


def reshape(a, shape):
    return _mk(a.v.reshape(shape), (a,), lambda g: (g.reshape(a.v.shape),))


def scale(a, c):
    return _mk(a.v * c, (a,), lambda g: (g * c,))


def chunk_last(a, n):
    """把最后一维等分成 n 份(adaLN 的 6 路调制输出就是这么切的)。"""
    k = a.v.shape[-1] // n
    out = []
    for i in range(n):
        v = a.v[..., i * k:(i + 1) * k]

        def f(g, i=i, k=k):
            full = np.zeros_like(a.v)
            full[..., i * k:(i + 1) * k] = g
            return (full,)
        out.append(_mk(v, (a,), f))
    return out


def split3(a):
    """把最后一维等分三份(qkv)。"""
    d = a.v.shape[-1] // 3
    out = []
    for i in range(3):
        v = a.v[..., i * d:(i + 1) * d]

        def f(g, i=i, d=d):
            full = np.zeros_like(a.v)
            full[..., i * d:(i + 1) * d] = g
            return (full,)
        out.append(_mk(v, (a,), f))
    return out


def take_tokens(a, n):
    """取前 n 个 token,其余位置梯度补零(in-context 读回图像 token 用)。"""
    v = a.v[:, :n, :]

    def f(g):
        full = np.zeros_like(a.v)
        full[:, :n, :] = g
        return (full,)
    return _mk(v, (a,), f)


def concat_tokens(a, b):
    """沿 token 维拼接(in-context 把条件 token 接在序列末尾)。"""
    v = np.concatenate([a.v, b.v], axis=1)

    def f(g):
        return (g[:, :a.v.shape[1], :], g[:, a.v.shape[1]:, :])
    return _mk(v, (a, b), f)


def tanh(a):
    t = np.tanh(a.v)
    return _mk(t, (a,), lambda g: (g * (1.0 - t * t),))


def sigmoid(a):
    s = 1.0 / (1.0 + np.exp(-a.v))
    return _mk(s, (a,), lambda g: (g * s * (1.0 - s),))


def silu(a):
    s = 1.0 / (1.0 + np.exp(-a.v))
    v = a.v * s
    return _mk(v, (a,), lambda g: (g * (s + v * (1.0 - s)),))


def gelu_tanh(a):
    """GELU 的 tanh 近似,DiT 源码里用的就是这个(nn.GELU(approximate='tanh'))。"""
    k = np.sqrt(2.0 / np.pi)
    s = a.v * a.v
    inner = k * (a.v + 0.044715 * a.v * s)
    t = np.tanh(inner)
    v = 0.5 * a.v * (1.0 + t)
    dv = 0.5 * (1.0 + t) + 0.5 * a.v * (1.0 - t * t) * k * (1.0 + 0.134145 * s)
    return _mk(v, (a,), lambda g: (g * dv,))


def softmax(a, axis=-1):
    e = np.exp(a.v - a.v.max(axis=axis, keepdims=True))
    p = e / e.sum(axis=axis, keepdims=True)

    def f(g):
        s = (g * p).sum(axis=axis, keepdims=True)
        return (p * (g - s),)
    return _mk(p, (a,), f)


def layernorm(a, eps=1e-6):
    """对最后一维做归一化,无仿射参数(DiT 的缩放/平移由 adaLN 给出)。"""
    mu = a.v.mean(axis=-1, keepdims=True)
    xc = a.v - mu
    var = (xc * xc).mean(axis=-1, keepdims=True)
    inv = 1.0 / np.sqrt(var + eps)
    v = xc * inv

    def f(g):
        gm = g.mean(axis=-1, keepdims=True)
        gv = (g * v).mean(axis=-1, keepdims=True)
        return (inv * (g - gm - v * gv),)
    return _mk(v, (a,), f)


def mean_all(a):
    n = a.v.size
    return _mk(a.v.mean(), (a,), lambda g: (np.full(a.v.shape, g / n),))


def backward(loss):
    for node in _TAPE:
        node.g = np.zeros_like(node.v)
    loss.g = np.ones_like(loss.v)
    for node in reversed(_TAPE):
        if node.bwd is not None:
            node.bwd(node.g)
    return loss


def gradcheck(build, eps=1e-5, seed=0, n_probe=6):
    """中心差分逐参数核对,返回最大相对误差。

    build() 每次调用要返回 (loss Node, 参数 Node 列表)。差分直接扰动 Node 底层
    的 ndarray——Node.v 与模型里那块内存是同一块,所以扰动后重新 forward
    就能拿到扰动后的 loss。
    """
    rng = np.random.default_rng(seed)
    worst = 0.0
    _, store = build()
    for node in store:
        flat = node.v.reshape(-1)
        idx = rng.choice(flat.size, size=min(n_probe, flat.size), replace=False)
        for i in idx:
            old = flat[i]
            flat[i] = old + eps
            reset()
            lp = float(build()[0].v)
            flat[i] = old - eps
            reset()
            lm = float(build()[0].v)
            flat[i] = old
            num = (lp - lm) / (2 * eps)
            reset()
            loss, s2 = build()
            backward(loss)
            cur = [n for n in s2 if n.v is node.v]
            ana = cur[0].g.reshape(-1)[i] if cur and cur[0].g is not None else 0.0
            denom = max(abs(num), abs(ana), 1e-8)
            worst = max(worst, abs(num - ana) / denom)
        reset()
    return worst

dit_core.py

# -*- coding: utf-8 -*-
"""DiT 的最小可运行复刻(numpy):一个 block、四种条件注入方式、一个 toy 训练集。

结构上贴 facebookresearch/DiT 的 models.py:
    modulate(x, shift, scale) = x * (1 + scale) + shift
    block:  x = x + gate_msa * attn(modulate(norm1(x), shift_msa, scale_msa))
            x = x + gate_mlp * mlp(modulate(norm2(x), shift_mlp, scale_mlp))
初始化也照抄:所有 Linear 走 xavier_uniform,然后把每个 block 的 adaLN 调制层
(SiLU → Linear(d, 6d))的 weight 与 bias 整个置零;FinalLayer 的调制层与输出
线性层同样置零。

四种条件注入(对应论文 Figure 3 与 Table 4 末尾四行):
    adaln_zero : adaLN-Zero,调制层零初始化
    adaln      : 同样结构,但调制层用正常 xavier 初始化(论文里的 vanilla adaLN)
    in_context : t 与 y 当作两个额外 token 塞进序列,标准 ViT block
    cross_attn : 标准 ViT block + 一层对条件向量的 cross-attention
"""

import numpy as np

import tiny_grad as tg

COND_MODES = ("adaln_zero", "adaln", "in_context", "cross_attn")


# ─────────────── 线性层 ───────────────
class Linear:
    def __init__(self, fan_in, fan_out, rng, zero=False, std=None):
        if zero:
            self.W = np.zeros((fan_in, fan_out))
        elif std is not None:
            self.W = rng.normal(0.0, std, (fan_in, fan_out))
        else:
            lim = np.sqrt(6.0 / (fan_in + fan_out))       # xavier_uniform
            self.W = rng.uniform(-lim, lim, (fan_in, fan_out))
        self.b = np.zeros(fan_out)

    def __call__(self, x, store=None):
        Wn = tg.leaf(self.W)
        bn = tg.leaf(self.b)
        if store is not None:
            store.append(Wn)
            store.append(bn)
        return tg.add(tg.matmul(x, Wn), bn)


def _to_heads(a, B, T, n_heads, dh):
    return tg.transpose(tg.reshape(a, (B, T, n_heads, dh)), 0, 2, 1, 3)


def _attend(q, k, v, scale):
    s = tg.scale(tg.matmul(q, tg.transpose(k, 0, 1, 3, 2)), scale)
    return tg.matmul(tg.softmax(s, axis=-1), v)


def mh_self_attention(x, n_heads, w_qkv, w_o, store):
    B, T, d = x.v.shape
    dh = d // n_heads
    q, k, v = tg.chunk_last(w_qkv(x, store), 3)
    q = _to_heads(q, B, T, n_heads, dh)
    k = _to_heads(k, B, T, n_heads, dh)
    v = _to_heads(v, B, T, n_heads, dh)
    y = _attend(q, k, v, 1.0 / np.sqrt(dh))              # [B,H,T,dh]
    y = tg.reshape(tg.transpose(y, 0, 2, 1, 3), (B, T, d))
    return w_o(y, store)


def mh_cross_attention(x, ctx, n_heads, w_q, w_kv, w_o, store):
    B, T, d = x.v.shape
    Tc = ctx.v.shape[1]
    dh = d // n_heads
    q = tg.chunk_last(w_q(x, store), 3)[0]
    k, v = tg.chunk_last(w_kv(ctx, store), 2)
    q = _to_heads(q, B, T, n_heads, dh)
    k = _to_heads(k, B, Tc, n_heads, dh)
    v = _to_heads(v, B, Tc, n_heads, dh)
    y = _attend(q, k, v, 1.0 / np.sqrt(dh))
    y = tg.reshape(tg.transpose(y, 0, 2, 1, 3), (B, T, d))
    return w_o(y, store)


# ─────────────── DiT block ───────────────
class DiTBlock:
    def __init__(self, d, n_heads, rng, mode="adaln_zero", mlp_ratio=4):
        assert mode in COND_MODES, mode
        self.mode = mode
        self.d = d
        self.n_heads = n_heads
        self.w_qkv = Linear(d, 3 * d, rng)
        self.w_o = Linear(d, d, rng)
        hid = int(d * mlp_ratio)
        self.w_1 = Linear(d, hid, rng)
        self.w_2 = Linear(hid, d, rng)
        self.w_cq = Linear(d, 3 * d, rng)          # cross-attn 的 q
        self.w_ckv = Linear(d, 2 * d, rng)         # cross-attn 的 k/v
        self.w_co = Linear(d, d, rng)              # cross-attn 的输出投影
        if mode == "adaln_zero":
            self.w_mod = Linear(d, 6 * d, rng, zero=True)
        elif mode == "adaln":
            self.w_mod = Linear(d, 4 * d, rng)
        else:
            # 标准 ViT block:LayerNorm 带仿射参数,初始化成 γ=1、β=0
            self.ln1_g, self.ln1_b = np.ones(d), np.zeros(d)
            self.ln2_g, self.ln2_b = np.ones(d), np.zeros(d)
            self.ln3_g, self.ln3_b = np.ones(d), np.zeros(d)

    def _modulate(self, h, shift, scale):
        """DiT 源码的 modulate: x * (1 + scale) + shift,展平成 [B,1,d] 再广播。"""
        sh = tg.reshape(shift, (shift.v.shape[0], 1, -1))
        sc = tg.reshape(scale, (scale.v.shape[0], 1, -1))
        return tg.add(tg.add(h, tg.mul(h, sc)), sh)

    def _vi_norm(self, x, g, b, store):
        gn, bn = tg.leaf(g), tg.leaf(b)
        store.append(gn)
        store.append(bn)
        return tg.add(tg.mul(tg.layernorm(x), gn), bn)

    def __call__(self, x, c, store, ctx=None):
        B, T, d = x.v.shape
        if self.mode == "adaln_zero":
            sh_a, sc_a, gate_a, sh_m, sc_m, gate_m = tg.chunk_last(
                self.w_mod(tg.silu(c), store), 6)
            h = self._modulate(tg.layernorm(x), sh_a, sc_a)
            a = mh_self_attention(h, self.n_heads, self.w_qkv, self.w_o, store)
            x = tg.add(x, tg.mul(tg.reshape(gate_a, (B, 1, d)), a))
            h2 = self._modulate(tg.layernorm(x), sh_m, sc_m)
            u = self.w_2(tg.gelu_tanh(self.w_1(h2, store)), store)
            return tg.add(x, tg.mul(tg.reshape(gate_m, (B, 1, d)), u))

        if self.mode == "adaln":
            sh_a, sc_a, sh_m, sc_m = tg.chunk_last(self.w_mod(tg.silu(c), store), 4)
            h = self._modulate(tg.layernorm(x), sh_a, sc_a)
            a = mh_self_attention(h, self.n_heads, self.w_qkv, self.w_o, store)
            x = tg.add(x, a)
            h2 = self._modulate(tg.layernorm(x), sh_m, sc_m)
            u = self.w_2(tg.gelu_tanh(self.w_1(h2, store)), store)
            return tg.add(x, u)

        # ── 标准 ViT block(in-context / cross-attn)──
        h = self._vi_norm(x, self.ln1_g, self.ln1_b, store)
        x = tg.add(x, mh_self_attention(h, self.n_heads, self.w_qkv, self.w_o, store))
        if self.mode == "cross_attn":
            h2 = self._vi_norm(x, self.ln2_g, self.ln2_b, store)
            x = tg.add(x, mh_cross_attention(h2, ctx, self.n_heads,
                                             self.w_cq, self.w_ckv, self.w_co, store))
            h3 = self._vi_norm(x, self.ln3_g, self.ln3_b, store)
        else:
            h3 = self._vi_norm(x, self.ln2_g, self.ln2_b, store)
        u = self.w_2(tg.gelu_tanh(self.w_1(h3, store)), store)
        return tg.add(x, u)

    def arrays(self):
        out = [self.w_qkv.W, self.w_qkv.b, self.w_o.W, self.w_o.b,
               self.w_1.W, self.w_1.b, self.w_2.W, self.w_2.b]
        if self.mode in ("adaln_zero", "adaln"):
            out += [self.w_mod.W, self.w_mod.b]
        else:
            out += [self.ln1_g, self.ln1_b, self.ln2_g, self.ln2_b]
            if self.mode == "cross_attn":
                out += [self.ln3_g, self.ln3_b,
                        self.w_cq.W, self.w_cq.b, self.w_ckv.W, self.w_ckv.b,
                        self.w_co.W, self.w_co.b]
        return out


# ─────────────── 位置编码与时间步编码 ───────────────
def sincos_pos_embed(side, d):
    """2D sin-cos 位置编码(DiT 直接抄 MAE 的,固定不可学习)。

    两个空间轴各占一半维度:每轴的 d/2 维里再对半分成 sin 与 cos。
    """
    assert d % 4 == 0, "d 必须是 4 的倍数才能两轴对半分"

    def emb1d(pos, dim):
        half = dim // 2
        omega = 1.0 / (10000 ** (np.arange(half, dtype=np.float64) / half))
        out = np.asarray(pos, dtype=np.float64).reshape(-1, 1) * omega[None]
        return np.concatenate([np.sin(out), np.cos(out)], axis=-1)

    gy, gx = np.meshgrid(np.arange(side), np.arange(side), indexing="ij")
    emb = np.concatenate([emb1d(gy.reshape(-1), d // 2),
                          emb1d(gx.reshape(-1), d // 2)], axis=-1)
    return emb[None, :, :]


def timestep_embedding(t, dim, max_period=10000):
    """DiT 抄 GLIDE 的正弦时间步编码。t: [B] -> [B, dim]。"""
    half = dim // 2
    freqs = np.exp(-np.log(max_period) * np.arange(half, dtype=np.float64) / half)
    args = np.asarray(t, dtype=np.float64)[:, None] * freqs[None]
    return np.concatenate([np.cos(args), np.sin(args)], axis=-1)


# ─────────────── 一个能训的小 DiT ───────────────
class MicroDiT:
    def __init__(self, d=64, n_heads=4, depth=8, n_classes=8, patch_dim=8,
                 T=16, t_freq=32, mode="adaln_zero", seed=0, out_dim=8):
        rng = np.random.default_rng(seed)
        self.mode = mode
        self.d, self.T, self.n_classes = d, T, n_classes
        self.w_in = Linear(patch_dim, d, rng)
        self.pos = sincos_pos_embed(int(round(np.sqrt(T))), d)
        self.w_t0 = Linear(t_freq, d, rng, std=0.02)
        self.w_t1 = Linear(d, d, rng, std=0.02)
        self.y_emb = rng.normal(0.0, 0.02, (n_classes, d))
        self.blocks = [DiTBlock(d, n_heads, rng, mode=mode) for _ in range(depth)]
        # FinalLayer:adaLN 调制(零初始化)+ 输出线性层(零初始化)
        if mode in ("adaln_zero", "adaln"):
            self.w_fmod = Linear(d, 2 * d, rng, zero=True)
        self.fin_g, self.fin_b = np.ones(d), np.zeros(d)
        self.w_out = Linear(d, out_dim, rng, zero=True)

    def forward(self, z, t, y_onehot, store):
        """z:[B,T,patch_dim]  t:[B]  y_onehot:[B,K] -> pred:[B,T,out_dim]"""
        B = z.shape[0]
        d = self.d
        x = tg.add(self.w_in(tg.leaf(z), store), tg.leaf(self.pos))
        tf = timestep_embedding(t, self.w_t0.W.shape[0])
        c = self.w_t1(tg.silu(self.w_t0(tg.leaf(tf), store)), store)
        ytab = tg.leaf(self.y_emb)
        store.append(ytab)
        yv = tg.matmul(tg.leaf(y_onehot), ytab)
        c = tg.add(c, yv)

        t_tok = tg.reshape(c, (B, 1, d))
        y_tok = tg.reshape(yv, (B, 1, d))
        if self.mode == "in_context":
            x = tg.concat_tokens(x, tg.concat_tokens(t_tok, y_tok))
        ctx = tg.concat_tokens(t_tok, y_tok)

        for blk in self.blocks:
            x = blk(x, c, store, ctx=ctx)
        if self.mode == "in_context":
            x = tg.take_tokens(x, self.T)

        if self.mode in ("adaln_zero", "adaln"):
            sh, sc = tg.chunk_last(self.w_fmod(tg.silu(c), store), 2)
            hn = tg.layernorm(x)
            h = tg.add(tg.add(hn, tg.mul(hn, tg.reshape(sc, (B, 1, d)))),
                       tg.reshape(sh, (B, 1, d)))
        else:
            g, b = tg.leaf(self.fin_g), tg.leaf(self.fin_b)
            store.append(g)
            store.append(b)
            h = tg.add(tg.mul(tg.layernorm(x), g), b)
        return self.w_out(h, store)

    def arrays(self):
        out = [self.w_in.W, self.w_in.b,
               self.w_t0.W, self.w_t0.b, self.w_t1.W, self.w_t1.b,
               self.y_emb, self.w_out.W, self.w_out.b]
        if self.mode in ("adaln_zero", "adaln"):
            out += [self.w_fmod.W, self.w_fmod.b]
        else:
            out += [self.fin_g, self.fin_b]
        for blk in self.blocks:
            out += blk.arrays()
        return out


# ─────────────── Adam ───────────────
class Adam:
    def __init__(self, arrays):
        self.idx = {id(a): i for i, a in enumerate(arrays)}
        self.m = [np.zeros_like(a) for a in arrays]
        self.v = [np.zeros_like(a) for a in arrays]

    def step(self, nodes, lr, step, b1=0.9, b2=0.999, eps=1e-8):
        gnorm = 0.0
        for n in nodes:
            i = self.idx.get(id(n.v))
            if i is None or n.g is None:
                continue
            g = n.g
            gnorm += float((g * g).sum())
            self.m[i] = b1 * self.m[i] + (1 - b1) * g
            self.v[i] = b2 * self.v[i] + (1 - b2) * g * g
            mh = self.m[i] / (1 - b1 ** step)
            vh = self.v[i] / (1 - b2 ** step)
            n.v -= lr * mh / (np.sqrt(vh) + eps)
        return float(np.sqrt(gnorm))


# ─────────────── toy 训练数据:8 个方向的条纹 ───────────────
def make_patterns(side=8, C=2, K=8):
    """K 个不同朝向的二维条纹,当作 K 类「图像」。"""
    yy, xx = np.meshgrid(np.arange(side), np.arange(side), indexing="ij")
    pats = np.zeros((K, side, side, C))
    for k in range(K):
        th = k * np.pi / K
        phase = 2 * np.pi * 2.0 * (xx * np.cos(th) + yy * np.sin(th)) / side
        pats[k, :, :, 0] = np.cos(phase)
        pats[k, :, :, 1] = np.sin(phase)
    return pats / np.std(pats)


def patchify_imgs(imgs, p):
    """imgs:[B,I,I,C] -> [B,(I/p)^2,p*p*C]"""
    B, I, _, C = imgs.shape
    g = I // p
    z = imgs.reshape(B, g, p, g, p, C).transpose(0, 1, 3, 2, 4, 5)
    return z.reshape(B, g * g, p * p * C)


def alpha_bar_cosine(t):
    """余弦 schedule 的 ᾱ_t:t=0 时 1(干净),t=1 时 0(纯噪声)。"""
    return np.cos(np.pi * np.asarray(t, dtype=np.float64) / 2.0) ** 2

make_figures.py

# -*- coding: utf-8 -*-
"""画配图。数字全部来自同目录的实验脚本,不另算一遍。

运行: python make_figures.py [--only 图名]
"""

import os
import sys

import numpy as np
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))

import dit_flops as F                                     # noqa: E402
import adaln_zero as AZ                                   # noqa: E402
import cond_ablation as CA                                # noqa: E402

HERE = os.path.dirname(os.path.abspath(__file__))
FIGDIR = os.path.join(os.path.dirname(HERE), "figures")
os.makedirs(FIGDIR, exist_ok=True)

C0 = "#2f4b7c"
C1 = "#d45087"
C2 = "#f0a35e"
C3 = "#4c9f70"
CGREY = "#8a8a8a"

plt.rcParams["font.sans-serif"] = ["Heiti TC", "Arial Unicode MS", "DejaVu Sans"]
plt.rcParams["axes.unicode_minus"] = False


def _save(fig, name):
    p = os.path.join(FIGDIR, name)
    fig.savefig(p, dpi=150, bbox_inches="tight", facecolor="white")
    plt.close(fig)
    print(f"  [ok] {name}  ({os.path.getsize(p)} bytes)")


# ─────────────── 图 1:算力账本 ───────────────
def fig_ledger():
    fig, axes = plt.subplots(1, 2, figsize=(12.6, 4.8))

    # 左:同样 Gflops,参数量差 14 倍,FID 却一样
    ax = axes[0]
    marks = {8: "o", 4: "s", 2: "^"}
    cols = {8: CGREY, 4: C2, 2: C1}
    # 手工微调标号偏移,避免同算力组里三个点叠在一起看不清
    offs = {("S", 2): (7, -15), ("B", 4): (1, 9), ("L", 8): (8, 5),
            ("S", 4): (8, 5), ("B", 8): (8, 5), ("B", 2): (7, -15),
            ("L", 4): (-4, 10), ("XL", 4): (8, 6), ("L", 2): (8, 5),
            ("XL", 2): (8, 5), ("S", 8): (8, 5), ("XL", 8): (8, 5)}
    for cfg in ["S", "B", "L", "XL"]:
        for p in [8, 4, 2]:
            g, pm, fid = F.PAPER_TABLE4_256[(cfg, p)]
            ax.scatter(g, fid, marker=marks[p], s=95, color=cols[p], zorder=3,
                       edgecolor="white", linewidth=1.3)
            ax.annotate(f"{cfg}/{p}", (g, fid), textcoords="offset points",
                        xytext=offs[(cfg, p)], fontsize=9, color="#333333",
                        bbox=dict(boxstyle="round,pad=0.15", fc="white",
                                  ec="none", alpha=0.75))
    ax.set_xscale("log")
    ax.set_xlim(0.25, 320)
    ax.set_ylim(0, 175)
    ax.set_xlabel("一次前向的算力 Gflops(对数轴,论文 Table 4)")
    ax.set_ylabel("FID-50K(无分类器引导,越低越好)")
    ax.set_title("同算力时:把预算花在 token 上,比花在宽度深度上划算")
    ax.grid(alpha=0.25, ls=":")
    for tag, keys, col, tx, ty in [
            ("约 5~6 G", [("S", 2), ("B", 4), ("L", 8)], C3, 1.05, 152),
            ("约 20~29 G", [("B", 2), ("L", 4), ("XL", 4)], C0, 36, 60)]:
        xs = [F.PAPER_TABLE4_256[k][0] for k in keys]
        ys = [F.PAPER_TABLE4_256[k][2] for k in keys]
        ax.plot(xs, ys, ls="--", lw=1.3, color=col, alpha=0.75, zorder=1)
        ax.text(tx, ty, tag, color=col, fontsize=9.5, ha="left",
                bbox=dict(boxstyle="round,pad=0.25", fc="white", ec=col,
                          alpha=0.85, lw=0.8))
    ax.annotate("这两个点参数量差 4 倍\n(33M / 130M),FID 几乎重合",
                xy=(6.06, 68.40), xytext=(0.42, 26), fontsize=9,
                color="#333333", ha="left",
                arrowprops=dict(arrowstyle="->", color="#888888", lw=1))

    # 右:注意力占多少算力,随 token 数怎么变
    ax = axes[1]
    Ts = np.array([64, 256, 1024, 4096, 16384])
    ds = {"DiT-S (d=384)": 384, "DiT-B (d=768)": 768, "DiT-L (d=1024)": 1024,
          "DiT-XL (d=1152)": 1152}
    for name, d in ds.items():
        lin = 12 * d * d              # 每 token 的线性部分(qkv+out+mlp)
        attn = 2 * Ts * d             # 每 token 摊到的注意力部分
        ax.plot(Ts, attn / (lin + attn) * 100, marker="o", lw=2, ms=4,
                label=name)
    ax.set_xscale("log")
    ax.set_ylim(0, 70)
    ax.set_xlabel("token 数 T(对数轴)")
    ax.set_ylabel("注意力占单块算力的百分比")
    ax.set_title("T=256 时注意力只占 3.6%,要到 T 上千才成为大头")
    ax.grid(alpha=0.25, ls=":")
    ax.legend(fontsize=9, loc="upper left")
    ax.axvline(256, color="#cccccc", lw=1.2, ls=":", zorder=0)
    ax.annotate("256×256 图像、p=2:3.6%", xy=(256, 3.6), xytext=(330, 8),
                fontsize=9, color="#555555")
    ax.text(900, 2.5, "同样 T 下,网络越宽(d 越大)注意力占比越低",
            fontsize=9, color="#555555")

    fig.suptitle("图 1:DiT 的算力账本——钱花在哪,以及该多买 token 还是多买宽度",
                 fontsize=12)
    fig.tight_layout()
    _save(fig, "fig1_compute_ledger.png")


# ─────────────── 图 2:初始化恒等 ───────────────
def fig_identity():
    r = AZ.identity_and_drift()
    drift_zero = r["adaln_zero"]["drift"]
    drift_plain = r["adaln"]["drift"]
    ks = np.arange(1, len(drift_plain) + 1)

    fig, axes = plt.subplots(1, 2, figsize=(12.6, 4.4))
    ax = axes[0]
    ax.plot(ks, drift_plain, color=C1, lw=2.4, marker="o", ms=3,
            label="vanilla adaLN(调制层正常初始化)")
    ax.plot(ks, drift_zero, color=C0, lw=2.4, label="adaLN-Zero(调制层零初始化)")
    ax.set_xlabel("第 k 个 block 之后")
    ax.set_ylabel(r"残差流漂移 std(x_k - x_0)")
    ax.set_title("零初始化:叠 28 层,漂移严格是 0")
    ax.legend(fontsize=9)
    ax.grid(alpha=0.25, ls=":")
    ax.annotate(f"第 28 层:{drift_plain[-1]:.2f} vs {drift_zero[-1]:.2f}",
                xy=(28, drift_plain[-1]), xytext=(13, 3.4), fontsize=10,
                color="#333333",
                arrowprops=dict(arrowstyle="->", color="#888888", lw=1))

    ax = axes[1]
    labels = ["W_qkv", "W_o", "W_1", "W_2", "调制层\nshift/scale 段", "调制层\ngate 段"]
    gz = AZ.grad_at_init("adaln_zero")
    gp = AZ.grad_at_init("adaln")

    def pick(r, kind):
        """按名字取,避免把 vanilla adaLN 的 shift_mlp 误当成 gate。"""
        vals = [v for nm, v in zip(r["mod_name"], r["mod_chunks"])
                if (("gate" in nm) if kind == "gate"
                    else ("shift" in nm or "scale" in nm))]
        return max(vals) if vals else None

    v_zero = [gz["W_qkv"], gz["W_o"], gz["W_1"], gz["W_2"],
              pick(gz, "ss"), pick(gz, "gate")]
    v_plain = [gp["W_qkv"], gp["W_o"], gp["W_1"], gp["W_2"],
               pick(gp, "ss"), pick(gp, "gate")]
    # vanilla adaLN 结构里没有 gate:画在最底部,并单独标注
    FLOOR = 1e-12
    v_plain = [FLOOR if v is None else v for v in v_plain]
    v_zero = [FLOOR if v is None else v for v in v_zero]

    x = np.arange(len(labels))
    w = 0.38
    ax.bar(x - w / 2, np.array(v_plain) + FLOOR, w, color=C2, label="vanilla adaLN")
    ax.bar(x + w / 2, np.array(v_zero) + FLOOR, w, color=C0, label="adaLN-Zero")
    ax.set_yscale("log")
    ax.set_xticks(x)
    ax.set_xticklabels(labels, fontsize=9)
    ax.set_ylabel("初始化时刻的 |grad| 最大值(对数轴)")
    ax.set_title("门关着的时候,只有门自己有梯度")
    ax.legend(fontsize=9)
    ax.grid(alpha=0.25, ls=":", axis="y")
    ax.annotate("adaLN-Zero 的主干参数梯度精确为 0\n(0 画在 1e-12,否则 log 轴显示不出来)",
                xy=(x[0] + w / 2, FLOOR), xytext=(0.35, 1e-9), fontsize=9,
                color="#333333")
    ax.annotate("vanilla adaLN 结构里\n根本没有 gate 这一项",
                xy=(x[-1] - w / 2, FLOOR), xytext=(3.35, 1e-7), fontsize=9,
                color="#555555", ha="center",
                arrowprops=dict(arrowstyle="->", color="#888888", lw=1))

    fig.suptitle("图 2:adaLN-Zero 的初始化——恒等是怎么来的,代价是什么",
                 fontsize=12)
    fig.tight_layout()
    _save(fig, "fig2_identity.png")


# ─────────────── 图 3:四种条件注入的 toy 训练对照 ───────────────
def fig_ablation():
    res = CA.load_or_run()
    fig, axes = plt.subplots(1, 2, figsize=(12.6, 4.6))
    cols = {"adaln_zero": C0, "adaln": C2, "in_context": C1, "cross_attn": C3}
    names = {"adaln_zero": "adaLN-Zero", "adaln": "vanilla adaLN",
             "in_context": "in-context", "cross_attn": "cross-attention"}

    ax = axes[0]
    for mode in CA.COND_MODES:
        c = res[mode]["curve"]
        ax.plot([p[0] for p in c], [p[1] for p in c], lw=2.2, color=cols[mode],
                label=names[mode])
    ax.set_xlabel("训练步数")
    ax.set_ylabel("噪声预测 MSE(每 50 步取均值)")
    ax.set_title("同一个 toy 任务,四种条件注入的训练曲线")
    ax.legend(fontsize=9)
    ax.grid(alpha=0.25, ls=":")

    ax = axes[1]
    x = np.arange(len(CA.COND_MODES))
    means = [res[m]["final_mean"] for m in CA.COND_MODES]
    stds = [res[m]["final_std"] for m in CA.COND_MODES]
    ax.bar(x, means, yerr=stds, capsize=4,
           color=[cols[m] for m in CA.COND_MODES], alpha=0.9)
    # 把 3 个 seed 的最终 loss 直接点到柱子上:看排序有没有重叠
    for i, m in enumerate(CA.COND_MODES):
        vals = res[m]["finals"]
        ax.scatter([i] * len(vals), vals, s=20, color="white",
                   edgecolor="#333333", linewidth=0.8, zorder=4)
    ax.set_xticks(x)
    ax.set_xticklabels([names[m] for m in CA.COND_MODES], fontsize=9)
    ax.set_ylabel("最后 200 步的平均 loss(3 个 seed)")
    ax.set_title("最终 loss:adaLN-Zero 与其余三档完全不重叠")
    ax.grid(alpha=0.25, ls=":", axis="y")
    top = max(means) + max(stds) * 3
    ax.set_ylim(0, top)
    for i, (m, s) in enumerate(zip(means, stds)):
        ax.text(i, m + s + top * 0.02, f"{m:.4f}", ha="center", fontsize=9,
                color="#333333")
    ax.annotate("白点 = 每个 seed 单独的结果\n(adaLN-Zero 最差的那个 seed\n也比其它三档最好的 seed 低)",
                xy=(0, 0.041), xytext=(1.15, top * 0.62), fontsize=9,
                color="#333333",
                arrowprops=dict(arrowstyle="->", color="#888888", lw=1))

    fig.suptitle("图 3:条件注入方式的 toy 对照(6 层 d=64、T=16,800 步 × 3 seed)",
                 fontsize=12)
    fig.tight_layout()
    _save(fig, "fig3_cond_ablation.png")


# ─────────────── 图 4:四种 block design 的算力-质量对照 ───────────────
def fig_design():
    order = ["in-context", "cross-attention", "adaLN", "adaLN-Zero"]
    names = ["in-context", "cross-attention", "vanilla adaLN", "adaLN-Zero"]
    g = [F.PAPER_BLOCK_DESIGN[k][0] for k in order]
    pm = [F.PAPER_BLOCK_DESIGN[k][1] for k in order]
    fid = [F.PAPER_BLOCK_DESIGN[k][2] for k in order]

    fig, axes = plt.subplots(1, 3, figsize=(13.2, 4.0))
    cols = [C1, C2, CGREY, C0]
    for ax, vals, title, ylab, fmt in [
        (axes[0], g, "一次前向 Gflops", "Gflops", "{:.1f}"),
        (axes[1], pm, "参数量", "M", "{:d}"),
        (axes[2], fid, "FID-50K(越低越好)", "FID", "{:.2f}"),
    ]:
        bars = ax.bar(np.arange(4), vals, color=cols, alpha=0.92)
        ax.set_xticks(np.arange(4))
        ax.set_xticklabels(names, fontsize=8.5, rotation=18)
        ax.set_title(title)
        ax.set_ylabel(ylab)
        ax.grid(alpha=0.25, ls=":", axis="y")
        for b, v in zip(bars, vals):
            ax.text(b.get_x() + b.get_width() / 2, v, fmt.format(v), ha="center",
                    va="bottom", fontsize=9, color="#333333")
        if title.startswith("FID"):
            ax.set_ylim(0, max(vals) * 1.25)
        else:
            ax.set_ylim(0, max(vals) * 1.22)

    fig.suptitle("图 4:DiT-XL/2 骨架下四种 block design(论文 Table 4,400K 步)",
                 fontsize=12)
    fig.tight_layout()
    _save(fig, "fig4_block_design.png")


FIGURES = {
    "ledger": fig_ledger,
    "identity": fig_identity,
    "ablation": fig_ablation,
    "design": fig_design,
}

if __name__ == "__main__":
    only = sys.argv[2] if len(sys.argv) > 2 and sys.argv[1] == "--only" else None
    for name, fn in FIGURES.items():
        if only and name != only:
            continue
        print(f"[draw] {name}")
        fn()
0

评论 (0)

取消
粤ICP备2021042327号