update v0.0.4

This commit is contained in:
LouieStark
2024-03-31 13:08:41 +08:00
parent 35aada8ce8
commit bf53829530
106 changed files with 6927 additions and 889 deletions
Binary file not shown.

After

Width:  |  Height:  |  Size: 1.3 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 4.4 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 14 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 120 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 117 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.7 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 6.2 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 39 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 17 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 118 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 19 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 4.5 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 20 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 123 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 45 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 94 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.5 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.3 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 100 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 104 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 103 KiB

+2
View File
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
from . import impls
+260
View File
@@ -0,0 +1,260 @@
ENV:
USE_PL: False
# SET GLOBAL SYSTEM
SOLVER:
# NAME DESCRIPTION: TYPE: default: 'TrainValSolver'
NAME: TrainValSolver
# RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: ''
RESUME_FROM:
# MAX_EPOCHS DESCRIPTION: Max epochs for training. TYPE: int default: 10
MAX_EPOCHS: 200
# NUM_FOLDS DESCRIPTION: Num folds for training. TYPE: int default: 0
NUM_FOLDS: 1
# WORK_DIR DESCRIPTION: Save dir of the training log or model. TYPE: str default: ''
WORK_DIR: ./exp12/
LOG_FILE: std_log.txt
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
EVAL_INTERVAL: 1
ACCU_STEP: 1
# DO_FINAL_EVAL DESCRIPTION: If do final evaluation or not. TYPE: bool default: False
DO_FINAL_EVAL: True
# SAVE_EVAL_DATA DESCRIPTION: If save the evaluation data or not. TYPE: bool default: False
SAVE_EVAL_DATA: True
# EXTRA_KEYS DESCRIPTION: The extra keys for metric. TYPE: list default: []
EXTRA_KEYS: []
# TRAIN_DATA DESCRIPTION: Train data config. TYPE: default: ''
TRAIN_DATA:
# NAME DESCRIPTION: TYPE: default: 'ImageClassifyPublicDataset'
NAME: ImageClassifyExampleDataset
# DATASET DESCRIPTION: the public dataset name TYPE: str default: 'cifar10'
DATASET: cifar10
# DATA_ROOT DESCRIPTION: the download data save path TYPE: str default: ''
DATA_ROOT: cifar10
# MODE DESCRIPTION: test TYPE: str default: test
MODE: train
# PIN_MEMORY DESCRIPTION: pin_memory for data loader TYPE: bool default: False
PIN_MEMORY: True
# BATCH_SIZE DESCRIPTION: batch size for data TYPE: int default: 4
BATCH_SIZE: 96
# NUM_WORKERS DESCRIPTION: num workers for fetching data! TYPE: int default: 1
NUM_WORKERS: 4
# TRANSFORMS DESCRIPTION: TYPE: default:
TRANSFORMS:
# - DESCRIPTION: TYPE: default:
- # NAME DESCRIPTION: TYPE: default: 'RandomResizedCrop'
NAME: RandomResizedCrop
SIZE: 32
# RATIO DESCRIPTION: ratio TYPE: list default: [0.75, 1.3333333333333333]
RATIO: [0.75, 1.33]
# SCALE DESCRIPTION: scale TYPE: list default: [0.08, 1.0]
SCALE: [0.8, 1.0]
# INTERPOLATION DESCRIPTION: interpolation TYPE: str default: 'blilinear'
INTERPOLATION: bilinear
# INPUT_KEY DESCRIPTION: input key TYPE: str default: 'img'
INPUT_KEY: img
# OUTPUT_KEY DESCRIPTION: input key TYPE: str default: 'img'
OUTPUT_KEY: img
# BACKEND DESCRIPTION: backend, choose from pillow, cv2, torchvision TYPE: str default: 'pillow'
BACKEND: pillow
- # NAME DESCRIPTION: TYPE: default: 'RandomHorizontalFlip'
NAME: RandomHorizontalFlip
# P DESCRIPTION: P TYPE: float default: 0.5
P: 0.5
# INPUT_KEY DESCRIPTION: input key TYPE: str default: 'img'
INPUT_KEY: img
# OUTPUT_KEY DESCRIPTION: input key TYPE: str default: 'img'
OUTPUT_KEY: img
# BACKEND DESCRIPTION: backend, choose from pillow, cv2, torchvision TYPE: str default: 'pillow'
BACKEND: pillow
- # NAME DESCRIPTION: TYPE: default: 'ImageToTensor'
NAME: ImageToTensor
# INPUT_KEY DESCRIPTION: input key TYPE: str default: 'img'
INPUT_KEY: img
# OUTPUT_KEY DESCRIPTION: input key TYPE: str default: 'img'
OUTPUT_KEY: img
# BACKEND DESCRIPTION: backend, choose from pillow, cv2, torchvision TYPE: str default: 'pillow'
BACKEND: pillow
- # NAME DESCRIPTION: TYPE: default: 'Normalize'
NAME: Normalize
# MEAN DESCRIPTION: mean TYPE: list default: []
MEAN: [0.4914, 0.4822, 0.4465]
# STD DESCRIPTION: std TYPE: list default: []
STD: [0.2023, 0.1994, 0.2010]
# INPUT_KEY DESCRIPTION: input key TYPE: str default: 'img'
INPUT_KEY: img
# OUTPUT_KEY DESCRIPTION: input key TYPE: str default: 'img'
OUTPUT_KEY: img
# BACKEND DESCRIPTION: backend, choose from pillow, cv2, torchvision TYPE: str default: 'pillow'
BACKEND: pillow
- NAME: ToTensor
# KEYS DESCRIPTION: keys TYPE: list default: []
KEYS: ["img", "label"]
- # NAME DESCRIPTION: TYPE: default: 'Select'
NAME: Select
# KEYS DESCRIPTION: keys TYPE: list default: []
KEYS: ["img", "label"]
# META_KEYS DESCRIPTION: meta keys TYPE: list default: []
META_KEYS: []
# EVAL_DATA DESCRIPTION: Eval data config. TYPE: default: ''
EVAL_DATA:
# NAME DESCRIPTION: TYPE: default: 'ImageClassifyPublicDataset'
NAME: ImageClassifyPublicDataset
# DATASET DESCRIPTION: the public dataset name TYPE: str default: 'cifar10'
DATASET: cifar10
# DATA_ROOT DESCRIPTION: the download data save path TYPE: str default: ''
DATA_ROOT: ./local_data/cifar10
# MODE DESCRIPTION: test TYPE: str default: test
MODE: test
# PIN_MEMORY DESCRIPTION: pin_memory for data loader TYPE: bool default: False
PIN_MEMORY: True
# BATCH_SIZE DESCRIPTION: batch size for data TYPE: int default: 4
BATCH_SIZE: 96
# NUM_WORKERS DESCRIPTION: num workers for fetching data! TYPE: int default: 1
NUM_WORKERS: 4
# TRANSFORMS DESCRIPTION: TYPE: default:
TRANSFORMS:
# - DESCRIPTION: TYPE: default:
- # NAME DESCRIPTION: TYPE: default: 'Resize'
NAME: Resize
SIZE: 32
# INTERPOLATION DESCRIPTION: interpolation TYPE: str default: 'blilinear'
INTERPOLATION: bilinear
# INPUT_KEY DESCRIPTION: input key TYPE: str default: 'img'
INPUT_KEY: img
# OUTPUT_KEY DESCRIPTION: input key TYPE: str default: 'img'
OUTPUT_KEY: img
# BACKEND DESCRIPTION: backend, choose from pillow, cv2, torchvision TYPE: str default: 'pillow'
BACKEND: pillow
- # NAME DESCRIPTION: TYPE: default: 'ImageToTensor'
NAME: ImageToTensor
# INPUT_KEY DESCRIPTION: input key TYPE: str default: 'img'
INPUT_KEY: img
# OUTPUT_KEY DESCRIPTION: input key TYPE: str default: 'img'
OUTPUT_KEY: img
# BACKEND DESCRIPTION: backend, choose from pillow, cv2, torchvision TYPE: str default: 'pillow'
BACKEND: pillow
- # NAME DESCRIPTION: TYPE: default: 'Normalize'
NAME: Normalize
# MEAN DESCRIPTION: mean TYPE: list default: []
MEAN: [0.4914, 0.4822, 0.4465]
# STD DESCRIPTION: std TYPE: list default: []
STD: [0.2023, 0.1994, 0.2010]
# INPUT_KEY DESCRIPTION: input key TYPE: str default: 'img'
INPUT_KEY: img
# OUTPUT_KEY DESCRIPTION: input key TYPE: str default: 'img'
OUTPUT_KEY: img
# BACKEND DESCRIPTION: backend, choose from pillow, cv2, torchvision TYPE: str default: 'pillow'
BACKEND: pillow
- NAME: ToTensor
# KEYS DESCRIPTION: keys TYPE: list default: []
KEYS: ["img", "label"]
- # NAME DESCRIPTION: TYPE: default: 'Select'
NAME: Select
# KEYS DESCRIPTION: keys TYPE: list default: []
KEYS: ["img", "label"]
# META_KEYS DESCRIPTION: meta keys TYPE: list default: []
META_KEYS: []
# TRAIN_HOOKS DESCRIPTION: TYPE: default: ''
TRAIN_HOOKS:
- # NAME DESCRIPTION: TYPE: default: 'LogHook'
NAME: LogHook
# LOG_INTERVAL DESCRIPTION: the interval for log print! TYPE: int default: 10
LOG_INTERVAL: 10
# EVAL_HOOKS DESCRIPTION: TYPE: default: ''
EVAL_HOOKS:
- # NAME DESCRIPTION: TYPE: default: 'LogHook'
NAME: LogHook
# LOG_INTERVAL DESCRIPTION: the interval for log print! TYPE: int default: 10
LOG_INTERVAL: 10
# TEST_HOOKS DESCRIPTION: TYPE: default: ''
MODEL:
# NAME DESCRIPTION: TYPE: default: 'Classifier'
NAME: Classifier
# ACT_NAME DESCRIPTION: the activation function for logits, select from [softmax, sigmoid]! TYPE: str default: 'softmax'
ACT_NAME: softmax
# FREEZE_BN DESCRIPTION: if freeze bn of not TYPE: bool default: False
FREEZE_BN: False
# BACKBONE DESCRIPTION: TYPE: default: ''
BACKBONE:
# NAME DESCRIPTION: TYPE: default: 'ResNet'
NAME: ResNet
# DEPTH DESCRIPTION: the depth of network for resnet! TYPE: int default: 18
DEPTH: 18
# PRETRAINED DESCRIPTION: if load the official pretrained model or not. TYPE: bool default: False
PRETRAINED: false
#
KERNEL_SIZE: 3
# USE_RELU DESCRIPTION: use relu or not! TYPE: bool default: True
USE_RELU: True
# USE_MAXPOOL DESCRIPTION: use maxpool or not! TYPE: bool default: True
USE_MAXPOOL: false
# FIRST_CONV_STRIDE DESCRIPTION: first conv stride 1 or 2! TYPE: int default: 1
FIRST_CONV_STRIDE: 1
# FIRST_MAX_POOL_STRIDE DESCRIPTION: first max pool stride 1 or 2! TYPE: int default: 1
FIRST_MAX_POOL_STRIDE: 1
# NECK DESCRIPTION: TYPE: default: ''
NECK:
# NAME DESCRIPTION: TYPE: default: 'GlobalAveragePooling'
NAME: GlobalAveragePooling
# DIM DESCRIPTION: GlobalAveragePooling dim! TYPE: int default: 2
DIM: 2
# HEAD DESCRIPTION: TYPE: default: ''
HEAD:
# NAME DESCRIPTION: TYPE: default: 'ClassifierHead'
NAME: ClassifierHead
# DIM DESCRIPTION: representation dim! TYPE: int default: 512
DIM: 512
# NUM_CLASSES DESCRIPTION: number of classes. TYPE: int default: 10
NUM_CLASSES: 10
# DROPOUT_RATE DESCRIPTION: dropout rate, default 0. TYPE: float default: 0.0
DROPOUT_RATE: 0.0
METRIC:
# NAME DESCRIPTION: TYPE: default: 'AccuracyMetric'
NAME: AccuracyMetric
# TOPK DESCRIPTION: topk accuracy! TYPE: int default: 1
TOPK: 1
# LOSS DESCRIPTION: TYPE: default: ''
LOSS:
# NAME DESCRIPTION: TYPE: default: 'CrossEntropy'
NAME: CrossEntropy
# REDUCE DESCRIPTION: reduce is False, returns a loss per batch element instead and ignores :attr: size_average. Default: True TYPE: NoneType default: None
# REDUCE: None
# SIZE_AVERAGE DESCRIPTION: Deprecated (see :attr: reduction). By default,the losses are averaged over each loss element in the batch. Note that forsome losses, there are multiple elements per sample. If the field :attr: size_averageis set to False, the losses are instead summed for each minibatch. Ignoredwhen :attr: reduce is False. Default: True TYPE: NoneType default: None
# SIZE_AVERAGE: None
# IGNORE_INDEX DESCRIPTION: Specifies a target value that is ignoredand does not contribute to the input gradient. When :attr: size_average isTrue, the loss is averaged over non-ignored targets. Note that:attr: ignore_index is only applicable when the target contains class indices. TYPE: int default: -100
# IGNORE_INDEX: -100
# REDUCTION DESCRIPTION: Specifies the reduction to apply to the output:'none' | 'mean' | 'sum'. 'none': no reduction willbe applied, 'mean': the weighted mean of the output is taken,'sum': the output will be summed. Note: :attr: size_averageand :attr:`reduce` are in the process of being deprecated, and inthe meantime, specifying either of those two args will override:attr:`reduction`. Default: 'mean' TYPE: str default: 'mean'
# REDUCTION: mean
# LABEL_SMOOTHING DESCRIPTION: A float in [0.0, 1.0]. Specifies the amountof smoothing when computing the loss, where 0.0 means no smoothing. TYPE: float default: 0.0
# LABEL_SMOOTHING: 0.0
# OPTIMIZER DESCRIPTION: TYPE: default: ''
OPTIMIZER:
# NAME DESCRIPTION: TYPE: default: 'SGD'
NAME: SGD
# LEARNING_RATE DESCRIPTION: the initial learning rate! TYPE: float default: 0.1
LEARNING_RATE: 0.01
# MOMENTUM DESCRIPTION: the momentum! TYPE: int default: 0
MOMENTUM: 0.9
# DAMPENING DESCRIPTION: the dampening! TYPE: int default: 0
DAMPENING: 0
# WEIGHT_DECAY DESCRIPTION: the weight decay! TYPE: int default: 0
WEIGHT_DECAY: 5e-4
# NESTEROV DESCRIPTION: the nesterov! TYPE: bool default: False
NESTEROV: False
# LR_SCHEDULER DESCRIPTION: TYPE: default: ''
LR_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'CosineAnnealingLR'
NAME: CosineAnnealingLR
# T_MAX DESCRIPTION: the T max! TYPE: float default: 1.0
T_MAX: 200.0
# ETA_MIN DESCRIPTION: the eta min! TYPE: int default: 0
ETA_MIN: 0
# LAST_EPOCH DESCRIPTION: the last epoch! TYPE: int default: -1
LAST_EPOCH: -1
# METRICS DESCRIPTION: TYPE: default: ''
METRICS:
- # NAME DESCRIPTION: TYPE: default: 'AccuracyMetric'
NAME: AccuracyMetric
# TOPK DESCRIPTION: topk accuracy! TYPE: int default: 1
TOPK: 1
KEYS: ["logits", "label"]
+2
View File
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
from . import dataset
+2
View File
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
from .classifier_dataset import ImageClassifyExampleDataset
@@ -0,0 +1,80 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import numpy as np
import torchvision
from scepter.modules.data.dataset.base_dataset import BaseDataset
from scepter.modules.data.dataset.registry import DATASETS
from scepter.modules.utils.config import dict_to_yaml
@DATASETS.register_class()
class ImageClassifyExampleDataset(BaseDataset):
"""
Dataset for image classification wrapper
Args:
json_path (str): json file which contains all instances, should be a list of dict
which contains img_path and gt_label
image_dir (str or None): image directory, if None, img_path in json_path will be considered as absolute path
classes (list[str] or None): image class description
"""
para_dict = {
'DATASET': {
'value': 'cifar10',
'description': 'the public dataset name'
},
'DATA_ROOT': {
'value': '',
'description': 'the download data save path'
}
}
para_dict.update(BaseDataset.para_dict)
def __init__(self, cfg, logger=None):
super(ImageClassifyExampleDataset, self).__init__(cfg, logger=logger)
self.dataset_name = cfg.DATASET
self.data_root = cfg.DATA_ROOT
self.phase = cfg.MODE
if self.dataset_name == 'cifar10':
self.dataset = torchvision.datasets.CIFAR10(
root=self.data_root,
train=self.phase == 'train',
download=True)
def __len__(self) -> int:
return len(self.dataset)
def _get(self, index: int):
img, target = self.dataset.__getitem__(index)
ret = {
'meta': {},
'label': np.asarray(target, dtype=np.int64),
'img': img
}
return ret
def worker_init_fn(self, worker_id, num_workers=1):
super(ImageClassifyExampleDataset,
self).worker_init_fn(worker_id, num_workers=num_workers)
@staticmethod
def get_config_template():
'''
{ "ENV" :
{ "description" : "",
"A" : {
"value": 1.0,
"description": ""
}
}
}
:return:
'''
return dict_to_yaml('modename_DATA',
__class__.__name__,
ImageClassifyExampleDataset.para_dict,
set_name=True)
+7
View File
@@ -0,0 +1,7 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.tools.run_train import run
if __name__ == '__main__':
run()
+127 -16
View File
@@ -9,8 +9,8 @@
</p>
## 📖 Table of Contents
- [Introduction](#-introduction)
- [News](#-news)
- [Introduction](#-introduction)
- [Installation](#%EF%B8%8F-installation)
- [Getting Started](#-getting-started)
- [SCEPTER Studio](#%EF%B8%8F-scepter-studio)
@@ -18,6 +18,15 @@
- [Features](#-features)
- [Learn More](#-learn-more)
- [License](#license)
- [Acknowledgement](#acknowledgement)
## 🎉 News
- [2024.03]: We optimize the training UI and checkpoint management. New [LAR-Gen](https://arxiv.org/abs/2403.19534) model has been added on SCEPTER Studio, supporting `zoom-out`, `virtual try on`, `inpainting`.
- [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.
- [2023.12]: We release [🪄SCEPTER](https://github.com/modelscope/scepter/) library.
## 📝 Introduction
@@ -28,7 +37,7 @@ Main Feature:
- Task:
- Text-to-image generation
- Controllable image synthesis
- Image editing (TODO)
- Image editing
- Training / Inference:
- Distribute: DDP / FSDP / FairScale / Xformers
- File system: Local / Http / OSS / Modelscope
@@ -40,15 +49,9 @@ Main Feature:
Currently supported approaches (and counting):
1. SD Series: [Stable Diffusion v1.5](https://huggingface.co/runwayml/stable-diffusion-v1-5) / [Stable Diffusion v2.1](https://huggingface.co/runwayml/stable-diffusion-v1-5) / [Stable Diffusion XL](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0)
2. SCEdit: [SCEdit: Efficient and Controllable Image Diffusion Generation via Skip Connection Editing](https://arxiv.org/abs/2312.11392) [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=SCEdit&color=red&logo=arxiv)](https://arxiv.org/abs/2312.11392) [![Page link](https://img.shields.io/badge/Page-SCEdit-Gree)](https://scedit.github.io/)
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.
- [2023.12]: We release [🪄SCEPTER](https://github.com/modelscope/scepter/) library.
2. SCEdit(CVPR2024): [SCEdit: Efficient and Controllable Image Diffusion Generation via Skip Connection Editing](https://arxiv.org/abs/2312.11392) [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=SCEdit&color=red&logo=arxiv)](https://arxiv.org/abs/2312.11392) [![Page link](https://img.shields.io/badge/Page-SCEdit-Gree)](https://scedit.github.io/)
3. Res-Tuning(NeurIPS2023 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/)
4. LAR-Gen: [Locate, Assign, Refine: Taming Customized Image Inpainting with Text-Subject Guidance](https://arxiv.org/abs/2403.19534) [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=LARGen&color=red&logo=arxiv)](https://arxiv.org/abs/2403.19534) [![Page link](https://img.shields.io/badge/Page-LARGen-Gree)](https://ali-vilab.github.io/largen-page/)
## 🛠️ Installation
@@ -94,7 +97,7 @@ For the data format used by SCEPTER Studio, please refer to [3D_example_csv.zip]
To facilitate starting training in command-line mode, you can use a dataset in text format, please refer to [3D_example_txt.zip](https://modelscope.cn/api/v1/models/damo/scepter/repo?Revision=master&FilePath=datasets/3D_example_txt.zip)
```shell
mkdir -p cache/dataset/ && wget 'https://modelscope.cn/api/v1/models/damo/scepter_scedit/repo?Revision=master&FilePath=dataset/3D_example_txt.zip' -O cache/dataset/3D_example_txt.zip && unzip cache/dataset/3D_example_txt.zip -d cache/dataset/ && rm cache/dataset/3D_example_txt.zip
mkdir -p cache/datasets/ && wget 'https://modelscope.cn/api/v1/models/damo/scepter_scedit/repo?Revision=master&FilePath=dataset/3D_example_txt.zip' -O cache/datasets/3D_example_txt.zip && unzip cache/datasets/3D_example_txt.zip -d cache/datasets/ && rm cache/datasets/3D_example_txt.zip
```
### Training
@@ -165,6 +168,14 @@ 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
```
### Customize Modules
Refer to `example`, build your modules in `example/impls`.
```python
cd example
python run.py --cfg example.yaml
```
## 🖥️ SCEPTER Studio
### Launch
@@ -181,16 +192,94 @@ 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.
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.
* LAR-Gen: we release `zoom-out`, `virtual try on`, `inpainting(text guided)`, `inpainting(text + reference image guided)` image editing capabilities.
Please note that the **Data Preprocess** button must be clicked before clicking the **Generate** button.
<p align="center">
<img src="./asset/images/largen.gif">
</p>
### 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
### LAR-Gen: Zoom-Out
<table>
<tr>
<td><strong>Origin Image</strong><br>Prompt: a temple on fire</td>
<td><strong>Zoom-Out</strong><br>CenterAround:0.75</td>
<td><strong>Zoom-Out</strong><br>CenterAround:0.75</td>
<td><strong>Zoom-Out</strong><br>CenterAround:0.75</td>
<td><strong>Zoom-Out</strong><br>CenterAround:0.75</td>
</tr>
<tr>
<td><img src="asset/images/zoom_out/ex1_scene_im.jpeg" width="240"></td>
<td><img src="asset/images/zoom_out/ex1_zoom_out1.jpeg" width="240"></td>
<td><img src="./asset/images/zoom_out/ex1_zoom_out2.jpeg" width="240"></td>
<td><img src="./asset/images/zoom_out/ex1_zoom_out3.jpeg" width="240"></td>
<td><img src="./asset/images/zoom_out/ex1_zoom_out4.jpeg" width="240"></td>
</tr>
</table>
### LAR-Gen: Virtual-Try-on
<table>
<tr>
<td><strong>Model Image</strong></td>
<td><strong>Model Mask</strong></td>
<td><strong>Clothing Image</strong></td>
<td><strong>Clothing Mask</strong></td>
<td><strong>Try-on Output</strong></td>
</tr>
<tr>
<td><img src="asset/images/virtual_try_on/model.jpg" width="240"></td>
<td><img src="asset/images/virtual_try_on/ex2_scene_mask.jpg" width="240"></td>
<td><img src="asset/images/virtual_try_on/tshirt.jpg" width="240"></td>
<td><img src="asset/images/virtual_try_on/ex2_subject_mask.jpg" width="240"></td>
<td><img src="asset/images/virtual_try_on/try_on_out.jpeg" width="240"></td>
</tr>
</table>
### LAR-Gen: Inpainting(Text-guided)
<table>
<tr>
<td><strong>Origin Image</strong><br>Prompt: a blue and white porcelain</td>
<td><strong>Inpainting Mask1</strong></td>
<td><strong>Inpainting Output1</strong></td>
<td><strong>Inpainting Mask2</strong><br>Prompt: a clock</td>
<td><strong>Inpainting Output2</strong></td>
</tr>
<tr>
<td><img src="asset/images/inpainting_text/ex3_scene_im.jpg" width="240"></td>
<td><img src="asset/images/inpainting_text/ex3_scene_mask.jpg" width="240"></td>
<td><img src="asset/images/inpainting_text/inpainting_text.jpeg" width="240"></td>
<td><img src="asset/images/inpainting_text/ex3_scene_mask2.jpg" width="240"></td>
<td><img src="asset/images/inpainting_text/inpainting_text2.jpeg" width="240"></td>
</tr>
</table>
### LAR-Gen: Inpainting(Text + Reference Image Guided)
<table>
<tr>
<td><strong>Origin Image</strong><br>Prompt: a dog wearing sunglasses</td>
<td><strong>Origin Mask</strong></td>
<td><strong>Reference Image</strong></td>
<td><strong>Reference Mask</strong></td>
<td><strong>Inpainting Output</strong></td>
</tr>
<tr>
<td><img src="asset/images/inpainting_text_ref/ex4_scene_im.jpg" width="240"></td>
<td><img src="asset/images/inpainting_text_ref/ex4_scene_mask.jpg" width="240"></td>
<td><img src="asset/images/inpainting_text_ref/ex4_subject_im.jpg" width="240"></td>
<td><img src="asset/images/inpainting_text_ref/ex4_subject_mask.jpg" width="240"></td>
<td><img src="asset/images/inpainting_text_ref/inpainting_text_ref.jpeg" width="240"></td>
</tr>
</table>
### Dragon Year Special: Dragon Tuner
<table>
@@ -244,22 +333,31 @@ We deploy a work studio on Modelscope that includes only the inference tab, plea
| SD 2.1 | 🪄 | 🪄 | 🪄 | 🪄 | 🪄 |
| SD XL | 🪄 | 🪄 | 🪄 | 🪄 | 🪄 |
### Image Editing
- LAR-Gen
| **Model** | **Locate** | **Assign** | **Refine** |
|:---------:|:----------:|:----------:|:----------:|
| SD XL | 🪄 | 🪄 | ⏳ |
### Model URL
- ✅ indicates support for both training and inference.
- 🪄 denotes that the model has been published.
- ⏳ denotes that the module has not been integrated currently.
- More models will be released in the future.
| Model | URL |
|--------|------------------------------------------------------------------------------------------------------------------------------------------------|
| SCEdit | [ModelScope](https://modelscope.cn/models/damo/scepter_scedit/summary) [HuggingFace](https://huggingface.co/scepter-studio/scepter_scedit) |
| SCEdit | [ModelScope](https://modelscope.cn/models/iic/scepter_scedit/summary) [HuggingFace](https://huggingface.co/scepter-studio/scepter_scedit) |
| LAR-Gen | [ModelScope](https://www.modelscope.cn/models/iic/LARGEN/summary) |
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.
## 🔍 Learn More
- [Alibaba TongYi Vision Intelligence Lab](https://github.com/damo-vilab)
- [Alibaba TongYi Vision Intelligence Lab](https://github.com/ali-vilab)
Discover more about open-source projects on image generation, video generation, and editing tasks.
@@ -271,7 +369,20 @@ PS: Scripts running within the SCEPTER framework will automatically fetch and lo
SWIFT (Scalable lightWeight Infrastructure for Fine-Tuning) is an extensible framwork designed to faciliate lightweight model fine-tuning and inference.
## BibTeX
If our work is useful for your research, please consider citing:
```bibtex
@misc{scepter,
title = {SCEPTER, https://github.com/modelscope/scepter},
author = {SCEPTER},
year = {2023}
}
```
## License
This project is licensed under the [Apache License (Version 2.0)](https://github.com/modelscope/modelscope/blob/master/LICENSE).
## Acknowledgement
Thanks to [Stability-AI](https://github.com/Stability-AI), [SWIFT library](https://github.com/modelscope/swift/) and [Fooocus](https://github.com/lllyasviel/Fooocus) for their awesome work.
+3
View File
@@ -1,3 +1,5 @@
albumentations
bezier
einops
modelscope
ms-swift>=1.5.2
@@ -6,6 +8,7 @@ open_clip_torch
opencv-python
opencv_transforms>=0.0.6
oss2>=2.15.0
pycocotools
pyyaml>=5.3.1
scikit-image
torchsde
+1
View File
@@ -1,3 +1,4 @@
git+https://github.com/cocodataset/panopticapi.git
torch==2.0.1
torchvision==0.15.2
xformers==0.0.21
+1
View File
@@ -1,2 +1,3 @@
gradio>=3.47.1,<4.0.0
imagehash
psutil
@@ -79,6 +79,7 @@ EXTENSION_PARAS:
MANTRA_BOOK: scepter/methods/studio/extensions/mantra_book/mantra_book.yaml
OFFICIAL_TUNERS: scepter/methods/studio/extensions/tuners/official_tuners.yaml
OFFICIAL_CONTROLLERS: scepter/methods/studio/extensions/controllers/official_controllers.yaml
TUNER_MANAGER: scepter/methods/studio/tuner_manager/tuner_manager.yaml
CONTROLABLE_ANNOTATORS:
-
NAME: "CannyAnnotator"
@@ -0,0 +1,268 @@
NAME: LARGEN
IS_DEFAULT: True
DEFAULT_PARAS:
PARAS:
RESOLUTIONS: [[1024, 1024]]
INPUT:
IMAGE:
ORIGINAL_SIZE_AS_TUPLE: [1024, 1024]
TARGET_SIZE_AS_TUPLE: [1024, 1024]
AESTHETIC_SCORE: 6.0
NEGATIVE_AESTHETIC_SCORE: 2.5
PROMPT: ""
NEGATIVE_PROMPT: ""
PROMPT_PREFIX: ""
CROP_COORDS_TOP_LEFT: [0, 0]
SAMPLE: ddim
SAMPLE_STEPS: 50
GUIDE_SCALE: 7.5
GUIDE_RESCALE: 0.5
DISCRETIZATION: trailing
REFINE_SAMPLE: ddim
REFINE_GUIDE_SCALE: 7.5
REFINE_GUIDE_RESCALE: 0.5
REFINE_DISCRETIZATION: trailing
OUTPUT:
LATENT:
BEFORE_REFINE_IMAGES:
IMAGES:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
-
NAME: encode
DTYPE: float32
INPUT: ["IMAGE"]
-
NAME: decode
DTYPE: float32
INPUT: ["LATENT"]
PARAS:
# SCALE_FACTOR DESCRIPTION: The vae embeding scale. TYPE: float default: 0.18215
SCALE_FACTOR: 0.13025
SIZE_FACTOR: 8
DIFFUSION_MODEL:
FUNCTION:
-
NAME: forward
DTYPE: float16
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE", "GUIDE_RESCALE", "DISCRETIZATION"]
COND_STAGE_MODEL:
FUNCTION:
-
NAME: encode
DTYPE: float16
INPUT: ["ORIGINAL_SIZE_AS_TUPLE", "CROP_COORDS_TOP_LEFT", "PROMPT", "NEGATIVE_PROMPT"]
REFINER_MODEL:
FUNCTION:
-
NAME: forward
DTYPE: float16
INPUT: ["SAMPLE_STEPS", "REFINE_SAMPLE", "REFINE_GUIDE_SCALE", "REFINE_GUIDE_RESCALE", "REFINE_DISCRETIZATION"]
REFINER_COND_MODEL:
FUNCTION:
-
NAME: encode
DTYPE: float16
INPUT: ["ORIGINAL_SIZE_AS_TUPLE", "AESTHETIC_SCORE", "NEGATIVE_AESTHETIC_SCORE", "CROP_COORDS_TOP_LEFT", "PROMPT", "NEGATIVE_PROMPT"]
MODEL:
PRETRAINED_MODEL: ms://damo/LARGEN@models/largen_ckpt_s22k.pth
# SCHEDULE_ARGS DESCRIPTION: TYPE: default: ''
SCHEDULE:
PARAMETERIZATION: "eps"
TIMESTEPS: 1000
ZERO_TERMINAL_SNR: False
SCHEDULE_ARGS:
# NAME DESCRIPTION: TYPE: default: ''
NAME: "scaled_linear"
BETA_MIN: 0.00085
BETA_MAX: 0.0120
# DIFFUSION_MODEL DESCRIPTION: TYPE: default: ''
DIFFUSION_MODEL:
# NAME DESCRIPTION: TYPE: default: 'DiffusionUNetXL'
NAME: LargenUNetXL
# PRETRAINED_MODEL DESCRIPTION: Whole model's pretrained model path. TYPE: NoneType default: None
PRETRAINED_MODEL:
# IN_CHANNELS DESCRIPTION: Unet channels for input, considering the input image's channels. TYPE: int default: 4
IN_CHANNELS: 9
# OUT_CHANNELS DESCRIPTION: Unet channels for output, considering the input image's channels. TYPE: int default: 4
OUT_CHANNELS: 4
# NUM_RES_BLOCKS DESCRIPTION: The blocks's number of res. TYPE: int default: 2
NUM_RES_BLOCKS: 2
# MODEL_CHANNELS DESCRIPTION: base channel count for the model. TYPE: int default: 320
MODEL_CHANNELS: 320
# ATTENTION_RESOLUTIONS DESCRIPTION: A collection of downsample rates at which attention will take place. May be a set, list, or tuple. For example, if this contains 4, then at 4x downsampling, attentio will be used. TYPE: list default: [4, 2]
ATTENTION_RESOLUTIONS: [4, 2]
# DROPOUT DESCRIPTION: The dropout rate. TYPE: int default: 0
DROPOUT: 0
# CHANNEL_MULT DESCRIPTION: channel multiplier for each level of the UNet. TYPE: list default: [1, 2, 4]
CHANNEL_MULT: [1, 2, 4]
# CONV_RESAMPLE DESCRIPTION: Use conv to resample when downsample. TYPE: bool default: True
CONV_RESAMPLE: True
# DIMS DESCRIPTION: The Conv dims which 2 represent Conv2D. TYPE: int default: 2
DIMS: 2
# NUM_CLASSES DESCRIPTION: The class num for class guided setting, also can be set as continuous. TYPE: str default: 'sequential'
NUM_CLASSES: sequential
# USE_CHECKPOINT DESCRIPTION: Use gradient checkpointing to reduce memory usage. TYPE: bool default: False
USE_CHECKPOINT: False
# NUM_HEADS DESCRIPTION: The number of attention heads in each attention layer. TYPE: int default: -1
NUM_HEADS: -1
# NUM_HEADS_CHANNELS DESCRIPTION: If specified, ignore num_heads and instead use a fixed channel width per attention head. TYPE: int default: 64
NUM_HEADS_CHANNELS: 64
# USE_SCALE_SHIFT_NORM DESCRIPTION: The scale and shift for the outnorm of RESBLOCK, use a FiLM-like conditioning mechanism. TYPE: bool default: False
USE_SCALE_SHIFT_NORM: False
# RESBLOCK_UPDOWN DESCRIPTION: Use residual blocks for up/downsampling, if False use Conv. TYPE: bool default: False
RESBLOCK_UPDOWN: False
# USE_NEW_ATTENTION_ORDER DESCRIPTION: Whether use new attention(qkv before split heads or not) or not. TYPE: bool default: True
USE_NEW_ATTENTION_ORDER: True
# USE_SPATIAL_TRANSFORMER DESCRIPTION: Custom transformer which support the context, if context_dim is not None, the parameter must set True TYPE: bool default: True
USE_SPATIAL_TRANSFORMER: True
# TRANSFORMER_DEPTH DESCRIPTION: Custom transformer's depth, valid when USE_SPATIAL_TRANSFORMER is True. TYPE: list default: [1, 2, 10]
TRANSFORMER_DEPTH: [1, 2, 10]
# TRANSFORMER_DEPTH_MIDDLE DESCRIPTION: Custom transformer's depth of middle block, If set None, use TRANSFORMER_DEPTH last value. TYPE: NoneType default: None
# TRANSFORMER_DEPTH_MIDDLE: None
# CONTEXT_DIM DESCRIPTION: Custom context info, if set, USE_SPATIAL_TRANSFORMER also set True. TYPE: int default: 2048
CONTEXT_DIM: 2048
# DISABLE_SELF_ATTENTIONS DESCRIPTION: Whether disable the self-attentions on some level, should be a list, [False, True, ...] TYPE: NoneType default: None
# DISABLE_SELF_ATTENTIONS: None
# NUM_ATTENTION_BLOCKS DESCRIPTION: The number of attention blocks for attention layer. TYPE: NoneType default: None
# NUM_ATTENTION_BLOCKS: None
# DISABLE_MIDDLE_SELF_ATTN DESCRIPTION: Whether disable the self-attentions in middle blocks. TYPE: bool default: False
DISABLE_MIDDLE_SELF_ATTN: False
# USE_LINEAR_IN_TRANSFORMER DESCRIPTION: Custom transformer's parameter, valid when USE_SPATIAL_TRANSFORMER is True. TYPE: bool default: True
USE_LINEAR_IN_TRANSFORMER: True
# ADM_IN_CHANNELS DESCRIPTION: Used when num_classes == 'sequential' or 'timestep'. TYPE: int default: 2816
ADM_IN_CHANNELS: 2816
# USE_SENTENCE_EMB DESCRIPTION: Used sentence emb or not, default False. TYPE: bool default: False
USE_SENTENCE_EMB: False
# USE_WORD_MAPPING DESCRIPTION: Used word mapping or not, default False. TYPE: bool default: False
USE_WORD_MAPPING: False
TRANSFORMER_BLOCK_TYPE: att_v2
IMAGE_SCALE: 1.0
USE_REFINE: False
# FIRST_STAGE_MODEL DESCRIPTION: TYPE: default: ''
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
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 DESCRIPTION: TYPE: default: ''
COND_STAGE_MODEL:
# NAME DESCRIPTION: TYPE: default: 'GeneralConditioner'
NAME: GeneralConditioner
USE_GRAD: False
# EMBEDDERS DESCRIPTION: TYPE: default: ''
EMBEDDERS:
-
# NAME DESCRIPTION: TYPE: default: 'FrozenCLIPEmbedder'
NAME: FrozenCLIPEmbedder
# PRETRAINED_MODEL DESCRIPTION: TYPE: str default: ''
PRETRAINED_MODEL: ms://AI-ModelScope/clip-vit-large-patch14
TOKENIZER_PATH: ms://AI-ModelScope/clip-vit-large-patch14
# MAX_LENGTH DESCRIPTION: TYPE: int default: 77
MAX_LENGTH: 77
# FREEZE DESCRIPTION: TYPE: bool default: True
FREEZE: True
# LAYER DESCRIPTION: TYPE: str default: 'last'
LAYER: hidden
# LAYER_IDX DESCRIPTION: TYPE: NoneType default: None
LAYER_IDX: 11
# USE_FINAL_LAYER_NORM DESCRIPTION: TYPE: bool default: False
USE_FINAL_LAYER_NORM: False
UCG_RATE: 0.0
INPUT_KEYS: ["prompt"]
LEGACY_UCG_VALUE:
-
# NAME DESCRIPTION: TYPE: default: 'FrozenOpenCLIPEmbedder2'
NAME: FrozenOpenCLIPEmbedder2
# ARCH DESCRIPTION: TYPE: str default: 'ViT-H-14'
ARCH: ViT-bigG-14
# MAX_LENGTH DESCRIPTION: TYPE: int default: 77
MAX_LENGTH: 77
# FREEZE DESCRIPTION: TYPE: bool default: True
FREEZE: True
# ALWAYS_RETURN_POOLED DESCRIPTION: Whether always return pooled results or not ,default False. TYPE: bool default: False
ALWAYS_RETURN_POOLED: True
# LEGACY DESCRIPTION: Whether use legacy returnd feature or not ,default True. TYPE: bool default: True
LEGACY: False
# LAYER DESCRIPTION: TYPE: str default: 'last'
LAYER: penultimate
UCG_RATE: 0.0
INPUT_KEYS: ["prompt"]
LEGACY_UCG_VALUE:
-
# NAME DESCRIPTION: TYPE: default: 'ConcatTimestepEmbedderND'
NAME: ConcatTimestepEmbedderND
# OUT_DIM DESCRIPTION: Output dim TYPE: int default: 256
OUT_DIM: 256
UCG_RATE: 0.0
INPUT_KEYS: ["original_size_as_tuple"]
LEGACY_UCG_VALUE:
-
# NAME DESCRIPTION: TYPE: default: 'ConcatTimestepEmbedderND'
NAME: ConcatTimestepEmbedderND
# OUT_DIM DESCRIPTION: Output dim TYPE: int default: 256
OUT_DIM: 256
UCG_RATE: 0.0
INPUT_KEYS: ["crop_coords_top_left"]
LEGACY_UCG_VALUE:
-
# NAME DESCRIPTION: TYPE: default: 'ConcatTimestepEmbedderND'
NAME: ConcatTimestepEmbedderND
# OUT_DIM DESCRIPTION: Output dim TYPE: int default: 256
OUT_DIM: 256
UCG_RATE: 0.0
INPUT_KEYS: ["target_size_as_tuple"]
LEGACY_UCG_VALUE:
-
NAME: IPAdapterPlusEmbedder
CLIP_DIR: ms://damo/LARGEN@models/clip_encoder/
PRETRAINED_MODEL: ms://damo/LARGEN@models/ip-adapter-plus_sdxl_vit-h.bin
INPUT_KEYS: [ "ref_ip", "ref_detail" ]
IN_DIM: 1280
HEADS: 20
CROSSATTN_DIM: 2048
-
NAME: TransparentEmbedder
INPUT_KEYS: [ "tar_x0", "tar_mask_latent" ]
-
NAME: NoiseConcatEmbedder
INPUT_KEYS: [ "tar_mask_latent", "masked_x0" ]
-
NAME: TransparentEmbedder
INPUT_KEYS: [ "ref_x0" ]
-
NAME: TransparentEmbedder
INPUT_KEYS: [ "task" ]
-
NAME: TransparentEmbedder
INPUT_KEYS: [ "image_scale" ]
+6 -2
View File
@@ -49,11 +49,11 @@ BANNER: |
<div class="qr-codes">
<div class="qr-code-container">
<img src="https://modelscope.cn/api/v1/models/damo/scepter/repo?Revision=master&FilePath=assets/scepter_studio/ms_scepter_studio_qr.png" alt="ms_scepter_studio_qr">
<div class="caption">Modelscope Studio</div>
<div class="caption"><a href="https://www.modelscope.cn/studios/iic/scepter_studio">Modelscope Studio</a></div>
</div>
<div class="qr-code-container">
<img src="https://modelscope.cn/api/v1/models/damo/scepter/repo?Revision=master&FilePath=assets/scepter_studio/scepter_github_qr.png" alt="scepter_github_qr">
<div class="caption">Github</div>
<div class="caption"><a href="https://github.com/modelscope/scepter">Github</a></div>
</div>
</div>
</div>
@@ -79,6 +79,10 @@ INTERFACE:
NAME_EN: Train
IFID: self_train
CONFIG: scepter/methods/studio/self_train/self_train.yaml
- NAME: 模型管理
NAME_EN: Tuner Management
IFID: tuner_manager
CONFIG: scepter/methods/studio/tuner_manager/tuner_manager.yaml
- NAME: 推理
NAME_EN: Inference
IFID: inference
@@ -11,13 +11,13 @@ META:
DEFAULT_SAMPLER: "dpmpp_2s_ancestral"
DEFAULT_SAMPLE_STEPS: 40
INFERENCE_N_PROMPT: ""
RESOLUTION: 1024
RESOLUTION: [1024, 1024]
PARAS:
-
TRAIN_BATCH_SIZE: 2
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: 1024
RESOLUTION: [1024, 1024]
MEMORY: 29000
EPOCHS: 50
SAVE_INTERVAL: 25
@@ -29,7 +29,7 @@ META:
TRAIN_BATCH_SIZE: 2
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: 1024
RESOLUTION: [1024, 1024]
MEMORY: 29000
EPOCHS: 50
SAVE_INTERVAL: 25
@@ -41,7 +41,7 @@ META:
TRAIN_BATCH_SIZE: 2
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: 1024
RESOLUTION: [1024, 1024]
MEMORY: 29000
EPOCHS: 200
SAVE_INTERVAL: 25
@@ -53,7 +53,7 @@ META:
TRAIN_BATCH_SIZE: 2
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: 1024
RESOLUTION: [1024, 1024]
MEMORY: 29000
EPOCHS: 200
SAVE_INTERVAL: 25
@@ -65,7 +65,7 @@ META:
TRAIN_BATCH_SIZE: 2
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: 1024
RESOLUTION: [1024, 1024]
MEMORY: 29000
EPOCHS: 50
SAVE_INTERVAL: 25
@@ -146,10 +146,12 @@ SOLVER:
MAX_EPOCHS: -1
# NUM_FOLDS DESCRIPTION: Num folds for training. TYPE: int default: 1
NUM_FOLDS: 1
#
EVAL_INTERVAL: -1
# WORK_DIR DESCRIPTION: Save dir of the training log or model. TYPE: str default: ''
WORK_DIR:
# LOG_FILE DESCRIPTION: Save log path. TYPE: str default: ''
LOG_FILE: stg_log.txt
LOG_FILE: std_log.txt
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/data"
@@ -586,15 +588,44 @@ SOLVER:
INPUT_KEY: [ 'img', 'img_original_size_as_tuple', 'img_target_size_as_tuple', 'img_crop_coords_top_left']
OUTPUT_KEY: [ 'image', 'original_size_as_tuple', 'target_size_as_tuple', 'crop_coords_top_left' ]
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
PROMPT_FILE:
PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ]
IMAGE_SIZE: [ 1024, 1024 ]
FIELDS: [ "prompt" ]
DELIMITER: '#;#'
PROMPT_PREFIX: ''
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
TRANSFORMS:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: BackwardHook
# GRADIENT_CLIP: 1.0
-
NAME: BackwardHook
PRIORITY: 0
- NAME: LogHook
LOG_INTERVAL: 50
-
NAME: LogHook
LOG_INTERVAL: 10
SHOW_GPU_MEM: True
-
NAME: TensorboardLogHook
-
NAME: CheckpointHook
SAVE_LAST: True
INTERVAL: 10000
PRIORITY: 200
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
#
EVAL_HOOKS:
-
NAME: ProbeDataHook
PROB_INTERVAL: 100
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
SAVE_PROBE_PREFIX: 'image'
@@ -8,3 +8,13 @@ SAMPLERS:
NAME: 'dpmpp_2m_sde'
-
NAME: 'dpmpp_2s_ancestral'
TRAIN_PARAS:
RESOLUTIONS:
VALUES: [[256, 256], [320, 180], [180, 320],
[512, 512], [640, 360], [360, 640],
[768, 768], [960, 540], [540, 960],
[1024, 1024], [1280, 720], [720, 1280]]
DEFAULT: [1024, 1024]
EVAL_PROMPTS:
- a boy wearing a jacket
- a dog running on the lawn
@@ -10,13 +10,13 @@ META:
DEFAULT_SAMPLER: "ddim"
DEFAULT_SAMPLE_STEPS: 40
INFERENCE_N_PROMPT: ""
RESOLUTION: 512
RESOLUTION: [512, 512]
PARAS:
-
TRAIN_BATCH_SIZE: 2
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: 512
RESOLUTION: [512, 512]
MEMORY: 29000
EPOCHS: 50
SAVE_INTERVAL: 25
@@ -28,7 +28,7 @@ META:
TRAIN_BATCH_SIZE: 2
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: 512
RESOLUTION: [512, 512]
MEMORY: 29000
EPOCHS: 50
SAVE_INTERVAL: 25
@@ -40,7 +40,7 @@ META:
TRAIN_BATCH_SIZE: 2
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: 512
RESOLUTION: [512, 512]
MEMORY: 29000
EPOCHS: 200
SAVE_INTERVAL: 25
@@ -53,7 +53,7 @@ META:
TRAIN_BATCH_SIZE: 2
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: 512
RESOLUTION: [512, 512]
MEMORY: 29000
EPOCHS: 200
SAVE_INTERVAL: 25
@@ -66,7 +66,7 @@ META:
TRAIN_BATCH_SIZE: 4
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: 512
RESOLUTION: [512, 512]
MEMORY: 29000
EPOCHS: 50
SAVE_INTERVAL: 25
@@ -135,6 +135,7 @@ SOLVER:
MAX_EPOCHS: -1
NUM_FOLDS: 1
ACCU_STEP: 1
EVAL_INTERVAL: -1
#
WORK_DIR:
LOG_FILE: std_log.txt
@@ -298,15 +299,44 @@ SOLVER:
KEYS: [ 'image', 'prompt' ]
META_KEYS: [ 'data_key' ]
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
PROMPT_FILE:
PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ]
IMAGE_SIZE: [ 512, 512 ]
FIELDS: [ "prompt" ]
DELIMITER: '#;#'
PROMPT_PREFIX: ''
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
TRANSFORMS:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: BackwardHook
# GRADIENT_CLIP: 1.0
-
NAME: BackwardHook
PRIORITY: 0
- NAME: LogHook
LOG_INTERVAL: 50
-
NAME: LogHook
LOG_INTERVAL: 10
SHOW_GPU_MEM: True
-
NAME: TensorboardLogHook
-
NAME: CheckpointHook
SAVE_LAST: True
INTERVAL: 10000
PRIORITY: 200
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
#
EVAL_HOOKS:
-
NAME: ProbeDataHook
PROB_INTERVAL: 100
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
SAVE_PROBE_PREFIX: 'image'
@@ -10,13 +10,13 @@ META:
DEFAULT_SAMPLER: "ddim"
DEFAULT_SAMPLE_STEPS: 40
INFERENCE_N_PROMPT: ""
RESOLUTION: 768
RESOLUTION: [768, 768]
PARAS:
-
TRAIN_BATCH_SIZE: 2
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: 768
RESOLUTION: [768, 768]
MEMORY: 29000
EPOCHS: 50
SAVE_INTERVAL: 25
@@ -28,7 +28,7 @@ META:
TRAIN_BATCH_SIZE: 2
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: 768
RESOLUTION: [768, 768]
MEMORY: 29000
EPOCHS: 50
SAVE_INTERVAL: 25
@@ -40,7 +40,7 @@ META:
TRAIN_BATCH_SIZE: 2
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: 768
RESOLUTION: [768, 768]
MEMORY: 29000
EPOCHS: 200
SAVE_INTERVAL: 25
@@ -78,6 +78,7 @@ SOLVER:
MAX_EPOCHS: -1
NUM_FOLDS: 1
ACCU_STEP: 1
EVAL_INTERVAL: -1
#
WORK_DIR:
LOG_FILE: std_log.txt
@@ -240,15 +241,44 @@ SOLVER:
KEYS: [ 'image', 'prompt' ]
META_KEYS: [ 'data_key' ]
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
PROMPT_FILE:
PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ]
IMAGE_SIZE: [ 768, 768 ]
FIELDS: [ "prompt" ]
DELIMITER: '#;#'
PROMPT_PREFIX: ''
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
TRANSFORMS:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: BackwardHook
# GRADIENT_CLIP: 1.0
-
NAME: BackwardHook
PRIORITY: 0
- NAME: LogHook
LOG_INTERVAL: 50
-
NAME: LogHook
LOG_INTERVAL: 10
SHOW_GPU_MEM: True
-
NAME: TensorboardLogHook
-
NAME: CheckpointHook
SAVE_LAST: True
INTERVAL: 10000
PRIORITY: 200
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
#
EVAL_HOOKS:
-
NAME: ProbeDataHook
PROB_INTERVAL: 100
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
SAVE_PROBE_PREFIX: 'image'
@@ -0,0 +1,2 @@
WORK_DIR: "tuner_manager"
TUNER_LIST_YAML: "tuner_list.yaml"
+12 -7
View File
@@ -244,12 +244,18 @@ class Text2ImageDataset(BaseDataset):
image_size = [image_size, image_size]
assert isinstance(image_size, Iterable) and len(image_size) == 2
prompt_file = cfg.PROMPT_FILE
with FS.get_object(prompt_file) as local_data:
if cfg.PROMPT_FILE is not None and cfg.PROMPT_FILE != '':
prompt_file = cfg.PROMPT_FILE
with FS.get_object(prompt_file) as local_data:
rows = [
i.split(delimiter,
len(fields) - 1)
for i in local_data.decode('utf-8').strip().split('\n')
]
else:
rows = [
i.split(delimiter,
len(fields) - 1)
for i in local_data.decode('utf-8').strip().split('\n')
len(fields) - 1) for i in cfg.PROMPT_DATA
]
self.items = list()
@@ -263,10 +269,9 @@ class Text2ImageDataset(BaseDataset):
item['meta']['img_path'] = os.path.join(path_prefix, value)
elif key in ['width', 'height']:
item['meta'][key] = int(value)
elif key != 'meta':
item[key] = value
else:
continue
item['meta'][key] = value
self.items.append(item)
if use_num > 0:
self.items = self.items[:use_num]
@@ -2,18 +2,23 @@
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import os
import warnings
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
try:
from swift import SwiftModel
except Exception:
warnings.warn('Import swift failed, please check it.')
class ControlInference():
def __init__(self, logger=None):
@@ -474,10 +474,14 @@ class DiffusionInference():
enabled=dtype == 'float16',
dtype=getattr(torch, dtype)):
if self.tokenizer:
if not hasattr(get_model(self.cond_stage_model), 'tokenizer'):
setattr(get_model(self.cond_stage_model), 'tokenizer',
self.tokenizer)
context = getattr(get_model(self.cond_stage_model),
function_name)(batch['tokens'])
null_context = getattr(get_model(self.cond_stage_model),
function_name)(batch_uc['tokens'])
else:
context = getattr(get_model(self.cond_stage_model),
function_name)(batch)
@@ -558,12 +562,13 @@ class DiffusionInference():
seed=seed,
condition_fn=None,
clamp=None,
sharpness=value_input.get('sharpness', 0.0),
percentile=None,
t_max=None,
t_min=None,
discard_penultimate_step=None,
intermediate_callback=intermediate_callback,
cat_uc=cat_uc,
cat_uc=value_input.get('cat_uc', cat_uc),
**kwargs)
self.dynamic_unload(self.diffusion_model,
@@ -0,0 +1,595 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import os.path
import random
from collections import OrderedDict
import torch
import torchvision.transforms.functional as TF
from PIL.Image import Image
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
from scepter.modules.model.network.diffusion.schedules import noise_schedule
from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, MODELS,
TOKENIZERS)
from scepter.modules.model.utils.data_utils import crop_back
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
def get_model(model_tuple):
assert 'model' in model_tuple
return model_tuple['model']
class LargenInference():
'''
define vae, unet, text-encoder, tuner, refiner components
support to load the components dynamicly.
create and load model when run this model at the first time.
'''
def __init__(self, logger=None):
self.logger = logger
self.loaded_model = {}
self.loaded_model_name = [
'diffusion_model', 'first_stage_model', 'cond_stage_model'
]
def init_from_cfg(self, cfg):
self.name = cfg.NAME
self.is_default = cfg.get('IS_DEFAULT', False)
module_paras = self.load_default(cfg.get('DEFAULT_PARAS', None))
assert cfg.have('MODEL')
cfg.MODEL = self.redefine_paras(cfg.MODEL)
self.diffusion = self.load_schedule(cfg.MODEL.SCHEDULE)
self.diffusion_model = self.infer_model(
cfg.MODEL.DIFFUSION_MODEL, module_paras.get(
'DIFFUSION_MODEL',
None)) if cfg.MODEL.have('DIFFUSION_MODEL') else None
self.first_stage_model = self.infer_model(
cfg.MODEL.FIRST_STAGE_MODEL,
module_paras.get(
'FIRST_STAGE_MODEL',
None)) if cfg.MODEL.have('FIRST_STAGE_MODEL') else None
self.cond_stage_model = self.infer_model(
cfg.MODEL.COND_STAGE_MODEL,
module_paras.get(
'COND_STAGE_MODEL',
None)) if cfg.MODEL.have('COND_STAGE_MODEL') else None
self.refiner_cond_model = self.infer_model(
cfg.MODEL.REFINER_COND_MODEL,
module_paras.get(
'REFINER_COND_MODEL',
None)) if cfg.MODEL.have('REFINER_COND_MODEL') else None
self.refiner_diffusion_model = self.infer_model(
cfg.MODEL.REFINER_MODEL, module_paras.get(
'REFINER_MODEL',
None)) if cfg.MODEL.have('REFINER_MODEL') else None
self.tokenizer = TOKENIZERS.build(
cfg.MODEL.TOKENIZER,
logger=self.logger) if cfg.MODEL.have('TOKENIZER') else None
if self.tokenizer is not None:
self.cond_stage_model['cfg'].KWARGS = {
'vocab_size': self.tokenizer.vocab_size
}
def redefine_paras(self, cfg):
if cfg.get('PRETRAINED_MODEL', None):
assert FS.isfile(cfg.PRETRAINED_MODEL)
with FS.get_from(cfg.PRETRAINED_MODEL,
wait_finish=True) as local_path:
if local_path.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(local_path)
else:
sd = torch.load(local_path, map_location='cpu')
if 'model' in sd:
sd = sd['model']
first_stage_model_path = os.path.join(
os.path.dirname(local_path), 'first_stage_model.pth')
cond_stage_model_path = os.path.join(
os.path.dirname(local_path), 'cond_stage_model.pth')
diffusion_model_path = os.path.join(
os.path.dirname(local_path), 'diffusion_model.pth')
if (not os.path.exists(first_stage_model_path)
or not os.path.exists(cond_stage_model_path)
or not os.path.exists(diffusion_model_path)):
self.logger.info(
'Now read the whole model and rearrange the modules, it may take several mins.'
)
first_stage_model = OrderedDict()
cond_stage_model = OrderedDict()
diffusion_model = OrderedDict()
for k, v in sd.items():
if k.startswith('first_stage_model.'):
first_stage_model[k.replace(
'first_stage_model.', '')] = v
elif k.startswith('conditioner.'):
cond_stage_model[k.replace('conditioner.', '')] = v
elif k.startswith('cond_stage_model.'):
if k.startswith('cond_stage_model.model.'):
cond_stage_model[k.replace(
'cond_stage_model.model.', '')] = v
else:
cond_stage_model[k.replace(
'cond_stage_model.', '')] = v
elif k.startswith('model.diffusion_model.'):
diffusion_model[k.replace('model.diffusion_model.',
'')] = v
elif k.startswith('model.'):
diffusion_model[k.replace('model.', '')] = v
else:
continue
if cfg.have('FIRST_STAGE_MODEL'):
with open(first_stage_model_path + 'cache', 'wb') as f:
torch.save(first_stage_model, f)
os.rename(first_stage_model_path + 'cache',
first_stage_model_path)
self.logger.info(
'First stage model has been processed.')
if cfg.have('COND_STAGE_MODEL'):
with open(cond_stage_model_path + 'cache', 'wb') as f:
torch.save(cond_stage_model, f)
os.rename(cond_stage_model_path + 'cache',
cond_stage_model_path)
self.logger.info(
'Cond stage model has been processed.')
if cfg.have('DIFFUSION_MODEL'):
with open(diffusion_model_path + 'cache', 'wb') as f:
torch.save(diffusion_model, f)
os.rename(diffusion_model_path + 'cache',
diffusion_model_path)
self.logger.info('Diffusion model has been processed.')
if not cfg.FIRST_STAGE_MODEL.get('PRETRAINED_MODEL', None):
cfg.FIRST_STAGE_MODEL.PRETRAINED_MODEL = first_stage_model_path
else:
cfg.FIRST_STAGE_MODEL.RELOAD_MODEL = first_stage_model_path
if not cfg.COND_STAGE_MODEL.get('PRETRAINED_MODEL', None):
cfg.COND_STAGE_MODEL.PRETRAINED_MODEL = cond_stage_model_path
else:
cfg.COND_STAGE_MODEL.RELOAD_MODEL = cond_stage_model_path
if not cfg.DIFFUSION_MODEL.get('PRETRAINED_MODEL', None):
cfg.DIFFUSION_MODEL.PRETRAINED_MODEL = diffusion_model_path
else:
cfg.DIFFUSION_MODEL.RELOAD_MODEL = diffusion_model_path
return cfg
def init_from_modules(self, modules):
for k, v in modules.items():
self.__setattr__(k, v)
def infer_model(self, cfg, module_paras=None):
module = {
'model': None,
'cfg': cfg,
'device': 'offline',
'name': cfg.NAME,
'function_info': {},
'paras': {}
}
if module_paras is None:
return module
function_info = {}
paras = {
k.lower(): v
for k, v in module_paras.get('PARAS', {}).items()
}
for function in module_paras.get('FUNCTION', []):
input_dict = {}
for inp in function.get('INPUT', []):
if inp.lower() in self.input:
input_dict[inp.lower()] = self.input[inp.lower()]
function_info[function.NAME] = {
'dtype': function.get('DTYPE', 'float32'),
'input': input_dict
}
module['paras'] = paras
module['function_info'] = function_info
return module
def init_from_ckpt(self, path, model, ignore_keys=list()):
if path.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(path)
else:
sd = torch.load(path, map_location='cpu')
new_sd = OrderedDict()
for k, v in sd.items():
ignored = False
for ik in ignore_keys:
if ik in k:
if we.rank == 0:
self.logger.info(
'Ignore key {} from state_dict.'.format(k))
ignored = True
break
if not ignored:
new_sd[k] = v
missing, unexpected = model.load_state_dict(new_sd, strict=False)
if we.rank == 0:
self.logger.info(
f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys'
)
if len(missing) > 0:
self.logger.info(f'Missing Keys:\n {missing}')
if len(unexpected) > 0:
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
def load(self, module):
if module['device'] == 'offline':
if module['cfg'].NAME in MODELS.class_map:
model = MODELS.build(module['cfg'], logger=self.logger).eval()
elif module['cfg'].NAME in BACKBONES.class_map:
model = BACKBONES.build(module['cfg'],
logger=self.logger).eval()
elif module['cfg'].NAME in EMBEDDERS.class_map:
model = EMBEDDERS.build(module['cfg'],
logger=self.logger).eval()
else:
raise NotImplementedError
if module['cfg'].get('RELOAD_MODEL', None):
self.init_from_ckpt(module['cfg'].RELOAD_MODEL, model)
module['model'] = model
module['device'] = 'cpu'
if module['device'] == 'cpu':
module['device'] = we.device_id
module['model'] = module['model'].to(we.device_id)
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:
self.paras = cfg.PARAS
self.input = {k.lower(): v for k, v in cfg.INPUT.items()}
self.output = {k.lower(): v for k, v in cfg.OUTPUT.items()}
module_paras = cfg.MODULES_PARAS
return module_paras
def load_schedule(self, cfg):
parameterization = cfg.get('PARAMETERIZATION', 'eps')
assert parameterization in [
'eps', 'x0', 'v'
], 'currently only supporting "eps" and "x0" and "v"'
num_timesteps = cfg.get('TIMESTEPS', 1000)
schedule_args = {
k.lower(): v
for k, v in cfg.get('SCHEDULE_ARGS', {
'NAME': 'logsnr_cosine_interp',
'SCALE_MIN': 2.0,
'SCALE_MAX': 4.0
}).items()
}
zero_terminal_snr = cfg.get('ZERO_TERMINAL_SNR', False)
if zero_terminal_snr:
assert parameterization == 'v', 'Now zero_terminal_snr only support v-prediction mode.'
sigmas = noise_schedule(schedule=schedule_args.pop('name'),
n=num_timesteps,
zero_terminal_snr=zero_terminal_snr,
**schedule_args)
diffusion = GaussianDiffusion(sigmas=sigmas,
prediction_type=parameterization)
return diffusion
def get_batch(self, value_dict, num_samples=1):
batch = {}
batch_uc = {}
N = num_samples
device = we.device_id
for key in value_dict:
if key == 'prompt':
if not self.tokenizer:
batch['prompt'] = value_dict['prompt']
batch_uc['prompt'] = value_dict['negative_prompt']
else:
batch['tokens'] = self.tokenizer(value_dict['prompt']).to(
we.device_id)
batch_uc['tokens'] = self.tokenizer(
value_dict['negative_prompt']).to(we.device_id)
elif key == 'original_size_as_tuple':
batch['original_size_as_tuple'] = (torch.tensor(
value_dict['original_size_as_tuple']).to(device).repeat(
N, 1))
elif key == 'crop_coords_top_left':
batch['crop_coords_top_left'] = (torch.tensor(
value_dict['crop_coords_top_left']).to(device).repeat(
N, 1))
elif key == 'aesthetic_score':
batch['aesthetic_score'] = (torch.tensor(
[value_dict['aesthetic_score']]).to(device).repeat(N, 1))
batch_uc['aesthetic_score'] = (torch.tensor([
value_dict['negative_aesthetic_score']
]).to(device).repeat(N, 1))
elif key == 'target_size_as_tuple':
batch['target_size_as_tuple'] = (torch.tensor(
value_dict['target_size_as_tuple']).to(device).repeat(
N, 1))
elif key == 'image':
batch[key] = self.load_image(value_dict[key], num_samples=N)
else:
batch[key] = value_dict[key]
for key in batch.keys():
if key not in batch_uc and isinstance(batch[key], torch.Tensor):
batch_uc[key] = torch.clone(batch[key])
return batch, batch_uc
def load_image(self, image, num_samples=1):
if isinstance(image, torch.Tensor):
pass
elif isinstance(image, Image):
pass
elif isinstance(image, Image):
pass
def get_function_info(self, module, function_name=None):
all_function = module['function_info']
if function_name in all_function:
return function_name, all_function[function_name]['dtype']
if function_name is None and len(all_function) == 1:
for k, v in all_function.items():
return k, v['dtype']
def encode_first_stage(self, x, **kwargs):
_, dtype = self.get_function_info(self.first_stage_model, 'encode')
with torch.autocast('cuda',
enabled=dtype == 'float16',
dtype=getattr(torch, dtype)):
z = get_model(self.first_stage_model).encode(x)
return self.first_stage_model['paras']['scale_factor'] * z
def decode_first_stage(self, z):
_, dtype = self.get_function_info(self.first_stage_model, 'encode')
with torch.autocast('cuda',
enabled=dtype == 'float16',
dtype=getattr(torch, dtype)):
z = 1. / self.first_stage_model['paras']['scale_factor'] * z
return get_model(self.first_stage_model).decode(z)
@torch.no_grad()
def __call__(self,
input,
num_samples=1,
intermediate_callback=None,
refine_strength=0,
img_to_img_strength=0,
cat_uc=True,
tuner_model=None,
control_model=None,
**kwargs):
value_input = copy.deepcopy(self.input)
value_input.update(input)
print(value_input)
height, width = value_input['target_size_as_tuple']
value_output = copy.deepcopy(self.output)
batch, batch_uc = self.get_batch(value_input, num_samples=1)
# first stage encode
task = kwargs.get('largen_task', 'Text_Guided_Inpainting')
image_scale = kwargs.get('largen_image_scale', 1.0)
tar_image = kwargs.get('largen_tar_image', None)
tar_mask = kwargs.get('largen_tar_mask', None)
masked_image = kwargs.get('largen_masked_image', None)
ref_image = kwargs.get('largen_ref_image', None)
ref_mask = kwargs.get('largen_ref_mask', None)
ref_clip = kwargs.get('largen_ref_clip', None)
base_image = kwargs.get('largen_base_image', None)
extra_sizes = kwargs.get('largen_extra_sizes', None)
bbox_yyxx = kwargs.get('largen_bbox_yyxx', None)
device = we.device_id
tar_image = tar_image.to(device)
tar_mask = tar_mask.to(device)
masked_image = masked_image.to(device)
if 'Subject' in task:
ref_image = ref_image.to(device)
ref_mask = ref_mask.to(device)
ref_clip = ref_clip.to(device)
self.dynamic_load(self.first_stage_model, 'first_stage_model')
tar_x0 = self.encode_first_stage(tar_image)
masked_x0 = self.encode_first_stage(masked_image)
b, _, h, w = tar_x0.shape
tar_mask_latent = TF.resize(tar_mask, (h, w), antialias=True)
tar_mask_latent = (tar_mask_latent > 0.5).float()
batch.update({
'tar_x0': tar_x0,
'tar_mask_latent': tar_mask_latent,
'masked_x0': masked_x0,
'task': task
})
batch_uc.update({
'tar_x0': tar_x0,
'tar_mask_latent': tar_mask_latent,
'masked_x0': masked_x0,
'task': task
})
if 'Subject' in task and ref_image is not None:
ref_x0 = self.encode_first_stage(ref_image)
batch.update({
'ref_ip': ref_clip,
'ref_detail': ref_clip,
'ref_x0': ref_x0,
'ref_mask': ref_mask,
'image_scale': image_scale,
})
batch_uc.update({
'ref_ip': torch.zeros_like(ref_clip),
'ref_detail': ref_clip,
'ref_x0': ref_x0,
'ref_mask': ref_mask,
'image_scale': image_scale,
})
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=True)
# cond stage
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
function_name, dtype = self.get_function_info(self.cond_stage_model)
with torch.autocast('cuda',
enabled=dtype == 'float16',
dtype=getattr(torch, dtype)):
context = getattr(get_model(self.cond_stage_model),
function_name)(batch)
null_context = getattr(get_model(self.cond_stage_model),
function_name)(batch_uc)
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=True)
# get noise
seed = kwargs.pop('seed', -1)
g = torch.Generator(device=we.device_id)
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
g.manual_seed(seed)
if 'seed' in value_output:
value_output['seed'] = seed
for sample_id in range(num_samples):
if self.diffusion_model is not None:
noise = torch.empty(
1,
4,
height // self.first_stage_model['paras']['size_factor'],
width // self.first_stage_model['paras']['size_factor'],
device=we.device_id).normal_(generator=g)
self.dynamic_load(self.diffusion_model, 'diffusion_model')
# UNet use input n_prompt
function_name, dtype = self.get_function_info(
self.diffusion_model)
with torch.autocast('cuda',
enabled=dtype == 'float16',
dtype=getattr(torch, dtype)):
latent = self.diffusion.sample(
noise=noise,
x=None,
denoising_strength=1.0,
refine_strength=refine_strength,
solver=value_input.get('sample', 'ddim'),
model=get_model(self.diffusion_model),
model_kwargs=[{
'cond': context
}, {
'cond': null_context
}],
steps=value_input.get('sample_steps', 50),
guide_scale=value_input.get('guide_scale', 7.5),
guide_rescale=value_input.get('guide_rescale', 0.5),
discretization=value_input.get('discretization',
'trailing'),
show_progress=True,
seed=seed,
condition_fn=None,
clamp=None,
percentile=None,
t_max=None,
t_min=None,
discard_penultimate_step=None,
intermediate_callback=intermediate_callback,
cat_uc=cat_uc,
**kwargs)
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=True)
if 'latent' in value_output:
if value_output['latent'] is None or (
isinstance(value_output['latent'], list)
and len(value_output['latent']) < 1):
value_output['latent'] = []
value_output['latent'].append(latent)
self.dynamic_load(self.first_stage_model, 'first_stage_model')
x_samples = self.decode_first_stage(latent).float()
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=True)
images = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0)
if base_image is not None:
stitch_images = []
for img in images:
stitch_img = crop_back(img, copy.deepcopy(base_image),
extra_sizes, bbox_yyxx)
stitch_images.append(stitch_img)
images = torch.stack(stitch_images, dim=0)
if 'images' in value_output:
if value_output['images'] is None or (
isinstance(value_output['images'], list)
and len(value_output['images']) < 1):
value_output['images'] = []
value_output['images'].append(images)
for k, v in value_output.items():
if isinstance(v, list):
value_output[k] = torch.cat(v, dim=0)
if isinstance(v, torch.Tensor):
value_output[k] = v.cpu()
return value_output
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.backbone.autoencoder.ae_module import (Decoder,
Encoder)
Encoder,
RDecoder)
@@ -3,6 +3,7 @@
import torch
import torch.nn as nn
from einops import repeat
from torch.utils.checkpoint import checkpoint
from scepter.modules.model.backbone.autoencoder.ae_utils import (
@@ -244,6 +245,7 @@ class Decoder(BaseModel):
# compute in_ch_mult, block_in and curr_res at lowest res
block_in = self.ch * self.ch_mult[self.num_resolutions - 1]
self.block_in = block_in
curr_res = 1
# z to block_in
self.conv_in = torch.nn.Conv2d(self.z_channels,
@@ -340,3 +342,48 @@ class Decoder(BaseModel):
__class__.__name__,
Decoder.para_dict,
set_name=True)
@BACKBONES.register_class()
class RDecoder(Decoder):
def construct_model(self):
super().construct_model()
self.resize_level = nn.Sequential(
nn.Linear(self.block_in, self.block_in),
nn.SiLU(),
nn.Linear(self.block_in, self.block_in),
)
def forward(self, z, rembed=None):
# timestep embedding
temb = None
h = self.conv_in(z)
bs, channel, hdim, wdim = h.size()
if rembed is not None:
rembed = self.resize_level(rembed)
rembed = repeat(rembed, 'b e-> b e hd wd', hd=hdim, wd=wdim)
h = h + rembed
# middle
if not self.use_checkpoint:
h = self.mid_upsclae_transform(h, temb)
else:
h = checkpoint(self.mid_upsclae_transform, h, temb)
# end
if self.give_pre_end:
return h
h = self.norm_out(h)
h = nonlinearity(h)
h = self.conv_out(h)
if self.tanh_out:
h = torch.tanh(h)
return h
@staticmethod
def get_config_template():
return dict_to_yaml('BACKBONE',
__class__.__name__,
Decoder.para_dict,
set_name=True)
@@ -1,3 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.backbone.unet.unet_module import DiffusionUNet
from scepter.modules.model.backbone.unet.unet_module import (DiffusionUNet,
DiffusionUNetXL,
LargenUNetXL)
@@ -8,8 +8,9 @@ import torch
import torch.nn as nn
from scepter.modules.model.backbone.unet.unet_utils import (
Downsample, ResBlock, SpatialTransformer, Timestep,
TimestepEmbedSequential, Upsample, conv_nd, linear, normalization,
BasicTransformerBlock, Downsample, ResBlock, SpatialTransformer,
SpatialTransformerV2, Timestep, TimestepEmbedSequential,
TransformerBlockV2, Upsample, conv_nd, linear, normalization,
timestep_embedding, zero_module)
from scepter.modules.model.base_model import BaseModel
from scepter.modules.model.registry import BACKBONES
@@ -951,3 +952,443 @@ class DiffusionUNetXL(DiffusionUNet):
__class__.__name__,
DiffusionUNetXL.para_dict,
set_name=True)
@BACKBONES.register_class()
class LargenUNetXL(DiffusionUNetXL):
para_dict = {
'TRANSFORMER_BLOCK_TYPE': {
'value': 'att_v1'
},
'IMAGE_SCALE': {
'value': 0.0,
},
}
para_dict.update(DiffusionUNetXL.para_dict)
def __init__(self, cfg, logger):
super().__init__(cfg, logger=logger)
self.init_params(cfg)
self.construct_network()
def init_params(self, cfg):
super().init_params(cfg)
self.transformer_block_type = cfg.get('TRANSFORMER_BLOCK_TYPE',
'att_v1')
TRANSFORMER_BLOCKS = {
'att_v1': BasicTransformerBlock,
'att_v2': TransformerBlockV2,
}
assert self.transformer_block_type in list(TRANSFORMER_BLOCKS.keys())
self.transformer_block = TRANSFORMER_BLOCKS[
self.transformer_block_type]
self.image_scale = cfg.get('IMAGE_SCALE', 0.0)
self.use_refine = cfg.get('USE_REFINE', False)
def construct_network(self):
in_channels = self.in_channels
model_channels = self.model_channels
out_channels = self.out_channels
attention_resolutions = self.attention_resolutions
channel_mult = self.channel_mult
num_classes = self.num_classes
num_heads = self.num_heads
num_head_channels = self.num_head_channels
dims = self.dims
dropout = self.dropout
use_checkpoint = self.use_checkpoint
use_scale_shift_norm = self.use_scale_shift_norm
disable_self_attentions = self.disable_self_attentions
disable_middle_self_attn = self.disable_middle_self_attn
transformer_depth = self.transformer_depth
transformer_depth_middle = self.transformer_depth_middle
context_dim = self.context_dim
use_linear_in_transformer = self.use_linear_in_transformer
resblock_updown = self.resblock_updown
conv_resample = self.conv_resample
adm_in_channels = self.adm_in_channels
transformer_block = self.transformer_block
time_embed_dim = model_channels * 4
self.time_embed = nn.Sequential(
linear(model_channels, time_embed_dim),
nn.SiLU(),
linear(time_embed_dim, time_embed_dim),
)
if self.num_classes is not None:
if isinstance(self.num_classes, int):
self.label_emb = nn.Embedding(num_classes, time_embed_dim)
elif self.num_classes == 'continuous':
print('setting up linear c_adm embedding layer')
self.label_emb = nn.Linear(1, time_embed_dim)
elif self.num_classes == 'timestep':
self.label_emb = nn.Sequential(
Timestep(model_channels),
nn.Sequential(
linear(model_channels, time_embed_dim),
nn.SiLU(),
linear(time_embed_dim, time_embed_dim),
),
)
elif self.num_classes == 'sequential':
assert adm_in_channels is not None
self.label_emb = nn.Sequential(
nn.Sequential(
linear(adm_in_channels, time_embed_dim),
nn.SiLU(),
linear(time_embed_dim, time_embed_dim),
))
else:
raise ValueError()
self.input_blocks = nn.ModuleList([
TimestepEmbedSequential(
conv_nd(dims, in_channels, model_channels, 3, padding=1))
])
self._feature_size = model_channels
input_block_chans = [model_channels]
input_down_flag = [False]
ch = model_channels
ds = 1
for level, mult in enumerate(channel_mult):
for nr in range(self.num_res_blocks[level]):
layers = [
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=mult * model_channels,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
)
]
ch = mult * model_channels
if ds in attention_resolutions:
if num_head_channels == -1:
dim_head = ch // num_heads
else:
num_heads = ch // num_head_channels
dim_head = num_head_channels
disabled_sa = disable_self_attentions[level] if exists(
disable_self_attentions) else False
layers.append(
SpatialTransformerV2(
ch,
num_heads,
dim_head,
transformer_block=transformer_block,
depth=transformer_depth[level],
context_dim=context_dim,
disable_self_attn=disabled_sa,
use_linear=use_linear_in_transformer,
use_checkpoint=use_checkpoint))
self.input_blocks.append(TimestepEmbedSequential(*layers))
self._feature_size += ch
input_block_chans.append(ch)
input_down_flag.append(False)
if level != len(channel_mult) - 1:
out_ch = ch
self.input_blocks.append(
TimestepEmbedSequential(
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=out_ch,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
down=True,
) if resblock_updown else Downsample(
ch, conv_resample, dims=dims, out_channels=out_ch))
)
ch = out_ch
input_block_chans.append(ch)
input_down_flag.append(True)
ds *= 2
self._feature_size += ch
self._input_block_chans = copy.deepcopy(input_block_chans)
self._input_down_flag = input_down_flag
if num_head_channels == -1:
dim_head = ch // num_heads
else:
num_heads = ch // num_head_channels
dim_head = num_head_channels
self.middle_block = TimestepEmbedSequential(
ResBlock(
ch,
time_embed_dim,
dropout,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
),
SpatialTransformerV2(ch,
num_heads,
dim_head,
transformer_block=transformer_block,
depth=transformer_depth_middle,
context_dim=context_dim,
disable_self_attn=disable_middle_self_attn,
use_linear=use_linear_in_transformer,
use_checkpoint=use_checkpoint),
ResBlock(
ch,
time_embed_dim,
dropout,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
),
)
self._feature_size += ch
self._middle_block_chans = [ch]
self._output_block_chans = []
self.output_blocks = nn.ModuleList([])
for level, mult in list(enumerate(channel_mult))[::-1]:
for i in range(self.num_res_blocks[level] + 1):
ich = input_block_chans.pop()
layers = [
ResBlock(
ch + ich,
time_embed_dim,
dropout,
out_channels=model_channels * mult,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
)
]
ch = model_channels * mult
if ds in attention_resolutions:
if num_head_channels == -1:
dim_head = ch // num_heads
else:
num_heads = ch // num_head_channels
dim_head = num_head_channels
disabled_sa = disable_self_attentions[level] if exists(
disable_self_attentions) else False
layers.append(
SpatialTransformerV2(
ch,
num_heads,
dim_head,
transformer_block=transformer_block,
depth=transformer_depth[level],
context_dim=context_dim,
disable_self_attn=disabled_sa,
use_linear=use_linear_in_transformer,
use_checkpoint=use_checkpoint))
if level and i == self.num_res_blocks[level]:
out_ch = ch
layers.append(
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=out_ch,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
up=True,
) if resblock_updown else Upsample(
ch, conv_resample, dims=dims, out_channels=out_ch))
ds //= 2
self.output_blocks.append(TimestepEmbedSequential(*layers))
self._feature_size += ch
self._output_block_chans.append(ch)
self.out = nn.Sequential(
normalization(ch),
nn.SiLU(),
zero_module(
conv_nd(dims, model_channels, out_channels, 3, padding=1)),
)
if self.use_refine:
self.ref_time_embed = copy.deepcopy(self.time_embed)
self.ref_label_emb = copy.deepcopy(self.label_emb)
self.ref_input_blocks = copy.deepcopy(self.input_blocks)
self.ref_input_blocks[0] = TimestepEmbedSequential(
conv_nd(dims, 4, model_channels, 3, padding=1))
self.ref_middle_block = copy.deepcopy(self.middle_block)
self.ref_output_blocks = copy.deepcopy(self.output_blocks)
def load_pretrained_model(self, pretrained_model):
if pretrained_model is not None:
with FS.get_from(pretrained_model,
wait_finish=True) as local_model:
self.init_from_ckpt(local_model, ignore_keys=self.ignore_keys)
def init_from_ckpt(self, path, ignore_keys=list()):
if path.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(path)
else:
sd = torch.load(path, map_location='cpu')
new_sd = OrderedDict()
for k, v in sd.items():
ignored = False
for ik in ignore_keys:
if ik in k:
if we.rank == 0:
self.logger.info(
'Ignore key {} from state_dict.'.format(k))
ignored = True
break
if not ignored:
if k == 'input_blocks.0.0.weight':
if we.rank == 0:
self.logger.info(
'Partial initial key {} from state_dict.'.format(
k))
new_v = torch.empty(320, self.in_channels, 3, 3)
nn.init.zeros_(new_v)
new_v[:, :v.shape[1]] = v
new_sd[k] = new_v
if self.use_refine:
new_sd['ref_' + k] = v
else:
new_sd[k] = v
if self.use_refine:
new_sd['ref_' + k] = v
missing, unexpected = self.load_state_dict(new_sd, strict=False)
if we.rank == 0:
self.logger.info(
f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys'
)
if len(missing) > 0:
self.logger.info(f'Missing Keys:\n {missing}')
if len(unexpected) > 0:
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
def forward(self, x, t=None, cond=dict(), **kwargs):
t_emb = timestep_embedding(t,
self.model_channels,
repeat_only=False,
legacy=True)
emb = self.time_embed(t_emb)
if isinstance(cond, dict):
if 'y' in cond:
assert self.num_classes is not None
emb = emb + self.label_emb(cond['y'])
if self.use_refine:
ref_emb = self.ref_time_embed(t_emb)
assert 'null_y' in cond
cond_y = cond['y'].clone()
cond_y[:, :cond['null_y'].shape[1]] = cond['null_y']
ref_emb = ref_emb + self.ref_label_emb(cond_y)
if 'concat' in cond:
c = cond['concat']
x = torch.cat([x, c], dim=1)
context = cond.get('crossattn', None)
img_context = cond.get('img_crossattn', None)
task = cond['task']
image_scale = cond.get('image_scale', self.image_scale)
if 'Subject' in task and img_context is not None:
ip_enc_scale = image_scale
ip_dec_scale = image_scale
num_img_tokens = img_context.shape[1]
context = torch.cat([context, img_context], dim=1)
else:
ip_enc_scale = None
ip_dec_scale = None
num_img_tokens = None
ref = cond.get('ref_xt', None)
ref_context = cond.get('ref_crossattn', None)
else:
raise TypeError
hs = []
refs = []
h = x
if self.use_refine:
assert ref is not None and ref_context is not None
for i, (ref_module, module) in enumerate(
zip(self.ref_input_blocks, self.input_blocks)):
ref = ref_module(ref, ref_emb, ref_context, caching=None)
h = module(h,
emb,
context,
caching=None,
scale=ip_enc_scale,
num_img_token=num_img_tokens)
refs.append(ref)
hs.append(h)
ref = self.ref_middle_block(ref,
ref_emb,
ref_context,
caching=None)
h = self.middle_block(h,
emb,
context,
caching=None,
scale=ip_enc_scale,
num_img_token=num_img_tokens)
for i, (ref_module, module) in enumerate(
zip(self.ref_output_blocks, self.output_blocks)):
cache = []
ref = torch.cat([ref, refs.pop()], dim=1)
ref = ref_module(ref,
ref_emb,
ref_context,
caching='write',
cache=cache)
h = torch.cat([h, hs.pop()], dim=1)
h = module(h,
emb,
context,
caching='read',
cache=cache,
scale=ip_dec_scale,
num_img_token=num_img_tokens)
else:
for module in self.input_blocks:
h = module(h,
emb,
context,
caching=None,
scale=ip_enc_scale,
num_img_token=num_img_tokens)
hs.append(h)
h = self.middle_block(h,
emb,
context,
caching=None,
scale=ip_enc_scale,
num_img_token=num_img_tokens)
for module in self.output_blocks:
h = torch.cat([h, hs.pop()], dim=1)
h = module(h,
emb,
context,
caching=None,
scale=ip_dec_scale,
num_img_token=num_img_tokens)
out = self.out(h)
return out
@staticmethod
def get_config_template():
return dict_to_yaml('BACKBONE',
__class__.__name__,
LargenUNetXL.para_dict,
set_name=True)
@@ -10,6 +10,7 @@ import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision.transforms.functional as TF
from einops import rearrange, repeat
from packaging import version
@@ -171,12 +172,14 @@ class TimestepEmbedSequential(nn.Sequential, TimestepBlock):
A sequential module that passes timestep embeddings to the children that
support it as an extra input.
"""
def forward(self, x, emb, context=None, target_size=None):
def forward(self, x, emb, context=None, target_size=None, **kwargs):
for layer in self:
if isinstance(layer, TimestepBlock):
x = layer(x, emb)
elif isinstance(layer, SpatialTransformer):
x = layer(x, context)
elif isinstance(layer, SpatialTransformerV2):
x = layer(x, context, **kwargs)
elif isinstance(layer, Upsample):
x = layer(x, target_size)
else:
@@ -864,6 +867,92 @@ class MemoryEfficientCrossAttention(nn.Module):
return self.to_out(out)
class XFormersMHA_IP(nn.Module):
def __init__(self,
query_dim,
context_dim=None,
heads=8,
dim_head=64,
dropout=0.0):
super().__init__()
inner_dim = dim_head * heads
context_dim = default(context_dim, query_dim)
self.heads = heads
self.dim_head = dim_head
self.to_q = nn.Linear(query_dim, inner_dim, bias=False)
self.to_k = nn.Linear(context_dim, inner_dim, bias=False)
self.to_v = nn.Linear(context_dim, inner_dim, bias=False)
self.to_k_ip = nn.Linear(context_dim, inner_dim, bias=False)
self.to_v_ip = nn.Linear(context_dim, inner_dim, bias=False)
self.to_out = nn.Sequential(nn.Linear(inner_dim, query_dim),
nn.Dropout(dropout))
self.attention_op = None
def forward(self,
x,
context=None,
mask=None,
scale=None,
num_img_token=None):
q = self.to_q(x)
context = default(context, x)
if scale is not None and num_img_token is not None:
eos = context.shape[1] - num_img_token
txt_context = context[:, :eos, :]
img_context = context[:, eos:, :]
k = self.to_k(txt_context)
v = self.to_v(txt_context)
k_i = self.to_k_ip(img_context)
v_i = self.to_v_ip(img_context)
b, _, _ = q.shape
q, k, v, k_i, v_i = map(
lambda t: t.unsqueeze(3).reshape(b, t.shape[
1], self.heads, self.dim_head).permute(0, 2, 1, 3).reshape(
b * self.heads, t.shape[1], self.dim_head).contiguous(
),
(q, k, v, k_i, v_i),
)
# actually compute the attention, what we cannot get enough of
txt_out = xformers.ops.memory_efficient_attention(
q, k, v, attn_bias=None, op=self.attention_op)
img_out = xformers.ops.memory_efficient_attention(
q, k_i, v_i, attn_bias=None, op=self.attention_op)
out = txt_out + scale * img_out
else:
k = self.to_k(context)
v = self.to_v(context)
b, _, _ = q.shape
q, k, v = map(
lambda t: t.unsqueeze(3).reshape(b, t.shape[
1], self.heads, self.dim_head).permute(0, 2, 1, 3).reshape(
b * self.heads, t.shape[1], self.dim_head).contiguous(
),
(q, k, v),
)
out = xformers.ops.memory_efficient_attention(q,
k,
v,
attn_bias=None,
op=self.attention_op)
# TODO: Use this directly in the attention operation, as a bias
if exists(mask):
raise NotImplementedError
out = (out.unsqueeze(0).reshape(
b, self.heads, out.shape[1],
self.dim_head).permute(0, 2, 1,
3).reshape(b, out.shape[1],
self.heads * self.dim_head))
return self.to_out(out)
class BasicTransformerBlock(nn.Module):
def __init__(self,
dim,
@@ -908,6 +997,65 @@ class BasicTransformerBlock(nn.Module):
return x
class TransformerBlockV2(nn.Module):
def __init__(self,
query_dim,
n_heads,
d_head,
dropout=0.,
context_dim=None,
gated_ff=True,
use_checkpoint=False,
disable_self_attn=False):
super().__init__()
self.disable_self_attn = disable_self_attn
self.attn1 = MemoryEfficientCrossAttention(query_dim=query_dim,
heads=n_heads,
dim_head=d_head,
dropout=dropout,
context_dim=None)
self.ff = FeedForward(query_dim, dropout=dropout, glu=gated_ff)
self.attn2 = XFormersMHA_IP(query_dim=query_dim,
heads=n_heads,
dim_head=d_head,
context_dim=context_dim)
self.norm1 = nn.LayerNorm(query_dim)
self.norm2 = nn.LayerNorm(query_dim)
self.norm3 = nn.LayerNorm(query_dim)
self.use_checkpoint = use_checkpoint
def forward(self,
x,
context,
caching=None,
cache=None,
scale=None,
num_img_token=None,
**kwargs):
y = self.norm1(x)
if caching == 'write':
assert isinstance(cache, list)
cache.append(y)
x = self.attn1(y, context=None) + x
elif caching == 'read':
assert isinstance(cache, list) and len(cache) > 0
c = cache.pop(0)
self_ctx = torch.cat([y, c], dim=1)
x = self.attn1(y, context=self_ctx) + x
elif caching is None:
x = self.attn1(y, context=None) + x
else:
assert False
x = self.attn2(self.norm2(x),
context=context,
scale=scale,
num_img_token=num_img_token) + x
x = self.ff(self.norm3(x)) + x
return x
class SpatialTransformer(nn.Module):
"""
Transformer block for image-like data.
@@ -1003,3 +1151,108 @@ class SpatialTransformer(nn.Module):
if not self.use_linear:
x = self.proj_out(x)
return x + x_in
class SpatialTransformerV2(nn.Module):
"""
Transformer block for image-like data.
First, project the input (aka embedding)
and reshape to b, t, d.
Then apply standard transformer action.
Finally, reshape to image
NEW: use_linear for more efficiency instead of the 1x1 convs
"""
def __init__(self,
in_channels,
n_heads,
d_head,
transformer_block,
depth=1,
dropout=0.,
context_dim=None,
disable_self_attn=False,
use_linear=False,
use_checkpoint=True):
super().__init__()
if exists(context_dim) and not isinstance(context_dim, list):
context_dim = [context_dim]
if exists(context_dim) and not isinstance(context_dim, (list)):
context_dim = [context_dim]
if exists(context_dim) and isinstance(context_dim, list):
if depth != len(context_dim):
print(
f'WARNING: {self.__class__.__name__}: Found context dims {context_dim} of'
f" depth {len(context_dim)}, which does not match the specified 'depth' of"
f' {depth}. Setting context_dim to {depth * [context_dim[0]]} now.'
)
# depth does not match context dims.
assert all(
map(lambda x: x == context_dim[0], context_dim)
), 'need homogenous context_dim to match depth automatically'
context_dim = depth * [context_dim[0]]
elif context_dim is None:
context_dim = [None] * depth
self.in_channels = in_channels
inner_dim = n_heads * d_head
self.norm = normalization(in_channels)
if not use_linear:
self.proj_in = nn.Conv2d(in_channels,
inner_dim,
kernel_size=1,
stride=1,
padding=0)
else:
self.proj_in = nn.Linear(in_channels, inner_dim)
self.transformer_blocks = nn.ModuleList([
transformer_block(inner_dim,
n_heads,
d_head,
dropout=dropout,
context_dim=context_dim[d],
disable_self_attn=disable_self_attn,
use_checkpoint=use_checkpoint)
for d in range(depth)
])
if not use_linear:
self.proj_out = zero_module(
nn.Conv2d(inner_dim,
in_channels,
kernel_size=1,
stride=1,
padding=0))
else:
self.proj_out = zero_module(nn.Linear(in_channels, inner_dim))
self.use_linear = use_linear
def forward(self, x, context=None, **kwargs):
# note: if no context is given, cross-attention defaults to self-attention
if not isinstance(context, list):
context = [context]
b, c, h, w = x.shape
ref_mask = kwargs.pop('ref_mask', None)
if ref_mask is not None:
ref_mask = TF.resize(ref_mask, (h, w), antialias=True)
ref_mask = (ref_mask > 0.5).float()
ref_mask = rearrange(ref_mask, 'b c h w -> b (h w) c').contiguous()
x_in = x
x = self.norm(x)
if not self.use_linear:
x = self.proj_in(x)
x = rearrange(x, 'b c h w -> b (h w) c').contiguous()
if self.use_linear:
x = self.proj_in(x)
for i, block in enumerate(self.transformer_blocks):
if i > 0 and len(context) == 1:
i = 0 # use same context for each block
x = block(x, context=context[i], ref_mask=ref_mask, **kwargs)
if self.use_linear:
x = self.proj_out(x)
x = rearrange(x, 'b (h w) c -> b c h w', h=h, w=w).contiguous()
if not self.use_linear:
x = self.proj_out(x)
return x + x_in
@@ -1,17 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import math
import torch
import torch.nn as nn
import torch.nn.functional
from einops import rearrange
from scepter.modules.model.backbone.video.init_helper import (
_init_transformer_weights, trunc_normal_)
from scepter.modules.model.registry import BACKBONES, BRICKS, STEMS
from scepter.modules.utils.config import dict_to_yaml
'''
The implementations of vivit as https://arxiv.org/abs/2103.15691.
The following setting alined the proposed model in the paper above.
@@ -39,6 +27,18 @@ TimesFormer:
complexity: (n_h * n_w) ** 2 + O(attn_temp)
'''
import math
import torch
import torch.nn as nn
import torch.nn.functional
from einops import rearrange
from scepter.modules.model.backbone.video.init_helper import (
_init_transformer_weights, trunc_normal_)
from scepter.modules.model.registry import BACKBONES, BRICKS, STEMS
from scepter.modules.utils.config import dict_to_yaml
@BACKBONES.register_class()
class VideoTransformer(nn.Module):
+4 -1
View File
@@ -1,7 +1,10 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.embedder.embedder import (ConcatTimestepEmbedderND,
FrozenCLIPEmbedder,
FrozenOpenCLIPEmbedder,
FrozenOpenCLIPEmbedder2,
GeneralConditioner)
GeneralConditioner,
IPAdapterPlusEmbedder,
RefCrossEmbedder)
File diff suppressed because it is too large Load Diff
+137 -31
View File
@@ -22,9 +22,10 @@ from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from .base_embedder import BaseEmbedder
from .resampler import Resampler
try:
from transformers import CLIPTextModel, CLIPTokenizer
from transformers import CLIPTextModel, CLIPTokenizer, CLIPVisionModelWithProjection
except Exception as e:
warnings.warn(
f'Import transformers error, please deal with this problem: {e}')
@@ -513,6 +514,93 @@ class ConcatTimestepEmbedderND(BaseEmbedder):
set_name=True)
@EMBEDDERS.register_class()
class IPAdapterPlusEmbedder(BaseEmbedder):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
with FS.get_dir_to_local_dir(cfg.CLIP_DIR,
wait_finish=True) as local_path:
self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(
local_path)
self.image_proj_model = Resampler(
dim=self.cfg.get('IN_DIM', 768),
depth=self.cfg.get('DEPTH', 4),
dim_head=64,
heads=self.cfg.get('HEADS', 12),
num_queries=self.cfg.get('NUM_TOKENS', 16),
embedding_dim=self.image_encoder.config.hidden_size,
output_dim=self.cfg.get('CROSSATTN_DIM', 768),
ff_mult=4,
)
with FS.get_from(cfg.PRETRAINED_MODEL, wait_finish=True) as local_path:
ckpt = torch.load(local_path, map_location='cpu')
self.image_proj_model.load_state_dict(ckpt['image_proj'],
strict=True)
self.patch_projector = nn.Linear(self.image_encoder.config.hidden_size,
self.cfg.get('CROSSATTN_DIM', 768))
def encode(self, ref_ip, ref_detail):
encoder_output = self.image_encoder(ref_ip, output_hidden_states=True)
image_prompt_embeds = self.image_proj_model(
encoder_output.hidden_states[-2])
encoder_output_2 = self.image_encoder(ref_detail,
output_hidden_states=True)
image_patch_embeds = self.patch_projector(
encoder_output_2.last_hidden_state)
out = {
'img_crossattn': image_prompt_embeds,
'ref_crossattn': image_patch_embeds,
}
return out
def forward(self, ref_ip, ref_detail):
return self.encode(ref_ip, ref_detail)
class RefCrossEmbedder(BaseEmbedder):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
with FS.get_dir_to_local_dir(cfg.CLIP_DIR,
wait_finish=True) as local_path:
self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(
local_path)
self.patch_projector = nn.Linear(self.image_encoder.config.hidden_size,
self.cfg.get('CROSSATTN_DIM', 768))
def encode(self, img):
encoder_output = self.image_encoder(img, output_hidden_states=True)
image_patch_embeds = self.patch_projector(
encoder_output.last_hidden_state)
out = {
'ref_crossattn': image_patch_embeds,
}
return out
def forward(self, img):
return self.encode(img)
@EMBEDDERS.register_class()
class TransparentEmbedder(BaseEmbedder):
def forward(self, *args):
out = dict()
for key, val in zip(self.input_keys, args):
out[key] = val
return out
@EMBEDDERS.register_class()
class NoiseConcatEmbedder(BaseEmbedder):
def forward(self, *args):
return {'concat': torch.cat(args, dim=1)}
@EMBEDDERS.register_class()
class GeneralConditioner(BaseEmbedder):
OUTPUT_DIM2KEYS = {2: 'y', 3: 'crossattn', 4: 'concat', 5: 'concat'}
@@ -598,42 +686,60 @@ class GeneralConditioner(BaseEmbedder):
with embedding_context():
if hasattr(embedder, 'input_key') and (embedder.input_key
is not None):
if embedder.input_key not in batch:
continue
if embedder.legacy_ucg_val is not None:
batch = self.possibly_get_ucg_val(embedder, batch)
emb_out = embedder(batch[embedder.input_key])
elif hasattr(embedder, 'input_keys'):
if any([k not in batch for k in embedder.input_keys]):
continue
emb_out = embedder(
*[batch[k] for k in embedder.input_keys])
assert isinstance(
emb_out, (torch.Tensor, list, tuple)
), f'encoder outputs must be tensors or a sequence, but got {type(emb_out)}'
if not isinstance(emb_out, (list, tuple)):
emb_out = [emb_out]
for emb in emb_out:
# print("emb.shape", emb.shape)
# print("emb.input_keys", embedder.input_keys)
out_key = self.OUTPUT_DIM2KEYS[emb.dim()]
if embedder.ucg_rate > 0.0 and embedder.legacy_ucg_val is None:
emb = (expand_dims_like(
torch.bernoulli(
(1.0 - embedder.ucg_rate) *
torch.ones(emb.shape[0], device=emb.device)),
emb,
) * emb)
if (hasattr(embedder, 'input_keys')):
if np.sum(
np.array([
key in force_zero_embeddings
for key in embedder.input_keys
])) > 0:
emb = torch.zeros_like(emb)
if out_key in output:
output[out_key] = torch.cat((output[out_key], emb),
self.KEY2CATDIM[out_key])
else:
output[out_key] = emb
# if "y" in output:
# print("out.shape", output["y"].shape)
if isinstance(emb_out, dict):
for key, val in emb_out.items():
if key in output:
# 重复出现的key必须在(y, crossattn, concat)中,否则raise error
assert key in self.KEY2CATDIM
output[key] = torch.cat([output[key], val],
dim=self.KEY2CATDIM[key])
else:
output[key] = val
else:
assert isinstance(
emb_out, (torch.Tensor, list, tuple)
), f'encoder outputs must be tensors or a sequence, but got {type(emb_out)}'
if not isinstance(emb_out, (list, tuple)):
emb_out = [emb_out]
for emb in emb_out:
# 根据emb的维度,判断该cond归属于 (y, concat, crossattn)中的哪一种
out_key = self.OUTPUT_DIM2KEYS[emb.dim()]
if embedder.ucg_rate > 0.0 and embedder.legacy_ucg_val is None:
emb = (expand_dims_like(
torch.bernoulli(
(1.0 - embedder.ucg_rate) *
torch.ones(emb.shape[0], device=emb.device)),
emb,
) * emb)
if (hasattr(embedder, 'input_keys')):
if np.sum(
np.array([
key in force_zero_embeddings
for key in embedder.input_keys
])) > 0:
emb = torch.zeros_like(emb)
if out_key in output:
output[out_key] = torch.cat((output[out_key], emb),
self.KEY2CATDIM[out_key])
else:
output[out_key] = emb
return output
def get_unconditional_conditioning(self,
+160
View File
@@ -0,0 +1,160 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import math
import torch
import torch.nn as nn
from einops import rearrange
from einops.layers.torch import Rearrange
# FFN
def FeedForward(dim, mult=4):
inner_dim = int(dim * mult)
return nn.Sequential(
nn.LayerNorm(dim),
nn.Linear(dim, inner_dim, bias=False),
nn.GELU(),
nn.Linear(inner_dim, dim, bias=False),
)
def reshape_tensor(x, heads):
bs, length, width = x.shape
# (bs, length, width) --> (bs, length, n_heads, dim_per_head)
x = x.view(bs, length, heads, -1)
# (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head)
x = x.transpose(1, 2)
# (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head)
x = x.reshape(bs, heads, length, -1)
return x
class PerceiverAttention(nn.Module):
def __init__(self, *, dim, dim_head=64, heads=8):
super().__init__()
self.scale = dim_head**-0.5
self.dim_head = dim_head
self.heads = heads
inner_dim = dim_head * heads
self.norm1 = nn.LayerNorm(dim)
self.norm2 = nn.LayerNorm(dim)
self.to_q = nn.Linear(dim, inner_dim, bias=False)
self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False)
self.to_out = nn.Linear(inner_dim, dim, bias=False)
def forward(self, x, latents):
"""
Args:
x (torch.Tensor): image features
shape (b, n1, D)
latent (torch.Tensor): latent features
shape (b, n2, D)
"""
x = self.norm1(x)
latents = self.norm2(latents)
b, l, _ = latents.shape
q = self.to_q(latents)
kv_input = torch.cat((x, latents), dim=-2)
k, v = self.to_kv(kv_input).chunk(2, dim=-1)
q = reshape_tensor(q, self.heads)
k = reshape_tensor(k, self.heads)
v = reshape_tensor(v, self.heads)
# attention
scale = 1 / math.sqrt(math.sqrt(self.dim_head))
weight = (q * scale) @ (k * scale).transpose(
-2, -1) # More stable with f16 than dividing afterwards
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
out = weight @ v
out = out.permute(0, 2, 1, 3).reshape(b, l, -1)
return self.to_out(out)
class Resampler(nn.Module):
def __init__(
self,
dim=1024,
depth=8,
dim_head=64,
heads=16,
num_queries=8,
embedding_dim=768,
output_dim=1024,
ff_mult=4,
max_seq_len: int = 257, # CLIP tokens + CLS token
apply_pos_emb: bool = False,
num_latents_mean_pooled:
int = 0, # number of latents derived from mean pooled representation of the sequence
):
super().__init__()
self.pos_emb = nn.Embedding(max_seq_len,
embedding_dim) if apply_pos_emb else None
self.latents = nn.Parameter(
torch.randn(1, num_queries, dim) / dim**0.5)
self.proj_in = nn.Linear(embedding_dim, dim)
self.proj_out = nn.Linear(dim, output_dim)
self.norm_out = nn.LayerNorm(output_dim)
self.to_latents_from_mean_pooled_seq = (nn.Sequential(
nn.LayerNorm(dim),
nn.Linear(dim, dim * num_latents_mean_pooled),
Rearrange('b (n d) -> b n d', n=num_latents_mean_pooled),
) if num_latents_mean_pooled > 0 else None)
self.layers = nn.ModuleList([])
for _ in range(depth):
self.layers.append(
nn.ModuleList([
PerceiverAttention(dim=dim, dim_head=dim_head,
heads=heads),
FeedForward(dim=dim, mult=ff_mult),
]))
def forward(self, x):
if self.pos_emb is not None:
n, device = x.shape[1], x.device
pos_emb = self.pos_emb(torch.arange(n, device=device))
x = x + pos_emb
latents = self.latents.repeat(x.size(0), 1, 1)
x = self.proj_in(x)
if self.to_latents_from_mean_pooled_seq:
meanpooled_seq = masked_mean(x,
dim=1,
mask=torch.ones(x.shape[:2],
device=x.device,
dtype=torch.bool))
meanpooled_latents = self.to_latents_from_mean_pooled_seq(
meanpooled_seq)
latents = torch.cat((meanpooled_latents, latents), dim=-2)
for attn, ff in self.layers:
latents = attn(x, latents) + latents
latents = ff(latents) + latents
latents = self.proj_out(latents)
return self.norm_out(latents)
def masked_mean(t, *, dim, mask=None):
if mask is None:
return t.mean(dim=dim)
denom = mask.sum(dim=dim, keepdim=True)
mask = rearrange(mask, 'b n -> b n 1')
masked_t = t.masked_fill(~mask, 0.0)
return masked_t.sum(dim=dim) / denom.clamp(min=1e-5)
+6 -3
View File
@@ -1,5 +1,8 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.head.classifier_head import (
ClassifierHead, CosineLinearHead, TransformerHead, TransformerHeadx2,
VideoClassifierHead, VideoClassifierHeadx2)
from scepter.modules.model.head.classifier_head import (ClassifierHead,
CosineLinearHead,
TransformerHead,
TransformerHeadx2,
VideoClassifierHead,
VideoClassifierHeadx2)
+2 -3
View File
@@ -1,5 +1,4 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.metric.classification import (AccuracyMetric,
EnsembleAccuracyMetric
)
from scepter.modules.model.metric.classification import (
AccuracyMetric, EnsembleAccuracyMetric)
@@ -13,8 +13,7 @@ from .schedules import karras_schedule
from .solvers import (sample_ddim, sample_dpm_2, sample_dpm_2_ancestral,
sample_dpmpp_2m, sample_dpmpp_2m_sde,
sample_dpmpp_2s_ancestral, sample_dpmpp_sde,
sample_euler, sample_euler_ancestral, sample_heun,
sample_img2img_euler, sample_img2img_euler_ancestral)
sample_euler, sample_euler_ancestral, sample_heun)
__all__ = ['GaussianDiffusion']
@@ -27,6 +26,148 @@ def _i(tensor, t, x):
return tensor[t.to(tensor.device)].view(shape).to(x.device)
def _unpack_2d_ks(kernel_size):
if isinstance(kernel_size, int):
ky = kx = kernel_size
else:
assert len(
kernel_size) == 2, '2D Kernel size should have a length of 2.'
ky, kx = kernel_size
ky = int(ky)
kx = int(kx)
return ky, kx
def _compute_zero_padding(kernel_size):
ky, kx = _unpack_2d_ks(kernel_size)
return (ky - 1) // 2, (kx - 1) // 2
def _bilateral_blur(
input,
guidance,
kernel_size,
sigma_color,
sigma_space,
border_type='reflect',
color_distance_type='l1',
):
if isinstance(sigma_color, torch.Tensor):
sigma_color = sigma_color.to(device=input.device,
dtype=input.dtype).view(-1, 1, 1, 1, 1)
ky, kx = _unpack_2d_ks(kernel_size)
pad_y, pad_x = _compute_zero_padding(kernel_size)
padded_input = torch.nn.functional.pad(input, (pad_x, pad_x, pad_y, pad_y),
mode=border_type)
unfolded_input = padded_input.unfold(2, ky, 1).unfold(3, kx, 1).flatten(
-2) # (B, C, H, W, Ky x Kx)
if guidance is None:
guidance = input
unfolded_guidance = unfolded_input
else:
padded_guidance = torch.nn.functional.pad(guidance,
(pad_x, pad_x, pad_y, pad_y),
mode=border_type)
unfolded_guidance = padded_guidance.unfold(2, ky, 1).unfold(
3, kx, 1).flatten(-2) # (B, C, H, W, Ky x Kx)
diff = unfolded_guidance - guidance.unsqueeze(-1)
if color_distance_type == 'l1':
color_distance_sq = diff.abs().sum(1, keepdim=True).square()
elif color_distance_type == 'l2':
color_distance_sq = diff.square().sum(1, keepdim=True)
else:
raise ValueError('color_distance_type only acceps l1 or l2')
color_kernel = (-0.5 / sigma_color**2 *
color_distance_sq).exp() # (B, 1, H, W, Ky x Kx)
space_kernel = get_gaussian_kernel2d(kernel_size,
sigma_space,
device=input.device,
dtype=input.dtype)
space_kernel = space_kernel.view(-1, 1, 1, 1, kx * ky)
kernel = space_kernel * color_kernel
out = (unfolded_input * kernel).sum(-1) / kernel.sum(-1)
return out
def get_gaussian_kernel1d(
kernel_size,
sigma,
force_even,
*,
device=None,
dtype=None,
):
return gaussian(kernel_size, sigma, device=device, dtype=dtype)
def gaussian(window_size, sigma, *, device=None, dtype=None):
batch_size = sigma.shape[0]
x = (torch.arange(window_size, device=sigma.device, dtype=sigma.dtype) -
window_size // 2).expand(batch_size, -1)
if window_size % 2 == 0:
x = x + 0.5
gauss = torch.exp(-x.pow(2.0) / (2 * sigma.pow(2.0)))
return gauss / gauss.sum(-1, keepdim=True)
def get_gaussian_kernel2d(
kernel_size,
sigma,
force_even=False,
*,
device=None,
dtype=None,
):
sigma = torch.Tensor([[sigma, sigma]]).to(device=device, dtype=dtype)
ksize_y, ksize_x = _unpack_2d_ks(kernel_size)
sigma_y, sigma_x = sigma[:, 0, None], sigma[:, 1, None]
kernel_y = get_gaussian_kernel1d(ksize_y,
sigma_y,
force_even,
device=device,
dtype=dtype)[..., None]
kernel_x = get_gaussian_kernel1d(ksize_x,
sigma_x,
force_even,
device=device,
dtype=dtype)[..., None]
return kernel_y * kernel_x.view(-1, 1, ksize_x)
def adaptive_anisotropic_filter(x, g=None):
if g is None:
g = x
s, m = torch.std_mean(g, dim=(1, 2, 3), keepdim=True)
s = s + 1e-5
guidance = (g - m) / s
y = _bilateral_blur(x,
guidance,
kernel_size=(13, 13),
sigma_color=3.0,
sigma_space=3.0,
border_type='reflect',
color_distance_type='l1')
return y
class GaussianDiffusion(object):
def __init__(self, sigmas, prediction_type='eps'):
assert prediction_type in {'x0', 'eps', 'v'}
@@ -53,6 +194,7 @@ class GaussianDiffusion(object):
guide_scale=None,
guide_rescale=None,
clamp=None,
sharpness=0.0,
percentile=None,
cat_uc=False,
**kwargs):
@@ -79,54 +221,99 @@ class GaussianDiffusion(object):
# prediction
if guide_scale is None:
assert isinstance(model_kwargs, dict)
out = model(xt, t=t, **model_kwargs, **kwargs)
if isinstance(model_kwargs, dict):
out = model(xt, t=t, **model_kwargs, **kwargs)
elif isinstance(model_kwargs, list) and len(model_kwargs) > 0:
out = model(xt, t=t, **model_kwargs[0], **kwargs)
else:
raise Exception('Error')
else:
# classifier-free guidance (arXiv:2207.12598)
# model_kwargs[0]: conditional kwargs
# model_kwargs[1]: non-conditional kwargs
assert isinstance(model_kwargs, list) and len(model_kwargs) == 2
if guide_scale == 1.:
out = model(xt, t=t, **model_kwargs[0], **kwargs)
else:
if cat_uc:
def parse_model_kwargs(prev_value, value):
if isinstance(value, torch.Tensor):
prev_value = torch.cat([prev_value, value], dim=0)
elif isinstance(value, dict):
for k, v in value.items():
prev_value[k] = parse_model_kwargs(
prev_value[k], v)
elif isinstance(value, list):
for idx, v in enumerate(value):
prev_value[idx] = parse_model_kwargs(
prev_value[idx], v)
return prev_value
all_model_kwargs = copy.deepcopy(model_kwargs[0])
for model_kwarg in model_kwargs[1:]:
for key, value in model_kwarg.items():
all_model_kwargs[key] = parse_model_kwargs(
all_model_kwargs[key], value)
all_out = model(xt.repeat(2, 1, 1, 1),
t=t.repeat(2),
**all_model_kwargs,
**kwargs)
y_out, u_out = all_out.chunk(2)
assert isinstance(model_kwargs, list) and len(model_kwargs) >= 2
if isinstance(guide_scale, float) or isinstance(guide_scale, int):
assert len(model_kwargs) == 2
if guide_scale == 1.:
out = model(xt, t=t, **model_kwargs[0], **kwargs)
else:
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)
if cat_uc:
# rescale the output according to arXiv:2305.08891
if guide_rescale is not None:
assert guide_rescale >= 0 and guide_rescale <= 1
ratio = (y_out.flatten(1).std(dim=1) /
(out.flatten(1).std(dim=1) +
1e-12)).view((-1, ) + (1, ) * (y_out.ndim - 1))
out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0
def parse_model_kwargs(prev_value, value):
if isinstance(value, torch.Tensor):
prev_value = torch.cat([prev_value, value],
dim=0)
elif isinstance(value, dict):
for k, v in value.items():
prev_value[k] = parse_model_kwargs(
prev_value[k], v)
elif isinstance(value, list):
for idx, v in enumerate(value):
prev_value[idx] = parse_model_kwargs(
prev_value[idx], v)
return prev_value
all_model_kwargs = copy.deepcopy(model_kwargs[0])
for model_kwarg in model_kwargs[1:]:
for key, value in model_kwarg.items():
all_model_kwargs[key] = parse_model_kwargs(
all_model_kwargs[key], value)
all_out = model(xt.repeat(2, 1, 1, 1),
t=t.repeat(2),
**all_model_kwargs,
**kwargs)
y_out, u_out = all_out.chunk(2)
else:
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
u_out = model(xt, t=t, **model_kwargs[1], **kwargs)
# todo sharpness
# sharpness sampling
if sharpness is not None and sharpness > 0:
positive_x0 = alphas * xt - sigmas * y_out
negative_x0 = alphas * xt - sigmas * u_out
positive_eps = xt - positive_x0
negative_eps = xt - negative_x0
global_diffusion_progress = (
1 - t / 999.0).detach().cpu().numpy().tolist()[0]
alpha = 0.001 * sharpness * global_diffusion_progress
positive_eps_degraded = adaptive_anisotropic_filter(
x=positive_eps, g=positive_x0)
positive_eps_degraded_weighted = positive_eps_degraded * alpha + positive_eps * (
1.0 - alpha)
final_eps = negative_eps + guide_scale * (
positive_eps_degraded_weighted - negative_eps)
final_x0 = xt - final_eps
out = (alphas * xt - final_x0) / sigmas
else:
out = u_out + guide_scale * (y_out - u_out)
elif isinstance(guide_scale, dict):
assert len(model_kwargs) == 3
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
m_out = model(xt, t=t, **model_kwargs[1], **kwargs)
u_out = model(xt, t=t, **model_kwargs[2], **kwargs)
out = u_out + guide_scale['image'] * (
m_out - u_out) + guide_scale['text'] * (y_out - m_out)
elif isinstance(guide_scale, list):
assert len(guide_scale) == len(model_kwargs) - 1
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
outs = [y_out]
for i in range(1, len(model_kwargs)):
outs.append(model(xt, t=t, **model_kwargs[i], **kwargs))
out = outs[-1]
for i in range(len(guide_scale)):
out += guide_scale[i] * (outs[-i - 2] - outs[-i - 1])
# rescale the output according to arXiv:2305.08891
if guide_rescale is not None and guide_rescale > 0.0:
assert guide_rescale >= 0 and guide_rescale <= 1
ratio = (
y_out.flatten(1).std(dim=1) /
(out.flatten(1).std(dim=1) + 1e-12)).view((-1, ) + (1, ) *
(y_out.ndim - 1))
out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0
# compute x0
if self.prediction_type == 'x0':
x0 = out
@@ -197,6 +384,7 @@ class GaussianDiffusion(object):
guide_scale=None,
guide_rescale=None,
clamp=None,
sharpness=0.0,
percentile=None,
solver='euler_a',
steps=20,
@@ -209,12 +397,16 @@ class GaussianDiffusion(object):
seed=-1,
intermediate_callback=None,
cat_uc=False,
add_noise=False,
free_steps=None,
step_offset=None,
**kwargs):
# sanity check
assert isinstance(steps, (int, torch.LongTensor))
assert t_max is None or (t_max > 0 and t_max <= self.num_timesteps - 1)
assert t_min is None or (t_min >= 0 and t_min < self.num_timesteps - 1)
assert discretization in (None, 'leading', 'linspace', 'trailing')
assert discretization in (None, 'leading', 'linspace', 'trailing',
'free')
assert discard_penultimate_step in (None, True, False)
assert return_intermediate in (None, 'x0', 'xt')
@@ -255,17 +447,50 @@ class GaussianDiffusion(object):
def model_fn(xt, sigma):
# denoising
t = self._sigma_to_t(sigma).repeat(len(xt)).round().long()
x0 = self.denoise(xt,
t,
None,
model,
model_kwargs,
guide_scale,
guide_rescale,
clamp,
percentile,
cat_uc=cat_uc,
**kwargs)[-2]
if isinstance(
model_kwargs[0]['cond'], dict) and \
'tar_x0' in model_kwargs[0]['cond'] and \
'tar_mask_latent' in model_kwargs[0]['cond']:
tar_x0 = model_kwargs[0]['cond']['tar_x0']
tar_mask = model_kwargs[0]['cond']['tar_mask_latent']
tar_xt = self.diffuse(x0=tar_x0, t=t)
xt = tar_xt * (1.0 - tar_mask) + xt * tar_mask
if isinstance(model_kwargs[0]['cond'],
dict) and 'ref_x0' in model_kwargs[0]['cond']:
model_kwargs[0]['cond']['ref_xt'] = self.diffuse(
x0=model_kwargs[0]['cond']['ref_x0'], t=t)
model_kwargs[1]['cond']['ref_xt'] = self.diffuse(
x0=model_kwargs[1]['cond']['ref_x0'], t=t)
if solver in ('onestep', 'multistep', 'multistep2', 'multistep3'):
x0 = self.denoise(xt,
t,
None,
model,
model_kwargs,
guide_scale,
guide_rescale,
clamp,
sharpness,
percentile,
cat_uc=cat_uc,
**kwargs)[-3]
else:
x0 = self.denoise(xt,
t,
None,
model,
model_kwargs,
guide_scale,
guide_rescale,
clamp,
sharpness,
percentile,
cat_uc=cat_uc,
**kwargs)[-2]
# collect intermediate outputs
if return_intermediate == 'xt':
@@ -291,10 +516,14 @@ class GaussianDiffusion(object):
elif discretization == 'trailing':
steps = torch.arange(t_max, t_min - 1,
-((t_max - t_min + 1) / steps))
elif discretization == 'free':
steps = torch.tensor(free_steps)
else:
raise NotImplementedError(
f'{discretization} discretization not implemented')
steps = steps.clamp_(t_min, t_max)
elif isinstance(steps, list):
steps = torch.tensor(steps)
steps = torch.as_tensor(steps,
dtype=torch.float32,
device=noise.device)
@@ -335,6 +564,23 @@ class GaussianDiffusion(object):
if discard_penultimate_step:
sigmas = torch.cat([sigmas[:-2], sigmas[-1:]])
kwargs['seed'] = seed
# add noise to x0
if add_noise:
if 'dm_steps' in kwargs:
if step_offset:
add_noise_step = -kwargs['dm_steps'] + step_offset
if add_noise_step < 0:
noise = self.diffuse(
noise,
torch.full((noise.shape[0], 1),
steps[add_noise_step],
dtype=torch.int))
else:
noise = self.diffuse(
noise,
torch.full((noise.shape[0], 1),
steps[-kwargs['dm_steps'] - 1],
dtype=torch.int))
# sampling
x0 = solver_fn(noise,
model_fn,
@@ -373,168 +619,6 @@ class GaussianDiffusion(object):
| torch.isinf(log_sigma)] = float('inf')
return log_sigma.exp()
@torch.no_grad()
def stochastic_encode(self, x0, t, steps):
# fast, but does not allow for exact reconstruction
# t serves as an index to gather the correct alphas
t_max = None
t_min = None
# discretization method
discretization = 'trailing' if self.prediction_type == 'v' else 'leading'
# timesteps
if isinstance(steps, int):
t_max = self.num_timesteps - 1 if t_max is None else t_max
t_min = 0 if t_min is None else t_min
steps = discretize_timesteps(t_max, t_min, steps, discretization)
steps = torch.as_tensor(steps).round().long().flip(0).to(x0.device)
# steps = torch.as_tensor(steps).round().long().to(x0.device)
# self.alphas_bar = torch.cumprod(1 - self.sigmas ** 2, dim=0)
# print('sigma: ', self.sigmas, len(self.sigmas))
# print('alpha_bar: ', self.alphas_bar, len(self.alphas_bar))
# print('steps: ', steps, len(steps))
# sqrt_alphas_cumprod = torch.sqrt(self.alphas_bar).to(x0.device)[steps]
# sqrt_one_minus_alphas_cumprod = torch.sqrt(1 - self.alphas_bar).to(x0.device)[steps]
sqrt_alphas_cumprod = self.alphas.to(x0.device)[steps]
sqrt_one_minus_alphas_cumprod = self.sigmas.to(x0.device)[steps]
# print('sigma: ', self.sigmas, len(self.sigmas))
# print('alpha: ', self.alphas, len(self.alphas))
# print('steps: ', steps, len(steps))
noise = torch.randn_like(x0)
return (
extract_into_tensor(sqrt_alphas_cumprod, t, x0.shape) * x0 +
extract_into_tensor(sqrt_one_minus_alphas_cumprod, t, x0.shape) *
noise)
@torch.no_grad()
def sample_img2img(self,
x,
noise,
model,
denoising_strength=1,
model_kwargs={},
condition_fn=None,
guide_scale=None,
guide_rescale=None,
clamp=None,
percentile=None,
solver='euler_a',
steps=20,
t_max=None,
t_min=None,
discretization=None,
discard_penultimate_step=None,
return_intermediate=None,
show_progress=False,
seed=-1,
**kwargs):
# sanity check
assert isinstance(steps, (int, torch.LongTensor))
assert t_max is None or (t_max > 0 and t_max <= self.num_timesteps - 1)
assert t_min is None or (t_min >= 0 and t_min < self.num_timesteps - 1)
assert discretization in (None, 'leading', 'linspace', 'trailing')
assert discard_penultimate_step in (None, True, False)
assert return_intermediate in (None, 'x0', 'xt')
# function of diffusion solver
solver_fn = {
'euler_ancestral': sample_img2img_euler_ancestral,
'euler': sample_img2img_euler,
}[solver]
# options
schedule = 'karras' if 'karras' in solver else None
discretization = discretization or 'linspace'
seed = seed if seed >= 0 else random.randint(0, 2**31)
if isinstance(steps, torch.LongTensor):
discard_penultimate_step = False
if discard_penultimate_step is None:
discard_penultimate_step = True if solver in (
'dpm2', 'dpm2_ancestral', 'dpmpp_2m_sde', 'dpm2_karras',
'dpm2_ancestral_karras', 'dpmpp_2m_sde_karras') else False
# function for denoising xt to get x0
intermediates = []
def get_scalings(sigma):
c_out = -sigma
c_in = 1 / (sigma**2 + 1.**2)**0.5
return c_out, c_in
def model_fn(xt, sigma):
# denoising
c_out, c_in = get_scalings(sigma)
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,
**kwargs)[-2]
# collect intermediate outputs
if return_intermediate == 'xt':
intermediates.append(xt)
elif return_intermediate == 'x0':
intermediates.append(x0)
return xt + x0 * c_out
# get timesteps
if isinstance(steps, int):
steps += 1 if discard_penultimate_step else 0
t_max = self.num_timesteps - 1 if t_max is None else t_max
t_min = 0 if t_min is None else t_min
# discretize timesteps
if discretization == 'leading':
steps = torch.arange(t_min, t_max + 1,
(t_max - t_min + 1) / steps).flip(0)
elif discretization == 'linspace':
steps = torch.linspace(t_max, t_min, steps)
elif discretization == 'trailing':
steps = torch.arange(t_max, t_min - 1,
-((t_max - t_min + 1) / steps))
else:
raise NotImplementedError(
f'{discretization} discretization not implemented')
steps = steps.clamp_(t_min, t_max)
steps = torch.as_tensor(steps, dtype=torch.float32, device=x.device)
# get sigmas
sigmas = self._t_to_sigma(steps)
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
t_enc = int(min(denoising_strength, 0.999) * len(steps))
sigmas = sigmas[len(steps) - t_enc - 1:]
noise = x + noise * sigmas[0]
if schedule == 'karras':
if sigmas[0] == float('inf'):
sigmas = karras_schedule(
n=len(steps) - 1,
sigma_min=sigmas[sigmas > 0].min().item(),
sigma_max=sigmas[sigmas < float('inf')].max().item(),
rho=7.).to(sigmas)
sigmas = torch.cat([
sigmas.new_tensor([float('inf')]), sigmas,
sigmas.new_zeros([1])
])
else:
sigmas = karras_schedule(
n=len(steps),
sigma_min=sigmas[sigmas > 0].min().item(),
sigma_max=sigmas.max().item(),
rho=7.).to(sigmas)
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
if discard_penultimate_step:
sigmas = torch.cat([sigmas[:-2], sigmas[-1:]])
# sampling
x0 = solver_fn(noise,
model_fn,
sigmas,
seed=seed,
show_progress=show_progress,
**kwargs)
return (x0, intermediates) if return_intermediate is not None else x0
def extract_into_tensor(a, t, x_shape):
b, *_ = t.shape
+148
View File
@@ -0,0 +1,148 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import cv2
import numpy as np
import torchvision.transforms.functional as TF
def get_bbox_from_mask(mask):
h, w = mask.shape[0], mask.shape[1]
if mask.sum() < 10:
return 0, h, 0, w
rows = np.any(mask, axis=1)
cols = np.any(mask, axis=0)
y1, y2 = np.where(rows)[0][[0, -1]]
x1, x2 = np.where(cols)[0][[0, -1]]
return (y1, y2, x1, x2)
def pad_to_square(image, pad_value=255, random=False):
H, W = image.shape[0], image.shape[1]
if H == W:
return image, 0, 0
padd = abs(H - W)
if random:
padd_1 = int(np.random.randint(0, padd))
else:
padd_1 = int(padd / 2)
padd_2 = padd - padd_1
if H > W:
pad_param = ((0, 0), (padd_1, padd_2), (0, 0))
else:
pad_param = ((padd_1, padd_2), (0, 0), (0, 0))
# print(pad_param, pad_value)
image = np.pad(image, pad_param, 'constant', constant_values=pad_value)
return image, padd_1, padd_2
def box_in_box(small_box, big_box):
y1, y2, x1, x2 = small_box
y1_b, _, x1_b, _ = big_box
y1, y2, x1, x2 = y1 - y1_b, y2 - y1_b, x1 - x1_b, x2 - x1_b
return (y1, y2, x1, x2)
def box2squre(image, box):
H, W = image.shape[0], image.shape[1]
y1, y2, x1, x2 = box
cx = (x1 + x2) // 2
cy = (y1 + y2) // 2
h, w = y2 - y1, x2 - x1
if h >= w:
x1 = cx - h // 2
x2 = x1 + h
else:
y1 = cy - w // 2
y2 = y1 + w
x1 = max(0, x1)
x2 = min(W, x2)
y1 = max(0, y1)
y2 = min(H, y2)
return (y1, y2, x1, x2)
def expand_bbox(mask,
yyxx,
ratio=1.0,
min_crop=0,
expand_type='center',
to_square=False):
y1, y2, x1, x2 = yyxx
h = y2 - y1 + 1
w = x2 - x1 + 1
H, W = mask.shape[0], mask.shape[1]
xc, yc = 0.5 * (x1 + x2), 0.5 * (y1 + y2)
def expand(k):
if isinstance(ratio, tuple) or isinstance(ratio, list):
r = np.random.uniform(*ratio)
k = k * r
else:
k = ratio * k
return k
new_h = expand(h)
new_w = expand(w)
new_h = max(new_h, min_crop)
new_w = max(new_w, min_crop)
if to_square:
if new_w / new_h < 0.334:
new_w = new_w + 1.0 / 3.0 * new_h
elif new_h / new_w < 0.334:
new_h = new_h + 1.0 / 3.0 * new_w
if expand_type == 'center':
x1 = max(0, int(xc - new_w * 0.5))
x2 = min(W, int(xc + new_w * 0.5))
y1 = max(0, int(yc - new_h * 0.5))
y2 = min(H, int(yc + new_h * 0.5))
else:
x1 = max(0, min(x1,
int(x2 - new_w * np.random.uniform(w / new_w, 1.0))))
x2 = min(W, max(x2, x1 + new_w))
y1 = max(0, min(y1,
int(y2 - new_h * np.random.uniform(h / new_h, 1.0))))
y2 = min(H, max(y2, y1 + new_h))
return (int(y1), int(y2), int(x1), int(x2))
def crop_back(pred, tar_image, extra_sizes, tar_box_yyxx_crop):
H1, W1, H2, W2, pad1, pad2 = extra_sizes
y1, y2, x1, x2 = tar_box_yyxx_crop
pred = TF.resize(pred, (H2, W2), antialias=True)
# if W1 == H1:
# tar_image[:, y1:y2, x1:x2] = pred
# return tar_image
# if W1 < W2:
# pad1 = int((W2 - W1) / 2)
# pad2 = W2 - W1 - pad1
# pred = pred[:, :,pad1:-pad2]
# else:
# pad1 = int((H2 - H1) / 2)
# pad2 = H2 - H1 - pad1
# pred = pred[:, pad1:-pad2, :]
if W1 < W2:
# pad width
assert H1 == H2 and (pad1 + W1) == (W2 - pad2)
pred = pred[:, :, pad1 + 2:(W2 - pad2 - 2)]
tar_image[:, y1:y2, x1 + 2:x2 - 2] = pred
elif H1 < H2:
# pad height
assert W1 == W2 and (pad1 + H1) == (H2 - pad2)
pred = pred[:, pad1 + 2:(H2 - pad2 - 2), :]
tar_image[:, y1 + 2:y2 - 2, x1:x2] = pred
else:
tar_image[:, y1:y2, x1:x2] = pred
return tar_image
def save_image(image, save_path):
image = cv2.cvtColor(image, cv2.COLOR_RGB2BGR)
cv2.imwrite(save_path, image)
+6 -3
View File
@@ -1,7 +1,10 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.opt.optimizers.official_optimizers import (
ASGD, LBFGS, SGD, Adadelta, Adagrad, Adam, Adamax, AdamW, RMSprop, Rprop,
SparseAdam)
from scepter.modules.opt.optimizers.official_optimizers import (ASGD, LBFGS,
SGD, Adadelta,
Adagrad, Adam,
Adamax, AdamW,
RMSprop, Rprop,
SparseAdam)
from scepter.modules.opt.optimizers.registry import OPTIMIZERS
+6 -3
View File
@@ -333,12 +333,12 @@ class LatentDiffusionSolver(BaseSolver):
data_iter = iter(self.datas[self._mode].dataloader)
self.print_memory_status()
for step in range(self.max_steps):
if 'eval' in self._mode_set and (step % self.eval_interval == 0
or step == self.max_steps - 1):
if 'eval' in self._mode_set and (self.eval_interval > 0 and
step % self.eval_interval == 0):
self.run_eval()
self.train_mode()
self.before_iter(self.hooks_dict[self._mode])
batch_data = next(data_iter)
self.before_iter(self.hooks_dict[self._mode])
if self.sample_args:
batch_data.update(self.sample_args.get_lowercase_dict())
if 'meta' in batch_data:
@@ -364,6 +364,9 @@ class LatentDiffusionSolver(BaseSolver):
self.after_iter(self.hooks_dict[self._mode])
if we.debug:
self.print_trainable_params_status(prefix='model.')
if 'eval' in self._mode_set and (self.eval_interval > 0
and step == self.max_steps - 1):
self.run_eval()
self.after_all_iter(self.hooks_dict[self._mode])
@torch.no_grad()
+4 -1
View File
@@ -4,12 +4,14 @@
from scepter.modules.solver.hooks.backward import BackwardHook
from scepter.modules.solver.hooks.checkpoint import CheckpointHook
from scepter.modules.solver.hooks.data_probe import ProbeDataHook
from scepter.modules.solver.hooks.ema import ModelEmaHook
from scepter.modules.solver.hooks.hook import Hook
from scepter.modules.solver.hooks.log import LogHook, TensorboardLogHook
from scepter.modules.solver.hooks.lr import LrHook
from scepter.modules.solver.hooks.registry import HOOKS
from scepter.modules.solver.hooks.safetensors import SafetensorsHook
from scepter.modules.solver.hooks.sampler import DistSamplerHook
"""
Normally, hooks have priorities, below we recommend priority that runs fine (low score MEANS high priority)
BackwardHook: 0
@@ -47,5 +49,6 @@ after solve:
__all__ = [
'HOOKS', 'BackwardHook', 'CheckpointHook', 'Hook', 'LrHook', 'LogHook',
'TensorboardLogHook', 'DistSamplerHook', 'ProbeDataHook', 'SafetensorsHook'
'TensorboardLogHook', 'DistSamplerHook', 'ProbeDataHook',
'SafetensorsHook', 'ModelEmaHook'
]
+3 -1
View File
@@ -147,7 +147,9 @@ class CheckpointHook(Hook):
if self.save_last and solver.total_iter == solver.max_steps - 1:
with FS.get_fs_client(save_path) as client:
last_path = osp.join(solver.work_dir, 'checkpoint.pth')
last_path = osp.join(
solver.work_dir,
f'checkpoints/{self.save_name_prefix}-last')
client.make_link(last_path, save_path)
self.last_ckpt = save_path
+35 -12
View File
@@ -30,7 +30,10 @@ class ProbeDataHook(Hook):
def __init__(self, cfg, logger=None):
super(ProbeDataHook, self).__init__(cfg, logger=logger)
self.priority = cfg.get('PRIORITY', _DEFAULT_PROBE_PRIORITY)
self.log_interval = cfg.get('PROB_INTERVAL', 1000)
self.prob_interval = cfg.get('PROB_INTERVAL', 1000)
self.save_name_prefix = cfg.get('SAVE_NAME_PREFIX', 'step')
self.save_probe_prefix = cfg.get('SAVE_PROBE_PREFIX', None)
self.save_last = cfg.get('SAVE_LAST', False)
def before_all_iter(self, solver):
pass
@@ -39,19 +42,23 @@ class ProbeDataHook(Hook):
pass
def after_iter(self, solver):
if solver.mode == 'train' and solver.total_iter % self.log_interval == 0:
if solver.mode == 'train' and solver.total_iter % self.prob_interval == 0:
probe_dict = solver.probe_data
if we.rank == 0:
save_folder = os.path.join(
solver.work_dir,
f'{solver.mode}_probe/step_{solver.total_iter}')
f'{solver.mode}_probe/{self.save_probe_prefix}-{solver.total_iter}'
)
ret_data = {}
for k, v in probe_dict.items():
ret_one = v.to_log(
os.path.join(
if self.save_probe_prefix is not None:
ret_prefix = os.path.join(save_folder,
self.save_probe_prefix)
else:
ret_prefix = os.path.join(
save_folder,
k.replace('/', '_') +
f'_step_{solver.total_iter}'))
k.replace('/', '_') + f'_step_{solver.total_iter}')
ret_one = v.to_log(ret_prefix)
if (isinstance(ret_one, list)
or isinstance(ret_one, dict)) and len(ret_one) < 1:
continue
@@ -71,13 +78,19 @@ class ProbeDataHook(Hook):
if we.rank == 0:
step = solver._total_iter[
'train'] if 'train' in solver._total_iter else 0
save_folder = os.path.join(solver.work_dir,
f'{solver.mode}_probe/step_{step}')
save_folder = os.path.join(
solver.work_dir,
f'{solver.mode}_probe/{self.save_name_prefix}-{step}')
ret_data = {}
for k, v in probe_dict.items():
ret_one = v.to_log(
os.path.join(save_folder,
k.replace('/', '_') + f'_step_{step}'))
if self.save_probe_prefix is not None:
ret_prefix = os.path.join(save_folder,
self.save_probe_prefix)
else:
ret_prefix = os.path.join(
save_folder,
k.replace('/', '_') + f'_step_{step}')
ret_one = v.to_log(ret_prefix)
if (isinstance(ret_one, list)
or isinstance(ret_one, dict)) and len(ret_one) < 1:
continue
@@ -87,6 +100,16 @@ class ProbeDataHook(Hook):
json.dump(ret_data,
open(local_path, 'w'),
ensure_ascii=False)
if self.save_last and step == solver.max_steps:
with FS.get_fs_client(save_folder) as client:
last_save_folder = os.path.join(
solver.work_dir,
f'{solver.mode}_probe/{self.save_name_prefix}-last'
)
print(last_save_folder, save_folder)
client.make_link(last_save_folder, save_folder)
solver.clear_probe()
torch.cuda.synchronize()
barrier()
+60
View File
@@ -0,0 +1,60 @@
# -*- coding: utf-8 -*-
import torch
from torch.distributed.fsdp import (FullStateDictConfig,
FullyShardedDataParallel, StateDictType)
from scepter.modules.solver.hooks.hook import Hook
from scepter.modules.solver.hooks.registry import HOOKS
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
@HOOKS.register_class()
class ModelEmaHook(Hook):
para_dict = [{
'PRIORITY': {
'value': 100,
'description': 'the priority for processing!'
},
'BETA': {
'value': 0.9999,
'description': ''
},
}]
def __init__(self, cfg, logger=None):
super(ModelEmaHook, self).__init__(cfg, logger=logger)
self.priority = cfg.get('PRIORITY', 100)
self.beta = cfg.get('BETA', 0.9999)
def after_iter(self, solver):
if solver.model.use_ema:
model_ema = solver.model.model_ema
model = solver.model.model
self.ema(model_ema, model, use_fsdp=solver.use_fsdp)
@torch.no_grad()
def ema(self, net_ema, net, use_fsdp=True):
if we.is_distributed:
if use_fsdp:
save_policy = FullStateDictConfig(offload_to_cpu=False,
rank0_only=False)
with FullyShardedDataParallel.state_dict_type(
net, StateDictType.FULL_STATE_DICT, save_policy):
nonema_state = net.state_dict()
elif hasattr(net, 'module'):
nonema_state = net.module.state_dict()
else:
nonema_state = net.state_dict()
else:
nonema_state = net.state_dict()
for k, v in net_ema.named_parameters():
v.copy_(nonema_state[k].lerp(v, self.beta))
@staticmethod
def get_config_template():
return dict_to_yaml('hook',
__class__.__name__,
ModelEmaHook.para_dict,
set_name=True)
+10 -5
View File
@@ -10,11 +10,16 @@ import torchvision.transforms as transforms
import torchvision.transforms.functional as TF
from scepter.modules.transform.registry import TRANSFORMS
from scepter.modules.transform.utils import (
BACKEND_CV2, BACKEND_PILLOW, BACKEND_TORCHVISION, INPUT_CV2_TYPE_WARNING,
INPUT_PIL_TYPE_WARNING, INPUT_TENSOR_TYPE_WARNING, INTERPOLATION_STYLE,
INTERPOLATION_STYLE_CV2, TORCHVISION_CAPABILITY, is_cv2_image,
is_pil_image, is_tensor)
from scepter.modules.transform.utils import (BACKEND_CV2, BACKEND_PILLOW,
BACKEND_TORCHVISION,
INPUT_CV2_TYPE_WARNING,
INPUT_PIL_TYPE_WARNING,
INPUT_TENSOR_TYPE_WARNING,
INTERPOLATION_STYLE,
INTERPOLATION_STYLE_CV2,
TORCHVISION_CAPABILITY,
is_cv2_image, is_pil_image,
is_tensor)
from scepter.modules.utils.config import dict_to_yaml
if TORCHVISION_CAPABILITY:
+3 -1
View File
@@ -6,13 +6,15 @@ import os
import cv2
import numpy as np
import torch
from PIL import Image
from PIL import Image, ImageFile
from scepter.modules.transform.registry import TRANSFORMS
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import DATA_FS as FS
ImageFile.LOAD_TRUNCATED_IMAGES = True
def pillow_convert(image, rgb_order):
if image.mode != rgb_order:
+3
View File
@@ -603,3 +603,6 @@ class Config(object):
return cfg_new
else:
return cfg
def pop(self, name):
self.cfg_dict.pop(name)
+1 -1
View File
@@ -4,10 +4,10 @@ import io
from io import BytesIO
import onnx
import onnxruntime
import torch
from torch.onnx import OperatorExportTypes
import onnxruntime
from scepter.modules.utils.distribute import we
type_map = {
@@ -293,6 +293,9 @@ class LocalFs(BaseFs):
return True
def get_logging_handler(self, target_logging_path):
dirname = os.path.dirname(target_logging_path)
if not os.path.exists(dirname):
os.makedirs(dirname, exist_ok=True)
return logging.FileHandler(target_logging_path)
def put_dir_from_local_dir(self,
@@ -303,11 +306,11 @@ class LocalFs(BaseFs):
target_dir = self.reconstruct_path(target_dir)
if local_dir == target_dir:
return True
# cp -f local_dir/* target_dir/*
if not osp.exists(target_dir):
status = os.system(f'mkdir -p {target_dir}')
if status != 0:
return False
# # cp -f local_dir/* target_dir/*
# if not osp.exists(target_dir):
# status = os.system(f'mkdir -p {target_dir}')
# if status != 0:
# return False
try:
shutil.copytree(local_dir, target_dir, symlinks=True)
except Exception:
+5
View File
@@ -1,5 +1,6 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import io
import os
import threading
import time
@@ -327,6 +328,10 @@ class FileSystem(object):
target_path_list):
if local_path is None or target_path is None:
flg = False
elif isinstance(local_path, io.BytesIO):
flg = FS.put_object(local_path.getvalue(), target_path)
elif isinstance(local_path, bytes):
flg = FS.put_object(local_path, target_path)
elif self.exists(local_path):
local_cache = self.get_from(local_path,
local_path + f'{time.time()}',
+44
View File
@@ -0,0 +1,44 @@
# -*- coding: utf-8 -*-
import os.path as osp
from tqdm import tqdm
from scepter.modules.utils.file_system import FS
def init_1level_llfs(list_file,
max_lines=1024,
index_name='index',
delimiter='\n'):
r"""Construct large-list-file index.
"""
index_dir = osp.splitext(list_file)[0]
print(list_file)
with FS.get_from(list_file, wait_finish=True) as local_path:
print(local_path)
_num_split, index_stack, save_files = 0, [], []
with open(local_path, 'r', buffering=1000000) as f:
for line in tqdm(f):
index_stack.append(line.strip())
if len(index_stack) >= max_lines:
save_file = f'{index_dir}/{index_name}/{_num_split + 1:09d}.txt'
with FS.put_to(save_file) as cache_path:
with open(cache_path, 'w') as f_w:
f_w.write('\n'.join(index_stack))
index_stack = []
_num_split += 1
save_files.append(save_file)
if len(index_stack) > 0:
save_file = f'{index_dir}/{index_name}/{_num_split + 1:06d}.txt'
with FS.put_to(save_file) as cache_path:
with open(cache_path, 'w') as f_w:
f_w.write('\n'.join(index_stack))
save_files.append(save_file)
# output meta-file
index_file = osp.join(index_dir, f'{index_name}.txt')
with FS.put_to(index_file) as cache_path:
with open(cache_path, 'w') as f_w:
f_w.write('\n'.join(save_files))
return index_file
+2 -1
View File
@@ -263,7 +263,8 @@ class ProbeData():
url = FS.get_url(one_path,
lifecycle=3600 * 365 * 24).replace(
'.oss-internal.aliyun-inc.',
'.oss.aliyuncs.')
'.oss.aliyuncs.').replace(
'-internal', '')
one_rank += (
f'<td align="center"><input type="image" src="{url}" >'
f'<br><font size="4"><strong>{save_id}-{idx}|{one_label}<strong></font><br/></td>'
+76 -8
View File
@@ -16,6 +16,7 @@ from scepter.studio.inference.inference_ui.component_names import \
from scepter.studio.inference.inference_ui.control_ui import ControlUI
from scepter.studio.inference.inference_ui.diffusion_ui import DiffusionUI
from scepter.studio.inference.inference_ui.gallery_ui import GalleryUI
from scepter.studio.inference.inference_ui.largen_ui import LargenUI
from scepter.studio.inference.inference_ui.mantra_ui import MantraUI
from scepter.studio.inference.inference_ui.model_manage_ui import ModelManageUI
from scepter.studio.inference.inference_ui.refiner_ui import RefinerUI
@@ -23,7 +24,7 @@ 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)]
('control', ControlUI), ('refiner', RefinerUI), ('largen', LargenUI)]
class InferenceUI():
@@ -54,6 +55,19 @@ class InferenceUI():
cfg_general.EXTENSION_PARAS.OFFICIAL_CONTROLLERS))
cfg_general.CONTROLLERS = official_controllers.CONTROLLERS
# customized tuners
tuner_manager = Config(
cfg_file=os.path.join(os.path.dirname(scepter.dirname),
cfg_general.EXTENSION_PARAS.TUNER_MANAGER))
tuner_manager = os.path.join(root_work_dir, tuner_manager.WORK_DIR,
tuner_manager.TUNER_LIST_YAML)
if FS.exists(tuner_manager):
with FS.get_from(tuner_manager) as local_path:
custom_tuners = Config(cfg_file=local_path)
cfg_general.CUSTOM_TUNERS = custom_tuners.get('TUNERS', [])
else:
cfg_general.CUSTOM_TUNERS = []
pipe_manager = PipelineManager()
config_list = glob(os.path.join(config_dir, '*/*_pro.yaml'),
recursive=True)
@@ -68,6 +82,12 @@ class InferenceUI():
for one_controller in cfg_general.CONTROLLERS:
pipe_manager.register_controllers(one_controller)
for one_tuner in cfg_general.CUSTOM_TUNERS:
pipe_manager.register_tuner(
one_tuner,
name=one_tuner.NAME_ZH if language == 'zh' else one_tuner.NAME,
is_customized=True)
self.model_manage_ui = ModelManageUI(cfg_general,
pipe_manager,
is_debug=is_debug,
@@ -88,7 +108,8 @@ class InferenceUI():
self.tab_ui_kwargs[f'{name}_ui'] = ui
self.__setattr__(f'{name}_ui', ui)
self.check_box_controlled_tabs = ['mantra', 'tuner', 'control']
self.check_box_controlled_tabs = ['mantra', 'tuner', 'control', 'largen']
self.pipe_manager = pipe_manager
assert len(self.component_names.check_box_for_setting) == len(
self.check_box_controlled_tabs)
@@ -96,6 +117,7 @@ class InferenceUI():
# create model
self.model_manage_ui.create_ui()
self.gallery_ui.create_ui()
self.infer_info = gr.State(value=None)
# create tabs
def create_tab(name, ui):
@@ -123,20 +145,65 @@ class InferenceUI():
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)
gallery_ui=self.gallery_ui,
manager=manager)
def change_setting_tab(check_box, *args):
def change_setting_tab(check_box, default_diffusion_model, *args):
selected_tab = 'diffusion_ui'
ui_tabs_state = [False] * len(args)
largen_index = self.check_box_controlled_tabs.index('largen')
largen_key = self.component_names.check_box_for_setting[largen_index]
largen_status = args[largen_index]
for key in check_box:
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'
new_check_box_value = check_box
for key in check_box:
i = self.component_names.check_box_for_setting.index(key)
if ui_tabs_state[i] != args[i]:
if i in [largen_index]:
new_check_box_value = [key]
for j in range(len(ui_tabs_state)):
ui_tabs_state[j] = j == i
else:
new_check_box_value = [
k for k in check_box
if k not in [largen_key]
]
ui_tabs_state[largen_index] = False
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
if ui_tabs_state[largen_index]:
diffusion_model = gr.Dropdown(
label=self.model_manage_ui.component_names.diffusion_model,
choices=self.model_manage_ui.
default_choices['diffusion_model']['choices'],
value='LARGEN_LargenUNetXL',
interactive=False)
elif largen_status:
diffusion_model = gr.Dropdown(
label=self.model_manage_ui.component_names.diffusion_model,
choices=self.model_manage_ui.
default_choices['diffusion_model']['choices'],
value=self.model_manage_ui.
default_choices['diffusion_model']['default'],
interactive=True)
else:
diffusion_model = gr.Dropdown(
label=self.model_manage_ui.component_names.diffusion_model,
choices=self.model_manage_ui.
default_choices['diffusion_model']['choices'],
value=default_diffusion_model,
interactive=True)
return gr.CheckboxGroup(
choices=self.component_names.check_box_for_setting,
value=new_check_box_value,
show_label=False), gr.update(
selected=selected_tab), *ui_tabs_state, *ui_tabs_updates, diffusion_model
gr_states = [
self.tab_ui[name].state for name in self.check_box_controlled_tabs
@@ -146,8 +213,9 @@ class InferenceUI():
]
self.check_box_for_setting.change(
change_setting_tab,
inputs=[self.check_box_for_setting, *gr_states],
outputs=[self.setting_tab, *gr_states, *gr_tabs],
inputs=[self.check_box_for_setting, self.model_manage_ui.diffusion_model, *gr_states],
outputs=[self.check_box_for_setting, self.setting_tab, *gr_states,
*gr_tabs, self.model_manage_ui.diffusion_model],
queue=False)
@@ -1,6 +1,7 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.inference.diffusion_inference import DiffusionInference
from scepter.modules.inference.largen_inference import LargenInference
from scepter.modules.utils.logger import get_logger
@@ -95,7 +96,12 @@ class PipelineManager():
pass
def register_pipeline(self, cfg):
new_inference = DiffusionInference(logger=self.logger)
pipeline_name = cfg.NAME
if 'LARGEN' in pipeline_name:
PipelineBuilder = LargenInference
else:
PipelineBuilder = DiffusionInference
new_inference = PipelineBuilder(logger=self.logger)
new_inference.init_from_cfg(cfg)
self.contruct_models_index(cfg.NAME, new_inference)
@@ -1,16 +1,17 @@
# -*- 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
from scepter.modules.utils.file_system import FS
def download_image(image):
if image is not None:
client = FS.get_fs_client(image)
if client.tmp_dir.startswith("/home"):
if client.tmp_dir.startswith('/home'):
name = get_md5(image)
local_path = FS.get_from(image, f"/tmp/gradio/scepter_examples/{name}")
local_path = FS.get_from(image,
f'/tmp/gradio/scepter_examples/{name}')
else:
local_path = FS.get_from(image)
return local_path
@@ -23,21 +24,23 @@ class InferenceUIName():
if language == 'en':
self.advance_block_name = 'Advance Setting'
self.check_box_for_setting = [
'Use Mantra', 'Use Tuners', 'Use Controller'
'Use Mantra', 'Use Tuners', 'Use Controller', 'LAR-Gen'
]
self.diffusion_paras = 'Generation Setting'
self.mantra_paras = 'Mantra Book'
self.tuner_paras = 'Tuners'
self.control_paras = 'Controlable Generation'
self.refiner_paras = 'Refiner Setting'
self.largen_paras = 'LAR-Gen'
elif language == 'zh':
self.advance_block_name = '生成选项'
self.check_box_for_setting = ['使用咒语', '使用微调', '使用控制']
self.check_box_for_setting = ['使用咒语', '使用微调', '使用控制', 'LAR-Gen']
self.diffusion_paras = '生成参数设置'
self.mantra_paras = '咒语书'
self.tuner_paras = '微调模型'
self.control_paras = '可控生成'
self.refiner_paras = 'Refine设置'
self.largen_paras = 'LAR-Gen'
class ModelManageUIName():
@@ -216,6 +219,7 @@ class TunerUIName():
self.example_block_name = 'Examples'
self.examples = [[['Pencil Sketch Drawing'], 'a girl in a jacket'],
[['Flat 2D Art'], 'a cat']]
self.save_button = 'Save'
elif language == 'zh':
self.tuner_model = '微调模型'
@@ -231,6 +235,7 @@ class TunerUIName():
self.example_block_name = '样例'
self.examples = [[['铅笔素描'], 'a girl in a jacket'],
[['扁平2D艺术'], 'a cat']]
self.save_button = '保存'
class ControlUIName():
@@ -342,3 +347,155 @@ class ControlUIName():
self.advance_block_name = '高级设置'
self.control_scale = '控制强度'
self.example_block_name = '样例'
class LargenUIName():
def __init__(self, language='en'):
self.tasks = [
'Text_Guided_Outpainting', 'Subject_Guided_Inpainting',
'Text_Guided_Inpainting', 'Text_Subject_Guided_Inpainting'
]
if language == 'en':
self.apps = [
'Zoom Out', 'Virtual Try On', 'Inpainting (text guided)',
'Inpainting (text + reference image guided)'
]
self.dropdown_name = 'Application'
self.subject_image = 'Reference Image'
self.subject_mask = 'Reference Mask'
self.scene_image = 'Scene Image'
self.scene_mask = 'Scene Mask'
self.prompt = 'Prompt'
self.masked_image = 'Masked Image'
self.preprocess = 'Input Preprocess'
self.button_name = 'Data Preprocess'
self.direction = (
'Instruction: \n\n'
'For customized data: \n\n'
'a.1) Select the task; \n\n'
'a.1) Upload scene image; \n\n'
'a.2) Upload reference image (if needed); \n\n'
'a.3) Use the brush tool to cover the areas on the scene'
'image and reference image (to generate corresponding scene'
'mask and reference mask); \n\n'
'a.4) Click Data Preprocess button; \n\n'
'a.5) Input text prompt; \n\n'
'For example data: \n\n'
'b.1) click example row \n\n'
'Finally, click Generate button and get the output image!')
self.out_direction_label = 'Out Direction'
self.out_directions = [
'CenterAround',
'RightDown',
'LeftDown',
'RightUp',
'LeftUp',
]
elif language == 'zh':
self.apps = ['图像扩展', '虚拟试衣', '图像补全(文本引导)', '图像补全(文本+参考图引导)']
self.dropdown_name = '应用'
self.subject_image = '参考图片'
self.subject_mask = '参考图掩码'
self.scene_image = '背景图片'
self.scene_mask = '背景图掩码'
self.prompt = '提示文本'
self.masked_image = '掩码图片'
self.preprocess = '输入图片预处理'
self.button_name = '数据预处理'
self.direction = ('使用说明:\n'
'针对自定义数据:\n'
'a.1)上传背景图片(待编辑);\n'
'a.2)上传参考图片(如需要);\n'
'a.3)使用笔刷功能涂抹图像中的特定位置(获取对应的掩码图像);\n'
'a.4)点击数据预处理按钮;\n'
'a.5)输入文本提示;\n'
'针对提供的样例数据: \n'
'b.1)点击样例数据 \n'
'最后,点击生成按钮获取生成图片')
self.out_direction_label = '扩展方向'
self.out_directions = [
'中心向外',
'右下',
'左下',
'右上',
'左上',
]
self.examples = [
[
self.apps[0],
'a temple on fire',
download_image(
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex1_scene_im.png' # noqa
),
None,
None,
None,
1.0,
0.75,
'CenterAround',
1024,
1024
],
[
self.apps[1],
'',
download_image(
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex2_scene_im.jpg' # noqa
),
download_image(
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex2_scene_mask.png' # noqa
),
download_image(
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex2_subject_im.jpg' # noqa
),
download_image(
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex2_subject_mask.jpg' # noqa
),
1.0,
0.0,
'',
1024,
1024
],
[
self.apps[2],
'a blue and white porcelain',
download_image(
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex3_scene_im.png' # noqa
),
download_image(
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex3_scene_mask.png' # noqa
),
None,
None,
1.0,
0.0,
'',
1024,
1024
],
[
self.apps[3],
'a dog wearing sunglasses',
download_image(
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex4_scene_im.png' # noqa
),
download_image(
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex4_scene_mask.png' # noqa
),
download_image(
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex4_subject_im.png' # noqa
),
download_image(
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex4_subject_mask.png' # noqa
),
0.45,
0.0,
'',
1024,
1024
],
]
@@ -6,6 +6,7 @@ import gradio as gr
import numpy as np
from PIL import Image
from scepter.modules.utils.file_system import FS
from scepter.studio.inference.inference_ui.component_names import GalleryUIName
from scepter.studio.utils.uibase import UIBase
@@ -15,6 +16,9 @@ class GalleryUI(UIBase):
self.pipe_manager = pipe_manager
self.component_names = GalleryUIName(language)
self.cfg = cfg
self.work_dir = cfg.WORK_DIR
self.local_work_dir, _ = FS.map_to_local(self.work_dir)
os.makedirs(self.local_work_dir, exist_ok=True)
def create_ui(self, *args, **kwargs):
with gr.Group():
@@ -41,6 +45,7 @@ class GalleryUI(UIBase):
container=False,
autofocus=True,
elem_classes='type_row',
submit_on_enter=True,
lines=1)
with gr.Column(scale=3, min_width=0):
@@ -87,6 +92,19 @@ class GalleryUI(UIBase):
style_template,
style_negative_template,
image_seed,
largen_state,
largen_task,
largen_image_scale,
largen_tar_image,
largen_tar_mask,
largen_masked_image,
largen_ref_image,
largen_ref_mask,
largen_ref_clip,
largen_base_image,
largen_extra_sizes,
largen_bbox_yyxx,
largen_history,
show_jpeg_image=True):
if control_state and control_cond_image is None:
raise gr.Error(self.component_names.control_err1)
@@ -111,9 +129,8 @@ class GalleryUI(UIBase):
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]):
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)
@@ -163,6 +180,23 @@ class GalleryUI(UIBase):
pipeline_input['refine_guide_rescale'] = refine_guide_rescale
else:
refine_strength = 0
if largen_state:
largen_cfg = {
'largen_task': largen_task,
'largen_image_scale': largen_image_scale,
'largen_tar_image': largen_tar_image,
'largen_tar_mask': largen_tar_mask,
'largen_ref_image': largen_ref_image,
'largen_ref_mask': largen_ref_mask,
'largen_masked_image': largen_masked_image,
'largen_ref_clip': largen_ref_clip,
'largen_base_image': largen_base_image,
'largen_extra_sizes': largen_extra_sizes,
'largen_bbox_yyxx': largen_bbox_yyxx,
}
else:
largen_cfg = {}
results = current_pipeline(
pipeline_input,
num_samples=image_number,
@@ -177,7 +211,8 @@ class GalleryUI(UIBase):
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))
seed=int(image_seed),
**largen_cfg)
images = []
before_images = []
if 'images' in results:
@@ -198,27 +233,34 @@ class GalleryUI(UIBase):
if 'seed' in results:
print(results['seed'])
print(images, before_images)
largen_history.extend(images)
if len(largen_history) > 10:
largen_history = largen_history[-10:]
if show_jpeg_image:
save_list = []
for i, img in enumerate(images):
save_image = os.path.join(self.cfg.WORK_DIR,
save_image = os.path.join(self.local_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,
largen_history,
gr.update(value=largen_history),
)
def generate_image(self, *args, **kwargs):
gallery_result = self.generate_gallery(*args, **kwargs)
before_refine_panel, before_refine_gallery, output_gallery = gallery_result
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):
mantra_ui, tuner_ui, refiner_ui, control_ui, largen_ui,
**kwargs):
self.gen_inputs = [
self.prompt, mantra_ui.state, tuner_ui.state, control_ui.state,
@@ -237,12 +279,20 @@ class GalleryUI(UIBase):
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
mantra_ui.style_negative_template, diffusion_ui.image_seed,
largen_ui.state, largen_ui.task, largen_ui.image_scale,
largen_ui.tar_image, largen_ui.tar_mask, largen_ui.masked_image,
largen_ui.ref_image, largen_ui.ref_mask, largen_ui.ref_clip,
largen_ui.base_image, largen_ui.extra_sizes, largen_ui.bbox_yyxx,
largen_ui.image_history
]
self.gen_outputs = [
self.before_refine_panel, self.before_refine_gallery,
self.output_gallery
self.before_refine_panel,
self.before_refine_gallery,
self.output_gallery,
largen_ui.image_history,
largen_ui.gallery,
]
self.generate_button.click(self.generate_gallery,
@@ -0,0 +1,468 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import albumentations as A
import cv2
import gradio as gr
import numpy as np
import torch
import torchvision.transforms as T
import torchvision.transforms.functional as TF
from PIL import Image
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.model.utils.data_utils import (box2squre, expand_bbox,
get_bbox_from_mask,
pad_to_square)
from scepter.modules.utils.distribute import we
from scepter.studio.inference.inference_ui.component_names import LargenUIName
from scepter.studio.utils.uibase import UIBase
refresh_symbol = '\U0001f504' # 🔄
class LargenUI(UIBase):
def __init__(self, cfg, pipe_manager, is_debug=False, language='en'):
self.cfg = cfg
self.pipe_manager = pipe_manager
self.component_names = LargenUIName(language)
def load_annotator(self, annotator):
if annotator['device'] == 'offline':
annotator['model'] = ANNOTATORS.build(annotator['cfg'])
annotator['device'] = 'cpu'
if annotator['device'] == 'cpu':
annotator['model'] = annotator['model'].to(we.device_id)
annotator['device'] = we.device_id
return annotator
def unload_annotator(self, annotator):
if not annotator['device'] == 'offline' and not annotator[
'device'] == 'cpu':
annotator['model'] = annotator['model'].to('cpu')
annotator['device'] = 'cpu'
return annotator
def create_ui(self, *args, **kwargs):
self.state = gr.State(value=False)
self.tar_image = gr.State(value=None)
self.tar_mask = gr.State(value=None)
self.ref_image = gr.State(value=None)
self.ref_mask = gr.State(value=None)
self.ref_clip = gr.State(value=None)
self.task = gr.State(value=self.component_names.tasks[0])
self.masked_image = gr.State(value=None)
self.base_image = gr.State(value=None)
self.extra_sizes = gr.State(value=None)
self.bbox_yyxx = gr.State(value=None)
self.image_history = gr.State(value=[])
with gr.Column(visible=False) as self.tab:
with gr.Row():
self.select_app = gr.Dropdown(
label=self.component_names.dropdown_name,
choices=self.component_names.apps,
value=self.component_names.apps[0],
type='index')
with gr.Row(equal_height=True):
with gr.Column(scale=2, min_width=0):
with gr.Row():
with gr.Column(scale=1, min_width=0):
self.scene_image = gr.Image(
label=self.component_names.scene_image,
type='pil',
tool='sketch',
source='upload',
height=400,
interactive=True)
self.cache_button = gr.Button(value='Use Last Generated Image', visible=True)
with gr.Column(scale=1, min_width=0):
self.subject_image = gr.Image(
label=self.component_names.subject_image,
type='pil',
tool='sketch',
source='upload',
interactive=True,
height=400,
visible=False)
self.gallery = gr.Gallery(label='Image History', value=[], columns=1, rows=1, height=500)
self.clear_button = gr.Button(value='Clear History', visible=True)
with gr.Column(scale=1, min_width=0):
self.image_scale = gr.Slider(label='Image Strength',
minimum=0.0,
maximum=1.0,
value=1.0,
visible=False)
self.image_ratio = gr.Slider(minimum=0.5,
maximum=1.0,
value=0.75,
label='Image Resize Ratio',
visible=True)
self.out_direction = gr.Dropdown(label=self.component_names.out_direction_label,
choices=self.component_names.out_directions,
value=self.component_names.out_directions[0],
visible=True)
self.proc_button = gr.Button(value=self.component_names.button_name)
self.proc_status = gr.Markdown(value='', visible=False)
self.task_desc = gr.Markdown(self.component_names.direction, visible=True)
self.eg = gr.Column(visible=True)
def set_callbacks(self, model_manage_ui, diffusion_ui, **kwargs):
def example_data_process(select_app_id, prompt, scene_image, scene_mask, subject_image,
subject_mask, image_scale, image_ratio, out_direction,
output_height, output_width):
task = self.component_names.tasks[select_app_id]
if scene_mask is not None:
scene_mask = (scene_mask > 128).astype(np.uint8)
if subject_mask is not None:
subject_mask = (subject_mask > 128).astype(np.uint8)
if task == 'Text_Guided_Inpainting':
data = self.data_preprocess_inpaint(scene_image,
scene_mask,
None,
None,
False,
1.3,
output_height, output_width)
elif task == 'Subject_Guided_Inpainting':
data = self.data_preprocess_inpaint(scene_image,
scene_mask,
subject_image,
subject_mask,
False,
1.3,
output_height,
output_width)
elif task == 'Text_Subject_Guided_Inpainting':
data = self.data_preprocess_inpaint(scene_image,
scene_mask,
subject_image,
subject_mask,
True,
1.3,
output_height,
output_width)
elif task == 'Text_Guided_Outpainting':
data = self.data_preprocess_outpaint(scene_image,
out_direction,
image_ratio,
output_height,
output_width)
subject_image_show = None if subject_image is None else Image.fromarray(subject_image.astype(np.uint8))
return *data, gr.update(value='Data Process Succeed!', visible=True), \
gr.update(value=Image.fromarray(scene_image.astype(np.uint8))), \
gr.update(value=subject_image_show), \
task, gr.update(value=self.component_names.apps[select_app_id]), \
gr.update(value=prompt), gr.update(value=image_scale), gr.update(value=image_ratio)
gallery_ui = kwargs.pop('gallery_ui')
with self.eg:
self.scene_image_eg = gr.Image(label=self.component_names.scene_image, type='numpy', visible=False) # noqa
self.scene_mask_eg = gr.Image(label=self.component_names.scene_mask, type='numpy', image_mode='L', visible=False) # noqa
self.subject_image_eg = gr.Image(label=self.component_names.subject_image, type='numpy', visible=False) # noqa
self.subject_mask_eg = gr.Image(label=self.component_names.subject_mask, type='numpy', image_mode='L', visible=False) # noqa
self.prompt = gr.Textbox(label=self.component_names.prompt, visible=False)
self.examples = gr.Examples(
examples=self.component_names.examples,
inputs=[self.select_app, self.prompt,
self.scene_image_eg, self.scene_mask_eg,
self.subject_image_eg, self.subject_mask_eg,
self.image_scale,
self.image_ratio,
self.out_direction,
diffusion_ui.output_height,
diffusion_ui.output_width,
],
outputs=[self.tar_image,
self.tar_mask,
self.masked_image,
self.ref_image,
self.ref_mask,
self.ref_clip,
self.base_image,
self.extra_sizes,
self.bbox_yyxx,
self.proc_status,
self.scene_image,
self.subject_image,
self.task,
self.select_app,
gallery_ui.prompt,
self.image_scale,
self.image_ratio,
],
fn=example_data_process,
cache_examples=False,
run_on_click=True)
def change_app(select_app_id):
select_task = self.component_names.tasks[select_app_id]
return gr.update(visible=('Subject' in select_task)), \
gr.update(visible=('Subject' in select_task)), \
gr.update(visible=('Outpainting' in select_task)), \
gr.update(visible=('Outpainting' in select_task)), select_task
self.select_app.change(
change_app,
inputs=[self.select_app],
outputs=[
self.subject_image,
self.image_scale,
self.image_ratio,
self.out_direction,
self.task
],
queue=False
)
def read_gallery_image(gallery):
if len(gallery) == 0:
last_image = None
else:
last_image = gallery[-1]['name']
return gr.update(value=last_image)
self.cache_button.click(read_gallery_image,
inputs=[self.gallery],
outputs=[self.scene_image])
def clear_gallery(image_history, gallery):
image_history.clear()
gallery.clear()
return image_history, gallery
self.clear_button.click(
fn=clear_gallery,
inputs=[self.image_history, self.gallery],
outputs=[self.image_history, self.gallery]
)
def data_process(scene_image, subject_image, task, image_ratio, out_direction, output_height, output_width):
tar_image = scene_image['image'].convert('RGB')
tar_mask = scene_image['mask'].convert('L')
tar_image = np.asarray(tar_image)
tar_mask = np.asarray(tar_mask)
tar_mask = np.where(tar_mask > 128, 1, 0).astype(np.uint8)
if task == 'Text_Guided_Inpainting':
data = self.data_preprocess_inpaint(tar_image,
tar_mask,
None,
None,
False,
1.3,
output_height,
output_width)
elif task == 'Subject_Guided_Inpainting':
ref_image = subject_image['image'].convert('RGB')
ref_mask = subject_image['mask'].convert('L')
ref_image = np.asarray(ref_image)
ref_mask = np.asarray(ref_mask)
ref_mask = np.where(ref_mask > 128, 1, 0).astype(np.uint8)
data = self.data_preprocess_inpaint(tar_image,
tar_mask,
ref_image,
ref_mask,
False,
1.3,
output_height,
output_width)
elif task == 'Text_Subject_Guided_Inpainting':
ref_image = subject_image['image'].convert('RGB')
ref_mask = subject_image['mask'].convert('L')
ref_image = np.asarray(ref_image)
ref_mask = np.asarray(ref_mask)
ref_mask = np.where(ref_mask > 128, 1, 0).astype(np.uint8)
data = self.data_preprocess_inpaint(tar_image,
tar_mask,
ref_image,
ref_mask,
True,
1.3,
output_height,
output_width)
elif task == 'Text_Guided_Outpainting':
data = self.data_preprocess_outpaint(tar_image,
out_direction,
image_ratio,
output_height,
output_width)
return *data, gr.update(value='Data Process Succeed!', visible=True)
self.proc_button.click(data_process,
inputs=[
self.scene_image,
self.subject_image,
self.task,
self.image_ratio,
self.out_direction,
diffusion_ui.output_height,
diffusion_ui.output_width
],
outputs=[
self.tar_image,
self.tar_mask,
self.masked_image,
self.ref_image,
self.ref_mask,
self.ref_clip,
self.base_image,
self.extra_sizes,
self.bbox_yyxx,
self.proc_status,
])
def data_preprocess_inpaint(self,
tar_image,
tar_mask,
ref_image,
ref_mask,
use_rectangle_mask,
tar_crop_ratio,
output_height,
output_width):
tar_mask = np.expand_dims(tar_mask, 2).astype(np.float32)
# Zoom-In
tar_yyxx = get_bbox_from_mask(tar_mask)
tar_yyxx_crop = expand_bbox(tar_mask, tar_yyxx, ratio=tar_crop_ratio)
tar_yyxx_crop = box2squre(tar_mask, tar_yyxx_crop)
y1, y2, x1, x2 = tar_yyxx_crop
crop_tar_image = tar_image[y1:y2, x1:x2, :]
crop_tar_mask = tar_mask[y1:y2, x1:x2, :]
# 得到pad之前的HW
H1, W1 = crop_tar_image.shape[:2]
# 对tar_mask进行一些预处理
if use_rectangle_mask:
tar_bbox_yyxx = get_bbox_from_mask(crop_tar_mask)
y1, y2, x1, x2 = tar_bbox_yyxx
crop_tar_mask[y1:y2, x1:x2] = 1
crop_tar_image, pad1, pad2 = pad_to_square(crop_tar_image.astype(np.uint8), pad_value=0)
crop_tar_mask, _, _ = pad_to_square(crop_tar_mask, pad_value=0)
# 得到pad之后的HW
H2, W2 = crop_tar_image.shape[:2]
aug_tar_image = cv2.resize(crop_tar_image.astype(np.uint8), (output_width, output_height))
aug_tar_mask = cv2.resize(crop_tar_mask, (output_width, output_height))
final_tar_image = TF.to_tensor(aug_tar_image)
final_tar_image = TF.normalize(final_tar_image, mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
final_tar_mask = TF.to_tensor((aug_tar_mask > 0.5).astype(np.float32))
masked_image = final_tar_image.clone()
masked_image = masked_image * (1 - final_tar_mask)
final_tar_image = final_tar_image.unsqueeze(0)
final_tar_mask = final_tar_mask.unsqueeze(0)
masked_image = masked_image.unsqueeze(0)
if ref_image is not None and ref_mask is not None:
ref_mask = np.expand_dims(ref_mask, 2).astype(np.float32)
# background-free
ref_image = ref_image * ref_mask + np.ones_like(ref_image) * 255. * (1 - ref_mask)
ref_yyxx = get_bbox_from_mask(ref_mask)
y1, y2, x1, x2 = ref_yyxx
crop_ref_image_i = ref_image[y1:y2, x1:x2, :]
crop_ref_mask_i = ref_mask[y1:y2, x1:x2, :]
# 通过ref_expand_ratio这个值pad subject image, 从而调整clip输入图像中物体的大小
# ref_expand_ratio越大,对应物体越小
h, w = crop_ref_mask_i.shape[:2]
ref_expand_size = int(max(h, w) * 1.02)
pad_op = A.PadIfNeeded(ref_expand_size, ref_expand_size,
border_mode=cv2.BORDER_CONSTANT,
value=(255, 255, 255), mask_value=0)
out = pad_op(image=crop_ref_image_i, mask=crop_ref_mask_i)
crop_ref_image = out['image']
to_clip_input = T.Compose([
T.ToTensor(),
T.Resize((224, 224)),
T.Normalize(mean=(0.48145466, 0.4578275, 0.40821073), std=(0.26862954, 0.26130258, 0.27577711)),
])
ref_clip = to_clip_input(crop_ref_image.astype(np.uint8))
output_size = max(output_height, output_width)
ref_resize_op = A.Compose([
A.LongestMaxSize(output_size),
A.PadIfNeeded(output_size, output_size,
border_mode=cv2.BORDER_CONSTANT,
value=(255, 255, 255), mask_value=0),
])
aug_out = ref_resize_op(image=crop_ref_image_i.astype(np.uint8), mask=crop_ref_mask_i)
aug_ref_image = aug_out['image']
aug_ref_mask = aug_out['mask']
final_ref_image = TF.to_tensor(aug_ref_image)
final_ref_image = TF.normalize(final_ref_image, mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
final_ref_mask = TF.to_tensor(aug_ref_mask)
final_ref_image = final_ref_image.unsqueeze(0)
final_ref_mask = final_ref_mask.unsqueeze(0)
ref_clip = ref_clip.unsqueeze(0)
else:
final_ref_image = None
final_ref_mask = None
ref_clip = None
return final_tar_image, final_tar_mask, masked_image, final_ref_image, final_ref_mask, ref_clip, \
TF.to_tensor(tar_image), torch.LongTensor([H1, W1, H2, W2, pad1, pad2]), torch.LongTensor(tar_yyxx_crop)
def data_preprocess_outpaint(self,
tar_image,
direction,
img_ratio,
output_height,
output_width):
oh, ow = output_height, output_width
h, w = tar_image.shape[:2]
ratio = max(h/(oh*img_ratio), w/(ow*img_ratio))
ih, iw = int(h / ratio), int(w / ratio)
masked_image = np.zeros((oh, ow, 3), dtype=np.uint8)
mask = np.zeros((oh, ow, 1))
if direction in ['CenterAround', '中心向外']:
y1, x1 = (oh-ih)//2, (ow-iw)//2
elif direction in ['RightDown', '右下']:
y1, x1 = 0, 0
elif direction in ['LeftDown', '左下']:
y1, x1 = 0, ow-iw
elif direction in ['RightUp', '右上']:
y1, x1 = oh-ih, 0
elif direction in ['LeftUp', '左上']:
y1, x1 = oh-ih, ow-iw
else:
y1, x1 = 0, 0
tar_image = cv2.resize(tar_image.astype(np.uint8), (iw, ih))
masked_image[y1:y1+ih, x1:x1+iw] = tar_image
mask[y1+5:y1+ih-5, x1+5:x1+iw-5] = 1
final_tar_image = TF.to_tensor(masked_image.astype(np.uint8))
final_tar_image = TF.normalize(final_tar_image, mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
final_tar_mask = TF.to_tensor(((1.0-mask) > 0.5).astype(np.float32))
masked_image = final_tar_image.clone()
masked_image = masked_image * (1 - final_tar_mask)
final_tar_image = final_tar_image.unsqueeze(0)
final_tar_mask = final_tar_mask.unsqueeze(0)
masked_image = masked_image.unsqueeze(0)
return final_tar_image, final_tar_mask, masked_image, None, None, None, None, None, None
@@ -47,14 +47,15 @@ class MantraUI(UIBase):
def create_ui(self, *args, **kwargs):
self.state = gr.State(value=False)
with gr.Column(visible=False) as self.tab:
with gr.Row():
with gr.Column(equal_height=True, visible=False) as self.tab:
with gr.Row(scale=1):
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],
choices=self.all_styles.get(
self.default_pipeline, []),
value=None,
multiselect=True,
interactive=True)
@@ -146,11 +146,20 @@ class ModelManageUI(UIBase):
continue
model_name = f"{now_pipeline}_{module['name']}"
all_module_name[module_name] = model_name
tunner_choices = []
if now_pipeline in self.default_choices['tuners']:
tunner_choices = self.default_choices['tuners'][now_pipeline][
'choices']
else:
tunner_choices = []
custom_tunner_choices = []
custom_tunner_default = []
if now_pipeline in self.default_choices.get(
'customized_tuners', []):
custom_tunner_choices = self.default_choices[
'customized_tuners'][now_pipeline]['choices']
custom_tunner_default = self.default_choices[
'customized_tuners'][now_pipeline]['default']
if isinstance(custom_tunner_default, str):
custom_tunner_default = [custom_tunner_default]
if now_pipeline in self.default_choices[
'controllers'] and control_mode in self.default_choices[
@@ -179,9 +188,11 @@ class ModelManageUI(UIBase):
gr.Dropdown(value=all_module_name['first_stage_model']),
gr.Dropdown(value=all_module_name['cond_stage_model']),
gr.Dropdown(choices=tunner_choices, value=[]),
gr.Dropdown(choices=custom_tunner_choices,
value=custom_tunner_default),
gr.Dropdown(choices=controller_choices,
value=controller_default),
gr.Dropdown(choices=mantra_ui.all_styles[now_pipeline],
gr.Dropdown(choices=mantra_ui.all_styles.get(now_pipeline, []),
value=[]),
gr.Textbox(choices=cur_paras.NEGATIVE_PROMPT.get('VALUES', []),
value=cur_paras.NEGATIVE_PROMPT.get('DEFAULT', '')),
@@ -206,10 +217,11 @@ class ModelManageUI(UIBase):
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
tuner_ui.custom_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)
@@ -21,17 +21,21 @@ class TunerUI(UIBase):
'default']
self.default_pipeline = pipe_manager.model_level_info[
default_diffusion_model]['pipeline'][0]
self.tunner_choices = []
if self.default_pipeline in self.default_choices['tuners']:
self.tunner_choices = self.default_choices['tuners'][
self.default_pipeline]['choices']
self.tunner_default = self.default_choices['tuners'][
self.default_pipeline]['default']
else:
self.tunner_choices = []
self.custom_tuner_choices = []
if self.default_pipeline in self.default_choices.get(
'customized_tuners', []):
self.custom_tuner_choices = self.default_choices[
'customized_tuners'][self.default_pipeline]['choices']
self.tunner_default = None
self.component_names = TunerUIName(language)
self.cfg_tuners = cfg.TUNERS
self.cfg_tuners = cfg.TUNERS + cfg.CUSTOM_TUNERS
self.name_level_tuners = {}
for one_tuner in tqdm(self.cfg_tuners):
if one_tuner.BASE_MODEL not in self.name_level_tuners:
@@ -47,8 +51,8 @@ class TunerUI(UIBase):
def create_ui(self, *args, **kwargs):
self.state = gr.State(value=False)
with gr.Column(visible=False) as self.tab:
with gr.Row():
with gr.Column(equal_height=True, visible=False) as self.tab:
with gr.Row(scale=1):
with gr.Column(variant='panel', scale=1, min_width=0):
with gr.Group(visible=True):
with gr.Row(equal_height=True):
@@ -63,10 +67,16 @@ class TunerUI(UIBase):
self.custom_tuner_model = gr.Dropdown(
label=self.component_names.
custom_tuner_model,
choices=[],
choices=self.custom_tuner_choices,
value=None,
multiselect=True,
interactive=True)
self.save_button = gr.Button(
label=self.component_names.save_button,
value=self.component_names.save_button,
elem_classes='type_row',
elem_id='save_button',
visible=True)
with gr.Row(equal_height=True):
with gr.Column(scale=1):
self.tuner_type = gr.Text(
@@ -110,6 +120,7 @@ class TunerUI(UIBase):
label=self.component_names.example_block_name, open=True)
def set_callbacks(self, model_manage_ui, **kwargs):
manager = kwargs.pop('manager')
gallery_ui = kwargs.pop('gallery_ui')
with self.example_block:
gr.Examples(examples=self.component_names.examples,
@@ -121,7 +132,7 @@ class TunerUI(UIBase):
now_pipeline = diffusion_model_info['pipeline'][0]
tuner_info = {}
if tuner_model is not None and len(tuner_model) > 0:
tuner_info = self.name_level_tuners[now_pipeline].get(
tuner_info = self.name_level_tuners.get(now_pipeline, {}).get(
tuner_model[-1], {})
if tuner_info.get(
'IMAGE_PATH',
@@ -141,3 +152,45 @@ class TunerUI(UIBase):
self.tuner_example, self.tuner_prompt_example
],
queue=False)
self.custom_tuner_model.change(
tuner_model_change,
inputs=[self.custom_tuner_model, model_manage_ui.diffusion_model],
outputs=[
self.tuner_type, self.base_model, self.tuner_desc,
self.tuner_example, self.tuner_prompt_example
],
queue=False)
def save_customized_tuner(tuner_model, diffusion_model):
diffusion_model_info = self.pipe_manager.model_level_info[
diffusion_model]
now_pipeline = diffusion_model_info['pipeline'][0]
tuner_info = {}
if tuner_model is not None and len(tuner_model) > 0:
tuner_info = self.name_level_tuners.get(now_pipeline, {}).get(
tuner_model[-1], {})
if tuner_info.get(
'IMAGE_PATH',
None) and not os.path.exists(tuner_info.IMAGE_PATH):
tuner_info.IMAGE_PATH = FS.get_from(tuner_info.IMAGE_PATH)
return (gr.Tabs(selected='tuner_manager'),
gr.Text(value=tuner_info.NAME), gr.Text(value=''),
gr.Text(value=tuner_info.get('TUNER_TYPE', '')),
gr.Text(value=tuner_info.get('BASE_MODEL', '')),
gr.Text(value=tuner_info.get('DESCRIPTION', '')),
gr.Image(value=tuner_info.get('IMAGE_PATH', None)),
gr.Text(value=tuner_info.get('PROMPT_EXAMPLE', '')))
self.save_button.click(
save_customized_tuner,
inputs=[self.custom_tuner_model, model_manage_ui.diffusion_model],
outputs=[
manager.tabs, manager.tuner_manager.info_ui.tuner_name,
manager.tuner_manager.info_ui.new_name,
manager.tuner_manager.info_ui.tuner_type,
manager.tuner_manager.info_ui.base_model,
manager.tuner_manager.info_ui.tuner_desc,
manager.tuner_manager.info_ui.tuner_example,
manager.tuner_manager.info_ui.tuner_prompt_example
])
@@ -24,7 +24,6 @@ refresh_symbol = '\U0001f504' # 🔄
class CreateDatasetUI(UIBase):
def __init__(self, cfg, is_debug=False, language='en'):
self.work_dir = cfg.WORK_DIR
self.dir_list = FS.walk_dir(self.work_dir, recurse=False)
self.cache_file = {}
self.meta_dict = {}
self.dataset_list = self.load_history()
@@ -71,7 +70,8 @@ class CreateDatasetUI(UIBase):
json.dump(save_meta, open(local_path, 'w'))
return meta_file
def construct_meta(self, cursor, file_list, dataset_folder, user_name):
def construct_meta(self, cursor, file_list, dataset_folder, user_name,
login_user_name):
'''
{
"dataset_name": "xxxx",
@@ -86,11 +86,17 @@ class CreateDatasetUI(UIBase):
save_file_list = os.path.join(dataset_folder, 'file.csv')
save_file_list = self.write_file_list(file_list, save_file_list)
meta = {
'dataset_name': user_name,
'cursor': cursor,
'file_list': file_list,
'train_csv': train_csv,
'save_file_list': save_file_list
'dataset_name':
user_name if login_user_name == '' or login_user_name is None else
'_'.join([login_user_name, user_name]),
'cursor':
cursor,
'file_list':
file_list,
'train_csv':
train_csv,
'save_file_list':
save_file_list
}
self.save_meta(meta, dataset_folder)
return meta
@@ -180,7 +186,7 @@ class CreateDatasetUI(UIBase):
file_folder = None
train_list = None
hit_dir = None
raw_list = []
raw_list = {}
mac_osx = os.path.join(local_dataset_folder, '__MACOSX')
if os.path.exists(mac_osx):
res = os.popen(f"rm -rf '{mac_osx}'")
@@ -191,28 +197,45 @@ class CreateDatasetUI(UIBase):
res = res.readlines()
continue
if FS.isdir(one_dir):
sub_dir = FS.walk_dir(one_dir)
for one_s_dir in sub_dir:
if FS.isdir(one_s_dir) and one_s_dir.split(
one_dir)[1].replace('/', '') == 'images':
file_folder = one_s_dir
hit_dir = one_dir
if FS.isfile(one_s_dir) and one_s_dir.split(
one_dir)[1].replace('/', '') == 'train.csv':
train_list = one_s_dir
if file_folder is not None and train_list is not None:
break
if (one_s_dir.endswith('.jpg')
or one_s_dir.endswith('.jpeg')
or one_s_dir.endswith('.png')
or one_s_dir.endswith('.webp')):
raw_list.append(one_s_dir)
if one_dir.endswith('images') or one_dir.endswith('images/'):
file_folder = one_dir
hit_dir = one_dir
else:
sub_dir = FS.walk_dir(one_dir)
for one_s_dir in sub_dir:
if FS.isdir(one_s_dir) and one_s_dir.split(
one_dir)[1].replace('/', '') == 'images':
file_folder = one_s_dir
hit_dir = one_dir
if FS.isfile(one_s_dir) and one_s_dir.split(
one_dir)[1].replace('/', '') == 'train.csv':
train_list = one_s_dir
if file_folder is not None and train_list is not None:
break
if (one_s_dir.endswith('.jpg')
or one_s_dir.endswith('.jpeg')
or one_s_dir.endswith('.png')
or one_s_dir.endswith('.webp')):
file_name, surfix = os.path.splitext(one_s_dir)
txt_file = file_name + '.txt'
if os.path.exists(txt_file):
raw_list[one_s_dir] = txt_file
else:
raw_list[one_s_dir] = None
elif one_dir.endswith('train.csv'):
train_list = one_dir
else:
if (one_dir.endswith('.jpg') or one_dir.endswith('.jpeg')
or one_dir.endswith('.png')
or one_dir.endswith('.webp')):
raw_list.append(one_dir)
file_name, surfix = os.path.splitext(one_dir)
txt_file = file_name + '.txt'
if os.path.exists(txt_file):
raw_list[one_dir] = txt_file
else:
raw_list[one_dir] = None
if file_folder is not None and train_list is not None:
break
if file_folder is None and len(raw_list) < 1:
raise gr.Error(
"images folder or train.csv doesn't exists, or nothing exists in your zip"
@@ -222,17 +245,21 @@ class CreateDatasetUI(UIBase):
if file_folder is not None:
_ = FS.get_dir_to_local_dir(file_folder, new_file_folder)
elif len(raw_list) > 0:
raw_list = list(set(raw_list))
raw_list = [[k, v] for k, v in raw_list.items()]
for img_id, cur_image in enumerate(raw_list):
_, surfix = os.path.splitext(cur_image)
image_name, surfix = os.path.splitext(cur_image[0])
if cur_image[1] is not None and os.path.exists(cur_image[1]):
prompt = open(cur_image[1], 'r').read()
else:
prompt = image_name.split('/')[-1]
try:
os.rename(
os.path.abspath(cur_image),
f'{new_file_folder}/{get_md5(cur_image)}{surfix}')
os.path.abspath(cur_image[0]),
f'{new_file_folder}/{get_md5(cur_image[0])}{surfix}')
raw_list[img_id] = [
os.path.join('images',
f'{get_md5(cur_image)}{surfix}'),
cur_image.split('/')[-1]
f'{get_md5(cur_image[0])}{surfix}'),
prompt
]
except Exception as e:
print(e)
@@ -251,13 +278,14 @@ class CreateDatasetUI(UIBase):
res = res.readlines()
if not os.path.exists(new_train_list):
raise gr.Error(f'{str(res)}')
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
if not file_folder == hit_dir:
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
@@ -317,21 +345,25 @@ class CreateDatasetUI(UIBase):
return True, file_path
return False, file_path
def load_history(self):
def load_history(self, login_user_name=''):
dataset_list = []
self.dir_list = FS.walk_dir(self.work_dir, recurse=False)
for one_dir in self.dir_list:
if FS.isdir(one_dir):
meta_file = os.path.join(one_dir, 'meta.json')
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, multi_thread=True)
if not FS.exists(
os.path.join(local_dataset_folder, 'meta.json')):
local_dataset_folder = FS.get_dir_to_local_dir(
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
meta_data['work_dir'] = one_dir
dataset_list.append(meta_data['dataset_name'])
self.meta_dict[meta_data['dataset_name']] = meta_data
if meta_data['dataset_name'].startswith(login_user_name):
dataset_list.append(meta_data['dataset_name'])
self.meta_dict[meta_data['dataset_name']] = meta_data
return dataset_list
def create_ui(self):
@@ -396,7 +428,7 @@ class CreateDatasetUI(UIBase):
self.file_panel = file_panel
self.modify_panel = modify_panel
def set_callbacks(self, gallery_dataset, export_dataset):
def set_callbacks(self, gallery_dataset, export_dataset, manager):
def show_dataset_panel():
return (gr.Column(visible=False), gr.Column(visible=True),
gr.Column(visible=True),
@@ -419,17 +451,19 @@ class CreateDatasetUI(UIBase):
datetime.datetime.now())
return data_name
def refresh():
return gr.Dropdown(value=self.dataset_list[-1]
if len(self.dataset_list) > 0 else '',
choices=self.dataset_list)
def refresh(login_user_name):
dataset_list = self.load_history(login_user_name=login_user_name)
return gr.Dropdown(
value=dataset_list[-1] if len(dataset_list) > 0 else '',
choices=dataset_list)
self.refresh_dataset_name.click(refresh,
inputs=[manager.user_name],
outputs=[self.dataset_name],
queue=False)
def confirm_create_dataset(user_name, create_mode, file_url, file_path,
panel_state):
panel_state, login_user_name):
if user_name.strip() == '' or ' ' in user_name or '/' in user_name:
raise gr.Error(self.components_name.illegal_data_name_err1)
@@ -442,7 +476,9 @@ 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 == 3 and not file_url.strip() == '':
file_name, surfix = os.path.splitext(file_url.split('?')[0])
if 'oss' in file_url:
file_url = file_url.split('?')[0]
file_name, surfix = os.path.splitext(file_url)
save_file = os.path.join(self.work_dir, f'{user_name}{surfix}')
local_path, _ = FS.map_to_local(save_file)
res = os.popen(f"wget -c '{file_url}' -O '{local_path}'")
@@ -490,7 +526,7 @@ class CreateDatasetUI(UIBase):
cursor = 0 if len(file_list) > 0 else -1
meta = self.construct_meta(cursor, file_list, dataset_folder,
user_name)
user_name, login_user_name)
meta['local_work_dir'] = local_dataset_folder
meta['work_dir'] = dataset_folder
@@ -498,13 +534,13 @@ class CreateDatasetUI(UIBase):
self.meta_dict[meta['dataset_name']] = meta
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),
)
return (gr.Checkbox(value=True, visible=False),
gr.Dropdown(value=meta['dataset_name'],
choices=self.dataset_list),
gr.Text(value=meta['dataset_name']))
def clear_file():
return gr.Text(visible=True)
return gr.Text(visible=False)
# Click Create
self.btn_create_datasets.click(show_dataset_panel, [], [
@@ -545,9 +581,9 @@ class CreateDatasetUI(UIBase):
# Click Confirm
self.confirm_data_button.click(confirm_create_dataset, [
self.user_data_name, self.create_mode, self.file_path_url,
self.file_path, self.panel_state
], [self.panel_state, self.dataset_name],
queue=True)
self.file_path, self.panel_state, manager.user_name
], [self.panel_state, self.dataset_name, self.user_data_name],
queue=False)
def show_edit_panel(panel_state, data_name):
if panel_state:
@@ -568,7 +604,7 @@ class CreateDatasetUI(UIBase):
],
queue=False)
def modify_data_name(user_name, prev_data_name):
def modify_data_name(user_name, prev_data_name, login_user_name):
print(
f'Current file name {prev_data_name}, new file name {user_name}.'
)
@@ -600,7 +636,8 @@ class CreateDatasetUI(UIBase):
raise gr.Error(self.components_name.illegal_data_err3)
cursor = ori_meta['cursor']
meta = self.construct_meta(cursor, file_list,
dataset_folder, user_name)
dataset_folder, user_name,
login_user_name)
meta['local_work_dir'] = local_dataset_folder
meta['work_dir'] = dataset_folder
@@ -624,7 +661,10 @@ class CreateDatasetUI(UIBase):
self.modify_data_button.click(
modify_data_name,
inputs=[self.user_data_name, self.user_data_name_state],
inputs=[
self.user_data_name, self.user_data_name_state,
manager.user_name
],
outputs=[self.user_data_name_state, self.dataset_name],
queue=False)
@@ -647,3 +687,14 @@ class CreateDatasetUI(UIBase):
gallery_dataset.gallery_state
],
queue=False)
def login_user_name_change(login_user_name):
dataset_list = self.load_history(login_user_name=login_user_name)
return gr.Dropdown(
value=dataset_list[-1] if len(dataset_list) > 0 else '',
choices=dataset_list)
manager.user_name.change(login_user_name_change,
inputs=[manager.user_name],
outputs=[self.dataset_name],
queue=False)
+1 -1
View File
@@ -46,7 +46,7 @@ class PreprocessUI():
def set_callbacks(self, manager):
self.create_dataset.set_callbacks(self.dataset_gallery,
self.export_dataset)
self.export_dataset, manager)
self.dataset_gallery.set_callbacks(self.create_dataset)
self.export_dataset.set_callbacks(self.create_dataset, manager)
@@ -8,6 +8,7 @@ import numpy as np
import torch
from scepter.modules.solver.hooks.checkpoint import CheckpointHook
from scepter.modules.solver.hooks.data_probe import ProbeDataHook
from scepter.modules.solver.registry import SOLVERS
from scepter.modules.utils.config import Config
from scepter.modules.utils.distribute import we
@@ -27,6 +28,7 @@ def run_task(cfg):
solver.set_up_pre()
solver.set_up()
ori_steps = solver.max_steps
if 'train' in solver.datas:
dataset = solver.datas['train'].dataset
if hasattr(dataset, 'real_number'):
@@ -48,6 +50,19 @@ def run_task(cfg):
f'checkpoint save interval is changed from {ori_interval} '
f'to {hook.interval} according to the setting epoches '
f'interval {ori_interval}')
if 'eval' in solver.hooks_dict:
for hook in solver.hooks_dict['eval']:
if isinstance(hook, ProbeDataHook):
ori_interval = hook.prob_interval
hook.prob_interval = int(hook.prob_interval *
solver.max_steps /
cfg.SOLVER.MAX_EPOCHS)
std_logger.info(
f'prob interval is changed from {ori_interval} '
f'to {hook.prob_interval} according to the setting epoches '
f'interval {ori_interval}')
solver.eval_interval = hook.prob_interval
# size 为无限的时候,使用默认值。
solver.solve()
@@ -0,0 +1,6 @@
# -*- coding: utf-8 -*-
import time
for i in range(180):
time.sleep(1)
print('sleep', i)
@@ -0,0 +1,336 @@
# -*- coding: utf-8 -*-
import json
import os
import subprocess
import threading
import time
import psutil
import torch
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.logger import as_time
class TaskStatus():
def __init__(self):
self.pid = -1
self.retcode = -999
self.error_msg = None
self.out_msg = None
def __repr__(self):
return f'Process {self.pid} retcode {self.retcode}, error msg: {self.error_msg}.'
def kill_job(pid):
try:
current_process = psutil.Process(pid)
children = current_process.children(recursive=True)
for child in children:
child.terminate()
current_process.terminate()
except Exception:
try:
os.system(f'kill -9 {pid}')
except Exception:
pass
class Trainer():
def __init__(self, run_script, status_message):
self.run_script = run_script
self.status_message = status_message
self.proc = None
def __call__(self, task_name):
torch.cuda.empty_cache()
cmd = f'PYTHONPATH=. python {self.run_script} ' \
f'--cfg={task_name}/train.yaml'
# cmd = [f"python {self.run_script}"]
print(cmd)
try:
# self.proc = subprocess.Popen(cmd, shell=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True)
self.proc = subprocess.Popen(cmd, shell=True)
self.status_message.pid = self.proc.pid
self.status_message.retcode = self.proc.wait(
) # self.proc.wait(3600*24*14)
# self.status_message.error_msg = self.proc.stderr.read()
self.status_message.error_msg = ''
except Exception:
self.status_message.retcode = -2
# self.status_message.error_msg = self.proc.stderr.read()
self.status_message.error_msg = ''
def terminate(self):
print(f'Terminate {self.proc.pid} ...')
kill_job(self.proc.pid)
self.proc.terminate()
print(f'Terminate {self.proc.pid} success.')
class TrainManager():
'''
Manage all the submitted training tasks and control the status of tasks.
When instance this class, we should start the threading, which used to get the task name
and start the training task.
'''
def __init__(self, run_script, work_dir):
self.task_queue = []
self.runing_tasks = {}
self.run_script = run_script
self.work_dir = work_dir
def task_dispatch():
while True:
all_running_tasks = [k for k in self.runing_tasks]
for k in all_running_tasks:
now_task = self.runing_tasks[k]
# print(now_task["train_status"].pid, psutil.pid_exists(now_task["train_status"].pid))
if (not psutil.pid_exists(now_task['train_status'].pid) and
not now_task['train_status'].retcode in (-1, 0)):
now_task['train_status'].retcode = -2
# print(now_task["train_status"])
if not now_task['train_status'].retcode == -999:
status_file = os.path.join(self.work_dir, k,
'status.json')
if FS.exists(status_file):
with FS.get_from(status_file) as local_status:
task_status = json.load(open(
local_status, 'r'))
code = now_task['train_status'].retcode
task_status['code'] = code
task_status['end_time'] = time.time()
duration = task_status['end_time'] - task_status[
'start_time']
if code == 0:
message = f'''
Training completed! \n
Save in [ {k} ] \n
Detail log please export the log. \n
Take time [ {duration:.4f}s ] \n
{self.check_memory()}
'''
task_status['msg'] = message
task_status['status'] = 'success'
else:
err_msg = now_task['train_status'].error_msg
message = f'''
Training failed! \n
Error msg: {err_msg[-1000:]} \n
Take time [ {duration:.4f}s ] \n
{self.check_memory()}
'''
task_status['msg'] = message
task_status['status'] = 'failed'
with FS.put_to(status_file) as local_path:
json.dump(task_status,
open(local_path, 'w'),
ensure_ascii=False)
train_ins = now_task['train_ins']
train_ins.terminate()
kill_job(now_task['train_status'].pid)
self.runing_tasks.pop(k)
# print(len(self.task_queue))
if len(self.task_queue) > 0 and len(self.runing_tasks) == 0:
task_name = self.task_queue.pop(0)
print(f'start task {task_name}')
status_message = TaskStatus()
train_ins = Trainer(self.run_script, status_message)
train_thread = threading.Thread(target=train_ins,
args=(os.path.join(
self.work_dir,
task_name), ))
train_thread.start()
time.sleep(10)
self.runing_tasks[task_name] = {
'train_ins': train_ins,
'train_status': status_message
}
status_file = os.path.join(self.work_dir, task_name,
'status.json')
with FS.get_from(status_file) as local_status:
task_status = json.load(open(local_status, 'r'))
task_status['status'] = 'running'
for _ in range(10):
if status_message.pid > 0:
task_status['pid'] = status_message.pid
break
time.sleep(10)
task_status['update_time'] = time.time()
with FS.put_to(status_file) as local_path:
json.dump(task_status,
open(local_path, 'w'),
ensure_ascii=False)
time.sleep(5)
self.task_manage = threading.Thread(target=task_dispatch, daemon=True)
self.task_manage.start()
def check_memory(self):
# Check Cuda Memory
mem_msg = ''
if torch.cuda.is_available():
for device_id in range(torch.cuda.device_count()):
free_mem, total_mem = torch.cuda.mem_get_info(device_id)
free_mem = free_mem / (1024**3)
total_mem = total_mem / (1024**3)
mem_msg += f'GPU {device_id}: free mem {free_mem:.3f}G, total mem {total_mem:.3f}G \n'
else:
mem_msg += 'GPU is not available!'
return mem_msg
def get_status(self, task_name):
if task_name is None:
return ''
status_file = os.path.join(self.work_dir, task_name, 'status.json')
if not FS.exists(status_file):
return ''
with FS.get_from(status_file) as local_status:
task_status = json.load(open(local_status, 'r'))
return task_status['status']
def get_log(self, task_name):
if task_name is None:
return ''
status_file = os.path.join(self.work_dir, task_name, 'status.json')
if not FS.exists(status_file):
return ''
with FS.get_from(status_file) as local_status:
task_status = json.load(open(local_status, 'r'))
time.sleep(2)
log_msg = ''
if task_status['status'] == 'queue':
start_time = task_status['start_time']
if task_name in self.task_queue:
log_msg += f'Task status: Queuing. Still have {self.task_queue.index(task_name) + 1} tasks.\n'
else:
log_msg += 'Task status: Queuing. Still have 0 tasks.\n'
log_msg += f'Have waited for {as_time(time.time() - start_time)}\n\n'
log_msg += f'Memory status: {self.check_memory()}\n'
elif task_status['status'] == 'running':
start_time = task_status['start_time']
update_time = task_status['update_time']
log_msg += (
f'Task status: Running. Have run for {as_time(time.time() - update_time)} after waiting'
f'for {as_time(update_time - start_time)}.\n\n')
std_log = os.path.join(self.work_dir, task_name, 'std_log.txt')
out_log = os.path.join(self.work_dir, task_name,
'output_std_log.txt')
output_msg = []
if os.path.exists(std_log):
fp_w = open(out_log, 'w')
output_msg.append(f'Model {task_name} Start Training....')
with open(std_log, 'r') as f:
for line in f:
line = line.strip()
if line == '':
continue
if 'OSS_AK' in line or 'OSS_AK'.lower() in line:
continue
if 'OSS_SK' in line or 'OSS_SK'.lower() in line:
continue
if 'Restored' in line or 'Restored'.lower() in line:
continue
# if 'Stage' in line:
output_msg.append(line)
fp_w.write(f'{line}\n')
fp_w.close()
else:
fp_w = open(out_log, 'w')
fp_w.close()
log_msg += f'Memory status: {self.check_memory()}\n\n'
log_msg += 'Recent output as follows:\n\n'
log_msg += '\n\n'.join(output_msg)
elif task_status['status'] == 'success':
start_time = task_status['start_time']
update_time = task_status['update_time']
end_time = task_status['end_time']
output_msg = []
out_log = os.path.join(self.work_dir, task_name,
'output_std_log.txt')
std_log = os.path.join(self.work_dir, task_name, 'std_log.txt')
if os.path.exists(std_log):
fp_w = open(out_log, 'w')
with open(std_log, 'r') as f:
for line in f:
line = line.strip()
if line == '':
continue
if 'OSS_AK' in line or 'OSS_AK'.lower() in line:
continue
if 'OSS_SK' in line or 'OSS_SK'.lower() in line:
continue
if 'Restored' in line or 'Restored'.lower() in line:
continue
output_msg.append(line)
fp_w.write(f'{line}\n')
fp_w.close()
else:
fp_w = open(out_log, 'w')
fp_w.close()
log_msg += (
f'Task status: Success. Have run for {as_time(end_time - update_time)} after waiting'
f'for {as_time(update_time - start_time)}.\n\n')
log_msg += f'Memory status: {self.check_memory()}\n\n'
log_msg += '\n\n'.join(output_msg)
elif task_status['status'] == 'failed':
start_time = task_status['start_time']
update_time = task_status['update_time']
end_time = task_status['end_time']
err_msg = task_status['msg']
log_msg += (
f'Task status: Failed. Have run for {as_time(end_time - update_time)} after waiting'
f'for {as_time(update_time - start_time)}.\n\n')
log_msg += f'Memory status: {self.check_memory()}\n\n'
log_msg += 'The error msg is as follows:\n\n'
log_msg += f'{err_msg}'
return log_msg
def start_task(self, task_name):
if task_name is None:
return
self.task_queue.append(task_name)
task_status = {
'status': 'queue',
'position': len(self.task_queue),
'start_time': time.time(),
'end_time': time.time()
}
status_file = os.path.join(self.work_dir, task_name, 'status.json')
with FS.put_to(status_file) as local_path:
json.dump(task_status, open(local_path, 'w'), ensure_ascii=False)
def stop_task(self, task_name):
if task_name in self.runing_tasks:
task_info = self.runing_tasks.pop(task_name)
train_status = task_info['train_status']
train_ins = task_info['train_ins']
kill_job(train_status.pid)
train_ins.terminate()
elif task_name in self.task_queue:
self.task_queue.remove(task_name)
else:
pass
# modify the task's status met
def __del__(self):
for task_name in self.task_queue:
status_file = os.path.join(self.work_dir, task_name, 'status.json')
with FS.get_from(status_file) as local_status:
task_status = json.load(open(local_status, 'r'))
task_status['status'] = 'failed'
task_status['msg'] = 'Main process has been killed.'
with FS.put_to(status_file) as local_path:
json.dump(task_status,
open(local_path, 'w'),
ensure_ascii=False)
pass
if __name__ == '__main__':
train_ins = TrainManager('', '')
for i in range(5):
train_ins.start_task(f'{i}')
time.sleep(300)
+8 -8
View File
@@ -7,7 +7,7 @@ import gradio as gr
import scepter
from scepter.modules.utils.config import Config
from scepter.modules.utils.file_system import FS
from scepter.studio.self_train.self_train_ui.inference_ui import InferenceUI
from scepter.studio.self_train.self_train_ui.model_ui import ModelUI
from scepter.studio.self_train.self_train_ui.trainer_ui import TrainerUI
from scepter.studio.self_train.utils.config_parser import get_all_config
from scepter.studio.utils.env import init_env
@@ -34,20 +34,20 @@ class SelfTrainUI():
BASE_CFG_VALUE,
is_debug=is_debug,
language=language)
self.inference_ui = InferenceUI(cfg_general,
BASE_CFG_VALUE,
is_debug=is_debug,
language=language)
self.model_ui = ModelUI(cfg_general,
BASE_CFG_VALUE,
is_debug=is_debug,
language=language)
def create_ui(self):
with gr.Row():
self.trainer_ui.create_ui()
with gr.Row():
self.inference_ui.create_ui()
self.model_ui.create_ui()
def set_callbacks(self, manager):
self.trainer_ui.set_callbacks(self.inference_ui)
self.inference_ui.set_callbacks(self.trainer_ui, manager)
self.trainer_ui.set_callbacks(self.model_ui, manager)
self.model_ui.set_callbacks(self.trainer_ui, manager)
if __name__ == '__main__':
@@ -1,11 +1,12 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
# For dataset manager
class InferenceUIName():
class ModelUIName():
def __init__(self, language='en'):
if language == 'en':
self.output_model_block = 'Model Output'
self.output_model_name = 'Output Model Name'
self.output_ckpt_name = 'Output Ckpt Name'
self.test_prompt = 'Test Prompt'
self.test_prefix = 'Test Prefix'
self.test_n_prompt = 'Negative Prompt'
@@ -21,15 +22,23 @@ class InferenceUIName():
self.extra_model_gbtn = 'Add Model'
self.refresh_model_gbtn = 'Refresh Model'
self.go_to_inference = 'Go to inference'
self.btn_export_log = 'Export Log'
self.export_file = 'Log File'
self.log_block = 'Training Log...'
self.gallery_block = 'Gallery Log...'
self.eval_gallery = 'Eval Gallery'
# Error or Warning
self.inference_err1 = 'Inference failed, please try again.'
self.inference_err2 = 'Test prompt is empty.'
self.inference_err3 = "Doesn't surpport this base model"
self.inference_err4 = "This model maybe not finish training, because model doesn't exist."
self.model_err3 = "Doesn't surpport this base model"
self.model_err4 = "This model maybe not finish training, because model doesn't exist."
self.model_err5 = "Model {} doesn't exist."
self.training_warn1 = 'No log message util now.'
elif language == 'zh':
self.output_model_block = '模型产出'
self.output_model_name = '产出模型名称'
self.output_model_name = '产出名称'
self.output_ckpt_name = '产出检查点名称'
self.test_prompt = '测试提示词'
self.test_prefix = '测试前缀'
self.test_n_prompt = '负向提示词'
@@ -44,12 +53,20 @@ class InferenceUIName():
self.extra_model_gtxt = '额外模型'
self.extra_model_gbtn = '添加模型'
self.refresh_model_gbtn = '刷新模型'
self.btn_export_log = '导出日志'
self.export_file = '日志文件'
self.log_block = '训练日志...'
self.training_button = '开始训练'
self.gallery_block = '图像日志...'
self.eval_gallery = '评测图像'
# Error or Warning
self.inference_err1 = '推理失败,请重试。'
self.inference_err2 = '测试提示词为空。'
self.inference_err3 = '不支持的基础模型'
self.model_err3 = '不支持的基础模型'
self.go_to_inference = '使用模型'
self.inference_err4 = '模型可能没有训练完成或者模型不存在'
self.model_err4 = '模型可能没有训练完成或者模型不存在'
self.model_err5 = '模型{}不存在'
self.training_warn1 = '暂时没有日志文件;任务启动中或失败!'
class TrainerUIName():
@@ -80,7 +97,8 @@ class TrainerUIName():
self.base_model = 'Base Model'
self.tuner_name = 'Fine-tuning Method'
self.base_model_revision = 'Model Version Number'
self.resolution = 'Resolution'
self.resolution_height = 'Resolution Height'
self.resolution_width = 'Resolution Width'
self.train_epoch = 'Number of Training Epochs'
self.learning_rate = 'Learning Rate'
self.save_interval = 'Save Interval'
@@ -89,9 +107,8 @@ class TrainerUIName():
self.replace_keywords = 'Trigger Keywords'
self.work_name = 'Save Model Name (refresh to get a random value)'
self.push_to_hub = 'Push to hub'
self.log_block = 'Training Log...'
self.training_button = 'Start Training'
self.eval_prompts = 'Eval Prompts'
# Error or Warning
self.training_err1 = 'CUDA is unavailable.'
self.training_err2 = 'Currently insufficient VRAM, training failed!'
@@ -121,7 +138,8 @@ class TrainerUIName():
self.base_model = '基础模型'
self.tuner_name = '微调方法'
self.base_model_revision = '模型版本号'
self.resolution = '分辨率'
self.resolution_height = '训练高度'
self.resolution_width = '训练宽度'
self.train_epoch = '训练轮数'
self.learning_rate = '学习率'
self.save_interval = '存储间隔'
@@ -130,8 +148,7 @@ class TrainerUIName():
self.replace_keywords = '触发关键词'
self.work_name = '保存模型名称(刷新获得随机值)'
self.push_to_hub = '推送魔搭社区'
self.log_block = '训练日志...'
self.training_button = '开始训练'
self.eval_prompts = '评测文本'
# Error or Warning
self.training_err1 = 'CUDA不可用.'
self.training_err2 = '目前显存不足,训练失败!'
@@ -1,145 +0,0 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os
import gradio as gr
from scepter.modules.utils.config import Config
from scepter.modules.utils.file_system import FS
from scepter.studio.self_train.self_train_ui.component_names import \
InferenceUIName
from scepter.studio.self_train.utils.config_parser import (
get_base_model_list, get_inference_para_by_model_version)
from scepter.studio.utils.uibase import UIBase
class InferenceUI(UIBase):
def __init__(self, cfg, all_cfg_value, is_debug=False, language='en'):
self.BASE_CFG_VALUE = all_cfg_value
self.language = language
self.base_model_info = get_base_model_list(self.BASE_CFG_VALUE)
self.work_dir, _ = FS.map_to_local(cfg.WORK_DIR)
os.makedirs(self.work_dir, exist_ok=True)
self.model_list = []
# self.model_list.extend(self.base_model_info.get('model_choices', []))
have_model_list = []
if not self.work_dir.endswith('/'):
self.work_dir += '/'
for one_dir in FS.walk_dir(self.work_dir):
if one_dir.startswith(self.work_dir):
one_dir = one_dir[len(self.work_dir):]
if not os.path.isdir(os.path.join(self.work_dir, one_dir)):
continue
if len(one_dir.split('/')) > 1:
continue
if '@' in one_dir and os.path.exists(
os.path.join(self.work_dir, one_dir, 'checkpoint.pth')):
if len(one_dir.split('@')) > 4:
have_model_list.append([one_dir, one_dir.split('@')[-1]])
have_model_list.sort(key=lambda x: -int(x[-1][:-3]))
self.model_list.extend([v[0].split('/')[-1]
for v in have_model_list][:50])
self.infer_para_data = get_inference_para_by_model_version(
self.BASE_CFG_VALUE, self.base_model_info.get('model_name', []),
self.base_model_info.get('version_name', []))
self.is_debug = is_debug
self.component_names = InferenceUIName(language)
def create_ui(self, *args, **kwargs):
with gr.Box():
gr.Markdown(self.component_names.output_model_block)
with gr.Row():
with gr.Column(scale=2, min_width=0):
self.output_model_name = gr.Dropdown(
label=self.component_names.output_model_name,
choices=self.model_list,
value=self.base_model_info.get('model_default', ''),
interactive=True)
with gr.Column(scale=1, min_width=0):
self.refresh_model_gbtn = gr.Button(
self.component_names.refresh_model_gbtn)
with gr.Row():
with gr.Column(scale=2, min_width=0):
self.extra_model_gtxt = gr.Text(
label=self.component_names.extra_model_gtxt,
show_label=False,
placeholder='Add Extra Model')
with gr.Column(scale=1, min_width=0):
self.extra_model_gbtn = gr.Button(
self.component_names.extra_model_gbtn)
with gr.Row():
with gr.Column(scale=1, min_width=0):
self.go_to_inferece_btn = gr.Button(
self.component_names.go_to_inference)
def set_callbacks(self, trainer_ui, manager):
self.manager = manager
def add_model(model_name):
if model_name not in self.model_list:
self.model_list.append(model_name)
return '', gr.Dropdown(choices=self.model_list)
def refresh_model():
return gr.Dropdown(choices=self.model_list)
self.extra_model_gbtn.click(
fn=add_model,
inputs=[self.extra_model_gtxt],
outputs=[self.extra_model_gtxt, self.output_model_name],
queue=False)
self.refresh_model_gbtn.click(fn=refresh_model,
inputs=[],
outputs=[self.output_model_name],
queue=False)
def go_to_inferece(output_model):
output_model_path = os.path.join(self.work_dir, output_model)
_, _, base_model, _, resolution, _ = output_model_path.split('@')
tuner_cfg = Config(cfg_dict={}, load=False)
tuner_cfg.NAME = output_model
tuner_cfg.NAME_ZH = output_model
tuner_cfg.BASE_MODEL = base_model
model_path = os.path.join(output_model_path, 'checkpoint.pth')
if not os.path.exists(model_path):
gr.Error(self.component_names.inference_err4)
tuner_cfg.MODEL_PATH = model_path
self.manager.inference.model_manage_ui.pipe_manager.register_tuner(
tuner_cfg,
name=tuner_cfg.NAME_ZH
if self.language == 'zh' else tuner_cfg.NAME,
is_customized=True)
pipeline_level_modules = self.manager.inference.model_manage_ui.pipe_manager.pipeline_level_modules
if tuner_cfg.BASE_MODEL not in pipeline_level_modules:
gr.Error(self.component_names.inference_err3 +
tuner_cfg.BASE_MODEL)
pipeline_ins = pipeline_level_modules[tuner_cfg.BASE_MODEL]
diffusion_model = f"{tuner_cfg.BASE_MODEL}_{pipeline_ins.diffusion_model['name']}"
default_choices = self.manager.inference.model_manage_ui.pipe_manager.module_level_choices
if 'customized_tuners' in default_choices and tuner_cfg.BASE_MODEL in default_choices[
'customized_tuners']:
tunner_choices = default_choices['customized_tuners'][
tuner_cfg.BASE_MODEL]['choices']
tunner_default = default_choices['customized_tuners'][
tuner_cfg.BASE_MODEL]['default']
if not isinstance(tunner_default, list):
tunner_default = [tunner_default]
else:
tunner_choices = []
tunner_default = ''
return (gr.Tabs(selected='inference'),
gr.Dropdown(choices=tunner_choices, value=tunner_default),
gr.Dropdown(value=diffusion_model),
gr.Tabs(selected='tuner_ui'))
self.go_to_inferece_btn.click(
go_to_inferece,
inputs=[self.output_model_name],
outputs=[
manager.tabs, manager.inference.tuner_ui.custom_tuner_model,
manager.inference.model_manage_ui.diffusion_model,
manager.inference.setting_tab
],
queue=False)
@@ -0,0 +1,459 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import json
import os
import queue
import time
import gradio as gr
import yaml
from scepter.modules.utils.config import Config
from scepter.modules.utils.file_system import FS
from scepter.studio.self_train.self_train_ui.component_names import ModelUIName
from scepter.studio.self_train.utils.config_parser import get_base_model_list
from scepter.studio.utils.uibase import UIBase
refresh_symbol = '\U0001f504' # 🔄
delete_symbol = '\U0001f5d1' # 🗑️
add_symbol = '\U00002795' # ➕
confirm_symbol = '\U00002714' # ✔️
class ModelUI(UIBase):
def __init__(self, cfg, all_cfg_value, is_debug=False, language='en'):
self.BASE_CFG_VALUE = all_cfg_value
self.language = language
self.base_model_info = get_base_model_list(self.BASE_CFG_VALUE)
self.work_dir, _ = FS.map_to_local(cfg.WORK_DIR)
os.makedirs(self.work_dir, exist_ok=True)
self.default_model_name = 'step-last'
self.old_default_model_name = 'checkpoint.pth'
self.delete_folder_queue = queue.Queue()
self.model_list = []
# self.model_list.extend(self.base_model_info.get('model_choices', []))
have_model_list = []
if not self.work_dir.endswith('/'):
self.work_dir += '/'
for one_dir in FS.walk_dir(self.work_dir):
if one_dir.startswith(self.work_dir):
one_dir = one_dir[len(self.work_dir):]
if not os.path.isdir(os.path.join(self.work_dir, one_dir)):
continue
if len(one_dir.split('/')) > 1:
continue
if ('@' in one_dir
and (os.path.exists(
os.path.join(self.work_dir, one_dir, 'checkpoints',
self.default_model_name))
or os.path.exists(
os.path.join(self.work_dir, one_dir,
self.old_default_model_name)))):
if len(one_dir.split('@')) > 4:
have_model_list.append([one_dir, one_dir.split('@')[-1]])
have_model_list.sort(key=lambda x: -int(x[-1][:-3]))
self.model_list.extend([v[0].split('/')[-1]
for v in have_model_list][:50])
self.is_debug = is_debug
self.component_names = ModelUIName(language)
def get_ckpt_list(self, output_model):
all_ckpt_list = []
if output_model is None or output_model == '':
return all_ckpt_list
output_model_path = os.path.join(self.work_dir, output_model)
all_ckpt_path = os.path.join(output_model_path, 'checkpoints')
if os.path.exists(all_ckpt_path):
for name in os.listdir(all_ckpt_path):
if name == self.default_model_name:
continue
path = os.path.join(all_ckpt_path, name)
if os.path.isdir(path):
all_ckpt_list.append(name)
all_ckpt_list = sorted(all_ckpt_list,
key=lambda x: int(x.split('-')[-1]))
if os.path.exists(os.path.join(all_ckpt_path,
self.default_model_name)):
all_ckpt_list.append(self.default_model_name)
return all_ckpt_list
def get_gallery_list(self, model_name, ckpt_name):
ckpt_probe_dir = os.path.join(self.work_dir, model_name, 'eval_probe',
ckpt_name, 'image')
all_gallery_list = []
if os.path.exists(ckpt_probe_dir):
for name in os.listdir(ckpt_probe_dir):
path = os.path.join(ckpt_probe_dir, name)
all_gallery_list.append(path)
return all_gallery_list
def create_ui(self, *args, **kwargs):
with gr.Box():
gr.Markdown(self.component_names.output_model_block)
with gr.Row(variant='panel', equal_height=True):
with gr.Column(scale=7, min_width=0):
self.output_model_name = gr.Dropdown(
label=self.component_names.output_model_name,
choices=self.model_list,
value=self.base_model_info.get('model_default', ''),
show_label=False,
container=False,
interactive=True)
with gr.Column(scale=2, min_width=0):
self.output_ckpt_name = gr.Dropdown(
label=self.component_names.output_ckpt_name,
value='',
choices=[],
show_label=False,
container=False,
interactive=True)
with gr.Column(scale=1, min_width=0):
self.refresh_model_gbtn = gr.Button(value=refresh_symbol)
with gr.Column(scale=1, min_width=0):
self.add_model_gbtn = gr.Button(value=add_symbol)
with gr.Column(scale=1, min_width=0):
self.delete_model_gbtn = gr.Button(value=delete_symbol)
with gr.Row(variant='panel', equal_height=True,
visible=False) as self.extra_model_panel:
with gr.Column(scale=4, min_width=0):
self.extra_model_txt = gr.Text(
label=self.component_names.extra_model_gtxt,
show_label=False,
container=False,
placeholder='Add Extra Model')
with gr.Column(scale=1, min_width=0):
self.confirm_add = gr.Button(value=confirm_symbol)
with gr.Row(variant='panel', equal_height=True):
with gr.Column(scale=3, min_width=0):
with gr.Box():
self.log_message = gr.Text(
placeholder='Please select model'
'or press export button.',
autoscroll=True,
lines=10,
label=self.component_names.log_block)
with gr.Column(scale=1, min_width=0,
visible=False) as self.export_log_panel:
# with gr.Column(scale=1, min_width=0):
self.export_log = gr.Button(
value=self.component_names.btn_export_log)
# with gr.Column(scale=1, min_width=0):
self.export_url = gr.File(
label=self.component_names.export_file,
visible=False,
value=None,
interactive=False,
show_label=True)
with gr.Row(variant='panel', equal_height=True):
with gr.Accordion(label=self.component_names.gallery_block,
open=True):
self.eval_gallery = gr.Gallery(
label=self.component_names.eval_gallery,
value=[],
preview=True,
selected_index=None)
with gr.Row():
with gr.Column(scale=1, min_width=0):
self.go_to_inferece_btn = gr.Button(
self.component_names.go_to_inference)
def set_callbacks(self, trainer_ui, manager):
self.manager = manager
def model_name_change(model_name):
if model_name is None:
return '', gr.Column(), '', []
message = trainer_ui.trainer_ins.get_log(model_name)
status = trainer_ui.trainer_ins.get_status(model_name)
ckpt_list = self.get_ckpt_list(model_name)
ckpt_value = ckpt_list[-1] if len(ckpt_list) > 0 else ''
return (message, gr.Column(visible=status in ('running',
'success')),
gr.Dropdown(choices=ckpt_list, value=ckpt_value))
self.output_model_name.change(fn=model_name_change,
inputs=[self.output_model_name],
outputs=[
self.log_message,
self.export_log_panel,
self.output_ckpt_name
],
queue=False)
def ckpt_name_change(model_name, ckpt_value):
if ckpt_value is not None and len(ckpt_value) > 0:
gallery_value = self.get_gallery_list(model_name, ckpt_value)
else:
gallery_value = []
if len(gallery_value) > 0:
return gr.Gallery(value=gallery_value,
preview=True,
selected_index=0)
else:
return gr.Gallery(value=gallery_value,
preview=True,
selected_index=None)
self.output_ckpt_name.change(
fn=ckpt_name_change,
inputs=[self.output_model_name, self.output_ckpt_name],
outputs=[self.eval_gallery],
queue=False)
def add_model():
return gr.Row(visible=True)
self.add_model_gbtn.click(fn=add_model,
inputs=[],
outputs=[self.extra_model_panel],
queue=False)
def confirm_add(model_name):
model_folder = os.path.join(self.work_dir, model_name)
have_model = os.path.exists(model_folder)
if not have_model:
gr.Error(self.component_names.model_err5.format(model_name))
if model_name not in self.model_list and have_model:
self.model_list.append(model_name)
return gr.Row(visible=False), gr.Dropdown(choices=self.model_list,
value=model_name)
self.confirm_add.click(
fn=confirm_add,
inputs=[self.extra_model_txt],
outputs=[self.extra_model_panel, self.output_model_name],
queue=False)
def refresh_model(model_name):
if not self.delete_folder_queue.empty():
try:
del_folder = self.delete_folder_queue.get_nowait()
if os.path.exists(del_folder):
os.system(f'rm -rf {del_folder}')
self.delete_folder_queue.put_nowait(del_folder)
except Exception:
pass
message = trainer_ui.trainer_ins.get_log(model_name)
status = trainer_ui.trainer_ins.get_status(model_name)
ckpt_list = self.get_ckpt_list(model_name)
ckpt_value = ckpt_list[-1] if len(ckpt_list) > 0 else ''
ret_gallery = ckpt_name_change(model_name, ckpt_value)
return (message, gr.Column(visible=status in ('running',
'success')),
gr.Dropdown(choices=ckpt_list,
value=ckpt_value), ret_gallery)
self.refresh_model_gbtn.click(
fn=refresh_model,
inputs=[self.output_model_name],
outputs=[
self.log_message,
self.export_log_panel,
# self.output_model_name,
self.output_ckpt_name,
self.eval_gallery
],
queue=False)
def delete_model(model_name):
index = 0
trainer_ui.trainer_ins.stop_task(model_name)
if model_name in self.model_list:
index = self.model_list.index(model_name)
self.model_list.remove(model_name)
folder = os.path.join(self.work_dir, model_name)
for _ in range(1):
if os.path.exists(folder):
try:
os.system(f'rm -rf {folder}')
except Exception:
time.sleep(2)
else:
break
self.delete_folder_queue.put_nowait(folder)
if index <= len(self.model_list) - 1:
model_name = self.model_list[index]
elif len(self.model_list) > 0:
model_name = self.model_list[0]
else:
model_name = None
return gr.Dropdown(choices=self.model_list, value=model_name)
self.delete_model_gbtn.click(fn=delete_model,
inputs=[self.output_model_name],
outputs=[self.output_model_name])
def go_to_inferece(output_model, output_ckpt_name):
params_path = os.path.join(self.work_dir, output_model,
'params.json')
if os.path.exists(params_path):
params_info = json.loads(open(params_path).read())
assert params_info['work_name'] == output_model
# base_model = params_info['base_model']
base_model_revision = params_info['base_model_revision']
tuner_name = params_info['tuner_name']
model_path = os.path.join(self.work_dir, output_model,
'checkpoints', output_ckpt_name)
eval_prompts = params_info[
'eval_prompts'] if 'eval_prompts' in params_info else []
image_dir = os.path.join(self.work_dir, output_model,
'eval_probe', output_ckpt_name,
'image')
image_path = [
os.path.join(image_dir, name)
for name in os.listdir(image_dir)
]
else:
_, base_model, base_model_revision, tuner_name, _ = output_model.split(
'@', 4)
model_path = os.path.join(self.work_dir, output_model,
self.old_default_model_name)
output_ckpt_name = self.old_default_model_name
image_path = []
eval_prompts = []
params_info = {}
if isinstance(image_path, list):
image_path = image_path[0] if len(image_path) > 0 else None
if isinstance(eval_prompts, list):
eval_prompts = eval_prompts[0] if len(eval_prompts) > 0 else ''
cfg_file = os.path.join(self.work_dir, output_model,
f'meta_{output_ckpt_name}.yaml')
output_model = output_model + '@' + output_ckpt_name
tuner_dict = {
'NAME': output_model,
'NAME_ZH': output_model,
# 'BASE_MODEL': base_model,
'BASE_MODEL': base_model_revision,
'TUNER_TYPE': tuner_name,
'DESCRIPTION': '',
'MODEL_PATH': model_path,
'IMAGE_PATH': image_path,
'PROMPT_EXAMPLE': eval_prompts,
'SOURCE': 'self_train',
'CKPT_NAME': output_ckpt_name,
'PARAMS': params_info
}
tuner_cfg = Config(cfg_dict=tuner_dict, load=False)
if not os.path.exists(model_path):
gr.Error(self.component_names.model_err4)
self.manager.inference.model_manage_ui.pipe_manager.register_tuner(
tuner_cfg,
name=tuner_cfg.NAME_ZH
if self.language == 'zh' else tuner_cfg.NAME,
is_customized=True)
pipeline_level_modules = self.manager.inference.model_manage_ui.pipe_manager.pipeline_level_modules
if tuner_cfg.BASE_MODEL not in pipeline_level_modules:
gr.Error(self.component_names.model_err3 +
tuner_cfg.BASE_MODEL)
pipeline_ins = pipeline_level_modules[tuner_cfg.BASE_MODEL]
diffusion_model = f"{tuner_cfg.BASE_MODEL}_{pipeline_ins.diffusion_model['name']}"
default_choices = self.manager.inference.model_manage_ui.pipe_manager.module_level_choices
if 'customized_tuners' in default_choices:
if tuner_cfg.BASE_MODEL not in default_choices[
'customized_tuners']:
default_choices['customized_tuners'] = {}
tunner_choices = default_choices['customized_tuners'][
tuner_cfg.BASE_MODEL]['choices']
tunner_default = default_choices['customized_tuners'][
tuner_cfg.BASE_MODEL]['default']
if not isinstance(tunner_default, list):
tunner_default = [tunner_default]
else:
tunner_choices = []
tunner_default = []
with open(cfg_file, 'w') as f_out:
yaml.dump(copy.deepcopy(tuner_cfg.cfg_dict),
f_out,
encoding='utf-8',
allow_unicode=True,
default_flow_style=False)
# example_image = tuner_cfg.get('IMAGE_PATH', None)
# if isinstance(example_image, list) and len(example_image)>0:
# example_image = example_image[0]
#
# prompt_example = tuner_cfg.get('PROMPT_EXAMPLE', None)
# if isinstance(prompt_example, list) and len(prompt_example)>0:
# prompt_example = prompt_example[0]
base_model = tuner_cfg.get('BASE_MODEL', '')
if not base_model == '':
if base_model not in manager.inference.tuner_ui.name_level_tuners:
manager.inference.tuner_ui.name_level_tuners[
base_model] = {}
manager.inference.tuner_ui.name_level_tuners[base_model][
output_model] = tuner_cfg
return (
gr.Tabs(selected='inference'), cfg_file,
gr.Tabs(selected='tuner_ui'),
gr.CheckboxGroup(
value='使用微调' if self.language == 'zh' else 'Use Tuners'),
gr.Dropdown(value=diffusion_model),
gr.Dropdown(choices=tunner_choices, value=tunner_default)
# gr.Text(value=tuner_cfg.get('TUNER_TYPE', '')),
# gr.Text(value=tuner_cfg.get('BASE_MODEL', '')),
# gr.Image(value=example_image),
# gr.Text(value=tuner_cfg.get('DESCRIPTION', '')),
# gr.Text(value=prompt_example)
)
self.go_to_inferece_btn.click(
go_to_inferece,
inputs=[self.output_model_name, self.output_ckpt_name],
outputs=[
manager.tabs, manager.inference.infer_info,
manager.inference.setting_tab,
manager.inference.check_box_for_setting,
manager.inference.model_manage_ui.diffusion_model,
manager.inference.tuner_ui.custom_tuner_model
# manager.inference.tuner_ui.tuner_type,
# manager.inference.tuner_ui.base_model,
# manager.inference.tuner_ui.tuner_example,
# manager.inference.tuner_ui.tuner_desc,
# manager.inference.tuner_ui.tuner_prompt_example
],
queue=False)
def export_train_log(model_name):
current_log_folder = os.path.join(self.work_dir, model_name)
_ = trainer_ui.trainer_ins.get_log(model_name)
out_log = os.path.join(current_log_folder, 'output_std_log.txt')
if os.path.exists(out_log):
zip_path = f'{trainer_ui.current_train_model}_log.zip'
cache_folder = f'{trainer_ui.current_train_model}_log'
if not os.path.exists(
os.path.join(current_log_folder, 'tensorboard')):
os.makedirs(os.path.join(current_log_folder,
'tensorboard'),
exist_ok=True)
res = os.popen(
f"cd {current_log_folder} && mkdir -p '{cache_folder}' "
f"&& cp -rf tensorboard '{cache_folder}/tensorboard' "
f"&& cp -rf 'output_std_log.txt' '{cache_folder}/output_std_log.txt' "
f"&& zip -r '{zip_path}' '{cache_folder}'/* "
f"&& rm -rf '{cache_folder}'")
print(res.readlines())
return gr.File(value=os.path.join(current_log_folder,
zip_path),
visible=True)
else:
gr.Error(self.component_names.training_warn1)
return gr.File()
self.export_log.click(export_train_log,
inputs=[self.output_model_name],
outputs=[self.export_url],
queue=False)
@@ -2,9 +2,10 @@
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import datetime
import json
import os
import random
import time
from collections import OrderedDict
import gradio as gr
import torch
@@ -12,12 +13,12 @@ import yaml
import scepter
from scepter.modules.utils.file_system import FS
from scepter.studio.self_train.scripts.trainer import TrainManager
from scepter.studio.self_train.self_train_ui.component_names import \
TrainerUIName
from scepter.studio.self_train.utils.config_parser import (
get_default, get_values_by_model, get_values_by_model_version,
get_values_by_model_version_tuner,
get_values_by_model_version_tuner_resolution)
get_values_by_model_version_tuner)
from scepter.studio.utils.uibase import UIBase
@@ -30,11 +31,17 @@ def print_memory_status(is_debug):
return gpu_mem
def is_basic_or_container_type(obj):
basic_and_container_types = (int, float, str, bool, complex, list, tuple,
dict, set)
return isinstance(obj, basic_and_container_types)
refresh_symbol = '\U0001f504' # 🔄
def get_work_name(model, version, tuner, resolution):
model_prefix = f'Swift@{model}@{version}@{tuner}@{resolution}'
def get_work_name(model, version, tuner):
model_prefix = f'Swift@{model}@{version}@{tuner}'
return model_prefix + '@' + '{0:%Y%m%d%H%M%S%f}'.format(
datetime.datetime.now()) + ''.join(
[str(random.randint(1, 10)) for i in range(3)])
@@ -44,12 +51,26 @@ class TrainerUI(UIBase):
def __init__(self, cfg, all_cfg_value, is_debug=False, language='en'):
self.BASE_CFG_VALUE = all_cfg_value
self.para_data = get_default(self.BASE_CFG_VALUE)
self.train_para_data = cfg.TRAIN_PARAS
self.run_script = os.path.join(os.path.dirname(scepter.dirname),
cfg.SCRIPT_DIR, 'run_task.py')
self.work_dir_pre, _ = FS.map_to_local(cfg.WORK_DIR)
self.is_debug = is_debug
self.current_train_model = None
self.trainer_ins = TrainManager(self.run_script, self.work_dir_pre)
self.component_names = TrainerUIName(language=language)
self.h_level_dict = {}
for hw_tuple in self.train_para_data.RESOLUTIONS.get('VALUES', []):
h, w = hw_tuple
if h not in self.h_level_dict:
self.h_level_dict[h] = []
self.h_level_dict[h].append(w)
self.h_level_dict = OrderedDict(
sorted(self.h_level_dict.items(),
key=lambda x: int(x[0]),
reverse=False))
def create_ui(self):
with gr.Box():
with gr.Row(variant='panel', equal_height=True):
@@ -88,15 +109,6 @@ class TrainerUI(UIBase):
'model_default', ''),
label=self.component_names.base_model,
interactive=True)
with gr.Column(scale=1, min_width=0):
self.tuner_name = gr.Dropdown(
choices=self.para_data.get(
'tuner_choices', []),
value=self.para_data.get(
'tuner_default', ''),
label=self.component_names.tuner_name,
interactive=True)
with gr.Row():
with gr.Column(scale=1, min_width=0):
self.base_model_revision = gr.Dropdown(
choices=self.para_data.get(
@@ -107,13 +119,35 @@ class TrainerUI(UIBase):
base_model_revision,
interactive=True)
with gr.Row():
with gr.Column(scale=1, min_width=0):
self.resolution = gr.Dropdown(
self.tuner_name = gr.Dropdown(
choices=self.para_data.get(
'resolution_choices', []),
'tuner_choices', []),
value=self.para_data.get(
'resolution_default', 1024),
label=self.component_names.resolution,
'tuner_default', ''),
label=self.component_names.tuner_name,
interactive=True)
with gr.Row():
with gr.Column(scale=1, min_width=0):
self.resolution_height = gr.Dropdown(
choices=list(self.h_level_dict.keys()),
value=self.para_data.get(
'RESOLUTION', 1024)[0],
label=self.component_names.
resolution_height,
allow_custom_value=True,
interactive=True)
with gr.Column(scale=1, min_width=0):
self.resolution_width = gr.Dropdown(
choices=self.h_level_dict[
self.resolution_height.value],
value=self.para_data.get(
'RESOLUTION', 1024)[1],
label=self.component_names.
resolution_width,
allow_custom_value=True,
interactive=True)
@@ -161,19 +195,36 @@ class TrainerUI(UIBase):
value='')
with gr.Row():
with gr.Column(scale=5, min_width=0):
self.work_name = gr.Text(
label=self.component_names.work_name,
value=None,
interactive=False)
with gr.Column(scale=1, min_width=0):
self.work_name_button = gr.Button(
value=refresh_symbol)
with gr.Column(scale=2, min_width=0):
self.push_to_hub = gr.Checkbox(
label=self.component_names.push_to_hub,
value=False,
visible=False)
self.eval_prompts = gr.Dropdown(
value=None,
choices=self.train_para_data.get(
'EVAL_PROMPTS', []),
label=self.component_names.eval_prompts,
interactive=True,
multiselect=True,
allow_custom_value=True,
max_choices=20)
with gr.Row(variant='panel', equal_height=True):
with gr.Box():
with gr.Row(variant='panel', equal_height=True):
gr.Markdown(self.component_names.work_name)
with gr.Row(variant='panel', equal_height=True):
with gr.Column(scale=6, min_width=0):
self.work_name = gr.Text(value=None,
container=False,
interactive=False)
with gr.Column(scale=2, min_width=0):
self.work_name_button = gr.Button(
value=refresh_symbol, size='lg')
with gr.Row(variant='panel', visible=False, equal_height=True):
with gr.Column(scale=2, min_width=0):
self.push_to_hub = gr.Checkbox(
label=self.component_names.push_to_hub,
value=False,
visible=False)
with gr.Row(variant='panel', equal_height=True):
self.examples = gr.Examples(
@@ -193,15 +244,10 @@ class TrainerUI(UIBase):
self.data_type, self.ms_data_space, self.ms_data_name,
self.ms_data_subname
])
with gr.Row(variant='panel', equal_height=True):
with gr.Box():
gr.Markdown(self.component_names.log_block)
self.training_message = gr.Markdown()
with gr.Row(variant='panel', equal_height=True):
self.training_button = gr.Button()
def set_callbacks(self, inference_ui):
def set_callbacks(self, inference_ui, manager):
def change_data_type(data_type):
if data_type == self.component_names.data_type_choices[0]:
return gr.Box(visible=False)
@@ -217,7 +263,7 @@ class TrainerUI(UIBase):
inputs=[
self.base_model,
self.base_model_revision,
self.tuner_name, self.resolution
self.tuner_name
],
outputs=[self.work_name],
queue=False)
@@ -245,8 +291,12 @@ class TrainerUI(UIBase):
interactive=True), \
gr.Dropdown(value=ret_data.get('tuner_default', ''), choices=ret_data.get('tuner_choices', []),
interactive=True), \
gr.Dropdown(value=ret_data.get('resolution_default', 1024),
choices=ret_data.get('resolution_choices', []), interactive=True)
gr.Dropdown(value=ret_data.get('RESOLUTION', 1024)[0],
choices=list(self.h_level_dict.keys()),
interactive=True), \
gr.Dropdown(value=ret_data.get('RESOLUTION', 1024)[1],
choices=self.h_level_dict[ret_data.get('RESOLUTION', 1024)[0]],
interactive=True)
self.base_model.change(fn=change_train_value_by_model,
inputs=[self.base_model],
@@ -255,7 +305,8 @@ class TrainerUI(UIBase):
self.save_interval, self.train_batch_size,
self.prompt_prefix,
self.base_model_revision, self.tuner_name,
self.resolution
self.resolution_height,
self.resolution_width
],
queue=False)
@@ -282,10 +333,15 @@ class TrainerUI(UIBase):
ret_data.get('SAVE_INTERVAL', 10), \
ret_data.get('TRAIN_BATCH_SIZE', 4), \
ret_data.get('TRAIN_PREFIX', ''), \
gr.Dropdown(value=ret_data.get('tuner_default', ''), choices=ret_data.get('tuner_choices', []),
gr.Dropdown(value=ret_data.get('tuner_default', ''),
choices=ret_data.get('tuner_choices', []),
interactive=True), \
gr.Dropdown(value=ret_data.get('resolution_default', 1024),
choices=ret_data.get('resolution_choices', []), interactive=True)
gr.Dropdown(value=ret_data.get('RESOLUTION', 1024)[0],
choices=list(self.h_level_dict.keys()),
interactive=True), \
gr.Dropdown(value=ret_data.get('RESOLUTION', 1024)[1],
choices=self.h_level_dict[ret_data.get('RESOLUTION', 1024)[0]],
interactive=True)
#
self.base_model_revision.change(
@@ -294,11 +350,10 @@ class TrainerUI(UIBase):
outputs=[
self.train_epoch, self.learning_rate, self.save_interval,
self.train_batch_size, self.prompt_prefix, self.tuner_name,
self.resolution
self.resolution_height, self.resolution_width
],
queue=False)
#
#
def change_train_value_by_model_version_tuner(base_model,
base_model_revision,
@@ -323,8 +378,12 @@ class TrainerUI(UIBase):
ret_data.get('SAVE_INTERVAL', 10), \
ret_data.get('TRAIN_BATCH_SIZE', 4), \
ret_data.get('TRAIN_PREFIX', ''), \
gr.Dropdown(value=ret_data.get('resolution_default', 1024),
choices=ret_data.get('resolution_choices', []), interactive=True)
gr.Dropdown(value=ret_data.get('RESOLUTION', 1024)[0],
choices=list(self.h_level_dict.keys()),
interactive=True), \
gr.Dropdown(value=ret_data.get('RESOLUTION', 1024)[1],
choices=self.h_level_dict[ret_data.get('RESOLUTION', 1024)[0]],
interactive=True)
#
self.tuner_name.change(fn=change_train_value_by_model_version_tuner,
@@ -335,56 +394,30 @@ class TrainerUI(UIBase):
outputs=[
self.train_epoch, self.learning_rate,
self.save_interval, self.train_batch_size,
self.prompt_prefix, self.resolution
self.prompt_prefix, self.resolution_height,
self.resolution_width
],
queue=False)
#
def change_train_value_by_model_version_tuner_resolution(
base_model, base_model_revision, tuner_name, resolution):
'''
Changes to the base model will affect the training parameters,
and it is best to define the related default values in the YAML.
Training Iterations:
Learning Rate:
Training Prefix Used:
Training Batch Size:
Supported Resolutions:
Supported Fine-tuning Methods:
Supported Fine-tuning Methods:
Prefix for Saved Model:
'''
ret_data = get_values_by_model_version_tuner_resolution(
self.BASE_CFG_VALUE, base_model, base_model_revision,
tuner_name, resolution)
print('change_train_value_by_model_version_tuner_resolution',
ret_data)
# work_name = get_work_name(base_model, base_model_revision,
# tuner_name, resolution)
return ret_data.get('EPOCHS', 10), \
ret_data.get('LEARNING_RATE', 0.0001), \
ret_data.get('SAVE_INTERVAL', 10), \
ret_data.get('TRAIN_BATCH_SIZE', 4), \
ret_data.get('TRAIN_PREFIX', '')
def change_resolution(h):
if h not in self.h_level_dict:
return gr.Dropdown()
all_choices = self.h_level_dict[h]
default = all_choices[0]
return gr.Dropdown(choices=all_choices, value=default)
#
self.resolution.change(
fn=change_train_value_by_model_version_tuner_resolution,
inputs=[
self.base_model, self.base_model_revision, self.tuner_name,
self.resolution
],
outputs=[
self.train_epoch, self.learning_rate, self.save_interval,
self.train_batch_size, self.prompt_prefix
],
queue=False)
self.resolution_height.change(change_resolution,
inputs=[self.resolution_height],
outputs=[self.resolution_width],
queue=False)
def run_train(work_name, data_type, ms_data_space, ms_data_name,
ms_data_subname, base_model, base_model_revision,
tuner_name, resolution, train_epoch, learning_rate,
save_interval, train_batch_size, prompt_prefix,
replace_keywords, push_to_hub):
tuner_name, resolution_height, resolution_width,
train_epoch, learning_rate, save_interval,
train_batch_size, prompt_prefix, replace_keywords,
push_to_hub, eval_prompts):
# Check Cuda
if not torch.cuda.is_available() and not self.is_debug:
raise gr.Error(self.component_names.training_err1)
@@ -392,10 +425,25 @@ class TrainerUI(UIBase):
if work_name == 'custom' or work_name is None or work_name == '':
raise gr.Error(self.component_names.training_err4)
work_dir = os.path.join(self.work_dir_pre, work_name)
if not os.path.exists(work_dir):
os.makedirs(work_dir)
else:
self.current_train_model = work_name
if os.path.exists(work_dir) or os.path.exists(
f'.flag/{work_name}.tmp'):
raise gr.Error(self.component_names.training_err4)
else:
os.makedirs(work_dir)
os.makedirs('.flag', exist_ok=True)
with open(f'.flag/{work_name}.tmp', 'w') as f:
f.write('new line.')
# save params
with open(os.path.join(work_dir, 'params.json'), 'w') as f_out:
json.dump(
{
key: val
for key, val in locals().items()
if is_basic_or_container_type(val)
},
f_out,
ensure_ascii=False)
if push_to_hub:
model_id = work_name
@@ -404,28 +452,23 @@ class TrainerUI(UIBase):
else:
hub_model_id = ''
# Check Cuda Memory
if torch.cuda.is_available() and not self.is_debug:
device = torch.device('cuda:0')
required_memory_bytes = 4 * (1024**3)
try:
tensor = torch.empty( # noqa
(required_memory_bytes // 4, ), device=device
) # create 4GB tensor to check the memory if enough
del tensor
except RuntimeError:
raise gr.Error(self.component_names.training_err2)
# Check Instance Valid
if ms_data_name is None:
raise gr.Error(self.component_names.training_err3)
st_time = time.time()
def prepare_data(data_cfg):
def prepare_train_data(data_cfg):
data_cfg['BATCH_SIZE'] = int(train_batch_size)
data_cfg['PROMPT_PREFIX'] = prompt_prefix
data_cfg['REPLACE_KEYWORDS'] = replace_keywords
for trans in data_cfg['TRANSFORMS']:
if trans['NAME'] in [
'Resize', 'FlexibleResize', 'CenterCrop',
'FlexibleCenterCrop'
]:
trans['SIZE'] = [
int(resolution_height),
int(resolution_width)
]
if data_type in self.component_names.data_type_choices:
if ms_data_name.startswith(
'http') or ms_data_name.endswith('zip'):
@@ -482,7 +525,19 @@ class TrainerUI(UIBase):
data_cfg['MS_REMAP_KEYS'] = {'Text': 'Prompt'}
else:
data_cfg['MS_REMAP_KEYS'] = None
data_cfg['OUTPUT_SIZE'] = int(resolution)
data_cfg['OUTPUT_SIZE'] = [
int(resolution_height),
int(resolution_width)
]
return data_cfg
def prepare_eval_data(data_cfg):
data_cfg['PROMPT_PREFIX'] = prompt_prefix
data_cfg['IMAGE_SIZE'] = [
int(resolution_height),
int(resolution_width)
]
data_cfg['PROMPT_DATA'] = eval_prompts
return data_cfg
def prepare_train_config():
@@ -490,7 +545,7 @@ class TrainerUI(UIBase):
current_model_info = self.BASE_CFG_VALUE[base_model][
base_model_revision]
modify_para = current_model_info['modify_para']
cfg = current_model_info['config_value']
cfg = copy.deepcopy(current_model_info['config_value'])
if isinstance(modify_para, dict) and tuner_name in modify_para:
modify_c = modify_para[tuner_name]
if isinstance(modify_c, dict) and 'TRAIN' in modify_c:
@@ -518,18 +573,33 @@ class TrainerUI(UIBase):
cfg['SOLVER']['MAX_EPOCHS'] = int(train_epoch)
cfg['SOLVER']['TRAIN_DATA']['BATCH_SIZE'] = int(
train_batch_size)
cfg['SOLVER']['TUNER'] = current_model_info[
'tuner_para'][tuner_name] if isinstance(
current_model_info['tuner_para'],
dict) and tuner_name in current_model_info[
'tuner_para'] else None
cfg['SOLVER']['TRAIN_DATA'] = prepare_data(
if 'TUNER' in cfg['SOLVER']:
cfg['SOLVER']['TUNER'] = current_model_info['tuner_para'][
tuner_name] if isinstance(
current_model_info['tuner_para'],
dict) and tuner_name in current_model_info[
'tuner_para'] else None
cfg['SOLVER']['TRAIN_DATA'] = prepare_train_data(
cfg['SOLVER']['TRAIN_DATA'])
if eval_prompts is not None and len(eval_prompts) > 0:
cfg['SOLVER']['EVAL_DATA'] = prepare_eval_data(
cfg['SOLVER']['EVAL_DATA'])
else:
cfg['SOLVER'].pop('EVAL_DATA')
if 'SAMPLE_ARGS' in cfg['SOLVER']:
cfg['SOLVER']['SAMPLE_ARGS']['IMAGE_SIZE'] = [
int(resolution_height),
int(resolution_width)
]
for hook in cfg['SOLVER']['TRAIN_HOOKS']:
if hook['NAME'] == 'CheckpointHook':
hook['INTERVAL'] = save_interval
hook['PUSH_TO_HUB'] = push_to_hub
hook['HUB_MODEL_ID'] = hub_model_id
if 'EVAL_HOOKS' in cfg['SOLVER']:
for hook in cfg['SOLVER']['EVAL_HOOKS']:
if hook['NAME'] == 'ProbeDataHook':
hook['PROB_INTERVAL'] = save_interval
with open(cfg_file, 'w') as f_out:
yaml.dump(cfg,
@@ -539,48 +609,33 @@ class TrainerUI(UIBase):
default_flow_style=False)
return cfg_file
cfg = prepare_train_config()
def train_fn(cfg_file):
torch.cuda.empty_cache()
cmd = f'PYTHONPATH=. python {self.run_script} ' \
f'--cfg={cfg_file} 2> {self.work_dir_pre}/std_out.txt'
print(cmd)
if not self.is_debug:
res = os.system(cmd)
else:
res = 0
if res != 0:
error_info = '\n'.join(
open(f'{self.work_dir_pre}/std_out.txt',
'r').read().split('\n')[-20:])
raise gr.Error(
f'{self.component_names.training_err5} ({error_info}) '
)
train_fn(cfg)
before_kill_inference = self.trainer_ins.check_memory()
for k, v in manager.inference.pipe_manager.pipeline_level_modules.items(
):
if hasattr(v, 'dynamic_unload'):
v.dynamic_unload(name='all')
after_kill_inference = self.trainer_ins.check_memory()
message = f'GPU info: {before_kill_inference}. \n\n'
message += f'After unloading inference models, the GPU info: {after_kill_inference}. \n\n'
_ = prepare_train_config()
self.trainer_ins.start_task(work_name)
message += self.trainer_ins.get_log(work_name)
if work_name not in inference_ui.model_list:
inference_ui.model_list.append(work_name)
message = f'''
Training completed! \n
Save in [ {work_name} ] \n
Take time [ {time.time() - st_time:.4f}s ] \n
Mem: [ {print_memory_status(self.is_debug)} ]
'''
print(message)
return message, gr.Dropdown.update(choices=inference_ui.model_list,
value=work_name)
gr.Info('Start Training!' + message)
return gr.Dropdown.update(choices=inference_ui.model_list,
value=work_name)
self.training_button.click(
run_train,
inputs=[
self.work_name, self.data_type, self.ms_data_space,
self.ms_data_name, self.ms_data_subname, self.base_model,
self.base_model_revision, self.tuner_name, self.resolution,
self.base_model_revision, self.tuner_name,
self.resolution_height, self.resolution_width,
self.train_epoch, self.learning_rate, self.save_interval,
self.train_batch_size, self.prompt_prefix,
self.replace_keywords, self.push_to_hub
self.replace_keywords, self.push_to_hub, self.eval_prompts
],
outputs=[self.training_message, inference_ui.output_model_name],
outputs=[inference_ui.output_model_name],
queue=True)
@@ -10,7 +10,6 @@ paras_keys = [
'MEMORY', 'EPOCHS', 'SAVE_INTERVAL', 'EPSEC', 'LEARNING_RATE',
'IS_DEFAULT', 'TUNER'
]
control_paras_keys = ['CONTROL_MODE', 'RESOLUTION', 'IS_DEFAULT']
@@ -25,10 +24,9 @@ def build_meta_index(meta_cfg, config_file):
f'Para key {key} not defined in {config_file} META/RARAS[{idx}]'
)
assert key in para
tuner_type[para['TUNER'] + '@' + str(para['RESOLUTION'])] = para
tuner_type[para['TUNER']] = para
if para['IS_DEFAULT']:
tuner_type['default'] = para['TUNER'] + '@' + str(
para['RESOLUTION'])
tuner_type['default'] = para['TUNER']
tuner_type['choices'] = list(tuner_type.keys())
if 'default' in tuner_type['choices']:
@@ -37,14 +35,6 @@ def build_meta_index(meta_cfg, config_file):
tuner_type['default'] = tuner_type['choices'][0] if len(
tuner_type['choices']) > 0 else ''
# 重新组织选项
choices = {}
for t_type in tuner_type['choices']:
t_name, resolution = t_type.split('@')
if t_name not in choices:
choices[t_name] = []
choices[t_name].append(int(resolution))
tuner_type['choices'] = choices
tuner_paras = meta_cfg.get('TUNERS', None)
return tuner_type, tuner_paras
@@ -60,13 +50,10 @@ def build_meta_index_control(meta_cfg, config_file):
f'Para key {key} not defined in {config_file} META/RARAS[{idx}]'
)
assert key in para
control_type[para['CONTROL_MODE'] + '@' +
str(para['RESOLUTION'])] = para
control_type[para['CONTROL_MODE']] = para
# control_type[para["CONTROL_MODE"]] = para
if para['IS_DEFAULT']:
control_type['default'] = para['CONTROL_MODE'] + '@' + str(
para['RESOLUTION'])
# control_type["default"] = para["CONTROL_MODE"]
control_type['default'] = para['CONTROL_MODE']
control_type['choices'] = list(control_type.keys())
if 'default' in control_type['choices']:
@@ -75,14 +62,6 @@ def build_meta_index_control(meta_cfg, config_file):
control_type['default'] = control_type['choices'][0] if len(
control_type['choices']) > 0 else ''
# 重新组织选项
choices = {}
for t_type in control_type['choices']:
t_name, resolution = t_type.split('@')
if t_name not in choices:
choices[t_name] = []
choices[t_name].append(int(resolution))
control_type['choices'] = choices
return control_type, paras
@@ -166,26 +145,19 @@ def get_default(config_dict):
return ret_data
ret_data['version_choices'] = default_version_cfg['choices']
ret_data['version_default'] = default_version_cfg['default']
default_tuner_cfg = default_version_cfg.get(default_version_cfg['default'],
default_model_cfg = default_version_cfg.get(default_version_cfg['default'],
None)
if default_tuner_cfg is None:
if default_model_cfg is None:
return ret_data
if 'tuner_type' in default_tuner_cfg and default_tuner_cfg['tuner_type'][
if 'tuner_type' in default_model_cfg and default_model_cfg['tuner_type'][
'default'] != '':
default_tuner_cfg = default_tuner_cfg['tuner_type']
default_tuner_cfg = default_model_cfg['tuner_type']
else:
return ret_data
ret_data['tuner_choices'] = list(default_tuner_cfg['choices'].keys())
ret_data['tuner_choices'] = default_tuner_cfg['choices']
defalt_t_type = default_tuner_cfg['default']
ret_data['tuner_default'] = defalt_t_type
type_paras = default_tuner_cfg.get(defalt_t_type, None)
default_t_n = defalt_t_type.split('@')[0]
default_r_n = int(defalt_t_type.split('@')[1])
ret_data['resolution_choices'] = default_tuner_cfg['choices'].get(
default_t_n, [])
ret_data['tuner_default'] = default_t_n
ret_data['resolution_default'] = default_r_n
if type_paras is not None:
ret_data.update(type_paras)
return ret_data
@@ -198,21 +170,15 @@ def get_values_by_model(config_dict, model_name):
return ret_data
ret_data['version_choices'] = version_cfg['choices']
ret_data['version_default'] = version_cfg['default']
default_tuner_cfg = version_cfg.get(version_cfg['default'], None)
if default_tuner_cfg is None:
default_model_cfg = version_cfg.get(version_cfg['default'], None)
if default_model_cfg is None:
return ret_data
default_tuner_cfg = default_tuner_cfg['tuner_type']
ret_data['tuner_choices'] = list(default_tuner_cfg['choices'].keys())
default_tuner_cfg = default_model_cfg['tuner_type']
ret_data['tuner_choices'] = default_tuner_cfg['choices']
defalt_t_type = default_tuner_cfg['default']
ret_data['tuner_default'] = defalt_t_type
type_paras = default_tuner_cfg.get(defalt_t_type, None)
default_t_n = defalt_t_type.split('@')[0]
default_r_n = int(defalt_t_type.split('@')[1])
ret_data['resolution_choices'] = default_tuner_cfg['choices'].get(
default_t_n, [])
ret_data['tuner_default'] = default_t_n
ret_data['resolution_default'] = default_r_n
if type_paras is not None:
ret_data.update(type_paras)
return ret_data
@@ -226,18 +192,12 @@ def get_values_by_model_version(config_dict, model_name, version):
tuner_cfg = version_cfg.get(version, None)
if tuner_cfg is None:
return ret_data
default_tuner_cfg = tuner_cfg['tuner_type']
ret_data['tuner_choices'] = list(default_tuner_cfg['choices'].keys())
ret_data['tuner_choices'] = default_tuner_cfg['choices']
defalt_t_type = default_tuner_cfg['default']
ret_data['tuner_default'] = defalt_t_type
type_paras = default_tuner_cfg.get(defalt_t_type, None)
default_t_n = defalt_t_type.split('@')[0]
default_r_n = int(defalt_t_type.split('@')[1])
ret_data['resolution_choices'] = default_tuner_cfg['choices'].get(
default_t_n, [])
ret_data['tuner_default'] = default_t_n
ret_data['resolution_default'] = default_r_n
if type_paras is not None:
ret_data.update(type_paras)
return ret_data
@@ -249,34 +209,11 @@ def get_values_by_model_version_tuner(config_dict, model_name, version,
version_cfg = config_dict.get(model_name, None)
if version_cfg is None:
return ret_data
tuner_cfg = version_cfg.get(version, None)
if tuner_cfg is None:
model_cfg = version_cfg.get(version, None)
if model_cfg is None:
return ret_data
tuner_cfg = tuner_cfg['tuner_type']
ret_data['resolution_choices'] = tuner_cfg['choices'].get(tuner_name, [])
if len(ret_data['resolution_choices']) > 0:
t_type = '{}@{}'.format(tuner_name, ret_data['resolution_choices'][0])
ret_data['resolution_default'] = ret_data['resolution_choices'][0]
type_paras = tuner_cfg.get(t_type, None)
if type_paras is not None:
ret_data.update(type_paras)
return ret_data
def get_values_by_model_version_tuner_resolution(config_dict, model_name,
version, tuner_name,
resolution):
ret_data = {}
version_cfg = config_dict.get(model_name, None)
if version_cfg is None:
return ret_data
tuner_cfg = version_cfg.get(version, None)
if tuner_cfg is None:
return ret_data
tuner_cfg = tuner_cfg['tuner_type']
type_paras = tuner_cfg.get('{}@{}'.format(tuner_name, resolution), None)
tuner_cfg = model_cfg['tuner_type']
type_paras = tuner_cfg.get(tuner_name, None)
if type_paras is not None:
ret_data.update(type_paras)
return ret_data
@@ -371,15 +308,8 @@ def get_control_default(config_dict):
ret_data['control_choices'] = list(default_control_cfg['choices'].keys())
defalt_t_type = default_control_cfg['default']
ret_data['control_default'] = defalt_t_type
type_paras = default_control_cfg.get(defalt_t_type, None)
default_t_n = defalt_t_type.split('@')[0]
default_r_n = int(defalt_t_type.split('@')[1])
ret_data['resolution_choices'] = default_control_cfg['choices'].get(
default_t_n, [])
ret_data['control_default'] = default_t_n
ret_data['resolution_default'] = default_r_n
if type_paras is not None:
ret_data.update(type_paras)
return ret_data
@@ -0,0 +1,322 @@
# -*- coding: utf-8 -*-
import os
from collections import OrderedDict
import gradio as gr
from scepter.modules.utils.config import Config
from scepter.modules.utils.file_system import FS
from scepter.studio.tuner_manager.manager_ui.component_names import \
TunerManagerNames
from scepter.studio.tuner_manager.utils.dict import (delete_2level_dict,
update_2level_dict)
from scepter.studio.tuner_manager.utils.path import is_valid_filename
from scepter.studio.tuner_manager.utils.yaml import save_yaml
from scepter.studio.utils.uibase import UIBase
class BrowserUI(UIBase):
def __init__(self, cfg, language='en'):
self.work_dir = cfg.WORK_DIR
self.yaml = os.path.join(self.work_dir, cfg.TUNER_LIST_YAML)
if not FS.exists(self.yaml):
self.saved_tuners = []
with FS.put_to(self.yaml) as local_path:
save_yaml({'TUNERS': self.saved_tuners}, local_path)
else:
with FS.get_from(self.yaml) as local_path:
self.saved_tuners = Config.get_plain_cfg(
Config(cfg_file=local_path).TUNERS)
self.saved_tuners_to_category()
self.component_names = TunerManagerNames(language)
self.language = language
def saved_tuners_to_category(self):
self.saved_tuners_category = OrderedDict()
for tuner in self.saved_tuners:
first_level = f"{tuner['BASE_MODEL']}-{tuner['TUNER_TYPE']}"
second_level = f"{tuner['NAME']}"
update_2level_dict(self.saved_tuners_category,
{first_level: {
second_level: tuner
}})
def category_to_saved_tuners(self):
self.saved_tuners = []
for k, v in self.saved_tuners_category.items():
for kk, vv in v.items():
self.saved_tuners.append(vv)
def get_choices_and_values(self):
diffusion_models_choice = list(self.saved_tuners_category.keys())
diffusion_model = diffusion_models_choice[0] if len(
diffusion_models_choice) > 0 else None
tuner_models_choice = []
tuner_model = None
if diffusion_model:
tuner_models_choice = list(
self.saved_tuners_category.get(diffusion_model, {}).keys())
tuner_model = tuner_models_choice[0] if len(
tuner_models_choice) > 0 else None
return diffusion_models_choice, diffusion_model, tuner_models_choice, tuner_model
def create_ui(self, *args, **kwargs):
diffusion_models_choice, diffusion_model, tuner_models_choice, tuner_model = self.get_choices_and_values(
)
with gr.Column():
with gr.Box():
gr.Markdown(self.component_names.browser_block_name)
with gr.Row(variant='panel', equal_height=True):
with gr.Column(scale=4, min_width=0):
self.diffusion_models = gr.Dropdown(
label=self.component_names.base_models,
choices=diffusion_models_choice,
value=diffusion_model,
multiselect=False,
interactive=True)
with gr.Column(scale=4, min_width=0):
self.tuner_models = gr.Dropdown(
label=self.component_names.tuner_name,
choices=tuner_models_choice,
value=tuner_model,
multiselect=False,
interactive=True)
with gr.Column(scale=1, min_width=0):
self.save_button = gr.Button(
label='Save',
value=self.component_names.save_symbol,
elem_classes='type_row',
elem_id='save_button',
visible=True)
self.delete_button = gr.Button(
label='Delete',
value=self.component_names.delete_symbol,
elem_classes='type_row',
elem_id='delete_button',
visible=False)
with gr.Column(scale=1, min_width=0):
self.refresh_button = gr.Button(
label='Delete',
value=self.component_names.refresh_symbol,
elem_classes='type_row',
elem_id='refresh_button',
visible=True)
def check_new_name(self, tuner_name):
if tuner_name.strip() == '':
return False, f"'{tuner_name}' is whitespace! Only support 'a-z', 'A-Z', '0-9' and '_'."
if not is_valid_filename(tuner_name):
return False, f"'{tuner_name}' is not a valid tuner name! Only support 'a-z', 'A-Z', '0-9' and '_'."
for tuner in self.saved_tuners:
if tuner_name == tuner['NAME']:
return False, f"Tuner name '{tuner_name}' has been taken!"
return True, 'legal'
def save_tuner(self, src_path, sub_dir, tuner_name, tuner_example):
tar_path = os.path.join(self.work_dir, sub_dir)
if not FS.exists(tar_path):
FS.make_dir(tar_path)
tar_path = os.path.join(tar_path, tuner_name)
FS.put_dir_from_local_dir(src_path, tar_path)
tuner_example_path = None
if tuner_example is not None:
from PIL import Image
tuner_example_path = os.path.join(tar_path, f'{tuner_name}.jpg')
tuner_example = Image.fromarray(tuner_example)
with FS.put_to(tuner_example_path) as local_path:
tuner_example.save(local_path)
return tar_path, tuner_example_path
def add_tuner(self, new_tuner, manager, now_diffusion_model):
self.saved_tuners.append(new_tuner)
self.saved_tuners_to_category()
with FS.put_to(self.yaml) as local_path:
save_yaml({'TUNERS': self.saved_tuners}, local_path)
# register to pipeline
new_tuner = Config(cfg_dict=new_tuner, load=False)
manager.inference.model_manage_ui.pipe_manager.register_tuner(
new_tuner,
name=new_tuner.NAME_ZH
if self.language == 'zh' else new_tuner.NAME,
is_customized=True)
# update choices
pipe_manager = manager.inference.model_manage_ui.pipe_manager
now_pipeline = pipe_manager.model_level_info[now_diffusion_model][
'pipeline'][0]
default_choices = pipe_manager.module_level_choices
custom_tunner_choices = []
if 'customized_tuners' in default_choices and now_pipeline in default_choices[
'customized_tuners']:
custom_tunner_choices = default_choices['customized_tuners'][
now_pipeline]['choices']
# update tuner ui name_level_tuners
name_level_tuners = manager.inference.tuner_ui.name_level_tuners
if new_tuner.BASE_MODEL not in name_level_tuners:
name_level_tuners[new_tuner.BASE_MODEL] = {}
if self.language == 'zh':
name_level_tuners[new_tuner.BASE_MODEL][
new_tuner.NAME_ZH] = new_tuner
else:
name_level_tuners[new_tuner.BASE_MODEL][new_tuner.NAME] = new_tuner
return custom_tunner_choices
def delete_tuner(self, first_level, second_level, manager,
now_diffusion_model):
self.saved_tuners_category, del_tuner = delete_2level_dict(
self.saved_tuners_category, first_level, second_level)
self.category_to_saved_tuners()
save_yaml({'TUNERS': self.saved_tuners}, self.yaml)
# update choices
pipe_manager = manager.inference.model_manage_ui.pipe_manager
now_pipeline = pipe_manager.model_level_info[now_diffusion_model][
'pipeline'][0]
default_choices = pipe_manager.module_level_choices
custom_tuner_choices = []
if 'customized_tuners' in default_choices and now_pipeline in default_choices[
'customized_tuners']:
custom_tuner_choices = default_choices['customized_tuners'][
now_pipeline]['choices']
# update tuner ui name_level_tuners
del_tuner = Config(cfg_dict=del_tuner, load=False)
name_level_tuners = manager.inference.tuner_ui.name_level_tuners
if del_tuner.BASE_MODEL in name_level_tuners:
if self.language == 'zh':
del name_level_tuners[del_tuner.BASE_MODEL][del_tuner.NAME_ZH]
else:
del name_level_tuners[del_tuner.BASE_MODEL][del_tuner.NAME]
return custom_tuner_choices
def set_callbacks(self, manager, info_ui):
def refresh_browser():
diffusion_models_choice, diffusion_model, tuner_models_choice, tuner_model = self.get_choices_and_values(
)
return (gr.Dropdown(choices=diffusion_models_choice,
value=diffusion_model),
gr.Dropdown(choices=tuner_models_choice,
value=tuner_model))
self.refresh_button.click(
refresh_browser,
inputs=[],
outputs=[self.diffusion_models, self.tuner_models])
def diffusion_model_change(diffusion_model):
choices = list(
self.saved_tuners_category.get(diffusion_model, {}).keys())
return gr.Dropdown(choices=choices,
value=choices[-1] if len(choices) > 0 else None)
self.diffusion_models.change(diffusion_model_change,
inputs=[self.diffusion_models],
outputs=[self.tuner_models],
queue=True)
def tuner_model_change(tuner_model, diffusion_model):
tuner_info = {}
if tuner_model is not None:
tuner_info = self.saved_tuners_category[diffusion_model][
tuner_model]
image_path = tuner_info.get('IMAGE_PATH', None)
if image_path is not None:
image_path = FS.get_from(image_path)
return (gr.Text(value=tuner_info.get('NAME', '')),
gr.Text(value=tuner_info.get('NAME', ''),
interactive=True),
gr.Text(value=tuner_info.get('TUNER_TYPE', '')),
gr.Text(value=tuner_info.get('BASE_MODEL', '')),
gr.Text(value=tuner_info.get('DESCRIPTION', ''),
interactive=True), gr.Image(value=image_path),
gr.Text(value=tuner_info.get('PROMPT_EXAMPLE', ''),
interactive=True))
self.tuner_models.change(
tuner_model_change,
inputs=[self.tuner_models, self.diffusion_models],
outputs=[
info_ui.tuner_name, info_ui.new_name, info_ui.tuner_type,
info_ui.base_model, info_ui.tuner_desc, info_ui.tuner_example,
info_ui.tuner_prompt_example
],
queue=False)
def save_tuner(tuner_name, new_name, tuner_desc, tuner_example,
tuner_prompt_example, now_diffusion_model, info_path):
is_legal, msg = self.check_new_name(new_name)
if not is_legal:
gr.Info('Save failed because ' + msg)
return (gr.Dropdown(), gr.Text(), gr.Dropdown())
info = Config(cfg_file=info_path)
model_dir = info.MODEL_PATH
sub_dir = f'{info.BASE_MODEL}-{info.TUNER_TYPE}'
model_dir, tuner_example = self.save_tuner(model_dir, sub_dir,
new_name, tuner_example)
# config info update
new_tuner = {
'NAME': new_name,
'NAME_ZH': new_name,
'SOURCE': 'self_train',
'DESCRIPTION': tuner_desc,
'BASE_MODEL': info.BASE_MODEL,
'MODEL_PATH': model_dir,
'IMAGE_PATH': tuner_example,
'TUNER_TYPE': info.TUNER_TYPE,
'PROMPT_EXAMPLE': tuner_prompt_example
}
custom_tuner_choices = self.add_tuner(new_tuner, manager,
now_diffusion_model)
return (gr.Dropdown(choices=list(
self.saved_tuners_category.keys()),
value=sub_dir),
gr.Dropdown(choices=list(
self.saved_tuners_category.get(sub_dir, {}).keys()),
value=new_name), gr.Text(value=new_name),
gr.Dropdown(choices=custom_tuner_choices))
self.save_button.click(
save_tuner,
inputs=[
info_ui.tuner_name, info_ui.new_name, info_ui.tuner_desc,
info_ui.tuner_example, info_ui.tuner_prompt_example,
manager.inference.model_manage_ui.diffusion_model,
manager.inference.infer_info
],
outputs=[
self.diffusion_models, self.tuner_models, info_ui.tuner_name,
manager.inference.tuner_ui.custom_tuner_model
],
queue=True)
def delete_tuner(tuner_name, tuner_type, base_model,
now_diffusion_model):
first_level = f'{base_model}-{tuner_type}'
second_level = f'{tuner_name}'
custom_tuner_choices = self.delete_tuner(first_level, second_level,
manager,
now_diffusion_model)
return (gr.Dropdown(
choices=list(self.saved_tuners_category.keys()),
value=None), gr.Dropdown(choices=custom_tuner_choices))
self.delete_button.click(
delete_tuner,
inputs=[
info_ui.tuner_name, info_ui.tuner_type, info_ui.base_model,
manager.inference.model_manage_ui.diffusion_model
],
outputs=[
self.diffusion_models,
manager.inference.tuner_ui.custom_tuner_model
],
queue=True)
@@ -0,0 +1,38 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
# For dataset manager
class TunerManagerNames():
def __init__(self, language='en'):
self.save_symbol = '\U0001F4BE' # 💾
self.delete_symbol = '\U0001f5d1' # 🗑️
self.refresh_symbol = '\U0001f504' # 🔄
if language == 'en':
self.browser_block_name = 'Tuner Browser'
self.base_models = 'Base Model-Tuner Type'
self.tuner_models = 'Tuner Name'
self.info_block_name = 'Tuner Info'
self.tuner_name = 'Tuner Name'
self.rename = 'Rename'
self.tuner_type = 'Tuner Type'
self.base_model_name = 'Base Model Name'
self.tuner_desc = 'Tuner Description'
self.tuner_example = 'Results Example'
self.tuner_prompt_example = 'Prompt Example'
self.save = 'save changes'
self.delete = 'Delete'
elif language == 'zh':
self.browser_block_name = '微调模型查找'
self.base_models = '基模型-微调类型'
self.tuner_models = '微调模型名称'
self.info_block_name = '微调模型详情'
self.tuner_name = '微调模型名称'
self.rename = '重命名'
self.tuner_type = '微调模型类型'
self.base_model_name = '基模型名称'
self.tuner_desc = '微调模型描述'
self.tuner_example = '示例结果'
self.tuner_prompt_example = '示例提示词'
self.save = '保存修改'
self.delete = '删除'
@@ -0,0 +1,64 @@
# -*- coding: utf-8 -*-
import gradio as gr
from scepter.studio.tuner_manager.manager_ui.component_names import \
TunerManagerNames
from scepter.studio.utils.uibase import UIBase
class InfoUI(UIBase):
def __init__(self, cfg, language='en'):
self.component_names = TunerManagerNames(language)
def create_ui(self, *args, **kwargs):
with gr.Column():
with gr.Box():
gr.Markdown(self.component_names.info_block_name)
with gr.Row(variant='panel', 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_name = gr.Text(
value='',
label=self.component_names.tuner_name,
interactive=False)
with gr.Column(scale=1):
self.new_name = gr.Text(
value='',
label=self.component_names.rename)
with gr.Row(equal_height=True):
with gr.Column(scale=1):
self.tuner_type = gr.Text(
value='',
label=self.component_names.tuner_type,
interactive=False)
with gr.Column(scale=1):
self.base_model = gr.Text(
value='',
label=self.component_names.
base_model_name,
interactive=False)
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=True)
with gr.Row(equal_height=True):
self.tuner_prompt_example = gr.Text(
value='',
label=self.component_names.
tuner_prompt_example,
lines=2)
def set_callbacks(self, manager):
pass
@@ -0,0 +1,33 @@
# -*- coding: utf-8 -*-
import os.path
from scepter.modules.utils.config import Config
from scepter.modules.utils.file_system import FS
from scepter.studio.tuner_manager.manager_ui.browser_ui import BrowserUI
from scepter.studio.tuner_manager.manager_ui.info_ui import InfoUI
from scepter.studio.utils.env import init_env
class TunerManagerUI():
def __init__(self,
cfg_general_file,
is_debug=False,
language='en',
root_work_dir='./'):
cfg_general = Config(cfg_file=cfg_general_file)
cfg_general.WORK_DIR = os.path.join(root_work_dir,
cfg_general.WORK_DIR)
if not FS.exists(cfg_general.WORK_DIR):
FS.make_dir(cfg_general.WORK_DIR)
cfg_general = init_env(cfg_general)
self.info_ui = InfoUI(cfg_general, language=language)
self.browser_ui = BrowserUI(cfg_general, language=language)
def create_ui(self):
self.browser_ui.create_ui()
self.info_ui.create_ui()
def set_callbacks(self, manager):
self.info_ui.set_callbacks(manager)
self.browser_ui.set_callbacks(manager, self.info_ui)
@@ -0,0 +1,16 @@
# -*- coding: utf-8 -*-
def update_2level_dict(d, new_dict):
for first, v in new_dict.items():
if first in d:
d[first].update(v)
else:
d[first] = v
return d
def delete_2level_dict(d, first_key, second_key):
first = d.pop(first_key)
second = first.pop(second_key)
if len(first) > 0:
d.update({first_key: first})
return d, second
@@ -0,0 +1,9 @@
# -*- coding: utf-8 -*-
import re
def is_valid_filename(filename):
if re.match('^[A-Za-z0-9_@]+$', filename):
return True
else:
return False
@@ -0,0 +1,11 @@
# -*- coding: utf-8 -*-
import yaml
def save_yaml(data, file_path):
with open(file_path, 'w') as f_out:
yaml.dump(data,
f_out,
encoding='utf-8',
allow_unicode=True,
default_flow_style=False)

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