Skip to content
Published on

人人可懂的 AI 第6篇 — 用111万参数的扩散模型从单词画出数字

分享
Authors

引言 — 第5篇留下的作业

第5篇的上色把形状保留得很好,颜色却发灰。原因不在模型而在损失。黑白汽车可能是红也可能是蓝,而 L1 要求"最小化绝对误差",于是 押多个正解的中间值——灰色——成了最安全的策略

本篇的问题同样是多解的。对应 "three" 这个单词的手写 3 有成千上万种。若像第5篇那样用回归损失让它"直接预测 3 的图像",得到的会是所有 3 的平均:一团模糊的污迹。

扩散模型正面绕开了这个问题。 它不一次性给出答案,而是把生成变成从噪声里一点点清除的过程。

核心想法 — 破坏容易,还原困难

把干净图像逐步加噪直到变成纯噪点,很容易,一行公式就够。难的是反方向。

扩散的洞见是: 把一个难问题拆成 400 个简单问题 。从噪点一步造出图像很难,但"让它稍微干净一点"是可学的。重复 400 次即可。

Forward — 加噪的过程

T = 400
beta = torch.linspace(1e-4, 0.02, T)   # 每一步要混入的噪声量
alpha = 1 - beta
abar = torch.cumprod(alpha, 0)         # 累积乘积

abar[t] 表示"到时刻 t 为止原图还剩多少"。t=0 时接近 1 (原样),t=399 时接近 0 (纯噪点)。

这里出现了扩散的第一个机关。 不必按顺序走完 400 步,可以一次跳到任意 t。

t = torch.randint(0, T, (256,), device=dev)
eps = torch.randn_like(x0)                          # 标准正态噪声
a = abar[t][:, None, None, None]
xt = a.sqrt() * x0 + (1 - a).sqrt() * eps           # 一行完成 t 步加噪

正是这一行让训练变得高效。每个批次抽取不同的 t,所有时刻并行学习。

训练目标 — 猜噪声,而非图像

loss = F.mse_loss(model(xt, t.float(), w), eps)

模型预测的不是原图 x0,而是 被混入的噪声 eps

数学上二者等价:知道 xtt 就能从 eps 算出 x0,反之亦然。但实践中预测噪声要好得多。

原因在于目标的性质。若让它预测 x0,当 t 很大、几乎只剩噪点时,模型实际上要在没有信息的情况下编造图像 —— 第5篇的均值回归就在这里重演。而 eps 在任何 t 上都是标准正态分布 。目标的尺度和分布恒定,训练自然稳定。

加条件 — 用单词指挥画面

把时刻 t 和单词在同一向量空间里相加,注入 U-Net 的中部。

class CondUNet(nn.Module):
    def __init__(self):
        super().__init__()
        self.temb = nn.Sequential(nn.Linear(64, 128), nn.GELU(), nn.Linear(128, 128))
        self.wemb = nn.Embedding(10, 128)          # zero~nine
        self.e1 = block(1, 48); self.e2 = block(48, 96); self.e3 = block(96, 192)
        self.cond3 = nn.Linear(128, 192)
        self.d2 = block(192 + 96, 96); self.d1 = block(96 + 48, 48)
        self.out = nn.Conv2d(48, 1, 1)

    def forward(self, x, t, w):
        half = torch.exp(-np.log(10000) * torch.arange(32, device=x.device) / 32)
        te = torch.cat([torch.sin(t[:, None] * half), torch.cos(t[:, None] * half)], 1)
        c = self.temb(te) + self.wemb(w)           # 时刻 + 单词
        s1 = self.e1(x)
        s2 = self.e2(F.max_pool2d(s1, 2))
        h = self.e3(F.max_pool2d(s2, 2)) + self.cond3(c)[:, :, None, None]
        ...

时刻嵌入用正弦余弦,与 Transformer 的位置编码同出一辙。直接喂整数 t 会给网络一个难以处理的尺度,而把它铺展到不同频率的正弦波上,相邻时刻便拥有相近的向量。

注入位置选在 最深层 (e3 的输出) 也是有意的。那里分辨率最小 (8×8)、每通道信息密度最高,在此施加条件会影响后续整个解码器。

self.cond3(c)[:, :, None, None] 里的两个 None 把 (批, 通道) 的向量变成 (批, 通道, 1, 1),以便向空间维度广播。也就是说 条件被均等地加到每一个位置上

Reverse — 倒着走 400 次

