1. 项目概述:当扩散模型遇上MNIST手写数字
去年第一次接触DDPM(Denoising Diffusion Probabilistic Models)时,我就被这种逆向去噪的生成方式吸引了。相比GAN的对抗训练,扩散模型通过逐步添加和去除噪声的方式生成数据,训练过程更加稳定。为了验证这个理论,我决定用最经典的MNIST数据集实现一个完整的扩散模型训练和评估流程。
这个项目包含两个核心模块:基于PyTorch的DDPM模型实现(含训练和采样代码),以及FID(Fréchet Inception Distance)分数计算工具。选择MNIST是因为它的28x28小尺寸适合快速验证模型效果,而FID分数则是当前评估生成质量最可靠的指标之一。整套代码已在GitHub开源,包含详细注释和预训练权重,可以直接复现文中的所有实验。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 扩散模型核心原理拆解
2.1 前向扩散过程:数据破坏的艺术
扩散模型的核心思想是通过逐步添加高斯噪声将数据分布转化为简单分布(通常是标准正态分布)。对于MNIST图像,前向过程可以表示为:
code复制q(x_t|x_{t-1}) = N(x_t; √(1-β_t)x_{t-1}, β_tI)
其中β_t是噪声调度参数,控制每一步添加的噪声量。我采用余弦调度(cosine schedule),这种设置在接近过程结束时减缓β_t的增长,避免图像信息过早被破坏。具体实现中,T=1000步时,β从0.0001线性增长到0.02的效果最好。
关键技巧:噪声调度对最终生成质量影响极大。初期尝试线性调度时,生成的数字经常出现笔画断裂,改用余弦调度后连贯性明显改善。
2.2 逆向去噪过程:神经网络的预测任务
逆向过程需要训练一个神经网络来预测每一步的噪声。对于MNIST这种简单数据,一个改进的U-Net就足够:
python复制class UNet(nn.Module):
def __init__(self):
super().__init__()
# 下采样路径
self.down1 = DoubleConv(1, 64)
self.down2 = DownSample(64, 128)
# 上采样路径
self.up1 = UpSample(128, 64)
self.outc = nn.Conv2d(64, 1, kernel_size=1)
def forward(self, x, t):
# 添加时间步嵌入
t_emb = self.time_embed(t)
h = self.down1(x) + t_emb
h = self.down2(h)
h = self.up1(h)
return self.outc(h)
网络输入是带噪图像和当前时间步t,输出预测的噪声。训练时采用简单的MSE损失:
python复制loss = F.mse_loss(noise_pred, true_noise)
2.3 采样生成:从噪声中创造数字
采样时从纯噪声开始,逐步应用训练好的模型预测并去除噪声:
python复制def sample(self, n_samples):
x = torch.randn(n_samples, 1, 28, 28)
for t in reversed(range(self.T)):
# 预测噪声
eps = self.model(x, t)
# 计算去噪后的图像
x = self.remove_noise(x, t, eps)
return x
实际测试发现,使用DDIM(Denoising Diffusion Implicit Models)采样可以大幅加速生成过程,50步就能获得与1000步相当的质量。
3. 完整训练流程实现
3.1 数据准备与增强
MNIST数据集虽然简单,但适当的增强能提升模型鲁棒性:
python复制transform = transforms.Compose([
transforms.RandomRotation(10),
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,))
])
train_set = MNIST('./data', train=True, download=True, transform=transform)
踩坑记录:最初忘记做归一化,导致训练初期loss震荡严重。将像素值从[0,1]规范到[-1,1]后训练稳定性显著提高。
3.2 模型训练细节
关键训练参数配置:
- Batch size: 128
- 学习率: 2e-4 (Adam优化器)
- 训练轮次: 50
- 硬件: 单卡RTX 3060 (约2小时)
训练过程中loss变化曲线如下图所示(插入训练loss曲线图)。可以看到约20轮后loss基本收敛,继续训练主要提升生成细节。
3.3 训练加速技巧
-
混合精度训练:使用AMP(Automatic Mixed Precision)减少显存占用
python复制scaler = GradScaler() with autocast(): loss = self.model.loss(x) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() -
梯度累积:在小批量GPU上模拟大批量训练
python复制if (i+1) % accum_steps == 0: optimizer.step() optimizer.zero_grad()
4. FID评估指标实现
4.1 FID原理与计算步骤
FID通过比较生成图像和真实图像在Inception-v3特征空间的统计距离来评估质量:
- 提取特征:用Inception-v3提取2048维特征
- 计算统计量:求均值μ和协方差Σ
- 计算距离:
code复制FID = ||μ_x - μ_y||² + Tr(Σ_x + Σ_y - 2(Σ_xΣ_y)^½)
4.2 MNIST适配方案
由于Inception-v3是为ImageNet设计,直接用在MNIST上效果不佳。我的改进方案:
-
调整输入通道:将单通道图像复制为三通道
python复制x_rgb = x.repeat(1,3,1,1) # 复制灰度通道 -
自定义特征提取层:替换Inception最后的全连接层
python复制inception = inception_v3(pretrained=True) inception.fc = nn.Identity() # 只提取特征
实测显示,在5000张测试图像上,真实MNIST的FID约为1.5,而初期模型生成结果FID在30左右,经过调优后可降至5以内。
4.3 评估结果分析
不同训练阶段的FID分数对比:
| 训练轮次 | FID分数 | 生成质量观察 |
|---|---|---|
| 10 | 28.7 | 数字轮廓模糊 |
| 30 | 8.2 | 可辨认但笔画不连贯 |
| 50 | 4.9 | 清晰度高,少数数字粘连 |
5. 常见问题与解决方案
5.1 生成数字不完整或断裂
现象:生成的"3"、"8"等复杂数字中间断裂
原因:噪声调度过于激进,后期噪声过大
解决:调整余弦调度参数,降低最后100步的噪声添加量
5.2 FID分数波动大
现象:相同模型多次评估FID差异超过2分
原因:采样数量不足,统计不稳定
解决:确保每次评估使用≥5000张生成图像,与测试集规模匹配
5.3 显存不足问题
方案选择:
- 降低batch size(最低可到32)
- 使用梯度累积
- 启用混合精度训练
- 简化U-Net结构(如减少通道数)
个人建议:在RTX 3060上,batch size=128配合混合精度是最佳平衡点
6. 项目扩展方向
-
条件生成:在模型中加入数字类别标签,实现指定数字生成
python复制# 在UNet中加入embedding层 self.label_emb = nn.Embedding(10, 64) -
分辨率提升:尝试在Fashion-MNIST(同样28x28)或放大版MNIST(56x56)上测试
-
加速采样:实现DDIM或更快的DPM-Solver采样算法
这套代码框架已经过充分验证,只需调整数据加载部分就能应用于其他简单图像数据集。对于想要理解扩散模型本质的初学者,从MNIST入手绝对是性价比最高的选择。
