前言
扩散模型(Diffusion Models)是当前图像生成领域最先进的技术之一,Stable Diffusion、DALL-E 3等均基于此。本教程将带你从零搭建一个简单的扩散模型,并训练它生成手写数字(MNIST数据集)。
扩散模型原理
前向过程
逐步向图像添加高斯噪声,直到变成纯噪声。数学上,给定原始图像 (x_0),经过 (T) 步噪声添加得到 (x_T):
[ q(x_t | x_{t-1}) = \mathcal{N}(x_t; \sqrt{1-\beta_t} x_{t-1}, \beta_t I) ]
逆向过程
学习一个神经网络 (\epsilon_\theta) 来预测噪声,从而逐步去噪:
[ p_\theta(x_{t-1} | x_t) = \mathcal{N}(x_{t-1}; \mu_\theta(x_t, t), \Sigma_\theta(x_t, t)) ]
环境准备
pip install torch torchvision matplotlib tqdm
代码实现
1. 定义噪声调度器
import torch
import torch.nn as nn
class NoiseScheduler:
def __init__(self, T=1000, beta_start=1e-4, beta_end=0.02):
self.T = T
self.betas = torch.linspace(beta_start, beta_end, T)
self.alphas = 1 - self.betas
self.alpha_bars = torch.cumprod(self.alphas, dim=0)
2. 构建UNet模型
扩散模型通常使用UNet架构,包含下采样和上采样路径,并加入时间嵌入:
class UNet(nn.Module):
def __init__(self):
super().__init__()
# 时间嵌入
self.time_embed = nn.Sequential(
nn.Linear(1, 128),
nn.ReLU(),
nn.Linear(128, 128)
)
# 下采样
self.down1 = nn.Conv2d(1, 64, 3, padding=1)
self.down2 = nn.Conv2d(64, 128, 3, stride=2, padding=1)
# 中间
self.mid = nn.Conv2d(128, 128, 3, padding=1)
# 上采样
self.up2 = nn.ConvTranspose2d(128, 64, 4, stride=2, padding=1)
self.up1 = nn.Conv2d(64, 1, 3, padding=1)
# 时间条件投影
self.time_proj = nn.Linear(128, 128)
def forward(self, x, t):
t_emb = self.time_embed(t.unsqueeze(1))
# 下采样
d1 = self.down1(x)
d2 = self.down2(torch.relu(d1))
# 中间
m = self.mid(torch.relu(d2))
m = m + self.time_proj(t_emb).view(-1, 128, 1, 1)
# 上采样
u2 = self.up2(torch.relu(m))
u1 = self.up1(torch.relu(u2 + d1))
return u1
3. 训练循环
def train(model, dataloader, scheduler, epochs=10):
optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)
for epoch in range(epochs):
for batch in dataloader:
x0 = batch[0]
t = torch.randint(0, scheduler.T, (x0.size(0),))
noise = torch.randn_like(x0)
xt = scheduler.add_noise(x0, noise, t)
pred_noise = model(xt, t)
loss = nn.MSELoss()(pred_noise, noise)
optimizer.zero_grad()
loss.backward()
optimizer.step()
print(f"Epoch {epoch+1}, Loss: {loss.item():.4f}")
4. 采样生成
def sample(model, scheduler, num_images=16):
model.eval()
x = torch.randn(num_images, 1, 28, 28)
for t in reversed(range(scheduler.T)):
t_tensor = torch.full((num_images,), t, dtype=torch.long)
pred_noise = model(x, t_tensor)
x = scheduler.remove_noise(x, pred_noise, t)
return x
训练与结果
在MNIST上训练10个epoch后,生成的数字样本如下(示例图略)。可以看到模型能够生成清晰可辨的手写数字。
优化技巧
- 学习率调度:使用余弦退火或线性衰减
- 梯度裁剪:防止梯度爆炸
- EMA:指数移动平均模型参数
- 混合精度训练:使用AMP加速
扩展方向
- 使用更大的数据集(如CIFAR-10、ImageNet)
- 加入文本条件(如CLIP嵌入)
- 尝试潜在扩散模型(LDM)提高效率
结语
通过本教程,你已掌握扩散模型的核心实现。虽然这是一个简化版本,但原理与工业级模型一致。继续探索,你也能创建自己的图像生成应用!