Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8a9b791e3f | ||
|
|
ee7fe888f2 | ||
|
|
92c849412e | ||
|
|
cdba82baf8 | ||
|
|
565c7957d8 | ||
|
|
e00c23d09a | ||
|
|
bf53829530 | ||
|
|
35aada8ce8 | ||
|
|
d3ce651bf7 | ||
|
|
7a58c91940 | ||
|
|
4e1606af2d | ||
|
|
01c03683e8 | ||
|
|
9adb273e4b | ||
|
|
2249ff37c9 |
|
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,15 +9,24 @@
|
|||||||
</p>
|
</p>
|
||||||
|
|
||||||
## 📖 Table of Contents
|
## 📖 Table of Contents
|
||||||
- [Introduction](#-introduction)
|
|
||||||
- [News](#-news)
|
- [News](#-news)
|
||||||
- [Installation](#-Installation)
|
- [Introduction](#-introduction)
|
||||||
|
- [Installation](#%EF%B8%8F-installation)
|
||||||
- [Getting Started](#-getting-started)
|
- [Getting Started](#-getting-started)
|
||||||
- [SCEPTER Studio](#-scepter-studio)
|
- [SCEPTER Studio](#%EF%B8%8F-scepter-studio)
|
||||||
- [Gallery](#-gallery)
|
- [Gallery](#%EF%B8%8F-gallery)
|
||||||
- [Features](#-features)
|
- [Features](#-features)
|
||||||
- [Learn More](#-learn-more)
|
- [Learn More](#-learn-more)
|
||||||
- [License](#license)
|
- [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
|
## 📝 Introduction
|
||||||
|
|
||||||
@@ -28,7 +37,7 @@ Main Feature:
|
|||||||
- Task:
|
- Task:
|
||||||
- Text-to-image generation
|
- Text-to-image generation
|
||||||
- Controllable image synthesis
|
- Controllable image synthesis
|
||||||
- Image editing (TODO)
|
- Image editing
|
||||||
- Training / Inference:
|
- Training / Inference:
|
||||||
- Distribute: DDP / FSDP / FairScale / Xformers
|
- Distribute: DDP / FSDP / FairScale / Xformers
|
||||||
- File system: Local / Http / OSS / Modelscope
|
- File system: Local / Http / OSS / Modelscope
|
||||||
@@ -40,15 +49,9 @@ Main Feature:
|
|||||||
Currently supported approaches (and counting):
|
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)
|
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/)
|
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(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/)
|
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/)
|
||||||
## 🎉 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.
|
|
||||||
|
|
||||||
## 🛠️ Installation
|
## 🛠️ 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)
|
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
|
```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
|
### Training
|
||||||
@@ -165,6 +168,14 @@ python scepter/tools/run_inference.py --cfg scepter/methods/scedit/ctr/sd21_768_
|
|||||||
python scepter/tools/run_inference.py --cfg scepter/methods/scedit/ctr/sd21_768_sce_ctr_pose.yaml --num_samples 1 --prompt 'super mario' --save_folder 'test_mario_pose' --image_size 768 --task control --image 'asset/images/pose_source.png' --control_mode source --pretrained_model ms://damo/scepter_scedit@controllable_model/SD2.1/pose_control/0_SwiftSCETuning/pytorch_model.bin # pose
|
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
|
## 🖥️ SCEPTER Studio
|
||||||
|
|
||||||
### Launch
|
### Launch
|
||||||
@@ -172,15 +183,103 @@ python scepter/tools/run_inference.py --cfg scepter/methods/scedit/ctr/sd21_768_
|
|||||||
To fully experience **SCEPTER Studio**, you can launch the following command line:
|
To fully experience **SCEPTER Studio**, you can launch the following command line:
|
||||||
|
|
||||||
```shell
|
```shell
|
||||||
|
pip install scepter
|
||||||
|
python -m scepter.tools.webui
|
||||||
|
```
|
||||||
|
or run after clone repo code
|
||||||
|
```shell
|
||||||
|
git clone https://github.com/modelscope/scepter.git
|
||||||
PYTHONPATH=. python scepter/tools/webui.py --cfg scepter/methods/studio/scepter_ui.yaml
|
PYTHONPATH=. python scepter/tools/webui.py --cfg scepter/methods/studio/scepter_ui.yaml
|
||||||
```
|
```
|
||||||
|
|
||||||
|
The startup of **SCEPTER Studio** eliminates the need for manual downloading and organizing of models; it will automatically load the corresponding models and store them in a local directory.
|
||||||
|
Depending on the network and hardware situation, the initial startup usually requires 15-60 minutes, primarily involving the download and processing of SDv1.5, SDv2.1, and SDXL models.
|
||||||
|
Therefore, subsequent startups will become much faster (about one minute) as downloading is no longer required.
|
||||||
|
|
||||||
|
* 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
|
### 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)
|
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
|
## 🖼️ 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
|
### Dragon Year Special: Dragon Tuner
|
||||||
|
|
||||||
<table>
|
<table>
|
||||||
@@ -234,22 +333,31 @@ We deploy a work studio on Modelscope that includes only the inference tab, plea
|
|||||||
| SD 2.1 | 🪄 | 🪄 | 🪄 | 🪄 | 🪄 |
|
| SD 2.1 | 🪄 | 🪄 | 🪄 | 🪄 | 🪄 |
|
||||||
| SD XL | 🪄 | 🪄 | 🪄 | 🪄 | 🪄 |
|
| SD XL | 🪄 | 🪄 | 🪄 | 🪄 | 🪄 |
|
||||||
|
|
||||||
|
### Image Editing
|
||||||
|
- LAR-Gen
|
||||||
|
|
||||||
|
| **Model** | **Locate** | **Assign** | **Refine** |
|
||||||
|
|:---------:|:----------:|:----------:|:----------:|
|
||||||
|
| SD XL | 🪄 | 🪄 | ⏳ |
|
||||||
|
|
||||||
### Model URL
|
### Model URL
|
||||||
|
|
||||||
- ✅ indicates support for both training and inference.
|
- ✅ indicates support for both training and inference.
|
||||||
- 🪄 denotes that the model has been published.
|
- 🪄 denotes that the model has been published.
|
||||||
|
- ⏳ denotes that the module has not been integrated currently.
|
||||||
- More models will be released in the future.
|
- More models will be released in the future.
|
||||||
|
|
||||||
| Model | URL |
|
| 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.
|
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
|
## 🔍 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.
|
Discover more about open-source projects on image generation, video generation, and editing tasks.
|
||||||
|
|
||||||
@@ -261,7 +369,20 @@ PS: Scripts running within the SCEPTER framework will automatically fetch and lo
|
|||||||
|
|
||||||
SWIFT (Scalable lightWeight Infrastructure for Fine-Tuning) is an extensible framwork designed to faciliate lightweight model fine-tuning and inference.
|
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
|
## License
|
||||||
|
|
||||||
This project is licensed under the [Apache License (Version 2.0)](https://github.com/modelscope/modelscope/blob/master/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
|
einops
|
||||||
modelscope
|
modelscope
|
||||||
ms-swift>=1.5.2
|
ms-swift>=1.5.2
|
||||||
@@ -6,6 +8,7 @@ open_clip_torch
|
|||||||
opencv-python
|
opencv-python
|
||||||
opencv_transforms>=0.0.6
|
opencv_transforms>=0.0.6
|
||||||
oss2>=2.15.0
|
oss2>=2.15.0
|
||||||
|
pycocotools
|
||||||
pyyaml>=5.3.1
|
pyyaml>=5.3.1
|
||||||
scikit-image
|
scikit-image
|
||||||
torchsde
|
torchsde
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
git+https://github.com/cocodataset/panopticapi.git
|
||||||
torch==2.0.1
|
torch==2.0.1
|
||||||
torchvision==0.15.2
|
torchvision==0.15.2
|
||||||
xformers==0.0.21
|
xformers==0.0.21
|
||||||
|
|||||||
@@ -1,2 +1,3 @@
|
|||||||
gradio>=3.47.1,<4.0.0
|
gradio>=3.47.1,<4.0.0
|
||||||
imagehash
|
imagehash
|
||||||
|
psutil
|
||||||
|
|||||||
@@ -79,6 +79,7 @@ EXTENSION_PARAS:
|
|||||||
MANTRA_BOOK: scepter/methods/studio/extensions/mantra_book/mantra_book.yaml
|
MANTRA_BOOK: scepter/methods/studio/extensions/mantra_book/mantra_book.yaml
|
||||||
OFFICIAL_TUNERS: scepter/methods/studio/extensions/tuners/official_tuners.yaml
|
OFFICIAL_TUNERS: scepter/methods/studio/extensions/tuners/official_tuners.yaml
|
||||||
OFFICIAL_CONTROLLERS: scepter/methods/studio/extensions/controllers/official_controllers.yaml
|
OFFICIAL_CONTROLLERS: scepter/methods/studio/extensions/controllers/official_controllers.yaml
|
||||||
|
TUNER_MANAGER: scepter/methods/studio/tuner_manager/tuner_manager.yaml
|
||||||
CONTROLABLE_ANNOTATORS:
|
CONTROLABLE_ANNOTATORS:
|
||||||
-
|
-
|
||||||
NAME: "CannyAnnotator"
|
NAME: "CannyAnnotator"
|
||||||
|
|||||||
@@ -0,0 +1,268 @@
|
|||||||
|
NAME: LARGEN
|
||||||
|
IS_DEFAULT: False
|
||||||
|
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-codes">
|
||||||
<div class="qr-code-container">
|
<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">
|
<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>
|
||||||
<div class="qr-code-container">
|
<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">
|
<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>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
@@ -79,6 +79,10 @@ INTERFACE:
|
|||||||
NAME_EN: Train
|
NAME_EN: Train
|
||||||
IFID: self_train
|
IFID: self_train
|
||||||
CONFIG: scepter/methods/studio/self_train/self_train.yaml
|
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: 推理
|
||||||
NAME_EN: Inference
|
NAME_EN: Inference
|
||||||
IFID: inference
|
IFID: inference
|
||||||
|
|||||||
@@ -11,13 +11,13 @@ META:
|
|||||||
DEFAULT_SAMPLER: "dpmpp_2s_ancestral"
|
DEFAULT_SAMPLER: "dpmpp_2s_ancestral"
|
||||||
DEFAULT_SAMPLE_STEPS: 40
|
DEFAULT_SAMPLE_STEPS: 40
|
||||||
INFERENCE_N_PROMPT: ""
|
INFERENCE_N_PROMPT: ""
|
||||||
RESOLUTION: 1024
|
RESOLUTION: [1024, 1024]
|
||||||
PARAS:
|
PARAS:
|
||||||
-
|
-
|
||||||
TRAIN_BATCH_SIZE: 2
|
TRAIN_BATCH_SIZE: 2
|
||||||
TRAIN_PREFIX: ""
|
TRAIN_PREFIX: ""
|
||||||
TRAIN_N_PROMPT: ""
|
TRAIN_N_PROMPT: ""
|
||||||
RESOLUTION: 1024
|
RESOLUTION: [1024, 1024]
|
||||||
MEMORY: 29000
|
MEMORY: 29000
|
||||||
EPOCHS: 50
|
EPOCHS: 50
|
||||||
SAVE_INTERVAL: 25
|
SAVE_INTERVAL: 25
|
||||||
@@ -29,7 +29,7 @@ META:
|
|||||||
TRAIN_BATCH_SIZE: 2
|
TRAIN_BATCH_SIZE: 2
|
||||||
TRAIN_PREFIX: ""
|
TRAIN_PREFIX: ""
|
||||||
TRAIN_N_PROMPT: ""
|
TRAIN_N_PROMPT: ""
|
||||||
RESOLUTION: 1024
|
RESOLUTION: [1024, 1024]
|
||||||
MEMORY: 29000
|
MEMORY: 29000
|
||||||
EPOCHS: 50
|
EPOCHS: 50
|
||||||
SAVE_INTERVAL: 25
|
SAVE_INTERVAL: 25
|
||||||
@@ -41,7 +41,7 @@ META:
|
|||||||
TRAIN_BATCH_SIZE: 2
|
TRAIN_BATCH_SIZE: 2
|
||||||
TRAIN_PREFIX: ""
|
TRAIN_PREFIX: ""
|
||||||
TRAIN_N_PROMPT: ""
|
TRAIN_N_PROMPT: ""
|
||||||
RESOLUTION: 1024
|
RESOLUTION: [1024, 1024]
|
||||||
MEMORY: 29000
|
MEMORY: 29000
|
||||||
EPOCHS: 200
|
EPOCHS: 200
|
||||||
SAVE_INTERVAL: 25
|
SAVE_INTERVAL: 25
|
||||||
@@ -53,7 +53,7 @@ META:
|
|||||||
TRAIN_BATCH_SIZE: 2
|
TRAIN_BATCH_SIZE: 2
|
||||||
TRAIN_PREFIX: ""
|
TRAIN_PREFIX: ""
|
||||||
TRAIN_N_PROMPT: ""
|
TRAIN_N_PROMPT: ""
|
||||||
RESOLUTION: 1024
|
RESOLUTION: [1024, 1024]
|
||||||
MEMORY: 29000
|
MEMORY: 29000
|
||||||
EPOCHS: 200
|
EPOCHS: 200
|
||||||
SAVE_INTERVAL: 25
|
SAVE_INTERVAL: 25
|
||||||
@@ -65,7 +65,7 @@ META:
|
|||||||
TRAIN_BATCH_SIZE: 2
|
TRAIN_BATCH_SIZE: 2
|
||||||
TRAIN_PREFIX: ""
|
TRAIN_PREFIX: ""
|
||||||
TRAIN_N_PROMPT: ""
|
TRAIN_N_PROMPT: ""
|
||||||
RESOLUTION: 1024
|
RESOLUTION: [1024, 1024]
|
||||||
MEMORY: 29000
|
MEMORY: 29000
|
||||||
EPOCHS: 50
|
EPOCHS: 50
|
||||||
SAVE_INTERVAL: 25
|
SAVE_INTERVAL: 25
|
||||||
@@ -146,10 +146,12 @@ SOLVER:
|
|||||||
MAX_EPOCHS: -1
|
MAX_EPOCHS: -1
|
||||||
# NUM_FOLDS DESCRIPTION: Num folds for training. TYPE: int default: 1
|
# NUM_FOLDS DESCRIPTION: Num folds for training. TYPE: int default: 1
|
||||||
NUM_FOLDS: 1
|
NUM_FOLDS: 1
|
||||||
|
#
|
||||||
|
EVAL_INTERVAL: -1
|
||||||
# WORK_DIR DESCRIPTION: Save dir of the training log or model. TYPE: str default: ''
|
# WORK_DIR DESCRIPTION: Save dir of the training log or model. TYPE: str default: ''
|
||||||
WORK_DIR:
|
WORK_DIR:
|
||||||
# LOG_FILE DESCRIPTION: Save log path. TYPE: str default: ''
|
# LOG_FILE DESCRIPTION: Save log path. TYPE: str default: ''
|
||||||
LOG_FILE: stg_log.txt
|
LOG_FILE: std_log.txt
|
||||||
FILE_SYSTEM:
|
FILE_SYSTEM:
|
||||||
NAME: "ModelscopeFs"
|
NAME: "ModelscopeFs"
|
||||||
TEMP_DIR: "./cache/data"
|
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']
|
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' ]
|
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:
|
TRAIN_HOOKS:
|
||||||
- NAME: BackwardHook
|
-
|
||||||
# GRADIENT_CLIP: 1.0
|
NAME: BackwardHook
|
||||||
PRIORITY: 0
|
PRIORITY: 0
|
||||||
- NAME: LogHook
|
-
|
||||||
LOG_INTERVAL: 50
|
NAME: LogHook
|
||||||
|
LOG_INTERVAL: 10
|
||||||
SHOW_GPU_MEM: True
|
SHOW_GPU_MEM: True
|
||||||
|
-
|
||||||
|
NAME: TensorboardLogHook
|
||||||
-
|
-
|
||||||
NAME: CheckpointHook
|
NAME: CheckpointHook
|
||||||
SAVE_LAST: True
|
|
||||||
INTERVAL: 10000
|
INTERVAL: 10000
|
||||||
PRIORITY: 200
|
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_2m_sde'
|
||||||
-
|
-
|
||||||
NAME: 'dpmpp_2s_ancestral'
|
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_SAMPLER: "ddim"
|
||||||
DEFAULT_SAMPLE_STEPS: 40
|
DEFAULT_SAMPLE_STEPS: 40
|
||||||
INFERENCE_N_PROMPT: ""
|
INFERENCE_N_PROMPT: ""
|
||||||
RESOLUTION: 512
|
RESOLUTION: [512, 512]
|
||||||
PARAS:
|
PARAS:
|
||||||
-
|
-
|
||||||
TRAIN_BATCH_SIZE: 2
|
TRAIN_BATCH_SIZE: 2
|
||||||
TRAIN_PREFIX: ""
|
TRAIN_PREFIX: ""
|
||||||
TRAIN_N_PROMPT: ""
|
TRAIN_N_PROMPT: ""
|
||||||
RESOLUTION: 512
|
RESOLUTION: [512, 512]
|
||||||
MEMORY: 29000
|
MEMORY: 29000
|
||||||
EPOCHS: 50
|
EPOCHS: 50
|
||||||
SAVE_INTERVAL: 25
|
SAVE_INTERVAL: 25
|
||||||
@@ -28,7 +28,7 @@ META:
|
|||||||
TRAIN_BATCH_SIZE: 2
|
TRAIN_BATCH_SIZE: 2
|
||||||
TRAIN_PREFIX: ""
|
TRAIN_PREFIX: ""
|
||||||
TRAIN_N_PROMPT: ""
|
TRAIN_N_PROMPT: ""
|
||||||
RESOLUTION: 512
|
RESOLUTION: [512, 512]
|
||||||
MEMORY: 29000
|
MEMORY: 29000
|
||||||
EPOCHS: 50
|
EPOCHS: 50
|
||||||
SAVE_INTERVAL: 25
|
SAVE_INTERVAL: 25
|
||||||
@@ -40,7 +40,7 @@ META:
|
|||||||
TRAIN_BATCH_SIZE: 2
|
TRAIN_BATCH_SIZE: 2
|
||||||
TRAIN_PREFIX: ""
|
TRAIN_PREFIX: ""
|
||||||
TRAIN_N_PROMPT: ""
|
TRAIN_N_PROMPT: ""
|
||||||
RESOLUTION: 512
|
RESOLUTION: [512, 512]
|
||||||
MEMORY: 29000
|
MEMORY: 29000
|
||||||
EPOCHS: 200
|
EPOCHS: 200
|
||||||
SAVE_INTERVAL: 25
|
SAVE_INTERVAL: 25
|
||||||
@@ -53,7 +53,7 @@ META:
|
|||||||
TRAIN_BATCH_SIZE: 2
|
TRAIN_BATCH_SIZE: 2
|
||||||
TRAIN_PREFIX: ""
|
TRAIN_PREFIX: ""
|
||||||
TRAIN_N_PROMPT: ""
|
TRAIN_N_PROMPT: ""
|
||||||
RESOLUTION: 512
|
RESOLUTION: [512, 512]
|
||||||
MEMORY: 29000
|
MEMORY: 29000
|
||||||
EPOCHS: 200
|
EPOCHS: 200
|
||||||
SAVE_INTERVAL: 25
|
SAVE_INTERVAL: 25
|
||||||
@@ -66,7 +66,7 @@ META:
|
|||||||
TRAIN_BATCH_SIZE: 4
|
TRAIN_BATCH_SIZE: 4
|
||||||
TRAIN_PREFIX: ""
|
TRAIN_PREFIX: ""
|
||||||
TRAIN_N_PROMPT: ""
|
TRAIN_N_PROMPT: ""
|
||||||
RESOLUTION: 512
|
RESOLUTION: [512, 512]
|
||||||
MEMORY: 29000
|
MEMORY: 29000
|
||||||
EPOCHS: 50
|
EPOCHS: 50
|
||||||
SAVE_INTERVAL: 25
|
SAVE_INTERVAL: 25
|
||||||
@@ -135,6 +135,7 @@ SOLVER:
|
|||||||
MAX_EPOCHS: -1
|
MAX_EPOCHS: -1
|
||||||
NUM_FOLDS: 1
|
NUM_FOLDS: 1
|
||||||
ACCU_STEP: 1
|
ACCU_STEP: 1
|
||||||
|
EVAL_INTERVAL: -1
|
||||||
#
|
#
|
||||||
WORK_DIR:
|
WORK_DIR:
|
||||||
LOG_FILE: std_log.txt
|
LOG_FILE: std_log.txt
|
||||||
@@ -298,15 +299,44 @@ SOLVER:
|
|||||||
KEYS: [ 'image', 'prompt' ]
|
KEYS: [ 'image', 'prompt' ]
|
||||||
META_KEYS: [ 'data_key' ]
|
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:
|
TRAIN_HOOKS:
|
||||||
- NAME: BackwardHook
|
-
|
||||||
# GRADIENT_CLIP: 1.0
|
NAME: BackwardHook
|
||||||
PRIORITY: 0
|
PRIORITY: 0
|
||||||
- NAME: LogHook
|
-
|
||||||
LOG_INTERVAL: 50
|
NAME: LogHook
|
||||||
|
LOG_INTERVAL: 10
|
||||||
SHOW_GPU_MEM: True
|
SHOW_GPU_MEM: True
|
||||||
|
-
|
||||||
|
NAME: TensorboardLogHook
|
||||||
-
|
-
|
||||||
NAME: CheckpointHook
|
NAME: CheckpointHook
|
||||||
SAVE_LAST: True
|
|
||||||
INTERVAL: 10000
|
INTERVAL: 10000
|
||||||
PRIORITY: 200
|
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_SAMPLER: "ddim"
|
||||||
DEFAULT_SAMPLE_STEPS: 40
|
DEFAULT_SAMPLE_STEPS: 40
|
||||||
INFERENCE_N_PROMPT: ""
|
INFERENCE_N_PROMPT: ""
|
||||||
RESOLUTION: 768
|
RESOLUTION: [768, 768]
|
||||||
PARAS:
|
PARAS:
|
||||||
-
|
-
|
||||||
TRAIN_BATCH_SIZE: 2
|
TRAIN_BATCH_SIZE: 2
|
||||||
TRAIN_PREFIX: ""
|
TRAIN_PREFIX: ""
|
||||||
TRAIN_N_PROMPT: ""
|
TRAIN_N_PROMPT: ""
|
||||||
RESOLUTION: 768
|
RESOLUTION: [768, 768]
|
||||||
MEMORY: 29000
|
MEMORY: 29000
|
||||||
EPOCHS: 50
|
EPOCHS: 50
|
||||||
SAVE_INTERVAL: 25
|
SAVE_INTERVAL: 25
|
||||||
@@ -28,7 +28,7 @@ META:
|
|||||||
TRAIN_BATCH_SIZE: 2
|
TRAIN_BATCH_SIZE: 2
|
||||||
TRAIN_PREFIX: ""
|
TRAIN_PREFIX: ""
|
||||||
TRAIN_N_PROMPT: ""
|
TRAIN_N_PROMPT: ""
|
||||||
RESOLUTION: 768
|
RESOLUTION: [768, 768]
|
||||||
MEMORY: 29000
|
MEMORY: 29000
|
||||||
EPOCHS: 50
|
EPOCHS: 50
|
||||||
SAVE_INTERVAL: 25
|
SAVE_INTERVAL: 25
|
||||||
@@ -40,7 +40,7 @@ META:
|
|||||||
TRAIN_BATCH_SIZE: 2
|
TRAIN_BATCH_SIZE: 2
|
||||||
TRAIN_PREFIX: ""
|
TRAIN_PREFIX: ""
|
||||||
TRAIN_N_PROMPT: ""
|
TRAIN_N_PROMPT: ""
|
||||||
RESOLUTION: 768
|
RESOLUTION: [768, 768]
|
||||||
MEMORY: 29000
|
MEMORY: 29000
|
||||||
EPOCHS: 200
|
EPOCHS: 200
|
||||||
SAVE_INTERVAL: 25
|
SAVE_INTERVAL: 25
|
||||||
@@ -78,6 +78,7 @@ SOLVER:
|
|||||||
MAX_EPOCHS: -1
|
MAX_EPOCHS: -1
|
||||||
NUM_FOLDS: 1
|
NUM_FOLDS: 1
|
||||||
ACCU_STEP: 1
|
ACCU_STEP: 1
|
||||||
|
EVAL_INTERVAL: -1
|
||||||
#
|
#
|
||||||
WORK_DIR:
|
WORK_DIR:
|
||||||
LOG_FILE: std_log.txt
|
LOG_FILE: std_log.txt
|
||||||
@@ -240,15 +241,44 @@ SOLVER:
|
|||||||
KEYS: [ 'image', 'prompt' ]
|
KEYS: [ 'image', 'prompt' ]
|
||||||
META_KEYS: [ 'data_key' ]
|
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:
|
TRAIN_HOOKS:
|
||||||
- NAME: BackwardHook
|
-
|
||||||
# GRADIENT_CLIP: 1.0
|
NAME: BackwardHook
|
||||||
PRIORITY: 0
|
PRIORITY: 0
|
||||||
- NAME: LogHook
|
-
|
||||||
LOG_INTERVAL: 50
|
NAME: LogHook
|
||||||
|
LOG_INTERVAL: 10
|
||||||
SHOW_GPU_MEM: True
|
SHOW_GPU_MEM: True
|
||||||
|
-
|
||||||
|
NAME: TensorboardLogHook
|
||||||
-
|
-
|
||||||
NAME: CheckpointHook
|
NAME: CheckpointHook
|
||||||
SAVE_LAST: True
|
|
||||||
INTERVAL: 10000
|
INTERVAL: 10000
|
||||||
PRIORITY: 200
|
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]
|
image_size = [image_size, image_size]
|
||||||
assert isinstance(image_size, Iterable) and len(image_size) == 2
|
assert isinstance(image_size, Iterable) and len(image_size) == 2
|
||||||
|
|
||||||
prompt_file = cfg.PROMPT_FILE
|
if cfg.PROMPT_FILE is not None and cfg.PROMPT_FILE != '':
|
||||||
with FS.get_object(prompt_file) as local_data:
|
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 = [
|
rows = [
|
||||||
i.split(delimiter,
|
i.split(delimiter,
|
||||||
len(fields) - 1)
|
len(fields) - 1) for i in cfg.PROMPT_DATA
|
||||||
for i in local_data.decode('utf-8').strip().split('\n')
|
|
||||||
]
|
]
|
||||||
|
|
||||||
self.items = list()
|
self.items = list()
|
||||||
@@ -263,10 +269,9 @@ class Text2ImageDataset(BaseDataset):
|
|||||||
item['meta']['img_path'] = os.path.join(path_prefix, value)
|
item['meta']['img_path'] = os.path.join(path_prefix, value)
|
||||||
elif key in ['width', 'height']:
|
elif key in ['width', 'height']:
|
||||||
item['meta'][key] = int(value)
|
item['meta'][key] = int(value)
|
||||||
elif key != 'meta':
|
|
||||||
item[key] = value
|
|
||||||
else:
|
else:
|
||||||
continue
|
item['meta'][key] = value
|
||||||
|
|
||||||
self.items.append(item)
|
self.items.append(item)
|
||||||
if use_num > 0:
|
if use_num > 0:
|
||||||
self.items = self.items[:use_num]
|
self.items = self.items[:use_num]
|
||||||
|
|||||||
@@ -2,18 +2,23 @@
|
|||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
import copy
|
import copy
|
||||||
import os
|
import os
|
||||||
|
import warnings
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torchvision.transforms as TT
|
import torchvision.transforms as TT
|
||||||
from PIL.Image import Image
|
from PIL.Image import Image
|
||||||
from swift import SwiftModel
|
|
||||||
|
|
||||||
from scepter.modules.model.registry import TUNERS
|
from scepter.modules.model.registry import TUNERS
|
||||||
from scepter.modules.utils.config import Config
|
from scepter.modules.utils.config import Config
|
||||||
from scepter.modules.utils.distribute import we
|
from scepter.modules.utils.distribute import we
|
||||||
from scepter.modules.utils.file_system import FS
|
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():
|
class ControlInference():
|
||||||
def __init__(self, logger=None):
|
def __init__(self, logger=None):
|
||||||
|
|||||||
@@ -395,7 +395,7 @@ class DiffusionInference():
|
|||||||
return self.first_stage_model['paras']['scale_factor'] * z
|
return self.first_stage_model['paras']['scale_factor'] * z
|
||||||
|
|
||||||
def decode_first_stage(self, z):
|
def decode_first_stage(self, z):
|
||||||
_, dtype = self.get_function_info(self.first_stage_model, 'encode')
|
_, dtype = self.get_function_info(self.first_stage_model, 'decode')
|
||||||
with torch.autocast('cuda',
|
with torch.autocast('cuda',
|
||||||
enabled=dtype == 'float16',
|
enabled=dtype == 'float16',
|
||||||
dtype=getattr(torch, dtype)):
|
dtype=getattr(torch, dtype)):
|
||||||
@@ -474,10 +474,14 @@ class DiffusionInference():
|
|||||||
enabled=dtype == 'float16',
|
enabled=dtype == 'float16',
|
||||||
dtype=getattr(torch, dtype)):
|
dtype=getattr(torch, dtype)):
|
||||||
if self.tokenizer:
|
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),
|
context = getattr(get_model(self.cond_stage_model),
|
||||||
function_name)(batch['tokens'])
|
function_name)(batch['tokens'])
|
||||||
null_context = getattr(get_model(self.cond_stage_model),
|
null_context = getattr(get_model(self.cond_stage_model),
|
||||||
function_name)(batch_uc['tokens'])
|
function_name)(batch_uc['tokens'])
|
||||||
|
|
||||||
else:
|
else:
|
||||||
context = getattr(get_model(self.cond_stage_model),
|
context = getattr(get_model(self.cond_stage_model),
|
||||||
function_name)(batch)
|
function_name)(batch)
|
||||||
@@ -558,12 +562,13 @@ class DiffusionInference():
|
|||||||
seed=seed,
|
seed=seed,
|
||||||
condition_fn=None,
|
condition_fn=None,
|
||||||
clamp=None,
|
clamp=None,
|
||||||
|
sharpness=value_input.get('sharpness', 0.0),
|
||||||
percentile=None,
|
percentile=None,
|
||||||
t_max=None,
|
t_max=None,
|
||||||
t_min=None,
|
t_min=None,
|
||||||
discard_penultimate_step=None,
|
discard_penultimate_step=None,
|
||||||
intermediate_callback=intermediate_callback,
|
intermediate_callback=intermediate_callback,
|
||||||
cat_uc=cat_uc,
|
cat_uc=value_input.get('cat_uc', cat_uc),
|
||||||
**kwargs)
|
**kwargs)
|
||||||
|
|
||||||
self.dynamic_unload(self.diffusion_model,
|
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 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
from scepter.modules.model.backbone.autoencoder.ae_module import (Decoder,
|
from scepter.modules.model.backbone.autoencoder.ae_module import (Decoder,
|
||||||
Encoder)
|
Encoder,
|
||||||
|
RDecoder)
|
||||||
|
|||||||
@@ -3,6 +3,7 @@
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
from einops import repeat
|
||||||
from torch.utils.checkpoint import checkpoint
|
from torch.utils.checkpoint import checkpoint
|
||||||
|
|
||||||
from scepter.modules.model.backbone.autoencoder.ae_utils import (
|
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
|
# compute in_ch_mult, block_in and curr_res at lowest res
|
||||||
block_in = self.ch * self.ch_mult[self.num_resolutions - 1]
|
block_in = self.ch * self.ch_mult[self.num_resolutions - 1]
|
||||||
|
self.block_in = block_in
|
||||||
curr_res = 1
|
curr_res = 1
|
||||||
# z to block_in
|
# z to block_in
|
||||||
self.conv_in = torch.nn.Conv2d(self.z_channels,
|
self.conv_in = torch.nn.Conv2d(self.z_channels,
|
||||||
@@ -340,3 +342,48 @@ class Decoder(BaseModel):
|
|||||||
__class__.__name__,
|
__class__.__name__,
|
||||||
Decoder.para_dict,
|
Decoder.para_dict,
|
||||||
set_name=True)
|
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 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# 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
|
import torch.nn as nn
|
||||||
|
|
||||||
from scepter.modules.model.backbone.unet.unet_utils import (
|
from scepter.modules.model.backbone.unet.unet_utils import (
|
||||||
Downsample, ResBlock, SpatialTransformer, Timestep,
|
BasicTransformerBlock, Downsample, ResBlock, SpatialTransformer,
|
||||||
TimestepEmbedSequential, Upsample, conv_nd, linear, normalization,
|
SpatialTransformerV2, Timestep, TimestepEmbedSequential,
|
||||||
|
TransformerBlockV2, Upsample, conv_nd, linear, normalization,
|
||||||
timestep_embedding, zero_module)
|
timestep_embedding, zero_module)
|
||||||
from scepter.modules.model.base_model import BaseModel
|
from scepter.modules.model.base_model import BaseModel
|
||||||
from scepter.modules.model.registry import BACKBONES
|
from scepter.modules.model.registry import BACKBONES
|
||||||
@@ -951,3 +952,443 @@ class DiffusionUNetXL(DiffusionUNet):
|
|||||||
__class__.__name__,
|
__class__.__name__,
|
||||||
DiffusionUNetXL.para_dict,
|
DiffusionUNetXL.para_dict,
|
||||||
set_name=True)
|
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
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
import torchvision.transforms.functional as TF
|
||||||
from einops import rearrange, repeat
|
from einops import rearrange, repeat
|
||||||
from packaging import version
|
from packaging import version
|
||||||
|
|
||||||
@@ -171,12 +172,14 @@ class TimestepEmbedSequential(nn.Sequential, TimestepBlock):
|
|||||||
A sequential module that passes timestep embeddings to the children that
|
A sequential module that passes timestep embeddings to the children that
|
||||||
support it as an extra input.
|
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:
|
for layer in self:
|
||||||
if isinstance(layer, TimestepBlock):
|
if isinstance(layer, TimestepBlock):
|
||||||
x = layer(x, emb)
|
x = layer(x, emb)
|
||||||
elif isinstance(layer, SpatialTransformer):
|
elif isinstance(layer, SpatialTransformer):
|
||||||
x = layer(x, context)
|
x = layer(x, context)
|
||||||
|
elif isinstance(layer, SpatialTransformerV2):
|
||||||
|
x = layer(x, context, **kwargs)
|
||||||
elif isinstance(layer, Upsample):
|
elif isinstance(layer, Upsample):
|
||||||
x = layer(x, target_size)
|
x = layer(x, target_size)
|
||||||
else:
|
else:
|
||||||
@@ -864,6 +867,92 @@ class MemoryEfficientCrossAttention(nn.Module):
|
|||||||
return self.to_out(out)
|
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):
|
class BasicTransformerBlock(nn.Module):
|
||||||
def __init__(self,
|
def __init__(self,
|
||||||
dim,
|
dim,
|
||||||
@@ -908,6 +997,65 @@ class BasicTransformerBlock(nn.Module):
|
|||||||
return x
|
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):
|
class SpatialTransformer(nn.Module):
|
||||||
"""
|
"""
|
||||||
Transformer block for image-like data.
|
Transformer block for image-like data.
|
||||||
@@ -1003,3 +1151,108 @@ class SpatialTransformer(nn.Module):
|
|||||||
if not self.use_linear:
|
if not self.use_linear:
|
||||||
x = self.proj_out(x)
|
x = self.proj_out(x)
|
||||||
return x + x_in
|
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 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# 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 implementations of vivit as https://arxiv.org/abs/2103.15691.
|
||||||
The following setting alined the proposed model in the paper above.
|
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)
|
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()
|
@BACKBONES.register_class()
|
||||||
class VideoTransformer(nn.Module):
|
class VideoTransformer(nn.Module):
|
||||||
|
|||||||
@@ -1,7 +1,10 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
|
|
||||||
from scepter.modules.model.embedder.embedder import (ConcatTimestepEmbedderND,
|
from scepter.modules.model.embedder.embedder import (ConcatTimestepEmbedderND,
|
||||||
FrozenCLIPEmbedder,
|
FrozenCLIPEmbedder,
|
||||||
FrozenOpenCLIPEmbedder,
|
FrozenOpenCLIPEmbedder,
|
||||||
FrozenOpenCLIPEmbedder2,
|
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 scepter.modules.utils.file_system import FS
|
||||||
|
|
||||||
from .base_embedder import BaseEmbedder
|
from .base_embedder import BaseEmbedder
|
||||||
|
from .resampler import Resampler
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from transformers import CLIPTextModel, CLIPTokenizer
|
from transformers import CLIPTextModel, CLIPTokenizer, CLIPVisionModelWithProjection
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
warnings.warn(
|
warnings.warn(
|
||||||
f'Import transformers error, please deal with this problem: {e}')
|
f'Import transformers error, please deal with this problem: {e}')
|
||||||
@@ -513,6 +514,93 @@ class ConcatTimestepEmbedderND(BaseEmbedder):
|
|||||||
set_name=True)
|
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()
|
@EMBEDDERS.register_class()
|
||||||
class GeneralConditioner(BaseEmbedder):
|
class GeneralConditioner(BaseEmbedder):
|
||||||
OUTPUT_DIM2KEYS = {2: 'y', 3: 'crossattn', 4: 'concat', 5: 'concat'}
|
OUTPUT_DIM2KEYS = {2: 'y', 3: 'crossattn', 4: 'concat', 5: 'concat'}
|
||||||
@@ -598,42 +686,58 @@ class GeneralConditioner(BaseEmbedder):
|
|||||||
with embedding_context():
|
with embedding_context():
|
||||||
if hasattr(embedder, 'input_key') and (embedder.input_key
|
if hasattr(embedder, 'input_key') and (embedder.input_key
|
||||||
is not None):
|
is not None):
|
||||||
|
if embedder.input_key not in batch:
|
||||||
|
continue
|
||||||
if embedder.legacy_ucg_val is not None:
|
if embedder.legacy_ucg_val is not None:
|
||||||
batch = self.possibly_get_ucg_val(embedder, batch)
|
batch = self.possibly_get_ucg_val(embedder, batch)
|
||||||
emb_out = embedder(batch[embedder.input_key])
|
emb_out = embedder(batch[embedder.input_key])
|
||||||
elif hasattr(embedder, 'input_keys'):
|
elif hasattr(embedder, 'input_keys'):
|
||||||
|
if any([k not in batch for k in embedder.input_keys]):
|
||||||
|
continue
|
||||||
emb_out = embedder(
|
emb_out = embedder(
|
||||||
*[batch[k] for k in embedder.input_keys])
|
*[batch[k] for k in embedder.input_keys])
|
||||||
assert isinstance(
|
|
||||||
emb_out, (torch.Tensor, list, tuple)
|
if isinstance(emb_out, dict):
|
||||||
), f'encoder outputs must be tensors or a sequence, but got {type(emb_out)}'
|
for key, val in emb_out.items():
|
||||||
if not isinstance(emb_out, (list, tuple)):
|
if key in output:
|
||||||
emb_out = [emb_out]
|
assert key in self.KEY2CATDIM
|
||||||
for emb in emb_out:
|
output[key] = torch.cat([output[key], val],
|
||||||
# print("emb.shape", emb.shape)
|
dim=self.KEY2CATDIM[key])
|
||||||
# print("emb.input_keys", embedder.input_keys)
|
else:
|
||||||
out_key = self.OUTPUT_DIM2KEYS[emb.dim()]
|
output[key] = val
|
||||||
if embedder.ucg_rate > 0.0 and embedder.legacy_ucg_val is None:
|
else:
|
||||||
emb = (expand_dims_like(
|
assert isinstance(
|
||||||
torch.bernoulli(
|
emb_out, (torch.Tensor, list, tuple)
|
||||||
(1.0 - embedder.ucg_rate) *
|
), f'encoder outputs must be tensors or a sequence, but got {type(emb_out)}'
|
||||||
torch.ones(emb.shape[0], device=emb.device)),
|
|
||||||
emb,
|
if not isinstance(emb_out, (list, tuple)):
|
||||||
) * emb)
|
emb_out = [emb_out]
|
||||||
if (hasattr(embedder, 'input_keys')):
|
|
||||||
if np.sum(
|
for emb in emb_out:
|
||||||
np.array([
|
out_key = self.OUTPUT_DIM2KEYS[emb.dim()]
|
||||||
key in force_zero_embeddings
|
|
||||||
for key in embedder.input_keys
|
if embedder.ucg_rate > 0.0 and embedder.legacy_ucg_val is None:
|
||||||
])) > 0:
|
emb = (expand_dims_like(
|
||||||
emb = torch.zeros_like(emb)
|
torch.bernoulli(
|
||||||
if out_key in output:
|
(1.0 - embedder.ucg_rate) *
|
||||||
output[out_key] = torch.cat((output[out_key], emb),
|
torch.ones(emb.shape[0], device=emb.device)),
|
||||||
self.KEY2CATDIM[out_key])
|
emb,
|
||||||
else:
|
) * emb)
|
||||||
output[out_key] = emb
|
|
||||||
# if "y" in output:
|
if (hasattr(embedder, 'input_keys')):
|
||||||
# print("out.shape", output["y"].shape)
|
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
|
return output
|
||||||
|
|
||||||
def get_unconditional_conditioning(self,
|
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 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
from scepter.modules.model.head.classifier_head import (
|
from scepter.modules.model.head.classifier_head import (ClassifierHead,
|
||||||
ClassifierHead, CosineLinearHead, TransformerHead, TransformerHeadx2,
|
CosineLinearHead,
|
||||||
VideoClassifierHead, VideoClassifierHeadx2)
|
TransformerHead,
|
||||||
|
TransformerHeadx2,
|
||||||
|
VideoClassifierHead,
|
||||||
|
VideoClassifierHeadx2)
|
||||||
|
|||||||
@@ -52,7 +52,7 @@ class DiagonalGaussianDistribution(object):
|
|||||||
dim=dims)
|
dim=dims)
|
||||||
|
|
||||||
def mode(self):
|
def mode(self):
|
||||||
print('*** use DiagonalGaussianDistribution.mode() ***')
|
# print('*** use DiagonalGaussianDistribution.mode() ***')
|
||||||
return self.mean
|
return self.mean
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -13,8 +13,7 @@ from .schedules import karras_schedule
|
|||||||
from .solvers import (sample_ddim, sample_dpm_2, sample_dpm_2_ancestral,
|
from .solvers import (sample_ddim, sample_dpm_2, sample_dpm_2_ancestral,
|
||||||
sample_dpmpp_2m, sample_dpmpp_2m_sde,
|
sample_dpmpp_2m, sample_dpmpp_2m_sde,
|
||||||
sample_dpmpp_2s_ancestral, sample_dpmpp_sde,
|
sample_dpmpp_2s_ancestral, sample_dpmpp_sde,
|
||||||
sample_euler, sample_euler_ancestral, sample_heun,
|
sample_euler, sample_euler_ancestral, sample_heun)
|
||||||
sample_img2img_euler, sample_img2img_euler_ancestral)
|
|
||||||
|
|
||||||
__all__ = ['GaussianDiffusion']
|
__all__ = ['GaussianDiffusion']
|
||||||
|
|
||||||
@@ -27,6 +26,148 @@ def _i(tensor, t, x):
|
|||||||
return tensor[t.to(tensor.device)].view(shape).to(x.device)
|
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):
|
class GaussianDiffusion(object):
|
||||||
def __init__(self, sigmas, prediction_type='eps'):
|
def __init__(self, sigmas, prediction_type='eps'):
|
||||||
assert prediction_type in {'x0', 'eps', 'v'}
|
assert prediction_type in {'x0', 'eps', 'v'}
|
||||||
@@ -53,6 +194,7 @@ class GaussianDiffusion(object):
|
|||||||
guide_scale=None,
|
guide_scale=None,
|
||||||
guide_rescale=None,
|
guide_rescale=None,
|
||||||
clamp=None,
|
clamp=None,
|
||||||
|
sharpness=0.0,
|
||||||
percentile=None,
|
percentile=None,
|
||||||
cat_uc=False,
|
cat_uc=False,
|
||||||
**kwargs):
|
**kwargs):
|
||||||
@@ -79,54 +221,99 @@ class GaussianDiffusion(object):
|
|||||||
|
|
||||||
# prediction
|
# prediction
|
||||||
if guide_scale is None:
|
if guide_scale is None:
|
||||||
assert isinstance(model_kwargs, dict)
|
if isinstance(model_kwargs, dict):
|
||||||
out = model(xt, t=t, **model_kwargs, **kwargs)
|
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:
|
else:
|
||||||
# classifier-free guidance (arXiv:2207.12598)
|
# classifier-free guidance (arXiv:2207.12598)
|
||||||
# model_kwargs[0]: conditional kwargs
|
# model_kwargs[0]: conditional kwargs
|
||||||
# model_kwargs[1]: non-conditional kwargs
|
# model_kwargs[1]: non-conditional kwargs
|
||||||
assert isinstance(model_kwargs, list) and len(model_kwargs) == 2
|
assert isinstance(model_kwargs, list) and len(model_kwargs) >= 2
|
||||||
|
if isinstance(guide_scale, float) or isinstance(guide_scale, int):
|
||||||
if guide_scale == 1.:
|
assert len(model_kwargs) == 2
|
||||||
out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
if guide_scale == 1.:
|
||||||
else:
|
out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
||||||
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)
|
|
||||||
else:
|
else:
|
||||||
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
if cat_uc:
|
||||||
u_out = model(xt, t=t, **model_kwargs[1], **kwargs)
|
|
||||||
out = u_out + guide_scale * (y_out - u_out)
|
|
||||||
|
|
||||||
# rescale the output according to arXiv:2305.08891
|
def parse_model_kwargs(prev_value, value):
|
||||||
if guide_rescale is not None:
|
if isinstance(value, torch.Tensor):
|
||||||
assert guide_rescale >= 0 and guide_rescale <= 1
|
prev_value = torch.cat([prev_value, value],
|
||||||
ratio = (y_out.flatten(1).std(dim=1) /
|
dim=0)
|
||||||
(out.flatten(1).std(dim=1) +
|
elif isinstance(value, dict):
|
||||||
1e-12)).view((-1, ) + (1, ) * (y_out.ndim - 1))
|
for k, v in value.items():
|
||||||
out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0
|
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
|
# compute x0
|
||||||
if self.prediction_type == 'x0':
|
if self.prediction_type == 'x0':
|
||||||
x0 = out
|
x0 = out
|
||||||
@@ -197,6 +384,7 @@ class GaussianDiffusion(object):
|
|||||||
guide_scale=None,
|
guide_scale=None,
|
||||||
guide_rescale=None,
|
guide_rescale=None,
|
||||||
clamp=None,
|
clamp=None,
|
||||||
|
sharpness=0.0,
|
||||||
percentile=None,
|
percentile=None,
|
||||||
solver='euler_a',
|
solver='euler_a',
|
||||||
steps=20,
|
steps=20,
|
||||||
@@ -209,12 +397,16 @@ class GaussianDiffusion(object):
|
|||||||
seed=-1,
|
seed=-1,
|
||||||
intermediate_callback=None,
|
intermediate_callback=None,
|
||||||
cat_uc=False,
|
cat_uc=False,
|
||||||
|
add_noise=False,
|
||||||
|
free_steps=None,
|
||||||
|
step_offset=None,
|
||||||
**kwargs):
|
**kwargs):
|
||||||
# sanity check
|
# sanity check
|
||||||
assert isinstance(steps, (int, torch.LongTensor))
|
assert isinstance(steps, (int, torch.LongTensor))
|
||||||
assert t_max is None or (t_max > 0 and t_max <= self.num_timesteps - 1)
|
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 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 discard_penultimate_step in (None, True, False)
|
||||||
assert return_intermediate in (None, 'x0', 'xt')
|
assert return_intermediate in (None, 'x0', 'xt')
|
||||||
|
|
||||||
@@ -255,17 +447,50 @@ class GaussianDiffusion(object):
|
|||||||
def model_fn(xt, sigma):
|
def model_fn(xt, sigma):
|
||||||
# denoising
|
# denoising
|
||||||
t = self._sigma_to_t(sigma).repeat(len(xt)).round().long()
|
t = self._sigma_to_t(sigma).repeat(len(xt)).round().long()
|
||||||
x0 = self.denoise(xt,
|
|
||||||
t,
|
if isinstance(
|
||||||
None,
|
model_kwargs[0]['cond'], dict) and \
|
||||||
model,
|
'tar_x0' in model_kwargs[0]['cond'] and \
|
||||||
model_kwargs,
|
'tar_mask_latent' in model_kwargs[0]['cond']:
|
||||||
guide_scale,
|
tar_x0 = model_kwargs[0]['cond']['tar_x0']
|
||||||
guide_rescale,
|
tar_mask = model_kwargs[0]['cond']['tar_mask_latent']
|
||||||
clamp,
|
|
||||||
percentile,
|
tar_xt = self.diffuse(x0=tar_x0, t=t)
|
||||||
cat_uc=cat_uc,
|
xt = tar_xt * (1.0 - tar_mask) + xt * tar_mask
|
||||||
**kwargs)[-2]
|
|
||||||
|
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
|
# collect intermediate outputs
|
||||||
if return_intermediate == 'xt':
|
if return_intermediate == 'xt':
|
||||||
@@ -291,10 +516,14 @@ class GaussianDiffusion(object):
|
|||||||
elif discretization == 'trailing':
|
elif discretization == 'trailing':
|
||||||
steps = torch.arange(t_max, t_min - 1,
|
steps = torch.arange(t_max, t_min - 1,
|
||||||
-((t_max - t_min + 1) / steps))
|
-((t_max - t_min + 1) / steps))
|
||||||
|
elif discretization == 'free':
|
||||||
|
steps = torch.tensor(free_steps)
|
||||||
else:
|
else:
|
||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
f'{discretization} discretization not implemented')
|
f'{discretization} discretization not implemented')
|
||||||
steps = steps.clamp_(t_min, t_max)
|
steps = steps.clamp_(t_min, t_max)
|
||||||
|
elif isinstance(steps, list):
|
||||||
|
steps = torch.tensor(steps)
|
||||||
steps = torch.as_tensor(steps,
|
steps = torch.as_tensor(steps,
|
||||||
dtype=torch.float32,
|
dtype=torch.float32,
|
||||||
device=noise.device)
|
device=noise.device)
|
||||||
@@ -335,6 +564,23 @@ class GaussianDiffusion(object):
|
|||||||
if discard_penultimate_step:
|
if discard_penultimate_step:
|
||||||
sigmas = torch.cat([sigmas[:-2], sigmas[-1:]])
|
sigmas = torch.cat([sigmas[:-2], sigmas[-1:]])
|
||||||
kwargs['seed'] = seed
|
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
|
# sampling
|
||||||
x0 = solver_fn(noise,
|
x0 = solver_fn(noise,
|
||||||
model_fn,
|
model_fn,
|
||||||
@@ -373,168 +619,6 @@ class GaussianDiffusion(object):
|
|||||||
| torch.isinf(log_sigma)] = float('inf')
|
| torch.isinf(log_sigma)] = float('inf')
|
||||||
return log_sigma.exp()
|
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):
|
def extract_into_tensor(a, t, x_shape):
|
||||||
b, *_ = t.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 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
|
|
||||||
from scepter.modules.opt.optimizers.official_optimizers import (
|
from scepter.modules.opt.optimizers.official_optimizers import (ASGD, LBFGS,
|
||||||
ASGD, LBFGS, SGD, Adadelta, Adagrad, Adam, Adamax, AdamW, RMSprop, Rprop,
|
SGD, Adadelta,
|
||||||
SparseAdam)
|
Adagrad, Adam,
|
||||||
|
Adamax, AdamW,
|
||||||
|
RMSprop, Rprop,
|
||||||
|
SparseAdam)
|
||||||
from scepter.modules.opt.optimizers.registry import OPTIMIZERS
|
from scepter.modules.opt.optimizers.registry import OPTIMIZERS
|
||||||
|
|||||||
@@ -333,12 +333,12 @@ class LatentDiffusionSolver(BaseSolver):
|
|||||||
data_iter = iter(self.datas[self._mode].dataloader)
|
data_iter = iter(self.datas[self._mode].dataloader)
|
||||||
self.print_memory_status()
|
self.print_memory_status()
|
||||||
for step in range(self.max_steps):
|
for step in range(self.max_steps):
|
||||||
if 'eval' in self._mode_set and (step % self.eval_interval == 0
|
if 'eval' in self._mode_set and (self.eval_interval > 0 and
|
||||||
or step == self.max_steps - 1):
|
step % self.eval_interval == 0):
|
||||||
self.run_eval()
|
self.run_eval()
|
||||||
self.train_mode()
|
self.train_mode()
|
||||||
self.before_iter(self.hooks_dict[self._mode])
|
|
||||||
batch_data = next(data_iter)
|
batch_data = next(data_iter)
|
||||||
|
self.before_iter(self.hooks_dict[self._mode])
|
||||||
if self.sample_args:
|
if self.sample_args:
|
||||||
batch_data.update(self.sample_args.get_lowercase_dict())
|
batch_data.update(self.sample_args.get_lowercase_dict())
|
||||||
if 'meta' in batch_data:
|
if 'meta' in batch_data:
|
||||||
@@ -364,6 +364,9 @@ class LatentDiffusionSolver(BaseSolver):
|
|||||||
self.after_iter(self.hooks_dict[self._mode])
|
self.after_iter(self.hooks_dict[self._mode])
|
||||||
if we.debug:
|
if we.debug:
|
||||||
self.print_trainable_params_status(prefix='model.')
|
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])
|
self.after_all_iter(self.hooks_dict[self._mode])
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
|
|||||||
@@ -4,6 +4,7 @@
|
|||||||
from scepter.modules.solver.hooks.backward import BackwardHook
|
from scepter.modules.solver.hooks.backward import BackwardHook
|
||||||
from scepter.modules.solver.hooks.checkpoint import CheckpointHook
|
from scepter.modules.solver.hooks.checkpoint import CheckpointHook
|
||||||
from scepter.modules.solver.hooks.data_probe import ProbeDataHook
|
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.hook import Hook
|
||||||
from scepter.modules.solver.hooks.log import LogHook, TensorboardLogHook
|
from scepter.modules.solver.hooks.log import LogHook, TensorboardLogHook
|
||||||
from scepter.modules.solver.hooks.lr import LrHook
|
from scepter.modules.solver.hooks.lr import LrHook
|
||||||
@@ -47,5 +48,6 @@ after solve:
|
|||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
'HOOKS', 'BackwardHook', 'CheckpointHook', 'Hook', 'LrHook', 'LogHook',
|
'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:
|
if self.save_last and solver.total_iter == solver.max_steps - 1:
|
||||||
with FS.get_fs_client(save_path) as client:
|
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)
|
client.make_link(last_path, save_path)
|
||||||
self.last_ckpt = save_path
|
self.last_ckpt = save_path
|
||||||
|
|
||||||
|
|||||||
@@ -30,7 +30,10 @@ class ProbeDataHook(Hook):
|
|||||||
def __init__(self, cfg, logger=None):
|
def __init__(self, cfg, logger=None):
|
||||||
super(ProbeDataHook, self).__init__(cfg, logger=logger)
|
super(ProbeDataHook, self).__init__(cfg, logger=logger)
|
||||||
self.priority = cfg.get('PRIORITY', _DEFAULT_PROBE_PRIORITY)
|
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):
|
def before_all_iter(self, solver):
|
||||||
pass
|
pass
|
||||||
@@ -39,19 +42,23 @@ class ProbeDataHook(Hook):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
def after_iter(self, solver):
|
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
|
probe_dict = solver.probe_data
|
||||||
if we.rank == 0:
|
if we.rank == 0:
|
||||||
save_folder = os.path.join(
|
save_folder = os.path.join(
|
||||||
solver.work_dir,
|
solver.work_dir,
|
||||||
f'{solver.mode}_probe/step_{solver.total_iter}')
|
f'{solver.mode}_probe/{self.save_name_prefix}-{solver.total_iter}'
|
||||||
|
)
|
||||||
ret_data = {}
|
ret_data = {}
|
||||||
for k, v in probe_dict.items():
|
for k, v in probe_dict.items():
|
||||||
ret_one = v.to_log(
|
if self.save_probe_prefix is not None:
|
||||||
os.path.join(
|
ret_prefix = os.path.join(save_folder,
|
||||||
|
self.save_probe_prefix)
|
||||||
|
else:
|
||||||
|
ret_prefix = os.path.join(
|
||||||
save_folder,
|
save_folder,
|
||||||
k.replace('/', '_') +
|
k.replace('/', '_') + f'_step_{solver.total_iter}')
|
||||||
f'_step_{solver.total_iter}'))
|
ret_one = v.to_log(ret_prefix)
|
||||||
if (isinstance(ret_one, list)
|
if (isinstance(ret_one, list)
|
||||||
or isinstance(ret_one, dict)) and len(ret_one) < 1:
|
or isinstance(ret_one, dict)) and len(ret_one) < 1:
|
||||||
continue
|
continue
|
||||||
@@ -71,13 +78,19 @@ class ProbeDataHook(Hook):
|
|||||||
if we.rank == 0:
|
if we.rank == 0:
|
||||||
step = solver._total_iter[
|
step = solver._total_iter[
|
||||||
'train'] if 'train' in solver._total_iter else 0
|
'train'] if 'train' in solver._total_iter else 0
|
||||||
save_folder = os.path.join(solver.work_dir,
|
save_folder = os.path.join(
|
||||||
f'{solver.mode}_probe/step_{step}')
|
solver.work_dir,
|
||||||
|
f'{solver.mode}_probe/{self.save_name_prefix}-{step}')
|
||||||
ret_data = {}
|
ret_data = {}
|
||||||
for k, v in probe_dict.items():
|
for k, v in probe_dict.items():
|
||||||
ret_one = v.to_log(
|
if self.save_probe_prefix is not None:
|
||||||
os.path.join(save_folder,
|
ret_prefix = os.path.join(save_folder,
|
||||||
k.replace('/', '_') + f'_step_{step}'))
|
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)
|
if (isinstance(ret_one, list)
|
||||||
or isinstance(ret_one, dict)) and len(ret_one) < 1:
|
or isinstance(ret_one, dict)) and len(ret_one) < 1:
|
||||||
continue
|
continue
|
||||||
@@ -87,6 +100,16 @@ class ProbeDataHook(Hook):
|
|||||||
json.dump(ret_data,
|
json.dump(ret_data,
|
||||||
open(local_path, 'w'),
|
open(local_path, 'w'),
|
||||||
ensure_ascii=False)
|
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()
|
solver.clear_probe()
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
barrier()
|
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
|
import torchvision.transforms.functional as TF
|
||||||
|
|
||||||
from scepter.modules.transform.registry import TRANSFORMS
|
from scepter.modules.transform.registry import TRANSFORMS
|
||||||
from scepter.modules.transform.utils import (
|
from scepter.modules.transform.utils import (BACKEND_CV2, BACKEND_PILLOW,
|
||||||
BACKEND_CV2, BACKEND_PILLOW, BACKEND_TORCHVISION, INPUT_CV2_TYPE_WARNING,
|
BACKEND_TORCHVISION,
|
||||||
INPUT_PIL_TYPE_WARNING, INPUT_TENSOR_TYPE_WARNING, INTERPOLATION_STYLE,
|
INPUT_CV2_TYPE_WARNING,
|
||||||
INTERPOLATION_STYLE_CV2, TORCHVISION_CAPABILITY, is_cv2_image,
|
INPUT_PIL_TYPE_WARNING,
|
||||||
is_pil_image, is_tensor)
|
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
|
from scepter.modules.utils.config import dict_to_yaml
|
||||||
|
|
||||||
if TORCHVISION_CAPABILITY:
|
if TORCHVISION_CAPABILITY:
|
||||||
|
|||||||
@@ -6,13 +6,15 @@ import os
|
|||||||
import cv2
|
import cv2
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
from PIL import Image
|
from PIL import Image, ImageFile
|
||||||
|
|
||||||
from scepter.modules.transform.registry import TRANSFORMS
|
from scepter.modules.transform.registry import TRANSFORMS
|
||||||
from scepter.modules.utils.config import dict_to_yaml
|
from scepter.modules.utils.config import dict_to_yaml
|
||||||
from scepter.modules.utils.distribute import we
|
from scepter.modules.utils.distribute import we
|
||||||
from scepter.modules.utils.file_system import DATA_FS as FS
|
from scepter.modules.utils.file_system import DATA_FS as FS
|
||||||
|
|
||||||
|
ImageFile.LOAD_TRUNCATED_IMAGES = True
|
||||||
|
|
||||||
|
|
||||||
def pillow_convert(image, rgb_order):
|
def pillow_convert(image, rgb_order):
|
||||||
if image.mode != rgb_order:
|
if image.mode != rgb_order:
|
||||||
|
|||||||
@@ -603,3 +603,6 @@ class Config(object):
|
|||||||
return cfg_new
|
return cfg_new
|
||||||
else:
|
else:
|
||||||
return cfg
|
return cfg
|
||||||
|
|
||||||
|
def pop(self, name):
|
||||||
|
self.cfg_dict.pop(name)
|
||||||
|
|||||||
@@ -4,10 +4,10 @@ import io
|
|||||||
from io import BytesIO
|
from io import BytesIO
|
||||||
|
|
||||||
import onnx
|
import onnx
|
||||||
|
import onnxruntime
|
||||||
import torch
|
import torch
|
||||||
from torch.onnx import OperatorExportTypes
|
from torch.onnx import OperatorExportTypes
|
||||||
|
|
||||||
import onnxruntime
|
|
||||||
from scepter.modules.utils.distribute import we
|
from scepter.modules.utils.distribute import we
|
||||||
|
|
||||||
type_map = {
|
type_map = {
|
||||||
|
|||||||
@@ -293,6 +293,9 @@ class LocalFs(BaseFs):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
def get_logging_handler(self, target_logging_path):
|
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)
|
return logging.FileHandler(target_logging_path)
|
||||||
|
|
||||||
def put_dir_from_local_dir(self,
|
def put_dir_from_local_dir(self,
|
||||||
@@ -303,11 +306,11 @@ class LocalFs(BaseFs):
|
|||||||
target_dir = self.reconstruct_path(target_dir)
|
target_dir = self.reconstruct_path(target_dir)
|
||||||
if local_dir == target_dir:
|
if local_dir == target_dir:
|
||||||
return True
|
return True
|
||||||
# cp -f local_dir/* target_dir/*
|
# # cp -f local_dir/* target_dir/*
|
||||||
if not osp.exists(target_dir):
|
# if not osp.exists(target_dir):
|
||||||
status = os.system(f'mkdir -p {target_dir}')
|
# status = os.system(f'mkdir -p {target_dir}')
|
||||||
if status != 0:
|
# if status != 0:
|
||||||
return False
|
# return False
|
||||||
try:
|
try:
|
||||||
shutil.copytree(local_dir, target_dir, symlinks=True)
|
shutil.copytree(local_dir, target_dir, symlinks=True)
|
||||||
except Exception:
|
except Exception:
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
|
import io
|
||||||
import os
|
import os
|
||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
@@ -327,6 +328,10 @@ class FileSystem(object):
|
|||||||
target_path_list):
|
target_path_list):
|
||||||
if local_path is None or target_path is None:
|
if local_path is None or target_path is None:
|
||||||
flg = False
|
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):
|
elif self.exists(local_path):
|
||||||
local_cache = self.get_from(local_path,
|
local_cache = self.get_from(local_path,
|
||||||
local_path + f'{time.time()}',
|
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,
|
url = FS.get_url(one_path,
|
||||||
lifecycle=3600 * 365 * 24).replace(
|
lifecycle=3600 * 365 * 24).replace(
|
||||||
'.oss-internal.aliyun-inc.',
|
'.oss-internal.aliyun-inc.',
|
||||||
'.oss.aliyuncs.')
|
'.oss.aliyuncs.').replace(
|
||||||
|
'-internal', '')
|
||||||
one_rank += (
|
one_rank += (
|
||||||
f'<td align="center"><input type="image" src="{url}" >'
|
f'<td align="center"><input type="image" src="{url}" >'
|
||||||
f'<br><font size="4"><strong>{save_id}-{idx}|{one_label}<strong></font><br/></td>'
|
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.control_ui import ControlUI
|
||||||
from scepter.studio.inference.inference_ui.diffusion_ui import DiffusionUI
|
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.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.mantra_ui import MantraUI
|
||||||
from scepter.studio.inference.inference_ui.model_manage_ui import ModelManageUI
|
from scepter.studio.inference.inference_ui.model_manage_ui import ModelManageUI
|
||||||
from scepter.studio.inference.inference_ui.refiner_ui import RefinerUI
|
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
|
from scepter.studio.utils.env import init_env
|
||||||
|
|
||||||
UI_MAP = [('diffusion', DiffusionUI), ('mantra', MantraUI), ('tuner', TunerUI),
|
UI_MAP = [('diffusion', DiffusionUI), ('mantra', MantraUI), ('tuner', TunerUI),
|
||||||
('control', ControlUI), ('refiner', RefinerUI)]
|
('control', ControlUI), ('refiner', RefinerUI), ('largen', LargenUI)]
|
||||||
|
|
||||||
|
|
||||||
class InferenceUI():
|
class InferenceUI():
|
||||||
@@ -54,6 +55,19 @@ class InferenceUI():
|
|||||||
cfg_general.EXTENSION_PARAS.OFFICIAL_CONTROLLERS))
|
cfg_general.EXTENSION_PARAS.OFFICIAL_CONTROLLERS))
|
||||||
cfg_general.CONTROLLERS = official_controllers.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()
|
pipe_manager = PipelineManager()
|
||||||
config_list = glob(os.path.join(config_dir, '*/*_pro.yaml'),
|
config_list = glob(os.path.join(config_dir, '*/*_pro.yaml'),
|
||||||
recursive=True)
|
recursive=True)
|
||||||
@@ -68,6 +82,12 @@ class InferenceUI():
|
|||||||
for one_controller in cfg_general.CONTROLLERS:
|
for one_controller in cfg_general.CONTROLLERS:
|
||||||
pipe_manager.register_controllers(one_controller)
|
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,
|
self.model_manage_ui = ModelManageUI(cfg_general,
|
||||||
pipe_manager,
|
pipe_manager,
|
||||||
is_debug=is_debug,
|
is_debug=is_debug,
|
||||||
@@ -88,7 +108,8 @@ class InferenceUI():
|
|||||||
self.tab_ui_kwargs[f'{name}_ui'] = ui
|
self.tab_ui_kwargs[f'{name}_ui'] = ui
|
||||||
self.__setattr__(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(
|
assert len(self.component_names.check_box_for_setting) == len(
|
||||||
self.check_box_controlled_tabs)
|
self.check_box_controlled_tabs)
|
||||||
|
|
||||||
@@ -96,6 +117,7 @@ class InferenceUI():
|
|||||||
# create model
|
# create model
|
||||||
self.model_manage_ui.create_ui()
|
self.model_manage_ui.create_ui()
|
||||||
self.gallery_ui.create_ui()
|
self.gallery_ui.create_ui()
|
||||||
|
self.infer_info = gr.State(value=None)
|
||||||
|
|
||||||
# create tabs
|
# create tabs
|
||||||
def create_tab(name, ui):
|
def create_tab(name, ui):
|
||||||
@@ -123,20 +145,65 @@ class InferenceUI():
|
|||||||
for name, ui in self.tab_ui_kwargs.items():
|
for name, ui in self.tab_ui_kwargs.items():
|
||||||
ui.set_callbacks(self.model_manage_ui,
|
ui.set_callbacks(self.model_manage_ui,
|
||||||
**self.tab_ui_kwargs,
|
**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'
|
selected_tab = 'diffusion_ui'
|
||||||
ui_tabs_state = [False] * len(args)
|
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:
|
for key in check_box:
|
||||||
i = self.component_names.check_box_for_setting.index(key)
|
i = self.component_names.check_box_for_setting.index(key)
|
||||||
ui_tabs_state[i] = True
|
ui_tabs_state[i] = True
|
||||||
if ui_tabs_state[i] != args[i]:
|
if ui_tabs_state[i] != args[i]:
|
||||||
selected_tab = self.check_box_controlled_tabs[i] + '_ui'
|
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]
|
ui_tabs_updates = [gr.update(visible=v) for v in ui_tabs_state]
|
||||||
|
|
||||||
return gr.update(
|
if ui_tabs_state[largen_index]:
|
||||||
selected=selected_tab), *ui_tabs_state, *ui_tabs_updates
|
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 = [
|
gr_states = [
|
||||||
self.tab_ui[name].state for name in self.check_box_controlled_tabs
|
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(
|
self.check_box_for_setting.change(
|
||||||
change_setting_tab,
|
change_setting_tab,
|
||||||
inputs=[self.check_box_for_setting, *gr_states],
|
inputs=[self.check_box_for_setting, self.model_manage_ui.diffusion_model, *gr_states],
|
||||||
outputs=[self.setting_tab, *gr_states, *gr_tabs],
|
outputs=[self.check_box_for_setting, self.setting_tab, *gr_states,
|
||||||
|
*gr_tabs, self.model_manage_ui.diffusion_model],
|
||||||
queue=False)
|
queue=False)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
from scepter.modules.inference.diffusion_inference import DiffusionInference
|
from scepter.modules.inference.diffusion_inference import DiffusionInference
|
||||||
|
from scepter.modules.inference.largen_inference import LargenInference
|
||||||
from scepter.modules.utils.logger import get_logger
|
from scepter.modules.utils.logger import get_logger
|
||||||
|
|
||||||
|
|
||||||
@@ -95,7 +96,12 @@ class PipelineManager():
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
def register_pipeline(self, cfg):
|
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)
|
new_inference.init_from_cfg(cfg)
|
||||||
self.contruct_models_index(cfg.NAME, new_inference)
|
self.contruct_models_index(cfg.NAME, new_inference)
|
||||||
|
|
||||||
|
|||||||
@@ -1,16 +1,17 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
# For dataset manager
|
# For dataset manager
|
||||||
from scepter.modules.utils.file_system import FS
|
|
||||||
from scepter.modules.utils.directory import get_md5
|
from scepter.modules.utils.directory import get_md5
|
||||||
|
from scepter.modules.utils.file_system import FS
|
||||||
|
|
||||||
|
|
||||||
def download_image(image):
|
def download_image(image):
|
||||||
if image is not None:
|
if image is not None:
|
||||||
client = FS.get_fs_client(image)
|
client = FS.get_fs_client(image)
|
||||||
if client.tmp_dir.startswith("/home"):
|
if client.tmp_dir.startswith('/home'):
|
||||||
name = get_md5(image)
|
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:
|
else:
|
||||||
local_path = FS.get_from(image)
|
local_path = FS.get_from(image)
|
||||||
return local_path
|
return local_path
|
||||||
@@ -23,21 +24,23 @@ class InferenceUIName():
|
|||||||
if language == 'en':
|
if language == 'en':
|
||||||
self.advance_block_name = 'Advance Setting'
|
self.advance_block_name = 'Advance Setting'
|
||||||
self.check_box_for_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.diffusion_paras = 'Generation Setting'
|
||||||
self.mantra_paras = 'Mantra Book'
|
self.mantra_paras = 'Mantra Book'
|
||||||
self.tuner_paras = 'Tuners'
|
self.tuner_paras = 'Tuners'
|
||||||
self.control_paras = 'Controlable Generation'
|
self.control_paras = 'Controlable Generation'
|
||||||
self.refiner_paras = 'Refiner Setting'
|
self.refiner_paras = 'Refiner Setting'
|
||||||
|
self.largen_paras = 'LAR-Gen'
|
||||||
elif language == 'zh':
|
elif language == 'zh':
|
||||||
self.advance_block_name = '生成选项'
|
self.advance_block_name = '生成选项'
|
||||||
self.check_box_for_setting = ['使用咒语', '使用微调', '使用控制']
|
self.check_box_for_setting = ['使用咒语', '使用微调', '使用控制', 'LAR-Gen']
|
||||||
self.diffusion_paras = '生成参数设置'
|
self.diffusion_paras = '生成参数设置'
|
||||||
self.mantra_paras = '咒语书'
|
self.mantra_paras = '咒语书'
|
||||||
self.tuner_paras = '微调模型'
|
self.tuner_paras = '微调模型'
|
||||||
self.control_paras = '可控生成'
|
self.control_paras = '可控生成'
|
||||||
self.refiner_paras = 'Refine设置'
|
self.refiner_paras = 'Refine设置'
|
||||||
|
self.largen_paras = 'LAR-Gen'
|
||||||
|
|
||||||
|
|
||||||
class ModelManageUIName():
|
class ModelManageUIName():
|
||||||
@@ -216,6 +219,7 @@ class TunerUIName():
|
|||||||
self.example_block_name = 'Examples'
|
self.example_block_name = 'Examples'
|
||||||
self.examples = [[['Pencil Sketch Drawing'], 'a girl in a jacket'],
|
self.examples = [[['Pencil Sketch Drawing'], 'a girl in a jacket'],
|
||||||
[['Flat 2D Art'], 'a cat']]
|
[['Flat 2D Art'], 'a cat']]
|
||||||
|
self.save_button = 'Save'
|
||||||
|
|
||||||
elif language == 'zh':
|
elif language == 'zh':
|
||||||
self.tuner_model = '微调模型'
|
self.tuner_model = '微调模型'
|
||||||
@@ -231,6 +235,7 @@ class TunerUIName():
|
|||||||
self.example_block_name = '样例'
|
self.example_block_name = '样例'
|
||||||
self.examples = [[['铅笔素描'], 'a girl in a jacket'],
|
self.examples = [[['铅笔素描'], 'a girl in a jacket'],
|
||||||
[['扁平2D艺术'], 'a cat']]
|
[['扁平2D艺术'], 'a cat']]
|
||||||
|
self.save_button = '保存'
|
||||||
|
|
||||||
|
|
||||||
class ControlUIName():
|
class ControlUIName():
|
||||||
@@ -342,3 +347,155 @@ class ControlUIName():
|
|||||||
self.advance_block_name = '高级设置'
|
self.advance_block_name = '高级设置'
|
||||||
self.control_scale = '控制强度'
|
self.control_scale = '控制强度'
|
||||||
self.example_block_name = '样例'
|
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
|
import numpy as np
|
||||||
from PIL import Image
|
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.inference.inference_ui.component_names import GalleryUIName
|
||||||
from scepter.studio.utils.uibase import UIBase
|
from scepter.studio.utils.uibase import UIBase
|
||||||
|
|
||||||
@@ -15,6 +16,9 @@ class GalleryUI(UIBase):
|
|||||||
self.pipe_manager = pipe_manager
|
self.pipe_manager = pipe_manager
|
||||||
self.component_names = GalleryUIName(language)
|
self.component_names = GalleryUIName(language)
|
||||||
self.cfg = cfg
|
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):
|
def create_ui(self, *args, **kwargs):
|
||||||
with gr.Group():
|
with gr.Group():
|
||||||
@@ -41,6 +45,7 @@ class GalleryUI(UIBase):
|
|||||||
container=False,
|
container=False,
|
||||||
autofocus=True,
|
autofocus=True,
|
||||||
elem_classes='type_row',
|
elem_classes='type_row',
|
||||||
|
submit_on_enter=True,
|
||||||
lines=1)
|
lines=1)
|
||||||
|
|
||||||
with gr.Column(scale=3, min_width=0):
|
with gr.Column(scale=3, min_width=0):
|
||||||
@@ -87,6 +92,19 @@ class GalleryUI(UIBase):
|
|||||||
style_template,
|
style_template,
|
||||||
style_negative_template,
|
style_negative_template,
|
||||||
image_seed,
|
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):
|
show_jpeg_image=True):
|
||||||
if control_state and control_cond_image is None:
|
if control_state and control_cond_image is None:
|
||||||
raise gr.Error(self.component_names.control_err1)
|
raise gr.Error(self.component_names.control_err1)
|
||||||
@@ -111,9 +129,8 @@ class GalleryUI(UIBase):
|
|||||||
for tuner_m in tuner_model:
|
for tuner_m in tuner_model:
|
||||||
if tuner_m is None or tuner_m == '':
|
if tuner_m is None or tuner_m == '':
|
||||||
continue
|
continue
|
||||||
if (now_pipeline in self.pipe_manager.model_level_info['tuners']
|
if now_pipeline in self.pipe_manager.model_level_info['tuners'] and \
|
||||||
and tuner_m in self.pipe_manager.model_level_info['tuners']
|
tuner_m in self.pipe_manager.model_level_info['tuners'][now_pipeline]:
|
||||||
[now_pipeline]):
|
|
||||||
tuner_m = self.pipe_manager.model_level_info['tuners'][
|
tuner_m = self.pipe_manager.model_level_info['tuners'][
|
||||||
now_pipeline][tuner_m]['model_info']
|
now_pipeline][tuner_m]['model_info']
|
||||||
used_tuner_model.append(tuner_m)
|
used_tuner_model.append(tuner_m)
|
||||||
@@ -163,6 +180,23 @@ class GalleryUI(UIBase):
|
|||||||
pipeline_input['refine_guide_rescale'] = refine_guide_rescale
|
pipeline_input['refine_guide_rescale'] = refine_guide_rescale
|
||||||
else:
|
else:
|
||||||
refine_strength = 0
|
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(
|
results = current_pipeline(
|
||||||
pipeline_input,
|
pipeline_input,
|
||||||
num_samples=image_number,
|
num_samples=image_number,
|
||||||
@@ -177,7 +211,8 @@ class GalleryUI(UIBase):
|
|||||||
if tuner_state or control_state else None,
|
if tuner_state or control_state else None,
|
||||||
control_cond_image=control_cond_image if control_state else None,
|
control_cond_image=control_cond_image if control_state else None,
|
||||||
crop_type=crop_type if control_state else None,
|
crop_type=crop_type if control_state else None,
|
||||||
seed=int(image_seed))
|
seed=int(image_seed),
|
||||||
|
**largen_cfg)
|
||||||
images = []
|
images = []
|
||||||
before_images = []
|
before_images = []
|
||||||
if 'images' in results:
|
if 'images' in results:
|
||||||
@@ -198,27 +233,34 @@ class GalleryUI(UIBase):
|
|||||||
if 'seed' in results:
|
if 'seed' in results:
|
||||||
print(results['seed'])
|
print(results['seed'])
|
||||||
print(images, before_images)
|
print(images, before_images)
|
||||||
|
largen_history.extend(images)
|
||||||
|
if len(largen_history) > 10:
|
||||||
|
largen_history = largen_history[-10:]
|
||||||
if show_jpeg_image:
|
if show_jpeg_image:
|
||||||
save_list = []
|
save_list = []
|
||||||
for i, img in enumerate(images):
|
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')
|
f'cur_gallery_{i}.jpg')
|
||||||
img.save(save_image)
|
img.save(save_image)
|
||||||
save_list.append(save_image)
|
save_list.append(save_image)
|
||||||
images = save_list
|
images = save_list
|
||||||
|
|
||||||
return (
|
return (
|
||||||
gr.Column(visible=len(before_images) > 0),
|
gr.Column(visible=len(before_images) > 0),
|
||||||
before_images,
|
before_images,
|
||||||
images,
|
images,
|
||||||
|
largen_history,
|
||||||
|
gr.update(value=largen_history),
|
||||||
)
|
)
|
||||||
|
|
||||||
def generate_image(self, *args, **kwargs):
|
def generate_image(self, *args, **kwargs):
|
||||||
gallery_result = self.generate_gallery(*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])
|
return (before_refine_panel, before_refine_gallery, output_gallery[0])
|
||||||
|
|
||||||
def set_callbacks(self, inference_ui, model_manage_ui, diffusion_ui,
|
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.gen_inputs = [
|
||||||
self.prompt, mantra_ui.state, tuner_ui.state, control_ui.state,
|
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_strength, refiner_ui.refine_sampler,
|
||||||
refiner_ui.refine_discretization, refiner_ui.refine_guide_scale,
|
refiner_ui.refine_discretization, refiner_ui.refine_guide_scale,
|
||||||
refiner_ui.refine_guide_rescale, mantra_ui.style_template,
|
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.gen_outputs = [
|
||||||
self.before_refine_panel, self.before_refine_gallery,
|
self.before_refine_panel,
|
||||||
self.output_gallery
|
self.before_refine_gallery,
|
||||||
|
self.output_gallery,
|
||||||
|
largen_ui.image_history,
|
||||||
|
largen_ui.gallery,
|
||||||
]
|
]
|
||||||
|
|
||||||
self.generate_button.click(self.generate_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):
|
def create_ui(self, *args, **kwargs):
|
||||||
self.state = gr.State(value=False)
|
self.state = gr.State(value=False)
|
||||||
with gr.Column(visible=False) as self.tab:
|
with gr.Column(equal_height=True, visible=False) as self.tab:
|
||||||
with gr.Row():
|
with gr.Row(scale=1):
|
||||||
with gr.Column(scale=1):
|
with gr.Column(scale=1):
|
||||||
with gr.Group(visible=True):
|
with gr.Group(visible=True):
|
||||||
with gr.Row(equal_height=True):
|
with gr.Row(equal_height=True):
|
||||||
self.style = gr.Dropdown(
|
self.style = gr.Dropdown(
|
||||||
label=self.component_names.mantra_styles,
|
label=self.component_names.mantra_styles,
|
||||||
choices=self.all_styles[self.default_pipeline],
|
choices=self.all_styles.get(
|
||||||
|
self.default_pipeline, []),
|
||||||
value=None,
|
value=None,
|
||||||
multiselect=True,
|
multiselect=True,
|
||||||
interactive=True)
|
interactive=True)
|
||||||
|
|||||||
@@ -146,11 +146,20 @@ class ModelManageUI(UIBase):
|
|||||||
continue
|
continue
|
||||||
model_name = f"{now_pipeline}_{module['name']}"
|
model_name = f"{now_pipeline}_{module['name']}"
|
||||||
all_module_name[module_name] = model_name
|
all_module_name[module_name] = model_name
|
||||||
|
tunner_choices = []
|
||||||
if now_pipeline in self.default_choices['tuners']:
|
if now_pipeline in self.default_choices['tuners']:
|
||||||
tunner_choices = self.default_choices['tuners'][now_pipeline][
|
tunner_choices = self.default_choices['tuners'][now_pipeline][
|
||||||
'choices']
|
'choices']
|
||||||
else:
|
custom_tunner_choices = []
|
||||||
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[
|
if now_pipeline in self.default_choices[
|
||||||
'controllers'] and control_mode 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['first_stage_model']),
|
||||||
gr.Dropdown(value=all_module_name['cond_stage_model']),
|
gr.Dropdown(value=all_module_name['cond_stage_model']),
|
||||||
gr.Dropdown(choices=tunner_choices, value=[]),
|
gr.Dropdown(choices=tunner_choices, value=[]),
|
||||||
|
gr.Dropdown(choices=custom_tunner_choices,
|
||||||
|
value=custom_tunner_default),
|
||||||
gr.Dropdown(choices=controller_choices,
|
gr.Dropdown(choices=controller_choices,
|
||||||
value=controller_default),
|
value=controller_default),
|
||||||
gr.Dropdown(choices=mantra_ui.all_styles[now_pipeline],
|
gr.Dropdown(choices=mantra_ui.all_styles.get(now_pipeline, []),
|
||||||
value=[]),
|
value=[]),
|
||||||
gr.Textbox(choices=cur_paras.NEGATIVE_PROMPT.get('VALUES', []),
|
gr.Textbox(choices=cur_paras.NEGATIVE_PROMPT.get('VALUES', []),
|
||||||
value=cur_paras.NEGATIVE_PROMPT.get('DEFAULT', '')),
|
value=cur_paras.NEGATIVE_PROMPT.get('DEFAULT', '')),
|
||||||
@@ -206,10 +217,11 @@ class ModelManageUI(UIBase):
|
|||||||
outputs=[
|
outputs=[
|
||||||
self.diffusion_state, self.first_stage_model,
|
self.diffusion_state, self.first_stage_model,
|
||||||
self.cond_stage_model, tuner_ui.tuner_model,
|
self.cond_stage_model, tuner_ui.tuner_model,
|
||||||
control_ui.control_model, mantra_ui.style,
|
tuner_ui.custom_tuner_model, control_ui.control_model,
|
||||||
diffusion_ui.negative_prompt, diffusion_ui.prompt_prefix,
|
mantra_ui.style, diffusion_ui.negative_prompt,
|
||||||
diffusion_ui.output_height, diffusion_ui.sampler,
|
diffusion_ui.prompt_prefix, diffusion_ui.output_height,
|
||||||
diffusion_ui.discretization, diffusion_ui.sample_steps,
|
diffusion_ui.sampler, diffusion_ui.discretization,
|
||||||
diffusion_ui.guide_scale, diffusion_ui.guide_rescale
|
diffusion_ui.sample_steps, diffusion_ui.guide_scale,
|
||||||
|
diffusion_ui.guide_rescale
|
||||||
],
|
],
|
||||||
queue=True)
|
queue=True)
|
||||||
|
|||||||
@@ -21,17 +21,21 @@ class TunerUI(UIBase):
|
|||||||
'default']
|
'default']
|
||||||
self.default_pipeline = pipe_manager.model_level_info[
|
self.default_pipeline = pipe_manager.model_level_info[
|
||||||
default_diffusion_model]['pipeline'][0]
|
default_diffusion_model]['pipeline'][0]
|
||||||
|
self.tunner_choices = []
|
||||||
if self.default_pipeline in self.default_choices['tuners']:
|
if self.default_pipeline in self.default_choices['tuners']:
|
||||||
self.tunner_choices = self.default_choices['tuners'][
|
self.tunner_choices = self.default_choices['tuners'][
|
||||||
self.default_pipeline]['choices']
|
self.default_pipeline]['choices']
|
||||||
self.tunner_default = self.default_choices['tuners'][
|
self.tunner_default = self.default_choices['tuners'][
|
||||||
self.default_pipeline]['default']
|
self.default_pipeline]['default']
|
||||||
else:
|
self.custom_tuner_choices = []
|
||||||
self.tunner_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.tunner_default = None
|
||||||
self.component_names = TunerUIName(language)
|
self.component_names = TunerUIName(language)
|
||||||
self.cfg_tuners = cfg.TUNERS
|
self.cfg_tuners = cfg.TUNERS + cfg.CUSTOM_TUNERS
|
||||||
self.name_level_tuners = {}
|
self.name_level_tuners = {}
|
||||||
for one_tuner in tqdm(self.cfg_tuners):
|
for one_tuner in tqdm(self.cfg_tuners):
|
||||||
if one_tuner.BASE_MODEL not in self.name_level_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):
|
def create_ui(self, *args, **kwargs):
|
||||||
self.state = gr.State(value=False)
|
self.state = gr.State(value=False)
|
||||||
with gr.Column(visible=False) as self.tab:
|
with gr.Column(equal_height=True, visible=False) as self.tab:
|
||||||
with gr.Row():
|
with gr.Row(scale=1):
|
||||||
with gr.Column(variant='panel', scale=1, min_width=0):
|
with gr.Column(variant='panel', scale=1, min_width=0):
|
||||||
with gr.Group(visible=True):
|
with gr.Group(visible=True):
|
||||||
with gr.Row(equal_height=True):
|
with gr.Row(equal_height=True):
|
||||||
@@ -63,10 +67,16 @@ class TunerUI(UIBase):
|
|||||||
self.custom_tuner_model = gr.Dropdown(
|
self.custom_tuner_model = gr.Dropdown(
|
||||||
label=self.component_names.
|
label=self.component_names.
|
||||||
custom_tuner_model,
|
custom_tuner_model,
|
||||||
choices=[],
|
choices=self.custom_tuner_choices,
|
||||||
value=None,
|
value=None,
|
||||||
multiselect=True,
|
multiselect=True,
|
||||||
interactive=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.Row(equal_height=True):
|
||||||
with gr.Column(scale=1):
|
with gr.Column(scale=1):
|
||||||
self.tuner_type = gr.Text(
|
self.tuner_type = gr.Text(
|
||||||
@@ -110,6 +120,7 @@ class TunerUI(UIBase):
|
|||||||
label=self.component_names.example_block_name, open=True)
|
label=self.component_names.example_block_name, open=True)
|
||||||
|
|
||||||
def set_callbacks(self, model_manage_ui, **kwargs):
|
def set_callbacks(self, model_manage_ui, **kwargs):
|
||||||
|
manager = kwargs.pop('manager')
|
||||||
gallery_ui = kwargs.pop('gallery_ui')
|
gallery_ui = kwargs.pop('gallery_ui')
|
||||||
with self.example_block:
|
with self.example_block:
|
||||||
gr.Examples(examples=self.component_names.examples,
|
gr.Examples(examples=self.component_names.examples,
|
||||||
@@ -121,7 +132,7 @@ class TunerUI(UIBase):
|
|||||||
now_pipeline = diffusion_model_info['pipeline'][0]
|
now_pipeline = diffusion_model_info['pipeline'][0]
|
||||||
tuner_info = {}
|
tuner_info = {}
|
||||||
if tuner_model is not None and len(tuner_model) > 0:
|
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], {})
|
tuner_model[-1], {})
|
||||||
if tuner_info.get(
|
if tuner_info.get(
|
||||||
'IMAGE_PATH',
|
'IMAGE_PATH',
|
||||||
@@ -141,3 +152,45 @@ class TunerUI(UIBase):
|
|||||||
self.tuner_example, self.tuner_prompt_example
|
self.tuner_example, self.tuner_prompt_example
|
||||||
],
|
],
|
||||||
queue=False)
|
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):
|
class CreateDatasetUI(UIBase):
|
||||||
def __init__(self, cfg, is_debug=False, language='en'):
|
def __init__(self, cfg, is_debug=False, language='en'):
|
||||||
self.work_dir = cfg.WORK_DIR
|
self.work_dir = cfg.WORK_DIR
|
||||||
self.dir_list = FS.walk_dir(self.work_dir, recurse=False)
|
|
||||||
self.cache_file = {}
|
self.cache_file = {}
|
||||||
self.meta_dict = {}
|
self.meta_dict = {}
|
||||||
self.dataset_list = self.load_history()
|
self.dataset_list = self.load_history()
|
||||||
@@ -71,7 +70,8 @@ class CreateDatasetUI(UIBase):
|
|||||||
json.dump(save_meta, open(local_path, 'w'))
|
json.dump(save_meta, open(local_path, 'w'))
|
||||||
return meta_file
|
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",
|
"dataset_name": "xxxx",
|
||||||
@@ -86,11 +86,17 @@ class CreateDatasetUI(UIBase):
|
|||||||
save_file_list = os.path.join(dataset_folder, 'file.csv')
|
save_file_list = os.path.join(dataset_folder, 'file.csv')
|
||||||
save_file_list = self.write_file_list(file_list, save_file_list)
|
save_file_list = self.write_file_list(file_list, save_file_list)
|
||||||
meta = {
|
meta = {
|
||||||
'dataset_name': user_name,
|
'dataset_name':
|
||||||
'cursor': cursor,
|
user_name if login_user_name == '' or login_user_name is None else
|
||||||
'file_list': file_list,
|
'_'.join([login_user_name, user_name]),
|
||||||
'train_csv': train_csv,
|
'cursor':
|
||||||
'save_file_list': save_file_list
|
cursor,
|
||||||
|
'file_list':
|
||||||
|
file_list,
|
||||||
|
'train_csv':
|
||||||
|
train_csv,
|
||||||
|
'save_file_list':
|
||||||
|
save_file_list
|
||||||
}
|
}
|
||||||
self.save_meta(meta, dataset_folder)
|
self.save_meta(meta, dataset_folder)
|
||||||
return meta
|
return meta
|
||||||
@@ -180,7 +186,7 @@ class CreateDatasetUI(UIBase):
|
|||||||
file_folder = None
|
file_folder = None
|
||||||
train_list = None
|
train_list = None
|
||||||
hit_dir = None
|
hit_dir = None
|
||||||
raw_list = []
|
raw_list = {}
|
||||||
mac_osx = os.path.join(local_dataset_folder, '__MACOSX')
|
mac_osx = os.path.join(local_dataset_folder, '__MACOSX')
|
||||||
if os.path.exists(mac_osx):
|
if os.path.exists(mac_osx):
|
||||||
res = os.popen(f"rm -rf '{mac_osx}'")
|
res = os.popen(f"rm -rf '{mac_osx}'")
|
||||||
@@ -191,28 +197,45 @@ class CreateDatasetUI(UIBase):
|
|||||||
res = res.readlines()
|
res = res.readlines()
|
||||||
continue
|
continue
|
||||||
if FS.isdir(one_dir):
|
if FS.isdir(one_dir):
|
||||||
sub_dir = FS.walk_dir(one_dir)
|
if one_dir.endswith('images') or one_dir.endswith('images/'):
|
||||||
for one_s_dir in sub_dir:
|
file_folder = one_dir
|
||||||
if FS.isdir(one_s_dir) and one_s_dir.split(
|
hit_dir = one_dir
|
||||||
one_dir)[1].replace('/', '') == 'images':
|
else:
|
||||||
file_folder = one_s_dir
|
sub_dir = FS.walk_dir(one_dir)
|
||||||
hit_dir = one_dir
|
for one_s_dir in sub_dir:
|
||||||
if FS.isfile(one_s_dir) and one_s_dir.split(
|
if FS.isdir(one_s_dir) and one_s_dir.split(
|
||||||
one_dir)[1].replace('/', '') == 'train.csv':
|
one_dir)[1].replace('/', '') == 'images':
|
||||||
train_list = one_s_dir
|
file_folder = one_s_dir
|
||||||
if file_folder is not None and train_list is not None:
|
hit_dir = one_dir
|
||||||
break
|
if FS.isfile(one_s_dir) and one_s_dir.split(
|
||||||
if (one_s_dir.endswith('.jpg')
|
one_dir)[1].replace('/', '') == 'train.csv':
|
||||||
or one_s_dir.endswith('.jpeg')
|
train_list = one_s_dir
|
||||||
or one_s_dir.endswith('.png')
|
if file_folder is not None and train_list is not None:
|
||||||
or one_s_dir.endswith('.webp')):
|
break
|
||||||
raw_list.append(one_s_dir)
|
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:
|
else:
|
||||||
if (one_dir.endswith('.jpg') or one_dir.endswith('.jpeg')
|
if (one_dir.endswith('.jpg') or one_dir.endswith('.jpeg')
|
||||||
or one_dir.endswith('.png')
|
or one_dir.endswith('.png')
|
||||||
or one_dir.endswith('.webp')):
|
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:
|
if file_folder is None and len(raw_list) < 1:
|
||||||
raise gr.Error(
|
raise gr.Error(
|
||||||
"images folder or train.csv doesn't exists, or nothing exists in your zip"
|
"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:
|
if file_folder is not None:
|
||||||
_ = FS.get_dir_to_local_dir(file_folder, new_file_folder)
|
_ = FS.get_dir_to_local_dir(file_folder, new_file_folder)
|
||||||
elif len(raw_list) > 0:
|
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):
|
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:
|
try:
|
||||||
os.rename(
|
os.rename(
|
||||||
os.path.abspath(cur_image),
|
os.path.abspath(cur_image[0]),
|
||||||
f'{new_file_folder}/{get_md5(cur_image)}{surfix}')
|
f'{new_file_folder}/{get_md5(cur_image[0])}{surfix}')
|
||||||
raw_list[img_id] = [
|
raw_list[img_id] = [
|
||||||
os.path.join('images',
|
os.path.join('images',
|
||||||
f'{get_md5(cur_image)}{surfix}'),
|
f'{get_md5(cur_image[0])}{surfix}'),
|
||||||
cur_image.split('/')[-1]
|
prompt
|
||||||
]
|
]
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(e)
|
print(e)
|
||||||
@@ -251,13 +278,14 @@ class CreateDatasetUI(UIBase):
|
|||||||
res = res.readlines()
|
res = res.readlines()
|
||||||
if not os.path.exists(new_train_list):
|
if not os.path.exists(new_train_list):
|
||||||
raise gr.Error(f'{str(res)}')
|
raise gr.Error(f'{str(res)}')
|
||||||
try:
|
if not file_folder == hit_dir:
|
||||||
res = os.popen(f"rm -rf '{hit_dir}/images/*'")
|
try:
|
||||||
_ = res.readlines()
|
res = os.popen(f"rm -rf '{hit_dir}/images/*'")
|
||||||
res = os.popen(f"rm -rf '{hit_dir}'")
|
_ = res.readlines()
|
||||||
_ = res.readlines()
|
res = os.popen(f"rm -rf '{hit_dir}'")
|
||||||
except Exception:
|
_ = res.readlines()
|
||||||
pass
|
except Exception:
|
||||||
|
pass
|
||||||
file_list = self.load_train_csv(new_train_list, data_folder)
|
file_list = self.load_train_csv(new_train_list, data_folder)
|
||||||
return file_list
|
return file_list
|
||||||
|
|
||||||
@@ -317,21 +345,25 @@ class CreateDatasetUI(UIBase):
|
|||||||
return True, file_path
|
return True, file_path
|
||||||
return False, file_path
|
return False, file_path
|
||||||
|
|
||||||
def load_history(self):
|
def load_history(self, login_user_name=''):
|
||||||
dataset_list = []
|
dataset_list = []
|
||||||
|
self.dir_list = FS.walk_dir(self.work_dir, recurse=False)
|
||||||
for one_dir in self.dir_list:
|
for one_dir in self.dir_list:
|
||||||
if FS.isdir(one_dir):
|
if FS.isdir(one_dir):
|
||||||
meta_file = os.path.join(one_dir, 'meta.json')
|
meta_file = os.path.join(one_dir, 'meta.json')
|
||||||
if FS.exists(meta_file):
|
if FS.exists(meta_file):
|
||||||
local_dataset_folder, _ = FS.map_to_local(one_dir)
|
local_dataset_folder, _ = FS.map_to_local(one_dir)
|
||||||
local_dataset_folder = FS.get_dir_to_local_dir(
|
if not FS.exists(
|
||||||
one_dir, local_dataset_folder, multi_thread=True)
|
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(
|
meta_data = self.load_meta(
|
||||||
os.path.join(local_dataset_folder, 'meta.json'))
|
os.path.join(local_dataset_folder, 'meta.json'))
|
||||||
meta_data['local_work_dir'] = local_dataset_folder
|
meta_data['local_work_dir'] = local_dataset_folder
|
||||||
meta_data['work_dir'] = one_dir
|
meta_data['work_dir'] = one_dir
|
||||||
dataset_list.append(meta_data['dataset_name'])
|
if meta_data['dataset_name'].startswith(login_user_name):
|
||||||
self.meta_dict[meta_data['dataset_name']] = meta_data
|
dataset_list.append(meta_data['dataset_name'])
|
||||||
|
self.meta_dict[meta_data['dataset_name']] = meta_data
|
||||||
return dataset_list
|
return dataset_list
|
||||||
|
|
||||||
def create_ui(self):
|
def create_ui(self):
|
||||||
@@ -396,7 +428,7 @@ class CreateDatasetUI(UIBase):
|
|||||||
self.file_panel = file_panel
|
self.file_panel = file_panel
|
||||||
self.modify_panel = modify_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():
|
def show_dataset_panel():
|
||||||
return (gr.Column(visible=False), gr.Column(visible=True),
|
return (gr.Column(visible=False), gr.Column(visible=True),
|
||||||
gr.Column(visible=True),
|
gr.Column(visible=True),
|
||||||
@@ -419,17 +451,19 @@ class CreateDatasetUI(UIBase):
|
|||||||
datetime.datetime.now())
|
datetime.datetime.now())
|
||||||
return data_name
|
return data_name
|
||||||
|
|
||||||
def refresh():
|
def refresh(login_user_name):
|
||||||
return gr.Dropdown(value=self.dataset_list[-1]
|
dataset_list = self.load_history(login_user_name=login_user_name)
|
||||||
if len(self.dataset_list) > 0 else '',
|
return gr.Dropdown(
|
||||||
choices=self.dataset_list)
|
value=dataset_list[-1] if len(dataset_list) > 0 else '',
|
||||||
|
choices=dataset_list)
|
||||||
|
|
||||||
self.refresh_dataset_name.click(refresh,
|
self.refresh_dataset_name.click(refresh,
|
||||||
|
inputs=[manager.user_name],
|
||||||
outputs=[self.dataset_name],
|
outputs=[self.dataset_name],
|
||||||
queue=False)
|
queue=False)
|
||||||
|
|
||||||
def confirm_create_dataset(user_name, create_mode, file_url, file_path,
|
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:
|
if user_name.strip() == '' or ' ' in user_name or '/' in user_name:
|
||||||
raise gr.Error(self.components_name.illegal_data_name_err1)
|
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:
|
if not file_url.strip() == '' and file_path is not None:
|
||||||
raise gr.Error(self.components_name.illegal_data_name_err4)
|
raise gr.Error(self.components_name.illegal_data_name_err4)
|
||||||
if create_mode == 3 and not file_url.strip() == '':
|
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}')
|
save_file = os.path.join(self.work_dir, f'{user_name}{surfix}')
|
||||||
local_path, _ = FS.map_to_local(save_file)
|
local_path, _ = FS.map_to_local(save_file)
|
||||||
res = os.popen(f"wget -c '{file_url}' -O '{local_path}'")
|
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
|
cursor = 0 if len(file_list) > 0 else -1
|
||||||
meta = self.construct_meta(cursor, file_list, dataset_folder,
|
meta = self.construct_meta(cursor, file_list, dataset_folder,
|
||||||
user_name)
|
user_name, login_user_name)
|
||||||
|
|
||||||
meta['local_work_dir'] = local_dataset_folder
|
meta['local_work_dir'] = local_dataset_folder
|
||||||
meta['work_dir'] = dataset_folder
|
meta['work_dir'] = dataset_folder
|
||||||
@@ -498,13 +534,13 @@ class CreateDatasetUI(UIBase):
|
|||||||
self.meta_dict[meta['dataset_name']] = meta
|
self.meta_dict[meta['dataset_name']] = meta
|
||||||
if meta['dataset_name'] not in self.dataset_list:
|
if meta['dataset_name'] not in self.dataset_list:
|
||||||
self.dataset_list.append(meta['dataset_name'])
|
self.dataset_list.append(meta['dataset_name'])
|
||||||
return (
|
return (gr.Checkbox(value=True, visible=False),
|
||||||
gr.Checkbox(value=True, visible=False),
|
gr.Dropdown(value=meta['dataset_name'],
|
||||||
gr.Dropdown(value=user_name, choices=self.dataset_list),
|
choices=self.dataset_list),
|
||||||
)
|
gr.Text(value=meta['dataset_name']))
|
||||||
|
|
||||||
def clear_file():
|
def clear_file():
|
||||||
return gr.Text(visible=True)
|
return gr.Text(visible=False)
|
||||||
|
|
||||||
# Click Create
|
# Click Create
|
||||||
self.btn_create_datasets.click(show_dataset_panel, [], [
|
self.btn_create_datasets.click(show_dataset_panel, [], [
|
||||||
@@ -545,9 +581,9 @@ class CreateDatasetUI(UIBase):
|
|||||||
# Click Confirm
|
# Click Confirm
|
||||||
self.confirm_data_button.click(confirm_create_dataset, [
|
self.confirm_data_button.click(confirm_create_dataset, [
|
||||||
self.user_data_name, self.create_mode, self.file_path_url,
|
self.user_data_name, self.create_mode, self.file_path_url,
|
||||||
self.file_path, self.panel_state
|
self.file_path, self.panel_state, manager.user_name
|
||||||
], [self.panel_state, self.dataset_name],
|
], [self.panel_state, self.dataset_name, self.user_data_name],
|
||||||
queue=True)
|
queue=False)
|
||||||
|
|
||||||
def show_edit_panel(panel_state, data_name):
|
def show_edit_panel(panel_state, data_name):
|
||||||
if panel_state:
|
if panel_state:
|
||||||
@@ -568,7 +604,7 @@ class CreateDatasetUI(UIBase):
|
|||||||
],
|
],
|
||||||
queue=False)
|
queue=False)
|
||||||
|
|
||||||
def modify_data_name(user_name, prev_data_name):
|
def modify_data_name(user_name, prev_data_name, login_user_name):
|
||||||
print(
|
print(
|
||||||
f'Current file name {prev_data_name}, new file name {user_name}.'
|
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)
|
raise gr.Error(self.components_name.illegal_data_err3)
|
||||||
cursor = ori_meta['cursor']
|
cursor = ori_meta['cursor']
|
||||||
meta = self.construct_meta(cursor, file_list,
|
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['local_work_dir'] = local_dataset_folder
|
||||||
meta['work_dir'] = dataset_folder
|
meta['work_dir'] = dataset_folder
|
||||||
|
|
||||||
@@ -624,7 +661,10 @@ class CreateDatasetUI(UIBase):
|
|||||||
|
|
||||||
self.modify_data_button.click(
|
self.modify_data_button.click(
|
||||||
modify_data_name,
|
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],
|
outputs=[self.user_data_name_state, self.dataset_name],
|
||||||
queue=False)
|
queue=False)
|
||||||
|
|
||||||
@@ -647,3 +687,14 @@ class CreateDatasetUI(UIBase):
|
|||||||
gallery_dataset.gallery_state
|
gallery_dataset.gallery_state
|
||||||
],
|
],
|
||||||
queue=False)
|
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):
|
def set_callbacks(self, manager):
|
||||||
self.create_dataset.set_callbacks(self.dataset_gallery,
|
self.create_dataset.set_callbacks(self.dataset_gallery,
|
||||||
self.export_dataset)
|
self.export_dataset, manager)
|
||||||
self.dataset_gallery.set_callbacks(self.create_dataset)
|
self.dataset_gallery.set_callbacks(self.create_dataset)
|
||||||
self.export_dataset.set_callbacks(self.create_dataset, manager)
|
self.export_dataset.set_callbacks(self.create_dataset, manager)
|
||||||
|
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import numpy as np
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from scepter.modules.solver.hooks.checkpoint import CheckpointHook
|
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.solver.registry import SOLVERS
|
||||||
from scepter.modules.utils.config import Config
|
from scepter.modules.utils.config import Config
|
||||||
from scepter.modules.utils.distribute import we
|
from scepter.modules.utils.distribute import we
|
||||||
@@ -27,6 +28,7 @@ def run_task(cfg):
|
|||||||
solver.set_up_pre()
|
solver.set_up_pre()
|
||||||
solver.set_up()
|
solver.set_up()
|
||||||
ori_steps = solver.max_steps
|
ori_steps = solver.max_steps
|
||||||
|
|
||||||
if 'train' in solver.datas:
|
if 'train' in solver.datas:
|
||||||
dataset = solver.datas['train'].dataset
|
dataset = solver.datas['train'].dataset
|
||||||
if hasattr(dataset, 'real_number'):
|
if hasattr(dataset, 'real_number'):
|
||||||
@@ -48,7 +50,20 @@ def run_task(cfg):
|
|||||||
f'checkpoint save interval is changed from {ori_interval} '
|
f'checkpoint save interval is changed from {ori_interval} '
|
||||||
f'to {hook.interval} according to the setting epoches '
|
f'to {hook.interval} according to the setting epoches '
|
||||||
f'interval {ori_interval}')
|
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()
|
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
|
import scepter
|
||||||
from scepter.modules.utils.config import Config
|
from scepter.modules.utils.config import Config
|
||||||
from scepter.modules.utils.file_system import FS
|
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.self_train_ui.trainer_ui import TrainerUI
|
||||||
from scepter.studio.self_train.utils.config_parser import get_all_config
|
from scepter.studio.self_train.utils.config_parser import get_all_config
|
||||||
from scepter.studio.utils.env import init_env
|
from scepter.studio.utils.env import init_env
|
||||||
@@ -34,20 +34,20 @@ class SelfTrainUI():
|
|||||||
BASE_CFG_VALUE,
|
BASE_CFG_VALUE,
|
||||||
is_debug=is_debug,
|
is_debug=is_debug,
|
||||||
language=language)
|
language=language)
|
||||||
self.inference_ui = InferenceUI(cfg_general,
|
self.model_ui = ModelUI(cfg_general,
|
||||||
BASE_CFG_VALUE,
|
BASE_CFG_VALUE,
|
||||||
is_debug=is_debug,
|
is_debug=is_debug,
|
||||||
language=language)
|
language=language)
|
||||||
|
|
||||||
def create_ui(self):
|
def create_ui(self):
|
||||||
with gr.Row():
|
with gr.Row():
|
||||||
self.trainer_ui.create_ui()
|
self.trainer_ui.create_ui()
|
||||||
with gr.Row():
|
with gr.Row():
|
||||||
self.inference_ui.create_ui()
|
self.model_ui.create_ui()
|
||||||
|
|
||||||
def set_callbacks(self, manager):
|
def set_callbacks(self, manager):
|
||||||
self.trainer_ui.set_callbacks(self.inference_ui)
|
self.trainer_ui.set_callbacks(self.model_ui, manager)
|
||||||
self.inference_ui.set_callbacks(self.trainer_ui, manager)
|
self.model_ui.set_callbacks(self.trainer_ui, manager)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
|
|||||||
@@ -1,11 +1,12 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
# For dataset manager
|
# For dataset manager
|
||||||
class InferenceUIName():
|
class ModelUIName():
|
||||||
def __init__(self, language='en'):
|
def __init__(self, language='en'):
|
||||||
if language == 'en':
|
if language == 'en':
|
||||||
self.output_model_block = 'Model Output'
|
self.output_model_block = 'Model Output'
|
||||||
self.output_model_name = 'Output Model Name'
|
self.output_model_name = 'Output Model Name'
|
||||||
|
self.output_ckpt_name = 'Output Ckpt Name'
|
||||||
self.test_prompt = 'Test Prompt'
|
self.test_prompt = 'Test Prompt'
|
||||||
self.test_prefix = 'Test Prefix'
|
self.test_prefix = 'Test Prefix'
|
||||||
self.test_n_prompt = 'Negative Prompt'
|
self.test_n_prompt = 'Negative Prompt'
|
||||||
@@ -21,15 +22,23 @@ class InferenceUIName():
|
|||||||
self.extra_model_gbtn = 'Add Model'
|
self.extra_model_gbtn = 'Add Model'
|
||||||
self.refresh_model_gbtn = 'Refresh Model'
|
self.refresh_model_gbtn = 'Refresh Model'
|
||||||
self.go_to_inference = 'Go to inference'
|
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
|
# Error or Warning
|
||||||
self.inference_err1 = 'Inference failed, please try again.'
|
self.inference_err1 = 'Inference failed, please try again.'
|
||||||
self.inference_err2 = 'Test prompt is empty.'
|
self.inference_err2 = 'Test prompt is empty.'
|
||||||
self.inference_err3 = "Doesn't surpport this base model"
|
self.model_err3 = "Doesn't surpport this base model"
|
||||||
self.inference_err4 = "This model maybe not finish training, because model doesn't exist."
|
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':
|
elif language == 'zh':
|
||||||
self.output_model_block = '模型产出'
|
self.output_model_block = '模型产出'
|
||||||
self.output_model_name = '产出模型名称'
|
self.output_model_name = '产出名称'
|
||||||
|
self.output_ckpt_name = '产出检查点名称'
|
||||||
self.test_prompt = '测试提示词'
|
self.test_prompt = '测试提示词'
|
||||||
self.test_prefix = '测试前缀'
|
self.test_prefix = '测试前缀'
|
||||||
self.test_n_prompt = '负向提示词'
|
self.test_n_prompt = '负向提示词'
|
||||||
@@ -44,12 +53,20 @@ class InferenceUIName():
|
|||||||
self.extra_model_gtxt = '额外模型'
|
self.extra_model_gtxt = '额外模型'
|
||||||
self.extra_model_gbtn = '添加模型'
|
self.extra_model_gbtn = '添加模型'
|
||||||
self.refresh_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
|
# Error or Warning
|
||||||
self.inference_err1 = '推理失败,请重试。'
|
self.inference_err1 = '推理失败,请重试。'
|
||||||
self.inference_err2 = '测试提示词为空。'
|
self.inference_err2 = '测试提示词为空。'
|
||||||
self.inference_err3 = '不支持的基础模型'
|
self.model_err3 = '不支持的基础模型'
|
||||||
self.go_to_inference = '使用模型'
|
self.go_to_inference = '使用模型'
|
||||||
self.inference_err4 = '模型可能没有训练完成或者模型不存在'
|
self.model_err4 = '模型可能没有训练完成或者模型不存在'
|
||||||
|
self.model_err5 = '模型{}不存在'
|
||||||
|
self.training_warn1 = '暂时没有日志文件;任务启动中或失败!'
|
||||||
|
|
||||||
|
|
||||||
class TrainerUIName():
|
class TrainerUIName():
|
||||||
@@ -80,7 +97,8 @@ class TrainerUIName():
|
|||||||
self.base_model = 'Base Model'
|
self.base_model = 'Base Model'
|
||||||
self.tuner_name = 'Fine-tuning Method'
|
self.tuner_name = 'Fine-tuning Method'
|
||||||
self.base_model_revision = 'Model Version Number'
|
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.train_epoch = 'Number of Training Epochs'
|
||||||
self.learning_rate = 'Learning Rate'
|
self.learning_rate = 'Learning Rate'
|
||||||
self.save_interval = 'Save Interval'
|
self.save_interval = 'Save Interval'
|
||||||
@@ -89,9 +107,8 @@ class TrainerUIName():
|
|||||||
self.replace_keywords = 'Trigger Keywords'
|
self.replace_keywords = 'Trigger Keywords'
|
||||||
self.work_name = 'Save Model Name (refresh to get a random value)'
|
self.work_name = 'Save Model Name (refresh to get a random value)'
|
||||||
self.push_to_hub = 'Push to hub'
|
self.push_to_hub = 'Push to hub'
|
||||||
self.log_block = 'Training Log...'
|
|
||||||
self.training_button = 'Start Training'
|
self.training_button = 'Start Training'
|
||||||
|
self.eval_prompts = 'Eval Prompts'
|
||||||
# Error or Warning
|
# Error or Warning
|
||||||
self.training_err1 = 'CUDA is unavailable.'
|
self.training_err1 = 'CUDA is unavailable.'
|
||||||
self.training_err2 = 'Currently insufficient VRAM, training failed!'
|
self.training_err2 = 'Currently insufficient VRAM, training failed!'
|
||||||
@@ -121,7 +138,8 @@ class TrainerUIName():
|
|||||||
self.base_model = '基础模型'
|
self.base_model = '基础模型'
|
||||||
self.tuner_name = '微调方法'
|
self.tuner_name = '微调方法'
|
||||||
self.base_model_revision = '模型版本号'
|
self.base_model_revision = '模型版本号'
|
||||||
self.resolution = '分辨率'
|
self.resolution_height = '训练高度'
|
||||||
|
self.resolution_width = '训练宽度'
|
||||||
self.train_epoch = '训练轮数'
|
self.train_epoch = '训练轮数'
|
||||||
self.learning_rate = '学习率'
|
self.learning_rate = '学习率'
|
||||||
self.save_interval = '存储间隔'
|
self.save_interval = '存储间隔'
|
||||||
@@ -130,8 +148,7 @@ class TrainerUIName():
|
|||||||
self.replace_keywords = '触发关键词'
|
self.replace_keywords = '触发关键词'
|
||||||
self.work_name = '保存模型名称(刷新获得随机值)'
|
self.work_name = '保存模型名称(刷新获得随机值)'
|
||||||
self.push_to_hub = '推送魔搭社区'
|
self.push_to_hub = '推送魔搭社区'
|
||||||
self.log_block = '训练日志...'
|
self.eval_prompts = '评测文本'
|
||||||
self.training_button = '开始训练'
|
|
||||||
# Error or Warning
|
# Error or Warning
|
||||||
self.training_err1 = 'CUDA不可用.'
|
self.training_err1 = 'CUDA不可用.'
|
||||||
self.training_err2 = '目前显存不足,训练失败!'
|
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,473 @@
|
|||||||
|
# -*- 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 ''
|
||||||
|
if ckpt_value is not None and len(ckpt_value) > 0:
|
||||||
|
gallery_value = self.get_gallery_list(model_name, ckpt_value)
|
||||||
|
else:
|
||||||
|
gallery_value = []
|
||||||
|
select_index = 0 if len(gallery_value) > 0 else None
|
||||||
|
return (message, gr.Column(visible=status in ('running',
|
||||||
|
'success')),
|
||||||
|
gr.Dropdown(choices=ckpt_list, value=ckpt_value),
|
||||||
|
gr.Gallery(value=gallery_value,
|
||||||
|
preview=True,
|
||||||
|
selected_index=select_index)
|
||||||
|
)
|
||||||
|
|
||||||
|
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,
|
||||||
|
self.eval_gallery
|
||||||
|
],
|
||||||
|
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=self.model_list, value=model_name),
|
||||||
|
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')
|
||||||
|
if os.path.exists(image_dir):
|
||||||
|
image_path = [
|
||||||
|
os.path.join(image_dir, name)
|
||||||
|
for name in os.listdir(image_dir)
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
image_path = []
|
||||||
|
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.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
import copy
|
import copy
|
||||||
import datetime
|
import datetime
|
||||||
|
import json
|
||||||
import os
|
import os
|
||||||
import random
|
import random
|
||||||
import time
|
from collections import OrderedDict
|
||||||
|
|
||||||
import gradio as gr
|
import gradio as gr
|
||||||
import torch
|
import torch
|
||||||
@@ -12,12 +13,12 @@ import yaml
|
|||||||
|
|
||||||
import scepter
|
import scepter
|
||||||
from scepter.modules.utils.file_system import FS
|
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 \
|
from scepter.studio.self_train.self_train_ui.component_names import \
|
||||||
TrainerUIName
|
TrainerUIName
|
||||||
from scepter.studio.self_train.utils.config_parser import (
|
from scepter.studio.self_train.utils.config_parser import (
|
||||||
get_default, get_values_by_model, get_values_by_model_version,
|
get_default, get_values_by_model, get_values_by_model_version,
|
||||||
get_values_by_model_version_tuner,
|
get_values_by_model_version_tuner)
|
||||||
get_values_by_model_version_tuner_resolution)
|
|
||||||
from scepter.studio.utils.uibase import UIBase
|
from scepter.studio.utils.uibase import UIBase
|
||||||
|
|
||||||
|
|
||||||
@@ -30,11 +31,17 @@ def print_memory_status(is_debug):
|
|||||||
return gpu_mem
|
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' # 🔄
|
refresh_symbol = '\U0001f504' # 🔄
|
||||||
|
|
||||||
|
|
||||||
def get_work_name(model, version, tuner, resolution):
|
def get_work_name(model, version, tuner):
|
||||||
model_prefix = f'Swift@{model}@{version}@{tuner}@{resolution}'
|
model_prefix = f'Swift@{model}@{version}@{tuner}'
|
||||||
return model_prefix + '@' + '{0:%Y%m%d%H%M%S%f}'.format(
|
return model_prefix + '@' + '{0:%Y%m%d%H%M%S%f}'.format(
|
||||||
datetime.datetime.now()) + ''.join(
|
datetime.datetime.now()) + ''.join(
|
||||||
[str(random.randint(1, 10)) for i in range(3)])
|
[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'):
|
def __init__(self, cfg, all_cfg_value, is_debug=False, language='en'):
|
||||||
self.BASE_CFG_VALUE = all_cfg_value
|
self.BASE_CFG_VALUE = all_cfg_value
|
||||||
self.para_data = get_default(self.BASE_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),
|
self.run_script = os.path.join(os.path.dirname(scepter.dirname),
|
||||||
cfg.SCRIPT_DIR, 'run_task.py')
|
cfg.SCRIPT_DIR, 'run_task.py')
|
||||||
self.work_dir_pre, _ = FS.map_to_local(cfg.WORK_DIR)
|
self.work_dir_pre, _ = FS.map_to_local(cfg.WORK_DIR)
|
||||||
self.is_debug = is_debug
|
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.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):
|
def create_ui(self):
|
||||||
with gr.Box():
|
with gr.Box():
|
||||||
with gr.Row(variant='panel', equal_height=True):
|
with gr.Row(variant='panel', equal_height=True):
|
||||||
@@ -88,15 +109,6 @@ class TrainerUI(UIBase):
|
|||||||
'model_default', ''),
|
'model_default', ''),
|
||||||
label=self.component_names.base_model,
|
label=self.component_names.base_model,
|
||||||
interactive=True)
|
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):
|
with gr.Column(scale=1, min_width=0):
|
||||||
self.base_model_revision = gr.Dropdown(
|
self.base_model_revision = gr.Dropdown(
|
||||||
choices=self.para_data.get(
|
choices=self.para_data.get(
|
||||||
@@ -107,13 +119,35 @@ class TrainerUI(UIBase):
|
|||||||
base_model_revision,
|
base_model_revision,
|
||||||
interactive=True)
|
interactive=True)
|
||||||
|
|
||||||
|
with gr.Row():
|
||||||
with gr.Column(scale=1, min_width=0):
|
with gr.Column(scale=1, min_width=0):
|
||||||
self.resolution = gr.Dropdown(
|
self.tuner_name = gr.Dropdown(
|
||||||
choices=self.para_data.get(
|
choices=self.para_data.get(
|
||||||
'resolution_choices', []),
|
'tuner_choices', []),
|
||||||
value=self.para_data.get(
|
value=self.para_data.get(
|
||||||
'resolution_default', 1024),
|
'tuner_default', ''),
|
||||||
label=self.component_names.resolution,
|
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,
|
allow_custom_value=True,
|
||||||
interactive=True)
|
interactive=True)
|
||||||
|
|
||||||
@@ -161,19 +195,36 @@ class TrainerUI(UIBase):
|
|||||||
value='')
|
value='')
|
||||||
|
|
||||||
with gr.Row():
|
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):
|
with gr.Column(scale=1, min_width=0):
|
||||||
self.work_name_button = gr.Button(
|
self.eval_prompts = gr.Dropdown(
|
||||||
value=refresh_symbol)
|
value=None,
|
||||||
with gr.Column(scale=2, min_width=0):
|
choices=self.train_para_data.get(
|
||||||
self.push_to_hub = gr.Checkbox(
|
'EVAL_PROMPTS', []),
|
||||||
label=self.component_names.push_to_hub,
|
label=self.component_names.eval_prompts,
|
||||||
value=False,
|
interactive=True,
|
||||||
visible=False)
|
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):
|
with gr.Row(variant='panel', equal_height=True):
|
||||||
self.examples = gr.Examples(
|
self.examples = gr.Examples(
|
||||||
@@ -193,15 +244,10 @@ class TrainerUI(UIBase):
|
|||||||
self.data_type, self.ms_data_space, self.ms_data_name,
|
self.data_type, self.ms_data_space, self.ms_data_name,
|
||||||
self.ms_data_subname
|
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):
|
with gr.Row(variant='panel', equal_height=True):
|
||||||
self.training_button = gr.Button()
|
self.training_button = gr.Button()
|
||||||
|
|
||||||
def set_callbacks(self, inference_ui):
|
def set_callbacks(self, inference_ui, manager):
|
||||||
def change_data_type(data_type):
|
def change_data_type(data_type):
|
||||||
if data_type == self.component_names.data_type_choices[0]:
|
if data_type == self.component_names.data_type_choices[0]:
|
||||||
return gr.Box(visible=False)
|
return gr.Box(visible=False)
|
||||||
@@ -217,7 +263,7 @@ class TrainerUI(UIBase):
|
|||||||
inputs=[
|
inputs=[
|
||||||
self.base_model,
|
self.base_model,
|
||||||
self.base_model_revision,
|
self.base_model_revision,
|
||||||
self.tuner_name, self.resolution
|
self.tuner_name
|
||||||
],
|
],
|
||||||
outputs=[self.work_name],
|
outputs=[self.work_name],
|
||||||
queue=False)
|
queue=False)
|
||||||
@@ -245,8 +291,12 @@ class TrainerUI(UIBase):
|
|||||||
interactive=True), \
|
interactive=True), \
|
||||||
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), \
|
interactive=True), \
|
||||||
gr.Dropdown(value=ret_data.get('resolution_default', 1024),
|
gr.Dropdown(value=ret_data.get('RESOLUTION', 1024)[0],
|
||||||
choices=ret_data.get('resolution_choices', []), interactive=True)
|
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,
|
self.base_model.change(fn=change_train_value_by_model,
|
||||||
inputs=[self.base_model],
|
inputs=[self.base_model],
|
||||||
@@ -255,7 +305,8 @@ class TrainerUI(UIBase):
|
|||||||
self.save_interval, self.train_batch_size,
|
self.save_interval, self.train_batch_size,
|
||||||
self.prompt_prefix,
|
self.prompt_prefix,
|
||||||
self.base_model_revision, self.tuner_name,
|
self.base_model_revision, self.tuner_name,
|
||||||
self.resolution
|
self.resolution_height,
|
||||||
|
self.resolution_width
|
||||||
],
|
],
|
||||||
queue=False)
|
queue=False)
|
||||||
|
|
||||||
@@ -282,10 +333,15 @@ class TrainerUI(UIBase):
|
|||||||
ret_data.get('SAVE_INTERVAL', 10), \
|
ret_data.get('SAVE_INTERVAL', 10), \
|
||||||
ret_data.get('TRAIN_BATCH_SIZE', 4), \
|
ret_data.get('TRAIN_BATCH_SIZE', 4), \
|
||||||
ret_data.get('TRAIN_PREFIX', ''), \
|
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), \
|
interactive=True), \
|
||||||
gr.Dropdown(value=ret_data.get('resolution_default', 1024),
|
gr.Dropdown(value=ret_data.get('RESOLUTION', 1024)[0],
|
||||||
choices=ret_data.get('resolution_choices', []), interactive=True)
|
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(
|
self.base_model_revision.change(
|
||||||
@@ -294,11 +350,10 @@ class TrainerUI(UIBase):
|
|||||||
outputs=[
|
outputs=[
|
||||||
self.train_epoch, self.learning_rate, self.save_interval,
|
self.train_epoch, self.learning_rate, self.save_interval,
|
||||||
self.train_batch_size, self.prompt_prefix, self.tuner_name,
|
self.train_batch_size, self.prompt_prefix, self.tuner_name,
|
||||||
self.resolution
|
self.resolution_height, self.resolution_width
|
||||||
],
|
],
|
||||||
queue=False)
|
queue=False)
|
||||||
|
|
||||||
#
|
|
||||||
#
|
#
|
||||||
def change_train_value_by_model_version_tuner(base_model,
|
def change_train_value_by_model_version_tuner(base_model,
|
||||||
base_model_revision,
|
base_model_revision,
|
||||||
@@ -323,8 +378,12 @@ class TrainerUI(UIBase):
|
|||||||
ret_data.get('SAVE_INTERVAL', 10), \
|
ret_data.get('SAVE_INTERVAL', 10), \
|
||||||
ret_data.get('TRAIN_BATCH_SIZE', 4), \
|
ret_data.get('TRAIN_BATCH_SIZE', 4), \
|
||||||
ret_data.get('TRAIN_PREFIX', ''), \
|
ret_data.get('TRAIN_PREFIX', ''), \
|
||||||
gr.Dropdown(value=ret_data.get('resolution_default', 1024),
|
gr.Dropdown(value=ret_data.get('RESOLUTION', 1024)[0],
|
||||||
choices=ret_data.get('resolution_choices', []), interactive=True)
|
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,
|
self.tuner_name.change(fn=change_train_value_by_model_version_tuner,
|
||||||
@@ -335,56 +394,30 @@ class TrainerUI(UIBase):
|
|||||||
outputs=[
|
outputs=[
|
||||||
self.train_epoch, self.learning_rate,
|
self.train_epoch, self.learning_rate,
|
||||||
self.save_interval, self.train_batch_size,
|
self.save_interval, self.train_batch_size,
|
||||||
self.prompt_prefix, self.resolution
|
self.prompt_prefix, self.resolution_height,
|
||||||
|
self.resolution_width
|
||||||
],
|
],
|
||||||
queue=False)
|
queue=False)
|
||||||
|
|
||||||
#
|
def change_resolution(h):
|
||||||
def change_train_value_by_model_version_tuner_resolution(
|
if h not in self.h_level_dict:
|
||||||
base_model, base_model_revision, tuner_name, resolution):
|
return gr.Dropdown()
|
||||||
'''
|
all_choices = self.h_level_dict[h]
|
||||||
Changes to the base model will affect the training parameters,
|
default = all_choices[0]
|
||||||
and it is best to define the related default values in the YAML.
|
return gr.Dropdown(choices=all_choices, value=default)
|
||||||
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', '')
|
|
||||||
|
|
||||||
#
|
self.resolution_height.change(change_resolution,
|
||||||
self.resolution.change(
|
inputs=[self.resolution_height],
|
||||||
fn=change_train_value_by_model_version_tuner_resolution,
|
outputs=[self.resolution_width],
|
||||||
inputs=[
|
queue=False)
|
||||||
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)
|
|
||||||
|
|
||||||
def run_train(work_name, data_type, ms_data_space, ms_data_name,
|
def run_train(work_name, data_type, ms_data_space, ms_data_name,
|
||||||
ms_data_subname, base_model, base_model_revision,
|
ms_data_subname, base_model, base_model_revision,
|
||||||
tuner_name, resolution, train_epoch, learning_rate,
|
tuner_name, resolution_height, resolution_width,
|
||||||
save_interval, train_batch_size, prompt_prefix,
|
train_epoch, learning_rate, save_interval,
|
||||||
replace_keywords, push_to_hub):
|
train_batch_size, prompt_prefix, replace_keywords,
|
||||||
|
push_to_hub, eval_prompts):
|
||||||
|
|
||||||
# Check Cuda
|
# Check Cuda
|
||||||
if not torch.cuda.is_available() and not self.is_debug:
|
if not torch.cuda.is_available() and not self.is_debug:
|
||||||
raise gr.Error(self.component_names.training_err1)
|
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 == '':
|
if work_name == 'custom' or work_name is None or work_name == '':
|
||||||
raise gr.Error(self.component_names.training_err4)
|
raise gr.Error(self.component_names.training_err4)
|
||||||
work_dir = os.path.join(self.work_dir_pre, work_name)
|
work_dir = os.path.join(self.work_dir_pre, work_name)
|
||||||
if not os.path.exists(work_dir):
|
self.current_train_model = work_name
|
||||||
os.makedirs(work_dir)
|
if os.path.exists(work_dir) or os.path.exists(
|
||||||
else:
|
f'.flag/{work_name}.tmp'):
|
||||||
raise gr.Error(self.component_names.training_err4)
|
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:
|
if push_to_hub:
|
||||||
model_id = work_name
|
model_id = work_name
|
||||||
@@ -404,28 +452,23 @@ class TrainerUI(UIBase):
|
|||||||
else:
|
else:
|
||||||
hub_model_id = ''
|
hub_model_id = ''
|
||||||
|
|
||||||
# Check Cuda Memory
|
|
||||||
if torch.cuda.is_available() and not self.is_debug:
|
|
||||||
device = torch.device('cuda:0')
|
|
||||||
required_memory_bytes = 40 * (1024**3)
|
|
||||||
try:
|
|
||||||
tensor = torch.empty( # noqa
|
|
||||||
(required_memory_bytes // 4, ), device=device
|
|
||||||
) # create 18GB tensor to check the memory if enough
|
|
||||||
del tensor
|
|
||||||
except RuntimeError:
|
|
||||||
raise gr.Error(self.component_names.training_err2)
|
|
||||||
|
|
||||||
# Check Instance Valid
|
# Check Instance Valid
|
||||||
if ms_data_name is None:
|
if ms_data_name is None:
|
||||||
raise gr.Error(self.component_names.training_err3)
|
raise gr.Error(self.component_names.training_err3)
|
||||||
|
|
||||||
st_time = time.time()
|
def prepare_train_data(data_cfg):
|
||||||
|
|
||||||
def prepare_data(data_cfg):
|
|
||||||
data_cfg['BATCH_SIZE'] = int(train_batch_size)
|
data_cfg['BATCH_SIZE'] = int(train_batch_size)
|
||||||
data_cfg['PROMPT_PREFIX'] = prompt_prefix
|
data_cfg['PROMPT_PREFIX'] = prompt_prefix
|
||||||
data_cfg['REPLACE_KEYWORDS'] = replace_keywords
|
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 data_type in self.component_names.data_type_choices:
|
||||||
if ms_data_name.startswith(
|
if ms_data_name.startswith(
|
||||||
'http') or ms_data_name.endswith('zip'):
|
'http') or ms_data_name.endswith('zip'):
|
||||||
@@ -482,7 +525,19 @@ class TrainerUI(UIBase):
|
|||||||
data_cfg['MS_REMAP_KEYS'] = {'Text': 'Prompt'}
|
data_cfg['MS_REMAP_KEYS'] = {'Text': 'Prompt'}
|
||||||
else:
|
else:
|
||||||
data_cfg['MS_REMAP_KEYS'] = None
|
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
|
return data_cfg
|
||||||
|
|
||||||
def prepare_train_config():
|
def prepare_train_config():
|
||||||
@@ -490,7 +545,7 @@ class TrainerUI(UIBase):
|
|||||||
current_model_info = self.BASE_CFG_VALUE[base_model][
|
current_model_info = self.BASE_CFG_VALUE[base_model][
|
||||||
base_model_revision]
|
base_model_revision]
|
||||||
modify_para = current_model_info['modify_para']
|
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:
|
if isinstance(modify_para, dict) and tuner_name in modify_para:
|
||||||
modify_c = modify_para[tuner_name]
|
modify_c = modify_para[tuner_name]
|
||||||
if isinstance(modify_c, dict) and 'TRAIN' in modify_c:
|
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']['MAX_EPOCHS'] = int(train_epoch)
|
||||||
cfg['SOLVER']['TRAIN_DATA']['BATCH_SIZE'] = int(
|
cfg['SOLVER']['TRAIN_DATA']['BATCH_SIZE'] = int(
|
||||||
train_batch_size)
|
train_batch_size)
|
||||||
cfg['SOLVER']['TUNER'] = current_model_info[
|
if 'TUNER' in cfg['SOLVER']:
|
||||||
'tuner_para'][tuner_name] if isinstance(
|
cfg['SOLVER']['TUNER'] = current_model_info['tuner_para'][
|
||||||
current_model_info['tuner_para'],
|
tuner_name] if isinstance(
|
||||||
dict) and tuner_name in current_model_info[
|
current_model_info['tuner_para'],
|
||||||
'tuner_para'] else None
|
dict) and tuner_name in current_model_info[
|
||||||
cfg['SOLVER']['TRAIN_DATA'] = prepare_data(
|
'tuner_para'] else None
|
||||||
|
cfg['SOLVER']['TRAIN_DATA'] = prepare_train_data(
|
||||||
cfg['SOLVER']['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']:
|
for hook in cfg['SOLVER']['TRAIN_HOOKS']:
|
||||||
if hook['NAME'] == 'CheckpointHook':
|
if hook['NAME'] == 'CheckpointHook':
|
||||||
hook['INTERVAL'] = save_interval
|
hook['INTERVAL'] = save_interval
|
||||||
hook['PUSH_TO_HUB'] = push_to_hub
|
hook['PUSH_TO_HUB'] = push_to_hub
|
||||||
hook['HUB_MODEL_ID'] = hub_model_id
|
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:
|
with open(cfg_file, 'w') as f_out:
|
||||||
yaml.dump(cfg,
|
yaml.dump(cfg,
|
||||||
@@ -539,48 +609,33 @@ class TrainerUI(UIBase):
|
|||||||
default_flow_style=False)
|
default_flow_style=False)
|
||||||
return cfg_file
|
return cfg_file
|
||||||
|
|
||||||
cfg = prepare_train_config()
|
before_kill_inference = self.trainer_ins.check_memory()
|
||||||
|
for k, v in manager.inference.pipe_manager.pipeline_level_modules.items(
|
||||||
def train_fn(cfg_file):
|
):
|
||||||
torch.cuda.empty_cache()
|
if hasattr(v, 'dynamic_unload'):
|
||||||
cmd = f'PYTHONPATH=. python {self.run_script} ' \
|
v.dynamic_unload(name='all')
|
||||||
f'--cfg={cfg_file} 2> {self.work_dir_pre}/std_out.txt'
|
after_kill_inference = self.trainer_ins.check_memory()
|
||||||
print(cmd)
|
message = f'GPU info: {before_kill_inference}. \n\n'
|
||||||
if not self.is_debug:
|
message += f'After unloading inference models, the GPU info: {after_kill_inference}. \n\n'
|
||||||
res = os.system(cmd)
|
_ = prepare_train_config()
|
||||||
else:
|
self.trainer_ins.start_task(work_name)
|
||||||
res = 0
|
message += self.trainer_ins.get_log(work_name)
|
||||||
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)
|
|
||||||
|
|
||||||
if work_name not in inference_ui.model_list:
|
if work_name not in inference_ui.model_list:
|
||||||
inference_ui.model_list.append(work_name)
|
inference_ui.model_list.append(work_name)
|
||||||
message = f'''
|
gr.Info('Start Training!' + message)
|
||||||
Training completed! \n
|
return gr.Dropdown.update(choices=inference_ui.model_list,
|
||||||
Save in [ {work_name} ] \n
|
value=work_name)
|
||||||
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)
|
|
||||||
|
|
||||||
self.training_button.click(
|
self.training_button.click(
|
||||||
run_train,
|
run_train,
|
||||||
inputs=[
|
inputs=[
|
||||||
self.work_name, self.data_type, self.ms_data_space,
|
self.work_name, self.data_type, self.ms_data_space,
|
||||||
self.ms_data_name, self.ms_data_subname, self.base_model,
|
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_epoch, self.learning_rate, self.save_interval,
|
||||||
self.train_batch_size, self.prompt_prefix,
|
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)
|
queue=True)
|
||||||
|
|||||||
@@ -10,7 +10,6 @@ paras_keys = [
|
|||||||
'MEMORY', 'EPOCHS', 'SAVE_INTERVAL', 'EPSEC', 'LEARNING_RATE',
|
'MEMORY', 'EPOCHS', 'SAVE_INTERVAL', 'EPSEC', 'LEARNING_RATE',
|
||||||
'IS_DEFAULT', 'TUNER'
|
'IS_DEFAULT', 'TUNER'
|
||||||
]
|
]
|
||||||
|
|
||||||
control_paras_keys = ['CONTROL_MODE', 'RESOLUTION', 'IS_DEFAULT']
|
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}]'
|
f'Para key {key} not defined in {config_file} META/RARAS[{idx}]'
|
||||||
)
|
)
|
||||||
assert key in para
|
assert key in para
|
||||||
tuner_type[para['TUNER'] + '@' + str(para['RESOLUTION'])] = para
|
tuner_type[para['TUNER']] = para
|
||||||
if para['IS_DEFAULT']:
|
if para['IS_DEFAULT']:
|
||||||
tuner_type['default'] = para['TUNER'] + '@' + str(
|
tuner_type['default'] = para['TUNER']
|
||||||
para['RESOLUTION'])
|
|
||||||
|
|
||||||
tuner_type['choices'] = list(tuner_type.keys())
|
tuner_type['choices'] = list(tuner_type.keys())
|
||||||
if 'default' in tuner_type['choices']:
|
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['default'] = tuner_type['choices'][0] if len(
|
||||||
tuner_type['choices']) > 0 else ''
|
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)
|
tuner_paras = meta_cfg.get('TUNERS', None)
|
||||||
return tuner_type, tuner_paras
|
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}]'
|
f'Para key {key} not defined in {config_file} META/RARAS[{idx}]'
|
||||||
)
|
)
|
||||||
assert key in para
|
assert key in para
|
||||||
control_type[para['CONTROL_MODE'] + '@' +
|
control_type[para['CONTROL_MODE']] = para
|
||||||
str(para['RESOLUTION'])] = para
|
|
||||||
# control_type[para["CONTROL_MODE"]] = para
|
# control_type[para["CONTROL_MODE"]] = para
|
||||||
if para['IS_DEFAULT']:
|
if para['IS_DEFAULT']:
|
||||||
control_type['default'] = para['CONTROL_MODE'] + '@' + str(
|
control_type['default'] = para['CONTROL_MODE']
|
||||||
para['RESOLUTION'])
|
|
||||||
# control_type["default"] = para["CONTROL_MODE"]
|
|
||||||
|
|
||||||
control_type['choices'] = list(control_type.keys())
|
control_type['choices'] = list(control_type.keys())
|
||||||
if 'default' in control_type['choices']:
|
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['default'] = control_type['choices'][0] if len(
|
||||||
control_type['choices']) > 0 else ''
|
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
|
return control_type, paras
|
||||||
|
|
||||||
|
|
||||||
@@ -166,26 +145,19 @@ def get_default(config_dict):
|
|||||||
return ret_data
|
return ret_data
|
||||||
ret_data['version_choices'] = default_version_cfg['choices']
|
ret_data['version_choices'] = default_version_cfg['choices']
|
||||||
ret_data['version_default'] = default_version_cfg['default']
|
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)
|
None)
|
||||||
if default_tuner_cfg is None:
|
if default_model_cfg is None:
|
||||||
return ret_data
|
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'] != '':
|
||||||
default_tuner_cfg = default_tuner_cfg['tuner_type']
|
default_tuner_cfg = default_model_cfg['tuner_type']
|
||||||
else:
|
else:
|
||||||
return ret_data
|
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']
|
defalt_t_type = default_tuner_cfg['default']
|
||||||
|
ret_data['tuner_default'] = defalt_t_type
|
||||||
type_paras = default_tuner_cfg.get(defalt_t_type, None)
|
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:
|
if type_paras is not None:
|
||||||
ret_data.update(type_paras)
|
ret_data.update(type_paras)
|
||||||
return ret_data
|
return ret_data
|
||||||
@@ -198,21 +170,15 @@ def get_values_by_model(config_dict, model_name):
|
|||||||
return ret_data
|
return ret_data
|
||||||
ret_data['version_choices'] = version_cfg['choices']
|
ret_data['version_choices'] = version_cfg['choices']
|
||||||
ret_data['version_default'] = version_cfg['default']
|
ret_data['version_default'] = version_cfg['default']
|
||||||
default_tuner_cfg = version_cfg.get(version_cfg['default'], None)
|
default_model_cfg = version_cfg.get(version_cfg['default'], None)
|
||||||
if default_tuner_cfg is None:
|
if default_model_cfg is None:
|
||||||
return ret_data
|
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']
|
defalt_t_type = default_tuner_cfg['default']
|
||||||
|
ret_data['tuner_default'] = defalt_t_type
|
||||||
type_paras = default_tuner_cfg.get(defalt_t_type, None)
|
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:
|
if type_paras is not None:
|
||||||
ret_data.update(type_paras)
|
ret_data.update(type_paras)
|
||||||
return ret_data
|
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)
|
tuner_cfg = version_cfg.get(version, None)
|
||||||
if tuner_cfg is None:
|
if tuner_cfg is None:
|
||||||
return ret_data
|
return ret_data
|
||||||
|
|
||||||
default_tuner_cfg = tuner_cfg['tuner_type']
|
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']
|
defalt_t_type = default_tuner_cfg['default']
|
||||||
|
ret_data['tuner_default'] = defalt_t_type
|
||||||
type_paras = default_tuner_cfg.get(defalt_t_type, None)
|
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:
|
if type_paras is not None:
|
||||||
ret_data.update(type_paras)
|
ret_data.update(type_paras)
|
||||||
return ret_data
|
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)
|
version_cfg = config_dict.get(model_name, None)
|
||||||
if version_cfg is None:
|
if version_cfg is None:
|
||||||
return ret_data
|
return ret_data
|
||||||
tuner_cfg = version_cfg.get(version, None)
|
model_cfg = version_cfg.get(version, None)
|
||||||
if tuner_cfg is None:
|
if model_cfg is None:
|
||||||
return ret_data
|
return ret_data
|
||||||
tuner_cfg = tuner_cfg['tuner_type']
|
tuner_cfg = model_cfg['tuner_type']
|
||||||
|
type_paras = tuner_cfg.get(tuner_name, None)
|
||||||
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)
|
|
||||||
if type_paras is not None:
|
if type_paras is not None:
|
||||||
ret_data.update(type_paras)
|
ret_data.update(type_paras)
|
||||||
return ret_data
|
return ret_data
|
||||||
@@ -371,15 +308,8 @@ def get_control_default(config_dict):
|
|||||||
|
|
||||||
ret_data['control_choices'] = list(default_control_cfg['choices'].keys())
|
ret_data['control_choices'] = list(default_control_cfg['choices'].keys())
|
||||||
defalt_t_type = default_control_cfg['default']
|
defalt_t_type = default_control_cfg['default']
|
||||||
|
ret_data['control_default'] = defalt_t_type
|
||||||
type_paras = default_control_cfg.get(defalt_t_type, None)
|
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:
|
if type_paras is not None:
|
||||||
ret_data.update(type_paras)
|
ret_data.update(type_paras)
|
||||||
return ret_data
|
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 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
import argparse
|
import argparse
|
||||||
|
import importlib
|
||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
#
|
#
|
||||||
from scepter.modules.utils.registry import REGISTRY_LIST
|
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))
|
sys.path.insert(0, os.path.abspath(os.curdir))
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,7 +1,9 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
import argparse
|
import argparse
|
||||||
|
import importlib
|
||||||
import os
|
import os
|
||||||
|
import sys
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
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.file_system import FS
|
||||||
from scepter.modules.utils.logger import get_logger
|
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):
|
def run_task(cfg):
|
||||||
std_logger = get_logger(name='scepter')
|
std_logger = get_logger(name='scepter')
|
||||||
|
|||||||
@@ -1,12 +1,22 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||||
import argparse
|
import argparse
|
||||||
|
import importlib
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
|
||||||
from scepter.modules.solver.registry import SOLVERS
|
from scepter.modules.solver.registry import SOLVERS
|
||||||
from scepter.modules.utils.config import Config
|
from scepter.modules.utils.config import Config
|
||||||
from scepter.modules.utils.distribute import we
|
from scepter.modules.utils.distribute import we
|
||||||
from scepter.modules.utils.logger import get_logger
|
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):
|
def run_task(cfg):
|
||||||
std_logger = get_logger(name='scepter')
|
std_logger = get_logger(name='scepter')
|
||||||
@@ -18,19 +28,21 @@ def run_task(cfg):
|
|||||||
|
|
||||||
def update_config(cfg):
|
def update_config(cfg):
|
||||||
if hasattr(cfg.args, 'learning_rate') and cfg.args.learning_rate:
|
if hasattr(cfg.args, 'learning_rate') and cfg.args.learning_rate:
|
||||||
print(
|
if cfg.SOLVER.OPTIMIZER.get('LEARNING_RATE', None) is not None:
|
||||||
f'learning_rate change from {cfg.SOLVER.OPTIMIZER.LEARNING_RATE} to {cfg.args.learning_rate}'
|
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)
|
cfg.SOLVER.OPTIMIZER.LEARNING_RATE = float(cfg.args.learning_rate)
|
||||||
if hasattr(cfg.args, 'max_steps') and cfg.args.max_steps:
|
if hasattr(cfg.args, 'max_steps') and cfg.args.max_steps:
|
||||||
print(
|
if cfg.SOLVER.get('MAX_STEPS', None) is not None:
|
||||||
f'max_steps change from {cfg.SOLVER.MAX_STEPS} to {cfg.args.max_steps}'
|
print(
|
||||||
)
|
f'max_steps change from {cfg.SOLVER.MAX_STEPS} to {cfg.args.max_steps}'
|
||||||
|
)
|
||||||
cfg.SOLVER.MAX_STEPS = int(cfg.args.max_steps)
|
cfg.SOLVER.MAX_STEPS = int(cfg.args.max_steps)
|
||||||
return cfg
|
return cfg
|
||||||
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
def run():
|
||||||
parser = argparse.ArgumentParser(description='Argparser for Scepter:\n')
|
parser = argparse.ArgumentParser(description='Argparser for Scepter:\n')
|
||||||
parser.add_argument('--learning_rate',
|
parser.add_argument('--learning_rate',
|
||||||
dest='learning_rate',
|
dest='learning_rate',
|
||||||
@@ -44,3 +56,7 @@ if __name__ == '__main__':
|
|||||||
cfg = Config(load=True, parser_ins=parser)
|
cfg = Config(load=True, parser_ins=parser)
|
||||||
cfg = update_config(cfg)
|
cfg = update_config(cfg)
|
||||||
we.init_env(cfg, logger=None, fn=run_task)
|
we.init_env(cfg, logger=None, fn=run_task)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
run()
|
||||||
|
|||||||