用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上跑了整整两天,但当看到第一张由噪声逐渐"浮现"出的人脸时,那种成就感绝对值得等待。

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