diff --git a/__init__.py b/__init__.py index e69de29..cc26a06 100644 --- a/__init__.py +++ b/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/docs/zh_cn/scepter/utils/utils.md b/docs/zh_cn/scepter/utils/utils.md index be8917a..3140b43 100644 --- a/docs/zh_cn/scepter/utils/utils.md +++ b/docs/zh_cn/scepter/utils/utils.md @@ -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}")) ```
diff --git a/readme.md b/readme.md index 7ed2242..a21623e 100644 --- a/readme.md +++ b/readme.md @@ -14,6 +14,7 @@ - [Installation](#-Installation) - [Getting Started](#-getting-started) - [SCEPTER Studio](#-scepter-studio) +- [Gallery](#-gallery) - [Features](#-features) - [Learn More](#-learn-more) - [License](#license) @@ -43,6 +44,7 @@ Currently supported approaches (and counting): 3. Res-Tuning(TODO): [Res-Tuning: A Flexible and Efficient Tuning Paradigm via Unbinding Tuner from Backbone](https://arxiv.org/abs/2310.19859) [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=ResTuning&color=red&logo=arxiv)](https://arxiv.org/abs/2310.19859) [![Page link](https://img.shields.io/badge/Page-ResTuning-Gree)](https://res-tuning.github.io/) ## 🎉 News +- [2024.02]: We release new SCEdit controllable image synthesis models for SD v2.1 and SD XL. Multiple strategies applied to accelerate inference time for SCEPTER Studio. - [2024.01]: We release **SCEPTER Studio**, an integrated toolkit for data management, model training and inference based on [Gradio](https://www.gradio.app/). - [2024.01]: [SCEdit](https://arxiv.org/abs/2312.11392) support controllable image synthesis for training and inference. - [2023.12]: We propose [SCEdit](https://arxiv.org/abs/2312.11392), an efficient and controllable generation framework. @@ -56,13 +58,17 @@ Currently supported approaches (and counting): conda env create -f environment.yaml conda activate scepter ``` +- We recommend installing the specific version of PyTorch and accelerate toolbox [xFormers](https://pypi.org/project/xformers/). You can install these recommended version by pip: + +```shell +pip install -r requirements/recommended.txt +``` - Install SCEPTER by the `pip` command: ```shell pip install scepter ``` -- PS: We recommend installing PyTorch follwing [official documentation](https://pytorch.org/get-started/locally/) ## 🚀 Getting Started @@ -159,7 +165,6 @@ python scepter/tools/run_inference.py --cfg scepter/methods/scedit/ctr/sd21_768_ python scepter/tools/run_inference.py --cfg scepter/methods/scedit/ctr/sd21_768_sce_ctr_pose.yaml --num_samples 1 --prompt 'super mario' --save_folder 'test_mario_pose' --image_size 768 --task control --image 'asset/images/pose_source.png' --control_mode source --pretrained_model ms://damo/scepter_scedit@controllable_model/SD2.1/pose_control/0_SwiftSCETuning/pytorch_model.bin # pose ``` - ## 🖥️ SCEPTER Studio ### Launch @@ -167,24 +172,49 @@ python scepter/tools/run_inference.py --cfg scepter/methods/scedit/ctr/sd21_768_ To fully experience **SCEPTER Studio**, you can launch the following command line: ```shell -pip install scepter -python -m scepter.tools.webui -``` -or run after clone repo code -```shell -git clone https://github.com/modelscope/scepter.git PYTHONPATH=. python scepter/tools/webui.py --cfg scepter/methods/studio/scepter_ui.yaml ``` -The startup of **SCEPTER Studio** eliminates the need for manual downloading and organizing of models; it will automatically load the corresponding models and store them in a local directory. -Depending on the network and hardware situation, the initial startup usually requires 15-60 minutes, primarily involving the download and processing of SDv1.5, SDv2.1, and SDXL models. -Therefore, subsequent startups will become much faster (about one minute) as downloading is no longer required. - - ### Modelscope Studio We deploy a work studio on Modelscope that includes only the inference tab, please refer to [ms_scepter_studio](https://www.modelscope.cn/studios/damo/scepter_studio/summary) +## 🖼️ Gallery + +### Dragon Year Special: Dragon Tuner + + + + + + + + + + + + + + +
Gold Dragon TunerSloppy Dragon TunerRed Dragon Tuner
+ Papercraft Mantra
Azure Dragon Tuner
+ Pose Control
+ +### Text Effect Image + + + + + + + + + + + + + + +
Conditional ImageMidas Control
"Race track, top view"
Midas Control
+ Watercolor Mantra
"white lilies"
Midas Control
+ Dragon Tuner
"Spring Festival, Chinese dragon"
+ ## ✨ 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. diff --git a/requirements/framework.txt b/requirements/framework.txt index 3e7a72e..3e12b56 100644 --- a/requirements/framework.txt +++ b/requirements/framework.txt @@ -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 diff --git a/requirements/recommended.txt b/requirements/recommended.txt new file mode 100644 index 0000000..d25969c --- /dev/null +++ b/requirements/recommended.txt @@ -0,0 +1,3 @@ +torch==2.0.1 +torchvision==0.15.2 +xformers==0.0.21 diff --git a/scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml b/scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml index 7918211..207114d 100644 --- a/scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml +++ b/scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml @@ -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 diff --git a/scepter/methods/examples/generation/stable_diffusion_1.5_512_lora.yaml b/scepter/methods/examples/generation/stable_diffusion_1.5_512_lora.yaml index 4df6747..8e8af4d 100644 --- a/scepter/methods/examples/generation/stable_diffusion_1.5_512_lora.yaml +++ b/scepter/methods/examples/generation/stable_diffusion_1.5_512_lora.yaml @@ -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 diff --git a/scepter/methods/examples/generation/stable_diffusion_2.1_512.yaml b/scepter/methods/examples/generation/stable_diffusion_2.1_512.yaml new file mode 100644 index 0000000..f148649 --- /dev/null +++ b/scepter/methods/examples/generation/stable_diffusion_2.1_512.yaml @@ -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 diff --git a/scepter/methods/examples/generation/stable_diffusion_2.1_512_lora.yaml b/scepter/methods/examples/generation/stable_diffusion_2.1_512_lora.yaml new file mode 100644 index 0000000..2ea41b4 --- /dev/null +++ b/scepter/methods/examples/generation/stable_diffusion_2.1_512_lora.yaml @@ -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 diff --git a/scepter/methods/examples/generation/stable_diffusion_2.1_768.yaml b/scepter/methods/examples/generation/stable_diffusion_2.1_768.yaml index 0d9a1cf..08b3707 100644 --- a/scepter/methods/examples/generation/stable_diffusion_2.1_768.yaml +++ b/scepter/methods/examples/generation/stable_diffusion_2.1_768.yaml @@ -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 diff --git a/scepter/methods/examples/generation/stable_diffusion_2.1_768_lora.yaml b/scepter/methods/examples/generation/stable_diffusion_2.1_768_lora.yaml index 90becb2..e3434c7 100644 --- a/scepter/methods/examples/generation/stable_diffusion_2.1_768_lora.yaml +++ b/scepter/methods/examples/generation/stable_diffusion_2.1_768_lora.yaml @@ -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 diff --git a/scepter/methods/scedit/ctr/sd21_768_sce_ctr_canny.yaml b/scepter/methods/scedit/ctr/sd21_768_sce_ctr_canny.yaml index a40cb3e..af07087 100644 --- a/scepter/methods/scedit/ctr/sd21_768_sce_ctr_canny.yaml +++ b/scepter/methods/scedit/ctr/sd21_768_sce_ctr_canny.yaml @@ -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 diff --git a/scepter/methods/scedit/ctr/sd21_768_sce_ctr_pose.yaml b/scepter/methods/scedit/ctr/sd21_768_sce_ctr_pose.yaml index 932e22f..bce9527 100644 --- a/scepter/methods/scedit/ctr/sd21_768_sce_ctr_pose.yaml +++ b/scepter/methods/scedit/ctr/sd21_768_sce_ctr_pose.yaml @@ -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 diff --git a/scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_canny.yaml b/scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_canny.yaml new file mode 100644 index 0000000..d853348 --- /dev/null +++ b/scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_canny.yaml @@ -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 diff --git a/scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_color.yaml b/scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_color.yaml index cdf7fa3..7a95e2d 100644 --- a/scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_color.yaml +++ b/scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_color.yaml @@ -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 diff --git a/scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_color_datatxt.yaml b/scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_color_datatxt.yaml index 4a15c00..40dddec 100644 --- a/scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_color_datatxt.yaml +++ b/scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_color_datatxt.yaml @@ -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 diff --git a/scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_depth.yaml b/scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_depth.yaml index e4b33ab..af556a6 100644 --- a/scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_depth.yaml +++ b/scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_depth.yaml @@ -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 diff --git a/scepter/methods/studio/extensions/controllers/official_controllers.yaml b/scepter/methods/studio/extensions/controllers/official_controllers.yaml index 63a8c8d..836bfed 100644 --- a/scepter/methods/studio/extensions/controllers/official_controllers.yaml +++ b/scepter/methods/studio/extensions/controllers/official_controllers.yaml @@ -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 diff --git a/scepter/methods/studio/extensions/tuners/official_tuners.yaml b/scepter/methods/studio/extensions/tuners/official_tuners.yaml index 99ee032..d195271 100644 --- a/scepter/methods/studio/extensions/tuners/official_tuners.yaml +++ b/scepter/methods/studio/extensions/tuners/official_tuners.yaml @@ -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: 夸张漫画 diff --git a/scepter/methods/studio/home/home.yaml b/scepter/methods/studio/home/home.yaml index e760d92..61ebabc 100644 --- a/scepter/methods/studio/home/home.yaml +++ b/scepter/methods/studio/home/home.yaml @@ -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; } @@ -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; } @@ -92,4 +92,4 @@ GUIDE_INFO:
Train & Inference Video
- \ No newline at end of file + diff --git a/scepter/methods/studio/inference/inference.yaml b/scepter/methods/studio/inference/inference.yaml index 33a78b3..542b616 100644 --- a/scepter/methods/studio/inference/inference.yaml +++ b/scepter/methods/studio/inference/inference.yaml @@ -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 diff --git a/scepter/modules/annotator/__init__.py b/scepter/modules/annotator/__init__.py index 46a3367..f14a40f 100644 --- a/scepter/modules/annotator/__init__.py +++ b/scepter/modules/annotator/__init__.py @@ -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 diff --git a/scepter/modules/annotator/canny.py b/scepter/modules/annotator/canny.py index 6048970..c0f7a18 100644 --- a/scepter/modules/annotator/canny.py +++ b/scepter/modules/annotator/canny.py @@ -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 diff --git a/scepter/modules/annotator/color.py b/scepter/modules/annotator/color.py index 02b406d..40690e3 100644 --- a/scepter/modules/annotator/color.py +++ b/scepter/modules/annotator/color.py @@ -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 diff --git a/scepter/modules/annotator/hed.py b/scepter/modules/annotator/hed.py index fbfd84c..b06f4f9 100644 --- a/scepter/modules/annotator/hed.py +++ b/scepter/modules/annotator/hed.py @@ -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. diff --git a/scepter/modules/annotator/identity.py b/scepter/modules/annotator/identity.py new file mode 100644 index 0000000..f142a67 --- /dev/null +++ b/scepter/modules/annotator/identity.py @@ -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) diff --git a/scepter/modules/annotator/invert.py b/scepter/modules/annotator/invert.py new file mode 100644 index 0000000..afdf571 --- /dev/null +++ b/scepter/modules/annotator/invert.py @@ -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) diff --git a/scepter/modules/annotator/midas_op.py b/scepter/modules/annotator/midas_op.py index 09a883d..af9e47a 100644 --- a/scepter/modules/annotator/midas_op.py +++ b/scepter/modules/annotator/midas_op.py @@ -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 diff --git a/scepter/modules/annotator/mlsd_op.py b/scepter/modules/annotator/mlsd_op.py index fbab416..ec5f094 100644 --- a/scepter/modules/annotator/mlsd_op.py +++ b/scepter/modules/annotator/mlsd_op.py @@ -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 diff --git a/scepter/modules/annotator/openpose.py b/scepter/modules/annotator/openpose.py index 0c3b608..69320ec 100644 --- a/scepter/modules/annotator/openpose.py +++ b/scepter/modules/annotator/openpose.py @@ -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 diff --git a/scepter/modules/annotator/utils.py b/scepter/modules/annotator/utils.py index 672e921..89a25bd 100644 --- a/scepter/modules/annotator/utils.py +++ b/scepter/modules/annotator/utils.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. import cv2 import numpy as np diff --git a/scepter/modules/data/dataset/base_dataset.py b/scepter/modules/data/dataset/base_dataset.py index 0211d6d..fa9dc8d 100644 --- a/scepter/modules/data/dataset/base_dataset.py +++ b/scepter/modules/data/dataset/base_dataset.py @@ -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 diff --git a/scepter/modules/inference/__init__.py b/scepter/modules/inference/__init__.py index 6c8b984..d442d9c 100644 --- a/scepter/modules/inference/__init__.py +++ b/scepter/modules/inference/__init__.py @@ -1,2 +1,3 @@ # -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. from scepter.modules.inference.diffusion_inference import DiffusionInference diff --git a/scepter/modules/inference/control_inference.py b/scepter/modules/inference/control_inference.py new file mode 100644 index 0000000..47836ea --- /dev/null +++ b/scepter/modules/inference/control_inference.py @@ -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 diff --git a/scepter/modules/inference/diffusion_inference.py b/scepter/modules/inference/diffusion_inference.py index 0bd925c..901ff95 100644 --- a/scepter/modules/inference/diffusion_inference.py +++ b/scepter/modules/inference/diffusion_inference.py @@ -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 diff --git a/scepter/modules/inference/tuner_inference.py b/scepter/modules/inference/tuner_inference.py new file mode 100644 index 0000000..ea7797c --- /dev/null +++ b/scepter/modules/inference/tuner_inference.py @@ -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 diff --git a/scepter/modules/model/backbone/unet/unet_module.py b/scepter/modules/model/backbone/unet/unet_module.py index 9fcb98a..6c71fcc 100644 --- a/scepter/modules/model/backbone/unet/unet_module.py +++ b/scepter/modules/model/backbone/unet/unet_module.py @@ -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): diff --git a/scepter/modules/model/backbone/unet/unet_utils.py b/scepter/modules/model/backbone/unet/unet_utils.py index be0f5c1..69d270b 100644 --- a/scepter/modules/model/backbone/unet/unet_utils.py +++ b/scepter/modules/model/backbone/unet/unet_utils.py @@ -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) diff --git a/scepter/modules/model/embedder/embedder.py b/scepter/modules/model/embedder/embedder.py index 5cc1d38..1e1addf 100644 --- a/scepter/modules/model/embedder/embedder.py +++ b/scepter/modules/model/embedder/embedder.py @@ -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): diff --git a/scepter/modules/model/network/diffusion/diffusion.py b/scepter/modules/model/network/diffusion/diffusion.py index 5a23c56..4f8080a 100644 --- a/scepter/modules/model/network/diffusion/diffusion.py +++ b/scepter/modules/model/network/diffusion/diffusion.py @@ -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) diff --git a/scepter/modules/transform/__init__.py b/scepter/modules/transform/__init__.py index 9b2e1b4..eeb08c0 100644 --- a/scepter/modules/transform/__init__.py +++ b/scepter/modules/transform/__init__.py @@ -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, diff --git a/scepter/modules/transform/tensor.py b/scepter/modules/transform/tensor.py index f2df463..51568b5 100644 --- a/scepter/modules/transform/tensor.py +++ b/scepter/modules/transform/tensor.py @@ -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) diff --git a/scepter/modules/utils/distribute.py b/scepter/modules/utils/distribute.py index 439eda5..3fc81dc 100644 --- a/scepter/modules/utils/distribute.py +++ b/scepter/modules/utils/distribute.py @@ -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!' diff --git a/scepter/modules/utils/export_model.py b/scepter/modules/utils/export_model.py index 679e3ef..e499ec7 100644 --- a/scepter/modules/utils/export_model.py +++ b/scepter/modules/utils/export_model.py @@ -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 = { diff --git a/scepter/modules/utils/file_clients/__init__.py b/scepter/modules/utils/file_clients/__init__.py index 84656dc..52fcbd5 100644 --- a/scepter/modules/utils/file_clients/__init__.py +++ b/scepter/modules/utils/file_clients/__init__.py @@ -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 diff --git a/scepter/modules/utils/file_clients/aliyun_oss_fs.py b/scepter/modules/utils/file_clients/aliyun_oss_fs.py index 27de807..622710c 100644 --- a/scepter/modules/utils/file_clients/aliyun_oss_fs.py +++ b/scepter/modules/utils/file_clients/aliyun_oss_fs.py @@ -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]: diff --git a/scepter/modules/utils/file_clients/http_fs.py b/scepter/modules/utils/file_clients/http_fs.py index 3418334..7dcab7f 100644 --- a/scepter/modules/utils/file_clients/http_fs.py +++ b/scepter/modules/utils/file_clients/http_fs.py @@ -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 diff --git a/scepter/modules/utils/file_clients/huggingface_fs.py b/scepter/modules/utils/file_clients/huggingface_fs.py new file mode 100644 index 0000000..150e1b1 --- /dev/null +++ b/scepter/modules/utils/file_clients/huggingface_fs.py @@ -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 diff --git a/scepter/modules/utils/file_clients/local_fs.py b/scepter/modules/utils/file_clients/local_fs.py index 7cd3d05..eae15ab 100644 --- a/scepter/modules/utils/file_clients/local_fs.py +++ b/scepter/modules/utils/file_clients/local_fs.py @@ -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: diff --git a/scepter/modules/utils/file_clients/modelscope_fs.py b/scepter/modules/utils/file_clients/modelscope_fs.py index 7d20b98..33a9f51 100644 --- a/scepter/modules/utils/file_clients/modelscope_fs.py +++ b/scepter/modules/utils/file_clients/modelscope_fs.py @@ -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]: diff --git a/scepter/modules/utils/file_system.py b/scepter/modules/utils/file_system.py index b4b3c10..11cb420 100644 --- a/scepter/modules/utils/file_system.py +++ b/scepter/modules/utils/file_system.py @@ -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. diff --git a/scepter/studio/__init__.py b/scepter/studio/__init__.py index e69de29..cc26a06 100644 --- a/scepter/studio/__init__.py +++ b/scepter/studio/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/studio/home/__init__.py b/scepter/studio/home/__init__.py index e69de29..cc26a06 100644 --- a/scepter/studio/home/__init__.py +++ b/scepter/studio/home/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/studio/home/home.py b/scepter/studio/home/home.py index c272c05..9b87e80 100644 --- a/scepter/studio/home/home.py +++ b/scepter/studio/home/home.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. import gradio as gr diff --git a/scepter/studio/home/home_ui/__init__.py b/scepter/studio/home/home_ui/__init__.py index e69de29..cc26a06 100644 --- a/scepter/studio/home/home_ui/__init__.py +++ b/scepter/studio/home/home_ui/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/studio/home/home_ui/component_names.py b/scepter/studio/home/home_ui/component_names.py index 1d01b7a..9fffd68 100644 --- a/scepter/studio/home/home_ui/component_names.py +++ b/scepter/studio/home/home_ui/component_names.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. class DescUIName(): diff --git a/scepter/studio/home/home_ui/desc_ui.py b/scepter/studio/home/home_ui/desc_ui.py index 22b15b4..1748a83 100644 --- a/scepter/studio/home/home_ui/desc_ui.py +++ b/scepter/studio/home/home_ui/desc_ui.py @@ -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 diff --git a/scepter/studio/home/home_ui/guide_ui.py b/scepter/studio/home/home_ui/guide_ui.py index 981fcfa..54105a6 100644 --- a/scepter/studio/home/home_ui/guide_ui.py +++ b/scepter/studio/home/home_ui/guide_ui.py @@ -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 diff --git a/scepter/studio/inference/__init__.py b/scepter/studio/inference/__init__.py index e69de29..cc26a06 100644 --- a/scepter/studio/inference/__init__.py +++ b/scepter/studio/inference/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/studio/inference/inference.py b/scepter/studio/inference/inference.py index 5f8e2ff..789799a 100644 --- a/scepter/studio/inference/inference.py +++ b/scepter/studio/inference/inference.py @@ -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) diff --git a/scepter/studio/inference/inference_manager/__init__.py b/scepter/studio/inference/inference_manager/__init__.py index e69de29..cc26a06 100644 --- a/scepter/studio/inference/inference_manager/__init__.py +++ b/scepter/studio/inference/inference_manager/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/studio/inference/inference_manager/infer_runer.py b/scepter/studio/inference/inference_manager/infer_runer.py index 1007525..382e737 100644 --- a/scepter/studio/inference/inference_manager/infer_runer.py +++ b/scepter/studio/inference/inference_manager/infer_runer.py @@ -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 diff --git a/scepter/studio/inference/inference_ui/__init__.py b/scepter/studio/inference/inference_ui/__init__.py index e69de29..cc26a06 100644 --- a/scepter/studio/inference/inference_ui/__init__.py +++ b/scepter/studio/inference/inference_ui/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/studio/inference/inference_ui/component_names.py b/scepter/studio/inference/inference_ui/component_names.py index 715d10b..98c9872 100644 --- a/scepter/studio/inference/inference_ui/component_names.py +++ b/scepter/studio/inference/inference_ui/component_names.py @@ -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 = '样例' diff --git a/scepter/studio/inference/inference_ui/control_ui.py b/scepter/studio/inference/inference_ui/control_ui.py index 1610cf3..ee6d7f1 100644 --- a/scepter/studio/inference/inference_ui/control_ui.py +++ b/scepter/studio/inference/inference_ui/control_ui.py @@ -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: diff --git a/scepter/studio/inference/inference_ui/diffusion_ui.py b/scepter/studio/inference/inference_ui/diffusion_ui.py index 2c4d74d..e4f0f34 100644 --- a/scepter/studio/inference/inference_ui/diffusion_ui.py +++ b/scepter/studio/inference/inference_ui/diffusion_ui.py @@ -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)) diff --git a/scepter/studio/inference/inference_ui/gallery_ui.py b/scepter/studio/inference/inference_ui/gallery_ui.py index 9f4fb4b..2c50936 100644 --- a/scepter/studio/inference/inference_ui/gallery_ui.py +++ b/scepter/studio/inference/inference_ui/gallery_ui.py @@ -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) diff --git a/scepter/studio/inference/inference_ui/mantra_ui.py b/scepter/studio/inference/inference_ui/mantra_ui.py index 407bb3c..1be6c42 100644 --- a/scepter/studio/inference/inference_ui/mantra_ui.py +++ b/scepter/studio/inference/inference_ui/mantra_ui.py @@ -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 = [] diff --git a/scepter/studio/inference/inference_ui/model_manage_ui.py b/scepter/studio/inference/inference_ui/model_manage_ui.py index fb962c9..c3f4643 100644 --- a/scepter/studio/inference/inference_ui/model_manage_ui.py +++ b/scepter/studio/inference/inference_ui/model_manage_ui.py @@ -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) diff --git a/scepter/studio/inference/inference_ui/refiner_ui.py b/scepter/studio/inference/inference_ui/refiner_ui.py index 7158a46..4b27aac 100644 --- a/scepter/studio/inference/inference_ui/refiner_ui.py +++ b/scepter/studio/inference/inference_ui/refiner_ui.py @@ -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 diff --git a/scepter/studio/inference/inference_ui/tuner_ui.py b/scepter/studio/inference/inference_ui/tuner_ui.py index 488de78..4ca4f06 100644 --- a/scepter/studio/inference/inference_ui/tuner_ui.py +++ b/scepter/studio/inference/inference_ui/tuner_ui.py @@ -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] diff --git a/scepter/studio/preprocess/__init__.py b/scepter/studio/preprocess/__init__.py index e69de29..cc26a06 100644 --- a/scepter/studio/preprocess/__init__.py +++ b/scepter/studio/preprocess/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/studio/preprocess/caption_editor_ui/__init__.py b/scepter/studio/preprocess/caption_editor_ui/__init__.py index e69de29..cc26a06 100644 --- a/scepter/studio/preprocess/caption_editor_ui/__init__.py +++ b/scepter/studio/preprocess/caption_editor_ui/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/studio/preprocess/caption_editor_ui/component_names.py b/scepter/studio/preprocess/caption_editor_ui/component_names.py index 4d64f3a..608fd1c 100644 --- a/scepter/studio/preprocess/caption_editor_ui/component_names.py +++ b/scepter/studio/preprocess/caption_editor_ui/component_names.py @@ -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(): diff --git a/scepter/studio/preprocess/caption_editor_ui/create_dataset_ui.py b/scepter/studio/preprocess/caption_editor_ui/create_dataset_ui.py index f1b7dc9..ed3302a 100644 --- a/scepter/studio/preprocess/caption_editor_ui/create_dataset_ui.py +++ b/scepter/studio/preprocess/caption_editor_ui/create_dataset_ui.py @@ -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'] diff --git a/scepter/studio/preprocess/caption_editor_ui/dataset_gallery_ui.py b/scepter/studio/preprocess/caption_editor_ui/dataset_gallery_ui.py index 30d4e13..df8a730 100644 --- a/scepter/studio/preprocess/caption_editor_ui/dataset_gallery_ui.py +++ b/scepter/studio/preprocess/caption_editor_ui/dataset_gallery_ui.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. from __future__ import annotations import os.path diff --git a/scepter/studio/preprocess/caption_editor_ui/export_dataset_ui.py b/scepter/studio/preprocess/caption_editor_ui/export_dataset_ui.py index d4d1ef1..71be654 100644 --- a/scepter/studio/preprocess/caption_editor_ui/export_dataset_ui.py +++ b/scepter/studio/preprocess/caption_editor_ui/export_dataset_ui.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. from __future__ import annotations import os diff --git a/scepter/studio/preprocess/preprocess.py b/scepter/studio/preprocess/preprocess.py index 27dbbe2..78f1920 100644 --- a/scepter/studio/preprocess/preprocess.py +++ b/scepter/studio/preprocess/preprocess.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. import os.path import gradio as gr diff --git a/scepter/studio/self_train/__init__.py b/scepter/studio/self_train/__init__.py index e69de29..cc26a06 100644 --- a/scepter/studio/self_train/__init__.py +++ b/scepter/studio/self_train/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/studio/self_train/scripts/__init__.py b/scepter/studio/self_train/scripts/__init__.py index e69de29..cc26a06 100644 --- a/scepter/studio/self_train/scripts/__init__.py +++ b/scepter/studio/self_train/scripts/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/studio/self_train/scripts/run_task.py b/scepter/studio/self_train/scripts/run_task.py index c6ad59d..9003fc1 100644 --- a/scepter/studio/self_train/scripts/run_task.py +++ b/scepter/studio/self_train/scripts/run_task.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. import argparse import os diff --git a/scepter/studio/self_train/self_train.py b/scepter/studio/self_train/self_train.py index 9465f8a..dc375ae 100644 --- a/scepter/studio/self_train/self_train.py +++ b/scepter/studio/self_train/self_train.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. import os import gradio as gr diff --git a/scepter/studio/self_train/self_train_ui/__init__.py b/scepter/studio/self_train/self_train_ui/__init__.py index e69de29..cc26a06 100644 --- a/scepter/studio/self_train/self_train_ui/__init__.py +++ b/scepter/studio/self_train/self_train_ui/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/studio/self_train/self_train_ui/component_names.py b/scepter/studio/self_train/self_train_ui/component_names.py index 4b61343..9e94bb2 100644 --- a/scepter/studio/self_train/self_train_ui/component_names.py +++ b/scepter/studio/self_train/self_train_ui/component_names.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. # For dataset manager class InferenceUIName(): def __init__(self, language='en'): diff --git a/scepter/studio/self_train/self_train_ui/inference_ui.py b/scepter/studio/self_train/self_train_ui/inference_ui.py index 5d1f630..e7c0dde 100644 --- a/scepter/studio/self_train/self_train_ui/inference_ui.py +++ b/scepter/studio/self_train/self_train_ui/inference_ui.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. import os import gradio as gr diff --git a/scepter/studio/self_train/self_train_ui/trainer_ui.py b/scepter/studio/self_train/self_train_ui/trainer_ui.py index f00681d..a8317f5 100644 --- a/scepter/studio/self_train/self_train_ui/trainer_ui.py +++ b/scepter/studio/self_train/self_train_ui/trainer_ui.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. import copy import datetime import os diff --git a/scepter/studio/self_train/utils/__init__.py b/scepter/studio/self_train/utils/__init__.py index e69de29..cc26a06 100644 --- a/scepter/studio/self_train/utils/__init__.py +++ b/scepter/studio/self_train/utils/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/studio/self_train/utils/config_parser.py b/scepter/studio/self_train/utils/config_parser.py index 5d2c350..519a377 100644 --- a/scepter/studio/self_train/utils/config_parser.py +++ b/scepter/studio/self_train/utils/config_parser.py @@ -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) diff --git a/scepter/studio/utils/__init__.py b/scepter/studio/utils/__init__.py index e69de29..cc26a06 100644 --- a/scepter/studio/utils/__init__.py +++ b/scepter/studio/utils/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/studio/utils/env.py b/scepter/studio/utils/env.py index d6755f6..cfa7214 100644 --- a/scepter/studio/utils/env.py +++ b/scepter/studio/utils/env.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. from scepter.modules.utils.file_system import FS diff --git a/scepter/studio/utils/file.py b/scepter/studio/utils/file.py index 2ace679..0d3751a 100644 --- a/scepter/studio/utils/file.py +++ b/scepter/studio/utils/file.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. import os diff --git a/scepter/studio/utils/singleton.py b/scepter/studio/utils/singleton.py index 3ea30e0..942b20f 100644 --- a/scepter/studio/utils/singleton.py +++ b/scepter/studio/utils/singleton.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. class Singleton(object): @classmethod def get_instance(cls, cfg, **kwargs): diff --git a/scepter/studio/utils/uibase.py b/scepter/studio/utils/uibase.py index 3b81f72..ae0aeef 100644 --- a/scepter/studio/utils/uibase.py +++ b/scepter/studio/utils/uibase.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. from scepter.studio.utils.singleton import Singleton diff --git a/scepter/tools/run_inference.py b/scepter/tools/run_inference.py index 7376e33..fcb5b33 100644 --- a/scepter/tools/run_inference.py +++ b/scepter/tools/run_inference.py @@ -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}') diff --git a/scepter/tools/webui.py b/scepter/tools/webui.py index d08c364..d084f7a 100644 --- a/scepter/tools/webui.py +++ b/scepter/tools/webui.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. import argparse import datetime import os diff --git a/scepter/version.py b/scepter/version.py index 1c5ccae..5a5ead3 100644 --- a/scepter/version.py +++ b/scepter/version.py @@ -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]) diff --git a/tests/tools/test_annotators.py b/tests/tools/test_annotators.py index 631daac..89b42c0 100644 --- a/tests/tools/test_annotators.py +++ b/tests/tools/test_annotators.py @@ -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') diff --git a/tests/tools/test_train.py b/tests/tools/test_train.py index 0752644..4d390f0 100644 --- a/tests/tools/test_train.py +++ b/tests/tools/test_train.py @@ -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 ' diff --git a/tests/utils/test_fs.py b/tests/utils/test_fs.py index 3286b65..4023e53 100644 --- a/tests/utils/test_fs.py +++ b/tests/utils/test_fs.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()