Compare commits

...
27 Commits
Author SHA1 Message Date
LouieStark aac85fa94f update readme 2024-11-05 19:44:54 +08:00
LouieStark 3f267aaea2 update example 2024-11-05 15:49:00 +08:00
LouieStark 53357f95d6 update chatbot example 2024-11-05 14:36:52 +08:00
LouieStark cef93bdbfe update instr 2024-11-04 16:23:19 +08:00
LouieStark b886400e06 add instruction 2024-11-04 14:34:07 +08:00
LouieStark f98adabeb3 update readme 2024-11-01 23:32:28 +08:00
LouieStark fe3e11b49e update chatbot 2024-11-01 21:13:21 +08:00
LouieStark e6b43f19f6 update chatbot 2024-11-01 17:10:56 +08:00
LouieStark 986349deac fix chatbot bug 2024-11-01 16:38:46 +08:00
LouieStark 3f047be43c update readme and yaml 2024-11-01 11:37:53 +08:00
LouieStark e8d8e63cba update v1.2.0 2024-11-01 10:15:27 +08:00
jiangzeyinzi eac03e9856 Merge pull request #52 from modelscope/v1.1.0_dev
update v1.1.0
2024-10-23 10:19:29 +08:00
jiangzeyinzi 379b94ab4f update 2024-10-23 10:15:31 +08:00
jiangzeyinzi 342c6b8a15 Merge branch 'v1.1.0_dev' of https://github.com/modelscope/scepter into v1.1.0_dev 2024-10-21 11:59:55 +08:00
jiangzeyinzi 4ddcb08c7b update 2024-10-21 11:59:41 +08:00
jiangzeyinzi b2169a1597 Update readme.md 2024-10-21 11:58:43 +08:00
jiangzeyinzi cffd54a02a update 2024-10-21 11:36:02 +08:00
zeyinzi.jzyz 0bba2c319d update v1.1.0 2024-10-21 00:35:53 +08:00
LouieStark 7d6451efad update project page url 2024-10-01 23:45:22 +08:00
LouieStark a5decd17aa update readme 2024-09-30 14:22:14 +08:00
jiangzeyinzi 8a14866562 Merge pull request #46 from yaosheng216/patch-4
Update ldm_sce.py
2024-09-27 09:49:27 +08:00
jiangzeyinzi 5a94f6c4a2 Merge pull request #45 from yaosheng216/patch-3
Update sd15_512_sce_ctr_hed.yaml
2024-09-27 09:48:43 +08:00
Great 4ad0f6038c Update ldm_sce.py 2024-09-27 09:39:22 +08:00
Great 4f56623627 Update sd15_512_sce_ctr_hed.yaml 2024-09-27 09:36:27 +08:00
jiangzeyinzi 30afc0a1de Merge pull request #39 from yaosheng216/patch-2
Update readme.md
2024-07-18 17:51:53 +08:00
Great aeabef75b4 Update readme.md 2024-07-18 17:40:28 +08:00
jiangzeyinzi 9376149415 Merge pull request #38 from modelscope/v1.0.3_dev
v1.0.3
2024-07-18 17:31:44 +08:00
180 changed files with 23658 additions and 1480 deletions
+2 -2
View File
@@ -9,12 +9,12 @@
*.bin
*.idea
*.csv
cache
build
dist
dev
scepter.egg-info
.readthedocs.yml
1.9
#MANIFEST.in
*resources
*.ipynb_checkpoints*
*.vscode
Binary file not shown.

After

Width:  |  Height:  |  Size: 74 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 36 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 44 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 140 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 99 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 97 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 48 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 56 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 88 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 97 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 62 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 53 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 378 KiB

+163
View File
@@ -0,0 +1,163 @@
{
"last_node_id": 4,
"last_link_id": 3,
"nodes": [
{
"id": 2,
"type": "PreviewImage",
"pos": {
"0": 937,
"1": 243
},
"size": {
"0": 284.03680419921875,
"1": 246
},
"flags": {
"collapsed": false
},
"order": 2,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 1
}
],
"outputs": [],
"properties": {
"Node name for S&R": "PreviewImage"
},
"widgets_values": []
},
{
"id": 4,
"type": "ParameterNode",
"pos": {
"0": 0,
"1": 244
},
"size": {
"0": 311.9849853515625,
"1": 236.4765625
},
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "Result",
"type": "CONDITIONING",
"links": [
3
]
}
],
"properties": {
"Node name for S&R": "ParameterNode"
},
"widgets_values": [
"ddim",
50,
5,
0.5,
"trailing",
1024,
1024,
2024
]
},
{
"id": 1,
"type": "ModelNode",
"pos": {
"0": 430,
"1": 243
},
"size": {
"0": 400,
"1": 234
},
"flags": {},
"order": 1,
"mode": 0,
"inputs": [
{
"name": "parameters",
"type": "CONDITIONING",
"link": 3,
"shape": 7
},
{
"name": "mantras",
"type": "CONDITIONING",
"link": null,
"shape": 7
},
{
"name": "tuners",
"type": "CONDITIONING",
"link": null,
"shape": 7
},
{
"name": "controls",
"type": "CONDITIONING",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
1
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "ModelNode"
},
"widgets_values": [
"SD_XL1.0",
"ModelScope",
"a cat",
""
]
}
],
"links": [
[
1,
1,
0,
2,
0,
"IMAGE"
],
[
3,
4,
0,
1,
0,
"CONDITIONING"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.9229599817706423,
"offset": [
128.91466506592707,
43.07074577975898
]
}
},
"version": 0.4
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 291 KiB

+202
View File
@@ -0,0 +1,202 @@
{
"last_node_id": 5,
"last_link_id": 4,
"nodes": [
{
"id": 2,
"type": "PreviewImage",
"pos": {
"0": 937,
"1": 243
},
"size": {
"0": 284.03680419921875,
"1": 246
},
"flags": {
"collapsed": false
},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 1
}
],
"outputs": [],
"properties": {
"Node name for S&R": "PreviewImage"
},
"widgets_values": []
},
{
"id": 4,
"type": "ParameterNode",
"pos": {
"0": 0,
"1": 244
},
"size": {
"0": 311.9849853515625,
"1": 236.4765625
},
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "Result",
"type": "CONDITIONING",
"links": [
3
]
}
],
"properties": {
"Node name for S&R": "ParameterNode"
},
"widgets_values": [
"ddim",
50,
5,
0.5,
"trailing",
1024,
1024,
2024
]
},
{
"id": 1,
"type": "ModelNode",
"pos": {
"0": 430,
"1": 243
},
"size": {
"0": 400,
"1": 234
},
"flags": {},
"order": 2,
"mode": 0,
"inputs": [
{
"name": "parameters",
"type": "CONDITIONING",
"link": 3,
"shape": 7
},
{
"name": "mantras",
"type": "CONDITIONING",
"link": 4,
"shape": 7
},
{
"name": "tuners",
"type": "CONDITIONING",
"link": null,
"shape": 7
},
{
"name": "controls",
"type": "CONDITIONING",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
1
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "ModelNode"
},
"widgets_values": [
"SD_XL1.0",
"ModelScope",
"a cat",
""
]
},
{
"id": 5,
"type": "MantrasNode",
"pos": {
"0": 3,
"1": 547
},
"size": {
"0": 315,
"1": 58
},
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "Result",
"type": "CONDITIONING",
"links": [
4
]
}
],
"properties": {
"Node name for S&R": "MantrasNode"
},
"widgets_values": [
"Flat 2D Art"
]
}
],
"links": [
[
1,
1,
0,
2,
0,
"IMAGE"
],
[
3,
4,
0,
1,
0,
"CONDITIONING"
],
[
4,
5,
0,
1,
1,
"CONDITIONING"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 1.0152559799477068,
"offset": [
58.35814123566077,
-65.12749162217047
]
}
},
"version": 0.4
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 230 KiB

+242
View File
@@ -0,0 +1,242 @@
{
"last_node_id": 6,
"last_link_id": 5,
"nodes": [
{
"id": 2,
"type": "PreviewImage",
"pos": {
"0": 937,
"1": 243
},
"size": {
"0": 284.03680419921875,
"1": 246
},
"flags": {
"collapsed": false
},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 1
}
],
"outputs": [],
"properties": {
"Node name for S&R": "PreviewImage"
},
"widgets_values": []
},
{
"id": 4,
"type": "ParameterNode",
"pos": {
"0": 0,
"1": 244
},
"size": {
"0": 311.9849853515625,
"1": 236.4765625
},
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "Result",
"type": "CONDITIONING",
"links": [
3
]
}
],
"properties": {
"Node name for S&R": "ParameterNode"
},
"widgets_values": [
"ddim",
50,
5,
0.5,
"trailing",
1024,
1024,
2024
]
},
{
"id": 5,
"type": "MantrasNode",
"pos": {
"0": 3,
"1": 547
},
"size": {
"0": 315,
"1": 58
},
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "Result",
"type": "CONDITIONING",
"links": [
4
]
}
],
"properties": {
"Node name for S&R": "MantrasNode"
},
"widgets_values": [
"Flat 2D Art"
]
},
{
"id": 1,
"type": "ModelNode",
"pos": {
"0": 430,
"1": 243
},
"size": {
"0": 400,
"1": 234
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "parameters",
"type": "CONDITIONING",
"link": 3,
"shape": 7
},
{
"name": "mantras",
"type": "CONDITIONING",
"link": 4,
"shape": 7
},
{
"name": "tuners",
"type": "CONDITIONING",
"link": 5,
"shape": 7
},
{
"name": "controls",
"type": "CONDITIONING",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
1
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "ModelNode"
},
"widgets_values": [
"SD_XL1.0",
"ModelScope",
"a cat",
""
]
},
{
"id": 6,
"type": "TunerNode",
"pos": {
"0": 8,
"1": 680
},
"size": {
"0": 315,
"1": 82
},
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "Result",
"type": "CONDITIONING",
"links": [
5
]
}
],
"properties": {
"Node name for S&R": "TunerNode"
},
"widgets_values": [
"SD_XL1.0_Pencil Sketch Drawing",
1
]
}
],
"links": [
[
1,
1,
0,
2,
0,
"IMAGE"
],
[
3,
4,
0,
1,
0,
"CONDITIONING"
],
[
4,
5,
0,
1,
1,
"CONDITIONING"
],
[
5,
6,
0,
1,
2,
"CONDITIONING"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.6934334949441366,
"offset": [
450.0923948318967,
30.783653239467935
]
}
},
"version": 0.4
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 355 KiB

@@ -0,0 +1,366 @@
{
"last_node_id": 11,
"last_link_id": 9,
"nodes": [
{
"id": 1,
"type": "ModelNode",
"pos": {
"0": 439,
"1": 123
},
"size": {
"0": 400,
"1": 234
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "parameters",
"type": "CONDITIONING",
"link": 3,
"shape": 7
},
{
"name": "mantras",
"type": "CONDITIONING",
"link": 4,
"shape": 7
},
{
"name": "tuners",
"type": "CONDITIONING",
"link": 5,
"shape": 7
},
{
"name": "controls",
"type": "CONDITIONING",
"link": 6,
"shape": 7
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
1
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "ModelNode"
},
"widgets_values": [
"SD_XL1.0",
"ModelScope",
"a cat",
""
]
},
{
"id": 5,
"type": "MantrasNode",
"pos": {
"0": 13,
"1": 336
},
"size": {
"0": 315,
"1": 58
},
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "Result",
"type": "CONDITIONING",
"links": [
4
]
}
],
"properties": {
"Node name for S&R": "MantrasNode"
},
"widgets_values": [
"Flat 2D Art"
]
},
{
"id": 6,
"type": "TunerNode",
"pos": {
"0": 12,
"1": 455
},
"size": {
"0": 315,
"1": 82
},
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "Result",
"type": "CONDITIONING",
"links": [
5
]
}
],
"properties": {
"Node name for S&R": "TunerNode"
},
"widgets_values": [
"SD_XL1.0_Pencil Sketch Drawing",
1
]
},
{
"id": 11,
"type": "NoteNode",
"pos": {
"0": 845,
"1": 506
},
"size": {
"0": 398.3059387207031,
"1": 210.83267211914062
},
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {
"Node name for S&R": "NoteNode"
},
"widgets_values": [
"This is a sample for quickly setting up ComfyUI using the Scepter open-source library:\n\n1) The first run involves downloading the models. By default, we will automatically pull models from ModelScope. The initial download may take some time, and you can also adjust the model source to change the model address.\n\n2) Currently, it supports various base models, basic settings for some inference hyperparameters, mantra settings, tuning model settings, and conditional generation node.\n\n3) In the future, we will gradually integrate interesting features to enrich the use cases.\n\n"
]
},
{
"id": 8,
"type": "LoadImage",
"pos": {
"0": 471,
"1": 451
},
"size": {
"0": 315,
"1": 314
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
7
]
},
{
"name": "MASK",
"type": "MASK",
"links": null
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"cat.jpg",
"image"
]
},
{
"id": 7,
"type": "ControlNode",
"pos": {
"0": 15,
"1": 600
},
"size": {
"0": 330,
"1": 198
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "source_image",
"type": "IMAGE",
"link": 7
}
],
"outputs": [
{
"name": "Result",
"type": "CONDITIONING",
"links": [
6
]
},
{
"name": "Control Image",
"type": "IMAGE",
"links": [],
"slot_index": 1
}
],
"properties": {
"Node name for S&R": "ControlNode"
},
"widgets_values": [
"SD_XL1.0_color",
"Color",
"CenterCrop",
1,
1024,
1024
]
},
{
"id": 2,
"type": "PreviewImage",
"pos": {
"0": 959,
"1": 121
},
"size": {
"0": 284.03680419921875,
"1": 246
},
"flags": {
"collapsed": false
},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 1
}
],
"outputs": [],
"properties": {
"Node name for S&R": "PreviewImage"
},
"widgets_values": []
},
{
"id": 4,
"type": "ParameterNode",
"pos": {
"0": 13,
"1": 43
},
"size": {
"0": 311.9849853515625,
"1": 236.4765625
},
"flags": {},
"order": 4,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "Result",
"type": "CONDITIONING",
"links": [
3
]
}
],
"properties": {
"Node name for S&R": "ParameterNode"
},
"widgets_values": [
"ddim",
50,
5,
0.5,
"trailing",
1024,
1024,
2024
]
}
],
"links": [
[
1,
1,
0,
2,
0,
"IMAGE"
],
[
3,
4,
0,
1,
0,
"CONDITIONING"
],
[
4,
5,
0,
1,
1,
"CONDITIONING"
],
[
5,
6,
0,
1,
2,
"CONDITIONING"
],
[
6,
7,
0,
1,
3,
"CONDITIONING"
],
[
7,
8,
0,
7,
0,
"IMAGE"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.7627768444385667,
"offset": [
254.62860058494985,
87.05734194144193
]
}
},
"version": 0.4
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 634 KiB

+22
View File
@@ -0,0 +1,22 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os
import shutil
import subprocess
import sys
if sys.argv[0] == 'install.py':
sys.path.append('.') # for portable version
source_folder = os.path.join(os.path.dirname(__file__), "scepter/workflow")
current_dir = os.path.dirname(__file__)
destination_folder = os.path.join(os.path.dirname(current_dir), "ComfyUI-Scepter")
if not os.path.exists(destination_folder):
shutil.copytree(source_folder, destination_folder)
print(f"{os.path.abspath(source_folder)} copy to {os.path.abspath(destination_folder)} success!")
else:
print(f"{os.path.abspath(destination_folder)} exist.")
# pip install scepter
subprocess.check_call([sys.executable, '-m', 'pip', 'install', 'scepter'])
+165 -39
View File
@@ -14,10 +14,14 @@ SCEPTER integrates popular community-driven implementations as well as proprieta
SCEPTER offers 3 core components:
- [Generative training and inference framework](#tutorials)
- [Easy implementation of popular approaches](#currently-supported-approaches)
- [Interactive user interface: SCEPTER Studio](#launch)
- [Interactive user interface: SCEPTER Studio & Comfy UI](#launch)
## 🎉 News
- [🔥🔥🔥2024.10]: We are pleased to announce the release of the code for [ACE](https://arxiv.org/abs/2410.00086), supporting Customized Training / Comfy UI Workflow / gradio-based ChatBot Interface. The detailed documents can be found at [ACE repo](https://github.com/ali-vilab/ACE.git).
- [2024.10]: Support for inference and tuning with [FLUX](https://huggingface.co/black-forest-labs/FLUX.1-dev), as well as for building [ComfyUI](https://github.com/comfyanonymous/ComfyUI) workflows using this framework.
- [2024.09]: We introduce **ACE**, an **A**ll-round **C**reator and **E**ditor adept at executing a diverse array of image editing tasks tailored to your specifications. Built upon the cutting-edge Diffusion Transformer architecture, ACE has been extensively trained on a comprehensive dataset to seamlessly interpret and execute any natural language instruction. For further information, please consult the [project page](https://ali-vilab.github.io/ace-page/).
- [2024.07]: Support the inference and training of open-source generative models based on the [DiT](https://arxiv.org/abs/2212.09748) architecture, such as [SD3](https://arxiv.org/pdf/2403.03206) and [PixArt](https://arxiv.org/abs/2310.00426).
- [2024.05]: Introducing SCEPTER v1, supporting customized image edit tasks! Simply provide 10 image pairs, SCEPTER will tune an edit tuner for your own Image-to-Image tasks, like `Clay Style`, `De-Text`, `Segmentation`, etc.
- [2024.04]: New [StyleBooth](https://ali-vilab.github.io/stylebooth-page/) demo on SCEPTER Studio for`Text-Based Style Editing`.
- [2024.03]: We optimize the training UI and checkpoint management. New [LAR-Gen](https://arxiv.org/abs/2403.19534) model has been added on SCEPTER Studio, supporting `zoom-out`, `virtual try on`, `inpainting`.
@@ -30,45 +34,149 @@ SCEPTER offers 3 core components:
## 🖼 Gallery for Recent Works
### Edit Tuners
### ACE
Simply provide 10 image pairs, SCEPTER will tune an edit tuner for your own Image-to-Image tasks, like `Clay Style`, `De-Text`, `Segmentation`, etc.
Try our official few-shot datasets: [De-Text](https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=datasets%2Fdetext.zip), [Image2Hed](https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=datasets%2Fhed_pair.zip), [Image2Depth](https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=datasets%2Fimage2depth.zip), [Depth2Image](https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=datasets%2Fdepth2image.zip).
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.
[![Watch the demo](https://ali-vilab.github.io/ace-page/static/images/tasks.png)](https://ali-vilab.github.io/ace-page/)
#### 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`.
##### Prepare datasets
Please find the dataset class located in `scepter/modules/data/dataset/ms_dataset.py`,
designed to facilitate end-to-end training using an open-source toy dataset.
Download a dataset zip file from [modelscope](https://www.modelscope.cn/models/iic/scepter/resolve/master/datasets/hed_pair.zip), and then extract its contents into the `cache/datasets/` directory.
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.
##### 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
PYTHONPATH=. python scepter/tools/run_train.py --cfg scepter/methods/edit/dit_ace_0.6b_512.yaml
```
#### ACE Chat Bot
We have developed a chatbot interface utilizing Gradio, designed to convert user input in natural language into visually captivating images that align semantically with the specified instructions. You can easily access this functionality by launching Scepter Studio with the following command:
```bash
PYTHONPATH=. python scepter/tools/webui.py --cfg scepter/methods/studio/scepter_ui.yaml --language zh
```
Upon starting, you will find a "ChatBot" tab within the Gradio application, which serves as a chat-based interface to handle any requests related to image editing or generation.
#### ACE ComfyUI Workflow
![Workflow](https://github.com/ali-vilab/ace-page/raw/main/assets/comfyui/ace_example.jpg)
<table><tbody>
<tr>
<th align="center" colspan="4">Clay Style<br>Prompt: "Convert this image into clay style"</th>
<th align="center" colspan="4">ACE Workflow Examples</th>
</tr>
<tr>
<td><img src="asset/images/edit_tuner/vermeer.jpeg" width="300"></td>
<td><img src="asset/images/edit_tuner/clay_vermeer.jpeg" width="300"></td>
<td><img src="asset/images/edit_tuner/cat_512.jpg" width="300"></td>
<td><img src="asset/images/edit_tuner/clay_cat.jpeg" width="300"></td>
<th align="center" colspan="1">Control</th>
<th align="center" colspan="1">Semantic</th>
<th align="center" colspan="1">Element</th>
</tr>
<tr>
<th align="center" colspan="2">De-Text<br>Prompt: "Remove the texts"</th>
<th align="center" colspan="2">Image2Hed<br>Prompt: "Convert to an edge map"</th>
</tr>
<tr>
<td><img src="asset/images/edit_tuner/text.jpg" width="300"></td>
<td><img src="asset/images/edit_tuner/detext.jpeg" width="300"></td>
<td><img src="asset/images/edit_tuner/cat_512.jpg" width="300"></td>
<td><img src="asset/images/edit_tuner/hed.jpeg" width="300"></td>
</tr>
<tr>
<th align="center" colspan="2">Image2Depth<br>Prompt: "Calculate the depth map"</th>
<th align="center" colspan="2">Depth2Image<br>Prompt: "Convert depth map into color image"</th>
</tr>
<tr>
<td><img src="asset/images/edit_tuner/house.jpg" width="300"></td>
<td><img src="asset/images/edit_tuner/image2depth.jpeg" width="300"></td>
<td><img src="asset/images/edit_tuner/depth.jpg" width="300"></td>
<td><img src="asset/images/edit_tuner/depth2image.jpeg" width="300"></td>
<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>
Note: Left image is input and right image is output.
### FLUX Tuners
<table><tbody>
<tr>
<th align="center" colspan="3">Yarn Style</th>
<th align="center" colspan="3">Soft Watercolor Style</th>
</tr>
<tr>
<td><img src="asset/images/flux_tuner/flux_tuner_2_1.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_2_2.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_2_3.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_1_1.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_1_2.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_1_3.webp" width="200"></td>
</tr>
<tr>
<th align="center" colspan="3">Travel Style</th>
<th align="center" colspan="3">WuKong Style</th>
</tr>
<tr>
<td><img src="asset/images/flux_tuner/flux_tuner_3_1.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_3_2.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_3_3.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_4_1.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_4_2.webp" width="200"></td>
<td><img src="asset/images/flux_tuner/flux_tuner_4_3.webp" width="200"></td>
</tr>
</tbody>
</table>
### ComfyUI Workflow
![Workflow](asset/workflow/workflow.jpg)
<table><tbody>
<tr>
<th align="center" colspan="4">Example Workflow Case</th>
</tr>
<tr>
<th align="center" colspan="1">Base</th>
<th align="center" colspan="1">+Mantra</th>
<th align="center" colspan="1">+Tuner</th>
<th align="center" colspan="1">+Control</th>
</tr>
<tr>
<td>
<a href="asset/workflow/sdxl_base.json" target="_blank">
<img src="asset/workflow/sdxl_base.jpg" width="200">
</a>
</td>
<td>
<a href="asset/workflow/sdxl_base_mantra.json" target="_blank">
<img src="asset/workflow/sdxl_base_mantra.jpg" width="200">
</a>
</td>
<td>
<a href="asset/workflow/sdxl_base_mantra_tuner.json" target="_blank">
<img src="asset/workflow/sdxl_base_mantra_tuner.jpg" width="200">
</a>
</td>
<td>
<a href="asset/workflow/sdxl_base_mantra_tuner_control.json" target="_blank">
<img src="asset/workflow/sdxl_base_mantra_tuner_control.jpg" width="200">
</a>
</td>
</tr>
</tbody>
</table>
## 🛠️ Installation
@@ -102,16 +210,18 @@ pip install scepter
### Currently supported approaches
| Tasks | Methods | Links |
|:----------------------------:|:--------------------------------------------:|:------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
| Text-to-image generation | SD v1.5 | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
| Text-to-image generation | SD v2.1 | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
| Text-to-image generation | SD-XL | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0) |
| Efficient Tuning | LoRA | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=LoRA&color=red&logo=arxiv)](https://arxiv.org/abs/2106.09685) |
| Efficient Tuning | Res-Tuning(NeurIPS23) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=Res-Tuing&color=red&logo=arxiv)](https://arxiv.org/abs/2310.19859) [![Page link](https://img.shields.io/badge/Page-ResTuning-Gree)](https://res-tuning.github.io/) |
| Controllable image synthesis | [🌟SCEdit(CVPR24)](docs/en/tasks/scedit.md) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=SCEdit&color=red&logo=arxiv)](https://arxiv.org/abs/2312.11392) [![Page link](https://img.shields.io/badge/Page-SCEdit-Gree)](https://scedit.github.io/) |
| Image editing | [🌟LAR-Gen](docs/en/tasks/largen.md) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=LARGen&color=red&logo=arxiv)](https://arxiv.org/abs/2403.19534) [![Page link](https://img.shields.io/badge/Page-LARGen-Gree)](https://ali-vilab.github.io/largen-page/) |
| Image editing | [🌟StyleBooth](docs/en/tasks/stylebooth.md) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=StyleBooth&color=red&logo=arxiv)](https://arxiv.org/abs/2404.12154) [![Page link](https://img.shields.io/badge/Page-StyleBooth-Gree)](https://ali-vilab.github.io/stylebooth-page/) |
| Tasks | Methods | Links |
|:----------------------------:|:----------------------------------------------:|:------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
| Text-to-image Generation | SD v1.5 | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
| Text-to-image Generation | SD v2.1 | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
| Text-to-image Generation | SD-XL | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0) |
| Text-to-image Generation | FLUX | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/black-forest-labs/FLUX.1-dev) |
| Efficient Tuning | LoRA | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=LoRA&color=red&logo=arxiv)](https://arxiv.org/abs/2106.09685) |
| Efficient Tuning | Res-Tuning(NeurIPS23) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=Res-Tuing&color=red&logo=arxiv)](https://arxiv.org/abs/2310.19859) [![Page link](https://img.shields.io/badge/Page-ResTuning-Gree)](https://res-tuning.github.io/) |
| Controllable Image Synthesis | [🌟SCEdit(CVPR24)](docs/en/tasks/scedit.md) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=SCEdit&color=red&logo=arxiv)](https://arxiv.org/abs/2312.11392) [![Page link](https://img.shields.io/badge/Page-SCEdit-Gree)](https://scedit.github.io/) |
| Image Editing | [🌟LAR-Gen](docs/en/tasks/largen.md) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=LARGen&color=red&logo=arxiv)](https://arxiv.org/abs/2403.19534) [![Page link](https://img.shields.io/badge/Page-LARGen-Gree)](https://ali-vilab.github.io/largen-page/) |
| Image Editing | [🌟StyleBooth](docs/en/tasks/stylebooth.md) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=StyleBooth&color=red&logo=arxiv)](https://arxiv.org/abs/2404.12154) [![Page link](https://img.shields.io/badge/Page-StyleBooth-Gree)](https://ali-vilab.github.io/stylebooth-page/) |
| Image Generation and Editing | [🌟ACE](https://ali-vilab.github.io/ace-page/) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=ACE&color=red&logo=arxiv)](https://arxiv.org/abs/2410.00086) [![Page link](https://img.shields.io/badge/Page-ACE-Gree)](https://ali-vilab.github.io/ace-page/) [![Demo link](https://img.shields.io/badge/Demo-ACE-purple)](https://huggingface.co/spaces/scepter-studio/ACE-Chat) <br> [![ModelScope link](https://img.shields.io/badge/ModelScope-Model-blue)](https://www.modelscope.cn/models/iic/ACE-0.6B-512px) [![HuggingFace link](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Model-yellow)](https://huggingface.co/scepter-studio/ACE-0.6B-512px) |
## 🖥️ SCEPTER Studio
@@ -145,6 +255,22 @@ Therefore, subsequent startups will become much faster (about one minute) as dow
We deploy a work studio on Modelscope that includes only the inference tab, please refer to [ms_scepter_studio](https://www.modelscope.cn/studios/iic/scepter_studio/summary) and [hf_scepter_studio](https://huggingface.co/spaces/modelscope/scepter_studio)
## ⚙️️ ComfyUI Workflow
### Launch
Manually install by moving custom_nodes to ComfyUI.
```shell
cd path/to/scepter
pip install -e .
cp -r path/to/scepter/workflow/ path/to/ComfyUI/custom_nodes/ComfyUI-Scepter
cd path/to/ComfyUI
python main.py
```
In addition, we also support installation and usage through the ComfyUI Manager.
## 🔍 Learn More
- [Alibaba TongYi Vision Intelligence Lab](https://github.com/ali-vilab)
@@ -177,4 +303,4 @@ This project is licensed under the [Apache License (Version 2.0)](https://github
## Acknowledgement
Thanks to [Stability-AI](https://github.com/Stability-AI), [SWIFT library](https://github.com/modelscope/swift/) and [Fooocus](https://github.com/lllyasviel/Fooocus) for their awesome work.
Thanks to [Stability-AI](https://github.com/Stability-AI), [SWIFT library](https://github.com/modelscope/swift/), [Fooocus](https://github.com/lllyasviel/Fooocus) and [ComfyUI](https://github.com/comfyanonymous/ComfyUI) for their awesome work.
+3 -2
View File
@@ -2,8 +2,8 @@ albumentations
beautifulsoup4
bezier
einops
modelscope==1.14.0
ms-swift>=2.0.1
modelscope
ms-swift
numpy
open_clip_torch
opencv-python
@@ -14,3 +14,4 @@ pyyaml>=5.3.1
scikit-image
torchsde
transformers
scikit-learn
+2 -1
View File
@@ -1,5 +1,6 @@
bitsandbytes
gradio
gradio==4.44.1
gradio_imageslider
imagehash
psutil
tiktoken
+161
View File
@@ -0,0 +1,161 @@
ENV:
BACKEND: nccl
SEED: 2024
#
SOLVER:
NAME: ACESolver
RESUME_FROM:
LOAD_MODEL_ONLY: True
USE_FSDP: False
SHARDING_STRATEGY:
USE_AMP: True
DTYPE: float16
CHANNELS_LAST: True
MAX_STEPS: 500
MAX_EPOCHS: -1
NUM_FOLDS: 1
ACCU_STEP: 1
EVAL_INTERVAL: 50
RESCALE_LR: False
#
WORK_DIR: ./cache/save_data/ace_0.6b_512
LOG_FILE: std_log.txt
#
FILE_SYSTEM:
- NAME: "HuggingfaceFs"
TEMP_DIR: ./cache/cache_data
- NAME: "LocalFs"
TEMP_DIR: ./cache/cache_data
- NAME: "ModelscopeFs"
TEMP_DIR: ./cache/cache_data
#
MODEL:
NAME: LatentDiffusionACE
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
DECODER_BIAS: 0.5
DEFAULT_N_PROMPT:
USE_EMA: True
EVAL_EMA: False
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
USE_TEXT_POS_EMBEDDINGS: True
#
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: eps
MIN_SNR_GAMMA:
NOISE_SCHEDULER:
NAME: LinearScheduler
NUM_TIMESTEPS: 1000
BETA_MIN: 0.0001
BETA_MAX: 0.02
#
DIFFUSION_MODEL:
NAME: ACE
PRETRAINED_MODEL: ms://iic/ACE-0.6B-512px@models/dit/ace_0.6b_512px.pth
IGNORE_KEYS: [ ]
PATCH_SIZE: 2
IN_CHANNELS: 4
HIDDEN_SIZE: 1152
DEPTH: 28
NUM_HEADS: 16
MLP_RATIO: 4.0
PRED_SIGMA: True
DROP_PATH: 0.0
WINDOW_DIZE: 0
Y_CHANNELS: 4096
MAX_SEQ_LEN: 1024
QK_NORM: True
USE_GRAD_CHECKPOINT: True
ATTENTION_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
PRETRAINED_MODEL: ms://iic/ACE-0.6B-512px@models/vae/vae.bin
IGNORE_KEYS: []
#
ENCODER:
NAME: Encoder
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 4
DOUBLE_Z: True
DROPOUT: 0.0
RESAMP_WITH_CONV: True
#
DECODER:
NAME: Decoder
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 4
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://iic/ACE-0.6B-512px@models/text_encoder/t5-v1_1-xxl/
TOKENIZER_PATH: ms://iic/ACE-0.6B-512px@models/tokenizer/t5-v1_1-xxl
LENGTH: 120
T5_DTYPE: bfloat16
ADDED_IDENTIFIER: [ '{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
CLEAN: whitespace
USE_GRAD: False
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 20
GUIDE_SCALE: 4.5
GUIDE_RESCALE: 0.5
#
OPTIMIZER:
NAME: AdamW
LEARNING_RATE: 1e-7
EPS: 1e-10
WEIGHT_DECAY: 5e-4
#
TRAIN_DATA:
NAME: ImageTextPairMSDatasetForACE
MODE: train
MS_DATASET_NAME: cache/datasets/hed_pair
MS_DATASET_NAMESPACE: ""
MS_DATASET_SPLIT: "train"
MS_DATASET_SUBNAME: ""
PROMPT_PREFIX: ""
REPLACE_STYLE: False
MAX_SEQ_LEN: 1024
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 1
SAMPLER:
NAME: LoopSampler
#
TRAIN_HOOKS:
-
NAME: BackwardHook
PRIORITY: 0
-
NAME: LogHook
LOG_INTERVAL: 50
-
NAME: CheckpointHook
INTERVAL: 100
-
NAME: ProbeDataHook
PROB_INTERVAL: 100
@@ -0,0 +1,303 @@
ENV:
BACKEND: nccl
SEED: 166666
SOLVER:
# NAME DESCRIPTION: TYPE: default: 'LatentUfitSolver'
NAME: LatentDiffusionSolver
# MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
MAX_STEPS: 100000
# USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False
USE_AMP: True
# DTYPE DESCRIPTION: The precision for training. TYPE: str default: 'float32'
DTYPE: bfloat16
# USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
USE_FAIRSCALE: False
# USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
USE_FSDP: True
# LOAD_MODEL_ONLY DESCRIPTION: Only load the model rather than the optimizer and schedule, default is False. TYPE: bool default: False
LOAD_MODEL_ONLY: False
# RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: ''
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_flux_dev_1024_lora
LOG_FILE: std_log.txt
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
EVAL_INTERVAL: 100
# LOG_TRAIN_NUM DESCRIPTION: The number samples used to log in training phase. TYPE: int default: -1
LOG_TRAIN_NUM: 16
# FSDP_REDUCE_DTYPE DESCRIPTION: The dtype of reduce in FSDP. TYPE: str default: 'float16'
FSDP_REDUCE_DTYPE: float32
# FSDP_BUFFER_DTYPE DESCRIPTION: The dtype of buffer in FSDP. TYPE: str default: 'float16'
FSDP_BUFFER_DTYPE: float32
# FSDP_SHARD_MODULES DESCRIPTION: The modules to be sharded in FSDP. TYPE: list default: ['model']
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ] #
SAVE_MODULES: [ 'model'] #
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
FREEZE:
TUNER:
- NAME: SwiftLoRA
R: 4
LORA_ALPHA: 4
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "(model.double_blocks.*(.qkv|.proj|.img_mod.lin|.txt_mod.lin))|(model.single_blocks.*(.linear1|.linear2|.modulation.lin))$"
##
MODEL:
NAME: LatentDiffusionFlux
PARAMETERIZATION: rf
TIMESTEPS: 1000
MIN_SNR_GAMMA:
ZERO_TERMINAL_SNR: False
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
NOISE_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
NAME: FlowMatchSigmaScheduler
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
WEIGHTING_SCHEME: logit_normal
SHIFT: 3.0
# LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
LOGIT_MEAN: 0.0
# LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
LOGIT_STD: 1.0
# MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
MODE_SCALE: 1.29
SAMPLER_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
NAME: FlowMatchFluxShiftScheduler
# SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
SHIFT: False
# SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
SIGMOID_SCALE: 1
# BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
BASE_SHIFT: 0.5
# MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
MAX_SHIFT: 1.15
#
DIFFUSION_MODEL:
# NAME DESCRIPTION: TYPE: default: 'Flux'
NAME: Flux
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@flux1-dev.safetensors
# IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
IN_CHANNELS: 64
# HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
HIDDEN_SIZE: 3072
# NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
NUM_HEADS: 24
# AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
AXES_DIM: [ 16, 56, 56 ]
# THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
THETA: 10000
# VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
VEC_IN_DIM: 768
# GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
GUIDANCE_EMBED: True
# CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
CONTEXT_IN_DIM: 4096
# MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
MLP_RATIO: 4.0
# QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
QKV_BIAS: True
# DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
DEPTH: 19
# DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
DEPTH_SINGLE_BLOCKS: 38
USE_GRAD_CHECKPOINT: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLFlux
EMBED_DIM: 16
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@ae.safetensors
IGNORE_KEYS: [ ]
BATCH_SIZE: 8
USE_CONV: False
SCALE_FACTOR: 0.3611
SHIFT_FACTOR: 0.1159
#
ENCODER:
NAME: Encoder
USE_CHECKPOINT: True
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 16
DOUBLE_Z: True
DROPOUT: 0.0
RESAMP_WITH_CONV: True
#
DECODER:
NAME: Decoder
USE_CHECKPOINT: True
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 16
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
NAME: T5PlusClipFluxEmbedder
# T5_MODEL DESCRIPTION: TYPE: default: ''
T5_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: T5EncoderModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder_2/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: T5Tokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer_2/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 512
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: last_hidden_state
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: False
CLEAN: whitespace
# CLIP_MODEL DESCRIPTION: TYPE: default: ''
CLIP_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: CLIPTextModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: CLIPTokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 77
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: pooler_output
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: True
CLEAN: whitespace
USE_GRAD_CHECKPOINT: True
#
SAMPLE_ARGS:
SAMPLE_STEPS: 50
SAMPLER: flow_eluer
SEED: 2024
IMAGE_SIZE: [ 1024, 1024 ]
GUIDE_SCALE: 3.5
#
OPTIMIZER:
NAME: AdamW
LEARNING_RATE: 4e-4
BETAS: [ 0.9, 0.999 ]
EPS: 1e-8
WEIGHT_DECAY: 1e-2
AMSGRAD: False
#
TRAIN_DATA:
NAME: ImageTextPairMSDataset
MODE: train
MS_DATASET_NAME: style_custom_dataset
MS_DATASET_NAMESPACE: damo
MS_DATASET_SUBNAME: 3D
PROMPT_PREFIX: ""
MS_DATASET_SPLIT: train
MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
REPLACE_STYLE: False
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
SAMPLER:
NAME: LoopSampler
TRANSFORMS:
- NAME: LoadImageFromFile
RGB_ORDER: RGB
BACKEND: pillow
- NAME: FlexibleResize
INTERPOLATION: bilinear
SIZE: [ 1024, 1024 ]
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: FlexibleCenterCrop
SIZE: [ 1024, 1024 ]
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: [ 'data_key' ]
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
PROMPT_FILE:
PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ]
IMAGE_SIZE: [ 1024, 1024 ]
FIELDS: [ "prompt" ]
DELIMITER: '#;#'
PROMPT_PREFIX: ''
PIN_MEMORY: True
BATCH_SIZE: 2
NUM_WORKERS: 4
TRANSFORMS:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
# GRADIENT_CLIP: 1.0
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
-
NAME: TensorboardLogHook
-
NAME: CheckpointHook
INTERVAL: 10000
PRIORITY: 200
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
DISABLE_SNAPSHOT: True
EVAL_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
@@ -0,0 +1,302 @@
ENV:
BACKEND: nccl
SEED: 166666
SOLVER:
# NAME DESCRIPTION: TYPE: default: 'LatentUfitSolver'
NAME: LatentDiffusionSolver
# MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
MAX_STEPS: 100000
# USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False
USE_AMP: True
# DTYPE DESCRIPTION: The precision for training. TYPE: str default: 'float32'
DTYPE: bfloat16
# USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
USE_FAIRSCALE: False
# USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
USE_FSDP: True
# LOAD_MODEL_ONLY DESCRIPTION: Only load the model rather than the optimizer and schedule, default is False. TYPE: bool default: False
LOAD_MODEL_ONLY: False
# RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: ''
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_flux_schnell_1024_lora
LOG_FILE: std_log.txt
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
EVAL_INTERVAL: 100
# LOG_TRAIN_NUM DESCRIPTION: The number samples used to log in training phase. TYPE: int default: -1
LOG_TRAIN_NUM: 16
# FSDP_REDUCE_DTYPE DESCRIPTION: The dtype of reduce in FSDP. TYPE: str default: 'float16'
FSDP_REDUCE_DTYPE: float32
# FSDP_BUFFER_DTYPE DESCRIPTION: The dtype of buffer in FSDP. TYPE: str default: 'float16'
FSDP_BUFFER_DTYPE: float32
# FSDP_SHARD_MODULES DESCRIPTION: The modules to be sharded in FSDP. TYPE: list default: ['model']
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ] #
SAVE_MODULES: [ 'model'] #
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
FREEZE:
TUNER:
- NAME: SwiftLoRA
R: 4
LORA_ALPHA: 4
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "(model.double_blocks.*(.qkv|.proj|.img_mod.lin|.txt_mod.lin))|(model.single_blocks.*(.linear1|.linear2|.modulation.lin))$"
##
MODEL:
NAME: LatentDiffusionFlux
PARAMETERIZATION: rf
TIMESTEPS: 1000
MIN_SNR_GAMMA:
ZERO_TERMINAL_SNR: False
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
NOISE_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
NAME: FlowMatchSigmaScheduler
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
WEIGHTING_SCHEME: logit_normal
SHIFT: 3.0
# LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
LOGIT_MEAN: 0.0
# LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
LOGIT_STD: 1.0
# MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
MODE_SCALE: 1.29
SAMPLER_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
NAME: FlowMatchFluxShiftScheduler
# SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
SHIFT: False
# SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
SIGMOID_SCALE: 1
# BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
BASE_SHIFT: 0.5
# MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
MAX_SHIFT: 1.15
#
DIFFUSION_MODEL:
# NAME DESCRIPTION: TYPE: default: 'Flux'
NAME: Flux
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-schnell@flux1-schnell.safetensors
# IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
IN_CHANNELS: 64
# HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
HIDDEN_SIZE: 3072
# NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
NUM_HEADS: 24
# AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
AXES_DIM: [ 16, 56, 56 ]
# THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
THETA: 10000
# VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
VEC_IN_DIM: 768
# GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
GUIDANCE_EMBED: False
# CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
CONTEXT_IN_DIM: 4096
# MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
MLP_RATIO: 4.0
# QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
QKV_BIAS: True
# DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
DEPTH: 19
# DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
DEPTH_SINGLE_BLOCKS: 38
USE_GRAD_CHECKPOINT: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLFlux
EMBED_DIM: 16
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-schnell@ae.safetensors
IGNORE_KEYS: [ ]
BATCH_SIZE: 8
USE_CONV: False
SCALE_FACTOR: 0.3611
SHIFT_FACTOR: 0.1159
#
ENCODER:
NAME: Encoder
USE_CHECKPOINT: True
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 16
DOUBLE_Z: True
DROPOUT: 0.0
RESAMP_WITH_CONV: True
#
DECODER:
NAME: Decoder
USE_CHECKPOINT: True
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 16
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
NAME: T5PlusClipFluxEmbedder
# T5_MODEL DESCRIPTION: TYPE: default: ''
T5_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: T5EncoderModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder_2/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: T5Tokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer_2/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 256
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: last_hidden_state
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: False
CLEAN: whitespace
# CLIP_MODEL DESCRIPTION: TYPE: default: ''
CLIP_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: CLIPTextModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: CLIPTokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 77
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: pooler_output
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: True
CLEAN: whitespace
#
SAMPLE_ARGS:
SAMPLE_STEPS: 4
SAMPLER: flow_eluer
SEED: 2024
IMAGE_SIZE: [ 1024, 1024 ]
GUIDE_SCALE: 3.5
#
OPTIMIZER:
NAME: AdamW
LEARNING_RATE: 4e-4
BETAS: [ 0.9, 0.999 ]
EPS: 1e-8
WEIGHT_DECAY: 1e-2
AMSGRAD: False
#
TRAIN_DATA:
NAME: ImageTextPairMSDataset
MODE: train
MS_DATASET_NAME: style_custom_dataset
MS_DATASET_NAMESPACE: damo
MS_DATASET_SUBNAME: 3D
PROMPT_PREFIX: ""
MS_DATASET_SPLIT: train
MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
REPLACE_STYLE: False
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
SAMPLER:
NAME: LoopSampler
TRANSFORMS:
- NAME: LoadImageFromFile
RGB_ORDER: RGB
BACKEND: pillow
- NAME: FlexibleResize
INTERPOLATION: bilinear
SIZE: [ 1024, 1024 ]
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: FlexibleCenterCrop
SIZE: [ 1024, 1024 ]
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: [ 'data_key' ]
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
PROMPT_FILE:
PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ]
IMAGE_SIZE: [ 1024, 1024 ]
FIELDS: [ "prompt" ]
DELIMITER: '#;#'
PROMPT_PREFIX: ''
PIN_MEMORY: True
BATCH_SIZE: 2
NUM_WORKERS: 4
TRANSFORMS:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
# GRADIENT_CLIP: 1.0
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
-
NAME: TensorboardLogHook
-
NAME: CheckpointHook
INTERVAL: 10000
PRIORITY: 200
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
DISABLE_SNAPSHOT: True
EVAL_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
@@ -168,13 +168,13 @@ SOLVER:
RGB_ORDER: RGB
BACKEND: pillow
- NAME: Resize
SIZE: 768
SIZE: 512
INTERPOLATION: bilinear
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: CenterCrop
SIZE: 768
SIZE: 512
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
@@ -216,13 +216,13 @@ SOLVER:
RGB_ORDER: RGB
BACKEND: pillow
- NAME: Resize
SIZE: 768
SIZE: 512
INTERPOLATION: bilinear
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: CenterCrop
SIZE: 768
SIZE: 512
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
@@ -0,0 +1,25 @@
WORK_DIR: chatbot
FILE_SYSTEM:
- NAME: LocalFs
TEMP_DIR: ./cache/cache_data
- NAME: ModelscopeFs
TEMP_DIR: ./cache/cache_data
- NAME: HuggingfaceFs
TEMP_DIR: ./cache/cache_data
#
ENABLE_I2V: False
#
MODEL:
EDIT_MODEL:
MODEL_CFG_DIR: scepter/methods/studio/chatbot/models/
DEFAULT: ace_0.6b_512
I2V:
MODEL_NAME: CogVideoX-5b-I2V
MODEL_DIR: ms://ZhipuAI/CogVideoX-5b-I2V/
CAPTIONER:
MODEL_NAME: InternVL2-2B
MODEL_DIR: ms://OpenGVLab/InternVL2-2B/
PROMPT: '<image>\nThis image is the first frame of a video. Based on this image, please imagine what changes may occur in the next few seconds of the video. Please output brief description, such as "a dog running" or "a person turns to left". No more than 30 words.'
ENHANCER:
MODEL_NAME: Meta-Llama-3.1-8B-Instruct
MODEL_DIR: ms://LLM-Research/Meta-Llama-3.1-8B-Instruct/
@@ -0,0 +1,127 @@
NAME: ACE_0.6B_512
IS_DEFAULT: False
DEFAULT_PARAS:
PARAS:
#
INPUT:
INPUT_IMAGE:
INPUT_MASK:
TASK:
PROMPT: ""
NEGATIVE_PROMPT: ""
OUTPUT_HEIGHT: 512
OUTPUT_WIDTH: 512
SAMPLER: ddim
SAMPLE_STEPS: 20
GUIDE_SCALE: 4.5
GUIDE_RESCALE: 0.5
SEED: -1
TAR_INDEX: 0
OUTPUT:
LATENT:
IMAGES:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
- NAME: encode
DTYPE: float16
INPUT: ["IMAGE"]
- NAME: decode
DTYPE: float16
INPUT: ["LATENT"]
#
DIFFUSION_MODEL:
FUNCTION:
- NAME: forward
DTYPE: float16
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE"]
#
COND_STAGE_MODEL:
FUNCTION:
- NAME: encode_list
DTYPE: bfloat16
INPUT: ["PROMPT"]
#
MODEL:
NAME: LatentDiffusionACE
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
DECODER_BIAS: 0.5
DEFAULT_N_PROMPT: ""
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
USE_TEXT_POS_EMBEDDINGS: True
#
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: eps
MIN_SNR_GAMMA:
NOISE_SCHEDULER:
NAME: LinearScheduler
NUM_TIMESTEPS: 1000
BETA_MIN: 0.0001
BETA_MAX: 0.02
#
DIFFUSION_MODEL:
NAME: ACE
PRETRAINED_MODEL: ms://iic/ACE-0.6B-512px@models/dit/ace_0.6b_512px.pth
IGNORE_KEYS: [ ]
PATCH_SIZE: 2
IN_CHANNELS: 4
HIDDEN_SIZE: 1152
DEPTH: 28
NUM_HEADS: 16
MLP_RATIO: 4.0
PRED_SIGMA: True
DROP_PATH: 0.0
WINDOW_DIZE: 0
Y_CHANNELS: 4096
MAX_SEQ_LEN: 1024
QK_NORM: True
USE_GRAD_CHECKPOINT: True
ATTENTION_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
PRETRAINED_MODEL: ms://iic/ACE-0.6B-512px@models/vae/vae.bin
IGNORE_KEYS: []
#
ENCODER:
NAME: Encoder
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 4
DOUBLE_Z: True
DROPOUT: 0.0
RESAMP_WITH_CONV: True
#
DECODER:
NAME: Decoder
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 4
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://iic/ACE-0.6B-512px@models/text_encoder/t5-v1_1-xxl/
TOKENIZER_PATH: ms://iic/ACE-0.6B-512px@models/tokenizer/t5-v1_1-xxl
LENGTH: 120
T5_DTYPE: bfloat16
ADDED_IDENTIFIER: [ '{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
CLEAN: whitespace
USE_GRAD: False
@@ -0,0 +1,201 @@
NAME: FLUX1.0_DEV
IS_DEFAULT: False
DEFAULT_PARAS:
PARAS:
RESOLUTIONS: [[1024, 1024]]
INPUT:
IMAGE:
ORIGINAL_SIZE_AS_TUPLE: [1024, 1024]
TARGET_SIZE_AS_TUPLE: [1024, 1024]
PROMPT: ""
NEGATIVE_PROMPT:
DEFAULT: ""
VISIBLE: False
PROMPT_PREFIX: ""
SAMPLE:
VALUES: ["flow_eluer"]
DEFAULT: "flow_eluer"
SAMPLE_STEPS: 50
GUIDE_SCALE: 3.5
GUIDE_RESCALE:
DEFAULT: 0.0
VISIBLE: False
DISCRETIZATION:
VALUES: []
DEFAULT:
VISIBLE: False
OUTPUT:
LATENT:
IMAGES:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
-
NAME: encode
DTYPE: bfloat16
INPUT: ["IMAGE"]
-
NAME: decode
DTYPE: bfloat16
INPUT: ["LATENT"]
PARAS:
SCALE_FACTOR: 1.5305
SHIFT_FACTOR: 0.0609
SIZE_FACTOR: 8
DIFFUSION_MODEL:
FUNCTION:
-
NAME: forward
DTYPE: bfloat16
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE"]
COND_STAGE_MODEL:
FUNCTION:
-
NAME: encode
DTYPE: bfloat16
INPUT: ["PROMPT"]
#
MODEL:
NAME: LatentDiffusionFlux
PARAMETERIZATION: rf
TIMESTEPS: 1000
MIN_SNR_GAMMA:
ZERO_TERMINAL_SNR: False
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
NOISE_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
NAME: FlowMatchSigmaScheduler
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
WEIGHTING_SCHEME: logit_normal
SHIFT: 3.0
# LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
LOGIT_MEAN: 0.0
# LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
LOGIT_STD: 1.0
# MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
MODE_SCALE: 1.29
#
DIFFUSION_MODEL:
# NAME DESCRIPTION: TYPE: default: 'Flux'
NAME: Flux
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@flux1-dev.safetensors
# IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
IN_CHANNELS: 64
# HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
HIDDEN_SIZE: 3072
# NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
NUM_HEADS: 24
# AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
AXES_DIM: [ 16, 56, 56 ]
# THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
THETA: 10000
# VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
VEC_IN_DIM: 768
# GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
GUIDANCE_EMBED: True
# CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
CONTEXT_IN_DIM: 4096
# MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
MLP_RATIO: 4.0
# QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
QKV_BIAS: True
# DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
DEPTH: 19
# DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
DEPTH_SINGLE_BLOCKS: 38
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLFlux
EMBED_DIM: 16
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@ae.safetensors
IGNORE_KEYS: [ ]
BATCH_SIZE: 8
USE_CONV: False
SCALE_FACTOR: 0.3611
SHIFT_FACTOR: 0.1159
#
ENCODER:
NAME: Encoder
USE_CHECKPOINT: True
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 16
DOUBLE_Z: True
DROPOUT: 0.0
RESAMP_WITH_CONV: True
#
DECODER:
NAME: Decoder
USE_CHECKPOINT: True
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 16
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
NAME: T5PlusClipFluxEmbedder
# T5_MODEL DESCRIPTION: TYPE: default: ''
T5_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: T5EncoderModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder_2/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: T5Tokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer_2/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 512
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: last_hidden_state
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: False
CLEAN: whitespace
# CLIP_MODEL DESCRIPTION: TYPE: default: ''
CLIP_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: CLIPTextModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: CLIPTokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 77
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: pooler_output
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: True
CLEAN: whitespace
@@ -0,0 +1,212 @@
NAME: FLUX1.0_SCHNELL
IS_DEFAULT: False
DEFAULT_PARAS:
PARAS:
RESOLUTIONS: [[1024, 1024]]
INPUT:
IMAGE:
ORIGINAL_SIZE_AS_TUPLE: [1024, 1024]
TARGET_SIZE_AS_TUPLE: [1024, 1024]
PROMPT: ""
NEGATIVE_PROMPT:
DEFAULT: ""
VISIBLE: False
PROMPT_PREFIX: ""
SAMPLE:
VALUES: ["flow_eluer"]
DEFAULT: "flow_eluer"
SAMPLE_STEPS: 4
GUIDE_SCALE: 3.5
GUIDE_RESCALE:
DEFAULT: 0.0
VISIBLE: False
DISCRETIZATION:
VALUES: []
DEFAULT:
VISIBLE: False
OUTPUT:
LATENT:
IMAGES:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
-
NAME: encode
DTYPE: bfloat16
INPUT: ["IMAGE"]
-
NAME: decode
DTYPE: bfloat16
INPUT: ["LATENT"]
PARAS:
SCALE_FACTOR: 1.5305
SHIFT_FACTOR: 0.0609
SIZE_FACTOR: 8
DIFFUSION_MODEL:
FUNCTION:
-
NAME: forward
DTYPE: bfloat16
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE"]
COND_STAGE_MODEL:
FUNCTION:
-
NAME: encode
DTYPE: bfloat16
INPUT: ["PROMPT"]
#
MODEL:
NAME: LatentDiffusionFlux
PARAMETERIZATION: rf
TIMESTEPS: 1000
MIN_SNR_GAMMA:
ZERO_TERMINAL_SNR: False
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
NOISE_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
NAME: FlowMatchSigmaScheduler
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
WEIGHTING_SCHEME: logit_normal
SHIFT: 3.0
# LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
LOGIT_MEAN: 0.0
# LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
LOGIT_STD: 1.0
# MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
MODE_SCALE: 1.29
SAMPLER_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
NAME: FlowMatchFluxShiftScheduler
# SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
SHIFT: False
# SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
SIGMOID_SCALE: 1
# BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
BASE_SHIFT: 0.5
# MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
MAX_SHIFT: 1.15
#
DIFFUSION_MODEL:
# NAME DESCRIPTION: TYPE: default: 'Flux'
NAME: Flux
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-schnell@flux1-schnell.safetensors
# IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
IN_CHANNELS: 64
# HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
HIDDEN_SIZE: 3072
# NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
NUM_HEADS: 24
# AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
AXES_DIM: [ 16, 56, 56 ]
# THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
THETA: 10000
# VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
VEC_IN_DIM: 768
# GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
GUIDANCE_EMBED: False
# CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
CONTEXT_IN_DIM: 4096
# MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
MLP_RATIO: 4.0
# QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
QKV_BIAS: True
# DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
DEPTH: 19
# DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
DEPTH_SINGLE_BLOCKS: 38
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLFlux
EMBED_DIM: 16
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-schnell@ae.safetensors
IGNORE_KEYS: [ ]
BATCH_SIZE: 8
USE_CONV: False
SCALE_FACTOR: 0.3611
SHIFT_FACTOR: 0.1159
#
ENCODER:
NAME: Encoder
USE_CHECKPOINT: True
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 16
DOUBLE_Z: True
DROPOUT: 0.0
RESAMP_WITH_CONV: True
#
DECODER:
NAME: Decoder
USE_CHECKPOINT: True
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 16
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
NAME: T5PlusClipFluxEmbedder
# T5_MODEL DESCRIPTION: TYPE: default: ''
T5_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: T5EncoderModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder_2/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: T5Tokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer_2/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 256
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: last_hidden_state
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: False
CLEAN: whitespace
# CLIP_MODEL DESCRIPTION: TYPE: default: ''
CLIP_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: CLIPTextModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: CLIPTokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 77
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: pooler_output
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: True
CLEAN: whitespace
@@ -6,65 +6,81 @@ DIFFUSION_PARAS:
'dpm2_karras', 'dpm2_ancestral_karras', 'dpmpp_2s_ancestral_karras', 'dpmpp_2m_karras',
'dpmpp_sde_karras', 'dpmpp_2m_sde_karras']
DEFAULT: 'dpmpp_2s_ancestral'
VISIBLE: True
NEGATIVE_PROMPT:
DEFAULT:
VISIBLE: True
PROMPT_PREFIX:
DEFAULT:
VISIBLE: True
SAMPLES:
MIN: 1
MAX: 4
DEFAULT: 1
VISIBLE: True
SAMPLE_STEPS:
MIN: 1
MAX: 100
DEFAULT: 30
VISIBLE: True
GUIDE_SCALE:
MIN: 0
MAX: 10
DEFAULT: 5.0
VISIBLE: True
GUIDE_RESCALE:
MIN: 0
MAX: 1.0
DEFAULT: 0.5
VISIBLE: True
DISCRETIZATION:
VALUES: ["trailing", "leading", "linspace"]
DEFAULT: "linspace"
VISIBLE: True
REFINE_SAMPLERS:
VALUES: [ 'ddim', 'euler', 'euler_ancestral', 'heun', 'dpm2',
'dpm2_ancestral', 'dpmpp_2m', 'dpmpp_sde', 'dpmpp_2m_sde', 'dpmpp_2s_ancestral',
'dpm2_karras', 'dpm2_ancestral_karras', 'dpmpp_2s_ancestral_karras', 'dpmpp_2m_karras',
'dpmpp_sde_karras', 'dpmpp_2m_sde_karras' ]
DEFAULT: 'dpmpp_2s_ancestral'
VISIBLE: True
REFINE_SAMPLE_STEPS:
MIN: 0
MAX: 100
DEFAULT: 30
VISIBLE: True
REFINE_GUIDE_SCALE:
MIN: 0
MAX: 10
DEFAULT: 5.0
VISIBLE: True
REFINE_GUIDE_RESCALE:
MIN: 0
MAX: 1.0
DEFAULT: 0.5
VISIBLE: True
REFINE_DISCRETIZATION:
VALUES: [ "trailing", "leading", "linspace" ]
DEFAULT: "linspace"
VISIBLE: True
AESTHETIC_SCORE:
MIN: 0.0
MAX: 10.0
DEFAULT: 6.0
VISIBLE: True
NEGATIVE_AESTHETIC_SCORE:
MIN: 0.0
MAX: 10.0
DEFAULT: 2.5
VISIBLE: True
REFINE_STRENGTH:
MIN: 0
MAX: 1.0
DEFAULT: 0.15
RESOLUTIONS:
VALUES: [ [512, 512], [768, 768],
[704, 1408], [704, 1344], [768, 1344],
VALUES: [
[512, 512], [768, 768],
[704, 1408], [704, 1344], [768, 1344],
[720, 1280],
[768, 1280], [832, 1216], [832, 1152],
[896, 1152], [896, 1088], [960, 1088],
@@ -74,8 +90,13 @@ DIFFUSION_PARAS:
[1280, 768],
[1344, 768], [1344, 704], [1408, 704],
[1472, 704], [1536, 640], [1600, 640],
[1664, 576], [1728, 576]]
[1664, 576], [1728, 576],
[2048, 2048], [2048, 1920], [1920, 2048],
[1536, 2560], [2560, 1536], [2560, 1440],
[2560, 1440]
]
DEFAULT: [1024, 1024]
VISIBLE: True
EXTENSION_PARAS:
MANTRA_BOOK: scepter/methods/studio/extensions/mantra_book/mantra_book.yaml
OFFICIAL_TUNERS: scepter/methods/studio/extensions/tuners/official_tuners.yaml
+361 -82
View File
@@ -2,11 +2,10 @@ WORK_DIR: datasets
EXPORT_DIR: export_datasets
FILE_SYSTEM:
-
# NAME DESCRIPTION: TYPE: default: ''
NAME: LocalFs
AUTO_CLEAN: False
PROCESSORS:
# Caption processor
- NAME: BlipImageBase
TYPE: caption
MODEL_PATH: ms://cubeai/blip-image-captioning-base
@@ -15,6 +14,81 @@ PROCESSORS:
PARAS:
- LANGUAGE_NAME: English
LANGUAGE_ZH_NAME: 英语
- NAME: InternVL15
TYPE: caption
MODEL_PATH: ms://AI-ModelScope/InternVL-Chat-V1-5
DEVICE: "gpu"
MEMORY: 49968
PARAS:
- PROMPT: 用中文描述这张图片
LANGUAGE_NAME: Chinese
LANGUAGE_ZH_NAME: 中文
- PROMPT: Generate the caption in English
LANGUAGE_NAME: English
LANGUAGE_ZH_NAME: 英语
- NAME: QWVLQuantize
TYPE: caption
DEVICE: "gpu"
MEMORY: 7885
MODEL_PATH: ms://qwen/Qwen-VL:v1.0.3
PARAS:
- PROMPT: 用中文描述这张图片
LANGUAGE_NAME: Chinese
LANGUAGE_ZH_NAME: 中文
MAX_NEW_TOKENS:
VALUE: 1024
MAX: 2048
STEP: 128
MIN: 256
MIN_NEW_TOKENS:
VALUE: 16
MAX: 1024
STEP: 16
MIN: 0
NUM_BEAMS:
VALUE: 1
MAX: 12
STEP: 1
MIN: 1
REPETITION_PENALTY:
VALUE: 1.0
MAX: 100.0
STEP: 1.0
MIN: 1.0
TEMPERATURE:
VALUE: 1.0
MAX: 100.0
STEP: 1.0
MIN: 1.0
- PROMPT: Generate the caption in English
LANGUAGE_NAME: English
LANGUAGE_ZH_NAME: 英语
MAX_NEW_TOKENS:
VALUE: 1024
MAX: 2048
STEP: 128
MIN: 256
MIN_NEW_TOKENS:
VALUE: 16
MAX: 1024
STEP: 16
MIN: 0
NUM_BEAMS:
VALUE: 1
MAX: 12
STEP: 1
MIN: 1
REPETITION_PENALTY:
VALUE: 1.0
MAX: 100.0
STEP: 1.0
MIN: 1.0
TEMPERATURE:
VALUE: 1.0
MAX: 100.0
STEP: 1.0
MIN: 1.0
- NAME: QWVL
TYPE: caption
MODEL_PATH: ms://qwen/Qwen-VL:v1.0.3
@@ -77,75 +151,14 @@ PROCESSORS:
MAX: 100.0
STEP: 1.0
MIN: 1.0
-
NAME: QWVLQuantize
TYPE: caption
DEVICE: "gpu"
MEMORY: 7885
MODEL_PATH: ms://qwen/Qwen-VL:v1.0.3
PARAS:
- PROMPT: 用中文描述这张图片
LANGUAGE_NAME: Chinese
LANGUAGE_ZH_NAME: 中文
MAX_NEW_TOKENS:
VALUE: 1024
MAX: 2048
STEP: 128
MIN: 256
MIN_NEW_TOKENS:
VALUE: 16
MAX: 1024
STEP: 16
MIN: 0
NUM_BEAMS:
VALUE: 1
MAX: 12
STEP: 1
MIN: 1
REPETITION_PENALTY:
VALUE: 1.0
MAX: 100.0
STEP: 1.0
MIN: 1.0
TEMPERATURE:
VALUE: 1.0
MAX: 100.0
STEP: 1.0
MIN: 1.0
- PROMPT: Generate the caption in English
LANGUAGE_NAME: English
LANGUAGE_ZH_NAME: 英语
MAX_NEW_TOKENS:
VALUE: 1024
MAX: 2048
STEP: 128
MIN: 256
MIN_NEW_TOKENS:
VALUE: 16
MAX: 1024
STEP: 16
MIN: 0
NUM_BEAMS:
VALUE: 1
MAX: 12
STEP: 1
MIN: 1
REPETITION_PENALTY:
VALUE: 1.0
MAX: 100.0
STEP: 1.0
MIN: 1.0
TEMPERATURE:
VALUE: 1.0
MAX: 100.0
STEP: 1.0
MIN: 1.0
-
NAME: CenterCrop
# Simple processor
- NAME: CenterCrop
TYPE: image
DEVICE: "cpu"
MEMORY: 10
PARAS:
CAPTION_INTERACTIVE: False
HEIGHT_RATIO:
VALUE: 1
MAX: 20
@@ -156,18 +169,284 @@ PROCESSORS:
MAX: 20
STEP: 1
MIN: 1
# - NAME: PaddingCrop
# TYPE: image
# DEVICE: "cpu"
# MEMORY: 10
# PARAS:
# HEIGHT_RATIO:
# VALUE: 3
# MAX: 25
# STEP: 1
# MIN: 1
# WIDTH_RATIO:
# VALUE: 4
# MAX: 20
# STEP: 1
# MIN: 1
- NAME: ChangeSample
TYPE: image
DEVICE: "cpu"
MEMORY: 10
PARAS:
CAPTION_INTERACTIVE: False
SRC_IMAGE_INTERACTIVE: True
SRC_IMAGE_MASK_INTERACTIVE: True
TARGET_IMAGE_INTERACTIVE: True
PREVIEW_BTN_VISIBLE: False
- NAME: MaskEditSample
TYPE: image
DEVICE: "cpu"
MEMORY: 10
PARAS:
SRC_IMAGE_INTERACTIVE: True
SRC_IMAGE_TOOL: sketch
CAPTION_INTERACTIVE: False
- NAME: SourceMaskSample
TYPE: image
DEVICE: "cpu"
MEMORY: 10
PARAS:
SRC_IMAGE_INTERACTIVE: True
SRC_IMAGE_TOOL: sketch
CAPTION_INTERACTIVE: False
- NAME: MaskSwapEditSample
TYPE: image
DEVICE: "cpu"
MEMORY: 10
PARAS:
PREVIEW_BTN_VISIBLE: False
SRC_IMAGE_INTERACTIVE: True
SRC_IMAGE_TOOL: sketch
CAPTION_INTERACTIVE: False
- NAME: SwapMaskSwapEditSample
TYPE: image
DEVICE: "cpu"
MEMORY: 10
PARAS:
PREVIEW_BTN_VISIBLE: False
SRC_IMAGE_INTERACTIVE: True
SRC_IMAGE_TOOL: sketch
CAPTION_INTERACTIVE: False
- NAME: SwapSample
TYPE: image
DEVICE: "cpu"
MEMORY: 10
PARAS:
PREVIEW_BTN_VISIBLE: True
SRC_IMAGE_INTERACTIVE: True
TARGET_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
- NAME: DefaultMaskSample
TYPE: image
DEVICE: "cpu"
MEMORY: 10
PARAS:
PREVIEW_BTN_VISIBLE: True
SRC_IMAGE_INTERACTIVE: False
TARGET_IMAGE_INTERACTIVE: False
CAPTION_INTERACTIVE: False
# - NAME: PaddingCrop
# TYPE: image
# DEVICE: "cpu"
# MEMORY: 10
# PARAS:
# HEIGHT_RATIO:
# VALUE: 3
# MAX: 25
# STEP: 1
# MIN: 1
# WIDTH_RATIO:
# VALUE: 4
# MAX: 20
# STEP: 1
# MIN: 1
# Annotator processor
- NAME: CannyExtractor
TYPE: image
DEVICE: "cpu"
MEMORY: 10
MODEL:
NAME: "CannyAnnotator"
LOW_THRESHOLD: 100
HIGH_THRESHOLD: 200
PARAS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
- NAME: ColorExtractor
TYPE: image
DEVICE: "cpu"
MEMORY: 10
MODEL:
NAME: "ColorAnnotator"
RATIO: 64
PARAS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
- NAME: InfoDrawContourExtractor
TYPE: image
DEVICE: "gpu"
MEMORY: 10
MODEL:
NAME: "InfoDrawContourAnnotator"
INPUT_NC: 3
OUTPUT_NC: 1
N_RESIDUAL_BLOCKS: 3
SIGMOID: True
PRETRAINED_MODEL: "ms://iic/scepter_annotator@annotator/ckpts/informative_drawing_contour_style.pth"
PARAS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
- NAME: DegradationExtractor
TYPE: image
DEVICE: "cpu"
MEMORY: 10
MODEL:
NAME: "DegradationAnnotator"
RANDOM_DEGRADATION: True
PARAMS:
gaussian_noise: { }
resize: { 'scale': [ 0.4, 0.8 ] }
jpeg: { 'jpeg_level': [ 25, 75 ] }
gaussian_blur: { 'kernel_size': [ 7, 9, 11, 13, 15 ], 'sigma': [ 0.9, 1.8 ] }
PARAS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
- NAME: MidasExtractor
TYPE: image
DEVICE: "gpu"
MEMORY: 10
MODEL:
NAME: "MidasDetector"
PRETRAINED_MODEL: "ms://iic/scepter_annotator@annotator/ckpts/dpt_hybrid-midas-501f0c75.pt"
PARAS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
- NAME: DoodleExtractor
TYPE: image
DEVICE: "gpu"
MEMORY: 10
MODEL:
NAME: "DoodleAnnotator"
PROCESSOR_TYPE: "pidinet_sketch"
PROCESSOR_CFG:
- NAME: "PiDiAnnotator"
PRETRAINED_MODEL: "ms://iic/scepter_annotator@annotator/ckpts/table5_pidinet.pth"
- NAME: "SketchAnnotator"
PRETRAINED_MODEL: "ms://iic/scepter_annotator@annotator/ckpts/sketch_simplification_gan.pth"
PARAS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
- NAME: GrayExtractor
TYPE: image
DEVICE: "cpu"
MEMORY: 10
MODEL:
NAME: "GrayAnnotator"
PARAS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
- NAME: InpaintingExtractor
TYPE: image
DEVICE: "cpu"
MEMORY: 10
MODEL:
NAME: "InpaintingAnnotator"
RETURN_MASK: False
PARAS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
- NAME: InpaintingSourceExtractor
TYPE: image
DEVICE: "cpu"
MEMORY: 10
PARAS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
- NAME: OpenposeExtractor
TYPE: image
DEVICE: "gpu"
MEMORY: 10
MODEL:
NAME: "OpenposeAnnotator"
BODY_MODEL_PATH: "ms://iic/scepter_annotator@annotator/ckpts/body_pose_model.pth"
HAND_MODEL_PATH: "ms://iic/scepter_annotator@annotator/ckpts/hand_pose_model.pth"
PARAS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
- NAME: OutpaintingExtractor
TYPE: image
DEVICE: "cpu"
MEMORY: 10
MODEL:
NAME: "OutpaintingAnnotator"
RETURN_MASK: False
KEEP_PADDING_RATIO: 1
RANDOM_CFG:
DIRECTION_RANGE: [ 'left', 'right', 'up', 'down' ]
RATIO_RANGE: [ 0.1, 0.7 ]
PARAS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
USE_MASK_VISIBLE: True
PREVIEW_BTN_VISIBLE: False
- NAME: OutpaintingResize
TYPE: image
DEVICE: "cpu"
MEMORY: 10
MODEL:
NAME: "OutpaintingResize"
PARAS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
PREVIEW_BTN_VISIBLE: False
- NAME: InfoDrawAnimeAnnotator
TYPE: image
DEVICE: "gpu"
MEMORY: 10
MODEL:
NAME: "InfoDrawAnimeAnnotator"
INPUT_NC: 3
OUTPUT_NC: 1
N_RESIDUAL_BLOCKS: 3
SIGMOID: True
PRETRAINED_MODEL: "ms://iic/scepter_annotator@annotator/ckpts/informative_drawing_anime_style.pth"
PARAS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
- NAME: ESAMExtractor
TYPE: image
DEVICE: "gpu"
MEMORY: 10
MODEL:
NAME: "ESAMAnnotator"
PRETRAINED_MODEL: "ms://iic/scepter_annotator@annotator/ckpts/efficient_sam_vits.pt"
SAVE_MODE: 'P'
GRID_SIZE: 32
USE_DOMINANT_COLOR: True
RETURN_MASK: False
PARAS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
- NAME: InvertExtractor
TYPE: image
DEVICE: "cpu"
MEMORY: 10
MODEL:
NAME: "InvertAnnotator"
PARAS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
- NAME: LamaExtractor
TYPE: image
DEVICE: "cpu"
MEMORY: 10
MODEL:
NAME: "LamaAnnotator"
PRETRAINED_MODEL: "ms:///iic/cv_fft_inpainting_lama/"
PARAS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
+4
View File
@@ -87,3 +87,7 @@ INTERFACE:
NAME_EN: Inference
IFID: inference
CONFIG: scepter/methods/studio/inference/inference.yaml
- NAME: 对话式编辑
NAME_EN: ChatBot
IFID: ChatBot
CONFIG: scepter/methods/studio/chatbot/chatbot.yaml
@@ -0,0 +1,343 @@
ENV:
BACKEND: nccl
META:
VERSION: 'FLUX1.0_DEV'
DESCRIPTION: "flux 1.0 dev"
IS_DEFAULT: False
IS_SHARE: True
INFERENCE_PARAS:
INFERENCE_BATCH_SIZE: 1
INFERENCE_PREFIX: ""
DEFAULT_SAMPLER: "flow_euler"
DEFAULT_SAMPLE_STEPS: 50
INFERENCE_N_PROMPT: ""
RESOLUTION: [1024, 1024]
PARAS:
-
TRAIN_BATCH_SIZE: 1
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: [1024, 1024]
MEMORY: 89000
EPOCHS: 50
SAVE_INTERVAL: 25
EPSEC: 0.818
LEARNING_RATE: 4e-4
IS_DEFAULT: False
TUNER: FULL
-
TRAIN_BATCH_SIZE: 1
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: [1024, 1024]
MEMORY: 89000
EPOCHS: 50
SAVE_INTERVAL: 25
EPSEC: 0.818
LEARNING_RATE: 4e-4
IS_DEFAULT: True
TUNER: LORA
#
TUNERS:
LORA:
-
NAME: SwiftLoRA
R: 4
LORA_ALPHA: 4
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "(model.double_blocks.*(.qkv|.proj|.img_mod.lin|.txt_mod.lin))|(model.single_blocks.*(.linear1|.linear2|.modulation.lin))$"
#
SOLVER:
NAME: LatentDiffusionSolver
# MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
MAX_STEPS: 100000
# USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False
USE_AMP: True
# DTYPE DESCRIPTION: The precision for training. TYPE: str default: 'float32'
DTYPE: bfloat16
# USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
USE_FAIRSCALE: False
# USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
USE_FSDP: True
# LOAD_MODEL_ONLY DESCRIPTION: Only load the model rather than the optimizer and schedule, default is False. TYPE: bool default: False
LOAD_MODEL_ONLY: False
# RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: ''
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_flux_dev_1024_lora
LOG_FILE: std_log.txt
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
EVAL_INTERVAL: 100
# LOG_TRAIN_NUM DESCRIPTION: The number samples used to log in training phase. TYPE: int default: -1
LOG_TRAIN_NUM: 16
# FSDP_REDUCE_DTYPE DESCRIPTION: The dtype of reduce in FSDP. TYPE: str default: 'float16'
FSDP_REDUCE_DTYPE: float32
# FSDP_BUFFER_DTYPE DESCRIPTION: The dtype of buffer in FSDP. TYPE: str default: 'float16'
FSDP_BUFFER_DTYPE: float32
# FSDP_SHARD_MODULES DESCRIPTION: The modules to be sharded in FSDP. TYPE: list default: ['model']
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ] #
SAVE_MODULES: [ 'model'] #
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
FREEZE:
#
TUNER:
#
MODEL:
NAME: LatentDiffusionFlux
PARAMETERIZATION: rf
TIMESTEPS: 1000
MIN_SNR_GAMMA:
ZERO_TERMINAL_SNR: False
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
NOISE_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
NAME: FlowMatchSigmaScheduler
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
WEIGHTING_SCHEME: logit_normal
SHIFT: 3.0
# LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
LOGIT_MEAN: 0.0
# LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
LOGIT_STD: 1.0
# MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
MODE_SCALE: 1.29
SAMPLER_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
NAME: FlowMatchFluxShiftScheduler
# SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
SHIFT: False
# SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
SIGMOID_SCALE: 1
# BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
BASE_SHIFT: 0.5
# MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
MAX_SHIFT: 1.15
#
DIFFUSION_MODEL:
# NAME DESCRIPTION: TYPE: default: 'Flux'
NAME: Flux
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@flux1-dev.safetensors
# IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
IN_CHANNELS: 64
# HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
HIDDEN_SIZE: 3072
# NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
NUM_HEADS: 24
# AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
AXES_DIM: [ 16, 56, 56 ]
# THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
THETA: 10000
# VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
VEC_IN_DIM: 768
# GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
GUIDANCE_EMBED: False
# CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
CONTEXT_IN_DIM: 4096
# MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
MLP_RATIO: 4.0
# QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
QKV_BIAS: True
# DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
DEPTH: 19
# DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
DEPTH_SINGLE_BLOCKS: 38
USE_GRAD_CHECKPOINT: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLFlux
EMBED_DIM: 16
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@ae.safetensors
IGNORE_KEYS: [ ]
BATCH_SIZE: 8
USE_CONV: False
SCALE_FACTOR: 0.3611
SHIFT_FACTOR: 0.1159
#
ENCODER:
NAME: Encoder
USE_CHECKPOINT: True
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 16
DOUBLE_Z: True
DROPOUT: 0.0
RESAMP_WITH_CONV: True
#
DECODER:
NAME: Decoder
USE_CHECKPOINT: True
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 16
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
NAME: T5PlusClipFluxEmbedder
# T5_MODEL DESCRIPTION: TYPE: default: ''
T5_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: T5EncoderModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder_2/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: T5Tokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer_2/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 512
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: last_hidden_state
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: False
CLEAN: whitespace
# CLIP_MODEL DESCRIPTION: TYPE: default: ''
CLIP_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: CLIPTextModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: CLIPTokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 77
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: pooler_output
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: True
CLEAN: whitespace
#
SAMPLE_ARGS:
SAMPLE_STEPS: 50
SAMPLER: flow_eluer
SEED: 2024
IMAGE_SIZE: [ 1024, 1024 ]
SHIFT: True
GUIDE_SCALE: 3.5
#
OPTIMIZER:
NAME: AdamW
LEARNING_RATE: 4e-4
BETAS: [ 0.9, 0.999 ]
EPS: 1e-8
WEIGHT_DECAY: 1e-2
AMSGRAD: False
#
TRAIN_DATA:
NAME: ImageTextPairMSDataset
MODE: train
MS_DATASET_NAME: style_custom_dataset
MS_DATASET_NAMESPACE: damo
MS_DATASET_SUBNAME: 3D
PROMPT_PREFIX: ""
MS_DATASET_SPLIT: train
MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
REPLACE_STYLE: False
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
SAMPLER:
NAME: LoopSampler
TRANSFORMS:
- NAME: LoadImageFromFile
RGB_ORDER: RGB
BACKEND: pillow
- NAME: FlexibleResize
INTERPOLATION: bilinear
SIZE: [ 1024, 1024 ]
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: FlexibleCenterCrop
SIZE: [ 1024, 1024 ]
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: [ 'data_key' ]
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
PROMPT_FILE:
PROMPT_DATA: [ "a cat holds a blackboard that writes \"hello world\"", "a dog running on the lawn" ]
IMAGE_SIZE: [ 1024, 1024 ]
FIELDS: [ "prompt" ]
DELIMITER: '#;#'
PROMPT_PREFIX: ''
PIN_MEMORY: True
BATCH_SIZE: 2
NUM_WORKERS: 4
TRANSFORMS:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
# GRADIENT_CLIP: 1.0
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
- NAME: CheckpointHook
INTERVAL: 10000
PRIORITY: 200
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
DISABLE_SNAPSHOT: True
EVAL_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
SAVE_PROBE_PREFIX: 'image'
@@ -0,0 +1,342 @@
ENV:
BACKEND: nccl
META:
VERSION: 'FLUX1.0_SCHNELL'
DESCRIPTION: "flux 1.0 schnell"
IS_DEFAULT: False
IS_SHARE: True
INFERENCE_PARAS:
INFERENCE_BATCH_SIZE: 1
INFERENCE_PREFIX: ""
DEFAULT_SAMPLER: "flow_euler"
DEFAULT_SAMPLE_STEPS: 4
INFERENCE_N_PROMPT: ""
RESOLUTION: [1024, 1024]
PARAS:
-
TRAIN_BATCH_SIZE: 1
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: [1024, 1024]
MEMORY: 89000
EPOCHS: 50
SAVE_INTERVAL: 25
EPSEC: 0.818
LEARNING_RATE: 4e-4
IS_DEFAULT: False
TUNER: FULL
-
TRAIN_BATCH_SIZE: 1
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: [1024, 1024]
MEMORY: 89000
EPOCHS: 50
SAVE_INTERVAL: 25
EPSEC: 0.818
LEARNING_RATE: 4e-4
IS_DEFAULT: True
TUNER: LORA
#
TUNERS:
LORA:
-
NAME: SwiftLoRA
R: 4
LORA_ALPHA: 4
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "(model.double_blocks.*(.qkv|.proj|.img_mod.lin|.txt_mod.lin))|(model.single_blocks.*(.linear1|.linear2|.modulation.lin))$"
#
SOLVER:
NAME: LatentDiffusionSolver
# MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
MAX_STEPS: 100000
# USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False
USE_AMP: True
# DTYPE DESCRIPTION: The precision for training. TYPE: str default: 'float32'
DTYPE: bfloat16
# USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
USE_FAIRSCALE: False
# USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
USE_FSDP: True
# LOAD_MODEL_ONLY DESCRIPTION: Only load the model rather than the optimizer and schedule, default is False. TYPE: bool default: False
LOAD_MODEL_ONLY: False
# RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: ''
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_flux_schnell_1024_lora
LOG_FILE: std_log.txt
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
EVAL_INTERVAL: 100
# LOG_TRAIN_NUM DESCRIPTION: The number samples used to log in training phase. TYPE: int default: -1
LOG_TRAIN_NUM: 16
# FSDP_REDUCE_DTYPE DESCRIPTION: The dtype of reduce in FSDP. TYPE: str default: 'float16'
FSDP_REDUCE_DTYPE: float32
# FSDP_BUFFER_DTYPE DESCRIPTION: The dtype of buffer in FSDP. TYPE: str default: 'float16'
FSDP_BUFFER_DTYPE: float32
# FSDP_SHARD_MODULES DESCRIPTION: The modules to be sharded in FSDP. TYPE: list default: ['model']
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ] #
SAVE_MODULES: [ 'model'] #
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
FREEZE:
#
TUNER:
#
MODEL:
NAME: LatentDiffusionFlux
PARAMETERIZATION: rf
TIMESTEPS: 1000
MIN_SNR_GAMMA:
ZERO_TERMINAL_SNR: False
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
NOISE_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
NAME: FlowMatchSigmaScheduler
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
WEIGHTING_SCHEME: logit_normal
SHIFT: 3.0
# LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
LOGIT_MEAN: 0.0
# LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
LOGIT_STD: 1.0
# MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
MODE_SCALE: 1.29
SAMPLER_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
NAME: FlowMatchFluxShiftScheduler
# SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
SHIFT: False
# SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
SIGMOID_SCALE: 1
# BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
BASE_SHIFT: 0.5
# MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
MAX_SHIFT: 1.15
#
DIFFUSION_MODEL:
# NAME DESCRIPTION: TYPE: default: 'Flux'
NAME: Flux
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-schnell@flux1-schnell.safetensors
# IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
IN_CHANNELS: 64
# HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
HIDDEN_SIZE: 3072
# NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
NUM_HEADS: 24
# AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
AXES_DIM: [ 16, 56, 56 ]
# THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
THETA: 10000
# VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
VEC_IN_DIM: 768
# GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
GUIDANCE_EMBED: False
# CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
CONTEXT_IN_DIM: 4096
# MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
MLP_RATIO: 4.0
# QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
QKV_BIAS: True
# DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
DEPTH: 19
# DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
DEPTH_SINGLE_BLOCKS: 38
USE_GRAD_CHECKPOINT: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLFlux
EMBED_DIM: 16
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-schnell@ae.safetensors
IGNORE_KEYS: [ ]
BATCH_SIZE: 8
USE_CONV: False
SCALE_FACTOR: 0.3611
SHIFT_FACTOR: 0.1159
#
ENCODER:
NAME: Encoder
USE_CHECKPOINT: True
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 16
DOUBLE_Z: True
DROPOUT: 0.0
RESAMP_WITH_CONV: True
#
DECODER:
NAME: Decoder
USE_CHECKPOINT: True
CH: 128
OUT_CH: 3
NUM_RES_BLOCKS: 2
IN_CHANNELS: 3
ATTN_RESOLUTIONS: [ ]
CH_MULT: [ 1, 2, 4, 4 ]
Z_CHANNELS: 16
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
NAME: T5PlusClipFluxEmbedder
# T5_MODEL DESCRIPTION: TYPE: default: ''
T5_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: T5EncoderModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder_2/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: T5Tokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer_2/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 256
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: last_hidden_state
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: False
CLEAN: whitespace
# CLIP_MODEL DESCRIPTION: TYPE: default: ''
CLIP_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_MODEL_CLS: CLIPTextModel
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder/
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
HF_TOKENIZER_CLS: CLIPTokenizer
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 77
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
OUTPUT_KEY: pooler_output
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
D_TYPE: bfloat16
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
BATCH_INFER: True
CLEAN: whitespace
#
SAMPLE_ARGS:
SAMPLE_STEPS: 4
SAMPLER: flow_eluer
SEED: 2024
IMAGE_SIZE: [ 1024, 1024 ]
GUIDE_SCALE: 3.5
#
OPTIMIZER:
NAME: AdamW
LEARNING_RATE: 4e-4
BETAS: [ 0.9, 0.999 ]
EPS: 1e-8
WEIGHT_DECAY: 1e-2
AMSGRAD: False
#
TRAIN_DATA:
NAME: ImageTextPairMSDataset
MODE: train
MS_DATASET_NAME: style_custom_dataset
MS_DATASET_NAMESPACE: damo
MS_DATASET_SUBNAME: 3D
PROMPT_PREFIX: ""
MS_DATASET_SPLIT: train
MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
REPLACE_STYLE: False
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
SAMPLER:
NAME: LoopSampler
TRANSFORMS:
- NAME: LoadImageFromFile
RGB_ORDER: RGB
BACKEND: pillow
- NAME: FlexibleResize
INTERPOLATION: bilinear
SIZE: [ 1024, 1024 ]
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: FlexibleCenterCrop
SIZE: [ 1024, 1024 ]
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: [ 'data_key' ]
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
PROMPT_FILE:
PROMPT_DATA: [ "a cat holds a blackboard that writes \"hello world\"", "a dog running on the lawn" ]
IMAGE_SIZE: [ 1024, 1024 ]
FIELDS: [ "prompt" ]
DELIMITER: '#;#'
PROMPT_PREFIX: ''
PIN_MEMORY: True
BATCH_SIZE: 2
NUM_WORKERS: 4
TRANSFORMS:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
# GRADIENT_CLIP: 1.0
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
- NAME: CheckpointHook
INTERVAL: 10000
PRIORITY: 200
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
DISABLE_SNAPSHOT: True
EVAL_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
SAVE_PROBE_PREFIX: 'image'
@@ -13,3 +13,5 @@ BASE_MODEL_VERSION:
TUNER_TYPE: [ 'LORA', 'SCE', 'FULL' ]
- BASE_MODEL: 'EDIT'
TUNER_TYPE: [ 'LORA', 'SCE', 'FULL' ]
- BASE_MODEL: 'FLUX1.0_DEV'
TUNER_TYPE: [ 'LORA' ]
+12
View File
@@ -3,9 +3,21 @@
from scepter.modules.annotator.base_annotator import GeneralAnnotator
from scepter.modules.annotator.canny import CannyAnnotator
from scepter.modules.annotator.color import ColorAnnotator
from scepter.modules.annotator.degradation import DegradationAnnotator
from scepter.modules.annotator.doodle import DoodleAnnotator
from scepter.modules.annotator.gray import GrayAnnotator
from scepter.modules.annotator.hed import HedAnnotator
from scepter.modules.annotator.identity import IdentityAnnotator
from scepter.modules.annotator.informative_drawing import (
InfoDrawAnimeAnnotator, InfoDrawContourAnnotator,
InfoDrawOpenSketchAnnotator)
from scepter.modules.annotator.inpainting import InpaintingAnnotator
from scepter.modules.annotator.invert import InvertAnnotator
from scepter.modules.annotator.midas_op import MidasDetector
from scepter.modules.annotator.mlsd_op import MLSDdetector
from scepter.modules.annotator.openpose import OpenposeAnnotator
from scepter.modules.annotator.outpainting import OutpaintingAnnotator, OutpaintingResize
from scepter.modules.annotator.pidinet import PiDiAnnotator
from scepter.modules.annotator.segmentation import ESAMAnnotator
from scepter.modules.annotator.sketch import SketchAnnotator
from scepter.modules.annotator.lama import LamaAnnotator
+135
View File
@@ -0,0 +1,135 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import math
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
def gaussian_noise_op(im, v):
from basicsr.data.degradations import random_add_gaussian_noise
noise_level = v.get('noise_level', [10, 20])
out = random_add_gaussian_noise(
im,
sigma_range=noise_level,
clip=True,
rounds=False,
gray_prob=0.4,
)
out = np.clip(out, 0.0, 1.0)
return out
def resize_op(im, v):
scale = v.get('scale', [0.5, 0.8])
h, w = im.shape[:2]
scale = random.uniform(scale[0], scale[1])
h_, w_ = int(h * scale), int(w * scale)
mode = v.get('mode', 'nearest')
if mode == 'nearest':
interpolation = cv2.INTER_NEAREST
elif mode == 'bilinear':
interpolation = cv2.INTER_LINEAR
elif mode == 'bicubic':
interpolation = cv2.INTER_CUBIC
else:
interpolation = cv2.INTER_NEAREST
im = cv2.resize(im, (w_, h_), interpolation=interpolation)
out = cv2.resize(im, (w, h), interpolation=interpolation)
out = np.clip(out, 0.0, 1.0)
return out
def jpeg_op(im, v):
from basicsr.data.degradations import add_jpg_compression
jpeg_level = v.get('jpeg_level', [50, 75])
v = int(random.uniform(jpeg_level[0], jpeg_level[1]))
out = add_jpg_compression(im, v)
out = np.clip(out, 0.0, 1.0)
return out
def gaussian_blur_op(im, v):
from basicsr.data.degradations import random_mixed_kernels
kernel_range = v.get('kernel_size', [7, 9])
kernel_size = random.choice(kernel_range)
kernel_size = min(int(kernel_size) // 2 * 2 + 1, 21)
blur_sigma = v.get('sigma', [0.9, 1.0])
kernel = random_mixed_kernels(
('iso', 'aniso', 'generalized_iso', 'generalized_aniso', 'plateau_iso',
'plateau_aniso'), (0.45, 0.25, 0.12, 0.03, 0.12, 0.03),
kernel_size,
blur_sigma,
blur_sigma, [-math.pi, math.pi], [0.5, 2.0], [1, 1.5],
noise_range=None)
pad_size = (21 - kernel_size) // 2
kernel = np.pad(kernel, ((pad_size, pad_size), (pad_size, pad_size)))
out = cv2.filter2D(im, -1, kernel)
out = np.clip(out, 0.0, 1.0)
return out
@ANNOTATORS.register_class()
class DegradationAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.params = cfg.get('PARAMS', {
'gaussian_noise': {},
'resize': {},
'jpeg': {},
'gaussian_blur': {},
})
if not isinstance(self.params, dict):
self.params = Config.get_dict(self.params)
self.random_degradation = cfg.get('RANDOM_DEGRADATION', False)
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.'
if np.max(image) > 1.0:
image = (image / 255.).astype(np.float32)
degradation_list = list(self.params.keys())
if self.random_degradation:
random.shuffle(degradation_list)
for degradation_type in degradation_list:
if degradation_type == 'gaussian_noise':
image = gaussian_noise_op(image, self.params[degradation_type])
elif degradation_type == 'resize':
image = resize_op(image, self.params[degradation_type])
elif degradation_type == 'jpeg':
image = jpeg_op(image, self.params[degradation_type])
elif degradation_type == 'gaussian_blur':
image = gaussian_blur_op(image, self.params[degradation_type])
else:
raise NotImplementedError(
f'ERROR: degradation_type: {degradation_type} is invalid.')
image = (image * 255.0).astype(np.uint8)
assert len(image.shape) < 4
return image
@staticmethod
def get_config_template():
return dict_to_yaml('ANNOTATORS',
__class__.__name__,
DegradationAnnotator.para_dict,
set_name=True)
+42
View File
@@ -0,0 +1,42 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from abc import ABCMeta
import torch
from scepter.modules.annotator.base_annotator import BaseAnnotator
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.utils.config import dict_to_yaml
@ANNOTATORS.register_class()
class DoodleAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.processor_type = cfg.get('PROCESSOR_TYPE', 'pidinet_sketch')
processor_cfg = cfg.get('PROCESSOR_CFG', None)
if self.processor_type == 'pidinet_sketch':
self.pidinet_ins = ANNOTATORS.build(processor_cfg[0])
self.sketch_ins = ANNOTATORS.build(processor_cfg[1])
else:
raise 'Unsurpport PROCESSOR for DoodleAnnotator'
@torch.no_grad()
@torch.inference_mode()
@torch.autocast('cuda', enabled=False)
def forward(self, image):
if self.processor_type == 'pidinet_sketch':
pidinet_res = self.pidinet_ins(image)
sketch_res = self.sketch_ins(pidinet_res)
doodle_res = sketch_res
else:
raise 'Unsurpport PROCESSOR for DoodleAnnotator'
return doodle_res
@staticmethod
def get_config_template():
return dict_to_yaml('ANNOTATORS',
__class__.__name__,
DoodleAnnotator.para_dict,
set_name=True)
+38
View File
@@ -0,0 +1,38 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
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 dict_to_yaml
@ANNOTATORS.register_class()
class GrayAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
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.'
gray_map = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)
return gray_map[..., None].repeat(3, axis=2)
@staticmethod
def get_config_template():
return dict_to_yaml('ANNOTATORS',
__class__.__name__,
GrayAnnotator.para_dict,
set_name=True)
@@ -0,0 +1,175 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from abc import ABCMeta
import numpy as np
import torch
import torch.nn as nn
from einops import rearrange
from torchvision.transforms import InterpolationMode
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.distribute import we
from scepter.modules.utils.file_system import FS
norm_layer = nn.InstanceNorm2d
class ResidualBlock(nn.Module):
def __init__(self, in_features):
super(ResidualBlock, self).__init__()
conv_block = [
nn.ReflectionPad2d(1),
nn.Conv2d(in_features, in_features, 3),
norm_layer(in_features),
nn.ReLU(inplace=True),
nn.ReflectionPad2d(1),
nn.Conv2d(in_features, in_features, 3),
norm_layer(in_features)
]
self.conv_block = nn.Sequential(*conv_block)
def forward(self, x):
return x + self.conv_block(x)
class ContourInference(nn.Module):
def __init__(self, input_nc, output_nc, n_residual_blocks=9, sigmoid=True):
super(ContourInference, self).__init__()
# Initial convolution block
model0 = [
nn.ReflectionPad2d(3),
nn.Conv2d(input_nc, 64, 7),
norm_layer(64),
nn.ReLU(inplace=True)
]
self.model0 = nn.Sequential(*model0)
# Downsampling
model1 = []
in_features = 64
out_features = in_features * 2
for _ in range(2):
model1 += [
nn.Conv2d(in_features, out_features, 3, stride=2, padding=1),
norm_layer(out_features),
nn.ReLU(inplace=True)
]
in_features = out_features
out_features = in_features * 2
self.model1 = nn.Sequential(*model1)
model2 = []
# Residual blocks
for _ in range(n_residual_blocks):
model2 += [ResidualBlock(in_features)]
self.model2 = nn.Sequential(*model2)
# Upsampling
model3 = []
out_features = in_features // 2
for _ in range(2):
model3 += [
nn.ConvTranspose2d(in_features,
out_features,
3,
stride=2,
padding=1,
output_padding=1),
norm_layer(out_features),
nn.ReLU(inplace=True)
]
in_features = out_features
out_features = in_features // 2
self.model3 = nn.Sequential(*model3)
# Output layer
model4 = [nn.ReflectionPad2d(3), nn.Conv2d(64, output_nc, 7)]
if sigmoid:
model4 += [nn.Sigmoid()]
self.model4 = nn.Sequential(*model4)
def forward(self, x, cond=None):
out = self.model0(x)
out = self.model1(out)
out = self.model2(out)
out = self.model3(out)
out = self.model4(out)
return out
@ANNOTATORS.register_class()
class InfoDrawContourAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
input_nc = cfg.get('INPUT_NC', 3)
output_nc = cfg.get('OUTPUT_NC', 1)
n_residual_blocks = cfg.get('N_RESIDUAL_BLOCKS', 3)
sigmoid = cfg.get('SIGMOID', True)
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
self.model = ContourInference(input_nc, output_nc, n_residual_blocks,
sigmoid)
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
self.model.load_state_dict(torch.load(local_path))
self.model = self.model.eval().requires_grad_(False).to(we.device_id)
@torch.no_grad()
@torch.inference_mode()
@torch.autocast('cuda', enabled=False)
def forward(self, image):
is_batch = False if len(image.shape) == 3 else True
if isinstance(image, torch.Tensor):
if len(image.shape) == 3:
image = rearrange(image, 'h w c -> 1 c h w')
B, C, H, W = image.shape
elif len(image.shape) == 4:
B, C, H, W = image.shape
else:
raise "Unsurpport input image's shape"
elif isinstance(image, np.ndarray):
image = torch.from_numpy(image.copy()).float()
if len(image.shape) == 3:
image = rearrange(image, 'h w c -> 1 c h w')
B, C, H, W = image.shape
elif len(image.shape) == 4:
B, C, H, W = image.shape
else:
raise "Unsurpport input image's shape"
else:
raise "Unsurpport input image's type"
image = image.float().div(255).to(we.device_id)
contour_map = self.model(image)
contour_map = (contour_map.squeeze(dim=1) * 255.0).clip(
0, 255).cpu().numpy().astype(np.uint8)
contour_map = contour_map[..., None].repeat(3, -1)
if not is_batch:
contour_map = contour_map.squeeze()
return contour_map
@staticmethod
def get_config_template():
return dict_to_yaml('ANNOTATORS',
__class__.__name__,
InfoDrawContourAnnotator.para_dict,
set_name=True)
@ANNOTATORS.register_class()
class InfoDrawAnimeAnnotator(InfoDrawContourAnnotator):
pass
@ANNOTATORS.register_class()
class InfoDrawOpenSketchAnnotator(InfoDrawContourAnnotator):
pass
+275
View File
@@ -0,0 +1,275 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import random
from abc import ABCMeta
from enum import Enum
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
def invert_image(im):
im_arr = np.array(im)
mask_1 = im_arr == 0
mask_2 = im_arr == 255
im_arr[mask_1] = 255
im_arr[mask_2] = 0
new_im = im_arr
return new_im
class DrawMethod(Enum):
LINE = 'line'
CIRCLE = 'circle'
SQUARE = 'square'
def make_random_irregular_mask(shape,
max_angle=4,
max_len=60,
max_width=20,
min_times=0,
max_times=10,
draw_method=DrawMethod.LINE):
draw_method = DrawMethod(draw_method)
height, width = shape
mask = np.zeros((height, width), np.float32)
times = np.random.randint(min_times, max_times + 1)
for i in range(times):
start_x = np.random.randint(width)
start_y = np.random.randint(height)
for j in range(1 + np.random.randint(5)):
angle = 0.01 + np.random.randint(max_angle)
if i % 2 == 0:
angle = 2 * 3.1415926 - angle
length = 10 + np.random.randint(max_len)
brush_w = 5 + np.random.randint(max_width)
end_x = np.clip(
(start_x + length * np.sin(angle)).astype(np.int32), 0, width)
end_y = np.clip(
(start_y + length * np.cos(angle)).astype(np.int32), 0, height)
if draw_method == DrawMethod.LINE:
cv2.line(mask, (start_x, start_y), (end_x, end_y), 1.0,
brush_w)
elif draw_method == DrawMethod.CIRCLE:
cv2.circle(mask, (start_x, start_y),
radius=brush_w,
color=1.,
thickness=-1)
elif draw_method == DrawMethod.SQUARE:
radius = brush_w // 2
mask[start_y - radius:start_y + radius,
start_x - radius:start_x + radius] = 1
start_x, start_y = end_x, end_y
return mask[None, ...]
class RandomIrregularMaskGenerator:
def __init__(self,
max_angle=4,
max_len=60,
max_width=20,
min_times=0,
max_times=10,
ramp_kwargs=None,
draw_method=DrawMethod.LINE):
self.max_angle = max_angle
self.max_len = max_len
self.max_width = max_width
self.min_times = min_times
self.max_times = max_times
self.draw_method = draw_method
self.ramp = None
# self.ramp = LinearRamp(**ramp_kwargs) if ramp_kwargs is not None else None
def __call__(self, img, iter_i=None, raw_image=None):
coef = self.ramp(iter_i) if (self.ramp is not None) and (
iter_i is not None) else 1
cur_max_len = int(max(1, self.max_len * coef))
cur_max_width = int(max(1, self.max_width * coef))
cur_max_times = int(self.min_times + 1 +
(self.max_times - self.min_times) * coef)
return make_random_irregular_mask(img.shape[1:],
max_angle=self.max_angle,
max_len=cur_max_len,
max_width=cur_max_width,
min_times=self.min_times,
max_times=cur_max_times,
draw_method=self.draw_method)
def make_random_rectangle_mask(shape,
margin=10,
bbox_min_size=30,
bbox_max_size=100,
min_times=0,
max_times=3):
height, width = shape
mask = np.zeros((height, width), np.float32)
bbox_max_size = min(bbox_max_size, height - margin * 2, width - margin * 2)
times = np.random.randint(min_times, max_times + 1)
for i in range(times):
box_width = np.random.randint(bbox_min_size, bbox_max_size)
box_height = np.random.randint(bbox_min_size, bbox_max_size)
start_x = np.random.randint(margin, width - margin - box_width + 1)
start_y = np.random.randint(margin, height - margin - box_height + 1)
mask[start_y:start_y + box_height, start_x:start_x + box_width] = 1
return mask[None, ...]
class RandomRectangleMaskGenerator:
def __init__(self,
margin=10,
bbox_min_size=30,
bbox_max_size=100,
min_times=0,
max_times=3,
ramp_kwargs=None):
self.margin = margin
self.bbox_min_size = bbox_min_size
self.bbox_max_size = bbox_max_size
self.min_times = min_times
self.max_times = max_times
self.ramp = None
# self.ramp = LinearRamp(**ramp_kwargs) if ramp_kwargs is not None else None
def __call__(self, img, iter_i=None, raw_image=None):
coef = self.ramp(iter_i) if (self.ramp is not None) and (
iter_i is not None) else 1
cur_bbox_max_size = int(self.bbox_min_size + 1 +
(self.bbox_max_size - self.bbox_min_size) *
coef)
cur_max_times = int(self.min_times +
(self.max_times - self.min_times) * coef)
return make_random_rectangle_mask(img.shape[1:],
margin=self.margin,
bbox_min_size=self.bbox_min_size,
bbox_max_size=cur_bbox_max_size,
min_times=self.min_times,
max_times=cur_max_times)
class MixedMaskGenerator:
def __init__(self,
irregular_proba=1 / 3,
irregular_kwargs=None,
box_proba=1 / 3,
box_kwargs=None,
invert_proba=0):
self.probas = []
self.gens = []
if irregular_proba > 0:
self.probas.append(irregular_proba)
if irregular_kwargs is None:
irregular_kwargs = {}
else:
irregular_kwargs = dict(irregular_kwargs)
irregular_kwargs['draw_method'] = DrawMethod.LINE
self.gens.append(RandomIrregularMaskGenerator(**irregular_kwargs))
if box_proba > 0:
self.probas.append(box_proba)
if box_kwargs is None:
box_kwargs = {}
self.gens.append(RandomRectangleMaskGenerator(**box_kwargs))
self.probas = np.array(self.probas, dtype='float32')
self.probas /= self.probas.sum()
self.invert_proba = invert_proba
def __call__(self, img, iter_i=None, raw_image=None):
kind = np.random.choice(len(self.probas), p=self.probas)
gen = self.gens[kind]
result = gen(img, iter_i=iter_i, raw_image=raw_image)
if self.invert_proba > 0 and random.random() < self.invert_proba:
result = 1 - result
return result
@ANNOTATORS.register_class()
class InpaintingAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.mask_cfg = cfg.get(
'MASK_CFG', {
'irregular_proba': 0.5,
'irregular_kwargs': {
'min_times': 4,
'max_times': 10,
'max_width': 100,
'max_angle': 4,
'max_len': 200
},
'box_proba': 0.5,
'box_kwargs': {
'margin': 0,
'bbox_min_size': 30,
'bbox_max_size': 150,
'max_times': 5,
'min_times': 1
}
})
self.mask_cfg = Config.get_dict(self.mask_cfg) if not isinstance(
self.mask_cfg, dict) else self.mask_cfg
self.mask_generator = MixedMaskGenerator(**self.mask_cfg)
self.return_mask = cfg.get('RETURN_MASK', False)
self.return_invert = cfg.get('RETURN_INVERT', True)
self.mask_color = cfg.get('MASK_COLOR', 0)
def forward(self,
image,
mask=None,
return_mask=None,
return_invert=None,
mask_color=None):
return_mask = return_mask if return_mask is not None else self.return_mask
return_invert = return_invert if return_invert is not None else self.return_invert
mask_color = mask_color if mask_color is not None else self.mask_color
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.'
if mask is not None:
if mask_color:
image[np.array(mask) == 255] = mask_color
else:
image[np.array(mask) == 255] = 0
else:
img = np.transpose(image, (2, 0, 1))
mask = self.mask_generator(img)
mask = (np.transpose(mask,
(1, 2, 0)).squeeze(-1) * 255).astype(np.uint8)
if return_invert:
mask = invert_image(mask)
colored_mask = np.zeros_like(image)
if mask_color: colored_mask[:] = mask_color
image = np.where(mask[:, :, np.newaxis] == 255, colored_mask,
image)
if return_mask:
ret_data = {'image': np.array(image), 'mask': np.array(mask)}
else:
ret_data = np.array(image)
return ret_data
@staticmethod
def get_config_template():
return dict_to_yaml('ANNOTATORS',
__class__.__name__,
InpaintingAnnotator.para_dict,
set_name=True)
+106
View File
@@ -0,0 +1,106 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from abc import ABCMeta
import numpy as np
import cv2
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 dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
def dilate_mask(mask, dilate_factor=15):
mask = mask.astype(np.uint8)
mask = cv2.dilate(mask,
np.ones((dilate_factor, dilate_factor), np.uint8),
iterations=1)
return mask
@ANNOTATORS.register_class()
class LamaAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
from modelscope.pipelines.builder import PIPELINES
from modelscope.pipelines.cv import ImageInpaintingPipeline
from modelscope.pipelines import pipeline
from modelscope.utils.constant import Tasks
from modelscope.metainfo import Pipelines
from modelscope.models.cv.image_inpainting.refinement import refine_predict
from torch.utils.data._utils.collate import default_collate
@PIPELINES.register_module(Tasks.image_inpainting,
module_name=Pipelines.image_inpainting +
'-v2')
class ImageInpaintingPipelineV2(ImageInpaintingPipeline):
def perform_inference(self, data):
px_budget = 9000000
batch = default_collate([data])
if self.refine:
assert 'unpad_to_size' in batch, 'Unpadded size is required for the refinement'
assert 'cuda' in str(
self.device), 'GPU is required for refinement'
gpu_ids = str(self.device).split(':')[-1]
cur_res = refine_predict(batch,
self.infer_model,
gpu_ids=gpu_ids,
modulo=self.pad_out_to_modulo,
n_iters=15,
lr=0.002,
min_side=512,
max_scales=3,
px_budget=px_budget)
cur_res = cur_res[0].permute(1, 2,
0).detach().cpu().numpy()
else:
with torch.no_grad():
batch = self.move_to_device(batch, self.device)
batch['mask'] = (batch['mask'] > 0) * 1
batch = self.infer_model(batch)
cur_res = batch['inpainted'][0].permute(
1, 2, 0).detach().cpu().numpy()
unpad_to_size = batch.get('unpad_to_size', None)
if unpad_to_size is not None:
orig_height, orig_width = unpad_to_size
cur_res = cur_res[:orig_height, :orig_width]
cur_res = np.clip(cur_res * 255, 0, 255).astype('uint8')
cur_res = cv2.cvtColor(cur_res, cv2.COLOR_RGB2BGR)
return cur_res
lama_model_dir = FS.get_dir_to_local_dir(cfg.PRETRAINED_MODEL)
self.lama_model = pipeline(Tasks.image_inpainting,
model=lama_model_dir,
pipeline_name=Pipelines.image_inpainting +
'-v2',
refine=True,
device='cuda:{}'.format(we.device_id))
def forward(self, image, mask):
mask = dilate_mask(mask, dilate_factor=19)
input_mask = Image.fromarray(mask)
mask_expanded = np.tile(np.expand_dims(mask, axis=-1), (1, 1, 3))
input_image_np = np.array(image)
input_image_np[mask_expanded == 255] = 0
input_image = Image.fromarray(input_image_np)
input = {
'img': input_image,
'mask': input_mask,
}
result = self.lama_model(input)
output_img = result['output_img']
return output_img[..., ::-1]
@staticmethod
def get_config_template():
return dict_to_yaml('ANNOTATORS',
__class__.__name__,
LamaAnnotator.para_dict,
set_name=True)
+4 -3
View File
@@ -9,9 +9,10 @@ import os
from abc import ABCMeta
from collections import OrderedDict
import numpy as np
import cv2
import matplotlib
import numpy as np
import torch
import torch.nn as nn
from PIL import Image
@@ -478,7 +479,7 @@ class Hand(object):
map_ori = heatmap_avg[:, :, part]
one_heatmap = gaussian_filter(map_ori, sigma=3)
binary = np.ascontiguousarray(one_heatmap > thre, dtype=np.uint8)
# 全部小于阈值
if np.sum(binary) == 0:
all_peaks.append([0, 0])
continue
@@ -788,7 +789,7 @@ class OpenposeAnnotator(BaseAnnotator, metaclass=ABCMeta):
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
image = image[:, :, ::-1]
candidate, subset = self.body_estimation(image)
canvas = np.zeros_like(image)
canvas = np.zeros_like(image, order='C') # to check
canvas = draw_bodypose(canvas, candidate, subset)
if self.use_hand:
hands_list = handDetect(candidate, subset, image)
+199
View File
@@ -0,0 +1,199 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import math
import random
from abc import ABCMeta
import numpy as np
import torch
from PIL import Image, ImageDraw
from scepter.modules.annotator.base_annotator import BaseAnnotator
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.utils.config import dict_to_yaml
@ANNOTATORS.register_class()
class OutpaintingAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.mask_blur = cfg.get('MASK_BLUR', 0)
self.random_cfg = cfg.get('RANDOM_CFG', None)
self.return_mask = cfg.get('RETURN_MASK', False)
self.return_source = cfg.get('RETURN_SOURCE', True)
self.keep_padding_ratio = cfg.get('KEEP_PADDING_RATIO', 64)
self.mask_color = cfg.get('MASK_COLOR', 0)
def get_box(self, mask):
locs = np.where(mask == 255)
if len(locs) < 1 or locs[0].shape[0] < 1 or locs[1].shape[0] < 1:
return None
left, right = np.min(locs[1]), np.max(locs[1])
top, bottom = np.min(locs[0]), np.max(locs[0])
return [left, top, right, bottom]
def forward(self,
image,
ratio=0.3,
mask=None,
direction=['left', 'right', 'up', 'down'],
return_mask=None,
return_source=None,
mask_color=None):
return_mask = return_mask if return_mask is not None else self.return_mask
return_source = return_source if return_source is not None else self.return_source
mask_color = mask_color if mask_color is not None else self.mask_color
if isinstance(image, Image.Image):
image = image
elif isinstance(image, torch.Tensor):
image = Image.fromarray(image.detach().cpu().numpy())
elif isinstance(image, np.ndarray):
image = Image.fromarray(image.copy())
else:
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
if self.random_cfg:
direction_range = self.random_cfg.get(
'DIRECTION_RANGE', ['left', 'right', 'up', 'down'])
ratio_range = self.random_cfg.get('RATIO_RANGE', [0.0, 1.0])
direction = random.sample(
direction_range,
random.choice(list(range(1,
len(direction_range) + 1))))
ratio = random.uniform(ratio_range[0], ratio_range[1])
if mask is None:
init_image = image
src_width, src_height = init_image.width, init_image.height
left = int(ratio * src_width) if 'left' in direction else 0
right = int(ratio * src_width) if 'right' in direction else 0
up = int(ratio * src_height) if 'up' in direction else 0
down = int(ratio * src_height) if 'down' in direction else 0
# print(direction, ratio, left, right, up, down)
tar_width = math.ceil(
(src_width + left + right) /
self.keep_padding_ratio) * self.keep_padding_ratio
tar_height = math.ceil(
(src_height + up + down) /
self.keep_padding_ratio) * self.keep_padding_ratio
if left > 0:
left = left * (tar_width - src_width) // (left + right)
if right > 0:
right = tar_width - src_width - left
if up > 0:
up = up * (tar_height - src_height) // (up + down)
if down > 0:
down = tar_height - src_height - up
if mask_color is not None:
img = Image.new('RGB', (tar_width, tar_height),
color=mask_color)
else:
img = Image.new('RGB', (tar_width, tar_height))
img.paste(init_image, (left, up))
mask = Image.new('L', (img.width, img.height), 'white')
draw = ImageDraw.Draw(mask)
draw.rectangle(
(left + (self.mask_blur * 2 if left > 0 else 0), up +
(self.mask_blur * 2 if up > 0 else 0), mask.width - right -
(self.mask_blur * 2 if right > 0 else 0), mask.height - down -
(self.mask_blur * 2 if down > 0 else 0)),
fill='black')
else:
bbox = self.get_box(np.array(mask))
if bbox is None:
img = image
mask = mask
init_image = image
else:
mask = Image.new('L', (image.width, image.height), 'white')
mask_zero = Image.new('L',
(bbox[2] - bbox[0], bbox[3] - bbox[1]),
'black')
mask.paste(mask_zero, (bbox[0], bbox[1]))
crop_image = image.crop(bbox)
init_image = Image.new('RGB', (image.width, image.height),
'black')
init_image.paste(crop_image, (bbox[0], bbox[1]))
img = image
if return_mask:
if return_source:
ret_data = {
'src_image': np.array(init_image),
'image': np.array(img),
'mask': np.array(mask)
}
else:
ret_data = {'image': np.array(img), 'mask': np.array(mask)}
else:
if return_source:
ret_data = {
'src_image': np.array(init_image),
'image': np.array(img)
}
else:
ret_data = np.array(img)
return ret_data
@staticmethod
def get_config_template():
return dict_to_yaml('ANNOTATORS',
__class__.__name__,
OutpaintingAnnotator.para_dict,
set_name=True)
@ANNOTATORS.register_class()
class OutpaintingResize(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
def get_box(self, mask):
locs = np.where(mask == 0)
if len(locs) < 1 or locs[0].shape[0] < 1 or locs[1].shape[0] < 1:
return None
left, right = np.min(locs[1]), np.max(locs[1])
top, bottom = np.min(locs[0]), np.max(locs[0])
return [left, top, right, bottom]
def forward(self, image, target_image, mask=None):
if isinstance(image, Image.Image):
image = image
elif isinstance(image, torch.Tensor):
image = Image.fromarray(image.detach().cpu().numpy())
elif isinstance(image, np.ndarray):
image = Image.fromarray(image.copy())
else:
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
if isinstance(target_image, Image.Image):
target_image = target_image
elif isinstance(target_image, torch.Tensor):
target_image = Image.fromarray(target_image.detach().cpu().numpy())
elif isinstance(target_image, np.ndarray):
target_image = Image.fromarray(target_image.copy())
else:
raise f'Unsurpport datatype{type(target_image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
bbox = self.get_box(np.array(mask))
if bbox is None:
init_image = image
else:
paste_img = image.resize((bbox[2] - bbox[0], bbox[3] - bbox[1]))
init_image = Image.new('RGB',
(target_image.width, target_image.height),
'black')
init_image.paste(paste_img, (bbox[0], bbox[1]))
ret_data = {'src_image': np.array(init_image)}
return ret_data
@staticmethod
def get_config_template():
return dict_to_yaml('ANNOTATORS',
__class__.__name__,
OutpaintingResize.para_dict,
set_name=True)
+934
View File
@@ -0,0 +1,934 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import math
from abc import ABCMeta
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
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.distribute import we
from scepter.modules.utils.file_system import FS
CONFIGS = {
'baseline': {
'layer0': 'cv',
'layer1': 'cv',
'layer2': 'cv',
'layer3': 'cv',
'layer4': 'cv',
'layer5': 'cv',
'layer6': 'cv',
'layer7': 'cv',
'layer8': 'cv',
'layer9': 'cv',
'layer10': 'cv',
'layer11': 'cv',
'layer12': 'cv',
'layer13': 'cv',
'layer14': 'cv',
'layer15': 'cv',
},
'c-v15': {
'layer0': 'cd',
'layer1': 'cv',
'layer2': 'cv',
'layer3': 'cv',
'layer4': 'cv',
'layer5': 'cv',
'layer6': 'cv',
'layer7': 'cv',
'layer8': 'cv',
'layer9': 'cv',
'layer10': 'cv',
'layer11': 'cv',
'layer12': 'cv',
'layer13': 'cv',
'layer14': 'cv',
'layer15': 'cv',
},
'a-v15': {
'layer0': 'ad',
'layer1': 'cv',
'layer2': 'cv',
'layer3': 'cv',
'layer4': 'cv',
'layer5': 'cv',
'layer6': 'cv',
'layer7': 'cv',
'layer8': 'cv',
'layer9': 'cv',
'layer10': 'cv',
'layer11': 'cv',
'layer12': 'cv',
'layer13': 'cv',
'layer14': 'cv',
'layer15': 'cv',
},
'r-v15': {
'layer0': 'rd',
'layer1': 'cv',
'layer2': 'cv',
'layer3': 'cv',
'layer4': 'cv',
'layer5': 'cv',
'layer6': 'cv',
'layer7': 'cv',
'layer8': 'cv',
'layer9': 'cv',
'layer10': 'cv',
'layer11': 'cv',
'layer12': 'cv',
'layer13': 'cv',
'layer14': 'cv',
'layer15': 'cv',
},
'cvvv4': {
'layer0': 'cd',
'layer1': 'cv',
'layer2': 'cv',
'layer3': 'cv',
'layer4': 'cd',
'layer5': 'cv',
'layer6': 'cv',
'layer7': 'cv',
'layer8': 'cd',
'layer9': 'cv',
'layer10': 'cv',
'layer11': 'cv',
'layer12': 'cd',
'layer13': 'cv',
'layer14': 'cv',
'layer15': 'cv',
},
'avvv4': {
'layer0': 'ad',
'layer1': 'cv',
'layer2': 'cv',
'layer3': 'cv',
'layer4': 'ad',
'layer5': 'cv',
'layer6': 'cv',
'layer7': 'cv',
'layer8': 'ad',
'layer9': 'cv',
'layer10': 'cv',
'layer11': 'cv',
'layer12': 'ad',
'layer13': 'cv',
'layer14': 'cv',
'layer15': 'cv',
},
'rvvv4': {
'layer0': 'rd',
'layer1': 'cv',
'layer2': 'cv',
'layer3': 'cv',
'layer4': 'rd',
'layer5': 'cv',
'layer6': 'cv',
'layer7': 'cv',
'layer8': 'rd',
'layer9': 'cv',
'layer10': 'cv',
'layer11': 'cv',
'layer12': 'rd',
'layer13': 'cv',
'layer14': 'cv',
'layer15': 'cv',
},
'cccv4': {
'layer0': 'cd',
'layer1': 'cd',
'layer2': 'cd',
'layer3': 'cv',
'layer4': 'cd',
'layer5': 'cd',
'layer6': 'cd',
'layer7': 'cv',
'layer8': 'cd',
'layer9': 'cd',
'layer10': 'cd',
'layer11': 'cv',
'layer12': 'cd',
'layer13': 'cd',
'layer14': 'cd',
'layer15': 'cv',
},
'aaav4': {
'layer0': 'ad',
'layer1': 'ad',
'layer2': 'ad',
'layer3': 'cv',
'layer4': 'ad',
'layer5': 'ad',
'layer6': 'ad',
'layer7': 'cv',
'layer8': 'ad',
'layer9': 'ad',
'layer10': 'ad',
'layer11': 'cv',
'layer12': 'ad',
'layer13': 'ad',
'layer14': 'ad',
'layer15': 'cv',
},
'rrrv4': {
'layer0': 'rd',
'layer1': 'rd',
'layer2': 'rd',
'layer3': 'cv',
'layer4': 'rd',
'layer5': 'rd',
'layer6': 'rd',
'layer7': 'cv',
'layer8': 'rd',
'layer9': 'rd',
'layer10': 'rd',
'layer11': 'cv',
'layer12': 'rd',
'layer13': 'rd',
'layer14': 'rd',
'layer15': 'cv',
},
'c16': {
'layer0': 'cd',
'layer1': 'cd',
'layer2': 'cd',
'layer3': 'cd',
'layer4': 'cd',
'layer5': 'cd',
'layer6': 'cd',
'layer7': 'cd',
'layer8': 'cd',
'layer9': 'cd',
'layer10': 'cd',
'layer11': 'cd',
'layer12': 'cd',
'layer13': 'cd',
'layer14': 'cd',
'layer15': 'cd',
},
'a16': {
'layer0': 'ad',
'layer1': 'ad',
'layer2': 'ad',
'layer3': 'ad',
'layer4': 'ad',
'layer5': 'ad',
'layer6': 'ad',
'layer7': 'ad',
'layer8': 'ad',
'layer9': 'ad',
'layer10': 'ad',
'layer11': 'ad',
'layer12': 'ad',
'layer13': 'ad',
'layer14': 'ad',
'layer15': 'ad',
},
'r16': {
'layer0': 'rd',
'layer1': 'rd',
'layer2': 'rd',
'layer3': 'rd',
'layer4': 'rd',
'layer5': 'rd',
'layer6': 'rd',
'layer7': 'rd',
'layer8': 'rd',
'layer9': 'rd',
'layer10': 'rd',
'layer11': 'rd',
'layer12': 'rd',
'layer13': 'rd',
'layer14': 'rd',
'layer15': 'rd',
},
'carv4': {
'layer0': 'cd',
'layer1': 'ad',
'layer2': 'rd',
'layer3': 'cv',
'layer4': 'cd',
'layer5': 'ad',
'layer6': 'rd',
'layer7': 'cv',
'layer8': 'cd',
'layer9': 'ad',
'layer10': 'rd',
'layer11': 'cv',
'layer12': 'cd',
'layer13': 'ad',
'layer14': 'rd',
'layer15': 'cv'
}
}
def create_conv_func(op_type):
assert op_type in ['cv', 'cd', 'ad',
'rd'], 'unknown op type: %s' % str(op_type)
if op_type == 'cv':
return F.conv2d
if op_type == 'cd':
def func(x,
weights,
bias=None,
stride=1,
padding=0,
dilation=1,
groups=1):
assert dilation in [1,
2], 'dilation for cd_conv should be in 1 or 2'
assert weights.size(2) == 3 and weights.size(3) == 3, \
'kernel size for cd_conv should be 3x3'
assert padding == dilation, 'padding for cd_conv set wrong'
weights_c = weights.sum(dim=[2, 3], keepdim=True)
yc = F.conv2d(x,
weights_c,
stride=stride,
padding=0,
groups=groups)
y = F.conv2d(x,
weights,
bias,
stride=stride,
padding=padding,
dilation=dilation,
groups=groups)
return y - yc
return func
elif op_type == 'ad':
def func(x,
weights,
bias=None,
stride=1,
padding=0,
dilation=1,
groups=1):
assert dilation in [1,
2], 'dilation for ad_conv should be in 1 or 2'
assert weights.size(2) == 3 and weights.size(3) == 3, \
'kernel size for ad_conv should be 3x3'
assert padding == dilation, 'padding for ad_conv set wrong'
shape = weights.shape
weights = weights.view(shape[0], shape[1], -1)
# clock-wise
weights_conv = (
weights -
weights[:, :, [3, 0, 1, 6, 4, 2, 7, 8, 5]]).view(shape)
y = F.conv2d(x,
weights_conv,
bias,
stride=stride,
padding=padding,
dilation=dilation,
groups=groups)
return y
return func
elif op_type == 'rd':
def func(x,
weights,
bias=None,
stride=1,
padding=0,
dilation=1,
groups=1):
assert dilation in [1,
2], 'dilation for rd_conv should be in 1 or 2'
assert weights.size(2) == 3 and weights.size(3) == 3, \
'kernel size for rd_conv should be 3x3'
padding = 2 * dilation
shape = weights.shape
if weights.is_cuda:
buffer = torch.cuda.FloatTensor(shape[0], shape[1],
5 * 5).fill_(0)
else:
buffer = torch.zeros(shape[0], shape[1], 5 * 5)
weights = weights.view(shape[0], shape[1], -1)
buffer[:, :, [0, 2, 4, 10, 14, 20, 22, 24]] = weights[:, :, 1:]
buffer[:, :, [6, 7, 8, 11, 13, 16, 17, 18]] = -weights[:, :, 1:]
buffer[:, :, 12] = 0
buffer = buffer.view(shape[0], shape[1], 5, 5)
y = F.conv2d(x,
buffer,
bias,
stride=stride,
padding=padding,
dilation=dilation,
groups=groups)
return y
return func
else:
print('impossible to be here unless you force that', flush=True)
return None
def config_model(model):
model_options = list(CONFIGS.keys())
assert model in model_options, \
'unrecognized model, please choose from %s' % str(model_options)
pdcs = []
for i in range(16):
layer_name = 'layer%d' % i
op = CONFIGS[model][layer_name]
pdcs.append(create_conv_func(op))
return pdcs
def config_model_converted(model):
model_options = list(CONFIGS.keys())
assert model in model_options, \
'unrecognized model, please choose from %s' % str(model_options)
pdcs = []
for i in range(16):
layer_name = 'layer%d' % i
op = CONFIGS[model][layer_name]
pdcs.append(op)
return pdcs
def convert_pdc(op, weight):
if op == 'cv':
return weight
elif op == 'cd':
shape = weight.shape
weight_c = weight.sum(dim=[2, 3])
weight = weight.view(shape[0], shape[1], -1)
weight[:, :, 4] = weight[:, :, 4] - weight_c
weight = weight.view(shape)
return weight
elif op == 'ad':
shape = weight.shape
weight = weight.view(shape[0], shape[1], -1)
weight_conv = (weight -
weight[:, :, [3, 0, 1, 6, 4, 2, 7, 8, 5]]).view(shape)
return weight_conv
elif op == 'rd':
shape = weight.shape
buffer = torch.zeros(shape[0], shape[1], 5 * 5, device=weight.device)
weight = weight.view(shape[0], shape[1], -1)
buffer[:, :, [0, 2, 4, 10, 14, 20, 22, 24]] = weight[:, :, 1:]
buffer[:, :, [6, 7, 8, 11, 13, 16, 17, 18]] = -weight[:, :, 1:]
buffer = buffer.view(shape[0], shape[1], 5, 5)
return buffer
raise ValueError('wrong op {}'.format(str(op)))
def convert_pidinet(state_dict, config):
pdcs = config_model_converted(config)
new_dict = {}
for pname, p in state_dict.items():
if 'init_block.weight' in pname:
new_dict[pname] = convert_pdc(pdcs[0], p)
elif 'block1_1.conv1.weight' in pname:
new_dict[pname] = convert_pdc(pdcs[1], p)
elif 'block1_2.conv1.weight' in pname:
new_dict[pname] = convert_pdc(pdcs[2], p)
elif 'block1_3.conv1.weight' in pname:
new_dict[pname] = convert_pdc(pdcs[3], p)
elif 'block2_1.conv1.weight' in pname:
new_dict[pname] = convert_pdc(pdcs[4], p)
elif 'block2_2.conv1.weight' in pname:
new_dict[pname] = convert_pdc(pdcs[5], p)
elif 'block2_3.conv1.weight' in pname:
new_dict[pname] = convert_pdc(pdcs[6], p)
elif 'block2_4.conv1.weight' in pname:
new_dict[pname] = convert_pdc(pdcs[7], p)
elif 'block3_1.conv1.weight' in pname:
new_dict[pname] = convert_pdc(pdcs[8], p)
elif 'block3_2.conv1.weight' in pname:
new_dict[pname] = convert_pdc(pdcs[9], p)
elif 'block3_3.conv1.weight' in pname:
new_dict[pname] = convert_pdc(pdcs[10], p)
elif 'block3_4.conv1.weight' in pname:
new_dict[pname] = convert_pdc(pdcs[11], p)
elif 'block4_1.conv1.weight' in pname:
new_dict[pname] = convert_pdc(pdcs[12], p)
elif 'block4_2.conv1.weight' in pname:
new_dict[pname] = convert_pdc(pdcs[13], p)
elif 'block4_3.conv1.weight' in pname:
new_dict[pname] = convert_pdc(pdcs[14], p)
elif 'block4_4.conv1.weight' in pname:
new_dict[pname] = convert_pdc(pdcs[15], p)
else:
new_dict[pname] = p
return new_dict
class Conv2d(nn.Module):
def __init__(self,
pdc,
in_channels,
out_channels,
kernel_size,
stride=1,
padding=0,
dilation=1,
groups=1,
bias=False):
super().__init__()
if in_channels % groups != 0:
raise ValueError('in_channels must be divisible by groups')
if out_channels % groups != 0:
raise ValueError('out_channels must be divisible by groups')
self.in_channels = in_channels
self.out_channels = out_channels
self.kernel_size = kernel_size
self.stride = stride
self.padding = padding
self.dilation = dilation
self.groups = groups
self.weight = nn.Parameter(
torch.Tensor(out_channels, in_channels // groups, kernel_size,
kernel_size))
if bias:
self.bias = nn.Parameter(torch.Tensor(out_channels))
else:
self.register_parameter('bias', None)
self.reset_parameters()
self.pdc = pdc
def reset_parameters(self):
nn.init.kaiming_uniform_(self.weight, a=math.sqrt(5))
if self.bias is not None:
fan_in, _ = nn.init._calculate_fan_in_and_fan_out(self.weight)
bound = 1 / math.sqrt(fan_in)
nn.init.uniform_(self.bias, -bound, bound)
def forward(self, input):
return self.pdc(input, self.weight, self.bias, self.stride,
self.padding, self.dilation, self.groups)
class CSAM(nn.Module):
"""
Compact Spatial Attention Module
"""
def __init__(self, channels):
super().__init__()
mid_channels = 4
self.relu1 = nn.ReLU()
self.conv1 = nn.Conv2d(channels,
mid_channels,
kernel_size=1,
padding=0)
self.conv2 = nn.Conv2d(mid_channels,
1,
kernel_size=3,
padding=1,
bias=False)
self.sigmoid = nn.Sigmoid()
nn.init.constant_(self.conv1.bias, 0)
def forward(self, x):
y = self.relu1(x)
y = self.conv1(y)
y = self.conv2(y)
y = self.sigmoid(y)
return x * y
class CDCM(nn.Module):
"""
Compact Dilation Convolution based Module
"""
def __init__(self, in_channels, out_channels):
super().__init__()
self.relu1 = nn.ReLU()
self.conv1 = nn.Conv2d(in_channels,
out_channels,
kernel_size=1,
padding=0)
self.conv2_1 = nn.Conv2d(out_channels,
out_channels,
kernel_size=3,
dilation=5,
padding=5,
bias=False)
self.conv2_2 = nn.Conv2d(out_channels,
out_channels,
kernel_size=3,
dilation=7,
padding=7,
bias=False)
self.conv2_3 = nn.Conv2d(out_channels,
out_channels,
kernel_size=3,
dilation=9,
padding=9,
bias=False)
self.conv2_4 = nn.Conv2d(out_channels,
out_channels,
kernel_size=3,
dilation=11,
padding=11,
bias=False)
nn.init.constant_(self.conv1.bias, 0)
def forward(self, x):
x = self.relu1(x)
x = self.conv1(x)
x1 = self.conv2_1(x)
x2 = self.conv2_2(x)
x3 = self.conv2_3(x)
x4 = self.conv2_4(x)
return x1 + x2 + x3 + x4
class MapReduce(nn.Module):
"""
Reduce feature maps into a single edge map
"""
def __init__(self, channels):
super().__init__()
self.conv = nn.Conv2d(channels, 1, kernel_size=1, padding=0)
nn.init.constant_(self.conv.bias, 0)
def forward(self, x):
return self.conv(x)
class PDCBlock(nn.Module):
def __init__(self, pdc, inplane, ouplane, stride=1):
super().__init__()
self.stride = stride
self.stride = stride
if self.stride > 1:
self.pool = nn.MaxPool2d(kernel_size=2, stride=2)
self.shortcut = nn.Conv2d(inplane,
ouplane,
kernel_size=1,
padding=0)
self.conv1 = Conv2d(pdc,
inplane,
inplane,
kernel_size=3,
padding=1,
groups=inplane,
bias=False)
self.relu2 = nn.ReLU()
self.conv2 = nn.Conv2d(inplane,
ouplane,
kernel_size=1,
padding=0,
bias=False)
def forward(self, x):
if self.stride > 1:
x = self.pool(x)
y = self.conv1(x)
y = self.relu2(y)
y = self.conv2(y)
if self.stride > 1:
x = self.shortcut(x)
y = y + x
return y
class PDCBlock_converted(nn.Module):
"""
CPDC, APDC can be converted to vanilla 3x3 convolution
RPDC can be converted to vanilla 5x5 convolution
"""
def __init__(self, pdc, inplane, ouplane, stride=1):
super().__init__()
self.stride = stride
if self.stride > 1:
self.pool = nn.MaxPool2d(kernel_size=2, stride=2)
self.shortcut = nn.Conv2d(inplane,
ouplane,
kernel_size=1,
padding=0)
if pdc == 'rd':
self.conv1 = nn.Conv2d(inplane,
inplane,
kernel_size=5,
padding=2,
groups=inplane,
bias=False)
else:
self.conv1 = nn.Conv2d(inplane,
inplane,
kernel_size=3,
padding=1,
groups=inplane,
bias=False)
self.relu2 = nn.ReLU()
self.conv2 = nn.Conv2d(inplane,
ouplane,
kernel_size=1,
padding=0,
bias=False)
def forward(self, x):
if self.stride > 1:
x = self.pool(x)
y = self.conv1(x)
y = self.relu2(y)
y = self.conv2(y)
if self.stride > 1:
x = self.shortcut(x)
y = y + x
return y
class PiDiNet(nn.Module):
def __init__(self,
inplane,
pdcs,
dil=None,
sa=False,
convert=False,
mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]):
super().__init__()
self.sa = sa
if dil is not None:
assert isinstance(dil, int), 'dil should be an int'
self.dil = dil
self.mean = mean
self.std = std
self.fuseplanes = []
self.inplane = inplane
if convert:
if pdcs[0] == 'rd':
init_kernel_size = 5
init_padding = 2
else:
init_kernel_size = 3
init_padding = 1
self.init_block = nn.Conv2d(3,
self.inplane,
kernel_size=init_kernel_size,
padding=init_padding,
bias=False)
block_class = PDCBlock_converted
else:
self.init_block = Conv2d(pdcs[0],
3,
self.inplane,
kernel_size=3,
padding=1)
block_class = PDCBlock
self.block1_1 = block_class(pdcs[1], self.inplane, self.inplane)
self.block1_2 = block_class(pdcs[2], self.inplane, self.inplane)
self.block1_3 = block_class(pdcs[3], self.inplane, self.inplane)
self.fuseplanes.append(self.inplane) # C
inplane = self.inplane
self.inplane = self.inplane * 2
self.block2_1 = block_class(pdcs[4], inplane, self.inplane, stride=2)
self.block2_2 = block_class(pdcs[5], self.inplane, self.inplane)
self.block2_3 = block_class(pdcs[6], self.inplane, self.inplane)
self.block2_4 = block_class(pdcs[7], self.inplane, self.inplane)
self.fuseplanes.append(self.inplane) # 2C
inplane = self.inplane
self.inplane = self.inplane * 2
self.block3_1 = block_class(pdcs[8], inplane, self.inplane, stride=2)
self.block3_2 = block_class(pdcs[9], self.inplane, self.inplane)
self.block3_3 = block_class(pdcs[10], self.inplane, self.inplane)
self.block3_4 = block_class(pdcs[11], self.inplane, self.inplane)
self.fuseplanes.append(self.inplane) # 4C
self.block4_1 = block_class(pdcs[12],
self.inplane,
self.inplane,
stride=2)
self.block4_2 = block_class(pdcs[13], self.inplane, self.inplane)
self.block4_3 = block_class(pdcs[14], self.inplane, self.inplane)
self.block4_4 = block_class(pdcs[15], self.inplane, self.inplane)
self.fuseplanes.append(self.inplane) # 4C
self.conv_reduces = nn.ModuleList()
if self.sa and self.dil is not None:
self.attentions = nn.ModuleList()
self.dilations = nn.ModuleList()
for i in range(4):
self.dilations.append(CDCM(self.fuseplanes[i], self.dil))
self.attentions.append(CSAM(self.dil))
self.conv_reduces.append(MapReduce(self.dil))
elif self.sa:
self.attentions = nn.ModuleList()
for i in range(4):
self.attentions.append(CSAM(self.fuseplanes[i]))
self.conv_reduces.append(MapReduce(self.fuseplanes[i]))
elif self.dil is not None:
self.dilations = nn.ModuleList()
for i in range(4):
self.dilations.append(CDCM(self.fuseplanes[i], self.dil))
self.conv_reduces.append(MapReduce(self.dil))
else:
for i in range(4):
self.conv_reduces.append(MapReduce(self.fuseplanes[i]))
self.classifier = nn.Conv2d(4, 1, kernel_size=1) # has bias
nn.init.constant_(self.classifier.weight, 0.25)
nn.init.constant_(self.classifier.bias, 0)
def get_weights(self):
conv_weights = []
bn_weights = []
relu_weights = []
for pname, p in self.named_parameters():
if 'bn' in pname:
bn_weights.append(p)
elif 'relu' in pname:
relu_weights.append(p)
else:
conv_weights.append(p)
return conv_weights, bn_weights, relu_weights
def forward(self, x):
"""x: [B, 3, H, W] within range [0, 1].
"""
x = (x - x.new_tensor(self.mean).view(1, -1, 1, 1)) / \
x.new_tensor(self.std).view(1, -1, 1, 1)
h, w = x.size()[2:]
x = self.init_block(x)
x1 = self.block1_1(x)
x1 = self.block1_2(x1)
x1 = self.block1_3(x1)
x2 = self.block2_1(x1)
x2 = self.block2_2(x2)
x2 = self.block2_3(x2)
x2 = self.block2_4(x2)
x3 = self.block3_1(x2)
x3 = self.block3_2(x3)
x3 = self.block3_3(x3)
x3 = self.block3_4(x3)
x4 = self.block4_1(x3)
x4 = self.block4_2(x4)
x4 = self.block4_3(x4)
x4 = self.block4_4(x4)
x_fuses = []
if self.sa and self.dil is not None:
for i, xi in enumerate([x1, x2, x3, x4]):
x_fuses.append(self.attentions[i](self.dilations[i](xi)))
elif self.sa:
for i, xi in enumerate([x1, x2, x3, x4]):
x_fuses.append(self.attentions[i](xi))
elif self.dil is not None:
for i, xi in enumerate([x1, x2, x3, x4]):
x_fuses.append(self.dilations[i](xi))
else:
x_fuses = [x1, x2, x3, x4]
e1 = self.conv_reduces[0](x_fuses[0])
e1 = F.interpolate(e1, (h, w), mode='bilinear', align_corners=False)
e2 = self.conv_reduces[1](x_fuses[1])
e2 = F.interpolate(e2, (h, w), mode='bilinear', align_corners=False)
e3 = self.conv_reduces[2](x_fuses[2])
e3 = F.interpolate(e3, (h, w), mode='bilinear', align_corners=False)
e4 = self.conv_reduces[3](x_fuses[3])
e4 = F.interpolate(e4, (h, w), mode='bilinear', align_corners=False)
outputs = [e1, e2, e3, e4]
output = self.classifier(torch.cat(outputs, dim=1))
outputs.append(output)
outputs = [torch.sigmoid(r) for r in outputs]
return outputs[-1]
@ANNOTATORS.register_class()
class PiDiAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
vanilla_cnn = cfg.get('VANILLA_CNN', True)
pdcs = config_model_converted(
'carv4') if vanilla_cnn else config_model('carv4')
self.model = PiDiNet(60, pdcs, dil=24, sa=True,
convert=vanilla_cnn).eval()
if pretrained_model:
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
state = torch.load(local_path,
map_location='cpu')['state_dict']
if vanilla_cnn:
state = convert_pidinet(state, 'carv4')
state = {
k[len('module.'):] if k.startswith('module.') else k: v
for k, v in state.items()
}
self.model.load_state_dict(state)
@torch.no_grad()
@torch.inference_mode()
@torch.autocast('cuda', enabled=False)
def forward(self, image, return_grayscale=False):
is_batch = False if len(image.shape) == 3 else True
if isinstance(image, torch.Tensor):
if len(image.shape) == 3:
image = rearrange(image, 'h w c -> 1 c h w')
elif len(image.shape) == 4:
image = rearrange(image, 'b h w c -> b c h w')
else:
raise "Unsurpport input image's shape"
elif isinstance(image, np.ndarray):
image = torch.from_numpy(image.copy()).float()
if len(image.shape) == 3:
image = rearrange(image, 'h w c -> 1 c h w')
elif len(image.shape) == 4:
image = rearrange(image, 'b h w c -> b c h w')
else:
raise "Unsurpport input image's shape"
else:
raise "Unsurpport input image's type"
image = image.float().div(255)
image = image.to(we.device_id)
edge = self.model(image)
edge = edge.squeeze(dim=1)
edge = 255 - (edge * 255.0).clip(0, 255) # return white background
edge = edge.cpu().numpy()
edge = edge.astype(np.uint8)
if not is_batch:
edge = edge.squeeze()
if not return_grayscale:
edge = edge[..., None].repeat(3, -1)
return edge
@staticmethod
def get_config_template():
return dict_to_yaml('ANNOTATORS',
__class__.__name__,
PiDiAnnotator.para_dict,
set_name=True)
+1 -1
View File
@@ -15,7 +15,7 @@ def build_annotator(cfg, registry, logger=None, *args, **kwargs):
raise TypeError(f'Config must be type dict, got {type(cfg)}')
if cfg.have('PRETRAINED_MODEL'):
pretrain_cfg = cfg.PRETRAINED_MODEL
if pretrain_cfg is not None and not isinstance(pretrain_cfg, (str)):
if pretrain_cfg is not None and not isinstance(pretrain_cfg, (str, list)):
raise TypeError('Pretrain parameter must be a string')
else:
pretrain_cfg = None
+384
View File
@@ -0,0 +1,384 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import random
from abc import ABCMeta
import cv2
import numpy as np
import torch
import torchvision.transforms as T
from PIL import Image
from pycocotools import mask as mask_utils
from scipy import ndimage
from sklearn.cluster import KMeans
from torchvision.ops.boxes import batched_nms
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.distribute import we
from scepter.modules.utils.file_system import FS
def find_dominant_color(image, k=1):
pixels = image.reshape((-1, 3))
mask = (pixels != [0, 0, 0]).all(axis=1)
pixels = pixels[mask]
try:
kmeans = KMeans(n_clusters=k, n_init='auto')
kmeans.fit(pixels)
dominant_color = kmeans.cluster_centers_.astype(int)[0]
except Exception:
dominant_color = np.array([255, 255, 255])
return dominant_color
def cv2_resize_crop(image, resize_size, crop_size):
resize_height, resize_width = resize_size
crop_height, crop_width = crop_size
resized_image = cv2.resize(image, (resize_width, resize_height))
center_x, center_y = resize_width // 2, resize_height // 2
crop_start_x = max(center_x - crop_width // 2, 0)
crop_start_y = max(center_y - crop_height // 2, 0)
crop_end_x = crop_start_x + crop_width
crop_end_y = crop_start_y + crop_height
crop_end_x = min(crop_end_x, resize_width)
crop_end_y = min(crop_end_y, resize_height)
center_cropped_image = resized_image[crop_start_y:crop_end_y,
crop_start_x:crop_end_x]
return center_cropped_image
@ANNOTATORS.register_class()
class ESAMAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
try:
from efficient_sam.efficient_sam import build_efficient_sam
except Exception:
raise NotImplementedError(
'Please install efficient_sam and segment_anything modules.')
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
if pretrained_model:
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
self.efficient_sam_module = build_efficient_sam(
encoder_patch_embed_dim=384,
encoder_num_heads=6,
checkpoint=local_path).eval().to(we.device_id)
self.GRID_SIZE = cfg.get('GRID_SIZE', 16)
self.save_mode = cfg.get('SAVE_MODE', 'P')
self.use_dominant_color = cfg.get('USE_DOMINANT_COLOR', False)
self.return_mask = cfg.get('RETURN_MASK', False)
@torch.no_grad()
def get_predictions_given_embeddings_and_queries(self, img, points,
point_labels, model):
from segment_anything.utils.amg import calculate_stability_score
predicted_masks, predicted_iou = [], []
bs = 128
num = int(float(self.GRID_SIZE * self.GRID_SIZE) /
bs) if self.GRID_SIZE * self.GRID_SIZE % bs == 0 else int(
float(self.GRID_SIZE * self.GRID_SIZE) / bs) + 1
for i in range(num):
predicted_mask_item, predicted_iou_item = model(
img[None, ...], points[:, i * bs:(i + 1) * bs, ...],
point_labels[:, i * bs:(i + 1) * bs, :])
predicted_masks.append(predicted_mask_item)
predicted_iou.append(predicted_iou_item)
torch.cuda.empty_cache()
# # predicted: torch.Size([1, 1024, 3]) torch.Size([1, 1024, 3, 512, 512])
# print('predicted: ', predicted_iou.size(), predicted_masks.size())
predicted_masks = torch.cat(predicted_masks, dim=1)
predicted_iou = torch.cat(predicted_iou, dim=1)
sorted_ids = torch.argsort(predicted_iou, dim=-1, descending=True)
predicted_iou_scores = torch.take_along_dim(predicted_iou,
sorted_ids,
dim=2)
predicted_masks = torch.take_along_dim(predicted_masks,
sorted_ids[..., None, None],
dim=2)
predicted_masks = predicted_masks[0]
iou = predicted_iou_scores[0, :, 0]
index_iou = iou > 0.7
iou_ = iou[index_iou]
masks = predicted_masks[index_iou]
score = calculate_stability_score(masks, 0.0, 1.0)
score = score[:, 0]
index = score > 0.9
masks = masks[index]
iou_ = iou_[index]
masks = torch.ge(masks, 0.0)
return masks, iou_
def singel_mask_to_rle(self, mask):
rle = mask_utils.encode(
np.array(mask[:, :, None], order='F', dtype='uint8'))[0]
rle['counts'] = rle['counts'].decode('utf-8')
return rle
def process_small_region(self, rles):
from segment_anything.utils.amg import rle_to_mask, remove_small_regions, \
batched_mask_to_box, mask_to_rle_pytorch
new_masks = []
scores = []
min_area = 100
nms_thresh = 0.7
for rle in rles:
mask = rle_to_mask(rle[0])
mask, changed = remove_small_regions(mask, min_area, mode='holes')
unchanged = not changed
mask, changed = remove_small_regions(mask,
min_area,
mode='islands')
unchanged = unchanged and not changed
new_masks.append(torch.as_tensor(mask).unsqueeze(0))
# Give score=0 to changed masks and score=1 to unchanged masks
# so NMS will prefer ones that didn't need postprocessing
scores.append(float(unchanged))
# Recalculate boxes and remove any new duplicates
masks = torch.cat(new_masks, dim=0).to(we.device_id)
boxes = batched_mask_to_box(masks)
keep_by_nms = batched_nms(
boxes.float(),
torch.as_tensor(scores).to(we.device_id),
torch.zeros_like(boxes[:, 0]), # categories
iou_threshold=nms_thresh,
)
# Only recalculate RLEs for masks that have changed
for i_mask in keep_by_nms:
if scores[i_mask] == 0.0:
mask_torch = masks[i_mask].unsqueeze(0)
rles[i_mask] = mask_to_rle_pytorch(mask_torch)
masks = [rle_to_mask(rles[i][0]) for i in keep_by_nms]
return masks
def run_everything_ours(self, img_tensor, model):
from segment_anything.utils.amg import mask_to_rle_pytorch
img_tensor = img_tensor.squeeze(0)
_, original_image_h, original_image_w = img_tensor.shape
xy = []
for i in range(self.GRID_SIZE):
curr_x = 0.5 + i / self.GRID_SIZE * original_image_w
for j in range(self.GRID_SIZE):
curr_y = 0.5 + j / self.GRID_SIZE * original_image_h
xy.append([curr_x, curr_y])
xy = torch.from_numpy(np.array(xy))
points = xy
num_pts = xy.shape[0]
point_labels = torch.ones(num_pts, 1)
with torch.no_grad():
predicted_masks, predicted_iou = self.get_predictions_given_embeddings_and_queries(
img_tensor,
points.reshape(1, num_pts, 1, 2).to(we.device_id),
point_labels.reshape(1, num_pts, 1).to(we.device_id),
model,
)
# print('predicted_masks: ', predicted_masks[0][0:1].dtype, predicted_masks[0][0:1].device)
rle = [mask_to_rle_pytorch(m[0:1]) for m in predicted_masks]
# transform to numpy
size, counts = [], []
for rle_item in rle:
size.append(rle_item[0]['size'])
counts += rle_item[0]['counts']
counts += '#'
predicted_masks = self.process_small_region(rle)
return predicted_masks
def forward(self, image, return_mask=None):
return_mask = return_mask if return_mask is not None else self.return_mask
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.'
h, w = image.shape[:2]
max_rate = max(float(w) / 1024.0, float(h) / 1024.0)
w_ori = int(float(w) / max_rate)
h_ori = int(float(h) / max_rate)
# image = T.ToTensor()(T.Resize((h_ori, w_ori))(Image.fromarray(image)))
# image_pad = T.Pad((0, 0, 1024 - w_ori, 1024 - h_ori))(image)
image_pad = T.Pad((0, 0, 1024 - w_ori, 1024 - h_ori))(T.Resize(
(h_ori, w_ori))(Image.fromarray(image)))
input_image = T.ToTensor()(image_pad)
input_image = input_image.unsqueeze(0).to(we.device_id)
mask_efficient_sam_vits = self.run_everything_ours(
input_image, self.efficient_sam_module)
annos = []
mask_efficient_sam_vits = sorted(list(mask_efficient_sam_vits),
key=lambda m: int(m.sum()),
reverse=True)
mask_efficient_sam_vits = mask_efficient_sam_vits[:256]
for mask in mask_efficient_sam_vits:
mask_item = mask_utils.encode(
np.array(mask[:, :, None], order='F', dtype='uint8'))[0]
mask_item['counts'] = mask_item['counts'].decode('utf-8')
mask_area = int(mask.sum())
annos.append({'mask': mask_item, 'mask_area': mask_area})
annos = sorted(annos, key=lambda x: x['mask_area'], reverse=True)
seg_img = None
dominant_palette = []
image_pad_np = np.array(image_pad)
for idx, anno in enumerate(annos):
color = idx
if idx > 255:
break
mask = np.array(mask_utils.decode(anno['mask'])).astype(np.uint8)
h, w = mask.shape[:2]
if seg_img is None:
seg_img = np.ones((h, w, 3)) * 255
if self.use_dominant_color:
masked_image = cv2.bitwise_and(image_pad_np,
image_pad_np,
mask=mask)
dominant_color = find_dominant_color(masked_image).tolist()
dominant_palette.append(dominant_color)
seg_img[mask.astype(bool)] = [color, color, color]
seg_img = Image.fromarray(seg_img.astype(np.uint8)).convert('L')
resize_rate = max(float(h_ori) / 1024.0, float(w_ori) / 1024.0)
h_new = int(float(h_ori) / resize_rate)
w_new = int(float(w_ori) / resize_rate)
seg_img = seg_img.crop((0, 0, w_new, h_new))
if self.save_mode == 'P':
palette = []
for i in range(256):
if not self.use_dominant_color:
palette_item = [random.randint(0, 255) for _ in range(3)]
else:
palette_item = dominant_palette[i] if i < len(
dominant_palette) else [255, 255, 255]
palette += palette_item
seg_img = seg_img.convert('P')
seg_img.putpalette(palette)
seg_rgb_img = seg_img.convert('RGB')
if return_mask:
return {
'image': np.array(seg_rgb_img),
'mask': np.array(seg_img)
}
else:
return np.array(seg_rgb_img)
else:
return np.array(seg_img)
@staticmethod
def get_config_template():
return dict_to_yaml('ANNOTATORS',
__class__.__name__,
ESAMAnnotator.para_dict,
set_name=True)
@ANNOTATORS.register_class()
class SAMAnnotatorDraw(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
from segment_anything import sam_model_registry, SamPredictor
from segment_anything.utils.transforms import ResizeLongestSide
self.transform = ResizeLongestSide(1024)
self.task_type = cfg.get('TASK_TYPE', 'input_box')
self.sam_model = cfg.get('SAM_MODEL', 'vit_b')
pretrained_model = cfg.get('PRETRAINED_MODEL', 'sam_vit_b_01ec64.pth')
if pretrained_model:
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
seg_model = sam_model_registry[self.sam_model](
checkpoint=local_path).eval().to(we.device_id)
self.sam_predictor = SamPredictor(seg_model)
def forward(self,
image,
input_box=None,
mask=None,
task_type=None,
multimask_output=False):
task_type = task_type if task_type is not None else self.task_type
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.'
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':
scribble = mask.transpose(2, 1, 0)[0]
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 = {
'point_coords': point_coords,
'point_labels': point_labels
}
elif task_type == 'mask_box':
scribble = mask.transpose(2, 1, 0)[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()
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}
self.sam_predictor.set_image(image)
masks, scores, logits = self.sam_predictor.predict(
**sample, multimask_output=True)
index = np.argmax(scores)
ret_data = {
'mask': (masks[index] * 255).astype(np.uint8),
'score': scores[index]
}
return ret_data
@staticmethod
def get_config_template():
return dict_to_yaml('ANNOTATORS',
__class__.__name__,
SAMAnnotatorDraw.para_dict,
set_name=True)
+157
View File
@@ -0,0 +1,157 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from abc import ABCMeta
import numpy as np
import torch
import torch.nn as nn
from einops import rearrange
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.distribute import we
from scepter.modules.utils.file_system import FS
class SketchNet(nn.Module):
def __init__(self, mean, std):
assert isinstance(mean, float) and isinstance(std, float)
super().__init__()
self.mean = mean
self.std = std
# layers
self.layers = nn.Sequential(nn.Conv2d(1, 48, 5, 2, 2),
nn.ReLU(inplace=True),
nn.Conv2d(48, 128, 3, 1, 1),
nn.ReLU(inplace=True),
nn.Conv2d(128, 128, 3, 1, 1),
nn.ReLU(inplace=True),
nn.Conv2d(128, 128, 3, 2, 1),
nn.ReLU(inplace=True),
nn.Conv2d(128, 256, 3, 1, 1),
nn.ReLU(inplace=True),
nn.Conv2d(256, 256, 3, 1, 1),
nn.ReLU(inplace=True),
nn.Conv2d(256, 256, 3, 2, 1),
nn.ReLU(inplace=True),
nn.Conv2d(256, 512, 3, 1, 1),
nn.ReLU(inplace=True),
nn.Conv2d(512, 1024, 3, 1, 1),
nn.ReLU(inplace=True),
nn.Conv2d(1024, 1024, 3, 1, 1),
nn.ReLU(inplace=True),
nn.Conv2d(1024, 1024, 3, 1, 1),
nn.ReLU(inplace=True),
nn.Conv2d(1024, 1024, 3, 1, 1),
nn.ReLU(inplace=True),
nn.Conv2d(1024, 512, 3, 1, 1),
nn.ReLU(inplace=True),
nn.Conv2d(512, 256, 3, 1, 1),
nn.ReLU(inplace=True),
nn.ConvTranspose2d(256, 256, 4, 2, 1),
nn.ReLU(inplace=True),
nn.Conv2d(256, 256, 3, 1, 1),
nn.ReLU(inplace=True),
nn.Conv2d(256, 128, 3, 1, 1),
nn.ReLU(inplace=True),
nn.ConvTranspose2d(128, 128, 4, 2, 1),
nn.ReLU(inplace=True),
nn.Conv2d(128, 128, 3, 1, 1),
nn.ReLU(inplace=True),
nn.Conv2d(128, 48, 3, 1, 1),
nn.ReLU(inplace=True),
nn.ConvTranspose2d(48, 48, 4, 2, 1),
nn.ReLU(inplace=True),
nn.Conv2d(48, 24, 3, 1, 1),
nn.ReLU(inplace=True),
nn.Conv2d(24, 1, 3, 1, 1), nn.Sigmoid())
def forward(self, x):
"""x: [B, 1, H, W] within range [0, 1]. Sketch pixels in dark color.
"""
x = (x - self.mean) / self.std
return self.layers(x)
@ANNOTATORS.register_class()
class SketchAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
self.model = SketchNet(mean=0.9664114577640158,
std=0.0858381272736797).eval()
if pretrained_model:
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
state = torch.load(local_path, map_location='cpu')
self.model.load_state_dict(state)
@torch.no_grad()
@torch.inference_mode()
@torch.autocast('cuda', enabled=False)
def forward(self, image):
is_batch = False if len(image.shape) == 3 else True
if isinstance(image, torch.Tensor):
if len(image.shape) == 3:
if torch.equal(image[:, :, 0], image[:, :, 1]) and torch.equal(
image[:, :, 1], image[:, :, 2]):
image = image[:, :, 0]
else:
raise "Unsurpport input image's shape and each channel is different"
elif len(image.shape) == 4:
if (torch.equal(image[:, :, :, 0], image[:, :, :, 1])
and torch.equal(image[:, :, :, 1], image[:, :, :, 2])):
image = image[:, :, :, 0]
else:
raise "Unsurpport input image's shape and each channel is different"
if len(image.shape) == 2:
image = rearrange(image, 'h w -> 1 h w')
B, H, W = image.shape
elif len(image.shape) == 3:
B, H, W = image.shape
else:
raise "Unsurpport input image's shape"
elif isinstance(image, np.ndarray):
image = image.copy()
if len(image.shape) == 3:
if np.array_equal(image[:, :, 0],
image[:, :, 1]) and np.array_equal(
image[:, :, 1], image[:, :, 2]):
image = image[:, :, 0]
else:
raise "Unsurpport input image's shape and each channel is different"
elif len(image.shape) == 4:
if (np.array_equal(image[:, :, :, 0], image[:, :, :, 1]) and
np.array_equal(image[:, :, :, 1], image[:, :, :, 2])):
image = image[:, :, :, 0]
else:
raise "Unsurpport input image's shape and each channel is different"
image = torch.from_numpy(image).float()
if len(image.shape) == 2:
image = rearrange(image, 'h w -> 1 1 h w')
elif len(image.shape) == 3:
image = rearrange(image, 'b h w -> b 1 h w')
else:
raise "Unsurpport input image's shape"
else:
raise "Unsurpport input image's type"
image = image.float().div(255)
image = image.to(we.device_id)
edge = self.model(image)
edge = edge.squeeze(dim=1)
edge = (edge * 255.0).clip(0, 255)
edge = edge.cpu().numpy()
edge = edge.astype(np.uint8)
if not is_batch:
edge = edge.squeeze()
return edge[..., None].repeat(3, -1)
@staticmethod
def get_config_template():
return dict_to_yaml('ANNOTATORS',
__class__.__name__,
SketchAnnotator.para_dict,
set_name=True)
+2 -1
View File
@@ -7,5 +7,6 @@ from scepter.modules.data.dataset.dataset import (Image2ImageDataset,
ImageTextPairDataset,
Text2ImageDataset)
from scepter.modules.data.dataset.ms_dataset import (
ImageTextPairFolderDataset, ImageTextPairMSDataset)
ImageTextPairFolderDataset, ImageTextPairMSDataset,
ImageTextPairMSDatasetForACE)
from scepter.modules.data.dataset.registry import DATASETS
@@ -82,6 +82,8 @@ class BaseDataset(Dataset, metaclass=ABCMeta):
overwrite=False)
self.worker_id = worker_id
self.logger = self.worker_logger
self.local_we["seed"] += (worker_id + we.rank)
self.seed = self.local_we["seed"]
we.set_env(self.local_we)
@abstractmethod
+259 -3
View File
@@ -1,16 +1,28 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import io
import math
import numbers
import os
import sys
from collections import defaultdict
import numpy as np
import torch
import torchvision.transforms as T
from PIL import Image
from torchvision.transforms.functional import InterpolationMode
from scepter.modules.data.dataset.base_dataset import BaseDataset
from scepter.modules.data.dataset.registry import DATASETS
from scepter.modules.transform.io import pillow_convert
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
Image.MAX_IMAGE_PIXELS = None
@DATASETS.register_class()
class ImageTextPairMSDataset(BaseDataset):
@@ -102,7 +114,7 @@ class ImageTextPairMSDataset(BaseDataset):
self.output_size = [self.output_size, self.output_size]
# Use modelscope dataset
if not ms_dataset_name:
raise (
raise ValueError(
'Your must set MS_DATASET_NAME as modelscope dataset or your local dataset orignized '
'as modelscope dataset.')
if FS.exists(ms_dataset_name):
@@ -125,7 +137,7 @@ class ImageTextPairMSDataset(BaseDataset):
split=ms_dataset_split,
download_mode=DownloadMode.FORCE_REDOWNLOAD)
except Exception as sec_e:
raise f'Load Modelscope dataset failed {sec_e}.'
raise ValueError(f'Load Modelscope dataset failed {sec_e}.')
if ms_remap_keys:
self.data = self.data.remap_columns(ms_remap_keys.get_dict())
@@ -245,7 +257,7 @@ class ImageTextPairFolderDataset(BaseDataset):
self.output_size = [self.output_size, self.output_size]
# Use modelscope dataset
if not data_folder or not FS.exists(data_folder):
raise ('Your must set datafolder for local dataset.')
raise ValueError('Your must set datafolder for local dataset.')
data_folder = FS.get_dir_to_local_dir(data_folder)
all_lines = open(os.path.join(data_folder, 'train.csv'),
'r').read().split('\n')
@@ -311,3 +323,247 @@ class ImageTextPairFolderDataset(BaseDataset):
__class__.__name__,
ImageTextPairMSDataset.para_dict,
set_name=True)
@DATASETS.register_class()
class ImageTextPairMSDatasetForACE(BaseDataset):
para_dict = {
'MS_DATASET_NAME': {
'value': '',
'description': 'Modelscope dataset name.'
},
'MS_DATASET_NAMESPACE': {
'value': '',
'description': 'Modelscope dataset namespace.'
},
'MS_DATASET_SUBNAME': {
'value': '',
'description': 'Modelscope dataset subname.'
},
'MS_DATASET_SPLIT': {
'value': '',
'description':
'Modelscope dataset split set name, default is train.'
},
'MS_REMAP_KEYS': {
'value':
None,
'description':
'Modelscope dataset header of list file, the default is Target:FILE; '
'If your file is not this header, please set this field, which is a map dict.'
"For example, { 'Image:FILE': 'Target:FILE' } will replace the filed Image:FILE to Target:FILE"
},
'MS_REMAP_PATH': {
'value':
None,
'description':
'When modelscope dataset name is not None, that means you use the dataset from modelscope,'
' default is None. But if you want to use the datalist from modelscope and the file from '
'local device, you can use this field to set the root path of your images. '
},
'TRIGGER_WORDS': {
'value':
'',
'description':
'The words used to describe the common features of your data, especially when you customize a '
'tuner. Use these words you can get what you want.'
},
'REPLACE_STYLE': {
'value':
False,
'description':
'Whether use the MS_DATASET_SUBNAME to replace the word in your description, default is False.'
},
'HIGHLIGHT_KEYWORDS': {
'value':
'',
'description':
'The keywords you want to highlight in prompt, which will be replace by <HIGHLIGHT_KEYWORDS>.'
},
'KEYWORDS_SIGN': {
'value':
'',
'description':
'The keywords sign you want to add, which is like <{HIGHLIGHT_KEYWORDS}{KEYWORDS_SIGN}>'
},
'OUTPUT_SIZE': {
'value':
None,
'description':
'If you use the FlexibleResize transforms, this filed will output the image_size as [h, w],'
'which will be used to set the output size of images used to train the model.'
},
}
def __init__(self, cfg, logger=None):
super().__init__(cfg=cfg, logger=logger)
from modelscope import MsDataset
from modelscope.utils.constant import DownloadMode
ms_dataset_name = cfg.get('MS_DATASET_NAME', None)
ms_dataset_namespace = cfg.get('MS_DATASET_NAMESPACE', None)
ms_dataset_subname = cfg.get('MS_DATASET_SUBNAME', None)
ms_dataset_split = cfg.get('MS_DATASET_SPLIT', 'train')
ms_remap_keys = cfg.get('MS_REMAP_KEYS', None)
ms_remap_path = cfg.get('MS_REMAP_PATH', None)
self.max_seq_len = cfg.get('MAX_SEQ_LEN', 1024)
self.max_aspect_ratio = cfg.get('MAX_ASPECT_RATIO', 4)
self.d = cfg.get('DOWNSAMPLE_RATIO', 16)
self.replace_style = cfg.get('REPLACE_STYLE', False)
self.trigger_words = cfg.get('TRIGGER_WORDS', '')
self.replace_keywords = cfg.get('HIGHLIGHT_KEYWORDS', '')
self.keywords_sign = cfg.get('KEYWORDS_SIGN', '')
self.add_indicator = cfg.get('ADD_INDICATOR', False)
# Use modelscope dataset
if not ms_dataset_name:
raise ValueError(
'Your must set MS_DATASET_NAME as modelscope dataset or your local dataset orignized '
'as modelscope dataset.')
if FS.exists(ms_dataset_name):
ms_dataset_name = FS.get_dir_to_local_dir(ms_dataset_name)
self.ms_dataset_name = ms_dataset_name
# ms_remap_path = ms_dataset_name
try:
self.data = MsDataset.load(str(ms_dataset_name),
namespace=ms_dataset_namespace,
subset_name=ms_dataset_subname,
split=ms_dataset_split)
except Exception:
self.logger.info(
"Load Modelscope dataset failed, retry with download_mode='force_redownload'."
)
try:
self.data = MsDataset.load(
str(ms_dataset_name),
namespace=ms_dataset_namespace,
subset_name=ms_dataset_subname,
split=ms_dataset_split,
download_mode=DownloadMode.FORCE_REDOWNLOAD)
except Exception as sec_e:
raise ValueError(f'Load Modelscope dataset failed {sec_e}.')
if ms_remap_keys:
self.data = self.data.remap_columns(ms_remap_keys.get_dict())
if ms_remap_path:
def map_func(example):
return {
k: os.path.join(ms_remap_path, v)
if k.endswith(':FILE') else v
for k, v in example.items()
}
self.data = self.data.ds_instance.map(map_func)
self.transforms = T.Compose([
T.ToTensor(),
T.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
])
def __len__(self):
if self.mode == 'train':
return sys.maxsize
else:
return len(self.data)
def _get(self, index: int):
current_data = self.data[index % len(self.data)]
tar_image_path = current_data.get('Target:FILE', '')
src_image_path = current_data.get('Source:FILE', '')
style = current_data.get('Style', '')
prompt = current_data.get('Prompt', current_data.get('prompt', ''))
if self.replace_style and not style == '':
prompt = prompt.replace(style, f'<{self.keywords_sign}>')
elif not self.replace_keywords.strip() == '':
prompt = prompt.replace(
self.replace_keywords,
'<' + self.replace_keywords + f'{self.keywords_sign}>')
if not self.trigger_words == '':
prompt = self.trigger_words.strip() + ' ' + prompt
src_image = self.load_image(self.ms_dataset_name,
src_image_path,
cvt_type='RGB')
tar_image = self.load_image(self.ms_dataset_name,
tar_image_path,
cvt_type='RGB')
src_image = self.image_preprocess(src_image)
tar_image = self.image_preprocess(tar_image)
tar_image = self.transforms(tar_image)
src_image = self.transforms(src_image)
src_mask = torch.ones_like(src_image[[0]])
tar_mask = torch.ones_like(tar_image[[0]])
if self.add_indicator:
if '{image}' not in prompt:
prompt = '{image}, ' + prompt
return {
'edit_image': [src_image],
'edit_image_mask': [src_mask],
'image': tar_image,
'image_mask': tar_mask,
'prompt': [prompt],
}
def load_image(self, prefix, img_path, cvt_type=None):
if img_path is None or img_path == '':
return None
img_path = os.path.join(prefix, img_path)
with FS.get_object(img_path) as image_bytes:
image = Image.open(io.BytesIO(image_bytes))
if cvt_type is not None:
image = pillow_convert(image, cvt_type)
return image
def image_preprocess(self,
img,
size=None,
interpolation=InterpolationMode.BILINEAR):
H, W = img.height, img.width
if H / W > self.max_aspect_ratio:
img = T.CenterCrop((self.max_aspect_ratio * W, W))(img)
elif W / H > self.max_aspect_ratio:
img = T.CenterCrop((H, self.max_aspect_ratio * H))(img)
if size is None:
# resize image for max_seq_len, while keep the aspect ratio
H, W = img.height, img.width
scale = min(
1.0,
math.sqrt(self.max_seq_len / ((H / self.d) * (W / self.d))))
rH = int(
H * scale) // self.d * self.d # ensure divisible by self.d
rW = int(W * scale) // self.d * self.d
else:
rH, rW = size
img = T.Resize((rH, rW), interpolation=interpolation,
antialias=True)(img)
return np.array(img, dtype=np.uint8)
@staticmethod
def get_config_template():
return dict_to_yaml('DATASet',
__class__.__name__,
ImageTextPairMSDatasetForACE.para_dict,
set_name=True)
@staticmethod
def collate_fn(batch):
collect = defaultdict(list)
for sample in batch:
for k, v in sample.items():
collect[k].append(v)
new_batch = dict()
for k, v in collect.items():
if all([i is None for i in v]):
new_batch[k] = None
else:
new_batch[k] = v
return new_batch
+3 -1
View File
@@ -258,6 +258,8 @@ class DataObject(object):
if sampler_name == 'MixtureOfSamplers':
subsampler_configs = self.data_sampler_config.get(
'SUB_SAMPLERS', [])
keep_order = self.data_sampler_config.get(
'KEEP_ORDER', False)
subsamplers = list()
subsampler_probs = list()
for ssconfig in subsampler_configs:
@@ -277,7 +279,7 @@ class DataObject(object):
subsampler_probs.append(prob)
self.batch_sampler = MixtureOfSamplers(subsamplers,
subsampler_probs, rank,
seed)
seed, keep_order = keep_order)
elif sampler_name == 'MultiLevelBatchSampler':
self.batch_sampler = self._instantiate_multi_level_batch_sampler(
self.data_sampler_config, self.batch_size, rank, seed)
+5 -2
View File
@@ -523,12 +523,15 @@ class MultiLevelBatchSampler(BaseSampler):
class MixtureOfSamplers(BaseSampler):
para_dict = {'SUB_SAMPLERS': []}
def __init__(self, samplers, probabilities, rank=0, seed=8888):
def __init__(self, samplers, probabilities, rank=0, seed=8888, keep_order = False):
self.samplers = samplers
self.iterators = [iter(u) for u in samplers]
self.probabilities = probabilities
self.seed = seed
self.rng = np.random.default_rng(seed + rank)
if keep_order:
self.rng = np.random.default_rng(seed)
else:
self.rng = np.random.default_rng(seed + rank)
def __iter__(self):
while True:
+376
View File
@@ -0,0 +1,376 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import math
import random
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision.transforms.functional as TF
from PIL import Image
from scepter.modules.model.registry import DIFFUSIONS
from scepter.modules.model.utils.basic_utils import check_list_of_list
from scepter.modules.model.utils.basic_utils import \
pack_imagelist_into_tensor_v2 as pack_imagelist_into_tensor
from scepter.modules.model.utils.basic_utils import (
to_device, unpack_tensor_into_imagelist)
from scepter.modules.utils.distribute import we
from scepter.modules.utils.logger import get_logger
from .diffusion_inference import DiffusionInference, get_model
def process_edit_image(images,
masks,
tasks,
max_seq_len=1024,
max_aspect_ratio=4,
d=16,
**kwargs):
if not isinstance(images, list):
images = [images]
if not isinstance(masks, list):
masks = [masks]
if not isinstance(tasks, list):
tasks = [tasks]
img_tensors = []
mask_tensors = []
for img, mask, task in zip(images, masks, tasks):
if mask is None or mask == '':
mask = Image.new('L', img.size, 0)
W, H = img.size
if H / W > max_aspect_ratio:
img = TF.center_crop(img, [int(max_aspect_ratio * W), W])
mask = TF.center_crop(mask, [int(max_aspect_ratio * W), W])
elif W / H > max_aspect_ratio:
img = TF.center_crop(img, [H, int(max_aspect_ratio * H)])
mask = TF.center_crop(mask, [H, int(max_aspect_ratio * H)])
H, W = img.height, img.width
scale = min(1.0, math.sqrt(max_seq_len / ((H / d) * (W / d))))
rH = int(H * scale) // d * d # ensure divisible by self.d
rW = int(W * scale) // d * d
img = TF.resize(img, (rH, rW),
interpolation=TF.InterpolationMode.BICUBIC)
mask = TF.resize(mask, (rH, rW),
interpolation=TF.InterpolationMode.NEAREST_EXACT)
mask = np.asarray(mask)
mask = np.where(mask > 128, 1, 0)
mask = mask.astype(
np.float32) if np.any(mask) else np.ones_like(mask).astype(
np.float32)
img_tensor = TF.to_tensor(img).to(we.device_id)
img_tensor = TF.normalize(img_tensor,
mean=[0.5, 0.5, 0.5],
std=[0.5, 0.5, 0.5])
mask_tensor = TF.to_tensor(mask).to(we.device_id)
if task in ['inpainting', 'Try On', 'Inpainting']:
mask_indicator = mask_tensor.repeat(3, 1, 1)
img_tensor[mask_indicator == 1] = -1.0
img_tensors.append(img_tensor)
mask_tensors.append(mask_tensor)
return img_tensors, mask_tensors
class TextEmbedding(nn.Module):
def __init__(self, embedding_shape):
super().__init__()
self.pos = nn.Parameter(data=torch.zeros(embedding_shape))
class ACEInference(DiffusionInference):
def __init__(self, logger=None):
if logger is None:
logger = get_logger(name='scepter')
self.logger = logger
self.loaded_model = {}
self.loaded_model_name = [
'diffusion_model', 'first_stage_model', 'cond_stage_model'
]
def init_from_cfg(self, cfg):
self.name = cfg.NAME
self.is_default = cfg.get('IS_DEFAULT', False)
module_paras = self.load_default(cfg.get('DEFAULT_PARAS', None))
assert cfg.have('MODEL')
self.diffusion_model = self.infer_model(
cfg.MODEL.DIFFUSION_MODEL, module_paras.get(
'DIFFUSION_MODEL',
None)) if cfg.MODEL.have('DIFFUSION_MODEL') else None
self.first_stage_model = self.infer_model(
cfg.MODEL.FIRST_STAGE_MODEL,
module_paras.get(
'FIRST_STAGE_MODEL',
None)) if cfg.MODEL.have('FIRST_STAGE_MODEL') else None
self.cond_stage_model = self.infer_model(
cfg.MODEL.COND_STAGE_MODEL,
module_paras.get(
'COND_STAGE_MODEL',
None)) if cfg.MODEL.have('COND_STAGE_MODEL') else None
self.diffusion = DIFFUSIONS.build(cfg.MODEL.DIFFUSION,
logger=self.logger)
self.interpolate_func = lambda x: (F.interpolate(
x.unsqueeze(0),
scale_factor=1 / self.size_factor,
mode='nearest-exact') if x is not None else None)
self.text_indentifers = cfg.MODEL.get('TEXT_IDENTIFIER', [])
self.use_text_pos_embeddings = cfg.MODEL.get('USE_TEXT_POS_EMBEDDINGS',
False)
if self.use_text_pos_embeddings:
self.text_position_embeddings = TextEmbedding(
(10, 4096)).eval().requires_grad_(False).to(we.device_id)
else:
self.text_position_embeddings = None
self.max_seq_len = cfg.MODEL.DIFFUSION_MODEL.MAX_SEQ_LEN
self.scale_factor = cfg.get('SCALE_FACTOR', 0.18215)
self.size_factor = cfg.get('SIZE_FACTOR', 8)
self.decoder_bias = cfg.get('DECODER_BIAS', 0)
self.default_n_prompt = cfg.get('DEFAULT_N_PROMPT', '')
@torch.no_grad()
def encode_first_stage(self, x, **kwargs):
_, dtype = self.get_function_info(self.first_stage_model, 'encode')
with torch.autocast('cuda',
enabled=(dtype != 'float32'),
dtype=getattr(torch, dtype)):
z = [
self.scale_factor * get_model(self.first_stage_model)._encode(
i.unsqueeze(0).to(getattr(torch, dtype))) for i in x
]
return z
@torch.no_grad()
def decode_first_stage(self, z):
_, dtype = self.get_function_info(self.first_stage_model, 'decode')
with torch.autocast('cuda',
enabled=(dtype != 'float32'),
dtype=getattr(torch, dtype)):
x = [
get_model(self.first_stage_model)._decode(
1. / self.scale_factor * i.to(getattr(torch, dtype)))
for i in z
]
return x
@torch.no_grad()
def __call__(self,
image=None,
mask=None,
prompt='',
task=None,
negative_prompt='',
output_height=512,
output_width=512,
sampler='ddim',
sample_steps=20,
guide_scale=4.5,
guide_rescale=0.5,
seed=-1,
history_io=None,
tar_index=0,
**kwargs):
input_image, input_mask = image, mask
g = torch.Generator(device=we.device_id)
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
g.manual_seed(int(seed))
if input_image is not None:
# assert isinstance(input_image, list) and isinstance(input_mask, list)
if task is None:
task = [''] * len(input_image)
if not isinstance(prompt, list):
prompt = [prompt] * len(input_image)
if history_io is not None and len(history_io) > 0:
his_image, his_maks, his_prompt, his_task = history_io[
'image'], history_io['mask'], history_io[
'prompt'], history_io['task']
assert len(his_image) == len(his_maks) == len(
his_prompt) == len(his_task)
input_image = his_image + input_image
input_mask = his_maks + input_mask
task = his_task + task
prompt = his_prompt + [prompt[-1]]
prompt = [
pp.replace('{image}', f'{{image{i}}}') if i > 0 else pp
for i, pp in enumerate(prompt)
]
edit_image, edit_image_mask = process_edit_image(
input_image, input_mask, task, max_seq_len=self.max_seq_len)
image, image_mask = edit_image[tar_index], edit_image_mask[
tar_index]
edit_image, edit_image_mask = [edit_image], [edit_image_mask]
else:
edit_image = edit_image_mask = [[]]
image = torch.zeros(
size=[3, int(output_height),
int(output_width)])
image_mask = torch.ones(
size=[1, int(output_height),
int(output_width)])
if not isinstance(prompt, list):
prompt = [prompt]
image, image_mask, prompt = [image], [image_mask], [prompt]
assert check_list_of_list(prompt) and check_list_of_list(
edit_image) and check_list_of_list(edit_image_mask)
# Assign Negative Prompt
if isinstance(negative_prompt, list):
negative_prompt = negative_prompt[0]
assert isinstance(negative_prompt, str)
n_prompt = copy.deepcopy(prompt)
for nn_p_id, nn_p in enumerate(n_prompt):
assert isinstance(nn_p, list)
n_prompt[nn_p_id][-1] = negative_prompt
ctx, null_ctx = {}, {}
# Get Noise Shape
self.dynamic_load(self.first_stage_model, 'first_stage_model')
image = to_device(image)
x = self.encode_first_stage(image)
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=True)
noise = [
torch.empty(*i.shape, device=we.device_id).normal_(generator=g)
for i in x
]
noise, x_shapes = pack_imagelist_into_tensor(noise)
ctx['x_shapes'] = null_ctx['x_shapes'] = x_shapes
image_mask = to_device(image_mask, strict=False)
cond_mask = [self.interpolate_func(i) for i in image_mask
] if image_mask is not None else [None] * len(image)
ctx['x_mask'] = null_ctx['x_mask'] = cond_mask
# Encode Prompt
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
function_name, dtype = self.get_function_info(self.cond_stage_model)
cont, cont_mask = getattr(get_model(self.cond_stage_model),
function_name)(prompt)
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont,
cont_mask)
null_cont, null_cont_mask = getattr(get_model(self.cond_stage_model),
function_name)(n_prompt)
null_cont, null_cont_mask = self.cond_stage_embeddings(
prompt, edit_image, null_cont, null_cont_mask)
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=False)
ctx['crossattn'] = cont
null_ctx['crossattn'] = null_cont
# Encode Edit Images
self.dynamic_load(self.first_stage_model, 'first_stage_model')
edit_image = [to_device(i, strict=False) for i in edit_image]
edit_image_mask = [to_device(i, strict=False) for i in edit_image_mask]
e_img, e_mask = [], []
for u, m in zip(edit_image, edit_image_mask):
if u is None:
continue
if m is None:
m = [None] * len(u)
e_img.append(self.encode_first_stage(u, **kwargs))
e_mask.append([self.interpolate_func(i) for i in m])
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=True)
null_ctx['edit'] = ctx['edit'] = e_img
null_ctx['edit_mask'] = ctx['edit_mask'] = e_mask
# Diffusion Process
self.dynamic_load(self.diffusion_model, 'diffusion_model')
function_name, dtype = self.get_function_info(self.diffusion_model)
with torch.autocast('cuda',
enabled=dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, dtype)):
latent = self.diffusion.sample(
noise=noise,
sampler=sampler,
model=get_model(self.diffusion_model),
model_kwargs=[{
'cond':
ctx,
'mask':
cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
}, {
'cond':
null_ctx,
'mask':
null_cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
}] if guide_scale is not None and guide_scale > 1 else {
'cond':
null_ctx,
'mask':
cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
},
steps=sample_steps,
show_progress=True,
seed=seed,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
return_intermediate=None,
**kwargs)
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=False)
# Decode to Pixel Space
self.dynamic_load(self.first_stage_model, 'first_stage_model')
samples = unpack_tensor_into_imagelist(latent, x_shapes)
x_samples = self.decode_first_stage(samples)
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=False)
imgs = [
torch.clamp((x_i + 1.0) / 2.0 + self.decoder_bias / 255,
min=0.0,
max=1.0).squeeze(0).permute(1, 2, 0).cpu().numpy()
for x_i in x_samples
]
imgs = [Image.fromarray((img * 255).astype(np.uint8)) for img in imgs]
return imgs
def cond_stage_embeddings(self, prompt, edit_image, cont, cont_mask):
if self.use_text_pos_embeddings and not torch.sum(
self.text_position_embeddings.pos) > 0:
identifier_cont, _ = getattr(get_model(self.cond_stage_model),
'encode')(self.text_indentifers,
return_mask=True)
self.text_position_embeddings.load_state_dict(
{'pos': identifier_cont[:, 0, :]})
cont_, cont_mask_ = [], []
for pp, edit, c, cm in zip(prompt, edit_image, cont, cont_mask):
if isinstance(pp, list):
cont_.append([c[-1], *c] if len(edit) > 0 else [c[-1]])
cont_mask_.append([cm[-1], *cm] if len(edit) > 0 else [cm[-1]])
else:
raise NotImplementedError
return cont_, cont_mask_
+10 -6
View File
@@ -13,12 +13,6 @@ from scepter.modules.utils.config import Config
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
try:
from swift import SwiftModel
except Exception:
warnings.warn('Import swift failed, please check it.')
class ControlInference():
def __init__(self, logger=None):
self.logger = logger
@@ -26,6 +20,11 @@ class ControlInference():
# @classmethod
def unregister_controllers(self, control_model_ins, diffusion_model):
try:
from swift import SwiftModel
except Exception:
warnings.warn('Import swift failed, please check it.')
self.logger.info('Unloading control model')
if isinstance(diffusion_model['model'], SwiftModel):
if (hasattr(diffusion_model['model'].base_model, 'control_blocks')
@@ -42,6 +41,11 @@ class ControlInference():
# @classmethod
def register_controllers(self, control_model_ins, diffusion_model):
try:
from swift import SwiftModel
except Exception:
warnings.warn('Import swift failed, please check it.')
self.logger.info('Loading control model')
if control_model_ins is None or control_model_ins == '':
self.unregister_controllers(control_model_ins, diffusion_model)
@@ -11,7 +11,7 @@ from PIL.Image import Image
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
from scepter.modules.model.network.diffusion.schedules import noise_schedule
from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, MODELS,
TOKENIZERS)
TOKENIZERS, DIFFUSIONS)
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from scepter.studio.utils.env import get_available_memory
@@ -49,7 +49,10 @@ class DiffusionInference():
assert cfg.have('MODEL')
if self.is_redefine_paras:
cfg.MODEL = self.redefine_paras(cfg.MODEL)
self.diffusion = self.load_schedule(cfg.MODEL.SCHEDULE)
if 'DIFFUSION' in cfg.MODEL:
self.diffusion = DIFFUSIONS.build(cfg.MODEL.DIFFUSION, logger=self.logger)
else:
self.diffusion = self.load_schedule(cfg.MODEL.SCHEDULE)
self.diffusion_model = self.infer_model(
cfg.MODEL.DIFFUSION_MODEL, module_paras.get(
'DIFFUSION_MODEL',
@@ -84,14 +87,16 @@ class DiffusionInference():
def redefine_paras(self, cfg):
if cfg.get('PRETRAINED_MODEL', None):
assert FS.isfile(cfg.PRETRAINED_MODEL)
with FS.get_from(cfg.PRETRAINED_MODEL,
wait_finish=True) as local_path:
if local_path.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(local_path)
else:
sd = torch.load(local_path, map_location='cpu')
if 'weights_only' in torch.load.__code__.co_varnames:
sd = torch.load(local_path, map_location='cpu', weights_only=True)
else:
sd = torch.load(local_path, map_location='cpu')
first_stage_model_path = os.path.join(
os.path.dirname(local_path), 'first_stage_model.pth')
cond_stage_model_path = os.path.join(
@@ -311,7 +316,7 @@ class DiffusionInference():
module_paras = {}
if cfg is not None:
self.paras = cfg.PARAS
self.input = {k.lower(): v for k, v in cfg.INPUT.items()}
self.input = {k.lower(): dict(v).get('DEFAULT', None) if isinstance(v, (dict, OrderedDict)) else v for k, v in cfg.INPUT.items()}
self.output = {k.lower(): v for k, v in cfg.OUTPUT.items()}
module_paras = cfg.MODULES_PARAS
return module_paras
+216
View File
@@ -0,0 +1,216 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import math
import random
import torch
from scepter.modules.utils.distribute import we
from .control_inference import ControlInference
from .diffusion_inference import DiffusionInference, get_model
from .tuner_inference import TunerInference
from scepter.modules.model.registry import DIFFUSIONS, TOKENIZERS
class FluxInference(DiffusionInference):
def __init__(self, logger=None):
self.logger = logger
self.is_redefine_paras = False
self.loaded_model = {}
self.loaded_model_name = [
'diffusion_model', 'first_stage_model', 'cond_stage_model'
]
self.tuner_infer = TunerInference(self.logger)
self.control_infer = ControlInference(self.logger)
def init_from_cfg(self, cfg):
self.name = cfg.NAME
self.is_default = cfg.get('IS_DEFAULT', False)
module_paras = self.load_default(cfg.get('DEFAULT_PARAS', None))
assert cfg.have('MODEL')
if self.is_redefine_paras:
cfg.MODEL = self.redefine_paras(cfg.MODEL)
self.diffusion_model = self.infer_model(
cfg.MODEL.DIFFUSION_MODEL, module_paras.get(
'DIFFUSION_MODEL',
None)) if cfg.MODEL.have('DIFFUSION_MODEL') else None
self.first_stage_model = self.infer_model(
cfg.MODEL.FIRST_STAGE_MODEL,
module_paras.get(
'FIRST_STAGE_MODEL',
None)) if cfg.MODEL.have('FIRST_STAGE_MODEL') else None
self.cond_stage_model = self.infer_model(
cfg.MODEL.COND_STAGE_MODEL,
module_paras.get(
'COND_STAGE_MODEL',
None)) if cfg.MODEL.have('COND_STAGE_MODEL') else None
self.refiner_cond_model = self.infer_model(
cfg.MODEL.REFINER_COND_MODEL,
module_paras.get(
'REFINER_COND_MODEL',
None)) if cfg.MODEL.have('REFINER_COND_MODEL') else None
self.refiner_diffusion_model = self.infer_model(
cfg.MODEL.REFINER_MODEL, module_paras.get(
'REFINER_MODEL',
None)) if cfg.MODEL.have('REFINER_MODEL') else None
self.tokenizer = TOKENIZERS.build(
cfg.MODEL.TOKENIZER,
logger=self.logger) if cfg.MODEL.have('TOKENIZER') else None
if self.tokenizer is not None:
self.cond_stage_model['cfg'].KWARGS = {
'vocab_size': self.tokenizer.vocab_size
}
self.diffusion = DIFFUSIONS.build(cfg.MODEL.DIFFUSION, logger=self.logger) \
if cfg.MODEL.have('DIFFUSION') else None
assert self.diffusion is not None
@torch.no_grad()
def encode_first_stage(self, x, **kwargs):
_, dtype = self.get_function_info(self.first_stage_model, 'encode')
with torch.autocast('cuda',
enabled= dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, dtype)):
z = get_model(self.first_stage_model).encode(x)
if isinstance(z, (tuple, list)):
z = z[0]
return z
@torch.no_grad()
def decode_first_stage(self, z):
_, dtype = self.get_function_info(self.first_stage_model, 'decode')
with torch.autocast('cuda',
enabled=dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, dtype)):
return get_model(self.first_stage_model).decode(z)
@torch.no_grad()
def __call__(self,
input,
num_samples=1,
cat_uc=True,
tuner_model=None,
control_model=None,
**kwargs):
value_input = copy.deepcopy(self.input)
value_input.update(input)
print(value_input)
height, width = value_input['target_size_as_tuple']
value_output = copy.deepcopy(self.output)
# register tuner
if tuner_model is not None and tuner_model != '' and len(
tuner_model) > 0:
if not isinstance(tuner_model, list):
tuner_model = [tuner_model]
self.dynamic_load(self.diffusion_model, 'diffusion_model')
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
self.tuner_infer.register_tuner(tuner_model, self.diffusion_model,
self.cond_stage_model)
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=True)
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=True)
# cond stage
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
function_name, dtype = self.get_function_info(self.cond_stage_model)
with torch.autocast('cuda',
enabled=dtype == 'float16',
dtype=getattr(torch, dtype)):
ctx = getattr(get_model(self.cond_stage_model),
function_name)(value_input['prompt'])
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=True)
# get noise
seed = kwargs.pop('seed', -1)
g = torch.Generator(device=we.device_id)
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
g.manual_seed(seed)
if 'seed' in value_output:
value_output['seed'] = seed
for sample_id in range(num_samples):
if self.diffusion_model is not None:
noise = torch.randn(
num_samples,
16,
# allow for packing
2 * math.ceil(height / 16),
2 * math.ceil(width / 16),
device=we.device_id,
dtype=getattr(torch, dtype),
generator=torch.Generator(device=we.device_id).manual_seed(seed),
)
self.dynamic_load(self.diffusion_model, 'diffusion_model')
# UNet use input n_prompt
function_name, dtype = self.get_function_info(
self.diffusion_model)
with torch.autocast('cuda',
enabled= dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, dtype)):
solver_sample = value_input.get('sample', 'flow_eluer')
sample_steps = value_input.get('sample_steps', 20)
guide_scale = value_input.get('guide_scale', 3.5)
if guide_scale is not None:
guide_scale = torch.full((noise.shape[0],), guide_scale, device=noise.device,
dtype=noise.dtype)
else:
guide_scale = None
latent = self.diffusion.sample(
noise=noise,
sampler=solver_sample,
model=get_model(self.diffusion_model),
model_kwargs={"cond": ctx, "guidance": guide_scale},
steps=sample_steps,
show_progress=True,
guide_scale=guide_scale,
return_intermediate=None,
**kwargs)
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=True)
if 'latent' in value_output:
if value_output['latent'] is None or (
isinstance(value_output['latent'], list)
and len(value_output['latent']) < 1):
value_output['latent'] = []
value_output['latent'].append(latent)
self.dynamic_load(self.first_stage_model, 'first_stage_model')
x_samples = self.decode_first_stage(latent).float()
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=True)
images = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0)
if 'images' in value_output:
if value_output['images'] is None or (
isinstance(value_output['images'], list)
and len(value_output['images']) < 1):
value_output['images'] = []
value_output['images'].append(images)
for k, v in value_output.items():
if isinstance(v, list):
value_output[k] = torch.cat(v, dim=0)
if isinstance(v, torch.Tensor):
value_output[k] = v.cpu()
# unregister tuner
if tuner_model is not None and tuner_model != '' and len(
tuner_model) > 0:
self.tuner_infer.unregister_tuner(tuner_model,
self.diffusion_model,
self.cond_stage_model)
# unregister control
if control_model is not None and control_model != '':
self.control_infer.unregister_controllers(control_model,
self.diffusion_model)
return value_output
@@ -40,8 +40,10 @@ class LargenInference(DiffusionInference):
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(local_path)
else:
sd = torch.load(local_path, map_location='cpu')
if 'weights_only' in torch.load.__code__.co_varnames:
sd = torch.load(local_path, map_location='cpu', weights_only=True)
else:
sd = torch.load(local_path, map_location='cpu')
if 'model' in sd:
sd = sd['model']
@@ -1,20 +1,12 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import os.path
import random
from collections import OrderedDict
import torch
import torch.nn.functional as F
from PIL.Image import Image
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
from scepter.modules.model.network.diffusion.schedules import noise_schedule
from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, MODELS,
TOKENIZERS)
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from scepter.studio.utils.env import get_available_memory
from .control_inference import ControlInference
from .diffusion_inference import DiffusionInference, get_model
@@ -84,9 +76,11 @@ class PixArtInference(DiffusionInference):
function_name)(value_input['prompt'],
return_mask=True)
context['crossattn'] = cont.float()
self.dynamic_load(self.diffusion_model, 'diffusion_model')
null_context['crossattn'] = get_model(
self.diffusion_model).y_embedder.y_embedding[None].repeat(
num_samples, 1, 1)
self.dynamic_unload(self.diffusion_model, 'diffusion_model')
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=True)
+1 -3
View File
@@ -3,10 +3,8 @@
import copy
import random
import gradio as gr
import torch
import torch.nn.functional as F
import torchvision.transforms.functional as TF
from scepter.modules.model.network.diffusion.diffusion import \
GaussianDiffusionRF
from scepter.modules.utils.distribute import we
+14 -5
View File
@@ -13,10 +13,6 @@ try:
from peft.utils import CONFIG_NAME, SAFETENSORS_WEIGHTS_NAME, WEIGHTS_NAME
except Exception as e:
warnings.warn(f'Import peft error, please deal with this problem: {e}')
try:
from swift import Swift, SwiftModel
except Exception as e:
warnings.warn(f'Import swift error, please deal with this problem: {e}')
class TunerInference():
@@ -27,6 +23,11 @@ class TunerInference():
# @classmethod
def unregister_tuner(self, tuner_model_list, diffusion_model,
cond_stage_model):
try:
from swift import SwiftModel
except Exception as e:
warnings.warn(f'Import swift error, please deal with this problem: {e}')
self.logger.info('Unloading tuner model')
if isinstance(diffusion_model['model'], SwiftModel):
for adapter_name in diffusion_model['model'].adapters:
@@ -41,6 +42,11 @@ class TunerInference():
# @classmethod
def register_tuner(self, tuner_model_list, diffusion_model,
cond_stage_model):
try:
from swift import Swift
except Exception as e:
warnings.warn(f'Import swift error, please deal with this problem: {e}')
self.logger.info('Loading tuner model')
if len(tuner_model_list) < 1:
self.unregister_tuner(tuner_model_list, diffusion_model,
@@ -137,7 +143,10 @@ class TunerInference():
state_dict = {}
is_bin_file = True
if os.path.isfile(bin_file):
state_dict = torch.load(bin_file)
if 'weights_only' in torch.load.__code__.co_varnames:
state_dict = torch.load(bin_file, weights_only=True)
else:
state_dict = torch.load(bin_file)
elif os.path.isfile(safe_file):
is_bin_file = False
from safetensors.torch import \
+1 -1
View File
@@ -2,4 +2,4 @@
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model import (backbone, embedder, head, loss, metric,
neck, network, tokenizer, tuner)
neck, network, tokenizer, tuner, diffusion)
+2 -2
View File
@@ -1,4 +1,4 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.backbone import (autoencoder, image, mmdit, pixart,
unet, utils, video)
from scepter.modules.model.backbone import (ace, autoencoder, flux, image,
mmdit, pixart, unet, utils, video)
@@ -0,0 +1,3 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from .ace import ACE
+372
View File
@@ -0,0 +1,372 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import re
from collections import OrderedDict
from functools import partial
import torch
import torch.nn as nn
from einops import rearrange
from torch.nn.utils.rnn import pad_sequence
from torch.utils.checkpoint import checkpoint_sequential
from scepter.modules.model.backbone.transformer.layers import (Mlp,
T2IFinalLayer,
TimestepEmbedder
)
from scepter.modules.model.backbone.transformer.patchify import PatchEmbed
from scepter.modules.model.backbone.transformer.pos_embed import rope_params
from scepter.modules.model.base_model import BaseModel
from scepter.modules.model.registry import BACKBONES
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.file_system import FS
from .layers import ACEBlock
@BACKBONES.register_class()
class ACE(BaseModel):
para_dict = {
'PATCH_SIZE': {
'value': 2,
'description': ''
},
'IN_CHANNELS': {
'value': 4,
'description': ''
},
'HIDDEN_SIZE': {
'value': 1152,
'description': ''
},
'DEPTH': {
'value': 28,
'description': ''
},
'NUM_HEADS': {
'value': 16,
'description': ''
},
'MLP_RATIO': {
'value': 4.0,
'description': ''
},
'PRED_SIGMA': {
'value': True,
'description': ''
},
'DROP_PATH': {
'value': 0.,
'description': ''
},
'WINDOW_SIZE': {
'value': 0,
'description': ''
},
'WINDOW_BLOCK_INDEXES': {
'value': None,
'description': ''
},
'Y_CHANNELS': {
'value': 4096,
'description': ''
},
'ATTENTION_BACKEND': {
'value': None,
'description': ''
},
'QK_NORM': {
'value': True,
'description': 'Whether to use RMSNorm for query and key.',
},
}
para_dict.update(BaseModel.para_dict)
def __init__(self, cfg, logger):
super().__init__(cfg, logger=logger)
self.window_block_indexes = cfg.get('WINDOW_BLOCK_INDEXES', None)
if self.window_block_indexes is None:
self.window_block_indexes = []
self.pred_sigma = cfg.get('PRED_SIGMA', True)
self.in_channels = cfg.get('IN_CHANNELS', 4)
self.out_channels = self.in_channels * 2 if self.pred_sigma else self.in_channels
self.patch_size = cfg.get('PATCH_SIZE', 2)
self.num_heads = cfg.get('NUM_HEADS', 16)
self.hidden_size = cfg.get('HIDDEN_SIZE', 1152)
self.y_channels = cfg.get('Y_CHANNELS', 4096)
self.drop_path = cfg.get('DROP_PATH', 0.)
self.depth = cfg.get('DEPTH', 28)
self.mlp_ratio = cfg.get('MLP_RATIO', 4.0)
self.use_grad_checkpoint = cfg.get('USE_GRAD_CHECKPOINT', False)
self.attention_backend = cfg.get('ATTENTION_BACKEND', None)
self.max_seq_len = cfg.get('MAX_SEQ_LEN', 1024)
self.qk_norm = cfg.get('QK_NORM', False)
self.ignore_keys = cfg.get('IGNORE_KEYS', [])
assert (self.hidden_size % self.num_heads
) == 0 and (self.hidden_size // self.num_heads) % 2 == 0
d = self.hidden_size // self.num_heads
self.freqs = torch.cat(
[
rope_params(self.max_seq_len, d - 4 * (d // 6)), # T (~1/3)
rope_params(self.max_seq_len, 2 * (d // 6)), # H (~1/3)
rope_params(self.max_seq_len, 2 * (d // 6)) # W (~1/3)
],
dim=1)
# init embedder
self.x_embedder = PatchEmbed(self.patch_size,
self.in_channels + 1,
self.hidden_size,
bias=True,
flatten=False)
self.t_embedder = TimestepEmbedder(self.hidden_size)
self.y_embedder = Mlp(in_features=self.y_channels,
hidden_features=self.hidden_size,
out_features=self.hidden_size,
act_layer=lambda: nn.GELU(approximate='tanh'),
drop=0)
self.t_block = nn.Sequential(
nn.SiLU(),
nn.Linear(self.hidden_size, 6 * self.hidden_size, bias=True))
# init blocks
drop_path = [
x.item() for x in torch.linspace(0, self.drop_path, self.depth)
]
self.blocks = nn.ModuleList([
ACEBlock(self.hidden_size,
self.num_heads,
mlp_ratio=self.mlp_ratio,
drop_path=drop_path[i],
window_size=self.window_size
if i in self.window_block_indexes else 0,
backend=self.attention_backend,
use_condition=True,
qk_norm=self.qk_norm) for i in range(self.depth)
])
self.final_layer = T2IFinalLayer(self.hidden_size, self.patch_size,
self.out_channels)
self.initialize_weights()
def load_pretrained_model(self, pretrained_model):
if pretrained_model:
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
model = torch.load(local_path, map_location='cpu')
if 'state_dict' in model:
model = model['state_dict']
new_ckpt = OrderedDict()
for k, v in model.items():
if self.ignore_keys is not None:
if (isinstance(self.ignore_keys, str) and re.match(self.ignore_keys, k)) or \
(isinstance(self.ignore_keys, list) and k in self.ignore_keys):
continue
k = k.replace('.cross_attn.q_linear.', '.cross_attn.q.')
k = k.replace('.cross_attn.proj.',
'.cross_attn.o.').replace(
'.attn.proj.', '.attn.o.')
if '.cross_attn.kv_linear.' in k:
k_p, v_p = torch.split(v, v.shape[0] // 2)
new_ckpt[k.replace('.cross_attn.kv_linear.',
'.cross_attn.k.')] = k_p
new_ckpt[k.replace('.cross_attn.kv_linear.',
'.cross_attn.v.')] = v_p
elif '.attn.qkv.' in k:
q_p, k_p, v_p = torch.split(v, v.shape[0] // 3)
new_ckpt[k.replace('.attn.qkv.', '.attn.q.')] = q_p
new_ckpt[k.replace('.attn.qkv.', '.attn.k.')] = k_p
new_ckpt[k.replace('.attn.qkv.', '.attn.v.')] = v_p
elif 'y_embedder.y_proj.' in k:
new_ckpt[k.replace('y_embedder.y_proj.',
'y_embedder.')] = v
elif k in ('x_embedder.proj.weight'):
model_p = self.state_dict()[k]
if v.shape != model_p.shape:
model_p.zero_()
model_p[:, :4, :, :].copy_(v)
new_ckpt[k] = torch.nn.parameter.Parameter(model_p)
else:
new_ckpt[k] = v
elif k in ('x_embedder.proj.bias'):
new_ckpt[k] = v
else:
new_ckpt[k] = v
missing, unexpected = self.load_state_dict(new_ckpt,
strict=False)
print(
f'Restored from {pretrained_model} with {len(missing)} missing and {len(unexpected)} unexpected keys'
)
if len(missing) > 0:
print(f'Missing Keys:\n {missing}')
if len(unexpected) > 0:
print(f'\nUnexpected Keys:\n {unexpected}')
def forward(self,
x,
t=None,
cond=dict(),
mask=None,
text_position_embeddings=None,
gc_seg=-1,
**kwargs):
if self.freqs.device != x.device:
self.freqs = self.freqs.to(x.device)
if isinstance(cond, dict):
context = cond.get('crossattn', None)
else:
context = cond
if text_position_embeddings is not None:
# default use the text_position_embeddings in state_dict
# if state_dict doesn't including this key, use the arg: text_position_embeddings
proj_position_embeddings = self.y_embedder(
text_position_embeddings)
else:
proj_position_embeddings = None
ctx_batch, txt_lens = [], []
if mask is not None and isinstance(mask, list):
for ctx, ctx_mask in zip(context, mask):
for frame_id, one_ctx in enumerate(zip(ctx, ctx_mask)):
u, m = one_ctx
t_len = m.flatten().sum() # l
u = u[:t_len]
u = self.y_embedder(u)
if frame_id == 0:
u = u + proj_position_embeddings[
len(ctx) -
1] if proj_position_embeddings is not None else u
else:
u = u + proj_position_embeddings[
frame_id -
1] if proj_position_embeddings is not None else u
ctx_batch.append(u)
txt_lens.append(t_len)
else:
raise TypeError
y = torch.cat(ctx_batch, dim=0)
txt_lens = torch.LongTensor(txt_lens).to(x.device, non_blocking=True)
batch_frames = []
for u, shape, m in zip(x, cond['x_shapes'], cond['x_mask']):
u = u[:, :shape[0] * shape[1]].view(-1, shape[0], shape[1])
m = torch.ones_like(u[[0], :, :]) if m is None else m.squeeze(0)
batch_frames.append([torch.cat([u, m], dim=0).unsqueeze(0)])
if 'edit' in cond:
for i, (edit, edit_mask) in enumerate(
zip(cond['edit'], cond['edit_mask'])):
if edit is None:
continue
for u, m in zip(edit, edit_mask):
u = u.squeeze(0)
m = torch.ones_like(
u[[0], :, :]) if m is None else m.squeeze(0)
batch_frames[i].append(
torch.cat([u, m], dim=0).unsqueeze(0))
patch_batch, shape_batch, self_x_len, cross_x_len = [], [], [], []
for frames in batch_frames:
patches, patch_shapes = [], []
self_x_len.append(0)
for frame_id, u in enumerate(frames):
u = self.x_embedder(u)
h, w = u.size(2), u.size(3)
u = rearrange(u, '1 c h w -> (h w) c')
if frame_id == 0:
u = u + proj_position_embeddings[
len(frames) -
1] if proj_position_embeddings is not None else u
else:
u = u + proj_position_embeddings[
frame_id -
1] if proj_position_embeddings is not None else u
patches.append(u)
patch_shapes.append([h, w])
cross_x_len.append(h * w) # b*s, 1
self_x_len[-1] += h * w # b, 1
# u = torch.cat(patches, dim=0)
patch_batch.extend(patches)
shape_batch.append(
torch.LongTensor(patch_shapes).to(x.device, non_blocking=True))
# repeat t to align with x
t = torch.cat([t[i].repeat(l) for i, l in enumerate(self_x_len)])
self_x_len, cross_x_len = (torch.LongTensor(self_x_len).to(
x.device, non_blocking=True), torch.LongTensor(cross_x_len).to(
x.device, non_blocking=True))
# x = pad_sequence(tuple(patch_batch), batch_first=True) # b, s*max(cl), c
x = torch.cat(patch_batch, dim=0)
x_shapes = pad_sequence(tuple(shape_batch),
batch_first=True) # b, max(len(frames)), 2
t = self.t_embedder(t) # (N, D)
t0 = self.t_block(t)
# y = self.y_embedder(context)
kwargs = dict(y=y,
t=t0,
x_shapes=x_shapes,
self_x_len=self_x_len,
cross_x_len=cross_x_len,
freqs=self.freqs,
txt_lens=txt_lens)
if self.use_grad_checkpoint and gc_seg >= 0:
x = checkpoint_sequential(
functions=[partial(block, **kwargs) for block in self.blocks],
segments=gc_seg if gc_seg > 0 else len(self.blocks),
input=x,
use_reentrant=False)
else:
for block in self.blocks:
x = block(x, **kwargs)
x = self.final_layer(x, t) # b*s*n, d
outs, cur_length = [], 0
p = self.patch_size
for seq_length, shape in zip(self_x_len, shape_batch):
x_i = x[cur_length:cur_length + seq_length]
h, w = shape[0].tolist()
u = x_i[:h * w].view(h, w, p, p, -1)
u = rearrange(u, 'h w p q c -> (h p w q) c'
) # dump into sequence for following tensor ops
cur_length = cur_length + seq_length
outs.append(u)
x = pad_sequence(tuple(outs), batch_first=True).permute(0, 2, 1)
if self.pred_sigma:
return x.chunk(2, dim=1)[0]
else:
return x
def initialize_weights(self):
# Initialize transformer layers:
def _basic_init(module):
if isinstance(module, nn.Linear):
torch.nn.init.xavier_uniform_(module.weight)
if module.bias is not None:
nn.init.constant_(module.bias, 0)
self.apply(_basic_init)
# Initialize patch_embed like nn.Linear (instead of nn.Conv2d):
w = self.x_embedder.proj.weight.data
nn.init.xavier_uniform_(w.view([w.shape[0], -1]))
# Initialize timestep embedding MLP:
nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02)
nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02)
nn.init.normal_(self.t_block[1].weight, std=0.02)
# Initialize caption embedding MLP:
if hasattr(self, 'y_embedder'):
nn.init.normal_(self.y_embedder.fc1.weight, std=0.02)
nn.init.normal_(self.y_embedder.fc2.weight, std=0.02)
# Zero-out adaLN modulation layers
for block in self.blocks:
nn.init.constant_(block.cross_attn.o.weight, 0)
nn.init.constant_(block.cross_attn.o.bias, 0)
# Zero-out output layers:
nn.init.constant_(self.final_layer.linear.weight, 0)
nn.init.constant_(self.final_layer.linear.bias, 0)
@property
def dtype(self):
return next(self.parameters()).dtype
@staticmethod
def get_config_template():
return dict_to_yaml('BACKBONE',
__class__.__name__,
ACE.para_dict,
set_name=True)
@@ -0,0 +1,205 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import math
import warnings
import torch
import torch.nn as nn
from scepter.modules.model.backbone.transformer.attention import RMSNorm
from scepter.modules.model.backbone.transformer.layers import (DropPath, Mlp,
modulate)
from scepter.modules.model.backbone.transformer.pos_embed import \
rope_apply_multires as rope_apply
try:
from flash_attn import (flash_attn_varlen_func)
FLASHATTN_IS_AVAILABLE = True
except ImportError as e:
FLASHATTN_IS_AVAILABLE = False
flash_attn_varlen_func = None
warnings.warn(f'{e}')
class ACEBlock(nn.Module):
def __init__(self,
hidden_size,
num_heads,
mlp_ratio=4.0,
drop_path=0.,
window_size=0,
backend=None,
use_condition=True,
qk_norm=False,
**block_kwargs):
super().__init__()
self.hidden_size = hidden_size
self.use_condition = use_condition
self.norm1 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.attn = MultiHeadAttention(hidden_size,
num_heads=num_heads,
qkv_bias=True,
backend=backend,
qk_norm=qk_norm,
**block_kwargs)
if self.use_condition:
self.cross_attn = MultiHeadAttention(hidden_size,
context_dim=hidden_size,
num_heads=num_heads,
qkv_bias=True,
backend=backend,
qk_norm=qk_norm,
**block_kwargs)
self.norm2 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
# to be compatible with lower version pytorch
approx_gelu = lambda: nn.GELU(approximate='tanh')
self.mlp = Mlp(in_features=hidden_size,
hidden_features=int(hidden_size * mlp_ratio),
act_layer=approx_gelu,
drop=0)
self.drop_path = DropPath(
drop_path) if drop_path > 0. else nn.Identity()
self.window_size = window_size
self.scale_shift_table = nn.Parameter(
torch.randn(6, hidden_size) / hidden_size**0.5)
def forward(self, x, y, t, **kwargs):
B = x.size(0)
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
self.scale_shift_table[None] + t.reshape(B, 6, -1)).chunk(6, dim=1)
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
shift_msa.squeeze(1), scale_msa.squeeze(1), gate_msa.squeeze(1),
shift_mlp.squeeze(1), scale_mlp.squeeze(1), gate_mlp.squeeze(1))
x = x + self.drop_path(gate_msa * self.attn(
modulate(self.norm1(x), shift_msa, scale_msa, unsqueeze=False), **
kwargs))
if self.use_condition:
x = x + self.cross_attn(x, context=y, **kwargs)
x = x + self.drop_path(gate_mlp * self.mlp(
modulate(self.norm2(x), shift_mlp, scale_mlp, unsqueeze=False)))
return x
class MultiHeadAttention(nn.Module):
def __init__(self,
dim,
context_dim=None,
num_heads=None,
head_dim=None,
attn_drop=0.0,
qkv_bias=False,
dropout=0.0,
backend=None,
qk_norm=False,
eps=1e-6,
**block_kwargs):
super().__init__()
# consider head_dim first, then num_heads
num_heads = dim // head_dim if head_dim else num_heads
head_dim = dim // num_heads
assert num_heads * head_dim == dim
context_dim = context_dim or dim
self.dim = dim
self.context_dim = context_dim
self.num_heads = num_heads
self.head_dim = head_dim
self.scale = math.pow(head_dim, -0.25)
# layers
self.q = nn.Linear(dim, dim, bias=qkv_bias)
self.k = nn.Linear(context_dim, dim, bias=qkv_bias)
self.v = nn.Linear(context_dim, dim, bias=qkv_bias)
self.o = nn.Linear(dim, dim)
self.norm_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
self.dropout = nn.Dropout(dropout)
self.attention_op = None
self.attn_drop = nn.Dropout(attn_drop)
self.backend = backend
assert self.backend in ('flash_attn', 'xformer_attn', 'pytorch_attn',
None)
if FLASHATTN_IS_AVAILABLE and self.backend in ('flash_attn', None):
self.backend = 'flash_attn'
self.softmax_scale = block_kwargs.get('softmax_scale', None)
self.causal = block_kwargs.get('causal', False)
self.window_size = block_kwargs.get('window_size', (-1, -1))
self.deterministic = block_kwargs.get('deterministic', False)
else:
raise NotImplementedError
def flash_attn(self, x, context=None, **kwargs):
'''
The implementation will be very slow when mask is not None,
because we need rearange the x/context features according to mask.
Args:
x:
context:
mask:
**kwargs:
Returns: x
'''
dtype = kwargs.get('dtype', torch.float16)
def half(x):
return x if x.dtype in [torch.float16, torch.bfloat16
] else x.to(dtype)
x_shapes = kwargs['x_shapes']
freqs = kwargs['freqs']
self_x_len = kwargs['self_x_len']
cross_x_len = kwargs['cross_x_len']
txt_lens = kwargs['txt_lens']
n, d = self.num_heads, self.head_dim
if context is None:
# self-attn
q = self.norm_q(self.q(x)).view(-1, n, d)
k = self.norm_q(self.k(x)).view(-1, n, d)
v = self.v(x).view(-1, n, d)
q = rope_apply(q, self_x_len, x_shapes, freqs, pad=False)
k = rope_apply(k, self_x_len, x_shapes, freqs, pad=False)
q_lens = k_lens = self_x_len
else:
# cross-attn
q = self.norm_q(self.q(x)).view(-1, n, d)
k = self.norm_q(self.k(context)).view(-1, n, d)
v = self.v(context).view(-1, n, d)
q_lens = cross_x_len
k_lens = txt_lens
cu_seqlens_q = torch.cat([q_lens.new_zeros([1]),
q_lens]).cumsum(0, dtype=torch.int32)
cu_seqlens_k = torch.cat([k_lens.new_zeros([1]),
k_lens]).cumsum(0, dtype=torch.int32)
max_seqlen_q = q_lens.max()
max_seqlen_k = k_lens.max()
out_dtype = q.dtype
q, k, v = half(q), half(k), half(v)
x = flash_attn_varlen_func(q,
k,
v,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
max_seqlen_q=max_seqlen_q,
max_seqlen_k=max_seqlen_k,
dropout_p=self.attn_drop.p,
softmax_scale=self.softmax_scale,
causal=self.causal,
window_size=self.window_size,
deterministic=self.deterministic)
x = x.type(out_dtype)
x = x.reshape(-1, n * d)
x = self.o(x)
x = self.dropout(x)
return x
def forward(self, x, context=None, **kwargs):
x = getattr(self, self.backend)(x, context=context, **kwargs)
return x
@@ -0,0 +1,3 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from .flux import Flux
+251
View File
@@ -0,0 +1,251 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import math
from functools import partial
import torch
from einops import rearrange, repeat
from scepter.modules.model.base_model import BaseModel
from scepter.modules.model.registry import BACKBONES
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from torch import Tensor, nn
from torch.utils.checkpoint import checkpoint_sequential
from .layers import (DoubleStreamBlock, EmbedND, LastLayer, MLPEmbedder,
SingleStreamBlock, timestep_embedding)
@BACKBONES.register_class()
class Flux(BaseModel):
"""
Transformer backbone Diffusion model with RoPE.
"""
para_dict = {
'IN_CHANNELS': {
'value': 64,
'description': "model's input channels."
},
'OUT_CHANNELS': {
'value': 64,
'description': "model's output channels."
},
'HIDDEN_SIZE': {
'value': 1024,
'description': "model's hidden size."
},
'NUM_HEADS': {
'value': 16,
'description': 'number of heads in the transformer.'
},
'AXES_DIM': {
'value': [16, 56, 56],
'description': 'dimensions of the axes of the positional encoding.'
},
'THETA': {
'value': 10_000,
'description': 'theta for positional encoding.'
},
'VEC_IN_DIM': {
'value': 768,
'description': 'dimension of the vector input.'
},
'GUIDANCE_EMBED': {
'value': False,
'description': 'whether to use guidance embedding.'
},
'CONTEXT_IN_DIM': {
'value': 4096,
'description': 'dimension of the context input.'
},
'MLP_RATIO': {
'value': 4.0,
'description': 'ratio of mlp hidden size to hidden size.'
},
'QKV_BIAS': {
'value': True,
'description': 'whether to use bias in qkv projection.'
},
'DEPTH': {
'value': 19,
'description': 'number of transformer blocks.'
},
'DEPTH_SINGLE_BLOCKS': {
'value':
38,
'description':
'number of transformer blocks in the single stream block.'
},
'USE_GRAD_CHECKPOINT': {
'value': False,
'description': 'whether to use gradient checkpointing.'
}
}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.in_channels = cfg.IN_CHANNELS
self.out_channels = cfg.get('OUT_CHANNELS', self.in_channels)
hidden_size = cfg.get('HIDDEN_SIZE', 1024)
num_heads = cfg.get('NUM_HEADS', 16)
axes_dim = cfg.AXES_DIM
theta = cfg.THETA
vec_in_dim = cfg.VEC_IN_DIM
self.guidance_embed = cfg.GUIDANCE_EMBED
context_in_dim = cfg.CONTEXT_IN_DIM
mlp_ratio = cfg.MLP_RATIO
qkv_bias = cfg.QKV_BIAS
depth = cfg.DEPTH
depth_single_blocks = cfg.DEPTH_SINGLE_BLOCKS
self.use_grad_checkpoint = cfg.get('USE_GRAD_CHECKPOINT', False)
if hidden_size % num_heads != 0:
raise ValueError(
f"Hidden size {hidden_size} must be divisible by num_heads {num_heads}"
)
pe_dim = hidden_size // num_heads
if sum(axes_dim) != pe_dim:
raise ValueError(
f"Got {axes_dim} but expected positional dim {pe_dim}")
self.hidden_size = hidden_size
self.num_heads = num_heads
self.pe_embedder = EmbedND(dim=pe_dim, theta=theta, axes_dim=axes_dim)
self.img_in = nn.Linear(self.in_channels, self.hidden_size, bias=True)
self.time_in = MLPEmbedder(in_dim=256, hidden_dim=self.hidden_size)
self.vector_in = MLPEmbedder(vec_in_dim, self.hidden_size)
self.guidance_in = (MLPEmbedder(in_dim=256,
hidden_dim=self.hidden_size)
if self.guidance_embed else nn.Identity())
self.txt_in = nn.Linear(context_in_dim, self.hidden_size)
self.double_blocks = nn.ModuleList([
DoubleStreamBlock(
self.hidden_size,
self.num_heads,
mlp_ratio=mlp_ratio,
qkv_bias=qkv_bias,
) for _ in range(depth)
])
self.single_blocks = nn.ModuleList([
SingleStreamBlock(self.hidden_size,
self.num_heads,
mlp_ratio=mlp_ratio)
for _ in range(depth_single_blocks)
])
self.final_layer = LastLayer(self.hidden_size, 1, self.out_channels)
def prepare_input(self, x, context, y, x_shape=None):
# x.shape [6, 16, 16, 16] target is [6, 16, 768, 1360]
bs, c, h, w = x.shape
x = rearrange(x, 'b c (h ph) (w pw) -> b (h w) (c ph pw)', ph=2, pw=2)
x_id = torch.zeros(h // 2, w // 2, 3)
x_id[..., 1] = x_id[..., 1] + torch.arange(h // 2)[:, None]
x_id[..., 2] = x_id[..., 2] + torch.arange(w // 2)[None, :]
x_ids = repeat(x_id, 'h w c -> b (h w) c', b=bs)
txt_ids = torch.zeros(bs, context.shape[1], 3)
return x, x_ids.to(x), context.to(x), txt_ids.to(x), y.to(x), h, w
def unpack(self, x: Tensor, height: int, width: int) -> Tensor:
return rearrange(
x,
'b (h w) (c ph pw) -> b c (h ph) (w pw)',
h=math.ceil(height / 2),
w=math.ceil(width / 2),
ph=2,
pw=2,
)
def load_pretrained_model(self, pretrained_model):
if next(self.parameters()).device.type == 'meta':
map_location = we.device_id
else:
map_location = 'cpu'
if pretrained_model is not None:
with FS.get_from(pretrained_model,
wait_finish=True) as local_model:
if local_model.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(local_model, device=map_location)
else:
sd = torch.load(local_model, map_location=map_location)
missing, unexpected = self.load_state_dict(sd,
strict=False,
assign=True)
self.logger.info(
f'Restored from {pretrained_model} with {len(missing)} missing and {len(unexpected)} unexpected keys'
)
if len(missing) > 0:
self.logger.info(f'Missing Keys:\n {missing}') # noqa
if len(unexpected) > 0:
self.logger.info(f'\nUnexpected Keys:\n {unexpected}') # noqa
def forward(self,
x: Tensor,
t: Tensor,
cond: dict = {},
guidance: Tensor | None = None,
gc_seg: int = 0) -> Tensor:
x, x_ids, txt, txt_ids, y, h, w = self.prepare_input(
x, cond['context'], cond['y'])
# running on sequences img
x = self.img_in(x)
vec = self.time_in(timestep_embedding(t, 256))
if self.guidance_embed:
if guidance is None:
raise ValueError(
"Didn't get guidance strength for guidance distilled model."
)
vec = vec + self.guidance_in(timestep_embedding(guidance, 256))
vec = vec + self.vector_in(y)
txt = self.txt_in(txt)
ids = torch.cat((txt_ids, x_ids), dim=1)
pe = self.pe_embedder(ids)
kwargs = dict(
vec=vec,
pe=pe,
txt_length=txt.shape[1],
)
x = torch.cat((txt, x), 1)
if self.use_grad_checkpoint and gc_seg >= 0:
x = checkpoint_sequential(
functions=[
partial(block, **kwargs) for block in self.double_blocks
],
segments=gc_seg if gc_seg > 0 else len(self.double_blocks),
input=x,
use_reentrant=False)
else:
for block in self.double_blocks:
x = block(x, **kwargs)
kwargs = dict(
vec=vec,
pe=pe,
)
if self.use_grad_checkpoint and gc_seg >= 0:
x = checkpoint_sequential(
functions=[
partial(block, **kwargs) for block in self.single_blocks
],
segments=gc_seg if gc_seg > 0 else len(self.single_blocks),
input=x,
use_reentrant=False)
else:
for block in self.single_blocks:
x = block(x, **kwargs)
x = x[:, txt.shape[1]:, ...]
x = self.final_layer(
x, vec) # (N, T, patch_size ** 2 * out_channels) 6 64 64
x = self.unpack(x, h, w)
return x
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
Flux.para_dict,
set_name=True)
@@ -0,0 +1,362 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from __future__ import annotations
import math
from dataclasses import dataclass
import torch
from einops import rearrange, repeat
from torch import Tensor, nn
def attention(q: Tensor,
k: Tensor,
v: Tensor,
pe: Tensor,
mask: Tensor | None = None) -> Tensor:
q, k = apply_rope(q, k, pe)
x = torch.nn.functional.scaled_dot_product_attention(q,
k,
v,
attn_mask=mask)
x = torch.nan_to_num(x, nan=0.0, posinf=1e10, neginf=-1e10)
x = rearrange(x, 'B H L D -> B L (H D)')
return x
def rope(pos: Tensor, dim: int, theta: int) -> Tensor:
assert dim % 2 == 0
scale = torch.arange(0, dim, 2, dtype=torch.float64,
device=pos.device) / dim
omega = 1.0 / (theta**scale)
out = torch.einsum('...n,d->...nd', pos, omega)
out = torch.stack(
[torch.cos(out), -torch.sin(out),
torch.sin(out),
torch.cos(out)],
dim=-1)
out = rearrange(out, 'b n d (i j) -> b n d i j', i=2, j=2)
return out.float()
def apply_rope(xq: Tensor, xk: Tensor,
freqs_cis: Tensor) -> tuple[Tensor, Tensor]:
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)
xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]
xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]
return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(
*xk.shape).type_as(xk)
class EmbedND(nn.Module):
def __init__(self, dim: int, theta: int, axes_dim: list[int]):
super().__init__()
self.dim = dim
self.theta = theta
self.axes_dim = axes_dim
def forward(self, ids: Tensor) -> Tensor:
n_axes = ids.shape[-1]
emb = torch.cat(
[
rope(ids[..., i], self.axes_dim[i], self.theta)
for i in range(n_axes)
],
dim=-3,
)
return emb.unsqueeze(1)
def timestep_embedding(t: Tensor,
dim,
max_period=10000,
time_factor: float = 1000.0):
"""
Create sinusoidal timestep embeddings.
:param t: a 1-D Tensor of N indices, one per batch element.
These may be fractional.
:param dim: the dimension of the output.
:param max_period: controls the minimum frequency of the embeddings.
:return: an (N, D) Tensor of positional embeddings.
"""
t = time_factor * t
half = dim // 2
freqs = torch.exp(-math.log(max_period) *
torch.arange(start=0, end=half, dtype=torch.float32) /
half).to(t.device)
args = t[:, None].float() * freqs[None]
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
embedding = torch.cat(
[embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
if torch.is_floating_point(t):
embedding = embedding.to(t)
return embedding
class MLPEmbedder(nn.Module):
def __init__(self, in_dim: int, hidden_dim: int):
super().__init__()
self.in_layer = nn.Linear(in_dim, hidden_dim, bias=True)
self.silu = nn.SiLU()
self.out_layer = nn.Linear(hidden_dim, hidden_dim, bias=True)
def forward(self, x: Tensor) -> Tensor:
return self.out_layer(self.silu(self.in_layer(x)))
class RMSNorm(torch.nn.Module):
def __init__(self, dim: int):
super().__init__()
self.scale = nn.Parameter(torch.ones(dim))
def forward(self, x: Tensor):
x_dtype = x.dtype
x = x.float()
rrms = torch.rsqrt(torch.mean(x**2, dim=-1, keepdim=True) + 1e-6)
return (x * rrms).to(dtype=x_dtype) * self.scale
class QKNorm(torch.nn.Module):
def __init__(self, dim: int):
super().__init__()
self.query_norm = RMSNorm(dim)
self.key_norm = RMSNorm(dim)
def forward(self, q: Tensor, k: Tensor,
v: Tensor) -> tuple[Tensor, Tensor]:
q = self.query_norm(q)
k = self.key_norm(k)
return q.to(v), k.to(v)
class SelfAttention(nn.Module):
def __init__(self, dim: int, num_heads: int = 8, qkv_bias: bool = False):
super().__init__()
self.num_heads = num_heads
head_dim = dim // num_heads
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
self.norm = QKNorm(head_dim)
self.proj = nn.Linear(dim, dim)
def forward(self,
x: Tensor,
pe: Tensor,
mask: Tensor | None = None) -> Tensor:
qkv = self.qkv(x)
q, k, v = rearrange(qkv,
'B L (K H D) -> K B H L D',
K=3,
H=self.num_heads)
q, k = self.norm(q, k, v)
x = attention(q, k, v, pe=pe, mask=mask)
x = self.proj(x)
return x
@dataclass
class ModulationOut:
shift: Tensor
scale: Tensor
gate: Tensor
class Modulation(nn.Module):
def __init__(self, dim: int, double: bool):
super().__init__()
self.is_double = double
self.multiplier = 6 if double else 3
self.lin = nn.Linear(dim, self.multiplier * dim, bias=True)
def forward(self,
vec: Tensor) -> tuple[ModulationOut, ModulationOut | None]:
out = self.lin(nn.functional.silu(vec))[:,
None, :].chunk(self.multiplier,
dim=-1)
return (
ModulationOut(*out[:3]),
ModulationOut(*out[3:]) if self.is_double else None,
)
class DoubleStreamBlock(nn.Module):
def __init__(self,
hidden_size: int,
num_heads: int,
mlp_ratio: float,
qkv_bias: bool = False):
super().__init__()
mlp_hidden_dim = int(hidden_size * mlp_ratio)
self.num_heads = num_heads
self.hidden_size = hidden_size
self.img_mod = Modulation(hidden_size, double=True)
self.img_norm1 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.img_attn = SelfAttention(dim=hidden_size,
num_heads=num_heads,
qkv_bias=qkv_bias)
self.img_norm2 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.img_mlp = nn.Sequential(
nn.Linear(hidden_size, mlp_hidden_dim, bias=True),
nn.GELU(approximate='tanh'),
nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
)
self.txt_mod = Modulation(hidden_size, double=True)
self.txt_norm1 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.txt_attn = SelfAttention(dim=hidden_size,
num_heads=num_heads,
qkv_bias=qkv_bias)
self.txt_norm2 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.txt_mlp = nn.Sequential(
nn.Linear(hidden_size, mlp_hidden_dim, bias=True),
nn.GELU(approximate='tanh'),
nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
)
def forward(self,
x: Tensor,
vec: Tensor,
pe: Tensor,
mask: Tensor = None,
txt_length=None):
img_mod1, img_mod2 = self.img_mod(vec)
txt_mod1, txt_mod2 = self.txt_mod(vec)
txt, img = x[:, :txt_length], x[:, txt_length:]
# prepare image for attention
img_modulated = self.img_norm1(img)
img_modulated = (1 + img_mod1.scale) * img_modulated + img_mod1.shift
img_qkv = self.img_attn.qkv(img_modulated)
img_q, img_k, img_v = rearrange(img_qkv,
'B L (K H D) -> K B H L D',
K=3,
H=self.num_heads)
img_q, img_k = self.img_attn.norm(img_q, img_k, img_v)
# prepare txt for attention
txt_modulated = self.txt_norm1(txt)
txt_modulated = (1 + txt_mod1.scale) * txt_modulated + txt_mod1.shift
txt_qkv = self.txt_attn.qkv(txt_modulated)
txt_q, txt_k, txt_v = rearrange(txt_qkv,
'B L (K H D) -> K B H L D',
K=3,
H=self.num_heads)
txt_q, txt_k = self.txt_attn.norm(txt_q, txt_k, txt_v)
# run actual attention
q = torch.cat((txt_q, img_q), dim=2)
k = torch.cat((txt_k, img_k), dim=2)
v = torch.cat((txt_v, img_v), dim=2)
if mask is not None:
mask = repeat(mask, 'B L S-> B H L S', H=self.num_heads)
attn = attention(q, k, v, pe=pe, mask=mask)
txt_attn, img_attn = attn[:, :txt.shape[1]], attn[:, txt.shape[1]:]
# calculate the img bloks
img = img + img_mod1.gate * self.img_attn.proj(img_attn)
img = img + img_mod2.gate * self.img_mlp(
(1 + img_mod2.scale) * self.img_norm2(img) + img_mod2.shift)
# calculate the txt bloks
txt = txt + txt_mod1.gate * self.txt_attn.proj(txt_attn)
txt = txt + txt_mod2.gate * self.txt_mlp(
(1 + txt_mod2.scale) * self.txt_norm2(txt) + txt_mod2.shift)
x = torch.cat((txt, img), 1)
return x
class SingleStreamBlock(nn.Module):
"""
A DiT block with parallel linear layers as described in
https://arxiv.org/abs/2302.05442 and adapted modulation interface.
"""
def __init__(
self,
hidden_size: int,
num_heads: int,
mlp_ratio: float = 4.0,
qk_scale: float | None = None,
):
super().__init__()
self.hidden_dim = hidden_size
self.num_heads = num_heads
head_dim = hidden_size // num_heads
self.scale = qk_scale or head_dim**-0.5
self.mlp_hidden_dim = int(hidden_size * mlp_ratio)
# qkv and mlp_in
self.linear1 = nn.Linear(hidden_size,
hidden_size * 3 + self.mlp_hidden_dim)
# proj and mlp_out
self.linear2 = nn.Linear(hidden_size + self.mlp_hidden_dim,
hidden_size)
self.norm = QKNorm(head_dim)
self.hidden_size = hidden_size
self.pre_norm = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.mlp_act = nn.GELU(approximate='tanh')
self.modulation = Modulation(hidden_size, double=False)
def forward(self,
x: Tensor,
vec: Tensor,
pe: Tensor,
mask: Tensor = None) -> Tensor:
mod, _ = self.modulation(vec)
x_mod = (1 + mod.scale) * self.pre_norm(x) + mod.shift
qkv, mlp = torch.split(self.linear1(x_mod),
[3 * self.hidden_size, self.mlp_hidden_dim],
dim=-1)
q, k, v = rearrange(qkv,
'B L (K H D) -> K B H L D',
K=3,
H=self.num_heads)
q, k = self.norm(q, k, v)
if mask is not None:
mask = repeat(mask, 'B L S-> B H L S', H=self.num_heads)
# compute attention
attn = attention(q, k, v, pe=pe, mask=mask)
# compute activation in mlp stream, cat again and run second linear layer
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
return x + mod.gate * output
class LastLayer(nn.Module):
def __init__(self, hidden_size: int, patch_size: int, out_channels: int):
super().__init__()
self.norm_final = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.linear = nn.Linear(hidden_size,
patch_size * patch_size * out_channels,
bias=True)
self.adaLN_modulation = nn.Sequential(
nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True))
def forward(self, x: Tensor, vec: Tensor) -> Tensor:
shift, scale = self.adaLN_modulation(vec).chunk(2, dim=1)
x = (1 + scale[:, None, :]) * self.norm_final(x) + shift[:, None, :]
x = self.linear(x)
return x
@@ -1,2 +1,3 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from .sd3 import MMDiT
+4 -10
View File
@@ -5,12 +5,11 @@
# diffusers: https://github.com/huggingface/diffusers
# ComfyUI: https://github.com/comfyanonymous/ComfyUI
import logging
import math
import re
from collections import OrderedDict
from functools import partial
from typing import Dict, Optional
from typing import Optional
import numpy as np
import torch
@@ -26,7 +25,7 @@ try:
import xformers
import xformers.ops
XFORMERS_IS_AVAILBLE = True
except:
except Exception:
XFORMERS_IS_AVAILBLE = False
BROKEN_XFORMERS = False
@@ -35,7 +34,7 @@ try:
# XFormers bug confirmed on all versions from 0.0.21 to 0.0.26 (q with bs bigger than 65535 gives CUDA error)
BROKEN_XFORMERS = x_vers.startswith(
'0.0.2') and not x_vers.startswith('0.0.20')
except:
except Exception:
pass
@@ -1145,7 +1144,7 @@ class MMDiT(BaseModel):
for k, v in model.items():
if self.ignore_keys is not None:
if (isinstance(self.ignore_keys, str) and re.match(self.ignore_keys, k)) or \
(isinstance(self.ignore_keys, list) and k in self.ignore_keys):
(isinstance(self.ignore_keys, list) and k in self.ignore_keys):
ignore_ckpt[k] = v
continue
k = k.replace('model.diffusion_model.', '')
@@ -1185,11 +1184,6 @@ class MMDiT(BaseModel):
spatial_pos_embed = spatial_pos_embed[:, top:top + h, left:left + w, :]
spatial_pos_embed = rearrange(spatial_pos_embed,
'1 h w c -> 1 (h w) c')
# print(spatial_pos_embed, top, left, h, w)
# # t = get_2d_sincos_pos_embed_torch(self.hidden_size, w, h, 7.875, 7.875, device=device) #matches exactly for 1024 res
# t = get_2d_sincos_pos_embed_torch(self.hidden_size, w, h, 7.5, 7.5, device=device) #scales better
# # print(t)
# return t
return spatial_pos_embed
def unpatchify(self, x, hw=None):
@@ -1,2 +1,3 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from .pixart_alpha import PixArt
@@ -60,8 +60,8 @@ class FinalLayer(nn.Module):
self.linear = nn.Linear(hidden_size,
patch_size * patch_size * out_channels,
bias=True)
self.adaLN_modulation = nn.Sequential(
nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True))
self.adaLN_modulation = nn.Sequential(nn.SiLU(),
nn.Linear(hidden_size, 2 * hidden_size, bias=True))
def forward(self, x, c):
shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
@@ -537,7 +537,7 @@ class FullAttention(nn.Module):
k_img, k_txt = self.k_img_norm(k_img).view(
b, -1, n * d), self.k_txt_norm(k_txt).view(b, -1, n * d)
### add position
# add position
q_img, k_img = apply_2d_rope(q_img, k_img, padded_pos_index, n, d)
# support varying length
@@ -152,7 +152,6 @@ class SizeEmbedder(TimestepEmbedder):
@property
def dtype(self):
# 返回模型参数的数据类型
return next(self.parameters()).dtype
@@ -301,3 +300,28 @@ class Mlp(nn.Module):
x = self.fc2(x)
x = self.drop(x)
return x
class T2IFinalLayer(nn.Module):
"""
The final layer of PixArt.
"""
def __init__(self, hidden_size, patch_size, out_channels):
super().__init__()
self.norm_final = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.linear = nn.Linear(hidden_size,
patch_size * patch_size * out_channels,
bias=True)
self.scale_shift_table = nn.Parameter(
torch.randn(2, hidden_size) / hidden_size**0.5)
self.out_channels = out_channels
def forward(self, x, t):
shift, scale = (self.scale_shift_table[None] + t[:, None]).chunk(2,
dim=1)
shift, scale = shift.squeeze(1), scale.squeeze(1)
x = modulate(self.norm_final(x), shift, scale)
x = self.linear(x)
return x
@@ -9,6 +9,12 @@ from typing import Iterable
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from torch import Tensor
from torch.cuda import amp
from torch.nn.utils.rnn import pad_sequence
def _ntuple(n):
@@ -117,7 +123,6 @@ def apply_2d_rope(xq,
# xq_.shape = [b, seq_len, dim // 2, 2]
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 2)
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 2)
# 转为复数域
xq_ = torch.view_as_complex(
xq_) # [b, seq_len, dim // 2, 2]=>xq.shape = [b, seq_len, dim]
xk_ = torch.view_as_complex(xk_)
@@ -127,3 +132,218 @@ def apply_2d_rope(xq,
2) # point_wise mul, then flatten eg[[1,2],[3,4],[5,6]]->[1,2,3,4,5,6]
xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(2)
return xq_out.type_as(xq), xk_out.type_as(xk)
def sinusoidal_embedding_1d(dim, position):
# preprocess
assert dim % 2 == 0
half = dim // 2
position = position.type(torch.float64)
# calculation
sinusoid = torch.outer(
position, torch.pow(10000, -torch.arange(half).to(position).div(half)))
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
return x.float()
def frame_pad(x, seq_len, shapes):
max_h, max_w = np.max(shapes, 0)
frames = []
cur_len = 0
for h, w in shapes:
frame_len = h * w
frames.append(
F.pad(
x[cur_len:cur_len + frame_len].view(h, w, -1),
(0, 0, 0, max_w - w, 0, max_h - h)) # .view(max_h * max_w, -1)
)
cur_len += frame_len
if cur_len >= seq_len:
break
return torch.stack(frames)
def frame_unpad(x, shapes):
max_h, max_w = np.max(shapes, 0)
x = rearrange(x, '(b h w) n c -> b h w n c', h=max_h, w=max_w)
frames = []
for i, (h, w) in enumerate(shapes):
if i >= len(x):
break
frames.append(rearrange(x[i, :h, :w], 'h w n c -> (h w) n c'))
return torch.concat(frames)
@amp.autocast(enabled=False)
def rope_params(max_seq_len, dim, theta=10000):
"""
Precompute the frequency tensor for complex exponentials.
"""
assert dim % 2 == 0
freqs = torch.outer(
torch.arange(max_seq_len),
1.0 / torch.pow(theta,
torch.arange(0, dim, 2).to(torch.float64).div(dim)))
freqs = torch.polar(torch.ones_like(freqs), freqs)
return freqs
@amp.autocast(enabled=False)
def rope_apply(x, grid_sizes, freqs):
"""
x: [B, L, N, C].
grid_sizes: [B, 3].
freqs: [M, C // 2].
"""
n, c = x.size(2), x.size(3) // 2
# split freqs
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
# loop over samples
output = []
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
seq_len = f * h * w
# precompute multipliers
x_i = torch.view_as_complex(x[i, :seq_len].to(torch.float64).reshape(
seq_len, n, -1, 2))
freqs_i = torch.cat([
freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
],
dim=-1).reshape(seq_len, 1, -1)
# apply rotary embedding
x_i = torch.view_as_real(x_i * freqs_i).flatten(2).type_as(x)
x_i = torch.cat([x_i, x[i, seq_len:]])
# append to collection
output.append(x_i)
return torch.stack(output)
@amp.autocast(enabled=False)
def rope_apply_multires_pad(x, x_lens, x_shapes, freqs, pad=True):
"""
x: [B, L, N, C].
x_lens: [B].
x_shapes: [B, F, 2].
freqs: [M, C // 2].
"""
n, c = x.size(2), x.size(3) // 2
# split freqs
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
# loop over samples
output = []
for i, (seq_len,
shapes) in enumerate(zip(x_lens.tolist(), x_shapes.tolist())):
x_i = frame_pad(x[i], seq_len, shapes) # f, h, w, c
f, h, w = x_i.shape[:3]
pad_seq_len = f * h * w
# precompute multipliers
x_i = torch.view_as_complex(
x_i.to(torch.float64).reshape(pad_seq_len, n, -1, 2))
freqs_i = torch.cat([
freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
],
dim=-1).reshape(pad_seq_len, 1, -1)
# apply rotary embedding
x_i = torch.view_as_real(x_i * freqs_i).flatten(2).type_as(x)
x_i = frame_unpad(x_i, shapes)
if pad:
x_i = torch.cat([x_i, x[i, seq_len:]])
# append to collection
output.append(x_i)
return torch.stack(output) if pad else torch.concat(output)
@amp.autocast(enabled=False)
def rope_apply_multires(x, x_lens, x_shapes, freqs, pad=True):
"""
x: [B*L, N, C].
x_lens: [B].
x_shapes: [B, F, 2].
freqs: [M, C // 2].
"""
n, c = x.size(1), x.size(2) // 2
# split freqs
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
# loop over samples
output = []
st = 0
for i, (seq_len,
shapes) in enumerate(zip(x_lens.tolist(), x_shapes.tolist())):
x_i = frame_pad(x[st:st + seq_len], seq_len, shapes) # f, h, w, c
f, h, w = x_i.shape[:3]
pad_seq_len = f * h * w
# precompute multipliers
x_i = torch.view_as_complex(
x_i.to(torch.float64).reshape(pad_seq_len, n, -1, 2))
freqs_i = torch.cat([
freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
],
dim=-1).reshape(pad_seq_len, 1, -1)
# apply rotary embedding
x_i = torch.view_as_real(x_i * freqs_i).flatten(2).type_as(x)
x_i = frame_unpad(x_i, shapes)
# append to collection
output.append(x_i)
st += seq_len
return pad_sequence(output) if pad else torch.concat(output)
def rope(pos: Tensor, dim: int, theta: int) -> Tensor:
assert dim % 2 == 0
scale = torch.arange(0, dim, 2, dtype=torch.float64,
device=pos.device) / dim
omega = 1.0 / (theta**scale)
out = torch.einsum('...n,d->...nd', pos, omega)
out = torch.stack(
[torch.cos(out), -torch.sin(out),
torch.sin(out),
torch.cos(out)],
dim=-1)
out = rearrange(out, 'b n d (i j) -> b n d i j', i=2, j=2)
return out.float()
def apply_rope(xq: Tensor, xk: Tensor,
freqs_cis: Tensor) -> tuple[Tensor, Tensor]:
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)
xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]
xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]
return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(
*xk.shape).type_as(xk)
class EmbedND(nn.Module):
def __init__(self, dim: int, theta: int, axes_dim: list[int]):
super().__init__()
self.dim = dim
self.theta = theta
self.axes_dim = axes_dim
def forward(self, ids: Tensor) -> Tensor:
n_axes = ids.shape[-1]
emb = torch.cat(
[
rope(ids[..., i], self.axes_dim[i], self.theta)
for i in range(n_axes)
],
dim=-3,
)
return emb.unsqueeze(1)
+7 -1
View File
@@ -43,7 +43,13 @@ class BaseModel(nn.Module):
self._dist_data[key][k] += v
else:
self._dist_data[key][k] = v
def collect_probe(self):
probe_data_dict = self._probe_data
for k, v in self._modules.items():
if isinstance(getattr(self, k), BaseModel):
for kk, vv in getattr(self, k).collect_probe().items():
probe_data_dict[f'{k}/{kk}'] = vv
return probe_data_dict
def probe_data(self):
gather_probe_data = gather_data(self._probe_data)
_dist_data_list = gather_data([self._dist_data])
@@ -0,0 +1,7 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from .diffusions import BaseDiffusion, DiffusionFluxRF
from .samplers import BaseDiffusionSampler, DDIMSampler, FlowEluerSampler
from .schedules import (BaseNoiseScheduler, FlowMatchShiftScheduler,
ScaledLinearScheduler)
@@ -0,0 +1,316 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import math
import os
from collections import OrderedDict
import torch
from tqdm import trange
from scepter.modules.model.registry import (DIFFUSION_SAMPLERS, DIFFUSIONS,
NOISE_SCHEDULERS)
from scepter.modules.utils.config import Config, dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
@DIFFUSIONS.register_class()
class BaseDiffusion(object):
para_dict = {
'NOISE_SCHEDULER': {},
'SAMPLER_SCHEDULER': {},
'MIN_SNR_GAMMA': {
'value': None,
'description': 'The minimum SNR gamma value for the loss function.'
},
'PREDICTION_TYPE': {
'value': 'eps',
'description':
'The type of prediction to use for the loss function.'
}
}
def __init__(self, cfg, logger=None):
super(BaseDiffusion, self).__init__()
self.logger = logger
self.cfg = cfg
self.init_params()
def init_params(self):
self.min_snr_gamma = self.cfg.get('MIN_SNR_GAMMA', None)
self.prediction_type = self.cfg.get('PREDICTION_TYPE', 'eps')
self.noise_scheduler = NOISE_SCHEDULERS.build(self.cfg.NOISE_SCHEDULER,
logger=self.logger)
self.sampler_scheduler = NOISE_SCHEDULERS.build(self.cfg.get(
'SAMPLER_SCHEDULER', self.cfg.NOISE_SCHEDULER),
logger=self.logger)
self.num_timesteps = self.noise_scheduler.num_timesteps
if self.cfg.have('WORK_DIR') and we.rank == 0:
schedule_visualization = os.path.join(self.cfg.WORK_DIR,
'noise_schedule.png')
with FS.put_to(schedule_visualization) as local_path:
self.noise_scheduler.plot_noise_sampling_map(local_path)
schedule_visualization = os.path.join(self.cfg.WORK_DIR,
'sampler_schedule.png')
with FS.put_to(schedule_visualization) as local_path:
self.sampler_scheduler.plot_noise_sampling_map(local_path)
def sample(self,
noise,
model,
model_kwargs={},
steps=20,
sampler=None,
use_dynamic_cfg=False,
guide_scale=None,
guide_rescale=None,
show_progress=False,
return_intermediate=None,
intermediate_callback=None,
**kwargs):
assert isinstance(steps, (int, torch.LongTensor))
assert return_intermediate in (None, 'x0', 'xt')
assert isinstance(sampler, (str, dict, Config))
intermediates = []
def callback_fn(x_t, t, sigma=None, alpha=None):
timestamp = t
t = t.repeat(len(x_t)).round().long().to(x_t.device)
sigma = sigma.repeat(len(x_t), *([1] * (len(sigma.shape) - 1)))
alpha = alpha.repeat(len(x_t), *([1] * (len(alpha.shape) - 1)))
if guide_scale is None or guide_scale == 1.0:
out = model(x=x_t, t=t, **model_kwargs)
else:
if use_dynamic_cfg:
guidance_scale = 1 + guide_scale * (
(1 - math.cos(math.pi * (
(steps - timestamp.item()) / steps)**5.0)) / 2)
else:
guidance_scale = guide_scale
y_out = model(x=x_t, t=t, **model_kwargs[0])
u_out = model(x=x_t, t=t, **model_kwargs[1])
out = u_out + guidance_scale * (y_out - u_out)
if guide_rescale is not None and guide_rescale > 0.0:
ratio = (
y_out.flatten(1).std(dim=1) /
(out.flatten(1).std(dim=1) + 1e-12)).view((-1, ) + (1, ) *
(y_out.ndim - 1))
out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0
if self.prediction_type == 'x0':
x0 = out
elif self.prediction_type == 'eps':
x0 = (x_t - sigma * out) / alpha
elif self.prediction_type == 'v':
x0 = alpha * x_t - sigma * out
else:
raise NotImplementedError(
f'prediction_type {self.prediction_type} not implemented')
# print("torch.sum(y_out):", torch.sum(y_out), "torch.sum(u_out):", torch.sum(u_out), "torch.sum(out):",
# torch.sum(out), "torch.sum(x0):", torch.sum(x0), "sigmas", sigma, "alphas", alpha)
return x0
sampler_ins = self.get_sampler(sampler)
# this is ignored for schnell
sampler_output = sampler_ins.preprare_sampler(
noise,
steps=steps,
prediction_type=self.prediction_type,
scheduler_ins=self.sampler_scheduler,
callback_fn=callback_fn)
for _ in trange(steps, disable=not show_progress):
trange.desc = sampler_output.msg
sampler_output = sampler_ins.step(sampler_output)
if return_intermediate == 'x_0':
intermediates.append(sampler_output.x_0)
elif return_intermediate == 'x_t':
intermediates.append(sampler_output.x_t)
if intermediate_callback is not None:
intermediate_callback(intermediates[-1])
return (sampler_output.x_0, intermediates
) if return_intermediate is not None else sampler_output.x_0
def loss(self,
x_0,
model,
model_kwargs={},
reduction='mean',
noise=None,
**kwargs):
# use noise scheduler to add noise
if noise is None:
noise = torch.randn_like(x_0)
schedule_output = self.noise_scheduler.add_noise(x_0, noise, **kwargs)
x_t, t, sigma, alpha = schedule_output.x_t, schedule_output.t, schedule_output.sigma, schedule_output.alpha
out = model(x=x_t, t=t, **model_kwargs)
# mse loss
target = {
'eps': noise,
'x0': x_0,
'v': alpha * noise - sigma * x_0
}[self.prediction_type]
loss = (out - target).pow(2)
if reduction == 'mean':
loss = loss.flatten(1).mean(dim=1)
if self.min_snr_gamma is not None:
alphas = self.noise_scheduler.alphas.to(x_0.device)[t]
sigmas = self.noise_scheduler.sigmas.pow(2).to(x_0.device)[t]
snrs = (alphas / sigmas).clamp(min=1e-20)
min_snrs = snrs.clamp(max=self.min_snr_gamma)
weights = min_snrs / snrs
else:
weights = 1
loss = loss * weights
return loss
def get_sampler(self, sampler):
if isinstance(sampler, str):
if sampler not in DIFFUSION_SAMPLERS.class_map:
if self.logger is not None:
self.logger.info(
f'{sampler} not in the defined samplers list {DIFFUSION_SAMPLERS.class_map.keys()}'
)
else:
print(
f'{sampler} not in the defined samplers list {DIFFUSION_SAMPLERS.class_map.keys()}'
)
return None
sampler_cfg = Config(cfg_dict={'NAME': sampler}, load=False)
sampler_ins = DIFFUSION_SAMPLERS.build(sampler_cfg,
logger=self.logger)
elif isinstance(sampler, (Config, dict, OrderedDict)):
if isinstance(sampler, (dict, OrderedDict)):
sampler = Config(
cfg_dict={k.upper(): v
for k, v in dict(sampler).items()},
load=False)
sampler_ins = DIFFUSION_SAMPLERS.build(sampler, logger=self.logger)
else:
raise NotImplementedError
return sampler_ins
def __repr__(self) -> str:
return f'{self.__class__.__name__}' + ' ' + super().__repr__()
@staticmethod
def get_config_template():
return dict_to_yaml('DIFFUSIONS',
__class__.__name__,
BaseDiffusion.para_dict,
set_name=True)
@DIFFUSIONS.register_class()
class DiffusionFluxRF(BaseDiffusion):
para_dict = {
'PREDICTION_TYPE': {
'value': 'raw',
'description':
'The type of prediction to use for the loss function.'
}
}
para_dict.update(BaseDiffusion.para_dict)
def __init__(self, cfg, logger=None):
super(DiffusionFluxRF, self).__init__(cfg, logger=logger)
self.prediction_type = self.cfg.get('PREDICTION_TYPE', 'raw')
def loss(self,
x_0,
model,
model_kwargs={},
reduction='mean',
noise=None,
**kwargs):
if noise is None:
noise = torch.randn_like(x_0)
schedule_output = self.noise_scheduler.add_noise(x_0, noise, **kwargs)
x_t, t, sigma = schedule_output.x_t, schedule_output.t, schedule_output.sigma
out = model(x=x_t, t=sigma, **model_kwargs)
# raw
if self.prediction_type == 'raw':
target = noise - x_0
out = out
elif self.prediction_type == 'sigma_scaled':
target = x_0
out = out * (-sigma) + x_t
else:
raise NotImplementedError
loss = (target - out)**2
if reduction == 'mean':
loss = loss.flatten(1).mean(dim=1)
if self.min_snr_gamma is not None:
alphas = self.noise_scheduler.alphas.to(x_0.device)[t]
sigmas = self.noise_scheduler.sigmas.pow(2).to(x_0.device)[t]
snrs = (alphas / sigmas).clamp(min=1e-20)
min_snrs = snrs.clamp(max=self.min_snr_gamma)
weights = min_snrs / snrs
else:
weights = 1
loss = loss * weights
return loss
@torch.no_grad()
def sample(self,
noise,
model,
model_kwargs={},
steps=20,
sampler=None,
show_progress=False,
return_intermediate=None,
intermediate_callback=None,
**kwargs):
# sanity check
assert isinstance(steps, (int, torch.LongTensor))
assert return_intermediate in (None, 'x0', 'xt')
assert isinstance(sampler, (str, dict, Config))
intermediates = []
def callback_fn(x_t, t, sigma=None, alpha=None):
sigma = torch.full((x_t.shape[0], ),
sigma,
dtype=x_t.dtype,
device=x_t.device)
x_0 = model(x=x_t, t=sigma, **model_kwargs)
return x_0
sampler_ins = self.get_sampler(sampler)
# this is ignored for schnell
sampler_output = sampler_ins.preprare_sampler(
noise,
steps=steps,
prediction_type=self.prediction_type,
scheduler_ins=self.sampler_scheduler,
callback_fn=callback_fn)
for _ in trange(steps, disable=not show_progress):
trange.desc = sampler_output.msg
sampler_output = sampler_ins.step(sampler_output)
if return_intermediate == 'x_0':
intermediates.append(sampler_output.x_0)
elif return_intermediate == 'x_t':
intermediates.append(sampler_output.x_t)
if intermediate_callback is not None:
intermediate_callback(intermediates[-1])
return (sampler_output.x_0, intermediates
) if return_intermediate is not None else sampler_output.x_t
@staticmethod
def get_config_template():
return dict_to_yaml('DIFFUSIONS',
__class__.__name__,
DiffusionFluxRF.para_dict,
set_name=True)
+254
View File
@@ -0,0 +1,254 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from dataclasses import dataclass
import torch
from scepter.modules.model.registry import DIFFUSION_SAMPLERS
from scepter.modules.utils.config import dict_to_yaml
from .util import _i
@dataclass
class SamplerOutput(object):
callback_fn: callable
prediction_type: str
alphas: torch.Tensor
betas: torch.Tensor
sigmas: torch.Tensor
alphas_init: torch.Tensor
betas_init: torch.Tensor
sigmas_init: torch.Tensor
ts: torch.Tensor
x_t: torch.Tensor
x_0: torch.Tensor
step: int
msg: str
def add_custom_field(self, key: str, value) -> None:
self.__setattr__(key, value)
@DIFFUSION_SAMPLERS.register_class('base')
class BaseDiffusionSampler(object):
para_dict = {}
def __init__(self, cfg, logger=None):
super(BaseDiffusionSampler, self).__init__()
self.logger = logger
self.cfg = cfg
self.init_params()
def init_params(self):
self.discretization_type = self.cfg.get('DISCRETIZATION_TYPE',
'linspace')
self.discard_penultimate_step = self.cfg.get(
'DISCARD_PENULTIMATE_STEP', False)
self.free_steps = self.cfg.get('FREE_STEPS', None)
self.t_max = self.cfg.get('T_MAX', None)
self.t_min = self.cfg.get('T_MIN', None)
def discretization(self, steps=20, num_timesteps=1000, **kwargs):
# get timesteps
if isinstance(steps, int):
steps += 1 if self.discard_penultimate_step else 0
t_max = num_timesteps - 1 if self.t_max is None else self.t_max
t_min = 0 if self.t_min is None else self.t_min
# discretize timesteps
if self.discretization_type == 'leading':
steps = torch.arange(t_min, t_max + 1,
(t_max - t_min + 1) / steps).flip(0)
elif self.discretization_type == 'linspace':
steps = torch.linspace(t_max, t_min, steps)
elif self.discretization_type == 'trailing':
steps = torch.arange(t_max, t_min - 1,
-((t_max - t_min + 1) / steps))
elif self.discretization_type == 'free':
steps = torch.tensor(self.free_steps)
else:
raise NotImplementedError(
f'{self.discretization_type} discretization not implemented'
)
steps = steps.clamp_(t_min, t_max)
elif isinstance(steps, list):
steps = torch.tensor(steps)
timesteps = torch.as_tensor(steps, dtype=torch.float32)
return timesteps
def preprare_sampler(self,
noise,
steps=20,
scheduler_ins=None,
prediction_type='',
sigmas=None,
betas=None,
alphas=None,
callback_fn=None,
**kwargs):
'''
1. Control the model's inputs and outputs externally in the solver by callback_fn,
and perform the conversion between x0 and xt internally within the solver.
2. The function callback_fn use the x0 and xt as the default inputs and also give me
the x0 and xt as output. The other inputs will be set in kwargs.
3. The basic parameters of the schedule should be set manually.
4. To ensure the safety of threading, use the instance of SamplerOutput as the manager,
which manage all necessary information.
'''
num_timesteps = scheduler_ins.num_timesteps if scheduler_ins is not None else 1000
timestamps = self.discretization(steps,
num_timesteps=num_timesteps,
**kwargs)
alphas = scheduler_ins.t_to_alpha(
timestamps, **kwargs) if scheduler_ins is not None else alphas
betas = scheduler_ins.t_to_beta(
timestamps, **kwargs) if scheduler_ins is not None else betas
sigmas = scheduler_ins.t_to_sigma(
timestamps, **kwargs) if scheduler_ins is not None else sigmas
alphas_init = scheduler_ins.t_to_alpha_init(
timestamps, **kwargs) if scheduler_ins is not None else alphas
betas_init = scheduler_ins.t_to_beta_init(
timestamps, **kwargs) if scheduler_ins is not None else betas
sigmas_init = scheduler_ins.t_to_sigma_init(
timestamps, **kwargs) if scheduler_ins is not None else sigmas
output = SamplerOutput(callback_fn=callback_fn,
prediction_type=prediction_type,
alphas=alphas,
betas=betas,
sigmas=sigmas,
alphas_init=alphas_init,
betas_init=betas_init,
sigmas_init=sigmas_init,
ts=timestamps,
x_t=noise,
x_0=noise,
step=0,
msg='step 0')
return output
def step(self, sampler_ouput):
raise NotImplementedError(
'DiffusionSampler step function not implemented')
def __repr__(self) -> str:
return f'{self.__class__.__name__}' + ' ' + super().__repr__()
@staticmethod
def get_config_template():
return dict_to_yaml('DIFFUSION_SAMPLERS',
__class__.__name__,
BaseDiffusionSampler.para_dict,
set_name=True)
@DIFFUSION_SAMPLERS.register_class('eluer')
class EulerSampler(BaseDiffusionSampler):
def step(self, sampler_ouput):
pass
@DIFFUSION_SAMPLERS.register_class('ddim')
class DDIMSampler(BaseDiffusionSampler):
def init_params(self):
super().init_params()
self.eta = self.cfg.get('ETA', 0.)
self.discretization_type = self.cfg.get('DISCRETIZATION_TYPE',
'trailing')
def preprare_sampler(self,
noise,
steps=20,
scheduler_ins=None,
prediction_type='',
sigmas=None,
betas=None,
alphas=None,
callback_fn=None,
**kwargs):
output = super().preprare_sampler(noise, steps, scheduler_ins,
prediction_type, sigmas, betas,
alphas, callback_fn, **kwargs)
sigmas = output.sigmas
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
sigmas_vp = (sigmas**2 / (1 + sigmas**2))**0.5
sigmas_vp[sigmas == float('inf')] = 1.
output.add_custom_field('sigmas_vp', sigmas_vp)
return output
def step(self, sampler_output):
x_t = sampler_output.x_t
step = sampler_output.step
t = sampler_output.ts[step]
sigmas_vp = sampler_output.sigmas_vp.to(x_t.device)
alpha_init = _i(sampler_output.alphas_init, step, x_t[:1])
sigma_init = _i(sampler_output.sigmas_init, step, x_t[:1])
x = sampler_output.callback_fn(x_t, t, sigma_init, alpha_init)
noise_factor = self.eta * (sigmas_vp[step + 1]**2 /
sigmas_vp[step]**2 *
(1 - (1 - sigmas_vp[step]**2) /
(1 - sigmas_vp[step + 1]**2)))
d = (x_t - (1 - sigmas_vp[step]**2)**0.5 * x) / sigmas_vp[step]
x = (1 - sigmas_vp[step + 1] ** 2) ** 0.5 * x + \
(sigmas_vp[step + 1] ** 2 - noise_factor ** 2) ** 0.5 * d
sampler_output.x_0 = x
if sigmas_vp[step + 1] > 0:
x += noise_factor * torch.randn_like(x)
sampler_output.x_t = x
sampler_output.step += 1
sampler_output.msg = f'step {step}'
return sampler_output
@DIFFUSION_SAMPLERS.register_class('flow_eluer')
class FlowEluerSampler(BaseDiffusionSampler):
def preprare_sampler(self,
noise,
steps=20,
scheduler_ins=None,
prediction_type='',
sigmas=None,
betas=None,
alphas=None,
callback_fn=None,
**kwargs):
if noise.ndim == 3:
seq_len = noise.shape[2] // 4
else:
n, _, h, w = noise.shape
seq_len = (h // 2 * w // 2)
kwargs['seq_len'] = seq_len
output = super().preprare_sampler(noise, steps, scheduler_ins,
prediction_type, sigmas, betas,
alphas, callback_fn, **kwargs)
return output
def step(self, sampler_output):
step = sampler_output.step
x_t = sampler_output.x_t
sigma_curr, sigma_prev = sampler_output.sigmas[
step], sampler_output.sigmas[step + 1]
prediction_type = sampler_output.prediction_type
assert prediction_type in ('raw', 'sigma_scaled')
t = sampler_output.ts[step]
x_0 = sampler_output.callback_fn(x_t, t, sigma_curr)
x_t = x_t + (sigma_prev - sigma_curr) * x_0
sampler_output.x_0 = x_0
sampler_output.x_t = x_t
sampler_output.step += 1
sampler_output.msg = f'step {step}, sigma_curr: {sigma_curr}, sigma_prev: {sigma_prev}'
return sampler_output
def discretization(self, steps=20, num_timesteps=1000, **kwargs):
# extra step for zero
timesteps = torch.linspace(num_timesteps, 0, steps + 1)
return timesteps
@staticmethod
def get_config_template():
return dict_to_yaml('DIFFUSION_SAMPLERS',
__class__.__name__,
FlowEluerSampler.para_dict,
set_name=True)
@@ -0,0 +1,635 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import math
from dataclasses import dataclass, field
from typing import Callable
import numpy as np
import torch
from torch import Tensor
from scepter.modules.model.registry import NOISE_SCHEDULERS
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.math_plot import plot_multi_curves
from .util import _i
@dataclass
class ScheduleOutput(object):
x_t: torch.Tensor
x_0: torch.Tensor
t: torch.Tensor
sigma: torch.Tensor
alpha: torch.Tensor
custom_fields: dict = field(default_factory=dict)
def add_custom_field(self, key: str, value) -> None:
self.__setattr__(key, value)
@NOISE_SCHEDULERS.register_class()
class BaseNoiseScheduler(object):
para_dict = {
'NUM_TIMESTEPS': {
'value': 1000,
'description': 'The number of timesteps for sampling.'
},
}
def __init__(self, cfg, logger=None):
super(BaseNoiseScheduler, self).__init__()
self.logger = logger
self.cfg = cfg
self.init_params()
self.get_schedule()
def init_params(self):
self.num_timesteps = self.cfg.get('NUM_TIMESTEPS', 1000)
self._sample_steps = torch.arange(self.num_timesteps,
dtype=torch.float32)
self._sigmas, self._betas, self._alphas, self._timesteps = None, None, None, None
def check_function(self):
try:
predict_timestamps = self.sigma_to_t(self.sigmas)
predict_sigmas = self.t_to_sigma(self._timesteps)
diff_sigmas = torch.sum(torch.abs(predict_sigmas - self.sigmas))
diff_timestamps = torch.sum(
torch.abs(predict_timestamps - self._timesteps))
if diff_sigmas > 1e-3 or diff_timestamps > 1:
self.logger.info(
f'The noise scheduler {self.__class__.__name__} is not correct, '
f'please check the function sigma_to_t or t_to_sigma.'
f'Info: diff sigmas {diff_sigmas}, diff timestamps {diff_timestamps}'
)
raise 'The noise scheduler checked failed.'
else:
self.logger.info(
f'The noise scheduler {self.__class__.__name__} is checked and passed.'
)
except Exception as e:
if isinstance(e, NotImplementedError):
self.logger.info(
'Not implemented function sigma_to_t or t_to_sigma, skip check.'
)
else:
self.logger.info(
f'The noise scheduler {self.__class__.__name__} is not correct, '
f'please check the function sigma_to_t or t_to_sigma. Error: {e}'
)
raise e
def get_schedule(self):
raise NotImplementedError(
'NoiseScheduler get_schedule function not implemented')
def square_betas_to_sigmas(self, square_betas):
return torch.sqrt(1 - torch.cumprod(1 - square_betas, dim=0))
def sigmas_to_square_betas(self, sigmas):
square_alphas = 1 - sigmas**2
betas = 1 - torch.cat(
[square_alphas[:1], square_alphas[1:] / square_alphas[:-1]])
return betas
def sigma_to_t(self, sigma, **kwargs):
if sigma == float('inf'):
t = torch.full_like(sigma, len(self._sigmas) - 1)
else:
log_sigmas = torch.sqrt(self._sigmas**2 /
(1 - self._sigmas**2)).log().to(sigma)
log_sigma = sigma.log()
dists = log_sigma - log_sigmas[:, None]
low_idx = dists.ge(0).cumsum(dim=0).argmax(dim=0).clamp(
max=log_sigmas.shape[0] - 2)
high_idx = low_idx + 1
low, high = log_sigmas[low_idx], log_sigmas[high_idx]
w = (low - log_sigma) / (low - high)
w = w.clamp(0, 1)
t = (1 - w) * low_idx + w * high_idx
t = t.view(sigma.shape)
if t.ndim == 0:
t = t.unsqueeze(0)
return t
def t_to_sigma(self, t, **kwargs):
t = t.float()
low_idx, high_idx, w = t.floor().long(), t.ceil().long(), t.frac()
log_sigmas = torch.sqrt(self.sigmas**2 /
(1 - self.sigmas**2)).log().to(t)
log_sigma = (1 - w) * log_sigmas[low_idx] + w * log_sigmas[high_idx]
log_sigma[torch.isnan(log_sigma)
| torch.isinf(log_sigma)] = float('inf')
return log_sigma.exp()
def t_to_alpha(self, t, **kwargs):
sigma = self.t_to_sigma(t)
square_beta = self.sigmas_to_square_betas(sigma)
return torch.sqrt(1 - square_beta)
def t_to_beta(self, t, **kwargs):
sigma = self.t_to_sigma(t)
square_beta = self.sigmas_to_square_betas(sigma)
return torch.sqrt(square_beta)
def add_noise(self, x_0, noise=None, t=None, **kwargs):
if t is None:
t = torch.randint(0,
self.num_timesteps, (x_0.shape[0], ),
device=x_0.device).long()
alpha = _i(self.alphas, t, x_0)
sigma = _i(self.sigmas, t, x_0)
x_t = alpha * x_0 + sigma * noise
return ScheduleOutput(x_0=x_0, x_t=x_t, t=t, alpha=alpha, sigma=sigma)
def t_to_alpha_init(self, t, **kwargs):
indices = t.long()
indices[indices >= self.num_timesteps] = self.num_timesteps - 1
timesteps = self.timesteps.to(t)[indices]
step_indices = [(self.timesteps.to(t) == t).nonzero().item()
for t in timesteps]
alpha = self.alphas[step_indices].flatten().to(t)
return alpha
def t_to_beta_init(self, t, **kwargs):
indices = t.long()
indices[indices >= self.num_timesteps] = self.num_timesteps - 1
timesteps = self.timesteps.to(t)[indices]
step_indices = [(self.timesteps.to(t) == t).nonzero().item()
for t in timesteps]
beta = self.betas[step_indices].flatten().to(t)
return beta
def t_to_sigma_init(self, t, **kwargs):
indices = t.long()
indices[indices >= self.num_timesteps] = self.num_timesteps - 1
timesteps = self.timesteps.to(t)[indices]
step_indices = [(self.timesteps.to(t) == t).nonzero().item()
for t in timesteps]
sigma = self.sigmas[step_indices].flatten().to(t)
return sigma
def rescale_zero_terminal_snr(self, alphas_cumprod):
"""
Rescales betas to have zero terminal SNR Based on https://arxiv.org/pdf/2305.08891.pdf (Algorithm 1)
Args:
betas (`torch.Tensor`):
the betas that the scheduler is being initialized with.
Returns:
`torch.Tensor`: rescaled betas with zero terminal SNR
"""
alphas_bar_sqrt = alphas_cumprod.sqrt()
# Store old values.
alphas_bar_sqrt_0 = alphas_bar_sqrt[0].clone()
alphas_bar_sqrt_T = alphas_bar_sqrt[-1].clone()
# Shift so the last timestep is zero.
alphas_bar_sqrt -= alphas_bar_sqrt_T
# Scale so the first timestep is back to the old value.
alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 -
alphas_bar_sqrt_T)
# Convert alphas_bar_sqrt to betas
alphas_bar = alphas_bar_sqrt**2 # Revert sqrt
return alphas_bar
@property
def sigmas(self):
return self._sigmas
@property
def betas(self):
return self._betas
@property
def alphas(self):
return self._alphas
@property
def timesteps(self):
return self._timesteps
# plot the noise sampling map
def plot_noise_sampling_map(self, save_path):
y = [{
'data': self._sigmas.cpu().numpy(),
'label': 'sigmas'
}, {
'data': self._betas.cpu().numpy(),
'label': 'betas'
}, {
'data': self._alphas.cpu().numpy(),
'label': 'alphas'
}, {
'data': self._timesteps.cpu().numpy() / self.num_timesteps,
'label': 'timesteps'
}]
plot_multi_curves(
x=self._sample_steps.cpu().numpy(),
y=y,
x_label='timesteps',
y_label=None,
title=f"{self.__class__.__name__}'s noise sampling map",
save_path=save_path)
return save_path
def __repr__(self) -> str:
return f'{self.__class__.__name__}' + ' ' + super().__repr__()
@staticmethod
def get_config_template():
return dict_to_yaml('NOISE_SCHEDULER',
__class__.__name__,
BaseNoiseScheduler.para_dict,
set_name=True)
@NOISE_SCHEDULERS.register_class()
class ScaledLinearScheduler(BaseNoiseScheduler):
para_dict = {}
def init_params(self):
super().init_params()
self.beta_min = self.cfg.get('BETA_MIN', 0.00085)
self.beta_max = self.cfg.get('BETA_MAX', 0.012)
self.snr_shift_scale = self.cfg.get('SNR_SHIFT_SCALE', None)
self.rescale_betas_zero_snr = self.cfg.get('RESCALE_BETAS_ZERO_SNR',
False)
def square_betas_to_sigmas(self,
square_betas,
snr_shift_scale=None,
rescale_betas_zero_snr=False):
if snr_shift_scale is not None or rescale_betas_zero_snr:
alphas_cumprod = torch.cumprod(1 - square_betas, dim=0)
if snr_shift_scale is not None and snr_shift_scale > 0:
alphas_cumprod = alphas_cumprod / (
snr_shift_scale + (1 - snr_shift_scale) * alphas_cumprod)
if rescale_betas_zero_snr:
alphas_cumprod = self.rescale_zero_terminal_snr(alphas_cumprod)
return torch.sqrt(1 - alphas_cumprod)
else:
return torch.sqrt(1 - torch.cumprod(1 - square_betas, dim=0))
def get_schedule(self):
square_betas = torch.linspace(self.beta_min**0.5,
self.beta_max**0.5,
self.num_timesteps,
dtype=torch.float32)**2
self._sigmas = self.square_betas_to_sigmas(square_betas,
self.snr_shift_scale,
self.rescale_betas_zero_snr)
self._betas = torch.sqrt(square_betas)
self._alphas = torch.sqrt(1 - self._sigmas**2)
self._timesteps = torch.arange(len(self._sigmas), dtype=torch.float32)
@NOISE_SCHEDULERS.register_class()
class LinearScheduler(BaseNoiseScheduler):
para_dict = {}
def init_params(self):
super().init_params()
self.beta_min = self.cfg.get('BETA_MIN', 0.00085)
self.beta_max = self.cfg.get('BETA_MAX', 0.012)
def betas_to_sigmas(self, betas):
return torch.sqrt(1 - torch.cumprod(1 - betas, dim=0))
def get_schedule(self):
betas = torch.linspace(self.beta_min,
self.beta_max,
self.num_timesteps,
dtype=torch.float32)
sigmas = self.betas_to_sigmas(betas)
self._sigmas = sigmas
self._betas = betas
self._alphas = torch.sqrt(1 - sigmas**2)
self._timesteps = torch.arange(len(sigmas), dtype=torch.float32)
@NOISE_SCHEDULERS.register_class()
class FlowMatchUniformScheduler(BaseNoiseScheduler):
def get_schedule(self):
timesteps = np.linspace(1,
self.num_timesteps,
self.num_timesteps,
dtype=np.float32).copy()
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
self._timesteps = timesteps
self._sigmas = self.t_to_sigma(timesteps)
self._betas = torch.sqrt(self.sigmas_to_square_betas(self._sigmas))
self._alphas = torch.sqrt(1 - self.betas**2)
def add_noise(self, x_0, noise=None, t=None, **kwargs):
if t is None:
t = torch.rand(
(x_0.shape[0], ), device=x_0.device) * self.num_timesteps
sigma = self.t_to_sigma(t)
shape = (x_0.size(0), ) + (1, ) * (x_0.ndim - 1)
x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
return ScheduleOutput(x_0=x_0,
x_t=x_t,
t=t,
sigma=sigma,
alpha=self.t_to_alpha(t))
def sigma_to_t(self, sigma, **kwargs):
return sigma * self.num_timesteps
def t_to_sigma(self, t, **kwargs):
return t / self.num_timesteps
@staticmethod
def get_config_template():
return dict_to_yaml('NOISE_SCHEDULER',
__class__.__name__,
FlowMatchUniformScheduler.para_dict,
set_name=True)
@NOISE_SCHEDULERS.register_class()
class FlowMatchSigmoidScheduler(FlowMatchUniformScheduler):
para_dict = {
'SIGMOID_SCALE': {
'value': 1,
'description': 'The scale for the sigmoid function.'
}
}
def init_params(self):
super().init_params()
self.sigmoid_scale = self.cfg.get('SIGMOID_SCALE', 1)
def sigma_to_t(self, sigma, **kwargs):
t = -torch.log(1 / sigma - 1) / self.sigmoid_scale
return t * self.num_timesteps
def t_to_sigma(self, t, **kwargs):
return torch.sigmoid(self.sigmoid_scale * t / self.num_timesteps)
@staticmethod
def get_config_template():
return dict_to_yaml('NOISE_SCHEDULER',
__class__.__name__,
FlowMatchSigmoidScheduler.para_dict,
set_name=True)
@NOISE_SCHEDULERS.register_class()
class FlowMatchShiftScheduler(FlowMatchUniformScheduler):
para_dict = {
'SHIFT': {
'value': 3,
'description': 'The shift factor for the timestamp.'
},
'SIGMOID_SCALE': {
'value': 1,
'description': 'The scale for the sigmoid function.'
}
}
def init_params(self):
super().init_params()
self.shift = self.cfg.get('SHIFT', 3)
self.sigmoid_scale = self.cfg.get('SIGMOID_SCALE', 1)
def add_noise(self, x_0, noise=None, t=None, **kwargs):
if t is None:
logits_norm = torch.randn(x_0.shape[0], device=x_0.device)
logits_norm = logits_norm * self.sigmoid_scale # larger scale for more uniform sampling
t = logits_norm.sigmoid() * self.num_timesteps
sigma = self.t_to_sigma(t)
shape = (x_0.size(0), ) + (1, ) * (x_0.ndim - 1)
x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
return ScheduleOutput(x_0=x_0,
x_t=x_t,
t=t,
sigma=sigma,
alpha=self.t_to_alpha(t))
def sigma_to_t(self, sigma, **kwargs):
t = sigma / (sigma - self.shift * sigma + self.shift)
return t * self.num_timesteps
def t_to_sigma(self, t, **kwargs):
t = t / self.num_timesteps
return (t * self.shift) / (1 + (self.shift - 1) * t)
@staticmethod
def get_config_template():
return dict_to_yaml('NOISE_SCHEDULER',
__class__.__name__,
FlowMatchShiftScheduler.para_dict,
set_name=True)
@NOISE_SCHEDULERS.register_class()
class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler):
para_dict = {
'SHIFT': {
'value': True,
'description': 'Use timestamp shift or not, default is True.'
},
'SIGMOID_SCALE': {
'value': 1,
'description':
'The scale of sigmoid function for sampling timesteps.'
},
'BASE_SHIFT': {
'value': 0.5,
'description': 'The base shift factor for the timestamp.'
},
'MAX_SHIFT': {
'value': 1.15,
'description': 'The max shift factor for the timestamp.'
}
}
def init_params(self):
super().init_params()
self.shift = self.cfg.get('SHIFT', True)
self.sigmoid_scale = self.cfg.get('SIGMOID_SCALE', 1)
self.base_shift = self.cfg.get('BASE_SHIFT', 0.5)
self.max_shift = self.cfg.get('MAX_SHIFT', 1.15)
def time_shift(self, mu: float, sigma_scale: float, t: Tensor):
return math.exp(mu) / (math.exp(mu) + (1 / t - 1)**sigma_scale)
def sigma_shift(self, mu: float, sigma_scale: float, sigma: Tensor):
return 1 / (torch.pow(
(1 - sigma) * math.exp(mu) / sigma, sigma_scale) + 1)
def get_lin_function(self,
x1: float = 256,
y1: float = 0.5,
x2: float = 4096,
y2: float = 1.15) -> Callable[[float], float]:
m = (y2 - y1) / (x2 - x1)
b = y1 - m * x1
return lambda x: m * x + b
def add_noise(self, x_0, noise=None, t=None, **kwargs):
if x_0.ndim == 3:
seq_len = x_0.shape[2] // 4
else:
n, _, h, w = x_0.shape
seq_len = (h // 2 * w // 2)
if t is None:
logits_norm = torch.randn(x_0.shape[0], device=x_0.device)
logits_norm = logits_norm * self.sigmoid_scale # larger scale for more uniform sampling
t = logits_norm.sigmoid() * self.num_timesteps
sigma = self.t_to_sigma(t, seq_len=seq_len)
shape = (x_0.size(0), ) + (1, ) * (x_0.ndim - 1)
x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
return ScheduleOutput(x_0=x_0,
x_t=x_t,
t=t,
sigma=sigma,
alpha=self.t_to_alpha(t))
def sigma_to_t(self, sigma, **kwargs):
seq_len = kwargs.get('seq_len', 256)
if self.shift:
mu = self.get_lin_function(y1=self.base_shift,
y2=self.max_shift)(seq_len)
sigma = self.sigma_shift(mu, 1.0, sigma)
t = torch.as_tensor(sigma, dtype=torch.float32)
return t * self.num_timesteps
def t_to_sigma(self, t, **kwargs):
seq_len = kwargs.get('seq_len', 256)
t = t / self.num_timesteps
if self.shift:
mu = self.get_lin_function(y1=self.base_shift,
y2=self.max_shift)(seq_len)
t = self.time_shift(mu, 1.0, t)
sigma = torch.as_tensor(t, dtype=torch.float32)
return sigma
@staticmethod
def get_config_template():
return dict_to_yaml('NOISE_SCHEDULER',
__class__.__name__,
FlowMatchFluxShiftScheduler.para_dict,
set_name=True)
@NOISE_SCHEDULERS.register_class()
class FlowMatchSigmaScheduler(FlowMatchUniformScheduler):
para_dict = {
'WEIGHTING_SCHEME': {
'value':
'logit_normal',
'description':
'The weighting scheme for sampling timesteps, '
"choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']."
},
'SHIFT': {
'value': 3.0,
'description': 'The shift factor for the timestamp.'
},
'LOGIT_MEAN': {
'value':
0.0,
'description':
'The mean of the logit distribution for sampling timesteps.'
},
'LOGIT_STD': {
'value':
1.0,
'description':
'The standard deviation of the logit distribution for sampling timesteps.'
},
'MODE_SCALE': {
'value':
1.29,
'description':
'The scale factor for the mode of the logit distribution for sampling timesteps.'
}
}
def init_params(self):
super().init_params()
self.weighting_scheme = self.cfg.get('WEIGHTING_SCHEME',
'logit_normal')
self.logit_mean = self.cfg.get('LOGIT_MEAN', 0.0)
self.logit_std = self.cfg.get('LOGIT_STD', 1.0)
self.mode_scale = self.cfg.get('MODE_SCALE', 1.29)
self.shift = self.cfg.get('SHIFT', 1.0)
def get_schedule(self):
timesteps = np.linspace(1,
self.num_timesteps,
self.num_timesteps,
dtype=np.float32).copy()
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
self._timesteps = timesteps
timesteps = timesteps / self.num_timesteps
self._sigmas = self.shift * timesteps / (1 +
(self.shift - 1) * timesteps)
self._betas = torch.sqrt(self.sigmas_to_square_betas(self._sigmas))
self._alphas = torch.sqrt(1 - self.betas**2)
def add_noise(self, x_0, noise=None, t=None, **kwargs):
if t is None:
if self.weighting_scheme == 'logit_normal':
t = torch.normal(mean=self.logit_mean,
std=self.logit_std,
size=(x_0.shape[0], ),
device=x_0.device)
else:
t = torch.rand(x_0.shape[0], device=x_0.device)
t = self.compute_density_for_timestep_sampling(
t) * self.num_timesteps
sigma = self.t_to_sigma(t)
shape = (x_0.size(0), ) + (1, ) * (x_0.ndim - 1)
x_t = (1 - sigma.view(shape)) * x_0 + sigma.view(shape) * noise
return ScheduleOutput(x_0=x_0,
x_t=x_t,
t=t,
sigma=sigma,
alpha=self.t_to_alpha(t))
def compute_density_for_timestep_sampling(self, t):
"""Compute the density for sampling the timesteps when doing SD3 training.
Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.
SD3 paper reference: https://arxiv.org/abs/2403.03206v1.
"""
if self.weighting_scheme == 'logit_normal':
# See 3.1 in the SD3 paper ($rf/lognorm(0.00,1.00)$).
t = torch.nn.functional.sigmoid(t)
elif self.weighting_scheme == 'mode':
t = 1 - t - self.mode_scale * (torch.cos(math.pi * t / 2)**2 - 1 +
t)
return t
def sigma_to_t(self, sigma, **kwargs):
raise NotImplementedError
def t_to_sigma(self, t, **kwargs):
indices = t.long()
indices[indices >= self.num_timesteps] = self.num_timesteps - 1
timesteps = self.timesteps.to(t)[indices]
step_indices = [(self.timesteps.to(t) == t).nonzero().item()
for t in timesteps]
sigma = self.sigmas[step_indices].flatten().to(t)
return sigma
@staticmethod
def get_config_template():
return dict_to_yaml('NOISE_SCHEDULER',
__class__.__name__,
FlowMatchSigmaScheduler.para_dict,
set_name=True)
if __name__ == '__main__':
from scepter.modules.utils.config import Config
cfg = Config(cfg_dict={
'NAME': 'FlowMatchShiftScheduler',
'SHIFT': 1.15
},
load=False)
scheduler = NOISE_SCHEDULERS.build(cfg)
+13
View File
@@ -0,0 +1,13 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import torch
def _i(tensor, t, x):
"""
Index tensor using t and format the output according to x.
"""
shape = (x.size(0), ) + (1, ) * (x.ndim - 1)
if isinstance(t, torch.Tensor):
t = t.to(tensor.device)
return tensor[t].view(shape).to(x.device)
@@ -5,3 +5,4 @@ from scepter.modules.model.embedder.embedder import (
ConcatTimestepEmbedderND, FrozenCLIPEmbedder, FrozenCLIPEmbedder2,
FrozenOpenCLIPEmbedder, FrozenOpenCLIPEmbedder2, GeneralConditioner,
IPAdapterPlusEmbedder, RefCrossEmbedder, SD3TextEmbedder, T5EmbedderHF)
from scepter.modules.model.embedder.flux_embedder import HFEmbedder
+105 -16
View File
@@ -9,8 +9,11 @@ import numpy as np
import open_clip
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.utils.dlpack
from einops import rearrange
from torch.utils.checkpoint import checkpoint
from scepter.modules.model.backbone.unet.unet_utils import Timestep
from scepter.modules.model.embedder.base_embedder import BaseEmbedder
from scepter.modules.model.embedder.resampler import Resampler
@@ -21,7 +24,6 @@ from scepter.modules.model.utils.basic_utils import expand_dims_like
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from torch.utils.checkpoint import checkpoint
try:
from transformers import (CLIPTextModel, CLIPTokenizer,
@@ -831,22 +833,30 @@ class T5EmbedderHF(BaseEmbedder):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
pretrained_path = cfg.get('PRETRAINED_MODEL', None)
t5_dtype = cfg.get('T5_DTYPE', None)
self.t5_dtype = cfg.get('T5_DTYPE', 'float32')
assert pretrained_path
with FS.get_dir_to_local_dir(pretrained_path,
wait_finish=True) as local_path:
if t5_dtype is not None:
self.model = T5EncoderModel.from_pretrained(
local_path, torch_dtype=getattr(torch, t5_dtype))
else:
self.model = T5EncoderModel.from_pretrained(local_path)
self.model = T5EncoderModel.from_pretrained(
local_path,
torch_dtype=getattr(
torch,
'float' if self.t5_dtype == 'float32' else self.t5_dtype))
tokenizer_path = cfg.get('TOKENIZER_PATH', None)
self.length = cfg.get('LENGTH', 77)
self.use_grad = cfg.get('USE_GRAD', False)
self.clean = cfg.get('CLEAN', 'whitespace')
self.added_identifier = cfg.get('ADDED_IDENTIFIER', None)
if tokenizer_path:
self.tokenize_kargs = {'return_tensors': 'pt'}
with FS.get_dir_to_local_dir(tokenizer_path,
wait_finish=True) as local_path:
self.tokenizer = AutoTokenizer.from_pretrained(local_path)
if self.added_identifier is not None and isinstance(
self.added_identifier, list):
self.tokenizer = AutoTokenizer.from_pretrained(local_path)
else:
self.tokenizer = AutoTokenizer.from_pretrained(local_path)
if self.length is not None:
self.tokenize_kargs.update({
'padding': 'max_length',
@@ -868,12 +878,15 @@ class T5EmbedderHF(BaseEmbedder):
param.requires_grad = False
# encode && encode_text
def forward(self, tokens, return_mask=False):
def forward(self, tokens, return_mask=False, use_mask=True):
# tokenization
embedding_context = nullcontext if self.use_grad else torch.no_grad
with embedding_context():
x = self.model(tokens.input_ids.to(we.device_id),
tokens.attention_mask.to(we.device_id))
if use_mask:
x = self.model(tokens.input_ids.to(we.device_id),
tokens.attention_mask.to(we.device_id))
else:
x = self.model(tokens.input_ids.to(we.device_id))
x = x.last_hidden_state
# if not self.return_pooled:
# return x.detach()
@@ -882,7 +895,7 @@ class T5EmbedderHF(BaseEmbedder):
if return_mask:
return x.detach() + 0.0, tokens.attention_mask.to(we.device_id)
else:
return x.detach() + 0.0
return x.detach() + 0.0, None
def pool(self, x, tokens):
# take features from the eot embedding (eot_token is the highest number in each sequence)
@@ -897,7 +910,7 @@ class T5EmbedderHF(BaseEmbedder):
elif self.clean == 'canonicalize':
text = canonicalize(basic_clean(text))
elif self.clean == 'heavy':
text = heavy_clean(heavy_clean(text))
text = heavy_clean(basic_clean(text))
return text
def encode_text(self,
@@ -907,14 +920,90 @@ class T5EmbedderHF(BaseEmbedder):
return_mask=False):
return self(tokens, return_mask=return_mask)
def encode(self, text, return_mask=False):
def encode(self, text, return_mask=False, use_mask=True):
if isinstance(text, str):
text = [text]
if self.clean:
text = [self._clean(u) for u in text]
assert self.tokenizer is not None
tokens = self.tokenizer(text, **self.tokenize_kargs)
return self(tokens, return_mask=return_mask)
cont, mask = [], []
with torch.autocast(device_type='cuda',
enabled=self.t5_dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, self.t5_dtype)):
for tt in text:
tokens = self.tokenizer([tt], **self.tokenize_kargs)
one_cont, one_mask = self(tokens,
return_mask=return_mask,
use_mask=use_mask)
cont.append(one_cont)
mask.append(one_mask)
if return_mask:
return torch.cat(cont, dim=0), torch.cat(mask, dim=0)
else:
return torch.cat(cont, dim=0)
def encode_longlist(self, text_list, return_mask=True):
text_max_len = max([len(p) for p in text_list]) * self.length
cont_list, cont_mask_list = [], []
for pp in text_list:
cont, cont_mask = self.encode(pp, return_mask=return_mask)
cont_channel, cont_dim = cont.shape[0] * cont.shape[1], cont.shape[
2]
cont = cont.view(cont_channel, cont_dim)
cont_mask_channel = cont_mask.shape[0] * cont_mask.shape[1]
cont_mask = cont_mask.view(cont_mask_channel)
select_cont = cont[cont_mask == 1]
select_cont_mask, _ = torch.sort(cont_mask, dim=0, descending=True)
if select_cont.shape[0] != text_max_len:
select_cont = F.pad(
select_cont,
(0, 0, 0, text_max_len - select_cont.shape[0]))
if select_cont_mask.shape[0] != text_max_len:
select_cont_mask = F.pad(
select_cont_mask,
(0, text_max_len - select_cont_mask.shape[0]))
cont_list.append(select_cont)
cont_mask_list.append(select_cont_mask)
return torch.stack(cont_list), torch.stack(cont_mask_list)
def encode_longlist_v1(self, text_list, return_mask=True):
cont_list = []
max_len = 0
for pp in text_list:
cont, cont_mask = self.encode(pp, return_mask=True)
txt_lens = cont_mask.flatten(start_dim=1).sum(dim=-1)
pp_cont = torch.cat(
[c[:txt_len] for c, txt_len in zip(cont, txt_lens)], dim=0)
max_len = pp_cont.size(0) if pp_cont.size(0) > max_len else max_len
cont_list.append(pp_cont)
cont = torch.cat([
torch.cat([c, c.new_zeros(max_len - c.size(0), c.size(1))],
dim=0).unsqueeze(0) for c in cont_list
],
dim=0)
if return_mask:
cont_mask = torch.cat([
torch.cat(
[c.new_ones(c.size(0)),
c.new_zeros(max_len - c.size(0))],
dim=-1).unsqueeze(0) for c in cont_list
],
dim=0).type(torch.long, non_blocking=True)
return cont, cont_mask
else:
return cont
def encode_list(self, text_list, return_mask=True):
cont_list = []
mask_list = []
for pp in text_list:
cont, cont_mask = self.encode(pp, return_mask=return_mask)
cont_list.append(cont)
mask_list.append(cont_mask)
if return_mask:
return cont_list, mask_list
else:
return cont_list
@staticmethod
def get_config_template():
@@ -0,0 +1,171 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import torch
import transformers
from scepter.modules.model.embedder.base_embedder import BaseEmbedder
from scepter.modules.model.registry import EMBEDDERS
from scepter.modules.model.tokenizer.tokenizer_component import (
basic_clean, canonicalize, whitespace_clean)
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.file_system import FS
@EMBEDDERS.register_class()
class HFEmbedder(BaseEmbedder):
para_dict = {
'HF_MODEL_CLS': {
'value': None,
'description': 'huggingface cls in transfomer'
},
'MODEL_PATH': {
'value': None,
'description': 'model folder path'
},
'HF_TOKENIZER_CLS': {
'value': None,
'description': 'huggingface cls in transfomer'
},
'TOKENIZER_PATH': {
'value': None,
'description': 'tokenizer folder path'
},
'MAX_LENGTH': {
'value': 77,
'description': 'max length of input'
},
'OUTPUT_KEY': {
'value': 'last_hidden_state',
'description': 'output key'
},
'D_TYPE': {
'value': 'float',
'description': 'dtype'
},
'BATCH_INFER': {
'value': False,
'description': 'batch infer'
}
}
para_dict.update(BaseEmbedder.para_dict)
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
hf_model_cls = cfg.get('HF_MODEL_CLS', None)
model_path = cfg.get('MODEL_PATH', None)
hf_tokenizer_cls = cfg.get('HF_TOKENIZER_CLS', None)
tokenizer_path = cfg.get('TOKENIZER_PATH', None)
self.max_length = cfg.get('MAX_LENGTH', 77)
self.output_key = cfg.get('OUTPUT_KEY', 'last_hidden_state')
self.d_type = cfg.get('D_TYPE', 'float')
self.clean = cfg.get('CLEAN', 'whitespace')
self.batch_infer = cfg.get('BATCH_INFER', False)
torch_dtype = getattr(torch, self.d_type)
assert hf_model_cls is not None and hf_tokenizer_cls is not None
assert model_path is not None and tokenizer_path is not None
with FS.get_dir_to_local_dir(tokenizer_path,
wait_finish=True) as local_path:
self.tokenizer = getattr(transformers,
hf_tokenizer_cls).from_pretrained(
local_path,
max_length=self.max_length,
torch_dtype=torch_dtype)
with FS.get_dir_to_local_dir(model_path,
wait_finish=True) as local_path:
self.hf_module = getattr(transformers,
hf_model_cls).from_pretrained(
local_path, torch_dtype=torch_dtype)
self.hf_module = self.hf_module.eval().requires_grad_(False)
def forward(self, text: list[str], return_mask=False):
batch_encoding = self.tokenizer(
text,
truncation=True,
max_length=self.max_length,
return_length=False,
return_overflowing_tokens=False,
padding='max_length',
return_tensors='pt',
)
outputs = self.hf_module(
input_ids=batch_encoding['input_ids'].to(self.hf_module.device),
attention_mask=None,
output_hidden_states=False,
)
if return_mask:
return outputs[
self.output_key], batch_encoding['attention_mask'].to(
self.hf_module.device)
else:
return outputs[self.output_key], None
def encode(self, text, return_mask=False):
if isinstance(text, str):
text = [text]
if self.clean:
text = [self._clean(u) for u in text]
if not self.batch_infer:
cont, mask = [], []
for tt in text:
one_cont, one_mask = self([tt], return_mask=return_mask)
cont.append(one_cont)
mask.append(one_mask)
if return_mask:
return torch.cat(cont, dim=0), torch.cat(mask, dim=0)
else:
return torch.cat(cont, dim=0)
else:
ret_data = self(text, return_mask=return_mask)
if return_mask:
return ret_data
else:
return ret_data[0]
def _clean(self, text):
if self.clean == 'whitespace':
text = whitespace_clean(basic_clean(text))
elif self.clean == 'lower':
text = whitespace_clean(basic_clean(text)).lower()
elif self.clean == 'canonicalize':
text = canonicalize(basic_clean(text))
return text
@staticmethod
def get_config_template():
return dict_to_yaml('EMBEDDER',
__class__.__name__,
HFEmbedder.para_dict,
set_name=True)
@EMBEDDERS.register_class()
class T5PlusClipFluxEmbedder(BaseEmbedder):
"""
Uses the OpenCLIP transformer encoder for text
"""
para_dict = {'T5_MODEL': {}, 'CLIP_MODEL': {}}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.t5_model = EMBEDDERS.build(cfg.T5_MODEL, logger=logger)
self.clip_model = EMBEDDERS.build(cfg.CLIP_MODEL, logger=logger)
def encode(self, text):
t5_embeds = self.t5_model.encode(text, return_mask=False)
clip_embeds = self.clip_model.encode(text, return_mask=False)
# change embedding strategy here
return {
'context': t5_embeds,
'y': clip_embeds,
}
@staticmethod
def get_config_template():
return dict_to_yaml('EMBEDDER',
__class__.__name__,
T5PlusClipFluxEmbedder.para_dict,
set_name=True)
+2 -1
View File
@@ -5,4 +5,5 @@ 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_sce, ldm_sd3, ldm_xl,
ldm_flux)
@@ -1,18 +1,20 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import numbers
import re
from collections import OrderedDict
import numpy as np
import torch
from scepter.modules.model.network.train_module import TrainModule
from scepter.modules.model.registry import BACKBONES, LOSSES, MODELS
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from einops import repeat
import math
import torch.nn.functional as F
class DiagonalGaussianDistribution(object):
def __init__(self, mean, logvar, deterministic=False):
self.mean = mean
@@ -77,6 +79,11 @@ class AutoencoderKL(TrainModule):
'value': 16,
'description': ''
},
'SCALE_FACTOR': {
'value': None,
'description':
'if is not None, will used to scale the latent space.'
},
}
def __init__(self, cfg, logger=None):
@@ -89,6 +96,7 @@ class AutoencoderKL(TrainModule):
self.ignore_keys = self.cfg.get('IGNORE_KEYS', [])
self.batch_size = self.cfg.get('BATCH_SIZE', 16)
self.use_conv = self.cfg.get('USE_CONV', True)
self.scale_factor = self.cfg.get('SCALE_FACTOR', None)
self.construct_network()
self.init_network()
@@ -169,6 +177,7 @@ class AutoencoderKL(TrainModule):
return z
def _encode(self, x, return_mom=False):
h = self.encoder(x)
moments = self.conv1(h)
if return_mom:
@@ -176,9 +185,15 @@ class AutoencoderKL(TrainModule):
mean, logvar = torch.chunk(moments, 2, dim=1)
posterior = DiagonalGaussianDistribution(mean, logvar)
z = posterior.sample()
if self.scale_factor is not None and isinstance(
self.scale_factor, numbers.Number):
z = self.scale_factor * z
return z
def _decode(self, z):
if self.scale_factor is not None and isinstance(
self.scale_factor, numbers.Number):
z = z / self.scale_factor
z = self.conv2(z)
dec = self.decoder(z)
return dec
@@ -268,13 +283,274 @@ class AutoencoderKL(TrainModule):
AutoencoderKL.para_dict,
set_name=True)
def timestep_embedding(timesteps, dim, max_period=10000, repeat_only=False):
"""
Create sinusoidal timestep embeddings.
:param timesteps: a 1-D Tensor of N indices, one per batch element.
These may be fractional.
:param dim: the dimension of the output.
:param max_period: controls the minimum frequency of the embeddings.
:return: an [N x dim] Tensor of positional embeddings.
"""
if not repeat_only:
half = dim // 2
freqs = torch.exp(
-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half
).to(device=timesteps.device)
args = torch.mm(timesteps.float().unsqueeze(1), freqs.unsqueeze(0)).view(timesteps.shape[0], len(freqs))
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
else:
embedding = repeat(timesteps, 'b -> b d', d=dim)
return embedding
if __name__ == '__main__':
import argparse
from scepter.modules.utils.config import Config
from scepter.modules.utils.logger import get_logger
std_logger = get_logger(name='scepter')
parser = argparse.ArgumentParser(description='Argparser for Scepter:\n')
cfg = Config(load=True, parser_ins=parser)
model = AutoencoderKL(cfg, logger=std_logger)
model.load_pretrained_model(cfg.PRETRAINED_MODEL)
@MODELS.register_class()
class AutoencoderKLFlux(TrainModule):
para_dict = {
"ENCODER": {},
"DECODER": {},
"LOSS": {},
"EMBED_DIM": {
"value": 4,
"description": ""
},
"PRETRAINED_MODEL": {
"value": None,
"description": ""
},
"IGNORE_KEYS": {
"value": [],
"description": ""
},
"BATCH_SIZE": {
"value": 16,
"description": ""
},
}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.encoder_cfg = self.cfg.ENCODER
self.decoder_cfg = self.cfg.DECODER
self.loss_cfg = self.cfg.get("LOSS", None)
self.embed_dim = self.cfg.get("EMBED_DIM", 4)
self.pretrained_model = self.cfg.get("PRETRAINED_MODEL", None)
#
self.ignore_keys = self.cfg.get("IGNORE_KEYS", [])
self.batch_size = self.cfg.get("BATCH_SIZE", 16)
self.resize_nx = self.cfg.get("RESIZE_NX", 1)
self.use_rembed = self.cfg.get("USE_REMBED", True)
self.use_conv = self.cfg.get('USE_CONV', True)
self.scale_factor = self.cfg.get('SCALE_FACTOR', None)
self.shift_factor = self.cfg.get('SHIFT_FACTOR', None)
self.construct_network()
def construct_network(self):
z_channels = self.encoder_cfg.Z_CHANNELS
self.encoder = BACKBONES.build(self.encoder_cfg, logger=self.logger)
self.decoder = BACKBONES.build(self.decoder_cfg, logger=self.logger)
self.conv1 = torch.nn.Conv2d(2 * z_channels, 2 * self.embed_dim, 1) if self.use_conv else torch.nn.Identity()
self.conv2 = torch.nn.Conv2d(self.embed_dim, z_channels,1) if self.use_conv else torch.nn.Identity()
# freeze encoder
for param in self.encoder.parameters():
param.requires_grad = False
for param in self.conv1.parameters():
param.requires_grad = False
for param in self.conv2.parameters():
param.requires_grad = False
def load_pretrained_model(self, pretrained_model):
if pretrained_model is not None:
with FS.get_from(pretrained_model, wait_finish=True) as local_model:
self.init_from_ckpt(local_model)
def init_from_ckpt(self, path, ignore_keys=list()):
if path.find('.safetensors') > -1:
from safetensors import safe_open
sd = OrderedDict()
with safe_open(path, framework="pt", device='cpu') as f:
for k in f.keys():
sd[k] = f.get_tensor(k)
else:
sd = torch.load(path, map_location="cpu")
if path.find('.pt') > -1 and 'state_dict' in sd:
sd = sd['state_dict']
elif path.find('.ckpt') > -1 and 'state_dict' in sd:
sd = sd['state_dict']
elif path.find('.pth') > -1 and 'model' in sd:
sd = sd['model']
new_sd = OrderedDict()
for k, v in sd.items():
ignored = False
for ik in ignore_keys:
if ik in k:
if we.rank == 0:
self.logger.info("ignore key {} from state_dict.".format(k))
ignored = True
break
k = k.replace("post_quant_conv", "conv2") if "post_quant_conv" in k else k
k = k.replace("quant_conv", "conv1") if "quant_conv" in k else k
if not ignored:
new_sd[k] = v
missing, unexpected = self.load_state_dict(new_sd, strict=False)
if we.rank == 0:
self.logger.info(f"Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys")
if len(missing) > 0:
self.logger.info(f"Missing Keys:\n {missing}")
if len(unexpected) > 0:
self.logger.info(f"\nUnexpected Keys:\n {unexpected}")
@torch.no_grad()
def encode(self, x, sample_posterior = True):
h = self.encoder(x)
moments = self.conv1(h)
mean, logvar = torch.chunk(moments, 2, dim=1)
posterior = DiagonalGaussianDistribution(mean, logvar)
if sample_posterior:
z = posterior.sample()
else:
z = posterior.mode()
if self.shift_factor is not None and isinstance(self.shift_factor, numbers.Number):
z = z - self.shift_factor
if self.scale_factor is not None and isinstance(self.scale_factor, numbers.Number):
z = self.scale_factor * z
return z, posterior
def decode(self, z, **kwargs):
b, c, h, w = z.size()
if kwargs.get('resize_nx', None) is not None and self.use_rembed:
resize_nx = kwargs['resize_nx']
if not torch.is_tensor(resize_nx):
resize_nx = torch.full((b,), resize_nx, device=we.device_id, dtype=z.dtype)
rembed = timestep_embedding(resize_nx, dim=self.decoder_cfg.CH_MULT[-1] * self.decoder_cfg.CH)
else:
rembed = None
if self.scale_factor is not None and isinstance(self.scale_factor, numbers.Number):
z = z / self.scale_factor
if self.shift_factor is not None and isinstance(self.shift_factor, numbers.Number):
z = z + self.shift_factor
z = self.conv2(z)
# add grad;
if rembed is not None:
dec = self.decoder(z, rembed)
else:
dec = self.decoder(z)
return dec
def share_forward(self, image=None, sample_posterior=True, **kwargs):
# rembed: resize embedding
if image is not None:
z, posterior = self.encode(image, sample_posterior = sample_posterior)
else:
latent = kwargs.pop("latent", None)
assert latent is not None
z = latent
posterior = None
if self.shift_factor is not None and isinstance(self.scale_factor, numbers.Number):
z = z - self.shift_factor
if self.scale_factor is not None and isinstance(self.scale_factor, numbers.Number):
z = self.scale_factor * z
dec = self.decode(z, **kwargs)
return dec, posterior
def forward(self, **kwargs):
if self.training:
ret = self.forward_train(**kwargs)
else:
ret = self.forward_test(**kwargs)
return ret
def forward_train(self,
image=None,
gt_image=None,
sample_posterior=True,
optimizer_idx=0,
global_step=0,
**kwargs):
if gt_image is None:
gt_image = copy.deepcopy(image)
reconstructions, posterior = self.share_forward(image, sample_posterior, **kwargs)
ret = {}
if optimizer_idx == 0:
# train encoder+decoder+logvar
aeloss, log_dict_ae = self.loss(gt_image, reconstructions, posterior, optimizer_idx,
global_step, last_layer=self.get_last_layer(), split="train")
# self.logger.info(f"aeloss: {aeloss.detach().cpu().item()}, ")
ret["loss"] = aeloss
ret.update(log_dict_ae)
if optimizer_idx == 1:
# train the discriminator
discloss, log_dict_disc = self.loss(gt_image, reconstructions, posterior, optimizer_idx,
global_step, last_layer=self.get_last_layer(), split="train")
# self.logger.info(f"discloss: {discloss.detach().cpu().item()}, ")
ret["loss"] = discloss
ret.update(log_dict_disc)
return ret
@torch.no_grad()
def forward_test(self,
image=None,
gt_image=None,
sample_posterior=True,
**kwargs):
resize_nx_ = 1
if image is not None:
b, c, h, w = image.size()
if kwargs.get('resize_ex'):
resize_nx_ = kwargs.pop('resize_ex')
image = F.interpolate(image, (int(float(h) / resize_nx_), int(float(w) / resize_nx_)), mode='bicubic')
image = F.interpolate(image, (h, w), mode='bicubic')
kwargs["resize_nx"] = resize_nx_
elif kwargs.get('resize_nx', None) is not None:
resize_nx_ = kwargs['resize_nx']
if gt_image is None:
if image is not None:
gt_image = copy.deepcopy(image)
# kwargs["resize_nx"] = resize_nx_
reconstructions, posterior = self.share_forward(image, sample_posterior, **kwargs)
reconstructions = torch.clamp((reconstructions + 1.0) / 2.0, min=0.0, max=1.0)
if gt_image is not None:
gt_image = torch.clamp((gt_image + 1.0) / 2.0, min=0.0, max=1.0)
else:
gt_image = [None for _ in range(reconstructions.shape[0])]
if image is not None:
lr_image = torch.clamp((image + 1.0) / 2.0, min=0.0, max=1.0)
else:
lr_image = [None for _ in range(reconstructions.shape[0])]
ret = list()
if torch.is_tensor(resize_nx_):
resize_nx = [nx.item() for nx, in zip(resize_nx_.cpu())]
else:
resize_nx = [resize_nx_ for _ in range(reconstructions.size(0))]
for img, ori, gt_img, nx in zip(reconstructions, lr_image, gt_image, resize_nx):
ret.append({
"prompt": "",
"n_prompt": "",
"image": img,
"lr_image": ori,
"gt_image": gt_img,
"resize_nx": nx
})
return ret
def get_last_layer(self):
if hasattr(self.decoder, 'conv_out'):
return self.decoder.conv_out.weight
else:
return self.decoder.head[-1].weight
@staticmethod
def get_config_template():
return dict_to_yaml("MODEL", __class__.__name__, AutoencoderKLFlux.para_dict, set_name=True)
@@ -1,6 +1,7 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.network.ldm.ldm import LatentDiffusion
from scepter.modules.model.network.ldm.ldm_ace import LatentDiffusionACE
from scepter.modules.model.network.ldm.ldm_edit import LatentDiffusionEdit
from scepter.modules.model.network.ldm.ldm_pixart import LatentDiffusionPixart
from scepter.modules.model.network.ldm.ldm_sce import (
+16 -8
View File
@@ -10,7 +10,7 @@ import torch
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
from scepter.modules.model.network.diffusion.schedules import noise_schedule
from scepter.modules.model.network.train_module import TrainModule
from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, LOSSES,
from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, LOSSES, DIFFUSIONS,
MODELS, TOKENIZERS)
from scepter.modules.model.utils.basic_utils import count_params, default
from scepter.modules.utils.config import dict_to_yaml
@@ -112,13 +112,18 @@ class LatentDiffusion(TrainModule):
if self.zero_terminal_snr:
assert self.parameterization == 'v', 'Now zero_terminal_snr only support v-prediction mode.'
self.sigmas = noise_schedule(schedule=self.schedule_args.pop('name'),
n=self.num_timesteps,
zero_terminal_snr=self.zero_terminal_snr,
**self.schedule_args)
self.diffusion = GaussianDiffusion(
sigmas=self.sigmas, prediction_type=self.parameterization)
diffusion_cfg = self.cfg.get("DIFFUSION", None)
if diffusion_cfg is not None:
if self.cfg.have("WORK_DIR"):
diffusion_cfg.WORK_DIR = self.cfg.WORK_DIR
self.diffusion = DIFFUSIONS.build(diffusion_cfg, logger=self.logger)
else:
self.sigmas = noise_schedule(schedule=self.schedule_args.pop('name'),
n=self.num_timesteps,
zero_terminal_snr=self.zero_terminal_snr,
**self.schedule_args)
self.diffusion = GaussianDiffusion(
sigmas=self.sigmas, prediction_type=self.parameterization)
self.pretrained_model = self.cfg.get('PRETRAINED_MODEL', None)
self.ignore_keys = self.cfg.get('IGNORE_KEYS', [])
@@ -131,6 +136,7 @@ class LatentDiffusion(TrainModule):
self.scale_factor = self.cfg.get('SCALE_FACTOR', 0.18215)
self.size_factor = self.cfg.get('SIZE_FACTOR', 8)
self.decoder_bias = self.cfg.get("DECODER_BIAS", 0)
self.default_n_prompt = self.cfg.get('DEFAULT_N_PROMPT', '')
self.default_n_prompt = '' if self.default_n_prompt is None else self.default_n_prompt
self.p_zero = self.cfg.get('P_ZERO', 0.0)
@@ -140,6 +146,7 @@ class LatentDiffusion(TrainModule):
if self.train_n_prompt is None:
self.train_n_prompt = ''
self.use_ema = self.cfg.get('USE_EMA', False)
self.eval_ema = self.cfg.get('EVAL_EMA', False)
self.model_ema_config = self.cfg.get('DIFFUSION_MODEL_EMA', None)
def construct_network(self):
@@ -290,6 +297,7 @@ class LatentDiffusion(TrainModule):
@torch.no_grad()
@torch.autocast('cuda', dtype=torch.float16)
def forward_test(self,
image=None,
prompt=None,
n_prompt=None,
sampler='ddim',
@@ -0,0 +1,351 @@
# -*- 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 import nn
from scepter.modules.model.network.ldm import LatentDiffusion
from scepter.modules.model.registry import MODELS
from scepter.modules.model.utils.basic_utils import check_list_of_list
from scepter.modules.model.utils.basic_utils import \
pack_imagelist_into_tensor_v2 as pack_imagelist_into_tensor
from scepter.modules.model.utils.basic_utils import (
to_device, unpack_tensor_into_imagelist)
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
class TextEmbedding(nn.Module):
def __init__(self, embedding_shape):
super().__init__()
self.pos = nn.Parameter(data=torch.zeros(embedding_shape))
@MODELS.register_class()
class LatentDiffusionACE(LatentDiffusion):
para_dict = LatentDiffusion.para_dict
para_dict['DECODER_BIAS'] = {'value': 0, 'description': ''}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.interpolate_func = lambda x: (F.interpolate(
x.unsqueeze(0),
scale_factor=1 / self.size_factor,
mode='nearest-exact') if x is not None else None)
self.text_indentifers = cfg.get('TEXT_IDENTIFIER', [])
self.use_text_pos_embeddings = cfg.get('USE_TEXT_POS_EMBEDDINGS',
False)
if self.use_text_pos_embeddings:
self.text_position_embeddings = TextEmbedding(
(10, 4096)).eval().requires_grad_(False)
else:
self.text_position_embeddings = None
self.logger.info(self.model)
@torch.no_grad()
def encode_first_stage(self, x, **kwargs):
return [
self.scale_factor *
self.first_stage_model._encode(i.unsqueeze(0).to(torch.float16))
for i in x
]
@torch.no_grad()
def decode_first_stage(self, z):
return [
self.first_stage_model._decode(1. / self.scale_factor *
i.to(torch.float16)) for i in z
]
def cond_stage_embeddings(self, prompt, edit_image, cont, cont_mask):
if self.use_text_pos_embeddings and not torch.sum(
self.text_position_embeddings.pos) > 0:
identifier_cont, identifier_cont_mask = getattr(
self.cond_stage_model, 'encode')(self.text_indentifers,
return_mask=True)
self.text_position_embeddings.load_state_dict(
{'pos': identifier_cont[:, 0, :]})
cont_, cont_mask_ = [], []
for pp, edit, c, cm in zip(prompt, edit_image, cont, cont_mask):
if isinstance(pp, list):
cont_.append([c[-1], *c] if len(edit) > 0 else [c[-1]])
cont_mask_.append([cm[-1], *cm] if len(edit) > 0 else [cm[-1]])
else:
raise NotImplementedError
return cont_, cont_mask_
def limit_batch_data(self, batch_data_list, log_num):
if log_num and log_num > 0:
batch_data_list_limited = []
for sub_data in batch_data_list:
if sub_data is not None:
sub_data = sub_data[:log_num]
batch_data_list_limited.append(sub_data)
return batch_data_list_limited
else:
return batch_data_list
def forward_train(self,
edit_image=[],
edit_image_mask=[],
image=None,
image_mask=None,
noise=None,
prompt=[],
**kwargs):
'''
Args:
edit_image: list of list of edit_image
edit_image_mask: list of list of edit_image_mask
image: target image
image_mask: target image mask
noise: default is None, generate automaticly
prompt: list of list of text
**kwargs:
Returns:
'''
assert check_list_of_list(prompt) and check_list_of_list(
edit_image) and check_list_of_list(edit_image_mask)
assert len(edit_image) == len(edit_image_mask) == len(prompt)
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 = {}
# process image
image = to_device(image)
x_start = self.encode_first_stage(image, **kwargs)
x_start, x_shapes = pack_imagelist_into_tensor(x_start) # B, C, L
n, _, _ = x_start.shape
t = torch.randint(0, self.num_timesteps, (n, ),
device=x_start.device).long()
context['x_shapes'] = x_shapes
# process image mask
image_mask = to_device(image_mask, strict=False)
context['x_mask'] = [self.interpolate_func(i) for i in image_mask
] if image_mask is not None else [None] * n
# process text
# with torch.autocast(device_type="cuda", enabled=True, dtype=torch.bfloat16):
prompt_ = [[pp] if isinstance(pp, str) else pp for pp in prompt]
try:
cont, cont_mask = getattr(self.cond_stage_model,
'encode_list')(prompt_, return_mask=True)
except Exception as e:
print(e, prompt_)
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont,
cont_mask)
context['crossattn'] = cont
# process edit image & edit image mask
edit_image = [to_device(i, strict=False) for i in edit_image]
edit_image_mask = [to_device(i, strict=False) for i in edit_image_mask]
e_img, e_mask = [], []
for u, m in zip(edit_image, edit_image_mask):
if m is None:
m = [None] * len(u) if u is not None else [None]
e_img.append(
self.encode_first_stage(u, **kwargs) if u is not None else u)
e_mask.append([
self.interpolate_func(i) if i is not None else None for i in m
])
context['edit'], context['edit_mask'] = e_img, e_mask
# process loss
loss = self.diffusion.loss(
x_0=x_start,
t=t,
noise=noise,
model=self.model,
model_kwargs={
'cond':
context,
'mask':
cont_mask,
'gc_seg':
gc_seg,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
},
**kwargs)
loss = loss.mean()
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
return ret
@torch.no_grad()
def forward_test(self,
edit_image=[],
edit_image_mask=[],
image=None,
image_mask=None,
prompt=[],
n_prompt=[],
sampler='ddim',
sample_steps=20,
guide_scale=4.5,
guide_rescale=0.5,
log_num=-1,
seed=2024,
**kwargs):
assert check_list_of_list(prompt) and check_list_of_list(
edit_image) and check_list_of_list(edit_image_mask)
assert len(edit_image) == len(edit_image_mask) == len(prompt)
assert self.cond_stage_model is not None
# gc_seg is unused
kwargs.pop('gc_seg', -1)
# prepare data
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],
log_num)
g = torch.Generator(device=we.device_id)
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
g.manual_seed(seed)
n_prompt = copy.deepcopy(prompt)
# only modify the last prompt to be zero
for nn_p_id, nn_p in enumerate(n_prompt):
if isinstance(nn_p, str):
n_prompt[nn_p_id] = ['']
elif isinstance(nn_p, list):
n_prompt[nn_p_id][-1] = ''
else:
raise NotImplementedError
# process image
image = to_device(image)
x = self.encode_first_stage(image, **kwargs)
noise = [
torch.empty(*i.shape, device=we.device_id).normal_(generator=g)
for i in x
]
noise, x_shapes = pack_imagelist_into_tensor(noise)
context['x_shapes'] = null_context['x_shapes'] = x_shapes
# process image mask
image_mask = to_device(image_mask, strict=False)
cond_mask = [self.interpolate_func(i) for i in image_mask
] if image_mask is not None else [None] * len(image)
context['x_mask'] = null_context['x_mask'] = cond_mask
# process text
# with torch.autocast(device_type="cuda", enabled=True, dtype=torch.bfloat16):
prompt_ = [[pp] if isinstance(pp, str) else pp for pp in prompt]
cont, cont_mask = getattr(self.cond_stage_model,
'encode_list')(prompt_, return_mask=True)
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont,
cont_mask)
null_cont, null_cont_mask = getattr(self.cond_stage_model,
'encode_list')(n_prompt,
return_mask=True)
null_cont, null_cont_mask = self.cond_stage_embeddings(
prompt, edit_image, null_cont, null_cont_mask)
context['crossattn'] = cont
null_context['crossattn'] = null_cont
# processe edit image & edit image mask
edit_image = [to_device(i, strict=False) for i in edit_image]
edit_image_mask = [to_device(i, strict=False) for i in edit_image_mask]
e_img, e_mask = [], []
for u, m in zip(edit_image, edit_image_mask):
if u is None:
continue
if m is None:
m = [None] * len(u)
e_img.append(self.encode_first_stage(u, **kwargs))
e_mask.append([self.interpolate_func(i) for i in m])
null_context['edit'] = context['edit'] = e_img
null_context['edit_mask'] = context['edit_mask'] = e_mask
# process sample
model = self.model_ema if self.use_ema and self.eval_ema else self.model
embedding_context = model.no_sync if isinstance(model, torch.distributed.fsdp.FullyShardedDataParallel) \
else nullcontext
with embedding_context():
samples = self.diffusion.sample(
sampler=sampler,
noise=noise,
model=model,
model_kwargs=[{
'cond':
context,
'mask':
cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
}, {
'cond':
null_context,
'mask':
null_cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
}] if guide_scale is not None and guide_scale > 1 else {
'cond':
context,
'mask':
cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
},
steps=sample_steps,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
show_progress=True,
**kwargs)
samples = unpack_tensor_into_imagelist(samples, x_shapes)
x_samples = self.decode_first_stage(samples)
outputs = list()
for i in range(len(prompt)):
rec_img = torch.clamp(
(x_samples[i] + 1.0) / 2.0 + self.decoder_bias / 255,
min=0.0,
max=1.0)
rec_img = rec_img.squeeze(0)
edit_imgs, edit_img_masks = [], []
if edit_image is not None and edit_image[i] is not None:
if edit_image_mask[i] is None:
edit_image_mask[i] = [None] * len(edit_image[i])
for edit_img, edit_mask in zip(edit_image[i],
edit_image_mask[i]):
edit_img = torch.clamp((edit_img + 1.0) / 2.0,
min=0.0,
max=1.0)
edit_imgs.append(edit_img.squeeze(0))
if edit_mask is None:
edit_mask = torch.ones_like(edit_img[[0], :, :])
edit_img_masks.append(edit_mask)
one_tup = {
'reconstruct_image': rec_img,
'instruction': prompt[i],
'edit_image': edit_imgs if len(edit_imgs) > 0 else None,
'edit_mask': edit_img_masks if len(edit_imgs) > 0 else None
}
if image is not None:
if image_mask is None:
image_mask = [None] * len(image)
ori_img = torch.clamp((image[i] + 1.0) / 2.0, min=0.0, max=1.0)
one_tup['target_image'] = ori_img.squeeze(0)
one_tup['target_mask'] = image_mask[i] if image_mask[
i] is not None else torch.ones_like(ori_img[[0], :, :])
outputs.append(one_tup)
return outputs
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
LatentDiffusionACE.para_dict,
set_name=True)
@@ -111,6 +111,7 @@ class LatentDiffusionEdit(LatentDiffusion):
@torch.no_grad()
@torch.autocast('cuda', dtype=torch.float16)
def forward_test(self,
image=None,
prompt=None,
n_prompt=None,
sampler='ddim',
@@ -0,0 +1,220 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import math
import numbers
import random
import torch
from scepter.modules.model.network.ldm import LatentDiffusion
from scepter.modules.model.registry import MODELS, BACKBONES, LOSSES, TOKENIZERS, EMBEDDERS, DIFFUSIONS
from scepter.modules.model.utils.basic_utils import disabled_train
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.model.utils.basic_utils import count_params
@MODELS.register_class()
class LatentDiffusionFlux(LatentDiffusion):
para_dict = LatentDiffusion.para_dict
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.guide_scale = cfg.get('GUIDE_SCALE', 3.5)
def init_params(self):
self.parameterization = self.cfg.get('PARAMETERIZATION', 'rf')
assert self.parameterization in [
'eps', 'x0', 'v', 'rf'
], 'currently only supporting "eps" and "x0" and "v" and "rf"'
diffusion_cfg = self.cfg.get("DIFFUSION", None)
assert diffusion_cfg is not None
if self.cfg.have("WORK_DIR"):
diffusion_cfg.WORK_DIR = self.cfg.WORK_DIR
self.diffusion = DIFFUSIONS.build(diffusion_cfg, logger=self.logger)
self.pretrained_model = self.cfg.get('PRETRAINED_MODEL', None)
self.ignore_keys = self.cfg.get('IGNORE_KEYS', [])
self.model_config = self.cfg.DIFFUSION_MODEL
self.first_stage_config = self.cfg.FIRST_STAGE_MODEL
self.cond_stage_config = self.cfg.COND_STAGE_MODEL
self.tokenizer_config = self.cfg.get('TOKENIZER', None)
self.loss_config = self.cfg.get('LOSS', None)
self.scale_factor = self.cfg.get('SCALE_FACTOR', 0.18215)
self.size_factor = self.cfg.get('SIZE_FACTOR', 16)
self.default_n_prompt = self.cfg.get('DEFAULT_N_PROMPT', '')
self.default_n_prompt = '' if self.default_n_prompt is None else self.default_n_prompt
self.p_zero = self.cfg.get('P_ZERO', 0.0)
self.train_n_prompt = self.cfg.get('TRAIN_N_PROMPT', '')
if self.default_n_prompt is None:
self.default_n_prompt = ''
if self.train_n_prompt is None:
self.train_n_prompt = ''
self.use_ema = self.cfg.get('USE_EMA', False)
self.model_ema_config = self.cfg.get('DIFFUSION_MODEL_EMA', None)
def construct_network(self):
# embedding_context = torch.device("meta") if self.model_config.get("PRETRAINED_MODEL", None) else nullcontext()
# with embedding_context:
self.model = BACKBONES.build(self.model_config, logger=self.logger).to(torch.bfloat16)
self.logger.info('all parameters:{}'.format(count_params(self.model)))
if self.use_ema:
if self.model_ema_config:
self.model_ema = BACKBONES.build(self.model_ema_config,
logger=self.logger)
else:
self.model_ema = copy.deepcopy(self.model)
self.model_ema = self.model_ema.eval()
for param in self.model_ema.parameters():
param.requires_grad = False
if self.loss_config:
self.loss = LOSSES.build(self.loss_config, logger=self.logger)
if self.tokenizer_config is not None:
self.tokenizer = TOKENIZERS.build(self.tokenizer_config,
logger=self.logger)
if self.first_stage_config:
self.first_stage_model = MODELS.build(self.first_stage_config,
logger=self.logger)
self.first_stage_model = self.first_stage_model.eval()
self.first_stage_model.train = disabled_train
for param in self.first_stage_model.parameters():
param.requires_grad = False
else:
self.first_stage_model = None
if self.tokenizer_config is not None:
self.cond_stage_config.KWARGS = {
'vocab_size': self.tokenizer.vocab_size
}
if self.cond_stage_config == '__is_unconditional__':
print(
f'Training {self.__class__.__name__} as an unconditional model.'
)
self.cond_stage_model = None
else:
model = EMBEDDERS.build(self.cond_stage_config, logger=self.logger)
self.cond_stage_model = model.eval().requires_grad_(False)
self.cond_stage_model.train = disabled_train
def noise_sample(self, num_samples, h, w, seed, dtype = torch.bfloat16):
noise = torch.randn(
num_samples,
16,
# allow for packing
2 * math.ceil(h / 16),
2 * math.ceil(w / 16),
device=we.device_id,
dtype=dtype,
generator=torch.Generator(device=we.device_id).manual_seed(seed),
)
return noise
def forward_train(self, image=None, noise=None, prompt=None, **kwargs):
x_start = self.encode_first_stage(image, **kwargs)
if prompt and self.cond_stage_model:
ctx = getattr(self.cond_stage_model, 'encode')(prompt)
else:
assert False
if 'index' in kwargs:
kwargs.pop('index')
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": ctx, "guidance": guide_scale},
noise=noise,
**kwargs)
loss = loss.mean()
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
return ret
@torch.no_grad()
def forward_test(self,
image=None,
prompt=None,
sampler='flow_eluer',
sample_steps=20,
seed=2023,
guide_scale=4.5,
guide_rescale=0.0,
show_process=False,
**kwargs):
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
if isinstance(prompt, str):
prompt = [prompt]
assert isinstance(prompt, list)
num_samples = len(prompt)
if prompt and self.cond_stage_model:
ctx = getattr(self.cond_stage_model, 'encode')(prompt)
else:
assert False
if 'index' in kwargs:
kwargs.pop('index')
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(
num_samples,
height,
width,
seed
)
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
samples = self.diffusion.sample(
noise=noise,
sampler=sampler,
model=self.model,
model_kwargs= {"cond": ctx, "guidance": guide_scale},
steps=sample_steps,
show_progress=True,
guide_scale = guide_scale,
return_intermediate=None,
**kwargs).float()
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
x_samples = self.decode_first_stage(samples).float()
x_samples = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0)
outputs = list()
for i, (p, img) in enumerate(zip(prompt, x_samples)):
one_tup = {'prompt': str(p), 'n_prompt': '', 'image': img}
outputs.append(one_tup)
return outputs
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
LatentDiffusionFlux.para_dict,
set_name=True)
@torch.no_grad()
def encode_first_stage(self, x, **kwargs):
z = self.first_stage_model.encode(x)
if isinstance(z, (tuple, list)):
z = z[0]
return z
@torch.no_grad()
def decode_first_stage(self, z):
return self.first_stage_model.decode(z)
@@ -3,20 +3,15 @@
import copy
import numbers
import random
from collections import OrderedDict
import torch
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
from scepter.modules.model.network.diffusion.schedules import noise_schedule
from scepter.modules.model.network.ldm import LatentDiffusion
from scepter.modules.model.network.train_module import TrainModule
from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, LOSSES,
MODELS, TOKENIZERS)
from scepter.modules.model.utils.basic_utils import count_params, default
from scepter.modules.model.utils.basic_utils import count_params
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
def disabled_train(self, mode=True):
@@ -134,6 +129,7 @@ class LatentDiffusionPixart(LatentDiffusion):
@torch.no_grad()
def forward_test(self,
image=None,
prompt=None,
label=None,
sampler='ddim',
+4 -1
View File
@@ -105,7 +105,10 @@ class LatentDiffusionSCEControl(LatentDiffusion):
hints = []
for ctr in control:
hint = self.control_processor(ctr)
hint = TT.ToTensor()(hint)
if len(hint.shape) == 4:
hint = TT.ToTensor()(hint.squeeze())
else:
hint = TT.ToTensor()(hint)
hints.append(hint)
hints = torch.stack(hints).to(control.device)
return hints
@@ -136,6 +136,7 @@ class LatentDiffusionSD3(LatentDiffusion):
@torch.no_grad()
def forward_test(self,
image=None,
prompt=None,
sampler='ddim',
sample_steps=20,

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