LAB 03

Lab 3 解析:条件图像生成(DiT + VAE + 隐空间扩散)

把 Lab 2 的两行训练循环搬到 1024 维的 MNIST 上:classifier-free guidance 负责「听懂指令」,diffusion transformer 负责「装得下图像」,VAE 负责「把维度降下来」。15 处 TODO,一个完整的 latent diffusion 系统。

原始文件:labs/lab_three.ipynb 对应讲义:§5–§6 + 附录 D 本页解答已在 RTX 5080 上完整训练跑通

0. 本 lab 导读

Lab 2 把 flow matching 的公式链跑通了,但跑的是二维玩具数据、用的是一个三层 MLP、生成的是「随便什么样本」。Lab 3 要在这三件事上同时升级:

维度Lab 2Lab 3需要的新工具
数据$\R^2$ 上的高斯混合MNIST,$1\times32\times32=1024$ 维——
控制无条件生成「生成一个 8」——条件生成classifier-free guidance(Part 2)
架构MLPdiffusion transformerFourier 编码 / patch / 注意力 / adaLN-Zero(Part 3)
空间直接在数据空间先压到 $128\times4\times4$ 的隐空间VAE(Part 4)+ latent diffusion(Part 5)

但要强调一件让人安心的事:训练循环的那四行,和你在 Lab 2 Q3.1 写下的一模一样——采 $z$、采 $t$、采 $x_t$、回归 $u_t(x\mid z)$。Lab 3 的 15 个 TODO 里只有两个(Q2.2 和 Q5.1)在动这条循环,而且动的只是「$z$ 从哪来」和「$y$ 要不要丢掉」。剩下的 13 个全部是架构工程:怎么把一个函数 $u^\theta:\R^{1024}\times[0,1]\times\set{0,\dots,10}\to\R^{1024}$ 用张量算子搭出来。

速览:本 lab 有 15 处需要你写代码

Part 2 · Classifier-Free Guidance(2 处)——对应 Lecture 3-B 的引导理论。

  • Q2.2 CFGTrainer.get_train_loss:带标签丢弃(label dropout)的条件 flow matching 损失。四行采样 + 一行 MSE。
  • Q2.3 MLPConditionalVectorField.forward:把 $x$、类别嵌入、时间三者拼接后过 MLP,用来在二维高斯混合上做正确性自检。

Part 3 · Diffusion Transformer(5 处)——对应 Lecture 4 §6.1 的架构讨论。

  • Q3.1 FourierEncoder:把标量 $t$ 升维成随机 Fourier 特征。
  • Q3.2 Patchifier:图像 → token 序列,一行卷积 + 一行 Rearrange。
  • Q3.3 MHA + DiffusionTransformerLayer + DiffusionTransformer:本 lab 最重的一题,从零手写多头自注意力,并实现 adaLN-Zero 条件调制。
  • Q3.4 Depatchifier:token 序列 → 图像,Q3.2 的逆操作。
  • Q3.5 DiffusionTransformerFlowModel:把上面四块串成一个 $u_t^\theta(x\mid y)$。

Part 4 · 变分自编码器(7 处)——对应 Lecture 4 §6–§7 的 ELBO 推导。

  • Q4.1 ResidualBlock、Q4.2 AttnBlock:两块积木。
  • Q4.3 EncoderBlock、Q4.4 Encoder:$1\times32\times32\to128\times4\times4$。
  • Q4.5 DecoderBlock、Q4.6 Decoder:镜像回去。
  • Q4.7 VAE.compute_loss:重构项 + KL 项,对应讲义那条完整的 $\beta$-VAE 损失。

Part 5 · 隐空间扩散(1 处)

  • Q5.1 LatentCFGTrainer.get_train_loss:与 Q2.2 只差一处——训练数据换成冻结 VAE 编码后的隐变量。

预计耗时:读题 + 推导约 2 小时;写代码 3–5 小时(Q3.3 一题就可能占掉一半);四段训练在单张现代 GPU 上合计约 1.5–2.5 小时(3000 + 20000 + 5000 + 10000 步)。这是三个 lab 里唯一一个「训练时间不可忽略」的,建议先把所有形状检查跑通再开训。

学习建议:这份解析的正确用法

原题面在 Q3.3 里写了一段很直白的话:

「我们强烈建议你不要用 ChatGPT、Gemini、Claude 或任何其它大语言模型来替你写这段代码。这个 lab 是选修的,你只会剥夺自己一次精心设计过的、亲手把机器学习做一遍的机会。」

这个要求在 Q3.3 上尤其成立——从零手写一次多头注意力,是这个 lab 唯一无法被替代的收获。所以本页的定位是写完之后的对照与查错:

  1. 先自己写。卡住了就只读「题目在问什么」和「数学依据」两小节,跳过代码。
  2. 每写完一个模块,立刻跑形状与不变量检查(第 7 节给了一套现成的),不要攒到最后。20000 步跑完才发现 rearrange 写反了,是这个 lab 里最贵的错误。
  3. 训完之后把你的图和本页的真实运行结果图逐张对比。图对不上,先看每题末尾的「常见错误」。
  4. 最后看「参考实现」确认细节。

另外提前说明两处本页会重点展开、但题面没讲的东西:(一)Q2.2 的官方提示里有一条过期信息(关于 $t$ 的形状),照抄会直接报错;(二)Part 5 训出来的隐空间 DiT,损失几乎不下降但样本是好的,这是本 lab 最容易被误判成「没训起来」的地方,第 6 节会用一个可以手算的下界把它解释清楚。

MNIST 上的条件概率路径,t 从 0 到 1 的五个时刻
Part 1 的热身图,也是理解整个 lab 的锚点:条件概率路径 $p_t(\cdot\mid z)=\N(\alpha_t z,\beta_t^2 I)$ 在 MNIST 上的样子,从左到右 $t=0,0.25,0.5,0.75,1$。$t=0$ 是纯高斯噪声,$t=1$ 是数据本身,中间是线性插值加噪。这张图验证的是:Lab 2 里在 $\R^2$ 上写的 sample_conditional_path 一个字都不用改就能作用在 b 1 32 32 上——因为 $\alpha_t,\beta_t$ 是标量,逐元素广播与维度无关。同时它也直观地说明了任务难度:模型要学会的,是把最左边那团噪声逆着推回最右边那些结构极强的笔画。

1. 骨架的三处升级:从 Lab 2 到 Lab 3

Part 0 和 Part 1 是「回收站」,把 Lab 1/2 的类原样搬了过来。但有三处改动是刻意的,看懂它们能省掉后面一半的调试时间。

1.1 升级一:Sampleable 变成 LabeledSampleable

Lab 2 的分布只需要能采样:sample(n) -> x。Lab 3 要做条件生成,所以采样必须同时返回标签:

class LabeledSampleable(ABC):
    @abstractmethod
    def sample(self, num_samples: int) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
        """
        Returns:
            - samples: b d   (或 b c h w)
            - labels:  b
        """

注意标签的形状是 (b,),是一维整数张量,因为它要直接喂给 nn.Embedding。而 IsotropicGaussian(源分布 $p_{\text{init}}$)仍然只实现无标签的 Sampleable——高斯噪声谈不上「类别」,这个区分是有意义的。

1.2 升级二:时间 $t$ 的形状从 (b, 1) 变成 (b,)

这是 Lab 2 → Lab 3 最容易踩的一个坑,而且它贯穿全 lab。Lab 2 里数据是 (b, dim),时间约定成 (b, 1),广播天然对齐。Lab 3 的数据是 (b, c, h, w),如果还用 (b, 1),广播就会把时间维当成通道维,结果完全错乱。

官方的解法是:对外统一把 $t$ 定成 (b,),进到路径类内部再按数据的秩自己撑开:

class GaussianConditionalProbabilityPath(ConditionalProbabilityPath):
    def __init__(self, p_data, p_simple_shape, alpha, beta):
        ...
        # p_simple_shape=[1,32,32] 时,这行等价于 Rearrange('b -> b 1 1 1')
        # p_simple_shape=[2]       时,等价于 Rearrange('b -> b 1')
        self.rearrange_scalar = Rearrange(f'b -> b{" 1" * len(p_simple_shape)}')

    def sample_conditional_path(self, z, t):
        alpha_t = self.rearrange_scalar(self.alpha(t))   # (b, 1, 1, 1)
        beta_t  = self.rearrange_scalar(self.beta(t))    # (b, 1, 1, 1)
        return alpha_t * z + beta_t * torch.randn_like(z)

这个设计非常漂亮:同一份路径代码,既能作用在 (b, 2) 的玩具数据上(Sanity Check 2.4),也能作用在 (b, 1, 32, 32) 的图像上(Part 3),还能作用在 (b, 128, 4, 4) 的隐变量上(Part 5)——因为「秩」这个信息在构造时就被烘焙进 rearrange_scalar 里了。

题面里的一处过期提示

Q2.2 的官方 Hint 第 3 条写着:

「You can sample $t\sim\mathcal{U}[0,1]$ using torch.rand(batch_size, 1, 1, 1).」

照抄会直接报错。 这条提示是从某个更早的版本遗留下来的:在那个版本里路径类不做 rearrange_scalar,需要调用方自己把 $t$ 撑成四维。但在你手上这份 notebook 里,Alpha.__call__、Beta.__call__、sample_conditional_path、conditional_vector_field 的 docstring 全都写着 t: b,内部也全都做了 rearrange_scalar。传 (b,1,1,1) 进去会被再撑成 (b,1,1,1,1,1,1),与 z 的 (b,1,32,32) 广播时维度对不上而崩掉。

正确写法是 torch.rand(batch_size)。 判断依据不要看 Hint,要看 Alpha.dt 的实现——它里面写着 t = t.unsqueeze(1) 然后 vmap(jacrev(self))(t),最后 .view(-1),这只有在输入是一维时才成立。

1.3 升级三:MNISTSampler 与一处性能改写

官方的 MNISTSampler 长这样(简化):每次 sample 都用 torch.randperm 取一批下标,然后对每一张图调用 PIL 的 Resize((32,32)) 和 ToTensor()。这在功能上没问题,但在 20000 步 × batch 256 的训练里,等于要做 512 万次单张 PIL 变换,会成为纯 CPU 瓶颈——GPU 大部分时间在等数据。

本页的实测版本把它改成了一次性预处理并缓存到 GPU:

class MNISTSampler(nn.Module, LabeledSampleable):
    def __init__(self):
        super().__init__()
        dataset = datasets.MNIST(
            root='./data', train=True, download=True,
            transform=transforms.Compose([
                transforms.Resize((32, 32)),
                transforms.ToTensor(),
                transforms.Normalize((0.1305,), (0.2891,)),
            ])
        )
        loader = torch.utils.data.DataLoader(dataset, batch_size=2048,
                                             shuffle=False, num_workers=4)
        xs, ys = [], []
        for xb, yb in loader:
            xs.append(xb); ys.append(yb)
        self.register_buffer('images', torch.cat(xs))                 # (60000, 1, 32, 32)
        self.register_buffer('labels', torch.cat(ys).to(torch.int64)) # (60000,)
        self.dummy = nn.Buffer(torch.zeros(1))

    def sample(self, num_samples: int):
        idx = torch.randint(0, self.images.shape[0], (num_samples,),
                            device=self.images.device)               # (b,)
        # 高级索引本身会返回副本;再 clone 一次,确保调用方的原地改写
        # (CFG 的 label dropout 会写 y[mask] = 10)不会污染缓存的数据集
        return self.images[idx].clone(), self.labels[idx].clone()

整个数据集常驻显存的代价是 $60000\times1\times32\times32\times4\ \text{B}\approx245$ MB,对 16 GB 显存完全不是问题,换来的是训练吞吐提升约一个数量级。

这个改动改变数学吗?只改变了一点,而且是往好的方向:randperm 是无放回抽样,randint 是有放回抽样。理论上我们要的是从 $\data$ 独立同分布采一个 mini-batch,有放回才是那个理想设定;无放回反而引入了 batch 内的负相关。对 SGD 的收敛性没有实际影响,工程上更简单。

.clone() 不是多余的

CFG 的训练循环里有一行原地写:y[xi < self.eta] = self.null_label。在官方实现里 y 每次都是新造的张量,写坏了也无所谓;但一旦你把数据集缓存成 buffer,self.labels[idx] 返回的虽然是副本(高级索引会复制),语义上却很容易在后续重构中变成视图。加一个 .clone() 是零成本的保险——否则你会看到一个极其诡异的现象:训着训着数据集里的标签越来越多变成 10,模型逐渐退化成纯无条件模型。

更稳妥的做法是在 trainer 里写 y = y.clone() 或者用非原地的 torch.where(xi < eta, null, y)。

1.4 顺带一提:Trainer 加了什么

Lab 3 的 Trainer 基类比 Lab 2 多了三样东西,都值得知道:

  • 线性 warmup:前 warmup_steps(默认 500)步把学习率从 0 线性升到 lr。transformer 对初期的大梯度非常敏感,没有 warmup 很容易在前几百步崩掉。
  • checkpoint 回调:每 ckpt_every 步存一次权重 + 采一批样本存图。强烈建议用起来——20000 步要跑几十分钟,能在第 1000 步就看出「样本是不是完全的噪声」,比跑完再看划算得多。
  • model_size_b:打印模型大小。DiT 约 40.3 MiB,VAE 约 6.2 MiB,可以用来快速核对你的架构参数有没有写错(比如把 dim 写成 512 会让大小翻四倍)。

2. Part 2 · Classifier-Free Guidance

先把理论理顺,代码只有五行。Lecture 3-B 已经完整讲过引导,这里只重做一遍与代码直接对应的那条推导链。

2.0 从条件生成到 CFG:三步推导

第一步:条件生成本身不需要新理论。固定一个标签 $y$,把 $\data(\cdot\mid y)$ 当成新的数据分布,Lab 2 的整套 flow matching 原封不动就能用:

$$ \mathcal{L}^{\text{guided}}_{\text{CFM}}(\theta;y)=\E_{z\sim \data(\cdot\mid y),\,t\sim U[0,1],\,x\sim p_t(\cdot\mid z)}\norm{u_t^\theta(x\mid y)-u_t^{\text{ref}}(x\mid z)}^2 $$

再把 $y$ 也放进期望里(而不是固定住),就得到了「一个网络同时学会所有类别」的目标。到这一步理论上就已经完事了:$u^\theta_t(x\mid y)$ 的最优解就是条件边际向量场 $u_t(x\mid y)$,用它做 ODE 采样得到的就是 $\data(\cdot\mid y)$。

第二步:但人们发现「更像那一类」比「正确采样」更讨喜。回忆高斯路径下向量场与得分的线性关系(Lab 2 Q3.3 的那条):对 $a_t=\frac{\dot\alpha_t}{\alpha_t}$、$b_t=-\frac{\dot\beta_t\beta_t\alpha_t-\dot\alpha_t\beta_t^2}{\alpha_t}$,有

$$ u_t(x\mid y)=a_t x+b_t\,\grad\log p_t(x\mid y). $$

对条件得分用一次贝叶斯:$p_t(x\mid y)=\frac{p_t(x)p_t(y\mid x)}{p_t(y)}$,两边取 $\log$ 再对 $x$ 求梯度,$\log p_t(y)$ 与 $x$ 无关直接消失:

$$ \grad\log p_t(x\mid y)=\grad\log p_t(x)+\grad\log p_t(y\mid x). $$

代回去:

$$ u_t(x\mid y)=\underbrace{a_tx+b_t\grad\log p_t(x)}_{=\,u_t(x)\ \text{(无条件向量场)}}+\;b_t\grad\log p_t(y\mid x). $$

第二项 $\grad\log p_t(y\mid x)$ 就是一个「带噪分类器」:它指向「让当前样本更像类别 $y$」的方向。把这一项的权重放大 $w$ 倍,就得到引导向量场:

$$ \tilde u_t(x\mid y)=u_t(x)+w\,b_t\grad\log p_t(y\mid x). $$

第三步:消掉分类器。上面那条恒等式反解出 $b_t\grad\log p_t(y\mid x)=u_t(x\mid y)-u_t(x)$,代进去:

$$ \boxed{\ \tilde u_t(x\mid y)=(1-w)\,u_t(x\mid\varnothing)+w\,u_t(x\mid y)\ } $$

