diff --git a/asset/images/inpainting_text/ex3_scene_im.jpg b/asset/images/inpainting_text/ex3_scene_im.jpg new file mode 100644 index 0000000..ef76b8f Binary files /dev/null and b/asset/images/inpainting_text/ex3_scene_im.jpg differ diff --git a/asset/images/inpainting_text/ex3_scene_mask.jpg b/asset/images/inpainting_text/ex3_scene_mask.jpg new file mode 100644 index 0000000..70ec5a5 Binary files /dev/null and b/asset/images/inpainting_text/ex3_scene_mask.jpg differ diff --git a/asset/images/inpainting_text/ex3_scene_mask2.jpg b/asset/images/inpainting_text/ex3_scene_mask2.jpg new file mode 100644 index 0000000..159bd6f Binary files /dev/null and b/asset/images/inpainting_text/ex3_scene_mask2.jpg differ diff --git a/asset/images/inpainting_text/inpainting_text.jpeg b/asset/images/inpainting_text/inpainting_text.jpeg new file mode 100644 index 0000000..174dba0 Binary files /dev/null and b/asset/images/inpainting_text/inpainting_text.jpeg differ diff --git a/asset/images/inpainting_text/inpainting_text2.jpeg b/asset/images/inpainting_text/inpainting_text2.jpeg new file mode 100644 index 0000000..7f827ad Binary files /dev/null and b/asset/images/inpainting_text/inpainting_text2.jpeg differ diff --git a/asset/images/inpainting_text_ref/ex4_scene_im.jpg b/asset/images/inpainting_text_ref/ex4_scene_im.jpg new file mode 100644 index 0000000..6233409 Binary files /dev/null and b/asset/images/inpainting_text_ref/ex4_scene_im.jpg differ diff --git a/asset/images/inpainting_text_ref/ex4_scene_mask.jpg b/asset/images/inpainting_text_ref/ex4_scene_mask.jpg new file mode 100644 index 0000000..9c1fe96 Binary files /dev/null and b/asset/images/inpainting_text_ref/ex4_scene_mask.jpg differ diff --git a/asset/images/inpainting_text_ref/ex4_subject_im.jpg b/asset/images/inpainting_text_ref/ex4_subject_im.jpg new file mode 100644 index 0000000..27f0835 Binary files /dev/null and b/asset/images/inpainting_text_ref/ex4_subject_im.jpg differ diff --git a/asset/images/inpainting_text_ref/ex4_subject_mask.jpg b/asset/images/inpainting_text_ref/ex4_subject_mask.jpg new file mode 100644 index 0000000..e3124cf Binary files /dev/null and b/asset/images/inpainting_text_ref/ex4_subject_mask.jpg differ diff --git a/asset/images/inpainting_text_ref/inpainting_text_ref.jpeg b/asset/images/inpainting_text_ref/inpainting_text_ref.jpeg new file mode 100644 index 0000000..9f61239 Binary files /dev/null and b/asset/images/inpainting_text_ref/inpainting_text_ref.jpeg differ diff --git a/asset/images/largen.gif b/asset/images/largen.gif new file mode 100644 index 0000000..bd07415 Binary files /dev/null and b/asset/images/largen.gif differ diff --git a/asset/images/virtual_try_on/ex2_scene_mask.jpg b/asset/images/virtual_try_on/ex2_scene_mask.jpg new file mode 100644 index 0000000..f800425 Binary files /dev/null and b/asset/images/virtual_try_on/ex2_scene_mask.jpg differ diff --git a/asset/images/virtual_try_on/ex2_subject_mask.jpg b/asset/images/virtual_try_on/ex2_subject_mask.jpg new file mode 100644 index 0000000..58678c0 Binary files /dev/null and b/asset/images/virtual_try_on/ex2_subject_mask.jpg differ diff --git a/asset/images/virtual_try_on/model.jpg b/asset/images/virtual_try_on/model.jpg new file mode 100644 index 0000000..d0b589e Binary files /dev/null and b/asset/images/virtual_try_on/model.jpg differ diff --git a/asset/images/virtual_try_on/try_on_out.jpeg b/asset/images/virtual_try_on/try_on_out.jpeg new file mode 100644 index 0000000..35c81e8 Binary files /dev/null and b/asset/images/virtual_try_on/try_on_out.jpeg differ diff --git a/asset/images/virtual_try_on/tshirt.jpg b/asset/images/virtual_try_on/tshirt.jpg new file mode 100644 index 0000000..462db84 Binary files /dev/null and b/asset/images/virtual_try_on/tshirt.jpg differ diff --git a/asset/images/zoom_out/ex1_scene_im.jpeg b/asset/images/zoom_out/ex1_scene_im.jpeg new file mode 100644 index 0000000..089e453 Binary files /dev/null and b/asset/images/zoom_out/ex1_scene_im.jpeg differ diff --git a/asset/images/zoom_out/ex1_zoom_out1.jpeg b/asset/images/zoom_out/ex1_zoom_out1.jpeg new file mode 100644 index 0000000..ccd5cf4 Binary files /dev/null and b/asset/images/zoom_out/ex1_zoom_out1.jpeg differ diff --git a/asset/images/zoom_out/ex1_zoom_out2.jpeg b/asset/images/zoom_out/ex1_zoom_out2.jpeg new file mode 100644 index 0000000..25e1c10 Binary files /dev/null and b/asset/images/zoom_out/ex1_zoom_out2.jpeg differ diff --git a/asset/images/zoom_out/ex1_zoom_out3.jpeg b/asset/images/zoom_out/ex1_zoom_out3.jpeg new file mode 100644 index 0000000..662f996 Binary files /dev/null and b/asset/images/zoom_out/ex1_zoom_out3.jpeg differ diff --git a/asset/images/zoom_out/ex1_zoom_out4.jpeg b/asset/images/zoom_out/ex1_zoom_out4.jpeg new file mode 100644 index 0000000..8911c1d Binary files /dev/null and b/asset/images/zoom_out/ex1_zoom_out4.jpeg differ diff --git a/example/__init__.py b/example/__init__.py new file mode 100644 index 0000000..c2df617 --- /dev/null +++ b/example/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +from . import impls diff --git a/example/example.yaml b/example/example.yaml new file mode 100644 index 0000000..59cef89 --- /dev/null +++ b/example/example.yaml @@ -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"] diff --git a/example/impls/__init__.py b/example/impls/__init__.py new file mode 100644 index 0000000..945fb86 --- /dev/null +++ b/example/impls/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +from . import dataset diff --git a/example/impls/dataset/__init__.py b/example/impls/dataset/__init__.py new file mode 100644 index 0000000..1635bb3 --- /dev/null +++ b/example/impls/dataset/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +from .classifier_dataset import ImageClassifyExampleDataset diff --git a/example/impls/dataset/classifier_dataset.py b/example/impls/dataset/classifier_dataset.py new file mode 100644 index 0000000..4d8334a --- /dev/null +++ b/example/impls/dataset/classifier_dataset.py @@ -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) diff --git a/example/run.py b/example/run.py new file mode 100644 index 0000000..4ed5f26 --- /dev/null +++ b/example/run.py @@ -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() diff --git a/readme.md b/readme.md index 5b8f7b1..44afbc2 100644 --- a/readme.md +++ b/readme.md @@ -9,8 +9,8 @@
## 📖 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) [](https://arxiv.org/abs/2312.11392) [](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) [](https://arxiv.org/abs/2310.19859) [](https://res-tuning.github.io/) - -## 🎉 News -- [2024.02]: We release new SCEdit controllable image synthesis models for SD v2.1 and SD XL. Multiple strategies applied to accelerate inference time for SCEPTER Studio. -- [2024.01]: We release **SCEPTER Studio**, an integrated toolkit for data management, model training and inference based on [Gradio](https://www.gradio.app/). -- [2024.01]: [SCEdit](https://arxiv.org/abs/2312.11392) support controllable image synthesis for training and inference. -- [2023.12]: We propose [SCEdit](https://arxiv.org/abs/2312.11392), an efficient and controllable generation framework. -- [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) [](https://arxiv.org/abs/2312.11392) [](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) [](https://arxiv.org/abs/2310.19859) [](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) [](https://arxiv.org/abs/2403.19534) [](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. +
+
+
| Origin Image Prompt: a temple on fire |
+ Zoom-Out CenterAround:0.75 |
+ Zoom-Out CenterAround:0.75 |
+ Zoom-Out CenterAround:0.75 |
+ Zoom-Out CenterAround:0.75 |
+
![]() |
+ ![]() |
+ ![]() |
+ ![]() |
+ ![]() |
+
| Model Image | +Model Mask | +Clothing Image | +Clothing Mask | +Try-on Output | +
![]() |
+ ![]() |
+ ![]() |
+ ![]() |
+ ![]() |
+
| Origin Image Prompt: a blue and white porcelain |
+ Inpainting Mask1 | +Inpainting Output1 | +Inpainting Mask2 Prompt: a clock |
+ Inpainting Output2 | +
![]() |
+ ![]() |
+ ![]() |
+ ![]() |
+ ![]() |
+
| Origin Image Prompt: a dog wearing sunglasses |
+ Origin Mask | +Reference Image | +Reference Mask | +Inpainting Output | +
![]() |
+ ![]() |
+ ![]() |
+ ![]() |
+ ![]() |
+
| '
f' {save_id}-{idx}|{one_label} | '
diff --git a/scepter/studio/inference/inference.py b/scepter/studio/inference/inference.py
index 789799a..d71dbc7 100644
--- a/scepter/studio/inference/inference.py
+++ b/scepter/studio/inference/inference.py
@@ -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)
diff --git a/scepter/studio/inference/inference_manager/infer_runer.py b/scepter/studio/inference/inference_manager/infer_runer.py
index 382e737..deb4c75 100644
--- a/scepter/studio/inference/inference_manager/infer_runer.py
+++ b/scepter/studio/inference/inference_manager/infer_runer.py
@@ -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)
diff --git a/scepter/studio/inference/inference_ui/component_names.py b/scepter/studio/inference/inference_ui/component_names.py
index 98c9872..3e5b2a3 100644
--- a/scepter/studio/inference/inference_ui/component_names.py
+++ b/scepter/studio/inference/inference_ui/component_names.py
@@ -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
+ ],
+ ]
diff --git a/scepter/studio/inference/inference_ui/gallery_ui.py b/scepter/studio/inference/inference_ui/gallery_ui.py
index 2c50936..74b3544 100644
--- a/scepter/studio/inference/inference_ui/gallery_ui.py
+++ b/scepter/studio/inference/inference_ui/gallery_ui.py
@@ -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,
diff --git a/scepter/studio/inference/inference_ui/largen_ui.py b/scepter/studio/inference/inference_ui/largen_ui.py
new file mode 100644
index 0000000..4012ad6
--- /dev/null
+++ b/scepter/studio/inference/inference_ui/largen_ui.py
@@ -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
diff --git a/scepter/studio/inference/inference_ui/mantra_ui.py b/scepter/studio/inference/inference_ui/mantra_ui.py
index 1be6c42..993a020 100644
--- a/scepter/studio/inference/inference_ui/mantra_ui.py
+++ b/scepter/studio/inference/inference_ui/mantra_ui.py
@@ -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)
diff --git a/scepter/studio/inference/inference_ui/model_manage_ui.py b/scepter/studio/inference/inference_ui/model_manage_ui.py
index c3f4643..5286673 100644
--- a/scepter/studio/inference/inference_ui/model_manage_ui.py
+++ b/scepter/studio/inference/inference_ui/model_manage_ui.py
@@ -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)
diff --git a/scepter/studio/inference/inference_ui/tuner_ui.py b/scepter/studio/inference/inference_ui/tuner_ui.py
index 4ca4f06..0c93973 100644
--- a/scepter/studio/inference/inference_ui/tuner_ui.py
+++ b/scepter/studio/inference/inference_ui/tuner_ui.py
@@ -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
+ ])
diff --git a/scepter/studio/preprocess/caption_editor_ui/create_dataset_ui.py b/scepter/studio/preprocess/caption_editor_ui/create_dataset_ui.py
index ed3302a..7128030 100644
--- a/scepter/studio/preprocess/caption_editor_ui/create_dataset_ui.py
+++ b/scepter/studio/preprocess/caption_editor_ui/create_dataset_ui.py
@@ -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)
diff --git a/scepter/studio/preprocess/preprocess.py b/scepter/studio/preprocess/preprocess.py
index 78f1920..fe54798 100644
--- a/scepter/studio/preprocess/preprocess.py
+++ b/scepter/studio/preprocess/preprocess.py
@@ -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)
diff --git a/scepter/studio/self_train/scripts/run_task.py b/scepter/studio/self_train/scripts/run_task.py
index 9003fc1..a9352f9 100644
--- a/scepter/studio/self_train/scripts/run_task.py
+++ b/scepter/studio/self_train/scripts/run_task.py
@@ -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()
diff --git a/scepter/studio/self_train/scripts/sleep.py b/scepter/studio/self_train/scripts/sleep.py
new file mode 100644
index 0000000..efbbb0e
--- /dev/null
+++ b/scepter/studio/self_train/scripts/sleep.py
@@ -0,0 +1,6 @@
+# -*- coding: utf-8 -*-
+import time
+
+for i in range(180):
+ time.sleep(1)
+ print('sleep', i)
diff --git a/scepter/studio/self_train/scripts/trainer.py b/scepter/studio/self_train/scripts/trainer.py
new file mode 100644
index 0000000..f4b4988
--- /dev/null
+++ b/scepter/studio/self_train/scripts/trainer.py
@@ -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)
diff --git a/scepter/studio/self_train/self_train.py b/scepter/studio/self_train/self_train.py
index dc375ae..d2d4727 100644
--- a/scepter/studio/self_train/self_train.py
+++ b/scepter/studio/self_train/self_train.py
@@ -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__':
diff --git a/scepter/studio/self_train/self_train_ui/component_names.py b/scepter/studio/self_train/self_train_ui/component_names.py
index 9e94bb2..034ae9a 100644
--- a/scepter/studio/self_train/self_train_ui/component_names.py
+++ b/scepter/studio/self_train/self_train_ui/component_names.py
@@ -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 = '目前显存不足,训练失败!'
diff --git a/scepter/studio/self_train/self_train_ui/inference_ui.py b/scepter/studio/self_train/self_train_ui/inference_ui.py
deleted file mode 100644
index e7c0dde..0000000
--- a/scepter/studio/self_train/self_train_ui/inference_ui.py
+++ /dev/null
@@ -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)
diff --git a/scepter/studio/self_train/self_train_ui/model_ui.py b/scepter/studio/self_train/self_train_ui/model_ui.py
new file mode 100644
index 0000000..d2de96d
--- /dev/null
+++ b/scepter/studio/self_train/self_train_ui/model_ui.py
@@ -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)
diff --git a/scepter/studio/self_train/self_train_ui/trainer_ui.py b/scepter/studio/self_train/self_train_ui/trainer_ui.py
index bd54217..1cfd846 100644
--- a/scepter/studio/self_train/self_train_ui/trainer_ui.py
+++ b/scepter/studio/self_train/self_train_ui/trainer_ui.py
@@ -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)
diff --git a/scepter/studio/self_train/utils/config_parser.py b/scepter/studio/self_train/utils/config_parser.py
index 519a377..e9add78 100644
--- a/scepter/studio/self_train/utils/config_parser.py
+++ b/scepter/studio/self_train/utils/config_parser.py
@@ -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
diff --git a/scepter/studio/tuner_manager/__init__.py b/scepter/studio/tuner_manager/__init__.py
new file mode 100644
index 0000000..e69de29
diff --git a/scepter/studio/tuner_manager/manager_ui/__init__.py b/scepter/studio/tuner_manager/manager_ui/__init__.py
new file mode 100644
index 0000000..e69de29
diff --git a/scepter/studio/tuner_manager/manager_ui/browser_ui.py b/scepter/studio/tuner_manager/manager_ui/browser_ui.py
new file mode 100644
index 0000000..ba7f179
--- /dev/null
+++ b/scepter/studio/tuner_manager/manager_ui/browser_ui.py
@@ -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)
diff --git a/scepter/studio/tuner_manager/manager_ui/component_names.py b/scepter/studio/tuner_manager/manager_ui/component_names.py
new file mode 100644
index 0000000..a103cfa
--- /dev/null
+++ b/scepter/studio/tuner_manager/manager_ui/component_names.py
@@ -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 = '删除'
diff --git a/scepter/studio/tuner_manager/manager_ui/info_ui.py b/scepter/studio/tuner_manager/manager_ui/info_ui.py
new file mode 100644
index 0000000..16c1c3c
--- /dev/null
+++ b/scepter/studio/tuner_manager/manager_ui/info_ui.py
@@ -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
diff --git a/scepter/studio/tuner_manager/tuner_manager.py b/scepter/studio/tuner_manager/tuner_manager.py
new file mode 100644
index 0000000..185b65a
--- /dev/null
+++ b/scepter/studio/tuner_manager/tuner_manager.py
@@ -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)
diff --git a/scepter/studio/tuner_manager/utils/__init__.py b/scepter/studio/tuner_manager/utils/__init__.py
new file mode 100644
index 0000000..e69de29
diff --git a/scepter/studio/tuner_manager/utils/dict.py b/scepter/studio/tuner_manager/utils/dict.py
new file mode 100644
index 0000000..5ed6dc4
--- /dev/null
+++ b/scepter/studio/tuner_manager/utils/dict.py
@@ -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
diff --git a/scepter/studio/tuner_manager/utils/path.py b/scepter/studio/tuner_manager/utils/path.py
new file mode 100644
index 0000000..cb7b58c
--- /dev/null
+++ b/scepter/studio/tuner_manager/utils/path.py
@@ -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
diff --git a/scepter/studio/tuner_manager/utils/yaml.py b/scepter/studio/tuner_manager/utils/yaml.py
new file mode 100644
index 0000000..cf9813d
--- /dev/null
+++ b/scepter/studio/tuner_manager/utils/yaml.py
@@ -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)
diff --git a/scepter/tools/helper.py b/scepter/tools/helper.py
index 4dadeb1..c385172 100644
--- a/scepter/tools/helper.py
+++ b/scepter/tools/helper.py
@@ -1,12 +1,19 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import argparse
+import importlib
import os
import sys
#
from scepter.modules.utils.registry import REGISTRY_LIST
+if os.path.exists('__init__.py'):
+ package_name = 'scepter_ext'
+ spec = importlib.util.spec_from_file_location(package_name, '__init__.py')
+ package = importlib.util.module_from_spec(spec)
+ sys.modules[package_name] = package
+ spec.loader.exec_module(package)
sys.path.insert(0, os.path.abspath(os.curdir))
diff --git a/scepter/tools/run_inference.py b/scepter/tools/run_inference.py
index fcb5b33..5bb9bdc 100644
--- a/scepter/tools/run_inference.py
+++ b/scepter/tools/run_inference.py
@@ -1,7 +1,9 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import argparse
+import importlib
import os
+import sys
import numpy as np
import torch
@@ -16,6 +18,13 @@ from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.logger import get_logger
+if os.path.exists('__init__.py'):
+ package_name = 'scepter_ext'
+ spec = importlib.util.spec_from_file_location(package_name, '__init__.py')
+ package = importlib.util.module_from_spec(spec)
+ sys.modules[package_name] = package
+ spec.loader.exec_module(package)
+
def run_task(cfg):
std_logger = get_logger(name='scepter')
diff --git a/scepter/tools/run_train.py b/scepter/tools/run_train.py
index e369035..4cd54ff 100644
--- a/scepter/tools/run_train.py
+++ b/scepter/tools/run_train.py
@@ -1,12 +1,22 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import argparse
+import importlib
+import os
+import sys
from scepter.modules.solver.registry import SOLVERS
from scepter.modules.utils.config import Config
from scepter.modules.utils.distribute import we
from scepter.modules.utils.logger import get_logger
+if os.path.exists('__init__.py'):
+ package_name = 'scepter_ext'
+ spec = importlib.util.spec_from_file_location(package_name, '__init__.py')
+ package = importlib.util.module_from_spec(spec)
+ sys.modules[package_name] = package
+ spec.loader.exec_module(package)
+
def run_task(cfg):
std_logger = get_logger(name='scepter')
@@ -18,19 +28,21 @@ def run_task(cfg):
def update_config(cfg):
if hasattr(cfg.args, 'learning_rate') and cfg.args.learning_rate:
- print(
- f'learning_rate change from {cfg.SOLVER.OPTIMIZER.LEARNING_RATE} to {cfg.args.learning_rate}'
- )
+ if cfg.SOLVER.OPTIMIZER.get('LEARNING_RATE', None) is not None:
+ print(
+ f'learning_rate change from {cfg.SOLVER.OPTIMIZER.LEARNING_RATE} to {cfg.args.learning_rate}'
+ )
cfg.SOLVER.OPTIMIZER.LEARNING_RATE = float(cfg.args.learning_rate)
if hasattr(cfg.args, 'max_steps') and cfg.args.max_steps:
- print(
- f'max_steps change from {cfg.SOLVER.MAX_STEPS} to {cfg.args.max_steps}'
- )
+ if cfg.SOLVER.get('MAX_STEPS', None) is not None:
+ print(
+ f'max_steps change from {cfg.SOLVER.MAX_STEPS} to {cfg.args.max_steps}'
+ )
cfg.SOLVER.MAX_STEPS = int(cfg.args.max_steps)
return cfg
-if __name__ == '__main__':
+def run():
parser = argparse.ArgumentParser(description='Argparser for Scepter:\n')
parser.add_argument('--learning_rate',
dest='learning_rate',
@@ -44,3 +56,7 @@ if __name__ == '__main__':
cfg = Config(load=True, parser_ins=parser)
cfg = update_config(cfg)
we.init_env(cfg, logger=None, fn=run_task)
+
+
+if __name__ == '__main__':
+ run()
diff --git a/scepter/tools/webui.py b/scepter/tools/webui.py
index d084f7a..c6b9682 100644
--- a/scepter/tools/webui.py
+++ b/scepter/tools/webui.py
@@ -2,8 +2,10 @@
# Copyright (c) Alibaba, Inc. and its affiliates.
import argparse
import datetime
+import importlib
import os
import random
+import sys
import gradio as gr
@@ -12,6 +14,13 @@ from scepter.modules.utils.config import Config
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.logger import get_logger, init_logger
+if os.path.exists('__init__.py'):
+ package_name = 'scepter_ext'
+ spec = importlib.util.spec_from_file_location(package_name, '__init__.py')
+ package = importlib.util.module_from_spec(spec)
+ sys.modules[package_name] = package
+ spec.loader.exec_module(package)
+
def prepare(config):
if 'FILE_SYSTEM' in config:
@@ -57,6 +66,13 @@ if __name__ == '__main__':
default='en',
help='Now we only support english(en) and chinese(zh)')
args = parser.parse_args()
+ if not os.path.exists(args.config):
+ print(
+ f"{args.config} doesn't exist, find this file in {os.path.dirname(scepter.dirname)}"
+ )
+ args.config = os.path.join(os.path.dirname(scepter.dirname),
+ args.config)
+ assert os.path.exists(args.config)
config = Config(load=True, cfg_file=args.config)
prepare(config)
@@ -74,24 +90,34 @@ if __name__ == '__main__':
interface = None
if ifid == 'home':
from scepter.studio.home.home import HomeUI
+
interface = HomeUI(info['CONFIG'],
is_debug=args.debug,
language=args.language,
root_work_dir=config.WORK_DIR)
if ifid == 'preprocess':
from scepter.studio.preprocess.preprocess import PreprocessUI
+
interface = PreprocessUI(info['CONFIG'],
is_debug=args.debug,
language=args.language,
root_work_dir=config.WORK_DIR)
if ifid == 'self_train':
from scepter.studio.self_train.self_train import SelfTrainUI
+
interface = SelfTrainUI(info['CONFIG'],
is_debug=args.debug,
language=args.language,
root_work_dir=config.WORK_DIR)
+ if ifid == 'tuner_manager':
+ from scepter.studio.tuner_manager.tuner_manager import TunerManagerUI
+ interface = TunerManagerUI(info['CONFIG'],
+ is_debug=args.debug,
+ language=args.language,
+ root_work_dir=config.WORK_DIR)
if ifid == 'inference':
from scepter.studio.inference.inference import InferenceUI
+
interface = InferenceUI(info['CONFIG'],
is_debug=args.debug,
language=args.language,
@@ -109,6 +135,8 @@ if __name__ == '__main__':
gr.Markdown(
f"