Compare commits
48
Commits
v1.3.0_dev
...
v1.4.1_dev
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6c8af8d7a8 | ||
|
|
467652bd69 | ||
|
|
ae59b96f1d | ||
|
|
2c75835035 | ||
|
|
5cbcd3ee04 | ||
|
|
0f12e4db73 | ||
|
|
a9b6337ae2 | ||
|
|
8526ae0234 | ||
|
|
1cd5604e7b | ||
|
|
bbb8f35f49 | ||
|
|
825e8c1cdb | ||
|
|
ff3ccd6050 | ||
|
|
9ad6f9dc5c | ||
|
|
755b39b968 | ||
|
|
3427d630a2 | ||
|
|
d9cd203be6 | ||
|
|
0a10447558 | ||
|
|
32b48d2f08 | ||
|
|
2122221697 | ||
|
|
e591d2e4cb | ||
|
|
33b8adda82 | ||
|
|
1ef6b4d4ec | ||
|
|
cc3e6868ce | ||
|
|
1758959cb9 | ||
|
|
1d7b829c17 | ||
|
|
59c5fadd77 | ||
|
|
cb23134173 | ||
|
|
c4a593b02f | ||
|
|
7dde2741fb | ||
|
|
3713445817 | ||
|
|
5af3ae0eeb | ||
|
|
c3376064fd | ||
|
|
ff7b6c1891 | ||
|
|
db65a08cea | ||
|
|
3083d509de | ||
|
|
740a6c87e0 | ||
|
|
9d5b3a7101 | ||
|
|
c4b4d88e75 | ||
|
|
043222de49 | ||
|
|
171c7ce0b3 | ||
|
|
11c69121cd | ||
|
|
d7dbdc5292 | ||
|
|
711c10a68c | ||
|
|
adda36e39d | ||
|
|
1da1864993 | ||
|
|
5eac362325 | ||
|
|
9ae3bca43d | ||
|
|
ca034ef765 |
@@ -0,0 +1,24 @@
|
|||||||
|
name: Publish to Comfy registry
|
||||||
|
on:
|
||||||
|
workflow_dispatch:
|
||||||
|
push:
|
||||||
|
branches:
|
||||||
|
- main
|
||||||
|
- master
|
||||||
|
paths:
|
||||||
|
- "scepter/workflow/pyproject.toml"
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
publish-node:
|
||||||
|
name: Publish Custom Node to registry
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
# if this is a forked repository. Skipping the workflow.
|
||||||
|
if: github.event.repository.fork == false
|
||||||
|
steps:
|
||||||
|
- name: Check out code
|
||||||
|
uses: actions/checkout@v4
|
||||||
|
- name: Publish Custom Node
|
||||||
|
uses: Comfy-Org/publish-node-action@main
|
||||||
|
with:
|
||||||
|
## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||||
|
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||||
@@ -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}
|
||||||
|
}
|
||||||
|
```
|
||||||
@@ -1,10 +0,0 @@
|
|||||||
name: scepter
|
|
||||||
channels:
|
|
||||||
- defaults
|
|
||||||
dependencies:
|
|
||||||
- python==3.8
|
|
||||||
- pip>=20.3
|
|
||||||
- numpy>=1.23.1
|
|
||||||
- pip:
|
|
||||||
- -r requirements/recommended.txt
|
|
||||||
- -r requirements.txt
|
|
||||||
@@ -18,10 +18,8 @@ SCEPTER offers 3 core components:
|
|||||||
|
|
||||||
|
|
||||||
## 🎉 News
|
## 🎉 News
|
||||||
- [🔥🔥🔥2024.11]: We're excited to announce the upcoming release of the [ACE-0.6b-1024px](https://huggingface.co/scepter-studio/ACE-0.6B-1024px) model,
|
- [🔥🔥🔥 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/).
|
||||||
which significantly enhances image generation quality compared with [ACE-0.6b-512px](https://huggingface.co/scepter-studio/ACE-0.6B-512px). The detailed documents can be found at [ACE repo](https://github.com/ali-vilab/ACE.git).
|
- [2024.11]: Supports video files, video annotation, caption translation in data management, and inference & training of the [CogVideoX](https://arxiv.org/abs/2408.06072).
|
||||||
At the same time, based on the editing results of ACE, combined with the powerful text-to-image capabilities of the [FLUX-dev](https://huggingface.co/black-forest-labs/FLUX.1-dev) model through SDEdit as an image quality refiner, the quality of image editing can be further enhanced.
|
|
||||||
- [🔥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]: 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.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.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/).
|
||||||
@@ -35,123 +33,65 @@ At the same time, based on the editing results of ACE, combined with the powerfu
|
|||||||
- [2023.12]: We propose [SCEdit](https://arxiv.org/abs/2312.11392), an efficient and controllable generation framework.
|
- [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.
|
- [2023.12]: We release [🪄SCEPTER](https://github.com/modelscope/scepter/) library.
|
||||||
|
|
||||||
|
[//]: # (## 🖼 Gallery for Recent Works)
|
||||||
|
|
||||||
|
[//]: # ()
|
||||||
|
[//]: # (### FLUX Tuners)
|
||||||
|
|
||||||
|
[//]: # ()
|
||||||
|
[//]: # (<table><tbody>)
|
||||||
|
|
||||||
## 🪄ACE
|
[//]: # ( <tr>)
|
||||||
|
|
||||||
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.
|
[//]: # ( <th align="center" colspan="3">Yarn Style</th>)
|
||||||
|
|
||||||
[](https://ali-vilab.github.io/ace-page/)
|
[//]: # ( <th align="center" colspan="3">Soft Watercolor Style</th>)
|
||||||
|
|
||||||
### ACE Models
|
[//]: # ( </tr>)
|
||||||
| **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) | |
|
|
||||||
| ACE-12B-FLUX-dev | Coming Soon |
|
|
||||||
### ACE 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 `scepter/methods/edit/dit_ace_0.6b_512.yaml`.
|
[//]: # ( <tr>)
|
||||||
|
|
||||||
#### Prepare datasets
|
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_2_1.webp" width="200"></td>)
|
||||||
|
|
||||||
Please find the dataset class located in `scepter/modules/data/dataset/ms_dataset.py`,
|
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_2_2.webp" width="200"></td>)
|
||||||
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 `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_3.webp" width="200"></td>)
|
||||||
|
|
||||||
#### Prepare initial weight
|
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_1_1.webp" width="200"></td>)
|
||||||
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").
|
[//]: # ( <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>)
|
||||||
|
|
||||||
#### Start training
|
[//]: # ( </tr>)
|
||||||
|
|
||||||
You can easily start training procedure by executing the following command:
|
[//]: # ( <tr>)
|
||||||
```bash
|
|
||||||
# ACE-0.6B-512px
|
|
||||||
PYTHONPATH=. python scepter/tools/run_train.py --cfg scepter/methods/edit/dit_ace_0.6b_512.yaml
|
|
||||||
# ACE-0.6B-1024px
|
|
||||||
PYTHONPATH=. python scepter/tools/run_train.py --cfg scepter/methods/edit/dit_ace_0.6b_1024.yaml
|
|
||||||
```
|
|
||||||
|
|
||||||
### ACE Chat Bot
|
[//]: # ( <th align="center" colspan="3">Travel Style</th>)
|
||||||
|
|
||||||
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:
|
[//]: # ( <th align="center" colspan="3">WuKong Style</th>)
|
||||||
```bash
|
|
||||||
PYTHONPATH=. python scepter/tools/webui.py --cfg scepter/methods/studio/scepter_ui.yaml --language zh --tab chatbot
|
|
||||||
```
|
|
||||||
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.
|
|
||||||
|
|
||||||
### ACE ComfyUI Workflow
|
[//]: # ( </tr>)
|
||||||
|
|
||||||

|
[//]: # ( <tr>)
|
||||||
|
|
||||||
<table><tbody>
|
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_3_1.webp" width="200"></td>)
|
||||||
<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>
|
|
||||||
|
|
||||||
## 🖼 Gallery for Recent Works
|
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_3_2.webp" width="200"></td>)
|
||||||
|
|
||||||
### FLUX Tuners
|
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_3_3.webp" width="200"></td>)
|
||||||
|
|
||||||
<table><tbody>
|
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_4_1.webp" width="200"></td>)
|
||||||
<tr>
|
|
||||||
<th align="center" colspan="3">Yarn Style</th>
|
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_4_2.webp" width="200"></td>)
|
||||||
<th align="center" colspan="3">Soft Watercolor Style</th>
|
|
||||||
</tr>
|
[//]: # ( <td><img src="asset/images/flux_tuner/flux_tuner_4_3.webp" width="200"></td>)
|
||||||
<tr>
|
|
||||||
<td><img src="asset/images/flux_tuner/flux_tuner_2_1.webp" width="200"></td>
|
[//]: # ( </tr>)
|
||||||
<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>
|
[//]: # (</tbody>)
|
||||||
<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>
|
[//]: # (</table>)
|
||||||
<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>
|
|
||||||
|
|
||||||
### ComfyUI Workflow
|
### ComfyUI Workflow
|
||||||
|
|
||||||
@@ -195,13 +135,6 @@ Upon starting, you will find a "ChatBot" tab within the Gradio application, whic
|
|||||||
|
|
||||||
## 🛠️ Installation
|
## 🛠️ Installation
|
||||||
|
|
||||||
- Create new environment with `conda` command:
|
|
||||||
|
|
||||||
```shell
|
|
||||||
conda env create -f environment.yaml
|
|
||||||
conda activate scepter
|
|
||||||
```
|
|
||||||
|
|
||||||
- Install with `pip` command:
|
- Install with `pip` command:
|
||||||
|
|
||||||
We recommend installing the specific version of PyTorch and accelerate toolbox [xFormers](https://pypi.org/project/xformers/). You can install these recommended version by pip:
|
We recommend installing the specific version of PyTorch and accelerate toolbox [xFormers](https://pypi.org/project/xformers/). You can install these recommended version by pip:
|
||||||
@@ -225,18 +158,19 @@ pip install scepter
|
|||||||
|
|
||||||
### Currently supported approaches
|
### Currently supported approaches
|
||||||
|
|
||||||
| Tasks | Methods | Links |
|
| Tasks | Methods | Links |
|
||||||
|:----------------------------:|:----------------------------------------------:|:------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
|
|:----------------------------:|:------------------------------------------------:|:------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
|
||||||
| Text-to-image Generation | SD v1.5 | [](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
|
| 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 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 | 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) |
|
| 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 | LoRA | [](https://arxiv.org/abs/2106.09685) |
|
||||||
| Efficient Tuning | Res-Tuning(NeurIPS23) | [](https://arxiv.org/abs/2310.19859) [](https://res-tuning.github.io/) |
|
| 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/) |
|
| 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 | [🌟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 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-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
|
## 🖥️ SCEPTER Studio
|
||||||
|
|||||||
@@ -1,6 +1,5 @@
|
|||||||
albumentations
|
albumentations
|
||||||
beautifulsoup4
|
beautifulsoup4
|
||||||
bezier
|
|
||||||
einops
|
einops
|
||||||
modelscope[framework]
|
modelscope[framework]
|
||||||
ms-swift
|
ms-swift
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
git+https://github.com/cocodataset/panopticapi.git
|
git+https://github.com/cocodataset/panopticapi.git
|
||||||
torch==2.4.1
|
torch==2.4.1
|
||||||
torchvision==.19.1
|
torchvision==0.19.1
|
||||||
flash-attn==2.5.8
|
flash-attn==2.5.8
|
||||||
xformers==0.0.28
|
xformers==0.0.28
|
||||||
+23
-13
@@ -1,18 +1,28 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# 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__ = [
|
import sys
|
||||||
utils, transform, data, model, solver, version_info, opt, '__version__',
|
sys.modules[__name__] = LazyImportModule(
|
||||||
'dirname'
|
__name__,
|
||||||
]
|
globals()['__file__'],
|
||||||
|
_import_structure,
|
||||||
|
module_spec=__spec__,
|
||||||
|
extra_objects={},
|
||||||
|
)
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -183,6 +183,10 @@ SOLVER:
|
|||||||
PIN_MEMORY: True
|
PIN_MEMORY: True
|
||||||
BATCH_SIZE: 1
|
BATCH_SIZE: 1
|
||||||
NUM_WORKERS: 4
|
NUM_WORKERS: 4
|
||||||
|
NUM_FRAMES: 49
|
||||||
|
FPS: 8
|
||||||
|
HEIGHT: 480
|
||||||
|
WIDTH: 720
|
||||||
PROMPT_PREFIX: 'DISNEY '
|
PROMPT_PREFIX: 'DISNEY '
|
||||||
SAMPLER:
|
SAMPLER:
|
||||||
NAME: MixtureOfSamplers
|
NAME: MixtureOfSamplers
|
||||||
|
|||||||
@@ -188,6 +188,10 @@ SOLVER:
|
|||||||
PIN_MEMORY: True
|
PIN_MEMORY: True
|
||||||
BATCH_SIZE: 1
|
BATCH_SIZE: 1
|
||||||
NUM_WORKERS: 0
|
NUM_WORKERS: 0
|
||||||
|
NUM_FRAMES: 49
|
||||||
|
FPS: 8
|
||||||
|
HEIGHT: 480
|
||||||
|
WIDTH: 720
|
||||||
PROMPT_PREFIX: 'DISNEY '
|
PROMPT_PREFIX: 'DISNEY '
|
||||||
DATA_TYPE: 'i2v'
|
DATA_TYPE: 'i2v'
|
||||||
SAMPLER:
|
SAMPLER:
|
||||||
|
|||||||
@@ -185,6 +185,10 @@ SOLVER:
|
|||||||
PIN_MEMORY: True
|
PIN_MEMORY: True
|
||||||
BATCH_SIZE: 1
|
BATCH_SIZE: 1
|
||||||
NUM_WORKERS: 4
|
NUM_WORKERS: 4
|
||||||
|
NUM_FRAMES: 49
|
||||||
|
FPS: 8
|
||||||
|
HEIGHT: 480
|
||||||
|
WIDTH: 720
|
||||||
PROMPT_PREFIX: 'DISNEY '
|
PROMPT_PREFIX: 'DISNEY '
|
||||||
DELIMITER: '#;#'
|
DELIMITER: '#;#'
|
||||||
FIELDS: [ 'video_path', 'prompt' ]
|
FIELDS: [ 'video_path', 'prompt' ]
|
||||||
|
|||||||
@@ -148,4 +148,5 @@ MODEL:
|
|||||||
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
|
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
|
||||||
LENGTH: 226
|
LENGTH: 226
|
||||||
CLEAN:
|
CLEAN:
|
||||||
USE_GRAD: False
|
USE_GRAD: False
|
||||||
|
T5_DTYPE: bfloat16
|
||||||
@@ -150,4 +150,5 @@ MODEL:
|
|||||||
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
|
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
|
||||||
LENGTH: 226
|
LENGTH: 226
|
||||||
CLEAN:
|
CLEAN:
|
||||||
USE_GRAD: False
|
USE_GRAD: False
|
||||||
|
T5_DTYPE: bfloat16
|
||||||
@@ -1,4 +1,5 @@
|
|||||||
WORK_DIR: "inference"
|
WORK_DIR: "inference"
|
||||||
|
SKIP_EXAMPLES: True
|
||||||
DIFFUSION_PARAS:
|
DIFFUSION_PARAS:
|
||||||
SAMPLE:
|
SAMPLE:
|
||||||
VALUES: ['ddim', 'euler', 'euler_ancestral', 'heun', 'dpm2',
|
VALUES: ['ddim', 'euler', 'euler_ancestral', 'heun', 'dpm2',
|
||||||
|
|||||||
@@ -1,4 +1,23 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
from scepter.modules import (data, inference, model, opt, solver, transform,
|
from typing import TYPE_CHECKING
|
||||||
utils)
|
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,64 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
from scepter.modules.annotator.base_annotator import GeneralAnnotator
|
from typing import TYPE_CHECKING
|
||||||
from scepter.modules.annotator.canny import CannyAnnotator
|
from scepter.modules.utils.import_utils import LazyImportModule
|
||||||
from scepter.modules.annotator.color import ColorAnnotator
|
|
||||||
from scepter.modules.annotator.degradation import DegradationAnnotator
|
if TYPE_CHECKING:
|
||||||
from scepter.modules.annotator.doodle import DoodleAnnotator
|
from scepter.modules.annotator.base_annotator import GeneralAnnotator
|
||||||
from scepter.modules.annotator.gray import GrayAnnotator
|
from scepter.modules.annotator.canny import CannyAnnotator
|
||||||
from scepter.modules.annotator.hed import HedAnnotator
|
from scepter.modules.annotator.color import ColorAnnotator
|
||||||
from scepter.modules.annotator.identity import IdentityAnnotator
|
from scepter.modules.annotator.degradation import DegradationAnnotator
|
||||||
from scepter.modules.annotator.informative_drawing import (
|
from scepter.modules.annotator.doodle import DoodleAnnotator
|
||||||
InfoDrawAnimeAnnotator, InfoDrawContourAnnotator,
|
from scepter.modules.annotator.gray import GrayAnnotator
|
||||||
InfoDrawOpenSketchAnnotator)
|
from scepter.modules.annotator.hed import HedAnnotator
|
||||||
from scepter.modules.annotator.inpainting import InpaintingAnnotator
|
from scepter.modules.annotator.identity import IdentityAnnotator
|
||||||
from scepter.modules.annotator.invert import InvertAnnotator
|
from scepter.modules.annotator.informative_drawing import (
|
||||||
from scepter.modules.annotator.midas_op import MidasDetector
|
InfoDrawAnimeAnnotator, InfoDrawContourAnnotator,
|
||||||
from scepter.modules.annotator.mlsd_op import MLSDdetector
|
InfoDrawOpenSketchAnnotator)
|
||||||
from scepter.modules.annotator.openpose import OpenposeAnnotator
|
from scepter.modules.annotator.inpainting import InpaintingAnnotator
|
||||||
from scepter.modules.annotator.outpainting import OutpaintingAnnotator, OutpaintingResize
|
from scepter.modules.annotator.invert import InvertAnnotator
|
||||||
from scepter.modules.annotator.pidinet import PiDiAnnotator
|
from scepter.modules.annotator.midas_op import MidasDetector
|
||||||
from scepter.modules.annotator.segmentation import ESAMAnnotator
|
from scepter.modules.annotator.mlsd_op import MLSDdetector
|
||||||
from scepter.modules.annotator.sketch import SketchAnnotator
|
from scepter.modules.annotator.openpose import OpenposeAnnotator
|
||||||
from scepter.modules.annotator.lama import LamaAnnotator
|
from scepter.modules.annotator.outpainting import OutpaintingAnnotator, OutpaintingResize
|
||||||
|
from scepter.modules.annotator.pidinet import PiDiAnnotator
|
||||||
|
from scepter.modules.annotator.segmentation import ESAMAnnotator
|
||||||
|
from scepter.modules.annotator.sketch import SketchAnnotator
|
||||||
|
from scepter.modules.annotator.lama import LamaAnnotator
|
||||||
|
from scepter.modules.annotator.mask_aug import MaskAugAnnotator, MaskDrawAnnotator, MaskLayoutAnnotator
|
||||||
|
from scepter.modules.annotator.raft import RAFTAnnotator, RAFTVisAnnotator
|
||||||
|
else:
|
||||||
|
_import_structure = {
|
||||||
|
'base_annotator': ['GeneralAnnotator'],
|
||||||
|
'canny': ['CannyAnnotator'],
|
||||||
|
'color': ['ColorAnnotator'],
|
||||||
|
'degradation': ['DegradationAnnotator'],
|
||||||
|
'doodle': ['DoodleAnnotator'],
|
||||||
|
'gray': ['GrayAnnotator'],
|
||||||
|
'hed': ['HedAnnotator'],
|
||||||
|
'identity': ['IdentityAnnotator'],
|
||||||
|
'informative_drawing': ['InfoDrawAnimeAnnotator',
|
||||||
|
'InfoDrawContourAnnotator',
|
||||||
|
'InfoDrawOpenSketchAnnotator'],
|
||||||
|
'inpainting': ['InpaintingAnnotator'],
|
||||||
|
'invert': ['InvertAnnotator'],
|
||||||
|
'midas_op': ['MidasDetector'],
|
||||||
|
'mlsd_op': ['MLSDdetector'],
|
||||||
|
'openpose': ['OpenposeAnnotator'],
|
||||||
|
'outpainting': ['OutpaintingAnnotator', 'OutpaintingResize'],
|
||||||
|
'pidinet': ['PiDiAnnotator'],
|
||||||
|
'segmentation': ['ESAMAnnotator'],
|
||||||
|
'sketch': ['SketchAnnotator'],
|
||||||
|
'lama': ['LamaAnnotator'],
|
||||||
|
'mask_aug': ['MaskAugAnnotator', 'MaskDrawAnnotator', 'MaskLayoutAnnotator'],
|
||||||
|
'raft': ['RAFTAnnotator', 'RAFTVisAnnotator'],
|
||||||
|
}
|
||||||
|
|
||||||
|
import sys
|
||||||
|
sys.modules[__name__] = LazyImportModule(
|
||||||
|
__name__,
|
||||||
|
globals()['__file__'],
|
||||||
|
_import_structure,
|
||||||
|
module_spec=__spec__,
|
||||||
|
extra_objects={},
|
||||||
|
)
|
||||||
|
|||||||
@@ -0,0 +1,2 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
@@ -0,0 +1,127 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
|
import cv2
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
import onnxruntime
|
||||||
|
|
||||||
|
def nms(boxes, scores, nms_thr):
|
||||||
|
"""Single class NMS implemented in Numpy."""
|
||||||
|
x1 = boxes[:, 0]
|
||||||
|
y1 = boxes[:, 1]
|
||||||
|
x2 = boxes[:, 2]
|
||||||
|
y2 = boxes[:, 3]
|
||||||
|
|
||||||
|
areas = (x2 - x1 + 1) * (y2 - y1 + 1)
|
||||||
|
order = scores.argsort()[::-1]
|
||||||
|
|
||||||
|
keep = []
|
||||||
|
while order.size > 0:
|
||||||
|
i = order[0]
|
||||||
|
keep.append(i)
|
||||||
|
xx1 = np.maximum(x1[i], x1[order[1:]])
|
||||||
|
yy1 = np.maximum(y1[i], y1[order[1:]])
|
||||||
|
xx2 = np.minimum(x2[i], x2[order[1:]])
|
||||||
|
yy2 = np.minimum(y2[i], y2[order[1:]])
|
||||||
|
|
||||||
|
w = np.maximum(0.0, xx2 - xx1 + 1)
|
||||||
|
h = np.maximum(0.0, yy2 - yy1 + 1)
|
||||||
|
inter = w * h
|
||||||
|
ovr = inter / (areas[i] + areas[order[1:]] - inter)
|
||||||
|
|
||||||
|
inds = np.where(ovr <= nms_thr)[0]
|
||||||
|
order = order[inds + 1]
|
||||||
|
|
||||||
|
return keep
|
||||||
|
|
||||||
|
def multiclass_nms(boxes, scores, nms_thr, score_thr):
|
||||||
|
"""Multiclass NMS implemented in Numpy. Class-aware version."""
|
||||||
|
final_dets = []
|
||||||
|
num_classes = scores.shape[1]
|
||||||
|
for cls_ind in range(num_classes):
|
||||||
|
cls_scores = scores[:, cls_ind]
|
||||||
|
valid_score_mask = cls_scores > score_thr
|
||||||
|
if valid_score_mask.sum() == 0:
|
||||||
|
continue
|
||||||
|
else:
|
||||||
|
valid_scores = cls_scores[valid_score_mask]
|
||||||
|
valid_boxes = boxes[valid_score_mask]
|
||||||
|
keep = nms(valid_boxes, valid_scores, nms_thr)
|
||||||
|
if len(keep) > 0:
|
||||||
|
cls_inds = np.ones((len(keep), 1)) * cls_ind
|
||||||
|
dets = np.concatenate(
|
||||||
|
[valid_boxes[keep], valid_scores[keep, None], cls_inds], 1
|
||||||
|
)
|
||||||
|
final_dets.append(dets)
|
||||||
|
if len(final_dets) == 0:
|
||||||
|
return None
|
||||||
|
return np.concatenate(final_dets, 0)
|
||||||
|
|
||||||
|
def demo_postprocess(outputs, img_size, p6=False):
|
||||||
|
grids = []
|
||||||
|
expanded_strides = []
|
||||||
|
strides = [8, 16, 32] if not p6 else [8, 16, 32, 64]
|
||||||
|
|
||||||
|
hsizes = [img_size[0] // stride for stride in strides]
|
||||||
|
wsizes = [img_size[1] // stride for stride in strides]
|
||||||
|
|
||||||
|
for hsize, wsize, stride in zip(hsizes, wsizes, strides):
|
||||||
|
xv, yv = np.meshgrid(np.arange(wsize), np.arange(hsize))
|
||||||
|
grid = np.stack((xv, yv), 2).reshape(1, -1, 2)
|
||||||
|
grids.append(grid)
|
||||||
|
shape = grid.shape[:2]
|
||||||
|
expanded_strides.append(np.full((*shape, 1), stride))
|
||||||
|
|
||||||
|
grids = np.concatenate(grids, 1)
|
||||||
|
expanded_strides = np.concatenate(expanded_strides, 1)
|
||||||
|
outputs[..., :2] = (outputs[..., :2] + grids) * expanded_strides
|
||||||
|
outputs[..., 2:4] = np.exp(outputs[..., 2:4]) * expanded_strides
|
||||||
|
|
||||||
|
return outputs
|
||||||
|
|
||||||
|
def preprocess(img, input_size, swap=(2, 0, 1)):
|
||||||
|
if len(img.shape) == 3:
|
||||||
|
padded_img = np.ones((input_size[0], input_size[1], 3), dtype=np.uint8) * 114
|
||||||
|
else:
|
||||||
|
padded_img = np.ones(input_size, dtype=np.uint8) * 114
|
||||||
|
|
||||||
|
r = min(input_size[0] / img.shape[0], input_size[1] / img.shape[1])
|
||||||
|
resized_img = cv2.resize(
|
||||||
|
img,
|
||||||
|
(int(img.shape[1] * r), int(img.shape[0] * r)),
|
||||||
|
interpolation=cv2.INTER_LINEAR,
|
||||||
|
).astype(np.uint8)
|
||||||
|
padded_img[: int(img.shape[0] * r), : int(img.shape[1] * r)] = resized_img
|
||||||
|
|
||||||
|
padded_img = padded_img.transpose(swap)
|
||||||
|
padded_img = np.ascontiguousarray(padded_img, dtype=np.float32)
|
||||||
|
return padded_img, r
|
||||||
|
|
||||||
|
def inference_detector(session, oriImg):
|
||||||
|
input_shape = (640,640)
|
||||||
|
img, ratio = preprocess(oriImg, input_shape)
|
||||||
|
|
||||||
|
ort_inputs = {session.get_inputs()[0].name: img[None, :, :, :]}
|
||||||
|
output = session.run(None, ort_inputs)
|
||||||
|
predictions = demo_postprocess(output[0], input_shape)[0]
|
||||||
|
|
||||||
|
boxes = predictions[:, :4]
|
||||||
|
scores = predictions[:, 4:5] * predictions[:, 5:]
|
||||||
|
|
||||||
|
boxes_xyxy = np.ones_like(boxes)
|
||||||
|
boxes_xyxy[:, 0] = boxes[:, 0] - boxes[:, 2]/2.
|
||||||
|
boxes_xyxy[:, 1] = boxes[:, 1] - boxes[:, 3]/2.
|
||||||
|
boxes_xyxy[:, 2] = boxes[:, 0] + boxes[:, 2]/2.
|
||||||
|
boxes_xyxy[:, 3] = boxes[:, 1] + boxes[:, 3]/2.
|
||||||
|
boxes_xyxy /= ratio
|
||||||
|
dets = multiclass_nms(boxes_xyxy, scores, nms_thr=0.45, score_thr=0.1)
|
||||||
|
if dets is not None:
|
||||||
|
final_boxes, final_scores, final_cls_inds = dets[:, :4], dets[:, 4], dets[:, 5]
|
||||||
|
isscore = final_scores>0.3
|
||||||
|
iscat = final_cls_inds == 0
|
||||||
|
isbbox = [ i and j for (i, j) in zip(isscore, iscat)]
|
||||||
|
final_boxes = final_boxes[isbbox]
|
||||||
|
else:
|
||||||
|
final_boxes = np.array([])
|
||||||
|
|
||||||
|
return final_boxes
|
||||||
@@ -0,0 +1,362 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
|
from typing import List, Tuple
|
||||||
|
|
||||||
|
import cv2
|
||||||
|
import numpy as np
|
||||||
|
import onnxruntime as ort
|
||||||
|
|
||||||
|
def preprocess(
|
||||||
|
img: np.ndarray, out_bbox, input_size: Tuple[int, int] = (192, 256)
|
||||||
|
) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
|
||||||
|
"""Do preprocessing for RTMPose model inference.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
img (np.ndarray): Input image in shape.
|
||||||
|
input_size (tuple): Input image size in shape (w, h).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple:
|
||||||
|
- resized_img (np.ndarray): Preprocessed image.
|
||||||
|
- center (np.ndarray): Center of image.
|
||||||
|
- scale (np.ndarray): Scale of image.
|
||||||
|
"""
|
||||||
|
# get shape of image
|
||||||
|
img_shape = img.shape[:2]
|
||||||
|
out_img, out_center, out_scale = [], [], []
|
||||||
|
if len(out_bbox) == 0:
|
||||||
|
out_bbox = [[0, 0, img_shape[1], img_shape[0]]]
|
||||||
|
for i in range(len(out_bbox)):
|
||||||
|
x0 = out_bbox[i][0]
|
||||||
|
y0 = out_bbox[i][1]
|
||||||
|
x1 = out_bbox[i][2]
|
||||||
|
y1 = out_bbox[i][3]
|
||||||
|
bbox = np.array([x0, y0, x1, y1])
|
||||||
|
|
||||||
|
# get center and scale
|
||||||
|
center, scale = bbox_xyxy2cs(bbox, padding=1.25)
|
||||||
|
|
||||||
|
# do affine transformation
|
||||||
|
resized_img, scale = top_down_affine(input_size, scale, center, img)
|
||||||
|
|
||||||
|
# normalize image
|
||||||
|
mean = np.array([123.675, 116.28, 103.53])
|
||||||
|
std = np.array([58.395, 57.12, 57.375])
|
||||||
|
resized_img = (resized_img - mean) / std
|
||||||
|
|
||||||
|
out_img.append(resized_img)
|
||||||
|
out_center.append(center)
|
||||||
|
out_scale.append(scale)
|
||||||
|
|
||||||
|
return out_img, out_center, out_scale
|
||||||
|
|
||||||
|
|
||||||
|
def inference(sess: ort.InferenceSession, img: np.ndarray) -> np.ndarray:
|
||||||
|
"""Inference RTMPose model.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
sess (ort.InferenceSession): ONNXRuntime session.
|
||||||
|
img (np.ndarray): Input image in shape.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
outputs (np.ndarray): Output of RTMPose model.
|
||||||
|
"""
|
||||||
|
all_out = []
|
||||||
|
# build input
|
||||||
|
for i in range(len(img)):
|
||||||
|
input = [img[i].transpose(2, 0, 1)]
|
||||||
|
|
||||||
|
# build output
|
||||||
|
sess_input = {sess.get_inputs()[0].name: input}
|
||||||
|
sess_output = []
|
||||||
|
for out in sess.get_outputs():
|
||||||
|
sess_output.append(out.name)
|
||||||
|
|
||||||
|
# run model
|
||||||
|
outputs = sess.run(sess_output, sess_input)
|
||||||
|
all_out.append(outputs)
|
||||||
|
|
||||||
|
return all_out
|
||||||
|
|
||||||
|
|
||||||
|
def postprocess(outputs: List[np.ndarray],
|
||||||
|
model_input_size: Tuple[int, int],
|
||||||
|
center: Tuple[int, int],
|
||||||
|
scale: Tuple[int, int],
|
||||||
|
simcc_split_ratio: float = 2.0
|
||||||
|
) -> Tuple[np.ndarray, np.ndarray]:
|
||||||
|
"""Postprocess for RTMPose model output.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
outputs (np.ndarray): Output of RTMPose model.
|
||||||
|
model_input_size (tuple): RTMPose model Input image size.
|
||||||
|
center (tuple): Center of bbox in shape (x, y).
|
||||||
|
scale (tuple): Scale of bbox in shape (w, h).
|
||||||
|
simcc_split_ratio (float): Split ratio of simcc.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple:
|
||||||
|
- keypoints (np.ndarray): Rescaled keypoints.
|
||||||
|
- scores (np.ndarray): Model predict scores.
|
||||||
|
"""
|
||||||
|
all_key = []
|
||||||
|
all_score = []
|
||||||
|
for i in range(len(outputs)):
|
||||||
|
# use simcc to decode
|
||||||
|
simcc_x, simcc_y = outputs[i]
|
||||||
|
keypoints, scores = decode(simcc_x, simcc_y, simcc_split_ratio)
|
||||||
|
|
||||||
|
# rescale keypoints
|
||||||
|
keypoints = keypoints / model_input_size * scale[i] + center[i] - scale[i] / 2
|
||||||
|
all_key.append(keypoints[0])
|
||||||
|
all_score.append(scores[0])
|
||||||
|
|
||||||
|
return np.array(all_key), np.array(all_score)
|
||||||
|
|
||||||
|
|
||||||
|
def bbox_xyxy2cs(bbox: np.ndarray,
|
||||||
|
padding: float = 1.) -> Tuple[np.ndarray, np.ndarray]:
|
||||||
|
"""Transform the bbox format from (x,y,w,h) into (center, scale)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
bbox (ndarray): Bounding box(es) in shape (4,) or (n, 4), formatted
|
||||||
|
as (left, top, right, bottom)
|
||||||
|
padding (float): BBox padding factor that will be multilied to scale.
|
||||||
|
Default: 1.0
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple: A tuple containing center and scale.
|
||||||
|
- np.ndarray[float32]: Center (x, y) of the bbox in shape (2,) or
|
||||||
|
(n, 2)
|
||||||
|
- np.ndarray[float32]: Scale (w, h) of the bbox in shape (2,) or
|
||||||
|
(n, 2)
|
||||||
|
"""
|
||||||
|
# convert single bbox from (4, ) to (1, 4)
|
||||||
|
dim = bbox.ndim
|
||||||
|
if dim == 1:
|
||||||
|
bbox = bbox[None, :]
|
||||||
|
|
||||||
|
# get bbox center and scale
|
||||||
|
x1, y1, x2, y2 = np.hsplit(bbox, [1, 2, 3])
|
||||||
|
center = np.hstack([x1 + x2, y1 + y2]) * 0.5
|
||||||
|
scale = np.hstack([x2 - x1, y2 - y1]) * padding
|
||||||
|
|
||||||
|
if dim == 1:
|
||||||
|
center = center[0]
|
||||||
|
scale = scale[0]
|
||||||
|
|
||||||
|
return center, scale
|
||||||
|
|
||||||
|
|
||||||
|
def _fix_aspect_ratio(bbox_scale: np.ndarray,
|
||||||
|
aspect_ratio: float) -> np.ndarray:
|
||||||
|
"""Extend the scale to match the given aspect ratio.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
scale (np.ndarray): The image scale (w, h) in shape (2, )
|
||||||
|
aspect_ratio (float): The ratio of ``w/h``
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
np.ndarray: The reshaped image scale in (2, )
|
||||||
|
"""
|
||||||
|
w, h = np.hsplit(bbox_scale, [1])
|
||||||
|
bbox_scale = np.where(w > h * aspect_ratio,
|
||||||
|
np.hstack([w, w / aspect_ratio]),
|
||||||
|
np.hstack([h * aspect_ratio, h]))
|
||||||
|
return bbox_scale
|
||||||
|
|
||||||
|
|
||||||
|
def _rotate_point(pt: np.ndarray, angle_rad: float) -> np.ndarray:
|
||||||
|
"""Rotate a point by an angle.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
pt (np.ndarray): 2D point coordinates (x, y) in shape (2, )
|
||||||
|
angle_rad (float): rotation angle in radian
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
np.ndarray: Rotated point in shape (2, )
|
||||||
|
"""
|
||||||
|
sn, cs = np.sin(angle_rad), np.cos(angle_rad)
|
||||||
|
rot_mat = np.array([[cs, -sn], [sn, cs]])
|
||||||
|
return rot_mat @ pt
|
||||||
|
|
||||||
|
|
||||||
|
def _get_3rd_point(a: np.ndarray, b: np.ndarray) -> np.ndarray:
|
||||||
|
"""To calculate the affine matrix, three pairs of points are required. This
|
||||||
|
function is used to get the 3rd point, given 2D points a & b.
|
||||||
|
|
||||||
|
The 3rd point is defined by rotating vector `a - b` by 90 degrees
|
||||||
|
anticlockwise, using b as the rotation center.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
a (np.ndarray): The 1st point (x,y) in shape (2, )
|
||||||
|
b (np.ndarray): The 2nd point (x,y) in shape (2, )
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
np.ndarray: The 3rd point.
|
||||||
|
"""
|
||||||
|
direction = a - b
|
||||||
|
c = b + np.r_[-direction[1], direction[0]]
|
||||||
|
return c
|
||||||
|
|
||||||
|
|
||||||
|
def get_warp_matrix(center: np.ndarray,
|
||||||
|
scale: np.ndarray,
|
||||||
|
rot: float,
|
||||||
|
output_size: Tuple[int, int],
|
||||||
|
shift: Tuple[float, float] = (0., 0.),
|
||||||
|
inv: bool = False) -> np.ndarray:
|
||||||
|
"""Calculate the affine transformation matrix that can warp the bbox area
|
||||||
|
in the input image to the output size.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
center (np.ndarray[2, ]): Center of the bounding box (x, y).
|
||||||
|
scale (np.ndarray[2, ]): Scale of the bounding box
|
||||||
|
wrt [width, height].
|
||||||
|
rot (float): Rotation angle (degree).
|
||||||
|
output_size (np.ndarray[2, ] | list(2,)): Size of the
|
||||||
|
destination heatmaps.
|
||||||
|
shift (0-100%): Shift translation ratio wrt the width/height.
|
||||||
|
Default (0., 0.).
|
||||||
|
inv (bool): Option to inverse the affine transform direction.
|
||||||
|
(inv=False: src->dst or inv=True: dst->src)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
np.ndarray: A 2x3 transformation matrix
|
||||||
|
"""
|
||||||
|
shift = np.array(shift)
|
||||||
|
src_w = scale[0]
|
||||||
|
dst_w = output_size[0]
|
||||||
|
dst_h = output_size[1]
|
||||||
|
|
||||||
|
# compute transformation matrix
|
||||||
|
rot_rad = np.deg2rad(rot)
|
||||||
|
src_dir = _rotate_point(np.array([0., src_w * -0.5]), rot_rad)
|
||||||
|
dst_dir = np.array([0., dst_w * -0.5])
|
||||||
|
|
||||||
|
# get four corners of the src rectangle in the original image
|
||||||
|
src = np.zeros((3, 2), dtype=np.float32)
|
||||||
|
src[0, :] = center + scale * shift
|
||||||
|
src[1, :] = center + src_dir + scale * shift
|
||||||
|
src[2, :] = _get_3rd_point(src[0, :], src[1, :])
|
||||||
|
|
||||||
|
# get four corners of the dst rectangle in the input image
|
||||||
|
dst = np.zeros((3, 2), dtype=np.float32)
|
||||||
|
dst[0, :] = [dst_w * 0.5, dst_h * 0.5]
|
||||||
|
dst[1, :] = np.array([dst_w * 0.5, dst_h * 0.5]) + dst_dir
|
||||||
|
dst[2, :] = _get_3rd_point(dst[0, :], dst[1, :])
|
||||||
|
|
||||||
|
if inv:
|
||||||
|
warp_mat = cv2.getAffineTransform(np.float32(dst), np.float32(src))
|
||||||
|
else:
|
||||||
|
warp_mat = cv2.getAffineTransform(np.float32(src), np.float32(dst))
|
||||||
|
|
||||||
|
return warp_mat
|
||||||
|
|
||||||
|
|
||||||
|
def top_down_affine(input_size: dict, bbox_scale: dict, bbox_center: dict,
|
||||||
|
img: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
|
||||||
|
"""Get the bbox image as the model input by affine transform.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
input_size (dict): The input size of the model.
|
||||||
|
bbox_scale (dict): The bbox scale of the img.
|
||||||
|
bbox_center (dict): The bbox center of the img.
|
||||||
|
img (np.ndarray): The original image.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple: A tuple containing center and scale.
|
||||||
|
- np.ndarray[float32]: img after affine transform.
|
||||||
|
- np.ndarray[float32]: bbox scale after affine transform.
|
||||||
|
"""
|
||||||
|
w, h = input_size
|
||||||
|
warp_size = (int(w), int(h))
|
||||||
|
|
||||||
|
# reshape bbox to fixed aspect ratio
|
||||||
|
bbox_scale = _fix_aspect_ratio(bbox_scale, aspect_ratio=w / h)
|
||||||
|
|
||||||
|
# get the affine matrix
|
||||||
|
center = bbox_center
|
||||||
|
scale = bbox_scale
|
||||||
|
rot = 0
|
||||||
|
warp_mat = get_warp_matrix(center, scale, rot, output_size=(w, h))
|
||||||
|
|
||||||
|
# do affine transform
|
||||||
|
img = cv2.warpAffine(img, warp_mat, warp_size, flags=cv2.INTER_LINEAR)
|
||||||
|
|
||||||
|
return img, bbox_scale
|
||||||
|
|
||||||
|
|
||||||
|
def get_simcc_maximum(simcc_x: np.ndarray,
|
||||||
|
simcc_y: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
|
||||||
|
"""Get maximum response location and value from simcc representations.
|
||||||
|
|
||||||
|
Note:
|
||||||
|
instance number: N
|
||||||
|
num_keypoints: K
|
||||||
|
heatmap height: H
|
||||||
|
heatmap width: W
|
||||||
|
|
||||||
|
Args:
|
||||||
|
simcc_x (np.ndarray): x-axis SimCC in shape (K, Wx) or (N, K, Wx)
|
||||||
|
simcc_y (np.ndarray): y-axis SimCC in shape (K, Wy) or (N, K, Wy)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple:
|
||||||
|
- locs (np.ndarray): locations of maximum heatmap responses in shape
|
||||||
|
(K, 2) or (N, K, 2)
|
||||||
|
- vals (np.ndarray): values of maximum heatmap responses in shape
|
||||||
|
(K,) or (N, K)
|
||||||
|
"""
|
||||||
|
N, K, Wx = simcc_x.shape
|
||||||
|
simcc_x = simcc_x.reshape(N * K, -1)
|
||||||
|
simcc_y = simcc_y.reshape(N * K, -1)
|
||||||
|
|
||||||
|
# get maximum value locations
|
||||||
|
x_locs = np.argmax(simcc_x, axis=1)
|
||||||
|
y_locs = np.argmax(simcc_y, axis=1)
|
||||||
|
locs = np.stack((x_locs, y_locs), axis=-1).astype(np.float32)
|
||||||
|
max_val_x = np.amax(simcc_x, axis=1)
|
||||||
|
max_val_y = np.amax(simcc_y, axis=1)
|
||||||
|
|
||||||
|
# get maximum value across x and y axis
|
||||||
|
mask = max_val_x > max_val_y
|
||||||
|
max_val_x[mask] = max_val_y[mask]
|
||||||
|
vals = max_val_x
|
||||||
|
locs[vals <= 0.] = -1
|
||||||
|
|
||||||
|
# reshape
|
||||||
|
locs = locs.reshape(N, K, 2)
|
||||||
|
vals = vals.reshape(N, K)
|
||||||
|
|
||||||
|
return locs, vals
|
||||||
|
|
||||||
|
|
||||||
|
def decode(simcc_x: np.ndarray, simcc_y: np.ndarray,
|
||||||
|
simcc_split_ratio) -> Tuple[np.ndarray, np.ndarray]:
|
||||||
|
"""Modulate simcc distribution with Gaussian.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
simcc_x (np.ndarray[K, Wx]): model predicted simcc in x.
|
||||||
|
simcc_y (np.ndarray[K, Wy]): model predicted simcc in y.
|
||||||
|
simcc_split_ratio (int): The split ratio of simcc.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple: A tuple containing center and scale.
|
||||||
|
- np.ndarray[float32]: keypoints in shape (K, 2) or (n, K, 2)
|
||||||
|
- np.ndarray[float32]: scores in shape (K,) or (n, K)
|
||||||
|
"""
|
||||||
|
keypoints, scores = get_simcc_maximum(simcc_x, simcc_y)
|
||||||
|
keypoints /= simcc_split_ratio
|
||||||
|
|
||||||
|
return keypoints, scores
|
||||||
|
|
||||||
|
|
||||||
|
def inference_pose(session, out_bbox, oriImg):
|
||||||
|
h, w = session.get_inputs()[0].shape[2:]
|
||||||
|
model_input_size = (w, h)
|
||||||
|
resized_img, center, scale = preprocess(oriImg, out_bbox, model_input_size)
|
||||||
|
outputs = inference(session, resized_img)
|
||||||
|
keypoints, scores = postprocess(outputs, model_input_size, center, scale)
|
||||||
|
|
||||||
|
return keypoints, scores
|
||||||
@@ -0,0 +1,299 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
|
import math
|
||||||
|
import numpy as np
|
||||||
|
import matplotlib
|
||||||
|
import cv2
|
||||||
|
|
||||||
|
|
||||||
|
eps = 0.01
|
||||||
|
|
||||||
|
|
||||||
|
def smart_resize(x, s):
|
||||||
|
Ht, Wt = s
|
||||||
|
if x.ndim == 2:
|
||||||
|
Ho, Wo = x.shape
|
||||||
|
Co = 1
|
||||||
|
else:
|
||||||
|
Ho, Wo, Co = x.shape
|
||||||
|
if Co == 3 or Co == 1:
|
||||||
|
k = float(Ht + Wt) / float(Ho + Wo)
|
||||||
|
return cv2.resize(x, (int(Wt), int(Ht)), interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4)
|
||||||
|
else:
|
||||||
|
return np.stack([smart_resize(x[:, :, i], s) for i in range(Co)], axis=2)
|
||||||
|
|
||||||
|
|
||||||
|
def smart_resize_k(x, fx, fy):
|
||||||
|
if x.ndim == 2:
|
||||||
|
Ho, Wo = x.shape
|
||||||
|
Co = 1
|
||||||
|
else:
|
||||||
|
Ho, Wo, Co = x.shape
|
||||||
|
Ht, Wt = Ho * fy, Wo * fx
|
||||||
|
if Co == 3 or Co == 1:
|
||||||
|
k = float(Ht + Wt) / float(Ho + Wo)
|
||||||
|
return cv2.resize(x, (int(Wt), int(Ht)), interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4)
|
||||||
|
else:
|
||||||
|
return np.stack([smart_resize_k(x[:, :, i], fx, fy) for i in range(Co)], axis=2)
|
||||||
|
|
||||||
|
|
||||||
|
def padRightDownCorner(img, stride, padValue):
|
||||||
|
h = img.shape[0]
|
||||||
|
w = img.shape[1]
|
||||||
|
|
||||||
|
pad = 4 * [None]
|
||||||
|
pad[0] = 0 # up
|
||||||
|
pad[1] = 0 # left
|
||||||
|
pad[2] = 0 if (h % stride == 0) else stride - (h % stride) # down
|
||||||
|
pad[3] = 0 if (w % stride == 0) else stride - (w % stride) # right
|
||||||
|
|
||||||
|
img_padded = img
|
||||||
|
pad_up = np.tile(img_padded[0:1, :, :]*0 + padValue, (pad[0], 1, 1))
|
||||||
|
img_padded = np.concatenate((pad_up, img_padded), axis=0)
|
||||||
|
pad_left = np.tile(img_padded[:, 0:1, :]*0 + padValue, (1, pad[1], 1))
|
||||||
|
img_padded = np.concatenate((pad_left, img_padded), axis=1)
|
||||||
|
pad_down = np.tile(img_padded[-2:-1, :, :]*0 + padValue, (pad[2], 1, 1))
|
||||||
|
img_padded = np.concatenate((img_padded, pad_down), axis=0)
|
||||||
|
pad_right = np.tile(img_padded[:, -2:-1, :]*0 + padValue, (1, pad[3], 1))
|
||||||
|
img_padded = np.concatenate((img_padded, pad_right), axis=1)
|
||||||
|
|
||||||
|
return img_padded, pad
|
||||||
|
|
||||||
|
|
||||||
|
def transfer(model, model_weights):
|
||||||
|
transfered_model_weights = {}
|
||||||
|
for weights_name in model.state_dict().keys():
|
||||||
|
transfered_model_weights[weights_name] = model_weights['.'.join(weights_name.split('.')[1:])]
|
||||||
|
return transfered_model_weights
|
||||||
|
|
||||||
|
|
||||||
|
def draw_bodypose(canvas, candidate, subset):
|
||||||
|
H, W, C = canvas.shape
|
||||||
|
candidate = np.array(candidate)
|
||||||
|
subset = np.array(subset)
|
||||||
|
|
||||||
|
stickwidth = 4
|
||||||
|
|
||||||
|
limbSeq = [[2, 3], [2, 6], [3, 4], [4, 5], [6, 7], [7, 8], [2, 9], [9, 10], \
|
||||||
|
[10, 11], [2, 12], [12, 13], [13, 14], [2, 1], [1, 15], [15, 17], \
|
||||||
|
[1, 16], [16, 18], [3, 17], [6, 18]]
|
||||||
|
|
||||||
|
colors = [[255, 0, 0], [255, 85, 0], [255, 170, 0], [255, 255, 0], [170, 255, 0], [85, 255, 0], [0, 255, 0], \
|
||||||
|
[0, 255, 85], [0, 255, 170], [0, 255, 255], [0, 170, 255], [0, 85, 255], [0, 0, 255], [85, 0, 255], \
|
||||||
|
[170, 0, 255], [255, 0, 255], [255, 0, 170], [255, 0, 85]]
|
||||||
|
|
||||||
|
for i in range(17):
|
||||||
|
for n in range(len(subset)):
|
||||||
|
index = subset[n][np.array(limbSeq[i]) - 1]
|
||||||
|
if -1 in index:
|
||||||
|
continue
|
||||||
|
Y = candidate[index.astype(int), 0] * float(W)
|
||||||
|
X = candidate[index.astype(int), 1] * float(H)
|
||||||
|
mX = np.mean(X)
|
||||||
|
mY = np.mean(Y)
|
||||||
|
length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5
|
||||||
|
angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1]))
|
||||||
|
polygon = cv2.ellipse2Poly((int(mY), int(mX)), (int(length / 2), stickwidth), int(angle), 0, 360, 1)
|
||||||
|
cv2.fillConvexPoly(canvas, polygon, colors[i])
|
||||||
|
|
||||||
|
canvas = (canvas * 0.6).astype(np.uint8)
|
||||||
|
|
||||||
|
for i in range(18):
|
||||||
|
for n in range(len(subset)):
|
||||||
|
index = int(subset[n][i])
|
||||||
|
if index == -1:
|
||||||
|
continue
|
||||||
|
x, y = candidate[index][0:2]
|
||||||
|
x = int(x * W)
|
||||||
|
y = int(y * H)
|
||||||
|
cv2.circle(canvas, (int(x), int(y)), 4, colors[i], thickness=-1)
|
||||||
|
|
||||||
|
return canvas
|
||||||
|
|
||||||
|
|
||||||
|
def draw_handpose(canvas, all_hand_peaks):
|
||||||
|
H, W, C = canvas.shape
|
||||||
|
|
||||||
|
edges = [[0, 1], [1, 2], [2, 3], [3, 4], [0, 5], [5, 6], [6, 7], [7, 8], [0, 9], [9, 10], \
|
||||||
|
[10, 11], [11, 12], [0, 13], [13, 14], [14, 15], [15, 16], [0, 17], [17, 18], [18, 19], [19, 20]]
|
||||||
|
|
||||||
|
for peaks in all_hand_peaks:
|
||||||
|
peaks = np.array(peaks)
|
||||||
|
|
||||||
|
for ie, e in enumerate(edges):
|
||||||
|
x1, y1 = peaks[e[0]]
|
||||||
|
x2, y2 = peaks[e[1]]
|
||||||
|
x1 = int(x1 * W)
|
||||||
|
y1 = int(y1 * H)
|
||||||
|
x2 = int(x2 * W)
|
||||||
|
y2 = int(y2 * H)
|
||||||
|
if x1 > eps and y1 > eps and x2 > eps and y2 > eps:
|
||||||
|
cv2.line(canvas, (x1, y1), (x2, y2), matplotlib.colors.hsv_to_rgb([ie / float(len(edges)), 1.0, 1.0]) * 255, thickness=2)
|
||||||
|
|
||||||
|
for i, keyponit in enumerate(peaks):
|
||||||
|
x, y = keyponit
|
||||||
|
x = int(x * W)
|
||||||
|
y = int(y * H)
|
||||||
|
if x > eps and y > eps:
|
||||||
|
cv2.circle(canvas, (x, y), 4, (0, 0, 255), thickness=-1)
|
||||||
|
return canvas
|
||||||
|
|
||||||
|
|
||||||
|
def draw_facepose(canvas, all_lmks):
|
||||||
|
H, W, C = canvas.shape
|
||||||
|
for lmks in all_lmks:
|
||||||
|
lmks = np.array(lmks)
|
||||||
|
for lmk in lmks:
|
||||||
|
x, y = lmk
|
||||||
|
x = int(x * W)
|
||||||
|
y = int(y * H)
|
||||||
|
if x > eps and y > eps:
|
||||||
|
cv2.circle(canvas, (x, y), 3, (255, 255, 255), thickness=-1)
|
||||||
|
return canvas
|
||||||
|
|
||||||
|
|
||||||
|
# detect hand according to body pose keypoints
|
||||||
|
# please refer to https://github.com/CMU-Perceptual-Computing-Lab/openpose/blob/master/src/openpose/hand/handDetector.cpp
|
||||||
|
def handDetect(candidate, subset, oriImg):
|
||||||
|
# right hand: wrist 4, elbow 3, shoulder 2
|
||||||
|
# left hand: wrist 7, elbow 6, shoulder 5
|
||||||
|
ratioWristElbow = 0.33
|
||||||
|
detect_result = []
|
||||||
|
image_height, image_width = oriImg.shape[0:2]
|
||||||
|
for person in subset.astype(int):
|
||||||
|
# if any of three not detected
|
||||||
|
has_left = np.sum(person[[5, 6, 7]] == -1) == 0
|
||||||
|
has_right = np.sum(person[[2, 3, 4]] == -1) == 0
|
||||||
|
if not (has_left or has_right):
|
||||||
|
continue
|
||||||
|
hands = []
|
||||||
|
#left hand
|
||||||
|
if has_left:
|
||||||
|
left_shoulder_index, left_elbow_index, left_wrist_index = person[[5, 6, 7]]
|
||||||
|
x1, y1 = candidate[left_shoulder_index][:2]
|
||||||
|
x2, y2 = candidate[left_elbow_index][:2]
|
||||||
|
x3, y3 = candidate[left_wrist_index][:2]
|
||||||
|
hands.append([x1, y1, x2, y2, x3, y3, True])
|
||||||
|
# right hand
|
||||||
|
if has_right:
|
||||||
|
right_shoulder_index, right_elbow_index, right_wrist_index = person[[2, 3, 4]]
|
||||||
|
x1, y1 = candidate[right_shoulder_index][:2]
|
||||||
|
x2, y2 = candidate[right_elbow_index][:2]
|
||||||
|
x3, y3 = candidate[right_wrist_index][:2]
|
||||||
|
hands.append([x1, y1, x2, y2, x3, y3, False])
|
||||||
|
|
||||||
|
for x1, y1, x2, y2, x3, y3, is_left in hands:
|
||||||
|
# pos_hand = pos_wrist + ratio * (pos_wrist - pos_elbox) = (1 + ratio) * pos_wrist - ratio * pos_elbox
|
||||||
|
# handRectangle.x = posePtr[wrist*3] + ratioWristElbow * (posePtr[wrist*3] - posePtr[elbow*3]);
|
||||||
|
# handRectangle.y = posePtr[wrist*3+1] + ratioWristElbow * (posePtr[wrist*3+1] - posePtr[elbow*3+1]);
|
||||||
|
# const auto distanceWristElbow = getDistance(poseKeypoints, person, wrist, elbow);
|
||||||
|
# const auto distanceElbowShoulder = getDistance(poseKeypoints, person, elbow, shoulder);
|
||||||
|
# handRectangle.width = 1.5f * fastMax(distanceWristElbow, 0.9f * distanceElbowShoulder);
|
||||||
|
x = x3 + ratioWristElbow * (x3 - x2)
|
||||||
|
y = y3 + ratioWristElbow * (y3 - y2)
|
||||||
|
distanceWristElbow = math.sqrt((x3 - x2) ** 2 + (y3 - y2) ** 2)
|
||||||
|
distanceElbowShoulder = math.sqrt((x2 - x1) ** 2 + (y2 - y1) ** 2)
|
||||||
|
width = 1.5 * max(distanceWristElbow, 0.9 * distanceElbowShoulder)
|
||||||
|
# x-y refers to the center --> offset to topLeft point
|
||||||
|
# handRectangle.x -= handRectangle.width / 2.f;
|
||||||
|
# handRectangle.y -= handRectangle.height / 2.f;
|
||||||
|
x -= width / 2
|
||||||
|
y -= width / 2 # width = height
|
||||||
|
# overflow the image
|
||||||
|
if x < 0: x = 0
|
||||||
|
if y < 0: y = 0
|
||||||
|
width1 = width
|
||||||
|
width2 = width
|
||||||
|
if x + width > image_width: width1 = image_width - x
|
||||||
|
if y + width > image_height: width2 = image_height - y
|
||||||
|
width = min(width1, width2)
|
||||||
|
# the max hand box value is 20 pixels
|
||||||
|
if width >= 20:
|
||||||
|
detect_result.append([int(x), int(y), int(width), is_left])
|
||||||
|
|
||||||
|
'''
|
||||||
|
return value: [[x, y, w, True if left hand else False]].
|
||||||
|
width=height since the network require squared input.
|
||||||
|
x, y is the coordinate of top left
|
||||||
|
'''
|
||||||
|
return detect_result
|
||||||
|
|
||||||
|
|
||||||
|
# Written by Lvmin
|
||||||
|
def faceDetect(candidate, subset, oriImg):
|
||||||
|
# left right eye ear 14 15 16 17
|
||||||
|
detect_result = []
|
||||||
|
image_height, image_width = oriImg.shape[0:2]
|
||||||
|
for person in subset.astype(int):
|
||||||
|
has_head = person[0] > -1
|
||||||
|
if not has_head:
|
||||||
|
continue
|
||||||
|
|
||||||
|
has_left_eye = person[14] > -1
|
||||||
|
has_right_eye = person[15] > -1
|
||||||
|
has_left_ear = person[16] > -1
|
||||||
|
has_right_ear = person[17] > -1
|
||||||
|
|
||||||
|
if not (has_left_eye or has_right_eye or has_left_ear or has_right_ear):
|
||||||
|
continue
|
||||||
|
|
||||||
|
head, left_eye, right_eye, left_ear, right_ear = person[[0, 14, 15, 16, 17]]
|
||||||
|
|
||||||
|
width = 0.0
|
||||||
|
x0, y0 = candidate[head][:2]
|
||||||
|
|
||||||
|
if has_left_eye:
|
||||||
|
x1, y1 = candidate[left_eye][:2]
|
||||||
|
d = max(abs(x0 - x1), abs(y0 - y1))
|
||||||
|
width = max(width, d * 3.0)
|
||||||
|
|
||||||
|
if has_right_eye:
|
||||||
|
x1, y1 = candidate[right_eye][:2]
|
||||||
|
d = max(abs(x0 - x1), abs(y0 - y1))
|
||||||
|
width = max(width, d * 3.0)
|
||||||
|
|
||||||
|
if has_left_ear:
|
||||||
|
x1, y1 = candidate[left_ear][:2]
|
||||||
|
d = max(abs(x0 - x1), abs(y0 - y1))
|
||||||
|
width = max(width, d * 1.5)
|
||||||
|
|
||||||
|
if has_right_ear:
|
||||||
|
x1, y1 = candidate[right_ear][:2]
|
||||||
|
d = max(abs(x0 - x1), abs(y0 - y1))
|
||||||
|
width = max(width, d * 1.5)
|
||||||
|
|
||||||
|
x, y = x0, y0
|
||||||
|
|
||||||
|
x -= width
|
||||||
|
y -= width
|
||||||
|
|
||||||
|
if x < 0:
|
||||||
|
x = 0
|
||||||
|
|
||||||
|
if y < 0:
|
||||||
|
y = 0
|
||||||
|
|
||||||
|
width1 = width * 2
|
||||||
|
width2 = width * 2
|
||||||
|
|
||||||
|
if x + width > image_width:
|
||||||
|
width1 = image_width - x
|
||||||
|
|
||||||
|
if y + width > image_height:
|
||||||
|
width2 = image_height - y
|
||||||
|
|
||||||
|
width = min(width1, width2)
|
||||||
|
|
||||||
|
if width >= 20:
|
||||||
|
detect_result.append([int(x), int(y), int(width)])
|
||||||
|
|
||||||
|
return detect_result
|
||||||
|
|
||||||
|
|
||||||
|
# get max index of 2d array
|
||||||
|
def npmax(array):
|
||||||
|
arrayindex = array.argmax(1)
|
||||||
|
arrayvalue = array.max(1)
|
||||||
|
i = arrayvalue.argmax()
|
||||||
|
j = arrayindex[i]
|
||||||
|
return i, j
|
||||||
@@ -0,0 +1,80 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
|
import cv2
|
||||||
|
import numpy as np
|
||||||
|
import onnxruntime as ort
|
||||||
|
from .onnxdet import inference_detector
|
||||||
|
from .onnxpose import inference_pose
|
||||||
|
|
||||||
|
def HWC3(x):
|
||||||
|
assert x.dtype == np.uint8
|
||||||
|
if x.ndim == 2:
|
||||||
|
x = x[:, :, None]
|
||||||
|
assert x.ndim == 3
|
||||||
|
H, W, C = x.shape
|
||||||
|
assert C == 1 or C == 3 or C == 4
|
||||||
|
if C == 3:
|
||||||
|
return x
|
||||||
|
if C == 1:
|
||||||
|
return np.concatenate([x, x, x], axis=2)
|
||||||
|
if C == 4:
|
||||||
|
color = x[:, :, 0:3].astype(np.float32)
|
||||||
|
alpha = x[:, :, 3:4].astype(np.float32) / 255.0
|
||||||
|
y = color * alpha + 255.0 * (1.0 - alpha)
|
||||||
|
y = y.clip(0, 255).astype(np.uint8)
|
||||||
|
return y
|
||||||
|
|
||||||
|
|
||||||
|
def resize_image(input_image, resolution):
|
||||||
|
H, W, C = input_image.shape
|
||||||
|
H = float(H)
|
||||||
|
W = float(W)
|
||||||
|
k = float(resolution) / min(H, W)
|
||||||
|
H *= k
|
||||||
|
W *= k
|
||||||
|
H = int(np.round(H / 64.0)) * 64
|
||||||
|
W = int(np.round(W / 64.0)) * 64
|
||||||
|
img = cv2.resize(input_image, (W, H), interpolation=cv2.INTER_LANCZOS4 if k > 1 else cv2.INTER_AREA)
|
||||||
|
return img
|
||||||
|
|
||||||
|
class Wholebody:
|
||||||
|
def __init__(self, onnx_det, onnx_pose, device = 'cuda:0'):
|
||||||
|
|
||||||
|
providers = ['CPUExecutionProvider'
|
||||||
|
] if device == 'cpu' else ['CUDAExecutionProvider']
|
||||||
|
# onnx_det = 'annotator/ckpts/yolox_l.onnx'
|
||||||
|
# onnx_pose = 'annotator/ckpts/dw-ll_ucoco_384.onnx'
|
||||||
|
|
||||||
|
self.session_det = ort.InferenceSession(path_or_bytes=onnx_det, providers=providers)
|
||||||
|
self.session_pose = ort.InferenceSession(path_or_bytes=onnx_pose, providers=providers)
|
||||||
|
|
||||||
|
def __call__(self, ori_img):
|
||||||
|
det_result = inference_detector(self.session_det, ori_img)
|
||||||
|
keypoints, scores = inference_pose(self.session_pose, det_result, ori_img)
|
||||||
|
|
||||||
|
keypoints_info = np.concatenate(
|
||||||
|
(keypoints, scores[..., None]), axis=-1)
|
||||||
|
# compute neck joint
|
||||||
|
neck = np.mean(keypoints_info[:, [5, 6]], axis=1)
|
||||||
|
# neck score when visualizing pred
|
||||||
|
neck[:, 2:4] = np.logical_and(
|
||||||
|
keypoints_info[:, 5, 2:4] > 0.3,
|
||||||
|
keypoints_info[:, 6, 2:4] > 0.3).astype(int)
|
||||||
|
new_keypoints_info = np.insert(
|
||||||
|
keypoints_info, 17, neck, axis=1)
|
||||||
|
mmpose_idx = [
|
||||||
|
17, 6, 8, 10, 7, 9, 12, 14, 16, 13, 15, 2, 1, 4, 3
|
||||||
|
]
|
||||||
|
openpose_idx = [
|
||||||
|
1, 2, 3, 4, 6, 7, 8, 9, 10, 12, 13, 14, 15, 16, 17
|
||||||
|
]
|
||||||
|
new_keypoints_info[:, openpose_idx] = \
|
||||||
|
new_keypoints_info[:, mmpose_idx]
|
||||||
|
keypoints_info = new_keypoints_info
|
||||||
|
|
||||||
|
keypoints, scores = keypoints_info[
|
||||||
|
..., :2], keypoints_info[..., 2]
|
||||||
|
|
||||||
|
return keypoints, scores, det_result
|
||||||
|
|
||||||
|
|
||||||
@@ -0,0 +1,203 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
|
# Openpose
|
||||||
|
# Original from CMU https://github.com/CMU-Perceptual-Computing-Lab/openpose
|
||||||
|
# 2nd Edited by https://github.com/Hzzone/pytorch-openpose
|
||||||
|
# 3rd Edited by ControlNet
|
||||||
|
# 4th Edited by ControlNet (added face and correct hands)
|
||||||
|
|
||||||
|
# ``` requirements for cuda 12.1:
|
||||||
|
# onnxruntime==1.19
|
||||||
|
# onnxruntime-gpu==1.19
|
||||||
|
# ```
|
||||||
|
|
||||||
|
import os
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
|
import cv2
|
||||||
|
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||||
|
from scepter.modules.annotator.dwpose import util
|
||||||
|
from scepter.modules.annotator.dwpose.wholebody import (HWC3, Wholebody,
|
||||||
|
resize_image)
|
||||||
|
from scepter.modules.annotator.registry import ANNOTATORS
|
||||||
|
from scepter.modules.utils.distribute import we
|
||||||
|
from scepter.modules.utils.file_system import FS
|
||||||
|
|
||||||
|
os.environ['KMP_DUPLICATE_LIB_OK'] = 'TRUE'
|
||||||
|
|
||||||
|
|
||||||
|
def draw_pose(pose, H, W, use_hand=False, use_body=False, use_face=False):
|
||||||
|
bodies = pose['bodies']
|
||||||
|
faces = pose['faces']
|
||||||
|
hands = pose['hands']
|
||||||
|
candidate = bodies['candidate']
|
||||||
|
subset = bodies['subset']
|
||||||
|
canvas = np.zeros(shape=(H, W, 3), dtype=np.uint8)
|
||||||
|
|
||||||
|
if use_body:
|
||||||
|
canvas = util.draw_bodypose(canvas, candidate, subset)
|
||||||
|
if use_hand:
|
||||||
|
canvas = util.draw_handpose(canvas, hands)
|
||||||
|
if use_face:
|
||||||
|
canvas = util.draw_facepose(canvas, faces)
|
||||||
|
|
||||||
|
return canvas
|
||||||
|
|
||||||
|
|
||||||
|
@ANNOTATORS.register_class()
|
||||||
|
class DWposeAnnotator(BaseAnnotator):
|
||||||
|
def __init__(self, cfg, logger=None):
|
||||||
|
super().__init__(cfg, logger=logger)
|
||||||
|
with FS.get_from(cfg['DETECTION_MODEL'],
|
||||||
|
wait_finish=True) as onnx_det, FS.get_from(
|
||||||
|
cfg['POSE_MODEL'], wait_finish=True) as onnx_pose:
|
||||||
|
self.pose_estimation = Wholebody(onnx_det,
|
||||||
|
onnx_pose,
|
||||||
|
device=f'cuda:{we.device_id}')
|
||||||
|
self.resize_size = cfg.get('RESIZE_SIZE', 1024)
|
||||||
|
self.use_body = cfg.get('USE_BODY', True)
|
||||||
|
self.use_face = cfg.get('USE_FACE', True)
|
||||||
|
self.use_hand = cfg.get('USE_HAND', True)
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
@torch.inference_mode
|
||||||
|
def forward(self, image):
|
||||||
|
if isinstance(image, Image.Image):
|
||||||
|
image = np.array(image)
|
||||||
|
elif isinstance(image, torch.Tensor):
|
||||||
|
image = image.detach().cpu().numpy()
|
||||||
|
elif isinstance(image, np.ndarray):
|
||||||
|
image = image.copy()
|
||||||
|
else:
|
||||||
|
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||||
|
|
||||||
|
input_image = HWC3(image[..., ::-1])
|
||||||
|
return self.process(resize_image(input_image, self.resize_size),
|
||||||
|
image.shape[:2])
|
||||||
|
|
||||||
|
def process(self, ori_img, ori_shape):
|
||||||
|
ori_h, ori_w = ori_shape
|
||||||
|
ori_img = ori_img.copy()
|
||||||
|
H, W, C = ori_img.shape
|
||||||
|
with torch.no_grad():
|
||||||
|
candidate, subset, det_result = self.pose_estimation(ori_img)
|
||||||
|
nums, keys, locs = candidate.shape
|
||||||
|
candidate[..., 0] /= float(W)
|
||||||
|
candidate[..., 1] /= float(H)
|
||||||
|
body = candidate[:, :18].copy()
|
||||||
|
body = body.reshape(nums * 18, locs)
|
||||||
|
score = subset[:, :18]
|
||||||
|
for i in range(len(score)):
|
||||||
|
for j in range(len(score[i])):
|
||||||
|
if score[i][j] > 0.3:
|
||||||
|
score[i][j] = int(18 * i + j)
|
||||||
|
else:
|
||||||
|
score[i][j] = -1
|
||||||
|
|
||||||
|
un_visible = subset < 0.3
|
||||||
|
candidate[un_visible] = -1
|
||||||
|
|
||||||
|
foot = candidate[:, 18:24]
|
||||||
|
|
||||||
|
faces = candidate[:, 24:92]
|
||||||
|
|
||||||
|
hands = candidate[:, 92:113]
|
||||||
|
hands = np.vstack([hands, candidate[:, 113:]])
|
||||||
|
|
||||||
|
bodies = dict(candidate=body, subset=score)
|
||||||
|
pose = dict(bodies=bodies, hands=hands, faces=faces)
|
||||||
|
|
||||||
|
ret_data = {}
|
||||||
|
if self.use_body:
|
||||||
|
detected_map_body = draw_pose(pose, H, W, use_body=True)
|
||||||
|
detected_map_body = cv2.resize(
|
||||||
|
detected_map_body[..., ::-1], (ori_w, ori_h),
|
||||||
|
interpolation=cv2.INTER_LANCZOS4
|
||||||
|
if ori_h * ori_w > H * W else cv2.INTER_AREA)
|
||||||
|
ret_data['detected_map_body'] = detected_map_body
|
||||||
|
|
||||||
|
if self.use_face:
|
||||||
|
detected_map_face = draw_pose(pose, H, W, use_face=True)
|
||||||
|
detected_map_face = cv2.resize(
|
||||||
|
detected_map_face[..., ::-1], (ori_w, ori_h),
|
||||||
|
interpolation=cv2.INTER_LANCZOS4
|
||||||
|
if ori_h * ori_w > H * W else cv2.INTER_AREA)
|
||||||
|
ret_data['detected_map_face'] = detected_map_face
|
||||||
|
|
||||||
|
if self.use_body and self.use_face:
|
||||||
|
detected_map_bodyface = draw_pose(pose,
|
||||||
|
H,
|
||||||
|
W,
|
||||||
|
use_body=True,
|
||||||
|
use_face=True)
|
||||||
|
detected_map_bodyface = cv2.resize(
|
||||||
|
detected_map_bodyface[..., ::-1], (ori_w, ori_h),
|
||||||
|
interpolation=cv2.INTER_LANCZOS4
|
||||||
|
if ori_h * ori_w > H * W else cv2.INTER_AREA)
|
||||||
|
ret_data['detected_map_bodyface'] = detected_map_bodyface
|
||||||
|
|
||||||
|
if self.use_hand and self.use_body and self.use_face:
|
||||||
|
detected_map_handbodyface = draw_pose(pose,
|
||||||
|
H,
|
||||||
|
W,
|
||||||
|
use_hand=True,
|
||||||
|
use_body=True,
|
||||||
|
use_face=True)
|
||||||
|
detected_map_handbodyface = cv2.resize(
|
||||||
|
detected_map_handbodyface[..., ::-1], (ori_w, ori_h),
|
||||||
|
interpolation=cv2.INTER_LANCZOS4
|
||||||
|
if ori_h * ori_w > H * W else cv2.INTER_AREA)
|
||||||
|
ret_data[
|
||||||
|
'detected_map_handbodyface'] = detected_map_handbodyface
|
||||||
|
|
||||||
|
# convert_size
|
||||||
|
if det_result.shape[0] > 0:
|
||||||
|
w_ratio, h_ratio = ori_w / W, ori_h / H
|
||||||
|
det_result[..., ::2] *= h_ratio
|
||||||
|
det_result[..., 1::2] *= w_ratio
|
||||||
|
det_result = det_result.astype(np.int32)
|
||||||
|
# for det_tup in det_result:
|
||||||
|
# cv2.rectangle(detected_map, det_tup[2:].tolist(), det_tup[:2].tolist(), color=(255, 0, 0), thickness=3)
|
||||||
|
return ret_data, det_result
|
||||||
|
|
||||||
|
|
||||||
|
@ANNOTATORS.register_class()
|
||||||
|
class DWposeBodyAnnotator(DWposeAnnotator):
|
||||||
|
def __init__(self, cfg, logger=None):
|
||||||
|
super().__init__(cfg, logger=logger)
|
||||||
|
self.use_body, self.use_face, self.use_hand = True, False, False
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
@torch.inference_mode
|
||||||
|
def forward(self, image):
|
||||||
|
ret_data, det_result = super().forward(image)
|
||||||
|
return ret_data['detected_map_body']
|
||||||
|
|
||||||
|
|
||||||
|
@ANNOTATORS.register_class()
|
||||||
|
class DWposeFaceAnnotator(DWposeAnnotator):
|
||||||
|
def __init__(self, cfg, logger=None):
|
||||||
|
super().__init__(cfg, logger=logger)
|
||||||
|
self.use_body, self.use_face, self.use_hand = False, True, False
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
@torch.inference_mode
|
||||||
|
def forward(self, image):
|
||||||
|
ret_data, det_result = super().forward(image)
|
||||||
|
return ret_data['detected_map_face']
|
||||||
|
|
||||||
|
|
||||||
|
@ANNOTATORS.register_class()
|
||||||
|
class DWposeBodyFaceAnnotator(DWposeAnnotator):
|
||||||
|
def __init__(self, cfg, logger=None):
|
||||||
|
super().__init__(cfg, logger=logger)
|
||||||
|
self.use_body, self.use_face, self.use_hand = True, True, False
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
@torch.inference_mode
|
||||||
|
def forward(self, image):
|
||||||
|
ret_data, det_result = super().forward(image)
|
||||||
|
return ret_data['detected_map_bodyface']
|
||||||
@@ -0,0 +1,63 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
|
import os
|
||||||
|
from abc import ABCMeta
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
|
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||||
|
from scepter.modules.annotator.registry import ANNOTATORS
|
||||||
|
from scepter.modules.utils.distribute import we
|
||||||
|
from scepter.modules.utils.file_system import FS
|
||||||
|
|
||||||
|
|
||||||
|
@ANNOTATORS.register_class()
|
||||||
|
class FaceAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||||
|
para_dict = {}
|
||||||
|
|
||||||
|
def __init__(self, cfg, logger=None):
|
||||||
|
super().__init__(cfg, logger=logger)
|
||||||
|
from insightface.app import FaceAnalysis
|
||||||
|
local_path = FS.map_to_local(cfg.PRETRAINED_MODEL)[0]
|
||||||
|
local_model_path = os.path.join(local_path, 'models', cfg.MODEL_NAME)
|
||||||
|
FS.get_dir_to_local_dir(cfg.PRETRAINED_MODEL, local_model_path)
|
||||||
|
self.model = FaceAnalysis(name=cfg.MODEL_NAME, root=local_path, providers=['CUDAExecutionProvider', 'CPUExecutionProvider'])
|
||||||
|
self.model.prepare(ctx_id=we.device_id, det_size=(640, 640))
|
||||||
|
|
||||||
|
def forward(self, image=None):
|
||||||
|
|
||||||
|
if isinstance(image, Image.Image):
|
||||||
|
image = np.array(image)
|
||||||
|
elif isinstance(image, torch.Tensor):
|
||||||
|
image = image.detach().cpu().numpy()
|
||||||
|
elif isinstance(image, np.ndarray):
|
||||||
|
image = image.copy()
|
||||||
|
else:
|
||||||
|
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||||
|
|
||||||
|
# [dict_keys(['bbox', 'kps', 'det_score', 'landmark_3d_68', 'pose', 'landmark_2d_106', 'gender', 'age', 'embedding'])]
|
||||||
|
faces = self.model.get(image)
|
||||||
|
return faces
|
||||||
|
|
||||||
|
|
||||||
|
@ANNOTATORS.register_class()
|
||||||
|
class FaceMaskAnnotator(FaceAnnotator):
|
||||||
|
|
||||||
|
def __init__(self, cfg, logger=None):
|
||||||
|
super().__init__(cfg, logger=logger)
|
||||||
|
self.multi_face = cfg.get('MULTI_FACE', True)
|
||||||
|
|
||||||
|
def forward(self, image=None):
|
||||||
|
faces = super().forward(image)
|
||||||
|
if len(faces) > 0:
|
||||||
|
if not self.multi_face:
|
||||||
|
faces = faces[:1]
|
||||||
|
mask = np.zeros_like(image[:, :, 0])
|
||||||
|
for face in faces:
|
||||||
|
x_min, y_min, x_max, y_max = face['bbox'].tolist()
|
||||||
|
mask[int(y_min): int(y_max) + 1, int(x_min): int(x_max) + 1] = 255
|
||||||
|
return mask
|
||||||
|
else:
|
||||||
|
return np.zeros_like(image[:, :, 0])
|
||||||
@@ -0,0 +1,58 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
|
import random
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||||
|
from scepter.modules.annotator.registry import ANNOTATORS
|
||||||
|
from scepter.modules.utils.config import Config
|
||||||
|
|
||||||
|
|
||||||
|
@ANNOTATORS.register_class()
|
||||||
|
class FrameReferenceAnnotator(BaseAnnotator):
|
||||||
|
para_dict = {}
|
||||||
|
|
||||||
|
def __init__(self, cfg, logger=None):
|
||||||
|
super().__init__(cfg, logger=logger)
|
||||||
|
# first / last / firstlast / random
|
||||||
|
self.ref_cfg = cfg.get('REF_CFG', [{"mode": "first", "proba": 0.1},
|
||||||
|
{"mode": "last", "proba": 0.1},
|
||||||
|
{"mode": "firstlast", "proba": 0.1},
|
||||||
|
{"mode": "random", "proba": 0.1}])
|
||||||
|
self.ref_num = cfg.get('REF_NUM', 1)
|
||||||
|
self.ref_cfg = Config.get_dict(self.ref_cfg) if isinstance(
|
||||||
|
self.ref_cfg, Config) else self.ref_cfg
|
||||||
|
self.ref_color = cfg.get('REF_COLOR', 127.5)
|
||||||
|
|
||||||
|
def forward(self, frames, ref_cfg=None, ref_num=None):
|
||||||
|
ref_cfg = ref_cfg if ref_cfg is not None else self.ref_cfg
|
||||||
|
ref_cfg = [ref_cfg] if not isinstance(ref_cfg, list) else ref_cfg
|
||||||
|
probas = [item['proba'] if 'proba' in item else 1.0 / len(ref_cfg) for item in ref_cfg]
|
||||||
|
sel_ref_cfg = random.choices(ref_cfg, weights=probas, k=1)[0]
|
||||||
|
mode = sel_ref_cfg['mode'] if 'mode' in sel_ref_cfg else 'original'
|
||||||
|
ref_num = int(ref_num) if ref_num is not None else self.ref_num
|
||||||
|
|
||||||
|
frame_num = len(frames)
|
||||||
|
frame_num_range = list(range(frame_num))
|
||||||
|
if mode == "first":
|
||||||
|
sel_idx = frame_num_range[:ref_num]
|
||||||
|
elif mode == "last":
|
||||||
|
sel_idx = frame_num_range[-ref_num:]
|
||||||
|
elif mode == "firstlast":
|
||||||
|
sel_idx = frame_num_range[:ref_num] + frame_num_range[-ref_num:]
|
||||||
|
elif mode == "random":
|
||||||
|
sel_idx = random.sample(frame_num_range, ref_num)
|
||||||
|
else:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
out_frames, out_masks = [], []
|
||||||
|
for i in range(frame_num):
|
||||||
|
if i in sel_idx:
|
||||||
|
out_frame = frames[i]
|
||||||
|
out_mask = np.zeros_like(frames[i][:, :, 0])
|
||||||
|
else:
|
||||||
|
out_frame = np.ones_like(frames[i]) * self.ref_color
|
||||||
|
out_mask = np.ones_like(frames[i][:, :, 0]) * 255
|
||||||
|
out_frames.append(out_frame)
|
||||||
|
out_masks.append(out_mask)
|
||||||
|
return out_frames, out_masks
|
||||||
@@ -114,7 +114,7 @@ class HedAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
|||||||
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
|
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
|
||||||
if pretrained_model:
|
if pretrained_model:
|
||||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
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.no_grad()
|
||||||
@torch.inference_mode()
|
@torch.inference_mode()
|
||||||
|
|||||||
@@ -120,7 +120,7 @@ class InfoDrawContourAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
|||||||
self.model = ContourInference(input_nc, output_nc, n_residual_blocks,
|
self.model = ContourInference(input_nc, output_nc, n_residual_blocks,
|
||||||
sigmoid)
|
sigmoid)
|
||||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
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)
|
self.model = self.model.eval().requires_grad_(False).to(we.device_id)
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
|
|||||||
@@ -0,0 +1,450 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
|
import random
|
||||||
|
from abc import ABCMeta
|
||||||
|
from functools import partial
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from PIL import Image, ImageDraw
|
||||||
|
|
||||||
|
import cv2
|
||||||
|
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||||
|
from scepter.modules.annotator.registry import ANNOTATORS
|
||||||
|
from scepter.modules.utils.config import Config
|
||||||
|
from scepter.modules.utils.file_system import FS
|
||||||
|
from scipy import ndimage
|
||||||
|
from scipy.spatial import ConvexHull
|
||||||
|
from skimage.draw import polygon
|
||||||
|
|
||||||
|
|
||||||
|
@ANNOTATORS.register_class()
|
||||||
|
class MaskDrawAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||||
|
para_dict = {}
|
||||||
|
|
||||||
|
def __init__(self, cfg, logger=None):
|
||||||
|
super().__init__(cfg, logger=logger)
|
||||||
|
self.task_type = cfg.get('TASK_TYPE', 'input_box')
|
||||||
|
|
||||||
|
def forward(self, mask=None, image=None, input_box=None, task_type=None):
|
||||||
|
task_type = task_type if task_type is not None else self.task_type
|
||||||
|
|
||||||
|
if mask is not None:
|
||||||
|
if isinstance(mask, Image.Image):
|
||||||
|
mask = np.array(mask)
|
||||||
|
elif isinstance(mask, torch.Tensor):
|
||||||
|
mask = mask.detach().cpu().numpy()
|
||||||
|
elif isinstance(mask, np.ndarray):
|
||||||
|
mask = mask.copy()
|
||||||
|
else:
|
||||||
|
raise f'Unsurpport datatype{type(mask)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||||
|
|
||||||
|
if image is not None:
|
||||||
|
if isinstance(image, Image.Image):
|
||||||
|
image = np.array(image)
|
||||||
|
elif isinstance(image, torch.Tensor):
|
||||||
|
image = image.detach().cpu().numpy()
|
||||||
|
elif isinstance(image, np.ndarray):
|
||||||
|
image = image.copy()
|
||||||
|
else:
|
||||||
|
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||||
|
|
||||||
|
mask_shape = mask.shape
|
||||||
|
if task_type == 'mask_point':
|
||||||
|
scribble = mask.transpose(1, 0)
|
||||||
|
labeled_array, num_features = ndimage.label(scribble >= 255)
|
||||||
|
centers = ndimage.center_of_mass(scribble, labeled_array,
|
||||||
|
range(1, num_features + 1))
|
||||||
|
centers = np.array(centers)
|
||||||
|
out_mask = np.zeros(mask_shape, dtype=np.uint8)
|
||||||
|
hull = ConvexHull(centers)
|
||||||
|
hull_vertices = centers[hull.vertices]
|
||||||
|
rr, cc = polygon(hull_vertices[:, 1], hull_vertices[:, 0],
|
||||||
|
mask_shape)
|
||||||
|
out_mask[rr, cc] = 255
|
||||||
|
elif task_type == 'mask_box':
|
||||||
|
scribble = mask.transpose(1, 0)
|
||||||
|
labeled_array, num_features = ndimage.label(scribble >= 255)
|
||||||
|
centers = ndimage.center_of_mass(scribble, labeled_array,
|
||||||
|
range(1, num_features + 1))
|
||||||
|
centers = np.array(centers)
|
||||||
|
# (x1, y1, x2, y2)
|
||||||
|
x_min = centers[:, 0].min()
|
||||||
|
x_max = centers[:, 0].max()
|
||||||
|
y_min = centers[:, 1].min()
|
||||||
|
y_max = centers[:, 1].max()
|
||||||
|
out_mask = np.zeros(mask_shape, dtype=np.uint8)
|
||||||
|
out_mask[int(y_min):int(y_max) + 1,
|
||||||
|
int(x_min):int(x_max) + 1] = 255
|
||||||
|
if image is not None:
|
||||||
|
out_image = image[int(y_min):int(y_max) + 1,
|
||||||
|
int(x_min):int(x_max) + 1]
|
||||||
|
elif task_type == 'input_box':
|
||||||
|
if isinstance(input_box, list):
|
||||||
|
input_box = np.array(input_box)
|
||||||
|
x_min, y_min, x_max, y_max = input_box
|
||||||
|
out_mask = np.zeros(mask_shape, dtype=np.uint8)
|
||||||
|
out_mask[int(y_min):int(y_max) + 1,
|
||||||
|
int(x_min):int(x_max) + 1] = 255
|
||||||
|
if image is not None:
|
||||||
|
out_image = image[int(y_min):int(y_max) + 1,
|
||||||
|
int(x_min):int(x_max) + 1]
|
||||||
|
elif task_type == 'mask':
|
||||||
|
out_mask = mask
|
||||||
|
else:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
if image is not None:
|
||||||
|
return out_image, out_mask
|
||||||
|
else:
|
||||||
|
return out_mask
|
||||||
|
|
||||||
|
|
||||||
|
@ANNOTATORS.register_class()
|
||||||
|
class MaskAugAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||||
|
para_dict = {}
|
||||||
|
|
||||||
|
def __init__(self, cfg, logger=None):
|
||||||
|
super().__init__(cfg, logger=logger)
|
||||||
|
# original / original_expand / hull / hull_expand / bbox / bbox_expand
|
||||||
|
self.mask_cfg = cfg.get('MASK_CFG', [{
|
||||||
|
'mode': 'original',
|
||||||
|
'proba': 0.1
|
||||||
|
}, {
|
||||||
|
'mode': 'original_expand',
|
||||||
|
'proba': 0.1
|
||||||
|
}, {
|
||||||
|
'mode': 'hull',
|
||||||
|
'proba': 0.1
|
||||||
|
}, {
|
||||||
|
'mode': 'hull_expand',
|
||||||
|
'proba': 0.1,
|
||||||
|
'kwargs': {
|
||||||
|
'expand_rate': 0.2
|
||||||
|
}
|
||||||
|
}, {
|
||||||
|
'mode': 'bbox',
|
||||||
|
'proba': 0.1
|
||||||
|
}, {
|
||||||
|
'mode': 'bbox_expand',
|
||||||
|
'proba': 0.1,
|
||||||
|
'kwargs': {
|
||||||
|
'min_expand_rate': 0.2,
|
||||||
|
'max_expand_rate': 0.5
|
||||||
|
}
|
||||||
|
}])
|
||||||
|
self.mask_cfg = Config.get_dict(self.mask_cfg) if isinstance(
|
||||||
|
self.mask_cfg, Config) else self.mask_cfg
|
||||||
|
|
||||||
|
def forward(self, mask, mask_cfg=None):
|
||||||
|
mask_cfg = mask_cfg if mask_cfg is not None else self.mask_cfg
|
||||||
|
if not isinstance(mask, list):
|
||||||
|
is_batch = False
|
||||||
|
masks = [mask]
|
||||||
|
else:
|
||||||
|
is_batch = True
|
||||||
|
masks = mask
|
||||||
|
|
||||||
|
mask_func = self.get_mask_func(mask_cfg)
|
||||||
|
# print(mask_func)
|
||||||
|
aug_masks = []
|
||||||
|
for submask in masks:
|
||||||
|
mask = self.get_mask(submask)
|
||||||
|
valid, large, h, w, bbox = self.get_mask_info(mask)
|
||||||
|
# print(valid, large, h, w, bbox)
|
||||||
|
if valid:
|
||||||
|
mask = mask_func(mask, bbox, h, w)
|
||||||
|
else:
|
||||||
|
mask = mask.astype(np.uint8)
|
||||||
|
aug_masks.append(mask)
|
||||||
|
return aug_masks if is_batch else aug_masks[0]
|
||||||
|
|
||||||
|
def get_mask(self, mask):
|
||||||
|
if isinstance(mask, Image.Image):
|
||||||
|
mask = np.array(mask)
|
||||||
|
elif isinstance(mask, torch.Tensor):
|
||||||
|
mask = mask.detach().cpu().numpy()
|
||||||
|
elif isinstance(mask, np.ndarray):
|
||||||
|
mask = mask.copy()
|
||||||
|
else:
|
||||||
|
raise f'Unsurpport datatype{type(mask)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||||
|
return mask
|
||||||
|
|
||||||
|
def get_mask_info(self, mask):
|
||||||
|
h, w = mask.shape
|
||||||
|
locs = mask.nonzero()
|
||||||
|
valid = True
|
||||||
|
if len(locs) < 1 or locs[0].shape[0] < 1 or locs[1].shape[0] < 1:
|
||||||
|
valid = False
|
||||||
|
return valid, False, h, w, [0, 0, 0, 0]
|
||||||
|
|
||||||
|
left, right = np.min(locs[1]), np.max(locs[1])
|
||||||
|
top, bottom = np.min(locs[0]), np.max(locs[0])
|
||||||
|
bbox = [left, top, right, bottom]
|
||||||
|
|
||||||
|
large = False
|
||||||
|
if (right - left + 1) * (bottom - top + 1) > 0.9 * h * w:
|
||||||
|
large = True
|
||||||
|
return valid, large, h, w, bbox
|
||||||
|
|
||||||
|
def get_expand_params(self, mask_kwargs):
|
||||||
|
if 'expand_rate' in mask_kwargs:
|
||||||
|
expand_rate = mask_kwargs['expand_rate']
|
||||||
|
elif 'min_expand_rate' in mask_kwargs and 'max_expand_rate' in mask_kwargs:
|
||||||
|
expand_rate = random.uniform(mask_kwargs['min_expand_rate'],
|
||||||
|
mask_kwargs['max_expand_rate'])
|
||||||
|
else:
|
||||||
|
expand_rate = 0.3
|
||||||
|
|
||||||
|
if 'expand_iters' in mask_kwargs:
|
||||||
|
expand_iters = mask_kwargs['expand_iters']
|
||||||
|
else:
|
||||||
|
expand_iters = random.randint(1, 10)
|
||||||
|
|
||||||
|
if 'expand_lrtp' in mask_kwargs:
|
||||||
|
expand_lrtp = mask_kwargs['expand_lrtp']
|
||||||
|
else:
|
||||||
|
expand_lrtp = [
|
||||||
|
random.random(),
|
||||||
|
random.random(),
|
||||||
|
random.random(),
|
||||||
|
random.random()
|
||||||
|
]
|
||||||
|
|
||||||
|
return expand_rate, expand_iters, expand_lrtp
|
||||||
|
|
||||||
|
def get_mask_func(self, mask_cfg):
|
||||||
|
if not isinstance(mask_cfg, list):
|
||||||
|
mask_cfg = [mask_cfg]
|
||||||
|
probas = [
|
||||||
|
item['proba'] if 'proba' in item else 1.0 / len(mask_cfg)
|
||||||
|
for item in mask_cfg
|
||||||
|
]
|
||||||
|
sel_mask_cfg = random.choices(mask_cfg, weights=probas, k=1)[0]
|
||||||
|
mode = sel_mask_cfg['mode'] if 'mode' in sel_mask_cfg else 'original'
|
||||||
|
mask_kwargs = sel_mask_cfg[
|
||||||
|
'kwargs'] if 'kwargs' in sel_mask_cfg else {}
|
||||||
|
|
||||||
|
if mode == 'random':
|
||||||
|
mode = random.choice([
|
||||||
|
'original', 'original_expand', 'hull', 'hull_expand', 'bbox',
|
||||||
|
'bbox_expand'
|
||||||
|
])
|
||||||
|
if mode == 'original':
|
||||||
|
mask_func = partial(self.generate_mask)
|
||||||
|
elif mode == 'original_expand':
|
||||||
|
expand_rate, expand_iters, expand_lrtp = self.get_expand_params(
|
||||||
|
mask_kwargs)
|
||||||
|
mask_func = partial(self.generate_mask,
|
||||||
|
expand_rate=expand_rate,
|
||||||
|
expand_iters=expand_iters,
|
||||||
|
expand_lrtp=expand_lrtp)
|
||||||
|
elif mode == 'hull':
|
||||||
|
clockwise = random.choice([
|
||||||
|
True, False
|
||||||
|
]) if 'clockwise' not in mask_kwargs else mask_kwargs['clockwise']
|
||||||
|
mask_func = partial(self.generate_hull_mask, clockwise=clockwise)
|
||||||
|
elif mode == 'hull_expand':
|
||||||
|
expand_rate, expand_iters, expand_lrtp = self.get_expand_params(
|
||||||
|
mask_kwargs)
|
||||||
|
clockwise = random.choice([
|
||||||
|
True, False
|
||||||
|
]) if 'clockwise' not in mask_kwargs else mask_kwargs['clockwise']
|
||||||
|
mask_func = partial(self.generate_hull_mask,
|
||||||
|
clockwise=clockwise,
|
||||||
|
expand_rate=expand_rate,
|
||||||
|
expand_iters=expand_iters,
|
||||||
|
expand_lrtp=expand_lrtp)
|
||||||
|
elif mode == 'bbox':
|
||||||
|
mask_func = partial(self.generate_bbox_mask)
|
||||||
|
elif mode == 'bbox_expand':
|
||||||
|
expand_rate, expand_iters, expand_lrtp = self.get_expand_params(
|
||||||
|
mask_kwargs)
|
||||||
|
mask_func = partial(self.generate_bbox_mask,
|
||||||
|
expand_rate=expand_rate,
|
||||||
|
expand_iters=expand_iters,
|
||||||
|
expand_lrtp=expand_lrtp)
|
||||||
|
else:
|
||||||
|
raise NotImplementedError
|
||||||
|
return mask_func
|
||||||
|
|
||||||
|
def generate_mask(self,
|
||||||
|
mask,
|
||||||
|
bbox,
|
||||||
|
h,
|
||||||
|
w,
|
||||||
|
expand_rate=None,
|
||||||
|
expand_iters=None,
|
||||||
|
expand_lrtp=None):
|
||||||
|
bin_mask = mask.astype(np.uint8)
|
||||||
|
if expand_rate:
|
||||||
|
bin_mask = self.rand_expand_mask(bin_mask, bbox, h, w, expand_rate,
|
||||||
|
expand_iters, expand_lrtp)
|
||||||
|
return bin_mask
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def rand_expand_mask(mask,
|
||||||
|
bbox,
|
||||||
|
h,
|
||||||
|
w,
|
||||||
|
expand_rate=None,
|
||||||
|
expand_iters=None,
|
||||||
|
expand_lrtp=None):
|
||||||
|
expand_rate = 0.3 if expand_rate is None else expand_rate
|
||||||
|
expand_iters = random.randint(
|
||||||
|
1, 10) if expand_iters is None else expand_iters
|
||||||
|
expand_lrtp = [
|
||||||
|
random.random(),
|
||||||
|
random.random(),
|
||||||
|
random.random(),
|
||||||
|
random.random()
|
||||||
|
] if expand_lrtp is None else expand_lrtp
|
||||||
|
# print('iters', expand_iters, 'expand_rate', expand_rate, 'expand_lrtp', expand_lrtp)
|
||||||
|
# mask = np.squeeze(mask)
|
||||||
|
left, top, right, bottom = bbox
|
||||||
|
# mask expansion
|
||||||
|
box_w = (right - left + 1) * expand_rate
|
||||||
|
box_h = (bottom - top + 1) * expand_rate
|
||||||
|
left_, right_ = int(
|
||||||
|
expand_lrtp[0] * min(box_w, left / 2) / expand_iters), int(
|
||||||
|
expand_lrtp[1] * min(box_w, (w - right) / 2) / expand_iters)
|
||||||
|
top_, bottom_ = int(
|
||||||
|
expand_lrtp[2] * min(box_h, top / 2) / expand_iters), int(
|
||||||
|
expand_lrtp[3] * min(box_h, (h - bottom) / 2) / expand_iters)
|
||||||
|
kernel_size = max(left_, right_, top_, bottom_)
|
||||||
|
if kernel_size > 0:
|
||||||
|
kernel = np.zeros((kernel_size * 2, kernel_size * 2),
|
||||||
|
dtype=np.uint8)
|
||||||
|
new_left, new_right = kernel_size - right_, kernel_size + left_
|
||||||
|
new_top, new_bottom = kernel_size - bottom_, kernel_size + top_
|
||||||
|
kernel[new_top:new_bottom + 1, new_left:new_right + 1] = 1
|
||||||
|
mask = mask.astype(np.uint8)
|
||||||
|
mask = cv2.dilate(mask, kernel,
|
||||||
|
iterations=expand_iters).astype(np.uint8)
|
||||||
|
# mask = new_mask - (mask / 2).astype(np.uint8)
|
||||||
|
# mask = np.expand_dims(mask, axis=-1)
|
||||||
|
return mask
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _convexhull(image, clockwise):
|
||||||
|
# print('clockwise', clockwise)
|
||||||
|
contours, hierarchy = cv2.findContours(image, 2, 1)
|
||||||
|
cnt = np.concatenate(contours) # merge all regions
|
||||||
|
hull = cv2.convexHull(cnt, clockwise=clockwise)
|
||||||
|
hull = np.squeeze(hull, axis=1).astype(np.float32).tolist()
|
||||||
|
hull = [tuple(x) for x in hull]
|
||||||
|
return hull # b, 1, 2
|
||||||
|
|
||||||
|
def generate_hull_mask(self,
|
||||||
|
mask,
|
||||||
|
bbox,
|
||||||
|
h,
|
||||||
|
w,
|
||||||
|
clockwise=None,
|
||||||
|
expand_rate=None,
|
||||||
|
expand_iters=None,
|
||||||
|
expand_lrtp=None):
|
||||||
|
clockwise = random.choice([True, False
|
||||||
|
]) if clockwise is None else clockwise
|
||||||
|
hull = self._convexhull(mask, clockwise)
|
||||||
|
mask_img = Image.new('L', (w, h), 0)
|
||||||
|
pt_list = hull
|
||||||
|
mask_img_draw = ImageDraw.Draw(mask_img)
|
||||||
|
mask_img_draw.polygon(pt_list, fill=255)
|
||||||
|
bin_mask = np.array(mask_img).astype(np.uint8)
|
||||||
|
if expand_rate:
|
||||||
|
bin_mask = self.rand_expand_mask(bin_mask, bbox, h, w, expand_rate,
|
||||||
|
expand_iters, expand_lrtp)
|
||||||
|
return bin_mask
|
||||||
|
|
||||||
|
def generate_bbox_mask(self,
|
||||||
|
mask,
|
||||||
|
bbox,
|
||||||
|
h,
|
||||||
|
w,
|
||||||
|
expand_rate=None,
|
||||||
|
expand_iters=None,
|
||||||
|
expand_lrtp=None):
|
||||||
|
left, top, right, bottom = bbox
|
||||||
|
bin_mask = np.zeros((h, w), dtype=np.uint8)
|
||||||
|
bin_mask[top:bottom + 1, left:right + 1] = 255
|
||||||
|
if expand_rate:
|
||||||
|
bin_mask = self.rand_expand_mask(bin_mask, bbox, h, w, expand_rate,
|
||||||
|
expand_iters, expand_lrtp)
|
||||||
|
return bin_mask
|
||||||
|
|
||||||
|
|
||||||
|
@ANNOTATORS.register_class()
|
||||||
|
class MaskLayoutAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||||
|
para_dict = {}
|
||||||
|
|
||||||
|
def __init__(self, cfg, logger=None):
|
||||||
|
super().__init__(cfg, logger=logger)
|
||||||
|
ram_tag_color = cfg.get('RAM_TAG_COLOR', None)
|
||||||
|
default_color = cfg.get('DEFAULT_COLOR', [0, 0, 0])
|
||||||
|
self.use_aug = cfg.get('USE_AUG', False)
|
||||||
|
self.color_dict = {'default': tuple(default_color)}
|
||||||
|
if ram_tag_color is not None:
|
||||||
|
with FS.get_object(ram_tag_color) as object:
|
||||||
|
lines = object.decode('utf-8').strip().split('\n')
|
||||||
|
lines = [id_name_color.split('#;#') for id_name_color in lines]
|
||||||
|
self.color_dict.update({
|
||||||
|
id_name_color[1]: tuple(eval(id_name_color[2]))
|
||||||
|
for id_name_color in lines
|
||||||
|
})
|
||||||
|
if self.use_aug:
|
||||||
|
mask_aug_dict = {'NAME': 'MaskAugAnnotator'}
|
||||||
|
mask_aug_cfg = Config(cfg_dict=mask_aug_dict, load=False)
|
||||||
|
self.mask_aug_anno = ANNOTATORS.build(mask_aug_cfg)
|
||||||
|
|
||||||
|
def find_contours(self, mask):
|
||||||
|
# @mask: gray cv2 image
|
||||||
|
# contours, hier = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE)
|
||||||
|
contours, hier = cv2.findContours(mask, cv2.RETR_EXTERNAL,
|
||||||
|
cv2.CHAIN_APPROX_SIMPLE)
|
||||||
|
return contours
|
||||||
|
|
||||||
|
def draw_contours(self, canvas, contour, color):
|
||||||
|
canvas = np.ascontiguousarray(canvas, dtype=np.uint8)
|
||||||
|
canvas = cv2.drawContours(canvas, contour, -1, color, thickness=3)
|
||||||
|
return canvas
|
||||||
|
|
||||||
|
def get_mask(self, mask):
|
||||||
|
if isinstance(mask, Image.Image):
|
||||||
|
mask = np.array(mask)
|
||||||
|
elif isinstance(mask, torch.Tensor):
|
||||||
|
mask = mask.detach().cpu().numpy()
|
||||||
|
elif isinstance(mask, np.ndarray):
|
||||||
|
mask = mask.copy()
|
||||||
|
else:
|
||||||
|
raise f'Unsurpport datatype{type(mask)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||||
|
return mask
|
||||||
|
|
||||||
|
def forward(self, mask=None, color=None, label=None, mask_cfg=None):
|
||||||
|
if not isinstance(mask, list):
|
||||||
|
is_batch = False
|
||||||
|
mask = [mask]
|
||||||
|
else:
|
||||||
|
is_batch = True
|
||||||
|
|
||||||
|
if label is not None and label in self.color_dict:
|
||||||
|
color = self.color_dict[label]
|
||||||
|
elif color is not None:
|
||||||
|
color = color
|
||||||
|
else:
|
||||||
|
color = self.color_dict['default']
|
||||||
|
|
||||||
|
ret_data = []
|
||||||
|
for sub_mask in mask:
|
||||||
|
sub_mask = self.get_mask(sub_mask)
|
||||||
|
if self.use_aug:
|
||||||
|
sub_mask = self.mask_aug_anno(sub_mask, mask_cfg)
|
||||||
|
canvas = np.ones((sub_mask.shape[0], sub_mask.shape[1], 3)) * 255
|
||||||
|
contour = self.find_contours(sub_mask)
|
||||||
|
frame = self.draw_contours(canvas, contour, color)
|
||||||
|
ret_data.append(frame)
|
||||||
|
|
||||||
|
if is_batch:
|
||||||
|
return ret_data
|
||||||
|
else:
|
||||||
|
return ret_data[0]
|
||||||
@@ -10,7 +10,7 @@ class BaseModel(torch.nn.Module):
|
|||||||
Args:
|
Args:
|
||||||
path (str): file path
|
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:
|
if 'optimizer' in parameters:
|
||||||
parameters = parameters['model']
|
parameters = parameters['model']
|
||||||
|
|||||||
@@ -29,7 +29,7 @@ class MLSDdetector(BaseAnnotator, metaclass=ABCMeta):
|
|||||||
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
|
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
|
||||||
if pretrained_model:
|
if pretrained_model:
|
||||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
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.model = model.eval()
|
||||||
self.thr_v = cfg.get('THR_V', 0.1)
|
self.thr_v = cfg.get('THR_V', 0.1)
|
||||||
self.thr_d = cfg.get('THR_D', 0.1)
|
self.thr_d = cfg.get('THR_D', 0.1)
|
||||||
|
|||||||
@@ -423,7 +423,7 @@ class Hand(object):
|
|||||||
self.model = handpose_model()
|
self.model = handpose_model()
|
||||||
if torch.cuda.is_available():
|
if torch.cuda.is_available():
|
||||||
self.model = self.model.to(device)
|
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.load_state_dict(model_dict)
|
||||||
self.model.eval()
|
self.model.eval()
|
||||||
self.device = device
|
self.device = device
|
||||||
@@ -503,7 +503,7 @@ class Body(object):
|
|||||||
self.model = bodypose_model()
|
self.model = bodypose_model()
|
||||||
if torch.cuda.is_available():
|
if torch.cuda.is_available():
|
||||||
self.model = self.model.to(device)
|
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.load_state_dict(model_dict)
|
||||||
self.model.eval()
|
self.model.eval()
|
||||||
self.device = device
|
self.device = device
|
||||||
|
|||||||
@@ -98,9 +98,15 @@ class OutpaintingAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
|||||||
draw.rectangle(
|
draw.rectangle(
|
||||||
(left + (self.mask_blur * 2 if left > 0 else 0), up +
|
(left + (self.mask_blur * 2 if left > 0 else 0), up +
|
||||||
(self.mask_blur * 2 if up > 0 else 0), mask.width - right -
|
(self.mask_blur * 2 if up > 0 else 0), mask.width - right -
|
||||||
(self.mask_blur * 2 if right > 0 else 0), mask.height - down -
|
(self.mask_blur * 2 if right > 0 else 0) - 1, mask.height - down -
|
||||||
(self.mask_blur * 2 if down > 0 else 0)),
|
(self.mask_blur * 2 if down > 0 else 0) - 1),
|
||||||
fill='black')
|
fill='black')
|
||||||
|
# draw.rectangle(
|
||||||
|
# (left + (self.mask_blur * 2 if left > 0 else 0), up +
|
||||||
|
# (self.mask_blur * 2 if up > 0 else 0), left + src_width -
|
||||||
|
# (self.mask_blur * 2 if right > 0 else 0), up + src_height -
|
||||||
|
# (self.mask_blur * 2 if down > 0 else 0)),
|
||||||
|
# fill='black')
|
||||||
else:
|
else:
|
||||||
bbox = self.get_box(np.array(mask))
|
bbox = self.get_box(np.array(mask))
|
||||||
if bbox is None:
|
if bbox is None:
|
||||||
|
|||||||
@@ -882,7 +882,7 @@ class PiDiAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
|||||||
if pretrained_model:
|
if pretrained_model:
|
||||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
||||||
state = torch.load(local_path,
|
state = torch.load(local_path,
|
||||||
map_location='cpu')['state_dict']
|
map_location='cpu', weights_only=True)['state_dict']
|
||||||
if vanilla_cnn:
|
if vanilla_cnn:
|
||||||
state = convert_pidinet(state, 'carv4')
|
state = convert_pidinet(state, 'carv4')
|
||||||
state = {
|
state = {
|
||||||
|
|||||||
@@ -0,0 +1,62 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
|
import torch
|
||||||
|
import random
|
||||||
|
import numpy as np
|
||||||
|
import argparse
|
||||||
|
|
||||||
|
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||||
|
from scepter.modules.annotator.registry import ANNOTATORS
|
||||||
|
from scepter.modules.utils.config import Config
|
||||||
|
from scepter.modules.utils.distribute import we
|
||||||
|
from scepter.modules.utils.file_system import FS
|
||||||
|
|
||||||
|
try:
|
||||||
|
from raft import RAFT
|
||||||
|
from raft.utils.utils import InputPadder
|
||||||
|
from raft.utils import flow_viz
|
||||||
|
except:
|
||||||
|
import warnings
|
||||||
|
warnings.warn("ignore raft import, please pip install raft.")
|
||||||
|
|
||||||
|
|
||||||
|
@ANNOTATORS.register_class()
|
||||||
|
class RAFTAnnotator(BaseAnnotator):
|
||||||
|
para_dict = {}
|
||||||
|
|
||||||
|
def __init__(self, cfg, logger=None):
|
||||||
|
super().__init__(cfg, logger=logger)
|
||||||
|
params = {
|
||||||
|
"small": False,
|
||||||
|
"mixed_precision": False,
|
||||||
|
"alternate_corr": False
|
||||||
|
}
|
||||||
|
params = argparse.Namespace(**params)
|
||||||
|
model = RAFT(params)
|
||||||
|
if cfg.PRETRAINED_MODEL is not None:
|
||||||
|
with FS.get_from(cfg.PRETRAINED_MODEL,
|
||||||
|
wait_finish=True) as local_path:
|
||||||
|
model.load_state_dict({k.replace('module.', ''): v for k, v in torch.load(local_path, map_location="cpu", weights_only=True).items()})
|
||||||
|
self.model = model.to(we.device_id).eval()
|
||||||
|
|
||||||
|
def forward(self, frames):
|
||||||
|
# frames / RGB
|
||||||
|
frames = [torch.from_numpy(frame.astype(np.uint8)).permute(2, 0, 1).float()[None].to(we.device_id) for frame in frames]
|
||||||
|
flow_up_list, flow_up_vis_list = [], []
|
||||||
|
with torch.no_grad():
|
||||||
|
for i, (image1, image2) in enumerate(zip(frames[:-1], frames[1:])):
|
||||||
|
padder = InputPadder(image1.shape)
|
||||||
|
image1, image2 = padder.pad(image1, image2)
|
||||||
|
flow_low, flow_up = self.model(image1, image2, iters=20, test_mode=True)
|
||||||
|
flow_up = flow_up[0].permute(1, 2, 0).cpu().numpy()
|
||||||
|
flow_up_vis = flow_viz.flow_to_image(flow_up)
|
||||||
|
flow_up_list.append(flow_up)
|
||||||
|
flow_up_vis_list.append(flow_up_vis)
|
||||||
|
return flow_up_list, flow_up_vis_list # RGB
|
||||||
|
|
||||||
|
|
||||||
|
@ANNOTATORS.register_class()
|
||||||
|
class RAFTVisAnnotator(RAFTAnnotator):
|
||||||
|
def forward(self, frames):
|
||||||
|
flow_up_list, flow_up_vis_list = super().forward(frames)
|
||||||
|
return flow_up_vis_list[:1] + flow_up_vis_list
|
||||||
@@ -0,0 +1,95 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
|
import random
|
||||||
|
from abc import ABCMeta
|
||||||
|
|
||||||
|
import cv2
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
|
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||||
|
from scepter.modules.annotator.registry import ANNOTATORS
|
||||||
|
from scepter.modules.utils.config import Config, dict_to_yaml
|
||||||
|
|
||||||
|
|
||||||
|
@ANNOTATORS.register_class()
|
||||||
|
class RegionCanvasAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||||
|
para_dict = {}
|
||||||
|
|
||||||
|
def __init__(self, cfg, logger=None):
|
||||||
|
super().__init__(cfg, logger=logger)
|
||||||
|
self.scale_range = cfg.get('SCALE_RANGE', [0.75, 1.0])
|
||||||
|
self.canvas_value = cfg.get('CANVAS_VALUE', 255)
|
||||||
|
self.use_resize = cfg.get('USE_RESIZE', True)
|
||||||
|
self.use_canvas = cfg.get('USE_CANVAS', True)
|
||||||
|
self.use_aug = cfg.get('USE_AUG', False)
|
||||||
|
if self.use_aug:
|
||||||
|
mask_aug_dict = {'NAME': 'MaskAugAnnotator'}
|
||||||
|
mask_aug_cfg = Config(cfg_dict=mask_aug_dict, load=False)
|
||||||
|
self.mask_aug_anno = ANNOTATORS.build(mask_aug_cfg)
|
||||||
|
|
||||||
|
|
||||||
|
def forward(self,
|
||||||
|
image,
|
||||||
|
mask,
|
||||||
|
mask_cfg=None):
|
||||||
|
if isinstance(image, Image.Image):
|
||||||
|
image = np.array(image)
|
||||||
|
elif isinstance(image, torch.Tensor):
|
||||||
|
image = image.detach().cpu().numpy()
|
||||||
|
elif isinstance(image, np.ndarray):
|
||||||
|
image = image
|
||||||
|
else:
|
||||||
|
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||||
|
|
||||||
|
mask = np.array(mask).astype(np.uint8)
|
||||||
|
image_h, image_w = image.shape[:2]
|
||||||
|
|
||||||
|
if self.use_aug:
|
||||||
|
mask = self.mask_aug_anno(mask, mask_cfg)
|
||||||
|
|
||||||
|
# get region with white bg
|
||||||
|
image[np.array(mask) == 0] = self.canvas_value
|
||||||
|
x, y, w, h = cv2.boundingRect(mask)
|
||||||
|
region_crop = image[y:y + h, x:x + w]
|
||||||
|
|
||||||
|
if self.use_resize:
|
||||||
|
# resize region
|
||||||
|
scale_min, scale_max = self.scale_range
|
||||||
|
scale_factor = random.uniform(scale_min, scale_max)
|
||||||
|
new_w, new_h = int(image_w * scale_factor), int(image_h * scale_factor)
|
||||||
|
obj_scale_factor = min(new_w/w, new_h/h)
|
||||||
|
|
||||||
|
new_w = int(w * obj_scale_factor)
|
||||||
|
new_h = int(h * obj_scale_factor)
|
||||||
|
region_crop_resized = cv2.resize(region_crop, (new_w, new_h), interpolation=cv2.INTER_AREA)
|
||||||
|
else:
|
||||||
|
region_crop_resized = region_crop
|
||||||
|
|
||||||
|
if self.use_canvas:
|
||||||
|
# plot region into canvas
|
||||||
|
new_canvas = np.ones_like(image) * self.canvas_value
|
||||||
|
max_x = max(0, image_w - new_w)
|
||||||
|
max_y = max(0, image_h - new_h)
|
||||||
|
new_x = random.randint(0, max_x)
|
||||||
|
new_y = random.randint(0, max_y)
|
||||||
|
|
||||||
|
new_canvas[new_y:new_y + new_h, new_x:new_x + new_w] = region_crop_resized
|
||||||
|
else:
|
||||||
|
new_canvas = region_crop_resized
|
||||||
|
return new_canvas
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def get_config_template():
|
||||||
|
return dict_to_yaml('ANNOTATORS',
|
||||||
|
__class__.__name__,
|
||||||
|
RegionCanvasAnnotator.para_dict,
|
||||||
|
set_name=True)
|
||||||
|
|
||||||
|
|
||||||
|
@ANNOTATORS.register_class()
|
||||||
|
class RegionCanvasCropAnnotator(RegionCanvasAnnotator):
|
||||||
|
def __init__(self, cfg, logger=None):
|
||||||
|
super().__init__(cfg, logger=logger)
|
||||||
|
self.use_resize, self.use_canvas = False, False
|
||||||
@@ -10,7 +10,11 @@ import torchvision.transforms as T
|
|||||||
from PIL import Image
|
from PIL import Image
|
||||||
from pycocotools import mask as mask_utils
|
from pycocotools import mask as mask_utils
|
||||||
from scipy import ndimage
|
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 torchvision.ops.boxes import batched_nms
|
||||||
|
|
||||||
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||||
|
|||||||
@@ -86,7 +86,7 @@ class SketchAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
|||||||
std=0.0858381272736797).eval()
|
std=0.0858381272736797).eval()
|
||||||
if pretrained_model:
|
if pretrained_model:
|
||||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
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)
|
self.model.load_state_dict(state)
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
|
|||||||
@@ -0,0 +1,153 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
|
from abc import ABCMeta
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from PIL import Image
|
||||||
|
from scipy import ndimage
|
||||||
|
try:
|
||||||
|
from sklearn.cluster import KMeans
|
||||||
|
except:
|
||||||
|
import warnings
|
||||||
|
warnings.warn("ignore sklearn import, please pip install scikit-learn.")
|
||||||
|
|
||||||
|
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||||
|
from scepter.modules.annotator.registry import ANNOTATORS
|
||||||
|
from scepter.modules.utils.config import dict_to_yaml
|
||||||
|
from scepter.modules.utils.file_system import FS
|
||||||
|
import pycocotools.mask as mask_utils
|
||||||
|
|
||||||
|
|
||||||
|
def single_mask_to_rle(mask):
|
||||||
|
rle = mask_utils.encode(np.array(mask[:, :, None], order="F", dtype="uint8"))[0]
|
||||||
|
rle["counts"] = rle["counts"].decode("utf-8")
|
||||||
|
return rle
|
||||||
|
|
||||||
|
def single_rle_to_mask(rle):
|
||||||
|
mask = np.array(mask_utils.decode(rle)).astype(np.uint8)
|
||||||
|
return mask
|
||||||
|
|
||||||
|
def single_mask_to_xyxy(mask):
|
||||||
|
bbox = np.zeros((4), dtype=int)
|
||||||
|
rows, cols = np.where(np.array(mask))
|
||||||
|
if len(rows) > 0 and len(cols) > 0:
|
||||||
|
x_min, x_max = np.min(cols), np.max(cols)
|
||||||
|
y_min, y_max = np.min(rows), np.max(rows)
|
||||||
|
bbox[:] = [x_min, y_min, x_max, y_max]
|
||||||
|
return bbox.tolist()
|
||||||
|
|
||||||
|
@ANNOTATORS.register_class()
|
||||||
|
class SAM2DrawVideoAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||||
|
para_dict = {}
|
||||||
|
|
||||||
|
def __init__(self, cfg, logger=None):
|
||||||
|
super().__init__(cfg, logger=logger)
|
||||||
|
self.task_type = cfg.get('TASK_TYPE', 'input_box')
|
||||||
|
from sam2.build_sam import build_sam2_video_predictor
|
||||||
|
config_path = FS.get_from(cfg.CONFIG_PATH, local_path=cfg.CONFIG_LOCAL_PATH, wait_finish=True)
|
||||||
|
pretrained_model = FS.get_from(cfg.PRETRAINED_MODEL, wait_finish=True)
|
||||||
|
self.video_predictor = build_sam2_video_predictor(config_path, pretrained_model, fill_hole_area=0)
|
||||||
|
|
||||||
|
def forward(self,
|
||||||
|
video,
|
||||||
|
input_box=None,
|
||||||
|
mask=None,
|
||||||
|
task_type=None):
|
||||||
|
task_type = task_type if task_type is not None else self.task_type
|
||||||
|
|
||||||
|
if mask is not None:
|
||||||
|
if isinstance(mask, Image.Image):
|
||||||
|
mask = np.array(mask)
|
||||||
|
elif isinstance(mask, torch.Tensor):
|
||||||
|
mask = mask.detach().cpu().numpy()
|
||||||
|
elif isinstance(mask, np.ndarray):
|
||||||
|
mask = mask.copy()
|
||||||
|
else:
|
||||||
|
raise f'Unsurpport datatype{type(mask)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||||
|
|
||||||
|
if task_type == 'mask_point':
|
||||||
|
if len(mask.shape) == 3:
|
||||||
|
scribble = mask.transpose(2, 1, 0)[0]
|
||||||
|
else:
|
||||||
|
scribble = mask.transpose(1, 0) # (H, W) -> (W, H)
|
||||||
|
labeled_array, num_features = ndimage.label(scribble >= 255)
|
||||||
|
centers = ndimage.center_of_mass(scribble, labeled_array,
|
||||||
|
range(1, num_features + 1))
|
||||||
|
point_coords = np.array(centers)
|
||||||
|
point_labels = np.array([1] * len(centers))
|
||||||
|
sample = {
|
||||||
|
'points': point_coords,
|
||||||
|
'labels': point_labels
|
||||||
|
}
|
||||||
|
elif task_type == 'mask_box':
|
||||||
|
if len(mask.shape) == 3:
|
||||||
|
scribble = mask.transpose(2, 1, 0)[0]
|
||||||
|
else:
|
||||||
|
scribble = mask.transpose(1, 0) # (H, W) -> (W, H)
|
||||||
|
labeled_array, num_features = ndimage.label(scribble >= 255)
|
||||||
|
centers = ndimage.center_of_mass(scribble, labeled_array,
|
||||||
|
range(1, num_features + 1))
|
||||||
|
centers = np.array(centers)
|
||||||
|
# (x1, y1, x2, y2)
|
||||||
|
x_min = centers[:, 0].min()
|
||||||
|
x_max = centers[:, 0].max()
|
||||||
|
y_min = centers[:, 1].min()
|
||||||
|
y_max = centers[:, 1].max()
|
||||||
|
bbox = np.array([x_min, y_min, x_max, y_max])
|
||||||
|
sample = {'box': bbox}
|
||||||
|
elif task_type == 'input_box':
|
||||||
|
if isinstance(input_box, list):
|
||||||
|
input_box = np.array(input_box)
|
||||||
|
sample = {'box': input_box}
|
||||||
|
elif task_type == 'mask':
|
||||||
|
sample = {'mask': mask}
|
||||||
|
else:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
ann_frame_idx = 0
|
||||||
|
object_id = 0
|
||||||
|
with (torch.inference_mode(), torch.autocast("cuda", dtype=torch.bfloat16)):
|
||||||
|
|
||||||
|
inference_state = self.video_predictor.init_state(video_path=video)
|
||||||
|
|
||||||
|
if task_type in ['mask_point', 'mask_box', 'input_box']:
|
||||||
|
_, out_obj_ids, out_mask_logits = self.video_predictor.add_new_points_or_box(
|
||||||
|
inference_state=inference_state,
|
||||||
|
frame_idx=ann_frame_idx,
|
||||||
|
obj_id=object_id,
|
||||||
|
**sample
|
||||||
|
)
|
||||||
|
elif task_type in ['mask']:
|
||||||
|
_, out_obj_ids, out_mask_logits = self.video_predictor.add_new_mask(
|
||||||
|
inference_state=inference_state,
|
||||||
|
frame_idx=ann_frame_idx,
|
||||||
|
obj_id=object_id,
|
||||||
|
**sample
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
video_segments = {} # video_segments contains the per-frame segmentation results
|
||||||
|
for out_frame_idx, out_obj_ids, out_mask_logits in self.video_predictor.propagate_in_video(inference_state):
|
||||||
|
frame_segments = {}
|
||||||
|
for i, out_obj_id in enumerate(out_obj_ids):
|
||||||
|
mask = (out_mask_logits[i] > 0.0).cpu().numpy().squeeze(0)
|
||||||
|
frame_segments[out_obj_id] = {
|
||||||
|
"mask": single_mask_to_rle(mask),
|
||||||
|
"mask_area": int(mask.sum()),
|
||||||
|
"mask_box": single_mask_to_xyxy(mask),
|
||||||
|
}
|
||||||
|
video_segments[out_frame_idx] = frame_segments
|
||||||
|
|
||||||
|
ret_data = {
|
||||||
|
"annotations": video_segments
|
||||||
|
}
|
||||||
|
return ret_data
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def get_config_template():
|
||||||
|
return dict_to_yaml('ANNOTATORS',
|
||||||
|
__class__.__name__,
|
||||||
|
SAM2DrawVideoAnnotator.para_dict,
|
||||||
|
set_name=True)
|
||||||
@@ -1,4 +1,21 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# 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 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# 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,
|
if TYPE_CHECKING:
|
||||||
ImageClassifyPublicDataset,
|
from scepter.modules.data.dataset.base_dataset import BaseDataset
|
||||||
ImageTextPairDataset,
|
from scepter.modules.data.dataset.dataset import (Image2ImageDataset,
|
||||||
Text2ImageDataset)
|
ImageClassifyPublicDataset,
|
||||||
from scepter.modules.data.dataset.ms_dataset import (
|
ImageTextPairDataset,
|
||||||
ImageTextPairFolderDataset, ImageTextPairMSDataset)
|
Text2ImageDataset)
|
||||||
from scepter.modules.data.dataset.registry import DATASETS
|
from scepter.modules.data.dataset.ms_dataset import (
|
||||||
from scepter.modules.data.dataset.video_gen_dataset import VideoGenDataset
|
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)
|
overwrite=False)
|
||||||
self.worker_id = worker_id
|
self.worker_id = worker_id
|
||||||
self.logger = self.worker_logger
|
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"]
|
self.seed = self.local_we["seed"]
|
||||||
we.set_env(self.local_we)
|
we.set_env(self.local_we)
|
||||||
|
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
import numbers
|
import numbers
|
||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
|
import copy
|
||||||
from collections.abc import Iterable
|
from collections.abc import Iterable
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|||||||
@@ -386,6 +386,11 @@ class ImageTextPairMSDatasetForACE(BaseDataset):
|
|||||||
'description':
|
'description':
|
||||||
'The keywords sign you want to add, which is like <{HIGHLIGHT_KEYWORDS}{KEYWORDS_SIGN}>'
|
'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': {
|
'OUTPUT_SIZE': {
|
||||||
'value':
|
'value':
|
||||||
None,
|
None,
|
||||||
@@ -414,6 +419,8 @@ class ImageTextPairMSDatasetForACE(BaseDataset):
|
|||||||
self.replace_keywords = cfg.get('HIGHLIGHT_KEYWORDS', '')
|
self.replace_keywords = cfg.get('HIGHLIGHT_KEYWORDS', '')
|
||||||
self.keywords_sign = cfg.get('KEYWORDS_SIGN', '')
|
self.keywords_sign = cfg.get('KEYWORDS_SIGN', '')
|
||||||
self.add_indicator = cfg.get('ADD_INDICATOR', False)
|
self.add_indicator = cfg.get('ADD_INDICATOR', False)
|
||||||
|
|
||||||
|
self.align_size = cfg.get('ALIGN_SIZE', False)
|
||||||
# Use modelscope dataset
|
# Use modelscope dataset
|
||||||
if not ms_dataset_name:
|
if not ms_dataset_name:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -492,7 +499,7 @@ class ImageTextPairMSDatasetForACE(BaseDataset):
|
|||||||
tar_image_path,
|
tar_image_path,
|
||||||
cvt_type='RGB')
|
cvt_type='RGB')
|
||||||
src_image = self.image_preprocess(src_image)
|
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)
|
tar_image = self.transforms(tar_image)
|
||||||
src_image = self.transforms(src_image)
|
src_image = self.transforms(src_image)
|
||||||
@@ -501,13 +508,13 @@ class ImageTextPairMSDatasetForACE(BaseDataset):
|
|||||||
if self.add_indicator:
|
if self.add_indicator:
|
||||||
if '{image}' not in prompt:
|
if '{image}' not in prompt:
|
||||||
prompt = '{image}, ' + prompt
|
prompt = '{image}, ' + prompt
|
||||||
|
|
||||||
return {
|
return {
|
||||||
'edit_image': [src_image],
|
'src_image_list': [src_image],
|
||||||
'edit_image_mask': [src_mask],
|
'src_mask_list': [src_mask],
|
||||||
'image': tar_image,
|
'image': tar_image,
|
||||||
'image_mask': tar_mask,
|
'image_mask': tar_mask,
|
||||||
'prompt': [prompt],
|
'prompt': [prompt],
|
||||||
|
'edit_id': [0]
|
||||||
}
|
}
|
||||||
|
|
||||||
def load_image(self, prefix, img_path, cvt_type=None):
|
def load_image(self, prefix, img_path, cvt_type=None):
|
||||||
|
|||||||
@@ -304,9 +304,10 @@ class DataObject(object):
|
|||||||
delimiter = sampler_config.get('DELIMITER', ',')
|
delimiter = sampler_config.get('DELIMITER', ',')
|
||||||
path_prefix = sampler_config.get('PATH_PREFIX', '')
|
path_prefix = sampler_config.get('PATH_PREFIX', '')
|
||||||
prompt_prefix = sampler_config.get('PROMPT_PREFIX', '')
|
prompt_prefix = sampler_config.get('PROMPT_PREFIX', '')
|
||||||
|
oss_prefix = sampler_config.get('OSS_PREFIX', '')
|
||||||
return MultiLevelBatchSampler(batch_size, index_file, image_size,
|
return MultiLevelBatchSampler(batch_size, index_file, image_size,
|
||||||
fields, delimiter, path_prefix,
|
fields, delimiter, path_prefix,
|
||||||
prompt_prefix, rank, seed)
|
prompt_prefix, oss_prefix, rank, seed)
|
||||||
|
|
||||||
|
|
||||||
def build_dataset_config(cfg, registry, logger=None, *args, **kwargs):
|
def build_dataset_config(cfg, registry, logger=None, *args, **kwargs):
|
||||||
@@ -337,8 +338,12 @@ def build_dataset_config(cfg, registry, logger=None, *args, **kwargs):
|
|||||||
f'registry must be type Registry, got {type(registry)}')
|
f'registry must be type Registry, got {type(registry)}')
|
||||||
|
|
||||||
cfg = deep_copy(cfg)
|
cfg = deep_copy(cfg)
|
||||||
|
|
||||||
req_type = cfg.get('NAME')
|
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):
|
if isinstance(req_type, str):
|
||||||
req_type_entry = registry.get(req_type)
|
req_type_entry = registry.get(req_type)
|
||||||
if req_type_entry is None:
|
if req_type_entry is None:
|
||||||
|
|||||||
@@ -1,44 +1,46 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
import io
|
import io
|
||||||
|
import os
|
||||||
import random
|
import random
|
||||||
import sys
|
import sys
|
||||||
import os
|
|
||||||
import warnings
|
import warnings
|
||||||
|
|
||||||
import torch
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
import torch
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
from scepter.modules.utils.distribute import we
|
|
||||||
from scepter.modules.data.dataset import DATASETS, BaseDataset
|
from scepter.modules.data.dataset import DATASETS, BaseDataset
|
||||||
|
from scepter.modules.utils.distribute import we
|
||||||
from scepter.modules.utils.file_system import FS
|
from scepter.modules.utils.file_system import FS
|
||||||
|
|
||||||
try:
|
try:
|
||||||
import decord
|
import decord
|
||||||
decord.bridge.set_bridge("torch")
|
decord.bridge.set_bridge('torch')
|
||||||
except ImportError:
|
except ImportError:
|
||||||
warnings.warn(
|
warnings.warn(
|
||||||
"The `decord` package is required for loading the video dataset. Install with `pip install decord`"
|
'The `decord` package is required for loading the video dataset. Install with `pip install decord`'
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@DATASETS.register_class()
|
@DATASETS.register_class()
|
||||||
class VideoGenDataset(BaseDataset):
|
class VideoGenDataset(BaseDataset):
|
||||||
def __init__(self, cfg, logger = None):
|
def __init__(self, cfg, logger=None):
|
||||||
super().__init__(cfg, logger=logger)
|
super().__init__(cfg, logger=logger)
|
||||||
self.prompt_prefix = cfg.get('PROMPT_PREFIX', '')
|
self.prompt_prefix = cfg.get('PROMPT_PREFIX', '')
|
||||||
self.path_prefix = cfg.get('PATH_PREFIX', '')
|
self.path_prefix = cfg.get('PATH_PREFIX', '')
|
||||||
self.p_zero = cfg.get('P_ZERO', 0.0)
|
self.p_zero = cfg.get('P_ZERO', 0.0)
|
||||||
self.max_num_frames = cfg.get("NUM_FRAMES", 49)
|
self.max_num_frames = cfg.get('NUM_FRAMES', 49)
|
||||||
self.fps = cfg.get("FPS", 8)
|
self.fps = cfg.get('FPS', 8)
|
||||||
self.height = cfg.get("HEIGHT", 480)
|
self.height = cfg.get('HEIGHT', 480)
|
||||||
self.width = cfg.get("WIDTH", 720)
|
self.width = cfg.get('WIDTH', 720)
|
||||||
self.skip_frames_start = cfg.get("SKIP_FRAMES_START", 0)
|
self.skip_frames_start = cfg.get('SKIP_FRAMES_START', 0)
|
||||||
self.skip_frames_end = cfg.get("SKIP_FRAMES_END", 0)
|
self.skip_frames_end = cfg.get('SKIP_FRAMES_END', 0)
|
||||||
self.data_type = cfg.get('DATA_TYPE', 't2v')
|
self.data_type = cfg.get('DATA_TYPE', 't2v')
|
||||||
|
|
||||||
def worker_init_fn(self, worker_id, num_workers=1):
|
def worker_init_fn(self, worker_id, num_workers=1):
|
||||||
super().worker_init_fn(worker_id, num_workers=num_workers)
|
super().worker_init_fn(worker_id, num_workers=num_workers)
|
||||||
randseed = np.random.randint(0, 2 ** 32 - num_workers - 1)
|
randseed = np.random.randint(0, 2**32 - num_workers - 1)
|
||||||
workerseed = randseed + worker_id
|
workerseed = randseed + worker_id
|
||||||
random.seed(workerseed)
|
random.seed(workerseed)
|
||||||
np.random.seed(workerseed)
|
np.random.seed(workerseed)
|
||||||
@@ -46,7 +48,9 @@ class VideoGenDataset(BaseDataset):
|
|||||||
def _preprocess_video_data(self, video_path):
|
def _preprocess_video_data(self, video_path):
|
||||||
|
|
||||||
with FS.get_object(video_path) as video_data:
|
with FS.get_object(video_path) as video_data:
|
||||||
video_reader = decord.VideoReader(io.BytesIO(video_data), width=self.width, height=self.height)
|
video_reader = decord.VideoReader(io.BytesIO(video_data),
|
||||||
|
width=self.width,
|
||||||
|
height=self.height)
|
||||||
video_num_frames = len(video_reader)
|
video_num_frames = len(video_reader)
|
||||||
|
|
||||||
start_frame = min(self.skip_frames_start, video_num_frames)
|
start_frame = min(self.skip_frames_start, video_num_frames)
|
||||||
@@ -54,13 +58,16 @@ class VideoGenDataset(BaseDataset):
|
|||||||
if end_frame <= start_frame:
|
if end_frame <= start_frame:
|
||||||
frames = video_reader.get_batch([start_frame])
|
frames = video_reader.get_batch([start_frame])
|
||||||
elif end_frame - start_frame <= self.max_num_frames:
|
elif end_frame - start_frame <= self.max_num_frames:
|
||||||
frames = video_reader.get_batch(list(range(start_frame, end_frame)))
|
frames = video_reader.get_batch(list(range(start_frame,
|
||||||
|
end_frame)))
|
||||||
else:
|
else:
|
||||||
indices = list(range(start_frame, end_frame, (end_frame - start_frame) // self.max_num_frames))
|
indices = list(
|
||||||
|
range(start_frame, end_frame,
|
||||||
|
(end_frame - start_frame) // self.max_num_frames))
|
||||||
frames = video_reader.get_batch(indices)
|
frames = video_reader.get_batch(indices)
|
||||||
|
|
||||||
# Ensure that we don't go over the limit
|
# Ensure that we don't go over the limit
|
||||||
frames = frames[: self.max_num_frames]
|
frames = frames[:self.max_num_frames]
|
||||||
selected_num_frames = frames.shape[0]
|
selected_num_frames = frames.shape[0]
|
||||||
|
|
||||||
# Choose first (4k + 1) frames as this is how many is required by the VAE
|
# Choose first (4k + 1) frames as this is how many is required by the VAE
|
||||||
@@ -73,14 +80,16 @@ class VideoGenDataset(BaseDataset):
|
|||||||
|
|
||||||
# Training transforms
|
# Training transforms
|
||||||
frames = frames.float().div_(127.5).sub_(1.)
|
frames = frames.float().div_(127.5).sub_(1.)
|
||||||
frames = frames.permute(3, 0, 1, 2).contiguous() # [C, F, H, W]
|
frames = frames.permute(3, 0, 1, 2).contiguous() # [C, F, H, W]
|
||||||
return frames
|
return frames
|
||||||
|
|
||||||
def _parse_index(self, index):
|
def _parse_index(self, index):
|
||||||
meta = dict()
|
meta = dict()
|
||||||
for key, value in zip(index[-1], index[:-1]):
|
for key, value in zip(index[-1], index[:-1]):
|
||||||
if key in ['oss_key', 'path', 'video_path']:
|
if key in ['oss_key', 'path', 'video_path', 'target_video_path']:
|
||||||
meta['video_path'] = value
|
meta['video_path'] = value
|
||||||
|
elif key in ['source_video_path', 'src_video_path']:
|
||||||
|
meta['src_video_path'] = value
|
||||||
elif key in ['prompt', 'caption', 'text']:
|
elif key in ['prompt', 'caption', 'text']:
|
||||||
meta['prompt'] = value
|
meta['prompt'] = value
|
||||||
elif key in ['width', 'height']:
|
elif key in ['width', 'height']:
|
||||||
@@ -104,8 +113,13 @@ class VideoGenDataset(BaseDataset):
|
|||||||
'prompt': prompt,
|
'prompt': prompt,
|
||||||
'meta': meta,
|
'meta': meta,
|
||||||
}
|
}
|
||||||
if self.data_type == 'i2v':
|
if 'i2v' in self.data_type:
|
||||||
item['image'] = item['video'][:, :1, :, :]
|
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
|
return item
|
||||||
|
|
||||||
def __len__(self):
|
def __len__(self):
|
||||||
@@ -122,7 +136,6 @@ class VideoGenDataset(BaseDataset):
|
|||||||
return collect
|
return collect
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
@DATASETS.register_class()
|
@DATASETS.register_class()
|
||||||
class VideoGenDatasetOTF(VideoGenDataset):
|
class VideoGenDatasetOTF(VideoGenDataset):
|
||||||
def __init__(self, cfg, logger=None):
|
def __init__(self, cfg, logger=None):
|
||||||
@@ -135,8 +148,11 @@ class VideoGenDatasetOTF(VideoGenDataset):
|
|||||||
from scepter.modules.model.registry import MODELS
|
from scepter.modules.model.registry import MODELS
|
||||||
model_cfg = cfg.get('MODEL', None)
|
model_cfg = cfg.get('MODEL', None)
|
||||||
if model_cfg is not None:
|
if model_cfg is not None:
|
||||||
self.model = MODELS.build(cfg.MODEL, logger=logger).eval().requires_grad_(False).to(we.device_id)
|
self.model = MODELS.build(
|
||||||
self.items = self.parse_data(self.data_file, self.delimiter, self.fields)
|
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:
|
if self.use_num and self.use_num > 0:
|
||||||
self.items = self.items[:self.use_num]
|
self.items = self.items[:self.use_num]
|
||||||
self.data = self.encode(self.items)
|
self.data = self.encode(self.items)
|
||||||
@@ -169,11 +185,13 @@ class VideoGenDatasetOTF(VideoGenDataset):
|
|||||||
return items
|
return items
|
||||||
|
|
||||||
def encode(self, items):
|
def encode(self, items):
|
||||||
self.logger.info("Start to encode video data [{}]!".format(len(items)))
|
self.logger.info('Start to encode video data [{}]!'.format(len(items)))
|
||||||
for item in tqdm(items):
|
for item in tqdm(items):
|
||||||
video_path = os.path.join(self.path_prefix, item.get('video_path', ''))
|
video_path = os.path.join(self.path_prefix,
|
||||||
|
item.get('video_path', ''))
|
||||||
video = self._preprocess_video_data(video_path)
|
video = self._preprocess_video_data(video_path)
|
||||||
latent = self.model.encode_first_stage(video.unsqueeze(0).to(we.device_id)).squeeze(0)
|
latent = self.model.encode_first_stage(
|
||||||
|
video.unsqueeze(0).to(we.device_id)).squeeze(0)
|
||||||
item['video_latent'] = latent.detach().cpu()
|
item['video_latent'] = latent.detach().cpu()
|
||||||
item['video'] = video
|
item['video'] = video
|
||||||
if self.data_type == 'i2v':
|
if self.data_type == 'i2v':
|
||||||
@@ -181,4 +199,4 @@ class VideoGenDatasetOTF(VideoGenDataset):
|
|||||||
return items
|
return items
|
||||||
|
|
||||||
def _get(self, index):
|
def _get(self, index):
|
||||||
return self.data[index % self.real_number]
|
return self.data[index % self.real_number]
|
||||||
|
|||||||
@@ -1,9 +1,31 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# 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
|
if TYPE_CHECKING:
|
||||||
from scepter.modules.data.sampler.sampler import (
|
from scepter.modules.data.sampler.base_sampler import BaseSampler
|
||||||
EvalDistributedSampler, LoopSampler, MixtureOfSamplers,
|
from scepter.modules.data.sampler.registry import SAMPLERS
|
||||||
MultiFoldDistributedSampler, MultiLevelBatchSampler,
|
from scepter.modules.data.sampler.sampler import (
|
||||||
MultiLevelBatchSamplerMultiSource, ResolutionBatchSampler)
|
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)}')
|
f'registry must be type Registry, got {type(registry)}')
|
||||||
|
|
||||||
cfg = deep_copy(cfg)
|
cfg = deep_copy(cfg)
|
||||||
|
|
||||||
req_type = cfg.get('NAME')
|
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):
|
if isinstance(req_type, str):
|
||||||
req_type_entry = registry.get(req_type)
|
req_type_entry = registry.get(req_type)
|
||||||
if req_type_entry is None:
|
if req_type_entry is None:
|
||||||
|
|||||||
@@ -79,6 +79,7 @@ class MultiLevelBatchSamplerMultiSource(BaseSampler):
|
|||||||
self.num_fields = len(self.fields)
|
self.num_fields = len(self.fields)
|
||||||
self.delimiter = cfg.get('DELIMITER', ',')
|
self.delimiter = cfg.get('DELIMITER', ',')
|
||||||
self.path_prefix = cfg.get('PATH_PREFIX', '')
|
self.path_prefix = cfg.get('PATH_PREFIX', '')
|
||||||
|
oss_prefix = cfg.get('OSS_PREFIX', '')
|
||||||
common_prob = cfg.get('PROB', 1)
|
common_prob = cfg.get('PROB', 1)
|
||||||
sub_data_weights = cfg.get('SUB_DATA_WEIGHTS', None)
|
sub_data_weights = cfg.get('SUB_DATA_WEIGHTS', None)
|
||||||
sub_data_weights = {} if sub_data_weights is None else sub_data_weights.get_dict(
|
sub_data_weights = {} if sub_data_weights is None else sub_data_weights.get_dict(
|
||||||
@@ -137,7 +138,7 @@ class MultiLevelBatchSamplerMultiSource(BaseSampler):
|
|||||||
f"{p * common_prob} and samples'num: {sub_data['total']} in this cluster."
|
f"{p * common_prob} and samples'num: {sub_data['total']} in this cluster."
|
||||||
)
|
)
|
||||||
self.rng = np.random.default_rng(self.seed + we.rank)
|
self.rng = np.random.default_rng(self.seed + we.rank)
|
||||||
self.oss_prefix = '/'.join(index_file.split('/')[:3])
|
self.oss_prefix = '/'.join(index_file.split('/')[:3]) if (oss_prefix is None or oss_prefix == '') and index_file.startswith('oss') else oss_prefix
|
||||||
self.index_dir = os.path.dirname(index_file)
|
self.index_dir = os.path.dirname(index_file)
|
||||||
|
|
||||||
def __iter__(self):
|
def __iter__(self):
|
||||||
@@ -434,6 +435,7 @@ class MultiLevelBatchSampler(BaseSampler):
|
|||||||
delimiter=',',
|
delimiter=',',
|
||||||
path_prefix='',
|
path_prefix='',
|
||||||
prompt_prefix='',
|
prompt_prefix='',
|
||||||
|
oss_prefix='',
|
||||||
rank=0,
|
rank=0,
|
||||||
seed=8888):
|
seed=8888):
|
||||||
self.batch_size = batch_size
|
self.batch_size = batch_size
|
||||||
@@ -457,7 +459,7 @@ class MultiLevelBatchSampler(BaseSampler):
|
|||||||
'index_level': 1,
|
'index_level': 1,
|
||||||
'num_fields': self.num_fields
|
'num_fields': self.num_fields
|
||||||
}
|
}
|
||||||
self.oss_prefix = '/'.join(index_file.split('/')[:3])
|
self.oss_prefix = '/'.join(index_file.split('/')[:3]) if (oss_prefix is None or oss_prefix == '') and index_file.startswith('oss') else oss_prefix
|
||||||
self.index_dir = os.path.dirname(index_file)
|
self.index_dir = os.path.dirname(index_file)
|
||||||
|
|
||||||
def __iter__(self):
|
def __iter__(self):
|
||||||
|
|||||||
@@ -1,4 +1,21 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# 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 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# 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={},
|
||||||
|
)
|
||||||
|
|||||||
@@ -11,7 +11,6 @@ import torch
|
|||||||
from scepter.modules.utils.file_system import FS
|
from scepter.modules.utils.file_system import FS
|
||||||
from scepter.modules.utils.distribute import we
|
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 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 .diffusion_inference import DiffusionInference, get_model
|
||||||
from .tuner_inference import TunerInference
|
from .tuner_inference import TunerInference
|
||||||
|
|
||||||
@@ -27,9 +26,13 @@ class CogVideoXInference(DiffusionInference):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def decode_first_stage(self, latents):
|
def decode_first_stage(self, latents):
|
||||||
latents = latents.permute(0, 2, 1, 3, 4)
|
_, dtype = self.get_function_info(self.first_stage_model, 'decode')
|
||||||
latents = 1 / self.first_stage_model['paras']['scaling_factor_image'] * latents
|
with torch.autocast('cuda',
|
||||||
frames = get_model(self.first_stage_model).decode(latents)
|
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
|
return frames
|
||||||
|
|
||||||
def _prepare_rotary_positional_embeddings(
|
def _prepare_rotary_positional_embeddings(
|
||||||
@@ -121,8 +124,9 @@ class CogVideoXInference(DiffusionInference):
|
|||||||
)
|
)
|
||||||
function_name, dtype = self.get_function_info(
|
function_name, dtype = self.get_function_info(
|
||||||
self.diffusion_model)
|
self.diffusion_model)
|
||||||
|
|
||||||
with torch.autocast('cuda',
|
with torch.autocast('cuda',
|
||||||
enabled=dtype=='bfloat16',
|
enabled=dtype in ('float16', 'bfloat16'),
|
||||||
dtype=getattr(torch, dtype)):
|
dtype=getattr(torch, dtype)):
|
||||||
solver_sample = value_input.get('sample', 'ddim')
|
solver_sample = value_input.get('sample', 'ddim')
|
||||||
sample_steps = value_input.get('sample_steps', 50)
|
sample_steps = value_input.get('sample_steps', 50)
|
||||||
@@ -143,7 +147,6 @@ class CogVideoXInference(DiffusionInference):
|
|||||||
}],
|
}],
|
||||||
steps=sample_steps,
|
steps=sample_steps,
|
||||||
show_progress=True,
|
show_progress=True,
|
||||||
use_dynamic_cfg=True,
|
|
||||||
guide_scale=guide_scale,
|
guide_scale=guide_scale,
|
||||||
guide_rescale=guide_rescale,
|
guide_rescale=guide_rescale,
|
||||||
return_intermediate=None,
|
return_intermediate=None,
|
||||||
@@ -151,7 +154,6 @@ class CogVideoXInference(DiffusionInference):
|
|||||||
self.dynamic_unload(self.diffusion_model,
|
self.dynamic_unload(self.diffusion_model,
|
||||||
'diffusion_model',
|
'diffusion_model',
|
||||||
skip_loaded=True)
|
skip_loaded=True)
|
||||||
|
|
||||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||||
x_samples = self.decode_first_stage(latent).float() # [B, C, F, H, W]
|
x_samples = self.decode_first_stage(latent).float() # [B, C, F, H, W]
|
||||||
self.dynamic_unload(self.first_stage_model,
|
self.dynamic_unload(self.first_stage_model,
|
||||||
|
|||||||
@@ -97,7 +97,7 @@ class DiffusionInference():
|
|||||||
if 'weights_only' in torch.load.__code__.co_varnames:
|
if 'weights_only' in torch.load.__code__.co_varnames:
|
||||||
sd = torch.load(local_path, map_location='cpu', weights_only=True)
|
sd = torch.load(local_path, map_location='cpu', weights_only=True)
|
||||||
else:
|
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(
|
first_stage_model_path = os.path.join(
|
||||||
os.path.dirname(local_path), 'first_stage_model.pth')
|
os.path.dirname(local_path), 'first_stage_model.pth')
|
||||||
cond_stage_model_path = os.path.join(
|
cond_stage_model_path = os.path.join(
|
||||||
@@ -203,7 +203,7 @@ class DiffusionInference():
|
|||||||
from safetensors.torch import load_file as load_safetensors
|
from safetensors.torch import load_file as load_safetensors
|
||||||
sd = load_safetensors(path)
|
sd = load_safetensors(path)
|
||||||
else:
|
else:
|
||||||
sd = torch.load(path, map_location='cpu')
|
sd = torch.load(path, map_location='cpu', weights_only=True)
|
||||||
|
|
||||||
new_sd = OrderedDict()
|
new_sd = OrderedDict()
|
||||||
for k, v in sd.items():
|
for k, v in sd.items():
|
||||||
@@ -230,16 +230,22 @@ class DiffusionInference():
|
|||||||
|
|
||||||
def load(self, module):
|
def load(self, module):
|
||||||
if module['device'] == 'offline':
|
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()
|
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'],
|
model = BACKBONES.build(module['cfg'],
|
||||||
logger=self.logger).eval()
|
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'],
|
model = EMBEDDERS.build(module['cfg'],
|
||||||
logger=self.logger).eval()
|
logger=self.logger).eval()
|
||||||
else:
|
else:
|
||||||
raise NotImplementedError
|
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):
|
if module['cfg'].get('RELOAD_MODEL', None):
|
||||||
self.init_from_ckpt(module['cfg'].RELOAD_MODEL, model)
|
self.init_from_ckpt(module['cfg'].RELOAD_MODEL, model)
|
||||||
module['model'] = model
|
module['model'] = model
|
||||||
@@ -268,8 +274,9 @@ class DiffusionInference():
|
|||||||
module['device'] = 'cpu'
|
module['device'] = 'cpu'
|
||||||
else:
|
else:
|
||||||
module['device'] = 'offline'
|
module['device'] = 'offline'
|
||||||
torch.cuda.empty_cache()
|
if torch.cuda.is_available():
|
||||||
torch.cuda.ipc_collect()
|
torch.cuda.empty_cache()
|
||||||
|
torch.cuda.ipc_collect()
|
||||||
return module
|
return module
|
||||||
|
|
||||||
def dynamic_load(self, module=None, name=''):
|
def dynamic_load(self, module=None, name=''):
|
||||||
|
|||||||
@@ -43,7 +43,7 @@ class LargenInference(DiffusionInference):
|
|||||||
if 'weights_only' in torch.load.__code__.co_varnames:
|
if 'weights_only' in torch.load.__code__.co_varnames:
|
||||||
sd = torch.load(local_path, map_location='cpu', weights_only=True)
|
sd = torch.load(local_path, map_location='cpu', weights_only=True)
|
||||||
else:
|
else:
|
||||||
sd = torch.load(local_path, map_location='cpu')
|
sd = torch.load(local_path, map_location='cpu', weights_only=True)
|
||||||
if 'model' in sd:
|
if 'model' in sd:
|
||||||
sd = sd['model']
|
sd = sd['model']
|
||||||
|
|
||||||
|
|||||||
@@ -144,9 +144,9 @@ class TunerInference():
|
|||||||
is_bin_file = True
|
is_bin_file = True
|
||||||
if os.path.isfile(bin_file):
|
if os.path.isfile(bin_file):
|
||||||
if 'weights_only' in torch.load.__code__.co_varnames:
|
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:
|
else:
|
||||||
state_dict = torch.load(bin_file)
|
state_dict = torch.load(bin_file, map_location="cpu")
|
||||||
elif os.path.isfile(safe_file):
|
elif os.path.isfile(safe_file):
|
||||||
is_bin_file = False
|
is_bin_file = False
|
||||||
from safetensors.torch import \
|
from safetensors.torch import \
|
||||||
|
|||||||
@@ -1,5 +1,23 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# 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 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
from scepter.modules.model.backbone import (ace, autoencoder, flux, image, cogvideox,
|
from typing import TYPE_CHECKING
|
||||||
mmdit, pixart, unet, utils, video)
|
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):
|
def load_pretrained_model(self, pretrained_model):
|
||||||
if pretrained_model:
|
if pretrained_model:
|
||||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
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:
|
if 'state_dict' in model:
|
||||||
model = model['state_dict']
|
model = model['state_dict']
|
||||||
new_ckpt = OrderedDict()
|
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)
|
||||||
@@ -48,6 +48,8 @@ class CogVideoXTransformer3DModel(BaseModel):
|
|||||||
Whether to flip the sin to cos in the time embedding.
|
Whether to flip the sin to cos in the time embedding.
|
||||||
time_embed_dim (`int`, defaults to `512`):
|
time_embed_dim (`int`, defaults to `512`):
|
||||||
Output dimension of timestep embeddings.
|
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`):
|
text_embed_dim (`int`, defaults to `4096`):
|
||||||
Input dimension of text embeddings from the text encoder.
|
Input dimension of text embeddings from the text encoder.
|
||||||
num_layers (`int`, defaults to `30`):
|
num_layers (`int`, defaults to `30`):
|
||||||
@@ -98,6 +100,7 @@ class CogVideoXTransformer3DModel(BaseModel):
|
|||||||
flip_sin_to_cos = cfg.get("FLIP_SIN_TO_COS", True)
|
flip_sin_to_cos = cfg.get("FLIP_SIN_TO_COS", True)
|
||||||
freq_shift = cfg.get("FREQ_SHIFT", 0)
|
freq_shift = cfg.get("FREQ_SHIFT", 0)
|
||||||
time_embed_dim = cfg.get("TIME_EMBED_DIM", 512)
|
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)
|
text_embed_dim = cfg.get("TEXT_EMBED_DIM", 4096)
|
||||||
num_layers = cfg.get("NUM_LAYERS", 30)
|
num_layers = cfg.get("NUM_LAYERS", 30)
|
||||||
dropout = cfg.get("DROPOUT", 0.0)
|
dropout = cfg.get("DROPOUT", 0.0)
|
||||||
@@ -106,6 +109,8 @@ class CogVideoXTransformer3DModel(BaseModel):
|
|||||||
sample_height = cfg.get("SAMPLE_HEIGHT", 60)
|
sample_height = cfg.get("SAMPLE_HEIGHT", 60)
|
||||||
sample_frames = cfg.get("SAMPLE_FRAMES", 49)
|
sample_frames = cfg.get("SAMPLE_FRAMES", 49)
|
||||||
patch_size = cfg.get("PATCH_SIZE", 2)
|
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)
|
temporal_compression_ratio = cfg.get("TEMPORAL_COMPRESSION_RATIO", 4)
|
||||||
max_text_seq_length = cfg.get("MAX_TEXT_SEQ_LENGTH", 226)
|
max_text_seq_length = cfg.get("MAX_TEXT_SEQ_LENGTH", 226)
|
||||||
activation_fn = cfg.get("ACTIVATION_FN", "gelu-approximate")
|
activation_fn = cfg.get("ACTIVATION_FN", "gelu-approximate")
|
||||||
@@ -119,6 +124,7 @@ class CogVideoXTransformer3DModel(BaseModel):
|
|||||||
self.gradient_checkpointing = cfg.get("GRADIENT_CHECKPOINTING", False)
|
self.gradient_checkpointing = cfg.get("GRADIENT_CHECKPOINTING", False)
|
||||||
inner_dim = num_attention_heads * attention_head_dim
|
inner_dim = num_attention_heads * attention_head_dim
|
||||||
self.patch_size = patch_size
|
self.patch_size = patch_size
|
||||||
|
self.patch_size_t = patch_size_t
|
||||||
self.use_rotary_positional_embeddings = use_rotary_positional_embeddings
|
self.use_rotary_positional_embeddings = use_rotary_positional_embeddings
|
||||||
|
|
||||||
if not use_rotary_positional_embeddings and use_learned_positional_embeddings:
|
if not use_rotary_positional_embeddings and use_learned_positional_embeddings:
|
||||||
@@ -131,10 +137,11 @@ class CogVideoXTransformer3DModel(BaseModel):
|
|||||||
# 1. Patch embedding
|
# 1. Patch embedding
|
||||||
self.patch_embed = CogVideoXPatchEmbed(
|
self.patch_embed = CogVideoXPatchEmbed(
|
||||||
patch_size=patch_size,
|
patch_size=patch_size,
|
||||||
|
patch_size_t=patch_size_t,
|
||||||
in_channels=in_channels,
|
in_channels=in_channels,
|
||||||
embed_dim=inner_dim,
|
embed_dim=inner_dim,
|
||||||
text_embed_dim=text_embed_dim,
|
text_embed_dim=text_embed_dim,
|
||||||
bias=True,
|
bias=patch_bias,
|
||||||
sample_width=sample_width,
|
sample_width=sample_width,
|
||||||
sample_height=sample_height,
|
sample_height=sample_height,
|
||||||
sample_frames=sample_frames,
|
sample_frames=sample_frames,
|
||||||
@@ -147,10 +154,18 @@ class CogVideoXTransformer3DModel(BaseModel):
|
|||||||
)
|
)
|
||||||
self.embedding_dropout = nn.Dropout(dropout)
|
self.embedding_dropout = nn.Dropout(dropout)
|
||||||
|
|
||||||
# 2. Time embeddings
|
# 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_proj = Timesteps(inner_dim, flip_sin_to_cos, freq_shift)
|
||||||
self.time_embedding = TimestepEmbedding(inner_dim, time_embed_dim, timestep_activation_fn)
|
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
|
# 3. Define spatio-temporal transformers blocks
|
||||||
self.transformer_blocks = nn.ModuleList(
|
self.transformer_blocks = nn.ModuleList(
|
||||||
[
|
[
|
||||||
@@ -178,7 +193,15 @@ class CogVideoXTransformer3DModel(BaseModel):
|
|||||||
norm_eps=norm_eps,
|
norm_eps=norm_eps,
|
||||||
chunk_dim=1,
|
chunk_dim=1,
|
||||||
)
|
)
|
||||||
self.proj_out = nn.Linear(inner_dim, patch_size * patch_size * out_channels)
|
|
||||||
|
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(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -186,6 +209,7 @@ class CogVideoXTransformer3DModel(BaseModel):
|
|||||||
t: Union[int, float, torch.LongTensor] = None,
|
t: Union[int, float, torch.LongTensor] = None,
|
||||||
cond: torch.Tensor = None,
|
cond: torch.Tensor = None,
|
||||||
timestep_cond: Optional[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,
|
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||||
**kwargs
|
**kwargs
|
||||||
):
|
):
|
||||||
@@ -208,6 +232,12 @@ class CogVideoXTransformer3DModel(BaseModel):
|
|||||||
t_emb = t_emb.to(dtype=encoder_hidden_states.dtype)
|
t_emb = t_emb.to(dtype=encoder_hidden_states.dtype)
|
||||||
emb = self.time_embedding(t_emb, timestep_cond)
|
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
|
# 2. Patch embedding
|
||||||
hidden_states = self.patch_embed(encoder_hidden_states, hidden_states)
|
hidden_states = self.patch_embed(encoder_hidden_states, hidden_states)
|
||||||
hidden_states = self.embedding_dropout(hidden_states)
|
hidden_states = self.embedding_dropout(hidden_states)
|
||||||
@@ -261,8 +291,16 @@ class CogVideoXTransformer3DModel(BaseModel):
|
|||||||
# - It is okay to `channels` use for CogVideoX-2b and CogVideoX-5b (number of input channels is equal to output 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)
|
# - 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 = self.patch_size
|
||||||
output = hidden_states.reshape(batch_size, num_frames, height // p, width // p, -1, p, p)
|
p_t = self.patch_size_t
|
||||||
output = output.permute(0, 1, 4, 2, 5, 3, 6).flatten(5, 6).flatten(3, 4)
|
|
||||||
|
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
|
return output
|
||||||
|
|
||||||
@@ -277,7 +315,7 @@ class CogVideoXTransformer3DModel(BaseModel):
|
|||||||
from safetensors.torch import load_file as load_safetensors
|
from safetensors.torch import load_file as load_safetensors
|
||||||
ckpt = load_safetensors(local_model)
|
ckpt = load_safetensors(local_model)
|
||||||
else:
|
else:
|
||||||
ckpt = torch.load(local_model, map_location='cpu')
|
ckpt = torch.load(local_model, map_location='cpu', weights_only=True)
|
||||||
ckpt_all.update(ckpt)
|
ckpt_all.update(ckpt)
|
||||||
missing, unexpected = self.load_state_dict(ckpt_all, strict=False)
|
missing, unexpected = self.load_state_dict(ckpt_all, strict=False)
|
||||||
if we.rank == 0:
|
if we.rank == 0:
|
||||||
@@ -309,9 +347,9 @@ if __name__ == "__main__":
|
|||||||
FS.init_fs_client(file_sys)
|
FS.init_fs_client(file_sys)
|
||||||
model = BACKBONES.build(cfg.DIFFUSION_MODEL, logger=get_logger()).eval().requires_grad_(False).to('cuda').to(torch.bfloat16)
|
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))
|
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))
|
encoder_hidden_states = torch.load(FS.get_from(cfg.ENCODER_HIDDEN_STATES), weights_only=True)
|
||||||
timestep = torch.load(FS.get_from(cfg.TIMESTEP))
|
timestep = torch.load(FS.get_from(cfg.TIMESTEP), weights_only=True)
|
||||||
timestep_cond = None
|
timestep_cond = None
|
||||||
image_rotary_emb = None
|
image_rotary_emb = None
|
||||||
attention_kwargs = None
|
attention_kwargs = None
|
||||||
|
|||||||
@@ -178,6 +178,7 @@ class CogVideoXPatchEmbed(nn.Module):
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
patch_size: int = 2,
|
patch_size: int = 2,
|
||||||
|
patch_size_t: Optional[int] = None,
|
||||||
in_channels: int = 16,
|
in_channels: int = 16,
|
||||||
embed_dim: int = 1920,
|
embed_dim: int = 1920,
|
||||||
text_embed_dim: int = 4096,
|
text_embed_dim: int = 4096,
|
||||||
@@ -195,6 +196,7 @@ class CogVideoXPatchEmbed(nn.Module):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
self.patch_size = patch_size
|
self.patch_size = patch_size
|
||||||
|
self.patch_size_t = patch_size_t
|
||||||
self.embed_dim = embed_dim
|
self.embed_dim = embed_dim
|
||||||
self.sample_height = sample_height
|
self.sample_height = sample_height
|
||||||
self.sample_width = sample_width
|
self.sample_width = sample_width
|
||||||
@@ -206,9 +208,15 @@ class CogVideoXPatchEmbed(nn.Module):
|
|||||||
self.use_positional_embeddings = use_positional_embeddings
|
self.use_positional_embeddings = use_positional_embeddings
|
||||||
self.use_learned_positional_embeddings = use_learned_positional_embeddings
|
self.use_learned_positional_embeddings = use_learned_positional_embeddings
|
||||||
|
|
||||||
self.proj = nn.Conv2d(
|
if patch_size_t is None:
|
||||||
in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size, bias=bias
|
# 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)
|
self.text_proj = nn.Linear(text_embed_dim, embed_dim)
|
||||||
|
|
||||||
if use_positional_embeddings or use_learned_positional_embeddings:
|
if use_positional_embeddings or use_learned_positional_embeddings:
|
||||||
@@ -247,12 +255,24 @@ class CogVideoXPatchEmbed(nn.Module):
|
|||||||
"""
|
"""
|
||||||
text_embeds = self.text_proj(text_embeds)
|
text_embeds = self.text_proj(text_embeds)
|
||||||
|
|
||||||
batch, num_frames, channels, height, width = image_embeds.shape
|
batch_size, num_frames, channels, height, width = image_embeds.shape
|
||||||
image_embeds = image_embeds.reshape(-1, channels, height, width)
|
|
||||||
image_embeds = self.proj(image_embeds)
|
if self.patch_size_t is None:
|
||||||
image_embeds = image_embeds.view(batch, num_frames, *image_embeds.shape[1:])
|
image_embeds = image_embeds.reshape(-1, channels, height, width)
|
||||||
image_embeds = image_embeds.flatten(3).transpose(2, 3) # [batch, num_frames, height x width, channels]
|
image_embeds = self.proj(image_embeds)
|
||||||
image_embeds = image_embeds.flatten(1, 2) # [batch, num_frames x height x width, channels]
|
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(
|
embeds = torch.cat(
|
||||||
[text_embeds, image_embeds], dim=1
|
[text_embeds, image_embeds], dim=1
|
||||||
|
|||||||
@@ -459,7 +459,14 @@ def get_1d_rotary_pos_embed(
|
|||||||
|
|
||||||
|
|
||||||
def get_3d_rotary_pos_embed(
|
def get_3d_rotary_pos_embed(
|
||||||
embed_dim, crops_coords, grid_size, temporal_size, theta: int = 10000, use_real: bool = True
|
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]]:
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
"""
|
"""
|
||||||
RoPE for video tokens with 3D structure.
|
RoPE for video tokens with 3D structure.
|
||||||
@@ -475,17 +482,30 @@ def get_3d_rotary_pos_embed(
|
|||||||
The size of the temporal dimension.
|
The size of the temporal dimension.
|
||||||
theta (`float`):
|
theta (`float`):
|
||||||
Scaling factor for frequency computation.
|
Scaling factor for frequency computation.
|
||||||
|
grid_type (`str`):
|
||||||
|
Whether to use "linspace" or "slice" to compute grids.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
`torch.Tensor`: positional embedding with shape `(temporal_size * grid_size[0] * grid_size[1], embed_dim/2)`.
|
`torch.Tensor`: positional embedding with shape `(temporal_size * grid_size[0] * grid_size[1], embed_dim/2)`.
|
||||||
"""
|
"""
|
||||||
if use_real is not True:
|
if use_real is not True:
|
||||||
raise ValueError(" `use_real = False` is not currently supported for get_3d_rotary_pos_embed")
|
raise ValueError(" `use_real = False` is not currently supported for get_3d_rotary_pos_embed")
|
||||||
start, stop = crops_coords
|
|
||||||
grid_size_h, grid_size_w = grid_size
|
if grid_type == "linspace":
|
||||||
grid_h = np.linspace(start[0], stop[0], grid_size_h, endpoint=False, dtype=np.float32)
|
start, stop = crops_coords
|
||||||
grid_w = np.linspace(start[1], stop[1], grid_size_w, endpoint=False, dtype=np.float32)
|
grid_size_h, grid_size_w = grid_size
|
||||||
grid_t = np.linspace(0, temporal_size, temporal_size, endpoint=False, dtype=np.float32)
|
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
|
# Compute dimensions for each axis
|
||||||
dim_t = embed_dim // 4
|
dim_t = embed_dim // 4
|
||||||
@@ -521,6 +541,12 @@ def get_3d_rotary_pos_embed(
|
|||||||
t_cos, t_sin = freqs_t # both t_cos and t_sin has shape: temporal_size, dim_t
|
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
|
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
|
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)
|
cos = combine_time_height_width(t_cos, h_cos, w_cos)
|
||||||
sin = combine_time_height_width(t_sin, h_sin, w_sin)
|
sin = combine_time_height_width(t_sin, h_sin, w_sin)
|
||||||
return cos, sin
|
return cos, sin
|
||||||
|
|||||||
@@ -1,3 +1,3 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# 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 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# 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
|
import math
|
||||||
|
from collections import OrderedDict
|
||||||
from functools import partial
|
from functools import partial
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -15,8 +18,6 @@ from torch.utils.checkpoint import checkpoint_sequential
|
|||||||
from torch.nn.utils.rnn import pad_sequence
|
from torch.nn.utils.rnn import pad_sequence
|
||||||
from .layers import (DoubleStreamBlock, EmbedND, LastLayer, MLPEmbedder,
|
from .layers import (DoubleStreamBlock, EmbedND, LastLayer, MLPEmbedder,
|
||||||
SingleStreamBlock, timestep_embedding)
|
SingleStreamBlock, timestep_embedding)
|
||||||
|
|
||||||
|
|
||||||
@BACKBONES.register_class()
|
@BACKBONES.register_class()
|
||||||
class Flux(BaseModel):
|
class Flux(BaseModel):
|
||||||
"""
|
"""
|
||||||
@@ -98,7 +99,14 @@ class Flux(BaseModel):
|
|||||||
qkv_bias = cfg.QKV_BIAS
|
qkv_bias = cfg.QKV_BIAS
|
||||||
depth = cfg.DEPTH
|
depth = cfg.DEPTH
|
||||||
depth_single_blocks = cfg.DEPTH_SINGLE_BLOCKS
|
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:
|
if hidden_size % num_heads != 0:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -119,85 +127,350 @@ class Flux(BaseModel):
|
|||||||
if self.guidance_embed else nn.Identity())
|
if self.guidance_embed else nn.Identity())
|
||||||
self.txt_in = nn.Linear(context_in_dim, self.hidden_size)
|
self.txt_in = nn.Linear(context_in_dim, self.hidden_size)
|
||||||
|
|
||||||
self.double_blocks = nn.ModuleList([
|
self.double_blocks = nn.ModuleList(
|
||||||
DoubleStreamBlock(
|
[
|
||||||
self.hidden_size,
|
DoubleStreamBlock(
|
||||||
self.num_heads,
|
self.hidden_size,
|
||||||
mlp_ratio=mlp_ratio,
|
self.num_heads,
|
||||||
qkv_bias=qkv_bias,
|
mlp_ratio=mlp_ratio,
|
||||||
) for _ in range(depth)
|
qkv_bias=qkv_bias,
|
||||||
])
|
backend=self.attn_backend
|
||||||
|
)
|
||||||
|
for _ in range(depth)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
self.single_blocks = nn.ModuleList([
|
self.single_blocks = nn.ModuleList(
|
||||||
SingleStreamBlock(self.hidden_size,
|
[
|
||||||
self.num_heads,
|
SingleStreamBlock(self.hidden_size, self.num_heads, mlp_ratio=mlp_ratio, backend=self.attn_backend)
|
||||||
mlp_ratio=mlp_ratio)
|
for _ in range(depth_single_blocks)
|
||||||
for _ in range(depth_single_blocks)
|
]
|
||||||
])
|
)
|
||||||
|
|
||||||
self.final_layer = LastLayer(self.hidden_size, 1, self.out_channels)
|
self.final_layer = LastLayer(self.hidden_size, 1, self.out_channels)
|
||||||
|
|
||||||
def prepare_input(self, x, context, y, x_shape=None):
|
def prepare_input(self, x, context, y, x_shape=None):
|
||||||
# x.shape [6, 16, 16, 16] target is [6, 16, 768, 1360]
|
# x.shape [6, 16, 16, 16] target is [6, 16, 768, 1360]
|
||||||
bs, c, h, w = x.shape
|
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 = torch.zeros(h // 2, w // 2, 3)
|
||||||
x_id[..., 1] = x_id[..., 1] + torch.arange(h // 2)[:, None]
|
x_id[..., 1] = x_id[..., 1] + torch.arange(h // 2)[:, None]
|
||||||
x_id[..., 2] = x_id[..., 2] + torch.arange(w // 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)
|
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
|
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:
|
def unpack(self, x: Tensor, height: int, width: int) -> Tensor:
|
||||||
return rearrange(
|
return rearrange(
|
||||||
x,
|
x,
|
||||||
'b (h w) (c ph pw) -> b c (h ph) (w pw)',
|
"b (h w) (c ph pw) -> b c (h ph) (w pw)",
|
||||||
h=math.ceil(height / 2),
|
h=math.ceil(height/2),
|
||||||
w=math.ceil(width / 2),
|
w=math.ceil(width/2),
|
||||||
ph=2,
|
ph=2,
|
||||||
pw=2,
|
pw=2,
|
||||||
)
|
)
|
||||||
|
|
||||||
def load_pretrained_model(self, pretrained_model):
|
def merge_diffuser_lora(self, ori_sd, lora_sd, scale=1.0):
|
||||||
if next(self.parameters()).device.type == 'meta':
|
key_map = {
|
||||||
map_location = we.device_id
|
"single_blocks.{}.linear1.weight": {"key_list": [
|
||||||
else:
|
["transformer.single_transformer_blocks.{}.attn.to_q.lora_A.weight",
|
||||||
map_location = 'cpu'
|
"transformer.single_transformer_blocks.{}.attn.to_q.lora_B.weight", [0, 3072]],
|
||||||
if pretrained_model is not None:
|
["transformer.single_transformer_blocks.{}.attn.to_k.lora_A.weight",
|
||||||
with FS.get_from(pretrained_model,
|
"transformer.single_transformer_blocks.{}.attn.to_k.lora_B.weight", [3072, 6144]],
|
||||||
wait_finish=True) as local_model:
|
["transformer.single_transformer_blocks.{}.attn.to_v.lora_A.weight",
|
||||||
if local_model.endswith('safetensors'):
|
"transformer.single_transformer_blocks.{}.attn.to_v.lora_B.weight", [6144, 9216]],
|
||||||
from safetensors.torch import load_file as load_safetensors
|
["transformer.single_transformer_blocks.{}.proj_mlp.lora_A.weight",
|
||||||
sd = load_safetensors(local_model, device=map_location)
|
"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:
|
else:
|
||||||
sd = torch.load(local_model, map_location=map_location)
|
print("unsurpport keys: ", k)
|
||||||
missing, unexpected = self.load_state_dict(sd,
|
self.logger.info(f"merge_blackforest_lora loads lora'parameters lora-paras: \n"
|
||||||
strict=False,
|
f"cover-{len(cover_lora_keys)} vs total {len(lora_sd)} \n"
|
||||||
assign=True)
|
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(
|
self.logger.info(
|
||||||
f'Restored from {pretrained_model} with {len(missing)} missing and {len(unexpected)} unexpected keys'
|
f'Restored from {pretrained_model} with {len(missing)} missing and {len(unexpected)} unexpected keys'
|
||||||
)
|
)
|
||||||
if len(missing) > 0:
|
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:
|
if len(unexpected) > 0:
|
||||||
self.logger.info(f'\nUnexpected Keys:\n {unexpected}') # noqa
|
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
|
||||||
|
|
||||||
def forward(self,
|
def forward(
|
||||||
x: Tensor,
|
self,
|
||||||
t: Tensor,
|
x: Tensor,
|
||||||
cond: dict = {},
|
t: Tensor,
|
||||||
guidance: Tensor | None = None,
|
cond: dict = {},
|
||||||
gc_seg: int = 0) -> Tensor:
|
guidance: Tensor | None = None,
|
||||||
x, x_ids, txt, txt_ids, y, h, w = self.prepare_input(
|
gc_seg: int = 0
|
||||||
x, cond['context'], cond['y'])
|
) -> Tensor:
|
||||||
|
x, x_ids, txt, txt_ids, y, h, w = self.prepare_input(x, cond["context"], cond["y"])
|
||||||
# running on sequences img
|
# running on sequences img
|
||||||
x = self.img_in(x)
|
x = self.img_in(x)
|
||||||
vec = self.time_in(timestep_embedding(t, 256))
|
vec = self.time_in(timestep_embedding(t, 256))
|
||||||
if self.guidance_embed:
|
if self.guidance_embed:
|
||||||
if guidance is None:
|
if guidance is None:
|
||||||
raise ValueError(
|
raise ValueError("Didn't get guidance strength for guidance distilled model.")
|
||||||
"Didn't get guidance strength for guidance distilled model."
|
|
||||||
)
|
|
||||||
vec = vec + self.guidance_in(timestep_embedding(guidance, 256))
|
vec = vec + self.guidance_in(timestep_embedding(guidance, 256))
|
||||||
vec = vec + self.vector_in(y)
|
vec = vec + self.vector_in(y)
|
||||||
txt = self.txt_in(txt)
|
txt = self.txt_in(txt)
|
||||||
@@ -211,12 +484,11 @@ class Flux(BaseModel):
|
|||||||
x = torch.cat((txt, x), 1)
|
x = torch.cat((txt, x), 1)
|
||||||
if self.use_grad_checkpoint and gc_seg >= 0:
|
if self.use_grad_checkpoint and gc_seg >= 0:
|
||||||
x = checkpoint_sequential(
|
x = checkpoint_sequential(
|
||||||
functions=[
|
functions=[partial(block, **kwargs) for block in self.double_blocks],
|
||||||
partial(block, **kwargs) for block in self.double_blocks
|
|
||||||
],
|
|
||||||
segments=gc_seg if gc_seg > 0 else len(self.double_blocks),
|
segments=gc_seg if gc_seg > 0 else len(self.double_blocks),
|
||||||
input=x,
|
input=x,
|
||||||
use_reentrant=False)
|
use_reentrant=False
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
for block in self.double_blocks:
|
for block in self.double_blocks:
|
||||||
x = block(x, **kwargs)
|
x = block(x, **kwargs)
|
||||||
@@ -228,18 +500,16 @@ class Flux(BaseModel):
|
|||||||
|
|
||||||
if self.use_grad_checkpoint and gc_seg >= 0:
|
if self.use_grad_checkpoint and gc_seg >= 0:
|
||||||
x = checkpoint_sequential(
|
x = checkpoint_sequential(
|
||||||
functions=[
|
functions=[partial(block, **kwargs) for block in self.single_blocks],
|
||||||
partial(block, **kwargs) for block in self.single_blocks
|
|
||||||
],
|
|
||||||
segments=gc_seg if gc_seg > 0 else len(self.single_blocks),
|
segments=gc_seg if gc_seg > 0 else len(self.single_blocks),
|
||||||
input=x,
|
input=x,
|
||||||
use_reentrant=False)
|
use_reentrant=False
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
for block in self.single_blocks:
|
for block in self.single_blocks:
|
||||||
x = block(x, **kwargs)
|
x = block(x, **kwargs)
|
||||||
x = x[:, txt.shape[1]:, ...]
|
x = x[:, txt.shape[1] :, ...]
|
||||||
x = self.final_layer(
|
x = self.final_layer(x, vec) # (N, T, patch_size ** 2 * out_channels) 6 64 64
|
||||||
x, vec) # (N, T, patch_size ** 2 * out_channels) 6 64 64
|
|
||||||
x = self.unpack(x, h, w)
|
x = self.unpack(x, h, w)
|
||||||
return x
|
return x
|
||||||
|
|
||||||
@@ -249,11 +519,13 @@ class Flux(BaseModel):
|
|||||||
__class__.__name__,
|
__class__.__name__,
|
||||||
Flux.para_dict,
|
Flux.para_dict,
|
||||||
set_name=True)
|
set_name=True)
|
||||||
|
|
||||||
@BACKBONES.register_class()
|
@BACKBONES.register_class()
|
||||||
class FluxMR(Flux):
|
class FluxMR(Flux):
|
||||||
def prepare_input(self, x, cond):
|
def prepare_input(self, x, cond):
|
||||||
context, y = cond["context"].to(x), cond["y"].to(x)
|
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 = [], []
|
batch_frames, batch_frames_ids = [], []
|
||||||
for ix, shape in zip(x, cond["x_shapes"]):
|
for ix, shape in zip(x, cond["x_shapes"]):
|
||||||
# unpack image from sequence
|
# unpack image from sequence
|
||||||
@@ -319,7 +591,7 @@ class FluxMR(Flux):
|
|||||||
x, x_ids, txt, txt_ids, y, mask_x, mask_txt, seq_length_list = self.prepare_input(x, cond)
|
x, x_ids, txt, txt_ids, y, mask_x, mask_txt, seq_length_list = self.prepare_input(x, cond)
|
||||||
# running on sequences img
|
# running on sequences img
|
||||||
vec = self.time_in(timestep_embedding(t, 256))
|
vec = self.time_in(timestep_embedding(t, 256))
|
||||||
if self.guidance_embed:
|
if self.guidance_embed and guidance[-1] >= 0:
|
||||||
if guidance is None:
|
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.guidance_in(timestep_embedding(guidance, 256))
|
||||||
@@ -371,7 +643,170 @@ class FluxMR(Flux):
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def get_config_template():
|
def get_config_template():
|
||||||
return dict_to_yaml('BACKBONE',
|
return dict_to_yaml('MODEL',
|
||||||
__class__.__name__,
|
__class__.__name__,
|
||||||
FluxMR.para_dict,
|
FluxMR.para_dict,
|
||||||
set_name=True)
|
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,5 +1,7 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# 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
|
from __future__ import annotations
|
||||||
|
|
||||||
import math
|
import math
|
||||||
@@ -351,7 +353,7 @@ class SingleStreamBlock(nn.Module):
|
|||||||
if mask is not None:
|
if mask is not None:
|
||||||
mask = repeat(mask, 'B L S-> B H L S', H=self.num_heads)
|
mask = repeat(mask, 'B L S-> B H L S', H=self.num_heads)
|
||||||
# compute attention
|
# 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
|
# compute activation in mlp stream, cat again and run second linear layer
|
||||||
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
|
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
|
||||||
return x + mod.gate * output
|
return x + mod.gate * output
|
||||||
|
|||||||
@@ -74,7 +74,7 @@ class VisualTransformer(BaseModel):
|
|||||||
with FS.get_from(self.pretrain_path,
|
with FS.get_from(self.pretrain_path,
|
||||||
wait_finish=True) as local_file:
|
wait_finish=True) as local_file:
|
||||||
logger.info(f'Loading checkpoint from {self.pretrain_path}')
|
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:
|
if not use_proj:
|
||||||
visual_pre.pop('proj')
|
visual_pre.pop('proj')
|
||||||
if visual_pre['conv1.weight'].dtype == torch.float16:
|
if visual_pre['conv1.weight'].dtype == torch.float16:
|
||||||
@@ -145,7 +145,7 @@ class SomeFTVisualTransformer(BaseModel):
|
|||||||
with FS.get_from(self.pretrain_path,
|
with FS.get_from(self.pretrain_path,
|
||||||
wait_finish=True) as local_file:
|
wait_finish=True) as local_file:
|
||||||
logger.info(f'Loading checkpoint from {self.pretrain_path}')
|
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)
|
state_dict_update = self.reformat_state_dict(visual_pre)
|
||||||
self.visual.load_state_dict(state_dict_update, strict=True)
|
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
|
from safetensors.torch import load_file as load_safetensors
|
||||||
model = load_safetensors(local_path)
|
model = load_safetensors(local_path)
|
||||||
else:
|
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:
|
if 'state_dict' in model:
|
||||||
model = model['state_dict']
|
model = model['state_dict']
|
||||||
new_ckpt = OrderedDict()
|
new_ckpt = OrderedDict()
|
||||||
|
|||||||
@@ -354,7 +354,7 @@ class PixArt(BaseModel):
|
|||||||
def load_pretrained_model(self, pretrained_model):
|
def load_pretrained_model(self, pretrained_model):
|
||||||
if pretrained_model:
|
if pretrained_model:
|
||||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
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:
|
if 'state_dict' in model:
|
||||||
model = model['state_dict']
|
model = model['state_dict']
|
||||||
new_ckpt = OrderedDict()
|
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
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from torch.cuda import amp
|
from torch import amp
|
||||||
from torch.nn import functional as F
|
from torch.nn import functional as F
|
||||||
from torch.nn.utils.rnn import pad_sequence
|
from torch.nn.utils.rnn import pad_sequence
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
@@ -440,7 +440,7 @@ def multi_head_varlen_attention(q_img,
|
|||||||
k = k.type(flash_dtype)
|
k = k.type(flash_dtype)
|
||||||
v = v.type(flash_dtype)
|
v = v.type(flash_dtype)
|
||||||
|
|
||||||
with amp.autocast():
|
with amp.autocast("cuda"):
|
||||||
x = flash_attn_varlen_func(q=q,
|
x = flash_attn_varlen_func(q=q,
|
||||||
k=k,
|
k=k,
|
||||||
v=v,
|
v=v,
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ import torch.nn as nn
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from einops import rearrange
|
from einops import rearrange
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
from torch.cuda import amp
|
from torch import amp
|
||||||
from torch.nn.utils.rnn import pad_sequence
|
from torch.nn.utils.rnn import pad_sequence
|
||||||
|
|
||||||
|
|
||||||
@@ -175,7 +175,7 @@ def frame_unpad(x, shapes):
|
|||||||
return torch.concat(frames)
|
return torch.concat(frames)
|
||||||
|
|
||||||
|
|
||||||
@amp.autocast(enabled=False)
|
@amp.autocast("cuda", enabled=False)
|
||||||
def rope_params(max_seq_len, dim, theta=10000):
|
def rope_params(max_seq_len, dim, theta=10000):
|
||||||
"""
|
"""
|
||||||
Precompute the frequency tensor for complex exponentials.
|
Precompute the frequency tensor for complex exponentials.
|
||||||
@@ -189,7 +189,7 @@ def rope_params(max_seq_len, dim, theta=10000):
|
|||||||
return freqs
|
return freqs
|
||||||
|
|
||||||
|
|
||||||
@amp.autocast(enabled=False)
|
@amp.autocast("cuda", enabled=False)
|
||||||
def rope_apply(x, grid_sizes, freqs):
|
def rope_apply(x, grid_sizes, freqs):
|
||||||
"""
|
"""
|
||||||
x: [B, L, N, C].
|
x: [B, L, N, C].
|
||||||
@@ -225,7 +225,7 @@ def rope_apply(x, grid_sizes, freqs):
|
|||||||
return torch.stack(output)
|
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):
|
def rope_apply_multires_pad(x, x_lens, x_shapes, freqs, pad=True):
|
||||||
"""
|
"""
|
||||||
x: [B, L, N, C].
|
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)
|
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):
|
def rope_apply_multires(x, x_lens, x_shapes, freqs, pad=True):
|
||||||
"""
|
"""
|
||||||
x: [B*L, N, C].
|
x: [B*L, N, C].
|
||||||
|
|||||||
@@ -459,7 +459,7 @@ class DiffusionUNet(BaseModel):
|
|||||||
from safetensors.torch import load_file as load_safetensors
|
from safetensors.torch import load_file as load_safetensors
|
||||||
sd = load_safetensors(path)
|
sd = load_safetensors(path)
|
||||||
else:
|
else:
|
||||||
sd = torch.load(path, map_location='cpu')
|
sd = torch.load(path, map_location='cpu', weights_only=True)
|
||||||
|
|
||||||
new_sd = OrderedDict()
|
new_sd = OrderedDict()
|
||||||
for k, v in sd.items():
|
for k, v in sd.items():
|
||||||
@@ -1231,7 +1231,7 @@ class LargenUNetXL(DiffusionUNetXL):
|
|||||||
from safetensors.torch import load_file as load_safetensors
|
from safetensors.torch import load_file as load_safetensors
|
||||||
sd = load_safetensors(path)
|
sd = load_safetensors(path)
|
||||||
else:
|
else:
|
||||||
sd = torch.load(path, map_location='cpu')
|
sd = torch.load(path, map_location='cpu', weights_only=True)
|
||||||
|
|
||||||
new_sd = OrderedDict()
|
new_sd = OrderedDict()
|
||||||
for k, v in sd.items():
|
for k, v in sd.items():
|
||||||
|
|||||||
@@ -3,8 +3,10 @@
|
|||||||
import copy
|
import copy
|
||||||
|
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
||||||
from scepter.modules.utils.config import dict_to_yaml
|
from scepter.modules.utils.config import dict_to_yaml
|
||||||
from scepter.modules.utils.distribute import gather_data, we
|
from scepter.modules.utils.distribute import gather_data, we
|
||||||
|
from scepter.modules.utils.model import get_parameter_dtype
|
||||||
from scepter.modules.utils.probe import (ProbeData, merge_gathered_probe,
|
from scepter.modules.utils.probe import (ProbeData, merge_gathered_probe,
|
||||||
register_data)
|
register_data)
|
||||||
|
|
||||||
@@ -43,13 +45,15 @@ class BaseModel(nn.Module):
|
|||||||
self._dist_data[key][k] += v
|
self._dist_data[key][k] += v
|
||||||
else:
|
else:
|
||||||
self._dist_data[key][k] = v
|
self._dist_data[key][k] = v
|
||||||
|
|
||||||
def collect_probe(self):
|
def collect_probe(self):
|
||||||
probe_data_dict = self._probe_data
|
probe_data_dict = self._probe_data
|
||||||
for k, v in self._modules.items():
|
for k, v in self._modules.items():
|
||||||
if isinstance(getattr(self, k), BaseModel):
|
if isinstance(getattr(self, k), BaseModel):
|
||||||
for kk, vv in getattr(self, k).collect_probe().items():
|
for kk, vv in getattr(self, k).collect_probe().items():
|
||||||
probe_data_dict[f'{k}/{kk}'] = vv
|
probe_data_dict[f'{k}/{kk}'] = vv
|
||||||
return probe_data_dict
|
return probe_data_dict
|
||||||
|
|
||||||
def probe_data(self):
|
def probe_data(self):
|
||||||
gather_probe_data = gather_data(self._probe_data)
|
gather_probe_data = gather_data(self._probe_data)
|
||||||
_dist_data_list = gather_data([self._dist_data])
|
_dist_data_list = gather_data([self._dist_data])
|
||||||
@@ -97,6 +101,13 @@ class BaseModel(nn.Module):
|
|||||||
self._probe_data = {}
|
self._probe_data = {}
|
||||||
return ret_data
|
return ret_data
|
||||||
|
|
||||||
|
@property
|
||||||
|
def model_dtype(self):
|
||||||
|
"""
|
||||||
|
`torch.dtype`: The dtype of the module (assuming that all the module parameters have the same dtype).
|
||||||
|
"""
|
||||||
|
return get_parameter_dtype(self)
|
||||||
|
|
||||||
def clear_probe(self):
|
def clear_probe(self):
|
||||||
self._probe_data.clear()
|
self._probe_data.clear()
|
||||||
|
|
||||||
|
|||||||
@@ -1,7 +1,27 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# 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
|
if TYPE_CHECKING:
|
||||||
from .schedules import (BaseNoiseScheduler, FlowMatchShiftScheduler,
|
from .diffusions import BaseDiffusion, DiffusionFluxRF
|
||||||
ScaledLinearScheduler)
|
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={},
|
||||||
|
)
|
||||||
|
|||||||
@@ -34,6 +34,7 @@ class BaseDiffusion(object):
|
|||||||
|
|
||||||
def init_params(self):
|
def init_params(self):
|
||||||
self.prediction_type = self.cfg.get('PREDICTION_TYPE', 'eps')
|
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,
|
self.noise_scheduler = NOISE_SCHEDULERS.build(self.cfg.NOISE_SCHEDULER,
|
||||||
logger=self.logger)
|
logger=self.logger)
|
||||||
self.sampler_scheduler = NOISE_SCHEDULERS.build(self.cfg.get(
|
self.sampler_scheduler = NOISE_SCHEDULERS.build(self.cfg.get(
|
||||||
@@ -56,7 +57,6 @@ class BaseDiffusion(object):
|
|||||||
model_kwargs={},
|
model_kwargs={},
|
||||||
steps=20,
|
steps=20,
|
||||||
sampler=None,
|
sampler=None,
|
||||||
use_dynamic_cfg=False,
|
|
||||||
guide_scale=None,
|
guide_scale=None,
|
||||||
guide_rescale=None,
|
guide_rescale=None,
|
||||||
show_progress=False,
|
show_progress=False,
|
||||||
@@ -79,7 +79,7 @@ class BaseDiffusion(object):
|
|||||||
if guide_scale is None or guide_scale == 1.0:
|
if guide_scale is None or guide_scale == 1.0:
|
||||||
out = model(x=x_t, t=t, **model_kwargs)
|
out = model(x=x_t, t=t, **model_kwargs)
|
||||||
else:
|
else:
|
||||||
if use_dynamic_cfg:
|
if self.use_dynamic_cfg:
|
||||||
guidance_scale = 1 + guide_scale * (
|
guidance_scale = 1 + guide_scale * (
|
||||||
(1 - math.cos(math.pi * (
|
(1 - math.cos(math.pi * (
|
||||||
(steps - timestamp.item()) / steps)**5.0)) / 2)
|
(steps - timestamp.item()) / steps)**5.0)) / 2)
|
||||||
@@ -158,14 +158,16 @@ class BaseDiffusion(object):
|
|||||||
|
|
||||||
def get_sampler(self, sampler):
|
def get_sampler(self, sampler):
|
||||||
if isinstance(sampler, str):
|
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:
|
if self.logger is not None:
|
||||||
self.logger.info(
|
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:
|
else:
|
||||||
print(
|
print(
|
||||||
f'{sampler} not in the defined samplers list {DIFFUSION_SAMPLERS.class_map.keys()}'
|
f'{sampler} not in the defined samplers list.'
|
||||||
)
|
)
|
||||||
return None
|
return None
|
||||||
sampler_cfg = Config(cfg_dict={'NAME': sampler}, load=False)
|
sampler_cfg = Config(cfg_dict={'NAME': sampler}, load=False)
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
import math
|
import math
|
||||||
|
import random
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Callable
|
from typing import Callable
|
||||||
|
|
||||||
@@ -30,7 +31,7 @@ class ScheduleOutput(object):
|
|||||||
|
|
||||||
@NOISE_SCHEDULERS.register_class()
|
@NOISE_SCHEDULERS.register_class()
|
||||||
class BaseNoiseScheduler(object):
|
class BaseNoiseScheduler(object):
|
||||||
'''
|
r'''
|
||||||
In the diffusion model, the parameters related to the noise schedule are alpha, beta,
|
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
|
and sigma. The following are the definitions of the above three parameters, which should
|
||||||
be the basic property for the instance of noise scheduler.
|
be the basic property for the instance of noise scheduler.
|
||||||
@@ -483,6 +484,14 @@ class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler):
|
|||||||
'MAX_SHIFT': {
|
'MAX_SHIFT': {
|
||||||
'value': 1.15,
|
'value': 1.15,
|
||||||
'description': 'The max shift factor for the timestamp.'
|
'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.'
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -492,6 +501,23 @@ class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler):
|
|||||||
self.sigmoid_scale = self.cfg.get('SIGMOID_SCALE', 1)
|
self.sigmoid_scale = self.cfg.get('SIGMOID_SCALE', 1)
|
||||||
self.base_shift = self.cfg.get('BASE_SHIFT', 0.5)
|
self.base_shift = self.cfg.get('BASE_SHIFT', 0.5)
|
||||||
self.max_shift = self.cfg.get('MAX_SHIFT', 1.15)
|
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):
|
def time_shift(self, mu: float, sigma_scale: float, t: Tensor):
|
||||||
return math.exp(mu) / (math.exp(mu) + (1 / t - 1)**sigma_scale)
|
return math.exp(mu) / (math.exp(mu) + (1 / t - 1)**sigma_scale)
|
||||||
@@ -516,11 +542,22 @@ class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler):
|
|||||||
n, _, h, w = x_0.shape
|
n, _, h, w = x_0.shape
|
||||||
seq_len = (h // 2 * w // 2)
|
seq_len = (h // 2 * w // 2)
|
||||||
if t is None:
|
if t is None:
|
||||||
logits_norm = torch.randn(x_0.shape[0], device=x_0.device)
|
if self.pre_t_sample:
|
||||||
logits_norm = logits_norm * self.sigmoid_scale # larger scale for more uniform sampling
|
timestep_indices = torch.randint(
|
||||||
t = logits_norm.sigmoid() * self.num_timesteps
|
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)
|
sigma = self.t_to_sigma(t, seq_len=seq_len)
|
||||||
shape = (x_0.size(0), ) + (1, ) * (x_0.ndim - 1)
|
shape = (x_0.size(0), ) + (1, ) * (x_0.ndim - 1)
|
||||||
|
# print(sigma)
|
||||||
x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
|
x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
|
||||||
return ScheduleOutput(x_0=x_0,
|
return ScheduleOutput(x_0=x_0,
|
||||||
x_t=x_t,
|
x_t=x_t,
|
||||||
|
|||||||
@@ -1,8 +1,30 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# 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,
|
if TYPE_CHECKING:
|
||||||
FrozenOpenCLIPEmbedder, FrozenOpenCLIPEmbedder2, GeneralConditioner,
|
from scepter.modules.model.embedder.embedder import (
|
||||||
IPAdapterPlusEmbedder, RefCrossEmbedder, SD3TextEmbedder, T5EmbedderHF)
|
ConcatTimestepEmbedderND, FrozenCLIPEmbedder, FrozenCLIPEmbedder2,
|
||||||
from scepter.modules.model.embedder.flux_embedder import HFEmbedder
|
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 autocast(f, enabled=True):
|
||||||
def do_autocast(*args, **kwargs):
|
def do_autocast(*args, **kwargs):
|
||||||
with torch.cuda.amp.autocast(
|
with torch.amp.autocast(
|
||||||
|
"cuda",
|
||||||
enabled=enabled,
|
enabled=enabled,
|
||||||
dtype=torch.get_autocast_gpu_dtype(),
|
dtype=torch.get_autocast_gpu_dtype(),
|
||||||
cache_enabled=torch.is_autocast_cache_enabled(),
|
cache_enabled=torch.is_autocast_cache_enabled(),
|
||||||
@@ -239,7 +240,7 @@ class FrozenOpenCLIPEmbedder(BaseEmbedder):
|
|||||||
if cfg.PRETRAINED_MODEL is not None:
|
if cfg.PRETRAINED_MODEL is not None:
|
||||||
with FS.get_from(cfg.PRETRAINED_MODEL,
|
with FS.get_from(cfg.PRETRAINED_MODEL,
|
||||||
wait_finish=True) as local_path:
|
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.model = model
|
||||||
|
|
||||||
self.use_grad = cfg.get('USE_GRAD', False)
|
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:
|
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'],
|
self.image_proj_model.load_state_dict(ckpt['image_proj'],
|
||||||
strict=True)
|
strict=True)
|
||||||
|
|
||||||
@@ -645,7 +646,7 @@ class GeneralConditioner(BaseEmbedder):
|
|||||||
from safetensors.torch import load_file as load_safetensors
|
from safetensors.torch import load_file as load_safetensors
|
||||||
sd = load_safetensors(path)
|
sd = load_safetensors(path)
|
||||||
else:
|
else:
|
||||||
sd = torch.load(path, map_location='cpu')
|
sd = torch.load(path, map_location='cpu', weights_only=True)
|
||||||
new_sd = OrderedDict()
|
new_sd = OrderedDict()
|
||||||
for k, v in sd.items():
|
for k, v in sd.items():
|
||||||
ignored = False
|
ignored = False
|
||||||
|
|||||||
@@ -1,5 +1,7 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# 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 torch
|
||||||
import transformers
|
import transformers
|
||||||
from scepter.modules.model.embedder.base_embedder import BaseEmbedder
|
from scepter.modules.model.embedder.base_embedder import BaseEmbedder
|
||||||
@@ -51,59 +53,53 @@ class HFEmbedder(BaseEmbedder):
|
|||||||
def __init__(self, cfg, logger=None):
|
def __init__(self, cfg, logger=None):
|
||||||
super().__init__(cfg, logger=logger)
|
super().__init__(cfg, logger=logger)
|
||||||
hf_model_cls = cfg.get('HF_MODEL_CLS', None)
|
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)
|
hf_tokenizer_cls = cfg.get('HF_TOKENIZER_CLS', None)
|
||||||
tokenizer_path = cfg.get('TOKENIZER_PATH', None)
|
tokenizer_path = cfg.get('TOKENIZER_PATH', None)
|
||||||
self.max_length = cfg.get('MAX_LENGTH', 77)
|
self.max_length = cfg.get('MAX_LENGTH', 77)
|
||||||
self.output_key = cfg.get('OUTPUT_KEY', 'last_hidden_state')
|
self.output_key = cfg.get("OUTPUT_KEY", "last_hidden_state")
|
||||||
self.d_type = cfg.get('D_TYPE', 'float')
|
self.d_type = cfg.get("D_TYPE", "float")
|
||||||
self.clean = cfg.get('CLEAN', 'whitespace')
|
self.clean = cfg.get("CLEAN", "whitespace")
|
||||||
self.batch_infer = cfg.get('BATCH_INFER', False)
|
self.batch_infer = cfg.get("BATCH_INFER", False)
|
||||||
|
self.added_identifier = cfg.get('ADDED_IDENTIFIER', None)
|
||||||
torch_dtype = getattr(torch, self.d_type)
|
torch_dtype = getattr(torch, self.d_type)
|
||||||
|
|
||||||
assert hf_model_cls is not None and hf_tokenizer_cls is not None
|
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
|
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,
|
with FS.get_dir_to_local_dir(model_path, wait_finish=True) as local_path:
|
||||||
wait_finish=True) as local_path:
|
self.hf_module = getattr(transformers, hf_model_cls).from_pretrained(local_path, torch_dtype = torch_dtype)
|
||||||
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)
|
|
||||||
|
|
||||||
self.hf_module = self.hf_module.eval().requires_grad_(False)
|
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(
|
batch_encoding = self.tokenizer(
|
||||||
text,
|
text,
|
||||||
truncation=True,
|
truncation=True,
|
||||||
max_length=self.max_length,
|
max_length=self.max_length,
|
||||||
return_length=False,
|
return_length=False,
|
||||||
return_overflowing_tokens=False,
|
return_overflowing_tokens=False,
|
||||||
padding='max_length',
|
padding="max_length",
|
||||||
return_tensors='pt',
|
return_tensors="pt",
|
||||||
)
|
)
|
||||||
|
|
||||||
outputs = self.hf_module(
|
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,
|
attention_mask=None,
|
||||||
output_hidden_states=False,
|
output_hidden_states=False,
|
||||||
)
|
)
|
||||||
if return_mask:
|
if return_mask:
|
||||||
return outputs[
|
return outputs[self.output_key], batch_encoding['attention_mask'].to(self.hf_module.device)
|
||||||
self.output_key], batch_encoding['attention_mask'].to(
|
|
||||||
self.hf_module.device)
|
|
||||||
else:
|
else:
|
||||||
return outputs[self.output_key], None
|
return outputs[self.output_key], None
|
||||||
|
|
||||||
def encode(self, text, return_mask=False):
|
def encode(self, text, return_mask = False):
|
||||||
if isinstance(text, str):
|
if isinstance(text, str):
|
||||||
text = [text]
|
text = [text]
|
||||||
if self.clean:
|
if self.clean:
|
||||||
@@ -119,12 +115,36 @@ class HFEmbedder(BaseEmbedder):
|
|||||||
else:
|
else:
|
||||||
return torch.cat(cont, dim=0)
|
return torch.cat(cont, dim=0)
|
||||||
else:
|
else:
|
||||||
ret_data = self(text, return_mask=return_mask)
|
ret_data = self(text, return_mask = return_mask)
|
||||||
if return_mask:
|
if return_mask:
|
||||||
return ret_data
|
return ret_data
|
||||||
else:
|
else:
|
||||||
return ret_data[0]
|
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):
|
def _clean(self, text):
|
||||||
if self.clean == 'whitespace':
|
if self.clean == 'whitespace':
|
||||||
text = whitespace_clean(basic_clean(text))
|
text = whitespace_clean(basic_clean(text))
|
||||||
@@ -133,7 +153,6 @@ class HFEmbedder(BaseEmbedder):
|
|||||||
elif self.clean == 'canonicalize':
|
elif self.clean == 'canonicalize':
|
||||||
text = canonicalize(basic_clean(text))
|
text = canonicalize(basic_clean(text))
|
||||||
return text
|
return text
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def get_config_template():
|
def get_config_template():
|
||||||
return dict_to_yaml('EMBEDDER',
|
return dict_to_yaml('EMBEDDER',
|
||||||
@@ -141,28 +160,49 @@ class HFEmbedder(BaseEmbedder):
|
|||||||
HFEmbedder.para_dict,
|
HFEmbedder.para_dict,
|
||||||
set_name=True)
|
set_name=True)
|
||||||
|
|
||||||
|
|
||||||
@EMBEDDERS.register_class()
|
@EMBEDDERS.register_class()
|
||||||
class T5PlusClipFluxEmbedder(BaseEmbedder):
|
class T5PlusClipFluxEmbedder(BaseEmbedder):
|
||||||
"""
|
"""
|
||||||
Uses the OpenCLIP transformer encoder for text
|
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):
|
def __init__(self, cfg, logger=None):
|
||||||
super().__init__(cfg, logger=logger)
|
super().__init__(cfg, logger=logger)
|
||||||
self.t5_model = EMBEDDERS.build(cfg.T5_MODEL, logger=logger)
|
self.t5_model = EMBEDDERS.build(cfg.T5_MODEL, logger=logger)
|
||||||
self.clip_model = EMBEDDERS.build(cfg.CLIP_MODEL, logger=logger)
|
self.clip_model = EMBEDDERS.build(cfg.CLIP_MODEL, logger=logger)
|
||||||
|
|
||||||
def encode(self, text):
|
def encode(self, text, return_mask = False):
|
||||||
t5_embeds = self.t5_model.encode(text, return_mask=False)
|
t5_embeds = self.t5_model.encode(text, return_mask = return_mask)
|
||||||
clip_embeds = self.clip_model.encode(text, return_mask=False)
|
clip_embeds = self.clip_model.encode(text, return_mask = return_mask)
|
||||||
# change embedding strategy here
|
# change embedding strategy here
|
||||||
return {
|
return {
|
||||||
'context': t5_embeds,
|
'context': t5_embeds,
|
||||||
'y': clip_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
|
@staticmethod
|
||||||
def get_config_template():
|
def get_config_template():
|
||||||
return dict_to_yaml('EMBEDDER',
|
return dict_to_yaml('EMBEDDER',
|
||||||
|
|||||||
@@ -1,5 +1,25 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
from scepter.modules.model.head.classifier_head import (
|
from typing import TYPE_CHECKING
|
||||||
ClassifierHead, CosineLinearHead, TransformerHead, TransformerHeadx2,
|
from scepter.modules.utils.import_utils import LazyImportModule
|
||||||
VideoClassifierHead, VideoClassifierHeadx2)
|
|
||||||
|
|
||||||
|
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 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
from scepter.modules.model.loss.base_losses import CrossEntropy
|
from typing import TYPE_CHECKING
|
||||||
from scepter.modules.model.loss.rec_loss import MinSNRLoss, ReconstructLoss
|
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 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
from scepter.modules.model.metric.classification import (AccuracyMetric,
|
from typing import TYPE_CHECKING
|
||||||
EnsembleAccuracyMetric
|
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 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
from scepter.modules.model.neck.global_average_pooling import \
|
from typing import TYPE_CHECKING
|
||||||
GlobalAveragePooling
|
from scepter.modules.utils.import_utils import LazyImportModule
|
||||||
from scepter.modules.model.neck.identity import Identity
|
|
||||||
|
|
||||||
|
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 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
from scepter.modules.model.network.autoencoder import ae_kl
|
from typing import TYPE_CHECKING
|
||||||
from scepter.modules.model.network.classifier import Classifier
|
from scepter.modules.utils.import_utils import LazyImportModule
|
||||||
from scepter.modules.model.network.diffusion import (diffusion, schedules,
|
|
||||||
solvers)
|
|
||||||
from scepter.modules.model.network.ldm import (ldm, ldm_edit, ldm_pixart,
|
if TYPE_CHECKING:
|
||||||
ldm_sce, ldm_sd3, ldm_xl,
|
from scepter.modules.model.network.autoencoder import ae_kl
|
||||||
ldm_flux)
|
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,4 +1,23 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# 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.model.network.autoencoder.ae_kl_cogvideox import AutoencoderKLCogVideoX
|
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():
|
for k in f.keys():
|
||||||
sd[k] = f.get_tensor(k)
|
sd[k] = f.get_tensor(k)
|
||||||
else:
|
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:
|
if path.find('.pt') > -1 and 'state_dict' in sd:
|
||||||
sd = sd['state_dict']
|
sd = sd['state_dict']
|
||||||
elif path.find('.ckpt') > -1 and 'state_dict' in sd:
|
elif path.find('.ckpt') > -1 and 'state_dict' in sd:
|
||||||
@@ -373,7 +373,7 @@ class AutoencoderKLFlux(TrainModule):
|
|||||||
for k in f.keys():
|
for k in f.keys():
|
||||||
sd[k] = f.get_tensor(k)
|
sd[k] = f.get_tensor(k)
|
||||||
else:
|
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:
|
if path.find('.pt') > -1 and 'state_dict' in sd:
|
||||||
sd = sd['state_dict']
|
sd = sd['state_dict']
|
||||||
elif path.find('.ckpt') > -1 and 'state_dict' in sd:
|
elif path.find('.ckpt') > -1 and 'state_dict' in sd:
|
||||||
|
|||||||
@@ -1591,7 +1591,7 @@ class AutoencoderKLCogVideoX(TrainModule):
|
|||||||
from safetensors.torch import load_file as load_safetensors
|
from safetensors.torch import load_file as load_safetensors
|
||||||
ckpt = load_safetensors(local_model)
|
ckpt = load_safetensors(local_model)
|
||||||
else:
|
else:
|
||||||
ckpt = torch.load(local_model, map_location='cpu')
|
ckpt = torch.load(local_model, map_location='cpu', weights_only=True)
|
||||||
missing, unexpected = self.load_state_dict(ckpt, strict=False)
|
missing, unexpected = self.load_state_dict(ckpt, strict=False)
|
||||||
if we.rank == 0:
|
if we.rank == 0:
|
||||||
self.logger.info(
|
self.logger.info(
|
||||||
|
|||||||
@@ -1,4 +1,22 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
from scepter.modules.model.network.diffusion import (diffusion, schedules,
|
from typing import TYPE_CHECKING
|
||||||
solvers)
|
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,
|
percentile=None,
|
||||||
cat_uc=False,
|
cat_uc=False,
|
||||||
**kwargs):
|
**kwargs):
|
||||||
"""
|
r"""
|
||||||
Apply one step of denoising from the posterior distribution q(x_s | x_t, x0).
|
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
|
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
|
distribution p(x_s | x_t, \hat{x}_0 == f(x_t)). # noqa
|
||||||
|
|||||||
@@ -1,15 +1,44 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
from scepter.modules.model.network.ldm.ldm import LatentDiffusion
|
from typing import TYPE_CHECKING
|
||||||
from scepter.modules.model.network.ldm.ldm_ace import (LatentDiffusionACE,
|
from scepter.modules.utils.import_utils import LazyImportModule
|
||||||
LatentDiffusionACERefiner)
|
|
||||||
from scepter.modules.model.network.ldm.ldm_edit import LatentDiffusionEdit
|
|
||||||
from scepter.modules.model.network.ldm.ldm_pixart import LatentDiffusionPixart
|
if TYPE_CHECKING:
|
||||||
from scepter.modules.model.network.ldm.ldm_sce import (
|
from scepter.modules.model.network.ldm.ldm import LatentDiffusion
|
||||||
LatentDiffusionSCEControl, LatentDiffusionSCETuning,
|
from scepter.modules.model.network.ldm.ldm_ace import (LatentDiffusionACE,
|
||||||
LatentDiffusionXLSCEControl, LatentDiffusionXLSCETuning)
|
LatentDiffusionACERefiner)
|
||||||
from scepter.modules.model.network.ldm.ldm_sd3 import LatentDiffusionSD3
|
from scepter.modules.model.network.ldm.ldm_edit import LatentDiffusionEdit
|
||||||
from scepter.modules.model.network.ldm.ldm_xl import LatentDiffusionXL
|
from scepter.modules.model.network.ldm.ldm_pixart import LatentDiffusionPixart
|
||||||
from scepter.modules.model.network.ldm.ldm_cogvideox import LatentDiffusionCogVideoX
|
from scepter.modules.model.network.ldm.ldm_sce import (
|
||||||
from scepter.modules.model.network.ldm.ldm_flux import (LatentDiffusionFlux,
|
LatentDiffusionSCEControl, LatentDiffusionSCETuning,
|
||||||
LatentDiffusionFluxMR)
|
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
|
from safetensors.torch import load_file as load_safetensors
|
||||||
sd = load_safetensors(path)
|
sd = load_safetensors(path)
|
||||||
else:
|
else:
|
||||||
sd = torch.load(path, map_location='cpu')
|
sd = torch.load(path, map_location='cpu',weights_only=True)
|
||||||
new_sd = OrderedDict()
|
new_sd = OrderedDict()
|
||||||
for k, v in sd.items():
|
for k, v in sd.items():
|
||||||
ignored = False
|
ignored = False
|
||||||
|
|||||||
@@ -95,8 +95,8 @@ class LatentDiffusionACE(LatentDiffusion):
|
|||||||
return batch_data_list
|
return batch_data_list
|
||||||
|
|
||||||
def forward_train(self,
|
def forward_train(self,
|
||||||
edit_image=[],
|
src_image_list=[],
|
||||||
edit_image_mask=[],
|
src_mask_list=[],
|
||||||
image=None,
|
image=None,
|
||||||
image_mask=None,
|
image_mask=None,
|
||||||
noise=None,
|
noise=None,
|
||||||
@@ -114,8 +114,8 @@ class LatentDiffusionACE(LatentDiffusion):
|
|||||||
Returns:
|
Returns:
|
||||||
'''
|
'''
|
||||||
assert check_list_of_list(prompt) and check_list_of_list(
|
assert check_list_of_list(prompt) and check_list_of_list(
|
||||||
edit_image) and check_list_of_list(edit_image_mask)
|
src_image_list) and check_list_of_list(src_mask_list)
|
||||||
assert len(edit_image) == len(edit_image_mask) == len(prompt)
|
assert len(src_image_list) == len(src_mask_list) == len(prompt)
|
||||||
assert self.cond_stage_model is not None
|
assert self.cond_stage_model is not None
|
||||||
gc_seg = kwargs.pop('gc_seg', [])
|
gc_seg = kwargs.pop('gc_seg', [])
|
||||||
gc_seg = int(gc_seg[0]) if len(gc_seg) > 0 else 0
|
gc_seg = int(gc_seg[0]) if len(gc_seg) > 0 else 0
|
||||||
@@ -143,13 +143,13 @@ class LatentDiffusionACE(LatentDiffusion):
|
|||||||
'encode_list_of_list')(prompt_, return_mask=True)
|
'encode_list_of_list')(prompt_, return_mask=True)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(e, prompt_)
|
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)
|
cont_mask)
|
||||||
context['crossattn'] = cont
|
context['crossattn'] = cont
|
||||||
|
|
||||||
# process edit image & edit image mask
|
# process edit image & edit image mask
|
||||||
edit_image = [to_device(i, strict=False) for i in edit_image]
|
edit_image = [to_device(i, strict=False) for i in src_image_list]
|
||||||
edit_image_mask = [to_device(i, strict=False) for i in edit_image_mask]
|
edit_image_mask = [to_device(i, strict=False) for i in src_mask_list]
|
||||||
e_img, e_mask = [], []
|
e_img, e_mask = [], []
|
||||||
for u, m in zip(edit_image, edit_image_mask):
|
for u, m in zip(edit_image, edit_image_mask):
|
||||||
if m is None:
|
if m is None:
|
||||||
@@ -185,8 +185,8 @@ class LatentDiffusionACE(LatentDiffusion):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def forward_test(self,
|
def forward_test(self,
|
||||||
edit_image=[],
|
src_image_list=[],
|
||||||
edit_image_mask=[],
|
src_mask_list=[],
|
||||||
image=None,
|
image=None,
|
||||||
image_mask=None,
|
image_mask=None,
|
||||||
prompt=[],
|
prompt=[],
|
||||||
@@ -200,8 +200,8 @@ class LatentDiffusionACE(LatentDiffusion):
|
|||||||
**kwargs):
|
**kwargs):
|
||||||
|
|
||||||
assert check_list_of_list(prompt) and check_list_of_list(
|
assert check_list_of_list(prompt) and check_list_of_list(
|
||||||
edit_image) and check_list_of_list(edit_image_mask)
|
src_image_list) and check_list_of_list(src_mask_list)
|
||||||
assert len(edit_image) == len(edit_image_mask) == len(prompt)
|
assert len(src_image_list) == len(src_mask_list) == len(prompt)
|
||||||
assert self.cond_stage_model is not None
|
assert self.cond_stage_model is not None
|
||||||
# gc_seg is unused
|
# gc_seg is unused
|
||||||
kwargs.pop('gc_seg', -1)
|
kwargs.pop('gc_seg', -1)
|
||||||
@@ -209,7 +209,7 @@ class LatentDiffusionACE(LatentDiffusion):
|
|||||||
context, null_context = {}, {}
|
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 = 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)
|
log_num)
|
||||||
g = torch.Generator(device=we.device_id)
|
g = torch.Generator(device=we.device_id)
|
||||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||||
@@ -368,8 +368,8 @@ class LatentDiffusionACERefiner(LatentDiffusionACE):
|
|||||||
self.enhence_sampler_cfg = None
|
self.enhence_sampler_cfg = None
|
||||||
|
|
||||||
def forward_sample(self,
|
def forward_sample(self,
|
||||||
edit_image=[],
|
src_image_list=[],
|
||||||
edit_mask=[],
|
src_mask_list=[],
|
||||||
noise=None,
|
noise=None,
|
||||||
cond_mask=[],
|
cond_mask=[],
|
||||||
x_shapes=[],
|
x_shapes=[],
|
||||||
@@ -414,15 +414,15 @@ class LatentDiffusionACERefiner(LatentDiffusionACE):
|
|||||||
# with torch.autocast(device_type="cuda", enabled=True, dtype=torch.bfloat16):
|
# 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 = getattr(self.cond_stage_model, 'encode_list')(prompt, return_mask=True)
|
||||||
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont, cont_mask)
|
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 = getattr(self.cond_stage_model, 'encode_list')(n_prompt, return_mask=True)
|
||||||
null_cont, null_cont_mask = self.cond_stage_embeddings(prompt, edit_image, null_cont, null_cont_mask)
|
null_cont, null_cont_mask = self.cond_stage_embeddings(prompt, src_image_list, null_cont, null_cont_mask)
|
||||||
context['crossattn'] = cont
|
context['crossattn'] = cont
|
||||||
null_context['crossattn'] = null_cont
|
null_context['crossattn'] = null_cont
|
||||||
|
|
||||||
|
|
||||||
null_context['edit'] = context['edit'] = edit_image
|
null_context['edit'] = context['edit'] = src_image_list
|
||||||
null_context['edit_mask'] = context['edit_mask'] = edit_mask
|
null_context['edit_mask'] = context['edit_mask'] = src_mask_list
|
||||||
|
|
||||||
# process sample
|
# process sample
|
||||||
model = self.model_ema if self.use_ema and self.eval_ema else self.model
|
model = self.model_ema if self.use_ema and self.eval_ema else self.model
|
||||||
@@ -478,8 +478,8 @@ class LatentDiffusionACERefiner(LatentDiffusionACE):
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def forward_test(self,
|
def forward_test(self,
|
||||||
edit_image=[],
|
src_image_list=[],
|
||||||
edit_image_mask=[],
|
src_mask_list=[],
|
||||||
image=None,
|
image=None,
|
||||||
image_mask=None,
|
image_mask=None,
|
||||||
prompt=[],
|
prompt=[],
|
||||||
@@ -493,13 +493,13 @@ class LatentDiffusionACERefiner(LatentDiffusionACE):
|
|||||||
enhance_scale=0.99,
|
enhance_scale=0.99,
|
||||||
log_num=-1,
|
log_num=-1,
|
||||||
**kwargs):
|
**kwargs):
|
||||||
assert check_list_of_list(prompt) and check_list_of_list(edit_image) and check_list_of_list(edit_image_mask)
|
assert check_list_of_list(prompt) and check_list_of_list(src_image_list) and check_list_of_list(src_mask_list)
|
||||||
assert len(edit_image) == len(edit_image_mask) == len(prompt)
|
assert len(src_image_list) == len(src_mask_list) == len(prompt)
|
||||||
assert self.cond_stage_model is not None
|
assert self.cond_stage_model is not None
|
||||||
# gc_seg is unused
|
# gc_seg is unused
|
||||||
kwargs.pop("gc_seg", -1)
|
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, edit_image, edit_image_mask = self.limit_batch_data(
|
||||||
[prompt, n_prompt, image, image_mask, edit_image, edit_image_mask], log_num)
|
[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]
|
prompt = [[pp] if isinstance(pp, str) else pp for pp in prompt]
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -26,19 +26,27 @@ class LatentDiffusionCogVideoX(LatentDiffusion):
|
|||||||
self.use_rotary_positional_embeddings = self.model_config.get('USE_ROTARY_POSITIONAL_EMBEDDINGS', False)
|
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.attention_head_dim = self.model_config.get('ATTENTION_HEAD_DIM', 64)
|
||||||
self.patch_size = self.model_config.get('PATCH_SIZE', 2)
|
self.patch_size = self.model_config.get('PATCH_SIZE', 2)
|
||||||
self.sample_height = self.first_stage_config.get('SAMPLE_HEIGHT', 480)
|
self.patch_size_t = self.model_config.get('PATCH_SIZE_T', None)
|
||||||
self.sample_width = self.first_stage_config.get('SAMPLE_WIDTH', 720)
|
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.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):
|
def construct_network(self):
|
||||||
super().construct_network()
|
super().construct_network()
|
||||||
self.model = self.model.to(getattr(torch, self.model_config.DTYPE))
|
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()
|
@torch.no_grad()
|
||||||
def encode_first_stage(self, x, **kwargs):
|
def encode_first_stage(self, x, **kwargs):
|
||||||
if isinstance(x, list):
|
if isinstance(x, list):
|
||||||
x = torch.stack(x, dim=0) # [B, C, F, H, W]
|
x = torch.stack(x, dim=0) # [B, C, F, H, W]
|
||||||
latents = self.scaling_factor_image * self.first_stage_model.encode(x).sample()
|
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
|
return latents
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
@@ -78,18 +86,36 @@ class LatentDiffusionCogVideoX(LatentDiffusion):
|
|||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
grid_height = height // (self.scale_factor_spatial * self.patch_size)
|
grid_height = height // (self.scale_factor_spatial * self.patch_size)
|
||||||
grid_width = width // (self.scale_factor_spatial * self.patch_size)
|
grid_width = width // (self.scale_factor_spatial * self.patch_size)
|
||||||
base_size_width = self.sample_width // (self.scale_factor_spatial * self.patch_size)
|
|
||||||
base_size_height = self.sample_height // (self.scale_factor_spatial * self.patch_size)
|
|
||||||
|
|
||||||
grid_crops_coords = get_resize_crop_region_for_grid(
|
p = self.patch_size
|
||||||
(grid_height, grid_width), base_size_width, base_size_height
|
p_t = self.patch_size_t
|
||||||
)
|
|
||||||
freqs_cos, freqs_sin = get_3d_rotary_pos_embed(
|
base_size_width = self.sample_width // p
|
||||||
embed_dim=self.attention_head_dim,
|
base_size_height = self.sample_height // p
|
||||||
crops_coords=grid_crops_coords,
|
|
||||||
grid_size=(grid_height, grid_width),
|
if p_t is None:
|
||||||
temporal_size=num_frames,
|
# 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_cos = freqs_cos.to(device=device)
|
||||||
freqs_sin = freqs_sin.to(device=device)
|
freqs_sin = freqs_sin.to(device=device)
|
||||||
@@ -121,19 +147,21 @@ class LatentDiffusionCogVideoX(LatentDiffusion):
|
|||||||
else:
|
else:
|
||||||
image_latent = None
|
image_latent = None
|
||||||
|
|
||||||
height, width = image_size
|
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 = (
|
image_rotary_emb = (
|
||||||
self._prepare_rotary_positional_embeddings(height, width, noise.size(1), we.device_id)
|
self._prepare_rotary_positional_embeddings(height=height, width=width, num_frames=noise.size(1), device=we.device_id)
|
||||||
if self.use_rotary_positional_embeddings
|
if self.use_rotary_positional_embeddings
|
||||||
else None
|
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,
|
loss = self.diffusion.loss(x_0=x_start,
|
||||||
t=t,
|
t=t,
|
||||||
model=self.model,
|
model=self.model,
|
||||||
model_kwargs={"cond": cont,
|
model_kwargs={"cond": cont,
|
||||||
'image_latent': image_latent,
|
'image_latent': image_latent,
|
||||||
'image_rotary_emb': image_rotary_emb},
|
'image_rotary_emb': image_rotary_emb,
|
||||||
|
'ofs': ofs_emb},
|
||||||
noise=noise,
|
noise=noise,
|
||||||
**kwargs)
|
**kwargs)
|
||||||
loss = loss.mean()
|
loss = loss.mean()
|
||||||
@@ -160,6 +188,7 @@ class LatentDiffusionCogVideoX(LatentDiffusion):
|
|||||||
image_size = [480, 720]
|
image_size = [480, 720]
|
||||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||||
generator = torch.Generator().manual_seed(seed)
|
generator = torch.Generator().manual_seed(seed)
|
||||||
|
# generator = torch.Generator(we.device_id).manual_seed(seed)
|
||||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||||
num_samples = len(prompt)
|
num_samples = len(prompt)
|
||||||
n_prompt = default(n_prompt, [self.default_n_prompt] * len(prompt))
|
n_prompt = default(n_prompt, [self.default_n_prompt] * len(prompt))
|
||||||
@@ -169,14 +198,21 @@ class LatentDiffusionCogVideoX(LatentDiffusion):
|
|||||||
cont = getattr(self.cond_stage_model, 'encode')(prompt, return_mask=False, use_mask=False)
|
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)
|
null_cont = getattr(self.cond_stage_model, 'encode')(n_prompt, return_mask=False, use_mask=False)
|
||||||
|
|
||||||
height, width = image_size
|
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)
|
noise = self.noise_sample(num_samples, num_frames, height, width, generator)
|
||||||
image_rotary_emb = (
|
image_rotary_emb = (
|
||||||
self._prepare_rotary_positional_embeddings(height, width, noise.size(1), we.device_id)
|
self._prepare_rotary_positional_embeddings(height, width, noise.size(1), we.device_id)
|
||||||
if self.use_rotary_positional_embeddings
|
if self.use_rotary_positional_embeddings
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
|
|
||||||
image_latent, image = self.get_image_latent(image, video, noise) if image is not None else (None, 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,
|
samples = self.diffusion.sample(noise=noise,
|
||||||
sampler=sampler,
|
sampler=sampler,
|
||||||
@@ -185,19 +221,21 @@ class LatentDiffusionCogVideoX(LatentDiffusion):
|
|||||||
'cond': cont,
|
'cond': cont,
|
||||||
'image_latent': image_latent,
|
'image_latent': image_latent,
|
||||||
'image_rotary_emb': image_rotary_emb,
|
'image_rotary_emb': image_rotary_emb,
|
||||||
|
'ofs': ofs_emb
|
||||||
}, {
|
}, {
|
||||||
'cond': null_cont,
|
'cond': null_cont,
|
||||||
'image_latent': image_latent,
|
'image_latent': image_latent,
|
||||||
'image_rotary_emb': image_rotary_emb,
|
'image_rotary_emb': image_rotary_emb,
|
||||||
|
'ofs': ofs_emb
|
||||||
}],
|
}],
|
||||||
steps=sample_steps,
|
steps=sample_steps,
|
||||||
show_progress=True,
|
show_progress=True,
|
||||||
use_dynamic_cfg=True,
|
|
||||||
guide_scale=guide_scale,
|
guide_scale=guide_scale,
|
||||||
guide_rescale=guide_rescale,
|
guide_rescale=guide_rescale,
|
||||||
return_intermediate=None,
|
return_intermediate=None,
|
||||||
**kwargs).float()
|
**kwargs).float()
|
||||||
|
|
||||||
|
samples = samples[:, additional_frames:]
|
||||||
x_frames = self.decode_first_stage(samples).float()
|
x_frames = self.decode_first_stage(samples).float()
|
||||||
|
|
||||||
outputs = []
|
outputs = []
|
||||||
|
|||||||
@@ -14,16 +14,13 @@ from scepter.modules.model.utils.basic_utils import disabled_train, check_list_o
|
|||||||
from scepter.modules.utils.config import dict_to_yaml
|
from scepter.modules.utils.config import dict_to_yaml
|
||||||
from scepter.modules.utils.distribute import we
|
from scepter.modules.utils.distribute import we
|
||||||
from scepter.modules.model.utils.basic_utils import count_params
|
from scepter.modules.model.utils.basic_utils import count_params
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
@MODELS.register_class()
|
@MODELS.register_class()
|
||||||
class LatentDiffusionFlux(LatentDiffusion):
|
class LatentDiffusionFlux(LatentDiffusion):
|
||||||
para_dict = LatentDiffusion.para_dict
|
para_dict = LatentDiffusion.para_dict
|
||||||
|
|
||||||
def __init__(self, cfg, logger=None):
|
def __init__(self, cfg, logger=None):
|
||||||
super().__init__(cfg, logger=logger)
|
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):
|
def init_params(self):
|
||||||
self.parameterization = self.cfg.get('PARAMETERIZATION', 'rf')
|
self.parameterization = self.cfg.get('PARAMETERIZATION', 'rf')
|
||||||
@@ -271,7 +268,7 @@ class LatentDiffusionFluxMR(LatentDiffusionFlux):
|
|||||||
guide_scale=3.5,
|
guide_scale=3.5,
|
||||||
show_process=True,
|
show_process=True,
|
||||||
x = None,
|
x = None,
|
||||||
reverse_scale = 0.,
|
reverse_scale = -1.,
|
||||||
**kwargs
|
**kwargs
|
||||||
):
|
):
|
||||||
noise, x_shapes = pack_imagelist_into_tensor(noise)
|
noise, x_shapes = pack_imagelist_into_tensor(noise)
|
||||||
@@ -377,9 +374,167 @@ class LatentDiffusionFluxMR(LatentDiffusionFlux):
|
|||||||
zu = zu[0]
|
zu = zu[0]
|
||||||
return zu
|
return zu
|
||||||
|
|
||||||
z = [run_one_image(u.unsqueeze(0) if u.dim == 3 else u) for u in x]
|
z = [run_one_image(u.unsqueeze(0) if u.dim() == 3 else u) for u in x]
|
||||||
return z
|
return z
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def decode_first_stage(self, z):
|
def decode_first_stage(self, z):
|
||||||
return [self.first_stage_model.decode(zu) for zu in 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)
|
||||||
|
|||||||
@@ -80,7 +80,7 @@ class LatentDiffusionXL(LatentDiffusion):
|
|||||||
from safetensors.torch import load_file as load_safetensors
|
from safetensors.torch import load_file as load_safetensors
|
||||||
sd = load_safetensors(path)
|
sd = load_safetensors(path)
|
||||||
else:
|
else:
|
||||||
sd = torch.load(path, map_location='cpu')
|
sd = torch.load(path, map_location='cpu', weights_only=True)
|
||||||
new_sd = OrderedDict()
|
new_sd = OrderedDict()
|
||||||
for k, v in sd.items():
|
for k, v in sd.items():
|
||||||
ignored = False
|
ignored = False
|
||||||
|
|||||||
@@ -1,5 +1,7 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
|
import torch
|
||||||
|
|
||||||
from scepter.modules.utils.config import Config
|
from scepter.modules.utils.config import Config
|
||||||
from scepter.modules.utils.registry import Registry, build_from_config
|
from scepter.modules.utils.registry import Registry, build_from_config
|
||||||
|
|
||||||
@@ -15,17 +17,25 @@ def build_model(cfg, registry, logger=None, *args, **kwargs):
|
|||||||
raise TypeError(f'Config must be type dict, got {type(cfg)}')
|
raise TypeError(f'Config must be type dict, got {type(cfg)}')
|
||||||
if cfg.have('PRETRAINED_MODEL'):
|
if cfg.have('PRETRAINED_MODEL'):
|
||||||
pretrain_cfg = cfg.PRETRAINED_MODEL
|
pretrain_cfg = cfg.PRETRAINED_MODEL
|
||||||
if pretrain_cfg is not None and not isinstance(pretrain_cfg, (str, list)):
|
if pretrain_cfg is not None and not isinstance(pretrain_cfg,
|
||||||
|
(str, list)):
|
||||||
raise TypeError('Pretrain parameter must be a string or list')
|
raise TypeError('Pretrain parameter must be a string or list')
|
||||||
else:
|
else:
|
||||||
pretrain_cfg = None
|
pretrain_cfg = None
|
||||||
|
if cfg.get('MODEL_DTYPE', None):
|
||||||
|
default_dtype = getattr(torch, cfg.MODEL_DTYPE)
|
||||||
|
ori_default_dtype = torch.get_default_dtype()
|
||||||
|
torch.set_default_dtype(default_dtype)
|
||||||
|
device = cfg.get('DEVICE', None)
|
||||||
model = build_from_config(cfg, registry, logger=logger, *args, **kwargs)
|
model = build_from_config(cfg, registry, logger=logger, *args, **kwargs)
|
||||||
|
if cfg.get('MODEL_DTYPE', None):
|
||||||
|
torch.set_default_dtype(ori_default_dtype)
|
||||||
if pretrain_cfg is not None:
|
if pretrain_cfg is not None:
|
||||||
if hasattr(model, 'load_pretrained_model'):
|
if hasattr(model, 'load_pretrained_model'):
|
||||||
model.load_pretrained_model(pretrain_cfg)
|
model.load_pretrained_model(pretrain_cfg)
|
||||||
return model
|
return model
|
||||||
|
|
||||||
|
|
||||||
def build_diffusion(cfg, registry, logger=None, *args, **kwargs):
|
def build_diffusion(cfg, registry, logger=None, *args, **kwargs):
|
||||||
""" After build model, load pretrained model if exists key `pretrain`.
|
""" After build model, load pretrained model if exists key `pretrain`.
|
||||||
|
|
||||||
@@ -37,11 +47,13 @@ def build_diffusion(cfg, registry, logger=None, *args, **kwargs):
|
|||||||
raise TypeError(f'Config must be type dict, got {type(cfg)}')
|
raise TypeError(f'Config must be type dict, got {type(cfg)}')
|
||||||
return build_from_config(cfg, registry, logger=logger, *args, **kwargs)
|
return build_from_config(cfg, registry, logger=logger, *args, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
def build_scheduler(cfg, registry, logger=None, *args, **kwargs):
|
def build_scheduler(cfg, registry, logger=None, *args, **kwargs):
|
||||||
if not isinstance(cfg, Config):
|
if not isinstance(cfg, Config):
|
||||||
raise TypeError(f'Config must be type dict, got {type(cfg)}')
|
raise TypeError(f'Config must be type dict, got {type(cfg)}')
|
||||||
return build_from_config(cfg, registry, logger=logger, *args, **kwargs)
|
return build_from_config(cfg, registry, logger=logger, *args, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
def build_diffusion_sampler(cfg, registry, logger=None, *args, **kwargs):
|
def build_diffusion_sampler(cfg, registry, logger=None, *args, **kwargs):
|
||||||
if not isinstance(cfg, Config):
|
if not isinstance(cfg, Config):
|
||||||
raise TypeError(f'Config must be type dict, got {type(cfg)}')
|
raise TypeError(f'Config must be type dict, got {type(cfg)}')
|
||||||
@@ -49,7 +61,7 @@ def build_diffusion_sampler(cfg, registry, logger=None, *args, **kwargs):
|
|||||||
|
|
||||||
|
|
||||||
MODELS = Registry('MODELS', build_func=build_model)
|
MODELS = Registry('MODELS', build_func=build_model)
|
||||||
TOKENIZERS = Registry('TOKENIZER', build_func=build_model)
|
TOKENIZERS = Registry('TOKENIZERS', build_func=build_model)
|
||||||
EMBEDDERS = Registry('EMBEDDERS', build_func=build_model)
|
EMBEDDERS = Registry('EMBEDDERS', build_func=build_model)
|
||||||
BACKBONES = Registry('BACKBONES', build_func=build_model)
|
BACKBONES = Registry('BACKBONES', build_func=build_model)
|
||||||
NECKS = Registry('NECKS', build_func=build_model)
|
NECKS = Registry('NECKS', build_func=build_model)
|
||||||
@@ -60,7 +72,7 @@ LOSSES = Registry('LOSSES', build_func=build_model)
|
|||||||
TUNERS = Registry('TUNERS', build_func=build_model)
|
TUNERS = Registry('TUNERS', build_func=build_model)
|
||||||
|
|
||||||
# reigister cls for diffusion.
|
# reigister cls for diffusion.
|
||||||
|
|
||||||
DIFFUSIONS = Registry('DIFFUSIONS', build_func=build_diffusion)
|
DIFFUSIONS = Registry('DIFFUSIONS', build_func=build_diffusion)
|
||||||
NOISE_SCHEDULERS = Registry('NOISE_SCHEDULERS', build_func=build_diffusion)
|
NOISE_SCHEDULERS = Registry('NOISE_SCHEDULERS', build_func=build_diffusion)
|
||||||
DIFFUSION_SAMPLERS = Registry('DIFFUSION_SAMPLERS', build_func=build_diffusion_sampler)
|
DIFFUSION_SAMPLERS = Registry('DIFFUSION_SAMPLERS',
|
||||||
|
build_func=build_diffusion_sampler)
|
||||||
|
|||||||
@@ -1,6 +1,25 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
from scepter.modules.model.tokenizer.base_tokenizer import BaseTokenizer
|
from typing import TYPE_CHECKING
|
||||||
from scepter.modules.model.tokenizer.tokenizer import (ClipTokenizer,
|
from scepter.modules.utils.import_utils import LazyImportModule
|
||||||
HuggingfaceTokenizer,
|
|
||||||
OpenClipTokenizer)
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from scepter.modules.model.tokenizer.base_tokenizer import BaseTokenizer
|
||||||
|
from scepter.modules.model.tokenizer.tokenizer import (ClipTokenizer,
|
||||||
|
HuggingfaceTokenizer,
|
||||||
|
OpenClipTokenizer)
|
||||||
|
else:
|
||||||
|
_import_structure = {
|
||||||
|
'base_tokenizer': ['BaseTokenizer'],
|
||||||
|
'tokenizer': ['ClipTokenizer', 'HuggingfaceTokenizer', 'OpenClipTokenizer']
|
||||||
|
}
|
||||||
|
|
||||||
|
import sys
|
||||||
|
sys.modules[__name__] = LazyImportModule(
|
||||||
|
__name__,
|
||||||
|
globals()['__file__'],
|
||||||
|
_import_structure,
|
||||||
|
module_spec=__spec__,
|
||||||
|
extra_objects={},
|
||||||
|
)
|
||||||
|
|||||||
@@ -152,9 +152,9 @@ def heavy_clean(text):
|
|||||||
text = re.sub(r'[\"\']{2,}', r'"', text) # """AUSVERKAUFT"""
|
text = re.sub(r'[\"\']{2,}', r'"', text) # """AUSVERKAUFT"""
|
||||||
text = re.sub(r'[\.]{2,}', r' ', text) # """AUSVERKAUFT"""
|
text = re.sub(r'[\.]{2,}', r' ', text) # """AUSVERKAUFT"""
|
||||||
text = re.sub(
|
text = re.sub(
|
||||||
re.compile(r'[' + '#®•©™&@·º½¾¿¡§~' + '\)' + '\(' + '\]' + # noqa
|
re.compile(r'[' + '#®•©™&@·º½¾¿¡§~' + r'\)' + r'\(' + r'\]' + # noqa
|
||||||
'\[' + # noqa
|
r'\[' + # noqa
|
||||||
'\}' + '\{' + '\|' + '\\' + '\/' + '\*' + # noqa
|
r'\}' + r'\{' + r'\|' + '\\' + r'\/' + r'\*' + # noqa
|
||||||
r']{1,}'), # noqa
|
r']{1,}'), # noqa
|
||||||
r' ',
|
r' ',
|
||||||
text) # ***AUSVERKAUFT***, #AUSVERKAUFT
|
text) # ***AUSVERKAUFT***, #AUSVERKAUFT
|
||||||
|
|||||||
@@ -1,5 +1,25 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
from scepter.modules.model.tuner import sce
|
from typing import TYPE_CHECKING
|
||||||
from scepter.modules.model.tuner.swift_tuner import (SwiftPart, SwiftAdapter, SwiftFull,
|
from scepter.modules.utils.import_utils import LazyImportModule
|
||||||
SwiftLoRA, SwiftSCETuning)
|
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from scepter.modules.model.tuner import sce
|
||||||
|
from scepter.modules.model.tuner.swift_tuner import (SwiftPart, SwiftAdapter, SwiftFull,
|
||||||
|
SwiftLoRA, SwiftSCETuning)
|
||||||
|
else:
|
||||||
|
_import_structure = {
|
||||||
|
'tuner': ['sce'],
|
||||||
|
'swift_tuner': ['SwiftPart', 'SwiftAdapter', 'SwiftFull',
|
||||||
|
'SwiftLoRA', 'SwiftSCETuning']
|
||||||
|
}
|
||||||
|
|
||||||
|
import sys
|
||||||
|
sys.modules[__name__] = LazyImportModule(
|
||||||
|
__name__,
|
||||||
|
globals()['__file__'],
|
||||||
|
_import_structure,
|
||||||
|
module_spec=__spec__,
|
||||||
|
extra_objects={},
|
||||||
|
)
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user