从DDPM到Stable Diffusion:扩散模型数学推导与代码实现全解析
2026/9/19 17:58:24 网站建设 项目流程

扩散模型这两年火得一塌糊涂,但真正沉下心把它的数学推导和代码实现从头撸一遍的人其实不多。我最初接触DDPM的时候,看论文觉得道理挺简单——加噪、去噪、学一个反向过程,但真到动手复现的时候,前向过程的噪声调度怎么设计、反向过程的均值方差怎么推、Unet里的时间步嵌入到底怎么加,这些问题一个接一个地冒出来。这篇内容就是把我自己从看论文到跑通代码的完整过程整理出来,既讲清楚数学上的来龙去脉,也给出可以直接跑的代码实现。不管你是刚入门想搞明白扩散模型到底怎么回事,还是已经用过Stable Diffusion想深入理解底层机制,应该都能从中找到有用的东西。

1. 扩散模型到底在解决什么问题

1.1 从生成模型的大背景说起

生成模型的核心任务就一件事:给定一批训练数据,学出一个模型,让它能生成跟训练数据相似的新样本。听起来简单,但怎么做这件事有很多种思路。GAN走的是对抗博弈的路子,生成器和判别器互相较劲;VAE走的是隐变量加变分推断的路子;而扩散模型走的是另一条路——它把生成过程拆成很多个很小的去噪步骤,每一步只做一点点修正,最终从纯噪声里“雕刻”出一张清晰的图。

这个思路其实很符合直觉。你想象一块大理石,雕塑家不是一刀就雕出成品,而是一点点去掉多余的部分。扩散模型也是这样,从一团完全随机的噪声开始,每一步去掉一点点噪声,经过几百上千步之后,剩下的就是一张有意义的图像。

我第一次看到这个思路的时候觉得挺笨的——为什么要分那么多步?一步到位不行吗?后来想明白了,一步到位就是GAN在做的事,但GAN的训练极不稳定,容易模式崩塌。扩散模型用很多小步换来了训练的稳定性和生成质量的上限。这就像爬山,一步登天很难,但分成一千个小台阶,每步只走一点点,就容易多了。

1.2 扩散模型的核心直觉

扩散模型分两个过程:前向过程和反向过程。

前向过程也叫扩散过程,就是不断往一张清晰图像上加高斯噪声,每一步加一点点,经过T步之后,图像就变成纯噪声了。这个过程是固定的,不需要学习,就是一个马尔可夫链。

反向过程也叫去噪过程,就是从纯噪声出发,一步步去掉噪声,最终恢复出清晰图像。这个过程是需要学习的,我们要训练一个神经网络来预测每一步应该去掉多少噪声。

关键洞察在于:如果我们能学会反向过程,那就可以从纯噪声开始,一步步去噪,生成全新的图像。这就是扩散模型生成样本的方式。

注意:前向过程是固定的加噪过程,不涉及任何学习;反向过程才是需要训练的部分。很多人初学时会搞混这两个过程的关系。

1.3 为什么扩散模型能work

从数学上看,扩散模型的训练目标可以推导出一个非常简洁的形式:预测加入的噪声。这个推导过程涉及变分下界(ELBO)的分解和重参数化技巧,后面会详细展开。

从直觉上看,扩散模型之所以能work,是因为它把“生成一张图”这个极其复杂的分布建模问题,拆解成了很多个“去一点噪声”的简单问题。每个简单问题用一个神经网络来学,学起来容易得多。而且因为每一步的变化很小,所以可以用高斯分布来近似每一步的反向过程,这让数学处理变得可行。

另一个重要原因是,扩散模型的训练目标本质上是一个去噪自编码器的变体。去噪自编码器本身就是一种很有效的表示学习方法,扩散模型把这个思想推到了极致——不是去一个固定程度的噪声,而是在所有噪声水平上都学会去噪。

2. 前向过程:从清晰图像到纯噪声

2.1 前向过程的数学定义

