Compare commits

...
23 Commits
Author SHA1 Message Date
mcj 0010d4282a Merge pull request #18 from modelscope/v0.0.4_dev
V0.0.4 dev
2024-04-10 09:41:45 +08:00
LouieStark 8a9b791e3f fix bug 2024-04-09 14:02:28 +08:00
LouieStark ee7fe888f2 delete clip.py 2024-04-02 18:27:47 +08:00
LouieStark 36b259b4b4 Merge pull request #14 from modelscope/v0.0.4_dev
fix largen default
2024-04-02 16:19:30 +08:00
LouieStark 92c849412e fix largen default 2024-04-02 16:15:25 +08:00
LouieStark 2a7e026f84 Merge pull request #13 from modelscope/v0.0.4_dev
V0.0.4 dev
2024-04-01 12:07:24 +08:00
LouieStark cdba82baf8 update readme 2024-04-01 12:06:48 +08:00
LouieStark 565c7957d8 fix bug 2024-04-01 11:46:03 +08:00
LouieStark e00c23d09a fix error 2024-03-31 19:26:14 +08:00
LouieStark bf53829530 update v0.0.4 2024-03-31 13:08:41 +08:00
zeyinzi.jzyz 35aada8ce8 Fix required_memory 2024-02-29 20:12:08 +08:00
Zhen Han d3ce651bf7 Update readme.md 2024-02-07 21:08:16 +08:00
Zhen Han 7a58c91940 Update readme.md 2024-02-07 20:24:19 +08:00
Zhen Han 4e1606af2d Merge pull request #6 from modelscope/v0.0.3_dev
V0.0.3 dev
2024-02-07 20:18:55 +08:00
Zhen Han 09459c11b7 Update readme.md 2024-02-07 20:16:58 +08:00
hanzhn d9a48268d5 update example cache dir 2024-02-07 20:02:58 +08:00
jiangzeyinzi 9f1847501d fix ctr null 2024-02-07 19:46:48 +08:00
Zhen Han 8214227098 Update readme.md 2024-02-07 18:49:21 +08:00
Zhen Han 2e69b2b116 Update readme.md 2024-02-07 15:35:12 +08:00
hanzhn 3440ec7c38 update v0.0.3 2024-02-07 14:57:28 +08:00
hanzhn 9999e0e1f9 v0.0.3 2024-02-06 17:58:30 +08:00
mcj 01c03683e8 Merge pull request #5 from eltociear/patch-1
Update readme.md
2024-01-25 09:45:20 +08:00
Ikko Eltociear Ashimine 9adb273e4b Update readme.md
approches -> approaches
2024-01-25 02:52:48 +09:00
170 changed files with 8795 additions and 1615 deletions
+2
View File
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
Binary file not shown.

After

Width:  |  Height:  |  Size: 121 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 16 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 19 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 120 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 117 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 121 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 22 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 39 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 17 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 118 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 15 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 20 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 49 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 45 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 34 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 129 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 101 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 100 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 104 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 103 KiB

+1 -1
View File
@@ -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/>
+2
View File
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
+3
View File
@@ -0,0 +1,3 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from .classifier_dataset import ImageClassifyExampleDataset
+260
View File
@@ -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"]
+80
View File
@@ -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)
+7
View File
@@ -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()
+174 -23
View File
@@ -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) [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=SCEdit&color=red&logo=arxiv)](https://arxiv.org/abs/2312.11392) [![Page link](https://img.shields.io/badge/Page-SCEdit-Gree)](https://scedit.github.io/)
3. Res-Tuning(TODO): [Res-Tuning: A Flexible and Efficient Tuning Paradigm via Unbinding Tuner from Backbone](https://arxiv.org/abs/2310.19859) [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=ResTuning&color=red&logo=arxiv)](https://arxiv.org/abs/2310.19859) [![Page link](https://img.shields.io/badge/Page-ResTuning-Gree)](https://res-tuning.github.io/)
## 🎉 News
- [2024.01]: We release **SCEPTER Studio**, an integrated toolkit for data management, model training and inference based on [Gradio](https://www.gradio.app/).
- [2024.01]: [SCEdit](https://arxiv.org/abs/2312.11392) support controllable image synthesis for training and inference.
- [2023.12]: We propose [SCEdit](https://arxiv.org/abs/2312.11392), an efficient and controllable generation framework.
- [2023.12]: We release [🪄SCEPTER](https://github.com/modelscope/scepter/) library.
2. SCEdit(CVPR2024): [SCEdit: Efficient and Controllable Image Diffusion Generation via Skip Connection Editing](https://arxiv.org/abs/2312.11392) [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=SCEdit&color=red&logo=arxiv)](https://arxiv.org/abs/2312.11392) [![Page link](https://img.shields.io/badge/Page-SCEdit-Gree)](https://scedit.github.io/)
3. Res-Tuning(NeurIPS2023 TODO): [Res-Tuning: A Flexible and Efficient Tuning Paradigm via Unbinding Tuner from Backbone](https://arxiv.org/abs/2310.19859) [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=ResTuning&color=red&logo=arxiv)](https://arxiv.org/abs/2310.19859) [![Page link](https://img.shields.io/badge/Page-ResTuning-Gree)](https://res-tuning.github.io/)
4. LAR-Gen: [Locate, Assign, Refine: Taming Customized Image Inpainting with Text-Subject Guidance](https://arxiv.org/abs/2403.19534) [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=LARGen&color=red&logo=arxiv)](https://arxiv.org/abs/2403.19534) [![Page link](https://img.shields.io/badge/Page-LARGen-Gree)](https://ali-vilab.github.io/largen-page/)
## 🛠️ Installation
@@ -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.
+5 -1
View File
@@ -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
+4
View File
@@ -0,0 +1,4 @@
git+https://github.com/cocodataset/panopticapi.git
torch==2.0.1
torchvision==0.15.2
xformers==0.0.21
+1
View File
@@ -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: 夸张漫画
+11 -11
View File
@@ -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" ]
+6 -2
View File
@@ -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"
+3
View File
@@ -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
+24 -4
View File
@@ -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
+17 -2
View File
@@ -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
View File
@@ -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.
+25
View File
@@ -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)
+25
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import cv2
import numpy as np
+1 -3
View File
@@ -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
+12 -7
View File
@@ -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
View File
@@ -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
+121 -294
View File
@@ -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):
+4 -1
View File
@@ -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)
+141 -31
View File
@@ -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,
+160
View File
@@ -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)
+6 -3
View File
@@ -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
+138
View File
@@ -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)
+6 -3
View File
@@ -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
+6 -3
View File
@@ -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()
+3 -1
View File
@@ -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'
]
+3 -1
View File
@@ -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
+35 -12
View File
@@ -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()
+60
View File
@@ -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)
+2 -1
View File
@@ -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 -5
View File
@@ -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:
+3 -1
View File
@@ -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:
+84
View File
@@ -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)
+3
View File
@@ -603,3 +603,6 @@ class Config(object):
return cfg_new
else:
return cfg
def pop(self, name):
self.cfg_dict.pop(name)
+1
View File
@@ -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]:
+13 -1
View File
@@ -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

Some files were not shown because too many files have changed in this diff Show More