AI百事通

从零搭建AI图像生成器:基于扩散模型的实战教程

📅 2026-07-15📰 ai_generated👁 2 次阅读
从零搭建AI图像生成器:基于扩散模型的实战教程
扩散模型图像生成深度学习PyTorch教程

前言

扩散模型(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后,生成的数字样本如下(示例图略)。可以看到模型能够生成清晰可辨的手写数字。

优化技巧

  1. 学习率调度:使用余弦退火或线性衰减
  2. 梯度裁剪:防止梯度爆炸
  3. EMA:指数移动平均模型参数
  4. 混合精度训练:使用AMP加速

扩展方向

  • 使用更大的数据集(如CIFAR-10、ImageNet)
  • 加入文本条件(如CLIP嵌入)
  • 尝试潜在扩散模型(LDM)提高效率

结语

通过本教程,你已掌握扩散模型的核心实现。虽然这是一个简化版本,但原理与工业级模型一致。继续探索,你也能创建自己的图像生成应用!