前向过程定义一个马尔可夫链,从真实数据 $x_0$ 出发,逐步加噪:

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

其中 $\beta_t$ 是第 $t$ 步的噪声方差,通常从 $\beta_1=10^{-4}$ 线性增加到 $\beta_T=0.02$。这个 $\beta_t$ 的调度很关键,后面会详细讨论。

这个式子的意思是:每一步的新图像 $x_t$ 是上一步图像 $x_{t-1}$ 乘以一个缩放因子 $\sqrt{1-\beta_t}$,再加上方差为 $\beta_t$ 的高斯噪声。缩放因子保证图像的方差不会爆炸。

2.2 重参数化:一步到位加噪

虽然前向过程是逐步的,但有一个非常重要的性质:我们可以直接从 $x_0$ 采样出任意时刻的 $x_t$,不需要一步步迭代。这个性质叫重参数化。

令 $\alpha_t = 1 - \beta_t$,$\bar{\alpha}t = \prod{s=1}^{t} \alpha_s$,则有:

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

也就是说:

$$x_t = \sqrt{\bar{\alpha}_t} x_0 + \sqrt{1-\bar{\alpha}_t} \epsilon, \quad \epsilon \sim \mathcal{N}(0, \mathbf{I})$$

这个公式极其重要,是整个扩散模型训练的基础。它告诉我们:任意时刻的 $x_t$ 都可以写成原始图像 $x_0$ 和一个标准高斯噪声 $\epsilon$ 的线性组合。系数 $\sqrt{\bar{\alpha}_t}$ 和 $\sqrt{1-\bar{\alpha}_t}$ 是固定的,只取决于时间步 $t$。

推导过程用到了高斯分布的可加性:两个独立高斯分布的和仍然是高斯分布,均值和方差分别相加。具体推导如下:

$$ \begin{aligned} x_t &= \sqrt{\alpha_t} x_{t-1} + \sqrt{1-\alpha_t} \epsilon_{t-1} \ &= \sqrt{\alpha_t}(\sqrt{\alpha_{t-1}} x_{t-2} + \sqrt{1-\alpha_{t-1}} \epsilon_{t-2}) + \sqrt{1-\alpha_t} \epsilon_{t-1} \ &= \sqrt{\alpha_t \alpha_{t-1}} x_{t-2} + \sqrt{\alpha_t(1-\alpha_{t-1})} \epsilon_{t-2} + \sqrt{1-\alpha_t} \epsilon_{t-1} \end{aligned} $$

后面两项都是独立高斯噪声,合并后方差为 $\alpha_t(1-\alpha_{t-1}) + (1-\alpha_t) = 1 - \alpha_t \alpha_{t-1}$。依此类推,最终得到 $x_t = \sqrt{\bar{\alpha}_t} x_0 + \sqrt{1-\bar{\alpha}_t} \epsilon$。

2.3 噪声调度的选择

$\beta_t$ 的调度直接影响模型质量。DDPM原论文用的是线性调度,从 $10^{-4}$ 到 $0.02$,T=1000。但后来的工作发现,线性调度在低分辨率图像上还行,在高分辨率图像上会丢失太多信息。

改进方案是余弦调度:

$$\bar{\alpha}_t = \frac{f(t)}{f(0)}, \quad f(t) = \cos\left(\frac{t/T + s}{1+s} \cdot \frac{\pi}{2}\right)^2$$

其中 $s$ 是一个小的偏移量,通常取0.008。余弦调度的好处是在中间时刻加噪速度更均匀,不会在早期就把图像信息破坏得太厉害。

我实测下来,对于64x64以下的小图,线性调度够用;对于256x256以上的图,余弦调度明显更好。这个选择不是玄学,背后有信息论的解释——余弦调度让信噪比在时间轴上分布更均匀。

3. 反向过程:从噪声恢复图像

3.1 反向过程的数学形式

反向过程也是马尔可夫链,从 $x_T \sim \mathcal{N}(0, \mathbf{I})$ 出发,逐步去噪:

