ARTICLE DETAIL

资讯详情

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

纯PyTorch实现DDPM:MNIST手写数字生成全流程拆解

纯PyTorch实现DDPM:MNIST手写数字生成全流程拆解 简介面向深度学习入门者与扩散模型研究者提供一套可直接运行的DDPM实现资源包涵盖从理论到代码的完整链路完整复现了PyTorch环境下的训练与采样流程。整个包体共27个文件约22.21MB包含8个Python源码文件、预训练权重、MNIST数据集、实验结果图以及说明文档数据、代码和文档分开放置模块划分清晰解压后即可对照运行。已有218人学习下载。内容从数据集获取与处理入手覆盖DDPM类设计、U-Net去噪网络构建、训练算法与采样细节并在MNIST上完成了训练与采样生成图像在视觉上与真实样本接近验证了模型有效性同时探讨了不同网络架构对生成质量的影响并总结了复现中的关键点和注意事项。其中训练与采样脚本均可直接运行配合生成的图像示例能快速验证模型效果针对常见环境与参数问题也有说明无论用于课程作业、课题复现还是进一步研究都具备实用价值适合需要快速上手扩散模型或在此基础上改进网络结构的开发者直接参考。 前阵子复现DDPMDenoising Diffusion Probabilistic Models时我翻了很多资料发现一个尴尬的情况讲原理的文章一大堆但真正能直接跑通的PyTorch实现却不多要么是论文公式的堆砌要么依赖一堆重型扩展库反而把核心逻辑盖住了。折腾了几天我把一份精简但完整的可运行源码整理了出来——纯PyTorch实现不依赖任何第三方扩散模型库在MNIST上几十个epoch就能看到清晰的手写数字生成效果。这篇文章就是这份代码的完整拆解先讲清楚DDPM到底在干什么再把环境搭建、代码实现、训练采样一条龙跑通最后把我会踩的坑全列出来。这套东西适合两类人一是有PyTorch基础、想从原理层面搞懂扩散模型的学生或工程师二是想快速在自己的小数据集上验证DDPM效果、需要一份干净代码做二次改动的开发者。整个项目用的技术栈非常基础PyTorch torchvision numpy硬件上有一块普通显卡就行实在没GPU用CPU也能跑通全流程。1. DDPM干了什么事从噪声里“雕”出图像1.1 前向过程到底在做什么先看扩散模型的“破坏”环节。DDPM定义了一个前向过程给一张干净图像 x0 逐步添加高斯噪声一共加 T 步通常取1000步每一步的噪声强度由一个逐渐增大的方差 β_t 控制。公式写出来是这样q(x_t | x_{t-1}) N(x_t; √(1-β_t)·x_{t-1}, β_t·I)这个式子看着唬人其实理解起来很简单。你可以想象一杯清水每步滴入一滴墨汁前几步水还大致透明到后面整杯水彻底变黑。β_t 就是“每步滴多少墨”的节奏。实际操作中我们不需要一步一步去加噪因为高斯分布的叠加性让我们可以“一步到位”。用重参数化技巧从 x0 直接算任意时刻 t 的噪声图x_t √(ᾱ_t)·x0 √(1-ᾱ_t)·ε其中 ε 是标准高斯噪声ᾱ_t 是前 t 步 α 的累积乘积α_t 1 - β_t。这个式子的好处在于训练时不用真的循环1000次加噪随机抽一个时间步 t用一条公式就能算出来。这也是整个DDPM实现里最关键的“加速器”。1.2 反向过程让网络学会“倒放”如果能把前向过程倒过来把纯噪声一步步还原成图像就实现了生成。但前向过程的每一步都是“破坏”信息在丢失严格意义上的逆过程无法直接算。DDPM的解决办法是训练一个神经网络来近似这个逆过程。用公式说就是我们希望学一个 p_θ(x_{t-1}|x_t)它也是一个高斯分布均值和方差由网络预测。实际训练中网络要做的事更简单输入当前的噪声图 x_t 和步数 t输出它预测的噪声 ε_θ。训练目标就是让预测噪声和真实加入的噪声尽可能一致。这个过程可以类比成有人当着你的面把照片复印了1000次每次复印都变得更模糊你看惯了各种模糊程度下的照片后看到一张中等模糊的图就能大致猜出刚才复印时加进去的是什么干扰。时间步 t 也要作为网络输入这一点常被新手忽略。因为模型拿到的图 x_t噪声程度不同同一张模糊图如果不知道它被“复印”了多少次就不可能准确预测噪声。所以网络里必须有一个把 t 编码成向量的模块通常用正弦位置编码和Transformer里的位置编码同源。1.3 损失函数为什么只有MSEDDPM训练时用的损失函数简单到让人怀疑L E_{t, x0, ε} [ || ε - ε_θ(x_t, t) ||² ]也就是均方误差。这个形式其实是变分下界经过一系列化简之后的结果。完整的推导过程涉及KL散度、马尔可夫链、熵很长但对于工程实现来说我们只需要抓住核心结论训练目标就是让网络准确预测出前向过程加入的噪声。训练流程可以总结成三步随机抽一个时间步 t用 q_sample 给图像加噪让网络预测噪声并计算MSE。明白了这三个环节实现DDPM的代码框架已经清晰了一大半。2. 环境准备先把PyTorch跑起来2.1 推荐版本组合与安装命令我这次用的组合是 Python 3.10 PyTorch 2.1.2 CUDA 12.1这也是目前兼容性比较稳的一套。PyTorch 2.x 和 1.13 在核心API上对DDPM这类代码没什么区别关键是torch和CUDA版本别打架。推荐用Anaconda先建独立环境避免污染其他项目的依赖conda create -n ddpm python3.10 -y conda activate ddpm pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121 pip install numpy matplotlib tqdm如果你用的是国内网络pip默认源拉大文件容易中断可以临时换成清华源或者直接给pip加-i参数。装完之后验证一下import torch print(torch.__version__) print(torch.cuda.is_available())输出True就说明GPU环境正常。注意这里有个容易踩的坑有些人装的是CPU版PyTorchtorch.cuda.is_available()返回False还以为是显卡坏了。先确认安装命令里的cu121或cu118字样别拿CPU版凑合。2.2 没有GPU能不能跑完全能跑。MNIST的图片只有28x28模型也很小CPU训练一个epoch也就一两分钟。我在一台没有独显的笔记本上测试过50个epoch大概要一个多小时但结果一样能出图。如果你只有CPU安装时用默认pip源装CPU版本就可以pip install torch torchvision需要注意一点CPU训练时dataloader的num_workers设成0避免Windows上多进程加载数据报错。Mac用户还可以用MPS加速代码里加一行判断device torch.device(cuda if torch.cuda.is_available() else (mps if torch.backends.mps.is_available() else cpu))2.3 数据集的准备训练数据直接用torchvision自带的MNIST一份FashionMNIST也一样代码不用改。下载后会存到本地./data目录过程很简单transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)), ]) dataset torchvision.datasets.MNIST(root./data, trainTrue, downloadTrue, transformtransform) dataloader torch.utils.data.DataLoader(dataset, batch_size128, shuffleTrue, num_workers4)这里Normalize((0.5,), (0.5,))的效果是把像素值从 [0,1] 映射到 [-1,1]。为什么要这么做因为DDPM前向加噪使用的是标准高斯噪声均值是0标准的图像范围 [-1,1] 和噪声分布更匹配也方便最后采样时对输出做裁剪。如果忘了归一化训练过程会非常不稳定生成结果也容易发灰。3. 代码实现一个能在MNIST上跑通的完整流程3.1 噪声调度定义扩散的节奏先说噪声调度。最经典的是线性调度β 从 0.0001 线性增长到 0.02共1000步。为什么这样选因为加噪前几步幅度要小保留图像结构好让网络学习后面可以大一些快速把图像破坏成噪声。实现如下import torch import torch.nn.functional as F def linear_beta_schedule(timesteps1000, beta_start1e-4, beta_end0.02): return torch.linspace(beta_start, beta_end, timesteps) timesteps 1000 betas linear_beta_schedule() alphas 1.0 - betas alphas_cumprod torch.cumprod(alphas, dim0) alphas_cumprod_prev F.pad(alphas_cumprod[:-1], (1, 0), value1.0) sqrt_alphas_cumprod torch.sqrt(alphas_cumprod) sqrt_one_minus_alphas_cumprod torch.sqrt(1.0 - alphas_cumprod)这几个预计算量在训练前算好放在内存里全程用不到也没关系后面加噪和采样都要引用。alphas_cumprod_prev用来算后验分布的方差采样时会用到。3.2 前向加噪一条公式算任意步基于前面推导的重参数化公式加噪函数写出来就四行def q_sample(x_start, t, noise): sqrt_alpha_bar sqrt_alphas_cumprod[t].view(-1, 1, 1, 1) sqrt_one_minus_bar sqrt_one_minus_alphas_cumprod[t].view(-1, 1, 1, 1) return sqrt_alpha_bar * x_start sqrt_one_minus_bar * noise注意这里t是一个batch的随机时间步形状是(batch,)索引出来的系数也要保持batch维度所以要用view(-1, 1, 1, 1)扩展成和图像张量形状匹配的维度。这一步如果写错了训练时会出现形状不匹配的报错这是新手最容易遇到的问题之一。3.3 轻量级UNet扩散模型的骨干网络DDPM官网实现用的是UNet ResNet Block Attention的组合。MNIST这种小图不需要那么重我写了一个简化版只保留UNet的encoder-decoder结构和时间步注入训练速度更快效果也不差时间编码模块class TimeEmbedding(nn.Module): def __init__(self, dim): super().__init__() self.dim dim def forward(self, t): half self.dim // 2 freqs torch.exp(-torch.log(torch.tensor(10000.0)) * torch.arange(half).to(t.device) / half) args t[:, None] * freqs[None, :] return torch.cat([torch.sin(args), torch.cos(args)], dim-1)基础卷积块负责“特征提取 时间步信息注入”class Block(nn.Module): def __init__(self, in_ch, out_ch, time_ch): super().__init__() self.conv1 nn.Conv2d(in_ch, out_ch, 3, padding1) self.conv2 nn.Conv2d(out_ch, out_ch, 3, padding1) self.mlp nn.Linear(time_ch, out_ch) self.norm1 nn.GroupNorm(8, out_ch) self.norm2 nn.GroupNorm(8, out_ch) def forward(self, x, t): h F.silu(self.norm1(self.conv1(x))) h h self.mlp(t)[:, :, None, None] return F.silu(self.norm2(self.conv2(h)))self.mlp(t)把时间embedding映射到和特征图通道数一致的向量然后通过广播加到特征图上。这就是时间步信息进入网络的方式不复杂但作用很大。简化的UNet主体class SimpleUNet(nn.Module): def __init__(self, in_ch1, base_ch64, time_dim256): super().__init__() self.time_mlp nn.Sequential( TimeEmbedding(time_dim), nn.Linear(time_dim, time_dim), nn.SiLU(), nn.Linear(time_dim, time_dim), ) self.inc Block(in_ch, base_ch, time_dim) self.down1 Block(base_ch, base_ch * 2, time_dim) self.down2 Block(base_ch * 2, base_ch * 4, time_dim) self.up1 Block(base_ch * 4, base_ch * 2, time_dim) self.up2 Block(base_ch * 2, base_ch, time_dim) self.outc nn.Conv2d(base_ch, in_ch, 3, padding1) self.pool nn.MaxPool2d(2) self.upsample nn.Upsample(scale_factor2, modenearest) def forward(self, x, t): t self.time_mlp(t) x1 self.inc(x, t) x2 self.down1(self.pool(x1), t) x3 self.down2(self.pool(x2), t) h self.up1(self.upsample(x3) x2, t) h self.up2(self.upsample(h) x1, t) return self.outc(h)这个结构里encoder两层下采样decoder两层上采样中间有跳跃连接。比起完整版UNet少了很多层但在MNIST这种低分辨率数据集上完全够用。你要是第一次实现DDPM我建议就用这种小网络训练快迭代调参的效率也高。3.4 训练循环核心只有五步训练过程的代码量其实很少核心步骤写在注释里model SimpleUNet().to(device) optimizer torch.optim.AdamW(model.parameters(), lr1e-4) for epoch in range(epochs): for x0, _ in dataloader: x0 x0.to(device) batch x0.size(0) t torch.randint(0, timesteps, (batch,), devicedevice) noise torch.randn_like(x0) x_t q_sample(x0, t, noise) pred_noise model(x_t, t) loss F.mse_loss(pred_noise, noise) optimizer.zero_grad() loss.backward() optimizer.step()就这么简单随机抽时间步加噪预测噪声算MSE。每轮迭代都在重复这个过程没有任何花哨操作。训练时建议用一个很小的batch size跑两步确认一下loss在下降再去跑完整训练避免一开始就发现 loss为nan 还得排查半天。3.5 采样生成从纯噪声开始一步步“雕”出图像采样就是前向过程的逆过程从纯噪声 x_T 开始循环 T 步每步根据网络预测的噪声算出去噪后的均值再采样一次torch.no_grad() def p_sample(model, x, t_index): t torch.full((x.size(0),), t_index, devicex.device, dtypetorch.long) pred_noise model(x, t) alpha_bar alphas_cumprod[t_index] alpha alphas[t_index] beta betas[t_index] alpha_bar_prev alphas_cumprod_prev[t_index] # 估计x0并裁剪到 [-1,1] x0 (x - torch.sqrt(1 - alpha_bar) * pred_noise) / torch.sqrt(alpha_bar) x0 torch.clamp(x0, -1.0, 1.0) if t_index 0: return x0 posterior_var beta * (1 - alpha_bar_prev) / (1 - alpha_bar) mean (x - (beta / torch.sqrt(1 - alpha_bar)) * pred_noise) / torch.sqrt(alpha) noise torch.randn_like(x) return mean torch.sqrt(posterior_var) * noise torch.no_grad() def p_sample_loop(model, shape): x torch.randn(shape, devicenext(model.parameters()).device) for i in reversed(range(timesteps)): x p_sample(model, x, i) return x采样的关键点是不要跳步一步一步从1000走回0。每一步都在做“去噪 增加随机性”只有在最后一步不加噪声。你可以试着把循环改成每50步采一次样把中间结果打印出来会非常清楚地看到图像从噪声到数字的演化过程那种体验比看任何论文图都直观。4. 训练实录与调参心得4.1 一组可以直接抄的训练参数我这里给一份实测能稳定出图的参数组合在单张RTX 3060上跑大约半小时参数值说明timesteps1000扩散步数论文标准值batch size128显存不够就降到64learning rate1e-4AdamW不需要太高epochs60MNIST上60轮足够看到清晰效果beta schedule线性 1e-4 → 0.02默认选择优化器AdamW比Adam更稳最好加weight decay显存占用在12GB之间很多老显卡都能跑。batch size再大的话MNIST这种小图也吃不了多少显存但如果你换到CIFAR-10分辨率从28涨到32通道数变成3显存占用会涨很多这时候需要把batch size降下来。4.2 三个让生成质量质变的小技巧第一个技巧是EMA指数移动平均。训练时同时维护一份模型权重的滑动平均采样时用这份平均权重而不是实时权重。这个技巧在扩散模型里几乎是标配原论文也用了。实现很简单训练完取权重和模型的影子权重做一个加权和多写十行代码生成效果的稳定性提升一个档次。第二个技巧是采样时对预测出的 x0 做clamp。因为网络预测有误差反向过程中估算出的 x0 偶尔会超出 [-1,1] 范围如果不裁剪误差会在迭代中被不断放大最后生成图出现明显的噪点和条纹。代码里我就是在p_sample函数中先算x0然后clamp再往下走。第三个技巧是换余弦调度。线性调度在中间步数噪声增长过快余弦调度更平滑被后续的很多扩散模型采用。实现只需改一行def cosine_beta_schedule(timesteps, s0.008): steps torch.arange(timesteps 1, dtypetorch.float32) / timesteps alpha_bar torch.cos((steps s) / (1 s) * torch.pi / 2) ** 2 betas 1 - alpha_bar[1:] / alpha_bar[:-1] return torch.clip(betas, 1e-8, 0.999)我用同样的训练配置对比过余弦调度生成的数字边缘更干净尤其是训练epoch数不足的时候优势更明显。4.3 训练过程中看什么指标生成模型不好用单纯的loss判断好坏因为loss降低到一定程度后肉眼效果仍可能有明显差异。我的习惯是每5个epoch采样一批图存下来直接看生成效果。损失如果一直在降生成的数字从一团噪点慢慢成形就是正常的。另一个观察点是验证集上的loss如果过拟合验证loss会回升这时候可以加数据增强或者提前停止。5. 常见问题与排查技巧实录现象可能原因排查与解决loss变成nan学习率过大或数据没归一化检查图像是否归一化到[-1,1]把lr降到1e-5试跑几步生成图全黑或全白采样循环方向写反或没clamp确认从T循环到0别正着循环加上x0裁剪生成图像有横向/纵向条纹时间步输入形状错误广播维度不匹配检查t是否view(-1)展平后再传入模型显存不足batch size太大或模型通道数过多降batch size把base_ch从64降到32训练loss下降但生成效果很差训练epoch数不够或模型太浅增加epoch加一层down/up或者加attention首次加载数据集非常慢downloads没有自动下载或网络问题手动下载MNIST放到./data/MNIST/raw/下换CIFAR-10后报维度错误输入通道数还是1把in_ch改成3确认归一化参数这里面最让我印象深刻的一个坑是模型里时间步t忘记展平。torch.randint生成的t形状是(batch,)但经过DataLoader或某些中间操作后可能变成(batch, 1)一旦形状不对UNet里的t[:, None]就会广播出错。这种bug不会让程序崩溃但会静默地把时间步信息弄乱训练出来的模型生成效果特别差。遇到生成效果莫名糟糕的情况优先检查t的shape。另外换数据集时有个细节FashionMNIST和MNIST分辨率一样直接替换下载接口就行如果换到自己的灰度图数据集需要resize到统一尺寸最好是偶数方便UNet下采样。CIFAR-10这种三通道图把in_ch改成3就能跑其余逻辑完全不动。最后再分享一个小技巧训练结束后不要只保存模型权重把betas、alphas、alphas_cumprod这些常量也一起存成.pt文件。否则下次采样时你很可能忘了当初用的哪种beta schedule或者timesteps改了没同步导致采样结果完全不对。我习惯这样保存torch.save({ model: model.state_dict(), ema_model: ema_model.state_dict(), betas: betas, alphas: alphas, alphas_cumprod: alphas_cumprod, alphas_cumprod_prev: alphas_cumprod_prev, }, ddpm_mnist.pt)这一个小习惯帮我省了无数次重新训练的时间。扩散模型这个东西第一次完整跑通之后你就会发现它并没有论文里写得那么高不可攀。核心代码加起来不到200行原理也很朴素学会从噪声中预测噪声然后一步一步走回去。希望这份可运行源码能让你少走几个弯路早点把自己的第一张扩散图像生成出来。本文还有配套的精品资源点击获取
返回列表