Compare commits
22
Commits
v1.2.0_dev
...
v1.4.0_dev
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c4b4d88e75 | ||
|
|
043222de49 | ||
|
|
d7dbdc5292 | ||
|
|
711c10a68c | ||
|
|
448cdba522 | ||
|
|
adda36e39d | ||
|
|
1da1864993 | ||
|
|
5eac362325 | ||
|
|
9ae3bca43d | ||
|
|
ca034ef765 | ||
|
|
2a29446d45 | ||
|
|
cd33b4ab15 | ||
|
|
7d7943fed3 | ||
|
|
d48b2f110f | ||
|
|
82486adf38 | ||
|
|
a683061c6f | ||
|
|
4ec0492897 | ||
|
|
02c0ba9757 | ||
|
|
82132ff3a1 | ||
|
|
73984c4f9e | ||
|
|
edb46a615c | ||
|
|
d9b207cf5b |
@@ -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 | [](https://huggingface.co/spaces/scepter-studio/ACE-Chat)<br>[](https://www.modelscope.cn/models/iic/ACE-0.6B-512px) [](https://huggingface.co/scepter-studio/ACE-0.6B-512px) |
|
||||
| ACE-0.6B-1024px | [](https://huggingface.co/spaces/scepter-studio/ACE-Refiner-Chat)<br>[](https://www.modelscope.cn/models/iic/ACE-0.6B-1024px) [](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.
|
||||
|
||||

|
||||
|
||||
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.
|
||||
|
||||

|
||||
|
||||
|
||||
We compared the generation and editing performance of different models on several tasks, as shown as following.
|
||||

|
||||
|
||||
|
||||
## 🔥 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
|
||||
|
||||

|
||||
|
||||
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}
|
||||
}
|
||||
```
|
||||
@@ -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>)
|
||||
|
||||
[](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>)
|
||||
|
||||

|
||||
[//]: # ( </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
|
||||
|
||||
@@ -210,18 +165,19 @@ pip install scepter
|
||||
|
||||
### Currently supported approaches
|
||||
|
||||
| Tasks | Methods | Links |
|
||||
|:----------------------------:|:----------------------------------------------:|:------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
|
||||
| Text-to-image Generation | SD v1.5 | [](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
|
||||
| Text-to-image Generation | SD v2.1 | [](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
|
||||
| Text-to-image Generation | SD-XL | [](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0) |
|
||||
| Text-to-image Generation | FLUX | [](https://huggingface.co/black-forest-labs/FLUX.1-dev) |
|
||||
| Efficient Tuning | LoRA | [](https://arxiv.org/abs/2106.09685) |
|
||||
| Efficient Tuning | Res-Tuning(NeurIPS23) | [](https://arxiv.org/abs/2310.19859) [](https://res-tuning.github.io/) |
|
||||
| Controllable Image Synthesis | [🌟SCEdit(CVPR24)](docs/en/tasks/scedit.md) | [](https://arxiv.org/abs/2312.11392) [](https://scedit.github.io/) |
|
||||
| Image Editing | [🌟LAR-Gen](docs/en/tasks/largen.md) | [](https://arxiv.org/abs/2403.19534) [](https://ali-vilab.github.io/largen-page/) |
|
||||
| Image Editing | [🌟StyleBooth](docs/en/tasks/stylebooth.md) | [](https://arxiv.org/abs/2404.12154) [](https://ali-vilab.github.io/stylebooth-page/) |
|
||||
| Image Generation and Editing | [🌟ACE](https://ali-vilab.github.io/ace-page/) | [](https://arxiv.org/abs/2410.00086) [](https://ali-vilab.github.io/ace-page/) [](https://huggingface.co/spaces/scepter-studio/ACE-Chat) <br> [](https://www.modelscope.cn/models/iic/ACE-0.6B-512px) [](https://huggingface.co/scepter-studio/ACE-0.6B-512px) |
|
||||
| Tasks | Methods | Links |
|
||||
|:----------------------------:|:------------------------------------------------:|:------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
|
||||
| Text-to-image Generation | SD v1.5 | [](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
|
||||
| Text-to-image Generation | SD v2.1 | [](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
|
||||
| Text-to-image Generation | SD-XL | [](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0) |
|
||||
| Text-to-image Generation | FLUX | [](https://huggingface.co/black-forest-labs/FLUX.1-dev) |
|
||||
| Efficient Tuning | LoRA | [](https://arxiv.org/abs/2106.09685) |
|
||||
| Efficient Tuning | Res-Tuning(NeurIPS23) | [](https://arxiv.org/abs/2310.19859) [](https://res-tuning.github.io/) |
|
||||
| Controllable Image Synthesis | [🌟SCEdit(CVPR24)](docs/en/tasks/scedit.md) | [](https://arxiv.org/abs/2312.11392) [](https://scedit.github.io/) |
|
||||
| Image Editing | [🌟LAR-Gen](docs/en/tasks/largen.md) | [](https://arxiv.org/abs/2403.19534) [](https://ali-vilab.github.io/largen-page/) |
|
||||
| Image Editing | [🌟StyleBooth](docs/en/tasks/stylebooth.md) | [](https://arxiv.org/abs/2404.12154) [](https://ali-vilab.github.io/stylebooth-page/) |
|
||||
| Image Generation and Editing | [🌟ACE](https://ali-vilab.github.io/ace-page/) | [](https://arxiv.org/abs/2410.00086) [](https://ali-vilab.github.io/ace-page/) [](https://huggingface.co/spaces/scepter-studio/ACE-Chat) <br> [](https://www.modelscope.cn/models/iic/ACE-0.6B-512px) [](https://huggingface.co/scepter-studio/ACE-0.6B-512px) |
|
||||
| Image Generation and Editing | [🌟ACE++](https://ali-vilab.github.io/ACE_plus_page/) | [](https://arxiv.org/abs/2501.02487) [](https://ali-vilab.github.io/ACE_plus_page/) [](https://huggingface.co/spaces/scepter-studio/ACE-Plus) <br> [](https://www.modelscope.cn/models/iic/ACE_Plus/summary) [](https://huggingface.co/ali-vilab/ACE_Plus/tree/main) |
|
||||
|
||||
|
||||
## 🖥️ SCEPTER Studio
|
||||
@@ -258,18 +214,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
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ albumentations
|
||||
beautifulsoup4
|
||||
bezier
|
||||
einops
|
||||
modelscope
|
||||
modelscope[framework]
|
||||
ms-swift
|
||||
numpy
|
||||
open_clip_torch
|
||||
@@ -12,6 +12,7 @@ oss2>=2.15.0
|
||||
pycocotools
|
||||
pyyaml>=5.3.1
|
||||
scikit-image
|
||||
scikit-learn
|
||||
sentencepiece
|
||||
torchsde
|
||||
transformers
|
||||
scikit-learn
|
||||
@@ -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==.19.1
|
||||
flash-attn==2.5.8
|
||||
xformers==0.0.28
|
||||
@@ -1,5 +1,5 @@
|
||||
bitsandbytes
|
||||
gradio==4.44.1
|
||||
gradio
|
||||
gradio_imageslider
|
||||
imagehash
|
||||
psutil
|
||||
|
||||
+23
-13
@@ -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={},
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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"
|
||||
|
||||
@@ -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={},
|
||||
)
|
||||
|
||||
@@ -1,23 +1,60 @@
|
||||
# -*- 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
|
||||
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'],
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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']
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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={},
|
||||
)
|
||||
|
||||
@@ -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={},
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -337,8 +337,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]
|
||||
@@ -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={},
|
||||
)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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={},
|
||||
)
|
||||
|
||||
@@ -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={},
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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']
|
||||
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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={},
|
||||
)
|
||||
|
||||
@@ -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={},
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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={},
|
||||
)
|
||||
|
||||
@@ -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':
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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={},
|
||||
)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
# -*- 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 torch
|
||||
import transformers
|
||||
from scepter.modules.model.embedder.base_embedder import BaseEmbedder
|
||||
@@ -51,59 +53,53 @@ class HFEmbedder(BaseEmbedder):
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
hf_model_cls = cfg.get('HF_MODEL_CLS', None)
|
||||
model_path = cfg.get('MODEL_PATH', None)
|
||||
model_path = cfg.get("MODEL_PATH", None)
|
||||
hf_tokenizer_cls = cfg.get('HF_TOKENIZER_CLS', None)
|
||||
tokenizer_path = cfg.get('TOKENIZER_PATH', None)
|
||||
self.max_length = cfg.get('MAX_LENGTH', 77)
|
||||
self.output_key = cfg.get('OUTPUT_KEY', 'last_hidden_state')
|
||||
self.d_type = cfg.get('D_TYPE', 'float')
|
||||
self.clean = cfg.get('CLEAN', 'whitespace')
|
||||
self.batch_infer = cfg.get('BATCH_INFER', False)
|
||||
self.output_key = cfg.get("OUTPUT_KEY", "last_hidden_state")
|
||||
self.d_type = cfg.get("D_TYPE", "float")
|
||||
self.clean = cfg.get("CLEAN", "whitespace")
|
||||
self.batch_infer = cfg.get("BATCH_INFER", False)
|
||||
self.added_identifier = cfg.get('ADDED_IDENTIFIER', None)
|
||||
torch_dtype = getattr(torch, self.d_type)
|
||||
|
||||
assert hf_model_cls is not None and hf_tokenizer_cls is not None
|
||||
assert model_path is not None and tokenizer_path is not None
|
||||
with FS.get_dir_to_local_dir(tokenizer_path, wait_finish=True) as local_path:
|
||||
self.tokenizer = getattr(transformers, hf_tokenizer_cls).from_pretrained(local_path,
|
||||
max_length = self.max_length,
|
||||
torch_dtype = torch_dtype,
|
||||
additional_special_tokens=self.added_identifier)
|
||||
|
||||
with FS.get_dir_to_local_dir(tokenizer_path,
|
||||
wait_finish=True) as local_path:
|
||||
self.tokenizer = getattr(transformers,
|
||||
hf_tokenizer_cls).from_pretrained(
|
||||
local_path,
|
||||
max_length=self.max_length,
|
||||
torch_dtype=torch_dtype)
|
||||
with FS.get_dir_to_local_dir(model_path, wait_finish=True) as local_path:
|
||||
self.hf_module = getattr(transformers, hf_model_cls).from_pretrained(local_path, torch_dtype = torch_dtype)
|
||||
|
||||
with FS.get_dir_to_local_dir(model_path,
|
||||
wait_finish=True) as local_path:
|
||||
self.hf_module = getattr(transformers,
|
||||
hf_model_cls).from_pretrained(
|
||||
local_path, torch_dtype=torch_dtype)
|
||||
|
||||
self.hf_module = self.hf_module.eval().requires_grad_(False)
|
||||
|
||||
def forward(self, text: list[str], return_mask=False):
|
||||
def forward(self, text: list[str], return_mask = False):
|
||||
batch_encoding = self.tokenizer(
|
||||
text,
|
||||
truncation=True,
|
||||
max_length=self.max_length,
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
padding='max_length',
|
||||
return_tensors='pt',
|
||||
padding="max_length",
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
outputs = self.hf_module(
|
||||
input_ids=batch_encoding['input_ids'].to(self.hf_module.device),
|
||||
input_ids=batch_encoding["input_ids"].to(self.hf_module.device),
|
||||
attention_mask=None,
|
||||
output_hidden_states=False,
|
||||
)
|
||||
if return_mask:
|
||||
return outputs[
|
||||
self.output_key], batch_encoding['attention_mask'].to(
|
||||
self.hf_module.device)
|
||||
return outputs[self.output_key], batch_encoding['attention_mask'].to(self.hf_module.device)
|
||||
else:
|
||||
return outputs[self.output_key], None
|
||||
|
||||
def encode(self, text, return_mask=False):
|
||||
def encode(self, text, return_mask = False):
|
||||
if isinstance(text, str):
|
||||
text = [text]
|
||||
if self.clean:
|
||||
@@ -119,12 +115,36 @@ class HFEmbedder(BaseEmbedder):
|
||||
else:
|
||||
return torch.cat(cont, dim=0)
|
||||
else:
|
||||
ret_data = self(text, return_mask=return_mask)
|
||||
ret_data = self(text, return_mask = return_mask)
|
||||
if return_mask:
|
||||
return ret_data
|
||||
else:
|
||||
return ret_data[0]
|
||||
|
||||
def encode_list(self, text_list, return_mask=True):
|
||||
cont_list = []
|
||||
mask_list = []
|
||||
for pp in text_list:
|
||||
cont = self.encode(pp, return_mask=return_mask)
|
||||
cont_list.append(cont[0]) if return_mask else cont_list.append(cont)
|
||||
mask_list.append(cont[1]) if return_mask else mask_list.append(None)
|
||||
if return_mask:
|
||||
return cont_list, mask_list
|
||||
else:
|
||||
return cont_list
|
||||
|
||||
def encode_list_of_list(self, text_list, return_mask=True):
|
||||
cont_list = []
|
||||
mask_list = []
|
||||
for pp in text_list:
|
||||
cont = self.encode_list(pp, return_mask=return_mask)
|
||||
cont_list.append(cont[0]) if return_mask else cont_list.append(cont)
|
||||
mask_list.append(cont[1]) if return_mask else mask_list.append(None)
|
||||
if return_mask:
|
||||
return cont_list, mask_list
|
||||
else:
|
||||
return cont_list
|
||||
|
||||
def _clean(self, text):
|
||||
if self.clean == 'whitespace':
|
||||
text = whitespace_clean(basic_clean(text))
|
||||
@@ -133,7 +153,6 @@ class HFEmbedder(BaseEmbedder):
|
||||
elif self.clean == 'canonicalize':
|
||||
text = canonicalize(basic_clean(text))
|
||||
return text
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('EMBEDDER',
|
||||
@@ -141,28 +160,49 @@ class HFEmbedder(BaseEmbedder):
|
||||
HFEmbedder.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@EMBEDDERS.register_class()
|
||||
class T5PlusClipFluxEmbedder(BaseEmbedder):
|
||||
"""
|
||||
Uses the OpenCLIP transformer encoder for text
|
||||
"""
|
||||
para_dict = {'T5_MODEL': {}, 'CLIP_MODEL': {}}
|
||||
para_dict = {
|
||||
'T5_MODEL': {},
|
||||
'CLIP_MODEL': {}
|
||||
}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.t5_model = EMBEDDERS.build(cfg.T5_MODEL, logger=logger)
|
||||
self.clip_model = EMBEDDERS.build(cfg.CLIP_MODEL, logger=logger)
|
||||
|
||||
def encode(self, text):
|
||||
t5_embeds = self.t5_model.encode(text, return_mask=False)
|
||||
clip_embeds = self.clip_model.encode(text, return_mask=False)
|
||||
def encode(self, text, return_mask = False):
|
||||
t5_embeds = self.t5_model.encode(text, return_mask = return_mask)
|
||||
clip_embeds = self.clip_model.encode(text, return_mask = return_mask)
|
||||
# change embedding strategy here
|
||||
return {
|
||||
'context': t5_embeds,
|
||||
'y': clip_embeds,
|
||||
}
|
||||
|
||||
def encode_list(self, text, return_mask = False):
|
||||
t5_embeds = self.t5_model.encode_list(text, return_mask = return_mask)
|
||||
clip_embeds = self.clip_model.encode_list(text, return_mask = return_mask)
|
||||
# change embedding strategy here
|
||||
return {
|
||||
'context': t5_embeds,
|
||||
'y': clip_embeds,
|
||||
}
|
||||
|
||||
def encode_list_of_list(self, text, return_mask = False):
|
||||
t5_embeds = self.t5_model.encode_list_of_list(text, return_mask = return_mask)
|
||||
clip_embeds = self.clip_model.encode_list_of_list(text, return_mask = return_mask)
|
||||
# change embedding strategy here
|
||||
return {
|
||||
'context': t5_embeds,
|
||||
'y': clip_embeds,
|
||||
}
|
||||
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('EMBEDDER',
|
||||
|
||||
@@ -1,5 +1,25 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.head.classifier_head import (
|
||||
ClassifierHead, CosineLinearHead, TransformerHead, TransformerHeadx2,
|
||||
VideoClassifierHead, VideoClassifierHeadx2)
|
||||
from typing import TYPE_CHECKING
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from scepter.modules.model.head.classifier_head import (
|
||||
ClassifierHead, CosineLinearHead, TransformerHead, TransformerHeadx2,
|
||||
VideoClassifierHead, VideoClassifierHeadx2)
|
||||
else:
|
||||
_import_structure = {
|
||||
'classifier_head': ['ClassifierHead', 'CosineLinearHead',
|
||||
'TransformerHead', 'TransformerHeadx2',
|
||||
'VideoClassifierHead', 'VideoClassifierHeadx2']
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
@@ -1,4 +1,23 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.loss.base_losses import CrossEntropy
|
||||
from scepter.modules.model.loss.rec_loss import MinSNRLoss, ReconstructLoss
|
||||
from typing import TYPE_CHECKING
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from scepter.modules.model.loss.base_losses import CrossEntropy
|
||||
from scepter.modules.model.loss.rec_loss import MinSNRLoss, ReconstructLoss
|
||||
else:
|
||||
_import_structure = {
|
||||
'base_losses': ['CrossEntropy'],
|
||||
'rec_loss': ['MinSNRLoss', 'ReconstructLoss']
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
@@ -1,5 +1,23 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.metric.classification import (AccuracyMetric,
|
||||
EnsembleAccuracyMetric
|
||||
)
|
||||
from typing import TYPE_CHECKING
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from scepter.modules.model.metric.classification import (AccuracyMetric,
|
||||
EnsembleAccuracyMetric
|
||||
)
|
||||
else:
|
||||
_import_structure = {
|
||||
'classification': ['AccuracyMetric', 'EnsembleAccuracyMetric']
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
@@ -1,5 +1,24 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.neck.global_average_pooling import \
|
||||
GlobalAveragePooling
|
||||
from scepter.modules.model.neck.identity import Identity
|
||||
from typing import TYPE_CHECKING
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from scepter.modules.model.neck.global_average_pooling import \
|
||||
GlobalAveragePooling
|
||||
from scepter.modules.model.neck.identity import Identity
|
||||
else:
|
||||
_import_structure = {
|
||||
'global_average_pooling': ['GlobalAveragePooling'],
|
||||
'identity': ['Identity']
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
@@ -1,9 +1,32 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.network.autoencoder import ae_kl
|
||||
from scepter.modules.model.network.classifier import Classifier
|
||||
from scepter.modules.model.network.diffusion import (diffusion, schedules,
|
||||
solvers)
|
||||
from scepter.modules.model.network.ldm import (ldm, ldm_edit, ldm_pixart,
|
||||
ldm_sce, ldm_sd3, ldm_xl,
|
||||
ldm_flux)
|
||||
from typing import TYPE_CHECKING
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from scepter.modules.model.network.autoencoder import ae_kl
|
||||
from scepter.modules.model.network.classifier import Classifier
|
||||
from scepter.modules.model.network.diffusion import (diffusion, schedules,
|
||||
solvers)
|
||||
from scepter.modules.model.network.ldm import (ldm, ldm_edit, ldm_pixart,
|
||||
ldm_sce, ldm_sd3, ldm_xl,
|
||||
ldm_flux)
|
||||
else:
|
||||
_import_structure = {
|
||||
'autoencoder': ['ae_kl'],
|
||||
'classifier': ['Classifier'],
|
||||
'diffusion': ['diffusion', 'schedules', 'solvers'],
|
||||
'ldm': ['ldm', 'ldm_edit', 'ldm_pixart',
|
||||
'ldm_sce', 'ldm_sd3', 'ldm_xl',
|
||||
'ldm_flux']
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
@@ -1,3 +1,23 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.network.autoencoder.ae_kl import AutoencoderKL
|
||||
from typing import TYPE_CHECKING
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from scepter.modules.model.network.autoencoder.ae_kl import AutoencoderKL
|
||||
from scepter.modules.model.network.autoencoder.ae_kl_cogvideox import AutoencoderKLCogVideoX
|
||||
else:
|
||||
_import_structure = {
|
||||
'ae_kl': ['AutoencoderKL'],
|
||||
'ae_kl_cogvideox': ['AutoencoderKLCogVideoX']
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
@@ -129,7 +129,7 @@ class AutoencoderKL(TrainModule):
|
||||
for k in f.keys():
|
||||
sd[k] = f.get_tensor(k)
|
||||
else:
|
||||
sd = torch.load(path, map_location='cpu')
|
||||
sd = torch.load(path, map_location='cpu', weights_only=True)
|
||||
if path.find('.pt') > -1 and 'state_dict' in sd:
|
||||
sd = sd['state_dict']
|
||||
elif path.find('.ckpt') > -1 and 'state_dict' in sd:
|
||||
@@ -373,7 +373,7 @@ class AutoencoderKLFlux(TrainModule):
|
||||
for k in f.keys():
|
||||
sd[k] = f.get_tensor(k)
|
||||
else:
|
||||
sd = torch.load(path, map_location="cpu")
|
||||
sd = torch.load(path, map_location="cpu", weights_only=True)
|
||||
if path.find('.pt') > -1 and 'state_dict' in sd:
|
||||
sd = sd['state_dict']
|
||||
elif path.find('.ckpt') > -1 and 'state_dict' in sd:
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,4 +1,22 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.network.diffusion import (diffusion, schedules,
|
||||
solvers)
|
||||
from typing import TYPE_CHECKING
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from scepter.modules.model.network.diffusion import (diffusion, schedules,
|
||||
solvers)
|
||||
else:
|
||||
_import_structure = {
|
||||
'diffusion': ['diffusion', 'schedules', 'solvers']
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
@@ -239,7 +239,7 @@ class GaussianDiffusion(object):
|
||||
percentile=None,
|
||||
cat_uc=False,
|
||||
**kwargs):
|
||||
"""
|
||||
r"""
|
||||
Apply one step of denoising from the posterior distribution q(x_s | x_t, x0).
|
||||
Since x0 is not available, estimate the denoising results using the learned
|
||||
distribution p(x_s | x_t, \hat{x}_0 == f(x_t)). # noqa
|
||||
|
||||
@@ -1,11 +1,44 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.network.ldm.ldm import LatentDiffusion
|
||||
from scepter.modules.model.network.ldm.ldm_ace import LatentDiffusionACE
|
||||
from scepter.modules.model.network.ldm.ldm_edit import LatentDiffusionEdit
|
||||
from scepter.modules.model.network.ldm.ldm_pixart import LatentDiffusionPixart
|
||||
from scepter.modules.model.network.ldm.ldm_sce import (
|
||||
LatentDiffusionSCEControl, LatentDiffusionSCETuning,
|
||||
LatentDiffusionXLSCEControl, LatentDiffusionXLSCETuning)
|
||||
from scepter.modules.model.network.ldm.ldm_sd3 import LatentDiffusionSD3
|
||||
from scepter.modules.model.network.ldm.ldm_xl import LatentDiffusionXL
|
||||
from typing import TYPE_CHECKING
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from scepter.modules.model.network.ldm.ldm import LatentDiffusion
|
||||
from scepter.modules.model.network.ldm.ldm_ace import (LatentDiffusionACE,
|
||||
LatentDiffusionACERefiner)
|
||||
from scepter.modules.model.network.ldm.ldm_edit import LatentDiffusionEdit
|
||||
from scepter.modules.model.network.ldm.ldm_pixart import LatentDiffusionPixart
|
||||
from scepter.modules.model.network.ldm.ldm_sce import (
|
||||
LatentDiffusionSCEControl, LatentDiffusionSCETuning,
|
||||
LatentDiffusionXLSCEControl, LatentDiffusionXLSCETuning)
|
||||
from scepter.modules.model.network.ldm.ldm_sd3 import LatentDiffusionSD3
|
||||
from scepter.modules.model.network.ldm.ldm_xl import LatentDiffusionXL
|
||||
from scepter.modules.model.network.ldm.ldm_cogvideox import LatentDiffusionCogVideoX
|
||||
from scepter.modules.model.network.ldm.ldm_flux import (LatentDiffusionFlux,
|
||||
LatentDiffusionFluxMR)
|
||||
from scepter.modules.model.network.ldm.ldm_ace_plus import LatentDiffusionACEPlus
|
||||
else:
|
||||
_import_structure = {
|
||||
'ldm': ['LatentDiffusion'],
|
||||
'ldm_ace': ['LatentDiffusionACE', 'LatentDiffusionACERefiner'],
|
||||
'ldm_edit': ['LatentDiffusionEdit'],
|
||||
'ldm_pixart': ['LatentDiffusionPixart'],
|
||||
'ldm_sce': ['LatentDiffusionSCEControl', 'LatentDiffusionSCETuning',
|
||||
'LatentDiffusionXLSCEControl', 'LatentDiffusionXLSCETuning'],
|
||||
'ldm_sd3': ['LatentDiffusionSD3'],
|
||||
'ldm_xl': ['LatentDiffusionXL'],
|
||||
'ldm_cogvideox': ['LatentDiffusionCogVideoX'],
|
||||
'ldm_flux': ['LatentDiffusionFlux', 'LatentDiffusionFluxMR'],
|
||||
'ldm_ace_plus': ['LatentDiffusionACEPlus'],
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
@@ -200,7 +200,7 @@ class LatentDiffusion(TrainModule):
|
||||
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
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import math
|
||||
import random
|
||||
from contextlib import nullcontext
|
||||
|
||||
@@ -10,6 +11,7 @@ from torch import nn
|
||||
|
||||
from scepter.modules.model.network.ldm import LatentDiffusion
|
||||
from scepter.modules.model.registry import MODELS
|
||||
import torchvision.transforms as T
|
||||
from scepter.modules.model.utils.basic_utils import check_list_of_list
|
||||
from scepter.modules.model.utils.basic_utils import \
|
||||
pack_imagelist_into_tensor_v2 as pack_imagelist_into_tensor
|
||||
@@ -67,10 +69,10 @@ class LatentDiffusionACE(LatentDiffusion):
|
||||
if self.use_text_pos_embeddings and not torch.sum(
|
||||
self.text_position_embeddings.pos) > 0:
|
||||
identifier_cont, identifier_cont_mask = getattr(
|
||||
self.cond_stage_model, 'encode')(self.text_indentifers,
|
||||
self.cond_stage_model, 'encode_list_of_list')(self.text_indentifers,
|
||||
return_mask=True)
|
||||
self.text_position_embeddings.load_state_dict(
|
||||
{'pos': identifier_cont[:, 0, :]})
|
||||
{'pos': torch.cat( [one_id[0][0, :].unsqueeze(0) for one_id in identifier_cont], dim=0)})
|
||||
cont_, cont_mask_ = [], []
|
||||
for pp, edit, c, cm in zip(prompt, edit_image, cont, cont_mask):
|
||||
if isinstance(pp, list):
|
||||
@@ -93,8 +95,8 @@ class LatentDiffusionACE(LatentDiffusion):
|
||||
return batch_data_list
|
||||
|
||||
def forward_train(self,
|
||||
edit_image=[],
|
||||
edit_image_mask=[],
|
||||
src_image_list=[],
|
||||
src_mask_list=[],
|
||||
image=None,
|
||||
image_mask=None,
|
||||
noise=None,
|
||||
@@ -112,8 +114,8 @@ class LatentDiffusionACE(LatentDiffusion):
|
||||
Returns:
|
||||
'''
|
||||
assert check_list_of_list(prompt) and check_list_of_list(
|
||||
edit_image) and check_list_of_list(edit_image_mask)
|
||||
assert len(edit_image) == len(edit_image_mask) == len(prompt)
|
||||
src_image_list) and check_list_of_list(src_mask_list)
|
||||
assert len(src_image_list) == len(src_mask_list) == len(prompt)
|
||||
assert self.cond_stage_model is not None
|
||||
gc_seg = kwargs.pop('gc_seg', [])
|
||||
gc_seg = int(gc_seg[0]) if len(gc_seg) > 0 else 0
|
||||
@@ -138,16 +140,16 @@ class LatentDiffusionACE(LatentDiffusion):
|
||||
prompt_ = [[pp] if isinstance(pp, str) else pp for pp in prompt]
|
||||
try:
|
||||
cont, cont_mask = getattr(self.cond_stage_model,
|
||||
'encode_list')(prompt_, return_mask=True)
|
||||
'encode_list_of_list')(prompt_, return_mask=True)
|
||||
except Exception as e:
|
||||
print(e, prompt_)
|
||||
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont,
|
||||
cont, cont_mask = self.cond_stage_embeddings(prompt, src_image_list, cont,
|
||||
cont_mask)
|
||||
context['crossattn'] = cont
|
||||
|
||||
# process edit image & edit image mask
|
||||
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]
|
||||
edit_image = [to_device(i, strict=False) for i in src_image_list]
|
||||
edit_image_mask = [to_device(i, strict=False) for i in src_mask_list]
|
||||
e_img, e_mask = [], []
|
||||
for u, m in zip(edit_image, edit_image_mask):
|
||||
if m is None:
|
||||
@@ -183,8 +185,8 @@ class LatentDiffusionACE(LatentDiffusion):
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_test(self,
|
||||
edit_image=[],
|
||||
edit_image_mask=[],
|
||||
src_image_list=[],
|
||||
src_mask_list=[],
|
||||
image=None,
|
||||
image_mask=None,
|
||||
prompt=[],
|
||||
@@ -198,8 +200,8 @@ class LatentDiffusionACE(LatentDiffusion):
|
||||
**kwargs):
|
||||
|
||||
assert check_list_of_list(prompt) and check_list_of_list(
|
||||
edit_image) and check_list_of_list(edit_image_mask)
|
||||
assert len(edit_image) == len(edit_image_mask) == len(prompt)
|
||||
src_image_list) and check_list_of_list(src_mask_list)
|
||||
assert len(src_image_list) == len(src_mask_list) == len(prompt)
|
||||
assert self.cond_stage_model is not None
|
||||
# gc_seg is unused
|
||||
kwargs.pop('gc_seg', -1)
|
||||
@@ -207,7 +209,7 @@ class LatentDiffusionACE(LatentDiffusion):
|
||||
context, null_context = {}, {}
|
||||
|
||||
prompt, n_prompt, image, image_mask, edit_image, edit_image_mask = self.limit_batch_data(
|
||||
[prompt, n_prompt, image, image_mask, edit_image, edit_image_mask],
|
||||
[prompt, n_prompt, image, image_mask, src_image_list, src_mask_list],
|
||||
log_num)
|
||||
g = torch.Generator(device=we.device_id)
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||
@@ -240,11 +242,11 @@ class LatentDiffusionACE(LatentDiffusion):
|
||||
# with torch.autocast(device_type="cuda", enabled=True, dtype=torch.bfloat16):
|
||||
prompt_ = [[pp] if isinstance(pp, str) else pp for pp in prompt]
|
||||
cont, cont_mask = getattr(self.cond_stage_model,
|
||||
'encode_list')(prompt_, return_mask=True)
|
||||
'encode_list_of_list')(prompt_, return_mask=True)
|
||||
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont,
|
||||
cont_mask)
|
||||
null_cont, null_cont_mask = getattr(self.cond_stage_model,
|
||||
'encode_list')(n_prompt,
|
||||
'encode_list_of_list')(n_prompt,
|
||||
return_mask=True)
|
||||
null_cont, null_cont_mask = self.cond_stage_embeddings(
|
||||
prompt, edit_image, null_cont, null_cont_mask)
|
||||
@@ -349,3 +351,254 @@ class LatentDiffusionACE(LatentDiffusion):
|
||||
__class__.__name__,
|
||||
LatentDiffusionACE.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@MODELS.register_class()
|
||||
class LatentDiffusionACERefiner(LatentDiffusionACE):
|
||||
def init_params(self):
|
||||
super().init_params()
|
||||
self.enhence_model_cfg = self.cfg.get("ENHENCE_MODEL", None)
|
||||
self.enhence_sampler_cfg = self.cfg.get("ENHENCE_SAMPLER_CFG", {})
|
||||
def construct_network(self):
|
||||
super().construct_network()
|
||||
if self.enhence_model_cfg:
|
||||
self.enhence_model = MODELS.build(self.enhence_model_cfg, logger=self.logger).eval().requires_grad_(False)
|
||||
self.enhence_sampler_cfg = {key.lower(): value for key, value in self.enhence_sampler_cfg.items()}
|
||||
else:
|
||||
self.enhence_model = None
|
||||
self.enhence_sampler_cfg = None
|
||||
|
||||
def forward_sample(self,
|
||||
src_image_list=[],
|
||||
src_mask_list=[],
|
||||
noise=None,
|
||||
cond_mask=[],
|
||||
x_shapes=[],
|
||||
prompt=[],
|
||||
n_prompt=[],
|
||||
sampler='ddim',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=4.5,
|
||||
guide_rescale=0.5,
|
||||
discretization='trailing',
|
||||
**kwargs
|
||||
):
|
||||
'''
|
||||
Args:
|
||||
edit_image: list of list of edit_image
|
||||
edit_image_mask: list of list of edit_image_mask
|
||||
image: target image
|
||||
image_mask: target image mask
|
||||
prompt: list of list of text
|
||||
n_prompt: list of list of text
|
||||
sampler:
|
||||
sample_steps:
|
||||
seed:
|
||||
guide_scale:
|
||||
guide_rescale:
|
||||
discretization:
|
||||
log_num:
|
||||
**kwargs:
|
||||
|
||||
Returns:
|
||||
|
||||
'''
|
||||
|
||||
# prepare data
|
||||
context, null_context = {}, {}
|
||||
context['x_shapes'] = null_context['x_shapes'] = x_shapes
|
||||
# process image mask
|
||||
|
||||
context['x_mask'] = null_context['x_mask'] = cond_mask
|
||||
# process text
|
||||
# with torch.autocast(device_type="cuda", enabled=True, dtype=torch.bfloat16):
|
||||
|
||||
cont, cont_mask = getattr(self.cond_stage_model, 'encode_list')(prompt, return_mask=True)
|
||||
cont, cont_mask = self.cond_stage_embeddings(prompt, src_image_list, cont, cont_mask)
|
||||
null_cont, null_cont_mask = getattr(self.cond_stage_model, 'encode_list')(n_prompt, return_mask=True)
|
||||
null_cont, null_cont_mask = self.cond_stage_embeddings(prompt, src_image_list, null_cont, null_cont_mask)
|
||||
context['crossattn'] = cont
|
||||
null_context['crossattn'] = null_cont
|
||||
|
||||
|
||||
null_context['edit'] = context['edit'] = src_image_list
|
||||
null_context['edit_mask'] = context['edit_mask'] = src_mask_list
|
||||
|
||||
# process sample
|
||||
model = self.model_ema if self.use_ema and self.eval_ema else self.model
|
||||
embedding_context = model.no_sync if isinstance(model, torch.distributed.fsdp.FullyShardedDataParallel) \
|
||||
else nullcontext
|
||||
with embedding_context():
|
||||
samples = self.diffusion.sample(solver=sampler,
|
||||
noise=noise,
|
||||
model=model,
|
||||
model_kwargs=[{
|
||||
'cond': context,
|
||||
'mask': cont_mask,
|
||||
'text_position_embeddings': self.text_position_embeddings.pos if hasattr(
|
||||
self.text_position_embeddings, 'pos') else None
|
||||
}, {
|
||||
'cond': null_context,
|
||||
'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': context,
|
||||
'mask': cont_mask,
|
||||
'text_position_embeddings': self.text_position_embeddings.pos if hasattr(
|
||||
self.text_position_embeddings, 'pos') else None
|
||||
},
|
||||
cat_uc=False,
|
||||
steps=sample_steps,
|
||||
guide_scale=guide_scale,
|
||||
guide_rescale=guide_rescale,
|
||||
discretization=discretization,
|
||||
show_progress=True,
|
||||
seed=seed,
|
||||
condition_fn=None,
|
||||
clamp=None,
|
||||
percentile=None,
|
||||
t_max=None,
|
||||
t_min=None,
|
||||
discard_penultimate_step=None,
|
||||
return_intermediate=None,
|
||||
**kwargs)
|
||||
|
||||
samples = unpack_tensor_into_imagelist(samples, x_shapes)
|
||||
x_samples = self.decode_first_stage(samples)
|
||||
return x_samples
|
||||
|
||||
def upscale_resize(self, image, interpolation=T.InterpolationMode.BILINEAR):
|
||||
_, c, H, W = image.shape
|
||||
scale = max(1.0, math.sqrt(4096 / ((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 forward_test(self,
|
||||
src_image_list=[],
|
||||
src_mask_list=[],
|
||||
image=None,
|
||||
image_mask=None,
|
||||
prompt=[],
|
||||
n_prompt=[],
|
||||
sampler='ddim',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=4.5,
|
||||
guide_rescale=0.5,
|
||||
discretization='trailing',
|
||||
enhance_scale=0.99,
|
||||
log_num=-1,
|
||||
**kwargs):
|
||||
assert check_list_of_list(prompt) and check_list_of_list(src_image_list) and check_list_of_list(src_mask_list)
|
||||
assert len(src_image_list) == len(src_mask_list) == len(prompt)
|
||||
assert self.cond_stage_model is not None
|
||||
# gc_seg is unused
|
||||
kwargs.pop("gc_seg", -1)
|
||||
prompt, n_prompt, image, image_mask, edit_image, edit_image_mask = self.limit_batch_data(
|
||||
[prompt, n_prompt, image, image_mask, src_image_list, src_mask_list], log_num)
|
||||
|
||||
prompt = [[pp] if isinstance(pp, str) else pp for pp in prompt]
|
||||
|
||||
g = torch.Generator(device=we.device_id)
|
||||
seed = seed if seed >= 0 else random.randint(0, 2 ** 32 - 1)
|
||||
g.manual_seed(seed)
|
||||
n_prompt = copy.deepcopy(prompt)
|
||||
# only modify the last prompt to be zero
|
||||
for nn_p_id, nn_p in enumerate(n_prompt):
|
||||
if isinstance(nn_p, str):
|
||||
n_prompt[nn_p_id] = [""]
|
||||
elif isinstance(nn_p, list):
|
||||
n_prompt[nn_p_id][-1] = ""
|
||||
else:
|
||||
raise NotImplementedError
|
||||
# process image
|
||||
image = to_device(image)
|
||||
x = self.encode_first_stage(image, **kwargs)
|
||||
noise = [torch.empty(*i.shape, device=we.device_id).normal_(generator=g) for i in x]
|
||||
noise, x_shapes = pack_imagelist_into_tensor(noise)
|
||||
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)
|
||||
|
||||
# processe edit image & edit image mask
|
||||
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])
|
||||
|
||||
x_samples = self.forward_sample(
|
||||
edit_image=e_img,
|
||||
edit_mask=e_mask,
|
||||
noise=noise,
|
||||
cond_mask=cond_mask,
|
||||
x_shapes=x_shapes,
|
||||
prompt=prompt,
|
||||
n_prompt=n_prompt,
|
||||
sampler=sampler,
|
||||
sample_steps=sample_steps,
|
||||
seed=seed,
|
||||
guide_scale=guide_scale,
|
||||
guide_rescale=guide_rescale,
|
||||
discretization='trailing',
|
||||
**kwargs)
|
||||
|
||||
if self.enhence_model and enhance_scale > 0:
|
||||
x_samples = [self.upscale_resize(x) for x in x_samples]
|
||||
x_start = self.enhence_model.encode_first_stage(x_samples, **kwargs)
|
||||
noise = []
|
||||
for i, x in enumerate(x_start):
|
||||
noise_ = self.enhence_model.noise_sample(1, x_samples[i].shape[2], x_samples[i].shape[3], seed)
|
||||
noise.append(noise_)
|
||||
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
|
||||
x_samples = self.enhence_model.forward_sample(noise = noise,
|
||||
x = x_start,
|
||||
reverse_scale = enhance_scale,
|
||||
prompt =[kwargs.pop("enhance_prompt", "") for _ in noise],
|
||||
**self.enhence_sampler_cfg)
|
||||
outputs = list()
|
||||
for i in range(len(prompt)):
|
||||
rec_img = torch.clamp((x_samples[i].float() + 1.0) / 2.0 + self.decoder_bias / 255, min=0.0, max=1.0)
|
||||
rec_img = rec_img.squeeze(0)
|
||||
edit_imgs, edit_img_masks = [], []
|
||||
if edit_image is not None and edit_image[i] is not None:
|
||||
if edit_image_mask[i] is None:
|
||||
edit_image_mask[i] = [None] * len(edit_image[i])
|
||||
for edit_img, edit_mask in zip(edit_image[i], edit_image_mask[i]):
|
||||
edit_img = torch.clamp((edit_img + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
edit_imgs.append(edit_img.squeeze(0))
|
||||
if edit_mask is None:
|
||||
edit_mask = torch.ones_like(edit_img[[0], :, :])
|
||||
edit_img_masks.append(edit_mask)
|
||||
one_tup = {
|
||||
'reconstruct_image': rec_img,
|
||||
'instruction': prompt[i],
|
||||
'edit_image': edit_imgs if len(edit_imgs) > 0 else None,
|
||||
'edit_mask': edit_img_masks if len(edit_imgs) > 0 else None
|
||||
}
|
||||
if image is not None:
|
||||
if image_mask is None:
|
||||
image_mask = [None] * len(image)
|
||||
ori_img = torch.clamp((image[i] + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
one_tup['target_image'] = ori_img.squeeze(0)
|
||||
one_tup['target_mask'] = image_mask[i] if image_mask[i] is not None else torch.ones_like(
|
||||
ori_img[[0], :, :])
|
||||
outputs.append(one_tup)
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODEL',
|
||||
__class__.__name__,
|
||||
LatentDiffusionACERefiner.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@@ -0,0 +1,371 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import random
|
||||
from contextlib import nullcontext
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel
|
||||
|
||||
from einops import rearrange
|
||||
from scepter.modules.model.network.ldm import LatentDiffusionFluxMR
|
||||
from scepter.modules.model.registry import MODELS
|
||||
from scepter.modules.model.utils.basic_utils import (
|
||||
check_list_of_list, limit_batch_data, pack_imagelist_into_tensor,
|
||||
to_device, unpack_tensor_into_imagelist)
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
|
||||
|
||||
@MODELS.register_class()
|
||||
class LatentDiffusionACEPlus(LatentDiffusionFluxMR):
|
||||
para_dict = {}
|
||||
para_dict.update(LatentDiffusionFluxMR.para_dict)
|
||||
|
||||
def resize_func(self, x, size):
|
||||
if x is None:
|
||||
return x
|
||||
return F.interpolate(x.unsqueeze(0), size=size, mode='nearest-exact')
|
||||
|
||||
def parse_ref_and_edit(
|
||||
self,
|
||||
src_image,
|
||||
src_image_mask,
|
||||
text_embedding,
|
||||
# text_mask,
|
||||
edit_id):
|
||||
edit_image = []
|
||||
edit_mask = []
|
||||
ref_image = []
|
||||
ref_mask = []
|
||||
ref_context = []
|
||||
ref_y = []
|
||||
ref_id = []
|
||||
txt = []
|
||||
txt_y = []
|
||||
for sample_id, (
|
||||
one_src,
|
||||
one_src_mask,
|
||||
one_text_embedding,
|
||||
one_text_y,
|
||||
# one_text_mask,
|
||||
one_edit_id) in enumerate(
|
||||
zip(
|
||||
src_image,
|
||||
src_image_mask,
|
||||
text_embedding['context'],
|
||||
text_embedding['y'],
|
||||
# text_mask,
|
||||
edit_id)):
|
||||
ref_id.append([i for i in range(len(one_src))])
|
||||
if hasattr(self,
|
||||
'ref_cond_stage_model') and self.ref_cond_stage_model:
|
||||
ref_image.append(
|
||||
self.ref_cond_stage_model.encode_list([
|
||||
((i + 1.0) / 2.0 * 255).type(torch.uint8)
|
||||
for i in one_src
|
||||
]))
|
||||
else:
|
||||
ref_image.append(one_src)
|
||||
ref_mask.append(one_src_mask)
|
||||
# process edit image & edit image mask
|
||||
current_edit_image = to_device([one_src[i] for i in one_edit_id],
|
||||
strict=False)
|
||||
current_edit_image = [
|
||||
v.squeeze(0)
|
||||
for v in self.encode_first_stage(current_edit_image)
|
||||
]
|
||||
current_edit_image_mask = to_device(
|
||||
[one_src_mask[i] for i in one_edit_id], strict=False)
|
||||
current_edit_image_mask = [
|
||||
self.reshape_func(m).squeeze(0)
|
||||
for m in current_edit_image_mask
|
||||
]
|
||||
|
||||
edit_image.append(current_edit_image)
|
||||
edit_mask.append(current_edit_image_mask)
|
||||
ref_context.append(one_text_embedding[:len(ref_id[-1])])
|
||||
ref_y.append(one_text_y[:len(ref_id[-1])])
|
||||
if not sum(len(src_) for src_ in src_image) > 0:
|
||||
ref_image = None
|
||||
ref_context = None
|
||||
ref_y = None
|
||||
for sample_id, (one_text_embedding, one_text_y) in enumerate(
|
||||
zip(text_embedding['context'], text_embedding['y'])):
|
||||
txt.append(one_text_embedding[-1].squeeze(0))
|
||||
txt_y.append(one_text_y[-1])
|
||||
return {
|
||||
'edit': edit_image,
|
||||
'edit_mask': edit_mask,
|
||||
'edit_id': edit_id,
|
||||
'ref_context': ref_context,
|
||||
'ref_y': ref_y,
|
||||
'context': txt,
|
||||
'y': txt_y,
|
||||
'ref_x': ref_image,
|
||||
'ref_mask': ref_mask,
|
||||
'ref_id': ref_id
|
||||
}
|
||||
|
||||
def reshape_func(self, mask):
|
||||
mask = mask.to(torch.bfloat16)
|
||||
mask = mask.view((-1, mask.shape[-2], mask.shape[-1]))
|
||||
mask = rearrange(
|
||||
mask,
|
||||
'c (h ph) (w pw) -> c (ph pw) h w',
|
||||
ph=8,
|
||||
pw=8,
|
||||
)
|
||||
return mask
|
||||
|
||||
def forward_train(self,
|
||||
src_image_list=[],
|
||||
src_mask_list=[],
|
||||
edit_id=[],
|
||||
image=None,
|
||||
image_mask=None,
|
||||
noise=None,
|
||||
prompt=[],
|
||||
**kwargs):
|
||||
'''
|
||||
Args:
|
||||
src_image: list of list of src_image
|
||||
src_image_mask: list of list of src_image_mask
|
||||
image: target image
|
||||
image_mask: target image mask
|
||||
noise: default is None, generate automaticly
|
||||
ref_prompt: list of list of text
|
||||
prompt: list of text
|
||||
**kwargs:
|
||||
Returns:
|
||||
'''
|
||||
assert check_list_of_list(src_image_list) and check_list_of_list(
|
||||
src_mask_list)
|
||||
assert self.cond_stage_model is not None
|
||||
|
||||
gc_seg = kwargs.pop('gc_seg', [])
|
||||
gc_seg = int(gc_seg[0]) if len(gc_seg) > 0 else 0
|
||||
align = kwargs.pop('align', [])
|
||||
prompt_ = [[pp] if isinstance(pp, str) else pp for pp in prompt]
|
||||
if len(align) < 1:
|
||||
align = [0] * len(prompt_)
|
||||
context = getattr(self.cond_stage_model,
|
||||
'encode_list_of_list')(prompt_)
|
||||
guide_scale = self.guide_scale
|
||||
if guide_scale is not None:
|
||||
guide_scale = torch.full((len(prompt_), ),
|
||||
guide_scale,
|
||||
device=we.device_id)
|
||||
else:
|
||||
guide_scale = None
|
||||
# image and image_mask
|
||||
# print("is list of list", check_list_of_list(image))
|
||||
if check_list_of_list(image):
|
||||
image = [to_device(ix) for ix in image]
|
||||
x_start = [self.encode_first_stage(ix, **kwargs) for ix in image]
|
||||
noise = [[torch.randn_like(ii) for ii in ix] for ix in x_start]
|
||||
x_start = [torch.cat(ix, dim=-1) for ix in x_start]
|
||||
noise = [torch.cat(ix, dim=-1) for ix in noise]
|
||||
|
||||
noise, _ = pack_imagelist_into_tensor(noise)
|
||||
|
||||
image_mask = [to_device(im, strict=False) for im in image_mask]
|
||||
x_mask = [[self.reshape_func(i).squeeze(0)
|
||||
for i in im] if im is not None else [None] * len(ix)
|
||||
for ix, im in zip(image, image_mask)]
|
||||
x_mask = [torch.cat(im, dim=-1) for im in x_mask]
|
||||
else:
|
||||
image = to_device(image)
|
||||
x_start = self.encode_first_stage(image, **kwargs)
|
||||
image_mask = to_device(image_mask, strict=False)
|
||||
x_mask = [self.reshape_func(i).squeeze(0) for i in image_mask
|
||||
] if image_mask is not None else [None] * len(image)
|
||||
loss_mask, _ = pack_imagelist_into_tensor(
|
||||
tuple(
|
||||
torch.ones_like(ix, dtype=torch.bool, device=ix.device)
|
||||
for ix in x_start))
|
||||
x_start, x_shapes = pack_imagelist_into_tensor(x_start)
|
||||
context['x_shapes'] = x_shapes
|
||||
context['align'] = align
|
||||
# process image mask
|
||||
|
||||
context['x_mask'] = x_mask
|
||||
ref_edit_context = self.parse_ref_and_edit(src_image_list,
|
||||
src_mask_list, context,
|
||||
edit_id)
|
||||
context.update(ref_edit_context)
|
||||
|
||||
teacher_context = copy.deepcopy(context)
|
||||
teacher_context['context'] = torch.cat(teacher_context['context'],
|
||||
dim=0)
|
||||
teacher_context['y'] = torch.cat(teacher_context['y'], dim=0)
|
||||
loss = self.diffusion.loss(x_0=x_start,
|
||||
model=self.model,
|
||||
model_kwargs={
|
||||
'cond': context,
|
||||
'gc_seg': gc_seg,
|
||||
'guidance': guide_scale
|
||||
},
|
||||
noise=noise,
|
||||
reduction='none',
|
||||
**kwargs)
|
||||
loss = loss[loss_mask].mean()
|
||||
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
|
||||
return ret
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_test(self,
|
||||
src_image_list=[],
|
||||
src_mask_list=[],
|
||||
edit_id=[],
|
||||
image=None,
|
||||
image_mask=None,
|
||||
prompt=[],
|
||||
sampler='flow_euler',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=3.5,
|
||||
guide_rescale=0.0,
|
||||
show_process=False,
|
||||
log_num=-1,
|
||||
**kwargs):
|
||||
outputs = self.forward_editing(src_image_list=src_image_list,
|
||||
src_mask_list=src_mask_list,
|
||||
edit_id=edit_id,
|
||||
image=image,
|
||||
image_mask=image_mask,
|
||||
prompt=prompt,
|
||||
sampler=sampler,
|
||||
sample_steps=sample_steps,
|
||||
seed=seed,
|
||||
guide_scale=guide_scale,
|
||||
guide_rescale=guide_rescale,
|
||||
show_process=show_process,
|
||||
log_num=log_num,
|
||||
**kwargs)
|
||||
return outputs
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_editing(self,
|
||||
src_image_list=[],
|
||||
src_mask_list=[],
|
||||
edit_id=[],
|
||||
image=None,
|
||||
image_mask=None,
|
||||
prompt=[],
|
||||
sampler='flow_euler',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=3.5,
|
||||
log_num=-1,
|
||||
**kwargs):
|
||||
# gc_seg is unused
|
||||
prompt, image, image_mask, src_image, src_image_mask, edit_id = limit_batch_data(
|
||||
[
|
||||
prompt, image, image_mask, src_image_list, src_mask_list,
|
||||
edit_id
|
||||
], log_num)
|
||||
assert check_list_of_list(src_image) and check_list_of_list(
|
||||
src_image_mask)
|
||||
assert self.cond_stage_model is not None
|
||||
align = kwargs.pop('align', [])
|
||||
prompt_ = [[pp] if isinstance(pp, str) else pp for pp in prompt]
|
||||
if len(align) < 1:
|
||||
align = [0] * len(prompt_)
|
||||
context = getattr(self.cond_stage_model,
|
||||
'encode_list_of_list')(prompt_)
|
||||
guide_scale = guide_scale or self.guide_scale
|
||||
if guide_scale is not None:
|
||||
guide_scale = torch.full((len(prompt), ),
|
||||
guide_scale,
|
||||
device=we.device_id)
|
||||
else:
|
||||
guide_scale = None
|
||||
# image and image_mask
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||
if image is not None:
|
||||
if check_list_of_list(image):
|
||||
image = [torch.cat(ix, dim=-1) for ix in image]
|
||||
image_mask = [torch.cat(im, dim=-1) for im in image_mask]
|
||||
noise = [
|
||||
self.noise_sample(1, ix.shape[1], ix.shape[2], seed)
|
||||
for ix in image
|
||||
]
|
||||
else:
|
||||
height, width = kwargs.pop('height'), kwargs.pop('width')
|
||||
noise = [self.noise_sample(1, height, width, seed) for _ in prompt]
|
||||
noise, x_shapes = pack_imagelist_into_tensor(noise)
|
||||
context['x_shapes'] = x_shapes
|
||||
context['align'] = align
|
||||
# process image mask
|
||||
image_mask = to_device(image_mask, strict=False)
|
||||
x_mask = [self.reshape_func(i).squeeze(0) for i in image_mask]
|
||||
context['x_mask'] = x_mask
|
||||
ref_edit_context = self.parse_ref_and_edit(src_image, src_image_mask,
|
||||
context, edit_id)
|
||||
context.update(ref_edit_context)
|
||||
# UNet use input n_prompt
|
||||
# model = self.model_ema if self.use_ema and self.eval_ema else self.model
|
||||
# import pdb;pdb.set_trace()
|
||||
model = self.model
|
||||
embedding_context = model.no_sync if isinstance(model, FullyShardedDataParallel) \
|
||||
else nullcontext
|
||||
with embedding_context():
|
||||
samples = self.diffusion.sample(noise=noise,
|
||||
sampler=sampler,
|
||||
model=self.model,
|
||||
model_kwargs={
|
||||
'cond': context,
|
||||
'guidance': guide_scale,
|
||||
'gc_seg': -1
|
||||
},
|
||||
steps=sample_steps,
|
||||
show_progress=True,
|
||||
guide_scale=guide_scale,
|
||||
return_intermediate=None,
|
||||
**kwargs).float()
|
||||
samples = unpack_tensor_into_imagelist(samples, x_shapes)
|
||||
with torch.autocast(device_type='cuda', dtype=torch.bfloat16):
|
||||
x_samples = self.decode_first_stage(samples)
|
||||
outputs = list()
|
||||
for i in range(len(prompt)):
|
||||
rec_img = torch.clamp((x_samples[i].float() + 1.0) / 2.0,
|
||||
min=0.0,
|
||||
max=1.0)
|
||||
rec_img = rec_img.squeeze(0)
|
||||
edit_imgs, edit_img_masks = [], []
|
||||
if src_image is not None and src_image[i] is not None:
|
||||
if src_image_mask[i] is None:
|
||||
src_image_mask[i] = [None] * len(src_image[i])
|
||||
for edit_img, edit_mask in zip(src_image[i],
|
||||
src_image_mask[i]):
|
||||
edit_img = torch.clamp((edit_img.float() + 1.0) / 2.0,
|
||||
min=0.0,
|
||||
max=1.0)
|
||||
edit_imgs.append(edit_img.squeeze(0))
|
||||
if edit_mask is None:
|
||||
edit_mask = torch.ones_like(edit_img[[0], :, :])
|
||||
edit_img_masks.append(edit_mask)
|
||||
one_tup = {
|
||||
'reconstruct_image': rec_img,
|
||||
'instruction': prompt[i],
|
||||
'edit_image': edit_imgs if len(edit_imgs) > 0 else None,
|
||||
'edit_mask': edit_img_masks if len(edit_imgs) > 0 else None
|
||||
}
|
||||
if image is not None:
|
||||
if image_mask is None:
|
||||
image_mask = [None] * len(image)
|
||||
ori_img = torch.clamp((image[i] + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
one_tup['target_image'] = ori_img.squeeze(0)
|
||||
one_tup['target_mask'] = image_mask[i] if image_mask[
|
||||
i] is not None else torch.ones_like(ori_img[[0], :, :])
|
||||
outputs.append(one_tup)
|
||||
return outputs
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODEL',
|
||||
__class__.__name__,
|
||||
LatentDiffusionACEPlus.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,263 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import random
|
||||
import torch
|
||||
from typing import Tuple
|
||||
|
||||
from scepter.modules.model.network.ldm import LatentDiffusion
|
||||
from scepter.modules.model.registry import MODELS
|
||||
from scepter.modules.model.utils.basic_utils import default
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
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
|
||||
|
||||
|
||||
@MODELS.register_class()
|
||||
class LatentDiffusionCogVideoX(LatentDiffusion):
|
||||
para_dict = LatentDiffusion.para_dict
|
||||
|
||||
def init_params(self):
|
||||
super().init_params()
|
||||
self.latent_channels = self.model_config.get('LATENT_CHANNELS', self.model_config.IN_CHANNELS)
|
||||
self.scale_factor_spatial = self.cfg.get('SCALE_FACTOR_SPATIAL', 8)
|
||||
self.scale_factor_temporal = self.cfg.get('SCALE_FACTOR_TEMPORAL', 4)
|
||||
self.scaling_factor_image = self.cfg.get('SCALING_FACTOR_IMAGE', 0.7)
|
||||
self.use_rotary_positional_embeddings = self.model_config.get('USE_ROTARY_POSITIONAL_EMBEDDINGS', False)
|
||||
self.attention_head_dim = self.model_config.get('ATTENTION_HEAD_DIM', 64)
|
||||
self.patch_size = self.model_config.get('PATCH_SIZE', 2)
|
||||
self.patch_size_t = self.model_config.get('PATCH_SIZE_T', None)
|
||||
self.ofs_embed_dim = self.model_config.get('OFS_EMBED_DIM', None)
|
||||
self.sample_height = self.first_stage_config.get('SAMPLE_HEIGHT', 60)
|
||||
self.sample_width = self.first_stage_config.get('SAMPLE_WIDTH', 90)
|
||||
self.noised_image_dropout = self.cfg.get('NOISED_IMAGE_DROPOUT', 0.05)
|
||||
self.invert_scale_latents = self.cfg.get('INVERT_SCALE_LATENTS', False)
|
||||
|
||||
def construct_network(self):
|
||||
super().construct_network()
|
||||
self.model = self.model.to(getattr(torch, self.model_config.DTYPE))
|
||||
self.first_stage_model = self.first_stage_model.to(getattr(torch, self.first_stage_config.DTYPE))
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_first_stage(self, x, **kwargs):
|
||||
if isinstance(x, list):
|
||||
x = torch.stack(x, dim=0) # [B, C, F, H, W]
|
||||
image_latents = self.first_stage_model.encode(x).sample()
|
||||
if not self.invert_scale_latents:
|
||||
latents = self.scaling_factor_image * image_latents
|
||||
else:
|
||||
latents = 1 / self.scaling_factor_image * image_latents
|
||||
return latents
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, latents):
|
||||
latents = latents.permute(0, 2, 1, 3, 4) # [batch_size, num_channels, num_frames, height, width]
|
||||
latents = 1 / self.scaling_factor_image * latents
|
||||
frames = self.first_stage_model.decode(latents)
|
||||
return frames
|
||||
|
||||
def get_image_latent(self, image, video, noise):
|
||||
latent = torch.zeros_like(noise)
|
||||
if isinstance(image, list):
|
||||
image = torch.stack(image, dim=0) # [B, C, F, H, W]
|
||||
if len(image.shape) == 4: # [B, C, H, W]
|
||||
image = image.unsqueeze(2) # [B, C, F, H, W]
|
||||
image_latent = self.encode_first_stage(image) # [B, C, F, H, W]
|
||||
image_latent = image_latent.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
|
||||
latent[:, :1, :, :, :] = image_latent
|
||||
return latent, image
|
||||
|
||||
def noise_sample(self, batch_size, num_frames, height, width, generator, dtype=torch.bfloat16):
|
||||
shape = (batch_size,
|
||||
(num_frames - 1) // self.scale_factor_temporal + 1,
|
||||
self.latent_channels,
|
||||
height // self.scale_factor_spatial,
|
||||
width // self.scale_factor_spatial
|
||||
)
|
||||
noise = torch.randn(shape, generator=generator, dtype=dtype, device='cpu').to(we.device_id)
|
||||
return noise
|
||||
|
||||
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.scale_factor_spatial * self.patch_size)
|
||||
grid_width = width // (self.scale_factor_spatial * self.patch_size)
|
||||
|
||||
p = self.patch_size
|
||||
p_t = self.patch_size_t
|
||||
|
||||
base_size_width = self.sample_width // p
|
||||
base_size_height = self.sample_height // p
|
||||
|
||||
if p_t is None:
|
||||
# CogVideoX 1.0
|
||||
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.attention_head_dim,
|
||||
crops_coords=grid_crops_coords,
|
||||
grid_size=(grid_height, grid_width),
|
||||
temporal_size=num_frames,
|
||||
)
|
||||
else:
|
||||
# CogVideoX 1.5
|
||||
base_num_frames = (num_frames + p_t - 1) // p_t
|
||||
|
||||
freqs_cos, freqs_sin = get_3d_rotary_pos_embed(
|
||||
embed_dim=self.attention_head_dim,
|
||||
crops_coords=None,
|
||||
grid_size=(grid_height, grid_width),
|
||||
temporal_size=base_num_frames,
|
||||
grid_type="slice",
|
||||
max_size=(base_size_height, base_size_width),
|
||||
)
|
||||
|
||||
freqs_cos = freqs_cos.to(device=device)
|
||||
freqs_sin = freqs_sin.to(device=device)
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
|
||||
def forward_train(self, video=None, video_latent=None, image=None, noise=None, prompt=None, image_size=None, **kwargs):
|
||||
# video: [B, C, F, H, W]
|
||||
if image_size is None: image_size = [480, 720]
|
||||
if video_latent is not None:
|
||||
x_start = torch.stack(video_latent)
|
||||
else:
|
||||
x_start = self.encode_first_stage(video, **kwargs)
|
||||
x_start = x_start.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
|
||||
t = torch.randint(low=0, high=self.num_timesteps, size=(len(video),), device=we.device_id)
|
||||
|
||||
if prompt and self.cond_stage_model:
|
||||
with torch.autocast(device_type='cuda', enabled=True, dtype=torch.bfloat16):
|
||||
cont = getattr(self.cond_stage_model, 'encode')(prompt, return_mask=False, use_mask=False)
|
||||
|
||||
if noise is None:
|
||||
noise = torch.randn_like(x_start)
|
||||
|
||||
if image is not None:
|
||||
if random.random() < self.noised_image_dropout:
|
||||
image_latent = torch.zeros_like(noise)
|
||||
else:
|
||||
image_latent, _ = self.get_image_latent(image, video, noise)
|
||||
else:
|
||||
image_latent = None
|
||||
|
||||
height, width = image_size[0] if isinstance(image_size, list) and all(isinstance(elem, list) for elem in image_size) else image_size
|
||||
image_rotary_emb = (
|
||||
self._prepare_rotary_positional_embeddings(height=height, width=width, num_frames=noise.size(1), device=we.device_id)
|
||||
if self.use_rotary_positional_embeddings
|
||||
else None
|
||||
)
|
||||
ofs_emb = None if self.ofs_embed_dim is None else image_latent.new_full((1,), fill_value=2.0)
|
||||
|
||||
loss = self.diffusion.loss(x_0=x_start,
|
||||
t=t,
|
||||
model=self.model,
|
||||
model_kwargs={"cond": cont,
|
||||
'image_latent': image_latent,
|
||||
'image_rotary_emb': image_rotary_emb,
|
||||
'ofs': ofs_emb},
|
||||
noise=noise,
|
||||
**kwargs)
|
||||
loss = loss.mean()
|
||||
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
|
||||
return ret
|
||||
|
||||
@torch.no_grad()
|
||||
@torch.autocast('cuda', dtype=torch.bfloat16)
|
||||
def forward_test(self,
|
||||
video=None,
|
||||
image=None,
|
||||
prompt=None,
|
||||
n_prompt=None,
|
||||
sampler='ddim',
|
||||
sample_steps=50,
|
||||
seed=42,
|
||||
guide_scale=6.0,
|
||||
guide_rescale=0.0,
|
||||
num_frames=49,
|
||||
image_size=None,
|
||||
show_process=False,
|
||||
**kwargs):
|
||||
if image_size is None:
|
||||
image_size = [480, 720]
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||
generator = torch.Generator().manual_seed(seed)
|
||||
# generator = torch.Generator(we.device_id).manual_seed(seed)
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
num_samples = len(prompt)
|
||||
n_prompt = default(n_prompt, [self.default_n_prompt] * len(prompt))
|
||||
|
||||
if prompt and self.cond_stage_model:
|
||||
with torch.autocast(device_type='cuda', enabled=True, dtype=torch.bfloat16):
|
||||
cont = getattr(self.cond_stage_model, 'encode')(prompt, return_mask=False, use_mask=False)
|
||||
null_cont = getattr(self.cond_stage_model, 'encode')(n_prompt, return_mask=False, use_mask=False)
|
||||
|
||||
height, width = image_size[0] if isinstance(image_size, list) and all(isinstance(elem, list) for elem in image_size) else image_size
|
||||
latent_frames = (num_frames - 1) // self.scale_factor_temporal + 1
|
||||
additional_frames = 0
|
||||
if self.patch_size_t is not None and latent_frames % self.patch_size_t != 0:
|
||||
additional_frames = self.patch_size_t - latent_frames % self.patch_size_t
|
||||
num_frames += additional_frames * self.scale_factor_temporal
|
||||
noise = self.noise_sample(num_samples, num_frames, height, width, generator)
|
||||
image_rotary_emb = (
|
||||
self._prepare_rotary_positional_embeddings(height, width, noise.size(1), we.device_id)
|
||||
if self.use_rotary_positional_embeddings
|
||||
else None
|
||||
)
|
||||
|
||||
image_latent, image = self.get_image_latent(image, video, noise) if image is not None else (None, None)
|
||||
ofs_emb = None if self.ofs_embed_dim is None else image_latent.new_full((1,), fill_value=2.0)
|
||||
|
||||
samples = self.diffusion.sample(noise=noise,
|
||||
sampler=sampler,
|
||||
model=self.model,
|
||||
model_kwargs=[{
|
||||
'cond': cont,
|
||||
'image_latent': image_latent,
|
||||
'image_rotary_emb': image_rotary_emb,
|
||||
'ofs': ofs_emb
|
||||
}, {
|
||||
'cond': null_cont,
|
||||
'image_latent': image_latent,
|
||||
'image_rotary_emb': image_rotary_emb,
|
||||
'ofs': ofs_emb
|
||||
}],
|
||||
steps=sample_steps,
|
||||
show_progress=True,
|
||||
guide_scale=guide_scale,
|
||||
guide_rescale=guide_rescale,
|
||||
return_intermediate=None,
|
||||
**kwargs).float()
|
||||
|
||||
samples = samples[:, additional_frames:]
|
||||
x_frames = self.decode_first_stage(samples).float()
|
||||
|
||||
outputs = []
|
||||
for batch_idx in range(num_samples):
|
||||
rec_video = torch.clamp(x_frames[batch_idx] / 2 + 0.5, min=0.0, max=1.0)
|
||||
one_tup = {
|
||||
'reconstruct_video': rec_video.squeeze(0).float(),
|
||||
'instruction': prompt[batch_idx]
|
||||
}
|
||||
if image is not None:
|
||||
ori_image = torch.clamp(image[batch_idx] / 2 + 0.5, min=0.0, max=1.0)
|
||||
one_tup['edit_image'] = ori_image
|
||||
if video is not None:
|
||||
ori_video = torch.clamp(video[batch_idx] / 2 + 0.5, min=0.0, max=1.0)
|
||||
one_tup['target_video'] = ori_video.squeeze(0)
|
||||
outputs.append(one_tup)
|
||||
return outputs
|
||||
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODEL',
|
||||
__class__.__name__,
|
||||
LatentDiffusionCogVideoX.para_dict,
|
||||
set_name=True)
|
||||
@@ -4,22 +4,23 @@ import copy
|
||||
import math
|
||||
import numbers
|
||||
import random
|
||||
from contextlib import nullcontext
|
||||
|
||||
import torch
|
||||
from scepter.modules.model.network.ldm import LatentDiffusion
|
||||
from scepter.modules.model.registry import MODELS, BACKBONES, LOSSES, TOKENIZERS, EMBEDDERS, DIFFUSIONS
|
||||
from scepter.modules.model.utils.basic_utils import disabled_train
|
||||
from scepter.modules.model.utils.basic_utils import disabled_train, check_list_of_list, to_device, \
|
||||
pack_imagelist_into_tensor, unpack_tensor_into_imagelist, limit_batch_data
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.model.utils.basic_utils import count_params
|
||||
|
||||
|
||||
@MODELS.register_class()
|
||||
class LatentDiffusionFlux(LatentDiffusion):
|
||||
para_dict = LatentDiffusion.para_dict
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.guide_scale = cfg.get('GUIDE_SCALE', 3.5)
|
||||
self.guide_scale = cfg.get('GUIDE_SCALE', 1.0)
|
||||
|
||||
def init_params(self):
|
||||
self.parameterization = self.cfg.get('PARAMETERIZATION', 'rf')
|
||||
@@ -137,7 +138,7 @@ class LatentDiffusionFlux(LatentDiffusion):
|
||||
def forward_test(self,
|
||||
image=None,
|
||||
prompt=None,
|
||||
sampler='flow_eluer',
|
||||
sampler='flow_euler',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=4.5,
|
||||
@@ -218,3 +219,322 @@ class LatentDiffusionFlux(LatentDiffusion):
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, z):
|
||||
return self.first_stage_model.decode(z)
|
||||
|
||||
@MODELS.register_class()
|
||||
class LatentDiffusionFluxMR(LatentDiffusionFlux):
|
||||
para_dict = {
|
||||
}
|
||||
para_dict.update(LatentDiffusion.para_dict)
|
||||
def forward_train(self,
|
||||
image=None,
|
||||
noise=None,
|
||||
prompt=[],
|
||||
**kwargs):
|
||||
if check_list_of_list(prompt):
|
||||
prompt = [pp[0] for pp in prompt]
|
||||
assert self.cond_stage_model is not None
|
||||
gc_seg = kwargs.pop("gc_seg", [])
|
||||
gc_seg = int(gc_seg[0]) if len(gc_seg) > 0 else 0
|
||||
context = getattr(self.cond_stage_model, 'encode')(prompt)
|
||||
|
||||
image = to_device(image)
|
||||
x_start = self.encode_first_stage(image, **kwargs)
|
||||
loss_mask, _ = pack_imagelist_into_tensor(tuple(torch.ones_like(ix, dtype=torch.bool, device=ix.device) for ix in x_start))
|
||||
x_start, x_shapes = pack_imagelist_into_tensor(x_start)
|
||||
context['x_shapes'] = x_shapes
|
||||
guide_scale = self.guide_scale
|
||||
if guide_scale is not None:
|
||||
guide_scale = torch.full((x_start.shape[0],), guide_scale, device=x_start.device, dtype=x_start.dtype)
|
||||
else:
|
||||
guide_scale = None
|
||||
loss = self.diffusion.loss(x_0=x_start,
|
||||
model=self.model,
|
||||
model_kwargs={"cond": context,
|
||||
"gc_seg": gc_seg,
|
||||
"guidance": guide_scale},
|
||||
noise=None,
|
||||
reduction='none',
|
||||
**kwargs)
|
||||
loss = loss[loss_mask].mean()
|
||||
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
|
||||
return ret
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_sample(self,
|
||||
noise = None,
|
||||
prompt=None,
|
||||
sampler='flow_euler',
|
||||
sample_steps=20,
|
||||
guide_scale=3.5,
|
||||
show_process=True,
|
||||
x = None,
|
||||
reverse_scale = -1.,
|
||||
**kwargs
|
||||
):
|
||||
noise, x_shapes = pack_imagelist_into_tensor(noise)
|
||||
if x is not None:
|
||||
x, _ = pack_imagelist_into_tensor(x)
|
||||
context = getattr(self.cond_stage_model, 'encode')(prompt)
|
||||
context["x_shapes"] = x_shapes
|
||||
guide_scale = guide_scale or self.guide_scale
|
||||
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
|
||||
# UNet use input n_prompt
|
||||
model = self.model_ema if self.use_ema and self.eval_ema else self.model
|
||||
embedding_context = model.no_sync if isinstance(model, torch.distributed.fsdp.FullyShardedDataParallel) \
|
||||
else nullcontext
|
||||
with embedding_context():
|
||||
x_samples = self.diffusion.sample(
|
||||
noise=noise,
|
||||
sampler=sampler,
|
||||
model=self.model,
|
||||
model_kwargs={"cond": context, "guidance": guide_scale, "gc_seg": -1},
|
||||
steps=sample_steps,
|
||||
show_progress=True,
|
||||
guide_scale=guide_scale,
|
||||
return_intermediate=None,
|
||||
reverse_scale = reverse_scale,
|
||||
x = x,
|
||||
**kwargs).float()
|
||||
x_samples = unpack_tensor_into_imagelist(x_samples, x_shapes)
|
||||
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
|
||||
x_samples = self.decode_first_stage(x_samples)
|
||||
return x_samples
|
||||
@torch.no_grad()
|
||||
def forward_test(self,
|
||||
image=None,
|
||||
prompt=[],
|
||||
sampler='flow_euler',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=3.5,
|
||||
guide_rescale=0.0,
|
||||
show_process=True,
|
||||
log_num = -1,
|
||||
**kwargs):
|
||||
|
||||
if check_list_of_list(prompt):
|
||||
prompt = [pp[0] for pp in prompt]
|
||||
assert self.cond_stage_model is not None
|
||||
# gc_seg is unused
|
||||
prompt, image = limit_batch_data([prompt, image], log_num)
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||
|
||||
if 'index' in kwargs:
|
||||
kwargs.pop('index')
|
||||
if image is not None:
|
||||
noise = [self.noise_sample(1, ix.shape[1], ix.shape[2], seed) for ix in image]
|
||||
else:
|
||||
image_size = None
|
||||
if 'meta' in kwargs:
|
||||
meta = kwargs.pop('meta')
|
||||
if 'image_size' in meta:
|
||||
h = int(meta['image_size'][0][0])
|
||||
w = int(meta['image_size'][1][0])
|
||||
image_size = [h, w]
|
||||
if 'image_size' in kwargs:
|
||||
image_size = kwargs.pop('image_size')
|
||||
if isinstance(image_size, numbers.Number):
|
||||
image_size = [image_size, image_size]
|
||||
if image_size is None:
|
||||
image_size = [1024, 1024]
|
||||
height, width = image_size
|
||||
noise = [self.noise_sample(1, height, width, seed) for _ in prompt]
|
||||
|
||||
x_samples = self.forward_sample(
|
||||
prompt=prompt,
|
||||
sampler=sampler,
|
||||
sample_steps=sample_steps,
|
||||
guide_scale=guide_scale,
|
||||
show_process=show_process,
|
||||
noise=noise,
|
||||
)
|
||||
|
||||
|
||||
outputs = list()
|
||||
for i in range(len(prompt)):
|
||||
rec_img = torch.clamp((x_samples[i].float() + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
rec_img = rec_img.squeeze(0)
|
||||
one_tup = {'prompt': prompt[i], 'n_prompt': '', 'image': rec_img}
|
||||
outputs.append(one_tup)
|
||||
return outputs
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODEL',
|
||||
__class__.__name__,
|
||||
LatentDiffusionFlux.para_dict,
|
||||
set_name=True)
|
||||
@torch.no_grad()
|
||||
def encode_first_stage(self, x, **kwargs):
|
||||
def run_one_image(u):
|
||||
zu = 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
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, z):
|
||||
return [self.first_stage_model.decode(zu) for zu in z]
|
||||
|
||||
@MODELS.register_class()
|
||||
class LatentDiffusionFluxMRRedux(LatentDiffusionFluxMR):
|
||||
para_dict = {
|
||||
}
|
||||
para_dict.update(LatentDiffusionFluxMR.para_dict)
|
||||
|
||||
def init_params(self):
|
||||
super().init_params()
|
||||
self.redux_adapter_cfg = self.cfg.get("REDUX_ADAPTER", None)
|
||||
|
||||
|
||||
def construct_network(self):
|
||||
super().construct_network()
|
||||
if self.redux_adapter_cfg is not None:
|
||||
self.redux_adapter = EMBEDDERS.build(self.redux_adapter_cfg, logger=self.logger).eval().requires_grad_(False)
|
||||
|
||||
def forward_train(self,
|
||||
image=None,
|
||||
noise=None,
|
||||
prompt=[],
|
||||
**kwargs):
|
||||
if check_list_of_list(prompt):
|
||||
prompt = [pp[0] for pp in prompt]
|
||||
assert self.cond_stage_model is not None
|
||||
gc_seg = kwargs.pop("gc_seg", [])
|
||||
gc_seg = int(gc_seg[0]) if len(gc_seg) > 0 else 0
|
||||
context = getattr(self.cond_stage_model, 'encode')(prompt)
|
||||
|
||||
image = to_device(image)
|
||||
x_start = self.encode_first_stage(image, **kwargs)
|
||||
loss_mask, _ = pack_imagelist_into_tensor(tuple(torch.ones_like(ix, dtype=torch.bool, device=ix.device) for ix in x_start))
|
||||
x_start, x_shapes = pack_imagelist_into_tensor(x_start)
|
||||
context['x_shapes'] = x_shapes
|
||||
guide_scale = self.guide_scale
|
||||
if guide_scale is not None:
|
||||
guide_scale = torch.full((x_start.shape[0],), guide_scale, device=x_start.device, dtype=x_start.dtype)
|
||||
else:
|
||||
guide_scale = None
|
||||
loss = self.diffusion.loss(x_0=x_start,
|
||||
model=self.model,
|
||||
model_kwargs={"cond": context,
|
||||
"gc_seg": gc_seg,
|
||||
"guidance": guide_scale},
|
||||
noise=None,
|
||||
reduction='none',
|
||||
**kwargs)
|
||||
loss = loss[loss_mask].mean()
|
||||
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
|
||||
return ret
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_sample(self,
|
||||
noise = None,
|
||||
prompt=None,
|
||||
sampler='flow_euler',
|
||||
sample_steps=20,
|
||||
guide_scale=3.5,
|
||||
show_process=True,
|
||||
x = None,
|
||||
reverse_scale = -1.,
|
||||
**kwargs
|
||||
):
|
||||
noise, x_shapes = pack_imagelist_into_tensor(noise)
|
||||
if x is not None:
|
||||
x, _ = pack_imagelist_into_tensor(x)
|
||||
context = getattr(self.cond_stage_model, 'encode')(prompt)
|
||||
context["x_shapes"] = x_shapes
|
||||
guide_scale = guide_scale or self.guide_scale
|
||||
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
|
||||
# UNet use input n_prompt
|
||||
model = self.model_ema if self.use_ema and self.eval_ema else self.model
|
||||
embedding_context = model.no_sync if isinstance(model, torch.distributed.fsdp.FullyShardedDataParallel) \
|
||||
else nullcontext
|
||||
with embedding_context():
|
||||
x_samples = self.diffusion.sample(
|
||||
noise=noise,
|
||||
sampler=sampler,
|
||||
model=self.model,
|
||||
model_kwargs={"cond": context, "guidance": guide_scale, "gc_seg": -1},
|
||||
steps=sample_steps,
|
||||
show_progress=True,
|
||||
guide_scale=guide_scale,
|
||||
return_intermediate=None,
|
||||
reverse_scale = reverse_scale,
|
||||
x = x,
|
||||
**kwargs).float()
|
||||
x_samples = unpack_tensor_into_imagelist(x_samples, x_shapes)
|
||||
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
|
||||
x_samples = self.decode_first_stage(x_samples)
|
||||
return x_samples
|
||||
@torch.no_grad()
|
||||
def forward_test(self,
|
||||
image=None,
|
||||
prompt=[],
|
||||
sampler='flow_euler',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=3.5,
|
||||
guide_rescale=0.0,
|
||||
show_process=True,
|
||||
log_num = -1,
|
||||
**kwargs):
|
||||
|
||||
if check_list_of_list(prompt):
|
||||
prompt = [pp[0] for pp in prompt]
|
||||
assert self.cond_stage_model is not None
|
||||
# gc_seg is unused
|
||||
prompt, image = limit_batch_data([prompt, image], log_num)
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||
|
||||
if 'index' in kwargs:
|
||||
kwargs.pop('index')
|
||||
if image is not None:
|
||||
noise = [self.noise_sample(1, ix.shape[1], ix.shape[2], seed) for ix in image]
|
||||
else:
|
||||
image_size = None
|
||||
if 'meta' in kwargs:
|
||||
meta = kwargs.pop('meta')
|
||||
if 'image_size' in meta:
|
||||
h = int(meta['image_size'][0][0])
|
||||
w = int(meta['image_size'][1][0])
|
||||
image_size = [h, w]
|
||||
if 'image_size' in kwargs:
|
||||
image_size = kwargs.pop('image_size')
|
||||
if isinstance(image_size, numbers.Number):
|
||||
image_size = [image_size, image_size]
|
||||
if image_size is None:
|
||||
image_size = [1024, 1024]
|
||||
height, width = image_size
|
||||
noise = [self.noise_sample(1, height, width, seed) for _ in prompt]
|
||||
|
||||
x_samples = self.forward_sample(
|
||||
prompt=prompt,
|
||||
sampler=sampler,
|
||||
sample_steps=sample_steps,
|
||||
guide_scale=guide_scale,
|
||||
show_process=show_process,
|
||||
noise=noise,
|
||||
)
|
||||
|
||||
|
||||
outputs = list()
|
||||
for i in range(len(prompt)):
|
||||
rec_img = torch.clamp((x_samples[i].float() + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
rec_img = rec_img.squeeze(0)
|
||||
one_tup = {'prompt': prompt[i], 'n_prompt': '', 'image': rec_img}
|
||||
outputs.append(one_tup)
|
||||
return outputs
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODEL',
|
||||
__class__.__name__,
|
||||
LatentDiffusionFluxMR.para_dict,
|
||||
set_name=True)
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user