DDPM实战:从零构建PyTorch去噪扩散模型
1. 从理论到代码:为什么你需要亲手实现一个DDPM
如果你对扩散模型(Diffusion Model)感兴趣,甚至已经看过几篇原理讲解,感觉“懂了”,但一打开GitHub上那些动辄几千行的开源项目,看到里面复杂的模块、继承关系和工程化封装,瞬间又觉得无从下手——那么,你现在的感受,我完全理解。几年前我第一次接触DDPM(Denoising Diffusion Probabilistic Models)时,也是这种感觉。原理公式推导得头头是道,但代码就是看不懂,更别提自己从头写一个了。
这篇文章就是为你准备的。我们不谈复杂的数学推导,那些网上已经有很多优秀的资料。我们只聚焦一件事:如何用PyTorch,从零开始,一行一行地构建一个能真正跑起来、能生成图片的DDPM模型。 我的目标是,让你读完这篇文章,不仅能理解代码的每一部分在做什么,更能自己动手,搭建出一个可以训练MNIST手写数字的扩散模型。
为什么一定要“从零构建”?因为只有亲手把每个模块像搭积木一样组装起来,你才能真正掌握扩散模型的核心运作机制。你会明白,那些看起来神秘的“前向过程”、“反向采样”,本质上就是一系列张量操作和系数提取;你会清楚,训练一个扩散模型,数据是如何流动的,损失是如何计算的。这种实践带来的理解,远比只看论文或别人的代码要深刻得多。
我将会用一个高度模块化、结构清晰的GaussianDiffusion类作为主线,带你逐步实现。这个类封装了扩散模型训练和采样的所有核心逻辑。我们会从两个最基础的工具函数extract和EMA讲起,然后深入到GaussianDiffusion的初始化、前向加噪、损失计算,最后完成反向采样生成图像的全过程。过程中,我会分享我实际编码时踩过的坑和调试技巧,比如系数计算容易出错的地方、训练时Loss不下降怎么办、采样效果不好如何调整等。
准备好了吗?让我们打开编辑器,开始这场从理论到实战的旅程。相信我,当你看到自己写的代码从一片随机噪声中逐步“画”出一个清晰的数字时,那种成就感是无与伦比的。
2. 搭建基石:两个不可或缺的辅助工具
在开始构建核心的扩散模型类之前,我们需要先准备好两个“瑞士军刀”式的小工具。它们代码量不大,但在整个DDPM的实现中会被反复使用,理解它们能让你后面的路走得非常顺畅。
2.1 核心工具:extract函数
这个函数可能是你理解DDPM代码的第一个小门槛。它的作用,我用一个生活化的比喻来解释:想象你有一长串按照顺序排列的调料包(比如糖、盐、醋、酱油……),每个调料包对应做菜的一个步骤。现在你要做第t步,就需要从这长串里精准地拿出第t个调料包,并且把它变成适合你锅里菜量的形状(比如从一大包变成一小勺)。
在DDPM中,我们有一系列预先计算好的系数张量,比如sqrt_alphas_cumprod、sqrt_one_minus_alphas_cumprod等,它们的长度等于扩散步数T(比如1000步)。在训练或采样的每一步,我们都需要根据当前的时间步t(一个0到T-1的整数),从这些长系数张量中取出对应的那个值,并且把它扩展成和当前图像数据x一样的形状,以便进行逐元素的乘法运算。
这就是extract函数干的事情。我们来看看代码:
def extract(a, t, x_shape):
"""
从a中提取t位置的数据并且reshape成x_shape的形状返回
:param a: 系数张量,形状为 (num_timesteps,)
:param t: 时间步索引张量,形状为 (batch_size,)
:param x_shape: 目标数据的形状,例如 (batch_size, channels, height, width)
:return: 提取并reshape后的系数,形状与x_shape匹配
"""
b, *_ = t.shape
out = a.gather(-1, t)
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
逐行拆解:
b, *_ = t.shape:这行代码获取了批次大小b。*_表示忽略t张量后面的其他维度(这里t通常是(b,)的一维张量)。out = a.gather(-1, t):这是关键操作。a.gather(dim, index)会沿着dim维度,根据index张量里的索引值,从a中收集数据。这里dim=-1表示最后一个维度。假设a的形状是(1000,),t是[5, 12, 999](b=3),那么gather操作就会取出a[5],a[12],a[999],out的形状就是(3,)。return out.reshape(b, *((1,) * (len(x_shape) - 1))):接下来我们要把取出的这b个值,变成能和输入数据x(形状为x_shape)进行广播计算的形状。x_shape可能是(b, c, h, w)。len(x_shape)-1的值是3(代表c, h, w三个维度)。(1,)*3得到(1,1,1)。最终reshape成(b, 1, 1, 1)。这样,当我们用这个系数乘以形状为(b, c, h, w)的x时,PyTorch的广播机制会自动让这个系数作用于x的每一个通道、每一个像素点,完美匹配。
一个具体例子: 假设batch_size=2,图像是2x28x28的MNIST图片(x_shape = (2, 1, 28, 28)),当前时间步t = tensor([100, 200])。sqrt_alphas_cumprod是一个有1000个元素的张量。extract(sqrt_alphas_cumprod, t, x_shape)会取出第100和第200个系数,然后将其形状变为(2, 1, 1, 1)。这样我们就可以直接用这个形状的系数去乘x了。
2.2 训练稳定器:EMA(指数移动平均)
EMA不是扩散模型独有的,它在很多深度学习模型训练中都有应用,用来平滑模型权重,获得一个更稳定、泛化性更好的“影子模型”。你可以把它理解为你模型权重的一个“慢速追随者”。每次你的主模型(current_model)更新后,EMA模型(ema_model)的权重并不会立刻变成主模型的样子,而是会缓慢地向主模型靠拢,融合了历史权重信息。
为什么在DDPM中特别有用?因为扩散模型的训练过程相对“嘈杂”,每一步预测的噪声都有波动。使用EMA模型在采样(生成)时,往往能获得更清晰、更稳定的结果。它的实现非常直观:
class EMA():
"""EMA(指数移动平均)优化器"""
def __init__(self, decay):
# decay是衰减率,通常设置为0.999或0.9999,越接近1,历史权重占比越大,更新越平滑。
self.decay = decay
def update_average(self, old, new):
# 更新单个权重值:新影子权重 = decay * 旧影子权重 + (1 - decay) * 新主权重
if old is None:
return new
return old * self.decay + (1 - self.decay) * new
def update_model_average(self, ema_model, current_model):
# 遍历模型的所有参数,对整个模型的权重进行EMA更新
for current_params, ema_params in zip(current_model.parameters(), ema_model.parameters()):
old, new = ema_params.data, current_params.data
ema_params.data = self.update_average(old, new)
使用技巧: 在实际训练中,我们通常不会一开始就启用EMA。因为模型初期权重变化剧烈,过早使用EMA可能会拖慢学习速度。常见的做法是设置一个ema_start步数(比如2000步),在达到这个步数之前,EMA模型直接复制主模型的权重;在这之后,再开始进行EMA平滑更新。我们会在后面的GaussianDiffusion类中看到这个逻辑。
3. 构建核心:GaussianDiffusion类的初始化
现在,我们进入正题,开始构建DDPM的核心类GaussianDiffusion。这个类将像一个总控制器,管理着扩散模型训练和生成的所有环节。初始化是这个类最复杂也最关键的一步,因为它要预先计算好扩散过程所需的所有数学系数。
3.1 构造函数参数详解
我们先来看__init__方法需要哪些参数,这能帮你理解搭建一个扩散模型需要准备什么:
def __init__(self, model, input_shape, input_channels, betas, device,
num_class=None, loss_type="l2", ema_decay=0.9999,
ema_start=2000, ema_update_rate=1):
model: 这是你的神经网络,通常是一个U-Net。它的任务是接收一个带噪声的图片x_t和时间步t,预测出加入的噪声epsilon。注意:这个网络的输出形状必须和输入图片x_t完全一致。input_shape: 输入数据的空间形状。这里有个容易混淆的点:在PyTorch中,图像张量默认是(batch, channels, height, width)。但在这个实现里,input_shape指的是(height, width),即(H, W)。如果你的数据是(batch, channels, height, width),那么input_shape就应该传入(height, width)这个元组。input_channels: 图像的通道数C。对于MNIST灰度图是1,对于RGB彩色图是3。betas: 这是扩散模型超参数的核心,一个长度为T(扩散总步数)的一维数组或列表。它定义了每一步噪声的方差大小,通常是从一个很小的值(如0.0001)线性增长到一个较大的值(如0.02)。这个序列的设定直接影响生成质量。device: 训练设备,cuda或cpu。扩散模型计算量大,强烈建议使用GPU。num_class: 用于条件生成,比如你想生成指定数字(0-9)的MNIST图片。如果是无条件生成,设为None即可。loss_type: 损失函数类型,"l1"(L1损失)或"l2"(均方误差MSE)。原始DDPM论文使用的是简单的MSE损失,效果就很好。ema_decay,ema_start,ema_update_rate: 这三个是EMA相关的参数,控制平滑更新的强度、开始步数和更新频率。
3.2 核心系数计算与注册
初始化函数里最“硬核”的部分,就是根据传入的betas,计算并注册一系列在后续公式中反复使用的系数。这些系数一旦算好就固定不变,前向和反向过程直接查表使用,极大提高了效率。
# 1. 计算alpha系列
alphas = 1.0 - betas # alpha_t = 1 - beta_t
alphas_cumprod = np.cumprod(alphas) # alpha_bar_t = 累乘(alpha_1 ... alpha_t)
# 2. 将numpy数组转为PyTorch张量,并放到指定设备上
to_torch = partial(torch.tensor, dtype=torch.float32, device=device)
# 3. 使用register_buffer注册为模型的“缓冲区”
# 这些张量会成为模型的一部分,随模型移动设备,但不参与梯度更新。
self.register_buffer("betas", to_torch(betas))
self.register_buffer("alphas", to_torch(alphas))
self.register_buffer("alphas_cumprod", to_torch(alphas_cumprod))
# 4. 计算并注册派生系数
self.register_buffer("sqrt_alphas_cumprod", to_torch(np.sqrt(alphas_cumprod))) # sqrt(alpha_bar_t)
self.register_buffer("sqrt_one_minus_alphas_cumprod", to_torch(np.sqrt(1. - alphas_cumprod))) # sqrt(1 - alpha_bar_t)
self.register_buffer("reciprocal_sqrt_alphas", to_torch(np.sqrt(1. / alphas))) # sqrt(1 / alpha_t)
self.register_buffer("remove_noise_coeff", to_torch(betas / np.sqrt(1. - alphas_cumprod))) # 反向去噪公式中的系数
self.register_buffer("sigma", to_torch(np.sqrt(betas))) # 反向采样时添加的噪声标准差
这些系数分别有什么用?
sqrt_alphas_cumprod和sqrt_one_minus_alphas_cumprod: 用于前向加噪过程。根据公式x_t = sqrt(alpha_bar_t) * x_0 + sqrt(1-alpha_bar_t) * epsilon,我们可以直接用这两个系数,将原始图片x_0和随机噪声epsilon混合,得到第t步的带噪图片x_t。reciprocal_sqrt_alphas和remove_noise_coeff: 用于反向去噪过程。根据公式x_{t-1} = 1/sqrt(alpha_t) * (x_t - beta_t/sqrt(1-alpha_bar_t) * epsilon_theta) + sigma_t * z,我们需要用remove_noise_coeff乘以预测的噪声epsilon_theta,然后用reciprocal_sqrt_alphas进行缩放。sigma: 同样是反向过程公式的一部分,当t > 0时,需要在去噪后添加一点随机噪声z,其标准差就是sigma_t(即sqrt(beta_t))。
把这些系数一次性算好存起来,后面无论是训练还是采样,都只需要做简单的extract和乘加运算,代码会非常简洁高效。这是实现DDPM的一个关键技巧。
4. 前向过程:加噪与损失计算
理解了初始化时准备的那些“弹药”,我们现在来看DDPM是如何训练的。训练的核心思想其实非常直观:教一个网络如何“去噪”。
4.1 加噪函数 perturb_x
这个函数实现了前向扩散过程的核心公式。给定一张干净图片x_0、一个时间步t和一张随机噪声noise,它计算出第t步的带噪图片x_t。
def perturb_x(self, x, t, noise):
"""从x0获取加噪的xt的噪声图"""
return (
extract(self.sqrt_alphas_cumprod, t, x.shape) * x +
extract(self.sqrt_one_minus_alphas_cumprod, t, x.shape) * noise
)
代码解读:
x: 这里是原始干净图片x_0,形状为(batch, channels, height, width)。t: 一个整数张量,形状为(batch,),表示这一批图片中每个样本对应的扩散时间步。注意,在训练时,我们会对每个批次中的每张图片随机采样一个不同的t。noise: 从标准正态分布中采样的随机噪声,形状和x完全相同。extract(...) * x: 根据t取出对应的sqrt(alpha_bar_t),并将其形状扩展为(batch, 1, 1, 1),然后与x_0相乘。这部分保留了原始图像的信息。extract(...) * noise: 同样,根据t取出对应的sqrt(1 - alpha_bar_t),与噪声noise相乘。这部分是添加的噪声。- 两者相加,就得到了
x_t。当t很大时,sqrt(alpha_bar_t)接近0,sqrt(1 - alpha_bar_t)接近1,x_t就几乎完全是随机噪声了。
这个过程是确定性的,不需要学习。它模拟了数据在T步内逐渐被噪声破坏的过程。
4.2 损失计算 get_losses 与 forward
有了加噪的x_t,我们就可以计算损失了。DDPM的损失函数设计得很巧妙:它让神经网络model去预测我们加入的噪声epsilon。
def get_losses(self, x, t, y):
"""
:param x: 原始干净图片 x0
:param t: 时间步
:param y: 条件标签(如无条件生成则为None)
:return: 预测噪声与真实噪声之间的损失
"""
# 1. 生成与x同形状的标准高斯噪声
noise = torch.randn_like(x)
# 2. 根据公式,计算加噪后的图片 xt
perturbed_x = self.perturb_x(x, t, noise)
# 3. 让神经网络预测噪声。输入是加噪图xt和时间步t。
estimated_noise = self.model(perturbed_x, t, y) # y可能用于条件生成
# 4. 计算预测噪声和真实噪声之间的差异
if self.loss_type == "l1":
loss = F.l1_loss(estimated_noise, noise)
elif self.loss_type == "l2":
loss = F.mse_loss(estimated_noise, noise)
return loss
def forward(self, x, y=None):
B, N, C = x.shape # 注意这里的形状假设,N是展平后的像素数?通常我们处理(B, C, H, W)
device = x.device
# 为批次中的每张图片随机采样一个时间步t
t = torch.randint(0, self.num_timesteps, (B,), device=device)
return self.get_losses(x, t, y)
这里有几个非常重要的细节和易错点:
- 时间步
t的输入:神经网络model需要知道当前处理的是第几步的加噪图片。因此,我们需要将时间步t编码后(例如通过正弦位置编码)作为额外信息输入网络。这是扩散模型网络设计的一个关键点,通常我们会实现一个TimestepEmbedding模块。 - 损失的意义:这个简单的MSE损失(预测噪声 vs 真实噪声)是DDPM工作的核心。通过最小化这个损失,网络逐渐学会了如何从任意噪声水平
t的图片x_t中,估计出被加入的噪声。一旦网络学会了这个,在生成时,我们就可以从纯噪声x_T开始,一步步使用网络预测的噪声来“去噪”,最终得到干净图片x_0。 forward函数的形状:注意示例代码中B, N, C = x.shape,这暗示输入x的形状可能是(batch, height*width, channels),即展平后的形式。但在图像处理中,我们更常用(batch, channels, height, width)。你需要根据自己数据的实际形状和网络结构来调整。一致性是关键,确保perturb_x、model的输入输出、损失计算等所有环节的形状都匹配。
训练循环就是不断地调用这个forward函数,计算损失,然后反向传播更新model的参数。同时,每隔一定步数,调用update_ema函数来更新EMA模型的权重。
5. 反向过程:从噪声中采样生成图像
训练完成后,最激动人心的部分来了:如何用训练好的模型,从一片随机噪声中“创造”出新的图像?这就是反向采样过程,它像是把前向加噪的录像带倒着播放。
5.1 单步去噪 remove_noise
这是采样循环中最核心的一步,它根据公式,利用网络预测的噪声,从x_t计算出x_{t-1}。
@torch.no_grad() # 采样时不需要计算梯度,节省内存和计算资源
def remove_noise(self, x, t, y, use_ema=True):
"""
从xt中获取xt-1
:param x: 当前时刻的带噪图像 xt
:param t: 当前时间步
:param y: 条件标签
:param use_ema: 是否使用更平滑的EMA模型进行预测
:return: 去噪后的图像 xt-1
"""
# 选择使用EMA模型还是原始模型进行噪声预测
model_to_use = self.ema_model if use_ema else self.model
# 调用模型预测噪声 epsilon_theta
predicted_noise = model_to_use(x, t, y)
# DDPM论文中的去噪公式:
# x_{t-1} = 1 / sqrt(alpha_t) * ( x_t - beta_t / sqrt(1 - alpha_bar_t) * epsilon_theta ) + sigma_t * z
# 其中 z ~ N(0, I),当 t > 0 时添加。
# 计算公式的主体部分:x_t - coeff * predicted_noise
x_t_minus_noise = x - extract(self.remove_noise_coeff, t, x.shape) * predicted_noise
# 乘以系数 1 / sqrt(alpha_t)
x_prev = x_t_minus_noise * extract(self.reciprocal_sqrt_alphas, t, x.shape)
return x_prev
代码逻辑解析:
@torch.no_grad():这是一个装饰器,表示在这个函数中的所有计算都不会构建计算图,不保存梯度。这在推理(采样)阶段至关重要,可以大幅减少GPU内存占用。- 模型选择:我们提供了一个
use_ema开关。在训练后期,EMA模型通常比原始模型更稳定,生成的图像质量更好。所以默认在采样时使用ema_model。 - 核心计算:代码完全对应了论文中的去噪公式。
remove_noise_coeff对应beta_t / sqrt(1 - alpha_bar_t),reciprocal_sqrt_alphas对应1 / sqrt(alpha_t)。通过extract函数取出对应时间步t的系数,然后进行张量运算。 - 注意:这个函数返回的
x_prev,还没有添加随机噪声sigma_t * z。添加随机噪声的步骤是在外层的采样循环中控制的,这是因为在最后一步(t=0)时,我们不再添加噪声。
5.2 完整采样循环 sample
现在,我们将remove_noise函数放入一个从T到0的循环中,就构成了完整的图像生成流程。
@torch.no_grad()
def sample(self, batch_size, device, y=None, use_ema=True):
"""
从DDPM中采样生成最终图像
:param batch_size: 要生成的图片数量
:param device: 生成设备
:param y: 条件标签(用于条件生成)
:param use_ema: 是否使用EMA模型
:return: 生成的图片,形状为 (batch_size, channels, height, width)
"""
# 1. 从标准高斯分布中采样初始噪声 x_T
x = torch.randn(batch_size, self.input_channels, *self.input_shape, device=device)
# 注意:这里我调整了形状顺序为 (B, C, H, W),更符合PyTorch惯例。
# 2. 从T-1步开始,逐步去噪,直到0步
for t in tqdm(range(self.num_timesteps - 1, -1, -1), desc="Sampling: ", total=self.num_timesteps):
# 为批次中所有样本创建相同的时间步t
t_batch = torch.tensor([t], device=device).repeat(batch_size)
# 3. 调用remove_noise,得到去噪后的图像(此时还未加噪声z)
x = self.remove_noise(x, t_batch, y, use_ema=use_ema)
# 4. 如果当前不是最后一步(t>0),则添加随机噪声
if t > 0:
# 获取当前步的噪声标准差 sigma_t
sigma_t = extract(self.sigma, t_batch, x.shape)
# 采样随机噪声 z ~ N(0, I)
z = torch.randn_like(x)
# 按照公式添加噪声
x = x + sigma_t * z
# 5. 循环结束,x 就是生成的最终图像 x_0
return x.cpu().detach() # 移回CPU并脱离计算图,方便后续保存或显示
采样过程可视化:
你可以把sample函数想象成一个雕刻家的过程。一开始,你只有一块完全随机的大理石(x_T,纯噪声)。雕刻家(我们的模型)根据一个从粗到细的计划(从t=T-1到t=0),一步步凿掉多余的石头(噪声)。在每一步(t),雕刻家先根据当前石头的形状(x_t)和计划步骤(t),决定凿掉哪部分(remove_noise)。但为了保持一点随机性,让每次雕刻的作品都独一无二,在凿完之后(除了最后一步),他还会轻轻地、随机地敲打一下石头表面(添加噪声sigma_t * z)。经过T个步骤,一块精美的雕像(x_0,清晰的图像)就诞生了。
实用技巧:
tqdm进度条:采样过程通常需要很多步(如1000步),使用tqdm可以直观地看到生成进度。sample_diffusion_sequence:这个函数与sample几乎一样,区别在于它把每一步去噪后的x_t都保存到一个列表里并返回。这非常有用!你可以通过可视化这个序列,看到图像是如何从噪声一步步变得清晰的,这对于调试模型和理解扩散过程非常有帮助。- 批次采样:
batch_size可以大于1,一次性生成多张图片。但要注意,生成过程是串行的,且需要为每张图片存储中间状态,显存占用会随batch_size线性增长。
6. 实战演练:训练一个MNIST数字生成器
理论说了这么多,是时候动手了。让我们把上面的代码片段组装起来,并补全神经网络部分,训练一个能生成手写数字的DDPM模型。
6.1 构建一个简单的噪声预测网络
我们不需要一开始就实现复杂的U-Net。为了快速验证流程,我们可以用一个非常简单的多层感知机(MLP)来作为噪声预测模型。只要它的输入输出形状匹配,就能工作。
import torch.nn as nn
import torch.nn.functional as F
import math
class SimpleNoisePredictor(nn.Module):
"""一个简单的全连接网络,用于预测噪声。适用于展平后的图像数据。"""
def __init__(self, input_dim, hidden_dim=512, timestep_embed_dim=128):
super().__init__()
self.input_dim = input_dim # 展平后图像的维度,例如 28*28=784
self.timestep_embed_dim = timestep_embed_dim
# 时间步编码层
self.timestep_encoder = nn.Sequential(
nn.Linear(timestep_embed_dim, hidden_dim),
nn.SiLU(),
nn.Linear(hidden_dim, hidden_dim),
)
# 主网络
self.net = nn.Sequential(
nn.Linear(input_dim + hidden_dim, hidden_dim), # 将图像特征和时间编码拼接
nn.SiLU(),
nn.Linear(hidden_dim, hidden_dim),
nn.SiLU(),
nn.Linear(hidden_dim, hidden_dim),
nn.SiLU(),
nn.Linear(hidden_dim, input_dim), # 输出预测的噪声,形状与输入图像相同
)
# 初始化时间步正弦位置编码
self.register_buffer('pos_embedding', self._build_pos_embedding(timestep_embed_dim))
def _build_pos_embedding(self, dim, max_period=10000):
# 创建正弦位置编码,用于编码时间步t
half = dim // 2
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half)
return freqs # 这里返回的是频率,实际编码在forward中完成
def timestep_embed(self, t, dim):
# t: (batch,)
half = dim // 2
# 将t转换为与freqs相同设备
t = t.float().to(self.pos_embedding.device)
args = t[:, None].float() * self.pos_embedding[None, :]
embedding = torch.cat([torch.sin(args), torch.cos(args)], dim=-1)
if dim % 2:
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
return embedding
def forward(self, x, t, y=None):
# x: (batch, input_dim) 展平后的图像
# t: (batch,) 时间步
# 1. 对时间步t进行编码
t_emb = self.timestep_embed(t, self.timestep_embed_dim)
t_emb = self.timestep_encoder(t_emb) # (batch, hidden_dim)
# 2. 将图像特征和时间编码拼接
x = x.view(x.size(0), -1) # 确保x是展平的
combined = torch.cat([x, t_emb], dim=-1) # (batch, input_dim + hidden_dim)
# 3. 通过网络预测噪声
predicted_noise = self.net(combined) # (batch, input_dim)
# 将输出reshape回图像的空间形状(如果需要的话,这里输出是展平的)
# 在这个简单例子中,我们直接返回展平的噪声,在损失计算时与同样展平的真实噪声比较。
return predicted_noise
这个网络的关键点在于时间步编码。扩散模型需要知道当前处理的是哪个时间步的加噪图片,因此我们必须将标量时间步t编码成一个高维向量,并与图像特征融合。这里使用了Transformer中常见的正弦余弦位置编码。
6.2 组装训练流程
现在,我们有了GaussianDiffusion类和SimpleNoisePredictor网络。接下来就是标准的PyTorch训练循环。
import torch
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
from tqdm import tqdm
# 1. 准备数据
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,)) # MNIST灰度图,归一化到[-1, 1]
])
train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform)
train_loader = DataLoader(train_dataset, batch_size=128, shuffle=True, num_workers=4)
# 2. 定义模型和扩散过程参数
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
img_shape = (28, 28) # MNIST图像大小
channels = 1
input_dim = 28 * 28 # 展平后的维度
# 定义beta调度表,线性从0.0001到0.02
num_timesteps = 1000
betas = torch.linspace(0.0001, 0.02, num_timesteps).to(device)
# 实例化噪声预测网络和扩散模型
model = SimpleNoisePredictor(input_dim=input_dim).to(device)
diffusion = GaussianDiffusion(
model=model,
input_shape=img_shape, # 注意:这里传入的是(H, W)
input_channels=channels,
betas=betas,
device=device,
loss_type='l2'
).to(device)
# 3. 定义优化器
optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)
# 4. 训练循环
num_epochs = 50
for epoch in range(num_epochs):
epoch_loss = 0.0
pbar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{num_epochs}')
for batch_idx, (images, _) in enumerate(pbar):
# 将图像数据移动到设备,并展平
x = images.to(device) # (batch, 1, 28, 28)
x_flat = x.view(x.size(0), -1) # (batch, 784) 为了适配我们的简单网络
# 前向传播,计算损失
loss = diffusion(x_flat) # diffusion.forward会调用get_losses
# 反向传播和优化
optimizer.zero_grad()
loss.backward()
optimizer.step()
# 更新EMA模型
diffusion.update_ema()
epoch_loss += loss.item()
pbar.set_postfix({'loss': loss.item()})
avg_loss = epoch_loss / len(train_loader)
print(f'Epoch {epoch+1} Average Loss: {avg_loss:.6f}')
# 每隔几个epoch,采样并保存生成的图片,观察训练效果
if (epoch + 1) % 10 == 0:
with torch.no_grad():
# 注意:采样时,扩散模型期望输入形状是 (B, C, H, W)
# 我们的SimpleNoisePredictor需要展平输入,但diffusion.sample内部会处理好。
# 我们需要调整diffusion的初始化,使其内部处理与网络匹配。
# 更简单的做法:修改网络,使其接受 (B, C, H, W) 输入。
# 这里为了示例,我们假设已经调整好了。
sampled_images = diffusion.sample(batch_size=16, device=device, use_ema=True)
# sampled_images 形状应为 (16, 1, 28, 28)
# ... 保存或显示sampled_images的代码 ...
训练注意事项:
- 数据形状:这是最大的坑之一。确保你的数据形状、网络输入输出形状、扩散模型内部计算形状三者完全一致。上面的例子中,简单网络处理展平数据,而扩散模型的一些计算(如
perturb_x)可能期望空间维度。你需要调整网络或数据预处理来匹配。一个更稳妥的方法是,让噪声预测网络直接处理(B, C, H, W)形状的数据,比如使用一个微型的CNN。 - Loss下降:DDPM训练初期Loss下降很快,但很快会进入平台期。不要慌,这是正常的。扩散模型需要很长时间来学习精细的去噪。只要Loss在缓慢下降或小幅震荡,就让它继续训练。
- 采样观察:一定要定期(比如每10或20个epoch)采样生成图片,这是判断模型是否在学习的唯一可靠方法。如果生成的图片始终是噪声,可能是代码有bug、学习率不合适或训练时间不够。
- 显存:即使对于MNIST这样的小图片,
T=1000的扩散模型训练也会占用不少显存。如果遇到CUDA out of memory,可以尝试减小batch_size。
6.3 从简单MLP到U-Net
一旦你用简单网络跑通了整个流程,理解了数据流向和训练感觉,就可以将SimpleNoisePredictor替换成更强大的U-Net了。U-Net的架构非常适合图像去噪任务,因为它有编码器-解码器结构,能捕捉多尺度特征。网上有很多DDPM U-Net的实现,核心要点是:
- 在每一层(或关键层)注入时间步嵌入(
timestep_embed)。 - 使用自注意力机制(特别是在低分辨率层)来建模全局依赖。
- 可能使用条件归一化(如GroupNorm with adaptive affine)来融合时间信息。
替换网络结构后,你可能会发现生成质量有显著提升。但训练复杂度也会增加,需要更仔细地调参。
7. 避坑指南与经验分享
最后,结合我自己的实战经验,分享一些在实现和训练DDPM时容易遇到的问题和技巧。这些“坑”很多是教程里不会细说的,但却能决定你的项目成败。
1. 系数计算错误
这是最隐蔽的bug。betas、alphas_cumprod等系数的计算必须和论文公式严格一致。一个常见的错误是sqrt_one_minus_alphas_cumprod计算成1 - sqrt_alphas_cumprod(这是错误的,应该是sqrt(1 - alphas_cumprod))。建议将计算系数的代码单独拿出来,用一个小脚本,对比几个时间步t的手动计算结果和代码输出,确保完全一致。
2. 时间步嵌入不当
网络预测不准,很多时候问题出在时间步t没有有效地传递给网络。确保你的timestep_embed函数能产生区分度足够的不同时间步编码。可以将不同t的编码向量画出来看看,它们应该是平滑变化的。另外,时间步编码的维度不能太小,通常128或256维是安全的起点。
3. 损失不下降或生成全是噪声
- 检查数据归一化:DDPM通常假设输入数据在
[-1, 1]范围内。确保你的数据预处理(transforms.Normalize)是正确的。 - 检查噪声预测目标:在
get_losses函数里,确保你计算的是predicted_noise和true_noise的损失,而不是predicted_noise和x或perturbed_x的损失。 - 学习率:尝试不同的学习率,比如
1e-4,5e-5,2e-5。扩散模型有时对学习率很敏感。 - 网络容量:你的网络可能太简单,无法学习复杂的去噪映射。尝试增加网络深度或宽度。
- 训练步数:扩散模型就是需要长时间训练。对于MNIST,可能几千个epoch才开始有模糊的形状,上万epoch才能清晰。要有耐心。
4. 采样结果模糊 这是DDPM的一个已知特点,尤其是使用MSE损失时,它倾向于生成“平均”的、看起来模糊的图像。可以尝试:
- 使用
ema_model进行采样,通常更清晰。 - 调整采样过程,使用更少的步数(如DDIM采样)或修改噪声调度(
betas),但这属于进阶话题。 - 尝试使用
L1损失,有时能产生更锐利的结果。
5. 显存爆炸
- 降低
batch_size:这是最直接有效的方法。 - 使用梯度累积:如果想让有效
batch_size更大,可以累积多个小批次的梯度后再更新。 - 混合精度训练:使用
torch.cuda.amp进行自动混合精度训练,可以显著减少显存占用并加速训练。 - 检查
sample_diffusion_sequence:这个函数会保存所有中间结果,如果T很大(如1000),会占用大量内存。仅在需要可视化时使用。
6. 调试技巧
- 可视化前向过程:取一张训练图片,手动调用
perturb_x,对不同的t(如0, 10, 50, 100, 500, 999)进行加噪并显示。你应该能看到图片从清晰逐渐变成纯噪声。 - 可视化采样过程:使用
sample_diffusion_sequence,把生成的全过程(从噪声到清晰图)保存为GIF。这能直观地看到模型是否在有效地去噪。 - 检查中间变量:在训练和采样时,打印或记录关键张量(如
sqrt_alphas_cumprod[t],predicted_noise的均值/标准差)的统计信息,确保它们在合理的范围内。
实现DDPM就像解一道精密的数学题,每一步都需要准确无误。但一旦你亲手把它搭建起来并看到它成功运行,那种对扩散模型原理豁然开朗的感觉,以及对代码掌控力的提升,是仅仅阅读论文无法比拟的。希望这份详细的指南能帮你跨过从理论到实践的门槛,开启你的生成式AI创作之旅。如果在实现过程中遇到具体问题,多检查数据流、多打印中间状态,耐心调试,你一定能成功。
更多推荐
所有评论(0)