Compare commits

..
58 Commits
Author SHA1 Message Date
jiangzeyinzi 9f5b4399bf Merge pull request #105 from yaosheng216/patch-24
Update constant.py
2025-04-03 14:00:15 +08:00
Great ddba40d0d0 Update constant.py 2025-04-03 11:56:35 +08:00
jiangzeyinzi dd4e2a7f42 Merge pull request #103 from modelscope/v1.4.1_dev
update v1.4.1
2025-04-02 19:30:49 +08:00
zeyinzi.jzyz 6c8af8d7a8 update v1.4.1 2025-04-02 19:27:43 +08:00
jiangzeyinzi 467652bd69 Merge pull request #91 from yaosheng216/patch-23
Update calculator_node.py
2025-02-17 11:27:48 +08:00
Great ae59b96f1d Update calculator_node.py 2025-02-17 11:23:24 +08:00
jiangzeyinzi 2c75835035 Merge pull request #90 from yaosheng216/patch-22
Update pyproject.toml
2025-02-17 10:37:18 +08:00
Great 5cbcd3ee04 Update pyproject.toml 2025-02-17 10:34:30 +08:00
jiangzeyinzi 0f12e4db73 Merge pull request #87 from yaosheng216/patch-19
Update scepter_workflow.yaml
2025-02-14 16:47:01 +08:00
jiangzeyinzi a9b6337ae2 Merge pull request #89 from yaosheng216/patch-21
Create calculator_node.py
2025-02-14 16:46:49 +08:00
jiangzeyinzi 8526ae0234 Merge pull request #88 from yaosheng216/patch-20
Update __init__.py
2025-02-14 16:46:34 +08:00
Great 1cd5604e7b Create calculator_node.py 2025-02-14 16:45:00 +08:00
Great bbb8f35f49 Update __init__.py 2025-02-14 16:43:11 +08:00
Great 825e8c1cdb Update scepter_workflow.yaml 2025-02-14 16:41:14 +08:00
jiangzeyinzi ff3ccd6050 Merge pull request #85 from yaosheng216/patch-17
Update pyproject.toml
2025-02-13 17:50:21 +08:00
jiangzeyinzi 9ad6f9dc5c Merge pull request #86 from yaosheng216/patch-18
Update publish.yml
2025-02-13 17:50:05 +08:00
Great 755b39b968 Update publish.yml 2025-02-13 17:47:13 +08:00
Great 3427d630a2 Rename scepter/workflow/config/pyproject.toml to scepter/workflow/pyproject.toml 2025-02-13 17:46:21 +08:00
Great d9cd203be6 Update pyproject.toml 2025-02-13 17:42:26 +08:00
jiangzeyinzi 0a10447558 Add files via upload 2025-02-13 17:41:35 +08:00
jiangzeyinzi 32b48d2f08 Merge pull request #84 from yaosheng216/patch-16
Update pyproject.toml
2025-02-13 17:37:46 +08:00
Great 2122221697 Update pyproject.toml 2025-02-13 17:37:00 +08:00
jiangzeyinzi e591d2e4cb Merge pull request #83 from yaosheng216/patch-15
Update and rename pyproject.toml to scepter/workflow/config/pyproject…
2025-02-13 17:30:26 +08:00
jiangzeyinzi 33b8adda82 Merge pull request #81 from yaosheng216/patch-13
Update publish.yml
2025-02-13 17:30:13 +08:00
Great 1ef6b4d4ec Update pyproject.toml 2025-02-13 17:27:24 +08:00
Great cc3e6868ce Update pyproject.toml 2025-02-13 17:26:09 +08:00
Great 1758959cb9 Update and rename pyproject.toml to scepter/workflow/config/pyproject.toml 2025-02-13 17:20:20 +08:00
Great 1d7b829c17 Update publish.yml 2025-02-13 17:16:32 +08:00
jiangzeyinzi 59c5fadd77 Update pyproject.toml 2025-02-11 13:31:29 +08:00
jiangzeyinzi cb23134173 Merge pull request #80 from yaosheng216/patch-12
Update pyproject.toml
2025-02-11 11:01:14 +08:00
jiangzeyinzi c4a593b02f Merge pull request #79 from yaosheng216/patch-11
Update framework.txt
2025-02-11 11:01:02 +08:00
Great 7dde2741fb Update pyproject.toml 2025-02-11 10:56:14 +08:00
Great 3713445817 Update framework.txt 2025-02-11 10:54:32 +08:00
jiangzeyinzi 5af3ae0eeb Merge pull request #72 from ComfyNodePRs/publish
Add Github Action for Publishing to Comfy Registry
2025-02-10 14:58:39 +08:00
jiangzeyinzi c3376064fd Merge pull request #71 from ComfyNodePRs/pyproject
Add pyproject.toml for Custom Node Registry
2025-02-10 14:58:29 +08:00
jiangzeyinzi ff7b6c1891 Merge pull request #76 from yaosheng216/patch-9
Update self_train.py
2025-02-10 14:42:16 +08:00
jiangzeyinzi db65a08cea Merge pull request #77 from yaosheng216/patch-8
Update webui.py
2025-02-10 14:42:05 +08:00
Great 3083d509de Update self_train.py 2025-02-10 14:36:08 +08:00
Great 740a6c87e0 Update webui.py 2025-02-10 14:34:06 +08:00
mcj 9d5b3a7101 Merge pull request #75 from modelscope/v1.4.0_dev
V1.4.0 dev
2025-02-10 09:55:33 +08:00
jiangzeyinzi c4b4d88e75 update readme 2025-02-10 09:01:29 +08:00
jiangzeyinzi 043222de49 update 1.4.0 2025-02-03 13:36:44 +08:00
snomiao 171c7ce0b3 chore(publish): Add Github Action for Publishing to Comfy Registry 2024-12-25 06:52:15 +00:00
snomiao 11c69121cd chore(pyproject): Add pyproject.toml for Custom Node Registry 2024-12-25 06:52:15 +00:00
皓童 d7dbdc5292 modify default path 2024-12-05 16:05:30 +08:00
mcj 711c10a68c Merge pull request #69 from modelscope/v1.3.0_dev
add init file for chatbot
2024-12-05 15:23:37 +08:00
皓童 448cdba522 add init file for chatbot 2024-12-05 15:20:00 +08:00
jiangzeyinzi adda36e39d Merge pull request #66 from yaosheng216/patch-7
Update model_node.py
2024-11-26 14:58:52 +08:00
jiangzeyinzi 1da1864993 Merge pull request #65 from yaosheng216/patch-6
Update parameter_node.py
2024-11-26 14:58:36 +08:00
Great 5eac362325 Update model_node.py 2024-11-26 14:57:20 +08:00
Great 9ae3bca43d Update parameter_node.py 2024-11-26 14:55:29 +08:00
jiangzeyinzi ca034ef765 Merge pull request #64 from modelscope/v1.3.0_dev
V1.3.0 dev
2024-11-26 12:13:08 +08:00
皓童 2a29446d45 modify ace inference and ace yaml 2024-11-25 14:24:54 +08:00
maochaojie cd33b4ab15 Merge branch 'v1.3.0_dev' of https://github.com/modelscope/scepter into v1.3.0_dev 2024-11-21 15:42:05 +08:00
maochaojie 7d7943fed3 modify yaml and workflow 2024-11-21 15:41:45 +08:00
jiangzeyinzi d48b2f110f Merge pull request #63 from yaosheng216/patch-5
Update model_node.py
2024-11-20 13:26:33 +08:00
Great 82486adf38 Update model_node.py 2024-11-20 13:24:28 +08:00
maochaojie a683061c6f upgrade from 1.2.0 to 1.3.0 2024-11-19 19:20:02 +08:00
198 changed files with 17917 additions and 1877 deletions
+24
View File
@@ -0,0 +1,24 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
- master
paths:
- "scepter/workflow/pyproject.toml"
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
# if this is a forked repository. Skipping the workflow.
if: github.event.repository.fork == false
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+184
View File
@@ -0,0 +1,184 @@
<p align="center">
<h2 align="center"><img src="https://raw.githubusercontent.com/ali-vilab/ACE/refs/heads/main/assets/figures/icon.png" height=16> : All-round Creator and Editor Following <br> Instructions via Diffusion Transformer</h2>
<p align="center">
<a href="https://arxiv.org/abs/2410.00086"><img src='https://img.shields.io/badge/arXiv-ACE-red' alt='Paper PDF'></a>
<a href='https://ali-vilab.github.io/ace-page'><img src='https://img.shields.io/badge/Project_Page-ACE-blue' alt='Project Page'></a>
<a href='https://github.com/modelscope/scepter'><img src='https://img.shields.io/badge/Scepter-ACE-green'></a>
<a href='https://huggingface.co/spaces/scepter-studio/ACE-Chat'><img src='https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Space-orange'></a>
<a href='https://huggingface.co/scepter-studio/ACE-0.6B-512px'><img src='https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Model-orange'></a>
<a href='https://www.modelscope.cn/models/iic/ACE-0.6B-512px'><img src='https://img.shields.io/badge/ModelScope-Model-purple'></a>
<br>
<strong>Zhen Han*</strong>
·
<strong>Zeyinzi Jiang*</strong>
·
<strong>Yulin Pan*</strong>
·
<strong>Jingfeng Zhang*</strong>
·
<strong>Chaojie Mao*</strong>
<br>
<strong>Chenwei Xie</strong>
·
<strong>Yu Liu</strong>
·
<strong>Jingren Zhou</strong>
<br>
Tongyi Lab, Alibaba Group
</p>
<table align="center">
<tr>
<td>
<img src="https://raw.githubusercontent.com/ali-vilab/ACE/refs/heads/main/assets/figures/teaser.png">
</td>
</tr>
</table>
## 🚀 Installation
Install the necessary packages with `pip`:
```bash
pip install -r requirements.txt
```
## 🔥 ACE Models
| **Model** | **Status** |
|:----------------:|:---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------:|
| ACE-0.6B-512px | [![Demo link](https://img.shields.io/badge/Demo-ACE_Chat-purple)](https://huggingface.co/spaces/scepter-studio/ACE-Chat)<br>[![ModelScope link](https://img.shields.io/badge/ModelScope-Model-blue)](https://www.modelscope.cn/models/iic/ACE-0.6B-512px) [![HuggingFace link](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Model-yellow)](https://huggingface.co/scepter-studio/ACE-0.6B-512px) |
| ACE-0.6B-1024px | [![Demo link](https://img.shields.io/badge/Demo-ACE_Refiner_Chat-purple)](https://huggingface.co/spaces/scepter-studio/ACE-Refiner-Chat)<br>[![ModelScope link](https://img.shields.io/badge/ModelScope-Model-blue)](https://www.modelscope.cn/models/iic/ACE-0.6B-1024px) [![HuggingFace link](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Model-yellow)](https://huggingface.co/scepter-studio/ACE-0.6B-1024px) | |
## 🖼 Model Performance Visualization
The current model's parameters scale of ACE is 0.6B, which imposes certain limitations on the quality of image generation. [FLUX.1-Dev](https://huggingface.co/black-forest-labs/FLUX.1-dev), on the other hand,
has a significant advantage in text-to-image generation quality. By using SDEdit, we can effectively leverage the generative capabilities of FLUX to further enhance the image results generated by ACE. Based on the above considerations, we have designed the ACE-Refiner pipeline, as shown in the diagram below.
![ACE_REFINER](https://raw.githubusercontent.com/ali-vilab/ACE/refs/heads/main/assets/ace_method/ace_refiner_process.webp)
As shown in the figure below, when the strength
σ of the generated image is high, the generated image will suffer from fidelity loss compared to the original image. Conversely, lower
σ does not significantly improve the image quality. Therefore, users can make a trade-off between fidelity to the generated result and the image quality based on their own needs.
Users can set the value of "REFINER_SCALE" in the configuration file `config/inference_config/models/ace_0.6b_1024_refiner.yaml`.
We recommend that users use the advance options in the [webui-demo](#-chat-bot-) for effect verification.
![ACE_REFINER_EXAMPLE](https://raw.githubusercontent.com/ali-vilab/ACE/refs/heads/main/assets/ace_method/ace_refiner.webp)
We compared the generation and editing performance of different models on several tasks, as shown as following.
![Samples](https://raw.githubusercontent.com/ali-vilab/ACE/refs/heads/main/assets/ace_method/samples_compare.webp)
## 🔥 Training
We offer a demonstration training YAML that enables the end-to-end training of ACE using a toy dataset. For a comprehensive overview of the hyperparameter configurations, please consult `config/ace_0.6b_512_train.yaml`.
### Prepare datasets
Please find the dataset class located in `modules/data/dataset/dataset.py`,
designed to facilitate end-to-end training using an open-source toy dataset.
Download a dataset zip file from [modelscope](https://www.modelscope.cn/models/iic/scepter/resolve/master/datasets/hed_pair.zip), and then extract its contents into the `cache/datasets/` directory.
Should you wish to prepare your own datasets, we recommend consulting `modules/data/dataset/dataset.py` for detailed guidance on the required data format.
### Prepare initial weight
The ACE checkpoint has been uploaded to both ModelScope and HuggingFace platforms:
* [ModelScope](https://www.modelscope.cn/models/iic/ACE-0.6B-512px)
* [HuggingFace](https://huggingface.co/scepter-studio/ACE-0.6B-512px)
In the provided training YAML configuration, we have designated the Modelscope URL as the default checkpoint URL. Should you wish to transition to Hugging Face, you can effortlessly achieve this by modifying the PRETRAINED_MODEL value within the YAML file (replace the prefix "ms://iic" to "hf://scepter-studio").
### Start training
You can easily start training procedure by executing the following command:
```bash
# ACE-0.6B-512px
PYTHONPATH=. python tools/run_train.py --cfg config/ace_0.6b_512_train.yaml
# ACE-0.6B-1024px
PYTHONPATH=. python tools/run_train.py --cfg config/ace_0.6b_1024_train.yaml
```
## 🚀 Inference
We provide a simple inference demo that allows users to generate images from text descriptions.
```bash
PYTHONPATH=. python tools/run_inference.py --cfg config/inference_config/models/ace_0.6b_512.yaml --instruction "make the boy cry, his eyes filled with tears" --seed 199999 --input_image examples/input_images/example0.webp
```
We recommend runing the examples for quick testing. Running the following command will run the example inference and the results will be saved in `examples/output_images/`.
```bash
PYTHONPATH=. python tools/run_inference.py --cfg config/inference_config/models/ace_0.6b_512.yaml
```
## 💬 Chat Bot
We have developed an chatbot UI utilizing Gradio, designed to transform user input in natural language into visually stunning images that align semantically with the provided instructions. Users can effortlessly initiate the chatbot app by executing the following command:
```bash
python chatbot/run_gradio.py --cfg chatbot/config/chatbot_ui.yaml --server_port 2024
```
<table align="center">
<tr>
<td>
<img src="https://raw.githubusercontent.com/ali-vilab/ACE/refs/heads/main/assets/videos/demo_chat.gif">
</td>
</tr>
</table>
## ⚙️️ ComfyUI Workflow
![Workflow](https://raw.githubusercontent.com/ali-vilab/ACE/refs/heads/main/assets/comfyui/ace_example.jpg)
We support the use of ACE in the ComfyUI Workflow through the following methods:
1) Automatic installation directly via the ComfyUI Manager by searching for the **ComfyUI-Scepter** node.
2) Manually install by moving custom_nodes from Scepter to ComfyUI.
```shell
git clone https://github.com/modelscope/scepter.git
cd path/to/scepter
pip install -e .
cp -r path/to/scepter/workflow/ path/to/ComfyUI/custom_nodes/ComfyUI-Scepter
cd path/to/ComfyUI
python main.py
```
**Note**: You can use the nodes by dragging the sample images below into ComfyUI. Additionally, our nodes can automatically pull models from ModelScope or HuggingFace by selecting the *model_source* field, or you can place the already downloaded models in a local path.
<table><tbody>
<tr>
<th align="center" colspan="4">ACE Workflow Examples</th>
</tr>
<tr>
<th align="center" colspan="1">Control</th>
<th align="center" colspan="1">Semantic</th>
<th align="center" colspan="1">Element</th>
</tr>
<tr>
<td>
<a href="https://raw.githubusercontent.com/ali-vilab/ACE/refs/heads/main/assets/comfyui/ace_control.png" target="_blank">
<img src="https://raw.githubusercontent.com/ali-vilab/ACE/refs/heads/main/assets/comfyui/ace_control.png" width="200">
</a>
</td>
<td>
<a href="https://raw.githubusercontent.com/ali-vilab/ACE/refs/heads/main/assets/comfyui/ace_semantic.png" target="_blank">
<img src="https://raw.githubusercontent.com/ali-vilab/ACE/refs/heads/main/assets/comfyui/ace_semantic.png" width="200">
</a>
</td>
<td>
<a href="https://raw.githubusercontent.com/ali-vilab/ACE/refs/heads/main/assets/comfyui/ace_element.png" target="_blank">
<img src="https://raw.githubusercontent.com/ali-vilab/ACE/refs/heads/main/assets/comfyui/ace_element.png" width="200">
</a>
</td>
</tr>
</tbody>
</table>
## 📝 Citation
```bibtex
@article{han2024ace,
title={ACE: All-round Creator and Editor Following Instructions via Diffusion Transformer},
author={Han, Zhen and Jiang, Zeyinzi and Pan, Yulin and Zhang, Jingfeng and Mao, Chaojie and Xie, Chenwei and Liu, Yu and Zhou, Jingren},
journal={arXiv preprint arXiv:2410.00086},
year={2024}
}
```
-10
View File
@@ -1,10 +0,0 @@
name: scepter
channels:
- defaults
dependencies:
- python==3.8
- pip>=20.3
- numpy>=1.23.1
- pip:
- -r requirements/recommended.txt
- -r requirements.txt
+59 -108
View File
@@ -18,7 +18,9 @@ SCEPTER offers 3 core components:
## 🎉 News
- [🔥🔥🔥2024.10]: We are pleased to announce the release of the code for [ACE](https://arxiv.org/abs/2410.00086), supporting Customized Training / Comfy UI Workflow / gradio-based ChatBot Interface. The detailed documents can be found at [ACE repo](https://github.com/ali-vilab/ACE.git).
- [🔥🔥🔥 2025.01]: We report ACE++, an instruction-based diffusion framework that tackles various image generation and editing tasks. The code and paper is available on [ACE++](https://ali-vilab.github.io/ACE_plus_page/).
- [2024.11]: Supports video files, video annotation, caption translation in data management, and inference & training of the [CogVideoX](https://arxiv.org/abs/2408.06072).
- [2024.10]: We are pleased to announce the release of the code for [ACE](https://arxiv.org/abs/2410.00086), supporting Customized Training / Comfy UI Workflow / gradio-based ChatBot Interface.
- [2024.10]: Support for inference and tuning with [FLUX](https://huggingface.co/black-forest-labs/FLUX.1-dev), as well as for building [ComfyUI](https://github.com/comfyanonymous/ComfyUI) workflows using this framework.
- [2024.09]: We introduce **ACE**, an **A**ll-round **C**reator and **E**ditor adept at executing a diverse array of image editing tasks tailored to your specifications. Built upon the cutting-edge Diffusion Transformer architecture, ACE has been extensively trained on a comprehensive dataset to seamlessly interpret and execute any natural language instruction. For further information, please consult the [project page](https://ali-vilab.github.io/ace-page/).
- [2024.07]: Support the inference and training of open-source generative models based on the [DiT](https://arxiv.org/abs/2212.09748) architecture, such as [SD3](https://arxiv.org/pdf/2403.03206) and [PixArt](https://arxiv.org/abs/2310.00426).
@@ -31,112 +33,65 @@ SCEPTER offers 3 core components:
- [2023.12]: We propose [SCEdit](https://arxiv.org/abs/2312.11392), an efficient and controllable generation framework.
- [2023.12]: We release [🪄SCEPTER](https://github.com/modelscope/scepter/) library.
[//]: # (## 🖼 Gallery for Recent Works)
## 🖼 Gallery for Recent Works
[//]: # ()
[//]: # (### FLUX Tuners)
### ACE
[//]: # ()
[//]: # (<table><tbody>)
ACE is a unified foundational model framework that supports a wide range of visual generation tasks. By defining CU for unifying multi-modal inputs across different tasks and incorporating long-context CU, we introduce historical contextual information into visual generation tasks, paving the way for ChatGPT-like dialog systems in visual generation.
[//]: # ( <tr>)
[![Watch the demo](https://ali-vilab.github.io/ace-page/static/images/tasks.png)](https://ali-vilab.github.io/ace-page/)
[//]: # ( <th align="center" colspan="3">Yarn Style</th>)
#### ACE Training
[//]: # ( <th align="center" colspan="3">Soft Watercolor Style</th>)
We offer a demonstration training YAML that enables the end-to-end training of ACE using a toy dataset. For a comprehensive overview of the hyperparameter configurations, please consult `scepter/methods/edit/dit_ace_0.6b_512.yaml`.
[//]: # ( </tr>)
##### Prepare datasets
[//]: # ( <tr>)
Please find the dataset class located in `scepter/modules/data/dataset/ms_dataset.py`,
designed to facilitate end-to-end training using an open-source toy dataset.
Download a dataset zip file from [modelscope](https://www.modelscope.cn/models/iic/scepter/resolve/master/datasets/hed_pair.zip), and then extract its contents into the `cache/datasets/` directory.
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_2_1.webp" width="200"></td>)
Should you wish to prepare your own datasets, we recommend consulting `scepter/modules/data/dataset/ms_dataset.py` for detailed guidance on the required data format.
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_2_2.webp" width="200"></td>)
##### Prepare initial weight
The ACE checkpoint has been uploaded to both ModelScope and HuggingFace platforms:
* [ModelScope](https://www.modelscope.cn/models/iic/ACE-0.6B-512px)
* [HuggingFace](https://huggingface.co/scepter-studio/ACE-0.6B-512px)
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_2_3.webp" width="200"></td>)
In the provided training YAML configuration, we have designated the Modelscope URL as the default checkpoint URL. Should you wish to transition to Hugging Face, you can effortlessly achieve this by modifying the PRETRAINED_MODEL value within the YAML file (replace the prefix "ms://iic" to "hf://scepter-studio").
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_1_1.webp" width="200"></td>)
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_1_2.webp" width="200"></td>)
##### Start training
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_1_3.webp" width="200"></td>)
You can easily start training procedure by executing the following command:
```bash
PYTHONPATH=. python scepter/tools/run_train.py --cfg scepter/methods/edit/dit_ace_0.6b_512.yaml
```
[//]: # ( </tr>)
#### ACE Chat Bot
[//]: # ( <tr>)
We have developed a chatbot interface utilizing Gradio, designed to convert user input in natural language into visually captivating images that align semantically with the specified instructions. You can easily access this functionality by launching Scepter Studio with the following command:
```bash
PYTHONPATH=. python scepter/tools/webui.py --cfg scepter/methods/studio/scepter_ui.yaml --language zh
```
Upon starting, you will find a "ChatBot" tab within the Gradio application, which serves as a chat-based interface to handle any requests related to image editing or generation.
[//]: # ( <th align="center" colspan="3">Travel Style</th>)
#### ACE ComfyUI Workflow
[//]: # ( <th align="center" colspan="3">WuKong Style</th>)
![Workflow](https://github.com/ali-vilab/ace-page/raw/main/assets/comfyui/ace_example.jpg)
[//]: # ( </tr>)
<table><tbody>
<tr>
<th align="center" colspan="4">ACE Workflow Examples</th>
</tr>
<tr>
<th align="center" colspan="1">Control</th>
<th align="center" colspan="1">Semantic</th>
<th align="center" colspan="1">Element</th>
</tr>
<tr>
<td>
<a href="https://github.com/ali-vilab/ace-page/raw/main/assets/comfyui/ace_control.png" target="_blank">
<img src="https://github.com/ali-vilab/ace-page/raw/main/assets/comfyui/ace_control.png" width="200">
</a>
</td>
<td>
<a href="https://github.com/ali-vilab/ace-page/raw/main/assets/comfyui/ace_semantic.png" target="_blank">
<img src="https://github.com/ali-vilab/ace-page/raw/main/assets/comfyui/ace_semantic.png" width="200">
</a>
</td>
<td>
<a href="https://github.com/ali-vilab/ace-page/raw/main/assets/comfyui/ace_element.png" target="_blank">
<img src="https://github.com/ali-vilab/ace-page/raw/main/assets/comfyui/ace_element.png" width="200">
</a>
</td>
</tr>
</tbody>
</table>
[//]: # ( <tr>)
### FLUX Tuners
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_3_1.webp" width="200"></td>)
<table><tbody>
<tr>
<th align="center" colspan="3">Yarn Style</th>
<th align="center" colspan="3">Soft Watercolor Style</th>
</tr>
<tr>
<td><img src="asset/images/flux_tuner/flux_tuner_2_1.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_2_2.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_2_3.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_1_1.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_1_2.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_1_3.webp" width="200"></td>
</tr>
<tr>
<th align="center" colspan="3">Travel Style</th>
<th align="center" colspan="3">WuKong Style</th>
</tr>
<tr>
<td><img src="asset/images/flux_tuner/flux_tuner_3_1.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_3_2.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_3_3.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_4_1.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_4_2.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_4_3.webp" width="200"></td>
</tr>
</tbody>
</table>
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_3_2.webp" width="200"></td>)
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_3_3.webp" width="200"></td>)
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_4_1.webp" width="200"></td>)
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_4_2.webp" width="200"></td>)
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_4_3.webp" width="200"></td>)
[//]: # ( </tr>)
[//]: # (</tbody>)
[//]: # (</table>)
### ComfyUI Workflow
@@ -180,13 +135,6 @@ Upon starting, you will find a "ChatBot" tab within the Gradio application, whic
## 🛠️ Installation
- Create new environment with `conda` command:
```shell
conda env create -f environment.yaml
conda activate scepter
```
- Install with `pip` command:
We recommend installing the specific version of PyTorch and accelerate toolbox [xFormers](https://pypi.org/project/xformers/). You can install these recommended version by pip:
@@ -210,18 +158,19 @@ pip install scepter
### Currently supported approaches
| Tasks | Methods | Links |
|:----------------------------:|:----------------------------------------------:|:------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
| Text-to-image Generation | SD v1.5 | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
| Text-to-image Generation | SD v2.1 | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
| Text-to-image Generation | SD-XL | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0) |
| Text-to-image Generation | FLUX | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/black-forest-labs/FLUX.1-dev) |
| Efficient Tuning | LoRA | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=LoRA&color=red&logo=arxiv)](https://arxiv.org/abs/2106.09685) |
| Efficient Tuning | Res-Tuning(NeurIPS23) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=Res-Tuing&color=red&logo=arxiv)](https://arxiv.org/abs/2310.19859) [![Page link](https://img.shields.io/badge/Page-ResTuning-Gree)](https://res-tuning.github.io/) |
| Controllable Image Synthesis | [🌟SCEdit(CVPR24)](docs/en/tasks/scedit.md) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=SCEdit&color=red&logo=arxiv)](https://arxiv.org/abs/2312.11392) [![Page link](https://img.shields.io/badge/Page-SCEdit-Gree)](https://scedit.github.io/) |
| Image Editing | [🌟LAR-Gen](docs/en/tasks/largen.md) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=LARGen&color=red&logo=arxiv)](https://arxiv.org/abs/2403.19534) [![Page link](https://img.shields.io/badge/Page-LARGen-Gree)](https://ali-vilab.github.io/largen-page/) |
| Image Editing | [🌟StyleBooth](docs/en/tasks/stylebooth.md) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=StyleBooth&color=red&logo=arxiv)](https://arxiv.org/abs/2404.12154) [![Page link](https://img.shields.io/badge/Page-StyleBooth-Gree)](https://ali-vilab.github.io/stylebooth-page/) |
| Image Generation and Editing | [🌟ACE](https://ali-vilab.github.io/ace-page/) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=ACE&color=red&logo=arxiv)](https://arxiv.org/abs/2410.00086) [![Page link](https://img.shields.io/badge/Page-ACE-Gree)](https://ali-vilab.github.io/ace-page/) [![Demo link](https://img.shields.io/badge/Demo-ACE-purple)](https://huggingface.co/spaces/scepter-studio/ACE-Chat) <br> [![ModelScope link](https://img.shields.io/badge/ModelScope-Model-blue)](https://www.modelscope.cn/models/iic/ACE-0.6B-512px) [![HuggingFace link](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Model-yellow)](https://huggingface.co/scepter-studio/ACE-0.6B-512px) |
| Tasks | Methods | Links |
|:----------------------------:|:------------------------------------------------:|:------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
| Text-to-image Generation | SD v1.5 | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
| Text-to-image Generation | SD v2.1 | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
| Text-to-image Generation | SD-XL | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0) |
| Text-to-image Generation | FLUX | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/black-forest-labs/FLUX.1-dev) |
| Efficient Tuning | LoRA | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=LoRA&color=red&logo=arxiv)](https://arxiv.org/abs/2106.09685) |
| Efficient Tuning | Res-Tuning(NeurIPS23) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=Res-Tuing&color=red&logo=arxiv)](https://arxiv.org/abs/2310.19859) [![Page link](https://img.shields.io/badge/Page-ResTuning-Gree)](https://res-tuning.github.io/) |
| Controllable Image Synthesis | [🌟SCEdit(CVPR24)](docs/en/tasks/scedit.md) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=SCEdit&color=red&logo=arxiv)](https://arxiv.org/abs/2312.11392) [![Page link](https://img.shields.io/badge/Page-SCEdit-Gree)](https://scedit.github.io/) |
| Image Editing | [🌟LAR-Gen](docs/en/tasks/largen.md) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=LARGen&color=red&logo=arxiv)](https://arxiv.org/abs/2403.19534) [![Page link](https://img.shields.io/badge/Page-LARGen-Gree)](https://ali-vilab.github.io/largen-page/) |
| Image Editing | [🌟StyleBooth](docs/en/tasks/stylebooth.md) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=StyleBooth&color=red&logo=arxiv)](https://arxiv.org/abs/2404.12154) [![Page link](https://img.shields.io/badge/Page-StyleBooth-Gree)](https://ali-vilab.github.io/stylebooth-page/) |
| Image Generation and Editing | [🌟ACE](https://ali-vilab.github.io/ace-page/) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=ACE&color=red&logo=arxiv)](https://arxiv.org/abs/2410.00086) [![Page link](https://img.shields.io/badge/Page-ACE-Gree)](https://ali-vilab.github.io/ace-page/) [![Demo link](https://img.shields.io/badge/Demo-ACE-purple)](https://huggingface.co/spaces/scepter-studio/ACE-Chat) <br> [![ModelScope link](https://img.shields.io/badge/ModelScope-Model-blue)](https://www.modelscope.cn/models/iic/ACE-0.6B-512px) [![HuggingFace link](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Model-yellow)](https://huggingface.co/scepter-studio/ACE-0.6B-512px) |
| Image Generation and Editing | [🌟ACE++](https://ali-vilab.github.io/ACE_plus_page/) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=ACEPlus&color=red&logo=arxiv)](https://arxiv.org/abs/2501.02487) [![Page link](https://img.shields.io/badge/Page-ACE++-Gree)](https://ali-vilab.github.io/ACE_plus_page/) [![Demo link](https://img.shields.io/badge/Demo-ACE++-purple)](https://huggingface.co/spaces/scepter-studio/ACE-Plus) <br> [![ModelScope link](https://img.shields.io/badge/ModelScope-Model-blue)](https://www.modelscope.cn/models/iic/ACE_Plus/summary) [![HuggingFace link](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Model-yellow)](https://huggingface.co/ali-vilab/ACE_Plus/tree/main) |
## 🖥️ SCEPTER Studio
@@ -258,18 +207,20 @@ We deploy a work studio on Modelscope that includes only the inference tab, plea
## ⚙️️ ComfyUI Workflow
### Launch
We support the use of all models in the ComfyUI Workflow through the following methods:
Manually install by moving custom_nodes to ComfyUI.
1) Automatic installation directly via the ComfyUI Manager by searching for the **ComfyUI-Scepter** node.
2) Manually install by moving custom_nodes from Scepter to ComfyUI.
```shell
git clone https://github.com/modelscope/scepter.git
cd path/to/scepter
pip install -e .
cp -r path/to/scepter/workflow/ path/to/ComfyUI/custom_nodes/ComfyUI-Scepter
cd path/to/ComfyUI
python main.py
```
In addition, we also support installation and usage through the ComfyUI Manager.
**Note**: You can use the nodes by dragging the sample images into ComfyUI. Additionally, our nodes can automatically pull models from ModelScope or HuggingFace by selecting the *model_source* field, or you can place the already downloaded models in a local path.
## 🔍 Learn More
+3 -3
View File
@@ -1,8 +1,7 @@
albumentations
beautifulsoup4
bezier
einops
modelscope
modelscope[framework]
ms-swift
numpy
open_clip_torch
@@ -12,6 +11,7 @@ oss2>=2.15.0
pycocotools
pyyaml>=5.3.1
scikit-image
scikit-learn
sentencepiece
torchsde
transformers
scikit-learn
+4 -3
View File
@@ -1,4 +1,5 @@
git+https://github.com/cocodataset/panopticapi.git
torch==2.0.1
torchvision==0.15.2
xformers==0.0.21
torch==2.4.1
torchvision==0.19.1
flash-attn==2.5.8
xformers==0.0.28
+1 -1
View File
@@ -1,5 +1,5 @@
bitsandbytes
gradio==4.44.1
gradio
gradio_imageslider
imagehash
psutil
+23 -13
View File
@@ -1,18 +1,28 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os
from typing import TYPE_CHECKING
from scepter.modules.utils.import_utils import LazyImportModule
import scepter
from scepter.modules import data, model, opt, solver, transform, utils
from scepter.tools.helper import get_module_list as module_list
from scepter.tools.helper import \
get_module_object_config as configures_by_objects
from scepter.tools.helper import get_module_objects as objects_by_module
from scepter.version import __version__, version_info
dirname = os.path.dirname(scepter.__file__)
if TYPE_CHECKING:
from scepter.modules import data, model, opt, solver, transform, utils
from scepter.tools.helper import get_module_list as module_list
from scepter.tools.helper import \
get_module_object_config as configures_by_objects
from scepter.tools.helper import get_module_objects as objects_by_module
from scepter.version import __version__, version_info
else:
_import_structure = {
'modules': ['data', 'model', 'opt', 'solver', 'transform', 'utils'],
'helper': ['get_module_list', 'get_module_object_config', 'get_module_objects'],
'version': ['__version__', 'version_info']
}
__all__ = [
utils, transform, data, model, solver, version_info, opt, '__version__',
'dirname'
]
import sys
sys.modules[__name__] = LazyImportModule(
__name__,
globals()['__file__'],
_import_structure,
module_spec=__spec__,
extra_objects={},
)
+161
View File
@@ -0,0 +1,161 @@
ENV:
BACKEND: nccl
SEED: 2024
#
SOLVER:
NAME: ACESolver
RESUME_FROM:
LOAD_MODEL_ONLY: True
USE_FSDP: False
SHARDING_STRATEGY:
USE_AMP: True
DTYPE: float16
CHANNELS_LAST: True
MAX_STEPS: 500
MAX_EPOCHS: -1
NUM_FOLDS: 1
ACCU_STEP: 1
EVAL_INTERVAL: 50
RESCALE_LR: False
#
WORK_DIR: ./cache/save_data/ace_0.6b_1024
LOG_FILE: std_log.txt
#
FILE_SYSTEM:
- NAME: "HuggingfaceFs"
TEMP_DIR: ./cache/cache_data
- NAME: "LocalFs"
TEMP_DIR: ./cache/cache_data
- NAME: "ModelscopeFs"
TEMP_DIR: ./cache/cache_data
#
MODEL:
NAME: LatentDiffusionACE
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
DECODER_BIAS: 0.5
DEFAULT_N_PROMPT:
USE_EMA: True
EVAL_EMA: False
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
USE_TEXT_POS_EMBEDDINGS: True
#
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: eps
MIN_SNR_GAMMA:
NOISE_SCHEDULER:
NAME: LinearScheduler
NUM_TIMESTEPS: 1000
BETA_MIN: 0.0001
BETA_MAX: 0.02
#
DIFFUSION_MODEL:
NAME: ACE
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/dit/ace_0.6b_1024px.pth
IGNORE_KEYS: [ ]
PATCH_SIZE: 2
IN_CHANNELS: 4
HIDDEN_SIZE: 1152
DEPTH: 28
NUM_HEADS: 16
MLP_RATIO: 4.0
PRED_SIGMA: True
DROP_PATH: 0.0
WINDOW_DIZE: 0
Y_CHANNELS: 4096
MAX_SEQ_LEN: 4096
QK_NORM: True
USE_GRAD_CHECKPOINT: True
ATTENTION_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/vae/vae.bin
IGNORE_KEYS: []
#
ENCODER:
NAME: Encoder
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 4
DOUBLE_Z: True
DROPOUT: 0.0
RESAMP_WITH_CONV: True
#
DECODER:
NAME: Decoder
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 4
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/text_encoder/t5-v1_1-xxl/
TOKENIZER_PATH: ms://iic/ACE-0.6B-1024px@models/tokenizer/t5-v1_1-xxl
LENGTH: 120
T5_DTYPE: bfloat16
ADDED_IDENTIFIER: [ '{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
CLEAN: whitespace
USE_GRAD: False
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 20
GUIDE_SCALE: 4.5
GUIDE_RESCALE: 0.5
#
OPTIMIZER:
NAME: AdamW
LEARNING_RATE: 1e-7
EPS: 1e-10
WEIGHT_DECAY: 5e-4
#
TRAIN_DATA:
NAME: ImageTextPairMSDatasetForACE
MODE: train
MS_DATASET_NAME: cache/datasets/hed_pair
MS_DATASET_NAMESPACE: ""
MS_DATASET_SPLIT: "train"
MS_DATASET_SUBNAME: ""
PROMPT_PREFIX: ""
REPLACE_STYLE: False
MAX_SEQ_LEN: 4096
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 1
SAMPLER:
NAME: LoopSampler
#
TRAIN_HOOKS:
-
NAME: BackwardHook
PRIORITY: 0
-
NAME: LogHook
LOG_INTERVAL: 50
-
NAME: CheckpointHook
INTERVAL: 100
-
NAME: ProbeDataHook
PROB_INTERVAL: 100
@@ -0,0 +1,277 @@
ENV:
BACKEND: nccl
SEED: 42
TENSOR_PARALLEL_SIZE: 1
PIPELINE_PARALLEL_SIZE: 1
SYS_ENVS:
TORCH_CUDNN_V8_API_ENABLED: '1'
TOKENIZERS_PARALLELISM: 'false'
TF_CPP_MIN_LOG_LEVEL: '3'
PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True'
#
SOLVER:
NAME: LatentDiffusionVideoSolver
MAX_STEPS: 2000
USE_AMP: True
DTYPE: bfloat16
USE_FAIRSCALE: False
USE_FSDP: True
LOAD_MODEL_ONLY: False
ENABLE_GRADSCALER: False
USE_SCALER: False
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_cogvideox1.5_5b_i2v_lora
LOG_FILE: std_log.txt
EVAL_INTERVAL: 100
LOG_TRAIN_NUM: 4
FPS: 16
SHARDING_STRATEGY: full_shard
FSDP_REDUCE_DTYPE: float32
FSDP_BUFFER_DTYPE: float32
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model']
SAVE_MODULES: [ 'model', 'cond_stage_model.model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
TUNER:
- NAME: SwiftLoRA
R: 64
LORA_ALPHA: 64
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$"
#
MODEL:
NAME: LatentDiffusionCogVideoX
PRETRAINED_MODEL:
PARAMETERIZATION: v
TIMESTEPS: 1000
MIN_SNR_GAMMA: 3.0
ZERO_TERMINAL_SNR: True
SCALE_FACTOR_SPATIAL: 8
SCALE_FACTOR_TEMPORAL: 4
SCALING_FACTOR_IMAGE: 0.7
NOISED_IMAGE_DROPOUT: 0.05
INVERT_SCALE_LATENTS: True
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: v
USE_DYNAMIC_CFG: False
NOISE_SCHEDULER:
NAME: ScaledLinearScheduler
BETA_MIN: 0.00085
BETA_MAX: 0.012
SNR_SHIFT_SCALE: 1.0
RESCALE_BETAS_ZERO_SNR: True
DIFFUSION_SAMPLERS:
NAME: DDIMSampler
DISCRETIZATION_TYPE: trailing
ETA: 0.0
#
DIFFUSION_MODEL:
NAME: CogVideoXTransformer3DModel
DTYPE: bfloat16
PRETRAINED_MODEL: # 5b-I2V diff
- ms://ZhipuAI/CogVideoX1.5-5B-I2V@transformer/diffusion_pytorch_model-00001-of-00003.safetensors
- ms://ZhipuAI/CogVideoX1.5-5B-I2V@transformer/diffusion_pytorch_model-00002-of-00003.safetensors
- ms://ZhipuAI/CogVideoX1.5-5B-I2V@transformer/diffusion_pytorch_model-00003-of-00003.safetensors
NUM_ATTENTION_HEADS: 48
ATTENTION_HEAD_DIM: 64
IN_CHANNELS: 32
LATENT_CHANNELS: 16
OUT_CHANNELS: 16
FLIP_SIN_TO_COS: True
FREQ_SHIFT: 0
TIME_EMBED_DIM: 512
TEXT_EMBED_DIM: 4096
OFS_EMBED_DIM: 512 # v1.5 diff
NUM_LAYERS: 42
DROPOUT: 0.0
ATTENTION_BIAS: True
SAMPLE_WIDTH: 300
SAMPLE_HEIGHT: 300
SAMPLE_FRAMES: 81
PATCH_SIZE: 2
PATCH_SIZE_T: 2 # v1.5 diff
PATCH_BIAS: False # v1.5 diff
TEMPORAL_COMPRESSION_RATIO: 4
MAX_TEXT_SEQ_LENGTH: 224
ACTIVATION_FN: "gelu-approximate"
TIMESTEP_ACTIVATION_FN: "silu"
NORM_ELEMENTWISE_AFFINE: True
NORM_EPS: 1e-5
SPATIAL_INTERPOLATION_SCALE: 1.875
TEMPORAL_INTERPOLATION_SCALE: 1.0
USE_ROTARY_POSITIONAL_EMBEDDINGS: True
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
GRADIENT_CHECKPOINTING: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://ZhipuAI/CogVideoX1.5-5B-I2V@vae/diffusion_pytorch_model.safetensors
SAMPLE_HEIGHT: 768
SAMPLE_WIDTH: 1360
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: True
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
LENGTH: 224
CLEAN:
USE_GRAD: False
T5_DTYPE: bfloat16
#
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 42
GUIDE_SCALE: 6.0
GUIDE_RESCALE: 0.0
NUM_FRAMES: 81
IMAGE_SIZE: [768, 1360]
#
OPTIMIZER:
NAME: Adam
LEARNING_RATE: 1e-3
BETAS: [ 0.9, 0.95 ]
EPS: 1e-8
WEIGHT_DECAY: 0.0
AMSGRAD: False
#
# LR_SCHEDULER:
# NAME: StepAnnealingLR
# WARMUP_STEPS: 200
# TOTAL_STEPS: 2000
# DECAY_MODE: 'cosine'
#
TRAIN_DATA:
NAME: VideoGenDataset
MODE: train
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 0
NUM_FRAMES: 85
FPS: 16
HEIGHT: 768
WIDTH: 1360
PROMPT_PREFIX: 'DISNEY '
DATA_TYPE: 'i2v'
SAMPLER:
NAME: MixtureOfSamplers
SUB_SAMPLERS:
- NAME: MultiLevelBatchSampler
PROB: 1.0
FIELDS: [ "video_path", "prompt" ]
DELIMITER: '#;#'
PATH_PREFIX: cache/datasets/Disney-VideoGeneration-Dataset/
INDEX_FILE: cache/datasets/Disney-VideoGeneration-Dataset/index.jsonl
TRANSFORMS:
- NAME: Select
KEYS: [ "video", "image", "prompt" ]
META_KEYS: [ ]
#
# EVAL_DATA:
# NAME: Text2ImageDataset
# MODE: eval
# PROMPT_FILE:
# PROMPT_DATA: [ "A cat running.#;#asset/images/edit_tuner/cat_512.jpg" ]
# FIELDS: [ "prompt", "img_path" ]
# DELIMITER: '#;#'
# PROMPT_PREFIX: ''
# PIN_MEMORY: True
# BATCH_SIZE: 1
# USE_NUM: 8
# NUM_WORKERS: 0
# IMAGE_SIZE: [768, 1360]
# TRANSFORMS:
# - NAME: LoadImageFromFileList
# FILE_KEYS: [ 'img_path' ]
# RGB_ORDER: RGB
# BACKEND: pillow
# - NAME: FlexibleResize
# INTERPOLATION: bilinear
# SIZE: [768, 1360]
# INPUT_KEY: [ 'img' ]
# OUTPUT_KEY: [ 'img' ]
# BACKEND: pillow
# - NAME: FlexibleCenterCrop
# SIZE: [768, 1360]
# INPUT_KEY: [ 'img' ]
# OUTPUT_KEY: [ 'img' ]
# BACKEND: pillow
# - NAME: ImageToTensor
# INPUT_KEY: [ 'img' ]
# OUTPUT_KEY: [ 'img' ]
# BACKEND: pillow
# - NAME: Normalize
# MEAN: [ 0.5, 0.5, 0.5 ]
# STD: [ 0.5, 0.5, 0.5 ]
# INPUT_KEY: [ 'img' ]
# OUTPUT_KEY: [ 'image' ]
# BACKEND: torchvision
# - NAME: Select
# KEYS: [ 'image', 'prompt' ]
# META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
PRIORITY: 20
- NAME: CheckpointHook
INTERVAL: 1000
PRIORITY: 40
#
# EVAL_HOOKS:
# - NAME: ProbeDataHook
# PROB_INTERVAL: 100
# PRIORITY: 0
@@ -0,0 +1,248 @@
ENV:
BACKEND: nccl
SEED: 42
TENSOR_PARALLEL_SIZE: 1
PIPELINE_PARALLEL_SIZE: 1
SYS_ENVS:
TORCH_CUDNN_V8_API_ENABLED: '1'
TOKENIZERS_PARALLELISM: 'false'
TF_CPP_MIN_LOG_LEVEL: '3'
PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True'
#
SOLVER:
NAME: LatentDiffusionVideoSolver
MAX_STEPS: 2000
USE_AMP: True
DTYPE: bfloat16
USE_FAIRSCALE: False
USE_FSDP: True
LOAD_MODEL_ONLY: False
ENABLE_GRADSCALER: False
USE_SCALER: False
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_cogvideox1.5_5b_lora
LOG_FILE: std_log.txt
EVAL_INTERVAL: 100
LOG_TRAIN_NUM: 4
FPS: 16
SHARDING_STRATEGY: full_shard
FSDP_REDUCE_DTYPE: float32
FSDP_BUFFER_DTYPE: float32
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model']
SAVE_MODULES: [ 'model', 'cond_stage_model.model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
TUNER:
- NAME: SwiftLoRA
R: 64
LORA_ALPHA: 64
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$"
#
MODEL:
NAME: LatentDiffusionCogVideoX
PRETRAINED_MODEL:
PARAMETERIZATION: v
TIMESTEPS: 1000
MIN_SNR_GAMMA: 3.0
ZERO_TERMINAL_SNR: True
SCALE_FACTOR_SPATIAL: 8
SCALE_FACTOR_TEMPORAL: 4
SCALING_FACTOR_IMAGE: 0.7
INVERT_SCALE_LATENTS: True
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: v
USE_DYNAMIC_CFG: False
NOISE_SCHEDULER:
NAME: ScaledLinearScheduler
BETA_MIN: 0.00085
BETA_MAX: 0.012
SNR_SHIFT_SCALE: 1.0
RESCALE_BETAS_ZERO_SNR: True
DIFFUSION_SAMPLERS:
NAME: DDIMSampler
DISCRETIZATION_TYPE: trailing
ETA: 0.0
#
DIFFUSION_MODEL:
NAME: CogVideoXTransformer3DModel
DTYPE: bfloat16
PRETRAINED_MODEL:
- ms://ZhipuAI/CogVideoX1.5-5B@transformer/diffusion_pytorch_model-00001-of-00003.safetensors
- ms://ZhipuAI/CogVideoX1.5-5B@transformer/diffusion_pytorch_model-00002-of-00003.safetensors
- ms://ZhipuAI/CogVideoX1.5-5B@transformer/diffusion_pytorch_model-00003-of-00003.safetensors
NUM_ATTENTION_HEADS: 48
ATTENTION_HEAD_DIM: 64
IN_CHANNELS: 16
OUT_CHANNELS: 16
FLIP_SIN_TO_COS: True
FREQ_SHIFT: 0
TIME_EMBED_DIM: 512
TEXT_EMBED_DIM: 4096
NUM_LAYERS: 42
DROPOUT: 0.0
ATTENTION_BIAS: True
SAMPLE_WIDTH: 300
SAMPLE_HEIGHT: 300
SAMPLE_FRAMES: 81
PATCH_SIZE: 2
PATCH_SIZE_T: 2 # v1.5 diff
PATCH_BIAS: False # v1.5 diff
TEMPORAL_COMPRESSION_RATIO: 4
MAX_TEXT_SEQ_LENGTH: 224
ACTIVATION_FN: "gelu-approximate"
TIMESTEP_ACTIVATION_FN: "silu"
NORM_ELEMENTWISE_AFFINE: True
NORM_EPS: 1e-5
SPATIAL_INTERPOLATION_SCALE: 1.875
TEMPORAL_INTERPOLATION_SCALE: 1.0
USE_ROTARY_POSITIONAL_EMBEDDINGS: True
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
GRADIENT_CHECKPOINTING: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://ZhipuAI/CogVideoX1.5-5B@vae/diffusion_pytorch_model.safetensors
SAMPLE_HEIGHT: 768
SAMPLE_WIDTH: 1360
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: True
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
LENGTH: 224
CLEAN:
USE_GRAD: False
T5_DTYPE: bfloat16
#
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 42
GUIDE_SCALE: 6.0
GUIDE_RESCALE: 0.0
NUM_FRAMES: 81
IMAGE_SIZE: [768, 1360]
#
OPTIMIZER:
NAME: Adam
LEARNING_RATE: 1e-3
BETAS: [ 0.9, 0.95 ]
EPS: 1e-8
WEIGHT_DECAY: 0.0
AMSGRAD: False
#
# LR_SCHEDULER:
# NAME: StepAnnealingLR
# WARMUP_STEPS: 200
# TOTAL_STEPS: 2000
# DECAY_MODE: 'cosine'
#
TRAIN_DATA:
NAME: VideoGenDataset
MODE: train
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 0
NUM_FRAMES: 85
FPS: 16
HEIGHT: 768
WIDTH: 1360
PROMPT_PREFIX: 'DISNEY '
SAMPLER:
NAME: MixtureOfSamplers
SUB_SAMPLERS:
- NAME: MultiLevelBatchSampler
PROB: 1.0
FIELDS: [ "video_path", "prompt" ]
DELIMITER: '#;#'
PATH_PREFIX: cache/datasets/Disney-VideoGeneration-Dataset/
INDEX_FILE: cache/datasets/Disney-VideoGeneration-Dataset/index.jsonl
TRANSFORMS:
- NAME: Select
KEYS: [ 'video', "prompt" ]
META_KEYS: [ ]
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
PROMPT_FILE:
PROMPT_DATA: [ "A girl riding a bike." ]
IMAGE_SIZE: [ 768, 1360 ]
FIELDS: [ "prompt" ]
DELIMITER: '#;#'
PROMPT_PREFIX: 'DISNEY ' # ''
PIN_MEMORY: True
BATCH_SIZE: 1
USE_NUM: 8
NUM_WORKERS: 0
TRANSFORMS:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
PRIORITY: 20
- NAME: CheckpointHook
INTERVAL: 1000
PRIORITY: 40
#
EVAL_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
@@ -0,0 +1,239 @@
ENV:
BACKEND: nccl
SEED: 42
TENSOR_PARALLEL_SIZE: 1
PIPELINE_PARALLEL_SIZE: 1
SYS_ENVS:
TORCH_CUDNN_V8_API_ENABLED: '1'
TOKENIZERS_PARALLELISM: 'false'
TF_CPP_MIN_LOG_LEVEL: '3'
PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True'
#
SOLVER:
NAME: LatentDiffusionVideoSolver
MAX_STEPS: 2000
USE_AMP: True
DTYPE: bfloat16
USE_FAIRSCALE: False
USE_FSDP: True
LOAD_MODEL_ONLY: False
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_cogvideox_2b_lora
LOG_FILE: std_log.txt
EVAL_INTERVAL: 100
LOG_TRAIN_NUM: 4
ENABLE_GRADSCALER: False
USE_SCALER: False
FPS: 8
SHARDING_STRATEGY: full_shard
FSDP_REDUCE_DTYPE: float32
FSDP_BUFFER_DTYPE: float32
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model']
SAVE_MODULES: [ 'model', 'cond_stage_model.model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
TUNER:
- NAME: SwiftLoRA
R: 64
LORA_ALPHA: 64
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$"
#
MODEL:
NAME: LatentDiffusionCogVideoX
PRETRAINED_MODEL:
PARAMETERIZATION: v
TIMESTEPS: 1000
MIN_SNR_GAMMA: 3.0
ZERO_TERMINAL_SNR: True
SCALE_FACTOR_SPATIAL: 8
SCALE_FACTOR_TEMPORAL: 4
SCALING_FACTOR_IMAGE: 1.15258426
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: v
NOISE_SCHEDULER:
NAME: ScaledLinearScheduler
BETA_MIN: 0.00085
BETA_MAX: 0.012
SNR_SHIFT_SCALE: 3.0
RESCALE_BETAS_ZERO_SNR: True
DIFFUSION_SAMPLERS:
NAME: DDIMSampler
DISCRETIZATION_TYPE: trailing
ETA: 0.0
#
DIFFUSION_MODEL:
NAME: CogVideoXTransformer3DModel
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@transformer/diffusion_pytorch_model.safetensors
NUM_ATTENTION_HEADS: 30
ATTENTION_HEAD_DIM: 64
IN_CHANNELS: 16
OUT_CHANNELS: 16
FLIP_SIN_TO_COS: True
FREQ_SHIFT: 0
TIME_EMBED_DIM: 512
TEXT_EMBED_DIM: 4096
NUM_LAYERS: 30
DROPOUT: 0.0
ATTENTION_BIAS: True
SAMPLE_WIDTH: 90
SAMPLE_HEIGHT: 60
SAMPLE_FRAMES: 49
PATCH_SIZE: 2
TEMPORAL_COMPRESSION_RATIO: 4
MAX_TEXT_SEQ_LENGTH: 226
ACTIVATION_FN: "gelu-approximate"
TIMESTEP_ACTIVATION_FN: "silu"
NORM_ELEMENTWISE_AFFINE: True
NORM_EPS: 1e-5
SPATIAL_INTERPOLATION_SCALE: 1.875
TEMPORAL_INTERPOLATION_SCALE: 1.0
USE_ROTARY_POSITIONAL_EMBEDDINGS: False
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
GRADIENT_CHECKPOINTING: False
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@vae/diffusion_pytorch_model.safetensors
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: False
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: False
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
LENGTH: 226
CLEAN:
USE_GRAD: False
#
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 42
GUIDE_SCALE: 6.0
GUIDE_RESCALE: 0.0
NUM_FRAMES: 49
#
OPTIMIZER:
NAME: Adam
LEARNING_RATE: 1e-3
BETAS: [ 0.9, 0.95 ]
EPS: 1e-8
WEIGHT_DECAY: 0.0
AMSGRAD: False
#
# LR_SCHEDULER:
# NAME: StepAnnealingLR
# WARMUP_STEPS: 200
# TOTAL_STEPS: 2000
# DECAY_MODE: 'cosine'
#
TRAIN_DATA:
NAME: VideoGenDataset
MODE: train
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
NUM_FRAMES: 49
FPS: 8
HEIGHT: 480
WIDTH: 720
PROMPT_PREFIX: 'DISNEY '
SAMPLER:
NAME: MixtureOfSamplers
SUB_SAMPLERS:
- NAME: MultiLevelBatchSampler
PROB: 1.0
FIELDS: [ "video_path", "prompt" ]
DELIMITER: '#;#'
PATH_PREFIX: cache/datasets/Disney-VideoGeneration-Dataset/
INDEX_FILE: cache/datasets/Disney-VideoGeneration-Dataset/index.jsonl
TRANSFORMS:
- NAME: Select
KEYS: [ 'video', "prompt" ]
META_KEYS: [ ]
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
PROMPT_FILE:
PROMPT_DATA: [ "A girl riding a bike.", "A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance." ]
IMAGE_SIZE: [ 480, 720 ]
FIELDS: [ "prompt" ]
DELIMITER: '#;#'
PROMPT_PREFIX: 'DISNEY '
PIN_MEMORY: True
BATCH_SIZE: 1
USE_NUM: 8
NUM_WORKERS: 4
TRANSFORMS:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
PRIORITY: 20
- NAME: CheckpointHook
INTERVAL: 1000
PRIORITY: 40
#
EVAL_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
@@ -0,0 +1,270 @@
ENV:
BACKEND: nccl
SEED: 42
TENSOR_PARALLEL_SIZE: 1
PIPELINE_PARALLEL_SIZE: 1
SYS_ENVS:
TORCH_CUDNN_V8_API_ENABLED: '1'
TOKENIZERS_PARALLELISM: 'false'
TF_CPP_MIN_LOG_LEVEL: '3'
PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True'
#
SOLVER:
NAME: LatentDiffusionVideoSolver
MAX_STEPS: 2000
USE_AMP: True
DTYPE: bfloat16
USE_FAIRSCALE: False
USE_FSDP: True
LOAD_MODEL_ONLY: False
ENABLE_GRADSCALER: False
USE_SCALER: False
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_cogvideox_5b_i2v_lora
LOG_FILE: std_log.txt
EVAL_INTERVAL: 100
LOG_TRAIN_NUM: 4
FPS: 8
SHARDING_STRATEGY: full_shard
FSDP_REDUCE_DTYPE: float32
FSDP_BUFFER_DTYPE: float32
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model']
SAVE_MODULES: [ 'model', 'cond_stage_model.model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
TUNER:
- NAME: SwiftLoRA
R: 64
LORA_ALPHA: 64
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$"
#
MODEL:
NAME: LatentDiffusionCogVideoX
PRETRAINED_MODEL:
PARAMETERIZATION: v
TIMESTEPS: 1000
MIN_SNR_GAMMA: 3.0
ZERO_TERMINAL_SNR: True
SCALE_FACTOR_SPATIAL: 8
SCALE_FACTOR_TEMPORAL: 4
SCALING_FACTOR_IMAGE: 0.7 # 5b diff
NOISED_IMAGE_DROPOUT: 0.05
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: v
NOISE_SCHEDULER:
NAME: ScaledLinearScheduler
BETA_MIN: 0.00085
BETA_MAX: 0.012
SNR_SHIFT_SCALE: 1.0 # 5b diff
RESCALE_BETAS_ZERO_SNR: True
DIFFUSION_SAMPLERS:
NAME: DDIMSampler
DISCRETIZATION_TYPE: trailing
ETA: 0.0
#
DIFFUSION_MODEL:
NAME: CogVideoXTransformer3DModel
DTYPE: bfloat16
PRETRAINED_MODEL: # 5b-I2V diff
- ms://AI-ModelScope/CogVideoX-5b-I2V@transformer/diffusion_pytorch_model-00001-of-00003.safetensors
- ms://AI-ModelScope/CogVideoX-5b-I2V@transformer/diffusion_pytorch_model-00002-of-00003.safetensors
- ms://AI-ModelScope/CogVideoX-5b-I2V@transformer/diffusion_pytorch_model-00003-of-00003.safetensors
NUM_ATTENTION_HEADS: 48 # 5b diff
ATTENTION_HEAD_DIM: 64
IN_CHANNELS: 32 # 5b-I2V diff
LATENT_CHANNELS: 16
OUT_CHANNELS: 16
FLIP_SIN_TO_COS: True
FREQ_SHIFT: 0
TIME_EMBED_DIM: 512
TEXT_EMBED_DIM: 4096
NUM_LAYERS: 42 # 5b diff
DROPOUT: 0.0
ATTENTION_BIAS: True
SAMPLE_WIDTH: 90
SAMPLE_HEIGHT: 60
SAMPLE_FRAMES: 49
PATCH_SIZE: 2
TEMPORAL_COMPRESSION_RATIO: 4
MAX_TEXT_SEQ_LENGTH: 226
ACTIVATION_FN: "gelu-approximate"
TIMESTEP_ACTIVATION_FN: "silu"
NORM_ELEMENTWISE_AFFINE: True
NORM_EPS: 1e-5
SPATIAL_INTERPOLATION_SCALE: 1.875
TEMPORAL_INTERPOLATION_SCALE: 1.0
USE_ROTARY_POSITIONAL_EMBEDDINGS: True # 5b diff
USE_LEARNED_POSITIONAL_EMBEDDINGS: True # 5b-I2V diff
GRADIENT_CHECKPOINTING: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b-I2V@vae/diffusion_pytorch_model.safetensors # 5b diff
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: True
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
LENGTH: 226
CLEAN:
USE_GRAD: False
#
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 42
GUIDE_SCALE: 6.0
GUIDE_RESCALE: 0.0
NUM_FRAMES: 49
#
OPTIMIZER:
NAME: Adam
LEARNING_RATE: 1e-3
BETAS: [ 0.9, 0.95 ]
EPS: 1e-8
WEIGHT_DECAY: 0.0
AMSGRAD: False
#
# LR_SCHEDULER:
# NAME: StepAnnealingLR
# WARMUP_STEPS: 200
# TOTAL_STEPS: 2000
# DECAY_MODE: 'cosine'
#
TRAIN_DATA:
NAME: VideoGenDataset
MODE: train
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 0
NUM_FRAMES: 49
FPS: 8
HEIGHT: 480
WIDTH: 720
PROMPT_PREFIX: 'DISNEY '
DATA_TYPE: 'i2v'
SAMPLER:
NAME: MixtureOfSamplers
SUB_SAMPLERS:
- NAME: MultiLevelBatchSampler
PROB: 1.0
FIELDS: [ "video_path", "prompt" ]
DELIMITER: '#;#'
PATH_PREFIX: cache/datasets/Disney-VideoGeneration-Dataset/
INDEX_FILE: cache/datasets/Disney-VideoGeneration-Dataset/index.jsonl
TRANSFORMS:
- NAME: Select
KEYS: [ "video", "image", "prompt" ]
META_KEYS: [ ]
#
# EVAL_DATA:
# NAME: Text2ImageDataset
# MODE: eval
# PROMPT_FILE:
# PROMPT_DATA: [ "A cat running.#;#asset/images/edit_tuner/cat_512.jpg" ]
# FIELDS: [ "prompt", "img_path" ]
# DELIMITER: '#;#'
# PROMPT_PREFIX: ''
# PIN_MEMORY: True
# BATCH_SIZE: 1
# USE_NUM: 8
# NUM_WORKERS: 0
# IMAGE_SIZE: [ 480, 720 ]
# TRANSFORMS:
# - NAME: LoadImageFromFileList
# FILE_KEYS: [ 'img_path' ]
# RGB_ORDER: RGB
# BACKEND: pillow
# - NAME: FlexibleResize
# INTERPOLATION: bilinear
# SIZE: [ 480, 720 ]
# INPUT_KEY: [ 'img' ]
# OUTPUT_KEY: [ 'img' ]
# BACKEND: pillow
# - NAME: FlexibleCenterCrop
# SIZE: [ 480, 720 ]
# INPUT_KEY: [ 'img' ]
# OUTPUT_KEY: [ 'img' ]
# BACKEND: pillow
# - NAME: ImageToTensor
# INPUT_KEY: [ 'img' ]
# OUTPUT_KEY: [ 'img' ]
# BACKEND: pillow
# - NAME: Normalize
# MEAN: [ 0.5, 0.5, 0.5 ]
# STD: [ 0.5, 0.5, 0.5 ]
# INPUT_KEY: [ 'img' ]
# OUTPUT_KEY: [ 'image' ]
# BACKEND: torchvision
# - NAME: Select
# KEYS: [ 'image', 'prompt' ]
# META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
PRIORITY: 20
- NAME: CheckpointHook
INTERVAL: 1000
PRIORITY: 40
#
# EVAL_HOOKS:
# - NAME: ProbeDataHook
# PROB_INTERVAL: 100
# PRIORITY: 0
@@ -0,0 +1,277 @@
ENV:
BACKEND: nccl
SEED: 42
TENSOR_PARALLEL_SIZE: 1
PIPELINE_PARALLEL_SIZE: 1
SYS_ENVS:
TORCH_CUDNN_V8_API_ENABLED: '1'
TOKENIZERS_PARALLELISM: 'false'
TF_CPP_MIN_LOG_LEVEL: '3'
PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True'
#
SOLVER:
NAME: LatentDiffusionVideoSolver
MAX_STEPS: 2000
USE_AMP: True
DTYPE: bfloat16
USE_FAIRSCALE: False
USE_FSDP: True
LOAD_MODEL_ONLY: False
ENABLE_GRADSCALER: False
USE_SCALER: False
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_cogvideox_5b_lora
LOG_FILE: std_log.txt
EVAL_INTERVAL: 100
LOG_TRAIN_NUM: 4
FPS: 8
SHARDING_STRATEGY: full_shard
FSDP_REDUCE_DTYPE: float32
FSDP_BUFFER_DTYPE: float32
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model']
SAVE_MODULES: [ 'model', 'cond_stage_model.model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
TUNER:
- NAME: SwiftLoRA
R: 64
LORA_ALPHA: 64
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$"
#
MODEL:
NAME: LatentDiffusionCogVideoX
PRETRAINED_MODEL:
PARAMETERIZATION: v
TIMESTEPS: 1000
MIN_SNR_GAMMA: 3.0
ZERO_TERMINAL_SNR: True
SCALE_FACTOR_SPATIAL: 8
SCALE_FACTOR_TEMPORAL: 4
SCALING_FACTOR_IMAGE: 0.7 # 5b diff
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: v
NOISE_SCHEDULER:
NAME: ScaledLinearScheduler
BETA_MIN: 0.00085
BETA_MAX: 0.012
SNR_SHIFT_SCALE: 1.0 # 5b diff
RESCALE_BETAS_ZERO_SNR: True
DIFFUSION_SAMPLERS:
NAME: DDIMSampler
DISCRETIZATION_TYPE: trailing
ETA: 0.0
#
DIFFUSION_MODEL:
NAME: CogVideoXTransformer3DModel
DTYPE: bfloat16
PRETRAINED_MODEL: # 5b diff
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00001-of-00002.safetensors
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00002-of-00002.safetensors
NUM_ATTENTION_HEADS: 48 # 5b diff
ATTENTION_HEAD_DIM: 64
IN_CHANNELS: 16
OUT_CHANNELS: 16
FLIP_SIN_TO_COS: True
FREQ_SHIFT: 0
TIME_EMBED_DIM: 512
TEXT_EMBED_DIM: 4096
NUM_LAYERS: 42 # 5b diff
DROPOUT: 0.0
ATTENTION_BIAS: True
SAMPLE_WIDTH: 90
SAMPLE_HEIGHT: 60
SAMPLE_FRAMES: 49
PATCH_SIZE: 2
TEMPORAL_COMPRESSION_RATIO: 4
MAX_TEXT_SEQ_LENGTH: 226
ACTIVATION_FN: "gelu-approximate"
TIMESTEP_ACTIVATION_FN: "silu"
NORM_ELEMENTWISE_AFFINE: True
NORM_EPS: 1e-5
SPATIAL_INTERPOLATION_SCALE: 1.875
TEMPORAL_INTERPOLATION_SCALE: 1.0
USE_ROTARY_POSITIONAL_EMBEDDINGS: True # 5b diff
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
GRADIENT_CHECKPOINTING: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b@vae/diffusion_pytorch_model.safetensors # 5b diff
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: True
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
LENGTH: 226
CLEAN:
USE_GRAD: False
#
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 42
GUIDE_SCALE: 6.0
GUIDE_RESCALE: 0.0
NUM_FRAMES: 49
#
OPTIMIZER:
NAME: Adam
LEARNING_RATE: 1e-3
BETAS: [ 0.9, 0.95 ]
EPS: 1e-8
WEIGHT_DECAY: 0.0
AMSGRAD: False
#
# LR_SCHEDULER:
# NAME: StepAnnealingLR
# WARMUP_STEPS: 200
# TOTAL_STEPS: 2000
# DECAY_MODE: 'cosine'
#
TRAIN_DATA:
NAME: VideoGenDatasetOTF
MODE: train
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
NUM_FRAMES: 49
FPS: 8
HEIGHT: 480
WIDTH: 720
PROMPT_PREFIX: 'DISNEY '
DELIMITER: '#;#'
FIELDS: [ 'video_path', 'prompt' ]
PATH_PREFIX: cache/datasets/Disney-VideoGeneration-Dataset/
DATA_FILE: cache/datasets/Disney-VideoGeneration-Dataset/index.jsonl
SAMPLER:
NAME: LoopSampler
TRANSFORMS:
- NAME: Select
KEYS: [ 'video', 'video_latent', "prompt" ]
META_KEYS: [ ]
MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b@vae/diffusion_pytorch_model.safetensors
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: True
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
PROMPT_FILE:
PROMPT_DATA: [ "A girl riding a bike.", "A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance." ]
IMAGE_SIZE: [ 480, 720 ]
FIELDS: [ "prompt" ]
DELIMITER: '#;#'
PROMPT_PREFIX: 'DISNEY '
PIN_MEMORY: True
BATCH_SIZE: 1
USE_NUM: 8
NUM_WORKERS: 4
TRANSFORMS:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
PRIORITY: 20
- NAME: CheckpointHook
INTERVAL: 1000
PRIORITY: 40
#
EVAL_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
@@ -2,33 +2,22 @@ ENV:
BACKEND: nccl
SEED: 166666
SOLVER:
# NAME DESCRIPTION: TYPE: default: 'LatentUfitSolver'
NAME: LatentDiffusionSolver
# MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
MAX_STEPS: 100000
# USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False
USE_AMP: True
# DTYPE DESCRIPTION: The precision for training. TYPE: str default: 'float32'
DTYPE: bfloat16
# USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
USE_FAIRSCALE: False
# USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
USE_FSDP: True
# LOAD_MODEL_ONLY DESCRIPTION: Only load the model rather than the optimizer and schedule, default is False. TYPE: bool default: False
LOAD_MODEL_ONLY: False
# RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: ''
ENABLE_GRADSCALER: False
USE_SCALER: False
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_flux_dev_1024_lora
LOG_FILE: std_log.txt
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
EVAL_INTERVAL: 100
# LOG_TRAIN_NUM DESCRIPTION: The number samples used to log in training phase. TYPE: int default: -1
LOG_TRAIN_NUM: 16
# FSDP_REDUCE_DTYPE DESCRIPTION: The dtype of reduce in FSDP. TYPE: str default: 'float16'
FSDP_REDUCE_DTYPE: float32
# FSDP_BUFFER_DTYPE DESCRIPTION: The dtype of buffer in FSDP. TYPE: str default: 'float16'
FSDP_BUFFER_DTYPE: float32
# FSDP_SHARD_MODULES DESCRIPTION: The modules to be sharded in FSDP. TYPE: list default: ['model']
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ] #
SAVE_MODULES: [ 'model'] #
TRAIN_MODULES: ['model']
@@ -58,61 +47,36 @@ SOLVER:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
NOISE_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
NAME: FlowMatchSigmaScheduler
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
WEIGHTING_SCHEME: logit_normal
SHIFT: 3.0
# LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
LOGIT_MEAN: 0.0
# LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
LOGIT_STD: 1.0
# MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
MODE_SCALE: 1.29
SAMPLER_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
NAME: FlowMatchFluxShiftScheduler
# SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
SHIFT: False
# SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
SIGMOID_SCALE: 1
# BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
BASE_SHIFT: 0.5
# MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
MAX_SHIFT: 1.15
#
DIFFUSION_MODEL:
# NAME DESCRIPTION: TYPE: default: 'Flux'
NAME: Flux
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@flux1-dev.safetensors
# IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
IN_CHANNELS: 64
# HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
HIDDEN_SIZE: 3072
# NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
NUM_HEADS: 24
# AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
AXES_DIM: [ 16, 56, 56 ]
# THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
THETA: 10000
# VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
VEC_IN_DIM: 768
# GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
GUIDANCE_EMBED: True
# CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
CONTEXT_IN_DIM: 4096
# MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
MLP_RATIO: 4.0
# QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
QKV_BIAS: True
# DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
DEPTH: 19
# DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
DEPTH_SINGLE_BLOCKS: 38
USE_GRAD_CHECKPOINT: True
@@ -157,55 +121,34 @@ SOLVER:
TANH_OUT: False
#
COND_STAGE_MODEL:
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
NAME: T5PlusClipFluxEmbedder
# T5_MODEL DESCRIPTION: TYPE: default: ''
T5_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: T5EncoderModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder_2/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: T5Tokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer_2/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 512
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: last_hidden_state
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: False
CLEAN: whitespace
# CLIP_MODEL DESCRIPTION: TYPE: default: ''
CLIP_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: CLIPTextModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: CLIPTokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 77
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: pooler_output
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: True
CLEAN: whitespace
USE_GRAD_CHECKPOINT: True
#
SAMPLE_ARGS:
SAMPLE_STEPS: 50
SAMPLER: flow_eluer
SAMPLER: flow_euler
SEED: 2024
IMAGE_SIZE: [ 1024, 1024 ]
GUIDE_SCALE: 3.5
@@ -2,35 +2,24 @@ ENV:
BACKEND: nccl
SEED: 166666
SOLVER:
# NAME DESCRIPTION: TYPE: default: 'LatentUfitSolver'
NAME: LatentDiffusionSolver
# MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
MAX_STEPS: 100000
# USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False
USE_AMP: True
# DTYPE DESCRIPTION: The precision for training. TYPE: str default: 'float32'
DTYPE: bfloat16
# USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
USE_FAIRSCALE: False
# USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
USE_FSDP: True
# LOAD_MODEL_ONLY DESCRIPTION: Only load the model rather than the optimizer and schedule, default is False. TYPE: bool default: False
LOAD_MODEL_ONLY: False
# RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: ''
ENABLE_GRADSCALER: False
USE_SCALER: False
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_flux_schnell_1024_lora
LOG_FILE: std_log.txt
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
EVAL_INTERVAL: 100
# LOG_TRAIN_NUM DESCRIPTION: The number samples used to log in training phase. TYPE: int default: -1
LOG_TRAIN_NUM: 16
# FSDP_REDUCE_DTYPE DESCRIPTION: The dtype of reduce in FSDP. TYPE: str default: 'float16'
FSDP_REDUCE_DTYPE: float32
# FSDP_BUFFER_DTYPE DESCRIPTION: The dtype of buffer in FSDP. TYPE: str default: 'float16'
FSDP_BUFFER_DTYPE: float32
# FSDP_SHARD_MODULES DESCRIPTION: The modules to be sharded in FSDP. TYPE: list default: ['model']
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ] #
SAVE_MODULES: [ 'model'] #
SAVE_MODULES: [ 'model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
@@ -58,12 +47,9 @@ SOLVER:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
NOISE_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
NAME: FlowMatchSigmaScheduler
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
WEIGHTING_SCHEME: logit_normal
@@ -157,54 +143,33 @@ SOLVER:
TANH_OUT: False
#
COND_STAGE_MODEL:
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
NAME: T5PlusClipFluxEmbedder
# T5_MODEL DESCRIPTION: TYPE: default: ''
T5_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: T5EncoderModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder_2/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: T5Tokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer_2/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 256
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: last_hidden_state
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: False
CLEAN: whitespace
# CLIP_MODEL DESCRIPTION: TYPE: default: ''
CLIP_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: CLIPTextModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: CLIPTokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 77
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: pooler_output
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: True
CLEAN: whitespace
#
SAMPLE_ARGS:
SAMPLE_STEPS: 4
SAMPLER: flow_eluer
SAMPLER: flow_euler
SEED: 2024
IMAGE_SIZE: [ 1024, 1024 ]
GUIDE_SCALE: 3.5
+1 -1
View File
@@ -8,11 +8,11 @@ FILE_SYSTEM:
TEMP_DIR: ./cache/cache_data
#
ENABLE_I2V: False
SKIP_EXAMPLES: True
#
MODEL:
EDIT_MODEL:
MODEL_CFG_DIR: scepter/methods/studio/chatbot/models/
DEFAULT: ace_0.6b_512
I2V:
MODEL_NAME: CogVideoX-5b-I2V
MODEL_DIR: ms://ZhipuAI/CogVideoX-5b-I2V/
@@ -0,0 +1,128 @@
NAME: ACE_0.6B_1024
IS_DEFAULT: False
USE_DYNAMIC_MODEL: True
DEFAULT_PARAS:
PARAS:
#
INPUT:
INPUT_IMAGE:
INPUT_MASK:
TASK:
PROMPT: ""
NEGATIVE_PROMPT: ""
OUTPUT_HEIGHT: 1024
OUTPUT_WIDTH: 1024
SAMPLER: ddim
SAMPLE_STEPS: 50
GUIDE_SCALE: 4.5
GUIDE_RESCALE: 0.5
SEED: -1
TAR_INDEX: 0
OUTPUT:
LATENT:
IMAGES:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
- NAME: encode
DTYPE: float16
INPUT: ["IMAGE"]
- NAME: decode
DTYPE: float16
INPUT: ["LATENT"]
#
DIFFUSION_MODEL:
FUNCTION:
- NAME: forward
DTYPE: float16
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE"]
#
COND_STAGE_MODEL:
FUNCTION:
- NAME: encode_list_of_list
DTYPE: bfloat16
INPUT: ["PROMPT"]
#
MODEL:
NAME: LatentDiffusionACE
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
DECODER_BIAS: 0.5
DEFAULT_N_PROMPT: ""
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
USE_TEXT_POS_EMBEDDINGS: True
#
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: eps
MIN_SNR_GAMMA:
NOISE_SCHEDULER:
NAME: LinearScheduler
NUM_TIMESTEPS: 1000
BETA_MIN: 0.0001
BETA_MAX: 0.02
#
DIFFUSION_MODEL:
NAME: ACE
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/dit/ace_0.6b_1024px.pth
IGNORE_KEYS: [ ]
PATCH_SIZE: 2
IN_CHANNELS: 4
HIDDEN_SIZE: 1152
DEPTH: 28
NUM_HEADS: 16
MLP_RATIO: 4.0
PRED_SIGMA: True
DROP_PATH: 0.0
WINDOW_DIZE: 0
Y_CHANNELS: 4096
MAX_SEQ_LEN: 4096
QK_NORM: True
USE_GRAD_CHECKPOINT: True
ATTENTION_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/vae/vae.bin
IGNORE_KEYS: []
#
ENCODER:
NAME: Encoder
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 4
DOUBLE_Z: True
DROPOUT: 0.0
RESAMP_WITH_CONV: True
#
DECODER:
NAME: Decoder
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 4
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/text_encoder/t5-v1_1-xxl/
TOKENIZER_PATH: ms://iic/ACE-0.6B-1024px@models/tokenizer/t5-v1_1-xxl
LENGTH: 120
T5_DTYPE: bfloat16
ADDED_IDENTIFIER: [ '{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
CLEAN: whitespace
USE_GRAD: False
@@ -0,0 +1,284 @@
NAME: ACE_0.6B_1024_REFINER
IS_DEFAULT: False
USE_DYNAMIC_MODEL: True
DEFAULT_PARAS:
PARAS:
#
INPUT:
INPUT_IMAGE:
INPUT_MASK:
TASK:
PROMPT: ""
NEGATIVE_PROMPT: ""
OUTPUT_HEIGHT: 1024
OUTPUT_WIDTH: 1024
SAMPLER: ddim
SAMPLE_STEPS: 50
GUIDE_SCALE: 4.5
GUIDE_RESCALE: 0.5
SEED: -1
TAR_INDEX: 0
REFINER_SCALE: 0.2
USE_ACE: True
#REFINER_PROMPT: "High Resolution, Sharpness, Clarity, Detail Enhancement, Noise Reduction, HD, 4k, Image Restoration, HDR"
REFINER_PROMPT: "High Resolution, Sharpness, Clarity, Detail Enhancement, Noise Reduction, HD, 4k, Image Restoration, HDR"
OUTPUT:
LATENT:
IMAGES:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
- NAME: encode
DTYPE: float16
INPUT: ["IMAGE"]
- NAME: decode
DTYPE: float16
INPUT: ["LATENT"]
#
DIFFUSION_MODEL:
FUNCTION:
- NAME: forward
DTYPE: float16
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE"]
#
COND_STAGE_MODEL:
FUNCTION:
- NAME: encode_list_of_list
DTYPE: bfloat16
INPUT: ["PROMPT"]
#
MODEL:
NAME: LatentDiffusionACE
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
DECODER_BIAS: 0.5
DEFAULT_N_PROMPT: ""
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
USE_TEXT_POS_EMBEDDINGS: True
#
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: eps
MIN_SNR_GAMMA:
NOISE_SCHEDULER:
NAME: LinearScheduler
NUM_TIMESTEPS: 1000
BETA_MIN: 0.0001
BETA_MAX: 0.02
#
DIFFUSION_MODEL:
NAME: ACE
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/dit/ace_0.6b_1024px.pth
IGNORE_KEYS: [ ]
PATCH_SIZE: 2
IN_CHANNELS: 4
HIDDEN_SIZE: 1152
DEPTH: 28
NUM_HEADS: 16
MLP_RATIO: 4.0
PRED_SIGMA: True
DROP_PATH: 0.0
WINDOW_DIZE: 0
Y_CHANNELS: 4096
MAX_SEQ_LEN: 4096
QK_NORM: True
USE_GRAD_CHECKPOINT: True
ATTENTION_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/vae/vae.bin
IGNORE_KEYS: []
#
ENCODER:
NAME: Encoder
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 4
DOUBLE_Z: True
DROPOUT: 0.0
RESAMP_WITH_CONV: True
#
DECODER:
NAME: Decoder
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 4
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/text_encoder/t5-v1_1-xxl/
TOKENIZER_PATH: ms://iic/ACE-0.6B-1024px@models/tokenizer/t5-v1_1-xxl
LENGTH: 120
T5_DTYPE: bfloat16
ADDED_IDENTIFIER: [ '{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
CLEAN: whitespace
USE_GRAD: False
ACE_PROMPT: [
"A cute cartoon rabbit holding a whiteboard that says 'ACE Refiner', standing in a sunny meadow filled with flowers, with a big smile and bright colors.",
"A beautiful young woman with long flowing hair, wearing a summer dress, holding a whiteboard that reads 'ACE Refiner' while sitting on a park bench surrounded by cherry blossoms.",
"An adorable cartoon cat wearing oversized glasses, holding a whiteboard that says 'ACE Refiner', perched on a stack of colorful books in a cozy library setting.",
"A charming girl with pigtails, wearing a cute school uniform, enthusiastically holding a whiteboard that has 'ACE Refiner' written on it, in a bright and cheerful classroom full of educational posters.",
"A friendly cartoon dog with floppy ears, sitting in front of a doghouse, proudly holding a whiteboard that says 'ACE Refiner', with a playful expression and a blue sky in the background.",
"A cute anime girl with big expressive eyes, dressed in a colorful outfit, holding a whiteboard that reads 'ACE Refiner' in a fantastical landscape filled with mythical creatures.",
"A vibrant cartoon fox holding a whiteboard that says 'ACE Refiner', standing on a rock by a sparkling stream, surrounded by lush greenery and butterflies.",
"A stylish young woman in a business outfit, smiling as she holds a whiteboard written with 'ACE Refiner', in a modern office filled with plants and natural light.",
"A cute cartoon unicorn holding a sparkling whiteboard that says 'ACE Refiner', frolicking in a magical forest, with rainbows and stars in the background.",
"A happy family, consisting of a cute little girl and her playful puppy, holding a whiteboard that says 'ACE Refiner', together in their backyard on a sunny day."
]
REFINER_MODEL:
NAME: ""
IS_DEFAULT: False
DEFAULT_PARAS:
PARAS:
RESOLUTIONS: [ [ 1024, 1024 ] ]
INPUT:
INPUT_IMAGE:
INPUT_MASK:
TASK:
PROMPT: ""
NEGATIVE_PROMPT: ""
OUTPUT_HEIGHT: 1024
OUTPUT_WIDTH: 1024
SAMPLER: flow_euler
SAMPLE_STEPS: 30
GUIDE_SCALE: 3.5
GUIDE_RESCALE:
OUTPUT:
LATENT:
IMAGES:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
- NAME: encode
DTYPE: bfloat16
INPUT: [ "IMAGE" ]
- NAME: decode
DTYPE: bfloat16
INPUT: [ "LATENT" ]
PARAS:
SCALE_FACTOR: 1.5305
SHIFT_FACTOR: 0.0609
SIZE_FACTOR: 8
DIFFUSION_MODEL:
FUNCTION:
- NAME: forward
DTYPE: bfloat16
INPUT: [ "SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE" ]
COND_STAGE_MODEL:
FUNCTION:
- NAME: encode
DTYPE: bfloat16
INPUT: [ "PROMPT" ]
MODEL:
DIFFUSION:
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
NOISE_SCHEDULER:
NAME: FlowMatchSigmaScheduler
WEIGHTING_SCHEME: logit_normal
SHIFT: 3.0
LOGIT_MEAN: 0.0
LOGIT_STD: 1.0
MODE_SCALE: 1.29
DIFFUSION_MODEL:
NAME: FluxMR
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@flux1-dev.safetensors
IN_CHANNELS: 64
OUT_CHANNELS: 64
HIDDEN_SIZE: 3072
NUM_HEADS: 24
AXES_DIM: [ 16, 56, 56 ]
THETA: 10000
VEC_IN_DIM: 768
GUIDANCE_EMBED: True
CONTEXT_IN_DIM: 4096
MLP_RATIO: 4.0
QKV_BIAS: True
DEPTH: 19
DEPTH_SINGLE_BLOCKS: 38
USE_GRAD_CHECKPOINT: True
ATTN_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLFlux
EMBED_DIM: 16
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@ae.safetensors
IGNORE_KEYS: [ ]
BATCH_SIZE: 8
USE_CONV: False
SCALE_FACTOR: 0.3611
SHIFT_FACTOR: 0.1159
#
ENCODER:
NAME: Encoder
USE_CHECKPOINT: False
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 16
DOUBLE_Z: True
DROPOUT: 0.0
RESAMP_WITH_CONV: True
#
DECODER:
NAME: Decoder
USE_CHECKPOINT: False
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 16
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
NAME: T5PlusClipFluxEmbedder
T5_MODEL:
NAME: HFEmbedder
HF_MODEL_CLS: T5EncoderModel
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder_2/
HF_TOKENIZER_CLS: T5Tokenizer
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer_2/
MAX_LENGTH: 512
OUTPUT_KEY: last_hidden_state
D_TYPE: bfloat16
BATCH_INFER: False
CLEAN: whitespace
CLIP_MODEL:
NAME: HFEmbedder
HF_MODEL_CLS: CLIPTextModel
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder/
HF_TOKENIZER_CLS: CLIPTokenizer
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer/
MAX_LENGTH: 77
OUTPUT_KEY: pooler_output
D_TYPE: bfloat16
BATCH_INFER: True
CLEAN: whitespace
@@ -1,5 +1,6 @@
NAME: ACE_0.6B_512
IS_DEFAULT: False
IS_DEFAULT: True
USE_DYNAMIC_MODEL: True
DEFAULT_PARAS:
PARAS:
#
@@ -39,7 +40,7 @@ DEFAULT_PARAS:
#
COND_STAGE_MODEL:
FUNCTION:
- NAME: encode_list
- NAME: encode_list_of_list
DTYPE: bfloat16
INPUT: ["PROMPT"]
#
@@ -0,0 +1,152 @@
NAME: COGVIDEOX_2B
IS_DEFAULT: False
DEFAULT_PARAS:
PARAS:
RESOLUTIONS: [[480, 720]]
INPUT:
IMAGE:
ORIGINAL_SIZE_AS_TUPLE: [480, 720]
TARGET_SIZE_AS_TUPLE: [480, 720]
PROMPT: ""
NEGATIVE_PROMPT: ""
PROMPT_PREFIX: ""
SAMPLE: ddim
SAMPLE_STEPS: 50
GUIDE_SCALE: 6.0
GUIDE_RESCALE: 0.0
DISCRETIZATION: trailing
NUM_FRAMES:
DEFAULT: 49
VISIBLE: True
FPS:
DEFAULT: 8
VISIBLE: True
OUTPUT:
VIDEOS:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
-
NAME: decode
DTYPE: bfloat16
INPUT: ["LATENT"]
PARAS:
SCALING_FACTOR_IMAGE: 1.15258426
DIFFUSION_MODEL:
FUNCTION:
-
NAME: forward
DTYPE: bfloat16
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE", "GUIDE_RESCALE", "DISCRETIZATION", "NUM_FRAMES", "FPS"]
PARAS:
USE_ROTARY_POSITIONAL_EMBEDDINGS: False
PATCH_SIZE: 2
LATENT_CHANNELS: 16
SCALE_FACTOR_SPATIAL: 8
SCALE_FACTOR_TEMPORAL: 4
ATTENTION_HEAD_DIM: 64
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
COND_STAGE_MODEL:
FUNCTION:
-
NAME: encode
DTYPE: bfloat16
INPUT: ["PROMPT"]
#
MODEL:
PRETRAINED_MODEL:
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: v
NOISE_SCHEDULER:
NAME: ScaledLinearScheduler
BETA_MIN: 0.00085
BETA_MAX: 0.012
SNR_SHIFT_SCALE: 3.0
RESCALE_BETAS_ZERO_SNR: True
DIFFUSION_SAMPLERS:
NAME: DDIMSampler
DISCRETIZATION_TYPE: trailing
ETA: 0.0
#
DIFFUSION_MODEL:
NAME: CogVideoXTransformer3DModel
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@transformer/diffusion_pytorch_model.safetensors
NUM_ATTENTION_HEADS: 30
ATTENTION_HEAD_DIM: 64
IN_CHANNELS: 16
OUT_CHANNELS: 16
FLIP_SIN_TO_COS: True
FREQ_SHIFT: 0
TIME_EMBED_DIM: 512
TEXT_EMBED_DIM: 4096
NUM_LAYERS: 30
DROPOUT: 0.0
ATTENTION_BIAS: True
SAMPLE_WIDTH: 90
SAMPLE_HEIGHT: 60
SAMPLE_FRAMES: 49
PATCH_SIZE: 2
TEMPORAL_COMPRESSION_RATIO: 4
MAX_TEXT_SEQ_LENGTH: 226
ACTIVATION_FN: "gelu-approximate"
TIMESTEP_ACTIVATION_FN: "silu"
NORM_ELEMENTWISE_AFFINE: True
NORM_EPS: 1e-5
SPATIAL_INTERPOLATION_SCALE: 1.875
TEMPORAL_INTERPOLATION_SCALE: 1.0
USE_ROTARY_POSITIONAL_EMBEDDINGS: False
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
GRADIENT_CHECKPOINTING: False
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@vae/diffusion_pytorch_model.safetensors
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: False
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: False
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
LENGTH: 226
CLEAN:
USE_GRAD: False
T5_DTYPE: bfloat16
@@ -0,0 +1,154 @@
NAME: COGVIDEOX_5B
IS_DEFAULT: False
DEFAULT_PARAS:
PARAS:
RESOLUTIONS: [[480, 720]]
INPUT:
IMAGE:
ORIGINAL_SIZE_AS_TUPLE: [480, 720]
TARGET_SIZE_AS_TUPLE: [480, 720]
PROMPT: ""
NEGATIVE_PROMPT: ""
PROMPT_PREFIX: ""
SAMPLE: ddim
SAMPLE_STEPS: 50
GUIDE_SCALE: 6.0
GUIDE_RESCALE: 0.0
DISCRETIZATION: trailing
NUM_FRAMES:
DEFAULT: 49
VISIBLE: True
FPS:
DEFAULT: 8
VISIBLE: True
OUTPUT:
VIDEOS:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
-
NAME: decode
DTYPE: bfloat16
INPUT: ["LATENT"]
PARAS:
SCALING_FACTOR_IMAGE: 0.7 # 5b diff
DIFFUSION_MODEL:
FUNCTION:
-
NAME: forward
DTYPE: bfloat16
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE", "GUIDE_RESCALE", "DISCRETIZATION", "NUM_FRAMES", "FPS"]
PARAS:
USE_ROTARY_POSITIONAL_EMBEDDINGS: True
PATCH_SIZE: 2
LATENT_CHANNELS: 16
SCALE_FACTOR_SPATIAL: 8
SCALE_FACTOR_TEMPORAL: 4
ATTENTION_HEAD_DIM: 64
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
COND_STAGE_MODEL:
FUNCTION:
-
NAME: encode
DTYPE: bfloat16
INPUT: ["PROMPT"]
#
MODEL:
PRETRAINED_MODEL:
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: v
NOISE_SCHEDULER:
NAME: ScaledLinearScheduler
BETA_MIN: 0.00085
BETA_MAX: 0.012
SNR_SHIFT_SCALE: 1.0 # 5b diff
RESCALE_BETAS_ZERO_SNR: True
DIFFUSION_SAMPLERS:
NAME: DDIMSampler
DISCRETIZATION_TYPE: trailing
ETA: 0.0
#
DIFFUSION_MODEL:
NAME: CogVideoXTransformer3DModel
DTYPE: bfloat16
PRETRAINED_MODEL: # 5b diff
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00001-of-00002.safetensors
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00002-of-00002.safetensors
NUM_ATTENTION_HEADS: 48 # 5b diff
ATTENTION_HEAD_DIM: 64
IN_CHANNELS: 16
OUT_CHANNELS: 16
FLIP_SIN_TO_COS: True
FREQ_SHIFT: 0
TIME_EMBED_DIM: 512
TEXT_EMBED_DIM: 4096
NUM_LAYERS: 42 # 5b diff
DROPOUT: 0.0
ATTENTION_BIAS: True
SAMPLE_WIDTH: 90
SAMPLE_HEIGHT: 60
SAMPLE_FRAMES: 49
PATCH_SIZE: 2
TEMPORAL_COMPRESSION_RATIO: 4
MAX_TEXT_SEQ_LENGTH: 226
ACTIVATION_FN: "gelu-approximate"
TIMESTEP_ACTIVATION_FN: "silu"
NORM_ELEMENTWISE_AFFINE: True
NORM_EPS: 1e-5
SPATIAL_INTERPOLATION_SCALE: 1.875
TEMPORAL_INTERPOLATION_SCALE: 1.0
USE_ROTARY_POSITIONAL_EMBEDDINGS: True # 5b diff
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
GRADIENT_CHECKPOINTING: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b@vae/diffusion_pytorch_model.safetensors # 5b diff
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: True
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
LENGTH: 226
CLEAN:
USE_GRAD: False
T5_DTYPE: bfloat16
@@ -13,8 +13,8 @@ DEFAULT_PARAS:
VISIBLE: False
PROMPT_PREFIX: ""
SAMPLE:
VALUES: ["flow_eluer"]
DEFAULT: "flow_eluer"
VALUES: ["flow_euler"]
DEFAULT: "flow_euler"
SAMPLE_STEPS: 50
GUIDE_SCALE: 3.5
GUIDE_RESCALE:
@@ -13,8 +13,8 @@ DEFAULT_PARAS:
VISIBLE: False
PROMPT_PREFIX: ""
SAMPLE:
VALUES: ["flow_eluer"]
DEFAULT: "flow_eluer"
VALUES: ["flow_euler"]
DEFAULT: "flow_euler"
SAMPLE_STEPS: 4
GUIDE_SCALE: 3.5
GUIDE_RESCALE:
@@ -1,4 +1,5 @@
WORK_DIR: "inference"
SKIP_EXAMPLES: True
DIFFUSION_PARAS:
SAMPLE:
VALUES: ['ddim', 'euler', 'euler_ancestral', 'heun', 'dpm2',
@@ -18,6 +19,16 @@ DIFFUSION_PARAS:
MAX: 4
DEFAULT: 1
VISIBLE: True
NUM_FRAMES:
MIN: 1
MAX: 100
DEFAULT: 49
VISIBLE: False
FPS:
MIN: 1
MAX: 50
DEFAULT: 8
VISIBLE: False
SAMPLE_STEPS:
MIN: 1
MAX: 100
@@ -93,7 +104,8 @@ DIFFUSION_PARAS:
[1664, 576], [1728, 576],
[2048, 2048], [2048, 1920], [1920, 2048],
[1536, 2560], [2560, 1536], [2560, 1440],
[2560, 1440]
[2560, 1440],
[480, 720], [720, 480]
]
DEFAULT: [1024, 1024]
VISIBLE: True
@@ -450,3 +450,28 @@ PROCESSORS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
VIDEO_PROCESSORS:
- NAME: CogVLM2Llama3Caption
TYPE: caption
MODEL_PATH: ms://ZhipuAI/cogvlm2-llama3-caption
DEVICE: "gpu"
MEMORY: 20000
PROMPT: Please describe this video in detail.
TEMPERATURE: 0.1
MAX_NEW_TOKENS: 2048
PAD_TOKEN_ID: 128002
TOP_K: 1
TOP_P: 0.1
TRANSLATION_PROCESSORS:
- NAME: OpusMtZhEn
TYPE: caption
MODEL_PATH: ms://cubeai/trans-opus-mt-zh-en
DEVICE: "gpu"
MEMORY: 5000
- NAME: OpusMtEnZh
TYPE: caption
MODEL_PATH: ms://cubeai/trans-opus-mt-en-zh
DEVICE: "gpu"
MEMORY: 5000
+1 -1
View File
@@ -89,5 +89,5 @@ INTERFACE:
CONFIG: scepter/methods/studio/inference/inference.yaml
- NAME: 对话式编辑
NAME_EN: ChatBot
IFID: ChatBot
IFID: chatbot
CONFIG: scepter/methods/studio/chatbot/chatbot.yaml
@@ -0,0 +1,315 @@
ENV:
BACKEND: nccl
SEED: 42
TENSOR_PARALLEL_SIZE: 1
PIPELINE_PARALLEL_SIZE: 1
SYS_ENVS:
TORCH_CUDNN_V8_API_ENABLED: '1'
TOKENIZERS_PARALLELISM: 'false'
TF_CPP_MIN_LOG_LEVEL: '3'
PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True'
META:
VERSION: 'COGVIDEOX_2B'
DESCRIPTION: "cogvideox 2b"
IS_DEFAULT: False
IS_SHARE: True
INFERENCE_PARAS:
INFERENCE_BATCH_SIZE: 1
INFERENCE_PREFIX: ""
DEFAULT_SAMPLER: "ddim"
DEFAULT_SAMPLE_STEPS: 50
INFERENCE_N_PROMPT: ""
RESOLUTION: [ 480, 720 ]
PARAS:
- TRAIN_BATCH_SIZE: 1
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: [ 480, 720 ]
MEMORY: 89000
EPOCHS: 50
SAVE_INTERVAL: 25
EPSEC: 0.818
LEARNING_RATE: 4e-4
IS_DEFAULT: False
TUNER: FULL
- TRAIN_BATCH_SIZE: 1
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: [ 480, 720 ]
MEMORY: 89000
EPOCHS: 50
SAVE_INTERVAL: 25
EPSEC: 0.818
LEARNING_RATE: 4e-4
IS_DEFAULT: True
TUNER: LORA
#
TUNERS:
LORA:
- NAME: SwiftLoRA
R: 64
LORA_ALPHA: 64
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$"
#
SOLVER:
NAME: LatentDiffusionVideoSolver
MAX_STEPS: 2000
USE_AMP: True
DTYPE: bfloat16
USE_FAIRSCALE: False
USE_FSDP: True
LOAD_MODEL_ONLY: False
ENABLE_GRADSCALER: False
USE_SCALER: False
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_cogvideox_2b_lora
LOG_FILE: std_log.txt
EVAL_INTERVAL: 100
LOG_TRAIN_NUM: 4
FPS: 8
SHARDING_STRATEGY: full_shard
FSDP_REDUCE_DTYPE: float32
FSDP_BUFFER_DTYPE: float32
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model']
SAVE_MODULES: [ 'model', 'cond_stage_model.model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
TUNER:
#
MODEL:
NAME: LatentDiffusionCogVideoX
PRETRAINED_MODEL:
PARAMETERIZATION: v
TIMESTEPS: 1000
MIN_SNR_GAMMA: 3.0
ZERO_TERMINAL_SNR: True
SCALE_FACTOR_SPATIAL: 8
SCALE_FACTOR_TEMPORAL: 4
SCALING_FACTOR_IMAGE: 1.15258426
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: v
NOISE_SCHEDULER:
NAME: ScaledLinearScheduler
BETA_MIN: 0.00085
BETA_MAX: 0.012
SNR_SHIFT_SCALE: 3.0
RESCALE_BETAS_ZERO_SNR: True
DIFFUSION_SAMPLERS:
NAME: DDIMSampler
DISCRETIZATION_TYPE: trailing
ETA: 0.0
#
DIFFUSION_MODEL:
NAME: CogVideoXTransformer3DModel
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@transformer/diffusion_pytorch_model.safetensors
NUM_ATTENTION_HEADS: 30
ATTENTION_HEAD_DIM: 64
IN_CHANNELS: 16
OUT_CHANNELS: 16
FLIP_SIN_TO_COS: True
FREQ_SHIFT: 0
TIME_EMBED_DIM: 512
TEXT_EMBED_DIM: 4096
NUM_LAYERS: 30
DROPOUT: 0.0
ATTENTION_BIAS: True
SAMPLE_WIDTH: 90
SAMPLE_HEIGHT: 60
SAMPLE_FRAMES: 49
PATCH_SIZE: 2
TEMPORAL_COMPRESSION_RATIO: 4
MAX_TEXT_SEQ_LENGTH: 226
ACTIVATION_FN: "gelu-approximate"
TIMESTEP_ACTIVATION_FN: "silu"
NORM_ELEMENTWISE_AFFINE: True
NORM_EPS: 1e-5
SPATIAL_INTERPOLATION_SCALE: 1.875
TEMPORAL_INTERPOLATION_SCALE: 1.0
USE_ROTARY_POSITIONAL_EMBEDDINGS: False
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
GRADIENT_CHECKPOINTING: False
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@vae/diffusion_pytorch_model.safetensors
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: False
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: False
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
LENGTH: 226
CLEAN:
USE_GRAD: False
#
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 42
GUIDE_SCALE: 6.0
GUIDE_RESCALE: 0.0
NUM_FRAMES: 49
#
OPTIMIZER:
NAME: Adam
LEARNING_RATE: 1e-3
BETAS: [ 0.9, 0.95 ]
EPS: 1e-8
WEIGHT_DECAY: 0.0
AMSGRAD: False
#
# LR_SCHEDULER:
# NAME: StepAnnealingLR
# WARMUP_STEPS: 200
# TOTAL_STEPS: 2000
# DECAY_MODE: 'cosine'
#
TRAIN_DATA:
NAME: VideoGenDatasetOTF
MODE: train
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
PROMPT_PREFIX: ''
DELIMITER: '#;#'
FIELDS: [ 'video_path', 'width', 'height', 'prompt' ]
PATH_PREFIX:
DATA_FILE:
SAMPLER:
NAME: LoopSampler
TRANSFORMS:
- NAME: Select
KEYS: [ 'video', 'video_latent', "prompt" ]
META_KEYS: [ ]
MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@vae/diffusion_pytorch_model.safetensors
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: True
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
PROMPT_FILE:
PROMPT_DATA: [ "A girl riding a bike.", "A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance." ]
IMAGE_SIZE: [ 480, 720 ]
FIELDS: [ "prompt" ]
DELIMITER: '#;#'
PROMPT_PREFIX: ''
PIN_MEMORY: True
BATCH_SIZE: 1
# USE_NUM: 8
NUM_WORKERS: 4
TRANSFORMS:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
PRIORITY: 20
- NAME: CheckpointHook
INTERVAL: 1000
PRIORITY: 40
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
DISABLE_SNAPSHOT: True
#
EVAL_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
SAVE_PROBE_PREFIX: 'image'
@@ -0,0 +1,317 @@
ENV:
BACKEND: nccl
SEED: 42
TENSOR_PARALLEL_SIZE: 1
PIPELINE_PARALLEL_SIZE: 1
SYS_ENVS:
TORCH_CUDNN_V8_API_ENABLED: '1'
TOKENIZERS_PARALLELISM: 'false'
TF_CPP_MIN_LOG_LEVEL: '3'
PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True'
META:
VERSION: 'COGVIDEOX_5B'
DESCRIPTION: "cogvideox 5b"
IS_DEFAULT: False
IS_SHARE: True
INFERENCE_PARAS:
INFERENCE_BATCH_SIZE: 1
INFERENCE_PREFIX: ""
DEFAULT_SAMPLER: "ddim"
DEFAULT_SAMPLE_STEPS: 50
INFERENCE_N_PROMPT: ""
RESOLUTION: [ 480, 720 ]
PARAS:
- TRAIN_BATCH_SIZE: 1
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: [ 480, 720 ]
MEMORY: 89000
EPOCHS: 50
SAVE_INTERVAL: 25
EPSEC: 0.818
LEARNING_RATE: 4e-4
IS_DEFAULT: False
TUNER: FULL
- TRAIN_BATCH_SIZE: 1
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: [ 480, 720 ]
MEMORY: 89000
EPOCHS: 50
SAVE_INTERVAL: 25
EPSEC: 0.818
LEARNING_RATE: 4e-4
IS_DEFAULT: True
TUNER: LORA
#
TUNERS:
LORA:
- NAME: SwiftLoRA
R: 64
LORA_ALPHA: 64
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$"
#
SOLVER:
NAME: LatentDiffusionVideoSolver
MAX_STEPS: 2000
USE_AMP: True
DTYPE: bfloat16
USE_FAIRSCALE: False
USE_FSDP: True
LOAD_MODEL_ONLY: False
ENABLE_GRADSCALER: False
USE_SCALER: False
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_cogvideox_5b_lora
LOG_FILE: std_log.txt
EVAL_INTERVAL: 100
LOG_TRAIN_NUM: 4
FPS: 8
SHARDING_STRATEGY: full_shard
FSDP_REDUCE_DTYPE: float32
FSDP_BUFFER_DTYPE: float32
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model']
SAVE_MODULES: [ 'model', 'cond_stage_model.model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
TUNER:
#
MODEL:
NAME: LatentDiffusionCogVideoX
PRETRAINED_MODEL:
PARAMETERIZATION: v
TIMESTEPS: 1000
MIN_SNR_GAMMA: 3.0
ZERO_TERMINAL_SNR: True
SCALE_FACTOR_SPATIAL: 8
SCALE_FACTOR_TEMPORAL: 4
SCALING_FACTOR_IMAGE: 0.7 # 5b diff
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: v
NOISE_SCHEDULER:
NAME: ScaledLinearScheduler
BETA_MIN: 0.00085
BETA_MAX: 0.012
SNR_SHIFT_SCALE: 1.0 # 5b diff
RESCALE_BETAS_ZERO_SNR: True
DIFFUSION_SAMPLERS:
NAME: DDIMSampler
DISCRETIZATION_TYPE: trailing
ETA: 0.0
#
DIFFUSION_MODEL:
NAME: CogVideoXTransformer3DModel
DTYPE: bfloat16
PRETRAINED_MODEL: # 5b diff
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00001-of-00002.safetensors
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00002-of-00002.safetensors
NUM_ATTENTION_HEADS: 48 # 5b diff
ATTENTION_HEAD_DIM: 64
IN_CHANNELS: 16
OUT_CHANNELS: 16
FLIP_SIN_TO_COS: True
FREQ_SHIFT: 0
TIME_EMBED_DIM: 512
TEXT_EMBED_DIM: 4096
NUM_LAYERS: 42 # 5b diff
DROPOUT: 0.0
ATTENTION_BIAS: True
SAMPLE_WIDTH: 90
SAMPLE_HEIGHT: 60
SAMPLE_FRAMES: 49
PATCH_SIZE: 2
TEMPORAL_COMPRESSION_RATIO: 4
MAX_TEXT_SEQ_LENGTH: 226
ACTIVATION_FN: "gelu-approximate"
TIMESTEP_ACTIVATION_FN: "silu"
NORM_ELEMENTWISE_AFFINE: True
NORM_EPS: 1e-5
SPATIAL_INTERPOLATION_SCALE: 1.875
TEMPORAL_INTERPOLATION_SCALE: 1.0
USE_ROTARY_POSITIONAL_EMBEDDINGS: True # 5b diff
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
GRADIENT_CHECKPOINTING: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b@vae/diffusion_pytorch_model.safetensors # 5b diff
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: True
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
LENGTH: 226
CLEAN:
USE_GRAD: False
#
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 42
GUIDE_SCALE: 6.0
GUIDE_RESCALE: 0.0
NUM_FRAMES: 49
#
OPTIMIZER:
NAME: Adam
LEARNING_RATE: 1e-3
BETAS: [ 0.9, 0.95 ]
EPS: 1e-8
WEIGHT_DECAY: 0.0
AMSGRAD: False
#
# LR_SCHEDULER:
# NAME: StepAnnealingLR
# WARMUP_STEPS: 200
# TOTAL_STEPS: 2000
# DECAY_MODE: 'cosine'
#
TRAIN_DATA:
NAME: VideoGenDatasetOTF
MODE: train
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
PROMPT_PREFIX: ''
DELIMITER: '#;#'
FIELDS: [ 'video_path', 'width', 'height', 'prompt' ]
PATH_PREFIX:
DATA_FILE:
SAMPLER:
NAME: LoopSampler
TRANSFORMS:
- NAME: Select
KEYS: [ 'video', 'video_latent', "prompt" ]
META_KEYS: [ ]
MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b@vae/diffusion_pytorch_model.safetensors
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: True
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
PROMPT_FILE:
PROMPT_DATA: [ "A girl riding a bike.", "A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance." ]
IMAGE_SIZE: [ 480, 720 ]
FIELDS: [ "prompt" ]
DELIMITER: '#;#'
PROMPT_PREFIX: ''
PIN_MEMORY: True
BATCH_SIZE: 1
# USE_NUM: 8
NUM_WORKERS: 4
TRANSFORMS:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
PRIORITY: 20
- NAME: CheckpointHook
INTERVAL: 1000
PRIORITY: 40
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
DISABLE_SNAPSHOT: True
#
EVAL_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
SAVE_PROBE_PREFIX: 'image'
@@ -3,7 +3,7 @@ ENV:
META:
VERSION: 'FLUX1.0_DEV'
DESCRIPTION: "flux 1.0 dev"
IS_DEFAULT: False
IS_DEFAULT: True
IS_SHARE: True
INFERENCE_PARAS:
INFERENCE_BATCH_SIZE: 1
@@ -50,43 +50,33 @@ META:
#
SOLVER:
NAME: LatentDiffusionSolver
# MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
MAX_STEPS: 100000
# USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False
USE_AMP: True
# DTYPE DESCRIPTION: The precision for training. TYPE: str default: 'float32'
DTYPE: bfloat16
# USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
USE_FAIRSCALE: False
# USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
USE_FSDP: True
# LOAD_MODEL_ONLY DESCRIPTION: Only load the model rather than the optimizer and schedule, default is False. TYPE: bool default: False
LOAD_MODEL_ONLY: False
# RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: ''
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_flux_dev_1024_lora
LOG_FILE: std_log.txt
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
EVAL_INTERVAL: 100
# LOG_TRAIN_NUM DESCRIPTION: The number samples used to log in training phase. TYPE: int default: -1
LOG_TRAIN_NUM: 16
# FSDP_REDUCE_DTYPE DESCRIPTION: The dtype of reduce in FSDP. TYPE: str default: 'float16'
ENABLE_GRADSCALER: False
USE_SCALER: False
FSDP_REDUCE_DTYPE: float32
# FSDP_BUFFER_DTYPE DESCRIPTION: The dtype of buffer in FSDP. TYPE: str default: 'float16'
FSDP_BUFFER_DTYPE: float32
# FSDP_SHARD_MODULES DESCRIPTION: The modules to be sharded in FSDP. TYPE: list default: ['model']
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ] #
SAVE_MODULES: [ 'model'] #
SAVE_MODULES: [ 'model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
FREEZE:
#
TUNER:
#
MODEL:
NAME: LatentDiffusionFlux
PARAMETERIZATION: rf
@@ -99,65 +89,39 @@ SOLVER:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
NOISE_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
NAME: FlowMatchSigmaScheduler
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
WEIGHTING_SCHEME: logit_normal
SHIFT: 3.0
# LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
LOGIT_MEAN: 0.0
# LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
LOGIT_STD: 1.0
# MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
MODE_SCALE: 1.29
SAMPLER_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
NAME: FlowMatchFluxShiftScheduler
# SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
SHIFT: False
# SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
SIGMOID_SCALE: 1
# BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
BASE_SHIFT: 0.5
# MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
MAX_SHIFT: 1.15
#
DIFFUSION_MODEL:
# NAME DESCRIPTION: TYPE: default: 'Flux'
NAME: Flux
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@flux1-dev.safetensors
# IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
IN_CHANNELS: 64
# HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
HIDDEN_SIZE: 3072
# NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
NUM_HEADS: 24
# AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
AXES_DIM: [ 16, 56, 56 ]
# THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
THETA: 10000
# VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
VEC_IN_DIM: 768
# GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
GUIDANCE_EMBED: False
# CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
CONTEXT_IN_DIM: 4096
# MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
MLP_RATIO: 4.0
# QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
QKV_BIAS: True
# DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
DEPTH: 19
# DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
DEPTH_SINGLE_BLOCKS: 38
USE_GRAD_CHECKPOINT: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLFlux
EMBED_DIM: 16
@@ -167,7 +131,7 @@ SOLVER:
USE_CONV: False
SCALE_FACTOR: 0.3611
SHIFT_FACTOR: 0.1159
#
ENCODER:
NAME: Encoder
USE_CHECKPOINT: True
@@ -181,7 +145,7 @@ SOLVER:
DOUBLE_Z: True
DROPOUT: 0.0
RESAMP_WITH_CONV: True
#
DECODER:
NAME: Decoder
USE_CHECKPOINT: True
@@ -196,61 +160,40 @@ SOLVER:
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
NAME: T5PlusClipFluxEmbedder
# T5_MODEL DESCRIPTION: TYPE: default: ''
T5_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: T5EncoderModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder_2/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: T5Tokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer_2/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 512
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: last_hidden_state
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: False
CLEAN: whitespace
# CLIP_MODEL DESCRIPTION: TYPE: default: ''
CLIP_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: CLIPTextModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: CLIPTokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 77
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: pooler_output
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: True
CLEAN: whitespace
#
SAMPLE_ARGS:
SAMPLE_STEPS: 50
SAMPLER: flow_eluer
SAMPLER: flow_euler
SEED: 2024
IMAGE_SIZE: [ 1024, 1024 ]
SHIFT: True
GUIDE_SCALE: 3.5
#
OPTIMIZER:
NAME: AdamW
LEARNING_RATE: 4e-4
@@ -258,7 +201,7 @@ SOLVER:
EPS: 1e-8
WEIGHT_DECAY: 1e-2
AMSGRAD: False
#
TRAIN_DATA:
NAME: ImageTextPairMSDataset
MODE: train
@@ -302,7 +245,7 @@ SOLVER:
- NAME: Select
KEYS: [ 'image', 'prompt' ]
META_KEYS: [ 'data_key' ]
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
@@ -319,13 +262,12 @@ SOLVER:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
# GRADIENT_CLIP: 1.0
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
@@ -50,43 +50,33 @@ META:
#
SOLVER:
NAME: LatentDiffusionSolver
# MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
MAX_STEPS: 100000
# USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False
USE_AMP: True
# DTYPE DESCRIPTION: The precision for training. TYPE: str default: 'float32'
DTYPE: bfloat16
# USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
USE_FAIRSCALE: False
# USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
USE_FSDP: True
# LOAD_MODEL_ONLY DESCRIPTION: Only load the model rather than the optimizer and schedule, default is False. TYPE: bool default: False
LOAD_MODEL_ONLY: False
# RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: ''
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_flux_schnell_1024_lora
LOG_FILE: std_log.txt
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
EVAL_INTERVAL: 100
# LOG_TRAIN_NUM DESCRIPTION: The number samples used to log in training phase. TYPE: int default: -1
LOG_TRAIN_NUM: 16
# FSDP_REDUCE_DTYPE DESCRIPTION: The dtype of reduce in FSDP. TYPE: str default: 'float16'
ENABLE_GRADSCALER: False
USE_SCALER: False
FSDP_REDUCE_DTYPE: float32
# FSDP_BUFFER_DTYPE DESCRIPTION: The dtype of buffer in FSDP. TYPE: str default: 'float16'
FSDP_BUFFER_DTYPE: float32
# FSDP_SHARD_MODULES DESCRIPTION: The modules to be sharded in FSDP. TYPE: list default: ['model']
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ] #
SAVE_MODULES: [ 'model'] #
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ]
SAVE_MODULES: [ 'model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
FREEZE:
#
TUNER:
#
MODEL:
NAME: LatentDiffusionFlux
PARAMETERIZATION: rf
@@ -99,65 +89,39 @@ SOLVER:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
NOISE_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
NAME: FlowMatchSigmaScheduler
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
WEIGHTING_SCHEME: logit_normal
SHIFT: 3.0
# LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
LOGIT_MEAN: 0.0
# LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
LOGIT_STD: 1.0
# MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
MODE_SCALE: 1.29
SAMPLER_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
NAME: FlowMatchFluxShiftScheduler
# SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
SHIFT: False
# SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
SIGMOID_SCALE: 1
# BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
BASE_SHIFT: 0.5
# MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
MAX_SHIFT: 1.15
#
DIFFUSION_MODEL:
# NAME DESCRIPTION: TYPE: default: 'Flux'
NAME: Flux
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-schnell@flux1-schnell.safetensors
# IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
IN_CHANNELS: 64
# HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
HIDDEN_SIZE: 3072
# NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
NUM_HEADS: 24
# AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
AXES_DIM: [ 16, 56, 56 ]
# THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
THETA: 10000
# VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
VEC_IN_DIM: 768
# GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
GUIDANCE_EMBED: False
# CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
CONTEXT_IN_DIM: 4096
# MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
MLP_RATIO: 4.0
# QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
QKV_BIAS: True
# DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
DEPTH: 19
# DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
DEPTH_SINGLE_BLOCKS: 38
USE_GRAD_CHECKPOINT: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLFlux
EMBED_DIM: 16
@@ -167,7 +131,7 @@ SOLVER:
USE_CONV: False
SCALE_FACTOR: 0.3611
SHIFT_FACTOR: 0.1159
#
ENCODER:
NAME: Encoder
USE_CHECKPOINT: True
@@ -181,7 +145,7 @@ SOLVER:
DOUBLE_Z: True
DROPOUT: 0.0
RESAMP_WITH_CONV: True
#
DECODER:
NAME: Decoder
USE_CHECKPOINT: True
@@ -196,60 +160,39 @@ SOLVER:
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
NAME: T5PlusClipFluxEmbedder
# T5_MODEL DESCRIPTION: TYPE: default: ''
T5_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: T5EncoderModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder_2/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: T5Tokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer_2/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 256
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: last_hidden_state
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: False
CLEAN: whitespace
# CLIP_MODEL DESCRIPTION: TYPE: default: ''
CLIP_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: CLIPTextModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: CLIPTokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 77
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: pooler_output
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: True
CLEAN: whitespace
#
SAMPLE_ARGS:
SAMPLE_STEPS: 4
SAMPLER: flow_eluer
SAMPLER: flow_euler
SEED: 2024
IMAGE_SIZE: [ 1024, 1024 ]
GUIDE_SCALE: 3.5
#
OPTIMIZER:
NAME: AdamW
LEARNING_RATE: 4e-4
@@ -257,7 +200,7 @@ SOLVER:
EPS: 1e-8
WEIGHT_DECAY: 1e-2
AMSGRAD: False
#
TRAIN_DATA:
NAME: ImageTextPairMSDataset
MODE: train
@@ -301,7 +244,7 @@ SOLVER:
- NAME: Select
KEYS: [ 'image', 'prompt' ]
META_KEYS: [ 'data_key' ]
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
@@ -318,13 +261,12 @@ SOLVER:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
# GRADIENT_CLIP: 1.0
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
@@ -339,4 +281,4 @@ SOLVER:
PROB_INTERVAL: 100
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
SAVE_PROBE_PREFIX: 'image'
SAVE_PROBE_PREFIX: 'image'
@@ -35,7 +35,7 @@ META:
SAVE_INTERVAL: 25
EPSEC: 0.818
LEARNING_RATE: 0.0001
IS_DEFAULT: False
IS_DEFAULT: True
TUNER: LORA
#
TUNERS:
@@ -13,8 +13,10 @@ TRAIN_PARAS:
VALUES: [[256, 256], [320, 180], [180, 320],
[512, 512], [640, 360], [360, 640],
[768, 768], [960, 540], [540, 960],
[1024, 1024], [1280, 720], [720, 1280]]
[1024, 1024], [1280, 720], [720, 1280],
[720, 480], [480, 720]]
DEFAULT: [1024, 1024]
EVAL_PROMPTS:
- a boy wearing a jacket
- a dog running on the lawn
SAVE_FILE_LOCAL_PATH: "cache/scepter_ui/datasets/train_data_from_list"
+21 -2
View File
@@ -1,4 +1,23 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules import (data, inference, model, opt, solver, transform,
utils)
from typing import TYPE_CHECKING
from scepter.modules.utils.import_utils import LazyImportModule
if TYPE_CHECKING:
from scepter.modules import (data, inference, model, opt, solver, transform,
utils)
else:
_import_structure = {
'modules': ['data', 'inference', 'model', 'opt', 'solver',
'transform', 'utils']
}
import sys
sys.modules[__name__] = LazyImportModule(
__name__,
globals()['__file__'],
_import_structure,
module_spec=__spec__,
extra_objects={},
)
+62 -21
View File
@@ -1,23 +1,64 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.annotator.base_annotator import GeneralAnnotator
from scepter.modules.annotator.canny import CannyAnnotator
from scepter.modules.annotator.color import ColorAnnotator
from scepter.modules.annotator.degradation import DegradationAnnotator
from scepter.modules.annotator.doodle import DoodleAnnotator
from scepter.modules.annotator.gray import GrayAnnotator
from scepter.modules.annotator.hed import HedAnnotator
from scepter.modules.annotator.identity import IdentityAnnotator
from scepter.modules.annotator.informative_drawing import (
InfoDrawAnimeAnnotator, InfoDrawContourAnnotator,
InfoDrawOpenSketchAnnotator)
from scepter.modules.annotator.inpainting import InpaintingAnnotator
from scepter.modules.annotator.invert import InvertAnnotator
from scepter.modules.annotator.midas_op import MidasDetector
from scepter.modules.annotator.mlsd_op import MLSDdetector
from scepter.modules.annotator.openpose import OpenposeAnnotator
from scepter.modules.annotator.outpainting import OutpaintingAnnotator, OutpaintingResize
from scepter.modules.annotator.pidinet import PiDiAnnotator
from scepter.modules.annotator.segmentation import ESAMAnnotator
from scepter.modules.annotator.sketch import SketchAnnotator
from scepter.modules.annotator.lama import LamaAnnotator
from typing import TYPE_CHECKING
from scepter.modules.utils.import_utils import LazyImportModule
if TYPE_CHECKING:
from scepter.modules.annotator.base_annotator import GeneralAnnotator
from scepter.modules.annotator.canny import CannyAnnotator
from scepter.modules.annotator.color import ColorAnnotator
from scepter.modules.annotator.degradation import DegradationAnnotator
from scepter.modules.annotator.doodle import DoodleAnnotator
from scepter.modules.annotator.gray import GrayAnnotator
from scepter.modules.annotator.hed import HedAnnotator
from scepter.modules.annotator.identity import IdentityAnnotator
from scepter.modules.annotator.informative_drawing import (
InfoDrawAnimeAnnotator, InfoDrawContourAnnotator,
InfoDrawOpenSketchAnnotator)
from scepter.modules.annotator.inpainting import InpaintingAnnotator
from scepter.modules.annotator.invert import InvertAnnotator
from scepter.modules.annotator.midas_op import MidasDetector
from scepter.modules.annotator.mlsd_op import MLSDdetector
from scepter.modules.annotator.openpose import OpenposeAnnotator
from scepter.modules.annotator.outpainting import OutpaintingAnnotator, OutpaintingResize
from scepter.modules.annotator.pidinet import PiDiAnnotator
from scepter.modules.annotator.segmentation import ESAMAnnotator
from scepter.modules.annotator.sketch import SketchAnnotator
from scepter.modules.annotator.lama import LamaAnnotator
from scepter.modules.annotator.mask_aug import MaskAugAnnotator, MaskDrawAnnotator, MaskLayoutAnnotator
from scepter.modules.annotator.raft import RAFTAnnotator, RAFTVisAnnotator
else:
_import_structure = {
'base_annotator': ['GeneralAnnotator'],
'canny': ['CannyAnnotator'],
'color': ['ColorAnnotator'],
'degradation': ['DegradationAnnotator'],
'doodle': ['DoodleAnnotator'],
'gray': ['GrayAnnotator'],
'hed': ['HedAnnotator'],
'identity': ['IdentityAnnotator'],
'informative_drawing': ['InfoDrawAnimeAnnotator',
'InfoDrawContourAnnotator',
'InfoDrawOpenSketchAnnotator'],
'inpainting': ['InpaintingAnnotator'],
'invert': ['InvertAnnotator'],
'midas_op': ['MidasDetector'],
'mlsd_op': ['MLSDdetector'],
'openpose': ['OpenposeAnnotator'],
'outpainting': ['OutpaintingAnnotator', 'OutpaintingResize'],
'pidinet': ['PiDiAnnotator'],
'segmentation': ['ESAMAnnotator'],
'sketch': ['SketchAnnotator'],
'lama': ['LamaAnnotator'],
'mask_aug': ['MaskAugAnnotator', 'MaskDrawAnnotator', 'MaskLayoutAnnotator'],
'raft': ['RAFTAnnotator', 'RAFTVisAnnotator'],
}
import sys
sys.modules[__name__] = LazyImportModule(
__name__,
globals()['__file__'],
_import_structure,
module_spec=__spec__,
extra_objects={},
)
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
+127
View File
@@ -0,0 +1,127 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import cv2
import numpy as np
import onnxruntime
def nms(boxes, scores, nms_thr):
"""Single class NMS implemented in Numpy."""
x1 = boxes[:, 0]
y1 = boxes[:, 1]
x2 = boxes[:, 2]
y2 = boxes[:, 3]
areas = (x2 - x1 + 1) * (y2 - y1 + 1)
order = scores.argsort()[::-1]
keep = []
while order.size > 0:
i = order[0]
keep.append(i)
xx1 = np.maximum(x1[i], x1[order[1:]])
yy1 = np.maximum(y1[i], y1[order[1:]])
xx2 = np.minimum(x2[i], x2[order[1:]])
yy2 = np.minimum(y2[i], y2[order[1:]])
w = np.maximum(0.0, xx2 - xx1 + 1)
h = np.maximum(0.0, yy2 - yy1 + 1)
inter = w * h
ovr = inter / (areas[i] + areas[order[1:]] - inter)
inds = np.where(ovr <= nms_thr)[0]
order = order[inds + 1]
return keep
def multiclass_nms(boxes, scores, nms_thr, score_thr):
"""Multiclass NMS implemented in Numpy. Class-aware version."""
final_dets = []
num_classes = scores.shape[1]
for cls_ind in range(num_classes):
cls_scores = scores[:, cls_ind]
valid_score_mask = cls_scores > score_thr
if valid_score_mask.sum() == 0:
continue
else:
valid_scores = cls_scores[valid_score_mask]
valid_boxes = boxes[valid_score_mask]
keep = nms(valid_boxes, valid_scores, nms_thr)
if len(keep) > 0:
cls_inds = np.ones((len(keep), 1)) * cls_ind
dets = np.concatenate(
[valid_boxes[keep], valid_scores[keep, None], cls_inds], 1
)
final_dets.append(dets)
if len(final_dets) == 0:
return None
return np.concatenate(final_dets, 0)
def demo_postprocess(outputs, img_size, p6=False):
grids = []
expanded_strides = []
strides = [8, 16, 32] if not p6 else [8, 16, 32, 64]
hsizes = [img_size[0] // stride for stride in strides]
wsizes = [img_size[1] // stride for stride in strides]
for hsize, wsize, stride in zip(hsizes, wsizes, strides):
xv, yv = np.meshgrid(np.arange(wsize), np.arange(hsize))
grid = np.stack((xv, yv), 2).reshape(1, -1, 2)
grids.append(grid)
shape = grid.shape[:2]
expanded_strides.append(np.full((*shape, 1), stride))
grids = np.concatenate(grids, 1)
expanded_strides = np.concatenate(expanded_strides, 1)
outputs[..., :2] = (outputs[..., :2] + grids) * expanded_strides
outputs[..., 2:4] = np.exp(outputs[..., 2:4]) * expanded_strides
return outputs
def preprocess(img, input_size, swap=(2, 0, 1)):
if len(img.shape) == 3:
padded_img = np.ones((input_size[0], input_size[1], 3), dtype=np.uint8) * 114
else:
padded_img = np.ones(input_size, dtype=np.uint8) * 114
r = min(input_size[0] / img.shape[0], input_size[1] / img.shape[1])
resized_img = cv2.resize(
img,
(int(img.shape[1] * r), int(img.shape[0] * r)),
interpolation=cv2.INTER_LINEAR,
).astype(np.uint8)
padded_img[: int(img.shape[0] * r), : int(img.shape[1] * r)] = resized_img
padded_img = padded_img.transpose(swap)
padded_img = np.ascontiguousarray(padded_img, dtype=np.float32)
return padded_img, r
def inference_detector(session, oriImg):
input_shape = (640,640)
img, ratio = preprocess(oriImg, input_shape)
ort_inputs = {session.get_inputs()[0].name: img[None, :, :, :]}
output = session.run(None, ort_inputs)
predictions = demo_postprocess(output[0], input_shape)[0]
boxes = predictions[:, :4]
scores = predictions[:, 4:5] * predictions[:, 5:]
boxes_xyxy = np.ones_like(boxes)
boxes_xyxy[:, 0] = boxes[:, 0] - boxes[:, 2]/2.
boxes_xyxy[:, 1] = boxes[:, 1] - boxes[:, 3]/2.
boxes_xyxy[:, 2] = boxes[:, 0] + boxes[:, 2]/2.
boxes_xyxy[:, 3] = boxes[:, 1] + boxes[:, 3]/2.
boxes_xyxy /= ratio
dets = multiclass_nms(boxes_xyxy, scores, nms_thr=0.45, score_thr=0.1)
if dets is not None:
final_boxes, final_scores, final_cls_inds = dets[:, :4], dets[:, 4], dets[:, 5]
isscore = final_scores>0.3
iscat = final_cls_inds == 0
isbbox = [ i and j for (i, j) in zip(isscore, iscat)]
final_boxes = final_boxes[isbbox]
else:
final_boxes = np.array([])
return final_boxes
@@ -0,0 +1,362 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from typing import List, Tuple
import cv2
import numpy as np
import onnxruntime as ort
def preprocess(
img: np.ndarray, out_bbox, input_size: Tuple[int, int] = (192, 256)
) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Do preprocessing for RTMPose model inference.
Args:
img (np.ndarray): Input image in shape.
input_size (tuple): Input image size in shape (w, h).
Returns:
tuple:
- resized_img (np.ndarray): Preprocessed image.
- center (np.ndarray): Center of image.
- scale (np.ndarray): Scale of image.
"""
# get shape of image
img_shape = img.shape[:2]
out_img, out_center, out_scale = [], [], []
if len(out_bbox) == 0:
out_bbox = [[0, 0, img_shape[1], img_shape[0]]]
for i in range(len(out_bbox)):
x0 = out_bbox[i][0]
y0 = out_bbox[i][1]
x1 = out_bbox[i][2]
y1 = out_bbox[i][3]
bbox = np.array([x0, y0, x1, y1])
# get center and scale
center, scale = bbox_xyxy2cs(bbox, padding=1.25)
# do affine transformation
resized_img, scale = top_down_affine(input_size, scale, center, img)
# normalize image
mean = np.array([123.675, 116.28, 103.53])
std = np.array([58.395, 57.12, 57.375])
resized_img = (resized_img - mean) / std
out_img.append(resized_img)
out_center.append(center)
out_scale.append(scale)
return out_img, out_center, out_scale
def inference(sess: ort.InferenceSession, img: np.ndarray) -> np.ndarray:
"""Inference RTMPose model.
Args:
sess (ort.InferenceSession): ONNXRuntime session.
img (np.ndarray): Input image in shape.
Returns:
outputs (np.ndarray): Output of RTMPose model.
"""
all_out = []
# build input
for i in range(len(img)):
input = [img[i].transpose(2, 0, 1)]
# build output
sess_input = {sess.get_inputs()[0].name: input}
sess_output = []
for out in sess.get_outputs():
sess_output.append(out.name)
# run model
outputs = sess.run(sess_output, sess_input)
all_out.append(outputs)
return all_out
def postprocess(outputs: List[np.ndarray],
model_input_size: Tuple[int, int],
center: Tuple[int, int],
scale: Tuple[int, int],
simcc_split_ratio: float = 2.0
) -> Tuple[np.ndarray, np.ndarray]:
"""Postprocess for RTMPose model output.
Args:
outputs (np.ndarray): Output of RTMPose model.
model_input_size (tuple): RTMPose model Input image size.
center (tuple): Center of bbox in shape (x, y).
scale (tuple): Scale of bbox in shape (w, h).
simcc_split_ratio (float): Split ratio of simcc.
Returns:
tuple:
- keypoints (np.ndarray): Rescaled keypoints.
- scores (np.ndarray): Model predict scores.
"""
all_key = []
all_score = []
for i in range(len(outputs)):
# use simcc to decode
simcc_x, simcc_y = outputs[i]
keypoints, scores = decode(simcc_x, simcc_y, simcc_split_ratio)
# rescale keypoints
keypoints = keypoints / model_input_size * scale[i] + center[i] - scale[i] / 2
all_key.append(keypoints[0])
all_score.append(scores[0])
return np.array(all_key), np.array(all_score)
def bbox_xyxy2cs(bbox: np.ndarray,
padding: float = 1.) -> Tuple[np.ndarray, np.ndarray]:
"""Transform the bbox format from (x,y,w,h) into (center, scale)
Args:
bbox (ndarray): Bounding box(es) in shape (4,) or (n, 4), formatted
as (left, top, right, bottom)
padding (float): BBox padding factor that will be multilied to scale.
Default: 1.0
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: Center (x, y) of the bbox in shape (2,) or
(n, 2)
- np.ndarray[float32]: Scale (w, h) of the bbox in shape (2,) or
(n, 2)
"""
# convert single bbox from (4, ) to (1, 4)
dim = bbox.ndim
if dim == 1:
bbox = bbox[None, :]
# get bbox center and scale
x1, y1, x2, y2 = np.hsplit(bbox, [1, 2, 3])
center = np.hstack([x1 + x2, y1 + y2]) * 0.5
scale = np.hstack([x2 - x1, y2 - y1]) * padding
if dim == 1:
center = center[0]
scale = scale[0]
return center, scale
def _fix_aspect_ratio(bbox_scale: np.ndarray,
aspect_ratio: float) -> np.ndarray:
"""Extend the scale to match the given aspect ratio.
Args:
scale (np.ndarray): The image scale (w, h) in shape (2, )
aspect_ratio (float): The ratio of ``w/h``
Returns:
np.ndarray: The reshaped image scale in (2, )
"""
w, h = np.hsplit(bbox_scale, [1])
bbox_scale = np.where(w > h * aspect_ratio,
np.hstack([w, w / aspect_ratio]),
np.hstack([h * aspect_ratio, h]))
return bbox_scale
def _rotate_point(pt: np.ndarray, angle_rad: float) -> np.ndarray:
"""Rotate a point by an angle.
Args:
pt (np.ndarray): 2D point coordinates (x, y) in shape (2, )
angle_rad (float): rotation angle in radian
Returns:
np.ndarray: Rotated point in shape (2, )
"""
sn, cs = np.sin(angle_rad), np.cos(angle_rad)
rot_mat = np.array([[cs, -sn], [sn, cs]])
return rot_mat @ pt
def _get_3rd_point(a: np.ndarray, b: np.ndarray) -> np.ndarray:
"""To calculate the affine matrix, three pairs of points are required. This
function is used to get the 3rd point, given 2D points a & b.
The 3rd point is defined by rotating vector `a - b` by 90 degrees
anticlockwise, using b as the rotation center.
Args:
a (np.ndarray): The 1st point (x,y) in shape (2, )
b (np.ndarray): The 2nd point (x,y) in shape (2, )
Returns:
np.ndarray: The 3rd point.
"""
direction = a - b
c = b + np.r_[-direction[1], direction[0]]
return c
def get_warp_matrix(center: np.ndarray,
scale: np.ndarray,
rot: float,
output_size: Tuple[int, int],
shift: Tuple[float, float] = (0., 0.),
inv: bool = False) -> np.ndarray:
"""Calculate the affine transformation matrix that can warp the bbox area
in the input image to the output size.
Args:
center (np.ndarray[2, ]): Center of the bounding box (x, y).
scale (np.ndarray[2, ]): Scale of the bounding box
wrt [width, height].
rot (float): Rotation angle (degree).
output_size (np.ndarray[2, ] | list(2,)): Size of the
destination heatmaps.
shift (0-100%): Shift translation ratio wrt the width/height.
Default (0., 0.).
inv (bool): Option to inverse the affine transform direction.
(inv=False: src->dst or inv=True: dst->src)
Returns:
np.ndarray: A 2x3 transformation matrix
"""
shift = np.array(shift)
src_w = scale[0]
dst_w = output_size[0]
dst_h = output_size[1]
# compute transformation matrix
rot_rad = np.deg2rad(rot)
src_dir = _rotate_point(np.array([0., src_w * -0.5]), rot_rad)
dst_dir = np.array([0., dst_w * -0.5])
# get four corners of the src rectangle in the original image
src = np.zeros((3, 2), dtype=np.float32)
src[0, :] = center + scale * shift
src[1, :] = center + src_dir + scale * shift
src[2, :] = _get_3rd_point(src[0, :], src[1, :])
# get four corners of the dst rectangle in the input image
dst = np.zeros((3, 2), dtype=np.float32)
dst[0, :] = [dst_w * 0.5, dst_h * 0.5]
dst[1, :] = np.array([dst_w * 0.5, dst_h * 0.5]) + dst_dir
dst[2, :] = _get_3rd_point(dst[0, :], dst[1, :])
if inv:
warp_mat = cv2.getAffineTransform(np.float32(dst), np.float32(src))
else:
warp_mat = cv2.getAffineTransform(np.float32(src), np.float32(dst))
return warp_mat
def top_down_affine(input_size: dict, bbox_scale: dict, bbox_center: dict,
img: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
"""Get the bbox image as the model input by affine transform.
Args:
input_size (dict): The input size of the model.
bbox_scale (dict): The bbox scale of the img.
bbox_center (dict): The bbox center of the img.
img (np.ndarray): The original image.
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: img after affine transform.
- np.ndarray[float32]: bbox scale after affine transform.
"""
w, h = input_size
warp_size = (int(w), int(h))
# reshape bbox to fixed aspect ratio
bbox_scale = _fix_aspect_ratio(bbox_scale, aspect_ratio=w / h)
# get the affine matrix
center = bbox_center
scale = bbox_scale
rot = 0
warp_mat = get_warp_matrix(center, scale, rot, output_size=(w, h))
# do affine transform
img = cv2.warpAffine(img, warp_mat, warp_size, flags=cv2.INTER_LINEAR)
return img, bbox_scale
def get_simcc_maximum(simcc_x: np.ndarray,
simcc_y: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
"""Get maximum response location and value from simcc representations.
Note:
instance number: N
num_keypoints: K
heatmap height: H
heatmap width: W
Args:
simcc_x (np.ndarray): x-axis SimCC in shape (K, Wx) or (N, K, Wx)
simcc_y (np.ndarray): y-axis SimCC in shape (K, Wy) or (N, K, Wy)
Returns:
tuple:
- locs (np.ndarray): locations of maximum heatmap responses in shape
(K, 2) or (N, K, 2)
- vals (np.ndarray): values of maximum heatmap responses in shape
(K,) or (N, K)
"""
N, K, Wx = simcc_x.shape
simcc_x = simcc_x.reshape(N * K, -1)
simcc_y = simcc_y.reshape(N * K, -1)
# get maximum value locations
x_locs = np.argmax(simcc_x, axis=1)
y_locs = np.argmax(simcc_y, axis=1)
locs = np.stack((x_locs, y_locs), axis=-1).astype(np.float32)
max_val_x = np.amax(simcc_x, axis=1)
max_val_y = np.amax(simcc_y, axis=1)
# get maximum value across x and y axis
mask = max_val_x > max_val_y
max_val_x[mask] = max_val_y[mask]
vals = max_val_x
locs[vals <= 0.] = -1
# reshape
locs = locs.reshape(N, K, 2)
vals = vals.reshape(N, K)
return locs, vals
def decode(simcc_x: np.ndarray, simcc_y: np.ndarray,
simcc_split_ratio) -> Tuple[np.ndarray, np.ndarray]:
"""Modulate simcc distribution with Gaussian.
Args:
simcc_x (np.ndarray[K, Wx]): model predicted simcc in x.
simcc_y (np.ndarray[K, Wy]): model predicted simcc in y.
simcc_split_ratio (int): The split ratio of simcc.
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: keypoints in shape (K, 2) or (n, K, 2)
- np.ndarray[float32]: scores in shape (K,) or (n, K)
"""
keypoints, scores = get_simcc_maximum(simcc_x, simcc_y)
keypoints /= simcc_split_ratio
return keypoints, scores
def inference_pose(session, out_bbox, oriImg):
h, w = session.get_inputs()[0].shape[2:]
model_input_size = (w, h)
resized_img, center, scale = preprocess(oriImg, out_bbox, model_input_size)
outputs = inference(session, resized_img)
keypoints, scores = postprocess(outputs, model_input_size, center, scale)
return keypoints, scores
+299
View File
@@ -0,0 +1,299 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import math
import numpy as np
import matplotlib
import cv2
eps = 0.01
def smart_resize(x, s):
Ht, Wt = s
if x.ndim == 2:
Ho, Wo = x.shape
Co = 1
else:
Ho, Wo, Co = x.shape
if Co == 3 or Co == 1:
k = float(Ht + Wt) / float(Ho + Wo)
return cv2.resize(x, (int(Wt), int(Ht)), interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4)
else:
return np.stack([smart_resize(x[:, :, i], s) for i in range(Co)], axis=2)
def smart_resize_k(x, fx, fy):
if x.ndim == 2:
Ho, Wo = x.shape
Co = 1
else:
Ho, Wo, Co = x.shape
Ht, Wt = Ho * fy, Wo * fx
if Co == 3 or Co == 1:
k = float(Ht + Wt) / float(Ho + Wo)
return cv2.resize(x, (int(Wt), int(Ht)), interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4)
else:
return np.stack([smart_resize_k(x[:, :, i], fx, fy) for i in range(Co)], axis=2)
def padRightDownCorner(img, stride, padValue):
h = img.shape[0]
w = img.shape[1]
pad = 4 * [None]
pad[0] = 0 # up
pad[1] = 0 # left
pad[2] = 0 if (h % stride == 0) else stride - (h % stride) # down
pad[3] = 0 if (w % stride == 0) else stride - (w % stride) # right
img_padded = img
pad_up = np.tile(img_padded[0:1, :, :]*0 + padValue, (pad[0], 1, 1))
img_padded = np.concatenate((pad_up, img_padded), axis=0)
pad_left = np.tile(img_padded[:, 0:1, :]*0 + padValue, (1, pad[1], 1))
img_padded = np.concatenate((pad_left, img_padded), axis=1)
pad_down = np.tile(img_padded[-2:-1, :, :]*0 + padValue, (pad[2], 1, 1))
img_padded = np.concatenate((img_padded, pad_down), axis=0)
pad_right = np.tile(img_padded[:, -2:-1, :]*0 + padValue, (1, pad[3], 1))
img_padded = np.concatenate((img_padded, pad_right), axis=1)
return img_padded, pad
def transfer(model, model_weights):
transfered_model_weights = {}
for weights_name in model.state_dict().keys():
transfered_model_weights[weights_name] = model_weights['.'.join(weights_name.split('.')[1:])]
return transfered_model_weights
def draw_bodypose(canvas, candidate, subset):
H, W, C = canvas.shape
candidate = np.array(candidate)
subset = np.array(subset)
stickwidth = 4
limbSeq = [[2, 3], [2, 6], [3, 4], [4, 5], [6, 7], [7, 8], [2, 9], [9, 10], \
[10, 11], [2, 12], [12, 13], [13, 14], [2, 1], [1, 15], [15, 17], \
[1, 16], [16, 18], [3, 17], [6, 18]]
colors = [[255, 0, 0], [255, 85, 0], [255, 170, 0], [255, 255, 0], [170, 255, 0], [85, 255, 0], [0, 255, 0], \
[0, 255, 85], [0, 255, 170], [0, 255, 255], [0, 170, 255], [0, 85, 255], [0, 0, 255], [85, 0, 255], \
[170, 0, 255], [255, 0, 255], [255, 0, 170], [255, 0, 85]]
for i in range(17):
for n in range(len(subset)):
index = subset[n][np.array(limbSeq[i]) - 1]
if -1 in index:
continue
Y = candidate[index.astype(int), 0] * float(W)
X = candidate[index.astype(int), 1] * float(H)
mX = np.mean(X)
mY = np.mean(Y)
length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5
angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1]))
polygon = cv2.ellipse2Poly((int(mY), int(mX)), (int(length / 2), stickwidth), int(angle), 0, 360, 1)
cv2.fillConvexPoly(canvas, polygon, colors[i])
canvas = (canvas * 0.6).astype(np.uint8)
for i in range(18):
for n in range(len(subset)):
index = int(subset[n][i])
if index == -1:
continue
x, y = candidate[index][0:2]
x = int(x * W)
y = int(y * H)
cv2.circle(canvas, (int(x), int(y)), 4, colors[i], thickness=-1)
return canvas
def draw_handpose(canvas, all_hand_peaks):
H, W, C = canvas.shape
edges = [[0, 1], [1, 2], [2, 3], [3, 4], [0, 5], [5, 6], [6, 7], [7, 8], [0, 9], [9, 10], \
[10, 11], [11, 12], [0, 13], [13, 14], [14, 15], [15, 16], [0, 17], [17, 18], [18, 19], [19, 20]]
for peaks in all_hand_peaks:
peaks = np.array(peaks)
for ie, e in enumerate(edges):
x1, y1 = peaks[e[0]]
x2, y2 = peaks[e[1]]
x1 = int(x1 * W)
y1 = int(y1 * H)
x2 = int(x2 * W)
y2 = int(y2 * H)
if x1 > eps and y1 > eps and x2 > eps and y2 > eps:
cv2.line(canvas, (x1, y1), (x2, y2), matplotlib.colors.hsv_to_rgb([ie / float(len(edges)), 1.0, 1.0]) * 255, thickness=2)
for i, keyponit in enumerate(peaks):
x, y = keyponit
x = int(x * W)
y = int(y * H)
if x > eps and y > eps:
cv2.circle(canvas, (x, y), 4, (0, 0, 255), thickness=-1)
return canvas
def draw_facepose(canvas, all_lmks):
H, W, C = canvas.shape
for lmks in all_lmks:
lmks = np.array(lmks)
for lmk in lmks:
x, y = lmk
x = int(x * W)
y = int(y * H)
if x > eps and y > eps:
cv2.circle(canvas, (x, y), 3, (255, 255, 255), thickness=-1)
return canvas
# detect hand according to body pose keypoints
# please refer to https://github.com/CMU-Perceptual-Computing-Lab/openpose/blob/master/src/openpose/hand/handDetector.cpp
def handDetect(candidate, subset, oriImg):
# right hand: wrist 4, elbow 3, shoulder 2
# left hand: wrist 7, elbow 6, shoulder 5
ratioWristElbow = 0.33
detect_result = []
image_height, image_width = oriImg.shape[0:2]
for person in subset.astype(int):
# if any of three not detected
has_left = np.sum(person[[5, 6, 7]] == -1) == 0
has_right = np.sum(person[[2, 3, 4]] == -1) == 0
if not (has_left or has_right):
continue
hands = []
#left hand
if has_left:
left_shoulder_index, left_elbow_index, left_wrist_index = person[[5, 6, 7]]
x1, y1 = candidate[left_shoulder_index][:2]
x2, y2 = candidate[left_elbow_index][:2]
x3, y3 = candidate[left_wrist_index][:2]
hands.append([x1, y1, x2, y2, x3, y3, True])
# right hand
if has_right:
right_shoulder_index, right_elbow_index, right_wrist_index = person[[2, 3, 4]]
x1, y1 = candidate[right_shoulder_index][:2]
x2, y2 = candidate[right_elbow_index][:2]
x3, y3 = candidate[right_wrist_index][:2]
hands.append([x1, y1, x2, y2, x3, y3, False])
for x1, y1, x2, y2, x3, y3, is_left in hands:
# pos_hand = pos_wrist + ratio * (pos_wrist - pos_elbox) = (1 + ratio) * pos_wrist - ratio * pos_elbox
# handRectangle.x = posePtr[wrist*3] + ratioWristElbow * (posePtr[wrist*3] - posePtr[elbow*3]);
# handRectangle.y = posePtr[wrist*3+1] + ratioWristElbow * (posePtr[wrist*3+1] - posePtr[elbow*3+1]);
# const auto distanceWristElbow = getDistance(poseKeypoints, person, wrist, elbow);
# const auto distanceElbowShoulder = getDistance(poseKeypoints, person, elbow, shoulder);
# handRectangle.width = 1.5f * fastMax(distanceWristElbow, 0.9f * distanceElbowShoulder);
x = x3 + ratioWristElbow * (x3 - x2)
y = y3 + ratioWristElbow * (y3 - y2)
distanceWristElbow = math.sqrt((x3 - x2) ** 2 + (y3 - y2) ** 2)
distanceElbowShoulder = math.sqrt((x2 - x1) ** 2 + (y2 - y1) ** 2)
width = 1.5 * max(distanceWristElbow, 0.9 * distanceElbowShoulder)
# x-y refers to the center --> offset to topLeft point
# handRectangle.x -= handRectangle.width / 2.f;
# handRectangle.y -= handRectangle.height / 2.f;
x -= width / 2
y -= width / 2 # width = height
# overflow the image
if x < 0: x = 0
if y < 0: y = 0
width1 = width
width2 = width
if x + width > image_width: width1 = image_width - x
if y + width > image_height: width2 = image_height - y
width = min(width1, width2)
# the max hand box value is 20 pixels
if width >= 20:
detect_result.append([int(x), int(y), int(width), is_left])
'''
return value: [[x, y, w, True if left hand else False]].
width=height since the network require squared input.
x, y is the coordinate of top left
'''
return detect_result
# Written by Lvmin
def faceDetect(candidate, subset, oriImg):
# left right eye ear 14 15 16 17
detect_result = []
image_height, image_width = oriImg.shape[0:2]
for person in subset.astype(int):
has_head = person[0] > -1
if not has_head:
continue
has_left_eye = person[14] > -1
has_right_eye = person[15] > -1
has_left_ear = person[16] > -1
has_right_ear = person[17] > -1
if not (has_left_eye or has_right_eye or has_left_ear or has_right_ear):
continue
head, left_eye, right_eye, left_ear, right_ear = person[[0, 14, 15, 16, 17]]
width = 0.0
x0, y0 = candidate[head][:2]
if has_left_eye:
x1, y1 = candidate[left_eye][:2]
d = max(abs(x0 - x1), abs(y0 - y1))
width = max(width, d * 3.0)
if has_right_eye:
x1, y1 = candidate[right_eye][:2]
d = max(abs(x0 - x1), abs(y0 - y1))
width = max(width, d * 3.0)
if has_left_ear:
x1, y1 = candidate[left_ear][:2]
d = max(abs(x0 - x1), abs(y0 - y1))
width = max(width, d * 1.5)
if has_right_ear:
x1, y1 = candidate[right_ear][:2]
d = max(abs(x0 - x1), abs(y0 - y1))
width = max(width, d * 1.5)
x, y = x0, y0
x -= width
y -= width
if x < 0:
x = 0
if y < 0:
y = 0
width1 = width * 2
width2 = width * 2
if x + width > image_width:
width1 = image_width - x
if y + width > image_height:
width2 = image_height - y
width = min(width1, width2)
if width >= 20:
detect_result.append([int(x), int(y), int(width)])
return detect_result
# get max index of 2d array
def npmax(array):
arrayindex = array.argmax(1)
arrayvalue = array.max(1)
i = arrayvalue.argmax()
j = arrayindex[i]
return i, j
@@ -0,0 +1,80 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import cv2
import numpy as np
import onnxruntime as ort
from .onnxdet import inference_detector
from .onnxpose import inference_pose
def HWC3(x):
assert x.dtype == np.uint8
if x.ndim == 2:
x = x[:, :, None]
assert x.ndim == 3
H, W, C = x.shape
assert C == 1 or C == 3 or C == 4
if C == 3:
return x
if C == 1:
return np.concatenate([x, x, x], axis=2)
if C == 4:
color = x[:, :, 0:3].astype(np.float32)
alpha = x[:, :, 3:4].astype(np.float32) / 255.0
y = color * alpha + 255.0 * (1.0 - alpha)
y = y.clip(0, 255).astype(np.uint8)
return y
def resize_image(input_image, resolution):
H, W, C = input_image.shape
H = float(H)
W = float(W)
k = float(resolution) / min(H, W)
H *= k
W *= k
H = int(np.round(H / 64.0)) * 64
W = int(np.round(W / 64.0)) * 64
img = cv2.resize(input_image, (W, H), interpolation=cv2.INTER_LANCZOS4 if k > 1 else cv2.INTER_AREA)
return img
class Wholebody:
def __init__(self, onnx_det, onnx_pose, device = 'cuda:0'):
providers = ['CPUExecutionProvider'
] if device == 'cpu' else ['CUDAExecutionProvider']
# onnx_det = 'annotator/ckpts/yolox_l.onnx'
# onnx_pose = 'annotator/ckpts/dw-ll_ucoco_384.onnx'
self.session_det = ort.InferenceSession(path_or_bytes=onnx_det, providers=providers)
self.session_pose = ort.InferenceSession(path_or_bytes=onnx_pose, providers=providers)
def __call__(self, ori_img):
det_result = inference_detector(self.session_det, ori_img)
keypoints, scores = inference_pose(self.session_pose, det_result, ori_img)
keypoints_info = np.concatenate(
(keypoints, scores[..., None]), axis=-1)
# compute neck joint
neck = np.mean(keypoints_info[:, [5, 6]], axis=1)
# neck score when visualizing pred
neck[:, 2:4] = np.logical_and(
keypoints_info[:, 5, 2:4] > 0.3,
keypoints_info[:, 6, 2:4] > 0.3).astype(int)
new_keypoints_info = np.insert(
keypoints_info, 17, neck, axis=1)
mmpose_idx = [
17, 6, 8, 10, 7, 9, 12, 14, 16, 13, 15, 2, 1, 4, 3
]
openpose_idx = [
1, 2, 3, 4, 6, 7, 8, 9, 10, 12, 13, 14, 15, 16, 17
]
new_keypoints_info[:, openpose_idx] = \
new_keypoints_info[:, mmpose_idx]
keypoints_info = new_keypoints_info
keypoints, scores = keypoints_info[
..., :2], keypoints_info[..., 2]
return keypoints, scores, det_result
+203
View File
@@ -0,0 +1,203 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
# Openpose
# Original from CMU https://github.com/CMU-Perceptual-Computing-Lab/openpose
# 2nd Edited by https://github.com/Hzzone/pytorch-openpose
# 3rd Edited by ControlNet
# 4th Edited by ControlNet (added face and correct hands)
# ``` requirements for cuda 12.1:
# onnxruntime==1.19
# onnxruntime-gpu==1.19
# ```
import os
import numpy as np
import torch
from PIL import Image
import cv2
from scepter.modules.annotator.base_annotator import BaseAnnotator
from scepter.modules.annotator.dwpose import util
from scepter.modules.annotator.dwpose.wholebody import (HWC3, Wholebody,
resize_image)
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
os.environ['KMP_DUPLICATE_LIB_OK'] = 'TRUE'
def draw_pose(pose, H, W, use_hand=False, use_body=False, use_face=False):
bodies = pose['bodies']
faces = pose['faces']
hands = pose['hands']
candidate = bodies['candidate']
subset = bodies['subset']
canvas = np.zeros(shape=(H, W, 3), dtype=np.uint8)
if use_body:
canvas = util.draw_bodypose(canvas, candidate, subset)
if use_hand:
canvas = util.draw_handpose(canvas, hands)
if use_face:
canvas = util.draw_facepose(canvas, faces)
return canvas
@ANNOTATORS.register_class()
class DWposeAnnotator(BaseAnnotator):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
with FS.get_from(cfg['DETECTION_MODEL'],
wait_finish=True) as onnx_det, FS.get_from(
cfg['POSE_MODEL'], wait_finish=True) as onnx_pose:
self.pose_estimation = Wholebody(onnx_det,
onnx_pose,
device=f'cuda:{we.device_id}')
self.resize_size = cfg.get('RESIZE_SIZE', 1024)
self.use_body = cfg.get('USE_BODY', True)
self.use_face = cfg.get('USE_FACE', True)
self.use_hand = cfg.get('USE_HAND', True)
@torch.no_grad()
@torch.inference_mode
def forward(self, image):
if isinstance(image, Image.Image):
image = np.array(image)
elif isinstance(image, torch.Tensor):
image = image.detach().cpu().numpy()
elif isinstance(image, np.ndarray):
image = image.copy()
else:
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
input_image = HWC3(image[..., ::-1])
return self.process(resize_image(input_image, self.resize_size),
image.shape[:2])
def process(self, ori_img, ori_shape):
ori_h, ori_w = ori_shape
ori_img = ori_img.copy()
H, W, C = ori_img.shape
with torch.no_grad():
candidate, subset, det_result = self.pose_estimation(ori_img)
nums, keys, locs = candidate.shape
candidate[..., 0] /= float(W)
candidate[..., 1] /= float(H)
body = candidate[:, :18].copy()
body = body.reshape(nums * 18, locs)
score = subset[:, :18]
for i in range(len(score)):
for j in range(len(score[i])):
if score[i][j] > 0.3:
score[i][j] = int(18 * i + j)
else:
score[i][j] = -1
un_visible = subset < 0.3
candidate[un_visible] = -1
foot = candidate[:, 18:24]
faces = candidate[:, 24:92]
hands = candidate[:, 92:113]
hands = np.vstack([hands, candidate[:, 113:]])
bodies = dict(candidate=body, subset=score)
pose = dict(bodies=bodies, hands=hands, faces=faces)
ret_data = {}
if self.use_body:
detected_map_body = draw_pose(pose, H, W, use_body=True)
detected_map_body = cv2.resize(
detected_map_body[..., ::-1], (ori_w, ori_h),
interpolation=cv2.INTER_LANCZOS4
if ori_h * ori_w > H * W else cv2.INTER_AREA)
ret_data['detected_map_body'] = detected_map_body
if self.use_face:
detected_map_face = draw_pose(pose, H, W, use_face=True)
detected_map_face = cv2.resize(
detected_map_face[..., ::-1], (ori_w, ori_h),
interpolation=cv2.INTER_LANCZOS4
if ori_h * ori_w > H * W else cv2.INTER_AREA)
ret_data['detected_map_face'] = detected_map_face
if self.use_body and self.use_face:
detected_map_bodyface = draw_pose(pose,
H,
W,
use_body=True,
use_face=True)
detected_map_bodyface = cv2.resize(
detected_map_bodyface[..., ::-1], (ori_w, ori_h),
interpolation=cv2.INTER_LANCZOS4
if ori_h * ori_w > H * W else cv2.INTER_AREA)
ret_data['detected_map_bodyface'] = detected_map_bodyface
if self.use_hand and self.use_body and self.use_face:
detected_map_handbodyface = draw_pose(pose,
H,
W,
use_hand=True,
use_body=True,
use_face=True)
detected_map_handbodyface = cv2.resize(
detected_map_handbodyface[..., ::-1], (ori_w, ori_h),
interpolation=cv2.INTER_LANCZOS4
if ori_h * ori_w > H * W else cv2.INTER_AREA)
ret_data[
'detected_map_handbodyface'] = detected_map_handbodyface
# convert_size
if det_result.shape[0] > 0:
w_ratio, h_ratio = ori_w / W, ori_h / H
det_result[..., ::2] *= h_ratio
det_result[..., 1::2] *= w_ratio
det_result = det_result.astype(np.int32)
# for det_tup in det_result:
# cv2.rectangle(detected_map, det_tup[2:].tolist(), det_tup[:2].tolist(), color=(255, 0, 0), thickness=3)
return ret_data, det_result
@ANNOTATORS.register_class()
class DWposeBodyAnnotator(DWposeAnnotator):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.use_body, self.use_face, self.use_hand = True, False, False
@torch.no_grad()
@torch.inference_mode
def forward(self, image):
ret_data, det_result = super().forward(image)
return ret_data['detected_map_body']
@ANNOTATORS.register_class()
class DWposeFaceAnnotator(DWposeAnnotator):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.use_body, self.use_face, self.use_hand = False, True, False
@torch.no_grad()
@torch.inference_mode
def forward(self, image):
ret_data, det_result = super().forward(image)
return ret_data['detected_map_face']
@ANNOTATORS.register_class()
class DWposeBodyFaceAnnotator(DWposeAnnotator):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.use_body, self.use_face, self.use_hand = True, True, False
@torch.no_grad()
@torch.inference_mode
def forward(self, image):
ret_data, det_result = super().forward(image)
return ret_data['detected_map_bodyface']
+63
View File
@@ -0,0 +1,63 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os
from abc import ABCMeta
import numpy as np
import torch
from PIL import Image
from scepter.modules.annotator.base_annotator import BaseAnnotator
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
@ANNOTATORS.register_class()
class FaceAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
from insightface.app import FaceAnalysis
local_path = FS.map_to_local(cfg.PRETRAINED_MODEL)[0]
local_model_path = os.path.join(local_path, 'models', cfg.MODEL_NAME)
FS.get_dir_to_local_dir(cfg.PRETRAINED_MODEL, local_model_path)
self.model = FaceAnalysis(name=cfg.MODEL_NAME, root=local_path, providers=['CUDAExecutionProvider', 'CPUExecutionProvider'])
self.model.prepare(ctx_id=we.device_id, det_size=(640, 640))
def forward(self, image=None):
if isinstance(image, Image.Image):
image = np.array(image)
elif isinstance(image, torch.Tensor):
image = image.detach().cpu().numpy()
elif isinstance(image, np.ndarray):
image = image.copy()
else:
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
# [dict_keys(['bbox', 'kps', 'det_score', 'landmark_3d_68', 'pose', 'landmark_2d_106', 'gender', 'age', 'embedding'])]
faces = self.model.get(image)
return faces
@ANNOTATORS.register_class()
class FaceMaskAnnotator(FaceAnnotator):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.multi_face = cfg.get('MULTI_FACE', True)
def forward(self, image=None):
faces = super().forward(image)
if len(faces) > 0:
if not self.multi_face:
faces = faces[:1]
mask = np.zeros_like(image[:, :, 0])
for face in faces:
x_min, y_min, x_max, y_max = face['bbox'].tolist()
mask[int(y_min): int(y_max) + 1, int(x_min): int(x_max) + 1] = 255
return mask
else:
return np.zeros_like(image[:, :, 0])
@@ -0,0 +1,58 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import random
import numpy as np
from scepter.modules.annotator.base_annotator import BaseAnnotator
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.utils.config import Config
@ANNOTATORS.register_class()
class FrameReferenceAnnotator(BaseAnnotator):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
# first / last / firstlast / random
self.ref_cfg = cfg.get('REF_CFG', [{"mode": "first", "proba": 0.1},
{"mode": "last", "proba": 0.1},
{"mode": "firstlast", "proba": 0.1},
{"mode": "random", "proba": 0.1}])
self.ref_num = cfg.get('REF_NUM', 1)
self.ref_cfg = Config.get_dict(self.ref_cfg) if isinstance(
self.ref_cfg, Config) else self.ref_cfg
self.ref_color = cfg.get('REF_COLOR', 127.5)
def forward(self, frames, ref_cfg=None, ref_num=None):
ref_cfg = ref_cfg if ref_cfg is not None else self.ref_cfg
ref_cfg = [ref_cfg] if not isinstance(ref_cfg, list) else ref_cfg
probas = [item['proba'] if 'proba' in item else 1.0 / len(ref_cfg) for item in ref_cfg]
sel_ref_cfg = random.choices(ref_cfg, weights=probas, k=1)[0]
mode = sel_ref_cfg['mode'] if 'mode' in sel_ref_cfg else 'original'
ref_num = int(ref_num) if ref_num is not None else self.ref_num
frame_num = len(frames)
frame_num_range = list(range(frame_num))
if mode == "first":
sel_idx = frame_num_range[:ref_num]
elif mode == "last":
sel_idx = frame_num_range[-ref_num:]
elif mode == "firstlast":
sel_idx = frame_num_range[:ref_num] + frame_num_range[-ref_num:]
elif mode == "random":
sel_idx = random.sample(frame_num_range, ref_num)
else:
raise NotImplementedError
out_frames, out_masks = [], []
for i in range(frame_num):
if i in sel_idx:
out_frame = frames[i]
out_mask = np.zeros_like(frames[i][:, :, 0])
else:
out_frame = np.ones_like(frames[i]) * self.ref_color
out_mask = np.ones_like(frames[i][:, :, 0]) * 255
out_frames.append(out_frame)
out_masks.append(out_mask)
return out_frames, out_masks
+1 -1
View File
@@ -114,7 +114,7 @@ class HedAnnotator(BaseAnnotator, metaclass=ABCMeta):
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
if pretrained_model:
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
self.netNetwork.load_state_dict(torch.load(local_path))
self.netNetwork.load_state_dict(torch.load(local_path, weights_only=True))
@torch.no_grad()
@torch.inference_mode()
@@ -120,7 +120,7 @@ class InfoDrawContourAnnotator(BaseAnnotator, metaclass=ABCMeta):
self.model = ContourInference(input_nc, output_nc, n_residual_blocks,
sigmoid)
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
self.model.load_state_dict(torch.load(local_path))
self.model.load_state_dict(torch.load(local_path, weights_only=True))
self.model = self.model.eval().requires_grad_(False).to(we.device_id)
@torch.no_grad()
+450
View File
@@ -0,0 +1,450 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import random
from abc import ABCMeta
from functools import partial
import numpy as np
import torch
from PIL import Image, ImageDraw
import cv2
from scepter.modules.annotator.base_annotator import BaseAnnotator
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.utils.config import Config
from scepter.modules.utils.file_system import FS
from scipy import ndimage
from scipy.spatial import ConvexHull
from skimage.draw import polygon
@ANNOTATORS.register_class()
class MaskDrawAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.task_type = cfg.get('TASK_TYPE', 'input_box')
def forward(self, mask=None, image=None, input_box=None, task_type=None):
task_type = task_type if task_type is not None else self.task_type
if mask is not None:
if isinstance(mask, Image.Image):
mask = np.array(mask)
elif isinstance(mask, torch.Tensor):
mask = mask.detach().cpu().numpy()
elif isinstance(mask, np.ndarray):
mask = mask.copy()
else:
raise f'Unsurpport datatype{type(mask)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
if image is not None:
if isinstance(image, Image.Image):
image = np.array(image)
elif isinstance(image, torch.Tensor):
image = image.detach().cpu().numpy()
elif isinstance(image, np.ndarray):
image = image.copy()
else:
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
mask_shape = mask.shape
if task_type == 'mask_point':
scribble = mask.transpose(1, 0)
labeled_array, num_features = ndimage.label(scribble >= 255)
centers = ndimage.center_of_mass(scribble, labeled_array,
range(1, num_features + 1))
centers = np.array(centers)
out_mask = np.zeros(mask_shape, dtype=np.uint8)
hull = ConvexHull(centers)
hull_vertices = centers[hull.vertices]
rr, cc = polygon(hull_vertices[:, 1], hull_vertices[:, 0],
mask_shape)
out_mask[rr, cc] = 255
elif task_type == 'mask_box':
scribble = mask.transpose(1, 0)
labeled_array, num_features = ndimage.label(scribble >= 255)
centers = ndimage.center_of_mass(scribble, labeled_array,
range(1, num_features + 1))
centers = np.array(centers)
# (x1, y1, x2, y2)
x_min = centers[:, 0].min()
x_max = centers[:, 0].max()
y_min = centers[:, 1].min()
y_max = centers[:, 1].max()
out_mask = np.zeros(mask_shape, dtype=np.uint8)
out_mask[int(y_min):int(y_max) + 1,
int(x_min):int(x_max) + 1] = 255
if image is not None:
out_image = image[int(y_min):int(y_max) + 1,
int(x_min):int(x_max) + 1]
elif task_type == 'input_box':
if isinstance(input_box, list):
input_box = np.array(input_box)
x_min, y_min, x_max, y_max = input_box
out_mask = np.zeros(mask_shape, dtype=np.uint8)
out_mask[int(y_min):int(y_max) + 1,
int(x_min):int(x_max) + 1] = 255
if image is not None:
out_image = image[int(y_min):int(y_max) + 1,
int(x_min):int(x_max) + 1]
elif task_type == 'mask':
out_mask = mask
else:
raise NotImplementedError
if image is not None:
return out_image, out_mask
else:
return out_mask
@ANNOTATORS.register_class()
class MaskAugAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
# original / original_expand / hull / hull_expand / bbox / bbox_expand
self.mask_cfg = cfg.get('MASK_CFG', [{
'mode': 'original',
'proba': 0.1
}, {
'mode': 'original_expand',
'proba': 0.1
}, {
'mode': 'hull',
'proba': 0.1
}, {
'mode': 'hull_expand',
'proba': 0.1,
'kwargs': {
'expand_rate': 0.2
}
}, {
'mode': 'bbox',
'proba': 0.1
}, {
'mode': 'bbox_expand',
'proba': 0.1,
'kwargs': {
'min_expand_rate': 0.2,
'max_expand_rate': 0.5
}
}])
self.mask_cfg = Config.get_dict(self.mask_cfg) if isinstance(
self.mask_cfg, Config) else self.mask_cfg
def forward(self, mask, mask_cfg=None):
mask_cfg = mask_cfg if mask_cfg is not None else self.mask_cfg
if not isinstance(mask, list):
is_batch = False
masks = [mask]
else:
is_batch = True
masks = mask
mask_func = self.get_mask_func(mask_cfg)
# print(mask_func)
aug_masks = []
for submask in masks:
mask = self.get_mask(submask)
valid, large, h, w, bbox = self.get_mask_info(mask)
# print(valid, large, h, w, bbox)
if valid:
mask = mask_func(mask, bbox, h, w)
else:
mask = mask.astype(np.uint8)
aug_masks.append(mask)
return aug_masks if is_batch else aug_masks[0]
def get_mask(self, mask):
if isinstance(mask, Image.Image):
mask = np.array(mask)
elif isinstance(mask, torch.Tensor):
mask = mask.detach().cpu().numpy()
elif isinstance(mask, np.ndarray):
mask = mask.copy()
else:
raise f'Unsurpport datatype{type(mask)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
return mask
def get_mask_info(self, mask):
h, w = mask.shape
locs = mask.nonzero()
valid = True
if len(locs) < 1 or locs[0].shape[0] < 1 or locs[1].shape[0] < 1:
valid = False
return valid, False, h, w, [0, 0, 0, 0]
left, right = np.min(locs[1]), np.max(locs[1])
top, bottom = np.min(locs[0]), np.max(locs[0])
bbox = [left, top, right, bottom]
large = False
if (right - left + 1) * (bottom - top + 1) > 0.9 * h * w:
large = True
return valid, large, h, w, bbox
def get_expand_params(self, mask_kwargs):
if 'expand_rate' in mask_kwargs:
expand_rate = mask_kwargs['expand_rate']
elif 'min_expand_rate' in mask_kwargs and 'max_expand_rate' in mask_kwargs:
expand_rate = random.uniform(mask_kwargs['min_expand_rate'],
mask_kwargs['max_expand_rate'])
else:
expand_rate = 0.3
if 'expand_iters' in mask_kwargs:
expand_iters = mask_kwargs['expand_iters']
else:
expand_iters = random.randint(1, 10)
if 'expand_lrtp' in mask_kwargs:
expand_lrtp = mask_kwargs['expand_lrtp']
else:
expand_lrtp = [
random.random(),
random.random(),
random.random(),
random.random()
]
return expand_rate, expand_iters, expand_lrtp
def get_mask_func(self, mask_cfg):
if not isinstance(mask_cfg, list):
mask_cfg = [mask_cfg]
probas = [
item['proba'] if 'proba' in item else 1.0 / len(mask_cfg)
for item in mask_cfg
]
sel_mask_cfg = random.choices(mask_cfg, weights=probas, k=1)[0]
mode = sel_mask_cfg['mode'] if 'mode' in sel_mask_cfg else 'original'
mask_kwargs = sel_mask_cfg[
'kwargs'] if 'kwargs' in sel_mask_cfg else {}
if mode == 'random':
mode = random.choice([
'original', 'original_expand', 'hull', 'hull_expand', 'bbox',
'bbox_expand'
])
if mode == 'original':
mask_func = partial(self.generate_mask)
elif mode == 'original_expand':
expand_rate, expand_iters, expand_lrtp = self.get_expand_params(
mask_kwargs)
mask_func = partial(self.generate_mask,
expand_rate=expand_rate,
expand_iters=expand_iters,
expand_lrtp=expand_lrtp)
elif mode == 'hull':
clockwise = random.choice([
True, False
]) if 'clockwise' not in mask_kwargs else mask_kwargs['clockwise']
mask_func = partial(self.generate_hull_mask, clockwise=clockwise)
elif mode == 'hull_expand':
expand_rate, expand_iters, expand_lrtp = self.get_expand_params(
mask_kwargs)
clockwise = random.choice([
True, False
]) if 'clockwise' not in mask_kwargs else mask_kwargs['clockwise']
mask_func = partial(self.generate_hull_mask,
clockwise=clockwise,
expand_rate=expand_rate,
expand_iters=expand_iters,
expand_lrtp=expand_lrtp)
elif mode == 'bbox':
mask_func = partial(self.generate_bbox_mask)
elif mode == 'bbox_expand':
expand_rate, expand_iters, expand_lrtp = self.get_expand_params(
mask_kwargs)
mask_func = partial(self.generate_bbox_mask,
expand_rate=expand_rate,
expand_iters=expand_iters,
expand_lrtp=expand_lrtp)
else:
raise NotImplementedError
return mask_func
def generate_mask(self,
mask,
bbox,
h,
w,
expand_rate=None,
expand_iters=None,
expand_lrtp=None):
bin_mask = mask.astype(np.uint8)
if expand_rate:
bin_mask = self.rand_expand_mask(bin_mask, bbox, h, w, expand_rate,
expand_iters, expand_lrtp)
return bin_mask
@staticmethod
def rand_expand_mask(mask,
bbox,
h,
w,
expand_rate=None,
expand_iters=None,
expand_lrtp=None):
expand_rate = 0.3 if expand_rate is None else expand_rate
expand_iters = random.randint(
1, 10) if expand_iters is None else expand_iters
expand_lrtp = [
random.random(),
random.random(),
random.random(),
random.random()
] if expand_lrtp is None else expand_lrtp
# print('iters', expand_iters, 'expand_rate', expand_rate, 'expand_lrtp', expand_lrtp)
# mask = np.squeeze(mask)
left, top, right, bottom = bbox
# mask expansion
box_w = (right - left + 1) * expand_rate
box_h = (bottom - top + 1) * expand_rate
left_, right_ = int(
expand_lrtp[0] * min(box_w, left / 2) / expand_iters), int(
expand_lrtp[1] * min(box_w, (w - right) / 2) / expand_iters)
top_, bottom_ = int(
expand_lrtp[2] * min(box_h, top / 2) / expand_iters), int(
expand_lrtp[3] * min(box_h, (h - bottom) / 2) / expand_iters)
kernel_size = max(left_, right_, top_, bottom_)
if kernel_size > 0:
kernel = np.zeros((kernel_size * 2, kernel_size * 2),
dtype=np.uint8)
new_left, new_right = kernel_size - right_, kernel_size + left_
new_top, new_bottom = kernel_size - bottom_, kernel_size + top_
kernel[new_top:new_bottom + 1, new_left:new_right + 1] = 1
mask = mask.astype(np.uint8)
mask = cv2.dilate(mask, kernel,
iterations=expand_iters).astype(np.uint8)
# mask = new_mask - (mask / 2).astype(np.uint8)
# mask = np.expand_dims(mask, axis=-1)
return mask
@staticmethod
def _convexhull(image, clockwise):
# print('clockwise', clockwise)
contours, hierarchy = cv2.findContours(image, 2, 1)
cnt = np.concatenate(contours) # merge all regions
hull = cv2.convexHull(cnt, clockwise=clockwise)
hull = np.squeeze(hull, axis=1).astype(np.float32).tolist()
hull = [tuple(x) for x in hull]
return hull # b, 1, 2
def generate_hull_mask(self,
mask,
bbox,
h,
w,
clockwise=None,
expand_rate=None,
expand_iters=None,
expand_lrtp=None):
clockwise = random.choice([True, False
]) if clockwise is None else clockwise
hull = self._convexhull(mask, clockwise)
mask_img = Image.new('L', (w, h), 0)
pt_list = hull
mask_img_draw = ImageDraw.Draw(mask_img)
mask_img_draw.polygon(pt_list, fill=255)
bin_mask = np.array(mask_img).astype(np.uint8)
if expand_rate:
bin_mask = self.rand_expand_mask(bin_mask, bbox, h, w, expand_rate,
expand_iters, expand_lrtp)
return bin_mask
def generate_bbox_mask(self,
mask,
bbox,
h,
w,
expand_rate=None,
expand_iters=None,
expand_lrtp=None):
left, top, right, bottom = bbox
bin_mask = np.zeros((h, w), dtype=np.uint8)
bin_mask[top:bottom + 1, left:right + 1] = 255
if expand_rate:
bin_mask = self.rand_expand_mask(bin_mask, bbox, h, w, expand_rate,
expand_iters, expand_lrtp)
return bin_mask
@ANNOTATORS.register_class()
class MaskLayoutAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
ram_tag_color = cfg.get('RAM_TAG_COLOR', None)
default_color = cfg.get('DEFAULT_COLOR', [0, 0, 0])
self.use_aug = cfg.get('USE_AUG', False)
self.color_dict = {'default': tuple(default_color)}
if ram_tag_color is not None:
with FS.get_object(ram_tag_color) as object:
lines = object.decode('utf-8').strip().split('\n')
lines = [id_name_color.split('#;#') for id_name_color in lines]
self.color_dict.update({
id_name_color[1]: tuple(eval(id_name_color[2]))
for id_name_color in lines
})
if self.use_aug:
mask_aug_dict = {'NAME': 'MaskAugAnnotator'}
mask_aug_cfg = Config(cfg_dict=mask_aug_dict, load=False)
self.mask_aug_anno = ANNOTATORS.build(mask_aug_cfg)
def find_contours(self, mask):
# @mask: gray cv2 image
# contours, hier = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE)
contours, hier = cv2.findContours(mask, cv2.RETR_EXTERNAL,
cv2.CHAIN_APPROX_SIMPLE)
return contours
def draw_contours(self, canvas, contour, color):
canvas = np.ascontiguousarray(canvas, dtype=np.uint8)
canvas = cv2.drawContours(canvas, contour, -1, color, thickness=3)
return canvas
def get_mask(self, mask):
if isinstance(mask, Image.Image):
mask = np.array(mask)
elif isinstance(mask, torch.Tensor):
mask = mask.detach().cpu().numpy()
elif isinstance(mask, np.ndarray):
mask = mask.copy()
else:
raise f'Unsurpport datatype{type(mask)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
return mask
def forward(self, mask=None, color=None, label=None, mask_cfg=None):
if not isinstance(mask, list):
is_batch = False
mask = [mask]
else:
is_batch = True
if label is not None and label in self.color_dict:
color = self.color_dict[label]
elif color is not None:
color = color
else:
color = self.color_dict['default']
ret_data = []
for sub_mask in mask:
sub_mask = self.get_mask(sub_mask)
if self.use_aug:
sub_mask = self.mask_aug_anno(sub_mask, mask_cfg)
canvas = np.ones((sub_mask.shape[0], sub_mask.shape[1], 3)) * 255
contour = self.find_contours(sub_mask)
frame = self.draw_contours(canvas, contour, color)
ret_data.append(frame)
if is_batch:
return ret_data
else:
return ret_data[0]
@@ -10,7 +10,7 @@ class BaseModel(torch.nn.Module):
Args:
path (str): file path
"""
parameters = torch.load(path, map_location=torch.device('cpu'))
parameters = torch.load(path, map_location=torch.device('cpu'), weights_only=True)
if 'optimizer' in parameters:
parameters = parameters['model']
+1 -1
View File
@@ -29,7 +29,7 @@ class MLSDdetector(BaseAnnotator, metaclass=ABCMeta):
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
if pretrained_model:
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
model.load_state_dict(torch.load(local_path), strict=True)
model.load_state_dict(torch.load(local_path, weights_only=True), strict=True)
self.model = model.eval()
self.thr_v = cfg.get('THR_V', 0.1)
self.thr_d = cfg.get('THR_D', 0.1)
+2 -2
View File
@@ -423,7 +423,7 @@ class Hand(object):
self.model = handpose_model()
if torch.cuda.is_available():
self.model = self.model.to(device)
model_dict = transfer(self.model, torch.load(model_path))
model_dict = transfer(self.model, torch.load(model_path, weights_only=True))
self.model.load_state_dict(model_dict)
self.model.eval()
self.device = device
@@ -503,7 +503,7 @@ class Body(object):
self.model = bodypose_model()
if torch.cuda.is_available():
self.model = self.model.to(device)
model_dict = transfer(self.model, torch.load(model_path))
model_dict = transfer(self.model, torch.load(model_path, weights_only=True))
self.model.load_state_dict(model_dict)
self.model.eval()
self.device = device
+8 -2
View File
@@ -98,9 +98,15 @@ class OutpaintingAnnotator(BaseAnnotator, metaclass=ABCMeta):
draw.rectangle(
(left + (self.mask_blur * 2 if left > 0 else 0), up +
(self.mask_blur * 2 if up > 0 else 0), mask.width - right -
(self.mask_blur * 2 if right > 0 else 0), mask.height - down -
(self.mask_blur * 2 if down > 0 else 0)),
(self.mask_blur * 2 if right > 0 else 0) - 1, mask.height - down -
(self.mask_blur * 2 if down > 0 else 0) - 1),
fill='black')
# draw.rectangle(
# (left + (self.mask_blur * 2 if left > 0 else 0), up +
# (self.mask_blur * 2 if up > 0 else 0), left + src_width -
# (self.mask_blur * 2 if right > 0 else 0), up + src_height -
# (self.mask_blur * 2 if down > 0 else 0)),
# fill='black')
else:
bbox = self.get_box(np.array(mask))
if bbox is None:
+1 -1
View File
@@ -882,7 +882,7 @@ class PiDiAnnotator(BaseAnnotator, metaclass=ABCMeta):
if pretrained_model:
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
state = torch.load(local_path,
map_location='cpu')['state_dict']
map_location='cpu', weights_only=True)['state_dict']
if vanilla_cnn:
state = convert_pidinet(state, 'carv4')
state = {
+62
View File
@@ -0,0 +1,62 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import torch
import random
import numpy as np
import argparse
from scepter.modules.annotator.base_annotator import BaseAnnotator
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.utils.config import Config
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
try:
from raft import RAFT
from raft.utils.utils import InputPadder
from raft.utils import flow_viz
except:
import warnings
warnings.warn("ignore raft import, please pip install raft.")
@ANNOTATORS.register_class()
class RAFTAnnotator(BaseAnnotator):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
params = {
"small": False,
"mixed_precision": False,
"alternate_corr": False
}
params = argparse.Namespace(**params)
model = RAFT(params)
if cfg.PRETRAINED_MODEL is not None:
with FS.get_from(cfg.PRETRAINED_MODEL,
wait_finish=True) as local_path:
model.load_state_dict({k.replace('module.', ''): v for k, v in torch.load(local_path, map_location="cpu", weights_only=True).items()})
self.model = model.to(we.device_id).eval()
def forward(self, frames):
# frames / RGB
frames = [torch.from_numpy(frame.astype(np.uint8)).permute(2, 0, 1).float()[None].to(we.device_id) for frame in frames]
flow_up_list, flow_up_vis_list = [], []
with torch.no_grad():
for i, (image1, image2) in enumerate(zip(frames[:-1], frames[1:])):
padder = InputPadder(image1.shape)
image1, image2 = padder.pad(image1, image2)
flow_low, flow_up = self.model(image1, image2, iters=20, test_mode=True)
flow_up = flow_up[0].permute(1, 2, 0).cpu().numpy()
flow_up_vis = flow_viz.flow_to_image(flow_up)
flow_up_list.append(flow_up)
flow_up_vis_list.append(flow_up_vis)
return flow_up_list, flow_up_vis_list # RGB
@ANNOTATORS.register_class()
class RAFTVisAnnotator(RAFTAnnotator):
def forward(self, frames):
flow_up_list, flow_up_vis_list = super().forward(frames)
return flow_up_vis_list[:1] + flow_up_vis_list
@@ -0,0 +1,95 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import random
from abc import ABCMeta
import cv2
import numpy as np
import torch
from PIL import Image
from scepter.modules.annotator.base_annotator import BaseAnnotator
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.utils.config import Config, dict_to_yaml
@ANNOTATORS.register_class()
class RegionCanvasAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.scale_range = cfg.get('SCALE_RANGE', [0.75, 1.0])
self.canvas_value = cfg.get('CANVAS_VALUE', 255)
self.use_resize = cfg.get('USE_RESIZE', True)
self.use_canvas = cfg.get('USE_CANVAS', True)
self.use_aug = cfg.get('USE_AUG', False)
if self.use_aug:
mask_aug_dict = {'NAME': 'MaskAugAnnotator'}
mask_aug_cfg = Config(cfg_dict=mask_aug_dict, load=False)
self.mask_aug_anno = ANNOTATORS.build(mask_aug_cfg)
def forward(self,
image,
mask,
mask_cfg=None):
if isinstance(image, Image.Image):
image = np.array(image)
elif isinstance(image, torch.Tensor):
image = image.detach().cpu().numpy()
elif isinstance(image, np.ndarray):
image = image
else:
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
mask = np.array(mask).astype(np.uint8)
image_h, image_w = image.shape[:2]
if self.use_aug:
mask = self.mask_aug_anno(mask, mask_cfg)
# get region with white bg
image[np.array(mask) == 0] = self.canvas_value
x, y, w, h = cv2.boundingRect(mask)
region_crop = image[y:y + h, x:x + w]
if self.use_resize:
# resize region
scale_min, scale_max = self.scale_range
scale_factor = random.uniform(scale_min, scale_max)
new_w, new_h = int(image_w * scale_factor), int(image_h * scale_factor)
obj_scale_factor = min(new_w/w, new_h/h)
new_w = int(w * obj_scale_factor)
new_h = int(h * obj_scale_factor)
region_crop_resized = cv2.resize(region_crop, (new_w, new_h), interpolation=cv2.INTER_AREA)
else:
region_crop_resized = region_crop
if self.use_canvas:
# plot region into canvas
new_canvas = np.ones_like(image) * self.canvas_value
max_x = max(0, image_w - new_w)
max_y = max(0, image_h - new_h)
new_x = random.randint(0, max_x)
new_y = random.randint(0, max_y)
new_canvas[new_y:new_y + new_h, new_x:new_x + new_w] = region_crop_resized
else:
new_canvas = region_crop_resized
return new_canvas
@staticmethod
def get_config_template():
return dict_to_yaml('ANNOTATORS',
__class__.__name__,
RegionCanvasAnnotator.para_dict,
set_name=True)
@ANNOTATORS.register_class()
class RegionCanvasCropAnnotator(RegionCanvasAnnotator):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.use_resize, self.use_canvas = False, False
+5 -1
View File
@@ -10,7 +10,11 @@ import torchvision.transforms as T
from PIL import Image
from pycocotools import mask as mask_utils
from scipy import ndimage
from sklearn.cluster import KMeans
try:
from sklearn.cluster import KMeans
except:
import warnings
warnings.warn("ignore sklearn import, please pip install scikit-learn.")
from torchvision.ops.boxes import batched_nms
from scepter.modules.annotator.base_annotator import BaseAnnotator
+1 -1
View File
@@ -86,7 +86,7 @@ class SketchAnnotator(BaseAnnotator, metaclass=ABCMeta):
std=0.0858381272736797).eval()
if pretrained_model:
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
state = torch.load(local_path, map_location='cpu')
state = torch.load(local_path, map_location='cpu', weights_only=True)
self.model.load_state_dict(state)
@torch.no_grad()
@@ -0,0 +1,153 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from abc import ABCMeta
import numpy as np
import torch
from PIL import Image
from scipy import ndimage
try:
from sklearn.cluster import KMeans
except:
import warnings
warnings.warn("ignore sklearn import, please pip install scikit-learn.")
from scepter.modules.annotator.base_annotator import BaseAnnotator
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.file_system import FS
import pycocotools.mask as mask_utils
def single_mask_to_rle(mask):
rle = mask_utils.encode(np.array(mask[:, :, None], order="F", dtype="uint8"))[0]
rle["counts"] = rle["counts"].decode("utf-8")
return rle
def single_rle_to_mask(rle):
mask = np.array(mask_utils.decode(rle)).astype(np.uint8)
return mask
def single_mask_to_xyxy(mask):
bbox = np.zeros((4), dtype=int)
rows, cols = np.where(np.array(mask))
if len(rows) > 0 and len(cols) > 0:
x_min, x_max = np.min(cols), np.max(cols)
y_min, y_max = np.min(rows), np.max(rows)
bbox[:] = [x_min, y_min, x_max, y_max]
return bbox.tolist()
@ANNOTATORS.register_class()
class SAM2DrawVideoAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.task_type = cfg.get('TASK_TYPE', 'input_box')
from sam2.build_sam import build_sam2_video_predictor
config_path = FS.get_from(cfg.CONFIG_PATH, local_path=cfg.CONFIG_LOCAL_PATH, wait_finish=True)
pretrained_model = FS.get_from(cfg.PRETRAINED_MODEL, wait_finish=True)
self.video_predictor = build_sam2_video_predictor(config_path, pretrained_model, fill_hole_area=0)
def forward(self,
video,
input_box=None,
mask=None,
task_type=None):
task_type = task_type if task_type is not None else self.task_type
if mask is not None:
if isinstance(mask, Image.Image):
mask = np.array(mask)
elif isinstance(mask, torch.Tensor):
mask = mask.detach().cpu().numpy()
elif isinstance(mask, np.ndarray):
mask = mask.copy()
else:
raise f'Unsurpport datatype{type(mask)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
if task_type == 'mask_point':
if len(mask.shape) == 3:
scribble = mask.transpose(2, 1, 0)[0]
else:
scribble = mask.transpose(1, 0) # (H, W) -> (W, H)
labeled_array, num_features = ndimage.label(scribble >= 255)
centers = ndimage.center_of_mass(scribble, labeled_array,
range(1, num_features + 1))
point_coords = np.array(centers)
point_labels = np.array([1] * len(centers))
sample = {
'points': point_coords,
'labels': point_labels
}
elif task_type == 'mask_box':
if len(mask.shape) == 3:
scribble = mask.transpose(2, 1, 0)[0]
else:
scribble = mask.transpose(1, 0) # (H, W) -> (W, H)
labeled_array, num_features = ndimage.label(scribble >= 255)
centers = ndimage.center_of_mass(scribble, labeled_array,
range(1, num_features + 1))
centers = np.array(centers)
# (x1, y1, x2, y2)
x_min = centers[:, 0].min()
x_max = centers[:, 0].max()
y_min = centers[:, 1].min()
y_max = centers[:, 1].max()
bbox = np.array([x_min, y_min, x_max, y_max])
sample = {'box': bbox}
elif task_type == 'input_box':
if isinstance(input_box, list):
input_box = np.array(input_box)
sample = {'box': input_box}
elif task_type == 'mask':
sample = {'mask': mask}
else:
raise NotImplementedError
ann_frame_idx = 0
object_id = 0
with (torch.inference_mode(), torch.autocast("cuda", dtype=torch.bfloat16)):
inference_state = self.video_predictor.init_state(video_path=video)
if task_type in ['mask_point', 'mask_box', 'input_box']:
_, out_obj_ids, out_mask_logits = self.video_predictor.add_new_points_or_box(
inference_state=inference_state,
frame_idx=ann_frame_idx,
obj_id=object_id,
**sample
)
elif task_type in ['mask']:
_, out_obj_ids, out_mask_logits = self.video_predictor.add_new_mask(
inference_state=inference_state,
frame_idx=ann_frame_idx,
obj_id=object_id,
**sample
)
else:
raise NotImplementedError
video_segments = {} # video_segments contains the per-frame segmentation results
for out_frame_idx, out_obj_ids, out_mask_logits in self.video_predictor.propagate_in_video(inference_state):
frame_segments = {}
for i, out_obj_id in enumerate(out_obj_ids):
mask = (out_mask_logits[i] > 0.0).cpu().numpy().squeeze(0)
frame_segments[out_obj_id] = {
"mask": single_mask_to_rle(mask),
"mask_area": int(mask.sum()),
"mask_box": single_mask_to_xyxy(mask),
}
video_segments[out_frame_idx] = frame_segments
ret_data = {
"annotations": video_segments
}
return ret_data
@staticmethod
def get_config_template():
return dict_to_yaml('ANNOTATORS',
__class__.__name__,
SAM2DrawVideoAnnotator.para_dict,
set_name=True)
+18 -1
View File
@@ -1,4 +1,21 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from typing import TYPE_CHECKING
from scepter.modules.utils.import_utils import LazyImportModule
from scepter.modules.data import dataset, sampler
if TYPE_CHECKING:
from scepter.modules.data import dataset, sampler
else:
_import_structure = {
'data': ['dataset', 'sampler']
}
import sys
sys.modules[__name__] = LazyImportModule(
__name__,
globals()['__file__'],
_import_structure,
module_spec=__spec__,
extra_objects={},
)
+32 -9
View File
@@ -1,12 +1,35 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from typing import TYPE_CHECKING
from scepter.modules.utils.import_utils import LazyImportModule
from scepter.modules.data.dataset.base_dataset import BaseDataset
from scepter.modules.data.dataset.dataset import (Image2ImageDataset,
ImageClassifyPublicDataset,
ImageTextPairDataset,
Text2ImageDataset)
from scepter.modules.data.dataset.ms_dataset import (
ImageTextPairFolderDataset, ImageTextPairMSDataset,
ImageTextPairMSDatasetForACE)
from scepter.modules.data.dataset.registry import DATASETS
if TYPE_CHECKING:
from scepter.modules.data.dataset.base_dataset import BaseDataset
from scepter.modules.data.dataset.dataset import (Image2ImageDataset,
ImageClassifyPublicDataset,
ImageTextPairDataset,
Text2ImageDataset)
from scepter.modules.data.dataset.ms_dataset import (
ImageTextPairFolderDataset, ImageTextPairMSDataset)
from scepter.modules.data.dataset.registry import DATASETS
from scepter.modules.data.dataset.video_gen_dataset import VideoGenDataset
else:
_import_structure = {
'base_dataset': ['BaseDataset'],
'dataset': ['Image2ImageDataset', 'ImageClassifyPublicDataset',
'ImageTextPairDataset', 'Text2ImageDataset'],
'ms_dataset': ['ImageTextPairFolderDataset',
'ImageTextPairMSDataset'],
'registry': ['DATASETS'],
'video_gen_dataset': ['VideoGenDataset']
}
import sys
sys.modules[__name__] = LazyImportModule(
__name__,
globals()['__file__'],
_import_structure,
module_spec=__spec__,
extra_objects={},
)
+1 -1
View File
@@ -82,7 +82,7 @@ class BaseDataset(Dataset, metaclass=ABCMeta):
overwrite=False)
self.worker_id = worker_id
self.logger = self.worker_logger
self.local_we["seed"] += (worker_id + we.rank)
self.local_we["seed"] += (worker_id + self.local_we['rank'] * 1234)
self.seed = self.local_we["seed"]
we.set_env(self.local_we)
+9 -1
View File
@@ -4,6 +4,7 @@
import numbers
import os
import sys
import copy
from collections.abc import Iterable
import numpy as np
@@ -242,6 +243,8 @@ class Text2ImageDataset(BaseDataset):
prompt_prefix = cfg.get('PROMPT_PREFIX', '')
path_prefix = cfg.get('PATH_PREFIX', '')
use_num = cfg.get('USE_NUM', -1)
meta_cfg = cfg.get('META_CFG', None)
meta_cfg = meta_cfg.get_lowercase_dict() if meta_cfg is not None else None
image_size = cfg.get('IMAGE_SIZE', 1024)
if isinstance(image_size, numbers.Number):
@@ -264,7 +267,12 @@ class Text2ImageDataset(BaseDataset):
self.items = list()
for i, row in enumerate(rows):
item = {'index': i, 'meta': {'image_size': image_size}}
if meta_cfg is not None:
meta_cfg_copy = copy.deepcopy(meta_cfg)
meta_cfg_copy['image_size'] = image_size
item = {'index': i, 'meta': meta_cfg_copy}
else:
item = {'index': i, 'meta': {'image_size': image_size}}
for key, value in zip(fields, row):
if key in ['prompt', 'caption', 'text']:
item['ori_prompt'] = value
+11 -4
View File
@@ -386,6 +386,11 @@ class ImageTextPairMSDatasetForACE(BaseDataset):
'description':
'The keywords sign you want to add, which is like <{HIGHLIGHT_KEYWORDS}{KEYWORDS_SIGN}>'
},
'ALIGN_SIZE': {
'value': False,
'description':
'Whether ensure the size align between the source image and target image.'
},
'OUTPUT_SIZE': {
'value':
None,
@@ -414,6 +419,8 @@ class ImageTextPairMSDatasetForACE(BaseDataset):
self.replace_keywords = cfg.get('HIGHLIGHT_KEYWORDS', '')
self.keywords_sign = cfg.get('KEYWORDS_SIGN', '')
self.add_indicator = cfg.get('ADD_INDICATOR', False)
self.align_size = cfg.get('ALIGN_SIZE', False)
# Use modelscope dataset
if not ms_dataset_name:
raise ValueError(
@@ -492,7 +499,7 @@ class ImageTextPairMSDatasetForACE(BaseDataset):
tar_image_path,
cvt_type='RGB')
src_image = self.image_preprocess(src_image)
tar_image = self.image_preprocess(tar_image)
tar_image = self.image_preprocess(tar_image, size = src_image.shape[:2] if self.align_size else None)
tar_image = self.transforms(tar_image)
src_image = self.transforms(src_image)
@@ -501,13 +508,13 @@ class ImageTextPairMSDatasetForACE(BaseDataset):
if self.add_indicator:
if '{image}' not in prompt:
prompt = '{image}, ' + prompt
return {
'edit_image': [src_image],
'edit_image_mask': [src_mask],
'src_image_list': [src_image],
'src_mask_list': [src_mask],
'image': tar_image,
'image_mask': tar_mask,
'prompt': [prompt],
'edit_id': [0]
}
def load_image(self, prefix, img_path, cvt_type=None):
+7 -2
View File
@@ -304,9 +304,10 @@ class DataObject(object):
delimiter = sampler_config.get('DELIMITER', ',')
path_prefix = sampler_config.get('PATH_PREFIX', '')
prompt_prefix = sampler_config.get('PROMPT_PREFIX', '')
oss_prefix = sampler_config.get('OSS_PREFIX', '')
return MultiLevelBatchSampler(batch_size, index_file, image_size,
fields, delimiter, path_prefix,
prompt_prefix, rank, seed)
prompt_prefix, oss_prefix, rank, seed)
def build_dataset_config(cfg, registry, logger=None, *args, **kwargs):
@@ -337,8 +338,12 @@ def build_dataset_config(cfg, registry, logger=None, *args, **kwargs):
f'registry must be type Registry, got {type(registry)}')
cfg = deep_copy(cfg)
req_type = cfg.get('NAME')
from scepter.modules.utils.import_utils import LazyImportModule
sig = (registry.name.upper(), req_type)
LazyImportModule.import_module(sig)
if isinstance(req_type, str):
req_type_entry = registry.get(req_type)
if req_type_entry is None:
@@ -0,0 +1,202 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import io
import os
import random
import sys
import warnings
import numpy as np
import torch
from tqdm import tqdm
from scepter.modules.data.dataset import DATASETS, BaseDataset
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
try:
import decord
decord.bridge.set_bridge('torch')
except ImportError:
warnings.warn(
'The `decord` package is required for loading the video dataset. Install with `pip install decord`'
)
@DATASETS.register_class()
class VideoGenDataset(BaseDataset):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.prompt_prefix = cfg.get('PROMPT_PREFIX', '')
self.path_prefix = cfg.get('PATH_PREFIX', '')
self.p_zero = cfg.get('P_ZERO', 0.0)
self.max_num_frames = cfg.get('NUM_FRAMES', 49)
self.fps = cfg.get('FPS', 8)
self.height = cfg.get('HEIGHT', 480)
self.width = cfg.get('WIDTH', 720)
self.skip_frames_start = cfg.get('SKIP_FRAMES_START', 0)
self.skip_frames_end = cfg.get('SKIP_FRAMES_END', 0)
self.data_type = cfg.get('DATA_TYPE', 't2v')
def worker_init_fn(self, worker_id, num_workers=1):
super().worker_init_fn(worker_id, num_workers=num_workers)
randseed = np.random.randint(0, 2**32 - num_workers - 1)
workerseed = randseed + worker_id
random.seed(workerseed)
np.random.seed(workerseed)
def _preprocess_video_data(self, video_path):
with FS.get_object(video_path) as video_data:
video_reader = decord.VideoReader(io.BytesIO(video_data),
width=self.width,
height=self.height)
video_num_frames = len(video_reader)
start_frame = min(self.skip_frames_start, video_num_frames)
end_frame = max(0, video_num_frames - self.skip_frames_end)
if end_frame <= start_frame:
frames = video_reader.get_batch([start_frame])
elif end_frame - start_frame <= self.max_num_frames:
frames = video_reader.get_batch(list(range(start_frame,
end_frame)))
else:
indices = list(
range(start_frame, end_frame,
(end_frame - start_frame) // self.max_num_frames))
frames = video_reader.get_batch(indices)
# Ensure that we don't go over the limit
frames = frames[:self.max_num_frames]
selected_num_frames = frames.shape[0]
# Choose first (4k + 1) frames as this is how many is required by the VAE
remainder = (3 + (selected_num_frames % 4)) % 4
if remainder != 0:
frames = frames[:-remainder]
selected_num_frames = frames.shape[0]
assert (selected_num_frames - 1) % 4 == 0
# Training transforms
frames = frames.float().div_(127.5).sub_(1.)
frames = frames.permute(3, 0, 1, 2).contiguous() # [C, F, H, W]
return frames
def _parse_index(self, index):
meta = dict()
for key, value in zip(index[-1], index[:-1]):
if key in ['oss_key', 'path', 'video_path', 'target_video_path']:
meta['video_path'] = value
elif key in ['source_video_path', 'src_video_path']:
meta['src_video_path'] = value
elif key in ['prompt', 'caption', 'text']:
meta['prompt'] = value
elif key in ['width', 'height']:
meta[key] = int(value)
else:
meta[key] = value
return meta
def _get(self, index):
meta = self._parse_index(index)
video_path = os.path.join(self.path_prefix, meta.get('video_path', ''))
video = self._preprocess_video_data(video_path)
prompt = self.prompt_prefix + meta.get('prompt', '')
if self.mode == 'train' and np.random.uniform() < self.p_zero:
prompt = ''
item = {
'video': video,
'prompt': prompt,
'meta': meta,
}
if 'i2v' in self.data_type:
item['image'] = item['video'][:, :1, :, :]
if 'v2v' in self.data_type:
src_video_path = os.path.join(self.path_prefix,
meta.get('src_video_path', ''))
src_video = self._preprocess_video_data(src_video_path)
item['src_video'] = src_video
return item
def __len__(self):
return sys.maxsize
@staticmethod
def collate_fn(batch):
collect = {}
for sample in batch:
for k, v in sample.items():
if k not in collect:
collect[k] = []
collect[k].append(v)
return collect
@DATASETS.register_class()
class VideoGenDatasetOTF(VideoGenDataset):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger)
self.data_file = cfg.DATA_FILE
self.delimiter = cfg.get('DELIMITER', '#;#')
self.fields = cfg.get('FIELDS', ['video_path', 'prompt'])
self.use_num = cfg.get('USE_NUM', -1)
from scepter.modules.model.registry import MODELS
model_cfg = cfg.get('MODEL', None)
if model_cfg is not None:
self.model = MODELS.build(
cfg.MODEL,
logger=logger).eval().requires_grad_(False).to(we.device_id)
self.items = self.parse_data(self.data_file, self.delimiter,
self.fields)
if self.use_num and self.use_num > 0:
self.items = self.items[:self.use_num]
self.data = self.encode(self.items)
self.real_number = len(self.data)
if model_cfg is not None:
self.model.to('cpu')
del self.model
torch.cuda.empty_cache()
def parse_data(self, data_file, delimiter, fields):
items = list()
with FS.get_object(data_file) as local_data:
rows = [
i.split(delimiter,
len(fields) - 1)
for i in local_data.decode('utf-8').strip().split('\n')
]
for i, row in enumerate(rows):
item = {}
for key, value in zip(self.fields, row):
if key in ['oss_key', 'path', 'video_path']:
item['video_path'] = value
elif key in ['prompt', 'caption', 'text']:
item['prompt'] = value
elif key in ['width', 'height']:
item[key] = int(value)
else:
item[key] = value
items.append(item)
return items
def encode(self, items):
self.logger.info('Start to encode video data [{}]!'.format(len(items)))
for item in tqdm(items):
video_path = os.path.join(self.path_prefix,
item.get('video_path', ''))
video = self._preprocess_video_data(video_path)
latent = self.model.encode_first_stage(
video.unsqueeze(0).to(we.device_id)).squeeze(0)
item['video_latent'] = latent.detach().cpu()
item['video'] = video
if self.data_type == 'i2v':
item['image'] = item['video'][:, :1, :, :]
return items
def _get(self, index):
return self.data[index % self.real_number]
+28 -6
View File
@@ -1,9 +1,31 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from typing import TYPE_CHECKING
from scepter.modules.utils.import_utils import LazyImportModule
from scepter.modules.data.sampler.base_sampler import BaseSampler
from scepter.modules.data.sampler.registry import SAMPLERS
from scepter.modules.data.sampler.sampler import (
EvalDistributedSampler, LoopSampler, MixtureOfSamplers,
MultiFoldDistributedSampler, MultiLevelBatchSampler,
MultiLevelBatchSamplerMultiSource, ResolutionBatchSampler)
if TYPE_CHECKING:
from scepter.modules.data.sampler.base_sampler import BaseSampler
from scepter.modules.data.sampler.registry import SAMPLERS
from scepter.modules.data.sampler.sampler import (
EvalDistributedSampler, LoopSampler, MixtureOfSamplers,
MultiFoldDistributedSampler, MultiLevelBatchSampler,
MultiLevelBatchSamplerMultiSource, ResolutionBatchSampler)
else:
_import_structure = {
'base_sampler': ['BaseSampler'],
'registry': ['SAMPLERS'],
'sampler': ['EvalDistributedSampler', 'LoopSampler',
'MixtureOfSamplers', 'MultiFoldDistributedSampler',
'MultiLevelBatchSampler', 'MultiLevelBatchSamplerMultiSource',
'ResolutionBatchSampler']
}
import sys
sys.modules[__name__] = LazyImportModule(
__name__,
globals()['__file__'],
_import_structure,
module_spec=__spec__,
extra_objects={},
)
+5 -1
View File
@@ -35,8 +35,12 @@ def build_sampler_config(cfg, registry, logger=None, **kwargs):
f'registry must be type Registry, got {type(registry)}')
cfg = deep_copy(cfg)
req_type = cfg.get('NAME')
from scepter.modules.utils.import_utils import LazyImportModule
sig = (registry.name.upper(), req_type)
LazyImportModule.import_module(sig)
if isinstance(req_type, str):
req_type_entry = registry.get(req_type)
if req_type_entry is None:
+4 -2
View File
@@ -79,6 +79,7 @@ class MultiLevelBatchSamplerMultiSource(BaseSampler):
self.num_fields = len(self.fields)
self.delimiter = cfg.get('DELIMITER', ',')
self.path_prefix = cfg.get('PATH_PREFIX', '')
oss_prefix = cfg.get('OSS_PREFIX', '')
common_prob = cfg.get('PROB', 1)
sub_data_weights = cfg.get('SUB_DATA_WEIGHTS', None)
sub_data_weights = {} if sub_data_weights is None else sub_data_weights.get_dict(
@@ -137,7 +138,7 @@ class MultiLevelBatchSamplerMultiSource(BaseSampler):
f"{p * common_prob} and samples'num: {sub_data['total']} in this cluster."
)
self.rng = np.random.default_rng(self.seed + we.rank)
self.oss_prefix = '/'.join(index_file.split('/')[:3])
self.oss_prefix = '/'.join(index_file.split('/')[:3]) if (oss_prefix is None or oss_prefix == '') and index_file.startswith('oss') else oss_prefix
self.index_dir = os.path.dirname(index_file)
def __iter__(self):
@@ -434,6 +435,7 @@ class MultiLevelBatchSampler(BaseSampler):
delimiter=',',
path_prefix='',
prompt_prefix='',
oss_prefix='',
rank=0,
seed=8888):
self.batch_size = batch_size
@@ -457,7 +459,7 @@ class MultiLevelBatchSampler(BaseSampler):
'index_level': 1,
'num_fields': self.num_fields
}
self.oss_prefix = '/'.join(index_file.split('/')[:3])
self.oss_prefix = '/'.join(index_file.split('/')[:3]) if (oss_prefix is None or oss_prefix == '') and index_file.startswith('oss') else oss_prefix
self.index_dir = os.path.dirname(index_file)
def __iter__(self):
+18 -1
View File
@@ -1,4 +1,21 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from typing import TYPE_CHECKING
from scepter.modules.utils.import_utils import LazyImportModule
from scepter.modules.data.utils.data_bucket import BucketManager
if TYPE_CHECKING:
from scepter.modules.data.utils.data_bucket import BucketManager
else:
_import_structure = {
'data_bucket': ['BucketManager']
}
import sys
sys.modules[__name__] = LazyImportModule(
__name__,
globals()['__file__'],
_import_structure,
module_spec=__spec__,
extra_objects={},
)
+37 -1
View File
@@ -1,3 +1,39 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.inference.diffusion_inference import DiffusionInference
from typing import TYPE_CHECKING
from scepter.modules.utils.import_utils import LazyImportModule
if TYPE_CHECKING:
from scepter.modules.inference.diffusion_inference import DiffusionInference
from scepter.modules.inference.ace_inference import ACEInference
from scepter.modules.inference.cogvideox_inference import CogVideoXInference
from scepter.modules.inference.control_inference import ControlInference
from scepter.modules.inference.flux_inference import FluxInference
from scepter.modules.inference.largen_inference import LargenInference
from scepter.modules.inference.pixart_inference import PixArtInference
from scepter.modules.inference.sd3_inference import SD3Inference
from scepter.modules.inference.stylebooth_inference import StyleboothInference
from scepter.modules.inference.tuner_inference import TunerInference
else:
_import_structure = {
'diffusion_inference': ['DiffusionInference'],
'ace_inference': ['ACEInference'],
'cogvideox_inference': ['CogVideoXInference'],
'control_inference': ['ControlInference'],
'flux_inference': ['FluxInference'],
'largen_inference': ['LargenInference'],
'pixart_inference': ['PixArtInference'],
'sd3_inference': ['SD3Inference'],
'stylebooth_inference': ['StyleboothInference'],
'tuner_inference': ['TunerInference']
}
import sys
sys.modules[__name__] = LazyImportModule(
__name__,
globals()['__file__'],
_import_structure,
module_spec=__spec__,
extra_objects={},
)
+281 -106
View File
@@ -10,7 +10,7 @@ import torch.nn as nn
import torch.nn.functional as F
import torchvision.transforms.functional as TF
from PIL import Image
import torchvision.transforms as T
from scepter.modules.model.registry import DIFFUSIONS
from scepter.modules.model.utils.basic_utils import check_list_of_list
from scepter.modules.model.utils.basic_utils import \
@@ -85,6 +85,138 @@ class TextEmbedding(nn.Module):
super().__init__()
self.pos = nn.Parameter(data=torch.zeros(embedding_shape))
class RefinerInference(DiffusionInference):
def init_from_cfg(self, cfg):
self.use_dynamic_model = cfg.get('USE_DYNAMIC_MODEL', True)
super().init_from_cfg(cfg)
self.diffusion = DIFFUSIONS.build(cfg.MODEL.DIFFUSION, logger=self.logger) \
if cfg.MODEL.have('DIFFUSION') else None
self.max_seq_length = cfg.MODEL.get("MAX_SEQ_LENGTH", 4096)
assert self.diffusion is not None
if not self.use_dynamic_model:
self.dynamic_load(self.first_stage_model, 'first_stage_model')
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
self.dynamic_load(self.diffusion_model, 'diffusion_model')
@torch.no_grad()
def encode_first_stage(self, x, **kwargs):
_, dtype = self.get_function_info(self.first_stage_model, 'encode')
with torch.autocast('cuda',
enabled=dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, dtype)):
def run_one_image(u):
zu = get_model(self.first_stage_model).encode(u)
if isinstance(zu, (tuple, list)):
zu = zu[0]
return zu
z = [run_one_image(u.unsqueeze(0) if u.dim == 3 else u) for u in x]
return z
def upscale_resize(self, image, interpolation=T.InterpolationMode.BILINEAR):
c, H, W = image.shape
scale = max(1.0, math.sqrt(self.max_seq_length / ((H / 16) * (W / 16))))
rH = int(H * scale) // 16 * 16 # ensure divisible by self.d
rW = int(W * scale) // 16 * 16
image = T.Resize((rH, rW), interpolation=interpolation, antialias=True)(image)
return image
@torch.no_grad()
def decode_first_stage(self, z):
_, dtype = self.get_function_info(self.first_stage_model, 'decode')
with torch.autocast('cuda',
enabled=dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, dtype)):
return [get_model(self.first_stage_model).decode(zu) for zu in z]
def noise_sample(self, num_samples, h, w, seed, device = None, dtype = torch.bfloat16):
noise = torch.randn(
num_samples,
16,
# allow for packing
2 * math.ceil(h / 16),
2 * math.ceil(w / 16),
device=device,
dtype=dtype,
generator=torch.Generator(device=device).manual_seed(seed),
)
return noise
def refine(self,
x_samples=None,
prompt=None,
reverse_scale=-1.,
seed = 2024,
**kwargs
):
print(prompt)
value_input = copy.deepcopy(self.input)
x_samples = [self.upscale_resize(x) for x in x_samples]
noise = []
for i, x in enumerate(x_samples):
noise_ = self.noise_sample(1, x.shape[1],
x.shape[2], seed,
device = x.device)
noise.append(noise_)
noise, x_shapes = pack_imagelist_into_tensor(noise)
if reverse_scale > 0:
self.dynamic_load(self.first_stage_model, 'first_stage_model')
x_samples = [x.unsqueeze(0) for x in x_samples]
x_start = self.encode_first_stage(x_samples, **kwargs)
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=not self.use_dynamic_model)
x_start, _ = pack_imagelist_into_tensor(x_start)
else:
x_start = None
# cond stage
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
function_name, dtype = self.get_function_info(self.cond_stage_model)
with torch.autocast('cuda',
enabled=dtype == 'float16',
dtype=getattr(torch, dtype)):
ctx = getattr(get_model(self.cond_stage_model),
function_name)(prompt)
ctx["x_shapes"] = x_shapes
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=not self.use_dynamic_model)
self.dynamic_load(self.diffusion_model, 'diffusion_model')
# UNet use input n_prompt
function_name, dtype = self.get_function_info(
self.diffusion_model)
with torch.autocast('cuda',
enabled=dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, dtype)):
solver_sample = value_input.get('sample', 'flow_euler')
sample_steps = value_input.get('sample_steps', 20)
guide_scale = value_input.get('guide_scale', 3.5)
if guide_scale is not None:
guide_scale = torch.full((noise.shape[0],), guide_scale, device=noise.device,
dtype=noise.dtype)
else:
guide_scale = None
latent = self.diffusion.sample(
noise=noise,
sampler=solver_sample,
model=get_model(self.diffusion_model),
model_kwargs={"cond": ctx, "guidance": guide_scale},
steps=sample_steps,
show_progress=True,
guide_scale=guide_scale,
return_intermediate=None,
reverse_scale=reverse_scale,
x=x_start,
**kwargs).float()
latent = unpack_tensor_into_imagelist(latent, x_shapes)
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=not self.use_dynamic_model)
self.dynamic_load(self.first_stage_model, 'first_stage_model')
x_samples = self.decode_first_stage(latent)
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=not self.use_dynamic_model)
return x_samples
class ACEInference(DiffusionInference):
def __init__(self, logger=None):
@@ -99,6 +231,7 @@ class ACEInference(DiffusionInference):
def init_from_cfg(self, cfg):
self.name = cfg.NAME
self.is_default = cfg.get('IS_DEFAULT', False)
self.use_dynamic_model = cfg.get('USE_DYNAMIC_MODEL', True)
module_paras = self.load_default(cfg.get('DEFAULT_PARAS', None))
assert cfg.have('MODEL')
@@ -116,9 +249,22 @@ class ACEInference(DiffusionInference):
module_paras.get(
'COND_STAGE_MODEL',
None)) if cfg.MODEL.have('COND_STAGE_MODEL') else None
self.refiner_model_cfg = cfg.get('REFINER_MODEL', None)
# self.refiner_scale = cfg.get('REFINER_SCALE', 0.)
# self.refiner_prompt = cfg.get('REFINER_PROMPT', "")
self.ace_prompt = cfg.get("ACE_PROMPT", [])
if self.refiner_model_cfg:
self.refiner_model_cfg.USE_DYNAMIC_MODEL = self.use_dynamic_model
self.refiner_module = RefinerInference(self.logger)
self.refiner_module.init_from_cfg(self.refiner_model_cfg)
else:
self.refiner_module = None
self.diffusion = DIFFUSIONS.build(cfg.MODEL.DIFFUSION,
logger=self.logger)
self.interpolate_func = lambda x: (F.interpolate(
x.unsqueeze(0),
scale_factor=1 / self.size_factor,
@@ -137,6 +283,10 @@ class ACEInference(DiffusionInference):
self.size_factor = cfg.get('SIZE_FACTOR', 8)
self.decoder_bias = cfg.get('DECODER_BIAS', 0)
self.default_n_prompt = cfg.get('DEFAULT_N_PROMPT', '')
if not self.use_dynamic_model:
self.dynamic_load(self.first_stage_model, 'first_stage_model')
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
self.dynamic_load(self.diffusion_model, 'diffusion_model')
@torch.no_grad()
def encode_first_stage(self, x, **kwargs):
@@ -163,6 +313,8 @@ class ACEInference(DiffusionInference):
]
return x
@torch.no_grad()
def __call__(self,
image=None,
@@ -184,7 +336,6 @@ class ACEInference(DiffusionInference):
g = torch.Generator(device=we.device_id)
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
g.manual_seed(int(seed))
if input_image is not None:
# assert isinstance(input_image, list) and isinstance(input_mask, list)
if task is None:
@@ -237,118 +388,142 @@ class ACEInference(DiffusionInference):
assert isinstance(nn_p, list)
n_prompt[nn_p_id][-1] = negative_prompt
ctx, null_ctx = {}, {}
# Get Noise Shape
self.dynamic_load(self.first_stage_model, 'first_stage_model')
is_txt_image = sum([len(e_i) for e_i in edit_image]) < 1
image = to_device(image)
x = self.encode_first_stage(image)
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=True)
noise = [
torch.empty(*i.shape, device=we.device_id).normal_(generator=g)
for i in x
]
noise, x_shapes = pack_imagelist_into_tensor(noise)
ctx['x_shapes'] = null_ctx['x_shapes'] = x_shapes
image_mask = to_device(image_mask, strict=False)
cond_mask = [self.interpolate_func(i) for i in image_mask
] if image_mask is not None else [None] * len(image)
ctx['x_mask'] = null_ctx['x_mask'] = cond_mask
refiner_scale = kwargs.pop("refiner_scale", 0.0)
refiner_prompt = kwargs.pop("refiner_prompt", "")
use_ace = kwargs.pop("use_ace", True)
# <= 0 use ace as the txt2img generator.
if use_ace and (not is_txt_image or refiner_scale <= 0):
ctx, null_ctx = {}, {}
# Get Noise Shape
self.dynamic_load(self.first_stage_model, 'first_stage_model')
x = self.encode_first_stage(image)
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=not self.use_dynamic_model)
noise = [
torch.empty(*i.shape, device=we.device_id).normal_(generator=g)
for i in x
]
noise, x_shapes = pack_imagelist_into_tensor(noise)
ctx['x_shapes'] = null_ctx['x_shapes'] = x_shapes
# Encode Prompt
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
function_name, dtype = self.get_function_info(self.cond_stage_model)
cont, cont_mask = getattr(get_model(self.cond_stage_model),
function_name)(prompt)
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont,
cont_mask)
null_cont, null_cont_mask = getattr(get_model(self.cond_stage_model),
function_name)(n_prompt)
null_cont, null_cont_mask = self.cond_stage_embeddings(
prompt, edit_image, null_cont, null_cont_mask)
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=False)
ctx['crossattn'] = cont
null_ctx['crossattn'] = null_cont
image_mask = to_device(image_mask, strict=False)
cond_mask = [self.interpolate_func(i) for i in image_mask
] if image_mask is not None else [None] * len(image)
ctx['x_mask'] = null_ctx['x_mask'] = cond_mask
# Encode Edit Images
self.dynamic_load(self.first_stage_model, 'first_stage_model')
edit_image = [to_device(i, strict=False) for i in edit_image]
edit_image_mask = [to_device(i, strict=False) for i in edit_image_mask]
e_img, e_mask = [], []
for u, m in zip(edit_image, edit_image_mask):
if u is None:
continue
if m is None:
m = [None] * len(u)
e_img.append(self.encode_first_stage(u, **kwargs))
e_mask.append([self.interpolate_func(i) for i in m])
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=True)
null_ctx['edit'] = ctx['edit'] = e_img
null_ctx['edit_mask'] = ctx['edit_mask'] = e_mask
# Encode Prompt
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
function_name, dtype = self.get_function_info(self.cond_stage_model)
cont, cont_mask = getattr(get_model(self.cond_stage_model),
function_name)(prompt)
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont,
cont_mask)
null_cont, null_cont_mask = getattr(get_model(self.cond_stage_model),
function_name)(n_prompt)
null_cont, null_cont_mask = self.cond_stage_embeddings(
prompt, edit_image, null_cont, null_cont_mask)
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=not self.use_dynamic_model)
ctx['crossattn'] = cont
null_ctx['crossattn'] = null_cont
# Diffusion Process
self.dynamic_load(self.diffusion_model, 'diffusion_model')
function_name, dtype = self.get_function_info(self.diffusion_model)
with torch.autocast('cuda',
enabled=dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, dtype)):
latent = self.diffusion.sample(
noise=noise,
sampler=sampler,
model=get_model(self.diffusion_model),
model_kwargs=[{
'cond':
ctx,
'mask':
cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
}, {
'cond':
null_ctx,
'mask':
null_cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
}] if guide_scale is not None and guide_scale > 1 else {
'cond':
null_ctx,
'mask':
cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
},
steps=sample_steps,
show_progress=True,
seed=seed,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
return_intermediate=None,
**kwargs)
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=False)
# Encode Edit Images
self.dynamic_load(self.first_stage_model, 'first_stage_model')
edit_image = [to_device(i, strict=False) for i in edit_image]
edit_image_mask = [to_device(i, strict=False) for i in edit_image_mask]
e_img, e_mask = [], []
for u, m in zip(edit_image, edit_image_mask):
if u is None:
continue
if m is None:
m = [None] * len(u)
e_img.append(self.encode_first_stage(u, **kwargs))
e_mask.append([self.interpolate_func(i) for i in m])
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=not self.use_dynamic_model)
null_ctx['edit'] = ctx['edit'] = e_img
null_ctx['edit_mask'] = ctx['edit_mask'] = e_mask
# Decode to Pixel Space
self.dynamic_load(self.first_stage_model, 'first_stage_model')
samples = unpack_tensor_into_imagelist(latent, x_shapes)
x_samples = self.decode_first_stage(samples)
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=False)
# Diffusion Process
self.dynamic_load(self.diffusion_model, 'diffusion_model')
function_name, dtype = self.get_function_info(self.diffusion_model)
with torch.autocast('cuda',
enabled=dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, dtype)):
latent = self.diffusion.sample(
noise=noise,
sampler=sampler,
model=get_model(self.diffusion_model),
model_kwargs=[{
'cond':
ctx,
'mask':
cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
}, {
'cond':
null_ctx,
'mask':
null_cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
}] if guide_scale is not None and guide_scale > 1 else {
'cond':
null_ctx,
'mask':
cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
},
steps=sample_steps,
show_progress=True,
seed=seed,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
return_intermediate=None,
**kwargs)
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=not self.use_dynamic_model)
# Decode to Pixel Space
self.dynamic_load(self.first_stage_model, 'first_stage_model')
samples = unpack_tensor_into_imagelist(latent, x_shapes)
x_samples = self.decode_first_stage(samples)
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=not self.use_dynamic_model)
x_samples = [x.squeeze(0) for x in x_samples]
else:
x_samples = image
if self.refiner_module and refiner_scale > 0:
if is_txt_image:
random.shuffle(self.ace_prompt)
input_refine_prompt = [self.ace_prompt[0] + refiner_prompt if p[0] == "" else p[0] for p in prompt]
input_refine_scale = -1.
else:
input_refine_prompt = [p[0].replace("{image}", "") + " " + refiner_prompt for p in prompt]
input_refine_scale = refiner_scale
print(input_refine_prompt)
x_samples = self.refiner_module.refine(x_samples,
reverse_scale = input_refine_scale,
prompt= input_refine_prompt,
seed=seed,
use_dynamic_model=self.use_dynamic_model)
imgs = [
torch.clamp((x_i + 1.0) / 2.0 + self.decoder_bias / 255,
torch.clamp((x_i.float() + 1.0) / 2.0 + self.decoder_bias / 255,
min=0.0,
max=1.0).squeeze(0).permute(1, 2, 0).cpu().numpy()
for x_i in x_samples
@@ -0,0 +1,183 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import numpy as np
from typing import Tuple
import random
import torch
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.distribute import we
from scepter.modules.model.backbone.cogvideox.utils import get_3d_rotary_pos_embed, get_resize_crop_region_for_grid
from .diffusion_inference import DiffusionInference, get_model
from .tuner_inference import TunerInference
class CogVideoXInference(DiffusionInference):
def __init__(self, logger=None):
self.logger = logger
self.is_redefine_paras = False
self.loaded_model = {}
self.loaded_model_name = [
'diffusion_model', 'first_stage_model', 'cond_stage_model'
]
self.tuner_infer = TunerInference(self.logger)
@torch.no_grad()
def decode_first_stage(self, latents):
_, dtype = self.get_function_info(self.first_stage_model, 'decode')
with torch.autocast('cuda',
enabled=dtype in ('bfloat16'),
dtype=getattr(torch, dtype)):
latents = latents.permute(0, 2, 1, 3, 4)
latents = 1 / self.first_stage_model['paras']['scaling_factor_image'] * latents
frames = get_model(self.first_stage_model).decode(latents)
return frames
def _prepare_rotary_positional_embeddings(
self,
height: int,
width: int,
num_frames: int,
device: torch.device,
) -> Tuple[torch.Tensor, torch.Tensor]:
grid_height = height // (self.diffusion_model['paras']['scale_factor_spatial'] * self.diffusion_model['paras']['patch_size'])
grid_width = width // (self.diffusion_model['paras']['scale_factor_spatial'] * self.diffusion_model['paras']['patch_size'])
base_size_width = self.diffusion_model['paras']['sample_width'] // (self.diffusion_model['paras']['scale_factor_spatial'] * self.diffusion_model['paras']['patch_size'])
base_size_height = self.diffusion_model['paras']['sample_height'] // (self.diffusion_model['paras']['scale_factor_spatial'] * self.diffusion_model['paras']['patch_size'])
grid_crops_coords = get_resize_crop_region_for_grid(
(grid_height, grid_width), base_size_width, base_size_height
)
freqs_cos, freqs_sin = get_3d_rotary_pos_embed(
embed_dim=self.diffusion_model['paras']['attention_head_dim'],
crops_coords=grid_crops_coords,
grid_size=(grid_height, grid_width),
temporal_size=num_frames,
)
freqs_cos = freqs_cos.to(device=device)
freqs_sin = freqs_sin.to(device=device)
return freqs_cos, freqs_sin
@torch.no_grad()
def __call__(self,
input,
num_samples=1,
cat_uc=True,
tuner_model=None,
**kwargs):
value_input = copy.deepcopy(self.input)
value_input.update(input)
print(value_input)
height, width = value_input['target_size_as_tuple']
value_output = copy.deepcopy(self.output)
# register tuner
if tuner_model is not None and tuner_model != '' and len(
tuner_model) > 0:
if not isinstance(tuner_model, list):
tuner_model = [tuner_model]
self.dynamic_load(self.diffusion_model, 'diffusion_model')
self.tuner_infer.register_tuner(tuner_model, self.diffusion_model,
cond_stage_model=None)
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=True)
# cond stage
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
function_name, dtype = self.get_function_info(self.cond_stage_model)
with torch.autocast(device_type='cuda', enabled=True, dtype=torch.bfloat16):
cont = getattr(get_model(self.cond_stage_model),
function_name)(value_input['prompt'], return_mask=False, use_mask=False)
null_cont = getattr(get_model(self.cond_stage_model),
function_name)(value_input['negative_prompt'] * num_samples, return_mask=False, use_mask=False)
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=True)
# get noise
seed = kwargs.pop('seed', -1)
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
generator = torch.Generator().manual_seed(seed)
if 'seed' in value_output:
value_output['seed'] = seed
for sample_id in range(num_samples):
if self.diffusion_model is not None:
noise_shape = (1,
(value_input['num_frames'] - 1) // self.diffusion_model['paras']['scale_factor_temporal'] + 1,
self.diffusion_model['paras']['latent_channels'],
height // self.diffusion_model['paras']['scale_factor_spatial'],
width // self.diffusion_model['paras']['scale_factor_spatial']
)
noise = torch.randn(noise_shape, generator=generator, dtype=getattr(torch, dtype), device='cpu').to(we.device_id)
self.dynamic_load(self.diffusion_model, 'diffusion_model')
image_rotary_emb = (
self._prepare_rotary_positional_embeddings(height, width, noise.size(1), we.device_id)
if self.diffusion_model['paras']['use_rotary_positional_embeddings']
else None
)
function_name, dtype = self.get_function_info(
self.diffusion_model)
with torch.autocast('cuda',
enabled=dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, dtype)):
solver_sample = value_input.get('sample', 'ddim')
sample_steps = value_input.get('sample_steps', 50)
guide_scale = value_input.get('guide_scale', 7.5)
guide_rescale = value_input.get('guide_rescale', 0.5)
latent = self.diffusion.sample(noise=noise,
sampler=solver_sample,
model=get_model(self.diffusion_model),
model_kwargs=[{
'cond': cont,
'image_latent': None,
'image_rotary_emb': image_rotary_emb,
}, {
'cond': null_cont,
'image_latent': None,
'image_rotary_emb': image_rotary_emb,
}],
steps=sample_steps,
show_progress=True,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
return_intermediate=None,
**kwargs).float()
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=True)
self.dynamic_load(self.first_stage_model, 'first_stage_model')
x_samples = self.decode_first_stage(latent).float() # [B, C, F, H, W]
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=True)
x_frames = torch.clamp(x_samples / 2 + 0.5, min=0.0, max=1.0)
if 'videos' in value_output:
if value_output['videos'] is None or (
isinstance(value_output['videos'], list)
and len(value_output['videos']) < 1):
value_output['videos'] = []
value_output['videos'].append(x_frames)
for k, v in value_output.items():
if isinstance(v, list):
value_output[k] = torch.cat(v, dim=0)
if isinstance(v, torch.Tensor):
value_output[k] = v.cpu()
# unregister tuner
if tuner_model is not None and tuner_model != '' and len(
tuner_model) > 0:
self.tuner_infer.unregister_tuner(tuner_model,
self.diffusion_model,
cond_stage_model=None)
return value_output
@@ -14,6 +14,7 @@ from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, MODELS,
TOKENIZERS, DIFFUSIONS)
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.config import Config
from scepter.studio.utils.env import get_available_memory
from .control_inference import ControlInference
@@ -96,7 +97,7 @@ class DiffusionInference():
if 'weights_only' in torch.load.__code__.co_varnames:
sd = torch.load(local_path, map_location='cpu', weights_only=True)
else:
sd = torch.load(local_path, map_location='cpu')
sd = torch.load(local_path, map_location='cpu', weights_only=True)
first_stage_model_path = os.path.join(
os.path.dirname(local_path), 'first_stage_model.pth')
cond_stage_model_path = os.path.join(
@@ -202,7 +203,7 @@ class DiffusionInference():
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(path)
else:
sd = torch.load(path, map_location='cpu')
sd = torch.load(path, map_location='cpu', weights_only=True)
new_sd = OrderedDict()
for k, v in sd.items():
@@ -229,16 +230,22 @@ class DiffusionInference():
def load(self, module):
if module['device'] == 'offline':
if module['cfg'].NAME in MODELS.class_map:
from scepter.modules.utils.import_utils import LazyImportModule
if (LazyImportModule.get_module_type(('MODELS', module['cfg'].NAME)) or
module['cfg'].NAME in MODELS.class_map):
model = MODELS.build(module['cfg'], logger=self.logger).eval()
elif module['cfg'].NAME in BACKBONES.class_map:
elif (LazyImportModule.get_module_type(('BACKBONES', module['cfg'].NAME)) or
module['cfg'].NAME in BACKBONES.class_map):
model = BACKBONES.build(module['cfg'],
logger=self.logger).eval()
elif module['cfg'].NAME in EMBEDDERS.class_map:
elif (LazyImportModule.get_module_type(('EMBEDDERS', module['cfg'].NAME)) or
module['cfg'].NAME in EMBEDDERS.class_map):
model = EMBEDDERS.build(module['cfg'],
logger=self.logger).eval()
else:
raise NotImplementedError
if 'DTYPE' in module['cfg'] and module['cfg']['DTYPE'] is not None:
model = model.to(getattr(torch, module['cfg'].DTYPE))
if module['cfg'].get('RELOAD_MODEL', None):
self.init_from_ckpt(module['cfg'].RELOAD_MODEL, model)
module['model'] = model
@@ -267,8 +274,9 @@ class DiffusionInference():
module['device'] = 'cpu'
else:
module['device'] = 'offline'
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
return module
def dynamic_load(self, module=None, name=''):
@@ -316,7 +324,8 @@ class DiffusionInference():
module_paras = {}
if cfg is not None:
self.paras = cfg.PARAS
self.input = {k.lower(): dict(v).get('DEFAULT', None) if isinstance(v, (dict, OrderedDict)) else v for k, v in cfg.INPUT.items()}
self.input_cfg = {k.lower(): v for k, v in cfg.INPUT.items()}
self.input = {k.lower(): dict(v).get('DEFAULT', None) if isinstance(v, (dict, OrderedDict, Config)) else v for k, v in cfg.INPUT.items()}
self.output = {k.lower(): v for k, v in cfg.OUTPUT.items()}
module_paras = cfg.MODULES_PARAS
return module_paras
+1 -1
View File
@@ -151,7 +151,7 @@ class FluxInference(DiffusionInference):
with torch.autocast('cuda',
enabled= dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, dtype)):
solver_sample = value_input.get('sample', 'flow_eluer')
solver_sample = value_input.get('sample', 'flow_euler')
sample_steps = value_input.get('sample_steps', 20)
guide_scale = value_input.get('guide_scale', 3.5)
if guide_scale is not None:
@@ -43,7 +43,7 @@ class LargenInference(DiffusionInference):
if 'weights_only' in torch.load.__code__.co_varnames:
sd = torch.load(local_path, map_location='cpu', weights_only=True)
else:
sd = torch.load(local_path, map_location='cpu')
sd = torch.load(local_path, map_location='cpu', weights_only=True)
if 'model' in sd:
sd = sd['model']
+4 -4
View File
@@ -29,11 +29,11 @@ class TunerInference():
warnings.warn(f'Import swift error, please deal with this problem: {e}')
self.logger.info('Unloading tuner model')
if isinstance(diffusion_model['model'], SwiftModel):
if diffusion_model is not None and isinstance(diffusion_model['model'], SwiftModel):
for adapter_name in diffusion_model['model'].adapters:
diffusion_model['model'].deactivate_adapter(adapter_name,
offload='cpu')
if isinstance(cond_stage_model['model'], SwiftModel):
if cond_stage_model is not None and isinstance(cond_stage_model['model'], SwiftModel):
for adapter_name in cond_stage_model['model'].adapters:
cond_stage_model['model'].deactivate_adapter(adapter_name,
offload='cpu')
@@ -144,9 +144,9 @@ class TunerInference():
is_bin_file = True
if os.path.isfile(bin_file):
if 'weights_only' in torch.load.__code__.co_varnames:
state_dict = torch.load(bin_file, weights_only=True)
state_dict = torch.load(bin_file, weights_only=True, map_location="cpu")
else:
state_dict = torch.load(bin_file)
state_dict = torch.load(bin_file, map_location="cpu")
elif os.path.isfile(safe_file):
is_bin_file = False
from safetensors.torch import \
+20 -2
View File
@@ -1,5 +1,23 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from typing import TYPE_CHECKING
from scepter.modules.utils.import_utils import LazyImportModule
from scepter.modules.model import (backbone, embedder, head, loss, metric,
neck, network, tokenizer, tuner, diffusion)
if TYPE_CHECKING:
from scepter.modules.model import (backbone, embedder, head, loss, metric,
neck, network, tokenizer, tuner, diffusion)
else:
_import_structure = {
'model': ['backbone', 'embedder', 'head', 'loss', 'metric',
'neck', 'network', 'tokenizer', 'tuner', 'diffusion']
}
import sys
sys.modules[__name__] = LazyImportModule(
__name__,
globals()['__file__'],
_import_structure,
module_spec=__spec__,
extra_objects={},
)
+21 -2
View File
@@ -1,4 +1,23 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.backbone import (ace, autoencoder, flux, image,
mmdit, pixart, unet, utils, video)
from typing import TYPE_CHECKING
from scepter.modules.utils.import_utils import LazyImportModule
if TYPE_CHECKING:
from scepter.modules.model.backbone import (ace, autoencoder, flux, image, cogvideox,
mmdit, pixart, unet, utils, video)
else:
_import_structure = {
'backbone': ['ace', 'autoencoder', 'flux', 'image', 'cogvideox',
'mmdit', 'pixart', 'unet', 'utils', 'video']
}
import sys
sys.modules[__name__] = LazyImportModule(
__name__,
globals()['__file__'],
_import_structure,
module_spec=__spec__,
extra_objects={},
)
+1 -1
View File
@@ -151,7 +151,7 @@ class ACE(BaseModel):
def load_pretrained_model(self, pretrained_model):
if pretrained_model:
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
model = torch.load(local_path, map_location='cpu')
model = torch.load(local_path, map_location='cpu', weights_only=True)
if 'state_dict' in model:
model = model['state_dict']
new_ckpt = OrderedDict()
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
@@ -0,0 +1,97 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import torch
from torch.nn.utils.rnn import pad_sequence
from einops import rearrange
from scepter.modules.model.backbone.flux import FluxMR
from scepter.modules.model.registry import BACKBONES
from scepter.modules.utils.config import dict_to_yaml
@BACKBONES.register_class()
class FluxMRACEPlus(FluxMR):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger)
def prepare_input(self, x, cond):
context, y = cond['context'], cond['y']
batch_frames, batch_frames_ids = [], []
for ix, shape, imask, ie, ie_mask in zip(x, cond['x_shapes'],
cond['x_mask'], cond['edit'],
cond['edit_mask']):
# unpack image from sequence
ix = ix[:, :shape[0] * shape[1]].view(-1, shape[0], shape[1])
imask = torch.ones_like(
ix[[0], :, :]) if imask is None else imask.squeeze(0)
if len(ie) > 0:
ie = [iie.squeeze(0) for iie in ie]
ie_mask = [
torch.ones(
(ix.shape[0] * 4, ix.shape[1],
ix.shape[2])) if iime is None else iime.squeeze(0)
for iime in ie_mask
]
ie = torch.cat(ie, dim=-1)
ie_mask = torch.cat(ie_mask, dim=-1)
else:
ie, ie_mask = torch.zeros_like(ix).to(x), torch.ones_like(
imask).to(x)
ix = torch.cat([ix, ie, ie_mask], dim=0)
c, h, w = ix.shape
ix = rearrange(ix,
'c (h ph) (w pw) -> (h w) (c ph pw)',
ph=2,
pw=2)
ix_id = torch.zeros(h // 2, w // 2, 3)
ix_id[..., 1] = ix_id[..., 1] + torch.arange(h // 2)[:, None]
ix_id[..., 2] = ix_id[..., 2] + torch.arange(w // 2)[None, :]
ix_id = rearrange(ix_id, 'h w c -> (h w) c')
batch_frames.append([ix])
batch_frames_ids.append([ix_id])
x_list, x_id_list, mask_x_list, x_seq_length = [], [], [], []
for frames, frame_ids in zip(batch_frames, batch_frames_ids):
proj_frames = []
for idx, one_frame in enumerate(frames):
one_frame = self.img_in(one_frame)
proj_frames.append(one_frame)
ix = torch.cat(proj_frames, dim=0)
if_id = torch.cat(frame_ids, dim=0)
x_list.append(ix)
x_id_list.append(if_id)
mask_x_list.append(
torch.ones(ix.shape[0]).to(ix.device,
non_blocking=True).bool())
x_seq_length.append(ix.shape[0])
# if len(x_list) < 1: import pdb;pdb.set_trace()
x = pad_sequence(tuple(x_list), batch_first=True)
x_ids = pad_sequence(tuple(x_id_list), batch_first=True).to(
x) # [b,pad_seq,2] pad (0.,0.) at dim2
mask_x = pad_sequence(tuple(mask_x_list), batch_first=True)
# import pdb;pdb.set_trace()
if isinstance(context, list):
txt_list, mask_txt_list, y_list = [], [], []
for sample_id, (ctx, yy) in enumerate(zip(context, y)):
txt_list.append(self.txt_in(ctx.to(x)))
mask_txt_list.append(
torch.ones(txt_list[-1].shape[0]).to(
ctx.device, non_blocking=True).bool())
y_list.append(yy.to(x))
txt = pad_sequence(tuple(txt_list), batch_first=True)
txt_ids = torch.zeros(txt.shape[0], txt.shape[1], 3).to(x)
mask_txt = pad_sequence(tuple(mask_txt_list), batch_first=True)
y = torch.cat(y_list, dim=0)
assert y.ndim == 2 and txt.ndim == 3
else:
txt = self.txt_in(context)
txt_ids = torch.zeros(context.shape[0], context.shape[1], 3).to(x)
mask_txt = torch.ones(context.shape[0], context.shape[1]).to(
x.device, non_blocking=True).bool()
return x, x_ids, txt, txt_ids, y, mask_x, mask_txt, x_seq_length
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
FluxMRACEPlus.para_dict,
set_name=True)
@@ -0,0 +1,3 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.backbone.cogvideox.cogvideox import CogVideoXTransformer3DModel
@@ -0,0 +1,357 @@
# -*- coding: utf-8 -*-
# Copyright 2024 The CogVideoX team, Tsinghua University & ZhipuAI and The HuggingFace Team.
# All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from collections import OrderedDict
from typing import Any, Dict, Optional, Tuple, Union
import torch
from torch import nn
from scepter.modules.model.base_model import BaseModel
from scepter.modules.model.registry import BACKBONES
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from .layers import CogVideoXBlock, CogVideoXPatchEmbed, TimestepEmbedding, Timesteps, AdaLayerNorm
@BACKBONES.register_class()
class CogVideoXTransformer3DModel(BaseModel):
"""
A Transformer model for video-like data in [CogVideoX](https://github.com/THUDM/CogVideo).
Parameters:
num_attention_heads (`int`, defaults to `30`):
The number of heads to use for multi-head attention.
attention_head_dim (`int`, defaults to `64`):
The number of channels in each head.
in_channels (`int`, defaults to `16`):
The number of channels in the input.
out_channels (`int`, *optional*, defaults to `16`):
The number of channels in the output.
flip_sin_to_cos (`bool`, defaults to `True`):
Whether to flip the sin to cos in the time embedding.
time_embed_dim (`int`, defaults to `512`):
Output dimension of timestep embeddings.
ofs_embed_dim (`int`, defaults to `512`):
Output dimension of "ofs" embeddings used in CogVideoX-5b-I2B in version 1.5
text_embed_dim (`int`, defaults to `4096`):
Input dimension of text embeddings from the text encoder.
num_layers (`int`, defaults to `30`):
The number of layers of Transformer blocks to use.
dropout (`float`, defaults to `0.0`):
The dropout probability to use.
attention_bias (`bool`, defaults to `True`):
Whether or not to use bias in the attention projection layers.
sample_width (`int`, defaults to `90`):
The width of the input latents.
sample_height (`int`, defaults to `60`):
The height of the input latents.
sample_frames (`int`, defaults to `49`):
The number of frames in the input latents. Note that this parameter was incorrectly initialized to 49
instead of 13 because CogVideoX processed 13 latent frames at once in its default and recommended settings,
but cannot be changed to the correct value to ensure backwards compatibility. To create a transformer with
K latent frames, the correct value to pass here would be: ((K - 1) * temporal_compression_ratio + 1).
patch_size (`int`, defaults to `2`):
The size of the patches to use in the patch embedding layer.
temporal_compression_ratio (`int`, defaults to `4`):
The compression ratio across the temporal dimension. See documentation for `sample_frames`.
max_text_seq_length (`int`, defaults to `226`):
The maximum sequence length of the input text embeddings.
activation_fn (`str`, defaults to `"gelu-approximate"`):
Activation function to use in feed-forward.
timestep_activation_fn (`str`, defaults to `"silu"`):
Activation function to use when generating the timestep embeddings.
norm_elementwise_affine (`bool`, defaults to `True`):
Whether or not to use elementwise affine in normalization layers.
norm_eps (`float`, defaults to `1e-5`):
The epsilon value to use in normalization layers.
spatial_interpolation_scale (`float`, defaults to `1.875`):
Scaling factor to apply in 3D positional embeddings across spatial dimensions.
temporal_interpolation_scale (`float`, defaults to `1.0`):
Scaling factor to apply in 3D positional embeddings across temporal dimensions.
"""
def __init__(
self,
cfg,
logger=None
):
super().__init__(cfg, logger=logger)
num_attention_heads = cfg.get("NUM_ATTENTION_HEADS", 30)
attention_head_dim = cfg.get("ATTENTION_HEAD_DIM", 64)
in_channels = cfg.get("IN_CHANNELS", 16)
out_channels = cfg.get("OUT_CHANNELS", 16)
flip_sin_to_cos = cfg.get("FLIP_SIN_TO_COS", True)
freq_shift = cfg.get("FREQ_SHIFT", 0)
time_embed_dim = cfg.get("TIME_EMBED_DIM", 512)
ofs_embed_dim = cfg.get("OFS_EMBED_DIM", None) # 1.5
text_embed_dim = cfg.get("TEXT_EMBED_DIM", 4096)
num_layers = cfg.get("NUM_LAYERS", 30)
dropout = cfg.get("DROPOUT", 0.0)
attention_bias = cfg.get("ATTENTION_BIAS", True)
sample_width = cfg.get("SAMPLE_WIDTH", 90)
sample_height = cfg.get("SAMPLE_HEIGHT", 60)
sample_frames = cfg.get("SAMPLE_FRAMES", 49)
patch_size = cfg.get("PATCH_SIZE", 2)
patch_size_t = cfg.get("PATCH_SIZE_T", None)
patch_bias = cfg.get("PATCH_BIAS", True)
temporal_compression_ratio = cfg.get("TEMPORAL_COMPRESSION_RATIO", 4)
max_text_seq_length = cfg.get("MAX_TEXT_SEQ_LENGTH", 226)
activation_fn = cfg.get("ACTIVATION_FN", "gelu-approximate")
timestep_activation_fn = cfg.get("TIMESTEP_ACTIVATION_FN", "silu")
norm_elementwise_affine = cfg.get("NORM_ELEMENTWISE_AFFINE", True)
norm_eps = cfg.get("NORM_EPS", 1e-5)
spatial_interpolation_scale = cfg.get("SPATIAL_INTERPOLATION_SCALE", 1.875)
temporal_interpolation_scale = cfg.get("TEMPORAL_INTERPOLATION_SCALE", 1.0)
use_rotary_positional_embeddings = cfg.get("USE_ROTARY_POSITIONAL_EMBEDDINGS", False)
use_learned_positional_embeddings = cfg.get("USE_LEARNED_POSITIONAL_EMBEDDINGS", False)
self.gradient_checkpointing = cfg.get("GRADIENT_CHECKPOINTING", False)
inner_dim = num_attention_heads * attention_head_dim
self.patch_size = patch_size
self.patch_size_t = patch_size_t
self.use_rotary_positional_embeddings = use_rotary_positional_embeddings
if not use_rotary_positional_embeddings and use_learned_positional_embeddings:
raise ValueError(
"There are no CogVideoX checkpoints available with disable rotary embeddings and learned positional "
"embeddings. If you're using a custom model and/or believe this should be supported, please open an "
"issue at https://github.com/huggingface/diffusers/issues."
)
# 1. Patch embedding
self.patch_embed = CogVideoXPatchEmbed(
patch_size=patch_size,
patch_size_t=patch_size_t,
in_channels=in_channels,
embed_dim=inner_dim,
text_embed_dim=text_embed_dim,
bias=patch_bias,
sample_width=sample_width,
sample_height=sample_height,
sample_frames=sample_frames,
temporal_compression_ratio=temporal_compression_ratio,
max_text_seq_length=max_text_seq_length,
spatial_interpolation_scale=spatial_interpolation_scale,
temporal_interpolation_scale=temporal_interpolation_scale,
use_positional_embeddings=not use_rotary_positional_embeddings,
use_learned_positional_embeddings=use_learned_positional_embeddings,
)
self.embedding_dropout = nn.Dropout(dropout)
# 2. Time embeddings and ofs embedding(Only CogVideoX1.5-5B I2V have)
self.time_proj = Timesteps(inner_dim, flip_sin_to_cos, freq_shift)
self.time_embedding = TimestepEmbedding(inner_dim, time_embed_dim, timestep_activation_fn)
self.ofs_proj = None
self.ofs_embedding = None
if ofs_embed_dim:
self.ofs_proj = Timesteps(ofs_embed_dim, flip_sin_to_cos, freq_shift)
self.ofs_embedding = TimestepEmbedding(
ofs_embed_dim, ofs_embed_dim, timestep_activation_fn
) # same as time embeddings, for ofs
# 3. Define spatio-temporal transformers blocks
self.transformer_blocks = nn.ModuleList(
[
CogVideoXBlock(
dim=inner_dim,
num_attention_heads=num_attention_heads,
attention_head_dim=attention_head_dim,
time_embed_dim=time_embed_dim,
dropout=dropout,
activation_fn=activation_fn,
attention_bias=attention_bias,
norm_elementwise_affine=norm_elementwise_affine,
norm_eps=norm_eps,
)
for _ in range(num_layers)
]
)
self.norm_final = nn.LayerNorm(inner_dim, norm_eps, norm_elementwise_affine)
# 4. Output blocks
self.norm_out = AdaLayerNorm(
embedding_dim=time_embed_dim,
output_dim=2 * inner_dim,
norm_elementwise_affine=norm_elementwise_affine,
norm_eps=norm_eps,
chunk_dim=1,
)
if patch_size_t is None:
# For CogVideox 1.0
output_dim = patch_size * patch_size * out_channels
else:
# For CogVideoX 1.5
output_dim = patch_size * patch_size * patch_size_t * out_channels
self.proj_out = nn.Linear(inner_dim, output_dim)
def forward(
self,
x: torch.Tensor = None,
t: Union[int, float, torch.LongTensor] = None,
cond: torch.Tensor = None,
timestep_cond: Optional[torch.Tensor] = None,
ofs: Optional[Union[int, float, torch.LongTensor]] = None,
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
**kwargs
):
if 'image_latent' in kwargs and kwargs['image_latent'] is not None:
hidden_states = torch.cat([x, kwargs['image_latent']], dim=2)
else:
hidden_states = x
timestep = t
encoder_hidden_states = cond
batch_size, num_frames, channels, height, width = hidden_states.shape
# 1. Time embedding
timesteps = timestep
t_emb = self.time_proj(timesteps)
# timesteps does not contain any weights and will always return f32 tensors
# but time_embedding might actually be running in fp16. so we need to cast here.
# there might be better ways to encapsulate this.
t_emb = t_emb.to(dtype=encoder_hidden_states.dtype)
emb = self.time_embedding(t_emb, timestep_cond)
if self.ofs_embedding is not None:
ofs_emb = self.ofs_proj(ofs)
ofs_emb = ofs_emb.to(dtype=hidden_states.dtype)
ofs_emb = self.ofs_embedding(ofs_emb)
emb = emb + ofs_emb
# 2. Patch embedding
hidden_states = self.patch_embed(encoder_hidden_states, hidden_states)
hidden_states = self.embedding_dropout(hidden_states)
text_seq_length = encoder_hidden_states.shape[1]
encoder_hidden_states = hidden_states[:, :text_seq_length]
hidden_states = hidden_states[:, text_seq_length:]
# 3. Transformer blocks
for i, block in enumerate(self.transformer_blocks):
if self.training and self.gradient_checkpointing:
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False}
hidden_states, encoder_hidden_states = torch.utils.checkpoint.checkpoint(
create_custom_forward(block),
hidden_states,
encoder_hidden_states,
emb,
image_rotary_emb,
**ckpt_kwargs,
)
else:
hidden_states, encoder_hidden_states = block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
temb=emb,
image_rotary_emb=image_rotary_emb,
)
if not self.use_rotary_positional_embeddings:
# CogVideoX-2B
hidden_states = self.norm_final(hidden_states)
else:
# CogVideoX-5B
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
hidden_states = self.norm_final(hidden_states)
hidden_states = hidden_states[:, text_seq_length:]
# 4. Final block
hidden_states = self.norm_out(hidden_states, temb=emb)
hidden_states = self.proj_out(hidden_states)
# 5. Unpatchify
# Note: we use `-1` instead of `channels`:
# - It is okay to `channels` use for CogVideoX-2b and CogVideoX-5b (number of input channels is equal to output channels)
# - However, for CogVideoX-5b-I2V also takes concatenated input image latents (number of input channels is twice the output channels)
p = self.patch_size
p_t = self.patch_size_t
if p_t is None:
output = hidden_states.reshape(batch_size, num_frames, height // p, width // p, -1, p, p)
output = output.permute(0, 1, 4, 2, 5, 3, 6).flatten(5, 6).flatten(3, 4)
else:
output = hidden_states.reshape(
batch_size, (num_frames + p_t - 1) // p_t, height // p, width // p, -1, p_t, p, p
)
output = output.permute(0, 1, 5, 4, 2, 6, 3, 7).flatten(6, 7).flatten(4, 5).flatten(1, 2)
return output
def load_pretrained_model(self, pretrained_model):
if pretrained_model is not None:
pretrained_model_list = [pretrained_model] if isinstance(pretrained_model, str) else pretrained_model
ckpt_all = OrderedDict()
for pretrained_model in pretrained_model_list:
with FS.get_from(pretrained_model,
wait_finish=True) as local_model:
if local_model.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
ckpt = load_safetensors(local_model)
else:
ckpt = torch.load(local_model, map_location='cpu', weights_only=True)
ckpt_all.update(ckpt)
missing, unexpected = self.load_state_dict(ckpt_all, strict=False)
if we.rank == 0:
self.logger.info(
f'Restored from {pretrained_model_list} with {len(missing)} missing and {len(unexpected)} unexpected keys'
)
if len(missing) > 0:
self.logger.info(f'Missing Keys:\n {missing}')
if len(unexpected) > 0:
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
CogVideoXTransformer3DModel.para_dict,
set_name=True)
if __name__ == "__main__":
import argparse
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.config import Config
from scepter.modules.utils.logger import get_logger
parser = argparse.ArgumentParser()
cfg = Config(parser_ins=parser)
for file_sys in cfg.FILE_SYSTEM:
FS.init_fs_client(file_sys)
model = BACKBONES.build(cfg.DIFFUSION_MODEL, logger=get_logger()).eval().requires_grad_(False).to('cuda').to(torch.bfloat16)
hidden_states = torch.load(FS.get_from(cfg.HIDDEN_STATES), weights_only=True)
encoder_hidden_states = torch.load(FS.get_from(cfg.ENCODER_HIDDEN_STATES), weights_only=True)
timestep = torch.load(FS.get_from(cfg.TIMESTEP), weights_only=True)
timestep_cond = None
image_rotary_emb = None
attention_kwargs = None
output = model(hidden_states, encoder_hidden_states, timestep, timestep_cond, image_rotary_emb, attention_kwargs)
print(output, torch.sum(output))
@@ -0,0 +1,574 @@
# -*- coding: utf-8 -*-
# Copyright 2024 The CogVideoX team, Tsinghua University & ZhipuAI and The HuggingFace Team.
# All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Optional, Tuple
import torch
from torch import nn
import torch.nn.functional as F
from .utils import get_activation, get_timestep_embedding, get_3d_sincos_pos_embed, apply_rotary_emb
from .utils import GELU, GEGLU, ApproximateGELU, SwiGLU
class TimestepEmbedding(nn.Module):
def __init__(
self,
in_channels: int,
time_embed_dim: int,
act_fn: str = "silu",
out_dim: int = None,
post_act_fn: Optional[str] = None,
cond_proj_dim=None,
sample_proj_bias=True,
):
super().__init__()
self.linear_1 = nn.Linear(in_channels, time_embed_dim, sample_proj_bias)
if cond_proj_dim is not None:
self.cond_proj = nn.Linear(cond_proj_dim, in_channels, bias=False)
else:
self.cond_proj = None
self.act = get_activation(act_fn)
if out_dim is not None:
time_embed_dim_out = out_dim
else:
time_embed_dim_out = time_embed_dim
self.linear_2 = nn.Linear(time_embed_dim, time_embed_dim_out, sample_proj_bias)
if post_act_fn is None:
self.post_act = None
else:
self.post_act = get_activation(post_act_fn)
def forward(self, sample, condition=None):
if condition is not None:
sample = sample + self.cond_proj(condition)
sample = self.linear_1(sample)
if self.act is not None:
sample = self.act(sample)
sample = self.linear_2(sample)
if self.post_act is not None:
sample = self.post_act(sample)
return sample
class Timesteps(nn.Module):
def __init__(self, num_channels: int, flip_sin_to_cos: bool, downscale_freq_shift: float, scale: int = 1):
super().__init__()
self.num_channels = num_channels
self.flip_sin_to_cos = flip_sin_to_cos
self.downscale_freq_shift = downscale_freq_shift
self.scale = scale
def forward(self, timesteps):
t_emb = get_timestep_embedding(
timesteps,
self.num_channels,
flip_sin_to_cos=self.flip_sin_to_cos,
downscale_freq_shift=self.downscale_freq_shift,
scale=self.scale,
)
return t_emb
class CogVideoXLayerNormZero(nn.Module):
def __init__(
self,
conditioning_dim: int,
embedding_dim: int,
elementwise_affine: bool = True,
eps: float = 1e-5,
bias: bool = True,
) -> None:
super().__init__()
self.silu = nn.SiLU()
self.linear = nn.Linear(conditioning_dim, 6 * embedding_dim, bias=bias)
self.norm = nn.LayerNorm(embedding_dim, eps=eps, elementwise_affine=elementwise_affine)
def forward(
self, hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor, temb: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
shift, scale, gate, enc_shift, enc_scale, enc_gate = self.linear(self.silu(temb)).chunk(6, dim=1)
hidden_states = self.norm(hidden_states) * (1 + scale)[:, None, :] + shift[:, None, :]
encoder_hidden_states = self.norm(encoder_hidden_states) * (1 + enc_scale)[:, None, :] + enc_shift[:, None, :]
return hidden_states, encoder_hidden_states, gate[:, None, :], enc_gate[:, None, :]
class AdaLayerNorm(nn.Module):
r"""
Norm layer modified to incorporate timestep embeddings.
Parameters:
embedding_dim (`int`): The size of each embedding vector.
num_embeddings (`int`, *optional*): The size of the embeddings dictionary.
output_dim (`int`, *optional*):
norm_elementwise_affine (`bool`, defaults to `False):
norm_eps (`bool`, defaults to `False`):
chunk_dim (`int`, defaults to `0`):
"""
def __init__(
self,
embedding_dim: int,
num_embeddings: Optional[int] = None,
output_dim: Optional[int] = None,
norm_elementwise_affine: bool = False,
norm_eps: float = 1e-5,
chunk_dim: int = 0,
):
super().__init__()
self.chunk_dim = chunk_dim
output_dim = output_dim or embedding_dim * 2
if num_embeddings is not None:
self.emb = nn.Embedding(num_embeddings, embedding_dim)
else:
self.emb = None
self.silu = nn.SiLU()
self.linear = nn.Linear(embedding_dim, output_dim)
self.norm = nn.LayerNorm(output_dim // 2, norm_eps, norm_elementwise_affine)
def forward(
self, x: torch.Tensor, timestep: Optional[torch.Tensor] = None, temb: Optional[torch.Tensor] = None
) -> torch.Tensor:
if self.emb is not None:
temb = self.emb(timestep)
temb = self.linear(self.silu(temb))
if self.chunk_dim == 1:
# This is a bit weird why we have the order of "shift, scale" here and "scale, shift" in the
# other if-branch. This branch is specific to CogVideoX for now.
shift, scale = temb.chunk(2, dim=1)
shift = shift[:, None, :]
scale = scale[:, None, :]
else:
scale, shift = temb.chunk(2, dim=0)
x = self.norm(x) * (1 + scale) + shift
return x
class CogVideoXPatchEmbed(nn.Module):
def __init__(
self,
patch_size: int = 2,
patch_size_t: Optional[int] = None,
in_channels: int = 16,
embed_dim: int = 1920,
text_embed_dim: int = 4096,
bias: bool = True,
sample_width: int = 90,
sample_height: int = 60,
sample_frames: int = 49,
temporal_compression_ratio: int = 4,
max_text_seq_length: int = 226,
spatial_interpolation_scale: float = 1.875,
temporal_interpolation_scale: float = 1.0,
use_positional_embeddings: bool = True,
use_learned_positional_embeddings: bool = True,
) -> None:
super().__init__()
self.patch_size = patch_size
self.patch_size_t = patch_size_t
self.embed_dim = embed_dim
self.sample_height = sample_height
self.sample_width = sample_width
self.sample_frames = sample_frames
self.temporal_compression_ratio = temporal_compression_ratio
self.max_text_seq_length = max_text_seq_length
self.spatial_interpolation_scale = spatial_interpolation_scale
self.temporal_interpolation_scale = temporal_interpolation_scale
self.use_positional_embeddings = use_positional_embeddings
self.use_learned_positional_embeddings = use_learned_positional_embeddings
if patch_size_t is None:
# CogVideoX 1.0 checkpoints
self.proj = nn.Conv2d(
in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size, bias=bias
)
else:
# CogVideoX 1.5 checkpoints
self.proj = nn.Linear(in_channels * patch_size * patch_size * patch_size_t, embed_dim)
self.text_proj = nn.Linear(text_embed_dim, embed_dim)
if use_positional_embeddings or use_learned_positional_embeddings:
persistent = use_learned_positional_embeddings
pos_embedding = self._get_positional_embeddings(sample_height, sample_width, sample_frames)
self.register_buffer("pos_embedding", pos_embedding, persistent=persistent)
def _get_positional_embeddings(self, sample_height: int, sample_width: int, sample_frames: int) -> torch.Tensor:
post_patch_height = sample_height // self.patch_size
post_patch_width = sample_width // self.patch_size
post_time_compression_frames = (sample_frames - 1) // self.temporal_compression_ratio + 1
num_patches = post_patch_height * post_patch_width * post_time_compression_frames
pos_embedding = get_3d_sincos_pos_embed(
self.embed_dim,
(post_patch_width, post_patch_height),
post_time_compression_frames,
self.spatial_interpolation_scale,
self.temporal_interpolation_scale,
)
pos_embedding = torch.from_numpy(pos_embedding).flatten(0, 1)
joint_pos_embedding = torch.zeros(
1, self.max_text_seq_length + num_patches, self.embed_dim, requires_grad=False
)
joint_pos_embedding.data[:, self.max_text_seq_length :].copy_(pos_embedding)
return joint_pos_embedding
def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor):
r"""
Args:
text_embeds (`torch.Tensor`):
Input text embeddings. Expected shape: (batch_size, seq_length, embedding_dim).
image_embeds (`torch.Tensor`):
Input image embeddings. Expected shape: (batch_size, num_frames, channels, height, width).
"""
text_embeds = self.text_proj(text_embeds)
batch_size, num_frames, channels, height, width = image_embeds.shape
if self.patch_size_t is None:
image_embeds = image_embeds.reshape(-1, channels, height, width)
image_embeds = self.proj(image_embeds)
image_embeds = image_embeds.view(batch_size, num_frames, *image_embeds.shape[1:])
image_embeds = image_embeds.flatten(3).transpose(2, 3) # [batch, num_frames, height x width, channels]
image_embeds = image_embeds.flatten(1, 2) # [batch, num_frames x height x width, channels]
else:
p = self.patch_size
p_t = self.patch_size_t
image_embeds = image_embeds.permute(0, 1, 3, 4, 2)
image_embeds = image_embeds.reshape(
batch_size, num_frames // p_t, p_t, height // p, p, width // p, p, channels
)
image_embeds = image_embeds.permute(0, 1, 3, 5, 7, 2, 4, 6).flatten(4, 7).flatten(1, 3)
image_embeds = self.proj(image_embeds)
embeds = torch.cat(
[text_embeds, image_embeds], dim=1
).contiguous() # [batch, seq_length + num_frames x height x width, channels]
if self.use_positional_embeddings or self.use_learned_positional_embeddings:
if self.use_learned_positional_embeddings and (self.sample_width != width or self.sample_height != height):
raise ValueError(
"It is currently not possible to generate videos at a different resolution that the defaults. This should only be the case with 'THUDM/CogVideoX-5b-I2V'."
"If you think this is incorrect, please open an issue at https://github.com/huggingface/diffusers/issues."
)
pre_time_compression_frames = (num_frames - 1) * self.temporal_compression_ratio + 1
if (
self.sample_height != height
or self.sample_width != width
or self.sample_frames != pre_time_compression_frames
):
pos_embedding = self._get_positional_embeddings(height, width, pre_time_compression_frames)
pos_embedding = pos_embedding.to(embeds.device, dtype=embeds.dtype)
else:
pos_embedding = self.pos_embedding
embeds = embeds + pos_embedding
return embeds
class FeedForward(nn.Module):
r"""
A feed-forward layer.
Parameters:
dim (`int`): The number of channels in the input.
dim_out (`int`, *optional*): The number of channels in the output. If not given, defaults to `dim`.
mult (`int`, *optional*, defaults to 4): The multiplier to use for the hidden dimension.
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
final_dropout (`bool` *optional*, defaults to False): Apply a final dropout.
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
"""
def __init__(
self,
dim: int,
dim_out: Optional[int] = None,
mult: int = 4,
dropout: float = 0.0,
activation_fn: str = "geglu",
final_dropout: bool = False,
inner_dim=None,
bias: bool = True,
):
super().__init__()
if inner_dim is None:
inner_dim = int(dim * mult)
dim_out = dim_out if dim_out is not None else dim
if activation_fn == "gelu":
act_fn = GELU(dim, inner_dim, bias=bias)
if activation_fn == "gelu-approximate":
act_fn = GELU(dim, inner_dim, approximate="tanh", bias=bias)
elif activation_fn == "geglu":
act_fn = GEGLU(dim, inner_dim, bias=bias)
elif activation_fn == "geglu-approximate":
act_fn = ApproximateGELU(dim, inner_dim, bias=bias)
elif activation_fn == "swiglu":
act_fn = SwiGLU(dim, inner_dim, bias=bias)
self.net = nn.ModuleList([])
# project in
self.net.append(act_fn)
# project dropout
self.net.append(nn.Dropout(dropout))
# project out
self.net.append(nn.Linear(inner_dim, dim_out, bias=bias))
# FF as used in Vision Transformer, MLP-Mixer, etc. have a final dropout
if final_dropout:
self.net.append(nn.Dropout(dropout))
def forward(self, hidden_states: torch.Tensor, *args, **kwargs) -> torch.Tensor:
if len(args) > 0 or kwargs.get("scale", None) is not None:
deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`."
print(deprecation_message)
for module in self.net:
hidden_states = module(hidden_states)
return hidden_states
class Attention(nn.Module):
def __init__(
self,
query_dim: int,
dim_head: int = 64,
heads: int = 8,
kv_heads: Optional[int] = None,
qk_norm: Optional[str] = None,
eps: float = 1e-5,
bias: bool = False,
out_bias: bool = True,
dropout: float = 0.0,
out_dim: int = None,
cross_attention_dim: Optional[int] = None,
):
super().__init__()
self.inner_dim = out_dim if out_dim is not None else dim_head * heads
self.inner_kv_dim = self.inner_dim if kv_heads is None else dim_head * kv_heads
self.query_dim = query_dim
self.cross_attention_dim = cross_attention_dim if cross_attention_dim is not None else query_dim
self.is_cross_attention = cross_attention_dim is not None
self.out_dim = out_dim if out_dim is not None else query_dim
self.heads = out_dim // dim_head if out_dim is not None else heads
if qk_norm is None:
self.norm_q = None
self.norm_k = None
elif qk_norm == "layer_norm":
self.norm_q = nn.LayerNorm(dim_head, eps=eps)
self.norm_k = nn.LayerNorm(dim_head, eps=eps)
self.to_q = nn.Linear(query_dim, self.inner_dim, bias=bias)
self.to_k = nn.Linear(self.cross_attention_dim, self.inner_kv_dim, bias=bias)
self.to_v = nn.Linear(self.cross_attention_dim, self.inner_kv_dim, bias=bias)
self.to_out = nn.ModuleList([])
self.to_out.append(nn.Linear(self.inner_dim, self.out_dim, bias=out_bias))
self.to_out.append(nn.Dropout(dropout))
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
image_rotary_emb: Optional[torch.Tensor] = None,
) -> torch.Tensor:
text_seq_length = encoder_hidden_states.size(1)
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
batch_size, sequence_length, _ = (
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
)
query = self.to_q(hidden_states)
key = self.to_k(hidden_states)
value = self.to_v(hidden_states)
inner_dim = key.shape[-1]
head_dim = inner_dim // self.heads
query = query.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
key = key.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
value = value.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
if self.norm_q is not None:
query = self.norm_q(query)
if self.norm_k is not None:
key = self.norm_k(key)
# Apply RoPE if needed
if image_rotary_emb is not None:
query[:, :, text_seq_length:] = apply_rotary_emb(query[:, :, text_seq_length:], image_rotary_emb)
if not self.is_cross_attention:
key[:, :, text_seq_length:] = apply_rotary_emb(key[:, :, text_seq_length:], image_rotary_emb)
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.heads * head_dim)
# linear proj
hidden_states = self.to_out[0](hidden_states)
# dropout
hidden_states = self.to_out[1](hidden_states)
encoder_hidden_states, hidden_states = hidden_states.split(
[text_seq_length, hidden_states.size(1) - text_seq_length], dim=1
)
return hidden_states, encoder_hidden_states
class CogVideoXBlock(nn.Module):
r"""
Transformer block used in [CogVideoX](https://github.com/THUDM/CogVideo) model.
Parameters:
dim (`int`):
The number of channels in the input and output.
num_attention_heads (`int`):
The number of heads to use for multi-head attention.
attention_head_dim (`int`):
The number of channels in each head.
time_embed_dim (`int`):
The number of channels in timestep embedding.
dropout (`float`, defaults to `0.0`):
The dropout probability to use.
activation_fn (`str`, defaults to `"gelu-approximate"`):
Activation function to be used in feed-forward.
attention_bias (`bool`, defaults to `False`):
Whether or not to use bias in attention projection layers.
qk_norm (`bool`, defaults to `True`):
Whether or not to use normalization after query and key projections in Attention.
norm_elementwise_affine (`bool`, defaults to `True`):
Whether to use learnable elementwise affine parameters for normalization.
norm_eps (`float`, defaults to `1e-5`):
Epsilon value for normalization layers.
final_dropout (`bool` defaults to `False`):
Whether to apply a final dropout after the last feed-forward layer.
ff_inner_dim (`int`, *optional*, defaults to `None`):
Custom hidden dimension of Feed-forward layer. If not provided, `4 * dim` is used.
ff_bias (`bool`, defaults to `True`):
Whether or not to use bias in Feed-forward layer.
attention_out_bias (`bool`, defaults to `True`):
Whether or not to use bias in Attention output projection layer.
"""
def __init__(
self,
dim: int,
num_attention_heads: int,
attention_head_dim: int,
time_embed_dim: int,
dropout: float = 0.0,
activation_fn: str = "gelu-approximate",
attention_bias: bool = False,
qk_norm: bool = True,
norm_elementwise_affine: bool = True,
norm_eps: float = 1e-5,
final_dropout: bool = True,
ff_inner_dim: Optional[int] = None,
ff_bias: bool = True,
attention_out_bias: bool = True,
):
super().__init__()
# 1. Self Attention
self.norm1 = CogVideoXLayerNormZero(time_embed_dim, dim, norm_elementwise_affine, norm_eps, bias=True)
self.attn1 = Attention(
query_dim=dim,
dim_head=attention_head_dim,
heads=num_attention_heads,
qk_norm="layer_norm" if qk_norm else None,
eps=1e-6,
bias=attention_bias,
out_bias=attention_out_bias
)
# 2. Feed Forward
self.norm2 = CogVideoXLayerNormZero(time_embed_dim, dim, norm_elementwise_affine, norm_eps, bias=True)
self.ff = FeedForward(
dim,
dropout=dropout,
activation_fn=activation_fn,
final_dropout=final_dropout,
inner_dim=ff_inner_dim,
bias=ff_bias,
)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
) -> torch.Tensor:
text_seq_length = encoder_hidden_states.size(1)
# norm & modulate
norm_hidden_states, norm_encoder_hidden_states, gate_msa, enc_gate_msa = self.norm1(
hidden_states, encoder_hidden_states, temb
)
# attention
attn_hidden_states, attn_encoder_hidden_states = self.attn1(
hidden_states=norm_hidden_states,
encoder_hidden_states=norm_encoder_hidden_states,
image_rotary_emb=image_rotary_emb,
)
hidden_states = hidden_states + gate_msa * attn_hidden_states
encoder_hidden_states = encoder_hidden_states + enc_gate_msa * attn_encoder_hidden_states
# norm & modulate
norm_hidden_states, norm_encoder_hidden_states, gate_ff, enc_gate_ff = self.norm2(
hidden_states, encoder_hidden_states, temb
)
# feed-forward
norm_hidden_states = torch.cat([norm_encoder_hidden_states, norm_hidden_states], dim=1)
ff_output = self.ff(norm_hidden_states)
hidden_states = hidden_states + gate_ff * ff_output[:, text_seq_length:]
encoder_hidden_states = encoder_hidden_states + enc_gate_ff * ff_output[:, :text_seq_length]
return hidden_states, encoder_hidden_states
@@ -0,0 +1,570 @@
# -*- coding: utf-8 -*-
# Copyright 2024 The CogVideoX team, Tsinghua University & ZhipuAI and The HuggingFace Team.
# All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import math
from typing import Optional, Tuple, Union, List
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
ACTIVATION_FUNCTIONS = {
"swish": nn.SiLU(),
"silu": nn.SiLU(),
"mish": nn.Mish(),
"gelu": nn.GELU(),
"relu": nn.ReLU(),
}
def get_activation(act_fn: str) -> nn.Module:
"""Helper function to get activation function from string.
Args:
act_fn (str): Name of activation function.
Returns:
nn.Module: Activation function.
"""
act_fn = act_fn.lower()
if act_fn in ACTIVATION_FUNCTIONS:
return ACTIVATION_FUNCTIONS[act_fn]
else:
raise ValueError(f"Unsupported activation function: {act_fn}")
class FP32SiLU(nn.Module):
r"""
SiLU activation function with input upcasted to torch.float32.
"""
def __init__(self):
super().__init__()
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
return F.silu(inputs.float(), inplace=False).to(inputs.dtype)
class GELU(nn.Module):
r"""
GELU activation function with tanh approximation support with `approximate="tanh"`.
Parameters:
dim_in (`int`): The number of channels in the input.
dim_out (`int`): The number of channels in the output.
approximate (`str`, *optional*, defaults to `"none"`): If `"tanh"`, use tanh approximation.
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
"""
def __init__(self, dim_in: int, dim_out: int, approximate: str = "none", bias: bool = True):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out, bias=bias)
self.approximate = approximate
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
if gate.device.type != "mps":
return F.gelu(gate, approximate=self.approximate)
# mps: gelu is not implemented for float16
return F.gelu(gate.to(dtype=torch.float32), approximate=self.approximate).to(dtype=gate.dtype)
def forward(self, hidden_states):
hidden_states = self.proj(hidden_states)
hidden_states = self.gelu(hidden_states)
return hidden_states
class GEGLU(nn.Module):
r"""
A [variant](https://arxiv.org/abs/2002.05202) of the gated linear unit activation function.
Parameters:
dim_in (`int`): The number of channels in the input.
dim_out (`int`): The number of channels in the output.
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
"""
def __init__(self, dim_in: int, dim_out: int, bias: bool = True):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out * 2, bias=bias)
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
if gate.device.type != "mps":
return F.gelu(gate)
# mps: gelu is not implemented for float16
return F.gelu(gate.to(dtype=torch.float32)).to(dtype=gate.dtype)
def forward(self, hidden_states, *args, **kwargs):
if len(args) > 0 or kwargs.get("scale", None) is not None:
deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`."
print("scale", "1.0.0", deprecation_message)
hidden_states = self.proj(hidden_states)
hidden_states, gate = hidden_states.chunk(2, dim=-1)
return hidden_states * self.gelu(gate)
class SwiGLU(nn.Module):
r"""
A [variant](https://arxiv.org/abs/2002.05202) of the gated linear unit activation function. It's similar to `GEGLU`
but uses SiLU / Swish instead of GeLU.
Parameters:
dim_in (`int`): The number of channels in the input.
dim_out (`int`): The number of channels in the output.
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
"""
def __init__(self, dim_in: int, dim_out: int, bias: bool = True):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out * 2, bias=bias)
self.activation = nn.SiLU()
def forward(self, hidden_states):
hidden_states = self.proj(hidden_states)
hidden_states, gate = hidden_states.chunk(2, dim=-1)
return hidden_states * self.activation(gate)
class ApproximateGELU(nn.Module):
r"""
The approximate form of the Gaussian Error Linear Unit (GELU). For more details, see section 2 of this
[paper](https://arxiv.org/abs/1606.08415).
Parameters:
dim_in (`int`): The number of channels in the input.
dim_out (`int`): The number of channels in the output.
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
"""
def __init__(self, dim_in: int, dim_out: int, bias: bool = True):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out, bias=bias)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.proj(x)
return x * torch.sigmoid(1.702 * x)
def randn_tensor(
shape: Union[Tuple, List],
generator: Optional[Union[List["torch.Generator"], "torch.Generator"]] = None,
device: Optional["torch.device"] = None,
dtype: Optional["torch.dtype"] = None,
layout: Optional["torch.layout"] = None,
):
"""A helper function to create random tensors on the desired `device` with the desired `dtype`. When
passing a list of generators, you can seed each batch size individually. If CPU generators are passed, the tensor
is always created on the CPU.
"""
# device on which tensor is created defaults to device
rand_device = device
batch_size = shape[0]
layout = layout or torch.strided
device = device or torch.device("cpu")
if generator is not None:
gen_device_type = generator.device.type if not isinstance(generator, list) else generator[0].device.type
if gen_device_type != device.type and gen_device_type == "cpu":
rand_device = "cpu"
if device != "mps":
print(
f"The passed generator was created on 'cpu' even though a tensor on {device} was expected."
f" Tensors will be created on 'cpu' and then moved to {device}. Note that one can probably"
f" slighly speed up this function by passing a generator that was created on the {device} device."
)
elif gen_device_type != device.type and gen_device_type == "cuda":
raise ValueError(f"Cannot generate a {device} tensor from a generator of type {gen_device_type}.")
# make sure generator list of length 1 is treated like a non-list
if isinstance(generator, list) and len(generator) == 1:
generator = generator[0]
if isinstance(generator, list):
shape = (1,) + shape[1:]
latents = [
torch.randn(shape, generator=generator[i], device=rand_device, dtype=dtype, layout=layout)
for i in range(batch_size)
]
latents = torch.cat(latents, dim=0).to(device)
else:
latents = torch.randn(shape, generator=generator, device=rand_device, dtype=dtype, layout=layout).to(device)
return latents
def get_timestep_embedding(
timesteps: torch.Tensor,
embedding_dim: int,
flip_sin_to_cos: bool = False,
downscale_freq_shift: float = 1,
scale: float = 1,
max_period: int = 10000,
):
"""
This matches the implementation in Denoising Diffusion Probabilistic Models: Create sinusoidal timestep embeddings.
Args
timesteps (torch.Tensor):
a 1-D Tensor of N indices, one per batch element. These may be fractional.
embedding_dim (int):
the dimension of the output.
flip_sin_to_cos (bool):
Whether the embedding order should be `cos, sin` (if True) or `sin, cos` (if False)
downscale_freq_shift (float):
Controls the delta between frequencies between dimensions
scale (float):
Scaling factor applied to the embeddings.
max_period (int):
Controls the maximum frequency of the embeddings
Returns
torch.Tensor: an [N x dim] Tensor of positional embeddings.
"""
assert len(timesteps.shape) == 1, "Timesteps should be a 1d-array"
half_dim = embedding_dim // 2
exponent = -math.log(max_period) * torch.arange(
start=0, end=half_dim, dtype=torch.float32, device=timesteps.device
)
exponent = exponent / (half_dim - downscale_freq_shift)
emb = torch.exp(exponent)
emb = timesteps[:, None].float() * emb[None, :]
# scale embeddings
emb = scale * emb
# concat sine and cosine embeddings
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1)
# flip sine and cosine embeddings
if flip_sin_to_cos:
emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1)
# zero pad
if embedding_dim % 2 == 1:
emb = torch.nn.functional.pad(emb, (0, 1, 0, 0))
return emb
def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
"""
embed_dim: output dimension for each position pos: a list of positions to be encoded: size (M,) out: (M, D)
"""
if embed_dim % 2 != 0:
raise ValueError("embed_dim must be divisible by 2")
omega = np.arange(embed_dim // 2, dtype=np.float64)
omega /= embed_dim / 2.0
omega = 1.0 / 10000**omega # (D/2,)
pos = pos.reshape(-1) # (M,)
out = np.einsum("m,d->md", pos, omega) # (M, D/2), outer product
emb_sin = np.sin(out) # (M, D/2)
emb_cos = np.cos(out) # (M, D/2)
emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D)
return emb
def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):
if embed_dim % 2 != 0:
raise ValueError("embed_dim must be divisible by 2")
# use half of dimensions to encode grid_h
emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2)
emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2)
emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D)
return emb
def get_3d_sincos_pos_embed(
embed_dim: int,
spatial_size: Union[int, Tuple[int, int]],
temporal_size: int,
spatial_interpolation_scale: float = 1.0,
temporal_interpolation_scale: float = 1.0,
) -> np.ndarray:
r"""
Args:
embed_dim (`int`):
spatial_size (`int` or `Tuple[int, int]`):
temporal_size (`int`):
spatial_interpolation_scale (`float`, defaults to 1.0):
temporal_interpolation_scale (`float`, defaults to 1.0):
"""
if embed_dim % 4 != 0:
raise ValueError("`embed_dim` must be divisible by 4")
if isinstance(spatial_size, int):
spatial_size = (spatial_size, spatial_size)
embed_dim_spatial = 3 * embed_dim // 4
embed_dim_temporal = embed_dim // 4
# 1. Spatial
grid_h = np.arange(spatial_size[1], dtype=np.float32) / spatial_interpolation_scale
grid_w = np.arange(spatial_size[0], dtype=np.float32) / spatial_interpolation_scale
grid = np.meshgrid(grid_w, grid_h) # here w goes first
grid = np.stack(grid, axis=0)
grid = grid.reshape([2, 1, spatial_size[1], spatial_size[0]])
pos_embed_spatial = get_2d_sincos_pos_embed_from_grid(embed_dim_spatial, grid)
# 2. Temporal
grid_t = np.arange(temporal_size, dtype=np.float32) / temporal_interpolation_scale
pos_embed_temporal = get_1d_sincos_pos_embed_from_grid(embed_dim_temporal, grid_t)
# 3. Concat
pos_embed_spatial = pos_embed_spatial[np.newaxis, :, :]
pos_embed_spatial = np.repeat(pos_embed_spatial, temporal_size, axis=0) # [T, H*W, D // 4 * 3]
pos_embed_temporal = pos_embed_temporal[:, np.newaxis, :]
pos_embed_temporal = np.repeat(pos_embed_temporal, spatial_size[0] * spatial_size[1], axis=1) # [T, H*W, D // 4]
pos_embed = np.concatenate([pos_embed_temporal, pos_embed_spatial], axis=-1) # [T, H*W, D]
return pos_embed
def apply_rotary_emb(
x: torch.Tensor,
freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor]],
use_real: bool = True,
use_real_unbind_dim: int = -1,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Apply rotary embeddings to input tensors using the given frequency tensor. This function applies rotary embeddings
to the given query or key 'x' tensors using the provided frequency tensor 'freqs_cis'. The input tensors are
reshaped as complex numbers, and the frequency tensor is reshaped for broadcasting compatibility. The resulting
tensors contain rotary embeddings and are returned as real tensors.
Args:
x (`torch.Tensor`):
Query or key tensor to apply rotary embeddings. [B, H, S, D] xk (torch.Tensor): Key tensor to apply
freqs_cis (`Tuple[torch.Tensor]`): Precomputed frequency tensor for complex exponentials. ([S, D], [S, D],)
Returns:
Tuple[torch.Tensor, torch.Tensor]: Tuple of modified query tensor and key tensor with rotary embeddings.
"""
if use_real:
cos, sin = freqs_cis # [S, D]
cos = cos[None, None]
sin = sin[None, None]
cos, sin = cos.to(x.device), sin.to(x.device)
if use_real_unbind_dim == -1:
# Used for flux, cogvideox, hunyuan-dit
x_real, x_imag = x.reshape(*x.shape[:-1], -1, 2).unbind(-1) # [B, S, H, D//2]
x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(3)
elif use_real_unbind_dim == -2:
# Used for Stable Audio
x_real, x_imag = x.reshape(*x.shape[:-1], 2, -1).unbind(-2) # [B, S, H, D//2]
x_rotated = torch.cat([-x_imag, x_real], dim=-1)
else:
raise ValueError(f"`use_real_unbind_dim={use_real_unbind_dim}` but should be -1 or -2.")
out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype)
return out
else:
# used for lumina
x_rotated = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
freqs_cis = freqs_cis.unsqueeze(2)
x_out = torch.view_as_real(x_rotated * freqs_cis).flatten(3)
return x_out.type_as(x)
def get_1d_rotary_pos_embed(
dim: int,
pos: Union[np.ndarray, int],
theta: float = 10000.0,
use_real=False,
linear_factor=1.0,
ntk_factor=1.0,
repeat_interleave_real=True,
freqs_dtype=torch.float32, # torch.float32, torch.float64 (flux)
):
"""
Precompute the frequency tensor for complex exponentials (cis) with given dimensions.
This function calculates a frequency tensor with complex exponentials using the given dimension 'dim' and the end
index 'end'. The 'theta' parameter scales the frequencies. The returned tensor contains complex values in complex64
data type.
Args:
dim (`int`): Dimension of the frequency tensor.
pos (`np.ndarray` or `int`): Position indices for the frequency tensor. [S] or scalar
theta (`float`, *optional*, defaults to 10000.0):
Scaling factor for frequency computation. Defaults to 10000.0.
use_real (`bool`, *optional*):
If True, return real part and imaginary part separately. Otherwise, return complex numbers.
linear_factor (`float`, *optional*, defaults to 1.0):
Scaling factor for the context extrapolation. Defaults to 1.0.
ntk_factor (`float`, *optional*, defaults to 1.0):
Scaling factor for the NTK-Aware RoPE. Defaults to 1.0.
repeat_interleave_real (`bool`, *optional*, defaults to `True`):
If `True` and `use_real`, real part and imaginary part are each interleaved with themselves to reach `dim`.
Otherwise, they are concateanted with themselves.
freqs_dtype (`torch.float32` or `torch.float64`, *optional*, defaults to `torch.float32`):
the dtype of the frequency tensor.
Returns:
`torch.Tensor`: Precomputed frequency tensor with complex exponentials. [S, D/2]
"""
assert dim % 2 == 0
if isinstance(pos, int):
pos = torch.arange(pos)
if isinstance(pos, np.ndarray):
pos = torch.from_numpy(pos) # type: ignore # [S]
theta = theta * ntk_factor
freqs = (
1.0
/ (theta ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=pos.device)[: (dim // 2)] / dim))
/ linear_factor
) # [D/2]
freqs = torch.outer(pos, freqs) # type: ignore # [S, D/2]
if use_real and repeat_interleave_real:
# flux, hunyuan-dit, cogvideox
freqs_cos = freqs.cos().repeat_interleave(2, dim=1).float() # [S, D]
freqs_sin = freqs.sin().repeat_interleave(2, dim=1).float() # [S, D]
return freqs_cos, freqs_sin
elif use_real:
# stable audio
freqs_cos = torch.cat([freqs.cos(), freqs.cos()], dim=-1).float() # [S, D]
freqs_sin = torch.cat([freqs.sin(), freqs.sin()], dim=-1).float() # [S, D]
return freqs_cos, freqs_sin
else:
# lumina
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64 # [S, D/2]
return freqs_cis
def get_3d_rotary_pos_embed(
embed_dim,
crops_coords,
grid_size,
temporal_size,
theta: int = 10000,
use_real: bool = True,
grid_type: str = "linspace",
max_size: Optional[Tuple[int, int]] = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
"""
RoPE for video tokens with 3D structure.
Args:
embed_dim: (`int`):
The embedding dimension size, corresponding to hidden_size_head.
crops_coords (`Tuple[int]`):
The top-left and bottom-right coordinates of the crop.
grid_size (`Tuple[int]`):
The grid size of the spatial positional embedding (height, width).
temporal_size (`int`):
The size of the temporal dimension.
theta (`float`):
Scaling factor for frequency computation.
grid_type (`str`):
Whether to use "linspace" or "slice" to compute grids.
Returns:
`torch.Tensor`: positional embedding with shape `(temporal_size * grid_size[0] * grid_size[1], embed_dim/2)`.
"""
if use_real is not True:
raise ValueError(" `use_real = False` is not currently supported for get_3d_rotary_pos_embed")
if grid_type == "linspace":
start, stop = crops_coords
grid_size_h, grid_size_w = grid_size
grid_h = np.linspace(start[0], stop[0], grid_size_h, endpoint=False, dtype=np.float32)
grid_w = np.linspace(start[1], stop[1], grid_size_w, endpoint=False, dtype=np.float32)
grid_t = np.arange(temporal_size, dtype=np.float32)
grid_t = np.linspace(0, temporal_size, temporal_size, endpoint=False, dtype=np.float32)
elif grid_type == "slice":
max_h, max_w = max_size
grid_size_h, grid_size_w = grid_size
grid_h = np.arange(max_h, dtype=np.float32)
grid_w = np.arange(max_w, dtype=np.float32)
grid_t = np.arange(temporal_size, dtype=np.float32)
else:
raise ValueError("Invalid value passed for `grid_type`.")
# Compute dimensions for each axis
dim_t = embed_dim // 4
dim_h = embed_dim // 8 * 3
dim_w = embed_dim // 8 * 3
# Temporal frequencies
freqs_t = get_1d_rotary_pos_embed(dim_t, grid_t, use_real=True)
# Spatial frequencies for height and width
freqs_h = get_1d_rotary_pos_embed(dim_h, grid_h, use_real=True)
freqs_w = get_1d_rotary_pos_embed(dim_w, grid_w, use_real=True)
# BroadCast and concatenate temporal and spaial frequencie (height and width) into a 3d tensor
def combine_time_height_width(freqs_t, freqs_h, freqs_w):
freqs_t = freqs_t[:, None, None, :].expand(
-1, grid_size_h, grid_size_w, -1
) # temporal_size, grid_size_h, grid_size_w, dim_t
freqs_h = freqs_h[None, :, None, :].expand(
temporal_size, -1, grid_size_w, -1
) # temporal_size, grid_size_h, grid_size_2, dim_h
freqs_w = freqs_w[None, None, :, :].expand(
temporal_size, grid_size_h, -1, -1
) # temporal_size, grid_size_h, grid_size_2, dim_w
freqs = torch.cat(
[freqs_t, freqs_h, freqs_w], dim=-1
) # temporal_size, grid_size_h, grid_size_w, (dim_t + dim_h + dim_w)
freqs = freqs.view(
temporal_size * grid_size_h * grid_size_w, -1
) # (temporal_size * grid_size_h * grid_size_w), (dim_t + dim_h + dim_w)
return freqs
t_cos, t_sin = freqs_t # both t_cos and t_sin has shape: temporal_size, dim_t
h_cos, h_sin = freqs_h # both h_cos and h_sin has shape: grid_size_h, dim_h
w_cos, w_sin = freqs_w # both w_cos and w_sin has shape: grid_size_w, dim_w
if grid_type == "slice":
t_cos, t_sin = t_cos[:temporal_size], t_sin[:temporal_size]
h_cos, h_sin = h_cos[:grid_size_h], h_sin[:grid_size_h]
w_cos, w_sin = w_cos[:grid_size_w], w_sin[:grid_size_w]
cos = combine_time_height_width(t_cos, h_cos, w_cos)
sin = combine_time_height_width(t_sin, h_sin, w_sin)
return cos, sin
def get_resize_crop_region_for_grid(src, tgt_width, tgt_height):
tw = tgt_width
th = tgt_height
h, w = src
r = h / w
if r > (th / tw):
resize_height = th
resize_width = int(round(th / h * w))
else:
resize_width = tw
resize_height = int(round(tw / w * h))
crop_top = int(round((th - resize_height) / 2.0))
crop_left = int(round((tw - resize_width) / 2.0))
return (crop_top, crop_left), (crop_top + resize_height, crop_left + resize_width)
@@ -1,3 +1,3 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from .flux import Flux
from .flux import Flux, FluxMR, FluxMRFill, FluxMRRedux, FluxMRControl
+624 -63
View File
@@ -1,6 +1,9 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
# This file contains code that is adapted from
# https://github.com/black-forest-labs/flux.git
import math
from collections import OrderedDict
from functools import partial
import torch
@@ -12,11 +15,9 @@ from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from torch import Tensor, nn
from torch.utils.checkpoint import checkpoint_sequential
from torch.nn.utils.rnn import pad_sequence
from .layers import (DoubleStreamBlock, EmbedND, LastLayer, MLPEmbedder,
SingleStreamBlock, timestep_embedding)
@BACKBONES.register_class()
class Flux(BaseModel):
"""
@@ -98,7 +99,14 @@ class Flux(BaseModel):
qkv_bias = cfg.QKV_BIAS
depth = cfg.DEPTH
depth_single_blocks = cfg.DEPTH_SINGLE_BLOCKS
self.use_grad_checkpoint = cfg.get('USE_GRAD_CHECKPOINT', False)
self.use_grad_checkpoint = cfg.get("USE_GRAD_CHECKPOINT", False)
self.attn_backend = cfg.get("ATTN_BACKEND", "pytorch")
self.cache_pretrain_model = cfg.get("CACHE_PRETRAIN_MODEL", False)
self.lora_model = cfg.get("DIFFUSERS_LORA_MODEL", None)
self.comfyui_lora_model = cfg.get("COMFYUI_LORA_MODEL", None)
self.swift_lora_model = cfg.get("SWIFT_LORA_MODEL", None)
self.blackforest_lora_model = cfg.get("BLACKFOREST_LORA_MODEL", None)
self.pretrain_adapter = cfg.get("PRETRAIN_ADAPTER", None)
if hidden_size % num_heads != 0:
raise ValueError(
@@ -119,85 +127,350 @@ class Flux(BaseModel):
if self.guidance_embed else nn.Identity())
self.txt_in = nn.Linear(context_in_dim, self.hidden_size)
self.double_blocks = nn.ModuleList([
DoubleStreamBlock(
self.hidden_size,
self.num_heads,
mlp_ratio=mlp_ratio,
qkv_bias=qkv_bias,
) for _ in range(depth)
])
self.double_blocks = nn.ModuleList(
[
DoubleStreamBlock(
self.hidden_size,
self.num_heads,
mlp_ratio=mlp_ratio,
qkv_bias=qkv_bias,
backend=self.attn_backend
)
for _ in range(depth)
]
)
self.single_blocks = nn.ModuleList([
SingleStreamBlock(self.hidden_size,
self.num_heads,
mlp_ratio=mlp_ratio)
for _ in range(depth_single_blocks)
])
self.single_blocks = nn.ModuleList(
[
SingleStreamBlock(self.hidden_size, self.num_heads, mlp_ratio=mlp_ratio, backend=self.attn_backend)
for _ in range(depth_single_blocks)
]
)
self.final_layer = LastLayer(self.hidden_size, 1, self.out_channels)
def prepare_input(self, x, context, y, x_shape=None):
# x.shape [6, 16, 16, 16] target is [6, 16, 768, 1360]
bs, c, h, w = x.shape
x = rearrange(x, 'b c (h ph) (w pw) -> b (h w) (c ph pw)', ph=2, pw=2)
x = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=2, pw=2)
x_id = torch.zeros(h // 2, w // 2, 3)
x_id[..., 1] = x_id[..., 1] + torch.arange(h // 2)[:, None]
x_id[..., 2] = x_id[..., 2] + torch.arange(w // 2)[None, :]
x_ids = repeat(x_id, 'h w c -> b (h w) c', b=bs)
x_ids = repeat(x_id, "h w c -> b (h w) c", b=bs)
txt_ids = torch.zeros(bs, context.shape[1], 3)
return x, x_ids.to(x), context.to(x), txt_ids.to(x), y.to(x), h, w
def unpack(self, x: Tensor, height: int, width: int) -> Tensor:
return rearrange(
x,
'b (h w) (c ph pw) -> b c (h ph) (w pw)',
h=math.ceil(height / 2),
w=math.ceil(width / 2),
"b (h w) (c ph pw) -> b c (h ph) (w pw)",
h=math.ceil(height/2),
w=math.ceil(width/2),
ph=2,
pw=2,
)
def load_pretrained_model(self, pretrained_model):
if next(self.parameters()).device.type == 'meta':
map_location = we.device_id
else:
map_location = 'cpu'
if pretrained_model is not None:
with FS.get_from(pretrained_model,
wait_finish=True) as local_model:
if local_model.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(local_model, device=map_location)
def merge_diffuser_lora(self, ori_sd, lora_sd, scale=1.0):
key_map = {
"single_blocks.{}.linear1.weight": {"key_list": [
["transformer.single_transformer_blocks.{}.attn.to_q.lora_A.weight",
"transformer.single_transformer_blocks.{}.attn.to_q.lora_B.weight", [0, 3072]],
["transformer.single_transformer_blocks.{}.attn.to_k.lora_A.weight",
"transformer.single_transformer_blocks.{}.attn.to_k.lora_B.weight", [3072, 6144]],
["transformer.single_transformer_blocks.{}.attn.to_v.lora_A.weight",
"transformer.single_transformer_blocks.{}.attn.to_v.lora_B.weight", [6144, 9216]],
["transformer.single_transformer_blocks.{}.proj_mlp.lora_A.weight",
"transformer.single_transformer_blocks.{}.proj_mlp.lora_B.weight", [9216, 21504]]
], "num": 38},
"single_blocks.{}.modulation.lin.weight": {"key_list": [
["transformer.single_transformer_blocks.{}.norm.linear.lora_A.weight",
"transformer.single_transformer_blocks.{}.norm.linear.lora_B.weight", [0, 9216]],
], "num": 38},
"single_blocks.{}.linear2.weight": {"key_list": [
["transformer.single_transformer_blocks.{}.proj_out.lora_A.weight",
"transformer.single_transformer_blocks.{}.proj_out.lora_B.weight", [0, 3072]],
], "num": 38},
"double_blocks.{}.txt_attn.qkv.weight": {"key_list": [
["transformer.transformer_blocks.{}.attn.add_q_proj.lora_A.weight",
"transformer.transformer_blocks.{}.attn.add_q_proj.lora_B.weight", [0, 3072]],
["transformer.transformer_blocks.{}.attn.add_k_proj.lora_A.weight",
"transformer.transformer_blocks.{}.attn.add_k_proj.lora_B.weight", [3072, 6144]],
["transformer.transformer_blocks.{}.attn.add_v_proj.lora_A.weight",
"transformer.transformer_blocks.{}.attn.add_v_proj.lora_B.weight", [6144, 9216]],
], "num": 19},
"double_blocks.{}.img_attn.qkv.weight": {"key_list": [
["transformer.transformer_blocks.{}.attn.to_q.lora_A.weight",
"transformer.transformer_blocks.{}.attn.to_q.lora_B.weight", [0, 3072]],
["transformer.transformer_blocks.{}.attn.to_k.lora_A.weight",
"transformer.transformer_blocks.{}.attn.to_k.lora_B.weight", [3072, 6144]],
["transformer.transformer_blocks.{}.attn.to_v.lora_A.weight",
"transformer.transformer_blocks.{}.attn.to_v.lora_B.weight", [6144, 9216]],
], "num": 19},
"double_blocks.{}.img_attn.proj.weight": {"key_list": [
["transformer.transformer_blocks.{}.attn.to_out.0.lora_A.weight",
"transformer.transformer_blocks.{}.attn.to_out.0.lora_B.weight", [0, 3072]]
], "num": 19},
"double_blocks.{}.txt_attn.proj.weight": {"key_list": [
["transformer.transformer_blocks.{}.attn.to_add_out.lora_A.weight",
"transformer.transformer_blocks.{}.attn.to_add_out.lora_B.weight", [0, 3072]]
], "num": 19},
"double_blocks.{}.img_mlp.0.weight": {"key_list": [
["transformer.transformer_blocks.{}.ff.net.0.proj.lora_A.weight",
"transformer.transformer_blocks.{}.ff.net.0.proj.lora_B.weight", [0, 12288]]
], "num": 19},
"double_blocks.{}.img_mlp.2.weight": {"key_list": [
["transformer.transformer_blocks.{}.ff.net.2.lora_A.weight",
"transformer.transformer_blocks.{}.ff.net.2.lora_B.weight", [0, 3072]]
], "num": 19},
"double_blocks.{}.txt_mlp.0.weight": {"key_list": [
["transformer.transformer_blocks.{}.ff_context.net.0.proj.lora_A.weight",
"transformer.transformer_blocks.{}.ff_context.net.0.proj.lora_B.weight", [0, 12288]]
], "num": 19},
"double_blocks.{}.txt_mlp.2.weight": {"key_list": [
["transformer.transformer_blocks.{}.ff_context.net.2.lora_A.weight",
"transformer.transformer_blocks.{}.ff_context.net.2.lora_B.weight", [0, 3072]]
], "num": 19},
"double_blocks.{}.img_mod.lin.weight": {"key_list": [
["transformer.transformer_blocks.{}.norm1.linear.lora_A.weight",
"transformer.transformer_blocks.{}.norm1.linear.lora_B.weight", [0, 18432]]
], "num": 19},
"double_blocks.{}.txt_mod.lin.weight": {"key_list": [
["transformer.transformer_blocks.{}.norm1_context.linear.lora_A.weight",
"transformer.transformer_blocks.{}.norm1_context.linear.lora_B.weight", [0, 18432]]
], "num": 19}
}
cover_lora_keys = set()
cover_ori_keys = set()
for k, v in key_map.items():
key_list = v["key_list"]
block_num = v["num"]
for block_id in range(block_num):
for k_list in key_list:
if k_list[0].format(block_id) in lora_sd and k_list[1].format(block_id) in lora_sd:
cover_lora_keys.add(k_list[0].format(block_id))
cover_lora_keys.add(k_list[1].format(block_id))
current_weight = torch.matmul(lora_sd[k_list[0].format(block_id)].permute(1, 0),
lora_sd[k_list[1].format(block_id)].permute(1, 0)).permute(1, 0)
ori_sd[k.format(block_id)][k_list[2][0]:k_list[2][1], ...] += scale * current_weight
cover_ori_keys.add(k.format(block_id))
# lora_sd.pop(k_list[0].format(block_id))
# lora_sd.pop(k_list[1].format(block_id))
self.logger.info(f"merge_blackforest_lora loads lora'parameters lora-paras: \n"
f"cover-{len(cover_lora_keys)} vs total {len(lora_sd)} \n"
f"cover ori-{len(cover_ori_keys)} vs total {len(ori_sd)}")
return ori_sd
def merge_swift_lora(self, ori_sd, lora_sd, scale = 1.0):
have_lora_keys = {}
for k, v in lora_sd.items():
k = k[len("model."):] if k.startswith("model.") else k
ori_key = k.split("lora")[0] + "weight"
if ori_key not in ori_sd:
raise f"{ori_key} should in the original statedict"
if ori_key not in have_lora_keys:
have_lora_keys[ori_key] = {}
if "lora_A" in k:
have_lora_keys[ori_key]["lora_A"] = v
elif "lora_B" in k:
have_lora_keys[ori_key]["lora_B"] = v
else:
raise NotImplementedError
self.logger.info(f"merge_swift_lora loads lora'parameters {len(have_lora_keys)}")
for key, v in have_lora_keys.items():
current_weight = torch.matmul(v["lora_A"].permute(1, 0), v["lora_B"].permute(1, 0)).permute(1, 0)
ori_sd[key] += scale * current_weight
return ori_sd
def merge_blackforest_lora(self, ori_sd, lora_sd, scale = 1.0):
have_lora_keys = {}
cover_lora_keys = set()
cover_ori_keys = set()
for k, v in lora_sd.items():
if "lora" in k:
ori_key = k.split("lora")[0] + "weight"
if ori_key not in ori_sd:
raise f"{ori_key} should in the original statedict"
if ori_key not in have_lora_keys:
have_lora_keys[ori_key] = {}
if "lora_A" in k:
have_lora_keys[ori_key]["lora_A"] = v
cover_lora_keys.add(k)
cover_ori_keys.add(ori_key)
elif "lora_B" in k:
have_lora_keys[ori_key]["lora_B"] = v
cover_lora_keys.add(k)
cover_ori_keys.add(ori_key)
else:
if k in ori_sd:
ori_sd[k] = v
cover_lora_keys.add(k)
cover_ori_keys.add(k)
else:
sd = torch.load(local_model, map_location=map_location)
missing, unexpected = self.load_state_dict(sd,
strict=False,
assign=True)
print("unsurpport keys: ", k)
self.logger.info(f"merge_blackforest_lora loads lora'parameters lora-paras: \n"
f"cover-{len(cover_lora_keys)} vs total {len(lora_sd)} \n"
f"cover ori-{len(cover_ori_keys)} vs total {len(ori_sd)}")
for key, v in have_lora_keys.items():
current_weight = torch.matmul(v["lora_A"].permute(1, 0), v["lora_B"].permute(1, 0)).permute(1, 0)
# print(key, ori_sd[key].shape, current_weight.shape)
ori_sd[key] += scale * current_weight
return ori_sd
def merge_comfyui_lora(self, ori_sd, lora_sd, scale = 1.0):
ori_key_map = {key.replace("_", ".") : key for key in ori_sd.keys()}
parse_ckpt = OrderedDict()
for k, v in lora_sd.items():
if "alpha" in k:
continue
k = k.replace("lora_unet_", "").replace("_", ".")
map_k = ori_key_map[k.split(".lora")[0] + ".weight"]
if map_k not in parse_ckpt:
parse_ckpt[map_k] = {}
if "lora.up" in k:
parse_ckpt[map_k]["lora_up"] = v
elif "lora.down" in k:
parse_ckpt[map_k]["lora_down"] = v
if self.cache_pretrain_model:
self.lora_dict[self.comfyui_lora_model] = {}
for key, v in parse_ckpt.items():
current_weight = torch.matmul(v["lora_down"].permute(1, 0), v["lora_up"].permute(1, 0)).permute(1, 0)
self.lora_dict[self.comfyui_lora_model] = current_weight
ori_sd[key] += scale * current_weight
return ori_sd
def easy_lora_merge(self, ori_sd, lora_sd, scale = 1.0):
for key, v in lora_sd.items():
ori_sd[key] += scale * v
return ori_sd
def load_pretrained_model(self, pretrained_model, lora_scale = 1.0):
if next(self.parameters()).device.type == 'meta':
map_location = torch.device(we.device_id)
safe_device = we.device_id
else:
map_location = "cpu"
safe_device = "cpu"
if pretrained_model is not None:
if not hasattr(self, "ckpt"):
with FS.get_from(pretrained_model, wait_finish=True) as local_model:
if local_model.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
ckpt = load_safetensors(local_model, device=safe_device)
else:
ckpt = torch.load(local_model, map_location=map_location, weights_only=True)
if "state_dict" in ckpt:
ckpt = ckpt["state_dict"]
if "model" in ckpt:
ckpt = ckpt["model"]["model"]
if self.cache_pretrain_model:
self.ckpt = ckpt
self.lora_dict = {}
else:
ckpt = self.ckpt
new_ckpt = OrderedDict()
for k, v in ckpt.items():
if k in ("img_in.weight"):
model_p = self.state_dict()[k]
if v.shape != model_p.shape:
expanded_state_dict_weight = torch.zeros_like(model_p, device=v.device)
slices = tuple(slice(0, dim) for dim in v.shape)
expanded_state_dict_weight[slices] = v
new_ckpt[k] = expanded_state_dict_weight
else:
new_ckpt[k] = v
else:
new_ckpt[k] = v
if self.lora_model is not None:
with FS.get_from(self.lora_model, wait_finish=True) as local_model:
if local_model.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
lora_sd = load_safetensors(local_model, device=safe_device)
else:
lora_sd = torch.load(local_model, map_location=map_location, weights_only=True)
new_ckpt = self.merge_diffuser_lora(new_ckpt, lora_sd, scale=lora_scale)
if self.swift_lora_model is not None:
if not isinstance(self.swift_lora_model, list):
self.swift_lora_model = [(self.swift_lora_model, 1.0)]
for lora_model in self.swift_lora_model:
if isinstance(lora_model, str):
lora_model = (lora_model, 1.0/len(self.swift_lora_model))
print(lora_model)
self.logger.info(f"load swift lora model: {lora_model}")
with FS.get_from(lora_model[0], wait_finish=True) as local_model:
if local_model.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
lora_sd = load_safetensors(local_model, device=safe_device)
else:
lora_sd = torch.load(local_model, map_location=map_location, weights_only=True)
new_ckpt = self.merge_swift_lora(new_ckpt, lora_sd, scale=lora_model[1])
if self.blackforest_lora_model is not None:
with FS.get_from(self.blackforest_lora_model, wait_finish=True) as local_model:
if local_model.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
lora_sd = load_safetensors(local_model, device=safe_device)
else:
lora_sd = torch.load(local_model, map_location=map_location, weights_only=True)
new_ckpt = self.merge_blackforest_lora(new_ckpt, lora_sd, scale=lora_scale)
if self.comfyui_lora_model is not None:
if hasattr(self, "current_lora") and self.current_lora == self.comfyui_lora_model:
return
if hasattr(self, "lora_dict") and self.comfyui_lora_model in self.lora_dict:
new_ckpt = self.easy_lora_merge(new_ckpt, self.lora_dict[self.comfyui_lora_model], scale=lora_scale)
else:
with FS.get_from(self.comfyui_lora_model, wait_finish=True) as local_model:
if local_model.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
lora_sd = load_safetensors(local_model, device=safe_device)
else:
lora_sd = torch.load(local_model, map_location=map_location, weights_only=True)
new_ckpt = self.merge_comfyui_lora(new_ckpt, lora_sd, scale=lora_scale)
if self.comfyui_lora_model:
self.current_lora = self.comfyui_lora_model
adapter_ckpt = {}
if self.pretrain_adapter is not None:
with FS.get_from(self.pretrain_adapter, wait_finish=True) as local_adapter:
if local_adapter.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
adapter_ckpt = load_safetensors(local_adapter, device=safe_device)
else:
adapter_ckpt = torch.load(local_adapter, map_location=map_location, weights_only=True)
new_ckpt.update(adapter_ckpt)
missing, unexpected = self.load_state_dict(new_ckpt, strict=False, assign=True)
self.logger.info(
f'Restored from {pretrained_model} with {len(missing)} missing and {len(unexpected)} unexpected keys'
)
if len(missing) > 0:
self.logger.info(f'Missing Keys:\n {missing}') # noqa
self.logger.info(f'Missing Keys:\n {missing}')
if len(unexpected) > 0:
self.logger.info(f'\nUnexpected Keys:\n {unexpected}') # noqa
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
def forward(self,
x: Tensor,
t: Tensor,
cond: dict = {},
guidance: Tensor | None = None,
gc_seg: int = 0) -> Tensor:
x, x_ids, txt, txt_ids, y, h, w = self.prepare_input(
x, cond['context'], cond['y'])
def forward(
self,
x: Tensor,
t: Tensor,
cond: dict = {},
guidance: Tensor | None = None,
gc_seg: int = 0
) -> Tensor:
x, x_ids, txt, txt_ids, y, h, w = self.prepare_input(x, cond["context"], cond["y"])
# running on sequences img
x = self.img_in(x)
vec = self.time_in(timestep_embedding(t, 256))
if self.guidance_embed:
if guidance is None:
raise ValueError(
"Didn't get guidance strength for guidance distilled model."
)
raise ValueError("Didn't get guidance strength for guidance distilled model.")
vec = vec + self.guidance_in(timestep_embedding(guidance, 256))
vec = vec + self.vector_in(y)
txt = self.txt_in(txt)
@@ -211,12 +484,11 @@ class Flux(BaseModel):
x = torch.cat((txt, x), 1)
if self.use_grad_checkpoint and gc_seg >= 0:
x = checkpoint_sequential(
functions=[
partial(block, **kwargs) for block in self.double_blocks
],
functions=[partial(block, **kwargs) for block in self.double_blocks],
segments=gc_seg if gc_seg > 0 else len(self.double_blocks),
input=x,
use_reentrant=False)
use_reentrant=False
)
else:
for block in self.double_blocks:
x = block(x, **kwargs)
@@ -228,24 +500,313 @@ class Flux(BaseModel):
if self.use_grad_checkpoint and gc_seg >= 0:
x = checkpoint_sequential(
functions=[
partial(block, **kwargs) for block in self.single_blocks
],
functions=[partial(block, **kwargs) for block in self.single_blocks],
segments=gc_seg if gc_seg > 0 else len(self.single_blocks),
input=x,
use_reentrant=False)
use_reentrant=False
)
else:
for block in self.single_blocks:
x = block(x, **kwargs)
x = x[:, txt.shape[1] :, ...]
x = self.final_layer(x, vec) # (N, T, patch_size ** 2 * out_channels) 6 64 64
x = self.unpack(x, h, w)
return x
@staticmethod
def get_config_template():
return dict_to_yaml('BACKBONE',
__class__.__name__,
Flux.para_dict,
set_name=True)
@BACKBONES.register_class()
class FluxMR(Flux):
def prepare_input(self, x, cond):
if isinstance(cond['context'], list):
context, y = torch.cat(cond["context"], dim=0).to(x), torch.cat(cond["y"], dim=0).to(x)
else:
context, y = cond['context'].to(x), cond['y'].to(x)
batch_frames, batch_frames_ids = [], []
for ix, shape in zip(x, cond["x_shapes"]):
# unpack image from sequence
ix = ix[:, :shape[0] * shape[1]].view(-1, shape[0], shape[1])
c, h, w = ix.shape
ix = rearrange(ix, "c (h ph) (w pw) -> (h w) (c ph pw)", ph=2, pw=2)
ix_id = torch.zeros(h // 2, w // 2, 3)
ix_id[..., 1] = ix_id[..., 1] + torch.arange(h // 2)[:, None]
ix_id[..., 2] = ix_id[..., 2] + torch.arange(w // 2)[None, :]
ix_id = rearrange(ix_id, "h w c -> (h w) c")
batch_frames.append([ix])
batch_frames_ids.append([ix_id])
x_list, x_id_list, mask_x_list, x_seq_length = [], [], [], []
for frames, frame_ids in zip(batch_frames, batch_frames_ids):
proj_frames = []
for idx, one_frame in enumerate(frames):
one_frame = self.img_in(one_frame)
proj_frames.append(one_frame)
ix = torch.cat(proj_frames, dim=0)
if_id = torch.cat(frame_ids, dim=0)
x_list.append(ix)
x_id_list.append(if_id)
mask_x_list.append(torch.ones(ix.shape[0]).to(ix.device, non_blocking=True).bool())
x_seq_length.append(ix.shape[0])
x = pad_sequence(tuple(x_list), batch_first=True)
x_ids = pad_sequence(tuple(x_id_list), batch_first=True).to(x) # [b,pad_seq,2] pad (0.,0.) at dim2
mask_x = pad_sequence(tuple(mask_x_list), batch_first=True)
txt = self.txt_in(context)
txt_ids = torch.zeros(context.shape[0], context.shape[1], 3).to(x)
mask_txt = torch.ones(context.shape[0], context.shape[1]).to(x.device, non_blocking=True).bool()
return x, x_ids, txt, txt_ids, y, mask_x, mask_txt, x_seq_length
def unpack(self, x: Tensor, cond: dict = None, x_seq_length: list = None) -> Tensor:
x_list = []
image_shapes = cond["x_shapes"]
for u, shape, seq_length in zip(x, image_shapes, x_seq_length):
height, width = shape
h, w = math.ceil(height / 2), math.ceil(width / 2)
u = rearrange(
u[seq_length-h*w:seq_length, ...],
"(h w) (c ph pw) -> (h ph w pw) c",
h=h,
w=w,
ph=2,
pw=2,
)
x_list.append(u)
x = pad_sequence(tuple(x_list), batch_first=True).permute(0, 2, 1)
return x
def forward(
self,
x: Tensor,
t: Tensor,
cond: dict = {},
guidance: Tensor | None = None,
gc_seg: int = 0,
**kwargs
) -> Tensor:
x, x_ids, txt, txt_ids, y, mask_x, mask_txt, seq_length_list = self.prepare_input(x, cond)
# running on sequences img
vec = self.time_in(timestep_embedding(t, 256))
if self.guidance_embed and guidance[-1] >= 0:
if guidance is None:
raise ValueError("Didn't get guidance strength for guidance distilled model.")
vec = vec + self.guidance_in(timestep_embedding(guidance, 256))
vec = vec + self.vector_in(y)
ids = torch.cat((txt_ids, x_ids), dim=1)
pe = self.pe_embedder(ids)
mask_aside = torch.cat((mask_txt, mask_x), dim=1)
mask = mask_aside[:, None, :] * mask_aside[:, :, None]
kwargs = dict(
vec=vec,
pe=pe,
mask=mask,
txt_length = txt.shape[1],
)
x = torch.cat((txt, x), 1)
if self.use_grad_checkpoint and gc_seg >= 0:
x = checkpoint_sequential(
functions=[partial(block, **kwargs) for block in self.double_blocks],
segments=gc_seg if gc_seg > 0 else len(self.double_blocks),
input=x,
use_reentrant=False
)
else:
for block in self.double_blocks:
x = block(x, **kwargs)
kwargs = dict(
vec=vec,
pe=pe,
mask=mask,
)
if self.use_grad_checkpoint and gc_seg >= 0:
x = checkpoint_sequential(
functions=[partial(block, **kwargs) for block in self.single_blocks],
segments=gc_seg if gc_seg > 0 else len(self.single_blocks),
input=x,
use_reentrant=False
)
else:
for block in self.single_blocks:
x = block(x, **kwargs)
x = x[:, txt.shape[1]:, ...]
x = self.final_layer(
x, vec) # (N, T, patch_size ** 2 * out_channels) 6 64 64
x = self.unpack(x, h, w)
x = self.final_layer(x, vec) # (N, T, patch_size ** 2 * out_channels) 6 64 64
x = self.unpack(x, cond, seq_length_list)
return x
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
Flux.para_dict,
FluxMR.para_dict,
set_name=True)
@BACKBONES.register_class()
class FluxMRFill(FluxMR):
def __init__(self, cfg, logger = None):
super().__init__(cfg, logger)
def prepare_input(self, x, cond):
context, y = cond["context"], cond["y"]
batch_frames, batch_frames_ids = [], []
for ix, shape, imask, ie, ie_mask in zip(x, cond["x_shapes"], cond["x_mask"],
cond["edit"], cond["edit_mask"]):
# unpack image from sequence
ix = ix[:, :shape[0] * shape[1]].view(-1, shape[0], shape[1])
imask = torch.ones_like(ix[[0], :, :]) if imask is None else imask.squeeze(0)
if len(ie) > 0:
ie = ie[0].squeeze(0)
ie_mask = torch.ones((ix.shape[0] * 4, ix.shape[1], ix.shape[2])) if ie_mask is None else ie_mask[0].squeeze(0)
else:
ie, ie_mask = torch.zeros_like(ix).to(x), torch.ones_like(imask).to(x)
ix = torch.cat([ix, ie, ie_mask], dim=0)
c, h, w = ix.shape
ix = rearrange(ix, "c (h ph) (w pw) -> (h w) (c ph pw)", ph=2, pw=2)
ix_id = torch.zeros(h // 2, w // 2, 3)
ix_id[..., 1] = ix_id[..., 1] + torch.arange(h // 2)[:, None]
ix_id[..., 2] = ix_id[..., 2] + torch.arange(w // 2)[None, :]
ix_id = rearrange(ix_id, "h w c -> (h w) c")
batch_frames.append([ix])
batch_frames_ids.append([ix_id])
x_list, x_id_list, mask_x_list, x_seq_length = [], [], [], []
for frames, frame_ids in zip(batch_frames, batch_frames_ids):
proj_frames = []
for idx, one_frame in enumerate(frames):
one_frame = self.img_in(one_frame)
proj_frames.append(one_frame)
ix = torch.cat(proj_frames, dim=0)
if_id = torch.cat(frame_ids, dim=0)
x_list.append(ix)
x_id_list.append(if_id)
mask_x_list.append(torch.ones(ix.shape[0]).to(ix.device, non_blocking=True).bool())
x_seq_length.append(ix.shape[0])
# if len(x_list) < 1: import pdb;pdb.set_trace()
x = pad_sequence(tuple(x_list), batch_first=True)
x_ids = pad_sequence(tuple(x_id_list), batch_first=True).to(x) # [b,pad_seq,2] pad (0.,0.) at dim2
mask_x = pad_sequence(tuple(mask_x_list), batch_first=True)
# import pdb;pdb.set_trace()
if isinstance(context, list):
txt_list, mask_txt_list, y_list = [], [], []
for sample_id, (ctx, yy) in enumerate(zip(context, y)):
txt_list.append(self.txt_in(ctx.to(x)))
mask_txt_list.append(torch.ones(txt_list[-1].shape[0]).to(ctx.device, non_blocking=True).bool())
y_list.append(yy.to(x))
txt = pad_sequence(tuple(txt_list), batch_first=True)
txt_ids = torch.zeros(txt.shape[0], txt.shape[1], 3).to(x)
mask_txt = pad_sequence(tuple(mask_txt_list), batch_first=True)
y = torch.cat(y_list, dim=0)
assert y.ndim == 2 and txt.ndim == 3
else:
txt = self.txt_in(context)
txt_ids = torch.zeros(context.shape[0], context.shape[1], 3).to(x)
mask_txt = torch.ones(context.shape[0], context.shape[1]).to(x.device, non_blocking=True).bool()
return x, x_ids, txt, txt_ids, y, mask_x, mask_txt, x_seq_length
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
FluxMRFill.para_dict,
set_name=True)
@BACKBONES.register_class()
class FluxMRRedux(FluxMR):
'''
ref_image_siglip + projector
'''
def __init__(self, cfg, logger = None):
super().__init__(cfg, logger)
self.redux_dim = cfg.get("REDUX_DIM", 1152)
self.context_in_dim = cfg.CONTEXT_IN_DIM
self.redux_up = nn.Linear(self.redux_dim, self.context_in_dim * 3)
self.redux_down = nn.Linear(self.context_in_dim * 3, self.context_in_dim)
def prepare_input(self, x, cond):
ref_x = cond.get("ref_x", None)
context, y = torch.cat(cond["context"], dim=0).to(x), torch.cat(cond["y"], dim=0).to(x)
if ref_x is not None:
ref_x = [torch.cat(ref_ix, dim=0).mean(dim=0, keepdim=True) for ref_ix in ref_x]
ref_x = self.redux_down(nn.functional.silu(self.redux_up(torch.cat(ref_x, dim=0))))
context = torch.cat((context, ref_x), dim=-2)
batch_frames, batch_frames_ids = [], []
for ix, shape in zip(x, cond["x_shapes"]):
# unpack image from sequence
ix = ix[:, :shape[0] * shape[1]].view(-1, shape[0], shape[1])
c, h, w = ix.shape
ix = rearrange(ix, "c (h ph) (w pw) -> (h w) (c ph pw)", ph=2, pw=2)
ix_id = torch.zeros(h // 2, w // 2, 3)
ix_id[..., 1] = ix_id[..., 1] + torch.arange(h // 2)[:, None]
ix_id[..., 2] = ix_id[..., 2] + torch.arange(w // 2)[None, :]
ix_id = rearrange(ix_id, "h w c -> (h w) c")
batch_frames.append([ix])
batch_frames_ids.append([ix_id])
x_list, x_id_list, mask_x_list, x_seq_length = [], [], [], []
for frames, frame_ids in zip(batch_frames, batch_frames_ids):
proj_frames = []
for idx, one_frame in enumerate(frames):
one_frame = self.img_in(one_frame)
proj_frames.append(one_frame)
ix = torch.cat(proj_frames, dim=0)
if_id = torch.cat(frame_ids, dim=0)
x_list.append(ix)
x_id_list.append(if_id)
mask_x_list.append(torch.ones(ix.shape[0]).to(ix.device, non_blocking=True).bool())
x_seq_length.append(ix.shape[0])
x = pad_sequence(tuple(x_list), batch_first=True)
x_ids = pad_sequence(tuple(x_id_list), batch_first=True).to(x) # [b,pad_seq,2] pad (0.,0.) at dim2
mask_x = pad_sequence(tuple(mask_x_list), batch_first=True)
txt = self.txt_in(context)
txt_ids = torch.zeros(context.shape[0], context.shape[1], 3).to(x)
mask_txt = torch.ones(context.shape[0], context.shape[1]).to(x.device, non_blocking=True).bool()
return x, x_ids, txt, txt_ids, y, mask_x, mask_txt, x_seq_length
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
FluxMRRedux.para_dict,
set_name=True)
@BACKBONES.register_class()
class FluxMRControl(FluxMR):
'''
cat([x, ie]) ensure the same size bettwn the x and ie
'''
def prepare_input(self, x, cond, *args, **kwargs ):
context, y = torch.cat(cond["context"], dim=0).to(x), torch.cat(cond["y"], dim=0).to(x)
x_list, x_id_list, mask_x_list, x_seq_length = [], [], [], []
for ix, shape, ie in zip(x, cond["x_shapes"], cond["edit"]):
ix = ix[:, :shape[0] * shape[1]].view(-1, shape[0], shape[1])
ix = torch.cat([ix, ie], dim=0)
c, h, w = ix.shape
ix = rearrange(ix, "c (h ph) (w pw) -> (h w) (c ph pw)", ph=2, pw=2)
ix_id = torch.zeros(h // 2, w // 2, 3)
ix_id[..., 1] = ix_id[..., 1] + torch.arange(h // 2)[:, None]
ix_id[..., 2] = ix_id[..., 2] + torch.arange(w // 2)[None, :]
ix_id = rearrange(ix_id, "h w c -> (h w) c")
x_list.append(self.img_in(ix))
x_id_list.append(ix_id)
mask_x_list.append(torch.ones(ix.shape[0]).to(ix.device, non_blocking=True).bool())
x_seq_length.append(ix.shape[0])
# if len(x_list) < 1: import pdb;pdb.set_trace()
x = pad_sequence(tuple(x_list), batch_first=True)
x_ids = pad_sequence(tuple(x_id_list), batch_first=True).to(x) # [b,pad_seq,2] pad (0.,0.) at dim2
mask_x = pad_sequence(tuple(mask_x_list), batch_first=True)
txt = self.txt_in(context)
txt_ids = torch.zeros(context.shape[0], context.shape[1], 3).to(x)
mask_txt = torch.ones(context.shape[0], context.shape[1]).to(x.device, non_blocking=True).bool()
return x, x_ids, txt, txt_ids, y, mask_x, mask_txt, x_seq_length
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
FluxMRControl.para_dict,
set_name=True)
+80 -64
View File
@@ -1,27 +1,71 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
# This file contains code that is adapted from
# https://github.com/black-forest-labs/flux.git
from __future__ import annotations
import math
from dataclasses import dataclass
from torch import Tensor, nn
import torch
from einops import rearrange, repeat
from torch import Tensor, nn
from torch import Tensor
from torch.nn.utils.rnn import pad_sequence
try:
from flash_attn import (
flash_attn_varlen_func
)
FLASHATTN_IS_AVAILABLE = True
except ImportError:
FLASHATTN_IS_AVAILABLE = False
flash_attn_varlen_func = None
def attention(q: Tensor,
k: Tensor,
v: Tensor,
pe: Tensor,
mask: Tensor | None = None) -> Tensor:
def attention(q: Tensor, k: Tensor, v: Tensor, pe: Tensor, mask: Tensor | None = None, backend = 'pytorch') -> Tensor:
q, k = apply_rope(q, k, pe)
x = torch.nn.functional.scaled_dot_product_attention(q,
k,
v,
attn_mask=mask)
x = torch.nan_to_num(x, nan=0.0, posinf=1e10, neginf=-1e10)
x = rearrange(x, 'B H L D -> B L (H D)')
if backend == 'pytorch':
if mask is not None and mask.dtype == torch.bool:
mask = torch.zeros_like(mask).to(q).masked_fill_(mask.logical_not(), -1e20)
x = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask)
# x = torch.nan_to_num(x, nan=0.0, posinf=1e10, neginf=-1e10)
x = rearrange(x, "B H L D -> B L (H D)")
elif backend == 'flash_attn':
# q: (B, H, L, D)
# k: (B, H, S, D) now L = S
# v: (B, H, S, D)
b, h, lq, d = q.shape
_, _, lk, _ = k.shape
q = rearrange(q, "B H L D -> B L H D")
k = rearrange(k, "B H S D -> B S H D")
v = rearrange(v, "B H S D -> B S H D")
if mask is None:
q_lens = torch.tensor([lq] * b, dtype=torch.int32).to(q.device, non_blocking=True)
k_lens = torch.tensor([lk] * b, dtype=torch.int32).to(k.device, non_blocking=True)
else:
q_lens = torch.sum(mask[:, 0, :, 0], dim=1).int()
k_lens = torch.sum(mask[:, 0, 0, :], dim=1).int()
q = torch.cat([q_v[:q_l] for q_v, q_l in zip(q, q_lens)])
k = torch.cat([k_v[:k_l] for k_v, k_l in zip(k, k_lens)])
v = torch.cat([v_v[:v_l] for v_v, v_l in zip(v, k_lens)])
cu_seqlens_q = torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(0, dtype=torch.int32)
cu_seqlens_k = torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(0, dtype=torch.int32)
max_seqlen_q = q_lens.max()
max_seqlen_k = k_lens.max()
x = flash_attn_varlen_func(
q,
k,
v,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
max_seqlen_q=max_seqlen_q,
max_seqlen_k=max_seqlen_k
)
x_list = [x[cu_seqlens_q[i]:cu_seqlens_q[i+1]] for i in range(b)]
x = pad_sequence(tuple(x_list), batch_first=True)
x = rearrange(x, "B L H D -> B L (H D)")
else:
raise NotImplementedError
return x
@@ -173,11 +217,8 @@ class Modulation(nn.Module):
self.multiplier = 6 if double else 3
self.lin = nn.Linear(dim, self.multiplier * dim, bias=True)
def forward(self,
vec: Tensor) -> tuple[ModulationOut, ModulationOut | None]:
out = self.lin(nn.functional.silu(vec))[:,
None, :].chunk(self.multiplier,
dim=-1)
def forward(self, vec: Tensor) -> tuple[ModulationOut, ModulationOut | None]:
out = self.lin(nn.functional.silu(vec))[:, None, :].chunk(self.multiplier, dim=-1)
return (
ModulationOut(*out[:3]),
@@ -186,56 +227,37 @@ class Modulation(nn.Module):
class DoubleStreamBlock(nn.Module):
def __init__(self,
hidden_size: int,
num_heads: int,
mlp_ratio: float,
qkv_bias: bool = False):
def __init__(self, hidden_size: int, num_heads: int, mlp_ratio: float, qkv_bias: bool = False, backend = 'pytorch'):
super().__init__()
mlp_hidden_dim = int(hidden_size * mlp_ratio)
self.num_heads = num_heads
self.hidden_size = hidden_size
self.img_mod = Modulation(hidden_size, double=True)
self.img_norm1 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.img_attn = SelfAttention(dim=hidden_size,
num_heads=num_heads,
qkv_bias=qkv_bias)
self.img_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
self.img_attn = SelfAttention(dim=hidden_size, num_heads=num_heads, qkv_bias=qkv_bias)
self.img_norm2 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.img_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
self.img_mlp = nn.Sequential(
nn.Linear(hidden_size, mlp_hidden_dim, bias=True),
nn.GELU(approximate='tanh'),
nn.GELU(approximate="tanh"),
nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
)
self.backend = backend
self.txt_mod = Modulation(hidden_size, double=True)
self.txt_norm1 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.txt_attn = SelfAttention(dim=hidden_size,
num_heads=num_heads,
qkv_bias=qkv_bias)
self.txt_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
self.txt_attn = SelfAttention(dim=hidden_size, num_heads=num_heads, qkv_bias=qkv_bias)
self.txt_norm2 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.txt_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
self.txt_mlp = nn.Sequential(
nn.Linear(hidden_size, mlp_hidden_dim, bias=True),
nn.GELU(approximate='tanh'),
nn.GELU(approximate="tanh"),
nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
)
def forward(self,
x: Tensor,
vec: Tensor,
pe: Tensor,
mask: Tensor = None,
txt_length=None):
def forward(self, x: Tensor, vec: Tensor, pe: Tensor, mask: Tensor = None, txt_length = None):
img_mod1, img_mod2 = self.img_mod(vec)
txt_mod1, txt_mod2 = self.txt_mod(vec)
@@ -245,19 +267,13 @@ class DoubleStreamBlock(nn.Module):
img_modulated = self.img_norm1(img)
img_modulated = (1 + img_mod1.scale) * img_modulated + img_mod1.shift
img_qkv = self.img_attn.qkv(img_modulated)
img_q, img_k, img_v = rearrange(img_qkv,
'B L (K H D) -> K B H L D',
K=3,
H=self.num_heads)
img_q, img_k, img_v = rearrange(img_qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
img_q, img_k = self.img_attn.norm(img_q, img_k, img_v)
# prepare txt for attention
txt_modulated = self.txt_norm1(txt)
txt_modulated = (1 + txt_mod1.scale) * txt_modulated + txt_mod1.shift
txt_qkv = self.txt_attn.qkv(txt_modulated)
txt_q, txt_k, txt_v = rearrange(txt_qkv,
'B L (K H D) -> K B H L D',
K=3,
H=self.num_heads)
txt_q, txt_k, txt_v = rearrange(txt_qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
txt_q, txt_k = self.txt_attn.norm(txt_q, txt_k, txt_v)
# run actual attention
@@ -266,18 +282,16 @@ class DoubleStreamBlock(nn.Module):
v = torch.cat((txt_v, img_v), dim=2)
if mask is not None:
mask = repeat(mask, 'B L S-> B H L S', H=self.num_heads)
attn = attention(q, k, v, pe=pe, mask=mask)
txt_attn, img_attn = attn[:, :txt.shape[1]], attn[:, txt.shape[1]:]
attn = attention(q, k, v, pe=pe, mask = mask, backend = self.backend)
txt_attn, img_attn = attn[:, : txt.shape[1]], attn[:, txt.shape[1] :]
# calculate the img bloks
img = img + img_mod1.gate * self.img_attn.proj(img_attn)
img = img + img_mod2.gate * self.img_mlp(
(1 + img_mod2.scale) * self.img_norm2(img) + img_mod2.shift)
img = img + img_mod2.gate * self.img_mlp((1 + img_mod2.scale) * self.img_norm2(img) + img_mod2.shift)
# calculate the txt bloks
txt = txt + txt_mod1.gate * self.txt_attn.proj(txt_attn)
txt = txt + txt_mod2.gate * self.txt_mlp(
(1 + txt_mod2.scale) * self.txt_norm2(txt) + txt_mod2.shift)
txt = txt + txt_mod2.gate * self.txt_mlp((1 + txt_mod2.scale) * self.txt_norm2(txt) + txt_mod2.shift)
x = torch.cat((txt, img), 1)
return x
@@ -293,6 +307,7 @@ class SingleStreamBlock(nn.Module):
num_heads: int,
mlp_ratio: float = 4.0,
qk_scale: float | None = None,
backend='pytorch'
):
super().__init__()
self.hidden_dim = hidden_size
@@ -317,6 +332,7 @@ class SingleStreamBlock(nn.Module):
self.mlp_act = nn.GELU(approximate='tanh')
self.modulation = Modulation(hidden_size, double=False)
self.backend = backend
def forward(self,
x: Tensor,
@@ -337,7 +353,7 @@ class SingleStreamBlock(nn.Module):
if mask is not None:
mask = repeat(mask, 'B L S-> B H L S', H=self.num_heads)
# compute attention
attn = attention(q, k, v, pe=pe, mask=mask)
attn = attention(q, k, v, pe=pe, mask=mask, backend=self.backend)
# compute activation in mlp stream, cat again and run second linear layer
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
return x + mod.gate * output
@@ -74,7 +74,7 @@ class VisualTransformer(BaseModel):
with FS.get_from(self.pretrain_path,
wait_finish=True) as local_file:
logger.info(f'Loading checkpoint from {self.pretrain_path}')
visual_pre = torch.load(local_file, map_location='cpu')
visual_pre = torch.load(local_file, map_location='cpu', weights_only=True)
if not use_proj:
visual_pre.pop('proj')
if visual_pre['conv1.weight'].dtype == torch.float16:
@@ -145,7 +145,7 @@ class SomeFTVisualTransformer(BaseModel):
with FS.get_from(self.pretrain_path,
wait_finish=True) as local_file:
logger.info(f'Loading checkpoint from {self.pretrain_path}')
visual_pre = torch.load(local_file, map_location='cpu')
visual_pre = torch.load(local_file, map_location='cpu', weights_only=True)
state_dict_update = self.reformat_state_dict(visual_pre)
self.visual.load_state_dict(state_dict_update, strict=True)
+1 -1
View File
@@ -1136,7 +1136,7 @@ class MMDiT(BaseModel):
from safetensors.torch import load_file as load_safetensors
model = load_safetensors(local_path)
else:
model = torch.load(local_path, map_location='cpu')
model = torch.load(local_path, map_location='cpu', weights_only=True)
if 'state_dict' in model:
model = model['state_dict']
new_ckpt = OrderedDict()
@@ -354,7 +354,7 @@ class PixArt(BaseModel):
def load_pretrained_model(self, pretrained_model):
if pretrained_model:
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
model = torch.load(local_path, map_location='cpu')
model = torch.load(local_path, map_location='cpu', weights_only=True)
if 'state_dict' in model:
model = model['state_dict']
new_ckpt = OrderedDict()
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
@@ -10,7 +10,7 @@ import warnings
import torch
import torch.nn as nn
from torch.cuda import amp
from torch import amp
from torch.nn import functional as F
from torch.nn.utils.rnn import pad_sequence
from tqdm import tqdm
@@ -440,7 +440,7 @@ def multi_head_varlen_attention(q_img,
k = k.type(flash_dtype)
v = v.type(flash_dtype)
with amp.autocast():
with amp.autocast("cuda"):
x = flash_attn_varlen_func(q=q,
k=k,
v=v,
@@ -13,7 +13,7 @@ import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from torch import Tensor
from torch.cuda import amp
from torch import amp
from torch.nn.utils.rnn import pad_sequence
@@ -175,7 +175,7 @@ def frame_unpad(x, shapes):
return torch.concat(frames)
@amp.autocast(enabled=False)
@amp.autocast("cuda", enabled=False)
def rope_params(max_seq_len, dim, theta=10000):
"""
Precompute the frequency tensor for complex exponentials.
@@ -189,7 +189,7 @@ def rope_params(max_seq_len, dim, theta=10000):
return freqs
@amp.autocast(enabled=False)
@amp.autocast("cuda", enabled=False)
def rope_apply(x, grid_sizes, freqs):
"""
x: [B, L, N, C].
@@ -225,7 +225,7 @@ def rope_apply(x, grid_sizes, freqs):
return torch.stack(output)
@amp.autocast(enabled=False)
@amp.autocast("cuda", enabled=False)
def rope_apply_multires_pad(x, x_lens, x_shapes, freqs, pad=True):
"""
x: [B, L, N, C].
@@ -267,7 +267,7 @@ def rope_apply_multires_pad(x, x_lens, x_shapes, freqs, pad=True):
return torch.stack(output) if pad else torch.concat(output)
@amp.autocast(enabled=False)
@amp.autocast("cuda", enabled=False)
def rope_apply_multires(x, x_lens, x_shapes, freqs, pad=True):
"""
x: [B*L, N, C].
@@ -459,7 +459,7 @@ class DiffusionUNet(BaseModel):
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(path)
else:
sd = torch.load(path, map_location='cpu')
sd = torch.load(path, map_location='cpu', weights_only=True)
new_sd = OrderedDict()
for k, v in sd.items():
@@ -1231,7 +1231,7 @@ class LargenUNetXL(DiffusionUNetXL):
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(path)
else:
sd = torch.load(path, map_location='cpu')
sd = torch.load(path, map_location='cpu', weights_only=True)
new_sd = OrderedDict()
for k, v in sd.items():
+12 -1
View File
@@ -3,8 +3,10 @@
import copy
import torch.nn as nn
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import gather_data, we
from scepter.modules.utils.model import get_parameter_dtype
from scepter.modules.utils.probe import (ProbeData, merge_gathered_probe,
register_data)
@@ -43,13 +45,15 @@ class BaseModel(nn.Module):
self._dist_data[key][k] += v
else:
self._dist_data[key][k] = v
def collect_probe(self):
probe_data_dict = self._probe_data
for k, v in self._modules.items():
if isinstance(getattr(self, k), BaseModel):
for kk, vv in getattr(self, k).collect_probe().items():
probe_data_dict[f'{k}/{kk}'] = vv
probe_data_dict[f'{k}/{kk}'] = vv
return probe_data_dict
def probe_data(self):
gather_probe_data = gather_data(self._probe_data)
_dist_data_list = gather_data([self._dist_data])
@@ -97,6 +101,13 @@ class BaseModel(nn.Module):
self._probe_data = {}
return ret_data
@property
def model_dtype(self):
"""
`torch.dtype`: The dtype of the module (assuming that all the module parameters have the same dtype).
"""
return get_parameter_dtype(self)
def clear_probe(self):
self._probe_data.clear()
+24 -4
View File
@@ -1,7 +1,27 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from typing import TYPE_CHECKING
from scepter.modules.utils.import_utils import LazyImportModule
from .diffusions import BaseDiffusion, DiffusionFluxRF
from .samplers import BaseDiffusionSampler, DDIMSampler, FlowEluerSampler
from .schedules import (BaseNoiseScheduler, FlowMatchShiftScheduler,
ScaledLinearScheduler)
if TYPE_CHECKING:
from .diffusions import BaseDiffusion, DiffusionFluxRF
from .samplers import BaseDiffusionSampler, DDIMSampler, FlowEluerSampler
from .schedules import (BaseNoiseScheduler, FlowMatchShiftScheduler,
ScaledLinearScheduler)
else:
_import_structure = {
'diffusions': ['BaseDiffusion', 'DiffusionFluxRF'],
'samplers': ['BaseDiffusionSampler', 'DDIMSampler', 'FlowEluerSampler'],
'schedules': ['BaseNoiseScheduler', 'FlowMatchShiftScheduler',
'ScaledLinearScheduler']
}
import sys
sys.modules[__name__] = LazyImportModule(
__name__,
globals()['__file__'],
_import_structure,
module_spec=__spec__,
extra_objects={},
)
+24 -44
View File
@@ -19,10 +19,6 @@ class BaseDiffusion(object):
para_dict = {
'NOISE_SCHEDULER': {},
'SAMPLER_SCHEDULER': {},
'MIN_SNR_GAMMA': {
'value': None,
'description': 'The minimum SNR gamma value for the loss function.'
},
'PREDICTION_TYPE': {
'value': 'eps',
'description':
@@ -37,8 +33,8 @@ class BaseDiffusion(object):
self.init_params()
def init_params(self):
self.min_snr_gamma = self.cfg.get('MIN_SNR_GAMMA', None)
self.prediction_type = self.cfg.get('PREDICTION_TYPE', 'eps')
self.use_dynamic_cfg = self.cfg.get('USE_DYNAMIC_CFG', False)
self.noise_scheduler = NOISE_SCHEDULERS.build(self.cfg.NOISE_SCHEDULER,
logger=self.logger)
self.sampler_scheduler = NOISE_SCHEDULERS.build(self.cfg.get(
@@ -61,28 +57,29 @@ class BaseDiffusion(object):
model_kwargs={},
steps=20,
sampler=None,
use_dynamic_cfg=False,
guide_scale=None,
guide_rescale=None,
show_progress=False,
return_intermediate=None,
intermediate_callback=None,
reverse_scale = -1.,
x = None,
**kwargs):
assert isinstance(steps, (int, torch.LongTensor))
assert return_intermediate in (None, 'x0', 'xt')
assert isinstance(sampler, (str, dict, Config))
intermediates = []
def callback_fn(x_t, t, sigma=None, alpha=None):
def callback_fn(x_t, t, sigma=None, alpha_bar=None):
timestamp = t
t = t.repeat(len(x_t)).round().long().to(x_t.device)
sigma = sigma.repeat(len(x_t), *([1] * (len(sigma.shape) - 1)))
alpha = alpha.repeat(len(x_t), *([1] * (len(alpha.shape) - 1)))
alpha_bar = alpha_bar.repeat(len(x_t), *([1] * (len(alpha_bar.shape) - 1)))
if guide_scale is None or guide_scale == 1.0:
out = model(x=x_t, t=t, **model_kwargs)
else:
if use_dynamic_cfg:
if self.use_dynamic_cfg:
guidance_scale = 1 + guide_scale * (
(1 - math.cos(math.pi * (
(steps - timestamp.item()) / steps)**5.0)) / 2)
@@ -101,15 +98,12 @@ class BaseDiffusion(object):
if self.prediction_type == 'x0':
x0 = out
elif self.prediction_type == 'eps':
x0 = (x_t - sigma * out) / alpha
x0 = (x_t - sigma * out) / alpha_bar
elif self.prediction_type == 'v':
x0 = alpha * x_t - sigma * out
x0 = alpha_bar * x_t - sigma * out
else:
raise NotImplementedError(
f'prediction_type {self.prediction_type} not implemented')
# print("torch.sum(y_out):", torch.sum(y_out), "torch.sum(u_out):", torch.sum(u_out), "torch.sum(out):",
# torch.sum(out), "torch.sum(x0):", torch.sum(x0), "sigmas", sigma, "alphas", alpha)
return x0
sampler_ins = self.get_sampler(sampler)
@@ -117,12 +111,14 @@ class BaseDiffusion(object):
# this is ignored for schnell
sampler_output = sampler_ins.preprare_sampler(
noise,
x = x,
steps=steps,
reverse_scale= reverse_scale,
prediction_type=self.prediction_type,
scheduler_ins=self.sampler_scheduler,
callback_fn=callback_fn)
for _ in trange(steps, disable=not show_progress):
for _ in trange(sampler_output.steps, disable=not show_progress):
trange.desc = sampler_output.msg
sampler_output = sampler_ins.step(sampler_output)
if return_intermediate == 'x_0':
@@ -145,42 +141,33 @@ class BaseDiffusion(object):
if noise is None:
noise = torch.randn_like(x_0)
schedule_output = self.noise_scheduler.add_noise(x_0, noise, **kwargs)
x_t, t, sigma, alpha = schedule_output.x_t, schedule_output.t, schedule_output.sigma, schedule_output.alpha
x_t, t, sigma, alpha_bar = schedule_output.x_t, schedule_output.t, schedule_output.sigma, schedule_output.alpha_bar
out = model(x=x_t, t=t, **model_kwargs)
# mse loss
target = {
'eps': noise,
'x0': x_0,
'v': alpha * noise - sigma * x_0
'v': alpha_bar * noise - sigma * x_0
}[self.prediction_type]
loss = (out - target).pow(2)
if reduction == 'mean':
loss = loss.flatten(1).mean(dim=1)
if self.min_snr_gamma is not None:
alphas = self.noise_scheduler.alphas.to(x_0.device)[t]
sigmas = self.noise_scheduler.sigmas.pow(2).to(x_0.device)[t]
snrs = (alphas / sigmas).clamp(min=1e-20)
min_snrs = snrs.clamp(max=self.min_snr_gamma)
weights = min_snrs / snrs
else:
weights = 1
loss = loss * weights
return loss
def get_sampler(self, sampler):
if isinstance(sampler, str):
if sampler not in DIFFUSION_SAMPLERS.class_map:
from scepter.modules.utils.import_utils import LazyImportModule
if (not LazyImportModule.get_module_type(('DIFFUSION_SAMPLERS', sampler))) and (
sampler not in DIFFUSION_SAMPLERS.class_map):
if self.logger is not None:
self.logger.info(
f'{sampler} not in the defined samplers list {DIFFUSION_SAMPLERS.class_map.keys()}'
f'{sampler} not in the defined samplers list.'
)
else:
print(
f'{sampler} not in the defined samplers list {DIFFUSION_SAMPLERS.class_map.keys()}'
f'{sampler} not in the defined samplers list.'
)
return None
sampler_cfg = Config(cfg_dict={'NAME': sampler}, load=False)
@@ -248,17 +235,6 @@ class DiffusionFluxRF(BaseDiffusion):
loss = (target - out)**2
if reduction == 'mean':
loss = loss.flatten(1).mean(dim=1)
if self.min_snr_gamma is not None:
alphas = self.noise_scheduler.alphas.to(x_0.device)[t]
sigmas = self.noise_scheduler.sigmas.pow(2).to(x_0.device)[t]
snrs = (alphas / sigmas).clamp(min=1e-20)
min_snrs = snrs.clamp(max=self.min_snr_gamma)
weights = min_snrs / snrs
else:
weights = 1
loss = loss * weights
return loss
@torch.no_grad()
@@ -271,6 +247,8 @@ class DiffusionFluxRF(BaseDiffusion):
show_progress=False,
return_intermediate=None,
intermediate_callback=None,
reverse_scale=-1.,
x=None,
**kwargs):
# sanity check
assert isinstance(steps, (int, torch.LongTensor))
@@ -278,7 +256,7 @@ class DiffusionFluxRF(BaseDiffusion):
assert isinstance(sampler, (str, dict, Config))
intermediates = []
def callback_fn(x_t, t, sigma=None, alpha=None):
def callback_fn(x_t, t, sigma=None, alpha_bar=None):
sigma = torch.full((x_t.shape[0], ),
sigma,
dtype=x_t.dtype,
@@ -291,12 +269,14 @@ class DiffusionFluxRF(BaseDiffusion):
# this is ignored for schnell
sampler_output = sampler_ins.preprare_sampler(
noise,
x=x,
steps=steps,
reverse_scale=reverse_scale,
prediction_type=self.prediction_type,
scheduler_ins=self.sampler_scheduler,
callback_fn=callback_fn)
for _ in trange(steps, disable=not show_progress):
for _ in trange(sampler_output.steps, disable=not show_progress):
trange.desc = sampler_output.msg
sampler_output = sampler_ins.step(sampler_output)
if return_intermediate == 'x_0':
+71 -17
View File
@@ -15,15 +15,18 @@ class SamplerOutput(object):
callback_fn: callable
prediction_type: str
alphas: torch.Tensor
alphas_bar: torch.Tensor
betas: torch.Tensor
sigmas: torch.Tensor
alphas_init: torch.Tensor
alphas_bar_init: torch.Tensor
betas_init: torch.Tensor
sigmas_init: torch.Tensor
ts: torch.Tensor
x_t: torch.Tensor
x_0: torch.Tensor
step: int
steps: int
msg: str
def add_custom_field(self, key: str, value) -> None:
@@ -49,7 +52,7 @@ class BaseDiffusionSampler(object):
self.t_max = self.cfg.get('T_MAX', None)
self.t_min = self.cfg.get('T_MIN', None)
def discretization(self, steps=20, num_timesteps=1000, **kwargs):
def discretization(self, steps=20, num_timesteps=1000, reverse_scale = -1., **kwargs):
# get timesteps
if isinstance(steps, int):
steps += 1 if self.discard_penultimate_step else 0
@@ -74,17 +77,23 @@ class BaseDiffusionSampler(object):
steps = steps.clamp_(t_min, t_max)
elif isinstance(steps, list):
steps = torch.tensor(steps)
timesteps = torch.as_tensor(steps, dtype=torch.float32)
return timesteps
if reverse_scale >=0:
img2img_step = int((1 - reverse_scale) * len(steps))
timesteps = torch.as_tensor(steps[img2img_step:], dtype=torch.float32)
return timesteps
return torch.as_tensor(steps, dtype=torch.float32)
def preprare_sampler(self,
noise,
x=None,
steps=20,
reverse_scale=-1.,
scheduler_ins=None,
prediction_type='',
sigmas=None,
betas=None,
alphas=None,
alphas_bar=None,
callback_fn=None,
**kwargs):
'''
@@ -96,36 +105,52 @@ class BaseDiffusionSampler(object):
4. To ensure the safety of threading, use the instance of SamplerOutput as the manager,
which manage all necessary information.
'''
if reverse_scale >= 0:
assert x is not None
num_timesteps = scheduler_ins.num_timesteps if scheduler_ins is not None else 1000
timestamps = self.discretization(steps,
num_timesteps=num_timesteps,
reverse_scale=reverse_scale,
**kwargs)
alphas = scheduler_ins.t_to_alpha(
timestamps, **kwargs) if scheduler_ins is not None else alphas
alphas_bar = scheduler_ins.t_to_alpha_bar(
timestamps, **kwargs) if scheduler_ins is not None else alphas_bar
betas = scheduler_ins.t_to_beta(
timestamps, **kwargs) if scheduler_ins is not None else betas
sigmas = scheduler_ins.t_to_sigma(
timestamps, **kwargs) if scheduler_ins is not None else sigmas
alphas_init = scheduler_ins.t_to_alpha_init(
timestamps, **kwargs) if scheduler_ins is not None else alphas
alphas_bar_init = scheduler_ins.t_to_alpha_bar_init(
timestamps, **kwargs) if scheduler_ins is not None else alphas_bar
betas_init = scheduler_ins.t_to_beta_init(
timestamps, **kwargs) if scheduler_ins is not None else betas
sigmas_init = scheduler_ins.t_to_sigma_init(
timestamps, **kwargs) if scheduler_ins is not None else sigmas
if reverse_scale >= 0:
x_t = x_0 = scheduler_ins.add_noise(x, noise=noise, t=timestamps[0].repeat(x.size(0)).to(x.device)).x_t if len(timestamps) > 0 else x
else:
x_t = x_0 = noise
# Consider the sigma's list is from sigma_ to zero. the steps equal to len(timestamps)
output = SamplerOutput(callback_fn=callback_fn,
prediction_type=prediction_type,
alphas=alphas,
alphas_bar=alphas_bar,
betas=betas,
sigmas=sigmas,
alphas_init=alphas_init,
alphas_bar_init=alphas_bar_init,
betas_init=betas_init,
sigmas_init=sigmas_init,
ts=timestamps,
x_t=noise,
x_0=noise,
x_t=x_t,
x_0=x_0,
step=0,
msg='step 0')
msg='step 0',
steps=len(timestamps) - 1)
return output
def step(self, sampler_ouput):
@@ -159,22 +184,35 @@ class DDIMSampler(BaseDiffusionSampler):
def preprare_sampler(self,
noise,
x=None,
steps=20,
reverse_scale = -1.,
scheduler_ins=None,
prediction_type='',
sigmas=None,
betas=None,
alphas=None,
alphas_bar=None,
callback_fn=None,
**kwargs):
output = super().preprare_sampler(noise, steps, scheduler_ins,
prediction_type, sigmas, betas,
alphas, callback_fn, **kwargs)
output = super().preprare_sampler(noise,
x = x,
steps = steps,
reverse_scale = reverse_scale,
scheduler_ins = scheduler_ins,
prediction_type = prediction_type,
sigmas = sigmas,
betas = betas,
alphas = alphas,
alphas_bar = alphas_bar,
callback_fn = callback_fn,
**kwargs)
sigmas = output.sigmas
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
sigmas_vp = (sigmas**2 / (1 + sigmas**2))**0.5
sigmas_vp[sigmas == float('inf')] = 1.
output.add_custom_field('sigmas_vp', sigmas_vp)
output.steps += 1
return output
def step(self, sampler_output):
@@ -182,10 +220,10 @@ class DDIMSampler(BaseDiffusionSampler):
step = sampler_output.step
t = sampler_output.ts[step]
sigmas_vp = sampler_output.sigmas_vp.to(x_t.device)
alpha_init = _i(sampler_output.alphas_init, step, x_t[:1])
alpha_bar_init = _i(sampler_output.alphas_bar_init, step, x_t[:1])
sigma_init = _i(sampler_output.sigmas_init, step, x_t[:1])
x = sampler_output.callback_fn(x_t, t, sigma_init, alpha_init)
x = sampler_output.callback_fn(x_t, t, sigma_init, alpha_bar_init)
noise_factor = self.eta * (sigmas_vp[step + 1]**2 /
sigmas_vp[step]**2 *
(1 - (1 - sigmas_vp[step]**2) /
@@ -202,16 +240,19 @@ class DDIMSampler(BaseDiffusionSampler):
return sampler_output
@DIFFUSION_SAMPLERS.register_class('flow_eluer')
@DIFFUSION_SAMPLERS.register_class('flow_euler')
class FlowEluerSampler(BaseDiffusionSampler):
def preprare_sampler(self,
noise,
x=None,
steps=20,
reverse_scale = -1.,
scheduler_ins=None,
prediction_type='',
sigmas=None,
betas=None,
alphas=None,
alphas_bar=None,
callback_fn=None,
**kwargs):
if noise.ndim == 3:
@@ -220,9 +261,18 @@ class FlowEluerSampler(BaseDiffusionSampler):
n, _, h, w = noise.shape
seq_len = (h // 2 * w // 2)
kwargs['seq_len'] = seq_len
output = super().preprare_sampler(noise, steps, scheduler_ins,
prediction_type, sigmas, betas,
alphas, callback_fn, **kwargs)
output = super().preprare_sampler(noise,
x = x,
steps = steps,
reverse_scale = reverse_scale,
scheduler_ins = scheduler_ins,
prediction_type = prediction_type,
sigmas = sigmas,
betas = betas,
alphas = alphas,
alphas_bar = alphas_bar,
callback_fn = callback_fn,
**kwargs)
return output
def step(self, sampler_output):
@@ -241,9 +291,13 @@ class FlowEluerSampler(BaseDiffusionSampler):
sampler_output.msg = f'step {step}, sigma_curr: {sigma_curr}, sigma_prev: {sigma_prev}'
return sampler_output
def discretization(self, steps=20, num_timesteps=1000, **kwargs):
def discretization(self, steps=20, num_timesteps=1000, reverse_scale=-1., **kwargs):
# extra step for zero
timesteps = torch.linspace(num_timesteps, 0, steps + 1)
if reverse_scale >= 0:
img2img_step = int((1 - reverse_scale) * len(timesteps))
timesteps = timesteps[img2img_step:]
return timesteps
return timesteps
@staticmethod
+92 -14
View File
@@ -1,6 +1,7 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import math
import random
from dataclasses import dataclass, field
from typing import Callable
@@ -21,7 +22,7 @@ class ScheduleOutput(object):
x_0: torch.Tensor
t: torch.Tensor
sigma: torch.Tensor
alpha: torch.Tensor
alpha_bar: torch.Tensor
custom_fields: dict = field(default_factory=dict)
def add_custom_field(self, key: str, value) -> None:
@@ -30,6 +31,21 @@ class ScheduleOutput(object):
@NOISE_SCHEDULERS.register_class()
class BaseNoiseScheduler(object):
r'''
In the diffusion model, the parameters related to the noise schedule are alpha, beta,
and sigma. The following are the definitions of the above three parameters, which should
be the basic property for the instance of noise scheduler.
\alpha_{t} = \sqrt{1 - \beta_{t}^2} \alpha is the strength of signal and \beta is the strength of noise
\sigma_{t} = \sqrt{1 - \overline\alpha} = \sqrt{1 - \prod_{i=1}^{t}\alpha^2_{i}} (P(x_{t}|x_{0}) ~ N(\overline\alpha x_{0}, \sigma^2))
\alpha_bar_{t} = \sqrt{\overline\alpha} = \sqrt{\prod_{i=1}^{t}\alpha^2_{i}} (P(x_{t}|x_{0}) ~ N(\overline\alpha x_{0}, \sigma^2))
where sigma_{t} is the var of p(x_{t-1}|x_{t}, x_{0}).
(reference to https://arxiv.org/abs/2010.02502)
let sigma transfer to beta:
square_\beta = 1 - \frac{1 - square_\sigma_{t}}{1 - square_\sigma_{t - 1 }}
'''
para_dict = {
'NUM_TIMESTEPS': {
'value': 1000,
@@ -48,7 +64,7 @@ class BaseNoiseScheduler(object):
self.num_timesteps = self.cfg.get('NUM_TIMESTEPS', 1000)
self._sample_steps = torch.arange(self.num_timesteps,
dtype=torch.float32)
self._sigmas, self._betas, self._alphas, self._timesteps = None, None, None, None
self._sigmas, self._betas, self._alphas, self._alphas_bar, self._timesteps = None, None, None, None, None
def check_function(self):
try:
@@ -128,6 +144,10 @@ class BaseNoiseScheduler(object):
square_beta = self.sigmas_to_square_betas(sigma)
return torch.sqrt(1 - square_beta)
def t_to_alpha_bar(self, t, **kwargs):
sigma = self.t_to_sigma(t)
return torch.sqrt(1 - sigma**2)
def t_to_beta(self, t, **kwargs):
sigma = self.t_to_sigma(t)
square_beta = self.sigmas_to_square_betas(sigma)
@@ -138,11 +158,11 @@ class BaseNoiseScheduler(object):
t = torch.randint(0,
self.num_timesteps, (x_0.shape[0], ),
device=x_0.device).long()
alpha = _i(self.alphas, t, x_0)
alpha = _i(self.alphas_bar, t, x_0)
sigma = _i(self.sigmas, t, x_0)
x_t = alpha * x_0 + sigma * noise
return ScheduleOutput(x_0=x_0, x_t=x_t, t=t, alpha=alpha, sigma=sigma)
return ScheduleOutput(x_0=x_0, x_t=x_t, t=t, alpha_bar=alpha, sigma=sigma)
def t_to_alpha_init(self, t, **kwargs):
indices = t.long()
@@ -153,6 +173,16 @@ class BaseNoiseScheduler(object):
alpha = self.alphas[step_indices].flatten().to(t)
return alpha
def t_to_alpha_bar_init(self, t, **kwargs):
indices = t.long()
indices[indices >= self.num_timesteps] = self.num_timesteps - 1
timesteps = self.timesteps.to(t)[indices]
step_indices = [(self.timesteps.to(t) == t).nonzero().item()
for t in timesteps]
alpha_bar = self.alphas_bar[step_indices].flatten().to(t)
return alpha_bar
def t_to_beta_init(self, t, **kwargs):
indices = t.long()
indices[indices >= self.num_timesteps] = self.num_timesteps - 1
@@ -205,6 +235,10 @@ class BaseNoiseScheduler(object):
def alphas(self):
return self._alphas
@property
def alphas_bar(self):
return self._alphas_bar
@property
def timesteps(self):
return self._timesteps
@@ -221,6 +255,10 @@ class BaseNoiseScheduler(object):
'data': self._alphas.cpu().numpy(),
'label': 'alphas'
}, {
'data': self._alphas_bar.cpu().numpy(),
'label': 'alphas_bar'
},
{
'data': self._timesteps.cpu().numpy() / self.num_timesteps,
'label': 'timesteps'
}]
@@ -280,7 +318,8 @@ class ScaledLinearScheduler(BaseNoiseScheduler):
self.snr_shift_scale,
self.rescale_betas_zero_snr)
self._betas = torch.sqrt(square_betas)
self._alphas = torch.sqrt(1 - self._sigmas**2)
self._alphas = torch.sqrt(1 - square_betas)
self._alphas_bar = torch.sqrt(1 - self._sigmas**2)
self._timesteps = torch.arange(len(self._sigmas), dtype=torch.float32)
@@ -304,7 +343,8 @@ class LinearScheduler(BaseNoiseScheduler):
sigmas = self.betas_to_sigmas(betas)
self._sigmas = sigmas
self._betas = betas
self._alphas = torch.sqrt(1 - sigmas**2)
self._alphas = torch.sqrt(1 - betas**2)
self._alphas_bar = torch.sqrt(1 - sigmas**2)
self._timesteps = torch.arange(len(sigmas), dtype=torch.float32)
@@ -319,7 +359,8 @@ class FlowMatchUniformScheduler(BaseNoiseScheduler):
self._timesteps = timesteps
self._sigmas = self.t_to_sigma(timesteps)
self._betas = torch.sqrt(self.sigmas_to_square_betas(self._sigmas))
self._alphas = torch.sqrt(1 - self.betas**2)
self._alphas = torch.sqrt(1 - self._betas**2)
self._alphas_bar = torch.sqrt(1 - self._sigmas ** 2)
def add_noise(self, x_0, noise=None, t=None, **kwargs):
if t is None:
@@ -332,7 +373,7 @@ class FlowMatchUniformScheduler(BaseNoiseScheduler):
x_t=x_t,
t=t,
sigma=sigma,
alpha=self.t_to_alpha(t))
alpha_bar=self.t_to_alpha_bar(t))
def sigma_to_t(self, sigma, **kwargs):
return sigma * self.num_timesteps
@@ -406,7 +447,7 @@ class FlowMatchShiftScheduler(FlowMatchUniformScheduler):
x_t=x_t,
t=t,
sigma=sigma,
alpha=self.t_to_alpha(t))
alpha_bar=self.t_to_alpha_bar(t))
def sigma_to_t(self, sigma, **kwargs):
t = sigma / (sigma - self.shift * sigma + self.shift)
@@ -443,6 +484,14 @@ class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler):
'MAX_SHIFT': {
'value': 1.15,
'description': 'The max shift factor for the timestamp.'
},
'PRE_T_SAMPLE': {
'value': False,
'description': 'Use pre-sampled timesteps or not, default is False.'
},
'PRE_T_SAMPLE_FOLD': {
'value': 1,
'description': 'The folds of pre-sampled timesteps.'
}
}
@@ -452,6 +501,23 @@ class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler):
self.sigmoid_scale = self.cfg.get('SIGMOID_SCALE', 1)
self.base_shift = self.cfg.get('BASE_SHIFT', 0.5)
self.max_shift = self.cfg.get('MAX_SHIFT', 1.15)
self.pre_t_sample = self.cfg.get('PRE_T_SAMPLE', False)
self.pre_t_sample_fold = self.cfg.get('PRE_T_SAMPLE_FOLD', 1)
if self.pre_t_sample:
t = torch.sigmoid(torch.randn((self.num_timesteps * self.pre_t_sample_fold,)))
# Scale and reverse the values to go from 1000 to 0
timesteps = ((1 - t) * 1000)
# Sort the timesteps in descending order
self.pre_sample_timesteps, _ = torch.sort(timesteps, descending=True)
else:
self.pre_sample_timesteps = None
@property
def pre_timesteps(self):
fold_id = random.randint(0, self.pre_t_sample_fold - 1)
# print("fold_id", fold_id)
return self.pre_sample_timesteps[fold_id::self.pre_t_sample_fold]
def time_shift(self, mu: float, sigma_scale: float, t: Tensor):
return math.exp(mu) / (math.exp(mu) + (1 / t - 1)**sigma_scale)
@@ -476,17 +542,28 @@ class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler):
n, _, h, w = x_0.shape
seq_len = (h // 2 * w // 2)
if t is None:
logits_norm = torch.randn(x_0.shape[0], device=x_0.device)
logits_norm = logits_norm * self.sigmoid_scale # larger scale for more uniform sampling
t = logits_norm.sigmoid() * self.num_timesteps
if self.pre_t_sample:
timestep_indices = torch.randint(
1,
self.num_timesteps - 1,
(x_0.shape[0],)
)
timestep_indices = timestep_indices.long()
t = [self.pre_timesteps[x.item()].to(x_0.device) for x in timestep_indices]
t = torch.stack(t, dim=0)
else:
logits_norm = torch.randn(x_0.shape[0], device=x_0.device)
logits_norm = logits_norm * self.sigmoid_scale # larger scale for more uniform sampling
t = logits_norm.sigmoid() * self.num_timesteps
sigma = self.t_to_sigma(t, seq_len=seq_len)
shape = (x_0.size(0), ) + (1, ) * (x_0.ndim - 1)
# print(sigma)
x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
return ScheduleOutput(x_0=x_0,
x_t=x_t,
t=t,
sigma=sigma,
alpha=self.t_to_alpha(t))
alpha_bar=self.t_to_alpha_bar(t))
def sigma_to_t(self, sigma, **kwargs):
seq_len = kwargs.get('seq_len', 256)
@@ -570,6 +647,7 @@ class FlowMatchSigmaScheduler(FlowMatchUniformScheduler):
(self.shift - 1) * timesteps)
self._betas = torch.sqrt(self.sigmas_to_square_betas(self._sigmas))
self._alphas = torch.sqrt(1 - self.betas**2)
self._alphas_bar = torch.sqrt(1 - self._sigmas ** 2)
def add_noise(self, x_0, noise=None, t=None, **kwargs):
if t is None:
@@ -589,7 +667,7 @@ class FlowMatchSigmaScheduler(FlowMatchUniformScheduler):
x_t=x_t,
t=t,
sigma=sigma,
alpha=self.t_to_alpha(t))
alpha_bar=self.t_to_alpha_bar(t))
def compute_density_for_timestep_sampling(self, t):
"""Compute the density for sampling the timesteps when doing SD3 training.
+27 -5
View File
@@ -1,8 +1,30 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from typing import TYPE_CHECKING
from scepter.modules.utils.import_utils import LazyImportModule
from scepter.modules.model.embedder.embedder import (
ConcatTimestepEmbedderND, FrozenCLIPEmbedder, FrozenCLIPEmbedder2,
FrozenOpenCLIPEmbedder, FrozenOpenCLIPEmbedder2, GeneralConditioner,
IPAdapterPlusEmbedder, RefCrossEmbedder, SD3TextEmbedder, T5EmbedderHF)
from scepter.modules.model.embedder.flux_embedder import HFEmbedder
if TYPE_CHECKING:
from scepter.modules.model.embedder.embedder import (
ConcatTimestepEmbedderND, FrozenCLIPEmbedder, FrozenCLIPEmbedder2,
FrozenOpenCLIPEmbedder, FrozenOpenCLIPEmbedder2, GeneralConditioner,
IPAdapterPlusEmbedder, RefCrossEmbedder, SD3TextEmbedder, T5EmbedderHF)
from scepter.modules.model.embedder.flux_embedder import HFEmbedder
else:
_import_structure = {
'embedder': ['ConcatTimestepEmbedderND', 'FrozenCLIPEmbedder',
'FrozenCLIPEmbedder2', 'FrozenOpenCLIPEmbedder',
'FrozenOpenCLIPEmbedder2', 'GeneralConditioner',
'IPAdapterPlusEmbedder', 'RefCrossEmbedder',
'SD3TextEmbedder', 'T5EmbedderHF'],
'flux_embedder': ['HFEmbedder']
}
import sys
sys.modules[__name__] = LazyImportModule(
__name__,
globals()['__file__'],
_import_structure,
module_spec=__spec__,
extra_objects={},
)
+35 -77
View File
@@ -36,7 +36,8 @@ except Exception as e:
def autocast(f, enabled=True):
def do_autocast(*args, **kwargs):
with torch.cuda.amp.autocast(
with torch.amp.autocast(
"cuda",
enabled=enabled,
dtype=torch.get_autocast_gpu_dtype(),
cache_enabled=torch.is_autocast_cache_enabled(),
@@ -239,7 +240,7 @@ class FrozenOpenCLIPEmbedder(BaseEmbedder):
if cfg.PRETRAINED_MODEL is not None:
with FS.get_from(cfg.PRETRAINED_MODEL,
wait_finish=True) as local_path:
model.load_state_dict(torch.load(local_path), strict=False)
model.load_state_dict(torch.load(local_path, weights_only=True), strict=False)
self.model = model
self.use_grad = cfg.get('USE_GRAD', False)
@@ -538,7 +539,7 @@ class IPAdapterPlusEmbedder(BaseEmbedder):
)
with FS.get_from(cfg.PRETRAINED_MODEL, wait_finish=True) as local_path:
ckpt = torch.load(local_path, map_location='cpu')
ckpt = torch.load(local_path, map_location='cpu', weights_only=True)
self.image_proj_model.load_state_dict(ckpt['image_proj'],
strict=True)
@@ -645,7 +646,7 @@ class GeneralConditioner(BaseEmbedder):
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(path)
else:
sd = torch.load(path, map_location='cpu')
sd = torch.load(path, map_location='cpu', weights_only=True)
new_sd = OrderedDict()
for k, v in sd.items():
ignored = False
@@ -832,22 +833,28 @@ class T5EmbedderHF(BaseEmbedder):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
pretrained_path = cfg.get('PRETRAINED_MODEL', None)
self.t5_dtype = cfg.get('T5_DTYPE', 'float32')
assert pretrained_path
with FS.get_dir_to_local_dir(pretrained_path,
wait_finish=True) as local_path:
self.model = T5EncoderModel.from_pretrained(
local_path,
torch_dtype=getattr(
torch,
'float' if self.t5_dtype == 'float32' else self.t5_dtype))
tokenizer_path = cfg.get('TOKENIZER_PATH', None)
self.length = cfg.get('LENGTH', 77)
self.t5_dtype = cfg.get('T5_DTYPE', 'bfloat16')
self.use_grad = cfg.get('USE_GRAD', False)
self.clean = cfg.get('CLEAN', 'whitespace')
self.added_identifier = cfg.get('ADDED_IDENTIFIER', None)
tokenizer_path = cfg.get('TOKENIZER_PATH', None)
pretrained_path = cfg.get('PRETRAINED_MODEL', None)
if pretrained_path:
with FS.get_dir_to_local_dir(pretrained_path,
wait_finish=True) as local_path:
if self.t5_dtype is not None:
self.model = T5EncoderModel.from_pretrained(
local_path,
torch_dtype=getattr(
torch,
'float' if self.t5_dtype == 'float32' else self.t5_dtype))
else:
self.model = T5EncoderModel.from_pretrained(local_path)
else:
self.model = None
if tokenizer_path:
self.tokenize_kargs = {'return_tensors': 'pt'}
with FS.get_dir_to_local_dir(tokenizer_path,
@@ -869,9 +876,6 @@ class T5EmbedderHF(BaseEmbedder):
self.tokenizer = None
self.tokenize_kargs = {}
self.use_grad = cfg.get('USE_GRAD', False)
self.clean = cfg.get('CLEAN', 'whitespace')
def freeze(self):
self.model = self.model.eval()
for param in self.parameters():
@@ -888,14 +892,10 @@ class T5EmbedderHF(BaseEmbedder):
else:
x = self.model(tokens.input_ids.to(we.device_id))
x = x.last_hidden_state
# if not self.return_pooled:
# return x.detach()
# else:
# return x.detach(), self.pool(x, tokens.input_ids)
if return_mask:
return x.detach() + 0.0, tokens.attention_mask.to(we.device_id)
else:
return x.detach() + 0.0, None
return x.detach() + 0.0
def pool(self, x, tokens):
# take features from the eot embedding (eot_token is the highest number in each sequence)
@@ -921,6 +921,15 @@ class T5EmbedderHF(BaseEmbedder):
return self(tokens, return_mask=return_mask)
def encode(self, text, return_mask=False, use_mask=True):
if isinstance(text, str):
text = [text]
if self.clean:
text = [self._clean(u) for u in text]
assert self.tokenizer is not None
tokens = self.tokenizer(text, **self.tokenize_kargs)
return self(tokens, return_mask=return_mask, use_mask=use_mask)
def encode_list(self, text, return_mask=False, use_mask=True):
if isinstance(text, str):
text = [text]
if self.clean:
@@ -942,62 +951,11 @@ class T5EmbedderHF(BaseEmbedder):
else:
return torch.cat(cont, dim=0)
def encode_longlist(self, text_list, return_mask=True):
text_max_len = max([len(p) for p in text_list]) * self.length
cont_list, cont_mask_list = [], []
for pp in text_list:
cont, cont_mask = self.encode(pp, return_mask=return_mask)
cont_channel, cont_dim = cont.shape[0] * cont.shape[1], cont.shape[
2]
cont = cont.view(cont_channel, cont_dim)
cont_mask_channel = cont_mask.shape[0] * cont_mask.shape[1]
cont_mask = cont_mask.view(cont_mask_channel)
select_cont = cont[cont_mask == 1]
select_cont_mask, _ = torch.sort(cont_mask, dim=0, descending=True)
if select_cont.shape[0] != text_max_len:
select_cont = F.pad(
select_cont,
(0, 0, 0, text_max_len - select_cont.shape[0]))
if select_cont_mask.shape[0] != text_max_len:
select_cont_mask = F.pad(
select_cont_mask,
(0, text_max_len - select_cont_mask.shape[0]))
cont_list.append(select_cont)
cont_mask_list.append(select_cont_mask)
return torch.stack(cont_list), torch.stack(cont_mask_list)
def encode_longlist_v1(self, text_list, return_mask=True):
cont_list = []
max_len = 0
for pp in text_list:
cont, cont_mask = self.encode(pp, return_mask=True)
txt_lens = cont_mask.flatten(start_dim=1).sum(dim=-1)
pp_cont = torch.cat(
[c[:txt_len] for c, txt_len in zip(cont, txt_lens)], dim=0)
max_len = pp_cont.size(0) if pp_cont.size(0) > max_len else max_len
cont_list.append(pp_cont)
cont = torch.cat([
torch.cat([c, c.new_zeros(max_len - c.size(0), c.size(1))],
dim=0).unsqueeze(0) for c in cont_list
],
dim=0)
if return_mask:
cont_mask = torch.cat([
torch.cat(
[c.new_ones(c.size(0)),
c.new_zeros(max_len - c.size(0))],
dim=-1).unsqueeze(0) for c in cont_list
],
dim=0).type(torch.long, non_blocking=True)
return cont, cont_mask
else:
return cont
def encode_list(self, text_list, return_mask=True):
def encode_list_of_list(self, text_list, return_mask=True, use_mask=True):
cont_list = []
mask_list = []
for pp in text_list:
cont, cont_mask = self.encode(pp, return_mask=return_mask)
cont, cont_mask = self.encode_list(pp, return_mask=return_mask, use_mask=use_mask)
cont_list.append(cont)
mask_list.append(cont_mask)
if return_mask:

Some files were not shown because too many files have changed in this diff Show More