Compare commits

..
24 Commits
Author SHA1 Message Date
Zhen Han 8076aae7da Merge pull request #28 from modelscope/v0.0.5_dev
V0.0.5 dev
2024-04-29 16:54:18 +08:00
靖渊 79e4910ce4 Merge branch 'main' into v0.0.5_dev 2024-04-29 16:04:23 +08:00
靖渊 6be496d5ac update v0.0.5.post1 2024-04-29 16:02:23 +08:00
jiangzeyinzi b5b2fd9974 Merge pull request #25 from modelscope/v0.0.5_dev
fix OOM bug
2024-04-23 17:46:53 +08:00
靖渊 f502f8063d fix OOM bug 2024-04-23 17:43:51 +08:00
mcj a3711e1d1d Merge pull request #24 from modelscope/v0.0.5_dev
update readme&OOM
2024-04-22 20:24:57 +08:00
靖渊 8ba0b4c673 update readme&OOM 2024-04-22 16:19:03 +08:00
mcj 4f4e516164 Merge pull request #22 from modelscope/v0.0.5_dev
update readme.md
2024-04-20 11:31:55 +08:00
Jingfeng727 88f0244516 update readme.md 2024-04-20 10:40:06 +08:00
hanzhn 5733918aec readme update 2024-04-19 21:56:35 +08:00
mcj da88802b68 Merge pull request #21 from modelscope/v0.0.5_dev
update v0.0.5
2024-04-19 11:33:41 +08:00
Jingfeng727 16d99a268f modify req&import format 2024-04-19 10:39:21 +08:00
Jingfeng727 fc46a63ac8 update v0.0.5 2024-04-18 15:53:37 +08:00
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
169 changed files with 14050 additions and 2284 deletions
+1
View File
@@ -1 +1,2 @@
recursive-include scepter *.yaml
recursive-include scepter *.md
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: 109 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 247 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 301 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 230 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 353 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 32 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 144 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 126 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 138 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 170 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 273 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 243 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 134 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 119 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 433 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 265 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 300 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 155 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

