注释扩散模型 – 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_size、dim_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_schedule、cosine_beta_schedule等)以在样本保真度和速度之间进行权衡。 - 使用 Huber 损失提升鲁棒性;代码同样支持 L1/L2。
- 周期性采样(
save_and_sample_every)对于监控模式崩溃或训练不稳定至关重要。 - 为加速推理,可考虑实现 DDIM 或随机采样器变体,或采用近期的“少步”求解器。
Sources
- OriginalThe Annotated Diffusion Model