1. 从一张图说起:扩散模型到底在干什么
第一次看到扩散模型(Diffusion Model)这个词,很多人会以为它和热力学里的扩散现象有什么直接关系。其实名字的来源确实借用了物理里“粒子从高浓度向低浓度扩散”的直觉——在模型里,它指的是把一张清晰的图像一步步“加噪”直到变成纯噪声,再训练一个网络把这个过程反过来,从纯噪声一步步“去噪”还原出图像。听起来绕,但核心思想就这么朴素。
我最早接触这块是在做图像生成相关项目的时候。当时主流的生成方案还是生成对抗网络(GAN),训练不稳定、模式坍塌这些问题让人头疼。扩散模型刚出来时,采样慢得离谱,一张图要跑上千步,但生成质量确实惊艳。后来DDIM、潜在扩散(Latent Diffusion)这些改进出来,采样步数压到几十步甚至几步,才真正具备了落地价值。Stable Diffusion就是潜在扩散模型的典型代表,它把扩散过程放到VAE压缩后的潜空间里做,算力需求直接降了一个数量级。
这篇文章我打算把扩散模型的原理和实现从头到尾捋一遍。不是那种只贴公式的论文解读,而是从“一个从业者要动手实现它需要知道什么”的角度来写。内容包括:前向加噪和反向去噪的数学形式、训练目标为什么是预测噪声、U-Net骨干网络的结构设计、时间步嵌入怎么做、采样加速的几种思路,最后给一份可以跑起来的PyTorch实现。适合有一定深度学习基础、想搞明白扩散模型内部机制、或者想自己动手训一个小模型的人。如果你只是想知道怎么用现成工具出图,那这篇可能偏底层了,但看完你会对提示词为什么有效、采样步数怎么选这些问题有更本质的理解。
2. 扩散模型的核心原理拆解
2.1 前向过程:把图像一步步“溶解”成噪声
前向过程(Forward Process)也叫扩散过程,是一个固定的马尔可夫链。给定一张图像 (x_0),我们定义一系列时间步 (t = 1, 2, ..., T),每一步都往图像里加一点高斯噪声:
[ q(x_t | x_{t-1}) = \mathcal{N}(x_t; \sqrt{1-\beta_t} x_{t-1}, \beta_t I) ]
这里的 (\beta_t) 是每一步的噪声方差,通常从 (10^{-4}) 线性增加到 (0.02),总共 (T=1000) 步。(\sqrt{1-\beta_t}) 这个系数是为了保持方差稳定——如果不乘这个系数,加噪过程中像素值的方差会越来越大,最后数值爆炸。
这个式子看起来要一步步算,但实际上有个闭式解,可以直接从 (x_0) 跳到任意 (x_t):
[ q(x_t | x_0) = \mathcal{N}(x_t; \sqrt{\bar{\alpha}_t} x_0, (1-\bar{\alpha}_t) I) ]
其中 (\alpha_t = 1 - \beta_t),(\bar{\alpha}t = \prod{s=1}^{t} \alpha_s)。这个闭式解是训练能高效进行的关键——我们不需要真的模拟1000步加噪,直接采样一个 (t),用公式一步算出 (x_t) 就行。
用重参数化技巧写出来就是:
[ x_t = \sqrt{\bar{\alpha}_t} x_0 + \sqrt{1-\bar{\alpha}_t} \epsilon, \quad \epsilon \sim \mathcal{N}(0, I) ]
我习惯把这个式子理解成“信号和噪声的加权混合”。当 (t) 很小的时候,(\sqrt{\bar{\alpha}_t}) 接近1,图像基本还是原样;当 (t) 接近 (T) 时,(\sqrt{\bar{\alpha}_t}) 接近0,图像就变成了纯高斯噪声。这个从信号到噪声的渐变过程,就是扩散模型名字的由来。
注意:(\beta_t) 的调度策略(noise schedule)对生成质量影响很大。早期用线性调度,后来余弦调度(cosine schedule)被证明效果更好,因为它让 (\bar{\alpha}_t) 在中间时间段下降得更平缓,避免图像信息过早被破坏。
2.2 反向过程:训练一个网络学会“去噪”
反向过程(Reverse Process)才是真正需要学习的地方。我们希望从纯噪声 (x_T \sim \mathcal{N}(0, I)) 出发,一步步去噪,最终得到一张清晰的图像 (x_0)。理论上,如果知道真实的反向分布 (q(x_{t-1}|x_t)),就能精确还原。但这个分布依赖于整个数据集,没法直接算。
于是我们用神经网络 (p_\theta) 来近似它:
[ p_\theta(x_{t-1} | x_t) = \mathcal{N}(x_{t-1}; \mu_\theta(x_t, t), \Sigma_\theta(x_t, t)) ]
关键洞察来了:当 (\beta_t) 足够小的时候,反向过程也可以近似为高斯分布。所以我们只需要让网络预测这个高斯分布的均值 (\mu_\theta) 和方差 (\Sigma_\theta)。
进一步推导可以发现,均值 (\mu_\theta) 可以写成:
[ \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) ]
也就是说,网络真正需要预测的是噪声 (\epsilon_\theta(x_t, t))。这就是为什么扩散模型的训练目标通常是“预测噪声”——不是直接预测图像,而是预测当前步加入的噪声,然后用它反推出干净的图像。
这个设计非常巧妙。直接预测 (x_0) 的话,网络要学的东西跨度太大,从纯噪声到清晰图像,难度很高。而预测噪声相当于让网络专注于“这一步的噪声长什么样”,任务更局部、更稳定。我实测下来,预测噪声的收敛速度和最终质量都明显优于直接预测 (x_0)。
2.3 训练目标:一个简化到极致的损失函数
有了上面的推导,训练目标就变得非常简洁。原始论文从变分下界(ELBO)出发推导,最后化简成一个均方误差:
[ L_{\text{simple}} = \mathbb{E}{t, x_0, \epsilon} \left[ | \epsilon - \epsilon\theta(x_t, t) |^2 \right] ]
训练流程就是:
- 从数据集采样一张图像 (x_0)
- 随机采样一个时间步 (t \sim \text{Uniform}(1, T))
- 采样噪声 (\epsilon \sim \mathcal{N}(0, I))
- 计算 (x_t = \sqrt{\bar{\alpha}_t} x_0 + \sqrt{1-\bar{\alpha}_t} \epsilon)
- 让网络预测 (\epsilon_\theta(x_t, t)),计算MSE损失,反向传播
就这么简单。没有对抗训练,没有复杂的损失平衡,就是一个回归问题。这也是扩散模型比GAN好训的根本原因——它把生成问题转化成了一个去噪自编码器的训练问题。
实操心得:虽然理论上 (t) 是均匀采样,但实际训练中可以对 (t) 做重要性采样,让模型更关注那些“去噪难度大”的时间步。我试过在中间时间段((t) 在300到700之间)加大采样权重,FID指标有轻微改善,但提升有限,不如把精力放在网络结构和采样策略上。
3. 网络架构与关键组件实现
3.1 U-Net骨干:为什么是它而不是Transformer
扩散模型的去噪网络 (\epsilon_\theta(x_t, t)) 需要满足两个要求:输入输出尺寸一致,且能捕捉多尺度特征。U-Net天然符合这两个条件。
U-Net的结构是编码器-解码器加跳跃连接。编码器逐层下采样,提取从细粒度到粗粒度的特征;解码器逐层上采样,恢复空间分辨率;跳跃连接把编码器同层的特征直接拼到解码器,保留细节信息。对于去噪任务来说,这个设计非常合适——低层特征帮助恢复纹理细节,高层特征帮助理解全局结构。
具体到扩散模型用的U-Net,和原始医学图像分割用的U-Net有几个区别:
- 时间步嵌入:每个残差块都要注入时间步信息,让网络知道当前是第几步去噪
- 自注意力层:在低分辨率层加入自注意力,捕捉长距离依赖
- 组归一化:用GroupNorm而不是BatchNorm,因为训练时batch size通常很小
我实现的时候用的配置是:基础通道数64,通道倍数(1, 2, 4, 8),每个分辨率层2个残差块,注意力分辨率设在16x16和8x8。这个配置在256x256图像上大约有50M参数,单卡24G显存可以跑batch size 16左右。
3.2 时间步嵌入:让网络知道“现在是第几步”
时间步 (t) 是一个标量,但网络需要它来调节每一层的特征。做法和Transformer的位置编码类似,用正弦函数生成一个高维向量:
import torch import math def timestep_embedding(timesteps, dim, max_period=10000): half = dim // 2 freqs = torch.exp( -math.log(max_period) * torch.arange(half, dtype=torch.float32) / half ).to(timesteps.device) args = timesteps[:, None].float() * freqs[None] embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) if dim % 2: embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) return embedding这个嵌入向量经过两层MLP后,加到每个残差块的特征上。为什么用正弦编码而不是直接用一个可学习的embedding?因为正弦编码有外推性,训练时见过 (t=1) 到 (T),推理时如果要用不同的步数调度,编码仍然合理。可学习embedding就没这个好处。
注意:时间步嵌入的维度要和残差块的通道数匹配。我一般设成基础通道数的4倍,比如基础通道64,嵌入维度就是256。太小了信息容量不够,太大了浪费参数。
3.3 残差块与注意力:去噪网络的基本单元
每个残差块的结构是:GroupNorm → SiLU激活 → 卷积 → 注入时间步嵌入 → GroupNorm → SiLU → 卷积 → 残差连接。时间步嵌入通过一个线性层投影后加到第一个卷积的输出上。
class ResBlock(nn.Module): def __init__(self, in_channels, out_channels, time_emb_dim): super().__init__() self.norm1 = nn.GroupNorm(32, 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(32, out_channels) self.conv2 = nn.Conv2d(out_channels, out_channels, 3, padding=1) self.skip = 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.skip(x)自注意力层用在16x16和8x8分辨率上。注意力计算是标准的QKV形式,但加了一个残差连接和归一化。我试过在更高分辨率也加注意力,显存直接爆了,而且收益不明显——高分辨率层的卷积已经能捕捉足够的局部信息。
3.4 潜在扩散:把计算搬到潜空间
直接在像素空间做扩散,256x256的图就要处理196608维的数据,训练和采样都慢。潜在扩散(Latent Diffusion)的思路是先用一个VAE把图像压缩到潜空间,比如8倍下采样后变成32x32x4,维度降到4096,计算量减少几十倍。
VAE的编码器把图像 (x) 映射到潜变量 (z = \mathcal{E}(x)),解码器从潜变量重建图像 (\hat{x} = \mathcal{D}(z))。扩散过程在 (z) 上进行,训练目标不变,只是 (x_0) 换成了 (z_0)。采样时从噪声生成 (z_0),再解码成图像。
这个设计的关键是VAE的重建质量要足够好,否则扩散模型生成的东西会被VAE的瓶颈限制。Stable Diffusion用的VAE下采样8倍,潜空间通道4,重建质量在感知上几乎无损。我自己训VAE的时候发现,KL正则的权重很关键——太大导致重建模糊,太小导致潜空间分布太散,扩散模型学起来困难。一般设 (10^{-6}) 到 (10^{-4}) 之间比较合适。
4. 完整实现与训练流程
4.1 环境准备与依赖
我用的环境是Python 3.10 + PyTorch 2.0 + CUDA 11.8。依赖不多:
pip install torch torchvision einops accelerate tensorboardeinops用来做张量重排,比原生permute可读性好很多。accelerate处理混合精度和分布式训练。数据我用的是CIFAR-10和CelebA-64做实验,前者32x32,后者64x64,单卡就能训。
4.2 数据加载与预处理
数据预处理很简单,归一化到[-1, 1]就行。扩散模型对数据增强不敏感,因为加噪过程本身就是一种强增强。我试过加随机翻转,效果没有明显变化。
from torchvision import datasets, transforms transform = transforms.Compose([ transforms.Resize(64), transforms.CenterCrop(64), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize([0.5]*3, [0.5]*3) ]) dataset = datasets.CelebA(root='./data', split='train', transform=transform, download=True) dataloader = torch.utils.data.DataLoader(dataset, batch_size=64, shuffle=True, num_workers=4)4.3 噪声调度与扩散参数
噪声调度我用余弦调度,比线性调度在低噪声区域更平缓:
def cosine_beta_schedule(timesteps, s=0.008): steps = timesteps + 1 x = torch.linspace(0, timesteps, steps) alphas_cumprod = torch.cos(((x / timesteps) + s) / (1 + s) * math.pi * 0.5) ** 2 alphas_cumprod = alphas_cumprod / alphas_cumprod[0] betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1]) return torch.clip(betas, 0.0001, 0.9999)预计算好 (\bar{\alpha}_t)、(\sqrt{\bar{\alpha}_t})、(\sqrt{1-\bar{\alpha}_t}) 这些系数,训练时直接查表,避免重复计算。
4.4 训练循环与关键参数
训练循环的核心就是前面说的五步。我用混合精度训练,显存占用减少约40%,速度提升30%左右。
scaler = torch.cuda.amp.GradScaler() for epoch in range(num_epochs): for batch in dataloader: x0 = batch[0].cuda() t = torch.randint(0, T, (x0.shape[0],), device=x0.device) noise = torch.randn_like(x0) xt = sqrt_alphas_cumprod[t][:, None, None, None] * x0 + \ sqrt_one_minus_alphas_cumprod[t][:, None, None, None] * noise with torch.cuda.amp.autocast(): noise_pred = model(xt, t) loss = F.mse_loss(noise_pred, noise) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() optimizer.zero_grad()关键参数:学习率用2e-4,AdamW优化器,weight decay 0.01,EMA衰减率0.9999。EMA对生成质量影响很大,我试过不用EMA,FID直接差一截。batch size 64,训练了大概200k步,在单张A100上跑了约3天。
实操心得:训练初期loss下降很快,但别高兴太早,那只是模型学会了预测均值。真正决定生成质量的是中后期,loss下降很慢但FID在持续改善。我一般每10k步采样一批图看看效果,比只看loss曲线靠谱。
4.5 采样:从噪声生成图像
训练完之后,采样就是从 (x_T \sim \mathcal{N}(0, I)) 出发,逐步去噪。标准DDPM采样要跑1000步,太慢。实际用DDIM采样,50步就能出不错的结果:
@torch.no_grad() def ddim_sample(model, shape, steps=50, eta=0.0): x = torch.randn(shape).cuda() timesteps = torch.linspace(T-1, 0, steps).long().cuda() for i in range(steps): t = timesteps[i] prev_t = timesteps[i+1] if i+1 < steps else -1 noise_pred = model(x, t.unsqueeze(0)) alpha_t = alphas_cumprod[t] alpha_prev = alphas_cumprod[prev_t] if prev_t >= 0 else torch.tensor(1.0) x0_pred = (x - torch.sqrt(1-alpha_t) * noise_pred) / torch.sqrt(alpha_t) x0_pred = x0_pred.clamp(-1, 1) sigma = eta * torch.sqrt((1-alpha_prev)/(1-alpha_t) * (1-alpha_t/alpha_prev)) c = torch.sqrt(1-alpha_prev-sigma**2) x = torch.sqrt(alpha_prev) * x0_pred + c * noise_pred if eta > 0: x = x + sigma * torch.randn_like(x) return xeta=0就是确定性DDIM,eta=1退化成DDPM。我一般用eta=0,50步,生成一张64x64的图在A100上约0.5秒。如果追求更高质量,可以加到100步,但边际收益递减明显。
5. 常见问题与排查实录
5.1 生成图像模糊或颜色失真
这是最常见的问题。排查顺序:
| 现象 | 可能原因 | 排查方法 | 解决 |
|---|---|---|---|
| 整体模糊 | 训练不充分 | 看FID是否还在下降 | 继续训练 |
| 颜色偏灰 | 归一化问题 | 检查数据预处理 | 确认归一化到[-1,1] |
| 局部模糊 | VAE重建瓶颈 | 单独测VAE重建 | 降低KL权重或增大潜空间 |
| 网格状伪影 | 上采样方式 | 检查U-Net上采样 | 用最近邻+卷积替代转置卷积 |
我踩过最坑的一次是颜色整体偏绿,查了半天发现是数据加载时通道顺序搞错了,RGB当BGR用了。这种低级错误反而最难发现,因为loss曲线看起来完全正常。
5.2 训练loss不下降或震荡
先检查学习率是不是太大。扩散模型对学习率比较敏感,2e-4是个比较稳的值,超过5e-4容易震荡。如果loss一直不降,检查时间步嵌入有没有正确注入——我见过有人忘了把时间步嵌入加到残差块里,网络根本不知道自己在第几步,loss卡在0.5左右下不去。
另一个常见问题是EMA没开或者衰减率设错。EMA衰减率一般设0.999到0.9999,太小了起不到平滑效果,太大了更新太慢。我习惯用0.9999,训练步数超过100k之后效果明显。
5.3 采样步数与质量的关系
很多人问采样步数怎么选。我的经验是:
- 20步以下:细节丢失明显,适合快速预览
- 50步:质量和速度的平衡点,日常用这个
- 100步:质量提升有限,除非做对比实验
- 200步以上:基本没有可见提升
DDIM的步数不需要和训练步数一致,这是它比DDPM灵活的地方。训练用1000步,采样用50步完全没问题。但要注意,采样步数太少时,eta=0的确定性采样可能陷入局部最优,适当加一点随机性(eta=0.2左右)有时能改善多样性。
5.4 显存不够怎么办
扩散模型训练显存占用主要来自激活值。几个有效的优化:
- 混合精度训练:省40%左右
- 梯度检查点:省60%但慢30%
- 减小batch size:最直接但影响BN统计(不过我们用GroupNorm,影响小)
- 潜在扩散:省一个数量级
我最早在2080Ti上训64x64的模型,11G显存,batch size只能开到8。后来换用梯度检查点,开到16,训练时间从5天缩到4天,还算划算。
6. 几个值得深挖的扩展方向
扩散模型这块发展太快,我列几个自己关注的方向。一是采样加速,除了DDIM还有DPM-Solver、一致性模型(Consistency Models),后者能做到一步生成,虽然质量还有差距但进步很快。二是条件生成,classifier-free guidance是目前的主流做法,训练时随机丢掉条件,采样时用引导系数控制条件强度,系数设7到10之间比较常见。三是与Transformer的结合,DiT(Diffusion Transformer)用Transformer替换U-Net,在ImageNet上已经刷到了很好的FID,而且扩展性更好,模型越大效果越好。
我自己最近在试的是把扩散模型用到非图像领域,比如音频生成和分子结构生成。核心思路是一样的,只是数据形式和网络结构要调整。音频用1D卷积或Transformer,分子用图神经网络。踩过的坑是不同领域的数据分布差异很大,噪声调度需要重新调,不能直接套图像的那套参数。
最后分享一个小技巧:如果你只是想快速验证一个想法,不用从头训。拿预训练的Stable Diffusion,冻结VAE和文本编码器,只微调U-Net的注意力层,用LoRA,几张图就能出效果。我试过用20张图微调一个特定风格,在A100上跑了15分钟就有模有样了。这个方法适合做风格迁移和小样本适配,比从头训划算太多。