$$p_\theta(x_{t-1} | x_t) = \mathcal{N}(x_{t-1}; \mu_\theta(x_t, t), \Sigma_\theta(x_t, t))$$

当 $\beta_t$ 足够小时,反向过程也可以用高斯分布来近似。这是扩散模型能work的关键假设之一。

DDPM的做法是固定方差 $\Sigma_\theta(x_t, t) = \sigma_t^2 \mathbf{I}$,其中 $\sigma_t^2 = \beta_t$ 或 $\sigma_t^2 = \tilde{\beta}t = \frac{1-\bar{\alpha}{t-1}}{1-\bar{\alpha}t} \beta_t$。然后只学习均值 $\mu\theta(x_t, t)$。

3.2 均值的推导

通过贝叶斯公式,后验 $q(x_{t-1} | x_t, x_0)$ 是可以精确计算的:

$$q(x_{t-1} | x_t, x_0) = \mathcal{N}(x_{t-1}; \tilde{\mu}_t(x_t, x_0), \tilde{\beta}_t \mathbf{I})$$

其中:

$$\tilde{\mu}t(x_t, x_0) = \frac{\sqrt{\bar{\alpha}{t-1}} \beta_t}{1-\bar{\alpha}t} x_0 + \frac{\sqrt{\alpha_t}(1-\bar{\alpha}{t-1})}{1-\bar{\alpha}_t} x_t$$

$$\tilde{\beta}t = \frac{1-\bar{\alpha}{t-1}}{1-\bar{\alpha}_t} \beta_t$$

这个后验均值是反向过程均值的“目标”。如果我们知道 $x_0$,就能算出最优的反向均值。

3.3 从预测 $x_0$ 到预测噪声

实际训练时,我们不是直接预测 $x_0$,而是预测噪声 $\epsilon$。原因在于:通过重参数化公式 $x_t = \sqrt{\bar{\alpha}_t} x_0 + \sqrt{1-\bar{\alpha}_t} \epsilon$,我们可以把 $x_0$ 表示为:

$$x_0 = \frac{x_t - \sqrt{1-\bar{\alpha}_t} \epsilon}{\sqrt{\bar{\alpha}_t}}$$

代入后验均值公式,得到:

$$\tilde{\mu}_t = \frac{1}{\sqrt{\alpha_t}} \left( x_t - \frac{\beta_t}{\sqrt{1-\bar{\alpha}_t}} \epsilon \right)$$

所以如果我们训练一个网络 $\epsilon_\theta(x_t, t)$ 来预测 $\epsilon$,那么反向均值就是:

$$\mu_\theta(x_t, t) = \frac{1}{\sqrt{\alpha_t}} \left( x_t - \frac{\beta_t}{\sqrt{1-\bar{\alpha}t}} \epsilon\theta(x_t, t) \right)$$

这就是DDPM的核心公式。训练目标也简化为:

$$L_{\text{simple}} = \mathbb{E}{t, x_0, \epsilon} \left[ | \epsilon - \epsilon\theta(\sqrt{\bar{\alpha}_t} x_0 + \sqrt{1-\bar{\alpha}_t} \epsilon, t) |^2 \right]$$

这个损失函数极其简洁:随机采样一个时间步 $t$,随机采样一个噪声 $\epsilon$,构造 $x_t$,让网络预测 $\epsilon$,然后算MSE。就这么简单。

提示:虽然理论上应该用变分下界(ELBO)作为损失,但DDPM发现简化后的MSE损失效果更好。这个“简化”去掉了ELBO中与时间步相关的权重项,让每个时间步的损失权重相同。

3.4 采样过程

训练好网络后,采样过程就是从 $x_T \sim \mathcal{N}(0, \mathbf{I})$ 出发,对 $t = T, T-1, \ldots, 1$ 迭代:

$$x_{t-1} = \frac{1}{\sqrt{\alpha_t}} \left( x_t - \frac{\beta_t}{\sqrt{1-\bar{\alpha}t}} \epsilon\theta(x_t, t) \right) + \sigma_t z$$

