Skip to content

扩散模型理论基础论文

扩散模型通过模拟前向加噪和逆向去噪过程,实现了高质量的图像生成,成为当前最重要的生成模型框架。本文梳理扩散模型的理论基础——从 DDPM 到分数匹配再到统一框架。

扩散模型理论

DDPM 基础

Ho et al. (2020) 的 DDPM 奠定了扩散模型的基础:

前向过程

逐步向数据添加高斯噪声:

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

任意步的直接加噪:

$$q(x_t | x_0) = \mathcal{N}(x_t; \sqrt{\bar{\alpha}_t} x_0, (1-\bar{\alpha}_t) I)$$

逆向过程

训练神经网络预测噪声,逐步去噪:

python
import torch
import torch.nn as nn

class DiffusionModel(nn.Module):
    """简化的 DDPM 扩散模型"""
    def __init__(self, unet, timesteps=1000, beta_start=1e-4, beta_end=0.02):
        super().__init__()
        self.unet = unet
        self.timesteps = timesteps

        # 线性噪声调度
        self.register_buffer(
            'betas', torch.linspace(beta_start, beta_end, timesteps)
        )
        alphas = 1.0 - self.betas
        self.register_buffer('alpha_bars', torch.cumprod(alphas, dim=0))

    def forward_process(self, x_0, t, noise=None):
        """前向加噪"""
        if noise is None:
            noise = torch.randn_like(x_0)
        alpha_bar = self.alpha_bars[t][:, None, None, None]
        return torch.sqrt(alpha_bar) * x_0 + torch.sqrt(1 - alpha_bar) * noise

    def training_loss(self, x_0):
        """训练损失:预测噪声"""
        t = torch.randint(0, self.timesteps, (x_0.shape[0],), device=x_0.device)
        noise = torch.randn_like(x_0)
        x_t = self.forward_process(x_0, t, noise)
        predicted_noise = self.unet(x_t, t)
        return nn.functional.mse_loss(predicted_noise, noise)

    @torch.no_grad()
    def sample(self, shape, device):
        """DDPM 采样"""
        x = torch.randn(shape, device=device)
        for t in reversed(range(self.timesteps)):
            t_batch = torch.full((shape[0],), t, device=device, dtype=torch.long)
            predicted_noise = self.unet(x, t_batch)

            alpha = 1.0 - self.betas[t]
            alpha_bar = self.alpha_bars[t]

            # 去噪一步
            x = (x - (1 - alpha) / torch.sqrt(1 - alpha_bar) * predicted_noise) / torch.sqrt(alpha)

            if t > 0:
                x += torch.sqrt(self.betas[t]) * torch.randn_like(x)
        return x

分数匹配视角

Song & Ermon (2019) 从分数匹配角度提供了另一种理解:

分数函数

分数函数是数据分布对数的梯度:$\nabla_x \log p(x)$

训练目标是学习分数函数:

$$\mathbb{E}{p(x)}\left[|\nabla_x \log p(x) - s\theta(x)|^2\right]$$

与 DDPM 的联系

  • DDPM 预测噪声 ε 等价于预测分数函数
  • 朗之万动力学采样等价于 DDPM 逆向过程
  • 两者是同一理论框架的不同参数化

统一视角

Score-Based Generative Model(Song et al., 2021)统一了 DDPM 和分数匹配:DDPM 的离散扩散是随机微分方程(SDE)的离散化,分数匹配学习的是 SDE 的漂移项。这一统一框架使得连续时间扩散模型成为可能。

重要理论结果

DDIM 加速采样

Song et al. (2020) 提出的 DDIM 将采样步数从 1000 减少到 50:

  • 将随机逆向过程替换为确定性映射
  • 保持生成质量的同时大幅加速
  • 支持在隐空间插值

最优传输

扩散模型可以解释为最优传输问题:

  • 前向过程是从数据分布到噪声分布的传输
  • 逆向过程学习最优传输的逆映射
  • Flow Matching 提供了更直接的最优传输路径

相关资源

最近更新