运行目标
一个逻辑完整的最小 DDPM:可训练、可保存 checkpoint、可从噪声采样。页面不展示伪造的 loss 或生成结果。
项目拆成 8 个可验证模块
01
schedules.py
与 预计算
02
forward.py
extract 与 q_sample
03
embeddings.py
正弦时间嵌入
04
blocks.py
残差卷积块
05
unet.py
简化噪声预测器
06
train.py
随机 与 MSE
07
sampling.py
单步与完整循环
08
sample.py
固定种子生成
先看 U-Net 里 shape 怎样流动
U-Net 不是扩散概率理论本身,而是常用的噪声预测网络。它必须接收时间步,因为同一个像素值在 与 时代表完全不同的噪声水平。
点击模块查看 shape
一个面向 MNIST 的教学版 U-Net
这张图说明数据流与张量尺寸,不在浏览器中运行大型网络。
tSinusoidal embedding → MLP
skip: down1 → up1 skip: down2 → up2
瓶颈
Res + Attention
[B, 256, 7, 7]低分辨率上聚合全局关系;注意力是常见工程模块,不是扩散概率理论本身。
本模块接收时间嵌入按模块读代码,不要一次吞下整份文件
从 q_sample 开始,每一块都先确认输入、输出与 shape,再切到工程版看归一化、AMP、EMA 与 checkpoint。八个模块按顺序可以组成完整训练与采样闭环。
01 · q_sample
调度、extract 与闭式加噪
输入x0 [B,C,H,W], t [B]输出xt 与 noise,均同 x0 shape
01def make_schedule(T: int):
02 betas = torch.linspace(1e-4, 2e-2, T)
03 alphas = 1.0 - betas
04 alpha_bars = torch.cumprod(alphas, dim=0)
05 return betas, alphas, alpha_bars
06
07def extract(values, t, x_shape):
08 # [T] + [B] -> [B,1,1,1]
09 out = values.gather(0, t)
10 return out.reshape(t.shape[0], *((1,) * (len(x_shape)-1)))
11
12def q_sample(x0, t, alpha_bars, noise=None):
13 noise = torch.randn_like(x0) if noise is None else noise
14 a_bar = extract(alpha_bars, t, x0.shape)
15 xt = a_bar.sqrt() * x0 + (1-a_bar).sqrt() * noise
16 return xt, noise
最容易错
把 β(方差)直接当成噪声标准差;或忘记把系数 reshape 为 [B,1,1,1]。
目录骨架
mini-ddpm/
├── models/{blocks.py, embeddings.py, unet.py}
├── diffusion/{schedules.py, forward.py, sampling.py}
├── train.py
├── sample.py
├── utils.py
└── requirements.txt想直接跑通:页面中的模块已经合并为一份可下载脚本,包含 MNIST 数据、Tiny U-Net、训练、checkpoint 与完整采样。默认 10 epochs;
--limit-batches 20 --steps 200 只用于冒烟检查,不能当作有效训练结果。设备提示: MNIST 的小型 U-Net 可在 CPU 上跑通流程,但完整训练会慢;GPU 路径应把模型、batch 和所有预计算 buffer 放在同一设备。EMA 与混合精度属于可选工程拓展,不是 DDPM 理论必需部分。
CPU 路径
先验证正确性
使用小 batch、少量时间步和少数数据,确认 shape、loss 能下降、采样循环不报错。不要把短暂 smoke test 当成训练结果。
GPU 路径
再扩大训练
统一 device,按显存调整 batch;可选 AMP 与 EMA。固定 seed 只能提升复现性,无法保证跨硬件逐位一致。