别再怕数学!用PyTorch手把手实现DDPM,从加噪到生成图片的完整代码解读
·
用PyTorch实战DDPM:零数学基础也能玩转图像生成
在咖啡馆里,我遇到一位刚入行AI的开发者小张。他盯着屏幕上密密麻麻的扩散模型公式,眉头紧锁:"这些数学推导看得我头大,但真的好想亲手实现一个能画画的AI..." 这场景让我想起三年前的自己。当时我发现, 真正理解AI模型的最佳方式不是死磕论文,而是动手实现它 。今天我们就用PyTorch,从零开始构建一个Denoising Diffusion Probabilistic Model(DDPM),完全避开数学公式的"恐吓",用代码对话这个神奇的图像生成模型。
1. 环境准备与数据加载
1.1 最小化依赖配置
我们只需要最基础的PyTorch生态工具包,创建一个干净的虚拟环境:
conda create -n ddpm python=3.8
conda activate ddpm
pip install torch torchvision matplotlib tqdm
1.2 数据加载策略
使用CIFAR-10作为示例数据集,但代码结构支持任意图像数据集。这里采用PyTorch的DataLoader实现高效加载:
from torchvision import datasets, transforms
def get_dataloader(batch_size=64):
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])
dataset = datasets.CIFAR10(
root='./data',
train=True,
download=True,
transform=transform
)
return DataLoader(dataset, batch_size=batch_size, shuffle=True)
提示:将图像像素值归一化到[-1,1]区间对扩散模型训练至关重要
2. 扩散过程的核心实现
2.1 噪声调度器设计
扩散模型的核心在于精心设计的噪声添加策略。我们实现一个线性调度器:
import torch
class NoiseScheduler:
def __init__(self, timesteps=1000, beta_start=1e-4, beta_end=0.02):
self.timesteps = timesteps
self.betas = torch.linspace(beta_start, beta_end, timesteps)
self.alphas = 1. - self.betas
self.alpha_bars = torch.cumprod(self.alphas, dim=0)
def add_noise(self, x_0, t):
noise = torch.randn_like(x_0)
sqrt_alpha_bar = torch.sqrt(self.alpha_bars[t])[:, None, None, None]
sqrt_one_minus_alpha_bar = torch.sqrt(1. - self.alpha_bars[t])[:, None, None, None]
return sqrt_alpha_bar * x_0 + sqrt_one_minus_alpha_bar * noise, noise
关键参数说明:
| 参数 | 作用 | 典型值 |
|---|---|---|
| timesteps | 扩散步数 | 1000 |
| beta_start | 初始噪声系数 | 1e-4 |
| beta_end | 最终噪声系数 | 0.02 |
2.2 可视化噪声添加过程
让我们观察图像如何逐步变成噪声:
import matplotlib.pyplot as plt
def plot_noising(scheduler, sample_img, steps=5):
plt.figure(figsize=(15,3))
for i in range(steps):
t = torch.tensor([i * (scheduler.timesteps//steps)])
noised, _ = scheduler.add_noise(sample_img, t)
plt.subplot(1, steps+1, i+1)
plt.imshow(noised[0].permute(1,2,0)*0.5+0.5)
plt.title(f't={t.item()}')
plt.show()
3. UNet噪声预测模型
3.1 构建DDPM专用UNet
不同于传统UNet,我们需要加入时间步嵌入:
from torch import nn
import math
class TimeEmbedding(nn.Module):
def __init__(self, dim):
super().__init__()
self.dim = dim
half_dim = dim // 2
emb = math.log(10000) / (half_dim - 1)
emb = torch.exp(torch.arange(half_dim, dtype=torch.float) * -emb)
self.register_buffer('emb', emb)
def forward(self, t):
emb = t.float()[:, None] * self.emb[None, :]
return torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1)
class UNetBlock(nn.Module):
def __init__(self, in_c, out_c, time_emb_dim):
super().__init__()
self.time_mlp = nn.Linear(time_emb_dim, out_c)
self.conv = nn.Sequential(
nn.Conv2d(in_c, out_c, 3, padding=1),
nn.BatchNorm2d(out_c),
nn.ReLU(),
nn.Conv2d(out_c, out_c, 3, padding=1),
nn.BatchNorm2d(out_c),
nn.ReLU()
)
def forward(self, x, t):
h = self.conv(x)
time_emb = self.time_mlp(t)[:,:,None,None]
return h + time_emb
3.2 完整UNet架构
实现一个对称的编码器-解码器结构:
class DDPM_UNet(nn.Module):
def __init__(self, channels=3, base_dim=64):
super().__init__()
time_dim = base_dim * 4
self.time_embed = TimeEmbedding(time_dim)
# 下采样
self.down1 = UNetBlock(channels, base_dim, time_dim)
self.down2 = UNetBlock(base_dim, base_dim*2, time_dim)
self.down3 = UNetBlock(base_dim*2, base_dim*4, time_dim)
# 瓶颈层
self.bottleneck = nn.Sequential(
nn.Conv2d(base_dim*4, base_dim*8, 3, padding=1),
nn.ReLU()
)
# 上采样
self.up1 = UNetBlock(base_dim*8 + base_dim*4, base_dim*4, time_dim)
self.up2 = UNetBlock(base_dim*4 + base_dim*2, base_dim*2, time_dim)
self.up3 = UNetBlock(base_dim*2 + base_dim, base_dim, time_dim)
self.final = nn.Conv2d(base_dim, channels, 1)
def forward(self, x, t):
t = self.time_embed(t)
# 编码器路径
d1 = self.down1(x, t)
d2 = self.down2(nn.MaxPool2d(2)(d1), t)
d3 = self.down3(nn.MaxPool2d(2)(d2), t)
# 瓶颈
bottleneck = self.bottleneck(nn.MaxPool2d(2)(d3))
# 解码器路径
u1 = self.up1(torch.cat([nn.Upsample(scale_factor=2)(bottleneck), d3], 1), t)
u2 = self.up2(torch.cat([nn.Upsample(scale_factor=2)(u1), d2], 1), t)
u3 = self.up3(torch.cat([nn.Upsample(scale_factor=2)(u2), d1], 1), t)
return self.final(u3)
4. 训练流程实现
4.1 损失函数与优化器
DDPM使用简单的MSE损失:
def train_step(model, scheduler, x_0, optimizer):
optimizer.zero_grad()
# 随机采样时间步
t = torch.randint(0, scheduler.timesteps, (x_0.shape[0],))
# 添加噪声
noised, noise = scheduler.add_noise(x_0, t)
# 预测噪声
pred_noise = model(noised, t)
# 计算损失
loss = nn.functional.mse_loss(pred_noise, noise)
loss.backward()
optimizer.step()
return loss.item()
4.2 训练循环封装
完整的训练流程封装:
from tqdm import tqdm
def train(model, dataloader, scheduler, epochs=50, lr=1e-3):
optimizer = torch.optim.Adam(model.parameters(), lr=lr)
model.train()
for epoch in range(epochs):
total_loss = 0
pbar = tqdm(dataloader)
for x, _ in pbar:
x = x.to(device)
loss = train_step(model, scheduler, x, optimizer)
total_loss += loss
pbar.set_description(f"Epoch {epoch+1} Loss: {loss:.4f}")
print(f"Epoch {epoch+1} Avg Loss: {total_loss/len(dataloader):.4f}")
# 每个epoch保存模型和生成样本
if (epoch+1) % 5 == 0:
torch.save(model.state_dict(), f"ddpm_epoch{epoch+1}.pth")
generate_samples(model, scheduler)
5. 图像生成与采样
5.1 反向扩散采样算法
从纯噪声逐步去噪生成图像:
@torch.no_grad()
def sample(model, scheduler, num_samples=16):
model.eval()
img_size = (num_samples, 3, 32, 32)
x_t = torch.randn(img_size).to(device)
for t in reversed(range(scheduler.timesteps)):
t_batch = torch.full((num_samples,), t, device=device)
pred_noise = model(x_t, t_batch)
alpha_t = scheduler.alphas[t]
alpha_bar_t = scheduler.alpha_bars[t]
beta_t = scheduler.betas[t]
if t > 0:
noise = torch.randn_like(x_t)
else:
noise = torch.zeros_like(x_t)
x_t = (1/torch.sqrt(alpha_t)) * (
x_t - ((1-alpha_t)/torch.sqrt(1-alpha_bar_t)) * pred_noise
) + torch.sqrt(beta_t) * noise
return torch.clamp(x_t, -1., 1.)
5.2 生成效果可视化
将生成的图像网格化显示:
def generate_samples(model, scheduler, n=16):
samples = sample(model, scheduler, n)
plt.figure(figsize=(10,10))
for i in range(n):
plt.subplot(4,4,i+1)
plt.imshow(samples[i].cpu().permute(1,2,0)*0.5+0.5)
plt.axis('off')
plt.tight_layout()
plt.savefig(f"generated_samples.png")
plt.show()
6. 模型优化技巧
6.1 学习率调度
使用余弦退火优化训练过程:
from torch.optim.lr_scheduler import CosineAnnealingLR
optimizer = torch.optim.Adam(model.parameters(), lr=2e-4)
scheduler = CosineAnnealingLR(optimizer, T_max=len(dataloader)*epochs)
6.2 混合精度训练
大幅减少显存占用并加速训练:
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
def train_step_amp(model, scheduler, x_0, optimizer):
optimizer.zero_grad()
t = torch.randint(0, scheduler.timesteps, (x_0.shape[0],))
with autocast():
noised, noise = scheduler.add_noise(x_0, t)
pred_noise = model(noised, t)
loss = nn.functional.mse_loss(pred_noise, noise)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
return loss.item()
6.3 模型检查点
实现训练中断恢复功能:
def save_checkpoint(model, optimizer, epoch, path):
torch.save({
'epoch': epoch,
'model_state': model.state_dict(),
'optimizer_state': optimizer.state_dict(),
}, path)
def load_checkpoint(model, optimizer, path):
checkpoint = torch.load(path)
model.load_state_dict(checkpoint['model_state'])
optimizer.load_state_dict(checkpoint['optimizer_state'])
return checkpoint['epoch']
7. 进阶改进方向
7.1 条件生成实现
添加类别标签实现可控生成:
class ConditionalUNet(DDPM_UNet):
def __init__(self, num_classes, **kwargs):
super().__init__(**kwargs)
self.label_emb = nn.Embedding(num_classes, self.time_embed.dim)
def forward(self, x, t, labels):
t_emb = self.time_embed(t)
label_emb = self.label_emb(labels)
t_emb += label_emb
return super().forward(x, t_emb)
7.2 扩散步数动态调整
实现自适应步长采样:
def dynamic_sampling(model, scheduler, num_steps=50):
model.eval()
x_t = torch.randn_like(x_0)
steps = torch.linspace(0, scheduler.timesteps-1, num_steps).long()
for t in reversed(steps):
t_batch = torch.full((x_t.shape[0],), t.item())
pred_noise = model(x_t, t_batch)
alpha_t = scheduler.alphas[t]
alpha_bar_t = scheduler.alpha_bars[t]
beta_t = scheduler.betas[t]
if t > 0:
noise = torch.randn_like(x_t)
else:
noise = torch.zeros_like(x_t)
x_t = (1/torch.sqrt(alpha_t)) * (
x_t - ((1-alpha_t)/torch.sqrt(1-alpha_bar_t)) * pred_noise
) + torch.sqrt(beta_t) * noise
return x_t
7.3 多尺度训练策略
提升高分辨率图像生成质量:
def multi_scale_training(model, dataloader, scales=[32, 64, 128]):
for x, _ in dataloader:
scale = random.choice(scales)
x_resized = F.interpolate(x, size=(scale, scale), mode='bilinear')
# 其余训练步骤相同...
在完成基础实现后,我强烈建议尝试在CelebA或LSUN卧室数据集上训练模型。记得准备好足够的GPU资源——我的第一次训练在Colab上跑了整整两天,但当看到第一张由噪声逐渐"浮现"出的人脸时,那种成就感绝对值得等待。
更多推荐




所有评论(0)