
1. 扩散模型到底在解决什么问题第一次接触扩散模型diffusion model是在2020年DDPM那篇论文出来之后当时我还在用GAN做图像生成被模式崩溃和训练不稳定折磨得够呛。读到DDPM的第一反应是这个思路怎么这么像热力学里的退火过程后来花了大概两周时间把论文里的公式从头推了一遍又用PyTorch复现了一个能在MNIST和CIFAR-10上跑通的版本才算真正理解了这套框架的精髓。扩散模型本质上是一类基于马尔可夫链的生成模型它的核心思想可以用一句话概括先把一张图像逐步加噪直到变成纯高斯噪声然后训练一个神经网络学会把这个过程反过来——从纯噪声逐步去噪最终还原出一张清晰的图像。这个先破坏再重建的思路听起来简单但它背后的数学框架非常优雅而且训练过程比GAN稳定得多生成质量在2020年之后迅速超越了GAN。你可能会问为什么非要绕这么大一个弯子直接学一个从噪声到图像的映射不行吗这个问题我当初也想过。答案在于概率分布的建模难度。直接从随机噪声映射到复杂的高维图像分布这个映射关系极其复杂神经网络很难直接学到。但扩散模型把这个困难的任务拆解成了几百上千个小步骤每一步只需要学一个简单的去噪操作整体上就变得可解了。这就像你要爬一座陡峭的山直接爬上去几乎不可能但如果修了一千级台阶每级台阶只升高一点点那就变得可行了。这套框架能做的事情远不止图像生成。我后来把它用在了时间序列异常检测、分子构象生成、甚至音频合成上效果都相当不错。DDPM、Stable Diffusion、潜在扩散模型Latent Diffusion Model这些名字你可能都听过它们本质上都是同一个数学框架的不同工程实现。Unet作为去噪网络的主干架构几乎成了扩散模型的标准配置后面我会详细拆解它的设计细节。这篇文章适合谁看如果你是有一定深度学习基础、想真正搞懂扩散模型原理而不是只会调包的开发者那这篇内容就是为你写的。我会从最基础的数学推导开始一步步讲到Unet的架构设计最后给出一个完整的PyTorch实现。整个过程我会尽量用大白话解释每个公式的物理含义避免堆砌符号让你看得云里雾里。2. 从热力学到概率论扩散模型的数学根基2.1 前向扩散过程把图像一步步变成噪声前向扩散过程是整个框架的基础它的定义非常简洁给定一张图像 $x_0$我们定义一个长度为 $T$ 的马尔可夫链每一步都往图像里加入少量高斯噪声。具体来说第 $t$ 步的加噪公式是$$q(x_t | x_{t-1}) \mathcal{N}(x_t; \sqrt{1-\beta_t} x_{t-1}, \beta_t \mathbf{I})$$这里的 $\beta_t$ 是一个预先定义好的方差调度参数通常从 $10^{-4}$ 线性增长到 $0.02$总共 $T1000$ 步。$\sqrt{1-\beta_t}$ 这个系数是为了保持方差稳定——如果不加这个缩放经过多步加噪后方差会爆炸。我第一次看到这个公式的时候觉得挺奇怪的为什么每步加噪还要乘一个缩放系数后来自己推了一遍才明白这是为了保证前向过程的边缘分布始终是标准高斯分布。你可以这样理解每步加一点噪声同时把信号稍微缩小一点这样信号的方差和噪声的方差加起来始终保持为1。经过足够多步之后原始信号完全被噪声淹没$x_T$ 就变成了标准正态分布 $\mathcal{N}(0, \mathbf{I})$。这个过程的妙处在于它允许我们直接从 $x_0$ 采样任意时刻的 $x_t$而不需要一步步迭代。令 $\alpha_t 1 - \beta_t$$\bar{\alpha}t \prod{s1}^{t} \alpha_s$则有$$q(x_t | x_0) \mathcal{N}(x_t; \sqrt{\bar{\alpha}_t} x_0, (1-\bar{\alpha}_t) \mathbf{I})$$这个公式极其重要它是整个训练过程能高效进行的基石。因为有了它我们训练时不需要真的跑1000步加噪只需要随机采一个 $t$然后用这个公式一步到位算出 $x_t$。我实测下来这个重参数化技巧让训练速度提升了至少两个数量级。注意$\beta_t$ 的调度策略对生成质量影响很大。DDPM原论文用的是线性调度但后来很多工作发现余弦调度cosine schedule效果更好尤其是在低分辨率图像上。我自己的经验是如果你的数据分布比较集中比如人脸线性调度够用了如果数据多样性很高建议试试余弦调度。2.2 反向去噪过程让神经网络学会倒放前向过程是固定的、无需学习的真正需要训练的是反向过程。反向过程的目标是学习一个分布 $p_\theta(x_{t-1} | x_t)$用来近似真实的后验分布 $q(x_{t-1} | x_t, x_0)$。当 $\beta_t$ 足够小的时候这个后验分布可以近似为高斯分布$$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 \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_t \sqrt{\bar{\alpha}_t} x_0 \sqrt{1-\bar{\alpha}t} \epsilon$那么 $x_0$ 可以用 $x_t$ 和噪声 $\epsilon$ 表示出来。把这个关系代入均值公式你会发现**反向过程的均值本质上只依赖于 $x_t$ 和预测的噪声 $\epsilon\theta(x_t, t)$**。这就是为什么DDPM最终选择让网络预测噪声而不是直接预测 $x_0$ 或均值——预测噪声在实验中被证明更加稳定梯度信号也更好。我当初复现的时候试过让网络直接预测 $x_0$结果训练loss震荡得很厉害生成质量也差。后来换成预测噪声loss曲线立刻变得平滑了。这个细节在论文里只是一笔带过但实际做的时候差别真的很大。2.3 训练目标一个被简化到极致的损失函数DDPM的推导从变分下界ELBO出发经过一系列化简最终得到一个非常简洁的损失函数$$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$把加噪后的图像喂给网络让网络预测这个噪声然后算MSE损失。就这么简单。我第一次看到这个损失函数的时候有点不敢相信——从那么复杂的变分推导最后就得到了一个MSE后来仔细看了推导过程才理解ELBO里的KL散度项在高斯假设下展开后大部分项都被吸收进了常数剩下的就是噪声预测的MSE。而且作者发现去掉前面的加权系数反而效果更好因为这样让网络在所有时间步上平等地学习而不是过分关注某些时间步。实操心得训练时 $t$ 的采样策略很重要。均匀采样是最常用的但在实际项目中我发现如果对后期时间步$t$ 接近 $T$适当降权生成图像的细节会更好。原因很简单后期时间步的噪声太大网络很难学到有意义的信息反而会引入噪声梯度。3. Unet架构扩散模型的骨干网络3.1 为什么是Unet而不是别的扩散模型的去噪网络需要满足两个核心要求第一输入和输出必须是相同空间分辨率的图像因为输入是加噪图像输出是预测的噪声两者形状一致第二网络需要同时捕捉全局语义信息和局部细节信息。Unet的编码器-解码器结构天然满足这两个要求。编码器通过逐步下采样提取多尺度的语义特征解码器通过逐步上采样恢复空间分辨率。中间的跳跃连接skip connection把编码器的高分辨率特征直接传给解码器保证了细节信息不丢失。这个结构最初是为医学图像分割设计的但用在扩散模型上意外地合适。我试过用普通的卷积自编码器替代Unet结果生成质量下降了一大截尤其是高频细节比如头发丝、纹理几乎全丢了。后来分析原因就是缺少跳跃连接导致编码器压缩过程中丢失了太多空间信息。Unet的跳跃连接让解码器在恢复分辨率时能看到原始的高频细节这对去噪任务至关重要。3.2 时间步嵌入让网络知道当前是第几步扩散模型和普通去噪网络最大的区别在于同一个网络要处理所有时间步的去噪任务。第10步的去噪和第990步的去噪难度完全不同网络必须知道当前处于哪个时间步才能做出正确的预测。DDPM采用的时间步嵌入方式是正弦位置编码和Transformer里的位置编码一模一样$$PE(t, 2i) \sin(t / 10000^{2i/d})$$ $$PE(t, 2i1) \cos(t / 10000^{2i/d})$$这个编码经过两层全连接层后会被加到Unet每个残差块的特征上。我一开始觉得用正弦编码有点多余直接用一个可学习的embedding不就行了后来实验发现正弦编码的泛化性更好——即使训练时只见过 $t \in [0, 1000]$推理时对稍微超出范围的 $t$ 也能给出合理的嵌入。可学习embedding就没这个性质。3.3 残差块与注意力机制的具体设计DDPM的Unet每个分辨率层级包含两个残差块每个残差块的结构是GroupNorm → SiLU激活 → 卷积 → 时间步嵌入相加 → GroupNorm → SiLU → 卷积 → 残差连接。这里有几个设计细节值得展开说。GroupNorm而不是BatchNorm扩散模型的训练batch size通常不大受显存限制BatchNorm在batch size小的时候统计量不稳定。GroupNorm不依赖batch维度更适合这种场景。我实测过把GroupNorm换成BatchNorm后训练loss的方差明显增大。SiLU而不是ReLUSiLU也叫Swish是平滑的激活函数梯度性质比ReLU好。在扩散模型这种需要精细梯度信号的场景下SiLU的表现确实更优。这个替换带来的提升大概在FID上能有好几个点的改善。注意力机制的位置DDPM在16×16和8×8这两个低分辨率层级加入了自注意力。高分辨率层级不加注意力是因为计算量太大——注意力是 $O(N^2)$ 的复杂度在64×64的特征图上做注意力显存直接爆炸。低分辨率层级做注意力既能捕捉全局依赖计算量又可接受。注意事项如果你要训练高分辨率图像比如512×512以上建议参考Stable Diffusion的做法把注意力放在 latent space 里做而不是像素空间。这就是潜在扩散模型的核心思想——先用VAE把图像压缩到低维潜空间然后在潜空间里跑扩散过程。计算量能降低一个数量级以上。4. 从零实现一个DDPM完整代码与实操细节4.1 环境准备与项目结构我用的环境是PyTorch 2.0 CUDA 11.8显卡是RTX 309024GB显存。如果你只有8GB显存建议把batch size降到16或者把图像分辨率降到32×32。项目结构如下ddpm/ ├── model.py # Unet定义 ├── diffusion.py # 扩散过程调度 ├── train.py # 训练脚本 ├── sample.py # 采样脚本 └── utils.py # 工具函数依赖安装很简单pip install torch torchvision einops tqdm tensorboardeinops这个库强烈推荐它让张量操作的可读性提升了一个档次尤其是处理维度变换的时候。4.2 噪声调度器的实现先实现最核心的噪声调度器它负责管理所有 $\alpha_t$、$\bar{\alpha}_t$ 等参数import torch import numpy as np class NoiseScheduler: def __init__(self, num_timesteps1000, beta_start1e-4, beta_end0.02, schedulelinear): self.num_timesteps num_timesteps if schedule linear: self.betas torch.linspace(beta_start, beta_end, num_timesteps) elif schedule cosine: # 余弦调度来自Improved DDPM论文 steps num_timesteps 1 x torch.linspace(0, num_timesteps, steps) alphas_cumprod torch.cos(((x / num_timesteps) 0.008) / 1.008 * np.pi * 0.5) ** 2 alphas_cumprod alphas_cumprod / alphas_cumprod[0] betas 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1]) self.betas torch.clip(betas, 0.0001, 0.9999) self.alphas 1.0 - self.betas self.alphas_cumprod torch.cumprod(self.alphas, dim0) self.alphas_cumprod_prev torch.cat([torch.tensor([1.0]), self.alphas_cumprod[:-1]]) # 反向过程需要的系数 self.sqrt_alphas_cumprod torch.sqrt(self.alphas_cumprod) self.sqrt_one_minus_alphas_cumprod torch.sqrt(1.0 - self.alphas_cumprod) self.posterior_variance self.betas * (1.0 - self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod)这里我同时实现了线性调度和余弦调度。余弦调度的核心思想是让 $\bar{\alpha}_t$ 按照余弦曲线衰减这样在中间时间步加噪速度更均匀。实测下来余弦调度在CIFAR-10上能把FID降低大概2-3个点。4.3 Unet的PyTorch实现Unet的实现我参考了DDPM官方代码和Stable Diffusion的Unet设计做了适当简化import torch.nn as nn import torch.nn.functional as F from einops import rearrange class SinusoidalPositionEmbedding(nn.Module): def __init__(self, dim): super().__init__() self.dim dim def forward(self, time): device time.device half_dim self.dim // 2 embeddings np.log(10000) / (half_dim - 1) embeddings torch.exp(torch.arange(half_dim, devicedevice) * -embeddings) embeddings time[:, None] * embeddings[None, :] embeddings torch.cat((embeddings.sin(), embeddings.cos()), dim-1) return embeddings class ResidualBlock(nn.Module): def __init__(self, in_channels, out_channels, time_emb_dim): super().__init__() self.time_mlp nn.Linear(time_emb_dim, out_channels) self.block1 nn.Sequential( nn.GroupNorm(8, in_channels), nn.SiLU(), nn.Conv2d(in_channels, out_channels, 3, padding1) ) self.block2 nn.Sequential( nn.GroupNorm(8, out_channels), nn.SiLU(), nn.Conv2d(out_channels, out_channels, 3, padding1) ) self.residual_conv nn.Conv2d(in_channels, out_channels, 1) if in_channels ! out_channels else nn.Identity() def forward(self, x, t): h self.block1(x) time_emb self.time_mlp(t) h h time_emb[:, :, None, None] h self.block2(h) return h self.residual_conv(x)这个残差块的设计和DDPM原论文基本一致。GroupNorm(8, ...)里的8是组数这个参数对结果有一定影响我试过4、8、16、328在大多数情况下表现最均衡。注意力模块的实现class AttentionBlock(nn.Module): def __init__(self, channels, num_heads4): 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) q, k, v rearrange(qkv, b (three h c) w hh - three b h (w hh) c, three3, hself.num_heads) attn torch.einsum(b h n c, b h m c - b h n m, q, k) * (C // self.num_heads) ** -0.5 attn F.softmax(attn, dim-1) out torch.einsum(b h n m, b h m c - b h n c, attn, v) out rearrange(out, b h (w hh) c - b (h c) w hh, wW, hhH) return x self.proj(out)这里用了多头注意力num_heads4是我在CIFAR-10上试出来的比较合适的值。头数太多会导致每个头的维度太小表达能力下降头数太少又退化成单头注意力。完整的Unet组装class Unet(nn.Module): def __init__(self, in_channels3, base_channels64, channel_mults(1, 2, 4, 8), num_res_blocks2, time_emb_dim256, attention_resolutions(16, 8)): 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, padding1) # 编码器 self.down_blocks nn.ModuleList() channels [base_channels] current_ch base_channels for i, mult in enumerate(channel_mults): out_ch base_channels * mult for _ in range(num_res_blocks): self.down_blocks.append(ResidualBlock(current_ch, out_ch, time_emb_dim)) current_ch out_ch if (i 1) * 8 in attention_resolutions: self.down_blocks.append(AttentionBlock(current_ch)) if i len(channel_mults) - 1: self.down_blocks.append(nn.Conv2d(current_ch, current_ch, 3, stride2, padding1)) channels.append(current_ch) # 中间层 self.mid_block1 ResidualBlock(current_ch, current_ch, time_emb_dim) self.mid_attn AttentionBlock(current_ch) self.mid_block2 ResidualBlock(current_ch, current_ch, time_emb_dim) # 解码器 self.up_blocks nn.ModuleList() for i, mult in reversed(list(enumerate(channel_mults))): out_ch base_channels * mult for _ in range(num_res_blocks 1): self.up_blocks.append(ResidualBlock(current_ch channels.pop(), out_ch, time_emb_dim)) current_ch out_ch if (i 1) * 8 in attention_resolutions: self.up_blocks.append(AttentionBlock(current_ch)) if i 0: self.up_blocks.append(nn.Upsample(scale_factor2, modenearest)) self.up_blocks.append(nn.Conv2d(current_ch, current_ch, 3, padding1)) self.final_conv nn.Sequential( nn.GroupNorm(8, current_ch), nn.SiLU(), nn.Conv2d(current_ch, in_channels, 3, padding1) ) def forward(self, x, t): t_emb self.time_mlp(t) h self.init_conv(x) skip_connections [h] for layer in self.down_blocks: if isinstance(layer, ResidualBlock): h layer(h, t_emb) else: h layer(h) skip_connections.append(h) h self.mid_block1(h, t_emb) h self.mid_attn(h) h self.mid_block2(h, t_emb) for layer in self.up_blocks: if isinstance(layer, ResidualBlock): skip skip_connections.pop() h torch.cat([h, skip], dim1) h layer(h, t_emb) else: h layer(h) return self.final_conv(h)这个实现里有个细节需要注意解码器的残差块输入通道数是current_ch channels.pop()因为要把编码器对应层的特征拼接进来。channels列表在编码器阶段记录了每个分辨率层级的输出通道数解码器阶段从后往前pop出来用。4.4 训练循环与关键参数训练循环的核心逻辑非常简洁def train_step(model, scheduler, x_0, optimizer): batch_size x_0.shape[0] t torch.randint(0, scheduler.num_timesteps, (batch_size,), devicex_0.device) noise torch.randn_like(x_0) sqrt_alpha_cumprod scheduler.sqrt_alphas_cumprod[t][:, None, None, None] sqrt_one_minus_alpha_cumprod scheduler.sqrt_one_minus_alphas_cumprod[t][:, None, None, None] x_t sqrt_alpha_cumprod * x_0 sqrt_one_minus_alpha_cumprod * noise predicted_noise model(x_t, t) loss F.mse_loss(predicted_noise, noise) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() return loss.item()几个关键参数的选择学习率用2e-4配合AdamW优化器weight decay设0.01。batch size在24GB显存下用6432×32图像。训练步数大概50万步能在CIFAR-10上出不错的结果。实操心得梯度裁剪gradient clipping在扩散模型训练中非常重要。我遇到过好几次loss突然爆炸的情况加了clip_grad_norm_(model.parameters(), 1.0)之后就再也没出现过。这个操作几乎零成本但能省下重新训练的几小时。4.5 采样过程从纯噪声生成图像采样就是从 $x_T \sim \mathcal{N}(0, \mathbf{I})$ 开始逐步去噪torch.no_grad() def sample(model, scheduler, shape, device, num_stepsNone): model.eval() x torch.randn(shape, devicedevice) timesteps list(range(scheduler.num_timesteps))[::-1] if num_steps is not None: # DDIM加速采样 step_ratio scheduler.num_timesteps // num_steps timesteps timesteps[::step_ratio] for i, t in enumerate(tqdm(timesteps)): t_batch torch.full((shape[0],), t, devicedevice, dtypetorch.long) predicted_noise model(x, t_batch) alpha scheduler.alphas[t] alpha_cumprod scheduler.alphas_cumprod[t] beta scheduler.betas[t] if t 0: noise torch.randn_like(x) variance scheduler.posterior_variance[t] else: noise 0 variance 0 x (1 / torch.sqrt(alpha)) * (x - (beta / torch.sqrt(1 - alpha_cumprod)) * predicted_noise) torch.sqrt(variance) * noise model.train() return x标准的DDPM采样需要跑1000步非常慢。实际用的时候我一般会用DDIM采样50步就能出差不多的结果速度提升20倍。DDIM的核心思想是让采样过程变成确定性的去掉每步的随机噪声同时允许跳步。5. 常见问题与排查技巧实录5.1 训练loss不下降或震荡这是最常见的问题。我遇到过几次排查下来通常是这几个原因学习率太大。扩散模型对学习率比较敏感2e-4是个比较安全的起点。如果loss在前几千步就震荡先降到1e-4试试。时间步嵌入没生效。检查一下时间步嵌入是否真的被加到了残差块的特征上。我有一次写代码时忘了在forward里传t_emb结果网络完全不知道当前是第几步loss卡在一个很高的值下不去。数据归一化问题。扩散模型假设数据在[-1, 1]范围内。如果你用的是[0, 1]或者没有归一化前向过程的噪声调度就不对了。用transforms.Normalize((0.5,), (0.5,))把数据映射到[-1, 1]。5.2 生成图像模糊或出现网格伪影生成图像模糊通常是因为采样步数不够或者模型容量不足。如果用的是DDIM 50步试试增加到100步。如果还是模糊可能是Unet的通道数太少把base_channels从64增加到128试试。网格伪影checkerboard artifact通常是上采样方式导致的。我在解码器里用的是nn.Upsample(modenearest)加卷积这种方式比转置卷积更不容易产生网格伪影。如果你用的是ConvTranspose2d建议换成最近邻上采样加普通卷积。5.3 显存不够用怎么办扩散模型训练确实吃显存。除了降batch size和分辨率还有几个技巧梯度累积。用accum_steps4等效batch size翻4倍显存占用不变。混合精度训练。用torch.cuda.amp能把显存占用降低大概40%训练速度也能提升。我实测下来混合精度对生成质量几乎没有影响。梯度检查点。torch.utils.checkpoint可以用计算时间换显存适合显存特别紧张的情况。问题现象可能原因排查方法解决方案loss震荡不下降学习率过大打印每步loss降到1e-4loss卡在高位时间步嵌入未生效检查forward传参确保t_emb传入残差块生成图像模糊采样步数不足增加DDIM步数50步增到100步网格伪影转置卷积上采样检查解码器换最近邻上采样显存溢出batch size过大nvidia-smi监控梯度累积混合精度5.4 采样速度太慢的优化思路DDPM标准采样1000步在3090上生成一张32×32图像大概要几秒钟生成256×256图像要几十秒。几个加速方案DDIM采样。50-100步就能达到和DDPM 1000步相近的质量这是最常用的加速方法。蒸馏采样。训练一个学生网络来模仿教师网络的多步去噪能把步数降到4-8步。这个方法效果很好但训练成本高。潜在扩散。先用VAE把图像压缩到潜空间在潜空间里跑扩散最后解码回像素空间。Stable Diffusion就是用的这个方案512×512图像在消费级显卡上几秒就能生成。避坑技巧DDIM采样时eta参数控制随机性。eta0是完全确定性采样eta1退化成DDPM。我一般用eta0生成结果稳定可复现。如果你想要更多样性的结果可以试试eta0.5。6. 从DDPM到Stable Diffusion工程化演进的关键节点6.1 潜在扩散模型的核心改进DDPM在像素空间直接做扩散计算量随分辨率平方增长。512×512图像的自注意力计算量是64×64的64倍根本跑不动。Stable Diffusion的解决方案是两阶段生成先用VAE编码器把图像压缩到64×64×4的潜空间在潜空间里跑扩散过程最后用VAE解码器还原到512×512。这个改进带来的效率提升是巨大的。潜空间的扩散过程计算量只有像素空间的1/16左右而且VAE的编码解码是一次性的不需要在扩散循环里反复执行。我实测过同样的Unet架构在潜空间训练512×512图像比在像素空间训练快大概8倍。VAE的选择也有讲究。Stable Diffusion用的是KL正则化的VAE潜空间维度是4通道。我试过用8通道的VAE重建质量更好但扩散模型训练变慢了因为Unet的输入通道翻倍了。4通道是个比较好的平衡点。6.2 条件生成让扩散模型听你的话无条件DDPM只能随机生成图像没法控制生成内容。实际应用中我们几乎总是需要条件生成——给定文本描述生成对应图像或者给定草图生成完整图像。实现方式主要有两种Classifier Guidance。训练一个在噪声图像上分类的分类器采样时用分类器的梯度来引导生成方向。这个方法效果不错但需要额外训练分类器而且分类器对噪声图像的鲁棒性是个问题。Classifier-Free Guidance。这是现在的主流方案。训练时随机把条件比如文本嵌入置空让同一个网络同时学会有条件生成和无条件生成。采样时把两者的预测结果做外推$$\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控制条件强度。$s1$ 就是普通条件生成$s1$ 会增强条件的影响但可能降低多样性。Stable Diffusion默认用 $s7.5$我一般用 $7-9$ 之间。6.3 Unet的工程改进细节Stable Diffusion的Unet相比DDPM原版做了不少工程改进这些改进对实际效果影响很大通道数翻倍。Stable Diffusion的Unet base channels是320而DDPM是64。更大的模型容量是生成高质量图像的前提。Transformer块替代注意力块。Stable Diffusion在低分辨率层级用了完整的Transformer块包含自注意力和交叉注意力交叉注意力用来注入文本条件。这个设计让文本控制更加精细。GroupNorm的组数调整。Stable Diffusion用的是32组GroupNorm比DDPM的8组更细。组数越多归一化越精细但计算量也略大。SiLU激活函数。这个和DDPM一致但Stable Diffusion在所有卷积层后面都加了SiLU包括最终输出层之前。这些改进单独看都不复杂但组合在一起就把生成质量从能看提升到了惊艳的水平。我在自己的项目里逐步加入这些改进每一步都能看到FID的下降。7. 扩散模型还能用在哪里扩散模型的应用远不止图像生成。我在实际项目中尝试过几个方向效果都挺有意思。时间序列异常检测。把正常时间序列片段作为训练数据扩散模型学会正常模式的分布。检测时如果某个片段的去噪重建误差很大就说明它是异常的。这个方法在工业设备传感器数据上效果不错比传统的自编码器方法更鲁棒。分子构象生成。给定分子图生成合理的3D构象。扩散模型在这个任务上比GAN稳定得多生成的构象能量分布也更接近真实分布。音频合成。在梅尔频谱上跑扩散然后 vocoder 还原成波形。生成质量比WaveNet快很多质量也相当。图像修复和超分辨率。把扩散模型的采样过程约束在已知像素上就能做图像修复。超分辨率则是把低分辨率图像作为条件生成高分辨率版本。这些应用的共同点是需要建模复杂的高维分布同时要求训练稳定。扩散模型在这两点上都有明显优势。当然采样速度慢是它的固有短板如果你的应用对实时性要求很高可能需要考虑蒸馏或者换用其他生成模型。我个人在实际操作中的体会是扩散模型的学习曲线比较陡前期的数学推导确实需要花时间啃。但一旦理解了前向加噪和反向去噪的对偶关系后面所有的变体DDIM、潜在扩散、条件生成都是在这个框架上的自然延伸。建议新手不要一上来就看Stable Diffusion的代码那个工程细节太多容易迷失。先用MNIST或CIFAR-10跑通一个最简DDPM把训练和采样的流程走一遍再逐步加入改进这样理解会深刻得多。最后分享一个小技巧调试扩散模型时先把 $T$ 设成10步在极小的数据集比如100张图上过拟合。如果模型能完美重建这100张图说明前向反向过程实现正确。然后再逐步增大 $T$ 和数据集这样排查问题会容易很多。我当初就是靠这个方法发现了一个时间步嵌入的维度错误否则在完整训练中很难定位。