@@ -0,0 +1,2 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
@@ -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/>
|
||||
|
||||
|
||||
@@ -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) [](https://arxiv.org/abs/2310.19859) [](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.
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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: 夸张漫画
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,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.
|
||||
|
||||
@@ -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)
|
||||
@@ -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,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,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,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,4 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
|
||||
@@ -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,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
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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!'
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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 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.
|
||||
|
||||
|
||||
class DescUIName():
|
||||
|
||||
@@ -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,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
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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,4 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import os.path
|
||||
|
||||
import gradio as gr
|
||||
|
||||
@@ -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,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)
|
||||
|
||||
@@ -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.utils.file_system import FS
|
||||
|
||||
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import os
|
||||
|
||||
|
||||
|
||||
@@ -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,4 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.studio.utils.singleton import Singleton
|
||||
|
||||
|
||||
|
||||
@@ -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,4 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import argparse
|
||||
import datetime
|
||||
import os
|
||||
|
||||
+1
-1
@@ -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])
|
||||
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -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 '
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user