其中 $z \sim \mathcal{N}(0, \mathbf{I})$ 当 $t > 1$,当 $t=1$ 时 $z=0$。$\sigma_t^2 = \tilde{\beta}_t$。

这个采样过程需要迭代T次,T通常取1000,所以生成一张图需要跑1000次网络前向。这也是扩散模型生成速度慢的根本原因。后来的DDIM、DPM-Solver等工作就是想办法减少采样步数。

4. Unet架构:噪声预测网络的设计

4.1 为什么选Unet

噪声预测网络需要输入一张带噪图像 $x_t$ 和时间步 $t$,输出预测的噪声 $\epsilon_\theta(x_t, t)$。输入和输出都是同尺寸的图像,这天然适合编码器-解码器结构。

Unet最初是为医学图像分割设计的,它的核心特点是跳跃连接:编码器每一层的特征图直接拼接到解码器对应层。这样既能捕获全局语义信息(通过下采样),又能保留精细的空间细节(通过跳跃连接)。

对于扩散模型来说,跳跃连接尤其重要。因为去噪任务需要同时考虑全局结构和局部细节——既要理解整张图的内容,又要精确地恢复每个像素。Unet的跳跃连接正好满足这个需求。

4.2 Unet的基本结构

DDPM用的Unet结构大致如下:

  • 输入层:一个卷积层,把输入图像的通道数映射到基础通道数
  • 下采样阶段:多个残差块+注意力块,每个阶段后跟一个下采样操作(stride=2的卷积或池化)
  • 中间层:残差块+注意力块,不改变分辨率
  • 上采样阶段:多个残差块+注意力块,每个阶段前跟一个上采样操作(最近邻插值+卷积)
  • 输出层:一个卷积层,把通道数映射回输入图像的通道数

每个残差块包含两组GroupNorm+SiLU+卷积,以及一个跳跃连接。时间步嵌入通过一个MLP映射后,加到每个残差块中。

4.3 时间步嵌入

时间步 $t$ 是一个标量,需要嵌入到网络中。DDPM用的是正弦位置编码,跟Transformer里的位置编码类似:

$$PE(t, 2i) = \sin(t / 10000^{2i/d})$$ $$PE(t, 2i+1) = \cos(t / 10000^{2i/d})$$

其中 $d$ 是嵌入维度。这个编码后的向量再通过两层MLP,然后加到每个残差块的特征图上。

为什么用正弦编码?因为它能让网络区分不同的时间步,而且对于相邻时间步,编码向量也是相邻的,这有助于网络学习平滑的去噪过程。

4.4 注意力机制

DDPM在16x16分辨率的特征图上加了自注意力。自注意力让每个位置都能看到其他所有位置的信息,有助于捕获全局依赖。

具体实现是标准的多头自注意力:把特征图reshape成序列,做QKV投影,算注意力权重,然后加权求和。加上残差连接和LayerNorm。

后来的工作(如Stable Diffusion)在多个分辨率上都加了注意力,并且用了更高效的实现(如Flash Attention)。

4.5 代码实现:一个精简版Unet

下面是一个可以直接跑的Unet实现,我把它拆成了几个模块,方便理解:

