共计 1822 个字符,预计需要花费 5 分钟才能阅读完成。
背景介绍
生成模型在过去几年取得了巨大进展,从最早的 GAN、VAE 到现在的扩散模型,每种方法都有其独特的优势。扩散模型之所以受到广泛关注,主要是因为它能够生成高质量的样本,并且在训练过程中更加稳定。与 GAN 相比,扩散模型避免了模式崩溃的问题;与 VAE 相比,扩散模型生成的样本质量更高。

数学原理
前向扩散过程
前向扩散过程可以看作是一个马尔可夫链,逐步向数据添加高斯噪声。给定一个数据点 $x_0$,前向扩散过程定义如下:
$$
q(x_t|x_{t-1}) = \mathcal{N}(x_t; \sqrt{1-\beta_t}x_{t-1}, \beta_t\mathbf{I})
$$
其中 $\beta_t$ 是噪声调度参数,控制每一步添加的噪声量。经过 $T$ 步扩散后,数据 $x_T$ 将近似服从标准高斯分布。
反向去噪过程
反向去噪过程的目标是从噪声中重建原始数据。这可以通过变分推断来实现,目标是最大化对数似然的下界:
$$
\log p_\theta(x_0) \geq \mathbb{E}{q(x \right]
$$}|x_0)} \left[\log \frac{p_\theta(x_{0:T})}{q(x_{1:T}|x_0)
通过优化这个下界,我们可以学习到一个能够逐步去噪的模型。
PyTorch 实现
噪声调度器
我们首先实现一个线性噪声调度器,用于控制每一步的噪声量:
import torch
def linear_beta_schedule(timesteps, beta_start=1e-4, beta_end=0.02):
return torch.linspace(beta_start, beta_end, timesteps)
UNet 架构
UNet 是扩散模型中常用的架构,用于预测噪声。以下是简化的 UNet 实现:
import torch.nn as nn
class UNet(nn.Module):
def __init__(self, dim):
super().__init__()
self.dim = dim
# 定义编码器和解码器层
self.encoder = nn.Sequential(nn.Conv2d(3, dim, 3, padding=1),
nn.ReLU(),
nn.Conv2d(dim, dim, 3, padding=1),
nn.ReLU())
self.decoder = nn.Sequential(nn.Conv2d(dim, dim, 3, padding=1),
nn.ReLU(),
nn.Conv2d(dim, 3, 3, padding=1)
)
def forward(self, x, t):
# 添加时间嵌入
h = self.encoder(x)
return self.decoder(h)
训练循环
训练循环的核心是计算噪声预测的损失:
def train_step(model, x0, noise_scheduler, optimizer):
# 随机选择时间步
t = torch.randint(0, noise_scheduler.timesteps, (x0.size(0),))
# 添加噪声
noise = torch.randn_like(x0)
xt = noise_scheduler.q_sample(x0, t, noise)
# 预测噪声
predicted_noise = model(xt, t)
# 计算损失
loss = nn.functional.mse_loss(predicted_noise, noise)
# 反向传播
optimizer.zero_grad()
loss.backward()
optimizer.step()
return loss.item()
优化技巧
训练稳定性
- 使用学习率调度器,如
torch.optim.lr_scheduler.CosineAnnealingLR - 应用梯度裁剪,防止梯度爆炸
采样加速
- 使用 DDIM(Denoising Diffusion Implicit Models)减少采样步数
- 应用知识蒸馏训练一个小型模型来近似原始模型
避坑指南
常见训练失败模式
- 损失不下降:检查噪声调度器和学习率设置
- 生成质量差:可能需要增加模型容量或调整训练步数
计算资源规划
- 扩散模型训练通常需要多 GPU 并行
- 采样阶段可能需要大量内存,建议使用梯度检查点技术
总结
扩散模型是一种强大的生成模型,能够生成高质量的样本。通过本文的数学推导和 PyTorch 实现,读者可以深入理解其工作原理,并快速应用到实际项目中。完整的 Colab Notebook 可以在 这里 找到。
