1. 项目概述:从零构建DDPM的动机与价值
最近在复现一些经典的生成模型,发现Denoising Diffusion Probabilistic Models(DDPM)虽然论文公式看着有点唬人,但当你真正动手把它从零搭出来,会发现其背后的思想异常清晰和优雅。很多朋友可能通过Stable Diffusion等应用已经体验了扩散模型的强大,但对其核心引擎——UNet网络以及时间嵌入、扩散调度等机制的理解可能还停留在“黑箱”调用层面。这次,我们就用PyTorch,不依赖任何高级扩散模型库,彻底拆解并复现一个完整的DDPM。这个过程不仅能让你深刻理解“噪声如何一步步变成图像”,更能让你掌握自定义扩散模型、调整生成过程的能力,比如修改采样步数、尝试不同的噪声调度器,甚至为UNet加入注意力机制。无论你是想深入AIGC领域,还是单纯对概率模型和深度学习结合感兴趣,这个从零开始的搭建之旅都会让你收获颇丰。
2. 核心理论拆解:DDPM的前向与逆向过程
要搭建模型,必须先吃透它的工作原理。DDPM包含两个核心过程:前向扩散过程和逆向去噪过程。
2.1 前向扩散过程:逐步添加噪声
前向过程是一个固定的马尔可夫链,它逐步向一张原始图片x0添加高斯噪声。这个过程是预先定义好的,不包含任何可学习的参数。在每一步t(从1到T,T是总步数,比如1000),我们根据上一步的数据x_{t-1}得到当前加噪后的数据x_t。其数学形式如下:
q(x_t | x_{t-1}) = N(x_t; sqrt(1 - β_t) * x_{t-1}, β_t * I)
这里的β_t是一个在0到1之间预先定义好的序列,称为噪声调度表。它决定了每一步添加的噪声量,通常随着t增大而增大。N表示高斯分布。这个公式的意思是:x_t的均值是sqrt(1 - β_t) * x_{t-1},方差是β_t。由于每一步都只依赖前一步,我们可以推导出一个非常实用的性质:可以从原始图像x0直接计算出任意中间时刻t的加噪图像x_t,而不需要一步步迭代。
q(x_t | x_0) = N(x_t; sqrt(ᾱ_t) * x_0, (1 - ᾱ_t) * I)
其中,α_t = 1 - β_t,ᾱ_t = Π_{s=1}^{t} α_s。这个性质是DDPM实现高效训练的关键。在代码中,这意味着我们可以在训练时,随机选择一个时间步t,然后直接用这个公式采样出对应的x_t。
注意:
β_t序列的设计至关重要。通常使用线性或余弦调度。线性调度可能导致在过程早期或晚期噪声变化过于剧烈,而余弦调度(如Improved DDPM中提出的)能让噪声添加更平滑,往往能带来更好的生成效果。我们复现时会实现这两种。
2.2 逆向去噪过程:神经网络学习“反扩散”
如果前向过程是把一幅画慢慢涂成纯噪声,那么逆向过程就是试图从纯噪声中一步步还原出那幅画。这是一个从x_T(纯高斯噪声)到x_0的生成过程。其每一步也定义为一个高斯分布:
p_θ(x_{t-1} | x_t) = N(x_{t-1}; μ_θ(x_t, t), Σ_θ(x_t, t))
这里的μ_θ和Σ_θ是由神经网络参数化的均值和方差。DDPM原文做了一个简化,将方差Σ_θ固定为与β_t相关的常数,只让神经网络学习均值μ_θ。那么,神经网络学的是什么呢?经过一番推导(这里不展开复杂公式),可以发现,预测均值μ_θ等价于预测在x_t和t的条件下,前向过程中所添加的噪声ε。
因此,DDPM的训练目标变得极其简洁:训练一个噪声预测网络ε_θ。给定任意时间步t和对应的加噪图像x_t,网络的目标是预测出添加到x_0上从而得到x_t的那个噪声ε。损失函数就是预测噪声和真实噪声之间的均方误差(MSE)。
L = E_{x_0, t, ε} [ || ε - ε_θ(x_t, t) ||^2 ]
其中ε是从标准高斯分布中采样的随机噪声。这个简单的目标使得训练非常稳定。
2.3 训练与采样算法
理解了上述理论,训练和采样的伪代码就一目了然:
训练循环:
- 从数据集中采样一个干净图像
x_0。 - 从
{1, ..., T}中均匀采样一个时间步t。 - 从标准高斯分布采样噪声
ε。 - 根据公式
x_t = sqrt(ᾱ_t) * x_0 + sqrt(1 - ᾱ_t) * ε计算加噪图像。 - 将
x_t和t输入噪声预测网络ε_θ,得到预测的噪声ε_θ。 - 计算
ε和ε_θ之间的MSE损失,反向传播更新网络参数。
采样(生成)循环:
- 从标准高斯分布采样一个随机噪声
x_T。 - 从
t = T到t = 1循环: a. 将当前的x_t和t输入网络ε_θ,得到预测噪声。 b. 根据预测噪声和公式,计算x_{t-1}的均值。 c. 根据固定方差采样一些额外噪声(用于随机性)。 d. 计算得到x_{t-1}。 - 循环结束后,
x_0即为生成的图像。
3. 核心模块一:时间步嵌入
在DDPM中,时间步t是一个标量,但我们需要将它转化为网络能够利用的条件信息。直接输入标量t效果很差,因此需要将其嵌入到一个高维向量空间。这里我们采用Transformer中提出的正弦位置编码的变体。
3.1 正弦位置编码原理
其思想是为每个时间步t生成一个唯一的高维向量,并且这个向量能反映时间的顺序关系(即t和t+1的嵌入向量是相似的)。公式如下:
对于嵌入向量的第i个维度:emb(t)[i] = sin(ω_i * t)如果i是偶数emb(t)[i] = cos(ω_i * t)如果i是奇数
其中,ω_i = 1 / (10000^(2i / d)),d是嵌入向量的总维度。
这种编码方式能确保不同时间步的嵌入具有区分度,同时其内积能反映时间步的接近程度。
3.2 PyTorch实现与集成
在我们的UNet中,时间步t是一个整数(例如250)。我们首先通过一个nn.Embedding层将其映射为一个初始向量,然后通过一个由线性层和SiLU激活函数组成的小型MLP,将其投影到与UNet中间特征图通道数相匹配的维度。这个最终的时间条件向量会被加到UNet的各个残差块的特征上,通常是通过特征图的通道维度相加或自适应组归一化(AdaGN)来实现。
import torch import torch.nn as nn import math class SinusoidalPositionEmbeddings(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 class TimeEmbedding(nn.Module): def __init__(self, time_dim, projection_dim): super().__init__() self.time_mlp = nn.Sequential( SinusoidalPositionEmbeddings(time_dim), nn.Linear(time_dim, projection_dim), nn.SiLU(), nn.Linear(projection_dim, projection_dim), ) def forward(self, t): return self.time_mlp(t)实操心得:嵌入维度
time_dim通常设置为256或512就足够了。projection_dim需要与UNet中应用时间条件的特征图通道数对齐。在实际添加时,我更喜欢使用“自适应组归一化”(AdaGN),它将时间嵌入向量通过线性层映射为组归一化(GroupNorm)的缩放因子gamma和偏移因子beta,然后应用于归一化后的特征上。这种方式比简单相加的条件注入方式更强大,能更有效地指导网络在不同时间步的行为。
4. 核心模块二:UNet网络架构
UNet是DDPM的“心脏”,负责根据带噪图像x_t和时间步t预测噪声ε。它是一个编码器-解码器结构,带有跳跃连接。
4.1 基础构建块:残差块与注意力块
我们的UNet由两种基本块堆叠而成:残差块和注意力块。
残差块:每个残差块包含两个卷积层,中间有组归一化和SiLU激活函数。时间条件信息(来自时间嵌入)通过AdaGN注入到第一个归一化层之后。跳跃连接确保梯度流动。
注意力块:为了提升模型对图像全局结构的建模能力,我们在UNet的底层(特征图分辨率较低时)插入自注意力或交叉注意力块。这里我们实现一个简单的单头自注意力机制。由于注意力机制的计算复杂度与特征图尺寸的平方成正比,因此只在下采样后的低分辨率特征上使用是计算可行的。
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 * 2) # 输出gamma和beta self.block1 = nn.Sequential( nn.GroupNorm(8, in_channels), nn.SiLU(), nn.Conv2d(in_channels, out_channels, 3, padding=1), ) self.block2 = nn.Sequential( nn.GroupNorm(8, out_channels), nn.SiLU(), 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.block1(x) # 自适应组归一化 gamma, beta = self.time_mlp(t_emb).chunk(2, dim=1) h = h * (gamma[:, :, None, None] + 1) + beta[:, :, None, None] # 第二部分 h = self.block2(h) # 残差连接 return h + self.residual_conv(x) class AttentionBlock(nn.Module): def __init__(self, channels): super().__init__() self.norm = nn.GroupNorm(8, channels) self.qkv = nn.Conv2d(channels, channels * 3, 1) self.proj_out = nn.Conv2d(channels, channels, 1) def forward(self, x): b, c, h, w = x.shape q, k, v = self.qkv(self.norm(x)).chunk(3, dim=1) # 重塑为 (b, c, h*w) 并转置k q = q.view(b, c, -1).transpose(1, 2) # (b, h*w, c) k = k.view(b, c, -1) # (b, c, h*w) v = v.view(b, c, -1).transpose(1, 2) # (b, h*w, c) # 注意力分数 attn = torch.bmm(q, k) * (c ** -0.5) # (b, h*w, h*w) attn = F.softmax(attn, dim=-1) # 加权求和 out = torch.bmm(attn, v) # (b, h*w, c) out = out.transpose(1, 2).view(b, c, h, w) # 恢复形状 return x + self.proj_out(out) # 残差连接4.2 完整的UNet组装
完整的UNet由下采样路径(编码器)和上采样路径(解码器)组成,中间有跳跃连接。下采样通过步长为2的卷积或池化实现,上采样通过转置卷积或最近邻插值+卷积实现。时间嵌入向量在每个分辨率级别的残差块中注入。
class UNet(nn.Module): def __init__(self, in_channels=3, out_channels=3, base_channels=64, time_emb_dim=256): super().__init__() self.time_embedding = TimeEmbedding(time_emb_dim, time_emb_dim*4) # 下采样 self.down1 = ResidualBlock(in_channels, base_channels, time_emb_dim) self.down2 = nn.Sequential( nn.Conv2d(base_channels, base_channels, 3, stride=2, padding=1), # 下采样 ResidualBlock(base_channels, base_channels*2, time_emb_dim), AttentionBlock(base_channels*2), # 在低分辨率特征上加注意力 ) # ... 可以继续添加更多下采样层 # 中间层 self.mid = nn.Sequential( ResidualBlock(base_channels*4, base_channels*4, time_emb_dim), AttentionBlock(base_channels*4), ResidualBlock(base_channels*4, base_channels*4, time_emb_dim), ) # 上采样 # ... 上采样层,与下采样对称,包含转置卷积和跳跃连接 self.up1 = nn.Sequential( ResidualBlock(base_channels*4 + base_channels*2, base_channels*2, time_emb_dim), # 跳跃连接拼接通道 AttentionBlock(base_channels*2), nn.ConvTranspose2d(base_channels*2, base_channels, 2, stride=2), # 上采样 ) self.up2 = ResidualBlock(base_channels*2, base_channels, time_emb_dim) # 再次拼接跳跃连接 self.out = nn.Sequential( nn.GroupNorm(8, base_channels), nn.SiLU(), nn.Conv2d(base_channels, out_channels, 3, padding=1), ) def forward(self, x, t): t_emb = self.time_embedding(t) # 下采样并保存特征用于跳跃连接 h1 = self.down1(x, t_emb) h2 = self.down2(h1, t_emb) # ... 中间层 h_mid = self.mid(h2, t_emb) # 上采样并拼接跳跃连接 h = self.up1(torch.cat([h_mid, h2], dim=1), t_emb) h = self.up2(torch.cat([h, h1], dim=1), t_emb) return self.out(h)注意事项:UNet的通道数配置(如
base_channels=64)需要根据你的计算资源和图像分辨率调整。对于64x64的图片,上述简化结构可能足够;对于256x256或更高分辨率,需要更深的网络和更多的通道数。跳跃连接是UNet的关键,它帮助解码器恢复在编码器中丢失的空间细节信息。
5. 核心模块三:扩散调度器
扩散调度器定义了前向过程中β_t序列,以及与之相关的α_t和ᾱ_t序列。它不参与训练,但在训练(计算x_t)和采样(计算x_{t-1})时被频繁使用。
5.1 线性调度与余弦调度
线性调度:这是DDPM原论文使用的方案。β_t从β_start(如0.0001)线性增长到β_end(如0.02)。β_t = β_start + (t/T) * (β_end - β_start)
余弦调度:由Improved DDPM论文提出,旨在改善线性调度在过程两端变化过快的问题。它直接定义ᾱ_t:ᾱ_t = f(t) / f(0), 其中f(t) = cos((t/T + s) / (1+s) * π/2)^2这里s是一个小偏移(如0.008),防止t接近T时ᾱ_t过小导致数值不稳定。
余弦调度通常能产生视觉质量更高、更平滑的生成样本。
5.2 调度器的实现与缓存
由于α_t,ᾱ_t等序列在训练和采样中需要反复使用,我们应在初始化时预先计算并缓存它们,避免重复计算。
import torch import numpy as np class DDPMScheduler: def __init__(self, num_timesteps=1000, beta_start=1e-4, beta_end=0.02, schedule='linear'): self.num_timesteps = num_timesteps self.schedule = schedule if schedule == 'linear': self.betas = torch.linspace(beta_start, beta_end, num_timesteps) elif schedule == 'cosine': steps = num_timesteps + 1 x = torch.linspace(0, num_timesteps, steps) alphas_cumprod = torch.cos(((x / num_timesteps) + 0.008) / 1.008 * torch.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) else: raise NotImplementedError self.alphas = 1. - self.betas self.alphas_cumprod = torch.cumprod(self.alphas, dim=0) # ᾱ_t self.sqrt_alphas_cumprod = torch.sqrt(self.alphas_cumprod) self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1. - self.alphas_cumprod) # 为采样过程计算参数 self.sqrt_recip_alphas = torch.sqrt(1.0 / self.alphas) self.posterior_variance = self.betas * (1. - self.alphas_cumprod[:-1]) / (1. - self.alphas_cumprod[1:]) def add_noise(self, original_samples, noise, timesteps): # 根据公式 x_t = sqrt(ᾱ_t) * x_0 + sqrt(1-ᾱ_t) * ε 添加噪声 sqrt_alpha_prod = self.sqrt_alphas_cumprod[timesteps].view(-1, 1, 1, 1) sqrt_one_minus_alpha_prod = self.sqrt_one_minus_alphas_cumprod[timesteps].view(-1, 1, 1, 1) noisy_samples = sqrt_alpha_prod * original_samples + sqrt_one_minus_alpha_prod * noise return noisy_samples def step(self, model_output, timestep, sample): # 根据预测的噪声 ε_θ,计算 x_{t-1} t = timestep beta_t = self.betas[t].view(-1, 1, 1, 1) sqrt_one_minus_alpha_cumprod_t = self.sqrt_one_minus_alphas_cumprod[t].view(-1, 1, 1, 1) sqrt_recip_alpha_t = self.sqrt_recip_alphas[t].view(-1, 1, 1, 1) # 公式:x_{t-1}的均值 = 1/sqrt(α_t) * (x_t - β_t/sqrt(1-ᾱ_t) * ε_θ) pred_original_sample = sqrt_recip_alpha_t * (sample - beta_t * model_output / sqrt_one_minus_alpha_cumprod_t) mean = pred_original_sample if t > 0: noise = torch.randn_like(sample) variance = (1 - self.alphas_cumprod[t-1]) / (1 - self.alphas_cumprod[t]) * self.betas[t] std = torch.sqrt(variance).view(-1, 1, 1, 1) else: std = 0. noise = 0. prev_sample = mean + std * noise return prev_sample实操心得:
add_noise函数用于训练时构造输入x_t。step函数用于采样时,根据网络预测的噪声,从x_t反推x_{t-1}。注意在t=0时,方差应为0,因为此时应得到确定的x_0。缓存所有张量到设备(CPU/GPU)上能显著加速训练和采样循环。
6. 训练流程完整实现
将上述所有模块组合起来,就构成了完整的训练流程。我们以在CIFAR-10(32x32)数据集上训练为例。
6.1 数据准备与加载
import torch from torch.utils.data import DataLoader from torchvision import datasets, transforms from torchvision.transforms import ToTensor, Lambda, Compose import torch.nn.functional as F # 数据预处理 transform = Compose([ transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)) # 将图像归一化到[-1, 1] ]) train_dataset = datasets.CIFAR10(root='./data', train=True, download=True, transform=transform) train_loader = DataLoader(train_dataset, batch_size=128, shuffle=True, num_workers=4)6.2 训练循环代码
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') model = UNet(in_channels=3, out_channels=3, base_channels=64, time_emb_dim=256).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4) scheduler = DDPMScheduler(num_timesteps=1000, schedule='cosine') mse_loss = nn.MSELoss() num_epochs = 200 gradient_accumulation_steps = 2 # 梯度累积,模拟更大batch size optimizer.zero_grad() for epoch in range(num_epochs): model.train() total_loss = 0 for step, (clean_images, _) in enumerate(train_loader): clean_images = clean_images.to(device) batch_size = clean_images.shape[0] # 1. 采样随机时间步和噪声 timesteps = torch.randint(0, scheduler.num_timesteps, (batch_size,), device=device).long() noise = torch.randn_like(clean_images) # 2. 根据时间步和噪声,对干净图像加噪,得到 x_t noisy_images = scheduler.add_noise(clean_images, noise, timesteps) # 3. 模型预测噪声 noise_pred = model(noisy_images, timesteps) # 4. 计算损失(预测噪声与真实噪声的MSE) loss = mse_loss(noise_pred, noise) loss = loss / gradient_accumulation_steps # 梯度累积 loss.backward() # 5. 梯度累积步骤完成后更新参数 if (step + 1) % gradient_accumulation_steps == 0: torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) # 梯度裁剪,防止爆炸 optimizer.step() optimizer.zero_grad() total_loss += loss.item() * gradient_accumulation_steps if step % 100 == 0: print(f"Epoch {epoch}, Step {step}, Loss: {loss.item() * gradient_accumulation_steps:.4f}") avg_loss = total_loss / len(train_loader) print(f"Epoch {epoch} finished. Average Loss: {avg_loss:.4f}") # 可选:每个epoch结束后保存一次模型检查点 if epoch % 10 == 0: torch.save({ 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'loss': avg_loss, }, f'ddpm_checkpoint_epoch_{epoch}.pth')注意事项:
- 归一化:输入图像被归一化到
[-1, 1],模型输出的噪声预测也在同一范围。在最终生成图像时,需要反归一化到[0, 1]。- 学习率:1e-4是一个比较安全的起点。可以使用学习率预热(Warmup)和余弦衰减(Cosine Annealing)来优化训练。
- 梯度累积:当GPU内存不足以支撑大的batch size时,梯度累积是有效的技巧。它通过多次前向传播累积梯度,再一次性更新参数,等效于增大了batch size。
- 梯度裁剪:扩散模型训练通常比较稳定,但梯度裁剪可以作为一个额外的安全措施,防止训练后期出现梯度爆炸。
7. 采样与图像生成
训练好模型后,我们就可以从随机噪声开始,运行逆向过程来生成图像。
7.1 采样循环实现
@torch.no_grad() def sample(model, scheduler, image_size, batch_size=16, channels=3, device='cuda'): """从随机噪声生成图像""" model.eval() # 1. 初始化随机噪声 x_T img = torch.randn((batch_size, channels, image_size, image_size), device=device) # 2. 从 T 到 1 循环采样 for t in reversed(range(scheduler.num_timesteps)): # 创建当前时间步的张量,形状为 (batch_size,) timesteps = torch.full((batch_size,), t, device=device, dtype=torch.long) # 3. 预测噪声 predicted_noise = model(img, timesteps) # 4. 使用调度器计算前一步的 x_{t-1} img = scheduler.step(predicted_noise, t, img) # 可选:显示中间过程(例如每100步保存一次) # if t % 100 == 0: # save_image(img, f'sample_step_{t}.png') # 5. 将生成的图像从 [-1, 1] 反归一化到 [0, 1] img = (img.clamp(-1, 1) + 1) / 2.0 return img # 使用示例 generated_images = sample(model, scheduler, image_size=32, batch_size=16, device=device) # 保存或显示图像 from torchvision.utils import save_image save_image(generated_images, 'generated_samples.png', nrow=4)7.2 加速采样技巧:DDIM
上述采样过程需要迭代完整的T步(如1000步),这很耗时。Denoising Diffusion Implicit Models (DDIM) 提出了一种在保持生成质量的同时,大幅减少采样步数的方法。其核心思想是定义一个非马尔科夫的逆向过程,允许跳过一些中间步骤。
DDIM的采样公式与DDPM不同,它允许我们定义一个子序列{τ_1, τ_2, ..., τ_S},其中S可以远小于T。采样时,我们只在这些子时间步上运行模型。在我们的代码中,只需实现一个DDIM调度器的step函数即可替换原来的采样循环。
class DDIMScheduler(DDPMScheduler): def step(self, model_output, timestep, sample, eta=0.0): # eta=0 对应DDIM确定性采样,eta=1 对应DDPM随机采样 t = timestep prev_t = t - self.num_timesteps // self.num_inference_steps # 假设我们定义了推理步数 alpha_prod_t = self.alphas_cumprod[t] alpha_prod_t_prev = self.alphas_cumprod[prev_t] if prev_t >= 0 else torch.tensor(1.0) beta_prod_t = 1 - alpha_prod_t beta_prod_t_prev = 1 - alpha_prod_t_prev # 预测 x_0 pred_original_sample = (sample - beta_prod_t ** 0.5 * model_output) / alpha_prod_t ** 0.5 # 计算 x_{t-1} 的方向 pred_sample_direction = (1 - alpha_prod_t_prev - eta ** 2 * beta_prod_t_prev) ** 0.5 * model_output # 计算 x_{t-1} prev_sample = alpha_prod_t_prev ** 0.5 * pred_original_sample + pred_sample_direction if eta > 0: noise = torch.randn_like(model_output) variance = (1 - alpha_prod_t_prev) / (1 - alpha_prod_t) * beta_prod_t std = eta * variance ** 0.5 prev_sample = prev_sample + std * noise return prev_sample使用DDIM,我们可以用50步甚至20步就获得与1000步DDPM采样相媲美的质量,极大提升了生成效率。
8. 常见问题、调试技巧与效果优化
在实际搭建和训练过程中,你肯定会遇到各种问题。下面是我踩过的一些坑和总结的经验。
8.1 训练不稳定或损失不下降
- 检查数据归一化:确保输入图像和模型输出在预期的范围内(通常是[-1,1])。一个常见的错误是输入了[0,1]的图像但没做归一化,或者输出层用了错误的激活函数(如Sigmoid)。
- 检查时间嵌入:确保时间步
t被正确嵌入并注入到UNet的每一层。可以打印中间特征图,看时间条件是否有效改变了特征分布。 - 学习率过高:扩散模型对学习率比较敏感。尝试从较低的学习率(如1e-5)开始,配合Warmup。
- 梯度爆炸/消失:使用梯度裁剪(
clip_grad_norm_)和检查模型初始化。UNet中的卷积层可以使用He初始化或Xavier初始化。 - 损失值范围:MSE损失在训练初期应该在0.9左右(因为预测随机噪声),然后缓慢下降。如果损失从一开始就非常大或非常小,可能是计算
x_t的公式有误。
8.2 生成的图像模糊或有噪声
- 训练不充分:扩散模型需要很长的训练时间才能收敛。在CIFAR-10上,可能需要200-500个epoch才能看到清晰的图像。确保训练了足够的轮数。
- 噪声调度问题:尝试从线性调度切换到余弦调度。余弦调度通常能产生更清晰、细节更丰富的图像。
- 模型容量不足:对于更大分辨率(如128x128)的图像,基础的64通道UNet可能不够。尝试增加通道数(如128)或加深网络层数。
- 采样步数不足:如果使用DDPM采样,确保步数足够(如1000步)。如果使用DDIM加速,可以尝试增加推理步数(如100步),或调整
eta参数(eta=0为确定性采样,通常更清晰;eta=1更随机)。
8.3 计算资源与性能优化
- 混合精度训练:使用
torch.cuda.amp进行自动混合精度训练,可以显著减少GPU内存占用并加快训练速度,尤其对于大型UNet模型。 - 梯度检查点:如果GPU内存严重不足,可以在UNet的某些层使用
torch.utils.checkpoint,以时间换空间。 - 多GPU训练:使用
nn.DataParallel或nn.DistributedDataParallel进行多卡训练,可以加快数据吞吐。
8.4 可视化与监控
- 监控损失曲线:使用TensorBoard或WandB记录训练损失。一个健康的训练曲线应该是平滑下降的。
- 定期采样:每隔一定训练步数或epoch,运行一次采样函数,将生成的图像保存下来。这是判断模型是否在学习的最直观方式。你可以观察到图像从噪声逐渐变得清晰的过程。
- 检查点管理:定期保存模型检查点,不仅保存模型参数,也保存优化器状态和当前epoch,方便从中断处恢复训练,或选择不同阶段的模型进行采样比较。
从零搭建DDPM是一个系统工程,涉及理论理解、模块实现、训练调试和效果优化多个环节。当你看到第一张由自己编写的代码生成的、清晰的图片时,那种成就感是无与伦比的。这个过程中积累的对扩散模型每个细节的掌控力,是直接调用高级API无法比拟的。希望这份详细的指南能帮助你顺利走完这段旅程,并为你后续探索更复杂的扩散模型(如条件生成、Latent Diffusion等)打下坚实的基础。