+17 -17
View File
@@ -14,8 +14,8 @@ Model modules are divided into backbones, necks, heads, loss, metrics, networks,
Subclass registration:
```python
from scepter.model.registry import BACKBONES
from scepter.model.base_model import BaseModel
from scepter.modules.model.registry import BACKBONES
from scepter.modules.model.base_model import BaseModel
@BACKBONES.register_class("ResNet")
@@ -25,8 +25,8 @@ class ResNet(BaseModel):
```
```python
from scepter.model.registry import NECKS
from scepter.model.base_model import BaseModel
from scepter.modules.model.registry import NECKS
from scepter.modules.model.base_model import BaseModel
@NECKS.register_class()
@@ -36,8 +36,8 @@ class GlobalAveragePooling(BaseModel):
```
```python
from scepter.model.registry import HEADS
from scepter.model.base_model import BaseModel
from scepter.modules.model.registry import HEADS
from scepter.modules.model.base_model import BaseModel
@HEADS.register_class()
@@ -47,7 +47,7 @@ class ClassifierHead(BaseModel):
```
```python
from scepter.model.registry import LOSSES
from scepter.modules.model.registry import LOSSES
import torch.nn as nn
@@ -59,7 +59,7 @@ class CrossEntropy(nn.Module):
Actual usage:
```python
from scepter.model.registry import BACKBONES, NECKS, HEADS, LOSSES
from scepter.modules.model.registry import BACKBONES, NECKS, HEADS, LOSSES
backbone = BACKBONES.build(cfg.BACKBONE, logger=logger)
neck = NECKS.build(cfg.NECK, logger=logger)
@@ -83,8 +83,8 @@ To be implemented specifically as needed;
Basic Usage Subclass registration:
```python
from scepter.model.metrics.registry import METRICS
from scepter.model.metrics.base_metric import BaseMetric
from scepter.modules.model.metrics.registry import METRICS
from scepter.modules.model.metrics.base_metric import BaseMetric
@METRICS.register_class("AccuracyMetric")
@@ -95,7 +95,7 @@ class AccuracyMetric(BaseMetric):
Actual usage:
```python
from scepter.model.metrics.registry import METRICS
from scepter.modules.model.metrics.registry import METRICS
metric = METRICS.build(cfgs, logger)
```
@@ -117,8 +117,8 @@ Typically takes logits and labels as well as other necessary variables as inputs
Subclass registration:
```python
from scepter.model.registry import TOKENIZERS
from scepter.model.tokenizers import BaseTokenizer
from scepter.modules.model.registry import TOKENIZERS
from scepter.modules.model.tokenizers import BaseTokenizer
@TOKENIZERS.register_class()
@@ -129,7 +129,7 @@ class BaseBertTokenizer(BaseTokenizer):
Actual usage:
```python
from scepter.model.registry import TOKENIZERS
from scepter.modules.model.registry import TOKENIZERS
tokenizer = TOKENIZERS.build(cfgs, logger)
```
@@ -147,8 +147,8 @@ Takes a list of texts that need tokenization as input and outputs token id seque
Subclass registration:
```python
from scepter.model.registry import MODELS
from scepter.model.networks.train_module import TrainModule
from scepter.modules.model.registry import MODELS
from scepter.modules.model.networks.train_module import TrainModule
@MODELS.register_class()
@@ -159,7 +159,7 @@ class Classifier(TrainModule):
Actual usage:
```python
from scepter.model.registry import MODELS
from scepter.modules.model.registry import MODELS
model = MODELS.build(self.cfg.MODEL, logger=self.logger)
```
+4 -4
View File
@@ -9,8 +9,8 @@
Usage when subclassing lr_schedulers:
```python
from scepter.opt.lr_schedulers import LR_SCHEDULERS
from scepter.opt.lr_schedulers.base_scheduler import BaseScheduler
from scepter.modules.opt.lr_schedulers import LR_SCHEDULERS
from scepter.modules.opt.lr_schedulers.base_scheduler import BaseScheduler
@LR_SCHEDULERS.register_class()
@@ -48,8 +48,8 @@ Sets up the schedule for the passed-in optimizer object;
Usage when subclassing optimizers:
```python
from scepter.opt.optimizers.base_optimizer import BaseOptimize
from scepter.opt.optimizers.registry import OPTIMIZERS
from scepter.modules.opt.optimizers.base_optimizer import BaseOptimize
from scepter.modules.opt.optimizers.registry import OPTIMIZERS
@OPTIMIZERS.register_class()
+5 -5
View File
@@ -6,17 +6,17 @@ This is the File System Module, designed to handle file transfer functionalities
The component currently supports three types of IO Handler:
1. scepter.utils.file_clients.AliyunOssFs
2. scepter.utils.file_clients.LocalFs
3. scepter.utils.file_clients.HttpFs
1. scepter.modules.utils.file_clients.AliyunOssFs
2. scepter.modules.utils.file_clients.LocalFs
3. scepter.modules.utils.file_clients.HttpFs
<hr/>
## Basic Usage
```python
from scepter.utils.file_system import FS
from scepter.utils.config import Config
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.config import Config
fs_cfg = Config(load=False, cfg_dict={
"NAME": "AliyunOssFs",
+39 -39
View File
@@ -4,18 +4,18 @@ Relies on SDKs, which are used to organize modules and SDKs that are frequently
## Overview
1. Parameter sdk (scepter.utils.config)
2. Path sdk (scepter.utils.directory)
3. PyTorch distributed sdk (scepter.utils.distribute)
4. Model export sdk (scepter.utils.export_model)
5. File system sdk (scepter.utils.file_system)
6. Logging sdk (scepter.utils.logger)
7. Video processing sdk (scepter.utils.video_reader), see the document (video_reader.md)
8. Module registration sdk (scepter.utils.registry)
9. Data sdk (scepter.utils.data)
10. Model sdk (scepter.utils.model)
11. Sampler sdk (scepter.utils.sampler)
12. Probing sdk (scepter.utils.probe)
1. Parameter sdk (scepter.modules.utils.config)
2. Path sdk (scepter.modules.utils.directory)
3. PyTorch distributed sdk (scepter.modules.utils.distribute)
4. Model export sdk (scepter.modules.utils.export_model)
5. File system sdk (scepter.modules.utils.file_system)
6. Logging sdk (scepter.modules.utils.logger)
7. Video processing sdk (scepter.modules.utils.video_reader), see the document (video_reader.md)
8. Module registration sdk (scepter.modules.utils.registry)
9. Data sdk (scepter.modules.utils.data)
10. Model sdk (scepter.modules.utils.model)
11. Sampler sdk (scepter.modules.utils.sampler)
12. Probing sdk (scepter.modules.utils.probe)
<hr/>
@@ -24,7 +24,7 @@ Relies on SDKs, which are used to organize modules and SDKs that are frequently
### Basic Usage
```python
from scepter.utils.config import Config
from scepter.modules.utils.config import Config
# Initialize Config object from a dict
fs_cfg = Config(load=False, cfg_dict={"NAME": "LocalFs"})
print(fs_cfg.NAME)
@@ -105,7 +105,7 @@ print(fs_cfg.args)
Some commonly used path functions
### Basic Usage
```python
from scepter.utils.directory import osp_path
from scepter.modules.utils.directory import osp_path
# Automatically join paths based on the path prefix
prefix = "xxxx"
data_file = "example_videos/1.mp4"
@@ -114,13 +114,13 @@ print(osp_path(prefix, data_file))
# Also outputs as xxxx/example_videos/1.mp4
data_file = "xxxx/example_videos/1.mp4"
print(osp_path(prefix, data_file))
from scepter.utils.directory import get_relative_folder
from scepter.modules.utils.directory import get_relative_folder
# Get the folder path at a specified level according to the path
# By default, the last level xxxx/example_videos/
print(get_relative_folder(data_file))
# The second last level xxxx/
print(get_relative_folder(data_file, keep_index=-2))
from scepter.utils.directory import get_md5
from scepter.modules.utils.directory import get_md5
# Get the md5 code of the text/path 34a447fb46d0b786a3999c9dad01d470
print(get_md5(data_file))
```
@@ -175,8 +175,8 @@ PyTorch distributed initialization SDK. By using this SDK, users can avoid focus
### Basic Usage
```python
from scepter.utils.distribute import we
from scepter.utils.config import Config
from scepter.modules.utils.distribute import we
from scepter.modules.utils.config import Config
cfg = Config(cfg_dict={}, load=False)
@@ -304,12 +304,12 @@ Since cloning is involved, this may cause additional GPU memory waste.
**Returns**
- **tensor** —— The output tensor on the CPU for process rank=0.
## 4. 模型导出sdk(scepter.utils.export_model)
## 4. 模型导出sdk(scepter.modules.utils.export_model)
APIs for exporting models to TorchScript/ONNX formats.
### Basic Usage
```python
from scepter.utils.export_model import save_develop_model_multi_io
from scepter.modules.utils.export_model import save_develop_model_multi_io
save_develop_model_multi_io(
model,
@@ -345,16 +345,16 @@ Supports importing and exporting models with multiple inputs and outputs
**Returns**
- **tensor** —— The output tensor on the CPU for process rank=0.
## 5. 文件系统sdk(scepter.utils.file_system)
## 5. 文件系统sdk(scepter.modules.utils.file_system)
Refer to [file_clients](file_clients.md)
## 6. Logging SDK(scepter.utils.logger)
## 6. Logging SDK(scepter.modules.utils.logger)
Used to instantiate a standard logging instance for printing information.
### Basic Usage
```python
from scepter.utils.logger import get_logger, init_logger
from scepter.modules.utils.logger import get_logger, init_logger
std_logger = get_logger(name="scepter")
init_logger(std_logger, log_file="", dist_launcher="pytorch")
@@ -405,14 +405,14 @@ Calculate the time remaining until completion based on the current usage time an
**Returns**
- **str** —— Formatted output.
## 7. Video Processing SDK (scepter.utils.video_reader)
## 7. Video Processing SDK (scepter.modules.utils.video_reader)
APIs for handling video reading.
### Basic Usage
```python
from scepter.utils.video_reader.frame_sampler import do_frame_sample
from scepter.utils.video_reader.video_reader import (
from scepter.modules.utils.video_reader.frame_sampler import do_frame_sample
from scepter.modules.utils.video_reader.video_reader import (
VideoReaderWrapper, EasyVideoReader, FramesReaderWrapper
)
```
@@ -554,14 +554,14 @@ Iterator, with each iteration returning a tensor of a segment.
**Returns**
- **tensor** —— The tensor of the video segment.
## 8. Module Registration SDK (scepter.utils.registry)
## 8. Module Registration SDK (scepter.modules.utils.registry)
Used for managing various registered classes.
### Basic Usage
```python
from scepter.utils.registry import Registry
from scepter.utils.config import Config
from scepter.modules.utils.registry import Registry
from scepter.modules.utils.config import Config
MODELS = Registry('MODELS')
@@ -614,14 +614,14 @@ Register a function
**Returns**
- **name** —— Registration name.
## 9. Data SDK(scepter.utils.data)
## 9. Data SDK(scepter.modules.utils.data)
Used for transferring data between devices
### Basic Usage
```python
import torch
from scepter.utils.data import transfer_data_to_numpy, transfer_data_to_cpu, transfer_data_to_cuda
from scepter.modules.utils.data import transfer_data_to_numpy, transfer_data_to_cpu, transfer_data_to_cuda
data = {"a": torch.Tensor([0])}
transfer_data_to_numpy(data)
@@ -668,7 +668,7 @@ Used for operations such as loading and evaluating models
```python
import torch
from scepter.utils.model import move_model_to_cpu, load_pretrained,
from scepter.modules.utils.model import move_model_to_cpu, load_pretrained,
count_params, init_weights
```
<hr/>
@@ -716,14 +716,14 @@ Initialize the parameters of the model modules.
**Parameters**
- **module** —— The torch.nn.Module model instance.
## 11. Sampler SDK(scepter.utils.sampler)
## 11. Sampler SDK(scepter.modules.utils.sampler)
Samplers are quite universal, and in most cases, custom development is not required. Here are provided several common types of sampler.
### Basic Usage
```python
import torch
from scepter.utils.sampler import MultiFoldDistributedSampler,
from scepter.modules.utils.sampler import MultiFoldDistributedSampler,
EvalDistributedSampler, MultiLevelBatchSampler, MixtureOfSamplers
```
<hr/>
@@ -830,17 +830,17 @@ A sampler for multi-level indexing of large-scale data.
Iterator, each iteration returns an index of a sample.
## 12. Prober SDK(scepter.utils.probe)
## 12. Prober SDK(scepter.modules.utils.probe)
Used for probing variable statistics of various components.
### Basic Usage
```python
import numpy as np
from scepter.model.base_model import BaseModel
from scepter.utils.config import Config
from scepter.utils.file_system import FS
from scepter.utils.probe import ProbeData
from scepter.modules.model.base_model import BaseModel
from scepter.modules.utils.config import Config
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.probe import ProbeData
class TestModel(BaseModel):
+133
View File
@@ -0,0 +1,133 @@
<h1 align="center"> Locate, Assign, Refine: Taming Customized Image Inpainting with Text-Subject Guidance </h1>
<p align="center">
<strong>Yulin Pan</strong>
·
<strong>Chaojie Mao</strong>
·
<strong>Zeyinzi Jiang</strong>
·
<strong>Zhen Han</strong>
·
<strong>Jingfeng Zhang</strong>
<br>
<a href="https://arxiv.org/abs/2403.19534"><img src="https://img.shields.io/static/v1?label=arXiv&message=LARGen&color=red&logo=arxiv"></a>
<a href="https://ali-vilab.github.io/largen-page/"><img src="https://img.shields.io/badge/Page-LARGen-Gree"></a>
</p>
LARGen is a unified image inpainting framework that supports text-guided, subject-guided and text-subject-guided inpainting simutaneously.
Four LARGen-based fantastic applications are now supported by SCEPTER Studio:
1. Zoom Out
2. Virtual Try On
3. Text-Guided Inpainting
4. Text-Subject-Guided Inpainting
## Basic Usage
Here's a demo showcasing the use of LARGen-based functions.
<p align="left">
<img src="https://raw.githubusercontent.com/ali-vilab/largen-page/main/public/images/largen.gif" width="1300">
</p>
## 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>
## Features
| **Model** | **Locate** | **Assign** | **Refine** |
|:---------:|:----------:|:----------:|:----------:|
| SD v1.5 | ⏳ | ⏳ | ⏳ |
| SD XL | 🪄 | 🪄 | ⏳ |
- 🪄 denotes that the feature has been supported.
- ⏳ denotes that the feature has not been integrated currently.
## Pretrained Models
| **Model** | **URL** |
|:----------:|:-------:|
| largen-sdxl-s22k | [ModelScope](https://www.modelscope.cn/models/iic/LARGEN/summary) |
## BibTeX
If our work is useful for your research, please consider citing:
```bibtex
@article{pan2024locate,
title={Locate, Assign, Refine: Taming Customized Image Inpainting with Text-Subject Guidance},
author={Pan, Yulin and Mao, Chaojie and Jiang, Zeyinzi and Han, Zhen and Zhang, Jingfeng},
journal={arXiv preprint arXiv:2403.19534},
year={2024}
}
```
+125
View File
@@ -0,0 +1,125 @@
<p align="center">
<h2 align="center">SCEdit: Efficient and Controllable Image Diffusion Generation via Skip Connection Editing</h2>
<h3 align="center">(CVPR 2024 Highlight)</h3>
<p align="center">
<strong>Zeyinzi Jiang</strong>
·
<strong>Chaojie Mao</strong>
·
<strong>Yulin Pan</strong>
·
<strong>Zhen Han</strong>
·
<strong>Jingfeng Zhang</strong>
<br>
<b>Alibaba Group</b>
<br>
<a href="https://arxiv.org/abs/2312.11392"><img src='https://img.shields.io/badge/arXiv-SCEdit-red' alt='Paper PDF'></a>
<a href='https://scedit.github.io/'><img src='https://img.shields.io/badge/Project_Page-SCEdit-green' alt='Project Page'></a>
<a href='https://github.com/modelscope/scepter'><img src='https://img.shields.io/badge/scepter-SCEdit-yellow'></a>
<a href='https://github.com/modelscope/swift'><img src='https://img.shields.io/badge/swift-SCEdit-blue'></a>
<br>
</p>
SCEdit is an efficient generative fine-tuning framework proposed by Alibaba TongYi Vision Intelligence Lab. This framework enhances the fine-tuning capabilities for text-to-image generation downstream tasks and enables quick adaptation to specific generative scenarios, **saving 30%-50% of training memory costs compared to LoRA**. Furthermore, it can be directly extended to controllable image generation tasks, **requiring only 7.9% of the parameters that ControlNet needs for conditional generation and saving 30% of memory usage**. It supports various conditional generation tasks including edge maps, depth maps, segmentation maps, poses, color maps, and image completion.
## Usage
### Text-to-Image Generation
```shell
# SD v1.5
python scepter/tools/run_train.py --cfg scepter/methods/scedit/t2i/sd15_512_sce_t2i.yaml
# SD v2.1
python scepter/tools/run_train.py --cfg scepter/methods/scedit/t2i/sd21_768_sce_t2i.yaml
# SD XL
python scepter/tools/run_train.py --cfg scepter/methods/scedit/t2i/sdxl_1024_sce_t2i.yaml
```
### Controllable Image Synthesis
```shell
# SD v1.5 + hed
python scepter/tools/run_train.py --cfg scepter/methods/scedit/ctr/sd15_512_sce_ctr_hed.yaml
# SD v2.1 + canny
python scepter/tools/run_train.py --cfg scepter/methods/scedit/ctr/sd21_768_sce_ctr_canny.yaml
# SD XL + depth
python scepter/tools/run_train.py --cfg scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_depth.yaml
```
### Gradio
```shell
python -m scepter.tools.webui # Then click [Use Tuners] or [Use Controller]
```
## Models
### Model URL
| Model | URL |
|--------|-------------------------------------------------------------------------------------------------------------------------------------------|
| SCEdit | [ModelScope](https://modelscope.cn/models/iic/scepter_scedit/summary) [HuggingFace](https://huggingface.co/scepter-studio/scepter_scedit) |
### Text-to-Image Generation
| **Model** | **SCEdit** |
|:---------:|:----------:|
| SD 1.5 | 🪄 |
| SD 2.1 | 🪄 |
| SD XL | 🪄 |
### Controllable Image Synthesis
| **Model** | **Canny** | **HED** | **Depth** | **Pose** | **Color** |
|:---------:|:---------:|:-------:|:---------:|:--------:|:---------:|
| SD 2.1 | 🪄 | 🪄 | 🪄 | 🪄 | 🪄 |
| SD XL | 🪄 | 🪄 | 🪄 | 🪄 | 🪄 |
## Application Gallery
### 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="../../../asset/images/scedit/tuner_gold_dragon.jpeg" width="300"></td>
<td><img src="../../../asset/images/scedit/tuner_sloppy_dragon.jpeg" width="300"></td>
<td><img src="../../../asset/images/scedit/tuner_mantra_papercraft_dragon.jpeg" width="300"></td>
<td><img src="../../../asset/images/scedit/tuner_pose.jpeg" 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="../../../asset/images/scedit/word_condition.png" width="300"></td>
<td><img src="../../../asset/images/scedit/word_race.jpeg" width="300"></td>
<td><img src="../../../asset/images/scedit/word_lilies.jpeg" width="300"></td>
<td><img src="../../../asset/images/scedit/word_festival.jpeg" width="300"></td>
</tr>
</table>
## BibTeX
```bibtex
@article{jiang2023scedit,
title = {SCEdit: Efficient and Controllable Image Diffusion Generation via Skip Connection Editing},
author = {Jiang, Zeyinzi and Mao, Chaojie and Pan, Yulin and Han, Zhen and Zhang, Jingfeng},
year = {2023},
journal = {arXiv preprint arXiv:2312.11392}
}
```
+103
View File
@@ -0,0 +1,103 @@
# StyleBooth: Image Style Editing with Multimodal Instruction
Zhen Han, Chaojie Mao, Zeyinzi Jiang, Yulin Pan, Jingfeng Zhang
Alibaba Group
[[paper](https://arxiv.org/abs/2404.12154)][[Model](https://modelscope.cn/models/iic/stylebooth/summary)] [[Dataset](https://modelscope.cn/models/iic/stylebooth/summary)]
## Abstract
Given an original image, image editing aims to generate an image that align with the provided instruction. The challenges are to accept multimodal inputs as instructions and a scarcity of high-quality training data, including crucial triplets of source/target image pairs and multimodal (text and image) instructions. In this paper, we focus on image style editing and present <strong>StyleBooth</strong>, a method that proposes a comprehensive framework for image editing and a feasible strategy for building a high-quality style editing dataset. We integrate encoded textual instruction and image exemplar as a unified condition for diffusion model, enabling the editing of original image following <strong>multimodal instructions</strong>. Furthermore, by <strong>iterative style-destyle tuning and editing</strong> and usability filtering, the StyleBooth dataset provides content-consistent stylized/plain image pairs in various categories of styles. To show the flexibility of StyleBooth, we conduct experiments on diverse tasks, such as textbased style editing, exemplar-based style editing and compositional style editing. The results demonstrate that the quality and variety of training data significantly enhance the ability to preserve content and improve the overall quality of generated images in editing tasks.
![head](https://ali-vilab.github.io/stylebooth-page/public/images/head.jpg "head")
## Gallery
<table>
<tr>
<td><strong>Origin Image</strong><br>Gold Dragon Tuner</td>
<td><strong>Graffiti Art</strong></td>
<td><strong>Adorable Kawaii</strong></td>
<td><strong>game-retro game</strong></td>
<td><strong>Vincent van Gogh</strong></td>
</tr>
<tr>
<td><img src="../../../asset/images/scedit/tuner_gold_dragon.jpeg" width="240"></td>
<td><img src="../../../asset/images/stylebooth/graffiti.jpeg" width="240"></td>
<td><img src="../../../asset/images/stylebooth/kawaii.jpeg" width="240"></td>
<td><img src="../../../asset/images/stylebooth/retrogame.jpeg" width="240"></td>
<td><img src="../../../asset/images/stylebooth/vangogh.jpeg" width="240"></td>
</tr>
</table>
## Features
| **Text-Based** | **Exemplar-Based** |
|:--------------:|:-----------------:|
| 🪄 | ⏳ |
- ✅ 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.
## Run StyleBooth
- Code implementation: See model configuration and code based on [🪄SCEPTER](https://github.com/modelscope/scepter).
- Demo: Try [🖥️SCEPTER Studio](https://github.com/modelscope/scepter/tree/main?tab=readme-ov-file#%EF%B8%8F-scepter-studio).
- Easy run:
Try the following example script to run StyleBooth modified from [tests/modules/test_diffusion_inference.py](https://github.com/modelscope/scepter/blob/main/tests/modules/test_diffusion_inference.py):
```python
# `pip install scepter>0.0.4` or
# clone newest SCEPTER and run `PYTHONPATH=./ python <this_script>` at the main branch root.
import os
import unittest
from PIL import Image
from torchvision.utils import save_image
from scepter.modules.inference.stylebooth_inference import StyleboothInference
from scepter.modules.utils.config import Config
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.logger import get_logger
class DiffusionInferenceTest(unittest.TestCase):
def setUp(self):
print(('Testing %s.%s' % (type(self).__name__, self._testMethodName)))
self.logger = get_logger(name='scepter')
config_file = 'scepter/methods/studio/scepter_ui.yaml'
cfg = Config(cfg_file=config_file)
if 'FILE_SYSTEM' in cfg:
for fs_info in cfg['FILE_SYSTEM']:
FS.init_fs_client(fs_info)
self.tmp_dir = './cache/save_data/diffusion_inference'
if not os.path.exists(self.tmp_dir):
os.makedirs(self.tmp_dir)
def tearDown(self):
super().tearDown()
# uncomment this line to skip this module.
# @unittest.skip('')
def test_stylebooth(self):
config_file = 'scepter/methods/studio/inference/edit/stylebooth_tb_pro.yaml'
cfg = Config(cfg_file=config_file)
diff_infer = StyleboothInference(logger=self.logger)
diff_infer.init_from_cfg(cfg)
output = diff_infer({'prompt': 'Let this image be in the style of sai-lowpoly'},
style_edit_image=Image.open('asset/images/inpainting_text_ref/ex4_scene_im.jpg'),
style_guide_scale_text=7.5,
style_guide_scale_image=0.5)
save_path = os.path.join(self.tmp_dir,
'stylebooth_test_lowpoly_cute_dog.png')
save_image(output['images'], save_path)
if __name__ == '__main__':
unittest.main()
```
+26
View File
@@ -0,0 +1,26 @@
<h1 align="center">Dataset Management</h1>
SCEPTER supports three types of dataset formats: TXT, CSV, and ModelScope.
Below are examples for each format, illustrating their details and basic usage.
## Modelscope Format
We use a [custom-stylized dataset](https://modelscope.cn/datasets/damo/style_custom_dataset/summary), which included classes 3D, anime, flat illustration, oil painting, sketch, and watercolor, each with 30 image-text pairs.
```python
# pip install modelscope
from modelscope.msdatasets import MsDataset
ms_train_dataset = MsDataset.load('style_custom_dataset', namespace='damo', subset_name='3D', split='train_short')
print(next(iter(ms_train_dataset)))
```
## CSV Format
For the data format used by SCEPTER Studio, please refer to [3D_example_csv.zip](https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=datasets/3D_example_csv.zip).
## TXT Format
To facilitate starting training in command-line mode, you can use a dataset in text format, please refer to [3D_example_txt.zip](https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=datasets/3D_example_txt.zip)
```shell
mkdir -p cache/datasets/ && wget 'https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=datasets/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
+45
View File
@@ -0,0 +1,45 @@
# Inference
In this tutorial, we'll cover the use of the scepter framework for convenient inference, including inference using the command line or specific method classes, and we'll give examples of inference methods for additional tasks.
## Command Line
Inference of SDXL generation models using the command line.
```shell
python scepter/tools/run_inference.py --cfg scepter/methods/examples/generation/stable_diffusion_xl_1024.yaml --prompt 'a cute dog' --save_folder 'inference' # generation on SD XL
```
## Class Instantiation
Inference of SD2.1 generation models using the class instantiation.
```python
from torchvision.utils import save_image
from scepter.modules.utils.config import Config
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.logger import get_logger
from scepter.modules.inference.diffusion_inference import DiffusionInference
# init file system - modelscope
FS.init_fs_client(Config(load=False, cfg_dict={'NAME': 'ModelscopeFs', 'TEMP_DIR': 'cache/data'}))
# init model config
logger = get_logger(name='scepter')
cfg = Config(cfg_file='scepter/methods/studio/inference/stable_diffusion/sd21_pro.yaml')
diff_infer = DiffusionInference(logger)
diff_infer.init_from_cfg(cfg)
# start inference
output = diff_infer({'prompt': 'a cute dog'})
save_image(output['images'], 'sd21_test_prompt_a_cute_dog.png')
```
## Additional Tasks
### Fine-tuned Model Inference
```shell
python scepter/tools/run_inference.py --cfg scepter/methods/scedit/t2i/sd15_512_sce_t2i_swift.yaml --pretrained_model 'cache/save_data/sd15_512_sce_t2i_swift/checkpoints/ldm_step-100.pth' --prompt 'A close up of a small rabbit wearing a hat and scarf' --save_folder 'trained_test_prompt_rabbit'
```
### Controllable Image Synthesis Inference
- SCEdit
```shell
python scepter/tools/run_inference.py --cfg scepter/methods/scedit/ctr/sd21_768_sce_ctr_canny.yaml --num_samples 1 --prompt 'a single flower is shown in front of a tree' --save_folder 'test_flower_canny' --image_size 768 --task control --image 'asset/images/flower.jpg' --control_mode canny --pretrained_model ms://damo/scepter_scedit@controllable_model/SD2.1/canny_control/0_SwiftSCETuning/pytorch_model.bin # canny
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
```
+84
View File
@@ -0,0 +1,84 @@
# Training
We provide a framework for training and validation.
The scripts below are just for illustration purposes. To achieve better results, you can modify the corresponding parameters as needed.
## Start Training
There are different ways to start a training:
- calling scepter/tools/run_train.py:
```bash
# calling at SCEPTER root:
PYTHONPATH=./ python scepter/tools/run_train.py --cfg [path-to-your-yaml]
# calling scepter library:
pip install scepter
python -m scepter.tools.run_train --cfg [path-to-your-yaml]
```
- calling your own script:
```bash
# calling at SCEPTER root:
PYTHONPATH=./ python [path-to-your-script] --cfg [path-to-your-yaml]
# calling scepter library:
pip install scepter
python [path-to-your-script] --cfg [path-to-your-yaml]
```
your scepter should be like:
```python
from scepter.tools.run_train import run
if __name__ == '__main__':
run()
```
## Popular Tasks
### Text-to-Image Generation
- SCEdit
```bash
python scepter/tools/run_train.py --cfg scepter/methods/scedit/t2i/sd15_512_sce_t2i.yaml # SD v1.5
python scepter/tools/run_train.py --cfg scepter/methods/scedit/t2i/sd21_768_sce_t2i.yaml # SD v2.1
python scepter/tools/run_train.py --cfg scepter/methods/scedit/t2i/sdxl_1024_sce_t2i.yaml # SD XL
```
- Existing Tuning Strategies
```bash
python scepter/tools/run_train.py --cfg scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml # fully-tuning on SD v1.5
python scepter/tools/run_train.py --cfg scepter/methods/examples/generation/stable_diffusion_2.1_768_lora.yaml # lora-tuning on SD v2.1
```
- Data Text Format
```bash
# Download the 3D_example_txt.zip as previously mentioned
python scepter/tools/run_train.py --cfg scepter/methods/scedit/t2i/sdxl_1024_sce_t2i_datatxt.yaml
```
### Controllable Image Synthesis
- SCEdit
The YAML configuration can be modified to combine different base models and conditions. The following is provided as an example.
```bash
python scepter/tools/run_train.py --cfg scepter/methods/scedit/ctr/sd15_512_sce_ctr_hed.yaml # SD v1.5 + hed
python scepter/tools/run_train.py --cfg scepter/methods/scedit/ctr/sd21_768_sce_ctr_canny.yaml # SD v2.1 + canny
python scepter/tools/run_train.py --cfg scepter/methods/scedit/ctr/sd21_768_sce_ctr_pose.yaml # SD v2.1 + pose
python scepter/tools/run_train.py --cfg scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_depth.yaml # SD XL + depth
python scepter/tools/run_train.py --cfg scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_color.yaml # SD XL + color
```
- Data Text Format
```bash
# Download the 3D_example_txt.zip as previously mentioned
python scepter/tools/run_train.py --cfg scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_color_datatxt.yaml
```
## Customize Modules
You can register your own Modules like DATASET, SAMPLERS, TRANSFORMS, MODELS, SOVLERS, HOOKS, OPTIMIZERS into SCEPTER.
Refer to `example/`, build the modules of your task in `example/{task}`.
```bash
cd example/classifier
python run.py --cfg classifier.yaml
```
+17 -17
View File
@@ -15,8 +15,8 @@
子类注册:
```python
from scepter.model.registry import BACKBONES
from scepter.model.base_model import BaseModel
from scepter.modules.model.registry import BACKBONES
from scepter.modules.model.base_model import BaseModel
@BACKBONES.register_class("ResNet")
@@ -26,8 +26,8 @@ class ResNet(BaseModel):
```
```python
from scepter.model.registry import NECKS
from scepter.model.base_model import BaseModel
from scepter.modules.model.registry import NECKS
from scepter.modules.model.base_model import BaseModel
@NECKS.register_class()
@@ -37,8 +37,8 @@ class GlobalAveragePooling(BaseModel):
```
```python
from scepter.model.registry import HEADS
from scepter.model.base_model import BaseModel
from scepter.modules.model.registry import HEADS
from scepter.modules.model.base_model import BaseModel
@HEADS.register_class()
@@ -48,7 +48,7 @@ class ClassifierHead(BaseModel):
```
```python
from scepter.model.registry import LOSSES
from scepter.modules.model.registry import LOSSES
import torch.nn as nn
@@ -60,7 +60,7 @@ class CrossEntropy(nn.Module):
实际调用:
```python
from scepter.model.registry import BACKBONES, NECKS, HEADS, LOSSES, TUNERS
from scepter.modules.model.registry import BACKBONES, NECKS, HEADS, LOSSES, TUNERS
backbone = BACKBONES.build(cfg.BACKBONE, logger=logger)
neck = NECKS.build(cfg.NECK, logger=logger)
@@ -85,8 +85,8 @@ tuner = TUNERS.build(cfg.TUNER, logger=logger)
子类注册:
```python
from scepter.model.metrics.registry import METRICS
from scepter.model.metrics.base_metric import BaseMetric
from scepter.modules.model.metrics.registry import METRICS
from scepter.modules.model.metrics.base_metric import BaseMetric
@METRICS.register_class("AccuracyMetric")
@@ -97,7 +97,7 @@ class AccuracyMetric(BaseMetric):
实际用法:
```python
from scepter.model.metrics.registry import METRICS
from scepter.modules.model.metrics.registry import METRICS
metric = METRICS.build(cfgs, logger)
```
@@ -119,8 +119,8 @@ metric = METRICS.build(cfgs, logger)
子类注册:
```python
from scepter.model.registry import TOKENIZERS
from scepter.model.tokenizers import BaseTokenizer
from scepter.modules.model.registry import TOKENIZERS
from scepter.modules.model.tokenizers import BaseTokenizer
@TOKENIZERS.register_class()
@@ -131,7 +131,7 @@ class BaseBertTokenizer(BaseTokenizer):
实际用法:
```python
from scepter.model.registry import TOKENIZERS
from scepter.modules.model.registry import TOKENIZERS
tokenizer = TOKENIZERS.build(cfgs, logger)
```
@@ -149,8 +149,8 @@ tokenizer = TOKENIZERS.build(cfgs, logger)
子类注册:
```python
from scepter.model.registry import MODELS
from scepter.model.networks.train_module import TrainModule
from scepter.modules.model.registry import MODELS
from scepter.modules.model.networks.train_module import TrainModule
@MODELS.register_class()
@@ -161,7 +161,7 @@ class Classifier(TrainModule):
实际用法:
```python
from scepter.model.registry import MODELS
from scepter.modules.model.registry import MODELS
model = MODELS.build(self.cfg.MODEL, logger=self.logger)
```
+4 -4
View File
@@ -9,8 +9,8 @@
子lr_schedulers继承时用法:
```python
from scepter.opt.lr_schedulers import LR_SCHEDULERS
from scepter.opt.lr_schedulers.base_scheduler import BaseScheduler
from scepter.modules.opt.lr_schedulers import LR_SCHEDULERS
from scepter.modules.opt.lr_schedulers.base_scheduler import BaseScheduler
@LR_SCHEDULERS.register_class()
@@ -48,8 +48,8 @@ lr_schedulers的基类,支持注册操作,可根据需要自定义;
子optimizers继承时用法:
```python
from scepter.opt.optimizers.base_optimizer import BaseOptimize
from scepter.opt.optimizers.registry import OPTIMIZERS
from scepter.modules.opt.optimizers.base_optimizer import BaseOptimize
from scepter.modules.opt.optimizers.registry import OPTIMIZERS
@OPTIMIZERS.register_class()
+6 -6
View File
@@ -6,10 +6,10 @@
支持3类文件IO Handler:
1. scepter.utils.file_clients.AliyunOssFs
2. scepter.utils.file_clients.LocalFs
3. scepter.utils.file_clients.HttpFs
4. scepter.utils.file_clients.ModelscopeFs
1. scepter.modules.utils.file_clients.AliyunOssFs
2. scepter.modules.utils.file_clients.LocalFs
3. scepter.modules.utils.file_clients.HttpFs
4. scepter.modules.utils.file_clients.ModelscopeFs
<hr/>
@@ -17,8 +17,8 @@
## 基础用法
```python
from scepter.utils.file_system import FS
from scepter.utils.config import Config
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.config import Config
fs_cfg = Config(load=False, cfg_dict={
"NAME": "AliyunOssFs",
+39 -39
View File
@@ -3,18 +3,18 @@
依赖SDK,该部分用于对框架全局经常复用的模块和sdk进行整理,并根据功能相关性进行聚合。
## 总览
1. 参数sdk(scepter.utils.config)
2. 路径sdk(scepter.utils.directory)
3. torch分布式sdk(scepter.utils.distribute)
4. 模型导出sdk(scepter.utils.export_model)
5. 文件系统sdk(scepter.utils.file_system)
6. 日志sdk(scepter.utils.logger)
7. 视频处理sdk(scepter.utils.video_reader),文档参考(video_reader.md)
8. 模块注册sdk(scepter.utils.registry)
9. 数据sdk(scepter.utils.data)
10. 模型sdk(scepter.utils.model)
11. 采样器sdk(scepter.utils.sampler)
12. 探针器sdk(scepter.utils.probe)
1. 参数sdk(scepter.modules.utils.config)
2. 路径sdk(scepter.modules.utils.directory)
3. torch分布式sdk(scepter.modules.utils.distribute)
4. 模型导出sdk(scepter.modules.utils.export_model)
5. 文件系统sdk(scepter.modules.utils.file_system)
6. 日志sdk(scepter.modules.utils.logger)
7. 视频处理sdk(scepter.modules.utils.video_reader),文档参考(video_reader.md)
8. 模块注册sdk(scepter.modules.utils.registry)
9. 数据sdk(scepter.modules.utils.data)
10. 模型sdk(scepter.modules.utils.model)
11. 采样器sdk(scepter.modules.utils.sampler)
12. 探针器sdk(scepter.modules.utils.probe)
<hr/>
@@ -23,7 +23,7 @@
### 基础用法
```python
from scepter.utils.config import Config
from scepter.modules.utils.config import Config
# 从一个dict对象 初始化 Config对象
fs_cfg = Config(load=False, cfg_dict={"NAME": "LocalFs"})
@@ -97,7 +97,7 @@ print(fs_cfg.args)
### 基础用法
```python
from scepter.utils.directory import osp_path
from scepter.modules.utils.directory import osp_path
# 根据路径前缀进行自动化路径拼接
prefix = "xxxx"
@@ -108,7 +108,7 @@ print(osp_path(prefix, data_file))
data_file = "xxxx/example_videos/1.mp4"
print(osp_path(prefix, data_file))
from scepter.utils.directory import get_relative_folder
from scepter.modules.utils.directory import get_relative_folder
# 根据路径获取指定层级的文件夹路径
# 默认最后一级 xxxx/example_videos/
@@ -116,7 +116,7 @@ print(get_relative_folder(data_file))
# 倒数第二级 xxxx/
print(get_relative_folder(data_file, keep_index=-2))
from scepter.utils.directory import get_md5
from scepter.modules.utils.directory import get_md5
# 获取文本/路径的md5码 34a447fb46d0b786a3999c9dad01d470
print(get_md5(data_file))
@@ -172,8 +172,8 @@ torch分布式初始化sdk,使用该sdk,可以让用户不要关注torch的
### 基础用法
```python
from scepter.utils.distribute import we
from scepter.utils.config import Config
from scepter.modules.utils.distribute import we
from scepter.modules.utils.config import Config
cfg = Config(cfg_dict={}, load=False)
@@ -304,12 +304,12 @@ we.init_env(cfg, fn, logger=None)
**Returns**
- **tensor** —— 输出的在进程rank=0上的cpu的tensor。
## 4. 模型导出sdk(scepter.utils.export_model)
## 4. 模型导出sdk(scepter.modules.utils.export_model)
用于模型导出为torchscript/Onnx格式的api。
### 基础用法
```python
from scepter.utils.export_model import save_develop_model_multi_io
from scepter.modules.utils.export_model import save_develop_model_multi_io
save_develop_model_multi_io(
model,
@@ -347,16 +347,16 @@ input_type 一一对应。
**Returns**
- **tensor** —— 输出的在进程rank=0上的cpu的tensor。
## 5. 文件系统sdk(scepter.utils.file_system)
## 5. 文件系统sdk(scepter.modules.utils.file_system)
参考[file_clients](file_clients.md)
## 6. 日志sdk(scepter.utils.logger)
## 6. 日志sdk(scepter.modules.utils.logger)
用于实例化一个标准的日志实例,用于打印信息。
### 基础用法
```python
from scepter.utils.logger import get_logger, init_logger
from scepter.modules.utils.logger import get_logger, init_logger
std_logger = get_logger(name="scepter")
init_logger(std_logger, log_file="", dist_launcher="pytorch")
@@ -407,14 +407,14 @@ init_logger(std_logger, log_file="", dist_launcher="pytorch")
**Returns**
- **str** —— 格式化的输出。
## 7. 视频处理sdk(scepter.utils.video_reader)
## 7. 视频处理sdk(scepter.modules.utils.video_reader)
用于处理视频读取的api。
### 基础用法
```python
from scepter.utils.video_reader.frame_sampler import do_frame_sample
from scepter.utils.video_reader.video_reader import (
from scepter.modules.utils.video_reader.frame_sampler import do_frame_sample
from scepter.modules.utils.video_reader.video_reader import (
VideoReaderWrapper, EasyVideoReader, FramesReaderWrapper
)
```
@@ -556,14 +556,14 @@ overlap: Union[float, Fraction, str] = Fraction(0), transforms: Optional[Callabl
**Returns**
- **tensor** —— 视频片段的tensor。
## 8. 模块注册sdk(scepter.utils.registry)
## 8. 模块注册sdk(scepter.modules.utils.registry)
用于管理各种注册的类。
### 基础用法
```python
from scepter.utils.registry import Registry
from scepter.utils.config import Config
from scepter.modules.utils.registry import Registry
from scepter.modules.utils.config import Config
MODELS = Registry('MODELS')
@@ -616,14 +616,14 @@ build目标类的实例
**Returns**
- **name** —— 注册名称。
## 9. 数据sdk(scepter.utils.data)
## 9. 数据sdk(scepter.modules.utils.data)
用于数据在设备间转移
### 基础用法
```python
import torch
from scepter.utils.data import transfer_data_to_numpy, transfer_data_to_cpu, transfer_data_to_cuda
from scepter.modules.utils.data import transfer_data_to_numpy, transfer_data_to_cpu, transfer_data_to_cuda
data = {"a": torch.Tensor([0])}
transfer_data_to_numpy(data)
@@ -670,7 +670,7 @@ transfer_data_to_cuda(data)
```python
import torch
from scepter.utils.model import move_model_to_cpu, load_pretrained,
from scepter.modules.utils.model import move_model_to_cpu, load_pretrained,
count_params, init_weights
```
<hr/>
@@ -718,14 +718,14 @@ from scepter.utils.model import move_model_to_cpu, load_pretrained,
**Parameters**
- **module** —— torch.nn.Module模型实例。
## 11. 采样器sdk(scepter.utils.sampler)
## 11. 采样器sdk(scepter.modules.utils.sampler)
采样器比较具有通用性,大多数情况下不会进行定制开发,这里提供了几类常用的sampler采样器。
### 基础用法
```python
import torch
from scepter.utils.sampler import MultiFoldDistributedSampler,
from scepter.modules.utils.sampler import MultiFoldDistributedSampler,
EvalDistributedSampler, MultiLevelBatchSampler, MixtureOfSamplers
```
<hr/>
@@ -832,17 +832,17 @@ from scepter.utils.sampler import MultiFoldDistributedSampler,
迭代器,每迭代一次得到一个样本的index
## 12. 探针器sdk(scepter.utils.probe)
## 12. 探针器sdk(scepter.modules.utils.probe)
用于探针各个组件的变量统计
### 基础用法
```python
import numpy as np
from scepter.model.base_model import BaseModel
from scepter.utils.config import Config
from scepter.utils.file_system import FS
from scepter.utils.probe import ProbeData
from scepter.modules.model.base_model import BaseModel
from scepter.modules.utils.config import Config
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.probe import ProbeData
class TestModel(BaseModel):
+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()
+92 -189
View File
@@ -8,48 +8,63 @@
<a href="https://github.com/modelscope/scepter/"><img src="https://img.shields.io/badge/scepter-Build from source-6FEBB9.svg"></a>
</p>
## 📖 Table of Contents
- [Introduction](#-introduction)
- [News](#-news)
- [Installation](#%EF%B8%8F-installation)
- [Getting Started](#-getting-started)
- [SCEPTER Studio](#%EF%B8%8F-scepter-studio)
- [Gallery](#%EF%B8%8F-gallery)
- [Features](#-features)
- [Learn More](#-learn-more)
- [License](#license)
🪄SCEPTER is an open-source code repository dedicated to generative training, fine-tuning, and inference, encompassing a suite of downstream tasks such as image generation, transfer, editing.
SCEPTER integrates popular community-driven implementations as well as proprietary methods by Tongyi Lab of Alibaba Group, offering a comprehensive toolkit for researchers and practitioners in the field of AIGC. This versatile library is designed to facilitate innovation and accelerate development in the rapidly evolving domain of generative models.
## 📝 Introduction
SCEPTER offers 3 core components:
- [Generative training and inference framework](#tutorials)
- [Easy implementation of popular approaches](#currently-supported-approaches)
- [Interactive user interface: SCEPTER Studio](#launch)
SCEPTER is an open-source code repository dedicated to generative training, fine-tuning, and inference, encompassing a suite of downstream tasks such as image generation, transfer, editing. It integrates popular community-driven implementations as well as proprietary methods by Tongyi Lab of Alibaba Group, offering a comprehensive toolkit for researchers and practitioners in the field of AIGC. This versatile library is designed to facilitate innovation and accelerate development in the rapidly evolving domain of generative models.
Main Feature:
- Task:
- Text-to-image generation
- Controllable image synthesis
- Image editing (TODO)
- Training / Inference:
- Distribute: DDP / FSDP / FairScale / Xformers
- File system: Local / Http / OSS / Modelscope
- Deploy:
- Data management
- Training
- Inference
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.04]: New [StyleBooth](https://ali-vilab.github.io/stylebooth-page/) demo on SCEPTER Studio for`Text-Based Style Editing`.
- [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.
## 🖼 Gallery for Recent Works
### StyleBooth
<table>
<tr>
<td><strong>Origin Image</strong><br>Gold Dragon Tuner</td>
<td><strong>Graffiti Art</strong></td>
<td><strong>Adorable Kawaii</strong></td>
<td><strong>game-retro game</strong></td>
<td><strong>Vincent van Gogh</strong></td>
</tr>
<tr>
<td><img src="asset/images/scedit/tuner_gold_dragon.jpeg" width="240"></td>
<td><img src="asset/images/stylebooth/graffiti.jpeg" width="240"></td>
<td><img src="asset/images/stylebooth/kawaii.jpeg" width="240"></td>
<td><img src="asset/images/stylebooth/retrogame.jpeg" width="240"></td>
<td><img src="asset/images/stylebooth/vangogh.jpeg" width="240"></td>
</tr>
</table>
<table>
<tr>
<td><strong>Origin Image</strong></td>
<td><strong>Lowpoly</strong></td>
<td><strong>Colored Pencil Art</strong></td>
<td><strong>Watercolor</strong></td>
<td><strong>misc-disco</strong></td>
</tr>
<tr>
<td><img src="asset/images/stylebooth/mountain.jpg" width="240"></td>
<td><img src="asset/images/stylebooth/lowpoly.jpg" width="240"></td>
<td><img src="asset/images/stylebooth/colorpencil.jpeg" width="240"></td>
<td><img src="asset/images/stylebooth/watercolor.jpeg" width="240"></td>
<td><img src="asset/images/stylebooth/disco.jpeg" width="240"></td>
</tr>
</table>
## 🛠️ Installation
- Create new environment
@@ -70,100 +85,32 @@ pip install -r requirements/recommended.txt
pip install scepter
```
## 🚀 Getting Started
## 🧩 Generative Framework
### Dataset
### Tutorials
#### Modelscope Format
| Documentation | Key Features |
|:---------------------------------------------------|:----------------------------------|
| [Train](docs/en/tutorials/train.md) | DDP / FSDP / FairScale / Xformers |
| [Inference](docs/en/tutorials/inference.md) | Dynamic load/unload |
| [Dataset Management](docs/en/tutorials/dataset.md) | Local / Http / OSS / Modelscope |
We use a [custom-stylized dataset](https://modelscope.cn/datasets/damo/style_custom_dataset/summary), which included classes 3D, anime, flat illustration, oil painting, sketch, and watercolor, each with 30 image-text pairs.
```python
# pip install modelscope
from modelscope.msdatasets import MsDataset
ms_train_dataset = MsDataset.load('style_custom_dataset', namespace='damo', subset_name='3D', split='train_short')
print(next(iter(ms_train_dataset)))
```
## 📝 Popular Approaches
#### CSV Format
### Currently supported approaches
For the data format used by SCEPTER Studio, please refer to [3D_example_csv.zip](https://modelscope.cn/api/v1/models/damo/scepter/repo?Revision=master&FilePath=datasets/3D_example_csv.zip).
| Tasks | Methods | Links |
|:----------------------------:|:--------------------------------------------:|:------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
| Text-to-image generation | SD v1.5 | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
| Text-to-image generation | SD v2.1 | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/runwayml/stable-diffusion-v1-5) |
| Text-to-image generation | SD-XL | [![Hugging Face Repo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Repo-blue)](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0) |
| Efficient Tuning | LoRA | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=LoRA&color=red&logo=arxiv)](https://arxiv.org/abs/2106.09685) |
| Efficient Tuning | Res-Tuning(NeurIPS23) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=Res-Tuing&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/) |
| Controllable image synthesis | [🌟SCEdit(CVPR24)](docs/en/tasks/scedit.md) | [![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/) |
| Image editing | [🌟LAR-Gen](docs/en/tasks/largen.md) | [![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/) |
| Image editing | [🌟StyleBooth](docs/en/tasks/stylebooth.md) | [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=StyleBooth&color=red&logo=arxiv)](https://arxiv.org/abs/2404.12154) [![Page link](https://img.shields.io/badge/Page-StyleBooth-Gree)](https://ali-vilab.github.io/stylebooth-page/) |
#### TXT Format
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
```
### Training
We provide a framework for training and inference, so the script below is just for illustration purposes. To achieve better results, you can modify the corresponding parameters as needed.
#### Text-to-Image Generation
- SCEdit
```python
python scepter/tools/run_train.py --cfg scepter/methods/scedit/t2i/sd15_512_sce_t2i.yaml # SD v1.5
python scepter/tools/run_train.py --cfg scepter/methods/scedit/t2i/sd21_768_sce_t2i.yaml # SD v2.1
python scepter/tools/run_train.py --cfg scepter/methods/scedit/t2i/sdxl_1024_sce_t2i.yaml # SD XL
```
- Existing Tuning Strategies
```python
python scepter/tools/run_train.py --cfg scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml # fully-tuning on SD v1.5
python scepter/tools/run_train.py --cfg scepter/methods/examples/generation/stable_diffusion_2.1_768_lora.yaml # lora-tuning on SD v2.1
```
- Data Text Format
```python
# Download the 3D_example_txt.zip as previously mentioned
python scepter/tools/run_train.py --cfg scepter/methods/scedit/t2i/sdxl_1024_sce_t2i_datatxt.yaml
```
#### Controllable Image Synthesis
- SCEdit
The YAML configuration can be modified to combine different base models and conditions. The following is provided as an example.
```python
python scepter/tools/run_train.py --cfg scepter/methods/scedit/ctr/sd15_512_sce_ctr_hed.yaml # SD v1.5 + hed
python scepter/tools/run_train.py --cfg scepter/methods/scedit/ctr/sd21_768_sce_ctr_canny.yaml # SD v2.1 + canny
python scepter/tools/run_train.py --cfg scepter/methods/scedit/ctr/sd21_768_sce_ctr_pose.yaml # SD v2.1 + pose
python scepter/tools/run_train.py --cfg scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_depth.yaml # SD XL + depth
python scepter/tools/run_train.py --cfg scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_color.yaml # SD XL + color
```
- Data Text Format
```python
# Download the 3D_example_txt.zip as previously mentioned
python scepter/tools/run_train.py --cfg scepter/methods/scedit/ctr/sdxl_1024_sce_ctr_color_datatxt.yaml
```
### Inference
#### Base Model Inference
```python
python scepter/tools/run_inference.py --cfg scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml --prompt 'a cute dog' --save_folder 'inference' # generation on SD v1.5
python scepter/tools/run_inference.py --cfg scepter/methods/examples/generation/stable_diffusion_2.1_768.yaml --prompt 'a cute dog' --save_folder 'inference' # generation on SD v2.1
python scepter/tools/run_inference.py --cfg scepter/methods/examples/generation/stable_diffusion_xl_1024.yaml --prompt 'a cute dog' --save_folder 'inference' # generation on SD XL
```
#### Fine-tuned Model Inference
```python
python scepter/tools/run_inference.py --cfg scepter/methods/scedit/t2i/sd15_512_sce_t2i_swift.yaml --pretrained_model 'cache/save_data/sd15_512_sce_t2i_swift/checkpoints/ldm_step-100.pth' --prompt 'A close up of a small rabbit wearing a hat and scarf' --save_folder 'trained_test_prompt_rabbit'
```
#### Controllable Image Synthesis Inference
- SCEdit
```python
python scepter/tools/run_inference.py --cfg scepter/methods/scedit/ctr/sd21_768_sce_ctr_canny.yaml --num_samples 1 --prompt 'a single flower is shown in front of a tree' --save_folder 'test_flower_canny' --image_size 768 --task control --image 'asset/images/flower.jpg' --control_mode canny --pretrained_model ms://damo/scepter_scedit@controllable_model/SD2.1/canny_control/0_SwiftSCETuning/pytorch_model.bin # canny
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
```
## 🖥️ SCEPTER Studio
@@ -181,85 +128,26 @@ 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.
### Modelscope Studio
To support the sharing and downloading of models,
please make sure that you have installed **zip** and Git Large File Storage (**git lfs**).
### Usage Demo
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)
| [Image Editing](https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=assets%2Fscepter_studio%2Fimage_editing_20240419.webm) | [Training](https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=assets%2Fscepter_studio%2Ftraining_20240419.webm) | [Model Sharing](https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=assets%2Fscepter_studio%2Fmodel_sharing_20240419.webm) | [Model Inference](https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=assets%2Fscepter_studio%2Fmodel_inference_20240419.webm) | [Data Management](https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=assets%2Fscepter_studio%2Fdata_management_20240419.webm) |
|:----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------:|:-----------------------------------------------------------------------------------------------------------------------------------------------------------------------------:|:-----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------:|:-------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------:|:--------------------------------------------:|
| <video src="https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=assets%2Fscepter_studio%2Fimage_editing_20240419.webm" width="240" controls></video> | <video src="https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=assets%2Fscepter_studio%2Ftraining_20240419.webm" width="240" controls></video> | <video src="https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=assets%2Fscepter_studio%2Fmodel_sharing_20240419.webm" width="240" controls></video> | <video src="https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=assets%2Fscepter_studio%2Fmodel_inference_20240419.webm" width="240" controls></video> | <video src="https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=assets%2Fscepter_studio%2Fdata_management_20240419.webm" width="240" controls></video> |
## 🖼️ Gallery
### Modelscope Studio & Huggingface Space
### 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
| **Model** | **SCEdit** | **Full** | **LoRA** |
|:---------:|:----------:|:--------:|:--------:|
| SD 1.5 | 🪄 | ✅ | ✅ |
| SD 2.1 | 🪄 | ✅ | ✅ |
| SD XL | 🪄 | ✅ | ✅ |
### Controllable Image Synthesis
- SCEdit
| **Model** | **Canny** | **HED** | **Depth** | **Pose** | **Color** |
|:---------:|:---------:|:-------:|:---------:|:--------:|:---------:|
| SD 1.5 | ✅ | ✅ | ✅ | ✅ | ✅ |
| SD 2.1 | 🪄 | 🪄 | 🪄 | 🪄 | 🪄 |
| SD XL | 🪄 | 🪄 | 🪄 | 🪄 | 🪄 |
### Model URL
- ✅ indicates support for both training and inference.
- 🪄 denotes that the model has been published.
- More models will be released in the future.
| Model | URL |
|--------|------------------------------------------------------------------------------------------------------------------------------------------------|
| SCEdit | [ModelScope](https://modelscope.cn/models/damo/scepter_scedit/summary) [HuggingFace](https://huggingface.co/scepter-studio/scepter_scedit) |
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.
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) and [hf_scepter_studio](https://huggingface.co/spaces/modelscope/scepter_studio)
## 🔍 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.
@@ -272,6 +160,21 @@ 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.
+4 -1
View File
@@ -1,11 +1,14 @@
albumentations
bezier
einops
modelscope
ms-swift>=1.5.2
ms-swift>=2.0.1
numpy
open_clip_torch
opencv-python
opencv_transforms>=0.0.6
oss2>=2.15.0
pycocotools
pyyaml>=5.3.1
scikit-image
torchsde
+1
View File
@@ -1,3 +1,4 @@
git+https://github.com/cocodataset/panopticapi.git
torch==2.0.1
torchvision==0.15.2
xformers==0.0.21
+4
View File
@@ -1,2 +1,6 @@
bitsandbytes
gradio>=3.47.1,<4.0.0
imagehash
psutil
tiktoken
transformers_stream_generator
@@ -0,0 +1,234 @@
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/sd15_512_textlora
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)$"
-
NAME: SwiftLoRA
R: 64
LORA_ALPHA: 64
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "cond_stage_model.*(q_proj|k_proj|v_proj|out_proj|mlp.fc1|mlp.fc2)$"
#
MODEL:
NAME: LatentDiffusion
PARAMETERIZATION: eps
TIMESTEPS: 1000
MIN_SNR_GAMMA:
ZERO_TERMINAL_SNR: False
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-v1-5@v1-5-pruned-emaonly.safetensors
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
# DEFAULT_N_PROMPT: 'lowres, error, worst quality, low quality, jpeg artifacts, ugly, duplicate, morbid, mutilated, out of frame, extra fingers, mutated hands, poorly drawn hands, poorly drawn face, mutation, deformed, blurry, dehydrated, bad anatomy, bad proportions, extra limbs, cloned face, disfigured, gross proportions, malformed limbs, missing arms, missing legs, extra arms, extra legs, fused fingers, too many fingers, long neck, username, watermark, signature'
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: 8
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: 768
DISABLE_MIDDLE_SELF_ATTN: False
USE_LINEAR_IN_TRANSFORMER: False
PRETRAINED_MODEL:
IGNORE_KEYS: []
#
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: ClipTokenizer
PRETRAINED_PATH: ms://AI-ModelScope/clip-vit-large-patch14
LENGTH: 77
CLEAN: True
#
COND_STAGE_MODEL:
NAME: FrozenCLIPEmbedder
FREEZE: True
LAYER: last
PRETRAINED_MODEL: ms://AI-ModelScope/clip-vit-large-patch14
USE_GRAD: True
#
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 2023
GUIDE_SCALE: 7.5
GUIDE_RESCALE: 0.5
DISCRETIZATION: trailing
IMAGE_SIZE: [512, 512]
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: 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,345 @@
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/sdxl_1024_textlora
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)$"
-
NAME: SwiftLoRA
R: 64
LORA_ALPHA: 64
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "cond_stage_model.embedders.0.*(q_proj|k_proj|v_proj|out_proj|mlp.fc1|mlp.fc2)$"
#
MODEL:
NAME: LatentDiffusionXL
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: 'lowres, error, worst quality, low quality, jpeg artifacts, ugly, duplicate, morbid, mutilated, out of frame, extra fingers, mutated hands, poorly drawn hands, poorly drawn face, mutation, deformed, blurry, dehydrated, bad anatomy, bad proportions, extra limbs, cloned face, disfigured, gross proportions, malformed limbs, missing arms, missing legs, extra arms, extra legs, fused fingers, too many fingers, long neck, username, watermark, signature'
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:
USE_GRAD: True
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
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 2023
GUIDE_SCALE: 5.0
GUIDE_RESCALE:
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
INTERPOLATION: bicubic
SIZE: 1024
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: FlexibleCropXL
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: [ 'img' ]
BACKEND: torchvision
- NAME: Select
KEYS: [ 'img', 'prompt', 'img_original_size_as_tuple', 'img_target_size_as_tuple', 'img_crop_coords_top_left' ]
META_KEYS: [ 'data_key', 'img_path' ]
- NAME: Rename
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: 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: [ 1024, 1024 ]
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,234 @@
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/sd15_512_textsce_t2i_swift
LOG_FILE: std_log.txt
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/data"
#
TUNER:
-
NAME: SwiftSCETuning
DIMS: [1280, 1280, 1280, 1280, 1280, 640, 640, 640, 320, 320, 320, 320]
TARGET_MODULES: model.lsc_identity\.\d+$
DOWN_RATIO: 1.0
TUNER_MODE: identity
-
NAME: SwiftLoRA
R: 64
LORA_ALPHA: 64
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "cond_stage_model.*(q_proj|k_proj|v_proj|out_proj|mlp.fc1|mlp.fc2)$"
#
MODEL:
NAME: LatentDiffusion
PARAMETERIZATION: eps
TIMESTEPS: 1000
MIN_SNR_GAMMA:
ZERO_TERMINAL_SNR: False
PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-v1-5@v1-5-pruned-emaonly.safetensors
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
# DEFAULT_N_PROMPT: 'lowres, error, worst quality, low quality, jpeg artifacts, ugly, duplicate, morbid, mutilated, out of frame, extra fingers, mutated hands, poorly drawn hands, poorly drawn face, mutation, deformed, blurry, dehydrated, bad anatomy, bad proportions, extra limbs, cloned face, disfigured, gross proportions, malformed limbs, missing arms, missing legs, extra arms, extra legs, fused fingers, too many fingers, long neck, username, watermark, signature'
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: 8
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: 768
DISABLE_MIDDLE_SELF_ATTN: False
USE_LINEAR_IN_TRANSFORMER: False
PRETRAINED_MODEL:
IGNORE_KEYS: []
#
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: ClipTokenizer
PRETRAINED_PATH: ms://AI-ModelScope/clip-vit-large-patch14
LENGTH: 77
CLEAN: True
#
COND_STAGE_MODEL:
NAME: FrozenCLIPEmbedder
FREEZE: True
LAYER: last
PRETRAINED_MODEL: ms://AI-ModelScope/clip-vit-large-patch14
USE_GRAD: True
#
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 2023
GUIDE_SCALE: 7.5
GUIDE_RESCALE: 0.5
DISCRETIZATION: trailing
IMAGE_SIZE: [512, 512]
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: 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,348 @@
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/sdxl_1024_textsce_t2i_swift
LOG_FILE: std_log.txt
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/data"
#
TUNER:
-
NAME: SwiftSCETuning
DIMS: [1280, 1280, 640, 640, 640, 320, 320, 320, 320]
TARGET_MODULES: model.lsc_identity\.\d+$
DOWN_RATIO: 1.0
TUNER_MODE: identity
-
NAME: SwiftLoRA
R: 64
LORA_ALPHA: 64
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "cond_stage_model.embedders.0.*(q_proj|k_proj|v_proj|out_proj|mlp.fc1|mlp.fc2)$"
#
MODEL:
NAME: LatentDiffusionXL
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: 'lowres, error, worst quality, low quality, jpeg artifacts, ugly, duplicate, morbid, mutilated, out of frame, extra fingers, mutated hands, poorly drawn hands, poorly drawn face, mutation, deformed, blurry, dehydrated, bad anatomy, bad proportions, extra limbs, cloned face, disfigured, gross proportions, malformed limbs, missing arms, missing legs, extra arms, extra legs, fused fingers, too many fingers, long neck, username, watermark, signature'
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:
USE_GRAD: True
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
#
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
INTERPOLATION: bicubic
SIZE: 1024
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: FlexibleCropXL
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: [ 'img' ]
BACKEND: torchvision
- NAME: Select
KEYS: [ 'img', 'prompt', 'img_original_size_as_tuple', 'img_target_size_as_tuple', 'img_crop_coords_top_left' ]
META_KEYS: [ 'data_key', 'img_path' ]
- NAME: Rename
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: 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: [ 1024, 1024 ]
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
@@ -1,7 +1,7 @@
TUNERS:
- NAME: Azure-Dragon
NAME_ZH: 青龙
SOURCE: wanx
SOURCE: scepter
DESCRIPTION: None
BASE_MODEL: SD_XL1.0
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/azure_dragon/
@@ -10,7 +10,7 @@ TUNERS:
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
SOURCE: scepter
DESCRIPTION: None
BASE_MODEL: SD_XL1.0
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/gold_dragon/
@@ -19,7 +19,7 @@ TUNERS:
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
SOURCE: scepter
DESCRIPTION: None
BASE_MODEL: SD_XL1.0
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/spring_festival_dragon/
@@ -28,7 +28,7 @@ TUNERS:
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
SOURCE: scepter
DESCRIPTION: None
BASE_MODEL: SD_XL1.0
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/red_dragon/
@@ -37,7 +37,7 @@ TUNERS:
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
SOURCE: scepter
DESCRIPTION: None
BASE_MODEL: SD_XL1.0
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/chinese_punk_dragon/
@@ -46,7 +46,7 @@ TUNERS:
PROMPT_EXAMPLE: uhd Image,Dragon,Chinese Dragon, Dunhuang Mural Style, Traditional Maritime Art Style
- NAME: Cute-Dragon
NAME_ZH: 喜庆龙
SOURCE: wanx
SOURCE: scepter
DESCRIPTION: None
BASE_MODEL: SD_XL1.0
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/cute_dragon/
@@ -55,7 +55,7 @@ TUNERS:
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
SOURCE: scepter
DESCRIPTION: None
BASE_MODEL: SD_XL1.0
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/baby_dragon/
@@ -64,7 +64,7 @@ TUNERS:
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
SOURCE: scepter
DESCRIPTION: None
BASE_MODEL: SD_XL1.0
MODEL_PATH: ms://damo/scepter_scedit@tuners_model/SD_XL1.0/sloppy_dragon/
@@ -0,0 +1,129 @@
NAME: EDIT
IS_DEFAULT: False
DEFAULT_PARAS:
PARAS:
RESOLUTIONS: [[1024, 1024]]
INPUT:
IMAGE:
PROMPT: ""
NEGATIVE_PROMPT: ""
TARGET_SIZE_AS_TUPLE: [1024, 1024]
PROMPT_PREFIX: ""
SAMPLE: ddim
SAMPLE_STEPS: 50
GUIDE_SCALE:
text: 7.5
image: 1.5
GUIDE_RESCALE: 0.5
DISCRETIZATION: trailing
OUTPUT:
LATENT:
IMAGES:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
-
NAME: encode
DTYPE: float16
INPUT: ["IMAGE"]
-
NAME: decode
DTYPE: float16
INPUT: ["LATENT"]
PARAS:
# SCALE_FACTOR DESCRIPTION: The vae embeding scale. TYPE: float default: 0.18215
SCALE_FACTOR: 0.18215
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_text
DTYPE: float16
INPUT: ["PROMPT", "NEGATIVE_PROMPT"]
MODEL:
PRETRAINED_MODEL: ms://damo/stylebooth@models/stylebooth-tb-5000-0.bin
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:
NAME: DiffusionUNet
PRETRAINED_PATH:
IN_CHANNELS: 8
OUT_CHANNELS: 4
MODEL_CHANNELS: 320
NUM_HEADS: 8
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: 768
DISABLE_MIDDLE_SELF_ATTN: False
USE_LINEAR_IN_TRANSFORMER: False
IGNORE_KEYS: []
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
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: ClipTokenizer
PRETRAINED_PATH: ms://AI-ModelScope/clip-vit-large-patch14
LENGTH: 77
CLEAN: True
#
COND_STAGE_MODEL:
NAME: FrozenCLIPEmbedder
FREEZE: True
USE_GRAD: False
LAYER: last
PRETRAINED_MODEL: ms://AI-ModelScope/clip-vit-large-patch14
@@ -79,6 +79,7 @@ EXTENSION_PARAS:
MANTRA_BOOK: scepter/methods/studio/extensions/mantra_book/mantra_book.yaml
OFFICIAL_TUNERS: scepter/methods/studio/extensions/tuners/official_tuners.yaml
OFFICIAL_CONTROLLERS: scepter/methods/studio/extensions/controllers/official_controllers.yaml
TUNER_MANAGER: scepter/methods/studio/tuner_manager/tuner_manager.yaml
CONTROLABLE_ANNOTATORS:
-
NAME: "CannyAnnotator"
@@ -0,0 +1,268 @@
NAME: LARGEN
IS_DEFAULT: 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" ]
@@ -5,3 +5,169 @@ FILE_SYSTEM:
# NAME DESCRIPTION: TYPE: default: ''
NAME: LocalFs
AUTO_CLEAN: False
PROCESSORS:
- NAME: BlipImageBase
TYPE: caption
MODEL_PATH: ms://cubeai/blip-image-captioning-base
DEVICE: "gpu"
MEMORY: 1200
PARAS:
- LANGUAGE_NAME: English
LANGUAGE_ZH_NAME: 英语
- NAME: QWVL
TYPE: caption
MODEL_PATH: ms://qwen/Qwen-VL:v1.0.3
DEVICE: "gpu"
MEMORY: 19968
PARAS:
- PROMPT: 用中文描述这张图片
LANGUAGE_NAME: Chinese
LANGUAGE_ZH_NAME: 中文
MAX_NEW_TOKENS:
VALUE: 1024
MAX: 2048
STEP: 128
MIN: 256
MIN_NEW_TOKENS:
VALUE: 16
MAX: 1024
STEP: 16
MIN: 0
NUM_BEAMS:
VALUE: 1
MAX: 12
STEP: 1
MIN: 1
REPETITION_PENALTY:
VALUE: 1.0
MAX: 100.0
STEP: 1.0
MIN: 1.0
TEMPERATURE:
VALUE: 1.0
MAX: 100.0
STEP: 1.0
MIN: 1.0
- PROMPT: Generate the caption in English
LANGUAGE_NAME: English
LANGUAGE_ZH_NAME: 英语
MAX_NEW_TOKENS:
VALUE: 1024
MAX: 2048
STEP: 128
MIN: 256
MIN_NEW_TOKENS:
VALUE: 16
MAX: 1024
STEP: 16
MIN: 0
NUM_BEAMS:
VALUE: 1
MAX: 12
STEP: 1
MIN: 1
REPETITION_PENALTY:
VALUE: 1.0
MAX: 100.0
STEP: 1.0
MIN: 1.0
TEMPERATURE:
VALUE: 1.0
MAX: 100.0
STEP: 1.0
MIN: 1.0
-
NAME: QWVLQuantize
TYPE: caption
DEVICE: "gpu"
MEMORY: 7885
MODEL_PATH: ms://qwen/Qwen-VL:v1.0.3
PARAS:
- PROMPT: 用中文描述这张图片
LANGUAGE_NAME: Chinese
LANGUAGE_ZH_NAME: 中文
MAX_NEW_TOKENS:
VALUE: 1024
MAX: 2048
STEP: 128
MIN: 256
MIN_NEW_TOKENS:
VALUE: 16
MAX: 1024
STEP: 16
MIN: 0
NUM_BEAMS:
VALUE: 1
MAX: 12
STEP: 1
MIN: 1
REPETITION_PENALTY:
VALUE: 1.0
MAX: 100.0
STEP: 1.0
MIN: 1.0
TEMPERATURE:
VALUE: 1.0
MAX: 100.0
STEP: 1.0
MIN: 1.0
- PROMPT: Generate the caption in English
LANGUAGE_NAME: English
LANGUAGE_ZH_NAME: 英语
MAX_NEW_TOKENS:
VALUE: 1024
MAX: 2048
STEP: 128
MIN: 256
MIN_NEW_TOKENS:
VALUE: 16
MAX: 1024
STEP: 16
MIN: 0
NUM_BEAMS:
VALUE: 1
MAX: 12
STEP: 1
MIN: 1
REPETITION_PENALTY:
VALUE: 1.0
MAX: 100.0
STEP: 1.0
MIN: 1.0
TEMPERATURE:
VALUE: 1.0
MAX: 100.0
STEP: 1.0
MIN: 1.0
-
NAME: CenterCrop
TYPE: image
DEVICE: "cpu"
MEMORY: 10
PARAS:
HEIGHT_RATIO:
VALUE: 1
MAX: 20
STEP: 1
MIN: 1
WIDTH_RATIO:
VALUE: 1
MAX: 20
STEP: 1
MIN: 1
# - NAME: PaddingCrop
# TYPE: image
# DEVICE: "cpu"
# MEMORY: 10
# PARAS:
# HEIGHT_RATIO:
# VALUE: 3
# MAX: 25
# STEP: 1
# MIN: 1
# WIDTH_RATIO:
# VALUE: 4
# MAX: 20
# STEP: 1
# MIN: 1
+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
@@ -5,19 +5,20 @@ META:
VERSION: 'SD_XL1.0'
DESCRIPTION: "Stable Diffusion XL1.0"
IS_DEFAULT: True
IS_SHARE: True
INFERENCE_PARAS:
INFERENCE_BATCH_SIZE: 1
INFERENCE_PREFIX: ""
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 +30,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 +42,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 +54,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 +66,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 +147,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"
@@ -530,7 +533,6 @@ SOLVER:
GUIDE_SCALE: 5.0
GUIDE_RESCALE:
DISCRETIZATION: linspace
IMAGE_SIZE: [ 1024, 1024]
RUN_TRAIN_N: False
# OPTIMIZER DESCRIPTION: TYPE: default: ''
OPTIMIZER:
@@ -586,15 +588,45 @@ 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'
DISABLE_SNAPSHOT: True
#
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
@@ -4,19 +4,20 @@ META:
VERSION: 'SD1.5'
DESCRIPTION: "Stable Diffusion v1.5"
IS_DEFAULT: False
IS_SHARE: True
INFERENCE_PARAS:
INFERENCE_BATCH_SIZE: 1
INFERENCE_PREFIX: ""
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 +29,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 +41,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 +54,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 +67,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 +136,7 @@ SOLVER:
MAX_EPOCHS: -1
NUM_FOLDS: 1
ACCU_STEP: 1
EVAL_INTERVAL: -1
#
WORK_DIR:
LOG_FILE: std_log.txt
@@ -243,7 +245,6 @@ SOLVER:
GUIDE_SCALE: 7.5
GUIDE_RESCALE:
DISCRETIZATION: trailing
IMAGE_SIZE: [512, 512]
RUN_TRAIN_N: False
#
OPTIMIZER:
@@ -273,14 +274,14 @@ SOLVER:
- NAME: LoadImageFromFile
RGB_ORDER: RGB
BACKEND: pillow
- NAME: Resize
SIZE: 512
- NAME: FlexibleResize
INTERPOLATION: bilinear
SIZE: [ 512, 512 ]
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: CenterCrop
SIZE: 512
- NAME: FlexibleCenterCrop
SIZE: [ 512, 512 ]
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
@@ -298,15 +299,45 @@ 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'
DISABLE_SNAPSHOT: True
#
EVAL_HOOKS:
-
NAME: ProbeDataHook
PROB_INTERVAL: 100
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
SAVE_PROBE_PREFIX: 'image'
@@ -4,19 +4,20 @@ META:
VERSION: 'SD2.1'
DESCRIPTION: "Stable Diffusion v2.1"
IS_DEFAULT: False
IS_SHARE: True
INFERENCE_PARAS:
INFERENCE_BATCH_SIZE: 1
INFERENCE_PREFIX: ""
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 +29,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 +41,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 +79,7 @@ SOLVER:
MAX_EPOCHS: -1
NUM_FOLDS: 1
ACCU_STEP: 1
EVAL_INTERVAL: -1
#
WORK_DIR:
LOG_FILE: std_log.txt
@@ -185,7 +187,6 @@ SOLVER:
GUIDE_SCALE: 7.5
GUIDE_RESCALE:
DISCRETIZATION: trailing
IMAGE_SIZE: [768, 768]
RUN_TRAIN_N: False
#
OPTIMIZER:
@@ -215,14 +216,14 @@ SOLVER:
- NAME: LoadImageFromFile
RGB_ORDER: RGB
BACKEND: pillow
- NAME: Resize
SIZE: 768
- NAME: FlexibleResize
INTERPOLATION: bilinear
SIZE: [ 768, 768 ]
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
- NAME: CenterCrop
SIZE: 768
- NAME: FlexibleCenterCrop
SIZE: [ 768, 768 ]
INPUT_KEY: [ 'img' ]
OUTPUT_KEY: [ 'img' ]
BACKEND: pillow
@@ -240,15 +241,45 @@ 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'
DISABLE_SNAPSHOT: True
#
EVAL_HOOKS:
-
NAME: ProbeDataHook
PROB_INTERVAL: 100
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
SAVE_PROBE_PREFIX: 'image'
@@ -0,0 +1,168 @@
---
frameworks:
- Pytorch
license: Apache License 2.0
tasks:
- efficient-diffusion-tuning
---
<p align="center">
<h2 align="center">{MODEL_NAME}</h2>
<p align="center">
<br>
<a href="https://github.com/modelscope/scepter/"><img src="https://img.shields.io/badge/powered by-scepter-6FEBB9.svg"></a>
<br>
</p>
## Model Introduction
{MODEL_DESCRIPTION}
## Model Parameters
<table>
<thead>
<tr>
<th rowspan="2">Base Model</th>
<th rowspan="2">Tuner Type</th>
<th colspan="4">Training Parameters</th>
</tr>
<tr>
<th>Batch Size</th>
<th>Epochs</th>
<th>Learning Rate</th>
<th>Resolution</th>
</tr>
</thead>
<tbody align="center">
<tr>
<td rowspan="8">{BASE_MODEL}</td>
<td>{TUNER_TYPE}</td>
<td>{TRAIN_BATCH_SIZE}</td>
<td>{TRAIN_EPOCH}</td>
<td>{LEARNING_RATE}</td>
<td>[{HEIGHT}, {WIDTH}]</td>
</tr>
</tbody>
</table>
<table>
<thead>
<tr>
<th>Data Type</th>
<th>Data Space</th>
<th>Data Name</th>
<th>Data Subset</th>
</tr>
</thead>
<tbody align="center">
<tr>
<td>{DATA_TYPE}</td>
<td>{MS_DATA_SPACE}</td>
<td>{MS_DATA_NAME}</td>
<td>{MS_DATA_SUBNAME}</td>
</tr>
</tbody>
</table>
## Model Performance
Given the input "{EVAL_PROMPT}," the following image may be generated:
![image]({IMAGE_PATH})
## Model Usage
### Command Line Execution
* Run using Scepter's SDK, taking care to use different configuration files in accordance with the different base models, as per the corresponding relationships shown below
<table>
<thead>
<tr>
<th rowspan="2">Base Model</th>
<th rowspan="1">LORA</th>
<th colspan="1">SCE</th>
<th colspan="1">TEXT_LORA</th>
<th colspan="1">TEXT_SCE</th>
</tr>
</thead>
<tbody align="center">
<tr>
<td rowspan="8">SD1.5</td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/examples/generation/stable_diffusion_1.5_512_lora.yaml">lora_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/scedit/t2i/sd15_512_sce_t2i_swift.yaml">sce_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/examples/generation/stable_diffusion_1.5_512_text_lora.yaml">text_lora_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/scedit/t2i/stable_diffusion_1.5_512_text_sce.yaml">text_sce_cfg</a></td>
</tr>
</tbody>
<tbody align="center">
<tr>
<td rowspan="8">SD2.1</td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/examples/generation/stable_diffusion_2.1_768_lora.yaml">lora_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/scedit/t2i/sd21_768_sce_t2i_swift.yaml">sce_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/examples/generation/stable_diffusion_2.1_768_text_lora.yaml">text_lora_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/scedit/t2i/sd21_768_text_sce_t2i_swift.yaml">text_sce_cfg</a></td>
</tr>
</tbody>
<tbody align="center">
<tr>
<td rowspan="8">SDXL</td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/examples/generation/stable_diffusion_xl_1024_lora.yaml">lora_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/scedit/t2i/sdxl_1024_sce_t2i_swift.yaml">sce_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/examples/generation/stable_diffusion_xl_1024_text_lora.yaml">text_lora_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/scedit/t2i/sdxl_1024_text_sce_t2i_swift.yaml">text_sce_cfg</a></td>
</tr>
</tbody>
</table>
* Running from Source Code
```shell
git clone https://github.com/modelscope/scepter.git
cd scepter
pip install -r requirements/recommended.txt
PYTHONPATH=. python scepter/tools/run_inference.py
--pretrained_model {this model folder}
--cfg {lora_cfg} or {sce_cfg} or {text_lora_cfg} or {text_sce_cfg}
--prompt '{EVAL_PROMPT}'
--save_folder 'inference'
```
* Running after Installing Scepter (Recommended)
```shell
pip install scepter
python -m scepter/tools/run_inference.py
--pretrained_model {this model folder}
--cfg {lora_cfg} or {sce_cfg} or {text_lora_cfg} or {text_sce_cfg}
--prompt '{EVAL_PROMPT}'
--save_folder 'inference'
```
### Running with Scepter Studio
```shell
pip install scepter
# Launch Scepter Studio
python -m scepter.tools.webui
```
* Refer to the following guides for model usage.
(video url)
## Model Reference
If you wish to use this model for your own purposes, please cite it as follows.
```bibtex
@misc{{MODEL_NAME},
title = {{MODEL_NAME}, {MODEL_URL}},
author = {{USER_NAME}},
year = {2024}
}
```
This model was trained using [Scepter Studio](https://github.com/modelscope/scepter); [Scepter](https://github.com/modelscope/scepter)
is an algorithm framework and toolbox developed by the Alibaba Tongyi Wanxiang Team. It provides a suite of tools and models for image generation, editing, fine-tuning, data processing, and more. If you find our work beneficial for your research,
please cite as follows.
```bibtex
@misc{scepter,
title = {SCEPTER, https://github.com/modelscope/scepter},
author = {SCEPTER},
year = {2023}
}
```
@@ -0,0 +1,168 @@
---
frameworks:
- Pytorch
license: Apache License 2.0
tasks:
- efficient-diffusion-tuning
---
<p align="center">
<h2 align="center">{MODEL_NAME}</h2>
<p align="center">
<br>
<a href="https://github.com/modelscope/scepter/"><img src="https://img.shields.io/badge/powered by-scepter-6FEBB9.svg"></a>
<br>
</p>
## 模型介绍
{MODEL_DESCRIPTION}
## 模型参数
<table>
<thead>
<tr>
<th rowspan="2">基础模型</th>
<th rowspan="2">微调类型</th>
<th colspan="4">训练参数</th>
</tr>
<tr>
<th>批次大小</th>
<th>轮数</th>
<th>学习率</th>
<th>分辨率</th>
</tr>
</thead>
<tbody align="center">
<tr>
<td rowspan="8">{BASE_MODEL}</td>
<td>{TUNER_TYPE}</td>
<td>{TRAIN_BATCH_SIZE}</td>
<td>{TRAIN_EPOCH}</td>
<td>{LEARNING_RATE}</td>
<td>[{HEIGHT}, {WIDTH}]</td>
</tr>
</tbody>
</table>
<table>
<thead>
<tr>
<th>数据类型</th>
<th>数据空间</th>
<th>数据名称</th>
<th>数据子集</th>
</tr>
</thead>
<tbody align="center">
<tr>
<td> {DATA_TYPE}</td>
<td>{MS_DATA_SPACE}</td>
<td>{MS_DATA_NAME}</td>
<td>{MS_DATA_SUBNAME}</td>
</tr>
</tbody>
</table>
## 模型效果
输入 "{EVAL_PROMPT}",可能会得到如下图像:
![image]({IMAGE_PATH})
## 模型使用
### 命令行运行
* 使用scepter的sdk进行运行,注意需要按照模型参数中基模型的不同使用不同的配置文件,其对应关系如下
<table>
<thead>
<tr>
<th rowspan="2">Base Model</th>
<th rowspan="1">LORA</th>
<th colspan="1">SCE</th>
<th colspan="1">TEXT_LORA</th>
<th colspan="1">TEXT_SCE</th>
</tr>
</thead>
<tbody align="center">
<tr>
<td rowspan="8">SD1.5</td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/examples/generation/stable_diffusion_1.5_512_lora.yaml">lora_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/scedit/t2i/sd15_512_sce_t2i_swift.yaml">sce_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/examples/generation/stable_diffusion_1.5_512_text_lora.yaml">text_lora_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/scedit/t2i/stable_diffusion_1.5_512_text_sce.yaml">text_sce_cfg</a></td>
</tr>
</tbody>
<tbody align="center">
<tr>
<td rowspan="8">SD2.1</td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/examples/generation/stable_diffusion_2.1_768_lora.yaml">lora_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/scedit/t2i/sd21_768_sce_t2i_swift.yaml">sce_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/examples/generation/stable_diffusion_2.1_768_text_lora.yaml">text_lora_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/scedit/t2i/sd21_768_text_sce_t2i_swift.yaml">text_sce_cfg</a></td>
</tr>
</tbody>
<tbody align="center">
<tr>
<td rowspan="8">SDXL</td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/examples/generation/stable_diffusion_xl_1024_lora.yaml">lora_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/scedit/t2i/sdxl_1024_sce_t2i_swift.yaml">sce_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/examples/generation/stable_diffusion_xl_1024_text_lora.yaml">text_lora_cfg</a></td>
<td><a href="https://github.com/modelscope/scepter/blob/main/scepter/methods/scedit/t2i/sdxl_1024_text_sce_t2i_swift.yaml">text_sce_cfg</a></td>
</tr>
</tbody>
</table>
* 从源码运行
```shell
git clone https://github.com/modelscope/scepter.git
cd scepter
pip install -r requirements/recommended.txt
PYTHONPATH=. python scepter/tools/run_inference.py
--pretrained_model {this model folder}
--cfg {lora_cfg} or {sce_cfg} or {text_lora_cfg} or {text_sce_cfg}
--prompt '{EVAL_PROMPT}'
--save_folder 'inference'
```
* 安装scepter后运行(推荐)
```shell
pip install scepter
python -m scepter/tools/run_inference.py
--pretrained_model {this model folder}
--cfg {lora_cfg} or {sce_cfg} or {text_lora_cfg} or {text_sce_cfg}
--prompt '{EVAL_PROMPT}'
--save_folder 'inference'
```
### 使用Scepter Studio运行
```shell
pip install scepter
启动scepter studio
python -m scepter.tools.webui
```
* 参考以下指南使用模型
## 模型引用
如果你想使用该模型应用于自己的场景,请按照如下方式引用该模型。
```bibtex
@misc{{MODEL_NAME},
title = {{MODEL_NAME}, {MODEL_URL}},
author = {{USER_NAME}},
year = {2024}
}
```
该模型是基于[Scepter Studio](https://github.com/modelscope/scepter)训练得到;[scepter](https://github.com/modelscope/scepter)
是由阿里巴巴通义万相团队开发的算法框架和工具箱,提供图像生成、编辑、微调、数据处理等一系列工具和模型。如果您觉得我们的工作有益于您的工作,
请按照如下方式引用。
```bibtex
@misc{scepter,
title = {SCEPTER, https://github.com/modelscope/scepter},
author = {SCEPTER},
year = {2023}
}
```
@@ -0,0 +1,14 @@
WORK_DIR: "tuner_manager"
SELF_TRAIN_DIR: "self_train"
EXPORT_DIR: "export_model"
TUNER_LIST_YAML: "tuner_list.yaml"
README_EN: "scepter/methods/studio/tuner_manager/readme_en.md"
README_ZH: "scepter/methods/studio/tuner_manager/readme_zh.md"
BASE_MODEL_VERSION:
- BASE_MODEL: 'SD_XL1.0'
TUNER_TYPE: [ 'TEXT_SCE', 'SCE', 'LORA', 'TEXT_LORA', 'FULL' ]
- BASE_MODEL: 'SD1.5'
TUNER_TYPE: [ 'TEXT_SCE', 'SCE', 'LORA', 'TEXT_LORA', 'FULL' ]
- BASE_MODEL: 'SD2.1'
TUNER_TYPE: [ 'SCE', 'LORA', 'FULL' ]
+17 -7
View File
@@ -101,6 +101,10 @@ class ImageTextPairDataset(BaseDataset):
'NEGTIVE_PROMPT': {
'value': '',
'description': 'The default negtive prompt',
},
'DATA_NUM': {
'value': '',
'description': '',
}
}
para_dict.update(BaseDataset.para_dict)
@@ -108,6 +112,7 @@ class ImageTextPairDataset(BaseDataset):
def __init__(self, cfg, logger=None):
super(ImageTextPairDataset, self).__init__(cfg, logger=logger)
self.p_zero = cfg.get('P_ZERO', 0.0)
self.real_number = cfg.get('DATA_NUM', None)
self._default_item = {
'meta': {},
'prompt':
@@ -244,12 +249,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 +274,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]
+6 -2
View File
@@ -288,8 +288,12 @@ class DataObject(object):
self.shuffle = False
self.data_sampler_config.SEED = seed
self.data_sampler_config.BATCH_SIZE = self.batch_size
self.sampler = SAMPLERS.build(self.data_sampler_config,
logger=self.logger)
sampler = SAMPLERS.build(self.data_sampler_config,
logger=self.logger)
if sampler_name.endswith('BatchSampler'):
self.batch_sampler = sampler
else:
self.sampler = sampler
def _instantiate_multi_level_batch_sampler(self, sampler_config,
batch_size, rank, seed):
+1 -1
View File
@@ -6,4 +6,4 @@ from scepter.modules.data.sampler.registry import SAMPLERS
from scepter.modules.data.sampler.sampler import (
EvalDistributedSampler, LoopSampler, MixtureOfSamplers,
MultiFoldDistributedSampler, MultiLevelBatchSampler,
MultiLevelBatchSamplerMultiSource)
MultiLevelBatchSamplerMultiSource, ResolutionBatchSampler)
+105 -1
View File
@@ -7,7 +7,7 @@ import numbers
import os
import sys
from collections.abc import Iterable
from typing import Optional
from typing import List, Optional
import numpy as np
import torch
@@ -15,6 +15,8 @@ import torch.distributed as dist
from scepter.modules.data.sampler.base_sampler import BaseSampler
from scepter.modules.data.sampler.registry import SAMPLERS
from scepter.modules.data.utils.data_bucket import (BucketBatchIndex,
BucketManager)
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.directory import osp_path
from scepter.modules.utils.distribute import we
@@ -589,3 +591,105 @@ class LoopSampler(BaseSampler):
__class__.__name__,
LoopSampler.para_dict,
set_name=True)
@SAMPLERS.register_class()
class ResolutionBatchSampler(BaseSampler):
para_dict = {}
def __init__(self, cfg, logger):
super().__init__(cfg, logger)
self.data_file = cfg.DATA_FILE
self.fields = cfg.get('FIELDS', [])
self.num_fields = len(self.fields)
self.delimiter = cfg.get('DELIMITER', ',')
self.path_prefix = cfg.get('PATH_PREFIX', '')
self.batch_size = cfg.BATCH_SIZE
max_reso = cfg.get('MAX_RESO', (1024, 1024))
min_bucket_reso = cfg.get('MIN_BUCKET_RESO', 256)
max_bucket_reso = cfg.get('MAX_BUCKET_RESO', 1024)
bucket_reso_steps = cfg.get('BUCKET_RESO_STEPS', 64)
bucket_no_upscale = cfg.get('BUCKET_NO_UPSCALE', False)
rank = we.rank
self.rng = np.random.default_rng(self.seed + rank)
assert 'img_path' in self.fields and 'width' in self.fields and 'height' in self.fields
self.bucket_manager = BucketManager(max_reso=max_reso,
min_size=min_bucket_reso,
max_size=max_bucket_reso,
reso_steps=bucket_reso_steps,
no_upscale=bucket_no_upscale)
if not bucket_no_upscale:
self.bucket_manager.make_buckets()
else:
self.logger.info(
'min_bucket_reso and max_bucket_reso are ignored if bucket_no_upscale is set, '
'because bucket reso is defined by image size automatically / bucket_no_upscale'
)
self.data_map = {}
img_path_idx, width_idx, height_idx = self.fields.index(
'img_path'), self.fields.index('width'), self.fields.index(
'height')
with FS.get_from(self.data_file) as local_path:
with open(local_path) as f:
for i, line in enumerate(f):
items = line.strip()
item_sp = items.split(self.delimiter, self.num_fields - 1)
img_path, width, height = item_sp[img_path_idx], int(
item_sp[width_idx]), int(item_sp[height_idx])
item_sp[img_path_idx] = os.path.join(
self.path_prefix, img_path)
bucket_reso, resized_size, ar_error = self.bucket_manager.select_bucket(
width, height)
self.bucket_manager.add_image(reso=bucket_reso, image=i)
self.data_map[i] = item_sp
for i, (reso, bucket) in enumerate(
zip(self.bucket_manager.resos, self.bucket_manager.buckets)):
count = len(bucket)
if count > 0:
# self.logger.info(f"bucket {i}: resolution {reso}, bucket {bucket}, count: {len(bucket)}")
self.logger.info(
f'bucket {i}: resolution {reso}, count: {len(bucket)}')
self.buckets_indices: List[BucketBatchIndex] = []
for bucket_index, (reso, bucket) in enumerate(
zip(self.bucket_manager.resos, self.bucket_manager.buckets)):
batch_count = int(math.ceil(len(bucket) / self.batch_size))
for batch_index in range(batch_count):
self.buckets_indices.append(
BucketBatchIndex(bucket_index, self.batch_size,
batch_index, reso))
self.shuffle_buckets()
def shuffle_buckets(self):
np.random.shuffle(self.buckets_indices)
self.bucket_manager.shuffle()
def __iter__(self):
while True:
index = self.rng.choice(len(self.buckets_indices))
bucket_reso = self.buckets_indices[index].bucket_reso
bucket_width, bucket_height = bucket_reso
bucket = self.bucket_manager.buckets[
self.buckets_indices[index].bucket_index]
batches = self.rng.choice(bucket, self.batch_size)
# image_index = self.buckets_indices[index].batch_index * self.batch_size
# batch = bucket[image_index : image_index + self.batch_size]
fields = self.fields + ['image_size', 'prompt_prefix']
batches = [
self.data_map[idx] + [[bucket_height, bucket_width], fields]
for idx in batches
]
yield batches
def __len__(self):
return sys.maxsize
@staticmethod
def get_config_template():
return dict_to_yaml('SAMPLERS',
__class__.__name__,
ResolutionBatchSampler.para_dict,
set_name=True)
+4
View File
@@ -0,0 +1,4 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.data.utils.data_bucket import BucketManager
+231
View File
@@ -0,0 +1,231 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import math
import random
from typing import List, NamedTuple
import numpy as np
def make_bucket_resolutions(max_reso,
min_size=256,
max_size=1024,
divisible=64):
max_width, max_height = max_reso
max_area = (max_width // divisible) * (max_height // divisible)
resos = set()
size = int(math.sqrt(max_area)) * divisible
resos.add((size, size))
size = min_size
while size <= max_size:
width = size
height = min(max_size, (max_area // (width // divisible)) * divisible)
resos.add((width, height))
resos.add((height, width))
# # make additional resos
# if width >= height and width - divisible >= min_size:
# resos.add((width - divisible, height))
# resos.add((height, width - divisible))
# if height >= width and height - divisible >= min_size:
# resos.add((width, height - divisible))
# resos.add((height - divisible, width))
size += divisible
resos = list(resos)
resos.sort()
return resos
class BucketBatchIndex(NamedTuple):
bucket_index: int
bucket_batch_size: int
batch_index: int
bucket_reso: List[int]
class BucketManager:
def __init__(self,
max_reso,
min_size=256,
max_size=1024,
reso_steps=64,
no_upscale=False) -> None:
self.no_upscale = no_upscale
if max_reso is None:
self.max_reso = None
self.max_area = None
else:
self.max_reso = max_reso
self.max_area = max_reso[0] * max_reso[1]
self.min_size = min_size
self.max_size = max_size
self.reso_steps = reso_steps
self.resos = []
self.reso_to_id = {}
self.buckets = []
def add_image(self, reso, image):
bucket_id = self.reso_to_id[reso]
self.buckets[bucket_id].append(image)
def shuffle(self):
for bucket in self.buckets:
random.shuffle(bucket)
def sort(self):
sorted_resos = self.resos.copy()
sorted_resos.sort()
sorted_buckets = []
sorted_reso_to_id = {}
for i, reso in enumerate(sorted_resos):
bucket_id = self.reso_to_id[reso]
sorted_buckets.append(self.buckets[bucket_id])
sorted_reso_to_id[reso] = i
self.resos = sorted_resos
self.buckets = sorted_buckets
self.reso_to_id = sorted_reso_to_id
def make_buckets(self):
resos = make_bucket_resolutions(self.max_reso, self.min_size,
self.max_size, self.reso_steps)
self.set_predefined_resos(resos)
def set_predefined_resos(self, resos):
self.predefined_resos = resos.copy()
self.predefined_resos_set = set(resos)
self.predefined_aspect_ratios = np.array([w / h for w, h in resos])
def add_if_new_reso(self, reso):
if reso not in self.reso_to_id:
bucket_id = len(self.resos)
self.reso_to_id[reso] = bucket_id
self.resos.append(reso)
self.buckets.append([])
# print(reso, bucket_id, len(self.buckets))
def round_to_steps(self, x):
x = int(x + 0.5)
return x - x % self.reso_steps
def select_bucket(self, image_width, image_height):
aspect_ratio = image_width / image_height
if not self.no_upscale:
reso = (image_width, image_height)
if reso in self.predefined_resos_set:
pass
else:
ar_errors = self.predefined_aspect_ratios - aspect_ratio
predefined_bucket_id = np.abs(ar_errors).argmin()
reso = self.predefined_resos[predefined_bucket_id]
ar_reso = reso[0] / reso[1]
if aspect_ratio > ar_reso:
scale = reso[1] / image_height
else:
scale = reso[0] / image_width
resized_size = (int(image_width * scale + 0.5),
int(image_height * scale + 0.5))
# print("use predef", image_width, image_height, reso, resized_size)
else:
if image_width * image_height > self.max_area:
resized_width = math.sqrt(self.max_area * aspect_ratio)
resized_height = self.max_area / resized_width
assert abs(resized_width / resized_height -
aspect_ratio) < 1e-2, 'aspect is illegal'
b_width_rounded = self.round_to_steps(resized_width)
b_height_in_wr = self.round_to_steps(b_width_rounded /
aspect_ratio)
ar_width_rounded = b_width_rounded / b_height_in_wr
b_height_rounded = self.round_to_steps(resized_height)
b_width_in_hr = self.round_to_steps(b_height_rounded *
aspect_ratio)
ar_height_rounded = b_width_in_hr / b_height_rounded
# print(b_width_rounded, b_height_in_wr, ar_width_rounded)
# print(b_width_in_hr, b_height_rounded, ar_height_rounded)
if abs(ar_width_rounded -
aspect_ratio) < abs(ar_height_rounded - aspect_ratio):
resized_size = (b_width_rounded,
int(b_width_rounded / aspect_ratio + 0.5))
else:
resized_size = (int(b_height_rounded * aspect_ratio + 0.5),
b_height_rounded)
# print(resized_size)
else:
resized_size = (image_width, image_height)
bucket_width = resized_size[0] - resized_size[0] % self.reso_steps
bucket_height = resized_size[1] - resized_size[1] % self.reso_steps
# print("use arbitrary", image_width, image_height, resized_size, bucket_width, bucket_height)
reso = (bucket_width, bucket_height)
self.add_if_new_reso(reso)
ar_error = (reso[0] / reso[1]) - aspect_ratio
return reso, resized_size, ar_error
if __name__ == '__main__':
image_size_list = [(256, 256), (512, 378), (378, 512), (1024, 1024),
(768, 1024), (768, 768), (256, 1024), (512, 512)]
image_path_list = [f'image_path_{i}' for i in range(len(image_size_list))]
max_reso = (512, 1024)
min_bucket_reso = 256
max_bucket_reso = 1024
bucket_reso_steps = 64
bucket_no_upscale = False
bucket_manager = BucketManager(max_reso=max_reso,
min_size=min_bucket_reso,
max_size=max_bucket_reso,
reso_steps=bucket_reso_steps,
no_upscale=bucket_no_upscale)
if not bucket_no_upscale:
bucket_manager.make_buckets()
else:
print(
'min_bucket_reso and max_bucket_reso are ignored if bucket_no_upscale is set, '
'because bucket reso is defined by image size automatically / bucket_no_upscale'
)
for i, (path, size) in enumerate(zip(image_path_list, image_size_list)):
image_width, image_height = size
bucket_reso, resized_size, ar_error = bucket_manager.select_bucket(
image_width, image_height)
print(i, size, bucket_reso, resized_size, ar_error)
bucket_manager.add_image(reso=bucket_reso, image=path)
for i, (reso, bucket) in enumerate(
zip(bucket_manager.resos, bucket_manager.buckets)):
count = len(bucket)
if count > 0:
print(
f'bucket {i}: resolution {reso}, bucket {bucket}, count: {len(bucket)}'
)
batch_size = 2
buckets_indices: List[BucketBatchIndex] = []
for bucket_index, bucket in enumerate(bucket_manager.buckets):
batch_count = int(math.ceil(len(bucket) / batch_size))
for batch_index in range(batch_count):
buckets_indices.append(
BucketBatchIndex(bucket_index, batch_size, batch_index))
def shuffle_buckets():
random.shuffle(buckets_indices)
bucket_manager.shuffle()
shuffle_buckets()
@@ -2,18 +2,23 @@
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import os
import warnings
import torch
import torch.nn as nn
import torchvision.transforms as TT
from PIL.Image import Image
from swift import SwiftModel
from scepter.modules.model.registry import TUNERS
from scepter.modules.utils.config import Config
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
try:
from swift import SwiftModel
except Exception:
warnings.warn('Import swift failed, please check it.')
class ControlInference():
def __init__(self, logger=None):
@@ -89,7 +94,7 @@ class ControlInference():
def get_control_input(self, control_model, control_cond_image, height,
width):
hints = []
if control_cond_image and control_model:
if control_cond_image is not None and control_model is not None:
if not isinstance(control_model, list):
control_model = [control_model]
if not isinstance(control_cond_image, list):
@@ -15,6 +15,7 @@ from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, MODELS,
TOKENIZERS)
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from scepter.studio.utils.env import get_available_memory
from .control_inference import ControlInference
from .tuner_inference import TunerInference
@@ -243,8 +244,22 @@ class DiffusionInference():
def unload(self, module):
if module is None:
return module
module['model'] = module['model'].to('cpu')
module['device'] = 'cpu'
mem = get_available_memory()
free_mem = int(mem['available'] / (1024**2))
total_mem = int(mem['total'] / (1024**2))
if free_mem < 0.5 * total_mem:
if module['model'] is not None:
module['model'] = module['model'].to('cpu')
del module['model']
module['model'] = None
module['device'] = 'offline'
print('delete module')
else:
if module['model'] is not None:
module['model'] = module['model'].to('cpu')
module['device'] = 'cpu'
else:
module['device'] = 'offline'
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
return module
@@ -262,7 +277,7 @@ class DiffusionInference():
module = self.load(module)
self.loaded_model[name] = module
return module
elif module['device'] == 'cpu':
elif module['device'] == 'cpu' or module['device'] == "offline":
module = self.load(module)
return module
else:
@@ -395,7 +410,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)):
@@ -474,10 +489,14 @@ class DiffusionInference():
enabled=dtype == 'float16',
dtype=getattr(torch, dtype)):
if self.tokenizer:
if not hasattr(get_model(self.cond_stage_model), 'tokenizer'):
setattr(get_model(self.cond_stage_model), 'tokenizer',
self.tokenizer)
context = getattr(get_model(self.cond_stage_model),
function_name)(batch['tokens'])
null_context = getattr(get_model(self.cond_stage_model),
function_name)(batch_uc['tokens'])
else:
context = getattr(get_model(self.cond_stage_model),
function_name)(batch)
@@ -558,12 +577,13 @@ class DiffusionInference():
seed=seed,
condition_fn=None,
clamp=None,
sharpness=value_input.get('sharpness', 0.0),
percentile=None,
t_max=None,
t_min=None,
discard_penultimate_step=None,
intermediate_callback=intermediate_callback,
cat_uc=cat_uc,
cat_uc=value_input.get('cat_uc', cat_uc),
**kwargs)
self.dynamic_unload(self.diffusion_model,
@@ -0,0 +1,310 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import os.path
import random
from collections import OrderedDict
import gradio as gr
import torch
import torchvision.transforms.functional as TF
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
from .diffusion_inference import DiffusionInference
def get_model(model_tuple):
assert 'model' in model_tuple
return model_tuple['model']
class LargenInference(DiffusionInference):
'''
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 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
@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,
largen_state=False,
**kwargs):
if not largen_state:
raise gr.Error('LARGEN model must be used with LAR-Gen settings')
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=False)
# 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=False)
# 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=False)
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,330 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import random
import gradio as gr
import torch
import torch.nn.functional as F
import torchvision.transforms.functional as TF
from scepter.modules.utils.distribute import we
from .control_inference import ControlInference
from .diffusion_inference import DiffusionInference
from .tuner_inference import TunerInference
def get_model(model_tuple):
assert 'model' in model_tuple
return model_tuple['model']
class StyleboothInference(DiffusionInference):
'''
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'
]
self.tuner_infer = TunerInference(self.logger)
self.control_infer = ControlInference(self.logger)
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':
batch['prompt'] = value_dict['prompt']
batch_uc['prompt'] = value_dict['negative_prompt']
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 encode_condition(self, data, data2=None, type='text'):
cond_stage_model = get_model(self.cond_stage_model)
assert hasattr(self, 'tokenizer')
with torch.autocast(device_type='cuda', enabled=False):
if type == 'image' and (
hasattr(cond_stage_model, 'build_new_tokens')
and not hasattr(cond_stage_model, 'new_tokens_to_ids')):
cond_stage_model.build_new_tokens(self.tokenizer)
if type == 'text':
text = self.tokenizer(data).to(we.device_id)
return cond_stage_model.encode_text(text)
elif type == 'image':
return cond_stage_model.encode_image(data)
elif type == 'hybrid':
text = self.tokenizer(data).to(we.device_id)
return cond_stage_model.encode_text(text, data2)
def process_edit_image(self, images, height, width):
if not isinstance(images, list):
images = [images]
tensors = []
for img in images:
w, h = img.size
if not h == height or not w == width:
scale = max(width / w, height / h)
new_size = (int(h * scale), int(w * scale))
img = TF.resize(img,
new_size,
interpolation=TF.InterpolationMode.BICUBIC)
img = TF.center_crop(img, (height, width))
tensor = TF.to_tensor(img).to(we.device_id)
tensors.append(tensor)
tensors = TF.normalize(torch.stack(tensors),
mean=[0.5, 0.5, 0.5],
std=[0.5, 0.5, 0.5])
return tensors
@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,
stylebooth_state=False,
style_edit_image=None,
style_exemplar_image=None,
style_guide_scale_text=None,
style_guide_scale_image=None,
**kwargs):
if not stylebooth_state:
raise gr.Error('EDIT model must be used with StyleBooth settings')
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)
# 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')
self.control_infer.register_controllers(control_model,
self.diffusion_model)
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=True)
# first stage encode
image = input.pop('image', None)
if image is not None and img_to_img_strength > 0:
# run image2image
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.dynamic_load(self.first_stage_model, 'first_stage_model')
input_latent = self.encode_first_stage(image)
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.dynamic_load(self.cond_stage_model, 'cond_stage_model')
context = {}
if style_exemplar_image is not None:
if not isinstance(style_exemplar_image, list):
style_exemplar_image = [style_exemplar_image]
style_exemplar_image = [
TF.resize(x, (224, 224),
interpolation=TF.InterpolationMode.BICUBIC)
for x in style_exemplar_image
]
style_exemplar_image = [
TF.to_tensor(x).to(we.device_id) for x in style_exemplar_image
]
style_exemplar_image = TF.normalize(
torch.stack(style_exemplar_image),
mean=[0.48145466, 0.4578275, 0.40821073],
std=[0.26862954, 0.26130258, 0.27577711])
image_feature = self.encode_condition(style_exemplar_image,
type='image')
context['crossattn'] = self.encode_condition(batch['prompt'],
image_feature,
type='hybrid')
else:
context['crossattn'] = self.encode_condition(batch['prompt'])
null_context = {}
null_context['crossattn'] = self.encode_condition(batch_uc['prompt'])
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=True)
model_kwargs = [{'cond': context}]
# style first stage encode
if style_edit_image is not None:
style_edit_image = self.process_edit_image(style_edit_image,
height, width)
self.dynamic_load(self.first_stage_model, 'first_stage_model')
cond_concat = self.encode_first_stage(style_edit_image)
cond_concat /= self.first_stage_model['paras']['scale_factor']
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=True)
context['concat'] = cond_concat
null_context['concat'] = torch.zeros_like(cond_concat)
mid_context = {}
mid_context.update(null_context)
mid_context.update({'concat': cond_concat})
model_kwargs.append({'cond': mid_context})
model_kwargs.append({'cond': null_context})
# 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=input_latent,
denoising_strength=img_to_img_strength
if input_latent is not None else 1.0,
refine_strength=refine_strength,
solver=value_input.get('sample', 'ddim'),
model=get_model(self.diffusion_model),
model_kwargs=model_kwargs,
steps=value_input.get('sample_steps', 50),
guide_scale={
'text': style_guide_scale_text,
'image': style_guide_scale_image
},
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,
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=value_input.get('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 '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()
# 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
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.backbone.autoencoder.ae_module import (Decoder,
Encoder)
Encoder,
RDecoder)
@@ -3,6 +3,7 @@
import torch
import torch.nn as nn
from einops import repeat
from torch.utils.checkpoint import checkpoint
from scepter.modules.model.backbone.autoencoder.ae_utils import (
@@ -244,6 +245,7 @@ class Decoder(BaseModel):
# compute in_ch_mult, block_in and curr_res at lowest res
block_in = self.ch * self.ch_mult[self.num_resolutions - 1]
self.block_in = block_in
curr_res = 1
# z to block_in
self.conv_in = torch.nn.Conv2d(self.z_channels,
@@ -340,3 +342,48 @@ class Decoder(BaseModel):
__class__.__name__,
Decoder.para_dict,
set_name=True)
@BACKBONES.register_class()
class RDecoder(Decoder):
def construct_model(self):
super().construct_model()
self.resize_level = nn.Sequential(
nn.Linear(self.block_in, self.block_in),
nn.SiLU(),
nn.Linear(self.block_in, self.block_in),
)
def forward(self, z, rembed=None):
# timestep embedding
temb = None
h = self.conv_in(z)
bs, channel, hdim, wdim = h.size()
if rembed is not None:
rembed = self.resize_level(rembed)
rembed = repeat(rembed, 'b e-> b e hd wd', hd=hdim, wd=wdim)
h = h + rembed
# middle
if not self.use_checkpoint:
h = self.mid_upsclae_transform(h, temb)
else:
h = checkpoint(self.mid_upsclae_transform, h, temb)
# end
if self.give_pre_end:
return h
h = self.norm_out(h)
h = nonlinearity(h)
h = self.conv_out(h)
if self.tanh_out:
h = torch.tanh(h)
return h
@staticmethod
def get_config_template():
return dict_to_yaml('BACKBONE',
__class__.__name__,
Decoder.para_dict,
set_name=True)
@@ -1,3 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.backbone.unet.unet_module import DiffusionUNet
from scepter.modules.model.backbone.unet.unet_module import (DiffusionUNet,
DiffusionUNetXL,
LargenUNetXL)
@@ -8,8 +8,9 @@ import torch
import torch.nn as nn
from scepter.modules.model.backbone.unet.unet_utils import (
Downsample, ResBlock, SpatialTransformer, Timestep,
TimestepEmbedSequential, Upsample, conv_nd, linear, normalization,
BasicTransformerBlock, Downsample, ResBlock, SpatialTransformer,
SpatialTransformerV2, Timestep, TimestepEmbedSequential,
TransformerBlockV2, Upsample, conv_nd, linear, normalization,
timestep_embedding, zero_module)
from scepter.modules.model.base_model import BaseModel
from scepter.modules.model.registry import BACKBONES
@@ -951,3 +952,443 @@ class DiffusionUNetXL(DiffusionUNet):
__class__.__name__,
DiffusionUNetXL.para_dict,
set_name=True)
@BACKBONES.register_class()
class LargenUNetXL(DiffusionUNetXL):
para_dict = {
'TRANSFORMER_BLOCK_TYPE': {
'value': 'att_v1'
},
'IMAGE_SCALE': {
'value': 0.0,
},
}
para_dict.update(DiffusionUNetXL.para_dict)
def __init__(self, cfg, logger):
super().__init__(cfg, logger=logger)
self.init_params(cfg)
self.construct_network()
def init_params(self, cfg):
super().init_params(cfg)
self.transformer_block_type = cfg.get('TRANSFORMER_BLOCK_TYPE',
'att_v1')
TRANSFORMER_BLOCKS = {
'att_v1': BasicTransformerBlock,
'att_v2': TransformerBlockV2,
}
assert self.transformer_block_type in list(TRANSFORMER_BLOCKS.keys())
self.transformer_block = TRANSFORMER_BLOCKS[
self.transformer_block_type]
self.image_scale = cfg.get('IMAGE_SCALE', 0.0)
self.use_refine = cfg.get('USE_REFINE', False)
def construct_network(self):
in_channels = self.in_channels
model_channels = self.model_channels
out_channels = self.out_channels
attention_resolutions = self.attention_resolutions
channel_mult = self.channel_mult
num_classes = self.num_classes
num_heads = self.num_heads
num_head_channels = self.num_head_channels
dims = self.dims
dropout = self.dropout
use_checkpoint = self.use_checkpoint
use_scale_shift_norm = self.use_scale_shift_norm
disable_self_attentions = self.disable_self_attentions
disable_middle_self_attn = self.disable_middle_self_attn
transformer_depth = self.transformer_depth
transformer_depth_middle = self.transformer_depth_middle
context_dim = self.context_dim
use_linear_in_transformer = self.use_linear_in_transformer
resblock_updown = self.resblock_updown
conv_resample = self.conv_resample
adm_in_channels = self.adm_in_channels
transformer_block = self.transformer_block
time_embed_dim = model_channels * 4
self.time_embed = nn.Sequential(
linear(model_channels, time_embed_dim),
nn.SiLU(),
linear(time_embed_dim, time_embed_dim),
)
if self.num_classes is not None:
if isinstance(self.num_classes, int):
self.label_emb = nn.Embedding(num_classes, time_embed_dim)
elif self.num_classes == 'continuous':
print('setting up linear c_adm embedding layer')
self.label_emb = nn.Linear(1, time_embed_dim)
elif self.num_classes == 'timestep':
self.label_emb = nn.Sequential(
Timestep(model_channels),
nn.Sequential(
linear(model_channels, time_embed_dim),
nn.SiLU(),
linear(time_embed_dim, time_embed_dim),
),
)
elif self.num_classes == 'sequential':
assert adm_in_channels is not None
self.label_emb = nn.Sequential(
nn.Sequential(
linear(adm_in_channels, time_embed_dim),
nn.SiLU(),
linear(time_embed_dim, time_embed_dim),
))
else:
raise ValueError()
self.input_blocks = nn.ModuleList([
TimestepEmbedSequential(
conv_nd(dims, in_channels, model_channels, 3, padding=1))
])
self._feature_size = model_channels
input_block_chans = [model_channels]
input_down_flag = [False]
ch = model_channels
ds = 1
for level, mult in enumerate(channel_mult):
for nr in range(self.num_res_blocks[level]):
layers = [
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=mult * model_channels,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
)
]
ch = mult * model_channels
if ds in attention_resolutions:
if num_head_channels == -1:
dim_head = ch // num_heads
else:
num_heads = ch // num_head_channels
dim_head = num_head_channels
disabled_sa = disable_self_attentions[level] if exists(
disable_self_attentions) else False
layers.append(
SpatialTransformerV2(
ch,
num_heads,
dim_head,
transformer_block=transformer_block,
depth=transformer_depth[level],
context_dim=context_dim,
disable_self_attn=disabled_sa,
use_linear=use_linear_in_transformer,
use_checkpoint=use_checkpoint))
self.input_blocks.append(TimestepEmbedSequential(*layers))
self._feature_size += ch
input_block_chans.append(ch)
input_down_flag.append(False)
if level != len(channel_mult) - 1:
out_ch = ch
self.input_blocks.append(
TimestepEmbedSequential(
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=out_ch,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
down=True,
) if resblock_updown else Downsample(
ch, conv_resample, dims=dims, out_channels=out_ch))
)
ch = out_ch
input_block_chans.append(ch)
input_down_flag.append(True)
ds *= 2
self._feature_size += ch
self._input_block_chans = copy.deepcopy(input_block_chans)
self._input_down_flag = input_down_flag
if num_head_channels == -1:
dim_head = ch // num_heads
else:
num_heads = ch // num_head_channels
dim_head = num_head_channels
self.middle_block = TimestepEmbedSequential(
ResBlock(
ch,
time_embed_dim,
dropout,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
),
SpatialTransformerV2(ch,
num_heads,
dim_head,
transformer_block=transformer_block,
depth=transformer_depth_middle,
context_dim=context_dim,
disable_self_attn=disable_middle_self_attn,
use_linear=use_linear_in_transformer,
use_checkpoint=use_checkpoint),
ResBlock(
ch,
time_embed_dim,
dropout,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
),
)
self._feature_size += ch
self._middle_block_chans = [ch]
self._output_block_chans = []
self.output_blocks = nn.ModuleList([])
for level, mult in list(enumerate(channel_mult))[::-1]:
for i in range(self.num_res_blocks[level] + 1):
ich = input_block_chans.pop()
layers = [
ResBlock(
ch + ich,
time_embed_dim,
dropout,
out_channels=model_channels * mult,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
)
]
ch = model_channels * mult
if ds in attention_resolutions:
if num_head_channels == -1:
dim_head = ch // num_heads
else:
num_heads = ch // num_head_channels
dim_head = num_head_channels
disabled_sa = disable_self_attentions[level] if exists(
disable_self_attentions) else False
layers.append(
SpatialTransformerV2(
ch,
num_heads,
dim_head,
transformer_block=transformer_block,
depth=transformer_depth[level],
context_dim=context_dim,
disable_self_attn=disabled_sa,
use_linear=use_linear_in_transformer,
use_checkpoint=use_checkpoint))
if level and i == self.num_res_blocks[level]:
out_ch = ch
layers.append(
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=out_ch,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
up=True,
) if resblock_updown else Upsample(
ch, conv_resample, dims=dims, out_channels=out_ch))
ds //= 2
self.output_blocks.append(TimestepEmbedSequential(*layers))
self._feature_size += ch
self._output_block_chans.append(ch)
self.out = nn.Sequential(
normalization(ch),
nn.SiLU(),
zero_module(
conv_nd(dims, model_channels, out_channels, 3, padding=1)),
)
if self.use_refine:
self.ref_time_embed = copy.deepcopy(self.time_embed)
self.ref_label_emb = copy.deepcopy(self.label_emb)
self.ref_input_blocks = copy.deepcopy(self.input_blocks)
self.ref_input_blocks[0] = TimestepEmbedSequential(
conv_nd(dims, 4, model_channels, 3, padding=1))
self.ref_middle_block = copy.deepcopy(self.middle_block)
self.ref_output_blocks = copy.deepcopy(self.output_blocks)
def load_pretrained_model(self, pretrained_model):
if pretrained_model is not None:
with FS.get_from(pretrained_model,
wait_finish=True) as local_model:
self.init_from_ckpt(local_model, ignore_keys=self.ignore_keys)
def init_from_ckpt(self, path, ignore_keys=list()):
if path.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(path)
else:
sd = torch.load(path, map_location='cpu')
new_sd = OrderedDict()
for k, v in sd.items():
ignored = False
for ik in ignore_keys:
if ik in k:
if we.rank == 0:
self.logger.info(
'Ignore key {} from state_dict.'.format(k))
ignored = True
break
if not ignored:
if k == 'input_blocks.0.0.weight':
if we.rank == 0:
self.logger.info(
'Partial initial key {} from state_dict.'.format(
k))
new_v = torch.empty(320, self.in_channels, 3, 3)
nn.init.zeros_(new_v)
new_v[:, :v.shape[1]] = v
new_sd[k] = new_v
if self.use_refine:
new_sd['ref_' + k] = v
else:
new_sd[k] = v
if self.use_refine:
new_sd['ref_' + k] = v
missing, unexpected = self.load_state_dict(new_sd, strict=False)
if we.rank == 0:
self.logger.info(
f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys'
)
if len(missing) > 0:
self.logger.info(f'Missing Keys:\n {missing}')
if len(unexpected) > 0:
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
def forward(self, x, t=None, cond=dict(), **kwargs):
t_emb = timestep_embedding(t,
self.model_channels,
repeat_only=False,
legacy=True)
emb = self.time_embed(t_emb)
if isinstance(cond, dict):
if 'y' in cond:
assert self.num_classes is not None
emb = emb + self.label_emb(cond['y'])
if self.use_refine:
ref_emb = self.ref_time_embed(t_emb)
assert 'null_y' in cond
cond_y = cond['y'].clone()
cond_y[:, :cond['null_y'].shape[1]] = cond['null_y']
ref_emb = ref_emb + self.ref_label_emb(cond_y)
if 'concat' in cond:
c = cond['concat']
x = torch.cat([x, c], dim=1)
context = cond.get('crossattn', None)
img_context = cond.get('img_crossattn', None)
task = cond['task']
image_scale = cond.get('image_scale', self.image_scale)
if 'Subject' in task and img_context is not None:
ip_enc_scale = image_scale
ip_dec_scale = image_scale
num_img_tokens = img_context.shape[1]
context = torch.cat([context, img_context], dim=1)
else:
ip_enc_scale = None
ip_dec_scale = None
num_img_tokens = None
ref = cond.get('ref_xt', None)
ref_context = cond.get('ref_crossattn', None)
else:
raise TypeError
hs = []
refs = []
h = x
if self.use_refine:
assert ref is not None and ref_context is not None
for i, (ref_module, module) in enumerate(
zip(self.ref_input_blocks, self.input_blocks)):
ref = ref_module(ref, ref_emb, ref_context, caching=None)
h = module(h,
emb,
context,
caching=None,
scale=ip_enc_scale,
num_img_token=num_img_tokens)
refs.append(ref)
hs.append(h)
ref = self.ref_middle_block(ref,
ref_emb,
ref_context,
caching=None)
h = self.middle_block(h,
emb,
context,
caching=None,
scale=ip_enc_scale,
num_img_token=num_img_tokens)
for i, (ref_module, module) in enumerate(
zip(self.ref_output_blocks, self.output_blocks)):
cache = []
ref = torch.cat([ref, refs.pop()], dim=1)
ref = ref_module(ref,
ref_emb,
ref_context,
caching='write',
cache=cache)
h = torch.cat([h, hs.pop()], dim=1)
h = module(h,
emb,
context,
caching='read',
cache=cache,
scale=ip_dec_scale,
num_img_token=num_img_tokens)
else:
for module in self.input_blocks:
h = module(h,
emb,
context,
caching=None,
scale=ip_enc_scale,
num_img_token=num_img_tokens)
hs.append(h)
h = self.middle_block(h,
emb,
context,
caching=None,
scale=ip_enc_scale,
num_img_token=num_img_tokens)
for module in self.output_blocks:
h = torch.cat([h, hs.pop()], dim=1)
h = module(h,
emb,
context,
caching=None,
scale=ip_dec_scale,
num_img_token=num_img_tokens)
out = self.out(h)
return out
@staticmethod
def get_config_template():
return dict_to_yaml('BACKBONE',
__class__.__name__,
LargenUNetXL.para_dict,
set_name=True)
@@ -10,10 +10,12 @@ 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 torch.utils.checkpoint import checkpoint
from scepter.modules.model.utils.basic_utils import checkpoint, default, exists
from scepter.modules.model.utils.basic_utils import default, exists
try:
import xformers
@@ -171,12 +173,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:
@@ -357,8 +361,10 @@ class ResBlock(TimestepBlock):
:param emb: an [N x emb_channels] Tensor of timestep embeddings.
:return: an [N x C x ...] Tensor of outputs.
"""
return checkpoint(self._forward, (x, emb), self.parameters(),
self.use_checkpoint)
if self.use_checkpoint:
return checkpoint(self._forward, x, emb)
else:
return self._forward(x, emb)
def _forward(self, x, emb):
if self.updown:
@@ -419,8 +425,10 @@ class AttentionBlock(nn.Module):
self.proj_out = zero_module(conv_nd(1, channels, channels, 1))
def forward(self, x):
return checkpoint(self._forward, (x, ), self.parameters(),
self.use_checkpoint)
if self.use_checkpoint:
return checkpoint(self._forward, x)
else:
return self._forward(x)
def _forward(self, x):
b, c, *spatial = x.shape
@@ -864,6 +872,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,
@@ -897,8 +991,11 @@ class BasicTransformerBlock(nn.Module):
self.use_checkpoint = use_checkpoint
def forward(self, x, context=None):
return checkpoint(self._forward, (x, context), self.parameters(),
self.use_checkpoint)
if self.use_checkpoint:
return checkpoint(self._forward, x, context)
else:
return self._forward(x, context)
def _forward(self, x, context=None):
x = self.attn1(self.norm1(x),
@@ -908,6 +1005,65 @@ class BasicTransformerBlock(nn.Module):
return x
class TransformerBlockV2(nn.Module):
def __init__(self,
query_dim,
n_heads,
d_head,
dropout=0.,
context_dim=None,
gated_ff=True,
use_checkpoint=False,
disable_self_attn=False):
super().__init__()
self.disable_self_attn = disable_self_attn
self.attn1 = MemoryEfficientCrossAttention(query_dim=query_dim,
heads=n_heads,
dim_head=d_head,
dropout=dropout,
context_dim=None)
self.ff = FeedForward(query_dim, dropout=dropout, glu=gated_ff)
self.attn2 = XFormersMHA_IP(query_dim=query_dim,
heads=n_heads,
dim_head=d_head,
context_dim=context_dim)
self.norm1 = nn.LayerNorm(query_dim)
self.norm2 = nn.LayerNorm(query_dim)
self.norm3 = nn.LayerNorm(query_dim)
self.use_checkpoint = use_checkpoint
def forward(self,
x,
context,
caching=None,
cache=None,
scale=None,
num_img_token=None,
**kwargs):
y = self.norm1(x)
if caching == 'write':
assert isinstance(cache, list)
cache.append(y)
x = self.attn1(y, context=None) + x
elif caching == 'read':
assert isinstance(cache, list) and len(cache) > 0
c = cache.pop(0)
self_ctx = torch.cat([y, c], dim=1)
x = self.attn1(y, context=self_ctx) + x
elif caching is None:
x = self.attn1(y, context=None) + x
else:
assert False
x = self.attn2(self.norm2(x),
context=context,
scale=scale,
num_img_token=num_img_token) + x
x = self.ff(self.norm3(x)) + x
return x
class SpatialTransformer(nn.Module):
"""
Transformer block for image-like data.
@@ -1003,3 +1159,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):
+5 -5
View File
@@ -1,7 +1,7 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.embedder.embedder import (ConcatTimestepEmbedderND,
FrozenCLIPEmbedder,
FrozenOpenCLIPEmbedder,
FrozenOpenCLIPEmbedder2,
GeneralConditioner)
from scepter.modules.model.embedder.embedder import (
ConcatTimestepEmbedderND, FrozenCLIPEmbedder, FrozenOpenCLIPEmbedder,
FrozenOpenCLIPEmbedder2, GeneralConditioner, IPAdapterPlusEmbedder,
RefCrossEmbedder)
+135 -31
View File
@@ -22,9 +22,10 @@ from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from .base_embedder import BaseEmbedder
from .resampler import Resampler
try:
from transformers import CLIPTextModel, CLIPTokenizer
from transformers import CLIPTextModel, CLIPTokenizer, CLIPVisionModelWithProjection
except Exception as e:
warnings.warn(
f'Import transformers error, please deal with this problem: {e}')
@@ -513,6 +514,93 @@ class ConcatTimestepEmbedderND(BaseEmbedder):
set_name=True)
@EMBEDDERS.register_class()
class IPAdapterPlusEmbedder(BaseEmbedder):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
with FS.get_dir_to_local_dir(cfg.CLIP_DIR,
wait_finish=True) as local_path:
self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(
local_path)
self.image_proj_model = Resampler(
dim=self.cfg.get('IN_DIM', 768),
depth=self.cfg.get('DEPTH', 4),
dim_head=64,
heads=self.cfg.get('HEADS', 12),
num_queries=self.cfg.get('NUM_TOKENS', 16),
embedding_dim=self.image_encoder.config.hidden_size,
output_dim=self.cfg.get('CROSSATTN_DIM', 768),
ff_mult=4,
)
with FS.get_from(cfg.PRETRAINED_MODEL, wait_finish=True) as local_path:
ckpt = torch.load(local_path, map_location='cpu')
self.image_proj_model.load_state_dict(ckpt['image_proj'],
strict=True)
self.patch_projector = nn.Linear(self.image_encoder.config.hidden_size,
self.cfg.get('CROSSATTN_DIM', 768))
def encode(self, ref_ip, ref_detail):
encoder_output = self.image_encoder(ref_ip, output_hidden_states=True)
image_prompt_embeds = self.image_proj_model(
encoder_output.hidden_states[-2])
encoder_output_2 = self.image_encoder(ref_detail,
output_hidden_states=True)
image_patch_embeds = self.patch_projector(
encoder_output_2.last_hidden_state)
out = {
'img_crossattn': image_prompt_embeds,
'ref_crossattn': image_patch_embeds,
}
return out
def forward(self, ref_ip, ref_detail):
return self.encode(ref_ip, ref_detail)
class RefCrossEmbedder(BaseEmbedder):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
with FS.get_dir_to_local_dir(cfg.CLIP_DIR,
wait_finish=True) as local_path:
self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(
local_path)
self.patch_projector = nn.Linear(self.image_encoder.config.hidden_size,
self.cfg.get('CROSSATTN_DIM', 768))
def encode(self, img):
encoder_output = self.image_encoder(img, output_hidden_states=True)
image_patch_embeds = self.patch_projector(
encoder_output.last_hidden_state)
out = {
'ref_crossattn': image_patch_embeds,
}
return out
def forward(self, img):
return self.encode(img)
@EMBEDDERS.register_class()
class TransparentEmbedder(BaseEmbedder):
def forward(self, *args):
out = dict()
for key, val in zip(self.input_keys, args):
out[key] = val
return out
@EMBEDDERS.register_class()
class NoiseConcatEmbedder(BaseEmbedder):
def forward(self, *args):
return {'concat': torch.cat(args, dim=1)}
@EMBEDDERS.register_class()
class GeneralConditioner(BaseEmbedder):
OUTPUT_DIM2KEYS = {2: 'y', 3: 'crossattn', 4: 'concat', 5: 'concat'}
@@ -598,42 +686,58 @@ class GeneralConditioner(BaseEmbedder):
with embedding_context():
if hasattr(embedder, 'input_key') and (embedder.input_key
is not None):
if embedder.input_key not in batch:
continue
if embedder.legacy_ucg_val is not None:
batch = self.possibly_get_ucg_val(embedder, batch)
emb_out = embedder(batch[embedder.input_key])
elif hasattr(embedder, 'input_keys'):
if any([k not in batch for k in embedder.input_keys]):
continue
emb_out = embedder(
*[batch[k] for k in embedder.input_keys])
assert isinstance(
emb_out, (torch.Tensor, list, tuple)
), f'encoder outputs must be tensors or a sequence, but got {type(emb_out)}'
if not isinstance(emb_out, (list, tuple)):
emb_out = [emb_out]
for emb in emb_out:
# print("emb.shape", emb.shape)
# print("emb.input_keys", embedder.input_keys)
out_key = self.OUTPUT_DIM2KEYS[emb.dim()]
if embedder.ucg_rate > 0.0 and embedder.legacy_ucg_val is None:
emb = (expand_dims_like(
torch.bernoulli(
(1.0 - embedder.ucg_rate) *
torch.ones(emb.shape[0], device=emb.device)),
emb,
) * emb)
if (hasattr(embedder, 'input_keys')):
if np.sum(
np.array([
key in force_zero_embeddings
for key in embedder.input_keys
])) > 0:
emb = torch.zeros_like(emb)
if out_key in output:
output[out_key] = torch.cat((output[out_key], emb),
self.KEY2CATDIM[out_key])
else:
output[out_key] = emb
# if "y" in output:
# print("out.shape", output["y"].shape)
if isinstance(emb_out, dict):
for key, val in emb_out.items():
if key in output:
assert key in self.KEY2CATDIM
output[key] = torch.cat([output[key], val],
dim=self.KEY2CATDIM[key])
else:
output[key] = val
else:
assert isinstance(
emb_out, (torch.Tensor, list, tuple)
), f'encoder outputs must be tensors or a sequence, but got {type(emb_out)}'
if not isinstance(emb_out, (list, tuple)):
emb_out = [emb_out]
for emb in emb_out:
out_key = self.OUTPUT_DIM2KEYS[emb.dim()]
if embedder.ucg_rate > 0.0 and embedder.legacy_ucg_val is None:
emb = (expand_dims_like(
torch.bernoulli(
(1.0 - embedder.ucg_rate) *
torch.ones(emb.shape[0], device=emb.device)),
emb,
) * emb)
if (hasattr(embedder, 'input_keys')):
if np.sum(
np.array([
key in force_zero_embeddings
for key in embedder.input_keys
])) > 0:
emb = torch.zeros_like(emb)
if out_key in output:
output[out_key] = torch.cat((output[out_key], emb),
self.KEY2CATDIM[out_key])
else:
output[out_key] = emb
return output
def get_unconditional_conditioning(self,
+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)
@@ -52,7 +52,6 @@ class DiagonalGaussianDistribution(object):
dim=dims)
def mode(self):
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,6 +194,7 @@ class GaussianDiffusion(object):
guide_scale=None,
guide_rescale=None,
clamp=None,
sharpness=0.0,
percentile=None,
cat_uc=False,
**kwargs):
@@ -79,54 +221,99 @@ class GaussianDiffusion(object):
# prediction
if guide_scale is None:
assert isinstance(model_kwargs, dict)
out = model(xt, t=t, **model_kwargs, **kwargs)
if isinstance(model_kwargs, dict):
out = model(xt, t=t, **model_kwargs, **kwargs)
elif isinstance(model_kwargs, list) and len(model_kwargs) > 0:
out = model(xt, t=t, **model_kwargs[0], **kwargs)
else:
raise Exception('Error')
else:
# classifier-free guidance (arXiv:2207.12598)
# model_kwargs[0]: conditional kwargs
# model_kwargs[1]: non-conditional kwargs
assert isinstance(model_kwargs, list) and len(model_kwargs) == 2
if guide_scale == 1.:
out = model(xt, t=t, **model_kwargs[0], **kwargs)
else:
if cat_uc:
def parse_model_kwargs(prev_value, value):
if isinstance(value, torch.Tensor):
prev_value = torch.cat([prev_value, value], dim=0)
elif isinstance(value, dict):
for k, v in value.items():
prev_value[k] = parse_model_kwargs(
prev_value[k], v)
elif isinstance(value, list):
for idx, v in enumerate(value):
prev_value[idx] = parse_model_kwargs(
prev_value[idx], v)
return prev_value
all_model_kwargs = copy.deepcopy(model_kwargs[0])
for model_kwarg in model_kwargs[1:]:
for key, value in model_kwarg.items():
all_model_kwargs[key] = parse_model_kwargs(
all_model_kwargs[key], value)
all_out = model(xt.repeat(2, 1, 1, 1),
t=t.repeat(2),
**all_model_kwargs,
**kwargs)
y_out, u_out = all_out.chunk(2)
assert isinstance(model_kwargs, list) and len(model_kwargs) >= 2
if isinstance(guide_scale, float) or isinstance(guide_scale, int):
assert len(model_kwargs) == 2
if guide_scale == 1.:
out = model(xt, t=t, **model_kwargs[0], **kwargs)
else:
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
u_out = model(xt, t=t, **model_kwargs[1], **kwargs)
out = u_out + guide_scale * (y_out - u_out)
if cat_uc:
# rescale the output according to arXiv:2305.08891
if guide_rescale is not None:
assert guide_rescale >= 0 and guide_rescale <= 1
ratio = (y_out.flatten(1).std(dim=1) /
(out.flatten(1).std(dim=1) +
1e-12)).view((-1, ) + (1, ) * (y_out.ndim - 1))
out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0
def parse_model_kwargs(prev_value, value):
if isinstance(value, torch.Tensor):
prev_value = torch.cat([prev_value, value],
dim=0)
elif isinstance(value, dict):
for k, v in value.items():
prev_value[k] = parse_model_kwargs(
prev_value[k], v)
elif isinstance(value, list):
for idx, v in enumerate(value):
prev_value[idx] = parse_model_kwargs(
prev_value[idx], v)
return prev_value
all_model_kwargs = copy.deepcopy(model_kwargs[0])
for model_kwarg in model_kwargs[1:]:
for key, value in model_kwarg.items():
all_model_kwargs[key] = parse_model_kwargs(
all_model_kwargs[key], value)
all_out = model(xt.repeat(2, 1, 1, 1),
t=t.repeat(2),
**all_model_kwargs,
**kwargs)
y_out, u_out = all_out.chunk(2)
else:
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
u_out = model(xt, t=t, **model_kwargs[1], **kwargs)
# todo sharpness
# sharpness sampling
if sharpness is not None and sharpness > 0:
positive_x0 = alphas * xt - sigmas * y_out
negative_x0 = alphas * xt - sigmas * u_out
positive_eps = xt - positive_x0
negative_eps = xt - negative_x0
global_diffusion_progress = (
1 - t / 999.0).detach().cpu().numpy().tolist()[0]
alpha = 0.001 * sharpness * global_diffusion_progress
positive_eps_degraded = adaptive_anisotropic_filter(
x=positive_eps, g=positive_x0)
positive_eps_degraded_weighted = positive_eps_degraded * alpha + positive_eps * (
1.0 - alpha)
final_eps = negative_eps + guide_scale * (
positive_eps_degraded_weighted - negative_eps)
final_x0 = xt - final_eps
out = (alphas * xt - final_x0) / sigmas
else:
out = u_out + guide_scale * (y_out - u_out)
elif isinstance(guide_scale, dict):
assert len(model_kwargs) == 3
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
m_out = model(xt, t=t, **model_kwargs[1], **kwargs)
u_out = model(xt, t=t, **model_kwargs[2], **kwargs)
out = u_out + guide_scale['image'] * (
m_out - u_out) + guide_scale['text'] * (y_out - m_out)
elif isinstance(guide_scale, list):
assert len(guide_scale) == len(model_kwargs) - 1
y_out = model(xt, t=t, **model_kwargs[0], **kwargs)
outs = [y_out]
for i in range(1, len(model_kwargs)):
outs.append(model(xt, t=t, **model_kwargs[i], **kwargs))
out = outs[-1]
for i in range(len(guide_scale)):
out += guide_scale[i] * (outs[-i - 2] - outs[-i - 1])
# rescale the output according to arXiv:2305.08891
if guide_rescale is not None and guide_rescale > 0.0:
assert guide_rescale >= 0 and guide_rescale <= 1
ratio = (
y_out.flatten(1).std(dim=1) /
(out.flatten(1).std(dim=1) + 1e-12)).view((-1, ) + (1, ) *
(y_out.ndim - 1))
out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0
# compute x0
if self.prediction_type == 'x0':
x0 = out
@@ -197,6 +384,7 @@ class GaussianDiffusion(object):
guide_scale=None,
guide_rescale=None,
clamp=None,
sharpness=0.0,
percentile=None,
solver='euler_a',
steps=20,
@@ -209,12 +397,16 @@ class GaussianDiffusion(object):
seed=-1,
intermediate_callback=None,
cat_uc=False,
add_noise=False,
free_steps=None,
step_offset=None,
**kwargs):
# sanity check
assert isinstance(steps, (int, torch.LongTensor))
assert t_max is None or (t_max > 0 and t_max <= self.num_timesteps - 1)
assert t_min is None or (t_min >= 0 and t_min < self.num_timesteps - 1)
assert discretization in (None, 'leading', 'linspace', 'trailing')
assert discretization in (None, 'leading', 'linspace', 'trailing',
'free')
assert discard_penultimate_step in (None, True, False)
assert return_intermediate in (None, 'x0', 'xt')
@@ -255,17 +447,51 @@ class GaussianDiffusion(object):
def model_fn(xt, sigma):
# denoising
t = self._sigma_to_t(sigma).repeat(len(xt)).round().long()
x0 = self.denoise(xt,
t,
None,
model,
model_kwargs,
guide_scale,
guide_rescale,
clamp,
percentile,
cat_uc=cat_uc,
**kwargs)[-2]
if isinstance(model_kwargs, list) and len(model_kwargs) == 2:
if isinstance(
model_kwargs[0]['cond'], dict) and \
'tar_x0' in model_kwargs[0]['cond'] and \
'tar_mask_latent' in model_kwargs[0]['cond']:
tar_x0 = model_kwargs[0]['cond']['tar_x0']
tar_mask = model_kwargs[0]['cond']['tar_mask_latent']
tar_xt = self.diffuse(x0=tar_x0, t=t)
xt = tar_xt * (1.0 - tar_mask) + xt * tar_mask
if isinstance(model_kwargs[0]['cond'],
dict) and 'ref_x0' in model_kwargs[0]['cond']:
model_kwargs[0]['cond']['ref_xt'] = self.diffuse(
x0=model_kwargs[0]['cond']['ref_x0'], t=t)
model_kwargs[1]['cond']['ref_xt'] = self.diffuse(
x0=model_kwargs[1]['cond']['ref_x0'], t=t)
if solver in ('onestep', 'multistep', 'multistep2', 'multistep3'):
x0 = self.denoise(xt,
t,
None,
model,
model_kwargs,
guide_scale,
guide_rescale,
clamp,
sharpness,
percentile,
cat_uc=cat_uc,
**kwargs)[-3]
else:
x0 = self.denoise(xt,
t,
None,
model,
model_kwargs,
guide_scale,
guide_rescale,
clamp,
sharpness,
percentile,
cat_uc=cat_uc,
**kwargs)[-2]
# collect intermediate outputs
if return_intermediate == 'xt':
@@ -291,10 +517,14 @@ class GaussianDiffusion(object):
elif discretization == 'trailing':
steps = torch.arange(t_max, t_min - 1,
-((t_max - t_min + 1) / steps))
elif discretization == 'free':
steps = torch.tensor(free_steps)
else:
raise NotImplementedError(
f'{discretization} discretization not implemented')
steps = steps.clamp_(t_min, t_max)
elif isinstance(steps, list):
steps = torch.tensor(steps)
steps = torch.as_tensor(steps,
dtype=torch.float32,
device=noise.device)
@@ -335,6 +565,23 @@ class GaussianDiffusion(object):
if discard_penultimate_step:
sigmas = torch.cat([sigmas[:-2], sigmas[-1:]])
kwargs['seed'] = seed
# add noise to x0
if add_noise:
if 'dm_steps' in kwargs:
if step_offset:
add_noise_step = -kwargs['dm_steps'] + step_offset
if add_noise_step < 0:
noise = self.diffuse(
noise,
torch.full((noise.shape[0], 1),
steps[add_noise_step],
dtype=torch.int))
else:
noise = self.diffuse(
noise,
torch.full((noise.shape[0], 1),
steps[-kwargs['dm_steps'] - 1],
dtype=torch.int))
# sampling
x0 = solver_fn(noise,
model_fn,
@@ -373,168 +620,6 @@ class GaussianDiffusion(object):
| torch.isinf(log_sigma)] = float('inf')
return log_sigma.exp()
@torch.no_grad()
def stochastic_encode(self, x0, t, steps):
# fast, but does not allow for exact reconstruction
# t serves as an index to gather the correct alphas
t_max = None
t_min = None
# discretization method
discretization = 'trailing' if self.prediction_type == 'v' else 'leading'
# timesteps
if isinstance(steps, int):
t_max = self.num_timesteps - 1 if t_max is None else t_max
t_min = 0 if t_min is None else t_min
steps = discretize_timesteps(t_max, t_min, steps, discretization)
steps = torch.as_tensor(steps).round().long().flip(0).to(x0.device)
# steps = torch.as_tensor(steps).round().long().to(x0.device)
# self.alphas_bar = torch.cumprod(1 - self.sigmas ** 2, dim=0)
# print('sigma: ', self.sigmas, len(self.sigmas))
# print('alpha_bar: ', self.alphas_bar, len(self.alphas_bar))
# print('steps: ', steps, len(steps))
# sqrt_alphas_cumprod = torch.sqrt(self.alphas_bar).to(x0.device)[steps]
# sqrt_one_minus_alphas_cumprod = torch.sqrt(1 - self.alphas_bar).to(x0.device)[steps]
sqrt_alphas_cumprod = self.alphas.to(x0.device)[steps]
sqrt_one_minus_alphas_cumprod = self.sigmas.to(x0.device)[steps]
# print('sigma: ', self.sigmas, len(self.sigmas))
# print('alpha: ', self.alphas, len(self.alphas))
# print('steps: ', steps, len(steps))
noise = torch.randn_like(x0)
return (
extract_into_tensor(sqrt_alphas_cumprod, t, x0.shape) * x0 +
extract_into_tensor(sqrt_one_minus_alphas_cumprod, t, x0.shape) *
noise)
@torch.no_grad()
def sample_img2img(self,
x,
noise,
model,
denoising_strength=1,
model_kwargs={},
condition_fn=None,
guide_scale=None,
guide_rescale=None,
clamp=None,
percentile=None,
solver='euler_a',
steps=20,
t_max=None,
t_min=None,
discretization=None,
discard_penultimate_step=None,
return_intermediate=None,
show_progress=False,
seed=-1,
**kwargs):
# sanity check
assert isinstance(steps, (int, torch.LongTensor))
assert t_max is None or (t_max > 0 and t_max <= self.num_timesteps - 1)
assert t_min is None or (t_min >= 0 and t_min < self.num_timesteps - 1)
assert discretization in (None, 'leading', 'linspace', 'trailing')
assert discard_penultimate_step in (None, True, False)
assert return_intermediate in (None, 'x0', 'xt')
# function of diffusion solver
solver_fn = {
'euler_ancestral': sample_img2img_euler_ancestral,
'euler': sample_img2img_euler,
}[solver]
# options
schedule = 'karras' if 'karras' in solver else None
discretization = discretization or 'linspace'
seed = seed if seed >= 0 else random.randint(0, 2**31)
if isinstance(steps, torch.LongTensor):
discard_penultimate_step = False
if discard_penultimate_step is None:
discard_penultimate_step = True if solver in (
'dpm2', 'dpm2_ancestral', 'dpmpp_2m_sde', 'dpm2_karras',
'dpm2_ancestral_karras', 'dpmpp_2m_sde_karras') else False
# function for denoising xt to get x0
intermediates = []
def get_scalings(sigma):
c_out = -sigma
c_in = 1 / (sigma**2 + 1.**2)**0.5
return c_out, c_in
def model_fn(xt, sigma):
# denoising
c_out, c_in = get_scalings(sigma)
t = self._sigma_to_t(sigma).repeat(len(xt)).round().long()
x0 = self.denoise(xt * c_in, t, None, model, model_kwargs,
guide_scale, guide_rescale, clamp, percentile,
**kwargs)[-2]
# collect intermediate outputs
if return_intermediate == 'xt':
intermediates.append(xt)
elif return_intermediate == 'x0':
intermediates.append(x0)
return xt + x0 * c_out
# get timesteps
if isinstance(steps, int):
steps += 1 if discard_penultimate_step else 0
t_max = self.num_timesteps - 1 if t_max is None else t_max
t_min = 0 if t_min is None else t_min
# discretize timesteps
if discretization == 'leading':
steps = torch.arange(t_min, t_max + 1,
(t_max - t_min + 1) / steps).flip(0)
elif discretization == 'linspace':
steps = torch.linspace(t_max, t_min, steps)
elif discretization == 'trailing':
steps = torch.arange(t_max, t_min - 1,
-((t_max - t_min + 1) / steps))
else:
raise NotImplementedError(
f'{discretization} discretization not implemented')
steps = steps.clamp_(t_min, t_max)
steps = torch.as_tensor(steps, dtype=torch.float32, device=x.device)
# get sigmas
sigmas = self._t_to_sigma(steps)
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
t_enc = int(min(denoising_strength, 0.999) * len(steps))
sigmas = sigmas[len(steps) - t_enc - 1:]
noise = x + noise * sigmas[0]
if schedule == 'karras':
if sigmas[0] == float('inf'):
sigmas = karras_schedule(
n=len(steps) - 1,
sigma_min=sigmas[sigmas > 0].min().item(),
sigma_max=sigmas[sigmas < float('inf')].max().item(),
rho=7.).to(sigmas)
sigmas = torch.cat([
sigmas.new_tensor([float('inf')]), sigmas,
sigmas.new_zeros([1])
])
else:
sigmas = karras_schedule(
n=len(steps),
sigma_min=sigmas[sigmas > 0].min().item(),
sigma_max=sigmas.max().item(),
rho=7.).to(sigmas)
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
if discard_penultimate_step:
sigmas = torch.cat([sigmas[:-2], sigmas[-1:]])
# sampling
x0 = solver_fn(noise,
model_fn,
sigmas,
seed=seed,
show_progress=show_progress,
**kwargs)
return (x0, intermediates) if return_intermediate is not None else x0
def extract_into_tensor(a, t, x_shape):
b, *_ = t.shape
+1 -1
View File
@@ -336,7 +336,7 @@ class LatentDiffusion(TrainModule):
h = int(meta['image_size'][0][0])
w = int(meta['image_size'][1][0])
image_size = [h, w]
if 'image_size' in kwargs:
if 'image_size' in kwargs and kwargs['image_size'] is not None:
image_size = kwargs.pop('image_size')
if isinstance(image_size, numbers.Number):
image_size = [image_size, image_size]

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