ml_mdm - Matryoshka Diffusion Models
ml_mdm 是一个用于高效训练高质量文本到图像扩散模型的 Python 包 —— 由 Luke Carlson、Jiatao Gu、Shuangfei Zhai 和 Navdeep Jaitly 向公众发布。
本软件项目随附于研究论文,Matryoshka Diffusion Models。
Jiatao Gu, Shuangfei Zhai, Yizhe Zhang, Josh Susskind, Navdeep Jaitly

目录
| 章节 | 描述 |
|---|---|
| 引言 | Matryoshka Diffusion Models 的简要概述 |
| 安装 | 使用 ml_mdm 开始训练模型和生成样本 |
| 预训练模型 | 下载我们预训练模型(64, 256, 1024)的链接 |
| Web 演示 | 使用我们的 Web UI 生成图像 |
| 代码库结构 | Python 模块的概述 |
| 概念 | 核心概念和设计原则。 |
| 教程 | 在 CC12m 上逐步训练 MDM 模型 |
引言
扩散模型是生成高质量图像和视频的事实标准方法,但由于计算和优化方面的挑战,学习高维模型仍然是一项艰巨的任务。
ml_mdm 是一个用于高分辨率图像和视频合成的端到端框架——它以我们的技术命名:Matryoshka Diffusion Models。
值得注意的是,我们可以在高达 1024x1024 像素的分辨率下训练单个像素空间模型,并使用仅包含 1200 万张图像的 CC12M 数据集展示了强大的零样本泛化能力。

安装
默认安装依赖项,如 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
预训练模型
我们已将模型检查点上传至:
- https://docs-assets.developer.apple.com/ml-research/models/mdm/flickr64/vis_model.pth
- https://docs-assets.developer.apple.com/ml-research/models/mdm/flickr256/vis_model.pth
- https://docs-assets.developer.apple.com/ml-research/models/mdm/flickr1024/vis_model.pth
注意:我们发布的是在从 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

代码库
1. /configs
| 模块 | 描述 |
|---|---|
configs.dataset_creation | 用于将数据集划分为训练-评估-验证流水线的配置文件 |
configs.datasets | 模型训练和评估阶段的数据集 |
configs.models | 不同分辨率模型的配置文件 |
2. /data
| module | description |
|---|---|
data |
|
3. /docs
| 模块 | 描述 |
|---|---|
docs |
|
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.ReaderConfig、ml_mdm.diffusion.DiffusionConfig。
ml_mdm.config 在 MODEL_REGISTRY、MODEL_CONFIG_REGISTRY、PIPELINE_REGISTRY 和 PIPELINE_CONFIG_REGISTRY 中存储名称到类的全局映射。
MODEL_REGISTRY 和 PIPELINE_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},
}