ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

从零实现DDPM:用PyTorch训练扩散模型生成CIFAR-10图像

从零实现DDPM:用PyTorch训练扩散模型生成CIFAR-10图像 我在过去的几个月里被GAN折腾得不轻。WGAN-GP、SAGAN、BigGAN这些我都试过模型结构加了所有能加的花活结果训练时还是会时不时出现判别器Loss直接飞升、生成器Loss清零这种让人抓狂的局面。后来我转向扩散模型用PyTorch从零写了第一版DDPM在CIFAR-10上只训练了一小会儿出来的图像就已经能看到清晰的轮廓——那一刻你就会明白为什么越来越多的人决定不再死磕GAN了。DDPMDenoising Diffusion Probabilistic Models去噪扩散概率模型是近年来生成模型领域的一个重要流派。它不再依赖生成器与判别器的对抗博弈而是用一条简单的马尔可夫链先在数据上逐步加噪声、再让神经网络逐步去噪硬生生把一张标准高斯噪声还原成一张真实感很强的图片。这篇文章我会把一个可运行的PyTorch实现从原理到代码完整过一遍数据集用CIFAR-10所有代码都能直接跑适合对GAN已经有点腻、想换个新思路的读者也适合正在入门生成模型、想彻底搞懂扩散模型底层逻辑的学习者。你不需要多深的数学背景只需要会用PyTorch做基础训练就能跟上。1. DDPM原理速通先搞清楚它到底在学什么1.1 正向过程给真实图像一点点加噪声在DDPM里我们先把一张真实图像当作x0。正向过程就是从x0出发按照一个预先设定的噪声调度beta_t逐步往图像里加入标准高斯噪声得到x1, x2, ..., xT。当T足够大比如默认的1000步xT几乎就是一张纯噪声完全看不出原始图像的内容。这里的加噪不是每一步都从上一张图采样而是可以直接用重参数化公式从x0一步跳到任意第t步的图像。这是DDPM训练效率高的关键之一。设alpha_t 1 - beta_talpha_bar_t表示从第1步到第t步的alpha连乘那么q(x_t | x_0) N(x_t; sqrt(alpha_bar_t) * x0, (1 - alpha_bar_t) * I)展开写就是x_t sqrt(alpha_bar_t) * x0 sqrt(1 - alpha_bar_t) * epsilon其中epsilon ~ N(0, I)。这个公式意味着我们不需要循环t步去慢慢加噪而是拿到x0后一次就能算出任意一步的噪声图训练效率非常高。在代码里实现这个公式非常简单但有一个细节容易错sqrt_alpha_bar_t和sqrt_one_minus_alpha_bar_t要预先算好并且要避免小数点精度引起的NaN。我习惯把所有调度参数放在一个列表里然后用torch.tensor一次性转到目标设备索引时保持维度一致后面会给完整的写法。1.2 逆向过程让网络学会“去噪”如果我们知道逆向分布q(x_{t-1} | x_t)就能从纯噪声一步步还原出真实图像。但这个逆向分布直接不可求因为需要知道原始数据分布。于是DDPM的做法是训练一个神经网络让它去拟合逆向分布。经过推导可以证明当beta_t足够小时每一步逆向分布也近似服从高斯分布。DDPM论文最核心的简化结论是训练目标可以等价地写成“让网络预测每一步加的噪声epsilon”。换句话说网络输入x_t和时间步t输出预测噪声epsilon_theta(x_t, t)损失函数就是MSE(epsilon_theta(x_t, t), epsilon)。这个损失简单得惊人。你可能会问为什么是预测噪声而不是直接预测x_{t-1}因为在这种参数化下目标分布的形状更平滑训练起来更稳这也是“别再死磕GAN”背后的一个重要原因——DDPM的训练目标本质上就是一个回归问题没有对抗博弈没有纳什均衡梯度自然不会忽大忽小。从另一个角度看DDPM的做法和“去噪”这个日常概念是一致的。照片拍糊了我们会想办法去除噪声DDPM则把这个过程反着玩先把照片彻底弄糊再让网络学怎么一步步还原。网络在1000个不同的噪声等级上都做了一遍预测噪声的题目学到的就不只是某一类去噪而是整条退化路径的逆过程所以生成时才能真正从纯噪声开始重建图像。1.3 为什么DDPM训练比GAN稳GAN最让人头疼的地方在于生成器和判别器是在玩一个零和博弈训练过程本质上是在找一个鞍点。一旦双方能力不匹配生成器就会陷入模式坍塌或者判别器收敛得太快导致生成器梯度消失。你想让两个网络维持一种微妙的动态平衡就得不断调学习率、调正则化、调容量稍有不慎模型就崩了。DDPM完全没有这个问题。它的核心任务是“给定第t步噪声图预测原始噪声”网络永远在做同一个回归任务只是输入噪声程度不同。你不需要担心模式坍塌不需要精心调节判别器与生成器的容量比例也不需要担心模型会突然“飞掉”。我从项目实践里感受到的差异非常明显DDPM基本一上来的Loss就在稳定下降只要数据管道和超参数没有重大错误第一天跑就能得到有结构的图像轮廓而GAN可能调了半天熵还是黑的。这背后还有个更深层的原因GAN的优化目标是非凸的极大极小问题收敛性高度依赖两个网络的博弈过程而DDPM的训练目标其实是变分下界的一个重参数化形式是标准的似然训练整个优化过程是单调、可控的。这也是为什么在近两年的图像生成应用里扩散模型逐渐成了更受信赖的底座而不是哪个GAN的变体。2. 环境准备与CIFAR-10数据管道2.1 PyTorch环境搭建与设备检查先说环境我这里用的是PyTorch 2.x以及配套的torchvision。若你是从零开始用conda建一个独立环境最省心conda create -n ddpm python3.10 conda activate ddpm pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install matplotlib tensorboard如果下载太慢可以把pip的源换成国内镜像。装完后先检查一下CUDA是否可用import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name() if torch.cuda.is_available() else CPU)如果打印出False就代表你目前是CPU环境。CPU也能跑CIFAR-10只是训练时间会拉得很长。建议先把batch size和模型通道数调小一点先跑通流程再考虑效率。下面所有代码我都会把device写成自动检测CPU也能运行。2.2 CIFAR-10数据加载与预处理CIFAR-10是32×32的彩色小图数据集共10类一共6万张。用它跑DDPM的好处有两个一是分辨率低模型能快速收敛二是torchvision自带了数据集不需要自己爬图。数据预处理要做的很简单随机水平翻转做数据增强转成Tensor把像素值归一化到[0, 1]再加一个标准化均值方差取0.5把数据变到[-1, 1]。这里要特别记住最后采样生成的图片是在[-1, 1]范围保存或可视化前必须反归一化否则看到的图像会发灰。很多人踩过这个坑一开始以为模型训练失败其实只是忘了把像素值映射回[0, 1]。import torch from torch.utils.data import DataLoader from torchvision import datasets, transforms transform transforms.Compose([ transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)), ]) train_dataset datasets.CIFAR10(root./data, trainTrue, downloadTrue, transformtransform) train_loader DataLoader(train_dataset, batch_size128, shuffleTrue, num_workers4, drop_lastTrue)2.3 超参数设定与项目结构DDPM有几个关键超参数扩散总步数T、beta调度、batch size、学习率和EMA衰减系数。我建议第一版直接用一套被验证过的参数先跑通再调不要一上来就自己发明配置。config { image_size: 32, channels: 3, T: 1000, beta_start: 1e-4, beta_end: 0.02, lr: 2e-4, batch_size: 128, epochs: 100, device: cuda if torch.cuda.is_available() else cpu, }beta选择线性调度是初学最稳妥的选择。当然你也可以用cosine调度它在细节纹理上往往更好一点但在CIFAR-10上两者差异不算大。项目文件我建议拆成config.py、model.py、ddpm.py、train.py、sample.py这样后续扩展到自己的数据集会方便很多。如果不拆文件至少把超参数、模型、加噪逻辑、采样逻辑分开放在不同的cell或函数里后面调试会省很多事。3. U-Net模型设计用PyTorch写一个能感知时间步的U-Net3.1 为什么扩散模型要用U-NetDDPM的目标是从带噪图像中预测噪声这就要求网络既能捕捉低频结构又能保留高频细节。U-Net的编码器-解码器结构配合跳跃连接天然适合这种任务。除此之外网络还必须知道当前“加了多大噪声”否则同样一个模糊色块在t100和t800时含义完全不同所以我们需要把时间步信息注入到网络内部。这里我给U-Net设计了三个核心组件残差块、自注意力块、时间嵌入。CIFAR-10分辨率低网络不需要太深三层下采样足够。网络越深不见得效果越好尤其在32×32这种小分辨率上太深的网络反而容易过拟合。3.2 时间嵌入与残差块实现时间嵌入和Transformer里的positional embedding一样用正弦余弦来编码时间步。简单来说把整数t映射成一个高维向量再通过两个线性层变成每一层可用的条件向量。import math import torch import torch.nn as nn def sinusoidal_embedding(t, dim): half dim // 2 freqs torch.exp(-math.log(10000) * torch.arange(half, devicet.device) / half) args t[:, None] * freqs[None, :] return torch.cat([torch.sin(args), torch.cos(args)], dim-1)残差块则参考DDPM原始实现两层卷积加GroupNorm中间把时间嵌入加到特征图上。注意时间嵌入要经过激活函数后再映射到与卷积输出相同的通道数否则注入效果很弱。import torch.nn.functional as F class ResBlock(nn.Module): def __init__(self, in_ch, out_ch, time_dim): super().__init__() self.norm1 nn.GroupNorm(8, in_ch) self.conv1 nn.Conv2d(in_ch, out_ch, 3, padding1) self.norm2 nn.GroupNorm(8, out_ch) self.conv2 nn.Conv2d(out_ch, out_ch, 3, padding1) self.time_proj nn.Linear(time_dim, out_ch) self.shortcut nn.Conv2d(in_ch, out_ch, 1) if in_ch ! out_ch else nn.Identity() def forward(self, x, t): h F.silu(self.norm1(x)) h self.conv1(h) h F.silu(self.norm2(h)) h self.conv2(h) h h self.time_proj(F.silu(t))[:, :, None, None] return h self.shortcut(x)这里有两个容易被忽视的细节。第一GroupNorm的组数我固定为8这是DDPM官方和多数实现中常见的配置。如果特征图通道数少于8GroupNorm会报错所以你把通道数调得很小时要一起调低组数。第二shortcut是恒等映射时直接用input通道数不一致时才用1×1卷积对齐千万别在每个ResBlock里都无条件加一个1×1卷积否则参数量会白白增加很多。3.3 自注意力块与完整U-Net代码在较低分辨率上加上自注意力可以让网络建立长距离依赖。CIFAR-10图像太小我通常只在16×16和8×8这两个分辨率上加注意力太高分辨率的注意力反而会增加显存开销收益却不明显。下面是自注意力块的实现做得非常简洁1×1卷积得到Q、K、V然后做标准缩放点积注意力输出再经过1×1卷积并残差连接。class AttnBlock(nn.Module): def __init__(self, ch): super().__init__() self.q nn.Conv2d(ch, ch, 1) self.k nn.Conv2d(ch, ch, 1) self.v nn.Conv2d(ch, ch, 1) self.proj nn.Conv2d(ch, ch, 1) def forward(self, x): B, C, H, W x.shape q self.q(x).view(B, C, H * W).permute(0, 2, 1) k self.k(x).view(B, C, H * W) v self.v(x).view(B, C, H * W).permute(0, 2, 1) attn torch.matmul(q, k) * (C ** -0.5) attn torch.softmax(attn, dim-1) h torch.matmul(attn, v) h h.permute(0, 2, 1).view(B, C, H, W) return self.proj(h) x下采样我用的是stride2的卷积上采样用最近邻插值加3×3卷积。这样的做法在DDPM中很常见既不会像反卷积那样产生棋盘伪影插值方式也足够平滑。class DownBlock(nn.Module): def __init__(self, in_ch, out_ch, time_dim, num_res2, attnFalse): super().__init__() self.res nn.ModuleList([ ResBlock(in_ch if i 0 else out_ch, out_ch, time_dim) for i in range(num_res) ]) self.attn AttnBlock(out_ch) if attn else nn.Identity() self.down nn.Conv2d(out_ch, out_ch, 3, stride2, padding1) def forward(self, x, t): for r in self.res: x r(x, t) x self.attn(x) return self.down(x) class UpBlock(nn.Module): def __init__(self, in_ch, skip_ch, out_ch, time_dim, num_res2, attnFalse): super().__init__() self.proj_skip nn.Conv2d(in_ch skip_ch, out_ch, 3, padding1) self.res nn.ModuleList([ ResBlock(out_ch, out_ch, time_dim) for _ in range(num_res) ]) self.attn AttnBlock(out_ch) if attn else nn.Identity() self.upsample nn.Sequential( nn.Upsample(scale_factor2, modenearest), nn.Conv2d(out_ch, out_ch, 3, padding1), ) def forward(self, x, skip, t): x torch.cat([x, skip], dim1) x self.proj_skip(x) for r in self.res: x r(x, t) x self.attn(x) return self.upsample(x)UpBlock里我先把编码器跳过来的特征和当前特征拼接再用一个卷积把通道对齐。这个设计比直接相加更稳因为不同层级的语义差异比较大强制相加可能损失信息。下面把完整UNet拼起来。ch_mult表示各层通道数相对于基础通道的倍数这里用(1, 2, 2)实际通道数为64、128、128。class UNet(nn.Module): def __init__(self, in_channels3, out_channels3, ch64, ch_mult(1, 2, 2), num_res_blocks2, time_dim256, attn_levels(8, 16)): super().__init__() self.time_dim time_dim self.time_mlp nn.Sequential( nn.Linear(time_dim, time_dim), nn.SiLU(), nn.Linear(time_dim, time_dim), nn.SiLU(), ) self.conv_in nn.Conv2d(in_channels, ch, 3, padding1) self.downs nn.ModuleList() self.ups nn.ModuleList() now_ch ch level_chs [] for i, mult in enumerate(ch_mult): out_ch ch * mult level_chs.append(out_ch) self.downs.append(DownBlock( now_ch, out_ch, time_dim, num_res_blocks, attn(32 // (2 ** i) in attn_levels), )) now_ch out_ch self.mid_res nn.ModuleList([ ResBlock(now_ch, now_ch, time_dim) for _ in range(num_res_blocks) ]) for i, skip_ch in enumerate(reversed(level_chs)): out_ch ch * ch_mult[len(ch_mult) - 1 - i] self.ups.append(UpBlock( now_ch, skip_ch, out_ch, time_dim, num_res_blocks, attn(32 // (2 ** (len(ch_mult) - 1 - i)) in attn_levels), )) now_ch out_ch self.norm_out nn.GroupNorm(8, ch) self.conv_out nn.Conv2d(ch, out_channels, 3, padding1) def forward(self, x, t): t sinusoidal_embedding(t, self.time_dim) t self.time_mlp(t) h self.conv_in(x) skips [] for down in self.downs: skips.append(h) h down(h, t) for res in self.mid_res: h res(h, t) for i, up in enumerate(self.ups): skip skips.pop() h up(h, skip, t) return self.conv_out(F.silu(self.norm_out(h)))这里有一个地方要解释一下attn_levels写的是(8, 16)表示在8×8和16×16分辨率上加注意力。32×32分辨率对应的i0不参与注意力16×16对应i18×8对应i2。如果你把网络改成更深层比如在64×64输入上训练记得同步调整attn_levels的数值否则注意力加错位置不会报错但效果和效率都会打折。4. 训练循环、EMA与采样生成4.1 训练前的调度参数准备在开训之前需要先把beta、sqrt_alpha_bar等参数准备好。这些参数会在加噪和采样中反复用到。手动计算最稳妥的方式如下T 1000 betas torch.linspace(1e-4, 0.02, T, dtypetorch.float32) alphas 1.0 - betas alpha_bar torch.cumprod(alphas, dim0) sqrt_alpha_bar torch.sqrt(alpha_bar) sqrt_one_minus_alpha_bar torch.sqrt(1.0 - alpha_bar)这段代码里的alpha_bar就是前面讲的累积连乘。需要注意的是所有预计算都要用float32不要用float64也不要中途切成半精度否则在时间步较大时可能会因为误差累积导致采样结果异常。另外一个细节torch.linspace在GPU上也可以直接用但我习惯先在CPU上算好使用时再通过.to(device)搬到目标设备这样各个设备之间结果完全一致。4.2 正向加噪函数的实现给定一张x0、时间步t和随机噪声epsilon可以直接构造加噪后的x_t。实战里我建议把t生成成每个样本独立的随机整数而不是整批次共用同一个t这样每个step模型能看到不同噪声程度的样本。def q_sample(x_start, t, noise): sqrt_alpha_bar_t sqrt_alpha_bar.to(t.device)[t, None, None, None] sqrt_one_minus_alpha_bar_t sqrt_one_minus_alpha_bar.to(t.device)[t, None, None, None] return sqrt_alpha_bar_t * x_start sqrt_one_minus_alpha_bar_t * noise这里索引出来是[B, 1, 1, 1]的形状会自动broadcast。不要直接索引成标量再相乘那样会丢失batch维度后续计算损失时会出现形状不匹配。这个函数会在训练循环的每一步被调用所以实现一定要轻量。4.3 完整的训练循环训练循环的核心可以压缩成五步取batch - 随机采样时间步 - 随机采样噪声 - 构造带噪图像 - 让网络预测噪声并计算MSE。下面是一个可直接训练的脚本核心代码包含EMA、学习率调度和梯度裁剪。import torch.nn.functional as F from torch.optim import AdamW from torch.optim.lr_scheduler import LambdaLR def cosine_lr(step, warmup1000, total100000): if step warmup: return step / warmup progress (step - warmup) / max(1, total - warmup) return 0.5 * (1.0 math.cos(math.pi * progress)) model UNet().to(config[device]) optimizer AdamW(model.parameters(), lrconfig[lr], weight_decay0.0) scheduler LambdaLR(optimizer, lr_lambdacosine_lr) ema_model UNet().to(config[device]) ema_model.load_state_dict(model.state_dict()) ema_decay 0.999 step_count 0 for epoch in range(config[epochs]): for images, _ in train_loader: images images.to(config[device]) t torch.randint(0, T, (images.shape[0],), deviceconfig[device]).long() noise torch.randn_like(images) x_t q_sample(images, t, noise) pred_noise model(x_t, t) loss F.mse_loss(pred_noise, noise) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() scheduler.step() with torch.no_grad(): for ema_p, p in zip(ema_model.parameters(), model.parameters()): ema_p.data.mul_(ema_decay).add_(p.data, alpha1 - ema_decay) if step_count % 500 0: print(fEpoch {epoch} step {step_count}: loss {loss.item():.4f}) step_count 1这段代码我特意加了梯度裁剪max_norm设为1.0。扩散模型在早期训练时偶尔会出现梯度过大导致的训练崩溃裁剪能大幅提升稳定性。EMA里指数移动平均的作用是让模型权重更平滑采样时用EMA权重通常比用原权重效果好很多这一点你在训练后期会明显感受到。这里再解释一个学习率设计为什么用cosine而不用固定学习率因为固定学习率在后期容易在最优解附近震荡cosine则能平滑地降低更新幅度让模型在训练后期更精细地收敛。和GAN相比DDPM不太需要把学习率调得很低但cosine调度确实能让FID再降一点。4.4 采样生成从纯噪声还原图像采样过程需要从xT开始按照训练时的逆向过程逐步去噪。DDPM的采样公式是x_{t-1} 1/sqrt(alpha_t) * (x_t - (1 - alpha_t) / sqrt(1 - alpha_bar_t) * eps_theta(x_t, t)) sigma_t * z当t 0时z是新的随机噪声当t 0时不加噪声。这里的sigma_t在原始论文里直接取sqrt(beta_t)。实现如下torch.no_grad() def ddpm_sample(model, n_samples, image_size, channels, device): model.eval() x torch.randn(n_samples, channels, image_size, image_size, devicedevice) for t in reversed(range(T)): t_tensor torch.full((x.shape[0],), t, devicedevice, dtypetorch.long) pred model(x, t_tensor) alpha_t alphas.to(device)[t] alpha_bar_t alpha_bar.to(device)[t] sqrt_alpha_t torch.sqrt(alpha_t) sqrt_one_minus_alpha_bar_t torch.sqrt(1.0 - alpha_bar_t) x 1 / sqrt_alpha_t * (x - (1 - alpha_t) / sqrt_one_minus_alpha_bar_t * pred) if t 0: sigma_t torch.sqrt(betas.to(device)[t]) x x sigma_t * torch.randn_like(x) return x这段代码能跑但完整跑1000步会有点慢尤其是CPU环境。如果你只是想快速看效果建议先实现DDIM采样把步数压到50步。这里也给一个最简DDIM采样的核心片段把上面采样循环里的“ sigma_t * z”去掉改用更小的步长重新计算系数就能大幅加速。torch.no_grad() def ddim_sample(model, n_samples, image_size, channels, device, sample_steps50): model.eval() times torch.linspace(0, T - 1, sample_steps, dtypetorch.long).flip(0) x torch.randn(n_samples, channels, image_size, image_size, devicedevice) for i in range(len(times)): t times[i].item() t_tensor torch.full((x.shape[0],), t, devicedevice, dtypetorch.long) pred model(x, t_tensor) alpha_bar_t alpha_bar.to(device)[t] if i len(times) - 1: t_next times[i 1].item() alpha_bar_next alpha_bar.to(device)[t_next] else: alpha_bar_next torch.ones_like(alpha_bar_t) x torch.sqrt(alpha_bar_next) * (x - torch.sqrt(1 - alpha_bar_t) * pred) / torch.sqrt(alpha_bar_t) x x torch.sqrt(1 - alpha_bar_next) * pred return xDDIM的优点是采样步数大幅减少代价是它不再严格对应DDPM的随机生成链而是把采样过程变成了确定性的ODE离散化。对追求生成速度的线上场景来说DDIM几乎是标配。不过第一次上手时我建议还是先跑纯DDPM因为DDIM对系数推导更敏感写错了不容易发现。4.5 如何保存图像和检查进度训练过程中我习惯每500个step记录一次loss每个epoch结束保存一组固定噪声种子下的采样图。这样你能直观看到训练进展。注意保存图片前要把[-1, 1]的数值反归一化到[0, 1]否则整个画面看起来会灰蒙蒙的。def save_images(images, path, nrow8, padding2): images (images.clamp(-1, 1) 1) / 2 grid torchvision.utils.make_grid(images, nrownrow, paddingpadding) torchvision.utils.save_image(grid, path)如果训练到后期发现生成的图虽然清晰但颜色和类别有些混淆可以适度增加总步数或者把beta_end调小到0.015附近。这个变量对最终色彩的饱和度影响很大建议你有时间多试几组。另一条路径是把模型通道从64提到128但CIFAR-10上这属于锦上添花先把基础版本跑通更重要。5. 常见问题与调试实录5.1 Loss不下降或下降极慢如果你刚跑起来发现Loss几乎不动第一件事检查数据管线确认输入图像不是全黑或全白确认归一化均值方差用的是(0.5, 0.5, 0.5)。第二个常见原因是学习率设置不合理扩散模型的默认学习率2e-4是经过大量实验验证的建议直接固定不要轻易改成1e-3。还有一个比较隐蔽的问题时间步t的dtype必须是long如果torch.randint输出后忘了加.long()索引计算会报错。即使不报错浮点索引在cuDNN下也可能产生不一致的数值。我调试时遇到过Loss在某个step后突然为NaN最后发现就是索引类型问题导致alpha_bar取到了负值。5.2 生成图片全是纯噪声这种情况通常不在训练而在采样阶段。常见原因有三个采样时t没有从T-1循环到0而是从中间某个位置开始采样公式里的系数写反了尤其要注意1/sqrt(alpha_t)和(1 - alpha_t) / sqrt(1 - alpha_bar_t)这两个系数不能混前一步生成的x没有作为下一步的输入而是重新从高斯噪声采样。有一个非常容易犯的错在采样循环里写成了x sqrt_alpha_bar * x sqrt_one_minus_alpha_bar * pred这等于又在加噪声所以当然无法还原图像。每次看到采样全噪声的提问八成是这里出了问题。你可以先打印一下x的均值和方差如果每一步均值都在0附近跳标准差始终接近1说明采样公式基本没对。5.3 显存不足或训练速度过慢CIFAR-10分辨率低在30系显卡上batch_size128跑64通道的U-Net显存大约占用5到7GB一般显卡都能应付。如果你显存不足优先把batch_size降到64还不行的再把U-Net的ch从64降到32效果略差但训练速度快很多。如果CPU训练建议把模型通道数降到32batch_size降到32并且把num_workers设为0否则Windows下DataLoader容易报错。CPU上完整训练几千步大概需要几小时体验肯定不如GPU但是用来验证代码逻辑是够了。还有一个常被忽略的点torch.backends.cudnn.benchmark True 可以给固定分辨率训练带来明显的加速尤其在快速迭代调试时建议加上。5.4 EMA权重与原始权重搞混训练时我同时维护了model和ema_model但采样时如果不小心加载了原始model权重生成的图会明显更毛糙。很多人跑完发现效果一般其实是忘用EMA。这里建议保存checkpoint时把两个权重都存进去名字区分清楚。torch.save({ model: model.state_dict(), ema_model: ema_model.state_dict(), optimizer: optimizer.state_dict(), step: step_count, }, fddpm_step_{step_count}.pt)加载时默认把ema_model作为最终生成模型这对稳定生成很有帮助。我在项目里实际对比过EMA权重生成的FID通常能降10%左右而且几乎不额外增加训练时间。EMA衰减系数一般取0.999如果你的训练总步数比较少可以适当降到0.995让EMA权重更新速度更快一些。5.5 不同随机种子影响很大DDPM对随机种子没有GAN那么敏感但要复现效果还是要把所有随机种子固定住。torch.manual_seed、np.random.seed、random.seed都要设。如果你发现同样的代码跑两次生成图风格有明显差异先检查是不是DataLoader shuffle的随机种子没固定。另外我习惯把采样固定噪声向量也搓出来存好这样每轮epoch对比时图像变化完全来自模型权重而不是随机噪声不同能更直观看到模型在逐步变好。做法很简单生成一批固定的z保存成pt文件每次采样都加载它。5.6 其他值得关注的经验还有一个容易被忽略的技巧在训练过程中动态修改时间步分布。默认情况t均匀采样自[0, T-1]这意味着模型对每个噪声等级的学习强度是一样的。实际上在采样时中低噪声等级对最终细节影响最大你可以把t的采样分布往中间偏一偏比如使用截断正态分布让模型更专注毛边的步骤。这个改造不影响训练流程只改t的采样方式很多二次实现都会这么做。还有一个是数据增强。CIFAR-10原始分辨率不高训练时随机翻转就够了不需要过度裁剪。如果你换成自己的数据集建议先统一resize到64×64或128×128再考虑是否加随机裁切。扩散模型对数据分布非常敏感训练集和测试集分布差太远时生成结果会明显发虚。最后再说一个我自己的习惯。无论跑GAN还是DDPM我都喜欢在项目目录里建一个debug_samples文件夹每个epoch结束就把当前生成图像丢进去。这样哪怕当天没有守在电脑前第二天打开文件夹也能像看连续剧一样看到模型是如何一步步从噪声里长出轮廓的。DDPM的好处就在这里训练过程没有太多玄学每一步梯度都在推动模型往前走。你不需要再和判别器斗智斗勇只需要按部就班地把去噪这个回归任务做深做透。CIFAR-10只是起点只要你把模型里的输入输出通道和图片尺寸改一改同样的代码完全可以迁移到灰度医学图像、真实照片或者其他生成任务上。如果你也在从GAN迁移过来的路上这篇文章和代码应该能帮你少走几段弯路。
返回列表