|
After Width: | Height: | Size: 121 KiB |
|
After Width: | Height: | Size: 16 KiB |
|
After Width: | Height: | Size: 19 KiB |
|
After Width: | Height: | Size: 120 KiB |
|
After Width: | Height: | Size: 117 KiB |
|
After Width: | Height: | Size: 121 KiB |
|
After Width: | Height: | Size: 22 KiB |
|
After Width: | Height: | Size: 39 KiB |
|
After Width: | Height: | Size: 17 KiB |
|
After Width: | Height: | Size: 118 KiB |
|
After Width: | Height: | Size: 15 KiB |
|
After Width: | Height: | Size: 20 KiB |
|
After Width: | Height: | Size: 49 KiB |
|
After Width: | Height: | Size: 45 KiB |
|
After Width: | Height: | Size: 34 KiB |
|
After Width: | Height: | Size: 129 KiB |
|
After Width: | Height: | Size: 101 KiB |
|
After Width: | Height: | Size: 100 KiB |
|
After Width: | Height: | Size: 104 KiB |
|
After Width: | Height: | Size: 103 KiB |
@@ -0,0 +1,2 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
@@ -0,0 +1,3 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from .classifier_dataset import ImageClassifyExampleDataset
|
||||
@@ -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"]
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -9,8 +9,8 @@
|
||||
</p>
|
||||
|
||||
## 📖 Table of Contents
|
||||
- [Introduction](#-introduction)
|
||||
- [News](#-news)
|
||||
- [Introduction](#-introduction)
|
||||
- [Installation](#%EF%B8%8F-installation)
|
||||
- [Getting Started](#-getting-started)
|
||||
- [SCEPTER Studio](#%EF%B8%8F-scepter-studio)
|
||||
@@ -18,6 +18,15 @@
|
||||
- [Features](#-features)
|
||||
- [Learn More](#-learn-more)
|
||||
- [License](#license)
|
||||
- [Acknowledgement](#acknowledgement)
|
||||
|
||||
## 🎉 News
|
||||
- [2024.03]: We optimize the training UI and checkpoint management. New [LAR-Gen](https://arxiv.org/abs/2403.19534) model has been added on SCEPTER Studio, supporting `zoom-out`, `virtual try on`, `inpainting`.
|
||||
- [2024.02]: We release new SCEdit controllable image synthesis models for SD v2.1 and SD XL. Multiple strategies applied to accelerate inference time for SCEPTER Studio.
|
||||
- [2024.01]: We release **SCEPTER Studio**, an integrated toolkit for data management, model training and inference based on [Gradio](https://www.gradio.app/).
|
||||
- [2024.01]: [SCEdit](https://arxiv.org/abs/2312.11392) support controllable image synthesis for training and inference.
|
||||
- [2023.12]: We propose [SCEdit](https://arxiv.org/abs/2312.11392), an efficient and controllable generation framework.
|
||||
- [2023.12]: We release [🪄SCEPTER](https://github.com/modelscope/scepter/) library.
|
||||
|
||||
## 📝 Introduction
|
||||
|
||||
@@ -28,7 +37,7 @@ Main Feature:
|
||||
- Task:
|
||||
- Text-to-image generation
|
||||
- Controllable image synthesis
|
||||
- Image editing (TODO)
|
||||
- Image editing
|
||||
- Training / Inference:
|
||||
- Distribute: DDP / FSDP / FairScale / Xformers
|
||||
- File system: Local / Http / OSS / Modelscope
|
||||
@@ -40,15 +49,9 @@ Main Feature:
|
||||
Currently supported approaches (and counting):
|
||||
|
||||
1. SD Series: [Stable Diffusion v1.5](https://huggingface.co/runwayml/stable-diffusion-v1-5) / [Stable Diffusion v2.1](https://huggingface.co/runwayml/stable-diffusion-v1-5) / [Stable Diffusion XL](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0)
|
||||
2. SCEdit: [SCEdit: Efficient and Controllable Image Diffusion Generation via Skip Connection Editing](https://arxiv.org/abs/2312.11392) [](https://arxiv.org/abs/2312.11392) [](https://scedit.github.io/)
|
||||
3. Res-Tuning(TODO): [Res-Tuning: A Flexible and Efficient Tuning Paradigm via Unbinding Tuner from Backbone](https://arxiv.org/abs/2310.19859) [](https://arxiv.org/abs/2310.19859) [](https://res-tuning.github.io/)
|
||||
|
||||
## 🎉 News
|
||||
- [2024.02]: We release new SCEdit controllable image synthesis models for SD v2.1 and SD XL. Multiple strategies applied to accelerate inference time for SCEPTER Studio.
|
||||
- [2024.01]: We release **SCEPTER Studio**, an integrated toolkit for data management, model training and inference based on [Gradio](https://www.gradio.app/).
|
||||
- [2024.01]: [SCEdit](https://arxiv.org/abs/2312.11392) support controllable image synthesis for training and inference.
|
||||
- [2023.12]: We propose [SCEdit](https://arxiv.org/abs/2312.11392), an efficient and controllable generation framework.
|
||||
- [2023.12]: We release [🪄SCEPTER](https://github.com/modelscope/scepter/) library.
|
||||
2. SCEdit(CVPR2024): [SCEdit: Efficient and Controllable Image Diffusion Generation via Skip Connection Editing](https://arxiv.org/abs/2312.11392) [](https://arxiv.org/abs/2312.11392) [](https://scedit.github.io/)
|
||||
3. Res-Tuning(NeurIPS2023 TODO): [Res-Tuning: A Flexible and Efficient Tuning Paradigm via Unbinding Tuner from Backbone](https://arxiv.org/abs/2310.19859) [](https://arxiv.org/abs/2310.19859) [](https://res-tuning.github.io/)
|
||||
4. LAR-Gen: [Locate, Assign, Refine: Taming Customized Image Inpainting with Text-Subject Guidance](https://arxiv.org/abs/2403.19534) [](https://arxiv.org/abs/2403.19534) [](https://ali-vilab.github.io/largen-page/)
|
||||
|
||||
## 🛠️ Installation
|
||||
|
||||
@@ -94,7 +97,7 @@ For the data format used by SCEPTER Studio, please refer to [3D_example_csv.zip]
|
||||
To facilitate starting training in command-line mode, you can use a dataset in text format, please refer to [3D_example_txt.zip](https://modelscope.cn/api/v1/models/damo/scepter/repo?Revision=master&FilePath=datasets/3D_example_txt.zip)
|
||||
|
||||
```shell
|
||||
mkdir -p cache/dataset/ && wget 'https://modelscope.cn/api/v1/models/damo/scepter_scedit/repo?Revision=master&FilePath=dataset/3D_example_txt.zip' -O cache/dataset/3D_example_txt.zip && unzip cache/dataset/3D_example_txt.zip -d cache/dataset/ && rm cache/dataset/3D_example_txt.zip
|
||||
mkdir -p cache/datasets/ && wget 'https://modelscope.cn/api/v1/models/damo/scepter_scedit/repo?Revision=master&FilePath=dataset/3D_example_txt.zip' -O cache/datasets/3D_example_txt.zip && unzip cache/datasets/3D_example_txt.zip -d cache/datasets/ && rm cache/datasets/3D_example_txt.zip
|
||||
```
|
||||
|
||||
### Training
|
||||
@@ -165,6 +168,15 @@ 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 the modules of your task in `example/{task}`.
|
||||
```python
|
||||
cd example/classifier
|
||||
python run.py --cfg classifier.yaml
|
||||
```
|
||||
|
||||
|
||||
|
||||
## 🖥️ SCEPTER Studio
|
||||
|
||||
### Launch
|
||||
@@ -181,16 +193,94 @@ git clone https://github.com/modelscope/scepter.git
|
||||
PYTHONPATH=. python scepter/tools/webui.py --cfg scepter/methods/studio/scepter_ui.yaml
|
||||
```
|
||||
|
||||
The startup of **SCEPTER Studio** eliminates the need for manual downloading and organizing of models; it will automatically load the corresponding models and store them in a local directory.
|
||||
Depending on the network and hardware situation, the initial startup usually requires 15-60 minutes, primarily involving the download and processing of SDv1.5, SDv2.1, and SDXL models.
|
||||
The startup of **SCEPTER Studio** eliminates the need for manual downloading and organizing of models; it will automatically load the corresponding models and store them in a local directory.
|
||||
Depending on the network and hardware situation, the initial startup usually requires 15-60 minutes, primarily involving the download and processing of SDv1.5, SDv2.1, and SDXL models.
|
||||
Therefore, subsequent startups will become much faster (about one minute) as downloading is no longer required.
|
||||
|
||||
* LAR-Gen: we release `zoom-out`, `virtual try on`, `inpainting(text guided)`, `inpainting(text + reference image guided)` image editing capabilities.
|
||||
Please note that the **Data Preprocess** button must be clicked before clicking the **Generate** button.
|
||||
<p align="center">
|
||||
<img src="https://raw.githubusercontent.com/ali-vilab/largen-page/main/public/images/largen.gif">
|
||||
</p>
|
||||
|
||||
### Modelscope Studio
|
||||
|
||||
We deploy a work studio on Modelscope that includes only the inference tab, please refer to [ms_scepter_studio](https://www.modelscope.cn/studios/damo/scepter_studio/summary)
|
||||
|
||||
## 🖼️ Gallery
|
||||
|
||||
### LAR-Gen: Zoom Out
|
||||
<table>
|
||||
<tr>
|
||||
<td><strong>Origin Image</strong><br>Prompt: a temple on fire</td>
|
||||
<td><strong>Zoom-Out</strong><br>CenterAround:0.75</td>
|
||||
<td><strong>Zoom-Out</strong><br>CenterAround:0.75</td>
|
||||
<td><strong>Zoom-Out</strong><br>CenterAround:0.75</td>
|
||||
<td><strong>Zoom-Out</strong><br>CenterAround:0.75</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><img src="asset/images/zoom_out/ex1_scene_im.jpg" width="240"></td>
|
||||
<td><img src="asset/images/zoom_out/ex1_zoom_out1.jpg" width="240"></td>
|
||||
<td><img src="./asset/images/zoom_out/ex1_zoom_out2.jpg" width="240"></td>
|
||||
<td><img src="./asset/images/zoom_out/ex1_zoom_out3.jpg" width="240"></td>
|
||||
<td><img src="./asset/images/zoom_out/ex1_zoom_out4.jpg" width="240"></td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
### LAR-Gen: Virtual Try-on
|
||||
<table>
|
||||
<tr>
|
||||
<td><strong>Model Image</strong></td>
|
||||
<td><strong>Model Mask</strong></td>
|
||||
<td><strong>Clothing Image</strong></td>
|
||||
<td><strong>Clothing Mask</strong></td>
|
||||
<td><strong>Try-on Output</strong></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><img src="asset/images/virtual_try_on/model.jpg" width="240"></td>
|
||||
<td><img src="asset/images/virtual_try_on/ex2_scene_mask.jpg" width="240"></td>
|
||||
<td><img src="asset/images/virtual_try_on/tshirt.jpg" width="240"></td>
|
||||
<td><img src="asset/images/virtual_try_on/ex2_subject_mask.jpg" width="240"></td>
|
||||
<td><img src="asset/images/virtual_try_on/try_on_out.jpg" width="240"></td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
### LAR-Gen: Inpainting (Text guided)
|
||||
<table>
|
||||
<tr>
|
||||
<td><strong>Origin Image</strong><br>Prompt: a blue and white porcelain</td>
|
||||
<td><strong>Inpainting Mask1</strong></td>
|
||||
<td><strong>Inpainting Output1</strong></td>
|
||||
<td><strong>Inpainting Mask2</strong><br>Prompt: a clock</td>
|
||||
<td><strong>Inpainting Output2</strong></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><img src="asset/images/inpainting_text/ex3_scene_im.jpg" width="240"></td>
|
||||
<td><img src="asset/images/inpainting_text/ex3_scene_mask.jpg" width="240"></td>
|
||||
<td><img src="asset/images/inpainting_text/inpainting_text.jpg" width="240"></td>
|
||||
<td><img src="asset/images/inpainting_text/ex3_scene_mask2.jpg" width="240"></td>
|
||||
<td><img src="asset/images/inpainting_text/inpainting_text2.jpg" width="240"></td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
### LAR-Gen: Inpainting (Text and Subject guided)
|
||||
<table>
|
||||
<tr>
|
||||
<td><strong>Origin Image</strong><br>Prompt: a dog wearing sunglasses</td>
|
||||
<td><strong>Origin Mask</strong></td>
|
||||
<td><strong>Reference Image</strong></td>
|
||||
<td><strong>Reference Mask</strong></td>
|
||||
<td><strong>Inpainting Output</strong></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><img src="asset/images/inpainting_text_ref/ex4_scene_im.jpg" width="240"></td>
|
||||
<td><img src="asset/images/inpainting_text_ref/ex4_scene_mask.jpg" width="240"></td>
|
||||
<td><img src="asset/images/inpainting_text_ref/ex4_subject_im.jpg" width="240"></td>
|
||||
<td><img src="asset/images/inpainting_text_ref/ex4_subject_mask.jpg" width="240"></td>
|
||||
<td><img src="asset/images/inpainting_text_ref/inpainting_text_ref.jpg" width="240"></td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
### Dragon Year Special: Dragon Tuner
|
||||
|
||||
<table>
|
||||
@@ -244,22 +334,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 +370,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.
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
albumentations
|
||||
bezier
|
||||
einops
|
||||
modelscope
|
||||
ms-swift>=1.5.2
|
||||
@@ -6,6 +8,7 @@ open_clip_torch
|
||||
opencv-python
|
||||
opencv_transforms>=0.0.6
|
||||
oss2>=2.15.0
|
||||
pycocotools
|
||||
pyyaml>=5.3.1
|
||||
scikit-image
|
||||
torchsde
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
git+https://github.com/cocodataset/panopticapi.git
|
||||
torch==2.0.1
|
||||
torchvision==0.15.2
|
||||
xformers==0.0.21
|
||||
|
||||
@@ -1,2 +1,3 @@
|
||||
gradio>=3.47.1,<4.0.0
|
||||
imagehash
|
||||
psutil
|
||||
|
||||
@@ -79,6 +79,7 @@ EXTENSION_PARAS:
|
||||
MANTRA_BOOK: scepter/methods/studio/extensions/mantra_book/mantra_book.yaml
|
||||
OFFICIAL_TUNERS: scepter/methods/studio/extensions/tuners/official_tuners.yaml
|
||||
OFFICIAL_CONTROLLERS: scepter/methods/studio/extensions/controllers/official_controllers.yaml
|
||||
TUNER_MANAGER: scepter/methods/studio/tuner_manager/tuner_manager.yaml
|
||||
CONTROLABLE_ANNOTATORS:
|
||||
-
|
||||
NAME: "CannyAnnotator"
|
||||
|
||||
@@ -0,0 +1,268 @@
|
||||
NAME: LARGEN
|
||||
IS_DEFAULT: True
|
||||
DEFAULT_PARAS:
|
||||
PARAS:
|
||||
RESOLUTIONS: [[1024, 1024]]
|
||||
INPUT:
|
||||
IMAGE:
|
||||
ORIGINAL_SIZE_AS_TUPLE: [1024, 1024]
|
||||
TARGET_SIZE_AS_TUPLE: [1024, 1024]
|
||||
AESTHETIC_SCORE: 6.0
|
||||
NEGATIVE_AESTHETIC_SCORE: 2.5
|
||||
PROMPT: ""
|
||||
NEGATIVE_PROMPT: ""
|
||||
PROMPT_PREFIX: ""
|
||||
CROP_COORDS_TOP_LEFT: [0, 0]
|
||||
SAMPLE: ddim
|
||||
SAMPLE_STEPS: 50
|
||||
GUIDE_SCALE: 7.5
|
||||
GUIDE_RESCALE: 0.5
|
||||
DISCRETIZATION: trailing
|
||||
REFINE_SAMPLE: ddim
|
||||
REFINE_GUIDE_SCALE: 7.5
|
||||
REFINE_GUIDE_RESCALE: 0.5
|
||||
REFINE_DISCRETIZATION: trailing
|
||||
OUTPUT:
|
||||
LATENT:
|
||||
BEFORE_REFINE_IMAGES:
|
||||
IMAGES:
|
||||
SEED:
|
||||
MODULES_PARAS:
|
||||
FIRST_STAGE_MODEL:
|
||||
FUNCTION:
|
||||
-
|
||||
NAME: encode
|
||||
DTYPE: float32
|
||||
INPUT: ["IMAGE"]
|
||||
-
|
||||
NAME: decode
|
||||
DTYPE: float32
|
||||
INPUT: ["LATENT"]
|
||||
PARAS:
|
||||
# SCALE_FACTOR DESCRIPTION: The vae embeding scale. TYPE: float default: 0.18215
|
||||
SCALE_FACTOR: 0.13025
|
||||
SIZE_FACTOR: 8
|
||||
DIFFUSION_MODEL:
|
||||
FUNCTION:
|
||||
-
|
||||
NAME: forward
|
||||
DTYPE: float16
|
||||
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE", "GUIDE_RESCALE", "DISCRETIZATION"]
|
||||
COND_STAGE_MODEL:
|
||||
FUNCTION:
|
||||
-
|
||||
NAME: encode
|
||||
DTYPE: float16
|
||||
INPUT: ["ORIGINAL_SIZE_AS_TUPLE", "CROP_COORDS_TOP_LEFT", "PROMPT", "NEGATIVE_PROMPT"]
|
||||
REFINER_MODEL:
|
||||
FUNCTION:
|
||||
-
|
||||
NAME: forward
|
||||
DTYPE: float16
|
||||
INPUT: ["SAMPLE_STEPS", "REFINE_SAMPLE", "REFINE_GUIDE_SCALE", "REFINE_GUIDE_RESCALE", "REFINE_DISCRETIZATION"]
|
||||
REFINER_COND_MODEL:
|
||||
FUNCTION:
|
||||
-
|
||||
NAME: encode
|
||||
DTYPE: float16
|
||||
INPUT: ["ORIGINAL_SIZE_AS_TUPLE", "AESTHETIC_SCORE", "NEGATIVE_AESTHETIC_SCORE", "CROP_COORDS_TOP_LEFT", "PROMPT", "NEGATIVE_PROMPT"]
|
||||
|
||||
MODEL:
|
||||
PRETRAINED_MODEL: ms://damo/LARGEN@models/largen_ckpt_s22k.pth
|
||||
# SCHEDULE_ARGS DESCRIPTION: TYPE: default: ''
|
||||
SCHEDULE:
|
||||
PARAMETERIZATION: "eps"
|
||||
TIMESTEPS: 1000
|
||||
ZERO_TERMINAL_SNR: False
|
||||
SCHEDULE_ARGS:
|
||||
# NAME DESCRIPTION: TYPE: default: ''
|
||||
NAME: "scaled_linear"
|
||||
BETA_MIN: 0.00085
|
||||
BETA_MAX: 0.0120
|
||||
# DIFFUSION_MODEL DESCRIPTION: TYPE: default: ''
|
||||
DIFFUSION_MODEL:
|
||||
# NAME DESCRIPTION: TYPE: default: 'DiffusionUNetXL'
|
||||
NAME: LargenUNetXL
|
||||
# PRETRAINED_MODEL DESCRIPTION: Whole model's pretrained model path. TYPE: NoneType default: None
|
||||
PRETRAINED_MODEL:
|
||||
# IN_CHANNELS DESCRIPTION: Unet channels for input, considering the input image's channels. TYPE: int default: 4
|
||||
IN_CHANNELS: 9
|
||||
# OUT_CHANNELS DESCRIPTION: Unet channels for output, considering the input image's channels. TYPE: int default: 4
|
||||
OUT_CHANNELS: 4
|
||||
# NUM_RES_BLOCKS DESCRIPTION: The blocks's number of res. TYPE: int default: 2
|
||||
NUM_RES_BLOCKS: 2
|
||||
# MODEL_CHANNELS DESCRIPTION: base channel count for the model. TYPE: int default: 320
|
||||
MODEL_CHANNELS: 320
|
||||
# ATTENTION_RESOLUTIONS DESCRIPTION: A collection of downsample rates at which attention will take place. May be a set, list, or tuple. For example, if this contains 4, then at 4x downsampling, attentio will be used. TYPE: list default: [4, 2]
|
||||
ATTENTION_RESOLUTIONS: [4, 2]
|
||||
# DROPOUT DESCRIPTION: The dropout rate. TYPE: int default: 0
|
||||
DROPOUT: 0
|
||||
# CHANNEL_MULT DESCRIPTION: channel multiplier for each level of the UNet. TYPE: list default: [1, 2, 4]
|
||||
CHANNEL_MULT: [1, 2, 4]
|
||||
# CONV_RESAMPLE DESCRIPTION: Use conv to resample when downsample. TYPE: bool default: True
|
||||
CONV_RESAMPLE: True
|
||||
# DIMS DESCRIPTION: The Conv dims which 2 represent Conv2D. TYPE: int default: 2
|
||||
DIMS: 2
|
||||
# NUM_CLASSES DESCRIPTION: The class num for class guided setting, also can be set as continuous. TYPE: str default: 'sequential'
|
||||
NUM_CLASSES: sequential
|
||||
# USE_CHECKPOINT DESCRIPTION: Use gradient checkpointing to reduce memory usage. TYPE: bool default: False
|
||||
USE_CHECKPOINT: False
|
||||
# NUM_HEADS DESCRIPTION: The number of attention heads in each attention layer. TYPE: int default: -1
|
||||
NUM_HEADS: -1
|
||||
# NUM_HEADS_CHANNELS DESCRIPTION: If specified, ignore num_heads and instead use a fixed channel width per attention head. TYPE: int default: 64
|
||||
NUM_HEADS_CHANNELS: 64
|
||||
# USE_SCALE_SHIFT_NORM DESCRIPTION: The scale and shift for the outnorm of RESBLOCK, use a FiLM-like conditioning mechanism. TYPE: bool default: False
|
||||
USE_SCALE_SHIFT_NORM: False
|
||||
# RESBLOCK_UPDOWN DESCRIPTION: Use residual blocks for up/downsampling, if False use Conv. TYPE: bool default: False
|
||||
RESBLOCK_UPDOWN: False
|
||||
# USE_NEW_ATTENTION_ORDER DESCRIPTION: Whether use new attention(qkv before split heads or not) or not. TYPE: bool default: True
|
||||
USE_NEW_ATTENTION_ORDER: True
|
||||
# USE_SPATIAL_TRANSFORMER DESCRIPTION: Custom transformer which support the context, if context_dim is not None, the parameter must set True TYPE: bool default: True
|
||||
USE_SPATIAL_TRANSFORMER: True
|
||||
# TRANSFORMER_DEPTH DESCRIPTION: Custom transformer's depth, valid when USE_SPATIAL_TRANSFORMER is True. TYPE: list default: [1, 2, 10]
|
||||
TRANSFORMER_DEPTH: [1, 2, 10]
|
||||
# TRANSFORMER_DEPTH_MIDDLE DESCRIPTION: Custom transformer's depth of middle block, If set None, use TRANSFORMER_DEPTH last value. TYPE: NoneType default: None
|
||||
# TRANSFORMER_DEPTH_MIDDLE: None
|
||||
# CONTEXT_DIM DESCRIPTION: Custom context info, if set, USE_SPATIAL_TRANSFORMER also set True. TYPE: int default: 2048
|
||||
CONTEXT_DIM: 2048
|
||||
# DISABLE_SELF_ATTENTIONS DESCRIPTION: Whether disable the self-attentions on some level, should be a list, [False, True, ...] TYPE: NoneType default: None
|
||||
# DISABLE_SELF_ATTENTIONS: None
|
||||
# NUM_ATTENTION_BLOCKS DESCRIPTION: The number of attention blocks for attention layer. TYPE: NoneType default: None
|
||||
# NUM_ATTENTION_BLOCKS: None
|
||||
# DISABLE_MIDDLE_SELF_ATTN DESCRIPTION: Whether disable the self-attentions in middle blocks. TYPE: bool default: False
|
||||
DISABLE_MIDDLE_SELF_ATTN: False
|
||||
# USE_LINEAR_IN_TRANSFORMER DESCRIPTION: Custom transformer's parameter, valid when USE_SPATIAL_TRANSFORMER is True. TYPE: bool default: True
|
||||
USE_LINEAR_IN_TRANSFORMER: True
|
||||
# ADM_IN_CHANNELS DESCRIPTION: Used when num_classes == 'sequential' or 'timestep'. TYPE: int default: 2816
|
||||
ADM_IN_CHANNELS: 2816
|
||||
# USE_SENTENCE_EMB DESCRIPTION: Used sentence emb or not, default False. TYPE: bool default: False
|
||||
USE_SENTENCE_EMB: False
|
||||
# USE_WORD_MAPPING DESCRIPTION: Used word mapping or not, default False. TYPE: bool default: False
|
||||
USE_WORD_MAPPING: False
|
||||
TRANSFORMER_BLOCK_TYPE: att_v2
|
||||
IMAGE_SCALE: 1.0
|
||||
USE_REFINE: False
|
||||
# FIRST_STAGE_MODEL DESCRIPTION: TYPE: default: ''
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKL
|
||||
EMBED_DIM: 4
|
||||
IGNORE_KEYS: [ ]
|
||||
BATCH_SIZE: 1
|
||||
#
|
||||
ENCODER:
|
||||
NAME: Encoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 4
|
||||
DOUBLE_Z: True
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
#
|
||||
DECODER:
|
||||
NAME: Decoder
|
||||
CH: 128
|
||||
OUT_CH: 3
|
||||
NUM_RES_BLOCKS: 2
|
||||
IN_CHANNELS: 3
|
||||
ATTN_RESOLUTIONS: [ ]
|
||||
CH_MULT: [ 1, 2, 4, 4 ]
|
||||
Z_CHANNELS: 4
|
||||
DROPOUT: 0.0
|
||||
RESAMP_WITH_CONV: True
|
||||
GIVE_PRE_END: False
|
||||
TANH_OUT: False
|
||||
# COND_STAGE_MODEL DESCRIPTION: TYPE: default: ''
|
||||
COND_STAGE_MODEL:
|
||||
# NAME DESCRIPTION: TYPE: default: 'GeneralConditioner'
|
||||
NAME: GeneralConditioner
|
||||
USE_GRAD: False
|
||||
# EMBEDDERS DESCRIPTION: TYPE: default: ''
|
||||
EMBEDDERS:
|
||||
-
|
||||
# NAME DESCRIPTION: TYPE: default: 'FrozenCLIPEmbedder'
|
||||
NAME: FrozenCLIPEmbedder
|
||||
# PRETRAINED_MODEL DESCRIPTION: TYPE: str default: ''
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/clip-vit-large-patch14
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/clip-vit-large-patch14
|
||||
# MAX_LENGTH DESCRIPTION: TYPE: int default: 77
|
||||
MAX_LENGTH: 77
|
||||
# FREEZE DESCRIPTION: TYPE: bool default: True
|
||||
FREEZE: True
|
||||
# LAYER DESCRIPTION: TYPE: str default: 'last'
|
||||
LAYER: hidden
|
||||
# LAYER_IDX DESCRIPTION: TYPE: NoneType default: None
|
||||
LAYER_IDX: 11
|
||||
# USE_FINAL_LAYER_NORM DESCRIPTION: TYPE: bool default: False
|
||||
USE_FINAL_LAYER_NORM: False
|
||||
UCG_RATE: 0.0
|
||||
INPUT_KEYS: ["prompt"]
|
||||
LEGACY_UCG_VALUE:
|
||||
-
|
||||
# NAME DESCRIPTION: TYPE: default: 'FrozenOpenCLIPEmbedder2'
|
||||
NAME: FrozenOpenCLIPEmbedder2
|
||||
# ARCH DESCRIPTION: TYPE: str default: 'ViT-H-14'
|
||||
ARCH: ViT-bigG-14
|
||||
# MAX_LENGTH DESCRIPTION: TYPE: int default: 77
|
||||
MAX_LENGTH: 77
|
||||
# FREEZE DESCRIPTION: TYPE: bool default: True
|
||||
FREEZE: True
|
||||
# ALWAYS_RETURN_POOLED DESCRIPTION: Whether always return pooled results or not ,default False. TYPE: bool default: False
|
||||
ALWAYS_RETURN_POOLED: True
|
||||
# LEGACY DESCRIPTION: Whether use legacy returnd feature or not ,default True. TYPE: bool default: True
|
||||
LEGACY: False
|
||||
# LAYER DESCRIPTION: TYPE: str default: 'last'
|
||||
LAYER: penultimate
|
||||
UCG_RATE: 0.0
|
||||
INPUT_KEYS: ["prompt"]
|
||||
LEGACY_UCG_VALUE:
|
||||
-
|
||||
# NAME DESCRIPTION: TYPE: default: 'ConcatTimestepEmbedderND'
|
||||
NAME: ConcatTimestepEmbedderND
|
||||
# OUT_DIM DESCRIPTION: Output dim TYPE: int default: 256
|
||||
OUT_DIM: 256
|
||||
UCG_RATE: 0.0
|
||||
INPUT_KEYS: ["original_size_as_tuple"]
|
||||
LEGACY_UCG_VALUE:
|
||||
-
|
||||
# NAME DESCRIPTION: TYPE: default: 'ConcatTimestepEmbedderND'
|
||||
NAME: ConcatTimestepEmbedderND
|
||||
# OUT_DIM DESCRIPTION: Output dim TYPE: int default: 256
|
||||
OUT_DIM: 256
|
||||
UCG_RATE: 0.0
|
||||
INPUT_KEYS: ["crop_coords_top_left"]
|
||||
LEGACY_UCG_VALUE:
|
||||
-
|
||||
# NAME DESCRIPTION: TYPE: default: 'ConcatTimestepEmbedderND'
|
||||
NAME: ConcatTimestepEmbedderND
|
||||
# OUT_DIM DESCRIPTION: Output dim TYPE: int default: 256
|
||||
OUT_DIM: 256
|
||||
UCG_RATE: 0.0
|
||||
INPUT_KEYS: ["target_size_as_tuple"]
|
||||
LEGACY_UCG_VALUE:
|
||||
-
|
||||
NAME: IPAdapterPlusEmbedder
|
||||
CLIP_DIR: ms://damo/LARGEN@models/clip_encoder/
|
||||
PRETRAINED_MODEL: ms://damo/LARGEN@models/ip-adapter-plus_sdxl_vit-h.bin
|
||||
INPUT_KEYS: [ "ref_ip", "ref_detail" ]
|
||||
IN_DIM: 1280
|
||||
HEADS: 20
|
||||
CROSSATTN_DIM: 2048
|
||||
-
|
||||
NAME: TransparentEmbedder
|
||||
INPUT_KEYS: [ "tar_x0", "tar_mask_latent" ]
|
||||
-
|
||||
NAME: NoiseConcatEmbedder
|
||||
INPUT_KEYS: [ "tar_mask_latent", "masked_x0" ]
|
||||
-
|
||||
NAME: TransparentEmbedder
|
||||
INPUT_KEYS: [ "ref_x0" ]
|
||||
-
|
||||
NAME: TransparentEmbedder
|
||||
INPUT_KEYS: [ "task" ]
|
||||
-
|
||||
NAME: TransparentEmbedder
|
||||
INPUT_KEYS: [ "image_scale" ]
|
||||
@@ -49,11 +49,11 @@ BANNER: |
|
||||
<div class="qr-codes">
|
||||
<div class="qr-code-container">
|
||||
<img src="https://modelscope.cn/api/v1/models/damo/scepter/repo?Revision=master&FilePath=assets/scepter_studio/ms_scepter_studio_qr.png" alt="ms_scepter_studio_qr">
|
||||
<div class="caption">Modelscope Studio</div>
|
||||
<div class="caption"><a href="https://www.modelscope.cn/studios/iic/scepter_studio">Modelscope Studio</a></div>
|
||||
</div>
|
||||
<div class="qr-code-container">
|
||||
<img src="https://modelscope.cn/api/v1/models/damo/scepter/repo?Revision=master&FilePath=assets/scepter_studio/scepter_github_qr.png" alt="scepter_github_qr">
|
||||
<div class="caption">Github</div>
|
||||
<div class="caption"><a href="https://github.com/modelscope/scepter">Github</a></div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
@@ -79,6 +79,10 @@ INTERFACE:
|
||||
NAME_EN: Train
|
||||
IFID: self_train
|
||||
CONFIG: scepter/methods/studio/self_train/self_train.yaml
|
||||
- NAME: 模型管理
|
||||
NAME_EN: Tuner Management
|
||||
IFID: tuner_manager
|
||||
CONFIG: scepter/methods/studio/tuner_manager/tuner_manager.yaml
|
||||
- NAME: 推理
|
||||
NAME_EN: Inference
|
||||
IFID: inference
|
||||
|
||||
@@ -11,13 +11,13 @@ META:
|
||||
DEFAULT_SAMPLER: "dpmpp_2s_ancestral"
|
||||
DEFAULT_SAMPLE_STEPS: 40
|
||||
INFERENCE_N_PROMPT: ""
|
||||
RESOLUTION: 1024
|
||||
RESOLUTION: [1024, 1024]
|
||||
PARAS:
|
||||
-
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 1024
|
||||
RESOLUTION: [1024, 1024]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -29,7 +29,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 1024
|
||||
RESOLUTION: [1024, 1024]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -41,7 +41,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 1024
|
||||
RESOLUTION: [1024, 1024]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 200
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -53,7 +53,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 1024
|
||||
RESOLUTION: [1024, 1024]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 200
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -65,7 +65,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 1024
|
||||
RESOLUTION: [1024, 1024]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -146,10 +146,12 @@ SOLVER:
|
||||
MAX_EPOCHS: -1
|
||||
# NUM_FOLDS DESCRIPTION: Num folds for training. TYPE: int default: 1
|
||||
NUM_FOLDS: 1
|
||||
#
|
||||
EVAL_INTERVAL: -1
|
||||
# WORK_DIR DESCRIPTION: Save dir of the training log or model. TYPE: str default: ''
|
||||
WORK_DIR:
|
||||
# LOG_FILE DESCRIPTION: Save log path. TYPE: str default: ''
|
||||
LOG_FILE: stg_log.txt
|
||||
LOG_FILE: std_log.txt
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
@@ -586,15 +588,44 @@ SOLVER:
|
||||
INPUT_KEY: [ 'img', 'img_original_size_as_tuple', 'img_target_size_as_tuple', 'img_crop_coords_top_left']
|
||||
OUTPUT_KEY: [ 'image', 'original_size_as_tuple', 'target_size_as_tuple', 'crop_coords_top_left' ]
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: Text2ImageDataset
|
||||
MODE: eval
|
||||
PROMPT_FILE:
|
||||
PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ]
|
||||
IMAGE_SIZE: [ 1024, 1024 ]
|
||||
FIELDS: [ "prompt" ]
|
||||
DELIMITER: '#;#'
|
||||
PROMPT_PREFIX: ''
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'index', 'prompt' ]
|
||||
META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
- NAME: BackwardHook
|
||||
# GRADIENT_CLIP: 1.0
|
||||
-
|
||||
NAME: BackwardHook
|
||||
PRIORITY: 0
|
||||
- NAME: LogHook
|
||||
LOG_INTERVAL: 50
|
||||
-
|
||||
NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
SHOW_GPU_MEM: True
|
||||
-
|
||||
NAME: TensorboardLogHook
|
||||
-
|
||||
NAME: CheckpointHook
|
||||
SAVE_LAST: True
|
||||
INTERVAL: 10000
|
||||
PRIORITY: 200
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
SAVE_PROBE_PREFIX: 'image'
|
||||
|
||||
@@ -8,3 +8,13 @@ SAMPLERS:
|
||||
NAME: 'dpmpp_2m_sde'
|
||||
-
|
||||
NAME: 'dpmpp_2s_ancestral'
|
||||
TRAIN_PARAS:
|
||||
RESOLUTIONS:
|
||||
VALUES: [[256, 256], [320, 180], [180, 320],
|
||||
[512, 512], [640, 360], [360, 640],
|
||||
[768, 768], [960, 540], [540, 960],
|
||||
[1024, 1024], [1280, 720], [720, 1280]]
|
||||
DEFAULT: [1024, 1024]
|
||||
EVAL_PROMPTS:
|
||||
- a boy wearing a jacket
|
||||
- a dog running on the lawn
|
||||
|
||||
@@ -10,13 +10,13 @@ META:
|
||||
DEFAULT_SAMPLER: "ddim"
|
||||
DEFAULT_SAMPLE_STEPS: 40
|
||||
INFERENCE_N_PROMPT: ""
|
||||
RESOLUTION: 512
|
||||
RESOLUTION: [512, 512]
|
||||
PARAS:
|
||||
-
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 512
|
||||
RESOLUTION: [512, 512]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -28,7 +28,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 512
|
||||
RESOLUTION: [512, 512]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -40,7 +40,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 512
|
||||
RESOLUTION: [512, 512]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 200
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -53,7 +53,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 512
|
||||
RESOLUTION: [512, 512]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 200
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -66,7 +66,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 4
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 512
|
||||
RESOLUTION: [512, 512]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -135,6 +135,7 @@ SOLVER:
|
||||
MAX_EPOCHS: -1
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: -1
|
||||
#
|
||||
WORK_DIR:
|
||||
LOG_FILE: std_log.txt
|
||||
@@ -298,15 +299,44 @@ SOLVER:
|
||||
KEYS: [ 'image', 'prompt' ]
|
||||
META_KEYS: [ 'data_key' ]
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: Text2ImageDataset
|
||||
MODE: eval
|
||||
PROMPT_FILE:
|
||||
PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ]
|
||||
IMAGE_SIZE: [ 512, 512 ]
|
||||
FIELDS: [ "prompt" ]
|
||||
DELIMITER: '#;#'
|
||||
PROMPT_PREFIX: ''
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'index', 'prompt' ]
|
||||
META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
- NAME: BackwardHook
|
||||
# GRADIENT_CLIP: 1.0
|
||||
-
|
||||
NAME: BackwardHook
|
||||
PRIORITY: 0
|
||||
- NAME: LogHook
|
||||
LOG_INTERVAL: 50
|
||||
-
|
||||
NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
SHOW_GPU_MEM: True
|
||||
-
|
||||
NAME: TensorboardLogHook
|
||||
-
|
||||
NAME: CheckpointHook
|
||||
SAVE_LAST: True
|
||||
INTERVAL: 10000
|
||||
PRIORITY: 200
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
SAVE_PROBE_PREFIX: 'image'
|
||||
|
||||
@@ -10,13 +10,13 @@ META:
|
||||
DEFAULT_SAMPLER: "ddim"
|
||||
DEFAULT_SAMPLE_STEPS: 40
|
||||
INFERENCE_N_PROMPT: ""
|
||||
RESOLUTION: 768
|
||||
RESOLUTION: [768, 768]
|
||||
PARAS:
|
||||
-
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 768
|
||||
RESOLUTION: [768, 768]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -28,7 +28,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 768
|
||||
RESOLUTION: [768, 768]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -40,7 +40,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 768
|
||||
RESOLUTION: [768, 768]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 200
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -78,6 +78,7 @@ SOLVER:
|
||||
MAX_EPOCHS: -1
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: -1
|
||||
#
|
||||
WORK_DIR:
|
||||
LOG_FILE: std_log.txt
|
||||
@@ -240,15 +241,44 @@ SOLVER:
|
||||
KEYS: [ 'image', 'prompt' ]
|
||||
META_KEYS: [ 'data_key' ]
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: Text2ImageDataset
|
||||
MODE: eval
|
||||
PROMPT_FILE:
|
||||
PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ]
|
||||
IMAGE_SIZE: [ 768, 768 ]
|
||||
FIELDS: [ "prompt" ]
|
||||
DELIMITER: '#;#'
|
||||
PROMPT_PREFIX: ''
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'index', 'prompt' ]
|
||||
META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
- NAME: BackwardHook
|
||||
# GRADIENT_CLIP: 1.0
|
||||
-
|
||||
NAME: BackwardHook
|
||||
PRIORITY: 0
|
||||
- NAME: LogHook
|
||||
LOG_INTERVAL: 50
|
||||
-
|
||||
NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
SHOW_GPU_MEM: True
|
||||
-
|
||||
NAME: TensorboardLogHook
|
||||
-
|
||||
NAME: CheckpointHook
|
||||
SAVE_LAST: True
|
||||
INTERVAL: 10000
|
||||
PRIORITY: 200
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
SAVE_PROBE_PREFIX: 'image'
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
WORK_DIR: "tuner_manager"
|
||||
TUNER_LIST_YAML: "tuner_list.yaml"
|
||||
@@ -244,12 +244,18 @@ class Text2ImageDataset(BaseDataset):
|
||||
image_size = [image_size, image_size]
|
||||
assert isinstance(image_size, Iterable) and len(image_size) == 2
|
||||
|
||||
prompt_file = cfg.PROMPT_FILE
|
||||
with FS.get_object(prompt_file) as local_data:
|
||||
if cfg.PROMPT_FILE is not None and cfg.PROMPT_FILE != '':
|
||||
prompt_file = cfg.PROMPT_FILE
|
||||
with FS.get_object(prompt_file) as local_data:
|
||||
rows = [
|
||||
i.split(delimiter,
|
||||
len(fields) - 1)
|
||||
for i in local_data.decode('utf-8').strip().split('\n')
|
||||
]
|
||||
else:
|
||||
rows = [
|
||||
i.split(delimiter,
|
||||
len(fields) - 1)
|
||||
for i in local_data.decode('utf-8').strip().split('\n')
|
||||
len(fields) - 1) for i in cfg.PROMPT_DATA
|
||||
]
|
||||
|
||||
self.items = list()
|
||||
@@ -263,10 +269,9 @@ class Text2ImageDataset(BaseDataset):
|
||||
item['meta']['img_path'] = os.path.join(path_prefix, value)
|
||||
elif key in ['width', 'height']:
|
||||
item['meta'][key] = int(value)
|
||||
elif key != 'meta':
|
||||
item[key] = value
|
||||
else:
|
||||
continue
|
||||
item['meta'][key] = value
|
||||
|
||||
self.items.append(item)
|
||||
if use_num > 0:
|
||||
self.items = self.items[:use_num]
|
||||
|
||||
@@ -2,18 +2,23 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import os
|
||||
import warnings
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torchvision.transforms as TT
|
||||
from PIL.Image import Image
|
||||
from swift import SwiftModel
|
||||
|
||||
from scepter.modules.model.registry import TUNERS
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
try:
|
||||
from swift import SwiftModel
|
||||
except Exception:
|
||||
warnings.warn('Import swift failed, please check it.')
|
||||
|
||||
|
||||
class ControlInference():
|
||||
def __init__(self, logger=None):
|
||||
|
||||
@@ -474,10 +474,14 @@ class DiffusionInference():
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
if self.tokenizer:
|
||||
if not hasattr(get_model(self.cond_stage_model), 'tokenizer'):
|
||||
setattr(get_model(self.cond_stage_model), 'tokenizer',
|
||||
self.tokenizer)
|
||||
context = getattr(get_model(self.cond_stage_model),
|
||||
function_name)(batch['tokens'])
|
||||
null_context = getattr(get_model(self.cond_stage_model),
|
||||
function_name)(batch_uc['tokens'])
|
||||
|
||||
else:
|
||||
context = getattr(get_model(self.cond_stage_model),
|
||||
function_name)(batch)
|
||||
@@ -558,12 +562,13 @@ class DiffusionInference():
|
||||
seed=seed,
|
||||
condition_fn=None,
|
||||
clamp=None,
|
||||
sharpness=value_input.get('sharpness', 0.0),
|
||||
percentile=None,
|
||||
t_max=None,
|
||||
t_min=None,
|
||||
discard_penultimate_step=None,
|
||||
intermediate_callback=intermediate_callback,
|
||||
cat_uc=cat_uc,
|
||||
cat_uc=value_input.get('cat_uc', cat_uc),
|
||||
**kwargs)
|
||||
|
||||
self.dynamic_unload(self.diffusion_model,
|
||||
|
||||
@@ -0,0 +1,595 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import os.path
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
import torchvision.transforms.functional as TF
|
||||
from PIL.Image import Image
|
||||
|
||||
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
|
||||
from scepter.modules.model.network.diffusion.schedules import noise_schedule
|
||||
from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, MODELS,
|
||||
TOKENIZERS)
|
||||
from scepter.modules.model.utils.data_utils import crop_back
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
|
||||
def get_model(model_tuple):
|
||||
assert 'model' in model_tuple
|
||||
return model_tuple['model']
|
||||
|
||||
|
||||
class LargenInference():
|
||||
'''
|
||||
define vae, unet, text-encoder, tuner, refiner components
|
||||
support to load the components dynamicly.
|
||||
create and load model when run this model at the first time.
|
||||
'''
|
||||
def __init__(self, logger=None):
|
||||
self.logger = logger
|
||||
self.loaded_model = {}
|
||||
self.loaded_model_name = [
|
||||
'diffusion_model', 'first_stage_model', 'cond_stage_model'
|
||||
]
|
||||
|
||||
def init_from_cfg(self, cfg):
|
||||
self.name = cfg.NAME
|
||||
self.is_default = cfg.get('IS_DEFAULT', False)
|
||||
module_paras = self.load_default(cfg.get('DEFAULT_PARAS', None))
|
||||
assert cfg.have('MODEL')
|
||||
cfg.MODEL = self.redefine_paras(cfg.MODEL)
|
||||
self.diffusion = self.load_schedule(cfg.MODEL.SCHEDULE)
|
||||
self.diffusion_model = self.infer_model(
|
||||
cfg.MODEL.DIFFUSION_MODEL, module_paras.get(
|
||||
'DIFFUSION_MODEL',
|
||||
None)) if cfg.MODEL.have('DIFFUSION_MODEL') else None
|
||||
self.first_stage_model = self.infer_model(
|
||||
cfg.MODEL.FIRST_STAGE_MODEL,
|
||||
module_paras.get(
|
||||
'FIRST_STAGE_MODEL',
|
||||
None)) if cfg.MODEL.have('FIRST_STAGE_MODEL') else None
|
||||
self.cond_stage_model = self.infer_model(
|
||||
cfg.MODEL.COND_STAGE_MODEL,
|
||||
module_paras.get(
|
||||
'COND_STAGE_MODEL',
|
||||
None)) if cfg.MODEL.have('COND_STAGE_MODEL') else None
|
||||
self.refiner_cond_model = self.infer_model(
|
||||
cfg.MODEL.REFINER_COND_MODEL,
|
||||
module_paras.get(
|
||||
'REFINER_COND_MODEL',
|
||||
None)) if cfg.MODEL.have('REFINER_COND_MODEL') else None
|
||||
self.refiner_diffusion_model = self.infer_model(
|
||||
cfg.MODEL.REFINER_MODEL, module_paras.get(
|
||||
'REFINER_MODEL',
|
||||
None)) if cfg.MODEL.have('REFINER_MODEL') else None
|
||||
self.tokenizer = TOKENIZERS.build(
|
||||
cfg.MODEL.TOKENIZER,
|
||||
logger=self.logger) if cfg.MODEL.have('TOKENIZER') else None
|
||||
|
||||
if self.tokenizer is not None:
|
||||
self.cond_stage_model['cfg'].KWARGS = {
|
||||
'vocab_size': self.tokenizer.vocab_size
|
||||
}
|
||||
|
||||
def redefine_paras(self, cfg):
|
||||
if cfg.get('PRETRAINED_MODEL', None):
|
||||
assert FS.isfile(cfg.PRETRAINED_MODEL)
|
||||
with FS.get_from(cfg.PRETRAINED_MODEL,
|
||||
wait_finish=True) as local_path:
|
||||
if local_path.endswith('safetensors'):
|
||||
from safetensors.torch import load_file as load_safetensors
|
||||
sd = load_safetensors(local_path)
|
||||
else:
|
||||
sd = torch.load(local_path, map_location='cpu')
|
||||
|
||||
if 'model' in sd:
|
||||
sd = sd['model']
|
||||
|
||||
first_stage_model_path = os.path.join(
|
||||
os.path.dirname(local_path), 'first_stage_model.pth')
|
||||
cond_stage_model_path = os.path.join(
|
||||
os.path.dirname(local_path), 'cond_stage_model.pth')
|
||||
diffusion_model_path = os.path.join(
|
||||
os.path.dirname(local_path), 'diffusion_model.pth')
|
||||
if (not os.path.exists(first_stage_model_path)
|
||||
or not os.path.exists(cond_stage_model_path)
|
||||
or not os.path.exists(diffusion_model_path)):
|
||||
self.logger.info(
|
||||
'Now read the whole model and rearrange the modules, it may take several mins.'
|
||||
)
|
||||
first_stage_model = OrderedDict()
|
||||
cond_stage_model = OrderedDict()
|
||||
diffusion_model = OrderedDict()
|
||||
for k, v in sd.items():
|
||||
if k.startswith('first_stage_model.'):
|
||||
first_stage_model[k.replace(
|
||||
'first_stage_model.', '')] = v
|
||||
elif k.startswith('conditioner.'):
|
||||
cond_stage_model[k.replace('conditioner.', '')] = v
|
||||
elif k.startswith('cond_stage_model.'):
|
||||
if k.startswith('cond_stage_model.model.'):
|
||||
cond_stage_model[k.replace(
|
||||
'cond_stage_model.model.', '')] = v
|
||||
else:
|
||||
cond_stage_model[k.replace(
|
||||
'cond_stage_model.', '')] = v
|
||||
elif k.startswith('model.diffusion_model.'):
|
||||
diffusion_model[k.replace('model.diffusion_model.',
|
||||
'')] = v
|
||||
elif k.startswith('model.'):
|
||||
diffusion_model[k.replace('model.', '')] = v
|
||||
else:
|
||||
continue
|
||||
if cfg.have('FIRST_STAGE_MODEL'):
|
||||
with open(first_stage_model_path + 'cache', 'wb') as f:
|
||||
torch.save(first_stage_model, f)
|
||||
os.rename(first_stage_model_path + 'cache',
|
||||
first_stage_model_path)
|
||||
self.logger.info(
|
||||
'First stage model has been processed.')
|
||||
if cfg.have('COND_STAGE_MODEL'):
|
||||
with open(cond_stage_model_path + 'cache', 'wb') as f:
|
||||
torch.save(cond_stage_model, f)
|
||||
os.rename(cond_stage_model_path + 'cache',
|
||||
cond_stage_model_path)
|
||||
self.logger.info(
|
||||
'Cond stage model has been processed.')
|
||||
if cfg.have('DIFFUSION_MODEL'):
|
||||
with open(diffusion_model_path + 'cache', 'wb') as f:
|
||||
torch.save(diffusion_model, f)
|
||||
os.rename(diffusion_model_path + 'cache',
|
||||
diffusion_model_path)
|
||||
self.logger.info('Diffusion model has been processed.')
|
||||
if not cfg.FIRST_STAGE_MODEL.get('PRETRAINED_MODEL', None):
|
||||
cfg.FIRST_STAGE_MODEL.PRETRAINED_MODEL = first_stage_model_path
|
||||
else:
|
||||
cfg.FIRST_STAGE_MODEL.RELOAD_MODEL = first_stage_model_path
|
||||
if not cfg.COND_STAGE_MODEL.get('PRETRAINED_MODEL', None):
|
||||
cfg.COND_STAGE_MODEL.PRETRAINED_MODEL = cond_stage_model_path
|
||||
else:
|
||||
cfg.COND_STAGE_MODEL.RELOAD_MODEL = cond_stage_model_path
|
||||
if not cfg.DIFFUSION_MODEL.get('PRETRAINED_MODEL', None):
|
||||
cfg.DIFFUSION_MODEL.PRETRAINED_MODEL = diffusion_model_path
|
||||
else:
|
||||
cfg.DIFFUSION_MODEL.RELOAD_MODEL = diffusion_model_path
|
||||
return cfg
|
||||
|
||||
def init_from_modules(self, modules):
|
||||
for k, v in modules.items():
|
||||
self.__setattr__(k, v)
|
||||
|
||||
def infer_model(self, cfg, module_paras=None):
|
||||
module = {
|
||||
'model': None,
|
||||
'cfg': cfg,
|
||||
'device': 'offline',
|
||||
'name': cfg.NAME,
|
||||
'function_info': {},
|
||||
'paras': {}
|
||||
}
|
||||
if module_paras is None:
|
||||
return module
|
||||
function_info = {}
|
||||
paras = {
|
||||
k.lower(): v
|
||||
for k, v in module_paras.get('PARAS', {}).items()
|
||||
}
|
||||
for function in module_paras.get('FUNCTION', []):
|
||||
input_dict = {}
|
||||
for inp in function.get('INPUT', []):
|
||||
if inp.lower() in self.input:
|
||||
input_dict[inp.lower()] = self.input[inp.lower()]
|
||||
function_info[function.NAME] = {
|
||||
'dtype': function.get('DTYPE', 'float32'),
|
||||
'input': input_dict
|
||||
}
|
||||
module['paras'] = paras
|
||||
module['function_info'] = function_info
|
||||
return module
|
||||
|
||||
def init_from_ckpt(self, path, model, ignore_keys=list()):
|
||||
if path.endswith('safetensors'):
|
||||
from safetensors.torch import load_file as load_safetensors
|
||||
sd = load_safetensors(path)
|
||||
else:
|
||||
sd = torch.load(path, map_location='cpu')
|
||||
|
||||
new_sd = OrderedDict()
|
||||
for k, v in sd.items():
|
||||
ignored = False
|
||||
for ik in ignore_keys:
|
||||
if ik in k:
|
||||
if we.rank == 0:
|
||||
self.logger.info(
|
||||
'Ignore key {} from state_dict.'.format(k))
|
||||
ignored = True
|
||||
break
|
||||
if not ignored:
|
||||
new_sd[k] = v
|
||||
|
||||
missing, unexpected = model.load_state_dict(new_sd, strict=False)
|
||||
if we.rank == 0:
|
||||
self.logger.info(
|
||||
f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys'
|
||||
)
|
||||
if len(missing) > 0:
|
||||
self.logger.info(f'Missing Keys:\n {missing}')
|
||||
if len(unexpected) > 0:
|
||||
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
|
||||
|
||||
def load(self, module):
|
||||
if module['device'] == 'offline':
|
||||
if module['cfg'].NAME in MODELS.class_map:
|
||||
model = MODELS.build(module['cfg'], logger=self.logger).eval()
|
||||
elif module['cfg'].NAME in BACKBONES.class_map:
|
||||
model = BACKBONES.build(module['cfg'],
|
||||
logger=self.logger).eval()
|
||||
elif module['cfg'].NAME in EMBEDDERS.class_map:
|
||||
model = EMBEDDERS.build(module['cfg'],
|
||||
logger=self.logger).eval()
|
||||
else:
|
||||
raise NotImplementedError
|
||||
if module['cfg'].get('RELOAD_MODEL', None):
|
||||
self.init_from_ckpt(module['cfg'].RELOAD_MODEL, model)
|
||||
module['model'] = model
|
||||
module['device'] = 'cpu'
|
||||
if module['device'] == 'cpu':
|
||||
module['device'] = we.device_id
|
||||
module['model'] = module['model'].to(we.device_id)
|
||||
return module
|
||||
|
||||
def unload(self, module):
|
||||
if module is None:
|
||||
return module
|
||||
module['model'] = module['model'].to('cpu')
|
||||
module['device'] = 'cpu'
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
return module
|
||||
|
||||
def dynamic_load(self, module=None, name=''):
|
||||
self.logger.info('Loading {} model'.format(name))
|
||||
if name == 'all':
|
||||
for subname in self.loaded_model_name:
|
||||
self.loaded_model[subname] = self.dynamic_load(
|
||||
getattr(self, subname), subname)
|
||||
elif name in self.loaded_model_name:
|
||||
if name in self.loaded_model:
|
||||
if module['cfg'] != self.loaded_model[name]['cfg']:
|
||||
self.unload(self.loaded_model[name])
|
||||
module = self.load(module)
|
||||
self.loaded_model[name] = module
|
||||
return module
|
||||
elif module['device'] == 'cpu':
|
||||
module = self.load(module)
|
||||
return module
|
||||
else:
|
||||
return module
|
||||
else:
|
||||
module = self.load(module)
|
||||
self.loaded_model[name] = module
|
||||
return module
|
||||
else:
|
||||
return self.load(module)
|
||||
|
||||
def dynamic_unload(self, module=None, name='', skip_loaded=False):
|
||||
self.logger.info('Unloading {} model'.format(name))
|
||||
if name == 'all':
|
||||
for name, module in self.loaded_model.items():
|
||||
module = self.unload(self.loaded_model[name])
|
||||
self.loaded_model[name] = module
|
||||
elif name in self.loaded_model_name:
|
||||
if name in self.loaded_model:
|
||||
if not skip_loaded:
|
||||
module = self.unload(self.loaded_model[name])
|
||||
self.loaded_model[name] = module
|
||||
else:
|
||||
self.unload(module)
|
||||
else:
|
||||
self.unload(module)
|
||||
|
||||
def load_default(self, cfg):
|
||||
module_paras = {}
|
||||
if cfg is not None:
|
||||
self.paras = cfg.PARAS
|
||||
self.input = {k.lower(): v for k, v in cfg.INPUT.items()}
|
||||
self.output = {k.lower(): v for k, v in cfg.OUTPUT.items()}
|
||||
module_paras = cfg.MODULES_PARAS
|
||||
return module_paras
|
||||
|
||||
def load_schedule(self, cfg):
|
||||
parameterization = cfg.get('PARAMETERIZATION', 'eps')
|
||||
assert parameterization in [
|
||||
'eps', 'x0', 'v'
|
||||
], 'currently only supporting "eps" and "x0" and "v"'
|
||||
num_timesteps = cfg.get('TIMESTEPS', 1000)
|
||||
|
||||
schedule_args = {
|
||||
k.lower(): v
|
||||
for k, v in cfg.get('SCHEDULE_ARGS', {
|
||||
'NAME': 'logsnr_cosine_interp',
|
||||
'SCALE_MIN': 2.0,
|
||||
'SCALE_MAX': 4.0
|
||||
}).items()
|
||||
}
|
||||
|
||||
zero_terminal_snr = cfg.get('ZERO_TERMINAL_SNR', False)
|
||||
if zero_terminal_snr:
|
||||
assert parameterization == 'v', 'Now zero_terminal_snr only support v-prediction mode.'
|
||||
sigmas = noise_schedule(schedule=schedule_args.pop('name'),
|
||||
n=num_timesteps,
|
||||
zero_terminal_snr=zero_terminal_snr,
|
||||
**schedule_args)
|
||||
diffusion = GaussianDiffusion(sigmas=sigmas,
|
||||
prediction_type=parameterization)
|
||||
return diffusion
|
||||
|
||||
def get_batch(self, value_dict, num_samples=1):
|
||||
batch = {}
|
||||
batch_uc = {}
|
||||
N = num_samples
|
||||
device = we.device_id
|
||||
for key in value_dict:
|
||||
if key == 'prompt':
|
||||
if not self.tokenizer:
|
||||
batch['prompt'] = value_dict['prompt']
|
||||
batch_uc['prompt'] = value_dict['negative_prompt']
|
||||
else:
|
||||
batch['tokens'] = self.tokenizer(value_dict['prompt']).to(
|
||||
we.device_id)
|
||||
batch_uc['tokens'] = self.tokenizer(
|
||||
value_dict['negative_prompt']).to(we.device_id)
|
||||
elif key == 'original_size_as_tuple':
|
||||
batch['original_size_as_tuple'] = (torch.tensor(
|
||||
value_dict['original_size_as_tuple']).to(device).repeat(
|
||||
N, 1))
|
||||
elif key == 'crop_coords_top_left':
|
||||
batch['crop_coords_top_left'] = (torch.tensor(
|
||||
value_dict['crop_coords_top_left']).to(device).repeat(
|
||||
N, 1))
|
||||
elif key == 'aesthetic_score':
|
||||
batch['aesthetic_score'] = (torch.tensor(
|
||||
[value_dict['aesthetic_score']]).to(device).repeat(N, 1))
|
||||
batch_uc['aesthetic_score'] = (torch.tensor([
|
||||
value_dict['negative_aesthetic_score']
|
||||
]).to(device).repeat(N, 1))
|
||||
|
||||
elif key == 'target_size_as_tuple':
|
||||
batch['target_size_as_tuple'] = (torch.tensor(
|
||||
value_dict['target_size_as_tuple']).to(device).repeat(
|
||||
N, 1))
|
||||
elif key == 'image':
|
||||
batch[key] = self.load_image(value_dict[key], num_samples=N)
|
||||
else:
|
||||
batch[key] = value_dict[key]
|
||||
|
||||
for key in batch.keys():
|
||||
if key not in batch_uc and isinstance(batch[key], torch.Tensor):
|
||||
batch_uc[key] = torch.clone(batch[key])
|
||||
return batch, batch_uc
|
||||
|
||||
def load_image(self, image, num_samples=1):
|
||||
if isinstance(image, torch.Tensor):
|
||||
pass
|
||||
elif isinstance(image, Image):
|
||||
pass
|
||||
elif isinstance(image, Image):
|
||||
pass
|
||||
|
||||
def get_function_info(self, module, function_name=None):
|
||||
all_function = module['function_info']
|
||||
if function_name in all_function:
|
||||
return function_name, all_function[function_name]['dtype']
|
||||
if function_name is None and len(all_function) == 1:
|
||||
for k, v in all_function.items():
|
||||
return k, v['dtype']
|
||||
|
||||
def encode_first_stage(self, x, **kwargs):
|
||||
_, dtype = self.get_function_info(self.first_stage_model, 'encode')
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
z = get_model(self.first_stage_model).encode(x)
|
||||
return self.first_stage_model['paras']['scale_factor'] * z
|
||||
|
||||
def decode_first_stage(self, z):
|
||||
_, dtype = self.get_function_info(self.first_stage_model, 'encode')
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
z = 1. / self.first_stage_model['paras']['scale_factor'] * z
|
||||
return get_model(self.first_stage_model).decode(z)
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(self,
|
||||
input,
|
||||
num_samples=1,
|
||||
intermediate_callback=None,
|
||||
refine_strength=0,
|
||||
img_to_img_strength=0,
|
||||
cat_uc=True,
|
||||
tuner_model=None,
|
||||
control_model=None,
|
||||
**kwargs):
|
||||
|
||||
value_input = copy.deepcopy(self.input)
|
||||
value_input.update(input)
|
||||
print(value_input)
|
||||
height, width = value_input['target_size_as_tuple']
|
||||
value_output = copy.deepcopy(self.output)
|
||||
batch, batch_uc = self.get_batch(value_input, num_samples=1)
|
||||
|
||||
# first stage encode
|
||||
task = kwargs.get('largen_task', 'Text_Guided_Inpainting')
|
||||
image_scale = kwargs.get('largen_image_scale', 1.0)
|
||||
tar_image = kwargs.get('largen_tar_image', None)
|
||||
tar_mask = kwargs.get('largen_tar_mask', None)
|
||||
masked_image = kwargs.get('largen_masked_image', None)
|
||||
ref_image = kwargs.get('largen_ref_image', None)
|
||||
ref_mask = kwargs.get('largen_ref_mask', None)
|
||||
ref_clip = kwargs.get('largen_ref_clip', None)
|
||||
|
||||
base_image = kwargs.get('largen_base_image', None)
|
||||
extra_sizes = kwargs.get('largen_extra_sizes', None)
|
||||
bbox_yyxx = kwargs.get('largen_bbox_yyxx', None)
|
||||
|
||||
device = we.device_id
|
||||
tar_image = tar_image.to(device)
|
||||
tar_mask = tar_mask.to(device)
|
||||
masked_image = masked_image.to(device)
|
||||
if 'Subject' in task:
|
||||
ref_image = ref_image.to(device)
|
||||
ref_mask = ref_mask.to(device)
|
||||
ref_clip = ref_clip.to(device)
|
||||
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
|
||||
tar_x0 = self.encode_first_stage(tar_image)
|
||||
masked_x0 = self.encode_first_stage(masked_image)
|
||||
b, _, h, w = tar_x0.shape
|
||||
tar_mask_latent = TF.resize(tar_mask, (h, w), antialias=True)
|
||||
tar_mask_latent = (tar_mask_latent > 0.5).float()
|
||||
|
||||
batch.update({
|
||||
'tar_x0': tar_x0,
|
||||
'tar_mask_latent': tar_mask_latent,
|
||||
'masked_x0': masked_x0,
|
||||
'task': task
|
||||
})
|
||||
batch_uc.update({
|
||||
'tar_x0': tar_x0,
|
||||
'tar_mask_latent': tar_mask_latent,
|
||||
'masked_x0': masked_x0,
|
||||
'task': task
|
||||
})
|
||||
|
||||
if 'Subject' in task and ref_image is not None:
|
||||
ref_x0 = self.encode_first_stage(ref_image)
|
||||
batch.update({
|
||||
'ref_ip': ref_clip,
|
||||
'ref_detail': ref_clip,
|
||||
'ref_x0': ref_x0,
|
||||
'ref_mask': ref_mask,
|
||||
'image_scale': image_scale,
|
||||
})
|
||||
batch_uc.update({
|
||||
'ref_ip': torch.zeros_like(ref_clip),
|
||||
'ref_detail': ref_clip,
|
||||
'ref_x0': ref_x0,
|
||||
'ref_mask': ref_mask,
|
||||
'image_scale': image_scale,
|
||||
})
|
||||
|
||||
self.dynamic_unload(self.first_stage_model,
|
||||
'first_stage_model',
|
||||
skip_loaded=True)
|
||||
|
||||
# cond stage
|
||||
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
|
||||
function_name, dtype = self.get_function_info(self.cond_stage_model)
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
context = getattr(get_model(self.cond_stage_model),
|
||||
function_name)(batch)
|
||||
null_context = getattr(get_model(self.cond_stage_model),
|
||||
function_name)(batch_uc)
|
||||
|
||||
self.dynamic_unload(self.cond_stage_model,
|
||||
'cond_stage_model',
|
||||
skip_loaded=True)
|
||||
|
||||
# get noise
|
||||
seed = kwargs.pop('seed', -1)
|
||||
g = torch.Generator(device=we.device_id)
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
|
||||
g.manual_seed(seed)
|
||||
if 'seed' in value_output:
|
||||
value_output['seed'] = seed
|
||||
for sample_id in range(num_samples):
|
||||
if self.diffusion_model is not None:
|
||||
noise = torch.empty(
|
||||
1,
|
||||
4,
|
||||
height // self.first_stage_model['paras']['size_factor'],
|
||||
width // self.first_stage_model['paras']['size_factor'],
|
||||
device=we.device_id).normal_(generator=g)
|
||||
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
# UNet use input n_prompt
|
||||
function_name, dtype = self.get_function_info(
|
||||
self.diffusion_model)
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
latent = self.diffusion.sample(
|
||||
noise=noise,
|
||||
x=None,
|
||||
denoising_strength=1.0,
|
||||
refine_strength=refine_strength,
|
||||
solver=value_input.get('sample', 'ddim'),
|
||||
model=get_model(self.diffusion_model),
|
||||
model_kwargs=[{
|
||||
'cond': context
|
||||
}, {
|
||||
'cond': null_context
|
||||
}],
|
||||
steps=value_input.get('sample_steps', 50),
|
||||
guide_scale=value_input.get('guide_scale', 7.5),
|
||||
guide_rescale=value_input.get('guide_rescale', 0.5),
|
||||
discretization=value_input.get('discretization',
|
||||
'trailing'),
|
||||
show_progress=True,
|
||||
seed=seed,
|
||||
condition_fn=None,
|
||||
clamp=None,
|
||||
percentile=None,
|
||||
t_max=None,
|
||||
t_min=None,
|
||||
discard_penultimate_step=None,
|
||||
intermediate_callback=intermediate_callback,
|
||||
cat_uc=cat_uc,
|
||||
**kwargs)
|
||||
|
||||
self.dynamic_unload(self.diffusion_model,
|
||||
'diffusion_model',
|
||||
skip_loaded=True)
|
||||
|
||||
if 'latent' in value_output:
|
||||
if value_output['latent'] is None or (
|
||||
isinstance(value_output['latent'], list)
|
||||
and len(value_output['latent']) < 1):
|
||||
value_output['latent'] = []
|
||||
value_output['latent'].append(latent)
|
||||
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
x_samples = self.decode_first_stage(latent).float()
|
||||
self.dynamic_unload(self.first_stage_model,
|
||||
'first_stage_model',
|
||||
skip_loaded=True)
|
||||
images = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0)
|
||||
if base_image is not None:
|
||||
stitch_images = []
|
||||
for img in images:
|
||||
stitch_img = crop_back(img, copy.deepcopy(base_image),
|
||||
extra_sizes, bbox_yyxx)
|
||||
stitch_images.append(stitch_img)
|
||||
images = torch.stack(stitch_images, dim=0)
|
||||
if 'images' in value_output:
|
||||
if value_output['images'] is None or (
|
||||
isinstance(value_output['images'], list)
|
||||
and len(value_output['images']) < 1):
|
||||
value_output['images'] = []
|
||||
value_output['images'].append(images)
|
||||
|
||||
for k, v in value_output.items():
|
||||
if isinstance(v, list):
|
||||
value_output[k] = torch.cat(v, dim=0)
|
||||
if isinstance(v, torch.Tensor):
|
||||
value_output[k] = v.cpu()
|
||||
|
||||
return value_output
|
||||
@@ -1,4 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.backbone.autoencoder.ae_module import (Decoder,
|
||||
Encoder)
|
||||
Encoder,
|
||||
RDecoder)
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import repeat
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
from scepter.modules.model.backbone.autoencoder.ae_utils import (
|
||||
@@ -244,6 +245,7 @@ class Decoder(BaseModel):
|
||||
|
||||
# compute in_ch_mult, block_in and curr_res at lowest res
|
||||
block_in = self.ch * self.ch_mult[self.num_resolutions - 1]
|
||||
self.block_in = block_in
|
||||
curr_res = 1
|
||||
# z to block_in
|
||||
self.conv_in = torch.nn.Conv2d(self.z_channels,
|
||||
@@ -340,3 +342,48 @@ class Decoder(BaseModel):
|
||||
__class__.__name__,
|
||||
Decoder.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@BACKBONES.register_class()
|
||||
class RDecoder(Decoder):
|
||||
def construct_model(self):
|
||||
super().construct_model()
|
||||
self.resize_level = nn.Sequential(
|
||||
nn.Linear(self.block_in, self.block_in),
|
||||
nn.SiLU(),
|
||||
nn.Linear(self.block_in, self.block_in),
|
||||
)
|
||||
|
||||
def forward(self, z, rembed=None):
|
||||
# timestep embedding
|
||||
temb = None
|
||||
h = self.conv_in(z)
|
||||
bs, channel, hdim, wdim = h.size()
|
||||
if rembed is not None:
|
||||
rembed = self.resize_level(rembed)
|
||||
rembed = repeat(rembed, 'b e-> b e hd wd', hd=hdim, wd=wdim)
|
||||
h = h + rembed
|
||||
|
||||
# middle
|
||||
if not self.use_checkpoint:
|
||||
h = self.mid_upsclae_transform(h, temb)
|
||||
else:
|
||||
h = checkpoint(self.mid_upsclae_transform, h, temb)
|
||||
|
||||
# end
|
||||
if self.give_pre_end:
|
||||
return h
|
||||
|
||||
h = self.norm_out(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv_out(h)
|
||||
if self.tanh_out:
|
||||
h = torch.tanh(h)
|
||||
return h
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('BACKBONE',
|
||||
__class__.__name__,
|
||||
Decoder.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.backbone.unet.unet_module import DiffusionUNet
|
||||
from scepter.modules.model.backbone.unet.unet_module import (DiffusionUNet,
|
||||
DiffusionUNetXL,
|
||||
LargenUNetXL)
|
||||
|
||||
@@ -8,8 +8,9 @@ import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from scepter.modules.model.backbone.unet.unet_utils import (
|
||||
Downsample, ResBlock, SpatialTransformer, Timestep,
|
||||
TimestepEmbedSequential, Upsample, conv_nd, linear, normalization,
|
||||
BasicTransformerBlock, Downsample, ResBlock, SpatialTransformer,
|
||||
SpatialTransformerV2, Timestep, TimestepEmbedSequential,
|
||||
TransformerBlockV2, Upsample, conv_nd, linear, normalization,
|
||||
timestep_embedding, zero_module)
|
||||
from scepter.modules.model.base_model import BaseModel
|
||||
from scepter.modules.model.registry import BACKBONES
|
||||
@@ -951,3 +952,443 @@ class DiffusionUNetXL(DiffusionUNet):
|
||||
__class__.__name__,
|
||||
DiffusionUNetXL.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@BACKBONES.register_class()
|
||||
class LargenUNetXL(DiffusionUNetXL):
|
||||
para_dict = {
|
||||
'TRANSFORMER_BLOCK_TYPE': {
|
||||
'value': 'att_v1'
|
||||
},
|
||||
'IMAGE_SCALE': {
|
||||
'value': 0.0,
|
||||
},
|
||||
}
|
||||
para_dict.update(DiffusionUNetXL.para_dict)
|
||||
|
||||
def __init__(self, cfg, logger):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.init_params(cfg)
|
||||
self.construct_network()
|
||||
|
||||
def init_params(self, cfg):
|
||||
super().init_params(cfg)
|
||||
self.transformer_block_type = cfg.get('TRANSFORMER_BLOCK_TYPE',
|
||||
'att_v1')
|
||||
TRANSFORMER_BLOCKS = {
|
||||
'att_v1': BasicTransformerBlock,
|
||||
'att_v2': TransformerBlockV2,
|
||||
}
|
||||
assert self.transformer_block_type in list(TRANSFORMER_BLOCKS.keys())
|
||||
self.transformer_block = TRANSFORMER_BLOCKS[
|
||||
self.transformer_block_type]
|
||||
self.image_scale = cfg.get('IMAGE_SCALE', 0.0)
|
||||
self.use_refine = cfg.get('USE_REFINE', False)
|
||||
|
||||
def construct_network(self):
|
||||
in_channels = self.in_channels
|
||||
model_channels = self.model_channels
|
||||
out_channels = self.out_channels
|
||||
attention_resolutions = self.attention_resolutions
|
||||
channel_mult = self.channel_mult
|
||||
num_classes = self.num_classes
|
||||
num_heads = self.num_heads
|
||||
num_head_channels = self.num_head_channels
|
||||
dims = self.dims
|
||||
dropout = self.dropout
|
||||
use_checkpoint = self.use_checkpoint
|
||||
use_scale_shift_norm = self.use_scale_shift_norm
|
||||
disable_self_attentions = self.disable_self_attentions
|
||||
disable_middle_self_attn = self.disable_middle_self_attn
|
||||
transformer_depth = self.transformer_depth
|
||||
transformer_depth_middle = self.transformer_depth_middle
|
||||
context_dim = self.context_dim
|
||||
use_linear_in_transformer = self.use_linear_in_transformer
|
||||
resblock_updown = self.resblock_updown
|
||||
conv_resample = self.conv_resample
|
||||
adm_in_channels = self.adm_in_channels
|
||||
transformer_block = self.transformer_block
|
||||
|
||||
time_embed_dim = model_channels * 4
|
||||
self.time_embed = nn.Sequential(
|
||||
linear(model_channels, time_embed_dim),
|
||||
nn.SiLU(),
|
||||
linear(time_embed_dim, time_embed_dim),
|
||||
)
|
||||
|
||||
if self.num_classes is not None:
|
||||
if isinstance(self.num_classes, int):
|
||||
self.label_emb = nn.Embedding(num_classes, time_embed_dim)
|
||||
elif self.num_classes == 'continuous':
|
||||
print('setting up linear c_adm embedding layer')
|
||||
self.label_emb = nn.Linear(1, time_embed_dim)
|
||||
elif self.num_classes == 'timestep':
|
||||
self.label_emb = nn.Sequential(
|
||||
Timestep(model_channels),
|
||||
nn.Sequential(
|
||||
linear(model_channels, time_embed_dim),
|
||||
nn.SiLU(),
|
||||
linear(time_embed_dim, time_embed_dim),
|
||||
),
|
||||
)
|
||||
elif self.num_classes == 'sequential':
|
||||
assert adm_in_channels is not None
|
||||
self.label_emb = nn.Sequential(
|
||||
nn.Sequential(
|
||||
linear(adm_in_channels, time_embed_dim),
|
||||
nn.SiLU(),
|
||||
linear(time_embed_dim, time_embed_dim),
|
||||
))
|
||||
else:
|
||||
raise ValueError()
|
||||
|
||||
self.input_blocks = nn.ModuleList([
|
||||
TimestepEmbedSequential(
|
||||
conv_nd(dims, in_channels, model_channels, 3, padding=1))
|
||||
])
|
||||
self._feature_size = model_channels
|
||||
input_block_chans = [model_channels]
|
||||
input_down_flag = [False]
|
||||
ch = model_channels
|
||||
ds = 1
|
||||
for level, mult in enumerate(channel_mult):
|
||||
for nr in range(self.num_res_blocks[level]):
|
||||
layers = [
|
||||
ResBlock(
|
||||
ch,
|
||||
time_embed_dim,
|
||||
dropout,
|
||||
out_channels=mult * model_channels,
|
||||
dims=dims,
|
||||
use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
)
|
||||
]
|
||||
ch = mult * model_channels
|
||||
if ds in attention_resolutions:
|
||||
if num_head_channels == -1:
|
||||
dim_head = ch // num_heads
|
||||
else:
|
||||
num_heads = ch // num_head_channels
|
||||
dim_head = num_head_channels
|
||||
disabled_sa = disable_self_attentions[level] if exists(
|
||||
disable_self_attentions) else False
|
||||
|
||||
layers.append(
|
||||
SpatialTransformerV2(
|
||||
ch,
|
||||
num_heads,
|
||||
dim_head,
|
||||
transformer_block=transformer_block,
|
||||
depth=transformer_depth[level],
|
||||
context_dim=context_dim,
|
||||
disable_self_attn=disabled_sa,
|
||||
use_linear=use_linear_in_transformer,
|
||||
use_checkpoint=use_checkpoint))
|
||||
self.input_blocks.append(TimestepEmbedSequential(*layers))
|
||||
self._feature_size += ch
|
||||
input_block_chans.append(ch)
|
||||
input_down_flag.append(False)
|
||||
if level != len(channel_mult) - 1:
|
||||
out_ch = ch
|
||||
self.input_blocks.append(
|
||||
TimestepEmbedSequential(
|
||||
ResBlock(
|
||||
ch,
|
||||
time_embed_dim,
|
||||
dropout,
|
||||
out_channels=out_ch,
|
||||
dims=dims,
|
||||
use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
down=True,
|
||||
) if resblock_updown else Downsample(
|
||||
ch, conv_resample, dims=dims, out_channels=out_ch))
|
||||
)
|
||||
ch = out_ch
|
||||
input_block_chans.append(ch)
|
||||
input_down_flag.append(True)
|
||||
ds *= 2
|
||||
self._feature_size += ch
|
||||
self._input_block_chans = copy.deepcopy(input_block_chans)
|
||||
self._input_down_flag = input_down_flag
|
||||
|
||||
if num_head_channels == -1:
|
||||
dim_head = ch // num_heads
|
||||
else:
|
||||
num_heads = ch // num_head_channels
|
||||
dim_head = num_head_channels
|
||||
self.middle_block = TimestepEmbedSequential(
|
||||
ResBlock(
|
||||
ch,
|
||||
time_embed_dim,
|
||||
dropout,
|
||||
dims=dims,
|
||||
use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
),
|
||||
SpatialTransformerV2(ch,
|
||||
num_heads,
|
||||
dim_head,
|
||||
transformer_block=transformer_block,
|
||||
depth=transformer_depth_middle,
|
||||
context_dim=context_dim,
|
||||
disable_self_attn=disable_middle_self_attn,
|
||||
use_linear=use_linear_in_transformer,
|
||||
use_checkpoint=use_checkpoint),
|
||||
ResBlock(
|
||||
ch,
|
||||
time_embed_dim,
|
||||
dropout,
|
||||
dims=dims,
|
||||
use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
),
|
||||
)
|
||||
self._feature_size += ch
|
||||
self._middle_block_chans = [ch]
|
||||
|
||||
self._output_block_chans = []
|
||||
self.output_blocks = nn.ModuleList([])
|
||||
for level, mult in list(enumerate(channel_mult))[::-1]:
|
||||
for i in range(self.num_res_blocks[level] + 1):
|
||||
ich = input_block_chans.pop()
|
||||
layers = [
|
||||
ResBlock(
|
||||
ch + ich,
|
||||
time_embed_dim,
|
||||
dropout,
|
||||
out_channels=model_channels * mult,
|
||||
dims=dims,
|
||||
use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
)
|
||||
]
|
||||
ch = model_channels * mult
|
||||
if ds in attention_resolutions:
|
||||
if num_head_channels == -1:
|
||||
dim_head = ch // num_heads
|
||||
else:
|
||||
num_heads = ch // num_head_channels
|
||||
dim_head = num_head_channels
|
||||
disabled_sa = disable_self_attentions[level] if exists(
|
||||
disable_self_attentions) else False
|
||||
layers.append(
|
||||
SpatialTransformerV2(
|
||||
ch,
|
||||
num_heads,
|
||||
dim_head,
|
||||
transformer_block=transformer_block,
|
||||
depth=transformer_depth[level],
|
||||
context_dim=context_dim,
|
||||
disable_self_attn=disabled_sa,
|
||||
use_linear=use_linear_in_transformer,
|
||||
use_checkpoint=use_checkpoint))
|
||||
if level and i == self.num_res_blocks[level]:
|
||||
out_ch = ch
|
||||
layers.append(
|
||||
ResBlock(
|
||||
ch,
|
||||
time_embed_dim,
|
||||
dropout,
|
||||
out_channels=out_ch,
|
||||
dims=dims,
|
||||
use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
up=True,
|
||||
) if resblock_updown else Upsample(
|
||||
ch, conv_resample, dims=dims, out_channels=out_ch))
|
||||
ds //= 2
|
||||
|
||||
self.output_blocks.append(TimestepEmbedSequential(*layers))
|
||||
|
||||
self._feature_size += ch
|
||||
self._output_block_chans.append(ch)
|
||||
|
||||
self.out = nn.Sequential(
|
||||
normalization(ch),
|
||||
nn.SiLU(),
|
||||
zero_module(
|
||||
conv_nd(dims, model_channels, out_channels, 3, padding=1)),
|
||||
)
|
||||
|
||||
if self.use_refine:
|
||||
self.ref_time_embed = copy.deepcopy(self.time_embed)
|
||||
self.ref_label_emb = copy.deepcopy(self.label_emb)
|
||||
self.ref_input_blocks = copy.deepcopy(self.input_blocks)
|
||||
self.ref_input_blocks[0] = TimestepEmbedSequential(
|
||||
conv_nd(dims, 4, model_channels, 3, padding=1))
|
||||
self.ref_middle_block = copy.deepcopy(self.middle_block)
|
||||
self.ref_output_blocks = copy.deepcopy(self.output_blocks)
|
||||
|
||||
def load_pretrained_model(self, pretrained_model):
|
||||
if pretrained_model is not None:
|
||||
with FS.get_from(pretrained_model,
|
||||
wait_finish=True) as local_model:
|
||||
self.init_from_ckpt(local_model, ignore_keys=self.ignore_keys)
|
||||
|
||||
def init_from_ckpt(self, path, ignore_keys=list()):
|
||||
if path.endswith('safetensors'):
|
||||
from safetensors.torch import load_file as load_safetensors
|
||||
sd = load_safetensors(path)
|
||||
else:
|
||||
sd = torch.load(path, map_location='cpu')
|
||||
|
||||
new_sd = OrderedDict()
|
||||
for k, v in sd.items():
|
||||
ignored = False
|
||||
for ik in ignore_keys:
|
||||
if ik in k:
|
||||
if we.rank == 0:
|
||||
self.logger.info(
|
||||
'Ignore key {} from state_dict.'.format(k))
|
||||
ignored = True
|
||||
break
|
||||
if not ignored:
|
||||
if k == 'input_blocks.0.0.weight':
|
||||
if we.rank == 0:
|
||||
self.logger.info(
|
||||
'Partial initial key {} from state_dict.'.format(
|
||||
k))
|
||||
new_v = torch.empty(320, self.in_channels, 3, 3)
|
||||
nn.init.zeros_(new_v)
|
||||
new_v[:, :v.shape[1]] = v
|
||||
new_sd[k] = new_v
|
||||
if self.use_refine:
|
||||
new_sd['ref_' + k] = v
|
||||
else:
|
||||
new_sd[k] = v
|
||||
if self.use_refine:
|
||||
new_sd['ref_' + k] = v
|
||||
|
||||
missing, unexpected = self.load_state_dict(new_sd, strict=False)
|
||||
if we.rank == 0:
|
||||
self.logger.info(
|
||||
f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys'
|
||||
)
|
||||
if len(missing) > 0:
|
||||
self.logger.info(f'Missing Keys:\n {missing}')
|
||||
if len(unexpected) > 0:
|
||||
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
|
||||
|
||||
def forward(self, x, t=None, cond=dict(), **kwargs):
|
||||
t_emb = timestep_embedding(t,
|
||||
self.model_channels,
|
||||
repeat_only=False,
|
||||
legacy=True)
|
||||
emb = self.time_embed(t_emb)
|
||||
|
||||
if isinstance(cond, dict):
|
||||
if 'y' in cond:
|
||||
assert self.num_classes is not None
|
||||
emb = emb + self.label_emb(cond['y'])
|
||||
if self.use_refine:
|
||||
ref_emb = self.ref_time_embed(t_emb)
|
||||
assert 'null_y' in cond
|
||||
cond_y = cond['y'].clone()
|
||||
cond_y[:, :cond['null_y'].shape[1]] = cond['null_y']
|
||||
ref_emb = ref_emb + self.ref_label_emb(cond_y)
|
||||
|
||||
if 'concat' in cond:
|
||||
c = cond['concat']
|
||||
x = torch.cat([x, c], dim=1)
|
||||
|
||||
context = cond.get('crossattn', None)
|
||||
img_context = cond.get('img_crossattn', None)
|
||||
|
||||
task = cond['task']
|
||||
image_scale = cond.get('image_scale', self.image_scale)
|
||||
if 'Subject' in task and img_context is not None:
|
||||
ip_enc_scale = image_scale
|
||||
ip_dec_scale = image_scale
|
||||
num_img_tokens = img_context.shape[1]
|
||||
context = torch.cat([context, img_context], dim=1)
|
||||
else:
|
||||
ip_enc_scale = None
|
||||
ip_dec_scale = None
|
||||
num_img_tokens = None
|
||||
|
||||
ref = cond.get('ref_xt', None)
|
||||
ref_context = cond.get('ref_crossattn', None)
|
||||
else:
|
||||
raise TypeError
|
||||
|
||||
hs = []
|
||||
refs = []
|
||||
h = x
|
||||
|
||||
if self.use_refine:
|
||||
assert ref is not None and ref_context is not None
|
||||
for i, (ref_module, module) in enumerate(
|
||||
zip(self.ref_input_blocks, self.input_blocks)):
|
||||
ref = ref_module(ref, ref_emb, ref_context, caching=None)
|
||||
h = module(h,
|
||||
emb,
|
||||
context,
|
||||
caching=None,
|
||||
scale=ip_enc_scale,
|
||||
num_img_token=num_img_tokens)
|
||||
refs.append(ref)
|
||||
hs.append(h)
|
||||
|
||||
ref = self.ref_middle_block(ref,
|
||||
ref_emb,
|
||||
ref_context,
|
||||
caching=None)
|
||||
h = self.middle_block(h,
|
||||
emb,
|
||||
context,
|
||||
caching=None,
|
||||
scale=ip_enc_scale,
|
||||
num_img_token=num_img_tokens)
|
||||
|
||||
for i, (ref_module, module) in enumerate(
|
||||
zip(self.ref_output_blocks, self.output_blocks)):
|
||||
cache = []
|
||||
ref = torch.cat([ref, refs.pop()], dim=1)
|
||||
ref = ref_module(ref,
|
||||
ref_emb,
|
||||
ref_context,
|
||||
caching='write',
|
||||
cache=cache)
|
||||
h = torch.cat([h, hs.pop()], dim=1)
|
||||
h = module(h,
|
||||
emb,
|
||||
context,
|
||||
caching='read',
|
||||
cache=cache,
|
||||
scale=ip_dec_scale,
|
||||
num_img_token=num_img_tokens)
|
||||
else:
|
||||
for module in self.input_blocks:
|
||||
h = module(h,
|
||||
emb,
|
||||
context,
|
||||
caching=None,
|
||||
scale=ip_enc_scale,
|
||||
num_img_token=num_img_tokens)
|
||||
hs.append(h)
|
||||
h = self.middle_block(h,
|
||||
emb,
|
||||
context,
|
||||
caching=None,
|
||||
scale=ip_enc_scale,
|
||||
num_img_token=num_img_tokens)
|
||||
for module in self.output_blocks:
|
||||
h = torch.cat([h, hs.pop()], dim=1)
|
||||
h = module(h,
|
||||
emb,
|
||||
context,
|
||||
caching=None,
|
||||
scale=ip_dec_scale,
|
||||
num_img_token=num_img_tokens)
|
||||
|
||||
out = self.out(h)
|
||||
return out
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('BACKBONE',
|
||||
__class__.__name__,
|
||||
LargenUNetXL.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@@ -10,6 +10,7 @@ import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms.functional as TF
|
||||
from einops import rearrange, repeat
|
||||
from packaging import version
|
||||
|
||||
@@ -171,12 +172,14 @@ class TimestepEmbedSequential(nn.Sequential, TimestepBlock):
|
||||
A sequential module that passes timestep embeddings to the children that
|
||||
support it as an extra input.
|
||||
"""
|
||||
def forward(self, x, emb, context=None, target_size=None):
|
||||
def forward(self, x, emb, context=None, target_size=None, **kwargs):
|
||||
for layer in self:
|
||||
if isinstance(layer, TimestepBlock):
|
||||
x = layer(x, emb)
|
||||
elif isinstance(layer, SpatialTransformer):
|
||||
x = layer(x, context)
|
||||
elif isinstance(layer, SpatialTransformerV2):
|
||||
x = layer(x, context, **kwargs)
|
||||
elif isinstance(layer, Upsample):
|
||||
x = layer(x, target_size)
|
||||
else:
|
||||
@@ -864,6 +867,92 @@ class MemoryEfficientCrossAttention(nn.Module):
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class XFormersMHA_IP(nn.Module):
|
||||
def __init__(self,
|
||||
query_dim,
|
||||
context_dim=None,
|
||||
heads=8,
|
||||
dim_head=64,
|
||||
dropout=0.0):
|
||||
super().__init__()
|
||||
inner_dim = dim_head * heads
|
||||
context_dim = default(context_dim, query_dim)
|
||||
|
||||
self.heads = heads
|
||||
self.dim_head = dim_head
|
||||
|
||||
self.to_q = nn.Linear(query_dim, inner_dim, bias=False)
|
||||
self.to_k = nn.Linear(context_dim, inner_dim, bias=False)
|
||||
self.to_v = nn.Linear(context_dim, inner_dim, bias=False)
|
||||
self.to_k_ip = nn.Linear(context_dim, inner_dim, bias=False)
|
||||
self.to_v_ip = nn.Linear(context_dim, inner_dim, bias=False)
|
||||
|
||||
self.to_out = nn.Sequential(nn.Linear(inner_dim, query_dim),
|
||||
nn.Dropout(dropout))
|
||||
self.attention_op = None
|
||||
|
||||
def forward(self,
|
||||
x,
|
||||
context=None,
|
||||
mask=None,
|
||||
scale=None,
|
||||
num_img_token=None):
|
||||
q = self.to_q(x)
|
||||
context = default(context, x)
|
||||
|
||||
if scale is not None and num_img_token is not None:
|
||||
eos = context.shape[1] - num_img_token
|
||||
txt_context = context[:, :eos, :]
|
||||
img_context = context[:, eos:, :]
|
||||
|
||||
k = self.to_k(txt_context)
|
||||
v = self.to_v(txt_context)
|
||||
k_i = self.to_k_ip(img_context)
|
||||
v_i = self.to_v_ip(img_context)
|
||||
|
||||
b, _, _ = q.shape
|
||||
q, k, v, k_i, v_i = map(
|
||||
lambda t: t.unsqueeze(3).reshape(b, t.shape[
|
||||
1], self.heads, self.dim_head).permute(0, 2, 1, 3).reshape(
|
||||
b * self.heads, t.shape[1], self.dim_head).contiguous(
|
||||
),
|
||||
(q, k, v, k_i, v_i),
|
||||
)
|
||||
|
||||
# actually compute the attention, what we cannot get enough of
|
||||
txt_out = xformers.ops.memory_efficient_attention(
|
||||
q, k, v, attn_bias=None, op=self.attention_op)
|
||||
img_out = xformers.ops.memory_efficient_attention(
|
||||
q, k_i, v_i, attn_bias=None, op=self.attention_op)
|
||||
out = txt_out + scale * img_out
|
||||
else:
|
||||
k = self.to_k(context)
|
||||
v = self.to_v(context)
|
||||
b, _, _ = q.shape
|
||||
q, k, v = map(
|
||||
lambda t: t.unsqueeze(3).reshape(b, t.shape[
|
||||
1], self.heads, self.dim_head).permute(0, 2, 1, 3).reshape(
|
||||
b * self.heads, t.shape[1], self.dim_head).contiguous(
|
||||
),
|
||||
(q, k, v),
|
||||
)
|
||||
out = xformers.ops.memory_efficient_attention(q,
|
||||
k,
|
||||
v,
|
||||
attn_bias=None,
|
||||
op=self.attention_op)
|
||||
|
||||
# TODO: Use this directly in the attention operation, as a bias
|
||||
if exists(mask):
|
||||
raise NotImplementedError
|
||||
out = (out.unsqueeze(0).reshape(
|
||||
b, self.heads, out.shape[1],
|
||||
self.dim_head).permute(0, 2, 1,
|
||||
3).reshape(b, out.shape[1],
|
||||
self.heads * self.dim_head))
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class BasicTransformerBlock(nn.Module):
|
||||
def __init__(self,
|
||||
dim,
|
||||
@@ -908,6 +997,65 @@ class BasicTransformerBlock(nn.Module):
|
||||
return x
|
||||
|
||||
|
||||
class TransformerBlockV2(nn.Module):
|
||||
def __init__(self,
|
||||
query_dim,
|
||||
n_heads,
|
||||
d_head,
|
||||
dropout=0.,
|
||||
context_dim=None,
|
||||
gated_ff=True,
|
||||
use_checkpoint=False,
|
||||
disable_self_attn=False):
|
||||
super().__init__()
|
||||
self.disable_self_attn = disable_self_attn
|
||||
self.attn1 = MemoryEfficientCrossAttention(query_dim=query_dim,
|
||||
heads=n_heads,
|
||||
dim_head=d_head,
|
||||
dropout=dropout,
|
||||
context_dim=None)
|
||||
self.ff = FeedForward(query_dim, dropout=dropout, glu=gated_ff)
|
||||
self.attn2 = XFormersMHA_IP(query_dim=query_dim,
|
||||
heads=n_heads,
|
||||
dim_head=d_head,
|
||||
context_dim=context_dim)
|
||||
self.norm1 = nn.LayerNorm(query_dim)
|
||||
self.norm2 = nn.LayerNorm(query_dim)
|
||||
self.norm3 = nn.LayerNorm(query_dim)
|
||||
self.use_checkpoint = use_checkpoint
|
||||
|
||||
def forward(self,
|
||||
x,
|
||||
context,
|
||||
caching=None,
|
||||
cache=None,
|
||||
scale=None,
|
||||
num_img_token=None,
|
||||
**kwargs):
|
||||
y = self.norm1(x)
|
||||
if caching == 'write':
|
||||
assert isinstance(cache, list)
|
||||
cache.append(y)
|
||||
x = self.attn1(y, context=None) + x
|
||||
elif caching == 'read':
|
||||
assert isinstance(cache, list) and len(cache) > 0
|
||||
c = cache.pop(0)
|
||||
self_ctx = torch.cat([y, c], dim=1)
|
||||
x = self.attn1(y, context=self_ctx) + x
|
||||
elif caching is None:
|
||||
x = self.attn1(y, context=None) + x
|
||||
else:
|
||||
assert False
|
||||
|
||||
x = self.attn2(self.norm2(x),
|
||||
context=context,
|
||||
scale=scale,
|
||||
num_img_token=num_img_token) + x
|
||||
x = self.ff(self.norm3(x)) + x
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class SpatialTransformer(nn.Module):
|
||||
"""
|
||||
Transformer block for image-like data.
|
||||
@@ -1003,3 +1151,108 @@ class SpatialTransformer(nn.Module):
|
||||
if not self.use_linear:
|
||||
x = self.proj_out(x)
|
||||
return x + x_in
|
||||
|
||||
|
||||
class SpatialTransformerV2(nn.Module):
|
||||
"""
|
||||
Transformer block for image-like data.
|
||||
First, project the input (aka embedding)
|
||||
and reshape to b, t, d.
|
||||
Then apply standard transformer action.
|
||||
Finally, reshape to image
|
||||
NEW: use_linear for more efficiency instead of the 1x1 convs
|
||||
"""
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
n_heads,
|
||||
d_head,
|
||||
transformer_block,
|
||||
depth=1,
|
||||
dropout=0.,
|
||||
context_dim=None,
|
||||
disable_self_attn=False,
|
||||
use_linear=False,
|
||||
use_checkpoint=True):
|
||||
super().__init__()
|
||||
if exists(context_dim) and not isinstance(context_dim, list):
|
||||
context_dim = [context_dim]
|
||||
|
||||
if exists(context_dim) and not isinstance(context_dim, (list)):
|
||||
context_dim = [context_dim]
|
||||
if exists(context_dim) and isinstance(context_dim, list):
|
||||
if depth != len(context_dim):
|
||||
print(
|
||||
f'WARNING: {self.__class__.__name__}: Found context dims {context_dim} of'
|
||||
f" depth {len(context_dim)}, which does not match the specified 'depth' of"
|
||||
f' {depth}. Setting context_dim to {depth * [context_dim[0]]} now.'
|
||||
)
|
||||
# depth does not match context dims.
|
||||
assert all(
|
||||
map(lambda x: x == context_dim[0], context_dim)
|
||||
), 'need homogenous context_dim to match depth automatically'
|
||||
context_dim = depth * [context_dim[0]]
|
||||
elif context_dim is None:
|
||||
context_dim = [None] * depth
|
||||
|
||||
self.in_channels = in_channels
|
||||
inner_dim = n_heads * d_head
|
||||
self.norm = normalization(in_channels)
|
||||
if not use_linear:
|
||||
self.proj_in = nn.Conv2d(in_channels,
|
||||
inner_dim,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0)
|
||||
else:
|
||||
self.proj_in = nn.Linear(in_channels, inner_dim)
|
||||
|
||||
self.transformer_blocks = nn.ModuleList([
|
||||
transformer_block(inner_dim,
|
||||
n_heads,
|
||||
d_head,
|
||||
dropout=dropout,
|
||||
context_dim=context_dim[d],
|
||||
disable_self_attn=disable_self_attn,
|
||||
use_checkpoint=use_checkpoint)
|
||||
for d in range(depth)
|
||||
])
|
||||
if not use_linear:
|
||||
self.proj_out = zero_module(
|
||||
nn.Conv2d(inner_dim,
|
||||
in_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0))
|
||||
else:
|
||||
self.proj_out = zero_module(nn.Linear(in_channels, inner_dim))
|
||||
self.use_linear = use_linear
|
||||
|
||||
def forward(self, x, context=None, **kwargs):
|
||||
# note: if no context is given, cross-attention defaults to self-attention
|
||||
if not isinstance(context, list):
|
||||
context = [context]
|
||||
b, c, h, w = x.shape
|
||||
|
||||
ref_mask = kwargs.pop('ref_mask', None)
|
||||
if ref_mask is not None:
|
||||
ref_mask = TF.resize(ref_mask, (h, w), antialias=True)
|
||||
ref_mask = (ref_mask > 0.5).float()
|
||||
ref_mask = rearrange(ref_mask, 'b c h w -> b (h w) c').contiguous()
|
||||
|
||||
x_in = x
|
||||
x = self.norm(x)
|
||||
if not self.use_linear:
|
||||
x = self.proj_in(x)
|
||||
x = rearrange(x, 'b c h w -> b (h w) c').contiguous()
|
||||
if self.use_linear:
|
||||
x = self.proj_in(x)
|
||||
for i, block in enumerate(self.transformer_blocks):
|
||||
if i > 0 and len(context) == 1:
|
||||
i = 0 # use same context for each block
|
||||
x = block(x, context=context[i], ref_mask=ref_mask, **kwargs)
|
||||
if self.use_linear:
|
||||
x = self.proj_out(x)
|
||||
x = rearrange(x, 'b (h w) c -> b c h w', h=h, w=w).contiguous()
|
||||
if not self.use_linear:
|
||||
x = self.proj_out(x)
|
||||
return x + x_in
|
||||
|
||||
@@ -1,17 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional
|
||||
from einops import rearrange
|
||||
|
||||
from scepter.modules.model.backbone.video.init_helper import (
|
||||
_init_transformer_weights, trunc_normal_)
|
||||
from scepter.modules.model.registry import BACKBONES, BRICKS, STEMS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
'''
|
||||
The implementations of vivit as https://arxiv.org/abs/2103.15691.
|
||||
The following setting alined the proposed model in the paper above.
|
||||
@@ -39,6 +27,18 @@ TimesFormer:
|
||||
complexity: (n_h * n_w) ** 2 + O(attn_temp)
|
||||
'''
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional
|
||||
from einops import rearrange
|
||||
|
||||
from scepter.modules.model.backbone.video.init_helper import (
|
||||
_init_transformer_weights, trunc_normal_)
|
||||
from scepter.modules.model.registry import BACKBONES, BRICKS, STEMS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
|
||||
|
||||
@BACKBONES.register_class()
|
||||
class VideoTransformer(nn.Module):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,58 @@ 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:
|
||||
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:
|
||||
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,
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1,138 @@
|
||||
# -*- 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 < 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)
|
||||
@@ -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
|
||||
|
||||
@@ -333,12 +333,12 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
data_iter = iter(self.datas[self._mode].dataloader)
|
||||
self.print_memory_status()
|
||||
for step in range(self.max_steps):
|
||||
if 'eval' in self._mode_set and (step % self.eval_interval == 0
|
||||
or step == self.max_steps - 1):
|
||||
if 'eval' in self._mode_set and (self.eval_interval > 0 and
|
||||
step % self.eval_interval == 0):
|
||||
self.run_eval()
|
||||
self.train_mode()
|
||||
self.before_iter(self.hooks_dict[self._mode])
|
||||
batch_data = next(data_iter)
|
||||
self.before_iter(self.hooks_dict[self._mode])
|
||||
if self.sample_args:
|
||||
batch_data.update(self.sample_args.get_lowercase_dict())
|
||||
if 'meta' in batch_data:
|
||||
@@ -364,6 +364,9 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
self.after_iter(self.hooks_dict[self._mode])
|
||||
if we.debug:
|
||||
self.print_trainable_params_status(prefix='model.')
|
||||
if 'eval' in self._mode_set and (self.eval_interval > 0
|
||||
and step == self.max_steps - 1):
|
||||
self.run_eval()
|
||||
self.after_all_iter(self.hooks_dict[self._mode])
|
||||
|
||||
@torch.no_grad()
|
||||
|
||||
@@ -4,6 +4,7 @@
|
||||
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
|
||||
@@ -47,5 +48,6 @@ after solve:
|
||||
|
||||
__all__ = [
|
||||
'HOOKS', 'BackwardHook', 'CheckpointHook', 'Hook', 'LrHook', 'LogHook',
|
||||
'TensorboardLogHook', 'DistSamplerHook', 'ProbeDataHook', 'SafetensorsHook'
|
||||
'TensorboardLogHook', 'DistSamplerHook', 'ProbeDataHook',
|
||||
'SafetensorsHook', 'ModelEmaHook'
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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_name_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()
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
import torch
|
||||
from torch.distributed.fsdp import (FullStateDictConfig,
|
||||
FullyShardedDataParallel, StateDictType)
|
||||
|
||||
from scepter.modules.solver.hooks.hook import Hook
|
||||
from scepter.modules.solver.hooks.registry import HOOKS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
|
||||
|
||||
@HOOKS.register_class()
|
||||
class ModelEmaHook(Hook):
|
||||
para_dict = [{
|
||||
'PRIORITY': {
|
||||
'value': 100,
|
||||
'description': 'the priority for processing!'
|
||||
},
|
||||
'BETA': {
|
||||
'value': 0.9999,
|
||||
'description': ''
|
||||
},
|
||||
}]
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(ModelEmaHook, self).__init__(cfg, logger=logger)
|
||||
self.priority = cfg.get('PRIORITY', 100)
|
||||
self.beta = cfg.get('BETA', 0.9999)
|
||||
|
||||
def after_iter(self, solver):
|
||||
if solver.model.use_ema:
|
||||
model_ema = solver.model.model_ema
|
||||
model = solver.model.model
|
||||
self.ema(model_ema, model, use_fsdp=solver.use_fsdp)
|
||||
|
||||
@torch.no_grad()
|
||||
def ema(self, net_ema, net, use_fsdp=True):
|
||||
if we.is_distributed:
|
||||
if use_fsdp:
|
||||
save_policy = FullStateDictConfig(offload_to_cpu=False,
|
||||
rank0_only=False)
|
||||
with FullyShardedDataParallel.state_dict_type(
|
||||
net, StateDictType.FULL_STATE_DICT, save_policy):
|
||||
nonema_state = net.state_dict()
|
||||
elif hasattr(net, 'module'):
|
||||
nonema_state = net.module.state_dict()
|
||||
else:
|
||||
nonema_state = net.state_dict()
|
||||
else:
|
||||
nonema_state = net.state_dict()
|
||||
|
||||
for k, v in net_ema.named_parameters():
|
||||
v.copy_(nonema_state[k].lerp(v, self.beta))
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('hook',
|
||||
__class__.__name__,
|
||||
ModelEmaHook.para_dict,
|
||||
set_name=True)
|
||||
@@ -10,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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -603,3 +603,6 @@ class Config(object):
|
||||
return cfg_new
|
||||
else:
|
||||
return cfg
|
||||
|
||||
def pop(self, name):
|
||||
self.cfg_dict.pop(name)
|
||||
|
||||
@@ -4,10 +4,10 @@ import io
|
||||
from io import BytesIO
|
||||
|
||||
import onnx
|
||||
import onnxruntime
|
||||
import torch
|
||||
from torch.onnx import OperatorExportTypes
|
||||
|
||||
import onnxruntime
|
||||
from scepter.modules.utils.distribute import we
|
||||
|
||||
type_map = {
|
||||
|
||||
@@ -293,6 +293,9 @@ class LocalFs(BaseFs):
|
||||
return True
|
||||
|
||||
def get_logging_handler(self, target_logging_path):
|
||||
dirname = os.path.dirname(target_logging_path)
|
||||
if not os.path.exists(dirname):
|
||||
os.makedirs(dirname, exist_ok=True)
|
||||
return logging.FileHandler(target_logging_path)
|
||||
|
||||
def put_dir_from_local_dir(self,
|
||||
@@ -303,11 +306,11 @@ class LocalFs(BaseFs):
|
||||
target_dir = self.reconstruct_path(target_dir)
|
||||
if local_dir == target_dir:
|
||||
return True
|
||||
# cp -f local_dir/* target_dir/*
|
||||
if not osp.exists(target_dir):
|
||||
status = os.system(f'mkdir -p {target_dir}')
|
||||
if status != 0:
|
||||
return False
|
||||
# # cp -f local_dir/* target_dir/*
|
||||
# if not osp.exists(target_dir):
|
||||
# status = os.system(f'mkdir -p {target_dir}')
|
||||
# if status != 0:
|
||||
# return False
|
||||
try:
|
||||
shutil.copytree(local_dir, target_dir, symlinks=True)
|
||||
except Exception:
|
||||
|
||||
@@ -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()}',
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
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
|
||||
@@ -263,7 +263,8 @@ class ProbeData():
|
||||
url = FS.get_url(one_path,
|
||||
lifecycle=3600 * 365 * 24).replace(
|
||||
'.oss-internal.aliyun-inc.',
|
||||
'.oss.aliyuncs.')
|
||||
'.oss.aliyuncs.').replace(
|
||||
'-internal', '')
|
||||
one_rank += (
|
||||
f'<td align="center"><input type="image" src="{url}" >'
|
||||
f'<br><font size="4"><strong>{save_id}-{idx}|{one_label}<strong></font><br/></td>'
|
||||
|
||||
@@ -16,6 +16,7 @@ from scepter.studio.inference.inference_ui.component_names import \
|
||||
from scepter.studio.inference.inference_ui.control_ui import ControlUI
|
||||
from scepter.studio.inference.inference_ui.diffusion_ui import DiffusionUI
|
||||
from scepter.studio.inference.inference_ui.gallery_ui import GalleryUI
|
||||
from scepter.studio.inference.inference_ui.largen_ui import LargenUI
|
||||
from scepter.studio.inference.inference_ui.mantra_ui import MantraUI
|
||||
from scepter.studio.inference.inference_ui.model_manage_ui import ModelManageUI
|
||||
from scepter.studio.inference.inference_ui.refiner_ui import RefinerUI
|
||||
@@ -23,7 +24,7 @@ from scepter.studio.inference.inference_ui.tuner_ui import TunerUI
|
||||
from scepter.studio.utils.env import init_env
|
||||
|
||||
UI_MAP = [('diffusion', DiffusionUI), ('mantra', MantraUI), ('tuner', TunerUI),
|
||||
('control', ControlUI), ('refiner', RefinerUI)]
|
||||
('control', ControlUI), ('refiner', RefinerUI), ('largen', LargenUI)]
|
||||
|
||||
|
||||
class InferenceUI():
|
||||
@@ -54,6 +55,19 @@ class InferenceUI():
|
||||
cfg_general.EXTENSION_PARAS.OFFICIAL_CONTROLLERS))
|
||||
cfg_general.CONTROLLERS = official_controllers.CONTROLLERS
|
||||
|
||||
# customized tuners
|
||||
tuner_manager = Config(
|
||||
cfg_file=os.path.join(os.path.dirname(scepter.dirname),
|
||||
cfg_general.EXTENSION_PARAS.TUNER_MANAGER))
|
||||
tuner_manager = os.path.join(root_work_dir, tuner_manager.WORK_DIR,
|
||||
tuner_manager.TUNER_LIST_YAML)
|
||||
if FS.exists(tuner_manager):
|
||||
with FS.get_from(tuner_manager) as local_path:
|
||||
custom_tuners = Config(cfg_file=local_path)
|
||||
cfg_general.CUSTOM_TUNERS = custom_tuners.get('TUNERS', [])
|
||||
else:
|
||||
cfg_general.CUSTOM_TUNERS = []
|
||||
|
||||
pipe_manager = PipelineManager()
|
||||
config_list = glob(os.path.join(config_dir, '*/*_pro.yaml'),
|
||||
recursive=True)
|
||||
@@ -68,6 +82,12 @@ class InferenceUI():
|
||||
for one_controller in cfg_general.CONTROLLERS:
|
||||
pipe_manager.register_controllers(one_controller)
|
||||
|
||||
for one_tuner in cfg_general.CUSTOM_TUNERS:
|
||||
pipe_manager.register_tuner(
|
||||
one_tuner,
|
||||
name=one_tuner.NAME_ZH if language == 'zh' else one_tuner.NAME,
|
||||
is_customized=True)
|
||||
|
||||
self.model_manage_ui = ModelManageUI(cfg_general,
|
||||
pipe_manager,
|
||||
is_debug=is_debug,
|
||||
@@ -88,7 +108,8 @@ class InferenceUI():
|
||||
self.tab_ui_kwargs[f'{name}_ui'] = ui
|
||||
self.__setattr__(f'{name}_ui', ui)
|
||||
|
||||
self.check_box_controlled_tabs = ['mantra', 'tuner', 'control']
|
||||
self.check_box_controlled_tabs = ['mantra', 'tuner', 'control', 'largen']
|
||||
self.pipe_manager = pipe_manager
|
||||
assert len(self.component_names.check_box_for_setting) == len(
|
||||
self.check_box_controlled_tabs)
|
||||
|
||||
@@ -96,6 +117,7 @@ class InferenceUI():
|
||||
# create model
|
||||
self.model_manage_ui.create_ui()
|
||||
self.gallery_ui.create_ui()
|
||||
self.infer_info = gr.State(value=None)
|
||||
|
||||
# create tabs
|
||||
def create_tab(name, ui):
|
||||
@@ -123,20 +145,65 @@ class InferenceUI():
|
||||
for name, ui in self.tab_ui_kwargs.items():
|
||||
ui.set_callbacks(self.model_manage_ui,
|
||||
**self.tab_ui_kwargs,
|
||||
gallery_ui=self.gallery_ui)
|
||||
gallery_ui=self.gallery_ui,
|
||||
manager=manager)
|
||||
|
||||
def change_setting_tab(check_box, *args):
|
||||
def change_setting_tab(check_box, default_diffusion_model, *args):
|
||||
selected_tab = 'diffusion_ui'
|
||||
ui_tabs_state = [False] * len(args)
|
||||
largen_index = self.check_box_controlled_tabs.index('largen')
|
||||
largen_key = self.component_names.check_box_for_setting[largen_index]
|
||||
largen_status = args[largen_index]
|
||||
for key in check_box:
|
||||
i = self.component_names.check_box_for_setting.index(key)
|
||||
ui_tabs_state[i] = True
|
||||
if ui_tabs_state[i] != args[i]:
|
||||
selected_tab = self.check_box_controlled_tabs[i] + '_ui'
|
||||
new_check_box_value = check_box
|
||||
for key in check_box:
|
||||
i = self.component_names.check_box_for_setting.index(key)
|
||||
if ui_tabs_state[i] != args[i]:
|
||||
if i in [largen_index]:
|
||||
new_check_box_value = [key]
|
||||
for j in range(len(ui_tabs_state)):
|
||||
ui_tabs_state[j] = j == i
|
||||
else:
|
||||
new_check_box_value = [
|
||||
k for k in check_box
|
||||
if k not in [largen_key]
|
||||
]
|
||||
ui_tabs_state[largen_index] = False
|
||||
|
||||
ui_tabs_updates = [gr.update(visible=v) for v in ui_tabs_state]
|
||||
|
||||
return gr.update(
|
||||
selected=selected_tab), *ui_tabs_state, *ui_tabs_updates
|
||||
if ui_tabs_state[largen_index]:
|
||||
diffusion_model = gr.Dropdown(
|
||||
label=self.model_manage_ui.component_names.diffusion_model,
|
||||
choices=self.model_manage_ui.
|
||||
default_choices['diffusion_model']['choices'],
|
||||
value='LARGEN_LargenUNetXL',
|
||||
interactive=False)
|
||||
elif largen_status:
|
||||
diffusion_model = gr.Dropdown(
|
||||
label=self.model_manage_ui.component_names.diffusion_model,
|
||||
choices=self.model_manage_ui.
|
||||
default_choices['diffusion_model']['choices'],
|
||||
value=self.model_manage_ui.
|
||||
default_choices['diffusion_model']['default'],
|
||||
interactive=True)
|
||||
else:
|
||||
diffusion_model = gr.Dropdown(
|
||||
label=self.model_manage_ui.component_names.diffusion_model,
|
||||
choices=self.model_manage_ui.
|
||||
default_choices['diffusion_model']['choices'],
|
||||
value=default_diffusion_model,
|
||||
interactive=True)
|
||||
|
||||
return gr.CheckboxGroup(
|
||||
choices=self.component_names.check_box_for_setting,
|
||||
value=new_check_box_value,
|
||||
show_label=False), gr.update(
|
||||
selected=selected_tab), *ui_tabs_state, *ui_tabs_updates, diffusion_model
|
||||
|
||||
gr_states = [
|
||||
self.tab_ui[name].state for name in self.check_box_controlled_tabs
|
||||
@@ -146,8 +213,9 @@ class InferenceUI():
|
||||
]
|
||||
self.check_box_for_setting.change(
|
||||
change_setting_tab,
|
||||
inputs=[self.check_box_for_setting, *gr_states],
|
||||
outputs=[self.setting_tab, *gr_states, *gr_tabs],
|
||||
inputs=[self.check_box_for_setting, self.model_manage_ui.diffusion_model, *gr_states],
|
||||
outputs=[self.check_box_for_setting, self.setting_tab, *gr_states,
|
||||
*gr_tabs, self.model_manage_ui.diffusion_model],
|
||||
queue=False)
|
||||
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.inference.diffusion_inference import DiffusionInference
|
||||
from scepter.modules.inference.largen_inference import LargenInference
|
||||
from scepter.modules.utils.logger import get_logger
|
||||
|
||||
|
||||
@@ -95,7 +96,12 @@ class PipelineManager():
|
||||
pass
|
||||
|
||||
def register_pipeline(self, cfg):
|
||||
new_inference = DiffusionInference(logger=self.logger)
|
||||
pipeline_name = cfg.NAME
|
||||
if 'LARGEN' in pipeline_name:
|
||||
PipelineBuilder = LargenInference
|
||||
else:
|
||||
PipelineBuilder = DiffusionInference
|
||||
new_inference = PipelineBuilder(logger=self.logger)
|
||||
new_inference.init_from_cfg(cfg)
|
||||
self.contruct_models_index(cfg.NAME, new_inference)
|
||||
|
||||
|
||||
@@ -1,16 +1,17 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
# For dataset manager
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.modules.utils.directory import get_md5
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
|
||||
def download_image(image):
|
||||
if image is not None:
|
||||
client = FS.get_fs_client(image)
|
||||
if client.tmp_dir.startswith("/home"):
|
||||
if client.tmp_dir.startswith('/home'):
|
||||
name = get_md5(image)
|
||||
local_path = FS.get_from(image, f"/tmp/gradio/scepter_examples/{name}")
|
||||
local_path = FS.get_from(image,
|
||||
f'/tmp/gradio/scepter_examples/{name}')
|
||||
else:
|
||||
local_path = FS.get_from(image)
|
||||
return local_path
|
||||
@@ -23,21 +24,23 @@ class InferenceUIName():
|
||||
if language == 'en':
|
||||
self.advance_block_name = 'Advance Setting'
|
||||
self.check_box_for_setting = [
|
||||
'Use Mantra', 'Use Tuners', 'Use Controller'
|
||||
'Use Mantra', 'Use Tuners', 'Use Controller', 'LAR-Gen'
|
||||
]
|
||||
self.diffusion_paras = 'Generation Setting'
|
||||
self.mantra_paras = 'Mantra Book'
|
||||
self.tuner_paras = 'Tuners'
|
||||
self.control_paras = 'Controlable Generation'
|
||||
self.refiner_paras = 'Refiner Setting'
|
||||
self.largen_paras = 'LAR-Gen'
|
||||
elif language == 'zh':
|
||||
self.advance_block_name = '生成选项'
|
||||
self.check_box_for_setting = ['使用咒语', '使用微调', '使用控制']
|
||||
self.check_box_for_setting = ['使用咒语', '使用微调', '使用控制', 'LAR-Gen']
|
||||
self.diffusion_paras = '生成参数设置'
|
||||
self.mantra_paras = '咒语书'
|
||||
self.tuner_paras = '微调模型'
|
||||
self.control_paras = '可控生成'
|
||||
self.refiner_paras = 'Refine设置'
|
||||
self.largen_paras = 'LAR-Gen'
|
||||
|
||||
|
||||
class ModelManageUIName():
|
||||
@@ -216,6 +219,7 @@ class TunerUIName():
|
||||
self.example_block_name = 'Examples'
|
||||
self.examples = [[['Pencil Sketch Drawing'], 'a girl in a jacket'],
|
||||
[['Flat 2D Art'], 'a cat']]
|
||||
self.save_button = 'Save'
|
||||
|
||||
elif language == 'zh':
|
||||
self.tuner_model = '微调模型'
|
||||
@@ -231,6 +235,7 @@ class TunerUIName():
|
||||
self.example_block_name = '样例'
|
||||
self.examples = [[['铅笔素描'], 'a girl in a jacket'],
|
||||
[['扁平2D艺术'], 'a cat']]
|
||||
self.save_button = '保存'
|
||||
|
||||
|
||||
class ControlUIName():
|
||||
@@ -342,3 +347,155 @@ class ControlUIName():
|
||||
self.advance_block_name = '高级设置'
|
||||
self.control_scale = '控制强度'
|
||||
self.example_block_name = '样例'
|
||||
|
||||
|
||||
class LargenUIName():
|
||||
def __init__(self, language='en'):
|
||||
|
||||
self.tasks = [
|
||||
'Text_Guided_Outpainting', 'Subject_Guided_Inpainting',
|
||||
'Text_Guided_Inpainting', 'Text_Subject_Guided_Inpainting'
|
||||
]
|
||||
|
||||
if language == 'en':
|
||||
self.apps = [
|
||||
'Zoom Out', 'Virtual Try On', 'Inpainting (text guided)',
|
||||
'Inpainting (text + reference image guided)'
|
||||
]
|
||||
self.dropdown_name = 'Application'
|
||||
self.subject_image = 'Reference Image'
|
||||
self.subject_mask = 'Reference Mask'
|
||||
self.scene_image = 'Scene Image'
|
||||
self.scene_mask = 'Scene Mask'
|
||||
self.prompt = 'Prompt'
|
||||
self.masked_image = 'Masked Image'
|
||||
self.preprocess = 'Input Preprocess'
|
||||
self.button_name = 'Data Preprocess'
|
||||
self.direction = (
|
||||
'Instruction: \n\n'
|
||||
'For customized data: \n\n'
|
||||
'a.1) Select the task; \n\n'
|
||||
'a.1) Upload scene image; \n\n'
|
||||
'a.2) Upload reference image (if needed); \n\n'
|
||||
'a.3) Use the brush tool to cover the areas on the scene'
|
||||
'image and reference image (to generate corresponding scene'
|
||||
'mask and reference mask); \n\n'
|
||||
'a.4) Click Data Preprocess button; \n\n'
|
||||
'a.5) Input text prompt; \n\n'
|
||||
'For example data: \n\n'
|
||||
'b.1) click example row \n\n'
|
||||
'Finally, click Generate button and get the output image!')
|
||||
self.out_direction_label = 'Out Direction'
|
||||
self.out_directions = [
|
||||
'CenterAround',
|
||||
'RightDown',
|
||||
'LeftDown',
|
||||
'RightUp',
|
||||
'LeftUp',
|
||||
]
|
||||
elif language == 'zh':
|
||||
self.apps = ['图像扩展', '虚拟试衣', '图像补全(文本引导)', '图像补全(文本+参考图引导)']
|
||||
self.dropdown_name = '应用'
|
||||
self.subject_image = '参考图片'
|
||||
self.subject_mask = '参考图掩码'
|
||||
self.scene_image = '背景图片'
|
||||
self.scene_mask = '背景图掩码'
|
||||
self.prompt = '提示文本'
|
||||
self.masked_image = '掩码图片'
|
||||
self.preprocess = '输入图片预处理'
|
||||
self.button_name = '数据预处理'
|
||||
self.direction = ('使用说明:\n'
|
||||
'针对自定义数据:\n'
|
||||
'a.1)上传背景图片(待编辑);\n'
|
||||
'a.2)上传参考图片(如需要);\n'
|
||||
'a.3)使用笔刷功能涂抹图像中的特定位置(获取对应的掩码图像);\n'
|
||||
'a.4)点击数据预处理按钮;\n'
|
||||
'a.5)输入文本提示;\n'
|
||||
'针对提供的样例数据: \n'
|
||||
'b.1)点击样例数据 \n'
|
||||
'最后,点击生成按钮获取生成图片')
|
||||
self.out_direction_label = '扩展方向'
|
||||
self.out_directions = [
|
||||
'中心向外',
|
||||
'右下',
|
||||
'左下',
|
||||
'右上',
|
||||
'左上',
|
||||
]
|
||||
|
||||
self.examples = [
|
||||
[
|
||||
self.apps[0],
|
||||
'a temple on fire',
|
||||
download_image(
|
||||
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex1_scene_im.png' # noqa
|
||||
),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
1.0,
|
||||
0.75,
|
||||
'CenterAround',
|
||||
1024,
|
||||
1024
|
||||
],
|
||||
[
|
||||
self.apps[1],
|
||||
'',
|
||||
download_image(
|
||||
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex2_scene_im.jpg' # noqa
|
||||
),
|
||||
download_image(
|
||||
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex2_scene_mask.png' # noqa
|
||||
),
|
||||
download_image(
|
||||
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex2_subject_im.jpg' # noqa
|
||||
),
|
||||
download_image(
|
||||
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex2_subject_mask.jpg' # noqa
|
||||
),
|
||||
1.0,
|
||||
0.0,
|
||||
'',
|
||||
1024,
|
||||
1024
|
||||
],
|
||||
[
|
||||
self.apps[2],
|
||||
'a blue and white porcelain',
|
||||
download_image(
|
||||
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex3_scene_im.png' # noqa
|
||||
),
|
||||
download_image(
|
||||
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex3_scene_mask.png' # noqa
|
||||
),
|
||||
None,
|
||||
None,
|
||||
1.0,
|
||||
0.0,
|
||||
'',
|
||||
1024,
|
||||
1024
|
||||
],
|
||||
[
|
||||
self.apps[3],
|
||||
'a dog wearing sunglasses',
|
||||
download_image(
|
||||
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex4_scene_im.png' # noqa
|
||||
),
|
||||
download_image(
|
||||
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex4_scene_mask.png' # noqa
|
||||
),
|
||||
download_image(
|
||||
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex4_subject_im.png' # noqa
|
||||
),
|
||||
download_image(
|
||||
'https://modelscope.cn/api/v1/models/iic/LARGEN/repo?Revision=master&FilePath=examples/ex4_subject_mask.png' # noqa
|
||||
),
|
||||
0.45,
|
||||
0.0,
|
||||
'',
|
||||
1024,
|
||||
1024
|
||||
],
|
||||
]
|
||||
|
||||
@@ -6,6 +6,7 @@ import gradio as gr
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.studio.inference.inference_ui.component_names import GalleryUIName
|
||||
from scepter.studio.utils.uibase import UIBase
|
||||
|
||||
@@ -15,6 +16,9 @@ class GalleryUI(UIBase):
|
||||
self.pipe_manager = pipe_manager
|
||||
self.component_names = GalleryUIName(language)
|
||||
self.cfg = cfg
|
||||
self.work_dir = cfg.WORK_DIR
|
||||
self.local_work_dir, _ = FS.map_to_local(self.work_dir)
|
||||
os.makedirs(self.local_work_dir, exist_ok=True)
|
||||
|
||||
def create_ui(self, *args, **kwargs):
|
||||
with gr.Group():
|
||||
@@ -41,6 +45,7 @@ class GalleryUI(UIBase):
|
||||
container=False,
|
||||
autofocus=True,
|
||||
elem_classes='type_row',
|
||||
submit_on_enter=True,
|
||||
lines=1)
|
||||
|
||||
with gr.Column(scale=3, min_width=0):
|
||||
@@ -87,6 +92,19 @@ class GalleryUI(UIBase):
|
||||
style_template,
|
||||
style_negative_template,
|
||||
image_seed,
|
||||
largen_state,
|
||||
largen_task,
|
||||
largen_image_scale,
|
||||
largen_tar_image,
|
||||
largen_tar_mask,
|
||||
largen_masked_image,
|
||||
largen_ref_image,
|
||||
largen_ref_mask,
|
||||
largen_ref_clip,
|
||||
largen_base_image,
|
||||
largen_extra_sizes,
|
||||
largen_bbox_yyxx,
|
||||
largen_history,
|
||||
show_jpeg_image=True):
|
||||
if control_state and control_cond_image is None:
|
||||
raise gr.Error(self.component_names.control_err1)
|
||||
@@ -111,9 +129,8 @@ class GalleryUI(UIBase):
|
||||
for tuner_m in tuner_model:
|
||||
if tuner_m is None or tuner_m == '':
|
||||
continue
|
||||
if (now_pipeline in self.pipe_manager.model_level_info['tuners']
|
||||
and tuner_m in self.pipe_manager.model_level_info['tuners']
|
||||
[now_pipeline]):
|
||||
if now_pipeline in self.pipe_manager.model_level_info['tuners'] and \
|
||||
tuner_m in self.pipe_manager.model_level_info['tuners'][now_pipeline]:
|
||||
tuner_m = self.pipe_manager.model_level_info['tuners'][
|
||||
now_pipeline][tuner_m]['model_info']
|
||||
used_tuner_model.append(tuner_m)
|
||||
@@ -163,6 +180,23 @@ class GalleryUI(UIBase):
|
||||
pipeline_input['refine_guide_rescale'] = refine_guide_rescale
|
||||
else:
|
||||
refine_strength = 0
|
||||
if largen_state:
|
||||
largen_cfg = {
|
||||
'largen_task': largen_task,
|
||||
'largen_image_scale': largen_image_scale,
|
||||
'largen_tar_image': largen_tar_image,
|
||||
'largen_tar_mask': largen_tar_mask,
|
||||
'largen_ref_image': largen_ref_image,
|
||||
'largen_ref_mask': largen_ref_mask,
|
||||
'largen_masked_image': largen_masked_image,
|
||||
'largen_ref_clip': largen_ref_clip,
|
||||
'largen_base_image': largen_base_image,
|
||||
'largen_extra_sizes': largen_extra_sizes,
|
||||
'largen_bbox_yyxx': largen_bbox_yyxx,
|
||||
}
|
||||
else:
|
||||
largen_cfg = {}
|
||||
|
||||
results = current_pipeline(
|
||||
pipeline_input,
|
||||
num_samples=image_number,
|
||||
@@ -177,7 +211,8 @@ class GalleryUI(UIBase):
|
||||
if tuner_state or control_state else None,
|
||||
control_cond_image=control_cond_image if control_state else None,
|
||||
crop_type=crop_type if control_state else None,
|
||||
seed=int(image_seed))
|
||||
seed=int(image_seed),
|
||||
**largen_cfg)
|
||||
images = []
|
||||
before_images = []
|
||||
if 'images' in results:
|
||||
@@ -198,27 +233,34 @@ class GalleryUI(UIBase):
|
||||
if 'seed' in results:
|
||||
print(results['seed'])
|
||||
print(images, before_images)
|
||||
largen_history.extend(images)
|
||||
if len(largen_history) > 10:
|
||||
largen_history = largen_history[-10:]
|
||||
if show_jpeg_image:
|
||||
save_list = []
|
||||
for i, img in enumerate(images):
|
||||
save_image = os.path.join(self.cfg.WORK_DIR,
|
||||
save_image = os.path.join(self.local_work_dir,
|
||||
f'cur_gallery_{i}.jpg')
|
||||
img.save(save_image)
|
||||
save_list.append(save_image)
|
||||
images = save_list
|
||||
|
||||
return (
|
||||
gr.Column(visible=len(before_images) > 0),
|
||||
before_images,
|
||||
images,
|
||||
largen_history,
|
||||
gr.update(value=largen_history),
|
||||
)
|
||||
|
||||
def generate_image(self, *args, **kwargs):
|
||||
gallery_result = self.generate_gallery(*args, **kwargs)
|
||||
before_refine_panel, before_refine_gallery, output_gallery = gallery_result
|
||||
before_refine_panel, before_refine_gallery, output_gallery, _ = gallery_result
|
||||
return (before_refine_panel, before_refine_gallery, output_gallery[0])
|
||||
|
||||
def set_callbacks(self, inference_ui, model_manage_ui, diffusion_ui,
|
||||
mantra_ui, tuner_ui, refiner_ui, control_ui, **kwargs):
|
||||
mantra_ui, tuner_ui, refiner_ui, control_ui, largen_ui,
|
||||
**kwargs):
|
||||
|
||||
self.gen_inputs = [
|
||||
self.prompt, mantra_ui.state, tuner_ui.state, control_ui.state,
|
||||
@@ -237,12 +279,20 @@ class GalleryUI(UIBase):
|
||||
refiner_ui.refine_strength, refiner_ui.refine_sampler,
|
||||
refiner_ui.refine_discretization, refiner_ui.refine_guide_scale,
|
||||
refiner_ui.refine_guide_rescale, mantra_ui.style_template,
|
||||
mantra_ui.style_negative_template, diffusion_ui.image_seed
|
||||
mantra_ui.style_negative_template, diffusion_ui.image_seed,
|
||||
largen_ui.state, largen_ui.task, largen_ui.image_scale,
|
||||
largen_ui.tar_image, largen_ui.tar_mask, largen_ui.masked_image,
|
||||
largen_ui.ref_image, largen_ui.ref_mask, largen_ui.ref_clip,
|
||||
largen_ui.base_image, largen_ui.extra_sizes, largen_ui.bbox_yyxx,
|
||||
largen_ui.image_history
|
||||
]
|
||||
|
||||
self.gen_outputs = [
|
||||
self.before_refine_panel, self.before_refine_gallery,
|
||||
self.output_gallery
|
||||
self.before_refine_panel,
|
||||
self.before_refine_gallery,
|
||||
self.output_gallery,
|
||||
largen_ui.image_history,
|
||||
largen_ui.gallery,
|
||||
]
|
||||
|
||||
self.generate_button.click(self.generate_gallery,
|
||||
|
||||
@@ -0,0 +1,462 @@
|
||||
# -*- 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, :]
|
||||
H1, W1 = crop_tar_image.shape[:2]
|
||||
|
||||
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)
|
||||
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, :]
|
||||
|
||||
h, w = crop_ref_mask_i.shape[:2]
|
||||
ref_expand_size = int(max(h, w) * 1.02)
|
||||
pad_op = A.PadIfNeeded(ref_expand_size, ref_expand_size,
|
||||
border_mode=cv2.BORDER_CONSTANT,
|
||||
value=(255, 255, 255), mask_value=0)
|
||||
out = pad_op(image=crop_ref_image_i, mask=crop_ref_mask_i)
|
||||
crop_ref_image = out['image']
|
||||
|
||||
to_clip_input = T.Compose([
|
||||
T.ToTensor(),
|
||||
T.Resize((224, 224)),
|
||||
T.Normalize(mean=(0.48145466, 0.4578275, 0.40821073), std=(0.26862954, 0.26130258, 0.27577711)),
|
||||
])
|
||||
ref_clip = to_clip_input(crop_ref_image.astype(np.uint8))
|
||||
|
||||
output_size = max(output_height, output_width)
|
||||
ref_resize_op = A.Compose([
|
||||
A.LongestMaxSize(output_size),
|
||||
A.PadIfNeeded(output_size, output_size,
|
||||
border_mode=cv2.BORDER_CONSTANT,
|
||||
value=(255, 255, 255), mask_value=0),
|
||||
])
|
||||
aug_out = ref_resize_op(image=crop_ref_image_i.astype(np.uint8), mask=crop_ref_mask_i)
|
||||
aug_ref_image = aug_out['image']
|
||||
aug_ref_mask = aug_out['mask']
|
||||
|
||||
final_ref_image = TF.to_tensor(aug_ref_image)
|
||||
final_ref_image = TF.normalize(final_ref_image, mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
|
||||
final_ref_mask = TF.to_tensor(aug_ref_mask)
|
||||
|
||||
final_ref_image = final_ref_image.unsqueeze(0)
|
||||
final_ref_mask = final_ref_mask.unsqueeze(0)
|
||||
ref_clip = ref_clip.unsqueeze(0)
|
||||
else:
|
||||
final_ref_image = None
|
||||
final_ref_mask = None
|
||||
ref_clip = None
|
||||
|
||||
return final_tar_image, final_tar_mask, masked_image, final_ref_image, final_ref_mask, ref_clip, \
|
||||
TF.to_tensor(tar_image), torch.LongTensor([H1, W1, H2, W2, pad1, pad2]), torch.LongTensor(tar_yyxx_crop)
|
||||
|
||||
def data_preprocess_outpaint(self,
|
||||
tar_image,
|
||||
direction,
|
||||
img_ratio,
|
||||
output_height,
|
||||
output_width):
|
||||
oh, ow = output_height, output_width
|
||||
h, w = tar_image.shape[:2]
|
||||
ratio = max(h/(oh*img_ratio), w/(ow*img_ratio))
|
||||
|
||||
ih, iw = int(h / ratio), int(w / ratio)
|
||||
|
||||
masked_image = np.zeros((oh, ow, 3), dtype=np.uint8)
|
||||
mask = np.zeros((oh, ow, 1))
|
||||
|
||||
if direction in ['CenterAround', '中心向外']:
|
||||
y1, x1 = (oh-ih)//2, (ow-iw)//2
|
||||
elif direction in ['RightDown', '右下']:
|
||||
y1, x1 = 0, 0
|
||||
elif direction in ['LeftDown', '左下']:
|
||||
y1, x1 = 0, ow-iw
|
||||
elif direction in ['RightUp', '右上']:
|
||||
y1, x1 = oh-ih, 0
|
||||
elif direction in ['LeftUp', '左上']:
|
||||
y1, x1 = oh-ih, ow-iw
|
||||
else:
|
||||
y1, x1 = 0, 0
|
||||
|
||||
tar_image = cv2.resize(tar_image.astype(np.uint8), (iw, ih))
|
||||
masked_image[y1:y1+ih, x1:x1+iw] = tar_image
|
||||
mask[y1+5:y1+ih-5, x1+5:x1+iw-5] = 1
|
||||
|
||||
final_tar_image = TF.to_tensor(masked_image.astype(np.uint8))
|
||||
final_tar_image = TF.normalize(final_tar_image, mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
|
||||
final_tar_mask = TF.to_tensor(((1.0-mask) > 0.5).astype(np.float32))
|
||||
|
||||
masked_image = final_tar_image.clone()
|
||||
masked_image = masked_image * (1 - final_tar_mask)
|
||||
|
||||
final_tar_image = final_tar_image.unsqueeze(0)
|
||||
final_tar_mask = final_tar_mask.unsqueeze(0)
|
||||
masked_image = masked_image.unsqueeze(0)
|
||||
|
||||
return final_tar_image, final_tar_mask, masked_image, None, None, None, None, None, None
|
||||
@@ -47,14 +47,15 @@ class MantraUI(UIBase):
|
||||
|
||||
def create_ui(self, *args, **kwargs):
|
||||
self.state = gr.State(value=False)
|
||||
with gr.Column(visible=False) as self.tab:
|
||||
with gr.Row():
|
||||
with gr.Column(equal_height=True, visible=False) as self.tab:
|
||||
with gr.Row(scale=1):
|
||||
with gr.Column(scale=1):
|
||||
with gr.Group(visible=True):
|
||||
with gr.Row(equal_height=True):
|
||||
self.style = gr.Dropdown(
|
||||
label=self.component_names.mantra_styles,
|
||||
choices=self.all_styles[self.default_pipeline],
|
||||
choices=self.all_styles.get(
|
||||
self.default_pipeline, []),
|
||||
value=None,
|
||||
multiselect=True,
|
||||
interactive=True)
|
||||
|
||||
@@ -146,11 +146,20 @@ class ModelManageUI(UIBase):
|
||||
continue
|
||||
model_name = f"{now_pipeline}_{module['name']}"
|
||||
all_module_name[module_name] = model_name
|
||||
tunner_choices = []
|
||||
if now_pipeline in self.default_choices['tuners']:
|
||||
tunner_choices = self.default_choices['tuners'][now_pipeline][
|
||||
'choices']
|
||||
else:
|
||||
tunner_choices = []
|
||||
custom_tunner_choices = []
|
||||
custom_tunner_default = []
|
||||
if now_pipeline in self.default_choices.get(
|
||||
'customized_tuners', []):
|
||||
custom_tunner_choices = self.default_choices[
|
||||
'customized_tuners'][now_pipeline]['choices']
|
||||
custom_tunner_default = self.default_choices[
|
||||
'customized_tuners'][now_pipeline]['default']
|
||||
if isinstance(custom_tunner_default, str):
|
||||
custom_tunner_default = [custom_tunner_default]
|
||||
|
||||
if now_pipeline in self.default_choices[
|
||||
'controllers'] and control_mode in self.default_choices[
|
||||
@@ -179,9 +188,11 @@ class ModelManageUI(UIBase):
|
||||
gr.Dropdown(value=all_module_name['first_stage_model']),
|
||||
gr.Dropdown(value=all_module_name['cond_stage_model']),
|
||||
gr.Dropdown(choices=tunner_choices, value=[]),
|
||||
gr.Dropdown(choices=custom_tunner_choices,
|
||||
value=custom_tunner_default),
|
||||
gr.Dropdown(choices=controller_choices,
|
||||
value=controller_default),
|
||||
gr.Dropdown(choices=mantra_ui.all_styles[now_pipeline],
|
||||
gr.Dropdown(choices=mantra_ui.all_styles.get(now_pipeline, []),
|
||||
value=[]),
|
||||
gr.Textbox(choices=cur_paras.NEGATIVE_PROMPT.get('VALUES', []),
|
||||
value=cur_paras.NEGATIVE_PROMPT.get('DEFAULT', '')),
|
||||
@@ -206,10 +217,11 @@ class ModelManageUI(UIBase):
|
||||
outputs=[
|
||||
self.diffusion_state, self.first_stage_model,
|
||||
self.cond_stage_model, tuner_ui.tuner_model,
|
||||
control_ui.control_model, mantra_ui.style,
|
||||
diffusion_ui.negative_prompt, diffusion_ui.prompt_prefix,
|
||||
diffusion_ui.output_height, diffusion_ui.sampler,
|
||||
diffusion_ui.discretization, diffusion_ui.sample_steps,
|
||||
diffusion_ui.guide_scale, diffusion_ui.guide_rescale
|
||||
tuner_ui.custom_tuner_model, control_ui.control_model,
|
||||
mantra_ui.style, diffusion_ui.negative_prompt,
|
||||
diffusion_ui.prompt_prefix, diffusion_ui.output_height,
|
||||
diffusion_ui.sampler, diffusion_ui.discretization,
|
||||
diffusion_ui.sample_steps, diffusion_ui.guide_scale,
|
||||
diffusion_ui.guide_rescale
|
||||
],
|
||||
queue=True)
|
||||
|
||||
@@ -21,17 +21,21 @@ class TunerUI(UIBase):
|
||||
'default']
|
||||
self.default_pipeline = pipe_manager.model_level_info[
|
||||
default_diffusion_model]['pipeline'][0]
|
||||
self.tunner_choices = []
|
||||
if self.default_pipeline in self.default_choices['tuners']:
|
||||
self.tunner_choices = self.default_choices['tuners'][
|
||||
self.default_pipeline]['choices']
|
||||
self.tunner_default = self.default_choices['tuners'][
|
||||
self.default_pipeline]['default']
|
||||
else:
|
||||
self.tunner_choices = []
|
||||
self.custom_tuner_choices = []
|
||||
if self.default_pipeline in self.default_choices.get(
|
||||
'customized_tuners', []):
|
||||
self.custom_tuner_choices = self.default_choices[
|
||||
'customized_tuners'][self.default_pipeline]['choices']
|
||||
|
||||
self.tunner_default = None
|
||||
self.component_names = TunerUIName(language)
|
||||
self.cfg_tuners = cfg.TUNERS
|
||||
self.cfg_tuners = cfg.TUNERS + cfg.CUSTOM_TUNERS
|
||||
self.name_level_tuners = {}
|
||||
for one_tuner in tqdm(self.cfg_tuners):
|
||||
if one_tuner.BASE_MODEL not in self.name_level_tuners:
|
||||
@@ -47,8 +51,8 @@ class TunerUI(UIBase):
|
||||
|
||||
def create_ui(self, *args, **kwargs):
|
||||
self.state = gr.State(value=False)
|
||||
with gr.Column(visible=False) as self.tab:
|
||||
with gr.Row():
|
||||
with gr.Column(equal_height=True, visible=False) as self.tab:
|
||||
with gr.Row(scale=1):
|
||||
with gr.Column(variant='panel', scale=1, min_width=0):
|
||||
with gr.Group(visible=True):
|
||||
with gr.Row(equal_height=True):
|
||||
@@ -63,10 +67,16 @@ class TunerUI(UIBase):
|
||||
self.custom_tuner_model = gr.Dropdown(
|
||||
label=self.component_names.
|
||||
custom_tuner_model,
|
||||
choices=[],
|
||||
choices=self.custom_tuner_choices,
|
||||
value=None,
|
||||
multiselect=True,
|
||||
interactive=True)
|
||||
self.save_button = gr.Button(
|
||||
label=self.component_names.save_button,
|
||||
value=self.component_names.save_button,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button',
|
||||
visible=True)
|
||||
with gr.Row(equal_height=True):
|
||||
with gr.Column(scale=1):
|
||||
self.tuner_type = gr.Text(
|
||||
@@ -110,6 +120,7 @@ class TunerUI(UIBase):
|
||||
label=self.component_names.example_block_name, open=True)
|
||||
|
||||
def set_callbacks(self, model_manage_ui, **kwargs):
|
||||
manager = kwargs.pop('manager')
|
||||
gallery_ui = kwargs.pop('gallery_ui')
|
||||
with self.example_block:
|
||||
gr.Examples(examples=self.component_names.examples,
|
||||
@@ -121,7 +132,7 @@ class TunerUI(UIBase):
|
||||
now_pipeline = diffusion_model_info['pipeline'][0]
|
||||
tuner_info = {}
|
||||
if tuner_model is not None and len(tuner_model) > 0:
|
||||
tuner_info = self.name_level_tuners[now_pipeline].get(
|
||||
tuner_info = self.name_level_tuners.get(now_pipeline, {}).get(
|
||||
tuner_model[-1], {})
|
||||
if tuner_info.get(
|
||||
'IMAGE_PATH',
|
||||
@@ -141,3 +152,45 @@ class TunerUI(UIBase):
|
||||
self.tuner_example, self.tuner_prompt_example
|
||||
],
|
||||
queue=False)
|
||||
|
||||
self.custom_tuner_model.change(
|
||||
tuner_model_change,
|
||||
inputs=[self.custom_tuner_model, model_manage_ui.diffusion_model],
|
||||
outputs=[
|
||||
self.tuner_type, self.base_model, self.tuner_desc,
|
||||
self.tuner_example, self.tuner_prompt_example
|
||||
],
|
||||
queue=False)
|
||||
|
||||
def save_customized_tuner(tuner_model, diffusion_model):
|
||||
diffusion_model_info = self.pipe_manager.model_level_info[
|
||||
diffusion_model]
|
||||
now_pipeline = diffusion_model_info['pipeline'][0]
|
||||
tuner_info = {}
|
||||
if tuner_model is not None and len(tuner_model) > 0:
|
||||
tuner_info = self.name_level_tuners.get(now_pipeline, {}).get(
|
||||
tuner_model[-1], {})
|
||||
if tuner_info.get(
|
||||
'IMAGE_PATH',
|
||||
None) and not os.path.exists(tuner_info.IMAGE_PATH):
|
||||
tuner_info.IMAGE_PATH = FS.get_from(tuner_info.IMAGE_PATH)
|
||||
return (gr.Tabs(selected='tuner_manager'),
|
||||
gr.Text(value=tuner_info.NAME), gr.Text(value=''),
|
||||
gr.Text(value=tuner_info.get('TUNER_TYPE', '')),
|
||||
gr.Text(value=tuner_info.get('BASE_MODEL', '')),
|
||||
gr.Text(value=tuner_info.get('DESCRIPTION', '')),
|
||||
gr.Image(value=tuner_info.get('IMAGE_PATH', None)),
|
||||
gr.Text(value=tuner_info.get('PROMPT_EXAMPLE', '')))
|
||||
|
||||
self.save_button.click(
|
||||
save_customized_tuner,
|
||||
inputs=[self.custom_tuner_model, model_manage_ui.diffusion_model],
|
||||
outputs=[
|
||||
manager.tabs, manager.tuner_manager.info_ui.tuner_name,
|
||||
manager.tuner_manager.info_ui.new_name,
|
||||
manager.tuner_manager.info_ui.tuner_type,
|
||||
manager.tuner_manager.info_ui.base_model,
|
||||
manager.tuner_manager.info_ui.tuner_desc,
|
||||
manager.tuner_manager.info_ui.tuner_example,
|
||||
manager.tuner_manager.info_ui.tuner_prompt_example
|
||||
])
|
||||
|
||||
@@ -24,7 +24,6 @@ refresh_symbol = '\U0001f504' # 🔄
|
||||
class CreateDatasetUI(UIBase):
|
||||
def __init__(self, cfg, is_debug=False, language='en'):
|
||||
self.work_dir = cfg.WORK_DIR
|
||||
self.dir_list = FS.walk_dir(self.work_dir, recurse=False)
|
||||
self.cache_file = {}
|
||||
self.meta_dict = {}
|
||||
self.dataset_list = self.load_history()
|
||||
@@ -71,7 +70,8 @@ class CreateDatasetUI(UIBase):
|
||||
json.dump(save_meta, open(local_path, 'w'))
|
||||
return meta_file
|
||||
|
||||
def construct_meta(self, cursor, file_list, dataset_folder, user_name):
|
||||
def construct_meta(self, cursor, file_list, dataset_folder, user_name,
|
||||
login_user_name):
|
||||
'''
|
||||
{
|
||||
"dataset_name": "xxxx",
|
||||
@@ -86,11 +86,17 @@ class CreateDatasetUI(UIBase):
|
||||
save_file_list = os.path.join(dataset_folder, 'file.csv')
|
||||
save_file_list = self.write_file_list(file_list, save_file_list)
|
||||
meta = {
|
||||
'dataset_name': user_name,
|
||||
'cursor': cursor,
|
||||
'file_list': file_list,
|
||||
'train_csv': train_csv,
|
||||
'save_file_list': save_file_list
|
||||
'dataset_name':
|
||||
user_name if login_user_name == '' or login_user_name is None else
|
||||
'_'.join([login_user_name, user_name]),
|
||||
'cursor':
|
||||
cursor,
|
||||
'file_list':
|
||||
file_list,
|
||||
'train_csv':
|
||||
train_csv,
|
||||
'save_file_list':
|
||||
save_file_list
|
||||
}
|
||||
self.save_meta(meta, dataset_folder)
|
||||
return meta
|
||||
@@ -180,7 +186,7 @@ class CreateDatasetUI(UIBase):
|
||||
file_folder = None
|
||||
train_list = None
|
||||
hit_dir = None
|
||||
raw_list = []
|
||||
raw_list = {}
|
||||
mac_osx = os.path.join(local_dataset_folder, '__MACOSX')
|
||||
if os.path.exists(mac_osx):
|
||||
res = os.popen(f"rm -rf '{mac_osx}'")
|
||||
@@ -191,28 +197,45 @@ class CreateDatasetUI(UIBase):
|
||||
res = res.readlines()
|
||||
continue
|
||||
if FS.isdir(one_dir):
|
||||
sub_dir = FS.walk_dir(one_dir)
|
||||
for one_s_dir in sub_dir:
|
||||
if FS.isdir(one_s_dir) and one_s_dir.split(
|
||||
one_dir)[1].replace('/', '') == 'images':
|
||||
file_folder = one_s_dir
|
||||
hit_dir = one_dir
|
||||
if FS.isfile(one_s_dir) and one_s_dir.split(
|
||||
one_dir)[1].replace('/', '') == 'train.csv':
|
||||
train_list = one_s_dir
|
||||
if file_folder is not None and train_list is not None:
|
||||
break
|
||||
if (one_s_dir.endswith('.jpg')
|
||||
or one_s_dir.endswith('.jpeg')
|
||||
or one_s_dir.endswith('.png')
|
||||
or one_s_dir.endswith('.webp')):
|
||||
raw_list.append(one_s_dir)
|
||||
if one_dir.endswith('images') or one_dir.endswith('images/'):
|
||||
file_folder = one_dir
|
||||
hit_dir = one_dir
|
||||
else:
|
||||
sub_dir = FS.walk_dir(one_dir)
|
||||
for one_s_dir in sub_dir:
|
||||
if FS.isdir(one_s_dir) and one_s_dir.split(
|
||||
one_dir)[1].replace('/', '') == 'images':
|
||||
file_folder = one_s_dir
|
||||
hit_dir = one_dir
|
||||
if FS.isfile(one_s_dir) and one_s_dir.split(
|
||||
one_dir)[1].replace('/', '') == 'train.csv':
|
||||
train_list = one_s_dir
|
||||
if file_folder is not None and train_list is not None:
|
||||
break
|
||||
if (one_s_dir.endswith('.jpg')
|
||||
or one_s_dir.endswith('.jpeg')
|
||||
or one_s_dir.endswith('.png')
|
||||
or one_s_dir.endswith('.webp')):
|
||||
file_name, surfix = os.path.splitext(one_s_dir)
|
||||
txt_file = file_name + '.txt'
|
||||
if os.path.exists(txt_file):
|
||||
raw_list[one_s_dir] = txt_file
|
||||
else:
|
||||
raw_list[one_s_dir] = None
|
||||
elif one_dir.endswith('train.csv'):
|
||||
train_list = one_dir
|
||||
else:
|
||||
if (one_dir.endswith('.jpg') or one_dir.endswith('.jpeg')
|
||||
or one_dir.endswith('.png')
|
||||
or one_dir.endswith('.webp')):
|
||||
raw_list.append(one_dir)
|
||||
|
||||
file_name, surfix = os.path.splitext(one_dir)
|
||||
txt_file = file_name + '.txt'
|
||||
if os.path.exists(txt_file):
|
||||
raw_list[one_dir] = txt_file
|
||||
else:
|
||||
raw_list[one_dir] = None
|
||||
if file_folder is not None and train_list is not None:
|
||||
break
|
||||
if file_folder is None and len(raw_list) < 1:
|
||||
raise gr.Error(
|
||||
"images folder or train.csv doesn't exists, or nothing exists in your zip"
|
||||
@@ -222,17 +245,21 @@ class CreateDatasetUI(UIBase):
|
||||
if file_folder is not None:
|
||||
_ = FS.get_dir_to_local_dir(file_folder, new_file_folder)
|
||||
elif len(raw_list) > 0:
|
||||
raw_list = list(set(raw_list))
|
||||
raw_list = [[k, v] for k, v in raw_list.items()]
|
||||
for img_id, cur_image in enumerate(raw_list):
|
||||
_, surfix = os.path.splitext(cur_image)
|
||||
image_name, surfix = os.path.splitext(cur_image[0])
|
||||
if cur_image[1] is not None and os.path.exists(cur_image[1]):
|
||||
prompt = open(cur_image[1], 'r').read()
|
||||
else:
|
||||
prompt = image_name.split('/')[-1]
|
||||
try:
|
||||
os.rename(
|
||||
os.path.abspath(cur_image),
|
||||
f'{new_file_folder}/{get_md5(cur_image)}{surfix}')
|
||||
os.path.abspath(cur_image[0]),
|
||||
f'{new_file_folder}/{get_md5(cur_image[0])}{surfix}')
|
||||
raw_list[img_id] = [
|
||||
os.path.join('images',
|
||||
f'{get_md5(cur_image)}{surfix}'),
|
||||
cur_image.split('/')[-1]
|
||||
f'{get_md5(cur_image[0])}{surfix}'),
|
||||
prompt
|
||||
]
|
||||
except Exception as e:
|
||||
print(e)
|
||||
@@ -251,13 +278,14 @@ class CreateDatasetUI(UIBase):
|
||||
res = res.readlines()
|
||||
if not os.path.exists(new_train_list):
|
||||
raise gr.Error(f'{str(res)}')
|
||||
try:
|
||||
res = os.popen(f"rm -rf '{hit_dir}/images/*'")
|
||||
_ = res.readlines()
|
||||
res = os.popen(f"rm -rf '{hit_dir}'")
|
||||
_ = res.readlines()
|
||||
except Exception:
|
||||
pass
|
||||
if not file_folder == hit_dir:
|
||||
try:
|
||||
res = os.popen(f"rm -rf '{hit_dir}/images/*'")
|
||||
_ = res.readlines()
|
||||
res = os.popen(f"rm -rf '{hit_dir}'")
|
||||
_ = res.readlines()
|
||||
except Exception:
|
||||
pass
|
||||
file_list = self.load_train_csv(new_train_list, data_folder)
|
||||
return file_list
|
||||
|
||||
@@ -317,21 +345,25 @@ class CreateDatasetUI(UIBase):
|
||||
return True, file_path
|
||||
return False, file_path
|
||||
|
||||
def load_history(self):
|
||||
def load_history(self, login_user_name=''):
|
||||
dataset_list = []
|
||||
self.dir_list = FS.walk_dir(self.work_dir, recurse=False)
|
||||
for one_dir in self.dir_list:
|
||||
if FS.isdir(one_dir):
|
||||
meta_file = os.path.join(one_dir, 'meta.json')
|
||||
if FS.exists(meta_file):
|
||||
local_dataset_folder, _ = FS.map_to_local(one_dir)
|
||||
local_dataset_folder = FS.get_dir_to_local_dir(
|
||||
one_dir, local_dataset_folder, multi_thread=True)
|
||||
if not FS.exists(
|
||||
os.path.join(local_dataset_folder, 'meta.json')):
|
||||
local_dataset_folder = FS.get_dir_to_local_dir(
|
||||
one_dir, local_dataset_folder, multi_thread=True)
|
||||
meta_data = self.load_meta(
|
||||
os.path.join(local_dataset_folder, 'meta.json'))
|
||||
meta_data['local_work_dir'] = local_dataset_folder
|
||||
meta_data['work_dir'] = one_dir
|
||||
dataset_list.append(meta_data['dataset_name'])
|
||||
self.meta_dict[meta_data['dataset_name']] = meta_data
|
||||
if meta_data['dataset_name'].startswith(login_user_name):
|
||||
dataset_list.append(meta_data['dataset_name'])
|
||||
self.meta_dict[meta_data['dataset_name']] = meta_data
|
||||
return dataset_list
|
||||
|
||||
def create_ui(self):
|
||||
@@ -396,7 +428,7 @@ class CreateDatasetUI(UIBase):
|
||||
self.file_panel = file_panel
|
||||
self.modify_panel = modify_panel
|
||||
|
||||
def set_callbacks(self, gallery_dataset, export_dataset):
|
||||
def set_callbacks(self, gallery_dataset, export_dataset, manager):
|
||||
def show_dataset_panel():
|
||||
return (gr.Column(visible=False), gr.Column(visible=True),
|
||||
gr.Column(visible=True),
|
||||
@@ -419,17 +451,19 @@ class CreateDatasetUI(UIBase):
|
||||
datetime.datetime.now())
|
||||
return data_name
|
||||
|
||||
def refresh():
|
||||
return gr.Dropdown(value=self.dataset_list[-1]
|
||||
if len(self.dataset_list) > 0 else '',
|
||||
choices=self.dataset_list)
|
||||
def refresh(login_user_name):
|
||||
dataset_list = self.load_history(login_user_name=login_user_name)
|
||||
return gr.Dropdown(
|
||||
value=dataset_list[-1] if len(dataset_list) > 0 else '',
|
||||
choices=dataset_list)
|
||||
|
||||
self.refresh_dataset_name.click(refresh,
|
||||
inputs=[manager.user_name],
|
||||
outputs=[self.dataset_name],
|
||||
queue=False)
|
||||
|
||||
def confirm_create_dataset(user_name, create_mode, file_url, file_path,
|
||||
panel_state):
|
||||
panel_state, login_user_name):
|
||||
if user_name.strip() == '' or ' ' in user_name or '/' in user_name:
|
||||
raise gr.Error(self.components_name.illegal_data_name_err1)
|
||||
|
||||
@@ -442,7 +476,9 @@ class CreateDatasetUI(UIBase):
|
||||
if not file_url.strip() == '' and file_path is not None:
|
||||
raise gr.Error(self.components_name.illegal_data_name_err4)
|
||||
if create_mode == 3 and not file_url.strip() == '':
|
||||
file_name, surfix = os.path.splitext(file_url.split('?')[0])
|
||||
if 'oss' in file_url:
|
||||
file_url = file_url.split('?')[0]
|
||||
file_name, surfix = os.path.splitext(file_url)
|
||||
save_file = os.path.join(self.work_dir, f'{user_name}{surfix}')
|
||||
local_path, _ = FS.map_to_local(save_file)
|
||||
res = os.popen(f"wget -c '{file_url}' -O '{local_path}'")
|
||||
@@ -490,7 +526,7 @@ class CreateDatasetUI(UIBase):
|
||||
|
||||
cursor = 0 if len(file_list) > 0 else -1
|
||||
meta = self.construct_meta(cursor, file_list, dataset_folder,
|
||||
user_name)
|
||||
user_name, login_user_name)
|
||||
|
||||
meta['local_work_dir'] = local_dataset_folder
|
||||
meta['work_dir'] = dataset_folder
|
||||
@@ -498,13 +534,13 @@ class CreateDatasetUI(UIBase):
|
||||
self.meta_dict[meta['dataset_name']] = meta
|
||||
if meta['dataset_name'] not in self.dataset_list:
|
||||
self.dataset_list.append(meta['dataset_name'])
|
||||
return (
|
||||
gr.Checkbox(value=True, visible=False),
|
||||
gr.Dropdown(value=user_name, choices=self.dataset_list),
|
||||
)
|
||||
return (gr.Checkbox(value=True, visible=False),
|
||||
gr.Dropdown(value=meta['dataset_name'],
|
||||
choices=self.dataset_list),
|
||||
gr.Text(value=meta['dataset_name']))
|
||||
|
||||
def clear_file():
|
||||
return gr.Text(visible=True)
|
||||
return gr.Text(visible=False)
|
||||
|
||||
# Click Create
|
||||
self.btn_create_datasets.click(show_dataset_panel, [], [
|
||||
@@ -545,9 +581,9 @@ class CreateDatasetUI(UIBase):
|
||||
# Click Confirm
|
||||
self.confirm_data_button.click(confirm_create_dataset, [
|
||||
self.user_data_name, self.create_mode, self.file_path_url,
|
||||
self.file_path, self.panel_state
|
||||
], [self.panel_state, self.dataset_name],
|
||||
queue=True)
|
||||
self.file_path, self.panel_state, manager.user_name
|
||||
], [self.panel_state, self.dataset_name, self.user_data_name],
|
||||
queue=False)
|
||||
|
||||
def show_edit_panel(panel_state, data_name):
|
||||
if panel_state:
|
||||
@@ -568,7 +604,7 @@ class CreateDatasetUI(UIBase):
|
||||
],
|
||||
queue=False)
|
||||
|
||||
def modify_data_name(user_name, prev_data_name):
|
||||
def modify_data_name(user_name, prev_data_name, login_user_name):
|
||||
print(
|
||||
f'Current file name {prev_data_name}, new file name {user_name}.'
|
||||
)
|
||||
@@ -600,7 +636,8 @@ class CreateDatasetUI(UIBase):
|
||||
raise gr.Error(self.components_name.illegal_data_err3)
|
||||
cursor = ori_meta['cursor']
|
||||
meta = self.construct_meta(cursor, file_list,
|
||||
dataset_folder, user_name)
|
||||
dataset_folder, user_name,
|
||||
login_user_name)
|
||||
meta['local_work_dir'] = local_dataset_folder
|
||||
meta['work_dir'] = dataset_folder
|
||||
|
||||
@@ -624,7 +661,10 @@ class CreateDatasetUI(UIBase):
|
||||
|
||||
self.modify_data_button.click(
|
||||
modify_data_name,
|
||||
inputs=[self.user_data_name, self.user_data_name_state],
|
||||
inputs=[
|
||||
self.user_data_name, self.user_data_name_state,
|
||||
manager.user_name
|
||||
],
|
||||
outputs=[self.user_data_name_state, self.dataset_name],
|
||||
queue=False)
|
||||
|
||||
@@ -647,3 +687,14 @@ class CreateDatasetUI(UIBase):
|
||||
gallery_dataset.gallery_state
|
||||
],
|
||||
queue=False)
|
||||
|
||||
def login_user_name_change(login_user_name):
|
||||
dataset_list = self.load_history(login_user_name=login_user_name)
|
||||
return gr.Dropdown(
|
||||
value=dataset_list[-1] if len(dataset_list) > 0 else '',
|
||||
choices=dataset_list)
|
||||
|
||||
manager.user_name.change(login_user_name_change,
|
||||
inputs=[manager.user_name],
|
||||
outputs=[self.dataset_name],
|
||||
queue=False)
|
||||
|
||||
@@ -46,7 +46,7 @@ class PreprocessUI():
|
||||
|
||||
def set_callbacks(self, manager):
|
||||
self.create_dataset.set_callbacks(self.dataset_gallery,
|
||||
self.export_dataset)
|
||||
self.export_dataset, manager)
|
||||
self.dataset_gallery.set_callbacks(self.create_dataset)
|
||||
self.export_dataset.set_callbacks(self.create_dataset, manager)
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ import numpy as np
|
||||
import torch
|
||||
|
||||
from scepter.modules.solver.hooks.checkpoint import CheckpointHook
|
||||
from scepter.modules.solver.hooks.data_probe import ProbeDataHook
|
||||
from scepter.modules.solver.registry import SOLVERS
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.distribute import we
|
||||
@@ -27,6 +28,7 @@ def run_task(cfg):
|
||||
solver.set_up_pre()
|
||||
solver.set_up()
|
||||
ori_steps = solver.max_steps
|
||||
|
||||
if 'train' in solver.datas:
|
||||
dataset = solver.datas['train'].dataset
|
||||
if hasattr(dataset, 'real_number'):
|
||||
@@ -48,7 +50,20 @@ 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}')
|
||||
# size 为无限的时候,使用默认值。
|
||||
|
||||
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
|
||||
|
||||
solver.solve()
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,7 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import time
|
||||
|
||||
for i in range(180):
|
||||
time.sleep(1)
|
||||
print('sleep', i)
|
||||
@@ -0,0 +1,337 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
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)
|
||||
@@ -7,7 +7,7 @@ import gradio as gr
|
||||
import scepter
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.studio.self_train.self_train_ui.inference_ui import InferenceUI
|
||||
from scepter.studio.self_train.self_train_ui.model_ui import ModelUI
|
||||
from scepter.studio.self_train.self_train_ui.trainer_ui import TrainerUI
|
||||
from scepter.studio.self_train.utils.config_parser import get_all_config
|
||||
from scepter.studio.utils.env import init_env
|
||||
@@ -34,20 +34,20 @@ class SelfTrainUI():
|
||||
BASE_CFG_VALUE,
|
||||
is_debug=is_debug,
|
||||
language=language)
|
||||
self.inference_ui = InferenceUI(cfg_general,
|
||||
BASE_CFG_VALUE,
|
||||
is_debug=is_debug,
|
||||
language=language)
|
||||
self.model_ui = ModelUI(cfg_general,
|
||||
BASE_CFG_VALUE,
|
||||
is_debug=is_debug,
|
||||
language=language)
|
||||
|
||||
def create_ui(self):
|
||||
with gr.Row():
|
||||
self.trainer_ui.create_ui()
|
||||
with gr.Row():
|
||||
self.inference_ui.create_ui()
|
||||
self.model_ui.create_ui()
|
||||
|
||||
def set_callbacks(self, manager):
|
||||
self.trainer_ui.set_callbacks(self.inference_ui)
|
||||
self.inference_ui.set_callbacks(self.trainer_ui, manager)
|
||||
self.trainer_ui.set_callbacks(self.model_ui, manager)
|
||||
self.model_ui.set_callbacks(self.trainer_ui, manager)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
@@ -1,11 +1,12 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
# For dataset manager
|
||||
class InferenceUIName():
|
||||
class ModelUIName():
|
||||
def __init__(self, language='en'):
|
||||
if language == 'en':
|
||||
self.output_model_block = 'Model Output'
|
||||
self.output_model_name = 'Output Model Name'
|
||||
self.output_ckpt_name = 'Output Ckpt Name'
|
||||
self.test_prompt = 'Test Prompt'
|
||||
self.test_prefix = 'Test Prefix'
|
||||
self.test_n_prompt = 'Negative Prompt'
|
||||
@@ -21,15 +22,23 @@ class InferenceUIName():
|
||||
self.extra_model_gbtn = 'Add Model'
|
||||
self.refresh_model_gbtn = 'Refresh Model'
|
||||
self.go_to_inference = 'Go to inference'
|
||||
self.btn_export_log = 'Export Log'
|
||||
self.export_file = 'Log File'
|
||||
self.log_block = 'Training Log...'
|
||||
self.gallery_block = 'Gallery Log...'
|
||||
self.eval_gallery = 'Eval Gallery'
|
||||
# Error or Warning
|
||||
self.inference_err1 = 'Inference failed, please try again.'
|
||||
self.inference_err2 = 'Test prompt is empty.'
|
||||
self.inference_err3 = "Doesn't surpport this base model"
|
||||
self.inference_err4 = "This model maybe not finish training, because model doesn't exist."
|
||||
self.model_err3 = "Doesn't surpport this base model"
|
||||
self.model_err4 = "This model maybe not finish training, because model doesn't exist."
|
||||
self.model_err5 = "Model {} doesn't exist."
|
||||
self.training_warn1 = 'No log message util now.'
|
||||
|
||||
elif language == 'zh':
|
||||
self.output_model_block = '模型产出'
|
||||
self.output_model_name = '产出模型名称'
|
||||
self.output_model_name = '产出名称'
|
||||
self.output_ckpt_name = '产出检查点名称'
|
||||
self.test_prompt = '测试提示词'
|
||||
self.test_prefix = '测试前缀'
|
||||
self.test_n_prompt = '负向提示词'
|
||||
@@ -44,12 +53,20 @@ class InferenceUIName():
|
||||
self.extra_model_gtxt = '额外模型'
|
||||
self.extra_model_gbtn = '添加模型'
|
||||
self.refresh_model_gbtn = '刷新模型'
|
||||
self.btn_export_log = '导出日志'
|
||||
self.export_file = '日志文件'
|
||||
self.log_block = '训练日志...'
|
||||
self.training_button = '开始训练'
|
||||
self.gallery_block = '图像日志...'
|
||||
self.eval_gallery = '评测图像'
|
||||
# Error or Warning
|
||||
self.inference_err1 = '推理失败,请重试。'
|
||||
self.inference_err2 = '测试提示词为空。'
|
||||
self.inference_err3 = '不支持的基础模型'
|
||||
self.model_err3 = '不支持的基础模型'
|
||||
self.go_to_inference = '使用模型'
|
||||
self.inference_err4 = '模型可能没有训练完成或者模型不存在'
|
||||
self.model_err4 = '模型可能没有训练完成或者模型不存在'
|
||||
self.model_err5 = '模型{}不存在'
|
||||
self.training_warn1 = '暂时没有日志文件;任务启动中或失败!'
|
||||
|
||||
|
||||
class TrainerUIName():
|
||||
@@ -80,7 +97,8 @@ class TrainerUIName():
|
||||
self.base_model = 'Base Model'
|
||||
self.tuner_name = 'Fine-tuning Method'
|
||||
self.base_model_revision = 'Model Version Number'
|
||||
self.resolution = 'Resolution'
|
||||
self.resolution_height = 'Resolution Height'
|
||||
self.resolution_width = 'Resolution Width'
|
||||
self.train_epoch = 'Number of Training Epochs'
|
||||
self.learning_rate = 'Learning Rate'
|
||||
self.save_interval = 'Save Interval'
|
||||
@@ -89,9 +107,8 @@ class TrainerUIName():
|
||||
self.replace_keywords = 'Trigger Keywords'
|
||||
self.work_name = 'Save Model Name (refresh to get a random value)'
|
||||
self.push_to_hub = 'Push to hub'
|
||||
self.log_block = 'Training Log...'
|
||||
self.training_button = 'Start Training'
|
||||
|
||||
self.eval_prompts = 'Eval Prompts'
|
||||
# Error or Warning
|
||||
self.training_err1 = 'CUDA is unavailable.'
|
||||
self.training_err2 = 'Currently insufficient VRAM, training failed!'
|
||||
@@ -121,7 +138,8 @@ class TrainerUIName():
|
||||
self.base_model = '基础模型'
|
||||
self.tuner_name = '微调方法'
|
||||
self.base_model_revision = '模型版本号'
|
||||
self.resolution = '分辨率'
|
||||
self.resolution_height = '训练高度'
|
||||
self.resolution_width = '训练宽度'
|
||||
self.train_epoch = '训练轮数'
|
||||
self.learning_rate = '学习率'
|
||||
self.save_interval = '存储间隔'
|
||||
@@ -130,8 +148,7 @@ class TrainerUIName():
|
||||
self.replace_keywords = '触发关键词'
|
||||
self.work_name = '保存模型名称(刷新获得随机值)'
|
||||
self.push_to_hub = '推送魔搭社区'
|
||||
self.log_block = '训练日志...'
|
||||
self.training_button = '开始训练'
|
||||
self.eval_prompts = '评测文本'
|
||||
# Error or Warning
|
||||
self.training_err1 = 'CUDA不可用.'
|
||||
self.training_err2 = '目前显存不足,训练失败!'
|
||||
|
||||
@@ -1,145 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import os
|
||||
|
||||
import gradio as gr
|
||||
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.studio.self_train.self_train_ui.component_names import \
|
||||
InferenceUIName
|
||||
from scepter.studio.self_train.utils.config_parser import (
|
||||
get_base_model_list, get_inference_para_by_model_version)
|
||||
from scepter.studio.utils.uibase import UIBase
|
||||
|
||||
|
||||
class InferenceUI(UIBase):
|
||||
def __init__(self, cfg, all_cfg_value, is_debug=False, language='en'):
|
||||
self.BASE_CFG_VALUE = all_cfg_value
|
||||
self.language = language
|
||||
self.base_model_info = get_base_model_list(self.BASE_CFG_VALUE)
|
||||
self.work_dir, _ = FS.map_to_local(cfg.WORK_DIR)
|
||||
os.makedirs(self.work_dir, exist_ok=True)
|
||||
self.model_list = []
|
||||
# self.model_list.extend(self.base_model_info.get('model_choices', []))
|
||||
have_model_list = []
|
||||
if not self.work_dir.endswith('/'):
|
||||
self.work_dir += '/'
|
||||
for one_dir in FS.walk_dir(self.work_dir):
|
||||
if one_dir.startswith(self.work_dir):
|
||||
one_dir = one_dir[len(self.work_dir):]
|
||||
if not os.path.isdir(os.path.join(self.work_dir, one_dir)):
|
||||
continue
|
||||
if len(one_dir.split('/')) > 1:
|
||||
continue
|
||||
if '@' in one_dir and os.path.exists(
|
||||
os.path.join(self.work_dir, one_dir, 'checkpoint.pth')):
|
||||
if len(one_dir.split('@')) > 4:
|
||||
have_model_list.append([one_dir, one_dir.split('@')[-1]])
|
||||
have_model_list.sort(key=lambda x: -int(x[-1][:-3]))
|
||||
self.model_list.extend([v[0].split('/')[-1]
|
||||
for v in have_model_list][:50])
|
||||
self.infer_para_data = get_inference_para_by_model_version(
|
||||
self.BASE_CFG_VALUE, self.base_model_info.get('model_name', []),
|
||||
self.base_model_info.get('version_name', []))
|
||||
self.is_debug = is_debug
|
||||
self.component_names = InferenceUIName(language)
|
||||
|
||||
def create_ui(self, *args, **kwargs):
|
||||
with gr.Box():
|
||||
gr.Markdown(self.component_names.output_model_block)
|
||||
with gr.Row():
|
||||
with gr.Column(scale=2, min_width=0):
|
||||
self.output_model_name = gr.Dropdown(
|
||||
label=self.component_names.output_model_name,
|
||||
choices=self.model_list,
|
||||
value=self.base_model_info.get('model_default', ''),
|
||||
interactive=True)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.refresh_model_gbtn = gr.Button(
|
||||
self.component_names.refresh_model_gbtn)
|
||||
with gr.Row():
|
||||
with gr.Column(scale=2, min_width=0):
|
||||
self.extra_model_gtxt = gr.Text(
|
||||
label=self.component_names.extra_model_gtxt,
|
||||
show_label=False,
|
||||
placeholder='Add Extra Model')
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.extra_model_gbtn = gr.Button(
|
||||
self.component_names.extra_model_gbtn)
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.go_to_inferece_btn = gr.Button(
|
||||
self.component_names.go_to_inference)
|
||||
|
||||
def set_callbacks(self, trainer_ui, manager):
|
||||
self.manager = manager
|
||||
|
||||
def add_model(model_name):
|
||||
if model_name not in self.model_list:
|
||||
self.model_list.append(model_name)
|
||||
return '', gr.Dropdown(choices=self.model_list)
|
||||
|
||||
def refresh_model():
|
||||
return gr.Dropdown(choices=self.model_list)
|
||||
|
||||
self.extra_model_gbtn.click(
|
||||
fn=add_model,
|
||||
inputs=[self.extra_model_gtxt],
|
||||
outputs=[self.extra_model_gtxt, self.output_model_name],
|
||||
queue=False)
|
||||
self.refresh_model_gbtn.click(fn=refresh_model,
|
||||
inputs=[],
|
||||
outputs=[self.output_model_name],
|
||||
queue=False)
|
||||
|
||||
def go_to_inferece(output_model):
|
||||
output_model_path = os.path.join(self.work_dir, output_model)
|
||||
_, _, base_model, _, resolution, _ = output_model_path.split('@')
|
||||
tuner_cfg = Config(cfg_dict={}, load=False)
|
||||
tuner_cfg.NAME = output_model
|
||||
tuner_cfg.NAME_ZH = output_model
|
||||
tuner_cfg.BASE_MODEL = base_model
|
||||
model_path = os.path.join(output_model_path, 'checkpoint.pth')
|
||||
if not os.path.exists(model_path):
|
||||
gr.Error(self.component_names.inference_err4)
|
||||
tuner_cfg.MODEL_PATH = model_path
|
||||
self.manager.inference.model_manage_ui.pipe_manager.register_tuner(
|
||||
tuner_cfg,
|
||||
name=tuner_cfg.NAME_ZH
|
||||
if self.language == 'zh' else tuner_cfg.NAME,
|
||||
is_customized=True)
|
||||
|
||||
pipeline_level_modules = self.manager.inference.model_manage_ui.pipe_manager.pipeline_level_modules
|
||||
if tuner_cfg.BASE_MODEL not in pipeline_level_modules:
|
||||
gr.Error(self.component_names.inference_err3 +
|
||||
tuner_cfg.BASE_MODEL)
|
||||
pipeline_ins = pipeline_level_modules[tuner_cfg.BASE_MODEL]
|
||||
diffusion_model = f"{tuner_cfg.BASE_MODEL}_{pipeline_ins.diffusion_model['name']}"
|
||||
|
||||
default_choices = self.manager.inference.model_manage_ui.pipe_manager.module_level_choices
|
||||
if 'customized_tuners' in default_choices and tuner_cfg.BASE_MODEL in default_choices[
|
||||
'customized_tuners']:
|
||||
tunner_choices = default_choices['customized_tuners'][
|
||||
tuner_cfg.BASE_MODEL]['choices']
|
||||
tunner_default = default_choices['customized_tuners'][
|
||||
tuner_cfg.BASE_MODEL]['default']
|
||||
if not isinstance(tunner_default, list):
|
||||
tunner_default = [tunner_default]
|
||||
else:
|
||||
tunner_choices = []
|
||||
tunner_default = ''
|
||||
return (gr.Tabs(selected='inference'),
|
||||
gr.Dropdown(choices=tunner_choices, value=tunner_default),
|
||||
gr.Dropdown(value=diffusion_model),
|
||||
gr.Tabs(selected='tuner_ui'))
|
||||
|
||||
self.go_to_inferece_btn.click(
|
||||
go_to_inferece,
|
||||
inputs=[self.output_model_name],
|
||||
outputs=[
|
||||
manager.tabs, manager.inference.tuner_ui.custom_tuner_model,
|
||||
manager.inference.model_manage_ui.diffusion_model,
|
||||
manager.inference.setting_tab
|
||||
],
|
||||
queue=False)
|
||||
@@ -0,0 +1,459 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import json
|
||||
import os
|
||||
import queue
|
||||
import time
|
||||
|
||||
import gradio as gr
|
||||
import yaml
|
||||
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.studio.self_train.self_train_ui.component_names import ModelUIName
|
||||
from scepter.studio.self_train.utils.config_parser import get_base_model_list
|
||||
from scepter.studio.utils.uibase import UIBase
|
||||
|
||||
refresh_symbol = '\U0001f504' # 🔄
|
||||
delete_symbol = '\U0001f5d1' # 🗑️
|
||||
add_symbol = '\U00002795' # ➕
|
||||
confirm_symbol = '\U00002714' # ✔️
|
||||
|
||||
|
||||
class ModelUI(UIBase):
|
||||
def __init__(self, cfg, all_cfg_value, is_debug=False, language='en'):
|
||||
self.BASE_CFG_VALUE = all_cfg_value
|
||||
self.language = language
|
||||
self.base_model_info = get_base_model_list(self.BASE_CFG_VALUE)
|
||||
self.work_dir, _ = FS.map_to_local(cfg.WORK_DIR)
|
||||
os.makedirs(self.work_dir, exist_ok=True)
|
||||
self.default_model_name = 'step-last'
|
||||
self.old_default_model_name = 'checkpoint.pth'
|
||||
self.delete_folder_queue = queue.Queue()
|
||||
self.model_list = []
|
||||
# self.model_list.extend(self.base_model_info.get('model_choices', []))
|
||||
have_model_list = []
|
||||
if not self.work_dir.endswith('/'):
|
||||
self.work_dir += '/'
|
||||
for one_dir in FS.walk_dir(self.work_dir):
|
||||
if one_dir.startswith(self.work_dir):
|
||||
one_dir = one_dir[len(self.work_dir):]
|
||||
if not os.path.isdir(os.path.join(self.work_dir, one_dir)):
|
||||
continue
|
||||
if len(one_dir.split('/')) > 1:
|
||||
continue
|
||||
if ('@' in one_dir
|
||||
and (os.path.exists(
|
||||
os.path.join(self.work_dir, one_dir, 'checkpoints',
|
||||
self.default_model_name))
|
||||
or os.path.exists(
|
||||
os.path.join(self.work_dir, one_dir,
|
||||
self.old_default_model_name)))):
|
||||
if len(one_dir.split('@')) > 4:
|
||||
have_model_list.append([one_dir, one_dir.split('@')[-1]])
|
||||
have_model_list.sort(key=lambda x: -int(x[-1][:-3]))
|
||||
self.model_list.extend([v[0].split('/')[-1]
|
||||
for v in have_model_list][:50])
|
||||
self.is_debug = is_debug
|
||||
self.component_names = ModelUIName(language)
|
||||
|
||||
def get_ckpt_list(self, output_model):
|
||||
all_ckpt_list = []
|
||||
if output_model is None or output_model == '':
|
||||
return all_ckpt_list
|
||||
output_model_path = os.path.join(self.work_dir, output_model)
|
||||
all_ckpt_path = os.path.join(output_model_path, 'checkpoints')
|
||||
if os.path.exists(all_ckpt_path):
|
||||
for name in os.listdir(all_ckpt_path):
|
||||
if name == self.default_model_name:
|
||||
continue
|
||||
path = os.path.join(all_ckpt_path, name)
|
||||
if os.path.isdir(path):
|
||||
all_ckpt_list.append(name)
|
||||
all_ckpt_list = sorted(all_ckpt_list,
|
||||
key=lambda x: int(x.split('-')[-1]))
|
||||
if os.path.exists(os.path.join(all_ckpt_path,
|
||||
self.default_model_name)):
|
||||
all_ckpt_list.append(self.default_model_name)
|
||||
return all_ckpt_list
|
||||
|
||||
def get_gallery_list(self, model_name, ckpt_name):
|
||||
ckpt_probe_dir = os.path.join(self.work_dir, model_name, 'eval_probe',
|
||||
ckpt_name, 'image')
|
||||
all_gallery_list = []
|
||||
if os.path.exists(ckpt_probe_dir):
|
||||
for name in os.listdir(ckpt_probe_dir):
|
||||
path = os.path.join(ckpt_probe_dir, name)
|
||||
all_gallery_list.append(path)
|
||||
return all_gallery_list
|
||||
|
||||
def create_ui(self, *args, **kwargs):
|
||||
with gr.Box():
|
||||
gr.Markdown(self.component_names.output_model_block)
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
with gr.Column(scale=7, min_width=0):
|
||||
self.output_model_name = gr.Dropdown(
|
||||
label=self.component_names.output_model_name,
|
||||
choices=self.model_list,
|
||||
value=self.base_model_info.get('model_default', ''),
|
||||
show_label=False,
|
||||
container=False,
|
||||
interactive=True)
|
||||
with gr.Column(scale=2, min_width=0):
|
||||
self.output_ckpt_name = gr.Dropdown(
|
||||
label=self.component_names.output_ckpt_name,
|
||||
value='',
|
||||
choices=[],
|
||||
show_label=False,
|
||||
container=False,
|
||||
interactive=True)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.refresh_model_gbtn = gr.Button(value=refresh_symbol)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.add_model_gbtn = gr.Button(value=add_symbol)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.delete_model_gbtn = gr.Button(value=delete_symbol)
|
||||
|
||||
with gr.Row(variant='panel', equal_height=True,
|
||||
visible=False) as self.extra_model_panel:
|
||||
with gr.Column(scale=4, min_width=0):
|
||||
self.extra_model_txt = gr.Text(
|
||||
label=self.component_names.extra_model_gtxt,
|
||||
show_label=False,
|
||||
container=False,
|
||||
placeholder='Add Extra Model')
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.confirm_add = gr.Button(value=confirm_symbol)
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
with gr.Column(scale=3, min_width=0):
|
||||
with gr.Box():
|
||||
self.log_message = gr.Text(
|
||||
placeholder='Please select model'
|
||||
'or press export button.',
|
||||
autoscroll=True,
|
||||
lines=10,
|
||||
label=self.component_names.log_block)
|
||||
with gr.Column(scale=1, min_width=0,
|
||||
visible=False) as self.export_log_panel:
|
||||
# with gr.Column(scale=1, min_width=0):
|
||||
self.export_log = gr.Button(
|
||||
value=self.component_names.btn_export_log)
|
||||
# with gr.Column(scale=1, min_width=0):
|
||||
self.export_url = gr.File(
|
||||
label=self.component_names.export_file,
|
||||
visible=False,
|
||||
value=None,
|
||||
interactive=False,
|
||||
show_label=True)
|
||||
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
with gr.Accordion(label=self.component_names.gallery_block,
|
||||
open=True):
|
||||
self.eval_gallery = gr.Gallery(
|
||||
label=self.component_names.eval_gallery,
|
||||
value=[],
|
||||
preview=True,
|
||||
selected_index=None)
|
||||
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.go_to_inferece_btn = gr.Button(
|
||||
self.component_names.go_to_inference)
|
||||
|
||||
def set_callbacks(self, trainer_ui, manager):
|
||||
self.manager = manager
|
||||
|
||||
def model_name_change(model_name):
|
||||
if model_name is None:
|
||||
return '', gr.Column(), '', []
|
||||
message = trainer_ui.trainer_ins.get_log(model_name)
|
||||
status = trainer_ui.trainer_ins.get_status(model_name)
|
||||
ckpt_list = self.get_ckpt_list(model_name)
|
||||
ckpt_value = ckpt_list[-1] if len(ckpt_list) > 0 else ''
|
||||
return (message, gr.Column(visible=status in ('running',
|
||||
'success')),
|
||||
gr.Dropdown(choices=ckpt_list, value=ckpt_value))
|
||||
|
||||
self.output_model_name.change(fn=model_name_change,
|
||||
inputs=[self.output_model_name],
|
||||
outputs=[
|
||||
self.log_message,
|
||||
self.export_log_panel,
|
||||
self.output_ckpt_name
|
||||
],
|
||||
queue=False)
|
||||
|
||||
def ckpt_name_change(model_name, ckpt_value):
|
||||
if ckpt_value is not None and len(ckpt_value) > 0:
|
||||
gallery_value = self.get_gallery_list(model_name, ckpt_value)
|
||||
else:
|
||||
gallery_value = []
|
||||
if len(gallery_value) > 0:
|
||||
return gr.Gallery(value=gallery_value,
|
||||
preview=True,
|
||||
selected_index=0)
|
||||
else:
|
||||
return gr.Gallery(value=gallery_value,
|
||||
preview=True,
|
||||
selected_index=None)
|
||||
|
||||
self.output_ckpt_name.change(
|
||||
fn=ckpt_name_change,
|
||||
inputs=[self.output_model_name, self.output_ckpt_name],
|
||||
outputs=[self.eval_gallery],
|
||||
queue=False)
|
||||
|
||||
def add_model():
|
||||
return gr.Row(visible=True)
|
||||
|
||||
self.add_model_gbtn.click(fn=add_model,
|
||||
inputs=[],
|
||||
outputs=[self.extra_model_panel],
|
||||
queue=False)
|
||||
|
||||
def confirm_add(model_name):
|
||||
model_folder = os.path.join(self.work_dir, model_name)
|
||||
have_model = os.path.exists(model_folder)
|
||||
if not have_model:
|
||||
gr.Error(self.component_names.model_err5.format(model_name))
|
||||
if model_name not in self.model_list and have_model:
|
||||
self.model_list.append(model_name)
|
||||
return gr.Row(visible=False), gr.Dropdown(choices=self.model_list,
|
||||
value=model_name)
|
||||
|
||||
self.confirm_add.click(
|
||||
fn=confirm_add,
|
||||
inputs=[self.extra_model_txt],
|
||||
outputs=[self.extra_model_panel, self.output_model_name],
|
||||
queue=False)
|
||||
|
||||
def refresh_model(model_name):
|
||||
if not self.delete_folder_queue.empty():
|
||||
try:
|
||||
del_folder = self.delete_folder_queue.get_nowait()
|
||||
if os.path.exists(del_folder):
|
||||
os.system(f'rm -rf {del_folder}')
|
||||
self.delete_folder_queue.put_nowait(del_folder)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
message = trainer_ui.trainer_ins.get_log(model_name)
|
||||
status = trainer_ui.trainer_ins.get_status(model_name)
|
||||
ckpt_list = self.get_ckpt_list(model_name)
|
||||
ckpt_value = ckpt_list[-1] if len(ckpt_list) > 0 else ''
|
||||
ret_gallery = ckpt_name_change(model_name, ckpt_value)
|
||||
return (message, gr.Column(visible=status in ('running',
|
||||
'success')),
|
||||
gr.Dropdown(choices=ckpt_list,
|
||||
value=ckpt_value), ret_gallery)
|
||||
|
||||
self.refresh_model_gbtn.click(
|
||||
fn=refresh_model,
|
||||
inputs=[self.output_model_name],
|
||||
outputs=[
|
||||
self.log_message,
|
||||
self.export_log_panel,
|
||||
# self.output_model_name,
|
||||
self.output_ckpt_name,
|
||||
self.eval_gallery
|
||||
],
|
||||
queue=False)
|
||||
|
||||
def delete_model(model_name):
|
||||
index = 0
|
||||
trainer_ui.trainer_ins.stop_task(model_name)
|
||||
if model_name in self.model_list:
|
||||
index = self.model_list.index(model_name)
|
||||
self.model_list.remove(model_name)
|
||||
folder = os.path.join(self.work_dir, model_name)
|
||||
for _ in range(1):
|
||||
if os.path.exists(folder):
|
||||
try:
|
||||
os.system(f'rm -rf {folder}')
|
||||
except Exception:
|
||||
time.sleep(2)
|
||||
else:
|
||||
break
|
||||
self.delete_folder_queue.put_nowait(folder)
|
||||
if index <= len(self.model_list) - 1:
|
||||
model_name = self.model_list[index]
|
||||
elif len(self.model_list) > 0:
|
||||
model_name = self.model_list[0]
|
||||
else:
|
||||
model_name = None
|
||||
|
||||
return gr.Dropdown(choices=self.model_list, value=model_name)
|
||||
|
||||
self.delete_model_gbtn.click(fn=delete_model,
|
||||
inputs=[self.output_model_name],
|
||||
outputs=[self.output_model_name])
|
||||
|
||||
def go_to_inferece(output_model, output_ckpt_name):
|
||||
params_path = os.path.join(self.work_dir, output_model,
|
||||
'params.json')
|
||||
if os.path.exists(params_path):
|
||||
params_info = json.loads(open(params_path).read())
|
||||
assert params_info['work_name'] == output_model
|
||||
# base_model = params_info['base_model']
|
||||
base_model_revision = params_info['base_model_revision']
|
||||
tuner_name = params_info['tuner_name']
|
||||
model_path = os.path.join(self.work_dir, output_model,
|
||||
'checkpoints', output_ckpt_name)
|
||||
eval_prompts = params_info[
|
||||
'eval_prompts'] if 'eval_prompts' in params_info else []
|
||||
image_dir = os.path.join(self.work_dir, output_model,
|
||||
'eval_probe', output_ckpt_name,
|
||||
'image')
|
||||
image_path = [
|
||||
os.path.join(image_dir, name)
|
||||
for name in os.listdir(image_dir)
|
||||
]
|
||||
else:
|
||||
_, base_model, base_model_revision, tuner_name, _ = output_model.split(
|
||||
'@', 4)
|
||||
model_path = os.path.join(self.work_dir, output_model,
|
||||
self.old_default_model_name)
|
||||
output_ckpt_name = self.old_default_model_name
|
||||
image_path = []
|
||||
eval_prompts = []
|
||||
params_info = {}
|
||||
|
||||
if isinstance(image_path, list):
|
||||
image_path = image_path[0] if len(image_path) > 0 else None
|
||||
if isinstance(eval_prompts, list):
|
||||
eval_prompts = eval_prompts[0] if len(eval_prompts) > 0 else ''
|
||||
cfg_file = os.path.join(self.work_dir, output_model,
|
||||
f'meta_{output_ckpt_name}.yaml')
|
||||
output_model = output_model + '@' + output_ckpt_name
|
||||
tuner_dict = {
|
||||
'NAME': output_model,
|
||||
'NAME_ZH': output_model,
|
||||
# 'BASE_MODEL': base_model,
|
||||
'BASE_MODEL': base_model_revision,
|
||||
'TUNER_TYPE': tuner_name,
|
||||
'DESCRIPTION': '',
|
||||
'MODEL_PATH': model_path,
|
||||
'IMAGE_PATH': image_path,
|
||||
'PROMPT_EXAMPLE': eval_prompts,
|
||||
'SOURCE': 'self_train',
|
||||
'CKPT_NAME': output_ckpt_name,
|
||||
'PARAMS': params_info
|
||||
}
|
||||
tuner_cfg = Config(cfg_dict=tuner_dict, load=False)
|
||||
|
||||
if not os.path.exists(model_path):
|
||||
gr.Error(self.component_names.model_err4)
|
||||
self.manager.inference.model_manage_ui.pipe_manager.register_tuner(
|
||||
tuner_cfg,
|
||||
name=tuner_cfg.NAME_ZH
|
||||
if self.language == 'zh' else tuner_cfg.NAME,
|
||||
is_customized=True)
|
||||
|
||||
pipeline_level_modules = self.manager.inference.model_manage_ui.pipe_manager.pipeline_level_modules
|
||||
if tuner_cfg.BASE_MODEL not in pipeline_level_modules:
|
||||
gr.Error(self.component_names.model_err3 +
|
||||
tuner_cfg.BASE_MODEL)
|
||||
pipeline_ins = pipeline_level_modules[tuner_cfg.BASE_MODEL]
|
||||
diffusion_model = f"{tuner_cfg.BASE_MODEL}_{pipeline_ins.diffusion_model['name']}"
|
||||
|
||||
default_choices = self.manager.inference.model_manage_ui.pipe_manager.module_level_choices
|
||||
if 'customized_tuners' in default_choices:
|
||||
if tuner_cfg.BASE_MODEL not in default_choices[
|
||||
'customized_tuners']:
|
||||
default_choices['customized_tuners'] = {}
|
||||
tunner_choices = default_choices['customized_tuners'][
|
||||
tuner_cfg.BASE_MODEL]['choices']
|
||||
tunner_default = default_choices['customized_tuners'][
|
||||
tuner_cfg.BASE_MODEL]['default']
|
||||
if not isinstance(tunner_default, list):
|
||||
tunner_default = [tunner_default]
|
||||
else:
|
||||
tunner_choices = []
|
||||
tunner_default = []
|
||||
|
||||
with open(cfg_file, 'w') as f_out:
|
||||
yaml.dump(copy.deepcopy(tuner_cfg.cfg_dict),
|
||||
f_out,
|
||||
encoding='utf-8',
|
||||
allow_unicode=True,
|
||||
default_flow_style=False)
|
||||
|
||||
# example_image = tuner_cfg.get('IMAGE_PATH', None)
|
||||
# if isinstance(example_image, list) and len(example_image)>0:
|
||||
# example_image = example_image[0]
|
||||
#
|
||||
# prompt_example = tuner_cfg.get('PROMPT_EXAMPLE', None)
|
||||
# if isinstance(prompt_example, list) and len(prompt_example)>0:
|
||||
# prompt_example = prompt_example[0]
|
||||
|
||||
base_model = tuner_cfg.get('BASE_MODEL', '')
|
||||
|
||||
if not base_model == '':
|
||||
if base_model not in manager.inference.tuner_ui.name_level_tuners:
|
||||
manager.inference.tuner_ui.name_level_tuners[
|
||||
base_model] = {}
|
||||
manager.inference.tuner_ui.name_level_tuners[base_model][
|
||||
output_model] = tuner_cfg
|
||||
|
||||
return (
|
||||
gr.Tabs(selected='inference'), cfg_file,
|
||||
gr.Tabs(selected='tuner_ui'),
|
||||
gr.CheckboxGroup(
|
||||
value='使用微调' if self.language == 'zh' else 'Use Tuners'),
|
||||
gr.Dropdown(value=diffusion_model),
|
||||
gr.Dropdown(choices=tunner_choices, value=tunner_default)
|
||||
# gr.Text(value=tuner_cfg.get('TUNER_TYPE', '')),
|
||||
# gr.Text(value=tuner_cfg.get('BASE_MODEL', '')),
|
||||
# gr.Image(value=example_image),
|
||||
# gr.Text(value=tuner_cfg.get('DESCRIPTION', '')),
|
||||
# gr.Text(value=prompt_example)
|
||||
)
|
||||
|
||||
self.go_to_inferece_btn.click(
|
||||
go_to_inferece,
|
||||
inputs=[self.output_model_name, self.output_ckpt_name],
|
||||
outputs=[
|
||||
manager.tabs, manager.inference.infer_info,
|
||||
manager.inference.setting_tab,
|
||||
manager.inference.check_box_for_setting,
|
||||
manager.inference.model_manage_ui.diffusion_model,
|
||||
manager.inference.tuner_ui.custom_tuner_model
|
||||
# manager.inference.tuner_ui.tuner_type,
|
||||
# manager.inference.tuner_ui.base_model,
|
||||
# manager.inference.tuner_ui.tuner_example,
|
||||
# manager.inference.tuner_ui.tuner_desc,
|
||||
# manager.inference.tuner_ui.tuner_prompt_example
|
||||
],
|
||||
queue=False)
|
||||
|
||||
def export_train_log(model_name):
|
||||
current_log_folder = os.path.join(self.work_dir, model_name)
|
||||
_ = trainer_ui.trainer_ins.get_log(model_name)
|
||||
out_log = os.path.join(current_log_folder, 'output_std_log.txt')
|
||||
if os.path.exists(out_log):
|
||||
zip_path = f'{trainer_ui.current_train_model}_log.zip'
|
||||
cache_folder = f'{trainer_ui.current_train_model}_log'
|
||||
if not os.path.exists(
|
||||
os.path.join(current_log_folder, 'tensorboard')):
|
||||
os.makedirs(os.path.join(current_log_folder,
|
||||
'tensorboard'),
|
||||
exist_ok=True)
|
||||
res = os.popen(
|
||||
f"cd {current_log_folder} && mkdir -p '{cache_folder}' "
|
||||
f"&& cp -rf tensorboard '{cache_folder}/tensorboard' "
|
||||
f"&& cp -rf 'output_std_log.txt' '{cache_folder}/output_std_log.txt' "
|
||||
f"&& zip -r '{zip_path}' '{cache_folder}'/* "
|
||||
f"&& rm -rf '{cache_folder}'")
|
||||
print(res.readlines())
|
||||
return gr.File(value=os.path.join(current_log_folder,
|
||||
zip_path),
|
||||
visible=True)
|
||||
else:
|
||||
gr.Error(self.component_names.training_warn1)
|
||||
return gr.File()
|
||||
|
||||
self.export_log.click(export_train_log,
|
||||
inputs=[self.output_model_name],
|
||||
outputs=[self.export_url],
|
||||
queue=False)
|
||||
@@ -2,9 +2,10 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import datetime
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
|
||||
import gradio as gr
|
||||
import torch
|
||||
@@ -12,12 +13,12 @@ import yaml
|
||||
|
||||
import scepter
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.studio.self_train.scripts.trainer import TrainManager
|
||||
from scepter.studio.self_train.self_train_ui.component_names import \
|
||||
TrainerUIName
|
||||
from scepter.studio.self_train.utils.config_parser import (
|
||||
get_default, get_values_by_model, get_values_by_model_version,
|
||||
get_values_by_model_version_tuner,
|
||||
get_values_by_model_version_tuner_resolution)
|
||||
get_values_by_model_version_tuner)
|
||||
from scepter.studio.utils.uibase import UIBase
|
||||
|
||||
|
||||
@@ -30,11 +31,17 @@ def print_memory_status(is_debug):
|
||||
return gpu_mem
|
||||
|
||||
|
||||
def is_basic_or_container_type(obj):
|
||||
basic_and_container_types = (int, float, str, bool, complex, list, tuple,
|
||||
dict, set)
|
||||
return isinstance(obj, basic_and_container_types)
|
||||
|
||||
|
||||
refresh_symbol = '\U0001f504' # 🔄
|
||||
|
||||
|
||||
def get_work_name(model, version, tuner, resolution):
|
||||
model_prefix = f'Swift@{model}@{version}@{tuner}@{resolution}'
|
||||
def get_work_name(model, version, tuner):
|
||||
model_prefix = f'Swift@{model}@{version}@{tuner}'
|
||||
return model_prefix + '@' + '{0:%Y%m%d%H%M%S%f}'.format(
|
||||
datetime.datetime.now()) + ''.join(
|
||||
[str(random.randint(1, 10)) for i in range(3)])
|
||||
@@ -44,12 +51,26 @@ class TrainerUI(UIBase):
|
||||
def __init__(self, cfg, all_cfg_value, is_debug=False, language='en'):
|
||||
self.BASE_CFG_VALUE = all_cfg_value
|
||||
self.para_data = get_default(self.BASE_CFG_VALUE)
|
||||
self.train_para_data = cfg.TRAIN_PARAS
|
||||
self.run_script = os.path.join(os.path.dirname(scepter.dirname),
|
||||
cfg.SCRIPT_DIR, 'run_task.py')
|
||||
self.work_dir_pre, _ = FS.map_to_local(cfg.WORK_DIR)
|
||||
self.is_debug = is_debug
|
||||
self.current_train_model = None
|
||||
self.trainer_ins = TrainManager(self.run_script, self.work_dir_pre)
|
||||
self.component_names = TrainerUIName(language=language)
|
||||
|
||||
self.h_level_dict = {}
|
||||
for hw_tuple in self.train_para_data.RESOLUTIONS.get('VALUES', []):
|
||||
h, w = hw_tuple
|
||||
if h not in self.h_level_dict:
|
||||
self.h_level_dict[h] = []
|
||||
self.h_level_dict[h].append(w)
|
||||
self.h_level_dict = OrderedDict(
|
||||
sorted(self.h_level_dict.items(),
|
||||
key=lambda x: int(x[0]),
|
||||
reverse=False))
|
||||
|
||||
def create_ui(self):
|
||||
with gr.Box():
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
@@ -88,15 +109,6 @@ class TrainerUI(UIBase):
|
||||
'model_default', ''),
|
||||
label=self.component_names.base_model,
|
||||
interactive=True)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.tuner_name = gr.Dropdown(
|
||||
choices=self.para_data.get(
|
||||
'tuner_choices', []),
|
||||
value=self.para_data.get(
|
||||
'tuner_default', ''),
|
||||
label=self.component_names.tuner_name,
|
||||
interactive=True)
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.base_model_revision = gr.Dropdown(
|
||||
choices=self.para_data.get(
|
||||
@@ -107,13 +119,35 @@ class TrainerUI(UIBase):
|
||||
base_model_revision,
|
||||
interactive=True)
|
||||
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.resolution = gr.Dropdown(
|
||||
self.tuner_name = gr.Dropdown(
|
||||
choices=self.para_data.get(
|
||||
'resolution_choices', []),
|
||||
'tuner_choices', []),
|
||||
value=self.para_data.get(
|
||||
'resolution_default', 1024),
|
||||
label=self.component_names.resolution,
|
||||
'tuner_default', ''),
|
||||
label=self.component_names.tuner_name,
|
||||
interactive=True)
|
||||
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.resolution_height = gr.Dropdown(
|
||||
choices=list(self.h_level_dict.keys()),
|
||||
value=self.para_data.get(
|
||||
'RESOLUTION', 1024)[0],
|
||||
label=self.component_names.
|
||||
resolution_height,
|
||||
allow_custom_value=True,
|
||||
interactive=True)
|
||||
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.resolution_width = gr.Dropdown(
|
||||
choices=self.h_level_dict[
|
||||
self.resolution_height.value],
|
||||
value=self.para_data.get(
|
||||
'RESOLUTION', 1024)[1],
|
||||
label=self.component_names.
|
||||
resolution_width,
|
||||
allow_custom_value=True,
|
||||
interactive=True)
|
||||
|
||||
@@ -161,19 +195,36 @@ class TrainerUI(UIBase):
|
||||
value='')
|
||||
|
||||
with gr.Row():
|
||||
with gr.Column(scale=5, min_width=0):
|
||||
self.work_name = gr.Text(
|
||||
label=self.component_names.work_name,
|
||||
value=None,
|
||||
interactive=False)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.work_name_button = gr.Button(
|
||||
value=refresh_symbol)
|
||||
with gr.Column(scale=2, min_width=0):
|
||||
self.push_to_hub = gr.Checkbox(
|
||||
label=self.component_names.push_to_hub,
|
||||
value=False,
|
||||
visible=False)
|
||||
self.eval_prompts = gr.Dropdown(
|
||||
value=None,
|
||||
choices=self.train_para_data.get(
|
||||
'EVAL_PROMPTS', []),
|
||||
label=self.component_names.eval_prompts,
|
||||
interactive=True,
|
||||
multiselect=True,
|
||||
allow_custom_value=True,
|
||||
max_choices=20)
|
||||
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
with gr.Box():
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
gr.Markdown(self.component_names.work_name)
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
with gr.Column(scale=6, min_width=0):
|
||||
self.work_name = gr.Text(value=None,
|
||||
container=False,
|
||||
interactive=False)
|
||||
with gr.Column(scale=2, min_width=0):
|
||||
self.work_name_button = gr.Button(
|
||||
value=refresh_symbol, size='lg')
|
||||
|
||||
with gr.Row(variant='panel', visible=False, equal_height=True):
|
||||
with gr.Column(scale=2, min_width=0):
|
||||
self.push_to_hub = gr.Checkbox(
|
||||
label=self.component_names.push_to_hub,
|
||||
value=False,
|
||||
visible=False)
|
||||
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
self.examples = gr.Examples(
|
||||
@@ -193,15 +244,10 @@ class TrainerUI(UIBase):
|
||||
self.data_type, self.ms_data_space, self.ms_data_name,
|
||||
self.ms_data_subname
|
||||
])
|
||||
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
with gr.Box():
|
||||
gr.Markdown(self.component_names.log_block)
|
||||
self.training_message = gr.Markdown()
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
self.training_button = gr.Button()
|
||||
|
||||
def set_callbacks(self, inference_ui):
|
||||
def set_callbacks(self, inference_ui, manager):
|
||||
def change_data_type(data_type):
|
||||
if data_type == self.component_names.data_type_choices[0]:
|
||||
return gr.Box(visible=False)
|
||||
@@ -217,7 +263,7 @@ class TrainerUI(UIBase):
|
||||
inputs=[
|
||||
self.base_model,
|
||||
self.base_model_revision,
|
||||
self.tuner_name, self.resolution
|
||||
self.tuner_name
|
||||
],
|
||||
outputs=[self.work_name],
|
||||
queue=False)
|
||||
@@ -245,8 +291,12 @@ class TrainerUI(UIBase):
|
||||
interactive=True), \
|
||||
gr.Dropdown(value=ret_data.get('tuner_default', ''), choices=ret_data.get('tuner_choices', []),
|
||||
interactive=True), \
|
||||
gr.Dropdown(value=ret_data.get('resolution_default', 1024),
|
||||
choices=ret_data.get('resolution_choices', []), interactive=True)
|
||||
gr.Dropdown(value=ret_data.get('RESOLUTION', 1024)[0],
|
||||
choices=list(self.h_level_dict.keys()),
|
||||
interactive=True), \
|
||||
gr.Dropdown(value=ret_data.get('RESOLUTION', 1024)[1],
|
||||
choices=self.h_level_dict[ret_data.get('RESOLUTION', 1024)[0]],
|
||||
interactive=True)
|
||||
|
||||
self.base_model.change(fn=change_train_value_by_model,
|
||||
inputs=[self.base_model],
|
||||
@@ -255,7 +305,8 @@ class TrainerUI(UIBase):
|
||||
self.save_interval, self.train_batch_size,
|
||||
self.prompt_prefix,
|
||||
self.base_model_revision, self.tuner_name,
|
||||
self.resolution
|
||||
self.resolution_height,
|
||||
self.resolution_width
|
||||
],
|
||||
queue=False)
|
||||
|
||||
@@ -282,10 +333,15 @@ class TrainerUI(UIBase):
|
||||
ret_data.get('SAVE_INTERVAL', 10), \
|
||||
ret_data.get('TRAIN_BATCH_SIZE', 4), \
|
||||
ret_data.get('TRAIN_PREFIX', ''), \
|
||||
gr.Dropdown(value=ret_data.get('tuner_default', ''), choices=ret_data.get('tuner_choices', []),
|
||||
gr.Dropdown(value=ret_data.get('tuner_default', ''),
|
||||
choices=ret_data.get('tuner_choices', []),
|
||||
interactive=True), \
|
||||
gr.Dropdown(value=ret_data.get('resolution_default', 1024),
|
||||
choices=ret_data.get('resolution_choices', []), interactive=True)
|
||||
gr.Dropdown(value=ret_data.get('RESOLUTION', 1024)[0],
|
||||
choices=list(self.h_level_dict.keys()),
|
||||
interactive=True), \
|
||||
gr.Dropdown(value=ret_data.get('RESOLUTION', 1024)[1],
|
||||
choices=self.h_level_dict[ret_data.get('RESOLUTION', 1024)[0]],
|
||||
interactive=True)
|
||||
|
||||
#
|
||||
self.base_model_revision.change(
|
||||
@@ -294,11 +350,10 @@ class TrainerUI(UIBase):
|
||||
outputs=[
|
||||
self.train_epoch, self.learning_rate, self.save_interval,
|
||||
self.train_batch_size, self.prompt_prefix, self.tuner_name,
|
||||
self.resolution
|
||||
self.resolution_height, self.resolution_width
|
||||
],
|
||||
queue=False)
|
||||
|
||||
#
|
||||
#
|
||||
def change_train_value_by_model_version_tuner(base_model,
|
||||
base_model_revision,
|
||||
@@ -323,8 +378,12 @@ class TrainerUI(UIBase):
|
||||
ret_data.get('SAVE_INTERVAL', 10), \
|
||||
ret_data.get('TRAIN_BATCH_SIZE', 4), \
|
||||
ret_data.get('TRAIN_PREFIX', ''), \
|
||||
gr.Dropdown(value=ret_data.get('resolution_default', 1024),
|
||||
choices=ret_data.get('resolution_choices', []), interactive=True)
|
||||
gr.Dropdown(value=ret_data.get('RESOLUTION', 1024)[0],
|
||||
choices=list(self.h_level_dict.keys()),
|
||||
interactive=True), \
|
||||
gr.Dropdown(value=ret_data.get('RESOLUTION', 1024)[1],
|
||||
choices=self.h_level_dict[ret_data.get('RESOLUTION', 1024)[0]],
|
||||
interactive=True)
|
||||
|
||||
#
|
||||
self.tuner_name.change(fn=change_train_value_by_model_version_tuner,
|
||||
@@ -335,56 +394,30 @@ class TrainerUI(UIBase):
|
||||
outputs=[
|
||||
self.train_epoch, self.learning_rate,
|
||||
self.save_interval, self.train_batch_size,
|
||||
self.prompt_prefix, self.resolution
|
||||
self.prompt_prefix, self.resolution_height,
|
||||
self.resolution_width
|
||||
],
|
||||
queue=False)
|
||||
|
||||
#
|
||||
def change_train_value_by_model_version_tuner_resolution(
|
||||
base_model, base_model_revision, tuner_name, resolution):
|
||||
'''
|
||||
Changes to the base model will affect the training parameters,
|
||||
and it is best to define the related default values in the YAML.
|
||||
Training Iterations:
|
||||
Learning Rate:
|
||||
Training Prefix Used:
|
||||
Training Batch Size:
|
||||
Supported Resolutions:
|
||||
Supported Fine-tuning Methods:
|
||||
Supported Fine-tuning Methods:
|
||||
Prefix for Saved Model:
|
||||
'''
|
||||
ret_data = get_values_by_model_version_tuner_resolution(
|
||||
self.BASE_CFG_VALUE, base_model, base_model_revision,
|
||||
tuner_name, resolution)
|
||||
print('change_train_value_by_model_version_tuner_resolution',
|
||||
ret_data)
|
||||
# work_name = get_work_name(base_model, base_model_revision,
|
||||
# tuner_name, resolution)
|
||||
return ret_data.get('EPOCHS', 10), \
|
||||
ret_data.get('LEARNING_RATE', 0.0001), \
|
||||
ret_data.get('SAVE_INTERVAL', 10), \
|
||||
ret_data.get('TRAIN_BATCH_SIZE', 4), \
|
||||
ret_data.get('TRAIN_PREFIX', '')
|
||||
def change_resolution(h):
|
||||
if h not in self.h_level_dict:
|
||||
return gr.Dropdown()
|
||||
all_choices = self.h_level_dict[h]
|
||||
default = all_choices[0]
|
||||
return gr.Dropdown(choices=all_choices, value=default)
|
||||
|
||||
#
|
||||
self.resolution.change(
|
||||
fn=change_train_value_by_model_version_tuner_resolution,
|
||||
inputs=[
|
||||
self.base_model, self.base_model_revision, self.tuner_name,
|
||||
self.resolution
|
||||
],
|
||||
outputs=[
|
||||
self.train_epoch, self.learning_rate, self.save_interval,
|
||||
self.train_batch_size, self.prompt_prefix
|
||||
],
|
||||
queue=False)
|
||||
self.resolution_height.change(change_resolution,
|
||||
inputs=[self.resolution_height],
|
||||
outputs=[self.resolution_width],
|
||||
queue=False)
|
||||
|
||||
def run_train(work_name, data_type, ms_data_space, ms_data_name,
|
||||
ms_data_subname, base_model, base_model_revision,
|
||||
tuner_name, resolution, train_epoch, learning_rate,
|
||||
save_interval, train_batch_size, prompt_prefix,
|
||||
replace_keywords, push_to_hub):
|
||||
tuner_name, resolution_height, resolution_width,
|
||||
train_epoch, learning_rate, save_interval,
|
||||
train_batch_size, prompt_prefix, replace_keywords,
|
||||
push_to_hub, eval_prompts):
|
||||
|
||||
# Check Cuda
|
||||
if not torch.cuda.is_available() and not self.is_debug:
|
||||
raise gr.Error(self.component_names.training_err1)
|
||||
@@ -392,10 +425,25 @@ class TrainerUI(UIBase):
|
||||
if work_name == 'custom' or work_name is None or work_name == '':
|
||||
raise gr.Error(self.component_names.training_err4)
|
||||
work_dir = os.path.join(self.work_dir_pre, work_name)
|
||||
if not os.path.exists(work_dir):
|
||||
os.makedirs(work_dir)
|
||||
else:
|
||||
self.current_train_model = work_name
|
||||
if os.path.exists(work_dir) or os.path.exists(
|
||||
f'.flag/{work_name}.tmp'):
|
||||
raise gr.Error(self.component_names.training_err4)
|
||||
else:
|
||||
os.makedirs(work_dir)
|
||||
os.makedirs('.flag', exist_ok=True)
|
||||
with open(f'.flag/{work_name}.tmp', 'w') as f:
|
||||
f.write('new line.')
|
||||
# save params
|
||||
with open(os.path.join(work_dir, 'params.json'), 'w') as f_out:
|
||||
json.dump(
|
||||
{
|
||||
key: val
|
||||
for key, val in locals().items()
|
||||
if is_basic_or_container_type(val)
|
||||
},
|
||||
f_out,
|
||||
ensure_ascii=False)
|
||||
|
||||
if push_to_hub:
|
||||
model_id = work_name
|
||||
@@ -404,28 +452,23 @@ class TrainerUI(UIBase):
|
||||
else:
|
||||
hub_model_id = ''
|
||||
|
||||
# Check Cuda Memory
|
||||
if torch.cuda.is_available() and not self.is_debug:
|
||||
device = torch.device('cuda:0')
|
||||
required_memory_bytes = 4 * (1024**3)
|
||||
try:
|
||||
tensor = torch.empty( # noqa
|
||||
(required_memory_bytes // 4, ), device=device
|
||||
) # create 4GB tensor to check the memory if enough
|
||||
del tensor
|
||||
except RuntimeError:
|
||||
raise gr.Error(self.component_names.training_err2)
|
||||
|
||||
# Check Instance Valid
|
||||
if ms_data_name is None:
|
||||
raise gr.Error(self.component_names.training_err3)
|
||||
|
||||
st_time = time.time()
|
||||
|
||||
def prepare_data(data_cfg):
|
||||
def prepare_train_data(data_cfg):
|
||||
data_cfg['BATCH_SIZE'] = int(train_batch_size)
|
||||
data_cfg['PROMPT_PREFIX'] = prompt_prefix
|
||||
data_cfg['REPLACE_KEYWORDS'] = replace_keywords
|
||||
for trans in data_cfg['TRANSFORMS']:
|
||||
if trans['NAME'] in [
|
||||
'Resize', 'FlexibleResize', 'CenterCrop',
|
||||
'FlexibleCenterCrop'
|
||||
]:
|
||||
trans['SIZE'] = [
|
||||
int(resolution_height),
|
||||
int(resolution_width)
|
||||
]
|
||||
if data_type in self.component_names.data_type_choices:
|
||||
if ms_data_name.startswith(
|
||||
'http') or ms_data_name.endswith('zip'):
|
||||
@@ -482,7 +525,19 @@ class TrainerUI(UIBase):
|
||||
data_cfg['MS_REMAP_KEYS'] = {'Text': 'Prompt'}
|
||||
else:
|
||||
data_cfg['MS_REMAP_KEYS'] = None
|
||||
data_cfg['OUTPUT_SIZE'] = int(resolution)
|
||||
data_cfg['OUTPUT_SIZE'] = [
|
||||
int(resolution_height),
|
||||
int(resolution_width)
|
||||
]
|
||||
return data_cfg
|
||||
|
||||
def prepare_eval_data(data_cfg):
|
||||
data_cfg['PROMPT_PREFIX'] = prompt_prefix
|
||||
data_cfg['IMAGE_SIZE'] = [
|
||||
int(resolution_height),
|
||||
int(resolution_width)
|
||||
]
|
||||
data_cfg['PROMPT_DATA'] = eval_prompts
|
||||
return data_cfg
|
||||
|
||||
def prepare_train_config():
|
||||
@@ -490,7 +545,7 @@ class TrainerUI(UIBase):
|
||||
current_model_info = self.BASE_CFG_VALUE[base_model][
|
||||
base_model_revision]
|
||||
modify_para = current_model_info['modify_para']
|
||||
cfg = current_model_info['config_value']
|
||||
cfg = copy.deepcopy(current_model_info['config_value'])
|
||||
if isinstance(modify_para, dict) and tuner_name in modify_para:
|
||||
modify_c = modify_para[tuner_name]
|
||||
if isinstance(modify_c, dict) and 'TRAIN' in modify_c:
|
||||
@@ -518,18 +573,33 @@ class TrainerUI(UIBase):
|
||||
cfg['SOLVER']['MAX_EPOCHS'] = int(train_epoch)
|
||||
cfg['SOLVER']['TRAIN_DATA']['BATCH_SIZE'] = int(
|
||||
train_batch_size)
|
||||
cfg['SOLVER']['TUNER'] = current_model_info[
|
||||
'tuner_para'][tuner_name] if isinstance(
|
||||
current_model_info['tuner_para'],
|
||||
dict) and tuner_name in current_model_info[
|
||||
'tuner_para'] else None
|
||||
cfg['SOLVER']['TRAIN_DATA'] = prepare_data(
|
||||
if 'TUNER' in cfg['SOLVER']:
|
||||
cfg['SOLVER']['TUNER'] = current_model_info['tuner_para'][
|
||||
tuner_name] if isinstance(
|
||||
current_model_info['tuner_para'],
|
||||
dict) and tuner_name in current_model_info[
|
||||
'tuner_para'] else None
|
||||
cfg['SOLVER']['TRAIN_DATA'] = prepare_train_data(
|
||||
cfg['SOLVER']['TRAIN_DATA'])
|
||||
if eval_prompts is not None and len(eval_prompts) > 0:
|
||||
cfg['SOLVER']['EVAL_DATA'] = prepare_eval_data(
|
||||
cfg['SOLVER']['EVAL_DATA'])
|
||||
else:
|
||||
cfg['SOLVER'].pop('EVAL_DATA')
|
||||
if 'SAMPLE_ARGS' in cfg['SOLVER']:
|
||||
cfg['SOLVER']['SAMPLE_ARGS']['IMAGE_SIZE'] = [
|
||||
int(resolution_height),
|
||||
int(resolution_width)
|
||||
]
|
||||
for hook in cfg['SOLVER']['TRAIN_HOOKS']:
|
||||
if hook['NAME'] == 'CheckpointHook':
|
||||
hook['INTERVAL'] = save_interval
|
||||
hook['PUSH_TO_HUB'] = push_to_hub
|
||||
hook['HUB_MODEL_ID'] = hub_model_id
|
||||
if 'EVAL_HOOKS' in cfg['SOLVER']:
|
||||
for hook in cfg['SOLVER']['EVAL_HOOKS']:
|
||||
if hook['NAME'] == 'ProbeDataHook':
|
||||
hook['PROB_INTERVAL'] = save_interval
|
||||
|
||||
with open(cfg_file, 'w') as f_out:
|
||||
yaml.dump(cfg,
|
||||
@@ -539,48 +609,33 @@ class TrainerUI(UIBase):
|
||||
default_flow_style=False)
|
||||
return cfg_file
|
||||
|
||||
cfg = prepare_train_config()
|
||||
|
||||
def train_fn(cfg_file):
|
||||
torch.cuda.empty_cache()
|
||||
cmd = f'PYTHONPATH=. python {self.run_script} ' \
|
||||
f'--cfg={cfg_file} 2> {self.work_dir_pre}/std_out.txt'
|
||||
print(cmd)
|
||||
if not self.is_debug:
|
||||
res = os.system(cmd)
|
||||
else:
|
||||
res = 0
|
||||
if res != 0:
|
||||
error_info = '\n'.join(
|
||||
open(f'{self.work_dir_pre}/std_out.txt',
|
||||
'r').read().split('\n')[-20:])
|
||||
raise gr.Error(
|
||||
f'{self.component_names.training_err5} ({error_info}) '
|
||||
)
|
||||
|
||||
train_fn(cfg)
|
||||
|
||||
before_kill_inference = self.trainer_ins.check_memory()
|
||||
for k, v in manager.inference.pipe_manager.pipeline_level_modules.items(
|
||||
):
|
||||
if hasattr(v, 'dynamic_unload'):
|
||||
v.dynamic_unload(name='all')
|
||||
after_kill_inference = self.trainer_ins.check_memory()
|
||||
message = f'GPU info: {before_kill_inference}. \n\n'
|
||||
message += f'After unloading inference models, the GPU info: {after_kill_inference}. \n\n'
|
||||
_ = prepare_train_config()
|
||||
self.trainer_ins.start_task(work_name)
|
||||
message += self.trainer_ins.get_log(work_name)
|
||||
if work_name not in inference_ui.model_list:
|
||||
inference_ui.model_list.append(work_name)
|
||||
message = f'''
|
||||
Training completed! \n
|
||||
Save in [ {work_name} ] \n
|
||||
Take time [ {time.time() - st_time:.4f}s ] \n
|
||||
Mem: [ {print_memory_status(self.is_debug)} ]
|
||||
'''
|
||||
print(message)
|
||||
return message, gr.Dropdown.update(choices=inference_ui.model_list,
|
||||
value=work_name)
|
||||
gr.Info('Start Training!' + message)
|
||||
return gr.Dropdown.update(choices=inference_ui.model_list,
|
||||
value=work_name)
|
||||
|
||||
self.training_button.click(
|
||||
run_train,
|
||||
inputs=[
|
||||
self.work_name, self.data_type, self.ms_data_space,
|
||||
self.ms_data_name, self.ms_data_subname, self.base_model,
|
||||
self.base_model_revision, self.tuner_name, self.resolution,
|
||||
self.base_model_revision, self.tuner_name,
|
||||
self.resolution_height, self.resolution_width,
|
||||
self.train_epoch, self.learning_rate, self.save_interval,
|
||||
self.train_batch_size, self.prompt_prefix,
|
||||
self.replace_keywords, self.push_to_hub
|
||||
self.replace_keywords, self.push_to_hub, self.eval_prompts
|
||||
],
|
||||
outputs=[self.training_message, inference_ui.output_model_name],
|
||||
outputs=[inference_ui.output_model_name],
|
||||
queue=True)
|
||||
|
||||
@@ -10,7 +10,6 @@ paras_keys = [
|
||||
'MEMORY', 'EPOCHS', 'SAVE_INTERVAL', 'EPSEC', 'LEARNING_RATE',
|
||||
'IS_DEFAULT', 'TUNER'
|
||||
]
|
||||
|
||||
control_paras_keys = ['CONTROL_MODE', 'RESOLUTION', 'IS_DEFAULT']
|
||||
|
||||
|
||||
@@ -25,10 +24,9 @@ def build_meta_index(meta_cfg, config_file):
|
||||
f'Para key {key} not defined in {config_file} META/RARAS[{idx}]'
|
||||
)
|
||||
assert key in para
|
||||
tuner_type[para['TUNER'] + '@' + str(para['RESOLUTION'])] = para
|
||||
tuner_type[para['TUNER']] = para
|
||||
if para['IS_DEFAULT']:
|
||||
tuner_type['default'] = para['TUNER'] + '@' + str(
|
||||
para['RESOLUTION'])
|
||||
tuner_type['default'] = para['TUNER']
|
||||
|
||||
tuner_type['choices'] = list(tuner_type.keys())
|
||||
if 'default' in tuner_type['choices']:
|
||||
@@ -37,14 +35,6 @@ def build_meta_index(meta_cfg, config_file):
|
||||
tuner_type['default'] = tuner_type['choices'][0] if len(
|
||||
tuner_type['choices']) > 0 else ''
|
||||
|
||||
# 重新组织选项
|
||||
choices = {}
|
||||
for t_type in tuner_type['choices']:
|
||||
t_name, resolution = t_type.split('@')
|
||||
if t_name not in choices:
|
||||
choices[t_name] = []
|
||||
choices[t_name].append(int(resolution))
|
||||
tuner_type['choices'] = choices
|
||||
tuner_paras = meta_cfg.get('TUNERS', None)
|
||||
return tuner_type, tuner_paras
|
||||
|
||||
@@ -60,13 +50,10 @@ def build_meta_index_control(meta_cfg, config_file):
|
||||
f'Para key {key} not defined in {config_file} META/RARAS[{idx}]'
|
||||
)
|
||||
assert key in para
|
||||
control_type[para['CONTROL_MODE'] + '@' +
|
||||
str(para['RESOLUTION'])] = para
|
||||
control_type[para['CONTROL_MODE']] = para
|
||||
# control_type[para["CONTROL_MODE"]] = para
|
||||
if para['IS_DEFAULT']:
|
||||
control_type['default'] = para['CONTROL_MODE'] + '@' + str(
|
||||
para['RESOLUTION'])
|
||||
# control_type["default"] = para["CONTROL_MODE"]
|
||||
control_type['default'] = para['CONTROL_MODE']
|
||||
|
||||
control_type['choices'] = list(control_type.keys())
|
||||
if 'default' in control_type['choices']:
|
||||
@@ -75,14 +62,6 @@ def build_meta_index_control(meta_cfg, config_file):
|
||||
control_type['default'] = control_type['choices'][0] if len(
|
||||
control_type['choices']) > 0 else ''
|
||||
|
||||
# 重新组织选项
|
||||
choices = {}
|
||||
for t_type in control_type['choices']:
|
||||
t_name, resolution = t_type.split('@')
|
||||
if t_name not in choices:
|
||||
choices[t_name] = []
|
||||
choices[t_name].append(int(resolution))
|
||||
control_type['choices'] = choices
|
||||
return control_type, paras
|
||||
|
||||
|
||||
@@ -166,26 +145,19 @@ def get_default(config_dict):
|
||||
return ret_data
|
||||
ret_data['version_choices'] = default_version_cfg['choices']
|
||||
ret_data['version_default'] = default_version_cfg['default']
|
||||
default_tuner_cfg = default_version_cfg.get(default_version_cfg['default'],
|
||||
default_model_cfg = default_version_cfg.get(default_version_cfg['default'],
|
||||
None)
|
||||
if default_tuner_cfg is None:
|
||||
if default_model_cfg is None:
|
||||
return ret_data
|
||||
if 'tuner_type' in default_tuner_cfg and default_tuner_cfg['tuner_type'][
|
||||
if 'tuner_type' in default_model_cfg and default_model_cfg['tuner_type'][
|
||||
'default'] != '':
|
||||
default_tuner_cfg = default_tuner_cfg['tuner_type']
|
||||
default_tuner_cfg = default_model_cfg['tuner_type']
|
||||
else:
|
||||
return ret_data
|
||||
ret_data['tuner_choices'] = list(default_tuner_cfg['choices'].keys())
|
||||
ret_data['tuner_choices'] = default_tuner_cfg['choices']
|
||||
defalt_t_type = default_tuner_cfg['default']
|
||||
ret_data['tuner_default'] = defalt_t_type
|
||||
type_paras = default_tuner_cfg.get(defalt_t_type, None)
|
||||
|
||||
default_t_n = defalt_t_type.split('@')[0]
|
||||
default_r_n = int(defalt_t_type.split('@')[1])
|
||||
|
||||
ret_data['resolution_choices'] = default_tuner_cfg['choices'].get(
|
||||
default_t_n, [])
|
||||
ret_data['tuner_default'] = default_t_n
|
||||
ret_data['resolution_default'] = default_r_n
|
||||
if type_paras is not None:
|
||||
ret_data.update(type_paras)
|
||||
return ret_data
|
||||
@@ -198,21 +170,15 @@ def get_values_by_model(config_dict, model_name):
|
||||
return ret_data
|
||||
ret_data['version_choices'] = version_cfg['choices']
|
||||
ret_data['version_default'] = version_cfg['default']
|
||||
default_tuner_cfg = version_cfg.get(version_cfg['default'], None)
|
||||
if default_tuner_cfg is None:
|
||||
default_model_cfg = version_cfg.get(version_cfg['default'], None)
|
||||
if default_model_cfg is None:
|
||||
return ret_data
|
||||
default_tuner_cfg = default_tuner_cfg['tuner_type']
|
||||
ret_data['tuner_choices'] = list(default_tuner_cfg['choices'].keys())
|
||||
|
||||
default_tuner_cfg = default_model_cfg['tuner_type']
|
||||
ret_data['tuner_choices'] = default_tuner_cfg['choices']
|
||||
defalt_t_type = default_tuner_cfg['default']
|
||||
ret_data['tuner_default'] = defalt_t_type
|
||||
type_paras = default_tuner_cfg.get(defalt_t_type, None)
|
||||
|
||||
default_t_n = defalt_t_type.split('@')[0]
|
||||
default_r_n = int(defalt_t_type.split('@')[1])
|
||||
|
||||
ret_data['resolution_choices'] = default_tuner_cfg['choices'].get(
|
||||
default_t_n, [])
|
||||
ret_data['tuner_default'] = default_t_n
|
||||
ret_data['resolution_default'] = default_r_n
|
||||
if type_paras is not None:
|
||||
ret_data.update(type_paras)
|
||||
return ret_data
|
||||
@@ -226,18 +192,12 @@ def get_values_by_model_version(config_dict, model_name, version):
|
||||
tuner_cfg = version_cfg.get(version, None)
|
||||
if tuner_cfg is None:
|
||||
return ret_data
|
||||
|
||||
default_tuner_cfg = tuner_cfg['tuner_type']
|
||||
ret_data['tuner_choices'] = list(default_tuner_cfg['choices'].keys())
|
||||
ret_data['tuner_choices'] = default_tuner_cfg['choices']
|
||||
defalt_t_type = default_tuner_cfg['default']
|
||||
ret_data['tuner_default'] = defalt_t_type
|
||||
type_paras = default_tuner_cfg.get(defalt_t_type, None)
|
||||
|
||||
default_t_n = defalt_t_type.split('@')[0]
|
||||
default_r_n = int(defalt_t_type.split('@')[1])
|
||||
|
||||
ret_data['resolution_choices'] = default_tuner_cfg['choices'].get(
|
||||
default_t_n, [])
|
||||
ret_data['tuner_default'] = default_t_n
|
||||
ret_data['resolution_default'] = default_r_n
|
||||
if type_paras is not None:
|
||||
ret_data.update(type_paras)
|
||||
return ret_data
|
||||
@@ -249,34 +209,11 @@ def get_values_by_model_version_tuner(config_dict, model_name, version,
|
||||
version_cfg = config_dict.get(model_name, None)
|
||||
if version_cfg is None:
|
||||
return ret_data
|
||||
tuner_cfg = version_cfg.get(version, None)
|
||||
if tuner_cfg is None:
|
||||
model_cfg = version_cfg.get(version, None)
|
||||
if model_cfg is None:
|
||||
return ret_data
|
||||
tuner_cfg = tuner_cfg['tuner_type']
|
||||
|
||||
ret_data['resolution_choices'] = tuner_cfg['choices'].get(tuner_name, [])
|
||||
|
||||
if len(ret_data['resolution_choices']) > 0:
|
||||
t_type = '{}@{}'.format(tuner_name, ret_data['resolution_choices'][0])
|
||||
ret_data['resolution_default'] = ret_data['resolution_choices'][0]
|
||||
type_paras = tuner_cfg.get(t_type, None)
|
||||
if type_paras is not None:
|
||||
ret_data.update(type_paras)
|
||||
return ret_data
|
||||
|
||||
|
||||
def get_values_by_model_version_tuner_resolution(config_dict, model_name,
|
||||
version, tuner_name,
|
||||
resolution):
|
||||
ret_data = {}
|
||||
version_cfg = config_dict.get(model_name, None)
|
||||
if version_cfg is None:
|
||||
return ret_data
|
||||
tuner_cfg = version_cfg.get(version, None)
|
||||
if tuner_cfg is None:
|
||||
return ret_data
|
||||
tuner_cfg = tuner_cfg['tuner_type']
|
||||
type_paras = tuner_cfg.get('{}@{}'.format(tuner_name, resolution), None)
|
||||
tuner_cfg = model_cfg['tuner_type']
|
||||
type_paras = tuner_cfg.get(tuner_name, None)
|
||||
if type_paras is not None:
|
||||
ret_data.update(type_paras)
|
||||
return ret_data
|
||||
@@ -371,15 +308,8 @@ def get_control_default(config_dict):
|
||||
|
||||
ret_data['control_choices'] = list(default_control_cfg['choices'].keys())
|
||||
defalt_t_type = default_control_cfg['default']
|
||||
ret_data['control_default'] = defalt_t_type
|
||||
type_paras = default_control_cfg.get(defalt_t_type, None)
|
||||
|
||||
default_t_n = defalt_t_type.split('@')[0]
|
||||
default_r_n = int(defalt_t_type.split('@')[1])
|
||||
|
||||
ret_data['resolution_choices'] = default_control_cfg['choices'].get(
|
||||
default_t_n, [])
|
||||
ret_data['control_default'] = default_t_n
|
||||
ret_data['resolution_default'] = default_r_n
|
||||
if type_paras is not None:
|
||||
ret_data.update(type_paras)
|
||||
return ret_data
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
@@ -0,0 +1,2 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
@@ -0,0 +1,323 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
|
||||
import gradio as gr
|
||||
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.studio.tuner_manager.manager_ui.component_names import \
|
||||
TunerManagerNames
|
||||
from scepter.studio.tuner_manager.utils.dict import (delete_2level_dict,
|
||||
update_2level_dict)
|
||||
from scepter.studio.tuner_manager.utils.path import is_valid_filename
|
||||
from scepter.studio.tuner_manager.utils.yaml import save_yaml
|
||||
from scepter.studio.utils.uibase import UIBase
|
||||
|
||||
|
||||
class BrowserUI(UIBase):
|
||||
def __init__(self, cfg, language='en'):
|
||||
self.work_dir = cfg.WORK_DIR
|
||||
self.yaml = os.path.join(self.work_dir, cfg.TUNER_LIST_YAML)
|
||||
if not FS.exists(self.yaml):
|
||||
self.saved_tuners = []
|
||||
with FS.put_to(self.yaml) as local_path:
|
||||
save_yaml({'TUNERS': self.saved_tuners}, local_path)
|
||||
else:
|
||||
with FS.get_from(self.yaml) as local_path:
|
||||
self.saved_tuners = Config.get_plain_cfg(
|
||||
Config(cfg_file=local_path).TUNERS)
|
||||
self.saved_tuners_to_category()
|
||||
self.component_names = TunerManagerNames(language)
|
||||
self.language = language
|
||||
|
||||
def saved_tuners_to_category(self):
|
||||
self.saved_tuners_category = OrderedDict()
|
||||
for tuner in self.saved_tuners:
|
||||
first_level = f"{tuner['BASE_MODEL']}-{tuner['TUNER_TYPE']}"
|
||||
second_level = f"{tuner['NAME']}"
|
||||
update_2level_dict(self.saved_tuners_category,
|
||||
{first_level: {
|
||||
second_level: tuner
|
||||
}})
|
||||
|
||||
def category_to_saved_tuners(self):
|
||||
self.saved_tuners = []
|
||||
for k, v in self.saved_tuners_category.items():
|
||||
for kk, vv in v.items():
|
||||
self.saved_tuners.append(vv)
|
||||
|
||||
def get_choices_and_values(self):
|
||||
diffusion_models_choice = list(self.saved_tuners_category.keys())
|
||||
diffusion_model = diffusion_models_choice[0] if len(
|
||||
diffusion_models_choice) > 0 else None
|
||||
tuner_models_choice = []
|
||||
tuner_model = None
|
||||
if diffusion_model:
|
||||
tuner_models_choice = list(
|
||||
self.saved_tuners_category.get(diffusion_model, {}).keys())
|
||||
tuner_model = tuner_models_choice[0] if len(
|
||||
tuner_models_choice) > 0 else None
|
||||
return diffusion_models_choice, diffusion_model, tuner_models_choice, tuner_model
|
||||
|
||||
def create_ui(self, *args, **kwargs):
|
||||
diffusion_models_choice, diffusion_model, tuner_models_choice, tuner_model = self.get_choices_and_values(
|
||||
)
|
||||
with gr.Column():
|
||||
with gr.Box():
|
||||
gr.Markdown(self.component_names.browser_block_name)
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
with gr.Column(scale=4, min_width=0):
|
||||
self.diffusion_models = gr.Dropdown(
|
||||
label=self.component_names.base_models,
|
||||
choices=diffusion_models_choice,
|
||||
value=diffusion_model,
|
||||
multiselect=False,
|
||||
interactive=True)
|
||||
with gr.Column(scale=4, min_width=0):
|
||||
self.tuner_models = gr.Dropdown(
|
||||
label=self.component_names.tuner_name,
|
||||
choices=tuner_models_choice,
|
||||
value=tuner_model,
|
||||
multiselect=False,
|
||||
interactive=True)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.save_button = gr.Button(
|
||||
label='Save',
|
||||
value=self.component_names.save_symbol,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button',
|
||||
visible=True)
|
||||
self.delete_button = gr.Button(
|
||||
label='Delete',
|
||||
value=self.component_names.delete_symbol,
|
||||
elem_classes='type_row',
|
||||
elem_id='delete_button',
|
||||
visible=False)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.refresh_button = gr.Button(
|
||||
label='Delete',
|
||||
value=self.component_names.refresh_symbol,
|
||||
elem_classes='type_row',
|
||||
elem_id='refresh_button',
|
||||
visible=True)
|
||||
|
||||
def check_new_name(self, tuner_name):
|
||||
if tuner_name.strip() == '':
|
||||
return False, f"'{tuner_name}' is whitespace! Only support 'a-z', 'A-Z', '0-9' and '_'."
|
||||
if not is_valid_filename(tuner_name):
|
||||
return False, f"'{tuner_name}' is not a valid tuner name! Only support 'a-z', 'A-Z', '0-9' and '_'."
|
||||
for tuner in self.saved_tuners:
|
||||
if tuner_name == tuner['NAME']:
|
||||
return False, f"Tuner name '{tuner_name}' has been taken!"
|
||||
return True, 'legal'
|
||||
|
||||
def save_tuner(self, src_path, sub_dir, tuner_name, tuner_example):
|
||||
tar_path = os.path.join(self.work_dir, sub_dir)
|
||||
if not FS.exists(tar_path):
|
||||
FS.make_dir(tar_path)
|
||||
tar_path = os.path.join(tar_path, tuner_name)
|
||||
|
||||
FS.put_dir_from_local_dir(src_path, tar_path)
|
||||
|
||||
tuner_example_path = None
|
||||
if tuner_example is not None:
|
||||
from PIL import Image
|
||||
tuner_example_path = os.path.join(tar_path, f'{tuner_name}.jpg')
|
||||
tuner_example = Image.fromarray(tuner_example)
|
||||
with FS.put_to(tuner_example_path) as local_path:
|
||||
tuner_example.save(local_path)
|
||||
return tar_path, tuner_example_path
|
||||
|
||||
def add_tuner(self, new_tuner, manager, now_diffusion_model):
|
||||
self.saved_tuners.append(new_tuner)
|
||||
self.saved_tuners_to_category()
|
||||
with FS.put_to(self.yaml) as local_path:
|
||||
save_yaml({'TUNERS': self.saved_tuners}, local_path)
|
||||
|
||||
# register to pipeline
|
||||
new_tuner = Config(cfg_dict=new_tuner, load=False)
|
||||
manager.inference.model_manage_ui.pipe_manager.register_tuner(
|
||||
new_tuner,
|
||||
name=new_tuner.NAME_ZH
|
||||
if self.language == 'zh' else new_tuner.NAME,
|
||||
is_customized=True)
|
||||
|
||||
# update choices
|
||||
pipe_manager = manager.inference.model_manage_ui.pipe_manager
|
||||
now_pipeline = pipe_manager.model_level_info[now_diffusion_model][
|
||||
'pipeline'][0]
|
||||
default_choices = pipe_manager.module_level_choices
|
||||
custom_tunner_choices = []
|
||||
if 'customized_tuners' in default_choices and now_pipeline in default_choices[
|
||||
'customized_tuners']:
|
||||
custom_tunner_choices = default_choices['customized_tuners'][
|
||||
now_pipeline]['choices']
|
||||
|
||||
# update tuner ui name_level_tuners
|
||||
name_level_tuners = manager.inference.tuner_ui.name_level_tuners
|
||||
if new_tuner.BASE_MODEL not in name_level_tuners:
|
||||
name_level_tuners[new_tuner.BASE_MODEL] = {}
|
||||
if self.language == 'zh':
|
||||
name_level_tuners[new_tuner.BASE_MODEL][
|
||||
new_tuner.NAME_ZH] = new_tuner
|
||||
else:
|
||||
name_level_tuners[new_tuner.BASE_MODEL][new_tuner.NAME] = new_tuner
|
||||
|
||||
return custom_tunner_choices
|
||||
|
||||
def delete_tuner(self, first_level, second_level, manager,
|
||||
now_diffusion_model):
|
||||
self.saved_tuners_category, del_tuner = delete_2level_dict(
|
||||
self.saved_tuners_category, first_level, second_level)
|
||||
self.category_to_saved_tuners()
|
||||
save_yaml({'TUNERS': self.saved_tuners}, self.yaml)
|
||||
|
||||
# update choices
|
||||
pipe_manager = manager.inference.model_manage_ui.pipe_manager
|
||||
now_pipeline = pipe_manager.model_level_info[now_diffusion_model][
|
||||
'pipeline'][0]
|
||||
default_choices = pipe_manager.module_level_choices
|
||||
custom_tuner_choices = []
|
||||
if 'customized_tuners' in default_choices and now_pipeline in default_choices[
|
||||
'customized_tuners']:
|
||||
custom_tuner_choices = default_choices['customized_tuners'][
|
||||
now_pipeline]['choices']
|
||||
|
||||
# update tuner ui name_level_tuners
|
||||
del_tuner = Config(cfg_dict=del_tuner, load=False)
|
||||
name_level_tuners = manager.inference.tuner_ui.name_level_tuners
|
||||
if del_tuner.BASE_MODEL in name_level_tuners:
|
||||
if self.language == 'zh':
|
||||
del name_level_tuners[del_tuner.BASE_MODEL][del_tuner.NAME_ZH]
|
||||
else:
|
||||
del name_level_tuners[del_tuner.BASE_MODEL][del_tuner.NAME]
|
||||
|
||||
return custom_tuner_choices
|
||||
|
||||
def set_callbacks(self, manager, info_ui):
|
||||
def refresh_browser():
|
||||
diffusion_models_choice, diffusion_model, tuner_models_choice, tuner_model = self.get_choices_and_values(
|
||||
)
|
||||
return (gr.Dropdown(choices=diffusion_models_choice,
|
||||
value=diffusion_model),
|
||||
gr.Dropdown(choices=tuner_models_choice,
|
||||
value=tuner_model))
|
||||
|
||||
self.refresh_button.click(
|
||||
refresh_browser,
|
||||
inputs=[],
|
||||
outputs=[self.diffusion_models, self.tuner_models])
|
||||
|
||||
def diffusion_model_change(diffusion_model):
|
||||
choices = list(
|
||||
self.saved_tuners_category.get(diffusion_model, {}).keys())
|
||||
return gr.Dropdown(choices=choices,
|
||||
value=choices[-1] if len(choices) > 0 else None)
|
||||
|
||||
self.diffusion_models.change(diffusion_model_change,
|
||||
inputs=[self.diffusion_models],
|
||||
outputs=[self.tuner_models],
|
||||
queue=True)
|
||||
|
||||
def tuner_model_change(tuner_model, diffusion_model):
|
||||
tuner_info = {}
|
||||
if tuner_model is not None:
|
||||
tuner_info = self.saved_tuners_category[diffusion_model][
|
||||
tuner_model]
|
||||
image_path = tuner_info.get('IMAGE_PATH', None)
|
||||
if image_path is not None:
|
||||
image_path = FS.get_from(image_path)
|
||||
return (gr.Text(value=tuner_info.get('NAME', '')),
|
||||
gr.Text(value=tuner_info.get('NAME', ''),
|
||||
interactive=True),
|
||||
gr.Text(value=tuner_info.get('TUNER_TYPE', '')),
|
||||
gr.Text(value=tuner_info.get('BASE_MODEL', '')),
|
||||
gr.Text(value=tuner_info.get('DESCRIPTION', ''),
|
||||
interactive=True), gr.Image(value=image_path),
|
||||
gr.Text(value=tuner_info.get('PROMPT_EXAMPLE', ''),
|
||||
interactive=True))
|
||||
|
||||
self.tuner_models.change(
|
||||
tuner_model_change,
|
||||
inputs=[self.tuner_models, self.diffusion_models],
|
||||
outputs=[
|
||||
info_ui.tuner_name, info_ui.new_name, info_ui.tuner_type,
|
||||
info_ui.base_model, info_ui.tuner_desc, info_ui.tuner_example,
|
||||
info_ui.tuner_prompt_example
|
||||
],
|
||||
queue=False)
|
||||
|
||||
def save_tuner(tuner_name, new_name, tuner_desc, tuner_example,
|
||||
tuner_prompt_example, now_diffusion_model, info_path):
|
||||
is_legal, msg = self.check_new_name(new_name)
|
||||
if not is_legal:
|
||||
gr.Info('Save failed because ' + msg)
|
||||
return (gr.Dropdown(), gr.Text(), gr.Dropdown())
|
||||
|
||||
info = Config(cfg_file=info_path)
|
||||
model_dir = info.MODEL_PATH
|
||||
sub_dir = f'{info.BASE_MODEL}-{info.TUNER_TYPE}'
|
||||
model_dir, tuner_example = self.save_tuner(model_dir, sub_dir,
|
||||
new_name, tuner_example)
|
||||
|
||||
# config info update
|
||||
new_tuner = {
|
||||
'NAME': new_name,
|
||||
'NAME_ZH': new_name,
|
||||
'SOURCE': 'self_train',
|
||||
'DESCRIPTION': tuner_desc,
|
||||
'BASE_MODEL': info.BASE_MODEL,
|
||||
'MODEL_PATH': model_dir,
|
||||
'IMAGE_PATH': tuner_example,
|
||||
'TUNER_TYPE': info.TUNER_TYPE,
|
||||
'PROMPT_EXAMPLE': tuner_prompt_example
|
||||
}
|
||||
custom_tuner_choices = self.add_tuner(new_tuner, manager,
|
||||
now_diffusion_model)
|
||||
|
||||
return (gr.Dropdown(choices=list(
|
||||
self.saved_tuners_category.keys()),
|
||||
value=sub_dir),
|
||||
gr.Dropdown(choices=list(
|
||||
self.saved_tuners_category.get(sub_dir, {}).keys()),
|
||||
value=new_name), gr.Text(value=new_name),
|
||||
gr.Dropdown(choices=custom_tuner_choices))
|
||||
|
||||
self.save_button.click(
|
||||
save_tuner,
|
||||
inputs=[
|
||||
info_ui.tuner_name, info_ui.new_name, info_ui.tuner_desc,
|
||||
info_ui.tuner_example, info_ui.tuner_prompt_example,
|
||||
manager.inference.model_manage_ui.diffusion_model,
|
||||
manager.inference.infer_info
|
||||
],
|
||||
outputs=[
|
||||
self.diffusion_models, self.tuner_models, info_ui.tuner_name,
|
||||
manager.inference.tuner_ui.custom_tuner_model
|
||||
],
|
||||
queue=True)
|
||||
|
||||
def delete_tuner(tuner_name, tuner_type, base_model,
|
||||
now_diffusion_model):
|
||||
first_level = f'{base_model}-{tuner_type}'
|
||||
second_level = f'{tuner_name}'
|
||||
custom_tuner_choices = self.delete_tuner(first_level, second_level,
|
||||
manager,
|
||||
now_diffusion_model)
|
||||
return (gr.Dropdown(
|
||||
choices=list(self.saved_tuners_category.keys()),
|
||||
value=None), gr.Dropdown(choices=custom_tuner_choices))
|
||||
|
||||
self.delete_button.click(
|
||||
delete_tuner,
|
||||
inputs=[
|
||||
info_ui.tuner_name, info_ui.tuner_type, info_ui.base_model,
|
||||
manager.inference.model_manage_ui.diffusion_model
|
||||
],
|
||||
outputs=[
|
||||
self.diffusion_models,
|
||||
manager.inference.tuner_ui.custom_tuner_model
|
||||
],
|
||||
queue=True)
|
||||
@@ -0,0 +1,38 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
# For dataset manager
|
||||
|
||||
|
||||
class TunerManagerNames():
|
||||
def __init__(self, language='en'):
|
||||
self.save_symbol = '\U0001F4BE' # 💾
|
||||
self.delete_symbol = '\U0001f5d1' # 🗑️
|
||||
self.refresh_symbol = '\U0001f504' # 🔄
|
||||
if language == 'en':
|
||||
self.browser_block_name = 'Tuner Browser'
|
||||
self.base_models = 'Base Model-Tuner Type'
|
||||
self.tuner_models = 'Tuner Name'
|
||||
self.info_block_name = 'Tuner Info'
|
||||
self.tuner_name = 'Tuner Name'
|
||||
self.rename = 'Rename'
|
||||
self.tuner_type = 'Tuner Type'
|
||||
self.base_model_name = 'Base Model Name'
|
||||
self.tuner_desc = 'Tuner Description'
|
||||
self.tuner_example = 'Results Example'
|
||||
self.tuner_prompt_example = 'Prompt Example'
|
||||
self.save = 'save changes'
|
||||
self.delete = 'Delete'
|
||||
elif language == 'zh':
|
||||
self.browser_block_name = '微调模型查找'
|
||||
self.base_models = '基模型-微调类型'
|
||||
self.tuner_models = '微调模型名称'
|
||||
self.info_block_name = '微调模型详情'
|
||||
self.tuner_name = '微调模型名称'
|
||||
self.rename = '重命名'
|
||||
self.tuner_type = '微调模型类型'
|
||||
self.base_model_name = '基模型名称'
|
||||
self.tuner_desc = '微调模型描述'
|
||||
self.tuner_example = '示例结果'
|
||||
self.tuner_prompt_example = '示例提示词'
|
||||
self.save = '保存修改'
|
||||
self.delete = '删除'
|
||||
@@ -0,0 +1,65 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import gradio as gr
|
||||
|
||||
from scepter.studio.tuner_manager.manager_ui.component_names import \
|
||||
TunerManagerNames
|
||||
from scepter.studio.utils.uibase import UIBase
|
||||
|
||||
|
||||
class InfoUI(UIBase):
|
||||
def __init__(self, cfg, language='en'):
|
||||
self.component_names = TunerManagerNames(language)
|
||||
|
||||
def create_ui(self, *args, **kwargs):
|
||||
with gr.Column():
|
||||
with gr.Box():
|
||||
gr.Markdown(self.component_names.info_block_name)
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
with gr.Column(variant='panel', scale=1, min_width=0):
|
||||
with gr.Group(visible=True):
|
||||
with gr.Row(equal_height=True):
|
||||
with gr.Column(scale=1):
|
||||
self.tuner_name = gr.Text(
|
||||
value='',
|
||||
label=self.component_names.tuner_name,
|
||||
interactive=False)
|
||||
with gr.Column(scale=1):
|
||||
self.new_name = gr.Text(
|
||||
value='',
|
||||
label=self.component_names.rename)
|
||||
with gr.Row(equal_height=True):
|
||||
with gr.Column(scale=1):
|
||||
self.tuner_type = gr.Text(
|
||||
value='',
|
||||
label=self.component_names.tuner_type,
|
||||
interactive=False)
|
||||
with gr.Column(scale=1):
|
||||
self.base_model = gr.Text(
|
||||
value='',
|
||||
label=self.component_names.
|
||||
base_model_name,
|
||||
interactive=False)
|
||||
with gr.Column(scale=1):
|
||||
self.tuner_desc = gr.Text(
|
||||
value='',
|
||||
label=self.component_names.tuner_desc,
|
||||
lines=4)
|
||||
with gr.Column(variant='panel', scale=1, min_width=0):
|
||||
with gr.Group(visible=True):
|
||||
with gr.Row(equal_height=True):
|
||||
self.tuner_example = gr.Image(
|
||||
label=self.component_names.tuner_example,
|
||||
source='upload',
|
||||
value=None,
|
||||
interactive=True)
|
||||
with gr.Row(equal_height=True):
|
||||
self.tuner_prompt_example = gr.Text(
|
||||
value='',
|
||||
label=self.component_names.
|
||||
tuner_prompt_example,
|
||||
lines=2)
|
||||
|
||||
def set_callbacks(self, manager):
|
||||
pass
|
||||
@@ -0,0 +1,34 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import os.path
|
||||
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.studio.tuner_manager.manager_ui.browser_ui import BrowserUI
|
||||
from scepter.studio.tuner_manager.manager_ui.info_ui import InfoUI
|
||||
from scepter.studio.utils.env import init_env
|
||||
|
||||
|
||||
class TunerManagerUI():
|
||||
def __init__(self,
|
||||
cfg_general_file,
|
||||
is_debug=False,
|
||||
language='en',
|
||||
root_work_dir='./'):
|
||||
cfg_general = Config(cfg_file=cfg_general_file)
|
||||
cfg_general.WORK_DIR = os.path.join(root_work_dir,
|
||||
cfg_general.WORK_DIR)
|
||||
if not FS.exists(cfg_general.WORK_DIR):
|
||||
FS.make_dir(cfg_general.WORK_DIR)
|
||||
|
||||
cfg_general = init_env(cfg_general)
|
||||
self.info_ui = InfoUI(cfg_general, language=language)
|
||||
self.browser_ui = BrowserUI(cfg_general, language=language)
|
||||
|
||||
def create_ui(self):
|
||||
self.browser_ui.create_ui()
|
||||
self.info_ui.create_ui()
|
||||
|
||||
def set_callbacks(self, manager):
|
||||
self.info_ui.set_callbacks(manager)
|
||||
self.browser_ui.set_callbacks(manager, self.info_ui)
|
||||
@@ -0,0 +1,2 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
@@ -0,0 +1,17 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
def update_2level_dict(d, new_dict):
|
||||
for first, v in new_dict.items():
|
||||
if first in d:
|
||||
d[first].update(v)
|
||||
else:
|
||||
d[first] = v
|
||||
return d
|
||||
|
||||
|
||||
def delete_2level_dict(d, first_key, second_key):
|
||||
first = d.pop(first_key)
|
||||
second = first.pop(second_key)
|
||||
if len(first) > 0:
|
||||
d.update({first_key: first})
|
||||
return d, second
|
||||
@@ -0,0 +1,10 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import re
|
||||
|
||||
|
||||
def is_valid_filename(filename):
|
||||
if re.match('^[A-Za-z0-9_@]+$', filename):
|
||||
return True
|
||||
else:
|
||||
return False
|
||||
@@ -0,0 +1,12 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
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)
|
||||
@@ -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))
|
||||
|
||||
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -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()
|
||||
|
||||