其中 $u_t(x)$ 被重写成 $u_t(x\mid\varnothing)$:把「无条件」当成一个特殊的类别 $\varnothing$。这就是 classifier-free guidance(无分类器引导)——不需要训分类器,只需要让同一个网络偶尔看不见标签。

三个特殊值,可以拿来当测试用例
  • $w=1$:$\tilde u_t=u_t(x\mid y)$,纯条件,无引导。此时 CFG 完全不起作用,采到的是真正的 $\data(\cdot\mid y)$。
  • $w=0$:$\tilde u_t=u_t(x\mid\varnothing)$,纯无条件,标签被完全忽略。
  • $y=\varnothing$:无论 $w$ 取多少,$(1-w)u_t(x\mid\varnothing)+w\,u_t(x\mid\varnothing)=u_t(x\mid\varnothing)$,引导强度完全失效。这条性质在第 3 节读 lab3-dit-samples.png 的最后一行时会用上。

前两条是本 lab 数值测试里的两项(误差均为 0),第三条是理解那张三联图的钥匙。

训练侧的改动只有一句话:以概率 $\eta$ 把标签 $y$ 换成 $\varnothing$。这样同一份参数既见过「带标签的样本」(学 $u_t(x\mid y)$),也见过「不带标签的样本」(学 $u_t(x\mid\varnothing)$)。本 lab 取 $\varnothing\triangleq10$,即 MNIST 十个数字之后的第十一个槽位。

2.1 Q2.2 · CFGTrainer.get_train_loss

题目在问什么。 实现下面这个期望的单批次蒙特卡洛估计:

$$ \mathcal{L}_{\text{CFM}}(\theta)=\E_{\square}\norm{u_t^\theta(x\mid y)-u_t^{\text{ref}}(x\mid z)}^2,\qquad \square=\begin{cases}(z,y)\sim\data(z,y)\\ y\leftarrow\varnothing\ \text{以概率}\ \eta\\ t\sim U[0,1]\\ x\sim p_t(\cdot\mid z)\end{cases} $$

数学依据。 与 Lab 2 的条件 flow matching 定理逐字对应,只多了第二行。回归目标是条件向量场 $u_t^{\text{ref}}(x\mid z)$(可解析计算),而 $L^2$ 最优解是它在给定 $(x,t,y)$ 下的条件期望,也就是我们真正想要的边际条件向量场 $u_t(x\mid y)$。标签丢弃只是把 $y$ 的边缘分布从 $p(y)$ 改成了 $(1-\eta)p(y)+\eta\delta_\varnothing$,不影响这个论证——在 $y=\varnothing$ 的那一支上,条件期望自动收敛到 $\E[u_t(x\mid z)\mid x,t]=u_t(x)$。

参考实现。

class CFGTrainer(Trainer):
    def __init__(self, path, eta: float, null_label: int, eps: float = 0.001, **kwargs):
        assert eta > 0 and eta < 1
        super().__init__(**kwargs)
        self.eta = eta
        self.eps = eps
        self.path = path
        self.null_label = null_label

    def get_train_loss(self, batch_size: int) -> torch.Tensor:
        # Step 1: 从 p_data 采 (z, y)
        z, y = self.path.p_data.sample(batch_size)   # z: (b, c, h, w),  y: (b,)

        # Step 2: 以概率 eta 把标签替换成空标签
        xi = torch.rand(y.shape[0], device=y.device)             # (b,)
        y[xi < self.eta] = self.null_label

        # Step 3: 采 t 和 x
        #   注意 t 的形状是 (b,),不是 (b,1,1,1)
        #   乘 (1 - eps) 是为了避开 t=1
        t = torch.rand(batch_size, device=z.device).to(z) * (1 - self.eps)  # (b,)
        x = self.path.sample_conditional_path(z, t)                        # (b, c, h, w)

        # Step 4: 回归
        ut_theta = self.model(x, t, y)                       # (b, c, h, w)
        ut_ref   = self.path.conditional_vector_field(x, z, t)  # (b, c, h, w)
        return torch.square(ut_theta - ut_ref).mean()

为什么这样写。

  • xi = torch.rand(y.shape[0]) 而不是 torch.rand(1)。 丢弃必须是逐样本独立的:一个 batch 里应该同时存在带标签样本和空标签样本。如果整批一起丢,梯度的方差会大很多,而且 batch norm 类统计量(这里没有,但 VAE 里有 GroupNorm)会在两种模式间来回跳。
  • t 的形状是 (b,)。 上一节已经详述:路径类内部会用 Rearrange('b -> b 1 1 1') 自己撑开。这一点在 Sanity Check 2.4 里同样成立(那里 p_simple_shape=[2],撑成 b 1),所以同一份 trainer 代码在二维玩具数据和图像上都能跑——这正是它被设计成 (b,) 的原因。
  • * (1 - self.eps) 是必须的,不是保险丝。 展开条件向量场: $$ u_t(x\mid z)=\left(\dot\alpha_t-\frac{\dot\beta_t}{\beta_t}\alpha_t\right)z+\frac{\dot\beta_t}{\beta_t}x $$ 本 lab 用 LinearAlpha/LinearBeta,即 $\alpha_t=t$、$\beta_t=1-t$,于是 $\frac{\dot\beta_t}{\beta_t}=\frac{-1}{1-t}$,在 $t\to1$ 时发散。虽然代数上 $u_t(x\mid z)$ 化简后等于 $z-\epsilon$(有限),但代码是按原式逐项算的,$t=1$ 会出现 $-\infty\cdot 0$,直接得到 nan。eps=1e-3 把采样区间截到 $[0,0.999]$,代价可以忽略($t$ 是连续均匀分布,$\{1\}$ 是零测集),收益是训练不会莫名其妙 NaN。采样时同样要截:torch.linspace(0, 0.999, num_timesteps)。
  • .to(z) 而不是只写 device=。 .to(z) 会把 dtype 也对齐到 z。如果你后面上 AMP / bf16 训练,这一行能省掉一次类型不匹配。
  • .mean() 而不是 .sum()。 对所有元素取平均,等价于「每维平均平方误差」。这让损失的数值不随图像分辨率变化,学习率也就不用跟着分辨率调。记住这个约定,第 6 节讨论损失下界时要用。
Q2.2 的五个常见错误
  • 按 Hint 写 torch.rand(batch_size, 1, 1, 1)——直接报广播错误。见上一节。
  • torch.rand 写成 torch.randn——题面自己都提醒了。$t$ 必须是 $[0,1]$ 上的均匀分布;用标准正态会采到负时间和大于 1 的时间,$\beta_t=1-t$ 变成负数或大于 1,路径彻底失效。症状是损失在几百步内爆到 $10^3$ 量级。
  • 丢标签时用了 > 而不是 <——变成以概率 $1-\eta$ 丢弃。$\eta=0.35$ 时你实际丢了 65%,条件分支训练不足,$w$ 越大样本越糊。这个错误不会报错,只会让结果变差,很难发现。
  • 把 z 和 x 搞反——conditional_vector_field(x, z, t) 的参数顺序是「先当前点,后条件变量」。写成 (z, x, t) 不会报错(形状一样),但回归目标完全错误。
  • 忘了 y 是原地修改的——见 1.3 节的 .clone() 讨论。

2.2 Q2.3 · MLPConditionalVectorField.forward

题目在问什么。 用最朴素的方式把 $(x,t,y)$ 三样输入喂进一个 MLP:拼接。这是自检用的小模型,跑在二维高斯混合上。

数学依据。 没有什么深刻的数学,但有一个重要的表示问题:$y$ 是离散的类别,不能当成实数直接拼进去(那样会隐含「类别 3 比类别 1 大 2」这种荒谬的序关系)。标准做法是查表——把每个类别映射到一个可学习的向量:

$$ y\in\set{0,1,\dots,C-1,\varnothing}\ \longmapsto\ e_y\in\R^{d_{\text{class}}},\qquad e\ \text{是可训练参数}. $$

参考实现。

class MLPConditionalVectorField(ConditionalVectorField):
  def __init__(self, dim: int, hidden_dim: int, class_dim: int, num_classes: int):
    super().__init__()
    self.mlp = MLP([dim + class_dim + 1, hidden_dim, hidden_dim, dim])
    self.class_embedding = nn.Embedding(num_classes + 1, class_dim)
    #                                   ^^^^^^^^^^^^^^^ 多出来的那一格就是空标签

  def forward(self, x: torch.Tensor, t: torch.Tensor, y: torch.Tensor):
      xyt = torch.cat([
          x,                          # (b, dim)
          self.class_embedding(y),    # (b, class_dim)
          t.unsqueeze(-1),            # (b, 1)
      ], dim=-1)                      # (b, dim + class_dim + 1)
      return self.mlp(xyt)            # (b, dim)

为什么这样写。

  • num_classes + 1 里的那个 +1,就是 $\varnothing$。 这是 CFG 落地的全部工程代价——一行代码里的一个加一。Sanity Check 里 num_classes=3(三个高斯模式),所以嵌入表有 4 行,null_label=3 指向最后一行。Part 3 的 DiT 里则写成 n_classes=11(10 个数字 + 1 个 $\varnothing$),是同一件事的另一种写法。如果你写成 nn.Embedding(num_classes, class_dim),训练会在第一次抽到空标签时抛 index out of range——这算是好消息,因为它会立刻暴露。
  • t.unsqueeze(-1)。 $t$ 进来是 (b,),拼接要求所有张量除拼接维外形状一致,所以要撑成 (b, 1)。这里的 $t$ 是原始标量,没有做 Fourier 编码——对二维玩具问题够用,但对图像就不够了,这正是 Q3.1 存在的理由。
  • 输入维度 dim + class_dim + 1。 三段拼接后的总宽度。写错这个数字会在第一次前向时报矩阵乘法维度不匹配,属于「好错误」。
Q2.3 的常见错误
  • 把 y 当浮点数直接 torch.cat——nn.Embedding 要求整型索引,传浮点会报错;而如果你绕过 embedding 直接拼 y.float().unsqueeze(-1),代码能跑,但模型被迫在类别上做插值,Sanity Check 的三个模式会糊成一团。
  • dim=-1 写成 dim=0——把三个张量沿 batch 维摞起来,形状变成 (3b, ...),报错或者算出完全错误的结果。
  • y 的 dtype 不是 int64——MNIST 的 label 默认就是 int64,但如果你自己构造标签时用了 torch.tensor([0,1,2]) 之外的写法(比如从 torch.zeros 转),要记得 .long()。

2.3 Sanity Check 2.4 · 在二维高斯混合上验证 CFG

Q2.2 + Q2.3 写完之后,题面给了一个只要三分钟的自检:三个模式的对称高斯混合(半径 2、标准差 0.2),$\eta=0.25$,$\varnothing=3$,训练 3000 步。务必先跑通这个再去碰 Part 3——它能在三分钟内验证你的训练循环是对的,而 DiT 要几十分钟。

MLP 自检模型的训练损失曲线
自检模型(0.259 MiB)3000 步的损失曲线,从 2.8910 降到 0.8401。这张图验证的是训练循环本身:损失单调下降并在 0.8 附近走平。注意它不会收敛到 0——这是 Lab 2 就讲过的道理,条件 flow matching 的损失有一个由数据分布决定的不可约下界(等于回归目标关于 $x_t$ 的条件方差)。如果你的曲线降到接近 0,说明回归目标写成了与输入平凡相关的东西;如果完全不降,先查 $t$ 的形状。
CFG 在二维高斯混合上的三联图:目标分布、三个模式的条件生成、无条件生成
Sanity Check 2.4 的三联图($w=1$)。左:目标分布,三个模式各占三分之一。中:给定三个不同标签各生成 250 个样本,用颜色区分——每种颜色只落在自己那一个模式上,没有一个样本跑错模式,这说明条件分支学对了,$u_t^\theta(x\mid y)$ 确实在把噪声导向 $\data(\cdot\mid y)$。右:把标签全设为 $\varnothing=3$,生成 750 个样本,三个模式都被覆盖且比例大致均等——这说明空标签分支学到的是完整的无条件分布,而不是某一类的残影。两张图合起来就是 CFG 训练成功的完整证据:一个网络,两种行为。
这张自检图能诊断出什么
  • 中间图颜色混杂 → 条件信息没传进去。查 class_embedding(y) 有没有真的被拼进 MLP 输入。
  • 右边图只覆盖一两个模式 → $\eta$ 太小(空标签样本太少)或者丢标签的比较符号写反了。
  • 右边图和中间图长得一样(只有一个模式) → null_label 传错,空标签指到了某个真实类别上。
  • 点云整体偏离目标位置 → 回归目标或路径写错,回去查 conditional_vector_field 的参数顺序。

3. Part 3 · 从零搭一个 Diffusion Transformer

这是本 lab 篇幅最大的一部分:五个 TODO 把一个完整的 DiT 搭出来。先看总装图,后面每一题都是它的一个零件:

阶段模块输入形状输出形状题号
时间编码FourierEncoder + MLP(b,)(b, 256)Q3.1
类别编码nn.Embedding(11, 256)(b,)(b, 256)Q3.5
切块Patchifier(b, 1, 32, 32)(b, 64, 256)Q3.2
主干DiffusionTransformer(8 层)(b, 64, 256) + (b, 256)(b, 64, 256)Q3.3
还原Depatchifier(b, 64, 256)(b, 1, 32, 32)Q3.4

其中 64 = (32/4)^2 是 token 数(patch size 取 4),256 是隐藏维度。整个网络实现的函数是 $u_t^\theta:\R^{1\times32\times32}\times[0,1]\times\set{0,\dots,10}\to\R^{1\times32\times32}$——输入输出同形,因为向量场和它作用的点住在同一个空间。

3.1 Q3.1 · FourierEncoder

题目在问什么。 把标量 $t\in[0,1]$ 映射成 $d$ 维向量:

$$ t^{\text{emb}}=\big[\cos(2\pi w_1 t),\ \dots,\ \cos(2\pi w_{d/2} t),\ \sin(2\pi w_1 t),\ \dots,\ \sin(2\pi w_{d/2} t)\big]^\top, $$

其中频率 $w_i\sim\N(0,1)$ 在初始化时随机抽取(本实现里做成 nn.Parameter,即可训练)。

数学依据:为什么标量 $t$ 不能直接喂进网络?

谱偏置与随机 Fourier 特征

带 ReLU/SiLU 的 MLP 有很强的谱偏置(spectral bias):它天然倾向于先学会输入的低频依赖,高频部分要花指数级更多的训练量。而向量场对 $t$ 的依赖恰恰是高频的——本 lab 的条件向量场展开后含有 $\frac{\dot\beta_t}{\beta_t}=\frac{-1}{1-t}$,在 $t\to1$ 附近变化极快。如果只给网络一个标量 $t$,它很难分辨 $t=0.98$ 和 $t=0.99$ 这两个在行为上差别巨大的时刻。

随机 Fourier 特征把这件事一次性解决:$\cos(2\pi w t)$ 与 $\sin(2\pi w t)$ 在 $w$ 较大时对 $t$ 的微小变化极其敏感。由三角恒等式

$$ \cos(2\pi w(t+\delta))=\cos(2\pi wt)\cos(2\pi w\delta)-\sin(2\pi wt)\sin(2\pi w\delta), $$

可以看出 $(\cos,\sin)$ 这一对共同构成了对平移的完整表示——这也是为什么必须两个都要,只留 $\cos$ 会丢掉符号信息($\cos$ 是偶函数,无法区分 $t$ 和 $-t$)。

更定量地说,这是 Rahimi & Recht 的随机特征方法:一组随机频率上的 $(\cos,\sin)$ 内积,是平移不变核的无偏估计——

$$ \E_{w\sim\N(0,1)}\big[\cos(2\pi w t)\cos(2\pi w s)+\sin(2\pi w t)\sin(2\pi w s)\big]=\E_w\big[\cos(2\pi w(t-s))\big]=e^{-2\pi^2(t-s)^2}. $$

也就是说,这个编码等价于给时间装了一个宽度约 $\frac{1}{2\pi}$ 的高斯核——相近的 $t$ 表示相似,远离的 $t$ 表示近乎正交,正是网络需要的。

参考实现。

