Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0010d4282a | ||
|
|
8a9b791e3f | ||
|
|
ee7fe888f2 | ||
|
|
36b259b4b4 | ||
|
|
92c849412e | ||
|
|
2a7e026f84 | ||
|
|
cdba82baf8 | ||
|
|
565c7957d8 | ||
|
|
e00c23d09a | ||
|
|
bf53829530 | ||
|
|
35aada8ce8 | ||
|
|
d3ce651bf7 | ||
|
|
7a58c91940 | ||
|
|
4e1606af2d | ||
|
|
09459c11b7 | ||
|
|
d9a48268d5 | ||
|
|
9f1847501d | ||
|
|
8214227098 | ||
|
|
2e69b2b116 | ||
|
|
3440ec7c38 | ||
|
|
9999e0e1f9 | ||
|
|
01c03683e8 | ||
|
|
9adb273e4b |
@@ -0,0 +1,2 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
|
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 |
@@ -914,7 +914,7 @@ data = {
|
||||
_model(data)
|
||||
probe = _model.probe_data()
|
||||
for key in probe:
|
||||
print(key, probe[key].to_log(prefix=f"xxx/dev_easytorch/{key}"))
|
||||
print(key, probe[key].to_log(prefix=f"xxx/{key}"))
|
||||
```
|
||||
<hr/>
|
||||
|
||||
|
||||
@@ -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,14 +9,24 @@
|
||||
</p>
|
||||
|
||||
## 📖 Table of Contents
|
||||
- [Introduction](#-introduction)
|
||||
- [News](#-news)
|
||||
- [Installation](#-Installation)
|
||||
- [Introduction](#-introduction)
|
||||
- [Installation](#%EF%B8%8F-installation)
|
||||
- [Getting Started](#-getting-started)
|
||||
- [SCEPTER Studio](#-scepter-studio)
|
||||
- [SCEPTER Studio](#%EF%B8%8F-scepter-studio)
|
||||
- [Gallery](#%EF%B8%8F-gallery)
|
||||
- [Features](#-features)
|
||||
- [Learn More](#-learn-more)
|
||||
- [License](#license)
|
||||
- [Acknowledgement](#acknowledgement)
|
||||
|
||||
## 🎉 News
|
||||
- [2024.03]: We optimize the training UI and checkpoint management. New [LAR-Gen](https://arxiv.org/abs/2403.19534) model has been added on SCEPTER Studio, supporting `zoom-out`, `virtual try on`, `inpainting`.
|
||||
- [2024.02]: We release new SCEdit controllable image synthesis models for SD v2.1 and SD XL. Multiple strategies applied to accelerate inference time for SCEPTER Studio.
|
||||
- [2024.01]: We release **SCEPTER Studio**, an integrated toolkit for data management, model training and inference based on [Gradio](https://www.gradio.app/).
|
||||
- [2024.01]: [SCEdit](https://arxiv.org/abs/2312.11392) support controllable image synthesis for training and inference.
|
||||
- [2023.12]: We propose [SCEdit](https://arxiv.org/abs/2312.11392), an efficient and controllable generation framework.
|
||||
- [2023.12]: We release [🪄SCEPTER](https://github.com/modelscope/scepter/) library.
|
||||
|
||||
## 📝 Introduction
|
||||
|
||||
@@ -27,7 +37,7 @@ Main Feature:
|
||||
- Task:
|
||||
- Text-to-image generation
|
||||
- Controllable image synthesis
|
||||
- Image editing (TODO)
|
||||
- Image editing
|
||||
- Training / Inference:
|
||||
- Distribute: DDP / FSDP / FairScale / Xformers
|
||||
- File system: Local / Http / OSS / Modelscope
|
||||
@@ -36,17 +46,12 @@ Main Feature:
|
||||
- Training
|
||||
- Inference
|
||||
|
||||
Currently supported approches (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)
|
||||
2. SCEdit: [SCEdit: Efficient and Controllable Image Diffusion Generation via Skip Connection Editing](https://arxiv.org/abs/2312.11392) [](https://arxiv.org/abs/2312.11392) [](https://scedit.github.io/)
|
||||
3. Res-Tuning(TODO): [Res-Tuning: A Flexible and Efficient Tuning Paradigm via Unbinding Tuner from Backbone](https://arxiv.org/abs/2310.19859) [](https://arxiv.org/abs/2310.19859) [](https://res-tuning.github.io/)
|
||||
|
||||
## 🎉 News
|
||||
- [2024.01]: We release **SCEPTER Studio**, an integrated toolkit for data management, model training and inference based on [Gradio](https://www.gradio.app/).
|
||||
- [2024.01]: [SCEdit](https://arxiv.org/abs/2312.11392) support controllable image synthesis for training and inference.
|
||||
- [2023.12]: We propose [SCEdit](https://arxiv.org/abs/2312.11392), an efficient and controllable generation framework.
|
||||
- [2023.12]: We release [🪄SCEPTER](https://github.com/modelscope/scepter/) library.
|
||||
2. SCEdit(CVPR2024): [SCEdit: Efficient and Controllable Image Diffusion Generation via Skip Connection Editing](https://arxiv.org/abs/2312.11392) [](https://arxiv.org/abs/2312.11392) [](https://scedit.github.io/)
|
||||
3. Res-Tuning(NeurIPS2023 TODO): [Res-Tuning: A Flexible and Efficient Tuning Paradigm via Unbinding Tuner from Backbone](https://arxiv.org/abs/2310.19859) [](https://arxiv.org/abs/2310.19859) [](https://res-tuning.github.io/)
|
||||
4. LAR-Gen: [Locate, Assign, Refine: Taming Customized Image Inpainting with Text-Subject Guidance](https://arxiv.org/abs/2403.19534) [](https://arxiv.org/abs/2403.19534) [](https://ali-vilab.github.io/largen-page/)
|
||||
|
||||
## 🛠️ Installation
|
||||
|
||||
@@ -56,13 +61,17 @@ Currently supported approches (and counting):
|
||||
conda env create -f environment.yaml
|
||||
conda activate scepter
|
||||
```
|
||||
- We recommend installing the specific version of PyTorch and accelerate toolbox [xFormers](https://pypi.org/project/xformers/). You can install these recommended version by pip:
|
||||
|
||||
```shell
|
||||
pip install -r requirements/recommended.txt
|
||||
```
|
||||
|
||||
- Install SCEPTER by the `pip` command:
|
||||
|
||||
```shell
|
||||
pip install scepter
|
||||
```
|
||||
- PS: We recommend installing PyTorch follwing [official documentation](https://pytorch.org/get-started/locally/)
|
||||
|
||||
## 🚀 Getting Started
|
||||
|
||||
@@ -88,7 +97,7 @@ For the data format used by SCEPTER Studio, please refer to [3D_example_csv.zip]
|
||||
To facilitate starting training in command-line mode, you can use a dataset in text format, please refer to [3D_example_txt.zip](https://modelscope.cn/api/v1/models/damo/scepter/repo?Revision=master&FilePath=datasets/3D_example_txt.zip)
|
||||
|
||||
```shell
|
||||
mkdir -p cache/dataset/ && wget 'https://modelscope.cn/api/v1/models/damo/scepter_scedit/repo?Revision=master&FilePath=dataset/3D_example_txt.zip' -O cache/dataset/3D_example_txt.zip && unzip cache/dataset/3D_example_txt.zip -d cache/dataset/ && rm cache/dataset/3D_example_txt.zip
|
||||
mkdir -p cache/datasets/ && wget 'https://modelscope.cn/api/v1/models/damo/scepter_scedit/repo?Revision=master&FilePath=dataset/3D_example_txt.zip' -O cache/datasets/3D_example_txt.zip && unzip cache/datasets/3D_example_txt.zip -d cache/datasets/ && rm cache/datasets/3D_example_txt.zip
|
||||
```
|
||||
|
||||
### Training
|
||||
@@ -159,6 +168,13 @@ python scepter/tools/run_inference.py --cfg scepter/methods/scedit/ctr/sd21_768_
|
||||
python scepter/tools/run_inference.py --cfg scepter/methods/scedit/ctr/sd21_768_sce_ctr_pose.yaml --num_samples 1 --prompt 'super mario' --save_folder 'test_mario_pose' --image_size 768 --task control --image 'asset/images/pose_source.png' --control_mode source --pretrained_model ms://damo/scepter_scedit@controllable_model/SD2.1/pose_control/0_SwiftSCETuning/pytorch_model.bin # pose
|
||||
```
|
||||
|
||||
### Customize Modules
|
||||
Refer to `example`, build the modules of your task in `example/{task}`.
|
||||
```python
|
||||
cd example/classifier
|
||||
python run.py --cfg classifier.yaml
|
||||
```
|
||||
|
||||
|
||||
## 🖥️ SCEPTER Studio
|
||||
|
||||
@@ -176,15 +192,128 @@ git clone https://github.com/modelscope/scepter.git
|
||||
PYTHONPATH=. python scepter/tools/webui.py --cfg scepter/methods/studio/scepter_ui.yaml
|
||||
```
|
||||
|
||||
The startup of **SCEPTER Studio** eliminates the need for manual downloading and organizing of models; it will automatically load the corresponding models and store them in a local directory.
|
||||
Depending on the network and hardware situation, the initial startup usually requires 15-60 minutes, primarily involving the download and processing of SDv1.5, SDv2.1, and SDXL models.
|
||||
The startup of **SCEPTER Studio** eliminates the need for manual downloading and organizing of models; it will automatically load the corresponding models and store them in a local directory.
|
||||
Depending on the network and hardware situation, the initial startup usually requires 15-60 minutes, primarily involving the download and processing of SDv1.5, SDv2.1, and SDXL models.
|
||||
Therefore, subsequent startups will become much faster (about one minute) as downloading is no longer required.
|
||||
|
||||
* LAR-Gen: we release `zoom-out`, `virtual try on`, `inpainting(text guided)`, `inpainting(text + reference image guided)` image editing capabilities.
|
||||
Please note that the **Data Preprocess** button must be clicked before clicking the **Generate** button.
|
||||
<p align="center">
|
||||
<img src="https://raw.githubusercontent.com/ali-vilab/largen-page/main/public/images/largen.gif">
|
||||
</p>
|
||||
|
||||
### Modelscope Studio
|
||||
|
||||
We deploy a work studio on Modelscope that includes only the inference tab, please refer to [ms_scepter_studio](https://www.modelscope.cn/studios/damo/scepter_studio/summary)
|
||||
|
||||
## 🖼️ Gallery
|
||||
|
||||
### LAR-Gen: Zoom Out
|
||||
<table>
|
||||
<tr>
|
||||
<td><strong>Origin Image</strong><br>Prompt: a temple on fire</td>
|
||||
<td><strong>Zoom-Out</strong><br>CenterAround:0.75</td>
|
||||
<td><strong>Zoom-Out</strong><br>CenterAround:0.75</td>
|
||||
<td><strong>Zoom-Out</strong><br>CenterAround:0.75</td>
|
||||
<td><strong>Zoom-Out</strong><br>CenterAround:0.75</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><img src="asset/images/zoom_out/ex1_scene_im.jpg" width="240"></td>
|
||||
<td><img src="asset/images/zoom_out/ex1_zoom_out1.jpg" width="240"></td>
|
||||
<td><img src="./asset/images/zoom_out/ex1_zoom_out2.jpg" width="240"></td>
|
||||
<td><img src="./asset/images/zoom_out/ex1_zoom_out3.jpg" width="240"></td>
|
||||
<td><img src="./asset/images/zoom_out/ex1_zoom_out4.jpg" width="240"></td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
### LAR-Gen: Virtual Try-on
|
||||
<table>
|
||||
<tr>
|
||||
<td><strong>Model Image</strong></td>
|
||||
<td><strong>Model Mask</strong></td>
|
||||
<td><strong>Clothing Image</strong></td>
|
||||
<td><strong>Clothing Mask</strong></td>
|
||||
<td><strong>Try-on Output</strong></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><img src="asset/images/virtual_try_on/model.jpg" width="240"></td>
|
||||
<td><img src="asset/images/virtual_try_on/ex2_scene_mask.jpg" width="240"></td>
|
||||
<td><img src="asset/images/virtual_try_on/tshirt.jpg" width="240"></td>
|
||||
<td><img src="asset/images/virtual_try_on/ex2_subject_mask.jpg" width="240"></td>
|
||||
<td><img src="asset/images/virtual_try_on/try_on_out.jpg" width="240"></td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
### LAR-Gen: Inpainting (Text guided)
|
||||
<table>
|
||||
<tr>
|
||||
<td><strong>Origin Image</strong><br>Prompt: a blue and white porcelain</td>
|
||||
<td><strong>Inpainting Mask1</strong></td>
|
||||
<td><strong>Inpainting Output1</strong></td>
|
||||
<td><strong>Inpainting Mask2</strong><br>Prompt: a clock</td>
|
||||
<td><strong>Inpainting Output2</strong></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><img src="asset/images/inpainting_text/ex3_scene_im.jpg" width="240"></td>
|
||||
<td><img src="asset/images/inpainting_text/ex3_scene_mask.jpg" width="240"></td>
|
||||
<td><img src="asset/images/inpainting_text/inpainting_text.jpg" width="240"></td>
|
||||
<td><img src="asset/images/inpainting_text/ex3_scene_mask2.jpg" width="240"></td>
|
||||
<td><img src="asset/images/inpainting_text/inpainting_text2.jpg" width="240"></td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
### LAR-Gen: Inpainting (Text and Subject guided)
|
||||
<table>
|
||||
<tr>
|
||||
<td><strong>Origin Image</strong><br>Prompt: a dog wearing sunglasses</td>
|
||||
<td><strong>Origin Mask</strong></td>
|
||||
<td><strong>Reference Image</strong></td>
|
||||
<td><strong>Reference Mask</strong></td>
|
||||
<td><strong>Inpainting Output</strong></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><img src="asset/images/inpainting_text_ref/ex4_scene_im.jpg" width="240"></td>
|
||||
<td><img src="asset/images/inpainting_text_ref/ex4_scene_mask.jpg" width="240"></td>
|
||||
<td><img src="asset/images/inpainting_text_ref/ex4_subject_im.jpg" width="240"></td>
|
||||
<td><img src="asset/images/inpainting_text_ref/ex4_subject_mask.jpg" width="240"></td>
|
||||
<td><img src="asset/images/inpainting_text_ref/inpainting_text_ref.jpg" width="240"></td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
### Dragon Year Special: Dragon Tuner
|
||||
|
||||
<table>
|
||||
<tr>
|
||||
<td><strong>Gold Dragon Tuner</strong></td>
|
||||
<td><strong>Sloppy Dragon Tuner</strong></td>
|
||||
<td><strong>Red Dragon Tuner</strong><br> + Papercraft Mantra</td>
|
||||
<td><strong>Azure Dragon Tuner</strong><br> + Pose Control</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><img src="https://github.com/hanzhn/datas/blob/main/scepter/readme/tuner_gold_dragon.jpeg?raw=true" width="300"></td>
|
||||
<td><img src="https://github.com/hanzhn/datas/blob/main/scepter/readme/tuner_sloppy_dragon.jpeg?raw=true" width="300"></td>
|
||||
<td><img src="https://github.com/hanzhn/datas/blob/main/scepter/readme/tuner_mantra_papercraft_dragon.jpeg?raw=true" width="300"></td>
|
||||
<td><img src="https://github.com/hanzhn/datas/blob/main/scepter/readme/tuner_pose.jpeg?raw=true" width="300"></td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
### Text Effect Image
|
||||
|
||||
<table>
|
||||
<tr>
|
||||
<td><strong>Conditional Image</strong></td>
|
||||
<td><strong>Midas Control</strong><br>"Race track, top view"</td>
|
||||
<td><strong>Midas Control</strong><br> + Watercolor Mantra<br>"white lilies"</td>
|
||||
<td><strong>Midas Control</strong><br> + Dragon Tuner<br>"Spring Festival, Chinese dragon"</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><img src="https://github.com/hanzhn/datas/blob/main/scepter/readme/word_condition.png?raw=true" width="300"></td>
|
||||
<td><img src="https://github.com/hanzhn/datas/blob/main/scepter/readme/word_race.jpeg?raw=true" width="300"></td>
|
||||
<td><img src="https://github.com/hanzhn/datas/blob/main/scepter/readme/word_lilies.jpeg?raw=true" width="300"></td>
|
||||
<td><img src="https://github.com/hanzhn/datas/blob/main/scepter/readme/word_festival.jpeg?raw=true" width="300"></td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
## ✨ Features
|
||||
|
||||
### Text-to-Image Generation
|
||||
@@ -201,25 +330,34 @@ We deploy a work studio on Modelscope that includes only the inference tab, plea
|
||||
| **Model** | **Canny** | **HED** | **Depth** | **Pose** | **Color** |
|
||||
|:---------:|:---------:|:-------:|:---------:|:--------:|:---------:|
|
||||
| SD 1.5 | ✅ | ✅ | ✅ | ✅ | ✅ |
|
||||
| SD 2.1 | 🪄 | ✅ | ✅ | 🪄 | 🪄 |
|
||||
| SD XL | ✅ | ✅ | ✅ | ✅ | ✅ |
|
||||
| SD 2.1 | 🪄 | 🪄 | 🪄 | 🪄 | 🪄 |
|
||||
| SD XL | 🪄 | 🪄 | 🪄 | 🪄 | 🪄 |
|
||||
|
||||
### Image Editing
|
||||
- LAR-Gen
|
||||
|
||||
| **Model** | **Locate** | **Assign** | **Refine** |
|
||||
|:---------:|:----------:|:----------:|:----------:|
|
||||
| SD XL | 🪄 | 🪄 | ⏳ |
|
||||
|
||||
### Model URL
|
||||
|
||||
- ✅ indicates support for both training and inference.
|
||||
- 🪄 denotes that the model has been published.
|
||||
- ⏳ denotes that the module has not been integrated currently.
|
||||
- More models will be released in the future.
|
||||
|
||||
| Model | URL |
|
||||
|--------|-------------------------------------------------------------------------------------|
|
||||
| SCEdit | [ModelCard](https://modelscope.cn/models/damo/scepter_scedit/summary) |
|
||||
| Model | URL |
|
||||
|--------|------------------------------------------------------------------------------------------------------------------------------------------------|
|
||||
| SCEdit | [ModelScope](https://modelscope.cn/models/iic/scepter_scedit/summary) [HuggingFace](https://huggingface.co/scepter-studio/scepter_scedit) |
|
||||
| LAR-Gen | [ModelScope](https://www.modelscope.cn/models/iic/LARGEN/summary) |
|
||||
|
||||
PS: Scripts running within the SCEPTER framework will automatically fetch and load models based on the required dependency files, eliminating the need for manual downloads.
|
||||
|
||||
|
||||
## 🔍 Learn More
|
||||
|
||||
- [Alibaba TongYi Vision Intelligence Lab](https://github.com/damo-vilab)
|
||||
- [Alibaba TongYi Vision Intelligence Lab](https://github.com/ali-vilab)
|
||||
|
||||
Discover more about open-source projects on image generation, video generation, and editing tasks.
|
||||
|
||||
@@ -231,7 +369,20 @@ PS: Scripts running within the SCEPTER framework will automatically fetch and lo
|
||||
|
||||
SWIFT (Scalable lightWeight Infrastructure for Fine-Tuning) is an extensible framwork designed to faciliate lightweight model fine-tuning and inference.
|
||||
|
||||
## BibTeX
|
||||
If our work is useful for your research, please consider citing:
|
||||
```bibtex
|
||||
@misc{scepter,
|
||||
title = {SCEPTER, https://github.com/modelscope/scepter},
|
||||
author = {SCEPTER},
|
||||
year = {2023}
|
||||
}
|
||||
```
|
||||
|
||||
|
||||
## License
|
||||
|
||||
This project is licensed under the [Apache License (Version 2.0)](https://github.com/modelscope/modelscope/blob/master/LICENSE).
|
||||
|
||||
## Acknowledgement
|
||||
Thanks to [Stability-AI](https://github.com/Stability-AI), [SWIFT library](https://github.com/modelscope/swift/) and [Fooocus](https://github.com/lllyasviel/Fooocus) for their awesome work.
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
albumentations
|
||||
bezier
|
||||
einops
|
||||
modelscope
|
||||
ms-swift>=1.5.2
|
||||
@@ -6,6 +8,8 @@ open_clip_torch
|
||||
opencv-python
|
||||
opencv_transforms>=0.0.6
|
||||
oss2>=2.15.0
|
||||
pycocotools
|
||||
pyyaml>=5.3.1
|
||||
scikit-image
|
||||
torchsde
|
||||
transformers
|
||||
xformers>=0.0.21
|
||||
|
||||
@@ -0,0 +1,4 @@
|
||||
git+https://github.com/cocodataset/panopticapi.git
|
||||
torch==2.0.1
|
||||
torchvision==0.15.2
|
||||
xformers==0.0.21
|
||||
@@ -1,2 +1,3 @@
|
||||
gradio>=3.47.1,<4.0.0
|
||||
imagehash
|
||||
psutil
|
||||
|
||||
@@ -117,7 +117,7 @@ SOLVER:
|
||||
SAMPLE_STEPS: 50
|
||||
SEED: 2023
|
||||
GUIDE_SCALE: 7.5
|
||||
GUIDE_RESCALE:
|
||||
GUIDE_RESCALE: 0.5
|
||||
DISCRETIZATION: trailing
|
||||
IMAGE_SIZE: [512, 512]
|
||||
RUN_TRAIN_N: False
|
||||
|
||||
@@ -125,7 +125,7 @@ SOLVER:
|
||||
SAMPLE_STEPS: 50
|
||||
SEED: 2023
|
||||
GUIDE_SCALE: 7.5
|
||||
GUIDE_RESCALE:
|
||||
GUIDE_RESCALE: 0.5
|
||||
DISCRETIZATION: trailing
|
||||
IMAGE_SIZE: [512, 512]
|
||||
RUN_TRAIN_N: False
|
||||
|
||||
@@ -0,0 +1,214 @@
|
||||
ENV:
|
||||
BACKEND: nccl
|
||||
SOLVER:
|
||||
NAME: LatentDiffusionSolver
|
||||
RESUME_FROM:
|
||||
LOAD_MODEL_ONLY: True
|
||||
USE_FSDP: False
|
||||
SHARDING_STRATEGY:
|
||||
USE_AMP: True
|
||||
DTYPE: float16
|
||||
CHANNELS_LAST: True
|
||||
MAX_STEPS: 2000
|
||||
MAX_EPOCHS: -1
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sd21_512_full
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusion
|
||||
PARAMETERIZATION: eps
|
||||
TIMESTEPS: 1000
|
||||
MIN_SNR_GAMMA:
|
||||
ZERO_TERMINAL_SNR: False
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-2-1-base@v2-1_512-ema-pruned.safetensors
|
||||
IGNORE_KEYS: [ ]
|
||||
SCALE_FACTOR: 0.18215
|
||||
SIZE_FACTOR: 8
|
||||
DEFAULT_N_PROMPT:
|
||||
SCHEDULE_ARGS:
|
||||
"NAME": "scaled_linear"
|
||||
"BETA_MIN": 0.00085
|
||||
"BETA_MAX": 0.012
|
||||
USE_EMA: False
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: DiffusionUNet
|
||||
IN_CHANNELS: 4
|
||||
OUT_CHANNELS: 4
|
||||
MODEL_CHANNELS: 320
|
||||
NUM_HEADS_CHANNELS: 64
|
||||
NUM_RES_BLOCKS: 2
|
||||
ATTENTION_RESOLUTIONS: [ 4, 2, 1 ]
|
||||
CHANNEL_MULT: [ 1, 2, 4, 4 ]
|
||||
CONV_RESAMPLE: True
|
||||
DIMS: 2
|
||||
USE_CHECKPOINT: False
|
||||
USE_SCALE_SHIFT_NORM: False
|
||||
RESBLOCK_UPDOWN: False
|
||||
USE_SPATIAL_TRANSFORMER: True
|
||||
TRANSFORMER_DEPTH: 1
|
||||
CONTEXT_DIM: 1024
|
||||
DISABLE_MIDDLE_SELF_ATTN: False
|
||||
USE_LINEAR_IN_TRANSFORMER: True
|
||||
PRETRAINED_MODEL:
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKL
|
||||
EMBED_DIM: 4
|
||||
PRETRAINED_MODEL:
|
||||
IGNORE_KEYS: [ ]
|
||||
BATCH_SIZE: 4
|
||||
#
|
||||
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
|
||||
#
|
||||
TOKENIZER:
|
||||
NAME: OpenClipTokenizer
|
||||
LENGTH: 77
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: FrozenOpenCLIPEmbedder
|
||||
ARCH: ViT-H-14
|
||||
PRETRAINED_MODEL:
|
||||
LAYER: penultimate
|
||||
#
|
||||
LOSS:
|
||||
NAME: ReconstructLoss
|
||||
LOSS_TYPE: l2
|
||||
#
|
||||
SAMPLE_ARGS:
|
||||
SAMPLER: ddim
|
||||
SAMPLE_STEPS: 50
|
||||
SEED: 2023
|
||||
GUIDE_SCALE: 7.5
|
||||
GUIDE_RESCALE:
|
||||
DISCRETIZATION: trailing
|
||||
IMAGE_SIZE: [512, 512]
|
||||
RUN_TRAIN_N: False
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.0064
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
AMSGRAD: False
|
||||
#
|
||||
TRAIN_DATA:
|
||||
NAME: ImageTextPairMSDataset
|
||||
MODE: train
|
||||
MS_DATASET_NAME: style_custom_dataset
|
||||
MS_DATASET_NAMESPACE: damo
|
||||
MS_DATASET_SUBNAME: 3D
|
||||
PROMPT_PREFIX: ""
|
||||
MS_DATASET_SPLIT: train_short
|
||||
MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
|
||||
REPLACE_STYLE: False
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
SAMPLER:
|
||||
NAME: LoopSampler
|
||||
TRANSFORMS:
|
||||
- NAME: LoadImageFromFile
|
||||
RGB_ORDER: RGB
|
||||
BACKEND: pillow
|
||||
- NAME: Resize
|
||||
SIZE: 512
|
||||
INTERPOLATION: bilinear
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: CenterCrop
|
||||
SIZE: 512
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: ImageToTensor
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: Normalize
|
||||
MEAN: [ 0.5, 0.5, 0.5 ]
|
||||
STD: [ 0.5, 0.5, 0.5 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'image' ]
|
||||
BACKEND: torchvision
|
||||
- NAME: Select
|
||||
KEYS: [ 'image', 'prompt' ]
|
||||
META_KEYS: [ 'data_key' ]
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: ImageTextPairMSDataset
|
||||
MODE: eval
|
||||
MS_DATASET_NAME: style_custom_dataset
|
||||
MS_DATASET_NAMESPACE: damo
|
||||
MS_DATASET_SUBNAME: 3D
|
||||
PROMPT_PREFIX: ""
|
||||
MS_REMAP_KEYS: { 'Image': 'Target:FILE' }
|
||||
MS_DATASET_SPLIT: test_short
|
||||
OUTPUT_SIZE: [512, 512]
|
||||
REPLACE_STYLE: False
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 4
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
-
|
||||
NAME: Select
|
||||
KEYS: ['prompt']
|
||||
META_KEYS: ['image_size']
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
-
|
||||
NAME: BackwardHook
|
||||
PRIORITY: 0
|
||||
-
|
||||
NAME: LogHook
|
||||
LOG_INTERVAL: 50
|
||||
-
|
||||
NAME: CheckpointHook
|
||||
INTERVAL: 1000
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
@@ -0,0 +1,223 @@
|
||||
ENV:
|
||||
BACKEND: nccl
|
||||
SOLVER:
|
||||
NAME: LatentDiffusionSolver
|
||||
RESUME_FROM:
|
||||
LOAD_MODEL_ONLY: True
|
||||
USE_FSDP: False
|
||||
SHARDING_STRATEGY:
|
||||
USE_AMP: True
|
||||
DTYPE: float16
|
||||
CHANNELS_LAST: True
|
||||
MAX_STEPS: 2000
|
||||
MAX_EPOCHS: -1
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sd21_512_lora
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
#
|
||||
TUNER:
|
||||
-
|
||||
NAME: SwiftLoRA
|
||||
R: 64
|
||||
LORA_ALPHA: 64
|
||||
LORA_DROPOUT: 0.0
|
||||
BIAS: "none"
|
||||
TARGET_MODULES: model.*(to_q|to_k|to_v|to_out.0|net.0.proj|net.2)$
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusion
|
||||
PARAMETERIZATION: eps
|
||||
TIMESTEPS: 1000
|
||||
MIN_SNR_GAMMA:
|
||||
ZERO_TERMINAL_SNR: False
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-2-1-base@v2-1_512-ema-pruned.safetensors
|
||||
IGNORE_KEYS: [ ]
|
||||
SCALE_FACTOR: 0.18215
|
||||
SIZE_FACTOR: 8
|
||||
DEFAULT_N_PROMPT:
|
||||
SCHEDULE_ARGS:
|
||||
"NAME": "scaled_linear"
|
||||
"BETA_MIN": 0.00085
|
||||
"BETA_MAX": 0.012
|
||||
USE_EMA: False
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: DiffusionUNet
|
||||
IN_CHANNELS: 4
|
||||
OUT_CHANNELS: 4
|
||||
MODEL_CHANNELS: 320
|
||||
NUM_HEADS_CHANNELS: 64
|
||||
NUM_RES_BLOCKS: 2
|
||||
ATTENTION_RESOLUTIONS: [ 4, 2, 1 ]
|
||||
CHANNEL_MULT: [ 1, 2, 4, 4 ]
|
||||
CONV_RESAMPLE: True
|
||||
DIMS: 2
|
||||
USE_CHECKPOINT: False
|
||||
USE_SCALE_SHIFT_NORM: False
|
||||
RESBLOCK_UPDOWN: False
|
||||
USE_SPATIAL_TRANSFORMER: True
|
||||
TRANSFORMER_DEPTH: 1
|
||||
CONTEXT_DIM: 1024
|
||||
DISABLE_MIDDLE_SELF_ATTN: False
|
||||
USE_LINEAR_IN_TRANSFORMER: True
|
||||
PRETRAINED_MODEL:
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKL
|
||||
EMBED_DIM: 4
|
||||
PRETRAINED_MODEL:
|
||||
IGNORE_KEYS: [ ]
|
||||
BATCH_SIZE: 4
|
||||
#
|
||||
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
|
||||
#
|
||||
TOKENIZER:
|
||||
NAME: OpenClipTokenizer
|
||||
LENGTH: 77
|
||||
#
|
||||
COND_STAGE_MODEL:
|
||||
NAME: FrozenOpenCLIPEmbedder
|
||||
ARCH: ViT-H-14
|
||||
PRETRAINED_MODEL:
|
||||
LAYER: penultimate
|
||||
#
|
||||
LOSS:
|
||||
NAME: ReconstructLoss
|
||||
LOSS_TYPE: l2
|
||||
#
|
||||
SAMPLE_ARGS:
|
||||
SAMPLER: ddim
|
||||
SAMPLE_STEPS: 50
|
||||
SEED: 2023
|
||||
GUIDE_SCALE: 7.5
|
||||
GUIDE_RESCALE:
|
||||
DISCRETIZATION: trailing
|
||||
IMAGE_SIZE: [512, 512]
|
||||
RUN_TRAIN_N: False
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.0064
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
AMSGRAD: False
|
||||
#
|
||||
TRAIN_DATA:
|
||||
NAME: ImageTextPairMSDataset
|
||||
MODE: train
|
||||
MS_DATASET_NAME: style_custom_dataset
|
||||
MS_DATASET_NAMESPACE: damo
|
||||
MS_DATASET_SUBNAME: 3D
|
||||
PROMPT_PREFIX: ""
|
||||
MS_DATASET_SPLIT: train_short
|
||||
MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
|
||||
REPLACE_STYLE: False
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
SAMPLER:
|
||||
NAME: LoopSampler
|
||||
TRANSFORMS:
|
||||
- NAME: LoadImageFromFile
|
||||
RGB_ORDER: RGB
|
||||
BACKEND: pillow
|
||||
- NAME: Resize
|
||||
SIZE: 512
|
||||
INTERPOLATION: bilinear
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: CenterCrop
|
||||
SIZE: 512
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: ImageToTensor
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: Normalize
|
||||
MEAN: [ 0.5, 0.5, 0.5 ]
|
||||
STD: [ 0.5, 0.5, 0.5 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'image' ]
|
||||
BACKEND: torchvision
|
||||
- NAME: Select
|
||||
KEYS: [ 'image', 'prompt' ]
|
||||
META_KEYS: [ 'data_key' ]
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: ImageTextPairMSDataset
|
||||
MODE: eval
|
||||
MS_DATASET_NAME: style_custom_dataset
|
||||
MS_DATASET_NAMESPACE: damo
|
||||
MS_DATASET_SUBNAME: 3D
|
||||
PROMPT_PREFIX: ""
|
||||
MS_REMAP_KEYS: { 'Image': 'Target:FILE' }
|
||||
MS_DATASET_SPLIT: test_short
|
||||
OUTPUT_SIZE: [512, 512]
|
||||
REPLACE_STYLE: False
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 4
|
||||
NUM_WORKERS: 4
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
#
|
||||
TRANSFORMS:
|
||||
-
|
||||
NAME: Select
|
||||
KEYS: ['prompt']
|
||||
META_KEYS: ['image_size']
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
-
|
||||
NAME: BackwardHook
|
||||
PRIORITY: 0
|
||||
-
|
||||
NAME: LogHook
|
||||
LOG_INTERVAL: 50
|
||||
-
|
||||
NAME: CheckpointHook
|
||||
INTERVAL: 1000
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
@@ -113,7 +113,7 @@ SOLVER:
|
||||
SAMPLE_STEPS: 50
|
||||
SEED: 2023
|
||||
GUIDE_SCALE: 7.5
|
||||
GUIDE_RESCALE:
|
||||
GUIDE_RESCALE: 0.5
|
||||
DISCRETIZATION: trailing
|
||||
IMAGE_SIZE: [768, 768]
|
||||
RUN_TRAIN_N: False
|
||||
|
||||
@@ -122,7 +122,7 @@ SOLVER:
|
||||
SAMPLE_STEPS: 50
|
||||
SEED: 2023
|
||||
GUIDE_SCALE: 7.5
|
||||
GUIDE_RESCALE:
|
||||
GUIDE_RESCALE: 0.5
|
||||
DISCRETIZATION: trailing
|
||||
IMAGE_SIZE: [768, 768]
|
||||
RUN_TRAIN_N: False
|
||||
|
||||
@@ -31,7 +31,7 @@ SOLVER:
|
||||
PARAMETERIZATION: v
|
||||
TIMESTEPS: 1000
|
||||
MIN_SNR_GAMMA:
|
||||
ZERO_TERMINAL_SNR: True
|
||||
ZERO_TERMINAL_SNR: False
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-2-1@v2-1_768-ema-pruned.safetensors
|
||||
IGNORE_KEYS: [ ]
|
||||
SCALE_FACTOR: 0.18215
|
||||
|
||||
@@ -31,7 +31,7 @@ SOLVER:
|
||||
PARAMETERIZATION: v
|
||||
TIMESTEPS: 1000
|
||||
MIN_SNR_GAMMA:
|
||||
ZERO_TERMINAL_SNR: True
|
||||
ZERO_TERMINAL_SNR: False
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-2-1@v2-1_768-ema-pruned.safetensors
|
||||
IGNORE_KEYS: [ ]
|
||||
SCALE_FACTOR: 0.18215
|
||||
|
||||
@@ -0,0 +1,378 @@
|
||||
ENV:
|
||||
BACKEND: nccl
|
||||
SOLVER:
|
||||
NAME: LatentDiffusionSolver
|
||||
RESUME_FROM:
|
||||
LOAD_MODEL_ONLY: True
|
||||
USE_FSDP: False
|
||||
SHARDING_STRATEGY:
|
||||
USE_AMP: True
|
||||
DTYPE: float16
|
||||
CHANNELS_LAST: True
|
||||
MAX_STEPS: 200
|
||||
MAX_EPOCHS: -1
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: 100
|
||||
#
|
||||
WORK_DIR: ./cache/save_data/sdxl_1024_sce_ctr_canny
|
||||
LOG_FILE: std_log.txt
|
||||
#
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
#
|
||||
FREEZE:
|
||||
FREEZE_PART: [ "first_stage_model", "cond_stage_model", "model" ]
|
||||
TRAIN_PART: [ "control_blocks" ]
|
||||
#
|
||||
MODEL:
|
||||
NAME: LatentDiffusionXLSCEControl
|
||||
PARAMETERIZATION: eps
|
||||
TIMESTEPS: 1000
|
||||
MIN_SNR_GAMMA:
|
||||
ZERO_TERMINAL_SNR: False
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-xl-base-1.0@sd_xl_base_1.0.safetensors
|
||||
IGNORE_KEYS: [ ]
|
||||
SCALE_FACTOR: 0.13025
|
||||
SIZE_FACTOR: 8
|
||||
DEFAULT_N_PROMPT:
|
||||
SCHEDULE_ARGS:
|
||||
"NAME": "scaled_linear"
|
||||
"BETA_MIN": 0.00085
|
||||
"BETA_MAX": 0.0120
|
||||
USE_EMA: False
|
||||
LOAD_REFINER: False
|
||||
#
|
||||
DIFFUSION_MODEL:
|
||||
NAME: DiffusionUNetXL
|
||||
PRETRAINED_MODEL:
|
||||
IN_CHANNELS: 4
|
||||
OUT_CHANNELS: 4
|
||||
NUM_RES_BLOCKS: 2
|
||||
MODEL_CHANNELS: 320
|
||||
ATTENTION_RESOLUTIONS: [ 4, 2 ]
|
||||
DROPOUT: 0
|
||||
CHANNEL_MULT: [ 1, 2, 4 ]
|
||||
CONV_RESAMPLE: True
|
||||
DIMS: 2
|
||||
NUM_CLASSES: sequential
|
||||
USE_CHECKPOINT: False
|
||||
NUM_HEADS: -1
|
||||
NUM_HEADS_CHANNELS: 64
|
||||
USE_SCALE_SHIFT_NORM: False
|
||||
RESBLOCK_UPDOWN: False
|
||||
USE_NEW_ATTENTION_ORDER: True
|
||||
USE_SPATIAL_TRANSFORMER: True
|
||||
TRANSFORMER_DEPTH: [ 1, 2, 10 ]
|
||||
CONTEXT_DIM: 2048
|
||||
DISABLE_MIDDLE_SELF_ATTN: False
|
||||
USE_LINEAR_IN_TRANSFORMER: True
|
||||
ADM_IN_CHANNELS: 2816
|
||||
USE_SENTENCE_EMB: False
|
||||
USE_WORD_MAPPING: False
|
||||
#
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKL
|
||||
EMBED_DIM: 4
|
||||
PRETRAINED_MODEL:
|
||||
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:
|
||||
NAME: GeneralConditioner
|
||||
PRETRAINED_MODEL:
|
||||
EMBEDDERS:
|
||||
-
|
||||
NAME: FrozenCLIPEmbedder
|
||||
PRETRAINED_MODEL: ms://AI-ModelScope/clip-vit-large-patch14
|
||||
TOKENIZER_PATH: ms://AI-ModelScope/clip-vit-large-patch14
|
||||
MAX_LENGTH: 77
|
||||
FREEZE: True
|
||||
LAYER: hidden
|
||||
LAYER_IDX: 11
|
||||
USE_FINAL_LAYER_NORM: False
|
||||
IS_TRAINABLE: False
|
||||
UCG_RATE: 0.0
|
||||
INPUT_KEYS: [ "prompt" ]
|
||||
LEGACY_UCG_VALUE:
|
||||
-
|
||||
NAME: FrozenOpenCLIPEmbedder2
|
||||
ARCH: ViT-bigG-14
|
||||
PRETRAINED_MODEL:
|
||||
MAX_LENGTH: 77
|
||||
FREEZE: True
|
||||
ALWAYS_RETURN_POOLED: True
|
||||
LEGACY: False
|
||||
LAYER: penultimate
|
||||
IS_TRAINABLE: False
|
||||
UCG_RATE: 0.0
|
||||
INPUT_KEYS: [ "prompt" ]
|
||||
LEGACY_UCG_VALUE:
|
||||
-
|
||||
NAME: ConcatTimestepEmbedderND
|
||||
OUT_DIM: 256
|
||||
IS_TRAINABLE: False
|
||||
UCG_RATE: 0.0
|
||||
INPUT_KEYS: [ "original_size_as_tuple" ]
|
||||
LEGACY_UCG_VALUE:
|
||||
-
|
||||
NAME: ConcatTimestepEmbedderND
|
||||
OUT_DIM: 256
|
||||
IS_TRAINABLE: False
|
||||
UCG_RATE: 0.0
|
||||
INPUT_KEYS: [ "crop_coords_top_left" ]
|
||||
LEGACY_UCG_VALUE:
|
||||
-
|
||||
NAME: ConcatTimestepEmbedderND
|
||||
OUT_DIM: 256
|
||||
IS_TRAINABLE: False
|
||||
UCG_RATE: 0.0
|
||||
INPUT_KEYS: [ "target_size_as_tuple" ]
|
||||
LEGACY_UCG_VALUE:
|
||||
#
|
||||
REFINER_MODEL:
|
||||
NAME: DiffusionUNetXL
|
||||
PRETRAINED_MODEL:
|
||||
IN_CHANNELS: 4
|
||||
OUT_CHANNELS: 4
|
||||
NUM_RES_BLOCKS: 2
|
||||
MODEL_CHANNELS: 384
|
||||
ATTENTION_RESOLUTIONS: [ 4, 2 ]
|
||||
DROPOUT: 0
|
||||
CHANNEL_MULT: [ 1, 2, 4, 4 ]
|
||||
CONV_RESAMPLE: True
|
||||
DIMS: 2
|
||||
NUM_CLASSES: sequential
|
||||
USE_CHECKPOINT: False
|
||||
NUM_HEADS: -1
|
||||
NUM_HEADS_CHANNELS: 64
|
||||
USE_SCALE_SHIFT_NORM: False
|
||||
RESBLOCK_UPDOWN: False
|
||||
USE_NEW_ATTENTION_ORDER: True
|
||||
USE_SPATIAL_TRANSFORMER: True
|
||||
TRANSFORMER_DEPTH: 4
|
||||
CONTEXT_DIM: [ 1280, 1280, 1280, 1280 ]
|
||||
DISABLE_MIDDLE_SELF_ATTN: False
|
||||
USE_LINEAR_IN_TRANSFORMER: True
|
||||
ADM_IN_CHANNELS: 2560
|
||||
USE_SENTENCE_EMB: False
|
||||
USE_WORD_MAPPING: False
|
||||
REFINER_COND_MODEL:
|
||||
NAME: GeneralConditioner
|
||||
PRETRAINED_MODEL:
|
||||
EMBEDDERS:
|
||||
-
|
||||
NAME: FrozenOpenCLIPEmbedder2
|
||||
ARCH: ViT-bigG-14
|
||||
PRETRAINED_MODEL:
|
||||
MAX_LENGTH: 77
|
||||
FREEZE: True
|
||||
ALWAYS_RETURN_POOLED: True
|
||||
LEGACY: False
|
||||
LAYER: penultimate
|
||||
IS_TRAINABLE: False
|
||||
UCG_RATE: 0.0
|
||||
INPUT_KEYS: [ "prompt" ]
|
||||
LEGACY_UCG_VALUE:
|
||||
-
|
||||
NAME: ConcatTimestepEmbedderND
|
||||
OUT_DIM: 256
|
||||
IS_TRAINABLE: False
|
||||
UCG_RATE: 0.0
|
||||
INPUT_KEYS: [ "original_size_as_tuple" ]
|
||||
LEGACY_UCG_VALUE:
|
||||
-
|
||||
NAME: ConcatTimestepEmbedderND
|
||||
OUT_DIM: 256
|
||||
IS_TRAINABLE: False
|
||||
UCG_RATE: 0.0
|
||||
INPUT_KEYS: [ "crop_coords_top_left" ]
|
||||
LEGACY_UCG_VALUE:
|
||||
-
|
||||
NAME: ConcatTimestepEmbedderND
|
||||
OUT_DIM: 256
|
||||
IS_TRAINABLE: False
|
||||
UCG_RATE: 0.0
|
||||
INPUT_KEYS: [ "aesthetic_score" ]
|
||||
LEGACY_UCG_VALUE:
|
||||
#
|
||||
LOSS:
|
||||
NAME: ReconstructLoss
|
||||
LOSS_TYPE: l2
|
||||
#
|
||||
CONTROL_MODEL:
|
||||
NAME: CSCTuners
|
||||
PRE_HINT_IN_CHANNELS: 3
|
||||
PRE_HINT_OUT_CHANNELS: 320
|
||||
DENSE_HINT_KERNAL: 3
|
||||
PRE_HINT_DIM_RATIO: 2.0
|
||||
SCALE: 1.0
|
||||
SC_TUNER_CFG:
|
||||
NAME: SCTuner
|
||||
TUNER_NAME: SCEAdapter
|
||||
DOWN_RATIO: 1.0
|
||||
CONTROL_ANNO:
|
||||
NAME: CannyAnnotator
|
||||
#
|
||||
SAMPLE_ARGS:
|
||||
SAMPLER: ddim
|
||||
SAMPLE_STEPS: 50
|
||||
SEED: 2023
|
||||
GUIDE_SCALE: 7.5
|
||||
GUIDE_RESCALE: 0.5
|
||||
DISCRETIZATION: trailing
|
||||
IMAGE_SIZE: [1024, 1024]
|
||||
RUN_TRAIN_N: False
|
||||
#
|
||||
OPTIMIZER:
|
||||
NAME: AdamW
|
||||
LEARNING_RATE: 0.064
|
||||
BETAS: [ 0.9, 0.999 ]
|
||||
EPS: 1e-8
|
||||
WEIGHT_DECAY: 1e-2
|
||||
AMSGRAD: False
|
||||
#
|
||||
TRAIN_DATA:
|
||||
NAME: ImageTextPairMSDataset
|
||||
MODE: train
|
||||
MS_DATASET_NAME: style_custom_dataset
|
||||
MS_DATASET_NAMESPACE: damo
|
||||
MS_DATASET_SUBNAME: 3D
|
||||
PROMPT_PREFIX: ""
|
||||
MS_DATASET_SPLIT: train_short
|
||||
MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
|
||||
REPLACE_STYLE: False
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
SAMPLER:
|
||||
NAME: LoopSampler
|
||||
TRANSFORMS:
|
||||
- NAME: LoadImageFromFile
|
||||
RGB_ORDER: RGB
|
||||
BACKEND: pillow
|
||||
- NAME: FlexibleResize
|
||||
SIZE: 1024
|
||||
INTERPOLATION: bilinear
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: FlexibleCropXL
|
||||
SIZE: 1024
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: ToNumpy
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'image_preprocess' ]
|
||||
- NAME: ImageToTensor
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: Normalize
|
||||
MEAN: [ 0.5, 0.5, 0.5 ]
|
||||
STD: [ 0.5, 0.5, 0.5 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: torchvision
|
||||
- NAME: Rename
|
||||
INPUT_KEY: [ 'img', 'image_preprocess', 'img_original_size_as_tuple', 'img_target_size_as_tuple', 'img_crop_coords_top_left' ]
|
||||
OUTPUT_KEY: [ 'image', 'image_preprocess', 'original_size_as_tuple', 'target_size_as_tuple', 'crop_coords_top_left' ]
|
||||
- NAME: Select
|
||||
KEYS: [ 'image', 'prompt', 'image_preprocess', 'original_size_as_tuple', 'target_size_as_tuple', 'crop_coords_top_left' ]
|
||||
META_KEYS: [ 'data_key' ]
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: ImageTextPairMSDataset
|
||||
MODE: eval
|
||||
MS_DATASET_NAME: style_custom_dataset
|
||||
MS_DATASET_NAMESPACE: damo
|
||||
MS_DATASET_SUBNAME: 3D
|
||||
PROMPT_PREFIX: ""
|
||||
MS_DATASET_SPLIT: train_short
|
||||
MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' }
|
||||
REPLACE_STYLE: False
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 10
|
||||
NUM_WORKERS: 4
|
||||
TRANSFORMS:
|
||||
- NAME: LoadImageFromFile
|
||||
RGB_ORDER: RGB
|
||||
BACKEND: pillow
|
||||
- NAME: Resize
|
||||
SIZE: 1024
|
||||
INTERPOLATION: bilinear
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: CenterCrop
|
||||
SIZE: 1024
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: ToNumpy
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'image_preprocess' ]
|
||||
- NAME: ImageToTensor
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: pillow
|
||||
- NAME: Normalize
|
||||
MEAN: [ 0.5, 0.5, 0.5 ]
|
||||
STD: [ 0.5, 0.5, 0.5 ]
|
||||
INPUT_KEY: [ 'img' ]
|
||||
OUTPUT_KEY: [ 'img' ]
|
||||
BACKEND: torchvision
|
||||
- NAME: Rename
|
||||
INPUT_KEY: [ 'img', 'image_preprocess' ]
|
||||
OUTPUT_KEY: [ 'image', 'image_preprocess' ]
|
||||
- NAME: Select
|
||||
KEYS: [ 'image', 'prompt', 'image_preprocess' ]
|
||||
META_KEYS: [ 'data_key' ]
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
-
|
||||
NAME: BackwardHook
|
||||
PRIORITY: 0
|
||||
-
|
||||
NAME: LogHook
|
||||
LOG_INTERVAL: 50
|
||||
-
|
||||
NAME: CheckpointHook
|
||||
INTERVAL: 100
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
@@ -231,8 +231,9 @@ SOLVER:
|
||||
CONTROL_MODEL:
|
||||
NAME: CSCTuners
|
||||
PRE_HINT_IN_CHANNELS: 3
|
||||
PRE_HINT_OUT_CHANNELS: 256
|
||||
PRE_HINT_OUT_CHANNELS: 320
|
||||
DENSE_HINT_KERNAL: 3
|
||||
PRE_HINT_DIM_RATIO: 2.0
|
||||
SCALE: 1.0
|
||||
SC_TUNER_CFG:
|
||||
NAME: SCTuner
|
||||
|
||||
@@ -231,8 +231,9 @@ SOLVER:
|
||||
CONTROL_MODEL:
|
||||
NAME: CSCTuners
|
||||
PRE_HINT_IN_CHANNELS: 3
|
||||
PRE_HINT_OUT_CHANNELS: 256
|
||||
PRE_HINT_OUT_CHANNELS: 320
|
||||
DENSE_HINT_KERNAL: 3
|
||||
PRE_HINT_DIM_RATIO: 2.0
|
||||
SCALE: 1.0
|
||||
SC_TUNER_CFG:
|
||||
NAME: SCTuner
|
||||
|
||||
@@ -231,8 +231,9 @@ SOLVER:
|
||||
CONTROL_MODEL:
|
||||
NAME: CSCTuners
|
||||
PRE_HINT_IN_CHANNELS: 3
|
||||
PRE_HINT_OUT_CHANNELS: 256
|
||||
PRE_HINT_OUT_CHANNELS: 320
|
||||
DENSE_HINT_KERNAL: 3
|
||||
PRE_HINT_DIM_RATIO: 2.0
|
||||
SCALE: 1.0
|
||||
SC_TUNER_CFG:
|
||||
NAME: SCTuner
|
||||
|
||||
@@ -1,19 +1,63 @@
|
||||
CONTROLLERS:
|
||||
# SD2.1
|
||||
- NAME: canny
|
||||
NAME_ZH:
|
||||
DESCRIPTION:
|
||||
BASE_MODEL: SD2.1
|
||||
TYPE: Canny
|
||||
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD2.1/canny_control/0_SwiftSCETuning
|
||||
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD2.1/canny_control/
|
||||
- NAME: openpose
|
||||
NAME_ZH:
|
||||
DESCRIPTION:
|
||||
BASE_MODEL: SD2.1
|
||||
TYPE: Openpose
|
||||
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD2.1/pose_control/0_SwiftSCETuning
|
||||
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD2.1/pose_control/
|
||||
- NAME: color
|
||||
NAME_ZH:
|
||||
DESCRIPTION:
|
||||
BASE_MODEL: SD2.1
|
||||
TYPE: Color
|
||||
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD2.1/color_control/0_SwiftSCETuning
|
||||
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD2.1/color_control/
|
||||
- NAME: hed
|
||||
NAME_ZH:
|
||||
DESCRIPTION:
|
||||
BASE_MODEL: SD2.1
|
||||
TYPE: Hed
|
||||
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD2.1/hed_control
|
||||
- NAME: depth
|
||||
NAME_ZH:
|
||||
DESCRIPTION:
|
||||
BASE_MODEL: SD2.1
|
||||
TYPE: Midas
|
||||
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD2.1/depth_control
|
||||
# SD_XL1.0
|
||||
- NAME: canny
|
||||
NAME_ZH:
|
||||
DESCRIPTION:
|
||||
BASE_MODEL: SD_XL1.0
|
||||
TYPE: Canny
|
||||
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD_XL1.0/canny_control
|
||||
- NAME: color
|
||||
NAME_ZH:
|
||||
DESCRIPTION:
|
||||
BASE_MODEL: SD_XL1.0
|
||||
TYPE: Color
|
||||
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD_XL1.0/color_control
|
||||
- NAME: depth
|
||||
NAME_ZH:
|
||||
DESCRIPTION:
|
||||
BASE_MODEL: SD_XL1.0
|
||||
TYPE: Midas
|
||||
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD_XL1.0/depth_control
|
||||
- NAME: hed
|
||||
NAME_ZH:
|
||||
DESCRIPTION:
|
||||
BASE_MODEL: SD_XL1.0
|
||||
TYPE: Hed
|
||||
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD_XL1.0/hed_control
|
||||
- NAME: openpose
|
||||
NAME_ZH:
|
||||
DESCRIPTION:
|
||||
BASE_MODEL: SD_XL1.0
|
||||
TYPE: Openpose
|
||||
MODEL_PATH: ms://damo/scepter_scedit@controllable_model/SD_XL1.0/pose_control
|
||||
|
||||
@@ -1,4 +1,76 @@
|
||||
TUNERS:
|
||||
- NAME: Azure-Dragon
|
||||
NAME_ZH: 青龙
|
||||
SOURCE: wanx
|
||||
DESCRIPTION: None
|
||||
BASE_MODEL: SD_XL1.0
|
||||
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/azure_dragon/
|
||||
IMAGE_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/azure_dragon/xl_azure_dragon.png
|
||||
TUNER_TYPE: SwiftSCE
|
||||
PROMPT_EXAMPLE: Azure Dragon, 8K, high quality,Ultra High Detail.One of the Four Divine Creatures in Charge of Water.
|
||||
- NAME: Gold-Dragon
|
||||
NAME_ZH: 金龙
|
||||
SOURCE: wanx
|
||||
DESCRIPTION: None
|
||||
BASE_MODEL: SD_XL1.0
|
||||
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/gold_dragon/
|
||||
IMAGE_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/gold_dragon/xl_gold_dragon.png
|
||||
TUNER_TYPE: SwiftSCE
|
||||
PROMPT_EXAMPLE: Chinese Gold Dragon in the clouds. Translucent Texture. Zbrush. Fuzzy Art. Exquisite Craftsmanship. 3D. 8K. Ultra High Detail
|
||||
- NAME: SpringFestival-Dragon
|
||||
NAME_ZH: 春节龙
|
||||
SOURCE: wanx
|
||||
DESCRIPTION: None
|
||||
BASE_MODEL: SD_XL1.0
|
||||
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/spring_festival_dragon/
|
||||
IMAGE_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/spring_festival_dragon/xl_spring_festival_dragon.png
|
||||
TUNER_TYPE: SwiftSCE
|
||||
PROMPT_EXAMPLE: Chinese dragon. Spring Festival.Festive.Street.Lanterns.32K.High quality.expressive, dramatic, dreamlike and mysterious, Surrealism
|
||||
- NAME: Red-Dragon
|
||||
NAME_ZH: 红龙
|
||||
SOURCE: wanx
|
||||
DESCRIPTION: None
|
||||
BASE_MODEL: SD_XL1.0
|
||||
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/red_dragon/
|
||||
IMAGE_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/red_dragon/xl_red_dragon.png
|
||||
TUNER_TYPE: SwiftSCE
|
||||
PROMPT_EXAMPLE: Traditional Red Dragon of China. Low Water Level. Studio Ghibli Style. Mural Illustration. White Background. High Detail
|
||||
- NAME: ChinesePunk-Dragon
|
||||
NAME_ZH: 中国朋克龙
|
||||
SOURCE: wanx
|
||||
DESCRIPTION: None
|
||||
BASE_MODEL: SD_XL1.0
|
||||
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/chinese_punk_dragon/
|
||||
IMAGE_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/chinese_punk_dragon/xl_chinese_punk_dragon.png
|
||||
TUNER_TYPE: SwiftSCE
|
||||
PROMPT_EXAMPLE: uhd Image,Dragon,Chinese Dragon, Dunhuang Mural Style, Traditional Maritime Art Style
|
||||
- NAME: Cute-Dragon
|
||||
NAME_ZH: 喜庆龙
|
||||
SOURCE: wanx
|
||||
DESCRIPTION: None
|
||||
BASE_MODEL: SD_XL1.0
|
||||
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/cute_dragon/
|
||||
IMAGE_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/cute_dragon/xl_kawaii_dragon.png
|
||||
TUNER_TYPE: SwiftSCE
|
||||
PROMPT_EXAMPLE: China Kawaii Dragon. Contest Winner. Minimalist Illustration. White Background. Flat Style. Digital Painting Style. Red. 32k uhd. Fun Comics. Fuzzy Art. Bold. Comic-Inspired Characters
|
||||
- NAME: Dragon-Baby
|
||||
NAME_ZH: 龙宝宝
|
||||
SOURCE: wanx
|
||||
DESCRIPTION: None
|
||||
BASE_MODEL: SD_XL1.0
|
||||
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/baby_dragon/
|
||||
IMAGE_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/baby_dragon/xl_baby_dragon.png
|
||||
TUNER_TYPE: SwiftSCE
|
||||
PROMPT_EXAMPLE: Warm Colors, Soft,Chinese Dragon Baby, Felt Style,Dragon Baby, Best Quality, 3D Doll, Macaron Tones, Glittering Big Eyes, Winter,Dragon
|
||||
- NAME: Sloppy-Dragon
|
||||
NAME_ZH: 潦草龙
|
||||
SOURCE: wanx
|
||||
DESCRIPTION: None
|
||||
BASE_MODEL: SD_XL1.0
|
||||
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/sloppy_dragon/
|
||||
IMAGE_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/sloppy_dragon/xl_sloppy_dragon.png
|
||||
TUNER_TYPE: SwiftSCE
|
||||
PROMPT_EXAMPLE: Messy Chinese Dragon,Cute, Wu Guanzhong, Rough
|
||||
-
|
||||
NAME: Caricature
|
||||
NAME_ZH: 夸张漫画
|
||||
|
||||
@@ -38,15 +38,15 @@ GUIDE_INFO:
|
||||
width: 100%;
|
||||
}
|
||||
.video-wrapper {
|
||||
width: 75%;
|
||||
width: 75%;
|
||||
}
|
||||
video {
|
||||
width: 100%;
|
||||
display: block;
|
||||
width: 100%;
|
||||
display: block;
|
||||
}
|
||||
.description {
|
||||
text-align: center;
|
||||
margin-top: 10px;
|
||||
text-align: center;
|
||||
margin-top: 10px;
|
||||
font-size: 0.8em;
|
||||
}
|
||||
</style>
|
||||
@@ -70,15 +70,15 @@ GUIDE_INFO:
|
||||
width: 100%;
|
||||
}
|
||||
.video-wrapper {
|
||||
width: 75%;
|
||||
width: 75%;
|
||||
}
|
||||
video {
|
||||
width: 100%;
|
||||
display: block;
|
||||
width: 100%;
|
||||
display: block;
|
||||
}
|
||||
.description {
|
||||
text-align: center;
|
||||
margin-top: 10px;
|
||||
text-align: center;
|
||||
margin-top: 10px;
|
||||
font-size: 0.8em;
|
||||
}
|
||||
</style>
|
||||
@@ -92,4 +92,4 @@ GUIDE_INFO:
|
||||
<div class="description">Train & Inference Video</div>
|
||||
</div>
|
||||
</div>
|
||||
</body>
|
||||
</body>
|
||||
|
||||
@@ -79,11 +79,10 @@ EXTENSION_PARAS:
|
||||
MANTRA_BOOK: scepter/methods/studio/extensions/mantra_book/mantra_book.yaml
|
||||
OFFICIAL_TUNERS: scepter/methods/studio/extensions/tuners/official_tuners.yaml
|
||||
OFFICIAL_CONTROLLERS: scepter/methods/studio/extensions/controllers/official_controllers.yaml
|
||||
TUNER_MANAGER: scepter/methods/studio/tuner_manager/tuner_manager.yaml
|
||||
CONTROLABLE_ANNOTATORS:
|
||||
-
|
||||
NAME: "CannyAnnotator"
|
||||
LOW_THRESHOLD: 100
|
||||
HIGH_THRESHOLD: 200
|
||||
TYPE: Canny
|
||||
IS_DEFAULT: True
|
||||
-
|
||||
@@ -100,8 +99,6 @@ CONTROLABLE_ANNOTATORS:
|
||||
-
|
||||
NAME: "MidasDetector"
|
||||
PRETRAINED_MODEL: "ms://damo/scepter_scedit@annotator/ckpts/dpt_hybrid-midas-501f0c75.pt"
|
||||
A: 6.2
|
||||
BG_TH: 0.1
|
||||
TYPE: Midas
|
||||
IS_DEFAULT: False
|
||||
-
|
||||
@@ -109,9 +106,6 @@ CONTROLABLE_ANNOTATORS:
|
||||
TYPE: Color
|
||||
IS_DEFAULT: False
|
||||
-
|
||||
NAME: "MLSDdetector"
|
||||
PRETRAINED_MODEL: "ms://damo/scepter_scedit@annotator/ckpts/mlsd_large_512_fp32.pth"
|
||||
THR_V: 0.1
|
||||
THR_D: 0.1
|
||||
TYPE: MLSD
|
||||
NAME: "InvertAnnotator"
|
||||
TYPE: Invert-Preprocess
|
||||
IS_DEFAULT: False
|
||||
|
||||
@@ -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-code-container">
|
||||
<img src="https://modelscope.cn/api/v1/models/damo/scepter/repo?Revision=master&FilePath=assets/scepter_studio/ms_scepter_studio_qr.png" alt="ms_scepter_studio_qr">
|
||||
<div class="caption">Modelscope Studio</div>
|
||||
<div class="caption"><a href="https://www.modelscope.cn/studios/iic/scepter_studio">Modelscope Studio</a></div>
|
||||
</div>
|
||||
<div class="qr-code-container">
|
||||
<img src="https://modelscope.cn/api/v1/models/damo/scepter/repo?Revision=master&FilePath=assets/scepter_studio/scepter_github_qr.png" alt="scepter_github_qr">
|
||||
<div class="caption">Github</div>
|
||||
<div class="caption"><a href="https://github.com/modelscope/scepter">Github</a></div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
@@ -79,6 +79,10 @@ INTERFACE:
|
||||
NAME_EN: Train
|
||||
IFID: self_train
|
||||
CONFIG: scepter/methods/studio/self_train/self_train.yaml
|
||||
- NAME: 模型管理
|
||||
NAME_EN: Tuner Management
|
||||
IFID: tuner_manager
|
||||
CONFIG: scepter/methods/studio/tuner_manager/tuner_manager.yaml
|
||||
- NAME: 推理
|
||||
NAME_EN: Inference
|
||||
IFID: inference
|
||||
|
||||
@@ -11,13 +11,13 @@ META:
|
||||
DEFAULT_SAMPLER: "dpmpp_2s_ancestral"
|
||||
DEFAULT_SAMPLE_STEPS: 40
|
||||
INFERENCE_N_PROMPT: ""
|
||||
RESOLUTION: 1024
|
||||
RESOLUTION: [1024, 1024]
|
||||
PARAS:
|
||||
-
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 1024
|
||||
RESOLUTION: [1024, 1024]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -29,7 +29,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 1024
|
||||
RESOLUTION: [1024, 1024]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -41,7 +41,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 1024
|
||||
RESOLUTION: [1024, 1024]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 200
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -53,7 +53,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 1024
|
||||
RESOLUTION: [1024, 1024]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 200
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -65,7 +65,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 1024
|
||||
RESOLUTION: [1024, 1024]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -146,10 +146,12 @@ SOLVER:
|
||||
MAX_EPOCHS: -1
|
||||
# NUM_FOLDS DESCRIPTION: Num folds for training. TYPE: int default: 1
|
||||
NUM_FOLDS: 1
|
||||
#
|
||||
EVAL_INTERVAL: -1
|
||||
# WORK_DIR DESCRIPTION: Save dir of the training log or model. TYPE: str default: ''
|
||||
WORK_DIR:
|
||||
# LOG_FILE DESCRIPTION: Save log path. TYPE: str default: ''
|
||||
LOG_FILE: stg_log.txt
|
||||
LOG_FILE: std_log.txt
|
||||
FILE_SYSTEM:
|
||||
NAME: "ModelscopeFs"
|
||||
TEMP_DIR: "./cache/data"
|
||||
@@ -586,15 +588,44 @@ SOLVER:
|
||||
INPUT_KEY: [ 'img', 'img_original_size_as_tuple', 'img_target_size_as_tuple', 'img_crop_coords_top_left']
|
||||
OUTPUT_KEY: [ 'image', 'original_size_as_tuple', 'target_size_as_tuple', 'crop_coords_top_left' ]
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: Text2ImageDataset
|
||||
MODE: eval
|
||||
PROMPT_FILE:
|
||||
PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ]
|
||||
IMAGE_SIZE: [ 1024, 1024 ]
|
||||
FIELDS: [ "prompt" ]
|
||||
DELIMITER: '#;#'
|
||||
PROMPT_PREFIX: ''
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'index', 'prompt' ]
|
||||
META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
- NAME: BackwardHook
|
||||
# GRADIENT_CLIP: 1.0
|
||||
-
|
||||
NAME: BackwardHook
|
||||
PRIORITY: 0
|
||||
- NAME: LogHook
|
||||
LOG_INTERVAL: 50
|
||||
-
|
||||
NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
SHOW_GPU_MEM: True
|
||||
-
|
||||
NAME: TensorboardLogHook
|
||||
-
|
||||
NAME: CheckpointHook
|
||||
SAVE_LAST: True
|
||||
INTERVAL: 10000
|
||||
PRIORITY: 200
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
SAVE_PROBE_PREFIX: 'image'
|
||||
|
||||
@@ -8,3 +8,13 @@ SAMPLERS:
|
||||
NAME: 'dpmpp_2m_sde'
|
||||
-
|
||||
NAME: 'dpmpp_2s_ancestral'
|
||||
TRAIN_PARAS:
|
||||
RESOLUTIONS:
|
||||
VALUES: [[256, 256], [320, 180], [180, 320],
|
||||
[512, 512], [640, 360], [360, 640],
|
||||
[768, 768], [960, 540], [540, 960],
|
||||
[1024, 1024], [1280, 720], [720, 1280]]
|
||||
DEFAULT: [1024, 1024]
|
||||
EVAL_PROMPTS:
|
||||
- a boy wearing a jacket
|
||||
- a dog running on the lawn
|
||||
|
||||
@@ -10,13 +10,13 @@ META:
|
||||
DEFAULT_SAMPLER: "ddim"
|
||||
DEFAULT_SAMPLE_STEPS: 40
|
||||
INFERENCE_N_PROMPT: ""
|
||||
RESOLUTION: 512
|
||||
RESOLUTION: [512, 512]
|
||||
PARAS:
|
||||
-
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 512
|
||||
RESOLUTION: [512, 512]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -28,7 +28,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 512
|
||||
RESOLUTION: [512, 512]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -40,7 +40,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 512
|
||||
RESOLUTION: [512, 512]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 200
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -53,7 +53,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 512
|
||||
RESOLUTION: [512, 512]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 200
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -66,7 +66,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 4
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 512
|
||||
RESOLUTION: [512, 512]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -135,6 +135,7 @@ SOLVER:
|
||||
MAX_EPOCHS: -1
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: -1
|
||||
#
|
||||
WORK_DIR:
|
||||
LOG_FILE: std_log.txt
|
||||
@@ -298,15 +299,44 @@ SOLVER:
|
||||
KEYS: [ 'image', 'prompt' ]
|
||||
META_KEYS: [ 'data_key' ]
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: Text2ImageDataset
|
||||
MODE: eval
|
||||
PROMPT_FILE:
|
||||
PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ]
|
||||
IMAGE_SIZE: [ 512, 512 ]
|
||||
FIELDS: [ "prompt" ]
|
||||
DELIMITER: '#;#'
|
||||
PROMPT_PREFIX: ''
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'index', 'prompt' ]
|
||||
META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
- NAME: BackwardHook
|
||||
# GRADIENT_CLIP: 1.0
|
||||
-
|
||||
NAME: BackwardHook
|
||||
PRIORITY: 0
|
||||
- NAME: LogHook
|
||||
LOG_INTERVAL: 50
|
||||
-
|
||||
NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
SHOW_GPU_MEM: True
|
||||
-
|
||||
NAME: TensorboardLogHook
|
||||
-
|
||||
NAME: CheckpointHook
|
||||
SAVE_LAST: True
|
||||
INTERVAL: 10000
|
||||
PRIORITY: 200
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
SAVE_PROBE_PREFIX: 'image'
|
||||
|
||||
@@ -10,13 +10,13 @@ META:
|
||||
DEFAULT_SAMPLER: "ddim"
|
||||
DEFAULT_SAMPLE_STEPS: 40
|
||||
INFERENCE_N_PROMPT: ""
|
||||
RESOLUTION: 768
|
||||
RESOLUTION: [768, 768]
|
||||
PARAS:
|
||||
-
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 768
|
||||
RESOLUTION: [768, 768]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -28,7 +28,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 768
|
||||
RESOLUTION: [768, 768]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 50
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -40,7 +40,7 @@ META:
|
||||
TRAIN_BATCH_SIZE: 2
|
||||
TRAIN_PREFIX: ""
|
||||
TRAIN_N_PROMPT: ""
|
||||
RESOLUTION: 768
|
||||
RESOLUTION: [768, 768]
|
||||
MEMORY: 29000
|
||||
EPOCHS: 200
|
||||
SAVE_INTERVAL: 25
|
||||
@@ -78,6 +78,7 @@ SOLVER:
|
||||
MAX_EPOCHS: -1
|
||||
NUM_FOLDS: 1
|
||||
ACCU_STEP: 1
|
||||
EVAL_INTERVAL: -1
|
||||
#
|
||||
WORK_DIR:
|
||||
LOG_FILE: std_log.txt
|
||||
@@ -240,15 +241,44 @@ SOLVER:
|
||||
KEYS: [ 'image', 'prompt' ]
|
||||
META_KEYS: [ 'data_key' ]
|
||||
#
|
||||
EVAL_DATA:
|
||||
NAME: Text2ImageDataset
|
||||
MODE: eval
|
||||
PROMPT_FILE:
|
||||
PROMPT_DATA: [ "a boy wearing a jacket", "a dog running on the lawn" ]
|
||||
IMAGE_SIZE: [ 768, 768 ]
|
||||
FIELDS: [ "prompt" ]
|
||||
DELIMITER: '#;#'
|
||||
PROMPT_PREFIX: ''
|
||||
PIN_MEMORY: True
|
||||
BATCH_SIZE: 1
|
||||
NUM_WORKERS: 4
|
||||
TRANSFORMS:
|
||||
- NAME: Select
|
||||
KEYS: [ 'index', 'prompt' ]
|
||||
META_KEYS: [ 'image_size' ]
|
||||
#
|
||||
TRAIN_HOOKS:
|
||||
- NAME: BackwardHook
|
||||
# GRADIENT_CLIP: 1.0
|
||||
-
|
||||
NAME: BackwardHook
|
||||
PRIORITY: 0
|
||||
- NAME: LogHook
|
||||
LOG_INTERVAL: 50
|
||||
-
|
||||
NAME: LogHook
|
||||
LOG_INTERVAL: 10
|
||||
SHOW_GPU_MEM: True
|
||||
-
|
||||
NAME: TensorboardLogHook
|
||||
-
|
||||
NAME: CheckpointHook
|
||||
SAVE_LAST: True
|
||||
INTERVAL: 10000
|
||||
PRIORITY: 200
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
#
|
||||
EVAL_HOOKS:
|
||||
-
|
||||
NAME: ProbeDataHook
|
||||
PROB_INTERVAL: 100
|
||||
SAVE_LAST: True
|
||||
SAVE_NAME_PREFIX: 'step'
|
||||
SAVE_PROBE_PREFIX: 'image'
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
WORK_DIR: "tuner_manager"
|
||||
TUNER_LIST_YAML: "tuner_list.yaml"
|
||||
@@ -1,8 +1,11 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.annotator.base_annotator import GeneralAnnotator
|
||||
from scepter.modules.annotator.canny import CannyAnnotator
|
||||
from scepter.modules.annotator.color import ColorAnnotator
|
||||
from scepter.modules.annotator.hed import HedAnnotator
|
||||
from scepter.modules.annotator.identity import IdentityAnnotator
|
||||
from scepter.modules.annotator.invert import InvertAnnotator
|
||||
from scepter.modules.annotator.midas_op import MidasDetector
|
||||
from scepter.modules.annotator.mlsd_op import MLSDdetector
|
||||
from scepter.modules.annotator.openpose import OpenposeAnnotator
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from abc import ABCMeta
|
||||
|
||||
import cv2
|
||||
@@ -19,20 +20,39 @@ class CannyAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.low_threshold = cfg.get('LOW_THRESHOLD', 100)
|
||||
self.high_threshold = cfg.get('HIGH_THRESHOLD', 200)
|
||||
self.random_cfg = cfg.get('RANDOM_CFG', None)
|
||||
|
||||
def forward(self, image):
|
||||
if isinstance(image, Image.Image):
|
||||
image = np.array(image)
|
||||
image = cv2.Canny(image, self.low_threshold, self.high_threshold)
|
||||
elif isinstance(image, torch.Tensor):
|
||||
image = image.detach().cpu().numpy()
|
||||
image = cv2.Canny(image, self.low_threshold, self.high_threshold)
|
||||
elif isinstance(image, np.ndarray):
|
||||
image = cv2.Canny(image.copy(), self.low_threshold,
|
||||
self.high_threshold)
|
||||
image = image.copy()
|
||||
else:
|
||||
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||
assert len(image.shape) < 4
|
||||
|
||||
if self.random_cfg is None:
|
||||
image = cv2.Canny(image, self.low_threshold, self.high_threshold)
|
||||
else:
|
||||
proba = self.random_cfg.get('PROBA', 1.0)
|
||||
if np.random.random() < proba:
|
||||
min_low_threshold = self.random_cfg.get(
|
||||
'MIN_LOW_THRESHOLD', 50)
|
||||
max_low_threshold = self.random_cfg.get(
|
||||
'MAX_LOW_THRESHOLD', 100)
|
||||
min_high_threshold = self.random_cfg.get(
|
||||
'MIN_HIGH_THRESHOLD', 200)
|
||||
max_high_threshold = self.random_cfg.get(
|
||||
'MAX_HIGH_THRESHOLD', 350)
|
||||
low_th = np.random.randint(min_low_threshold,
|
||||
max_low_threshold)
|
||||
high_th = np.random.randint(min_high_threshold,
|
||||
max_high_threshold)
|
||||
else:
|
||||
low_th, high_th = self.low_threshold, self.high_threshold
|
||||
image = cv2.Canny(image, low_th, high_th)
|
||||
return image[..., None].repeat(3, 2)
|
||||
|
||||
@staticmethod
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from abc import ABCMeta
|
||||
|
||||
import cv2
|
||||
@@ -18,6 +19,7 @@ class ColorAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.ratio = cfg.get('RATIO', 64)
|
||||
self.random_cfg = cfg.get('RANDOM_CFG', None)
|
||||
|
||||
def forward(self, image):
|
||||
if isinstance(image, Image.Image):
|
||||
@@ -29,8 +31,21 @@ class ColorAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
else:
|
||||
raise f'Unsurpport datatype{type(image)}, only surpport np.ndarray, torch.Tensor, Pillow Image.'
|
||||
h, w = image.shape[:2]
|
||||
ratio = self.ratio
|
||||
image = cv2.resize(image, (w // ratio, h // ratio),
|
||||
|
||||
if self.random_cfg is None:
|
||||
ratio = self.ratio
|
||||
else:
|
||||
proba = self.random_cfg.get('PROBA', 1.0)
|
||||
if np.random.random() < proba:
|
||||
if 'CHOICE_RATIO' in self.random_cfg:
|
||||
ratio = np.random.choice(self.random_cfg['CHOICE_RATIO'])
|
||||
else:
|
||||
min_ratio = self.random_cfg.get('MIN_RATIO', 48)
|
||||
max_ratio = self.random_cfg.get('MAX_RATIO', 96)
|
||||
ratio = np.random.randint(min_ratio, max_ratio)
|
||||
else:
|
||||
ratio = self.ratio
|
||||
image = cv2.resize(image, (int(w // ratio), int(h // ratio)),
|
||||
interpolation=cv2.INTER_CUBIC)
|
||||
image = cv2.resize(image, (w, h), interpolation=cv2.INTER_NEAREST)
|
||||
assert len(image.shape) < 4
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
# Please use this implementation in your products
|
||||
# This implementation may produce slightly different results from Saining Xie's official implementations,
|
||||
# but it generates smoother edges and is more suitable for ControlNet as well as other image-to-image translations.
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from abc import ABCMeta
|
||||
|
||||
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class IdentityAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
para_dict = {}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
|
||||
def forward(self, image):
|
||||
return image
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('ANNOTATORS',
|
||||
__class__.__name__,
|
||||
IdentityAnnotator.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,25 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from abc import ABCMeta
|
||||
|
||||
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
|
||||
|
||||
@ANNOTATORS.register_class()
|
||||
class InvertAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
para_dict = {}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
|
||||
def forward(self, image):
|
||||
return 255 - image
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('ANNOTATORS',
|
||||
__class__.__name__,
|
||||
InvertAnnotator.para_dict,
|
||||
set_name=True)
|
||||
@@ -1,4 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
# Midas Depth Estimation
|
||||
# From https://github.com/isl-org/MiDaS
|
||||
# MIT LICENSE
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
# MLSD Line Detection
|
||||
# From https://github.com/navervision/mlsd
|
||||
# Apache-2.0 license
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
# Openpose
|
||||
# Original from CMU https://github.com/CMU-Perceptual-Computing-Lab/openpose
|
||||
# 2nd Edited by https://github.com/Hzzone/pytorch-openpose
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
|
||||
@@ -1,14 +1,13 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import os
|
||||
from abc import ABCMeta, abstractmethod
|
||||
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from scepter.modules.transform.registry import TRANSFORMS, build_pipeline
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import set_random_seed, we
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.modules.utils.logger import get_logger
|
||||
from scepter.modules.utils.registry import old_python_version
|
||||
@@ -84,7 +83,6 @@ class BaseDataset(Dataset, metaclass=ABCMeta):
|
||||
overwrite=False)
|
||||
self.worker_id = worker_id
|
||||
self.logger = self.worker_logger
|
||||
set_random_seed(int(os.environ.get('ES_SEED', 2023)))
|
||||
we.set_env(self.local_we)
|
||||
|
||||
@abstractmethod
|
||||
|
||||
@@ -244,12 +244,18 @@ class Text2ImageDataset(BaseDataset):
|
||||
image_size = [image_size, image_size]
|
||||
assert isinstance(image_size, Iterable) and len(image_size) == 2
|
||||
|
||||
prompt_file = cfg.PROMPT_FILE
|
||||
with FS.get_object(prompt_file) as local_data:
|
||||
if cfg.PROMPT_FILE is not None and cfg.PROMPT_FILE != '':
|
||||
prompt_file = cfg.PROMPT_FILE
|
||||
with FS.get_object(prompt_file) as local_data:
|
||||
rows = [
|
||||
i.split(delimiter,
|
||||
len(fields) - 1)
|
||||
for i in local_data.decode('utf-8').strip().split('\n')
|
||||
]
|
||||
else:
|
||||
rows = [
|
||||
i.split(delimiter,
|
||||
len(fields) - 1)
|
||||
for i in local_data.decode('utf-8').strip().split('\n')
|
||||
len(fields) - 1) for i in cfg.PROMPT_DATA
|
||||
]
|
||||
|
||||
self.items = list()
|
||||
@@ -263,10 +269,9 @@ class Text2ImageDataset(BaseDataset):
|
||||
item['meta']['img_path'] = os.path.join(path_prefix, value)
|
||||
elif key in ['width', 'height']:
|
||||
item['meta'][key] = int(value)
|
||||
elif key != 'meta':
|
||||
item[key] = value
|
||||
else:
|
||||
continue
|
||||
item['meta'][key] = value
|
||||
|
||||
self.items.append(item)
|
||||
if use_num > 0:
|
||||
self.items = self.items[:use_num]
|
||||
|
||||
@@ -1,2 +1,3 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.inference.diffusion_inference import DiffusionInference
|
||||
|
||||
@@ -0,0 +1,117 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import os
|
||||
import warnings
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torchvision.transforms as TT
|
||||
from PIL.Image import Image
|
||||
|
||||
from scepter.modules.model.registry import TUNERS
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
try:
|
||||
from swift import SwiftModel
|
||||
except Exception:
|
||||
warnings.warn('Import swift failed, please check it.')
|
||||
|
||||
|
||||
class ControlInference():
|
||||
def __init__(self, logger=None):
|
||||
self.logger = logger
|
||||
self.is_register = False
|
||||
|
||||
# @classmethod
|
||||
def unregister_controllers(self, control_model_ins, diffusion_model):
|
||||
self.logger.info('Unloading control model')
|
||||
if isinstance(diffusion_model['model'], SwiftModel):
|
||||
if (hasattr(diffusion_model['model'].base_model, 'control_blocks')
|
||||
and diffusion_model['model'].base_model.control_blocks
|
||||
): # noqa
|
||||
del diffusion_model['model'].base_model.control_blocks
|
||||
diffusion_model['model'].base_model.control_blocks = None
|
||||
diffusion_model['model'].base_model.control_name = []
|
||||
else:
|
||||
del diffusion_model['model'].control_blocks
|
||||
diffusion_model['model'].control_blocks = None
|
||||
diffusion_model['model'].control_name = []
|
||||
self.is_register = False
|
||||
|
||||
# @classmethod
|
||||
def register_controllers(self, control_model_ins, diffusion_model):
|
||||
self.logger.info('Loading control model')
|
||||
if control_model_ins is None or control_model_ins == '':
|
||||
self.unregister_controllers(control_model_ins, diffusion_model)
|
||||
return
|
||||
if not isinstance(control_model_ins, list):
|
||||
control_model_ins = [control_model_ins]
|
||||
control_model = nn.ModuleList([])
|
||||
control_model_folder = []
|
||||
for one_control in control_model_ins:
|
||||
one_control_model_folder = one_control.MODEL_PATH
|
||||
control_model_folder.append(one_control_model_folder)
|
||||
have_list = getattr(diffusion_model['model'], 'control_name', [])
|
||||
if one_control_model_folder in have_list:
|
||||
ind = have_list.index(one_control_model_folder)
|
||||
csc_tuners = copy.deepcopy(
|
||||
diffusion_model['model'].control_blocks[ind])
|
||||
else:
|
||||
one_local_control_model = FS.get_dir_to_local_dir(
|
||||
one_control_model_folder)
|
||||
control_cfg = Config(cfg_file=os.path.join(
|
||||
one_local_control_model, '0_SwiftSCETuning',
|
||||
'configuration.json'))
|
||||
assert hasattr(control_cfg, 'CONTROL_MODEL')
|
||||
control_cfg.CONTROL_MODEL[
|
||||
'INPUT_BLOCK_CHANS'] = diffusion_model[
|
||||
'model']._input_block_chans
|
||||
control_cfg.CONTROL_MODEL['INPUT_DOWN_FLAG'] = diffusion_model[
|
||||
'model']._input_down_flag
|
||||
control_cfg.CONTROL_MODEL.PRETRAINED_MODEL = os.path.join(
|
||||
one_local_control_model, '0_SwiftSCETuning',
|
||||
'pytorch_model.bin')
|
||||
csc_tuners = TUNERS.build(control_cfg.CONTROL_MODEL,
|
||||
logger=self.logger)
|
||||
control_model.append(csc_tuners)
|
||||
|
||||
control_model.to(diffusion_model['device'])
|
||||
if isinstance(diffusion_model['model'], SwiftModel):
|
||||
del diffusion_model['model'].base_model.control_blocks
|
||||
diffusion_model['model'].base_model.control_blocks = control_model
|
||||
diffusion_model[
|
||||
'model'].base_model.control_name = control_model_folder
|
||||
else:
|
||||
del diffusion_model['model'].control_blocks
|
||||
diffusion_model['model'].control_blocks = control_model
|
||||
diffusion_model['model'].control_name = control_model_folder
|
||||
self.is_register = True
|
||||
|
||||
@classmethod
|
||||
def get_control_input(self, control_model, control_cond_image, height,
|
||||
width):
|
||||
hints = []
|
||||
if control_cond_image and control_model:
|
||||
if not isinstance(control_model, list):
|
||||
control_model = [control_model]
|
||||
if not isinstance(control_cond_image, list):
|
||||
control_cond_image = [control_cond_image]
|
||||
assert len(control_cond_image) == len(control_model)
|
||||
for img in control_cond_image:
|
||||
if isinstance(img, Image):
|
||||
w, h = img.size
|
||||
if not h == height or not w == width:
|
||||
img = TT.Resize(min(height, width))(img)
|
||||
img = TT.CenterCrop((height, width))(img)
|
||||
hint = TT.ToTensor()(img)
|
||||
hints.append(hint)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
if len(hints) > 0:
|
||||
hints = torch.stack(hints).to(we.device_id)
|
||||
else:
|
||||
hints = None
|
||||
return hints
|
||||
@@ -1,27 +1,24 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import hashlib
|
||||
import json
|
||||
import os.path
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms as TT
|
||||
from peft.utils import CONFIG_NAME, SAFETENSORS_WEIGHTS_NAME, WEIGHTS_NAME
|
||||
from PIL.Image import Image
|
||||
from swift import Swift, SwiftModel
|
||||
|
||||
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, TUNERS)
|
||||
from scepter.modules.utils.config import Config
|
||||
TOKENIZERS)
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
from .control_inference import ControlInference
|
||||
from .tuner_inference import TunerInference
|
||||
|
||||
|
||||
def get_model(model_tuple):
|
||||
assert 'model' in model_tuple
|
||||
@@ -36,6 +33,12 @@ class DiffusionInference():
|
||||
'''
|
||||
def __init__(self, logger=None):
|
||||
self.logger = logger
|
||||
self.loaded_model = {}
|
||||
self.loaded_model_name = [
|
||||
'diffusion_model', 'first_stage_model', 'cond_stage_model'
|
||||
]
|
||||
self.tuner_infer = TunerInference(self.logger)
|
||||
self.control_infer = ControlInference(self.logger)
|
||||
|
||||
def init_from_cfg(self, cfg):
|
||||
self.name = cfg.NAME
|
||||
@@ -76,248 +79,6 @@ class DiffusionInference():
|
||||
'vocab_size': self.tokenizer.vocab_size
|
||||
}
|
||||
|
||||
def register_tuner(self, tuner_model_list):
|
||||
if len(tuner_model_list) < 1:
|
||||
if isinstance(self.diffusion_model['model'], SwiftModel):
|
||||
for adapter_name in self.diffusion_model['model'].adapters:
|
||||
self.diffusion_model['model'].deactivate_adapter(
|
||||
adapter_name, offload='cpu')
|
||||
if isinstance(self.cond_stage_model['model'], SwiftModel):
|
||||
for adapter_name in self.cond_stage_model['model'].adapters:
|
||||
self.cond_stage_model['model'].deactivate_adapter(
|
||||
adapter_name, offload='cpu')
|
||||
return
|
||||
all_diffusion_tuner = {}
|
||||
all_cond_tuner = {}
|
||||
save_root_dir = '.cache_tuner'
|
||||
for tuner_model in tuner_model_list:
|
||||
tunner_model_folder = tuner_model.MODEL_PATH
|
||||
local_tuner_model = FS.get_dir_to_local_dir(tunner_model_folder)
|
||||
all_tuner_datas = os.listdir(local_tuner_model)
|
||||
cur_tuner_md5 = hashlib.md5(
|
||||
tunner_model_folder.encode('utf-8')).hexdigest()
|
||||
|
||||
local_diffusion_cache = os.path.join(
|
||||
save_root_dir, cur_tuner_md5 + '_' + 'diffusion')
|
||||
local_cond_cache = os.path.join(save_root_dir,
|
||||
cur_tuner_md5 + '_' + 'cond')
|
||||
|
||||
meta_file = os.path.join(save_root_dir,
|
||||
cur_tuner_md5 + '_meta.json')
|
||||
if not os.path.exists(meta_file):
|
||||
diffusion_tuner = {}
|
||||
cond_tuner = {}
|
||||
for sub in all_tuner_datas:
|
||||
sub_file = os.path.join(local_tuner_model, sub)
|
||||
config_file = os.path.join(sub_file, CONFIG_NAME)
|
||||
safe_file = os.path.join(sub_file,
|
||||
SAFETENSORS_WEIGHTS_NAME)
|
||||
bin_file = os.path.join(sub_file, WEIGHTS_NAME)
|
||||
if os.path.isdir(sub_file) and os.path.isfile(config_file):
|
||||
# diffusion or cond
|
||||
cfg = json.load(open(config_file, 'r'))
|
||||
if 'cond_stage_model.' in cfg['target_modules']:
|
||||
cond_cfg = copy.deepcopy(cfg)
|
||||
if 'cond_stage_model.*' in cond_cfg[
|
||||
'target_modules']:
|
||||
cond_cfg['target_modules'] = cond_cfg[
|
||||
'target_modules'].replace(
|
||||
'cond_stage_model.*', '.*')
|
||||
else:
|
||||
cond_cfg['target_modules'] = cond_cfg[
|
||||
'target_modules'].replace(
|
||||
'cond_stage_model.', '')
|
||||
if cond_cfg['target_modules'].startswith('*'):
|
||||
cond_cfg['target_modules'] = '.' + cond_cfg[
|
||||
'target_modules']
|
||||
os.makedirs(local_cond_cache + '_' + sub,
|
||||
exist_ok=True)
|
||||
cond_tuner[os.path.basename(local_cond_cache) +
|
||||
'_' + sub] = hashlib.md5(
|
||||
(local_cond_cache + '_' +
|
||||
sub).encode('utf-8')).hexdigest()
|
||||
os.makedirs(local_cond_cache + '_' + sub,
|
||||
exist_ok=True)
|
||||
|
||||
json.dump(
|
||||
cond_cfg,
|
||||
open(
|
||||
os.path.join(local_cond_cache + '_' + sub,
|
||||
CONFIG_NAME), 'w'))
|
||||
if 'model.' in cfg['target_modules'].replace(
|
||||
'cond_stage_model.', ''):
|
||||
diffusion_cfg = copy.deepcopy(cfg)
|
||||
if 'model.*' in diffusion_cfg['target_modules']:
|
||||
diffusion_cfg[
|
||||
'target_modules'] = diffusion_cfg[
|
||||
'target_modules'].replace(
|
||||
'model.*', '.*')
|
||||
else:
|
||||
diffusion_cfg[
|
||||
'target_modules'] = diffusion_cfg[
|
||||
'target_modules'].replace(
|
||||
'model.', '')
|
||||
if diffusion_cfg['target_modules'].startswith('*'):
|
||||
diffusion_cfg[
|
||||
'target_modules'] = '.' + diffusion_cfg[
|
||||
'target_modules']
|
||||
os.makedirs(local_diffusion_cache + '_' + sub,
|
||||
exist_ok=True)
|
||||
diffusion_tuner[
|
||||
os.path.basename(local_diffusion_cache) + '_' +
|
||||
sub] = hashlib.md5(
|
||||
(local_diffusion_cache + '_' +
|
||||
sub).encode('utf-8')).hexdigest()
|
||||
json.dump(
|
||||
diffusion_cfg,
|
||||
open(
|
||||
os.path.join(
|
||||
local_diffusion_cache + '_' + sub,
|
||||
CONFIG_NAME), 'w'))
|
||||
|
||||
state_dict = {}
|
||||
is_bin_file = True
|
||||
if os.path.isfile(bin_file):
|
||||
state_dict = torch.load(bin_file)
|
||||
elif os.path.isfile(safe_file):
|
||||
is_bin_file = False
|
||||
from safetensors.torch import \
|
||||
load_file as safe_load_file
|
||||
state_dict = safe_load_file(
|
||||
safe_file,
|
||||
device='cuda'
|
||||
if torch.cuda.is_available() else 'cpu')
|
||||
save_diffusion_state_dict = {}
|
||||
save_cond_state_dict = {}
|
||||
for key, value in state_dict.items():
|
||||
if key.startswith('model.'):
|
||||
save_diffusion_state_dict[
|
||||
key[len('model.'):].replace(
|
||||
sub,
|
||||
os.path.basename(local_diffusion_cache)
|
||||
+ '_' + sub)] = value
|
||||
elif key.startswith('cond_stage_model.'):
|
||||
save_cond_state_dict[
|
||||
key[len('cond_stage_model.'):].replace(
|
||||
sub,
|
||||
os.path.basename(local_cond_cache) +
|
||||
'_' + sub)] = value
|
||||
|
||||
if is_bin_file:
|
||||
if len(save_diffusion_state_dict) > 0:
|
||||
torch.save(
|
||||
save_diffusion_state_dict,
|
||||
os.path.join(
|
||||
local_diffusion_cache + '_' + sub,
|
||||
WEIGHTS_NAME))
|
||||
if len(save_cond_state_dict) > 0:
|
||||
torch.save(
|
||||
save_cond_state_dict,
|
||||
os.path.join(local_cond_cache + '_' + sub,
|
||||
WEIGHTS_NAME))
|
||||
else:
|
||||
from safetensors.torch import \
|
||||
save_file as safe_save_file
|
||||
if len(save_diffusion_state_dict) > 0:
|
||||
safe_save_file(
|
||||
save_diffusion_state_dict,
|
||||
os.path.join(
|
||||
local_diffusion_cache + '_' + sub,
|
||||
SAFETENSORS_WEIGHTS_NAME),
|
||||
metadata={'format': 'pt'})
|
||||
if len(save_cond_state_dict) > 0:
|
||||
safe_save_file(
|
||||
save_cond_state_dict,
|
||||
os.path.join(local_cond_cache + '_' + sub,
|
||||
SAFETENSORS_WEIGHTS_NAME),
|
||||
metadata={'format': 'pt'})
|
||||
json.dump(
|
||||
{
|
||||
'diffusion_tuner': diffusion_tuner,
|
||||
'cond_tuner': cond_tuner
|
||||
}, open(meta_file, 'w'))
|
||||
else:
|
||||
meta_conf = json.load(open(meta_file, 'r'))
|
||||
diffusion_tuner = meta_conf['diffusion_tuner']
|
||||
cond_tuner = meta_conf['cond_tuner']
|
||||
all_diffusion_tuner.update(diffusion_tuner)
|
||||
all_cond_tuner.update(cond_tuner)
|
||||
if len(all_diffusion_tuner) > 0:
|
||||
self.load(self.diffusion_model)
|
||||
self.diffusion_model['model'] = Swift.from_pretrained(
|
||||
self.diffusion_model['model'],
|
||||
save_root_dir,
|
||||
adapter_name=all_diffusion_tuner)
|
||||
self.diffusion_model['model'].set_active_adapters(
|
||||
list(all_diffusion_tuner.values()))
|
||||
self.unload(self.diffusion_model)
|
||||
if len(all_cond_tuner) > 0:
|
||||
self.load(self.cond_stage_model)
|
||||
self.cond_stage_model['model'] = Swift.from_pretrained(
|
||||
self.cond_stage_model['model'],
|
||||
save_root_dir,
|
||||
adapter_name=all_cond_tuner)
|
||||
self.cond_stage_model['model'].set_active_adapters(
|
||||
list(all_cond_tuner.values()))
|
||||
self.unload(self.cond_stage_model)
|
||||
|
||||
def register_controllers(self, control_model_ins):
|
||||
if control_model_ins is None or control_model_ins == '':
|
||||
if isinstance(self.diffusion_model['model'], SwiftModel):
|
||||
if (hasattr(self.diffusion_model['model'].base_model,
|
||||
'control_blocks') and
|
||||
self.diffusion_model['model'].base_model.control_blocks
|
||||
): # noqa
|
||||
del self.diffusion_model['model'].base_model.control_blocks
|
||||
self.diffusion_model[
|
||||
'model'].base_model.control_blocks = None
|
||||
self.diffusion_model['model'].base_model.control_name = []
|
||||
else:
|
||||
del self.diffusion_model['model'].control_blocks
|
||||
self.diffusion_model['model'].control_blocks = None
|
||||
self.diffusion_model['model'].control_name = []
|
||||
return
|
||||
if not isinstance(control_model_ins, list):
|
||||
control_model_ins = [control_model_ins]
|
||||
control_model = nn.ModuleList([])
|
||||
control_model_folder = []
|
||||
for one_control in control_model_ins:
|
||||
one_control_model_folder = one_control.MODEL_PATH
|
||||
control_model_folder.append(one_control_model_folder)
|
||||
have_list = getattr(self.diffusion_model['model'], 'control_name',
|
||||
[])
|
||||
if one_control_model_folder in have_list:
|
||||
ind = have_list.index(one_control_model_folder)
|
||||
csc_tuners = copy.deepcopy(
|
||||
self.diffusion_model['model'].control_blocks[ind])
|
||||
else:
|
||||
one_local_control_model = FS.get_dir_to_local_dir(
|
||||
one_control_model_folder)
|
||||
control_cfg = Config(cfg_file=os.path.join(
|
||||
one_local_control_model, 'configuration.json'))
|
||||
assert hasattr(control_cfg, 'CONTROL_MODEL')
|
||||
control_cfg.CONTROL_MODEL[
|
||||
'INPUT_BLOCK_CHANS'] = self.diffusion_model[
|
||||
'model']._input_block_chans
|
||||
control_cfg.CONTROL_MODEL[
|
||||
'INPUT_DOWN_FLAG'] = self.diffusion_model[
|
||||
'model']._input_down_flag
|
||||
control_cfg.CONTROL_MODEL.PRETRAINED_MODEL = os.path.join(
|
||||
one_local_control_model, 'pytorch_model.bin')
|
||||
csc_tuners = TUNERS.build(control_cfg.CONTROL_MODEL,
|
||||
logger=self.logger)
|
||||
control_model.append(csc_tuners)
|
||||
if isinstance(self.diffusion_model['model'], SwiftModel):
|
||||
del self.diffusion_model['model'].base_model.control_blocks
|
||||
self.diffusion_model[
|
||||
'model'].base_model.control_blocks = control_model
|
||||
self.diffusion_model[
|
||||
'model'].base_model.control_name = control_model_folder
|
||||
else:
|
||||
del self.diffusion_model['model'].control_blocks
|
||||
self.diffusion_model['model'].control_blocks = control_model
|
||||
self.diffusion_model['model'].control_name = control_model_folder
|
||||
|
||||
def redefine_paras(self, cfg):
|
||||
if cfg.get('PRETRAINED_MODEL', None):
|
||||
assert FS.isfile(cfg.PRETRAINED_MODEL)
|
||||
@@ -480,12 +241,55 @@ class DiffusionInference():
|
||||
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:
|
||||
@@ -591,7 +395,7 @@ class DiffusionInference():
|
||||
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')
|
||||
_, dtype = self.get_function_info(self.first_stage_model, 'decode')
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
@@ -616,35 +420,34 @@ class DiffusionInference():
|
||||
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)
|
||||
#
|
||||
if not isinstance(tuner_model, list):
|
||||
tuner_model = [tuner_model]
|
||||
for tuner in tuner_model:
|
||||
if tuner is None or tuner == '':
|
||||
tuner_model.remove(tuner)
|
||||
self.register_tuner(tuner_model)
|
||||
# control_cond_image
|
||||
control_cond_image = kwargs.pop('control_cond_image', None)
|
||||
# crop_type = kwargs.pop('crop_type', 'center_crop')
|
||||
hints = []
|
||||
if control_cond_image and control_model:
|
||||
if not isinstance(control_model, list):
|
||||
control_model = [control_model]
|
||||
if not isinstance(control_cond_image, list):
|
||||
control_cond_image = [control_cond_image]
|
||||
assert len(control_cond_image) == len(control_model)
|
||||
for img in control_cond_image:
|
||||
if isinstance(img, Image):
|
||||
w, h = img.size
|
||||
if not h == height or not w == width:
|
||||
img = TT.Resize(min(height, width))(img)
|
||||
img = TT.CenterCrop((height, width))(img)
|
||||
hint = TT.ToTensor()(img)
|
||||
hints.append(hint)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
if len(hints) > 0:
|
||||
hints = torch.stack(hints).to(we.device_id)
|
||||
|
||||
# register tuner
|
||||
if tuner_model is not None and tuner_model != '' and len(
|
||||
tuner_model) > 0:
|
||||
if not isinstance(tuner_model, list):
|
||||
tuner_model = [tuner_model]
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
|
||||
self.tuner_infer.register_tuner(tuner_model, self.diffusion_model,
|
||||
self.cond_stage_model)
|
||||
self.dynamic_unload(self.diffusion_model,
|
||||
'diffusion_model',
|
||||
skip_loaded=True)
|
||||
self.dynamic_unload(self.cond_stage_model,
|
||||
'cond_stage_model',
|
||||
skip_loaded=True)
|
||||
|
||||
# register control
|
||||
if control_model is not None and control_model != '':
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
hints = ControlInference.get_control_input(
|
||||
control_model, kwargs.pop('control_cond_image', None), height,
|
||||
width)
|
||||
self.control_infer.register_controllers(control_model,
|
||||
self.diffusion_model)
|
||||
self.dynamic_unload(self.diffusion_model,
|
||||
'diffusion_model',
|
||||
skip_loaded=True)
|
||||
else:
|
||||
hints = None
|
||||
|
||||
@@ -655,30 +458,38 @@ class DiffusionInference():
|
||||
b, c, ori_width, ori_height = image.shape
|
||||
if not (ori_width == width and ori_height == height):
|
||||
image = F.interpolate(image, (width, height), mode='bicubic')
|
||||
self.first_stage_model = self.load(self.first_stage_model)
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
input_latent = self.encode_first_stage(image)
|
||||
self.first_stage_model = self.unload(self.first_stage_model)
|
||||
self.dynamic_unload(self.first_stage_model,
|
||||
'first_stage_model',
|
||||
skip_loaded=True)
|
||||
else:
|
||||
input_latent = None
|
||||
if 'input_latent' in value_output and input_latent is not None:
|
||||
value_output['input_latent'] = input_latent
|
||||
# cond stage
|
||||
self.cond_stage_model = self.load(self.cond_stage_model)
|
||||
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)):
|
||||
if self.tokenizer:
|
||||
if not hasattr(get_model(self.cond_stage_model), 'tokenizer'):
|
||||
setattr(get_model(self.cond_stage_model), 'tokenizer',
|
||||
self.tokenizer)
|
||||
context = getattr(get_model(self.cond_stage_model),
|
||||
function_name)(batch['tokens'])
|
||||
null_context = getattr(get_model(self.cond_stage_model),
|
||||
function_name)(batch_uc['tokens'])
|
||||
|
||||
else:
|
||||
context = getattr(get_model(self.cond_stage_model),
|
||||
function_name)(batch)
|
||||
null_context = getattr(get_model(self.cond_stage_model),
|
||||
function_name)(batch_uc)
|
||||
self.cond_stage_model = self.unload(self.cond_stage_model)
|
||||
self.dynamic_unload(self.cond_stage_model,
|
||||
'cond_stage_model',
|
||||
skip_loaded=True)
|
||||
|
||||
if refine_strength > 0 and self.refiner_diffusion_model is not None:
|
||||
assert self.refiner_cond_model is not None
|
||||
@@ -703,9 +514,7 @@ class DiffusionInference():
|
||||
get_model(self.refiner_cond_model),
|
||||
function_name)(batch_uc)
|
||||
self.refiner_cond_model = self.unload(self.refiner_cond_model)
|
||||
self.load(self.diffusion_model)
|
||||
self.register_controllers(control_model)
|
||||
self.unload(self.diffusion_model)
|
||||
|
||||
# get noise
|
||||
seed = kwargs.pop('seed', -1)
|
||||
g = torch.Generator(device=we.device_id)
|
||||
@@ -721,8 +530,8 @@ class DiffusionInference():
|
||||
height // self.first_stage_model['paras']['size_factor'],
|
||||
width // self.first_stage_model['paras']['size_factor'],
|
||||
device=we.device_id).normal_(generator=g)
|
||||
#
|
||||
self.load(self.diffusion_model)
|
||||
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
# UNet use input n_prompt
|
||||
function_name, dtype = self.get_function_info(
|
||||
self.diffusion_model)
|
||||
@@ -753,14 +562,18 @@ class DiffusionInference():
|
||||
seed=seed,
|
||||
condition_fn=None,
|
||||
clamp=None,
|
||||
sharpness=value_input.get('sharpness', 0.0),
|
||||
percentile=None,
|
||||
t_max=None,
|
||||
t_min=None,
|
||||
discard_penultimate_step=None,
|
||||
intermediate_callback=intermediate_callback,
|
||||
cat_uc=cat_uc,
|
||||
cat_uc=value_input.get('cat_uc', cat_uc),
|
||||
**kwargs)
|
||||
self.diffusion_model = self.unload(self.diffusion_model)
|
||||
|
||||
self.dynamic_unload(self.diffusion_model,
|
||||
'diffusion_model',
|
||||
skip_loaded=True)
|
||||
|
||||
# apply refiner
|
||||
if refine_strength > 0 and self.refiner_diffusion_model is not None:
|
||||
@@ -828,10 +641,11 @@ class DiffusionInference():
|
||||
value_output['latent'] = []
|
||||
value_output['latent'].append(latent)
|
||||
|
||||
self.first_stage_model = self.load(self.first_stage_model)
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
x_samples = self.decode_first_stage(latent).float()
|
||||
self.first_stage_model = self.unload(self.first_stage_model)
|
||||
|
||||
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 'images' in value_output:
|
||||
if value_output['images'] is None or (
|
||||
@@ -845,4 +659,17 @@ class DiffusionInference():
|
||||
value_output[k] = torch.cat(v, dim=0)
|
||||
if isinstance(v, torch.Tensor):
|
||||
value_output[k] = v.cpu()
|
||||
|
||||
# unregister tuner
|
||||
if tuner_model is not None and tuner_model != '' and len(
|
||||
tuner_model) > 0:
|
||||
self.tuner_infer.unregister_tuner(tuner_model,
|
||||
self.diffusion_model,
|
||||
self.cond_stage_model)
|
||||
|
||||
# unregister control
|
||||
if control_model is not None and control_model != '':
|
||||
self.control_infer.unregister_controllers(control_model,
|
||||
self.diffusion_model)
|
||||
|
||||
return value_output
|
||||
|
||||
@@ -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
|
||||
@@ -0,0 +1,220 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import warnings
|
||||
|
||||
import torch
|
||||
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
try:
|
||||
from peft.utils import CONFIG_NAME, SAFETENSORS_WEIGHTS_NAME, WEIGHTS_NAME
|
||||
except Exception as e:
|
||||
warnings.warn(f'Import peft error, please deal with this problem: {e}')
|
||||
try:
|
||||
from swift import Swift, SwiftModel
|
||||
except Exception as e:
|
||||
warnings.warn(f'Import swift error, please deal with this problem: {e}')
|
||||
|
||||
|
||||
class TunerInference():
|
||||
def __init__(self, logger=None):
|
||||
self.logger = logger
|
||||
self.is_register = False
|
||||
|
||||
# @classmethod
|
||||
def unregister_tuner(self, tuner_model_list, diffusion_model,
|
||||
cond_stage_model):
|
||||
self.logger.info('Unloading tuner model')
|
||||
if isinstance(diffusion_model['model'], SwiftModel):
|
||||
for adapter_name in diffusion_model['model'].adapters:
|
||||
diffusion_model['model'].deactivate_adapter(adapter_name,
|
||||
offload='cpu')
|
||||
if isinstance(cond_stage_model['model'], SwiftModel):
|
||||
for adapter_name in cond_stage_model['model'].adapters:
|
||||
cond_stage_model['model'].deactivate_adapter(adapter_name,
|
||||
offload='cpu')
|
||||
return
|
||||
|
||||
# @classmethod
|
||||
def register_tuner(self, tuner_model_list, diffusion_model,
|
||||
cond_stage_model):
|
||||
self.logger.info('Loading tuner model')
|
||||
if len(tuner_model_list) < 1:
|
||||
self.unregister_tuner(tuner_model_list, diffusion_model,
|
||||
cond_stage_model)
|
||||
return
|
||||
all_diffusion_tuner = {}
|
||||
all_cond_tuner = {}
|
||||
save_root_dir = '.cache_tuner'
|
||||
for tuner_model in tuner_model_list:
|
||||
tunner_model_folder = tuner_model.MODEL_PATH
|
||||
local_tuner_model = FS.get_dir_to_local_dir(tunner_model_folder)
|
||||
all_tuner_datas = os.listdir(local_tuner_model)
|
||||
cur_tuner_md5 = hashlib.md5(
|
||||
tunner_model_folder.encode('utf-8')).hexdigest()
|
||||
|
||||
local_diffusion_cache = os.path.join(
|
||||
save_root_dir, cur_tuner_md5 + '_' + 'diffusion')
|
||||
local_cond_cache = os.path.join(save_root_dir,
|
||||
cur_tuner_md5 + '_' + 'cond')
|
||||
|
||||
meta_file = os.path.join(save_root_dir,
|
||||
cur_tuner_md5 + '_meta.json')
|
||||
if not os.path.exists(meta_file):
|
||||
diffusion_tuner = {}
|
||||
cond_tuner = {}
|
||||
for sub in all_tuner_datas:
|
||||
sub_file = os.path.join(local_tuner_model, sub)
|
||||
config_file = os.path.join(sub_file, CONFIG_NAME)
|
||||
safe_file = os.path.join(sub_file,
|
||||
SAFETENSORS_WEIGHTS_NAME)
|
||||
bin_file = os.path.join(sub_file, WEIGHTS_NAME)
|
||||
if os.path.isdir(sub_file) and os.path.isfile(config_file):
|
||||
# diffusion or cond
|
||||
cfg = json.load(open(config_file, 'r'))
|
||||
if 'cond_stage_model.' in cfg['target_modules']:
|
||||
cond_cfg = copy.deepcopy(cfg)
|
||||
if 'cond_stage_model.*' in cond_cfg[
|
||||
'target_modules']:
|
||||
cond_cfg['target_modules'] = cond_cfg[
|
||||
'target_modules'].replace(
|
||||
'cond_stage_model.*', '.*')
|
||||
else:
|
||||
cond_cfg['target_modules'] = cond_cfg[
|
||||
'target_modules'].replace(
|
||||
'cond_stage_model.', '')
|
||||
if cond_cfg['target_modules'].startswith('*'):
|
||||
cond_cfg['target_modules'] = '.' + cond_cfg[
|
||||
'target_modules']
|
||||
os.makedirs(local_cond_cache + '_' + sub,
|
||||
exist_ok=True)
|
||||
cond_tuner[os.path.basename(local_cond_cache) +
|
||||
'_' + sub] = hashlib.md5(
|
||||
(local_cond_cache + '_' +
|
||||
sub).encode('utf-8')).hexdigest()
|
||||
os.makedirs(local_cond_cache + '_' + sub,
|
||||
exist_ok=True)
|
||||
|
||||
json.dump(
|
||||
cond_cfg,
|
||||
open(
|
||||
os.path.join(local_cond_cache + '_' + sub,
|
||||
CONFIG_NAME), 'w'))
|
||||
if 'model.' in cfg['target_modules'].replace(
|
||||
'cond_stage_model.', ''):
|
||||
diffusion_cfg = copy.deepcopy(cfg)
|
||||
if 'model.*' in diffusion_cfg['target_modules']:
|
||||
diffusion_cfg[
|
||||
'target_modules'] = diffusion_cfg[
|
||||
'target_modules'].replace(
|
||||
'model.*', '.*')
|
||||
else:
|
||||
diffusion_cfg[
|
||||
'target_modules'] = diffusion_cfg[
|
||||
'target_modules'].replace(
|
||||
'model.', '')
|
||||
if diffusion_cfg['target_modules'].startswith('*'):
|
||||
diffusion_cfg[
|
||||
'target_modules'] = '.' + diffusion_cfg[
|
||||
'target_modules']
|
||||
os.makedirs(local_diffusion_cache + '_' + sub,
|
||||
exist_ok=True)
|
||||
diffusion_tuner[
|
||||
os.path.basename(local_diffusion_cache) + '_' +
|
||||
sub] = hashlib.md5(
|
||||
(local_diffusion_cache + '_' +
|
||||
sub).encode('utf-8')).hexdigest()
|
||||
json.dump(
|
||||
diffusion_cfg,
|
||||
open(
|
||||
os.path.join(
|
||||
local_diffusion_cache + '_' + sub,
|
||||
CONFIG_NAME), 'w'))
|
||||
|
||||
state_dict = {}
|
||||
is_bin_file = True
|
||||
if os.path.isfile(bin_file):
|
||||
state_dict = torch.load(bin_file)
|
||||
elif os.path.isfile(safe_file):
|
||||
is_bin_file = False
|
||||
from safetensors.torch import \
|
||||
load_file as safe_load_file
|
||||
state_dict = safe_load_file(
|
||||
safe_file,
|
||||
device='cuda'
|
||||
if torch.cuda.is_available() else 'cpu')
|
||||
save_diffusion_state_dict = {}
|
||||
save_cond_state_dict = {}
|
||||
for key, value in state_dict.items():
|
||||
if key.startswith('model.'):
|
||||
save_diffusion_state_dict[
|
||||
key[len('model.'):].replace(
|
||||
sub,
|
||||
os.path.basename(local_diffusion_cache)
|
||||
+ '_' + sub)] = value
|
||||
elif key.startswith('cond_stage_model.'):
|
||||
save_cond_state_dict[
|
||||
key[len('cond_stage_model.'):].replace(
|
||||
sub,
|
||||
os.path.basename(local_cond_cache) +
|
||||
'_' + sub)] = value
|
||||
|
||||
if is_bin_file:
|
||||
if len(save_diffusion_state_dict) > 0:
|
||||
torch.save(
|
||||
save_diffusion_state_dict,
|
||||
os.path.join(
|
||||
local_diffusion_cache + '_' + sub,
|
||||
WEIGHTS_NAME))
|
||||
if len(save_cond_state_dict) > 0:
|
||||
torch.save(
|
||||
save_cond_state_dict,
|
||||
os.path.join(local_cond_cache + '_' + sub,
|
||||
WEIGHTS_NAME))
|
||||
else:
|
||||
from safetensors.torch import \
|
||||
save_file as safe_save_file
|
||||
if len(save_diffusion_state_dict) > 0:
|
||||
safe_save_file(
|
||||
save_diffusion_state_dict,
|
||||
os.path.join(
|
||||
local_diffusion_cache + '_' + sub,
|
||||
SAFETENSORS_WEIGHTS_NAME),
|
||||
metadata={'format': 'pt'})
|
||||
if len(save_cond_state_dict) > 0:
|
||||
safe_save_file(
|
||||
save_cond_state_dict,
|
||||
os.path.join(local_cond_cache + '_' + sub,
|
||||
SAFETENSORS_WEIGHTS_NAME),
|
||||
metadata={'format': 'pt'})
|
||||
json.dump(
|
||||
{
|
||||
'diffusion_tuner': diffusion_tuner,
|
||||
'cond_tuner': cond_tuner
|
||||
}, open(meta_file, 'w'))
|
||||
else:
|
||||
meta_conf = json.load(open(meta_file, 'r'))
|
||||
diffusion_tuner = meta_conf['diffusion_tuner']
|
||||
cond_tuner = meta_conf['cond_tuner']
|
||||
all_diffusion_tuner.update(diffusion_tuner)
|
||||
all_cond_tuner.update(cond_tuner)
|
||||
if len(all_diffusion_tuner) > 0:
|
||||
|
||||
diffusion_model['model'] = Swift.from_pretrained(
|
||||
diffusion_model['model'],
|
||||
save_root_dir,
|
||||
adapter_name=all_diffusion_tuner)
|
||||
diffusion_model['model'].set_active_adapters(
|
||||
list(all_diffusion_tuner.values()))
|
||||
if len(all_cond_tuner) > 0:
|
||||
cond_stage_model['model'] = Swift.from_pretrained(
|
||||
cond_stage_model['model'],
|
||||
save_root_dir,
|
||||
adapter_name=all_cond_tuner)
|
||||
cond_stage_model['model'].set_active_adapters(
|
||||
list(all_cond_tuner.values()))
|
||||
self.is_register = True
|
||||
@@ -1,4 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.backbone.autoencoder.ae_module import (Decoder,
|
||||
Encoder)
|
||||
Encoder,
|
||||
RDecoder)
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import repeat
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
from scepter.modules.model.backbone.autoencoder.ae_utils import (
|
||||
@@ -244,6 +245,7 @@ class Decoder(BaseModel):
|
||||
|
||||
# compute in_ch_mult, block_in and curr_res at lowest res
|
||||
block_in = self.ch * self.ch_mult[self.num_resolutions - 1]
|
||||
self.block_in = block_in
|
||||
curr_res = 1
|
||||
# z to block_in
|
||||
self.conv_in = torch.nn.Conv2d(self.z_channels,
|
||||
@@ -340,3 +342,48 @@ class Decoder(BaseModel):
|
||||
__class__.__name__,
|
||||
Decoder.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@BACKBONES.register_class()
|
||||
class RDecoder(Decoder):
|
||||
def construct_model(self):
|
||||
super().construct_model()
|
||||
self.resize_level = nn.Sequential(
|
||||
nn.Linear(self.block_in, self.block_in),
|
||||
nn.SiLU(),
|
||||
nn.Linear(self.block_in, self.block_in),
|
||||
)
|
||||
|
||||
def forward(self, z, rembed=None):
|
||||
# timestep embedding
|
||||
temb = None
|
||||
h = self.conv_in(z)
|
||||
bs, channel, hdim, wdim = h.size()
|
||||
if rembed is not None:
|
||||
rembed = self.resize_level(rembed)
|
||||
rembed = repeat(rembed, 'b e-> b e hd wd', hd=hdim, wd=wdim)
|
||||
h = h + rembed
|
||||
|
||||
# middle
|
||||
if not self.use_checkpoint:
|
||||
h = self.mid_upsclae_transform(h, temb)
|
||||
else:
|
||||
h = checkpoint(self.mid_upsclae_transform, h, temb)
|
||||
|
||||
# end
|
||||
if self.give_pre_end:
|
||||
return h
|
||||
|
||||
h = self.norm_out(h)
|
||||
h = nonlinearity(h)
|
||||
h = self.conv_out(h)
|
||||
if self.tanh_out:
|
||||
h = torch.tanh(h)
|
||||
return h
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('BACKBONE',
|
||||
__class__.__name__,
|
||||
Decoder.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.backbone.unet.unet_module import DiffusionUNet
|
||||
from scepter.modules.model.backbone.unet.unet_module import (DiffusionUNet,
|
||||
DiffusionUNetXL,
|
||||
LargenUNetXL)
|
||||
|
||||
@@ -8,8 +8,9 @@ import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from scepter.modules.model.backbone.unet.unet_utils import (
|
||||
Downsample, ResBlock, SpatialTransformer, Timestep,
|
||||
TimestepEmbedSequential, Upsample, conv_nd, linear, normalization,
|
||||
BasicTransformerBlock, Downsample, ResBlock, SpatialTransformer,
|
||||
SpatialTransformerV2, Timestep, TimestepEmbedSequential,
|
||||
TransformerBlockV2, Upsample, conv_nd, linear, normalization,
|
||||
timestep_embedding, zero_module)
|
||||
from scepter.modules.model.base_model import BaseModel
|
||||
from scepter.modules.model.registry import BACKBONES
|
||||
@@ -484,7 +485,7 @@ class DiffusionUNet(BaseModel):
|
||||
if len(unexpected) > 0:
|
||||
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
|
||||
|
||||
def _forward_origin(self, x, emb, context, hint=None):
|
||||
def _forward_origin(self, x, emb, context, hint=None, **kwargs):
|
||||
hs = []
|
||||
h = x
|
||||
for module in self.input_blocks:
|
||||
@@ -492,13 +493,21 @@ class DiffusionUNet(BaseModel):
|
||||
hs.append(h)
|
||||
h = self.middle_block(h, emb, context)
|
||||
for m_id, module in enumerate(self.output_blocks):
|
||||
h = torch.cat([h, self.lsc_identity[m_id](hs.pop())], dim=1)
|
||||
skip_h = hs.pop()
|
||||
if 'tuner_scale' in kwargs and kwargs[
|
||||
'tuner_scale'] is not None and kwargs['tuner_scale'] < 1.0:
|
||||
tuner_scale = kwargs['tuner_scale']
|
||||
tuner_h = self.lsc_identity[m_id](skip_h) - skip_h
|
||||
h = torch.cat([h, skip_h + tuner_scale * tuner_h], dim=1)
|
||||
else:
|
||||
h = torch.cat([h, self.lsc_identity[m_id](skip_h)], dim=1)
|
||||
target_size = hs[-1].shape[-2:] if len(hs) > 0 else None
|
||||
h = module(h, emb, context, target_size)
|
||||
out = self.out(h)
|
||||
return out
|
||||
|
||||
def _forward_control(self, x, emb, context, hint, alpha=0.5):
|
||||
def _forward_control(self, x, emb, context, hint, **kwargs):
|
||||
control_scale = kwargs.pop('control_scale', 1.0)
|
||||
multi_csc_tuners = self.control_blocks
|
||||
# hints
|
||||
multi_hint_hs = []
|
||||
@@ -531,11 +540,11 @@ class DiffusionUNet(BaseModel):
|
||||
torch.zeros_like(tuner_h),
|
||||
atol=1e-6)):
|
||||
# csc-tuner
|
||||
skip_h_new = skip_h + multi_control_h
|
||||
skip_h_new = skip_h + control_scale * multi_control_h
|
||||
else:
|
||||
# csc-tuner + sc-tuner
|
||||
skip_h_new = skip_h + alpha * multi_control_h + (
|
||||
1 - alpha) * tuner_h
|
||||
tuner_scale = kwargs['tuner_scale']
|
||||
skip_h_new = skip_h + control_scale * multi_control_h + tuner_scale * tuner_h
|
||||
h = torch.cat([h, skip_h_new], dim=1)
|
||||
target_size = hs[-1].shape[-2:] if len(hs) > 0 else None
|
||||
h = module(h, emb, context, target_size)
|
||||
@@ -545,7 +554,6 @@ class DiffusionUNet(BaseModel):
|
||||
def forward(self, x, t=None, cond=dict(), **kwargs):
|
||||
t_emb = timestep_embedding(t, self.model_channels, repeat_only=False)
|
||||
emb = self.time_embed(t_emb)
|
||||
hint = None
|
||||
if isinstance(cond, dict):
|
||||
if 'y' in cond and cond['y'] is not None:
|
||||
assert self.num_classes is not None
|
||||
@@ -555,15 +563,19 @@ class DiffusionUNet(BaseModel):
|
||||
x = torch.cat([x, c], dim=1)
|
||||
if 'hint' in cond:
|
||||
hint = cond['hint']
|
||||
elif 'hint' in kwargs:
|
||||
hint = kwargs.pop('hint', None)
|
||||
else:
|
||||
hint = None
|
||||
context = cond.get('crossattn', None)
|
||||
else:
|
||||
context = cond
|
||||
hint = kwargs.pop('hint', None)
|
||||
|
||||
if self.control_blocks is not None:
|
||||
out = self._forward_control(x, emb, context, hint)
|
||||
if self.control_blocks is not None and hint is not None:
|
||||
out = self._forward_control(x, emb, context, hint, **kwargs)
|
||||
else:
|
||||
out = self._forward_origin(x, emb, context)
|
||||
out = self._forward_origin(x, emb, context, **kwargs)
|
||||
return out
|
||||
|
||||
@staticmethod
|
||||
@@ -822,7 +834,7 @@ class DiffusionUNetXL(DiffusionUNet):
|
||||
conv_nd(dims, model_channels, out_channels, 3, padding=1)),
|
||||
)
|
||||
|
||||
def _forward_origin(self, x, emb, context, hint=None):
|
||||
def _forward_origin(self, x, emb, context, hint=None, **kwargs):
|
||||
hs = []
|
||||
h = x
|
||||
for module in self.input_blocks:
|
||||
@@ -830,13 +842,21 @@ class DiffusionUNetXL(DiffusionUNet):
|
||||
hs.append(h)
|
||||
h = self.middle_block(h, emb, context)
|
||||
for m_id, module in enumerate(self.output_blocks):
|
||||
h = torch.cat([h, self.lsc_identity[m_id](hs.pop())], dim=1)
|
||||
skip_h = hs.pop()
|
||||
if 'tuner_scale' in kwargs and kwargs[
|
||||
'tuner_scale'] is not None and kwargs['tuner_scale'] < 1.0:
|
||||
tuner_scale = kwargs['tuner_scale']
|
||||
tuner_h = self.lsc_identity[m_id](skip_h) - skip_h
|
||||
h = torch.cat([h, skip_h + tuner_scale * tuner_h], dim=1)
|
||||
else:
|
||||
h = torch.cat([h, self.lsc_identity[m_id](skip_h)], dim=1)
|
||||
target_size = hs[-1].shape[-2:] if len(hs) > 0 else None
|
||||
h = module(h, emb, context, target_size)
|
||||
out = self.out(h)
|
||||
return out
|
||||
|
||||
def _forward_control(self, x, emb, context, hint, alpha=0.5):
|
||||
def _forward_control(self, x, emb, context, hint, **kwargs):
|
||||
control_scale = kwargs.pop('control_scale', 1.0)
|
||||
multi_csc_tuners = self.control_blocks
|
||||
# hints
|
||||
multi_hint_hs = []
|
||||
@@ -869,11 +889,11 @@ class DiffusionUNetXL(DiffusionUNet):
|
||||
torch.zeros_like(tuner_h),
|
||||
atol=1e-6)):
|
||||
# csc-tuner
|
||||
skip_h_new = skip_h + multi_control_h
|
||||
skip_h_new = skip_h + control_scale * multi_control_h
|
||||
else:
|
||||
# csc-tuner + sc-tuner
|
||||
skip_h_new = skip_h + alpha * multi_control_h + (
|
||||
1 - alpha) * tuner_h
|
||||
tuner_scale = kwargs['tuner_scale']
|
||||
skip_h_new = skip_h + control_scale * multi_control_h + tuner_scale * tuner_h
|
||||
h = torch.cat([h, skip_h_new], dim=1)
|
||||
target_size = hs[-1].shape[-2:] if len(hs) > 0 else None
|
||||
h = module(h, emb, context, target_size)
|
||||
@@ -886,7 +906,6 @@ class DiffusionUNetXL(DiffusionUNet):
|
||||
repeat_only=False,
|
||||
legacy=True)
|
||||
emb = self.time_embed(t_emb)
|
||||
hint = None
|
||||
if isinstance(cond, dict):
|
||||
if 'y' in cond:
|
||||
assert self.num_classes is not None
|
||||
@@ -896,15 +915,19 @@ class DiffusionUNetXL(DiffusionUNet):
|
||||
x = torch.cat([x, c], dim=1)
|
||||
if 'hint' in cond:
|
||||
hint = cond['hint']
|
||||
elif 'hint' in kwargs:
|
||||
hint = kwargs.pop('hint', None)
|
||||
else:
|
||||
hint = None
|
||||
context = cond.get('crossattn', None)
|
||||
else:
|
||||
context = cond
|
||||
hint = kwargs.pop('hint', None)
|
||||
|
||||
if self.control_blocks is not None:
|
||||
out = self._forward_control(x, emb, context, hint)
|
||||
if self.control_blocks is not None and hint is not None:
|
||||
out = self._forward_control(x, emb, context, hint, **kwargs)
|
||||
else:
|
||||
out = self._forward_origin(x, emb, context)
|
||||
out = self._forward_origin(x, emb, context, **kwargs)
|
||||
return out
|
||||
|
||||
def convert_to_fp16(self):
|
||||
@@ -929,3 +952,443 @@ class DiffusionUNetXL(DiffusionUNet):
|
||||
__class__.__name__,
|
||||
DiffusionUNetXL.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@BACKBONES.register_class()
|
||||
class LargenUNetXL(DiffusionUNetXL):
|
||||
para_dict = {
|
||||
'TRANSFORMER_BLOCK_TYPE': {
|
||||
'value': 'att_v1'
|
||||
},
|
||||
'IMAGE_SCALE': {
|
||||
'value': 0.0,
|
||||
},
|
||||
}
|
||||
para_dict.update(DiffusionUNetXL.para_dict)
|
||||
|
||||
def __init__(self, cfg, logger):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.init_params(cfg)
|
||||
self.construct_network()
|
||||
|
||||
def init_params(self, cfg):
|
||||
super().init_params(cfg)
|
||||
self.transformer_block_type = cfg.get('TRANSFORMER_BLOCK_TYPE',
|
||||
'att_v1')
|
||||
TRANSFORMER_BLOCKS = {
|
||||
'att_v1': BasicTransformerBlock,
|
||||
'att_v2': TransformerBlockV2,
|
||||
}
|
||||
assert self.transformer_block_type in list(TRANSFORMER_BLOCKS.keys())
|
||||
self.transformer_block = TRANSFORMER_BLOCKS[
|
||||
self.transformer_block_type]
|
||||
self.image_scale = cfg.get('IMAGE_SCALE', 0.0)
|
||||
self.use_refine = cfg.get('USE_REFINE', False)
|
||||
|
||||
def construct_network(self):
|
||||
in_channels = self.in_channels
|
||||
model_channels = self.model_channels
|
||||
out_channels = self.out_channels
|
||||
attention_resolutions = self.attention_resolutions
|
||||
channel_mult = self.channel_mult
|
||||
num_classes = self.num_classes
|
||||
num_heads = self.num_heads
|
||||
num_head_channels = self.num_head_channels
|
||||
dims = self.dims
|
||||
dropout = self.dropout
|
||||
use_checkpoint = self.use_checkpoint
|
||||
use_scale_shift_norm = self.use_scale_shift_norm
|
||||
disable_self_attentions = self.disable_self_attentions
|
||||
disable_middle_self_attn = self.disable_middle_self_attn
|
||||
transformer_depth = self.transformer_depth
|
||||
transformer_depth_middle = self.transformer_depth_middle
|
||||
context_dim = self.context_dim
|
||||
use_linear_in_transformer = self.use_linear_in_transformer
|
||||
resblock_updown = self.resblock_updown
|
||||
conv_resample = self.conv_resample
|
||||
adm_in_channels = self.adm_in_channels
|
||||
transformer_block = self.transformer_block
|
||||
|
||||
time_embed_dim = model_channels * 4
|
||||
self.time_embed = nn.Sequential(
|
||||
linear(model_channels, time_embed_dim),
|
||||
nn.SiLU(),
|
||||
linear(time_embed_dim, time_embed_dim),
|
||||
)
|
||||
|
||||
if self.num_classes is not None:
|
||||
if isinstance(self.num_classes, int):
|
||||
self.label_emb = nn.Embedding(num_classes, time_embed_dim)
|
||||
elif self.num_classes == 'continuous':
|
||||
print('setting up linear c_adm embedding layer')
|
||||
self.label_emb = nn.Linear(1, time_embed_dim)
|
||||
elif self.num_classes == 'timestep':
|
||||
self.label_emb = nn.Sequential(
|
||||
Timestep(model_channels),
|
||||
nn.Sequential(
|
||||
linear(model_channels, time_embed_dim),
|
||||
nn.SiLU(),
|
||||
linear(time_embed_dim, time_embed_dim),
|
||||
),
|
||||
)
|
||||
elif self.num_classes == 'sequential':
|
||||
assert adm_in_channels is not None
|
||||
self.label_emb = nn.Sequential(
|
||||
nn.Sequential(
|
||||
linear(adm_in_channels, time_embed_dim),
|
||||
nn.SiLU(),
|
||||
linear(time_embed_dim, time_embed_dim),
|
||||
))
|
||||
else:
|
||||
raise ValueError()
|
||||
|
||||
self.input_blocks = nn.ModuleList([
|
||||
TimestepEmbedSequential(
|
||||
conv_nd(dims, in_channels, model_channels, 3, padding=1))
|
||||
])
|
||||
self._feature_size = model_channels
|
||||
input_block_chans = [model_channels]
|
||||
input_down_flag = [False]
|
||||
ch = model_channels
|
||||
ds = 1
|
||||
for level, mult in enumerate(channel_mult):
|
||||
for nr in range(self.num_res_blocks[level]):
|
||||
layers = [
|
||||
ResBlock(
|
||||
ch,
|
||||
time_embed_dim,
|
||||
dropout,
|
||||
out_channels=mult * model_channels,
|
||||
dims=dims,
|
||||
use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
)
|
||||
]
|
||||
ch = mult * model_channels
|
||||
if ds in attention_resolutions:
|
||||
if num_head_channels == -1:
|
||||
dim_head = ch // num_heads
|
||||
else:
|
||||
num_heads = ch // num_head_channels
|
||||
dim_head = num_head_channels
|
||||
disabled_sa = disable_self_attentions[level] if exists(
|
||||
disable_self_attentions) else False
|
||||
|
||||
layers.append(
|
||||
SpatialTransformerV2(
|
||||
ch,
|
||||
num_heads,
|
||||
dim_head,
|
||||
transformer_block=transformer_block,
|
||||
depth=transformer_depth[level],
|
||||
context_dim=context_dim,
|
||||
disable_self_attn=disabled_sa,
|
||||
use_linear=use_linear_in_transformer,
|
||||
use_checkpoint=use_checkpoint))
|
||||
self.input_blocks.append(TimestepEmbedSequential(*layers))
|
||||
self._feature_size += ch
|
||||
input_block_chans.append(ch)
|
||||
input_down_flag.append(False)
|
||||
if level != len(channel_mult) - 1:
|
||||
out_ch = ch
|
||||
self.input_blocks.append(
|
||||
TimestepEmbedSequential(
|
||||
ResBlock(
|
||||
ch,
|
||||
time_embed_dim,
|
||||
dropout,
|
||||
out_channels=out_ch,
|
||||
dims=dims,
|
||||
use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
down=True,
|
||||
) if resblock_updown else Downsample(
|
||||
ch, conv_resample, dims=dims, out_channels=out_ch))
|
||||
)
|
||||
ch = out_ch
|
||||
input_block_chans.append(ch)
|
||||
input_down_flag.append(True)
|
||||
ds *= 2
|
||||
self._feature_size += ch
|
||||
self._input_block_chans = copy.deepcopy(input_block_chans)
|
||||
self._input_down_flag = input_down_flag
|
||||
|
||||
if num_head_channels == -1:
|
||||
dim_head = ch // num_heads
|
||||
else:
|
||||
num_heads = ch // num_head_channels
|
||||
dim_head = num_head_channels
|
||||
self.middle_block = TimestepEmbedSequential(
|
||||
ResBlock(
|
||||
ch,
|
||||
time_embed_dim,
|
||||
dropout,
|
||||
dims=dims,
|
||||
use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
),
|
||||
SpatialTransformerV2(ch,
|
||||
num_heads,
|
||||
dim_head,
|
||||
transformer_block=transformer_block,
|
||||
depth=transformer_depth_middle,
|
||||
context_dim=context_dim,
|
||||
disable_self_attn=disable_middle_self_attn,
|
||||
use_linear=use_linear_in_transformer,
|
||||
use_checkpoint=use_checkpoint),
|
||||
ResBlock(
|
||||
ch,
|
||||
time_embed_dim,
|
||||
dropout,
|
||||
dims=dims,
|
||||
use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
),
|
||||
)
|
||||
self._feature_size += ch
|
||||
self._middle_block_chans = [ch]
|
||||
|
||||
self._output_block_chans = []
|
||||
self.output_blocks = nn.ModuleList([])
|
||||
for level, mult in list(enumerate(channel_mult))[::-1]:
|
||||
for i in range(self.num_res_blocks[level] + 1):
|
||||
ich = input_block_chans.pop()
|
||||
layers = [
|
||||
ResBlock(
|
||||
ch + ich,
|
||||
time_embed_dim,
|
||||
dropout,
|
||||
out_channels=model_channels * mult,
|
||||
dims=dims,
|
||||
use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
)
|
||||
]
|
||||
ch = model_channels * mult
|
||||
if ds in attention_resolutions:
|
||||
if num_head_channels == -1:
|
||||
dim_head = ch // num_heads
|
||||
else:
|
||||
num_heads = ch // num_head_channels
|
||||
dim_head = num_head_channels
|
||||
disabled_sa = disable_self_attentions[level] if exists(
|
||||
disable_self_attentions) else False
|
||||
layers.append(
|
||||
SpatialTransformerV2(
|
||||
ch,
|
||||
num_heads,
|
||||
dim_head,
|
||||
transformer_block=transformer_block,
|
||||
depth=transformer_depth[level],
|
||||
context_dim=context_dim,
|
||||
disable_self_attn=disabled_sa,
|
||||
use_linear=use_linear_in_transformer,
|
||||
use_checkpoint=use_checkpoint))
|
||||
if level and i == self.num_res_blocks[level]:
|
||||
out_ch = ch
|
||||
layers.append(
|
||||
ResBlock(
|
||||
ch,
|
||||
time_embed_dim,
|
||||
dropout,
|
||||
out_channels=out_ch,
|
||||
dims=dims,
|
||||
use_checkpoint=use_checkpoint,
|
||||
use_scale_shift_norm=use_scale_shift_norm,
|
||||
up=True,
|
||||
) if resblock_updown else Upsample(
|
||||
ch, conv_resample, dims=dims, out_channels=out_ch))
|
||||
ds //= 2
|
||||
|
||||
self.output_blocks.append(TimestepEmbedSequential(*layers))
|
||||
|
||||
self._feature_size += ch
|
||||
self._output_block_chans.append(ch)
|
||||
|
||||
self.out = nn.Sequential(
|
||||
normalization(ch),
|
||||
nn.SiLU(),
|
||||
zero_module(
|
||||
conv_nd(dims, model_channels, out_channels, 3, padding=1)),
|
||||
)
|
||||
|
||||
if self.use_refine:
|
||||
self.ref_time_embed = copy.deepcopy(self.time_embed)
|
||||
self.ref_label_emb = copy.deepcopy(self.label_emb)
|
||||
self.ref_input_blocks = copy.deepcopy(self.input_blocks)
|
||||
self.ref_input_blocks[0] = TimestepEmbedSequential(
|
||||
conv_nd(dims, 4, model_channels, 3, padding=1))
|
||||
self.ref_middle_block = copy.deepcopy(self.middle_block)
|
||||
self.ref_output_blocks = copy.deepcopy(self.output_blocks)
|
||||
|
||||
def load_pretrained_model(self, pretrained_model):
|
||||
if pretrained_model is not None:
|
||||
with FS.get_from(pretrained_model,
|
||||
wait_finish=True) as local_model:
|
||||
self.init_from_ckpt(local_model, ignore_keys=self.ignore_keys)
|
||||
|
||||
def init_from_ckpt(self, path, ignore_keys=list()):
|
||||
if path.endswith('safetensors'):
|
||||
from safetensors.torch import load_file as load_safetensors
|
||||
sd = load_safetensors(path)
|
||||
else:
|
||||
sd = torch.load(path, map_location='cpu')
|
||||
|
||||
new_sd = OrderedDict()
|
||||
for k, v in sd.items():
|
||||
ignored = False
|
||||
for ik in ignore_keys:
|
||||
if ik in k:
|
||||
if we.rank == 0:
|
||||
self.logger.info(
|
||||
'Ignore key {} from state_dict.'.format(k))
|
||||
ignored = True
|
||||
break
|
||||
if not ignored:
|
||||
if k == 'input_blocks.0.0.weight':
|
||||
if we.rank == 0:
|
||||
self.logger.info(
|
||||
'Partial initial key {} from state_dict.'.format(
|
||||
k))
|
||||
new_v = torch.empty(320, self.in_channels, 3, 3)
|
||||
nn.init.zeros_(new_v)
|
||||
new_v[:, :v.shape[1]] = v
|
||||
new_sd[k] = new_v
|
||||
if self.use_refine:
|
||||
new_sd['ref_' + k] = v
|
||||
else:
|
||||
new_sd[k] = v
|
||||
if self.use_refine:
|
||||
new_sd['ref_' + k] = v
|
||||
|
||||
missing, unexpected = self.load_state_dict(new_sd, strict=False)
|
||||
if we.rank == 0:
|
||||
self.logger.info(
|
||||
f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys'
|
||||
)
|
||||
if len(missing) > 0:
|
||||
self.logger.info(f'Missing Keys:\n {missing}')
|
||||
if len(unexpected) > 0:
|
||||
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
|
||||
|
||||
def forward(self, x, t=None, cond=dict(), **kwargs):
|
||||
t_emb = timestep_embedding(t,
|
||||
self.model_channels,
|
||||
repeat_only=False,
|
||||
legacy=True)
|
||||
emb = self.time_embed(t_emb)
|
||||
|
||||
if isinstance(cond, dict):
|
||||
if 'y' in cond:
|
||||
assert self.num_classes is not None
|
||||
emb = emb + self.label_emb(cond['y'])
|
||||
if self.use_refine:
|
||||
ref_emb = self.ref_time_embed(t_emb)
|
||||
assert 'null_y' in cond
|
||||
cond_y = cond['y'].clone()
|
||||
cond_y[:, :cond['null_y'].shape[1]] = cond['null_y']
|
||||
ref_emb = ref_emb + self.ref_label_emb(cond_y)
|
||||
|
||||
if 'concat' in cond:
|
||||
c = cond['concat']
|
||||
x = torch.cat([x, c], dim=1)
|
||||
|
||||
context = cond.get('crossattn', None)
|
||||
img_context = cond.get('img_crossattn', None)
|
||||
|
||||
task = cond['task']
|
||||
image_scale = cond.get('image_scale', self.image_scale)
|
||||
if 'Subject' in task and img_context is not None:
|
||||
ip_enc_scale = image_scale
|
||||
ip_dec_scale = image_scale
|
||||
num_img_tokens = img_context.shape[1]
|
||||
context = torch.cat([context, img_context], dim=1)
|
||||
else:
|
||||
ip_enc_scale = None
|
||||
ip_dec_scale = None
|
||||
num_img_tokens = None
|
||||
|
||||
ref = cond.get('ref_xt', None)
|
||||
ref_context = cond.get('ref_crossattn', None)
|
||||
else:
|
||||
raise TypeError
|
||||
|
||||
hs = []
|
||||
refs = []
|
||||
h = x
|
||||
|
||||
if self.use_refine:
|
||||
assert ref is not None and ref_context is not None
|
||||
for i, (ref_module, module) in enumerate(
|
||||
zip(self.ref_input_blocks, self.input_blocks)):
|
||||
ref = ref_module(ref, ref_emb, ref_context, caching=None)
|
||||
h = module(h,
|
||||
emb,
|
||||
context,
|
||||
caching=None,
|
||||
scale=ip_enc_scale,
|
||||
num_img_token=num_img_tokens)
|
||||
refs.append(ref)
|
||||
hs.append(h)
|
||||
|
||||
ref = self.ref_middle_block(ref,
|
||||
ref_emb,
|
||||
ref_context,
|
||||
caching=None)
|
||||
h = self.middle_block(h,
|
||||
emb,
|
||||
context,
|
||||
caching=None,
|
||||
scale=ip_enc_scale,
|
||||
num_img_token=num_img_tokens)
|
||||
|
||||
for i, (ref_module, module) in enumerate(
|
||||
zip(self.ref_output_blocks, self.output_blocks)):
|
||||
cache = []
|
||||
ref = torch.cat([ref, refs.pop()], dim=1)
|
||||
ref = ref_module(ref,
|
||||
ref_emb,
|
||||
ref_context,
|
||||
caching='write',
|
||||
cache=cache)
|
||||
h = torch.cat([h, hs.pop()], dim=1)
|
||||
h = module(h,
|
||||
emb,
|
||||
context,
|
||||
caching='read',
|
||||
cache=cache,
|
||||
scale=ip_dec_scale,
|
||||
num_img_token=num_img_tokens)
|
||||
else:
|
||||
for module in self.input_blocks:
|
||||
h = module(h,
|
||||
emb,
|
||||
context,
|
||||
caching=None,
|
||||
scale=ip_enc_scale,
|
||||
num_img_token=num_img_tokens)
|
||||
hs.append(h)
|
||||
h = self.middle_block(h,
|
||||
emb,
|
||||
context,
|
||||
caching=None,
|
||||
scale=ip_enc_scale,
|
||||
num_img_token=num_img_tokens)
|
||||
for module in self.output_blocks:
|
||||
h = torch.cat([h, hs.pop()], dim=1)
|
||||
h = module(h,
|
||||
emb,
|
||||
context,
|
||||
caching=None,
|
||||
scale=ip_dec_scale,
|
||||
num_img_token=num_img_tokens)
|
||||
|
||||
out = self.out(h)
|
||||
return out
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('BACKBONE',
|
||||
__class__.__name__,
|
||||
LargenUNetXL.para_dict,
|
||||
set_name=True)
|
||||
|
||||
@@ -10,7 +10,9 @@ import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms.functional as TF
|
||||
from einops import rearrange, repeat
|
||||
from packaging import version
|
||||
|
||||
from scepter.modules.model.utils.basic_utils import checkpoint, default, exists
|
||||
|
||||
@@ -24,6 +26,12 @@ except Exception as e:
|
||||
|
||||
if find_loader('flash_attn'):
|
||||
FLASH_ATTN_IS_AVAILABLE = True
|
||||
import flash_attn
|
||||
if (not hasattr(flash_attn, '__version__')) or (version.parse(
|
||||
flash_attn.__version__) < version.parse('2.0')):
|
||||
from flash_attn.flash_attn_interface import flash_attn_unpadded_kvpacked_func
|
||||
else:
|
||||
from flash_attn.flash_attn_interface import flash_attn_varlen_kvpacked_func as flash_attn_unpadded_kvpacked_func
|
||||
else:
|
||||
FLASH_ATTN_IS_AVAILABLE = False
|
||||
|
||||
@@ -164,12 +172,14 @@ class TimestepEmbedSequential(nn.Sequential, TimestepBlock):
|
||||
A sequential module that passes timestep embeddings to the children that
|
||||
support it as an extra input.
|
||||
"""
|
||||
def forward(self, x, emb, context=None, target_size=None):
|
||||
def forward(self, x, emb, context=None, target_size=None, **kwargs):
|
||||
for layer in self:
|
||||
if isinstance(layer, TimestepBlock):
|
||||
x = layer(x, emb)
|
||||
elif isinstance(layer, SpatialTransformer):
|
||||
x = layer(x, context)
|
||||
elif isinstance(layer, SpatialTransformerV2):
|
||||
x = layer(x, context, **kwargs)
|
||||
elif isinstance(layer, Upsample):
|
||||
x = layer(x, target_size)
|
||||
else:
|
||||
@@ -607,8 +617,6 @@ class FlashattnMultiHeadAttention(nn.Module):
|
||||
and self.head_dim % 8 == 0 and self.head_dim <= 128
|
||||
and self.flash_dtype is not None):
|
||||
# flash implementation
|
||||
from flash_attn.flash_attn_interface import \
|
||||
flash_attn_unpadded_kvpacked_func
|
||||
dtype = q.dtype
|
||||
if dtype != self.flash_dtype:
|
||||
q = q.type(self.flash_dtype)
|
||||
@@ -859,6 +867,92 @@ class MemoryEfficientCrossAttention(nn.Module):
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class XFormersMHA_IP(nn.Module):
|
||||
def __init__(self,
|
||||
query_dim,
|
||||
context_dim=None,
|
||||
heads=8,
|
||||
dim_head=64,
|
||||
dropout=0.0):
|
||||
super().__init__()
|
||||
inner_dim = dim_head * heads
|
||||
context_dim = default(context_dim, query_dim)
|
||||
|
||||
self.heads = heads
|
||||
self.dim_head = dim_head
|
||||
|
||||
self.to_q = nn.Linear(query_dim, inner_dim, bias=False)
|
||||
self.to_k = nn.Linear(context_dim, inner_dim, bias=False)
|
||||
self.to_v = nn.Linear(context_dim, inner_dim, bias=False)
|
||||
self.to_k_ip = nn.Linear(context_dim, inner_dim, bias=False)
|
||||
self.to_v_ip = nn.Linear(context_dim, inner_dim, bias=False)
|
||||
|
||||
self.to_out = nn.Sequential(nn.Linear(inner_dim, query_dim),
|
||||
nn.Dropout(dropout))
|
||||
self.attention_op = None
|
||||
|
||||
def forward(self,
|
||||
x,
|
||||
context=None,
|
||||
mask=None,
|
||||
scale=None,
|
||||
num_img_token=None):
|
||||
q = self.to_q(x)
|
||||
context = default(context, x)
|
||||
|
||||
if scale is not None and num_img_token is not None:
|
||||
eos = context.shape[1] - num_img_token
|
||||
txt_context = context[:, :eos, :]
|
||||
img_context = context[:, eos:, :]
|
||||
|
||||
k = self.to_k(txt_context)
|
||||
v = self.to_v(txt_context)
|
||||
k_i = self.to_k_ip(img_context)
|
||||
v_i = self.to_v_ip(img_context)
|
||||
|
||||
b, _, _ = q.shape
|
||||
q, k, v, k_i, v_i = map(
|
||||
lambda t: t.unsqueeze(3).reshape(b, t.shape[
|
||||
1], self.heads, self.dim_head).permute(0, 2, 1, 3).reshape(
|
||||
b * self.heads, t.shape[1], self.dim_head).contiguous(
|
||||
),
|
||||
(q, k, v, k_i, v_i),
|
||||
)
|
||||
|
||||
# actually compute the attention, what we cannot get enough of
|
||||
txt_out = xformers.ops.memory_efficient_attention(
|
||||
q, k, v, attn_bias=None, op=self.attention_op)
|
||||
img_out = xformers.ops.memory_efficient_attention(
|
||||
q, k_i, v_i, attn_bias=None, op=self.attention_op)
|
||||
out = txt_out + scale * img_out
|
||||
else:
|
||||
k = self.to_k(context)
|
||||
v = self.to_v(context)
|
||||
b, _, _ = q.shape
|
||||
q, k, v = map(
|
||||
lambda t: t.unsqueeze(3).reshape(b, t.shape[
|
||||
1], self.heads, self.dim_head).permute(0, 2, 1, 3).reshape(
|
||||
b * self.heads, t.shape[1], self.dim_head).contiguous(
|
||||
),
|
||||
(q, k, v),
|
||||
)
|
||||
out = xformers.ops.memory_efficient_attention(q,
|
||||
k,
|
||||
v,
|
||||
attn_bias=None,
|
||||
op=self.attention_op)
|
||||
|
||||
# TODO: Use this directly in the attention operation, as a bias
|
||||
if exists(mask):
|
||||
raise NotImplementedError
|
||||
out = (out.unsqueeze(0).reshape(
|
||||
b, self.heads, out.shape[1],
|
||||
self.dim_head).permute(0, 2, 1,
|
||||
3).reshape(b, out.shape[1],
|
||||
self.heads * self.dim_head))
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class BasicTransformerBlock(nn.Module):
|
||||
def __init__(self,
|
||||
dim,
|
||||
@@ -903,6 +997,65 @@ class BasicTransformerBlock(nn.Module):
|
||||
return x
|
||||
|
||||
|
||||
class TransformerBlockV2(nn.Module):
|
||||
def __init__(self,
|
||||
query_dim,
|
||||
n_heads,
|
||||
d_head,
|
||||
dropout=0.,
|
||||
context_dim=None,
|
||||
gated_ff=True,
|
||||
use_checkpoint=False,
|
||||
disable_self_attn=False):
|
||||
super().__init__()
|
||||
self.disable_self_attn = disable_self_attn
|
||||
self.attn1 = MemoryEfficientCrossAttention(query_dim=query_dim,
|
||||
heads=n_heads,
|
||||
dim_head=d_head,
|
||||
dropout=dropout,
|
||||
context_dim=None)
|
||||
self.ff = FeedForward(query_dim, dropout=dropout, glu=gated_ff)
|
||||
self.attn2 = XFormersMHA_IP(query_dim=query_dim,
|
||||
heads=n_heads,
|
||||
dim_head=d_head,
|
||||
context_dim=context_dim)
|
||||
self.norm1 = nn.LayerNorm(query_dim)
|
||||
self.norm2 = nn.LayerNorm(query_dim)
|
||||
self.norm3 = nn.LayerNorm(query_dim)
|
||||
self.use_checkpoint = use_checkpoint
|
||||
|
||||
def forward(self,
|
||||
x,
|
||||
context,
|
||||
caching=None,
|
||||
cache=None,
|
||||
scale=None,
|
||||
num_img_token=None,
|
||||
**kwargs):
|
||||
y = self.norm1(x)
|
||||
if caching == 'write':
|
||||
assert isinstance(cache, list)
|
||||
cache.append(y)
|
||||
x = self.attn1(y, context=None) + x
|
||||
elif caching == 'read':
|
||||
assert isinstance(cache, list) and len(cache) > 0
|
||||
c = cache.pop(0)
|
||||
self_ctx = torch.cat([y, c], dim=1)
|
||||
x = self.attn1(y, context=self_ctx) + x
|
||||
elif caching is None:
|
||||
x = self.attn1(y, context=None) + x
|
||||
else:
|
||||
assert False
|
||||
|
||||
x = self.attn2(self.norm2(x),
|
||||
context=context,
|
||||
scale=scale,
|
||||
num_img_token=num_img_token) + x
|
||||
x = self.ff(self.norm3(x)) + x
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class SpatialTransformer(nn.Module):
|
||||
"""
|
||||
Transformer block for image-like data.
|
||||
@@ -998,3 +1151,108 @@ class SpatialTransformer(nn.Module):
|
||||
if not self.use_linear:
|
||||
x = self.proj_out(x)
|
||||
return x + x_in
|
||||
|
||||
|
||||
class SpatialTransformerV2(nn.Module):
|
||||
"""
|
||||
Transformer block for image-like data.
|
||||
First, project the input (aka embedding)
|
||||
and reshape to b, t, d.
|
||||
Then apply standard transformer action.
|
||||
Finally, reshape to image
|
||||
NEW: use_linear for more efficiency instead of the 1x1 convs
|
||||
"""
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
n_heads,
|
||||
d_head,
|
||||
transformer_block,
|
||||
depth=1,
|
||||
dropout=0.,
|
||||
context_dim=None,
|
||||
disable_self_attn=False,
|
||||
use_linear=False,
|
||||
use_checkpoint=True):
|
||||
super().__init__()
|
||||
if exists(context_dim) and not isinstance(context_dim, list):
|
||||
context_dim = [context_dim]
|
||||
|
||||
if exists(context_dim) and not isinstance(context_dim, (list)):
|
||||
context_dim = [context_dim]
|
||||
if exists(context_dim) and isinstance(context_dim, list):
|
||||
if depth != len(context_dim):
|
||||
print(
|
||||
f'WARNING: {self.__class__.__name__}: Found context dims {context_dim} of'
|
||||
f" depth {len(context_dim)}, which does not match the specified 'depth' of"
|
||||
f' {depth}. Setting context_dim to {depth * [context_dim[0]]} now.'
|
||||
)
|
||||
# depth does not match context dims.
|
||||
assert all(
|
||||
map(lambda x: x == context_dim[0], context_dim)
|
||||
), 'need homogenous context_dim to match depth automatically'
|
||||
context_dim = depth * [context_dim[0]]
|
||||
elif context_dim is None:
|
||||
context_dim = [None] * depth
|
||||
|
||||
self.in_channels = in_channels
|
||||
inner_dim = n_heads * d_head
|
||||
self.norm = normalization(in_channels)
|
||||
if not use_linear:
|
||||
self.proj_in = nn.Conv2d(in_channels,
|
||||
inner_dim,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0)
|
||||
else:
|
||||
self.proj_in = nn.Linear(in_channels, inner_dim)
|
||||
|
||||
self.transformer_blocks = nn.ModuleList([
|
||||
transformer_block(inner_dim,
|
||||
n_heads,
|
||||
d_head,
|
||||
dropout=dropout,
|
||||
context_dim=context_dim[d],
|
||||
disable_self_attn=disable_self_attn,
|
||||
use_checkpoint=use_checkpoint)
|
||||
for d in range(depth)
|
||||
])
|
||||
if not use_linear:
|
||||
self.proj_out = zero_module(
|
||||
nn.Conv2d(inner_dim,
|
||||
in_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0))
|
||||
else:
|
||||
self.proj_out = zero_module(nn.Linear(in_channels, inner_dim))
|
||||
self.use_linear = use_linear
|
||||
|
||||
def forward(self, x, context=None, **kwargs):
|
||||
# note: if no context is given, cross-attention defaults to self-attention
|
||||
if not isinstance(context, list):
|
||||
context = [context]
|
||||
b, c, h, w = x.shape
|
||||
|
||||
ref_mask = kwargs.pop('ref_mask', None)
|
||||
if ref_mask is not None:
|
||||
ref_mask = TF.resize(ref_mask, (h, w), antialias=True)
|
||||
ref_mask = (ref_mask > 0.5).float()
|
||||
ref_mask = rearrange(ref_mask, 'b c h w -> b (h w) c').contiguous()
|
||||
|
||||
x_in = x
|
||||
x = self.norm(x)
|
||||
if not self.use_linear:
|
||||
x = self.proj_in(x)
|
||||
x = rearrange(x, 'b c h w -> b (h w) c').contiguous()
|
||||
if self.use_linear:
|
||||
x = self.proj_in(x)
|
||||
for i, block in enumerate(self.transformer_blocks):
|
||||
if i > 0 and len(context) == 1:
|
||||
i = 0 # use same context for each block
|
||||
x = block(x, context=context[i], ref_mask=ref_mask, **kwargs)
|
||||
if self.use_linear:
|
||||
x = self.proj_out(x)
|
||||
x = rearrange(x, 'b (h w) c -> b c h w', h=h, w=w).contiguous()
|
||||
if not self.use_linear:
|
||||
x = self.proj_out(x)
|
||||
return x + x_in
|
||||
|
||||
@@ -1,17 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional
|
||||
from einops import rearrange
|
||||
|
||||
from scepter.modules.model.backbone.video.init_helper import (
|
||||
_init_transformer_weights, trunc_normal_)
|
||||
from scepter.modules.model.registry import BACKBONES, BRICKS, STEMS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
'''
|
||||
The implementations of vivit as https://arxiv.org/abs/2103.15691.
|
||||
The following setting alined the proposed model in the paper above.
|
||||
@@ -39,6 +27,18 @@ TimesFormer:
|
||||
complexity: (n_h * n_w) ** 2 + O(attn_temp)
|
||||
'''
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional
|
||||
from einops import rearrange
|
||||
|
||||
from scepter.modules.model.backbone.video.init_helper import (
|
||||
_init_transformer_weights, trunc_normal_)
|
||||
from scepter.modules.model.registry import BACKBONES, BRICKS, STEMS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
|
||||
|
||||
@BACKBONES.register_class()
|
||||
class VideoTransformer(nn.Module):
|
||||
|
||||
@@ -1,7 +1,10 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
from scepter.modules.model.embedder.embedder import (ConcatTimestepEmbedderND,
|
||||
FrozenCLIPEmbedder,
|
||||
FrozenOpenCLIPEmbedder,
|
||||
FrozenOpenCLIPEmbedder2,
|
||||
GeneralConditioner)
|
||||
GeneralConditioner,
|
||||
IPAdapterPlusEmbedder,
|
||||
RefCrossEmbedder)
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import warnings
|
||||
from collections import OrderedDict
|
||||
from contextlib import nullcontext
|
||||
from typing import Dict
|
||||
@@ -11,7 +12,6 @@ import torch.nn as nn
|
||||
import torch.utils.dlpack
|
||||
from einops import rearrange
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
from transformers import CLIPTextModel, CLIPTokenizer
|
||||
|
||||
# to check
|
||||
from scepter.modules.model.backbone.unet.unet_utils import Timestep
|
||||
@@ -22,6 +22,13 @@ from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
from .base_embedder import BaseEmbedder
|
||||
from .resampler import Resampler
|
||||
|
||||
try:
|
||||
from transformers import CLIPTextModel, CLIPTokenizer, CLIPVisionModelWithProjection
|
||||
except Exception as e:
|
||||
warnings.warn(
|
||||
f'Import transformers error, please deal with this problem: {e}')
|
||||
|
||||
|
||||
def autocast(f, enabled=True):
|
||||
@@ -507,6 +514,93 @@ class ConcatTimestepEmbedderND(BaseEmbedder):
|
||||
set_name=True)
|
||||
|
||||
|
||||
@EMBEDDERS.register_class()
|
||||
class IPAdapterPlusEmbedder(BaseEmbedder):
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
|
||||
with FS.get_dir_to_local_dir(cfg.CLIP_DIR,
|
||||
wait_finish=True) as local_path:
|
||||
self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(
|
||||
local_path)
|
||||
|
||||
self.image_proj_model = Resampler(
|
||||
dim=self.cfg.get('IN_DIM', 768),
|
||||
depth=self.cfg.get('DEPTH', 4),
|
||||
dim_head=64,
|
||||
heads=self.cfg.get('HEADS', 12),
|
||||
num_queries=self.cfg.get('NUM_TOKENS', 16),
|
||||
embedding_dim=self.image_encoder.config.hidden_size,
|
||||
output_dim=self.cfg.get('CROSSATTN_DIM', 768),
|
||||
ff_mult=4,
|
||||
)
|
||||
|
||||
with FS.get_from(cfg.PRETRAINED_MODEL, wait_finish=True) as local_path:
|
||||
ckpt = torch.load(local_path, map_location='cpu')
|
||||
self.image_proj_model.load_state_dict(ckpt['image_proj'],
|
||||
strict=True)
|
||||
|
||||
self.patch_projector = nn.Linear(self.image_encoder.config.hidden_size,
|
||||
self.cfg.get('CROSSATTN_DIM', 768))
|
||||
|
||||
def encode(self, ref_ip, ref_detail):
|
||||
encoder_output = self.image_encoder(ref_ip, output_hidden_states=True)
|
||||
image_prompt_embeds = self.image_proj_model(
|
||||
encoder_output.hidden_states[-2])
|
||||
encoder_output_2 = self.image_encoder(ref_detail,
|
||||
output_hidden_states=True)
|
||||
image_patch_embeds = self.patch_projector(
|
||||
encoder_output_2.last_hidden_state)
|
||||
out = {
|
||||
'img_crossattn': image_prompt_embeds,
|
||||
'ref_crossattn': image_patch_embeds,
|
||||
}
|
||||
return out
|
||||
|
||||
def forward(self, ref_ip, ref_detail):
|
||||
return self.encode(ref_ip, ref_detail)
|
||||
|
||||
|
||||
class RefCrossEmbedder(BaseEmbedder):
|
||||
def __init__(self, cfg, logger=None):
|
||||
super().__init__(cfg, logger=logger)
|
||||
|
||||
with FS.get_dir_to_local_dir(cfg.CLIP_DIR,
|
||||
wait_finish=True) as local_path:
|
||||
self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(
|
||||
local_path)
|
||||
|
||||
self.patch_projector = nn.Linear(self.image_encoder.config.hidden_size,
|
||||
self.cfg.get('CROSSATTN_DIM', 768))
|
||||
|
||||
def encode(self, img):
|
||||
encoder_output = self.image_encoder(img, output_hidden_states=True)
|
||||
image_patch_embeds = self.patch_projector(
|
||||
encoder_output.last_hidden_state)
|
||||
out = {
|
||||
'ref_crossattn': image_patch_embeds,
|
||||
}
|
||||
return out
|
||||
|
||||
def forward(self, img):
|
||||
return self.encode(img)
|
||||
|
||||
|
||||
@EMBEDDERS.register_class()
|
||||
class TransparentEmbedder(BaseEmbedder):
|
||||
def forward(self, *args):
|
||||
out = dict()
|
||||
for key, val in zip(self.input_keys, args):
|
||||
out[key] = val
|
||||
return out
|
||||
|
||||
|
||||
@EMBEDDERS.register_class()
|
||||
class NoiseConcatEmbedder(BaseEmbedder):
|
||||
def forward(self, *args):
|
||||
return {'concat': torch.cat(args, dim=1)}
|
||||
|
||||
|
||||
@EMBEDDERS.register_class()
|
||||
class GeneralConditioner(BaseEmbedder):
|
||||
OUTPUT_DIM2KEYS = {2: 'y', 3: 'crossattn', 4: 'concat', 5: 'concat'}
|
||||
@@ -592,42 +686,58 @@ class GeneralConditioner(BaseEmbedder):
|
||||
with embedding_context():
|
||||
if hasattr(embedder, 'input_key') and (embedder.input_key
|
||||
is not None):
|
||||
if embedder.input_key not in batch:
|
||||
continue
|
||||
if embedder.legacy_ucg_val is not None:
|
||||
batch = self.possibly_get_ucg_val(embedder, batch)
|
||||
emb_out = embedder(batch[embedder.input_key])
|
||||
elif hasattr(embedder, 'input_keys'):
|
||||
if any([k not in batch for k in embedder.input_keys]):
|
||||
continue
|
||||
emb_out = embedder(
|
||||
*[batch[k] for k in embedder.input_keys])
|
||||
assert isinstance(
|
||||
emb_out, (torch.Tensor, list, tuple)
|
||||
), f'encoder outputs must be tensors or a sequence, but got {type(emb_out)}'
|
||||
if not isinstance(emb_out, (list, tuple)):
|
||||
emb_out = [emb_out]
|
||||
for emb in emb_out:
|
||||
# print("emb.shape", emb.shape)
|
||||
# print("emb.input_keys", embedder.input_keys)
|
||||
out_key = self.OUTPUT_DIM2KEYS[emb.dim()]
|
||||
if embedder.ucg_rate > 0.0 and embedder.legacy_ucg_val is None:
|
||||
emb = (expand_dims_like(
|
||||
torch.bernoulli(
|
||||
(1.0 - embedder.ucg_rate) *
|
||||
torch.ones(emb.shape[0], device=emb.device)),
|
||||
emb,
|
||||
) * emb)
|
||||
if (hasattr(embedder, 'input_keys')):
|
||||
if np.sum(
|
||||
np.array([
|
||||
key in force_zero_embeddings
|
||||
for key in embedder.input_keys
|
||||
])) > 0:
|
||||
emb = torch.zeros_like(emb)
|
||||
if out_key in output:
|
||||
output[out_key] = torch.cat((output[out_key], emb),
|
||||
self.KEY2CATDIM[out_key])
|
||||
else:
|
||||
output[out_key] = emb
|
||||
# if "y" in output:
|
||||
# print("out.shape", output["y"].shape)
|
||||
|
||||
if isinstance(emb_out, dict):
|
||||
for key, val in emb_out.items():
|
||||
if key in output:
|
||||
assert key in self.KEY2CATDIM
|
||||
output[key] = torch.cat([output[key], val],
|
||||
dim=self.KEY2CATDIM[key])
|
||||
else:
|
||||
output[key] = val
|
||||
else:
|
||||
assert isinstance(
|
||||
emb_out, (torch.Tensor, list, tuple)
|
||||
), f'encoder outputs must be tensors or a sequence, but got {type(emb_out)}'
|
||||
|
||||
if not isinstance(emb_out, (list, tuple)):
|
||||
emb_out = [emb_out]
|
||||
|
||||
for emb in emb_out:
|
||||
out_key = self.OUTPUT_DIM2KEYS[emb.dim()]
|
||||
|
||||
if embedder.ucg_rate > 0.0 and embedder.legacy_ucg_val is None:
|
||||
emb = (expand_dims_like(
|
||||
torch.bernoulli(
|
||||
(1.0 - embedder.ucg_rate) *
|
||||
torch.ones(emb.shape[0], device=emb.device)),
|
||||
emb,
|
||||
) * emb)
|
||||
|
||||
if (hasattr(embedder, 'input_keys')):
|
||||
if np.sum(
|
||||
np.array([
|
||||
key in force_zero_embeddings
|
||||
for key in embedder.input_keys
|
||||
])) > 0:
|
||||
emb = torch.zeros_like(emb)
|
||||
|
||||
if out_key in output:
|
||||
output[out_key] = torch.cat((output[out_key], emb),
|
||||
self.KEY2CATDIM[out_key])
|
||||
else:
|
||||
output[out_key] = emb
|
||||
|
||||
return output
|
||||
|
||||
def get_unconditional_conditioning(self,
|
||||
|
||||
@@ -0,0 +1,160 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
from einops.layers.torch import Rearrange
|
||||
|
||||
|
||||
# FFN
|
||||
def FeedForward(dim, mult=4):
|
||||
inner_dim = int(dim * mult)
|
||||
return nn.Sequential(
|
||||
nn.LayerNorm(dim),
|
||||
nn.Linear(dim, inner_dim, bias=False),
|
||||
nn.GELU(),
|
||||
nn.Linear(inner_dim, dim, bias=False),
|
||||
)
|
||||
|
||||
|
||||
def reshape_tensor(x, heads):
|
||||
bs, length, width = x.shape
|
||||
# (bs, length, width) --> (bs, length, n_heads, dim_per_head)
|
||||
x = x.view(bs, length, heads, -1)
|
||||
# (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head)
|
||||
x = x.transpose(1, 2)
|
||||
# (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head)
|
||||
x = x.reshape(bs, heads, length, -1)
|
||||
return x
|
||||
|
||||
|
||||
class PerceiverAttention(nn.Module):
|
||||
def __init__(self, *, dim, dim_head=64, heads=8):
|
||||
super().__init__()
|
||||
self.scale = dim_head**-0.5
|
||||
self.dim_head = dim_head
|
||||
self.heads = heads
|
||||
inner_dim = dim_head * heads
|
||||
|
||||
self.norm1 = nn.LayerNorm(dim)
|
||||
self.norm2 = nn.LayerNorm(dim)
|
||||
|
||||
self.to_q = nn.Linear(dim, inner_dim, bias=False)
|
||||
self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False)
|
||||
self.to_out = nn.Linear(inner_dim, dim, bias=False)
|
||||
|
||||
def forward(self, x, latents):
|
||||
"""
|
||||
Args:
|
||||
x (torch.Tensor): image features
|
||||
shape (b, n1, D)
|
||||
latent (torch.Tensor): latent features
|
||||
shape (b, n2, D)
|
||||
"""
|
||||
x = self.norm1(x)
|
||||
latents = self.norm2(latents)
|
||||
|
||||
b, l, _ = latents.shape
|
||||
|
||||
q = self.to_q(latents)
|
||||
kv_input = torch.cat((x, latents), dim=-2)
|
||||
k, v = self.to_kv(kv_input).chunk(2, dim=-1)
|
||||
|
||||
q = reshape_tensor(q, self.heads)
|
||||
k = reshape_tensor(k, self.heads)
|
||||
v = reshape_tensor(v, self.heads)
|
||||
|
||||
# attention
|
||||
scale = 1 / math.sqrt(math.sqrt(self.dim_head))
|
||||
weight = (q * scale) @ (k * scale).transpose(
|
||||
-2, -1) # More stable with f16 than dividing afterwards
|
||||
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
|
||||
out = weight @ v
|
||||
|
||||
out = out.permute(0, 2, 1, 3).reshape(b, l, -1)
|
||||
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class Resampler(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim=1024,
|
||||
depth=8,
|
||||
dim_head=64,
|
||||
heads=16,
|
||||
num_queries=8,
|
||||
embedding_dim=768,
|
||||
output_dim=1024,
|
||||
ff_mult=4,
|
||||
max_seq_len: int = 257, # CLIP tokens + CLS token
|
||||
apply_pos_emb: bool = False,
|
||||
num_latents_mean_pooled:
|
||||
int = 0, # number of latents derived from mean pooled representation of the sequence
|
||||
):
|
||||
super().__init__()
|
||||
self.pos_emb = nn.Embedding(max_seq_len,
|
||||
embedding_dim) if apply_pos_emb else None
|
||||
|
||||
self.latents = nn.Parameter(
|
||||
torch.randn(1, num_queries, dim) / dim**0.5)
|
||||
|
||||
self.proj_in = nn.Linear(embedding_dim, dim)
|
||||
|
||||
self.proj_out = nn.Linear(dim, output_dim)
|
||||
self.norm_out = nn.LayerNorm(output_dim)
|
||||
|
||||
self.to_latents_from_mean_pooled_seq = (nn.Sequential(
|
||||
nn.LayerNorm(dim),
|
||||
nn.Linear(dim, dim * num_latents_mean_pooled),
|
||||
Rearrange('b (n d) -> b n d', n=num_latents_mean_pooled),
|
||||
) if num_latents_mean_pooled > 0 else None)
|
||||
|
||||
self.layers = nn.ModuleList([])
|
||||
for _ in range(depth):
|
||||
self.layers.append(
|
||||
nn.ModuleList([
|
||||
PerceiverAttention(dim=dim, dim_head=dim_head,
|
||||
heads=heads),
|
||||
FeedForward(dim=dim, mult=ff_mult),
|
||||
]))
|
||||
|
||||
def forward(self, x):
|
||||
if self.pos_emb is not None:
|
||||
n, device = x.shape[1], x.device
|
||||
pos_emb = self.pos_emb(torch.arange(n, device=device))
|
||||
x = x + pos_emb
|
||||
|
||||
latents = self.latents.repeat(x.size(0), 1, 1)
|
||||
|
||||
x = self.proj_in(x)
|
||||
|
||||
if self.to_latents_from_mean_pooled_seq:
|
||||
meanpooled_seq = masked_mean(x,
|
||||
dim=1,
|
||||
mask=torch.ones(x.shape[:2],
|
||||
device=x.device,
|
||||
dtype=torch.bool))
|
||||
meanpooled_latents = self.to_latents_from_mean_pooled_seq(
|
||||
meanpooled_seq)
|
||||
latents = torch.cat((meanpooled_latents, latents), dim=-2)
|
||||
|
||||
for attn, ff in self.layers:
|
||||
latents = attn(x, latents) + latents
|
||||
latents = ff(latents) + latents
|
||||
|
||||
latents = self.proj_out(latents)
|
||||
return self.norm_out(latents)
|
||||
|
||||
|
||||
def masked_mean(t, *, dim, mask=None):
|
||||
if mask is None:
|
||||
return t.mean(dim=dim)
|
||||
|
||||
denom = mask.sum(dim=dim, keepdim=True)
|
||||
mask = rearrange(mask, 'b n -> b n 1')
|
||||
masked_t = t.masked_fill(~mask, 0.0)
|
||||
|
||||
return masked_t.sum(dim=dim) / denom.clamp(min=1e-5)
|
||||
@@ -1,5 +1,8 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.head.classifier_head import (
|
||||
ClassifierHead, CosineLinearHead, TransformerHead, TransformerHeadx2,
|
||||
VideoClassifierHead, VideoClassifierHeadx2)
|
||||
from scepter.modules.model.head.classifier_head import (ClassifierHead,
|
||||
CosineLinearHead,
|
||||
TransformerHead,
|
||||
TransformerHeadx2,
|
||||
VideoClassifierHead,
|
||||
VideoClassifierHeadx2)
|
||||
|
||||
@@ -52,7 +52,7 @@ class DiagonalGaussianDistribution(object):
|
||||
dim=dims)
|
||||
|
||||
def mode(self):
|
||||
print('*** use DiagonalGaussianDistribution.mode() ***')
|
||||
# print('*** use DiagonalGaussianDistribution.mode() ***')
|
||||
return self.mean
|
||||
|
||||
|
||||
|
||||
@@ -13,8 +13,7 @@ from .schedules import karras_schedule
|
||||
from .solvers import (sample_ddim, sample_dpm_2, sample_dpm_2_ancestral,
|
||||
sample_dpmpp_2m, sample_dpmpp_2m_sde,
|
||||
sample_dpmpp_2s_ancestral, sample_dpmpp_sde,
|
||||
sample_euler, sample_euler_ancestral, sample_heun,
|
||||
sample_img2img_euler, sample_img2img_euler_ancestral)
|
||||
sample_euler, sample_euler_ancestral, sample_heun)
|
||||
|
||||
__all__ = ['GaussianDiffusion']
|
||||
|
||||
@@ -27,6 +26,148 @@ def _i(tensor, t, x):
|
||||
return tensor[t.to(tensor.device)].view(shape).to(x.device)
|
||||
|
||||
|
||||
def _unpack_2d_ks(kernel_size):
|
||||
if isinstance(kernel_size, int):
|
||||
ky = kx = kernel_size
|
||||
else:
|
||||
assert len(
|
||||
kernel_size) == 2, '2D Kernel size should have a length of 2.'
|
||||
ky, kx = kernel_size
|
||||
|
||||
ky = int(ky)
|
||||
kx = int(kx)
|
||||
return ky, kx
|
||||
|
||||
|
||||
def _compute_zero_padding(kernel_size):
|
||||
ky, kx = _unpack_2d_ks(kernel_size)
|
||||
return (ky - 1) // 2, (kx - 1) // 2
|
||||
|
||||
|
||||
def _bilateral_blur(
|
||||
input,
|
||||
guidance,
|
||||
kernel_size,
|
||||
sigma_color,
|
||||
sigma_space,
|
||||
border_type='reflect',
|
||||
color_distance_type='l1',
|
||||
):
|
||||
|
||||
if isinstance(sigma_color, torch.Tensor):
|
||||
sigma_color = sigma_color.to(device=input.device,
|
||||
dtype=input.dtype).view(-1, 1, 1, 1, 1)
|
||||
|
||||
ky, kx = _unpack_2d_ks(kernel_size)
|
||||
pad_y, pad_x = _compute_zero_padding(kernel_size)
|
||||
|
||||
padded_input = torch.nn.functional.pad(input, (pad_x, pad_x, pad_y, pad_y),
|
||||
mode=border_type)
|
||||
unfolded_input = padded_input.unfold(2, ky, 1).unfold(3, kx, 1).flatten(
|
||||
-2) # (B, C, H, W, Ky x Kx)
|
||||
|
||||
if guidance is None:
|
||||
guidance = input
|
||||
unfolded_guidance = unfolded_input
|
||||
else:
|
||||
padded_guidance = torch.nn.functional.pad(guidance,
|
||||
(pad_x, pad_x, pad_y, pad_y),
|
||||
mode=border_type)
|
||||
unfolded_guidance = padded_guidance.unfold(2, ky, 1).unfold(
|
||||
3, kx, 1).flatten(-2) # (B, C, H, W, Ky x Kx)
|
||||
|
||||
diff = unfolded_guidance - guidance.unsqueeze(-1)
|
||||
if color_distance_type == 'l1':
|
||||
color_distance_sq = diff.abs().sum(1, keepdim=True).square()
|
||||
elif color_distance_type == 'l2':
|
||||
color_distance_sq = diff.square().sum(1, keepdim=True)
|
||||
else:
|
||||
raise ValueError('color_distance_type only acceps l1 or l2')
|
||||
color_kernel = (-0.5 / sigma_color**2 *
|
||||
color_distance_sq).exp() # (B, 1, H, W, Ky x Kx)
|
||||
|
||||
space_kernel = get_gaussian_kernel2d(kernel_size,
|
||||
sigma_space,
|
||||
device=input.device,
|
||||
dtype=input.dtype)
|
||||
space_kernel = space_kernel.view(-1, 1, 1, 1, kx * ky)
|
||||
|
||||
kernel = space_kernel * color_kernel
|
||||
out = (unfolded_input * kernel).sum(-1) / kernel.sum(-1)
|
||||
return out
|
||||
|
||||
|
||||
def get_gaussian_kernel1d(
|
||||
kernel_size,
|
||||
sigma,
|
||||
force_even,
|
||||
*,
|
||||
device=None,
|
||||
dtype=None,
|
||||
):
|
||||
|
||||
return gaussian(kernel_size, sigma, device=device, dtype=dtype)
|
||||
|
||||
|
||||
def gaussian(window_size, sigma, *, device=None, dtype=None):
|
||||
|
||||
batch_size = sigma.shape[0]
|
||||
|
||||
x = (torch.arange(window_size, device=sigma.device, dtype=sigma.dtype) -
|
||||
window_size // 2).expand(batch_size, -1)
|
||||
|
||||
if window_size % 2 == 0:
|
||||
x = x + 0.5
|
||||
|
||||
gauss = torch.exp(-x.pow(2.0) / (2 * sigma.pow(2.0)))
|
||||
|
||||
return gauss / gauss.sum(-1, keepdim=True)
|
||||
|
||||
|
||||
def get_gaussian_kernel2d(
|
||||
kernel_size,
|
||||
sigma,
|
||||
force_even=False,
|
||||
*,
|
||||
device=None,
|
||||
dtype=None,
|
||||
):
|
||||
|
||||
sigma = torch.Tensor([[sigma, sigma]]).to(device=device, dtype=dtype)
|
||||
|
||||
ksize_y, ksize_x = _unpack_2d_ks(kernel_size)
|
||||
sigma_y, sigma_x = sigma[:, 0, None], sigma[:, 1, None]
|
||||
|
||||
kernel_y = get_gaussian_kernel1d(ksize_y,
|
||||
sigma_y,
|
||||
force_even,
|
||||
device=device,
|
||||
dtype=dtype)[..., None]
|
||||
kernel_x = get_gaussian_kernel1d(ksize_x,
|
||||
sigma_x,
|
||||
force_even,
|
||||
device=device,
|
||||
dtype=dtype)[..., None]
|
||||
|
||||
return kernel_y * kernel_x.view(-1, 1, ksize_x)
|
||||
|
||||
|
||||
def adaptive_anisotropic_filter(x, g=None):
|
||||
if g is None:
|
||||
g = x
|
||||
s, m = torch.std_mean(g, dim=(1, 2, 3), keepdim=True)
|
||||
s = s + 1e-5
|
||||
guidance = (g - m) / s
|
||||
y = _bilateral_blur(x,
|
||||
guidance,
|
||||
kernel_size=(13, 13),
|
||||
sigma_color=3.0,
|
||||
sigma_space=3.0,
|
||||
border_type='reflect',
|
||||
color_distance_type='l1')
|
||||
return y
|
||||
|
||||
|
||||
class GaussianDiffusion(object):
|
||||
def __init__(self, sigmas, prediction_type='eps'):
|
||||
assert prediction_type in {'x0', 'eps', 'v'}
|
||||
@@ -53,8 +194,10 @@ class GaussianDiffusion(object):
|
||||
guide_scale=None,
|
||||
guide_rescale=None,
|
||||
clamp=None,
|
||||
sharpness=0.0,
|
||||
percentile=None,
|
||||
cat_uc=False):
|
||||
cat_uc=False,
|
||||
**kwargs):
|
||||
"""
|
||||
Apply one step of denoising from the posterior distribution q(x_s | x_t, x0).
|
||||
Since x0 is not available, estimate the denoising results using the learned
|
||||
@@ -78,53 +221,99 @@ class GaussianDiffusion(object):
|
||||
|
||||
# prediction
|
||||
if guide_scale is None:
|
||||
assert isinstance(model_kwargs, dict)
|
||||
out = model(xt, t=t, **model_kwargs)
|
||||
if isinstance(model_kwargs, dict):
|
||||
out = model(xt, t=t, **model_kwargs, **kwargs)
|
||||
elif isinstance(model_kwargs, list) and len(model_kwargs) > 0:
|
||||
out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
||||
else:
|
||||
raise Exception('Error')
|
||||
else:
|
||||
# classifier-free guidance (arXiv:2207.12598)
|
||||
# model_kwargs[0]: conditional kwargs
|
||||
# model_kwargs[1]: non-conditional kwargs
|
||||
assert isinstance(model_kwargs, list) and len(model_kwargs) == 2
|
||||
|
||||
if guide_scale == 1.:
|
||||
out = model(xt, t=t, **model_kwargs[0])
|
||||
else:
|
||||
if cat_uc:
|
||||
|
||||
def parse_model_kwargs(prev_value, value):
|
||||
if isinstance(value, torch.Tensor):
|
||||
prev_value = torch.cat([prev_value, value], dim=0)
|
||||
elif isinstance(value, dict):
|
||||
for k, v in value.items():
|
||||
prev_value[k] = parse_model_kwargs(
|
||||
prev_value[k], v)
|
||||
elif isinstance(value, list):
|
||||
for idx, v in enumerate(value):
|
||||
prev_value[idx] = parse_model_kwargs(
|
||||
prev_value[idx], v)
|
||||
return prev_value
|
||||
|
||||
all_model_kwargs = copy.deepcopy(model_kwargs[0])
|
||||
for model_kwarg in model_kwargs[1:]:
|
||||
for key, value in model_kwarg.items():
|
||||
all_model_kwargs[key] = parse_model_kwargs(
|
||||
all_model_kwargs[key], value)
|
||||
all_out = model(xt.repeat(2, 1, 1, 1),
|
||||
t=t.repeat(2),
|
||||
**all_model_kwargs)
|
||||
y_out, u_out = all_out.chunk(2)
|
||||
assert isinstance(model_kwargs, list) and len(model_kwargs) >= 2
|
||||
if isinstance(guide_scale, float) or isinstance(guide_scale, int):
|
||||
assert len(model_kwargs) == 2
|
||||
if guide_scale == 1.:
|
||||
out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
||||
else:
|
||||
y_out = model(xt, t=t, **model_kwargs[0])
|
||||
u_out = model(xt, t=t, **model_kwargs[1])
|
||||
out = u_out + guide_scale * (y_out - u_out)
|
||||
if cat_uc:
|
||||
|
||||
# rescale the output according to arXiv:2305.08891
|
||||
if guide_rescale is not None:
|
||||
assert guide_rescale >= 0 and guide_rescale <= 1
|
||||
ratio = (y_out.flatten(1).std(dim=1) /
|
||||
(out.flatten(1).std(dim=1) +
|
||||
1e-12)).view((-1, ) + (1, ) * (y_out.ndim - 1))
|
||||
out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0
|
||||
def parse_model_kwargs(prev_value, value):
|
||||
if isinstance(value, torch.Tensor):
|
||||
prev_value = torch.cat([prev_value, value],
|
||||
dim=0)
|
||||
elif isinstance(value, dict):
|
||||
for k, v in value.items():
|
||||
prev_value[k] = parse_model_kwargs(
|
||||
prev_value[k], v)
|
||||
elif isinstance(value, list):
|
||||
for idx, v in enumerate(value):
|
||||
prev_value[idx] = parse_model_kwargs(
|
||||
prev_value[idx], v)
|
||||
return prev_value
|
||||
|
||||
all_model_kwargs = copy.deepcopy(model_kwargs[0])
|
||||
for model_kwarg in model_kwargs[1:]:
|
||||
for key, value in model_kwarg.items():
|
||||
all_model_kwargs[key] = parse_model_kwargs(
|
||||
all_model_kwargs[key], value)
|
||||
all_out = model(xt.repeat(2, 1, 1, 1),
|
||||
t=t.repeat(2),
|
||||
**all_model_kwargs,
|
||||
**kwargs)
|
||||
y_out, u_out = all_out.chunk(2)
|
||||
else:
|
||||
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
||||
u_out = model(xt, t=t, **model_kwargs[1], **kwargs)
|
||||
# todo sharpness
|
||||
# sharpness sampling
|
||||
if sharpness is not None and sharpness > 0:
|
||||
positive_x0 = alphas * xt - sigmas * y_out
|
||||
negative_x0 = alphas * xt - sigmas * u_out
|
||||
|
||||
positive_eps = xt - positive_x0
|
||||
negative_eps = xt - negative_x0
|
||||
|
||||
global_diffusion_progress = (
|
||||
1 - t / 999.0).detach().cpu().numpy().tolist()[0]
|
||||
alpha = 0.001 * sharpness * global_diffusion_progress
|
||||
positive_eps_degraded = adaptive_anisotropic_filter(
|
||||
x=positive_eps, g=positive_x0)
|
||||
positive_eps_degraded_weighted = positive_eps_degraded * alpha + positive_eps * (
|
||||
1.0 - alpha)
|
||||
|
||||
final_eps = negative_eps + guide_scale * (
|
||||
positive_eps_degraded_weighted - negative_eps)
|
||||
final_x0 = xt - final_eps
|
||||
out = (alphas * xt - final_x0) / sigmas
|
||||
else:
|
||||
out = u_out + guide_scale * (y_out - u_out)
|
||||
elif isinstance(guide_scale, dict):
|
||||
assert len(model_kwargs) == 3
|
||||
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
||||
m_out = model(xt, t=t, **model_kwargs[1], **kwargs)
|
||||
u_out = model(xt, t=t, **model_kwargs[2], **kwargs)
|
||||
out = u_out + guide_scale['image'] * (
|
||||
m_out - u_out) + guide_scale['text'] * (y_out - m_out)
|
||||
elif isinstance(guide_scale, list):
|
||||
assert len(guide_scale) == len(model_kwargs) - 1
|
||||
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
|
||||
outs = [y_out]
|
||||
for i in range(1, len(model_kwargs)):
|
||||
outs.append(model(xt, t=t, **model_kwargs[i], **kwargs))
|
||||
out = outs[-1]
|
||||
for i in range(len(guide_scale)):
|
||||
out += guide_scale[i] * (outs[-i - 2] - outs[-i - 1])
|
||||
|
||||
# rescale the output according to arXiv:2305.08891
|
||||
if guide_rescale is not None and guide_rescale > 0.0:
|
||||
assert guide_rescale >= 0 and guide_rescale <= 1
|
||||
ratio = (
|
||||
y_out.flatten(1).std(dim=1) /
|
||||
(out.flatten(1).std(dim=1) + 1e-12)).view((-1, ) + (1, ) *
|
||||
(y_out.ndim - 1))
|
||||
out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0
|
||||
# compute x0
|
||||
if self.prediction_type == 'x0':
|
||||
x0 = out
|
||||
@@ -195,6 +384,7 @@ class GaussianDiffusion(object):
|
||||
guide_scale=None,
|
||||
guide_rescale=None,
|
||||
clamp=None,
|
||||
sharpness=0.0,
|
||||
percentile=None,
|
||||
solver='euler_a',
|
||||
steps=20,
|
||||
@@ -207,12 +397,16 @@ class GaussianDiffusion(object):
|
||||
seed=-1,
|
||||
intermediate_callback=None,
|
||||
cat_uc=False,
|
||||
add_noise=False,
|
||||
free_steps=None,
|
||||
step_offset=None,
|
||||
**kwargs):
|
||||
# sanity check
|
||||
assert isinstance(steps, (int, torch.LongTensor))
|
||||
assert t_max is None or (t_max > 0 and t_max <= self.num_timesteps - 1)
|
||||
assert t_min is None or (t_min >= 0 and t_min < self.num_timesteps - 1)
|
||||
assert discretization in (None, 'leading', 'linspace', 'trailing')
|
||||
assert discretization in (None, 'leading', 'linspace', 'trailing',
|
||||
'free')
|
||||
assert discard_penultimate_step in (None, True, False)
|
||||
assert return_intermediate in (None, 'x0', 'xt')
|
||||
|
||||
@@ -253,16 +447,50 @@ class GaussianDiffusion(object):
|
||||
def model_fn(xt, sigma):
|
||||
# denoising
|
||||
t = self._sigma_to_t(sigma).repeat(len(xt)).round().long()
|
||||
x0 = self.denoise(xt,
|
||||
t,
|
||||
None,
|
||||
model,
|
||||
model_kwargs,
|
||||
guide_scale,
|
||||
guide_rescale,
|
||||
clamp,
|
||||
percentile,
|
||||
cat_uc=cat_uc)[-2]
|
||||
|
||||
if isinstance(
|
||||
model_kwargs[0]['cond'], dict) and \
|
||||
'tar_x0' in model_kwargs[0]['cond'] and \
|
||||
'tar_mask_latent' in model_kwargs[0]['cond']:
|
||||
tar_x0 = model_kwargs[0]['cond']['tar_x0']
|
||||
tar_mask = model_kwargs[0]['cond']['tar_mask_latent']
|
||||
|
||||
tar_xt = self.diffuse(x0=tar_x0, t=t)
|
||||
xt = tar_xt * (1.0 - tar_mask) + xt * tar_mask
|
||||
|
||||
if isinstance(model_kwargs[0]['cond'],
|
||||
dict) and 'ref_x0' in model_kwargs[0]['cond']:
|
||||
model_kwargs[0]['cond']['ref_xt'] = self.diffuse(
|
||||
x0=model_kwargs[0]['cond']['ref_x0'], t=t)
|
||||
model_kwargs[1]['cond']['ref_xt'] = self.diffuse(
|
||||
x0=model_kwargs[1]['cond']['ref_x0'], t=t)
|
||||
|
||||
if solver in ('onestep', 'multistep', 'multistep2', 'multistep3'):
|
||||
x0 = self.denoise(xt,
|
||||
t,
|
||||
None,
|
||||
model,
|
||||
model_kwargs,
|
||||
guide_scale,
|
||||
guide_rescale,
|
||||
clamp,
|
||||
sharpness,
|
||||
percentile,
|
||||
cat_uc=cat_uc,
|
||||
**kwargs)[-3]
|
||||
else:
|
||||
x0 = self.denoise(xt,
|
||||
t,
|
||||
None,
|
||||
model,
|
||||
model_kwargs,
|
||||
guide_scale,
|
||||
guide_rescale,
|
||||
clamp,
|
||||
sharpness,
|
||||
percentile,
|
||||
cat_uc=cat_uc,
|
||||
**kwargs)[-2]
|
||||
|
||||
# collect intermediate outputs
|
||||
if return_intermediate == 'xt':
|
||||
@@ -288,10 +516,14 @@ class GaussianDiffusion(object):
|
||||
elif discretization == 'trailing':
|
||||
steps = torch.arange(t_max, t_min - 1,
|
||||
-((t_max - t_min + 1) / steps))
|
||||
elif discretization == 'free':
|
||||
steps = torch.tensor(free_steps)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f'{discretization} discretization not implemented')
|
||||
steps = steps.clamp_(t_min, t_max)
|
||||
elif isinstance(steps, list):
|
||||
steps = torch.tensor(steps)
|
||||
steps = torch.as_tensor(steps,
|
||||
dtype=torch.float32,
|
||||
device=noise.device)
|
||||
@@ -332,6 +564,23 @@ class GaussianDiffusion(object):
|
||||
if discard_penultimate_step:
|
||||
sigmas = torch.cat([sigmas[:-2], sigmas[-1:]])
|
||||
kwargs['seed'] = seed
|
||||
# add noise to x0
|
||||
if add_noise:
|
||||
if 'dm_steps' in kwargs:
|
||||
if step_offset:
|
||||
add_noise_step = -kwargs['dm_steps'] + step_offset
|
||||
if add_noise_step < 0:
|
||||
noise = self.diffuse(
|
||||
noise,
|
||||
torch.full((noise.shape[0], 1),
|
||||
steps[add_noise_step],
|
||||
dtype=torch.int))
|
||||
else:
|
||||
noise = self.diffuse(
|
||||
noise,
|
||||
torch.full((noise.shape[0], 1),
|
||||
steps[-kwargs['dm_steps'] - 1],
|
||||
dtype=torch.int))
|
||||
# sampling
|
||||
x0 = solver_fn(noise,
|
||||
model_fn,
|
||||
@@ -370,168 +619,6 @@ class GaussianDiffusion(object):
|
||||
| torch.isinf(log_sigma)] = float('inf')
|
||||
return log_sigma.exp()
|
||||
|
||||
@torch.no_grad()
|
||||
def stochastic_encode(self, x0, t, steps):
|
||||
# fast, but does not allow for exact reconstruction
|
||||
# t serves as an index to gather the correct alphas
|
||||
|
||||
t_max = None
|
||||
t_min = None
|
||||
|
||||
# discretization method
|
||||
discretization = 'trailing' if self.prediction_type == 'v' else 'leading'
|
||||
|
||||
# timesteps
|
||||
if isinstance(steps, int):
|
||||
t_max = self.num_timesteps - 1 if t_max is None else t_max
|
||||
t_min = 0 if t_min is None else t_min
|
||||
steps = discretize_timesteps(t_max, t_min, steps, discretization)
|
||||
steps = torch.as_tensor(steps).round().long().flip(0).to(x0.device)
|
||||
# steps = torch.as_tensor(steps).round().long().to(x0.device)
|
||||
|
||||
# self.alphas_bar = torch.cumprod(1 - self.sigmas ** 2, dim=0)
|
||||
# print('sigma: ', self.sigmas, len(self.sigmas))
|
||||
# print('alpha_bar: ', self.alphas_bar, len(self.alphas_bar))
|
||||
# print('steps: ', steps, len(steps))
|
||||
# sqrt_alphas_cumprod = torch.sqrt(self.alphas_bar).to(x0.device)[steps]
|
||||
# sqrt_one_minus_alphas_cumprod = torch.sqrt(1 - self.alphas_bar).to(x0.device)[steps]
|
||||
|
||||
sqrt_alphas_cumprod = self.alphas.to(x0.device)[steps]
|
||||
sqrt_one_minus_alphas_cumprod = self.sigmas.to(x0.device)[steps]
|
||||
# print('sigma: ', self.sigmas, len(self.sigmas))
|
||||
# print('alpha: ', self.alphas, len(self.alphas))
|
||||
# print('steps: ', steps, len(steps))
|
||||
|
||||
noise = torch.randn_like(x0)
|
||||
return (
|
||||
extract_into_tensor(sqrt_alphas_cumprod, t, x0.shape) * x0 +
|
||||
extract_into_tensor(sqrt_one_minus_alphas_cumprod, t, x0.shape) *
|
||||
noise)
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_img2img(self,
|
||||
x,
|
||||
noise,
|
||||
model,
|
||||
denoising_strength=1,
|
||||
model_kwargs={},
|
||||
condition_fn=None,
|
||||
guide_scale=None,
|
||||
guide_rescale=None,
|
||||
clamp=None,
|
||||
percentile=None,
|
||||
solver='euler_a',
|
||||
steps=20,
|
||||
t_max=None,
|
||||
t_min=None,
|
||||
discretization=None,
|
||||
discard_penultimate_step=None,
|
||||
return_intermediate=None,
|
||||
show_progress=False,
|
||||
seed=-1,
|
||||
**kwargs):
|
||||
# sanity check
|
||||
assert isinstance(steps, (int, torch.LongTensor))
|
||||
assert t_max is None or (t_max > 0 and t_max <= self.num_timesteps - 1)
|
||||
assert t_min is None or (t_min >= 0 and t_min < self.num_timesteps - 1)
|
||||
assert discretization in (None, 'leading', 'linspace', 'trailing')
|
||||
assert discard_penultimate_step in (None, True, False)
|
||||
assert return_intermediate in (None, 'x0', 'xt')
|
||||
# function of diffusion solver
|
||||
solver_fn = {
|
||||
'euler_ancestral': sample_img2img_euler_ancestral,
|
||||
'euler': sample_img2img_euler,
|
||||
}[solver]
|
||||
# options
|
||||
schedule = 'karras' if 'karras' in solver else None
|
||||
discretization = discretization or 'linspace'
|
||||
seed = seed if seed >= 0 else random.randint(0, 2**31)
|
||||
if isinstance(steps, torch.LongTensor):
|
||||
discard_penultimate_step = False
|
||||
if discard_penultimate_step is None:
|
||||
discard_penultimate_step = True if solver in (
|
||||
'dpm2', 'dpm2_ancestral', 'dpmpp_2m_sde', 'dpm2_karras',
|
||||
'dpm2_ancestral_karras', 'dpmpp_2m_sde_karras') else False
|
||||
|
||||
# function for denoising xt to get x0
|
||||
intermediates = []
|
||||
|
||||
def get_scalings(sigma):
|
||||
c_out = -sigma
|
||||
c_in = 1 / (sigma**2 + 1.**2)**0.5
|
||||
return c_out, c_in
|
||||
|
||||
def model_fn(xt, sigma):
|
||||
# denoising
|
||||
c_out, c_in = get_scalings(sigma)
|
||||
t = self._sigma_to_t(sigma).repeat(len(xt)).round().long()
|
||||
|
||||
x0 = self.denoise(xt * c_in, t, None, model, model_kwargs,
|
||||
guide_scale, guide_rescale, clamp,
|
||||
percentile)[-2]
|
||||
# collect intermediate outputs
|
||||
if return_intermediate == 'xt':
|
||||
intermediates.append(xt)
|
||||
elif return_intermediate == 'x0':
|
||||
intermediates.append(x0)
|
||||
return xt + x0 * c_out
|
||||
|
||||
# get timesteps
|
||||
if isinstance(steps, int):
|
||||
steps += 1 if discard_penultimate_step else 0
|
||||
t_max = self.num_timesteps - 1 if t_max is None else t_max
|
||||
t_min = 0 if t_min is None else t_min
|
||||
# discretize timesteps
|
||||
if discretization == 'leading':
|
||||
steps = torch.arange(t_min, t_max + 1,
|
||||
(t_max - t_min + 1) / steps).flip(0)
|
||||
elif discretization == 'linspace':
|
||||
steps = torch.linspace(t_max, t_min, steps)
|
||||
elif discretization == 'trailing':
|
||||
steps = torch.arange(t_max, t_min - 1,
|
||||
-((t_max - t_min + 1) / steps))
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f'{discretization} discretization not implemented')
|
||||
steps = steps.clamp_(t_min, t_max)
|
||||
steps = torch.as_tensor(steps, dtype=torch.float32, device=x.device)
|
||||
# get sigmas
|
||||
sigmas = self._t_to_sigma(steps)
|
||||
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
|
||||
t_enc = int(min(denoising_strength, 0.999) * len(steps))
|
||||
sigmas = sigmas[len(steps) - t_enc - 1:]
|
||||
noise = x + noise * sigmas[0]
|
||||
|
||||
if schedule == 'karras':
|
||||
if sigmas[0] == float('inf'):
|
||||
sigmas = karras_schedule(
|
||||
n=len(steps) - 1,
|
||||
sigma_min=sigmas[sigmas > 0].min().item(),
|
||||
sigma_max=sigmas[sigmas < float('inf')].max().item(),
|
||||
rho=7.).to(sigmas)
|
||||
sigmas = torch.cat([
|
||||
sigmas.new_tensor([float('inf')]), sigmas,
|
||||
sigmas.new_zeros([1])
|
||||
])
|
||||
else:
|
||||
sigmas = karras_schedule(
|
||||
n=len(steps),
|
||||
sigma_min=sigmas[sigmas > 0].min().item(),
|
||||
sigma_max=sigmas.max().item(),
|
||||
rho=7.).to(sigmas)
|
||||
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
|
||||
if discard_penultimate_step:
|
||||
sigmas = torch.cat([sigmas[:-2], sigmas[-1:]])
|
||||
|
||||
# sampling
|
||||
x0 = solver_fn(noise,
|
||||
model_fn,
|
||||
sigmas,
|
||||
seed=seed,
|
||||
show_progress=show_progress,
|
||||
**kwargs)
|
||||
return (x0, intermediates) if return_intermediate is not None else x0
|
||||
|
||||
|
||||
def extract_into_tensor(a, t, x_shape):
|
||||
b, *_ = t.shape
|
||||
|
||||
@@ -0,0 +1,138 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torchvision.transforms.functional as TF
|
||||
|
||||
|
||||
def get_bbox_from_mask(mask):
|
||||
h, w = mask.shape[0], mask.shape[1]
|
||||
if mask.sum() < 10:
|
||||
return 0, h, 0, w
|
||||
rows = np.any(mask, axis=1)
|
||||
cols = np.any(mask, axis=0)
|
||||
y1, y2 = np.where(rows)[0][[0, -1]]
|
||||
x1, x2 = np.where(cols)[0][[0, -1]]
|
||||
return (y1, y2, x1, x2)
|
||||
|
||||
|
||||
def pad_to_square(image, pad_value=255, random=False):
|
||||
H, W = image.shape[0], image.shape[1]
|
||||
if H == W:
|
||||
return image, 0, 0
|
||||
|
||||
padd = abs(H - W)
|
||||
if random:
|
||||
padd_1 = int(np.random.randint(0, padd))
|
||||
else:
|
||||
padd_1 = int(padd / 2)
|
||||
padd_2 = padd - padd_1
|
||||
|
||||
if H > W:
|
||||
pad_param = ((0, 0), (padd_1, padd_2), (0, 0))
|
||||
else:
|
||||
pad_param = ((padd_1, padd_2), (0, 0), (0, 0))
|
||||
|
||||
# print(pad_param, pad_value)
|
||||
image = np.pad(image, pad_param, 'constant', constant_values=pad_value)
|
||||
return image, padd_1, padd_2
|
||||
|
||||
|
||||
def box_in_box(small_box, big_box):
|
||||
y1, y2, x1, x2 = small_box
|
||||
y1_b, _, x1_b, _ = big_box
|
||||
y1, y2, x1, x2 = y1 - y1_b, y2 - y1_b, x1 - x1_b, x2 - x1_b
|
||||
return (y1, y2, x1, x2)
|
||||
|
||||
|
||||
def box2squre(image, box):
|
||||
H, W = image.shape[0], image.shape[1]
|
||||
y1, y2, x1, x2 = box
|
||||
cx = (x1 + x2) // 2
|
||||
cy = (y1 + y2) // 2
|
||||
h, w = y2 - y1, x2 - x1
|
||||
|
||||
if h >= w:
|
||||
x1 = cx - h // 2
|
||||
x2 = x1 + h
|
||||
else:
|
||||
y1 = cy - w // 2
|
||||
y2 = y1 + w
|
||||
x1 = max(0, x1)
|
||||
x2 = min(W, x2)
|
||||
y1 = max(0, y1)
|
||||
y2 = min(H, y2)
|
||||
return (y1, y2, x1, x2)
|
||||
|
||||
|
||||
def expand_bbox(mask,
|
||||
yyxx,
|
||||
ratio=1.0,
|
||||
min_crop=0,
|
||||
expand_type='center',
|
||||
to_square=False):
|
||||
y1, y2, x1, x2 = yyxx
|
||||
h = y2 - y1 + 1
|
||||
w = x2 - x1 + 1
|
||||
|
||||
H, W = mask.shape[0], mask.shape[1]
|
||||
xc, yc = 0.5 * (x1 + x2), 0.5 * (y1 + y2)
|
||||
|
||||
def expand(k):
|
||||
if isinstance(ratio, tuple) or isinstance(ratio, list):
|
||||
r = np.random.uniform(*ratio)
|
||||
k = k * r
|
||||
else:
|
||||
k = ratio * k
|
||||
return k
|
||||
|
||||
new_h = expand(h)
|
||||
new_w = expand(w)
|
||||
new_h = max(new_h, min_crop)
|
||||
new_w = max(new_w, min_crop)
|
||||
|
||||
if to_square:
|
||||
if new_w / new_h < 0.334:
|
||||
new_w = new_w + 1.0 / 3.0 * new_h
|
||||
elif new_h / new_w < 0.334:
|
||||
new_h = new_h + 1.0 / 3.0 * new_w
|
||||
|
||||
if expand_type == 'center':
|
||||
x1 = max(0, int(xc - new_w * 0.5))
|
||||
x2 = min(W, int(xc + new_w * 0.5))
|
||||
y1 = max(0, int(yc - new_h * 0.5))
|
||||
y2 = min(H, int(yc + new_h * 0.5))
|
||||
else:
|
||||
x1 = max(0, min(x1,
|
||||
int(x2 - new_w * np.random.uniform(w / new_w, 1.0))))
|
||||
x2 = min(W, max(x2, x1 + new_w))
|
||||
y1 = max(0, min(y1,
|
||||
int(y2 - new_h * np.random.uniform(h / new_h, 1.0))))
|
||||
y2 = min(H, max(y2, y1 + new_h))
|
||||
|
||||
return (int(y1), int(y2), int(x1), int(x2))
|
||||
|
||||
|
||||
def crop_back(pred, tar_image, extra_sizes, tar_box_yyxx_crop):
|
||||
H1, W1, H2, W2, pad1, pad2 = extra_sizes
|
||||
y1, y2, x1, x2 = tar_box_yyxx_crop
|
||||
pred = TF.resize(pred, (H2, W2), antialias=True)
|
||||
|
||||
if W1 < W2:
|
||||
# pad width
|
||||
assert H1 == H2 and (pad1 + W1) == (W2 - pad2)
|
||||
pred = pred[:, :, pad1 + 2:(W2 - pad2 - 2)]
|
||||
tar_image[:, y1:y2, x1 + 2:x2 - 2] = pred
|
||||
elif H1 < H2:
|
||||
# pad height
|
||||
assert W1 == W2 and (pad1 + H1) == (H2 - pad2)
|
||||
pred = pred[:, pad1 + 2:(H2 - pad2 - 2), :]
|
||||
tar_image[:, y1 + 2:y2 - 2, x1:x2] = pred
|
||||
else:
|
||||
tar_image[:, y1:y2, x1:x2] = pred
|
||||
return tar_image
|
||||
|
||||
|
||||
def save_image(image, save_path):
|
||||
image = cv2.cvtColor(image, cv2.COLOR_RGB2BGR)
|
||||
cv2.imwrite(save_path, image)
|
||||
@@ -1,7 +1,10 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
from scepter.modules.opt.optimizers.official_optimizers import (
|
||||
ASGD, LBFGS, SGD, Adadelta, Adagrad, Adam, Adamax, AdamW, RMSprop, Rprop,
|
||||
SparseAdam)
|
||||
from scepter.modules.opt.optimizers.official_optimizers import (ASGD, LBFGS,
|
||||
SGD, Adadelta,
|
||||
Adagrad, Adam,
|
||||
Adamax, AdamW,
|
||||
RMSprop, Rprop,
|
||||
SparseAdam)
|
||||
from scepter.modules.opt.optimizers.registry import OPTIMIZERS
|
||||
|
||||
@@ -333,12 +333,12 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
data_iter = iter(self.datas[self._mode].dataloader)
|
||||
self.print_memory_status()
|
||||
for step in range(self.max_steps):
|
||||
if 'eval' in self._mode_set and (step % self.eval_interval == 0
|
||||
or step == self.max_steps - 1):
|
||||
if 'eval' in self._mode_set and (self.eval_interval > 0 and
|
||||
step % self.eval_interval == 0):
|
||||
self.run_eval()
|
||||
self.train_mode()
|
||||
self.before_iter(self.hooks_dict[self._mode])
|
||||
batch_data = next(data_iter)
|
||||
self.before_iter(self.hooks_dict[self._mode])
|
||||
if self.sample_args:
|
||||
batch_data.update(self.sample_args.get_lowercase_dict())
|
||||
if 'meta' in batch_data:
|
||||
@@ -364,6 +364,9 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
self.after_iter(self.hooks_dict[self._mode])
|
||||
if we.debug:
|
||||
self.print_trainable_params_status(prefix='model.')
|
||||
if 'eval' in self._mode_set and (self.eval_interval > 0
|
||||
and step == self.max_steps - 1):
|
||||
self.run_eval()
|
||||
self.after_all_iter(self.hooks_dict[self._mode])
|
||||
|
||||
@torch.no_grad()
|
||||
|
||||
@@ -4,6 +4,7 @@
|
||||
from scepter.modules.solver.hooks.backward import BackwardHook
|
||||
from scepter.modules.solver.hooks.checkpoint import CheckpointHook
|
||||
from scepter.modules.solver.hooks.data_probe import ProbeDataHook
|
||||
from scepter.modules.solver.hooks.ema import ModelEmaHook
|
||||
from scepter.modules.solver.hooks.hook import Hook
|
||||
from scepter.modules.solver.hooks.log import LogHook, TensorboardLogHook
|
||||
from scepter.modules.solver.hooks.lr import LrHook
|
||||
@@ -47,5 +48,6 @@ after solve:
|
||||
|
||||
__all__ = [
|
||||
'HOOKS', 'BackwardHook', 'CheckpointHook', 'Hook', 'LrHook', 'LogHook',
|
||||
'TensorboardLogHook', 'DistSamplerHook', 'ProbeDataHook', 'SafetensorsHook'
|
||||
'TensorboardLogHook', 'DistSamplerHook', 'ProbeDataHook',
|
||||
'SafetensorsHook', 'ModelEmaHook'
|
||||
]
|
||||
|
||||
@@ -147,7 +147,9 @@ class CheckpointHook(Hook):
|
||||
|
||||
if self.save_last and solver.total_iter == solver.max_steps - 1:
|
||||
with FS.get_fs_client(save_path) as client:
|
||||
last_path = osp.join(solver.work_dir, 'checkpoint.pth')
|
||||
last_path = osp.join(
|
||||
solver.work_dir,
|
||||
f'checkpoints/{self.save_name_prefix}-last')
|
||||
client.make_link(last_path, save_path)
|
||||
self.last_ckpt = save_path
|
||||
|
||||
|
||||
@@ -30,7 +30,10 @@ class ProbeDataHook(Hook):
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(ProbeDataHook, self).__init__(cfg, logger=logger)
|
||||
self.priority = cfg.get('PRIORITY', _DEFAULT_PROBE_PRIORITY)
|
||||
self.log_interval = cfg.get('PROB_INTERVAL', 1000)
|
||||
self.prob_interval = cfg.get('PROB_INTERVAL', 1000)
|
||||
self.save_name_prefix = cfg.get('SAVE_NAME_PREFIX', 'step')
|
||||
self.save_probe_prefix = cfg.get('SAVE_PROBE_PREFIX', None)
|
||||
self.save_last = cfg.get('SAVE_LAST', False)
|
||||
|
||||
def before_all_iter(self, solver):
|
||||
pass
|
||||
@@ -39,19 +42,23 @@ class ProbeDataHook(Hook):
|
||||
pass
|
||||
|
||||
def after_iter(self, solver):
|
||||
if solver.mode == 'train' and solver.total_iter % self.log_interval == 0:
|
||||
if solver.mode == 'train' and solver.total_iter % self.prob_interval == 0:
|
||||
probe_dict = solver.probe_data
|
||||
if we.rank == 0:
|
||||
save_folder = os.path.join(
|
||||
solver.work_dir,
|
||||
f'{solver.mode}_probe/step_{solver.total_iter}')
|
||||
f'{solver.mode}_probe/{self.save_name_prefix}-{solver.total_iter}'
|
||||
)
|
||||
ret_data = {}
|
||||
for k, v in probe_dict.items():
|
||||
ret_one = v.to_log(
|
||||
os.path.join(
|
||||
if self.save_probe_prefix is not None:
|
||||
ret_prefix = os.path.join(save_folder,
|
||||
self.save_probe_prefix)
|
||||
else:
|
||||
ret_prefix = os.path.join(
|
||||
save_folder,
|
||||
k.replace('/', '_') +
|
||||
f'_step_{solver.total_iter}'))
|
||||
k.replace('/', '_') + f'_step_{solver.total_iter}')
|
||||
ret_one = v.to_log(ret_prefix)
|
||||
if (isinstance(ret_one, list)
|
||||
or isinstance(ret_one, dict)) and len(ret_one) < 1:
|
||||
continue
|
||||
@@ -71,13 +78,19 @@ class ProbeDataHook(Hook):
|
||||
if we.rank == 0:
|
||||
step = solver._total_iter[
|
||||
'train'] if 'train' in solver._total_iter else 0
|
||||
save_folder = os.path.join(solver.work_dir,
|
||||
f'{solver.mode}_probe/step_{step}')
|
||||
save_folder = os.path.join(
|
||||
solver.work_dir,
|
||||
f'{solver.mode}_probe/{self.save_name_prefix}-{step}')
|
||||
ret_data = {}
|
||||
for k, v in probe_dict.items():
|
||||
ret_one = v.to_log(
|
||||
os.path.join(save_folder,
|
||||
k.replace('/', '_') + f'_step_{step}'))
|
||||
if self.save_probe_prefix is not None:
|
||||
ret_prefix = os.path.join(save_folder,
|
||||
self.save_probe_prefix)
|
||||
else:
|
||||
ret_prefix = os.path.join(
|
||||
save_folder,
|
||||
k.replace('/', '_') + f'_step_{step}')
|
||||
ret_one = v.to_log(ret_prefix)
|
||||
if (isinstance(ret_one, list)
|
||||
or isinstance(ret_one, dict)) and len(ret_one) < 1:
|
||||
continue
|
||||
@@ -87,6 +100,16 @@ class ProbeDataHook(Hook):
|
||||
json.dump(ret_data,
|
||||
open(local_path, 'w'),
|
||||
ensure_ascii=False)
|
||||
|
||||
if self.save_last and step == solver.max_steps:
|
||||
with FS.get_fs_client(save_folder) as client:
|
||||
last_save_folder = os.path.join(
|
||||
solver.work_dir,
|
||||
f'{solver.mode}_probe/{self.save_name_prefix}-last'
|
||||
)
|
||||
print(last_save_folder, save_folder)
|
||||
client.make_link(last_save_folder, save_folder)
|
||||
|
||||
solver.clear_probe()
|
||||
torch.cuda.synchronize()
|
||||
barrier()
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
import torch
|
||||
from torch.distributed.fsdp import (FullStateDictConfig,
|
||||
FullyShardedDataParallel, StateDictType)
|
||||
|
||||
from scepter.modules.solver.hooks.hook import Hook
|
||||
from scepter.modules.solver.hooks.registry import HOOKS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
|
||||
|
||||
@HOOKS.register_class()
|
||||
class ModelEmaHook(Hook):
|
||||
para_dict = [{
|
||||
'PRIORITY': {
|
||||
'value': 100,
|
||||
'description': 'the priority for processing!'
|
||||
},
|
||||
'BETA': {
|
||||
'value': 0.9999,
|
||||
'description': ''
|
||||
},
|
||||
}]
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(ModelEmaHook, self).__init__(cfg, logger=logger)
|
||||
self.priority = cfg.get('PRIORITY', 100)
|
||||
self.beta = cfg.get('BETA', 0.9999)
|
||||
|
||||
def after_iter(self, solver):
|
||||
if solver.model.use_ema:
|
||||
model_ema = solver.model.model_ema
|
||||
model = solver.model.model
|
||||
self.ema(model_ema, model, use_fsdp=solver.use_fsdp)
|
||||
|
||||
@torch.no_grad()
|
||||
def ema(self, net_ema, net, use_fsdp=True):
|
||||
if we.is_distributed:
|
||||
if use_fsdp:
|
||||
save_policy = FullStateDictConfig(offload_to_cpu=False,
|
||||
rank0_only=False)
|
||||
with FullyShardedDataParallel.state_dict_type(
|
||||
net, StateDictType.FULL_STATE_DICT, save_policy):
|
||||
nonema_state = net.state_dict()
|
||||
elif hasattr(net, 'module'):
|
||||
nonema_state = net.module.state_dict()
|
||||
else:
|
||||
nonema_state = net.state_dict()
|
||||
else:
|
||||
nonema_state = net.state_dict()
|
||||
|
||||
for k, v in net_ema.named_parameters():
|
||||
v.copy_(nonema_state[k].lerp(v, self.beta))
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('hook',
|
||||
__class__.__name__,
|
||||
ModelEmaHook.para_dict,
|
||||
set_name=True)
|
||||
@@ -16,7 +16,8 @@ from scepter.modules.transform.io import (LoadCvImageFromFile,
|
||||
from scepter.modules.transform.io_video import (DecodeVideoToTensor,
|
||||
LoadVideoFromFile)
|
||||
from scepter.modules.transform.registry import TRANSFORMS, build_pipeline
|
||||
from scepter.modules.transform.tensor import Rename, Select, ToNumpy, ToTensor
|
||||
from scepter.modules.transform.tensor import (Rename, RenameMeta, Select,
|
||||
TemplateStr, ToNumpy, ToTensor)
|
||||
from scepter.modules.transform.transform_xl import FlexibleCropXL
|
||||
from scepter.modules.transform.video import (AutoResizedCropVideo,
|
||||
CenterCropVideo, NormalizeVideo,
|
||||
|
||||
@@ -10,11 +10,16 @@ import torchvision.transforms as transforms
|
||||
import torchvision.transforms.functional as TF
|
||||
|
||||
from scepter.modules.transform.registry import TRANSFORMS
|
||||
from scepter.modules.transform.utils import (
|
||||
BACKEND_CV2, BACKEND_PILLOW, BACKEND_TORCHVISION, INPUT_CV2_TYPE_WARNING,
|
||||
INPUT_PIL_TYPE_WARNING, INPUT_TENSOR_TYPE_WARNING, INTERPOLATION_STYLE,
|
||||
INTERPOLATION_STYLE_CV2, TORCHVISION_CAPABILITY, is_cv2_image,
|
||||
is_pil_image, is_tensor)
|
||||
from scepter.modules.transform.utils import (BACKEND_CV2, BACKEND_PILLOW,
|
||||
BACKEND_TORCHVISION,
|
||||
INPUT_CV2_TYPE_WARNING,
|
||||
INPUT_PIL_TYPE_WARNING,
|
||||
INPUT_TENSOR_TYPE_WARNING,
|
||||
INTERPOLATION_STYLE,
|
||||
INTERPOLATION_STYLE_CV2,
|
||||
TORCHVISION_CAPABILITY,
|
||||
is_cv2_image, is_pil_image,
|
||||
is_tensor)
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
|
||||
if TORCHVISION_CAPABILITY:
|
||||
|
||||
@@ -6,13 +6,15 @@ import os
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
from PIL import Image, ImageFile
|
||||
|
||||
from scepter.modules.transform.registry import TRANSFORMS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import DATA_FS as FS
|
||||
|
||||
ImageFile.LOAD_TRUNCATED_IMAGES = True
|
||||
|
||||
|
||||
def pillow_convert(image, rgb_order):
|
||||
if image.mode != rgb_order:
|
||||
|
||||
@@ -252,3 +252,87 @@ class TensorToGPU(object):
|
||||
__class__.__name__,
|
||||
para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class RenameMeta(object):
|
||||
def __init__(self, cfg, logger=None):
|
||||
self.input_key = cfg.INPUT_KEY
|
||||
self.output_key = cfg.OUTPUT_KEY
|
||||
self.force = cfg.get('FORCE', False)
|
||||
|
||||
def __call__(self, item):
|
||||
if 'meta' in item:
|
||||
data = {}
|
||||
for idx, key in enumerate(self.input_key):
|
||||
data[self.output_key[idx]] = item['meta'][key]
|
||||
if not self.force:
|
||||
have_key_set = set(self.input_key)
|
||||
else:
|
||||
have_key_set = set(self.input_key + self.output_key)
|
||||
for k, v in item['meta'].items():
|
||||
if k not in have_key_set:
|
||||
data[k] = v
|
||||
item['meta'] = data
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
para_dict = [{
|
||||
'INPUT_KEY': {
|
||||
'value': [],
|
||||
'description':
|
||||
'The keys need to rename, the other keys are outputed by default.'
|
||||
},
|
||||
'OUTPUT_KEY': {
|
||||
'value': [],
|
||||
'description':
|
||||
'The keys need to rename, the other keys are outputed by default.'
|
||||
}
|
||||
}]
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class TemplateStr(object):
|
||||
def __init__(self, cfg, logger=None):
|
||||
self.template_str = cfg.get('TEMPLATE_STR', '')
|
||||
self.meta_template_str = cfg.get('META_TEMPLATE_STR', '')
|
||||
|
||||
def __call__(self, item):
|
||||
if self.template_str != '':
|
||||
for key, val in item.items():
|
||||
if isinstance(val, str) and f'{{{key}}}' in self.template_str:
|
||||
template = self.template_str
|
||||
val = template.replace(f'{{{key}}}', val)
|
||||
item[key] = val
|
||||
if self.meta_template_str != '' and 'meta' in item:
|
||||
for key, val in item['meta'].items():
|
||||
if isinstance(val,
|
||||
str) and f'{{{key}}}' in self.meta_template_str:
|
||||
template = self.meta_template_str
|
||||
val = template.replace(f'{{{key}}}', val)
|
||||
item['meta'][key] = val
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
para_dict = [{}]
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
para_dict,
|
||||
set_name=True)
|
||||
|
||||
@@ -603,3 +603,6 @@ class Config(object):
|
||||
return cfg_new
|
||||
else:
|
||||
return cfg
|
||||
|
||||
def pop(self, name):
|
||||
self.cfg_dict.pop(name)
|
||||
|
||||
@@ -440,6 +440,7 @@ class Workenv(object):
|
||||
def set_env(self, we_env):
|
||||
for k, v in we_env.items():
|
||||
setattr(self, k, v)
|
||||
set_random_seed(self.seed)
|
||||
|
||||
def __str__(self):
|
||||
environ_str = f'Now running in the distributed environment with world size {self.world_size}\n!'
|
||||
|
||||
@@ -2,5 +2,6 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.utils.file_clients.aliyun_oss_fs import AliyunOssFs
|
||||
from scepter.modules.utils.file_clients.http_fs import HttpFs
|
||||
from scepter.modules.utils.file_clients.huggingface_fs import HuggingfaceFs
|
||||
from scepter.modules.utils.file_clients.local_fs import LocalFs
|
||||
from scepter.modules.utils.file_clients.modelscope_fs import ModelscopeFs
|
||||
|
||||
@@ -490,15 +490,88 @@ class AliyunOssFs(BaseFs):
|
||||
meta_dict=copy.deepcopy(meta_dict)))
|
||||
return meta_dict
|
||||
|
||||
def _get_dir_multi(self,
|
||||
target_path,
|
||||
local_path,
|
||||
wait_finish=False,
|
||||
meta_dict={}):
|
||||
local_path = local_path.replace('/./', '/')
|
||||
os.makedirs(local_path, exist_ok=True)
|
||||
generator = self.walk_dir(target_path)
|
||||
single_file_name = []
|
||||
for file_name in generator:
|
||||
if file_name == target_path or file_name == target_path + '/':
|
||||
continue
|
||||
local_file_name = os.path.join(
|
||||
local_path,
|
||||
file_name.split(target_path)[-1]).replace('/./', '/')
|
||||
if not self.isdir(file_name):
|
||||
single_file_name.append((file_name, local_file_name))
|
||||
else:
|
||||
meta_dict.update(
|
||||
self._get_dir_multi(file_name,
|
||||
local_file_name,
|
||||
meta_dict=copy.deepcopy(meta_dict)))
|
||||
|
||||
data_quene = queue.Queue()
|
||||
batch_size = 20
|
||||
R = threading.Lock()
|
||||
|
||||
def get_one_object(target_path_list):
|
||||
if isinstance(target_path_list, tuple):
|
||||
target_path_list = [target_path_list]
|
||||
for target_path, local_path in target_path_list:
|
||||
if self.exists(target_path):
|
||||
etag, size = self.get_meta(target_path)
|
||||
if local_path in meta_dict and meta_dict[
|
||||
local_path] == etag:
|
||||
continue
|
||||
process_msg(f'Download {target_path} to {local_path}....')
|
||||
local_path = self.get_object_to_local_file(
|
||||
target_path, local_path, wait_finish=wait_finish)
|
||||
assert local_path is not None
|
||||
meta_dict[target_path] = etag
|
||||
else:
|
||||
local_path = None
|
||||
R.acquire()
|
||||
try:
|
||||
data_quene.put_nowait([target_path, local_path])
|
||||
except Exception:
|
||||
R.release()
|
||||
R.release()
|
||||
|
||||
while True:
|
||||
batch_list = single_file_name[:10 * batch_size]
|
||||
if len(batch_list) < 1:
|
||||
break
|
||||
single_file_name = single_file_name[10 * batch_size:]
|
||||
threading_list = []
|
||||
for i in range(batch_size):
|
||||
cur_batch = batch_list[i::batch_size]
|
||||
if isinstance(cur_batch, tuple):
|
||||
cur_batch = [cur_batch]
|
||||
t = threading.Thread(target=get_one_object, args=(cur_batch, ))
|
||||
t.daemon = True
|
||||
t.start()
|
||||
threading_list.append(t)
|
||||
[threading_t.join() for threading_t in threading_list]
|
||||
file_dict = {}
|
||||
while not data_quene.empty():
|
||||
target_path, local_path = data_quene.get_nowait()
|
||||
file_dict[target_path] = local_path
|
||||
return meta_dict
|
||||
|
||||
def get_dir_to_local_dir(self,
|
||||
target_path,
|
||||
local_path=None,
|
||||
wait_finish=False,
|
||||
timeout=3600,
|
||||
multi_thread=False,
|
||||
worker_id=-1) -> Optional[str]:
|
||||
if not self.isdir(target_path):
|
||||
self.logger.info(
|
||||
f"{target_path} is not directory or doesn't exist.")
|
||||
return None
|
||||
if not target_path.endswith('/'):
|
||||
target_path += '/'
|
||||
if local_path is None:
|
||||
@@ -533,9 +606,14 @@ class AliyunOssFs(BaseFs):
|
||||
meta_dict = json.load(open(check_file, 'r'))
|
||||
else:
|
||||
meta_dict = {}
|
||||
meta_dict = self._get_dir(target_path,
|
||||
local_path=local_path,
|
||||
meta_dict=copy.deepcopy(meta_dict))
|
||||
if multi_thread:
|
||||
meta_dict = self._get_dir_multi(target_path,
|
||||
local_path=local_path,
|
||||
meta_dict=copy.deepcopy(meta_dict))
|
||||
else:
|
||||
meta_dict = self._get_dir(target_path,
|
||||
local_path=local_path,
|
||||
meta_dict=copy.deepcopy(meta_dict))
|
||||
json.dump(meta_dict, open(check_file, 'w'))
|
||||
if is_tmp:
|
||||
self.add_temp_file(local_path)
|
||||
@@ -938,16 +1016,67 @@ class AliyunOssFs(BaseFs):
|
||||
continue
|
||||
yield osp.join(self._prefix, obj.key)
|
||||
|
||||
def put_dir_from_local_dir(self, local_dir, target_dir) -> bool:
|
||||
def put_dir_from_local_dir(self,
|
||||
local_dir,
|
||||
target_dir,
|
||||
multi_thread=False) -> bool:
|
||||
singe_file_names = []
|
||||
for folder, sub_folders, files in os.walk(local_dir):
|
||||
for file in files:
|
||||
file_abs_path = osp.join(folder, file)
|
||||
file_rel_path = osp.relpath(file_abs_path, local_dir)
|
||||
target_path = osp.join(target_dir, file_rel_path)
|
||||
singe_file_names.append((file_abs_path, target_path))
|
||||
|
||||
if not multi_thread:
|
||||
for file_abs_path, target_path in singe_file_names:
|
||||
status = self.put_object_from_local_file(
|
||||
file_abs_path, target_path)
|
||||
if not status:
|
||||
return False
|
||||
else:
|
||||
data_quene = queue.Queue()
|
||||
R = threading.Lock()
|
||||
batch_size = 20
|
||||
|
||||
def put_one_object(target_path_list):
|
||||
if isinstance(target_path_list, tuple):
|
||||
target_path_list = [target_path_list]
|
||||
for local_path, target_path in target_path_list:
|
||||
if local_path is None or target_path is None:
|
||||
flg = False
|
||||
elif os.path.exists(local_path):
|
||||
flg = self.put_object_from_local_file(
|
||||
local_path, target_path)
|
||||
else:
|
||||
flg = False
|
||||
R.acquire()
|
||||
try:
|
||||
data_quene.put_nowait([local_path, target_path, flg])
|
||||
except Exception:
|
||||
R.release()
|
||||
R.release()
|
||||
|
||||
while True:
|
||||
batch_list = singe_file_names[:10 * batch_size]
|
||||
if len(batch_list) < 1:
|
||||
break
|
||||
singe_file_names = singe_file_names[10 * batch_size:]
|
||||
threading_list = []
|
||||
for i in range(batch_size):
|
||||
cur_batch = batch_list[i::batch_size]
|
||||
if isinstance(cur_batch, tuple):
|
||||
cur_batch = [cur_batch]
|
||||
t = threading.Thread(target=put_one_object,
|
||||
args=(cur_batch, ))
|
||||
t.daemon = True
|
||||
t.start()
|
||||
threading_list.append(t)
|
||||
[threading_t.join() for threading_t in threading_list]
|
||||
while not data_quene.empty():
|
||||
local_path, target_path, flg = data_quene.get_nowait()
|
||||
if not flg:
|
||||
return False
|
||||
return True
|
||||
|
||||
def size(self, target_path) -> Optional[int]:
|
||||
|
||||
@@ -99,7 +99,10 @@ class HttpFs(BaseFs):
|
||||
def walk_dir(self, file_dir, recurse=True):
|
||||
raise NotImplementedError
|
||||
|
||||
def put_dir_from_local_dir(self, local_dir, target_dir) -> bool:
|
||||
def put_dir_from_local_dir(self,
|
||||
local_dir,
|
||||
target_dir,
|
||||
multi_thread=False) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
def size(self, target_path) -> Optional[int]:
|
||||
@@ -119,6 +122,15 @@ class HttpFs(BaseFs):
|
||||
delimiter=None) -> (Union[bytes, str, None], Optional[int]):
|
||||
raise NotImplementedError
|
||||
|
||||
def get_dir_to_local_dir(self,
|
||||
target_path,
|
||||
local_path=None,
|
||||
wait_finish=False,
|
||||
multi_thread=False,
|
||||
timeout=3600,
|
||||
worker_id=0) -> Optional[str]:
|
||||
raise NotImplementedError
|
||||
|
||||
def get_url(self, target_path, lifecycle=3600 * 100):
|
||||
return target_path
|
||||
|
||||
|
||||
@@ -0,0 +1,202 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import os
|
||||
import os.path as osp
|
||||
import urllib.parse as parse
|
||||
import urllib.request
|
||||
from typing import Optional, Union
|
||||
|
||||
from scepter.modules.utils.file_clients.base_fs import BaseFs
|
||||
from scepter.modules.utils.file_clients.registry import FILE_SYSTEMS
|
||||
|
||||
|
||||
@FILE_SYSTEMS.register_class()
|
||||
class HuggingfaceFs(BaseFs):
|
||||
para_dict = {
|
||||
'RETRY_TIMES': {
|
||||
'value': 10,
|
||||
'description': 'Retry get object times.'
|
||||
}
|
||||
}
|
||||
para_dict.update(BaseFs.para_dict)
|
||||
|
||||
def __init__(self, cfg, logger):
|
||||
super(HuggingfaceFs, self).__init__(cfg, logger=logger)
|
||||
retry_times = cfg.get('RETRY_TIMES', 10)
|
||||
self._retry_times = retry_times
|
||||
|
||||
def get_prefix(self) -> str:
|
||||
return 'hf://'
|
||||
|
||||
def support_write(self) -> bool:
|
||||
return False
|
||||
|
||||
def support_link(self) -> bool:
|
||||
return False
|
||||
|
||||
def basename(self, target_path) -> str:
|
||||
url = parse.unquote(target_path)
|
||||
url = url.split('?')[0]
|
||||
return osp.basename(url)
|
||||
|
||||
def get_object_to_local_file(self,
|
||||
target_path,
|
||||
local_path=None,
|
||||
wait_finish=False) -> Optional[str]:
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
key = osp.relpath(target_path, self.get_prefix())
|
||||
key, file_path = key.split('@', 1)
|
||||
|
||||
if ':' in key:
|
||||
key, revision = key.split(':', 1)
|
||||
else:
|
||||
revision = None
|
||||
|
||||
if local_path is None:
|
||||
local_path, is_tmp = self.map_to_local(key)
|
||||
else:
|
||||
is_tmp = False
|
||||
|
||||
if revision is not None:
|
||||
local_path = local_path + '_' + str(revision)
|
||||
|
||||
retry = 0
|
||||
while retry < self._retry_times:
|
||||
try:
|
||||
local_path = hf_hub_download(repo_id=key,
|
||||
revision=revision,
|
||||
filename=file_path,
|
||||
cache_dir=local_path)
|
||||
if osp.exists(local_path):
|
||||
break
|
||||
except Exception:
|
||||
retry += 1
|
||||
|
||||
if retry >= self._retry_times:
|
||||
return None
|
||||
|
||||
if is_tmp:
|
||||
self.add_temp_file(local_path)
|
||||
return local_path
|
||||
|
||||
def get_dir_to_local_dir(self,
|
||||
target_path,
|
||||
local_path=None,
|
||||
wait_finish=False,
|
||||
timeout=3600,
|
||||
worker_id=-1) -> Optional[str]:
|
||||
from huggingface_hub import snapshot_download
|
||||
assert target_path.startswith(self.get_prefix())
|
||||
|
||||
key = osp.relpath(target_path, self.get_prefix())
|
||||
if '@' not in key:
|
||||
key, ret_folder = key.split('@', 1)[0], ''
|
||||
else:
|
||||
at_level_folder = key.split('@')
|
||||
if len(at_level_folder) > 2:
|
||||
raise f'Target path should include only one @, but you give {len(at_level_folder)} @.'
|
||||
key, ret_folder = at_level_folder
|
||||
|
||||
if ':' in key:
|
||||
key, revision = key.split(':', 1)
|
||||
else:
|
||||
revision = None
|
||||
|
||||
if local_path is None:
|
||||
local_path, is_tmp = self.map_to_local(key)
|
||||
else:
|
||||
is_tmp = False
|
||||
|
||||
if revision is not None:
|
||||
local_path = local_path + '_' + str(revision)
|
||||
|
||||
retry = 0
|
||||
while retry < self._retry_times:
|
||||
try:
|
||||
local_path = snapshot_download(repo_id=key,
|
||||
revision=revision,
|
||||
cache_dir=local_path)
|
||||
if osp.exists(local_path):
|
||||
break
|
||||
except Exception:
|
||||
retry += 1
|
||||
|
||||
if retry >= self._retry_times:
|
||||
return None
|
||||
|
||||
if is_tmp:
|
||||
self.add_temp_file(local_path)
|
||||
if not ret_folder == '':
|
||||
local_path = os.path.join(local_path, ret_folder)
|
||||
return local_path
|
||||
|
||||
def get_object(self, target_path):
|
||||
try:
|
||||
local_data = open(self.get_object_to_local_file(target_path),
|
||||
'rb').read()
|
||||
except Exception as e:
|
||||
self.logger.error(f'Read {target_path} error {e}')
|
||||
local_data = None
|
||||
return local_data
|
||||
|
||||
def put_object(self, local_data, target_path):
|
||||
raise NotImplementedError
|
||||
|
||||
def put_object_from_local_file(self, local_path, target_path) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
def make_link(self, target_link_path, target_path) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
def make_dir(self, target_dir) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
def remove(self, target_path) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
def get_logging_handler(self, target_logging_path):
|
||||
raise NotImplementedError
|
||||
|
||||
def walk_dir(self, file_dir, recurse=True):
|
||||
raise NotImplementedError
|
||||
|
||||
def put_dir_from_local_dir(self, local_dir, target_dir) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
def size(self, target_path) -> Optional[int]:
|
||||
raise NotImplementedError
|
||||
|
||||
def get_object_chunk_list(self,
|
||||
target_path,
|
||||
chunk_num=1,
|
||||
delimiter=None) -> Optional[list]:
|
||||
raise NotImplementedError
|
||||
|
||||
def get_object_stream(
|
||||
self,
|
||||
target_path,
|
||||
start,
|
||||
size=10000,
|
||||
delimiter=None) -> (Union[bytes, str, None], Optional[int]):
|
||||
raise NotImplementedError
|
||||
|
||||
def get_url(self, target_path, lifecycle=3600 * 100):
|
||||
return target_path
|
||||
|
||||
def exists(self, target_path) -> bool:
|
||||
req = urllib.request.Request(target_path)
|
||||
req.get_method = lambda: 'HEAD'
|
||||
|
||||
try:
|
||||
urllib.request.urlopen(req)
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def isfile(self, target_path) -> bool:
|
||||
# Well for a http url, it should only be a file.
|
||||
return True
|
||||
|
||||
def isdir(self, target_path) -> bool:
|
||||
return False
|
||||