w86763777-pytorch-ddpm-be2787b:基于 PyTorch 的去噪扩散概率模型实现项目

ICLR ICLR 2023 poster Blurring Diffusion Models

Branch1Tags0
This repository is empty

去噪扩散概率模型

去噪扩散概率模型 [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

    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 实现

Introduction

ICLR ICLR 2023 poster Blurring Diffusion Models

Customize your domain