import torch import torch.nn as nn import torch.nn.functional as F import math class SinusoidalPositionEmbedding(nn.Module): def __init__(self, dim): super().__init__() self.dim = dim def forward(self, t): device = t.device half_dim = self.dim // 2 emb = math.log(10000) / (half_dim - 1) emb = torch.exp(torch.arange(half_dim, device=device) * -emb) emb = t[:, None] * emb[None, :] emb = torch.cat([emb.sin(), emb.cos()], dim=-1) return emb class ResidualBlock(nn.Module): def __init__(self, in_channels, out_channels, time_emb_dim): super().__init__() self.norm1 = nn.GroupNorm(8, in_channels) self.conv1 = nn.Conv2d(in_channels, out_channels, 3, padding=1) self.time_mlp = nn.Linear(time_emb_dim, out_channels) self.norm2 = nn.GroupNorm(8, out_channels) self.conv2 = nn.Conv2d(out_channels, out_channels, 3, padding=1) self.residual_conv = nn.Conv2d(in_channels, out_channels, 1) \ if in_channels != out_channels else nn.Identity() def forward(self, x, t_emb): h = self.conv1(F.silu(self.norm1(x))) h = h + self.time_mlp(F.silu(t_emb))[:, :, None, None] h = self.conv2(F.silu(self.norm2(h))) return h + self.residual_conv(x) class AttentionBlock(nn.Module): def __init__(self, channels, num_heads=4): super().__init__() self.num_heads = num_heads self.norm = nn.GroupNorm(8, channels) self.qkv = nn.Conv2d(channels, channels * 3, 1) self.proj = nn.Conv2d(channels, channels, 1) def forward(self, x): B, C, H, W = x.shape h = self.norm(x) qkv = self.qkv(h).reshape(B, 3, self.num_heads, C // self.num_heads, H * W) q, k, v = qkv[:, 0], qkv[:, 1], qkv[:, 2] attn = torch.einsum('bhcn,bhcm->bhnm', q, k) / math.sqrt(C // self.num_heads) attn = F.softmax(attn, dim=-1) out = torch.einsum('bhnm,bhcm->bhcn', attn, v) out = out.reshape(B, C, H, W) return x + self.proj(out) class Unet(nn.Module): def __init__(self, in_channels=3, base_channels=64, channel_mults=(1, 2, 4, 8), time_emb_dim=256): super().__init__() self.time_mlp = nn.Sequential( SinusoidalPositionEmbedding(base_channels), nn.Linear(base_channels, time_emb_dim), nn.SiLU(), nn.Linear(time_emb_dim, time_emb_dim) ) self.init_conv = nn.Conv2d(in_channels, base_channels, 3, padding=1) self.downs = nn.ModuleList() self.ups = nn.ModuleList() channels = [base_channels * m for m in channel_mults] # 下采样 prev_ch = base_channels for i, ch in enumerate(channels): self.downs.append(nn.ModuleList([ ResidualBlock(prev_ch, ch, time_emb_dim), ResidualBlock(ch, ch, time_emb_dim), AttentionBlock(ch) if i >= 2 else nn.Identity(), nn.Conv2d(ch, ch, 3, stride=2, padding=1) if i < len(channels) - 1 else nn.Identity() ])) prev_ch = ch # 中间层 self.mid = nn.ModuleList([ ResidualBlock(channels[-1], channels[-1], time_emb_dim), AttentionBlock(channels[-1]), ResidualBlock(channels[-1], channels[-1], time_emb_dim) ]) # 上采样 for i, ch in reversed(list(enumerate(channels))): prev_ch = channels[i] skip_ch = channels[i] self.ups.append(nn.ModuleList([ ResidualBlock(prev_ch + skip_ch, ch, time_emb_dim), ResidualBlock(ch, ch, time_emb_dim), AttentionBlock(ch) if i >= 2 else nn.Identity(), nn.ConvTranspose2d(ch, ch, 4, stride=2, padding=1) if i > 0 else nn.Identity() ])) self.out = nn.Sequential( nn.GroupNorm(8, base_channels), nn.SiLU(), nn.Conv2d(base_channels, in_channels, 3, padding=1) ) def forward(self, x, t): t_emb = self.time_mlp(t) h = self.init_conv(x) skips = [] for res1, res2, attn, down in self.downs: h = res1(h, t_emb) h = res2(h, t_emb) h = attn(h) skips.append(h) h = down(h) for res1, attn, res2 in self.mid: h = res1(h, t_emb) h = attn(h) h = res2(h, t_emb) for res1, res2, attn, up in self.ups: h = torch.cat([h, skips.pop()], dim=1) h = res1(h, t_emb) h = res2(h, t_emb) h = attn(h) h = up(h) return self.out(h)

