注释扩散模型 – DDPM 实现的详细演练

TL;DR

Hugging Face 发布了一个完整的、带注释的 PyTorch 笔记本,实现了原始的去噪扩散概率模型(DDPM)——一种从纯高斯噪声中学习去噪图像的神经网络——并展示了在低分辨率数据集上的训练以及高质量图像的采样。该资源阐明了数学原理、网络结构、训练循环和推理过程,使扩散模型对实践者更加易于使用。


扩散模型的作用(关键要点)

扩散模型学习一种逆向去噪过程,将随机高斯噪声转化为真实数据;前向过程按照预定义的时间表添加噪声,而 U‑Net 风格的神经网络在每个时间步预测添加的噪声,从而在固定步数后实现图像生成。


前向扩散:使用方差调度添加噪声

Conclusion: 前向过程是一个闭式高斯转移,可以直接采样任意噪声时间步,省去迭代添加噪声的需求。

前向扩散分布为

q(x_t | x_{t-1}) = N(x_t ; sqrt(1-β_t)·x_{t-1}, β_t·I)

其中方差调度 (β_1,…,β_T) 单调(例如线性、余弦、二次或 sigmoid)。通过重复应用此转移,边缘分布

q(x_t | x_0) = N(x_t ; sqrt(\bar α_t)·x_0, (1-\bar α_t)·I)

可以使用预先计算的 (\bar α_t = \prod_{s=1}^t (1-β_s)) 直接采样。此“良好属性”使得训练算法能够在一次步骤中抽取随机时间步 (t) 并对真实图像 (x_0) 进行腐蚀。


逆向扩散:通过噪声预测学习均值

Conclusion: 模型被训练以预测在时间步 (t) 添加的精确高斯噪声 (ε);预测的噪声随后用于计算逆向高斯分布的均值。

逆向条件被假设为高斯分布:

p_θ(x_{t-1} | x_t) = N(x_{t-1} ; μ_θ(x_t, t), σ_t^2·I)

DDPM 将 (σ_t^2) 固定为已知的前向方差,仅学习 (μ_θ)。通过以噪声预测器 (ε_θ) 重新参数化均值:

μ_θ(x_t, t) = (1/√α_t)·(x_t - (β_t/√(1-\bar α_t))·ε_θ(x_t, t))

训练损失简化为真实噪声 (ε) 与网络输出 (ε_θ) 之间的均方误差:

L_t = || ε - ε_θ(x_t, t) ||^2

在每个批次随机抽样 (t) 可得到变分下界的无偏估计。


网络架构:条件 U‑Net

Conclusion: DDPM 使用时间条件的 U‑Net,配备正弦位置嵌入、残差块、注意力机制和组归一化,以在每个空间分辨率上预测噪声。

关键组件:

  • Sinusoidal position embeddings 编码时间步 (t),并通过小型 MLP 注入到每个 ResNet 块中。
  • Weight‑standardized convolutions 与组归一化结合时提升训练稳定性。
  • ResNet blocks(两个卷积‑归一化‑SiLU 层,可选的来自时间嵌入的尺度‑平移)提供主要特征变换。
  • Attention modules(全多头或线性注意力)捕获长程依赖。
  • Group normalization 在注意力之前应用(PreNorm)。
  • Down/upsampling paths 将空间分辨率减半或加倍,同时保持通道深度,映射经典 U‑Net 结构。
  • Final head 将拼接的瓶颈特征映射回图像形状。

完整的 PyTorch 类 Unet 将这些组件组装起来,接受噪声图像张量和时间步张量,返回形状相同的噪声张量。


训练循环:随机时间步抽样与 Huber 损失

Conclusion: 训练通过对每个批次抽取随机时间步、使用闭式前向过程对输入进行腐蚀,并最小化真实噪声与预测噪声之间的 Huber 损失来进行。

伪代码(简化版):

for epoch in range(num_epochs):
    for batch in dataloader:
        t = torch.randint(0, T, (batch_size,)).to(device)          # uniform timestep
        loss = p_losses(model, batch, t, loss_type='huber')        # MSE/Huber on noise
        loss.backward()
        optimizer.step()

辅助函数 p_losses 调用 q_sample 获取 (x_t),随后计算相对于网络输出 model(x_t, t) 的损失。周期性采样(使用逆向过程)可可视化训练进度。


采样(推理):逆转扩散链

Conclusion: 生成从纯高斯噪声开始,迭代应用学习到的去噪步骤;每一步使用预测的噪声计算后验均值,并根据已知方差调度添加校准的高斯噪声。

算法(DDPM 论文中的算法 2):

x_T = torch.randn(shape)                     # start from noise
for t in reversed(range(T)):
    pred_noise = model(x_t, t)
    mean = (1/√α_t) * (x_t - β_t/√(1-\bar α_t) * pred_noise)
    if t > 0:
        x_{t-1} = mean + √posterior_variance_t * torch.randn_like(x_t)
    else:
        x_0 = mean

提供的笔记本在 p_sample_loop 中实现了该循环,并将去噪轨迹可视化为 GIF。


Fashion‑MNIST 端到端示例

Conclusion: 本教程在 28×28 Fashion‑MNIST 数据集上训练 DDPM,经过几千步训练后即可生成可辨认的服装图像。

  • 数据加载使用 🤗 datasets 库,并在加载时进行变换(随机水平翻转,缩放至 ([-1,1])。
  • 模型超参数:dim=image_sizedim_mults=(1,2,4)channels=1(用于灰度图像)。
  • 优化器:Adam,学习率 1e‑3。
  • 训练进行 6 个 epoch;损失快速下降至 0.05 以下。
  • 训练后采样得到清晰的 T 恤形状,验证实现可在低分辨率数据上工作。

该工作在扩散文献中的位置

Conclusion: 该带注释的实现复现了原始 DDPM(Ho 等,2020),并为后续众多进展提供了基础,如方差学习(Nichol 等,2021)、级联扩散(Ho 等,2021)、无分类器引导以及大规模文本到图像模型(DALL‑E 2、Imagen)。

博客文章中列出的关键后续工作:

  • Improved DDPM – 学习均值和方差,提升样本质量。
  • Cascaded Diffusion Models – 堆叠多个扩散模型以实现高分辨率合成。
  • Diffusion Models Beat GANs – 通过架构改进和分类器引导展示出更优的 FID 分数。
  • Classifier‑Free Guidance – 在条件生成时无需外部分类器。
  • DALL‑E 2 & ImageGen – 将扩散与 CLIP 嵌入或大型语言模型结合,实现文本条件的图像合成。

主要缺点仍是需要大量去噪步骤(通常约 1000 步),尽管近期研究(如高阶求解器)已将其降低至仅 10 步左右。


开发者的实用要点

  • 笔记本提供了可直接运行的参考实现;可更换数据集、增大 T,或将 U‑Net 替换为更大的主干网络以进行扩展。
  • 调整方差调度(linear_beta_schedulecosine_beta_schedule 等)以在样本保真度和速度之间进行权衡。
  • 使用 Huber 损失提升鲁棒性;代码同样支持 L1/L2。
  • 周期性采样(save_and_sample_every)对于监控模式崩溃或训练不稳定至关重要。
  • 为加速推理,可考虑实现 DDIM 或随机采样器变体,或采用近期的“少步”求解器。

Sources