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) [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=SCEdit&color=red&logo=arxiv)](https://arxiv.org/abs/2312.11392) [![Page link](https://img.shields.io/badge/Page-SCEdit-Gree)](https://scedit.github.io/) -3. Res-Tuning(TODO): [Res-Tuning: A Flexible and Efficient Tuning Paradigm via Unbinding Tuner from Backbone](https://arxiv.org/abs/2310.19859) [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=ResTuning&color=red&logo=arxiv)](https://arxiv.org/abs/2310.19859) [![Page link](https://img.shields.io/badge/Page-ResTuning-Gree)](https://res-tuning.github.io/) - -## 🎉 News -- [2024.02]: We release new SCEdit controllable image synthesis models for SD v2.1 and SD XL. Multiple strategies applied to accelerate inference time for SCEPTER Studio. -- [2024.01]: We release **SCEPTER Studio**, an integrated toolkit for data management, model training and inference based on [Gradio](https://www.gradio.app/). -- [2024.01]: [SCEdit](https://arxiv.org/abs/2312.11392) support controllable image synthesis for training and inference. -- [2023.12]: We propose [SCEdit](https://arxiv.org/abs/2312.11392), an efficient and controllable generation framework. -- [2023.12]: We release [🪄SCEPTER](https://github.com/modelscope/scepter/) library. +2. SCEdit(CVPR2024): [SCEdit: Efficient and Controllable Image Diffusion Generation via Skip Connection Editing](https://arxiv.org/abs/2312.11392) [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=SCEdit&color=red&logo=arxiv)](https://arxiv.org/abs/2312.11392) [![Page link](https://img.shields.io/badge/Page-SCEdit-Gree)](https://scedit.github.io/) +3. Res-Tuning(NeurIPS2023 TODO): [Res-Tuning: A Flexible and Efficient Tuning Paradigm via Unbinding Tuner from Backbone](https://arxiv.org/abs/2310.19859) [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=ResTuning&color=red&logo=arxiv)](https://arxiv.org/abs/2310.19859) [![Page link](https://img.shields.io/badge/Page-ResTuning-Gree)](https://res-tuning.github.io/) +4. LAR-Gen: [Locate, Assign, Refine: Taming Customized Image Inpainting with Text-Subject Guidance](https://arxiv.org/abs/2403.19534) [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=LARGen&color=red&logo=arxiv)](https://arxiv.org/abs/2403.19534) [![Page link](https://img.shields.io/badge/Page-LARGen-Gree)](https://ali-vilab.github.io/largen-page/) ## 🛠️ Installation @@ -94,7 +97,7 @@ For the data format used by SCEPTER Studio, please refer to [3D_example_csv.zip] To facilitate starting training in command-line mode, you can use a dataset in text format, please refer to [3D_example_txt.zip](https://modelscope.cn/api/v1/models/damo/scepter/repo?Revision=master&FilePath=datasets/3D_example_txt.zip) ```shell -mkdir -p cache/dataset/ && wget 'https://modelscope.cn/api/v1/models/damo/scepter_scedit/repo?Revision=master&FilePath=dataset/3D_example_txt.zip' -O cache/dataset/3D_example_txt.zip && unzip cache/dataset/3D_example_txt.zip -d cache/dataset/ && rm cache/dataset/3D_example_txt.zip +mkdir -p cache/datasets/ && wget 'https://modelscope.cn/api/v1/models/damo/scepter_scedit/repo?Revision=master&FilePath=dataset/3D_example_txt.zip' -O cache/datasets/3D_example_txt.zip && unzip cache/datasets/3D_example_txt.zip -d cache/datasets/ && rm cache/datasets/3D_example_txt.zip ``` ### Training @@ -165,6 +168,14 @@ python scepter/tools/run_inference.py --cfg scepter/methods/scedit/ctr/sd21_768_ python scepter/tools/run_inference.py --cfg scepter/methods/scedit/ctr/sd21_768_sce_ctr_pose.yaml --num_samples 1 --prompt 'super mario' --save_folder 'test_mario_pose' --image_size 768 --task control --image 'asset/images/pose_source.png' --control_mode source --pretrained_model ms://damo/scepter_scedit@controllable_model/SD2.1/pose_control/0_SwiftSCETuning/pytorch_model.bin # pose ``` +### Customize Modules +Refer to `example`, build your modules in `example/impls`. +```python +cd example +python run.py --cfg example.yaml +``` + + ## 🖥️ SCEPTER Studio ### Launch @@ -181,16 +192,94 @@ git clone https://github.com/modelscope/scepter.git PYTHONPATH=. python scepter/tools/webui.py --cfg scepter/methods/studio/scepter_ui.yaml ``` -The startup of **SCEPTER Studio** eliminates the need for manual downloading and organizing of models; it will automatically load the corresponding models and store them in a local directory. -Depending on the network and hardware situation, the initial startup usually requires 15-60 minutes, primarily involving the download and processing of SDv1.5, SDv2.1, and SDXL models. +The startup of **SCEPTER Studio** eliminates the need for manual downloading and organizing of models; it will automatically load the corresponding models and store them in a local directory. +Depending on the network and hardware situation, the initial startup usually requires 15-60 minutes, primarily involving the download and processing of SDv1.5, SDv2.1, and SDXL models. Therefore, subsequent startups will become much faster (about one minute) as downloading is no longer required. +* LAR-Gen: we release `zoom-out`, `virtual try on`, `inpainting(text guided)`, `inpainting(text + reference image guided)` image editing capabilities. +Please note that the **Data Preprocess** button must be clicked before clicking the **Generate** button. +

+ +

+ ### Modelscope Studio We deploy a work studio on Modelscope that includes only the inference tab, please refer to [ms_scepter_studio](https://www.modelscope.cn/studios/damo/scepter_studio/summary) ## 🖼️ Gallery +### LAR-Gen: Zoom-Out + + + + + + + + + + + + + + + +
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
+ +### LAR-Gen: Virtual-Try-on + + + + + + + + + + + + + + + +
Model ImageModel MaskClothing ImageClothing MaskTry-on Output
+ +### LAR-Gen: Inpainting(Text-guided) + + + + + + + + + + + + + + + +
Origin Image
Prompt: a blue and white porcelain
Inpainting Mask1Inpainting Output1Inpainting Mask2
Prompt: a clock
Inpainting Output2
+ +### LAR-Gen: Inpainting(Text + Reference Image Guided) + + + + + + + + + + + + + + + +
Origin Image
Prompt: a dog wearing sunglasses
Origin MaskReference ImageReference MaskInpainting Output
+ ### Dragon Year Special: Dragon Tuner @@ -244,22 +333,31 @@ We deploy a work studio on Modelscope that includes only the inference tab, plea | SD 2.1 | 🪄 | 🪄 | 🪄 | 🪄 | 🪄 | | SD XL | 🪄 | 🪄 | 🪄 | 🪄 | 🪄 | +### Image Editing +- LAR-Gen + +| **Model** | **Locate** | **Assign** | **Refine** | +|:---------:|:----------:|:----------:|:----------:| +| SD XL | 🪄 | 🪄 | ⏳ | + ### Model URL - ✅ indicates support for both training and inference. - 🪄 denotes that the model has been published. +- ⏳ denotes that the module has not been integrated currently. - More models will be released in the future. | Model | URL | |--------|------------------------------------------------------------------------------------------------------------------------------------------------| -| SCEdit | [ModelScope](https://modelscope.cn/models/damo/scepter_scedit/summary) [HuggingFace](https://huggingface.co/scepter-studio/scepter_scedit) | +| SCEdit | [ModelScope](https://modelscope.cn/models/iic/scepter_scedit/summary) [HuggingFace](https://huggingface.co/scepter-studio/scepter_scedit) | +| LAR-Gen | [ModelScope](https://www.modelscope.cn/models/iic/LARGEN/summary) | PS: Scripts running within the SCEPTER framework will automatically fetch and load models based on the required dependency files, eliminating the need for manual downloads. ## 🔍 Learn More -- [Alibaba TongYi Vision Intelligence Lab](https://github.com/damo-vilab) +- [Alibaba TongYi Vision Intelligence Lab](https://github.com/ali-vilab) Discover more about open-source projects on image generation, video generation, and editing tasks. @@ -271,7 +369,20 @@ PS: Scripts running within the SCEPTER framework will automatically fetch and lo SWIFT (Scalable lightWeight Infrastructure for Fine-Tuning) is an extensible framwork designed to faciliate lightweight model fine-tuning and inference. +## BibTeX +If our work is useful for your research, please consider citing: +```bibtex +@misc{scepter, + title = {SCEPTER, https://github.com/modelscope/scepter}, + author = {SCEPTER}, + year = {2023} +} +``` + ## License This project is licensed under the [Apache License (Version 2.0)](https://github.com/modelscope/modelscope/blob/master/LICENSE). + +## Acknowledgement +Thanks to [Stability-AI](https://github.com/Stability-AI), [SWIFT library](https://github.com/modelscope/swift/) and [Fooocus](https://github.com/lllyasviel/Fooocus) for their awesome work. diff --git a/requirements/framework.txt b/requirements/framework.txt index 3e12b56..3019164 100644 --- a/requirements/framework.txt +++ b/requirements/framework.txt @@ -1,3 +1,5 @@ +albumentations +bezier einops modelscope ms-swift>=1.5.2 @@ -6,6 +8,7 @@ open_clip_torch opencv-python opencv_transforms>=0.0.6 oss2>=2.15.0 +pycocotools pyyaml>=5.3.1 scikit-image torchsde diff --git a/requirements/recommended.txt b/requirements/recommended.txt index d25969c..e976f5c 100644 --- a/requirements/recommended.txt +++ b/requirements/recommended.txt @@ -1,3 +1,4 @@ +git+https://github.com/cocodataset/panopticapi.git torch==2.0.1 torchvision==0.15.2 xformers==0.0.21 diff --git a/requirements/scepter_studio.txt b/requirements/scepter_studio.txt index a83a6e0..029ab97 100644 --- a/requirements/scepter_studio.txt +++ b/requirements/scepter_studio.txt @@ -1,2 +1,3 @@ gradio>=3.47.1,<4.0.0 imagehash +psutil diff --git a/scepter/methods/studio/inference/inference.yaml b/scepter/methods/studio/inference/inference.yaml index 542b616..286a418 100644 --- a/scepter/methods/studio/inference/inference.yaml +++ b/scepter/methods/studio/inference/inference.yaml @@ -79,6 +79,7 @@ EXTENSION_PARAS: MANTRA_BOOK: scepter/methods/studio/extensions/mantra_book/mantra_book.yaml OFFICIAL_TUNERS: scepter/methods/studio/extensions/tuners/official_tuners.yaml OFFICIAL_CONTROLLERS: scepter/methods/studio/extensions/controllers/official_controllers.yaml + TUNER_MANAGER: scepter/methods/studio/tuner_manager/tuner_manager.yaml CONTROLABLE_ANNOTATORS: - NAME: "CannyAnnotator" diff --git a/scepter/methods/studio/inference/largen/largen_pro.yaml b/scepter/methods/studio/inference/largen/largen_pro.yaml new file mode 100644 index 0000000..aeaeca1 --- /dev/null +++ b/scepter/methods/studio/inference/largen/largen_pro.yaml @@ -0,0 +1,268 @@ +NAME: LARGEN +IS_DEFAULT: True +DEFAULT_PARAS: + PARAS: + RESOLUTIONS: [[1024, 1024]] + INPUT: + IMAGE: + ORIGINAL_SIZE_AS_TUPLE: [1024, 1024] + TARGET_SIZE_AS_TUPLE: [1024, 1024] + AESTHETIC_SCORE: 6.0 + NEGATIVE_AESTHETIC_SCORE: 2.5 + PROMPT: "" + NEGATIVE_PROMPT: "" + PROMPT_PREFIX: "" + CROP_COORDS_TOP_LEFT: [0, 0] + SAMPLE: ddim + SAMPLE_STEPS: 50 + GUIDE_SCALE: 7.5 + GUIDE_RESCALE: 0.5 + DISCRETIZATION: trailing + REFINE_SAMPLE: ddim + REFINE_GUIDE_SCALE: 7.5 + REFINE_GUIDE_RESCALE: 0.5 + REFINE_DISCRETIZATION: trailing + OUTPUT: + LATENT: + BEFORE_REFINE_IMAGES: + IMAGES: + SEED: + MODULES_PARAS: + FIRST_STAGE_MODEL: + FUNCTION: + - + NAME: encode + DTYPE: float32 + INPUT: ["IMAGE"] + - + NAME: decode + DTYPE: float32 + INPUT: ["LATENT"] + PARAS: + # SCALE_FACTOR DESCRIPTION: The vae embeding scale. TYPE: float default: 0.18215 + SCALE_FACTOR: 0.13025 + SIZE_FACTOR: 8 + DIFFUSION_MODEL: + FUNCTION: + - + NAME: forward + DTYPE: float16 + INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE", "GUIDE_RESCALE", "DISCRETIZATION"] + COND_STAGE_MODEL: + FUNCTION: + - + NAME: encode + DTYPE: float16 + INPUT: ["ORIGINAL_SIZE_AS_TUPLE", "CROP_COORDS_TOP_LEFT", "PROMPT", "NEGATIVE_PROMPT"] + REFINER_MODEL: + FUNCTION: + - + NAME: forward + DTYPE: float16 + INPUT: ["SAMPLE_STEPS", "REFINE_SAMPLE", "REFINE_GUIDE_SCALE", "REFINE_GUIDE_RESCALE", "REFINE_DISCRETIZATION"] + REFINER_COND_MODEL: + FUNCTION: + - + NAME: encode + DTYPE: float16 + INPUT: ["ORIGINAL_SIZE_AS_TUPLE", "AESTHETIC_SCORE", "NEGATIVE_AESTHETIC_SCORE", "CROP_COORDS_TOP_LEFT", "PROMPT", "NEGATIVE_PROMPT"] + +MODEL: + PRETRAINED_MODEL: ms://damo/LARGEN@models/largen_ckpt_s22k.pth + # SCHEDULE_ARGS DESCRIPTION: TYPE: default: '' + SCHEDULE: + PARAMETERIZATION: "eps" + TIMESTEPS: 1000 + ZERO_TERMINAL_SNR: False + SCHEDULE_ARGS: + # NAME DESCRIPTION: TYPE: default: '' + NAME: "scaled_linear" + BETA_MIN: 0.00085 + BETA_MAX: 0.0120 + # DIFFUSION_MODEL DESCRIPTION: TYPE: default: '' + DIFFUSION_MODEL: + # NAME DESCRIPTION: TYPE: default: 'DiffusionUNetXL' + NAME: LargenUNetXL + # PRETRAINED_MODEL DESCRIPTION: Whole model's pretrained model path. TYPE: NoneType default: None + PRETRAINED_MODEL: + # IN_CHANNELS DESCRIPTION: Unet channels for input, considering the input image's channels. TYPE: int default: 4 + IN_CHANNELS: 9 + # OUT_CHANNELS DESCRIPTION: Unet channels for output, considering the input image's channels. TYPE: int default: 4 + OUT_CHANNELS: 4 + # NUM_RES_BLOCKS DESCRIPTION: The blocks's number of res. TYPE: int default: 2 + NUM_RES_BLOCKS: 2 + # MODEL_CHANNELS DESCRIPTION: base channel count for the model. TYPE: int default: 320 + MODEL_CHANNELS: 320 + # ATTENTION_RESOLUTIONS DESCRIPTION: A collection of downsample rates at which attention will take place. May be a set, list, or tuple. For example, if this contains 4, then at 4x downsampling, attentio will be used. TYPE: list default: [4, 2] + ATTENTION_RESOLUTIONS: [4, 2] + # DROPOUT DESCRIPTION: The dropout rate. TYPE: int default: 0 + DROPOUT: 0 + # CHANNEL_MULT DESCRIPTION: channel multiplier for each level of the UNet. TYPE: list default: [1, 2, 4] + CHANNEL_MULT: [1, 2, 4] + # CONV_RESAMPLE DESCRIPTION: Use conv to resample when downsample. TYPE: bool default: True + CONV_RESAMPLE: True + # DIMS DESCRIPTION: The Conv dims which 2 represent Conv2D. TYPE: int default: 2 + DIMS: 2 + # NUM_CLASSES DESCRIPTION: The class num for class guided setting, also can be set as continuous. TYPE: str default: 'sequential' + NUM_CLASSES: sequential + # USE_CHECKPOINT DESCRIPTION: Use gradient checkpointing to reduce memory usage. TYPE: bool default: False + USE_CHECKPOINT: False + # NUM_HEADS DESCRIPTION: The number of attention heads in each attention layer. TYPE: int default: -1 + NUM_HEADS: -1 + # NUM_HEADS_CHANNELS DESCRIPTION: If specified, ignore num_heads and instead use a fixed channel width per attention head. TYPE: int default: 64 + NUM_HEADS_CHANNELS: 64 + # USE_SCALE_SHIFT_NORM DESCRIPTION: The scale and shift for the outnorm of RESBLOCK, use a FiLM-like conditioning mechanism. TYPE: bool default: False + USE_SCALE_SHIFT_NORM: False + # RESBLOCK_UPDOWN DESCRIPTION: Use residual blocks for up/downsampling, if False use Conv. TYPE: bool default: False + RESBLOCK_UPDOWN: False + # USE_NEW_ATTENTION_ORDER DESCRIPTION: Whether use new attention(qkv before split heads or not) or not. TYPE: bool default: True + USE_NEW_ATTENTION_ORDER: True + # USE_SPATIAL_TRANSFORMER DESCRIPTION: Custom transformer which support the context, if context_dim is not None, the parameter must set True TYPE: bool default: True + USE_SPATIAL_TRANSFORMER: True + # TRANSFORMER_DEPTH DESCRIPTION: Custom transformer's depth, valid when USE_SPATIAL_TRANSFORMER is True. TYPE: list default: [1, 2, 10] + TRANSFORMER_DEPTH: [1, 2, 10] + # TRANSFORMER_DEPTH_MIDDLE DESCRIPTION: Custom transformer's depth of middle block, If set None, use TRANSFORMER_DEPTH last value. TYPE: NoneType default: None + # TRANSFORMER_DEPTH_MIDDLE: None + # CONTEXT_DIM DESCRIPTION: Custom context info, if set, USE_SPATIAL_TRANSFORMER also set True. TYPE: int default: 2048 + CONTEXT_DIM: 2048 + # DISABLE_SELF_ATTENTIONS DESCRIPTION: Whether disable the self-attentions on some level, should be a list, [False, True, ...] TYPE: NoneType default: None + # DISABLE_SELF_ATTENTIONS: None + # NUM_ATTENTION_BLOCKS DESCRIPTION: The number of attention blocks for attention layer. TYPE: NoneType default: None + # NUM_ATTENTION_BLOCKS: None + # DISABLE_MIDDLE_SELF_ATTN DESCRIPTION: Whether disable the self-attentions in middle blocks. TYPE: bool default: False + DISABLE_MIDDLE_SELF_ATTN: False + # USE_LINEAR_IN_TRANSFORMER DESCRIPTION: Custom transformer's parameter, valid when USE_SPATIAL_TRANSFORMER is True. TYPE: bool default: True + USE_LINEAR_IN_TRANSFORMER: True + # ADM_IN_CHANNELS DESCRIPTION: Used when num_classes == 'sequential' or 'timestep'. TYPE: int default: 2816 + ADM_IN_CHANNELS: 2816 + # USE_SENTENCE_EMB DESCRIPTION: Used sentence emb or not, default False. TYPE: bool default: False + USE_SENTENCE_EMB: False + # USE_WORD_MAPPING DESCRIPTION: Used word mapping or not, default False. TYPE: bool default: False + USE_WORD_MAPPING: False + TRANSFORMER_BLOCK_TYPE: att_v2 + IMAGE_SCALE: 1.0 + USE_REFINE: False + # FIRST_STAGE_MODEL DESCRIPTION: TYPE: default: '' + FIRST_STAGE_MODEL: + NAME: AutoencoderKL + EMBED_DIM: 4 + IGNORE_KEYS: [ ] + BATCH_SIZE: 1 + # + ENCODER: + NAME: Encoder + CH: 128 + OUT_CH: 3 + NUM_RES_BLOCKS: 2 + IN_CHANNELS: 3 + ATTN_RESOLUTIONS: [ ] + CH_MULT: [ 1, 2, 4, 4 ] + Z_CHANNELS: 4 + DOUBLE_Z: True + DROPOUT: 0.0 + RESAMP_WITH_CONV: True + # + DECODER: + NAME: Decoder + CH: 128 + OUT_CH: 3 + NUM_RES_BLOCKS: 2 + IN_CHANNELS: 3 + ATTN_RESOLUTIONS: [ ] + CH_MULT: [ 1, 2, 4, 4 ] + Z_CHANNELS: 4 + DROPOUT: 0.0 + RESAMP_WITH_CONV: True + GIVE_PRE_END: False + TANH_OUT: False + # COND_STAGE_MODEL DESCRIPTION: TYPE: default: '' + COND_STAGE_MODEL: + # NAME DESCRIPTION: TYPE: default: 'GeneralConditioner' + NAME: GeneralConditioner + USE_GRAD: False + # EMBEDDERS DESCRIPTION: TYPE: default: '' + EMBEDDERS: + - + # NAME DESCRIPTION: TYPE: default: 'FrozenCLIPEmbedder' + NAME: FrozenCLIPEmbedder + # PRETRAINED_MODEL DESCRIPTION: TYPE: str default: '' + PRETRAINED_MODEL: ms://AI-ModelScope/clip-vit-large-patch14 + TOKENIZER_PATH: ms://AI-ModelScope/clip-vit-large-patch14 + # MAX_LENGTH DESCRIPTION: TYPE: int default: 77 + MAX_LENGTH: 77 + # FREEZE DESCRIPTION: TYPE: bool default: True + FREEZE: True + # LAYER DESCRIPTION: TYPE: str default: 'last' + LAYER: hidden + # LAYER_IDX DESCRIPTION: TYPE: NoneType default: None + LAYER_IDX: 11 + # USE_FINAL_LAYER_NORM DESCRIPTION: TYPE: bool default: False + USE_FINAL_LAYER_NORM: False + UCG_RATE: 0.0 + INPUT_KEYS: ["prompt"] + LEGACY_UCG_VALUE: + - + # NAME DESCRIPTION: TYPE: default: 'FrozenOpenCLIPEmbedder2' + NAME: FrozenOpenCLIPEmbedder2 + # ARCH DESCRIPTION: TYPE: str default: 'ViT-H-14' + ARCH: ViT-bigG-14 + # MAX_LENGTH DESCRIPTION: TYPE: int default: 77 + MAX_LENGTH: 77 + # FREEZE DESCRIPTION: TYPE: bool default: True + FREEZE: True + # ALWAYS_RETURN_POOLED DESCRIPTION: Whether always return pooled results or not ,default False. TYPE: bool default: False + ALWAYS_RETURN_POOLED: True + # LEGACY DESCRIPTION: Whether use legacy returnd feature or not ,default True. TYPE: bool default: True + LEGACY: False + # LAYER DESCRIPTION: TYPE: str default: 'last' + LAYER: penultimate + UCG_RATE: 0.0 + INPUT_KEYS: ["prompt"] + LEGACY_UCG_VALUE: + - + # NAME DESCRIPTION: TYPE: default: 'ConcatTimestepEmbedderND' + NAME: ConcatTimestepEmbedderND + # OUT_DIM DESCRIPTION: Output dim TYPE: int default: 256 + OUT_DIM: 256 + UCG_RATE: 0.0 + INPUT_KEYS: ["original_size_as_tuple"] + LEGACY_UCG_VALUE: + - + # NAME DESCRIPTION: TYPE: default: 'ConcatTimestepEmbedderND' + NAME: ConcatTimestepEmbedderND + # OUT_DIM DESCRIPTION: Output dim TYPE: int default: 256 + OUT_DIM: 256 + UCG_RATE: 0.0 + INPUT_KEYS: ["crop_coords_top_left"] + LEGACY_UCG_VALUE: + - + # NAME DESCRIPTION: TYPE: default: 'ConcatTimestepEmbedderND' + NAME: ConcatTimestepEmbedderND + # OUT_DIM DESCRIPTION: Output dim TYPE: int default: 256 + OUT_DIM: 256 + UCG_RATE: 0.0 + INPUT_KEYS: ["target_size_as_tuple"] + LEGACY_UCG_VALUE: + - + NAME: IPAdapterPlusEmbedder + CLIP_DIR: ms://damo/LARGEN@models/clip_encoder/ + PRETRAINED_MODEL: ms://damo/LARGEN@models/ip-adapter-plus_sdxl_vit-h.bin + INPUT_KEYS: [ "ref_ip", "ref_detail" ] + IN_DIM: 1280 + HEADS: 20 + CROSSATTN_DIM: 2048 + - + NAME: TransparentEmbedder + INPUT_KEYS: [ "tar_x0", "tar_mask_latent" ] + - + NAME: NoiseConcatEmbedder + INPUT_KEYS: [ "tar_mask_latent", "masked_x0" ] + - + NAME: TransparentEmbedder + INPUT_KEYS: [ "ref_x0" ] + - + NAME: TransparentEmbedder + INPUT_KEYS: [ "task" ] + - + NAME: TransparentEmbedder + INPUT_KEYS: [ "image_scale" ] diff --git a/scepter/methods/studio/scepter_ui.yaml b/scepter/methods/studio/scepter_ui.yaml index b3587e1..03eeb9c 100644 --- a/scepter/methods/studio/scepter_ui.yaml +++ b/scepter/methods/studio/scepter_ui.yaml @@ -49,11 +49,11 @@ BANNER: |
ms_scepter_studio_qr -
Modelscope Studio
+
scepter_github_qr -
Github
+
@@ -79,6 +79,10 @@ INTERFACE: NAME_EN: Train IFID: self_train CONFIG: scepter/methods/studio/self_train/self_train.yaml + - NAME: 模型管理 + NAME_EN: Tuner Management + IFID: tuner_manager + CONFIG: scepter/methods/studio/tuner_manager/tuner_manager.yaml - NAME: 推理 NAME_EN: Inference IFID: inference diff --git a/scepter/methods/studio/self_train/sd_xl/sdxl_pro.yaml b/scepter/methods/studio/self_train/sd_xl/sdxl_pro.yaml index dfee523..fddf18f 100644 --- a/scepter/methods/studio/self_train/sd_xl/sdxl_pro.yaml +++ b/scepter/methods/studio/self_train/sd_xl/sdxl_pro.yaml @@ -11,13 +11,13 @@ META: DEFAULT_SAMPLER: "dpmpp_2s_ancestral" DEFAULT_SAMPLE_STEPS: 40 INFERENCE_N_PROMPT: "" - RESOLUTION: 1024 + RESOLUTION: [1024, 1024] PARAS: - TRAIN_BATCH_SIZE: 2 TRAIN_PREFIX: "" TRAIN_N_PROMPT: "" - RESOLUTION: 1024 + RESOLUTION: [1024, 1024] MEMORY: 29000 EPOCHS: 50 SAVE_INTERVAL: 25 @@ -29,7 +29,7 @@ META: TRAIN_BATCH_SIZE: 2 TRAIN_PREFIX: "" TRAIN_N_PROMPT: "" - RESOLUTION: 1024 + RESOLUTION: [1024, 1024] MEMORY: 29000 EPOCHS: 50 SAVE_INTERVAL: 25 @@ -41,7 +41,7 @@ META: TRAIN_BATCH_SIZE: 2 TRAIN_PREFIX: "" TRAIN_N_PROMPT: "" - RESOLUTION: 1024 + RESOLUTION: [1024, 1024] MEMORY: 29000 EPOCHS: 200 SAVE_INTERVAL: 25 @@ -53,7 +53,7 @@ META: TRAIN_BATCH_SIZE: 2 TRAIN_PREFIX: "" TRAIN_N_PROMPT: "" - RESOLUTION: 1024 + RESOLUTION: [1024, 1024] MEMORY: 29000 EPOCHS: 200 SAVE_INTERVAL: 25 @@ -65,7 +65,7 @@ META: TRAIN_BATCH_SIZE: 2 TRAIN_PREFIX: "" TRAIN_N_PROMPT: "" - RESOLUTION: 1024 + RESOLUTION: [1024, 1024] MEMORY: 29000 EPOCHS: 50 SAVE_INTERVAL: 25 @@ -146,10 +146,12 @@ SOLVER: MAX_EPOCHS: -1 # NUM_FOLDS DESCRIPTION: Num folds for training. TYPE: int default: 1 NUM_FOLDS: 1 + # + EVAL_INTERVAL: -1 # WORK_DIR DESCRIPTION: Save dir of the training log or model. TYPE: str default: '' WORK_DIR: # LOG_FILE DESCRIPTION: Save log path. TYPE: str default: '' - LOG_FILE: stg_log.txt + LOG_FILE: std_log.txt FILE_SYSTEM: NAME: "ModelscopeFs" TEMP_DIR: "./cache/data" @@ -586,15 +588,44 @@ SOLVER: INPUT_KEY: [ 'img', 'img_original_size_as_tuple', 'img_target_size_as_tuple', 'img_crop_coords_top_left'] OUTPUT_KEY: [ 'image', 'original_size_as_tuple', 'target_size_as_tuple', 'crop_coords_top_left' ] # + EVAL_DATA: + NAME: Text2ImageDataset + MODE: eval + PROMPT_FILE: + PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ] + IMAGE_SIZE: [ 1024, 1024 ] + FIELDS: [ "prompt" ] + DELIMITER: '#;#' + PROMPT_PREFIX: '' + PIN_MEMORY: True + BATCH_SIZE: 1 + NUM_WORKERS: 4 + TRANSFORMS: + - NAME: Select + KEYS: [ 'index', 'prompt' ] + META_KEYS: [ 'image_size' ] + # TRAIN_HOOKS: - - NAME: BackwardHook - # GRADIENT_CLIP: 1.0 + - + NAME: BackwardHook PRIORITY: 0 - - NAME: LogHook - LOG_INTERVAL: 50 + - + NAME: LogHook + LOG_INTERVAL: 10 SHOW_GPU_MEM: True + - + NAME: TensorboardLogHook - NAME: CheckpointHook - SAVE_LAST: True INTERVAL: 10000 PRIORITY: 200 + SAVE_LAST: True + SAVE_NAME_PREFIX: 'step' + # + EVAL_HOOKS: + - + NAME: ProbeDataHook + PROB_INTERVAL: 100 + SAVE_LAST: True + SAVE_NAME_PREFIX: 'step' + SAVE_PROBE_PREFIX: 'image' diff --git a/scepter/methods/studio/self_train/self_train.yaml b/scepter/methods/studio/self_train/self_train.yaml index 0c5925f..d624818 100644 --- a/scepter/methods/studio/self_train/self_train.yaml +++ b/scepter/methods/studio/self_train/self_train.yaml @@ -8,3 +8,13 @@ SAMPLERS: NAME: 'dpmpp_2m_sde' - NAME: 'dpmpp_2s_ancestral' +TRAIN_PARAS: + RESOLUTIONS: + VALUES: [[256, 256], [320, 180], [180, 320], + [512, 512], [640, 360], [360, 640], + [768, 768], [960, 540], [540, 960], + [1024, 1024], [1280, 720], [720, 1280]] + DEFAULT: [1024, 1024] + EVAL_PROMPTS: + - a boy wearing a jacket + - a dog running on the lawn diff --git a/scepter/methods/studio/self_train/stable_diffusion/sd15_pro.yaml b/scepter/methods/studio/self_train/stable_diffusion/sd15_pro.yaml index ca7f4d1..f818291 100644 --- a/scepter/methods/studio/self_train/stable_diffusion/sd15_pro.yaml +++ b/scepter/methods/studio/self_train/stable_diffusion/sd15_pro.yaml @@ -10,13 +10,13 @@ META: DEFAULT_SAMPLER: "ddim" DEFAULT_SAMPLE_STEPS: 40 INFERENCE_N_PROMPT: "" - RESOLUTION: 512 + RESOLUTION: [512, 512] PARAS: - TRAIN_BATCH_SIZE: 2 TRAIN_PREFIX: "" TRAIN_N_PROMPT: "" - RESOLUTION: 512 + RESOLUTION: [512, 512] MEMORY: 29000 EPOCHS: 50 SAVE_INTERVAL: 25 @@ -28,7 +28,7 @@ META: TRAIN_BATCH_SIZE: 2 TRAIN_PREFIX: "" TRAIN_N_PROMPT: "" - RESOLUTION: 512 + RESOLUTION: [512, 512] MEMORY: 29000 EPOCHS: 50 SAVE_INTERVAL: 25 @@ -40,7 +40,7 @@ META: TRAIN_BATCH_SIZE: 2 TRAIN_PREFIX: "" TRAIN_N_PROMPT: "" - RESOLUTION: 512 + RESOLUTION: [512, 512] MEMORY: 29000 EPOCHS: 200 SAVE_INTERVAL: 25 @@ -53,7 +53,7 @@ META: TRAIN_BATCH_SIZE: 2 TRAIN_PREFIX: "" TRAIN_N_PROMPT: "" - RESOLUTION: 512 + RESOLUTION: [512, 512] MEMORY: 29000 EPOCHS: 200 SAVE_INTERVAL: 25 @@ -66,7 +66,7 @@ META: TRAIN_BATCH_SIZE: 4 TRAIN_PREFIX: "" TRAIN_N_PROMPT: "" - RESOLUTION: 512 + RESOLUTION: [512, 512] MEMORY: 29000 EPOCHS: 50 SAVE_INTERVAL: 25 @@ -135,6 +135,7 @@ SOLVER: MAX_EPOCHS: -1 NUM_FOLDS: 1 ACCU_STEP: 1 + EVAL_INTERVAL: -1 # WORK_DIR: LOG_FILE: std_log.txt @@ -298,15 +299,44 @@ SOLVER: KEYS: [ 'image', 'prompt' ] META_KEYS: [ 'data_key' ] # + EVAL_DATA: + NAME: Text2ImageDataset + MODE: eval + PROMPT_FILE: + PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ] + IMAGE_SIZE: [ 512, 512 ] + FIELDS: [ "prompt" ] + DELIMITER: '#;#' + PROMPT_PREFIX: '' + PIN_MEMORY: True + BATCH_SIZE: 1 + NUM_WORKERS: 4 + TRANSFORMS: + - NAME: Select + KEYS: [ 'index', 'prompt' ] + META_KEYS: [ 'image_size' ] + # TRAIN_HOOKS: - - NAME: BackwardHook - # GRADIENT_CLIP: 1.0 + - + NAME: BackwardHook PRIORITY: 0 - - NAME: LogHook - LOG_INTERVAL: 50 + - + NAME: LogHook + LOG_INTERVAL: 10 SHOW_GPU_MEM: True + - + NAME: TensorboardLogHook - NAME: CheckpointHook - SAVE_LAST: True INTERVAL: 10000 PRIORITY: 200 + SAVE_LAST: True + SAVE_NAME_PREFIX: 'step' + # + EVAL_HOOKS: + - + NAME: ProbeDataHook + PROB_INTERVAL: 100 + SAVE_LAST: True + SAVE_NAME_PREFIX: 'step' + SAVE_PROBE_PREFIX: 'image' diff --git a/scepter/methods/studio/self_train/stable_diffusion/sd21_pro.yaml b/scepter/methods/studio/self_train/stable_diffusion/sd21_pro.yaml index e05788f..257af14 100644 --- a/scepter/methods/studio/self_train/stable_diffusion/sd21_pro.yaml +++ b/scepter/methods/studio/self_train/stable_diffusion/sd21_pro.yaml @@ -10,13 +10,13 @@ META: DEFAULT_SAMPLER: "ddim" DEFAULT_SAMPLE_STEPS: 40 INFERENCE_N_PROMPT: "" - RESOLUTION: 768 + RESOLUTION: [768, 768] PARAS: - TRAIN_BATCH_SIZE: 2 TRAIN_PREFIX: "" TRAIN_N_PROMPT: "" - RESOLUTION: 768 + RESOLUTION: [768, 768] MEMORY: 29000 EPOCHS: 50 SAVE_INTERVAL: 25 @@ -28,7 +28,7 @@ META: TRAIN_BATCH_SIZE: 2 TRAIN_PREFIX: "" TRAIN_N_PROMPT: "" - RESOLUTION: 768 + RESOLUTION: [768, 768] MEMORY: 29000 EPOCHS: 50 SAVE_INTERVAL: 25 @@ -40,7 +40,7 @@ META: TRAIN_BATCH_SIZE: 2 TRAIN_PREFIX: "" TRAIN_N_PROMPT: "" - RESOLUTION: 768 + RESOLUTION: [768, 768] MEMORY: 29000 EPOCHS: 200 SAVE_INTERVAL: 25 @@ -78,6 +78,7 @@ SOLVER: MAX_EPOCHS: -1 NUM_FOLDS: 1 ACCU_STEP: 1 + EVAL_INTERVAL: -1 # WORK_DIR: LOG_FILE: std_log.txt @@ -240,15 +241,44 @@ SOLVER: KEYS: [ 'image', 'prompt' ] META_KEYS: [ 'data_key' ] # + EVAL_DATA: + NAME: Text2ImageDataset + MODE: eval + PROMPT_FILE: + PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ] + IMAGE_SIZE: [ 768, 768 ] + FIELDS: [ "prompt" ] + DELIMITER: '#;#' + PROMPT_PREFIX: '' + PIN_MEMORY: True + BATCH_SIZE: 1 + NUM_WORKERS: 4 + TRANSFORMS: + - NAME: Select + KEYS: [ 'index', 'prompt' ] + META_KEYS: [ 'image_size' ] + # TRAIN_HOOKS: - - NAME: BackwardHook - # GRADIENT_CLIP: 1.0 + - + NAME: BackwardHook PRIORITY: 0 - - NAME: LogHook - LOG_INTERVAL: 50 + - + NAME: LogHook + LOG_INTERVAL: 10 SHOW_GPU_MEM: True + - + NAME: TensorboardLogHook - NAME: CheckpointHook - SAVE_LAST: True INTERVAL: 10000 PRIORITY: 200 + SAVE_LAST: True + SAVE_NAME_PREFIX: 'step' + # + EVAL_HOOKS: + - + NAME: ProbeDataHook + PROB_INTERVAL: 100 + SAVE_LAST: True + SAVE_NAME_PREFIX: 'step' + SAVE_PROBE_PREFIX: 'image' diff --git a/scepter/methods/studio/tuner_manager/tuner_manager.yaml b/scepter/methods/studio/tuner_manager/tuner_manager.yaml new file mode 100644 index 0000000..c0daf65 --- /dev/null +++ b/scepter/methods/studio/tuner_manager/tuner_manager.yaml @@ -0,0 +1,2 @@ +WORK_DIR: "tuner_manager" +TUNER_LIST_YAML: "tuner_list.yaml" diff --git a/scepter/modules/data/dataset/dataset.py b/scepter/modules/data/dataset/dataset.py index 5fb7942..22e9537 100644 --- a/scepter/modules/data/dataset/dataset.py +++ b/scepter/modules/data/dataset/dataset.py @@ -244,12 +244,18 @@ class Text2ImageDataset(BaseDataset): image_size = [image_size, image_size] assert isinstance(image_size, Iterable) and len(image_size) == 2 - prompt_file = cfg.PROMPT_FILE - with FS.get_object(prompt_file) as local_data: + if cfg.PROMPT_FILE is not None and cfg.PROMPT_FILE != '': + prompt_file = cfg.PROMPT_FILE + with FS.get_object(prompt_file) as local_data: + rows = [ + i.split(delimiter, + len(fields) - 1) + for i in local_data.decode('utf-8').strip().split('\n') + ] + else: rows = [ i.split(delimiter, - len(fields) - 1) - for i in local_data.decode('utf-8').strip().split('\n') + len(fields) - 1) for i in cfg.PROMPT_DATA ] self.items = list() @@ -263,10 +269,9 @@ class Text2ImageDataset(BaseDataset): item['meta']['img_path'] = os.path.join(path_prefix, value) elif key in ['width', 'height']: item['meta'][key] = int(value) - elif key != 'meta': - item[key] = value else: - continue + item['meta'][key] = value + self.items.append(item) if use_num > 0: self.items = self.items[:use_num] diff --git a/scepter/modules/inference/control_inference.py b/scepter/modules/inference/control_inference.py index 47836ea..406e6e6 100644 --- a/scepter/modules/inference/control_inference.py +++ b/scepter/modules/inference/control_inference.py @@ -2,18 +2,23 @@ # Copyright (c) Alibaba, Inc. and its affiliates. import copy import os +import warnings import torch import torch.nn as nn import torchvision.transforms as TT from PIL.Image import Image -from swift import SwiftModel from scepter.modules.model.registry import TUNERS from scepter.modules.utils.config import Config from scepter.modules.utils.distribute import we from scepter.modules.utils.file_system import FS +try: + from swift import SwiftModel +except Exception: + warnings.warn('Import swift failed, please check it.') + class ControlInference(): def __init__(self, logger=None): diff --git a/scepter/modules/inference/diffusion_inference.py b/scepter/modules/inference/diffusion_inference.py index 901ff95..097c045 100644 --- a/scepter/modules/inference/diffusion_inference.py +++ b/scepter/modules/inference/diffusion_inference.py @@ -474,10 +474,14 @@ class DiffusionInference(): enabled=dtype == 'float16', dtype=getattr(torch, dtype)): if self.tokenizer: + if not hasattr(get_model(self.cond_stage_model), 'tokenizer'): + setattr(get_model(self.cond_stage_model), 'tokenizer', + self.tokenizer) context = getattr(get_model(self.cond_stage_model), function_name)(batch['tokens']) null_context = getattr(get_model(self.cond_stage_model), function_name)(batch_uc['tokens']) + else: context = getattr(get_model(self.cond_stage_model), function_name)(batch) @@ -558,12 +562,13 @@ class DiffusionInference(): seed=seed, condition_fn=None, clamp=None, + sharpness=value_input.get('sharpness', 0.0), percentile=None, t_max=None, t_min=None, discard_penultimate_step=None, intermediate_callback=intermediate_callback, - cat_uc=cat_uc, + cat_uc=value_input.get('cat_uc', cat_uc), **kwargs) self.dynamic_unload(self.diffusion_model, diff --git a/scepter/modules/inference/largen_inference.py b/scepter/modules/inference/largen_inference.py new file mode 100644 index 0000000..afbd454 --- /dev/null +++ b/scepter/modules/inference/largen_inference.py @@ -0,0 +1,595 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import copy +import os.path +import random +from collections import OrderedDict + +import torch +import torchvision.transforms.functional as TF +from PIL.Image import Image + +from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion +from scepter.modules.model.network.diffusion.schedules import noise_schedule +from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, MODELS, + TOKENIZERS) +from scepter.modules.model.utils.data_utils import crop_back +from scepter.modules.utils.distribute import we +from scepter.modules.utils.file_system import FS + + +def get_model(model_tuple): + assert 'model' in model_tuple + return model_tuple['model'] + + +class LargenInference(): + ''' + define vae, unet, text-encoder, tuner, refiner components + support to load the components dynamicly. + create and load model when run this model at the first time. + ''' + def __init__(self, logger=None): + self.logger = logger + self.loaded_model = {} + self.loaded_model_name = [ + 'diffusion_model', 'first_stage_model', 'cond_stage_model' + ] + + def init_from_cfg(self, cfg): + self.name = cfg.NAME + self.is_default = cfg.get('IS_DEFAULT', False) + module_paras = self.load_default(cfg.get('DEFAULT_PARAS', None)) + assert cfg.have('MODEL') + cfg.MODEL = self.redefine_paras(cfg.MODEL) + self.diffusion = self.load_schedule(cfg.MODEL.SCHEDULE) + self.diffusion_model = self.infer_model( + cfg.MODEL.DIFFUSION_MODEL, module_paras.get( + 'DIFFUSION_MODEL', + None)) if cfg.MODEL.have('DIFFUSION_MODEL') else None + self.first_stage_model = self.infer_model( + cfg.MODEL.FIRST_STAGE_MODEL, + module_paras.get( + 'FIRST_STAGE_MODEL', + None)) if cfg.MODEL.have('FIRST_STAGE_MODEL') else None + self.cond_stage_model = self.infer_model( + cfg.MODEL.COND_STAGE_MODEL, + module_paras.get( + 'COND_STAGE_MODEL', + None)) if cfg.MODEL.have('COND_STAGE_MODEL') else None + self.refiner_cond_model = self.infer_model( + cfg.MODEL.REFINER_COND_MODEL, + module_paras.get( + 'REFINER_COND_MODEL', + None)) if cfg.MODEL.have('REFINER_COND_MODEL') else None + self.refiner_diffusion_model = self.infer_model( + cfg.MODEL.REFINER_MODEL, module_paras.get( + 'REFINER_MODEL', + None)) if cfg.MODEL.have('REFINER_MODEL') else None + self.tokenizer = TOKENIZERS.build( + cfg.MODEL.TOKENIZER, + logger=self.logger) if cfg.MODEL.have('TOKENIZER') else None + + if self.tokenizer is not None: + self.cond_stage_model['cfg'].KWARGS = { + 'vocab_size': self.tokenizer.vocab_size + } + + def redefine_paras(self, cfg): + if cfg.get('PRETRAINED_MODEL', None): + assert FS.isfile(cfg.PRETRAINED_MODEL) + with FS.get_from(cfg.PRETRAINED_MODEL, + wait_finish=True) as local_path: + if local_path.endswith('safetensors'): + from safetensors.torch import load_file as load_safetensors + sd = load_safetensors(local_path) + else: + sd = torch.load(local_path, map_location='cpu') + + if 'model' in sd: + sd = sd['model'] + + first_stage_model_path = os.path.join( + os.path.dirname(local_path), 'first_stage_model.pth') + cond_stage_model_path = os.path.join( + os.path.dirname(local_path), 'cond_stage_model.pth') + diffusion_model_path = os.path.join( + os.path.dirname(local_path), 'diffusion_model.pth') + if (not os.path.exists(first_stage_model_path) + or not os.path.exists(cond_stage_model_path) + or not os.path.exists(diffusion_model_path)): + self.logger.info( + 'Now read the whole model and rearrange the modules, it may take several mins.' + ) + first_stage_model = OrderedDict() + cond_stage_model = OrderedDict() + diffusion_model = OrderedDict() + for k, v in sd.items(): + if k.startswith('first_stage_model.'): + first_stage_model[k.replace( + 'first_stage_model.', '')] = v + elif k.startswith('conditioner.'): + cond_stage_model[k.replace('conditioner.', '')] = v + elif k.startswith('cond_stage_model.'): + if k.startswith('cond_stage_model.model.'): + cond_stage_model[k.replace( + 'cond_stage_model.model.', '')] = v + else: + cond_stage_model[k.replace( + 'cond_stage_model.', '')] = v + elif k.startswith('model.diffusion_model.'): + diffusion_model[k.replace('model.diffusion_model.', + '')] = v + elif k.startswith('model.'): + diffusion_model[k.replace('model.', '')] = v + else: + continue + if cfg.have('FIRST_STAGE_MODEL'): + with open(first_stage_model_path + 'cache', 'wb') as f: + torch.save(first_stage_model, f) + os.rename(first_stage_model_path + 'cache', + first_stage_model_path) + self.logger.info( + 'First stage model has been processed.') + if cfg.have('COND_STAGE_MODEL'): + with open(cond_stage_model_path + 'cache', 'wb') as f: + torch.save(cond_stage_model, f) + os.rename(cond_stage_model_path + 'cache', + cond_stage_model_path) + self.logger.info( + 'Cond stage model has been processed.') + if cfg.have('DIFFUSION_MODEL'): + with open(diffusion_model_path + 'cache', 'wb') as f: + torch.save(diffusion_model, f) + os.rename(diffusion_model_path + 'cache', + diffusion_model_path) + self.logger.info('Diffusion model has been processed.') + if not cfg.FIRST_STAGE_MODEL.get('PRETRAINED_MODEL', None): + cfg.FIRST_STAGE_MODEL.PRETRAINED_MODEL = first_stage_model_path + else: + cfg.FIRST_STAGE_MODEL.RELOAD_MODEL = first_stage_model_path + if not cfg.COND_STAGE_MODEL.get('PRETRAINED_MODEL', None): + cfg.COND_STAGE_MODEL.PRETRAINED_MODEL = cond_stage_model_path + else: + cfg.COND_STAGE_MODEL.RELOAD_MODEL = cond_stage_model_path + if not cfg.DIFFUSION_MODEL.get('PRETRAINED_MODEL', None): + cfg.DIFFUSION_MODEL.PRETRAINED_MODEL = diffusion_model_path + else: + cfg.DIFFUSION_MODEL.RELOAD_MODEL = diffusion_model_path + return cfg + + def init_from_modules(self, modules): + for k, v in modules.items(): + self.__setattr__(k, v) + + def infer_model(self, cfg, module_paras=None): + module = { + 'model': None, + 'cfg': cfg, + 'device': 'offline', + 'name': cfg.NAME, + 'function_info': {}, + 'paras': {} + } + if module_paras is None: + return module + function_info = {} + paras = { + k.lower(): v + for k, v in module_paras.get('PARAS', {}).items() + } + for function in module_paras.get('FUNCTION', []): + input_dict = {} + for inp in function.get('INPUT', []): + if inp.lower() in self.input: + input_dict[inp.lower()] = self.input[inp.lower()] + function_info[function.NAME] = { + 'dtype': function.get('DTYPE', 'float32'), + 'input': input_dict + } + module['paras'] = paras + module['function_info'] = function_info + return module + + def init_from_ckpt(self, path, model, ignore_keys=list()): + if path.endswith('safetensors'): + from safetensors.torch import load_file as load_safetensors + sd = load_safetensors(path) + else: + sd = torch.load(path, map_location='cpu') + + new_sd = OrderedDict() + for k, v in sd.items(): + ignored = False + for ik in ignore_keys: + if ik in k: + if we.rank == 0: + self.logger.info( + 'Ignore key {} from state_dict.'.format(k)) + ignored = True + break + if not ignored: + new_sd[k] = v + + missing, unexpected = model.load_state_dict(new_sd, strict=False) + if we.rank == 0: + self.logger.info( + f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys' + ) + if len(missing) > 0: + self.logger.info(f'Missing Keys:\n {missing}') + if len(unexpected) > 0: + self.logger.info(f'\nUnexpected Keys:\n {unexpected}') + + def load(self, module): + if module['device'] == 'offline': + if module['cfg'].NAME in MODELS.class_map: + model = MODELS.build(module['cfg'], logger=self.logger).eval() + elif module['cfg'].NAME in BACKBONES.class_map: + model = BACKBONES.build(module['cfg'], + logger=self.logger).eval() + elif module['cfg'].NAME in EMBEDDERS.class_map: + model = EMBEDDERS.build(module['cfg'], + logger=self.logger).eval() + else: + raise NotImplementedError + if module['cfg'].get('RELOAD_MODEL', None): + self.init_from_ckpt(module['cfg'].RELOAD_MODEL, model) + module['model'] = model + module['device'] = 'cpu' + if module['device'] == 'cpu': + module['device'] = we.device_id + module['model'] = module['model'].to(we.device_id) + return module + + def unload(self, module): + if module is None: + return module + module['model'] = module['model'].to('cpu') + module['device'] = 'cpu' + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + return module + + def dynamic_load(self, module=None, name=''): + self.logger.info('Loading {} model'.format(name)) + if name == 'all': + for subname in self.loaded_model_name: + self.loaded_model[subname] = self.dynamic_load( + getattr(self, subname), subname) + elif name in self.loaded_model_name: + if name in self.loaded_model: + if module['cfg'] != self.loaded_model[name]['cfg']: + self.unload(self.loaded_model[name]) + module = self.load(module) + self.loaded_model[name] = module + return module + elif module['device'] == 'cpu': + module = self.load(module) + return module + else: + return module + else: + module = self.load(module) + self.loaded_model[name] = module + return module + else: + return self.load(module) + + def dynamic_unload(self, module=None, name='', skip_loaded=False): + self.logger.info('Unloading {} model'.format(name)) + if name == 'all': + for name, module in self.loaded_model.items(): + module = self.unload(self.loaded_model[name]) + self.loaded_model[name] = module + elif name in self.loaded_model_name: + if name in self.loaded_model: + if not skip_loaded: + module = self.unload(self.loaded_model[name]) + self.loaded_model[name] = module + else: + self.unload(module) + else: + self.unload(module) + + def load_default(self, cfg): + module_paras = {} + if cfg is not None: + self.paras = cfg.PARAS + self.input = {k.lower(): v for k, v in cfg.INPUT.items()} + self.output = {k.lower(): v for k, v in cfg.OUTPUT.items()} + module_paras = cfg.MODULES_PARAS + return module_paras + + def load_schedule(self, cfg): + parameterization = cfg.get('PARAMETERIZATION', 'eps') + assert parameterization in [ + 'eps', 'x0', 'v' + ], 'currently only supporting "eps" and "x0" and "v"' + num_timesteps = cfg.get('TIMESTEPS', 1000) + + schedule_args = { + k.lower(): v + for k, v in cfg.get('SCHEDULE_ARGS', { + 'NAME': 'logsnr_cosine_interp', + 'SCALE_MIN': 2.0, + 'SCALE_MAX': 4.0 + }).items() + } + + zero_terminal_snr = cfg.get('ZERO_TERMINAL_SNR', False) + if zero_terminal_snr: + assert parameterization == 'v', 'Now zero_terminal_snr only support v-prediction mode.' + sigmas = noise_schedule(schedule=schedule_args.pop('name'), + n=num_timesteps, + zero_terminal_snr=zero_terminal_snr, + **schedule_args) + diffusion = GaussianDiffusion(sigmas=sigmas, + prediction_type=parameterization) + return diffusion + + def get_batch(self, value_dict, num_samples=1): + batch = {} + batch_uc = {} + N = num_samples + device = we.device_id + for key in value_dict: + if key == 'prompt': + if not self.tokenizer: + batch['prompt'] = value_dict['prompt'] + batch_uc['prompt'] = value_dict['negative_prompt'] + else: + batch['tokens'] = self.tokenizer(value_dict['prompt']).to( + we.device_id) + batch_uc['tokens'] = self.tokenizer( + value_dict['negative_prompt']).to(we.device_id) + elif key == 'original_size_as_tuple': + batch['original_size_as_tuple'] = (torch.tensor( + value_dict['original_size_as_tuple']).to(device).repeat( + N, 1)) + elif key == 'crop_coords_top_left': + batch['crop_coords_top_left'] = (torch.tensor( + value_dict['crop_coords_top_left']).to(device).repeat( + N, 1)) + elif key == 'aesthetic_score': + batch['aesthetic_score'] = (torch.tensor( + [value_dict['aesthetic_score']]).to(device).repeat(N, 1)) + batch_uc['aesthetic_score'] = (torch.tensor([ + value_dict['negative_aesthetic_score'] + ]).to(device).repeat(N, 1)) + + elif key == 'target_size_as_tuple': + batch['target_size_as_tuple'] = (torch.tensor( + value_dict['target_size_as_tuple']).to(device).repeat( + N, 1)) + elif key == 'image': + batch[key] = self.load_image(value_dict[key], num_samples=N) + else: + batch[key] = value_dict[key] + + for key in batch.keys(): + if key not in batch_uc and isinstance(batch[key], torch.Tensor): + batch_uc[key] = torch.clone(batch[key]) + return batch, batch_uc + + def load_image(self, image, num_samples=1): + if isinstance(image, torch.Tensor): + pass + elif isinstance(image, Image): + pass + elif isinstance(image, Image): + pass + + def get_function_info(self, module, function_name=None): + all_function = module['function_info'] + if function_name in all_function: + return function_name, all_function[function_name]['dtype'] + if function_name is None and len(all_function) == 1: + for k, v in all_function.items(): + return k, v['dtype'] + + def encode_first_stage(self, x, **kwargs): + _, dtype = self.get_function_info(self.first_stage_model, 'encode') + with torch.autocast('cuda', + enabled=dtype == 'float16', + dtype=getattr(torch, dtype)): + z = get_model(self.first_stage_model).encode(x) + return self.first_stage_model['paras']['scale_factor'] * z + + def decode_first_stage(self, z): + _, dtype = self.get_function_info(self.first_stage_model, 'encode') + with torch.autocast('cuda', + enabled=dtype == 'float16', + dtype=getattr(torch, dtype)): + z = 1. / self.first_stage_model['paras']['scale_factor'] * z + return get_model(self.first_stage_model).decode(z) + + @torch.no_grad() + def __call__(self, + input, + num_samples=1, + intermediate_callback=None, + refine_strength=0, + img_to_img_strength=0, + cat_uc=True, + tuner_model=None, + control_model=None, + **kwargs): + + value_input = copy.deepcopy(self.input) + value_input.update(input) + print(value_input) + height, width = value_input['target_size_as_tuple'] + value_output = copy.deepcopy(self.output) + batch, batch_uc = self.get_batch(value_input, num_samples=1) + + # first stage encode + task = kwargs.get('largen_task', 'Text_Guided_Inpainting') + image_scale = kwargs.get('largen_image_scale', 1.0) + tar_image = kwargs.get('largen_tar_image', None) + tar_mask = kwargs.get('largen_tar_mask', None) + masked_image = kwargs.get('largen_masked_image', None) + ref_image = kwargs.get('largen_ref_image', None) + ref_mask = kwargs.get('largen_ref_mask', None) + ref_clip = kwargs.get('largen_ref_clip', None) + + base_image = kwargs.get('largen_base_image', None) + extra_sizes = kwargs.get('largen_extra_sizes', None) + bbox_yyxx = kwargs.get('largen_bbox_yyxx', None) + + device = we.device_id + tar_image = tar_image.to(device) + tar_mask = tar_mask.to(device) + masked_image = masked_image.to(device) + if 'Subject' in task: + ref_image = ref_image.to(device) + ref_mask = ref_mask.to(device) + ref_clip = ref_clip.to(device) + + self.dynamic_load(self.first_stage_model, 'first_stage_model') + + tar_x0 = self.encode_first_stage(tar_image) + masked_x0 = self.encode_first_stage(masked_image) + b, _, h, w = tar_x0.shape + tar_mask_latent = TF.resize(tar_mask, (h, w), antialias=True) + tar_mask_latent = (tar_mask_latent > 0.5).float() + + batch.update({ + 'tar_x0': tar_x0, + 'tar_mask_latent': tar_mask_latent, + 'masked_x0': masked_x0, + 'task': task + }) + batch_uc.update({ + 'tar_x0': tar_x0, + 'tar_mask_latent': tar_mask_latent, + 'masked_x0': masked_x0, + 'task': task + }) + + if 'Subject' in task and ref_image is not None: + ref_x0 = self.encode_first_stage(ref_image) + batch.update({ + 'ref_ip': ref_clip, + 'ref_detail': ref_clip, + 'ref_x0': ref_x0, + 'ref_mask': ref_mask, + 'image_scale': image_scale, + }) + batch_uc.update({ + 'ref_ip': torch.zeros_like(ref_clip), + 'ref_detail': ref_clip, + 'ref_x0': ref_x0, + 'ref_mask': ref_mask, + 'image_scale': image_scale, + }) + + self.dynamic_unload(self.first_stage_model, + 'first_stage_model', + skip_loaded=True) + + # cond stage + self.dynamic_load(self.cond_stage_model, 'cond_stage_model') + function_name, dtype = self.get_function_info(self.cond_stage_model) + with torch.autocast('cuda', + enabled=dtype == 'float16', + dtype=getattr(torch, dtype)): + context = getattr(get_model(self.cond_stage_model), + function_name)(batch) + null_context = getattr(get_model(self.cond_stage_model), + function_name)(batch_uc) + + self.dynamic_unload(self.cond_stage_model, + 'cond_stage_model', + skip_loaded=True) + + # get noise + seed = kwargs.pop('seed', -1) + g = torch.Generator(device=we.device_id) + seed = seed if seed >= 0 else random.randint(0, 2**32 - 1) + g.manual_seed(seed) + if 'seed' in value_output: + value_output['seed'] = seed + for sample_id in range(num_samples): + if self.diffusion_model is not None: + noise = torch.empty( + 1, + 4, + height // self.first_stage_model['paras']['size_factor'], + width // self.first_stage_model['paras']['size_factor'], + device=we.device_id).normal_(generator=g) + + self.dynamic_load(self.diffusion_model, 'diffusion_model') + # UNet use input n_prompt + function_name, dtype = self.get_function_info( + self.diffusion_model) + with torch.autocast('cuda', + enabled=dtype == 'float16', + dtype=getattr(torch, dtype)): + latent = self.diffusion.sample( + noise=noise, + x=None, + denoising_strength=1.0, + refine_strength=refine_strength, + solver=value_input.get('sample', 'ddim'), + model=get_model(self.diffusion_model), + model_kwargs=[{ + 'cond': context + }, { + 'cond': null_context + }], + steps=value_input.get('sample_steps', 50), + guide_scale=value_input.get('guide_scale', 7.5), + guide_rescale=value_input.get('guide_rescale', 0.5), + discretization=value_input.get('discretization', + 'trailing'), + show_progress=True, + seed=seed, + condition_fn=None, + clamp=None, + percentile=None, + t_max=None, + t_min=None, + discard_penultimate_step=None, + intermediate_callback=intermediate_callback, + cat_uc=cat_uc, + **kwargs) + + self.dynamic_unload(self.diffusion_model, + 'diffusion_model', + skip_loaded=True) + + if 'latent' in value_output: + if value_output['latent'] is None or ( + isinstance(value_output['latent'], list) + and len(value_output['latent']) < 1): + value_output['latent'] = [] + value_output['latent'].append(latent) + + self.dynamic_load(self.first_stage_model, 'first_stage_model') + x_samples = self.decode_first_stage(latent).float() + self.dynamic_unload(self.first_stage_model, + 'first_stage_model', + skip_loaded=True) + images = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0) + if base_image is not None: + stitch_images = [] + for img in images: + stitch_img = crop_back(img, copy.deepcopy(base_image), + extra_sizes, bbox_yyxx) + stitch_images.append(stitch_img) + images = torch.stack(stitch_images, dim=0) + if 'images' in value_output: + if value_output['images'] is None or ( + isinstance(value_output['images'], list) + and len(value_output['images']) < 1): + value_output['images'] = [] + value_output['images'].append(images) + + for k, v in value_output.items(): + if isinstance(v, list): + value_output[k] = torch.cat(v, dim=0) + if isinstance(v, torch.Tensor): + value_output[k] = v.cpu() + + return value_output diff --git a/scepter/modules/model/backbone/autoencoder/__init__.py b/scepter/modules/model/backbone/autoencoder/__init__.py index cfa8365..bf84469 100644 --- a/scepter/modules/model/backbone/autoencoder/__init__.py +++ b/scepter/modules/model/backbone/autoencoder/__init__.py @@ -1,4 +1,5 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. from scepter.modules.model.backbone.autoencoder.ae_module import (Decoder, - Encoder) + Encoder, + RDecoder) diff --git a/scepter/modules/model/backbone/autoencoder/ae_module.py b/scepter/modules/model/backbone/autoencoder/ae_module.py index 7c011b9..5ee854c 100644 --- a/scepter/modules/model/backbone/autoencoder/ae_module.py +++ b/scepter/modules/model/backbone/autoencoder/ae_module.py @@ -3,6 +3,7 @@ import torch import torch.nn as nn +from einops import repeat from torch.utils.checkpoint import checkpoint from scepter.modules.model.backbone.autoencoder.ae_utils import ( @@ -244,6 +245,7 @@ class Decoder(BaseModel): # compute in_ch_mult, block_in and curr_res at lowest res block_in = self.ch * self.ch_mult[self.num_resolutions - 1] + self.block_in = block_in curr_res = 1 # z to block_in self.conv_in = torch.nn.Conv2d(self.z_channels, @@ -340,3 +342,48 @@ class Decoder(BaseModel): __class__.__name__, Decoder.para_dict, set_name=True) + + +@BACKBONES.register_class() +class RDecoder(Decoder): + def construct_model(self): + super().construct_model() + self.resize_level = nn.Sequential( + nn.Linear(self.block_in, self.block_in), + nn.SiLU(), + nn.Linear(self.block_in, self.block_in), + ) + + def forward(self, z, rembed=None): + # timestep embedding + temb = None + h = self.conv_in(z) + bs, channel, hdim, wdim = h.size() + if rembed is not None: + rembed = self.resize_level(rembed) + rembed = repeat(rembed, 'b e-> b e hd wd', hd=hdim, wd=wdim) + h = h + rembed + + # middle + if not self.use_checkpoint: + h = self.mid_upsclae_transform(h, temb) + else: + h = checkpoint(self.mid_upsclae_transform, h, temb) + + # end + if self.give_pre_end: + return h + + h = self.norm_out(h) + h = nonlinearity(h) + h = self.conv_out(h) + if self.tanh_out: + h = torch.tanh(h) + return h + + @staticmethod + def get_config_template(): + return dict_to_yaml('BACKBONE', + __class__.__name__, + Decoder.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/unet/__init__.py b/scepter/modules/model/backbone/unet/__init__.py index 99d7eae..fb54fb4 100644 --- a/scepter/modules/model/backbone/unet/__init__.py +++ b/scepter/modules/model/backbone/unet/__init__.py @@ -1,3 +1,5 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.model.backbone.unet.unet_module import DiffusionUNet +from scepter.modules.model.backbone.unet.unet_module import (DiffusionUNet, + DiffusionUNetXL, + LargenUNetXL) diff --git a/scepter/modules/model/backbone/unet/unet_module.py b/scepter/modules/model/backbone/unet/unet_module.py index 6c71fcc..99caccb 100644 --- a/scepter/modules/model/backbone/unet/unet_module.py +++ b/scepter/modules/model/backbone/unet/unet_module.py @@ -8,8 +8,9 @@ import torch import torch.nn as nn from scepter.modules.model.backbone.unet.unet_utils import ( - Downsample, ResBlock, SpatialTransformer, Timestep, - TimestepEmbedSequential, Upsample, conv_nd, linear, normalization, + BasicTransformerBlock, Downsample, ResBlock, SpatialTransformer, + SpatialTransformerV2, Timestep, TimestepEmbedSequential, + TransformerBlockV2, Upsample, conv_nd, linear, normalization, timestep_embedding, zero_module) from scepter.modules.model.base_model import BaseModel from scepter.modules.model.registry import BACKBONES @@ -951,3 +952,443 @@ class DiffusionUNetXL(DiffusionUNet): __class__.__name__, DiffusionUNetXL.para_dict, set_name=True) + + +@BACKBONES.register_class() +class LargenUNetXL(DiffusionUNetXL): + para_dict = { + 'TRANSFORMER_BLOCK_TYPE': { + 'value': 'att_v1' + }, + 'IMAGE_SCALE': { + 'value': 0.0, + }, + } + para_dict.update(DiffusionUNetXL.para_dict) + + def __init__(self, cfg, logger): + super().__init__(cfg, logger=logger) + self.init_params(cfg) + self.construct_network() + + def init_params(self, cfg): + super().init_params(cfg) + self.transformer_block_type = cfg.get('TRANSFORMER_BLOCK_TYPE', + 'att_v1') + TRANSFORMER_BLOCKS = { + 'att_v1': BasicTransformerBlock, + 'att_v2': TransformerBlockV2, + } + assert self.transformer_block_type in list(TRANSFORMER_BLOCKS.keys()) + self.transformer_block = TRANSFORMER_BLOCKS[ + self.transformer_block_type] + self.image_scale = cfg.get('IMAGE_SCALE', 0.0) + self.use_refine = cfg.get('USE_REFINE', False) + + def construct_network(self): + in_channels = self.in_channels + model_channels = self.model_channels + out_channels = self.out_channels + attention_resolutions = self.attention_resolutions + channel_mult = self.channel_mult + num_classes = self.num_classes + num_heads = self.num_heads + num_head_channels = self.num_head_channels + dims = self.dims + dropout = self.dropout + use_checkpoint = self.use_checkpoint + use_scale_shift_norm = self.use_scale_shift_norm + disable_self_attentions = self.disable_self_attentions + disable_middle_self_attn = self.disable_middle_self_attn + transformer_depth = self.transformer_depth + transformer_depth_middle = self.transformer_depth_middle + context_dim = self.context_dim + use_linear_in_transformer = self.use_linear_in_transformer + resblock_updown = self.resblock_updown + conv_resample = self.conv_resample + adm_in_channels = self.adm_in_channels + transformer_block = self.transformer_block + + time_embed_dim = model_channels * 4 + self.time_embed = nn.Sequential( + linear(model_channels, time_embed_dim), + nn.SiLU(), + linear(time_embed_dim, time_embed_dim), + ) + + if self.num_classes is not None: + if isinstance(self.num_classes, int): + self.label_emb = nn.Embedding(num_classes, time_embed_dim) + elif self.num_classes == 'continuous': + print('setting up linear c_adm embedding layer') + self.label_emb = nn.Linear(1, time_embed_dim) + elif self.num_classes == 'timestep': + self.label_emb = nn.Sequential( + Timestep(model_channels), + nn.Sequential( + linear(model_channels, time_embed_dim), + nn.SiLU(), + linear(time_embed_dim, time_embed_dim), + ), + ) + elif self.num_classes == 'sequential': + assert adm_in_channels is not None + self.label_emb = nn.Sequential( + nn.Sequential( + linear(adm_in_channels, time_embed_dim), + nn.SiLU(), + linear(time_embed_dim, time_embed_dim), + )) + else: + raise ValueError() + + self.input_blocks = nn.ModuleList([ + TimestepEmbedSequential( + conv_nd(dims, in_channels, model_channels, 3, padding=1)) + ]) + self._feature_size = model_channels + input_block_chans = [model_channels] + input_down_flag = [False] + ch = model_channels + ds = 1 + for level, mult in enumerate(channel_mult): + for nr in range(self.num_res_blocks[level]): + layers = [ + ResBlock( + ch, + time_embed_dim, + dropout, + out_channels=mult * model_channels, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + ) + ] + ch = mult * model_channels + if ds in attention_resolutions: + if num_head_channels == -1: + dim_head = ch // num_heads + else: + num_heads = ch // num_head_channels + dim_head = num_head_channels + disabled_sa = disable_self_attentions[level] if exists( + disable_self_attentions) else False + + layers.append( + SpatialTransformerV2( + ch, + num_heads, + dim_head, + transformer_block=transformer_block, + depth=transformer_depth[level], + context_dim=context_dim, + disable_self_attn=disabled_sa, + use_linear=use_linear_in_transformer, + use_checkpoint=use_checkpoint)) + self.input_blocks.append(TimestepEmbedSequential(*layers)) + self._feature_size += ch + input_block_chans.append(ch) + input_down_flag.append(False) + if level != len(channel_mult) - 1: + out_ch = ch + self.input_blocks.append( + TimestepEmbedSequential( + ResBlock( + ch, + time_embed_dim, + dropout, + out_channels=out_ch, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + down=True, + ) if resblock_updown else Downsample( + ch, conv_resample, dims=dims, out_channels=out_ch)) + ) + ch = out_ch + input_block_chans.append(ch) + input_down_flag.append(True) + ds *= 2 + self._feature_size += ch + self._input_block_chans = copy.deepcopy(input_block_chans) + self._input_down_flag = input_down_flag + + if num_head_channels == -1: + dim_head = ch // num_heads + else: + num_heads = ch // num_head_channels + dim_head = num_head_channels + self.middle_block = TimestepEmbedSequential( + ResBlock( + ch, + time_embed_dim, + dropout, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + ), + SpatialTransformerV2(ch, + num_heads, + dim_head, + transformer_block=transformer_block, + depth=transformer_depth_middle, + context_dim=context_dim, + disable_self_attn=disable_middle_self_attn, + use_linear=use_linear_in_transformer, + use_checkpoint=use_checkpoint), + ResBlock( + ch, + time_embed_dim, + dropout, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + ), + ) + self._feature_size += ch + self._middle_block_chans = [ch] + + self._output_block_chans = [] + self.output_blocks = nn.ModuleList([]) + for level, mult in list(enumerate(channel_mult))[::-1]: + for i in range(self.num_res_blocks[level] + 1): + ich = input_block_chans.pop() + layers = [ + ResBlock( + ch + ich, + time_embed_dim, + dropout, + out_channels=model_channels * mult, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + ) + ] + ch = model_channels * mult + if ds in attention_resolutions: + if num_head_channels == -1: + dim_head = ch // num_heads + else: + num_heads = ch // num_head_channels + dim_head = num_head_channels + disabled_sa = disable_self_attentions[level] if exists( + disable_self_attentions) else False + layers.append( + SpatialTransformerV2( + ch, + num_heads, + dim_head, + transformer_block=transformer_block, + depth=transformer_depth[level], + context_dim=context_dim, + disable_self_attn=disabled_sa, + use_linear=use_linear_in_transformer, + use_checkpoint=use_checkpoint)) + if level and i == self.num_res_blocks[level]: + out_ch = ch + layers.append( + ResBlock( + ch, + time_embed_dim, + dropout, + out_channels=out_ch, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + up=True, + ) if resblock_updown else Upsample( + ch, conv_resample, dims=dims, out_channels=out_ch)) + ds //= 2 + + self.output_blocks.append(TimestepEmbedSequential(*layers)) + + self._feature_size += ch + self._output_block_chans.append(ch) + + self.out = nn.Sequential( + normalization(ch), + nn.SiLU(), + zero_module( + conv_nd(dims, model_channels, out_channels, 3, padding=1)), + ) + + if self.use_refine: + self.ref_time_embed = copy.deepcopy(self.time_embed) + self.ref_label_emb = copy.deepcopy(self.label_emb) + self.ref_input_blocks = copy.deepcopy(self.input_blocks) + self.ref_input_blocks[0] = TimestepEmbedSequential( + conv_nd(dims, 4, model_channels, 3, padding=1)) + self.ref_middle_block = copy.deepcopy(self.middle_block) + self.ref_output_blocks = copy.deepcopy(self.output_blocks) + + def load_pretrained_model(self, pretrained_model): + if pretrained_model is not None: + with FS.get_from(pretrained_model, + wait_finish=True) as local_model: + self.init_from_ckpt(local_model, ignore_keys=self.ignore_keys) + + def init_from_ckpt(self, path, ignore_keys=list()): + if path.endswith('safetensors'): + from safetensors.torch import load_file as load_safetensors + sd = load_safetensors(path) + else: + sd = torch.load(path, map_location='cpu') + + new_sd = OrderedDict() + for k, v in sd.items(): + ignored = False + for ik in ignore_keys: + if ik in k: + if we.rank == 0: + self.logger.info( + 'Ignore key {} from state_dict.'.format(k)) + ignored = True + break + if not ignored: + if k == 'input_blocks.0.0.weight': + if we.rank == 0: + self.logger.info( + 'Partial initial key {} from state_dict.'.format( + k)) + new_v = torch.empty(320, self.in_channels, 3, 3) + nn.init.zeros_(new_v) + new_v[:, :v.shape[1]] = v + new_sd[k] = new_v + if self.use_refine: + new_sd['ref_' + k] = v + else: + new_sd[k] = v + if self.use_refine: + new_sd['ref_' + k] = v + + missing, unexpected = self.load_state_dict(new_sd, strict=False) + if we.rank == 0: + self.logger.info( + f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys' + ) + if len(missing) > 0: + self.logger.info(f'Missing Keys:\n {missing}') + if len(unexpected) > 0: + self.logger.info(f'\nUnexpected Keys:\n {unexpected}') + + def forward(self, x, t=None, cond=dict(), **kwargs): + t_emb = timestep_embedding(t, + self.model_channels, + repeat_only=False, + legacy=True) + emb = self.time_embed(t_emb) + + if isinstance(cond, dict): + if 'y' in cond: + assert self.num_classes is not None + emb = emb + self.label_emb(cond['y']) + if self.use_refine: + ref_emb = self.ref_time_embed(t_emb) + assert 'null_y' in cond + cond_y = cond['y'].clone() + cond_y[:, :cond['null_y'].shape[1]] = cond['null_y'] + ref_emb = ref_emb + self.ref_label_emb(cond_y) + + if 'concat' in cond: + c = cond['concat'] + x = torch.cat([x, c], dim=1) + + context = cond.get('crossattn', None) + img_context = cond.get('img_crossattn', None) + + task = cond['task'] + image_scale = cond.get('image_scale', self.image_scale) + if 'Subject' in task and img_context is not None: + ip_enc_scale = image_scale + ip_dec_scale = image_scale + num_img_tokens = img_context.shape[1] + context = torch.cat([context, img_context], dim=1) + else: + ip_enc_scale = None + ip_dec_scale = None + num_img_tokens = None + + ref = cond.get('ref_xt', None) + ref_context = cond.get('ref_crossattn', None) + else: + raise TypeError + + hs = [] + refs = [] + h = x + + if self.use_refine: + assert ref is not None and ref_context is not None + for i, (ref_module, module) in enumerate( + zip(self.ref_input_blocks, self.input_blocks)): + ref = ref_module(ref, ref_emb, ref_context, caching=None) + h = module(h, + emb, + context, + caching=None, + scale=ip_enc_scale, + num_img_token=num_img_tokens) + refs.append(ref) + hs.append(h) + + ref = self.ref_middle_block(ref, + ref_emb, + ref_context, + caching=None) + h = self.middle_block(h, + emb, + context, + caching=None, + scale=ip_enc_scale, + num_img_token=num_img_tokens) + + for i, (ref_module, module) in enumerate( + zip(self.ref_output_blocks, self.output_blocks)): + cache = [] + ref = torch.cat([ref, refs.pop()], dim=1) + ref = ref_module(ref, + ref_emb, + ref_context, + caching='write', + cache=cache) + h = torch.cat([h, hs.pop()], dim=1) + h = module(h, + emb, + context, + caching='read', + cache=cache, + scale=ip_dec_scale, + num_img_token=num_img_tokens) + else: + for module in self.input_blocks: + h = module(h, + emb, + context, + caching=None, + scale=ip_enc_scale, + num_img_token=num_img_tokens) + hs.append(h) + h = self.middle_block(h, + emb, + context, + caching=None, + scale=ip_enc_scale, + num_img_token=num_img_tokens) + for module in self.output_blocks: + h = torch.cat([h, hs.pop()], dim=1) + h = module(h, + emb, + context, + caching=None, + scale=ip_dec_scale, + num_img_token=num_img_tokens) + + out = self.out(h) + return out + + @staticmethod + def get_config_template(): + return dict_to_yaml('BACKBONE', + __class__.__name__, + LargenUNetXL.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/unet/unet_utils.py b/scepter/modules/model/backbone/unet/unet_utils.py index 69d270b..50ab617 100644 --- a/scepter/modules/model/backbone/unet/unet_utils.py +++ b/scepter/modules/model/backbone/unet/unet_utils.py @@ -10,6 +10,7 @@ import numpy as np import torch import torch.nn as nn import torch.nn.functional as F +import torchvision.transforms.functional as TF from einops import rearrange, repeat from packaging import version @@ -171,12 +172,14 @@ class TimestepEmbedSequential(nn.Sequential, TimestepBlock): A sequential module that passes timestep embeddings to the children that support it as an extra input. """ - def forward(self, x, emb, context=None, target_size=None): + def forward(self, x, emb, context=None, target_size=None, **kwargs): for layer in self: if isinstance(layer, TimestepBlock): x = layer(x, emb) elif isinstance(layer, SpatialTransformer): x = layer(x, context) + elif isinstance(layer, SpatialTransformerV2): + x = layer(x, context, **kwargs) elif isinstance(layer, Upsample): x = layer(x, target_size) else: @@ -864,6 +867,92 @@ class MemoryEfficientCrossAttention(nn.Module): return self.to_out(out) +class XFormersMHA_IP(nn.Module): + def __init__(self, + query_dim, + context_dim=None, + heads=8, + dim_head=64, + dropout=0.0): + super().__init__() + inner_dim = dim_head * heads + context_dim = default(context_dim, query_dim) + + self.heads = heads + self.dim_head = dim_head + + self.to_q = nn.Linear(query_dim, inner_dim, bias=False) + self.to_k = nn.Linear(context_dim, inner_dim, bias=False) + self.to_v = nn.Linear(context_dim, inner_dim, bias=False) + self.to_k_ip = nn.Linear(context_dim, inner_dim, bias=False) + self.to_v_ip = nn.Linear(context_dim, inner_dim, bias=False) + + self.to_out = nn.Sequential(nn.Linear(inner_dim, query_dim), + nn.Dropout(dropout)) + self.attention_op = None + + def forward(self, + x, + context=None, + mask=None, + scale=None, + num_img_token=None): + q = self.to_q(x) + context = default(context, x) + + if scale is not None and num_img_token is not None: + eos = context.shape[1] - num_img_token + txt_context = context[:, :eos, :] + img_context = context[:, eos:, :] + + k = self.to_k(txt_context) + v = self.to_v(txt_context) + k_i = self.to_k_ip(img_context) + v_i = self.to_v_ip(img_context) + + b, _, _ = q.shape + q, k, v, k_i, v_i = map( + lambda t: t.unsqueeze(3).reshape(b, t.shape[ + 1], self.heads, self.dim_head).permute(0, 2, 1, 3).reshape( + b * self.heads, t.shape[1], self.dim_head).contiguous( + ), + (q, k, v, k_i, v_i), + ) + + # actually compute the attention, what we cannot get enough of + txt_out = xformers.ops.memory_efficient_attention( + q, k, v, attn_bias=None, op=self.attention_op) + img_out = xformers.ops.memory_efficient_attention( + q, k_i, v_i, attn_bias=None, op=self.attention_op) + out = txt_out + scale * img_out + else: + k = self.to_k(context) + v = self.to_v(context) + b, _, _ = q.shape + q, k, v = map( + lambda t: t.unsqueeze(3).reshape(b, t.shape[ + 1], self.heads, self.dim_head).permute(0, 2, 1, 3).reshape( + b * self.heads, t.shape[1], self.dim_head).contiguous( + ), + (q, k, v), + ) + out = xformers.ops.memory_efficient_attention(q, + k, + v, + attn_bias=None, + op=self.attention_op) + + # TODO: Use this directly in the attention operation, as a bias + if exists(mask): + raise NotImplementedError + out = (out.unsqueeze(0).reshape( + b, self.heads, out.shape[1], + self.dim_head).permute(0, 2, 1, + 3).reshape(b, out.shape[1], + self.heads * self.dim_head)) + return self.to_out(out) + + class BasicTransformerBlock(nn.Module): def __init__(self, dim, @@ -908,6 +997,65 @@ class BasicTransformerBlock(nn.Module): return x +class TransformerBlockV2(nn.Module): + def __init__(self, + query_dim, + n_heads, + d_head, + dropout=0., + context_dim=None, + gated_ff=True, + use_checkpoint=False, + disable_self_attn=False): + super().__init__() + self.disable_self_attn = disable_self_attn + self.attn1 = MemoryEfficientCrossAttention(query_dim=query_dim, + heads=n_heads, + dim_head=d_head, + dropout=dropout, + context_dim=None) + self.ff = FeedForward(query_dim, dropout=dropout, glu=gated_ff) + self.attn2 = XFormersMHA_IP(query_dim=query_dim, + heads=n_heads, + dim_head=d_head, + context_dim=context_dim) + self.norm1 = nn.LayerNorm(query_dim) + self.norm2 = nn.LayerNorm(query_dim) + self.norm3 = nn.LayerNorm(query_dim) + self.use_checkpoint = use_checkpoint + + def forward(self, + x, + context, + caching=None, + cache=None, + scale=None, + num_img_token=None, + **kwargs): + y = self.norm1(x) + if caching == 'write': + assert isinstance(cache, list) + cache.append(y) + x = self.attn1(y, context=None) + x + elif caching == 'read': + assert isinstance(cache, list) and len(cache) > 0 + c = cache.pop(0) + self_ctx = torch.cat([y, c], dim=1) + x = self.attn1(y, context=self_ctx) + x + elif caching is None: + x = self.attn1(y, context=None) + x + else: + assert False + + x = self.attn2(self.norm2(x), + context=context, + scale=scale, + num_img_token=num_img_token) + x + x = self.ff(self.norm3(x)) + x + + return x + + class SpatialTransformer(nn.Module): """ Transformer block for image-like data. @@ -1003,3 +1151,108 @@ class SpatialTransformer(nn.Module): if not self.use_linear: x = self.proj_out(x) return x + x_in + + +class SpatialTransformerV2(nn.Module): + """ + Transformer block for image-like data. + First, project the input (aka embedding) + and reshape to b, t, d. + Then apply standard transformer action. + Finally, reshape to image + NEW: use_linear for more efficiency instead of the 1x1 convs + """ + def __init__(self, + in_channels, + n_heads, + d_head, + transformer_block, + depth=1, + dropout=0., + context_dim=None, + disable_self_attn=False, + use_linear=False, + use_checkpoint=True): + super().__init__() + if exists(context_dim) and not isinstance(context_dim, list): + context_dim = [context_dim] + + if exists(context_dim) and not isinstance(context_dim, (list)): + context_dim = [context_dim] + if exists(context_dim) and isinstance(context_dim, list): + if depth != len(context_dim): + print( + f'WARNING: {self.__class__.__name__}: Found context dims {context_dim} of' + f" depth {len(context_dim)}, which does not match the specified 'depth' of" + f' {depth}. Setting context_dim to {depth * [context_dim[0]]} now.' + ) + # depth does not match context dims. + assert all( + map(lambda x: x == context_dim[0], context_dim) + ), 'need homogenous context_dim to match depth automatically' + context_dim = depth * [context_dim[0]] + elif context_dim is None: + context_dim = [None] * depth + + self.in_channels = in_channels + inner_dim = n_heads * d_head + self.norm = normalization(in_channels) + if not use_linear: + self.proj_in = nn.Conv2d(in_channels, + inner_dim, + kernel_size=1, + stride=1, + padding=0) + else: + self.proj_in = nn.Linear(in_channels, inner_dim) + + self.transformer_blocks = nn.ModuleList([ + transformer_block(inner_dim, + n_heads, + d_head, + dropout=dropout, + context_dim=context_dim[d], + disable_self_attn=disable_self_attn, + use_checkpoint=use_checkpoint) + for d in range(depth) + ]) + if not use_linear: + self.proj_out = zero_module( + nn.Conv2d(inner_dim, + in_channels, + kernel_size=1, + stride=1, + padding=0)) + else: + self.proj_out = zero_module(nn.Linear(in_channels, inner_dim)) + self.use_linear = use_linear + + def forward(self, x, context=None, **kwargs): + # note: if no context is given, cross-attention defaults to self-attention + if not isinstance(context, list): + context = [context] + b, c, h, w = x.shape + + ref_mask = kwargs.pop('ref_mask', None) + if ref_mask is not None: + ref_mask = TF.resize(ref_mask, (h, w), antialias=True) + ref_mask = (ref_mask > 0.5).float() + ref_mask = rearrange(ref_mask, 'b c h w -> b (h w) c').contiguous() + + x_in = x + x = self.norm(x) + if not self.use_linear: + x = self.proj_in(x) + x = rearrange(x, 'b c h w -> b (h w) c').contiguous() + if self.use_linear: + x = self.proj_in(x) + for i, block in enumerate(self.transformer_blocks): + if i > 0 and len(context) == 1: + i = 0 # use same context for each block + x = block(x, context=context[i], ref_mask=ref_mask, **kwargs) + if self.use_linear: + x = self.proj_out(x) + x = rearrange(x, 'b (h w) c -> b c h w', h=h, w=w).contiguous() + if not self.use_linear: + x = self.proj_out(x) + return x + x_in diff --git a/scepter/modules/model/backbone/video/video_transformer.py b/scepter/modules/model/backbone/video/video_transformer.py index 7dca97f..783cbde 100644 --- a/scepter/modules/model/backbone/video/video_transformer.py +++ b/scepter/modules/model/backbone/video/video_transformer.py @@ -1,17 +1,5 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. - -import math - -import torch -import torch.nn as nn -import torch.nn.functional -from einops import rearrange - -from scepter.modules.model.backbone.video.init_helper import ( - _init_transformer_weights, trunc_normal_) -from scepter.modules.model.registry import BACKBONES, BRICKS, STEMS -from scepter.modules.utils.config import dict_to_yaml ''' The implementations of vivit as https://arxiv.org/abs/2103.15691. The following setting alined the proposed model in the paper above. @@ -39,6 +27,18 @@ TimesFormer: complexity: (n_h * n_w) ** 2 + O(attn_temp) ''' +import math + +import torch +import torch.nn as nn +import torch.nn.functional +from einops import rearrange + +from scepter.modules.model.backbone.video.init_helper import ( + _init_transformer_weights, trunc_normal_) +from scepter.modules.model.registry import BACKBONES, BRICKS, STEMS +from scepter.modules.utils.config import dict_to_yaml + @BACKBONES.register_class() class VideoTransformer(nn.Module): diff --git a/scepter/modules/model/embedder/__init__.py b/scepter/modules/model/embedder/__init__.py index 20e2f21..959417a 100644 --- a/scepter/modules/model/embedder/__init__.py +++ b/scepter/modules/model/embedder/__init__.py @@ -1,7 +1,10 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. + from scepter.modules.model.embedder.embedder import (ConcatTimestepEmbedderND, FrozenCLIPEmbedder, FrozenOpenCLIPEmbedder, FrozenOpenCLIPEmbedder2, - GeneralConditioner) + GeneralConditioner, + IPAdapterPlusEmbedder, + RefCrossEmbedder) diff --git a/scepter/modules/model/embedder/clip.py b/scepter/modules/model/embedder/clip.py new file mode 100644 index 0000000..6884648 --- /dev/null +++ b/scepter/modules/model/embedder/clip.py @@ -0,0 +1,1059 @@ +# -*- coding: utf-8 -*- +"""Concise re-implementation of ``https://github.com/openai/CLIP'' and + ``https://github.com/mlfoundations/open_clip''. +""" +import math +from functools import partial +from importlib import find_loader + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from scepter.modules.model.base_model import BaseModel +from scepter.modules.model.embedder.xlm_roberta import \ + XLMRoberta # used in XLMRobertaCLIP (multilingual) +from scepter.modules.model.registry import EMBEDDERS +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we +from scepter.modules.utils.file_system import FS + + +def map_dtype(m, dtype=torch.float16): + if isinstance(m, (nn.Linear, nn.Conv2d)): + _ = m.to(dtype) + elif isinstance(m, LayerNorm): + _ = m.float() + elif hasattr(m, 'head') and isinstance(m.head, nn.Parameter): + p = getattr(m, 'head') + p.data = p.data.to(dtype) + + +class QuickGELU(nn.Module): + def forward(self, x): + return x * torch.sigmoid(1.702 * x) + + +class LayerNorm(nn.LayerNorm): + def forward(self, x): + return super().forward(x.float()).type_as(x) + + +class SelfAttention(nn.Module): + def __init__(self, + dim, + num_heads, + causal=False, + attn_dropout=0.0, + proj_dropout=0.0, + flash_dtype=torch.float16): + assert dim % num_heads == 0 + assert flash_dtype in (None, torch.float16, torch.bfloat16) + super().__init__() + self.dim = dim + self.num_heads = num_heads + self.head_dim = dim // num_heads + self.causal = causal + self.attn_dropout = attn_dropout + self.proj_dropout = proj_dropout + self.scale = math.pow(self.head_dim, -0.25) + self.flash_dtype = flash_dtype + + # layers + self.to_qkv = nn.Linear(dim, dim * 3) + self.proj = nn.Linear(dim, dim) + + def forward(self, x): + """x: [B, L, C]. + """ + b, s, c, n, d = *x.size(), self.num_heads, self.head_dim + + # compute query, key, value + qkv = self.to_qkv(x).view(b, s, 3, n, d) + + # compute attention + if x.device.type != 'cpu' and find_loader('flash_attn') and \ + self.flash_dtype is not None: + # flash implementation + from flash_attn.flash_attn_interface import ( + flash_attn_unpadded_qkvpacked_func, ) + dtype = qkv.dtype + if dtype != self.flash_dtype: + qkv = qkv.type(self.flash_dtype) + cu_seqlens = torch.arange(0, + b * s + 1, + s, + dtype=torch.int32, + device=x.device) + x = flash_attn_unpadded_qkvpacked_func( + qkv=qkv.reshape(-1, 3, n, d), + cu_seqlens=cu_seqlens, + max_seqlen=s, + dropout_p=self.attn_dropout if self.training else 0.0, + causal=self.causal, + return_attn_probs=False).reshape(b, s, n, d).type(dtype) + else: + # torch implementation + q, k, v = qkv.unbind(2) + attn = torch.einsum('binc,bjnc->bnij', q * self.scale, + k * self.scale) + if self.causal: + attn = attn.masked_fill( + torch.tril(attn.new_ones(1, 1, s, + s).float()).type_as(attn) == 0, + float('-inf')) + attn = F.softmax(attn.float(), dim=-1).type_as(attn) + x = torch.einsum('bnij,bjnc->binc', attn, v) + + # output + x = x.reshape(b, s, c) + x = self.proj(x) + x = F.dropout(x, self.proj_dropout, self.training) + return x + + +class AttentionBlock(nn.Module): + def __init__(self, + dim, + mlp_ratio, + num_heads, + causal=False, + activation='quick_gelu', + attn_dropout=0.0, + proj_dropout=0.0, + flash_dtype=torch.float16): + assert activation in ['quick_gelu', 'gelu'] + super().__init__() + self.dim = dim + self.mlp_ratio = mlp_ratio + self.num_heads = num_heads + self.causal = causal + self.flash_dtype = flash_dtype + + # layers + self.norm1 = LayerNorm(dim) + self.attn = SelfAttention(dim, num_heads, causal, attn_dropout, + proj_dropout, flash_dtype) + self.norm2 = LayerNorm(dim) + self.mlp = nn.Sequential( + nn.Linear(dim, int(dim * mlp_ratio)), + QuickGELU() if activation == 'quick_gelu' else nn.GELU(), + nn.Linear(int(dim * mlp_ratio), dim), nn.Dropout(proj_dropout)) + + def forward(self, x): + x = x + self.attn(self.norm1(x)) + x = x + self.mlp(self.norm2(x)) + return x + + +class VisionTransformer(nn.Module): + def __init__(self, + image_size=224, + patch_size=16, + dim=768, + mlp_ratio=4, + out_dim=512, + num_heads=12, + num_layers=12, + activation='quick_gelu', + attn_dropout=0.0, + proj_dropout=0.0, + embedding_dropout=0.0, + flash_dtype=torch.float16): + assert image_size % patch_size == 0 + super().__init__() + self.image_size = image_size + self.patch_size = patch_size + self.num_patches = (image_size // patch_size)**2 + self.dim = dim + self.mlp_ratio = mlp_ratio + self.out_dim = out_dim + self.num_heads = num_heads + self.num_layers = num_layers + self.flash_dtype = flash_dtype + + # embeddings + gain = 1.0 / math.sqrt(dim) + self.patch_embedding = nn.Conv2d(3, + dim, + kernel_size=patch_size, + stride=patch_size, + bias=False) + self.cls_embedding = nn.Parameter(gain * torch.randn(1, 1, dim)) + self.pos_embedding = nn.Parameter( + gain * torch.randn(1, self.num_patches + 1, dim)) + self.dropout = nn.Dropout(embedding_dropout) + + # transformer + self.pre_norm = LayerNorm(dim) + self.transformer = nn.Sequential(*[ + AttentionBlock(dim, mlp_ratio, num_heads, False, activation, + attn_dropout, proj_dropout, flash_dtype) + for _ in range(num_layers) + ]) + self.post_norm = LayerNorm(dim) + + # head + self.head = nn.Parameter(gain * torch.randn(dim, out_dim)) + + def forward(self, x): + b, dtype = x.size(0), self.head.dtype + x = x.type(dtype) + + # patch-embedding + x = self.patch_embedding(x).flatten(2).permute(0, 2, 1) + x = torch.cat([self.cls_embedding.repeat(b, 1, 1).type(dtype), x], + dim=1) + x = self.dropout(x + self.pos_embedding.type(dtype)) + x = self.pre_norm(x) + + # transformer + x = self.transformer(x) + + # head + x = self.post_norm(x) + x = torch.mm(x[:, 0, :], self.head) + return x + + def fp16(self, dtype=torch.float16): + return self.apply(partial(map_dtype, dtype=dtype)) + + +class TextTransformer(nn.Module): + def __init__(self, + vocab_size, + text_len, + dim=512, + mlp_ratio=4, + out_dim=512, + num_heads=8, + num_layers=12, + activation='quick_gelu', + attn_dropout=0.0, + proj_dropout=0.0, + embedding_dropout=0.0, + flash_dtype=torch.float16): + super().__init__() + self.vocab_size = vocab_size + self.text_len = text_len + self.dim = dim + self.mlp_ratio = mlp_ratio + self.out_dim = out_dim + self.num_heads = num_heads + self.num_layers = num_layers + self.flash_dtype = flash_dtype + + # embeddings + self.token_embedding = nn.Embedding(vocab_size, dim) + self.pos_embedding = nn.Parameter(0.01 * torch.randn(1, text_len, dim)) + self.dropout = nn.Dropout(embedding_dropout) + + # transformer + self.transformer = nn.Sequential(*[ + AttentionBlock(dim, mlp_ratio, num_heads, True, activation, + attn_dropout, proj_dropout, flash_dtype) + for _ in range(num_layers) + ]) + self.norm = LayerNorm(dim) + + # head + gain = 1.0 / math.sqrt(dim) + self.head = nn.Parameter(gain * torch.randn(dim, out_dim)) + + def forward(self, x): + eot, dtype = x.argmax(dim=-1), self.head.dtype + + # embeddings + x = self.dropout( + self.token_embedding(x).type(dtype) + + self.pos_embedding.type(dtype)) + + # transformer + x = self.transformer(x) + + # head + x = self.norm(x) + x = torch.mm(x[torch.arange(x.size(0)), eot], self.head) + return x + + def fp16(self, dtype=torch.float16): + return self.apply(partial(map_dtype, dtype=dtype)) + + +class CLIP(nn.Module): + def __init__(self, + embed_dim=512, + image_size=224, + patch_size=16, + vision_dim=768, + vision_mlp_ratio=4, + vision_heads=12, + vision_layers=12, + vocab_size=49408, + text_len=77, + text_dim=512, + text_mlp_ratio=4, + text_heads=8, + text_layers=12, + activation='quick_gelu', + attn_dropout=0.0, + proj_dropout=0.0, + embedding_dropout=0.0, + flash_dtype=torch.float16, + use_module=['visual', 'textual']): + assert flash_dtype in (None, torch.float16, torch.bfloat16) + super().__init__() + self.embed_dim = embed_dim + self.image_size = image_size + self.patch_size = patch_size + self.vision_dim = vision_dim + self.vision_mlp_ratio = vision_mlp_ratio + self.vision_heads = vision_heads + self.vision_layers = vision_layers + self.vocab_size = vocab_size + self.text_len = text_len + self.text_dim = text_dim + self.text_mlp_ratio = text_mlp_ratio + self.text_heads = text_heads + self.text_layers = text_layers + self.flash_dtype = flash_dtype + self.use_module = use_module + # models + if 'visual' in use_module: + self.visual = VisionTransformer( + image_size=image_size, + patch_size=patch_size, + dim=vision_dim, + mlp_ratio=vision_mlp_ratio, + out_dim=embed_dim, + num_heads=vision_heads, + num_layers=vision_layers, + activation=activation, + attn_dropout=attn_dropout, + proj_dropout=proj_dropout, + embedding_dropout=embedding_dropout, + flash_dtype=flash_dtype) + self.scale = math.sqrt(self.visual.out_dim) + else: + self.visual = nn.Identity() + if 'textual' in use_module: + self.textual = TextTransformer(vocab_size=vocab_size, + text_len=text_len, + dim=text_dim, + mlp_ratio=text_mlp_ratio, + out_dim=embed_dim, + num_heads=text_heads, + num_layers=text_layers, + activation=activation, + attn_dropout=attn_dropout, + proj_dropout=proj_dropout, + embedding_dropout=embedding_dropout, + flash_dtype=flash_dtype) + else: + self.textual = nn.Identity() + self.log_scale = nn.Parameter(math.log(1 / 0.07) * torch.ones([])) + + # initialize weights + self.init_weights() + + def forward(self, imgs, txt_tokens): + """imgs: [B, 3, H, W] of torch.float32. + mean: [0.48145466, 0.4578275, 0.40821073] + std: [0.26862954, 0.26130258, 0.27577711] + txt_tokens: [B, L] of torch.long. + Encoded by data.CLIPTokenizer. + """ + xi = self.visual(imgs) + xt = self.textual(txt_tokens) + return xi, xt + + def encode_image(self, x, skip_layers=0): + # clip inference + b, dtype = x.size(0), self.visual.head.dtype + x = x.type(dtype) + + # # patch-embedding + x = self.visual.patch_embedding(x).flatten(2).permute(0, 2, 1) + x = torch.cat( + [self.visual.cls_embedding.repeat(b, 1, 1).type(dtype), x], dim=1) + x = self.visual.dropout(x + self.visual.pos_embedding.type(dtype)) + x = self.visual.pre_norm(x) + + # # transformer + # assert skip_layers < 12 + if skip_layers == 0: + x = self.visual.transformer(x) + else: + for m in self.visual.transformer[:-skip_layers]: + x = m(x) + + # # head + x = self.visual.post_norm(x) + x = torch.mm(x[:, 0, :], self.visual.head) + x = self.scale * F.normalize(x, p=2, dim=1) + return x + + def encode_text(self): + pass + + def init_weights(self): + # embeddings + if 'textual' in self.use_module: + nn.init.normal_(self.textual.token_embedding.weight, std=0.02) + if 'visual' in self.use_module: + nn.init.normal_(self.visual.patch_embedding.weight, std=0.1) + + # attentions + for modality in self.use_module: + dim = self.vision_dim if modality == 'visual' else self.text_dim + transformer = getattr(self, modality).transformer + proj_gain = (1.0 / math.sqrt(dim)) * ( + 1.0 / math.sqrt(2 * len(transformer))) + attn_gain = 1.0 / math.sqrt(dim) + mlp_gain = 1.0 / math.sqrt(2.0 * dim) + for block in transformer: + nn.init.normal_(block.attn.to_qkv.weight, std=attn_gain) + nn.init.normal_(block.attn.proj.weight, std=proj_gain) + nn.init.normal_(block.mlp[0].weight, std=mlp_gain) + nn.init.normal_(block.mlp[2].weight, std=proj_gain) + + def param_groups(self): + groups = [{ + 'params': [ + p for n, p in self.named_parameters() + if 'norm' in n or n.endswith('bias') + ], + 'weight_decay': + 0.0 + }, { + 'params': [ + p for n, p in self.named_parameters() + if not ('norm' in n or n.endswith('bias')) + ] + }] + return groups + + def fp16(self, dtype=torch.float16): + return self.apply(partial(map_dtype, dtype=dtype)) + + def load_from_open_clip(self, checkpoint_or_path, **kwargs): + """Load and remap state-dict from open-clip. + """ + # load state-dict + device = next(self.parameters()).device + state = checkpoint_or_path + if isinstance(state, str): + state = torch.load(state, map_location=device) + + # reorder + prefix = [ + 'logit_scale', 'visual.', 'position', 'text_proj', 'token', + 'transformer.', 'ln_final.' + ] + state = type(state)([(k, v) for u in prefix for k, v in state.items() + if k.startswith(u)]) + + # convert to target keys + target = self.state_dict() + target = { + k: v.view(target[k].shape) + for k, v in zip(target.keys(), state.values()) + } + return self.load_state_dict(target, **kwargs) + + +class XLMRobertaWithHead(XLMRoberta): + def __init__(self, **kwargs): + self.out_dim = kwargs.pop('out_dim') + super().__init__(**kwargs) + + # head + mid_dim = (self.dim + self.out_dim) // 2 + self.head = nn.Sequential(nn.Linear(self.dim, mid_dim, bias=False), + nn.GELU(), + nn.Linear(mid_dim, self.out_dim, bias=False)) + + def forward(self, tokens): + # xlm-roberta + x = super().forward(tokens) + + # average pooling + mask = tokens.ne(self.pad_token).unsqueeze(-1).to(x) + x = (x * mask).sum(dim=1) / mask.sum(dim=1) + + # head + x = self.head(x) + return x + + +class XLMRobertaCLIP(nn.Module): + def __init__(self, + embed_dim=1024, + image_size=224, + patch_size=14, + vision_dim=1280, + vision_mlp_ratio=4, + vision_heads=16, + vision_layers=32, + activation='gelu', + vocab_size=250002, + max_text_len=514, + type_size=1, + pad_token=1, + text_dim=1024, + text_heads=16, + text_layers=24, + text_eps=1e-5, + text_dropout=0.1, + attn_dropout=0.0, + proj_dropout=0.0, + embedding_dropout=0.0, + flash_dtype=torch.float16, + use_module=['visual', 'textual']): + assert flash_dtype in (None, torch.float16, torch.bfloat16) + super().__init__() + self.embed_dim = embed_dim + self.image_size = image_size + self.patch_size = patch_size + self.vision_dim = vision_dim + self.vision_mlp_ratio = vision_mlp_ratio + self.vision_heads = vision_heads + self.vision_layers = vision_layers + self.activation = activation + self.vocab_size = vocab_size + self.max_text_len = max_text_len + self.type_size = type_size + self.pad_token = pad_token + self.text_dim = text_dim + self.text_heads = text_heads + self.text_layers = text_layers + self.text_eps = text_eps + self.flash_dtype = flash_dtype + + # models + if 'visual' in use_module: + self.visual = VisionTransformer( + image_size=image_size, + patch_size=patch_size, + dim=vision_dim, + mlp_ratio=vision_mlp_ratio, + out_dim=embed_dim, + num_heads=vision_heads, + num_layers=vision_layers, + activation=activation, + attn_dropout=attn_dropout, + proj_dropout=proj_dropout, + embedding_dropout=embedding_dropout, + flash_dtype=flash_dtype) + else: + self.visual = nn.Identity() + if 'textual' in use_module: + self.textual = XLMRobertaWithHead(vocab_size=vocab_size, + max_seq_len=max_text_len, + type_size=type_size, + pad_token=pad_token, + dim=text_dim, + out_dim=embed_dim, + num_heads=text_heads, + num_layers=text_layers, + dropout=text_dropout, + eps=text_eps) + else: + self.textual = nn.Identity() + self.log_scale = nn.Parameter(math.log(1 / 0.07) * torch.ones([])) + + def forward(self, imgs, txt_tokens): + """imgs: [B, 3, H, W] of torch.float32. + mean: [0.48145466, 0.4578275, 0.40821073] + std: [0.26862954, 0.26130258, 0.27577711] + txt_tokens: [B, L] of torch.long. + Encoded by data.CLIPTokenizer. + """ + xi = self.visual(imgs) + xt = self.textual(txt_tokens) + return xi, xt + + def param_groups(self): + groups = [{ + 'params': [ + p for n, p in self.named_parameters() + if 'norm' in n or n.endswith('bias') + ], + 'weight_decay': + 0.0 + }, { + 'params': [ + p for n, p in self.named_parameters() + if not ('norm' in n or n.endswith('bias')) + ] + }] + return groups + + def fp16(self, dtype=torch.float16): + return self.apply(partial(map_dtype, dtype=dtype)) + + def load_from_open_clip(self, checkpoint_or_path, **kwargs): + """Load and remap state-dict from open-clip. + """ + # load state-dict + device = next(self.parameters()).device + state = checkpoint_or_path + if isinstance(state, str): + state = torch.load(state, map_location=device) + if 'state_dict' in state: + state = state['state_dict'] + + # reorder + keys = [ + 'logit_scale', 'visual.', 'word_embeddings', + 'token_type_embeddings', 'position_embeddings', + 'embeddings.LayerNorm', 'encoder.layer.', 'text.proj' + ] + state = type(state)([(k, v) for u in keys for k, v in state.items() + if u in k]) + + # target state-dict + target = self.state_dict() + target = { + k: v.view(target[k].shape) + for k, v in zip(target.keys(), state.values()) + } + return self.load_state_dict(target, **kwargs) + + +def _clip(pretrained=False, pretrained_path=None, model_cls=CLIP, **kwargs): + model = model_cls(**kwargs) + if pretrained and pretrained_path: + pretrain_model = torch.load(pretrained_path, map_location='cpu') + key_str = ' '.join(list(pretrain_model.keys())) + have_load = False + if 'use_module' in kwargs and len(kwargs['use_module']) < 2: + for module_name in kwargs['use_module']: + if hasattr(model, module_name) and module_name not in key_str: + missing, unexpected = getattr( + model, module_name).load_state_dict(pretrain_model, + strict=False) + if we.rank == 0: + print(f'Restored from {pretrained_path} with' + '{len(missing)} missing and {len(unexpected)}' + 'unexpected keys') + if len(missing) > 0: + print(f'Missing Keys:\n {missing}') + if len(unexpected) > 0: + print(f'\nUnexpected Keys:\n {unexpected}') + have_load = True + if not have_load: + missing, unexpected = model.load_state_dict(pretrain_model, + strict=False) + if we.rank == 0: + print( + f'Restored from {pretrained_path} with {len(missing)} missing and {len(unexpected)} unexpected keys' + ) + if len(missing) > 0: + print(f'Missing Keys:\n {missing}') + if len(unexpected) > 0: + print(f'\nUnexpected Keys:\n {unexpected}') + return model + + +def clip_vit_b_32(**kwargs): + cfg = dict(embed_dim=512, + image_size=224, + patch_size=32, + vision_dim=768, + vision_heads=12, + vision_layers=12, + vocab_size=49408, + text_len=77, + text_dim=512, + text_heads=8, + text_layers=12, + activation='quick_gelu') + cfg.update(**kwargs) + return cfg, CLIP + + +def clip_vit_b_16(**kwargs): + cfg = dict(embed_dim=512, + image_size=224, + patch_size=16, + vision_dim=768, + vision_heads=12, + vision_layers=12, + vocab_size=49408, + text_len=77, + text_dim=512, + text_heads=8, + text_layers=12, + activation='quick_gelu') + cfg.update(**kwargs) + return cfg, CLIP + + +def clip_vit_l_14(**kwargs): + cfg = dict(embed_dim=768, + image_size=224, + patch_size=14, + vision_dim=1024, + vision_heads=16, + vision_layers=24, + vocab_size=49408, + text_len=77, + text_dim=768, + text_heads=12, + text_layers=12, + activation='quick_gelu') + cfg.update(**kwargs) + return cfg, CLIP + + +def clip_vit_l_14_336px(**kwargs): + cfg = dict(embed_dim=768, + image_size=336, + patch_size=14, + vision_dim=1024, + vision_heads=16, + vision_layers=24, + vocab_size=49408, + text_len=77, + text_dim=768, + text_heads=12, + text_layers=12, + activation='quick_gelu') + cfg.update(**kwargs) + return cfg, CLIP + + +def clip_vit_h_14(**kwargs): + cfg = dict(embed_dim=1024, + image_size=224, + patch_size=14, + vision_dim=1280, + vision_heads=16, + vision_layers=32, + vocab_size=49408, + text_len=77, + text_dim=1024, + text_heads=16, + text_layers=24, + activation='gelu') + cfg.update(**kwargs) + return cfg, CLIP + + +def clip_vit_g_14(**kwargs): + cfg = dict(embed_dim=1024, + image_size=224, + patch_size=14, + vision_dim=1408, + vision_mlp_ratio=4.3637, + vision_heads=16, + vision_layers=40, + vocab_size=49408, + text_len=77, + text_dim=1024, + text_heads=16, + text_layers=24, + activation='gelu') + cfg.update(**kwargs) + return cfg, CLIP + + +def clip_vit_bigG_14(**kwargs): + cfg = dict(embed_dim=1280, + image_size=224, + patch_size=14, + vision_dim=1664, + vision_mlp_ratio=4.9231, + vision_heads=16, + vision_layers=48, + vocab_size=49408, + text_len=77, + text_dim=1280, + text_heads=20, + text_layers=32, + activation='gelu') + cfg.update(**kwargs) + return cfg, CLIP + + +def clip_xlm_roberta_vit_h_14(**kwargs): + cfg = dict(embed_dim=1024, + image_size=224, + patch_size=14, + vision_dim=1280, + vision_mlp_ratio=4, + vision_heads=16, + vision_layers=32, + activation='gelu', + vocab_size=250002, + max_text_len=514, + type_size=1, + pad_token=1, + text_dim=1024, + text_heads=16, + text_layers=24, + text_eps=1e-5, + text_dropout=0.1, + attn_dropout=0.0, + proj_dropout=0.0, + embedding_dropout=0.0, + flash_dtype=torch.float16) + cfg.update(**kwargs) + return cfg, XLMRobertaCLIP + + +clip_functions = { + 'clip_vit_b_32': clip_vit_b_32, + 'clip_vit_b_16': clip_vit_b_16, + 'clip_vit_l_14': clip_vit_l_14, + 'clip_vit_l_14_336px': clip_vit_l_14_336px, + 'clip_vit_h_14': clip_vit_h_14, + 'clip_vit_g_14': clip_vit_g_14, + 'clip_vit_bigG_14': clip_vit_bigG_14, + 'clip_xlm_roberta_vit_h_14': clip_xlm_roberta_vit_h_14 +} + + +@EMBEDDERS.register_class() +class ClipEncoder(BaseModel): + para_dict = { + 'CLIP_FUNC': { + 'value': 'clip_vit_b_32', + 'description': + f'Select clip model from {list(clip_functions.keys())}' + }, + 'PRETRAINED': { + 'value': False, + 'description': 'Wether load from pretrained model or not.' + }, + 'USE_GRAD': { + 'value': False, + 'description': '' + }, + 'PRETRAINED_PATH': { + 'value': None, + 'description': 'Pretrained model load from.' + }, + 'USE_MODULE': { + 'value': ['visual', 'textual'], + 'description': + "Use module from visual or textual, default is ['visual', 'textual']." + }, + 'CLIP_SKIP': { + 'value': 2, + 'description': "Textuxl branch skip blocks' num. Default is 2." + }, + 'TOKEN_LENGTH': { + 'value': 77, + 'description': 'The input token length for text. Default is 77.' + }, + 'KWARGS': {} + } + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + clip_func = cfg.CLIP_FUNC + pretrained = cfg.get('PRETRAINED', False) + pretrained_path = cfg.get('PRETRAINED_PATH', None) + use_module = cfg.get('USE_MODULE', ['visual', 'textual']) + self.clip_skip = cfg.get('CLIP_SKIP', 2) + self.use_grad = cfg.get('USE_GRAD', False) + self.token_length = cfg.get('TOKEN_LENGTH', 77) + kwargs = {k.lower(): v for k, v in cfg.get('KWARGS', {}).items()} + if pretrained and pretrained_path: + local_path = FS.get_from(pretrained_path, wait_finish=True) + else: + local_path = None + assert clip_func in clip_functions + if clip_func in clip_functions: + conf, model_cls = clip_functions[clip_func](**kwargs) + conf['use_module'] = use_module + self.clip_model = _clip(pretrained, local_path, model_cls, **conf) + for module in use_module: + if hasattr(self.clip_model, module): + setattr(self, module, getattr(self.clip_model, module)) + + def encode_image(self, image): + if not self.use_grad: + with torch.no_grad(): + m = self.clip_model.visual + return m(image) + else: + m = self.clip_model.visual + return m(image) + + def encode_text(self, + tokens, + tokenizer=None, + append_sentence_embedding=True): + def fn(): + m = self.clip_model.textual + b, s = tokens.shape + mask = tokens.ne(m.pad_token).long() + # embeddings + x = m.token_embedding(tokens) + \ + m.type_embedding(torch.zeros_like(tokens)) + \ + m.pos_embedding(m.pad_token + torch.cumsum(mask, dim=1) * mask) + x = m.norm(x) + x = m.dropout(x) + + # blocks + for block in m.blocks[:-1]: + x = block(x, mask.view(b, 1, 1, s)) + words = x + + sentence = m.blocks[-1](x, mask.view(b, 1, 1, s)) + mask = tokens.ne(m.pad_token).unsqueeze(-1).to(sentence) + sentence = (sentence * mask).sum(dim=1) / mask.sum(dim=1) + sentence = m.head(sentence) + + return {'crossattn': words, 'y': sentence} + + if not self.use_grad: + with torch.no_grad(): + return fn() + else: + return fn() + + def dynamic_encode_text(self, + all_tokens, + tokenizer=None, + append_sentence_embedding=True): + ''' + m: clip model + t: tokenzer + tokens: tensor(1, N) + ''' + if tokenizer is None: + tokenizer = self.tokenizer + + def fn(): + m = self.clip_model.textual + ret_data = {'crossattn': [], 'y': []} + for tokens_id in range(all_tokens.shape[0]): + tokens = all_tokens[tokens_id] + text_len = self.token_length + device = tokens.device + dtype = m.type_embedding.weight.dtype + # special tokens + sos_emb, eos_emb, pad_emb = m.token_embedding( + torch.LongTensor([ + tokenizer.sos_token, tokenizer.eos_token, + tokenizer.pad_token + ]).to(device)).type(dtype).chunk(3) + + # get raw input tokens + tokens = list(tokens.cpu().numpy()) + while tokens[-1] == tokenizer.pad_token: + tokens = tokens[:-1] + tokens = tokens[1:-1] + embeds = m.token_embedding( + torch.LongTensor(tokens).to(device)).type(dtype) + + # split into chunks to support any-length text + chunk_embeds, chunk_tokens = [], [] + max_words = text_len - 2 + if len(tokens) == 0: + chunk = torch.cat([sos_emb, eos_emb]) + chunk = torch.cat( + [chunk, + pad_emb.repeat(text_len - len(chunk), 1)]) + chunk_embeds.append(chunk) + + chunk = torch.LongTensor([tokenizer.sos_token] + + [tokenizer.eos_token]) + chunk = torch.cat([ + chunk, + torch.LongTensor([tokenizer.pad_token] * + (text_len - len(chunk))) + ]) + chunk_tokens.append(chunk) + else: + while len(tokens) > 0: + # find splitting position + if len(tokens) <= max_words: + pos = len(tokens) + else: + pos = [ + i for i, u in enumerate(tokens[:max_words]) + if u == tokenizer.comma_token + ] + pos = max_words if len(pos) == 0 else pos[-1] + 1 + + # collect chunk + chunk = torch.cat([sos_emb, embeds[:pos], eos_emb]) + chunk = torch.cat( + [chunk, + pad_emb.repeat(text_len - len(chunk), 1)]) + chunk_embeds.append(chunk) + + chunk = torch.LongTensor([tokenizer.sos_token] + + tokens[:pos] + + [tokenizer.eos_token]) + chunk = torch.cat([ + chunk, + torch.LongTensor([tokenizer.pad_token] * + (text_len - len(chunk))) + ]) + chunk_tokens.append(chunk) + + # update + tokens = tokens[pos:] + embeds = embeds[pos:] + # loop over chunks + words = [] + sentences = [] + for i, (chunk_token, chunk_embed) in enumerate( + zip(chunk_tokens, chunk_embeds)): + chunk_token = chunk_token.unsqueeze(0).to(device) + chunk_embed = chunk_embed.unsqueeze(0).to(device) + # embeddings + mask = chunk_token.ne(tokenizer.pad_token).long() + x = chunk_embed.type(dtype) + m.type_embedding( + torch.zeros_like(chunk_token)) + m.pos_embedding( + m.pad_token + torch.cumsum(mask, dim=1) * mask) + + x = m.norm(x) + x = m.dropout(x) + blocks = m.blocks[:-(self.clip_skip - 1)] + for block in blocks: + x = block(x, mask.view(1, 1, 1, -1)) + print('word', torch.sum(x)) + words.append(x.clone()) + + # if append_sentence_embedding: + # last layers + blocks = m.blocks[-(self.clip_skip - 1):] + for block in blocks: + x = block(x, mask.view(1, 1, 1, -1)) + + # get global embedding + x = (x * mask.unsqueeze(2)).sum(dim=1) / mask.sum(dim=1) + x = m.head(x) + print('sentence', torch.sum(x)) + # output + sentences.append(x.unsqueeze(0)) + sentence = torch.cat(sentences, dim=0).mean(dim=0) + words = torch.cat(words, dim=1) + ret_data['crossattn'] = words + ret_data['y'] = sentence + # ret_data['y'].append(sentence) + # ret_data.append(torch.cat([sentence] + words, dim=1)) + # ret_data['crossattn'].append(torch.cat(words, dim=1)) + # ret_data['crossattn'] = torch.cat(ret_data['crossattn'], dim=0) + # ret_data['y'] = torch.cat(ret_data['y'], dim=0) + return ret_data + + if not self.use_grad: + with torch.no_grad(): + return fn() + else: + return fn() + + @staticmethod + def get_config_template(): + return dict_to_yaml('BACKBONE', + __class__.__name__, + ClipEncoder.para_dict, + set_name=True) diff --git a/scepter/modules/model/embedder/embedder.py b/scepter/modules/model/embedder/embedder.py index 1e1addf..76ef20f 100644 --- a/scepter/modules/model/embedder/embedder.py +++ b/scepter/modules/model/embedder/embedder.py @@ -22,9 +22,10 @@ from scepter.modules.utils.distribute import we from scepter.modules.utils.file_system import FS from .base_embedder import BaseEmbedder +from .resampler import Resampler try: - from transformers import CLIPTextModel, CLIPTokenizer + from transformers import CLIPTextModel, CLIPTokenizer, CLIPVisionModelWithProjection except Exception as e: warnings.warn( f'Import transformers error, please deal with this problem: {e}') @@ -513,6 +514,93 @@ class ConcatTimestepEmbedderND(BaseEmbedder): set_name=True) +@EMBEDDERS.register_class() +class IPAdapterPlusEmbedder(BaseEmbedder): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + + with FS.get_dir_to_local_dir(cfg.CLIP_DIR, + wait_finish=True) as local_path: + self.image_encoder = CLIPVisionModelWithProjection.from_pretrained( + local_path) + + self.image_proj_model = Resampler( + dim=self.cfg.get('IN_DIM', 768), + depth=self.cfg.get('DEPTH', 4), + dim_head=64, + heads=self.cfg.get('HEADS', 12), + num_queries=self.cfg.get('NUM_TOKENS', 16), + embedding_dim=self.image_encoder.config.hidden_size, + output_dim=self.cfg.get('CROSSATTN_DIM', 768), + ff_mult=4, + ) + + with FS.get_from(cfg.PRETRAINED_MODEL, wait_finish=True) as local_path: + ckpt = torch.load(local_path, map_location='cpu') + self.image_proj_model.load_state_dict(ckpt['image_proj'], + strict=True) + + self.patch_projector = nn.Linear(self.image_encoder.config.hidden_size, + self.cfg.get('CROSSATTN_DIM', 768)) + + def encode(self, ref_ip, ref_detail): + encoder_output = self.image_encoder(ref_ip, output_hidden_states=True) + image_prompt_embeds = self.image_proj_model( + encoder_output.hidden_states[-2]) + encoder_output_2 = self.image_encoder(ref_detail, + output_hidden_states=True) + image_patch_embeds = self.patch_projector( + encoder_output_2.last_hidden_state) + out = { + 'img_crossattn': image_prompt_embeds, + 'ref_crossattn': image_patch_embeds, + } + return out + + def forward(self, ref_ip, ref_detail): + return self.encode(ref_ip, ref_detail) + + +class RefCrossEmbedder(BaseEmbedder): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + + with FS.get_dir_to_local_dir(cfg.CLIP_DIR, + wait_finish=True) as local_path: + self.image_encoder = CLIPVisionModelWithProjection.from_pretrained( + local_path) + + self.patch_projector = nn.Linear(self.image_encoder.config.hidden_size, + self.cfg.get('CROSSATTN_DIM', 768)) + + def encode(self, img): + encoder_output = self.image_encoder(img, output_hidden_states=True) + image_patch_embeds = self.patch_projector( + encoder_output.last_hidden_state) + out = { + 'ref_crossattn': image_patch_embeds, + } + return out + + def forward(self, img): + return self.encode(img) + + +@EMBEDDERS.register_class() +class TransparentEmbedder(BaseEmbedder): + def forward(self, *args): + out = dict() + for key, val in zip(self.input_keys, args): + out[key] = val + return out + + +@EMBEDDERS.register_class() +class NoiseConcatEmbedder(BaseEmbedder): + def forward(self, *args): + return {'concat': torch.cat(args, dim=1)} + + @EMBEDDERS.register_class() class GeneralConditioner(BaseEmbedder): OUTPUT_DIM2KEYS = {2: 'y', 3: 'crossattn', 4: 'concat', 5: 'concat'} @@ -598,42 +686,60 @@ class GeneralConditioner(BaseEmbedder): with embedding_context(): if hasattr(embedder, 'input_key') and (embedder.input_key is not None): + if embedder.input_key not in batch: + continue if embedder.legacy_ucg_val is not None: batch = self.possibly_get_ucg_val(embedder, batch) emb_out = embedder(batch[embedder.input_key]) elif hasattr(embedder, 'input_keys'): + if any([k not in batch for k in embedder.input_keys]): + continue emb_out = embedder( *[batch[k] for k in embedder.input_keys]) - assert isinstance( - emb_out, (torch.Tensor, list, tuple) - ), f'encoder outputs must be tensors or a sequence, but got {type(emb_out)}' - if not isinstance(emb_out, (list, tuple)): - emb_out = [emb_out] - for emb in emb_out: - # print("emb.shape", emb.shape) - # print("emb.input_keys", embedder.input_keys) - out_key = self.OUTPUT_DIM2KEYS[emb.dim()] - if embedder.ucg_rate > 0.0 and embedder.legacy_ucg_val is None: - emb = (expand_dims_like( - torch.bernoulli( - (1.0 - embedder.ucg_rate) * - torch.ones(emb.shape[0], device=emb.device)), - emb, - ) * emb) - if (hasattr(embedder, 'input_keys')): - if np.sum( - np.array([ - key in force_zero_embeddings - for key in embedder.input_keys - ])) > 0: - emb = torch.zeros_like(emb) - if out_key in output: - output[out_key] = torch.cat((output[out_key], emb), - self.KEY2CATDIM[out_key]) - else: - output[out_key] = emb - # if "y" in output: - # print("out.shape", output["y"].shape) + + if isinstance(emb_out, dict): + for key, val in emb_out.items(): + if key in output: + # 重复出现的key必须在(y, crossattn, concat)中,否则raise error + assert key in self.KEY2CATDIM + output[key] = torch.cat([output[key], val], + dim=self.KEY2CATDIM[key]) + else: + output[key] = val + else: + assert isinstance( + emb_out, (torch.Tensor, list, tuple) + ), f'encoder outputs must be tensors or a sequence, but got {type(emb_out)}' + + if not isinstance(emb_out, (list, tuple)): + emb_out = [emb_out] + + for emb in emb_out: + # 根据emb的维度,判断该cond归属于 (y, concat, crossattn)中的哪一种 + out_key = self.OUTPUT_DIM2KEYS[emb.dim()] + + if embedder.ucg_rate > 0.0 and embedder.legacy_ucg_val is None: + emb = (expand_dims_like( + torch.bernoulli( + (1.0 - embedder.ucg_rate) * + torch.ones(emb.shape[0], device=emb.device)), + emb, + ) * emb) + + if (hasattr(embedder, 'input_keys')): + if np.sum( + np.array([ + key in force_zero_embeddings + for key in embedder.input_keys + ])) > 0: + emb = torch.zeros_like(emb) + + if out_key in output: + output[out_key] = torch.cat((output[out_key], emb), + self.KEY2CATDIM[out_key]) + else: + output[out_key] = emb + return output def get_unconditional_conditioning(self, diff --git a/scepter/modules/model/embedder/resampler.py b/scepter/modules/model/embedder/resampler.py new file mode 100644 index 0000000..b392d1e --- /dev/null +++ b/scepter/modules/model/embedder/resampler.py @@ -0,0 +1,160 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import math + +import torch +import torch.nn as nn +from einops import rearrange +from einops.layers.torch import Rearrange + + +# FFN +def FeedForward(dim, mult=4): + inner_dim = int(dim * mult) + return nn.Sequential( + nn.LayerNorm(dim), + nn.Linear(dim, inner_dim, bias=False), + nn.GELU(), + nn.Linear(inner_dim, dim, bias=False), + ) + + +def reshape_tensor(x, heads): + bs, length, width = x.shape + # (bs, length, width) --> (bs, length, n_heads, dim_per_head) + x = x.view(bs, length, heads, -1) + # (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head) + x = x.transpose(1, 2) + # (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head) + x = x.reshape(bs, heads, length, -1) + return x + + +class PerceiverAttention(nn.Module): + def __init__(self, *, dim, dim_head=64, heads=8): + super().__init__() + self.scale = dim_head**-0.5 + self.dim_head = dim_head + self.heads = heads + inner_dim = dim_head * heads + + self.norm1 = nn.LayerNorm(dim) + self.norm2 = nn.LayerNorm(dim) + + self.to_q = nn.Linear(dim, inner_dim, bias=False) + self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False) + self.to_out = nn.Linear(inner_dim, dim, bias=False) + + def forward(self, x, latents): + """ + Args: + x (torch.Tensor): image features + shape (b, n1, D) + latent (torch.Tensor): latent features + shape (b, n2, D) + """ + x = self.norm1(x) + latents = self.norm2(latents) + + b, l, _ = latents.shape + + q = self.to_q(latents) + kv_input = torch.cat((x, latents), dim=-2) + k, v = self.to_kv(kv_input).chunk(2, dim=-1) + + q = reshape_tensor(q, self.heads) + k = reshape_tensor(k, self.heads) + v = reshape_tensor(v, self.heads) + + # attention + scale = 1 / math.sqrt(math.sqrt(self.dim_head)) + weight = (q * scale) @ (k * scale).transpose( + -2, -1) # More stable with f16 than dividing afterwards + weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype) + out = weight @ v + + out = out.permute(0, 2, 1, 3).reshape(b, l, -1) + + return self.to_out(out) + + +class Resampler(nn.Module): + def __init__( + self, + dim=1024, + depth=8, + dim_head=64, + heads=16, + num_queries=8, + embedding_dim=768, + output_dim=1024, + ff_mult=4, + max_seq_len: int = 257, # CLIP tokens + CLS token + apply_pos_emb: bool = False, + num_latents_mean_pooled: + int = 0, # number of latents derived from mean pooled representation of the sequence + ): + super().__init__() + self.pos_emb = nn.Embedding(max_seq_len, + embedding_dim) if apply_pos_emb else None + + self.latents = nn.Parameter( + torch.randn(1, num_queries, dim) / dim**0.5) + + self.proj_in = nn.Linear(embedding_dim, dim) + + self.proj_out = nn.Linear(dim, output_dim) + self.norm_out = nn.LayerNorm(output_dim) + + self.to_latents_from_mean_pooled_seq = (nn.Sequential( + nn.LayerNorm(dim), + nn.Linear(dim, dim * num_latents_mean_pooled), + Rearrange('b (n d) -> b n d', n=num_latents_mean_pooled), + ) if num_latents_mean_pooled > 0 else None) + + self.layers = nn.ModuleList([]) + for _ in range(depth): + self.layers.append( + nn.ModuleList([ + PerceiverAttention(dim=dim, dim_head=dim_head, + heads=heads), + FeedForward(dim=dim, mult=ff_mult), + ])) + + def forward(self, x): + if self.pos_emb is not None: + n, device = x.shape[1], x.device + pos_emb = self.pos_emb(torch.arange(n, device=device)) + x = x + pos_emb + + latents = self.latents.repeat(x.size(0), 1, 1) + + x = self.proj_in(x) + + if self.to_latents_from_mean_pooled_seq: + meanpooled_seq = masked_mean(x, + dim=1, + mask=torch.ones(x.shape[:2], + device=x.device, + dtype=torch.bool)) + meanpooled_latents = self.to_latents_from_mean_pooled_seq( + meanpooled_seq) + latents = torch.cat((meanpooled_latents, latents), dim=-2) + + for attn, ff in self.layers: + latents = attn(x, latents) + latents + latents = ff(latents) + latents + + latents = self.proj_out(latents) + return self.norm_out(latents) + + +def masked_mean(t, *, dim, mask=None): + if mask is None: + return t.mean(dim=dim) + + denom = mask.sum(dim=dim, keepdim=True) + mask = rearrange(mask, 'b n -> b n 1') + masked_t = t.masked_fill(~mask, 0.0) + + return masked_t.sum(dim=dim) / denom.clamp(min=1e-5) diff --git a/scepter/modules/model/head/__init__.py b/scepter/modules/model/head/__init__.py index 9f70332..cfd2b15 100644 --- a/scepter/modules/model/head/__init__.py +++ b/scepter/modules/model/head/__init__.py @@ -1,5 +1,8 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.model.head.classifier_head import ( - ClassifierHead, CosineLinearHead, TransformerHead, TransformerHeadx2, - VideoClassifierHead, VideoClassifierHeadx2) +from scepter.modules.model.head.classifier_head import (ClassifierHead, + CosineLinearHead, + TransformerHead, + TransformerHeadx2, + VideoClassifierHead, + VideoClassifierHeadx2) diff --git a/scepter/modules/model/metric/__init__.py b/scepter/modules/model/metric/__init__.py index 155ba54..5e83c7b 100644 --- a/scepter/modules/model/metric/__init__.py +++ b/scepter/modules/model/metric/__init__.py @@ -1,5 +1,4 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.model.metric.classification import (AccuracyMetric, - EnsembleAccuracyMetric - ) +from scepter.modules.model.metric.classification import ( + AccuracyMetric, EnsembleAccuracyMetric) diff --git a/scepter/modules/model/network/diffusion/diffusion.py b/scepter/modules/model/network/diffusion/diffusion.py index 4f8080a..1eb7d12 100644 --- a/scepter/modules/model/network/diffusion/diffusion.py +++ b/scepter/modules/model/network/diffusion/diffusion.py @@ -13,8 +13,7 @@ from .schedules import karras_schedule from .solvers import (sample_ddim, sample_dpm_2, sample_dpm_2_ancestral, sample_dpmpp_2m, sample_dpmpp_2m_sde, sample_dpmpp_2s_ancestral, sample_dpmpp_sde, - sample_euler, sample_euler_ancestral, sample_heun, - sample_img2img_euler, sample_img2img_euler_ancestral) + sample_euler, sample_euler_ancestral, sample_heun) __all__ = ['GaussianDiffusion'] @@ -27,6 +26,148 @@ def _i(tensor, t, x): return tensor[t.to(tensor.device)].view(shape).to(x.device) +def _unpack_2d_ks(kernel_size): + if isinstance(kernel_size, int): + ky = kx = kernel_size + else: + assert len( + kernel_size) == 2, '2D Kernel size should have a length of 2.' + ky, kx = kernel_size + + ky = int(ky) + kx = int(kx) + return ky, kx + + +def _compute_zero_padding(kernel_size): + ky, kx = _unpack_2d_ks(kernel_size) + return (ky - 1) // 2, (kx - 1) // 2 + + +def _bilateral_blur( + input, + guidance, + kernel_size, + sigma_color, + sigma_space, + border_type='reflect', + color_distance_type='l1', +): + + if isinstance(sigma_color, torch.Tensor): + sigma_color = sigma_color.to(device=input.device, + dtype=input.dtype).view(-1, 1, 1, 1, 1) + + ky, kx = _unpack_2d_ks(kernel_size) + pad_y, pad_x = _compute_zero_padding(kernel_size) + + padded_input = torch.nn.functional.pad(input, (pad_x, pad_x, pad_y, pad_y), + mode=border_type) + unfolded_input = padded_input.unfold(2, ky, 1).unfold(3, kx, 1).flatten( + -2) # (B, C, H, W, Ky x Kx) + + if guidance is None: + guidance = input + unfolded_guidance = unfolded_input + else: + padded_guidance = torch.nn.functional.pad(guidance, + (pad_x, pad_x, pad_y, pad_y), + mode=border_type) + unfolded_guidance = padded_guidance.unfold(2, ky, 1).unfold( + 3, kx, 1).flatten(-2) # (B, C, H, W, Ky x Kx) + + diff = unfolded_guidance - guidance.unsqueeze(-1) + if color_distance_type == 'l1': + color_distance_sq = diff.abs().sum(1, keepdim=True).square() + elif color_distance_type == 'l2': + color_distance_sq = diff.square().sum(1, keepdim=True) + else: + raise ValueError('color_distance_type only acceps l1 or l2') + color_kernel = (-0.5 / sigma_color**2 * + color_distance_sq).exp() # (B, 1, H, W, Ky x Kx) + + space_kernel = get_gaussian_kernel2d(kernel_size, + sigma_space, + device=input.device, + dtype=input.dtype) + space_kernel = space_kernel.view(-1, 1, 1, 1, kx * ky) + + kernel = space_kernel * color_kernel + out = (unfolded_input * kernel).sum(-1) / kernel.sum(-1) + return out + + +def get_gaussian_kernel1d( + kernel_size, + sigma, + force_even, + *, + device=None, + dtype=None, +): + + return gaussian(kernel_size, sigma, device=device, dtype=dtype) + + +def gaussian(window_size, sigma, *, device=None, dtype=None): + + batch_size = sigma.shape[0] + + x = (torch.arange(window_size, device=sigma.device, dtype=sigma.dtype) - + window_size // 2).expand(batch_size, -1) + + if window_size % 2 == 0: + x = x + 0.5 + + gauss = torch.exp(-x.pow(2.0) / (2 * sigma.pow(2.0))) + + return gauss / gauss.sum(-1, keepdim=True) + + +def get_gaussian_kernel2d( + kernel_size, + sigma, + force_even=False, + *, + device=None, + dtype=None, +): + + sigma = torch.Tensor([[sigma, sigma]]).to(device=device, dtype=dtype) + + ksize_y, ksize_x = _unpack_2d_ks(kernel_size) + sigma_y, sigma_x = sigma[:, 0, None], sigma[:, 1, None] + + kernel_y = get_gaussian_kernel1d(ksize_y, + sigma_y, + force_even, + device=device, + dtype=dtype)[..., None] + kernel_x = get_gaussian_kernel1d(ksize_x, + sigma_x, + force_even, + device=device, + dtype=dtype)[..., None] + + return kernel_y * kernel_x.view(-1, 1, ksize_x) + + +def adaptive_anisotropic_filter(x, g=None): + if g is None: + g = x + s, m = torch.std_mean(g, dim=(1, 2, 3), keepdim=True) + s = s + 1e-5 + guidance = (g - m) / s + y = _bilateral_blur(x, + guidance, + kernel_size=(13, 13), + sigma_color=3.0, + sigma_space=3.0, + border_type='reflect', + color_distance_type='l1') + return y + + class GaussianDiffusion(object): def __init__(self, sigmas, prediction_type='eps'): assert prediction_type in {'x0', 'eps', 'v'} @@ -53,6 +194,7 @@ class GaussianDiffusion(object): guide_scale=None, guide_rescale=None, clamp=None, + sharpness=0.0, percentile=None, cat_uc=False, **kwargs): @@ -79,54 +221,99 @@ class GaussianDiffusion(object): # prediction if guide_scale is None: - assert isinstance(model_kwargs, dict) - out = model(xt, t=t, **model_kwargs, **kwargs) + if isinstance(model_kwargs, dict): + out = model(xt, t=t, **model_kwargs, **kwargs) + elif isinstance(model_kwargs, list) and len(model_kwargs) > 0: + out = model(xt, t=t, **model_kwargs[0], **kwargs) + else: + raise Exception('Error') else: # classifier-free guidance (arXiv:2207.12598) # model_kwargs[0]: conditional kwargs # model_kwargs[1]: non-conditional kwargs - assert isinstance(model_kwargs, list) and len(model_kwargs) == 2 - - if guide_scale == 1.: - out = model(xt, t=t, **model_kwargs[0], **kwargs) - else: - if cat_uc: - - def parse_model_kwargs(prev_value, value): - if isinstance(value, torch.Tensor): - prev_value = torch.cat([prev_value, value], dim=0) - elif isinstance(value, dict): - for k, v in value.items(): - prev_value[k] = parse_model_kwargs( - prev_value[k], v) - elif isinstance(value, list): - for idx, v in enumerate(value): - prev_value[idx] = parse_model_kwargs( - prev_value[idx], v) - return prev_value - - all_model_kwargs = copy.deepcopy(model_kwargs[0]) - for model_kwarg in model_kwargs[1:]: - for key, value in model_kwarg.items(): - all_model_kwargs[key] = parse_model_kwargs( - all_model_kwargs[key], value) - all_out = model(xt.repeat(2, 1, 1, 1), - t=t.repeat(2), - **all_model_kwargs, - **kwargs) - y_out, u_out = all_out.chunk(2) + assert isinstance(model_kwargs, list) and len(model_kwargs) >= 2 + if isinstance(guide_scale, float) or isinstance(guide_scale, int): + assert len(model_kwargs) == 2 + if guide_scale == 1.: + out = model(xt, t=t, **model_kwargs[0], **kwargs) else: - y_out = model(xt, t=t, **model_kwargs[0], **kwargs) - u_out = model(xt, t=t, **model_kwargs[1], **kwargs) - out = u_out + guide_scale * (y_out - u_out) + if cat_uc: - # rescale the output according to arXiv:2305.08891 - if guide_rescale is not None: - assert guide_rescale >= 0 and guide_rescale <= 1 - ratio = (y_out.flatten(1).std(dim=1) / - (out.flatten(1).std(dim=1) + - 1e-12)).view((-1, ) + (1, ) * (y_out.ndim - 1)) - out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0 + def parse_model_kwargs(prev_value, value): + if isinstance(value, torch.Tensor): + prev_value = torch.cat([prev_value, value], + dim=0) + elif isinstance(value, dict): + for k, v in value.items(): + prev_value[k] = parse_model_kwargs( + prev_value[k], v) + elif isinstance(value, list): + for idx, v in enumerate(value): + prev_value[idx] = parse_model_kwargs( + prev_value[idx], v) + return prev_value + + all_model_kwargs = copy.deepcopy(model_kwargs[0]) + for model_kwarg in model_kwargs[1:]: + for key, value in model_kwarg.items(): + all_model_kwargs[key] = parse_model_kwargs( + all_model_kwargs[key], value) + all_out = model(xt.repeat(2, 1, 1, 1), + t=t.repeat(2), + **all_model_kwargs, + **kwargs) + y_out, u_out = all_out.chunk(2) + else: + y_out = model(xt, t=t, **model_kwargs[0], **kwargs) + u_out = model(xt, t=t, **model_kwargs[1], **kwargs) + # todo sharpness + # sharpness sampling + if sharpness is not None and sharpness > 0: + positive_x0 = alphas * xt - sigmas * y_out + negative_x0 = alphas * xt - sigmas * u_out + + positive_eps = xt - positive_x0 + negative_eps = xt - negative_x0 + + global_diffusion_progress = ( + 1 - t / 999.0).detach().cpu().numpy().tolist()[0] + alpha = 0.001 * sharpness * global_diffusion_progress + positive_eps_degraded = adaptive_anisotropic_filter( + x=positive_eps, g=positive_x0) + positive_eps_degraded_weighted = positive_eps_degraded * alpha + positive_eps * ( + 1.0 - alpha) + + final_eps = negative_eps + guide_scale * ( + positive_eps_degraded_weighted - negative_eps) + final_x0 = xt - final_eps + out = (alphas * xt - final_x0) / sigmas + else: + out = u_out + guide_scale * (y_out - u_out) + elif isinstance(guide_scale, dict): + assert len(model_kwargs) == 3 + y_out = model(xt, t=t, **model_kwargs[0], **kwargs) + m_out = model(xt, t=t, **model_kwargs[1], **kwargs) + u_out = model(xt, t=t, **model_kwargs[2], **kwargs) + out = u_out + guide_scale['image'] * ( + m_out - u_out) + guide_scale['text'] * (y_out - m_out) + elif isinstance(guide_scale, list): + assert len(guide_scale) == len(model_kwargs) - 1 + y_out = model(xt, t=t, **model_kwargs[0], **kwargs) + outs = [y_out] + for i in range(1, len(model_kwargs)): + outs.append(model(xt, t=t, **model_kwargs[i], **kwargs)) + out = outs[-1] + for i in range(len(guide_scale)): + out += guide_scale[i] * (outs[-i - 2] - outs[-i - 1]) + + # rescale the output according to arXiv:2305.08891 + if guide_rescale is not None and guide_rescale > 0.0: + assert guide_rescale >= 0 and guide_rescale <= 1 + ratio = ( + y_out.flatten(1).std(dim=1) / + (out.flatten(1).std(dim=1) + 1e-12)).view((-1, ) + (1, ) * + (y_out.ndim - 1)) + out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0 # compute x0 if self.prediction_type == 'x0': x0 = out @@ -197,6 +384,7 @@ class GaussianDiffusion(object): guide_scale=None, guide_rescale=None, clamp=None, + sharpness=0.0, percentile=None, solver='euler_a', steps=20, @@ -209,12 +397,16 @@ class GaussianDiffusion(object): seed=-1, intermediate_callback=None, cat_uc=False, + add_noise=False, + free_steps=None, + step_offset=None, **kwargs): # sanity check assert isinstance(steps, (int, torch.LongTensor)) assert t_max is None or (t_max > 0 and t_max <= self.num_timesteps - 1) assert t_min is None or (t_min >= 0 and t_min < self.num_timesteps - 1) - assert discretization in (None, 'leading', 'linspace', 'trailing') + assert discretization in (None, 'leading', 'linspace', 'trailing', + 'free') assert discard_penultimate_step in (None, True, False) assert return_intermediate in (None, 'x0', 'xt') @@ -255,17 +447,50 @@ class GaussianDiffusion(object): def model_fn(xt, sigma): # denoising t = self._sigma_to_t(sigma).repeat(len(xt)).round().long() - x0 = self.denoise(xt, - t, - None, - model, - model_kwargs, - guide_scale, - guide_rescale, - clamp, - percentile, - cat_uc=cat_uc, - **kwargs)[-2] + + if isinstance( + model_kwargs[0]['cond'], dict) and \ + 'tar_x0' in model_kwargs[0]['cond'] and \ + 'tar_mask_latent' in model_kwargs[0]['cond']: + tar_x0 = model_kwargs[0]['cond']['tar_x0'] + tar_mask = model_kwargs[0]['cond']['tar_mask_latent'] + + tar_xt = self.diffuse(x0=tar_x0, t=t) + xt = tar_xt * (1.0 - tar_mask) + xt * tar_mask + + if isinstance(model_kwargs[0]['cond'], + dict) and 'ref_x0' in model_kwargs[0]['cond']: + model_kwargs[0]['cond']['ref_xt'] = self.diffuse( + x0=model_kwargs[0]['cond']['ref_x0'], t=t) + model_kwargs[1]['cond']['ref_xt'] = self.diffuse( + x0=model_kwargs[1]['cond']['ref_x0'], t=t) + + if solver in ('onestep', 'multistep', 'multistep2', 'multistep3'): + x0 = self.denoise(xt, + t, + None, + model, + model_kwargs, + guide_scale, + guide_rescale, + clamp, + sharpness, + percentile, + cat_uc=cat_uc, + **kwargs)[-3] + else: + x0 = self.denoise(xt, + t, + None, + model, + model_kwargs, + guide_scale, + guide_rescale, + clamp, + sharpness, + percentile, + cat_uc=cat_uc, + **kwargs)[-2] # collect intermediate outputs if return_intermediate == 'xt': @@ -291,10 +516,14 @@ class GaussianDiffusion(object): elif discretization == 'trailing': steps = torch.arange(t_max, t_min - 1, -((t_max - t_min + 1) / steps)) + elif discretization == 'free': + steps = torch.tensor(free_steps) else: raise NotImplementedError( f'{discretization} discretization not implemented') steps = steps.clamp_(t_min, t_max) + elif isinstance(steps, list): + steps = torch.tensor(steps) steps = torch.as_tensor(steps, dtype=torch.float32, device=noise.device) @@ -335,6 +564,23 @@ class GaussianDiffusion(object): if discard_penultimate_step: sigmas = torch.cat([sigmas[:-2], sigmas[-1:]]) kwargs['seed'] = seed + # add noise to x0 + if add_noise: + if 'dm_steps' in kwargs: + if step_offset: + add_noise_step = -kwargs['dm_steps'] + step_offset + if add_noise_step < 0: + noise = self.diffuse( + noise, + torch.full((noise.shape[0], 1), + steps[add_noise_step], + dtype=torch.int)) + else: + noise = self.diffuse( + noise, + torch.full((noise.shape[0], 1), + steps[-kwargs['dm_steps'] - 1], + dtype=torch.int)) # sampling x0 = solver_fn(noise, model_fn, @@ -373,168 +619,6 @@ class GaussianDiffusion(object): | torch.isinf(log_sigma)] = float('inf') return log_sigma.exp() - @torch.no_grad() - def stochastic_encode(self, x0, t, steps): - # fast, but does not allow for exact reconstruction - # t serves as an index to gather the correct alphas - - t_max = None - t_min = None - - # discretization method - discretization = 'trailing' if self.prediction_type == 'v' else 'leading' - - # timesteps - if isinstance(steps, int): - t_max = self.num_timesteps - 1 if t_max is None else t_max - t_min = 0 if t_min is None else t_min - steps = discretize_timesteps(t_max, t_min, steps, discretization) - steps = torch.as_tensor(steps).round().long().flip(0).to(x0.device) - # steps = torch.as_tensor(steps).round().long().to(x0.device) - - # self.alphas_bar = torch.cumprod(1 - self.sigmas ** 2, dim=0) - # print('sigma: ', self.sigmas, len(self.sigmas)) - # print('alpha_bar: ', self.alphas_bar, len(self.alphas_bar)) - # print('steps: ', steps, len(steps)) - # sqrt_alphas_cumprod = torch.sqrt(self.alphas_bar).to(x0.device)[steps] - # sqrt_one_minus_alphas_cumprod = torch.sqrt(1 - self.alphas_bar).to(x0.device)[steps] - - sqrt_alphas_cumprod = self.alphas.to(x0.device)[steps] - sqrt_one_minus_alphas_cumprod = self.sigmas.to(x0.device)[steps] - # print('sigma: ', self.sigmas, len(self.sigmas)) - # print('alpha: ', self.alphas, len(self.alphas)) - # print('steps: ', steps, len(steps)) - - noise = torch.randn_like(x0) - return ( - extract_into_tensor(sqrt_alphas_cumprod, t, x0.shape) * x0 + - extract_into_tensor(sqrt_one_minus_alphas_cumprod, t, x0.shape) * - noise) - - @torch.no_grad() - def sample_img2img(self, - x, - noise, - model, - denoising_strength=1, - model_kwargs={}, - condition_fn=None, - guide_scale=None, - guide_rescale=None, - clamp=None, - percentile=None, - solver='euler_a', - steps=20, - t_max=None, - t_min=None, - discretization=None, - discard_penultimate_step=None, - return_intermediate=None, - show_progress=False, - seed=-1, - **kwargs): - # sanity check - assert isinstance(steps, (int, torch.LongTensor)) - assert t_max is None or (t_max > 0 and t_max <= self.num_timesteps - 1) - assert t_min is None or (t_min >= 0 and t_min < self.num_timesteps - 1) - assert discretization in (None, 'leading', 'linspace', 'trailing') - assert discard_penultimate_step in (None, True, False) - assert return_intermediate in (None, 'x0', 'xt') - # function of diffusion solver - solver_fn = { - 'euler_ancestral': sample_img2img_euler_ancestral, - 'euler': sample_img2img_euler, - }[solver] - # options - schedule = 'karras' if 'karras' in solver else None - discretization = discretization or 'linspace' - seed = seed if seed >= 0 else random.randint(0, 2**31) - if isinstance(steps, torch.LongTensor): - discard_penultimate_step = False - if discard_penultimate_step is None: - discard_penultimate_step = True if solver in ( - 'dpm2', 'dpm2_ancestral', 'dpmpp_2m_sde', 'dpm2_karras', - 'dpm2_ancestral_karras', 'dpmpp_2m_sde_karras') else False - - # function for denoising xt to get x0 - intermediates = [] - - def get_scalings(sigma): - c_out = -sigma - c_in = 1 / (sigma**2 + 1.**2)**0.5 - return c_out, c_in - - def model_fn(xt, sigma): - # denoising - c_out, c_in = get_scalings(sigma) - t = self._sigma_to_t(sigma).repeat(len(xt)).round().long() - - x0 = self.denoise(xt * c_in, t, None, model, model_kwargs, - guide_scale, guide_rescale, clamp, percentile, - **kwargs)[-2] - # collect intermediate outputs - if return_intermediate == 'xt': - intermediates.append(xt) - elif return_intermediate == 'x0': - intermediates.append(x0) - return xt + x0 * c_out - - # get timesteps - if isinstance(steps, int): - steps += 1 if discard_penultimate_step else 0 - t_max = self.num_timesteps - 1 if t_max is None else t_max - t_min = 0 if t_min is None else t_min - # discretize timesteps - if discretization == 'leading': - steps = torch.arange(t_min, t_max + 1, - (t_max - t_min + 1) / steps).flip(0) - elif discretization == 'linspace': - steps = torch.linspace(t_max, t_min, steps) - elif discretization == 'trailing': - steps = torch.arange(t_max, t_min - 1, - -((t_max - t_min + 1) / steps)) - else: - raise NotImplementedError( - f'{discretization} discretization not implemented') - steps = steps.clamp_(t_min, t_max) - steps = torch.as_tensor(steps, dtype=torch.float32, device=x.device) - # get sigmas - sigmas = self._t_to_sigma(steps) - sigmas = torch.cat([sigmas, sigmas.new_zeros([1])]) - t_enc = int(min(denoising_strength, 0.999) * len(steps)) - sigmas = sigmas[len(steps) - t_enc - 1:] - noise = x + noise * sigmas[0] - - if schedule == 'karras': - if sigmas[0] == float('inf'): - sigmas = karras_schedule( - n=len(steps) - 1, - sigma_min=sigmas[sigmas > 0].min().item(), - sigma_max=sigmas[sigmas < float('inf')].max().item(), - rho=7.).to(sigmas) - sigmas = torch.cat([ - sigmas.new_tensor([float('inf')]), sigmas, - sigmas.new_zeros([1]) - ]) - else: - sigmas = karras_schedule( - n=len(steps), - sigma_min=sigmas[sigmas > 0].min().item(), - sigma_max=sigmas.max().item(), - rho=7.).to(sigmas) - sigmas = torch.cat([sigmas, sigmas.new_zeros([1])]) - if discard_penultimate_step: - sigmas = torch.cat([sigmas[:-2], sigmas[-1:]]) - - # sampling - x0 = solver_fn(noise, - model_fn, - sigmas, - seed=seed, - show_progress=show_progress, - **kwargs) - return (x0, intermediates) if return_intermediate is not None else x0 - def extract_into_tensor(a, t, x_shape): b, *_ = t.shape diff --git a/scepter/modules/model/utils/data_utils.py b/scepter/modules/model/utils/data_utils.py new file mode 100644 index 0000000..1377bcd --- /dev/null +++ b/scepter/modules/model/utils/data_utils.py @@ -0,0 +1,148 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import cv2 +import numpy as np +import torchvision.transforms.functional as TF + + +def get_bbox_from_mask(mask): + h, w = mask.shape[0], mask.shape[1] + if mask.sum() < 10: + return 0, h, 0, w + rows = np.any(mask, axis=1) + cols = np.any(mask, axis=0) + y1, y2 = np.where(rows)[0][[0, -1]] + x1, x2 = np.where(cols)[0][[0, -1]] + return (y1, y2, x1, x2) + + +def pad_to_square(image, pad_value=255, random=False): + H, W = image.shape[0], image.shape[1] + if H == W: + return image, 0, 0 + + padd = abs(H - W) + if random: + padd_1 = int(np.random.randint(0, padd)) + else: + padd_1 = int(padd / 2) + padd_2 = padd - padd_1 + + if H > W: + pad_param = ((0, 0), (padd_1, padd_2), (0, 0)) + else: + pad_param = ((padd_1, padd_2), (0, 0), (0, 0)) + + # print(pad_param, pad_value) + image = np.pad(image, pad_param, 'constant', constant_values=pad_value) + return image, padd_1, padd_2 + + +def box_in_box(small_box, big_box): + y1, y2, x1, x2 = small_box + y1_b, _, x1_b, _ = big_box + y1, y2, x1, x2 = y1 - y1_b, y2 - y1_b, x1 - x1_b, x2 - x1_b + return (y1, y2, x1, x2) + + +def box2squre(image, box): + H, W = image.shape[0], image.shape[1] + y1, y2, x1, x2 = box + cx = (x1 + x2) // 2 + cy = (y1 + y2) // 2 + h, w = y2 - y1, x2 - x1 + + if h >= w: + x1 = cx - h // 2 + x2 = x1 + h + else: + y1 = cy - w // 2 + y2 = y1 + w + x1 = max(0, x1) + x2 = min(W, x2) + y1 = max(0, y1) + y2 = min(H, y2) + return (y1, y2, x1, x2) + + +def expand_bbox(mask, + yyxx, + ratio=1.0, + min_crop=0, + expand_type='center', + to_square=False): + y1, y2, x1, x2 = yyxx + h = y2 - y1 + 1 + w = x2 - x1 + 1 + + H, W = mask.shape[0], mask.shape[1] + xc, yc = 0.5 * (x1 + x2), 0.5 * (y1 + y2) + + def expand(k): + if isinstance(ratio, tuple) or isinstance(ratio, list): + r = np.random.uniform(*ratio) + k = k * r + else: + k = ratio * k + return k + + new_h = expand(h) + new_w = expand(w) + new_h = max(new_h, min_crop) + new_w = max(new_w, min_crop) + + if to_square: + if new_w / new_h < 0.334: + new_w = new_w + 1.0 / 3.0 * new_h + elif new_h / new_w < 0.334: + new_h = new_h + 1.0 / 3.0 * new_w + + if expand_type == 'center': + x1 = max(0, int(xc - new_w * 0.5)) + x2 = min(W, int(xc + new_w * 0.5)) + y1 = max(0, int(yc - new_h * 0.5)) + y2 = min(H, int(yc + new_h * 0.5)) + else: + x1 = max(0, min(x1, + int(x2 - new_w * np.random.uniform(w / new_w, 1.0)))) + x2 = min(W, max(x2, x1 + new_w)) + y1 = max(0, min(y1, + int(y2 - new_h * np.random.uniform(h / new_h, 1.0)))) + y2 = min(H, max(y2, y1 + new_h)) + + return (int(y1), int(y2), int(x1), int(x2)) + + +def crop_back(pred, tar_image, extra_sizes, tar_box_yyxx_crop): + H1, W1, H2, W2, pad1, pad2 = extra_sizes + y1, y2, x1, x2 = tar_box_yyxx_crop + pred = TF.resize(pred, (H2, W2), antialias=True) + # if W1 == H1: + # tar_image[:, y1:y2, x1:x2] = pred + # return tar_image + # if W1 < W2: + # pad1 = int((W2 - W1) / 2) + # pad2 = W2 - W1 - pad1 + # pred = pred[:, :,pad1:-pad2] + # else: + # pad1 = int((H2 - H1) / 2) + # pad2 = H2 - H1 - pad1 + # pred = pred[:, pad1:-pad2, :] + if W1 < W2: + # pad width + assert H1 == H2 and (pad1 + W1) == (W2 - pad2) + pred = pred[:, :, pad1 + 2:(W2 - pad2 - 2)] + tar_image[:, y1:y2, x1 + 2:x2 - 2] = pred + elif H1 < H2: + # pad height + assert W1 == W2 and (pad1 + H1) == (H2 - pad2) + pred = pred[:, pad1 + 2:(H2 - pad2 - 2), :] + tar_image[:, y1 + 2:y2 - 2, x1:x2] = pred + else: + tar_image[:, y1:y2, x1:x2] = pred + return tar_image + + +def save_image(image, save_path): + image = cv2.cvtColor(image, cv2.COLOR_RGB2BGR) + cv2.imwrite(save_path, image) diff --git a/scepter/modules/opt/optimizers/__init__.py b/scepter/modules/opt/optimizers/__init__.py index 675bcbb..e51388a 100644 --- a/scepter/modules/opt/optimizers/__init__.py +++ b/scepter/modules/opt/optimizers/__init__.py @@ -1,7 +1,10 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -from scepter.modules.opt.optimizers.official_optimizers import ( - ASGD, LBFGS, SGD, Adadelta, Adagrad, Adam, Adamax, AdamW, RMSprop, Rprop, - SparseAdam) +from scepter.modules.opt.optimizers.official_optimizers import (ASGD, LBFGS, + SGD, Adadelta, + Adagrad, Adam, + Adamax, AdamW, + RMSprop, Rprop, + SparseAdam) from scepter.modules.opt.optimizers.registry import OPTIMIZERS diff --git a/scepter/modules/solver/diffusion_solver.py b/scepter/modules/solver/diffusion_solver.py index 3b5b3e9..4997a81 100644 --- a/scepter/modules/solver/diffusion_solver.py +++ b/scepter/modules/solver/diffusion_solver.py @@ -333,12 +333,12 @@ class LatentDiffusionSolver(BaseSolver): data_iter = iter(self.datas[self._mode].dataloader) self.print_memory_status() for step in range(self.max_steps): - if 'eval' in self._mode_set and (step % self.eval_interval == 0 - or step == self.max_steps - 1): + if 'eval' in self._mode_set and (self.eval_interval > 0 and + step % self.eval_interval == 0): self.run_eval() self.train_mode() - self.before_iter(self.hooks_dict[self._mode]) batch_data = next(data_iter) + self.before_iter(self.hooks_dict[self._mode]) if self.sample_args: batch_data.update(self.sample_args.get_lowercase_dict()) if 'meta' in batch_data: @@ -364,6 +364,9 @@ class LatentDiffusionSolver(BaseSolver): self.after_iter(self.hooks_dict[self._mode]) if we.debug: self.print_trainable_params_status(prefix='model.') + if 'eval' in self._mode_set and (self.eval_interval > 0 + and step == self.max_steps - 1): + self.run_eval() self.after_all_iter(self.hooks_dict[self._mode]) @torch.no_grad() diff --git a/scepter/modules/solver/hooks/__init__.py b/scepter/modules/solver/hooks/__init__.py index f2fca74..aa86e36 100644 --- a/scepter/modules/solver/hooks/__init__.py +++ b/scepter/modules/solver/hooks/__init__.py @@ -4,12 +4,14 @@ from scepter.modules.solver.hooks.backward import BackwardHook from scepter.modules.solver.hooks.checkpoint import CheckpointHook from scepter.modules.solver.hooks.data_probe import ProbeDataHook +from scepter.modules.solver.hooks.ema import ModelEmaHook from scepter.modules.solver.hooks.hook import Hook from scepter.modules.solver.hooks.log import LogHook, TensorboardLogHook from scepter.modules.solver.hooks.lr import LrHook from scepter.modules.solver.hooks.registry import HOOKS from scepter.modules.solver.hooks.safetensors import SafetensorsHook from scepter.modules.solver.hooks.sampler import DistSamplerHook + """ Normally, hooks have priorities, below we recommend priority that runs fine (low score MEANS high priority) BackwardHook: 0 @@ -47,5 +49,6 @@ after solve: __all__ = [ 'HOOKS', 'BackwardHook', 'CheckpointHook', 'Hook', 'LrHook', 'LogHook', - 'TensorboardLogHook', 'DistSamplerHook', 'ProbeDataHook', 'SafetensorsHook' + 'TensorboardLogHook', 'DistSamplerHook', 'ProbeDataHook', + 'SafetensorsHook', 'ModelEmaHook' ] diff --git a/scepter/modules/solver/hooks/checkpoint.py b/scepter/modules/solver/hooks/checkpoint.py index 45b91cc..91bcfe9 100644 --- a/scepter/modules/solver/hooks/checkpoint.py +++ b/scepter/modules/solver/hooks/checkpoint.py @@ -147,7 +147,9 @@ class CheckpointHook(Hook): if self.save_last and solver.total_iter == solver.max_steps - 1: with FS.get_fs_client(save_path) as client: - last_path = osp.join(solver.work_dir, 'checkpoint.pth') + last_path = osp.join( + solver.work_dir, + f'checkpoints/{self.save_name_prefix}-last') client.make_link(last_path, save_path) self.last_ckpt = save_path diff --git a/scepter/modules/solver/hooks/data_probe.py b/scepter/modules/solver/hooks/data_probe.py index a5b9657..ea98326 100644 --- a/scepter/modules/solver/hooks/data_probe.py +++ b/scepter/modules/solver/hooks/data_probe.py @@ -30,7 +30,10 @@ class ProbeDataHook(Hook): def __init__(self, cfg, logger=None): super(ProbeDataHook, self).__init__(cfg, logger=logger) self.priority = cfg.get('PRIORITY', _DEFAULT_PROBE_PRIORITY) - self.log_interval = cfg.get('PROB_INTERVAL', 1000) + self.prob_interval = cfg.get('PROB_INTERVAL', 1000) + self.save_name_prefix = cfg.get('SAVE_NAME_PREFIX', 'step') + self.save_probe_prefix = cfg.get('SAVE_PROBE_PREFIX', None) + self.save_last = cfg.get('SAVE_LAST', False) def before_all_iter(self, solver): pass @@ -39,19 +42,23 @@ class ProbeDataHook(Hook): pass def after_iter(self, solver): - if solver.mode == 'train' and solver.total_iter % self.log_interval == 0: + if solver.mode == 'train' and solver.total_iter % self.prob_interval == 0: probe_dict = solver.probe_data if we.rank == 0: save_folder = os.path.join( solver.work_dir, - f'{solver.mode}_probe/step_{solver.total_iter}') + f'{solver.mode}_probe/{self.save_probe_prefix}-{solver.total_iter}' + ) ret_data = {} for k, v in probe_dict.items(): - ret_one = v.to_log( - os.path.join( + if self.save_probe_prefix is not None: + ret_prefix = os.path.join(save_folder, + self.save_probe_prefix) + else: + ret_prefix = os.path.join( save_folder, - k.replace('/', '_') + - f'_step_{solver.total_iter}')) + k.replace('/', '_') + f'_step_{solver.total_iter}') + ret_one = v.to_log(ret_prefix) if (isinstance(ret_one, list) or isinstance(ret_one, dict)) and len(ret_one) < 1: continue @@ -71,13 +78,19 @@ class ProbeDataHook(Hook): if we.rank == 0: step = solver._total_iter[ 'train'] if 'train' in solver._total_iter else 0 - save_folder = os.path.join(solver.work_dir, - f'{solver.mode}_probe/step_{step}') + save_folder = os.path.join( + solver.work_dir, + f'{solver.mode}_probe/{self.save_name_prefix}-{step}') ret_data = {} for k, v in probe_dict.items(): - ret_one = v.to_log( - os.path.join(save_folder, - k.replace('/', '_') + f'_step_{step}')) + if self.save_probe_prefix is not None: + ret_prefix = os.path.join(save_folder, + self.save_probe_prefix) + else: + ret_prefix = os.path.join( + save_folder, + k.replace('/', '_') + f'_step_{step}') + ret_one = v.to_log(ret_prefix) if (isinstance(ret_one, list) or isinstance(ret_one, dict)) and len(ret_one) < 1: continue @@ -87,6 +100,16 @@ class ProbeDataHook(Hook): json.dump(ret_data, open(local_path, 'w'), ensure_ascii=False) + + if self.save_last and step == solver.max_steps: + with FS.get_fs_client(save_folder) as client: + last_save_folder = os.path.join( + solver.work_dir, + f'{solver.mode}_probe/{self.save_name_prefix}-last' + ) + print(last_save_folder, save_folder) + client.make_link(last_save_folder, save_folder) + solver.clear_probe() torch.cuda.synchronize() barrier() diff --git a/scepter/modules/solver/hooks/ema.py b/scepter/modules/solver/hooks/ema.py new file mode 100644 index 0000000..0d3b042 --- /dev/null +++ b/scepter/modules/solver/hooks/ema.py @@ -0,0 +1,60 @@ +# -*- coding: utf-8 -*- +import torch +from torch.distributed.fsdp import (FullStateDictConfig, + FullyShardedDataParallel, StateDictType) + +from scepter.modules.solver.hooks.hook import Hook +from scepter.modules.solver.hooks.registry import HOOKS +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we + + +@HOOKS.register_class() +class ModelEmaHook(Hook): + para_dict = [{ + 'PRIORITY': { + 'value': 100, + 'description': 'the priority for processing!' + }, + 'BETA': { + 'value': 0.9999, + 'description': '' + }, + }] + + def __init__(self, cfg, logger=None): + super(ModelEmaHook, self).__init__(cfg, logger=logger) + self.priority = cfg.get('PRIORITY', 100) + self.beta = cfg.get('BETA', 0.9999) + + def after_iter(self, solver): + if solver.model.use_ema: + model_ema = solver.model.model_ema + model = solver.model.model + self.ema(model_ema, model, use_fsdp=solver.use_fsdp) + + @torch.no_grad() + def ema(self, net_ema, net, use_fsdp=True): + if we.is_distributed: + if use_fsdp: + save_policy = FullStateDictConfig(offload_to_cpu=False, + rank0_only=False) + with FullyShardedDataParallel.state_dict_type( + net, StateDictType.FULL_STATE_DICT, save_policy): + nonema_state = net.state_dict() + elif hasattr(net, 'module'): + nonema_state = net.module.state_dict() + else: + nonema_state = net.state_dict() + else: + nonema_state = net.state_dict() + + for k, v in net_ema.named_parameters(): + v.copy_(nonema_state[k].lerp(v, self.beta)) + + @staticmethod + def get_config_template(): + return dict_to_yaml('hook', + __class__.__name__, + ModelEmaHook.para_dict, + set_name=True) diff --git a/scepter/modules/transform/image.py b/scepter/modules/transform/image.py index fea07d0..a0a7c5a 100644 --- a/scepter/modules/transform/image.py +++ b/scepter/modules/transform/image.py @@ -10,11 +10,16 @@ import torchvision.transforms as transforms import torchvision.transforms.functional as TF from scepter.modules.transform.registry import TRANSFORMS -from scepter.modules.transform.utils import ( - BACKEND_CV2, BACKEND_PILLOW, BACKEND_TORCHVISION, INPUT_CV2_TYPE_WARNING, - INPUT_PIL_TYPE_WARNING, INPUT_TENSOR_TYPE_WARNING, INTERPOLATION_STYLE, - INTERPOLATION_STYLE_CV2, TORCHVISION_CAPABILITY, is_cv2_image, - is_pil_image, is_tensor) +from scepter.modules.transform.utils import (BACKEND_CV2, BACKEND_PILLOW, + BACKEND_TORCHVISION, + INPUT_CV2_TYPE_WARNING, + INPUT_PIL_TYPE_WARNING, + INPUT_TENSOR_TYPE_WARNING, + INTERPOLATION_STYLE, + INTERPOLATION_STYLE_CV2, + TORCHVISION_CAPABILITY, + is_cv2_image, is_pil_image, + is_tensor) from scepter.modules.utils.config import dict_to_yaml if TORCHVISION_CAPABILITY: diff --git a/scepter/modules/transform/io.py b/scepter/modules/transform/io.py index b0b4151..b49cc02 100644 --- a/scepter/modules/transform/io.py +++ b/scepter/modules/transform/io.py @@ -6,13 +6,15 @@ import os import cv2 import numpy as np import torch -from PIL import Image +from PIL import Image, ImageFile from scepter.modules.transform.registry import TRANSFORMS from scepter.modules.utils.config import dict_to_yaml from scepter.modules.utils.distribute import we from scepter.modules.utils.file_system import DATA_FS as FS +ImageFile.LOAD_TRUNCATED_IMAGES = True + def pillow_convert(image, rgb_order): if image.mode != rgb_order: diff --git a/scepter/modules/utils/config.py b/scepter/modules/utils/config.py index d7bceb1..da19f47 100644 --- a/scepter/modules/utils/config.py +++ b/scepter/modules/utils/config.py @@ -603,3 +603,6 @@ class Config(object): return cfg_new else: return cfg + + def pop(self, name): + self.cfg_dict.pop(name) diff --git a/scepter/modules/utils/export_model.py b/scepter/modules/utils/export_model.py index e499ec7..679e3ef 100644 --- a/scepter/modules/utils/export_model.py +++ b/scepter/modules/utils/export_model.py @@ -4,10 +4,10 @@ import io from io import BytesIO import onnx +import onnxruntime import torch from torch.onnx import OperatorExportTypes -import onnxruntime from scepter.modules.utils.distribute import we type_map = { diff --git a/scepter/modules/utils/file_clients/local_fs.py b/scepter/modules/utils/file_clients/local_fs.py index eae15ab..8a97443 100644 --- a/scepter/modules/utils/file_clients/local_fs.py +++ b/scepter/modules/utils/file_clients/local_fs.py @@ -293,6 +293,9 @@ class LocalFs(BaseFs): return True def get_logging_handler(self, target_logging_path): + dirname = os.path.dirname(target_logging_path) + if not os.path.exists(dirname): + os.makedirs(dirname, exist_ok=True) return logging.FileHandler(target_logging_path) def put_dir_from_local_dir(self, @@ -303,11 +306,11 @@ class LocalFs(BaseFs): target_dir = self.reconstruct_path(target_dir) if local_dir == target_dir: return True - # cp -f local_dir/* target_dir/* - if not osp.exists(target_dir): - status = os.system(f'mkdir -p {target_dir}') - if status != 0: - return False + # # cp -f local_dir/* target_dir/* + # if not osp.exists(target_dir): + # status = os.system(f'mkdir -p {target_dir}') + # if status != 0: + # return False try: shutil.copytree(local_dir, target_dir, symlinks=True) except Exception: diff --git a/scepter/modules/utils/file_system.py b/scepter/modules/utils/file_system.py index 11cb420..042bfac 100644 --- a/scepter/modules/utils/file_system.py +++ b/scepter/modules/utils/file_system.py @@ -1,5 +1,6 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +import io import os import threading import time @@ -327,6 +328,10 @@ class FileSystem(object): target_path_list): if local_path is None or target_path is None: flg = False + elif isinstance(local_path, io.BytesIO): + flg = FS.put_object(local_path.getvalue(), target_path) + elif isinstance(local_path, bytes): + flg = FS.put_object(local_path, target_path) elif self.exists(local_path): local_cache = self.get_from(local_path, local_path + f'{time.time()}', diff --git a/scepter/modules/utils/index.py b/scepter/modules/utils/index.py new file mode 100644 index 0000000..9967a11 --- /dev/null +++ b/scepter/modules/utils/index.py @@ -0,0 +1,44 @@ +# -*- coding: utf-8 -*- +import os.path as osp + +from tqdm import tqdm + +from scepter.modules.utils.file_system import FS + + +def init_1level_llfs(list_file, + max_lines=1024, + index_name='index', + delimiter='\n'): + r"""Construct large-list-file index. + """ + index_dir = osp.splitext(list_file)[0] + print(list_file) + with FS.get_from(list_file, wait_finish=True) as local_path: + print(local_path) + _num_split, index_stack, save_files = 0, [], [] + with open(local_path, 'r', buffering=1000000) as f: + for line in tqdm(f): + index_stack.append(line.strip()) + if len(index_stack) >= max_lines: + save_file = f'{index_dir}/{index_name}/{_num_split + 1:09d}.txt' + with FS.put_to(save_file) as cache_path: + with open(cache_path, 'w') as f_w: + f_w.write('\n'.join(index_stack)) + index_stack = [] + _num_split += 1 + save_files.append(save_file) + + if len(index_stack) > 0: + save_file = f'{index_dir}/{index_name}/{_num_split + 1:06d}.txt' + with FS.put_to(save_file) as cache_path: + with open(cache_path, 'w') as f_w: + f_w.write('\n'.join(index_stack)) + save_files.append(save_file) + + # output meta-file + index_file = osp.join(index_dir, f'{index_name}.txt') + with FS.put_to(index_file) as cache_path: + with open(cache_path, 'w') as f_w: + f_w.write('\n'.join(save_files)) + return index_file diff --git a/scepter/modules/utils/probe.py b/scepter/modules/utils/probe.py index c4577a6..2dc7dcc 100644 --- a/scepter/modules/utils/probe.py +++ b/scepter/modules/utils/probe.py @@ -263,7 +263,8 @@ class ProbeData(): url = FS.get_url(one_path, lifecycle=3600 * 365 * 24).replace( '.oss-internal.aliyun-inc.', - '.oss.aliyuncs.') + '.oss.aliyuncs.').replace( + '-internal', '') one_rank += ( f'' 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"

{config.get('TITLE', 'scepter studio')}

" ) + setattr(tab_manager, 'user_name', + gr.Text(value='', visible=False, show_label=False)) with gr.Tabs(elem_id='tabs') as tabs: setattr(tab_manager, 'tabs', tabs) for interface, label, ifid in interfaces: @@ -116,11 +144,29 @@ if __name__ == '__main__': interface.create_ui() for interface, label, ifid in interfaces: interface.set_callbacks(tab_manager) + auth_info = {} + if config.have('AUTH_INFO'): + for auth_user in config.AUTH_INFO: + auth_info[auth_user.USER] = auth_user.PASSWD + + def check_auth(user_name, password): + if user_name in auth_info: + return auth_info[user_name] == password + else: + return False + + def init_value(req: gr.Request): + print(req.username, 'have login') + return gr.Text(value=req.username, visible=False) + + if len(auth_info) > 0: + demo.load(init_value, outputs=[tab_manager.user_name]) demo.queue(status_update_rate=1).launch( server_name=args.host if args.host else config['HOST'], - server_port=args.port if args.port else config['PORT'], + server_port=int(args.port) if args.port else config['PORT'], root_path=config['ROOT'], show_error=True, debug=True, - enable_queue=True) + enable_queue=True, + auth=check_auth if len(auth_info) > 0 else None) diff --git a/scepter/version.py b/scepter/version.py index 610cb1d..ac673ee 100644 --- a/scepter/version.py +++ b/scepter/version.py @@ -1,7 +1,7 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -__version__ = '0.0.3.post1' +__version__ = '0.0.4' version_info = tuple(int(x) for x in __version__.split('.')[0:3]) diff --git a/tests/tools/test_train.py b/tests/tools/test_train.py index 4d390f0..9aea821 100644 --- a/tests/tools/test_train.py +++ b/tests/tools/test_train.py @@ -9,6 +9,16 @@ class TrainTest(unittest.TestCase): def setUp(self): print(('Testing %s.%s' % (type(self).__name__, self._testMethodName))) self.tmp_dir = './cache/save_data' + if not os.path.exists(self.tmp_dir): + os.makedirs(self.tmp_dir) + self.data_dir = './cache/datasets' + if not os.path.exists(self.data_dir): + data_cmd = """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""" + os.system(data_cmd) def tearDown(self): super().tearDown()
' f'
{save_id}-{idx}|{one_label}