这个实现虽然精简,但包含了Unet的所有核心组件:残差块、时间步嵌入、注意力、跳跃连接。你可以直接拿它来训练DDPM。

注意:实际训练时,base_channels通常取128或256,channel_mults取(1,2,4,8)或(1,2,4,8,16)。这个精简版用64是为了方便在单卡上跑。

5. 训练与采样:完整代码实现

5.1 训练循环

训练DDPM的代码非常简洁,核心就是采样时间步、加噪、预测噪声、算MSE:

def train_step(model, x0, optimizer, device): model.train() batch_size = x0.shape[0] t = torch.randint(0, T, (batch_size,), device=device).long() noise = torch.randn_like(x0) sqrt_alpha_bar = extract(alphas_bar, t, x0.shape) sqrt_one_minus_alpha_bar = extract(1 - alphas_bar, t, x0.shape) xt = sqrt_alpha_bar * x0 + sqrt_one_minus_alpha_bar * noise pred_noise = model(xt, t) loss = F.mse_loss(pred_noise, noise) optimizer.zero_grad() loss.backward() optimizer.step() return loss.item()

其中extract函数把一维的系数数组按时间步索引扩展到跟图像同形状:

def extract(arr, t, shape): out = arr.gather(0, t) return out.reshape(t.shape[0], *([1] * (len(shape) - 1)))

5.2 采样循环

采样就是从纯噪声出发,逐步去噪:

@torch.no_grad() def sample(model, n_samples, device): model.eval() x = torch.randn(n_samples, 3, 32, 32, device=device) for t in reversed(range(T)): t_batch = torch.full((n_samples,), t, device=device, dtype=torch.long) pred_noise = model(x, t_batch) alpha_t = alphas[t] alpha_bar_t = alphas_bar[t] beta_t = betas[t] mean = (x - beta_t / torch.sqrt(1 - alpha_bar_t) * pred_noise) / torch.sqrt(alpha_t) if t > 0: sigma_t = torch.sqrt(betas[t]) z = torch.randn_like(x) x = mean + sigma_t * z else: x = mean return x

5.3 参数选择与调优经验

训练DDPM有几个关键参数需要调:

参数推荐值说明
T1000总时间步数,太小生成质量差,太大训练慢
β调度线性或余弦小图线性,大图余弦
学习率2e-4Adam优化器,配合warmup
batch size64-256越大越稳,但显存要求高
EMA decay0.9999对模型参数做指数移动平均,显著提升采样质量
梯度裁剪1.0防止梯度爆炸

EMA是我踩过的最大的坑。一开始没加EMA,训练loss降得很好,但采样出来的图全是噪声。后来加了EMA,采样质量立刻上了一个台阶。原因是扩散模型的训练目标本身噪声很大,模型参数会在最优解附近震荡,EMA相当于对参数做了平滑。

提示:EMA decay取0.9999意味着每步只更新万分之一的参数,看起来很小,但训练几十万步后效果显著。如果训练步数少,可以适当降低decay。

5.4 采样加速:DDIM

DDPM需要1000步采样,太慢了。DDIM(Denoising Diffusion Implicit Models)把采样过程变成确定性的,可以用更少的步数:

$$x_{t-1} = \sqrt{\bar{\alpha}{t-1}} \hat{x}0 + \sqrt{1-\bar{\alpha}{t-1}} \epsilon\theta(x_t, t)$$

其中 $\hat{x}_0 = \frac{x_t - \sqrt{1-\bar{\alpha}t} \epsilon\theta(x_t, t)}{\sqrt{\bar{\alpha}_t}}$。

DDIM可以用50步甚至20步就生成不错的图像。代价是多样性略有下降,但质量基本持平。

我实测下来,DDIM 50步和DDPM 1000步的FID差距在1以内,但速度快了20倍。所以实际部署时基本都用DDIM或更快的DPM-Solver。

6. 从DDPM到Stable Diffusion:潜在扩散模型

6.1 为什么要用潜在空间