class FourierEncoder(nn.Module):
    def __init__(self, dim: int):
        super().__init__()
        assert dim % 2 == 0
        self.half_dim = dim // 2
        self.weights = nn.Parameter(torch.randn(1, self.half_dim))   # (1, d/2)

    def forward(self, t: torch.Tensor) -> torch.Tensor:
        """ t: (b,)  ->  embeddings: (b, dim) """
        t = t.view(-1, 1)                            # (b, 1)
        freqs = 2 * math.pi * t * self.weights       # (b, 1) * (1, d/2) -> (b, d/2)
        cos_emb = torch.cos(freqs)                   # (b, d/2)
        sin_emb = torch.sin(freqs)                   # (b, d/2)
        return torch.cat([cos_emb, sin_emb], dim=-1) # (b, dim)

为什么这样写。

  • t.view(-1, 1) 加 weights 形状 (1, half_dim),靠广播一次得到 (b, half_dim)。 这是本题唯一的形状技巧:外积用广播实现,不需要 einsum。
  • 2 * math.pi 不能省。 省掉它并不会让模型不能训(相当于把 $w$ 的尺度缩小 $2\pi$ 倍),但会让频率分布偏低,时间分辨能力下降。既然公式写了就照写。
  • nn.Parameter 而非 register_buffer。 题面给的公式里 $w_i$ 是固定的随机数,但做成可训练参数只有好处:网络可以自己把频率调到任务需要的量级。注意这也意味着它会被 model.parameters() 计入,别对模型大小感到意外。
  • assert dim % 2 == 0。 因为要对半分给 $\cos$ 和 $\sin$。
一个零成本的自检不变量

由 $\cos^2\theta+\sin^2\theta=1$,输出张量每一行的平方和恒等于 half_dim,与 $t$ 无关:

fe = FourierEncoder(dim=64)
emb = fe(torch.rand(8))
assert torch.allclose((emb ** 2).sum(-1), torch.full((8,), 32.0))  # 32 = 64 / 2

再加一条端点检查:$t=0$ 时 $\cos=1,\sin=0$,所以 fe(torch.zeros(1)) 应当是 [1,...,1, 0,...,0]。这两条能抓出「$\cos/\sin$ 拼反」「漏乘 $t$」「维度切错」等几乎所有实现错误,写完立刻跑,一秒钟的事。

Q3.1 的常见错误
  • 输出维度是 2 * dim 而不是 dim——如果你写成 self.weights = nn.Parameter(torch.randn(1, dim)) 再拼接,输出就是 2*dim。后面 MLP 的输入维度对不上,会在 DiffusionTransformerFlowModel 里才炸,排查起来很绕。
  • 忘了 t.view(-1, 1)——t 是 (b,),直接乘 (1, half_dim) 会得到 (b, half_dim)……只在 b == 1 或 b == half_dim 时侥幸不报错,其余情况报广播错误。别赌。
  • 用 torch.arange 造固定的对数间隔频率——那是 transformer 原始论文的正弦位置编码,也能用,但和题面给的公式不是一回事,测试会不通过。

3.2 Q3.2 · Patchifier

题目在问什么。 把图像 (b, c, 32, 32) 变成 token 序列 (b, n, d),其中 $n=(32/p)^2$。两步:一层卷积把空间降到 $32/p$、通道升到 $d$;一次 Rearrange 把两个空间维压平成序列维。

数学依据:为什么 kernel = stride = patch 的卷积就是 patch embedding?

卷积与「切块后各自线性投影」的等价性

ViT/DiT 论文里对 patch embedding 的描述是:把图像切成 $p\times p$ 的不重叠小块,每块拉平成 $c p^2$ 维向量,再乘同一个矩阵 $W\in\R^{d\times cp^2}$ 得到 $d$ 维 token。

现在看 nn.Conv2d(c, d, kernel_size=p, stride=p) 在做什么。它的第 $(i,j)$ 个输出位置是

