Merge pull request #6 from modelscope/v0.0.3_dev

V0.0.3 dev
This commit is contained in:
Zhen Han
2024-02-07 20:18:55 +08:00
committed by GitHub
99 changed files with 2982 additions and 801 deletions
+2
View File
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
+1 -1
View File
@@ -914,7 +914,7 @@ data = {
_model(data)
probe = _model.probe_data()
for key in probe:
print(key, probe[key].to_log(prefix=f"xxx/dev_easytorch/{key}"))
print(key, probe[key].to_log(prefix=f"xxx/{key}"))
```
<hr/>
+48 -18
View File
@@ -14,6 +14,7 @@
- [Installation](#-Installation)
- [Getting Started](#-getting-started)
- [SCEPTER Studio](#-scepter-studio)
- [Gallery](#-gallery)
- [Features](#-features)
- [Learn More](#-learn-more)
- [License](#license)
@@ -43,6 +44,7 @@ Currently supported approaches (and counting):
3. Res-Tuning(TODO): [Res-Tuning: A Flexible and Efficient Tuning Paradigm via Unbinding Tuner from Backbone](https://arxiv.org/abs/2310.19859) [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=ResTuning&color=red&logo=arxiv)](https://arxiv.org/abs/2310.19859) [![Page link](https://img.shields.io/badge/Page-ResTuning-Gree)](https://res-tuning.github.io/)
## 🎉 News
- [2024.02]: We release new SCEdit controllable image synthesis models for SD v2.1 and SD XL. Multiple strategies applied to accelerate inference time for SCEPTER Studio.
- [2024.01]: We release **SCEPTER Studio**, an integrated toolkit for data management, model training and inference based on [Gradio](https://www.gradio.app/).
- [2024.01]: [SCEdit](https://arxiv.org/abs/2312.11392) support controllable image synthesis for training and inference.
- [2023.12]: We propose [SCEdit](https://arxiv.org/abs/2312.11392), an efficient and controllable generation framework.
@@ -56,13 +58,17 @@ Currently supported approaches (and counting):
conda env create -f environment.yaml
conda activate scepter
```
- We recommend installing the specific version of PyTorch and accelerate toolbox [xFormers](https://pypi.org/project/xformers/). You can install these recommended version by pip:
```shell
pip install -r requirements/recommended.txt
```
- Install SCEPTER by the `pip` command:
```shell
pip install scepter
```
- PS: We recommend installing PyTorch follwing [official documentation](https://pytorch.org/get-started/locally/)
## 🚀 Getting Started
@@ -159,7 +165,6 @@ python scepter/tools/run_inference.py --cfg scepter/methods/scedit/ctr/sd21_768_
python scepter/tools/run_inference.py --cfg scepter/methods/scedit/ctr/sd21_768_sce_ctr_pose.yaml --num_samples 1 --prompt 'super mario' --save_folder 'test_mario_pose' --image_size 768 --task control --image 'asset/images/pose_source.png' --control_mode source --pretrained_model ms://damo/scepter_scedit@controllable_model/SD2.1/pose_control/0_SwiftSCETuning/pytorch_model.bin # pose
```
## 🖥️ SCEPTER Studio
### Launch
@@ -167,24 +172,49 @@ python scepter/tools/run_inference.py --cfg scepter/methods/scedit/ctr/sd21_768_
To fully experience **SCEPTER Studio**, you can launch the following command line:
```shell
pip install scepter
python -m scepter.tools.webui
```
or run after clone repo code
```shell
git clone https://github.com/modelscope/scepter.git
PYTHONPATH=. python scepter/tools/webui.py --cfg scepter/methods/studio/scepter_ui.yaml
```
The startup of **SCEPTER Studio** eliminates the need for manual downloading and organizing of models; it will automatically load the corresponding models and store them in a local directory.
Depending on the network and hardware situation, the initial startup usually requires 15-60 minutes, primarily involving the download and processing of SDv1.5, SDv2.1, and SDXL models.
Therefore, subsequent startups will become much faster (about one minute) as downloading is no longer required.
### Modelscope Studio
We deploy a work studio on Modelscope that includes only the inference tab, please refer to [ms_scepter_studio](https://www.modelscope.cn/studios/damo/scepter_studio/summary)
## 🖼️ Gallery
### Dragon Year Special: Dragon Tuner
<table>
<tr>
<td><strong>Gold Dragon Tuner</strong></td>
<td><strong>Sloppy Dragon Tuner</strong></td>
<td><strong>Red Dragon Tuner</strong><br> + Papercraft Mantra</td>
<td><strong>Azure Dragon Tuner</strong><br> + Pose Control</td>
</tr>
<tr>
<td><img src="https://github.com/hanzhn/datas/blob/main/scepter/readme/tuner_gold_dragon.jpeg?raw=true" width="300"></td>
<td><img src="https://github.com/hanzhn/datas/blob/main/scepter/readme/tuner_sloppy_dragon.jpeg?raw=true" width="300"></td>
<td><img src="https://github.com/hanzhn/datas/blob/main/scepter/readme/tuner_mantra_papercraft_dragon.jpeg?raw=true" width="300"></td>
<td><img src="https://github.com/hanzhn/datas/blob/main/scepter/readme/tuner_pose.jpeg?raw=true" width="300"></td>
</tr>
</table>
### Text Effect Image
<table>
<tr>
<td><strong>Conditional Image</strong></td>
<td><strong>Midas Control</strong><br>"Race track, top view"</td>
<td><strong>Midas Control</strong><br> + Watercolor Mantra<br>"white lilies"</td>
<td><strong>Midas Control</strong><br> + Dragon Tuner<br>"Spring Festival, Chinese dragon"</td>
</tr>
<tr>
<td><img src="https://github.com/hanzhn/datas/blob/main/scepter/readme/word_condition.png?raw=true" width="300"></td>
<td><img src="https://github.com/hanzhn/datas/blob/main/scepter/readme/word_race.jpeg?raw=true" width="300"></td>
<td><img src="https://github.com/hanzhn/datas/blob/main/scepter/readme/word_lilies.jpeg?raw=true" width="300"></td>
<td><img src="https://github.com/hanzhn/datas/blob/main/scepter/readme/word_festival.jpeg?raw=true" width="300"></td>
</tr>
</table>
## ✨ Features
### Text-to-Image Generation
@@ -201,8 +231,8 @@ We deploy a work studio on Modelscope that includes only the inference tab, plea
| **Model** | **Canny** | **HED** | **Depth** | **Pose** | **Color** |
|:---------:|:---------:|:-------:|:---------:|:--------:|:---------:|
| SD 1.5 | ✅ | ✅ | ✅ | ✅ | ✅ |
| SD 2.1 | 🪄 | ✅ | ✅ | 🪄 | 🪄 |
| SD XL | ✅ | ✅ | ✅ | ✅ | ✅ |
| SD 2.1 | 🪄 | 🪄 | 🪄 | 🪄 | 🪄 |
| SD XL | 🪄 | 🪄 | 🪄 | 🪄 | 🪄 |
### Model URL
@@ -210,9 +240,9 @@ We deploy a work studio on Modelscope that includes only the inference tab, plea
- 🪄 denotes that the model has been published.
- More models will be released in the future.
| Model | URL |
|--------|-------------------------------------------------------------------------------------|
| SCEdit | [ModelCard](https://modelscope.cn/models/damo/scepter_scedit/summary) |
| Model | URL |
|--------|------------------------------------------------------------------------------------------------------------------------------------------------|
| SCEdit | [ModelScope](https://modelscope.cn/models/damo/scepter_scedit/summary) [HuggingFace](https://huggingface.co/scepter-studio/scepter_scedit) |
PS: Scripts running within the SCEPTER framework will automatically fetch and load models based on the required dependency files, eliminating the need for manual downloads.
+2 -1
View File
@@ -7,5 +7,6 @@ opencv-python
opencv_transforms>=0.0.6
oss2>=2.15.0
pyyaml>=5.3.1
scikit-image
torchsde
transformers
xformers>=0.0.21
+3
View File
@@ -0,0 +1,3 @@
torch==2.0.1
torchvision==0.15.2
xformers==0.0.21
@@ -117,7 +117,7 @@ SOLVER:
SAMPLE_STEPS: 50
SEED: 2023
GUIDE_SCALE: 7.5
GUIDE_RESCALE:
GUIDE_RESCALE: 0.5
DISCRETIZATION: trailing
IMAGE_SIZE: [512, 512]
RUN_TRAIN_N: False
@@ -125,7 +125,7 @@ SOLVER:
SAMPLE_STEPS: 50
SEED: 2023
GUIDE_SCALE: 7.5
GUIDE_RESCALE:
GUIDE_RESCALE: 0.5
DISCRETIZATION: trailing
IMAGE_SIZE: [512, 512]
RUN_TRAIN_N: False
@@ -0,0 +1,214 @@
ENV:
BACKEND: nccl
SOLVER:
NAME: LatentDiffusionSolver
RESUME_FROM:
LOAD_MODEL_ONLY: True
USE_FSDP: False
SHARDING_STRATEGY:
USE_AMP: True
DTYPE: float16
CHANNELS_LAST: True
MAX_STEPS: 2000
MAX_EPOCHS: -1
NUM_FOLDS: 1
ACCU_STEP: 1
EVAL_INTERVAL: 100
#
WORK_DIR: ./cache/save_data/sd21_512_full
LOG_FILE: std_log.txt
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/data"
#
MODEL:
NAME: LatentDiffusion
PARAMETERIZATION: eps
TIMESTEPS: 1000
MIN_SNR_GAMMA:
ZERO_TERMINAL_SNR: False
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-2-1-base@v2-1_512-ema-pruned.safetensors
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
DEFAULT_N_PROMPT:
SCHEDULE_ARGS:
"NAME": "scaled_linear"
"BETA_MIN": 0.00085
"BETA_MAX": 0.012
USE_EMA: False
#
DIFFUSION_MODEL:
NAME: DiffusionUNet
IN_CHANNELS: 4
OUT_CHANNELS: 4
MODEL_CHANNELS: 320
NUM_HEADS_CHANNELS: 64
NUM_RES_BLOCKS: 2
ATTENTION_RESOLUTIONS: [ 4, 2, 1 ]
CHANNEL_MULT: [ 1, 2, 4, 4 ]
CONV_RESAMPLE: True
DIMS: 2
USE_CHECKPOINT: False
USE_SCALE_SHIFT_NORM: False
RESBLOCK_UPDOWN: False
USE_SPATIAL_TRANSFORMER: True
TRANSFORMER_DEPTH: 1
CONTEXT_DIM: 1024
DISABLE_MIDDLE_SELF_ATTN: False
USE_LINEAR_IN_TRANSFORMER: True
PRETRAINED_MODEL:
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
BATCH_SIZE: 4
#
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
#
TOKENIZER:
NAME: OpenClipTokenizer
LENGTH: 77
#
COND_STAGE_MODEL:
NAME: FrozenOpenCLIPEmbedder
ARCH: ViT-H-14
PRETRAINED_MODEL:
LAYER: penultimate
#
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 2023
GUIDE_SCALE: 7.5
GUIDE_RESCALE:
DISCRETIZATION: trailing
IMAGE_SIZE: [512, 512]
RUN_TRAIN_N: False
#
OPTIMIZER:
NAME: AdamW
LEARNING_RATE: 0.0064
BETAS: [ 0.9, 0.999 ]
EPS: 1e-8
WEIGHT_DECAY: 1e-2
AMSGRAD: False
#
TRAIN_DATA:
NAME: ImageTextPairMSDataset
MODE: train
MS_DATASET_NAME: style_custom_dataset
MS_DATASET_NAMESPACE: damo
MS_DATASET_SUBNAME: 3D
PROMPT_PREFIX: ""
MS_DATASET_SPLIT: train_short
MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
REPLACE_STYLE: False
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
SAMPLER:
NAME: LoopSampler
TRANSFORMS:
- NAME: LoadImageFromFile
RGB_ORDER: RGB
BACKEND: pillow
- NAME: Resize
SIZE: 512
INTERPOLATION: bilinear
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: CenterCrop
SIZE: 512
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: ImageToTensor
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: Normalize
MEAN: [ 0.5, 0.5, 0.5 ]
STD: [ 0.5, 0.5, 0.5 ]
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'image' ]
BACKEND: torchvision
- NAME: Select
KEYS: [ 'image', 'prompt' ]
META_KEYS: [ 'data_key' ]
#
EVAL_DATA:
NAME: ImageTextPairMSDataset
MODE: eval
MS_DATASET_NAME: style_custom_dataset
MS_DATASET_NAMESPACE: damo
MS_DATASET_SUBNAME: 3D
PROMPT_PREFIX: ""
MS_REMAP_KEYS: { 'Image': 'Target:FILE' }
MS_DATASET_SPLIT: test_short
OUTPUT_SIZE: [512, 512]
REPLACE_STYLE: False
PIN_MEMORY: True
BATCH_SIZE: 4
NUM_WORKERS: 4
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/data"
#
TRANSFORMS:
-
NAME: Select
KEYS: ['prompt']
META_KEYS: ['image_size']
#
TRAIN_HOOKS:
-
NAME: BackwardHook
PRIORITY: 0
-
NAME: LogHook
LOG_INTERVAL: 50
-
NAME: CheckpointHook
INTERVAL: 1000
-
NAME: ProbeDataHook
PROB_INTERVAL: 100
#
EVAL_HOOKS:
-
NAME: ProbeDataHook
PROB_INTERVAL: 100
@@ -0,0 +1,223 @@
ENV:
BACKEND: nccl
SOLVER:
NAME: LatentDiffusionSolver
RESUME_FROM:
LOAD_MODEL_ONLY: True
USE_FSDP: False
SHARDING_STRATEGY:
USE_AMP: True
DTYPE: float16
CHANNELS_LAST: True
MAX_STEPS: 2000
MAX_EPOCHS: -1
NUM_FOLDS: 1
ACCU_STEP: 1
EVAL_INTERVAL: 100
#
WORK_DIR: ./cache/save_data/sd21_512_lora
LOG_FILE: std_log.txt
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/data"
#
TUNER:
-
NAME: SwiftLoRA
R: 64
LORA_ALPHA: 64
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: model.*(to_q|to_k|to_v|to_out.0|net.0.proj|net.2)$
#
MODEL:
NAME: LatentDiffusion
PARAMETERIZATION: eps
TIMESTEPS: 1000
MIN_SNR_GAMMA:
ZERO_TERMINAL_SNR: False
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-2-1-base@v2-1_512-ema-pruned.safetensors
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
DEFAULT_N_PROMPT:
SCHEDULE_ARGS:
"NAME": "scaled_linear"
"BETA_MIN": 0.00085
"BETA_MAX": 0.012
USE_EMA: False
#
DIFFUSION_MODEL:
NAME: DiffusionUNet
IN_CHANNELS: 4
OUT_CHANNELS: 4
MODEL_CHANNELS: 320
NUM_HEADS_CHANNELS: 64
NUM_RES_BLOCKS: 2
ATTENTION_RESOLUTIONS: [ 4, 2, 1 ]
CHANNEL_MULT: [ 1, 2, 4, 4 ]
CONV_RESAMPLE: True
DIMS: 2
USE_CHECKPOINT: False
USE_SCALE_SHIFT_NORM: False
RESBLOCK_UPDOWN: False
USE_SPATIAL_TRANSFORMER: True
TRANSFORMER_DEPTH: 1
CONTEXT_DIM: 1024
DISABLE_MIDDLE_SELF_ATTN: False
USE_LINEAR_IN_TRANSFORMER: True
PRETRAINED_MODEL:
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
BATCH_SIZE: 4
#
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
#
TOKENIZER:
NAME: OpenClipTokenizer
LENGTH: 77
#
COND_STAGE_MODEL:
NAME: FrozenOpenCLIPEmbedder
ARCH: ViT-H-14
PRETRAINED_MODEL:
LAYER: penultimate
#
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 2023
GUIDE_SCALE: 7.5
GUIDE_RESCALE:
DISCRETIZATION: trailing
IMAGE_SIZE: [512, 512]
RUN_TRAIN_N: False
#
OPTIMIZER:
NAME: AdamW
LEARNING_RATE: 0.0064
BETAS: [ 0.9, 0.999 ]
EPS: 1e-8
WEIGHT_DECAY: 1e-2
AMSGRAD: False
#
TRAIN_DATA:
NAME: ImageTextPairMSDataset
MODE: train
MS_DATASET_NAME: style_custom_dataset
MS_DATASET_NAMESPACE: damo
MS_DATASET_SUBNAME: 3D
PROMPT_PREFIX: ""
MS_DATASET_SPLIT: train_short
MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
REPLACE_STYLE: False
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
SAMPLER:
NAME: LoopSampler
TRANSFORMS:
- NAME: LoadImageFromFile
RGB_ORDER: RGB
BACKEND: pillow
- NAME: Resize
SIZE: 512
INTERPOLATION: bilinear
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: CenterCrop
SIZE: 512
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: ImageToTensor
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: Normalize
MEAN: [ 0.5, 0.5, 0.5 ]
STD: [ 0.5, 0.5, 0.5 ]
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'image' ]
BACKEND: torchvision
- NAME: Select
KEYS: [ 'image', 'prompt' ]
META_KEYS: [ 'data_key' ]
#
EVAL_DATA:
NAME: ImageTextPairMSDataset
MODE: eval
MS_DATASET_NAME: style_custom_dataset
MS_DATASET_NAMESPACE: damo
MS_DATASET_SUBNAME: 3D
PROMPT_PREFIX: ""
MS_REMAP_KEYS: { 'Image': 'Target:FILE' }
MS_DATASET_SPLIT: test_short
OUTPUT_SIZE: [512, 512]
REPLACE_STYLE: False
PIN_MEMORY: True
BATCH_SIZE: 4
NUM_WORKERS: 4
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/data"
#
TRANSFORMS:
-
NAME: Select
KEYS: ['prompt']
META_KEYS: ['image_size']
#
TRAIN_HOOKS:
-
NAME: BackwardHook
PRIORITY: 0
-
NAME: LogHook
LOG_INTERVAL: 50
-
NAME: CheckpointHook
INTERVAL: 1000
-
NAME: ProbeDataHook
PROB_INTERVAL: 100
#
EVAL_HOOKS:
-
NAME: ProbeDataHook
PROB_INTERVAL: 100
@@ -113,7 +113,7 @@ SOLVER:
SAMPLE_STEPS: 50
SEED: 2023
GUIDE_SCALE: 7.5
GUIDE_RESCALE:
GUIDE_RESCALE: 0.5
DISCRETIZATION: trailing
IMAGE_SIZE: [768, 768]
RUN_TRAIN_N: False
@@ -122,7 +122,7 @@ SOLVER:
SAMPLE_STEPS: 50
SEED: 2023
GUIDE_SCALE: 7.5
GUIDE_RESCALE:
GUIDE_RESCALE: 0.5
DISCRETIZATION: trailing
IMAGE_SIZE: [768, 768]
RUN_TRAIN_N: False
@@ -31,7 +31,7 @@ SOLVER:
PARAMETERIZATION: v
TIMESTEPS: 1000
MIN_SNR_GAMMA:
ZERO_TERMINAL_SNR: True
ZERO_TERMINAL_SNR: False
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-2-1@v2-1_768-ema-pruned.safetensors
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
@@ -31,7 +31,7 @@ SOLVER:
PARAMETERIZATION: v
TIMESTEPS: 1000
MIN_SNR_GAMMA:
ZERO_TERMINAL_SNR: True
ZERO_TERMINAL_SNR: False
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-2-1@v2-1_768-ema-pruned.safetensors
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
@@ -0,0 +1,378 @@
ENV:
BACKEND: nccl
SOLVER:
NAME: LatentDiffusionSolver
RESUME_FROM:
LOAD_MODEL_ONLY: True
USE_FSDP: False
SHARDING_STRATEGY:
USE_AMP: True
DTYPE: float16
CHANNELS_LAST: True
MAX_STEPS: 200
MAX_EPOCHS: -1
NUM_FOLDS: 1
ACCU_STEP: 1
EVAL_INTERVAL: 100
#
WORK_DIR: ./cache/save_data/sdxl_1024_sce_ctr_canny
LOG_FILE: std_log.txt
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/data"
#
FREEZE:
FREEZE_PART: [ "first_stage_model", "cond_stage_model", "model" ]
TRAIN_PART: [ "control_blocks" ]
#
MODEL:
NAME: LatentDiffusionXLSCEControl
PARAMETERIZATION: eps
TIMESTEPS: 1000
MIN_SNR_GAMMA:
ZERO_TERMINAL_SNR: False
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-xl-base-1.0@sd_xl_base_1.0.safetensors
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.13025
SIZE_FACTOR: 8
DEFAULT_N_PROMPT:
SCHEDULE_ARGS:
"NAME": "scaled_linear"
"BETA_MIN": 0.00085
"BETA_MAX": 0.0120
USE_EMA: False
LOAD_REFINER: False
#
DIFFUSION_MODEL:
NAME: DiffusionUNetXL
PRETRAINED_MODEL:
IN_CHANNELS: 4
OUT_CHANNELS: 4
NUM_RES_BLOCKS: 2
MODEL_CHANNELS: 320
ATTENTION_RESOLUTIONS: [ 4, 2 ]
DROPOUT: 0
CHANNEL_MULT: [ 1, 2, 4 ]
CONV_RESAMPLE: True
DIMS: 2
NUM_CLASSES: sequential
USE_CHECKPOINT: False
NUM_HEADS: -1
NUM_HEADS_CHANNELS: 64
USE_SCALE_SHIFT_NORM: False
RESBLOCK_UPDOWN: False
USE_NEW_ATTENTION_ORDER: True
USE_SPATIAL_TRANSFORMER: True
TRANSFORMER_DEPTH: [ 1, 2, 10 ]
CONTEXT_DIM: 2048
DISABLE_MIDDLE_SELF_ATTN: False
USE_LINEAR_IN_TRANSFORMER: True
ADM_IN_CHANNELS: 2816
USE_SENTENCE_EMB: False
USE_WORD_MAPPING: False
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
PRETRAINED_MODEL:
IGNORE_KEYS: []
BATCH_SIZE: 1
#
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: GeneralConditioner
PRETRAINED_MODEL:
EMBEDDERS:
-
NAME: FrozenCLIPEmbedder
PRETRAINED_MODEL: ms://AI-ModelScope/clip-vit-large-patch14
TOKENIZER_PATH: ms://AI-ModelScope/clip-vit-large-patch14
MAX_LENGTH: 77
FREEZE: True
LAYER: hidden
LAYER_IDX: 11
USE_FINAL_LAYER_NORM: False
IS_TRAINABLE: False
UCG_RATE: 0.0
INPUT_KEYS: [ "prompt" ]
LEGACY_UCG_VALUE:
-
NAME: FrozenOpenCLIPEmbedder2
ARCH: ViT-bigG-14
PRETRAINED_MODEL:
MAX_LENGTH: 77
FREEZE: True
ALWAYS_RETURN_POOLED: True
LEGACY: False
LAYER: penultimate
IS_TRAINABLE: False
UCG_RATE: 0.0
INPUT_KEYS: [ "prompt" ]
LEGACY_UCG_VALUE:
-
NAME: ConcatTimestepEmbedderND
OUT_DIM: 256
IS_TRAINABLE: False
UCG_RATE: 0.0
INPUT_KEYS: [ "original_size_as_tuple" ]
LEGACY_UCG_VALUE:
-
NAME: ConcatTimestepEmbedderND
OUT_DIM: 256
IS_TRAINABLE: False
UCG_RATE: 0.0
INPUT_KEYS: [ "crop_coords_top_left" ]
LEGACY_UCG_VALUE:
-
NAME: ConcatTimestepEmbedderND
OUT_DIM: 256
IS_TRAINABLE: False
UCG_RATE: 0.0
INPUT_KEYS: [ "target_size_as_tuple" ]
LEGACY_UCG_VALUE:
#
REFINER_MODEL:
NAME: DiffusionUNetXL
PRETRAINED_MODEL:
IN_CHANNELS: 4
OUT_CHANNELS: 4
NUM_RES_BLOCKS: 2
MODEL_CHANNELS: 384
ATTENTION_RESOLUTIONS: [ 4, 2 ]
DROPOUT: 0
CHANNEL_MULT: [ 1, 2, 4, 4 ]
CONV_RESAMPLE: True
DIMS: 2
NUM_CLASSES: sequential
USE_CHECKPOINT: False
NUM_HEADS: -1
NUM_HEADS_CHANNELS: 64
USE_SCALE_SHIFT_NORM: False
RESBLOCK_UPDOWN: False
USE_NEW_ATTENTION_ORDER: True
USE_SPATIAL_TRANSFORMER: True
TRANSFORMER_DEPTH: 4
CONTEXT_DIM: [ 1280, 1280, 1280, 1280 ]
DISABLE_MIDDLE_SELF_ATTN: False
USE_LINEAR_IN_TRANSFORMER: True
ADM_IN_CHANNELS: 2560
USE_SENTENCE_EMB: False
USE_WORD_MAPPING: False
REFINER_COND_MODEL:
NAME: GeneralConditioner
PRETRAINED_MODEL:
EMBEDDERS:
-
NAME: FrozenOpenCLIPEmbedder2
ARCH: ViT-bigG-14
PRETRAINED_MODEL:
MAX_LENGTH: 77
FREEZE: True
ALWAYS_RETURN_POOLED: True
LEGACY: False
LAYER: penultimate
IS_TRAINABLE: False
UCG_RATE: 0.0
INPUT_KEYS: [ "prompt" ]
LEGACY_UCG_VALUE:
-
NAME: ConcatTimestepEmbedderND
OUT_DIM: 256
IS_TRAINABLE: False
UCG_RATE: 0.0
INPUT_KEYS: [ "original_size_as_tuple" ]
LEGACY_UCG_VALUE:
-
NAME: ConcatTimestepEmbedderND
OUT_DIM: 256
IS_TRAINABLE: False
UCG_RATE: 0.0
INPUT_KEYS: [ "crop_coords_top_left" ]
LEGACY_UCG_VALUE:
-
NAME: ConcatTimestepEmbedderND
OUT_DIM: 256
IS_TRAINABLE: False
UCG_RATE: 0.0
INPUT_KEYS: [ "aesthetic_score" ]
LEGACY_UCG_VALUE:
#
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
CONTROL_MODEL:
NAME: CSCTuners
PRE_HINT_IN_CHANNELS: 3
PRE_HINT_OUT_CHANNELS: 320
DENSE_HINT_KERNAL: 3
PRE_HINT_DIM_RATIO: 2.0
SCALE: 1.0
SC_TUNER_CFG:
NAME: SCTuner
TUNER_NAME: SCEAdapter
DOWN_RATIO: 1.0
CONTROL_ANNO:
NAME: CannyAnnotator
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 2023
GUIDE_SCALE: 7.5
GUIDE_RESCALE: 0.5
DISCRETIZATION: trailing
IMAGE_SIZE: [1024, 1024]
RUN_TRAIN_N: False
#
OPTIMIZER:
NAME: AdamW
LEARNING_RATE: 0.064
BETAS: [ 0.9, 0.999 ]
EPS: 1e-8
WEIGHT_DECAY: 1e-2
AMSGRAD: False
#
TRAIN_DATA:
NAME: ImageTextPairMSDataset
MODE: train
MS_DATASET_NAME: style_custom_dataset
MS_DATASET_NAMESPACE: damo
MS_DATASET_SUBNAME: 3D
PROMPT_PREFIX: ""
MS_DATASET_SPLIT: train_short
MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
REPLACE_STYLE: False
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
SAMPLER:
NAME: LoopSampler
TRANSFORMS:
- NAME: LoadImageFromFile
RGB_ORDER: RGB
BACKEND: pillow
- NAME: FlexibleResize
SIZE: 1024
INTERPOLATION: bilinear
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: FlexibleCropXL
SIZE: 1024
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: ToNumpy
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'image_preprocess' ]
- 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: [ 'img' ]
BACKEND: torchvision
- NAME: Rename
INPUT_KEY: [ 'img', 'image_preprocess', 'img_original_size_as_tuple', 'img_target_size_as_tuple', 'img_crop_coords_top_left' ]
OUTPUT_KEY: [ 'image', 'image_preprocess', 'original_size_as_tuple', 'target_size_as_tuple', 'crop_coords_top_left' ]
- NAME: Select
KEYS: [ 'image', 'prompt', 'image_preprocess', 'original_size_as_tuple', 'target_size_as_tuple', 'crop_coords_top_left' ]
META_KEYS: [ 'data_key' ]
#
EVAL_DATA:
NAME: ImageTextPairMSDataset
MODE: eval
MS_DATASET_NAME: style_custom_dataset
MS_DATASET_NAMESPACE: damo
MS_DATASET_SUBNAME: 3D
PROMPT_PREFIX: ""
MS_DATASET_SPLIT: train_short
MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
REPLACE_STYLE: False
PIN_MEMORY: True
BATCH_SIZE: 10
NUM_WORKERS: 4
TRANSFORMS:
- NAME: LoadImageFromFile
RGB_ORDER: RGB
BACKEND: pillow
- NAME: Resize
SIZE: 1024
INTERPOLATION: bilinear
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: CenterCrop
SIZE: 1024
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: ToNumpy
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'image_preprocess' ]
- 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: [ 'img' ]
BACKEND: torchvision
- NAME: Rename
INPUT_KEY: [ 'img', 'image_preprocess' ]
OUTPUT_KEY: [ 'image', 'image_preprocess' ]
- NAME: Select
KEYS: [ 'image', 'prompt', 'image_preprocess' ]
META_KEYS: [ 'data_key' ]
#
TRAIN_HOOKS:
-
NAME: BackwardHook
PRIORITY: 0
-
NAME: LogHook
LOG_INTERVAL: 50
-
NAME: CheckpointHook
INTERVAL: 100
-
NAME: ProbeDataHook
PROB_INTERVAL: 100
#
EVAL_HOOKS:
-
NAME: ProbeDataHook
PROB_INTERVAL: 100
@@ -231,8 +231,9 @@ SOLVER:
CONTROL_MODEL:
NAME: CSCTuners
PRE_HINT_IN_CHANNELS: 3
PRE_HINT_OUT_CHANNELS: 256
PRE_HINT_OUT_CHANNELS: 320
DENSE_HINT_KERNAL: 3
PRE_HINT_DIM_RATIO: 2.0
SCALE: 1.0
SC_TUNER_CFG:
NAME: SCTuner
@@ -231,8 +231,9 @@ SOLVER:
CONTROL_MODEL:
NAME: CSCTuners
PRE_HINT_IN_CHANNELS: 3
PRE_HINT_OUT_CHANNELS: 256
PRE_HINT_OUT_CHANNELS: 320
DENSE_HINT_KERNAL: 3
PRE_HINT_DIM_RATIO: 2.0
SCALE: 1.0
SC_TUNER_CFG:
NAME: SCTuner
@@ -231,8 +231,9 @@ SOLVER:
CONTROL_MODEL:
NAME: CSCTuners
PRE_HINT_IN_CHANNELS: 3
PRE_HINT_OUT_CHANNELS: 256
PRE_HINT_OUT_CHANNELS: 320
DENSE_HINT_KERNAL: 3
PRE_HINT_DIM_RATIO: 2.0
SCALE: 1.0
SC_TUNER_CFG:
NAME: SCTuner
@@ -1,19 +1,63 @@
CONTROLLERS:
# SD2.1
- NAME: canny
NAME_ZH:
DESCRIPTION:
BASE_MODEL: SD2.1
TYPE: Canny
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD2.1/canny_control/0_SwiftSCETuning
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD2.1/canny_control/
- NAME: openpose
NAME_ZH:
DESCRIPTION:
BASE_MODEL: SD2.1
TYPE: Openpose
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD2.1/pose_control/0_SwiftSCETuning
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD2.1/pose_control/
- NAME: color
NAME_ZH:
DESCRIPTION:
BASE_MODEL: SD2.1
TYPE: Color
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD2.1/color_control/0_SwiftSCETuning
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD2.1/color_control/
- NAME: hed
NAME_ZH:
DESCRIPTION:
BASE_MODEL: SD2.1
TYPE: Hed
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD2.1/hed_control
- NAME: depth
NAME_ZH:
DESCRIPTION:
BASE_MODEL: SD2.1
TYPE: Midas
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD2.1/depth_control
# SD_XL1.0
- NAME: canny
NAME_ZH:
DESCRIPTION:
BASE_MODEL: SD_XL1.0
TYPE: Canny
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD_XL1.0/canny_control
- NAME: color
NAME_ZH:
DESCRIPTION:
BASE_MODEL: SD_XL1.0
TYPE: Color
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD_XL1.0/color_control
- NAME: depth
NAME_ZH:
DESCRIPTION:
BASE_MODEL: SD_XL1.0
TYPE: Midas
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD_XL1.0/depth_control
- NAME: hed
NAME_ZH:
DESCRIPTION:
BASE_MODEL: SD_XL1.0
TYPE: Hed
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD_XL1.0/hed_control
- NAME: openpose
NAME_ZH:
DESCRIPTION:
BASE_MODEL: SD_XL1.0
TYPE: Openpose
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD_XL1.0/pose_control
@@ -1,4 +1,76 @@
TUNERS:
- NAME: Azure-Dragon
NAME_ZH: 青龙
SOURCE: wanx
DESCRIPTION: None
BASE_MODEL: SD_XL1.0
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/azure_dragon/
IMAGE_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/azure_dragon/xl_azure_dragon.png
TUNER_TYPE: SwiftSCE
PROMPT_EXAMPLE: Azure Dragon, 8K, high quality,Ultra High Detail.One of the Four Divine Creatures in Charge of Water.
- NAME: Gold-Dragon
NAME_ZH: 金龙
SOURCE: wanx
DESCRIPTION: None
BASE_MODEL: SD_XL1.0
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/gold_dragon/
IMAGE_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/gold_dragon/xl_gold_dragon.png
TUNER_TYPE: SwiftSCE
PROMPT_EXAMPLE: Chinese Gold Dragon in the clouds. Translucent Texture. Zbrush. Fuzzy Art. Exquisite Craftsmanship. 3D. 8K. Ultra High Detail
- NAME: SpringFestival-Dragon
NAME_ZH: 春节龙
SOURCE: wanx
DESCRIPTION: None
BASE_MODEL: SD_XL1.0
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/spring_festival_dragon/
IMAGE_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/spring_festival_dragon/xl_spring_festival_dragon.png
TUNER_TYPE: SwiftSCE
PROMPT_EXAMPLE: Chinese dragon. Spring Festival.Festive.Street.Lanterns.32K.High quality.expressive, dramatic, dreamlike and mysterious, Surrealism
- NAME: Red-Dragon
NAME_ZH: 红龙
SOURCE: wanx
DESCRIPTION: None
BASE_MODEL: SD_XL1.0
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/red_dragon/
IMAGE_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/red_dragon/xl_red_dragon.png
TUNER_TYPE: SwiftSCE
PROMPT_EXAMPLE: Traditional Red Dragon of China. Low Water Level. Studio Ghibli Style. Mural Illustration. White Background. High Detail
- NAME: ChinesePunk-Dragon
NAME_ZH: 中国朋克龙
SOURCE: wanx
DESCRIPTION: None
BASE_MODEL: SD_XL1.0
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/chinese_punk_dragon/
IMAGE_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/chinese_punk_dragon/xl_chinese_punk_dragon.png
TUNER_TYPE: SwiftSCE
PROMPT_EXAMPLE: uhd Image,Dragon,Chinese Dragon, Dunhuang Mural Style, Traditional Maritime Art Style
- NAME: Cute-Dragon
NAME_ZH: 喜庆龙
SOURCE: wanx
DESCRIPTION: None
BASE_MODEL: SD_XL1.0
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/cute_dragon/
IMAGE_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/cute_dragon/xl_kawaii_dragon.png
TUNER_TYPE: SwiftSCE
PROMPT_EXAMPLE: China Kawaii Dragon. Contest Winner. Minimalist Illustration. White Background. Flat Style. Digital Painting Style. Red. 32k uhd. Fun Comics. Fuzzy Art. Bold. Comic-Inspired Characters
- NAME: Dragon-Baby
NAME_ZH: 龙宝宝
SOURCE: wanx
DESCRIPTION: None
BASE_MODEL: SD_XL1.0
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/baby_dragon/
IMAGE_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/baby_dragon/xl_baby_dragon.png
TUNER_TYPE: SwiftSCE
PROMPT_EXAMPLE: Warm Colors, Soft,Chinese Dragon Baby, Felt Style,Dragon Baby, Best Quality, 3D Doll, Macaron Tones, Glittering Big Eyes, Winter,Dragon
- NAME: Sloppy-Dragon
NAME_ZH: 潦草龙
SOURCE: wanx
DESCRIPTION: None
BASE_MODEL: SD_XL1.0
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/sloppy_dragon/
IMAGE_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/sloppy_dragon/xl_sloppy_dragon.png
TUNER_TYPE: SwiftSCE
PROMPT_EXAMPLE: Messy Chinese Dragon,Cute, Wu Guanzhong, Rough
-
NAME: Caricature
NAME_ZH: 夸张漫画
+11 -11
View File
@@ -38,15 +38,15 @@ GUIDE_INFO:
width: 100%;
}
.video-wrapper {
width: 75%;
width: 75%;
}
video {
width: 100%;
display: block;
width: 100%;
display: block;
}
.description {
text-align: center;
margin-top: 10px;
text-align: center;
margin-top: 10px;
font-size: 0.8em;
}
</style>
@@ -70,15 +70,15 @@ GUIDE_INFO:
width: 100%;
}
.video-wrapper {
width: 75%;
width: 75%;
}
video {
width: 100%;
display: block;
width: 100%;
display: block;
}
.description {
text-align: center;
margin-top: 10px;
text-align: center;
margin-top: 10px;
font-size: 0.8em;
}
</style>
@@ -92,4 +92,4 @@ GUIDE_INFO:
<div class="description">Train & Inference Video</div>
</div>
</div>
</body>
</body>
@@ -82,8 +82,6 @@ EXTENSION_PARAS:
CONTROLABLE_ANNOTATORS:
-
NAME: "CannyAnnotator"
LOW_THRESHOLD: 100
HIGH_THRESHOLD: 200
TYPE: Canny
IS_DEFAULT: True
-
@@ -100,8 +98,6 @@ CONTROLABLE_ANNOTATORS:
-
NAME: "MidasDetector"
PRETRAINED_MODEL: "ms://damo/scepter_scedit@annotator/ckpts/dpt_hybrid-midas-501f0c75.pt"
A: 6.2
BG_TH: 0.1
TYPE: Midas
IS_DEFAULT: False
-
@@ -109,9 +105,6 @@ CONTROLABLE_ANNOTATORS:
TYPE: Color
IS_DEFAULT: False
-
NAME: "MLSDdetector"
PRETRAINED_MODEL: "ms://damo/scepter_scedit@annotator/ckpts/mlsd_large_512_fp32.pth"
THR_V: 0.1
THR_D: 0.1
TYPE: MLSD
NAME: "InvertAnnotator"
TYPE: Invert-Preprocess
IS_DEFAULT: False
+3
View File
@@ -1,8 +1,11 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.annotator.base_annotator import GeneralAnnotator
from scepter.modules.annotator.canny import CannyAnnotator
from scepter.modules.annotator.color import ColorAnnotator
from scepter.modules.annotator.hed import HedAnnotator
from scepter.modules.annotator.identity import IdentityAnnotator
from scepter.modules.annotator.invert import InvertAnnotator
from scepter.modules.annotator.midas_op import MidasDetector
from scepter.modules.annotator.mlsd_op import MLSDdetector
from scepter.modules.annotator.openpose import OpenposeAnnotator
+24 -4
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from abc import ABCMeta
import cv2
@@ -19,20 +20,39 @@ class CannyAnnotator(BaseAnnotator, metaclass=ABCMeta):
super().__init__(cfg, logger=logger)
self.low_threshold = cfg.get('LOW_THRESHOLD', 100)
self.high_threshold = cfg.get('HIGH_THRESHOLD', 200)
self.random_cfg = cfg.get('RANDOM_CFG', None)
def forward(self, image):
if isinstance(image, Image.Image):
image = np.array(image)
image = cv2.Canny(image, self.low_threshold, self.high_threshold)
elif isinstance(image, torch.Tensor):
image = image.detach().cpu().numpy()
image = cv2.Canny(image, self.low_threshold, self.high_threshold)
elif isinstance(image, np.ndarray):
image = cv2.Canny(image.copy(), self.low_threshold,
self.high_threshold)
image = image.copy()
else:
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
assert len(image.shape) < 4
if self.random_cfg is None:
image = cv2.Canny(image, self.low_threshold, self.high_threshold)
else:
proba = self.random_cfg.get('PROBA', 1.0)
if np.random.random() < proba:
min_low_threshold = self.random_cfg.get(
'MIN_LOW_THRESHOLD', 50)
max_low_threshold = self.random_cfg.get(
'MAX_LOW_THRESHOLD', 100)
min_high_threshold = self.random_cfg.get(
'MIN_HIGH_THRESHOLD', 200)
max_high_threshold = self.random_cfg.get(
'MAX_HIGH_THRESHOLD', 350)
low_th = np.random.randint(min_low_threshold,
max_low_threshold)
high_th = np.random.randint(min_high_threshold,
max_high_threshold)
else:
low_th, high_th = self.low_threshold, self.high_threshold
image = cv2.Canny(image, low_th, high_th)
return image[..., None].repeat(3, 2)
@staticmethod
+17 -2
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from abc import ABCMeta
import cv2
@@ -18,6 +19,7 @@ class ColorAnnotator(BaseAnnotator, metaclass=ABCMeta):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.ratio = cfg.get('RATIO', 64)
self.random_cfg = cfg.get('RANDOM_CFG', None)
def forward(self, image):
if isinstance(image, Image.Image):
@@ -29,8 +31,21 @@ class ColorAnnotator(BaseAnnotator, metaclass=ABCMeta):
else:
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
h, w = image.shape[:2]
ratio = self.ratio
image = cv2.resize(image, (w // ratio, h // ratio),
if self.random_cfg is None:
ratio = self.ratio
else:
proba = self.random_cfg.get('PROBA', 1.0)
if np.random.random() < proba:
if 'CHOICE_RATIO' in self.random_cfg:
ratio = np.random.choice(self.random_cfg['CHOICE_RATIO'])
else:
min_ratio = self.random_cfg.get('MIN_RATIO', 48)
max_ratio = self.random_cfg.get('MAX_RATIO', 96)
ratio = np.random.randint(min_ratio, max_ratio)
else:
ratio = self.ratio
image = cv2.resize(image, (int(w // ratio), int(h // ratio)),
interpolation=cv2.INTER_CUBIC)
image = cv2.resize(image, (w, h), interpolation=cv2.INTER_NEAREST)
assert len(image.shape) < 4
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
# Please use this implementation in your products
# This implementation may produce slightly different results from Saining Xie's official implementations,
# but it generates smoother edges and is more suitable for ControlNet as well as other image-to-image translations.
+25
View File
@@ -0,0 +1,25 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from abc import ABCMeta
from scepter.modules.annotator.base_annotator import BaseAnnotator
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.utils.config import dict_to_yaml
@ANNOTATORS.register_class()
class IdentityAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
def forward(self, image):
return image
@staticmethod
def get_config_template():
return dict_to_yaml('ANNOTATORS',
__class__.__name__,
IdentityAnnotator.para_dict,
set_name=True)
+25
View File
@@ -0,0 +1,25 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from abc import ABCMeta
from scepter.modules.annotator.base_annotator import BaseAnnotator
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.utils.config import dict_to_yaml
@ANNOTATORS.register_class()
class InvertAnnotator(BaseAnnotator, metaclass=ABCMeta):
para_dict = {}
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
def forward(self, image):
return 255 - image
@staticmethod
def get_config_template():
return dict_to_yaml('ANNOTATORS',
__class__.__name__,
InvertAnnotator.para_dict,
set_name=True)
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
# Midas Depth Estimation
# From https://github.com/isl-org/MiDaS
# MIT LICENSE
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
# MLSD Line Detection
# From https://github.com/navervision/mlsd
# Apache-2.0 license
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
# Openpose
# Original from CMU https://github.com/CMU-Perceptual-Computing-Lab/openpose
# 2nd Edited by https://github.com/Hzzone/pytorch-openpose
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import cv2
import numpy as np
+1 -3
View File
@@ -1,14 +1,13 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os
from abc import ABCMeta, abstractmethod
from torch.utils.data import Dataset
from scepter.modules.transform.registry import TRANSFORMS, build_pipeline
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import set_random_seed, we
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.logger import get_logger
from scepter.modules.utils.registry import old_python_version
@@ -84,7 +83,6 @@ class BaseDataset(Dataset, metaclass=ABCMeta):
overwrite=False)
self.worker_id = worker_id
self.logger = self.worker_logger
set_random_seed(int(os.environ.get('ES_SEED', 2023)))
we.set_env(self.local_we)
@abstractmethod
+1
View File
@@ -1,2 +1,3 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.inference.diffusion_inference import DiffusionInference
@@ -0,0 +1,112 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import os
import torch
import torch.nn as nn
import torchvision.transforms as TT
from PIL.Image import Image
from swift import SwiftModel
from scepter.modules.model.registry import TUNERS
from scepter.modules.utils.config import Config
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
class ControlInference():
def __init__(self, logger=None):
self.logger = logger
self.is_register = False
# @classmethod
def unregister_controllers(self, control_model_ins, diffusion_model):
self.logger.info('Unloading control model')
if isinstance(diffusion_model['model'], SwiftModel):
if (hasattr(diffusion_model['model'].base_model, 'control_blocks')
and diffusion_model['model'].base_model.control_blocks
): # noqa
del diffusion_model['model'].base_model.control_blocks
diffusion_model['model'].base_model.control_blocks = None
diffusion_model['model'].base_model.control_name = []
else:
del diffusion_model['model'].control_blocks
diffusion_model['model'].control_blocks = None
diffusion_model['model'].control_name = []
self.is_register = False
# @classmethod
def register_controllers(self, control_model_ins, diffusion_model):
self.logger.info('Loading control model')
if control_model_ins is None or control_model_ins == '':
self.unregister_controllers(control_model_ins, diffusion_model)
return
if not isinstance(control_model_ins, list):
control_model_ins = [control_model_ins]
control_model = nn.ModuleList([])
control_model_folder = []
for one_control in control_model_ins:
one_control_model_folder = one_control.MODEL_PATH
control_model_folder.append(one_control_model_folder)
have_list = getattr(diffusion_model['model'], 'control_name', [])
if one_control_model_folder in have_list:
ind = have_list.index(one_control_model_folder)
csc_tuners = copy.deepcopy(
diffusion_model['model'].control_blocks[ind])
else:
one_local_control_model = FS.get_dir_to_local_dir(
one_control_model_folder)
control_cfg = Config(cfg_file=os.path.join(
one_local_control_model, '0_SwiftSCETuning',
'configuration.json'))
assert hasattr(control_cfg, 'CONTROL_MODEL')
control_cfg.CONTROL_MODEL[
'INPUT_BLOCK_CHANS'] = diffusion_model[
'model']._input_block_chans
control_cfg.CONTROL_MODEL['INPUT_DOWN_FLAG'] = diffusion_model[
'model']._input_down_flag
control_cfg.CONTROL_MODEL.PRETRAINED_MODEL = os.path.join(
one_local_control_model, '0_SwiftSCETuning',
'pytorch_model.bin')
csc_tuners = TUNERS.build(control_cfg.CONTROL_MODEL,
logger=self.logger)
control_model.append(csc_tuners)
control_model.to(diffusion_model['device'])
if isinstance(diffusion_model['model'], SwiftModel):
del diffusion_model['model'].base_model.control_blocks
diffusion_model['model'].base_model.control_blocks = control_model
diffusion_model[
'model'].base_model.control_name = control_model_folder
else:
del diffusion_model['model'].control_blocks
diffusion_model['model'].control_blocks = control_model
diffusion_model['model'].control_name = control_model_folder
self.is_register = True
@classmethod
def get_control_input(self, control_model, control_cond_image, height,
width):
hints = []
if control_cond_image and control_model:
if not isinstance(control_model, list):
control_model = [control_model]
if not isinstance(control_cond_image, list):
control_cond_image = [control_cond_image]
assert len(control_cond_image) == len(control_model)
for img in control_cond_image:
if isinstance(img, Image):
w, h = img.size
if not h == height or not w == width:
img = TT.Resize(min(height, width))(img)
img = TT.CenterCrop((height, width))(img)
hint = TT.ToTensor()(img)
hints.append(hint)
else:
raise NotImplementedError
if len(hints) > 0:
hints = torch.stack(hints).to(we.device_id)
else:
hints = None
return hints
+114 -292
View File
@@ -1,27 +1,24 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import hashlib
import json
import os.path
import random
from collections import OrderedDict
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision.transforms as TT
from peft.utils import CONFIG_NAME, SAFETENSORS_WEIGHTS_NAME, WEIGHTS_NAME
from PIL.Image import Image
from swift import Swift, SwiftModel
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, TUNERS)
from scepter.modules.utils.config import Config
TOKENIZERS)
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from .control_inference import ControlInference
from .tuner_inference import TunerInference
def get_model(model_tuple):
assert 'model' in model_tuple
@@ -36,6 +33,12 @@ class DiffusionInference():
'''
def __init__(self, logger=None):
self.logger = logger
self.loaded_model = {}
self.loaded_model_name = [
'diffusion_model', 'first_stage_model', 'cond_stage_model'
]
self.tuner_infer = TunerInference(self.logger)
self.control_infer = ControlInference(self.logger)
def init_from_cfg(self, cfg):
self.name = cfg.NAME
@@ -76,248 +79,6 @@ class DiffusionInference():
'vocab_size': self.tokenizer.vocab_size
}
def register_tuner(self, tuner_model_list):
if len(tuner_model_list) < 1:
if isinstance(self.diffusion_model['model'], SwiftModel):
for adapter_name in self.diffusion_model['model'].adapters:
self.diffusion_model['model'].deactivate_adapter(
adapter_name, offload='cpu')
if isinstance(self.cond_stage_model['model'], SwiftModel):
for adapter_name in self.cond_stage_model['model'].adapters:
self.cond_stage_model['model'].deactivate_adapter(
adapter_name, offload='cpu')
return
all_diffusion_tuner = {}
all_cond_tuner = {}
save_root_dir = '.cache_tuner'
for tuner_model in tuner_model_list:
tunner_model_folder = tuner_model.MODEL_PATH
local_tuner_model = FS.get_dir_to_local_dir(tunner_model_folder)
all_tuner_datas = os.listdir(local_tuner_model)
cur_tuner_md5 = hashlib.md5(
tunner_model_folder.encode('utf-8')).hexdigest()
local_diffusion_cache = os.path.join(
save_root_dir, cur_tuner_md5 + '_' + 'diffusion')
local_cond_cache = os.path.join(save_root_dir,
cur_tuner_md5 + '_' + 'cond')
meta_file = os.path.join(save_root_dir,
cur_tuner_md5 + '_meta.json')
if not os.path.exists(meta_file):
diffusion_tuner = {}
cond_tuner = {}
for sub in all_tuner_datas:
sub_file = os.path.join(local_tuner_model, sub)
config_file = os.path.join(sub_file, CONFIG_NAME)
safe_file = os.path.join(sub_file,
SAFETENSORS_WEIGHTS_NAME)
bin_file = os.path.join(sub_file, WEIGHTS_NAME)
if os.path.isdir(sub_file) and os.path.isfile(config_file):
# diffusion or cond
cfg = json.load(open(config_file, 'r'))
if 'cond_stage_model.' in cfg['target_modules']:
cond_cfg = copy.deepcopy(cfg)
if 'cond_stage_model.*' in cond_cfg[
'target_modules']:
cond_cfg['target_modules'] = cond_cfg[
'target_modules'].replace(
'cond_stage_model.*', '.*')
else:
cond_cfg['target_modules'] = cond_cfg[
'target_modules'].replace(
'cond_stage_model.', '')
if cond_cfg['target_modules'].startswith('*'):
cond_cfg['target_modules'] = '.' + cond_cfg[
'target_modules']
os.makedirs(local_cond_cache + '_' + sub,
exist_ok=True)
cond_tuner[os.path.basename(local_cond_cache) +
'_' + sub] = hashlib.md5(
(local_cond_cache + '_' +
sub).encode('utf-8')).hexdigest()
os.makedirs(local_cond_cache + '_' + sub,
exist_ok=True)
json.dump(
cond_cfg,
open(
os.path.join(local_cond_cache + '_' + sub,
CONFIG_NAME), 'w'))
if 'model.' in cfg['target_modules'].replace(
'cond_stage_model.', ''):
diffusion_cfg = copy.deepcopy(cfg)
if 'model.*' in diffusion_cfg['target_modules']:
diffusion_cfg[
'target_modules'] = diffusion_cfg[
'target_modules'].replace(
'model.*', '.*')
else:
diffusion_cfg[
'target_modules'] = diffusion_cfg[
'target_modules'].replace(
'model.', '')
if diffusion_cfg['target_modules'].startswith('*'):
diffusion_cfg[
'target_modules'] = '.' + diffusion_cfg[
'target_modules']
os.makedirs(local_diffusion_cache + '_' + sub,
exist_ok=True)
diffusion_tuner[
os.path.basename(local_diffusion_cache) + '_' +
sub] = hashlib.md5(
(local_diffusion_cache + '_' +
sub).encode('utf-8')).hexdigest()
json.dump(
diffusion_cfg,
open(
os.path.join(
local_diffusion_cache + '_' + sub,
CONFIG_NAME), 'w'))
state_dict = {}
is_bin_file = True
if os.path.isfile(bin_file):
state_dict = torch.load(bin_file)
elif os.path.isfile(safe_file):
is_bin_file = False
from safetensors.torch import \
load_file as safe_load_file
state_dict = safe_load_file(
safe_file,
device='cuda'
if torch.cuda.is_available() else 'cpu')
save_diffusion_state_dict = {}
save_cond_state_dict = {}
for key, value in state_dict.items():
if key.startswith('model.'):
save_diffusion_state_dict[
key[len('model.'):].replace(
sub,
os.path.basename(local_diffusion_cache)
+ '_' + sub)] = value
elif key.startswith('cond_stage_model.'):
save_cond_state_dict[
key[len('cond_stage_model.'):].replace(
sub,
os.path.basename(local_cond_cache) +
'_' + sub)] = value
if is_bin_file:
if len(save_diffusion_state_dict) > 0:
torch.save(
save_diffusion_state_dict,
os.path.join(
local_diffusion_cache + '_' + sub,
WEIGHTS_NAME))
if len(save_cond_state_dict) > 0:
torch.save(
save_cond_state_dict,
os.path.join(local_cond_cache + '_' + sub,
WEIGHTS_NAME))
else:
from safetensors.torch import \
save_file as safe_save_file
if len(save_diffusion_state_dict) > 0:
safe_save_file(
save_diffusion_state_dict,
os.path.join(
local_diffusion_cache + '_' + sub,
SAFETENSORS_WEIGHTS_NAME),
metadata={'format': 'pt'})
if len(save_cond_state_dict) > 0:
safe_save_file(
save_cond_state_dict,
os.path.join(local_cond_cache + '_' + sub,
SAFETENSORS_WEIGHTS_NAME),
metadata={'format': 'pt'})
json.dump(
{
'diffusion_tuner': diffusion_tuner,
'cond_tuner': cond_tuner
}, open(meta_file, 'w'))
else:
meta_conf = json.load(open(meta_file, 'r'))
diffusion_tuner = meta_conf['diffusion_tuner']
cond_tuner = meta_conf['cond_tuner']
all_diffusion_tuner.update(diffusion_tuner)
all_cond_tuner.update(cond_tuner)
if len(all_diffusion_tuner) > 0:
self.load(self.diffusion_model)
self.diffusion_model['model'] = Swift.from_pretrained(
self.diffusion_model['model'],
save_root_dir,
adapter_name=all_diffusion_tuner)
self.diffusion_model['model'].set_active_adapters(
list(all_diffusion_tuner.values()))
self.unload(self.diffusion_model)
if len(all_cond_tuner) > 0:
self.load(self.cond_stage_model)
self.cond_stage_model['model'] = Swift.from_pretrained(
self.cond_stage_model['model'],
save_root_dir,
adapter_name=all_cond_tuner)
self.cond_stage_model['model'].set_active_adapters(
list(all_cond_tuner.values()))
self.unload(self.cond_stage_model)
def register_controllers(self, control_model_ins):
if control_model_ins is None or control_model_ins == '':
if isinstance(self.diffusion_model['model'], SwiftModel):
if (hasattr(self.diffusion_model['model'].base_model,
'control_blocks') and
self.diffusion_model['model'].base_model.control_blocks
): # noqa
del self.diffusion_model['model'].base_model.control_blocks
self.diffusion_model[
'model'].base_model.control_blocks = None
self.diffusion_model['model'].base_model.control_name = []
else:
del self.diffusion_model['model'].control_blocks
self.diffusion_model['model'].control_blocks = None
self.diffusion_model['model'].control_name = []
return
if not isinstance(control_model_ins, list):
control_model_ins = [control_model_ins]
control_model = nn.ModuleList([])
control_model_folder = []
for one_control in control_model_ins:
one_control_model_folder = one_control.MODEL_PATH
control_model_folder.append(one_control_model_folder)
have_list = getattr(self.diffusion_model['model'], 'control_name',
[])
if one_control_model_folder in have_list:
ind = have_list.index(one_control_model_folder)
csc_tuners = copy.deepcopy(
self.diffusion_model['model'].control_blocks[ind])
else:
one_local_control_model = FS.get_dir_to_local_dir(
one_control_model_folder)
control_cfg = Config(cfg_file=os.path.join(
one_local_control_model, 'configuration.json'))
assert hasattr(control_cfg, 'CONTROL_MODEL')
control_cfg.CONTROL_MODEL[
'INPUT_BLOCK_CHANS'] = self.diffusion_model[
'model']._input_block_chans
control_cfg.CONTROL_MODEL[
'INPUT_DOWN_FLAG'] = self.diffusion_model[
'model']._input_down_flag
control_cfg.CONTROL_MODEL.PRETRAINED_MODEL = os.path.join(
one_local_control_model, 'pytorch_model.bin')
csc_tuners = TUNERS.build(control_cfg.CONTROL_MODEL,
logger=self.logger)
control_model.append(csc_tuners)
if isinstance(self.diffusion_model['model'], SwiftModel):
del self.diffusion_model['model'].base_model.control_blocks
self.diffusion_model[
'model'].base_model.control_blocks = control_model
self.diffusion_model[
'model'].base_model.control_name = control_model_folder
else:
del self.diffusion_model['model'].control_blocks
self.diffusion_model['model'].control_blocks = control_model
self.diffusion_model['model'].control_name = control_model_folder
def redefine_paras(self, cfg):
if cfg.get('PRETRAINED_MODEL', None):
assert FS.isfile(cfg.PRETRAINED_MODEL)
@@ -480,12 +241,55 @@ class DiffusionInference():
return module
def unload(self, module):
if module is None:
return module
module['model'] = module['model'].to('cpu')
module['device'] = 'cpu'
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
return module
def dynamic_load(self, module=None, name=''):
self.logger.info('Loading {} model'.format(name))
if name == 'all':
for subname in self.loaded_model_name:
self.loaded_model[subname] = self.dynamic_load(
getattr(self, subname), subname)
elif name in self.loaded_model_name:
if name in self.loaded_model:
if module['cfg'] != self.loaded_model[name]['cfg']:
self.unload(self.loaded_model[name])
module = self.load(module)
self.loaded_model[name] = module
return module
elif module['device'] == 'cpu':
module = self.load(module)
return module
else:
return module
else:
module = self.load(module)
self.loaded_model[name] = module
return module
else:
return self.load(module)
def dynamic_unload(self, module=None, name='', skip_loaded=False):
self.logger.info('Unloading {} model'.format(name))
if name == 'all':
for name, module in self.loaded_model.items():
module = self.unload(self.loaded_model[name])
self.loaded_model[name] = module
elif name in self.loaded_model_name:
if name in self.loaded_model:
if not skip_loaded:
module = self.unload(self.loaded_model[name])
self.loaded_model[name] = module
else:
self.unload(module)
else:
self.unload(module)
def load_default(self, cfg):
module_paras = {}
if cfg is not None:
@@ -616,35 +420,34 @@ class DiffusionInference():
height, width = value_input['target_size_as_tuple']
value_output = copy.deepcopy(self.output)
batch, batch_uc = self.get_batch(value_input, num_samples=1)
#
if not isinstance(tuner_model, list):
tuner_model = [tuner_model]
for tuner in tuner_model:
if tuner is None or tuner == '':
tuner_model.remove(tuner)
self.register_tuner(tuner_model)
# control_cond_image
control_cond_image = kwargs.pop('control_cond_image', None)
# crop_type = kwargs.pop('crop_type', 'center_crop')
hints = []
if control_cond_image and control_model:
if not isinstance(control_model, list):
control_model = [control_model]
if not isinstance(control_cond_image, list):
control_cond_image = [control_cond_image]
assert len(control_cond_image) == len(control_model)
for img in control_cond_image:
if isinstance(img, Image):
w, h = img.size
if not h == height or not w == width:
img = TT.Resize(min(height, width))(img)
img = TT.CenterCrop((height, width))(img)
hint = TT.ToTensor()(img)
hints.append(hint)
else:
raise NotImplementedError
if len(hints) > 0:
hints = torch.stack(hints).to(we.device_id)
# register tuner
if tuner_model is not None and tuner_model != '' and len(
tuner_model) > 0:
if not isinstance(tuner_model, list):
tuner_model = [tuner_model]
self.dynamic_load(self.diffusion_model, 'diffusion_model')
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
self.tuner_infer.register_tuner(tuner_model, self.diffusion_model,
self.cond_stage_model)
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=True)
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=True)
# register control
if control_model is not None and control_model != '':
self.dynamic_load(self.diffusion_model, 'diffusion_model')
hints = ControlInference.get_control_input(
control_model, kwargs.pop('control_cond_image', None), height,
width)
self.control_infer.register_controllers(control_model,
self.diffusion_model)
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=True)
else:
hints = None
@@ -655,15 +458,17 @@ class DiffusionInference():
b, c, ori_width, ori_height = image.shape
if not (ori_width == width and ori_height == height):
image = F.interpolate(image, (width, height), mode='bicubic')
self.first_stage_model = self.load(self.first_stage_model)
self.dynamic_load(self.first_stage_model, 'first_stage_model')
input_latent = self.encode_first_stage(image)
self.first_stage_model = self.unload(self.first_stage_model)
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=True)
else:
input_latent = None
if 'input_latent' in value_output and input_latent is not None:
value_output['input_latent'] = input_latent
# cond stage
self.cond_stage_model = self.load(self.cond_stage_model)
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',
@@ -678,7 +483,9 @@ class DiffusionInference():
function_name)(batch)
null_context = getattr(get_model(self.cond_stage_model),
function_name)(batch_uc)
self.cond_stage_model = self.unload(self.cond_stage_model)
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=True)
if refine_strength > 0 and self.refiner_diffusion_model is not None:
assert self.refiner_cond_model is not None
@@ -703,9 +510,7 @@ class DiffusionInference():
get_model(self.refiner_cond_model),
function_name)(batch_uc)
self.refiner_cond_model = self.unload(self.refiner_cond_model)
self.load(self.diffusion_model)
self.register_controllers(control_model)
self.unload(self.diffusion_model)
# get noise
seed = kwargs.pop('seed', -1)
g = torch.Generator(device=we.device_id)
@@ -721,8 +526,8 @@ class DiffusionInference():
height // self.first_stage_model['paras']['size_factor'],
width // self.first_stage_model['paras']['size_factor'],
device=we.device_id).normal_(generator=g)
#
self.load(self.diffusion_model)
self.dynamic_load(self.diffusion_model, 'diffusion_model')
# UNet use input n_prompt
function_name, dtype = self.get_function_info(
self.diffusion_model)
@@ -760,7 +565,10 @@ class DiffusionInference():
intermediate_callback=intermediate_callback,
cat_uc=cat_uc,
**kwargs)
self.diffusion_model = self.unload(self.diffusion_model)
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=True)
# apply refiner
if refine_strength > 0 and self.refiner_diffusion_model is not None:
@@ -828,10 +636,11 @@ class DiffusionInference():
value_output['latent'] = []
value_output['latent'].append(latent)
self.first_stage_model = self.load(self.first_stage_model)
self.dynamic_load(self.first_stage_model, 'first_stage_model')
x_samples = self.decode_first_stage(latent).float()
self.first_stage_model = self.unload(self.first_stage_model)
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=True)
images = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0)
if 'images' in value_output:
if value_output['images'] is None or (
@@ -845,4 +654,17 @@ class DiffusionInference():
value_output[k] = torch.cat(v, dim=0)
if isinstance(v, torch.Tensor):
value_output[k] = v.cpu()
# unregister tuner
if tuner_model is not None and tuner_model != '' and len(
tuner_model) > 0:
self.tuner_infer.unregister_tuner(tuner_model,
self.diffusion_model,
self.cond_stage_model)
# unregister control
if control_model is not None and control_model != '':
self.control_infer.unregister_controllers(control_model,
self.diffusion_model)
return value_output
@@ -0,0 +1,220 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import hashlib
import json
import os
import warnings
import torch
from scepter.modules.utils.file_system import FS
try:
from peft.utils import CONFIG_NAME, SAFETENSORS_WEIGHTS_NAME, WEIGHTS_NAME
except Exception as e:
warnings.warn(f'Import peft error, please deal with this problem: {e}')
try:
from swift import Swift, SwiftModel
except Exception as e:
warnings.warn(f'Import swift error, please deal with this problem: {e}')
class TunerInference():
def __init__(self, logger=None):
self.logger = logger
self.is_register = False
# @classmethod
def unregister_tuner(self, tuner_model_list, diffusion_model,
cond_stage_model):
self.logger.info('Unloading tuner model')
if 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):
for adapter_name in cond_stage_model['model'].adapters:
cond_stage_model['model'].deactivate_adapter(adapter_name,
offload='cpu')
return
# @classmethod
def register_tuner(self, tuner_model_list, diffusion_model,
cond_stage_model):
self.logger.info('Loading tuner model')
if len(tuner_model_list) < 1:
self.unregister_tuner(tuner_model_list, diffusion_model,
cond_stage_model)
return
all_diffusion_tuner = {}
all_cond_tuner = {}
save_root_dir = '.cache_tuner'
for tuner_model in tuner_model_list:
tunner_model_folder = tuner_model.MODEL_PATH
local_tuner_model = FS.get_dir_to_local_dir(tunner_model_folder)
all_tuner_datas = os.listdir(local_tuner_model)
cur_tuner_md5 = hashlib.md5(
tunner_model_folder.encode('utf-8')).hexdigest()
local_diffusion_cache = os.path.join(
save_root_dir, cur_tuner_md5 + '_' + 'diffusion')
local_cond_cache = os.path.join(save_root_dir,
cur_tuner_md5 + '_' + 'cond')
meta_file = os.path.join(save_root_dir,
cur_tuner_md5 + '_meta.json')
if not os.path.exists(meta_file):
diffusion_tuner = {}
cond_tuner = {}
for sub in all_tuner_datas:
sub_file = os.path.join(local_tuner_model, sub)
config_file = os.path.join(sub_file, CONFIG_NAME)
safe_file = os.path.join(sub_file,
SAFETENSORS_WEIGHTS_NAME)
bin_file = os.path.join(sub_file, WEIGHTS_NAME)
if os.path.isdir(sub_file) and os.path.isfile(config_file):
# diffusion or cond
cfg = json.load(open(config_file, 'r'))
if 'cond_stage_model.' in cfg['target_modules']:
cond_cfg = copy.deepcopy(cfg)
if 'cond_stage_model.*' in cond_cfg[
'target_modules']:
cond_cfg['target_modules'] = cond_cfg[
'target_modules'].replace(
'cond_stage_model.*', '.*')
else:
cond_cfg['target_modules'] = cond_cfg[
'target_modules'].replace(
'cond_stage_model.', '')
if cond_cfg['target_modules'].startswith('*'):
cond_cfg['target_modules'] = '.' + cond_cfg[
'target_modules']
os.makedirs(local_cond_cache + '_' + sub,
exist_ok=True)
cond_tuner[os.path.basename(local_cond_cache) +
'_' + sub] = hashlib.md5(
(local_cond_cache + '_' +
sub).encode('utf-8')).hexdigest()
os.makedirs(local_cond_cache + '_' + sub,
exist_ok=True)
json.dump(
cond_cfg,
open(
os.path.join(local_cond_cache + '_' + sub,
CONFIG_NAME), 'w'))
if 'model.' in cfg['target_modules'].replace(
'cond_stage_model.', ''):
diffusion_cfg = copy.deepcopy(cfg)
if 'model.*' in diffusion_cfg['target_modules']:
diffusion_cfg[
'target_modules'] = diffusion_cfg[
'target_modules'].replace(
'model.*', '.*')
else:
diffusion_cfg[
'target_modules'] = diffusion_cfg[
'target_modules'].replace(
'model.', '')
if diffusion_cfg['target_modules'].startswith('*'):
diffusion_cfg[
'target_modules'] = '.' + diffusion_cfg[
'target_modules']
os.makedirs(local_diffusion_cache + '_' + sub,
exist_ok=True)
diffusion_tuner[
os.path.basename(local_diffusion_cache) + '_' +
sub] = hashlib.md5(
(local_diffusion_cache + '_' +
sub).encode('utf-8')).hexdigest()
json.dump(
diffusion_cfg,
open(
os.path.join(
local_diffusion_cache + '_' + sub,
CONFIG_NAME), 'w'))
state_dict = {}
is_bin_file = True
if os.path.isfile(bin_file):
state_dict = torch.load(bin_file)
elif os.path.isfile(safe_file):
is_bin_file = False
from safetensors.torch import \
load_file as safe_load_file
state_dict = safe_load_file(
safe_file,
device='cuda'
if torch.cuda.is_available() else 'cpu')
save_diffusion_state_dict = {}
save_cond_state_dict = {}
for key, value in state_dict.items():
if key.startswith('model.'):
save_diffusion_state_dict[
key[len('model.'):].replace(
sub,
os.path.basename(local_diffusion_cache)
+ '_' + sub)] = value
elif key.startswith('cond_stage_model.'):
save_cond_state_dict[
key[len('cond_stage_model.'):].replace(
sub,
os.path.basename(local_cond_cache) +
'_' + sub)] = value
if is_bin_file:
if len(save_diffusion_state_dict) > 0:
torch.save(
save_diffusion_state_dict,
os.path.join(
local_diffusion_cache + '_' + sub,
WEIGHTS_NAME))
if len(save_cond_state_dict) > 0:
torch.save(
save_cond_state_dict,
os.path.join(local_cond_cache + '_' + sub,
WEIGHTS_NAME))
else:
from safetensors.torch import \
save_file as safe_save_file
if len(save_diffusion_state_dict) > 0:
safe_save_file(
save_diffusion_state_dict,
os.path.join(
local_diffusion_cache + '_' + sub,
SAFETENSORS_WEIGHTS_NAME),
metadata={'format': 'pt'})
if len(save_cond_state_dict) > 0:
safe_save_file(
save_cond_state_dict,
os.path.join(local_cond_cache + '_' + sub,
SAFETENSORS_WEIGHTS_NAME),
metadata={'format': 'pt'})
json.dump(
{
'diffusion_tuner': diffusion_tuner,
'cond_tuner': cond_tuner
}, open(meta_file, 'w'))
else:
meta_conf = json.load(open(meta_file, 'r'))
diffusion_tuner = meta_conf['diffusion_tuner']
cond_tuner = meta_conf['cond_tuner']
all_diffusion_tuner.update(diffusion_tuner)
all_cond_tuner.update(cond_tuner)
if len(all_diffusion_tuner) > 0:
diffusion_model['model'] = Swift.from_pretrained(
diffusion_model['model'],
save_root_dir,
adapter_name=all_diffusion_tuner)
diffusion_model['model'].set_active_adapters(
list(all_diffusion_tuner.values()))
if len(all_cond_tuner) > 0:
cond_stage_model['model'] = Swift.from_pretrained(
cond_stage_model['model'],
save_root_dir,
adapter_name=all_cond_tuner)
cond_stage_model['model'].set_active_adapters(
list(all_cond_tuner.values()))
self.is_register = True
@@ -484,7 +484,7 @@ class DiffusionUNet(BaseModel):
if len(unexpected) > 0:
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
def _forward_origin(self, x, emb, context, hint=None):
def _forward_origin(self, x, emb, context, hint=None, **kwargs):
hs = []
h = x
for module in self.input_blocks:
@@ -492,13 +492,21 @@ class DiffusionUNet(BaseModel):
hs.append(h)
h = self.middle_block(h, emb, context)
for m_id, module in enumerate(self.output_blocks):
h = torch.cat([h, self.lsc_identity[m_id](hs.pop())], dim=1)
skip_h = hs.pop()
if 'tuner_scale' in kwargs and kwargs[
'tuner_scale'] is not None and kwargs['tuner_scale'] < 1.0:
tuner_scale = kwargs['tuner_scale']
tuner_h = self.lsc_identity[m_id](skip_h) - skip_h
h = torch.cat([h, skip_h + tuner_scale * tuner_h], dim=1)
else:
h = torch.cat([h, self.lsc_identity[m_id](skip_h)], dim=1)
target_size = hs[-1].shape[-2:] if len(hs) > 0 else None
h = module(h, emb, context, target_size)
out = self.out(h)
return out
def _forward_control(self, x, emb, context, hint, alpha=0.5):
def _forward_control(self, x, emb, context, hint, **kwargs):
control_scale = kwargs.pop('control_scale', 1.0)
multi_csc_tuners = self.control_blocks
# hints
multi_hint_hs = []
@@ -531,11 +539,11 @@ class DiffusionUNet(BaseModel):
torch.zeros_like(tuner_h),
atol=1e-6)):
# csc-tuner
skip_h_new = skip_h + multi_control_h
skip_h_new = skip_h + control_scale * multi_control_h
else:
# csc-tuner + sc-tuner
skip_h_new = skip_h + alpha * multi_control_h + (
1 - alpha) * tuner_h
tuner_scale = kwargs['tuner_scale']
skip_h_new = skip_h + control_scale * multi_control_h + tuner_scale * tuner_h
h = torch.cat([h, skip_h_new], dim=1)
target_size = hs[-1].shape[-2:] if len(hs) > 0 else None
h = module(h, emb, context, target_size)
@@ -545,7 +553,6 @@ class DiffusionUNet(BaseModel):
def forward(self, x, t=None, cond=dict(), **kwargs):
t_emb = timestep_embedding(t, self.model_channels, repeat_only=False)
emb = self.time_embed(t_emb)
hint = None
if isinstance(cond, dict):
if 'y' in cond and cond['y'] is not None:
assert self.num_classes is not None
@@ -555,15 +562,19 @@ class DiffusionUNet(BaseModel):
x = torch.cat([x, c], dim=1)
if 'hint' in cond:
hint = cond['hint']
elif 'hint' in kwargs:
hint = kwargs.pop('hint', None)
else:
hint = None
context = cond.get('crossattn', None)
else:
context = cond
hint = kwargs.pop('hint', None)
if self.control_blocks is not None:
out = self._forward_control(x, emb, context, hint)
if self.control_blocks is not None and hint is not None:
out = self._forward_control(x, emb, context, hint, **kwargs)
else:
out = self._forward_origin(x, emb, context)
out = self._forward_origin(x, emb, context, **kwargs)
return out
@staticmethod
@@ -822,7 +833,7 @@ class DiffusionUNetXL(DiffusionUNet):
conv_nd(dims, model_channels, out_channels, 3, padding=1)),
)
def _forward_origin(self, x, emb, context, hint=None):
def _forward_origin(self, x, emb, context, hint=None, **kwargs):
hs = []
h = x
for module in self.input_blocks:
@@ -830,13 +841,21 @@ class DiffusionUNetXL(DiffusionUNet):
hs.append(h)
h = self.middle_block(h, emb, context)
for m_id, module in enumerate(self.output_blocks):
h = torch.cat([h, self.lsc_identity[m_id](hs.pop())], dim=1)
skip_h = hs.pop()
if 'tuner_scale' in kwargs and kwargs[
'tuner_scale'] is not None and kwargs['tuner_scale'] < 1.0:
tuner_scale = kwargs['tuner_scale']
tuner_h = self.lsc_identity[m_id](skip_h) - skip_h
h = torch.cat([h, skip_h + tuner_scale * tuner_h], dim=1)
else:
h = torch.cat([h, self.lsc_identity[m_id](skip_h)], dim=1)
target_size = hs[-1].shape[-2:] if len(hs) > 0 else None
h = module(h, emb, context, target_size)
out = self.out(h)
return out
def _forward_control(self, x, emb, context, hint, alpha=0.5):
def _forward_control(self, x, emb, context, hint, **kwargs):
control_scale = kwargs.pop('control_scale', 1.0)
multi_csc_tuners = self.control_blocks
# hints
multi_hint_hs = []
@@ -869,11 +888,11 @@ class DiffusionUNetXL(DiffusionUNet):
torch.zeros_like(tuner_h),
atol=1e-6)):
# csc-tuner
skip_h_new = skip_h + multi_control_h
skip_h_new = skip_h + control_scale * multi_control_h
else:
# csc-tuner + sc-tuner
skip_h_new = skip_h + alpha * multi_control_h + (
1 - alpha) * tuner_h
tuner_scale = kwargs['tuner_scale']
skip_h_new = skip_h + control_scale * multi_control_h + tuner_scale * tuner_h
h = torch.cat([h, skip_h_new], dim=1)
target_size = hs[-1].shape[-2:] if len(hs) > 0 else None
h = module(h, emb, context, target_size)
@@ -886,7 +905,6 @@ class DiffusionUNetXL(DiffusionUNet):
repeat_only=False,
legacy=True)
emb = self.time_embed(t_emb)
hint = None
if isinstance(cond, dict):
if 'y' in cond:
assert self.num_classes is not None
@@ -896,15 +914,19 @@ class DiffusionUNetXL(DiffusionUNet):
x = torch.cat([x, c], dim=1)
if 'hint' in cond:
hint = cond['hint']
elif 'hint' in kwargs:
hint = kwargs.pop('hint', None)
else:
hint = None
context = cond.get('crossattn', None)
else:
context = cond
hint = kwargs.pop('hint', None)
if self.control_blocks is not None:
out = self._forward_control(x, emb, context, hint)
if self.control_blocks is not None and hint is not None:
out = self._forward_control(x, emb, context, hint, **kwargs)
else:
out = self._forward_origin(x, emb, context)
out = self._forward_origin(x, emb, context, **kwargs)
return out
def convert_to_fp16(self):
@@ -11,6 +11,7 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange, repeat
from packaging import version
from scepter.modules.model.utils.basic_utils import checkpoint, default, exists
@@ -24,6 +25,12 @@ except Exception as e:
if find_loader('flash_attn'):
FLASH_ATTN_IS_AVAILABLE = True
import flash_attn
if (not hasattr(flash_attn, '__version__')) or (version.parse(
flash_attn.__version__) < version.parse('2.0')):
from flash_attn.flash_attn_interface import flash_attn_unpadded_kvpacked_func
else:
from flash_attn.flash_attn_interface import flash_attn_varlen_kvpacked_func as flash_attn_unpadded_kvpacked_func
else:
FLASH_ATTN_IS_AVAILABLE = False
@@ -607,8 +614,6 @@ class FlashattnMultiHeadAttention(nn.Module):
and self.head_dim % 8 == 0 and self.head_dim <= 128
and self.flash_dtype is not None):
# flash implementation
from flash_attn.flash_attn_interface import \
flash_attn_unpadded_kvpacked_func
dtype = q.dtype
if dtype != self.flash_dtype:
q = q.type(self.flash_dtype)
+7 -1
View File
@@ -1,5 +1,6 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import warnings
from collections import OrderedDict
from contextlib import nullcontext
from typing import Dict
@@ -11,7 +12,6 @@ import torch.nn as nn
import torch.utils.dlpack
from einops import rearrange
from torch.utils.checkpoint import checkpoint
from transformers import CLIPTextModel, CLIPTokenizer
# to check
from scepter.modules.model.backbone.unet.unet_utils import Timestep
@@ -23,6 +23,12 @@ from scepter.modules.utils.file_system import FS
from .base_embedder import BaseEmbedder
try:
from transformers import CLIPTextModel, CLIPTokenizer
except Exception as e:
warnings.warn(
f'Import transformers error, please deal with this problem: {e}')
def autocast(f, enabled=True):
def do_autocast(*args, **kwargs):
@@ -54,7 +54,8 @@ class GaussianDiffusion(object):
guide_rescale=None,
clamp=None,
percentile=None,
cat_uc=False):
cat_uc=False,
**kwargs):
"""
Apply one step of denoising from the posterior distribution q(x_s | x_t, x0).
Since x0 is not available, estimate the denoising results using the learned
@@ -79,7 +80,7 @@ class GaussianDiffusion(object):
# prediction
if guide_scale is None:
assert isinstance(model_kwargs, dict)
out = model(xt, t=t, **model_kwargs)
out = model(xt, t=t, **model_kwargs, **kwargs)
else:
# classifier-free guidance (arXiv:2207.12598)
# model_kwargs[0]: conditional kwargs
@@ -87,7 +88,7 @@ class GaussianDiffusion(object):
assert isinstance(model_kwargs, list) and len(model_kwargs) == 2
if guide_scale == 1.:
out = model(xt, t=t, **model_kwargs[0])
out = model(xt, t=t, **model_kwargs[0], **kwargs)
else:
if cat_uc:
@@ -111,11 +112,12 @@ class GaussianDiffusion(object):
all_model_kwargs[key], value)
all_out = model(xt.repeat(2, 1, 1, 1),
t=t.repeat(2),
**all_model_kwargs)
**all_model_kwargs,
**kwargs)
y_out, u_out = all_out.chunk(2)
else:
y_out = model(xt, t=t, **model_kwargs[0])
u_out = model(xt, t=t, **model_kwargs[1])
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
u_out = model(xt, t=t, **model_kwargs[1], **kwargs)
out = u_out + guide_scale * (y_out - u_out)
# rescale the output according to arXiv:2305.08891
@@ -262,7 +264,8 @@ class GaussianDiffusion(object):
guide_rescale,
clamp,
percentile,
cat_uc=cat_uc)[-2]
cat_uc=cat_uc,
**kwargs)[-2]
# collect intermediate outputs
if return_intermediate == 'xt':
@@ -467,8 +470,8 @@ class GaussianDiffusion(object):
t = self._sigma_to_t(sigma).repeat(len(xt)).round().long()
x0 = self.denoise(xt * c_in, t, None, model, model_kwargs,
guide_scale, guide_rescale, clamp,
percentile)[-2]
guide_scale, guide_rescale, clamp, percentile,
**kwargs)[-2]
# collect intermediate outputs
if return_intermediate == 'xt':
intermediates.append(xt)
+2 -1
View File
@@ -16,7 +16,8 @@ from scepter.modules.transform.io import (LoadCvImageFromFile,
from scepter.modules.transform.io_video import (DecodeVideoToTensor,
LoadVideoFromFile)
from scepter.modules.transform.registry import TRANSFORMS, build_pipeline
from scepter.modules.transform.tensor import Rename, Select, ToNumpy, ToTensor
from scepter.modules.transform.tensor import (Rename, RenameMeta, Select,
TemplateStr, ToNumpy, ToTensor)
from scepter.modules.transform.transform_xl import FlexibleCropXL
from scepter.modules.transform.video import (AutoResizedCropVideo,
CenterCropVideo, NormalizeVideo,
+84
View File
@@ -252,3 +252,87 @@ class TensorToGPU(object):
__class__.__name__,
para_dict,
set_name=True)
@TRANSFORMS.register_class()
class RenameMeta(object):
def __init__(self, cfg, logger=None):
self.input_key = cfg.INPUT_KEY
self.output_key = cfg.OUTPUT_KEY
self.force = cfg.get('FORCE', False)
def __call__(self, item):
if 'meta' in item:
data = {}
for idx, key in enumerate(self.input_key):
data[self.output_key[idx]] = item['meta'][key]
if not self.force:
have_key_set = set(self.input_key)
else:
have_key_set = set(self.input_key + self.output_key)
for k, v in item['meta'].items():
if k not in have_key_set:
data[k] = v
item['meta'] = data
return item
@staticmethod
def get_config_template():
'''
{ "ENV" :
{ "description" : "",
"A" : {
"value": 1.0,
"description": ""
}
}
}
:return:
'''
para_dict = [{
'INPUT_KEY': {
'value': [],
'description':
'The keys need to rename, the other keys are outputed by default.'
},
'OUTPUT_KEY': {
'value': [],
'description':
'The keys need to rename, the other keys are outputed by default.'
}
}]
return dict_to_yaml('TRANSFORM',
__class__.__name__,
para_dict,
set_name=True)
@TRANSFORMS.register_class()
class TemplateStr(object):
def __init__(self, cfg, logger=None):
self.template_str = cfg.get('TEMPLATE_STR', '')
self.meta_template_str = cfg.get('META_TEMPLATE_STR', '')
def __call__(self, item):
if self.template_str != '':
for key, val in item.items():
if isinstance(val, str) and f'{{{key}}}' in self.template_str:
template = self.template_str
val = template.replace(f'{{{key}}}', val)
item[key] = val
if self.meta_template_str != '' and 'meta' in item:
for key, val in item['meta'].items():
if isinstance(val,
str) and f'{{{key}}}' in self.meta_template_str:
template = self.meta_template_str
val = template.replace(f'{{{key}}}', val)
item['meta'][key] = val
return item
@staticmethod
def get_config_template():
para_dict = [{}]
return dict_to_yaml('TRANSFORM',
__class__.__name__,
para_dict,
set_name=True)
+1
View File
@@ -440,6 +440,7 @@ class Workenv(object):
def set_env(self, we_env):
for k, v in we_env.items():
setattr(self, k, v)
set_random_seed(self.seed)
def __str__(self):
environ_str = f'Now running in the distributed environment with world size {self.world_size}\n!'
+1 -1
View File
@@ -4,10 +4,10 @@ import io
from io import BytesIO
import onnx
import onnxruntime
import torch
from torch.onnx import OperatorExportTypes
import onnxruntime
from scepter.modules.utils.distribute import we
type_map = {
@@ -2,5 +2,6 @@
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.utils.file_clients.aliyun_oss_fs import AliyunOssFs
from scepter.modules.utils.file_clients.http_fs import HttpFs
from scepter.modules.utils.file_clients.huggingface_fs import HuggingfaceFs
from scepter.modules.utils.file_clients.local_fs import LocalFs
from scepter.modules.utils.file_clients.modelscope_fs import ModelscopeFs
@@ -490,15 +490,88 @@ class AliyunOssFs(BaseFs):
meta_dict=copy.deepcopy(meta_dict)))
return meta_dict
def _get_dir_multi(self,
target_path,
local_path,
wait_finish=False,
meta_dict={}):
local_path = local_path.replace('/./', '/')
os.makedirs(local_path, exist_ok=True)
generator = self.walk_dir(target_path)
single_file_name = []
for file_name in generator:
if file_name == target_path or file_name == target_path + '/':
continue
local_file_name = os.path.join(
local_path,
file_name.split(target_path)[-1]).replace('/./', '/')
if not self.isdir(file_name):
single_file_name.append((file_name, local_file_name))
else:
meta_dict.update(
self._get_dir_multi(file_name,
local_file_name,
meta_dict=copy.deepcopy(meta_dict)))
data_quene = queue.Queue()
batch_size = 20
R = threading.Lock()
def get_one_object(target_path_list):
if isinstance(target_path_list, tuple):
target_path_list = [target_path_list]
for target_path, local_path in target_path_list:
if self.exists(target_path):
etag, size = self.get_meta(target_path)
if local_path in meta_dict and meta_dict[
local_path] == etag:
continue
process_msg(f'Download {target_path} to {local_path}....')
local_path = self.get_object_to_local_file(
target_path, local_path, wait_finish=wait_finish)
assert local_path is not None
meta_dict[target_path] = etag
else:
local_path = None
R.acquire()
try:
data_quene.put_nowait([target_path, local_path])
except Exception:
R.release()
R.release()
while True:
batch_list = single_file_name[:10 * batch_size]
if len(batch_list) < 1:
break
single_file_name = single_file_name[10 * batch_size:]
threading_list = []
for i in range(batch_size):
cur_batch = batch_list[i::batch_size]
if isinstance(cur_batch, tuple):
cur_batch = [cur_batch]
t = threading.Thread(target=get_one_object, args=(cur_batch, ))
t.daemon = True
t.start()
threading_list.append(t)
[threading_t.join() for threading_t in threading_list]
file_dict = {}
while not data_quene.empty():
target_path, local_path = data_quene.get_nowait()
file_dict[target_path] = local_path
return meta_dict
def get_dir_to_local_dir(self,
target_path,
local_path=None,
wait_finish=False,
timeout=3600,
multi_thread=False,
worker_id=-1) -> Optional[str]:
if not self.isdir(target_path):
self.logger.info(
f"{target_path} is not directory or doesn't exist.")
return None
if not target_path.endswith('/'):
target_path += '/'
if local_path is None:
@@ -533,9 +606,14 @@ class AliyunOssFs(BaseFs):
meta_dict = json.load(open(check_file, 'r'))
else:
meta_dict = {}
meta_dict = self._get_dir(target_path,
local_path=local_path,
meta_dict=copy.deepcopy(meta_dict))
if multi_thread:
meta_dict = self._get_dir_multi(target_path,
local_path=local_path,
meta_dict=copy.deepcopy(meta_dict))
else:
meta_dict = self._get_dir(target_path,
local_path=local_path,
meta_dict=copy.deepcopy(meta_dict))
json.dump(meta_dict, open(check_file, 'w'))
if is_tmp:
self.add_temp_file(local_path)
@@ -938,16 +1016,67 @@ class AliyunOssFs(BaseFs):
continue
yield osp.join(self._prefix, obj.key)
def put_dir_from_local_dir(self, local_dir, target_dir) -> bool:
def put_dir_from_local_dir(self,
local_dir,
target_dir,
multi_thread=False) -> bool:
singe_file_names = []
for folder, sub_folders, files in os.walk(local_dir):
for file in files:
file_abs_path = osp.join(folder, file)
file_rel_path = osp.relpath(file_abs_path, local_dir)
target_path = osp.join(target_dir, file_rel_path)
singe_file_names.append((file_abs_path, target_path))
if not multi_thread:
for file_abs_path, target_path in singe_file_names:
status = self.put_object_from_local_file(
file_abs_path, target_path)
if not status:
return False
else:
data_quene = queue.Queue()
R = threading.Lock()
batch_size = 20
def put_one_object(target_path_list):
if isinstance(target_path_list, tuple):
target_path_list = [target_path_list]
for local_path, target_path in target_path_list:
if local_path is None or target_path is None:
flg = False
elif os.path.exists(local_path):
flg = self.put_object_from_local_file(
local_path, target_path)
else:
flg = False
R.acquire()
try:
data_quene.put_nowait([local_path, target_path, flg])
except Exception:
R.release()
R.release()
while True:
batch_list = singe_file_names[:10 * batch_size]
if len(batch_list) < 1:
break
singe_file_names = singe_file_names[10 * batch_size:]
threading_list = []
for i in range(batch_size):
cur_batch = batch_list[i::batch_size]
if isinstance(cur_batch, tuple):
cur_batch = [cur_batch]
t = threading.Thread(target=put_one_object,
args=(cur_batch, ))
t.daemon = True
t.start()
threading_list.append(t)
[threading_t.join() for threading_t in threading_list]
while not data_quene.empty():
local_path, target_path, flg = data_quene.get_nowait()
if not flg:
return False
return True
def size(self, target_path) -> Optional[int]:
+13 -1
View File
@@ -99,7 +99,10 @@ class HttpFs(BaseFs):
def walk_dir(self, file_dir, recurse=True):
raise NotImplementedError
def put_dir_from_local_dir(self, local_dir, target_dir) -> bool:
def put_dir_from_local_dir(self,
local_dir,
target_dir,
multi_thread=False) -> bool:
raise NotImplementedError
def size(self, target_path) -> Optional[int]:
@@ -119,6 +122,15 @@ class HttpFs(BaseFs):
delimiter=None) -> (Union[bytes, str, None], Optional[int]):
raise NotImplementedError
def get_dir_to_local_dir(self,
target_path,
local_path=None,
wait_finish=False,
multi_thread=False,
timeout=3600,
worker_id=0) -> Optional[str]:
raise NotImplementedError
def get_url(self, target_path, lifecycle=3600 * 100):
return target_path
@@ -0,0 +1,202 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os
import os.path as osp
import urllib.parse as parse
import urllib.request
from typing import Optional, Union
from scepter.modules.utils.file_clients.base_fs import BaseFs
from scepter.modules.utils.file_clients.registry import FILE_SYSTEMS
@FILE_SYSTEMS.register_class()
class HuggingfaceFs(BaseFs):
para_dict = {
'RETRY_TIMES': {
'value': 10,
'description': 'Retry get object times.'
}
}
para_dict.update(BaseFs.para_dict)
def __init__(self, cfg, logger):
super(HuggingfaceFs, self).__init__(cfg, logger=logger)
retry_times = cfg.get('RETRY_TIMES', 10)
self._retry_times = retry_times
def get_prefix(self) -> str:
return 'hf://'
def support_write(self) -> bool:
return False
def support_link(self) -> bool:
return False
def basename(self, target_path) -> str:
url = parse.unquote(target_path)
url = url.split('?')[0]
return osp.basename(url)
def get_object_to_local_file(self,
target_path,
local_path=None,
wait_finish=False) -> Optional[str]:
from huggingface_hub import hf_hub_download
key = osp.relpath(target_path, self.get_prefix())
key, file_path = key.split('@', 1)
if ':' in key:
key, revision = key.split(':', 1)
else:
revision = None
if local_path is None:
local_path, is_tmp = self.map_to_local(key)
else:
is_tmp = False
if revision is not None:
local_path = local_path + '_' + str(revision)
retry = 0
while retry < self._retry_times:
try:
local_path = hf_hub_download(repo_id=key,
revision=revision,
filename=file_path,
cache_dir=local_path)
if osp.exists(local_path):
break
except Exception:
retry += 1
if retry >= self._retry_times:
return None
if is_tmp:
self.add_temp_file(local_path)
return local_path
def get_dir_to_local_dir(self,
target_path,
local_path=None,
wait_finish=False,
timeout=3600,
worker_id=-1) -> Optional[str]:
from huggingface_hub import snapshot_download
assert target_path.startswith(self.get_prefix())
key = osp.relpath(target_path, self.get_prefix())
if '@' not in key:
key, ret_folder = key.split('@', 1)[0], ''
else:
at_level_folder = key.split('@')
if len(at_level_folder) > 2:
raise f'Target path should include only one @, but you give {len(at_level_folder)} @.'
key, ret_folder = at_level_folder
if ':' in key:
key, revision = key.split(':', 1)
else:
revision = None
if local_path is None:
local_path, is_tmp = self.map_to_local(key)
else:
is_tmp = False
if revision is not None:
local_path = local_path + '_' + str(revision)
retry = 0
while retry < self._retry_times:
try:
local_path = snapshot_download(repo_id=key,
revision=revision,
cache_dir=local_path)
if osp.exists(local_path):
break
except Exception:
retry += 1
if retry >= self._retry_times:
return None
if is_tmp:
self.add_temp_file(local_path)
if not ret_folder == '':
local_path = os.path.join(local_path, ret_folder)
return local_path
def get_object(self, target_path):
try:
local_data = open(self.get_object_to_local_file(target_path),
'rb').read()
except Exception as e:
self.logger.error(f'Read {target_path} error {e}')
local_data = None
return local_data
def put_object(self, local_data, target_path):
raise NotImplementedError
def put_object_from_local_file(self, local_path, target_path) -> bool:
raise NotImplementedError
def make_link(self, target_link_path, target_path) -> bool:
raise NotImplementedError
def make_dir(self, target_dir) -> bool:
raise NotImplementedError
def remove(self, target_path) -> bool:
raise NotImplementedError
def get_logging_handler(self, target_logging_path):
raise NotImplementedError
def walk_dir(self, file_dir, recurse=True):
raise NotImplementedError
def put_dir_from_local_dir(self, local_dir, target_dir) -> bool:
raise NotImplementedError
def size(self, target_path) -> Optional[int]:
raise NotImplementedError
def get_object_chunk_list(self,
target_path,
chunk_num=1,
delimiter=None) -> Optional[list]:
raise NotImplementedError
def get_object_stream(
self,
target_path,
start,
size=10000,
delimiter=None) -> (Union[bytes, str, None], Optional[int]):
raise NotImplementedError
def get_url(self, target_path, lifecycle=3600 * 100):
return target_path
def exists(self, target_path) -> bool:
req = urllib.request.Request(target_path)
req.get_method = lambda: 'HEAD'
try:
urllib.request.urlopen(req)
return True
except Exception:
return False
def isfile(self, target_path) -> bool:
# Well for a http url, it should only be a file.
return True
def isdir(self, target_path) -> bool:
return False
@@ -85,6 +85,7 @@ class LocalFs(BaseFs):
target_path,
local_path=None,
wait_finish=False,
multi_thread=False,
timeout=3600,
worker_id=0) -> Optional[str]:
if not self.isdir(target_path):
@@ -294,7 +295,10 @@ class LocalFs(BaseFs):
def get_logging_handler(self, target_logging_path):
return logging.FileHandler(target_logging_path)
def put_dir_from_local_dir(self, local_dir, target_dir) -> bool:
def put_dir_from_local_dir(self,
local_dir,
target_dir,
multi_thread=False) -> bool:
local_dir = self.reconstruct_path(local_dir)
target_dir = self.reconstruct_path(target_dir)
if local_dir == target_dir:
@@ -24,6 +24,8 @@ class ModelscopeFs(BaseFs):
super(ModelscopeFs, self).__init__(cfg, logger=logger)
retry_times = cfg.get('RETRY_TIMES', 10)
self._retry_times = retry_times
self._model_id_loaded = set()
self._model_file_loaded = set()
def get_prefix(self) -> str:
return 'ms://'
@@ -64,10 +66,16 @@ class ModelscopeFs(BaseFs):
retry = 0
while retry < self._retry_times:
try:
local_path = model_file_download(model_id=key,
revision=revision,
file_path=file_path,
cache_dir=local_path)
model_file = os.path.join(key, file_path)
if model_file in self._model_file_loaded:
local_path = os.path.join(local_path, model_file)
if not osp.exists(local_path):
self._model_file_loaded.remove(key)
else:
local_path = model_file_download(model_id=key,
revision=revision,
file_path=file_path,
cache_dir=local_path)
if osp.exists(local_path):
break
except Exception:
@@ -75,7 +83,7 @@ class ModelscopeFs(BaseFs):
if retry >= self._retry_times:
return None
self._model_file_loaded.add(model_file)
if is_tmp:
self.add_temp_file(local_path)
return local_path
@@ -84,6 +92,7 @@ class ModelscopeFs(BaseFs):
target_path,
local_path=None,
wait_finish=False,
multi_thread=False,
timeout=3600,
worker_id=-1) -> Optional[str]:
from modelscope.hub.snapshot_download import snapshot_download
@@ -114,9 +123,14 @@ class ModelscopeFs(BaseFs):
retry = 0
while retry < self._retry_times:
try:
local_path = snapshot_download(key,
revision=revision,
cache_dir=local_path)
if key in self._model_id_loaded:
local_path = os.path.join(local_path, key)
if not osp.exists(local_path):
self._model_id_loaded.remove(key)
else:
local_path = snapshot_download(key,
revision=revision,
cache_dir=local_path)
if osp.exists(local_path):
break
except Exception:
@@ -125,6 +139,7 @@ class ModelscopeFs(BaseFs):
if retry >= self._retry_times:
return None
self._model_id_loaded.add(key)
if is_tmp:
self.add_temp_file(local_path)
if not ret_folder == '':
@@ -161,7 +176,10 @@ class ModelscopeFs(BaseFs):
def walk_dir(self, file_dir, recurse=True):
raise NotImplementedError
def put_dir_from_local_dir(self, local_dir, target_dir) -> bool:
def put_dir_from_local_dir(self,
local_dir,
target_dir,
multi_thread=False) -> bool:
raise NotImplementedError
def size(self, target_path) -> Optional[int]:
+9 -2
View File
@@ -129,12 +129,14 @@ class FileSystem(object):
local_path=None,
wait_finish=False,
timeout=3600,
multi_thread=False,
worker_id=0):
with self.get_fs_client(target_path) as client:
local_path = client.get_dir_to_local_dir(target_path,
local_path=local_path,
wait_finish=wait_finish,
timeout=timeout,
multi_thread=multi_thread,
worker_id=worker_id)
if local_path is None:
raise ReadException(
@@ -189,7 +191,10 @@ class FileSystem(object):
local_path, is_tmp = client.map_to_local(target_path)
return local_path, is_tmp
def put_dir_from_local_dir(self, local_dir, target_dir):
def put_dir_from_local_dir(self,
local_dir,
target_dir,
multi_thread=False):
""" Upload all contents in local_dir to target_dir, keep the file tree.
Args:
@@ -200,7 +205,9 @@ class FileSystem(object):
Bool.
"""
with self.get_fs_client(target_dir) as client:
return client.put_dir_from_local_dir(local_dir, target_dir)
return client.put_dir_from_local_dir(local_dir,
target_dir,
multi_thread=multi_thread)
def walk_dir(self, target_dir, recurse=True):
""" Iterator to access the files of target dir.
+2
View File
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
+2
View File
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import gradio as gr
+2
View File
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
class DescUIName():
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import gradio as gr
from scepter.studio.home.home_ui.component_names import DescUIName
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import gradio as gr
from scepter.studio.home.home_ui.component_names import GuideUIName
+2
View File
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
+60 -71
View File
@@ -1,5 +1,7 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os
from collections import OrderedDict
from glob import glob
import gradio as gr
@@ -20,6 +22,9 @@ from scepter.studio.inference.inference_ui.refiner_ui import RefinerUI
from scepter.studio.inference.inference_ui.tuner_ui import TunerUI
from scepter.studio.utils.env import init_env
UI_MAP = [('diffusion', DiffusionUI), ('mantra', MantraUI), ('tuner', TunerUI),
('control', ControlUI), ('refiner', RefinerUI)]
class InferenceUI():
def __init__(self,
@@ -71,94 +76,78 @@ class InferenceUI():
pipe_manager,
is_debug=is_debug,
language=language)
self.diffusion_ui = DiffusionUI(cfg_general,
pipe_manager,
is_debug=is_debug,
language=language)
self.mantra_ui = MantraUI(cfg_general,
pipe_manager,
is_debug=is_debug,
language=language)
self.tuner_ui = TunerUI(cfg_general,
pipe_manager,
is_debug=is_debug,
language=language)
self.refiner_ui = RefinerUI(cfg_general,
pipe_manager,
is_debug=is_debug,
language=language)
self.control_ui = ControlUI(cfg_general,
pipe_manager,
is_debug=is_debug,
language=language)
self.component_names = InferenceUIName(language=language)
self.tab_ui = OrderedDict()
self.tab_ui_kwargs = OrderedDict()
for name, UI in UI_MAP:
ui = UI(cfg_general,
pipe_manager,
is_debug=is_debug,
language=language)
self.tab_ui[name] = ui
self.tab_ui_kwargs[f'{name}_ui'] = ui
self.__setattr__(f'{name}_ui', ui)
self.check_box_controlled_tabs = ['mantra', 'tuner', 'control']
assert len(self.component_names.check_box_for_setting) == len(
self.check_box_controlled_tabs)
def create_ui(self):
# create model
self.model_manage_ui.create_ui()
self.gallery_ui.create_ui()
with gr.Row(variant='panel', equal_height=True):
self.check_box_for_setting = gr.CheckboxGroup(
choices=self.component_names.check_box_for_setting,
show_label=False)
# create tabs
def create_tab(name, ui):
label = getattr(self.component_names, f'{name}_paras')
if name in ['refiner']:
ui.create_ui()
else:
with gr.TabItem(label=label, id=f'{name}_ui'):
ui.create_ui()
with gr.Row(variant='panel', equal_height=True):
with gr.Accordion(label=self.component_names.advance_block_name,
open=True):
self.check_box_for_setting = gr.CheckboxGroup(
choices=self.component_names.check_box_for_setting,
show_label=False)
with gr.Tabs() as self.setting_tab:
with gr.TabItem(label=self.component_names.diffusion_paras,
id='diffusion_ui'):
self.diffusion_ui.create_ui()
# 0
with gr.TabItem(label=self.component_names.mantra_paras,
id='mantra_ui',
visible=True) as self.mantra_tab:
self.mantra_ui.create_ui()
self.mantra_state = gr.State(value=False)
# 1
with gr.TabItem(label=self.component_names.tuner_paras,
id='tuner_ui',
visible=True) as self.tuner_tab:
self.tuner_ui.create_ui()
self.tuner_state = gr.State(value=False)
# 2
with gr.TabItem(label=self.component_names.contrl_paras,
id='control_ui',
visible=True) as self.control_tab:
self.control_ui.create_ui()
self.control_state = gr.State(value=False)
# 3
with gr.TabItem(label=self.component_names.refine_paras,
id='refiner_ui',
visible=True) as self.refine_tab:
self.refiner_ui.create_ui()
for name, ui in self.tab_ui.items():
create_tab(name, ui)
def set_callbacks(self, manager):
self.model_manage_ui.set_callbacks(self.diffusion_ui, self.tuner_ui,
self.control_ui, self.mantra_ui)
self.model_manage_ui.set_callbacks(**self.tab_ui_kwargs)
self.gallery_ui.set_callbacks(self, self.model_manage_ui,
self.diffusion_ui, self.mantra_ui,
self.tuner_ui, self.refiner_ui,
self.control_ui)
self.diffusion_ui.set_callbacks(self.model_manage_ui)
self.mantra_ui.set_callbacks(self.model_manage_ui)
self.tuner_ui.set_callbacks(self.model_manage_ui)
self.control_ui.set_callbacks(self.model_manage_ui, self.diffusion_ui)
self.refiner_ui.set_callbacks()
**self.tab_ui_kwargs)
for name, ui in self.tab_ui_kwargs.items():
ui.set_callbacks(self.model_manage_ui,
**self.tab_ui_kwargs,
gallery_ui=self.gallery_ui)
def change_setting_tab(check_box):
mantra_ui, tuner_ui, control_ui = False, False, False
def change_setting_tab(check_box, *args):
selected_tab = 'diffusion_ui'
ui_tabs_state = [False] * len(args)
for key in check_box:
if self.component_names.check_box_for_setting.index(key) == 0:
mantra_ui = True
if self.component_names.check_box_for_setting.index(key) == 1:
tuner_ui = True
if self.component_names.check_box_for_setting.index(key) == 2:
control_ui = True
return (mantra_ui, tuner_ui, control_ui)
i = self.component_names.check_box_for_setting.index(key)
ui_tabs_state[i] = True
if ui_tabs_state[i] != args[i]:
selected_tab = self.check_box_controlled_tabs[i] + '_ui'
ui_tabs_updates = [gr.update(visible=v) for v in ui_tabs_state]
return gr.update(
selected=selected_tab), *ui_tabs_state, *ui_tabs_updates
gr_states = [
self.tab_ui[name].state for name in self.check_box_controlled_tabs
]
gr_tabs = [
self.tab_ui[name].tab for name in self.check_box_controlled_tabs
]
self.check_box_for_setting.change(
change_setting_tab,
inputs=[self.check_box_for_setting],
outputs=[self.mantra_state, self.tuner_state, self.control_state],
inputs=[self.check_box_for_setting, *gr_states],
outputs=[self.setting_tab, *gr_states, *gr_tabs],
queue=False)
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.inference.diffusion_inference import DiffusionInference
from scepter.modules.utils.logger import get_logger
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
@@ -1,5 +1,21 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
# For dataset manager
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.directory import get_md5
def download_image(image):
if image is not None:
client = FS.get_fs_client(image)
if client.tmp_dir.startswith("/home"):
name = get_md5(image)
local_path = FS.get_from(image, f"/tmp/gradio/scepter_examples/{name}")
else:
local_path = FS.get_from(image)
return local_path
else:
return image
class InferenceUIName():
@@ -12,16 +28,16 @@ class InferenceUIName():
self.diffusion_paras = 'Generation Setting'
self.mantra_paras = 'Mantra Book'
self.tuner_paras = 'Tuners'
self.contrl_paras = 'Controlable Generation'
self.refine_paras = 'Refiner Setting'
self.control_paras = 'Controlable Generation'
self.refiner_paras = 'Refiner Setting'
elif language == 'zh':
self.advance_block_name = '生成选项'
self.check_box_for_setting = ['使用咒语', '使用微调', '使用控制']
self.diffusion_paras = '生成参数设置'
self.mantra_paras = '咒语书'
self.tuner_paras = '微调模型'
self.contrl_paras = '可控生成'
self.refine_paras = 'Refine设置'
self.control_paras = '可控生成'
self.refiner_paras = 'Refine设置'
class ModelManageUIName():
@@ -56,12 +72,14 @@ class GalleryUIName():
self.gallery_before_refine_output = 'Before Refine Output'
self.prompt_input = 'Type Prompt Here'
self.generate = 'Generate'
self.control_err1 = 'The conditional image does not exist'
elif language == 'zh':
self.gallery_block_name = '模型的输出'
self.gallery_diffusion_output = '扩散输出'
self.gallery_before_refine_output = '精炼前输出'
self.prompt_input = '在此输入提示'
self.generate = '生成'
self.control_err1 = '条件图片不存在'
class DiffusionUIName():
@@ -73,7 +91,6 @@ class DiffusionUIName():
self.image_number = 'Images Number'
self.resolutions_height = 'Output Height'
self.resolutions_width = 'Output Width'
self.negative_prompt = 'Negative Prompt'
self.negative_prompt_placeholder = 'Type Prompt Here'
self.negative_prompt_description = 'Describing what you do not want to see.'
@@ -83,6 +100,12 @@ class DiffusionUIName():
self.discretization = 'Discretization'
self.random_seed = 'Use Random Seed'
self.seed = 'Used Seed'
self.example_block_name = 'Prompt Examples'
self.examples = [
'dream dandelion', 'Mount Everest', 'a boy wearing a jacket',
'Spring, Birds, Cawing, Branches',
'Cyberpunk, Maiden, Heavy Machinery', 'A Frog'
]
elif language == 'zh':
@@ -100,6 +123,12 @@ class DiffusionUIName():
self.discretization = '离散化'
self.random_seed = '使用随机种子'
self.seed = '使用的种子'
self.example_block_name = '提示词样例'
self.examples = [
'dream dandelion', 'Mount Everest', 'a boy wearing a jacket',
'Spring, Birds, Cawing, Branches',
'Cyberpunk, Maiden, Heavy Machinery', 'A Frog'
]
class MantraUIName():
@@ -117,6 +146,13 @@ class MantraUIName():
self.style_negative_template = 'Mantra Negative Prompt Template'
self.style_example = 'Mantra Results Example'
self.style_example_prompt = 'Mantra Example Prompt'
self.example_block_name = 'Examples'
self.examples = [
[['Adorable 3D Character'], 'a girl'],
[['Watercolor 2'], 'a single flower'],
[['Action Figure'],
'a close up of a small rabbit wearing a hat and scarf']
]
elif language == 'zh':
self.mantra_styles = '咒语风格'
@@ -130,6 +166,12 @@ class MantraUIName():
self.style_negative_template = '咒语负向提示模板'
self.style_example = '咒语示例图'
self.style_example_prompt = '咒语示例提示词'
self.example_block_name = '样例'
self.examples = [
[['可爱的3D角色'], 'a girl'], [['水彩'], 'a single flower'],
[['动作人偶'],
'a close up of a small rabbit wearing a hat and scarf']
]
class RefinerUIName():
@@ -169,6 +211,11 @@ class TunerUIName():
self.tuner_prompt_example = 'Prompt Example'
self.base_model = 'Base Model Name'
self.custom_tuner_model = 'Customized Model'
self.advance_block_name = 'Advance Setting'
self.tuner_scale = 'Tuner Scale'
self.example_block_name = 'Examples'
self.examples = [[['Pencil Sketch Drawing'], 'a girl in a jacket'],
[['Flat 2D Art'], 'a cat']]
elif language == 'zh':
self.tuner_model = '微调模型'
@@ -179,10 +226,88 @@ class TunerUIName():
self.tuner_prompt_example = '示例提示词'
self.base_model = '基础模型'
self.custom_tuner_model = '自定义模型'
self.advance_block_name = '高级设置'
self.tuner_scale = '微调强度'
self.example_block_name = '样例'
self.examples = [[['铅笔素描'], 'a girl in a jacket'],
[['扁平2D艺术'], 'a cat']]
class ControlUIName():
def __init__(self, language='en'):
self.examples = [
[
'Canny',
download_image(
'https://modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=examples/control/canny_turtle.jpeg' # noqa
),
'sea turtle'
],
[
'Canny',
download_image(
'https://modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=examples/control/canny_starwar.jpeg' # noqa
),
'star wars stormtrooper with weapons'
],
[
'Hed',
download_image(
'https://modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=examples/control/hed_kingfisher.jpeg' # noqa
),
'a kingfisher coming out of the water, photorealistic hyperrealistic'
],
[
'Hed',
download_image(
'https://modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=examples/control/hed_lemon.jpeg' # noqa
),
'lemon and branches, simple background'
],
[
'Openpose',
download_image(
'https://modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=examples/control/pose_panda.jpeg' # noqa
),
'panda wearing pink suite sitting on iron throne with a sword'
],
[
'Openpose',
download_image(
'https://modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=examples/control/pose_girl.jpeg' # noqa
),
'a beautiful little girl walking on the grass, pixar style characters'
],
[
'Midas',
download_image(
'https://modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=examples/control/depth_rose.jpeg' # noqa
),
'beautiful red rose'
],
[
'Midas',
download_image(
'https://modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=examples/control/depth_c.jpg' # noqa
),
'three-dimensional letter C with fire'
],
[
'Color',
download_image(
'https://modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=examples/control/color_lilypad.jpeg' # noqa
),
'lilypad flower floating on a tiny pond surrounded by ferns morning'
],
[
'Color',
download_image(
'https://modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=examples/control/color_gooseberry.jpeg' # noqa
),
'gooseberry watercolor isolated'
]
]
if language == 'en':
self.source_image = 'Source Image'
@@ -190,25 +315,30 @@ class ControlUIName():
self.control_preprocessor = 'Control Preprocessor'
self.crop_type = 'Crop type'
self.direction = (
"Note: 1) After clicking 'Extract', you can extract "
'the conditional image from the original image on the left;\n'
'2) You can also directly transfer the conditional image on the right; \n'
'3) Or simply upload the original image and directly run the\n'
"'Conditional Inference' below.")
"1) After clicking 'Extract', you can extract the conditional image from the original image on the "
"left, then click the 'Generate' button above;\n"
"2) You can also directly transfer the conditional image on the right, then click the 'Generate' "
'button above;')
self.preprocess = 'Preprocess'
self.control_model = 'Generation Model'
self.control_err1 = "Condition preprocessor doesn't exist."
self.control_err2 = 'Condition preprocessor failed.'
self.advance_block_name = 'Advance Setting'
self.control_scale = 'Control Scale'
self.example_block_name = 'Examples'
elif language == 'zh':
self.preprocess = '条件预处理'
self.source_image = '源图片'
self.cond_image = '条件图片'
self.control_preprocessor = '图像预处理器'
self.crop_type = '抠图方式'
self.direction = ('注:1)点击【Extract】后可从左侧原始图像提取出条件图像;\n'
'2)也可直接传输右侧条件图像;\n'
'3)或只上传原始图像直接运行下方的【条件推理】')
self.direction = ('1)点击【Extract】后可从左侧原始图像提取出条件图像,再点击上方运行;\n'
'2)也可直接传输右侧条件图像,再点击上方运行;\n')
self.control_err1 = '预处理器不存在。'
self.control_err2 = '预处理失败。'
self.control_model = '生成模型'
self.advance_block_name = '高级设置'
self.control_scale = '控制强度'
self.example_block_name = '样例'
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import gradio as gr
import numpy as np
@@ -67,8 +68,9 @@ class ControlUI(UIBase):
return annotator
def create_ui(self, *args, **kwargs):
gr.Markdown(self.component_names.preprocess)
with gr.Group():
self.state = gr.State(value=False)
with gr.Column(visible=False) as self.tab:
# gr.Markdown(self.component_names.preprocess)
with gr.Row():
with gr.Column(scale=1, min_width=0):
self.source_image = gr.Image(
@@ -104,7 +106,27 @@ class ControlUI(UIBase):
self.cond_button = gr.Button('Extract')
gr.Markdown(self.component_names.direction)
def set_callbacks(self, model_manage_ui, diffusion_ui):
with gr.Accordion(label=self.component_names.advance_block_name,
open=False):
self.control_scale = gr.Slider(
label=self.component_names.control_scale,
minimum=0.0,
maximum=1.0,
step=0.05,
value=1.0,
interactive=True)
self.example_block = gr.Accordion(
label=self.component_names.example_block_name, open=True)
def set_callbacks(self, model_manage_ui, diffusion_ui, **kwargs):
gallery_ui = kwargs.pop('gallery_ui')
with self.example_block:
gr.Examples(
examples=self.component_names.examples,
inputs=[self.control_mode, self.cond_image, gallery_ui.prompt],
examples_per_page=20)
def extract_condition(source_image, control_mode, crop_type,
output_height, output_width):
if control_mode not in self.controlable_annotators:
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import random
@@ -67,6 +68,7 @@ class DiffusionUI(UIBase):
def create_ui(self, *args, **kwargs):
self.cur_paras = self.get_default(self.diffusion_paras,
self.default_input)
self.example_block = gr.Row(equal_height=True, visible=True)
with gr.Row(equal_height=True):
self.negative_prompt = gr.Textbox(
label=self.component_names.negative_prompt,
@@ -152,7 +154,13 @@ class DiffusionUI(UIBase):
with gr.Column(scale=1):
self.refresh_seed = gr.Button(value=refresh_symbol)
def set_callbacks(self, model_manage_ui):
def set_callbacks(self, model_manage_ui, **kwargs):
gallery_ui = kwargs.pop('gallery_ui')
with self.example_block:
gr.Examples(label=self.component_names.example_block_name,
examples=self.component_names.examples,
inputs=gallery_ui.prompt)
def random_checked(r):
value = -1
return (gr.Row(visible=not r), gr.Textbox(value=value))
@@ -1,4 +1,7 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os
import gradio as gr
import numpy as np
from PIL import Image
@@ -11,6 +14,7 @@ class GalleryUI(UIBase):
def __init__(self, cfg, pipe_manager, is_debug=False, language='en'):
self.pipe_manager = pipe_manager
self.component_names = GalleryUIName(language)
self.cfg = cfg
def create_ui(self, *args, **kwargs):
with gr.Group():
@@ -25,7 +29,9 @@ class GalleryUI(UIBase):
with gr.Column(scale=2, min_width=0):
self.output_gallery = gr.Gallery(
label=self.component_names.gallery_diffusion_output,
value=[])
value=[],
allow_preview=True,
preview=True)
with gr.Row(elem_classes='type_row'):
with gr.Column(scale=17):
self.prompt = gr.Textbox(
@@ -45,158 +51,206 @@ class GalleryUI(UIBase):
elem_id='generate_button',
visible=True)
def set_callbacks(self, inference_ui, model_manage_ui, diffusion_ui,
mantra_ui, tuner_ui, refiner_ui, control_ui):
def generate_image(
mantra_state, tuner_state, control_state, diffusion_model,
first_stage_model, cond_stage_model, refiner_cond_model,
refiner_diffusion_model, tuner_model, custom_tuner_model,
control_model, crop_type, control_cond_image, prompt,
negative_prompt, prompt_prefix, sample, discretization,
output_height, output_width, image_number, sample_steps,
guide_scale, guide_rescale, refine_state, refine_strength,
refine_sampler, refine_discretization, refine_guide_scale,
refine_guide_rescale, style_template, style_negative_template,
image_seed):
current_pipeline = self.pipe_manager.get_pipeline_given_modules({
'diffusion_model':
diffusion_model,
'first_stage_model':
first_stage_model,
'cond_stage_model':
cond_stage_model,
'refiner_cond_model':
refiner_cond_model,
'refiner_diffusion_model':
refiner_diffusion_model
})
now_pipeline = self.pipe_manager.model_level_info[diffusion_model][
'pipeline'][0]
used_tuner_model = []
if not isinstance(tuner_model, list):
tuner_model = [tuner_model]
for tuner_m in tuner_model:
if tuner_m is None or tuner_m == '':
continue
if (now_pipeline
in self.pipe_manager.model_level_info['tuners']
and tuner_m in self.pipe_manager.
model_level_info['tuners'][now_pipeline]):
tuner_m = self.pipe_manager.model_level_info['tuners'][
now_pipeline][tuner_m]['model_info']
used_tuner_model.append(tuner_m)
used_custom_tuner_model = []
if not isinstance(custom_tuner_model, list):
custom_tuner_model = [custom_tuner_model]
for tuner_m in custom_tuner_model:
if tuner_m is None or tuner_m == '':
continue
if (now_pipeline in
self.pipe_manager.model_level_info['customized_tuners']
and tuner_m in self.pipe_manager.
model_level_info['customized_tuners'][now_pipeline]):
tuner_m = self.pipe_manager.model_level_info[
'customized_tuners'][now_pipeline][tuner_m][
'model_info']
used_custom_tuner_model.append(tuner_m)
def generate_gallery(self,
prompt,
mantra_state,
tuner_state,
control_state,
refine_state,
diffusion_model,
first_stage_model,
cond_stage_model,
refiner_cond_model,
refiner_diffusion_model,
tuner_model,
tuner_scale,
custom_tuner_model,
control_model,
control_scale,
crop_type,
control_cond_image,
negative_prompt,
prompt_prefix,
sample,
discretization,
output_height,
output_width,
image_number,
sample_steps,
guide_scale,
guide_rescale,
refine_strength,
refine_sampler,
refine_discretization,
refine_guide_scale,
refine_guide_rescale,
style_template,
style_negative_template,
image_seed,
show_jpeg_image=True):
if control_state and control_cond_image is None:
raise gr.Error(self.component_names.control_err1)
current_pipeline = self.pipe_manager.get_pipeline_given_modules({
'diffusion_model':
diffusion_model,
'first_stage_model':
first_stage_model,
'cond_stage_model':
cond_stage_model,
'refiner_cond_model':
refiner_cond_model,
'refiner_diffusion_model':
refiner_diffusion_model
})
now_pipeline = self.pipe_manager.model_level_info[diffusion_model][
'pipeline'][0]
used_tuner_model = []
if not isinstance(tuner_model, list):
tuner_model = [tuner_model]
for tuner_m in tuner_model:
if tuner_m is None or tuner_m == '':
continue
if (now_pipeline in self.pipe_manager.model_level_info['tuners']
and tuner_m in self.pipe_manager.model_level_info['tuners']
[now_pipeline]):
tuner_m = self.pipe_manager.model_level_info['tuners'][
now_pipeline][tuner_m]['model_info']
used_tuner_model.append(tuner_m)
used_custom_tuner_model = []
if not isinstance(custom_tuner_model, list):
custom_tuner_model = [custom_tuner_model]
for tuner_m in custom_tuner_model:
if tuner_m is None or tuner_m == '':
continue
if (now_pipeline
in self.pipe_manager.model_level_info['controllers']
and control_model in self.pipe_manager.
model_level_info['controllers'][now_pipeline]):
control_model = self.pipe_manager.model_level_info[
'controllers'][now_pipeline][control_model]['model_info']
in self.pipe_manager.model_level_info['customized_tuners']
and tuner_m in self.pipe_manager.
model_level_info['customized_tuners'][now_pipeline]):
tuner_m = self.pipe_manager.model_level_info[
'customized_tuners'][now_pipeline][tuner_m]['model_info']
used_custom_tuner_model.append(tuner_m)
prompt_rephrased = style_template.replace(
'{prompt}', prompt
) if not style_template == '' and mantra_state else prompt
prompt_rephrased = f'{prompt_prefix}{prompt_rephrased}' if not prompt_prefix == '' else prompt_rephrased
negative_prompt_rephrased = negative_prompt + style_negative_template if mantra_state else negative_prompt
pipeline_input = {
'prompt': prompt_rephrased,
'negative_prompt': negative_prompt_rephrased,
'sample': sample,
'sample_steps': sample_steps,
'discretization': discretization,
'original_size_as_tuple':
[int(output_height), int(output_width)],
'target_size_as_tuple':
[int(output_height), int(output_width)],
'crop_coords_top_left': [0, 0],
'guide_scale': guide_scale,
'guide_rescale': guide_rescale,
}
if refine_state:
pipeline_input['refine_sampler'] = refine_sampler
pipeline_input['refine_discretization'] = refine_discretization
pipeline_input['refine_guide_scale'] = refine_guide_scale
pipeline_input['refine_guide_rescale'] = refine_guide_rescale
else:
refine_strength = 0
results = current_pipeline(
pipeline_input,
num_samples=image_number,
intermediate_callback=None,
refine_strength=refine_strength,
img_to_img_strength=0,
tuner_model=used_tuner_model +
used_custom_tuner_model if tuner_state else None,
control_model=control_model if control_state else None,
control_cond_image=control_cond_image
if control_state else None,
crop_type=crop_type if control_state else None,
seed=int(image_seed))
images = []
before_images = []
if 'images' in results:
images_tensor = results['images'] * 255
images = [
Image.fromarray(images_tensor[idx].permute(
1, 2, 0).cpu().numpy().astype(np.uint8))
for idx in range(images_tensor.shape[0])
]
if 'before_refine_images' in results and results[
'before_refine_images'] is not None:
before_refine_images_tensor = results[
'before_refine_images'] * 255
before_images = [
Image.fromarray(before_refine_images_tensor[idx].permute(
1, 2, 0).cpu().numpy().astype(np.uint8))
for idx in range(before_refine_images_tensor.shape[0])
]
if 'seed' in results:
print(results['seed'])
print(images, before_images)
return (
gr.Column(visible=len(before_images) > 0),
before_images,
images,
)
if (now_pipeline in self.pipe_manager.model_level_info['controllers']
and control_model in self.pipe_manager.
model_level_info['controllers'][now_pipeline]):
control_model = self.pipe_manager.model_level_info['controllers'][
now_pipeline][control_model]['model_info']
self.generate_button.click(
generate_image,
inputs=[
inference_ui.mantra_state, inference_ui.tuner_state,
inference_ui.control_state, model_manage_ui.diffusion_model,
model_manage_ui.first_stage_model,
model_manage_ui.cond_stage_model,
refiner_ui.refiner_cond_model,
refiner_ui.refiner_diffusion_model, tuner_ui.tuner_model,
tuner_ui.custom_tuner_model, control_ui.control_model,
control_ui.crop_type, control_ui.cond_image, self.prompt,
diffusion_ui.negative_prompt, diffusion_ui.prompt_prefix,
diffusion_ui.sampler, diffusion_ui.discretization,
diffusion_ui.output_height, diffusion_ui.output_width,
diffusion_ui.image_number, diffusion_ui.sample_steps,
diffusion_ui.guide_scale, diffusion_ui.guide_rescale,
refiner_ui.refine_state, refiner_ui.refine_strength,
refiner_ui.refine_sampler, refiner_ui.refine_discretization,
refiner_ui.refine_guide_scale, refiner_ui.refine_guide_rescale,
mantra_ui.style_template, mantra_ui.style_negative_template,
diffusion_ui.image_seed
],
outputs=[
self.before_refine_panel, self.before_refine_gallery,
self.output_gallery
],
queue=True)
prompt_rephrased = style_template.replace(
'{prompt}',
prompt) if not style_template == '' and mantra_state else prompt
prompt_rephrased = f'{prompt_prefix}{prompt_rephrased}' if not prompt_prefix == '' else prompt_rephrased
negative_prompt_rephrased = negative_prompt + style_negative_template if mantra_state else negative_prompt
pipeline_input = {
'prompt': prompt_rephrased,
'negative_prompt': negative_prompt_rephrased,
'sample': sample,
'sample_steps': sample_steps,
'discretization': discretization,
'original_size_as_tuple': [int(output_height),
int(output_width)],
'target_size_as_tuple': [int(output_height),
int(output_width)],
'crop_coords_top_left': [0, 0],
'guide_scale': guide_scale,
'guide_rescale': guide_rescale,
}
if refine_state:
pipeline_input['refine_sampler'] = refine_sampler
pipeline_input['refine_discretization'] = refine_discretization
pipeline_input['refine_guide_scale'] = refine_guide_scale
pipeline_input['refine_guide_rescale'] = refine_guide_rescale
else:
refine_strength = 0
results = current_pipeline(
pipeline_input,
num_samples=image_number,
intermediate_callback=None,
refine_strength=refine_strength,
img_to_img_strength=0,
tuner_model=used_tuner_model +
used_custom_tuner_model if tuner_state else None,
tuner_scale=tuner_scale if tuner_state or control_state else None,
control_model=control_model if control_state else None,
control_scale=control_scale
if tuner_state or control_state else None,
control_cond_image=control_cond_image if control_state else None,
crop_type=crop_type if control_state else None,
seed=int(image_seed))
images = []
before_images = []
if 'images' in results:
images_tensor = results['images'] * 255
images = [
Image.fromarray(images_tensor[idx].permute(
1, 2, 0).cpu().numpy().astype(np.uint8))
for idx in range(images_tensor.shape[0])
]
if 'before_refine_images' in results and results[
'before_refine_images'] is not None:
before_refine_images_tensor = results['before_refine_images'] * 255
before_images = [
Image.fromarray(before_refine_images_tensor[idx].permute(
1, 2, 0).cpu().numpy().astype(np.uint8))
for idx in range(before_refine_images_tensor.shape[0])
]
if 'seed' in results:
print(results['seed'])
print(images, before_images)
if show_jpeg_image:
save_list = []
for i, img in enumerate(images):
save_image = os.path.join(self.cfg.WORK_DIR,
f'cur_gallery_{i}.jpg')
img.save(save_image)
save_list.append(save_image)
images = save_list
return (
gr.Column(visible=len(before_images) > 0),
before_images,
images,
)
def generate_image(self, *args, **kwargs):
gallery_result = self.generate_gallery(*args, **kwargs)
before_refine_panel, before_refine_gallery, output_gallery = gallery_result
return (before_refine_panel, before_refine_gallery, output_gallery[0])
def set_callbacks(self, inference_ui, model_manage_ui, diffusion_ui,
mantra_ui, tuner_ui, refiner_ui, control_ui, **kwargs):
self.gen_inputs = [
self.prompt, mantra_ui.state, tuner_ui.state, control_ui.state,
refiner_ui.state, model_manage_ui.diffusion_model,
model_manage_ui.first_stage_model,
model_manage_ui.cond_stage_model, refiner_ui.refiner_cond_model,
refiner_ui.refiner_diffusion_model, tuner_ui.tuner_model,
tuner_ui.tuner_scale, tuner_ui.custom_tuner_model,
control_ui.control_model, control_ui.control_scale,
control_ui.crop_type, control_ui.cond_image,
diffusion_ui.negative_prompt, diffusion_ui.prompt_prefix,
diffusion_ui.sampler, diffusion_ui.discretization,
diffusion_ui.output_height, diffusion_ui.output_width,
diffusion_ui.image_number, diffusion_ui.sample_steps,
diffusion_ui.guide_scale, diffusion_ui.guide_rescale,
refiner_ui.refine_strength, refiner_ui.refine_sampler,
refiner_ui.refine_discretization, refiner_ui.refine_guide_scale,
refiner_ui.refine_guide_rescale, mantra_ui.style_template,
mantra_ui.style_negative_template, diffusion_ui.image_seed
]
self.gen_outputs = [
self.before_refine_panel, self.before_refine_gallery,
self.output_gallery
]
self.generate_button.click(self.generate_gallery,
inputs=self.gen_inputs,
outputs=self.gen_outputs,
queue=True)
self.prompt.submit(self.generate_gallery,
inputs=self.gen_inputs,
outputs=self.gen_outputs,
queue=True)
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os
import gradio as gr
@@ -45,64 +46,76 @@ class MantraUI(UIBase):
return name_level_style, all_styles
def create_ui(self, *args, **kwargs):
with gr.Row(equal_height=True):
with gr.Column(scale=1):
with gr.Group(visible=True):
with gr.Row(equal_height=True):
self.style = gr.Dropdown(
label=self.component_names.mantra_styles,
choices=self.all_styles[self.default_pipeline],
value=None,
multiselect=True,
interactive=True)
with gr.Row(equal_height=True):
with gr.Column(scale=1):
self.style_name = gr.Text(
self.state = gr.State(value=False)
with gr.Column(visible=False) as self.tab:
with gr.Row():
with gr.Column(scale=1):
with gr.Group(visible=True):
with gr.Row(equal_height=True):
self.style = gr.Dropdown(
label=self.component_names.mantra_styles,
choices=self.all_styles[self.default_pipeline],
value=None,
multiselect=True,
interactive=True)
with gr.Row(equal_height=True):
with gr.Column(scale=1):
self.style_name = gr.Text(
value='',
label=self.component_names.style_name)
with gr.Column(scale=1):
self.style_source = gr.Text(
value='',
label=self.component_names.style_source)
with gr.Column(scale=1):
self.style_desc = gr.Text(
value='',
label=self.component_names.style_desc)
with gr.Row(equal_height=True):
self.style_prompt = gr.Text(
value='',
label=self.component_names.style_name)
with gr.Column(scale=1):
self.style_source = gr.Text(
label=self.component_names.style_prompt,
lines=4)
with gr.Row(equal_height=True):
self.style_negative_prompt = gr.Text(
value='',
label=self.component_names.style_source)
with gr.Column(scale=1):
self.style_desc = gr.Text(
label=self.component_names.
style_negative_prompt,
lines=4)
with gr.Column(scale=1):
with gr.Group(visible=True):
with gr.Row(equal_height=True):
self.style_template = gr.Text(
value='',
label=self.component_names.style_desc)
with gr.Row(equal_height=True):
self.style_prompt = gr.Text(
value='',
label=self.component_names.style_prompt,
lines=4)
with gr.Row(equal_height=True):
self.style_negative_prompt = gr.Text(
value='',
label=self.component_names.style_negative_prompt,
lines=4)
with gr.Column(scale=1):
with gr.Group(visible=True):
with gr.Row(equal_height=True):
self.style_template = gr.Text(
value='',
label=self.component_names.style_template,
lines=2)
with gr.Row(equal_height=True):
self.style_negative_template = gr.Text(
value='',
label=self.component_names.style_negative_template,
lines=2)
with gr.Row(equal_height=True):
self.style_example = gr.Image(
label=self.component_names.style_example,
source='upload',
value=None,
interactive=False)
with gr.Row(equal_height=True):
self.style_example_prompt = gr.Text(
value='',
label=self.component_names.style_example_prompt,
lines=2)
label=self.component_names.style_template,
lines=2)
with gr.Row(equal_height=True):
self.style_negative_template = gr.Text(
value='',
label=self.component_names.
style_negative_template,
lines=2)
with gr.Row(equal_height=True):
self.style_example = gr.Image(
label=self.component_names.style_example,
source='upload',
value=None,
interactive=False)
with gr.Row(equal_height=True):
self.style_example_prompt = gr.Text(
value='',
label=self.component_names.
style_example_prompt,
lines=2)
self.example_block = gr.Accordion(
label=self.component_names.example_block_name, open=True)
def set_callbacks(self, model_manage_ui, **kwargs):
gallery_ui = kwargs.pop('gallery_ui')
with self.example_block:
gr.Examples(examples=self.component_names.examples,
inputs=[self.style, gallery_ui.prompt])
def set_callbacks(self, model_manage_ui):
def change_style(style, diffusion_model):
style_template = ''
style_negative_template = []
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import gradio as gr
@@ -16,6 +17,8 @@ class ModelManageUI(UIBase):
self.component_names = ModelManageUIName(language)
def create_ui(self, *args, **kwargs):
self.diffusion_state = gr.State(
value=self.default_choices['diffusion_model']['default'])
with gr.Group():
gr.Markdown(value=self.component_names.model_block_name)
with gr.Row(variant='panel', equal_height=True):
@@ -99,7 +102,8 @@ class ModelManageUI(UIBase):
# self.tuner_name = gr.Text(
# label='tuner_name')
def set_callbacks(self, diffusion_ui, tuner_ui, control_ui, mantra_ui):
def set_callbacks(self, diffusion_ui, tuner_ui, control_ui, mantra_ui,
**kwargs):
# def select_refine_tuner(all_select, evt: gr.SelectData):
# if 'Refiners' in all_select:
# refine_panel = gr.Row(visible=True)
@@ -122,12 +126,19 @@ class ModelManageUI(UIBase):
# self.refine_diffusion_panel, self.tuner_choice_panel,
# advance_ui.refine_tab, advance_ui.refine_state
# ])
def diffusion_model_change(diffusion_model, control_mode):
diffusion_model_info = self.pipe_manager.model_level_info[
diffusion_model]
now_pipeline = diffusion_model_info['pipeline'][0]
def diffusion_model_change(diffusion_state, diffusion_model,
control_mode):
if diffusion_state != diffusion_model:
last_pipline = self.pipe_manager.model_level_info[
diffusion_state]['pipeline'][0]
last_pipeline_ins = self.pipe_manager.pipeline_level_modules[
last_pipline]
last_pipeline_ins.dynamic_unload(name='all')
now_pipeline = self.pipe_manager.model_level_info[diffusion_model][
'pipeline'][0]
pipeline_ins = self.pipe_manager.pipeline_level_modules[
now_pipeline]
pipeline_ins.dynamic_load(name='all')
all_module_name = {}
for module_name in self.pipe_manager.module_list:
module = getattr(pipeline_ins, module_name)
@@ -164,9 +175,10 @@ class ModelManageUI(UIBase):
default_input)
diffusion_ui.cur_paras = cur_paras
return (
diffusion_model,
gr.Dropdown(value=all_module_name['first_stage_model']),
gr.Dropdown(value=all_module_name['cond_stage_model']),
gr.Dropdown(choices=tunner_choices, value=None),
gr.Dropdown(choices=tunner_choices, value=[]),
gr.Dropdown(choices=controller_choices,
value=controller_default),
gr.Dropdown(choices=mantra_ui.all_styles[now_pipeline],
@@ -187,14 +199,17 @@ class ModelManageUI(UIBase):
self.diffusion_model.change(
diffusion_model_change,
inputs=[self.diffusion_model, control_ui.control_mode],
outputs=[
self.first_stage_model, self.cond_stage_model,
tuner_ui.tuner_model, control_ui.control_model,
mantra_ui.style, diffusion_ui.negative_prompt,
diffusion_ui.prompt_prefix, diffusion_ui.output_height,
diffusion_ui.sampler, diffusion_ui.discretization,
diffusion_ui.sample_steps, diffusion_ui.guide_scale,
diffusion_ui.guide_rescale
inputs=[
self.diffusion_state, self.diffusion_model,
control_ui.control_mode
],
queue=False)
outputs=[
self.diffusion_state, self.first_stage_model,
self.cond_stage_model, tuner_ui.tuner_model,
control_ui.control_model, mantra_ui.style,
diffusion_ui.negative_prompt, diffusion_ui.prompt_prefix,
diffusion_ui.output_height, diffusion_ui.sampler,
diffusion_ui.discretization, diffusion_ui.sample_steps,
diffusion_ui.guide_scale, diffusion_ui.guide_rescale
],
queue=True)
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import gradio as gr
from scepter.studio.inference.inference_ui.component_names import RefinerUIName
@@ -19,8 +20,8 @@ class RefinerUI(UIBase):
return diffusion_paras
def create_ui(self, *args, **kwargs):
self.refine_state = gr.State(value=False)
with gr.Group(visible=False) as self.refine_tab:
self.state = gr.State(value=False)
with gr.Group(visible=False) as self.tab:
with gr.Row(equal_height=True):
with gr.Column(variant='panel', scale=1, min_width=0):
self.refiner_diffusion_model = gr.Dropdown(
@@ -86,5 +87,5 @@ class RefinerUI(UIBase):
'DEFAULT', 0.5),
interactive=True)
def set_callbacks(self):
def set_callbacks(self, model_manage_ui, **kwargs):
pass
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os
import gradio as gr
@@ -45,53 +46,75 @@ class TunerUI(UIBase):
one_tuner.NAME] = one_tuner
def create_ui(self, *args, **kwargs):
with gr.Row(equal_height=True):
with gr.Column(variant='panel', scale=1, min_width=0):
with gr.Group(visible=True):
with gr.Row(equal_height=True):
with gr.Column(scale=1):
self.tuner_model = gr.Dropdown(
label=self.component_names.tuner_model,
choices=self.tunner_choices,
self.state = gr.State(value=False)
with gr.Column(visible=False) as self.tab:
with gr.Row():
with gr.Column(variant='panel', scale=1, min_width=0):
with gr.Group(visible=True):
with gr.Row(equal_height=True):
with gr.Column(scale=1):
self.tuner_model = gr.Dropdown(
label=self.component_names.tuner_model,
choices=self.tunner_choices,
value=None,
multiselect=True,
interactive=True)
with gr.Column(scale=1):
self.custom_tuner_model = gr.Dropdown(
label=self.component_names.
custom_tuner_model,
choices=[],
value=None,
multiselect=True,
interactive=True)
with gr.Row(equal_height=True):
with gr.Column(scale=1):
self.tuner_type = gr.Text(
value='',
label=self.component_names.tuner_type)
with gr.Column(scale=1):
self.base_model = gr.Text(
value='',
label=self.component_names.base_model)
with gr.Column(scale=1):
self.tuner_desc = gr.Text(
value='',
label=self.component_names.tuner_desc,
lines=4)
with gr.Column(variant='panel', scale=1, min_width=0):
with gr.Group(visible=True):
with gr.Row(equal_height=True):
self.tuner_example = gr.Image(
label=self.component_names.tuner_example,
source='upload',
value=None,
multiselect=True,
interactive=True)
with gr.Column(scale=1):
self.custom_tuner_model = gr.Dropdown(
label=self.component_names.custom_tuner_model,
choices=[],
value=None,
multiselect=True,
interactive=True)
with gr.Row(equal_height=True):
with gr.Column(scale=1):
self.tuner_type = gr.Text(
interactive=False)
with gr.Row(equal_height=True):
self.tuner_prompt_example = gr.Text(
value='',
label=self.component_names.tuner_type)
with gr.Column(scale=1):
self.base_model = gr.Text(
value='',
label=self.component_names.base_model)
with gr.Column(scale=1):
self.tuner_desc = gr.Text(
value='',
label=self.component_names.tuner_desc,
lines=4)
with gr.Column(variant='panel', scale=1, min_width=0):
with gr.Group(visible=True):
with gr.Row(equal_height=True):
self.tuner_example = gr.Image(
label=self.component_names.tuner_example,
source='upload',
value=None,
interactive=False)
with gr.Row(equal_height=True):
self.tuner_prompt_example = gr.Text(
value='',
label=self.component_names.tuner_prompt_example,
lines=2)
label=self.component_names.
tuner_prompt_example,
lines=2)
with gr.Accordion(label=self.component_names.advance_block_name,
open=False):
self.tuner_scale = gr.Slider(
label=self.component_names.tuner_scale,
minimum=0.0,
maximum=1.0,
step=0.05,
value=1.0,
interactive=True)
self.example_block = gr.Accordion(
label=self.component_names.example_block_name, open=True)
def set_callbacks(self, model_manage_ui, **kwargs):
gallery_ui = kwargs.pop('gallery_ui')
with self.example_block:
gr.Examples(examples=self.component_names.examples,
inputs=[self.tuner_model, gallery_ui.prompt])
def set_callbacks(self, model_manage_ui):
def tuner_model_change(tuner_model, diffusion_model):
diffusion_model_info = self.pipe_manager.model_level_info[
diffusion_model]
+2
View File
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
# For dataset manager
class CreateDatasetUIName():
def __init__(self, language='en'):
@@ -44,6 +45,7 @@ class CreateDatasetUIName():
self.refresh_data_list_info1 = (
'The dataset name has been changed, '
'please refresh the list and try again.')
self.use_link = 'Use File Link'
elif language == 'zh':
self.dataset_name = '数据集'
self.btn_create_datasets = '新建'
@@ -76,6 +78,7 @@ class CreateDatasetUIName():
self.illegal_data_err3 = '文件解压失败,上传存储器失败!'
self.modify_data_name_err1 = '变更数据集名称失败!'
self.refresh_data_list_info1 = '该数据集名称发生了变更,请刷新列表试一下。'
self.use_link = '使用文件链接'
class DatasetGalleryUIName():
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from __future__ import annotations
import copy
@@ -219,8 +220,7 @@ class CreateDatasetUI(UIBase):
new_file_folder = f'{local_dataset_folder}/images'
os.makedirs(new_file_folder, exist_ok=True)
if file_folder is not None:
res = os.popen(f"mv '{file_folder}'/* '{new_file_folder}'")
res = res.readlines()
_ = FS.get_dir_to_local_dir(file_folder, new_file_folder)
elif len(raw_list) > 0:
raw_list = list(set(raw_list))
for img_id, cur_image in enumerate(raw_list):
@@ -251,8 +251,13 @@ class CreateDatasetUI(UIBase):
res = res.readlines()
if not os.path.exists(new_train_list):
raise gr.Error(f'{str(res)}')
res = os.popen(f"rm -rf '{hit_dir}'")
res = res.readlines()
try:
res = os.popen(f"rm -rf '{hit_dir}/images/*'")
_ = res.readlines()
res = os.popen(f"rm -rf '{hit_dir}'")
_ = res.readlines()
except Exception:
pass
file_list = self.load_train_csv(new_train_list, data_folder)
return file_list
@@ -320,7 +325,7 @@ class CreateDatasetUI(UIBase):
if FS.exists(meta_file):
local_dataset_folder, _ = FS.map_to_local(one_dir)
local_dataset_folder = FS.get_dir_to_local_dir(
one_dir, local_dataset_folder)
one_dir, local_dataset_folder, multi_thread=True)
meta_data = self.load_meta(
os.path.join(local_dataset_folder, 'meta.json'))
meta_data['local_work_dir'] = local_dataset_folder
@@ -361,10 +366,16 @@ class CreateDatasetUI(UIBase):
self.create_mode = gr.State(value=0)
with gr.Column(visible=False,
min_width=0) as file_panel:
self.use_link = gr.Checkbox(
label=self.components_name.use_link,
value=False,
visible=False)
self.file_path = gr.File(
label=self.components_name.zip_file,
min_width=0,
file_types=['.zip', '.txt', '.csv'])
file_types=['.zip', '.txt', '.csv'],
visible=False)
self.file_path_url = gr.Text(
label=self.components_name.zip_file_url,
value='',
@@ -397,8 +408,11 @@ class CreateDatasetUI(UIBase):
return (gr.Column(visible=True), gr.Column(visible=True),
gr.Column(visible=True),
gr.Checkbox(value=False, visible=False),
gr.Text(value=get_random_dataset_name(), interactive=True),
gr.File(value=None), gr.Text(value='', visible=False), 2)
gr.Text(value=get_random_dataset_name(),
interactive=True), gr.File(value=None,
visible=True),
gr.Text(value='', visible=False), 2,
gr.Checkbox(value=False, visible=True))
def get_random_dataset_name():
data_name = 'name-version-{0:%Y%m%d_%H_%M_%S}'.format(
@@ -427,12 +441,13 @@ class CreateDatasetUI(UIBase):
if not file_url.strip() == '' and file_path is not None:
raise gr.Error(self.components_name.illegal_data_name_err4)
if create_mode == 1 and not file_url.strip() == '':
if create_mode == 3 and not file_url.strip() == '':
file_name, surfix = os.path.splitext(file_url.split('?')[0])
save_file = os.path.join(self.work_dir, f'{user_name}{surfix}')
with FS.put_to(save_file) as local_path:
res = os.popen(f"wget '{file_url}' -O '{local_path}'")
res.readlines()
local_path, _ = FS.map_to_local(save_file)
res = os.popen(f"wget -c '{file_url}' -O '{local_path}'")
res.readlines()
FS.put_object_from_local_file(local_path, save_file)
if not FS.exists(save_file):
raise gr.Error(
f'{self.components_name.illegal_data_err1} {str(res)}')
@@ -468,7 +483,8 @@ class CreateDatasetUI(UIBase):
raise gr.Error(
f'{self.components_name.illegal_data_err2} {surfix}')
is_flag = FS.put_dir_from_local_dir(local_dataset_folder,
dataset_folder)
dataset_folder,
multi_thread=True)
if not is_flag:
raise gr.Error(f'{self.components_name.illegal_data_err3}')
@@ -480,7 +496,8 @@ class CreateDatasetUI(UIBase):
meta['work_dir'] = dataset_folder
self.meta_dict[meta['dataset_name']] = meta
self.dataset_list.append(meta['dataset_name'])
if meta['dataset_name'] not in self.dataset_list:
self.dataset_list.append(meta['dataset_name'])
return (
gr.Checkbox(value=True, visible=False),
gr.Dropdown(value=user_name, choices=self.dataset_list),
@@ -499,10 +516,23 @@ class CreateDatasetUI(UIBase):
self.btn_create_datasets_from_file.click(show_file_panel, [], [
self.file_panel, self.dataset_panel, self.btn_panel,
self.panel_state, self.user_data_name, self.file_path,
self.file_path_url, self.create_mode
self.file_path_url, self.create_mode, self.use_link
],
queue=False)
def use_link_change(use_link):
if use_link:
create_mode = 3
return (gr.File(value=None, visible=False),
gr.Text(value='', visible=True), create_mode)
else:
create_mode = 2
return (gr.File(value=None, visible=True),
gr.Text(value='', visible=False), create_mode)
self.use_link.change(
use_link_change, [self.use_link],
[self.file_path, self.file_path_url, self.create_mode])
# Click Refresh
self.random_data_button.click(get_random_dataset_name, [],
[self.user_data_name],
@@ -517,7 +547,7 @@ class CreateDatasetUI(UIBase):
self.user_data_name, self.create_mode, self.file_path_url,
self.file_path, self.panel_state
], [self.panel_state, self.dataset_name],
queue=False)
queue=True)
def show_edit_panel(panel_state, data_name):
if panel_state:
@@ -555,14 +585,17 @@ class CreateDatasetUI(UIBase):
local_dataset_folder, _ = FS.map_to_local(dataset_folder)
os.makedirs(local_dataset_folder, exist_ok=True)
is_flag = FS.get_dir_to_local_dir(ori_meta['work_dir'],
local_dataset_folder)
local_dataset_folder,
multi_thread=True)
file_list = ori_meta['file_list']
is_flag = FS.put_dir_from_local_dir(
local_dataset_folder, dataset_folder)
is_flag = FS.put_dir_from_local_dir(local_dataset_folder,
dataset_folder,
multi_thread=True)
if not is_flag:
raise gr.Error(self.components_name.illegal_data_err3)
is_flag = FS.put_dir_from_local_dir(
local_dataset_folder, dataset_folder)
is_flag = FS.put_dir_from_local_dir(local_dataset_folder,
dataset_folder,
multi_thread=True)
if not is_flag:
raise gr.Error(self.components_name.illegal_data_err3)
cursor = ori_meta['cursor']
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from __future__ import annotations
import os.path
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from __future__ import annotations
import os
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os.path
import gradio as gr
+2
View File
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import argparse
import os
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os
import gradio as gr
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
# For dataset manager
class InferenceUIName():
def __init__(self, language='en'):
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os
import gradio as gr
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import datetime
import os
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os
from glob import glob
@@ -367,7 +368,7 @@ def get_control_default(config_dict):
default_control_cfg = default_control_cfg['control_type']
else:
return ret_data
# import pdb; pdb.set_trace()
ret_data['control_choices'] = list(default_control_cfg['choices'].keys())
defalt_t_type = default_control_cfg['default']
type_paras = default_control_cfg.get(defalt_t_type, None)
+2
View File
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.utils.file_system import FS
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
class Singleton(object):
@classmethod
def get_instance(cls, cfg, **kwargs):
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.studio.utils.singleton import Singleton
+5 -9
View File
@@ -3,7 +3,6 @@
import argparse
import os
import cv2
import numpy as np
import torch
import torch.cuda.amp as amp
@@ -67,13 +66,12 @@ def run_task(cfg):
for idx, out in enumerate(ret):
img = out['image']
img = img.permute(1, 2, 0).cpu().numpy()
img = (img * 255).astype(np.uint8)
img = Image.fromarray((img * 255).astype(np.uint8))
filename = '{}_{}.png'.format('inference', idx)
save_file = os.path.join(save_folder, filename)
with FS.put_to(save_file) as local_path:
image = img.copy()
cv2.cvtColor(image, cv2.COLOR_RGB2BGR, image)
cv2.imwrite(local_path, image)
img.save(local_path)
std_logger.info(f'Processed {filename} save to {local_path}')
def run_task_control(cfg):
@@ -151,13 +149,11 @@ def run_task_control(cfg):
for name in ['image', 'hint']:
img = out[name]
img = img.permute(1, 2, 0).cpu().numpy()
img = (img * 255).astype(np.uint8)
img = Image.fromarray((img * 255).astype(np.uint8))
filename = '{}_{}_{}.png'.format('inference', name, idx)
save_file = os.path.join(save_folder, filename)
with FS.put_to(save_file) as local_path:
image = img.copy()
cv2.cvtColor(image, cv2.COLOR_RGB2BGR, image)
cv2.imwrite(local_path, image)
img.save(local_path)
std_logger.info(f'Processed {filename} save to {local_path}')
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import argparse
import datetime
import os
+1 -1
View File
@@ -1,7 +1,7 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
__version__ = '0.0.2'
__version__ = '0.0.3'
version_info = tuple(int(x) for x in __version__.split('.')[0:3])
+43 -1
View File
@@ -50,6 +50,28 @@ class AnnotatorTest(unittest.TestCase):
Image.fromarray(canny_image).save(
os.path.join(self.save_dir, 'sunflower_canny.png'))
@unittest.skip('')
def test_annotator_canny_random(self):
# canny
canny_dict = {
'NAME': 'CannyAnnotator',
'LOW_THRESHOLD': 100,
'HIGH_THRESHOLD': 200,
'RANDOM_CFG': {
'PROBA': 1.0,
'MIN_LOW_THRESHOLD': 50,
'MAX_LOW_THRESHOLD': 100,
'MIN_HIGH_THRESHOLD': 200,
'MAX_HIGH_THRESHOLD': 350
}
}
canny_anno = Config(cfg_dict=canny_dict, load=False)
canny_ins = ANNOTATORS.build(canny_anno).to(we.device_id)
canny_image = canny_ins(self.image)
print("canny's shape:", canny_image.shape)
Image.fromarray(canny_image).save(
os.path.join(self.save_dir, 'sunflower_canny_random.png'))
@unittest.skip('')
def test_annotator_hed(self):
# hed
@@ -129,6 +151,26 @@ class AnnotatorTest(unittest.TestCase):
Image.fromarray(color_image).save(
os.path.join(self.save_dir, 'sunflower_color.png'))
@unittest.skip('')
def test_annotator_color_random(self):
# color
color_dict = {
'NAME': 'ColorAnnotator',
'RATIO': 64,
'RANDOM_CFG': {
'PROBA': 1.0,
# 'MIN_RATIO': 64,
# 'MAX_RATIO': 128
'CHOICE_RATIO': [32, 64, 128]
}
}
color_anno = Config(cfg_dict=color_dict, load=False)
color_ins = ANNOTATORS.build(color_anno).to(we.device_id)
color_image = color_ins(self.image)
print("color's shape:", color_image.shape)
Image.fromarray(color_image).save(
os.path.join(self.save_dir, 'sunflower_color_random.png'))
@unittest.skip('')
def test_annotator_multi(self):
# multi annotators
@@ -185,7 +227,7 @@ class AnnotatorTest(unittest.TestCase):
Image.fromarray(save_image).save(
os.path.join(self.save_dir, f'sunflower_multi_{key}.png'))
# @unittest.skip('')
@unittest.skip('')
def test_annotator_processor(self):
from scepter.modules.annotator.utils import AnnotatorProcessor
anno_processor = AnnotatorProcessor(anno_type='hed')
+25 -7
View File
@@ -23,6 +23,14 @@ class TrainTest(unittest.TestCase):
os.path.exists(
os.path.join(self.tmp_dir, 'sd15_512_full/checkpoints')))
os.system(
'python scepter/tools/run_train.py '
'--cfg scepter/methods/examples/generation/stable_diffusion_2.1_512.yaml '
'--max_steps 100')
self.assertTrue(
os.path.exists(
os.path.join(self.tmp_dir, 'sd21_512_full/checkpoints')))
os.system(
'python scepter/tools/run_train.py '
'--cfg scepter/methods/examples/generation/stable_diffusion_2.1_768.yaml '
@@ -49,6 +57,14 @@ class TrainTest(unittest.TestCase):
os.path.exists(
os.path.join(self.tmp_dir, 'sd15_512_lora/checkpoints')))
os.system(
'python scepter/tools/run_train.py '
'--cfg scepter/methods/examples/generation/stable_diffusion_2.1_512_lora.yaml '
'--max_steps 100')
self.assertTrue(
os.path.exists(
os.path.join(self.tmp_dir, 'sd21_512_lora/checkpoints')))
os.system(
'python scepter/tools/run_train.py '
'--cfg scepter/methods/examples/generation/stable_diffusion_2.1_768_lora.yaml '
@@ -163,14 +179,16 @@ class TrainTest(unittest.TestCase):
os.path.join(self.tmp_dir,
'sdxl_1024_sce_ctr_color/checkpoints')))
# @unittest.skip('')
@unittest.skip('')
def test_generation_example_datatxt(self):
# os.system(
# 'python scepter/tools/run_train.py '
# '--cfg scepter/methods/scedit/t2i/sdxl_1024_sce_t2i_datatxt.yaml '
# '--max_steps 100'
# )
# self.assertTrue(os.path.exists(os.path.join(self.tmp_dir, 'sdxl_1024_sce_t2i_datatxt/checkpoints')))
os.system(
'python scepter/tools/run_train.py '
'--cfg scepter/methods/scedit/t2i/sdxl_1024_sce_t2i_datatxt.yaml '
'--max_steps 100')
self.assertTrue(
os.path.exists(
os.path.join(self.tmp_dir,
'sdxl_1024_sce_t2i_datatxt/checkpoints')))
os.system(
'python scepter/tools/run_train.py '
+85
View File
@@ -2,6 +2,7 @@
# Copyright (c) Alibaba, Inc. and its affiliates.
import os
import time
import unittest
from scepter.modules.utils.config import Config
@@ -15,6 +16,7 @@ class FSTest(unittest.TestCase):
def tearDown(self):
super().tearDown()
@unittest.skip('')
def test_modelscope(self):
fs_info = {'NAME': 'ModelscopeFs', 'TEMP_DIR': 'cache/data'}
config = Config(load=False, cfg_dict=fs_info)
@@ -50,6 +52,89 @@ class FSTest(unittest.TestCase):
print(f'Download from {path} to {local_path}')
self.assertTrue(os.path.exists(local_path))
@unittest.skip('')
def test_huggingface(self):
fs_info = {'NAME': 'HuggingfaceFs', 'TEMP_DIR': 'cache/data'}
config = Config(load=False, cfg_dict=fs_info)
FS.init_fs_client(config)
path = 'hf://runwayml/stable-diffusion-v1-5'
with FS.get_dir_to_local_dir(path, wait_finish=True) as local_path:
print(f'Download from {path} to {local_path}')
self.assertTrue(os.path.exists(local_path))
path = 'hf://stabilityai/stable-diffusion-2-1-base@text_encoder'
with FS.get_dir_to_local_dir(path, wait_finish=True) as local_path:
print(f'Download from {path} to {local_path}')
self.assertTrue(os.path.exists(local_path))
path = 'hf://stabilityai/stable-diffusion-xl-base-1.0@README.md'
with FS.get_from(path, wait_finish=True) as local_path:
print(f'Download from {path} to {local_path}')
self.assertTrue(os.path.exists(local_path))
# @unittest.skip('')
def test_scedit(self):
fs_info = {'NAME': 'ModelscopeFs', 'TEMP_DIR': 'cache/cache_data'}
config = Config(load=False, cfg_dict=fs_info)
FS.init_fs_client(config)
st_time = time.time()
path = 'ms://damo/scepter_scedit'
with FS.get_dir_to_local_dir(path, wait_finish=True) as local_path:
print(
f'Download from {path} to {local_path}, take {time.time()-st_time}'
)
self.assertTrue(os.path.exists(local_path))
st_time = time.time()
path = 'ms://damo/scepter_scedit'
with FS.get_dir_to_local_dir(path, wait_finish=True) as local_path:
print(
f'Download from {path} to {local_path}, take {time.time()-st_time}'
)
self.assertTrue(os.path.exists(local_path))
st_time = time.time()
path = 'ms://damo/scepter_scedit@controllable_model/SD2.1/canny_control/'
with FS.get_dir_to_local_dir(path, wait_finish=True) as local_path:
print(
f'Download from {path} to {local_path}, take {time.time()-st_time}'
)
self.assertTrue(os.path.exists(local_path))
st_time = time.time()
path = 'ms://damo/scepter_scedit@controllable_model/SD2.1/canny_control/'
with FS.get_dir_to_local_dir(path, wait_finish=True) as local_path:
print(
f'Download from {path} to {local_path}, take {time.time()-st_time}'
)
self.assertTrue(os.path.exists(local_path))
st_time = time.time()
path = 'ms://damo/scepter@mantra_images/SD_XL1.0/894f40ed44b37c3372e6a22b8ae577a4.png'
with FS.get_from(path, wait_finish=True) as local_path:
print(
f'Download from {path} to {local_path}, take {time.time()-st_time}'
)
self.assertTrue(os.path.exists(local_path))
st_time = time.time()
path = 'ms://damo/scepter@mantra_images/SD2.1/894f40ed44b37c3372e6a22b8ae577a4.png'
with FS.get_from(path, wait_finish=True) as local_path:
print(
f'Download from {path} to {local_path}, take {time.time()-st_time}'
)
self.assertTrue(os.path.exists(local_path))
st_time = time.time()
path = 'ms://damo/scepter@mantra_images/SD2.1/894f40ed44b37c3372e6a22b8ae577a4.png'
with FS.get_from(path, wait_finish=True) as local_path:
print(
f'Download from {path} to {local_path}, take {time.time()-st_time}'
)
self.assertTrue(os.path.exists(local_path))
if __name__ == '__main__':
unittest.main()