DDPM直接在像素空间操作,对于512x512的图,Unet的输入输出都是512x512x3,计算量巨大。Stable Diffusion的核心改进是把扩散过程搬到潜在空间:先用一个VAE把图像压缩到64x64x4的潜在表示,然后在潜在空间上做扩散。

这样做的好处是计算量降低了约48倍(512x512 vs 64x64,再考虑通道数),同时生成质量基本不受影响。因为VAE的编码器已经学到了图像的压缩表示,扩散模型只需要在这个压缩表示上建模。

6.2 Stable Diffusion的架构

Stable Diffusion包含三个主要组件:

  • VAE:编码器把图像压缩到潜在空间,解码器把潜在表示恢复成图像
  • Unet:在潜在空间上做去噪,条件包括文本嵌入和时间步
  • 文本编码器:通常是CLIP,把文本提示编码成嵌入向量

Unet的条件注入通过交叉注意力实现:Unet的中间层有交叉注意力块,Q来自图像特征,K和V来自文本嵌入。这样文本信息就能影响去噪过程。

6.3 条件生成与Classifier-Free Guidance

Stable Diffusion用Classifier-Free Guidance(CFG)来控制生成内容与提示的匹配程度:

$$\hat{\epsilon} = \epsilon_\theta(x_t, t, \emptyset) + s \cdot (\epsilon_\theta(x_t, t, c) - \epsilon_\theta(x_t, t, \emptyset))$$

其中 $s$ 是guidance scale,通常取7.5。$s$ 越大,生成内容越贴合提示,但多样性下降,过大还会导致颜色过饱和。

训练时,文本条件以一定概率(通常10%)被替换为空条件,这样同一个网络既能做有条件生成也能做无条件生成。

6.4 Unet的改进

Stable Diffusion的Unet相比DDPM有几个关键改进:

  • 加入了交叉注意力层,用于注入文本条件
  • 使用了更多的注意力分辨率
  • 用了更大的通道数和更多的层
  • 时间步嵌入用了更复杂的MLP

这些改进让Unet能处理更复杂的条件生成任务。

7. 实操中的常见问题与排查

7.1 训练不收敛

症状:loss不下降或震荡严重。

排查思路:

  • 检查噪声调度是否合理,$\bar{\alpha}_T$ 应该接近0
  • 检查学习率是否过大,DDPM推荐2e-4
  • 检查是否有梯度爆炸,加梯度裁剪
  • 检查数据归一化,图像应该归一化到[-1, 1]

我遇到过一次loss震荡,最后发现是数据归一化到了[0, 1]而不是[-1, 1]。因为扩散模型的前向过程假设数据是零均值的,[0, 1]的数据会导致加噪后的分布偏移。

7.2 采样结果全是噪声

症状:训练loss正常,但采样出来全是噪声。

排查思路:

  • 检查是否用了EMA,没用EMA很容易出现这个问题
  • 检查采样时的方差选择,$\sigma_t^2$ 应该用 $\tilde{\beta}_t$
  • 检查时间步嵌入是否正确注入
  • 检查采样步数是否足够

7.3 生成图像模糊

症状:生成的图像能看出轮廓但很模糊。

排查思路:

  • 增加训练步数
  • 增大模型容量
  • 检查是否过拟合,加数据增强
  • 尝试余弦噪声调度

7.4 常见问题速查表

问题可能原因解决方案
loss不下降学习率过大/数据未归一化调小学习率/归一化到[-1,1]
采样全是噪声未用EMA/方差选择错误加EMA/用$\tilde{\beta}_t$
生成模糊训练不足/模型太小增加步数/增大模型
颜色过饱和CFG scale过大降低guidance scale
显存不足batch size过大减小batch/用梯度累积
采样太慢步数太多用DDIM/DPM-Solver

7.5 独家避坑技巧

第一个技巧:训练初期先用小图(32x32)验证流程,跑通了再上大图。我一开始直接上256x256,调了一周都没收敛,后来换成32x32,半天就跑通了,然后再逐步放大。

