Compare commits
25
Commits
v1.1.0_dev
...
v1.3.0_dev
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
448cdba522 | ||
|
|
2a29446d45 | ||
|
|
cd33b4ab15 | ||
|
|
7d7943fed3 | ||
|
|
d48b2f110f | ||
|
|
82486adf38 | ||
|
|
a683061c6f | ||
|
|
4ec0492897 | ||
|
|
aac85fa94f | ||
|
|
3f267aaea2 | ||
|
|
53357f95d6 | ||
|
|
02c0ba9757 | ||
|
|
cef93bdbfe | ||
|
|
82132ff3a1 | ||
|
|
b886400e06 | ||
|
|
73984c4f9e | ||
|
|
f98adabeb3 | ||
|
|
edb46a615c | ||
|
|
fe3e11b49e | ||
|
|
e6b43f19f6 | ||
|
|
d9b207cf5b | ||
|
|
986349deac | ||
|
|
3f047be43c | ||
|
|
e8d8e63cba | ||
|
|
eac03e9856 |
+1
-2
@@ -9,13 +9,12 @@
|
||||
*.bin
|
||||
*.idea
|
||||
*.csv
|
||||
cache
|
||||
build
|
||||
dist
|
||||
dev
|
||||
scepter.egg-info
|
||||
.readthedocs.yml
|
||||
1.9
|
||||
#MANIFEST.in
|
||||
*resources
|
||||
*.ipynb_checkpoints*
|
||||
*.vscode
|
||||
|
||||
@@ -14,12 +14,17 @@ 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.11]: We're excited to announce the upcoming release of the [ACE-0.6b-1024px](https://huggingface.co/scepter-studio/ACE-0.6B-1024px) model,
|
||||
which significantly enhances image generation quality compared with [ACE-0.6b-512px](https://huggingface.co/scepter-studio/ACE-0.6B-512px). The detailed documents can be found at [ACE repo](https://github.com/ali-vilab/ACE.git).
|
||||
At the same time, based on the editing results of ACE, combined with the powerful text-to-image capabilities of the [FLUX-dev](https://huggingface.co/black-forest-labs/FLUX.1-dev) model through SDEdit as an image quality refiner, the quality of image editing can be further enhanced.
|
||||
- [🔥2024.11]: Supports video files, video annotation, caption translation in data management, and inference & training of the [CogVideoX](https://arxiv.org/abs/2408.06072).
|
||||
- [2024.10]: We are pleased to announce the release of the code for [ACE](https://arxiv.org/abs/2410.00086), supporting Customized Training / Comfy UI Workflow / gradio-based ChatBot Interface.
|
||||
- [2024.10]: Support for inference and tuning with [FLUX](https://huggingface.co/black-forest-labs/FLUX.1-dev), as well as for building [ComfyUI](https://github.com/comfyanonymous/ComfyUI) workflows using this framework.
|
||||
- [🔥2024.09]: We introduce **ACE**, an **A**ll-round **C**reator and **E**ditor adept at executing a diverse array of image editing tasks tailored to your specifications. Built upon the cutting-edge Diffusion Transformer architecture, ACE has been extensively trained on a comprehensive dataset to seamlessly interpret and execute any natural language instruction. For further information, please consult the [project page](https://ali-vilab.github.io/ace-page/).
|
||||
- [2024.09]: We introduce **ACE**, an **A**ll-round **C**reator and **E**ditor adept at executing a diverse array of image editing tasks tailored to your specifications. Built upon the cutting-edge Diffusion Transformer architecture, ACE has been extensively trained on a comprehensive dataset to seamlessly interpret and execute any natural language instruction. For further information, please consult the [project page](https://ali-vilab.github.io/ace-page/).
|
||||
- [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`.
|
||||
@@ -31,16 +36,93 @@ SCEPTER offers 3 core components:
|
||||
- [2023.12]: We release [🪄SCEPTER](https://github.com/modelscope/scepter/) library.
|
||||
|
||||
|
||||
|
||||
|
||||
## 🪄ACE
|
||||
|
||||
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.
|
||||
|
||||
[](https://ali-vilab.github.io/ace-page/)
|
||||
|
||||
### ACE Models
|
||||
| **Model** | **Status** |
|
||||
|:----------------:|:---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------:|
|
||||
| ACE-0.6B-512px | [](https://huggingface.co/spaces/scepter-studio/ACE-Chat)<br>[](https://www.modelscope.cn/models/iic/ACE-0.6B-512px) [](https://huggingface.co/scepter-studio/ACE-0.6B-512px) |
|
||||
| ACE-0.6B-1024px | [](https://huggingface.co/spaces/scepter-studio/ACE-Refiner-Chat)<br>[](https://www.modelscope.cn/models/iic/ACE-0.6B-1024px) [](https://huggingface.co/scepter-studio/ACE-0.6B-1024px) | |
|
||||
| ACE-12B-FLUX-dev | Coming Soon |
|
||||
### ACE Training
|
||||
|
||||
We offer a demonstration training YAML that enables the end-to-end training of ACE using a toy dataset. For a comprehensive overview of the hyperparameter configurations, please consult `scepter/methods/edit/dit_ace_0.6b_512.yaml`.
|
||||
|
||||
#### 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
|
||||
# ACE-0.6B-512px
|
||||
PYTHONPATH=. python scepter/tools/run_train.py --cfg scepter/methods/edit/dit_ace_0.6b_512.yaml
|
||||
# ACE-0.6B-1024px
|
||||
PYTHONPATH=. python scepter/tools/run_train.py --cfg scepter/methods/edit/dit_ace_0.6b_1024.yaml
|
||||
```
|
||||
|
||||
### ACE Chat Bot
|
||||
|
||||
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 --tab chatbot
|
||||
```
|
||||
Upon starting, you will find a "ChatBot" tab within the Gradio application, which serves as a chat-based interface to handle any requests related to image editing or generation.
|
||||
|
||||
### ACE ComfyUI Workflow
|
||||
|
||||

|
||||
|
||||
<table><tbody>
|
||||
<tr>
|
||||
<th align="center" colspan="4">ACE Workflow Examples</th>
|
||||
</tr>
|
||||
<tr>
|
||||
<th align="center" colspan="1">Control</th>
|
||||
<th align="center" colspan="1">Semantic</th>
|
||||
<th align="center" colspan="1">Element</th>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>
|
||||
<a href="https://github.com/ali-vilab/ace-page/raw/main/assets/comfyui/ace_control.png" target="_blank">
|
||||
<img src="https://github.com/ali-vilab/ace-page/raw/main/assets/comfyui/ace_control.png" width="200">
|
||||
</a>
|
||||
</td>
|
||||
<td>
|
||||
<a href="https://github.com/ali-vilab/ace-page/raw/main/assets/comfyui/ace_semantic.png" target="_blank">
|
||||
<img src="https://github.com/ali-vilab/ace-page/raw/main/assets/comfyui/ace_semantic.png" width="200">
|
||||
</a>
|
||||
</td>
|
||||
<td>
|
||||
<a href="https://github.com/ali-vilab/ace-page/raw/main/assets/comfyui/ace_element.png" target="_blank">
|
||||
<img src="https://github.com/ali-vilab/ace-page/raw/main/assets/comfyui/ace_element.png" width="200">
|
||||
</a>
|
||||
</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
## 🖼 Gallery for Recent Works
|
||||
|
||||
### <img src="https://github.com/ali-vilab/ace-page/raw/main/static/images/logo.png?raw=true" height=20> <img src="https://github.com/ali-vilab/ace-page/raw/main/static/images/icon.png?raw=true" height=20>
|
||||
|
||||
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.
|
||||
|
||||
[](https://ali-vilab.github.io/ace-page/)
|
||||
|
||||
### FLUX Tuners
|
||||
|
||||
<table><tbody>
|
||||
@@ -154,7 +236,7 @@ pip install scepter
|
||||
| Controllable Image Synthesis | [🌟SCEdit(CVPR24)](docs/en/tasks/scedit.md) | [](https://arxiv.org/abs/2312.11392) [](https://scedit.github.io/) |
|
||||
| Image Editing | [🌟LAR-Gen](docs/en/tasks/largen.md) | [](https://arxiv.org/abs/2403.19534) [](https://ali-vilab.github.io/largen-page/) |
|
||||
| Image Editing | [🌟StyleBooth](docs/en/tasks/stylebooth.md) | [](https://arxiv.org/abs/2404.12154) [](https://ali-vilab.github.io/stylebooth-page/) |
|
||||
| Image Generation and Editing | [🌟ACE](https://ali-vilab.github.io/ace-page/) | [](https://arxiv.org/abs/2410.00086) [](https://ali-vilab.github.io/ace-page/) |
|
||||
| Image Generation and Editing | [🌟ACE](https://ali-vilab.github.io/ace-page/) | [](https://arxiv.org/abs/2410.00086) [](https://ali-vilab.github.io/ace-page/) [](https://huggingface.co/spaces/scepter-studio/ACE-Chat) <br> [](https://www.modelscope.cn/models/iic/ACE-0.6B-512px) [](https://huggingface.co/scepter-studio/ACE-0.6B-512px) |
|
||||
|
||||
|
||||
## 🖥️ SCEPTER Studio
|
||||
@@ -191,18 +273,20 @@ We deploy a work studio on Modelscope that includes only the inference tab, plea
|
||||
|
||||
## ⚙️️ ComfyUI Workflow
|
||||
|
||||
### Launch
|
||||
We support the use of all models in the ComfyUI Workflow through the following methods:
|
||||
|
||||
Manually install by moving custom_nodes to ComfyUI.
|
||||
1) Automatic installation directly via the ComfyUI Manager by searching for the **ComfyUI-Scepter** node.
|
||||
2) Manually install by moving custom_nodes from Scepter to ComfyUI.
|
||||
```shell
|
||||
git clone https://github.com/modelscope/scepter.git
|
||||
cd path/to/scepter
|
||||
pip install -e .
|
||||
cp -r path/to/scepter/workflow/ path/to/ComfyUI/custom_nodes/ComfyUI-Scepter
|
||||
cd path/to/ComfyUI
|
||||
python main.py
|
||||
```
|
||||
Alternatively, we will support ComfyUI Manager shortly.
|
||||
|
||||
**Note**: You can use the nodes by dragging the sample images into ComfyUI. Additionally, our nodes can automatically pull models from ModelScope or HuggingFace by selecting the *model_source* field, or you can place the already downloaded models in a local path.
|
||||
|
||||
## 🔍 Learn More
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ albumentations
|
||||
beautifulsoup4
|
||||
bezier
|
||||
einops
|
||||
modelscope
|
||||
modelscope[framework]
|
||||
ms-swift
|
||||
numpy
|
||||
open_clip_torch
|
||||
@@ -12,6 +12,7 @@ oss2>=2.15.0
|
||||
pycocotools
|
||||
pyyaml>=5.3.1
|
||||
scikit-image
|
||||
scikit-learn
|
||||
sentencepiece
|
||||
torchsde
|
||||
transformers
|
||||
scikit-learn
|
||||
@@ -1,4 +1,5 @@
|
||||
git+https://github.com/cocodataset/panopticapi.git
|
||||
torch==2.0.1
|
||||
torchvision==0.15.2
|
||||
xformers==0.0.21
|
||||
torch==2.4.1
|
||||
torchvision==.19.1
|
||||
flash-attn==2.5.8
|
||||
xformers==0.0.28
|
||||
@@ -1,5 +1,6 @@
|
||||
bitsandbytes
|
||||
gradio
|
||||
gradio_imageslider
|
||||
imagehash
|
||||
psutil
|
||||
tiktoken
|
||||
|
||||
@@ -0,0 +1,161 @@
|
||||
ENV:
|
||||
BACKEND: nccl
|
||||
SEED: 2024
|
||||
#
|
||||
SOLVER:
|
||||
NAME: ACESolver
|
||||
RESUME_FROM:
|
||||
LOAD_MODEL_ONLY: True
|
||||
USE_FSDP: False
|
||||
SHARDING_STRATEGY:
|
||||
USE_AMP: True
|
||||
DTYPE: float16
|
||||
CHANNELS_LAST: True
|
||||
MAX_STEPS: 500
|
||||
MAX_EPOCHS: -1
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 50
|
||||
RESCALE_LR: False
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/ace_0.6b_1024
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
- NAME: "HuggingfaceFs"
|
||||
TEMP_DIR: ./cache/cache_data
|
||||
- NAME: "LocalFs"
|
||||
TEMP_DIR: ./cache/cache_data
|
||||
- NAME: "ModelscopeFs"
|
||||
TEMP_DIR: ./cache/cache_data
|
||||
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusionACE
|
||||
PRETRAINED_MODEL:
|
||||
IGNORE_KEYS: [ ]
|
||||
SCALE_FACTOR: 0.18215
|
||||
SIZE_FACTOR: 8
|
||||
DECODER_BIAS: 0.5
|
||||
DEFAULT_N_PROMPT:
|
||||
USE_EMA: True
|
||||
EVAL_EMA: False
|
||||
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
|
||||
USE_TEXT_POS_EMBEDDINGS: True
|
||||
#
|
||||
DIFFUSION:
|
||||
NAME: BaseDiffusion
|
||||
PREDICTION_TYPE: eps
|
||||
MIN_SNR_GAMMA:
|
||||
NOISE_SCHEDULER:
|
||||
NAME: LinearScheduler
|
||||
NUM_TIMESTEPS: 1000
|
||||
BETA_MIN: 0.0001
|
||||
BETA_MAX: 0.02
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: ACE
|
||||
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/dit/ace_0.6b_1024px.pth
|
||||
IGNORE_KEYS: [ ]
|
||||
PATCH_SIZE: 2
|
||||
IN_CHANNELS: 4
|
||||
HIDDEN_SIZE: 1152
|
||||
DEPTH: 28
|
||||
NUM_HEADS: 16
|
||||
MLP_RATIO: 4.0
|
||||
PRED_SIGMA: True
|
||||
DROP_PATH: 0.0
|
||||
WINDOW_DIZE: 0
|
||||
Y_CHANNELS: 4096
|
||||
MAX_SEQ_LEN: 4096
|
||||
QK_NORM: True
|
||||
USE_GRAD_CHECKPOINT: True
|
||||
ATTENTION_BACKEND: flash_attn
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKL
|
||||
EMBED_DIM: 4
|
||||
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/vae/vae.bin
|
||||
IGNORE_KEYS: []
|
||||
#
|
||||
ENCODER:
|
||||
NAME: Encoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 4
|
||||
DOUBLE_Z: True
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
#
|
||||
DECODER:
|
||||
NAME: Decoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 4
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
GIVE_PRE_END: False
|
||||
TANH_OUT: False
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: T5EmbedderHF
|
||||
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/text_encoder/t5-v1_1-xxl/
|
||||
TOKENIZER_PATH: ms://iic/ACE-0.6B-1024px@models/tokenizer/t5-v1_1-xxl
|
||||
LENGTH: 120
|
||||
T5_DTYPE: bfloat16
|
||||
ADDED_IDENTIFIER: [ '{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
|
||||
CLEAN: whitespace
|
||||
USE_GRAD: False
|
||||
LOSS:
|
||||
NAME: ReconstructLoss
|
||||
LOSS_TYPE: l2
|
||||
#
|
||||
SAMPLE_ARGS:
|
||||
SAMPLER: ddim
|
||||
SAMPLE_STEPS: 20
|
||||
GUIDE_SCALE: 4.5
|
||||
GUIDE_RESCALE: 0.5
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 1e-7
|
||||
EPS: 1e-10
|
||||
WEIGHT_DECAY: 5e-4
|
||||
#
|
||||
TRAIN_DATA:
|
||||
NAME: ImageTextPairMSDatasetForACE
|
||||
MODE: train
|
||||
MS_DATASET_NAME: cache/datasets/hed_pair
|
||||
MS_DATASET_NAMESPACE: ""
|
||||
MS_DATASET_SPLIT: "train"
|
||||
MS_DATASET_SUBNAME: ""
|
||||
PROMPT_PREFIX: ""
|
||||
REPLACE_STYLE: False
|
||||
MAX_SEQ_LEN: 4096
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 1
|
||||
SAMPLER:
|
||||
NAME: LoopSampler
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
-
|
||||
NAME: BackwardHook
|
||||
PRIORITY: 0
|
||||
-
|
||||
NAME: LogHook
|
||||
LOG_INTERVAL: 50
|
||||
-
|
||||
NAME: CheckpointHook
|
||||
INTERVAL: 100
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
@@ -0,0 +1,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,235 @@
|
||||
ENV:
|
||||
BACKEND: nccl
|
||||
SEED: 42
|
||||
TENSOR_PARALLEL_SIZE: 1
|
||||
PIPELINE_PARALLEL_SIZE: 1
|
||||
SYS_ENVS:
|
||||
TORCH_CUDNN_V8_API_ENABLED: '1'
|
||||
TOKENIZERS_PARALLELISM: 'false'
|
||||
TF_CPP_MIN_LOG_LEVEL: '3'
|
||||
PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True'
|
||||
#
|
||||
SOLVER:
|
||||
NAME: LatentDiffusionVideoSolver
|
||||
MAX_STEPS: 2000
|
||||
USE_AMP: True
|
||||
DTYPE: bfloat16
|
||||
USE_FAIRSCALE: False
|
||||
USE_FSDP: True
|
||||
LOAD_MODEL_ONLY: False
|
||||
RESUME_FROM:
|
||||
WORK_DIR: ./cache/save_data/dit_cogvideox_2b_lora
|
||||
LOG_FILE: std_log.txt
|
||||
EVAL_INTERVAL: 100
|
||||
LOG_TRAIN_NUM: 4
|
||||
ENABLE_GRADSCALER: False
|
||||
USE_SCALER: False
|
||||
FPS: 8
|
||||
SHARDING_STRATEGY: full_shard
|
||||
FSDP_REDUCE_DTYPE: float32
|
||||
FSDP_BUFFER_DTYPE: float32
|
||||
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model']
|
||||
SAVE_MODULES: [ 'model', 'cond_stage_model.model']
|
||||
TRAIN_MODULES: ['model']
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TUNER:
|
||||
- NAME: SwiftLoRA
|
||||
R: 64
|
||||
LORA_ALPHA: 64
|
||||
LORA_DROPOUT: 0.0
|
||||
BIAS: "none"
|
||||
TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$"
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusionCogVideoX
|
||||
PRETRAINED_MODEL:
|
||||
PARAMETERIZATION: v
|
||||
TIMESTEPS: 1000
|
||||
MIN_SNR_GAMMA: 3.0
|
||||
ZERO_TERMINAL_SNR: True
|
||||
SCALE_FACTOR_SPATIAL: 8
|
||||
SCALE_FACTOR_TEMPORAL: 4
|
||||
SCALING_FACTOR_IMAGE: 1.15258426
|
||||
IGNORE_KEYS: [ ]
|
||||
DEFAULT_N_PROMPT:
|
||||
USE_EMA: False
|
||||
EVAL_EMA: False
|
||||
DIFFUSION:
|
||||
NAME: BaseDiffusion
|
||||
PREDICTION_TYPE: v
|
||||
NOISE_SCHEDULER:
|
||||
NAME: ScaledLinearScheduler
|
||||
BETA_MIN: 0.00085
|
||||
BETA_MAX: 0.012
|
||||
SNR_SHIFT_SCALE: 3.0
|
||||
RESCALE_BETAS_ZERO_SNR: True
|
||||
DIFFUSION_SAMPLERS:
|
||||
NAME: DDIMSampler
|
||||
DISCRETIZATION_TYPE: trailing
|
||||
ETA: 0.0
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: CogVideoXTransformer3DModel
|
||||
DTYPE: bfloat16
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@transformer/diffusion_pytorch_model.safetensors
|
||||
NUM_ATTENTION_HEADS: 30
|
||||
ATTENTION_HEAD_DIM: 64
|
||||
IN_CHANNELS: 16
|
||||
OUT_CHANNELS: 16
|
||||
FLIP_SIN_TO_COS: True
|
||||
FREQ_SHIFT: 0
|
||||
TIME_EMBED_DIM: 512
|
||||
TEXT_EMBED_DIM: 4096
|
||||
NUM_LAYERS: 30
|
||||
DROPOUT: 0.0
|
||||
ATTENTION_BIAS: True
|
||||
SAMPLE_WIDTH: 90
|
||||
SAMPLE_HEIGHT: 60
|
||||
SAMPLE_FRAMES: 49
|
||||
PATCH_SIZE: 2
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
MAX_TEXT_SEQ_LENGTH: 226
|
||||
ACTIVATION_FN: "gelu-approximate"
|
||||
TIMESTEP_ACTIVATION_FN: "silu"
|
||||
NORM_ELEMENTWISE_AFFINE: True
|
||||
NORM_EPS: 1e-5
|
||||
SPATIAL_INTERPOLATION_SCALE: 1.875
|
||||
TEMPORAL_INTERPOLATION_SCALE: 1.0
|
||||
USE_ROTARY_POSITIONAL_EMBEDDINGS: False
|
||||
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
|
||||
GRADIENT_CHECKPOINTING: False
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKLCogVideoX
|
||||
DTYPE: bfloat16
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@vae/diffusion_pytorch_model.safetensors
|
||||
SAMPLE_HEIGHT: 480
|
||||
SAMPLE_WIDTH: 720
|
||||
USE_QUANT_CONV: False
|
||||
USE_POST_QUANT_CONV: False
|
||||
USE_SLICING: True
|
||||
USE_TILING: True
|
||||
GRADIENT_CHECKPOINTING: False
|
||||
ENCODER:
|
||||
NAME: CogVideoXEncoder3D
|
||||
IN_CHANNELS: 3
|
||||
OUT_CHANNELS: 16
|
||||
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: False
|
||||
DECODER:
|
||||
NAME: CogVideoXDecoder3D
|
||||
IN_CHANNELS: 16
|
||||
OUT_CHANNELS: 3
|
||||
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: False
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: T5EmbedderHF
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
|
||||
LENGTH: 226
|
||||
CLEAN:
|
||||
USE_GRAD: False
|
||||
#
|
||||
LOSS:
|
||||
NAME: ReconstructLoss
|
||||
LOSS_TYPE: l2
|
||||
#
|
||||
SAMPLE_ARGS:
|
||||
SAMPLER: ddim
|
||||
SAMPLE_STEPS: 50
|
||||
SEED: 42
|
||||
GUIDE_SCALE: 6.0
|
||||
GUIDE_RESCALE: 0.0
|
||||
NUM_FRAMES: 49
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: Adam
|
||||
LEARNING_RATE: 1e-3
|
||||
BETAS: [ 0.9, 0.95 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 0.0
|
||||
AMSGRAD: False
|
||||
#
|
||||
# LR_SCHEDULER:
|
||||
# NAME: StepAnnealingLR
|
||||
# WARMUP_STEPS: 200
|
||||
# TOTAL_STEPS: 2000
|
||||
# DECAY_MODE: 'cosine'
|
||||
#
|
||||
TRAIN_DATA:
|
||||
NAME: VideoGenDataset
|
||||
MODE: train
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
PROMPT_PREFIX: 'DISNEY '
|
||||
SAMPLER:
|
||||
NAME: MixtureOfSamplers
|
||||
SUB_SAMPLERS:
|
||||
- NAME: MultiLevelBatchSampler
|
||||
PROB: 1.0
|
||||
FIELDS: [ "video_path", "prompt" ]
|
||||
DELIMITER: '#;#'
|
||||
PATH_PREFIX: cache/datasets/Disney-VideoGeneration-Dataset/
|
||||
INDEX_FILE: cache/datasets/Disney-VideoGeneration-Dataset/index.jsonl
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'video', "prompt" ]
|
||||
META_KEYS: [ ]
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: Text2ImageDataset
|
||||
MODE: eval
|
||||
PROMPT_FILE:
|
||||
PROMPT_DATA: [ "A girl riding a bike.", "A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance." ]
|
||||
IMAGE_SIZE: [ 480, 720 ]
|
||||
FIELDS: [ "prompt" ]
|
||||
DELIMITER: '#;#'
|
||||
PROMPT_PREFIX: 'DISNEY '
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
USE_NUM: 8
|
||||
NUM_WORKERS: 4
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'index', 'prompt' ]
|
||||
META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
- NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
PRIORITY: 0
|
||||
- NAME: BackwardHook
|
||||
PRIORITY: 10
|
||||
- NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
PRIORITY: 20
|
||||
- NAME: CheckpointHook
|
||||
INTERVAL: 1000
|
||||
PRIORITY: 40
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
- NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
PRIORITY: 0
|
||||
@@ -0,0 +1,266 @@
|
||||
ENV:
|
||||
BACKEND: nccl
|
||||
SEED: 42
|
||||
TENSOR_PARALLEL_SIZE: 1
|
||||
PIPELINE_PARALLEL_SIZE: 1
|
||||
SYS_ENVS:
|
||||
TORCH_CUDNN_V8_API_ENABLED: '1'
|
||||
TOKENIZERS_PARALLELISM: 'false'
|
||||
TF_CPP_MIN_LOG_LEVEL: '3'
|
||||
PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True'
|
||||
#
|
||||
SOLVER:
|
||||
NAME: LatentDiffusionVideoSolver
|
||||
MAX_STEPS: 2000
|
||||
USE_AMP: True
|
||||
DTYPE: bfloat16
|
||||
USE_FAIRSCALE: False
|
||||
USE_FSDP: True
|
||||
LOAD_MODEL_ONLY: False
|
||||
ENABLE_GRADSCALER: False
|
||||
USE_SCALER: False
|
||||
RESUME_FROM:
|
||||
WORK_DIR: ./cache/save_data/dit_cogvideox_5b_i2v_lora
|
||||
LOG_FILE: std_log.txt
|
||||
EVAL_INTERVAL: 100
|
||||
LOG_TRAIN_NUM: 4
|
||||
FPS: 8
|
||||
SHARDING_STRATEGY: full_shard
|
||||
FSDP_REDUCE_DTYPE: float32
|
||||
FSDP_BUFFER_DTYPE: float32
|
||||
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model']
|
||||
SAVE_MODULES: [ 'model', 'cond_stage_model.model']
|
||||
TRAIN_MODULES: ['model']
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TUNER:
|
||||
- NAME: SwiftLoRA
|
||||
R: 64
|
||||
LORA_ALPHA: 64
|
||||
LORA_DROPOUT: 0.0
|
||||
BIAS: "none"
|
||||
TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$"
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusionCogVideoX
|
||||
PRETRAINED_MODEL:
|
||||
PARAMETERIZATION: v
|
||||
TIMESTEPS: 1000
|
||||
MIN_SNR_GAMMA: 3.0
|
||||
ZERO_TERMINAL_SNR: True
|
||||
SCALE_FACTOR_SPATIAL: 8
|
||||
SCALE_FACTOR_TEMPORAL: 4
|
||||
SCALING_FACTOR_IMAGE: 0.7 # 5b diff
|
||||
NOISED_IMAGE_DROPOUT: 0.05
|
||||
IGNORE_KEYS: [ ]
|
||||
DEFAULT_N_PROMPT:
|
||||
USE_EMA: False
|
||||
EVAL_EMA: False
|
||||
DIFFUSION:
|
||||
NAME: BaseDiffusion
|
||||
PREDICTION_TYPE: v
|
||||
NOISE_SCHEDULER:
|
||||
NAME: ScaledLinearScheduler
|
||||
BETA_MIN: 0.00085
|
||||
BETA_MAX: 0.012
|
||||
SNR_SHIFT_SCALE: 1.0 # 5b diff
|
||||
RESCALE_BETAS_ZERO_SNR: True
|
||||
DIFFUSION_SAMPLERS:
|
||||
NAME: DDIMSampler
|
||||
DISCRETIZATION_TYPE: trailing
|
||||
ETA: 0.0
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: CogVideoXTransformer3DModel
|
||||
DTYPE: bfloat16
|
||||
PRETRAINED_MODEL: # 5b-I2V diff
|
||||
- ms://AI-ModelScope/CogVideoX-5b-I2V@transformer/diffusion_pytorch_model-00001-of-00003.safetensors
|
||||
- ms://AI-ModelScope/CogVideoX-5b-I2V@transformer/diffusion_pytorch_model-00002-of-00003.safetensors
|
||||
- ms://AI-ModelScope/CogVideoX-5b-I2V@transformer/diffusion_pytorch_model-00003-of-00003.safetensors
|
||||
NUM_ATTENTION_HEADS: 48 # 5b diff
|
||||
ATTENTION_HEAD_DIM: 64
|
||||
IN_CHANNELS: 32 # 5b-I2V diff
|
||||
LATENT_CHANNELS: 16
|
||||
OUT_CHANNELS: 16
|
||||
FLIP_SIN_TO_COS: True
|
||||
FREQ_SHIFT: 0
|
||||
TIME_EMBED_DIM: 512
|
||||
TEXT_EMBED_DIM: 4096
|
||||
NUM_LAYERS: 42 # 5b diff
|
||||
DROPOUT: 0.0
|
||||
ATTENTION_BIAS: True
|
||||
SAMPLE_WIDTH: 90
|
||||
SAMPLE_HEIGHT: 60
|
||||
SAMPLE_FRAMES: 49
|
||||
PATCH_SIZE: 2
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
MAX_TEXT_SEQ_LENGTH: 226
|
||||
ACTIVATION_FN: "gelu-approximate"
|
||||
TIMESTEP_ACTIVATION_FN: "silu"
|
||||
NORM_ELEMENTWISE_AFFINE: True
|
||||
NORM_EPS: 1e-5
|
||||
SPATIAL_INTERPOLATION_SCALE: 1.875
|
||||
TEMPORAL_INTERPOLATION_SCALE: 1.0
|
||||
USE_ROTARY_POSITIONAL_EMBEDDINGS: True # 5b diff
|
||||
USE_LEARNED_POSITIONAL_EMBEDDINGS: True # 5b-I2V diff
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKLCogVideoX
|
||||
DTYPE: bfloat16
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b-I2V@vae/diffusion_pytorch_model.safetensors # 5b diff
|
||||
SAMPLE_HEIGHT: 480
|
||||
SAMPLE_WIDTH: 720
|
||||
USE_QUANT_CONV: False
|
||||
USE_POST_QUANT_CONV: False
|
||||
USE_SLICING: True
|
||||
USE_TILING: True
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
ENCODER:
|
||||
NAME: CogVideoXEncoder3D
|
||||
IN_CHANNELS: 3
|
||||
OUT_CHANNELS: 16
|
||||
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
DECODER:
|
||||
NAME: CogVideoXDecoder3D
|
||||
IN_CHANNELS: 16
|
||||
OUT_CHANNELS: 3
|
||||
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: T5EmbedderHF
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
|
||||
LENGTH: 226
|
||||
CLEAN:
|
||||
USE_GRAD: False
|
||||
#
|
||||
LOSS:
|
||||
NAME: ReconstructLoss
|
||||
LOSS_TYPE: l2
|
||||
#
|
||||
SAMPLE_ARGS:
|
||||
SAMPLER: ddim
|
||||
SAMPLE_STEPS: 50
|
||||
SEED: 42
|
||||
GUIDE_SCALE: 6.0
|
||||
GUIDE_RESCALE: 0.0
|
||||
NUM_FRAMES: 49
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: Adam
|
||||
LEARNING_RATE: 1e-3
|
||||
BETAS: [ 0.9, 0.95 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 0.0
|
||||
AMSGRAD: False
|
||||
#
|
||||
# LR_SCHEDULER:
|
||||
# NAME: StepAnnealingLR
|
||||
# WARMUP_STEPS: 200
|
||||
# TOTAL_STEPS: 2000
|
||||
# DECAY_MODE: 'cosine'
|
||||
#
|
||||
TRAIN_DATA:
|
||||
NAME: VideoGenDataset
|
||||
MODE: train
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 0
|
||||
PROMPT_PREFIX: 'DISNEY '
|
||||
DATA_TYPE: 'i2v'
|
||||
SAMPLER:
|
||||
NAME: MixtureOfSamplers
|
||||
SUB_SAMPLERS:
|
||||
- NAME: MultiLevelBatchSampler
|
||||
PROB: 1.0
|
||||
FIELDS: [ "video_path", "prompt" ]
|
||||
DELIMITER: '#;#'
|
||||
PATH_PREFIX: cache/datasets/Disney-VideoGeneration-Dataset/
|
||||
INDEX_FILE: cache/datasets/Disney-VideoGeneration-Dataset/index.jsonl
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ "video", "image", "prompt" ]
|
||||
META_KEYS: [ ]
|
||||
#
|
||||
# EVAL_DATA:
|
||||
# NAME: Text2ImageDataset
|
||||
# MODE: eval
|
||||
# PROMPT_FILE:
|
||||
# PROMPT_DATA: [ "A cat running.#;#asset/images/edit_tuner/cat_512.jpg" ]
|
||||
# FIELDS: [ "prompt", "img_path" ]
|
||||
# DELIMITER: '#;#'
|
||||
# PROMPT_PREFIX: ''
|
||||
# PIN_MEMORY: True
|
||||
# BATCH_SIZE: 1
|
||||
# USE_NUM: 8
|
||||
# NUM_WORKERS: 0
|
||||
# IMAGE_SIZE: [ 480, 720 ]
|
||||
# TRANSFORMS:
|
||||
# - NAME: LoadImageFromFileList
|
||||
# FILE_KEYS: [ 'img_path' ]
|
||||
# RGB_ORDER: RGB
|
||||
# BACKEND: pillow
|
||||
# - NAME: FlexibleResize
|
||||
# INTERPOLATION: bilinear
|
||||
# SIZE: [ 480, 720 ]
|
||||
# INPUT_KEY: [ 'img' ]
|
||||
# OUTPUT_KEY: [ 'img' ]
|
||||
# BACKEND: pillow
|
||||
# - NAME: FlexibleCenterCrop
|
||||
# SIZE: [ 480, 720 ]
|
||||
# INPUT_KEY: [ 'img' ]
|
||||
# OUTPUT_KEY: [ 'img' ]
|
||||
# BACKEND: pillow
|
||||
# - NAME: ImageToTensor
|
||||
# INPUT_KEY: [ 'img' ]
|
||||
# OUTPUT_KEY: [ 'img' ]
|
||||
# BACKEND: pillow
|
||||
# - NAME: Normalize
|
||||
# MEAN: [ 0.5, 0.5, 0.5 ]
|
||||
# STD: [ 0.5, 0.5, 0.5 ]
|
||||
# INPUT_KEY: [ 'img' ]
|
||||
# OUTPUT_KEY: [ 'image' ]
|
||||
# BACKEND: torchvision
|
||||
# - NAME: Select
|
||||
# KEYS: [ 'image', 'prompt' ]
|
||||
# META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
- NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
PRIORITY: 0
|
||||
- NAME: BackwardHook
|
||||
PRIORITY: 10
|
||||
- NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
PRIORITY: 20
|
||||
- NAME: CheckpointHook
|
||||
INTERVAL: 1000
|
||||
PRIORITY: 40
|
||||
#
|
||||
# EVAL_HOOKS:
|
||||
# - NAME: ProbeDataHook
|
||||
# PROB_INTERVAL: 100
|
||||
# PRIORITY: 0
|
||||
@@ -0,0 +1,273 @@
|
||||
ENV:
|
||||
BACKEND: nccl
|
||||
SEED: 42
|
||||
TENSOR_PARALLEL_SIZE: 1
|
||||
PIPELINE_PARALLEL_SIZE: 1
|
||||
SYS_ENVS:
|
||||
TORCH_CUDNN_V8_API_ENABLED: '1'
|
||||
TOKENIZERS_PARALLELISM: 'false'
|
||||
TF_CPP_MIN_LOG_LEVEL: '3'
|
||||
PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True'
|
||||
#
|
||||
SOLVER:
|
||||
NAME: LatentDiffusionVideoSolver
|
||||
MAX_STEPS: 2000
|
||||
USE_AMP: True
|
||||
DTYPE: bfloat16
|
||||
USE_FAIRSCALE: False
|
||||
USE_FSDP: True
|
||||
LOAD_MODEL_ONLY: False
|
||||
ENABLE_GRADSCALER: False
|
||||
USE_SCALER: False
|
||||
RESUME_FROM:
|
||||
WORK_DIR: ./cache/save_data/dit_cogvideox_5b_lora
|
||||
LOG_FILE: std_log.txt
|
||||
EVAL_INTERVAL: 100
|
||||
LOG_TRAIN_NUM: 4
|
||||
FPS: 8
|
||||
SHARDING_STRATEGY: full_shard
|
||||
FSDP_REDUCE_DTYPE: float32
|
||||
FSDP_BUFFER_DTYPE: float32
|
||||
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model']
|
||||
SAVE_MODULES: [ 'model', 'cond_stage_model.model']
|
||||
TRAIN_MODULES: ['model']
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TUNER:
|
||||
- NAME: SwiftLoRA
|
||||
R: 64
|
||||
LORA_ALPHA: 64
|
||||
LORA_DROPOUT: 0.0
|
||||
BIAS: "none"
|
||||
TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$"
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusionCogVideoX
|
||||
PRETRAINED_MODEL:
|
||||
PARAMETERIZATION: v
|
||||
TIMESTEPS: 1000
|
||||
MIN_SNR_GAMMA: 3.0
|
||||
ZERO_TERMINAL_SNR: True
|
||||
SCALE_FACTOR_SPATIAL: 8
|
||||
SCALE_FACTOR_TEMPORAL: 4
|
||||
SCALING_FACTOR_IMAGE: 0.7 # 5b diff
|
||||
IGNORE_KEYS: [ ]
|
||||
DEFAULT_N_PROMPT:
|
||||
USE_EMA: False
|
||||
EVAL_EMA: False
|
||||
DIFFUSION:
|
||||
NAME: BaseDiffusion
|
||||
PREDICTION_TYPE: v
|
||||
NOISE_SCHEDULER:
|
||||
NAME: ScaledLinearScheduler
|
||||
BETA_MIN: 0.00085
|
||||
BETA_MAX: 0.012
|
||||
SNR_SHIFT_SCALE: 1.0 # 5b diff
|
||||
RESCALE_BETAS_ZERO_SNR: True
|
||||
DIFFUSION_SAMPLERS:
|
||||
NAME: DDIMSampler
|
||||
DISCRETIZATION_TYPE: trailing
|
||||
ETA: 0.0
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: CogVideoXTransformer3DModel
|
||||
DTYPE: bfloat16
|
||||
PRETRAINED_MODEL: # 5b diff
|
||||
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00001-of-00002.safetensors
|
||||
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00002-of-00002.safetensors
|
||||
NUM_ATTENTION_HEADS: 48 # 5b diff
|
||||
ATTENTION_HEAD_DIM: 64
|
||||
IN_CHANNELS: 16
|
||||
OUT_CHANNELS: 16
|
||||
FLIP_SIN_TO_COS: True
|
||||
FREQ_SHIFT: 0
|
||||
TIME_EMBED_DIM: 512
|
||||
TEXT_EMBED_DIM: 4096
|
||||
NUM_LAYERS: 42 # 5b diff
|
||||
DROPOUT: 0.0
|
||||
ATTENTION_BIAS: True
|
||||
SAMPLE_WIDTH: 90
|
||||
SAMPLE_HEIGHT: 60
|
||||
SAMPLE_FRAMES: 49
|
||||
PATCH_SIZE: 2
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
MAX_TEXT_SEQ_LENGTH: 226
|
||||
ACTIVATION_FN: "gelu-approximate"
|
||||
TIMESTEP_ACTIVATION_FN: "silu"
|
||||
NORM_ELEMENTWISE_AFFINE: True
|
||||
NORM_EPS: 1e-5
|
||||
SPATIAL_INTERPOLATION_SCALE: 1.875
|
||||
TEMPORAL_INTERPOLATION_SCALE: 1.0
|
||||
USE_ROTARY_POSITIONAL_EMBEDDINGS: True # 5b diff
|
||||
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKLCogVideoX
|
||||
DTYPE: bfloat16
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b@vae/diffusion_pytorch_model.safetensors # 5b diff
|
||||
SAMPLE_HEIGHT: 480
|
||||
SAMPLE_WIDTH: 720
|
||||
USE_QUANT_CONV: False
|
||||
USE_POST_QUANT_CONV: False
|
||||
USE_SLICING: True
|
||||
USE_TILING: True
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
ENCODER:
|
||||
NAME: CogVideoXEncoder3D
|
||||
IN_CHANNELS: 3
|
||||
OUT_CHANNELS: 16
|
||||
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
DECODER:
|
||||
NAME: CogVideoXDecoder3D
|
||||
IN_CHANNELS: 16
|
||||
OUT_CHANNELS: 3
|
||||
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: T5EmbedderHF
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
|
||||
LENGTH: 226
|
||||
CLEAN:
|
||||
USE_GRAD: False
|
||||
#
|
||||
LOSS:
|
||||
NAME: ReconstructLoss
|
||||
LOSS_TYPE: l2
|
||||
#
|
||||
SAMPLE_ARGS:
|
||||
SAMPLER: ddim
|
||||
SAMPLE_STEPS: 50
|
||||
SEED: 42
|
||||
GUIDE_SCALE: 6.0
|
||||
GUIDE_RESCALE: 0.0
|
||||
NUM_FRAMES: 49
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: Adam
|
||||
LEARNING_RATE: 1e-3
|
||||
BETAS: [ 0.9, 0.95 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 0.0
|
||||
AMSGRAD: False
|
||||
#
|
||||
# LR_SCHEDULER:
|
||||
# NAME: StepAnnealingLR
|
||||
# WARMUP_STEPS: 200
|
||||
# TOTAL_STEPS: 2000
|
||||
# DECAY_MODE: 'cosine'
|
||||
#
|
||||
TRAIN_DATA:
|
||||
NAME: VideoGenDatasetOTF
|
||||
MODE: train
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
PROMPT_PREFIX: 'DISNEY '
|
||||
DELIMITER: '#;#'
|
||||
FIELDS: [ 'video_path', 'prompt' ]
|
||||
PATH_PREFIX: cache/datasets/Disney-VideoGeneration-Dataset/
|
||||
DATA_FILE: cache/datasets/Disney-VideoGeneration-Dataset/index.jsonl
|
||||
SAMPLER:
|
||||
NAME: LoopSampler
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'video', 'video_latent', "prompt" ]
|
||||
META_KEYS: [ ]
|
||||
MODEL:
|
||||
NAME: AutoencoderKLCogVideoX
|
||||
DTYPE: bfloat16
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b@vae/diffusion_pytorch_model.safetensors
|
||||
SAMPLE_HEIGHT: 480
|
||||
SAMPLE_WIDTH: 720
|
||||
USE_QUANT_CONV: False
|
||||
USE_POST_QUANT_CONV: False
|
||||
USE_SLICING: True
|
||||
USE_TILING: True
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
ENCODER:
|
||||
NAME: CogVideoXEncoder3D
|
||||
IN_CHANNELS: 3
|
||||
OUT_CHANNELS: 16
|
||||
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
DECODER:
|
||||
NAME: CogVideoXDecoder3D
|
||||
IN_CHANNELS: 16
|
||||
OUT_CHANNELS: 3
|
||||
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: Text2ImageDataset
|
||||
MODE: eval
|
||||
PROMPT_FILE:
|
||||
PROMPT_DATA: [ "A girl riding a bike.", "A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance." ]
|
||||
IMAGE_SIZE: [ 480, 720 ]
|
||||
FIELDS: [ "prompt" ]
|
||||
DELIMITER: '#;#'
|
||||
PROMPT_PREFIX: 'DISNEY '
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
USE_NUM: 8
|
||||
NUM_WORKERS: 4
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'index', 'prompt' ]
|
||||
META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
- NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
PRIORITY: 0
|
||||
- NAME: BackwardHook
|
||||
PRIORITY: 10
|
||||
- NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
PRIORITY: 20
|
||||
- NAME: CheckpointHook
|
||||
INTERVAL: 1000
|
||||
PRIORITY: 40
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
- NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
PRIORITY: 0
|
||||
@@ -2,33 +2,22 @@ ENV:
|
||||
BACKEND: nccl
|
||||
SEED: 166666
|
||||
SOLVER:
|
||||
# NAME DESCRIPTION: TYPE: default: 'LatentUfitSolver'
|
||||
NAME: LatentDiffusionSolver
|
||||
# MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
|
||||
MAX_STEPS: 100000
|
||||
# USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False
|
||||
USE_AMP: True
|
||||
# DTYPE DESCRIPTION: The precision for training. TYPE: str default: 'float32'
|
||||
DTYPE: bfloat16
|
||||
# USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
|
||||
USE_FAIRSCALE: False
|
||||
# USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
|
||||
USE_FSDP: True
|
||||
# LOAD_MODEL_ONLY DESCRIPTION: Only load the model rather than the optimizer and schedule, default is False. TYPE: bool default: False
|
||||
LOAD_MODEL_ONLY: False
|
||||
# RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: ''
|
||||
ENABLE_GRADSCALER: False
|
||||
USE_SCALER: False
|
||||
RESUME_FROM:
|
||||
WORK_DIR: ./cache/save_data/dit_flux_dev_1024_lora
|
||||
LOG_FILE: std_log.txt
|
||||
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
|
||||
EVAL_INTERVAL: 100
|
||||
# LOG_TRAIN_NUM DESCRIPTION: The number samples used to log in training phase. TYPE: int default: -1
|
||||
LOG_TRAIN_NUM: 16
|
||||
# FSDP_REDUCE_DTYPE DESCRIPTION: The dtype of reduce in FSDP. TYPE: str default: 'float16'
|
||||
FSDP_REDUCE_DTYPE: float32
|
||||
# FSDP_BUFFER_DTYPE DESCRIPTION: The dtype of buffer in FSDP. TYPE: str default: 'float16'
|
||||
FSDP_BUFFER_DTYPE: float32
|
||||
# FSDP_SHARD_MODULES DESCRIPTION: The modules to be sharded in FSDP. TYPE: list default: ['model']
|
||||
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ] #
|
||||
SAVE_MODULES: [ 'model'] #
|
||||
TRAIN_MODULES: ['model']
|
||||
@@ -58,61 +47,36 @@ SOLVER:
|
||||
USE_EMA: False
|
||||
EVAL_EMA: False
|
||||
DIFFUSION:
|
||||
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
|
||||
NAME: DiffusionFluxRF
|
||||
PREDICTION_TYPE: raw
|
||||
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
|
||||
NOISE_SCHEDULER:
|
||||
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
|
||||
NAME: FlowMatchSigmaScheduler
|
||||
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
|
||||
WEIGHTING_SCHEME: logit_normal
|
||||
SHIFT: 3.0
|
||||
# LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
|
||||
LOGIT_MEAN: 0.0
|
||||
# LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
|
||||
LOGIT_STD: 1.0
|
||||
# MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
|
||||
MODE_SCALE: 1.29
|
||||
SAMPLER_SCHEDULER:
|
||||
# NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
|
||||
NAME: FlowMatchFluxShiftScheduler
|
||||
# SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
|
||||
SHIFT: False
|
||||
# SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
|
||||
SIGMOID_SCALE: 1
|
||||
# BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
|
||||
BASE_SHIFT: 0.5
|
||||
# MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
|
||||
MAX_SHIFT: 1.15
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
# NAME DESCRIPTION: TYPE: default: 'Flux'
|
||||
NAME: Flux
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@flux1-dev.safetensors
|
||||
# IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
|
||||
IN_CHANNELS: 64
|
||||
# HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
|
||||
HIDDEN_SIZE: 3072
|
||||
# NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
|
||||
NUM_HEADS: 24
|
||||
# AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
|
||||
AXES_DIM: [ 16, 56, 56 ]
|
||||
# THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
|
||||
THETA: 10000
|
||||
# VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
|
||||
VEC_IN_DIM: 768
|
||||
# GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
|
||||
GUIDANCE_EMBED: True
|
||||
# CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
|
||||
CONTEXT_IN_DIM: 4096
|
||||
# MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
|
||||
MLP_RATIO: 4.0
|
||||
# QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
|
||||
QKV_BIAS: True
|
||||
# DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
|
||||
DEPTH: 19
|
||||
# DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
|
||||
DEPTH_SINGLE_BLOCKS: 38
|
||||
USE_GRAD_CHECKPOINT: True
|
||||
|
||||
@@ -157,55 +121,34 @@ SOLVER:
|
||||
TANH_OUT: False
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
|
||||
NAME: T5PlusClipFluxEmbedder
|
||||
# T5_MODEL DESCRIPTION: TYPE: default: ''
|
||||
T5_MODEL:
|
||||
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
|
||||
NAME: HFEmbedder
|
||||
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_MODEL_CLS: T5EncoderModel
|
||||
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
|
||||
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder_2/
|
||||
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_TOKENIZER_CLS: T5Tokenizer
|
||||
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer_2/
|
||||
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
|
||||
MAX_LENGTH: 512
|
||||
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
|
||||
OUTPUT_KEY: last_hidden_state
|
||||
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
|
||||
D_TYPE: bfloat16
|
||||
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
|
||||
BATCH_INFER: False
|
||||
CLEAN: whitespace
|
||||
# CLIP_MODEL DESCRIPTION: TYPE: default: ''
|
||||
CLIP_MODEL:
|
||||
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
|
||||
NAME: HFEmbedder
|
||||
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_MODEL_CLS: CLIPTextModel
|
||||
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
|
||||
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder/
|
||||
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_TOKENIZER_CLS: CLIPTokenizer
|
||||
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer/
|
||||
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
|
||||
MAX_LENGTH: 77
|
||||
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
|
||||
OUTPUT_KEY: pooler_output
|
||||
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
|
||||
D_TYPE: bfloat16
|
||||
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
|
||||
BATCH_INFER: True
|
||||
CLEAN: whitespace
|
||||
USE_GRAD_CHECKPOINT: True
|
||||
#
|
||||
SAMPLE_ARGS:
|
||||
SAMPLE_STEPS: 50
|
||||
SAMPLER: flow_eluer
|
||||
SAMPLER: flow_euler
|
||||
SEED: 2024
|
||||
IMAGE_SIZE: [ 1024, 1024 ]
|
||||
GUIDE_SCALE: 3.5
|
||||
|
||||
@@ -2,35 +2,24 @@ ENV:
|
||||
BACKEND: nccl
|
||||
SEED: 166666
|
||||
SOLVER:
|
||||
# NAME DESCRIPTION: TYPE: default: 'LatentUfitSolver'
|
||||
NAME: LatentDiffusionSolver
|
||||
# MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
|
||||
MAX_STEPS: 100000
|
||||
# USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False
|
||||
USE_AMP: True
|
||||
# DTYPE DESCRIPTION: The precision for training. TYPE: str default: 'float32'
|
||||
DTYPE: bfloat16
|
||||
# USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
|
||||
USE_FAIRSCALE: False
|
||||
# USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
|
||||
USE_FSDP: True
|
||||
# LOAD_MODEL_ONLY DESCRIPTION: Only load the model rather than the optimizer and schedule, default is False. TYPE: bool default: False
|
||||
LOAD_MODEL_ONLY: False
|
||||
# RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: ''
|
||||
ENABLE_GRADSCALER: False
|
||||
USE_SCALER: False
|
||||
RESUME_FROM:
|
||||
WORK_DIR: ./cache/save_data/dit_flux_schnell_1024_lora
|
||||
LOG_FILE: std_log.txt
|
||||
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
|
||||
EVAL_INTERVAL: 100
|
||||
# LOG_TRAIN_NUM DESCRIPTION: The number samples used to log in training phase. TYPE: int default: -1
|
||||
LOG_TRAIN_NUM: 16
|
||||
# FSDP_REDUCE_DTYPE DESCRIPTION: The dtype of reduce in FSDP. TYPE: str default: 'float16'
|
||||
FSDP_REDUCE_DTYPE: float32
|
||||
# FSDP_BUFFER_DTYPE DESCRIPTION: The dtype of buffer in FSDP. TYPE: str default: 'float16'
|
||||
FSDP_BUFFER_DTYPE: float32
|
||||
# FSDP_SHARD_MODULES DESCRIPTION: The modules to be sharded in FSDP. TYPE: list default: ['model']
|
||||
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ] #
|
||||
SAVE_MODULES: [ 'model'] #
|
||||
SAVE_MODULES: [ 'model']
|
||||
TRAIN_MODULES: ['model']
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
@@ -58,12 +47,9 @@ SOLVER:
|
||||
USE_EMA: False
|
||||
EVAL_EMA: False
|
||||
DIFFUSION:
|
||||
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
|
||||
NAME: DiffusionFluxRF
|
||||
PREDICTION_TYPE: raw
|
||||
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
|
||||
NOISE_SCHEDULER:
|
||||
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
|
||||
NAME: FlowMatchSigmaScheduler
|
||||
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
|
||||
WEIGHTING_SCHEME: logit_normal
|
||||
@@ -157,54 +143,33 @@ SOLVER:
|
||||
TANH_OUT: False
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
|
||||
NAME: T5PlusClipFluxEmbedder
|
||||
# T5_MODEL DESCRIPTION: TYPE: default: ''
|
||||
T5_MODEL:
|
||||
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
|
||||
NAME: HFEmbedder
|
||||
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_MODEL_CLS: T5EncoderModel
|
||||
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
|
||||
MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder_2/
|
||||
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_TOKENIZER_CLS: T5Tokenizer
|
||||
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer_2/
|
||||
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
|
||||
MAX_LENGTH: 256
|
||||
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
|
||||
OUTPUT_KEY: last_hidden_state
|
||||
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
|
||||
D_TYPE: bfloat16
|
||||
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
|
||||
BATCH_INFER: False
|
||||
CLEAN: whitespace
|
||||
# CLIP_MODEL DESCRIPTION: TYPE: default: ''
|
||||
CLIP_MODEL:
|
||||
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
|
||||
NAME: HFEmbedder
|
||||
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_MODEL_CLS: CLIPTextModel
|
||||
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
|
||||
MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder/
|
||||
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_TOKENIZER_CLS: CLIPTokenizer
|
||||
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer/
|
||||
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
|
||||
MAX_LENGTH: 77
|
||||
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
|
||||
OUTPUT_KEY: pooler_output
|
||||
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
|
||||
D_TYPE: bfloat16
|
||||
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
|
||||
BATCH_INFER: True
|
||||
CLEAN: whitespace
|
||||
#
|
||||
SAMPLE_ARGS:
|
||||
SAMPLE_STEPS: 4
|
||||
SAMPLER: flow_eluer
|
||||
SAMPLER: flow_euler
|
||||
SEED: 2024
|
||||
IMAGE_SIZE: [ 1024, 1024 ]
|
||||
GUIDE_SCALE: 3.5
|
||||
|
||||
@@ -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
|
||||
SKIP_EXAMPLES: True
|
||||
#
|
||||
MODEL:
|
||||
EDIT_MODEL:
|
||||
MODEL_CFG_DIR: scepter/methods/studio/chatbot/models/
|
||||
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,128 @@
|
||||
NAME: ACE_0.6B_1024
|
||||
IS_DEFAULT: False
|
||||
USE_DYNAMIC_MODEL: True
|
||||
DEFAULT_PARAS:
|
||||
PARAS:
|
||||
#
|
||||
INPUT:
|
||||
INPUT_IMAGE:
|
||||
INPUT_MASK:
|
||||
TASK:
|
||||
PROMPT: ""
|
||||
NEGATIVE_PROMPT: ""
|
||||
OUTPUT_HEIGHT: 1024
|
||||
OUTPUT_WIDTH: 1024
|
||||
SAMPLER: ddim
|
||||
SAMPLE_STEPS: 50
|
||||
GUIDE_SCALE: 4.5
|
||||
GUIDE_RESCALE: 0.5
|
||||
SEED: -1
|
||||
TAR_INDEX: 0
|
||||
OUTPUT:
|
||||
LATENT:
|
||||
IMAGES:
|
||||
SEED:
|
||||
MODULES_PARAS:
|
||||
FIRST_STAGE_MODEL:
|
||||
FUNCTION:
|
||||
- NAME: encode
|
||||
DTYPE: float16
|
||||
INPUT: ["IMAGE"]
|
||||
- NAME: decode
|
||||
DTYPE: float16
|
||||
INPUT: ["LATENT"]
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
FUNCTION:
|
||||
- NAME: forward
|
||||
DTYPE: float16
|
||||
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE"]
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
FUNCTION:
|
||||
- NAME: encode_list_of_list
|
||||
DTYPE: bfloat16
|
||||
INPUT: ["PROMPT"]
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusionACE
|
||||
PRETRAINED_MODEL:
|
||||
IGNORE_KEYS: [ ]
|
||||
SCALE_FACTOR: 0.18215
|
||||
SIZE_FACTOR: 8
|
||||
DECODER_BIAS: 0.5
|
||||
DEFAULT_N_PROMPT: ""
|
||||
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
|
||||
USE_TEXT_POS_EMBEDDINGS: True
|
||||
#
|
||||
DIFFUSION:
|
||||
NAME: BaseDiffusion
|
||||
PREDICTION_TYPE: eps
|
||||
MIN_SNR_GAMMA:
|
||||
NOISE_SCHEDULER:
|
||||
NAME: LinearScheduler
|
||||
NUM_TIMESTEPS: 1000
|
||||
BETA_MIN: 0.0001
|
||||
BETA_MAX: 0.02
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: ACE
|
||||
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/dit/ace_0.6b_1024px.pth
|
||||
IGNORE_KEYS: [ ]
|
||||
PATCH_SIZE: 2
|
||||
IN_CHANNELS: 4
|
||||
HIDDEN_SIZE: 1152
|
||||
DEPTH: 28
|
||||
NUM_HEADS: 16
|
||||
MLP_RATIO: 4.0
|
||||
PRED_SIGMA: True
|
||||
DROP_PATH: 0.0
|
||||
WINDOW_DIZE: 0
|
||||
Y_CHANNELS: 4096
|
||||
MAX_SEQ_LEN: 4096
|
||||
QK_NORM: True
|
||||
USE_GRAD_CHECKPOINT: True
|
||||
ATTENTION_BACKEND: flash_attn
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKL
|
||||
EMBED_DIM: 4
|
||||
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/vae/vae.bin
|
||||
IGNORE_KEYS: []
|
||||
#
|
||||
ENCODER:
|
||||
NAME: Encoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 4
|
||||
DOUBLE_Z: True
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
#
|
||||
DECODER:
|
||||
NAME: Decoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 4
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
GIVE_PRE_END: False
|
||||
TANH_OUT: False
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: T5EmbedderHF
|
||||
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/text_encoder/t5-v1_1-xxl/
|
||||
TOKENIZER_PATH: ms://iic/ACE-0.6B-1024px@models/tokenizer/t5-v1_1-xxl
|
||||
LENGTH: 120
|
||||
T5_DTYPE: bfloat16
|
||||
ADDED_IDENTIFIER: [ '{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
|
||||
CLEAN: whitespace
|
||||
USE_GRAD: False
|
||||
@@ -0,0 +1,284 @@
|
||||
NAME: ACE_0.6B_1024_REFINER
|
||||
IS_DEFAULT: False
|
||||
USE_DYNAMIC_MODEL: True
|
||||
DEFAULT_PARAS:
|
||||
PARAS:
|
||||
#
|
||||
INPUT:
|
||||
INPUT_IMAGE:
|
||||
INPUT_MASK:
|
||||
TASK:
|
||||
PROMPT: ""
|
||||
NEGATIVE_PROMPT: ""
|
||||
OUTPUT_HEIGHT: 1024
|
||||
OUTPUT_WIDTH: 1024
|
||||
SAMPLER: ddim
|
||||
SAMPLE_STEPS: 50
|
||||
GUIDE_SCALE: 4.5
|
||||
GUIDE_RESCALE: 0.5
|
||||
SEED: -1
|
||||
TAR_INDEX: 0
|
||||
REFINER_SCALE: 0.2
|
||||
USE_ACE: True
|
||||
#REFINER_PROMPT: "High Resolution, Sharpness, Clarity, Detail Enhancement, Noise Reduction, HD, 4k, Image Restoration, HDR"
|
||||
REFINER_PROMPT: "High Resolution, Sharpness, Clarity, Detail Enhancement, Noise Reduction, HD, 4k, Image Restoration, HDR"
|
||||
OUTPUT:
|
||||
LATENT:
|
||||
IMAGES:
|
||||
SEED:
|
||||
MODULES_PARAS:
|
||||
FIRST_STAGE_MODEL:
|
||||
FUNCTION:
|
||||
- NAME: encode
|
||||
DTYPE: float16
|
||||
INPUT: ["IMAGE"]
|
||||
- NAME: decode
|
||||
DTYPE: float16
|
||||
INPUT: ["LATENT"]
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
FUNCTION:
|
||||
- NAME: forward
|
||||
DTYPE: float16
|
||||
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE"]
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
FUNCTION:
|
||||
- NAME: encode_list_of_list
|
||||
DTYPE: bfloat16
|
||||
INPUT: ["PROMPT"]
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusionACE
|
||||
PRETRAINED_MODEL:
|
||||
IGNORE_KEYS: [ ]
|
||||
SCALE_FACTOR: 0.18215
|
||||
SIZE_FACTOR: 8
|
||||
DECODER_BIAS: 0.5
|
||||
DEFAULT_N_PROMPT: ""
|
||||
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
|
||||
USE_TEXT_POS_EMBEDDINGS: True
|
||||
#
|
||||
DIFFUSION:
|
||||
NAME: BaseDiffusion
|
||||
PREDICTION_TYPE: eps
|
||||
MIN_SNR_GAMMA:
|
||||
NOISE_SCHEDULER:
|
||||
NAME: LinearScheduler
|
||||
NUM_TIMESTEPS: 1000
|
||||
BETA_MIN: 0.0001
|
||||
BETA_MAX: 0.02
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: ACE
|
||||
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/dit/ace_0.6b_1024px.pth
|
||||
IGNORE_KEYS: [ ]
|
||||
PATCH_SIZE: 2
|
||||
IN_CHANNELS: 4
|
||||
HIDDEN_SIZE: 1152
|
||||
DEPTH: 28
|
||||
NUM_HEADS: 16
|
||||
MLP_RATIO: 4.0
|
||||
PRED_SIGMA: True
|
||||
DROP_PATH: 0.0
|
||||
WINDOW_DIZE: 0
|
||||
Y_CHANNELS: 4096
|
||||
MAX_SEQ_LEN: 4096
|
||||
QK_NORM: True
|
||||
USE_GRAD_CHECKPOINT: True
|
||||
ATTENTION_BACKEND: flash_attn
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKL
|
||||
EMBED_DIM: 4
|
||||
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/vae/vae.bin
|
||||
IGNORE_KEYS: []
|
||||
#
|
||||
ENCODER:
|
||||
NAME: Encoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 4
|
||||
DOUBLE_Z: True
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
#
|
||||
DECODER:
|
||||
NAME: Decoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 4
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
GIVE_PRE_END: False
|
||||
TANH_OUT: False
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: T5EmbedderHF
|
||||
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/text_encoder/t5-v1_1-xxl/
|
||||
TOKENIZER_PATH: ms://iic/ACE-0.6B-1024px@models/tokenizer/t5-v1_1-xxl
|
||||
LENGTH: 120
|
||||
T5_DTYPE: bfloat16
|
||||
ADDED_IDENTIFIER: [ '{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
|
||||
CLEAN: whitespace
|
||||
USE_GRAD: False
|
||||
|
||||
ACE_PROMPT: [
|
||||
"A cute cartoon rabbit holding a whiteboard that says 'ACE Refiner', standing in a sunny meadow filled with flowers, with a big smile and bright colors.",
|
||||
"A beautiful young woman with long flowing hair, wearing a summer dress, holding a whiteboard that reads 'ACE Refiner' while sitting on a park bench surrounded by cherry blossoms.",
|
||||
"An adorable cartoon cat wearing oversized glasses, holding a whiteboard that says 'ACE Refiner', perched on a stack of colorful books in a cozy library setting.",
|
||||
"A charming girl with pigtails, wearing a cute school uniform, enthusiastically holding a whiteboard that has 'ACE Refiner' written on it, in a bright and cheerful classroom full of educational posters.",
|
||||
"A friendly cartoon dog with floppy ears, sitting in front of a doghouse, proudly holding a whiteboard that says 'ACE Refiner', with a playful expression and a blue sky in the background.",
|
||||
"A cute anime girl with big expressive eyes, dressed in a colorful outfit, holding a whiteboard that reads 'ACE Refiner' in a fantastical landscape filled with mythical creatures.",
|
||||
"A vibrant cartoon fox holding a whiteboard that says 'ACE Refiner', standing on a rock by a sparkling stream, surrounded by lush greenery and butterflies.",
|
||||
"A stylish young woman in a business outfit, smiling as she holds a whiteboard written with 'ACE Refiner', in a modern office filled with plants and natural light.",
|
||||
"A cute cartoon unicorn holding a sparkling whiteboard that says 'ACE Refiner', frolicking in a magical forest, with rainbows and stars in the background.",
|
||||
"A happy family, consisting of a cute little girl and her playful puppy, holding a whiteboard that says 'ACE Refiner', together in their backyard on a sunny day."
|
||||
]
|
||||
REFINER_MODEL:
|
||||
NAME: ""
|
||||
IS_DEFAULT: False
|
||||
DEFAULT_PARAS:
|
||||
PARAS:
|
||||
RESOLUTIONS: [ [ 1024, 1024 ] ]
|
||||
INPUT:
|
||||
INPUT_IMAGE:
|
||||
INPUT_MASK:
|
||||
TASK:
|
||||
PROMPT: ""
|
||||
NEGATIVE_PROMPT: ""
|
||||
OUTPUT_HEIGHT: 1024
|
||||
OUTPUT_WIDTH: 1024
|
||||
SAMPLER: flow_euler
|
||||
SAMPLE_STEPS: 30
|
||||
GUIDE_SCALE: 3.5
|
||||
GUIDE_RESCALE:
|
||||
OUTPUT:
|
||||
LATENT:
|
||||
IMAGES:
|
||||
SEED:
|
||||
MODULES_PARAS:
|
||||
FIRST_STAGE_MODEL:
|
||||
FUNCTION:
|
||||
- NAME: encode
|
||||
DTYPE: bfloat16
|
||||
INPUT: [ "IMAGE" ]
|
||||
- NAME: decode
|
||||
DTYPE: bfloat16
|
||||
INPUT: [ "LATENT" ]
|
||||
PARAS:
|
||||
SCALE_FACTOR: 1.5305
|
||||
SHIFT_FACTOR: 0.0609
|
||||
SIZE_FACTOR: 8
|
||||
DIFFUSION_MODEL:
|
||||
FUNCTION:
|
||||
- NAME: forward
|
||||
DTYPE: bfloat16
|
||||
INPUT: [ "SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE" ]
|
||||
COND_STAGE_MODEL:
|
||||
FUNCTION:
|
||||
- NAME: encode
|
||||
DTYPE: bfloat16
|
||||
INPUT: [ "PROMPT" ]
|
||||
|
||||
MODEL:
|
||||
DIFFUSION:
|
||||
NAME: DiffusionFluxRF
|
||||
PREDICTION_TYPE: raw
|
||||
NOISE_SCHEDULER:
|
||||
NAME: FlowMatchSigmaScheduler
|
||||
WEIGHTING_SCHEME: logit_normal
|
||||
SHIFT: 3.0
|
||||
LOGIT_MEAN: 0.0
|
||||
LOGIT_STD: 1.0
|
||||
MODE_SCALE: 1.29
|
||||
DIFFUSION_MODEL:
|
||||
NAME: FluxMR
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@flux1-dev.safetensors
|
||||
IN_CHANNELS: 64
|
||||
OUT_CHANNELS: 64
|
||||
HIDDEN_SIZE: 3072
|
||||
NUM_HEADS: 24
|
||||
AXES_DIM: [ 16, 56, 56 ]
|
||||
THETA: 10000
|
||||
VEC_IN_DIM: 768
|
||||
GUIDANCE_EMBED: True
|
||||
CONTEXT_IN_DIM: 4096
|
||||
MLP_RATIO: 4.0
|
||||
QKV_BIAS: True
|
||||
DEPTH: 19
|
||||
DEPTH_SINGLE_BLOCKS: 38
|
||||
USE_GRAD_CHECKPOINT: True
|
||||
ATTN_BACKEND: flash_attn
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKLFlux
|
||||
EMBED_DIM: 16
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@ae.safetensors
|
||||
IGNORE_KEYS: [ ]
|
||||
BATCH_SIZE: 8
|
||||
USE_CONV: False
|
||||
SCALE_FACTOR: 0.3611
|
||||
SHIFT_FACTOR: 0.1159
|
||||
#
|
||||
ENCODER:
|
||||
NAME: Encoder
|
||||
USE_CHECKPOINT: False
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 16
|
||||
DOUBLE_Z: True
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
#
|
||||
DECODER:
|
||||
NAME: Decoder
|
||||
USE_CHECKPOINT: False
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 16
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
GIVE_PRE_END: False
|
||||
TANH_OUT: False
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: T5PlusClipFluxEmbedder
|
||||
T5_MODEL:
|
||||
NAME: HFEmbedder
|
||||
HF_MODEL_CLS: T5EncoderModel
|
||||
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder_2/
|
||||
HF_TOKENIZER_CLS: T5Tokenizer
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer_2/
|
||||
MAX_LENGTH: 512
|
||||
OUTPUT_KEY: last_hidden_state
|
||||
D_TYPE: bfloat16
|
||||
BATCH_INFER: False
|
||||
CLEAN: whitespace
|
||||
CLIP_MODEL:
|
||||
NAME: HFEmbedder
|
||||
HF_MODEL_CLS: CLIPTextModel
|
||||
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder/
|
||||
HF_TOKENIZER_CLS: CLIPTokenizer
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer/
|
||||
MAX_LENGTH: 77
|
||||
OUTPUT_KEY: pooler_output
|
||||
D_TYPE: bfloat16
|
||||
BATCH_INFER: True
|
||||
CLEAN: whitespace
|
||||
@@ -0,0 +1,128 @@
|
||||
NAME: ACE_0.6B_512
|
||||
IS_DEFAULT: True
|
||||
USE_DYNAMIC_MODEL: True
|
||||
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_of_list
|
||||
DTYPE: bfloat16
|
||||
INPUT: ["PROMPT"]
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusionACE
|
||||
PRETRAINED_MODEL:
|
||||
IGNORE_KEYS: [ ]
|
||||
SCALE_FACTOR: 0.18215
|
||||
SIZE_FACTOR: 8
|
||||
DECODER_BIAS: 0.5
|
||||
DEFAULT_N_PROMPT: ""
|
||||
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
|
||||
USE_TEXT_POS_EMBEDDINGS: True
|
||||
#
|
||||
DIFFUSION:
|
||||
NAME: BaseDiffusion
|
||||
PREDICTION_TYPE: eps
|
||||
MIN_SNR_GAMMA:
|
||||
NOISE_SCHEDULER:
|
||||
NAME: LinearScheduler
|
||||
NUM_TIMESTEPS: 1000
|
||||
BETA_MIN: 0.0001
|
||||
BETA_MAX: 0.02
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: ACE
|
||||
PRETRAINED_MODEL: ms://iic/ACE-0.6B-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,151 @@
|
||||
NAME: COGVIDEOX_2B
|
||||
IS_DEFAULT: False
|
||||
DEFAULT_PARAS:
|
||||
PARAS:
|
||||
RESOLUTIONS: [[480, 720]]
|
||||
INPUT:
|
||||
IMAGE:
|
||||
ORIGINAL_SIZE_AS_TUPLE: [480, 720]
|
||||
TARGET_SIZE_AS_TUPLE: [480, 720]
|
||||
PROMPT: ""
|
||||
NEGATIVE_PROMPT: ""
|
||||
PROMPT_PREFIX: ""
|
||||
SAMPLE: ddim
|
||||
SAMPLE_STEPS: 50
|
||||
GUIDE_SCALE: 6.0
|
||||
GUIDE_RESCALE: 0.0
|
||||
DISCRETIZATION: trailing
|
||||
NUM_FRAMES:
|
||||
DEFAULT: 49
|
||||
VISIBLE: True
|
||||
FPS:
|
||||
DEFAULT: 8
|
||||
VISIBLE: True
|
||||
OUTPUT:
|
||||
VIDEOS:
|
||||
SEED:
|
||||
MODULES_PARAS:
|
||||
FIRST_STAGE_MODEL:
|
||||
FUNCTION:
|
||||
-
|
||||
NAME: decode
|
||||
DTYPE: bfloat16
|
||||
INPUT: ["LATENT"]
|
||||
PARAS:
|
||||
SCALING_FACTOR_IMAGE: 1.15258426
|
||||
DIFFUSION_MODEL:
|
||||
FUNCTION:
|
||||
-
|
||||
NAME: forward
|
||||
DTYPE: bfloat16
|
||||
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE", "GUIDE_RESCALE", "DISCRETIZATION", "NUM_FRAMES", "FPS"]
|
||||
PARAS:
|
||||
USE_ROTARY_POSITIONAL_EMBEDDINGS: False
|
||||
PATCH_SIZE: 2
|
||||
LATENT_CHANNELS: 16
|
||||
SCALE_FACTOR_SPATIAL: 8
|
||||
SCALE_FACTOR_TEMPORAL: 4
|
||||
ATTENTION_HEAD_DIM: 64
|
||||
SAMPLE_HEIGHT: 480
|
||||
SAMPLE_WIDTH: 720
|
||||
COND_STAGE_MODEL:
|
||||
FUNCTION:
|
||||
-
|
||||
NAME: encode
|
||||
DTYPE: bfloat16
|
||||
INPUT: ["PROMPT"]
|
||||
#
|
||||
MODEL:
|
||||
PRETRAINED_MODEL:
|
||||
DIFFUSION:
|
||||
NAME: BaseDiffusion
|
||||
PREDICTION_TYPE: v
|
||||
NOISE_SCHEDULER:
|
||||
NAME: ScaledLinearScheduler
|
||||
BETA_MIN: 0.00085
|
||||
BETA_MAX: 0.012
|
||||
SNR_SHIFT_SCALE: 3.0
|
||||
RESCALE_BETAS_ZERO_SNR: True
|
||||
DIFFUSION_SAMPLERS:
|
||||
NAME: DDIMSampler
|
||||
DISCRETIZATION_TYPE: trailing
|
||||
ETA: 0.0
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: CogVideoXTransformer3DModel
|
||||
DTYPE: bfloat16
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@transformer/diffusion_pytorch_model.safetensors
|
||||
NUM_ATTENTION_HEADS: 30
|
||||
ATTENTION_HEAD_DIM: 64
|
||||
IN_CHANNELS: 16
|
||||
OUT_CHANNELS: 16
|
||||
FLIP_SIN_TO_COS: True
|
||||
FREQ_SHIFT: 0
|
||||
TIME_EMBED_DIM: 512
|
||||
TEXT_EMBED_DIM: 4096
|
||||
NUM_LAYERS: 30
|
||||
DROPOUT: 0.0
|
||||
ATTENTION_BIAS: True
|
||||
SAMPLE_WIDTH: 90
|
||||
SAMPLE_HEIGHT: 60
|
||||
SAMPLE_FRAMES: 49
|
||||
PATCH_SIZE: 2
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
MAX_TEXT_SEQ_LENGTH: 226
|
||||
ACTIVATION_FN: "gelu-approximate"
|
||||
TIMESTEP_ACTIVATION_FN: "silu"
|
||||
NORM_ELEMENTWISE_AFFINE: True
|
||||
NORM_EPS: 1e-5
|
||||
SPATIAL_INTERPOLATION_SCALE: 1.875
|
||||
TEMPORAL_INTERPOLATION_SCALE: 1.0
|
||||
USE_ROTARY_POSITIONAL_EMBEDDINGS: False
|
||||
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
|
||||
GRADIENT_CHECKPOINTING: False
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKLCogVideoX
|
||||
DTYPE: bfloat16
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@vae/diffusion_pytorch_model.safetensors
|
||||
SAMPLE_HEIGHT: 480
|
||||
SAMPLE_WIDTH: 720
|
||||
USE_QUANT_CONV: False
|
||||
USE_POST_QUANT_CONV: False
|
||||
USE_SLICING: True
|
||||
USE_TILING: True
|
||||
GRADIENT_CHECKPOINTING: False
|
||||
ENCODER:
|
||||
NAME: CogVideoXEncoder3D
|
||||
IN_CHANNELS: 3
|
||||
OUT_CHANNELS: 16
|
||||
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: False
|
||||
DECODER:
|
||||
NAME: CogVideoXDecoder3D
|
||||
IN_CHANNELS: 16
|
||||
OUT_CHANNELS: 3
|
||||
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: False
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: T5EmbedderHF
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
|
||||
LENGTH: 226
|
||||
CLEAN:
|
||||
USE_GRAD: False
|
||||
@@ -0,0 +1,153 @@
|
||||
NAME: COGVIDEOX_5B
|
||||
IS_DEFAULT: False
|
||||
DEFAULT_PARAS:
|
||||
PARAS:
|
||||
RESOLUTIONS: [[480, 720]]
|
||||
INPUT:
|
||||
IMAGE:
|
||||
ORIGINAL_SIZE_AS_TUPLE: [480, 720]
|
||||
TARGET_SIZE_AS_TUPLE: [480, 720]
|
||||
PROMPT: ""
|
||||
NEGATIVE_PROMPT: ""
|
||||
PROMPT_PREFIX: ""
|
||||
SAMPLE: ddim
|
||||
SAMPLE_STEPS: 50
|
||||
GUIDE_SCALE: 6.0
|
||||
GUIDE_RESCALE: 0.0
|
||||
DISCRETIZATION: trailing
|
||||
NUM_FRAMES:
|
||||
DEFAULT: 49
|
||||
VISIBLE: True
|
||||
FPS:
|
||||
DEFAULT: 8
|
||||
VISIBLE: True
|
||||
OUTPUT:
|
||||
VIDEOS:
|
||||
SEED:
|
||||
MODULES_PARAS:
|
||||
FIRST_STAGE_MODEL:
|
||||
FUNCTION:
|
||||
-
|
||||
NAME: decode
|
||||
DTYPE: bfloat16
|
||||
INPUT: ["LATENT"]
|
||||
PARAS:
|
||||
SCALING_FACTOR_IMAGE: 0.7 # 5b diff
|
||||
DIFFUSION_MODEL:
|
||||
FUNCTION:
|
||||
-
|
||||
NAME: forward
|
||||
DTYPE: bfloat16
|
||||
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE", "GUIDE_RESCALE", "DISCRETIZATION", "NUM_FRAMES", "FPS"]
|
||||
PARAS:
|
||||
USE_ROTARY_POSITIONAL_EMBEDDINGS: True
|
||||
PATCH_SIZE: 2
|
||||
LATENT_CHANNELS: 16
|
||||
SCALE_FACTOR_SPATIAL: 8
|
||||
SCALE_FACTOR_TEMPORAL: 4
|
||||
ATTENTION_HEAD_DIM: 64
|
||||
SAMPLE_HEIGHT: 480
|
||||
SAMPLE_WIDTH: 720
|
||||
COND_STAGE_MODEL:
|
||||
FUNCTION:
|
||||
-
|
||||
NAME: encode
|
||||
DTYPE: bfloat16
|
||||
INPUT: ["PROMPT"]
|
||||
#
|
||||
MODEL:
|
||||
PRETRAINED_MODEL:
|
||||
DIFFUSION:
|
||||
NAME: BaseDiffusion
|
||||
PREDICTION_TYPE: v
|
||||
NOISE_SCHEDULER:
|
||||
NAME: ScaledLinearScheduler
|
||||
BETA_MIN: 0.00085
|
||||
BETA_MAX: 0.012
|
||||
SNR_SHIFT_SCALE: 1.0 # 5b diff
|
||||
RESCALE_BETAS_ZERO_SNR: True
|
||||
DIFFUSION_SAMPLERS:
|
||||
NAME: DDIMSampler
|
||||
DISCRETIZATION_TYPE: trailing
|
||||
ETA: 0.0
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: CogVideoXTransformer3DModel
|
||||
DTYPE: bfloat16
|
||||
PRETRAINED_MODEL: # 5b diff
|
||||
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00001-of-00002.safetensors
|
||||
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00002-of-00002.safetensors
|
||||
NUM_ATTENTION_HEADS: 48 # 5b diff
|
||||
ATTENTION_HEAD_DIM: 64
|
||||
IN_CHANNELS: 16
|
||||
OUT_CHANNELS: 16
|
||||
FLIP_SIN_TO_COS: True
|
||||
FREQ_SHIFT: 0
|
||||
TIME_EMBED_DIM: 512
|
||||
TEXT_EMBED_DIM: 4096
|
||||
NUM_LAYERS: 42 # 5b diff
|
||||
DROPOUT: 0.0
|
||||
ATTENTION_BIAS: True
|
||||
SAMPLE_WIDTH: 90
|
||||
SAMPLE_HEIGHT: 60
|
||||
SAMPLE_FRAMES: 49
|
||||
PATCH_SIZE: 2
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
MAX_TEXT_SEQ_LENGTH: 226
|
||||
ACTIVATION_FN: "gelu-approximate"
|
||||
TIMESTEP_ACTIVATION_FN: "silu"
|
||||
NORM_ELEMENTWISE_AFFINE: True
|
||||
NORM_EPS: 1e-5
|
||||
SPATIAL_INTERPOLATION_SCALE: 1.875
|
||||
TEMPORAL_INTERPOLATION_SCALE: 1.0
|
||||
USE_ROTARY_POSITIONAL_EMBEDDINGS: True # 5b diff
|
||||
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKLCogVideoX
|
||||
DTYPE: bfloat16
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b@vae/diffusion_pytorch_model.safetensors # 5b diff
|
||||
SAMPLE_HEIGHT: 480
|
||||
SAMPLE_WIDTH: 720
|
||||
USE_QUANT_CONV: False
|
||||
USE_POST_QUANT_CONV: False
|
||||
USE_SLICING: True
|
||||
USE_TILING: True
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
ENCODER:
|
||||
NAME: CogVideoXEncoder3D
|
||||
IN_CHANNELS: 3
|
||||
OUT_CHANNELS: 16
|
||||
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
DECODER:
|
||||
NAME: CogVideoXDecoder3D
|
||||
IN_CHANNELS: 16
|
||||
OUT_CHANNELS: 3
|
||||
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: T5EmbedderHF
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
|
||||
LENGTH: 226
|
||||
CLEAN:
|
||||
USE_GRAD: False
|
||||
@@ -13,8 +13,8 @@ DEFAULT_PARAS:
|
||||
VISIBLE: False
|
||||
PROMPT_PREFIX: ""
|
||||
SAMPLE:
|
||||
VALUES: ["flow_eluer"]
|
||||
DEFAULT: "flow_eluer"
|
||||
VALUES: ["flow_euler"]
|
||||
DEFAULT: "flow_euler"
|
||||
SAMPLE_STEPS: 50
|
||||
GUIDE_SCALE: 3.5
|
||||
GUIDE_RESCALE:
|
||||
|
||||
@@ -13,8 +13,8 @@ DEFAULT_PARAS:
|
||||
VISIBLE: False
|
||||
PROMPT_PREFIX: ""
|
||||
SAMPLE:
|
||||
VALUES: ["flow_eluer"]
|
||||
DEFAULT: "flow_eluer"
|
||||
VALUES: ["flow_euler"]
|
||||
DEFAULT: "flow_euler"
|
||||
SAMPLE_STEPS: 4
|
||||
GUIDE_SCALE: 3.5
|
||||
GUIDE_RESCALE:
|
||||
|
||||
@@ -18,6 +18,16 @@ DIFFUSION_PARAS:
|
||||
MAX: 4
|
||||
DEFAULT: 1
|
||||
VISIBLE: True
|
||||
NUM_FRAMES:
|
||||
MIN: 1
|
||||
MAX: 100
|
||||
DEFAULT: 49
|
||||
VISIBLE: False
|
||||
FPS:
|
||||
MIN: 1
|
||||
MAX: 50
|
||||
DEFAULT: 8
|
||||
VISIBLE: False
|
||||
SAMPLE_STEPS:
|
||||
MIN: 1
|
||||
MAX: 100
|
||||
@@ -93,7 +103,8 @@ DIFFUSION_PARAS:
|
||||
[1664, 576], [1728, 576],
|
||||
[2048, 2048], [2048, 1920], [1920, 2048],
|
||||
[1536, 2560], [2560, 1536], [2560, 1440],
|
||||
[2560, 1440]
|
||||
[2560, 1440],
|
||||
[480, 720], [720, 480]
|
||||
]
|
||||
DEFAULT: [1024, 1024]
|
||||
VISIBLE: True
|
||||
|
||||
@@ -450,3 +450,28 @@ PROCESSORS:
|
||||
SRC_IMAGE_TOOL: sketch
|
||||
SRC_IMAGE_INTERACTIVE: True
|
||||
CAPTION_INTERACTIVE: False
|
||||
|
||||
VIDEO_PROCESSORS:
|
||||
- NAME: CogVLM2Llama3Caption
|
||||
TYPE: caption
|
||||
MODEL_PATH: ms://ZhipuAI/cogvlm2-llama3-caption
|
||||
DEVICE: "gpu"
|
||||
MEMORY: 20000
|
||||
PROMPT: Please describe this video in detail.
|
||||
TEMPERATURE: 0.1
|
||||
MAX_NEW_TOKENS: 2048
|
||||
PAD_TOKEN_ID: 128002
|
||||
TOP_K: 1
|
||||
TOP_P: 0.1
|
||||
|
||||
TRANSLATION_PROCESSORS:
|
||||
- NAME: OpusMtZhEn
|
||||
TYPE: caption
|
||||
MODEL_PATH: ms://cubeai/trans-opus-mt-zh-en
|
||||
DEVICE: "gpu"
|
||||
MEMORY: 5000
|
||||
- NAME: OpusMtEnZh
|
||||
TYPE: caption
|
||||
MODEL_PATH: ms://cubeai/trans-opus-mt-en-zh
|
||||
DEVICE: "gpu"
|
||||
MEMORY: 5000
|
||||
@@ -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,315 @@
|
||||
ENV:
|
||||
BACKEND: nccl
|
||||
SEED: 42
|
||||
TENSOR_PARALLEL_SIZE: 1
|
||||
PIPELINE_PARALLEL_SIZE: 1
|
||||
SYS_ENVS:
|
||||
TORCH_CUDNN_V8_API_ENABLED: '1'
|
||||
TOKENIZERS_PARALLELISM: 'false'
|
||||
TF_CPP_MIN_LOG_LEVEL: '3'
|
||||
PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True'
|
||||
META:
|
||||
VERSION: 'COGVIDEOX_2B'
|
||||
DESCRIPTION: "cogvideox 2b"
|
||||
IS_DEFAULT: False
|
||||
IS_SHARE: True
|
||||
INFERENCE_PARAS:
|
||||
INFERENCE_BATCH_SIZE: 1
|
||||
INFERENCE_PREFIX: ""
|
||||
DEFAULT_SAMPLER: "ddim"
|
||||
DEFAULT_SAMPLE_STEPS: 50
|
||||
INFERENCE_N_PROMPT: ""
|
||||
RESOLUTION: [ 480, 720 ]
|
||||
PARAS:
|
||||
- TRAIN_BATCH_SIZE: 1
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: [ 480, 720 ]
|
||||
MEMORY: 89000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
EPSEC: 0.818
|
||||
LEARNING_RATE: 4e-4
|
||||
IS_DEFAULT: False
|
||||
TUNER: FULL
|
||||
- TRAIN_BATCH_SIZE: 1
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: [ 480, 720 ]
|
||||
MEMORY: 89000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
EPSEC: 0.818
|
||||
LEARNING_RATE: 4e-4
|
||||
IS_DEFAULT: True
|
||||
TUNER: LORA
|
||||
#
|
||||
TUNERS:
|
||||
LORA:
|
||||
- NAME: SwiftLoRA
|
||||
R: 64
|
||||
LORA_ALPHA: 64
|
||||
LORA_DROPOUT: 0.0
|
||||
BIAS: "none"
|
||||
TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$"
|
||||
#
|
||||
SOLVER:
|
||||
NAME: LatentDiffusionVideoSolver
|
||||
MAX_STEPS: 2000
|
||||
USE_AMP: True
|
||||
DTYPE: bfloat16
|
||||
USE_FAIRSCALE: False
|
||||
USE_FSDP: True
|
||||
LOAD_MODEL_ONLY: False
|
||||
ENABLE_GRADSCALER: False
|
||||
USE_SCALER: False
|
||||
RESUME_FROM:
|
||||
WORK_DIR: ./cache/save_data/dit_cogvideox_2b_lora
|
||||
LOG_FILE: std_log.txt
|
||||
EVAL_INTERVAL: 100
|
||||
LOG_TRAIN_NUM: 4
|
||||
FPS: 8
|
||||
SHARDING_STRATEGY: full_shard
|
||||
FSDP_REDUCE_DTYPE: float32
|
||||
FSDP_BUFFER_DTYPE: float32
|
||||
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model']
|
||||
SAVE_MODULES: [ 'model', 'cond_stage_model.model']
|
||||
TRAIN_MODULES: ['model']
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TUNER:
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusionCogVideoX
|
||||
PRETRAINED_MODEL:
|
||||
PARAMETERIZATION: v
|
||||
TIMESTEPS: 1000
|
||||
MIN_SNR_GAMMA: 3.0
|
||||
ZERO_TERMINAL_SNR: True
|
||||
SCALE_FACTOR_SPATIAL: 8
|
||||
SCALE_FACTOR_TEMPORAL: 4
|
||||
SCALING_FACTOR_IMAGE: 1.15258426
|
||||
IGNORE_KEYS: [ ]
|
||||
DEFAULT_N_PROMPT:
|
||||
USE_EMA: False
|
||||
EVAL_EMA: False
|
||||
DIFFUSION:
|
||||
NAME: BaseDiffusion
|
||||
PREDICTION_TYPE: v
|
||||
NOISE_SCHEDULER:
|
||||
NAME: ScaledLinearScheduler
|
||||
BETA_MIN: 0.00085
|
||||
BETA_MAX: 0.012
|
||||
SNR_SHIFT_SCALE: 3.0
|
||||
RESCALE_BETAS_ZERO_SNR: True
|
||||
DIFFUSION_SAMPLERS:
|
||||
NAME: DDIMSampler
|
||||
DISCRETIZATION_TYPE: trailing
|
||||
ETA: 0.0
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: CogVideoXTransformer3DModel
|
||||
DTYPE: bfloat16
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@transformer/diffusion_pytorch_model.safetensors
|
||||
NUM_ATTENTION_HEADS: 30
|
||||
ATTENTION_HEAD_DIM: 64
|
||||
IN_CHANNELS: 16
|
||||
OUT_CHANNELS: 16
|
||||
FLIP_SIN_TO_COS: True
|
||||
FREQ_SHIFT: 0
|
||||
TIME_EMBED_DIM: 512
|
||||
TEXT_EMBED_DIM: 4096
|
||||
NUM_LAYERS: 30
|
||||
DROPOUT: 0.0
|
||||
ATTENTION_BIAS: True
|
||||
SAMPLE_WIDTH: 90
|
||||
SAMPLE_HEIGHT: 60
|
||||
SAMPLE_FRAMES: 49
|
||||
PATCH_SIZE: 2
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
MAX_TEXT_SEQ_LENGTH: 226
|
||||
ACTIVATION_FN: "gelu-approximate"
|
||||
TIMESTEP_ACTIVATION_FN: "silu"
|
||||
NORM_ELEMENTWISE_AFFINE: True
|
||||
NORM_EPS: 1e-5
|
||||
SPATIAL_INTERPOLATION_SCALE: 1.875
|
||||
TEMPORAL_INTERPOLATION_SCALE: 1.0
|
||||
USE_ROTARY_POSITIONAL_EMBEDDINGS: False
|
||||
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
|
||||
GRADIENT_CHECKPOINTING: False
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKLCogVideoX
|
||||
DTYPE: bfloat16
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@vae/diffusion_pytorch_model.safetensors
|
||||
SAMPLE_HEIGHT: 480
|
||||
SAMPLE_WIDTH: 720
|
||||
USE_QUANT_CONV: False
|
||||
USE_POST_QUANT_CONV: False
|
||||
USE_SLICING: True
|
||||
USE_TILING: True
|
||||
GRADIENT_CHECKPOINTING: False
|
||||
ENCODER:
|
||||
NAME: CogVideoXEncoder3D
|
||||
IN_CHANNELS: 3
|
||||
OUT_CHANNELS: 16
|
||||
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: False
|
||||
DECODER:
|
||||
NAME: CogVideoXDecoder3D
|
||||
IN_CHANNELS: 16
|
||||
OUT_CHANNELS: 3
|
||||
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: False
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: T5EmbedderHF
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
|
||||
LENGTH: 226
|
||||
CLEAN:
|
||||
USE_GRAD: False
|
||||
#
|
||||
LOSS:
|
||||
NAME: ReconstructLoss
|
||||
LOSS_TYPE: l2
|
||||
#
|
||||
SAMPLE_ARGS:
|
||||
SAMPLER: ddim
|
||||
SAMPLE_STEPS: 50
|
||||
SEED: 42
|
||||
GUIDE_SCALE: 6.0
|
||||
GUIDE_RESCALE: 0.0
|
||||
NUM_FRAMES: 49
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: Adam
|
||||
LEARNING_RATE: 1e-3
|
||||
BETAS: [ 0.9, 0.95 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 0.0
|
||||
AMSGRAD: False
|
||||
#
|
||||
# LR_SCHEDULER:
|
||||
# NAME: StepAnnealingLR
|
||||
# WARMUP_STEPS: 200
|
||||
# TOTAL_STEPS: 2000
|
||||
# DECAY_MODE: 'cosine'
|
||||
#
|
||||
TRAIN_DATA:
|
||||
NAME: VideoGenDatasetOTF
|
||||
MODE: train
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
PROMPT_PREFIX: ''
|
||||
DELIMITER: '#;#'
|
||||
FIELDS: [ 'video_path', 'width', 'height', 'prompt' ]
|
||||
PATH_PREFIX:
|
||||
DATA_FILE:
|
||||
SAMPLER:
|
||||
NAME: LoopSampler
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'video', 'video_latent', "prompt" ]
|
||||
META_KEYS: [ ]
|
||||
MODEL:
|
||||
NAME: AutoencoderKLCogVideoX
|
||||
DTYPE: bfloat16
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@vae/diffusion_pytorch_model.safetensors
|
||||
SAMPLE_HEIGHT: 480
|
||||
SAMPLE_WIDTH: 720
|
||||
USE_QUANT_CONV: False
|
||||
USE_POST_QUANT_CONV: False
|
||||
USE_SLICING: True
|
||||
USE_TILING: True
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
ENCODER:
|
||||
NAME: CogVideoXEncoder3D
|
||||
IN_CHANNELS: 3
|
||||
OUT_CHANNELS: 16
|
||||
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
DECODER:
|
||||
NAME: CogVideoXDecoder3D
|
||||
IN_CHANNELS: 16
|
||||
OUT_CHANNELS: 3
|
||||
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: Text2ImageDataset
|
||||
MODE: eval
|
||||
PROMPT_FILE:
|
||||
PROMPT_DATA: [ "A girl riding a bike.", "A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance." ]
|
||||
IMAGE_SIZE: [ 480, 720 ]
|
||||
FIELDS: [ "prompt" ]
|
||||
DELIMITER: '#;#'
|
||||
PROMPT_PREFIX: ''
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
# USE_NUM: 8
|
||||
NUM_WORKERS: 4
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'index', 'prompt' ]
|
||||
META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
- NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
PRIORITY: 0
|
||||
- NAME: BackwardHook
|
||||
PRIORITY: 10
|
||||
- NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
PRIORITY: 20
|
||||
- NAME: CheckpointHook
|
||||
INTERVAL: 1000
|
||||
PRIORITY: 40
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
DISABLE_SNAPSHOT: True
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
- NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
PRIORITY: 0
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
SAVE_PROBE_PREFIX: 'image'
|
||||
@@ -0,0 +1,317 @@
|
||||
ENV:
|
||||
BACKEND: nccl
|
||||
SEED: 42
|
||||
TENSOR_PARALLEL_SIZE: 1
|
||||
PIPELINE_PARALLEL_SIZE: 1
|
||||
SYS_ENVS:
|
||||
TORCH_CUDNN_V8_API_ENABLED: '1'
|
||||
TOKENIZERS_PARALLELISM: 'false'
|
||||
TF_CPP_MIN_LOG_LEVEL: '3'
|
||||
PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True'
|
||||
META:
|
||||
VERSION: 'COGVIDEOX_5B'
|
||||
DESCRIPTION: "cogvideox 5b"
|
||||
IS_DEFAULT: False
|
||||
IS_SHARE: True
|
||||
INFERENCE_PARAS:
|
||||
INFERENCE_BATCH_SIZE: 1
|
||||
INFERENCE_PREFIX: ""
|
||||
DEFAULT_SAMPLER: "ddim"
|
||||
DEFAULT_SAMPLE_STEPS: 50
|
||||
INFERENCE_N_PROMPT: ""
|
||||
RESOLUTION: [ 480, 720 ]
|
||||
PARAS:
|
||||
- TRAIN_BATCH_SIZE: 1
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: [ 480, 720 ]
|
||||
MEMORY: 89000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
EPSEC: 0.818
|
||||
LEARNING_RATE: 4e-4
|
||||
IS_DEFAULT: False
|
||||
TUNER: FULL
|
||||
- TRAIN_BATCH_SIZE: 1
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: [ 480, 720 ]
|
||||
MEMORY: 89000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
EPSEC: 0.818
|
||||
LEARNING_RATE: 4e-4
|
||||
IS_DEFAULT: True
|
||||
TUNER: LORA
|
||||
#
|
||||
TUNERS:
|
||||
LORA:
|
||||
- NAME: SwiftLoRA
|
||||
R: 64
|
||||
LORA_ALPHA: 64
|
||||
LORA_DROPOUT: 0.0
|
||||
BIAS: "none"
|
||||
TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$"
|
||||
#
|
||||
SOLVER:
|
||||
NAME: LatentDiffusionVideoSolver
|
||||
MAX_STEPS: 2000
|
||||
USE_AMP: True
|
||||
DTYPE: bfloat16
|
||||
USE_FAIRSCALE: False
|
||||
USE_FSDP: True
|
||||
LOAD_MODEL_ONLY: False
|
||||
ENABLE_GRADSCALER: False
|
||||
USE_SCALER: False
|
||||
RESUME_FROM:
|
||||
WORK_DIR: ./cache/save_data/dit_cogvideox_5b_lora
|
||||
LOG_FILE: std_log.txt
|
||||
EVAL_INTERVAL: 100
|
||||
LOG_TRAIN_NUM: 4
|
||||
FPS: 8
|
||||
SHARDING_STRATEGY: full_shard
|
||||
FSDP_REDUCE_DTYPE: float32
|
||||
FSDP_BUFFER_DTYPE: float32
|
||||
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model']
|
||||
SAVE_MODULES: [ 'model', 'cond_stage_model.model']
|
||||
TRAIN_MODULES: ['model']
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
TUNER:
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusionCogVideoX
|
||||
PRETRAINED_MODEL:
|
||||
PARAMETERIZATION: v
|
||||
TIMESTEPS: 1000
|
||||
MIN_SNR_GAMMA: 3.0
|
||||
ZERO_TERMINAL_SNR: True
|
||||
SCALE_FACTOR_SPATIAL: 8
|
||||
SCALE_FACTOR_TEMPORAL: 4
|
||||
SCALING_FACTOR_IMAGE: 0.7 # 5b diff
|
||||
IGNORE_KEYS: [ ]
|
||||
DEFAULT_N_PROMPT:
|
||||
USE_EMA: False
|
||||
EVAL_EMA: False
|
||||
DIFFUSION:
|
||||
NAME: BaseDiffusion
|
||||
PREDICTION_TYPE: v
|
||||
NOISE_SCHEDULER:
|
||||
NAME: ScaledLinearScheduler
|
||||
BETA_MIN: 0.00085
|
||||
BETA_MAX: 0.012
|
||||
SNR_SHIFT_SCALE: 1.0 # 5b diff
|
||||
RESCALE_BETAS_ZERO_SNR: True
|
||||
DIFFUSION_SAMPLERS:
|
||||
NAME: DDIMSampler
|
||||
DISCRETIZATION_TYPE: trailing
|
||||
ETA: 0.0
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: CogVideoXTransformer3DModel
|
||||
DTYPE: bfloat16
|
||||
PRETRAINED_MODEL: # 5b diff
|
||||
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00001-of-00002.safetensors
|
||||
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00002-of-00002.safetensors
|
||||
NUM_ATTENTION_HEADS: 48 # 5b diff
|
||||
ATTENTION_HEAD_DIM: 64
|
||||
IN_CHANNELS: 16
|
||||
OUT_CHANNELS: 16
|
||||
FLIP_SIN_TO_COS: True
|
||||
FREQ_SHIFT: 0
|
||||
TIME_EMBED_DIM: 512
|
||||
TEXT_EMBED_DIM: 4096
|
||||
NUM_LAYERS: 42 # 5b diff
|
||||
DROPOUT: 0.0
|
||||
ATTENTION_BIAS: True
|
||||
SAMPLE_WIDTH: 90
|
||||
SAMPLE_HEIGHT: 60
|
||||
SAMPLE_FRAMES: 49
|
||||
PATCH_SIZE: 2
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
MAX_TEXT_SEQ_LENGTH: 226
|
||||
ACTIVATION_FN: "gelu-approximate"
|
||||
TIMESTEP_ACTIVATION_FN: "silu"
|
||||
NORM_ELEMENTWISE_AFFINE: True
|
||||
NORM_EPS: 1e-5
|
||||
SPATIAL_INTERPOLATION_SCALE: 1.875
|
||||
TEMPORAL_INTERPOLATION_SCALE: 1.0
|
||||
USE_ROTARY_POSITIONAL_EMBEDDINGS: True # 5b diff
|
||||
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKLCogVideoX
|
||||
DTYPE: bfloat16
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b@vae/diffusion_pytorch_model.safetensors # 5b diff
|
||||
SAMPLE_HEIGHT: 480
|
||||
SAMPLE_WIDTH: 720
|
||||
USE_QUANT_CONV: False
|
||||
USE_POST_QUANT_CONV: False
|
||||
USE_SLICING: True
|
||||
USE_TILING: True
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
ENCODER:
|
||||
NAME: CogVideoXEncoder3D
|
||||
IN_CHANNELS: 3
|
||||
OUT_CHANNELS: 16
|
||||
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
DECODER:
|
||||
NAME: CogVideoXDecoder3D
|
||||
IN_CHANNELS: 16
|
||||
OUT_CHANNELS: 3
|
||||
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: T5EmbedderHF
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
|
||||
LENGTH: 226
|
||||
CLEAN:
|
||||
USE_GRAD: False
|
||||
#
|
||||
LOSS:
|
||||
NAME: ReconstructLoss
|
||||
LOSS_TYPE: l2
|
||||
#
|
||||
SAMPLE_ARGS:
|
||||
SAMPLER: ddim
|
||||
SAMPLE_STEPS: 50
|
||||
SEED: 42
|
||||
GUIDE_SCALE: 6.0
|
||||
GUIDE_RESCALE: 0.0
|
||||
NUM_FRAMES: 49
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: Adam
|
||||
LEARNING_RATE: 1e-3
|
||||
BETAS: [ 0.9, 0.95 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 0.0
|
||||
AMSGRAD: False
|
||||
#
|
||||
# LR_SCHEDULER:
|
||||
# NAME: StepAnnealingLR
|
||||
# WARMUP_STEPS: 200
|
||||
# TOTAL_STEPS: 2000
|
||||
# DECAY_MODE: 'cosine'
|
||||
#
|
||||
TRAIN_DATA:
|
||||
NAME: VideoGenDatasetOTF
|
||||
MODE: train
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
PROMPT_PREFIX: ''
|
||||
DELIMITER: '#;#'
|
||||
FIELDS: [ 'video_path', 'width', 'height', 'prompt' ]
|
||||
PATH_PREFIX:
|
||||
DATA_FILE:
|
||||
SAMPLER:
|
||||
NAME: LoopSampler
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'video', 'video_latent', "prompt" ]
|
||||
META_KEYS: [ ]
|
||||
MODEL:
|
||||
NAME: AutoencoderKLCogVideoX
|
||||
DTYPE: bfloat16
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b@vae/diffusion_pytorch_model.safetensors
|
||||
SAMPLE_HEIGHT: 480
|
||||
SAMPLE_WIDTH: 720
|
||||
USE_QUANT_CONV: False
|
||||
USE_POST_QUANT_CONV: False
|
||||
USE_SLICING: True
|
||||
USE_TILING: True
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
ENCODER:
|
||||
NAME: CogVideoXEncoder3D
|
||||
IN_CHANNELS: 3
|
||||
OUT_CHANNELS: 16
|
||||
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
DECODER:
|
||||
NAME: CogVideoXDecoder3D
|
||||
IN_CHANNELS: 16
|
||||
OUT_CHANNELS: 3
|
||||
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
|
||||
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
|
||||
LAYERS_PER_BLOCK: 3
|
||||
ACT_FN: "silu"
|
||||
NORM_EPS: 1e-6
|
||||
NORM_NUM_GROUPS: 32
|
||||
DROPOUT: 0.0
|
||||
PAD_MODE: "first"
|
||||
TEMPORAL_COMPRESSION_RATIO: 4
|
||||
GRADIENT_CHECKPOINTING: True
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: Text2ImageDataset
|
||||
MODE: eval
|
||||
PROMPT_FILE:
|
||||
PROMPT_DATA: [ "A girl riding a bike.", "A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance." ]
|
||||
IMAGE_SIZE: [ 480, 720 ]
|
||||
FIELDS: [ "prompt" ]
|
||||
DELIMITER: '#;#'
|
||||
PROMPT_PREFIX: ''
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
# USE_NUM: 8
|
||||
NUM_WORKERS: 4
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'index', 'prompt' ]
|
||||
META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
- NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
PRIORITY: 0
|
||||
- NAME: BackwardHook
|
||||
PRIORITY: 10
|
||||
- NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
PRIORITY: 20
|
||||
- NAME: CheckpointHook
|
||||
INTERVAL: 1000
|
||||
PRIORITY: 40
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
DISABLE_SNAPSHOT: True
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
- NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
PRIORITY: 0
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
SAVE_PROBE_PREFIX: 'image'
|
||||
@@ -3,7 +3,7 @@ ENV:
|
||||
META:
|
||||
VERSION: 'FLUX1.0_DEV'
|
||||
DESCRIPTION: "flux 1.0 dev"
|
||||
IS_DEFAULT: False
|
||||
IS_DEFAULT: True
|
||||
IS_SHARE: True
|
||||
INFERENCE_PARAS:
|
||||
INFERENCE_BATCH_SIZE: 1
|
||||
@@ -50,43 +50,33 @@ META:
|
||||
#
|
||||
SOLVER:
|
||||
NAME: LatentDiffusionSolver
|
||||
# MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
|
||||
MAX_STEPS: 100000
|
||||
# USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False
|
||||
USE_AMP: True
|
||||
# DTYPE DESCRIPTION: The precision for training. TYPE: str default: 'float32'
|
||||
DTYPE: bfloat16
|
||||
# USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
|
||||
USE_FAIRSCALE: False
|
||||
# USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
|
||||
USE_FSDP: True
|
||||
# LOAD_MODEL_ONLY DESCRIPTION: Only load the model rather than the optimizer and schedule, default is False. TYPE: bool default: False
|
||||
LOAD_MODEL_ONLY: False
|
||||
# RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: ''
|
||||
RESUME_FROM:
|
||||
WORK_DIR: ./cache/save_data/dit_flux_dev_1024_lora
|
||||
LOG_FILE: std_log.txt
|
||||
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
|
||||
EVAL_INTERVAL: 100
|
||||
# LOG_TRAIN_NUM DESCRIPTION: The number samples used to log in training phase. TYPE: int default: -1
|
||||
LOG_TRAIN_NUM: 16
|
||||
# FSDP_REDUCE_DTYPE DESCRIPTION: The dtype of reduce in FSDP. TYPE: str default: 'float16'
|
||||
ENABLE_GRADSCALER: False
|
||||
USE_SCALER: False
|
||||
FSDP_REDUCE_DTYPE: float32
|
||||
# FSDP_BUFFER_DTYPE DESCRIPTION: The dtype of buffer in FSDP. TYPE: str default: 'float16'
|
||||
FSDP_BUFFER_DTYPE: float32
|
||||
# FSDP_SHARD_MODULES DESCRIPTION: The modules to be sharded in FSDP. TYPE: list default: ['model']
|
||||
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ] #
|
||||
SAVE_MODULES: [ 'model'] #
|
||||
SAVE_MODULES: [ 'model']
|
||||
TRAIN_MODULES: ['model']
|
||||
#
|
||||
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
|
||||
FREEZE:
|
||||
#
|
||||
|
||||
TUNER:
|
||||
#
|
||||
|
||||
MODEL:
|
||||
NAME: LatentDiffusionFlux
|
||||
PARAMETERIZATION: rf
|
||||
@@ -99,65 +89,39 @@ SOLVER:
|
||||
USE_EMA: False
|
||||
EVAL_EMA: False
|
||||
DIFFUSION:
|
||||
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
|
||||
NAME: DiffusionFluxRF
|
||||
PREDICTION_TYPE: raw
|
||||
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
|
||||
NOISE_SCHEDULER:
|
||||
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
|
||||
NAME: FlowMatchSigmaScheduler
|
||||
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
|
||||
WEIGHTING_SCHEME: logit_normal
|
||||
SHIFT: 3.0
|
||||
# LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
|
||||
LOGIT_MEAN: 0.0
|
||||
# LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
|
||||
LOGIT_STD: 1.0
|
||||
# MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
|
||||
MODE_SCALE: 1.29
|
||||
SAMPLER_SCHEDULER:
|
||||
# NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
|
||||
NAME: FlowMatchFluxShiftScheduler
|
||||
# SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
|
||||
SHIFT: False
|
||||
# SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
|
||||
SIGMOID_SCALE: 1
|
||||
# BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
|
||||
BASE_SHIFT: 0.5
|
||||
# MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
|
||||
MAX_SHIFT: 1.15
|
||||
#
|
||||
|
||||
DIFFUSION_MODEL:
|
||||
# NAME DESCRIPTION: TYPE: default: 'Flux'
|
||||
NAME: Flux
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@flux1-dev.safetensors
|
||||
# IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
|
||||
IN_CHANNELS: 64
|
||||
# HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
|
||||
HIDDEN_SIZE: 3072
|
||||
# NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
|
||||
NUM_HEADS: 24
|
||||
# AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
|
||||
AXES_DIM: [ 16, 56, 56 ]
|
||||
# THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
|
||||
THETA: 10000
|
||||
# VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
|
||||
VEC_IN_DIM: 768
|
||||
# GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
|
||||
GUIDANCE_EMBED: False
|
||||
# CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
|
||||
CONTEXT_IN_DIM: 4096
|
||||
# MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
|
||||
MLP_RATIO: 4.0
|
||||
# QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
|
||||
QKV_BIAS: True
|
||||
# DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
|
||||
DEPTH: 19
|
||||
# DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
|
||||
DEPTH_SINGLE_BLOCKS: 38
|
||||
USE_GRAD_CHECKPOINT: True
|
||||
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKLFlux
|
||||
EMBED_DIM: 16
|
||||
@@ -167,7 +131,7 @@ SOLVER:
|
||||
USE_CONV: False
|
||||
SCALE_FACTOR: 0.3611
|
||||
SHIFT_FACTOR: 0.1159
|
||||
#
|
||||
|
||||
ENCODER:
|
||||
NAME: Encoder
|
||||
USE_CHECKPOINT: True
|
||||
@@ -181,7 +145,7 @@ SOLVER:
|
||||
DOUBLE_Z: True
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
#
|
||||
|
||||
DECODER:
|
||||
NAME: Decoder
|
||||
USE_CHECKPOINT: True
|
||||
@@ -196,61 +160,40 @@ SOLVER:
|
||||
RESAMP_WITH_CONV: True
|
||||
GIVE_PRE_END: False
|
||||
TANH_OUT: False
|
||||
#
|
||||
|
||||
COND_STAGE_MODEL:
|
||||
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
|
||||
NAME: T5PlusClipFluxEmbedder
|
||||
# T5_MODEL DESCRIPTION: TYPE: default: ''
|
||||
T5_MODEL:
|
||||
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
|
||||
NAME: HFEmbedder
|
||||
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_MODEL_CLS: T5EncoderModel
|
||||
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
|
||||
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder_2/
|
||||
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_TOKENIZER_CLS: T5Tokenizer
|
||||
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer_2/
|
||||
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
|
||||
MAX_LENGTH: 512
|
||||
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
|
||||
OUTPUT_KEY: last_hidden_state
|
||||
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
|
||||
D_TYPE: bfloat16
|
||||
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
|
||||
BATCH_INFER: False
|
||||
CLEAN: whitespace
|
||||
# CLIP_MODEL DESCRIPTION: TYPE: default: ''
|
||||
CLIP_MODEL:
|
||||
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
|
||||
NAME: HFEmbedder
|
||||
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_MODEL_CLS: CLIPTextModel
|
||||
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
|
||||
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder/
|
||||
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_TOKENIZER_CLS: CLIPTokenizer
|
||||
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer/
|
||||
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
|
||||
MAX_LENGTH: 77
|
||||
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
|
||||
OUTPUT_KEY: pooler_output
|
||||
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
|
||||
D_TYPE: bfloat16
|
||||
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
|
||||
BATCH_INFER: True
|
||||
CLEAN: whitespace
|
||||
#
|
||||
|
||||
SAMPLE_ARGS:
|
||||
SAMPLE_STEPS: 50
|
||||
SAMPLER: flow_eluer
|
||||
SAMPLER: flow_euler
|
||||
SEED: 2024
|
||||
IMAGE_SIZE: [ 1024, 1024 ]
|
||||
SHIFT: True
|
||||
GUIDE_SCALE: 3.5
|
||||
#
|
||||
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 4e-4
|
||||
@@ -258,7 +201,7 @@ SOLVER:
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
AMSGRAD: False
|
||||
#
|
||||
|
||||
TRAIN_DATA:
|
||||
NAME: ImageTextPairMSDataset
|
||||
MODE: train
|
||||
@@ -302,7 +245,7 @@ SOLVER:
|
||||
- NAME: Select
|
||||
KEYS: [ 'image', 'prompt' ]
|
||||
META_KEYS: [ 'data_key' ]
|
||||
#
|
||||
|
||||
EVAL_DATA:
|
||||
NAME: Text2ImageDataset
|
||||
MODE: eval
|
||||
@@ -319,13 +262,12 @@ SOLVER:
|
||||
- NAME: Select
|
||||
KEYS: [ 'index', 'prompt' ]
|
||||
META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
|
||||
TRAIN_HOOKS:
|
||||
- NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
PRIORITY: 0
|
||||
- NAME: BackwardHook
|
||||
# GRADIENT_CLIP: 1.0
|
||||
PRIORITY: 10
|
||||
- NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
|
||||
@@ -50,43 +50,33 @@ META:
|
||||
#
|
||||
SOLVER:
|
||||
NAME: LatentDiffusionSolver
|
||||
# MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
|
||||
MAX_STEPS: 100000
|
||||
# USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False
|
||||
USE_AMP: True
|
||||
# DTYPE DESCRIPTION: The precision for training. TYPE: str default: 'float32'
|
||||
DTYPE: bfloat16
|
||||
# USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
|
||||
USE_FAIRSCALE: False
|
||||
# USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
|
||||
USE_FSDP: True
|
||||
# LOAD_MODEL_ONLY DESCRIPTION: Only load the model rather than the optimizer and schedule, default is False. TYPE: bool default: False
|
||||
LOAD_MODEL_ONLY: False
|
||||
# RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: ''
|
||||
RESUME_FROM:
|
||||
WORK_DIR: ./cache/save_data/dit_flux_schnell_1024_lora
|
||||
LOG_FILE: std_log.txt
|
||||
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
|
||||
EVAL_INTERVAL: 100
|
||||
# LOG_TRAIN_NUM DESCRIPTION: The number samples used to log in training phase. TYPE: int default: -1
|
||||
LOG_TRAIN_NUM: 16
|
||||
# FSDP_REDUCE_DTYPE DESCRIPTION: The dtype of reduce in FSDP. TYPE: str default: 'float16'
|
||||
ENABLE_GRADSCALER: False
|
||||
USE_SCALER: False
|
||||
FSDP_REDUCE_DTYPE: float32
|
||||
# FSDP_BUFFER_DTYPE DESCRIPTION: The dtype of buffer in FSDP. TYPE: str default: 'float16'
|
||||
FSDP_BUFFER_DTYPE: float32
|
||||
# FSDP_SHARD_MODULES DESCRIPTION: The modules to be sharded in FSDP. TYPE: list default: ['model']
|
||||
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ] #
|
||||
SAVE_MODULES: [ 'model'] #
|
||||
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ]
|
||||
SAVE_MODULES: [ 'model']
|
||||
TRAIN_MODULES: ['model']
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/cache_data"
|
||||
#
|
||||
|
||||
FREEZE:
|
||||
#
|
||||
|
||||
TUNER:
|
||||
#
|
||||
|
||||
MODEL:
|
||||
NAME: LatentDiffusionFlux
|
||||
PARAMETERIZATION: rf
|
||||
@@ -99,65 +89,39 @@ SOLVER:
|
||||
USE_EMA: False
|
||||
EVAL_EMA: False
|
||||
DIFFUSION:
|
||||
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
|
||||
NAME: DiffusionFluxRF
|
||||
PREDICTION_TYPE: raw
|
||||
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
|
||||
NOISE_SCHEDULER:
|
||||
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
|
||||
NAME: FlowMatchSigmaScheduler
|
||||
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
|
||||
WEIGHTING_SCHEME: logit_normal
|
||||
SHIFT: 3.0
|
||||
# LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
|
||||
LOGIT_MEAN: 0.0
|
||||
# LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
|
||||
LOGIT_STD: 1.0
|
||||
# MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
|
||||
MODE_SCALE: 1.29
|
||||
SAMPLER_SCHEDULER:
|
||||
# NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
|
||||
NAME: FlowMatchFluxShiftScheduler
|
||||
# SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
|
||||
SHIFT: False
|
||||
# SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
|
||||
SIGMOID_SCALE: 1
|
||||
# BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
|
||||
BASE_SHIFT: 0.5
|
||||
# MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
|
||||
MAX_SHIFT: 1.15
|
||||
#
|
||||
|
||||
DIFFUSION_MODEL:
|
||||
# NAME DESCRIPTION: TYPE: default: 'Flux'
|
||||
NAME: Flux
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-schnell@flux1-schnell.safetensors
|
||||
# IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
|
||||
IN_CHANNELS: 64
|
||||
# HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
|
||||
HIDDEN_SIZE: 3072
|
||||
# NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16
|
||||
NUM_HEADS: 24
|
||||
# AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56]
|
||||
AXES_DIM: [ 16, 56, 56 ]
|
||||
# THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000
|
||||
THETA: 10000
|
||||
# VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768
|
||||
VEC_IN_DIM: 768
|
||||
# GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False
|
||||
GUIDANCE_EMBED: False
|
||||
# CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096
|
||||
CONTEXT_IN_DIM: 4096
|
||||
# MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0
|
||||
MLP_RATIO: 4.0
|
||||
# QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True
|
||||
QKV_BIAS: True
|
||||
# DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19
|
||||
DEPTH: 19
|
||||
# DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38
|
||||
DEPTH_SINGLE_BLOCKS: 38
|
||||
USE_GRAD_CHECKPOINT: True
|
||||
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKLFlux
|
||||
EMBED_DIM: 16
|
||||
@@ -167,7 +131,7 @@ SOLVER:
|
||||
USE_CONV: False
|
||||
SCALE_FACTOR: 0.3611
|
||||
SHIFT_FACTOR: 0.1159
|
||||
#
|
||||
|
||||
ENCODER:
|
||||
NAME: Encoder
|
||||
USE_CHECKPOINT: True
|
||||
@@ -181,7 +145,7 @@ SOLVER:
|
||||
DOUBLE_Z: True
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
#
|
||||
|
||||
DECODER:
|
||||
NAME: Decoder
|
||||
USE_CHECKPOINT: True
|
||||
@@ -196,60 +160,39 @@ SOLVER:
|
||||
RESAMP_WITH_CONV: True
|
||||
GIVE_PRE_END: False
|
||||
TANH_OUT: False
|
||||
#
|
||||
|
||||
COND_STAGE_MODEL:
|
||||
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
|
||||
NAME: T5PlusClipFluxEmbedder
|
||||
# T5_MODEL DESCRIPTION: TYPE: default: ''
|
||||
T5_MODEL:
|
||||
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
|
||||
NAME: HFEmbedder
|
||||
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_MODEL_CLS: T5EncoderModel
|
||||
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
|
||||
MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder_2/
|
||||
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_TOKENIZER_CLS: T5Tokenizer
|
||||
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer_2/
|
||||
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
|
||||
MAX_LENGTH: 256
|
||||
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
|
||||
OUTPUT_KEY: last_hidden_state
|
||||
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
|
||||
D_TYPE: bfloat16
|
||||
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
|
||||
BATCH_INFER: False
|
||||
CLEAN: whitespace
|
||||
# CLIP_MODEL DESCRIPTION: TYPE: default: ''
|
||||
CLIP_MODEL:
|
||||
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
|
||||
NAME: HFEmbedder
|
||||
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_MODEL_CLS: CLIPTextModel
|
||||
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
|
||||
MODEL_PATH: ms://AI-ModelScope/FLUX.1-schnell@text_encoder/
|
||||
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_TOKENIZER_CLS: CLIPTokenizer
|
||||
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-schnell@tokenizer/
|
||||
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
|
||||
MAX_LENGTH: 77
|
||||
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
|
||||
OUTPUT_KEY: pooler_output
|
||||
# D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16'
|
||||
D_TYPE: bfloat16
|
||||
# BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False
|
||||
BATCH_INFER: True
|
||||
CLEAN: whitespace
|
||||
#
|
||||
|
||||
SAMPLE_ARGS:
|
||||
SAMPLE_STEPS: 4
|
||||
SAMPLER: flow_eluer
|
||||
SAMPLER: flow_euler
|
||||
SEED: 2024
|
||||
IMAGE_SIZE: [ 1024, 1024 ]
|
||||
GUIDE_SCALE: 3.5
|
||||
#
|
||||
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 4e-4
|
||||
@@ -257,7 +200,7 @@ SOLVER:
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
AMSGRAD: False
|
||||
#
|
||||
|
||||
TRAIN_DATA:
|
||||
NAME: ImageTextPairMSDataset
|
||||
MODE: train
|
||||
@@ -301,7 +244,7 @@ SOLVER:
|
||||
- NAME: Select
|
||||
KEYS: [ 'image', 'prompt' ]
|
||||
META_KEYS: [ 'data_key' ]
|
||||
#
|
||||
|
||||
EVAL_DATA:
|
||||
NAME: Text2ImageDataset
|
||||
MODE: eval
|
||||
@@ -318,13 +261,12 @@ SOLVER:
|
||||
- NAME: Select
|
||||
KEYS: [ 'index', 'prompt' ]
|
||||
META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
|
||||
TRAIN_HOOKS:
|
||||
- NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
PRIORITY: 0
|
||||
- NAME: BackwardHook
|
||||
# GRADIENT_CLIP: 1.0
|
||||
PRIORITY: 10
|
||||
- NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
@@ -339,4 +281,4 @@ SOLVER:
|
||||
PROB_INTERVAL: 100
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
SAVE_PROBE_PREFIX: 'image'
|
||||
SAVE_PROBE_PREFIX: 'image'
|
||||
@@ -35,7 +35,7 @@ META:
|
||||
SAVE_INTERVAL: 25
|
||||
EPSEC: 0.818
|
||||
LEARNING_RATE: 0.0001
|
||||
IS_DEFAULT: False
|
||||
IS_DEFAULT: True
|
||||
TUNER: LORA
|
||||
#
|
||||
TUNERS:
|
||||
|
||||
@@ -13,8 +13,10 @@ TRAIN_PARAS:
|
||||
VALUES: [[256, 256], [320, 180], [180, 320],
|
||||
[512, 512], [640, 360], [360, 640],
|
||||
[768, 768], [960, 540], [540, 960],
|
||||
[1024, 1024], [1280, 720], [720, 1280]]
|
||||
[1024, 1024], [1280, 720], [720, 1280],
|
||||
[720, 480], [480, 720]]
|
||||
DEFAULT: [1024, 1024]
|
||||
EVAL_PROMPTS:
|
||||
- a boy wearing a jacket
|
||||
- a dog running on the lawn
|
||||
SAVE_FILE_LOCAL_PATH: "cache/scepter_ui/datasets/train_data_from_list"
|
||||
|
||||
@@ -1,20 +1,11 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
import math
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from abc import ABCMeta
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms as TT
|
||||
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
|
||||
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
|
||||
@@ -2,20 +2,17 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from abc import ABCMeta
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torchvision
|
||||
import torchvision.transforms as transforms
|
||||
from einops import rearrange
|
||||
from PIL import Image
|
||||
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
|
||||
from torchvision.transforms import InterpolationMode
|
||||
|
||||
norm_layer = nn.InstanceNorm2d
|
||||
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import math
|
||||
import random
|
||||
from abc import ABCMeta
|
||||
from enum import Enum
|
||||
@@ -8,7 +7,8 @@ from enum import Enum
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image, ImageDraw
|
||||
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
|
||||
@@ -226,7 +226,12 @@ class InpaintingAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
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):
|
||||
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
|
||||
@@ -247,18 +252,17 @@ class InpaintingAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
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)
|
||||
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)
|
||||
image = np.where(mask[:, :, np.newaxis] == 255, colored_mask,
|
||||
image)
|
||||
|
||||
if return_mask:
|
||||
ret_data = {
|
||||
'image': np.array(image),
|
||||
'mask': np.array(mask)
|
||||
}
|
||||
ret_data = {'image': np.array(image), 'mask': np.array(mask)}
|
||||
else:
|
||||
ret_data = np.array(image)
|
||||
return ret_data
|
||||
|
||||
@@ -1,27 +1,31 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from abc import ABCMeta
|
||||
|
||||
import torch
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
import cv2
|
||||
import torch
|
||||
from PIL import Image
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
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
|
||||
)
|
||||
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
|
||||
@@ -32,26 +36,29 @@ class LamaAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
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")
|
||||
@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'
|
||||
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()
|
||||
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)
|
||||
@@ -69,9 +76,13 @@ class LamaAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
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))
|
||||
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)
|
||||
@@ -93,6 +104,3 @@ class LamaAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
__class__.__name__,
|
||||
LamaAnnotator.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -789,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)
|
||||
|
||||
@@ -4,10 +4,10 @@ import math
|
||||
import random
|
||||
from abc import ABCMeta
|
||||
|
||||
import cv2
|
||||
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
|
||||
@@ -87,7 +87,8 @@ class OutpaintingAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
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)
|
||||
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))
|
||||
@@ -108,20 +109,30 @@ class OutpaintingAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
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_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 = 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)}
|
||||
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)}
|
||||
ret_data = {
|
||||
'src_image': np.array(init_image),
|
||||
'image': np.array(img)
|
||||
}
|
||||
else:
|
||||
ret_data = np.array(img)
|
||||
return ret_data
|
||||
@@ -133,6 +144,7 @@ class OutpaintingAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
OutpaintingAnnotator.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class OutpaintingResize(BaseAnnotator, metaclass=ABCMeta):
|
||||
para_dict = {}
|
||||
@@ -148,11 +160,7 @@ class OutpaintingResize(BaseAnnotator, metaclass=ABCMeta):
|
||||
top, bottom = np.min(locs[0]), np.max(locs[0])
|
||||
return [left, top, right, bottom]
|
||||
|
||||
def forward(self,
|
||||
image,
|
||||
target_image,
|
||||
mask=None
|
||||
):
|
||||
def forward(self, image, target_image, mask=None):
|
||||
if isinstance(image, Image.Image):
|
||||
image = image
|
||||
elif isinstance(image, torch.Tensor):
|
||||
@@ -175,8 +183,10 @@ class OutpaintingResize(BaseAnnotator, metaclass=ABCMeta):
|
||||
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')
|
||||
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
|
||||
|
||||
@@ -1,14 +1,13 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import math
|
||||
from abc import ABCMeta
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms as TT
|
||||
from einops import rearrange
|
||||
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import math
|
||||
import random
|
||||
from abc import ABCMeta
|
||||
|
||||
@@ -9,15 +8,16 @@ import numpy as np
|
||||
import torch
|
||||
import torchvision.transforms as T
|
||||
from PIL import Image
|
||||
from scipy import ndimage
|
||||
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
|
||||
from sklearn.cluster import KMeans
|
||||
from torchvision.ops.boxes import batched_nms
|
||||
|
||||
|
||||
def find_dominant_color(image, k=1):
|
||||
@@ -28,7 +28,7 @@ def find_dominant_color(image, k=1):
|
||||
kmeans = KMeans(n_clusters=k, n_init='auto')
|
||||
kmeans.fit(pixels)
|
||||
dominant_color = kmeans.cluster_centers_.astype(int)[0]
|
||||
except:
|
||||
except Exception:
|
||||
dominant_color = np.array([255, 255, 255])
|
||||
return dominant_color
|
||||
|
||||
@@ -62,16 +62,9 @@ class ESAMAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
super().__init__(cfg, logger=logger)
|
||||
try:
|
||||
from efficient_sam.efficient_sam import build_efficient_sam
|
||||
from segment_anything.utils.amg import (
|
||||
batched_mask_to_box,
|
||||
calculate_stability_score,
|
||||
mask_to_rle_pytorch,
|
||||
remove_small_regions,
|
||||
rle_to_mask,
|
||||
)
|
||||
except:
|
||||
except Exception:
|
||||
raise NotImplementedError(
|
||||
f'Please install efficient_sam and segment_anything modules.')
|
||||
'Please install efficient_sam and segment_anything modules.')
|
||||
|
||||
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
|
||||
if pretrained_model:
|
||||
@@ -294,8 +287,6 @@ class ESAMAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
set_name=True)
|
||||
|
||||
|
||||
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class SAMAnnotatorDraw(BaseAnnotator, metaclass=ABCMeta):
|
||||
para_dict = {}
|
||||
@@ -312,10 +303,16 @@ class SAMAnnotatorDraw(BaseAnnotator, metaclass=ABCMeta):
|
||||
|
||||
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)
|
||||
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):
|
||||
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):
|
||||
@@ -337,21 +334,25 @@ class SAMAnnotatorDraw(BaseAnnotator, metaclass=ABCMeta):
|
||||
else:
|
||||
raise f'Unsurpport datatype{type(mask)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||
|
||||
original_size = image.shape[:2]
|
||||
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))
|
||||
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}
|
||||
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 = ndimage.center_of_mass(scribble, labeled_array,
|
||||
range(1, num_features + 1))
|
||||
centers = np.array(centers)
|
||||
### (x1, y1, x2, y2)
|
||||
# (x1, y1, x2, y2)
|
||||
x_min = centers[:, 0].min()
|
||||
x_max = centers[:, 0].max()
|
||||
y_min = centers[:, 1].min()
|
||||
@@ -365,12 +366,13 @@ class SAMAnnotatorDraw(BaseAnnotator, metaclass=ABCMeta):
|
||||
sample = {'box': input_box}
|
||||
|
||||
self.sam_predictor.set_image(image)
|
||||
masks, scores, logits = self.sam_predictor.predict(**sample, multimask_output=True)
|
||||
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]
|
||||
'mask': (masks[index] * 255).astype(np.uint8),
|
||||
'score': scores[index]
|
||||
}
|
||||
return ret_data
|
||||
|
||||
@@ -379,4 +381,4 @@ class SAMAnnotatorDraw(BaseAnnotator, metaclass=ABCMeta):
|
||||
return dict_to_yaml('ANNOTATORS',
|
||||
__class__.__name__,
|
||||
SAMAnnotatorDraw.para_dict,
|
||||
set_name=True)
|
||||
set_name=True)
|
||||
|
||||
@@ -1,14 +1,11 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
import math
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from abc import ABCMeta
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms as TT
|
||||
from einops import rearrange
|
||||
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
|
||||
@@ -9,3 +9,4 @@ from scepter.modules.data.dataset.dataset import (Image2ImageDataset,
|
||||
from scepter.modules.data.dataset.ms_dataset import (
|
||||
ImageTextPairFolderDataset, ImageTextPairMSDataset)
|
||||
from scepter.modules.data.dataset.registry import DATASETS
|
||||
from scepter.modules.data.dataset.video_gen_dataset import VideoGenDataset
|
||||
@@ -242,6 +242,8 @@ class Text2ImageDataset(BaseDataset):
|
||||
prompt_prefix = cfg.get('PROMPT_PREFIX', '')
|
||||
path_prefix = cfg.get('PATH_PREFIX', '')
|
||||
use_num = cfg.get('USE_NUM', -1)
|
||||
meta_cfg = cfg.get('META_CFG', None)
|
||||
meta_cfg = meta_cfg.get_lowercase_dict() if meta_cfg is not None else None
|
||||
|
||||
image_size = cfg.get('IMAGE_SIZE', 1024)
|
||||
if isinstance(image_size, numbers.Number):
|
||||
@@ -264,7 +266,12 @@ class Text2ImageDataset(BaseDataset):
|
||||
|
||||
self.items = list()
|
||||
for i, row in enumerate(rows):
|
||||
item = {'index': i, 'meta': {'image_size': image_size}}
|
||||
if meta_cfg is not None:
|
||||
meta_cfg_copy = copy.deepcopy(meta_cfg)
|
||||
meta_cfg_copy['image_size'] = image_size
|
||||
item = {'index': i, 'meta': meta_cfg_copy}
|
||||
else:
|
||||
item = {'index': i, 'meta': {'image_size': image_size}}
|
||||
for key, value in zip(fields, row):
|
||||
if key in ['prompt', 'caption', 'text']:
|
||||
item['ori_prompt'] = value
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1,184 @@
|
||||
import io
|
||||
import random
|
||||
import sys
|
||||
import os
|
||||
import warnings
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.data.dataset import DATASETS, BaseDataset
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
try:
|
||||
import decord
|
||||
decord.bridge.set_bridge("torch")
|
||||
except ImportError:
|
||||
warnings.warn(
|
||||
"The `decord` package is required for loading the video dataset. Install with `pip install decord`"
|
||||
)
|
||||
|
||||
|
||||
@DATASETS.register_class()
|
||||
class VideoGenDataset(BaseDataset):
|
||||
def __init__(self, cfg, logger = None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.prompt_prefix = cfg.get('PROMPT_PREFIX', '')
|
||||
self.path_prefix = cfg.get('PATH_PREFIX', '')
|
||||
self.p_zero = cfg.get('P_ZERO', 0.0)
|
||||
self.max_num_frames = cfg.get("NUM_FRAMES", 49)
|
||||
self.fps = cfg.get("FPS", 8)
|
||||
self.height = cfg.get("HEIGHT", 480)
|
||||
self.width = cfg.get("WIDTH", 720)
|
||||
self.skip_frames_start = cfg.get("SKIP_FRAMES_START", 0)
|
||||
self.skip_frames_end = cfg.get("SKIP_FRAMES_END", 0)
|
||||
self.data_type = cfg.get('DATA_TYPE', 't2v')
|
||||
|
||||
def worker_init_fn(self, worker_id, num_workers=1):
|
||||
super().worker_init_fn(worker_id, num_workers=num_workers)
|
||||
randseed = np.random.randint(0, 2 ** 32 - num_workers - 1)
|
||||
workerseed = randseed + worker_id
|
||||
random.seed(workerseed)
|
||||
np.random.seed(workerseed)
|
||||
|
||||
def _preprocess_video_data(self, video_path):
|
||||
|
||||
with FS.get_object(video_path) as video_data:
|
||||
video_reader = decord.VideoReader(io.BytesIO(video_data), width=self.width, height=self.height)
|
||||
video_num_frames = len(video_reader)
|
||||
|
||||
start_frame = min(self.skip_frames_start, video_num_frames)
|
||||
end_frame = max(0, video_num_frames - self.skip_frames_end)
|
||||
if end_frame <= start_frame:
|
||||
frames = video_reader.get_batch([start_frame])
|
||||
elif end_frame - start_frame <= self.max_num_frames:
|
||||
frames = video_reader.get_batch(list(range(start_frame, end_frame)))
|
||||
else:
|
||||
indices = list(range(start_frame, end_frame, (end_frame - start_frame) // self.max_num_frames))
|
||||
frames = video_reader.get_batch(indices)
|
||||
|
||||
# Ensure that we don't go over the limit
|
||||
frames = frames[: self.max_num_frames]
|
||||
selected_num_frames = frames.shape[0]
|
||||
|
||||
# Choose first (4k + 1) frames as this is how many is required by the VAE
|
||||
remainder = (3 + (selected_num_frames % 4)) % 4
|
||||
if remainder != 0:
|
||||
frames = frames[:-remainder]
|
||||
selected_num_frames = frames.shape[0]
|
||||
|
||||
assert (selected_num_frames - 1) % 4 == 0
|
||||
|
||||
# Training transforms
|
||||
frames = frames.float().div_(127.5).sub_(1.)
|
||||
frames = frames.permute(3, 0, 1, 2).contiguous() # [C, F, H, W]
|
||||
return frames
|
||||
|
||||
def _parse_index(self, index):
|
||||
meta = dict()
|
||||
for key, value in zip(index[-1], index[:-1]):
|
||||
if key in ['oss_key', 'path', 'video_path']:
|
||||
meta['video_path'] = value
|
||||
elif key in ['prompt', 'caption', 'text']:
|
||||
meta['prompt'] = value
|
||||
elif key in ['width', 'height']:
|
||||
meta[key] = int(value)
|
||||
else:
|
||||
meta[key] = value
|
||||
return meta
|
||||
|
||||
def _get(self, index):
|
||||
meta = self._parse_index(index)
|
||||
|
||||
video_path = os.path.join(self.path_prefix, meta.get('video_path', ''))
|
||||
video = self._preprocess_video_data(video_path)
|
||||
|
||||
prompt = self.prompt_prefix + meta.get('prompt', '')
|
||||
if self.mode == 'train' and np.random.uniform() < self.p_zero:
|
||||
prompt = ''
|
||||
|
||||
item = {
|
||||
'video': video,
|
||||
'prompt': prompt,
|
||||
'meta': meta,
|
||||
}
|
||||
if self.data_type == 'i2v':
|
||||
item['image'] = item['video'][:, :1, :, :]
|
||||
return item
|
||||
|
||||
def __len__(self):
|
||||
return sys.maxsize
|
||||
|
||||
@staticmethod
|
||||
def collate_fn(batch):
|
||||
collect = {}
|
||||
for sample in batch:
|
||||
for k, v in sample.items():
|
||||
if k not in collect:
|
||||
collect[k] = []
|
||||
collect[k].append(v)
|
||||
return collect
|
||||
|
||||
|
||||
|
||||
@DATASETS.register_class()
|
||||
class VideoGenDatasetOTF(VideoGenDataset):
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger)
|
||||
self.data_file = cfg.DATA_FILE
|
||||
self.delimiter = cfg.get('DELIMITER', '#;#')
|
||||
self.fields = cfg.get('FIELDS', ['video_path', 'prompt'])
|
||||
self.use_num = cfg.get('USE_NUM', -1)
|
||||
|
||||
from scepter.modules.model.registry import MODELS
|
||||
model_cfg = cfg.get('MODEL', None)
|
||||
if model_cfg is not None:
|
||||
self.model = MODELS.build(cfg.MODEL, logger=logger).eval().requires_grad_(False).to(we.device_id)
|
||||
self.items = self.parse_data(self.data_file, self.delimiter, self.fields)
|
||||
if self.use_num and self.use_num > 0:
|
||||
self.items = self.items[:self.use_num]
|
||||
self.data = self.encode(self.items)
|
||||
self.real_number = len(self.data)
|
||||
if model_cfg is not None:
|
||||
self.model.to('cpu')
|
||||
del self.model
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def parse_data(self, data_file, delimiter, fields):
|
||||
items = list()
|
||||
with FS.get_object(data_file) as local_data:
|
||||
rows = [
|
||||
i.split(delimiter,
|
||||
len(fields) - 1)
|
||||
for i in local_data.decode('utf-8').strip().split('\n')
|
||||
]
|
||||
for i, row in enumerate(rows):
|
||||
item = {}
|
||||
for key, value in zip(self.fields, row):
|
||||
if key in ['oss_key', 'path', 'video_path']:
|
||||
item['video_path'] = value
|
||||
elif key in ['prompt', 'caption', 'text']:
|
||||
item['prompt'] = value
|
||||
elif key in ['width', 'height']:
|
||||
item[key] = int(value)
|
||||
else:
|
||||
item[key] = value
|
||||
items.append(item)
|
||||
return items
|
||||
|
||||
def encode(self, items):
|
||||
self.logger.info("Start to encode video data [{}]!".format(len(items)))
|
||||
for item in tqdm(items):
|
||||
video_path = os.path.join(self.path_prefix, item.get('video_path', ''))
|
||||
video = self._preprocess_video_data(video_path)
|
||||
latent = self.model.encode_first_stage(video.unsqueeze(0).to(we.device_id)).squeeze(0)
|
||||
item['video_latent'] = latent.detach().cpu()
|
||||
item['video'] = video
|
||||
if self.data_type == 'i2v':
|
||||
item['image'] = item['video'][:, :1, :, :]
|
||||
return items
|
||||
|
||||
def _get(self, index):
|
||||
return self.data[index % self.real_number]
|
||||
@@ -0,0 +1,551 @@
|
||||
# -*- 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
|
||||
import torchvision.transforms as T
|
||||
from scepter.modules.model.registry import DIFFUSIONS
|
||||
from scepter.modules.model.utils.basic_utils import check_list_of_list
|
||||
from scepter.modules.model.utils.basic_utils import \
|
||||
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 RefinerInference(DiffusionInference):
|
||||
def init_from_cfg(self, cfg):
|
||||
self.use_dynamic_model = cfg.get('USE_DYNAMIC_MODEL', True)
|
||||
super().init_from_cfg(cfg)
|
||||
self.diffusion = DIFFUSIONS.build(cfg.MODEL.DIFFUSION, logger=self.logger) \
|
||||
if cfg.MODEL.have('DIFFUSION') else None
|
||||
self.max_seq_length = cfg.MODEL.get("MAX_SEQ_LENGTH", 4096)
|
||||
assert self.diffusion is not None
|
||||
if not self.use_dynamic_model:
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
@torch.no_grad()
|
||||
def encode_first_stage(self, x, **kwargs):
|
||||
_, dtype = self.get_function_info(self.first_stage_model, 'encode')
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype in ('float16', 'bfloat16'),
|
||||
dtype=getattr(torch, dtype)):
|
||||
def run_one_image(u):
|
||||
zu = get_model(self.first_stage_model).encode(u)
|
||||
if isinstance(zu, (tuple, list)):
|
||||
zu = zu[0]
|
||||
return zu
|
||||
z = [run_one_image(u.unsqueeze(0) if u.dim == 3 else u) for u in x]
|
||||
return z
|
||||
def upscale_resize(self, image, interpolation=T.InterpolationMode.BILINEAR):
|
||||
c, H, W = image.shape
|
||||
scale = max(1.0, math.sqrt(self.max_seq_length / ((H / 16) * (W / 16))))
|
||||
rH = int(H * scale) // 16 * 16 # ensure divisible by self.d
|
||||
rW = int(W * scale) // 16 * 16
|
||||
image = T.Resize((rH, rW), interpolation=interpolation, antialias=True)(image)
|
||||
return image
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, z):
|
||||
_, dtype = self.get_function_info(self.first_stage_model, 'decode')
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype in ('float16', 'bfloat16'),
|
||||
dtype=getattr(torch, dtype)):
|
||||
return [get_model(self.first_stage_model).decode(zu) for zu in z]
|
||||
|
||||
def noise_sample(self, num_samples, h, w, seed, device = None, dtype = torch.bfloat16):
|
||||
noise = torch.randn(
|
||||
num_samples,
|
||||
16,
|
||||
# allow for packing
|
||||
2 * math.ceil(h / 16),
|
||||
2 * math.ceil(w / 16),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
generator=torch.Generator(device=device).manual_seed(seed),
|
||||
)
|
||||
return noise
|
||||
def refine(self,
|
||||
x_samples=None,
|
||||
prompt=None,
|
||||
reverse_scale=-1.,
|
||||
seed = 2024,
|
||||
**kwargs
|
||||
):
|
||||
print(prompt)
|
||||
value_input = copy.deepcopy(self.input)
|
||||
x_samples = [self.upscale_resize(x) for x in x_samples]
|
||||
|
||||
noise = []
|
||||
for i, x in enumerate(x_samples):
|
||||
noise_ = self.noise_sample(1, x.shape[1],
|
||||
x.shape[2], seed,
|
||||
device = x.device)
|
||||
noise.append(noise_)
|
||||
noise, x_shapes = pack_imagelist_into_tensor(noise)
|
||||
if reverse_scale > 0:
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
x_samples = [x.unsqueeze(0) for x in x_samples]
|
||||
x_start = self.encode_first_stage(x_samples, **kwargs)
|
||||
self.dynamic_unload(self.first_stage_model,
|
||||
'first_stage_model',
|
||||
skip_loaded=not self.use_dynamic_model)
|
||||
x_start, _ = pack_imagelist_into_tensor(x_start)
|
||||
else:
|
||||
x_start = None
|
||||
# cond stage
|
||||
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
|
||||
function_name, dtype = self.get_function_info(self.cond_stage_model)
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
ctx = getattr(get_model(self.cond_stage_model),
|
||||
function_name)(prompt)
|
||||
ctx["x_shapes"] = x_shapes
|
||||
self.dynamic_unload(self.cond_stage_model,
|
||||
'cond_stage_model',
|
||||
skip_loaded=not self.use_dynamic_model)
|
||||
|
||||
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
# UNet use input n_prompt
|
||||
function_name, dtype = self.get_function_info(
|
||||
self.diffusion_model)
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype in ('float16', 'bfloat16'),
|
||||
dtype=getattr(torch, dtype)):
|
||||
solver_sample = value_input.get('sample', 'flow_euler')
|
||||
sample_steps = value_input.get('sample_steps', 20)
|
||||
guide_scale = value_input.get('guide_scale', 3.5)
|
||||
if guide_scale is not None:
|
||||
guide_scale = torch.full((noise.shape[0],), guide_scale, device=noise.device,
|
||||
dtype=noise.dtype)
|
||||
else:
|
||||
guide_scale = None
|
||||
latent = self.diffusion.sample(
|
||||
noise=noise,
|
||||
sampler=solver_sample,
|
||||
model=get_model(self.diffusion_model),
|
||||
model_kwargs={"cond": ctx, "guidance": guide_scale},
|
||||
steps=sample_steps,
|
||||
show_progress=True,
|
||||
guide_scale=guide_scale,
|
||||
return_intermediate=None,
|
||||
reverse_scale=reverse_scale,
|
||||
x=x_start,
|
||||
**kwargs).float()
|
||||
latent = unpack_tensor_into_imagelist(latent, x_shapes)
|
||||
self.dynamic_unload(self.diffusion_model,
|
||||
'diffusion_model',
|
||||
skip_loaded=not self.use_dynamic_model)
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
x_samples = self.decode_first_stage(latent)
|
||||
self.dynamic_unload(self.first_stage_model,
|
||||
'first_stage_model',
|
||||
skip_loaded=not self.use_dynamic_model)
|
||||
return x_samples
|
||||
|
||||
|
||||
class ACEInference(DiffusionInference):
|
||||
def __init__(self, logger=None):
|
||||
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)
|
||||
self.use_dynamic_model = cfg.get('USE_DYNAMIC_MODEL', True)
|
||||
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.refiner_model_cfg = cfg.get('REFINER_MODEL', None)
|
||||
# self.refiner_scale = cfg.get('REFINER_SCALE', 0.)
|
||||
# self.refiner_prompt = cfg.get('REFINER_PROMPT', "")
|
||||
self.ace_prompt = cfg.get("ACE_PROMPT", [])
|
||||
if self.refiner_model_cfg:
|
||||
self.refiner_model_cfg.USE_DYNAMIC_MODEL = self.use_dynamic_model
|
||||
self.refiner_module = RefinerInference(self.logger)
|
||||
self.refiner_module.init_from_cfg(self.refiner_model_cfg)
|
||||
else:
|
||||
self.refiner_module = None
|
||||
|
||||
self.diffusion = DIFFUSIONS.build(cfg.MODEL.DIFFUSION,
|
||||
logger=self.logger)
|
||||
|
||||
|
||||
self.interpolate_func = lambda x: (F.interpolate(
|
||||
x.unsqueeze(0),
|
||||
scale_factor=1 / self.size_factor,
|
||||
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', '')
|
||||
if not self.use_dynamic_model:
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_first_stage(self, x, **kwargs):
|
||||
_, dtype = self.get_function_info(self.first_stage_model, 'encode')
|
||||
with torch.autocast('cuda',
|
||||
enabled=(dtype != '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
|
||||
|
||||
is_txt_image = sum([len(e_i) for e_i in edit_image]) < 1
|
||||
image = to_device(image)
|
||||
|
||||
refiner_scale = kwargs.pop("refiner_scale", 0.0)
|
||||
refiner_prompt = kwargs.pop("refiner_prompt", "")
|
||||
use_ace = kwargs.pop("use_ace", True)
|
||||
# <= 0 use ace as the txt2img generator.
|
||||
if use_ace and (not is_txt_image or refiner_scale <= 0):
|
||||
ctx, null_ctx = {}, {}
|
||||
# Get Noise Shape
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
x = self.encode_first_stage(image)
|
||||
self.dynamic_unload(self.first_stage_model,
|
||||
'first_stage_model',
|
||||
skip_loaded=not self.use_dynamic_model)
|
||||
noise = [
|
||||
torch.empty(*i.shape, device=we.device_id).normal_(generator=g)
|
||||
for i in x
|
||||
]
|
||||
noise, x_shapes = pack_imagelist_into_tensor(noise)
|
||||
ctx['x_shapes'] = null_ctx['x_shapes'] = x_shapes
|
||||
|
||||
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=not self.use_dynamic_model)
|
||||
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=not self.use_dynamic_model)
|
||||
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=not self.use_dynamic_model)
|
||||
|
||||
# Decode to Pixel Space
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
samples = unpack_tensor_into_imagelist(latent, x_shapes)
|
||||
x_samples = self.decode_first_stage(samples)
|
||||
self.dynamic_unload(self.first_stage_model,
|
||||
'first_stage_model',
|
||||
skip_loaded=not self.use_dynamic_model)
|
||||
x_samples = [x.squeeze(0) for x in x_samples]
|
||||
else:
|
||||
x_samples = image
|
||||
if self.refiner_module and refiner_scale > 0:
|
||||
if is_txt_image:
|
||||
random.shuffle(self.ace_prompt)
|
||||
input_refine_prompt = [self.ace_prompt[0] + refiner_prompt if p[0] == "" else p[0] for p in prompt]
|
||||
input_refine_scale = -1.
|
||||
else:
|
||||
input_refine_prompt = [p[0].replace("{image}", "") + " " + refiner_prompt for p in prompt]
|
||||
input_refine_scale = refiner_scale
|
||||
print(input_refine_prompt)
|
||||
|
||||
x_samples = self.refiner_module.refine(x_samples,
|
||||
reverse_scale = input_refine_scale,
|
||||
prompt= input_refine_prompt,
|
||||
seed=seed,
|
||||
use_dynamic_model=self.use_dynamic_model)
|
||||
|
||||
imgs = [
|
||||
torch.clamp((x_i.float() + 1.0) / 2.0 + self.decoder_bias / 255,
|
||||
min=0.0,
|
||||
max=1.0).squeeze(0).permute(1, 2, 0).cpu().numpy()
|
||||
for x_i in x_samples
|
||||
]
|
||||
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_
|
||||
@@ -0,0 +1,181 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import copy
|
||||
import numpy as np
|
||||
from typing import Tuple
|
||||
import random
|
||||
|
||||
import torch
|
||||
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.model.backbone.cogvideox.utils import get_3d_rotary_pos_embed, get_resize_crop_region_for_grid
|
||||
|
||||
from .diffusion_inference import DiffusionInference, get_model
|
||||
from .tuner_inference import TunerInference
|
||||
|
||||
class CogVideoXInference(DiffusionInference):
|
||||
def __init__(self, logger=None):
|
||||
self.logger = logger
|
||||
self.is_redefine_paras = False
|
||||
self.loaded_model = {}
|
||||
self.loaded_model_name = [
|
||||
'diffusion_model', 'first_stage_model', 'cond_stage_model'
|
||||
]
|
||||
self.tuner_infer = TunerInference(self.logger)
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, latents):
|
||||
latents = latents.permute(0, 2, 1, 3, 4)
|
||||
latents = 1 / self.first_stage_model['paras']['scaling_factor_image'] * latents
|
||||
frames = get_model(self.first_stage_model).decode(latents)
|
||||
return frames
|
||||
|
||||
def _prepare_rotary_positional_embeddings(
|
||||
self,
|
||||
height: int,
|
||||
width: int,
|
||||
num_frames: int,
|
||||
device: torch.device,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
|
||||
grid_height = height // (self.diffusion_model['paras']['scale_factor_spatial'] * self.diffusion_model['paras']['patch_size'])
|
||||
grid_width = width // (self.diffusion_model['paras']['scale_factor_spatial'] * self.diffusion_model['paras']['patch_size'])
|
||||
base_size_width = self.diffusion_model['paras']['sample_width'] // (self.diffusion_model['paras']['scale_factor_spatial'] * self.diffusion_model['paras']['patch_size'])
|
||||
base_size_height = self.diffusion_model['paras']['sample_height'] // (self.diffusion_model['paras']['scale_factor_spatial'] * self.diffusion_model['paras']['patch_size'])
|
||||
|
||||
grid_crops_coords = get_resize_crop_region_for_grid(
|
||||
(grid_height, grid_width), base_size_width, base_size_height
|
||||
)
|
||||
freqs_cos, freqs_sin = get_3d_rotary_pos_embed(
|
||||
embed_dim=self.diffusion_model['paras']['attention_head_dim'],
|
||||
crops_coords=grid_crops_coords,
|
||||
grid_size=(grid_height, grid_width),
|
||||
temporal_size=num_frames,
|
||||
)
|
||||
|
||||
freqs_cos = freqs_cos.to(device=device)
|
||||
freqs_sin = freqs_sin.to(device=device)
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(self,
|
||||
input,
|
||||
num_samples=1,
|
||||
cat_uc=True,
|
||||
tuner_model=None,
|
||||
**kwargs):
|
||||
value_input = copy.deepcopy(self.input)
|
||||
value_input.update(input)
|
||||
print(value_input)
|
||||
height, width = value_input['target_size_as_tuple']
|
||||
value_output = copy.deepcopy(self.output)
|
||||
|
||||
# register tuner
|
||||
if tuner_model is not None and tuner_model != '' and len(
|
||||
tuner_model) > 0:
|
||||
if not isinstance(tuner_model, list):
|
||||
tuner_model = [tuner_model]
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
self.tuner_infer.register_tuner(tuner_model, self.diffusion_model,
|
||||
cond_stage_model=None)
|
||||
self.dynamic_unload(self.diffusion_model,
|
||||
'diffusion_model',
|
||||
skip_loaded=True)
|
||||
|
||||
# cond stage
|
||||
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
|
||||
function_name, dtype = self.get_function_info(self.cond_stage_model)
|
||||
with torch.autocast(device_type='cuda', enabled=True, dtype=torch.bfloat16):
|
||||
cont = getattr(get_model(self.cond_stage_model),
|
||||
function_name)(value_input['prompt'], return_mask=False, use_mask=False)
|
||||
null_cont = getattr(get_model(self.cond_stage_model),
|
||||
function_name)(value_input['negative_prompt'] * num_samples, return_mask=False, use_mask=False)
|
||||
self.dynamic_unload(self.cond_stage_model,
|
||||
'cond_stage_model',
|
||||
skip_loaded=True)
|
||||
|
||||
# get noise
|
||||
seed = kwargs.pop('seed', -1)
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||
generator = torch.Generator().manual_seed(seed)
|
||||
if 'seed' in value_output:
|
||||
value_output['seed'] = seed
|
||||
for sample_id in range(num_samples):
|
||||
if self.diffusion_model is not None:
|
||||
noise_shape = (1,
|
||||
(value_input['num_frames'] - 1) // self.diffusion_model['paras']['scale_factor_temporal'] + 1,
|
||||
self.diffusion_model['paras']['latent_channels'],
|
||||
height // self.diffusion_model['paras']['scale_factor_spatial'],
|
||||
width // self.diffusion_model['paras']['scale_factor_spatial']
|
||||
)
|
||||
noise = torch.randn(noise_shape, generator=generator, dtype=getattr(torch, dtype), device='cpu').to(we.device_id)
|
||||
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
|
||||
image_rotary_emb = (
|
||||
self._prepare_rotary_positional_embeddings(height, width, noise.size(1), we.device_id)
|
||||
if self.diffusion_model['paras']['use_rotary_positional_embeddings']
|
||||
else None
|
||||
)
|
||||
function_name, dtype = self.get_function_info(
|
||||
self.diffusion_model)
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype=='bfloat16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
solver_sample = value_input.get('sample', 'ddim')
|
||||
sample_steps = value_input.get('sample_steps', 50)
|
||||
guide_scale = value_input.get('guide_scale', 7.5)
|
||||
guide_rescale = value_input.get('guide_rescale', 0.5)
|
||||
|
||||
latent = self.diffusion.sample(noise=noise,
|
||||
sampler=solver_sample,
|
||||
model=get_model(self.diffusion_model),
|
||||
model_kwargs=[{
|
||||
'cond': cont,
|
||||
'image_latent': None,
|
||||
'image_rotary_emb': image_rotary_emb,
|
||||
}, {
|
||||
'cond': null_cont,
|
||||
'image_latent': None,
|
||||
'image_rotary_emb': image_rotary_emb,
|
||||
}],
|
||||
steps=sample_steps,
|
||||
show_progress=True,
|
||||
use_dynamic_cfg=True,
|
||||
guide_scale=guide_scale,
|
||||
guide_rescale=guide_rescale,
|
||||
return_intermediate=None,
|
||||
**kwargs).float()
|
||||
self.dynamic_unload(self.diffusion_model,
|
||||
'diffusion_model',
|
||||
skip_loaded=True)
|
||||
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
x_samples = self.decode_first_stage(latent).float() # [B, C, F, H, W]
|
||||
self.dynamic_unload(self.first_stage_model,
|
||||
'first_stage_model',
|
||||
skip_loaded=True)
|
||||
|
||||
x_frames = torch.clamp(x_samples / 2 + 0.5, min=0.0, max=1.0)
|
||||
if 'videos' in value_output:
|
||||
if value_output['videos'] is None or (
|
||||
isinstance(value_output['videos'], list)
|
||||
and len(value_output['videos']) < 1):
|
||||
value_output['videos'] = []
|
||||
value_output['videos'].append(x_frames)
|
||||
|
||||
for k, v in value_output.items():
|
||||
if isinstance(v, list):
|
||||
value_output[k] = torch.cat(v, dim=0)
|
||||
if isinstance(v, torch.Tensor):
|
||||
value_output[k] = v.cpu()
|
||||
|
||||
# unregister tuner
|
||||
if tuner_model is not None and tuner_model != '' and len(
|
||||
tuner_model) > 0:
|
||||
self.tuner_infer.unregister_tuner(tuner_model,
|
||||
self.diffusion_model,
|
||||
cond_stage_model=None)
|
||||
return value_output
|
||||
@@ -11,9 +11,10 @@ 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.modules.utils.config import Config
|
||||
from scepter.studio.utils.env import get_available_memory
|
||||
|
||||
from .control_inference import ControlInference
|
||||
@@ -49,7 +50,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',
|
||||
@@ -313,7 +317,8 @@ class DiffusionInference():
|
||||
module_paras = {}
|
||||
if cfg is not None:
|
||||
self.paras = cfg.PARAS
|
||||
self.input = {k.lower(): dict(v).get('DEFAULT', None) if isinstance(v, (dict, OrderedDict)) else v for k, v in cfg.INPUT.items()}
|
||||
self.input_cfg = {k.lower(): v for k, v in cfg.INPUT.items()}
|
||||
self.input = {k.lower(): dict(v).get('DEFAULT', None) if isinstance(v, (dict, OrderedDict, Config)) else v for k, v in cfg.INPUT.items()}
|
||||
self.output = {k.lower(): v for k, v in cfg.OUTPUT.items()}
|
||||
module_paras = cfg.MODULES_PARAS
|
||||
return module_paras
|
||||
|
||||
@@ -151,7 +151,7 @@ class FluxInference(DiffusionInference):
|
||||
with torch.autocast('cuda',
|
||||
enabled= dtype in ('float16', 'bfloat16'),
|
||||
dtype=getattr(torch, dtype)):
|
||||
solver_sample = value_input.get('sample', 'flow_eluer')
|
||||
solver_sample = value_input.get('sample', 'flow_euler')
|
||||
sample_steps = value_input.get('sample_steps', 20)
|
||||
guide_scale = value_input.get('guide_scale', 3.5)
|
||||
if guide_scale is not None:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -29,11 +29,11 @@ class TunerInference():
|
||||
warnings.warn(f'Import swift error, please deal with this problem: {e}')
|
||||
|
||||
self.logger.info('Unloading tuner model')
|
||||
if isinstance(diffusion_model['model'], SwiftModel):
|
||||
if diffusion_model is not None and isinstance(diffusion_model['model'], SwiftModel):
|
||||
for adapter_name in diffusion_model['model'].adapters:
|
||||
diffusion_model['model'].deactivate_adapter(adapter_name,
|
||||
offload='cpu')
|
||||
if isinstance(cond_stage_model['model'], SwiftModel):
|
||||
if cond_stage_model is not None and isinstance(cond_stage_model['model'], SwiftModel):
|
||||
for adapter_name in cond_stage_model['model'].adapters:
|
||||
cond_stage_model['model'].deactivate_adapter(adapter_name,
|
||||
offload='cpu')
|
||||
|
||||
@@ -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, flux)
|
||||
from scepter.modules.model.backbone import (ace, autoencoder, flux, image, cogvideox,
|
||||
mmdit, pixart, unet, utils, video)
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from .ace import ACE
|
||||
@@ -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 scepter.modules.model.backbone.cogvideox.cogvideox import CogVideoXTransformer3DModel
|
||||
@@ -0,0 +1,319 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
# Copyright 2024 The CogVideoX team, Tsinghua University & ZhipuAI and The HuggingFace Team.
|
||||
# All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from collections import OrderedDict
|
||||
from typing import Any, Dict, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from scepter.modules.model.base_model import BaseModel
|
||||
from scepter.modules.model.registry import BACKBONES
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
from .layers import CogVideoXBlock, CogVideoXPatchEmbed, TimestepEmbedding, Timesteps, AdaLayerNorm
|
||||
|
||||
|
||||
@BACKBONES.register_class()
|
||||
class CogVideoXTransformer3DModel(BaseModel):
|
||||
"""
|
||||
A Transformer model for video-like data in [CogVideoX](https://github.com/THUDM/CogVideo).
|
||||
|
||||
Parameters:
|
||||
num_attention_heads (`int`, defaults to `30`):
|
||||
The number of heads to use for multi-head attention.
|
||||
attention_head_dim (`int`, defaults to `64`):
|
||||
The number of channels in each head.
|
||||
in_channels (`int`, defaults to `16`):
|
||||
The number of channels in the input.
|
||||
out_channels (`int`, *optional*, defaults to `16`):
|
||||
The number of channels in the output.
|
||||
flip_sin_to_cos (`bool`, defaults to `True`):
|
||||
Whether to flip the sin to cos in the time embedding.
|
||||
time_embed_dim (`int`, defaults to `512`):
|
||||
Output dimension of timestep embeddings.
|
||||
text_embed_dim (`int`, defaults to `4096`):
|
||||
Input dimension of text embeddings from the text encoder.
|
||||
num_layers (`int`, defaults to `30`):
|
||||
The number of layers of Transformer blocks to use.
|
||||
dropout (`float`, defaults to `0.0`):
|
||||
The dropout probability to use.
|
||||
attention_bias (`bool`, defaults to `True`):
|
||||
Whether or not to use bias in the attention projection layers.
|
||||
sample_width (`int`, defaults to `90`):
|
||||
The width of the input latents.
|
||||
sample_height (`int`, defaults to `60`):
|
||||
The height of the input latents.
|
||||
sample_frames (`int`, defaults to `49`):
|
||||
The number of frames in the input latents. Note that this parameter was incorrectly initialized to 49
|
||||
instead of 13 because CogVideoX processed 13 latent frames at once in its default and recommended settings,
|
||||
but cannot be changed to the correct value to ensure backwards compatibility. To create a transformer with
|
||||
K latent frames, the correct value to pass here would be: ((K - 1) * temporal_compression_ratio + 1).
|
||||
patch_size (`int`, defaults to `2`):
|
||||
The size of the patches to use in the patch embedding layer.
|
||||
temporal_compression_ratio (`int`, defaults to `4`):
|
||||
The compression ratio across the temporal dimension. See documentation for `sample_frames`.
|
||||
max_text_seq_length (`int`, defaults to `226`):
|
||||
The maximum sequence length of the input text embeddings.
|
||||
activation_fn (`str`, defaults to `"gelu-approximate"`):
|
||||
Activation function to use in feed-forward.
|
||||
timestep_activation_fn (`str`, defaults to `"silu"`):
|
||||
Activation function to use when generating the timestep embeddings.
|
||||
norm_elementwise_affine (`bool`, defaults to `True`):
|
||||
Whether or not to use elementwise affine in normalization layers.
|
||||
norm_eps (`float`, defaults to `1e-5`):
|
||||
The epsilon value to use in normalization layers.
|
||||
spatial_interpolation_scale (`float`, defaults to `1.875`):
|
||||
Scaling factor to apply in 3D positional embeddings across spatial dimensions.
|
||||
temporal_interpolation_scale (`float`, defaults to `1.0`):
|
||||
Scaling factor to apply in 3D positional embeddings across temporal dimensions.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
cfg,
|
||||
logger=None
|
||||
):
|
||||
super().__init__(cfg, logger=logger)
|
||||
num_attention_heads = cfg.get("NUM_ATTENTION_HEADS", 30)
|
||||
attention_head_dim = cfg.get("ATTENTION_HEAD_DIM", 64)
|
||||
in_channels = cfg.get("IN_CHANNELS", 16)
|
||||
out_channels = cfg.get("OUT_CHANNELS", 16)
|
||||
flip_sin_to_cos = cfg.get("FLIP_SIN_TO_COS", True)
|
||||
freq_shift = cfg.get("FREQ_SHIFT", 0)
|
||||
time_embed_dim = cfg.get("TIME_EMBED_DIM", 512)
|
||||
text_embed_dim = cfg.get("TEXT_EMBED_DIM", 4096)
|
||||
num_layers = cfg.get("NUM_LAYERS", 30)
|
||||
dropout = cfg.get("DROPOUT", 0.0)
|
||||
attention_bias = cfg.get("ATTENTION_BIAS", True)
|
||||
sample_width = cfg.get("SAMPLE_WIDTH", 90)
|
||||
sample_height = cfg.get("SAMPLE_HEIGHT", 60)
|
||||
sample_frames = cfg.get("SAMPLE_FRAMES", 49)
|
||||
patch_size = cfg.get("PATCH_SIZE", 2)
|
||||
temporal_compression_ratio = cfg.get("TEMPORAL_COMPRESSION_RATIO", 4)
|
||||
max_text_seq_length = cfg.get("MAX_TEXT_SEQ_LENGTH", 226)
|
||||
activation_fn = cfg.get("ACTIVATION_FN", "gelu-approximate")
|
||||
timestep_activation_fn = cfg.get("TIMESTEP_ACTIVATION_FN", "silu")
|
||||
norm_elementwise_affine = cfg.get("NORM_ELEMENTWISE_AFFINE", True)
|
||||
norm_eps = cfg.get("NORM_EPS", 1e-5)
|
||||
spatial_interpolation_scale = cfg.get("SPATIAL_INTERPOLATION_SCALE", 1.875)
|
||||
temporal_interpolation_scale = cfg.get("TEMPORAL_INTERPOLATION_SCALE", 1.0)
|
||||
use_rotary_positional_embeddings = cfg.get("USE_ROTARY_POSITIONAL_EMBEDDINGS", False)
|
||||
use_learned_positional_embeddings = cfg.get("USE_LEARNED_POSITIONAL_EMBEDDINGS", False)
|
||||
self.gradient_checkpointing = cfg.get("GRADIENT_CHECKPOINTING", False)
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
self.patch_size = patch_size
|
||||
self.use_rotary_positional_embeddings = use_rotary_positional_embeddings
|
||||
|
||||
if not use_rotary_positional_embeddings and use_learned_positional_embeddings:
|
||||
raise ValueError(
|
||||
"There are no CogVideoX checkpoints available with disable rotary embeddings and learned positional "
|
||||
"embeddings. If you're using a custom model and/or believe this should be supported, please open an "
|
||||
"issue at https://github.com/huggingface/diffusers/issues."
|
||||
)
|
||||
|
||||
# 1. Patch embedding
|
||||
self.patch_embed = CogVideoXPatchEmbed(
|
||||
patch_size=patch_size,
|
||||
in_channels=in_channels,
|
||||
embed_dim=inner_dim,
|
||||
text_embed_dim=text_embed_dim,
|
||||
bias=True,
|
||||
sample_width=sample_width,
|
||||
sample_height=sample_height,
|
||||
sample_frames=sample_frames,
|
||||
temporal_compression_ratio=temporal_compression_ratio,
|
||||
max_text_seq_length=max_text_seq_length,
|
||||
spatial_interpolation_scale=spatial_interpolation_scale,
|
||||
temporal_interpolation_scale=temporal_interpolation_scale,
|
||||
use_positional_embeddings=not use_rotary_positional_embeddings,
|
||||
use_learned_positional_embeddings=use_learned_positional_embeddings,
|
||||
)
|
||||
self.embedding_dropout = nn.Dropout(dropout)
|
||||
|
||||
# 2. Time embeddings
|
||||
self.time_proj = Timesteps(inner_dim, flip_sin_to_cos, freq_shift)
|
||||
self.time_embedding = TimestepEmbedding(inner_dim, time_embed_dim, timestep_activation_fn)
|
||||
|
||||
# 3. Define spatio-temporal transformers blocks
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
CogVideoXBlock(
|
||||
dim=inner_dim,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
time_embed_dim=time_embed_dim,
|
||||
dropout=dropout,
|
||||
activation_fn=activation_fn,
|
||||
attention_bias=attention_bias,
|
||||
norm_elementwise_affine=norm_elementwise_affine,
|
||||
norm_eps=norm_eps,
|
||||
)
|
||||
for _ in range(num_layers)
|
||||
]
|
||||
)
|
||||
self.norm_final = nn.LayerNorm(inner_dim, norm_eps, norm_elementwise_affine)
|
||||
|
||||
# 4. Output blocks
|
||||
self.norm_out = AdaLayerNorm(
|
||||
embedding_dim=time_embed_dim,
|
||||
output_dim=2 * inner_dim,
|
||||
norm_elementwise_affine=norm_elementwise_affine,
|
||||
norm_eps=norm_eps,
|
||||
chunk_dim=1,
|
||||
)
|
||||
self.proj_out = nn.Linear(inner_dim, patch_size * patch_size * out_channels)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor = None,
|
||||
t: Union[int, float, torch.LongTensor] = None,
|
||||
cond: torch.Tensor = None,
|
||||
timestep_cond: Optional[torch.Tensor] = None,
|
||||
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
**kwargs
|
||||
):
|
||||
if 'image_latent' in kwargs and kwargs['image_latent'] is not None:
|
||||
hidden_states = torch.cat([x, kwargs['image_latent']], dim=2)
|
||||
else:
|
||||
hidden_states = x
|
||||
timestep = t
|
||||
encoder_hidden_states = cond
|
||||
|
||||
batch_size, num_frames, channels, height, width = hidden_states.shape
|
||||
|
||||
# 1. Time embedding
|
||||
timesteps = timestep
|
||||
t_emb = self.time_proj(timesteps)
|
||||
|
||||
# timesteps does not contain any weights and will always return f32 tensors
|
||||
# but time_embedding might actually be running in fp16. so we need to cast here.
|
||||
# there might be better ways to encapsulate this.
|
||||
t_emb = t_emb.to(dtype=encoder_hidden_states.dtype)
|
||||
emb = self.time_embedding(t_emb, timestep_cond)
|
||||
|
||||
# 2. Patch embedding
|
||||
hidden_states = self.patch_embed(encoder_hidden_states, hidden_states)
|
||||
hidden_states = self.embedding_dropout(hidden_states)
|
||||
|
||||
text_seq_length = encoder_hidden_states.shape[1]
|
||||
encoder_hidden_states = hidden_states[:, :text_seq_length]
|
||||
hidden_states = hidden_states[:, text_seq_length:]
|
||||
|
||||
# 3. Transformer blocks
|
||||
for i, block in enumerate(self.transformer_blocks):
|
||||
if self.training and self.gradient_checkpointing:
|
||||
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs):
|
||||
return module(*inputs)
|
||||
|
||||
return custom_forward
|
||||
|
||||
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False}
|
||||
hidden_states, encoder_hidden_states = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(block),
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
emb,
|
||||
image_rotary_emb,
|
||||
**ckpt_kwargs,
|
||||
)
|
||||
else:
|
||||
hidden_states, encoder_hidden_states = block(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
temb=emb,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
)
|
||||
|
||||
if not self.use_rotary_positional_embeddings:
|
||||
# CogVideoX-2B
|
||||
hidden_states = self.norm_final(hidden_states)
|
||||
else:
|
||||
# CogVideoX-5B
|
||||
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
|
||||
hidden_states = self.norm_final(hidden_states)
|
||||
hidden_states = hidden_states[:, text_seq_length:]
|
||||
|
||||
# 4. Final block
|
||||
hidden_states = self.norm_out(hidden_states, temb=emb)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
# 5. Unpatchify
|
||||
# Note: we use `-1` instead of `channels`:
|
||||
# - It is okay to `channels` use for CogVideoX-2b and CogVideoX-5b (number of input channels is equal to output channels)
|
||||
# - However, for CogVideoX-5b-I2V also takes concatenated input image latents (number of input channels is twice the output channels)
|
||||
p = self.patch_size
|
||||
output = hidden_states.reshape(batch_size, num_frames, height // p, width // p, -1, p, p)
|
||||
output = output.permute(0, 1, 4, 2, 5, 3, 6).flatten(5, 6).flatten(3, 4)
|
||||
|
||||
return output
|
||||
|
||||
def load_pretrained_model(self, pretrained_model):
|
||||
if pretrained_model is not None:
|
||||
pretrained_model_list = [pretrained_model] if isinstance(pretrained_model, str) else pretrained_model
|
||||
ckpt_all = OrderedDict()
|
||||
for pretrained_model in pretrained_model_list:
|
||||
with FS.get_from(pretrained_model,
|
||||
wait_finish=True) as local_model:
|
||||
if local_model.endswith('safetensors'):
|
||||
from safetensors.torch import load_file as load_safetensors
|
||||
ckpt = load_safetensors(local_model)
|
||||
else:
|
||||
ckpt = torch.load(local_model, map_location='cpu')
|
||||
ckpt_all.update(ckpt)
|
||||
missing, unexpected = self.load_state_dict(ckpt_all, strict=False)
|
||||
if we.rank == 0:
|
||||
self.logger.info(
|
||||
f'Restored from {pretrained_model_list} with {len(missing)} missing and {len(unexpected)} unexpected keys'
|
||||
)
|
||||
if len(missing) > 0:
|
||||
self.logger.info(f'Missing Keys:\n {missing}')
|
||||
if len(unexpected) > 0:
|
||||
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODEL',
|
||||
__class__.__name__,
|
||||
CogVideoXTransformer3DModel.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.logger import get_logger
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
cfg = Config(parser_ins=parser)
|
||||
for file_sys in cfg.FILE_SYSTEM:
|
||||
FS.init_fs_client(file_sys)
|
||||
model = BACKBONES.build(cfg.DIFFUSION_MODEL, logger=get_logger()).eval().requires_grad_(False).to('cuda').to(torch.bfloat16)
|
||||
|
||||
hidden_states = torch.load(FS.get_from(cfg.HIDDEN_STATES))
|
||||
encoder_hidden_states = torch.load(FS.get_from(cfg.ENCODER_HIDDEN_STATES))
|
||||
timestep = torch.load(FS.get_from(cfg.TIMESTEP))
|
||||
timestep_cond = None
|
||||
image_rotary_emb = None
|
||||
attention_kwargs = None
|
||||
output = model(hidden_states, encoder_hidden_states, timestep, timestep_cond, image_rotary_emb, attention_kwargs)
|
||||
print(output, torch.sum(output))
|
||||
@@ -0,0 +1,554 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
# Copyright 2024 The CogVideoX team, Tsinghua University & ZhipuAI and The HuggingFace Team.
|
||||
# All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from .utils import get_activation, get_timestep_embedding, get_3d_sincos_pos_embed, apply_rotary_emb
|
||||
from .utils import GELU, GEGLU, ApproximateGELU, SwiGLU
|
||||
|
||||
|
||||
class TimestepEmbedding(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
time_embed_dim: int,
|
||||
act_fn: str = "silu",
|
||||
out_dim: int = None,
|
||||
post_act_fn: Optional[str] = None,
|
||||
cond_proj_dim=None,
|
||||
sample_proj_bias=True,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.linear_1 = nn.Linear(in_channels, time_embed_dim, sample_proj_bias)
|
||||
|
||||
if cond_proj_dim is not None:
|
||||
self.cond_proj = nn.Linear(cond_proj_dim, in_channels, bias=False)
|
||||
else:
|
||||
self.cond_proj = None
|
||||
|
||||
self.act = get_activation(act_fn)
|
||||
|
||||
if out_dim is not None:
|
||||
time_embed_dim_out = out_dim
|
||||
else:
|
||||
time_embed_dim_out = time_embed_dim
|
||||
self.linear_2 = nn.Linear(time_embed_dim, time_embed_dim_out, sample_proj_bias)
|
||||
|
||||
if post_act_fn is None:
|
||||
self.post_act = None
|
||||
else:
|
||||
self.post_act = get_activation(post_act_fn)
|
||||
|
||||
def forward(self, sample, condition=None):
|
||||
if condition is not None:
|
||||
sample = sample + self.cond_proj(condition)
|
||||
sample = self.linear_1(sample)
|
||||
|
||||
if self.act is not None:
|
||||
sample = self.act(sample)
|
||||
|
||||
sample = self.linear_2(sample)
|
||||
|
||||
if self.post_act is not None:
|
||||
sample = self.post_act(sample)
|
||||
return sample
|
||||
|
||||
|
||||
|
||||
class Timesteps(nn.Module):
|
||||
def __init__(self, num_channels: int, flip_sin_to_cos: bool, downscale_freq_shift: float, scale: int = 1):
|
||||
super().__init__()
|
||||
self.num_channels = num_channels
|
||||
self.flip_sin_to_cos = flip_sin_to_cos
|
||||
self.downscale_freq_shift = downscale_freq_shift
|
||||
self.scale = scale
|
||||
|
||||
def forward(self, timesteps):
|
||||
t_emb = get_timestep_embedding(
|
||||
timesteps,
|
||||
self.num_channels,
|
||||
flip_sin_to_cos=self.flip_sin_to_cos,
|
||||
downscale_freq_shift=self.downscale_freq_shift,
|
||||
scale=self.scale,
|
||||
)
|
||||
return t_emb
|
||||
|
||||
|
||||
class CogVideoXLayerNormZero(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
conditioning_dim: int,
|
||||
embedding_dim: int,
|
||||
elementwise_affine: bool = True,
|
||||
eps: float = 1e-5,
|
||||
bias: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.silu = nn.SiLU()
|
||||
self.linear = nn.Linear(conditioning_dim, 6 * embedding_dim, bias=bias)
|
||||
self.norm = nn.LayerNorm(embedding_dim, eps=eps, elementwise_affine=elementwise_affine)
|
||||
|
||||
def forward(
|
||||
self, hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor, temb: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
shift, scale, gate, enc_shift, enc_scale, enc_gate = self.linear(self.silu(temb)).chunk(6, dim=1)
|
||||
hidden_states = self.norm(hidden_states) * (1 + scale)[:, None, :] + shift[:, None, :]
|
||||
encoder_hidden_states = self.norm(encoder_hidden_states) * (1 + enc_scale)[:, None, :] + enc_shift[:, None, :]
|
||||
return hidden_states, encoder_hidden_states, gate[:, None, :], enc_gate[:, None, :]
|
||||
|
||||
|
||||
class AdaLayerNorm(nn.Module):
|
||||
r"""
|
||||
Norm layer modified to incorporate timestep embeddings.
|
||||
|
||||
Parameters:
|
||||
embedding_dim (`int`): The size of each embedding vector.
|
||||
num_embeddings (`int`, *optional*): The size of the embeddings dictionary.
|
||||
output_dim (`int`, *optional*):
|
||||
norm_elementwise_affine (`bool`, defaults to `False):
|
||||
norm_eps (`bool`, defaults to `False`):
|
||||
chunk_dim (`int`, defaults to `0`):
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
num_embeddings: Optional[int] = None,
|
||||
output_dim: Optional[int] = None,
|
||||
norm_elementwise_affine: bool = False,
|
||||
norm_eps: float = 1e-5,
|
||||
chunk_dim: int = 0,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.chunk_dim = chunk_dim
|
||||
output_dim = output_dim or embedding_dim * 2
|
||||
|
||||
if num_embeddings is not None:
|
||||
self.emb = nn.Embedding(num_embeddings, embedding_dim)
|
||||
else:
|
||||
self.emb = None
|
||||
|
||||
self.silu = nn.SiLU()
|
||||
self.linear = nn.Linear(embedding_dim, output_dim)
|
||||
self.norm = nn.LayerNorm(output_dim // 2, norm_eps, norm_elementwise_affine)
|
||||
|
||||
def forward(
|
||||
self, x: torch.Tensor, timestep: Optional[torch.Tensor] = None, temb: Optional[torch.Tensor] = None
|
||||
) -> torch.Tensor:
|
||||
if self.emb is not None:
|
||||
temb = self.emb(timestep)
|
||||
|
||||
temb = self.linear(self.silu(temb))
|
||||
|
||||
if self.chunk_dim == 1:
|
||||
# This is a bit weird why we have the order of "shift, scale" here and "scale, shift" in the
|
||||
# other if-branch. This branch is specific to CogVideoX for now.
|
||||
shift, scale = temb.chunk(2, dim=1)
|
||||
shift = shift[:, None, :]
|
||||
scale = scale[:, None, :]
|
||||
else:
|
||||
scale, shift = temb.chunk(2, dim=0)
|
||||
|
||||
x = self.norm(x) * (1 + scale) + shift
|
||||
return x
|
||||
|
||||
|
||||
class CogVideoXPatchEmbed(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: int = 2,
|
||||
in_channels: int = 16,
|
||||
embed_dim: int = 1920,
|
||||
text_embed_dim: int = 4096,
|
||||
bias: bool = True,
|
||||
sample_width: int = 90,
|
||||
sample_height: int = 60,
|
||||
sample_frames: int = 49,
|
||||
temporal_compression_ratio: int = 4,
|
||||
max_text_seq_length: int = 226,
|
||||
spatial_interpolation_scale: float = 1.875,
|
||||
temporal_interpolation_scale: float = 1.0,
|
||||
use_positional_embeddings: bool = True,
|
||||
use_learned_positional_embeddings: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.patch_size = patch_size
|
||||
self.embed_dim = embed_dim
|
||||
self.sample_height = sample_height
|
||||
self.sample_width = sample_width
|
||||
self.sample_frames = sample_frames
|
||||
self.temporal_compression_ratio = temporal_compression_ratio
|
||||
self.max_text_seq_length = max_text_seq_length
|
||||
self.spatial_interpolation_scale = spatial_interpolation_scale
|
||||
self.temporal_interpolation_scale = temporal_interpolation_scale
|
||||
self.use_positional_embeddings = use_positional_embeddings
|
||||
self.use_learned_positional_embeddings = use_learned_positional_embeddings
|
||||
|
||||
self.proj = nn.Conv2d(
|
||||
in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size, bias=bias
|
||||
)
|
||||
self.text_proj = nn.Linear(text_embed_dim, embed_dim)
|
||||
|
||||
if use_positional_embeddings or use_learned_positional_embeddings:
|
||||
persistent = use_learned_positional_embeddings
|
||||
pos_embedding = self._get_positional_embeddings(sample_height, sample_width, sample_frames)
|
||||
self.register_buffer("pos_embedding", pos_embedding, persistent=persistent)
|
||||
|
||||
def _get_positional_embeddings(self, sample_height: int, sample_width: int, sample_frames: int) -> torch.Tensor:
|
||||
post_patch_height = sample_height // self.patch_size
|
||||
post_patch_width = sample_width // self.patch_size
|
||||
post_time_compression_frames = (sample_frames - 1) // self.temporal_compression_ratio + 1
|
||||
num_patches = post_patch_height * post_patch_width * post_time_compression_frames
|
||||
|
||||
pos_embedding = get_3d_sincos_pos_embed(
|
||||
self.embed_dim,
|
||||
(post_patch_width, post_patch_height),
|
||||
post_time_compression_frames,
|
||||
self.spatial_interpolation_scale,
|
||||
self.temporal_interpolation_scale,
|
||||
)
|
||||
pos_embedding = torch.from_numpy(pos_embedding).flatten(0, 1)
|
||||
joint_pos_embedding = torch.zeros(
|
||||
1, self.max_text_seq_length + num_patches, self.embed_dim, requires_grad=False
|
||||
)
|
||||
joint_pos_embedding.data[:, self.max_text_seq_length :].copy_(pos_embedding)
|
||||
|
||||
return joint_pos_embedding
|
||||
|
||||
def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor):
|
||||
r"""
|
||||
Args:
|
||||
text_embeds (`torch.Tensor`):
|
||||
Input text embeddings. Expected shape: (batch_size, seq_length, embedding_dim).
|
||||
image_embeds (`torch.Tensor`):
|
||||
Input image embeddings. Expected shape: (batch_size, num_frames, channels, height, width).
|
||||
"""
|
||||
text_embeds = self.text_proj(text_embeds)
|
||||
|
||||
batch, num_frames, channels, height, width = image_embeds.shape
|
||||
image_embeds = image_embeds.reshape(-1, channels, height, width)
|
||||
image_embeds = self.proj(image_embeds)
|
||||
image_embeds = image_embeds.view(batch, num_frames, *image_embeds.shape[1:])
|
||||
image_embeds = image_embeds.flatten(3).transpose(2, 3) # [batch, num_frames, height x width, channels]
|
||||
image_embeds = image_embeds.flatten(1, 2) # [batch, num_frames x height x width, channels]
|
||||
|
||||
embeds = torch.cat(
|
||||
[text_embeds, image_embeds], dim=1
|
||||
).contiguous() # [batch, seq_length + num_frames x height x width, channels]
|
||||
|
||||
if self.use_positional_embeddings or self.use_learned_positional_embeddings:
|
||||
if self.use_learned_positional_embeddings and (self.sample_width != width or self.sample_height != height):
|
||||
raise ValueError(
|
||||
"It is currently not possible to generate videos at a different resolution that the defaults. This should only be the case with 'THUDM/CogVideoX-5b-I2V'."
|
||||
"If you think this is incorrect, please open an issue at https://github.com/huggingface/diffusers/issues."
|
||||
)
|
||||
|
||||
pre_time_compression_frames = (num_frames - 1) * self.temporal_compression_ratio + 1
|
||||
|
||||
if (
|
||||
self.sample_height != height
|
||||
or self.sample_width != width
|
||||
or self.sample_frames != pre_time_compression_frames
|
||||
):
|
||||
pos_embedding = self._get_positional_embeddings(height, width, pre_time_compression_frames)
|
||||
pos_embedding = pos_embedding.to(embeds.device, dtype=embeds.dtype)
|
||||
else:
|
||||
pos_embedding = self.pos_embedding
|
||||
|
||||
embeds = embeds + pos_embedding
|
||||
|
||||
return embeds
|
||||
|
||||
|
||||
class FeedForward(nn.Module):
|
||||
r"""
|
||||
A feed-forward layer.
|
||||
|
||||
Parameters:
|
||||
dim (`int`): The number of channels in the input.
|
||||
dim_out (`int`, *optional*): The number of channels in the output. If not given, defaults to `dim`.
|
||||
mult (`int`, *optional*, defaults to 4): The multiplier to use for the hidden dimension.
|
||||
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
|
||||
activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
|
||||
final_dropout (`bool` *optional*, defaults to False): Apply a final dropout.
|
||||
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
dim_out: Optional[int] = None,
|
||||
mult: int = 4,
|
||||
dropout: float = 0.0,
|
||||
activation_fn: str = "geglu",
|
||||
final_dropout: bool = False,
|
||||
inner_dim=None,
|
||||
bias: bool = True,
|
||||
):
|
||||
super().__init__()
|
||||
if inner_dim is None:
|
||||
inner_dim = int(dim * mult)
|
||||
dim_out = dim_out if dim_out is not None else dim
|
||||
|
||||
if activation_fn == "gelu":
|
||||
act_fn = GELU(dim, inner_dim, bias=bias)
|
||||
if activation_fn == "gelu-approximate":
|
||||
act_fn = GELU(dim, inner_dim, approximate="tanh", bias=bias)
|
||||
elif activation_fn == "geglu":
|
||||
act_fn = GEGLU(dim, inner_dim, bias=bias)
|
||||
elif activation_fn == "geglu-approximate":
|
||||
act_fn = ApproximateGELU(dim, inner_dim, bias=bias)
|
||||
elif activation_fn == "swiglu":
|
||||
act_fn = SwiGLU(dim, inner_dim, bias=bias)
|
||||
|
||||
self.net = nn.ModuleList([])
|
||||
# project in
|
||||
self.net.append(act_fn)
|
||||
# project dropout
|
||||
self.net.append(nn.Dropout(dropout))
|
||||
# project out
|
||||
self.net.append(nn.Linear(inner_dim, dim_out, bias=bias))
|
||||
# FF as used in Vision Transformer, MLP-Mixer, etc. have a final dropout
|
||||
if final_dropout:
|
||||
self.net.append(nn.Dropout(dropout))
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor, *args, **kwargs) -> torch.Tensor:
|
||||
if len(args) > 0 or kwargs.get("scale", None) is not None:
|
||||
deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`."
|
||||
print(deprecation_message)
|
||||
for module in self.net:
|
||||
hidden_states = module(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
query_dim: int,
|
||||
dim_head: int = 64,
|
||||
heads: int = 8,
|
||||
kv_heads: Optional[int] = None,
|
||||
qk_norm: Optional[str] = None,
|
||||
eps: float = 1e-5,
|
||||
bias: bool = False,
|
||||
out_bias: bool = True,
|
||||
dropout: float = 0.0,
|
||||
out_dim: int = None,
|
||||
cross_attention_dim: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.inner_dim = out_dim if out_dim is not None else dim_head * heads
|
||||
self.inner_kv_dim = self.inner_dim if kv_heads is None else dim_head * kv_heads
|
||||
self.query_dim = query_dim
|
||||
self.cross_attention_dim = cross_attention_dim if cross_attention_dim is not None else query_dim
|
||||
self.is_cross_attention = cross_attention_dim is not None
|
||||
self.out_dim = out_dim if out_dim is not None else query_dim
|
||||
self.heads = out_dim // dim_head if out_dim is not None else heads
|
||||
|
||||
if qk_norm is None:
|
||||
self.norm_q = None
|
||||
self.norm_k = None
|
||||
elif qk_norm == "layer_norm":
|
||||
self.norm_q = nn.LayerNorm(dim_head, eps=eps)
|
||||
self.norm_k = nn.LayerNorm(dim_head, eps=eps)
|
||||
|
||||
self.to_q = nn.Linear(query_dim, self.inner_dim, bias=bias)
|
||||
self.to_k = nn.Linear(self.cross_attention_dim, self.inner_kv_dim, bias=bias)
|
||||
self.to_v = nn.Linear(self.cross_attention_dim, self.inner_kv_dim, bias=bias)
|
||||
self.to_out = nn.ModuleList([])
|
||||
self.to_out.append(nn.Linear(self.inner_dim, self.out_dim, bias=out_bias))
|
||||
self.to_out.append(nn.Dropout(dropout))
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
image_rotary_emb: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
|
||||
text_seq_length = encoder_hidden_states.size(1)
|
||||
|
||||
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
|
||||
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
)
|
||||
|
||||
query = self.to_q(hidden_states)
|
||||
key = self.to_k(hidden_states)
|
||||
value = self.to_v(hidden_states)
|
||||
|
||||
inner_dim = key.shape[-1]
|
||||
head_dim = inner_dim // self.heads
|
||||
|
||||
query = query.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
|
||||
key = key.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
|
||||
value = value.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k(key)
|
||||
|
||||
# Apply RoPE if needed
|
||||
if image_rotary_emb is not None:
|
||||
query[:, :, text_seq_length:] = apply_rotary_emb(query[:, :, text_seq_length:], image_rotary_emb)
|
||||
if not self.is_cross_attention:
|
||||
key[:, :, text_seq_length:] = apply_rotary_emb(key[:, :, text_seq_length:], image_rotary_emb)
|
||||
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.heads * head_dim)
|
||||
|
||||
# linear proj
|
||||
hidden_states = self.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = self.to_out[1](hidden_states)
|
||||
|
||||
encoder_hidden_states, hidden_states = hidden_states.split(
|
||||
[text_seq_length, hidden_states.size(1) - text_seq_length], dim=1
|
||||
)
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
class CogVideoXBlock(nn.Module):
|
||||
r"""
|
||||
Transformer block used in [CogVideoX](https://github.com/THUDM/CogVideo) model.
|
||||
|
||||
Parameters:
|
||||
dim (`int`):
|
||||
The number of channels in the input and output.
|
||||
num_attention_heads (`int`):
|
||||
The number of heads to use for multi-head attention.
|
||||
attention_head_dim (`int`):
|
||||
The number of channels in each head.
|
||||
time_embed_dim (`int`):
|
||||
The number of channels in timestep embedding.
|
||||
dropout (`float`, defaults to `0.0`):
|
||||
The dropout probability to use.
|
||||
activation_fn (`str`, defaults to `"gelu-approximate"`):
|
||||
Activation function to be used in feed-forward.
|
||||
attention_bias (`bool`, defaults to `False`):
|
||||
Whether or not to use bias in attention projection layers.
|
||||
qk_norm (`bool`, defaults to `True`):
|
||||
Whether or not to use normalization after query and key projections in Attention.
|
||||
norm_elementwise_affine (`bool`, defaults to `True`):
|
||||
Whether to use learnable elementwise affine parameters for normalization.
|
||||
norm_eps (`float`, defaults to `1e-5`):
|
||||
Epsilon value for normalization layers.
|
||||
final_dropout (`bool` defaults to `False`):
|
||||
Whether to apply a final dropout after the last feed-forward layer.
|
||||
ff_inner_dim (`int`, *optional*, defaults to `None`):
|
||||
Custom hidden dimension of Feed-forward layer. If not provided, `4 * dim` is used.
|
||||
ff_bias (`bool`, defaults to `True`):
|
||||
Whether or not to use bias in Feed-forward layer.
|
||||
attention_out_bias (`bool`, defaults to `True`):
|
||||
Whether or not to use bias in Attention output projection layer.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
time_embed_dim: int,
|
||||
dropout: float = 0.0,
|
||||
activation_fn: str = "gelu-approximate",
|
||||
attention_bias: bool = False,
|
||||
qk_norm: bool = True,
|
||||
norm_elementwise_affine: bool = True,
|
||||
norm_eps: float = 1e-5,
|
||||
final_dropout: bool = True,
|
||||
ff_inner_dim: Optional[int] = None,
|
||||
ff_bias: bool = True,
|
||||
attention_out_bias: bool = True,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self Attention
|
||||
self.norm1 = CogVideoXLayerNormZero(time_embed_dim, dim, norm_elementwise_affine, norm_eps, bias=True)
|
||||
|
||||
self.attn1 = Attention(
|
||||
query_dim=dim,
|
||||
dim_head=attention_head_dim,
|
||||
heads=num_attention_heads,
|
||||
qk_norm="layer_norm" if qk_norm else None,
|
||||
eps=1e-6,
|
||||
bias=attention_bias,
|
||||
out_bias=attention_out_bias
|
||||
)
|
||||
|
||||
# 2. Feed Forward
|
||||
self.norm2 = CogVideoXLayerNormZero(time_embed_dim, dim, norm_elementwise_affine, norm_eps, bias=True)
|
||||
|
||||
self.ff = FeedForward(
|
||||
dim,
|
||||
dropout=dropout,
|
||||
activation_fn=activation_fn,
|
||||
final_dropout=final_dropout,
|
||||
inner_dim=ff_inner_dim,
|
||||
bias=ff_bias,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
) -> torch.Tensor:
|
||||
text_seq_length = encoder_hidden_states.size(1)
|
||||
|
||||
# norm & modulate
|
||||
norm_hidden_states, norm_encoder_hidden_states, gate_msa, enc_gate_msa = self.norm1(
|
||||
hidden_states, encoder_hidden_states, temb
|
||||
)
|
||||
|
||||
# attention
|
||||
attn_hidden_states, attn_encoder_hidden_states = self.attn1(
|
||||
hidden_states=norm_hidden_states,
|
||||
encoder_hidden_states=norm_encoder_hidden_states,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
)
|
||||
|
||||
hidden_states = hidden_states + gate_msa * attn_hidden_states
|
||||
encoder_hidden_states = encoder_hidden_states + enc_gate_msa * attn_encoder_hidden_states
|
||||
|
||||
# norm & modulate
|
||||
norm_hidden_states, norm_encoder_hidden_states, gate_ff, enc_gate_ff = self.norm2(
|
||||
hidden_states, encoder_hidden_states, temb
|
||||
)
|
||||
|
||||
# feed-forward
|
||||
norm_hidden_states = torch.cat([norm_encoder_hidden_states, norm_hidden_states], dim=1)
|
||||
ff_output = self.ff(norm_hidden_states)
|
||||
|
||||
hidden_states = hidden_states + gate_ff * ff_output[:, text_seq_length:]
|
||||
encoder_hidden_states = encoder_hidden_states + enc_gate_ff * ff_output[:, :text_seq_length]
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
@@ -0,0 +1,544 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
# Copyright 2024 The CogVideoX team, Tsinghua University & ZhipuAI and The HuggingFace Team.
|
||||
# All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import math
|
||||
from typing import Optional, Tuple, Union, List
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
ACTIVATION_FUNCTIONS = {
|
||||
"swish": nn.SiLU(),
|
||||
"silu": nn.SiLU(),
|
||||
"mish": nn.Mish(),
|
||||
"gelu": nn.GELU(),
|
||||
"relu": nn.ReLU(),
|
||||
}
|
||||
|
||||
|
||||
def get_activation(act_fn: str) -> nn.Module:
|
||||
"""Helper function to get activation function from string.
|
||||
|
||||
Args:
|
||||
act_fn (str): Name of activation function.
|
||||
|
||||
Returns:
|
||||
nn.Module: Activation function.
|
||||
"""
|
||||
|
||||
act_fn = act_fn.lower()
|
||||
if act_fn in ACTIVATION_FUNCTIONS:
|
||||
return ACTIVATION_FUNCTIONS[act_fn]
|
||||
else:
|
||||
raise ValueError(f"Unsupported activation function: {act_fn}")
|
||||
|
||||
|
||||
class FP32SiLU(nn.Module):
|
||||
r"""
|
||||
SiLU activation function with input upcasted to torch.float32.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
|
||||
return F.silu(inputs.float(), inplace=False).to(inputs.dtype)
|
||||
|
||||
|
||||
class GELU(nn.Module):
|
||||
r"""
|
||||
GELU activation function with tanh approximation support with `approximate="tanh"`.
|
||||
|
||||
Parameters:
|
||||
dim_in (`int`): The number of channels in the input.
|
||||
dim_out (`int`): The number of channels in the output.
|
||||
approximate (`str`, *optional*, defaults to `"none"`): If `"tanh"`, use tanh approximation.
|
||||
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
|
||||
"""
|
||||
|
||||
def __init__(self, dim_in: int, dim_out: int, approximate: str = "none", bias: bool = True):
|
||||
super().__init__()
|
||||
self.proj = nn.Linear(dim_in, dim_out, bias=bias)
|
||||
self.approximate = approximate
|
||||
|
||||
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
|
||||
if gate.device.type != "mps":
|
||||
return F.gelu(gate, approximate=self.approximate)
|
||||
# mps: gelu is not implemented for float16
|
||||
return F.gelu(gate.to(dtype=torch.float32), approximate=self.approximate).to(dtype=gate.dtype)
|
||||
|
||||
def forward(self, hidden_states):
|
||||
hidden_states = self.proj(hidden_states)
|
||||
hidden_states = self.gelu(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class GEGLU(nn.Module):
|
||||
r"""
|
||||
A [variant](https://arxiv.org/abs/2002.05202) of the gated linear unit activation function.
|
||||
|
||||
Parameters:
|
||||
dim_in (`int`): The number of channels in the input.
|
||||
dim_out (`int`): The number of channels in the output.
|
||||
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
|
||||
"""
|
||||
|
||||
def __init__(self, dim_in: int, dim_out: int, bias: bool = True):
|
||||
super().__init__()
|
||||
self.proj = nn.Linear(dim_in, dim_out * 2, bias=bias)
|
||||
|
||||
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
|
||||
if gate.device.type != "mps":
|
||||
return F.gelu(gate)
|
||||
# mps: gelu is not implemented for float16
|
||||
return F.gelu(gate.to(dtype=torch.float32)).to(dtype=gate.dtype)
|
||||
|
||||
def forward(self, hidden_states, *args, **kwargs):
|
||||
if len(args) > 0 or kwargs.get("scale", None) is not None:
|
||||
deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`."
|
||||
print("scale", "1.0.0", deprecation_message)
|
||||
hidden_states = self.proj(hidden_states)
|
||||
hidden_states, gate = hidden_states.chunk(2, dim=-1)
|
||||
return hidden_states * self.gelu(gate)
|
||||
|
||||
|
||||
class SwiGLU(nn.Module):
|
||||
r"""
|
||||
A [variant](https://arxiv.org/abs/2002.05202) of the gated linear unit activation function. It's similar to `GEGLU`
|
||||
but uses SiLU / Swish instead of GeLU.
|
||||
|
||||
Parameters:
|
||||
dim_in (`int`): The number of channels in the input.
|
||||
dim_out (`int`): The number of channels in the output.
|
||||
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
|
||||
"""
|
||||
|
||||
def __init__(self, dim_in: int, dim_out: int, bias: bool = True):
|
||||
super().__init__()
|
||||
self.proj = nn.Linear(dim_in, dim_out * 2, bias=bias)
|
||||
self.activation = nn.SiLU()
|
||||
|
||||
def forward(self, hidden_states):
|
||||
hidden_states = self.proj(hidden_states)
|
||||
hidden_states, gate = hidden_states.chunk(2, dim=-1)
|
||||
return hidden_states * self.activation(gate)
|
||||
|
||||
|
||||
class ApproximateGELU(nn.Module):
|
||||
r"""
|
||||
The approximate form of the Gaussian Error Linear Unit (GELU). For more details, see section 2 of this
|
||||
[paper](https://arxiv.org/abs/1606.08415).
|
||||
|
||||
Parameters:
|
||||
dim_in (`int`): The number of channels in the input.
|
||||
dim_out (`int`): The number of channels in the output.
|
||||
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
|
||||
"""
|
||||
|
||||
def __init__(self, dim_in: int, dim_out: int, bias: bool = True):
|
||||
super().__init__()
|
||||
self.proj = nn.Linear(dim_in, dim_out, bias=bias)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.proj(x)
|
||||
return x * torch.sigmoid(1.702 * x)
|
||||
|
||||
|
||||
def randn_tensor(
|
||||
shape: Union[Tuple, List],
|
||||
generator: Optional[Union[List["torch.Generator"], "torch.Generator"]] = None,
|
||||
device: Optional["torch.device"] = None,
|
||||
dtype: Optional["torch.dtype"] = None,
|
||||
layout: Optional["torch.layout"] = None,
|
||||
):
|
||||
"""A helper function to create random tensors on the desired `device` with the desired `dtype`. When
|
||||
passing a list of generators, you can seed each batch size individually. If CPU generators are passed, the tensor
|
||||
is always created on the CPU.
|
||||
"""
|
||||
# device on which tensor is created defaults to device
|
||||
rand_device = device
|
||||
batch_size = shape[0]
|
||||
|
||||
layout = layout or torch.strided
|
||||
device = device or torch.device("cpu")
|
||||
|
||||
if generator is not None:
|
||||
gen_device_type = generator.device.type if not isinstance(generator, list) else generator[0].device.type
|
||||
if gen_device_type != device.type and gen_device_type == "cpu":
|
||||
rand_device = "cpu"
|
||||
if device != "mps":
|
||||
print(
|
||||
f"The passed generator was created on 'cpu' even though a tensor on {device} was expected."
|
||||
f" Tensors will be created on 'cpu' and then moved to {device}. Note that one can probably"
|
||||
f" slighly speed up this function by passing a generator that was created on the {device} device."
|
||||
)
|
||||
elif gen_device_type != device.type and gen_device_type == "cuda":
|
||||
raise ValueError(f"Cannot generate a {device} tensor from a generator of type {gen_device_type}.")
|
||||
|
||||
# make sure generator list of length 1 is treated like a non-list
|
||||
if isinstance(generator, list) and len(generator) == 1:
|
||||
generator = generator[0]
|
||||
|
||||
if isinstance(generator, list):
|
||||
shape = (1,) + shape[1:]
|
||||
latents = [
|
||||
torch.randn(shape, generator=generator[i], device=rand_device, dtype=dtype, layout=layout)
|
||||
for i in range(batch_size)
|
||||
]
|
||||
latents = torch.cat(latents, dim=0).to(device)
|
||||
else:
|
||||
latents = torch.randn(shape, generator=generator, device=rand_device, dtype=dtype, layout=layout).to(device)
|
||||
|
||||
return latents
|
||||
|
||||
|
||||
def get_timestep_embedding(
|
||||
timesteps: torch.Tensor,
|
||||
embedding_dim: int,
|
||||
flip_sin_to_cos: bool = False,
|
||||
downscale_freq_shift: float = 1,
|
||||
scale: float = 1,
|
||||
max_period: int = 10000,
|
||||
):
|
||||
"""
|
||||
This matches the implementation in Denoising Diffusion Probabilistic Models: Create sinusoidal timestep embeddings.
|
||||
|
||||
Args
|
||||
timesteps (torch.Tensor):
|
||||
a 1-D Tensor of N indices, one per batch element. These may be fractional.
|
||||
embedding_dim (int):
|
||||
the dimension of the output.
|
||||
flip_sin_to_cos (bool):
|
||||
Whether the embedding order should be `cos, sin` (if True) or `sin, cos` (if False)
|
||||
downscale_freq_shift (float):
|
||||
Controls the delta between frequencies between dimensions
|
||||
scale (float):
|
||||
Scaling factor applied to the embeddings.
|
||||
max_period (int):
|
||||
Controls the maximum frequency of the embeddings
|
||||
Returns
|
||||
torch.Tensor: an [N x dim] Tensor of positional embeddings.
|
||||
"""
|
||||
assert len(timesteps.shape) == 1, "Timesteps should be a 1d-array"
|
||||
|
||||
half_dim = embedding_dim // 2
|
||||
exponent = -math.log(max_period) * torch.arange(
|
||||
start=0, end=half_dim, dtype=torch.float32, device=timesteps.device
|
||||
)
|
||||
exponent = exponent / (half_dim - downscale_freq_shift)
|
||||
|
||||
emb = torch.exp(exponent)
|
||||
emb = timesteps[:, None].float() * emb[None, :]
|
||||
|
||||
# scale embeddings
|
||||
emb = scale * emb
|
||||
|
||||
# concat sine and cosine embeddings
|
||||
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1)
|
||||
|
||||
# flip sine and cosine embeddings
|
||||
if flip_sin_to_cos:
|
||||
emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1)
|
||||
|
||||
# zero pad
|
||||
if embedding_dim % 2 == 1:
|
||||
emb = torch.nn.functional.pad(emb, (0, 1, 0, 0))
|
||||
return emb
|
||||
|
||||
|
||||
|
||||
def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
|
||||
"""
|
||||
embed_dim: output dimension for each position pos: a list of positions to be encoded: size (M,) out: (M, D)
|
||||
"""
|
||||
if embed_dim % 2 != 0:
|
||||
raise ValueError("embed_dim must be divisible by 2")
|
||||
|
||||
omega = np.arange(embed_dim // 2, dtype=np.float64)
|
||||
omega /= embed_dim / 2.0
|
||||
omega = 1.0 / 10000**omega # (D/2,)
|
||||
|
||||
pos = pos.reshape(-1) # (M,)
|
||||
out = np.einsum("m,d->md", pos, omega) # (M, D/2), outer product
|
||||
|
||||
emb_sin = np.sin(out) # (M, D/2)
|
||||
emb_cos = np.cos(out) # (M, D/2)
|
||||
|
||||
emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D)
|
||||
return emb
|
||||
|
||||
|
||||
def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):
|
||||
if embed_dim % 2 != 0:
|
||||
raise ValueError("embed_dim must be divisible by 2")
|
||||
|
||||
# use half of dimensions to encode grid_h
|
||||
emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2)
|
||||
emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2)
|
||||
|
||||
emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D)
|
||||
return emb
|
||||
|
||||
def get_3d_sincos_pos_embed(
|
||||
embed_dim: int,
|
||||
spatial_size: Union[int, Tuple[int, int]],
|
||||
temporal_size: int,
|
||||
spatial_interpolation_scale: float = 1.0,
|
||||
temporal_interpolation_scale: float = 1.0,
|
||||
) -> np.ndarray:
|
||||
r"""
|
||||
Args:
|
||||
embed_dim (`int`):
|
||||
spatial_size (`int` or `Tuple[int, int]`):
|
||||
temporal_size (`int`):
|
||||
spatial_interpolation_scale (`float`, defaults to 1.0):
|
||||
temporal_interpolation_scale (`float`, defaults to 1.0):
|
||||
"""
|
||||
if embed_dim % 4 != 0:
|
||||
raise ValueError("`embed_dim` must be divisible by 4")
|
||||
if isinstance(spatial_size, int):
|
||||
spatial_size = (spatial_size, spatial_size)
|
||||
|
||||
embed_dim_spatial = 3 * embed_dim // 4
|
||||
embed_dim_temporal = embed_dim // 4
|
||||
|
||||
# 1. Spatial
|
||||
grid_h = np.arange(spatial_size[1], dtype=np.float32) / spatial_interpolation_scale
|
||||
grid_w = np.arange(spatial_size[0], dtype=np.float32) / spatial_interpolation_scale
|
||||
grid = np.meshgrid(grid_w, grid_h) # here w goes first
|
||||
grid = np.stack(grid, axis=0)
|
||||
|
||||
grid = grid.reshape([2, 1, spatial_size[1], spatial_size[0]])
|
||||
pos_embed_spatial = get_2d_sincos_pos_embed_from_grid(embed_dim_spatial, grid)
|
||||
|
||||
# 2. Temporal
|
||||
grid_t = np.arange(temporal_size, dtype=np.float32) / temporal_interpolation_scale
|
||||
pos_embed_temporal = get_1d_sincos_pos_embed_from_grid(embed_dim_temporal, grid_t)
|
||||
|
||||
# 3. Concat
|
||||
pos_embed_spatial = pos_embed_spatial[np.newaxis, :, :]
|
||||
pos_embed_spatial = np.repeat(pos_embed_spatial, temporal_size, axis=0) # [T, H*W, D // 4 * 3]
|
||||
|
||||
pos_embed_temporal = pos_embed_temporal[:, np.newaxis, :]
|
||||
pos_embed_temporal = np.repeat(pos_embed_temporal, spatial_size[0] * spatial_size[1], axis=1) # [T, H*W, D // 4]
|
||||
|
||||
pos_embed = np.concatenate([pos_embed_temporal, pos_embed_spatial], axis=-1) # [T, H*W, D]
|
||||
return pos_embed
|
||||
|
||||
|
||||
def apply_rotary_emb(
|
||||
x: torch.Tensor,
|
||||
freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor]],
|
||||
use_real: bool = True,
|
||||
use_real_unbind_dim: int = -1,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Apply rotary embeddings to input tensors using the given frequency tensor. This function applies rotary embeddings
|
||||
to the given query or key 'x' tensors using the provided frequency tensor 'freqs_cis'. The input tensors are
|
||||
reshaped as complex numbers, and the frequency tensor is reshaped for broadcasting compatibility. The resulting
|
||||
tensors contain rotary embeddings and are returned as real tensors.
|
||||
|
||||
Args:
|
||||
x (`torch.Tensor`):
|
||||
Query or key tensor to apply rotary embeddings. [B, H, S, D] xk (torch.Tensor): Key tensor to apply
|
||||
freqs_cis (`Tuple[torch.Tensor]`): Precomputed frequency tensor for complex exponentials. ([S, D], [S, D],)
|
||||
|
||||
Returns:
|
||||
Tuple[torch.Tensor, torch.Tensor]: Tuple of modified query tensor and key tensor with rotary embeddings.
|
||||
"""
|
||||
if use_real:
|
||||
cos, sin = freqs_cis # [S, D]
|
||||
cos = cos[None, None]
|
||||
sin = sin[None, None]
|
||||
cos, sin = cos.to(x.device), sin.to(x.device)
|
||||
|
||||
if use_real_unbind_dim == -1:
|
||||
# Used for flux, cogvideox, hunyuan-dit
|
||||
x_real, x_imag = x.reshape(*x.shape[:-1], -1, 2).unbind(-1) # [B, S, H, D//2]
|
||||
x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(3)
|
||||
elif use_real_unbind_dim == -2:
|
||||
# Used for Stable Audio
|
||||
x_real, x_imag = x.reshape(*x.shape[:-1], 2, -1).unbind(-2) # [B, S, H, D//2]
|
||||
x_rotated = torch.cat([-x_imag, x_real], dim=-1)
|
||||
else:
|
||||
raise ValueError(f"`use_real_unbind_dim={use_real_unbind_dim}` but should be -1 or -2.")
|
||||
|
||||
out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype)
|
||||
|
||||
return out
|
||||
else:
|
||||
# used for lumina
|
||||
x_rotated = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
|
||||
freqs_cis = freqs_cis.unsqueeze(2)
|
||||
x_out = torch.view_as_real(x_rotated * freqs_cis).flatten(3)
|
||||
|
||||
return x_out.type_as(x)
|
||||
|
||||
|
||||
def get_1d_rotary_pos_embed(
|
||||
dim: int,
|
||||
pos: Union[np.ndarray, int],
|
||||
theta: float = 10000.0,
|
||||
use_real=False,
|
||||
linear_factor=1.0,
|
||||
ntk_factor=1.0,
|
||||
repeat_interleave_real=True,
|
||||
freqs_dtype=torch.float32, # torch.float32, torch.float64 (flux)
|
||||
):
|
||||
"""
|
||||
Precompute the frequency tensor for complex exponentials (cis) with given dimensions.
|
||||
|
||||
This function calculates a frequency tensor with complex exponentials using the given dimension 'dim' and the end
|
||||
index 'end'. The 'theta' parameter scales the frequencies. The returned tensor contains complex values in complex64
|
||||
data type.
|
||||
|
||||
Args:
|
||||
dim (`int`): Dimension of the frequency tensor.
|
||||
pos (`np.ndarray` or `int`): Position indices for the frequency tensor. [S] or scalar
|
||||
theta (`float`, *optional*, defaults to 10000.0):
|
||||
Scaling factor for frequency computation. Defaults to 10000.0.
|
||||
use_real (`bool`, *optional*):
|
||||
If True, return real part and imaginary part separately. Otherwise, return complex numbers.
|
||||
linear_factor (`float`, *optional*, defaults to 1.0):
|
||||
Scaling factor for the context extrapolation. Defaults to 1.0.
|
||||
ntk_factor (`float`, *optional*, defaults to 1.0):
|
||||
Scaling factor for the NTK-Aware RoPE. Defaults to 1.0.
|
||||
repeat_interleave_real (`bool`, *optional*, defaults to `True`):
|
||||
If `True` and `use_real`, real part and imaginary part are each interleaved with themselves to reach `dim`.
|
||||
Otherwise, they are concateanted with themselves.
|
||||
freqs_dtype (`torch.float32` or `torch.float64`, *optional*, defaults to `torch.float32`):
|
||||
the dtype of the frequency tensor.
|
||||
Returns:
|
||||
`torch.Tensor`: Precomputed frequency tensor with complex exponentials. [S, D/2]
|
||||
"""
|
||||
assert dim % 2 == 0
|
||||
|
||||
if isinstance(pos, int):
|
||||
pos = torch.arange(pos)
|
||||
if isinstance(pos, np.ndarray):
|
||||
pos = torch.from_numpy(pos) # type: ignore # [S]
|
||||
|
||||
theta = theta * ntk_factor
|
||||
freqs = (
|
||||
1.0
|
||||
/ (theta ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=pos.device)[: (dim // 2)] / dim))
|
||||
/ linear_factor
|
||||
) # [D/2]
|
||||
freqs = torch.outer(pos, freqs) # type: ignore # [S, D/2]
|
||||
if use_real and repeat_interleave_real:
|
||||
# flux, hunyuan-dit, cogvideox
|
||||
freqs_cos = freqs.cos().repeat_interleave(2, dim=1).float() # [S, D]
|
||||
freqs_sin = freqs.sin().repeat_interleave(2, dim=1).float() # [S, D]
|
||||
return freqs_cos, freqs_sin
|
||||
elif use_real:
|
||||
# stable audio
|
||||
freqs_cos = torch.cat([freqs.cos(), freqs.cos()], dim=-1).float() # [S, D]
|
||||
freqs_sin = torch.cat([freqs.sin(), freqs.sin()], dim=-1).float() # [S, D]
|
||||
return freqs_cos, freqs_sin
|
||||
else:
|
||||
# lumina
|
||||
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64 # [S, D/2]
|
||||
return freqs_cis
|
||||
|
||||
|
||||
def get_3d_rotary_pos_embed(
|
||||
embed_dim, crops_coords, grid_size, temporal_size, theta: int = 10000, use_real: bool = True
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""
|
||||
RoPE for video tokens with 3D structure.
|
||||
|
||||
Args:
|
||||
embed_dim: (`int`):
|
||||
The embedding dimension size, corresponding to hidden_size_head.
|
||||
crops_coords (`Tuple[int]`):
|
||||
The top-left and bottom-right coordinates of the crop.
|
||||
grid_size (`Tuple[int]`):
|
||||
The grid size of the spatial positional embedding (height, width).
|
||||
temporal_size (`int`):
|
||||
The size of the temporal dimension.
|
||||
theta (`float`):
|
||||
Scaling factor for frequency computation.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`: positional embedding with shape `(temporal_size * grid_size[0] * grid_size[1], embed_dim/2)`.
|
||||
"""
|
||||
if use_real is not True:
|
||||
raise ValueError(" `use_real = False` is not currently supported for get_3d_rotary_pos_embed")
|
||||
start, stop = crops_coords
|
||||
grid_size_h, grid_size_w = grid_size
|
||||
grid_h = np.linspace(start[0], stop[0], grid_size_h, endpoint=False, dtype=np.float32)
|
||||
grid_w = np.linspace(start[1], stop[1], grid_size_w, endpoint=False, dtype=np.float32)
|
||||
grid_t = np.linspace(0, temporal_size, temporal_size, endpoint=False, dtype=np.float32)
|
||||
|
||||
# Compute dimensions for each axis
|
||||
dim_t = embed_dim // 4
|
||||
dim_h = embed_dim // 8 * 3
|
||||
dim_w = embed_dim // 8 * 3
|
||||
|
||||
# Temporal frequencies
|
||||
freqs_t = get_1d_rotary_pos_embed(dim_t, grid_t, use_real=True)
|
||||
# Spatial frequencies for height and width
|
||||
freqs_h = get_1d_rotary_pos_embed(dim_h, grid_h, use_real=True)
|
||||
freqs_w = get_1d_rotary_pos_embed(dim_w, grid_w, use_real=True)
|
||||
|
||||
# BroadCast and concatenate temporal and spaial frequencie (height and width) into a 3d tensor
|
||||
def combine_time_height_width(freqs_t, freqs_h, freqs_w):
|
||||
freqs_t = freqs_t[:, None, None, :].expand(
|
||||
-1, grid_size_h, grid_size_w, -1
|
||||
) # temporal_size, grid_size_h, grid_size_w, dim_t
|
||||
freqs_h = freqs_h[None, :, None, :].expand(
|
||||
temporal_size, -1, grid_size_w, -1
|
||||
) # temporal_size, grid_size_h, grid_size_2, dim_h
|
||||
freqs_w = freqs_w[None, None, :, :].expand(
|
||||
temporal_size, grid_size_h, -1, -1
|
||||
) # temporal_size, grid_size_h, grid_size_2, dim_w
|
||||
|
||||
freqs = torch.cat(
|
||||
[freqs_t, freqs_h, freqs_w], dim=-1
|
||||
) # temporal_size, grid_size_h, grid_size_w, (dim_t + dim_h + dim_w)
|
||||
freqs = freqs.view(
|
||||
temporal_size * grid_size_h * grid_size_w, -1
|
||||
) # (temporal_size * grid_size_h * grid_size_w), (dim_t + dim_h + dim_w)
|
||||
return freqs
|
||||
|
||||
t_cos, t_sin = freqs_t # both t_cos and t_sin has shape: temporal_size, dim_t
|
||||
h_cos, h_sin = freqs_h # both h_cos and h_sin has shape: grid_size_h, dim_h
|
||||
w_cos, w_sin = freqs_w # both w_cos and w_sin has shape: grid_size_w, dim_w
|
||||
cos = combine_time_height_width(t_cos, h_cos, w_cos)
|
||||
sin = combine_time_height_width(t_sin, h_sin, w_sin)
|
||||
return cos, sin
|
||||
|
||||
|
||||
def get_resize_crop_region_for_grid(src, tgt_width, tgt_height):
|
||||
tw = tgt_width
|
||||
th = tgt_height
|
||||
h, w = src
|
||||
r = h / w
|
||||
if r > (th / tw):
|
||||
resize_height = th
|
||||
resize_width = int(round(th / h * w))
|
||||
else:
|
||||
resize_width = tw
|
||||
resize_height = int(round(tw / w * h))
|
||||
|
||||
crop_top = int(round((th - resize_height) / 2.0))
|
||||
crop_left = int(round((tw - resize_width) / 2.0))
|
||||
|
||||
return (crop_top, crop_left), (crop_top + resize_height, crop_left + resize_width)
|
||||
@@ -1 +1,3 @@
|
||||
from .flux import Flux
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from .flux import Flux
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import math
|
||||
from functools import partial
|
||||
|
||||
@@ -10,10 +12,10 @@ from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from torch import Tensor, nn
|
||||
from torch.utils.checkpoint import checkpoint_sequential
|
||||
from torch.nn.utils.rnn import pad_sequence
|
||||
from .layers import (DoubleStreamBlock, EmbedND, LastLayer, MLPEmbedder,
|
||||
SingleStreamBlock, timestep_embedding)
|
||||
|
||||
from .layers import (DoubleStreamBlock, EmbedND, LastLayer,
|
||||
MLPEmbedder, SingleStreamBlock,
|
||||
timestep_embedding)
|
||||
|
||||
@BACKBONES.register_class()
|
||||
class Flux(BaseModel):
|
||||
@@ -21,73 +23,72 @@ class Flux(BaseModel):
|
||||
Transformer backbone Diffusion model with RoPE.
|
||||
"""
|
||||
para_dict = {
|
||||
"IN_CHANNELS": {
|
||||
"value": 64,
|
||||
"description": "model's input channels."
|
||||
'IN_CHANNELS': {
|
||||
'value': 64,
|
||||
'description': "model's input channels."
|
||||
},
|
||||
"OUT_CHANNELS": {
|
||||
"value": 64,
|
||||
"description": "model's output channels."
|
||||
'OUT_CHANNELS': {
|
||||
'value': 64,
|
||||
'description': "model's output channels."
|
||||
},
|
||||
"HIDDEN_SIZE": {
|
||||
"value": 1024,
|
||||
"description": "model's hidden size."
|
||||
'HIDDEN_SIZE': {
|
||||
'value': 1024,
|
||||
'description': "model's hidden size."
|
||||
},
|
||||
"NUM_HEADS": {
|
||||
"value": 16,
|
||||
"description": "number of heads in the transformer."
|
||||
'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."
|
||||
'AXES_DIM': {
|
||||
'value': [16, 56, 56],
|
||||
'description': 'dimensions of the axes of the positional encoding.'
|
||||
},
|
||||
"THETA": {
|
||||
"value": 10_000,
|
||||
"description": "theta for positional encoding."
|
||||
'THETA': {
|
||||
'value': 10_000,
|
||||
'description': 'theta for positional encoding.'
|
||||
},
|
||||
"VEC_IN_DIM": {
|
||||
"value": 768,
|
||||
"description": "dimension of the vector input."
|
||||
'VEC_IN_DIM': {
|
||||
'value': 768,
|
||||
'description': 'dimension of the vector input.'
|
||||
},
|
||||
"GUIDANCE_EMBED": {
|
||||
"value": False,
|
||||
"description": "whether to use guidance embedding."
|
||||
'GUIDANCE_EMBED': {
|
||||
'value': False,
|
||||
'description': 'whether to use guidance embedding.'
|
||||
},
|
||||
"CONTEXT_IN_DIM": {
|
||||
"value": 4096,
|
||||
"description": "dimension of the context input."
|
||||
'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."
|
||||
'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."
|
||||
'QKV_BIAS': {
|
||||
'value': True,
|
||||
'description': 'whether to use bias in qkv projection.'
|
||||
},
|
||||
"DEPTH": {
|
||||
"value": 19,
|
||||
"description": "number of transformer blocks."
|
||||
'DEPTH': {
|
||||
'value': 19,
|
||||
'description': 'number of transformer blocks.'
|
||||
},
|
||||
"DEPTH_SINGLE_BLOCKS": {
|
||||
"value": 38,
|
||||
"description": "number of transformer blocks in the single stream block."
|
||||
'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."
|
||||
'USE_GRAD_CHECKPOINT': {
|
||||
'value': False,
|
||||
'description': 'whether to use gradient checkpointing.'
|
||||
}
|
||||
}
|
||||
def __init__(
|
||||
self,
|
||||
cfg,
|
||||
logger = None
|
||||
):
|
||||
|
||||
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)
|
||||
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
|
||||
@@ -97,7 +98,7 @@ class Flux(BaseModel):
|
||||
qkv_bias = cfg.QKV_BIAS
|
||||
depth = cfg.DEPTH
|
||||
depth_single_blocks = cfg.DEPTH_SINGLE_BLOCKS
|
||||
self.use_grad_checkpoint = cfg.get("USE_GRAD_CHECKPOINT", False)
|
||||
self.use_grad_checkpoint = cfg.get('USE_GRAD_CHECKPOINT', False)
|
||||
|
||||
if hidden_size % num_heads != 0:
|
||||
raise ValueError(
|
||||
@@ -105,56 +106,54 @@ class Flux(BaseModel):
|
||||
)
|
||||
pe_dim = hidden_size // num_heads
|
||||
if sum(axes_dim) != pe_dim:
|
||||
raise ValueError(f"Got {axes_dim} but expected positional 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.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.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.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.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 = rearrange(x, 'b c (h ph) (w pw) -> b (h w) (c ph pw)', ph=2, pw=2)
|
||||
x_id = torch.zeros(h // 2, w // 2, 3)
|
||||
x_id[..., 1] = x_id[..., 1] + torch.arange(h // 2)[:, None]
|
||||
x_id[..., 2] = x_id[..., 2] + torch.arange(w // 2)[None, :]
|
||||
x_ids = repeat(x_id, "h w c -> b (h w) c", b=bs)
|
||||
x_ids = repeat(x_id, 'h w c -> b (h w) c', b=bs)
|
||||
txt_ids = torch.zeros(bs, context.shape[1], 3)
|
||||
return x, x_ids.to(x), context.to(x), txt_ids.to(x), y.to(x), h, w
|
||||
|
||||
def unpack(self, x: Tensor, height: int, width: int) -> Tensor:
|
||||
return rearrange(
|
||||
x,
|
||||
"b (h w) (c ph pw) -> b c (h ph) (w pw)",
|
||||
h=math.ceil(height/2),
|
||||
w=math.ceil(width/2),
|
||||
'b (h w) (c ph pw) -> b c (h ph) (w pw)',
|
||||
h=math.ceil(height / 2),
|
||||
w=math.ceil(width / 2),
|
||||
ph=2,
|
||||
pw=2,
|
||||
)
|
||||
@@ -163,38 +162,42 @@ class Flux(BaseModel):
|
||||
if next(self.parameters()).device.type == 'meta':
|
||||
map_location = we.device_id
|
||||
else:
|
||||
map_location = "cpu"
|
||||
map_location = 'cpu'
|
||||
if pretrained_model is not None:
|
||||
with FS.get_from(pretrained_model, wait_finish=True) as local_model:
|
||||
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)
|
||||
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}')
|
||||
self.logger.info(f'Missing Keys:\n {missing}') # noqa
|
||||
if len(unexpected) > 0:
|
||||
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
|
||||
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"])
|
||||
def forward(self,
|
||||
x: Tensor,
|
||||
t: Tensor,
|
||||
cond: dict = {},
|
||||
guidance: Tensor | None = None,
|
||||
gc_seg: int = 0) -> Tensor:
|
||||
x, x_ids, txt, txt_ids, y, h, w = self.prepare_input(
|
||||
x, cond['context'], cond['y'])
|
||||
# running on sequences img
|
||||
x = self.img_in(x)
|
||||
vec = self.time_in(timestep_embedding(t, 256))
|
||||
if self.guidance_embed:
|
||||
if guidance is None:
|
||||
raise ValueError("Didn't get guidance strength for guidance distilled model.")
|
||||
raise ValueError(
|
||||
"Didn't get guidance strength for guidance distilled model."
|
||||
)
|
||||
vec = vec + self.guidance_in(timestep_embedding(guidance, 256))
|
||||
vec = vec + self.vector_in(y)
|
||||
txt = self.txt_in(txt)
|
||||
@@ -206,6 +209,134 @@ class Flux(BaseModel):
|
||||
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('BACKBONE',
|
||||
__class__.__name__,
|
||||
Flux.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@BACKBONES.register_class()
|
||||
class FluxMR(Flux):
|
||||
def prepare_input(self, x, cond):
|
||||
context, y = cond["context"].to(x), cond["y"].to(x)
|
||||
batch_frames, batch_frames_ids = [], []
|
||||
for ix, shape in zip(x, cond["x_shapes"]):
|
||||
# unpack image from sequence
|
||||
ix = ix[:, :shape[0] * shape[1]].view(-1, shape[0], shape[1])
|
||||
c, h, w = ix.shape
|
||||
ix = rearrange(ix, "c (h ph) (w pw) -> (h w) (c ph pw)", ph=2, pw=2)
|
||||
ix_id = torch.zeros(h // 2, w // 2, 3)
|
||||
ix_id[..., 1] = ix_id[..., 1] + torch.arange(h // 2)[:, None]
|
||||
ix_id[..., 2] = ix_id[..., 2] + torch.arange(w // 2)[None, :]
|
||||
ix_id = rearrange(ix_id, "h w c -> (h w) c")
|
||||
batch_frames.append([ix])
|
||||
batch_frames_ids.append([ix_id])
|
||||
|
||||
x_list, x_id_list, mask_x_list, x_seq_length = [], [], [], []
|
||||
for frames, frame_ids in zip(batch_frames, batch_frames_ids):
|
||||
proj_frames = []
|
||||
for idx, one_frame in enumerate(frames):
|
||||
one_frame = self.img_in(one_frame)
|
||||
proj_frames.append(one_frame)
|
||||
ix = torch.cat(proj_frames, dim=0)
|
||||
if_id = torch.cat(frame_ids, dim=0)
|
||||
x_list.append(ix)
|
||||
x_id_list.append(if_id)
|
||||
mask_x_list.append(torch.ones(ix.shape[0]).to(ix.device, non_blocking=True).bool())
|
||||
x_seq_length.append(ix.shape[0])
|
||||
x = pad_sequence(tuple(x_list), batch_first=True)
|
||||
x_ids = pad_sequence(tuple(x_id_list), batch_first=True).to(x) # [b,pad_seq,2] pad (0.,0.) at dim2
|
||||
mask_x = pad_sequence(tuple(mask_x_list), batch_first=True)
|
||||
|
||||
txt = self.txt_in(context)
|
||||
txt_ids = torch.zeros(context.shape[0], context.shape[1], 3).to(x)
|
||||
mask_txt = torch.ones(context.shape[0], context.shape[1]).to(x.device, non_blocking=True).bool()
|
||||
|
||||
return x, x_ids, txt, txt_ids, y, mask_x, mask_txt, x_seq_length
|
||||
|
||||
def unpack(self, x: Tensor, cond: dict = None, x_seq_length: list = None) -> Tensor:
|
||||
x_list = []
|
||||
image_shapes = cond["x_shapes"]
|
||||
for u, shape, seq_length in zip(x, image_shapes, x_seq_length):
|
||||
height, width = shape
|
||||
h, w = math.ceil(height / 2), math.ceil(width / 2)
|
||||
u = rearrange(
|
||||
u[seq_length-h*w:seq_length, ...],
|
||||
"(h w) (c ph pw) -> (h ph w pw) c",
|
||||
h=h,
|
||||
w=w,
|
||||
ph=2,
|
||||
pw=2,
|
||||
)
|
||||
x_list.append(u)
|
||||
x = pad_sequence(tuple(x_list), batch_first=True).permute(0, 2, 1)
|
||||
return x
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
t: Tensor,
|
||||
cond: dict = {},
|
||||
guidance: Tensor | None = None,
|
||||
gc_seg: int = 0,
|
||||
**kwargs
|
||||
) -> Tensor:
|
||||
x, x_ids, txt, txt_ids, y, mask_x, mask_txt, seq_length_list = self.prepare_input(x, cond)
|
||||
# running on sequences img
|
||||
vec = self.time_in(timestep_embedding(t, 256))
|
||||
if self.guidance_embed:
|
||||
if guidance is None:
|
||||
raise ValueError("Didn't get guidance strength for guidance distilled model.")
|
||||
vec = vec + self.guidance_in(timestep_embedding(guidance, 256))
|
||||
vec = vec + self.vector_in(y)
|
||||
ids = torch.cat((txt_ids, x_ids), dim=1)
|
||||
pe = self.pe_embedder(ids)
|
||||
|
||||
mask_aside = torch.cat((mask_txt, mask_x), dim=1)
|
||||
mask = mask_aside[:, None, :] * mask_aside[:, :, None]
|
||||
|
||||
kwargs = dict(
|
||||
vec=vec,
|
||||
pe=pe,
|
||||
mask=mask,
|
||||
txt_length = txt.shape[1],
|
||||
)
|
||||
x = torch.cat((txt, x), 1)
|
||||
if self.use_grad_checkpoint and gc_seg >= 0:
|
||||
x = checkpoint_sequential(
|
||||
functions=[partial(block, **kwargs) for block in self.double_blocks],
|
||||
@@ -220,6 +351,7 @@ class Flux(BaseModel):
|
||||
kwargs = dict(
|
||||
vec=vec,
|
||||
pe=pe,
|
||||
mask=mask,
|
||||
)
|
||||
|
||||
if self.use_grad_checkpoint and gc_seg >= 0:
|
||||
@@ -232,14 +364,14 @@ class Flux(BaseModel):
|
||||
else:
|
||||
for block in self.single_blocks:
|
||||
x = block(x, **kwargs)
|
||||
x = x[:, txt.shape[1] :, ...]
|
||||
x = x[:, txt.shape[1]:, ...]
|
||||
x = self.final_layer(x, vec) # (N, T, patch_size ** 2 * out_channels) 6 64 64
|
||||
x = self.unpack(x, h, w)
|
||||
x = self.unpack(x, cond, seq_length_list)
|
||||
return x
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODEL',
|
||||
return dict_to_yaml('BACKBONE',
|
||||
__class__.__name__,
|
||||
Flux.para_dict,
|
||||
FluxMR.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
@@ -6,32 +8,89 @@ from torch import Tensor, nn
|
||||
import torch
|
||||
from einops import rearrange, repeat
|
||||
from torch import Tensor
|
||||
from torch.nn.utils.rnn import pad_sequence
|
||||
|
||||
try:
|
||||
from flash_attn import (
|
||||
flash_attn_varlen_func
|
||||
)
|
||||
FLASHATTN_IS_AVAILABLE = True
|
||||
except ImportError:
|
||||
FLASHATTN_IS_AVAILABLE = False
|
||||
flash_attn_varlen_func = None
|
||||
|
||||
def attention(q: Tensor, k: Tensor, v: Tensor, pe: Tensor, mask: Tensor | None = None) -> Tensor:
|
||||
def attention(q: Tensor, k: Tensor, v: Tensor, pe: Tensor, mask: Tensor | None = None, backend = 'pytorch') -> Tensor:
|
||||
q, k = apply_rope(q, k, pe)
|
||||
x = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask)
|
||||
x = torch.nan_to_num(x, nan=0.0, posinf=1e10, neginf=-1e10)
|
||||
x = rearrange(x, "B H L D -> B L (H D)")
|
||||
if backend == 'pytorch':
|
||||
if mask is not None and mask.dtype == torch.bool:
|
||||
mask = torch.zeros_like(mask).to(q).masked_fill_(mask.logical_not(), -1e20)
|
||||
x = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask)
|
||||
# x = torch.nan_to_num(x, nan=0.0, posinf=1e10, neginf=-1e10)
|
||||
x = rearrange(x, "B H L D -> B L (H D)")
|
||||
elif backend == 'flash_attn':
|
||||
# q: (B, H, L, D)
|
||||
# k: (B, H, S, D) now L = S
|
||||
# v: (B, H, S, D)
|
||||
b, h, lq, d = q.shape
|
||||
_, _, lk, _ = k.shape
|
||||
q = rearrange(q, "B H L D -> B L H D")
|
||||
k = rearrange(k, "B H S D -> B S H D")
|
||||
v = rearrange(v, "B H S D -> B S H D")
|
||||
if mask is None:
|
||||
q_lens = torch.tensor([lq] * b, dtype=torch.int32).to(q.device, non_blocking=True)
|
||||
k_lens = torch.tensor([lk] * b, dtype=torch.int32).to(k.device, non_blocking=True)
|
||||
else:
|
||||
q_lens = torch.sum(mask[:, 0, :, 0], dim=1).int()
|
||||
k_lens = torch.sum(mask[:, 0, 0, :], dim=1).int()
|
||||
q = torch.cat([q_v[:q_l] for q_v, q_l in zip(q, q_lens)])
|
||||
k = torch.cat([k_v[:k_l] for k_v, k_l in zip(k, k_lens)])
|
||||
v = torch.cat([v_v[:v_l] for v_v, v_l in zip(v, k_lens)])
|
||||
cu_seqlens_q = torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(0, dtype=torch.int32)
|
||||
cu_seqlens_k = torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(0, dtype=torch.int32)
|
||||
max_seqlen_q = q_lens.max()
|
||||
max_seqlen_k = k_lens.max()
|
||||
|
||||
x = flash_attn_varlen_func(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
cu_seqlens_q=cu_seqlens_q,
|
||||
cu_seqlens_k=cu_seqlens_k,
|
||||
max_seqlen_q=max_seqlen_q,
|
||||
max_seqlen_k=max_seqlen_k
|
||||
)
|
||||
x_list = [x[cu_seqlens_q[i]:cu_seqlens_q[i+1]] for i in range(b)]
|
||||
x = pad_sequence(tuple(x_list), batch_first=True)
|
||||
x = rearrange(x, "B L H D -> B L (H D)")
|
||||
else:
|
||||
raise NotImplementedError
|
||||
return x
|
||||
|
||||
|
||||
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
|
||||
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)
|
||||
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]:
|
||||
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)
|
||||
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]):
|
||||
@@ -43,14 +102,20 @@ class EmbedND(nn.Module):
|
||||
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)],
|
||||
[
|
||||
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):
|
||||
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.
|
||||
@@ -61,14 +126,15 @@ def timestep_embedding(t: Tensor, dim, max_period=10000, time_factor: float = 10
|
||||
"""
|
||||
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
|
||||
)
|
||||
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)
|
||||
embedding = torch.cat(
|
||||
[embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||
if torch.is_floating_point(t):
|
||||
embedding = embedding.to(t)
|
||||
return embedding
|
||||
@@ -103,7 +169,8 @@ class QKNorm(torch.nn.Module):
|
||||
self.query_norm = RMSNorm(dim)
|
||||
self.key_norm = RMSNorm(dim)
|
||||
|
||||
def forward(self, q: Tensor, k: Tensor, v: Tensor) -> tuple[Tensor, Tensor]:
|
||||
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)
|
||||
@@ -119,9 +186,15 @@ class SelfAttention(nn.Module):
|
||||
self.norm = QKNorm(head_dim)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
|
||||
def forward(self, x: Tensor, pe: Tensor, mask: Tensor | None = None) -> Tensor:
|
||||
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, 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)
|
||||
@@ -152,7 +225,7 @@ class Modulation(nn.Module):
|
||||
|
||||
|
||||
class DoubleStreamBlock(nn.Module):
|
||||
def __init__(self, hidden_size: int, num_heads: int, mlp_ratio: float, qkv_bias: bool = False):
|
||||
def __init__(self, hidden_size: int, num_heads: int, mlp_ratio: float, qkv_bias: bool = False, backend = 'pytorch'):
|
||||
super().__init__()
|
||||
|
||||
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
||||
@@ -169,6 +242,8 @@ class DoubleStreamBlock(nn.Module):
|
||||
nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
|
||||
)
|
||||
|
||||
self.backend = backend
|
||||
|
||||
self.txt_mod = Modulation(hidden_size, double=True)
|
||||
self.txt_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.txt_attn = SelfAttention(dim=hidden_size, num_heads=num_heads, qkv_bias=qkv_bias)
|
||||
@@ -205,7 +280,7 @@ class DoubleStreamBlock(nn.Module):
|
||||
v = torch.cat((txt_v, img_v), dim=2)
|
||||
if mask is not None:
|
||||
mask = repeat(mask, 'B L S-> B H L S', H=self.num_heads)
|
||||
attn = attention(q, k, v, pe=pe, mask = mask)
|
||||
attn = attention(q, k, v, pe=pe, mask = mask, backend = self.backend)
|
||||
txt_attn, img_attn = attn[:, : txt.shape[1]], attn[:, txt.shape[1] :]
|
||||
|
||||
# calculate the img bloks
|
||||
@@ -224,13 +299,13 @@ 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,
|
||||
backend='pytorch'
|
||||
):
|
||||
super().__init__()
|
||||
self.hidden_dim = hidden_size
|
||||
@@ -240,29 +315,43 @@ class SingleStreamBlock(nn.Module):
|
||||
|
||||
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)
|
||||
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.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.pre_norm = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6)
|
||||
|
||||
self.mlp_act = nn.GELU(approximate="tanh")
|
||||
self.mlp_act = nn.GELU(approximate='tanh')
|
||||
self.modulation = Modulation(hidden_size, double=False)
|
||||
self.backend = backend
|
||||
|
||||
def forward(self, x: Tensor, vec: Tensor, pe: Tensor, mask: Tensor = None) -> Tensor:
|
||||
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)
|
||||
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, 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)
|
||||
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
|
||||
@@ -271,9 +360,14 @@ class SingleStreamBlock(nn.Module):
|
||||
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))
|
||||
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)
|
||||
|
||||
@@ -1,2 +1,3 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from .sd3 import MMDiT
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -9,6 +9,7 @@ import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
|
||||
from scepter.modules.model.backbone.transformer.attention import drop_path
|
||||
|
||||
|
||||
@@ -299,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
|
||||
|
||||
@@ -8,8 +8,13 @@ from itertools import repeat as iter_repeat
|
||||
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):
|
||||
@@ -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)
|
||||
|
||||
@@ -1,3 +1,7 @@
|
||||
from .samplers import BaseDiffusionSampler, FlowEluerSampler, DDIMSampler
|
||||
from .schedules import BaseNoiseScheduler, ScaledLinearScheduler, FlowMatchShiftScheduler
|
||||
from .diffusions import BaseDiffusion, DiffusionFluxRF
|
||||
# -*- 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)
|
||||
|
||||
@@ -1,28 +1,31 @@
|
||||
import os
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import math
|
||||
import torch
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
|
||||
from scepter.modules.utils.config import dict_to_yaml, Config
|
||||
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
|
||||
from scepter.modules.model.registry import DIFFUSIONS, NOISE_SCHEDULERS, DIFFUSION_SAMPLERS
|
||||
from tqdm import trange
|
||||
|
||||
|
||||
@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."
|
||||
'NOISE_SCHEDULER': {},
|
||||
'SAMPLER_SCHEDULER': {},
|
||||
'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
|
||||
@@ -30,39 +33,56 @@ class BaseDiffusion(object):
|
||||
self.init_params()
|
||||
|
||||
def init_params(self):
|
||||
self.min_snr_gamma = self.cfg.get("MIN_SNR_GAMMA", None)
|
||||
self.prediction_type = self.cfg.get("PREDICTION_TYPE", "eps")
|
||||
self.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),
|
||||
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")
|
||||
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")
|
||||
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):
|
||||
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,
|
||||
reverse_scale = -1.,
|
||||
x = None,
|
||||
**kwargs):
|
||||
assert isinstance(steps, (int, torch.LongTensor))
|
||||
assert return_intermediate in (None, 'x0', 'xt')
|
||||
assert isinstance(sampler, (str, dict, Config))
|
||||
intermediates = []
|
||||
|
||||
def callback_fn(x_t, t, sigma=None, alpha=None):
|
||||
def callback_fn(x_t, t, sigma=None, alpha_bar=None):
|
||||
timestamp = t
|
||||
t = t.repeat(len(x_t)).round().long().to(x_t.device)
|
||||
sigma = sigma.repeat(len(x_t), *([1] * (len(sigma.shape) - 1)))
|
||||
alpha = alpha.repeat(len(x_t), *([1] * (len(alpha.shape) - 1)))
|
||||
alpha_bar = alpha_bar.repeat(len(x_t), *([1] * (len(alpha_bar.shape) - 1)))
|
||||
|
||||
if guide_scale is None or guide_scale == 1.0:
|
||||
out = model(x=x_t, t=t, **model_kwargs)
|
||||
else:
|
||||
if use_dynamic_cfg:
|
||||
guidance_scale = 1 + guide_scale * ((1 - math.cos(math.pi * ((steps - timestamp.item()) / steps) ** 5.0)) / 2)
|
||||
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])
|
||||
@@ -78,15 +98,12 @@ class BaseDiffusion(object):
|
||||
if self.prediction_type == 'x0':
|
||||
x0 = out
|
||||
elif self.prediction_type == 'eps':
|
||||
x0 = (x_t - sigma * out) / alpha
|
||||
x0 = (x_t - sigma * out) / alpha_bar
|
||||
elif self.prediction_type == 'v':
|
||||
x0 = alpha * x_t - sigma * out
|
||||
x0 = alpha_bar * x_t - sigma * out
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f'prediction_type {self.prediction_type} not implemented')
|
||||
|
||||
# print("torch.sum(y_out):", torch.sum(y_out), "torch.sum(u_out):", torch.sum(u_out), "torch.sum(out):",
|
||||
# torch.sum(out), "torch.sum(x0):", torch.sum(x0), "sigmas", sigma, "alphas", alpha)
|
||||
return x0
|
||||
|
||||
sampler_ins = self.get_sampler(sampler)
|
||||
@@ -94,13 +111,14 @@ class BaseDiffusion(object):
|
||||
# this is ignored for schnell
|
||||
sampler_output = sampler_ins.preprare_sampler(
|
||||
noise,
|
||||
x = x,
|
||||
steps=steps,
|
||||
reverse_scale= reverse_scale,
|
||||
prediction_type=self.prediction_type,
|
||||
scheduler_ins=self.sampler_scheduler,
|
||||
callback_fn=callback_fn
|
||||
)
|
||||
callback_fn=callback_fn)
|
||||
|
||||
for _ in trange(steps, disable=not show_progress):
|
||||
for _ in trange(sampler_output.steps, disable=not show_progress):
|
||||
trange.desc = sampler_output.msg
|
||||
sampler_output = sampler_ins.step(sampler_output)
|
||||
if return_intermediate == 'x_0':
|
||||
@@ -109,53 +127,56 @@ class BaseDiffusion(object):
|
||||
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
|
||||
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):
|
||||
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)
|
||||
x_t, t, sigma, alpha = schedule_output.x_t, schedule_output.t, schedule_output.sigma, schedule_output.alpha
|
||||
schedule_output = self.noise_scheduler.add_noise(x_0, noise, **kwargs)
|
||||
x_t, t, sigma, alpha_bar = schedule_output.x_t, schedule_output.t, schedule_output.sigma, schedule_output.alpha_bar
|
||||
out = model(x=x_t, t=t, **model_kwargs)
|
||||
|
||||
# mse loss
|
||||
target = {
|
||||
'eps': noise,
|
||||
'x0': x_0,
|
||||
'v': alpha * noise - sigma * x_0
|
||||
'v': alpha_bar * noise - sigma * x_0
|
||||
}[self.prediction_type]
|
||||
|
||||
loss = (out - target).pow(2)
|
||||
if reduction == 'mean':
|
||||
loss = loss.flatten(1).mean(dim=1)
|
||||
|
||||
if self.min_snr_gamma is not None:
|
||||
alphas = self.noise_scheduler.alphas.to(x_0.device)[t]
|
||||
sigmas = self.noise_scheduler.sigmas.pow(2).to(x_0.device)[t]
|
||||
snrs = (alphas / sigmas).clamp(min=1e-20)
|
||||
min_snrs = snrs.clamp(max=self.min_snr_gamma)
|
||||
weights = min_snrs / snrs
|
||||
else:
|
||||
weights = 1
|
||||
|
||||
loss = loss * weights
|
||||
return loss
|
||||
|
||||
def get_sampler(self, sampler):
|
||||
if isinstance(sampler, str):
|
||||
if sampler not in DIFFUSION_SAMPLERS.class_map:
|
||||
if self.logger is not None:
|
||||
self.logger.info(f"{sampler} not in the defined samplers list {DIFFUSION_SAMPLERS.class_map.keys()}")
|
||||
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()}")
|
||||
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)
|
||||
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 = 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
|
||||
@@ -171,49 +192,47 @@ class BaseDiffusion(object):
|
||||
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."
|
||||
'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")
|
||||
self.prediction_type = self.cfg.get('PREDICTION_TYPE', 'raw')
|
||||
|
||||
def loss(self, x_0, model, model_kwargs={}, reduction='mean', noise=None, **kwargs):
|
||||
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)
|
||||
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":
|
||||
if self.prediction_type == 'raw':
|
||||
target = noise - x_0
|
||||
out = out
|
||||
elif self.prediction_type == "sigma_scaled":
|
||||
elif self.prediction_type == 'sigma_scaled':
|
||||
target = x_0
|
||||
out = out * (-sigma) + x_t
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
loss = (target - out) ** 2
|
||||
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()
|
||||
@@ -222,10 +241,12 @@ class DiffusionFluxRF(BaseDiffusion):
|
||||
model,
|
||||
model_kwargs={},
|
||||
steps=20,
|
||||
sampler = None,
|
||||
sampler=None,
|
||||
show_progress=False,
|
||||
return_intermediate=None,
|
||||
intermediate_callback=None,
|
||||
reverse_scale=-1.,
|
||||
x=None,
|
||||
**kwargs):
|
||||
# sanity check
|
||||
assert isinstance(steps, (int, torch.LongTensor))
|
||||
@@ -233,9 +254,12 @@ class DiffusionFluxRF(BaseDiffusion):
|
||||
assert isinstance(sampler, (str, dict, Config))
|
||||
intermediates = []
|
||||
|
||||
def callback_fn(x_t, t, sigma):
|
||||
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)
|
||||
def callback_fn(x_t, t, sigma=None, alpha_bar=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)
|
||||
@@ -243,13 +267,14 @@ class DiffusionFluxRF(BaseDiffusion):
|
||||
# this is ignored for schnell
|
||||
sampler_output = sampler_ins.preprare_sampler(
|
||||
noise,
|
||||
steps = steps,
|
||||
x=x,
|
||||
steps=steps,
|
||||
reverse_scale=reverse_scale,
|
||||
prediction_type=self.prediction_type,
|
||||
scheduler_ins = self.sampler_scheduler,
|
||||
callback_fn=callback_fn
|
||||
)
|
||||
scheduler_ins=self.sampler_scheduler,
|
||||
callback_fn=callback_fn)
|
||||
|
||||
for _ in trange(steps, disable=not show_progress):
|
||||
for _ in trange(sampler_output.steps, disable=not show_progress):
|
||||
trange.desc = sampler_output.msg
|
||||
sampler_output = sampler_ins.step(sampler_output)
|
||||
if return_intermediate == 'x_0':
|
||||
@@ -258,11 +283,12 @@ class DiffusionFluxRF(BaseDiffusion):
|
||||
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
|
||||
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)
|
||||
set_name=True)
|
||||
|
||||
@@ -1,32 +1,42 @@
|
||||
from dataclasses import dataclass, field
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
|
||||
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
|
||||
alphas_bar: torch.Tensor
|
||||
betas: torch.Tensor
|
||||
sigmas: torch.Tensor
|
||||
alphas_init: torch.Tensor
|
||||
alphas_bar_init: torch.Tensor
|
||||
betas_init: torch.Tensor
|
||||
sigmas_init: torch.Tensor
|
||||
ts: torch.Tensor
|
||||
x_t: torch.Tensor
|
||||
x_0: torch.Tensor
|
||||
step: int
|
||||
steps: int
|
||||
msg: str
|
||||
|
||||
def add_custom_field(self, key: str, value) -> None:
|
||||
self.__setattr__(key, value)
|
||||
|
||||
|
||||
@DIFFUSION_SAMPLERS.register_class("base")
|
||||
@DIFFUSION_SAMPLERS.register_class('base')
|
||||
class BaseDiffusionSampler(object):
|
||||
para_dict = {}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(BaseDiffusionSampler, self).__init__()
|
||||
self.logger = logger
|
||||
@@ -34,13 +44,15 @@ class BaseDiffusionSampler(object):
|
||||
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)
|
||||
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):
|
||||
def discretization(self, steps=20, num_timesteps=1000, reverse_scale = -1., **kwargs):
|
||||
# get timesteps
|
||||
if isinstance(steps, int):
|
||||
steps += 1 if self.discard_penultimate_step else 0
|
||||
@@ -60,15 +72,29 @@ class BaseDiffusionSampler(object):
|
||||
steps = torch.tensor(self.free_steps)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f'{self.discretization_type} discretization not implemented')
|
||||
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
|
||||
if reverse_scale >=0:
|
||||
img2img_step = int((1 - reverse_scale) * len(steps))
|
||||
timesteps = torch.as_tensor(steps[img2img_step:], dtype=torch.float32)
|
||||
return timesteps
|
||||
return torch.as_tensor(steps, dtype=torch.float32)
|
||||
|
||||
def preprare_sampler(self, noise, steps=20, scheduler_ins=None, prediction_type="",
|
||||
sigmas=None, betas=None, alphas=None, callback_fn = None,
|
||||
def preprare_sampler(self,
|
||||
noise,
|
||||
x=None,
|
||||
steps=20,
|
||||
reverse_scale=-1.,
|
||||
scheduler_ins=None,
|
||||
prediction_type='',
|
||||
sigmas=None,
|
||||
betas=None,
|
||||
alphas=None,
|
||||
alphas_bar=None,
|
||||
callback_fn=None,
|
||||
**kwargs):
|
||||
'''
|
||||
1. Control the model's inputs and outputs externally in the solver by callback_fn,
|
||||
@@ -79,34 +105,57 @@ class BaseDiffusionSampler(object):
|
||||
4. To ensure the safety of threading, use the instance of SamplerOutput as the manager,
|
||||
which manage all necessary information.
|
||||
'''
|
||||
if reverse_scale >= 0:
|
||||
assert x is not None
|
||||
num_timesteps = scheduler_ins.num_timesteps if scheduler_ins is not None else 1000
|
||||
timestamps = self.discretization(steps, num_timesteps=num_timesteps, **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
|
||||
timestamps = self.discretization(steps,
|
||||
num_timesteps=num_timesteps,
|
||||
reverse_scale=reverse_scale,
|
||||
**kwargs)
|
||||
alphas = scheduler_ins.t_to_alpha(
|
||||
timestamps, **kwargs) if scheduler_ins is not None else alphas
|
||||
alphas_bar = scheduler_ins.t_to_alpha_bar(
|
||||
timestamps, **kwargs) if scheduler_ins is not None else alphas_bar
|
||||
betas = scheduler_ins.t_to_beta(
|
||||
timestamps, **kwargs) if scheduler_ins is not None else betas
|
||||
sigmas = scheduler_ins.t_to_sigma(
|
||||
timestamps, **kwargs) if scheduler_ins is not None else sigmas
|
||||
alphas_init = scheduler_ins.t_to_alpha_init(
|
||||
timestamps, **kwargs) if scheduler_ins is not None else alphas
|
||||
|
||||
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=f"step 0"
|
||||
)
|
||||
alphas_bar_init = scheduler_ins.t_to_alpha_bar_init(
|
||||
timestamps, **kwargs) if scheduler_ins is not None else alphas_bar
|
||||
|
||||
betas_init = scheduler_ins.t_to_beta_init(
|
||||
timestamps, **kwargs) if scheduler_ins is not None else betas
|
||||
sigmas_init = scheduler_ins.t_to_sigma_init(
|
||||
timestamps, **kwargs) if scheduler_ins is not None else sigmas
|
||||
if reverse_scale >= 0:
|
||||
x_t = x_0 = scheduler_ins.add_noise(x, noise=noise, t=timestamps[0].repeat(x.size(0)).to(x.device)).x_t if len(timestamps) > 0 else x
|
||||
else:
|
||||
x_t = x_0 = noise
|
||||
# Consider the sigma's list is from sigma_ to zero. the steps equal to len(timestamps)
|
||||
output = SamplerOutput(callback_fn=callback_fn,
|
||||
prediction_type=prediction_type,
|
||||
alphas=alphas,
|
||||
alphas_bar=alphas_bar,
|
||||
betas=betas,
|
||||
sigmas=sigmas,
|
||||
alphas_init=alphas_init,
|
||||
alphas_bar_init=alphas_bar_init,
|
||||
betas_init=betas_init,
|
||||
sigmas_init=sigmas_init,
|
||||
ts=timestamps,
|
||||
x_t=x_t,
|
||||
x_0=x_0,
|
||||
step=0,
|
||||
msg='step 0',
|
||||
steps=len(timestamps) - 1)
|
||||
return output
|
||||
|
||||
def step(self, sampler_ouput):
|
||||
raise NotImplementedError(f'DiffusionSampler step function not implemented')
|
||||
raise NotImplementedError(
|
||||
'DiffusionSampler step function not implemented')
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f'{self.__class__.__name__}' + ' ' + super().__repr__()
|
||||
@@ -119,29 +168,51 @@ class BaseDiffusionSampler(object):
|
||||
set_name=True)
|
||||
|
||||
|
||||
@DIFFUSION_SAMPLERS.register_class("eluer")
|
||||
@DIFFUSION_SAMPLERS.register_class('eluer')
|
||||
class EulerSampler(BaseDiffusionSampler):
|
||||
def step(self, sampler_ouput):
|
||||
pass
|
||||
|
||||
|
||||
@DIFFUSION_SAMPLERS.register_class("ddim")
|
||||
@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")
|
||||
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,
|
||||
def preprare_sampler(self,
|
||||
noise,
|
||||
x=None,
|
||||
steps=20,
|
||||
reverse_scale = -1.,
|
||||
scheduler_ins=None,
|
||||
prediction_type='',
|
||||
sigmas=None,
|
||||
betas=None,
|
||||
alphas=None,
|
||||
alphas_bar=None,
|
||||
callback_fn=None,
|
||||
**kwargs):
|
||||
output = super().preprare_sampler(noise, steps, scheduler_ins, prediction_type, sigmas, betas, alphas, callback_fn, **kwargs)
|
||||
output = super().preprare_sampler(noise,
|
||||
x = x,
|
||||
steps = steps,
|
||||
reverse_scale = reverse_scale,
|
||||
scheduler_ins = scheduler_ins,
|
||||
prediction_type = prediction_type,
|
||||
sigmas = sigmas,
|
||||
betas = betas,
|
||||
alphas = alphas,
|
||||
alphas_bar = alphas_bar,
|
||||
callback_fn = callback_fn,
|
||||
**kwargs)
|
||||
sigmas = output.sigmas
|
||||
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
|
||||
sigmas_vp = (sigmas**2 / (1 + sigmas**2))**0.5
|
||||
sigmas_vp[sigmas == float('inf')] = 1.
|
||||
output.add_custom_field('sigmas_vp', sigmas_vp)
|
||||
output.steps += 1
|
||||
return output
|
||||
|
||||
def step(self, sampler_output):
|
||||
@@ -149,46 +220,68 @@ class DDIMSampler(BaseDiffusionSampler):
|
||||
step = sampler_output.step
|
||||
t = sampler_output.ts[step]
|
||||
sigmas_vp = sampler_output.sigmas_vp.to(x_t.device)
|
||||
alpha_init = _i(sampler_output.alphas_init, step, x_t[:1])
|
||||
alpha_bar_init = _i(sampler_output.alphas_bar_init, step, x_t[:1])
|
||||
sigma_init = _i(sampler_output.sigmas_init, step, x_t[:1])
|
||||
|
||||
x = sampler_output.callback_fn(x_t, t, sigma_init, alpha_init)
|
||||
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 = sampler_output.callback_fn(x_t, t, sigma_init, alpha_bar_init)
|
||||
noise_factor = self.eta * (sigmas_vp[step + 1]**2 /
|
||||
sigmas_vp[step]**2 *
|
||||
(1 - (1 - sigmas_vp[step]**2) /
|
||||
(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)
|
||||
# print("i:", step, "sigma_init:", sigma_init, "alpha_init", alpha_init, "sigmas_vp[i]", sigmas_vp[step], "torch.sum(x_0):", torch.sum(x_0), "torch.sum(x):", torch.sum(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")
|
||||
@DIFFUSION_SAMPLERS.register_class('flow_euler')
|
||||
class FlowEluerSampler(BaseDiffusionSampler):
|
||||
def preprare_sampler(self, noise, steps=20, scheduler_ins=None, prediction_type="",
|
||||
sigmas=None, betas=None, alphas=None, callback_fn = None,
|
||||
def preprare_sampler(self,
|
||||
noise,
|
||||
x=None,
|
||||
steps=20,
|
||||
reverse_scale = -1.,
|
||||
scheduler_ins=None,
|
||||
prediction_type='',
|
||||
sigmas=None,
|
||||
betas=None,
|
||||
alphas=None,
|
||||
alphas_bar=None,
|
||||
callback_fn=None,
|
||||
**kwargs):
|
||||
if noise.ndim == 3:
|
||||
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)
|
||||
kwargs['seq_len'] = seq_len
|
||||
output = super().preprare_sampler(noise,
|
||||
x = x,
|
||||
steps = steps,
|
||||
reverse_scale = reverse_scale,
|
||||
scheduler_ins = scheduler_ins,
|
||||
prediction_type = prediction_type,
|
||||
sigmas = sigmas,
|
||||
betas = betas,
|
||||
alphas = alphas,
|
||||
alphas_bar = alphas_bar,
|
||||
callback_fn = callback_fn,
|
||||
**kwargs)
|
||||
return output
|
||||
|
||||
def step(self, sampler_output):
|
||||
step = sampler_output.step
|
||||
x_t = sampler_output.x_t
|
||||
sigma_curr, sigma_prev = sampler_output.sigmas[step], sampler_output.sigmas[step + 1]
|
||||
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")
|
||||
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
|
||||
@@ -198,9 +291,13 @@ class FlowEluerSampler(BaseDiffusionSampler):
|
||||
sampler_output.msg = f'step {step}, sigma_curr: {sigma_curr}, sigma_prev: {sigma_prev}'
|
||||
return sampler_output
|
||||
|
||||
def discretization(self, steps=20, num_timesteps = 1000, **kwargs):
|
||||
def discretization(self, steps=20, num_timesteps=1000, reverse_scale=-1., **kwargs):
|
||||
# extra step for zero
|
||||
timesteps = torch.linspace(num_timesteps, 0, steps + 1)
|
||||
if reverse_scale >= 0:
|
||||
img2img_step = int((1 - reverse_scale) * len(timesteps))
|
||||
timesteps = timesteps[img2img_step:]
|
||||
return timesteps
|
||||
return timesteps
|
||||
|
||||
@staticmethod
|
||||
@@ -208,4 +305,4 @@ class FlowEluerSampler(BaseDiffusionSampler):
|
||||
return dict_to_yaml('DIFFUSION_SAMPLERS',
|
||||
__class__.__name__,
|
||||
FlowEluerSampler.para_dict,
|
||||
set_name=True)
|
||||
set_name=True)
|
||||
|
||||
@@ -1,24 +1,27 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import math
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.math_plot import plot_multi_curves
|
||||
from scepter.modules.model.registry import NOISE_SCHEDULERS
|
||||
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
|
||||
alpha_bar: torch.Tensor
|
||||
custom_fields: dict = field(default_factory=dict)
|
||||
|
||||
def add_custom_field(self, key: str, value) -> None:
|
||||
@@ -33,6 +36,7 @@ class BaseNoiseScheduler(object):
|
||||
be the basic property for the instance of noise scheduler.
|
||||
\alpha_{t} = \sqrt{1 - \beta_{t}^2} \alpha is the strength of signal and \beta is the strength of noise
|
||||
\sigma_{t} = \sqrt{1 - \overline\alpha} = \sqrt{1 - \prod_{i=1}^{t}\alpha^2_{i}} (P(x_{t}|x_{0}) ~ N(\overline\alpha x_{0}, \sigma^2))
|
||||
\alpha_bar_{t} = \sqrt{\overline\alpha} = \sqrt{\prod_{i=1}^{t}\alpha^2_{i}} (P(x_{t}|x_{0}) ~ N(\overline\alpha x_{0}, \sigma^2))
|
||||
|
||||
where sigma_{t} is the var of p(x_{t-1}|x_{t}, x_{0}).
|
||||
|
||||
@@ -42,54 +46,68 @@ class BaseNoiseScheduler(object):
|
||||
|
||||
'''
|
||||
para_dict = {
|
||||
"NUM_TIMESTEPS": {
|
||||
"value": 1000,
|
||||
"description": "The number of timesteps for sampling."
|
||||
'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()
|
||||
# self.check_function()
|
||||
|
||||
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
|
||||
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._alphas_bar, self._timesteps = None, None, None, None, None
|
||||
|
||||
def check_function(self):
|
||||
# for the same t, we should gurantee t_to_sigma and sigma_to_t is aligned
|
||||
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))
|
||||
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."
|
||||
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.")
|
||||
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.")
|
||||
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}")
|
||||
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(f'NoiseScheduler get_schedule function not implemented')
|
||||
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]])
|
||||
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)
|
||||
@@ -109,6 +127,7 @@ class BaseNoiseScheduler(object):
|
||||
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()
|
||||
@@ -124,33 +143,51 @@ class BaseNoiseScheduler(object):
|
||||
square_beta = self.sigmas_to_square_betas(sigma)
|
||||
return torch.sqrt(1 - square_beta)
|
||||
|
||||
def t_to_alpha_bar(self, t, **kwargs):
|
||||
sigma = self.t_to_sigma(t)
|
||||
return torch.sqrt(1 - sigma**2)
|
||||
|
||||
def t_to_beta(self, t, **kwargs):
|
||||
sigma = self.t_to_sigma(t)
|
||||
square_beta = self.sigmas_to_square_betas(sigma)
|
||||
return torch.sqrt(square_beta)
|
||||
|
||||
def add_noise(self, x_0, noise = None, t = None):
|
||||
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)
|
||||
t = torch.randint(0,
|
||||
self.num_timesteps, (x_0.shape[0], ),
|
||||
device=x_0.device).long()
|
||||
alpha = _i(self.alphas_bar, t, x_0)
|
||||
sigma = _i(self.sigmas, t, x_0)
|
||||
x_t = alpha * x_0 + sigma * noise
|
||||
|
||||
return ScheduleOutput(x_0 = x_0, x_t = x_t, t = t, alpha=alpha, sigma=sigma)
|
||||
return ScheduleOutput(x_0=x_0, x_t=x_t, t=t, alpha_bar=alpha, sigma=sigma)
|
||||
|
||||
def t_to_alpha_init(self, t, **kwargs):
|
||||
indices = t.long()
|
||||
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]
|
||||
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_alpha_bar_init(self, t, **kwargs):
|
||||
indices = t.long()
|
||||
indices[indices >= self.num_timesteps] = self.num_timesteps - 1
|
||||
timesteps = self.timesteps.to(t)[indices]
|
||||
step_indices = [(self.timesteps.to(t) == t).nonzero().item()
|
||||
for t in timesteps]
|
||||
alpha_bar = self.alphas_bar[step_indices].flatten().to(t)
|
||||
return alpha_bar
|
||||
|
||||
|
||||
def t_to_beta_init(self, t, **kwargs):
|
||||
indices = t.long()
|
||||
indices[indices >= self.num_timesteps] = self.num_timesteps - 1
|
||||
timesteps = self.timesteps.to(t)[indices]
|
||||
step_indices = [(self.timesteps.to(t) == t).nonzero().item() for t in timesteps]
|
||||
step_indices = [(self.timesteps.to(t) == t).nonzero().item()
|
||||
for t in timesteps]
|
||||
beta = self.betas[step_indices].flatten().to(t)
|
||||
return beta
|
||||
|
||||
@@ -158,7 +195,8 @@ class BaseNoiseScheduler(object):
|
||||
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]
|
||||
step_indices = [(self.timesteps.to(t) == t).nonzero().item()
|
||||
for t in timesteps]
|
||||
sigma = self.sigmas[step_indices].flatten().to(t)
|
||||
return sigma
|
||||
|
||||
@@ -178,9 +216,10 @@ class BaseNoiseScheduler(object):
|
||||
# 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)
|
||||
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
|
||||
alphas_bar = alphas_bar_sqrt**2 # Revert sqrt
|
||||
return alphas_bar
|
||||
|
||||
@property
|
||||
@@ -195,26 +234,40 @@ class BaseNoiseScheduler(object):
|
||||
def alphas(self):
|
||||
return self._alphas
|
||||
|
||||
@property
|
||||
def alphas_bar(self):
|
||||
return self._alphas_bar
|
||||
|
||||
@property
|
||||
def timesteps(self):
|
||||
return self._timesteps
|
||||
|
||||
# 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"}
|
||||
]
|
||||
y = [{
|
||||
'data': self._sigmas.cpu().numpy(),
|
||||
'label': 'sigmas'
|
||||
}, {
|
||||
'data': self._betas.cpu().numpy(),
|
||||
'label': 'betas'
|
||||
}, {
|
||||
'data': self._alphas.cpu().numpy(),
|
||||
'label': 'alphas'
|
||||
}, {
|
||||
'data': self._alphas_bar.cpu().numpy(),
|
||||
'label': 'alphas_bar'
|
||||
},
|
||||
{
|
||||
'data': self._timesteps.cpu().numpy() / self.num_timesteps,
|
||||
'label': 'timesteps'
|
||||
}]
|
||||
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
|
||||
)
|
||||
save_path=save_path)
|
||||
return save_path
|
||||
|
||||
def __repr__(self) -> str:
|
||||
@@ -231,18 +284,24 @@ class BaseNoiseScheduler(object):
|
||||
@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)
|
||||
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):
|
||||
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)
|
||||
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)
|
||||
@@ -250,36 +309,76 @@ class ScaledLinearScheduler(BaseNoiseScheduler):
|
||||
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)
|
||||
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._alphas = torch.sqrt(1 - square_betas)
|
||||
self._alphas_bar = torch.sqrt(1 - self._sigmas**2)
|
||||
self._timesteps = torch.arange(len(self._sigmas), dtype=torch.float32)
|
||||
|
||||
|
||||
@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 - betas**2)
|
||||
self._alphas_bar = 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 = 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)
|
||||
self._alphas = torch.sqrt(1 - self._betas**2)
|
||||
self._alphas_bar = torch.sqrt(1 - self._sigmas ** 2)
|
||||
|
||||
def add_noise(self, x_0, noise = None, t = None):
|
||||
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)
|
||||
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)
|
||||
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))
|
||||
return ScheduleOutput(x_0=x_0,
|
||||
x_t=x_t,
|
||||
t=t,
|
||||
sigma=sigma,
|
||||
alpha_bar=self.t_to_alpha_bar(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
|
||||
return t / self.num_timesteps
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
@@ -292,21 +391,22 @@ class FlowMatchUniformScheduler(BaseNoiseScheduler):
|
||||
@NOISE_SCHEDULERS.register_class()
|
||||
class FlowMatchSigmoidScheduler(FlowMatchUniformScheduler):
|
||||
para_dict = {
|
||||
"SIGMOID_SCALE": {
|
||||
"value": 1,
|
||||
"description": "The scale for the sigmoid function."
|
||||
'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)
|
||||
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
|
||||
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)
|
||||
return torch.sigmoid(self.sigmoid_scale * t / self.num_timesteps)
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
@@ -315,41 +415,46 @@ class FlowMatchSigmoidScheduler(FlowMatchUniformScheduler):
|
||||
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."
|
||||
'SHIFT': {
|
||||
'value': 3,
|
||||
'description': 'The shift factor for the timestamp.'
|
||||
},
|
||||
"SIGMOID_SCALE": {
|
||||
"value": 1,
|
||||
"description": "The scale for the sigmoid function."
|
||||
'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)
|
||||
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):
|
||||
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)
|
||||
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))
|
||||
return ScheduleOutput(x_0=x_0,
|
||||
x_t=x_t,
|
||||
t=t,
|
||||
sigma=sigma,
|
||||
alpha_bar=self.t_to_alpha_bar(t))
|
||||
|
||||
def sigma_to_t(self, sigma, **kwargs):
|
||||
t = sigma/(sigma - self.shift * sigma + self.shift)
|
||||
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)
|
||||
return (t * self.shift) / (1 + (self.shift - 1) * t)
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
@@ -358,47 +463,53 @@ class FlowMatchShiftScheduler(FlowMatchUniformScheduler):
|
||||
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."
|
||||
'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."
|
||||
'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."
|
||||
'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."
|
||||
'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.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)
|
||||
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)
|
||||
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]:
|
||||
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):
|
||||
def add_noise(self, x_0, noise=None, t=None, **kwargs):
|
||||
if x_0.ndim == 3:
|
||||
seq_len = x_0.shape[2] // 4
|
||||
else:
|
||||
@@ -409,23 +520,29 @@ class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler):
|
||||
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)
|
||||
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))
|
||||
return ScheduleOutput(x_0=x_0,
|
||||
x_t=x_t,
|
||||
t=t,
|
||||
sigma=sigma,
|
||||
alpha_bar=self.t_to_alpha_bar(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)
|
||||
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
|
||||
t = t / self.num_timesteps
|
||||
if self.shift:
|
||||
mu = self.get_lin_function(y1=self.base_shift, y2=self.max_shift)(seq_len)
|
||||
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
|
||||
@@ -437,71 +554,95 @@ class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler):
|
||||
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']."
|
||||
'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."
|
||||
'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_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."
|
||||
'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."
|
||||
'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)
|
||||
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 = 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._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)
|
||||
self._alphas = torch.sqrt(1 - self.betas**2)
|
||||
self._alphas_bar = torch.sqrt(1 - self._sigmas ** 2)
|
||||
|
||||
def add_noise(self, x_0, noise=None, t=None):
|
||||
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)
|
||||
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
|
||||
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)
|
||||
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))
|
||||
return ScheduleOutput(x_0=x_0,
|
||||
x_t=x_t,
|
||||
t=t,
|
||||
sigma=sigma,
|
||||
alpha_bar=self.t_to_alpha_bar(t))
|
||||
|
||||
def compute_density_for_timestep_sampling(self, t):
|
||||
"""Compute the density for sampling the timesteps when doing SD3 training.
|
||||
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":
|
||||
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)
|
||||
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):
|
||||
@@ -511,7 +652,8 @@ class FlowMatchSigmaScheduler(FlowMatchUniformScheduler):
|
||||
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]
|
||||
step_indices = [(self.timesteps.to(t) == t).nonzero().item()
|
||||
for t in timesteps]
|
||||
sigma = self.sigmas[step_indices].flatten().to(t)
|
||||
return sigma
|
||||
|
||||
@@ -526,8 +668,9 @@ class FlowMatchSigmaScheduler(FlowMatchUniformScheduler):
|
||||
if __name__ == '__main__':
|
||||
from scepter.modules.utils.config import Config
|
||||
cfg = Config(cfg_dict={
|
||||
"NAME": "FlowMatchShiftScheduler",
|
||||
"SHIFT": 1.15
|
||||
}, load=False)
|
||||
'NAME': 'FlowMatchShiftScheduler',
|
||||
'SHIFT': 1.15
|
||||
},
|
||||
load=False)
|
||||
|
||||
scheduler = NOISE_SCHEDULERS.build(cfg)
|
||||
|
||||
@@ -1,10 +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)
|
||||
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)
|
||||
return tensor[t].view(shape).to(x.device)
|
||||
|
||||
@@ -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,
|
||||
@@ -830,23 +832,37 @@ 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)
|
||||
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)
|
||||
tokenizer_path = cfg.get('TOKENIZER_PATH', None)
|
||||
self.length = cfg.get('LENGTH', 77)
|
||||
self.t5_dtype = cfg.get('T5_DTYPE', 'bfloat16')
|
||||
self.use_grad = cfg.get('USE_GRAD', False)
|
||||
self.clean = cfg.get('CLEAN', 'whitespace')
|
||||
self.added_identifier = cfg.get('ADDED_IDENTIFIER', None)
|
||||
tokenizer_path = cfg.get('TOKENIZER_PATH', None)
|
||||
pretrained_path = cfg.get('PRETRAINED_MODEL', None)
|
||||
|
||||
if pretrained_path:
|
||||
with FS.get_dir_to_local_dir(pretrained_path,
|
||||
wait_finish=True) as local_path:
|
||||
if self.t5_dtype is not None:
|
||||
self.model = T5EncoderModel.from_pretrained(
|
||||
local_path,
|
||||
torch_dtype=getattr(
|
||||
torch,
|
||||
'float' if self.t5_dtype == 'float32' else self.t5_dtype))
|
||||
else:
|
||||
self.model = T5EncoderModel.from_pretrained(local_path)
|
||||
else:
|
||||
self.model = None
|
||||
|
||||
if tokenizer_path:
|
||||
self.tokenize_kargs = {'return_tensors': 'pt'}
|
||||
with FS.get_dir_to_local_dir(tokenizer_path,
|
||||
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',
|
||||
@@ -859,26 +875,22 @@ class T5EmbedderHF(BaseEmbedder):
|
||||
self.tokenizer = None
|
||||
self.tokenize_kargs = {}
|
||||
|
||||
self.use_grad = cfg.get('USE_GRAD', False)
|
||||
self.clean = cfg.get('CLEAN', 'whitespace')
|
||||
|
||||
def freeze(self):
|
||||
self.model = self.model.eval()
|
||||
for param in self.parameters():
|
||||
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()
|
||||
# else:
|
||||
# return x.detach(), self.pool(x, tokens.input_ids)
|
||||
if return_mask:
|
||||
return x.detach() + 0.0, tokens.attention_mask.to(we.device_id)
|
||||
else:
|
||||
@@ -897,7 +909,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 +919,48 @@ 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)
|
||||
return self(tokens, return_mask=return_mask, use_mask=use_mask)
|
||||
|
||||
def encode_list(self, text, return_mask=False, use_mask=True):
|
||||
if isinstance(text, str):
|
||||
text = [text]
|
||||
if self.clean:
|
||||
text = [self._clean(u) for u in text]
|
||||
assert self.tokenizer is not None
|
||||
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_list_of_list(self, text_list, return_mask=True, use_mask=True):
|
||||
cont_list = []
|
||||
mask_list = []
|
||||
for pp in text_list:
|
||||
cont, cont_mask = self.encode_list(pp, return_mask=return_mask, use_mask=use_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():
|
||||
|
||||
@@ -1,99 +1,109 @@
|
||||
# -*- 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 whitespace_clean, basic_clean, canonicalize
|
||||
from scepter.modules.model.tokenizer.tokenizer_component import (
|
||||
basic_clean, canonicalize, whitespace_clean)
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
import transformers
|
||||
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"
|
||||
'HF_MODEL_CLS': {
|
||||
'value': None,
|
||||
'description': 'huggingface cls in transfomer'
|
||||
},
|
||||
"MODEL_PATH": {
|
||||
"value": None,
|
||||
"description": "model folder path"
|
||||
'MODEL_PATH': {
|
||||
'value': None,
|
||||
'description': 'model folder path'
|
||||
},
|
||||
"HF_TOKENIZER_CLS": {
|
||||
"value": None,
|
||||
"description": "huggingface cls in transfomer"
|
||||
'HF_TOKENIZER_CLS': {
|
||||
'value': None,
|
||||
'description': 'huggingface cls in transfomer'
|
||||
},
|
||||
|
||||
"TOKENIZER_PATH": {
|
||||
"value": None,
|
||||
"description": "tokenizer folder path"
|
||||
'TOKENIZER_PATH': {
|
||||
'value': None,
|
||||
'description': 'tokenizer folder path'
|
||||
},
|
||||
"MAX_LENGTH": {
|
||||
"value": 77,
|
||||
"description": "max length of input"
|
||||
'MAX_LENGTH': {
|
||||
'value': 77,
|
||||
'description': 'max length of input'
|
||||
},
|
||||
"OUTPUT_KEY": {
|
||||
"value": "last_hidden_state",
|
||||
"description": "output key"
|
||||
'OUTPUT_KEY': {
|
||||
'value': 'last_hidden_state',
|
||||
'description': 'output key'
|
||||
},
|
||||
"D_TYPE": {
|
||||
"value": "float",
|
||||
"description": "dtype"
|
||||
'D_TYPE': {
|
||||
'value': 'float',
|
||||
'description': 'dtype'
|
||||
},
|
||||
"BATCH_INFER": {
|
||||
"value": False,
|
||||
"description": "batch infer"
|
||||
'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)
|
||||
model_path = cfg.get('MODEL_PATH', None)
|
||||
hf_tokenizer_cls = cfg.get('HF_TOKENIZER_CLS', None)
|
||||
tokenizer_path = cfg.get('TOKENIZER_PATH', None)
|
||||
self.max_length = cfg.get('MAX_LENGTH', 77)
|
||||
self.output_key = cfg.get("OUTPUT_KEY", "last_hidden_state")
|
||||
self.d_type = cfg.get("D_TYPE", "float")
|
||||
self.clean = cfg.get("CLEAN", "whitespace")
|
||||
self.batch_infer = cfg.get("BATCH_INFER", False)
|
||||
self.output_key = cfg.get('OUTPUT_KEY', 'last_hidden_state')
|
||||
self.d_type = cfg.get('D_TYPE', 'float')
|
||||
self.clean = cfg.get('CLEAN', 'whitespace')
|
||||
self.batch_infer = cfg.get('BATCH_INFER', False)
|
||||
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)
|
||||
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):
|
||||
def forward(self, text: list[str], return_mask=False):
|
||||
batch_encoding = self.tokenizer(
|
||||
text,
|
||||
truncation=True,
|
||||
max_length=self.max_length,
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
padding="max_length",
|
||||
return_tensors="pt",
|
||||
padding='max_length',
|
||||
return_tensors='pt',
|
||||
)
|
||||
|
||||
outputs = self.hf_module(
|
||||
input_ids=batch_encoding["input_ids"].to(self.hf_module.device),
|
||||
input_ids=batch_encoding['input_ids'].to(self.hf_module.device),
|
||||
attention_mask=None,
|
||||
output_hidden_states=False,
|
||||
)
|
||||
if return_mask:
|
||||
return outputs[self.output_key], batch_encoding['attention_mask'].to(self.hf_module.device)
|
||||
return outputs[
|
||||
self.output_key], batch_encoding['attention_mask'].to(
|
||||
self.hf_module.device)
|
||||
else:
|
||||
return outputs[self.output_key], None
|
||||
|
||||
def encode(self, text, return_mask = False):
|
||||
def encode(self, text, return_mask=False):
|
||||
if isinstance(text, str):
|
||||
text = [text]
|
||||
if self.clean:
|
||||
@@ -109,13 +119,12 @@ class HFEmbedder(BaseEmbedder):
|
||||
else:
|
||||
return torch.cat(cont, dim=0)
|
||||
else:
|
||||
ret_data = self(text, return_mask = return_mask)
|
||||
ret_data = self(text, return_mask=return_mask)
|
||||
if return_mask:
|
||||
return ret_data
|
||||
else:
|
||||
return ret_data[0]
|
||||
|
||||
|
||||
def _clean(self, text):
|
||||
if self.clean == 'whitespace':
|
||||
text = whitespace_clean(basic_clean(text))
|
||||
@@ -124,6 +133,7 @@ class HFEmbedder(BaseEmbedder):
|
||||
elif self.clean == 'canonicalize':
|
||||
text = canonicalize(basic_clean(text))
|
||||
return text
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('EMBEDDER',
|
||||
@@ -131,15 +141,13 @@ class HFEmbedder(BaseEmbedder):
|
||||
HFEmbedder.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@EMBEDDERS.register_class()
|
||||
class T5PlusClipFluxEmbedder(BaseEmbedder):
|
||||
"""
|
||||
Uses the OpenCLIP transformer encoder for text
|
||||
"""
|
||||
para_dict = {
|
||||
'T5_MODEL': {},
|
||||
'CLIP_MODEL': {}
|
||||
}
|
||||
para_dict = {'T5_MODEL': {}, 'CLIP_MODEL': {}}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
@@ -147,8 +155,8 @@ class T5PlusClipFluxEmbedder(BaseEmbedder):
|
||||
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)
|
||||
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,
|
||||
@@ -160,4 +168,4 @@ class T5PlusClipFluxEmbedder(BaseEmbedder):
|
||||
return dict_to_yaml('EMBEDDER',
|
||||
__class__.__name__,
|
||||
T5PlusClipFluxEmbedder.para_dict,
|
||||
set_name=True)
|
||||
set_name=True)
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.network.autoencoder.ae_kl import AutoencoderKL
|
||||
from scepter.modules.model.network.autoencoder.ae_kl_cogvideox import AutoencoderKLCogVideoX
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,6 +1,8 @@
|
||||
# -*- 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,
|
||||
LatentDiffusionACERefiner)
|
||||
from scepter.modules.model.network.ldm.ldm_edit import LatentDiffusionEdit
|
||||
from scepter.modules.model.network.ldm.ldm_pixart import LatentDiffusionPixart
|
||||
from scepter.modules.model.network.ldm.ldm_sce import (
|
||||
@@ -8,3 +10,6 @@ from scepter.modules.model.network.ldm.ldm_sce import (
|
||||
LatentDiffusionXLSCEControl, LatentDiffusionXLSCETuning)
|
||||
from scepter.modules.model.network.ldm.ldm_sd3 import LatentDiffusionSD3
|
||||
from scepter.modules.model.network.ldm.ldm_xl import LatentDiffusionXL
|
||||
from scepter.modules.model.network.ldm.ldm_cogvideox import LatentDiffusionCogVideoX
|
||||
from scepter.modules.model.network.ldm.ldm_flux import (LatentDiffusionFlux,
|
||||
LatentDiffusionFluxMR)
|
||||
|
||||
@@ -0,0 +1,604 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import math
|
||||
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
|
||||
import torchvision.transforms as T
|
||||
from scepter.modules.model.utils.basic_utils import check_list_of_list
|
||||
from scepter.modules.model.utils.basic_utils import \
|
||||
pack_imagelist_into_tensor_v2 as pack_imagelist_into_tensor
|
||||
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_list_of_list')(self.text_indentifers,
|
||||
return_mask=True)
|
||||
self.text_position_embeddings.load_state_dict(
|
||||
{'pos': torch.cat( [one_id[0][0, :].unsqueeze(0) for one_id in identifier_cont], dim=0)})
|
||||
cont_, cont_mask_ = [], []
|
||||
for pp, edit, c, cm in zip(prompt, edit_image, cont, cont_mask):
|
||||
if isinstance(pp, list):
|
||||
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_of_list')(prompt_, return_mask=True)
|
||||
except Exception as e:
|
||||
print(e, prompt_)
|
||||
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont,
|
||||
cont_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_of_list')(prompt_, return_mask=True)
|
||||
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont,
|
||||
cont_mask)
|
||||
null_cont, null_cont_mask = getattr(self.cond_stage_model,
|
||||
'encode_list_of_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)
|
||||
|
||||
@MODELS.register_class()
|
||||
class LatentDiffusionACERefiner(LatentDiffusionACE):
|
||||
def init_params(self):
|
||||
super().init_params()
|
||||
self.enhence_model_cfg = self.cfg.get("ENHENCE_MODEL", None)
|
||||
self.enhence_sampler_cfg = self.cfg.get("ENHENCE_SAMPLER_CFG", {})
|
||||
def construct_network(self):
|
||||
super().construct_network()
|
||||
if self.enhence_model_cfg:
|
||||
self.enhence_model = MODELS.build(self.enhence_model_cfg, logger=self.logger).eval().requires_grad_(False)
|
||||
self.enhence_sampler_cfg = {key.lower(): value for key, value in self.enhence_sampler_cfg.items()}
|
||||
else:
|
||||
self.enhence_model = None
|
||||
self.enhence_sampler_cfg = None
|
||||
|
||||
def forward_sample(self,
|
||||
edit_image=[],
|
||||
edit_mask=[],
|
||||
noise=None,
|
||||
cond_mask=[],
|
||||
x_shapes=[],
|
||||
prompt=[],
|
||||
n_prompt=[],
|
||||
sampler='ddim',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=4.5,
|
||||
guide_rescale=0.5,
|
||||
discretization='trailing',
|
||||
**kwargs
|
||||
):
|
||||
'''
|
||||
Args:
|
||||
edit_image: list of list of edit_image
|
||||
edit_image_mask: list of list of edit_image_mask
|
||||
image: target image
|
||||
image_mask: target image mask
|
||||
prompt: list of list of text
|
||||
n_prompt: list of list of text
|
||||
sampler:
|
||||
sample_steps:
|
||||
seed:
|
||||
guide_scale:
|
||||
guide_rescale:
|
||||
discretization:
|
||||
log_num:
|
||||
**kwargs:
|
||||
|
||||
Returns:
|
||||
|
||||
'''
|
||||
|
||||
# prepare data
|
||||
context, null_context = {}, {}
|
||||
context['x_shapes'] = null_context['x_shapes'] = x_shapes
|
||||
# process image mask
|
||||
|
||||
context['x_mask'] = null_context['x_mask'] = cond_mask
|
||||
# process text
|
||||
# with torch.autocast(device_type="cuda", enabled=True, dtype=torch.bfloat16):
|
||||
|
||||
cont, cont_mask = getattr(self.cond_stage_model, 'encode_list')(prompt, return_mask=True)
|
||||
cont, cont_mask = self.cond_stage_embeddings(prompt, 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
|
||||
|
||||
|
||||
null_context['edit'] = context['edit'] = edit_image
|
||||
null_context['edit_mask'] = context['edit_mask'] = edit_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(solver=sampler,
|
||||
noise=noise,
|
||||
model=model,
|
||||
model_kwargs=[{
|
||||
'cond': context,
|
||||
'mask': cont_mask,
|
||||
'text_position_embeddings': self.text_position_embeddings.pos if hasattr(
|
||||
self.text_position_embeddings, 'pos') else None
|
||||
}, {
|
||||
'cond': null_context,
|
||||
'mask': null_cont_mask,
|
||||
'text_position_embeddings': self.text_position_embeddings.pos if hasattr(
|
||||
self.text_position_embeddings, 'pos') else None
|
||||
}] if guide_scale is not None and guide_scale > 1 else {
|
||||
'cond': context,
|
||||
'mask': cont_mask,
|
||||
'text_position_embeddings': self.text_position_embeddings.pos if hasattr(
|
||||
self.text_position_embeddings, 'pos') else None
|
||||
},
|
||||
cat_uc=False,
|
||||
steps=sample_steps,
|
||||
guide_scale=guide_scale,
|
||||
guide_rescale=guide_rescale,
|
||||
discretization=discretization,
|
||||
show_progress=True,
|
||||
seed=seed,
|
||||
condition_fn=None,
|
||||
clamp=None,
|
||||
percentile=None,
|
||||
t_max=None,
|
||||
t_min=None,
|
||||
discard_penultimate_step=None,
|
||||
return_intermediate=None,
|
||||
**kwargs)
|
||||
|
||||
samples = unpack_tensor_into_imagelist(samples, x_shapes)
|
||||
x_samples = self.decode_first_stage(samples)
|
||||
return x_samples
|
||||
|
||||
def upscale_resize(self, image, interpolation=T.InterpolationMode.BILINEAR):
|
||||
_, c, H, W = image.shape
|
||||
scale = max(1.0, math.sqrt(4096 / ((H / 16) * (W / 16))))
|
||||
rH = int(H * scale) // 16 * 16 # ensure divisible by self.d
|
||||
rW = int(W * scale) // 16 * 16
|
||||
image = T.Resize((rH, rW), interpolation=interpolation, antialias=True)(image)
|
||||
return image
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_test(self,
|
||||
edit_image=[],
|
||||
edit_image_mask=[],
|
||||
image=None,
|
||||
image_mask=None,
|
||||
prompt=[],
|
||||
n_prompt=[],
|
||||
sampler='ddim',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=4.5,
|
||||
guide_rescale=0.5,
|
||||
discretization='trailing',
|
||||
enhance_scale=0.99,
|
||||
log_num=-1,
|
||||
**kwargs):
|
||||
assert check_list_of_list(prompt) and check_list_of_list(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)
|
||||
prompt, n_prompt, image, image_mask, edit_image, edit_image_mask = self.limit_batch_data(
|
||||
[prompt, n_prompt, image, image_mask, edit_image, edit_image_mask], log_num)
|
||||
|
||||
prompt = [[pp] if isinstance(pp, str) else pp for pp in prompt]
|
||||
|
||||
g = torch.Generator(device=we.device_id)
|
||||
seed = seed if seed >= 0 else random.randint(0, 2 ** 32 - 1)
|
||||
g.manual_seed(seed)
|
||||
n_prompt = copy.deepcopy(prompt)
|
||||
# only modify the last prompt to be zero
|
||||
for nn_p_id, nn_p in enumerate(n_prompt):
|
||||
if isinstance(nn_p, str):
|
||||
n_prompt[nn_p_id] = [""]
|
||||
elif isinstance(nn_p, list):
|
||||
n_prompt[nn_p_id][-1] = ""
|
||||
else:
|
||||
raise NotImplementedError
|
||||
# process image
|
||||
image = to_device(image)
|
||||
x = self.encode_first_stage(image, **kwargs)
|
||||
noise = [torch.empty(*i.shape, device=we.device_id).normal_(generator=g) for i in x]
|
||||
noise, x_shapes = pack_imagelist_into_tensor(noise)
|
||||
image_mask = to_device(image_mask, strict=False)
|
||||
cond_mask = [self.interpolate_func(i) for i in image_mask] if image_mask is not None else [None] * len(image)
|
||||
|
||||
# processe edit image & edit image mask
|
||||
edit_image = [to_device(i, strict=False) for i in edit_image]
|
||||
edit_image_mask = [to_device(i, strict=False) for i in edit_image_mask]
|
||||
e_img, e_mask = [], []
|
||||
for u, m in zip(edit_image, edit_image_mask):
|
||||
if u is None:
|
||||
continue
|
||||
if m is None:
|
||||
m = [None] * len(u)
|
||||
e_img.append(self.encode_first_stage(u, **kwargs))
|
||||
e_mask.append([self.interpolate_func(i) for i in m])
|
||||
|
||||
x_samples = self.forward_sample(
|
||||
edit_image=e_img,
|
||||
edit_mask=e_mask,
|
||||
noise=noise,
|
||||
cond_mask=cond_mask,
|
||||
x_shapes=x_shapes,
|
||||
prompt=prompt,
|
||||
n_prompt=n_prompt,
|
||||
sampler=sampler,
|
||||
sample_steps=sample_steps,
|
||||
seed=seed,
|
||||
guide_scale=guide_scale,
|
||||
guide_rescale=guide_rescale,
|
||||
discretization='trailing',
|
||||
**kwargs)
|
||||
|
||||
if self.enhence_model and enhance_scale > 0:
|
||||
x_samples = [self.upscale_resize(x) for x in x_samples]
|
||||
x_start = self.enhence_model.encode_first_stage(x_samples, **kwargs)
|
||||
noise = []
|
||||
for i, x in enumerate(x_start):
|
||||
noise_ = self.enhence_model.noise_sample(1, x_samples[i].shape[2], x_samples[i].shape[3], seed)
|
||||
noise.append(noise_)
|
||||
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
|
||||
x_samples = self.enhence_model.forward_sample(noise = noise,
|
||||
x = x_start,
|
||||
reverse_scale = enhance_scale,
|
||||
prompt =[kwargs.pop("enhance_prompt", "") for _ in noise],
|
||||
**self.enhence_sampler_cfg)
|
||||
outputs = list()
|
||||
for i in range(len(prompt)):
|
||||
rec_img = torch.clamp((x_samples[i].float() + 1.0) / 2.0 + self.decoder_bias / 255, min=0.0, max=1.0)
|
||||
rec_img = rec_img.squeeze(0)
|
||||
edit_imgs, edit_img_masks = [], []
|
||||
if edit_image is not None and edit_image[i] is not None:
|
||||
if edit_image_mask[i] is None:
|
||||
edit_image_mask[i] = [None] * len(edit_image[i])
|
||||
for edit_img, edit_mask in zip(edit_image[i], edit_image_mask[i]):
|
||||
edit_img = torch.clamp((edit_img + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
edit_imgs.append(edit_img.squeeze(0))
|
||||
if edit_mask is None:
|
||||
edit_mask = torch.ones_like(edit_img[[0], :, :])
|
||||
edit_img_masks.append(edit_mask)
|
||||
one_tup = {
|
||||
'reconstruct_image': rec_img,
|
||||
'instruction': prompt[i],
|
||||
'edit_image': edit_imgs if len(edit_imgs) > 0 else None,
|
||||
'edit_mask': edit_img_masks if len(edit_imgs) > 0 else None
|
||||
}
|
||||
if image is not None:
|
||||
if image_mask is None:
|
||||
image_mask = [None] * len(image)
|
||||
ori_img = torch.clamp((image[i] + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
one_tup['target_image'] = ori_img.squeeze(0)
|
||||
one_tup['target_mask'] = image_mask[i] if image_mask[i] is not None else torch.ones_like(
|
||||
ori_img[[0], :, :])
|
||||
outputs.append(one_tup)
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODEL',
|
||||
__class__.__name__,
|
||||
LatentDiffusionACERefiner.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,225 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import random
|
||||
import torch
|
||||
from typing import Tuple
|
||||
|
||||
from scepter.modules.model.network.ldm import LatentDiffusion
|
||||
from scepter.modules.model.registry import MODELS
|
||||
from scepter.modules.model.utils.basic_utils import default
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.model.backbone.cogvideox.utils import get_3d_rotary_pos_embed, get_resize_crop_region_for_grid
|
||||
|
||||
|
||||
@MODELS.register_class()
|
||||
class LatentDiffusionCogVideoX(LatentDiffusion):
|
||||
para_dict = LatentDiffusion.para_dict
|
||||
|
||||
def init_params(self):
|
||||
super().init_params()
|
||||
self.latent_channels = self.model_config.get('LATENT_CHANNELS', self.model_config.IN_CHANNELS)
|
||||
self.scale_factor_spatial = self.cfg.get('SCALE_FACTOR_SPATIAL', 8)
|
||||
self.scale_factor_temporal = self.cfg.get('SCALE_FACTOR_TEMPORAL', 4)
|
||||
self.scaling_factor_image = self.cfg.get('SCALING_FACTOR_IMAGE', 0.7)
|
||||
self.use_rotary_positional_embeddings = self.model_config.get('USE_ROTARY_POSITIONAL_EMBEDDINGS', False)
|
||||
self.attention_head_dim = self.model_config.get('ATTENTION_HEAD_DIM', 64)
|
||||
self.patch_size = self.model_config.get('PATCH_SIZE', 2)
|
||||
self.sample_height = self.first_stage_config.get('SAMPLE_HEIGHT', 480)
|
||||
self.sample_width = self.first_stage_config.get('SAMPLE_WIDTH', 720)
|
||||
self.noised_image_dropout = self.cfg.get('NOISED_IMAGE_DROPOUT', 0.05)
|
||||
|
||||
def construct_network(self):
|
||||
super().construct_network()
|
||||
self.model = self.model.to(getattr(torch, self.model_config.DTYPE))
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_first_stage(self, x, **kwargs):
|
||||
if isinstance(x, list):
|
||||
x = torch.stack(x, dim=0) # [B, C, F, H, W]
|
||||
latents = self.scaling_factor_image * self.first_stage_model.encode(x).sample()
|
||||
return latents
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, latents):
|
||||
latents = latents.permute(0, 2, 1, 3, 4) # [batch_size, num_channels, num_frames, height, width]
|
||||
latents = 1 / self.scaling_factor_image * latents
|
||||
frames = self.first_stage_model.decode(latents)
|
||||
return frames
|
||||
|
||||
def get_image_latent(self, image, video, noise):
|
||||
latent = torch.zeros_like(noise)
|
||||
if isinstance(image, list):
|
||||
image = torch.stack(image, dim=0) # [B, C, F, H, W]
|
||||
if len(image.shape) == 4: # [B, C, H, W]
|
||||
image = image.unsqueeze(2) # [B, C, F, H, W]
|
||||
image_latent = self.encode_first_stage(image) # [B, C, F, H, W]
|
||||
image_latent = image_latent.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
|
||||
latent[:, :1, :, :, :] = image_latent
|
||||
return latent, image
|
||||
|
||||
def noise_sample(self, batch_size, num_frames, height, width, generator, dtype=torch.bfloat16):
|
||||
shape = (batch_size,
|
||||
(num_frames - 1) // self.scale_factor_temporal + 1,
|
||||
self.latent_channels,
|
||||
height // self.scale_factor_spatial,
|
||||
width // self.scale_factor_spatial
|
||||
)
|
||||
noise = torch.randn(shape, generator=generator, dtype=dtype, device='cpu').to(we.device_id)
|
||||
return noise
|
||||
|
||||
def _prepare_rotary_positional_embeddings(
|
||||
self,
|
||||
height: int,
|
||||
width: int,
|
||||
num_frames: int,
|
||||
device: torch.device,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
grid_height = height // (self.scale_factor_spatial * self.patch_size)
|
||||
grid_width = width // (self.scale_factor_spatial * self.patch_size)
|
||||
base_size_width = self.sample_width // (self.scale_factor_spatial * self.patch_size)
|
||||
base_size_height = self.sample_height // (self.scale_factor_spatial * self.patch_size)
|
||||
|
||||
grid_crops_coords = get_resize_crop_region_for_grid(
|
||||
(grid_height, grid_width), base_size_width, base_size_height
|
||||
)
|
||||
freqs_cos, freqs_sin = get_3d_rotary_pos_embed(
|
||||
embed_dim=self.attention_head_dim,
|
||||
crops_coords=grid_crops_coords,
|
||||
grid_size=(grid_height, grid_width),
|
||||
temporal_size=num_frames,
|
||||
)
|
||||
|
||||
freqs_cos = freqs_cos.to(device=device)
|
||||
freqs_sin = freqs_sin.to(device=device)
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
|
||||
def forward_train(self, video=None, video_latent=None, image=None, noise=None, prompt=None, image_size=None, **kwargs):
|
||||
# video: [B, C, F, H, W]
|
||||
if image_size is None: image_size = [480, 720]
|
||||
if video_latent is not None:
|
||||
x_start = torch.stack(video_latent)
|
||||
else:
|
||||
x_start = self.encode_first_stage(video, **kwargs)
|
||||
x_start = x_start.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
|
||||
t = torch.randint(low=0, high=self.num_timesteps, size=(len(video),), device=we.device_id)
|
||||
|
||||
if prompt and self.cond_stage_model:
|
||||
with torch.autocast(device_type='cuda', enabled=True, dtype=torch.bfloat16):
|
||||
cont = getattr(self.cond_stage_model, 'encode')(prompt, return_mask=False, use_mask=False)
|
||||
|
||||
if noise is None:
|
||||
noise = torch.randn_like(x_start)
|
||||
|
||||
if image is not None:
|
||||
if random.random() < self.noised_image_dropout:
|
||||
image_latent = torch.zeros_like(noise)
|
||||
else:
|
||||
image_latent, _ = self.get_image_latent(image, video, noise)
|
||||
else:
|
||||
image_latent = None
|
||||
|
||||
height, width = image_size
|
||||
image_rotary_emb = (
|
||||
self._prepare_rotary_positional_embeddings(height, width, noise.size(1), we.device_id)
|
||||
if self.use_rotary_positional_embeddings
|
||||
else None
|
||||
)
|
||||
|
||||
loss = self.diffusion.loss(x_0=x_start,
|
||||
t=t,
|
||||
model=self.model,
|
||||
model_kwargs={"cond": cont,
|
||||
'image_latent': image_latent,
|
||||
'image_rotary_emb': image_rotary_emb},
|
||||
noise=noise,
|
||||
**kwargs)
|
||||
loss = loss.mean()
|
||||
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
|
||||
return ret
|
||||
|
||||
@torch.no_grad()
|
||||
@torch.autocast('cuda', dtype=torch.bfloat16)
|
||||
def forward_test(self,
|
||||
video=None,
|
||||
image=None,
|
||||
prompt=None,
|
||||
n_prompt=None,
|
||||
sampler='ddim',
|
||||
sample_steps=50,
|
||||
seed=42,
|
||||
guide_scale=6.0,
|
||||
guide_rescale=0.0,
|
||||
num_frames=49,
|
||||
image_size=None,
|
||||
show_process=False,
|
||||
**kwargs):
|
||||
if image_size is None:
|
||||
image_size = [480, 720]
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||
generator = torch.Generator().manual_seed(seed)
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
num_samples = len(prompt)
|
||||
n_prompt = default(n_prompt, [self.default_n_prompt] * len(prompt))
|
||||
|
||||
if prompt and self.cond_stage_model:
|
||||
with torch.autocast(device_type='cuda', enabled=True, dtype=torch.bfloat16):
|
||||
cont = getattr(self.cond_stage_model, 'encode')(prompt, return_mask=False, use_mask=False)
|
||||
null_cont = getattr(self.cond_stage_model, 'encode')(n_prompt, return_mask=False, use_mask=False)
|
||||
|
||||
height, width = image_size
|
||||
noise = self.noise_sample(num_samples, num_frames, height, width, generator)
|
||||
image_rotary_emb = (
|
||||
self._prepare_rotary_positional_embeddings(height, width, noise.size(1), we.device_id)
|
||||
if self.use_rotary_positional_embeddings
|
||||
else None
|
||||
)
|
||||
image_latent, image = self.get_image_latent(image, video, noise) if image is not None else (None, None)
|
||||
|
||||
samples = self.diffusion.sample(noise=noise,
|
||||
sampler=sampler,
|
||||
model=self.model,
|
||||
model_kwargs=[{
|
||||
'cond': cont,
|
||||
'image_latent': image_latent,
|
||||
'image_rotary_emb': image_rotary_emb,
|
||||
}, {
|
||||
'cond': null_cont,
|
||||
'image_latent': image_latent,
|
||||
'image_rotary_emb': image_rotary_emb,
|
||||
}],
|
||||
steps=sample_steps,
|
||||
show_progress=True,
|
||||
use_dynamic_cfg=True,
|
||||
guide_scale=guide_scale,
|
||||
guide_rescale=guide_rescale,
|
||||
return_intermediate=None,
|
||||
**kwargs).float()
|
||||
|
||||
x_frames = self.decode_first_stage(samples).float()
|
||||
|
||||
outputs = []
|
||||
for batch_idx in range(num_samples):
|
||||
rec_video = torch.clamp(x_frames[batch_idx] / 2 + 0.5, min=0.0, max=1.0)
|
||||
one_tup = {
|
||||
'reconstruct_video': rec_video.squeeze(0).float(),
|
||||
'instruction': prompt[batch_idx]
|
||||
}
|
||||
if image is not None:
|
||||
ori_image = torch.clamp(image[batch_idx] / 2 + 0.5, min=0.0, max=1.0)
|
||||
one_tup['edit_image'] = ori_image
|
||||
if video is not None:
|
||||
ori_video = torch.clamp(video[batch_idx] / 2 + 0.5, min=0.0, max=1.0)
|
||||
one_tup['target_video'] = ori_video.squeeze(0)
|
||||
outputs.append(one_tup)
|
||||
return outputs
|
||||
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODEL',
|
||||
__class__.__name__,
|
||||
LatentDiffusionCogVideoX.para_dict,
|
||||
set_name=True)
|
||||
@@ -4,15 +4,19 @@ import copy
|
||||
import math
|
||||
import numbers
|
||||
import random
|
||||
from contextlib import nullcontext
|
||||
|
||||
import torch
|
||||
from scepter.modules.model.network.ldm import LatentDiffusion
|
||||
from scepter.modules.model.registry import MODELS, BACKBONES, LOSSES, TOKENIZERS, EMBEDDERS, DIFFUSIONS
|
||||
from scepter.modules.model.utils.basic_utils import disabled_train
|
||||
from scepter.modules.model.utils.basic_utils import disabled_train, check_list_of_list, to_device, \
|
||||
pack_imagelist_into_tensor, unpack_tensor_into_imagelist, limit_batch_data
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.model.utils.basic_utils import count_params
|
||||
|
||||
|
||||
|
||||
@MODELS.register_class()
|
||||
class LatentDiffusionFlux(LatentDiffusion):
|
||||
para_dict = LatentDiffusion.para_dict
|
||||
@@ -137,7 +141,7 @@ class LatentDiffusionFlux(LatentDiffusion):
|
||||
def forward_test(self,
|
||||
image=None,
|
||||
prompt=None,
|
||||
sampler='flow_eluer',
|
||||
sampler='flow_euler',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=4.5,
|
||||
@@ -218,3 +222,164 @@ class LatentDiffusionFlux(LatentDiffusion):
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, z):
|
||||
return self.first_stage_model.decode(z)
|
||||
|
||||
@MODELS.register_class()
|
||||
class LatentDiffusionFluxMR(LatentDiffusionFlux):
|
||||
para_dict = {
|
||||
}
|
||||
para_dict.update(LatentDiffusion.para_dict)
|
||||
def forward_train(self,
|
||||
image=None,
|
||||
noise=None,
|
||||
prompt=[],
|
||||
**kwargs):
|
||||
if check_list_of_list(prompt):
|
||||
prompt = [pp[0] for pp in prompt]
|
||||
assert self.cond_stage_model is not None
|
||||
gc_seg = kwargs.pop("gc_seg", [])
|
||||
gc_seg = int(gc_seg[0]) if len(gc_seg) > 0 else 0
|
||||
context = getattr(self.cond_stage_model, 'encode')(prompt)
|
||||
|
||||
image = to_device(image)
|
||||
x_start = self.encode_first_stage(image, **kwargs)
|
||||
loss_mask, _ = pack_imagelist_into_tensor(tuple(torch.ones_like(ix, dtype=torch.bool, device=ix.device) for ix in x_start))
|
||||
x_start, x_shapes = pack_imagelist_into_tensor(x_start)
|
||||
context['x_shapes'] = x_shapes
|
||||
guide_scale = self.guide_scale
|
||||
if guide_scale is not None:
|
||||
guide_scale = torch.full((x_start.shape[0],), guide_scale, device=x_start.device, dtype=x_start.dtype)
|
||||
else:
|
||||
guide_scale = None
|
||||
loss = self.diffusion.loss(x_0=x_start,
|
||||
model=self.model,
|
||||
model_kwargs={"cond": context,
|
||||
"gc_seg": gc_seg,
|
||||
"guidance": guide_scale},
|
||||
noise=None,
|
||||
reduction='none',
|
||||
**kwargs)
|
||||
loss = loss[loss_mask].mean()
|
||||
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
|
||||
return ret
|
||||
|
||||
@torch.no_grad()
|
||||
def forward_sample(self,
|
||||
noise = None,
|
||||
prompt=None,
|
||||
sampler='flow_euler',
|
||||
sample_steps=20,
|
||||
guide_scale=3.5,
|
||||
show_process=True,
|
||||
x = None,
|
||||
reverse_scale = 0.,
|
||||
**kwargs
|
||||
):
|
||||
noise, x_shapes = pack_imagelist_into_tensor(noise)
|
||||
if x is not None:
|
||||
x, _ = pack_imagelist_into_tensor(x)
|
||||
context = getattr(self.cond_stage_model, 'encode')(prompt)
|
||||
context["x_shapes"] = x_shapes
|
||||
guide_scale = guide_scale or self.guide_scale
|
||||
if guide_scale is not None:
|
||||
guide_scale = torch.full((noise.shape[0],), guide_scale, device=noise.device, dtype=noise.dtype)
|
||||
else:
|
||||
guide_scale = None
|
||||
# UNet use input n_prompt
|
||||
model = self.model_ema if self.use_ema and self.eval_ema else self.model
|
||||
embedding_context = model.no_sync if isinstance(model, torch.distributed.fsdp.FullyShardedDataParallel) \
|
||||
else nullcontext
|
||||
with embedding_context():
|
||||
x_samples = self.diffusion.sample(
|
||||
noise=noise,
|
||||
sampler=sampler,
|
||||
model=self.model,
|
||||
model_kwargs={"cond": context, "guidance": guide_scale, "gc_seg": -1},
|
||||
steps=sample_steps,
|
||||
show_progress=True,
|
||||
guide_scale=guide_scale,
|
||||
return_intermediate=None,
|
||||
reverse_scale = reverse_scale,
|
||||
x = x,
|
||||
**kwargs).float()
|
||||
x_samples = unpack_tensor_into_imagelist(x_samples, x_shapes)
|
||||
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
|
||||
x_samples = self.decode_first_stage(x_samples)
|
||||
return x_samples
|
||||
@torch.no_grad()
|
||||
def forward_test(self,
|
||||
image=None,
|
||||
prompt=[],
|
||||
sampler='flow_euler',
|
||||
sample_steps=20,
|
||||
seed=2023,
|
||||
guide_scale=3.5,
|
||||
guide_rescale=0.0,
|
||||
show_process=True,
|
||||
log_num = -1,
|
||||
**kwargs):
|
||||
|
||||
if check_list_of_list(prompt):
|
||||
prompt = [pp[0] for pp in prompt]
|
||||
assert self.cond_stage_model is not None
|
||||
# gc_seg is unused
|
||||
prompt, image = limit_batch_data([prompt, image], log_num)
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||
|
||||
if 'index' in kwargs:
|
||||
kwargs.pop('index')
|
||||
if image is not None:
|
||||
noise = [self.noise_sample(1, ix.shape[1], ix.shape[2], seed) for ix in image]
|
||||
else:
|
||||
image_size = None
|
||||
if 'meta' in kwargs:
|
||||
meta = kwargs.pop('meta')
|
||||
if 'image_size' in meta:
|
||||
h = int(meta['image_size'][0][0])
|
||||
w = int(meta['image_size'][1][0])
|
||||
image_size = [h, w]
|
||||
if 'image_size' in kwargs:
|
||||
image_size = kwargs.pop('image_size')
|
||||
if isinstance(image_size, numbers.Number):
|
||||
image_size = [image_size, image_size]
|
||||
if image_size is None:
|
||||
image_size = [1024, 1024]
|
||||
height, width = image_size
|
||||
noise = [self.noise_sample(1, height, width, seed) for _ in prompt]
|
||||
|
||||
x_samples = self.forward_sample(
|
||||
prompt=prompt,
|
||||
sampler=sampler,
|
||||
sample_steps=sample_steps,
|
||||
guide_scale=guide_scale,
|
||||
show_process=show_process,
|
||||
noise=noise,
|
||||
)
|
||||
|
||||
|
||||
outputs = list()
|
||||
for i in range(len(prompt)):
|
||||
rec_img = torch.clamp((x_samples[i].float() + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
rec_img = rec_img.squeeze(0)
|
||||
one_tup = {'prompt': prompt[i], 'n_prompt': '', 'image': rec_img}
|
||||
outputs.append(one_tup)
|
||||
return outputs
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('MODEL',
|
||||
__class__.__name__,
|
||||
LatentDiffusionFlux.para_dict,
|
||||
set_name=True)
|
||||
@torch.no_grad()
|
||||
def encode_first_stage(self, x, **kwargs):
|
||||
def run_one_image(u):
|
||||
zu = self.first_stage_model.encode(u)
|
||||
if isinstance(zu, (tuple, list)):
|
||||
zu = zu[0]
|
||||
return zu
|
||||
|
||||
z = [run_one_image(u.unsqueeze(0) if u.dim == 3 else u) for u in x]
|
||||
return z
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_first_stage(self, z):
|
||||
return [self.first_stage_model.decode(zu) for zu in z]
|
||||
@@ -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):
|
||||
|
||||
@@ -15,8 +15,8 @@ def build_model(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)):
|
||||
raise TypeError('Pretrain parameter must be a string')
|
||||
if pretrain_cfg is not None and not isinstance(pretrain_cfg, (str, list)):
|
||||
raise TypeError('Pretrain parameter must be a string or list')
|
||||
else:
|
||||
pretrain_cfg = None
|
||||
|
||||
|
||||
@@ -1,13 +1,14 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import open_clip
|
||||
from transformers import CLIPTokenizer as transformer_clip_tokenizer
|
||||
|
||||
from scepter.modules.model.registry import TOKENIZERS
|
||||
from scepter.modules.model.tokenizer import BaseTokenizer
|
||||
from scepter.modules.model.tokenizer.tokenizer_component import (
|
||||
basic_clean, canonicalize, heavy_clean, whitespace_clean)
|
||||
basic_clean, canonicalize, whitespace_clean)
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from transformers import CLIPTokenizer as transformer_clip_tokenizer
|
||||
|
||||
|
||||
@TOKENIZERS.register_class()
|
||||
@@ -31,13 +32,16 @@ class HuggingfaceTokenizer(BaseTokenizer):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.pretrained_path = cfg.get('PRETRAINED_PATH', 'xlm-roberta-large')
|
||||
self.length = cfg.get('LENGTH', 77)
|
||||
self.clean = cfg.get('CLEAN', True)
|
||||
self.clean = cfg.get('CLEAN', 'whitespace')
|
||||
assert self.clean in (None, 'whitespace', 'lower', 'canonicalize')
|
||||
|
||||
# init tokenizer
|
||||
from transformers import AutoTokenizer
|
||||
with FS.get_dir_to_local_dir(self.pretrained_path) as local_path:
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(local_path)
|
||||
self.vocab_size = len(self.tokenizer)
|
||||
|
||||
self.vocab_size = len(
|
||||
self.tokenizer) # self.vocab_size = self.tokenizer.vocab_size
|
||||
|
||||
# special tokens
|
||||
self.comma_token = self.tokenizer(',')['input_ids'][
|
||||
@@ -51,6 +55,7 @@ class HuggingfaceTokenizer(BaseTokenizer):
|
||||
|
||||
def __call__(self, sequence, **kwargs):
|
||||
# arguments
|
||||
return_mask = kwargs.pop('return_mask', False)
|
||||
_kwargs = {'return_tensors': 'pt'}
|
||||
if self.length is not None:
|
||||
_kwargs.update({
|
||||
@@ -64,9 +69,23 @@ class HuggingfaceTokenizer(BaseTokenizer):
|
||||
if isinstance(sequence, str):
|
||||
sequence = [sequence]
|
||||
if self.clean:
|
||||
sequence = [whitespace_clean(basic_clean(u)) for u in sequence]
|
||||
sequence = [self._clean(u) for u in sequence]
|
||||
tokens = self.tokenizer(sequence, **_kwargs)
|
||||
return tokens.input_ids
|
||||
|
||||
# output
|
||||
if return_mask:
|
||||
return tokens.input_ids, tokens.attention_mask
|
||||
else:
|
||||
return tokens.input_ids
|
||||
|
||||
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():
|
||||
|
||||
@@ -2,6 +2,11 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from inspect import isfunction
|
||||
|
||||
import torch
|
||||
from torch.nn.utils.rnn import pad_sequence
|
||||
|
||||
from scepter.modules.utils.distribute import we
|
||||
|
||||
|
||||
def exists(x):
|
||||
return x is not None
|
||||
@@ -45,3 +50,77 @@ def expand_dims_like(x, y):
|
||||
while x.dim() != y.dim():
|
||||
x = x.unsqueeze(-1)
|
||||
return x
|
||||
|
||||
|
||||
def unpack_tensor_into_imagelist(image_tensor, shapes):
|
||||
image_list = []
|
||||
for img, shape in zip(image_tensor, shapes):
|
||||
h, w = shape[0], shape[1]
|
||||
image_list.append(img[:, :h * w].view(1, -1, h, w))
|
||||
|
||||
return image_list
|
||||
|
||||
|
||||
def find_example(tensor_list, image_list):
|
||||
for i in tensor_list:
|
||||
if isinstance(i, torch.Tensor):
|
||||
return torch.zeros_like(i)
|
||||
for i in image_list:
|
||||
if isinstance(i, torch.Tensor):
|
||||
_, c, h, w = i.size()
|
||||
return torch.zeros_like(i.view(c, h * w).transpose(1, 0))
|
||||
return None
|
||||
|
||||
|
||||
def pack_imagelist_into_tensor_v2(image_list):
|
||||
# allow None
|
||||
example = None
|
||||
image_tensor, shapes = [], []
|
||||
for img in image_list:
|
||||
if img is None:
|
||||
example = find_example(image_tensor,
|
||||
image_list) if example is None else example
|
||||
image_tensor.append(example)
|
||||
shapes.append(None)
|
||||
continue
|
||||
_, c, h, w = img.size()
|
||||
image_tensor.append(img.view(c, h * w).transpose(1, 0)) # h*w, c
|
||||
shapes.append((h, w))
|
||||
|
||||
image_tensor = pad_sequence(image_tensor,
|
||||
batch_first=True).permute(0, 2, 1) # b, c, l
|
||||
return image_tensor, shapes
|
||||
|
||||
|
||||
def to_device(inputs, strict=True):
|
||||
if inputs is None:
|
||||
return None
|
||||
if strict:
|
||||
assert all(isinstance(i, torch.Tensor) for i in inputs)
|
||||
return [i.to(we.device_id) if i is not None else None for i in inputs]
|
||||
|
||||
|
||||
def check_list_of_list(ll):
|
||||
return isinstance(ll, list) and all(isinstance(i, list) for i in ll)
|
||||
|
||||
|
||||
def pack_imagelist_into_tensor(image_list):
|
||||
image_tensor, shapes = [], []
|
||||
for img in image_list:
|
||||
_, c, h, w = img.size()
|
||||
image_tensor.append(img.view(c, h * w).transpose(1, 0)) # h*w, c
|
||||
shapes.append((h, w))
|
||||
|
||||
image_tensor = pad_sequence(image_tensor, batch_first=True).permute(0, 2, 1) # b, c, l
|
||||
return image_tensor, shapes
|
||||
|
||||
def limit_batch_data(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
|
||||
@@ -4,3 +4,5 @@ from scepter.modules.solver import hooks
|
||||
from scepter.modules.solver.base_solver import BaseSolver
|
||||
from scepter.modules.solver.diffusion_solver import LatentDiffusionSolver
|
||||
from scepter.modules.solver.train_val_solver import TrainValSolver
|
||||
from scepter.modules.solver.ace_solver import ACESolver
|
||||
from scepter.modules.solver.diffusion_video_solver import LatentDiffusionVideoSolver
|
||||
@@ -0,0 +1,146 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import numpy as np
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from scepter.modules.utils.data import transfer_data_to_cuda
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.probe import ProbeData
|
||||
|
||||
from .diffusion_solver import LatentDiffusionSolver
|
||||
from .registry import SOLVERS
|
||||
|
||||
|
||||
@SOLVERS.register_class()
|
||||
class ACESolver(LatentDiffusionSolver):
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.log_train_num = cfg.get('LOG_TRAIN_NUM', -1)
|
||||
|
||||
def save_results(self, results):
|
||||
log_data, log_label = [], []
|
||||
for result in results:
|
||||
ret_images, ret_labels = [], []
|
||||
edit_image = result.get('edit_image', None)
|
||||
edit_mask = result.get('edit_mask', None)
|
||||
if edit_image is not None:
|
||||
for i, edit_img in enumerate(result['edit_image']):
|
||||
if edit_img is None:
|
||||
continue
|
||||
ret_images.append(
|
||||
(edit_img.permute(1, 2, 0).cpu().numpy() * 255).astype(
|
||||
np.uint8))
|
||||
ret_labels.append(f'edit_image{i}; ')
|
||||
if edit_mask is not None:
|
||||
ret_images.append(
|
||||
(edit_mask[i].permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
ret_labels.append(f'edit_mask{i}; ')
|
||||
|
||||
target_image = result.get('target_image', None)
|
||||
target_mask = result.get('target_mask', None)
|
||||
if target_image is not None:
|
||||
ret_images.append(
|
||||
(target_image.permute(1, 2, 0).cpu().numpy() * 255).astype(
|
||||
np.uint8))
|
||||
ret_labels.append('target_image; ')
|
||||
if target_mask is not None:
|
||||
ret_images.append(
|
||||
(target_mask.permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
ret_labels.append('target_mask; ')
|
||||
|
||||
reconstruct_image = result.get('reconstruct_image', None)
|
||||
if reconstruct_image is not None:
|
||||
ret_images.append(
|
||||
(reconstruct_image.permute(1, 2, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
ret_labels.append(f"{result['instruction']}")
|
||||
log_data.append(ret_images)
|
||||
log_label.append(ret_labels)
|
||||
return log_data, log_label
|
||||
|
||||
@torch.no_grad()
|
||||
def run_eval(self):
|
||||
self.eval_mode()
|
||||
self.before_all_iter(self.hooks_dict[self._mode])
|
||||
all_results = []
|
||||
for batch_idx, batch_data in tqdm(
|
||||
enumerate(self.datas[self._mode].dataloader)):
|
||||
self.before_iter(self.hooks_dict[self._mode])
|
||||
if self.sample_args:
|
||||
batch_data.update(self.sample_args.get_lowercase_dict())
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=self.use_amp,
|
||||
dtype=self.dtype):
|
||||
results = self.run_step_eval(transfer_data_to_cuda(batch_data),
|
||||
batch_idx,
|
||||
step=self.total_iter,
|
||||
rank=we.rank)
|
||||
all_results.extend(results)
|
||||
self.after_iter(self.hooks_dict[self._mode])
|
||||
log_data, log_label = self.save_results(all_results)
|
||||
self.register_probe({'eval_label': log_label})
|
||||
self.register_probe({
|
||||
'eval_image':
|
||||
ProbeData(log_data,
|
||||
is_image=True,
|
||||
build_html=True,
|
||||
build_label=log_label)
|
||||
})
|
||||
self.after_all_iter(self.hooks_dict[self._mode])
|
||||
|
||||
@torch.no_grad()
|
||||
def run_test(self):
|
||||
self.test_mode()
|
||||
self.before_all_iter(self.hooks_dict[self._mode])
|
||||
all_results = []
|
||||
for batch_idx, batch_data in tqdm(
|
||||
enumerate(self.datas[self._mode].dataloader)):
|
||||
self.before_iter(self.hooks_dict[self._mode])
|
||||
if self.sample_args:
|
||||
batch_data.update(self.sample_args.get_lowercase_dict())
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=self.use_amp,
|
||||
dtype=self.dtype):
|
||||
results = self.run_step_eval(transfer_data_to_cuda(batch_data),
|
||||
batch_idx,
|
||||
step=self.total_iter,
|
||||
rank=we.rank)
|
||||
all_results.extend(results)
|
||||
self.after_iter(self.hooks_dict[self._mode])
|
||||
log_data, log_label = self.save_results(all_results)
|
||||
self.register_probe({'test_label': log_label})
|
||||
self.register_probe({
|
||||
'test_image':
|
||||
ProbeData(log_data,
|
||||
is_image=True,
|
||||
build_html=True,
|
||||
build_label=log_label)
|
||||
})
|
||||
|
||||
self.after_all_iter(self.hooks_dict[self._mode])
|
||||
|
||||
@property
|
||||
def probe_data(self):
|
||||
if not we.debug and self.mode == 'train':
|
||||
batch_data = transfer_data_to_cuda(
|
||||
self.current_batch_data[self.mode])
|
||||
self.eval_mode()
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=self.use_amp,
|
||||
dtype=self.dtype):
|
||||
batch_data['log_num'] = self.log_train_num
|
||||
results = self.run_step_eval(batch_data)
|
||||
self.train_mode()
|
||||
log_data, log_label = self.save_results(results)
|
||||
self.register_probe({
|
||||
'train_image':
|
||||
ProbeData(log_data,
|
||||
is_image=True,
|
||||
build_html=True,
|
||||
build_label=log_label)
|
||||
})
|
||||
self.register_probe({'train_label': log_label})
|
||||
return super(LatentDiffusionSolver, self).probe_data
|
||||
@@ -8,6 +8,8 @@ from abc import ABCMeta
|
||||
from collections import OrderedDict, defaultdict
|
||||
|
||||
import torch
|
||||
from torch.nn.parallel import DistributedDataParallel
|
||||
|
||||
from scepter.modules.data.dataset import DATASETS
|
||||
from scepter.modules.model.base_model import BaseModel
|
||||
from scepter.modules.model.metric.registry import METRICS
|
||||
@@ -18,15 +20,14 @@ from scepter.modules.solver.hooks import HOOKS
|
||||
from scepter.modules.utils.config import Config, dict_to_yaml
|
||||
from scepter.modules.utils.data import transfer_data_to_cuda
|
||||
from scepter.modules.utils.directory import get_relative_folder, osp_path
|
||||
from scepter.modules.utils.distribute import (
|
||||
dist, gather_data, we, all_reduce,
|
||||
_serialize_to_tensor, broadcast, _unserialize_from_tensor,
|
||||
all_reduce, barrier)
|
||||
from scepter.modules.utils.distribute import (_serialize_to_tensor,
|
||||
_unserialize_from_tensor,
|
||||
all_reduce, barrier, broadcast,
|
||||
gather_data, we)
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.modules.utils.logger import get_logger, init_logger
|
||||
from scepter.modules.utils.probe import (ProbeData, merge_gathered_probe,
|
||||
register_data)
|
||||
from torch.nn.parallel import DistributedDataParallel
|
||||
|
||||
try:
|
||||
import pytorch_lightning as pl
|
||||
@@ -187,6 +188,7 @@ try:
|
||||
except Exception as e:
|
||||
warnings.warn(f'{e}')
|
||||
|
||||
|
||||
def async_str(text):
|
||||
broadcast_size = torch.zeros(1, dtype=torch.long).to(we.device_id)
|
||||
if we.rank == 0:
|
||||
@@ -196,7 +198,8 @@ def async_str(text):
|
||||
broadcast(text_tensor, src=0)
|
||||
else:
|
||||
broadcast(broadcast_size, src=0)
|
||||
text_tensor = torch.empty((broadcast_size[0],), dtype=torch.uint8).to(we.device_id)
|
||||
text_tensor = torch.empty((broadcast_size[0], ),
|
||||
dtype=torch.uint8).to(we.device_id)
|
||||
broadcast(text_tensor, src=0)
|
||||
text = _unserialize_from_tensor(text_tensor)
|
||||
return text
|
||||
@@ -271,8 +274,8 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
self.pl_dir = self.work_dir
|
||||
self.log_file = osp_path(self.work_dir, cfg.LOG_FILE)
|
||||
self.optimizer, self.lr_scheduler = None, None
|
||||
self.resume_from: str = cfg.get("RESUME_FROM", None)
|
||||
self.max_epochs: int = cfg.get("MAX_EPOCHS", -1)
|
||||
self.resume_from: str = cfg.get('RESUME_FROM', None)
|
||||
self.max_epochs: int = cfg.get('MAX_EPOCHS', -1)
|
||||
self.use_pl = we.use_pl
|
||||
self.train_precision = self.cfg.get('TRAIN_PRECISION', 32)
|
||||
self._mode_set = set()
|
||||
@@ -284,7 +287,7 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
if not self.use_pl:
|
||||
world_size = we.world_size
|
||||
if world_size > 1:
|
||||
self._num_folds: int = cfg.get("NUM_FOLDS", 1)
|
||||
self._num_folds: int = cfg.get('NUM_FOLDS', 1)
|
||||
if cfg.have('MODE'):
|
||||
self._mode_set.add(cfg.MODE)
|
||||
self._mode = cfg.MODE
|
||||
@@ -765,17 +768,21 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
save_folder = pre_save_paras['save_folder']
|
||||
save_probe_prefix = pre_save_paras['save_probe_prefix']
|
||||
step = pre_save_paras['step']
|
||||
save_image_postfix = pre_save_paras.get('save_image_postfix', 'jpg')
|
||||
save_video_postfix = pre_save_paras.get('save_video_postfix', 'mp4')
|
||||
save_image_postfix = pre_save_paras.get('save_image_postfix',
|
||||
'jpg')
|
||||
save_video_postfix = pre_save_paras.get('save_video_postfix',
|
||||
'mp4')
|
||||
for k, v in self.collect_probe.items():
|
||||
if save_probe_prefix is not None:
|
||||
ret_prefix = os.path.join(save_folder, save_probe_prefix)
|
||||
else:
|
||||
ret_prefix = os.path.join(save_folder, k.replace('/', '_') + f'_step_{step}')
|
||||
v.presave(prefix = ret_prefix,
|
||||
image_postfix = save_image_postfix,
|
||||
video_postfix = save_video_postfix,
|
||||
rank = we.rank)
|
||||
ret_prefix = os.path.join(
|
||||
save_folder,
|
||||
k.replace('/', '_') + f'_step_{step}')
|
||||
v.presave(prefix=ret_prefix,
|
||||
image_postfix=save_image_postfix,
|
||||
video_postfix=save_video_postfix,
|
||||
rank=we.rank)
|
||||
gather_probe_data = gather_data(self._probe_data[self.mode])
|
||||
_dist_data_list = gather_data([self._dist_data[self.mode] or {}])
|
||||
if not we.rank == 0:
|
||||
@@ -877,7 +884,7 @@ class BaseSolver(object, metaclass=ABCMeta):
|
||||
if we.is_distributed:
|
||||
value = value.data.clone()
|
||||
all_reduce(value, group=we.data_parallel_group)
|
||||
value = value/we.data_group_world_size
|
||||
value = value / we.data_group_world_size
|
||||
ret[key] = value
|
||||
else:
|
||||
ret[key] = value
|
||||
|
||||
@@ -37,12 +37,14 @@ sharding_strategy_map = {
|
||||
|
||||
def shard_model(model,
|
||||
device_id,
|
||||
process_group=None,
|
||||
param_dtype=torch.bfloat16,
|
||||
reduce_dtype=torch.float32,
|
||||
buffer_dtype=torch.float32,
|
||||
fsdp_group = ['blocks'],
|
||||
fsdp_group=['blocks'],
|
||||
sharding_strategy=ShardingStrategy.FULL_SHARD,
|
||||
sync_module_states=False):
|
||||
sync_module_states=False,
|
||||
use_orig_params=False):
|
||||
wrap_modules = []
|
||||
for module_name in fsdp_group:
|
||||
if hasattr(model, module_name):
|
||||
@@ -54,7 +56,7 @@ def shard_model(model,
|
||||
warnings.warn("Can't find module {} in model".format(module_name))
|
||||
return FSDP(
|
||||
module=model,
|
||||
process_group=None,
|
||||
process_group=process_group,
|
||||
sharding_strategy=sharding_strategy,
|
||||
auto_wrap_policy=partial(
|
||||
# size_based_auto_wrap_policy, min_num_params=int(1e6),
|
||||
@@ -64,7 +66,8 @@ def shard_model(model,
|
||||
reduce_dtype=reduce_dtype,
|
||||
buffer_dtype=buffer_dtype),
|
||||
device_id=device_id,
|
||||
sync_module_states=sync_module_states)
|
||||
sync_module_states=sync_module_states,
|
||||
use_orig_params=use_orig_params)
|
||||
|
||||
|
||||
def get_module(instance, sub_module):
|
||||
@@ -193,7 +196,11 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
self.logger.info('Use fsdp as the backend of ddp.')
|
||||
else:
|
||||
self.logger.info('Use default backend.')
|
||||
self.use_scaler = cfg.get('USE_SCALER', True)
|
||||
self.enable_gradscaler = cfg.get('ENABLE_GRADSCALER', False)
|
||||
self.use_orig_params = cfg.get('USE_ORIG_PARAMS', False)
|
||||
self.model_shard = cfg.get('SHARDING_STRATEGY', 'full_shard')
|
||||
self.sharding_size = cfg.get('SHARDING_SIZE', None)
|
||||
self.reduce_dtype = getattr(torch,
|
||||
cfg.get('FSDP_REDUCE_DTYPE', 'float32'))
|
||||
self.buffer_dtype = getattr(torch,
|
||||
@@ -218,6 +225,7 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
self.model_to_device()
|
||||
self.init_lr()
|
||||
self.init_opti()
|
||||
self.logger.info(self.model)
|
||||
|
||||
def construct_hook(self):
|
||||
# initialize data
|
||||
@@ -279,6 +287,7 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
|
||||
def init_opti(self):
|
||||
import torch.cuda.amp as amp
|
||||
import torch.distributed as dist
|
||||
|
||||
if we.is_distributed:
|
||||
if self.use_fairscale:
|
||||
@@ -296,6 +305,30 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
self.model = ShardedDataParallel(self.model, self.optimizer)
|
||||
elif self.use_fsdp:
|
||||
shard_fn = partial
|
||||
if self.model_shard == 'hybrid_shard' and self.sharding_size is not None and self.sharding_size > 1:
|
||||
if self.sharding_size > we.world_size:
|
||||
self.logger.info(f'Reset sharding_size ({self.sharding_size}) to world_size ({we.world_size})')
|
||||
sharding_size = min(self.sharding_size, we.world_size)
|
||||
assert we.world_size % sharding_size == 0
|
||||
# mesh to facilitate rank indexing
|
||||
mesh = torch.arange(we.world_size).view(-1, sharding_size)
|
||||
# sharding groups
|
||||
for ranks in mesh.tolist():
|
||||
group = dist.new_group(ranks=ranks)
|
||||
if we.rank in ranks:
|
||||
sharding_group = group
|
||||
# replication groups
|
||||
for ranks in mesh.t().tolist():
|
||||
group = dist.new_group(ranks=ranks)
|
||||
if we.rank in ranks:
|
||||
replication_group = group
|
||||
# fsdp group tuple
|
||||
fsdp_group = (sharding_group, replication_group)
|
||||
fsdp_rank0 = we.rank // sharding_size * sharding_size
|
||||
else:
|
||||
fsdp_group = None
|
||||
fsdp_rank0 = 0
|
||||
|
||||
if self.shard_modules is not None:
|
||||
for module in self.shard_modules:
|
||||
if isinstance(module, str):
|
||||
@@ -303,25 +336,29 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
if sub_module is not None:
|
||||
sub_module = shard_model(
|
||||
sub_module,
|
||||
process_group=fsdp_group,
|
||||
device_id=we.device_id,
|
||||
param_dtype=self.dtype,
|
||||
reduce_dtype=self.reduce_dtype,
|
||||
buffer_dtype=self.buffer_dtype,
|
||||
sharding_strategy=sharding_strategy_map[self.model_shard],
|
||||
sync_module_states=True)
|
||||
sync_module_states=True,
|
||||
use_orig_params=self.use_orig_params)
|
||||
set_module(self.model, module, sub_module)
|
||||
elif isinstance(module, (dict, Config)):
|
||||
sub_module = get_module(self.model, module["MODULE"])
|
||||
if sub_module is not None:
|
||||
sub_module = shard_model(
|
||||
sub_module,
|
||||
process_group=fsdp_group,
|
||||
device_id=we.device_id,
|
||||
param_dtype=self.dtype,
|
||||
reduce_dtype=self.reduce_dtype,
|
||||
buffer_dtype=self.buffer_dtype,
|
||||
fsdp_group=module.get("FSDP_GROUP", ["blocks"]),
|
||||
sharding_strategy=sharding_strategy_map[self.model_shard],
|
||||
sync_module_states=True)
|
||||
sync_module_states=module.get("SYNC_MODULE_STATES", True),
|
||||
use_orig_params=self.use_orig_params)
|
||||
set_module(self.model, module["MODULE"], sub_module)
|
||||
else:
|
||||
self.logger.warning(
|
||||
@@ -374,22 +411,24 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
logger=self.logger,
|
||||
optimizer=self.optimizer)
|
||||
|
||||
if self.cfg.DTYPE in ['float16']:
|
||||
if self.use_scaler and self.cfg.DTYPE in ['float16', 'bfloat16']:
|
||||
if we.is_distributed:
|
||||
if self.use_fairscale:
|
||||
from fairscale.optim.grad_scaler import ShardedGradScaler
|
||||
self.scaler = ShardedGradScaler(enabled=True)
|
||||
self.scaler = ShardedGradScaler(enabled=self.enable_gradscaler)
|
||||
elif self.use_fsdp:
|
||||
from torch.distributed.fsdp.sharded_grad_scaler import ShardedGradScaler
|
||||
self.scaler = ShardedGradScaler(enabled=True,
|
||||
self.scaler = ShardedGradScaler(enabled=self.enable_gradscaler,
|
||||
process_group=None)
|
||||
else:
|
||||
self.scaler = amp.GradScaler()
|
||||
else:
|
||||
self.scaler = amp.GradScaler(enabled=self.enable_gradscaler)
|
||||
elif self.cfg.DTYPE in ['float16']:
|
||||
self.scaler = amp.GradScaler()
|
||||
else:
|
||||
self.scaler = None
|
||||
else:
|
||||
self.scaler = None
|
||||
self.logger.info(self.model)
|
||||
|
||||
def load_checkpoint(self, checkpoint: dict):
|
||||
"""
|
||||
Load checkpoint function
|
||||
@@ -498,7 +537,7 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
model = self.model
|
||||
if self.save_modules is not None:
|
||||
for module in self.save_modules:
|
||||
current_module = get_module(self.model, module)
|
||||
current_module = get_module(model, module)
|
||||
if current_module is not None:
|
||||
ckpt['model'][module] = current_module.state_dict()
|
||||
else:
|
||||
@@ -510,12 +549,13 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
model = self.model
|
||||
if self.save_modules is not None:
|
||||
for module in self.save_modules:
|
||||
current_module = get_module(self.model, module)
|
||||
current_module = get_module(model, module)
|
||||
if current_module is not None:
|
||||
ckpt['model'][module] = current_module.state_dict()
|
||||
else:
|
||||
ckpt['model'] = model.state_dict()
|
||||
if self.optimizer and not self.use_fairscale:
|
||||
if (self.optimizer and not self.use_fairscale
|
||||
and self.save_modules and "optimizer" in self.save_modules):
|
||||
if self.use_fsdp and we.is_distributed:
|
||||
ckpt['optimizer'] = OrderedDict()
|
||||
for module in self.train_modules:
|
||||
@@ -579,9 +619,6 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
'batch_size': len(batch_data['prompt'])
|
||||
})
|
||||
self.current_batch_data[self.mode] = batch_data
|
||||
if self.sample_args:
|
||||
self.current_batch_data[self.mode].update(
|
||||
self.sample_args.get_lowercase_dict())
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=self.use_amp,
|
||||
dtype=self.dtype):
|
||||
@@ -708,9 +745,8 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
swift_cfg_dict[f'{t_id}_{cfg_name}'] = init_config
|
||||
if len(swift_cfg_dict) > 0:
|
||||
from swift import Swift
|
||||
model = Swift.prepare_model(self.model, config=swift_cfg_dict)
|
||||
|
||||
self.logger.info([(key, param.shape) for key, param in model.named_parameters() if param.requires_grad])
|
||||
model = Swift.prepare_model(self.model, config=swift_cfg_dict, autocast_adapter_dtype=False)
|
||||
self.logger.info([(key, param.shape) for key, param in model.named_parameters() if param.requires_grad])
|
||||
return model
|
||||
|
||||
def freeze(self, freeze_cfg, model=None):
|
||||
@@ -815,13 +851,15 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
@property
|
||||
def probe_data(self):
|
||||
if not we.debug and self.mode == 'train':
|
||||
batch_data = transfer_data_to_cuda(self.current_batch_data[self.mode])
|
||||
batch_data = self.current_batch_data[self.mode]
|
||||
if self.sample_args:
|
||||
batch_data.update(self.sample_args.get_lowercase_dict())
|
||||
self.eval_mode()
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=self.use_amp,
|
||||
dtype=self.dtype):
|
||||
batch_data['log_num'] = self.log_train_num
|
||||
results = self.run_step_eval(batch_data)
|
||||
results = self.run_step_eval(transfer_data_to_cuda(batch_data))
|
||||
images = batch_data['image'] if 'image' in batch_data else [None] * len(results)
|
||||
self.train_mode()
|
||||
log_data, log_label = [], []
|
||||
|
||||
@@ -0,0 +1,190 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.solver import LatentDiffusionSolver
|
||||
from scepter.modules.solver.registry import SOLVERS
|
||||
from scepter.modules.utils.data import transfer_data_to_cuda
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.probe import ProbeData
|
||||
|
||||
@SOLVERS.register_class()
|
||||
class LatentDiffusionVideoSolver(LatentDiffusionSolver):
|
||||
para_dict = LatentDiffusionSolver.para_dict
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.fps = cfg.get("FPS", 8)
|
||||
|
||||
def save_results(self, results):
|
||||
log_data, log_label = [], []
|
||||
for result in results:
|
||||
ret_videos, ret_labels = [], []
|
||||
if 'edit_video' in result:
|
||||
ret_videos.append((result['edit_video'].permute(1, 2, 3, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
ret_labels.append("left: edit video")
|
||||
if 'edit_image' in result:
|
||||
ret_videos.append((result['edit_image'].permute(1, 2, 3, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
ret_labels.append("left: edit image")
|
||||
if 'target_video' in result:
|
||||
if len(ret_videos) > 0:
|
||||
ret_labels.append("middle: target video")
|
||||
else:
|
||||
ret_labels.append("left: target video")
|
||||
ret_videos.append((result['target_video'].permute(1, 2, 3, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
|
||||
ret_videos.append((result['reconstruct_video'].permute(1, 2, 3, 0).cpu().numpy() *
|
||||
255).astype(np.uint8))
|
||||
ret_labels.append("right: generation video" + " Prompt: " + result['instruction'])
|
||||
|
||||
log_data.append(ret_videos)
|
||||
log_label.append(ret_labels)
|
||||
return log_data, log_label
|
||||
|
||||
def run_train(self):
|
||||
self.train_mode()
|
||||
self.before_all_iter(self.hooks_dict[self._mode])
|
||||
data_iter = iter(self.datas[self._mode].dataloader)
|
||||
self.print_memory_status()
|
||||
for step in range(self.max_steps):
|
||||
if 'eval' in self._mode_set and (self.eval_interval > 0 and
|
||||
step % self.eval_interval == 0):
|
||||
self.run_eval()
|
||||
self.train_mode()
|
||||
batch_data = next(data_iter)
|
||||
self.before_iter(self.hooks_dict[self._mode])
|
||||
if 'meta' in batch_data and isinstance(batch_data['meta'], dict):
|
||||
self.register_probe({
|
||||
'data_key':
|
||||
ProbeData(batch_data['meta'].get('data_key', []),
|
||||
view_distribute=True)
|
||||
})
|
||||
self.register_probe({
|
||||
'prompt': batch_data['prompt'],
|
||||
'batch_size': len(batch_data['prompt'])
|
||||
})
|
||||
self.current_batch_data[self.mode] = batch_data
|
||||
if self.sample_args:
|
||||
self.current_batch_data[self.mode].update(
|
||||
self.sample_args.get_lowercase_dict())
|
||||
batch_data = transfer_data_to_cuda(batch_data)
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=self.use_amp,
|
||||
dtype=self.dtype):
|
||||
results = self.run_step_train(
|
||||
batch_data,
|
||||
step,
|
||||
step=self.total_iter,
|
||||
rank=we.rank)
|
||||
self._iter_outputs[self._mode] = self._reduce_scalar(results)
|
||||
self.after_iter(self.hooks_dict[self._mode])
|
||||
if we.debug:
|
||||
self.print_trainable_params_status(prefix='model.')
|
||||
if 'eval' in self._mode_set and (self.eval_interval > 0
|
||||
and step == self.max_steps - 1):
|
||||
self.run_eval()
|
||||
self.train_mode()
|
||||
self.after_all_iter(self.hooks_dict[self._mode])
|
||||
|
||||
@torch.no_grad()
|
||||
def run_eval(self):
|
||||
self.eval_mode()
|
||||
self.before_all_iter(self.hooks_dict[self._mode])
|
||||
all_results = []
|
||||
for batch_idx, batch_data in tqdm(
|
||||
enumerate(self.datas[self._mode].dataloader)):
|
||||
self.before_iter(self.hooks_dict[self._mode])
|
||||
if self.sample_args:
|
||||
batch_data.update(self.sample_args.get_lowercase_dict())
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=self.use_amp,
|
||||
dtype=self.dtype):
|
||||
results = self.run_step_eval(transfer_data_to_cuda(batch_data),
|
||||
batch_idx,
|
||||
step=self.total_iter,
|
||||
rank=we.rank)
|
||||
all_results.extend(results)
|
||||
self.after_iter(self.hooks_dict[self._mode])
|
||||
log_data, log_label = self.save_results(all_results)
|
||||
self.register_probe({'eval_label': log_label})
|
||||
self.register_probe({
|
||||
'eval_video':
|
||||
ProbeData(log_data,
|
||||
is_image=False,
|
||||
is_video=True,
|
||||
fps=self.fps,
|
||||
build_html=True,
|
||||
build_label=log_label)
|
||||
})
|
||||
self.after_all_iter(self.hooks_dict[self._mode])
|
||||
|
||||
@torch.no_grad()
|
||||
def run_test(self):
|
||||
self.test_mode()
|
||||
self.before_all_iter(self.hooks_dict[self._mode])
|
||||
all_results = []
|
||||
for batch_idx, batch_data in tqdm(
|
||||
enumerate(self.datas[self._mode].dataloader)):
|
||||
self.before_iter(self.hooks_dict[self._mode])
|
||||
if self.sample_args:
|
||||
batch_data.update(self.sample_args.get_lowercase_dict())
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=self.use_amp,
|
||||
dtype=self.dtype):
|
||||
results = self.run_step_eval(transfer_data_to_cuda(batch_data),
|
||||
batch_idx,
|
||||
step=self.total_iter,
|
||||
rank=we.rank)
|
||||
all_results.extend(results)
|
||||
self.after_iter(self.hooks_dict[self._mode])
|
||||
log_data, log_label = self.save_results(all_results)
|
||||
self.register_probe({'test_label': log_label})
|
||||
self.register_probe({
|
||||
'test_video':
|
||||
ProbeData(log_data,
|
||||
is_image=False,
|
||||
is_video=True,
|
||||
fps=self.fps,
|
||||
build_html=True,
|
||||
build_label=log_label)
|
||||
})
|
||||
self.after_all_iter(self.hooks_dict[self._mode])
|
||||
|
||||
@property
|
||||
def probe_data(self):
|
||||
if not we.debug and self.mode == 'train':
|
||||
batch_data = self.current_batch_data[self.mode]
|
||||
if self.sample_args is not None:
|
||||
batch_data.update(self.sample_args.get_lowercase_dict())
|
||||
self.eval_mode()
|
||||
with torch.autocast(device_type='cuda',
|
||||
enabled=self.use_amp,
|
||||
dtype=self.dtype):
|
||||
batch_data['log_train_num'] = self.log_train_num
|
||||
all_results = self.run_step_eval(transfer_data_to_cuda(batch_data))
|
||||
self.train_mode()
|
||||
log_data, log_label = self.save_results(all_results)
|
||||
self.register_probe({
|
||||
'train_video':
|
||||
ProbeData(log_data,
|
||||
is_image=False,
|
||||
is_video=True,
|
||||
fps=self.fps,
|
||||
build_html=True,
|
||||
build_label=log_label)
|
||||
})
|
||||
self.register_probe({'train_label': log_label})
|
||||
return super(LatentDiffusionSolver, self).probe_data
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('SOLVER',
|
||||
__class__.__name__,
|
||||
LatentDiffusionVideoSolver.para_dict,
|
||||
set_name=True)
|
||||
@@ -4,13 +4,12 @@ import os
|
||||
import warnings
|
||||
|
||||
import torch
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
from scepter.modules.utils.distribute import we
|
||||
|
||||
from scepter.modules.solver.hooks.hook import Hook
|
||||
from scepter.modules.solver.hooks.registry import HOOKS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
_DEFAULT_BACKWARD_PRIORITY = 0
|
||||
|
||||
@@ -87,15 +86,21 @@ class BackwardHook(Hook):
|
||||
os.makedirs(self._local_log_dir, exist_ok=True)
|
||||
if self.do_profile:
|
||||
self.prof = torch.profiler.profile(
|
||||
schedule=torch.profiler.schedule(wait=self.wait, warmup=self.warmup, active=self.active, repeat=self.repeat),
|
||||
on_trace_ready=torch.profiler.tensorboard_trace_handler(self._local_log_dir),
|
||||
schedule=torch.profiler.schedule(wait=self.wait,
|
||||
warmup=self.warmup,
|
||||
active=self.active,
|
||||
repeat=self.repeat),
|
||||
on_trace_ready=torch.profiler.tensorboard_trace_handler(
|
||||
self._local_log_dir),
|
||||
record_shapes=True,
|
||||
with_stack=True)
|
||||
self.prof.start()
|
||||
solver.logger.info(f'Profiler start ...')
|
||||
solver.logger.info('Profiler start ...')
|
||||
solver.logger.info(f'Profiler: save to {self.log_dir}')
|
||||
|
||||
def profile(self, solver):
|
||||
if self.prof is None: return
|
||||
if self.prof is None:
|
||||
return
|
||||
if we.rank == 0 and self.do_profile:
|
||||
if self.profile_step < self.wait + self.warmup + self.active:
|
||||
self.prof.step()
|
||||
@@ -103,8 +108,10 @@ class BackwardHook(Hook):
|
||||
else:
|
||||
self.prof.stop()
|
||||
self.do_profile = False
|
||||
solver.logger.info(f'Profiler stop after {self.profile_step} steps')
|
||||
solver.logger.info(
|
||||
f'Profiler stop after {self.profile_step} steps')
|
||||
FS.put_dir_from_local_dir(self._local_log_dir, self.log_dir)
|
||||
|
||||
def grad_clip(self, parameters):
|
||||
torch.nn.utils.clip_grad_norm_(parameters=parameters,
|
||||
max_norm=self.gradient_clip,
|
||||
@@ -118,26 +125,27 @@ class BackwardHook(Hook):
|
||||
)
|
||||
return
|
||||
if solver.scaler is not None:
|
||||
solver.scaler.scale(solver.loss/self.accumulate_step).backward()
|
||||
if self.gradient_clip > 0:
|
||||
solver.scaler.unscale_(solver.optimizer)
|
||||
self.grad_clip(solver.train_parameters())
|
||||
solver.scaler.scale(solver.loss /
|
||||
self.accumulate_step).backward()
|
||||
self.current_step += 1
|
||||
# Suppose profiler run after backward, so we need to set backward_prev_step
|
||||
# as the previous one step before the backward step
|
||||
if self.current_step % self.accumulate_step == 0:
|
||||
if self.gradient_clip > 0:
|
||||
solver.scaler.unscale_(solver.optimizer)
|
||||
self.grad_clip(solver.train_parameters())
|
||||
self.profile(solver)
|
||||
solver.scaler.step(solver.optimizer)
|
||||
solver.scaler.update()
|
||||
solver.optimizer.zero_grad()
|
||||
else:
|
||||
(solver.loss/self.accumulate_step).backward()
|
||||
if self.gradient_clip > 0:
|
||||
self.grad_clip(solver.train_parameters())
|
||||
(solver.loss / self.accumulate_step).backward()
|
||||
self.current_step += 1
|
||||
# Suppose profiler run after backward, so we need to set backward_prev_step
|
||||
# as the previous one step before the backward step
|
||||
if self.current_step % self.accumulate_step == 0:
|
||||
if self.gradient_clip > 0:
|
||||
self.grad_clip(solver.train_parameters())
|
||||
self.profile(solver)
|
||||
solver.optimizer.step()
|
||||
solver.optimizer.zero_grad()
|
||||
|
||||
@@ -128,13 +128,24 @@ class CheckpointHook(Hook):
|
||||
solver.work_dir,
|
||||
'checkpoints/{}-{}'.format(self.save_name_prefix,
|
||||
solver.total_iter + 1))
|
||||
if we.rank == 0:
|
||||
local_folder, _ = FS.map_to_local(save_path)
|
||||
if hasattr(solver.model, 'module'):
|
||||
solver.model.module.save_pretrained(local_folder)
|
||||
else:
|
||||
solver.model.save_pretrained(local_folder)
|
||||
FS.put_dir_from_local_dir(local_folder, save_path)
|
||||
solver_model = solver.model.module if hasattr(solver.model, 'module') else solver.model
|
||||
if isinstance(solver_model.base_model.model, torch.distributed.fsdp.FullyShardedDataParallel):
|
||||
full_state_dict_config = torch.distributed.fsdp.FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
|
||||
with torch.distributed.fsdp.FullyShardedDataParallel.state_dict_type(solver_model.base_model, torch.distributed.fsdp.StateDictType.FULL_STATE_DICT, full_state_dict_config):
|
||||
state_dict = solver_model.base_model.state_dict()
|
||||
if we.rank == 0:
|
||||
state_dict_new = {}
|
||||
local_folder, _ = FS.map_to_local(save_path)
|
||||
for adapter_name in solver_model.adapters.keys():
|
||||
state_dict_adapter = solver_model.adapters[adapter_name].state_dict_callback(state_dict, adapter_name, replace_key=False)
|
||||
state_dict_new.update(state_dict_adapter)
|
||||
solver_model.save_pretrained(local_folder, state_dict=state_dict_new)
|
||||
FS.put_dir_from_local_dir(local_folder, save_path)
|
||||
else:
|
||||
if we.rank == 0:
|
||||
local_folder, _ = FS.map_to_local(save_path)
|
||||
solver_model.save_pretrained(local_folder)
|
||||
FS.put_dir_from_local_dir(local_folder, save_path)
|
||||
else:
|
||||
if hasattr(solver, 'save_pretrained'):
|
||||
save_path = osp.join(
|
||||
|
||||
@@ -117,6 +117,7 @@ class LogHook(Hook):
|
||||
super(LogHook, self).__init__(cfg, logger=logger)
|
||||
self.priority = cfg.get('PRIORITY', _DEFAULT_LOG_PRIORITY)
|
||||
self.log_interval = cfg.get('LOG_INTERVAL', 10)
|
||||
self.interval = cfg.get('INTERVAL', self.log_interval)
|
||||
self.show_gpu_mem = cfg.get('SHOW_GPU_MEM', False)
|
||||
self.log_agg_dict = defaultdict(LogAgg)
|
||||
|
||||
@@ -147,18 +148,18 @@ class LogHook(Hook):
|
||||
outputs['time'] = iter_time
|
||||
outputs['data_time'] = self.data_time
|
||||
if solver.mode in self.batch_size:
|
||||
outputs['throughput'] = int(self.batch_size[solver.mode] * we.world_size / iter_time * 86400)
|
||||
outputs['throughput'] = int(self.batch_size[solver.mode] * we.data_group_world_size / iter_time * 86400)
|
||||
log_agg.update(outputs, 1)
|
||||
log_agg = log_agg.aggregate(self.log_interval)
|
||||
log_agg = log_agg.aggregate(self.interval)
|
||||
if 'throughput' in log_agg:
|
||||
log_agg['throughput'] = f"{int(log_agg['throughput'][-1])}/day"
|
||||
if solver.mode in self.batch_size:
|
||||
log_agg['all_throughput'] = (solver.iter + 1) * we.world_size * self.batch_size[solver.mode]
|
||||
log_agg['all_throughput'] = (solver.iter + 1) * we.data_group_world_size * self.batch_size[solver.mode]
|
||||
|
||||
if self.show_gpu_mem:
|
||||
log_agg['nvidia-smi'] = str(print_memory_status()) +"MiB"
|
||||
|
||||
if (solver.iter + 1) % self.log_interval == 0:
|
||||
if (solver.iter + 1) % self.interval == 0:
|
||||
_print_iter_log(solver,
|
||||
log_agg,
|
||||
start_time=self.start_time,
|
||||
@@ -206,7 +207,7 @@ class LogHook(Hook):
|
||||
solver.logger.info(f'Current Epoch {mode} Summary:')
|
||||
log_agg = self.log_agg_dict[mode]
|
||||
_print_iter_log(solver,
|
||||
log_agg.aggregate(self.log_interval),
|
||||
log_agg.aggregate(self.interval),
|
||||
start_time=self.start_time,
|
||||
mode=mode)
|
||||
if not mode == 'train':
|
||||
@@ -242,6 +243,7 @@ class TensorboardLogHook(Hook):
|
||||
self.priority = cfg.get('PRIORITY', _DEFAULT_LOG_PRIORITY)
|
||||
self.log_dir = cfg.get('LOG_DIR', None)
|
||||
self.log_interval = cfg.get('LOG_INTERVAL', 1000)
|
||||
self.interval = cfg.get('INTERVAL', self.log_interval)
|
||||
self._local_log_dir = None
|
||||
self.writer: Optional[SummaryWriter] = None
|
||||
|
||||
@@ -286,7 +288,7 @@ class TensorboardLogHook(Hook):
|
||||
self.writer.add_scalar(f'{mode}/iter/{key}',
|
||||
value,
|
||||
global_step=solver.total_iter)
|
||||
if solver.total_iter % self.log_interval:
|
||||
if solver.total_iter % self.interval:
|
||||
self.writer.flush()
|
||||
# Put to remote file systems every epoch
|
||||
FS.put_dir_from_local_dir(self._local_log_dir, self.log_dir)
|
||||
|
||||
@@ -7,6 +7,7 @@ import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image, ImageFile
|
||||
|
||||
from scepter.modules.transform.registry import TRANSFORMS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
@@ -15,18 +16,18 @@ from scepter.modules.utils.file_system import DATA_FS as FS
|
||||
ImageFile.LOAD_TRUNCATED_IMAGES = True
|
||||
|
||||
|
||||
def pillow_convert(image, rgb_order):
|
||||
if image.mode != rgb_order:
|
||||
def pillow_convert(image, cvt_type):
|
||||
if image.mode != cvt_type:
|
||||
if image.mode == 'P':
|
||||
image = image.convert(f'{rgb_order}A')
|
||||
if image.mode == f'{rgb_order}A':
|
||||
bg = Image.new(rgb_order,
|
||||
image = image.convert(f'{cvt_type}A')
|
||||
if image.mode == f'{cvt_type}A':
|
||||
bg = Image.new(cvt_type,
|
||||
size=(image.width, image.height),
|
||||
color=(255, 255, 255))
|
||||
bg.paste(image, (0, 0), mask=image)
|
||||
image = bg
|
||||
else:
|
||||
image = image.convert('RGB')
|
||||
image = image.convert(cvt_type)
|
||||
return image
|
||||
|
||||
|
||||
|
||||
@@ -201,6 +201,11 @@ def broadcast(tensor, src, group=None, **kwargs):
|
||||
return dist.broadcast(tensor, src, group, **kwargs)
|
||||
|
||||
|
||||
def broadcast_object_list(object_list, src, group=None, **kwargs):
|
||||
if we.is_distributed:
|
||||
return dist.broadcast_object_list(object_list, src, group, **kwargs)
|
||||
|
||||
|
||||
def barrier():
|
||||
if we.is_distributed:
|
||||
dist.barrier()
|
||||
@@ -716,6 +721,7 @@ class Workenv(object):
|
||||
torch.backends.cudnn.benchmark = config.ENV.get(
|
||||
'CUDNN_BENCHMARK', False)
|
||||
fn(config)
|
||||
return
|
||||
else:
|
||||
import torch.multiprocessing as mp
|
||||
if 'MASTER_ADDR' not in os.environ:
|
||||
@@ -736,10 +742,13 @@ class Workenv(object):
|
||||
if self.is_distributed:
|
||||
self.backend = config.ENV.get('BACKEND', 'nccl')
|
||||
self.sync_bn = config.ENV.get('SYNC_BN', False)
|
||||
mp.spawn(mp_worker,
|
||||
nprocs=ngpus_per_node,
|
||||
args=(ngpus_per_node, config, fn, pmi_rank, world_size,
|
||||
self))
|
||||
spawn_join = config.ENV.get('SPAWN_JOIN', True)
|
||||
context = mp.spawn(mp_worker,
|
||||
nprocs=ngpus_per_node,
|
||||
join=spawn_join,
|
||||
args=(ngpus_per_node, config, fn, pmi_rank, world_size,
|
||||
self))
|
||||
return context
|
||||
|
||||
def get_env(self):
|
||||
ret_dict = {}
|
||||
|
||||
+145
-78
@@ -1,7 +1,6 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import json
|
||||
import os.path
|
||||
from io import BytesIO
|
||||
from numbers import Number
|
||||
@@ -9,6 +8,7 @@ from numbers import Number
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
|
||||
@@ -76,6 +76,7 @@ def merge_gathered_probe(all_gathered_data):
|
||||
all_gathered_data[key] = ProbeData(
|
||||
new_data,
|
||||
is_image=ret_data.is_image,
|
||||
is_video=ret_data.is_video,
|
||||
build_html=ret_data.build_html,
|
||||
build_label=ret_data.build_label,
|
||||
view_distribute=ret_data.view_distribute,
|
||||
@@ -96,6 +97,7 @@ def merge_gathered_probe(all_gathered_data):
|
||||
all_gathered_data[key] = ProbeData(
|
||||
ret_data.data,
|
||||
is_image=ret_data.is_image,
|
||||
is_video=ret_data.is_video,
|
||||
build_html=ret_data.build_html,
|
||||
build_label=ret_data.build_label,
|
||||
view_distribute=ret_data.view_distribute,
|
||||
@@ -106,7 +108,7 @@ def merge_gathered_probe(all_gathered_data):
|
||||
|
||||
|
||||
class MediaHandler():
|
||||
def __init__(self, batch_size = 10):
|
||||
def __init__(self, batch_size=10):
|
||||
self.file_list = []
|
||||
self.target_path_list = []
|
||||
self.target_status = {}
|
||||
@@ -116,20 +118,27 @@ class MediaHandler():
|
||||
self.file_list.append(source_file)
|
||||
self.target_path_list.append(target_path)
|
||||
if len(self.file_list) > 2 * self.batch_size:
|
||||
generator = FS.put_batch_objects_to(self.file_list, self.target_path_list, batch_size=self.batch_size)
|
||||
generator = FS.put_batch_objects_to(self.file_list,
|
||||
self.target_path_list,
|
||||
batch_size=self.batch_size)
|
||||
for local_path, target_path, flg in generator:
|
||||
self.target_status[target_path] = flg
|
||||
self.file_list.clear()
|
||||
self.target_path_list.clear()
|
||||
|
||||
def sync(self):
|
||||
if len(self.file_list) > 0:
|
||||
if len(self.file_list) > 4 * self.batch_size:
|
||||
generator = FS.put_batch_objects_to(self.file_list, self.target_path_list, batch_size=self.batch_size)
|
||||
generator = FS.put_batch_objects_to(self.file_list,
|
||||
self.target_path_list,
|
||||
batch_size=self.batch_size)
|
||||
for local_path, target_path, flg in generator:
|
||||
self.target_status[target_path] = flg
|
||||
else:
|
||||
for file_, target_path in zip(self.file_list, self.target_path_list):
|
||||
self.target_status[target_path] = FS.put_object(file_.getvalue(), target_path)
|
||||
for file_, target_path in zip(self.file_list,
|
||||
self.target_path_list):
|
||||
self.target_status[target_path] = FS.put_object(
|
||||
file_.getvalue(), target_path)
|
||||
self.file_list.clear()
|
||||
self.target_path_list.clear()
|
||||
|
||||
@@ -148,7 +157,7 @@ class ProbeData():
|
||||
build_html=False,
|
||||
build_label=None,
|
||||
view_distribute=False,
|
||||
is_presave = False):
|
||||
is_presave=False):
|
||||
''' Probe Data Initialize.
|
||||
We only support basic types such as [torch.Tensor, numpy.ndarray, number, str],
|
||||
or [dict, list] of [dict, list,
|
||||
@@ -243,7 +252,8 @@ class ProbeData():
|
||||
for v_idx, v_v in enumerate(v):
|
||||
if not check_legal_type(v_v):
|
||||
if isinstance(v_v, torch.Tensor):
|
||||
data[idx][v_idx] = v_v.detach().cpu().numpy()
|
||||
data[idx][v_idx] = v_v.detach().cpu(
|
||||
).numpy()
|
||||
self.basic_type = False
|
||||
elif isinstance(v_v, np.ndarray):
|
||||
data[idx][v_idx] = v_v
|
||||
@@ -296,25 +306,33 @@ class ProbeData():
|
||||
if extension.lower() in ['png']:
|
||||
return 'PNG'
|
||||
return 'JPEG'
|
||||
def save_one_video(self, file_path, videos, fps = 8):
|
||||
|
||||
def save_one_video(self, file_path, videos, fps=8):
|
||||
# write video
|
||||
import imageio
|
||||
try:
|
||||
writer = imageio.get_writer(file_path, fps=fps, format=".mp4", codec='libx264', quality=8)
|
||||
writer = imageio.get_writer(file_path,
|
||||
fps=fps,
|
||||
format='.mp4',
|
||||
codec='libx264',
|
||||
quality=8)
|
||||
for frame in videos:
|
||||
writer.append_data(frame)
|
||||
writer.close()
|
||||
return True
|
||||
except:
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def save_video(self, file_prefix, videos, video_postfix, fps = 8, rank = 0):
|
||||
def save_video(self, file_prefix, videos, video_postfix, fps=8, rank=0):
|
||||
if isinstance(videos, list):
|
||||
for video in videos:
|
||||
if isinstance(video, list):
|
||||
raise f"Only surpport one layer nested list."
|
||||
return [self.save_video(file_prefix + f'_{rank}_{idx}', v, video_postfix, fps) for idx, v in enumerate(videos)]
|
||||
raise 'Only surpport one layer nested list.'
|
||||
return [
|
||||
self.save_video(file_prefix + f'_{rank}_{idx}', v,
|
||||
video_postfix, fps)
|
||||
for idx, v in enumerate(videos)
|
||||
]
|
||||
np_shape = videos.shape
|
||||
# 4D
|
||||
shape_str = '_'.join([str(v) for v in np_shape])
|
||||
@@ -326,15 +344,22 @@ class ProbeData():
|
||||
file_list = []
|
||||
for idx in range(np_shape[0]):
|
||||
if videos[idx].shape[0] > 1:
|
||||
file_path = os.path.join(file_prefix, f'probe_{rank}_{idx}_[{shape_str}].{video_postfix}')
|
||||
file_path = os.path.join(
|
||||
file_prefix,
|
||||
f'probe_{rank}_{idx}_[{shape_str}].{video_postfix}'
|
||||
)
|
||||
byio = BytesIO()
|
||||
is_suc = self.save_one_video(byio, videos[idx], fps)
|
||||
if not is_suc:
|
||||
byio.write(b"")
|
||||
byio.write(b'')
|
||||
else:
|
||||
file_path = os.path.join(file_prefix, f'probe_{rank}_{idx}_[{shape_str}].{self.image_postfix}')
|
||||
file_path = os.path.join(
|
||||
file_prefix,
|
||||
f'probe_{rank}_{idx}_[{shape_str}].{self.image_postfix}'
|
||||
)
|
||||
byio = BytesIO()
|
||||
Image.fromarray(videos[idx][0]).save(byio, self.get_format(self.image_postfix))
|
||||
Image.fromarray(videos[idx][0]).save(
|
||||
byio, self.get_format(self.image_postfix))
|
||||
self.media_handler.append(byio, file_path)
|
||||
file_list.append(file_path)
|
||||
return file_list
|
||||
@@ -349,25 +374,32 @@ class ProbeData():
|
||||
byio = BytesIO()
|
||||
is_suc = self.save_one_video(byio, videos, fps)
|
||||
if not is_suc:
|
||||
byio.write(b"")
|
||||
byio.write(b'')
|
||||
else:
|
||||
file_path = file_prefix + f'_probe_{rank}_[{shape_str}].{self.image_postfix}'
|
||||
byio = BytesIO()
|
||||
Image.fromarray(videos[0]).save(byio, self.get_format(self.image_postfix))
|
||||
Image.fromarray(videos[0]).save(
|
||||
byio, self.get_format(self.image_postfix))
|
||||
self.media_handler.append(byio, file_path)
|
||||
return file_path
|
||||
else:
|
||||
videos = videos.reshape(list(videos.shape) + [1])
|
||||
return self.save_video(file_prefix, videos, video_postfix, fps = fps)
|
||||
return self.save_video(file_prefix,
|
||||
videos,
|
||||
video_postfix,
|
||||
fps=fps)
|
||||
else:
|
||||
raise f"Ensure your data's dim is BFWHC or FWHC, and channel is 1 or 3 for {file_prefix}"
|
||||
|
||||
def save_image(self, file_prefix, images, image_postfix, rank = 0):
|
||||
def save_image(self, file_prefix, images, image_postfix, rank=0):
|
||||
if isinstance(images, list):
|
||||
for image in images:
|
||||
if isinstance(image, list):
|
||||
raise f"Only surpport one layer nested list."
|
||||
return [self.save_image(file_prefix + f'_{rank}_{idx}', v, image_postfix) for idx, v in enumerate(images)]
|
||||
raise TypeError('Only surpport one layer nested list.')
|
||||
return [
|
||||
self.save_image(file_prefix + f'_{rank}_{idx}', v,
|
||||
image_postfix) for idx, v in enumerate(images)
|
||||
]
|
||||
np_shape = images.shape
|
||||
# 4D
|
||||
shape_str = '_'.join([str(v) for v in np_shape])
|
||||
@@ -378,9 +410,12 @@ class ProbeData():
|
||||
images = images.reshape(images.shape[:-1])
|
||||
file_list = []
|
||||
for idx in range(np_shape[0]):
|
||||
file_path = os.path.join(file_prefix, f'probe_{rank}_{idx}_[{shape_str}].{image_postfix}')
|
||||
file_path = os.path.join(
|
||||
file_prefix,
|
||||
f'probe_{rank}_{idx}_[{shape_str}].{image_postfix}')
|
||||
byio = BytesIO()
|
||||
Image.fromarray(images[idx, ...]).save(byio, self.get_format(self.image_postfix))
|
||||
Image.fromarray(images[idx, ...]).save(
|
||||
byio, self.get_format(self.image_postfix))
|
||||
self.media_handler.append(byio, file_path)
|
||||
file_list.append(file_path)
|
||||
return file_list
|
||||
@@ -392,7 +427,8 @@ class ProbeData():
|
||||
images = images.reshape(images.shape[:-1])
|
||||
file_path = file_prefix + f'_probe_{rank}_[{shape_str}].{image_postfix}'
|
||||
byio = BytesIO()
|
||||
Image.fromarray(images).save(byio, self.get_format(self.image_postfix))
|
||||
Image.fromarray(images).save(
|
||||
byio, self.get_format(self.image_postfix))
|
||||
self.media_handler.append(byio, file_path)
|
||||
return file_path
|
||||
else:
|
||||
@@ -401,13 +437,14 @@ class ProbeData():
|
||||
elif len(np_shape) == 2:
|
||||
file_path = file_prefix + f'_probe_{rank}_[{shape_str}].{image_postfix}'
|
||||
byio = BytesIO()
|
||||
Image.fromarray(images).save(byio, self.get_format(self.image_postfix))
|
||||
Image.fromarray(images).save(byio,
|
||||
self.get_format(self.image_postfix))
|
||||
self.media_handler.append(byio, file_path)
|
||||
return file_path
|
||||
else:
|
||||
raise f"Ensure your data's dim is BWHC or WHC or WH, and channel is 1 or 3 for {file_prefix}"
|
||||
|
||||
def save_npy(self, file_prefix, data, rank = 0):
|
||||
def save_npy(self, file_prefix, data, rank=0):
|
||||
shape_str = '_'.join([str(v) for v in data.shape])
|
||||
file_path = file_prefix + f'_{rank}_{shape_str}.npy'
|
||||
byio = BytesIO()
|
||||
@@ -420,8 +457,9 @@ class ProbeData():
|
||||
with FS.put_to(html_prefix) as local_path:
|
||||
with open(local_path, 'w') as f:
|
||||
f.writelines('<meta charset="utf-8">\n')
|
||||
f.writelines('<style>input{height:' + f'{height}px;' +
|
||||
'opacity:1.0;} textarea {font-size: 32px;}</style>\n')
|
||||
f.writelines(
|
||||
'<style>input{height:' + f'{height}px;' +
|
||||
'opacity:1.0;} textarea {font-size: 32px;}</style>\n')
|
||||
f.writelines('<br><hr/>\n')
|
||||
all_ranks = list()
|
||||
is_textarea = False
|
||||
@@ -433,16 +471,16 @@ class ProbeData():
|
||||
one_label = one_label.replace('<', '<').replace(
|
||||
'>', '>')
|
||||
try:
|
||||
url = FS.get_url(one_path,
|
||||
lifecycle=3600 * 365 * 24).replace(
|
||||
'.oss-internal.aliyun-inc.',
|
||||
'.oss.aliyuncs.').replace(
|
||||
'-internal', '')
|
||||
except:
|
||||
url = FS.get_url(
|
||||
one_path, lifecycle=3600 * 365 * 24).replace(
|
||||
'.oss-internal.aliyun-inc.',
|
||||
'.oss.aliyuncs.').replace('-internal', '')
|
||||
except Exception:
|
||||
url = one_path
|
||||
if len(one_label) > 10 and idx == len(save_path) - 1:
|
||||
is_textarea = True
|
||||
if self.is_video and one_path.endswith(self.video_postfix):
|
||||
if self.is_video and one_path.endswith(
|
||||
self.video_postfix):
|
||||
one_rank += f'<td align="center"><video height="{height}" controls="">'
|
||||
one_rank += f'<source src="{url}" type="video/mp4"></video>'
|
||||
if idx == len(save_path) - 1 and is_textarea:
|
||||
@@ -467,16 +505,27 @@ class ProbeData():
|
||||
def distribute(self):
|
||||
return self._distribute_dict
|
||||
|
||||
def save_one_media(self, idx, v, prefix_path, image_postfix, video_postfix, rank = 0):
|
||||
def save_one_media(self,
|
||||
idx,
|
||||
v,
|
||||
prefix_path,
|
||||
image_postfix,
|
||||
video_postfix,
|
||||
rank=0):
|
||||
ret_label = None
|
||||
if self.is_image:
|
||||
ret_medias = self.save_image(prefix_path, v,
|
||||
image_postfix, rank = rank)
|
||||
ret_medias = self.save_image(prefix_path,
|
||||
v,
|
||||
image_postfix,
|
||||
rank=rank)
|
||||
elif self.is_video:
|
||||
ret_medias = self.save_video(prefix_path, v,
|
||||
video_postfix, fps=self.fps, rank = rank)
|
||||
ret_medias = self.save_video(prefix_path,
|
||||
v,
|
||||
video_postfix,
|
||||
fps=self.fps,
|
||||
rank=rank)
|
||||
else:
|
||||
ret_data = self.save_npy(prefix_path, v, rank = rank)
|
||||
ret_data = self.save_npy(prefix_path, v, rank=rank)
|
||||
return ret_data, ret_label
|
||||
ret_data = ret_medias if isinstance(ret_medias, list) else [ret_medias]
|
||||
if self.build_html:
|
||||
@@ -484,14 +533,10 @@ class ProbeData():
|
||||
if isinstance(self.build_label, str):
|
||||
ret_label = [self.build_label for _ in ret_medias]
|
||||
elif isinstance(self.build_label[idx], list):
|
||||
assert len(self.build_label[idx]) == len(
|
||||
ret_medias)
|
||||
assert len(self.build_label[idx]) == len(ret_medias)
|
||||
ret_label = self.build_label[idx]
|
||||
else:
|
||||
ret_label = [
|
||||
self.build_label[idx]
|
||||
for _ in ret_medias
|
||||
]
|
||||
ret_label = [self.build_label[idx] for _ in ret_medias]
|
||||
else:
|
||||
if isinstance(self.build_label, str):
|
||||
ret_label = [self.build_label]
|
||||
@@ -499,7 +544,11 @@ class ProbeData():
|
||||
ret_label = [self.build_label[idx]]
|
||||
return ret_data, ret_label
|
||||
|
||||
def presave(self, prefix=None, image_postfix='jpg', video_postfix='mp4', rank = 0):
|
||||
def presave(self,
|
||||
prefix=None,
|
||||
image_postfix='jpg',
|
||||
video_postfix='mp4',
|
||||
rank=0):
|
||||
self.image_postfix = image_postfix
|
||||
self.video_postfix = video_postfix
|
||||
if isinstance(self.data, np.ndarray):
|
||||
@@ -507,11 +556,18 @@ class ProbeData():
|
||||
raise 'You should provide the save prefix for array sample.'
|
||||
# save jpg
|
||||
if self.is_image:
|
||||
ret_data = self.save_image(prefix, self.data, image_postfix, rank = rank)
|
||||
ret_data = self.save_image(prefix,
|
||||
self.data,
|
||||
image_postfix,
|
||||
rank=rank)
|
||||
elif self.is_video:
|
||||
ret_data = self.save_video(prefix, self.data, video_postfix, fps=self.fps, rank = rank)
|
||||
ret_data = self.save_video(prefix,
|
||||
self.data,
|
||||
video_postfix,
|
||||
fps=self.fps,
|
||||
rank=rank)
|
||||
else:
|
||||
ret_data = self.save_npy(prefix, self.data, rank = rank)
|
||||
ret_data = self.save_npy(prefix, self.data, rank=rank)
|
||||
self.media_handler.sync()
|
||||
self.media_handler.clear()
|
||||
if isinstance(ret_data, list):
|
||||
@@ -525,7 +581,7 @@ class ProbeData():
|
||||
ret_label.append(self.build_label)
|
||||
if not len(ret_data[0]) == len(ret_label[0]):
|
||||
raise f"The {prefix} label's length should be equal with 1st dim {self.data.shape[0]}."
|
||||
self.data = {"ret_data": ret_data, "ret_label": ret_label}
|
||||
self.data = {'ret_data': ret_data, 'ret_label': ret_label}
|
||||
else:
|
||||
self.data = ret_data
|
||||
self.is_presave = True
|
||||
@@ -535,16 +591,20 @@ class ProbeData():
|
||||
ret_label = []
|
||||
for idx, v in enumerate(self.data):
|
||||
prefix_path = os.path.join(prefix, f'{idx}')
|
||||
ret_one_data, ret_one_label = self.save_one_media(idx, v,
|
||||
prefix_path,
|
||||
image_postfix,
|
||||
video_postfix,
|
||||
rank=rank)
|
||||
ret_one_data, ret_one_label = self.save_one_media(
|
||||
idx,
|
||||
v,
|
||||
prefix_path,
|
||||
image_postfix,
|
||||
video_postfix,
|
||||
rank=rank)
|
||||
ret_data.append(ret_one_data)
|
||||
ret_label.append(ret_one_label) if ret_one_label is not None else ret_label
|
||||
ret_label.append(
|
||||
ret_one_label
|
||||
) if ret_one_label is not None else ret_label
|
||||
self.media_handler.sync()
|
||||
self.media_handler.clear()
|
||||
self.data = {"ret_data": ret_data, "ret_label": ret_label}
|
||||
self.data = {'ret_data': ret_data, 'ret_label': ret_label}
|
||||
self.is_presave = True
|
||||
elif isinstance(self.data, dict):
|
||||
if not self.basic_type:
|
||||
@@ -552,31 +612,38 @@ class ProbeData():
|
||||
ret_label = []
|
||||
for k, v in self.data.items():
|
||||
prefix_path = os.path.join(prefix, f'{k}_')
|
||||
ret_one_data, ret_one_label = self.save_one_media(k, v,
|
||||
prefix_path,
|
||||
image_postfix,
|
||||
video_postfix,
|
||||
rank = rank)
|
||||
ret_one_data, ret_one_label = self.save_one_media(
|
||||
k,
|
||||
v,
|
||||
prefix_path,
|
||||
image_postfix,
|
||||
video_postfix,
|
||||
rank=rank)
|
||||
ret_data.append(ret_one_data)
|
||||
ret_label.append(ret_one_label) if ret_one_label is not None else ret_label
|
||||
ret_label.append(
|
||||
ret_one_label
|
||||
) if ret_one_label is not None else ret_label
|
||||
self.media_handler.sync()
|
||||
self.media_handler.clear()
|
||||
self.data = {"ret_data": ret_data, "ret_label": ret_label}
|
||||
self.data = {'ret_data': ret_data, 'ret_label': ret_label}
|
||||
self.is_presave = True
|
||||
|
||||
def to_log(self, prefix=None, image_postfix='jpg', video_postfix='mp4', rank = 0):
|
||||
def to_log(self,
|
||||
prefix=None,
|
||||
image_postfix='jpg',
|
||||
video_postfix='mp4',
|
||||
rank=0):
|
||||
if not self.is_presave:
|
||||
self.presave(prefix, image_postfix, video_postfix, rank = rank)
|
||||
self.presave(prefix, image_postfix, video_postfix, rank=rank)
|
||||
if not self.is_presave:
|
||||
return self.data
|
||||
if isinstance(self.data, str):
|
||||
return self.data
|
||||
elif isinstance(self.data, dict):
|
||||
ret_data, ret_label = self.data["ret_data"], self.data["ret_label"]
|
||||
ret_data, ret_label = self.data['ret_data'], self.data['ret_label']
|
||||
if self.build_html:
|
||||
html_prefix = prefix + '_probe.html'
|
||||
html_file = self.save_html(html_prefix, ret_data,
|
||||
ret_label)
|
||||
html_file = self.save_html(html_prefix, ret_data, ret_label)
|
||||
return {'ori_file': ret_data, 'html': html_file}
|
||||
else:
|
||||
return {'ori_file': ret_data}
|
||||
@@ -584,15 +651,15 @@ class ProbeData():
|
||||
ret_data, ret_label = [], []
|
||||
for one_data in self.data:
|
||||
if isinstance(one_data, dict):
|
||||
one_ret_data, one_ret_label = one_data["ret_data"], one_data["ret_label"]
|
||||
one_ret_data, one_ret_label = one_data[
|
||||
'ret_data'], one_data['ret_label']
|
||||
ret_data.extend(one_ret_data)
|
||||
ret_label.extend(one_ret_label)
|
||||
elif isinstance(one_data, str):
|
||||
ret_data.append(one_data)
|
||||
if (self.is_image or self.is_video) and self.build_html:
|
||||
html_prefix = prefix + '_probe.html'
|
||||
html_file = self.save_html(html_prefix, ret_data,
|
||||
ret_label)
|
||||
html_file = self.save_html(html_prefix, ret_data, ret_label)
|
||||
return {'ori_file': ret_data, 'html': html_file}
|
||||
else:
|
||||
return {'ori_file': ret_data}
|
||||
|
||||
@@ -1,5 +1,10 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
from enum import Enum
|
||||
import os
|
||||
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
|
||||
class Media(Enum):
|
||||
@@ -7,16 +12,21 @@ class Media(Enum):
|
||||
IMAGE = 2
|
||||
VIDEO = 3
|
||||
AUDIO = 4
|
||||
IMAGE_PAIR = 5
|
||||
VIDEO_PAIR = 6
|
||||
|
||||
|
||||
class HtmlVisualization(object):
|
||||
def __init__(
|
||||
self,
|
||||
allow_annotation=False,
|
||||
slice_size=1000,
|
||||
align='center',
|
||||
width_scale='60%',
|
||||
title='Visualization',
|
||||
self,
|
||||
allow_annotation=False,
|
||||
slice_size=1000,
|
||||
align='center',
|
||||
width_scale='60%',
|
||||
title='Visualization',
|
||||
height=600,
|
||||
width=None,
|
||||
text_cols=40
|
||||
):
|
||||
self.content_list = []
|
||||
self.rows_meta = []
|
||||
@@ -27,48 +37,111 @@ class HtmlVisualization(object):
|
||||
self.title = title
|
||||
self.html_start = '<html>'
|
||||
self.html_head = f'<head><meta charset="utf-8"><title>{title}</title></head>'
|
||||
self.html_style = '''
|
||||
<style>
|
||||
table {
|
||||
border-collapse: collapse;
|
||||
}
|
||||
td {
|
||||
width: "{width_scale}";
|
||||
align: "{align}";
|
||||
margin: 0px;
|
||||
border: 0px;
|
||||
padding: 0px;
|
||||
vertical-align: top;
|
||||
}
|
||||
video {
|
||||
margin: 0px;
|
||||
border: 0px solid #ccc;
|
||||
padding: 0px;
|
||||
}
|
||||
textarea {
|
||||
margin: 0px;
|
||||
border: 0px;
|
||||
padding: 0px;
|
||||
resize: none;
|
||||
border: 1px solid #ccc;
|
||||
}
|
||||
</style>
|
||||
<script>
|
||||
function adjustHeight() {
|
||||
const textareas = document.querySelectorAll('textarea');
|
||||
textareas.forEach(textarea => {
|
||||
const td = textarea.parentNode;
|
||||
const tdHeight = td.clientHeight;
|
||||
textarea.style.height = tdHeight + 'px';
|
||||
});
|
||||
}
|
||||
window.onload = adjustHeight;
|
||||
window.onresize = adjustHeight;
|
||||
</script>
|
||||
self.height = height if height is not None else "600"
|
||||
self.width = width if width is not None else "auto"
|
||||
self.text_cols = text_cols if text_cols is not None else "auto"
|
||||
self.html_style = ('''
|
||||
<style> \n
|
||||
.container {
|
||||
display: flex;
|
||||
position: relative; \n
|
||||
overflow: hidden; \n
|
||||
justify-content: center; \n
|
||||
align-items: center; \n
|
||||
width: 600; \n
|
||||
height: {pair_height};
|
||||
border: 2px solid #ccc; \n
|
||||
} \n
|
||||
.image {
|
||||
display: flex;
|
||||
position: absolute; \n
|
||||
width: 100%; \n
|
||||
height: 100%; \n
|
||||
transition: 0.4s ease; \n
|
||||
}\n
|
||||
.image img { \n
|
||||
width: 100%; \n
|
||||
height: 100%; \n
|
||||
object-fit: contain; \n
|
||||
} \n
|
||||
|
||||
.video { \n
|
||||
display:flex; \n
|
||||
position:absolute; \n
|
||||
width:100%; \n
|
||||
height:100%; \n
|
||||
transition:0.4s ease; \n
|
||||
object-fit:contain; \n
|
||||
} \n
|
||||
|
||||
.slider {
|
||||
position: absolute; \n
|
||||
cursor: ew-resize; \n
|
||||
height: 100%; \n
|
||||
background-color: rgba(255, 255, 255, 0.5); \n
|
||||
z-index: 10; \n
|
||||
} \n
|
||||
textarea { \n
|
||||
margin: 0px; \n
|
||||
border: 0px; \n
|
||||
padding: 0px; \n
|
||||
resize: none; \n
|
||||
border: 1px solid #ccc; \n
|
||||
} \n
|
||||
.large-checkbox {transform: scale(2.5); margin-left: 20px; margin-bottom: 20px; vertical-align: middle;} \n
|
||||
</style> \n
|
||||
\n
|
||||
'''.replace('{width_scale}',
|
||||
self.width_scale).replace('{align}', self.align)
|
||||
self.html_body = '<body>{BODY}</body>\n'
|
||||
.replace('{pair_height}', f'{self.height}'))
|
||||
|
||||
self.html_body_script = '''
|
||||
<script>\n
|
||||
const containers = document.querySelectorAll('.container'); \n
|
||||
containers.forEach(container => {\n
|
||||
let isDragging = true;\n
|
||||
const slider = container.querySelector('.slider')\n
|
||||
const media2 = container.querySelector('#media2')\n
|
||||
|
||||
container.addEventListener('mousedown', () => {\n
|
||||
isDragging = true;\n
|
||||
});\n
|
||||
|
||||
|
||||
container.addEventListener('mouseup', () => {\n
|
||||
isDragging = true;\n
|
||||
});\n
|
||||
|
||||
container.addEventListener('mousemove', (event) => {\n
|
||||
if (!isDragging) return;\n
|
||||
const { clientX } = event;\n
|
||||
|
||||
const { left, width } = container.getBoundingClientRect();\n
|
||||
|
||||
let percentage = (clientX - left) / width * 100;\n
|
||||
|
||||
|
||||
// 限制百分比在0到100之间\n
|
||||
|
||||
percentage = Math.max(0, Math.min(100, percentage));\n
|
||||
|
||||
media2.style.clipPath = `inset(0 ${100 - percentage}% 0 0)`;\n
|
||||
|
||||
slider.style.left = `${percentage}%`;\n
|
||||
|
||||
console.info(slider.style.left);\n
|
||||
|
||||
});\n
|
||||
// 初始化滑块位置\n
|
||||
slider.style.left = '50%';\n
|
||||
});\n
|
||||
</script>\n
|
||||
'''
|
||||
|
||||
self.html_body = '<body>{BODY}\n' + self.html_body_script + '</body>\n'
|
||||
|
||||
self.html_end = '</html>'
|
||||
|
||||
self.html_script = '''
|
||||
<script>
|
||||
function saveSamples() {
|
||||
@@ -89,93 +162,139 @@ class HtmlVisualization(object):
|
||||
a.click();
|
||||
}
|
||||
</script>
|
||||
|
||||
'''
|
||||
self.label_button = (
|
||||
'<table><tr><td>' +
|
||||
"<button style='height: 50px;' type=\"button\" onclick=\"saveSamples()\">Save Samples</button>"
|
||||
+ '</td></tr></table>')
|
||||
'<table><tr><td>' +
|
||||
"<button style='height: 50px;' type=\"button\" onclick=\"saveSamples()\">Save Samples</button>"
|
||||
+ '</td></tr></table>')
|
||||
|
||||
def format_col(self,
|
||||
content='',
|
||||
label='',
|
||||
type=Media.TEXT,
|
||||
content_height=400,
|
||||
content_width=600):
|
||||
show_label=True,
|
||||
cols_span=1
|
||||
):
|
||||
if type == Media.TEXT:
|
||||
ret_str = '<td><textarea' # noqa: E501
|
||||
# if content_height is not None:
|
||||
# rows = f"rows={content_height//30}"
|
||||
# ret_str += f" {rows}"
|
||||
if content_width is not None:
|
||||
cols = f"cols={content_width//15}"
|
||||
ret_str = '<textarea' # noqa: E501
|
||||
if self.height is not None:
|
||||
rows = f"rows={self.height // 30}"
|
||||
ret_str += f" {rows}"
|
||||
if self.width is not None:
|
||||
cols = f"cols={self.text_cols * cols_span}"
|
||||
ret_str += f" {cols}"
|
||||
ret_str += f'>"{content}"</textarea></td>\n'
|
||||
sec_ret_str = f'<td align="center"><font size="3"><strong>{label}<strong></font></td>\n'
|
||||
return [ret_str, sec_ret_str]
|
||||
ret_str += f'>"{content}"</textarea>'
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
elif type == Media.IMAGE:
|
||||
ret_str = f'<td><img src="{content}"'
|
||||
if content_height is not None:
|
||||
height = f'height="{content_height}"'
|
||||
ret_str = f'<img src="{content}"'
|
||||
if self.height is not None:
|
||||
height = f'height="{self.height}"'
|
||||
ret_str += f" {height}"
|
||||
if content_width is not None:
|
||||
width = f'width="{content_width}"'
|
||||
if self.width is not None:
|
||||
width = f'width="{self.width}"'
|
||||
ret_str += f" {width}"
|
||||
ret_str += ' ></td>\n'
|
||||
sec_ret_str = f'<td align="center"><font size="3"><strong>{label}<strong></font></td>\n'
|
||||
return [ret_str, sec_ret_str]
|
||||
ret_str += ' >'
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
elif type == Media.VIDEO:
|
||||
ret_str = '<td><video' # noqa
|
||||
if content_height is not None:
|
||||
height = f'height="{content_height}"'
|
||||
ret_str = '<video' # noqa
|
||||
if self.height is not None:
|
||||
height = f'height="{self.width}"'
|
||||
ret_str += f" {height}"
|
||||
if content_width is not None:
|
||||
width = f'width="{content_width}"'
|
||||
if self.width is not None:
|
||||
width = f'width="{self.width}"'
|
||||
ret_str += f" {width}"
|
||||
ret_str += ' controls>'
|
||||
ret_str += f'<source src="{content}" type="video/mp4"></video></td>\n'
|
||||
sec_ret_str = f'<td align="center"><font size="3"><strong>{label}<strong></font></td>\n'
|
||||
return [ret_str, sec_ret_str]
|
||||
ret_str += ' preload="none" autoplay muted loop>'
|
||||
ret_str += f'<source src="{content}" type="video/mp4"></video>'
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
elif type == Media.AUDIO:
|
||||
ret_str = f'<td><audio src="{content}" controls></td>\n'
|
||||
sec_ret_str = f'<td align="center"><font size="3"><strong>{label}<strong></font></td>\n'
|
||||
return [ret_str, sec_ret_str]
|
||||
ret_str = f'<audio src="{content}" controls>'
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
elif type == Media.IMAGE_PAIR:
|
||||
assert isinstance(content, (list, tuple)) and len(content) == 2
|
||||
ret_str = f'\n'
|
||||
ret_str += f' <div class="container"'
|
||||
ret_str += (f'> \n'
|
||||
f' <div class="image" id="media1">'
|
||||
f' <img src="{content[1]}" alt="before">\n'
|
||||
f' </div>\n'
|
||||
f' <div class="image" id="media2" style="clip-path: inset(0 50% 0 0);">\n'
|
||||
f' <img src="{content[0]}" alt="after">\n'
|
||||
f' </div>\n'
|
||||
f' <div class="slider" id="slider"></div>\n'
|
||||
f'')
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
elif type == Media.VIDEO_PAIR:
|
||||
assert isinstance(content, (list, tuple)) and len(content) == 2
|
||||
ret_str = f'\n'
|
||||
ret_str += f' <div class="container"'
|
||||
ret_str += (f'> \n'
|
||||
f' <video autoplay muted loop class="video" id="media1"><source src="{content[1]}" type="video/mp4"></video>\n'
|
||||
f' <video autoplay muted loop class="video" id="media2" style="clip-path: inset(0 50% 0 0);"><source src="{content[0]}" type="video/mp4"></video>\n'
|
||||
f' <div class="slider" id="slider"></div>\n'
|
||||
f'</div>')
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
def format_row(self):
|
||||
sample_id = 0
|
||||
all_sample_html = []
|
||||
for one_content, one_row_meta in zip(self.content_list,
|
||||
self.rows_meta):
|
||||
one_row_str = '<table><tr>'
|
||||
one_row_str += '\n'.join([v[0] for v in one_content])
|
||||
if self.allow_annotation:
|
||||
row_meta = '#;#'.join(one_row_meta)
|
||||
one_row_str += (
|
||||
f'<td><input type="checkbox" class="large-checkbox" '
|
||||
f'id="sample{sample_id}" name="sample[]" value="{row_meta}"></td>\n'
|
||||
)
|
||||
one_row_str += '</tr><tr>'
|
||||
one_row_str += '\n'.join([v[1] for v in one_content]) # noqa
|
||||
if self.allow_annotation: # noqa
|
||||
one_row_str += f'<td></td>\n' # noqa
|
||||
one_row_str += '</tr></table>'
|
||||
if self.allow_annotation:
|
||||
one_row_str = f'<label for="sample{sample_id}">{one_row_str}</label>'
|
||||
all_sample_html.append(one_row_str)
|
||||
sample_id += 1
|
||||
if self.allow_annotation:
|
||||
ret_str = f'<label for="sample#sample_id#">{ret_str}</label>'
|
||||
|
||||
return '\n'.join(all_sample_html)
|
||||
if cols_span > 1:
|
||||
ret_str = f'<th colspan="{cols_span}">{ret_str}</th>\n'
|
||||
sec_ret_str = f'<th colspan="{cols_span}">{sec_ret_str}</th>\n' if not sec_ret_str == "" else sec_ret_str
|
||||
else:
|
||||
ret_str = f'<td>{ret_str}</td>\n'
|
||||
sec_ret_str = f'<td align="center">{sec_ret_str}</td>\n' if not sec_ret_str == "" else sec_ret_str
|
||||
return [ret_str, sec_ret_str]
|
||||
|
||||
def format_row(self):
|
||||
|
||||
all_sample_html = []
|
||||
current_content_list = copy.deepcopy(self.content_list)
|
||||
current_rows_meta = copy.deepcopy(self.rows_meta)
|
||||
|
||||
while len(current_content_list) > 0:
|
||||
sample_id = 0
|
||||
batch_content_list = current_content_list[:self.slice_size]
|
||||
current_content_list = current_content_list[self.slice_size:]
|
||||
batch_rows_meta = current_rows_meta[:self.slice_size]
|
||||
current_rows_meta = current_rows_meta[self.slice_size:]
|
||||
current_sample_html = []
|
||||
for one_content, one_row_meta in zip(batch_content_list,
|
||||
batch_rows_meta):
|
||||
one_row_str = '<tr>'
|
||||
if not self.allow_annotation:
|
||||
one_row_str += '\n'.join([v[0] for v in one_content])
|
||||
else:
|
||||
one_row_str += '\n'.join([v[0].replace('#sample_id#', f'{sample_id}') for v in one_content])
|
||||
row_meta = '#;#'.join(one_row_meta)
|
||||
one_row_str += (
|
||||
f'<td><input type="checkbox" class="large-checkbox" '
|
||||
f'id="sample{sample_id}" name="sample[]" value="{row_meta}"></td>\n'
|
||||
)
|
||||
one_row_str += '</tr><tr>'
|
||||
one_row_str += '\n'.join([v[1] for v in one_content]) # noqa
|
||||
if self.allow_annotation: # noqa
|
||||
one_row_str += f'<td></td>\n' # noqa
|
||||
one_row_str += '</tr>'
|
||||
# if self.allow_annotation:
|
||||
# one_row_str = f'<label for="sample{sample_id}">{one_row_str}</label>'
|
||||
current_sample_html.append(one_row_str)
|
||||
sample_id += 1
|
||||
all_sample_html.append("<table>" + '\n'.join(current_sample_html) + "</table>")
|
||||
|
||||
return all_sample_html
|
||||
|
||||
def add_record(self,
|
||||
content='',
|
||||
content,
|
||||
label='',
|
||||
type=Media.TEXT,
|
||||
row_id=1,
|
||||
col_id=1,
|
||||
cols_span=1,
|
||||
annotation_meta=None,
|
||||
content_height=None,
|
||||
content_width=None):
|
||||
show_label=True):
|
||||
if row_id >= len(self.content_list):
|
||||
self.content_list.append([])
|
||||
self.rows_meta.append([])
|
||||
@@ -186,7 +305,7 @@ class HtmlVisualization(object):
|
||||
raise RuntimeError(
|
||||
'col_id should be next number of the last col_id.')
|
||||
format_col = self.format_col(content, f"{row_id}-{col_id}: {label}",
|
||||
type, content_height, content_width)
|
||||
type, show_label=show_label, cols_span=cols_span)
|
||||
|
||||
annotation_meta = annotation_meta if annotation_meta else ''
|
||||
if col_id == len(self.content_list[row_id]):
|
||||
@@ -198,60 +317,30 @@ class HtmlVisualization(object):
|
||||
|
||||
def save_html(self, path):
|
||||
html_body = self.format_row()
|
||||
ret_html_list = [
|
||||
self.html_start, self.html_head, self.html_style,
|
||||
self.html_body.replace('{BODY}', html_body)
|
||||
]
|
||||
if self.allow_annotation:
|
||||
ret_html_list.append(self.label_button)
|
||||
ret_html_list.append(self.html_script)
|
||||
ret_html_list.append(self.html_end)
|
||||
ret_html = '\n'.join(ret_html_list)
|
||||
with open(path, 'w') as f:
|
||||
f.write(ret_html)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.file_system import FS
|
||||
FS.init_fs_client(Config(cfg_dict={}, load=False))
|
||||
|
||||
image_content_oss = '0_probe_0_[1024_2048_3].jpg'
|
||||
content_oss = '6UTWGRG1lx08iRBx5REA01041200dzcb0E010.mp4'
|
||||
caption = 'a little girl says hello.'
|
||||
|
||||
html_ins = HtmlVisualization(allow_annotation=True,
|
||||
slice_size=1000,
|
||||
title='Visualization',
|
||||
width_scale='100%')
|
||||
|
||||
for i in range(4):
|
||||
content_url = FS.get_url(content_oss, skip_check=True)
|
||||
html_ins.add_record(content=content_url,
|
||||
label='caption',
|
||||
type=Media.VIDEO,
|
||||
row_id=i,
|
||||
col_id=0,
|
||||
annotation_meta=None,
|
||||
content_height=600,
|
||||
content_width=None)
|
||||
html_ins.add_record(content=caption,
|
||||
label='caption',
|
||||
type=Media.TEXT,
|
||||
row_id=i,
|
||||
col_id=1,
|
||||
annotation_meta=None,
|
||||
content_height=600,
|
||||
content_width=750)
|
||||
image_content_url = FS.get_url(image_content_oss, skip_check=True)
|
||||
html_ins.add_record(content=image_content_url,
|
||||
label='caption',
|
||||
type=Media.IMAGE,
|
||||
row_id=i,
|
||||
col_id=2,
|
||||
annotation_meta=None,
|
||||
content_height=600,
|
||||
content_width=None)
|
||||
|
||||
with FS.put_to('visualize.html') as local_path:
|
||||
html_ins.save_html(local_path)
|
||||
if isinstance(html_body, list) and len(html_body) > 1:
|
||||
try:
|
||||
os.makedirs(path, exist_ok=True)
|
||||
except:
|
||||
print("Create folder path failed.")
|
||||
for html_id, one_html in enumerate(html_body):
|
||||
ret_html_list = [
|
||||
self.html_start, self.html_head, self.html_style,
|
||||
self.html_body.replace('{BODY}', one_html)
|
||||
]
|
||||
if self.allow_annotation:
|
||||
ret_html_list.append(self.label_button)
|
||||
ret_html_list.append(self.html_script)
|
||||
ret_html_list.append(self.html_end)
|
||||
ret_html = '\n'.join(ret_html_list)
|
||||
FS.put_object(ret_html.encode(), os.path.join(path, f"{html_id}.html"))
|
||||
else:
|
||||
ret_html_list = [
|
||||
self.html_start, self.html_head, self.html_style,
|
||||
self.html_body.replace('{BODY}', html_body[0])
|
||||
]
|
||||
if self.allow_annotation:
|
||||
ret_html_list.append(self.label_button)
|
||||
ret_html_list.append(self.html_script)
|
||||
ret_html_list.append(self.html_end)
|
||||
ret_html = '\n'.join(ret_html_list)
|
||||
FS.put_object(ret_html.encode(), path)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,380 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import os
|
||||
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def download_image(image, local_path=None):
|
||||
if not FS.exists(local_path):
|
||||
local_path = FS.get_from(image, local_path=local_path)
|
||||
return local_path
|
||||
|
||||
def blank_image():
|
||||
return Image.new('RGBA', (128, 128), (0, 0, 0, 0))
|
||||
|
||||
|
||||
|
||||
def get_examples(cache_dir):
|
||||
print('Downloading Examples ...')
|
||||
bl_img = blank_image()
|
||||
examples = [
|
||||
[
|
||||
'Facial Editing',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/e33edc106953.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/e33edc106953.png')), bl_img,
|
||||
bl_img, '{image} let the man smile', 6666
|
||||
],
|
||||
[
|
||||
'Facial Editing',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/5d2bcc91a3e9.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/5d2bcc91a3e9.png')), bl_img,
|
||||
bl_img, 'let the man in {image} wear sunglasses', 9999
|
||||
],
|
||||
[
|
||||
'Facial Editing',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/5d2bcc91a3e9.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/5d2bcc91a3e9.png')), bl_img,
|
||||
bl_img, 'let the man in {image} wear sunglasses', 9999
|
||||
],
|
||||
[
|
||||
'Facial Editing',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/3a52eac708bd.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/3a52eac708bd.png')), bl_img,
|
||||
bl_img, '{image} red hair', 9999
|
||||
],
|
||||
[
|
||||
'Facial Editing',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/3f4dc464a0ea.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/3f4dc464a0ea.png')), bl_img,
|
||||
bl_img, '{image} let the man serious', 99999
|
||||
],
|
||||
[
|
||||
'Controllable Generation',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/131ca90fd2a9.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/131ca90fd2a9.png')), bl_img, bl_img,
|
||||
'"A person sits contemplatively on the ground, surrounded by falling autumn leaves. Dressed in a green sweater and dark blue pants, they rest their chin on their hand, exuding a relaxed demeanor. Their stylish checkered slip-on shoes add a touch of flair, while a black purse lies in their lap. The backdrop of muted brown enhances the warm, cozy atmosphere of the scene." , generate the image that corresponds to the given scribble {image}.',
|
||||
613725
|
||||
],
|
||||
[
|
||||
'Render Text',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/33e9f27c2c48.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/33e9f27c2c48.png')),
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/33e9f27c2c48_mask.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/33e9f27c2c48_mask.png')), bl_img,
|
||||
'Put the text "C A T" at the position marked by mask in the {image}',
|
||||
6666
|
||||
],
|
||||
[
|
||||
'Style Transfer',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/9e73e7eeef55.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/9e73e7eeef55.png')), bl_img,
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/2e02975293d6.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/2e02975293d6.png')),
|
||||
'edit {image} based on the style of {image1} ', 99999
|
||||
],
|
||||
[
|
||||
'Outpainting',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/f2b22c08be3f.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/f2b22c08be3f.png')),
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/f2b22c08be3f_mask.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/f2b22c08be3f_mask.png')), bl_img,
|
||||
'Could the {image} be widened within the space designated by mask, while retaining the original?',
|
||||
6666
|
||||
],
|
||||
[
|
||||
'Image Segmentation',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/db3ebaa81899.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/db3ebaa81899.png')), bl_img,
|
||||
bl_img, '{image} Segmentation', 6666
|
||||
],
|
||||
[
|
||||
'Depth Estimation',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/f1927c4692ba.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/f1927c4692ba.png')), bl_img,
|
||||
bl_img, '{image} Depth Estimation', 6666
|
||||
],
|
||||
[
|
||||
'Pose Estimation',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/014e5bf3b4d1.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/014e5bf3b4d1.png')), bl_img,
|
||||
bl_img, '{image} distinguish the poses of the figures', 999999
|
||||
],
|
||||
[
|
||||
'Scribble Extraction',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/5f59a202f8ac.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/5f59a202f8ac.png')), bl_img,
|
||||
bl_img, 'Generate a scribble of {image}, please.', 6666
|
||||
],
|
||||
[
|
||||
'Mosaic',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/3a2f52361eea.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/3a2f52361eea.png')), bl_img,
|
||||
bl_img, 'Adapt {image} into a mosaic representation.', 6666
|
||||
],
|
||||
[
|
||||
'Edge map Extraction',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/b9d1e519d6e5.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/b9d1e519d6e5.png')), bl_img,
|
||||
bl_img, 'Get the edge-enhanced result for {image}.', 6666
|
||||
],
|
||||
[
|
||||
'Grayscale',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/c4ebbe2ba29b.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/c4ebbe2ba29b.png')), bl_img,
|
||||
bl_img, 'transform {image} into a black and white one', 6666
|
||||
],
|
||||
[
|
||||
'Contour Extraction',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/19652d0f6c4b.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/19652d0f6c4b.png')), bl_img, bl_img,
|
||||
'Would you be able to make a contour picture from {image} for me?',
|
||||
6666
|
||||
],
|
||||
[
|
||||
'Controllable Generation',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/249cda2844b7.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/249cda2844b7.png')), bl_img, bl_img,
|
||||
'Following the segmentation outcome in mask of {image}, develop a real-life image using the explanatory note in "a mighty cat lying on the bed”.',
|
||||
6666
|
||||
],
|
||||
[
|
||||
'Controllable Generation',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/411f6c4b8e6c.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/411f6c4b8e6c.png')), bl_img, bl_img,
|
||||
'use the depth map {image} and the text caption "a cut white cat" to create a corresponding graphic image',
|
||||
999999
|
||||
],
|
||||
[
|
||||
'Controllable Generation',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/a35c96ed137a.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/a35c96ed137a.png')), bl_img, bl_img,
|
||||
'help translate this posture schema {image} into a colored image based on the context I provided "A beautiful woman Climbing the climbing wall, wearing a harness and climbing gear, skillfully maneuvering up the wall with her back to the camera, with a safety rope."',
|
||||
3599999
|
||||
],
|
||||
[
|
||||
'Controllable Generation',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/dcb2fc86f1ce.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/dcb2fc86f1ce.png')), bl_img, bl_img,
|
||||
'Transform and generate an image using mosaic {image} and "Monarch butterflies gracefully perch on vibrant purple flowers, showcasing their striking orange and black wings in a lush garden setting." description',
|
||||
6666
|
||||
],
|
||||
[
|
||||
'Controllable Generation',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/4cd4ee494962.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/4cd4ee494962.png')), bl_img, bl_img,
|
||||
'make this {image} colorful as per the "beautiful sunflowers"',
|
||||
6666
|
||||
],
|
||||
[
|
||||
'Controllable Generation',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/a47e3a9cd166.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/a47e3a9cd166.png')), bl_img, bl_img,
|
||||
'Take the edge conscious {image} and the written guideline "A whimsical animated character is depicted holding a delectable cake adorned with blue and white frosting and a drizzle of chocolate. The character wears a yellow headband with a bow, matching a cozy yellow sweater. Her dark hair is styled in a braid, tied with a yellow ribbon. With a golden fork in hand, she stands ready to enjoy a slice, exuding an air of joyful anticipation. The scene is creatively rendered with a charming and playful aesthetic." and produce a realistic image.',
|
||||
613725
|
||||
],
|
||||
[
|
||||
'Controllable Generation',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/d890ed8a3ac2.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/d890ed8a3ac2.png')), bl_img, bl_img,
|
||||
'creating a vivid image based on {image} and description "This image features a delicious rectangular tart with a flaky, golden-brown crust. The tart is topped with evenly sliced tomatoes, layered over a creamy cheese filling. Aromatic herbs are sprinkled on top, adding a touch of green and enhancing the visual appeal. The background includes a soft, textured fabric and scattered white flowers, creating an elegant and inviting presentation. Bright red tomatoes in the upper right corner hint at the fresh ingredients used in the dish."',
|
||||
6666
|
||||
],
|
||||
[
|
||||
'Image Denoising',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/0844a686a179.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/0844a686a179.png')), bl_img, bl_img,
|
||||
'Eliminate noise interference in {image} and maximize the crispness to obtain superior high-definition quality',
|
||||
6666
|
||||
],
|
||||
[
|
||||
'Inpainting',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/fa91b6b7e59b.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/fa91b6b7e59b.png')),
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/fa91b6b7e59b_mask.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/fa91b6b7e59b_mask.png')), bl_img,
|
||||
'Ensure to overhaul the parts of the {image} indicated by the mask.',
|
||||
6666
|
||||
],
|
||||
[
|
||||
'Inpainting',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/632899695b26.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/632899695b26.png')),
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/632899695b26_mask.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/632899695b26_mask.png')), bl_img,
|
||||
'Refashion the mask portion of {image} in accordance with "A yellow egg with a smiling face painted on it"',
|
||||
6666
|
||||
],
|
||||
[
|
||||
'General Editing',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/354d17594afe.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/354d17594afe.png')), bl_img, bl_img,
|
||||
'{image} change the dog\'s posture to walking in the water, and change the background to green plants and a pond.',
|
||||
6666
|
||||
],
|
||||
[
|
||||
'General Editing',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/38946455752b.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/38946455752b.png')), bl_img, bl_img,
|
||||
'{image} change the color of the dress from white to red and the model\'s hair color red brown to blonde.Other parts remain unchanged',
|
||||
6669
|
||||
],
|
||||
[
|
||||
'Facial Editing',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/3ba5202f0cd8.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/3ba5202f0cd8.png')), bl_img, bl_img,
|
||||
'Keep the same facial feature in @3ba5202f0cd8, change the woman\'s clothing from a Blue denim jacket to a white turtleneck sweater and adjust her posture so that she is supporting her chin with both hands. Other aspects, such as background, hairstyle, facial expression, etc, remain unchanged.',
|
||||
99999
|
||||
],
|
||||
[
|
||||
'Facial Editing',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/369365b94725.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/369365b94725.png')), bl_img,
|
||||
bl_img, '{image} Make her looking at the camera', 6666
|
||||
],
|
||||
[
|
||||
'Facial Editing',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/92751f2e4a0e.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/92751f2e4a0e.png')), bl_img,
|
||||
bl_img, '{image} Remove the smile from his face', 9899999
|
||||
],
|
||||
[
|
||||
'Remove Text',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/8530a6711b2e.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/8530a6711b2e.png')), bl_img,
|
||||
bl_img, 'Aim to remove any textual element in {image}', 6666
|
||||
],
|
||||
[
|
||||
'Remove Text',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/c4d7fb28f8f6.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/c4d7fb28f8f6.png')),
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/c4d7fb28f8f6_mask.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/c4d7fb28f8f6_mask.png')), bl_img,
|
||||
'Rub out any text found in the mask sector of the {image}.', 6666
|
||||
],
|
||||
[
|
||||
'Remove Object',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/e2f318fa5e5b.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/e2f318fa5e5b.png')), bl_img, bl_img,
|
||||
'Remove the unicorn in this {image}, ensuring a smooth edit.',
|
||||
99999
|
||||
],
|
||||
[
|
||||
'Remove Object',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/1ae96d8aca00.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/1ae96d8aca00.png')),
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/1ae96d8aca00_mask.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/1ae96d8aca00_mask.png')),
|
||||
bl_img, 'Discard the contents of the mask area from {image}.', 99999
|
||||
],
|
||||
[
|
||||
'Add Object',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/80289f48e511.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/80289f48e511.png')),
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/80289f48e511_mask.png?raw=true',
|
||||
os.path.join(cache_dir,
|
||||
'examples/80289f48e511_mask.png')), bl_img,
|
||||
'add a Hot Air Balloon into the {image}, per the mask', 613725
|
||||
],
|
||||
[
|
||||
'Style Transfer',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/d725cb2009e8.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/d725cb2009e8.png')), bl_img,
|
||||
bl_img, 'Change the style of {image} to colored pencil style', 99999
|
||||
],
|
||||
[
|
||||
'Style Transfer',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/e0f48b3fd010.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/e0f48b3fd010.png')), bl_img,
|
||||
bl_img, 'make {image} to Walt Disney Animation style', 99999
|
||||
],
|
||||
[
|
||||
'Try On',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/ee4ca60b8c96.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/ee4ca60b8c96.png')),
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/ee4ca60b8c96_mask.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/ee4ca60b8c96_mask.png')),
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/ebe825bbfe3c.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/ebe825bbfe3c.png')),
|
||||
'Change the cloth in {image} to the one in {image1}', 99999
|
||||
],
|
||||
[
|
||||
'Workflow',
|
||||
download_image(
|
||||
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/cb85353c004b.png?raw=true',
|
||||
os.path.join(cache_dir, 'examples/cb85353c004b.png')), bl_img,
|
||||
bl_img, '<workflow> ice cream {image}', 99999
|
||||
],
|
||||
]
|
||||
print('Finish. Start building UI ...')
|
||||
return examples
|
||||
@@ -0,0 +1,95 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import torch
|
||||
import torchvision.transforms as T
|
||||
from PIL import Image
|
||||
from torchvision.transforms.functional import InterpolationMode
|
||||
|
||||
IMAGENET_MEAN = (0.485, 0.456, 0.406)
|
||||
IMAGENET_STD = (0.229, 0.224, 0.225)
|
||||
|
||||
|
||||
def build_transform(input_size):
|
||||
MEAN, STD = IMAGENET_MEAN, IMAGENET_STD
|
||||
transform = T.Compose([
|
||||
T.Lambda(lambda img: img.convert('RGB') if img.mode != 'RGB' else img),
|
||||
T.Resize((input_size, input_size),
|
||||
interpolation=InterpolationMode.BICUBIC),
|
||||
T.ToTensor(),
|
||||
T.Normalize(mean=MEAN, std=STD)
|
||||
])
|
||||
return transform
|
||||
|
||||
|
||||
def find_closest_aspect_ratio(aspect_ratio, target_ratios, width, height,
|
||||
image_size):
|
||||
best_ratio_diff = float('inf')
|
||||
best_ratio = (1, 1)
|
||||
area = width * height
|
||||
for ratio in target_ratios:
|
||||
target_aspect_ratio = ratio[0] / ratio[1]
|
||||
ratio_diff = abs(aspect_ratio - target_aspect_ratio)
|
||||
if ratio_diff < best_ratio_diff:
|
||||
best_ratio_diff = ratio_diff
|
||||
best_ratio = ratio
|
||||
elif ratio_diff == best_ratio_diff:
|
||||
if area > 0.5 * image_size * image_size * ratio[0] * ratio[1]:
|
||||
best_ratio = ratio
|
||||
return best_ratio
|
||||
|
||||
|
||||
def dynamic_preprocess(image,
|
||||
min_num=1,
|
||||
max_num=12,
|
||||
image_size=448,
|
||||
use_thumbnail=False):
|
||||
orig_width, orig_height = image.size
|
||||
aspect_ratio = orig_width / orig_height
|
||||
|
||||
# calculate the existing image aspect ratio
|
||||
target_ratios = set((i, j) for n in range(min_num, max_num + 1)
|
||||
for i in range(1, n + 1) for j in range(1, n + 1)
|
||||
if i * j <= max_num and i * j >= min_num)
|
||||
target_ratios = sorted(target_ratios, key=lambda x: x[0] * x[1])
|
||||
|
||||
# find the closest aspect ratio to the target
|
||||
target_aspect_ratio = find_closest_aspect_ratio(aspect_ratio,
|
||||
target_ratios, orig_width,
|
||||
orig_height, image_size)
|
||||
|
||||
# calculate the target width and height
|
||||
target_width = image_size * target_aspect_ratio[0]
|
||||
target_height = image_size * target_aspect_ratio[1]
|
||||
blocks = target_aspect_ratio[0] * target_aspect_ratio[1]
|
||||
|
||||
# resize the image
|
||||
resized_img = image.resize((target_width, target_height))
|
||||
processed_images = []
|
||||
for i in range(blocks):
|
||||
box = ((i % (target_width // image_size)) * image_size,
|
||||
(i // (target_width // image_size)) * image_size,
|
||||
((i % (target_width // image_size)) + 1) * image_size,
|
||||
((i // (target_width // image_size)) + 1) * image_size)
|
||||
# split the image
|
||||
split_img = resized_img.crop(box)
|
||||
processed_images.append(split_img)
|
||||
assert len(processed_images) == blocks
|
||||
if use_thumbnail and len(processed_images) != 1:
|
||||
thumbnail_img = image.resize((image_size, image_size))
|
||||
processed_images.append(thumbnail_img)
|
||||
return processed_images
|
||||
|
||||
|
||||
def load_image(image_file, input_size=448, max_num=12):
|
||||
if isinstance(image_file, str):
|
||||
image = Image.open(image_file).convert('RGB')
|
||||
else:
|
||||
image = image_file
|
||||
transform = build_transform(input_size=input_size)
|
||||
images = dynamic_preprocess(image,
|
||||
image_size=input_size,
|
||||
use_thumbnail=True,
|
||||
max_num=max_num)
|
||||
pixel_values = [transform(image) for image in images]
|
||||
pixel_values = torch.stack(pixel_values)
|
||||
return pixel_values
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user