ITADN
apple/ml-mdm
apple/ml-mdm · 文件 下载 ZIP
文件最后提交记录最后更新时间
README.md
以下内容由 AI 翻译,如有问题请点此提交 issue 反馈

ml_mdm - Matryoshka Diffusion Models

ml_mdm 是一个用于高效训练高质量文本到图像扩散模型的 Python 包 —— 由 Luke CarlsonJiatao GuShuangfei ZhaiNavdeep Jaitly 向公众发布。


本软件项目随附于研究论文,Matryoshka Diffusion Models

Jiatao Gu, Shuangfei Zhai, Yizhe Zhang, Josh Susskind, Navdeep Jaitly

[Paper] [BibTex]

mdm text to image outputs

目录

章节描述
引言Matryoshka Diffusion Models 的简要概述
安装使用 ml_mdm 开始训练模型和生成样本
预训练模型下载我们预训练模型(64, 256, 1024)的链接
Web 演示使用我们的 Web UI 生成图像
代码库结构Python 模块的概述
概念核心概念和设计原则。
教程在 CC12m 上逐步训练 MDM 模型

引言

扩散模型是生成高质量图像和视频的事实标准方法,但由于计算和优化方面的挑战,学习高维模型仍然是一项艰巨的任务。

ml_mdm 是一个用于高分辨率图像和视频合成的端到端框架——它以我们的技术命名:Matryoshka Diffusion Models

值得注意的是,我们可以在高达 1024x1024 像素的分辨率下训练单个像素空间模型,并使用仅包含 1200 万张图像的 CC12M 数据集展示了强大的零样本泛化能力。

mdm multi scale pipeline

安装

默认安装依赖项,如 pyproject.toml 中所述,经过选择,以便您即使在仅配备 CPU 的机器上也能安装此库。

用户已在 Python 3.9、3.10 以及 cuda_12、cuda-11.8 环境下运行过此代码库

> pip install -e .

开发者也应使用 pre-commit install 设置 pre-commit

运行测试用例

> pytest   # run test cases that can work with just cpu
> pytest  -m ''  # will run all test cases - including ones that require a gpu
> pytest -m gpu # run only gpu test cases

预训练模型

我们已将模型检查点上传至:

注意:我们发布的是在从 Flickr 收集的 5000 万文本-图像对上训练的模型。在本仓库中,我们提供了用于下载 CC12M 的脚本以及用于在 CC12M 数据上训练等效模型的配置。

您可以自由下载这些模型,或跳过以下内容以训练您自己的模型。一旦预训练模型下载到本地,您就可以在我们的 Web 演示中使用它,或将其作为参数传递给训练、采样等操作。

export ASSET_PATH=https://docs-assets.developer.apple.com/ml-research/models/mdm

curl $ASSET_PATH/flickr64/vis_model.pth --output vis_model_64x64.pth
curl $ASSET_PATH/flickr256/vis_model.pth --output vis_model_256x256.pth
curl $ASSET_PATH/flickr1024/vis_model.pth --output vis_model_1024x1024.pth

Web Demo

在(下载检查点之后)你可以使用以下命令运行你自己的 web demo 实例:

torchrun --standalone --nproc_per_node=1  ml_mdm/clis/generate_sample.py --port $YOUR_PORT

image

代码库

1. /configs

模块描述
configs.dataset_creation用于将数据集划分为训练-评估-验证流水线的配置文件
configs.datasets模型训练和评估阶段的数据集
configs.models不同分辨率模型的配置文件

2. /data

moduledescription
data
  • bert.vocab: 包含 token 及其关联向量表示的 BERT 训练词典
  • c4_wpm.vocab: 包含 token 及其关联向量表示的 C4 训练词典
  • cifar10.vocab: 包含 token 及其关联向量表示的 CIFAR10 训练词典
  • imagenet.vocab: 与 Imagenet 数据集相关的提示词
  • prompts_cc12m-64x64.tsv: 与 cc12m 数据集相关的提示词,用于 64x64 分辨率模型
  • prompts_cc12m-256x256.tsv: 与 cc12m 数据集相关的提示词,用于 256x256 分辨率模型
  • prompts_cifar10-32x32.tsv: 与 cifar10 数据集相关的提示词,用于 32x32 分辨率模型
  • prompts_cifar10-64x64.tsv: 与 cifar10 数据集相关的提示词,用于 64x64 分辨率模型
  • prompts_demo.tsv: 额外的演示提示词
  • prompts_imagenet-64px.tsv: 与 imagenet 数据集相关的提示词,用于 64x64 分辨率模型
  • prompts_WebImage-ALIGN-64px.tsv: 与 WebImage-ALIGN 数据集相关的提示词,用于 64x64 分辨率模型
  • t5.vocab: 包含 token 及其关联向量表示的 t5 训练词典
  • tokenizer_spm_32000_50m.vocab: 包含 token 及其关联向量表示的 SPM 训练词典