第二个技巧:保存检查点时同时保存EMA参数和原始参数。有时候EMA参数采样效果好,有时候原始参数更好,都留着方便对比。

第三个技巧:用wandb或tensorboard记录loss曲线和采样结果。扩散模型的loss曲线很平滑,看不出什么问题,必须看采样结果才能判断模型好坏。我一般每5000步采样一次,存成网格图。

第四个技巧:如果显存不够,可以用梯度累积模拟大batch。扩散模型对batch size比较敏感,小batch训练不稳定。梯度累积4次相当于batch size翻4倍。

第五个技巧:DDIM采样时,$\eta$ 参数控制随机性。$\eta=0$ 是确定性采样,$\eta=1$ 是DDPM采样。实际用 $\eta=0$ 效果就很好,而且可复现。

8. 扩散模型的扩展与改进方向

8.1 采样加速

DDIM之后,有一系列工作进一步加速采样。DPM-Solver把扩散方程的求解看成常微分方程,用高阶数值方法求解,可以用10-20步生成高质量图像。Consistency Models直接学习从噪声到图像的映射,支持一步生成。

这些方法的核心思想都是利用扩散过程的数学结构,用更聪明的数值方法替代朴素的迭代。

8.2 架构改进

Unet本身也在进化。DiT(Diffusion Transformer)用Transformer替代Unet,在ImageNet上取得了更好的效果。Transformer的scaling能力更强,随着模型增大,生成质量持续提升。

另一个方向是改进注意力机制,比如用线性注意力降低计算复杂度,或者用局部注意力减少计算量。

8.3 条件控制

除了文本条件,扩散模型还支持多种条件控制:

  • ControlNet:通过额外的网络注入空间条件(边缘、深度、姿态等)
  • IP-Adapter:注入图像条件,实现图像到图像的生成
  • LoRA:低秩适配,用少量参数微调模型

这些技术让扩散模型从“随机生成”变成了“可控生成”,大大扩展了应用场景。

8.4 应用场景

扩散模型的应用已经远远超出了图像生成:

  • 视频生成:Sora、Runway等用扩散模型生成视频
  • 3D生成:用扩散模型生成3D模型和场景
  • 音频生成:生成语音和音乐
  • 分子设计:生成新的分子结构
  • 地震数据:扩散模型用于地震数据去噪和重建

我最近在关注扩散模型在地震数据上的应用。地震数据本身噪声很大,传统去噪方法效果有限,扩散模型通过学习数据分布,能更好地分离信号和噪声。这个方向虽然小众,但很有潜力。

9. 一些个人体会

扩散模型是我见过的最“优雅”的生成模型之一。它的数学推导虽然涉及变分推断和马尔可夫链,但最终落地成一个极其简洁的MSE损失。这种“复杂理论、简单实现”的特点,让它既适合学术研究,也适合工程落地。

我刚开始学的时候,被那些公式吓到了,觉得肯定很难。但真正动手推了一遍之后发现,核心就是高斯分布的几个性质:可加性、重参数化、贝叶斯公式。把这些搞明白,剩下的就是工程问题了。

如果你也想入门扩散模型,我的建议是:先跑通一个最小的DDPM(32x32的CIFAR-10就够了),理解训练和采样的流程;然后逐步加组件(EMA、注意力、条件注入),看每个组件的作用;最后再去看Stable Diffusion的代码,会发现大部分东西都是相通的。

代码实现方面,我建议自己从零写一遍Unet和训练循环,不要直接抄现成的库。自己写一遍,踩一遍坑,比看十篇论文都管用。我当初就是自己写了一遍,才发现时间步嵌入的维度、跳跃连接的通道数匹配、注意力的reshape这些细节,看论文是注意不到的。

最后分享一个调试技巧:扩散模型的训练loss通常在0.1到0.5之间,如果loss降到0.01以下,大概率是过拟合了;如果loss一直在1以上,大概率是哪里配置错了。这个经验值不一定精确,但能帮你快速判断训练是否正常。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询