ICLR ICLR 2023 poster Blurring Diffusion Models
This repository is empty
Translated by AI, submit an issue feedback
去噪扩散概率模型
去噪扩散概率模型 [1] 的非官方 PyTorch 实现。
本实现基本遵循官方 TensorFlow 实现 [2] 中的大部分细节。我采用 PyTorch 的编码风格将 [2] 移植到 PyTorch,希望熟悉 PyTorch 的人员能够轻松理解每一个实现细节。
待办事项
- 数据集
- 功能
- 复现实验
环境要求
-
Python 3.6
-
软件包 升级 pip 以安装最新版 tensorboard
pip install -U pip setuptools pip install -r requirements.txt -
下载数据集的预计算统计数据:
为
cifar10.train.npz创建stats文件夹。stats └── cifar10.train.npz
从头开始训练
- 以 CIFAR10 为例:
python main.py --train \ --flagfile ./config/CIFAR10.txt - [可选] 覆盖参数
python main.py --train \ --flagfile ./config/CIFAR10.txt \ --batch_size 64 \ --logdir ./path/to/logdir - [可选] 选择 GPU ID
CUDA_VISIBLE_DEVICES=1 python main.py --train \ --flagfile ./config/CIFAR10.txt - [可选] 多 GPU 训练
CUDA_VISIBLE_DEVICES=0,1,2,3 python main.py --train \ --flagfile ./config/CIFAR10.txt \ --parallel
评估
flagfile.txt会自动保存到您的日志目录中。config/CIFAR10.txt的默认日志目录为./logs/DDPM_CIFAR10_EPS- 开始评估
python main.py \ --flagfile ./logs/DDPM_CIFAR10_EPS/flagfile.txt \ --notrain \ --eval - [可选] 多 GPU 评估
CUDA_VISIBLE_DEVICES=0,1,2,3 python main.py \ --flagfile ./logs/DDPM_CIFAR10_EPS/flagfile.txt \ --notrain \ --eval \ --parallel
复现实验
CIFAR10
- FID:3.249, inception 分数:9.475(0.174)

检查点可从我的 drive 下载。
参考文献
[1] 去噪扩散概率模型
[2] 官方 TensorFlow 实现