0基础学扩散模型:从数学原理到PyTorch实战

1次阅读
没有评论

共计 1640 个字符,预计需要花费 5 分钟才能阅读完成。

image.webp

扩散模型的核心思想

扩散模型的核心可以用两个过程概括:

0 基础学扩散模型:从数学原理到 PyTorch 实战

  1. 前向扩散 :通过逐步添加噪声将数据(如图片)变成随机噪声
  2. 反向去噪 :训练神经网络学习如何逆转这个过程

这其实就像把一杯清水慢慢滴入墨水(前向过程),再尝试用魔法把墨水重新分离出来(反向过程)。用数学语言来说,前向过程是一个马尔可夫链,每一步只与上一步有关:

$$q(x_t|x_{t-1}) = \mathcal{N}(x_t; \sqrt{1-\beta_t}x_{t-1}, \beta_t\mathbf{I})$$

前向过程数学详解

关键的超参数是噪声调度表 β_t(beta schedule),它控制着每一步添加多少噪声:

  • 当 β_t 很小时,前向过程需要很多步才能将数据变成噪声
  • 当 β_t 很大时,几步就能破坏数据

实践中常用线性调度或余弦调度:

# 线性噪声调度示例
def linear_beta_schedule(timesteps):
    beta_start = 0.0001
    beta_end = 0.02
    return torch.linspace(beta_start, beta_end, timesteps)

UNet 实现细节

反向过程的核心是一个 UNet,它需要:

  1. 处理不同时间步的输入
  2. 保持输出的尺寸与输入一致
  3. 使用残差连接避免梯度消失

关键实现技巧:

  • 使用正弦位置编码嵌入时间步信息
  • 在中间层添加自注意力机制
  • 使用 GroupNorm 代替 BatchNorm

完整训练代码

以下是 PyTorch Lightning 的训练框架核心部分:

class DiffusionModel(pl.LightningModule):
    def __init__(self, in_channels=3, model_channels=64):
        super().__init__()
        self.model = UNet(in_channels, model_channels)
        self.beta_schedule = linear_beta_schedule(1000)

    def forward(self, x, t):
        return self.model(x, t)

    def training_step(self, batch, batch_idx):
        x, _ = batch  # 假设 batch 来自 ImageDataset
        t = torch.randint(0, 1000, (x.shape[0],))
        noise = torch.randn_like(x)

        # 计算加噪后的图像
        sqrt_alpha = torch.sqrt(1 - self.beta_schedule[t])
        x_noisy = sqrt_alpha * x + (1 - sqrt_alpha) * noise

        # 预测噪声并计算损失
        pred_noise = self(x_noisy, t)
        loss = F.mse_loss(pred_noise, noise)
        return loss

实用调参技巧

  1. 噪声调度选择
  2. 简单任务用线性调度
  3. 高质量生成建议用余弦调度

  4. EMA 模型

    # 在 LightningModule 中添加
    def configure_optimizers(self):
        opt = torch.optim.Adam(self.parameters(), lr=1e-4)
        ema = EMA(self.model, decay=0.9999)  # 实现 EMA 类
        return [opt], [ema]

常见问题解决

  • 显存不足
  • 减小 batch size
  • 使用梯度累积
  • 尝试混合精度训练

  • 梯度爆炸

  • 添加梯度裁剪
  • 检查损失函数数值稳定性

  • FID 计算

  • 确保使用相同的数据预处理
  • 推荐使用 50000 个样本计算

延伸学习

  1. 原始论文:Denoising Diffusion Probabilistic Models
  2. 开源实现:GitHub 搜索 ”ddpm-pytorch”
  3. 进阶阅读:Score-Based Generative Modeling

通过这个框架,我在 CIFAR-10 上训练出了第一版可用的扩散模型,虽然生成效果还不够完美,但整个过程让我对生成模型有了更深的理解。建议读者先从小型数据集开始实验,逐步调整模型复杂度。

正文完
 0
评论(没有评论)