实际评估VQ-Diffusion时,我先确认它解决的具体问题:正式实施VQ-Diffusion。一旦进入日常自动化环节,输入边界、依赖和失败处理如果不清楚就很难稳定复用会直接影响交付,这也是我最关心的风险。与其反复读介绍,不如用一项范围明确的真实任务完成最小试跑,再依据配置时间、输出质量、异常信息和维护痕迹做取舍。我会把它列入愿意先做小范围验证并复查原始文档的团队的候选清单,而不是仅凭项目介绍直接纳入生产。
VQ-Diffusion(CVPR2022,口服)和
改进的 VQ-Diffusion
概述
这是论文的官方存储库:用于文本到图像合成的矢量量化扩散模型 和 改进的矢量量化扩散模型。
代码与https://github.com/cientgu/VQ-Diffusion,相同,已经提出的一些问题可以参考。
VQ-Diffusion 基于 VQ-VAE,其潜在空间由最近开发的去噪扩散概率模型 (DDPM) 的条件变体建模。与具有相似参数数量的自回归模型相比,它产生明显更好的文本到图像生成结果。与之前基于GAN-的方法相比,VQ-Diffusion可以处理更复杂的场景并大幅提高合成图像质量。
框架
与 Diffusers 库集成
VQ-Diffusion 现在也可在 扩散器中使用,并可通过 VQDiffusionPipeline 访问。 Diffusers 允许您只需几行代码即可测试 VQ-Diffusion。
您可以按如下方式安装扩散器:
pip install diffusers torch accelerate transformers
然后只需几行代码即可尝试该模型:
import torch
from diffusers import VQDiffusionPipeline
pipeline = VQDiffusionPipeline.from_pretrained("microsoft/vq-diffusion-ithq", torch_dtype=torch.float16, revision="fp16")
pipeline = pipeline.to("cuda")
image = pipeline("teddy bear playing in the pool").images[0]
# save image
image.save("./teddy_bear.png")
您可以在这里找到ITHQ检查点的型号卡。
要求
我们建议使用 docker。另外,您可以运行:
bash install_req.sh
数据准备
微软 COCO
│MSCOCO_Caption/
├──annotations/
│ ├── captions_train2014.json
│ ├── captions_val2014.json
├──train2014/
│ ├── train2014/
│ │ ├── COCO_train2014_000000000009.jpg
│ │ ├── ......
├──val2014/
│ ├── val2014/
│ │ ├── COCO_val2014_000000000042.jpg
│ │ ├── ......
CUB-200
│CUB-200/
├──images/
│ ├── 001.Black_footed_Albatross/
│ ├── 002.Laysan_Albatross
│ ├── ......
├──text/
│ ├── text/
│ │ ├── 001.Black_footed_Albatross/
│ │ ├── 002.Laysan_Albatross
│ │ ├── ......
├──train/
│ ├── filenames.pickle
├──test/
│ ├── filenames.pickle
ImageNet
│imagenet/
├──train/
│ ├── n01440764
│ │ ├── n01440764_10026.JPEG
│ │ ├── n01440764_10027.JPEG
│ │ ├── ......
│ ├── ......
├──val/
│ ├── n01440764
│ │ ├── ILSVRC2012_val_00000293.JPEG
│ │ ├── ILSVRC2012_val_00002138.JPEG
│ │ ├── ......
│ ├── ......
预训练模型
我们发布了四个文本到图像预训练模型,在 Conceptual Caption、MSCOCO、CUB200 和 LAION- human 数据集上进行训练。另外,我们发布了ImageNet预训练模型,并提供CLIP预训练模型以方便使用。这些应该放在 OUTPUT/pretrained_model/ 下。 这些预训练的模型文件可能很大,因为它们是训练检查点,其中包含梯度信息、优化器信息、ema模型等。
此外,我们在 ITHQ、ImageNet、Conceptual Caption 和 MSCOCO 数据集上发布了四个具有可学习分类器的预训练模型。
我们在FFHQ、OpenImages和ImageNet数据集上提供了VQVAE模型,这些模型来自Taming Transformer,我们在这里提供它们是为了方便。请将它们放在 OUTPUT/pretrained_model/taming_dvae/ 下。
为了支持 ITHQ 数据集,我们在 ITHQ 数据集上训练了一个新的 VQVAE 模型。
为了您的方便,我们提供了用于下载所有模型的脚本。您可以运行bash vqdiffusion_download_checkpoints.sh。
推理
要从野外文本生成图像:
from inference_VQ_Diffusion import VQ_Diffusion
VQ_Diffusion_model = VQ_Diffusion(config='configs/ithq.yaml', path='OUTPUT/pretrained_model/ithq_learnable.pth')
# Inference VQ-Diffusion
VQ_Diffusion_model.inference_generate_sample_with_condition("teddy bear playing in the pool", truncation_rate=0.86, save_root="RESULT", batch_size=4)
# Inference Improved VQ-Diffusion with learnable classifier-free sampling
VQ_Diffusion_model.inference_generate_sample_with_condition("teddy bear playing in the pool", truncation_rate=1.0, save_root="RESULT", batch_size=4, guidance_scale=5.0)
VQ_Diffusion_model.inference_generate_sample_with_condition("a long exposure photo of waterfall", truncation_rate=1.0, save_root="RESULT", batch_size=4, guidance_scale=5.0)
# Inference Improved VQ-Diffusion with fast/high-quality inference
VQ_Diffusion_model.inference_generate_sample_with_condition("a long exposure photo of waterfall", truncation_rate=0.86, save_root="RESULT", batch_size=4, infer_speed=0.5) # high-quality inference, 0.5x inference speed
VQ_Diffusion_model.inference_generate_sample_with_condition("a long exposure photo of waterfall", truncation_rate=0.86, save_root="RESULT", batch_size=4, infer_speed=2) # fast inference, 2x inference speed
# infer_speed shoule be float in [0.1, 10], larger infer_speed means faster inference and smaller infer_speed means slower inference
# Inference Improved VQ-Diffusion with purity sampling
VQ_Diffusion_model.inference_generate_sample_with_condition("a long exposure photo of waterfall", truncation_rate=0.86, save_root="RESULT", batch_size=4, prior_rule=2, prior_weight=1) # purity sampling
# Inference Improved VQ-Diffusion with both learnable classifier-free sampling and fast inference
VQ_Diffusion_model.inference_generate_sample_with_condition("a long exposure photo of waterfall", truncation_rate=1.0, save_root="RESULT", batch_size=4, guidance_scale=5.0, infer_speed=2) # classifier-free guidance and fast inference
要根据 MSCOCO/CUB/CC 数据集上的给定文本生成图像:
from inference_VQ_Diffusion import VQ_Diffusion
VQ_Diffusion_model = VQ_Diffusion(config='OUTPUT/pretrained_model/config_text.yaml', path='OUTPUT/pretrained_model/coco_learnable.pth')
# Inference VQ-Diffusion
VQ_Diffusion_model.inference_generate_sample_with_condition("A group of elephants walking in muddy water", truncation_rate=0.86, save_root="RESULT", batch_size=4)
# Inference Improved VQ-Diffusion with learnable classifier-free sampling
VQ_Diffusion_model.inference_generate_sample_with_condition("A group of elephants walking in muddy water", truncation_rate=1.0, save_root="RESULT", batch_size=4, guidance_scale=3.0)
您可以将 coco_learnable.pth 更改为其他预训练模型来测试不同的文本。
从给定的 ImageNet 类标签生成图像:
from inference_VQ_Diffusion import VQ_Diffusion
# Inference VQ-Diffusion
VQ_Diffusion_model = VQ_Diffusion(config='OUTPUT/pretrained_model/config_imagenet.yaml', path='OUTPUT/pretrained_model/imagenet_pretrained.pth')
VQ_Diffusion_model.inference_generate_sample_with_class(407, truncation_rate=0.86, save_root="RESULT", batch_size=4)
# Inference Improved VQ-Diffusion with classifier-free sampling
VQ_Diffusion_model = VQ_Diffusion(config='configs/imagenet.yaml', path='OUTPUT/pretrained_model/imagenet_learnable.pth', imagenet_cf=True)
VQ_Diffusion_model.inference_generate_sample_with_class(407, truncation_rate=0.94, save_root="RESULT", batch_size=8, guidance_scale=1.5)
培训
首先,将 data_root 更改为 configs/coco.yaml 或其他配置中的正确路径。
在 MSCOCO 数据集上训练 Text2Image 生成:
python running_command/run_train_coco.py
在 CUB200 数据集上训练 Text2Image 生成:
python running_command/run_train_cub.py
在 ImageNet 数据集上训练条件生成:
python running_command/run_train_imagenet.py
在 FFHQ 数据集上训练无条件生成:
python running_command/run_train_ffhq.py
使用可学习的无分类器在 MSCOCO 数据集上微调 Text2Image 生成:
python running_command/run_tune_coco.py
引用 VQ-Diffusion
如果您发现我们的代码对您的研究有帮助,请考虑引用:
@article{gu2021vector,
title={Vector Quantized Diffusion Model for Text-to-Image Synthesis},
author={Gu, Shuyang and Chen, Dong and Bao, Jianmin and Wen, Fang and Zhang, Bo and Chen, Dongdong and Yuan, Lu and Guo, Baining},
journal={arXiv preprint arXiv:2111.14822},
year={2021}
}