1. 从一张模糊噪点图到高清大图:DDPM到底在做什么
第一次接触DDPM(Denoising Diffusion Probabilistic Models,去噪扩散概率模型)的人,脑子里冒出来的第一个问题通常是:为什么给一张全是噪点的图,模型能一步步把它变回一张清晰的照片?这事听起来像变魔术,但拆开看,逻辑其实非常朴素。
DDPM的核心思想可以用一句话概括:把“加噪”这件事反过来做。训练阶段,我们拿一张干净图片,不断往上面撒噪声,撒到它变成纯高斯噪声;推理阶段,我们训练一个网络,让它学会从噪声里一步步“擦掉”噪声,最终还原出图像。整个过程像是一块干净的玻璃被逐渐糊上泥巴,而模型学的是怎么一层层把泥巴擦干净。
这个思路之所以在图像生成领域炸开,是因为它解决了此前GAN(生成对抗网络)训练不稳定、模式崩塌的老大难问题。GAN像两个人在博弈,一个造假一个鉴假,训练起来经常一边倒,生成器要么摆烂要么只生成少数几种图。DDPM不玩对抗,它就是一个回归问题——预测噪声。目标函数简单、训练稳定、生成质量高,这三点加在一起,让它成了图像生成大模型的主流路线之一。
适合谁来读这篇内容?如果你已经写过PyTorch的基础训练循环,知道卷积、注意力、残差连接大概是怎么回事,但没亲手搭过一个完整的扩散模型,那这篇就是给你准备的。我会从UNet结构、噪声调度、训练目标、采样过程一路讲到实操中会踩的坑,代码能直接抄,参数会解释为什么这么设。
提示:DDPM不是“一步到位”的生成模型。它的推理需要几十到上千步迭代,这也是后来DDIM、潜在扩散模型(LDM)等加速方案出现的根本原因。理解DDPM是理解整个扩散模型家族的地基。
2. 整体架构设计:为什么是UNet加噪声预测
2.1 扩散模型的前向过程与反向过程
DDPM把图像生成拆成两个马尔可夫链:前向扩散和反向去噪。
前向过程是固定的,不需要学习。给定一张干净图 (x_0),我们定义一系列时间步 (t=1,...,T),每一步按照一个方差调度 (\beta_t) 往图上加高斯噪声:
[ q(x_t | x_{t-1}) = \mathcal{N}(x_t; \sqrt{1-\beta_t} x_{t-1}, \beta_t I) ]
这个式子看着唬人,实际意思就是:每一步把原图乘一个略小于1的系数,再加一点随机噪声。当 (T) 足够大(通常1000步),(x_T) 就近似纯标准高斯噪声了。
反向过程则是我们要学的:从 (x_T) 出发,一步步预测并减去噪声,最终得到 (x_0)。网络要学的不是直接生成图像,而是预测每一步加入的噪声(\epsilon)。这就是“噪声预测”这个词的由来。
为什么预测噪声而不是直接预测 (x_0)?我试过两种方案,预测噪声的收敛速度和最终质量都更稳。原因在于:在 (t) 很大时,(x_t) 几乎全是噪声,直接预测 (x_0) 等于让网络从纯随机里猜原图,难度极高;而预测噪声相当于告诉网络“这一步我加了多少噪声”,任务更局部、更可学。
2.2 UNet作为骨干网络的合理性
DDPM的骨干网络几乎清一色用UNet,这不是偶然。UNet的结构是编码器-解码器加跳跃连接:编码器逐层下采样,提取从边缘、纹理到语义的抽象特征;解码器逐层上采样,恢复空间分辨率;跳跃连接把编码器同层的特征直接拼到解码器对应层,保留细节信息。
扩散模型的去噪任务需要同时具备两种能力:一是理解全局结构(比如这是一张人脸还是一辆车),二是恢复局部细节(比如眼睛的高光、车漆的反光)。UNet的下采样路径负责前者,上采样路径和跳跃连接负责后者。如果换成纯Transformer,全局建模强但局部细节恢复成本高;如果换成纯CNN堆叠,感受野有限,全局一致性差。UNet是这两者之间一个非常务实的平衡点。
2.3 时间步嵌入:让网络知道“现在第几步”
反向过程每一步的去噪难度不同。(t) 接近 (T) 时,输入几乎是纯噪声,网络需要做大幅度的结构决策;(t) 接近0时,输入已经很清晰,网络只需微调细节。如果网络不知道当前是第几步,它就没法调整自己的行为。
DDPM的做法是把时间步 (t) 编码成一个向量,注入到UNet的每个残差块里。常用的编码方式是正弦位置编码,和Transformer里的位置编码同源:
import math import torch import torch.nn as nn 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 = math.log(10000) / (half_dim - 1) embeddings = torch.exp(torch.arange(half_dim, device=device) * -embeddings) embeddings = time[:, None] * embeddings[None, :] embeddings = torch.cat((embeddings.sin(), embeddings.cos()), dim=-1) return embeddings这个编码的好处是:不同时间步得到的向量在空间中分布均匀,且相邻时间步的编码相似度高,网络能平滑地适应不同噪声水平。我实测过用可学习嵌入代替正弦编码,效果略差,尤其在训练步数较少时,正弦编码的泛化性更好。
注意:时间步嵌入的维度要和UNet残差块的特征维度匹配,通常通过一个两层MLP映射后再加到特征图上。如果维度不匹配,PyTorch会直接报广播错误,这是新手最容易卡住的地方之一。
3. 核心细节拆解:噪声调度、损失函数与UNet改进
3.1 噪声调度表的选型与计算
噪声调度决定了每一步加多少噪声。DDPM原论文用的是线性调度:(\beta_t) 从 (10^{-4}) 线性增加到 (0.02)。这个设置下,前向过程在 (t) 较大时噪声增长很快,导致反向过程后期(接近 (x_0))的步数“浪费”了——很多步都在处理几乎相同的噪声水平。
后来改进的余弦调度更受欢迎,它让噪声在中间时间段增长更平缓,信息破坏更均匀:
import numpy as np def cosine_beta_schedule(timesteps, s=0.008): steps = timesteps + 1 x = np.linspace(0, timesteps, steps) alphas_cumprod = np.cos(((x / timesteps) + s) / (1 + s) * np.pi * 0.5) ** 2 alphas_cumprod = alphas_cumprod / alphas_cumprod[0] betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1]) return np.clip(betas, 0.0001, 0.9999)这里 (\bar{\alpha}t = \prod{i=1}^{t}(1-\beta_i)) 是累积信号保留率。余弦调度的核心是让 (\bar{\alpha}_t) 按余弦曲线下降,而不是线性下降。我对比过两种调度在CIFAR-10上的FID,余弦调度通常能低2到3个点,且训练更稳定。
实际代码里,我们会预先计算好所有时间步的 (\bar{\alpha}_t)、(\sqrt{\bar{\alpha}_t})、(\sqrt{1-\bar{\alpha}_t}),训练时直接查表,避免重复计算。
3.2 简化损失函数:为什么只预测噪声就够了
DDPM原论文推导了一个变分下界(ELBO),但最终用的损失函数极其简单:
[ L_{\text{simple}} = \mathbb{E}{t, x_0, \epsilon} \left[ | \epsilon - \epsilon\theta(x_t, t) |^2 \right] ]
翻译成人话:随机选一个时间步 (t),随机采一个噪声 (\epsilon),把 (x_0) 加噪成 (x_t),让网络预测这个噪声,然后算均方误差。
这个简化损失丢掉了ELBO里的加权系数,但实验效果反而更好。原因我理解是:加权系数会让网络偏向某些时间步,而简化损失让每个时间步同等重要,网络被迫在所有噪声水平上都学好。这有点像老师布置作业,如果只重点批改某几道题,学生就只练那几道;均匀批改,学生整体能力更均衡。
代码实现就几行:
def p_losses(denoise_model, x_start, t, noise=None): if noise is None: noise = torch.randn_like(x_start) x_noisy = q_sample(x_start=x_start, t=t, noise=noise) predicted_noise = denoise_model(x_noisy, t) loss = F.mse_loss(noise, predicted_noise) return loss提示:
q_sample利用重参数化技巧,可以一步从 (x_0) 得到 (x_t),不需要循环加噪1000次。公式是 (x_t = \sqrt{\bar{\alpha}_t} x_0 + \sqrt{1-\bar{\alpha}_t} \epsilon)。这个技巧是DDPM训练效率的关键。
3.3 UNet模型改进的常见方向
原版DDPM的UNet比较朴素,后续改进主要集中在几个方向:
注意力机制的引入。在较低分辨率层(如16x16、8x8)加入自注意力,让网络捕捉长距离依赖。高分辨率层用卷积,因为注意力在像素级计算量太大。这个混合策略在潜在扩散模型里被进一步优化。
残差块的重设计。原版用两层卷积加GroupNorm和SiLU激活。改进版会加入缩放残差、Dropout、或者用FiLM(Feature-wise Linear Modulation)注入时间步信息,比简单相加更灵活。
下采样和上采样的方式。原版用步长卷积下采样、转置卷积上采样。后来发现用抗锯齿下采样(先低通滤波再采样)能减少混叠,生成质量更高。上采样则常用最近邻插值加卷积,比转置卷积更少棋盘伪影。
通道数的配置。原版CIFAR-10用64、128、256三档,每档两个残差块。高分辨率图像生成需要更多层和更多通道,但显存是硬约束。我的经验是:从64通道起步,每下采样一次翻倍,到512或1024封顶,再大就上梯度检查点或混合精度。
class ResidualBlock(nn.Module): def __init__(self, in_channels, out_channels, time_emb_dim, dropout=0.1): super().__init__() self.norm1 = nn.GroupNorm(8, in_channels) self.conv1 = nn.Conv2d(in_channels, out_channels, 3, padding=1) self.time_mlp = nn.Sequential( nn.SiLU(), 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.dropout = nn.Dropout(dropout) 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.norm1(x) h = F.silu(h) h = self.conv1(h) t_emb = self.time_mlp(t_emb)[:, :, None, None] h = h + t_emb h = self.norm2(h) h = F.silu(h) h = self.dropout(h) h = self.conv2(h) return h + self.residual_conv(x)这个残差块里,时间步嵌入通过MLP映射后直接加到特征图上,是最简单也最常用的做法。我试过用FiLM做乘性调制,在小数据集上提升不明显,但训练更慢,所以默认还是用加法。
4. 完整实操流程:从零训练一个DDPM
4.1 环境准备与依赖安装
先明确环境:Python 3.9以上,PyTorch 2.0以上,CUDA 11.8或12.1。显存建议至少8GB,训练CIFAR-10级别(32x32)的DDPM,batch size 64大概占6GB;如果要训64x64或128x128,显存需求翻倍。
pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install einops matplotlib tqdmeinops用来做张量重排,比permute和view组合更可读;tqdm看训练进度;matplotlib用来可视化采样结果。
4.2 数据加载与预处理
以CIFAR-10为例,标准化到[-1, 1]区间,这和DDPM的噪声假设匹配(标准高斯噪声均值为0,方差为1,图像也应在类似尺度)。
from torchvision import datasets, transforms from torch.utils.data import DataLoader transform = transforms.Compose([ transforms.ToTensor(), transforms.Lambda(lambda t: (t * 2) - 1) ]) dataset = datasets.CIFAR10(root='./data', train=True, download=True, transform=transform) dataloader = DataLoader(dataset, batch_size=64, shuffle=True, num_workers=4, drop_last=True)注意:
drop_last=True很重要。如果最后一个batch只有1张图,BatchNorm或GroupNorm在训练时会报错或统计量异常。我踩过这个坑,训练到最后一个batch突然崩掉,排查了半天才发现是batch size为1导致的。
4.3 定义前向扩散与采样函数
把噪声调度的预计算封装成一个类,训练和采样都从这里取参数。
class DiffusionScheduler: def __init__(self, timesteps=1000, schedule='cosine'): self.timesteps = timesteps if schedule == 'linear': self.betas = torch.linspace(1e-4, 0.02, timesteps) else: self.betas = torch.tensor(cosine_beta_schedule(timesteps), dtype=torch.float32) self.alphas = 1.0 - self.betas self.alphas_cumprod = torch.cumprod(self.alphas, dim=0) self.sqrt_alphas_cumprod = torch.sqrt(self.alphas_cumprod) self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - self.alphas_cumprod) def q_sample(self, x_start, t, noise=None): if noise is None: noise = torch.randn_like(x_start) sqrt_alpha = self.sqrt_alphas_cumprod[t].view(-1, 1, 1, 1) sqrt_one_minus = self.sqrt_one_minus_alphas_cumprod[t].view(-1, 1, 1, 1) return sqrt_alpha * x_start + sqrt_one_minus * noise采样函数是DDPM推理的核心。标准采样需要循环1000步,每步调用一次UNet:
@torch.no_grad() def p_sample(model, x, t, t_index, scheduler): betas_t = scheduler.betas[t].view(-1, 1, 1, 1) sqrt_one_minus_alphas_cumprod_t = scheduler.sqrt_one_minus_alphas_cumprod[t].view(-1, 1, 1, 1) sqrt_recip_alphas_t = (1.0 / torch.sqrt(scheduler.alphas[t])).view(-1, 1, 1, 1) predicted_noise = model(x, t) model_mean = sqrt_recip_alphas_t * (x - betas_t * predicted_noise / sqrt_one_minus_alphas_cumprod_t) if t_index == 0: return model_mean else: posterior_variance_t = betas_t noise = torch.randn_like(x) return model_mean + torch.sqrt(posterior_variance_t) * noise这个采样公式来自DDPM论文的推导,核心是贝叶斯后验均值。t_index == 0时不加噪声,因为最后一步要输出确定性的结果。
4.4 训练循环与关键参数
训练循环本身很标准,但有几个参数需要特别注意:
model = UNet(in_channels=3, base_channels=64, channel_mults=(1, 2, 4), num_res_blocks=2) optimizer = torch.optim.AdamW(model.parameters(), lr=2e-4, weight_decay=1e-4) scheduler = DiffusionScheduler(timesteps=1000, schedule='cosine') for epoch in range(200): for step, (images, _) in enumerate(dataloader): images = images.cuda() t = torch.randint(0, 1000, (images.shape[0],), device=images.device).long() loss = p_losses(model, images, t, scheduler) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step()学习率用2e-4,配合AdamW和权重衰减1e-4。梯度裁剪阈值设1.0,防止早期训练梯度爆炸。我试过不裁剪,前几百步loss会突然飙到NaN,裁剪后稳定很多。
EMA(指数移动平均)是DDPM训练的一个关键技巧。维护一份模型参数的滑动平均副本,采样时用EMA参数而不是原始参数,生成质量明显更高:
class EMA: def __init__(self, model, decay=0.9999): self.model = model self.decay = decay self.shadow = {k: v.clone().detach() for k, v in model.state_dict().items()} def update(self): for k, v in self.model.state_dict().items(): self.shadow[k] = self.decay * self.shadow[k] + (1 - self.decay) * v def apply(self): self.model.load_state_dict(self.shadow)EMA的decay设0.9999,意味着平均窗口大约10000步。训练步数少的时候(比如几千步),decay可以降到0.999,否则EMA参数更新太慢,采样效果反而差。
4.5 采样与结果可视化
训练完成后,从纯噪声开始采样:
@torch.no_grad() def sample(model, scheduler, image_size=32, batch_size=16, channels=3): model.eval() x = torch.randn(batch_size, channels, image_size, image_size).cuda() for t_index in reversed(range(scheduler.timesteps)): t = torch.full((batch_size,), t_index, device=x.device, dtype=torch.long) x = p_sample(model, x, t, t_index, scheduler) return x1000步采样在单张RTX 3090上大概需要20到30秒生成16张32x32图。如果觉得慢,可以用DDIM采样,50步就能达到接近的质量,原理是跳步采样,后面会提。
5. 常见问题与排查技巧实录
5.1 训练loss不下降或震荡
这是最常见的问题。排查顺序如下:
| 现象 | 可能原因 | 解决方法 |
|---|---|---|
| loss从一开始就NaN | 学习率过大或数据未归一化 | 降学习率到1e-4,检查数据是否在[-1,1] |
| loss震荡剧烈 | batch size太小或梯度爆炸 | 增大batch size,加梯度裁剪 |
| loss下降后突然飙升 | 某个batch数据异常 | 检查数据加载,加drop_last=True |
| loss长期在0.1以上不降 | 模型容量不足或时间步嵌入有问题 | 增大通道数,检查时间步嵌入维度 |
我遇到过一次loss卡在0.15左右不动,排查发现是时间步嵌入的MLP输出维度写错了,导致时间信息根本没注入到残差块里。网络不知道当前是第几步,自然学不好。这种bug不会报错,只能靠检查张量形状发现。
5.2 采样结果全是噪声或模糊
训练loss正常但采样出来是噪声,通常有几个原因:
采样步数不够。DDPM需要完整的1000步反向过程,如果只跑100步,噪声没去干净。检查采样循环是否从timesteps-1到0完整执行。
EMA参数未正确应用。如果采样时用的是原始参数而不是EMA参数,结果会差很多。确认ema.apply()在采样前被调用。
噪声调度不匹配。训练用余弦调度,采样用线性调度,参数对不上,结果必然崩。训练和采样必须用同一个scheduler实例。
模型处于train模式。如果忘了model.eval(),Dropout和BatchNorm会引入随机性,采样结果不稳定。这个坑我踩过不止一次。
5.3 显存不足的优化策略
显存不够时,按以下优先级优化:
- 降低batch size。最直接,但太小会影响训练稳定性,建议不低于16。
- 混合精度训练。用
torch.cuda.amp自动混合精度,显存占用减少约40%,速度提升20%到30%。 - 梯度检查点。用
torch.utils.checkpoint包装残差块,用计算时间换显存,显存减少约50%,速度降低约20%。 - 减少UNet通道数或层数。最后考虑,因为会直接影响生成质量。
from torch.cuda.amp import autocast, GradScaler scaler = GradScaler() with autocast(): loss = p_losses(model, images, t, scheduler) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()混合精度训练时,注意把噪声调度里的常数也转成float16,否则会出现类型不匹配。我一般把scheduler的所有预计算张量都注册为buffer,自动跟随模型dtype。
5.4 采样速度太慢的加速方案
DDPM的1000步采样是硬伤。实际项目中,我会用以下加速方案:
DDIM采样。把随机采样变成确定性采样,步数从1000降到50到100,质量损失很小。核心改动是采样公式里去掉随机噪声项,改用预测的 (x_0) 和方向向量组合。
潜在扩散模型(LDM)。先用VAE把图像压缩到潜在空间(比如32x32x4),在潜在空间上跑扩散,最后解码回像素空间。计算量减少几十倍,这也是Stable Diffusion的基础架构。
蒸馏和一致性模型。训练一个学生网络,直接学习从噪声到图像的映射,一步或几步生成。这是目前研究的热点,但训练复杂度高,适合有充足算力的场景。
提示:如果只是做实验验证想法,DDIM 50步足够看效果。如果要做产品级应用,LDM是更务实的选择。DDPM原版更适合作为理解原理的起点。
6. 从DDPM到潜在扩散:实际项目中的选型建议
6.1 什么场景用原版DDPM
原版DDPM适合以下场景:图像尺寸小(32x32到64x64)、训练数据量中等(几万到几十万张)、算力有限(单卡或双卡)、目标是验证算法或做研究。CIFAR-10、CelebA、MNIST这些基准数据集上,原版DDPM能跑出不错的结果,训练几天就能收敛。
如果图像尺寸到128x128以上,原版DDPM的显存和采样时间会变得难以接受。我试过在256x256上训原版DDPM,单卡24GB显存只能跑batch size 8,采样1000步要几分钟,实用性很低。
6.2 潜在扩散模型的改造要点
LDM的核心改动是引入一个VAE编码器-解码器。编码器把图像压缩成潜在表示,扩散过程在潜在空间进行,解码器把潜在表示还原成图像。
改造步骤:
- 训练或加载一个预训练VAE。编码器下采样8倍,通道数从3变成4。
- 把UNet的输入输出通道从3改成4,其他结构不变。
- 扩散过程的所有操作在潜在空间进行,图像尺寸变成原来的1/8。
- 采样完成后,用VAE解码器还原图像。
这个改造让计算量减少约64倍(8x8空间下采样),同时生成质量几乎不降。Stable Diffusion就是在这个框架上加了文本条件、交叉注意力等模块。
6.3 UNet使用时的注意事项
不管原版还是LDM,UNet的使用有几个通用注意事项:
GroupNorm的组数要能整除通道数。比如通道数64,GroupNorm组数设8或16都行,设7会报错。我习惯用8,小通道数时用4。
下采样和上采样的次数要匹配。编码器下采样3次,解码器就要上采样3次,否则特征图尺寸对不上,跳跃连接会报错。
注意力层的位置。一般放在16x16和8x8分辨率,太高分辨率加注意力显存吃不消。如果显存充足,可以在32x32也加一层,但收益递减。
时间步嵌入的维度。通常设为基础通道数的4倍,比如基础通道64,时间嵌入维度256。太小信息容量不够,太大增加参数量但收益有限。
6.4 训练数据量与模型容量的平衡
DDPM是数据饥渴型模型。CIFAR-10的5万张图,训到FID 10左右需要几十万步。如果数据量只有几千张,模型很容易过拟合,生成结果就是训练集的复制。
我的经验是:数据量少于1万张时,要么用强数据增强(随机裁剪、翻转、颜色抖动),要么用预训练模型微调,要么降低模型容量。从零训练一个UNet在几千张图上,效果通常不如微调一个在大型数据集上预训练的模型。
数据增强在DDPM里有个细节:增强后的图像要保证在[-1,1]范围内,且不能引入非高斯噪声。随机翻转和裁剪是安全的,颜色抖动要控制幅度,否则会干扰噪声预测任务。
7. 一些实操中攒下来的经验
训练DDPM最耗时的不是写代码,而是等结果。我一般会在训练脚本里加一个定期采样回调,每5000步生成16张图存到本地,用tqdm的进度条旁边打印当前loss和采样图的路径。这样不用等训练结束就能判断模型有没有学歪。
另一个实用技巧是从预训练权重初始化。如果要做类似数据集的生成任务,加载一个在大型数据集上训过的DDPM权重,冻结编码器部分,只微调解码器和时间嵌入,收敛速度快很多。我试过在CIFAR-10预训练权重上微调CelebA,10000步就能出可看的人脸,从零训练至少要50000步。
最后说一个采样时的细节:采样batch size不要设太大。虽然理论上可以并行生成很多张,但显存占用是线性的。我一般设16或32,生成一批存一批,避免OOM。如果要做大规模生成,写个循环分批采样,比一次性生成几百张更稳。
这个方向后续可以扩展的地方很多:条件生成(类别、文本、姿态)、超分辨率、图像编辑、视频生成。DDPM是这些应用的共同底座,把底座打牢,上面的东西学起来会快很多。