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。按 EscQ 可提前停止。无桌面环境使用:

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\),并令:

\[ \alpha_t = 1 - \beta_t, \qquad \bar\alpha_t = \prod_{s=1}^{t}\alpha_s \]

任意时刻的前向加噪可直接写为:

\[ x_t = \sqrt{\bar\alpha_t}x_0 + \sqrt{1-\bar\alpha_t}\epsilon, \qquad \epsilon \sim \mathcal N(0, I) \]

训练网络 \(\epsilon_\theta(x_t,t,c)\) 预测真实噪声,其中 \(c\) 是类别条件。优化目标为:

\[ \mathcal L = \mathbb E_{x_0,\epsilon,t,c} \left[\|\epsilon - \epsilon_\theta(x_t,t,c)\|_2^2\right] \]

推理时,网络预测噪声后计算反向均值:

\[ \mu_\theta(x_t,t,c) = \frac{1}{\sqrt{\alpha_t}} \left(x_t- \frac{\beta_t}{\sqrt{1-\bar\alpha_t}} \epsilon_\theta(x_t,t,c)\right) \]

再从该分布采样 \(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”。因此模型同时学习:

\[ \epsilon_\theta(x_t,t,c), \qquad \epsilon_\theta(x_t,t,\varnothing) \]

推理时以 guidance scale \(s\) 合成两次预测:

\[ \epsilon_{\mathrm{cfg}} = \epsilon_{\mathrm{uncond}} + s(\epsilon_{\mathrm{cond}}-\epsilon_{\mathrm{uncond}}) \]

这通常能提高类别一致性。

训练策略与监控

默认配置:

项目 默认值
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_steptrain/loss_epoch
  • train/learning_rate
  • train/gradient_norm
  • samples/class_grid_raw
  • samples/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 环境执行上述命令。

results matching ""

    No results matching ""