@torch.no_grad()
def sample(words):
    x = torch.randn(len(words), 1, 28, 28, device=dev)      # 从纯噪点出发
    w = torch.tensor([WORDS.index(s) for s in words], device=dev)
    for ti in reversed(range(T)):                            # 399 -> 0
        t = torch.full((len(words),), ti, device=dev, dtype=torch.float32)
        eps = model(x, t, w)                                 # 预测噪声
        a, ab = alpha[ti], abar[ti]
        x = (x - (1 - a) / (1 - ab).sqrt() * eps) / a.sqrt() # 去掉一步
        if ti > 0:
            x = x + beta[ti].sqrt() * torch.randn_like(x)    # 重新注入新噪声
    return ((x.clamp(-1, 1) + 1) / 2).squeeze(1).cpu().numpy()

最后一行的重新注入看起来很反常:好不容易去掉的噪声,为什么又加回去?

因为它正是多样性的来源。 每一步都注入新的随机数,所以同一个单词每次都会得到不同的 3 —— 与第5篇中回归被困在单一平均值上恰好相反。扩散是一种 在多个正解中按概率挑出一个 的结构。只有最后一步 (ti == 0) 略去注入,从而得到干净的结果。

训练日志

[    9.2s] params {"params": 1111777, "T": 400, "n": 60000}
[   10.1s] step      0 mse=1.1502
[   12.2s] step    100 mse=0.0807
[  959.7s] step  44200 mse=0.0324
[  960.6s] SUMMARY {"steps": 44215, "final_mse": 0.0320}

100 步、两秒之内,MSE 从 1.15 掉到 0.081。预测目标是标准正态分布,初始损失从 1 附近起步也很自然。随后平缓下降到 0.032。

结果

每行是 "zero" 到 "nine" 十个单词,两行是同一批单词的不同采样。

单词条件扩散生成结果 —— 每行为 zero 到 nine,两行是同一单词的不同采样

第一行 0 到 9 全部正确。 笔画的粗细与倾斜也很有手写的味道。

第二行错了两个。 "five" 的位置出现了接近 9 的形状,"nine" 的位置出现了接近 7 的形状。十中有八,80%。

对比两行便能看出扩散的性格。 同一个单词,字体却不同。 第一行的 2 与第二行的 2 弯曲程度不同,第一行的 7 与第二行的 7 有无横杠也不同。没有被平均抹平,每次都产出一份具体的手写字。这正是与第5篇发灰颜色相对照之处。

弄错的 5 和 9 在形态上本就相似。111 万参数、16 分钟训练,留下这样的混淆是合理的。

小结

项目
参数量1,111,777 (1.11M)
训练时间960.6秒 (16分钟)
步数44,215
扩散步数 T400
最终 MSE0.0320
生成正确率8/10

三句话概括。

扩散把 一个困难的生成问题拆成 400 个简单的去噪问题。 forward 是一行公式,训练通过跳到任意 t 实现并行。

模型 预测噪声而非图像。 目标在任何 t 上都是标准正态分布,训练因而稳定,也避开了直接预测 x0 时的均值回归。

reverse 过程中的 噪声再注入创造了多样性。 同一条件下每次给出不同结果,这是与回归模型的决定性差异。

不过本次实现是基础版。必须走完 400 步所以慢,也没有调节条件强度的旋钮。系列第11篇将用 Classifier-Free Guidance 和 DDIM 解决这两点,并用实测说明这两种技术并非彼此独立。

🧠 理解度自测

1. 为什么让模型预测噪声而不是原图?

二者数学上等价,但训练稳定性不同。噪声在任何时刻 t 上都是标准正态分布,目标的尺度与分布保持恒定。而让它预测原图,则在 t 很大、只剩噪点时模型必须凭空编造图像 —— 正是在这里,向多个正解的平均逃逸的问题会发生。

2. reverse 过程中为什么要重新注入噪声?

为了创造多样性。每步注入新的随机数,因此同一个条件单词每次都会得到不同的手写字 —— 结果图中两行同样的数字却字体不同即为证据。这正是与困在单一平均值上的回归模型的决定性差异。只有最后一步略去注入,以得到干净的输出。

3. 训练时为什么不必按顺序走完 400 步?

因为存在闭式解 xt = sqrt(abar[t]) * x0 + sqrt(1 - abar[t]) * eps,可一次算出任意 t 对应的加噪状态。于是每个批次可以随机抽取不同的 t,把所有时刻并行训练。

4. 为什么把条件向量注入 U-Net 最深层,并加上 [:, :, None, None]

最深层分辨率最小 (8×8)、每通道信息密度最高,在此施加条件会贯穿后续整个解码器。两个 None 把 (批, 通道) 的向量变为 (批, 通道, 1, 1) 以便向空间维度广播,从而把条件均等地加到每一个位置上。

参考资料