CIFAR-10 Conditional DDPM
这是一个使用 PyTorch 实现的 CIFAR-10 类别条件 DDPM 示例。train.py 训练带 classifier-free guidance(CFG)的条件去噪 U-Net;infer.py 从指定类别和高斯噪声开始,逐步可视化并保存去噪过程。
安装与快速开始
根目录的 environment.yaml 已包含 PyTorch、TensorBoard、OpenCV 及其他运行依赖。更新环境后运行:
python tests/ddpm/train.py --device cuda
训练会自动下载 CIFAR-10 到 tests/ddpm/data/,默认输出:
- TensorBoard event:
tests/ddpm/runs/ - 最终 checkpoint:
tests/ddpm/checkpoints/ddpm.pt - 周期 checkpoint:
tests/ddpm/checkpoints/ddpm_epoch_XXX.pt
启动 TensorBoard:
tensorboard --logdir tests/ddpm/runs
CPU smoke test:
python tests/ddpm/train.py \
--device cpu --epochs 1 --batch-size 2 --base-channels 8 \
--timesteps 2 --max-train-samples 2 --num-workers 0 \
--samples-per-class 1 --sample-every 1 --checkpoint-every 1
默认模型面向 GPU 训练。若显存不足,优先使用
--batch-size 64,而不是缩小模型。
推理与逐步可视化
指定类别生成一张猫图:
python tests/ddpm/infer.py \
--checkpoint tests/ddpm/checkpoints/ddpm.pt \
--class-name cat --device cuda --guidance-scale 2.0
推理默认执行完整 1,000 步 DDPM 采样,OpenCV 窗口会实时刷新,且每一步保存为 PNG。按 Esc 或 Q 可提前停止。无桌面环境使用:
python tests/ddpm/infer.py \
--checkpoint tests/ddpm/checkpoints/ddpm.pt \
--class-id 3 --no-window
采样完成后会在同一输出目录生成 denoising.gif。默认每 10 个去噪步骤保留一帧、以 20 FPS 循环播放;最终去噪帧会额外停留 10 个帧间隔。完整逐步 GIF 可使用:
python tests/ddpm/infer.py \
--checkpoint tests/ddpm/checkpoints/ddpm.pt \
--class-name cat --gif-frame-stride 1 --no-window
使用 --no-gif 可仅保存 PNG 帧。
快速预览可使用确定性 DDIM:
python tests/ddpm/infer.py \
--checkpoint tests/ddpm/checkpoints/ddpm.pt \
--class-name ship --sample-steps 100 --num-images 4 --no-window
--guidance-scale 控制类别约束:0 是无条件生成,1 是普通条件生成,推荐从 2.0 开始;过大可能降低图像多样性并产生伪影。
算法原理
DDPM 有两个方向相反的过程:
训练:真实图片 x0 → 加噪得到 xt → 网络预测噪声
推理:高斯噪声 xT → 逐步去噪 → 生成图片 x0
定义线性噪声日程 \(\beta_t\),并令:
任意时刻的前向加噪可直接写为:
训练网络 \(\epsilon_\theta(x_t,t,c)\) 预测真实噪声,其中 \(c\) 是类别条件。优化目标为:
推理时,网络预测噪声后计算反向均值:
再从该分布采样 \(x_{t-1}\),重复直到 \(x_0\)。
模型结构
xt [B, 3, 32, 32]
↓
32×32×128: 2 个条件残差块
↓
16×16×256: 2 个条件残差块 + self-attention
↓
8×8×256: 2 个条件残差块 + self-attention
↓
4×4×256: middle residual blocks
↑
U-Net 上采样路径 + skip connections
↓
预测噪声 εθ [B, 3, 32, 32]
时间步使用 sinusoidal embedding;类别使用 embedding。二者相加后投影到每个残差块的通道维度,使网络在所有分辨率上都知道当前噪声强度与目标类别。GroupNorm 减少对 batch size 的依赖;低分辨率 self-attention 用于建模图像的全局关系。
Classifier-Free Guidance
训练时,默认以 10% 概率将类别标签替换为“无条件 token”。因此模型同时学习:
推理时以 guidance scale \(s\) 合成两次预测:
这通常能提高类别一致性。
训练策略与监控
默认配置:
| 项目 | 默认值 |
|---|---|
| Epochs | 800 |
| Batch size | 128 |
| 扩散步数 | 1,000 |
| 学习率 | 2e-4 → 1e-5 |
| Warmup | 5 epochs |
| Optimizer | AdamW,weight decay 1e-4 |
| EMA decay | 0.9999 |
| CFG 条件丢弃率 | 0.1 |
学习率先线性 warmup,再按 cosine decay 下降。训练使用 AMP、GradScaler、梯度裁剪(最大范数 1.0)和 TF32。EMA 模型不参与反向传播,但用于采样和推理,通常比即时训练权重更稳定。
TensorBoard 记录:
train/loss_step、train/loss_epochtrain/learning_ratetrain/gradient_normsamples/class_grid_rawsamples/class_grid_ema
每 25 epoch 对每类固定噪声生成 4 张图。样本会最近邻放大 4 倍,并标注 CIFAR-10 类别名。判断质量时优先比较固定条件下的 EMA 样本,不要只依据噪声预测 MSE。
动态架构图
conditional_ddpm_animation.py 使用 ManimGL 动态展示前向扩散、条件 U-Net、MSE loss、CFG 与反向去噪。
在 tests/ddpm/ 目录渲染 720p MP4:
manimgl conditional_ddpm_animation.py ConditionalDDPMArchitecture \
-w -m --video_dir assets --file_name conditional_ddpm_architecture
渲染 GIF:
manimgl conditional_ddpm_animation.py ConditionalDDPMArchitecture \
-w -i -m --video_dir assets --file_name conditional_ddpm_architecture
如果 ManimGL 在编码阶段报出 libopenh264.so.5 缺失,应先用项目根目录的环境定义修复当前 Conda 环境:
mamba env update -n sketch -f environment.yaml
修复前可临时让 ManimGL 使用系统 FFmpeg(仅在系统的 /usr/bin/ffmpeg 可用时):
PATH=/usr/bin:$PATH manimgl conditional_ddpm_animation.py ConditionalDDPMArchitecture \
-w -i -m --video_dir assets --file_name conditional_ddpm_architecture
输出文件位于 tests/ddpm/assets/。
ManimGL 需要可用的 OpenGL 上下文;在无桌面 Linux 环境中,请通过 Xvfb 或具备 EGL/OpenGL 的 GPU 环境执行上述命令。