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: 265 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 300 KiB |
+17
-17
@@ -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)
|
||||
```
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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):
|
||||
|
||||
+17
-17
@@ -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)
|
||||
```
|
||||
|
||||
@@ -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,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",
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -21,6 +21,7 @@
|
||||
- [Acknowledgement](#acknowledgement)
|
||||
|
||||
## 🎉 News
|
||||
- [2024.04]: New [StyleBooth](https://ali-vilab.github.io/stylebooth-page/) demo on SCEPTER Studio, supporting `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/).
|
||||
@@ -90,14 +91,14 @@ print(next(iter(ms_train_dataset)))
|
||||
|
||||
#### CSV Format
|
||||
|
||||
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).
|
||||
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://modelscope.cn/api/v1/models/damo/scepter/repo?Revision=master&FilePath=datasets/3D_example_txt.zip)
|
||||
To facilitate starting training in command-line mode, you can use a dataset in text format, please refer to [3D_example_txt.zip](https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=datasets/3D_example_txt.zip)
|
||||
|
||||
```shell
|
||||
mkdir -p cache/datasets/ && wget 'https://modelscope.cn/api/v1/models/damo/scepter_scedit/repo?Revision=master&FilePath=dataset/3D_example_txt.zip' -O cache/datasets/3D_example_txt.zip && unzip cache/datasets/3D_example_txt.zip -d cache/datasets/ && rm cache/datasets/3D_example_txt.zip
|
||||
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
|
||||
```
|
||||
|
||||
### Training
|
||||
@@ -208,6 +209,25 @@ We deploy a work studio on Modelscope that includes only the inference tab, plea
|
||||
|
||||
## 🖼️ Gallery
|
||||
|
||||
### 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="https://github.com/hanzhn/datas/blob/main/scepter/readme/tuner_gold_dragon.jpeg?raw=true" 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>
|
||||
|
||||
|
||||
### LAR-Gen: Zoom Out
|
||||
<table>
|
||||
<tr>
|
||||
@@ -340,6 +360,12 @@ We deploy a work studio on Modelscope that includes only the inference tab, plea
|
||||
|:---------:|:----------:|:----------:|:----------:|
|
||||
| SD XL | 🪄 | 🪄 | ⏳ |
|
||||
|
||||
- StyleBooth
|
||||
|
||||
| **Text-Based** | **Exemplar-Based** |
|
||||
|:--------------:|:-----------------:|
|
||||
| 🪄 | ⏳ |
|
||||
|
||||
### Model URL
|
||||
|
||||
- ✅ indicates support for both training and inference.
|
||||
@@ -347,10 +373,11 @@ We deploy a work studio on Modelscope that includes only the inference tab, plea
|
||||
- ⏳ denotes that the module has not been integrated currently.
|
||||
- More models will be released in the future.
|
||||
|
||||
| Model | URL |
|
||||
|--------|------------------------------------------------------------------------------------------------------------------------------------------------|
|
||||
| Model | URL |
|
||||
|--------|-------------------------------------------------------------------------------------------------------------------------------------------|
|
||||
| SCEdit | [ModelScope](https://modelscope.cn/models/iic/scepter_scedit/summary) [HuggingFace](https://huggingface.co/scepter-studio/scepter_scedit) |
|
||||
| LAR-Gen | [ModelScope](https://www.modelscope.cn/models/iic/LARGEN/summary) |
|
||||
| LAR-Gen | [ModelScope](https://www.modelscope.cn/models/iic/LARGEN/summary) |
|
||||
| StyleBooth | [ModelScope](https://www.modelscope.cn/models/iic/stylebooth/summary) |
|
||||
|
||||
PS: Scripts running within the SCEPTER framework will automatically fetch and load models based on the required dependency files, eliminating the need for manual downloads.
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ albumentations
|
||||
bezier
|
||||
einops
|
||||
modelscope
|
||||
ms-swift>=1.5.2
|
||||
ms-swift>=2.0.1
|
||||
numpy
|
||||
open_clip_torch
|
||||
opencv-python
|
||||
|
||||
@@ -1,3 +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
|
||||
@@ -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
|
||||
|
||||
@@ -5,6 +5,7 @@ META:
|
||||
VERSION: 'SD_XL1.0'
|
||||
DESCRIPTION: "Stable Diffusion XL1.0"
|
||||
IS_DEFAULT: True
|
||||
IS_SHARE: True
|
||||
INFERENCE_PARAS:
|
||||
INFERENCE_BATCH_SIZE: 1
|
||||
INFERENCE_PREFIX: ""
|
||||
@@ -532,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:
|
||||
|
||||
@@ -4,6 +4,7 @@ META:
|
||||
VERSION: 'SD1.5'
|
||||
DESCRIPTION: "Stable Diffusion v1.5"
|
||||
IS_DEFAULT: False
|
||||
IS_SHARE: True
|
||||
INFERENCE_PARAS:
|
||||
INFERENCE_BATCH_SIZE: 1
|
||||
INFERENCE_PREFIX: ""
|
||||
@@ -244,7 +245,6 @@ SOLVER:
|
||||
GUIDE_SCALE: 7.5
|
||||
GUIDE_RESCALE:
|
||||
DISCRETIZATION: trailing
|
||||
IMAGE_SIZE: [512, 512]
|
||||
RUN_TRAIN_N: False
|
||||
#
|
||||
OPTIMIZER:
|
||||
@@ -274,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
|
||||
|
||||
@@ -4,6 +4,7 @@ META:
|
||||
VERSION: 'SD2.1'
|
||||
DESCRIPTION: "Stable Diffusion v2.1"
|
||||
IS_DEFAULT: False
|
||||
IS_SHARE: True
|
||||
INFERENCE_PARAS:
|
||||
INFERENCE_BATCH_SIZE: 1
|
||||
INFERENCE_PREFIX: ""
|
||||
@@ -186,7 +187,6 @@ SOLVER:
|
||||
GUIDE_SCALE: 7.5
|
||||
GUIDE_RESCALE:
|
||||
DISCRETIZATION: trailing
|
||||
IMAGE_SIZE: [768, 768]
|
||||
RUN_TRAIN_N: False
|
||||
#
|
||||
OPTIMIZER:
|
||||
@@ -216,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
|
||||
|
||||
@@ -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:
|
||||
|
||||

|
||||
|
||||
## 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}",可能会得到如下图像:
|
||||
|
||||

|
||||
|
||||
|
||||
## 模型使用
|
||||
### 命令行运行
|
||||
|
||||
* 使用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}
|
||||
}
|
||||
```
|
||||
@@ -1,2 +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' ]
|
||||
|
||||
@@ -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':
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -0,0 +1,4 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
from scepter.modules.data.utils.data_bucket import BucketManager
|
||||
@@ -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()
|
||||
@@ -94,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):
|
||||
|
||||
@@ -16,6 +16,7 @@ from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, MODELS,
|
||||
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,19 @@ 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:
|
||||
module['model'] = module['model'].to('cpu')
|
||||
module['device'] = 'cpu'
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
return module
|
||||
|
||||
@@ -16,6 +16,7 @@ from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, MODELS,
|
||||
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 scepter.studio.utils.env import get_available_memory
|
||||
|
||||
|
||||
def get_model(model_tuple):
|
||||
@@ -245,8 +246,19 @@ class LargenInference():
|
||||
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:
|
||||
module['model'] = module['model'].to('cpu')
|
||||
module['device'] = 'cpu'
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
return module
|
||||
|
||||
@@ -0,0 +1,663 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import os.path
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms.functional as TF
|
||||
from PIL.Image import Image
|
||||
|
||||
from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion
|
||||
from scepter.modules.model.network.diffusion.schedules import noise_schedule
|
||||
from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, MODELS,
|
||||
TOKENIZERS)
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
from ...studio.utils.env import get_available_memory
|
||||
from .control_inference import ControlInference
|
||||
from .tuner_inference import TunerInference
|
||||
|
||||
|
||||
def get_model(model_tuple):
|
||||
assert 'model' in model_tuple
|
||||
return model_tuple['model']
|
||||
|
||||
|
||||
class StyleboothInference():
|
||||
'''
|
||||
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 init_from_cfg(self, cfg):
|
||||
self.name = cfg.NAME
|
||||
self.is_default = cfg.get('IS_DEFAULT', False)
|
||||
module_paras = self.load_default(cfg.get('DEFAULT_PARAS', None))
|
||||
assert cfg.have('MODEL')
|
||||
cfg.MODEL = self.redefine_paras(cfg.MODEL)
|
||||
self.diffusion = self.load_schedule(cfg.MODEL.SCHEDULE)
|
||||
self.diffusion_model = self.infer_model(
|
||||
cfg.MODEL.DIFFUSION_MODEL, module_paras.get(
|
||||
'DIFFUSION_MODEL',
|
||||
None)) if cfg.MODEL.have('DIFFUSION_MODEL') else None
|
||||
self.first_stage_model = self.infer_model(
|
||||
cfg.MODEL.FIRST_STAGE_MODEL,
|
||||
module_paras.get(
|
||||
'FIRST_STAGE_MODEL',
|
||||
None)) if cfg.MODEL.have('FIRST_STAGE_MODEL') else None
|
||||
self.cond_stage_model = self.infer_model(
|
||||
cfg.MODEL.COND_STAGE_MODEL,
|
||||
module_paras.get(
|
||||
'COND_STAGE_MODEL',
|
||||
None)) if cfg.MODEL.have('COND_STAGE_MODEL') else None
|
||||
self.refiner_cond_model = self.infer_model(
|
||||
cfg.MODEL.REFINER_COND_MODEL,
|
||||
module_paras.get(
|
||||
'REFINER_COND_MODEL',
|
||||
None)) if cfg.MODEL.have('REFINER_COND_MODEL') else None
|
||||
self.refiner_diffusion_model = self.infer_model(
|
||||
cfg.MODEL.REFINER_MODEL, module_paras.get(
|
||||
'REFINER_MODEL',
|
||||
None)) if cfg.MODEL.have('REFINER_MODEL') else None
|
||||
self.tokenizer = TOKENIZERS.build(
|
||||
cfg.MODEL.TOKENIZER,
|
||||
logger=self.logger) if cfg.MODEL.have('TOKENIZER') else None
|
||||
|
||||
if self.tokenizer is not None:
|
||||
self.cond_stage_model['cfg'].KWARGS = {
|
||||
'vocab_size': self.tokenizer.vocab_size
|
||||
}
|
||||
|
||||
def redefine_paras(self, cfg):
|
||||
if cfg.get('PRETRAINED_MODEL', None):
|
||||
assert FS.isfile(cfg.PRETRAINED_MODEL)
|
||||
with FS.get_from(cfg.PRETRAINED_MODEL,
|
||||
wait_finish=True) as local_path:
|
||||
if local_path.endswith('safetensors'):
|
||||
from safetensors.torch import load_file as load_safetensors
|
||||
sd = load_safetensors(local_path)
|
||||
else:
|
||||
sd = torch.load(local_path, map_location='cpu')
|
||||
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
|
||||
else:
|
||||
continue
|
||||
if cfg.have('FIRST_STAGE_MODEL'):
|
||||
with open(first_stage_model_path + 'cache', 'wb') as f:
|
||||
torch.save(first_stage_model, f)
|
||||
os.rename(first_stage_model_path + 'cache',
|
||||
first_stage_model_path)
|
||||
self.logger.info(
|
||||
'First stage model has been processed.')
|
||||
if cfg.have('COND_STAGE_MODEL'):
|
||||
with open(cond_stage_model_path + 'cache', 'wb') as f:
|
||||
torch.save(cond_stage_model, f)
|
||||
os.rename(cond_stage_model_path + 'cache',
|
||||
cond_stage_model_path)
|
||||
self.logger.info(
|
||||
'Cond stage model has been processed.')
|
||||
if cfg.have('DIFFUSION_MODEL'):
|
||||
with open(diffusion_model_path + 'cache', 'wb') as f:
|
||||
torch.save(diffusion_model, f)
|
||||
os.rename(diffusion_model_path + 'cache',
|
||||
diffusion_model_path)
|
||||
self.logger.info('Diffusion model has been processed.')
|
||||
if not cfg.FIRST_STAGE_MODEL.get('PRETRAINED_MODEL', None):
|
||||
cfg.FIRST_STAGE_MODEL.PRETRAINED_MODEL = first_stage_model_path
|
||||
else:
|
||||
cfg.FIRST_STAGE_MODEL.RELOAD_MODEL = first_stage_model_path
|
||||
if not cfg.COND_STAGE_MODEL.get('PRETRAINED_MODEL', None):
|
||||
cfg.COND_STAGE_MODEL.PRETRAINED_MODEL = cond_stage_model_path
|
||||
else:
|
||||
cfg.COND_STAGE_MODEL.RELOAD_MODEL = cond_stage_model_path
|
||||
if not cfg.DIFFUSION_MODEL.get('PRETRAINED_MODEL', None):
|
||||
cfg.DIFFUSION_MODEL.PRETRAINED_MODEL = diffusion_model_path
|
||||
else:
|
||||
cfg.DIFFUSION_MODEL.RELOAD_MODEL = diffusion_model_path
|
||||
return cfg
|
||||
|
||||
def init_from_modules(self, modules):
|
||||
for k, v in modules.items():
|
||||
self.__setattr__(k, v)
|
||||
|
||||
def infer_model(self, cfg, module_paras=None):
|
||||
module = {
|
||||
'model': None,
|
||||
'cfg': cfg,
|
||||
'device': 'offline',
|
||||
'name': cfg.NAME,
|
||||
'function_info': {},
|
||||
'paras': {}
|
||||
}
|
||||
if module_paras is None:
|
||||
return module
|
||||
function_info = {}
|
||||
paras = {
|
||||
k.lower(): v
|
||||
for k, v in module_paras.get('PARAS', {}).items()
|
||||
}
|
||||
for function in module_paras.get('FUNCTION', []):
|
||||
input_dict = {}
|
||||
for inp in function.get('INPUT', []):
|
||||
if inp.lower() in self.input:
|
||||
input_dict[inp.lower()] = self.input[inp.lower()]
|
||||
function_info[function.NAME] = {
|
||||
'dtype': function.get('DTYPE', 'float32'),
|
||||
'input': input_dict
|
||||
}
|
||||
module['paras'] = paras
|
||||
module['function_info'] = function_info
|
||||
return module
|
||||
|
||||
def init_from_ckpt(self, path, model, ignore_keys=list()):
|
||||
if path.endswith('safetensors'):
|
||||
from safetensors.torch import load_file as load_safetensors
|
||||
sd = load_safetensors(path)
|
||||
else:
|
||||
sd = torch.load(path, map_location='cpu')
|
||||
|
||||
new_sd = OrderedDict()
|
||||
for k, v in sd.items():
|
||||
ignored = False
|
||||
for ik in ignore_keys:
|
||||
if ik in k:
|
||||
if we.rank == 0:
|
||||
self.logger.info(
|
||||
'Ignore key {} from state_dict.'.format(k))
|
||||
ignored = True
|
||||
break
|
||||
if not ignored:
|
||||
new_sd[k] = v
|
||||
|
||||
missing, unexpected = model.load_state_dict(new_sd, strict=False)
|
||||
if we.rank == 0:
|
||||
self.logger.info(
|
||||
f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys'
|
||||
)
|
||||
if len(missing) > 0:
|
||||
self.logger.info(f'Missing Keys:\n {missing}')
|
||||
if len(unexpected) > 0:
|
||||
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
|
||||
|
||||
def load(self, module):
|
||||
if module['device'] == 'offline':
|
||||
if module['cfg'].NAME in MODELS.class_map:
|
||||
model = MODELS.build(module['cfg'], logger=self.logger).eval()
|
||||
elif module['cfg'].NAME in BACKBONES.class_map:
|
||||
model = BACKBONES.build(module['cfg'],
|
||||
logger=self.logger).eval()
|
||||
elif module['cfg'].NAME in EMBEDDERS.class_map:
|
||||
model = EMBEDDERS.build(module['cfg'],
|
||||
logger=self.logger).eval()
|
||||
else:
|
||||
raise NotImplementedError
|
||||
if module['cfg'].get('RELOAD_MODEL', None):
|
||||
self.init_from_ckpt(module['cfg'].RELOAD_MODEL, model)
|
||||
module['model'] = model
|
||||
module['device'] = 'cpu'
|
||||
if module['device'] == 'cpu':
|
||||
module['device'] = we.device_id
|
||||
module['model'] = module['model'].to(we.device_id)
|
||||
return module
|
||||
|
||||
def unload(self, module):
|
||||
if module is None:
|
||||
return module
|
||||
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:
|
||||
module['model'] = module['model'].to('cpu')
|
||||
module['device'] = 'cpu'
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
return module
|
||||
|
||||
def dynamic_load(self, module=None, name=''):
|
||||
self.logger.info('Loading {} model'.format(name))
|
||||
if name == 'all':
|
||||
for subname in self.loaded_model_name:
|
||||
self.loaded_model[subname] = self.dynamic_load(
|
||||
getattr(self, subname), subname)
|
||||
elif name in self.loaded_model_name:
|
||||
if name in self.loaded_model:
|
||||
if module['cfg'] != self.loaded_model[name]['cfg']:
|
||||
self.unload(self.loaded_model[name])
|
||||
module = self.load(module)
|
||||
self.loaded_model[name] = module
|
||||
return module
|
||||
elif module['device'] == 'cpu':
|
||||
module = self.load(module)
|
||||
return module
|
||||
else:
|
||||
return module
|
||||
else:
|
||||
module = self.load(module)
|
||||
self.loaded_model[name] = module
|
||||
return module
|
||||
else:
|
||||
return self.load(module)
|
||||
|
||||
def dynamic_unload(self, module=None, name='', skip_loaded=False):
|
||||
self.logger.info('Unloading {} model'.format(name))
|
||||
if name == 'all':
|
||||
for name, module in self.loaded_model.items():
|
||||
module = self.unload(self.loaded_model[name])
|
||||
self.loaded_model[name] = module
|
||||
elif name in self.loaded_model_name:
|
||||
if name in self.loaded_model:
|
||||
if not skip_loaded:
|
||||
module = self.unload(self.loaded_model[name])
|
||||
self.loaded_model[name] = module
|
||||
else:
|
||||
self.unload(module)
|
||||
else:
|
||||
self.unload(module)
|
||||
|
||||
def load_default(self, cfg):
|
||||
module_paras = {}
|
||||
if cfg is not None:
|
||||
self.paras = cfg.PARAS
|
||||
self.input = {k.lower(): v for k, v in cfg.INPUT.items()}
|
||||
self.output = {k.lower(): v for k, v in cfg.OUTPUT.items()}
|
||||
module_paras = cfg.MODULES_PARAS
|
||||
return module_paras
|
||||
|
||||
def load_schedule(self, cfg):
|
||||
parameterization = cfg.get('PARAMETERIZATION', 'eps')
|
||||
assert parameterization in [
|
||||
'eps', 'x0', 'v'
|
||||
], 'currently only supporting "eps" and "x0" and "v"'
|
||||
num_timesteps = cfg.get('TIMESTEPS', 1000)
|
||||
|
||||
schedule_args = {
|
||||
k.lower(): v
|
||||
for k, v in cfg.get('SCHEDULE_ARGS', {
|
||||
'NAME': 'logsnr_cosine_interp',
|
||||
'SCALE_MIN': 2.0,
|
||||
'SCALE_MAX': 4.0
|
||||
}).items()
|
||||
}
|
||||
|
||||
zero_terminal_snr = cfg.get('ZERO_TERMINAL_SNR', False)
|
||||
if zero_terminal_snr:
|
||||
assert parameterization == 'v', 'Now zero_terminal_snr only support v-prediction mode.'
|
||||
sigmas = noise_schedule(schedule=schedule_args.pop('name'),
|
||||
n=num_timesteps,
|
||||
zero_terminal_snr=zero_terminal_snr,
|
||||
**schedule_args)
|
||||
diffusion = GaussianDiffusion(sigmas=sigmas,
|
||||
prediction_type=parameterization)
|
||||
return diffusion
|
||||
|
||||
def get_batch(self, value_dict, num_samples=1):
|
||||
batch = {}
|
||||
batch_uc = {}
|
||||
N = num_samples
|
||||
device = we.device_id
|
||||
for key in value_dict:
|
||||
if key == 'prompt':
|
||||
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 load_image(self, image, num_samples=1):
|
||||
if isinstance(image, torch.Tensor):
|
||||
pass
|
||||
elif isinstance(image, Image):
|
||||
pass
|
||||
elif isinstance(image, Image):
|
||||
pass
|
||||
|
||||
def get_function_info(self, module, function_name=None):
|
||||
all_function = module['function_info']
|
||||
if function_name in all_function:
|
||||
return function_name, all_function[function_name]['dtype']
|
||||
if function_name is None and len(all_function) == 1:
|
||||
for k, v in all_function.items():
|
||||
return k, v['dtype']
|
||||
|
||||
def encode_first_stage(self, x, **kwargs):
|
||||
_, dtype = self.get_function_info(self.first_stage_model, 'encode')
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
z = get_model(self.first_stage_model).encode(x)
|
||||
return self.first_stage_model['paras']['scale_factor'] * z
|
||||
|
||||
def decode_first_stage(self, z):
|
||||
_, dtype = self.get_function_info(self.first_stage_model, 'decode')
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
dtype=getattr(torch, dtype)):
|
||||
z = 1. / self.first_stage_model['paras']['scale_factor'] * z
|
||||
return get_model(self.first_stage_model).decode(z)
|
||||
|
||||
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,
|
||||
style_edit_image=None,
|
||||
style_exemplar_image=None,
|
||||
style_guide_scale_text=None,
|
||||
style_guide_scale_image=None,
|
||||
**kwargs):
|
||||
|
||||
value_input = copy.deepcopy(self.input)
|
||||
value_input.update(input)
|
||||
print(value_input)
|
||||
height, width = value_input['target_size_as_tuple']
|
||||
value_output = copy.deepcopy(self.output)
|
||||
batch, batch_uc = self.get_batch(value_input, num_samples=1)
|
||||
|
||||
# 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
|
||||
@@ -13,8 +13,9 @@ 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
|
||||
@@ -360,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:
|
||||
@@ -422,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
|
||||
@@ -986,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),
|
||||
|
||||
@@ -1,10 +1,7 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
from scepter.modules.model.embedder.embedder import (ConcatTimestepEmbedderND,
|
||||
FrozenCLIPEmbedder,
|
||||
FrozenOpenCLIPEmbedder,
|
||||
FrozenOpenCLIPEmbedder2,
|
||||
GeneralConditioner,
|
||||
IPAdapterPlusEmbedder,
|
||||
RefCrossEmbedder)
|
||||
from scepter.modules.model.embedder.embedder import (
|
||||
ConcatTimestepEmbedderND, FrozenCLIPEmbedder, FrozenOpenCLIPEmbedder,
|
||||
FrozenOpenCLIPEmbedder2, GeneralConditioner, IPAdapterPlusEmbedder,
|
||||
RefCrossEmbedder)
|
||||
|
||||
@@ -1,8 +1,5 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.head.classifier_head import (ClassifierHead,
|
||||
CosineLinearHead,
|
||||
TransformerHead,
|
||||
TransformerHeadx2,
|
||||
VideoClassifierHead,
|
||||
VideoClassifierHeadx2)
|
||||
from scepter.modules.model.head.classifier_head import (
|
||||
ClassifierHead, CosineLinearHead, TransformerHead, TransformerHeadx2,
|
||||
VideoClassifierHead, VideoClassifierHeadx2)
|
||||
|
||||
@@ -52,7 +52,6 @@ class DiagonalGaussianDistribution(object):
|
||||
dim=dims)
|
||||
|
||||
def mode(self):
|
||||
# print('*** use DiagonalGaussianDistribution.mode() ***')
|
||||
return self.mean
|
||||
|
||||
|
||||
|
||||
@@ -448,22 +448,23 @@ class GaussianDiffusion(object):
|
||||
# denoising
|
||||
t = self._sigma_to_t(sigma).repeat(len(xt)).round().long()
|
||||
|
||||
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']
|
||||
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
|
||||
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 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,
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -2,8 +2,6 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from inspect import isfunction
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def exists(x):
|
||||
return x is not None
|
||||
@@ -15,62 +13,6 @@ def default(val, d):
|
||||
return d() if isfunction(d) else d
|
||||
|
||||
|
||||
def checkpoint(func, inputs, params, flag):
|
||||
"""
|
||||
Evaluate a function without caching intermediate activations, allowing for
|
||||
reduced memory at the expense of extra compute in the backward pass.
|
||||
:param func: the function to evaluate.
|
||||
:param inputs: the argument sequence to pass to `func`.
|
||||
:param params: a sequence of parameters `func` depends on but does not
|
||||
explicitly take as arguments.
|
||||
:param flag: if False, disable gradient checkpointing.
|
||||
"""
|
||||
if flag:
|
||||
args = tuple(inputs) + tuple(params)
|
||||
return CheckpointFunction.apply(func, len(inputs), *args)
|
||||
else:
|
||||
return func(*inputs)
|
||||
|
||||
|
||||
class CheckpointFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, run_function, length, *args):
|
||||
ctx.run_function = run_function
|
||||
ctx.input_tensors = list(args[:length])
|
||||
ctx.input_params = list(args[length:])
|
||||
ctx.gpu_autocast_kwargs = {
|
||||
'enabled': torch.is_autocast_enabled(),
|
||||
'dtype': torch.get_autocast_gpu_dtype(),
|
||||
'cache_enabled': torch.is_autocast_cache_enabled()
|
||||
}
|
||||
with torch.no_grad():
|
||||
output_tensors = ctx.run_function(*ctx.input_tensors)
|
||||
return output_tensors
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, *output_grads):
|
||||
ctx.input_tensors = [
|
||||
x.detach().requires_grad_(True) for x in ctx.input_tensors
|
||||
]
|
||||
with torch.enable_grad(), \
|
||||
torch.cuda.amp.autocast(**ctx.gpu_autocast_kwargs):
|
||||
# Fixes a bug where the first op in run_function modifies the
|
||||
# Tensor storage in place, which is not allowed for detach()'d
|
||||
# Tensors.
|
||||
shallow_copies = [x.view_as(x) for x in ctx.input_tensors]
|
||||
output_tensors = ctx.run_function(*shallow_copies)
|
||||
input_grads = torch.autograd.grad(
|
||||
output_tensors,
|
||||
ctx.input_tensors + ctx.input_params,
|
||||
output_grads,
|
||||
allow_unused=True,
|
||||
)
|
||||
del ctx.input_tensors
|
||||
del ctx.input_params
|
||||
del output_tensors
|
||||
return (None, None) + input_grads
|
||||
|
||||
|
||||
def disabled_train(self, mode=True):
|
||||
"""Overwrite model.train with this function to make sure train/eval mode
|
||||
does not change anymore."""
|
||||
|
||||
@@ -1,10 +1,7 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
from scepter.modules.opt.optimizers.official_optimizers import (ASGD, LBFGS,
|
||||
SGD, Adadelta,
|
||||
Adagrad, Adam,
|
||||
Adamax, AdamW,
|
||||
RMSprop, Rprop,
|
||||
SparseAdam)
|
||||
from scepter.modules.opt.optimizers.official_optimizers import (
|
||||
ASGD, LBFGS, SGD, Adadelta, Adagrad, Adam, Adamax, AdamW, RMSprop, Rprop,
|
||||
SparseAdam)
|
||||
from scepter.modules.opt.optimizers.registry import OPTIMIZERS
|
||||
|
||||
@@ -10,16 +10,11 @@ import torchvision.transforms as transforms
|
||||
import torchvision.transforms.functional as TF
|
||||
|
||||
from scepter.modules.transform.registry import TRANSFORMS
|
||||
from scepter.modules.transform.utils import (BACKEND_CV2, BACKEND_PILLOW,
|
||||
BACKEND_TORCHVISION,
|
||||
INPUT_CV2_TYPE_WARNING,
|
||||
INPUT_PIL_TYPE_WARNING,
|
||||
INPUT_TENSOR_TYPE_WARNING,
|
||||
INTERPOLATION_STYLE,
|
||||
INTERPOLATION_STYLE_CV2,
|
||||
TORCHVISION_CAPABILITY,
|
||||
is_cv2_image, is_pil_image,
|
||||
is_tensor)
|
||||
from scepter.modules.transform.utils import (
|
||||
BACKEND_CV2, BACKEND_PILLOW, BACKEND_TORCHVISION, INPUT_CV2_TYPE_WARNING,
|
||||
INPUT_PIL_TYPE_WARNING, INPUT_TENSOR_TYPE_WARNING, INTERPOLATION_STYLE,
|
||||
INTERPOLATION_STYLE_CV2, TORCHVISION_CAPABILITY, is_cv2_image,
|
||||
is_pil_image, is_tensor)
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
|
||||
if TORCHVISION_CAPABILITY:
|
||||
|
||||
@@ -260,6 +260,7 @@ class RenameMeta(object):
|
||||
self.input_key = cfg.INPUT_KEY
|
||||
self.output_key = cfg.OUTPUT_KEY
|
||||
self.force = cfg.get('FORCE', False)
|
||||
self.move = cfg.get('MOVE', False)
|
||||
|
||||
def __call__(self, item):
|
||||
if 'meta' in item:
|
||||
@@ -270,10 +271,16 @@ class RenameMeta(object):
|
||||
have_key_set = set(self.input_key)
|
||||
else:
|
||||
have_key_set = set(self.input_key + self.output_key)
|
||||
for k, v in item['meta'].items():
|
||||
if k not in have_key_set:
|
||||
data[k] = v
|
||||
item['meta'] = data
|
||||
if not self.move:
|
||||
for k, v in item['meta'].items():
|
||||
if k not in have_key_set:
|
||||
data[k] = v
|
||||
item['meta'] = data
|
||||
else:
|
||||
for k, v in item.items():
|
||||
if k not in have_key_set:
|
||||
data[k] = v
|
||||
item.update(data)
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
|
||||
@@ -11,6 +11,13 @@ import yaml
|
||||
|
||||
from scepter.modules.utils.model import StdMsg
|
||||
|
||||
_SECURE_KEYWORDS = [
|
||||
'ENDPOINT', 'BUCKET', 'OSS_AK', 'OSS_SK', 'OSS', 'TOKEN', 'APPKEY'
|
||||
'SECRET', 'ACCESS_ID', 'ACCESS_KEY', 'PASSWORD', 'TEMP_DIR'
|
||||
] # -> "*****"
|
||||
|
||||
_SECURE_VALUEWORDS = ['oss://', 'oss-'] # -> "#####"
|
||||
|
||||
|
||||
def dict_to_yaml(module_name, name, json_config, set_name=False):
|
||||
'''
|
||||
@@ -517,8 +524,29 @@ class Config(object):
|
||||
def __repr__(self):
|
||||
return '{}\n'.format(self.dump())
|
||||
|
||||
def dump(self):
|
||||
return json.dumps(self.cfg_dict, indent=2)
|
||||
def dump(self, is_secure=False):
|
||||
if not is_secure:
|
||||
return json.dumps(self.cfg_dict, indent=2)
|
||||
else:
|
||||
|
||||
def make_secure(cfg):
|
||||
if isinstance(cfg, dict):
|
||||
for key, val in cfg.items():
|
||||
if key in _SECURE_KEYWORDS and type(val) is str:
|
||||
cfg[key] = '*****'
|
||||
else:
|
||||
cfg[key] = make_secure(cfg[key])
|
||||
elif isinstance(cfg, list):
|
||||
cfg = [make_secure(t) for t in cfg]
|
||||
elif isinstance(cfg, str):
|
||||
for sval in _SECURE_VALUEWORDS:
|
||||
if sval in cfg:
|
||||
cfg = '#####'
|
||||
return cfg
|
||||
|
||||
cfg_dict_copy = copy.deepcopy(self.cfg_dict)
|
||||
cfg_dict_copy = make_secure(cfg_dict_copy)
|
||||
return json.dumps(cfg_dict_copy, indent=2)
|
||||
|
||||
def deep_copy(self):
|
||||
return copy.deepcopy(self)
|
||||
@@ -604,5 +632,8 @@ class Config(object):
|
||||
else:
|
||||
return cfg
|
||||
|
||||
def __len__(self):
|
||||
return len(self.cfg_dict)
|
||||
|
||||
def pop(self, name):
|
||||
self.cfg_dict.pop(name)
|
||||
|
||||
@@ -131,7 +131,7 @@ class HttpFs(BaseFs):
|
||||
worker_id=0) -> Optional[str]:
|
||||
raise NotImplementedError
|
||||
|
||||
def get_url(self, target_path, lifecycle=3600 * 100):
|
||||
def get_url(self, target_path, set_public=False, lifecycle=3600 * 100):
|
||||
return target_path
|
||||
|
||||
def exists(self, target_path) -> bool:
|
||||
|
||||
@@ -181,7 +181,7 @@ class HuggingfaceFs(BaseFs):
|
||||
delimiter=None) -> (Union[bytes, str, None], Optional[int]):
|
||||
raise NotImplementedError
|
||||
|
||||
def get_url(self, target_path, lifecycle=3600 * 100):
|
||||
def get_url(self, target_path, set_public=False, lifecycle=3600 * 100):
|
||||
return target_path
|
||||
|
||||
def exists(self, target_path) -> bool:
|
||||
|
||||
@@ -199,7 +199,7 @@ class ModelscopeFs(BaseFs):
|
||||
delimiter=None) -> (Union[bytes, str, None], Optional[int]):
|
||||
raise NotImplementedError
|
||||
|
||||
def get_url(self, target_path, lifecycle=3600 * 100):
|
||||
def get_url(self, target_path, set_public=False, lifecycle=3600 * 100):
|
||||
return target_path
|
||||
|
||||
def exists(self, target_path) -> bool:
|
||||
|
||||
@@ -20,11 +20,13 @@ from scepter.studio.inference.inference_ui.largen_ui import LargenUI
|
||||
from scepter.studio.inference.inference_ui.mantra_ui import MantraUI
|
||||
from scepter.studio.inference.inference_ui.model_manage_ui import ModelManageUI
|
||||
from scepter.studio.inference.inference_ui.refiner_ui import RefinerUI
|
||||
from scepter.studio.inference.inference_ui.stylebooth_ui import StyleboothUI
|
||||
from scepter.studio.inference.inference_ui.tuner_ui import TunerUI
|
||||
from scepter.studio.utils.env import init_env
|
||||
|
||||
UI_MAP = [('diffusion', DiffusionUI), ('mantra', MantraUI), ('tuner', TunerUI),
|
||||
('control', ControlUI), ('refiner', RefinerUI), ('largen', LargenUI)]
|
||||
('control', ControlUI), ('refiner', RefinerUI), ('largen', LargenUI),
|
||||
('stylebooth', StyleboothUI)]
|
||||
|
||||
|
||||
class InferenceUI():
|
||||
@@ -108,14 +110,16 @@ class InferenceUI():
|
||||
self.tab_ui_kwargs[f'{name}_ui'] = ui
|
||||
self.__setattr__(f'{name}_ui', ui)
|
||||
|
||||
self.check_box_controlled_tabs = ['mantra', 'tuner', 'control', 'largen']
|
||||
self.check_box_controlled_tabs = [
|
||||
'mantra', 'tuner', 'control', 'largen', 'stylebooth'
|
||||
]
|
||||
self.pipe_manager = pipe_manager
|
||||
assert len(self.component_names.check_box_for_setting) == len(
|
||||
self.check_box_controlled_tabs)
|
||||
|
||||
def create_ui(self):
|
||||
# create model
|
||||
self.model_manage_ui.create_ui()
|
||||
self.model_manage_ui.create_ui(gallery_ui=self.gallery_ui)
|
||||
self.gallery_ui.create_ui()
|
||||
self.infer_info = gr.State(value=None)
|
||||
|
||||
@@ -123,10 +127,10 @@ class InferenceUI():
|
||||
def create_tab(name, ui):
|
||||
label = getattr(self.component_names, f'{name}_paras')
|
||||
if name in ['refiner']:
|
||||
ui.create_ui()
|
||||
ui.create_ui(gallery_ui=self.gallery_ui)
|
||||
else:
|
||||
with gr.TabItem(label=label, id=f'{name}_ui'):
|
||||
ui.create_ui()
|
||||
ui.create_ui(gallery_ui=self.gallery_ui)
|
||||
|
||||
with gr.Row(variant='panel', equal_height=True):
|
||||
with gr.Accordion(label=self.component_names.advance_block_name,
|
||||
@@ -140,8 +144,10 @@ class InferenceUI():
|
||||
|
||||
def set_callbacks(self, manager):
|
||||
self.model_manage_ui.set_callbacks(**self.tab_ui_kwargs)
|
||||
self.gallery_ui.set_callbacks(self, self.model_manage_ui,
|
||||
**self.tab_ui_kwargs)
|
||||
self.gallery_ui.set_callbacks(self,
|
||||
self.model_manage_ui,
|
||||
**self.tab_ui_kwargs,
|
||||
manager=manager)
|
||||
for name, ui in self.tab_ui_kwargs.items():
|
||||
ui.set_callbacks(self.model_manage_ui,
|
||||
**self.tab_ui_kwargs,
|
||||
@@ -152,8 +158,14 @@ class InferenceUI():
|
||||
selected_tab = 'diffusion_ui'
|
||||
ui_tabs_state = [False] * len(args)
|
||||
largen_index = self.check_box_controlled_tabs.index('largen')
|
||||
largen_key = self.component_names.check_box_for_setting[largen_index]
|
||||
largen_key = self.component_names.check_box_for_setting[
|
||||
largen_index]
|
||||
largen_status = args[largen_index]
|
||||
stylebooth_index = self.check_box_controlled_tabs.index(
|
||||
'stylebooth')
|
||||
stylebooth_key = self.component_names.check_box_for_setting[
|
||||
stylebooth_index]
|
||||
stylebooth_status = args[stylebooth_index]
|
||||
for key in check_box:
|
||||
i = self.component_names.check_box_for_setting.index(key)
|
||||
ui_tabs_state[i] = True
|
||||
@@ -163,16 +175,17 @@ class InferenceUI():
|
||||
for key in check_box:
|
||||
i = self.component_names.check_box_for_setting.index(key)
|
||||
if ui_tabs_state[i] != args[i]:
|
||||
if i in [largen_index]:
|
||||
if i in [largen_index, stylebooth_index]:
|
||||
new_check_box_value = [key]
|
||||
for j in range(len(ui_tabs_state)):
|
||||
ui_tabs_state[j] = j == i
|
||||
else:
|
||||
new_check_box_value = [
|
||||
k for k in check_box
|
||||
if k not in [largen_key]
|
||||
if k not in [largen_key, stylebooth_key]
|
||||
]
|
||||
ui_tabs_state[largen_index] = False
|
||||
ui_tabs_state[stylebooth_index] = False
|
||||
|
||||
ui_tabs_updates = [gr.update(visible=v) for v in ui_tabs_state]
|
||||
|
||||
@@ -183,7 +196,14 @@ class InferenceUI():
|
||||
default_choices['diffusion_model']['choices'],
|
||||
value='LARGEN_LargenUNetXL',
|
||||
interactive=False)
|
||||
elif largen_status:
|
||||
elif ui_tabs_state[stylebooth_index]:
|
||||
diffusion_model = gr.Dropdown(
|
||||
label=self.model_manage_ui.component_names.diffusion_model,
|
||||
choices=self.model_manage_ui.
|
||||
default_choices['diffusion_model']['choices'],
|
||||
value='EDIT_DiffusionUNet',
|
||||
interactive=False)
|
||||
elif largen_status or stylebooth_status:
|
||||
diffusion_model = gr.Dropdown(
|
||||
label=self.model_manage_ui.component_names.diffusion_model,
|
||||
choices=self.model_manage_ui.
|
||||
@@ -203,7 +223,8 @@ class InferenceUI():
|
||||
choices=self.component_names.check_box_for_setting,
|
||||
value=new_check_box_value,
|
||||
show_label=False), gr.update(
|
||||
selected=selected_tab), *ui_tabs_state, *ui_tabs_updates, diffusion_model
|
||||
selected=selected_tab
|
||||
), *ui_tabs_state, *ui_tabs_updates, diffusion_model
|
||||
|
||||
gr_states = [
|
||||
self.tab_ui[name].state for name in self.check_box_controlled_tabs
|
||||
@@ -213,9 +234,14 @@ class InferenceUI():
|
||||
]
|
||||
self.check_box_for_setting.change(
|
||||
change_setting_tab,
|
||||
inputs=[self.check_box_for_setting, self.model_manage_ui.diffusion_model, *gr_states],
|
||||
outputs=[self.check_box_for_setting, self.setting_tab, *gr_states,
|
||||
*gr_tabs, self.model_manage_ui.diffusion_model],
|
||||
inputs=[
|
||||
self.check_box_for_setting,
|
||||
self.model_manage_ui.diffusion_model, *gr_states
|
||||
],
|
||||
outputs=[
|
||||
self.check_box_for_setting, self.setting_tab, *gr_states,
|
||||
*gr_tabs, self.model_manage_ui.diffusion_model
|
||||
],
|
||||
queue=False)
|
||||
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.inference.diffusion_inference import DiffusionInference
|
||||
from scepter.modules.inference.largen_inference import LargenInference
|
||||
from scepter.modules.inference.stylebooth_inference import StyleboothInference
|
||||
from scepter.modules.utils.logger import get_logger
|
||||
|
||||
|
||||
@@ -99,6 +100,8 @@ class PipelineManager():
|
||||
pipeline_name = cfg.NAME
|
||||
if 'LARGEN' in pipeline_name:
|
||||
PipelineBuilder = LargenInference
|
||||
elif pipeline_name.startswith('EDIT'):
|
||||
PipelineBuilder = StyleboothInference
|
||||
else:
|
||||
PipelineBuilder = DiffusionInference
|
||||
new_inference = PipelineBuilder(logger=self.logger)
|
||||
|
||||
@@ -24,7 +24,8 @@ class InferenceUIName():
|
||||
if language == 'en':
|
||||
self.advance_block_name = 'Advance Setting'
|
||||
self.check_box_for_setting = [
|
||||
'Use Mantra', 'Use Tuners', 'Use Controller', 'LAR-Gen'
|
||||
'Use Mantra', 'Use Tuners', 'Use Controller', 'LAR-Gen',
|
||||
'StyleBooth'
|
||||
]
|
||||
self.diffusion_paras = 'Generation Setting'
|
||||
self.mantra_paras = 'Mantra Book'
|
||||
@@ -32,15 +33,19 @@ class InferenceUIName():
|
||||
self.control_paras = 'Controlable Generation'
|
||||
self.refiner_paras = 'Refiner Setting'
|
||||
self.largen_paras = 'LAR-Gen'
|
||||
self.stylebooth_paras = 'StyleBooth'
|
||||
elif language == 'zh':
|
||||
self.advance_block_name = '生成选项'
|
||||
self.check_box_for_setting = ['使用咒语', '使用微调', '使用控制', 'LAR-Gen']
|
||||
self.check_box_for_setting = [
|
||||
'使用咒语', '使用微调', '使用控制', 'LAR-Gen', 'StyleBooth'
|
||||
]
|
||||
self.diffusion_paras = '生成参数设置'
|
||||
self.mantra_paras = '咒语书'
|
||||
self.tuner_paras = '微调模型'
|
||||
self.control_paras = '可控生成'
|
||||
self.refiner_paras = 'Refine设置'
|
||||
self.largen_paras = 'LAR-Gen'
|
||||
self.stylebooth_paras = 'StyleBooth'
|
||||
|
||||
|
||||
class ModelManageUIName():
|
||||
@@ -499,3 +504,82 @@ class LargenUIName():
|
||||
1024
|
||||
],
|
||||
]
|
||||
|
||||
|
||||
class StyleboothUIName():
|
||||
def __init__(self, language='en'):
|
||||
if language == 'en':
|
||||
self.dropdown_name = 'Application'
|
||||
self.apps = ['Text-based Style Editing']
|
||||
# self.apps = ["Text-based Style Editing", "Exemplar-based Style Editing"]
|
||||
self.source_image = 'Source Image'
|
||||
self.exemplar_image = 'Exemplar Image'
|
||||
self.ins_format = 'Instruction Format (select or rewrite, en only)'
|
||||
self.style_format = '{} (select or rewrite, en only)'
|
||||
self.guide_scale_image = 'Guide Scale For Uncondition Image'
|
||||
self.guide_scale_text = 'Guide Scale For Uncondition Text'
|
||||
self.guide_rescale = 'Guide Rescale'
|
||||
self.resolution = 'Resolution of Short Edge'
|
||||
self.compose_button = 'Assemble Style Editing Instruction to Prompt'
|
||||
elif language == 'zh':
|
||||
self.dropdown_name = '应用'
|
||||
self.apps = ['根据文本编辑风格']
|
||||
# self.apps = ["根据文本编辑风格", "根据风格样例编辑风格"]
|
||||
self.source_image = '源图片'
|
||||
self.exemplar_image = '样例图片'
|
||||
self.ins_format = '指令模版 (选择或者新写,仅英文)'
|
||||
self.style_format = '{} (选择或者新写,仅英文)'
|
||||
self.guide_scale_image = '图片条件引导比例'
|
||||
self.guide_scale_text = '文本条件引导比例'
|
||||
self.guide_rescale = '引导缩放'
|
||||
self.resolution = '短边分辨率'
|
||||
self.compose_button = '组装风格编辑指令到Prompt栏'
|
||||
self.tb_ins_format_choice = [
|
||||
'Let this image be in the style of <style>',
|
||||
'Please edit this image to embody the characteristics of <style> style.',
|
||||
'Transform this image to reflect the distinct aesthetic of <style>.',
|
||||
'Can you infuse this image with the signature techniques representative of <style>?',
|
||||
'Adjust the visual elements of this image to emulate the <style> style.',
|
||||
'Reinterpret this image through the artistic lens of <style>.',
|
||||
'Apply the <style> style to this image to capture its unique essence.',
|
||||
'Modify this photograph to mirror the thematic qualities of <style>.',
|
||||
"I'd like you to rework this image to pay homage to the <style> movement.",
|
||||
'Ensure that this image adopts the brushwork and color palette typical of <style>.',
|
||||
'Give this image a makeover so that it aligns with the <style> stylistic approach.',
|
||||
'Retouch this image to channel the spirit and technique of <style>.',
|
||||
'Merge this image with the foundational elements of <style>.',
|
||||
'Re-envision this image to fit within the <style> genre.',
|
||||
'Adapt this image to exhibit the soft edges and vibrant light of <style>.',
|
||||
'Craft this image to resonate with the visual themes found in <style>.',
|
||||
]
|
||||
self.eb_ins_format_choice = [
|
||||
'Let this image be in the style of <image>',
|
||||
'Please match the aesthetic of this image to that of <image>.',
|
||||
'Adjust the current image to mimic the visual style of <image>.',
|
||||
'Edit this photo so that it reflects the artistic style found in <image>.',
|
||||
'Transform this picture to be stylistically similar to <image>.',
|
||||
'Recreate the ambiance and look of <image> in this one.',
|
||||
'Harmonize the visual elements of this image with those in <image>.',
|
||||
'Ensure the editing of this image captures the essence of the style in <image>.',
|
||||
'Can you make this image resonate with the artistic flair of <image>?',
|
||||
"I'd like this image to have the same feel and tone as the style reference in <image>.",
|
||||
'Retouch this image to align with the creative direction of <image>.',
|
||||
'Replicate the style characteristics of <image> onto this one.',
|
||||
'Adapt the visual theme of this image to be consistent with <image>.',
|
||||
"Conform this image's aesthetic to the distinctive look of <image>.",
|
||||
'Make over this image so it conforms to the stylistic cues of <image>.',
|
||||
"Modify this image's style to echo the artistic qualities of <image>.",
|
||||
]
|
||||
self.tb_target_style = [
|
||||
'Adorable 3D Character', 'Color Field Painting',
|
||||
'Colored Pencil Art', 'Graffiti Art', 'futuristic-retro futurism',
|
||||
'game-retro arcade', 'Simple Vector Art', 'Sketchup', 'mre-comic',
|
||||
'sai-comic book', 'sai-lowpoly', 'misc-stained glass',
|
||||
'misc-zentangle', 'papercraft-papercut collage', 'Pop Art 2',
|
||||
'photo-silhouette', 'Adorable Kawaii', 'sai-anime', 'sai-origami',
|
||||
'futuristic-vaporwave', 'game-retro game', 'misc-disco',
|
||||
'papercraft-flat papercut'
|
||||
]
|
||||
self.eb_target_style = ['<image>']
|
||||
self.tb_identifier = '<style>'
|
||||
self.eb_identifier = '<image>'
|
||||
|
||||
@@ -118,6 +118,14 @@ class ControlUI(UIBase):
|
||||
|
||||
self.example_block = gr.Accordion(
|
||||
label=self.component_names.example_block_name, open=True)
|
||||
gallery_ui = kwargs.pop('gallery_ui', None)
|
||||
gallery_ui.register_components({
|
||||
'control_state': self.state,
|
||||
'control_model': self.control_model,
|
||||
'control_scale': self.control_scale,
|
||||
'crop_type': self.crop_type,
|
||||
'control_cond_image': self.cond_image,
|
||||
})
|
||||
|
||||
def set_callbacks(self, model_manage_ui, diffusion_ui, **kwargs):
|
||||
gallery_ui = kwargs.pop('gallery_ui')
|
||||
|
||||
@@ -153,6 +153,20 @@ class DiffusionUI(UIBase):
|
||||
interactive=True)
|
||||
with gr.Column(scale=1):
|
||||
self.refresh_seed = gr.Button(value=refresh_symbol)
|
||||
gallery_ui = kwargs.pop('gallery_ui', None)
|
||||
gallery_ui.register_components({
|
||||
'negative_prompt': self.negative_prompt,
|
||||
'prompt_prefix': self.prompt_prefix,
|
||||
'sample': self.sampler,
|
||||
'discretization': self.discretization,
|
||||
'output_height': self.output_height,
|
||||
'output_width': self.output_width,
|
||||
'image_number': self.image_number,
|
||||
'sample_steps': self.sample_steps,
|
||||
'guide_scale': self.guide_scale,
|
||||
'guide_rescale': self.guide_rescale,
|
||||
'image_seed': self.image_seed,
|
||||
})
|
||||
|
||||
def set_callbacks(self, model_manage_ui, **kwargs):
|
||||
gallery_ui = kwargs.pop('gallery_ui')
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
|
||||
import gradio as gr
|
||||
import numpy as np
|
||||
@@ -19,6 +20,8 @@ class GalleryUI(UIBase):
|
||||
self.work_dir = cfg.WORK_DIR
|
||||
self.local_work_dir, _ = FS.map_to_local(self.work_dir)
|
||||
os.makedirs(self.local_work_dir, exist_ok=True)
|
||||
self.component_mapping = OrderedDict()
|
||||
self.manager = None
|
||||
|
||||
def create_ui(self, *args, **kwargs):
|
||||
with gr.Group():
|
||||
@@ -55,164 +58,157 @@ class GalleryUI(UIBase):
|
||||
elem_classes='type_row',
|
||||
elem_id='generate_button',
|
||||
visible=True)
|
||||
self.register_components({'prompt': self.prompt})
|
||||
|
||||
def generate_gallery(self,
|
||||
prompt,
|
||||
mantra_state,
|
||||
tuner_state,
|
||||
control_state,
|
||||
refine_state,
|
||||
diffusion_model,
|
||||
first_stage_model,
|
||||
cond_stage_model,
|
||||
refiner_cond_model,
|
||||
refiner_diffusion_model,
|
||||
tuner_model,
|
||||
tuner_scale,
|
||||
custom_tuner_model,
|
||||
control_model,
|
||||
control_scale,
|
||||
crop_type,
|
||||
control_cond_image,
|
||||
negative_prompt,
|
||||
prompt_prefix,
|
||||
sample,
|
||||
discretization,
|
||||
output_height,
|
||||
output_width,
|
||||
image_number,
|
||||
sample_steps,
|
||||
guide_scale,
|
||||
guide_rescale,
|
||||
refine_strength,
|
||||
refine_sampler,
|
||||
refine_discretization,
|
||||
refine_guide_scale,
|
||||
refine_guide_rescale,
|
||||
style_template,
|
||||
style_negative_template,
|
||||
image_seed,
|
||||
largen_state,
|
||||
largen_task,
|
||||
largen_image_scale,
|
||||
largen_tar_image,
|
||||
largen_tar_mask,
|
||||
largen_masked_image,
|
||||
largen_ref_image,
|
||||
largen_ref_mask,
|
||||
largen_ref_clip,
|
||||
largen_base_image,
|
||||
largen_extra_sizes,
|
||||
largen_bbox_yyxx,
|
||||
largen_history,
|
||||
show_jpeg_image=True):
|
||||
if control_state and control_cond_image is None:
|
||||
raise gr.Error(self.component_names.control_err1)
|
||||
def register_components(self, components):
|
||||
common_keys = self.component_mapping.keys() & components.keys()
|
||||
assert len(
|
||||
common_keys) == 0, f'Component key already exist: {common_keys}'
|
||||
self.component_mapping.update(components)
|
||||
|
||||
current_pipeline = self.pipe_manager.get_pipeline_given_modules({
|
||||
'diffusion_model':
|
||||
diffusion_model,
|
||||
'first_stage_model':
|
||||
first_stage_model,
|
||||
'cond_stage_model':
|
||||
cond_stage_model,
|
||||
'refiner_cond_model':
|
||||
refiner_cond_model,
|
||||
'refiner_diffusion_model':
|
||||
refiner_diffusion_model
|
||||
})
|
||||
now_pipeline = self.pipe_manager.model_level_info[diffusion_model][
|
||||
'pipeline'][0]
|
||||
used_tuner_model = []
|
||||
if not isinstance(tuner_model, list):
|
||||
tuner_model = [tuner_model]
|
||||
for tuner_m in tuner_model:
|
||||
if tuner_m is None or tuner_m == '':
|
||||
continue
|
||||
if now_pipeline in self.pipe_manager.model_level_info['tuners'] and \
|
||||
tuner_m in self.pipe_manager.model_level_info['tuners'][now_pipeline]:
|
||||
tuner_m = self.pipe_manager.model_level_info['tuners'][
|
||||
now_pipeline][tuner_m]['model_info']
|
||||
used_tuner_model.append(tuner_m)
|
||||
used_custom_tuner_model = []
|
||||
if not isinstance(custom_tuner_model, list):
|
||||
custom_tuner_model = [custom_tuner_model]
|
||||
for tuner_m in custom_tuner_model:
|
||||
if tuner_m is None or tuner_m == '':
|
||||
continue
|
||||
def generate_gallery(self, *args, show_jpeg_image=True):
|
||||
# unload all image preprocess model.
|
||||
if self.manager is not None:
|
||||
if (hasattr(self.manager, 'preprocess') and hasattr(
|
||||
self.manager.preprocess.dataset_gallery.processors_manager,
|
||||
'dynamic_unload')):
|
||||
self.manager.preprocess.dataset_gallery.processors_manager.dynamic_unload(
|
||||
)
|
||||
|
||||
def pipeline_init(args):
|
||||
return self.pipe_manager.get_pipeline_given_modules(args)
|
||||
|
||||
def control_init(args):
|
||||
control_state, tuner_state = args['control_state'], args[
|
||||
'tuner_state']
|
||||
control_cond_image = args.pop('control_cond_image')
|
||||
control_model = args.pop('control_model')
|
||||
|
||||
if control_state and control_cond_image is None:
|
||||
raise gr.Error(self.component_names.control_err1)
|
||||
|
||||
now_pipeline = self.pipe_manager.model_level_info[
|
||||
args['diffusion_model']]['pipeline'][0]
|
||||
if (now_pipeline
|
||||
in self.pipe_manager.model_level_info['customized_tuners']
|
||||
and tuner_m in self.pipe_manager.
|
||||
model_level_info['customized_tuners'][now_pipeline]):
|
||||
tuner_m = self.pipe_manager.model_level_info[
|
||||
'customized_tuners'][now_pipeline][tuner_m]['model_info']
|
||||
used_custom_tuner_model.append(tuner_m)
|
||||
in self.pipe_manager.model_level_info['controllers']
|
||||
and control_model in self.pipe_manager.
|
||||
model_level_info['controllers'][now_pipeline]):
|
||||
control_model = self.pipe_manager.model_level_info[
|
||||
'controllers'][now_pipeline][control_model]['model_info']
|
||||
|
||||
if (now_pipeline in self.pipe_manager.model_level_info['controllers']
|
||||
and control_model in self.pipe_manager.
|
||||
model_level_info['controllers'][now_pipeline]):
|
||||
control_model = self.pipe_manager.model_level_info['controllers'][
|
||||
now_pipeline][control_model]['model_info']
|
||||
args.update({
|
||||
'control_model':
|
||||
control_model if control_state else None,
|
||||
'control_scale':
|
||||
args.pop('control_scale')
|
||||
if tuner_state or control_state else None,
|
||||
'control_cond_image':
|
||||
control_cond_image if control_state else None,
|
||||
'crop_type':
|
||||
args.pop('crop_type') if control_state else None
|
||||
})
|
||||
|
||||
prompt_rephrased = style_template.replace(
|
||||
'{prompt}',
|
||||
prompt) if not style_template == '' and mantra_state else prompt
|
||||
prompt_rephrased = f'{prompt_prefix}{prompt_rephrased}' if not prompt_prefix == '' else prompt_rephrased
|
||||
negative_prompt_rephrased = negative_prompt + style_negative_template if mantra_state else negative_prompt
|
||||
pipeline_input = {
|
||||
'prompt': prompt_rephrased,
|
||||
'negative_prompt': negative_prompt_rephrased,
|
||||
'sample': sample,
|
||||
'sample_steps': sample_steps,
|
||||
'discretization': discretization,
|
||||
'original_size_as_tuple': [int(output_height),
|
||||
int(output_width)],
|
||||
'target_size_as_tuple': [int(output_height),
|
||||
int(output_width)],
|
||||
'crop_coords_top_left': [0, 0],
|
||||
'guide_scale': guide_scale,
|
||||
'guide_rescale': guide_rescale,
|
||||
}
|
||||
if refine_state:
|
||||
pipeline_input['refine_sampler'] = refine_sampler
|
||||
pipeline_input['refine_discretization'] = refine_discretization
|
||||
pipeline_input['refine_guide_scale'] = refine_guide_scale
|
||||
pipeline_input['refine_guide_rescale'] = refine_guide_rescale
|
||||
else:
|
||||
refine_strength = 0
|
||||
if largen_state:
|
||||
largen_cfg = {
|
||||
'largen_task': largen_task,
|
||||
'largen_image_scale': largen_image_scale,
|
||||
'largen_tar_image': largen_tar_image,
|
||||
'largen_tar_mask': largen_tar_mask,
|
||||
'largen_ref_image': largen_ref_image,
|
||||
'largen_ref_mask': largen_ref_mask,
|
||||
'largen_masked_image': largen_masked_image,
|
||||
'largen_ref_clip': largen_ref_clip,
|
||||
'largen_base_image': largen_base_image,
|
||||
'largen_extra_sizes': largen_extra_sizes,
|
||||
'largen_bbox_yyxx': largen_bbox_yyxx,
|
||||
def tuner_init(args):
|
||||
control_state, tuner_state = args['control_state'], args[
|
||||
'tuner_state']
|
||||
tuner_model = args.pop('tuner_model')
|
||||
custom_tuner_model = args.pop('custom_tuner_model')
|
||||
|
||||
now_pipeline = self.pipe_manager.model_level_info[
|
||||
args['diffusion_model']]['pipeline'][0]
|
||||
used_tuner_model = []
|
||||
if not isinstance(tuner_model, list):
|
||||
tuner_model = [tuner_model]
|
||||
for tuner_m in tuner_model:
|
||||
if tuner_m is None or tuner_m == '':
|
||||
continue
|
||||
if now_pipeline in self.pipe_manager.model_level_info['tuners'] and \
|
||||
tuner_m in self.pipe_manager.model_level_info['tuners'][now_pipeline]:
|
||||
tuner_m = self.pipe_manager.model_level_info['tuners'][
|
||||
now_pipeline][tuner_m]['model_info']
|
||||
used_tuner_model.append(tuner_m)
|
||||
used_custom_tuner_model = []
|
||||
if not isinstance(custom_tuner_model, list):
|
||||
custom_tuner_model = [custom_tuner_model]
|
||||
for tuner_m in custom_tuner_model:
|
||||
if tuner_m is None or tuner_m == '':
|
||||
continue
|
||||
if (now_pipeline in
|
||||
self.pipe_manager.model_level_info['customized_tuners']
|
||||
and tuner_m in self.pipe_manager.
|
||||
model_level_info['customized_tuners'][now_pipeline]):
|
||||
tuner_m = self.pipe_manager.model_level_info[
|
||||
'customized_tuners'][now_pipeline][tuner_m][
|
||||
'model_info']
|
||||
used_custom_tuner_model.append(tuner_m)
|
||||
args.update({
|
||||
'tuner_model':
|
||||
used_tuner_model +
|
||||
used_custom_tuner_model if tuner_state else None,
|
||||
'tuner_scale':
|
||||
args.pop('tuner_scale')
|
||||
if tuner_state or control_state else None,
|
||||
})
|
||||
|
||||
def input_init(args):
|
||||
mantra_state = args['mantra_state']
|
||||
prompt = args.pop('prompt')
|
||||
negative_prompt = args.pop('negative_prompt')
|
||||
prompt_prefix = args.pop('prompt_prefix')
|
||||
style_template = args.pop('style_template')
|
||||
size = [
|
||||
int(args.pop('output_height')),
|
||||
int(args.pop('output_width'))
|
||||
]
|
||||
|
||||
prompt_rephrased = style_template.replace(
|
||||
'{prompt}', prompt
|
||||
) if not style_template == '' and mantra_state else prompt
|
||||
|
||||
pipeline_input = {
|
||||
'prompt':
|
||||
f'{prompt_prefix}{prompt_rephrased}'
|
||||
if not prompt_prefix == '' else prompt_rephrased,
|
||||
'negative_prompt':
|
||||
negative_prompt + args.pop('style_negative_template')
|
||||
if mantra_state else negative_prompt,
|
||||
'sample':
|
||||
args.pop('sample'),
|
||||
'sample_steps':
|
||||
args.pop('sample_steps'),
|
||||
'discretization':
|
||||
args.pop('discretization'),
|
||||
'original_size_as_tuple':
|
||||
size,
|
||||
'target_size_as_tuple':
|
||||
size,
|
||||
'crop_coords_top_left': [0, 0],
|
||||
'guide_scale':
|
||||
args.pop('guide_scale'),
|
||||
'guide_rescale':
|
||||
args.pop('guide_rescale')
|
||||
}
|
||||
else:
|
||||
largen_cfg = {}
|
||||
args.update({'input': pipeline_input})
|
||||
|
||||
def appedix_init(args):
|
||||
args.update({
|
||||
'num_samples': args.pop('image_number'),
|
||||
'intermediate_callback': None,
|
||||
'img_to_img_strength': 0,
|
||||
'seed': int(args.pop('image_seed')),
|
||||
})
|
||||
|
||||
args = dict(zip(self.component_mapping.keys(), args))
|
||||
largen_history = args.pop('largen_history')
|
||||
|
||||
current_pipeline = pipeline_init(args)
|
||||
control_init(args)
|
||||
tuner_init(args)
|
||||
input_init(args)
|
||||
appedix_init(args)
|
||||
|
||||
results = current_pipeline(**args)
|
||||
|
||||
results = current_pipeline(
|
||||
pipeline_input,
|
||||
num_samples=image_number,
|
||||
intermediate_callback=None,
|
||||
refine_strength=refine_strength,
|
||||
img_to_img_strength=0,
|
||||
tuner_model=used_tuner_model +
|
||||
used_custom_tuner_model if tuner_state else None,
|
||||
tuner_scale=tuner_scale if tuner_state or control_state else None,
|
||||
control_model=control_model if control_state else None,
|
||||
control_scale=control_scale
|
||||
if tuner_state or control_state else None,
|
||||
control_cond_image=control_cond_image if control_state else None,
|
||||
crop_type=crop_type if control_state else None,
|
||||
seed=int(image_seed),
|
||||
**largen_cfg)
|
||||
images = []
|
||||
before_images = []
|
||||
if 'images' in results:
|
||||
@@ -233,9 +229,12 @@ class GalleryUI(UIBase):
|
||||
if 'seed' in results:
|
||||
print(results['seed'])
|
||||
print(images, before_images)
|
||||
largen_history.extend(images)
|
||||
if len(largen_history) > 10:
|
||||
largen_history = largen_history[-10:]
|
||||
|
||||
if args['largen_state']:
|
||||
largen_history.extend(images)
|
||||
if len(largen_history) > 5:
|
||||
largen_history = largen_history[-5:]
|
||||
|
||||
if show_jpeg_image:
|
||||
save_list = []
|
||||
for i, img in enumerate(images):
|
||||
@@ -258,35 +257,20 @@ class GalleryUI(UIBase):
|
||||
before_refine_panel, before_refine_gallery, output_gallery, _ = gallery_result
|
||||
return (before_refine_panel, before_refine_gallery, output_gallery[0])
|
||||
|
||||
def set_callbacks(self, inference_ui, model_manage_ui, diffusion_ui,
|
||||
mantra_ui, tuner_ui, refiner_ui, control_ui, largen_ui,
|
||||
def set_callbacks(self,
|
||||
inference_ui,
|
||||
model_manage_ui,
|
||||
diffusion_ui,
|
||||
mantra_ui,
|
||||
tuner_ui,
|
||||
refiner_ui,
|
||||
control_ui,
|
||||
largen_ui,
|
||||
manager=None,
|
||||
**kwargs):
|
||||
|
||||
self.gen_inputs = [
|
||||
self.prompt, mantra_ui.state, tuner_ui.state, control_ui.state,
|
||||
refiner_ui.state, model_manage_ui.diffusion_model,
|
||||
model_manage_ui.first_stage_model,
|
||||
model_manage_ui.cond_stage_model, refiner_ui.refiner_cond_model,
|
||||
refiner_ui.refiner_diffusion_model, tuner_ui.tuner_model,
|
||||
tuner_ui.tuner_scale, tuner_ui.custom_tuner_model,
|
||||
control_ui.control_model, control_ui.control_scale,
|
||||
control_ui.crop_type, control_ui.cond_image,
|
||||
diffusion_ui.negative_prompt, diffusion_ui.prompt_prefix,
|
||||
diffusion_ui.sampler, diffusion_ui.discretization,
|
||||
diffusion_ui.output_height, diffusion_ui.output_width,
|
||||
diffusion_ui.image_number, diffusion_ui.sample_steps,
|
||||
diffusion_ui.guide_scale, diffusion_ui.guide_rescale,
|
||||
refiner_ui.refine_strength, refiner_ui.refine_sampler,
|
||||
refiner_ui.refine_discretization, refiner_ui.refine_guide_scale,
|
||||
refiner_ui.refine_guide_rescale, mantra_ui.style_template,
|
||||
mantra_ui.style_negative_template, diffusion_ui.image_seed,
|
||||
largen_ui.state, largen_ui.task, largen_ui.image_scale,
|
||||
largen_ui.tar_image, largen_ui.tar_mask, largen_ui.masked_image,
|
||||
largen_ui.ref_image, largen_ui.ref_mask, largen_ui.ref_clip,
|
||||
largen_ui.base_image, largen_ui.extra_sizes, largen_ui.bbox_yyxx,
|
||||
largen_ui.image_history
|
||||
]
|
||||
|
||||
self.manager = manager
|
||||
self.gen_inputs = list(self.component_mapping.values())
|
||||
print(self.gen_inputs, len(self.gen_inputs))
|
||||
self.gen_outputs = [
|
||||
self.before_refine_panel,
|
||||
self.before_refine_gallery,
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import albumentations as A
|
||||
import cv2
|
||||
import gradio as gr
|
||||
import numpy as np
|
||||
@@ -10,6 +9,7 @@ import torchvision.transforms as T
|
||||
import torchvision.transforms.functional as TF
|
||||
from PIL import Image
|
||||
|
||||
import albumentations as A
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
from scepter.modules.model.utils.data_utils import (box2squre, expand_bbox,
|
||||
get_bbox_from_mask,
|
||||
@@ -61,10 +61,10 @@ class LargenUI(UIBase):
|
||||
with gr.Column(visible=False) as self.tab:
|
||||
with gr.Row():
|
||||
self.select_app = gr.Dropdown(
|
||||
label=self.component_names.dropdown_name,
|
||||
choices=self.component_names.apps,
|
||||
value=self.component_names.apps[0],
|
||||
type='index')
|
||||
label=self.component_names.dropdown_name,
|
||||
choices=self.component_names.apps,
|
||||
value=self.component_names.apps[0],
|
||||
type='index')
|
||||
with gr.Row(equal_height=True):
|
||||
with gr.Column(scale=2, min_width=0):
|
||||
with gr.Row():
|
||||
@@ -76,7 +76,8 @@ class LargenUI(UIBase):
|
||||
source='upload',
|
||||
height=400,
|
||||
interactive=True)
|
||||
self.cache_button = gr.Button(value='Use Last Generated Image', visible=True)
|
||||
self.cache_button = gr.Button(
|
||||
value='Use Last Generated Image', visible=True)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.subject_image = gr.Image(
|
||||
label=self.component_names.subject_image,
|
||||
@@ -87,8 +88,13 @@ class LargenUI(UIBase):
|
||||
height=400,
|
||||
visible=False)
|
||||
|
||||
self.gallery = gr.Gallery(label='Image History', value=[], columns=1, rows=1, height=500)
|
||||
self.clear_button = gr.Button(value='Clear History', visible=True)
|
||||
self.gallery = gr.Gallery(label='Image History',
|
||||
value=[],
|
||||
columns=1,
|
||||
rows=1,
|
||||
height=500)
|
||||
self.clear_button = gr.Button(value='Clear History',
|
||||
visible=True)
|
||||
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.image_scale = gr.Slider(label='Image Strength',
|
||||
@@ -101,21 +107,40 @@ class LargenUI(UIBase):
|
||||
value=0.75,
|
||||
label='Image Resize Ratio',
|
||||
visible=True)
|
||||
self.out_direction = gr.Dropdown(label=self.component_names.out_direction_label,
|
||||
choices=self.component_names.out_directions,
|
||||
value=self.component_names.out_directions[0],
|
||||
visible=True)
|
||||
self.out_direction = gr.Dropdown(
|
||||
label=self.component_names.out_direction_label,
|
||||
choices=self.component_names.out_directions,
|
||||
value=self.component_names.out_directions[0],
|
||||
visible=True)
|
||||
|
||||
self.proc_button = gr.Button(value=self.component_names.button_name)
|
||||
self.proc_button = gr.Button(
|
||||
value=self.component_names.button_name)
|
||||
self.proc_status = gr.Markdown(value='', visible=False)
|
||||
self.task_desc = gr.Markdown(self.component_names.direction, visible=True)
|
||||
self.task_desc = gr.Markdown(
|
||||
self.component_names.direction, visible=True)
|
||||
|
||||
self.eg = gr.Column(visible=True)
|
||||
gallery_ui = kwargs.pop('gallery_ui', None)
|
||||
gallery_ui.register_components({
|
||||
'largen_state': self.state,
|
||||
'largen_task': self.task,
|
||||
'largen_image_scale': self.image_scale,
|
||||
'largen_tar_image': self.tar_image,
|
||||
'largen_tar_mask': self.tar_mask,
|
||||
'largen_masked_image': self.masked_image,
|
||||
'largen_ref_image': self.ref_image,
|
||||
'largen_ref_mask': self.ref_mask,
|
||||
'largen_ref_clip': self.ref_clip,
|
||||
'largen_base_image': self.base_image,
|
||||
'largen_extra_sizes': self.extra_sizes,
|
||||
'largen_bbox_yyxx': self.bbox_yyxx,
|
||||
'largen_history': self.image_history
|
||||
})
|
||||
|
||||
def set_callbacks(self, model_manage_ui, diffusion_ui, **kwargs):
|
||||
|
||||
def example_data_process(select_app_id, prompt, scene_image, scene_mask, subject_image,
|
||||
subject_mask, image_scale, image_ratio, out_direction,
|
||||
def example_data_process(select_app_id, prompt, scene_image,
|
||||
scene_mask, subject_image, subject_mask,
|
||||
image_scale, image_ratio, out_direction,
|
||||
output_height, output_width):
|
||||
task = self.component_names.tasks[select_app_id]
|
||||
if scene_mask is not None:
|
||||
@@ -124,29 +149,20 @@ class LargenUI(UIBase):
|
||||
subject_mask = (subject_mask > 128).astype(np.uint8)
|
||||
|
||||
if task == 'Text_Guided_Inpainting':
|
||||
data = self.data_preprocess_inpaint(scene_image,
|
||||
scene_mask,
|
||||
None,
|
||||
None,
|
||||
False,
|
||||
1.3,
|
||||
output_height, output_width)
|
||||
data = self.data_preprocess_inpaint(scene_image, scene_mask,
|
||||
None, None, False, 1.3,
|
||||
output_height,
|
||||
output_width)
|
||||
elif task == 'Subject_Guided_Inpainting':
|
||||
data = self.data_preprocess_inpaint(scene_image,
|
||||
scene_mask,
|
||||
data = self.data_preprocess_inpaint(scene_image, scene_mask,
|
||||
subject_image,
|
||||
subject_mask,
|
||||
False,
|
||||
1.3,
|
||||
subject_mask, False, 1.3,
|
||||
output_height,
|
||||
output_width)
|
||||
elif task == 'Text_Subject_Guided_Inpainting':
|
||||
data = self.data_preprocess_inpaint(scene_image,
|
||||
scene_mask,
|
||||
data = self.data_preprocess_inpaint(scene_image, scene_mask,
|
||||
subject_image,
|
||||
subject_mask,
|
||||
True,
|
||||
1.3,
|
||||
subject_mask, True, 1.3,
|
||||
output_height,
|
||||
output_width)
|
||||
elif task == 'Text_Guided_Outpainting':
|
||||
@@ -156,7 +172,8 @@ class LargenUI(UIBase):
|
||||
output_height,
|
||||
output_width)
|
||||
|
||||
subject_image_show = None if subject_image is None else Image.fromarray(subject_image.astype(np.uint8))
|
||||
subject_image_show = None if subject_image is None else Image.fromarray(
|
||||
subject_image.astype(np.uint8))
|
||||
return *data, gr.update(value='Data Process Succeed!', visible=True), \
|
||||
gr.update(value=Image.fromarray(scene_image.astype(np.uint8))), \
|
||||
gr.update(value=subject_image_show), \
|
||||
@@ -165,43 +182,62 @@ class LargenUI(UIBase):
|
||||
|
||||
gallery_ui = kwargs.pop('gallery_ui')
|
||||
with self.eg:
|
||||
self.scene_image_eg = gr.Image(label=self.component_names.scene_image, type='numpy', visible=False) # noqa
|
||||
self.scene_mask_eg = gr.Image(label=self.component_names.scene_mask, type='numpy', image_mode='L', visible=False) # noqa
|
||||
self.subject_image_eg = gr.Image(label=self.component_names.subject_image, type='numpy', visible=False) # noqa
|
||||
self.subject_mask_eg = gr.Image(label=self.component_names.subject_mask, type='numpy', image_mode='L', visible=False) # noqa
|
||||
self.prompt = gr.Textbox(label=self.component_names.prompt, visible=False)
|
||||
self.examples = gr.Examples(
|
||||
examples=self.component_names.examples,
|
||||
inputs=[self.select_app, self.prompt,
|
||||
self.scene_image_eg, self.scene_mask_eg,
|
||||
self.subject_image_eg, self.subject_mask_eg,
|
||||
self.image_scale,
|
||||
self.image_ratio,
|
||||
self.out_direction,
|
||||
diffusion_ui.output_height,
|
||||
diffusion_ui.output_width,
|
||||
],
|
||||
outputs=[self.tar_image,
|
||||
self.tar_mask,
|
||||
self.masked_image,
|
||||
self.ref_image,
|
||||
self.ref_mask,
|
||||
self.ref_clip,
|
||||
self.base_image,
|
||||
self.extra_sizes,
|
||||
self.bbox_yyxx,
|
||||
self.proc_status,
|
||||
self.scene_image,
|
||||
self.subject_image,
|
||||
self.task,
|
||||
self.select_app,
|
||||
gallery_ui.prompt,
|
||||
self.image_scale,
|
||||
self.image_ratio,
|
||||
],
|
||||
fn=example_data_process,
|
||||
cache_examples=False,
|
||||
run_on_click=True)
|
||||
self.scene_image_eg = gr.Image(
|
||||
label=self.component_names.scene_image,
|
||||
type='numpy',
|
||||
visible=False) # noqa
|
||||
self.scene_mask_eg = gr.Image(
|
||||
label=self.component_names.scene_mask,
|
||||
type='numpy',
|
||||
image_mode='L',
|
||||
visible=False) # noqa
|
||||
self.subject_image_eg = gr.Image(
|
||||
label=self.component_names.subject_image,
|
||||
type='numpy',
|
||||
visible=False) # noqa
|
||||
self.subject_mask_eg = gr.Image(
|
||||
label=self.component_names.subject_mask,
|
||||
type='numpy',
|
||||
image_mode='L',
|
||||
visible=False) # noqa
|
||||
self.prompt = gr.Textbox(label=self.component_names.prompt,
|
||||
visible=False)
|
||||
self.examples = gr.Examples(examples=self.component_names.examples,
|
||||
inputs=[
|
||||
self.select_app,
|
||||
self.prompt,
|
||||
self.scene_image_eg,
|
||||
self.scene_mask_eg,
|
||||
self.subject_image_eg,
|
||||
self.subject_mask_eg,
|
||||
self.image_scale,
|
||||
self.image_ratio,
|
||||
self.out_direction,
|
||||
diffusion_ui.output_height,
|
||||
diffusion_ui.output_width,
|
||||
],
|
||||
outputs=[
|
||||
self.tar_image,
|
||||
self.tar_mask,
|
||||
self.masked_image,
|
||||
self.ref_image,
|
||||
self.ref_mask,
|
||||
self.ref_clip,
|
||||
self.base_image,
|
||||
self.extra_sizes,
|
||||
self.bbox_yyxx,
|
||||
self.proc_status,
|
||||
self.scene_image,
|
||||
self.subject_image,
|
||||
self.task,
|
||||
self.select_app,
|
||||
gallery_ui.prompt,
|
||||
self.image_scale,
|
||||
self.image_ratio,
|
||||
],
|
||||
fn=example_data_process,
|
||||
cache_examples=False,
|
||||
run_on_click=True)
|
||||
|
||||
def change_app(select_app_id):
|
||||
select_task = self.component_names.tasks[select_app_id]
|
||||
@@ -210,18 +246,14 @@ class LargenUI(UIBase):
|
||||
gr.update(visible=('Outpainting' in select_task)), \
|
||||
gr.update(visible=('Outpainting' in select_task)), select_task
|
||||
|
||||
self.select_app.change(
|
||||
change_app,
|
||||
inputs=[self.select_app],
|
||||
outputs=[
|
||||
self.subject_image,
|
||||
self.image_scale,
|
||||
self.image_ratio,
|
||||
self.out_direction,
|
||||
self.task
|
||||
],
|
||||
queue=False
|
||||
)
|
||||
self.select_app.change(change_app,
|
||||
inputs=[self.select_app],
|
||||
outputs=[
|
||||
self.subject_image, self.image_scale,
|
||||
self.image_ratio, self.out_direction,
|
||||
self.task
|
||||
],
|
||||
queue=False)
|
||||
|
||||
def read_gallery_image(gallery):
|
||||
if len(gallery) == 0:
|
||||
@@ -239,13 +271,12 @@ class LargenUI(UIBase):
|
||||
gallery.clear()
|
||||
return image_history, gallery
|
||||
|
||||
self.clear_button.click(
|
||||
fn=clear_gallery,
|
||||
inputs=[self.image_history, self.gallery],
|
||||
outputs=[self.image_history, self.gallery]
|
||||
)
|
||||
self.clear_button.click(fn=clear_gallery,
|
||||
inputs=[self.image_history, self.gallery],
|
||||
outputs=[self.image_history, self.gallery])
|
||||
|
||||
def data_process(scene_image, subject_image, task, image_ratio, out_direction, output_height, output_width):
|
||||
def data_process(scene_image, subject_image, task, image_ratio,
|
||||
out_direction, output_height, output_width):
|
||||
tar_image = scene_image['image'].convert('RGB')
|
||||
tar_mask = scene_image['mask'].convert('L')
|
||||
tar_image = np.asarray(tar_image)
|
||||
@@ -253,12 +284,8 @@ class LargenUI(UIBase):
|
||||
tar_mask = np.where(tar_mask > 128, 1, 0).astype(np.uint8)
|
||||
|
||||
if task == 'Text_Guided_Inpainting':
|
||||
data = self.data_preprocess_inpaint(tar_image,
|
||||
tar_mask,
|
||||
None,
|
||||
None,
|
||||
False,
|
||||
1.3,
|
||||
data = self.data_preprocess_inpaint(tar_image, tar_mask, None,
|
||||
None, False, 1.3,
|
||||
output_height,
|
||||
output_width)
|
||||
elif task == 'Subject_Guided_Inpainting':
|
||||
@@ -267,13 +294,9 @@ class LargenUI(UIBase):
|
||||
ref_image = np.asarray(ref_image)
|
||||
ref_mask = np.asarray(ref_mask)
|
||||
ref_mask = np.where(ref_mask > 128, 1, 0).astype(np.uint8)
|
||||
data = self.data_preprocess_inpaint(tar_image,
|
||||
tar_mask,
|
||||
ref_image,
|
||||
ref_mask,
|
||||
False,
|
||||
1.3,
|
||||
output_height,
|
||||
data = self.data_preprocess_inpaint(tar_image, tar_mask,
|
||||
ref_image, ref_mask, False,
|
||||
1.3, output_height,
|
||||
output_width)
|
||||
elif task == 'Text_Subject_Guided_Inpainting':
|
||||
ref_image = subject_image['image'].convert('RGB')
|
||||
@@ -281,29 +304,23 @@ class LargenUI(UIBase):
|
||||
ref_image = np.asarray(ref_image)
|
||||
ref_mask = np.asarray(ref_mask)
|
||||
ref_mask = np.where(ref_mask > 128, 1, 0).astype(np.uint8)
|
||||
data = self.data_preprocess_inpaint(tar_image,
|
||||
tar_mask,
|
||||
ref_image,
|
||||
ref_mask,
|
||||
True,
|
||||
1.3,
|
||||
output_height,
|
||||
data = self.data_preprocess_inpaint(tar_image, tar_mask,
|
||||
ref_image, ref_mask, True,
|
||||
1.3, output_height,
|
||||
output_width)
|
||||
elif task == 'Text_Guided_Outpainting':
|
||||
data = self.data_preprocess_outpaint(tar_image,
|
||||
out_direction,
|
||||
data = self.data_preprocess_outpaint(tar_image, out_direction,
|
||||
image_ratio,
|
||||
output_height,
|
||||
output_width)
|
||||
|
||||
return *data, gr.update(value='Data Process Succeed!', visible=True)
|
||||
return *data, gr.update(value='Data Process Succeed!',
|
||||
visible=True)
|
||||
|
||||
self.proc_button.click(data_process,
|
||||
inputs=[
|
||||
self.scene_image,
|
||||
self.subject_image,
|
||||
self.task,
|
||||
self.image_ratio,
|
||||
self.scene_image, self.subject_image,
|
||||
self.task, self.image_ratio,
|
||||
self.out_direction,
|
||||
diffusion_ui.output_height,
|
||||
diffusion_ui.output_width
|
||||
@@ -319,17 +336,11 @@ class LargenUI(UIBase):
|
||||
self.extra_sizes,
|
||||
self.bbox_yyxx,
|
||||
self.proc_status,
|
||||
])
|
||||
])
|
||||
|
||||
def data_preprocess_inpaint(self,
|
||||
tar_image,
|
||||
tar_mask,
|
||||
ref_image,
|
||||
ref_mask,
|
||||
use_rectangle_mask,
|
||||
tar_crop_ratio,
|
||||
output_height,
|
||||
output_width):
|
||||
def data_preprocess_inpaint(self, tar_image, tar_mask, ref_image, ref_mask,
|
||||
use_rectangle_mask, tar_crop_ratio,
|
||||
output_height, output_width):
|
||||
tar_mask = np.expand_dims(tar_mask, 2).astype(np.float32)
|
||||
|
||||
# Zoom-In
|
||||
@@ -346,15 +357,20 @@ class LargenUI(UIBase):
|
||||
y1, y2, x1, x2 = tar_bbox_yyxx
|
||||
crop_tar_mask[y1:y2, x1:x2] = 1
|
||||
|
||||
crop_tar_image, pad1, pad2 = pad_to_square(crop_tar_image.astype(np.uint8), pad_value=0)
|
||||
crop_tar_image, pad1, pad2 = pad_to_square(crop_tar_image.astype(
|
||||
np.uint8),
|
||||
pad_value=0)
|
||||
crop_tar_mask, _, _ = pad_to_square(crop_tar_mask, pad_value=0)
|
||||
H2, W2 = crop_tar_image.shape[:2]
|
||||
|
||||
aug_tar_image = cv2.resize(crop_tar_image.astype(np.uint8), (output_width, output_height))
|
||||
aug_tar_image = cv2.resize(crop_tar_image.astype(np.uint8),
|
||||
(output_width, output_height))
|
||||
aug_tar_mask = cv2.resize(crop_tar_mask, (output_width, output_height))
|
||||
|
||||
final_tar_image = TF.to_tensor(aug_tar_image)
|
||||
final_tar_image = TF.normalize(final_tar_image, mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
|
||||
final_tar_image = TF.normalize(final_tar_image,
|
||||
mean=[0.5, 0.5, 0.5],
|
||||
std=[0.5, 0.5, 0.5])
|
||||
final_tar_mask = TF.to_tensor((aug_tar_mask > 0.5).astype(np.float32))
|
||||
|
||||
masked_image = final_tar_image.clone()
|
||||
@@ -367,7 +383,8 @@ class LargenUI(UIBase):
|
||||
if ref_image is not None and ref_mask is not None:
|
||||
ref_mask = np.expand_dims(ref_mask, 2).astype(np.float32)
|
||||
# background-free
|
||||
ref_image = ref_image * ref_mask + np.ones_like(ref_image) * 255. * (1 - ref_mask)
|
||||
ref_image = ref_image * ref_mask + np.ones_like(
|
||||
ref_image) * 255. * (1 - ref_mask)
|
||||
|
||||
ref_yyxx = get_bbox_from_mask(ref_mask)
|
||||
y1, y2, x1, x2 = ref_yyxx
|
||||
@@ -377,32 +394,40 @@ class LargenUI(UIBase):
|
||||
|
||||
h, w = crop_ref_mask_i.shape[:2]
|
||||
ref_expand_size = int(max(h, w) * 1.02)
|
||||
pad_op = A.PadIfNeeded(ref_expand_size, ref_expand_size,
|
||||
pad_op = A.PadIfNeeded(ref_expand_size,
|
||||
ref_expand_size,
|
||||
border_mode=cv2.BORDER_CONSTANT,
|
||||
value=(255, 255, 255), mask_value=0)
|
||||
value=(255, 255, 255),
|
||||
mask_value=0)
|
||||
out = pad_op(image=crop_ref_image_i, mask=crop_ref_mask_i)
|
||||
crop_ref_image = out['image']
|
||||
|
||||
to_clip_input = T.Compose([
|
||||
T.ToTensor(),
|
||||
T.Resize((224, 224)),
|
||||
T.Normalize(mean=(0.48145466, 0.4578275, 0.40821073), std=(0.26862954, 0.26130258, 0.27577711)),
|
||||
T.Normalize(mean=(0.48145466, 0.4578275, 0.40821073),
|
||||
std=(0.26862954, 0.26130258, 0.27577711)),
|
||||
])
|
||||
ref_clip = to_clip_input(crop_ref_image.astype(np.uint8))
|
||||
|
||||
output_size = max(output_height, output_width)
|
||||
ref_resize_op = A.Compose([
|
||||
A.LongestMaxSize(output_size),
|
||||
A.PadIfNeeded(output_size, output_size,
|
||||
A.PadIfNeeded(output_size,
|
||||
output_size,
|
||||
border_mode=cv2.BORDER_CONSTANT,
|
||||
value=(255, 255, 255), mask_value=0),
|
||||
value=(255, 255, 255),
|
||||
mask_value=0),
|
||||
])
|
||||
aug_out = ref_resize_op(image=crop_ref_image_i.astype(np.uint8), mask=crop_ref_mask_i)
|
||||
aug_out = ref_resize_op(image=crop_ref_image_i.astype(np.uint8),
|
||||
mask=crop_ref_mask_i)
|
||||
aug_ref_image = aug_out['image']
|
||||
aug_ref_mask = aug_out['mask']
|
||||
|
||||
final_ref_image = TF.to_tensor(aug_ref_image)
|
||||
final_ref_image = TF.normalize(final_ref_image, mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
|
||||
final_ref_image = TF.normalize(final_ref_image,
|
||||
mean=[0.5, 0.5, 0.5],
|
||||
std=[0.5, 0.5, 0.5])
|
||||
final_ref_mask = TF.to_tensor(aug_ref_mask)
|
||||
|
||||
final_ref_image = final_ref_image.unsqueeze(0)
|
||||
@@ -416,15 +441,11 @@ class LargenUI(UIBase):
|
||||
return final_tar_image, final_tar_mask, masked_image, final_ref_image, final_ref_mask, ref_clip, \
|
||||
TF.to_tensor(tar_image), torch.LongTensor([H1, W1, H2, W2, pad1, pad2]), torch.LongTensor(tar_yyxx_crop)
|
||||
|
||||
def data_preprocess_outpaint(self,
|
||||
tar_image,
|
||||
direction,
|
||||
img_ratio,
|
||||
output_height,
|
||||
output_width):
|
||||
def data_preprocess_outpaint(self, tar_image, direction, img_ratio,
|
||||
output_height, output_width):
|
||||
oh, ow = output_height, output_width
|
||||
h, w = tar_image.shape[:2]
|
||||
ratio = max(h/(oh*img_ratio), w/(ow*img_ratio))
|
||||
ratio = max(h / (oh * img_ratio), w / (ow * img_ratio))
|
||||
|
||||
ih, iw = int(h / ratio), int(w / ratio)
|
||||
|
||||
@@ -432,25 +453,27 @@ class LargenUI(UIBase):
|
||||
mask = np.zeros((oh, ow, 1))
|
||||
|
||||
if direction in ['CenterAround', '中心向外']:
|
||||
y1, x1 = (oh-ih)//2, (ow-iw)//2
|
||||
y1, x1 = (oh - ih) // 2, (ow - iw) // 2
|
||||
elif direction in ['RightDown', '右下']:
|
||||
y1, x1 = 0, 0
|
||||
elif direction in ['LeftDown', '左下']:
|
||||
y1, x1 = 0, ow-iw
|
||||
y1, x1 = 0, ow - iw
|
||||
elif direction in ['RightUp', '右上']:
|
||||
y1, x1 = oh-ih, 0
|
||||
y1, x1 = oh - ih, 0
|
||||
elif direction in ['LeftUp', '左上']:
|
||||
y1, x1 = oh-ih, ow-iw
|
||||
y1, x1 = oh - ih, ow - iw
|
||||
else:
|
||||
y1, x1 = 0, 0
|
||||
|
||||
tar_image = cv2.resize(tar_image.astype(np.uint8), (iw, ih))
|
||||
masked_image[y1:y1+ih, x1:x1+iw] = tar_image
|
||||
mask[y1+5:y1+ih-5, x1+5:x1+iw-5] = 1
|
||||
masked_image[y1:y1 + ih, x1:x1 + iw] = tar_image
|
||||
mask[y1 + 5:y1 + ih - 5, x1 + 5:x1 + iw - 5] = 1
|
||||
|
||||
final_tar_image = TF.to_tensor(masked_image.astype(np.uint8))
|
||||
final_tar_image = TF.normalize(final_tar_image, mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
|
||||
final_tar_mask = TF.to_tensor(((1.0-mask) > 0.5).astype(np.float32))
|
||||
final_tar_image = TF.normalize(final_tar_image,
|
||||
mean=[0.5, 0.5, 0.5],
|
||||
std=[0.5, 0.5, 0.5])
|
||||
final_tar_mask = TF.to_tensor(((1.0 - mask) > 0.5).astype(np.float32))
|
||||
|
||||
masked_image = final_tar_image.clone()
|
||||
masked_image = masked_image * (1 - final_tar_mask)
|
||||
|
||||
@@ -111,6 +111,16 @@ class MantraUI(UIBase):
|
||||
self.example_block = gr.Accordion(
|
||||
label=self.component_names.example_block_name, open=True)
|
||||
|
||||
gallery_ui = kwargs.pop('gallery_ui', None)
|
||||
gallery_ui.register_components({
|
||||
'mantra_state':
|
||||
self.state,
|
||||
'style_template':
|
||||
self.style_template,
|
||||
'style_negative_template':
|
||||
self.style_negative_template,
|
||||
})
|
||||
|
||||
def set_callbacks(self, model_manage_ui, **kwargs):
|
||||
gallery_ui = kwargs.pop('gallery_ui')
|
||||
with self.example_block:
|
||||
|
||||
@@ -101,6 +101,15 @@ class ModelManageUI(UIBase):
|
||||
# open=False):
|
||||
# self.tuner_name = gr.Text(
|
||||
# label='tuner_name')
|
||||
gallery_ui = kwargs.pop('gallery_ui', None)
|
||||
gallery_ui.register_components({
|
||||
'diffusion_model':
|
||||
self.diffusion_model,
|
||||
'first_stage_model':
|
||||
self.first_stage_model,
|
||||
'cond_stage_model':
|
||||
self.cond_stage_model,
|
||||
})
|
||||
|
||||
def set_callbacks(self, diffusion_ui, tuner_ui, control_ui, mantra_ui,
|
||||
**kwargs):
|
||||
|
||||
@@ -0,0 +1,174 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import gradio as gr
|
||||
|
||||
from scepter.studio.inference.inference_ui.component_names import \
|
||||
StyleboothUIName
|
||||
from scepter.studio.utils.uibase import UIBase
|
||||
|
||||
|
||||
class StyleboothUI(UIBase):
|
||||
def __init__(self, cfg, pipe_manager, is_debug=False, language='en'):
|
||||
self.cfg = cfg
|
||||
self.pipe_manager = pipe_manager
|
||||
self.component_names = StyleboothUIName(language)
|
||||
|
||||
def create_ui(self, *args, **kwargs):
|
||||
self.state = gr.State(value=False)
|
||||
with gr.Column(equal_height=False, visible=False) as self.tab:
|
||||
with gr.Row():
|
||||
self.selected_app = gr.Dropdown(
|
||||
label=self.component_names.dropdown_name,
|
||||
choices=self.component_names.apps,
|
||||
value=self.component_names.apps[0],
|
||||
type='index',
|
||||
interactive=False)
|
||||
with gr.Row(equal_height=True):
|
||||
with gr.Column(variant='panel',
|
||||
scale=1,
|
||||
min_width=0,
|
||||
visible=True) as self.col1:
|
||||
self.edit_image = gr.Image(
|
||||
label=self.component_names.source_image,
|
||||
type='pil',
|
||||
tool='editor',
|
||||
interactive=True)
|
||||
with gr.Column(variant='panel',
|
||||
scale=1,
|
||||
min_width=0,
|
||||
visible=True) as self.col2:
|
||||
with gr.Group(visible=True):
|
||||
self.exemplar_image = gr.Image(
|
||||
label=self.component_names.exemplar_image,
|
||||
type='pil',
|
||||
interactive=True,
|
||||
visible=False)
|
||||
with gr.Row(equal_height=True):
|
||||
with gr.Group(visible=True):
|
||||
self.instruction_format = gr.Dropdown(
|
||||
label=self.component_names.ins_format,
|
||||
choices=self.component_names.
|
||||
tb_ins_format_choice,
|
||||
value=self.component_names.
|
||||
tb_ins_format_choice[0],
|
||||
multiselect=False,
|
||||
interactive=True,
|
||||
allow_custom_value=True)
|
||||
self.target_style = gr.Dropdown(
|
||||
label=self.component_names.style_format.
|
||||
format(self.component_names.tb_identifier),
|
||||
choices=self.component_names.
|
||||
tb_target_style,
|
||||
value=None,
|
||||
multiselect=False,
|
||||
interactive=True,
|
||||
allow_custom_value=True)
|
||||
self.compose_instruction = gr.Button(
|
||||
label=self.component_names.compose_button,
|
||||
value=self.component_names.compose_button,
|
||||
elem_classes='type_row',
|
||||
elem_id='push',
|
||||
visible=True)
|
||||
with gr.Column(variant='panel', scale=1, min_width=0):
|
||||
with gr.Group(visible=True):
|
||||
self.guide_scale_text = gr.Slider(
|
||||
label=self.component_names.guide_scale_text,
|
||||
minimum=1,
|
||||
maximum=10,
|
||||
step=0.5,
|
||||
value=7.5,
|
||||
interactive=True)
|
||||
self.guide_scale_image = gr.Slider(
|
||||
label=self.component_names.guide_scale_image,
|
||||
minimum=1,
|
||||
maximum=10,
|
||||
step=0.5,
|
||||
value=1.5,
|
||||
interactive=True)
|
||||
self.guide_rescale = gr.Slider(
|
||||
label=self.component_names.guide_rescale,
|
||||
minimum=0,
|
||||
maximum=1.0,
|
||||
step=0.1,
|
||||
value=0.5,
|
||||
interactive=True)
|
||||
# self.resolution = gr.Slider(
|
||||
# label=self.component_names.resolution,
|
||||
# minimum=384, maximum=768, step=32, value=512,
|
||||
# interactive=True)
|
||||
# gallery_ui = kwargs.pop("gallery_ui")
|
||||
# with gr.Column(visible=True):
|
||||
# with gr.Column(visible=True) as self.general_examples:
|
||||
# gr.Examples(
|
||||
# examples=self.component_names.general_examples,
|
||||
# inputs=[gallery_ui.prompt, self.source_image],
|
||||
# outputs=[gallery_ui.prompt, self.source_image],
|
||||
# fn=lambda x, y: (x, y),
|
||||
# cache_examples=False)
|
||||
# with gr.Column(visible=False) as self.tuner_examples:
|
||||
# gr.Examples(
|
||||
# examples=self.component_names.tuner_examples,
|
||||
# inputs=[self.editor_model, gallery_ui.prompt, self.source_image],
|
||||
# outputs=[self.editor_model, gallery_ui.prompt, self.source_image],
|
||||
# fn=lambda x, y, z: (x, y, z),
|
||||
# cache_examples=True)
|
||||
gallery_ui = kwargs.pop('gallery_ui', None)
|
||||
gallery_ui.register_components({
|
||||
'style_edit_image':
|
||||
self.edit_image,
|
||||
'style_exemplar_image':
|
||||
self.exemplar_image,
|
||||
'style_guide_scale_text':
|
||||
self.guide_scale_text,
|
||||
'style_guide_scale_image':
|
||||
self.guide_scale_image,
|
||||
})
|
||||
|
||||
def set_callbacks(self, model_manage_ui, diffusion_ui, gallery_ui,
|
||||
**kwargs):
|
||||
def app_change(app):
|
||||
selected = app
|
||||
# selected = None
|
||||
# for i, name in enumerate(self.component_names.apps):
|
||||
# if name == app:
|
||||
# selected = i
|
||||
# continue
|
||||
# assert selected in (0, 1)
|
||||
if not selected:
|
||||
format_choices = self.component_names.tb_ins_format_choice
|
||||
style_choices = self.component_names.tb_target_style
|
||||
else:
|
||||
format_choices = self.component_names.eb_ins_format_choice
|
||||
style_choices = self.component_names.eb_target_style
|
||||
|
||||
identifier = self.component_names.tb_identifier if selected == 0 else self.component_names.eb_identifier
|
||||
return (gr.update(value=None, visible=(selected != 0)),
|
||||
gr.update(choices=format_choices, value=format_choices[0]),
|
||||
gr.update(choices=style_choices,
|
||||
value=style_choices[0],
|
||||
interactive=(selected == 0),
|
||||
label=self.component_names.style_format.format(
|
||||
identifier)))
|
||||
|
||||
self.selected_app.change(app_change,
|
||||
inputs=[self.selected_app],
|
||||
outputs=[
|
||||
self.exemplar_image,
|
||||
self.instruction_format, self.target_style
|
||||
],
|
||||
queue=False)
|
||||
|
||||
def compose_instruction(format, style, app):
|
||||
if app == 1:
|
||||
return format
|
||||
identifier = self.component_names.tb_identifier if app == 0 else self.component_names.eb_identifier
|
||||
return format.replace(identifier, style)
|
||||
|
||||
self.compose_instruction.click(compose_instruction,
|
||||
inputs=[
|
||||
self.instruction_format,
|
||||
self.target_style, self.selected_app
|
||||
],
|
||||
outputs=[gallery_ui.prompt],
|
||||
queue=False)
|
||||
@@ -118,6 +118,17 @@ class TunerUI(UIBase):
|
||||
|
||||
self.example_block = gr.Accordion(
|
||||
label=self.component_names.example_block_name, open=True)
|
||||
gallery_ui = kwargs.pop('gallery_ui', None)
|
||||
gallery_ui.register_components({
|
||||
'tuner_state':
|
||||
self.state,
|
||||
'tuner_model':
|
||||
self.tuner_model,
|
||||
'tuner_scale':
|
||||
self.tuner_scale,
|
||||
'custom_tuner_model':
|
||||
self.custom_tuner_model,
|
||||
})
|
||||
|
||||
def set_callbacks(self, model_manage_ui, **kwargs):
|
||||
manager = kwargs.pop('manager')
|
||||
|
||||
@@ -4,32 +4,55 @@
|
||||
class CreateDatasetUIName():
|
||||
def __init__(self, language='en'):
|
||||
if language == 'en':
|
||||
self.dataset_name = 'All Dataset'
|
||||
self.btn_create_datasets = 'Create Dataset'
|
||||
self.user_data_name = 'Current Dataset Name'
|
||||
self.modify_data_button = 'Modify Name'
|
||||
self.confirm_data_button = 'Confirm'
|
||||
self.refresh_list_button = 'Refresh List'
|
||||
self.system_log = '<span style="color: blue;">System Log: {}</span> '
|
||||
self.btn_create_datasets = '\U00002795' # ➕
|
||||
self.get_data_name_button = '\U0001F3B2' # 🎲
|
||||
self.new_data_name = (
|
||||
'New dataset name, replace "name" and "version" with easy-to-remember identifiers.'
|
||||
f'Also get a random name by clicking {self.get_data_name_button}'
|
||||
)
|
||||
self.modify_data_button = '\U0001F4DD' # 📝
|
||||
self.confirm_data_button = '\U00002714' # ✔️
|
||||
self.cancel_create_button = '\U00002716' # ✖️
|
||||
self.refresh_list_button = '\U0001f504' # 🔄
|
||||
self.delete_dataset_button = '\U0001f5d1' # 🗑️
|
||||
self.dataset_name = (
|
||||
f'All Dataset,click{self.btn_create_datasets}to create new dataset,'
|
||||
f'click{self.delete_dataset_button}to delete this dataset.')
|
||||
self.dataset_type = 'Dataset Type'
|
||||
self.dataset_type_name = {
|
||||
'scepter_txt2img': 'Text2Image Generation',
|
||||
'scepter_img2img': 'Image Edit Generation'
|
||||
}
|
||||
self.user_data_name = (
|
||||
f'Current Dataset Name. Changes of dataset name take '
|
||||
f'effect after clicking {self.modify_data_button}')
|
||||
self.zip_file = 'Upload Dataset(Zip/Txt)'
|
||||
self.zip_file_url = 'Dataset Url'
|
||||
self.default_dataset_repo = 'https://www.modelscope.cn/api/v1/models/iic/scepter/'
|
||||
self.default_dataset_zip = \
|
||||
self.default_dataset_repo + 'repo?Revision=master&FilePath=datasets/3D_example_csv.zip'
|
||||
self.default_dataset_name = '3D_example'
|
||||
self.btn_create_datasets_from_file = 'Create Dataset From File'
|
||||
self.user_direction = (
|
||||
'### User Guide: \n' +
|
||||
f'* {self.btn_create_datasets} button is used to create a new dataset '
|
||||
"from scratch. Please make sure to modify the dataset's name and version. After creation, "
|
||||
". Please make sure to modify the dataset's name and version. After creation, "
|
||||
'you can upload images one by one. \n'
|
||||
f'* The {self.btn_create_datasets_from_file} button supports creating a new dataset from '
|
||||
f'* The "{self.btn_create_datasets_from_file}" button supports creating a new dataset from '
|
||||
'a file, currently supporting zip files. For zip files, the format should be consistent'
|
||||
" with the one used during training, ensuring it contains an 'images/' folder and a '"
|
||||
"train.csv' (which will use the image paths in this file); "
|
||||
'The first line is Target:FILE, Prompt, followed by the format of each line: image path, description.'
|
||||
'we also surpport the zip of '
|
||||
'one level subfolder of images whose format are in jpg, jpeg, png, webp.\n'
|
||||
'one level subfolder of images whose format are in jpg, jpeg, png, webp.'
|
||||
f'The ZIP example is: {self.default_dataset_zip}. \n' # noqa
|
||||
f'* If you have refreshed the page, please click the {self.refresh_list_button} '
|
||||
'button to ensure all previously created datasets are visible in the dropdown menu.\n'
|
||||
'* ZIP example: https://modelscope.cn/api/v1/models/damo/scepter/repo?Revision=master&FilePath=datasets/3D_example_csv.zip \n' # noqa
|
||||
'* For processing and training with large-scale data, it is recommended to use the command line.'
|
||||
)
|
||||
'* For processing and training with large-scale data(for example more than 10K samples), '
|
||||
'it is recommended to use the command line to train the model.'
|
||||
'* <span style="color: blue;">Please pay attention to the output of '
|
||||
'the system logs to help improve operations.</span> \n')
|
||||
# Error or Warning
|
||||
self.illegal_data_name_err1 = (
|
||||
'The data name is empty or contains illegal '
|
||||
@@ -39,22 +62,40 @@ class CreateDatasetUIName():
|
||||
self.illegal_data_name_err4 = 'Please do not upload files and set dataset links simultaneously.'
|
||||
self.illegal_data_name_err5 = 'Invalid dataset name, please switch datasets or create a new one.'
|
||||
self.illegal_data_err1 = 'File download failed'
|
||||
self.illegal_data_err2 = 'Illegal file format'
|
||||
self.illegal_data_err3 = 'File decompression failed, failed to upload to storage!'
|
||||
self.delete_data_err1 = 'The example dataset is not allowed delete!'
|
||||
self.modify_data_err1 = 'The example dataset is not allowed modify!'
|
||||
self.modify_data_name_err1 = 'Failed to change dataset name!'
|
||||
self.refresh_data_list_info1 = (
|
||||
'The dataset name has been changed, '
|
||||
'please refresh the list and try again.')
|
||||
self.use_link = 'Use File Link'
|
||||
elif language == 'zh':
|
||||
self.dataset_name = '数据集'
|
||||
self.btn_create_datasets = '新建'
|
||||
self.user_data_name = '当前数据集名称'
|
||||
self.modify_data_button = '修改数据集名称'
|
||||
self.confirm_data_button = '确认'
|
||||
self.refresh_list_button = '刷新列表'
|
||||
self.system_log = '<span style="color: blue;">系统日志: {}</span> '
|
||||
self.btn_create_datasets = '\U00002795' # ➕
|
||||
self.get_data_name_button = '\U0001F3B2' # 🎲
|
||||
self.new_data_name = ('新数据集名称,替换"name"和"version"为方便记忆名称.'
|
||||
f'可以通过点击{self.get_data_name_button}获取随机名称')
|
||||
self.modify_data_button = '\U0001F4DD' # 📝
|
||||
self.confirm_data_button = '\U00002714' # ✔️
|
||||
self.cancel_create_button = '\U00002716' # ✖️
|
||||
self.refresh_list_button = '\U0001f504' # 🔄
|
||||
self.delete_dataset_button = '\U0001f5d1' # 🗑️
|
||||
self.dataset_name = (f'数据集,点击{self.btn_create_datasets}新建数据集,'
|
||||
f'点击{self.delete_dataset_button}删除数据集')
|
||||
self.dataset_type = '数据集类型'
|
||||
self.dataset_type_name = {
|
||||
'scepter_txt2img': '文生图数据',
|
||||
'scepter_img2img': '图像编辑(图生图)数据'
|
||||
}
|
||||
|
||||
self.user_data_name = f'当前数据集名称,修改后点{self.modify_data_button}生效'
|
||||
self.zip_file = '上传数据集'
|
||||
self.zip_file_url = '数据集链接'
|
||||
self.default_dataset_repo = 'https://www.modelscope.cn/api/v1/models/iic/scepter/'
|
||||
self.default_dataset_zip = \
|
||||
self.default_dataset_repo + 'repo?Revision=master&FilePath=datasets/3D_example_csv.zip'
|
||||
self.default_dataset_name = '3D_example'
|
||||
self.btn_create_datasets_from_file = '从文件新建'
|
||||
self.user_direction = (
|
||||
'### 使用说明 \n' +
|
||||
@@ -63,19 +104,23 @@ class CreateDatasetUIName():
|
||||
f'* {self.btn_create_datasets_from_file} 按钮支持从文件中来新建数据集,目前支持zip文件,'
|
||||
'需要保证在文件夹外进行打包,并包含images/文件夹和train.csv(会使用该文件中的图片路径),首行为Target:FILE,Prompt,'
|
||||
'其次每行格式为:图片路径,描述;'
|
||||
'同时我们也支持图像文件的zip包,格式在jpg、jpeg、png或webp \n' +
|
||||
f'同时我们也支持图像文件的zip包,格式在jpg、jpeg、png或webp。数据ZIP样例路径:{self.default_dataset_zip}. \n'
|
||||
+
|
||||
f'* 如果刷新了页面,请点击{self.refresh_list_button} 按钮以确保所有以往创建的数据集在下拉框中可见。\n'
|
||||
'* ZIP样例路径:https://modelscope.cn/api/v1/models/damo/scepter/repo?Revision=master&FilePath=datasets/3D_example_csv.zip \n' # noqa
|
||||
'* 对于大规模数据的处理和训练,建议使用命令行形式')
|
||||
'* 对于大规模数据的处理和训练(数据规模大于1万),建议使用命令行形式\n'
|
||||
'* <span style="color: blue;">请注意观察系统日志的输出以帮助改进操作。</span> \n')
|
||||
# Error or Warning
|
||||
self.illegal_data_name_err1 = "数据名称为空或包含非法字符' '或者'/'"
|
||||
self.illegal_data_name_err2 = '请按照{name}-{version}-{randomstr}'
|
||||
self.illegal_data_name_err2 = '数据名称应该按照{name}-{version}-{randomstr}'
|
||||
self.illegal_data_name_err3 = "数据集名称中不要包含'.'"
|
||||
self.illegal_data_name_err4 = '请不要同时上传文件和设置数据集链接'
|
||||
self.illegal_data_name_err5 = '不合法的数据集名称,请切换数据集或新建数据集。'
|
||||
self.illegal_data_err1 = '文件下载失败'
|
||||
self.illegal_data_err2 = '非法的文件格式'
|
||||
self.illegal_data_err3 = '文件解压失败,上传存储器失败!'
|
||||
|
||||
self.delete_data_err1 = '示例数据集不允许删除!'
|
||||
self.modify_data_err1 = '示例数据集不允许修改!'
|
||||
|
||||
self.modify_data_name_err1 = '变更数据集名称失败!'
|
||||
self.refresh_data_list_info1 = '该数据集名称发生了变更,请刷新列表试一下。'
|
||||
self.use_link = '使用文件链接'
|
||||
@@ -84,27 +129,126 @@ class CreateDatasetUIName():
|
||||
class DatasetGalleryUIName():
|
||||
def __init__(self, language='en'):
|
||||
if language == 'en':
|
||||
self.system_log = '<span style="color: blue;">System Log: {}</span> '
|
||||
self.upload_image = 'Upload Image'
|
||||
self.upload_image_btn = 'Upload'
|
||||
self.upload_image_btn = '\U00002714' # ✔️
|
||||
self.cancel_upload_btn = '\U00002716' # ✖️
|
||||
self.image_caption = 'Image Caption'
|
||||
self.dataset_images = 'Dataset Images'
|
||||
|
||||
self.ori_caption = 'Original Caption'
|
||||
self.edit_caption = 'Editable Caption'
|
||||
self.btn_modify = 'Replace Caption'
|
||||
self.btn_delete = 'Delete Image'
|
||||
self.btn_modify = '\U0001F4DD' # 📝
|
||||
self.btn_delete = '\U0001f5d1' # 🗑️
|
||||
self.btn_add = '\U00002795' # ➕
|
||||
self.dataset_images = f'Original Images,click{self.btn_modify} into editable mode.'
|
||||
self.edit_caption = f'Editable Caption,click{self.btn_modify} into editable mode.'
|
||||
self.dataset_images = f'Original Images,click{self.btn_modify} into editable mode.'
|
||||
|
||||
self.ori_dataset = 'Original Data Height({}) * Width({}) and Image Format({})'
|
||||
self.edit_dataset = 'Editable Data Height({}) * Width({}) and Image Format({})'
|
||||
self.upload_image_info = 'Image Information: Height({}) * Width({}) and Image Format({})'
|
||||
|
||||
self.range_mode_name = [
|
||||
'Current sample', 'All samples', 'Samples in range'
|
||||
]
|
||||
self.samples_range = 'The samples range to be process.'
|
||||
self.samples_range_placeholder = (
|
||||
'"1,4,6" indicates to process 1st, 4th and 6th sample;'
|
||||
'"1-6" indicates to process samples from 1st to 6th.'
|
||||
'"1-4,6-8"indicates to process samples from 1st to 4th '
|
||||
'and from 6th to 8th.')
|
||||
self.set_range_name = 'Samples Range to be edited'
|
||||
self.btn_confirm_edit = '\U00002714' # ✔️
|
||||
self.btn_cancel_edit = '\U00002716' # ✖️
|
||||
self.btn_reset_edit = '\U000021BA' # ↺
|
||||
self.confirm_direction = (
|
||||
f'click{self.btn_confirm_edit} to apply all changes,'
|
||||
f'click{self.btn_reset_edit} to reset edited data,'
|
||||
f'click{self.btn_cancel_edit} to out of editing mode.')
|
||||
self.preprocess_choices = [
|
||||
'Image Preprocess', 'Caption Preprocess'
|
||||
]
|
||||
self.image_processor_type = 'Image Preprocessors'
|
||||
self.caption_processor_type = 'Caption Preprocessors'
|
||||
self.image_preprocess_btn = 'Run'
|
||||
self.caption_preprocess_btn = 'Run'
|
||||
self.caption_update_mode = 'Caption Update Mode'
|
||||
self.caption_update_choices = ['Append', 'Replace']
|
||||
|
||||
self.used_device = 'Used Device'
|
||||
self.used_memory = 'Used Memory'
|
||||
self.caption_language = "Caption's Language"
|
||||
self.advance_setting = 'Generation Setting'
|
||||
self.system_prompt = 'System Prompt'
|
||||
self.max_new_tokens = 'Max New Tokens'
|
||||
self.min_new_tokens = 'Min New Tokens'
|
||||
self.num_beams = 'Beams Num'
|
||||
self.repetition_penalty = 'Repetition Penalty'
|
||||
self.temperature = 'Temperature'
|
||||
|
||||
self.height_ratio = 'Height side scale'
|
||||
self.width_ratio = 'Width side scale'
|
||||
# Error or Warning
|
||||
self.delete_err1 = 'Deletion failed, the data is already empty.'
|
||||
|
||||
elif language == 'zh':
|
||||
self.system_log = '<span style="color: blue;">系统日志: {}</span> '
|
||||
self.upload_image = '上传图片'
|
||||
self.upload_image_btn = '上传'
|
||||
self.upload_image_btn = '\U00002714' # ✔️
|
||||
self.cancel_upload_btn = '\U00002716' # ✖️
|
||||
self.image_caption = '图片描述'
|
||||
self.dataset_images = '图片集'
|
||||
self.ori_caption = '原始描述'
|
||||
|
||||
# self.image_height = '高度'
|
||||
# self.image_width = '宽度'
|
||||
# self.image_format = '格式'
|
||||
|
||||
self.btn_modify = '\U0001F4DD' # 📝
|
||||
self.dataset_images = f'图片数据,点击{self.btn_modify}进入编辑模式'
|
||||
|
||||
self.btn_delete = '\U0001f5d1' # 🗑️
|
||||
self.btn_add = '\U00002795' # ➕
|
||||
|
||||
self.ori_caption = f'原始描述,点击{self.btn_modify}进入编辑模式'
|
||||
self.edit_caption = '编辑描述'
|
||||
self.btn_modify = '替换描述'
|
||||
self.btn_delete = '删除图片'
|
||||
self.batch_caption_generate = '处理范围'
|
||||
|
||||
self.ori_dataset = '原始数据 高({}) * 宽({}) 图像格式({})'
|
||||
self.edit_dataset = '可编辑数据 高({}) * 宽({}) 图像格式({})'
|
||||
self.upload_image_info = '图像信息 高({}) * 宽({})'
|
||||
|
||||
self.range_mode_name = ['当前样本', '全部样本', '指定范围']
|
||||
self.samples_range = '处理样本范围'
|
||||
self.samples_range_placeholder = (
|
||||
'"1,4,6"代表处理第1,4,6个样本;'
|
||||
'"1-6" 代表处理从第1个到第6个的全部样本;'
|
||||
'"1-4,6-8" 代表处理从第1个到第4个,第6到第8个样本。')
|
||||
self.set_range_name = '编辑数据范围'
|
||||
|
||||
self.btn_confirm_edit = '\U00002714' # ✔️
|
||||
self.btn_cancel_edit = '\U00002716' # ✖️
|
||||
self.btn_reset_edit = '\U000021BA' # ↺
|
||||
self.confirm_direction = (f'点击{self.btn_confirm_edit}使所有编辑内容生效,'
|
||||
f'点击{self.btn_cancel_edit}取消编辑,'
|
||||
f'点击{self.btn_reset_edit}重置数据,'
|
||||
f'修改编辑范围可以批量编辑不同范围的数据。')
|
||||
self.preprocess_choices = ['图像预处理', '描述生成']
|
||||
self.image_processor_type = '图像预处理器'
|
||||
self.caption_processor_type = '描述生成器'
|
||||
self.image_preprocess_btn = '运行'
|
||||
self.caption_preprocess_btn = '运行'
|
||||
self.caption_update_mode = '描述更新方式'
|
||||
self.caption_update_choices = ['追加', '替换']
|
||||
self.used_device = '使用设备'
|
||||
self.used_memory = '使用内存'
|
||||
self.caption_language = '描述语言'
|
||||
self.advance_setting = '生成设置'
|
||||
self.system_prompt = '系统提示'
|
||||
self.max_new_tokens = '描述最大长度'
|
||||
self.min_new_tokens = '描述最小长度'
|
||||
self.num_beams = 'Beams数'
|
||||
self.repetition_penalty = '重复惩罚'
|
||||
self.temperature = '温度系数'
|
||||
self.height_ratio = '高度比例'
|
||||
self.width_ratio = '宽度比例'
|
||||
# Error or Warning
|
||||
self.delete_err1 = '删除失败,数据已经为空了'
|
||||
|
||||
|
||||
class ExportDatasetUIName():
|
||||
@@ -115,7 +259,7 @@ class ExportDatasetUIName():
|
||||
self.export_file = 'Download Data'
|
||||
# Error or Warning
|
||||
self.export_err1 = 'The dataset is empty, export is not possible!'
|
||||
self.export_zip_err1 = 'Failed to compress the file!'
|
||||
|
||||
self.upload_err1 = 'Failed to compress the file!'
|
||||
self.go_to_train = 'Go to train...'
|
||||
elif language == 'zh':
|
||||
@@ -123,6 +267,66 @@ class ExportDatasetUIName():
|
||||
self.btn_export_list = '导出列表'
|
||||
self.export_file = '下载数据'
|
||||
self.export_err1 = '数据集为空,无法导出!'
|
||||
self.export_zip_err1 = '压缩文件失败!'
|
||||
|
||||
self.upload_err1 = '压缩文件失败!'
|
||||
self.go_to_train = '去训练...'
|
||||
|
||||
|
||||
class Text2ImageDataCardName():
|
||||
def __init__(self, language='en'):
|
||||
if language == 'en':
|
||||
self.illegal_data_err1 = (
|
||||
'The list supports only "," or "#;#" as delimiters. '
|
||||
'The four columns represent image path, width, height, '
|
||||
'and description, respectively.')
|
||||
self.illegal_data_err2 = 'Illegal file format'
|
||||
self.illegal_data_err3 = 'File decompression failed, failed to upload to storage!'
|
||||
self.illegal_data_err4 = 'Illegal width({}),height({})'
|
||||
self.illegal_data_err5 = (
|
||||
'The path should not contain "{}". '
|
||||
'It should be an OSS path (oss://) or the prefix '
|
||||
'can be omitted (xxx/xxx)."')
|
||||
self.illegal_data_err6 = 'Image download failed {}'
|
||||
self.illegal_data_err7 = 'Image upload failed {}'
|
||||
self.delete_err1 = 'Deletion failed, the data is already empty.'
|
||||
self.export_zip_err1 = 'Failed to compress the file!'
|
||||
elif language == 'zh':
|
||||
self.illegal_data_err1 = '列表只支持,或#;#作为分割符,四列分别为图像路径/宽/高/描述'
|
||||
self.illegal_data_err2 = '非法的文件格式'
|
||||
self.illegal_data_err3 = '文件解压失败,上传存储器失败!'
|
||||
self.illegal_data_err4 = '不合法的width({}),height({})'
|
||||
self.illegal_data_err5 = '路径不支持{},应该为oss路径(oss://)或者省略前缀(xxx/xxx)'
|
||||
self.illegal_data_err6 = '下载图像失败{}'
|
||||
self.illegal_data_err7 = '上传图像失败{}'
|
||||
self.delete_err1 = '删除失败,数据已经为空了'
|
||||
self.export_zip_err1 = '压缩文件失败!'
|
||||
|
||||
|
||||
class Image2ImageDataCardName():
|
||||
def __init__(self, language='en'):
|
||||
if language == 'en':
|
||||
self.illegal_data_err1 = (
|
||||
'The list supports only "," or "#;#" as delimiters. '
|
||||
'The four columns represent image path, width, height, '
|
||||
'and description, respectively.')
|
||||
self.illegal_data_err2 = 'Illegal file format'
|
||||
self.illegal_data_err3 = 'File decompression failed, failed to upload to storage!'
|
||||
self.illegal_data_err4 = 'Illegal width({}),height({})'
|
||||
self.illegal_data_err5 = (
|
||||
'The path should not contain "{}". '
|
||||
'It should be an OSS path (oss://) or the prefix '
|
||||
'can be omitted (xxx/xxx)."')
|
||||
self.illegal_data_err6 = 'Image download failed {}'
|
||||
self.illegal_data_err7 = 'Image upload failed {}'
|
||||
self.delete_err1 = 'Deletion failed, the data is already empty.'
|
||||
self.export_zip_err1 = 'Failed to compress the file!'
|
||||
elif language == 'zh':
|
||||
self.illegal_data_err1 = '列表只支持,或#;#作为分割符,四列分别为图像路径/宽/高/描述'
|
||||
self.illegal_data_err2 = '非法的文件格式'
|
||||
self.illegal_data_err3 = '文件解压失败,上传存储器失败!'
|
||||
self.illegal_data_err4 = '不合法的width({}),height({})'
|
||||
self.illegal_data_err5 = '路径不支持{},应该为oss路径(oss://)或者省略前缀(xxx/xxx)'
|
||||
self.illegal_data_err6 = '下载图像失败{}'
|
||||
self.illegal_data_err7 = '上传图像失败{}'
|
||||
self.delete_err1 = '删除失败,数据已经为空了'
|
||||
self.export_zip_err1 = '压缩文件失败!'
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -3,25 +3,28 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import urllib.parse as parse
|
||||
|
||||
import gradio as gr
|
||||
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.studio.preprocess.caption_editor_ui.component_names import \
|
||||
ExportDatasetUIName
|
||||
from scepter.studio.utils.uibase import UIBase
|
||||
|
||||
|
||||
class ExportDatasetUI(UIBase):
|
||||
def __init__(self, cfg, is_debug=False, language='en'):
|
||||
def __init__(self, cfg, is_debug=False, language='en', gallery_ins=None):
|
||||
self.dataset_name = ''
|
||||
self.work_dir = cfg.WORK_DIR
|
||||
self.export_folder = os.path.join(self.work_dir, cfg.EXPORT_DIR)
|
||||
self.component_names = ExportDatasetUIName(language)
|
||||
if gallery_ins is not None:
|
||||
self.default_dataset = gallery_ins.default_dataset
|
||||
else:
|
||||
self.default_dataset = None
|
||||
|
||||
def create_ui(self):
|
||||
with gr.Row(variant='panel', visible=False,
|
||||
with gr.Row(variant='panel',
|
||||
visible=self.default_dataset is not None,
|
||||
equal_height=True) as export_panel:
|
||||
self.data_state = gr.State(value=False)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
@@ -39,104 +42,34 @@ class ExportDatasetUI(UIBase):
|
||||
self.export_panel = export_panel
|
||||
|
||||
def set_callbacks(self, create_dataset, manager):
|
||||
def export_zip(dataset_name):
|
||||
meta = create_dataset.meta_dict[dataset_name]
|
||||
work_dir = meta['work_dir']
|
||||
local_work_dir = meta['local_work_dir']
|
||||
train_csv = os.path.join(work_dir, 'train.csv')
|
||||
if len(meta['file_list']) < 1:
|
||||
raise gr.Error(self.component_names.export_err1)
|
||||
train_csv = create_dataset.write_csv(meta['file_list'], train_csv,
|
||||
work_dir)
|
||||
_ = FS.get_from(train_csv, os.path.join(local_work_dir,
|
||||
'train.csv'))
|
||||
save_file_list = work_dir + '_file.csv'
|
||||
save_file_list = create_dataset.write_file_list(
|
||||
meta['file_list'], save_file_list)
|
||||
_ = FS.get_from(save_file_list,
|
||||
os.path.join(local_work_dir, 'file.csv'))
|
||||
zip_path = os.path.join(self.export_folder, f'{dataset_name}.zip')
|
||||
with FS.put_to(zip_path) as local_zip:
|
||||
res = os.popen(
|
||||
f"cd '{local_work_dir}' && mkdir -p '{dataset_name}' "
|
||||
f"&& cp -rf images '{dataset_name}/images' "
|
||||
f"&& cp -rf train.csv '{dataset_name}/train.csv' "
|
||||
f"&& zip -r '{os.path.abspath(local_zip)}' '{dataset_name}'/* "
|
||||
f"&& rm -rf '{dataset_name}'")
|
||||
print(res.readlines())
|
||||
|
||||
if not FS.exists(zip_path):
|
||||
raise gr.Error(self.component_names.export_zip_err1)
|
||||
create_dataset.save_meta(meta, work_dir)
|
||||
local_zip = FS.get_from(zip_path)
|
||||
def export_zip(dataset_type, dataset_name):
|
||||
dataset_type = create_dataset.get_trans_dataset_type(dataset_type)
|
||||
dataset_ins = create_dataset.dataset_dict[dataset_type][
|
||||
dataset_name]
|
||||
local_zip = dataset_ins.export_zip(self.export_folder)
|
||||
return gr.File(value=local_zip, visible=True)
|
||||
|
||||
self.export_to_zip.click(export_zip,
|
||||
inputs=[create_dataset.user_data_name],
|
||||
outputs=[self.export_url],
|
||||
queue=False)
|
||||
self.export_to_zip.click(
|
||||
export_zip,
|
||||
inputs=[create_dataset.dataset_type, create_dataset.dataset_name],
|
||||
outputs=[self.export_url],
|
||||
queue=False)
|
||||
|
||||
def export_csv(dataset_name):
|
||||
meta = create_dataset.meta_dict[dataset_name]
|
||||
work_dir = meta['work_dir']
|
||||
local_work_dir = meta['local_work_dir']
|
||||
train_csv = os.path.join(work_dir, 'train.csv')
|
||||
if len(meta['file_list']) < 1:
|
||||
raise gr.Error(self.component_names.export_err1)
|
||||
train_csv = create_dataset.write_csv(meta['file_list'], train_csv,
|
||||
work_dir)
|
||||
_ = FS.get_from(train_csv, os.path.join(local_work_dir,
|
||||
'train.csv'))
|
||||
save_file_list = os.path.join(work_dir, 'file.csv')
|
||||
save_file_list = create_dataset.write_file_list(
|
||||
meta['file_list'], save_file_list)
|
||||
local_file_csv = FS.get_from(
|
||||
save_file_list, os.path.join(local_work_dir, 'file.csv'))
|
||||
create_dataset.save_meta(meta, work_dir)
|
||||
is_flag = FS.put_object_from_local_file(
|
||||
local_file_csv,
|
||||
os.path.join(self.export_folder, dataset_name + '_file.csv'))
|
||||
if not is_flag:
|
||||
raise gr.Error(self.component_names.upload_err1)
|
||||
list_url = FS.get_url(os.path.join(self.export_folder,
|
||||
dataset_name + '_file.csv'),
|
||||
set_public=True)
|
||||
list_url = parse.unquote(list_url)
|
||||
if 'wulanchabu' in list_url:
|
||||
list_url = list_url.replace(
|
||||
'.cn-wulanchabu.oss-internal.aliyun-inc.',
|
||||
'.oss-cn-wulanchabu.aliyuncs.')
|
||||
else:
|
||||
list_url = list_url.replace('.oss-internal.aliyun-inc.',
|
||||
'.oss.aliyuncs.')
|
||||
if not list_url.split('/')[-1] == dataset_name + '_file.csv':
|
||||
list_url = os.path.join(os.path.dirname(list_url),
|
||||
dataset_name + '_file.csv')
|
||||
return gr.Text(value=list_url)
|
||||
|
||||
# self.export_to_list.click(export_csv,
|
||||
# inputs=[create_dataset.user_data_name],
|
||||
# outputs=[self.export_url])
|
||||
|
||||
def go_to_train(dataset_name):
|
||||
meta = create_dataset.meta_dict[dataset_name]
|
||||
work_dir = meta['work_dir']
|
||||
local_work_dir = meta['local_work_dir']
|
||||
train_csv = os.path.join(work_dir, 'train.csv')
|
||||
if len(meta['file_list']) < 1:
|
||||
raise gr.Error(self.component_names.export_err1)
|
||||
train_csv = create_dataset.write_csv(meta['file_list'], train_csv,
|
||||
work_dir)
|
||||
_ = FS.get_from(train_csv, os.path.join(local_work_dir,
|
||||
'train.csv'))
|
||||
save_file_list = work_dir + '_file.csv'
|
||||
_ = create_dataset.write_file_list(meta['file_list'],
|
||||
save_file_list)
|
||||
def go_to_train(dataset_type, dataset_name):
|
||||
dataset_type = create_dataset.get_trans_dataset_type(dataset_type)
|
||||
dataset_ins = create_dataset.dataset_dict[dataset_type][
|
||||
dataset_name]
|
||||
dataset_ins.update_dataset()
|
||||
return (gr.Tabs(selected='self_train'),
|
||||
gr.Textbox(value=os.path.abspath(local_work_dir)))
|
||||
gr.Textbox(
|
||||
value=os.path.abspath(dataset_ins.local_work_dir)),
|
||||
dataset_name)
|
||||
|
||||
self.go_to_train.click(
|
||||
go_to_train,
|
||||
inputs=[create_dataset.user_data_name],
|
||||
outputs=[manager.tabs, manager.self_train.trainer_ui.ms_data_name],
|
||||
inputs=[create_dataset.dataset_type, create_dataset.dataset_name],
|
||||
outputs=[
|
||||
manager.tabs, manager.self_train.trainer_ui.ms_data_name,
|
||||
manager.self_train.trainer_ui.ori_data_name
|
||||
],
|
||||
queue=False)
|
||||
|
||||
@@ -32,22 +32,26 @@ class PreprocessUI():
|
||||
self.create_dataset = CreateDatasetUI.get_instance(cfg_general,
|
||||
is_debug=is_debug,
|
||||
language=language)
|
||||
self.dataset_gallery = DatasetGalleryUI.get_instance(cfg_general,
|
||||
is_debug=is_debug,
|
||||
language=language)
|
||||
self.export_dataset = ExportDatasetUI.get_instance(cfg_general,
|
||||
is_debug=is_debug,
|
||||
language=language)
|
||||
self.dataset_gallery = DatasetGalleryUI.get_instance(
|
||||
cfg_general,
|
||||
is_debug=is_debug,
|
||||
language=language,
|
||||
create_ins=self.create_dataset)
|
||||
self.export_dataset = ExportDatasetUI.get_instance(
|
||||
cfg_general,
|
||||
is_debug=is_debug,
|
||||
language=language,
|
||||
gallery_ins=self.dataset_gallery)
|
||||
|
||||
def create_ui(self):
|
||||
self.create_dataset.create_ui()
|
||||
self.dataset_gallery.create_ui()
|
||||
self.export_dataset.create_ui()
|
||||
self.dataset_gallery.create_ui()
|
||||
|
||||
def set_callbacks(self, manager):
|
||||
self.create_dataset.set_callbacks(self.dataset_gallery,
|
||||
self.export_dataset, manager)
|
||||
self.dataset_gallery.set_callbacks(self.create_dataset)
|
||||
self.dataset_gallery.set_callbacks(self.create_dataset, manager)
|
||||
self.export_dataset.set_callbacks(self.create_dataset, manager)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
@@ -0,0 +1,132 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import torch
|
||||
|
||||
from scepter.studio.utils.env import get_available_memory
|
||||
|
||||
|
||||
class BaseCaptionProcessor(object):
|
||||
def __init__(self, cfg, language='en'):
|
||||
self.use_device = cfg.get('DEVICE', 'cpu')
|
||||
self.use_memory = cfg.get('MEMORY', 10)
|
||||
self.language = language
|
||||
self.system_paras = cfg.get('PARAS', [])
|
||||
self.language_level_paras = {}
|
||||
for sys_para in self.system_paras:
|
||||
if language == 'en':
|
||||
cur_lang = sys_para.get('LANGUAGE_NAME', None)
|
||||
else:
|
||||
cur_lang = sys_para.get('LANGUAGE_ZH_NAME', None)
|
||||
if cur_lang is not None:
|
||||
self.language_level_paras[cur_lang] = sys_para
|
||||
|
||||
def unload_model(self):
|
||||
mem = get_available_memory()
|
||||
free_mem = int(mem['available'] / (1024**2))
|
||||
total_mem = int(mem['total'] / (1024**2))
|
||||
self.delete_instance = False
|
||||
if free_mem < 0.5 * total_mem:
|
||||
self.delete_instance = True
|
||||
return True, ''
|
||||
|
||||
def load_model(self):
|
||||
is_flg, msg = self.check_memory()
|
||||
return is_flg, msg
|
||||
|
||||
@property
|
||||
def get_language_choice(self):
|
||||
language_choices = list(self.language_level_paras.keys())
|
||||
return language_choices
|
||||
|
||||
@property
|
||||
def get_language_default(self):
|
||||
language_choices = list(self.language_level_paras.keys())
|
||||
return language_choices[0] if len(language_choices) > 0 else None
|
||||
|
||||
def get_para_by_language(self, language):
|
||||
return self.language_level_paras.get(language, {})
|
||||
|
||||
def check_memory(self):
|
||||
mem_msg = ''
|
||||
if self.use_device == 'gpu':
|
||||
# Check Cuda Memory
|
||||
if torch.cuda.is_available():
|
||||
for device_id in range(torch.cuda.device_count()):
|
||||
free_mem, total_mem = torch.cuda.mem_get_info(device_id)
|
||||
free_mem = int(free_mem / (1024**2))
|
||||
total_mem = int(total_mem / (1024**2))
|
||||
if free_mem < self.use_memory:
|
||||
mem_msg += (
|
||||
f'Needed {self.use_memory}M, but free mem '
|
||||
f'is {free_mem:.3f}M(total is {total_mem})M \n')
|
||||
else:
|
||||
mem_msg += 'Needed GPU device, but this device is not available!'
|
||||
elif self.use_device == 'cpu':
|
||||
mem = get_available_memory()
|
||||
free_mem = int(mem['available'] / (1024**2))
|
||||
total_mem = int(mem['total'] / (1024**2))
|
||||
if free_mem < self.use_memory:
|
||||
mem_msg += (f'Needed {self.use_memory}M, but free mem '
|
||||
f'is {free_mem:.3f}M(total is {total_mem})M \n')
|
||||
if mem_msg == '':
|
||||
return True, mem_msg
|
||||
return False, mem_msg
|
||||
|
||||
def __call__(self, image, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class BaseImageProcessor(object):
|
||||
def __init__(self, cfg, language='en'):
|
||||
self.use_device = cfg.get('DEVICE', 'cpu')
|
||||
self.use_memory = cfg.get('MEMORY', 10)
|
||||
self.language = language
|
||||
self.system_paras = cfg.get('PARAS', {})
|
||||
self.language_level_paras = {}
|
||||
|
||||
def unload_model(self):
|
||||
mem = get_available_memory()
|
||||
free_mem = int(mem['available'] / (1024**2))
|
||||
total_mem = int(mem['total'] / (1024**2))
|
||||
self.delete_instance = False
|
||||
if free_mem < 0.5 * total_mem:
|
||||
self.delete_instance = True
|
||||
return True, ''
|
||||
|
||||
@property
|
||||
def system_para(self):
|
||||
return self.system_paras
|
||||
|
||||
def load_model(self):
|
||||
is_flg, msg = self.check_memory()
|
||||
return is_flg, msg
|
||||
|
||||
def check_memory(self):
|
||||
mem_msg = ''
|
||||
if self.use_device == 'gpu':
|
||||
# Check Cuda Memory
|
||||
if torch.cuda.is_available():
|
||||
for device_id in range(torch.cuda.device_count()):
|
||||
free_mem, total_mem = torch.cuda.mem_get_info(device_id)
|
||||
free_mem = int(free_mem / (1024**2))
|
||||
total_mem = int(total_mem / (1024**2))
|
||||
if free_mem < self.use_memory:
|
||||
mem_msg += (
|
||||
f'Needed {self.use_memory}M, but free mem '
|
||||
f'is {free_mem:.3f}M(total is {total_mem})M \n')
|
||||
else:
|
||||
mem_msg += 'Needed GPU device, but this device is not available!'
|
||||
elif self.use_device == 'cpu':
|
||||
mem = get_available_memory()
|
||||
free_mem = int(mem['available'] / (1024**2))
|
||||
total_mem = int(mem['total'] / (1024**2))
|
||||
if free_mem < self.use_memory:
|
||||
mem_msg += (f'Needed {self.use_memory}M, but free mem '
|
||||
f'is {free_mem:.3f}M(total is {total_mem})M \n')
|
||||
if mem_msg == '':
|
||||
return True, mem_msg
|
||||
return False, mem_msg
|
||||
|
||||
def __call__(self, image, **kwargs):
|
||||
raise NotImplementedError
|
||||
@@ -0,0 +1,258 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import numbers
|
||||
import re
|
||||
import time
|
||||
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.studio.preprocess.processors.base_processor import \
|
||||
BaseCaptionProcessor
|
||||
|
||||
__all__ = ['BlipImageBase', 'QWVL', 'QWVLQuantize']
|
||||
|
||||
|
||||
class BlipImageBase(BaseCaptionProcessor):
|
||||
def __init__(self, cfg, language='en'):
|
||||
super().__init__(cfg, language=language)
|
||||
self.model_path = cfg.MODEL_PATH
|
||||
self.model_info = {
|
||||
'device': 'offline',
|
||||
'model': None,
|
||||
'tokenizer': None
|
||||
}
|
||||
|
||||
def load_model(self):
|
||||
is_flg, msg = super().load_model()
|
||||
if not is_flg:
|
||||
return is_flg, msg
|
||||
if self.model_info['device'] == 'offline':
|
||||
model = None
|
||||
processor = None
|
||||
try:
|
||||
from transformers import BlipProcessor, BlipForConditionalGeneration
|
||||
local_model_dir = FS.get_dir_to_local_dir(self.model_path)
|
||||
processor = BlipProcessor.from_pretrained(local_model_dir)
|
||||
model = BlipForConditionalGeneration.from_pretrained(
|
||||
local_model_dir).to(we.device_id)
|
||||
except Exception as e:
|
||||
if model is not None:
|
||||
del model
|
||||
if processor is not None:
|
||||
del model
|
||||
return False, f"Load model error '{e}'"
|
||||
self.model_info['device'] = model.device
|
||||
self.model_info['model'] = model
|
||||
self.model_info['processor'] = processor
|
||||
elif self.model_info['device'] == 'cpu':
|
||||
try:
|
||||
self.model_info['model'].to(we.device_id)
|
||||
self.model_info['device'] = we.device_id
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
except Exception as e:
|
||||
del self.model_info['model']
|
||||
self.model_info['model'] = None
|
||||
self.model_info['device'] = 'offline'
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
return False, f"Load model error '{e}'"
|
||||
|
||||
return True, ''
|
||||
|
||||
def unload_model(self):
|
||||
super().unload_model()
|
||||
if self.delete_instance:
|
||||
self.model_info['device'] = 'offline'
|
||||
if self.model_info['model'] is not None:
|
||||
self.model_info['model'] = self.model_info['model'].to('cpu')
|
||||
del self.model_info['model']
|
||||
self.model_info['model'] = None
|
||||
elif (isinstance(self.model_info['device'], numbers.Number)
|
||||
or str(self.model_info['device']).startswith('cuda')):
|
||||
self.model_info['device'] = 'cpu'
|
||||
self.model_info['model'] = self.model_info['model'].to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
return True, ''
|
||||
|
||||
def __call__(self, image, prompt=None, **kwargs):
|
||||
raw_image = Image.open(image).convert('RGB')
|
||||
inputs = self.model_info['processor'](
|
||||
raw_image, return_tensors='pt').to(we.device_id)
|
||||
out = self.model_info['model'].generate(**inputs)
|
||||
return self.model_info['processor'].decode(out[0],
|
||||
skip_special_tokens=True)
|
||||
|
||||
|
||||
class QWVL(BaseCaptionProcessor):
|
||||
def __init__(self, cfg, language='en'):
|
||||
super().__init__(cfg, language=language)
|
||||
self.model_path = cfg.MODEL_PATH
|
||||
self.model_info = {
|
||||
'device': 'offline',
|
||||
'model': None,
|
||||
'tokenizer': None
|
||||
}
|
||||
|
||||
def load_model(self):
|
||||
is_flg, msg = super().load_model()
|
||||
if not is_flg:
|
||||
return is_flg, msg
|
||||
if self.model_info['device'] == 'offline':
|
||||
model = None
|
||||
try:
|
||||
from modelscope import (AutoModelForCausalLM, AutoTokenizer,
|
||||
GenerationConfig)
|
||||
local_model_dir = FS.get_dir_to_local_dir(self.model_path)
|
||||
# without quantization using 19.52G memory
|
||||
# with quantization using 7.7G memory
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
local_model_dir, trust_remote_code=True)
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
local_model_dir,
|
||||
device_map='auto',
|
||||
trust_remote_code=True,
|
||||
fp16=True).eval()
|
||||
model.generation_config = GenerationConfig.from_pretrained(
|
||||
local_model_dir, trust_remote_code=True)
|
||||
except Exception as e:
|
||||
if model is not None:
|
||||
del model
|
||||
return False, f"Load model error '{e}'"
|
||||
self.model_info['device'] = model.device
|
||||
self.model_info['model'] = model
|
||||
self.model_info['tokenizer'] = tokenizer
|
||||
elif self.model_info['device'] == 'cpu':
|
||||
try:
|
||||
self.model_info['model'].to(we.device_id)
|
||||
self.model_info['device'] = we.device_id
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
except Exception as e:
|
||||
del self.model_info['model']
|
||||
self.model_info['model'] = None
|
||||
self.model_info['device'] = 'offline'
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
return False, f"Load model error '{e}'"
|
||||
|
||||
return True, ''
|
||||
|
||||
def unload_model(self):
|
||||
super().unload_model()
|
||||
if self.delete_instance:
|
||||
if self.model_info['model'] is not None:
|
||||
self.model_info['model'] = self.model_info['model'].to('cpu')
|
||||
del self.model_info['model']
|
||||
self.model_info['model'] = None
|
||||
self.model_info['device'] = 'offline'
|
||||
elif (isinstance(self.model_info['device'], numbers.Number)
|
||||
or str(self.model_info['device']).startswith('cuda')):
|
||||
self.model_info['device'] = 'cpu'
|
||||
self.model_info['model'] = self.model_info['model'].to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
return True, ''
|
||||
|
||||
def __call__(self,
|
||||
image,
|
||||
prompt='Generate the caption in English',
|
||||
**kwargs):
|
||||
|
||||
torch.manual_seed(int(time.time()) % 100000)
|
||||
query = self.model_info['tokenizer'].from_list_format([
|
||||
{
|
||||
'image': image
|
||||
},
|
||||
{
|
||||
'text': prompt
|
||||
},
|
||||
])
|
||||
print(kwargs)
|
||||
|
||||
inputs = self.model_info['tokenizer'](query, return_tensors='pt')
|
||||
inputs = inputs.to(self.model_info['device'])
|
||||
pred = self.model_info['model'].generate(**inputs, **kwargs)
|
||||
response = self.model_info['tokenizer'].decode(
|
||||
pred.cpu()[0], skip_special_tokens=True)
|
||||
ret_caption = response.split(prompt)[-1]
|
||||
if ret_caption.startswith(','):
|
||||
ret_caption = ret_caption[1:]
|
||||
regex = re.compile(r'[' + '#®•©™&@·º½¾¿¡§~' + ')' + '(' + ']' + '[' +
|
||||
'}' + '{' + '|' + '\\' + '/' + '*' +
|
||||
r']{1,}') # noqa: E501
|
||||
ret_caption = re.sub(regex, r' ', ret_caption)
|
||||
regex = re.compile(r'^[\-\_]+')
|
||||
ret_caption = re.sub(regex, r'', ret_caption)
|
||||
return ret_caption
|
||||
|
||||
|
||||
class QWVLQuantize(QWVL):
|
||||
def load_model(self):
|
||||
is_flg, msg = super(QWVL, self).load_model()
|
||||
if not is_flg:
|
||||
return is_flg, msg
|
||||
self.model = None
|
||||
if self.model_info['device'] == 'offline':
|
||||
try:
|
||||
from transformers import BitsAndBytesConfig
|
||||
torch.manual_seed(int(time.time()))
|
||||
from modelscope import (AutoModelForCausalLM, AutoTokenizer,
|
||||
GenerationConfig)
|
||||
local_model_dir = FS.get_dir_to_local_dir(self.model_path)
|
||||
# without quantization using 19.52G memory
|
||||
# with quantization using 7.7G memory
|
||||
quantization_config = BitsAndBytesConfig(
|
||||
load_in_4bit=True,
|
||||
bnb_4bit_compute_dtype=torch.float16,
|
||||
bnb_4bit_quant_type='nf4',
|
||||
bnb_4bit_use_double_quant=True,
|
||||
llm_int8_skip_modules=['lm_head', 'attn_pool.attn'])
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
local_model_dir, trust_remote_code=True)
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
local_model_dir,
|
||||
device_map='auto',
|
||||
trust_remote_code=True,
|
||||
fp16=True,
|
||||
quantization_config=quantization_config).eval()
|
||||
model.generation_config = GenerationConfig.from_pretrained(
|
||||
local_model_dir, trust_remote_code=True)
|
||||
# model.to(we.device_id)
|
||||
except Exception as e:
|
||||
if self.model is not None:
|
||||
del self.model
|
||||
return False, f"Load model error '{e}'"
|
||||
self.model_info['device'] = model.device
|
||||
self.model_info['model'] = model
|
||||
self.model_info['tokenizer'] = tokenizer
|
||||
elif self.model_info['device'] == 'cpu':
|
||||
try:
|
||||
self.model_info['model'].to(we.device_id)
|
||||
self.model_info['device'] = we.device_id
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
except Exception as e:
|
||||
del self.model_info['model']
|
||||
self.model_info['model'] = None
|
||||
self.model_info['device'] = 'offline'
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
return False, f"Load model error '{e}'"
|
||||
|
||||
return True, ''
|
||||
|
||||
def unload_model(self):
|
||||
print(self.model_info['device'])
|
||||
if (isinstance(self.model_info['device'], numbers.Number)
|
||||
or str(self.model_info['device']).startswith('cuda')):
|
||||
del self.model_info['model']
|
||||
self.model_info['model'] = None
|
||||
self.model_info['device'] = 'offline'
|
||||
torch.cuda.empty_cache()
|
||||
return True, ''
|
||||
@@ -0,0 +1,38 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import torchvision.transforms as TT
|
||||
|
||||
from scepter.studio.preprocess.processors.base_processor import \
|
||||
BaseImageProcessor
|
||||
|
||||
__all__ = ['CenterCrop', 'PaddingCrop']
|
||||
|
||||
|
||||
class CenterCrop(BaseImageProcessor):
|
||||
def __init__(self, cfg, language='en'):
|
||||
super().__init__(cfg, language=language)
|
||||
|
||||
def __call__(self, image, **kwargs):
|
||||
w, h = image.size
|
||||
height_ratio = kwargs.get('height_ratio', 1)
|
||||
width_ratio = kwargs.get('width_ratio', 1)
|
||||
|
||||
output_h_align_height, output_h_align_width = h, int(h / height_ratio *
|
||||
width_ratio)
|
||||
if output_h_align_height * output_h_align_width <= w * h:
|
||||
output_height, output_width = output_h_align_height, output_h_align_width
|
||||
else:
|
||||
output_height, output_width = int(w / width_ratio *
|
||||
height_ratio), w
|
||||
image = TT.Resize(max(output_height, output_width))(image)
|
||||
image = TT.CenterCrop((output_height, output_width))(image)
|
||||
return image
|
||||
|
||||
|
||||
class PaddingCrop(BaseImageProcessor):
|
||||
def __init__(self, cfg, language='en'):
|
||||
super().__init__(cfg, language=language)
|
||||
|
||||
def get_caption(self, image, **kwargs):
|
||||
return image
|
||||
@@ -0,0 +1,57 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import scepter.studio.preprocess.processors.caption_processors as caption_processors
|
||||
import scepter.studio.preprocess.processors.image_processors as image_processors
|
||||
|
||||
model_dict = {'image': image_processors, 'caption': caption_processors}
|
||||
|
||||
|
||||
class ProcessorsManager():
|
||||
def __init__(self, cfg, language='en'):
|
||||
self.type_level_processor = {}
|
||||
for processor in cfg:
|
||||
name = processor.NAME
|
||||
type = processor.TYPE
|
||||
assert type in model_dict
|
||||
assert hasattr(model_dict[type], name)
|
||||
processor_ins = getattr(model_dict[type], name)(processor,
|
||||
language=language)
|
||||
if type not in self.type_level_processor:
|
||||
self.type_level_processor[type] = {}
|
||||
self.type_level_processor[type][name] = processor_ins
|
||||
|
||||
def dynamic_unload(self, type='all', name='all'):
|
||||
print('Unloading {} processor model'.format(name))
|
||||
if name == 'all':
|
||||
for module_type, module_dict in self.type_level_processor.items():
|
||||
for model_name, processor_ins in module_dict.items():
|
||||
if type == 'all' or type == module_type:
|
||||
processor_ins.unload_model()
|
||||
else:
|
||||
for module_type, module_dict in self.type_level_processor.items():
|
||||
for model_name, processor_ins in module_dict.items():
|
||||
if (type == 'all'
|
||||
or type == module_type) and model_name == name:
|
||||
processor_ins.unload_model()
|
||||
|
||||
def get_choices(self, type):
|
||||
return list(self.type_level_processor.get(type, {}).keys())
|
||||
|
||||
def get_default(self, type):
|
||||
processors_list = self.get_choices(type)
|
||||
return processors_list[0] if len(processors_list) > 0 else None
|
||||
|
||||
def get_default_device(self, type):
|
||||
processor_ins = self.get_processor(type, self.get_default(type))
|
||||
return processor_ins.use_device
|
||||
|
||||
def get_default_memory(self, type):
|
||||
processor_ins = self.get_processor(type, self.get_default(type))
|
||||
return f'{processor_ins.use_memory}M'
|
||||
|
||||
def get_processor(self, type, name):
|
||||
if type not in self.type_level_processor:
|
||||
return None
|
||||
if name in self.type_level_processor[type]:
|
||||
return self.type_level_processor[type].get(name, None)
|
||||
@@ -0,0 +1,2 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
File diff suppressed because it is too large
Load Diff
@@ -95,7 +95,7 @@ def save_config(cfg):
|
||||
cfg.args.cfg_file.split('/')[-1])
|
||||
with FS.put_to(config_path) as local_config_path:
|
||||
with open(local_config_path, 'w') as f_out:
|
||||
f_out.write(cfg.dump())
|
||||
f_out.write(cfg.dump(is_secure=True))
|
||||
|
||||
|
||||
def update_config(cfg):
|
||||
|
||||
@@ -84,21 +84,24 @@ class TrainerUIName():
|
||||
refresh the interface and then click [Refresh Model] at the bottom of the page.
|
||||
The trained model should appear in the [Output Model Name] if training was successful;
|
||||
if not, the training may be incomplete or have failed.
|
||||
- zip example: https://modelscope.cn/api/v1/models/damo/scepter/repo?Revision=master&FilePath=datasets/3D_example_csv.zip
|
||||
- zip example: https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=datasets/3D_example_csv.zip
|
||||
- For processing and training with large-scale data, it is recommended to use the command line.
|
||||
''' # noqa
|
||||
self.data_type_choices = ['Dataset zip', 'MaaS Dataset']
|
||||
self.data_type_value = 'Dataset zip'
|
||||
self.data_type_name = 'Data Source'
|
||||
self.ori_data_name = 'Data Name'
|
||||
self.ms_data_name_place_hold = 'Supports MaaS dataset/local/HTTP Zip package'
|
||||
self.ms_data_space = 'ModelScope Space'
|
||||
self.ms_data_subname = 'MaaS Dataset - Subset'
|
||||
self.training_block = 'Training Parameters'
|
||||
self.base_model = 'Base Model'
|
||||
self.tuner_name = 'Fine-tuning Method'
|
||||
self.tuner_name = 'Tuner Method'
|
||||
self.base_model_revision = 'Model Version Number'
|
||||
self.resolution_height = 'Resolution Height'
|
||||
self.resolution_width = 'Resolution Width'
|
||||
self.resolution_height_max = 'Resolution Height Max'
|
||||
self.resolution_width_max = 'Resolution Width Max'
|
||||
self.train_epoch = 'Number of Training Epochs'
|
||||
self.learning_rate = 'Learning Rate'
|
||||
self.save_interval = 'Save Interval'
|
||||
@@ -109,6 +112,13 @@ class TrainerUIName():
|
||||
self.push_to_hub = 'Push to hub'
|
||||
self.training_button = 'Start Training'
|
||||
self.eval_prompts = 'Eval Prompts'
|
||||
self.tuner_param = 'Tuner Parameters'
|
||||
self.enable_resolution_bucket = 'Enable Resolution Bucket'
|
||||
self.resolution_param = 'Resolution Parameters'
|
||||
self.min_bucket_resolution = 'Min Bucket Resolution'
|
||||
self.max_bucket_resolution = 'Max Bucket Resolution'
|
||||
self.bucket_resolution_steps = 'Bucket Resolution Steps'
|
||||
self.bucket_no_upscale = 'Bucket No Upscale'
|
||||
# Error or Warning
|
||||
self.training_err1 = 'CUDA is unavailable.'
|
||||
self.training_err2 = 'Currently insufficient VRAM, training failed!'
|
||||
@@ -125,12 +135,13 @@ class TrainerUIName():
|
||||
- 训练: 点击【开始训练】
|
||||
- 测试: 完成训练后点击【使用模型】
|
||||
- 注意:超时可能导致连接断开(出现Error),可以等差不多可能训完后,刷新界面再点击页面最后的[刷新模型],即可在[产出模型名称中]出现已经完成训练的模型,若不存在则没有完成训练或训练失败
|
||||
- ZIP样例:https://modelscope.cn/api/v1/models/damo/scepter/repo?Revision=master&FilePath=datasets/3D_example_csv.zip
|
||||
- ZIP样例:https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=datasets/3D_example_csv.zip
|
||||
- 对于大规模数据的处理和训练,建议使用命令行形式
|
||||
''' # noqa
|
||||
self.data_type_choices = ['数据集zip', 'MaaS数据集']
|
||||
self.data_type_value = '数据集zip'
|
||||
self.data_type_name = '数据集来源'
|
||||
self.ori_data_name = '数据集名称'
|
||||
self.ms_data_name_place_hold = '支持MaaS数据集/本地/Http Zip包'
|
||||
self.ms_data_space = 'ModelScope 空间'
|
||||
self.ms_data_subname = 'MaaS数据集-子集'
|
||||
@@ -140,6 +151,8 @@ class TrainerUIName():
|
||||
self.base_model_revision = '模型版本号'
|
||||
self.resolution_height = '训练高度'
|
||||
self.resolution_width = '训练宽度'
|
||||
self.resolution_height_max = '最大训练高度'
|
||||
self.resolution_width_max = '最大训练宽度'
|
||||
self.train_epoch = '训练轮数'
|
||||
self.learning_rate = '学习率'
|
||||
self.save_interval = '存储间隔'
|
||||
@@ -149,6 +162,13 @@ class TrainerUIName():
|
||||
self.work_name = '保存模型名称(刷新获得随机值)'
|
||||
self.push_to_hub = '推送魔搭社区'
|
||||
self.eval_prompts = '评测文本'
|
||||
self.tuner_param = '微调参数'
|
||||
self.enable_resolution_bucket = '开启分辨率分桶'
|
||||
self.resolution_param = '分辨率参数'
|
||||
self.min_bucket_resolution = '最小分桶分辨率'
|
||||
self.max_bucket_resolution = '最大分桶分辨率'
|
||||
self.bucket_resolution_steps = '分桶分辨率步长'
|
||||
self.bucket_no_upscale = '分桶分辨率不做放大'
|
||||
# Error or Warning
|
||||
self.training_err1 = 'CUDA不可用.'
|
||||
self.training_err2 = '目前显存不足,训练失败!'
|
||||
|
||||
@@ -181,8 +181,7 @@ class ModelUI(UIBase):
|
||||
gr.Dropdown(choices=ckpt_list, value=ckpt_value),
|
||||
gr.Gallery(value=gallery_value,
|
||||
preview=True,
|
||||
selected_index=select_index)
|
||||
)
|
||||
selected_index=select_index))
|
||||
|
||||
self.output_model_name.change(fn=model_name_change,
|
||||
inputs=[self.output_model_name],
|
||||
@@ -253,23 +252,22 @@ class ModelUI(UIBase):
|
||||
ckpt_list = self.get_ckpt_list(model_name)
|
||||
ckpt_value = ckpt_list[-1] if len(ckpt_list) > 0 else ''
|
||||
ret_gallery = ckpt_name_change(model_name, ckpt_value)
|
||||
return (message,
|
||||
gr.Column(visible=status in ('running', 'success')),
|
||||
return (message, gr.Column(visible=status in ('running',
|
||||
'success')),
|
||||
gr.Dropdown(choices=self.model_list, value=model_name),
|
||||
gr.Dropdown(choices=ckpt_list,
|
||||
value=ckpt_value), ret_gallery)
|
||||
|
||||
self.refresh_model_gbtn.click(
|
||||
fn=refresh_model,
|
||||
inputs=[self.output_model_name],
|
||||
outputs=[
|
||||
self.log_message,
|
||||
self.export_log_panel,
|
||||
self.output_model_name,
|
||||
self.output_ckpt_name,
|
||||
self.eval_gallery
|
||||
],
|
||||
queue=False)
|
||||
self.refresh_model_gbtn.click(fn=refresh_model,
|
||||
inputs=[self.output_model_name],
|
||||
outputs=[
|
||||
self.log_message,
|
||||
self.export_log_panel,
|
||||
self.output_model_name,
|
||||
self.output_ckpt_name,
|
||||
self.eval_gallery
|
||||
],
|
||||
queue=False)
|
||||
|
||||
def delete_model(model_name):
|
||||
index = 0
|
||||
@@ -306,7 +304,7 @@ class ModelUI(UIBase):
|
||||
if os.path.exists(params_path):
|
||||
params_info = json.loads(open(params_path).read())
|
||||
assert params_info['work_name'] == output_model
|
||||
# base_model = params_info['base_model']
|
||||
base_model = params_info['base_model']
|
||||
base_model_revision = params_info['base_model_revision']
|
||||
tuner_name = params_info['tuner_name']
|
||||
model_path = os.path.join(self.work_dir, output_model,
|
||||
@@ -341,18 +339,32 @@ class ModelUI(UIBase):
|
||||
f'meta_{output_ckpt_name}.yaml')
|
||||
output_model = output_model + '@' + output_ckpt_name
|
||||
tuner_dict = {
|
||||
'NAME': output_model,
|
||||
'NAME_ZH': output_model,
|
||||
'NAME':
|
||||
output_model,
|
||||
'NAME_ZH':
|
||||
output_model,
|
||||
# 'BASE_MODEL': base_model,
|
||||
'BASE_MODEL': base_model_revision,
|
||||
'TUNER_TYPE': tuner_name,
|
||||
'DESCRIPTION': '',
|
||||
'MODEL_PATH': model_path,
|
||||
'IMAGE_PATH': image_path,
|
||||
'PROMPT_EXAMPLE': eval_prompts,
|
||||
'SOURCE': 'self_train',
|
||||
'CKPT_NAME': output_ckpt_name,
|
||||
'PARAMS': params_info
|
||||
'BASE_MODEL':
|
||||
base_model_revision,
|
||||
'TUNER_TYPE':
|
||||
tuner_name,
|
||||
'DESCRIPTION':
|
||||
'',
|
||||
'MODEL_PATH':
|
||||
model_path,
|
||||
'IMAGE_PATH':
|
||||
image_path,
|
||||
'PROMPT_EXAMPLE':
|
||||
eval_prompts,
|
||||
'SOURCE':
|
||||
'self_train',
|
||||
'CKPT_NAME':
|
||||
output_ckpt_name,
|
||||
'PARAMS':
|
||||
params_info,
|
||||
'IS_SHARE':
|
||||
self.BASE_CFG_VALUE[base_model][base_model_revision]
|
||||
['is_share']
|
||||
}
|
||||
tuner_cfg = Config(cfg_dict=tuner_dict, load=False)
|
||||
|
||||
|
||||
@@ -47,6 +47,32 @@ def get_work_name(model, version, tuner):
|
||||
[str(random.randint(1, 10)) for i in range(3)])
|
||||
|
||||
|
||||
def judge_tuner_visible(tuner_name):
|
||||
lora_visible = tuner_name in ['LORA', 'TEXT_LORA']
|
||||
text_lora_visible = tuner_name in ['TEXT_SCE']
|
||||
sce_visible = tuner_name in ['SCE', 'TEXT_SCE']
|
||||
return lora_visible, text_lora_visible, sce_visible
|
||||
|
||||
|
||||
def update_tuner_cfg(tuner_name, tuner_cfg, **kwargs):
|
||||
update_info = {}
|
||||
if tuner_name in ['LORA', 'TEXT_LORA']:
|
||||
update_info = {
|
||||
'LORA_ALPHA': kwargs['lora_alpha'],
|
||||
'R': kwargs['lora_rank']
|
||||
}
|
||||
elif tuner_name == 'SCE' or (tuner_name == 'TEXT_SCE'
|
||||
and tuner_cfg['NAME'] == 'SwiftSCETuning'):
|
||||
update_info = {'DOWN_RATIO': kwargs['sce_ratio']}
|
||||
elif tuner_name == 'TEXT_SCE' and tuner_cfg['NAME'] == 'SwiftLoRA':
|
||||
update_info = {
|
||||
'LORA_ALPHA': kwargs['text_lora_alpha'],
|
||||
'R': kwargs['text_lora_rank']
|
||||
}
|
||||
tuner_cfg.update(update_info)
|
||||
return tuner_cfg
|
||||
|
||||
|
||||
class TrainerUI(UIBase):
|
||||
def __init__(self, cfg, all_cfg_value, is_debug=False, language='en'):
|
||||
self.BASE_CFG_VALUE = all_cfg_value
|
||||
@@ -81,6 +107,11 @@ class TrainerUI(UIBase):
|
||||
value=self.component_names.data_type_value,
|
||||
label=self.component_names.data_type_name,
|
||||
interactive=True)
|
||||
self.ori_data_name = gr.Textbox(
|
||||
label=self.component_names.ori_data_name,
|
||||
max_lines=1,
|
||||
placeholder=self.component_names.ori_data_name,
|
||||
interactive=True)
|
||||
self.ms_data_name = gr.Textbox(
|
||||
label=' or '.join(
|
||||
self.component_names.data_type_choices),
|
||||
@@ -130,7 +161,51 @@ class TrainerUI(UIBase):
|
||||
interactive=True)
|
||||
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
with gr.Accordion(
|
||||
label=self.component_names.tuner_param,
|
||||
open=False):
|
||||
if 'TUNER' in self.para_data:
|
||||
lora_visible, text_lora_visible, sce_visible = judge_tuner_visible(
|
||||
self.para_data['TUNER'])
|
||||
else:
|
||||
lora_visible, text_lora_visible, sce_visible = False, False, False
|
||||
with gr.Row(visible=lora_visible
|
||||
) as self.lora_param:
|
||||
self.lora_alpha = gr.Number(
|
||||
label='LoRA Alpha',
|
||||
value=self.para_data.get(
|
||||
'lora_alpha', 256),
|
||||
interactive=True)
|
||||
self.lora_rank = gr.Number(
|
||||
label='LoRA Rank',
|
||||
value=self.para_data.get(
|
||||
'lora_rank', 256),
|
||||
interactive=True)
|
||||
with gr.Row(visible=text_lora_visible
|
||||
) as self.text_lora_param:
|
||||
self.text_lora_alpha = gr.Number(
|
||||
label='Text LoRA Alpha',
|
||||
value=self.para_data.get(
|
||||
'text_lora_alpha', 256),
|
||||
interactive=True)
|
||||
self.text_lora_rank = gr.Number(
|
||||
label='Text LoRA Rank',
|
||||
value=self.para_data.get(
|
||||
'text_lora_rank', 256),
|
||||
interactive=True)
|
||||
with gr.Row(
|
||||
visible=sce_visible) as self.sce_param:
|
||||
self.sce_ratio = gr.Slider(
|
||||
label='SCE Ratio',
|
||||
minimum=0.2,
|
||||
maximum=2.0,
|
||||
step=0.1,
|
||||
value=self.para_data.get(
|
||||
'sce_ratio', 1.0),
|
||||
interactive=True)
|
||||
|
||||
with gr.Row():
|
||||
with gr.Column(scale=2, min_width=0):
|
||||
self.resolution_height = gr.Dropdown(
|
||||
choices=list(self.h_level_dict.keys()),
|
||||
value=self.para_data.get(
|
||||
@@ -140,7 +215,7 @@ class TrainerUI(UIBase):
|
||||
allow_custom_value=True,
|
||||
interactive=True)
|
||||
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
with gr.Column(scale=2, min_width=0):
|
||||
self.resolution_width = gr.Dropdown(
|
||||
choices=self.h_level_dict[
|
||||
self.resolution_height.value],
|
||||
@@ -151,6 +226,46 @@ class TrainerUI(UIBase):
|
||||
allow_custom_value=True,
|
||||
interactive=True)
|
||||
|
||||
with gr.Row():
|
||||
with gr.Accordion(label=self.component_names.
|
||||
resolution_param,
|
||||
open=False):
|
||||
self.enable_resolution_bucket = gr.Checkbox(
|
||||
value=False,
|
||||
container=True,
|
||||
interactive=True,
|
||||
label=self.component_names.
|
||||
enable_resolution_bucket)
|
||||
with gr.Column(
|
||||
visible=self.enable_resolution_bucket.
|
||||
value) as self.resolution_bucket_param:
|
||||
with gr.Row():
|
||||
self.min_bucket_resolution = gr.Number(
|
||||
label=self.component_names.
|
||||
min_bucket_resolution,
|
||||
value=self.para_data.get(
|
||||
'min_bucket_resolution', 256),
|
||||
interactive=True)
|
||||
self.max_bucket_resolution = gr.Number(
|
||||
label=self.component_names.
|
||||
max_bucket_resolution,
|
||||
value=self.para_data.get(
|
||||
'max_bucket_resolution', 1024),
|
||||
interactive=True)
|
||||
with gr.Row():
|
||||
self.bucket_resolution_steps = gr.Number(
|
||||
label=self.component_names.
|
||||
bucket_resolution_steps,
|
||||
value=self.para_data.get(
|
||||
'bucket_resolution_steps', 64),
|
||||
interactive=True)
|
||||
self.bucket_no_upscale = gr.Checkbox(
|
||||
value=False,
|
||||
container=True,
|
||||
interactive=True,
|
||||
label=self.component_names.
|
||||
bucket_no_upscale)
|
||||
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.train_epoch = gr.Number(
|
||||
@@ -232,7 +347,7 @@ class TrainerUI(UIBase):
|
||||
[
|
||||
self.component_names.data_type_choices[0],
|
||||
'',
|
||||
'https://modelscope.cn/api/v1/models/damo/scepter/repo?Revision=master&FilePath=datasets/3D_example_csv.zip', # noqa
|
||||
'https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=datasets/3D_example_csv.zip', # noqa
|
||||
''
|
||||
],
|
||||
[
|
||||
@@ -373,6 +488,8 @@ class TrainerUI(UIBase):
|
||||
ret_data = get_values_by_model_version_tuner(
|
||||
self.BASE_CFG_VALUE, base_model, base_model_revision,
|
||||
tuner_name)
|
||||
lora_visible, text_lora_visible, sce_visible = judge_tuner_visible(
|
||||
tuner_name)
|
||||
return ret_data.get('EPOCHS', 10), \
|
||||
ret_data.get('LEARNING_RATE', 0.0001), \
|
||||
ret_data.get('SAVE_INTERVAL', 10), \
|
||||
@@ -383,7 +500,10 @@ class TrainerUI(UIBase):
|
||||
interactive=True), \
|
||||
gr.Dropdown(value=ret_data.get('RESOLUTION', 1024)[1],
|
||||
choices=self.h_level_dict[ret_data.get('RESOLUTION', 1024)[0]],
|
||||
interactive=True)
|
||||
interactive=True), \
|
||||
gr.Row.update(visible=lora_visible), \
|
||||
gr.Row.update(visible=text_lora_visible), \
|
||||
gr.Row.update(visible=sce_visible)
|
||||
|
||||
#
|
||||
self.tuner_name.change(fn=change_train_value_by_model_version_tuner,
|
||||
@@ -395,7 +515,8 @@ class TrainerUI(UIBase):
|
||||
self.train_epoch, self.learning_rate,
|
||||
self.save_interval, self.train_batch_size,
|
||||
self.prompt_prefix, self.resolution_height,
|
||||
self.resolution_width
|
||||
self.resolution_width, self.lora_param,
|
||||
self.text_lora_param, self.sce_param
|
||||
],
|
||||
queue=False)
|
||||
|
||||
@@ -411,13 +532,36 @@ class TrainerUI(UIBase):
|
||||
outputs=[self.resolution_width],
|
||||
queue=False)
|
||||
|
||||
def run_train(work_name, data_type, ms_data_space, ms_data_name,
|
||||
ms_data_subname, base_model, base_model_revision,
|
||||
tuner_name, resolution_height, resolution_width,
|
||||
train_epoch, learning_rate, save_interval,
|
||||
train_batch_size, prompt_prefix, replace_keywords,
|
||||
push_to_hub, eval_prompts):
|
||||
def change_resolution_bucket(evt: gr.SelectData):
|
||||
is_selected = evt.selected
|
||||
return (
|
||||
gr.update(
|
||||
label=self.component_names.resolution_height_max if
|
||||
is_selected else self.component_names.resolution_height),
|
||||
gr.update(
|
||||
label=self.component_names.resolution_width_max
|
||||
if is_selected else self.component_names.resolution_width),
|
||||
gr.update(visible=is_selected))
|
||||
|
||||
self.enable_resolution_bucket.select(change_resolution_bucket,
|
||||
inputs=[],
|
||||
outputs=[
|
||||
self.resolution_height,
|
||||
self.resolution_width,
|
||||
self.resolution_bucket_param
|
||||
],
|
||||
queue=False)
|
||||
|
||||
def run_train(work_name, data_type, ori_data_name, ms_data_space,
|
||||
ms_data_name, ms_data_subname, base_model,
|
||||
base_model_revision, tuner_name, resolution_height,
|
||||
resolution_width, train_epoch, learning_rate,
|
||||
save_interval, train_batch_size, prompt_prefix,
|
||||
replace_keywords, push_to_hub, eval_prompts, lora_alpha,
|
||||
lora_rank, text_lora_alpha, text_lora_rank, sce_ratio,
|
||||
enable_resolution_bucket, min_bucket_resolution,
|
||||
max_bucket_resolution, bucket_resolution_steps,
|
||||
bucket_no_upscale):
|
||||
# Check Cuda
|
||||
if not torch.cuda.is_available() and not self.is_debug:
|
||||
raise gr.Error(self.component_names.training_err1)
|
||||
@@ -529,6 +673,49 @@ class TrainerUI(UIBase):
|
||||
int(resolution_height),
|
||||
int(resolution_width)
|
||||
]
|
||||
if enable_resolution_bucket:
|
||||
local_data_dir = data_cfg['MS_DATASET_NAME']
|
||||
local_file_list = os.path.join(local_data_dir,
|
||||
'file.txt')
|
||||
data_num = sum(1 for line in open(local_file_list))
|
||||
if os.path.exists(local_data_dir) and os.path.exists(
|
||||
local_file_list):
|
||||
data_cfg.update({
|
||||
'NAME': 'ImageTextPairDataset',
|
||||
'SAMPLER': {
|
||||
'NAME':
|
||||
'ResolutionBatchSampler',
|
||||
'DATA_FILE':
|
||||
local_file_list,
|
||||
'FIELDS':
|
||||
['img_path', 'width', 'height', 'prompt'],
|
||||
'DELIMITER':
|
||||
'#;#',
|
||||
'MAX_RESO': [
|
||||
int(resolution_width),
|
||||
int(resolution_height)
|
||||
],
|
||||
'MIN_BUCKET_RESO':
|
||||
int(min_bucket_resolution),
|
||||
'MAX_BUCKET_RESO':
|
||||
int(max_bucket_resolution),
|
||||
'BUCKET_RESO_STEPS':
|
||||
int(bucket_resolution_steps),
|
||||
'BUCKET_NO_UPSCALE':
|
||||
bucket_no_upscale
|
||||
},
|
||||
'DATA_NUM': data_num
|
||||
})
|
||||
for trans in data_cfg['TRANSFORMS']:
|
||||
if trans['NAME'] == 'Select':
|
||||
trans['META_KEYS'] = [
|
||||
'img_path', 'image_size'
|
||||
]
|
||||
else:
|
||||
raise Exception(
|
||||
'Cannot find right data format for resolution_bucket'
|
||||
)
|
||||
|
||||
return data_cfg
|
||||
|
||||
def prepare_eval_data(data_cfg):
|
||||
@@ -565,7 +752,6 @@ class TrainerUI(UIBase):
|
||||
v[c_k] = current_val
|
||||
current_val = v
|
||||
cfg = current_val
|
||||
|
||||
# update config
|
||||
cfg['SOLVER']['WORK_DIR'] = work_dir
|
||||
cfg['SOLVER']['OPTIMIZER']['LEARNING_RATE'] = float(
|
||||
@@ -574,11 +760,25 @@ class TrainerUI(UIBase):
|
||||
cfg['SOLVER']['TRAIN_DATA']['BATCH_SIZE'] = int(
|
||||
train_batch_size)
|
||||
if 'TUNER' in cfg['SOLVER']:
|
||||
cfg['SOLVER']['TUNER'] = current_model_info['tuner_para'][
|
||||
tuner_cfg_list = current_model_info['tuner_para'][
|
||||
tuner_name] if isinstance(
|
||||
current_model_info['tuner_para'],
|
||||
dict) and tuner_name in current_model_info[
|
||||
'tuner_para'] else None
|
||||
if tuner_cfg_list is not None:
|
||||
tuner_params = dict(
|
||||
lora_alpha=int(lora_alpha),
|
||||
lora_rank=int(lora_rank),
|
||||
text_lora_alpha=int(text_lora_alpha),
|
||||
text_lora_rank=int(text_lora_rank),
|
||||
sce_ratio=sce_ratio)
|
||||
tuner_cfg_list = [
|
||||
update_tuner_cfg(tuner_name, tuner_cfg,
|
||||
**tuner_params)
|
||||
for tuner_cfg in tuner_cfg_list
|
||||
]
|
||||
cfg['SOLVER']['TUNER'] = tuner_cfg_list
|
||||
|
||||
cfg['SOLVER']['TRAIN_DATA'] = prepare_train_data(
|
||||
cfg['SOLVER']['TRAIN_DATA'])
|
||||
if eval_prompts is not None and len(eval_prompts) > 0:
|
||||
@@ -586,7 +786,8 @@ class TrainerUI(UIBase):
|
||||
cfg['SOLVER']['EVAL_DATA'])
|
||||
else:
|
||||
cfg['SOLVER'].pop('EVAL_DATA')
|
||||
if 'SAMPLE_ARGS' in cfg['SOLVER']:
|
||||
if 'SAMPLE_ARGS' in cfg[
|
||||
'SOLVER'] and not enable_resolution_bucket:
|
||||
cfg['SOLVER']['SAMPLE_ARGS']['IMAGE_SIZE'] = [
|
||||
int(resolution_height),
|
||||
int(resolution_width)
|
||||
@@ -610,10 +811,17 @@ class TrainerUI(UIBase):
|
||||
return cfg_file
|
||||
|
||||
before_kill_inference = self.trainer_ins.check_memory()
|
||||
for k, v in manager.inference.pipe_manager.pipeline_level_modules.items(
|
||||
):
|
||||
if hasattr(v, 'dynamic_unload'):
|
||||
v.dynamic_unload(name='all')
|
||||
if hasattr(manager, 'inference'):
|
||||
for k, v in manager.inference.pipe_manager.pipeline_level_modules.items(
|
||||
):
|
||||
if hasattr(v, 'dynamic_unload'):
|
||||
v.dynamic_unload(name='all')
|
||||
if (hasattr(manager, 'preprocess') and hasattr(
|
||||
manager.preprocess.dataset_gallery.processors_manager,
|
||||
'dynamic_unload')):
|
||||
manager.preprocess.dataset_gallery.processors_manager.dynamic_unload(
|
||||
)
|
||||
|
||||
after_kill_inference = self.trainer_ins.check_memory()
|
||||
message = f'GPU info: {before_kill_inference}. \n\n'
|
||||
message += f'After unloading inference models, the GPU info: {after_kill_inference}. \n\n'
|
||||
@@ -629,13 +837,18 @@ class TrainerUI(UIBase):
|
||||
self.training_button.click(
|
||||
run_train,
|
||||
inputs=[
|
||||
self.work_name, self.data_type, self.ms_data_space,
|
||||
self.ms_data_name, self.ms_data_subname, self.base_model,
|
||||
self.base_model_revision, self.tuner_name,
|
||||
self.work_name, self.data_type, self.ori_data_name,
|
||||
self.ms_data_space, self.ms_data_name, self.ms_data_subname,
|
||||
self.base_model, self.base_model_revision, self.tuner_name,
|
||||
self.resolution_height, self.resolution_width,
|
||||
self.train_epoch, self.learning_rate, self.save_interval,
|
||||
self.train_batch_size, self.prompt_prefix,
|
||||
self.replace_keywords, self.push_to_hub, self.eval_prompts
|
||||
self.replace_keywords, self.push_to_hub, self.eval_prompts,
|
||||
self.lora_alpha, self.lora_rank, self.text_lora_alpha,
|
||||
self.text_lora_rank, self.sce_ratio,
|
||||
self.enable_resolution_bucket, self.min_bucket_resolution,
|
||||
self.max_bucket_resolution, self.bucket_resolution_steps,
|
||||
self.bucket_no_upscale
|
||||
],
|
||||
outputs=[inference_ui.output_model_name],
|
||||
queue=True)
|
||||
|
||||
@@ -96,7 +96,8 @@ def get_all_config(config_root, global_meta):
|
||||
'inference_para': inference_paras,
|
||||
'tuner_type': tuner_type,
|
||||
'tuner_para': tuner_para,
|
||||
'modify_para': modify_para
|
||||
'modify_para': modify_para,
|
||||
'is_share': 'IS_SHARE' in meta_cfg and meta_cfg['IS_SHARE']
|
||||
}
|
||||
if 'CONTROL_PARAS' in meta_cfg:
|
||||
control_type, control_para = build_meta_index_control(
|
||||
|
||||
@@ -4,6 +4,7 @@ import os
|
||||
from collections import OrderedDict
|
||||
|
||||
import gradio as gr
|
||||
from swift import push_to_hub
|
||||
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.file_system import FS
|
||||
@@ -19,6 +20,8 @@ from scepter.studio.utils.uibase import UIBase
|
||||
class BrowserUI(UIBase):
|
||||
def __init__(self, cfg, language='en'):
|
||||
self.work_dir = cfg.WORK_DIR
|
||||
self.train_dir = cfg.SELF_TRAIN_DIR
|
||||
self.base_model_tuner_methods = cfg.BASE_MODEL_VERSION
|
||||
self.yaml = os.path.join(self.work_dir, cfg.TUNER_LIST_YAML)
|
||||
if not FS.exists(self.yaml):
|
||||
self.saved_tuners = []
|
||||
@@ -31,6 +34,8 @@ class BrowserUI(UIBase):
|
||||
self.saved_tuners_to_category()
|
||||
self.component_names = TunerManagerNames(language)
|
||||
self.language = language
|
||||
self.export_folder = os.path.join(self.work_dir, cfg.EXPORT_DIR)
|
||||
self.readme_file = cfg.README_EN if self.language == 'en' else cfg.README_ZH
|
||||
|
||||
def saved_tuners_to_category(self):
|
||||
self.saved_tuners_category = OrderedDict()
|
||||
@@ -72,14 +77,14 @@ class BrowserUI(UIBase):
|
||||
self.diffusion_models = gr.Dropdown(
|
||||
label=self.component_names.base_models,
|
||||
choices=diffusion_models_choice,
|
||||
value=diffusion_model,
|
||||
value=None,
|
||||
multiselect=False,
|
||||
interactive=True)
|
||||
with gr.Column(scale=4, min_width=0):
|
||||
self.tuner_models = gr.Dropdown(
|
||||
label=self.component_names.tuner_name,
|
||||
choices=tuner_models_choice,
|
||||
value=tuner_model,
|
||||
value=None,
|
||||
multiselect=False,
|
||||
interactive=True)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
@@ -95,13 +100,147 @@ class BrowserUI(UIBase):
|
||||
elem_classes='type_row',
|
||||
elem_id='delete_button',
|
||||
visible=False)
|
||||
self.model_upload = gr.Button(
|
||||
label='ModelScope Upload',
|
||||
value=self.component_names.upload,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button',
|
||||
visible=True)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.refresh_button = gr.Button(
|
||||
label='Delete',
|
||||
label='Refresh',
|
||||
value=self.component_names.refresh_symbol,
|
||||
elem_classes='type_row',
|
||||
elem_id='refresh_button',
|
||||
visible=True)
|
||||
self.model_download = gr.Button(
|
||||
label='ModelScope Download',
|
||||
value=self.component_names.download,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button',
|
||||
visible=True)
|
||||
|
||||
with gr.Box(visible=False) as self.upload_setting:
|
||||
gr.Markdown(self.component_names.export_desc)
|
||||
with gr.Column(variant='panel'):
|
||||
with gr.Row(equal_height=True):
|
||||
with gr.Column(scale=9, min_width=0):
|
||||
self.import_src = gr.Dropdown(
|
||||
choices=['modelscope'],
|
||||
value='modelscope',
|
||||
label=None,
|
||||
show_label=False)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.ms_upload_close = gr.Button(
|
||||
label='Close MS Upload',
|
||||
value=self.component_names.close,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button')
|
||||
with gr.Row(equal_height=True):
|
||||
with gr.Column(scale=3, min_width=0):
|
||||
self.ms_sdk = gr.Text(
|
||||
label=self.component_names.ms_sdk,
|
||||
show_label=False,
|
||||
container=False,
|
||||
placeholder='ModelScope SDK Token',
|
||||
value='')
|
||||
with gr.Column(scale=3, min_width=0):
|
||||
self.ms_upload_username = gr.Text(
|
||||
label=self.component_names.ms_username,
|
||||
show_label=False,
|
||||
container=False,
|
||||
placeholder='ModelScope UserName',
|
||||
value='')
|
||||
with gr.Column(scale=3, min_width=0):
|
||||
self.model_private = gr.Checkbox(
|
||||
label=self.component_names.model_private,
|
||||
value=False)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.ms_upload_submit = gr.Button(
|
||||
label='Submit MS',
|
||||
value=self.component_names.ms_submit,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button')
|
||||
|
||||
with gr.Box(visible=False) as self.download_setting:
|
||||
gr.Markdown(self.component_names.import_desc)
|
||||
with gr.Column(variant='panel'):
|
||||
with gr.Row(equal_height=True):
|
||||
with gr.Column(scale=9, min_width=0):
|
||||
self.import_src = gr.Dropdown(
|
||||
choices=['modelscope', 'local'],
|
||||
value='modelscope',
|
||||
label=None,
|
||||
show_label=False)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.ms_download_close = gr.Button(
|
||||
label='Close MS Download',
|
||||
value=self.component_names.close,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button')
|
||||
with gr.Row(
|
||||
equal_height=True) as self.ms_import_setting:
|
||||
with gr.Column(scale=4.5, min_width=0):
|
||||
self.ms_modelid = gr.Text(
|
||||
label=self.component_names.ms_modelid,
|
||||
show_label=False,
|
||||
container=False,
|
||||
placeholder='ModelScope Model Path',
|
||||
value='')
|
||||
with gr.Column(scale=4.5, min_width=0):
|
||||
self.ms_download_username = gr.Text(
|
||||
label=self.component_names.ms_username,
|
||||
show_label=False,
|
||||
container=False,
|
||||
placeholder='ModelScope UserName',
|
||||
value='')
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.ms_download_submit = gr.Button(
|
||||
label='Submit MS',
|
||||
value=self.component_names.ms_submit,
|
||||
elem_classes='type_row',
|
||||
elem_id='save_button')
|
||||
with gr.Row(visible=False, equal_height=True
|
||||
) as self.local_import_setting:
|
||||
with gr.Column(scale=2, min_width=0):
|
||||
self.file_path = gr.File(
|
||||
label=self.component_names.zip_file,
|
||||
min_width=0,
|
||||
file_types=['.zip'],
|
||||
elem_classes='upload_zone')
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.upload_base_models = gr.Dropdown(
|
||||
label=self.component_names.ubase_model,
|
||||
choices=[
|
||||
base_model_version.BASE_MODEL
|
||||
for base_model_version in
|
||||
self.base_model_tuner_methods
|
||||
],
|
||||
value=self.base_model_tuner_methods[0].
|
||||
BASE_MODEL,
|
||||
multiselect=False,
|
||||
interactive=True)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.upload_tuner_type = gr.Dropdown(
|
||||
label=self.component_names.utuner_type,
|
||||
choices=self.base_model_tuner_methods[0].
|
||||
TUNER_TYPE,
|
||||
value=self.base_model_tuner_methods[0].
|
||||
TUNER_TYPE[0],
|
||||
multiselect=False,
|
||||
interactive=True)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.upload_tuner_name = gr.Text(
|
||||
label=self.component_names.utuner_name,
|
||||
show_label=False,
|
||||
container=False,
|
||||
placeholder='Upload Tuner Name',
|
||||
value='')
|
||||
self.local_upload_bt = gr.Button(
|
||||
label='Submit Local Model',
|
||||
value=self.component_names.ms_submit,
|
||||
elem_classes='type_row',
|
||||
elem_id='upload_button')
|
||||
|
||||
def check_new_name(self, tuner_name):
|
||||
if tuner_name.strip() == '':
|
||||
@@ -113,7 +252,8 @@ class BrowserUI(UIBase):
|
||||
return False, f"Tuner name '{tuner_name}' has been taken!"
|
||||
return True, 'legal'
|
||||
|
||||
def save_tuner(self, src_path, sub_dir, tuner_name, tuner_example):
|
||||
def save_tuner(self, src_path, sub_dir, tuner_name, tuner_desc,
|
||||
tuner_example, tuner_prompt_example):
|
||||
tar_path = os.path.join(self.work_dir, sub_dir)
|
||||
if not FS.exists(tar_path):
|
||||
FS.make_dir(tar_path)
|
||||
@@ -121,14 +261,74 @@ class BrowserUI(UIBase):
|
||||
|
||||
FS.put_dir_from_local_dir(src_path, tar_path)
|
||||
|
||||
# save image
|
||||
tuner_example_path = None
|
||||
if tuner_example is not None:
|
||||
from PIL import Image
|
||||
tuner_example_path = os.path.join(tar_path, f'{tuner_name}.jpg')
|
||||
tuner_example = Image.fromarray(tuner_example)
|
||||
with FS.put_to(tuner_example_path) as local_path:
|
||||
tuner_example.save(local_path)
|
||||
return tar_path, tuner_example_path
|
||||
tuner_example_path = os.path.join(tar_path, 'image.jpg')
|
||||
if not os.path.exists(tuner_example_path):
|
||||
tuner_example = Image.fromarray(tuner_example)
|
||||
with FS.put_to(tuner_example_path) as local_path:
|
||||
tuner_example.save(local_path)
|
||||
|
||||
# save param
|
||||
enable_share = True
|
||||
split_path = src_path.split('/')
|
||||
if split_path[-2] == 'checkpoints':
|
||||
src_dir = '/'.join(split_path[:-2])
|
||||
ckpt_name = split_path[-1]
|
||||
meta_read = f'{src_dir}/meta_{ckpt_name}.yaml'
|
||||
meta_save = f'{tar_path}/params.yaml'
|
||||
else:
|
||||
meta_read = f'{src_path}/params.yaml'
|
||||
meta_save = f'{tar_path}/params.yaml'
|
||||
|
||||
if os.path.exists(meta_read):
|
||||
meta = Config(cfg_file=meta_read)
|
||||
with FS.put_to(meta_save) as local_path:
|
||||
enable_share = Config.get_plain_cfg(meta.get('IS_SHARE', True))
|
||||
params = Config.get_plain_cfg(meta.get('PARAMS', {}))
|
||||
params['work_dir'] = ''
|
||||
params['work_name'] = ''
|
||||
save_yaml({'PARAMS': params}, local_path)
|
||||
|
||||
# rewrite readme
|
||||
with open(self.readme_file, 'r') as f:
|
||||
rc = f.read()
|
||||
|
||||
rc = rc.replace(r'{MODEL_NAME}', tuner_name)
|
||||
rc = rc.replace(r'{MODEL_DESCRIPTION}',
|
||||
tuner_desc if len(tuner_desc) > 0 else tuner_name)
|
||||
rc = rc.replace(r'{EVAL_PROMPT}', tuner_prompt_example)
|
||||
rc = rc.replace(r'{IMAGE_PATH}', './image.jpg')
|
||||
rc = rc.replace(r'{BASE_MODEL}',
|
||||
meta['PARAMS'].get('base_model_revision', ''))
|
||||
rc = rc.replace(r'{TUNER_TYPE}',
|
||||
meta['PARAMS'].get('tuner_name', ''))
|
||||
rc = rc.replace(r'{TRAIN_BATCH_SIZE}',
|
||||
str(meta['PARAMS'].get('train_batch_size', '')))
|
||||
rc = rc.replace(r'{TRAIN_EPOCH}',
|
||||
str(meta['PARAMS'].get('train_epoch', '')))
|
||||
rc = rc.replace(r'{LEARNING_RATE}',
|
||||
str(meta['PARAMS'].get('learning_rate', '')))
|
||||
rc = rc.replace(r'{HEIGHT}',
|
||||
str(meta['PARAMS'].get('resolution_height', '')))
|
||||
rc = rc.replace(r'{WIDTH}',
|
||||
str(meta['PARAMS'].get('resolution_width', '')))
|
||||
rc = rc.replace(r'{DATA_TYPE}',
|
||||
meta['PARAMS'].get('data_type', ''))
|
||||
rc = rc.replace(r'{MS_DATA_SPACE}',
|
||||
meta['PARAMS'].get('ms_data_space', ''))
|
||||
rc = rc.replace(r'{MS_DATA_NAME}',
|
||||
meta['PARAMS'].get('ms_data_name', ''))
|
||||
rc = rc.replace(r'{MS_DATA_SUBNAME}',
|
||||
meta['PARAMS'].get('ms_data_subname', ''))
|
||||
|
||||
with FS.put_to(os.path.join(tar_path, 'README.md')) as local_path:
|
||||
with open(local_path, 'w') as f:
|
||||
f.write(rc)
|
||||
|
||||
return tar_path, tuner_example_path, enable_share
|
||||
|
||||
def add_tuner(self, new_tuner, manager, now_diffusion_model):
|
||||
self.saved_tuners.append(new_tuner)
|
||||
@@ -196,6 +396,12 @@ class BrowserUI(UIBase):
|
||||
|
||||
return custom_tuner_choices
|
||||
|
||||
def update_tuner_info(self, base_model, tuner_name, update_items):
|
||||
self.saved_tuners_category[base_model][tuner_name].update(update_items)
|
||||
self.category_to_saved_tuners()
|
||||
with FS.put_to(self.yaml) as local_path:
|
||||
save_yaml({'TUNERS': self.saved_tuners}, local_path)
|
||||
|
||||
def set_callbacks(self, manager, info_ui):
|
||||
def refresh_browser():
|
||||
diffusion_models_choice, diffusion_model, tuner_models_choice, tuner_model = self.get_choices_and_values(
|
||||
@@ -222,10 +428,13 @@ class BrowserUI(UIBase):
|
||||
queue=True)
|
||||
|
||||
def tuner_model_change(tuner_model, diffusion_model):
|
||||
tuner_info = {}
|
||||
if tuner_model is not None:
|
||||
tuner_info = self.saved_tuners_category[diffusion_model][
|
||||
tuner_model]
|
||||
if tuner_model is None:
|
||||
# fix refresh bug
|
||||
return (gr.Text(), gr.Text(), gr.Text(), gr.Text(), gr.Text(),
|
||||
gr.Image(), gr.Text(), gr.Text())
|
||||
|
||||
tuner_info = self.saved_tuners_category[diffusion_model][
|
||||
tuner_model]
|
||||
image_path = tuner_info.get('IMAGE_PATH', None)
|
||||
if image_path is not None:
|
||||
image_path = FS.get_from(image_path)
|
||||
@@ -237,7 +446,9 @@ class BrowserUI(UIBase):
|
||||
gr.Text(value=tuner_info.get('DESCRIPTION', ''),
|
||||
interactive=True), gr.Image(value=image_path),
|
||||
gr.Text(value=tuner_info.get('PROMPT_EXAMPLE', ''),
|
||||
interactive=True))
|
||||
interactive=True),
|
||||
gr.Text(value=tuner_info.get('MODELSCOPE_URL', ''),
|
||||
interactive=False))
|
||||
|
||||
self.tuner_models.change(
|
||||
tuner_model_change,
|
||||
@@ -245,44 +456,61 @@ class BrowserUI(UIBase):
|
||||
outputs=[
|
||||
info_ui.tuner_name, info_ui.new_name, info_ui.tuner_type,
|
||||
info_ui.base_model, info_ui.tuner_desc, info_ui.tuner_example,
|
||||
info_ui.tuner_prompt_example
|
||||
info_ui.tuner_prompt_example, info_ui.ms_url
|
||||
],
|
||||
queue=False)
|
||||
|
||||
def save_tuner(tuner_name, new_name, tuner_desc, tuner_example,
|
||||
tuner_prompt_example, now_diffusion_model, info_path):
|
||||
tuner_prompt_example, base_model, tuner_type):
|
||||
is_legal, msg = self.check_new_name(new_name)
|
||||
if not is_legal:
|
||||
gr.Info('Save failed because ' + msg)
|
||||
return (gr.Dropdown(), gr.Text(), gr.Dropdown())
|
||||
return (gr.Dropdown(), gr.Dropdown(), gr.Text(), gr.Dropdown())
|
||||
|
||||
info = Config(cfg_file=info_path)
|
||||
model_dir = info.MODEL_PATH
|
||||
sub_dir = f'{info.BASE_MODEL}-{info.TUNER_TYPE}'
|
||||
model_dir, tuner_example = self.save_tuner(model_dir, sub_dir,
|
||||
new_name, tuner_example)
|
||||
sub_dir = f'{base_model}-{tuner_type}'
|
||||
if os.path.exists(os.path.join(self.train_dir, '@'.join(tuner_name.split('@')[:-1]))) \
|
||||
and len(tuner_name.split('@')[:-1]) > 0:
|
||||
steps = tuner_name.split('@')[-1]
|
||||
model_dir = os.path.join(self.train_dir,
|
||||
'@'.join(tuner_name.split('@')[:-1]),
|
||||
'checkpoints', steps)
|
||||
else:
|
||||
model_dir = self.saved_tuners_category.get(sub_dir, {}).get(
|
||||
tuner_name, {}).get('MODEL_PATH', '')
|
||||
if model_dir == '':
|
||||
gr.Error(self.component_names.model_err4 + tuner_name)
|
||||
|
||||
model_dir, tuner_example, enable_share = self.save_tuner(
|
||||
model_dir, sub_dir, new_name, tuner_desc, tuner_example,
|
||||
tuner_prompt_example)
|
||||
# config info update
|
||||
new_tuner = {
|
||||
'NAME': new_name,
|
||||
'NAME_ZH': new_name,
|
||||
'SOURCE': 'self_train',
|
||||
'DESCRIPTION': tuner_desc,
|
||||
'BASE_MODEL': info.BASE_MODEL,
|
||||
'BASE_MODEL': base_model,
|
||||
'MODEL_PATH': model_dir,
|
||||
'IMAGE_PATH': tuner_example,
|
||||
'TUNER_TYPE': info.TUNER_TYPE,
|
||||
'PROMPT_EXAMPLE': tuner_prompt_example
|
||||
'TUNER_TYPE': tuner_type,
|
||||
'PROMPT_EXAMPLE': tuner_prompt_example,
|
||||
'ENABLE_SHARE': enable_share,
|
||||
}
|
||||
pipeline_level_modules = manager.inference.model_manage_ui.pipe_manager.pipeline_level_modules
|
||||
if new_tuner['BASE_MODEL'] not in pipeline_level_modules:
|
||||
gr.Error(self.component_names.model_err3 +
|
||||
new_tuner['BASE_MODEL'])
|
||||
pipeline_ins = pipeline_level_modules[new_tuner['BASE_MODEL']]
|
||||
now_diffusion_model = f"{new_tuner['BASE_MODEL']}_{pipeline_ins.diffusion_model['name']}"
|
||||
|
||||
custom_tuner_choices = self.add_tuner(new_tuner, manager,
|
||||
now_diffusion_model)
|
||||
|
||||
return (gr.Dropdown(choices=list(
|
||||
self.saved_tuners_category.keys()),
|
||||
value=sub_dir),
|
||||
gr.Dropdown(choices=list(
|
||||
return (gr.update(choices=list(self.saved_tuners_category.keys()),
|
||||
value=sub_dir),
|
||||
gr.update(choices=list(
|
||||
self.saved_tuners_category.get(sub_dir, {}).keys()),
|
||||
value=new_name), gr.Text(value=new_name),
|
||||
value=new_name), gr.Text(value=new_name),
|
||||
gr.Dropdown(choices=custom_tuner_choices))
|
||||
|
||||
self.save_button.click(
|
||||
@@ -290,8 +518,7 @@ class BrowserUI(UIBase):
|
||||
inputs=[
|
||||
info_ui.tuner_name, info_ui.new_name, info_ui.tuner_desc,
|
||||
info_ui.tuner_example, info_ui.tuner_prompt_example,
|
||||
manager.inference.model_manage_ui.diffusion_model,
|
||||
manager.inference.infer_info
|
||||
info_ui.base_model, info_ui.tuner_type
|
||||
],
|
||||
outputs=[
|
||||
self.diffusion_models, self.tuner_models, info_ui.tuner_name,
|
||||
@@ -321,3 +548,263 @@ class BrowserUI(UIBase):
|
||||
manager.inference.tuner_ui.custom_tuner_model
|
||||
],
|
||||
queue=True)
|
||||
|
||||
def change_visible():
|
||||
return gr.update(visible=True), gr.update(visible=False)
|
||||
|
||||
self.model_upload.click(
|
||||
fn=change_visible,
|
||||
inputs=[],
|
||||
outputs=[self.upload_setting, self.download_setting],
|
||||
queue=False)
|
||||
self.model_download.click(
|
||||
fn=change_visible,
|
||||
inputs=[],
|
||||
outputs=[self.download_setting, self.upload_setting],
|
||||
queue=False)
|
||||
|
||||
def change_invisible():
|
||||
return gr.update(visible=False)
|
||||
|
||||
self.ms_upload_close.click(fn=change_invisible,
|
||||
inputs=[],
|
||||
outputs=[self.upload_setting],
|
||||
queue=False)
|
||||
self.ms_download_close.click(fn=change_invisible,
|
||||
inputs=[],
|
||||
outputs=[self.download_setting],
|
||||
queue=False)
|
||||
|
||||
def change_import_source(import_src):
|
||||
if import_src == 'modelscope':
|
||||
ms_visible = True
|
||||
local_visible = False
|
||||
elif import_src == 'local':
|
||||
ms_visible = False
|
||||
local_visible = True
|
||||
return gr.update(visible=ms_visible), gr.update(
|
||||
visible=local_visible)
|
||||
|
||||
self.import_src.change(
|
||||
fn=change_import_source,
|
||||
inputs=[self.import_src],
|
||||
outputs=[self.ms_import_setting, self.local_import_setting],
|
||||
queue=False)
|
||||
|
||||
def push_to_modelscope(ms_sdk, username, private, base_model_name,
|
||||
tuner_model_name):
|
||||
tuner = self.saved_tuners_category[base_model_name][
|
||||
tuner_model_name]
|
||||
|
||||
enable_share = tuner.get('ENABLE_SHARE', True)
|
||||
if enable_share:
|
||||
repo_name = f'{username}/{tuner_model_name}'
|
||||
ckpt_path = tuner['MODEL_PATH']
|
||||
ms_url = f'https://www.modelscope.cn/models/{repo_name}'
|
||||
|
||||
with open(os.path.join(ckpt_path, 'README.md'), 'r') as f:
|
||||
rc = f.read()
|
||||
rc = rc.replace(r'{MODEL_URL}', ms_url)
|
||||
rc = rc.replace(r'{USER_NAME}', username)
|
||||
with open(os.path.join(ckpt_path, 'README.md'), 'w') as f:
|
||||
f.write(rc)
|
||||
|
||||
push_status = push_to_hub(repo_name,
|
||||
ckpt_path,
|
||||
token=ms_sdk,
|
||||
private=private)
|
||||
if push_status:
|
||||
update_items = {'MODELSCOPE_URL': ms_url}
|
||||
self.update_tuner_info(base_model_name, tuner_model_name,
|
||||
update_items)
|
||||
gr.Info(
|
||||
'The tuner model has been uploaded to ModelScope Successfully!'
|
||||
)
|
||||
return update_items['MODELSCOPE_URL']
|
||||
else:
|
||||
gr.Info(
|
||||
'Error: The model failed to be uploaded to ModelScope!'
|
||||
)
|
||||
return ''
|
||||
else:
|
||||
gr.Info(
|
||||
'Error: The model is not allowed to be shared to ModelScope!'
|
||||
)
|
||||
return ''
|
||||
|
||||
self.ms_upload_submit.click(fn=push_to_modelscope,
|
||||
inputs=[
|
||||
self.ms_sdk, self.ms_upload_username,
|
||||
self.model_private,
|
||||
self.diffusion_models,
|
||||
self.tuner_models
|
||||
],
|
||||
outputs=[info_ui.ms_url],
|
||||
queue=True)
|
||||
|
||||
def pull_from_modelscope(modelid, username):
|
||||
tar_path = os.path.join(self.work_dir, 'modelscope')
|
||||
src_path = f'ms://{username}/{modelid}'
|
||||
FS.get_dir_to_local_dir(src_path, tar_path)
|
||||
gr.Info(
|
||||
'The tuner model has been downloaded from ModelScope Successfully!'
|
||||
)
|
||||
|
||||
tar_path = f'{tar_path}/{username}/{modelid}'
|
||||
meta_file = f'{tar_path}/params.yaml'
|
||||
meta = Config(cfg_file=meta_file)
|
||||
base_model = meta['PARAMS']['base_model_revision']
|
||||
tuner_type = meta['PARAMS']['tuner_name']
|
||||
tuner_prompt_example = meta['PARAMS']['eval_prompts'][0]
|
||||
tuner_category = f'{base_model}-{tuner_type}'
|
||||
new_name = f'modelscope@{username}@{modelid}'
|
||||
new_tuner = {
|
||||
'NAME': new_name,
|
||||
'NAME_ZH': new_name,
|
||||
'SOURCE': 'modelscope',
|
||||
'DESCRIPTION': '',
|
||||
'BASE_MODEL': base_model,
|
||||
'MODEL_PATH': tar_path,
|
||||
'IMAGE_PATH': f'{tar_path}/image.jpg',
|
||||
'TUNER_TYPE': tuner_type,
|
||||
'PROMPT_EXAMPLE': tuner_prompt_example
|
||||
}
|
||||
pipeline_level_modules = manager.inference.model_manage_ui.pipe_manager.pipeline_level_modules
|
||||
if new_tuner['BASE_MODEL'] not in pipeline_level_modules:
|
||||
gr.Error(self.component_names.model_err3 +
|
||||
new_tuner['BASE_MODEL'])
|
||||
pipeline_ins = pipeline_level_modules[new_tuner['BASE_MODEL']]
|
||||
now_diffusion_model = f"{new_tuner['BASE_MODEL']}_{pipeline_ins.diffusion_model['name']}"
|
||||
|
||||
custom_tuner_choices = self.add_tuner(new_tuner, manager,
|
||||
now_diffusion_model)
|
||||
|
||||
update_items = {
|
||||
'MODELSCOPE_URL':
|
||||
f'https://www.modelscope.cn/models/{username}/{modelid}'
|
||||
}
|
||||
self.update_tuner_info(tuner_category,
|
||||
new_name,
|
||||
update_items=update_items)
|
||||
|
||||
return (gr.update(choices=list(self.saved_tuners_category.keys()),
|
||||
value=tuner_category),
|
||||
gr.update(choices=list(
|
||||
self.saved_tuners_category.get(tuner_category,
|
||||
{}).keys()),
|
||||
value=new_name), gr.Text(value=new_name),
|
||||
gr.Dropdown(choices=custom_tuner_choices))
|
||||
|
||||
self.ms_download_submit.click(
|
||||
fn=pull_from_modelscope,
|
||||
inputs=[
|
||||
self.ms_modelid,
|
||||
self.ms_download_username,
|
||||
],
|
||||
outputs=[
|
||||
self.diffusion_models, self.tuner_models, info_ui.tuner_name,
|
||||
manager.inference.tuner_ui.custom_tuner_model
|
||||
],
|
||||
queue=True)
|
||||
|
||||
def change_tuner_type_by_model_version(base_model_revision):
|
||||
for base_model_tuner in self.base_model_tuner_methods:
|
||||
if base_model_tuner.BASE_MODEL == base_model_revision:
|
||||
return gr.Dropdown(value=base_model_tuner.TUNER_TYPE[0],
|
||||
choices=base_model_tuner.TUNER_TYPE,
|
||||
interactive=True)
|
||||
return gr.Dropdown(value='', choices=[], interactive=True)
|
||||
|
||||
self.upload_base_models.change(fn=change_tuner_type_by_model_version,
|
||||
inputs=[self.upload_base_models],
|
||||
outputs=[self.upload_tuner_type],
|
||||
queue=False)
|
||||
|
||||
def upload_zip(file_path, tuner_name, base_model, tuner_type):
|
||||
sub_dir = f'{base_model}-{tuner_type}'
|
||||
save_file = os.path.join(self.work_dir, sub_dir,
|
||||
f'{tuner_name}.zip')
|
||||
model_dir = os.path.join(self.work_dir, sub_dir, tuner_name)
|
||||
with FS.put_to(save_file) as local_zip:
|
||||
res = os.popen(f"cp '{file_path.name}' '{local_zip}'")
|
||||
res = res.readlines()
|
||||
with FS.get_from(save_file) as local_path:
|
||||
res = os.popen(f"unzip -o '{local_path}' -d '{model_dir}'")
|
||||
res = res.readlines()
|
||||
if not os.path.exists(model_dir):
|
||||
raise gr.Error(f'解压{save_file}失败{str(res)}')
|
||||
# find meta.yaml
|
||||
if os.path.exists(os.path.join(model_dir, 'meta.yaml')):
|
||||
meta = Config(cfg_file=os.path.join(model_dir, 'meta.yaml'))
|
||||
tuner_desc = meta.get('DESCRIPTION', '')
|
||||
tuner_example = meta.get('IMAGE_PATH', None)
|
||||
if tuner_example is not None:
|
||||
tuner_example = os.path.join(
|
||||
model_dir, os.path.basename(tuner_example))
|
||||
if os.path.exists(tuner_example):
|
||||
from PIL import Image
|
||||
tuner_example_path = os.path.join(
|
||||
model_dir, 'image.jpg')
|
||||
if not os.path.exists(tuner_example_path):
|
||||
tuner_example = Image.open(tuner_example)
|
||||
with FS.put_to(tuner_example_path) as local_path:
|
||||
tuner_example.save(local_path)
|
||||
tuner_example = tuner_example_path
|
||||
else:
|
||||
tuner_example = None
|
||||
tuner_prompt_example = meta.get('PROMPT_EXAMPLE', '')
|
||||
else:
|
||||
tuner_desc = ''
|
||||
tuner_example = None
|
||||
tuner_prompt_example = ''
|
||||
|
||||
if not FS.exists(model_dir):
|
||||
raise gr.Error(
|
||||
f'{self.component_names.illegal_data_err1}{str(res)}')
|
||||
# config info update
|
||||
new_tuner = {
|
||||
'NAME': tuner_name,
|
||||
'NAME_ZH': tuner_name,
|
||||
'SOURCE': 'self_train',
|
||||
'DESCRIPTION': tuner_desc,
|
||||
'BASE_MODEL': base_model,
|
||||
'MODEL_PATH': model_dir,
|
||||
'IMAGE_PATH': tuner_example,
|
||||
'TUNER_TYPE': tuner_type,
|
||||
'PROMPT_EXAMPLE': tuner_prompt_example
|
||||
}
|
||||
pipeline_level_modules = manager.inference.model_manage_ui.pipe_manager.pipeline_level_modules
|
||||
if new_tuner['BASE_MODEL'] not in pipeline_level_modules:
|
||||
gr.Error(self.component_names.model_err3 +
|
||||
new_tuner['BASE_MODEL'])
|
||||
pipeline_ins = pipeline_level_modules[new_tuner['BASE_MODEL']]
|
||||
now_diffusion_model = f"{new_tuner['BASE_MODEL']}_{pipeline_ins.diffusion_model['name']}"
|
||||
|
||||
custom_tuner_choices = self.add_tuner(new_tuner, manager,
|
||||
now_diffusion_model)
|
||||
|
||||
return (gr.Dropdown(choices=list(
|
||||
self.saved_tuners_category.keys()),
|
||||
value=sub_dir),
|
||||
gr.Dropdown(choices=list(
|
||||
self.saved_tuners_category.get(sub_dir, {}).keys()),
|
||||
value=tuner_name), gr.Text(value=tuner_name),
|
||||
gr.Text(value=tuner_type), gr.Text(value=base_model),
|
||||
gr.Text(value=tuner_desc),
|
||||
gr.Text(value=tuner_prompt_example),
|
||||
gr.Image(value=tuner_example),
|
||||
gr.Dropdown(choices=custom_tuner_choices))
|
||||
|
||||
self.local_upload_bt.click(
|
||||
upload_zip,
|
||||
inputs=[
|
||||
self.file_path, self.upload_tuner_name,
|
||||
self.upload_base_models, self.upload_tuner_type
|
||||
],
|
||||
outputs=[
|
||||
self.diffusion_models, self.tuner_models, info_ui.tuner_name,
|
||||
info_ui.tuner_type, info_ui.base_model, info_ui.tuner_desc,
|
||||
info_ui.tuner_prompt_example, info_ui.tuner_example,
|
||||
manager.inference.tuner_ui.custom_tuner_model
|
||||
],
|
||||
queue=False)
|
||||
|
||||
@@ -8,8 +8,14 @@ class TunerManagerNames():
|
||||
self.save_symbol = '\U0001F4BE' # 💾
|
||||
self.delete_symbol = '\U0001f5d1' # 🗑️
|
||||
self.refresh_symbol = '\U0001f504' # 🔄
|
||||
self.upload = '\U0001F517' # 🔗
|
||||
self.download = '\U00002795' # ➕
|
||||
self.ms_submit = '\U00002714' # ✔️
|
||||
self.close = '\U00002716' # ✖️
|
||||
|
||||
if language == 'en':
|
||||
self.browser_block_name = 'Tuner Browser'
|
||||
self.browser_block_name = 'Tuner Browser ' \
|
||||
'(\U0001F4BE: Save; \U0001f504: Refresh; \U0001F517: Export; \U00002795: Import)'
|
||||
self.base_models = 'Base Model-Tuner Type'
|
||||
self.tuner_models = 'Tuner Name'
|
||||
self.info_block_name = 'Tuner Info'
|
||||
@@ -20,10 +26,31 @@ class TunerManagerNames():
|
||||
self.tuner_desc = 'Tuner Description'
|
||||
self.tuner_example = 'Results Example'
|
||||
self.tuner_prompt_example = 'Prompt Example'
|
||||
self.model_err3 = "Doesn't surpport this base model"
|
||||
self.model_err4 = \
|
||||
"This model maybe not finish training, because model doesn't exist. Please save model first."
|
||||
self.model_err5 = 'Model name not registered locally.'
|
||||
self.go_to_inference = 'Go To Inference'
|
||||
self.save = 'save changes'
|
||||
self.delete = 'Delete'
|
||||
self.ms_sdk = 'ModelScope API Token'
|
||||
self.ms_username = 'ModelScope User Name'
|
||||
self.model_private = 'Model Private'
|
||||
self.ms_modelid = 'ModelScope Model ID'
|
||||
self.ms_url = 'ModelScope Model Url'
|
||||
self.ms_model_path = 'Hub Model ID'
|
||||
self.export_file = 'Download Model'
|
||||
self.export_zip_err1 = 'export model failure'
|
||||
self.zip_file = 'upload model'
|
||||
self.utuner_name = 'Upload Tuner Name'
|
||||
self.ubase_model = 'Upload Base Model'
|
||||
self.utuner_type = 'Upload Tuner Type'
|
||||
self.illegal_data_err1 = 'Upload File Format Error(not .zip)'
|
||||
self.download_to_local = 'Download to Local'
|
||||
self.export_desc = 'Model Export To ModelScope (\U00002714: Submit; \U00002716: Close)'
|
||||
self.import_desc = 'Model Import From **ModelScope/Local** (\U00002714: Submit; \U00002716: Close)'
|
||||
elif language == 'zh':
|
||||
self.browser_block_name = '微调模型查找'
|
||||
self.browser_block_name = '微调模型查找 (\U0001F4BE: 保存; \U0001f504: 刷新; \U0001F517: 导出; \U00002795: 导入)'
|
||||
self.base_models = '基模型-微调类型'
|
||||
self.tuner_models = '微调模型名称'
|
||||
self.info_block_name = '微调模型详情'
|
||||
@@ -34,5 +61,25 @@ class TunerManagerNames():
|
||||
self.tuner_desc = '微调模型描述'
|
||||
self.tuner_example = '示例结果'
|
||||
self.tuner_prompt_example = '示例提示词'
|
||||
self.model_err3 = '不支持的基础模型'
|
||||
self.model_err4 = '模型可能没有训练完成或者模型不存在,请先保存模型'
|
||||
self.model_err5 = '模型名未本地注册'
|
||||
self.go_to_inference = '使用模型'
|
||||
self.save = '保存修改'
|
||||
self.delete = '删除'
|
||||
self.ms_sdk = 'ModelScope API Token'
|
||||
self.ms_username = 'ModelScope用户名'
|
||||
self.model_private = '模型不公开'
|
||||
self.ms_modelid = 'ModelScope模型ID'
|
||||
self.ms_url = 'ModelScope模型地址'
|
||||
self.ms_model_path = 'MS模型地址'
|
||||
self.export_file = '下载数据'
|
||||
self.export_zip_err1 = '导出模型失败'
|
||||
self.zip_file = '上传模型'
|
||||
self.utuner_name = '上传微调模型名称'
|
||||
self.ubase_model = '上传基模型类型'
|
||||
self.utuner_type = '上传微调模型类型'
|
||||
self.illegal_data_err1 = '上传文件格式错误(not .zip)'
|
||||
self.download_to_local = '下载至本地'
|
||||
self.export_desc = '模型导出至modelscope (\U00002714: 提交; \U00002716: 关闭)'
|
||||
self.import_desc = '模型从 **modelscope/本地** 导入 (\U00002714: 提交; \U00002716: 关闭)'
|
||||
|
||||
@@ -1,8 +1,13 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import os
|
||||
|
||||
import gradio as gr
|
||||
import yaml
|
||||
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.studio.tuner_manager.manager_ui.component_names import \
|
||||
TunerManagerNames
|
||||
from scepter.studio.utils.uibase import UIBase
|
||||
@@ -11,6 +16,9 @@ from scepter.studio.utils.uibase import UIBase
|
||||
class InfoUI(UIBase):
|
||||
def __init__(self, cfg, language='en'):
|
||||
self.component_names = TunerManagerNames(language)
|
||||
self.work_dir = cfg.WORK_DIR
|
||||
self.export_folder = os.path.join(self.work_dir, cfg.EXPORT_DIR)
|
||||
self.language = language
|
||||
|
||||
def create_ui(self, *args, **kwargs):
|
||||
with gr.Column():
|
||||
@@ -60,6 +68,155 @@ class InfoUI(UIBase):
|
||||
label=self.component_names.
|
||||
tuner_prompt_example,
|
||||
lines=2)
|
||||
with gr.Row(equal_height=True):
|
||||
self.ms_url = gr.Text(
|
||||
value='',
|
||||
label=self.component_names.ms_url,
|
||||
lines=2)
|
||||
with gr.Row(equal_height=True):
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.go_to_inferece_btn = gr.Button(
|
||||
self.component_names.go_to_inference)
|
||||
with gr.Column(scale=1, min_width=0):
|
||||
self.local_download_bt = gr.Button(
|
||||
label='Download to Local Dir',
|
||||
value=self.component_names.
|
||||
download_to_local,
|
||||
# elem_classes='type_row',
|
||||
elem_id='save_button')
|
||||
self.export_url = gr.File(
|
||||
label=self.component_names.export_file,
|
||||
visible=False,
|
||||
value=None,
|
||||
interactive=False,
|
||||
show_label=True)
|
||||
|
||||
def set_callbacks(self, manager):
|
||||
pass
|
||||
def go_to_inferece(new_name, tuner_desc, tuner_prompt_example,
|
||||
tuner_type, base_model):
|
||||
sub_dir = f'{base_model}-{tuner_type}'
|
||||
tar_path = os.path.join(self.work_dir, sub_dir)
|
||||
model_dir = os.path.join(tar_path, new_name)
|
||||
if not os.path.exists(model_dir):
|
||||
tuner_list = Config(
|
||||
cfg_file=os.path.join(self.work_dir, 'tuner_list.yaml'))
|
||||
for tuner_item in tuner_list.get('TUNERS', []):
|
||||
if tuner_item.NAME == new_name:
|
||||
model_dir = tuner_item.MODEL_PATH
|
||||
tuner_example = os.path.join(model_dir, 'image.jpg')
|
||||
if not os.path.exists(tuner_example):
|
||||
tuner_example = None
|
||||
tuner_dict = {
|
||||
'NAME': new_name,
|
||||
'NAME_ZH': new_name,
|
||||
'SOURCE': 'self_train',
|
||||
'DESCRIPTION': tuner_desc,
|
||||
'BASE_MODEL': base_model,
|
||||
'MODEL_PATH': model_dir,
|
||||
'IMAGE_PATH': tuner_example,
|
||||
'TUNER_TYPE': tuner_type,
|
||||
'PROMPT_EXAMPLE': tuner_prompt_example
|
||||
}
|
||||
tuner_cfg = Config(cfg_dict=tuner_dict, load=False)
|
||||
cfg_file = os.path.join(model_dir, 'meta.yaml')
|
||||
if not os.path.exists(model_dir):
|
||||
gr.Error(self.component_names.model_err4)
|
||||
|
||||
pipeline_level_modules = manager.inference.model_manage_ui.pipe_manager.pipeline_level_modules
|
||||
if tuner_cfg.BASE_MODEL not in pipeline_level_modules:
|
||||
gr.Error(self.component_names.model_err3 +
|
||||
tuner_cfg.BASE_MODEL)
|
||||
pipeline_ins = pipeline_level_modules[tuner_cfg.BASE_MODEL]
|
||||
diffusion_model = f"{tuner_cfg.BASE_MODEL}_{pipeline_ins.diffusion_model['name']}"
|
||||
|
||||
default_choices = manager.inference.model_manage_ui.pipe_manager.module_level_choices
|
||||
if 'customized_tuners' in default_choices:
|
||||
if tuner_cfg.BASE_MODEL not in default_choices[
|
||||
'customized_tuners']:
|
||||
default_choices['customized_tuners'] = {}
|
||||
tunner_choices = default_choices['customized_tuners'][
|
||||
tuner_cfg.BASE_MODEL]['choices']
|
||||
tunner_default = tuner_cfg.NAME if self.language == 'en' else tuner_cfg.NAME_ZH
|
||||
if tunner_default not in tunner_choices:
|
||||
if self.language == 'zh':
|
||||
gr.Error(self.component_names.model_err5 +
|
||||
tuner_cfg.NAME_ZH)
|
||||
else:
|
||||
gr.Error(self.component_names.model_err5 +
|
||||
tuner_cfg.NAME)
|
||||
if not isinstance(tunner_default, list):
|
||||
tunner_default = [tunner_default]
|
||||
else:
|
||||
tunner_choices = []
|
||||
tunner_default = []
|
||||
|
||||
with open(cfg_file, 'w') as f_out:
|
||||
yaml.dump(copy.deepcopy(tuner_cfg.cfg_dict),
|
||||
f_out,
|
||||
encoding='utf-8',
|
||||
allow_unicode=True,
|
||||
default_flow_style=False)
|
||||
|
||||
base_model = tuner_cfg.get('BASE_MODEL', '')
|
||||
|
||||
if not base_model == '':
|
||||
if base_model not in manager.inference.tuner_ui.name_level_tuners:
|
||||
manager.inference.tuner_ui.name_level_tuners[
|
||||
base_model] = {}
|
||||
manager.inference.tuner_ui.name_level_tuners[base_model][
|
||||
tuner_cfg.NAME] = tuner_cfg
|
||||
|
||||
return (
|
||||
gr.Tabs(selected='inference'), cfg_file,
|
||||
gr.Tabs(selected='tuner_ui'),
|
||||
gr.CheckboxGroup(
|
||||
value='使用微调' if self.language == 'zh' else 'Use Tuners'),
|
||||
gr.Dropdown(value=diffusion_model),
|
||||
gr.Dropdown(choices=tunner_choices, value=tunner_default))
|
||||
|
||||
self.go_to_inferece_btn.click(
|
||||
go_to_inferece,
|
||||
inputs=[
|
||||
self.new_name, self.tuner_desc, self.tuner_prompt_example,
|
||||
self.tuner_type, self.base_model
|
||||
],
|
||||
outputs=[
|
||||
manager.tabs, manager.inference.infer_info,
|
||||
manager.inference.setting_tab,
|
||||
manager.inference.check_box_for_setting,
|
||||
manager.inference.model_manage_ui.diffusion_model,
|
||||
manager.inference.tuner_ui.custom_tuner_model
|
||||
],
|
||||
queue=False)
|
||||
|
||||
def export_zip(tuner_name, base_model, tuner_type):
|
||||
sub_dir = f'{base_model}-{tuner_type}'
|
||||
if os.path.exists(os.path.join(self.work_dir, sub_dir,
|
||||
tuner_name)):
|
||||
model_dir = os.path.join(self.work_dir, sub_dir, tuner_name)
|
||||
else:
|
||||
model_dir = ''
|
||||
tuner_list = Config(
|
||||
cfg_file=os.path.join(self.work_dir, 'tuner_list.yaml'))
|
||||
for tuner_item in tuner_list.get('TUNERS', []):
|
||||
if tuner_item.NAME == tuner_name:
|
||||
model_dir = tuner_item.MODEL_PATH
|
||||
break
|
||||
if not os.path.exists(model_dir) or model_dir == '':
|
||||
raise gr.Error(self.component_names.model_err4)
|
||||
zip_path = os.path.join(self.export_folder, f'{tuner_name}.zip')
|
||||
with FS.put_to(zip_path) as local_zip:
|
||||
res = os.popen(
|
||||
f"cd '{model_dir}' "
|
||||
f"&& zip -r '{os.path.abspath(local_zip)}' ./* ")
|
||||
print(res.readlines())
|
||||
if not FS.exists(zip_path):
|
||||
raise gr.Error(self.component_names.export_zip_err1)
|
||||
local_zip = FS.get_from(zip_path)
|
||||
return gr.File(value=local_zip, visible=True)
|
||||
|
||||
self.local_download_bt.click(
|
||||
export_zip,
|
||||
inputs=[self.tuner_name, self.base_model, self.tuner_type],
|
||||
outputs=[self.export_url],
|
||||
queue=False)
|
||||
|
||||
@@ -18,6 +18,8 @@ class TunerManagerUI():
|
||||
cfg_general = Config(cfg_file=cfg_general_file)
|
||||
cfg_general.WORK_DIR = os.path.join(root_work_dir,
|
||||
cfg_general.WORK_DIR)
|
||||
cfg_general.SELF_TRAIN_DIR = os.path.join(root_work_dir,
|
||||
cfg_general.SELF_TRAIN_DIR)
|
||||
if not FS.exists(cfg_general.WORK_DIR):
|
||||
FS.make_dir(cfg_general.WORK_DIR)
|
||||
|
||||
|
||||
@@ -15,3 +15,9 @@ def init_env(cfg_general):
|
||||
is_flag = FS.make_dir(work_dir)
|
||||
assert is_flag
|
||||
return cfg_general
|
||||
|
||||
|
||||
def get_available_memory():
|
||||
import psutil
|
||||
mem = psutil.virtual_memory()
|
||||
return {'total': mem.total, 'available': mem.available}
|
||||
|
||||
@@ -128,7 +128,11 @@ if __name__ == '__main__':
|
||||
interfaces.append((interface, name, ifid))
|
||||
setattr(tab_manager, ifid, interface)
|
||||
|
||||
with gr.Blocks() as demo:
|
||||
css = """
|
||||
.upload_zone { height: 100px; }
|
||||
"""
|
||||
|
||||
with gr.Blocks(css=css) as demo:
|
||||
if 'BANNER' in config:
|
||||
gr.HTML(config.BANNER)
|
||||
else:
|
||||
@@ -136,7 +140,7 @@ if __name__ == '__main__':
|
||||
f"<h2><center>{config.get('TITLE', 'scepter studio')}</center></h2>"
|
||||
)
|
||||
setattr(tab_manager, 'user_name',
|
||||
gr.Text(value='', visible=False, show_label=False))
|
||||
gr.Text(value='admin', visible=False, show_label=False))
|
||||
with gr.Tabs(elem_id='tabs') as tabs:
|
||||
setattr(tab_manager, 'tabs', tabs)
|
||||
for interface, label, ifid in interfaces:
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
__version__ = '0.0.4'
|
||||
__version__ = '0.0.5'
|
||||
|
||||
version_info = tuple(int(x) for x in __version__.split('.')[0:3])
|
||||
|
||||
|
||||
@@ -0,0 +1,158 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import os
|
||||
import unittest
|
||||
|
||||
import numpy as np
|
||||
import torchvision.transforms as TT
|
||||
import torchvision.transforms.functional as TF
|
||||
from PIL import Image
|
||||
from torchvision.utils import save_image
|
||||
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
from scepter.modules.inference.diffusion_inference import DiffusionInference
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.distribute import we
|
||||
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()
|
||||
|
||||
@unittest.skip('')
|
||||
def test_sd15(self):
|
||||
config_file = 'scepter/methods/studio/inference/stable_diffusion/sd15_pro.yaml'
|
||||
cfg = Config(cfg_file=config_file)
|
||||
diff_infer = DiffusionInference(logger=self.logger)
|
||||
diff_infer.init_from_cfg(cfg)
|
||||
output = diff_infer({'prompt': 'a cute dog'})
|
||||
save_path = os.path.join(self.tmp_dir,
|
||||
'sd15_test_prompt_a_cute_dog.png')
|
||||
save_image(output['images'], save_path)
|
||||
|
||||
@unittest.skip('')
|
||||
def test_sd21(self):
|
||||
config_file = 'scepter/methods/studio/inference/stable_diffusion/sd21_pro.yaml'
|
||||
cfg = Config(cfg_file=config_file)
|
||||
diff_infer = DiffusionInference(logger=self.logger)
|
||||
diff_infer.init_from_cfg(cfg)
|
||||
output = diff_infer({'prompt': 'a cute dog'})
|
||||
save_path = os.path.join(self.tmp_dir,
|
||||
'sd21_test_prompt_a_cute_dog.png')
|
||||
save_image(output['images'], save_path)
|
||||
|
||||
@unittest.skip('')
|
||||
def test_sdxl(self):
|
||||
config_file = 'scepter/methods/studio/inference/sdxl/sdxl1.0_pro.yaml'
|
||||
cfg = Config(cfg_file=config_file)
|
||||
diff_infer = DiffusionInference(logger=self.logger)
|
||||
diff_infer.init_from_cfg(cfg)
|
||||
output = diff_infer({'prompt': 'a cute dog'})
|
||||
save_path = os.path.join(self.tmp_dir,
|
||||
'sdxl_test_prompt_a_cute_dog.png')
|
||||
save_image(output['images'], save_path)
|
||||
|
||||
# @unittest.skip('')
|
||||
def test_sd15_scedit_t2i_2D(self):
|
||||
# init model
|
||||
config_file = 'scepter/methods/studio/inference/stable_diffusion/sd15_pro.yaml'
|
||||
cfg = Config(cfg_file=config_file)
|
||||
diff_infer = DiffusionInference(logger=self.logger)
|
||||
diff_infer.init_from_cfg(cfg)
|
||||
# load tuner model
|
||||
tuner_model = {
|
||||
'NAME': 'Flat 2D Art',
|
||||
'NAME_ZH': None,
|
||||
'DESCRIPTION': None,
|
||||
'BASE_MODEL': 'SD1.5',
|
||||
'IMAGE_PATH': None,
|
||||
'TUNER_TYPE': 'SwiftSCE',
|
||||
'MODEL_PATH':
|
||||
'ms://damo/scepter_scedit@tuners_model/SD1.5/Flat2DArt',
|
||||
'PROMPT_EXAMPLE': None
|
||||
}
|
||||
tuner_model = Config(cfg_dict=tuner_model, load=False)
|
||||
# prepare data
|
||||
input_data = {'prompt': 'a single flower is shown in front of a tree'}
|
||||
input_params = {
|
||||
'tuner_model': tuner_model,
|
||||
'tuner_scale': 1.0,
|
||||
'seed': 2024
|
||||
}
|
||||
output = diff_infer(input_data, **input_params)
|
||||
save_path = os.path.join(self.tmp_dir, 'sd15_flower_2d.png')
|
||||
save_image(output['images'], save_path)
|
||||
|
||||
# @unittest.skip('')
|
||||
def test_sdxl_scedit_ctr_canny(self):
|
||||
# init model
|
||||
config_file = 'scepter/methods/studio/inference/sdxl/sdxl1.0_pro.yaml'
|
||||
cfg = Config(cfg_file=config_file)
|
||||
diff_infer = DiffusionInference(logger=self.logger)
|
||||
diff_infer.init_from_cfg(cfg)
|
||||
# extract condition
|
||||
canny_dict = {
|
||||
'NAME': 'CannyAnnotator',
|
||||
'LOW_THRESHOLD': 100,
|
||||
'HIGH_THRESHOLD': 200
|
||||
}
|
||||
canny_anno = Config(cfg_dict=canny_dict, load=False)
|
||||
canny_ins = ANNOTATORS.build(canny_anno).to(we.device_id)
|
||||
output_height, output_width = 1024, 1024
|
||||
control_cond_image = Image.open('asset/images/flower.jpg')
|
||||
control_cond_image = TT.Resize(min(output_height,
|
||||
output_width))(control_cond_image)
|
||||
control_cond_image = TT.CenterCrop(
|
||||
(output_height, output_width))(control_cond_image)
|
||||
control_cond_image = canny_ins(np.array(control_cond_image))
|
||||
control_save_path = os.path.join(self.tmp_dir,
|
||||
'sdxl_flower_canny_preproccess.png')
|
||||
save_image(TF.to_tensor(control_cond_image), control_save_path)
|
||||
control_cond_image = Image.open(control_save_path)
|
||||
# load control model
|
||||
control_model = {
|
||||
'NAME':
|
||||
'canny',
|
||||
'NAME_ZH':
|
||||
None,
|
||||
'DESCRIPTION':
|
||||
None,
|
||||
'BASE_MODEL':
|
||||
'SD_XL1.0',
|
||||
'TYPE':
|
||||
'Canny',
|
||||
'MODEL_PATH':
|
||||
'ms://damo/scepter_scedit@controllable_model/SD_XL1.0/canny_control'
|
||||
}
|
||||
control_model = Config(cfg_dict=control_model, load=False)
|
||||
# prepare data
|
||||
input_data = {'prompt': 'a single flower is shown in front of a tree'}
|
||||
input_params = {
|
||||
'control_model': control_model,
|
||||
'control_cond_image': control_cond_image,
|
||||
'control_scale': 1.0,
|
||||
'crop_type': 'CenterCrop',
|
||||
'seed': 2024
|
||||
}
|
||||
output = diff_infer(input_data, **input_params)
|
||||
save_path = os.path.join(self.tmp_dir, 'sdxl_flower_canny.png')
|
||||
save_image(output['images'], save_path)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
@@ -13,8 +13,8 @@ class TrainTest(unittest.TestCase):
|
||||
os.makedirs(self.tmp_dir)
|
||||
self.data_dir = './cache/datasets'
|
||||
if not os.path.exists(self.data_dir):
|
||||
data_cmd = """mkdir -p cache/datasets/ && wget 'https://modelscope.cn/api/v1/models
|
||||
/damo/scepter_scedit/repo?Revision=master&FilePath=dataset/3D_example_txt.zip'
|
||||
data_cmd = """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"""
|
||||
|
||||
Reference in New Issue
Block a user