3. /docs

模块描述
docs
  • web_demo.png: 模型 Web 演示的截图

4. /ml_mdm

模块描述
ml_mdm.models核心模型实现
ml_mdm.diffusion模型流水线,例如 DDPM
ml_mdm.config使用 simple parsing 将配置数据类与相关的模型、流水线和 CLI 连接起来
ml_mdm.clis项目中所有的命令行工具,其中最相关的是 train_parallel.py
tests/单元测试和示例训练文件

5. /tests

模块描述
tests.test_files用于测试的示例文件

概念

ml_mdm.models

ml_mdm.models 子模块中,我们开源了以下实现:

  • U-Nets
  • Nested U-Nets

ml_mdm.config

ml_mdm.config 包含核心配置和 CLI 逻辑。此代码库中的许多模型、CLI 和函数通过传入 dataclass 对象进行配置。我们使用 SimpleParsing 动态创建 CLI,并允许通过 --config_path 传入 yaml config 表示。

本质上,simple_parsing 会将所有传入的 CLI 参数和 yaml 文件转换为整洁的配置类,例如 ml_mdm.reader.ReaderConfigml_mdm.diffusion.DiffusionConfig

ml_mdm.configMODEL_REGISTRYMODEL_CONFIG_REGISTRYPIPELINE_REGISTRYPIPELINE_CONFIG_REGISTRY 中存储名称到类的全局映射。

MODEL_REGISTRYPIPELINE_REGISTRY 存储的信息如下例所示:

*_CONFIG_REGISTRY[architecture name]["model"] = model name

*_CONFIG_REGISTRY[architecture name]["config"] = configuration class

MODEL_CONFIG_REGISTRY 和 PIPELINE_CONFIG_REGISTRY 存储的信息如下例所示:

*_CONFIG_REGISTRY[architecture name]["model"] = model name

*_CONFIG_REGISTRY[architecture name]["config"] = configuration class

architecture name 和 model name 通过函数参数 *names 传入 ml_mdm.config。其中 *names 指向 "architecture name"、"model name"

教程

使用预训练检查点生成您自己的图像

安装 ml_mdm 后,将这些检查点下载到仓库目录中。

curl https://docs-assets.developer.apple.com/ml-research/models/mdm/flickr64/vis_model.pth --output vis_model_64x64.pth
curl https://docs-assets.developer.apple.com/ml-research/models/mdm/flickr256/vis_model.pth --output vis_model_256x256.pth

Web 演示将使用相应的配置加载每个模型:

  • vis_model_64x64.pth 将使用 configs/models/cc12m_64x64.yaml 中的设置进行加载
  • vis_model_256x256.pth 将使用 configs/models/cc12m_256x256.yaml 中的设置进行加载
  • vis_model_1024x1024.pth 将使用 configs/models/cc12m_1024x1024.yaml 中的设置进行加载

在演示中,您可以更改各种设置并查看模型的内部结构。通过替换 $YOUR_PORT 来设置您想要使用的端口,然后运行:

torchrun --standalone --nproc_per_node=1  ml_mdm/clis/generate_sample.py --port $YOUR_PORT

在虚拟数据上训练

如果你只是想逐步了解训练模型和运行流水线的过程,而不想下载大型数据集,我们为你准备了一个最小示例。它使用了来自 tests/test_files/ 的虚拟数据

你可以随意尝试修改各种 --args,无论是直接在 cli 中修改,还是通过编辑 config yaml 文件

torchrun --standalone --nproc_per_node=1 ml_mdm/clis/train_parallel.py \
 --file-list=tests/test_files/sample_training_0.tsv \
 --multinode=0 \
  --output-dir=outputs    --config_path configs/models/cc12m_64x64.yaml \
  -num_diffusion_steps=10 \
	--num-training-steps=10

你应该会看到一个 outputs/vis_model_000100.pth 文件。现在让我们做一件更有意义的事情:

让我们在 CC12m 上训练一个 MDM 模型

1. 数据准备:

(可选)使用此示例参数下载 CC12m 的前 1K 个文件

该脚本基于 img2dataset 的 CC12M 脚本

curl https://storage.googleapis.com/conceptual_12m/cc12m.tsv | head -n 1000 > cc12m_index.tsv

# Add headers to the file
sed -i '1s/^/url\tcaption\n/'  cc12m_index.tsv

注意:如果你想要完整的 cc12m,请从调用中移除 | head -n 1000

然后准备并拆分为训练/验证集

此脚本需要 img2dataset,请运行 pip install '.[data_prep]' 或仅运行 pip install img2dataset

python3 -m ml_mdm.clis.scrape_cc12m \
  --cc12m_index cc12m_index.tsv \
  --cc12m_local_dir cc12m_download

运行此命令后,你将看到以下文件:

training.0.tsv # train index file
validation.tsv # validation index file
cc12m_download/
   00000.parquet  00000.tar  00000.tsv  00000_stats.json  validation.tsv
   00001.parquet ....

2. 训练

现在我们有了训练文件,我们可以选择一个模型配置并传递任何额外的训练参数:

# Modify torchrun arguments to fit your GPU setup
torchrun --standalone --nproc_per_node=8 ml_mdm/clis/train_parallel.py \
  --file-list=training_0.tsv \
  --multinode=0 --output-dir=/mnt/data/outputs \
  --config_path configs/models/cc12m_64x64.yaml \
  --num-training-steps=100   --warmup-steps 10

注意:configs/models/cc12m_64x64.yaml 包含更多参数,请查看以获取更多详细信息。

如果你已经下载了一个预训练模型,你可以将 --pretrained-vision-file 参数设置为其在磁盘上的位置

训练完成后,你可以在 --output-dir 参数定义的文件夹中找到模型:

2024-07-22:17:58:46,649 INFO     [model_ema.py:33] Saving EMA model file: /mnt/data/outputs/vis_model_000100.pth
2024-07-22:17:58:47,448 INFO     [unet.py:794] Saving model file: /mnt/data/outputs/vis_model_noema_000100.pth

3. 从模型采样

现在我们有了一个训练好的模型,我们可以从扩散模型生成样本:

torchrun --standalone --nproc_per_node=1 ml_mdm/clis/generate_batch.py \
  --config_path configs/models/cc12m_64x64.yaml \
  --min-examples 3 --test-file-list validation.tsv \
  --sample-image-size 64 --model-file /mnt/data/outputs/vis_model_000100.pth

如果你想跳过训练步骤,你可以更新 --model-file 以指向我们的一个预训练模型

数据集存储

为了长期存储,你可以选择将数据上传到 s3://{your_bucket}/datasets/{datasetname}/*.[tar,tsv]

然后更新 configs/datasets/cc12m.yaml 以指向你的 s3 路径。

# configs/datasets/cc12m.yaml
train:
  files:
    - s3://mlx/datasets/cc12m-64x64/images_00.*.tsv
eval:
  files:
    - s3://mlx/datasets/cc12m-64x64/validation.tsv
# configs/datasets/reader_config.yaml
reader_config:
  append_eos: true
  bucket: ${your_bucket} # add your s3 bucket
  endpoint_url: None # boto will automatically infer the endpoint

然后你可以使用我们的数据集下载助手:

python -m ml_mdm.clis.download_tar_from_index \
  --dataset_config_file configs/datasets/cc12m.yaml \
  --subset train --download_tar

python -m ml_mdm.clis.download_tar_from_index \
  --dataset_config_file configs/datasets/cc12m.yaml \
  --subset eval --download_tar

S3 数据集选择

请查看 configs/datasets/cc12m.yaml

该代码支持提供多个正则表达式。请注意,这些 正则表达式不是通配符 -- 它们是来自 python re 库的正则表达式。 因此,如果你只想使用 WebImage 中 1000 个 tar 文件中的 100 个进行训练,你可以 执行以下操作:

train:
  files:
    - s3://mlx/datasets/example-dataset-100M_64px/example-dataset-100M-00[0-1]..-[0-9]*-of-01000.tsv
eval:
  files:
    - s3://mlx/datasets/example-dataset-100M_64px/validation.tsv

你也可以混合搭配这些文件。因此,如果你想合并 CC12m 和 imagenet,你可以 创建一个包含以下内容的新的 yaml 文件:

train:
  files:
    - s3://mlx/datasets/imagenet-64px/imagenet-train-000??-of-00100.tsv
    - s3://mlx/datasets/cc12m-64x64/images_00.*.tsv
eval:
  files:
    - s3://mlx/datasets/cc12m-64x64/validation.tsv

数据集结构

S3 存储桶中包含一系列采用此格式的文件,请查看 ml_mdm/clis/scrape_cc12m.py 以生成您自己的文件。

2023-04-01 01:31:30   36147200 images_00000.tar
2023-05-10 11:34:49    1108424 images_00000.tsv
2023-04-01 01:31:26   36454400 images_00001.tar
2023-05-10 11:34:49    1109588 images_00001.tsv
2023-04-01 01:31:53   36116480 images_00002.tar
...

这些文件的最小表示形式可在 tests/test_files/ 中找到。

引用

如果您觉得我们的工作有用,请考虑引用我们:

@misc{gu2023matryoshkadiffusionmodels,
      title={Matryoshka Diffusion Models},
      author={Jiatao Gu and Shuangfei Zhai and Yizhe Zhang and Josh Susskind and Navdeep Jaitly},
      year={2023},
      eprint={2310.15111},
      archivePrefix={arXiv},
      primaryClass={cs.CV},
      url={https://arxiv.org/abs/2310.15111},
}