$$ \text{out}[:,i,j]=\sum_{c'=1}^{c}\sum_{a=0}^{p-1}\sum_{b=0}^{p-1}K[:,c',a,b]\cdot x[c',\,ip+a,\,jp+b]\;+\;\text{bias}. $$

因为 stride 等于 kernel size,感受野 $\{ip,\dots,ip+p-1\}\times\{jp,\dots,jp+p-1\}$ 恰好就是第 $(i,j)$ 个 patch,且不同 $(i,j)$ 的感受野互不重叠、恰好铺满图像。把卷积核 $K$ 沿 $(c',a,b)$ 拉平成矩阵 $W$,上式就逐字变成 $W\cdot\text{flatten}(\text{patch}_{ij})+\text{bias}$。

所以两者不是「近似」,而是同一个线性映射的两种写法。 用卷积写的好处是:一行代码、cuDNN 高度优化、不需要手动 unfold。

参考实现。

class Patchifier(nn.Module):
  def __init__(self, img_size: int, patch_size: int, c_in: int, dim: int):
    super().__init__()
    assert img_size % patch_size == 0, "Image size must be divisible by patch size"
    self.net = nn.Sequential(
        nn.Conv2d(c_in, dim, kernel_size=patch_size, stride=patch_size),
        #   (b, c, H, W) -> (b, d, H/p, W/p)
        Rearrange('b d h w -> b (h w) d'),
        #   -> (b, n, d),  n = (H/p) * (W/p)
    )

  def forward(self, x: torch.Tensor) -> torch.Tensor:
    return self.net(x)

为什么这样写。

  • c_in 是参数而不是硬编码的 1。 Part 5 会把同一个类用在 (b, 128, 4, 4) 的隐变量上(c_in=128、patch_size=1)。不要硬编码通道数,否则 Part 5 要重写。
  • Rearrange('b d h w -> b (h w) d') 的顺序至关重要。 括号里的 (h w) 表示「$h$ 是慢变维、$w$ 是快变维」,即 token 按行优先排列:第 0 个 token 是左上角,第 1 个是它右边一格。Depatchifier 必须用完全对应的写法折回去,否则图像会被转置或打乱。
  • assert img_size % patch_size == 0。 否则卷积会悄悄丢掉边缘像素(PyTorch 的 Conv2d 默认 floor),你会得到一个尺寸对不上的 token 数。
一个能抓出 rearrange 写反的检查

patch 划分必须是空间局部的:只修改图像左上角那个 $4\times4$ 块,应当只有一个 token 发生变化。

P = Patchifier(img_size=32, patch_size=4, c_in=1, dim=48)
x1 = torch.zeros(1, 1, 32, 32)
x2 = x1.clone(); x2[0, 0, 0:4, 0:4] = 1.0
d = (P(x1) - P(x2)).abs().sum(-1)[0]        # (n,)
assert (d > 1e-6).sum().item() == 1         # 只有 1 个 token 变了

这条检查在本 lab 的数值测试里是实打实通过的。它能抓出「用 reshape/view 代替 Rearrange 但维度顺序错了」这类经典错误——那种写法会让一个 patch 的改动扩散到多个 token 上。

Q3.2 的常见错误
  • 用 x.view(b, n, d) 直接压平——卷积输出的内存布局是 (b, d, h, w),通道维在前。直接 view 会把通道和空间搅在一起,得到语义完全错误的 token。必须先 permute/Rearrange。
  • 卷积写成 kernel_size=patch_size, stride=1——变成滑窗,输出 token 数是 $(32-p+1)^2$ 而不是 $(32/p)^2$,位置编码的长度立刻对不上。
  • 加了 padding——同上,会改变输出尺寸。这里不要 padding。
  • 写成 'b d h w -> b (w h) d'——列优先。代码能跑、测试的形状检查也能过,但图像会被转置。只有上面那条局部性检查 + 最终采样图能发现。

3.3 Q3.3 · MHA + DiffusionTransformerLayer + DiffusionTransformer

本 lab 最重的一题,一次要写三个类。题面建议自顶向下想、自底向上写,这里按后者组织。

3.3.1 MHA:从零手写多头自注意力

题目在问什么。 实现

$$ \text{MHA}(X)=W_O\,\text{concat}_{h=1}^{H}\Big(\softmax\!\Big(\tfrac{Q_hK_h^\top}{\sqrt{d_h}}\Big)V_h\Big),\qquad Q_h=XW_Q^{(h)},\ K_h=XW_K^{(h)},\ V_h=XW_V^{(h)}, $$

其中 $X\in\R^{n\times d}$ 是 token 序列,$H$ 是头数,$d_h=d/H$ 是每个头的维度。

为什么要除以 $\sqrt{d_h}$:一个方差论证

假设 $q,k\in\R^{d_h}$ 的各分量独立、均值 0、方差 1(初始化时大致如此)。那么点积

$$ q\cdot k=\sum_{i=1}^{d_h}q_ik_i $$

的均值是 $0$,方差是

$$ \Var\Big(\sum_i q_ik_i\Big)=\sum_i\Var(q_ik_i)=\sum_i\E[q_i^2]\E[k_i^2]=d_h, $$

即标准差为 $\sqrt{d_h}$。当 $d_h=32$ 时,logits 的典型幅度约 $\pm 5.7$;当 $d_h=128$ 时约 $\pm11$。而 $\softmax$ 在输入相差超过 $\sim10$ 时基本已经饱和成 one-hot——梯度接近 0,注意力层学不动。

除以 $\sqrt{d_h}$ 把方差拉回 1,logits 的典型幅度稳定在 $\pm1$ 附近,与头维度无关。这就是「scaled dot-product attention」里那个 scaled 的全部含义。

注意分母是 $\sqrt{d_h}$(头维度),不是 $\sqrt{d}$(总维度)。 写错会让缩放偏小 $\sqrt{H}$ 倍——模型仍然能训,但收敛更慢,而且与 PyTorch 官方实现对不上号。

参考实现。

class MHA(nn.Module):
  """ Multi-headed self-attention """
  def __init__(self, dim: int, heads: int):
    super().__init__()
    assert dim % heads == 0
    self.heads = heads
    self.head_dim = dim // heads
    self.scale = self.head_dim ** -0.5        # 1 / sqrt(d_h)
    self.qkv = nn.Linear(dim, 3 * dim, bias=False)
    self.proj = nn.Linear(dim, dim)

  def forward(self, x: torch.Tensor) -> torch.Tensor:
    """ x: (b, n, d) -> (b, n, d) """
    # 1. 一次线性层同时算出 q, k, v
    q, k, v = self.qkv(x).chunk(3, dim=-1)                       # 各 (b, n, d)

    # 2. 把「头」折进 batch 维:(b, n, h*e) -> (b*h, n, e)
    q = rearrange(q, 'b n (h e) -> (b h) n e', h=self.heads)     # (b*h, n, e)
    k = rearrange(k, 'b n (h e) -> (b h) n e', h=self.heads)
    v = rearrange(v, 'b n (h e) -> (b h) n e', h=self.heads)

    # 3. 缩放点积
    scores = torch.einsum('b i e, b j e -> b i j', q, k) * self.scale   # (b*h, n, n)
    attn = torch.softmax(scores, dim=-1)                               # 沿 key 维归一化

    # 4. 加权求和 value
    out = torch.einsum('b i j, b j e -> b i e', attn, v)               # (b*h, n, e)

    # 5. 头折回特征维
    out = rearrange(out, '(b h) n e -> b n (h e)', h=self.heads)       # (b, n, d)

    # 6. 输出投影
    return self.proj(out)

为什么这样写。

  • nn.Linear(dim, 3 * dim) 一次算 qkv。 数学上等价于三个独立的 nn.Linear(dim, dim),但只调一次 GEMM,对 GPU 友好得多。chunk(3, dim=-1) 沿最后一维切成三份。
  • bias=False 是有理由的。 qkv 之后紧接着的是 $\softmax(qk^\top)$,而 $q$ 上的常数偏置 $q\to q+b$ 会给所有 logits 加上 $b\cdot k_j$——这一项不被 softmax 消掉,所以并非完全冗余,但实践中它的作用可以被后面的 LayerNorm 与 proj 的 bias 覆盖。省掉能减参数、也是主流实现的做法。
  • rearrange('b n (h e) -> (b h) n e'):把头折进 batch。 这是本题最漂亮的一步。多头注意力本质上就是「$H$ 个互不相干的单头注意力并行跑」,所以把 $H$ 塞进 batch 维之后,后面的点积、softmax、加权求和全都是普通的三维张量运算,不需要写四维 einsum。折回去时用 '(b h) n e -> b n (h e)',einops 会自动保证与折进去时的顺序一致——这正是用 einops 而非手写 view/permute 的价值。
  • softmax(dim=-1):沿 key 维归一化。 scores 的形状是 (b*h, i, j),$i$ 是 query 下标、$j$ 是 key 下标。归一化必须沿 $j$,含义是「第 $i$ 个 query 对所有 key 的注意力权重之和为 1」。写成 dim=-2 是最隐蔽的一个 bug:形状完全一样,不会报错,但语义变成「每个 key 被所有 query 分配的权重之和为 1」,训练会明显变差却看不出原因。
  • 最后的 self.proj 不能省。 没有它,各个头的输出只是被拼在一起,头与头之间从不交换信息;输出投影负责把 $H$ 个头的结果混合成新的特征。
和 PyTorch 官方实现对拍(强烈建议做)

手写注意力最好的验证方式,是和 torch.nn.functional.scaled_dot_product_attention 逐元素比对:

mha = MHA(dim=64, heads=8).eval()
x = torch.randn(2, 10, 64)
with torch.no_grad():
    mine = mha(x)
    q, k, v = mha.qkv(x).chunk(3, dim=-1)
    q = q.view(2, 10, 8, 8).transpose(1, 2)     # (b, h, n, e)
    k = k.view(2, 10, 8, 8).transpose(1, 2)
    v = v.view(2, 10, 8, 8).transpose(1, 2)
    ref = torch.nn.functional.scaled_dot_product_attention(q, k, v)
    ref = mha.proj(ref.transpose(1, 2).reshape(2, 10, 64))
print((mine - ref).abs().max())     # 实测:8.94e-08

本页的实现实测最大误差 8.94e-08,就是 float32 的舍入噪声。注意这里参考实现里用的是 view(b, n, h, e).transpose(1, 2),而我们的实现用的是 rearrange('b n (h e) -> (b h) n e')——两者对「哪一段特征属于哪个头」的划分方式必须一致,才能对上。 这本身就是对 rearrange 写法的一次强验证:如果你写成 'b n (e h) -> (b h) n e',切分方式就变了,误差会是 $O(1)$ 而非 $10^{-8}$。

3.3.2 DiffusionTransformerLayer:adaLN-Zero 条件调制

题目在问什么。 实现 DiT 论文里的 adaLN-Zero block:

$$ \begin{aligned} (\gamma_1,\beta_1,\alpha_1,\gamma_2,\beta_2,\alpha_2)&=\text{MLP}_{\text{cond}}(c),\quad\text{各}\in\R^{b\times d}\\ h&=\text{LN}_1(x)\odot(1+\gamma_1)+\beta_1\\ x&\leftarrow x+\alpha_1\odot\text{MHA}(h)\\ h&=\text{LN}_2(x)\odot(1+\gamma_2)+\beta_2\\ x&\leftarrow x+\alpha_2\odot\text{FF}(h) \end{aligned} $$

数学依据:AdaLN 在做什么。 这一部分 Lecture 4 §6.1 已经完整讲过,这里只复述关键的一句:LayerNorm 先把每个 token 的特征向量标准化(减均值除标准差),故意抹掉「尺度」和「偏置」信息;AdaLN 紧接着用由条件 $c$ 决定的 $(\gamma,\beta)$ 把尺度和偏置重新写回去。于是时间和类别不是作为额外 token「混」在序列里,而是作为一个全局旋钮,逐通道地控制网络内部每一处激活的强弱。

写成 $x\cdot(1+\gamma)+\beta$ 而不是 $x\cdot\gamma+\beta$,是为了让 $\gamma=0$ 对应「恒等缩放」。这一点和下面的零初始化配合起来才有意义。

adaLN-Zero:三个必须同时做对的细节
  1. nn.LayerNorm(dim, elementwise_affine=False)——仿射参数必须关掉。 LayerNorm 默认自带可学习的 $(\text{weight},\text{bias})$,但在 adaLN 里缩放平移完全由条件网络产生。两者同时存在不会报错,但会造成参数冗余与优化路径的病态(两组参数相乘,尺度不唯一)。「adaptive LayerNorm」这个名字的字面意思就是「归一化的仿射部分是自适应的」,所以固有的那一份必须关。
  2. 六组参数、每组 (b, d),作用到 (b, n, d) 上要 unsqueeze(1)。 条件向量 $c$ 是每个样本一份,而不是每个 token 一份——同一张图的 64 个 token 共享同一组调制参数。所以要把 (b, d) 撑成 (b, 1, d) 再广播到 $n$ 个 token。忘了 unsqueeze(1) 会怎样? (b, n, d) * (b, d) 在 $n\ne d$ 时报错(好消息),在 $n = d$ 时静默地算错(把 token 维当成特征维)。
  3. final_init=True:条件 MLP 最后一层清零。 于是初始时六组参数全为 0,$\alpha_1=\alpha_2=0$,两条残差支路被完全「关掉」,整个 block 退化为恒等映射 $x\mapsto x$。这就是 adaLN-Zero 里 Zero 的含义。
为什么「初始即恒等」对深层网络这么重要

设第 $\ell$ 层为 $x_{\ell+1}=x_\ell+\alpha\,F_\ell(x_\ell)$。前向的方差满足(假设 $F_\ell$ 输出与输入近似独立)

$$ \Var(x_{\ell+1})\approx\Var(x_\ell)+\alpha^2\Var(F_\ell(x_\ell)). $$

$\alpha=1$ 时方差逐层线性累积,$L$ 层之后放大约 $L$ 倍;深度一大,输出的尺度就会失控,必须靠仔细的初始化或额外的 LayerNorm 压回来。而 $\alpha=0$ 时,网络初始就是一个恒等映射的堆叠——无论多少层,前向输出等于输入,反向梯度等于 1,不存在任何尺度爆炸或消失。训练开始后 $\alpha$ 由 0 缓慢长大,相当于网络自己决定何时以及以多大幅度启用每一层。

这个思想在文献里反复出现:ResNet 的 zero-init-$\gamma$、Fixup、ReZero、DiT 的 adaLN-Zero,本 lab Q4.1 里 ResidualBlock 的零初始化第二个卷积,全是同一招。

参考实现。

class DiffusionTransformerLayer(nn.Module):
  def __init__(self, dim: int, heads: int):
    super().__init__()
    # 仿射参数关掉:缩放平移全部交给条件网络产生
    self.norm1 = nn.LayerNorm(dim, elementwise_affine=False)
    self.norm2 = nn.LayerNorm(dim, elementwise_affine=False)

    # 由条件 c 一次预测 6 组调制参数;final_init=True 把最后一层权重和 bias 清零
    self.cond = MLP([dim, dim, 6 * dim], final_init=True)

    self.attn = MHA(dim, heads)
    self.ff   = MLP([dim, 4 * dim, dim])

  def forward(self, x: torch.Tensor, c: torch.Tensor) -> torch.Tensor:
    """ x: (b, n, d),  c: (b, d)  ->  (b, n, d) """
    # (b, 6d) -> 6 个 (b, d) -> 6 个 (b, 1, d)
    gamma1, beta1, alpha1, gamma2, beta2, alpha2 = \
        [p.unsqueeze(1) for p in self.cond(c).chunk(6, dim=-1)]

    # 注意力支路
    h = self.norm1(x) * (1 + gamma1) + beta1     # (b, n, d)
    x = x + alpha1 * self.attn(h)

    # 前馈支路
    h = self.norm2(x) * (1 + gamma2) + beta2     # (b, n, d)
    x = x + alpha2 * self.ff(h)
    return x

其中 MLP 的 final_init 是题面已经给好的:

class MLP(nn.Module):
  def __init__(self, dims, activation=torch.nn.SiLU, final_init: bool = False):
    ...
    if final_init:
      nn.init.zeros_(self.net[-1].weight)
      nn.init.zeros_(self.net[-1].bias)
一行代码验证 adaLN-Zero
layer = DiffusionTransformerLayer(dim=32, heads=4).eval()
x = torch.randn(4, 9, 32); c = torch.randn(4, 32)
print((layer(x, c) - x).abs().max())    # 实测:0.00e+00

初始化后的 block 必须是逐比特精确的恒等映射——不是「近似为 0」,而是严格 0,因为 $\alpha=0$ 让整条残差支路被乘没了。实测最大偏差 0.00e+00。这一条能同时验证:(a)final_init 生效了;(b)残差连接是 x + alpha * f(...) 而不是 alpha * (x + f(...));(c)六组参数的顺序没错(如果你把 $\alpha$ 和 $\gamma$ 的位置搞混,因为初始全 0,这条检查照样通过——所以还要补一条:把 cond 最后一层随机初始化后,改变 $c$ 必须改变输出)。

3.3.3 DiffusionTransformer:位置编码 + 堆叠

题目在问什么。 加位置编码,然后串 depth 层。

数学依据:为什么非要位置编码? 自注意力对 token 顺序是置换等变的:设 $P$ 是任意置换矩阵,则

$$ \text{MHA}(PX)=P\,\text{MHA}(X). $$

证明是一行:$Q,K,V$ 都是逐 token 的线性映射,所以 $(PX)W=P(XW)$;而 $\softmax\big((PQ)(PK)^\top/\sqrt{d_h}\big)(PV)=\softmax(PQK^\top P^\top/\sqrt{d_h})PV=P\,\softmax(QK^\top/\sqrt{d_h})V$。也就是说,如果把 64 个 image token 随机打乱,注意力算出来的结果也只是同样打乱一遍——网络根本不知道哪个 patch 在左上角。对图像来说这是灾难性的。

解决办法是在输入上加一个与位置有关的向量,把置换对称性打破。本 lab 用可学习位置编码:token 数是固定的(图像尺寸固定),所以直接学一张 $n\times d$ 的表就行,不需要正弦编码那样的外推能力。

参考实现。

class DiffusionTransformer(nn.Module):
  def __init__(self, depth: int, n_tokens: int, dim: int, **layer_kwargs):
    super().__init__()
    # 自注意力对 token 顺序置换等变,必须显式注入位置信息
    self.pos_embed = nn.Parameter(torch.randn(n_tokens, dim) * 0.02)   # (n, d)
    self.layers = nn.ModuleList([
        DiffusionTransformerLayer(dim=dim, **layer_kwargs) for _ in range(depth)
    ])

  def forward(self, x: torch.Tensor, c: torch.Tensor) -> torch.Tensor:
    """ x: (b, n, d),  c: (b, d) """
    x = x + self.pos_embed.unsqueeze(0)      # (1, n, d) 广播到 (b, n, d)
    for layer in self.layers:
        x = layer(x, c)
    return x

为什么这样写。

  • * 0.02 的初始化尺度。 位置编码是加到 patch embedding 上的,如果初始幅度和特征本身同量级(标准正态,尺度 1),会在训练早期严重干扰内容信息。0.02 是 GPT/ViT 系列的通用取值,含义是「一开始几乎不影响,让网络自己决定要多少位置信息」。用 torch.randn(n, dim) 不缩放也能训,但前几千步会明显更慢。
  • nn.Parameter 而不是 buffer。 它是要学的。
  • unsqueeze(0)。 (n, d) -> (1, n, d),广播到 batch。实际上 PyTorch 的广播规则会自动在左边补 1,所以直接写 x + self.pos_embed 也对——但显式写出来更不容易读错,题面也特意提醒了「be careful about broadcasting」。
  • nn.ModuleList 而不是 Python list。 用普通 list 装子模块,参数不会被 model.parameters() 收集,也不会随 .to(device) 迁移。症状是:模型大小打印出来只有几百 KB,训练时报「输入在 cuda 但权重在 cpu」,或者干脆损失一动不动。这是 PyTorch 新手最经典的坑之一。
  • 不能用 nn.Sequential。 每层要吃两个输入 (x, c),Sequential 只支持单输入。必须显式写循环。
Q3.3 的常见错误汇总
  • softmax 沿错了维——最隐蔽,不报错,只是变差。记住:dim=-1,沿 key。
  • 缩放用 $\sqrt{d}$ 而不是 $\sqrt{d_h}$——与官方实现对拍会发现误差 $O(1)$。
  • 六组调制参数忘了 unsqueeze(1)——$n\ne d$ 时报错,$n=d$ 时静默算错。DiT 默认 $n=64,d=256$,会报错,算是运气好。
  • LayerNorm 忘了 elementwise_affine=False——不报错,参数冗余,训练略变差。
  • 调制写成 x * gamma + beta(漏了 1 +)——配合零初始化,初始时 $\gamma=0$ 会把整个归一化输出清零,网络退化成只有偏置的常数函数,损失完全不降。这个错误症状很剧烈,反而好查。
  • 残差写成 x = alpha * (x + attn(h))——初始 $\alpha=0$ 会把 $x$ 本身也乘没,整个网络输出恒为 0,梯度全断。恒等映射检查能立刻抓到。
  • 子模块装进 Python list——参数丢失,见上。
  • 位置编码加在了 DiT 之外(比如 Patchifier 里)——功能上等价,不算错;但题面把它放在 DiffusionTransformer 里,因为 Part 5 会用 patch_size=1 复用同一份代码,token 数不同,放在 DiT 里由 n_tokens 参数控制更干净。

3.4 Q3.4 · Depatchifier

题目在问什么。 Patchifier 的逆操作:把 (b, n, d) 变回 (b, c_out, 32, 32)。题面给了四步配方:归一化 → MLP 映到 $f p^2$ 维 → Rearrange 折回空间 → 输出卷积。

数学依据。 没有新数学,但有一个必须理解的对偶关系。Patchifier 把每个 $p\times p\times c$ 的块压成一个 $d$ 维 token;Depatchifier 要把每个 $d$ 维 token 展回一个 $p\times p\times f$ 的块,再由最后一层卷积把 $f$ 个通道压到 $c_{\text{out}}$。所以中间那个 MLP 的输出维度必须是 $f\cdot p^2$——这个数字不是随便定的,它由「一个 token 要负责多少个像素」唯一确定。

参考实现。

class Depatchifier(nn.Module):
  def __init__(self, img_size: int, patch_size: int, dim: int,
               final_dim: int, c_out: int):
    super().__init__()
    assert img_size % patch_size == 0
    h = w = img_size // patch_size
    self.net = nn.Sequential(
        nn.RMSNorm(dim, elementwise_affine=False),                       # (b, n, d)
        MLP([dim, 4 * dim, final_dim * patch_size ** 2]),                # -> (b, n, f*p*p)
        Rearrange('b (h w) (f ph pw) -> b f (h ph) (w pw)',
                  h=h, w=w, f=final_dim, ph=patch_size, pw=patch_size),  # -> (b, f, H, W)
        nn.Conv2d(final_dim, c_out, kernel_size=3, padding=1),           # -> (b, c_out, H, W)
    )

  def forward(self, x: torch.Tensor) -> torch.Tensor:
    return self.net(x)

为什么这样写。

  • 开头的归一化。 DiT 主干最后一层的输出尺度是不受控的(残差累加),先归一化再映射能稳住训练。RMSNorm 和 LayerNorm 都行,前者省掉减均值这一步,稍快。同样要 elementwise_affine=False——理由和 adaLN 里一样,尺度信息紧接着由后面的 MLP 自由决定,不需要再多一组冗余参数。
  • Rearrange 的模式必须与 Patchifier 严格对偶。 逐段对照:
    • b (h w)——把序列维拆回两个空间维,$h$ 慢变、$w$ 快变,与 Patchifier 的 'b d h w -> b (h w) d' 一致。
    • (f ph pw)——把特征维拆成「通道、块内行、块内列」三段,顺序是 $f$ 最慢、$pw$ 最快。
    • b f (h ph) (w pw)——把「块坐标」和「块内坐标」交织成最终的像素坐标:第 $(i\cdot p+a)$ 行来自第 $i$ 个块的第 $a$ 行。
    这个交织顺序是整个函数的灵魂。如果写成 (ph h) 而不是 (h ph),得到的图像会是一个「马赛克重排」——所有块的第 0 行被拼在一起、所有块的第 1 行被拼在一起,看起来像被撕碎重贴。
  • h=h, w=w, f=..., ph=..., pw=... 这些关键字参数不能省。 einops 无法从形状唯一地推断出 (h w) 到底怎么拆($64=8\times8$ 也可以是 $4\times16$),必须显式告诉它。省掉会抛异常,属于「好错误」。
  • 最后一层 Conv2d(kernel_size=3, padding=1)。 为什么不直接让 MLP 输出 $c_{\text{out}}\cdot p^2$ 一步到位?因为 patch 之间是独立展开的,块与块的边界处会有不连续(每个 token 只管自己那 $4\times4$ 个像素)。一层 $3\times3$ 卷积跨过边界做局部平滑,能明显减轻方块状伪影。final_dim=10 就是给这层卷积留的「工作空间」通道数。
Q3.4 的常见错误
  • MLP 输出维度写成 final_dim * patch_size(少一次方)——Rearrange 会因为元素数对不上而抛异常。
  • (h ph) 写成 (ph h)——不报错,输出形状完全正确,但图像被撕碎。只能靠看图发现。 症状很好认:采样结果像是把数字切成 $4\times4$ 的小方块后随机重排。
  • c_out 硬编码成 1——Part 5 要输出 128 通道,会崩。
  • 忘了这是「向量场」不是「图像」——不要在最后加 sigmoid 或 tanh。$u_t^\theta$ 的取值范围是整个 $\R$,任何饱和激活都会把它截断。

3.5 Q3.5 · DiffusionTransformerFlowModel

题目在问什么。 把前四题拼成完整的 $u_t^\theta(x\mid y)$。

数学依据:为什么时间嵌入和类别嵌入是「相加」而不是「拼接」? 两者都被编码到同一个 $d=256$ 维空间,相加之后送进 adaLN 的条件 MLP。相加的合理性在于:条件 MLP 的第一层是线性的,所以

$$ W(e_t+e_y)+b=(We_t)+(We_y)+b, $$

也就是说相加后过线性层,与分别过线性层再相加,完全等价——这正是拼接后过一个(分块的)线性层能做到的事。相加只是把这个分块结构固定下来,省掉一半参数。代价是两路信号共用同一组坐标,可能互相干扰;但实践证明对「时间 + 类别」这种低信息量的条件完全够用。Stable Diffusion 3 和 DiT 原论文用的都是相加。

参考实现。

class DiffusionTransformerFlowModel(ConditionalVectorField):
  def __init__(self, img_size=32, patch_size=8, num_layers=12, c=1,
               dim=256, heads=4, final_dim=10, n_classes=11):
      super().__init__()
      # 0. 时间与类别嵌入
      self.time_embedder = nn.Sequential(
          FourierEncoder(dim),      # (b,) -> (b, dim)
          MLP([dim, dim, dim]),     # -> (b, dim)
      )
      self.y_embedder = nn.Embedding(n_classes, dim)   # 11 = 10 个数字 + 1 个空标签

      # 1. Patchifier
      self.patchifier = Patchifier(img_size=img_size, patch_size=patch_size,
                                   c_in=c, dim=dim)
      # 2. DiT 主干
      n_tokens = (img_size // patch_size) ** 2
      self.dit = DiffusionTransformer(depth=num_layers, n_tokens=n_tokens,
                                      dim=dim, heads=heads)
      # 3. Depatchifier
      self.depatchifier = Depatchifier(img_size=img_size, patch_size=patch_size,
                                       dim=dim, final_dim=final_dim, c_out=c)

  def forward(self, x, t, y):
    """ x: (b, c, h, w),  t: (b,),  y: (b,)  ->  (b, c, h, w) """
    c = self.time_embedder(t) + self.y_embedder(y)   # (b, dim)   <- 相加
    x = self.patchifier(x)                           # (b, n, dim)
    x = self.dit(x, c)                               # (b, n, dim)
    x = self.depatchifier(x)                         # (b, c, h, w)
    return x

为什么这样写。

  • FourierEncoder 后面还要跟一个 MLP。 Fourier 编码只是把标量升成了一组三角函数值,它本身是固定的、无学习能力的特征。后面的 MLP([dim, dim, dim]) 才让网络能把这些特征组合成任务需要的时间表示。少了它,adaLN 的条件 MLP 要直接从原始三角特征学起,收敛会慢很多。
  • n_classes=11 而不是 10。 又一次是 CFG 的那个 $\varnothing$。null_label=10,正好是最后一个索引。
  • 局部变量 c 覆盖了参数名 c(通道数)。 这是官方代码的写法,在 forward 里没问题(参数 c 只在 __init__ 用),但读代码时容易困惑:forward 里的 c 是条件向量 (b, dim),__init__ 里的 c 是通道数。
  • 实际训练用的超参和默认值不同。 默认签名是 patch_size=8, num_layers=12, heads=4,但训练时传的是 patch_size=4, num_layers=8, heads=8。$p=4$ 意味着 $n=64$ 个 token(而不是 16 个),空间分辨率高一倍,代价是注意力的计算量是 $O(n^2)$,涨了 16 倍。对 MNIST 这个尺寸完全吃得消,效果明显更好。
Q3.5 的常见错误
  • n_tokens 算成 img_size // patch_size(忘了平方)——位置编码表只有 8 行而 token 有 64 个,广播时报错。
  • 把 t 和 y 的嵌入拼接后送进 DiT——条件维度变成 2*dim,与 DiffusionTransformerLayer 里 MLP([dim, dim, 6*dim]) 的输入维度对不上。要么改成相加,要么把条件 MLP 的输入维度也改掉。
  • 忘记 c_in=c 传给 Patchifier——Part 3 里 $c=1$ 碰巧和默认值一致,能跑;Part 5 里 $c=128$ 就崩了。
  • final_dim 与 dim 混淆——final_dim=10 是 Depatchifier 内部倒数第二层的通道数,和类别数 10 没有任何关系,纯属数字巧合。

3.6 训练 DiT:跑出来应该长什么样

训练配置:$\eta=0.35$、$\varnothing=10$、20000 步、batch 256、学习率 $4\times10^{-4}$、warmup 500 步。模型 40.324 MiB。单张 RTX 5080 上约 30 分钟。

像素空间 DiT 的训练损失曲线
像素空间 DiT 的损失曲线:从 1.9788 降到 0.1906。这张图验证的是:一条健康的 flow matching 训练曲线应该是「陡降 → 长尾缓降 → 带噪声的平台」。前 1000 步是 warmup 加上模型学会「输出大致的均值」,之后是漫长的细节打磨。注意它不会到 0,也不会像分类任务那样有明显的收敛点——损失的绝对值在这里意义有限,真正的评判标准是采样图。下一节会说明为什么 0.19 这个数字其实包含了很多信息。
DiT 在三个引导强度下的采样结果三联图
本 lab 最值得盯着看的一张图。三个面板分别是引导强度 $w=1.0/3.0/5.0$,每个面板 11 行 × 10 列:前十行分别条件于数字 0–9,最后一行是空标签 $\varnothing$(无条件生成)。逐列对比可以直接看到「保真度–多样性权衡」:$w=1$(无引导,真正从 $\data(\cdot\mid y)$ 采样)数字身份全部正确,但笔迹粗细、倾斜、书写风格差异很大,个别样本潦草到接近误判;$w=3$ 明显规整,笔画干净,同一行里仍有可见的风格变化;$w=5$ 最干净利落,但同一行里的十个样本已经高度相似——多样性显著坍缩。这正是 Lecture 3-B 里「引导 = 从一个被锐化的分布 $\propto p_t(x)p_t(y\mid x)^w$ 中采样」的直接视觉证据:$w$ 越大,概率质量越集中到「最典型的那个 $y$」附近。
读懂最后一行:为什么 $\varnothing$ 那一行的质量不随 $w$ 变化

盯着三个面板的最后一行看,你会发现一件反直觉的事:它们的质量差不多,都是「随机的、有好有坏的数字」,看不出 $w=5$ 比 $w=1$ 更干净。这不是巧合,而是数学上的必然。

CFG 的漂移是 $(1-w)u_t(x\mid\varnothing)+w\,u_t(x\mid y)$。当 $y=\varnothing$ 时,两项里的网络调用完全相同,于是

$$ (1-w)\,u_t(x\mid\varnothing)+w\,u_t(x\mid\varnothing)=u_t(x\mid\varnothing), $$

与 $w$ 无关。三个面板的最后一行用的是同一个向量场,唯一的差别是 visualize_output 在每个 $w$ 的循环里重新采了一批初始噪声 x0——所以它们只是同一分布的三次独立抽样。

这条观察的诊断价值:如果你的最后一行随 $w$ 明显变化(比如 $w=5$ 时变成一堆噪声或全是同一个数字),说明 CFG 的组合公式写错了——最可能是把 unguided_y 写成了别的东西,或者 $w$ 的两个系数搞反成 $(1-w)u_y + w\,u_\varnothing$。

像素空间 DiT 在 step 1000/3000/6000/10000/19000 的采样演进
同一份 checkpoint 序列在固定 $w=3.0$ 下的采样演进(step 1000 / 3000 / 6000 / 10000 / 19000,行布局与上图相同)。这张图验证的是「模型多快学会」,也给出了中途排查的基准线。三个阶段泾渭分明:step 1000 还是不成形的噪声团,只有第二行的「1」因为笔画最简单而勉强可辨——但注意它已经有正确的明暗结构(中心亮、边缘黑),说明训练方向是对的;step 3000 是一道分水岭,十个数字全部正确且可辨认,也就是说「学会画数字」这件事只占了全部训练量的 15%;step 6000 之后身份不再出错,剩下的 17000 步全部花在笔画质量的精修上——笔触逐渐变粗、变匀、更接近真实书写,最后一行(无条件)的样本也从潦草变得工整。这解释了为什么损失曲线在 3000 步后就进入漫长的缓降段:损失的大头早就降完了,后面降的那一点点对应的正是这些肉眼可见的细节改善。
中途看不出效果时怎么排查

把 ckpt_every=1000 打开,训练过程中每 1000 步会存一张采样图。对照上面这张演进图,健康的进展是:

  • 1000 步:不成形的团块,但有清晰的明暗结构(中心亮、边缘黑),个别简单类别(1)已可辨认。
  • 3000 步:十个数字全部正确,笔画尚粗糙。这是最关键的检查点。
  • 6000 步以后:身份稳定,只剩笔画质量在改善。

如果 3000 步时还是纯噪声,立刻停下来,大概率是 $t$ 的形状、softmax 的维度或者 Rearrange 写错了。回去跑第 7 节的检查,不要指望「再多训几步就好了」——上面这条曲线说明,如果 3000 步都学不会画数字,20000 步也学不会。

4. Part 4 · 训练一个变分自编码器

七个 TODO,但难度分布很不均匀:Q4.1–Q4.6 是形状流水账(照配方搭积木,唯一的挑战是把通道数和分辨率对上),真正需要动脑的是 Q4.7 的损失函数。

先明确目标。VAE 由两半组成:

$$ \text{编码器}\ q_\phi(z\mid x)=\N\big(\mu_\phi(x),\ \sigma_\phi^2 I_k\big),\qquad \text{解码器}\ p_\theta(x\mid z)=\N\big(\mu_\theta(z),\ \sigma_\theta^2 I_d\big), $$

本 lab 的具体形状是 $x\in\R^{1\times32\times32}$($d=1024$)、$z\in\R^{128\times4\times4}$($k=2048$)。注意隐空间的元素数比图像还多——这在 latent diffusion 里很常见,压缩发生在空间维($32\times32\to4\times4$,降 64 倍),而通道维膨胀了 128 倍。真正的收益是让扩散模型面对 $4\times4$ 的「图像」,注意力的 token 数从 64 降到 16。

两个 $\sigma$ 都被参数化成可学习的共享标量 nn.Parameter(torch.zeros(())),存的是 $\log\sigma^2$ 而不是 $\sigma^2$——这样可以无约束地取任意实数,再用 exp 保证正性,避免了显式的正性约束。

4.1 Q4.1 · ResidualBlock

题目在问什么。 实现 $x\mapsto x+\text{Conv}_{1\times1}\big(\text{SiLU}(\text{Conv}_{3\times3}(\text{Norm}(x)))\big)$。

参考实现。

class ResidualBlock(nn.Module):
  def __init__(self, channels: int, act: nn.Module = nn.SiLU):
    super().__init__()
    self.norm  = nn.GroupNorm(1, channels)                        # 等价于图像上的 LayerNorm
    self.conv1 = nn.Conv2d(channels, channels, kernel_size=3, padding=1)
    self.act   = act()
    self.conv2 = nn.Conv2d(channels, channels, kernel_size=1)

    # 第二个卷积零初始化 => 初始时整个块是恒等映射
    nn.init.zeros_(self.conv2.weight)
    nn.init.zeros_(self.conv2.bias)

  def forward(self, x: torch.Tensor):
    x_skip = x                    # (b, c, h, w)
    x = self.norm(x)
    x = self.conv1(x)
    x = self.act(x)
    x = self.conv2(x)
    return x_skip + x             # (b, c, h, w)

为什么这样写。

  • nn.GroupNorm(1, channels) 为什么等价于 LayerNorm? GroupNorm 把 $C$ 个通道分成 $G$ 组,每组内跨「通道 + 空间」计算均值方差。取 $G=1$ 就是把整张特征图(所有通道、所有像素)当成一个整体归一化——这正是 LayerNorm 作用在 (C, H, W) 上的定义。相比 BatchNorm 的好处是不依赖 batch 统计量,训练和推理行为完全一致,batch size 小的时候也稳。nn.LayerNorm 需要显式给出 [C, H, W],而 VAE 里同一个 block 会用在不同分辨率上,用 GroupNorm 更省事。
  • padding=1 配 kernel_size=3。 保持空间尺寸不变($3\times3$ 卷积 + padding 1 是「same」卷积)。残差连接要求输入输出同形,少了 padding 会掉两个像素而报错。
  • 第二个卷积是 $1\times1$。 它的作用是「通道混合」而非「空间聚合」——空间感受野已经由 conv1 提供了。这也是「$3\times3$ 提特征、$1\times1$ 做投影」的经典瓶颈结构简化版。
  • 零初始化 conv2。 与 adaLN-Zero 同一个思想:初始时残差支路输出恒为 0,整个 block 是恒等映射。VAE 有 $4\times2=8$ 个残差块加上下采样,深度不算浅,这一手能明显稳住前几百步。
一行验证
rb = ResidualBlock(channels=8).eval()
x = torch.randn(2, 8, 16, 16)
print((rb(x) - x).abs().max())     # 实测:0.00e+00

和 adaLN-Zero 一样,这里也必须是严格 0。如果不是 0,检查:(a)nn.init.zeros_ 是不是只初始化了 weight 忘了 bias;(b)残差是不是写成了 x_skip + x 以外的形式;(c)零初始化是不是被后续的某个 reset_parameters() 覆盖了。

4.2 Q4.2 · AttnBlock

题目在问什么。 把空间位置当作 token,在特征图上做一次自注意力 + 前馈,各带残差。

数学依据:卷积的局限与注意力的补位。 卷积的感受野是局部且随深度线性增长的;即使堆到编码器最后一层($4\times4$),每个位置也只能看到有限邻域。而 VAE 需要把整张图的信息压进 $z$,全局交互是必要的。自注意力提供的正是「任意两个位置一步直连」的能力,代价是 $O((HW)^2)$ 的复杂度——这也是为什么 Stable Diffusion 的 VAE 只在最低分辨率层放注意力。

参考实现。

class AttnBlock(nn.Module):
  def __init__(self, channels: int):
    super().__init__()
    self.reshape1 = Rearrange('b c h w -> b (h w) c')
    self.norm1 = nn.LayerNorm(channels)
    self.attn  = MHA(channels, heads=1)     # heads=1:保证任意通道数都能整除
    self.norm2 = nn.LayerNorm(channels)
    self.ff    = MLP([channels, 4 * channels, channels])

  def forward(self, x: torch.Tensor):
    b, c, h, w = x.shape
    x = self.reshape1(x)                    # (b, h*w, c)

    x_skip = x                              # 注意力 + 残差
    x = self.attn(self.norm1(x))
    x = x_skip + x

    x_skip = x                              # 前馈 + 残差
    x = self.ff(self.norm2(x))
    x = x_skip + x

    return rearrange(x, 'b (h w) c -> b c h w', h=h, w=w)

为什么这样写。

  • 「像素即 token、通道即特征」。 'b c h w -> b (h w) c' 把 $H\times W$ 个空间位置变成序列,通道数当作 token 维度。于是 Q3.3 写的 MHA 一个字都不用改就能复用——这是本 lab 一处很好的架构设计,值得注意。
  • heads=1 不是偷懒。 MHA 里有 assert dim % heads == 0。VAE 的通道数依次是 16 / 32 / 64 / 128,如果写死 heads=4,16 通道那层每个头只有 4 维,表达力堪忧;写 heads=8 则 16 通道刚好整除但更窄。heads=1 对所有这些通道数都成立,且在 $4\times4=16$ 个 token 的规模上多头并没有优势。
  • 用 x.shape 解包出 h, w 再传给最后的 rearrange。 必须的:einops 没法从 (b, hw, c) 反推 $h$ 和 $w$。
  • Pre-norm 而非 post-norm。 归一化放在残差支路内部(x + f(norm(x))),而不是在残差之后(norm(x + f(x)))。前者的主干路径是一条无阻碍的恒等通路,梯度传播更好,是现代 transformer 的标准做法。
Q4.2 的常见错误
  • 'b c h w -> b c (h w)'——把序列维放到了最后,MHA 会把空间位置当成特征维、通道当成 token,语义完全错位。形状上 (b, c, hw) 也能过 nn.Linear(c, 3c)?不能——最后一维是 $hw$ 不是 $c$,会报错。算是好错误。
  • 返回时忘了变回 b c h w——下一个模块是卷积,会因为维度数不对而报错。
  • nn.LayerNorm(channels) 写成 nn.LayerNorm([h, w])——LayerNorm 归一化的是最后若干维,这里 token 张量的最后一维是通道,所以参数就是 channels。

4.3 Q4.3 · EncoderBlock 与 Q4.5 · DecoderBlock

这两题结构完全镜像,一起讲。题目在问什么:两个残差块 + 一个注意力块 + (可选的)下采样 / 上采样。

class EncoderBlock(nn.Module):
  def __init__(self, in_channels: int, downsample_channels: Optional[int] = None):
    super().__init__()
    self.res1 = ResidualBlock(in_channels)
    self.res2 = ResidualBlock(in_channels)
    self.attn = AttnBlock(in_channels)
    if downsample_channels is not None:
      # stride=2 的卷积同时完成「降采样」和「换通道数」
      self.downsample = nn.Conv2d(in_channels, downsample_channels,
                                  kernel_size=3, padding=1, stride=2)
    else:
      self.downsample = None

  def forward(self, x):
    x = self.res1(x); x = self.res2(x); x = self.attn(x)
    if self.downsample is not None:
      x = self.downsample(x)
    return x


class DecoderBlock(nn.Module):
  def __init__(self, in_channels: int, upsample_channels: Optional[int] = None):
    super().__init__()
    self.res1 = ResidualBlock(in_channels)
    self.res2 = ResidualBlock(in_channels)
    self.attn = AttnBlock(in_channels)
    if upsample_channels is not None:
      self.upsample = nn.Sequential(
        nn.Upsample(scale_factor=2, mode='nearest'),
        nn.Conv2d(in_channels, upsample_channels,
                  kernel_size=3, padding=1, stride=1),
      )
    else:
      self.upsample = None

  def forward(self, x):
    x = self.res1(x); x = self.res2(x); x = self.attn(x)
    if self.upsample is not None:
      x = self.upsample(x)
    return x

为什么这样写。

  • 下采样用 stride=2 的卷积,而不是池化。 池化是固定的(不可学习)降采样;stride=2 卷积让网络自己学「怎么降」,同时顺手把通道数从 in_channels 换成 downsample_channels,一步两用。kernel_size=3, padding=1, stride=2 作用在偶数尺寸上恰好减半:$\lfloor(32+2-3)/2\rfloor+1=16$。
  • Optional[int] = None 的语义是「最后一块不改变分辨率」。 见下一题的形状表。
  • 注意力放在残差块之后、采样之前。 这样注意力总是在「还没降维」的分辨率上做全局混合,信息损失更小。
为什么上采样用「最近邻 + 卷积」而不是 ConvTranspose2d?

答案是棋盘格伪影(checkerboard artifacts)。

转置卷积的工作方式是:把输入的每个像素乘以整个卷积核再摊到输出上,相邻输入的贡献区域重叠。当 kernel size 不能被 stride 整除时(最常见的 $k=3,s=2$),输出的不同位置接收到的贡献次数不同——比如某些位置被 2 个输入覆盖、相邻位置只被 1 个覆盖。这个不均匀性是周期性的,于是在输出上叠出一层规则的明暗方格。它在训练中很难被完全消掉,因为它是架构本身的性质而不是权重的问题。

「最近邻上采样 + 普通卷积」把两件事解耦:nn.Upsample 负责均匀地把每个像素复制成 $2\times2$(每个输出位置的来源数完全相同,不存在不均匀重叠),随后的 stride=1 卷积负责平滑和特征变换。代价是多一次内存搬运,收益是干净的输出。Stable Diffusion 的 VAE 解码器用的就是这个方案。

怎么验证你踩到了这个坑:把重构图放大看,如果看到规则的、周期为 2 或 4 像素的网格状纹理,而且它不随训练变淡,那就是棋盘格。

4.4 Q4.4 · Encoder

题目在问什么。 初始卷积 → 每个 hidden channel 一个 EncoderBlock(除最后一个外都下采样)→ 归一化 + $1\times1$ 卷积输出 z_mean,外加一个标量 logvar。

形状流水账(hidden_channels = [16, 32, 64, 128])。这是本题唯一真正要想清楚的东西:

步骤操作输出形状
输入——(b, 1, 32, 32)
init_conv$3\times3$,$1\to16$(b, 16, 32, 32)
blocks[0]EncoderBlock(16, 32),含下采样(b, 32, 16, 16)
blocks[1]EncoderBlock(32, 64),含下采样(b, 64, 8, 8)
blocks[2]EncoderBlock(64, 128),含下采样(b, 128, 4, 4)
blocks[3]EncoderBlock(128, None),无下采样(b, 128, 4, 4)
z_meanGroupNorm + $1\times1$ 卷积(b, 128, 4, 4)

参考实现。

class Encoder(nn.Module):
  def __init__(self, in_channels: int, hidden_channels: list[int]):
    super().__init__()
    self.init_conv = nn.Conv2d(in_channels, hidden_channels[0],
                               kernel_size=3, padding=1, stride=1)

    # 关键的一行:把 [16,32,64,128] 配对成 (16,32),(32,64),(64,128),(128,None)
    ch_in  = hidden_channels                 # [16, 32, 64, 128]
    ch_out = hidden_channels[1:] + [None]    # [32, 64, 128, None]
    self.blocks = nn.ModuleList([
        EncoderBlock(i, o) for i, o in zip(ch_in, ch_out)
    ])

    z_dim = hidden_channels[-1]
    self.z_mean = nn.Sequential(
      nn.GroupNorm(1, z_dim),
      nn.Conv2d(z_dim, z_dim, kernel_size=1, stride=1, padding=0),
    )
    self.logvar = nn.Parameter(torch.zeros(()))    # 标量!形状是 ()

  def forward(self, x: torch.Tensor):
    x = self.init_conv(x)
    for block in self.blocks:
      x = block(x)
    return self.z_mean(x), self.logvar

为什么这样写。

  • ch_out = hidden_channels[1:] + [None] 是整题的题眼。 它把「$n$ 个通道数」变成「$n$ 个 (输入, 输出) 对」,最后一个的输出是 None,正好触发 EncoderBlock 里「不下采样」的分支。四个 block、三次下采样,$32\to16\to8\to4$。
  • 为什么最后一块不下采样? 因为 $4\times4$ 就是目标分辨率了。多一个 block 是为了在最终分辨率上再做两轮残差 + 一次全局注意力,把信息充分整合后再输出 $\mu_\phi$。
  • self.logvar = nn.Parameter(torch.zeros(())) —— 注意是 torch.zeros(()) 不是 torch.zeros(1)。 前者形状是 torch.Size([])(真正的 0 维标量),后者是 torch.Size([1])。两者在广播时行为几乎一样,但测试里会检查 z_logvar.shape == torch.Size([])。语义上也确实是「一个数」:所有样本、所有维度共享同一个后验方差。
  • 共享标量方差意味着什么? 标准 VAE 里 $\sigma_\phi(x)$ 是输入相关的(每个 $x$ 有自己的不确定度)。这里退化成一个全局常数,等于说「编码器对所有输入的置信度一样」。这是 latent diffusion 里的常见简化(Lecture 4 §7.7 讨论过),好处是训练稳定得多,代价是模型在概率意义上更弱——但反正我们只是要一个「被轻度正则化过的自编码器」。
Q4.4 的常见错误
  • 用 zip(hidden_channels[:-1], hidden_channels[1:])——只造出 3 个 block,最后那个「不下采样的整合块」丢了。形状仍然对($4\times4$),但少了两个残差块和一次注意力,重构质量下降。
  • 把 blocks 放进 Python list——参数不被注册,见 Q3.3。
  • logvar 写成 nn.Parameter(torch.zeros(z_dim))——变成逐通道方差,不是题面要的标量。不会报错,但和 compute_loss 的 .mean() 组合起来含义会变。
  • init_conv 用了 stride=2——多降一次,最后得到 $2\times2$,Part 5 的 img_size=4 就对不上了。

4.5 Q4.6 · Decoder

题目在问什么。 Encoder 的镜像:每个 hidden channel 一个 DecoderBlock(除最后一个外都上采样)→ 归一化 + $1\times1$ 卷积输出 x_mean + 标量 logvar。注意没有 init_conv——输入已经是 128 通道了。

形状流水账。关键在于 VAE.__init__ 里传的是 list(reversed(hidden_channels)),即 [128, 64, 32, 16]:

步骤操作输出形状
输入——(b, 128, 4, 4)
blocks[0]DecoderBlock(128, 64),含上采样(b, 64, 8, 8)
blocks[1]DecoderBlock(64, 32),含上采样(b, 32, 16, 16)
blocks[2]DecoderBlock(32, 16),含上采样(b, 16, 32, 32)
blocks[3]DecoderBlock(16, None),无上采样(b, 16, 32, 32)
x_meanGroupNorm + $1\times1$ 卷积,$16\to1$(b, 1, 32, 32)
class Decoder(nn.Module):
  def __init__(self, out_channels: int, hidden_channels: list[int]):
    super().__init__()
    ch_in  = hidden_channels                 # [128, 64, 32, 16]
    ch_out = hidden_channels[1:] + [None]    # [64, 32, 16, None]
    self.blocks = nn.ModuleList([
        DecoderBlock(i, o) for i, o in zip(ch_in, ch_out)
    ])

    x_dim = hidden_channels[-1]              # 16
    self.x_mean = nn.Sequential(
      nn.GroupNorm(1, x_dim),
      nn.Conv2d(x_dim, out_channels, kernel_size=1, stride=1, padding=0),
    )
    self.logvar = nn.Parameter(torch.zeros(()))

  def forward(self, x: torch.Tensor):
    for block in self.blocks:
      x = block(x)
    return self.x_mean(x), self.logvar

为什么这样写。 结构与 Encoder 逐行对称,唯一要小心的是 reversed 是在 VAE.__init__ 里做的,不是在 Decoder 里:

class VAE(nn.Module):
  def __init__(self, data_channels: int, hidden_channels: list[int], beta: float = 0.1):
    super().__init__()
    self.beta = beta
    self._encoder = Encoder(data_channels, hidden_channels)                 # [16,32,64,128]
    self._decoder = Decoder(data_channels, list(reversed(hidden_channels))) # [128,64,32,16]

所以 Decoder 内部的写法和 Encoder 完全一样,不要在 Decoder 里再 reverse 一次——那会让通道数变成 $16\to32\to64\to128$,但输入是 128 通道,第一层就崩。

Q4.6 的常见错误
  • 在 Decoder 里又 reverse 一次——见上,立刻报错。
  • 误以为 Decoder 也需要 init_conv——不需要,多一层不会崩但破坏对称性。
  • 最后的 $1\times1$ 卷积输出通道写成 x_dim 而不是 out_channels——输出 16 通道的「图像」,与 x_true 的 1 通道做 MSE 时会广播成 (b, 16, 32, 32),不报错但完全错误。这是本题最阴的一个坑,靠形状测试才能抓到。
  • 忘了 Decoder 也有自己的 logvar——它和 Encoder 的是两个独立参数:前者是解码器的观测噪声 $\sigma_\theta$,后者是编码器的后验方差 $\sigma_\phi$,含义完全不同。

4.6 Q4.7 · VAE.compute_loss —— 本 Part 唯一需要动脑的一题

题目在问什么。 实现讲义那条完整的 $\beta$-VAE 损失(Lecture 4 §7.4 末尾给出,notebook 里标注为式 (83)/(85),不同版本编号略有出入):

$$ \mathcal{L}_{\text{VAE}}(\phi,\theta)= \underbrace{\frac{\norm{x-\mu_\theta(z)}^2}{2\sigma_\theta^2}+\frac{d}{2}\log\sigma_\theta^2}_{\text{重构 + 解码器置信度}} \;+\;\beta\underbrace{\frac{1}{2}\Big[\mathcal{K}\big(\sigma_\phi^2\big)+\norm{\mu_\phi(x)}^2\Big]}_{\text{KL 到先验}}, \qquad \mathcal{K}(\alpha)=\sum_i\big(\alpha_i-\log\alpha_i-1\big). $$

数学依据:两项各自从哪来。

重构项:高斯似然的负对数

解码器是 $p_\theta(x\mid z)=\N(\mu_\theta(z),\sigma_\theta^2 I_d)$,其负对数似然为

$$ -\log p_\theta(x\mid z)=\frac{\norm{x-\mu_\theta(z)}^2}{2\sigma_\theta^2}+\frac{d}{2}\log\sigma_\theta^2+\frac{d}{2}\log(2\pi). $$

最后一项是常数,扔掉。剩下两项的关系值得玩味:第一项希望 $\sigma_\theta$ 越大越好(把重构误差除小),第二项希望 $\sigma_\theta$ 越小越好。对 $\sigma_\theta^2$ 求导置零,最优解是 $\sigma_\theta^{2\star}=\frac{1}{d}\norm{x-\mu_\theta}^2$——即解码器方差会自动收敛到当前的均方重构误差。所以 $\log\sigma_\theta^2$ 这一项不是可有可无的正则,它让模型自动标定「我对重构有多确信」,等价于把重构损失的权重自适应地调成 $1/\text{MSE}$。

KL 项:两个对角高斯之间的闭式解

取先验 $p_{\text{prior}}=\N(0,I_k)$,后验 $q_\phi=\N(\mu,\diag(\sigma^2))$,则逐坐标可分解:

$$ \KL\big(\N(\mu_i,\sigma_i^2)\,\|\,\N(0,1)\big)=\frac{1}{2}\Big(\sigma_i^2+\mu_i^2-\log\sigma_i^2-1\Big). $$

求和得到

$$ \KL\big(q_\phi\,\|\,\N(0,I_k)\big)=\frac{1}{2}\Big[\underbrace{\sum_i(\sigma_i^2-\log\sigma_i^2-1)}_{=\,\mathcal{K}(\sigma^2)\ \text{:把方差推向 1}}+\underbrace{\norm{\mu}^2}_{\text{把均值推向 0}}\Big]. $$

两项的极小点分别在 $\sigma_i=1$ 和 $\mu_i=0$——KL 项做的事就是把编码分布往标准正态上按。注意 $\mathcal{K}(\alpha)=\alpha-\log\alpha-1\ge0$ 且仅在 $\alpha=1$ 取等(因为 $\log\alpha\le\alpha-1$),所以整个 KL 恒非负,这是一条可以直接写成断言的性质。

参考实现。

def compute_loss(self, z_mean, z_logvar, x_mean, x_logvar, x_true):
    eps = 1e-6

    # KL 项:把 (mu, sigma^2) 推向 (0, 1)
    kl_loss = self.beta * (z_mean.pow(2) + torch.exp(z_logvar) - z_logvar - 1).mean()

    # 重构项:高斯负对数似然(去掉常数)
    mse_term        = (x_true - x_mean).pow(2) / (torch.exp(x_logvar) + eps)
    confidence_term = x_logvar
    recon_loss      = (mse_term + confidence_term).mean()

    return kl_loss + recon_loss

以及 forward 里的重参数化:

def forward(self, x: torch.Tensor):
    z_mean, z_logvar = self.encode(x)                                # (b,128,4,4), ()
    z = z_mean + torch.exp(0.5 * z_logvar) * torch.randn_like(z_mean)  # 重参数化
    x_mean, x_logvar = self.decode(z)                                # (b,1,32,32), ()
    return z_mean, z_logvar, x_mean, x_logvar
官方实现与讲义公式的一处真实差异:.mean() 还是 .sum()?

讲义式子里两处都是按维度求和($\norm{\cdot}^2$ 和 $\sum_i$),而代码里两处都用了 .mean()——对张量的所有元素求平均。这不是笔误,但它确实改变了 $\beta$ 的含义,值得算清楚。

记 $d=1\times32\times32=1024$(图像维度)、$k=128\times4\times4=2048$(隐变量维度),$R$ 为讲义的重构项、$S=\mathcal{K}(\sigma_\phi^2)+\norm{\mu_\phi}^2$。逐项对照:

$$ \begin{aligned} \text{代码重构项}&=\frac{1}{d}\frac{\norm{x-\mu_\theta}^2}{\sigma_\theta^2}+\log\sigma_\theta^2=\frac{2}{d}\left[\frac{\norm{x-\mu_\theta}^2}{2\sigma_\theta^2}+\frac{d}{2}\log\sigma_\theta^2\right]=\frac{2}{d}R,\\[4pt] \text{代码 KL 项}&=\frac{\beta_{\text{代码}}}{k}S. \end{aligned} $$

所以 $\mathcal{L}_{\text{代码}}=\frac{2}{d}R+\frac{\beta_{\text{代码}}}{k}S$。整体乘以常数 $\frac{d}{2}$(不改变极小点)得到 $R+\frac{d\,\beta_{\text{代码}}}{2k}S$,与讲义的 $R+\frac{\beta_{\text{讲义}}}{2}S$ 对照,立刻读出换算关系:

$$ \boxed{\ \beta_{\text{讲义}}=\beta_{\text{代码}}\cdot\frac{d}{k}=10\times\frac{1024}{2048}=5.\ } $$

为什么要用 .mean()?因为这样损失的数值与分辨率无关:换成 $64\times64$ 的图,$d$ 变成 4 倍,用 .sum() 的话损失也变 4 倍,学习率得跟着调;用 .mean() 则不用动。代价是 $\beta$ 的数值不能直接和论文/讲义里的对照——必须先做上面这个换算。这是读 VAE 代码时最容易被绊住的地方之一,看到别人的 $\beta$ 与你的差几个数量级,八成就是归一化约定不同。

另外注意代码把重构项整体乘了 2(写成 $\frac{\norm{\cdot}^2}{\sigma^2}+\log\sigma^2$ 而非 $\frac{\norm{\cdot}^2}{2\sigma^2}+\frac{d}{2}\log\sigma^2$ 的按维度平均)。上面的换算已经把这个因子 2 算进去了——它相当于把有效 $\beta$ 减半,属于常数缩放,不改变最优点的位置。

那么 $\beta=10$ 和讲义说的「现代自编码器里 $\beta\ll1$」矛盾吗?

换算之后 $\beta_{\text{讲义}}=5$,确实仍然远大于 1。所以这不是单位问题,而是一个刻意的教学选择,理由有二:

  1. MNIST 太好重构了。 生产级 latent VAE(Stable Diffusion 那种)面对的是高分辨率自然图,重构项本身很难压下去,稍大的 $\beta$ 就会毁掉细节,所以必须 $\beta\ll1$。MNIST 是二值化的笔画,$16\to128$ 通道的编码器绰绰有余,可以用很强的 KL 换一个更规整的隐空间。
  2. 下游是扩散模型,隐空间的「形状」比重构精度更重要。 $\beta$ 大 ⇒ 隐分布更接近 $\N(0,I)$ ⇒ Part 5 的 flow matching 起点 $p_{\text{init}}=\N(0,I)$ 与终点分布的尺度天然匹配,不需要额外的全局缩放常数(Lecture 4 §7.7 提到实际系统里通常要乘一个这样的常数)。

但这个选择有代价,而且代价在 Part 5 里会以一种非常吓人的方式显现出来——第 6 节专门讲。

重参数化技巧:为什么不能直接对采样求导

我们要算的是 $\grad_\phi\,\E_{z\sim q_\phi(\cdot\mid x)}[f(z)]$。问题在于期望的分布本身依赖 $\phi$:

$$ \grad_\phi\int q_\phi(z\mid x)f(z)\ud z=\int \grad_\phi q_\phi(z\mid x)\,f(z)\ud z, $$

右边不是任何一个关于 $q_\phi$ 的期望,没法用「采样 → 求导」直接估计。更直白地说:torch.normal(mu, sigma) 这个操作在计算图里是一个断点,反向传播走不过去。

重参数化把随机性挪到与参数无关的地方:

$$ z=\mu_\phi(x)+\sigma_\phi\odot\epsilon,\qquad \epsilon\sim\N(0,I_k)\ \text{与}\ \phi\ \text{无关}. $$

于是 $\E_{z\sim q_\phi}[f(z)]=\E_{\epsilon\sim\N(0,I)}[f(\mu_\phi+\sigma_\phi\odot\epsilon)]$,期望的分布不再含 $\phi$,梯度可以直接换进去:

$$ \grad_\phi\,\E_\epsilon\big[f(\mu_\phi+\sigma_\phi\odot\epsilon)\big]=\E_\epsilon\big[\grad_\phi f(\mu_\phi+\sigma_\phi\odot\epsilon)\big]. $$

代码里就是那一行 z = z_mean + torch.exp(0.5 * z_logvar) * torch.randn_like(z_mean):randn_like 产生的 $\epsilon$ 是常数(不需要梯度),$\mu$ 和 $\sigma$ 都在计算图里。注意 0.5 * 不能少——存的是 $\log\sigma^2$,所以 $\sigma=\exp(\frac{1}{2}\log\sigma^2)$。漏掉这个 0.5 会让噪声方差变成 $\sigma^4$,训练初期 $\sigma\approx1$ 时看不出区别,训练后期会悄悄跑偏。

Q4.7 的常见错误
  • KL 项写成 z_logvar - torch.exp(z_logvar) ...(符号弄反)——KL 会变成负数并被无限最小化,编码器的方差爆炸。症状:损失一路跌到 $-\infty$。用「$q$ = 先验时 KL 必须恰好为 0」这条检查:μ=0, logvar=0 代进去,$0+1-0-1=0$。
  • 忘了 exp,直接用 z_logvar 当方差——不报错,但 KL 的最优点从 $\sigma^2=1$ 挪到了别处。
  • 重构项漏了 x_logvar 这一项——只剩 $\frac{\text{MSE}}{\sigma_\theta^2}$,模型会把 $\sigma_\theta\to\infty$ 让损失趋于 0,重构彻底摆烂。这个错误很致命且症状明确:损失快速趋近 0(而不是变负),重构图是一片灰。
  • 分母没加 eps——$\sigma_\theta^2$ 在训练后期可能变得很小,除法溢出成 inf。1e-6 是廉价保险。
  • 把 beta 乘到重构项上而不是 KL 项上——等价于用 $1/\beta$ 做 KL 权重,方向完全反了,会得到一个几乎不受约束的隐空间。
  • 对着「损失是负数」慌了——见下一节,这是正常的。

4.7 训练 VAE:跑出来应该长什么样

配置:hidden_channels=[16,32,64,128]、$\beta=10$、5000 步、batch 64、学习率 $10^{-3}$。模型 6.150 MiB,在 RTX 5080 上几分钟就跑完。

VAE 训练损失曲线,从 4.3 降到 -0.63
VAE 损失从 4.3153 降到 −0.6341。这张图验证的是:形状是典型的「悬崖 + 长缓坡」——前 200 步内从 4.3 直坠到 1 附近(模型迅速学会输出接近数据均值的东西 + $\sigma_\theta$ 快速自适应),之后是几千步的缓慢下降,对应重构细节的逐步改善。损失变成负数完全正常,不是 bug。原因是重构项里的 $\log\sigma_\theta^2$:当解码器足够准、$\sigma_\theta<1$ 时这一项为负,而它可以无下界地往负走(连续分布的对数密度本来就可以是任意大的正数,负对数似然自然可以任意负)。判断训练是否健康,看的是曲线形状和重构图,不是损失的符号。
VAE 重构质量在 step 250/1000/2000/3500/4750 的演进
VAE 重构质量的演进(step 250 / 1000 / 2000 / 3500 / 4750,每格上行是输入、下行是重构)。这张图把上面那条损失曲线翻译成了看得见的东西。对照着看:step 250(对应损失曲线刚从悬崖落到 1 附近)重构已经有正确的位置和亮度,但极度模糊,而且认错了数字——第 3、7 格的输入是「1」和「8」,重构出来更像别的东西,说明此时隐变量还没编码进足够的身份信息;step 1000 起身份基本全对,只是笔画偏胖、细节被抹平;step 2000 → 4750 是缓慢的锐化过程,到 4750 时重构与输入几乎无法区分,连「9」的开口大小、「5」的转折角度这类个体特征都保留了下来。这条演进曲线正好解释了 $\log\sigma_\theta^2$ 项的作用:重构从模糊到锐利的过程,同时也是解码器方差 $\sigma_\theta^2$ 自动收缩的过程(它会收敛到当前的均方重构误差),而 $\log\sigma_\theta^2$ 变负正是总损失最终落到 −0.63 的原因。
VAE 隐空间线性插值:从数字 1 平滑过渡到数字 8
隐空间线性插值:编码两张真实图得到 $z_1,z_2$,沿 $z_\lambda=(1-\lambda)z_1+\lambda z_2$ 取 10 个点解码。这张图验证的是 KL 正则真的起作用了。如果隐空间只是一个普通自编码器的隐空间,编码点会散落在一些孤立的区域,两点连线上的中间点落在训练时从未被覆盖的地方,解码出来是噪声或糊团。而这里从左边的「1」到右边的「8」,中间是一系列连续、始终像笔画的形状——第 3–5 帧可以看到一撇正在长出圈来。这说明 $\beta=10$ 的 KL 项成功地把编码分布压成了一片连通、稠密的区域。这正是 latent diffusion 需要的:扩散模型会在这个空间里到处采样,那里必须处处「有意义」。
如果你的插值图中间是糊的

调大 $\beta$。$\beta$ 太小时 KL 约束弱,隐空间会变成一堆孤岛,插值路径穿过「真空区」。反过来 $\beta$ 太大会触发后验塌缩(posterior collapse)——编码器索性无视 $x$,输出 $q_\phi\approx\N(0,I)$,KL 归零而信息全丢,症状是重构图与输入完全无关、所有输入都解码成同一个模糊数字。$\beta=10$(讲义约定下的 5)在这个规模上是个好平衡点。

5. Part 5 · 隐空间扩散

所有零件都造好了,这一 Part 只剩一件事:把 Part 3 的 DiT 搬进 Part 4 的隐空间。只有一个 TODO。

5.1 Q5.1 · LatentCFGTrainer.get_train_loss

题目在问什么。 与 Q2.2 的 CFGTrainer 唯一的区别:训练数据不是像素 $x$,而是 VAE 编码后的隐变量 $z$。而且 VAE 是冻结的(先单独训好,这一步不再更新它的参数)。

数学依据。 Latent diffusion 的整个想法可以写成一个分布的复合:

$$ p_{\text{model}}(x)=\int p_\theta(x\mid z)\,p_{\text{latent}}^{\text{flow}}(z)\ud z, $$

其中 $p_\theta(x\mid z)$ 是冻结的解码器,$p_{\text{latent}}^{\text{flow}}$ 是我们要训练的流模型在 $t=1$ 时的分布。训练目标是让 $p_{\text{latent}}^{\text{flow}}$ 匹配编码器诱导的聚合后验(aggregate posterior)

$$ p_{\text{latent}}(z)=\E_{x\sim\data}\big[q_\phi(z\mid x)\big]. $$

而「从 $p_{\text{latent}}$ 采样」的操作,恰好就是「采一张真实图 → 编码 → 重参数化采样」。所以 flow matching 的公式一个字都不用改,只需要把 p_data.sample() 换成这个两步过程。

参考实现。

def get_train_loss(self, batch_size: int) -> torch.Tensor:
    # Step 1: 从 MNIST 采样并编码到隐空间
    #         VAE 是冻结的,整个编码过程放在 no_grad 里
    with torch.no_grad():
      x, y = self.mnist.sample(batch_size)                # x: (b,1,32,32), y: (b,)
      z_mean, z_logvar = self.vae.encode(x)               # (b,128,4,4), ()
      # 用完整的后验采样(而非只取均值),让扩散模型见到编码器真实的隐分布
      z = z_mean + torch.exp(0.5 * z_logvar) * torch.randn_like(z_mean)

    # Step 2: 以概率 eta 把标签置为空标签
    xi = torch.rand(y.shape[0], device=y.device)
    y[xi < self.eta] = self.null_label

    # Step 3: 采 t 和 z_t
    t = torch.rand(batch_size, device=z.device).to(z) * (1 - self.eps)   # (b,)
    zt = self.path.sample_conditional_path(z, t)                         # (b,128,4,4)

    # Step 4: 回归
    ut_theta = self.model(zt, t, y)                          # (b,128,4,4)
    ut_ref   = self.path.conditional_vector_field(zt, z, t)  # (b,128,4,4)
    return torch.square(ut_theta - ut_ref).mean()

为什么这样写。

  • torch.no_grad() 是必须的,不只是省显存。 如果不加,编码器的计算图会被保留,loss.backward() 会把梯度一路传回 VAE 的参数。虽然优化器里只注册了 DiT 的参数(self.opt = AdamW(self.model.parameters())),VAE 不会真被更新,但你会白白多花显存和时间。更糟的情况是:如果你的 trainer 写法不小心把 VAE 也塞进了优化器,VAE 会边训边动,扩散模型追着一个移动的目标学,训练直接发散。加 no_grad 是最稳的做法。
  • 为什么采样而不是只用 z_mean? 因为扩散模型要匹配的是聚合后验 $p_{\text{latent}}$,那是带噪声的。只用均值会让扩散模型学到一个比真实隐分布「更瘦」的分布,采样时生成的 $z$ 落在解码器没见过的低方差区域,重构质量下降。这一步的噪声也起到轻微的数据增广作用。
  • path 的构造变了。 注意训练脚本里 GaussianConditionalProbabilityPath(p_data=None, p_simple_shape=[128, 4, 4], ...)——p_data 直接传 None,因为 trainer 不再通过 path 取数据(那条路被上面的两步采样取代了)。path 现在只负责三件事:p_simple(隐空间形状的高斯)、sample_conditional_path、conditional_vector_field。
  • DiT 的超参变了:img_size=4, patch_size=1, c=128。 patch_size=1 意味着每个空间位置就是一个 token,共 $4\times4=16$ 个 token(像素空间是 64 个)。Patchifier 退化成一个 $1\times1$ 卷积($128\to256$ 的逐位置线性投影)。这正是 Q3.2 里坚持不硬编码 c_in 的回报。
Q5.1 的常见错误
  • 忘了 torch.no_grad()——见上。
  • 忘了把 VAE 设成 eval()——本 lab 的 VAE 里没有 dropout 和 BatchNorm(用的是 GroupNorm/LayerNorm),所以影响不大,但养成习惯是好的。
  • 把标签丢弃放进了 no_grad 块外/内的错误位置——其实无所谓(整数张量本来就没梯度),但放外面更清晰。
  • 忘了改 p_simple_shape——还写 [1,32,32] 的话,path.p_simple.sample() 会给出错误形状的初始噪声,采样时崩掉。
  • patch_size 还写 4——$4\times4$ 的隐变量除以 patch 4 得到 1 个 token,注意力完全退化(自己看自己),模型只能学到逐位置的映射。不报错,效果显著变差。
  • 直接用 Part 3 训好的 DiT 权重——通道数从 1 变成 128,形状对不上。这里要重新初始化一个 DiT。

5.2 训练结果:一个必须解释清楚的现象

配置:10000 步、batch 256、$\eta=0.35$、学习率 $4\times10^{-4}$。模型 39.844 MiB。

隐空间 DiT 的训练损失曲线,从 2.0163 只降到 1.9431
隐空间 DiT 的损失:从 2.0163 只降到 1.9431——纵轴范围只有 $[1.91, 2.03]$,总共下降不到 5%。对比像素空间 DiT 的 $1.9788\to0.1906$(下降 90%),这条曲线看起来简直像是「没训起来」。但样本是好的(见下图)。这是本 lab 最容易导致误判、也最值得深挖的一个现象,下一节专门解释。
隐空间 DiT 在三个引导强度下的采样结果
隐空间 DiT 的采样结果,布局与 Part 3 完全一致(11 行 × 10 列,最后一行是空标签,三个面板对应 $w=1/3/5$)。这张图验证的是「损失数值不能跨空间比较」:一个损失只降了 5% 的模型,产生的样本质量与像素空间那个降了 90% 的模型在同一档次上。仔细比较可以看到隐空间样本的笔画略微更「圆润」——这是 VAE 解码器带来的平滑,是有损压缩的固有代价;同时 $w$ 增大带来的规整化与多样性坍缩趋势与像素空间完全一致,说明 CFG 在隐空间里工作方式没有任何变化。

6. 为什么隐空间的损失几乎不下降?

这一节是本页的核心结论之一。答案是:flow matching 的损失有一个不可约下界,而这个下界在隐空间里比在像素空间里高一个数量级。

6.1 不可约下界从哪来

条件期望是最优回归量,残差是条件方差

训练目标是

$$ \mathcal{L}(\theta)=\E_{t,z,\epsilon}\big\|u_t^\theta(x_t\mid y)-u_t(x_t\mid z)\big\|^2,\qquad x_t=\alpha_t z+\beta_t\epsilon. $$

注意回归的输入是 $(x_t,t,y)$,而目标 $u_t(x_t\mid z)$ 还依赖于 $z$——同一个 $x_t$ 可以由无数个不同的 $(z,\epsilon)$ 组合产生,它们的目标各不相同。对任意随机变量对 $(A,B)$ 与任意函数 $g$,有标准的偏差–方差分解

$$ \E\norm{g(A)-B}^2=\E\norm{g(A)-\E[B\mid A]}^2+\E\big[\tr\Cov(B\mid A)\big]. $$

第二项与 $g$ 无关。所以即使网络完美学到了最优解 $u_t^\theta=\E[u_t(x_t\mid z)\mid x_t,t,y]=u_t(x_t\mid y)$,损失也只能降到

$$ \mathcal{L}^\star=\E_{t,x_t}\big[\tr\Cov\big(u_t(x_t\mid z)\,\big|\,x_t,t,y\big)\big]. $$

这是数据分布本身的属性,与模型架构、训练时长都无关。(Lab 2 已经讲过这一条,这里是它的定量版本。)

6.2 把下界算出来

本 lab 用 LinearAlpha / LinearBeta,即 $\alpha_t=t$、$\beta_t=1-t$。把 $x_t=tz+(1-t)\epsilon$ 代入条件向量场:

$$ u_t(x_t\mid z)=\Big(\dot\alpha_t-\tfrac{\dot\beta_t}{\beta_t}\alpha_t\Big)z+\tfrac{\dot\beta_t}{\beta_t}x_t=\dot\alpha_t z+\dot\beta_t\epsilon=z-\epsilon. $$

所以回归目标就是 $z-\epsilon$,下界是 $\E_{t,x_t}\big[\Var(z-\epsilon\mid x_t)\big]$(这里按 .mean() 的约定,取逐维平均)。

白噪声情形:下界恰好是 $\frac{\pi\sigma}{2}$

先算一个可以完全手算的极端情形:设 $z$ 的各维独立、$z\sim\N(0,\sigma^2 I)$,与 $\epsilon\sim\N(0,I)$ 独立。此时 $(z-\epsilon, x_t)$ 联合高斯,条件方差有闭式:

$$ \Var(z-\epsilon\mid x_t)=\Var(z-\epsilon)-\frac{\Cov(z-\epsilon,\,x_t)^2}{\Var(x_t)} =(\sigma^2+1)-\frac{\big(t\sigma^2-(1-t)\big)^2}{t^2\sigma^2+(1-t)^2}. $$

通分化简,分子出奇地干净:

$$ \begin{aligned} &(\sigma^2+1)\big(t^2\sigma^2+(1-t)^2\big)-\big(t\sigma^2-(1-t)\big)^2\\ &\quad=(1-t)^2\sigma^2+t^2\sigma^2+2t(1-t)\sigma^2=\sigma^2\big[t+(1-t)\big]^2=\sigma^2. \end{aligned} $$

于是 $\Var(z-\epsilon\mid x_t)=\dfrac{\sigma^2}{t^2\sigma^2+(1-t)^2}$。对 $t\sim U[0,1]$ 取期望,令 $A=\sigma^2+1$:

$$ \mathcal{L}^\star=\sigma^2\int_0^1\frac{\ud t}{At^2-2t+1} =\sigma^2\cdot\frac{1}{\sigma}\left[\arctan\frac{At-1}{\sigma}\right]_0^1 =\sigma\left(\arctan\sigma+\arctan\frac{1}{\sigma}\right)=\frac{\pi\sigma}{2}, $$

最后一步用了恒等式 $\arctan\sigma+\arctan\frac{1}{\sigma}=\frac{\pi}{2}$(对 $\sigma>0$)。

$$ \boxed{\ \text{若 }z\sim\N(0,\sigma^2 I)\text{ 且各维独立,则 }\mathcal{L}^\star=\frac{\pi\sigma}{2}.\ } $$

特别地,$\sigma=1$ 时 $\mathcal{L}^\star=\frac{\pi}{2}\approx1.5708$。

6.3 两个空间的对比

像素空间 DiT隐空间 DiT
数据 $z$归一化后的 MNIST 像素VAE 编码后的隐变量
每维方差约 1(Normalize 保证)被 $\beta=10$ 的 KL 推向 1
维间相关性极强(大片黑背景、笔画连通)极弱(KL 项主动惩罚相关结构)
损失首 → 末1.9788 → 0.19062.0163 → 1.9431
「白噪声」参考下界$\pi/2\approx1.571$$\pi/2\approx1.571$

像素空间为什么能降到 0.19?因为 MNIST 高度结构化:图像的绝大部分是纯黑背景,笔画是连通的粗线条,相邻像素强相关。上面那个 $\frac{\pi\sigma}{2}$ 的公式只对各维独立的白噪声成立;当各维强相关时,观测到 $x_t$ 就等于同时观测到了 1024 个带噪线索,它们联合起来对 $z$ 的约束远比单维的情形强,条件方差因此小得多。0.19 远低于 1.571,量化地说明了「MNIST 离白噪声有多远」。

隐空间为什么停在 1.94?因为 $\beta=10$(讲义约定下的 5)的 KL 项正是在主动把隐变量推成白噪声:均值推向 0、方差推向 1、各维之间的冗余被压掉(KL 到 $\N(0,I)$ 的散度对任何偏离独立标准正态的结构都要收费)。于是隐变量非常接近上面那个理想化的假设,损失几乎就卡在理论下界上。

反过来用这个公式还能反推隐变量的尺度:由 $\mathcal{L}^\star=\frac{\pi\sigma}{2}$ 且实测 $\mathcal{L}=1.9431\ge\mathcal{L}^\star$,得

$$ \sigma\le\frac{2\times1.9431}{\pi}\approx1.237, $$

即隐变量每维标准差不超过约 1.24(等号在模型达到最优时取到)。这个数与 Lecture 4 §7.7 提到的现象一致:实际系统里隐分布的尺度并不严格是 1,所以工程实现通常要在编码结果上乘一个全局缩放常数,把经验标准差调到 1 附近再交给扩散模型——因为噪声调度 $\alpha_t,\beta_t$ 都是按「数据方差约为 1」设计的。本 lab 省略了这一步,代价就是那个略高于 $\pi/2$ 的损失平台。

隐空间 DiT 在 step 500/2000/4500/7000/9500 经冻结 VAE 解码后的采样演进
本节论点最直接的视觉佐证。隐空间 DiT 在 step 500 / 2000 / 4500 / 7000 / 9500 的采样(经冻结 VAE 解码,$w=3.0$,行布局同前)。把它和上面那条几乎水平的损失曲线并排看:损失从 2.0163 到 1.9431 总共只动了 0.073,而样本质量的改善是肉眼可见、逐格递进的——step 500 已经能画出笔画状的东西,但身份完全是乱的(每一行都不匹配它的标签);step 2000(损失大约刚跌到 1.94,此后基本不再下降)十个数字全部正确;而在损失已经完全走平的 step 2000 → 9500 之间,笔画仍在持续变干净、变规整,同类样本的一致性明显提升,最后一行的无条件样本也从破碎变得完整。结论很清楚:这 0.073 的下降里装着的信息量,远比它的数值看起来要多——因为损失的绝大部分(约 1.94)是不可约的条件方差,是模型无论如何都消不掉的常数底。用「损失降了百分之几」来判断隐空间扩散有没有训好,是完全错误的做法。
结论:损失的绝对值不能跨空间比较
  • 要看的不是损失降了多少,而是「相对于各自下界还剩多少」。像素空间 $0.1906$ 对应一个很低的下界(因为数据结构强);隐空间 $1.9431$ 对应一个接近 $\pi/2$ 的下界(因为数据接近白噪声)。后者其实离最优更近。
  • 一个可操作的判据:把你的隐变量的经验标准差 $\hat\sigma$ 量出来(z.std()),算 $\frac{\pi\hat\sigma}{2}$,和你的损失平台比。差得不多 = 训好了;差很多 = 真有问题。
  • 最终裁判永远是样本。本 lab 里两个模型的采样质量在同一档次,这才是「都训好了」的证据。
  • 这和 Lab 2 里「损失收敛但不到 0」是同一个道理,只是这里被放大到了一眼可见的程度。
顺带回答一个自然的疑问:那隐空间扩散图什么?

既然隐空间的目标「更难」(下界更高),为什么还要用它?三个理由:

  1. 计算量。本 lab 的隐空间 token 数是 16,像素空间是 64;注意力是 $O(n^2)$,差 16 倍。真实系统里($512\times512$ 图,8 倍下采样)差距是 $4096$ 倍。
  2. 「难」不等于「差」。下界高只是说残差方差大,不代表最优解学不到。扩散模型学的是条件期望,它在白噪声般的隐空间里照样能学好——只是损失读数更接近一个常数。
  3. 解耦。感知细节(笔画锐度、纹理)交给 VAE 解码器,语义结构交给扩散模型。这是 Stable Diffusion 系列的核心设计。

7. 如何验证自己的实现:一套写完就跑的检查

这个 lab 的训练动辄几十分钟,「写完全部代码再开训、跑完才发现形状错」是最贵的失败模式。下面这套检查全部是纯 CPU、秒级完成的,建议每写完一个模块就跑对应的那几条。括号里是本页实现的实测值。

7.1 不变量检查(能抓出绝大多数错误)

模块检查什么为什么它有效实测
FourierEncoder每行平方和 $=$ half_dim$\cos^2+\sin^2=1$,与 $t$ 无关误差 0
FourierEncoder$t=0$ 时输出 $[1,\dots,1,0,\dots,0]$端点值,能抓出 cos/sin 拼反通过
Patchifier只改一个 $4\times4$ 块 ⇒ 只有 1 个 token 变patch 必须是空间局部的,抓 rearrange 写反受影响 token 数 = 1
MHA与 F.scaled_dot_product_attention 对拍缩放、softmax 维度、头划分一次全验最大误差 8.94e-08
DiffusionTransformerLayer初始化后 $=$ 恒等映射adaLN-Zero 的定义性质最大偏差 0.00e+00
DiffusionTransformerLayer随机化 cond 末层后,改变 $c$ 必须改变输出补上上一条的盲区(防止条件被完全忽略)差异 > 0
ResidualBlock初始化后 $=$ 恒等映射零初始化 conv2 的定义性质最大偏差 0.00e+00
CFGVectorFieldODE$w=1$ 时 $=$ 纯条件向量场组合公式的边界情形误差 0
CFGVectorFieldODE$w=0$ 时 $=$ 纯无条件向量场同上,另一端误差 0
VAE.compute_loss$q=$ 先验时 KL 项恰为 0$0+1-0-1=0$,抓符号错误$<10^{-7}$
VAE.compute_lossKL 项恒非负(50 次随机测试)$\alpha-\log\alpha-1\ge0$最小值 $\ge-10^{-6}$
VAE.compute_loss完美重构 + $q=$ 先验 ⇒ 总损失 $=0$两项同时归零,端到端验证$<10^{-6}$

7.2 形状检查

模块输入期望输出
FourierEncoder(dim=64)(8,)(8, 64)
Patchifier(32, 4, 1, 48)(3, 1, 32, 32)(3, 64, 48)
Depatchifier(32, 4, 48, 8, 1)(3, 64, 48)(3, 1, 32, 32)
MHA(64, 8)(2, 10, 64)(2, 10, 64)
DiffusionTransformerFlowModel(5,1,16,16), (5,), (5,)(5, 1, 16, 16)
ResidualBlock(8) / AttnBlock(8)(2, 8, 16, 16)(2, 8, 16, 16)
VAE(1, [8,16,32,64]) 的 z_mean(2, 1, 32, 32)(2, 64, 4, 4)
VAE 的 x_mean / logvar同上(2,1,32,32) / torch.Size([])

本页的实现在这套检查上 22 项全部通过。注意几处「刻意做成小尺寸」的设计:用 img_size=16、num_layers=2、hidden_channels=[8,16,32,64] 就能在 CPU 上瞬间跑完,同时完整覆盖所有形状逻辑——不要用 img_size=32, num_layers=8 做单元测试,慢且没有额外收益。

7.3 推荐的验证顺序

  1. Q2.2 + Q2.3 → Sanity Check 2.4(3000 步,约 1–3 分钟)。这一步同时验证训练循环、CFG 组合公式、类别嵌入。不通过就绝不要往下走。
  2. Q3.1 → Q3.2 → Q3.3 → Q3.4,每写完一个立刻跑上表对应的行。特别是 MHA 的对拍,它一次能验四件事。
  3. Q3.5 写完先跑一次前向 + model_size_b,确认输出形状和 40.324 MiB 这个量级。
  4. 启动 DiT 训练,开 ckpt_every=1000,看第 1000/2000 步的图。3000 步还是纯噪声就停下来查。
  5. Q4.1–Q4.6,跑形状链路和两个恒等映射检查。
  6. Q4.7,跑三条损失性质检查。
  7. 启动 VAE 训练(几分钟),跑插值图。插值图不平滑就别继续——Part 5 的质量完全依赖 VAE。
  8. Q5.1,启动隐空间训练。看到损失只在 1.9x 徘徊时,不要慌,去看 checkpoint 采样图。

8. 训练实测数据汇总

模型参数量步数batch学习率损失 首 → 末
MLP sanity check($\R^2$ 上的 GMM)0.259 MiB3000250$10^{-3}$2.8910 → 0.8401
DiT(像素空间,$p=4$,8 层,$d=256$,8 头)40.324 MiB20000256$4\times10^{-4}$1.9788 → 0.1906
VAE(hidden $=[16,32,64,128]$,$\beta=10$)6.150 MiB500064$10^{-3}$4.3153 → −0.6341
Latent DiT($128\times4\times4$,$p=1$,8 层)39.844 MiB10000256$4\times10^{-4}$2.0163 → 1.9431

硬件:单张 RTX 5080(16 GB),PyTorch 2.11.0 + CUDA 12.8。四段训练合计约 1.5 小时(含 MNISTSampler 的缓存改写;不改写的话仅 DiT 一段就要成倍时间)。原题面说「~15 A100 minutes」是指 DiT 那一段。

三个数字的读法(容易读错,值得单独强调)
  • DiT 的 0.1906 和 Latent DiT 的 1.9431 不可比。见第 6 节:它们各自的不可约下界差了一个数量级。
  • VAE 的 −0.6341 是负的,正常。重构项含 $\log\sigma_\theta^2$,$\sigma_\theta<1$ 时为负且无下界。负对数似然对连续分布本来就可以任意负。
  • 两个 DiT 的参数量差 0.48 MiB,全在 Patchifier 和 Depatchifier 上。像素空间:$1\times4\times4\times256$ 的 patch 卷积;隐空间:$128\times1\times1\times256$ 的逐位置投影。主干(8 层 transformer)完全一样。这也从侧面说明「切块」这一步在参数量上有多便宜。

本 lab 小结

题号对象一句话要点
Q2.2CFGTrainer.get_train_loss四步采样 + MSE;$t$ 是 (b,);$\times(1-\epsilon)$ 避开 $\beta_t=0$
Q2.3MLPConditionalVectorField拼接 $[x,\,e_y,\,t]$;num_classes + 1 里的 $+1$ 就是 $\varnothing$
Q3.1FourierEncoder随机 Fourier 特征打破 MLP 的谱偏置,让网络分辨相近的 $t$
Q3.2Patchifierkernel $=$ stride $=$ patch 的卷积 $\equiv$ 切块 + 各自线性投影
Q3.3MHA / DiT layer / DiT$1/\sqrt{d_h}$ 稳住 softmax;adaLN-Zero 让初始 block 是恒等;位置编码打破置换等变
Q3.4DepatchifierRearrange 必须与 Patchifier 严格对偶;末尾卷积抹平块边界
Q3.5DiffusionTransformerFlowModel时间嵌入 $+$ 类别嵌入(相加),驱动每层的 adaLN
Q4.1ResidualBlockGroupNorm(1,C) $=$ 图像上的 LayerNorm;零初始化 ⇒ 初始恒等
Q4.2AttnBlock像素即 token、通道即特征;heads=1 保证任意通道数可整除
Q4.3 / Q4.5EncoderBlock / DecoderBlockstride-2 卷积下采样;最近邻 + 卷积上采样以避开棋盘格伪影
Q4.4 / Q4.6Encoder / Decoderch_out = hidden[1:] + [None] 是题眼;logvar 是共享标量
Q4.7VAE.compute_loss高斯 NLL + 闭式 KL;.mean() 约定下 $\beta_{\text{讲义}}=\beta_{\text{代码}}\cdot d/k$
Q5.1LatentCFGTrainer与 Q2.2 只差「数据来自冻结 VAE 的编码」,必须 no_grad
Lab 3 的五个 takeaway
  • 训练循环没变,变的全是架构。15 个 TODO 里只有 2 个在动损失,其余 13 个都在回答同一个工程问题:怎么把一个高维函数 $u_t^\theta(x\mid y)$ 用张量算子搭出来。这正是从 Lab 2 到真实系统的全部距离。
  • 「初始即恒等」是深层网络的通用护身符。adaLN-Zero 的 $\alpha=0$、ResidualBlock 的零初始化 conv2,是同一招的两次出场。它们都能用一行 assert (f(x) - x).abs().max() == 0 验证,而且必须严格为 0。
  • CFG 是一个纯推理时的旋钮,但需要训练时的配合。训练侧只多了一行标签丢弃,推理侧只多了一次网络调用和一个线性组合。三联图上肉眼可见的保真度–多样性权衡,全部来自这两处改动。
  • 损失的绝对值不能跨空间比较。flow matching 的不可约下界等于回归目标关于 $x_t$ 的条件方差;在白噪声般的隐空间里这个下界是 $\frac{\pi\sigma}{2}\approx1.57\sigma$,在强结构的像素空间里则低得多。看到隐空间损失卡在 1.94 不动,要去算下界,而不是怀疑代码。
  • 每写完一个模块就验一次形状与不变量。这个 lab 的所有致命错误(rearrange 写反、softmax 维度错、$t$ 的形状错)都能被秒级的检查抓到,而它们中的任何一个都能让你白跑 20000 步。

做完这三个 lab,你手上就有了一条完整的链路:Lab 1 的模拟器 → Lab 2 的 flow matching / score matching → Lab 3 的 CFG + DiT + VAE + latent diffusion。把 MNIST 换成更大的数据集、把 DiT 加深、把 VAE 换成带感知损失和对抗损失的版本,你得到的就是 Stable Diffusion 3 的骨架——Lecture 4 §6.3 逐条对照过这件事。

延伸阅读