diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..8a2eb67 --- /dev/null +++ b/.gitignore @@ -0,0 +1,22 @@ +*.pyc +*.pth +*.pt +*.pkl +*.ckpt +*.png +*.DS_Store +*__pycache__* +*.cache* +*.bin +*.idea +*.csv +*.txt +build +dist +dev +scepter.egg-info +.readthedocs.yml +1.9 +MANIFEST.in +*resources +*.ipynb_checkpoints* diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/docs/en/Makefile b/docs/en/Makefile new file mode 100644 index 0000000..ed88099 --- /dev/null +++ b/docs/en/Makefile @@ -0,0 +1,20 @@ +# Minimal makefile for Sphinx documentation +# + +# You can set these variables from the command line, and also +# from the environment for the first two. +SPHINXOPTS ?= +SPHINXBUILD ?= sphinx-build +SOURCEDIR = . +BUILDDIR = build + +# Put it first so that "make" without argument is like "make help". +help: + @$(SPHINXBUILD) -M help "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O) + +.PHONY: help Makefile + +# Catch-all target: route all unknown targets to Sphinx using the new +# "make mode" option. $(O) is meant as a shortcut for $(SPHINXOPTS). +%: Makefile + @$(SPHINXBUILD) -M $@ "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O) diff --git a/docs/en/conf.py b/docs/en/conf.py new file mode 100644 index 0000000..a713930 --- /dev/null +++ b/docs/en/conf.py @@ -0,0 +1,83 @@ +# -*- coding: utf-8 -*- +# Configuration file for the Sphinx documentation builder. +# +# This file only contains a selection of the most common options. For a full +# list see the documentation: +# https://www.sphinx-doc.org/en/master/usage/configuration.html + +# -- Path setup -------------------------------------------------------------- + +# If extensions (or modules to document with autodoc) are in another directory, +# add these directories to sys.path here. If the directory is relative to the +# documentation root, use os.path.abspath to make it absolute, like shown here. +# +# import os +# import sys +# sys.path.insert(0, os.path.abspath('.')) + +# -- Project information ----------------------------------------------------- + +project = 'scepter' +copyright = '2023, scepter' +author = 'scepter' + + +def get_version(): + version_path = '../../scepter/version.py' + with open(version_path) as f: + exec(compile(f.read(), version_path, 'exec')) + return locals()['__version__'] + + +release = get_version() + +# -- General configuration --------------------------------------------------- + +# Add any Sphinx extension module names here, as strings. They can be +# extensions coming with Sphinx (named 'sphinx.ext.*') or your custom +# ones. +extensions = [ + 'sphinx.ext.autodoc', + 'sphinx.ext.napoleon', + 'sphinx.ext.viewcode', + 'sphinx_copybutton', + 'recommonmark', + 'sphinx_markdown_tables', +] + +# Add any paths that contain templates here, relative to this directory. +templates_path = ['_templates'] + +# The language for content autogenerated by Sphinx. Refer to documentation +# for a list of supported languages. +# +# This is also used if you do content translation via gettext catalogs. +# Usually you set "language" from the command line for these cases. +language = 'zh_CN' + +master_doc = 'index' + +# List of patterns, relative to source directory, that match files and +# directories to ignore when looking for source files. +# This pattern also affects html_static_path and html_extra_path. +exclude_patterns = ['build'] + +# -- Options for HTML output ------------------------------------------------- + +source_suffix = { + '.rst': 'restructuredtext', + '.md': 'markdown', +} + +# The theme to use for HTML and HTML Help pages. See the documentation for +# a list of builtin themes. +html_theme = 'press' + +# Add any paths that contain custom static files (such as style sheets) here, +# relative to this directory. They are copied after the builtin static files, +# so a file named "default.css" will overwrite the builtin "default.css". +html_static_path = ['_static'] + +# import sphinx_rtd_theme +# html_theme = "sphinx_rtd_theme" +# html_theme_path = [sphinx_rtd_theme.get_html_theme_path()] diff --git a/docs/en/index.rst b/docs/en/index.rst new file mode 100644 index 0000000..a5ce60a --- /dev/null +++ b/docs/en/index.rst @@ -0,0 +1,32 @@ + +========================================== + +.. toctree:: + :maxdepth: 1 + :caption: Quick Start + + scepter/quick_start.md + +.. toctree:: + :maxdepth: 2 + :caption: Core Modules + + scepter/data.md + scepter/model.md + scepter/opt.md + scepter/solvers.md + scepter/tools.md + scepter/transforms.md + +.. toctree:: + :maxdepth: 2 + :caption: Utils Modules + + scepter/utils/file_clients.md + scepter/utils/utils.md + +.. toctree:: + :maxdepth: 2 + :caption: ChangeLogs + + scepter/changelog.md diff --git a/docs/en/make.bat b/docs/en/make.bat new file mode 100644 index 0000000..061f32f --- /dev/null +++ b/docs/en/make.bat @@ -0,0 +1,35 @@ +@ECHO OFF + +pushd %~dp0 + +REM Command file for Sphinx documentation + +if "%SPHINXBUILD%" == "" ( + set SPHINXBUILD=sphinx-build +) +set SOURCEDIR=source +set BUILDDIR=build + +if "%1" == "" goto help + +%SPHINXBUILD% >NUL 2>NUL +if errorlevel 9009 ( + echo. + echo.The 'sphinx-build' command was not found. Make sure you have Sphinx + echo.installed, then set the SPHINXBUILD environment variable to point + echo.to the full path of the 'sphinx-build' executable. Alternatively you + echo.may add the Sphinx directory to PATH. + echo. + echo.If you don't have Sphinx installed, grab it from + echo.https://www.sphinx-doc.org/ + exit /b 1 +) + +%SPHINXBUILD% -M %1 %SOURCEDIR% %BUILDDIR% %SPHINXOPTS% %O% +goto end + +:help +%SPHINXBUILD% -M help %SOURCEDIR% %BUILDDIR% %SPHINXOPTS% %O% + +:end +popd diff --git a/docs/en/scepter/dataset.md b/docs/en/scepter/dataset.md new file mode 100644 index 0000000..95041a3 --- /dev/null +++ b/docs/en/scepter/dataset.md @@ -0,0 +1,53 @@ +# Dataset Module (Dataset) +## Overview +Register individual dataset reading modules for each task by inheriting from BaseDataset, which encapsulates File System operations and a transform pipeline. +*** +## **scepter.modules.data.dataset.BaseDataset** +### function **\_\_init\_\_()** +#### Parameters (Input Parameters) +`(cfg: scepter.modules.utils.config.Config, logger = None) -> None` +#### config: +- MODE —— [train, test, eval]; +- TRANSFORMS —— For constructing a pipeline to process each sample. See docs.transforms for more details; +- FILE_SYSTEM —— File IO Handler, supports read and write operations for different file types. See docs.utils.file_clients for more details; +#### object (Internal) +- self.pipeline —— Transform pipeline +- self.fs_prefix —— The prefix of the instantiated fs_client +- self.local_we —— Local rank info when is_distributed + +### function **\_\_getitem\_\_()** +#### Parameters (Input Parameters): index +The type of index is determined by the sampler, refer to data/registry.py for details; +The default torch dataloader passes an integer type index as the dataset subscript; +Custom samplers can send customized items. Sampler definitions are found in scepter.modules.utils.sampler; +- The specific retrieval of a single data item is implemented by the _get() method with the index parameter; +- Outputs data that has been transformed by the pipeline (if there is a pipeline); + +### function **worker_init_fn()** +#### Parameters (Input Parameters): +`(worker_id, num_workers = 1)` +worker_id is the worker node ID in distributed training; +The dataloader instantiates num_workers worker processes at one time; +- For initializing the file reading system and setting parameters corresponding to multi-GPU workers; + +### function **\_get()** +#### Parameters (Input Parameters): index (passed in by __getitem__()) +An abstract method that must be concretely implemented by the custom dataset. It's used to read each data item in a batch according to the index; +*** +## Basic Usage +Subclass registration: +```python +from scepter.modules.data.dataset import BaseDataset +from scepter.modules.data.dataset import DATASETS +@DATASETS.register_class() +class XxxxxDataset(BaseDataset): + def __init__(self, cfg, logger=None): + super(XxxxxDataset, self).__init__(cfg, logger=logger) +``` +Actual usage of starting a dataset: + +```python +if self.cfg.have("TRAIN_DATA"): + train_data = DATASETS.build(self.cfg.TRAIN_DATA, logger=self.logger) +``` +Refer to scepter/modules/solver/diffusion_solver.py for more details. Specific parameter definitions are located under the TRAIN_DATA/EVAL_DATA/TEST_DATA sections in the configuration yaml. diff --git a/docs/en/scepter/model.md b/docs/en/scepter/model.md new file mode 100644 index 0000000..9ff4287 --- /dev/null +++ b/docs/en/scepter/model.md @@ -0,0 +1,181 @@ +# Model Modules (Model) +## Overview +Model modules are divided into backbones, necks, heads, loss, metrics, networks, tokenizers; +* backbones/necks:Generally the main modules for feature extraction (necks are not always present); +* heads: According to different task types, they take the features extracted by backbones and output the logits required for downstream tasks; +* loss: Used to calculate different types of loss; +* metric: Used to calculate various evaluation metrics; +* tokenizer: Used for tokenization; +* network:train和test模块,对数据集输入的batch整合上述模块进行最终loss和指标计算; +
+ +## **backbones/necks/heads/loss** +### Basic Usage +Subclass registration: + +```python +from scepter.model.registry import BACKBONES +from scepter.model.base_model import BaseModel + + +@BACKBONES.register_class("ResNet") +class ResNet(BaseModel): + def __init__(self, cfg, logger=None): + super(ResNet, self).__init__(cfg, logger=logger) +``` + +```python +from scepter.model.registry import NECKS +from scepter.model.base_model import BaseModel + + +@NECKS.register_class() +class GlobalAveragePooling(BaseModel): + def __init__(self, cfg, logger=None): + super(GlobalAveragePooling, self).__init__(cfg, logger=logger) +``` + +```python +from scepter.model.registry import HEADS +from scepter.model.base_model import BaseModel + + +@HEADS.register_class() +class ClassifierHead(BaseModel): + def __init__(self, cfg, logger=None): + super(ClassifierHead, self).__init__(cfg, logger=logger) +``` + +```python +from scepter.model.registry import LOSSES +import torch.nn as nn + + +@LOSSES.register_class() +class CrossEntropy(nn.Module): + def __init__(self, cfg, logger=None): + super(CrossEntropy, self).__init__(cfg, logger=logger) +``` +Actual usage: + +```python +from scepter.model.registry import BACKBONES, NECKS, HEADS, LOSSES + +backbone = BACKBONES.build(cfg.BACKBONE, logger=logger) +neck = NECKS.build(cfg.NECK, logger=logger) +head = HEADS.build(cfg.HEAD, logger=logger) +loss = LOSSES.build(cfg.LOSS, logger=logger) +``` + +### function **\_\_init\_\_()** +#### Parameters(input parameters) +(cfg: scepter.modules.utils.config.Config, logger = None) -> None + +Mainly used for initializing various layers of the model; + +### function **forward()** +To be implemented specifically as needed; +
+ + +## **metrics** +### Basic Usage +Basic Usage Subclass registration: + +```python +from scepter.model.metrics.registry import METRICS +from scepter.model.metrics.base_metric import BaseMetric + + +@METRICS.register_class("AccuracyMetric") +class AccuracyMetric(BaseMetric): + def __init__(self, cfg, logger=None): + super(CrossEntropy, self).__init__(cfg, logger=logger) +``` +Actual usage: + +```python +from scepter.model.metrics.registry import METRICS + +metric = METRICS.build(cfgs, logger) +``` + +### function **\_\_init\_\_()** +#### Parameters(input parameters) +(cfg: scepter.modules.utils.config.Config, logger = None) -> None + +Initializes the hyperparameters needed for calculating metrics, such as coefficients like topk; + +### function **\_\_call\_\_()** +@torch.no_grad() + +Typically takes logits and labels as well as other necessary variables as inputs and outputs calculated metrics; +
+ +## **tokenizers** +### Basic Usage +Subclass registration: + +```python +from scepter.model.registry import TOKENIZERS +from scepter.model.tokenizers import BaseTokenizer + + +@TOKENIZERS.register_class() +class BaseBertTokenizer(BaseTokenizer): + def __init__(self, cfg, logger=None): + super(BaseBertTokenizer, self).__init__(cfg, logger=logger) +``` +Actual usage: + +```python +from scepter.model.registry import TOKENIZERS + +tokenizer = TOKENIZERS.build(cfgs, logger) +``` +### function **\_\_init\_\_()** +#### Parameters(input parameters) +(cfg: scepter.modules.utils.config.Config, logger = None) -> None +None Used for initializing and loading the tokenizer object, such as BertTokenizer; + +### function **tokenize()** +Takes a list of texts that need tokenization as input and outputs token id sequences after tokenization, as well as other necessary elements like attention masks, type id lists, position id lists, etc; +
+ +## **networks** +### Basic Usage +Subclass registration: + +```python +from scepter.model.registry import MODELS +from scepter.model.networks.train_module import TrainModule + + +@MODELS.register_class() +class Classifier(TrainModule): + def __init__(self, cfg, logger=None): + super(Classifier, self).__init__(cfg, logger=logger) +``` +Actual usage: + +```python +from scepter.model.registry import MODELS + +model = MODELS.build(self.cfg.MODEL, logger=self.logger) +``` + +### function **\_\_init\_\_()** +#### Parameters(input parameters) +(cfg: scepter.modules.utils.config.Config, logger = None) -> None + +Combined with the above build method, initializes the backbone, neck, head, loss, metric, tokenizer modules needed for training; + +### function **forward_train()** +Takes data from a training batch, processes it through the backbone, neck, head, and loss to calculate the relevant loss; + +### function **forward_test()** +Takes data from a test batch, processes it through the backbone, neck, head, and metrics to calculate the relevant metrics; + +### function **forward** +The actual calling interface, used to dispatch tasks to forward_train()/forward_test(); +Other functions needed for training/testing can be customized under network;" diff --git a/docs/en/scepter/opt.md b/docs/en/scepter/opt.md new file mode 100644 index 0000000..a986393 --- /dev/null +++ b/docs/en/scepter/opt.md @@ -0,0 +1,83 @@ +# Optimizer (Optimizer) +## Overview +1. lr_schedulers +2. optimizers +
+ +## lr_schedulers +### Basic Usage +Usage when subclassing lr_schedulers: + +```python +from scepter.opt.lr_schedulers import LR_SCHEDULERS +from scepter.opt.lr_schedulers.base_scheduler import BaseScheduler + + +@LR_SCHEDULERS.register_class() +class XxxLR(BaseScheduler): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) +``` +Actual usage to start lr_scheduler (optimizer is necessary), refer to task/stable_diffusion/impls/solvers/diffusion_solver.py: +```python +if self.cfg.have("LR_SCHEDULER") and not self.optimizer is None: + self.lr_scheduler = LR_SCHEDULERS.build(self.cfg.LR_SCHEDULER, logger=self.logger, + optimizer=self.optimizer) +``` + +## **scepter.modules.opt.lr_schedulers.base_scheduler.BaseScheduler** +The base class for lr_schedulers, supports registration operations, can be customized as needed; + +### function **\_\_init\_\_()** +#### Parameters(input parameters) +(cfg: scepter.modules.utils.config.Config, logger = None) -> None +#### config(Common parameters, actual needs vary according to different schedulers, taking StepLR as an example): +* STEP_SIZE +* GAMMA +* LAST_EPOCH + +### function **\_\_call\_\_()** +#### Parameters(input parameters) +(optimizer: scepter.modules.opt.optimizers.OPTIMIZERS) -> None + +Sets up the schedule for the passed-in optimizer object; +
+ +## optimizers +### Basic Usage +Usage when subclassing optimizers: + +```python +from scepter.opt.optimizers.base_optimizer import BaseOptimize +from scepter.opt.optimizers.registry import OPTIMIZERS + + +@OPTIMIZERS.register_class() +class Xxx(BaseOptimize): + def __init__(self, cfg, logger=None): + super(Xxx, self).__init__(cfg, logger=logger) +``` +Actual usage to start optimizers, refer to task/stable_diffusion/impls/solvers/diffusion_solver.py, requires passing in train_parameters: +```python +if self.cfg.have("OPTIMIZER"): + self.optimizer = OPTIMIZERS.build(self.cfg.OPTIMIZER, logger=self.logger, + parameters=self.train_parameters()) +``` + +## **scepter.modules.opt.optimizers.base_optimizer.BaseOptimize** +The base class for optimizers, supports registration operations, can be customized as needed; + +### function **\_\_init\_\_()** +#### Parameters(input parameters) +(cfg: scepter.modules.utils.config.Config, logger = None) -> None +#### config(common parameters, actual needs vary according to different optimizers, taking SGD as an example): +* LEARNING_RATE +* MOMENTUM +* DAMPENING +* WEIGHT_DECAY +* NESTEROV + +### function **\_\_call\_\_()** +#### Parameters(input parameters) +(parameters:dict()) -> None +Inputs the train parameters that need gradient updates, in dict format; diff --git a/docs/en/scepter/quick_start.md b/docs/en/scepter/quick_start.md new file mode 100644 index 0000000..11b600a --- /dev/null +++ b/docs/en/scepter/quick_start.md @@ -0,0 +1,282 @@ +# Quick Start + +This chapter takes stable diffusion v1.5 as an example, demonstrating how to build a neural network architecture from scratch, as well as how to train and test a model based on that structure. + +# 1. Defining Model Structure with the Network Class + +The network class encompasses methods for defining, training, and testing the model. First, we initialize a network class and then define the required submodules: autoencoder, UNet, embedder, and loss. + + + +```python +@MODELS.register_class() +class LatentDiffusion(TrainModule): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.model_config = cfg.DIFFUSION_MODEL + self.first_stage_config = cfg.FIRST_STAGE_MODEL + self.cond_stage_config = cfg.COND_STAGE_MODEL + self.loss_config = cfg.get('LOSS', None) + + self.model = BACKBONES.build(self.model_config, logger=self.logger) + self.first_stage_model = MODELS.build(self.first_stage_config, + logger=self.logger) + self.cond_stage_model = EMBEDDERS.build(self.cond_stage_config, + logger=self.logger) + if self.loss_config: + self.loss = LOSSES.build(self.loss_config, logger=self.logger) + + # Other module definition. +``` + +# 2. Implementing the Training and Testing Code for the Network class + +Each network class relies on forward_train and forward_test functions to define its training and testing processes. In SD1.5, forward_train predicts the noise for each timestep t and conducts loss calculation. + +```python + def forward_train(self, image=None, noise=None, prompt=None, **kwargs): + x_start = self.encode_first_stage(image, **kwargs) + t = torch.randint(0, + self.num_timesteps, (x_start.shape[0], ), + device=x_start.device).long() + context = {} + if prompt and self.cond_stage_model: + zeros = (torch.rand(len(prompt)) < self.p_zero).numpy().tolist() + prompt = [ + self.train_n_prompt if zeros[idx] else p + for idx, p in enumerate(prompt) + ] + + with torch.autocast(device_type='cuda', enabled=False): + context = self.encode_condition( + self.tokenizer(prompt).to(we.device_id)) + + loss = self.diffusion.loss(x0=x_start, + t=t, + model=self.model, + model_kwargs={'cond': context}, + noise=noise) + loss = loss.mean() + ret = {'loss': loss, 'probe_data': {'prompt': prompt}} + return ret +``` + +The forward_test function is used during the inference stage to perform the full image denoising process. + +```python +@torch.no_grad() +@torch.autocast('cuda', dtype=torch.float16) +def forward_test(self, + prompt=None, + n_prompt=None, + sampler='ddim', + sample_steps=50, + seed=2023, + guide_scale=7.5, + guide_rescale=0.5, + discretization='trailing', + run_train_n=True, + **kwargs): + g = torch.Generator(device=we.device_id) + seed = seed if seed >= 0 else random.randint(0, 2**32 - 1) + g.manual_seed(seed) + num_samples = len(prompt) + + n_prompt = default(n_prompt, [self.default_n_prompt] * len(prompt)) + assert isinstance(prompt, list) and \ + isinstance(n_prompt, list) and \ + len(prompt) == len(n_prompt) + + context = self.encode_condition(self.tokenizer(prompt).to( + we.device_id), method='encode_text') + null_context = self.encode_condition(self.tokenizer(n_prompt).to( + we.device_id), method='encode_text') + + width, height = 512, 512 + noise = self.noise_sample(num_samples, width // self.size_factor, + height // self.size_factor, g) + # UNet use input n_prompt + samples = self.diffusion.sample(solver=sampler, + noise=noise, + model=self.model, + model_kwargs=[{ + 'cond': context + }, { + 'cond': null_context + }], + steps=sample_steps, + guide_scale=guide_scale, + guide_rescale=guide_rescale, + discretization=discretization, + show_progress=True, + seed=seed, + condition_fn=None, + clamp=None, + percentile=None, + t_max=None, + t_min=None, + discard_penultimate_step=None, + return_intermediate=None, + **kwargs) + x_samples = self.decode_first_stage(samples).float() + x_samples = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0) + + outputs = list() + for p, np, img in zip(prompt, n_prompt, x_samples): + one_tup = {'prompt': p, 'n_prompt': np, 'image': img} + outputs.append(one_tup) + + return outputs +``` + +# 3. Submodule Registration + +After the network class is fully implemented, it is necessary to ensure that all submodules used within the class are registered. For instance, with the embedder in SD1.5, to instantiate this embedder within the network's initialization method, we must first implement the embedder class and register it with Scepter. + +```python +@EMBEDDERS.register_class() +class FrozenCLIPEmbedder(BaseEmbedder): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + + with FS.get_dir_to_local_dir(cfg.PRETRAINED_MODEL, + wait_finish=True) as local_path: + self.tokenizer = CLIPTokenizer.from_pretrained(local_path) + self.transformer = CLIPTextModel.from_pretrained(local_path) + + self.use_grad = cfg.get('USE_GRAD', False) + self.freeze_flag = cfg.get('FREEZE', True) + if self.freeze_flag: + self.freeze() + + self.max_length = cfg.get('MAX_LENGTH', 77) + self.layer = cfg.get('LAYER', 'last') + self.layer_idx = cfg.get('LAYER_IDX', None) + self.use_final_layer_norm = cfg.get('USE_FINAL_LAYER_NORM', False) + assert self.layer in self.LAYERS + if self.layer == 'hidden': + assert self.layer_idx is not None + assert 0 <= abs(self.layer_idx) <= 12 + + def freeze(self): + self.transformer = self.transformer.eval() + for param in self.parameters(): + param.requires_grad = False + + # @torch.no_grad() + def _encode_text(self, + tokens, + tokenizer=None, + append_sentence_embedding=False): + outputs = self.transformer(input_ids=tokens, + output_hidden_states=self.layer == 'hidden') + if self.layer == 'last': + z = outputs.last_hidden_state + elif self.layer == 'pooled': + z = outputs.pooler_output[:, None, :] + if self.use_final_layer_norm: + z = self.transformer.text_model.final_layer_norm(z) + else: + z = outputs.hidden_states[self.layer_idx] + if self.use_final_layer_norm: + z = self.transformer.text_model.final_layer_norm(z) + return z + + def encode_text(self, + tokens, + tokenizer=None, + append_sentence_embedding=False): + if not self.use_grad: + with torch.no_grad(): + output = self._encode_text(tokens, tokenizer, + append_sentence_embedding) + else: + output = self._encode_text(tokens, tokenizer, + append_sentence_embedding) + return output + + @staticmethod + def get_config_template(): + return dict_to_yaml('MODELS', + __class__.__name__, + FrozenCLIPEmbedder.para_dict, + set_name=True) +``` + +# 4. Solver Registration + +The solver class encapsulates the complete process needed to train and test a network. For instance, to register a solver for training the SD1.5 model, you'd need to: + +1. Create one or more data loaders, corresponding to the construct_data method in the solver; +2. Instantiate a model, corresponding to the construct_model method in the solver; +3. (Optional) Define metrics, corresponding to the construct_metrics method in the solver; +4. (Optional) Define various hooks (HOOKS) used in training and testing, such as model saving, pre-trained parameter loading, logging, etc. + +```python +@SOLVERS.register_class() +class LatentDiffusionSolver(BaseSolver): + def set_up(self): + self.construct_data() + self.construct_model() + self.construct_metrics() + self.model_to_device() + self.init_opti() + + def load_checkpoint(self, checkpoint): + pass + + def save_checkpoint(self): + pass + + def solve(self): + self.before_solve() + if 'train' in self._mode_set: + self.run_train() + if 'test' in self._mode_set: + self.run_test() + self.after_solve() + + def run_train(self): + pass + + def run_eval(self): + pass + + def run_test(self): + pass +``` + +# 5. Training/Testing Hyperparameter Definitions + +After registering all required components, which include but are not limited to BACKBONE, NETWORK, EMBEDDER, SOLVER, and METRIC, it is necessary to configure some of the hyperparameters used. Scepter employs YAML files to define hyperparameters for each module; for specifics, refer to scepter/modules/examples/sd15/sd15_512_full.yaml. + +# 6. Model Training + +Execute training or batch testing by loading the required YAML file with the --cfg parameter. +```shell +# Multi-node Multi-GPU Training +# Spawn mode (default) +export CUDA_VISIBLE_DEVICES=0,1,2,3 +export WORLD_SIZE=1 +python -W ignore scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml +# torchrun mode +torchrun --nproc_per_node 4 scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml --launcher torchrun +# pytorch_lightning mode, ENV.USE_PL=true +python -W ignore scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml +# Single-GPU Training +# Spawn mode (default) +export CUDA_VISIBLE_DEVICES=0 +export WORLD_SIZE=1 +python -W ignore scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml +# torchrun mode +python -W ignore scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml --launcher torchrun +# pytorch_lightning mode, ENV.USE_PL=true +python scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml +``` + +# 7. Model Inference + +Single inference can be realized by customizing the run_inference.py file, with reference to the inference method of SD1.5. +```shell +python -W ignore scepter/run_inference.py --prompt "a woman" --n_prompt "" --num_samples 4 --pretrained_model "path/to/your/pretrained/model" +``` diff --git a/docs/en/scepter/sampler.md b/docs/en/scepter/sampler.md new file mode 100644 index 0000000..f59490a --- /dev/null +++ b/docs/en/scepter/sampler.md @@ -0,0 +1,130 @@ +# Sampler Module (Sampler) +## Overview +A sampler defines the method for selecting the necessary data for training, validation, and testing. +Samplers are fairly generic and in most cases do not require custom development. Here, several commonly used samplers are provided. +Some samplers, such as MixtureOfSamplers, encapsulate a SubSampler. + +Of course, samplers also support custom development and registration. Developers just need to inherit from BaseSampler and register their implemented sampler class. +## Basic Usage +Predefined samplers: +- TorchDefault +- LoopSampler +- MixtureOfSamplers +- MultiFoldDistributedSampler +- EvalDistributedSampler +- MultiLevelBatchSampler +- MultiLevelBatchSamplerMultiSource +```python +from scepter.modules.data.sampler import ( + LoopSampler, + MixtureOfSamplers, + MultiFoldDistributedSampler, + EvalDistributedSampler, + MultiLevelBatchSampler, + MultiLevelBatchSamplerMultiSource) +``` +Custom sampler, taking LoopSampler as an example: +```python +@SAMPLERS.register_class() +class LoopSampler(BaseSampler): + para_dict = {} + def __init__(self, cfg, logger): + super().__init__(cfg, logger) + rank = we.rank + self.rng = np.random.default_rng(self.seed + rank) + def __iter__(self): + while True: + yield self.rng.choice(sys.maxsize) + def __len__(self): + return sys.maxsize + @staticmethod + def get_config_template(): + return dict_to_yaml('SAMPLERS', + __class__.__name__, + LoopSampler.para_dict, + set_name=True) +``` +Implement the iter and len methods to perform sampling. Related configuration parameters can be read from the cfg parameter in init. +
+ +## LoopSampler +An infinite looping sampler for simple data. +#### function **LoopSampler.__init__** +(cfg: scepter.modules.utils.config.Config, logger=None) -> None + +#### function **LoopSampler.__iter__** +Iterator, each iteration yields the index of a sample. + +## MixtureOfSamplers +A sampler for large-scale data with multi-level indexing. +#### function **MixtureOfSamplers.__init__** +(samplers: list[sampler], probabilities: list[float], rank: int = 0, seed: int = 8888) + +**Parameters** +- **samplers** —— List of samplers for mixing. +- **probabilities** —— Probabilities associated with each sampler. +- **rank** —— The rank representing the current process number. +- **seed** —— The seed for random sampling, obtained globally from data.registry. + +#### function **MultiLevelBatchSampler.__iter__** +() +Iterator, each iteration yields the index of a sample. + +## MultiFoldDistributedSampler +A multi-fold sampler that supports repeating data several times within one epoch. +#### function **MultiFoldDistributedSampler.__init__** +(dataset: torch.utils.data.Dataset, num_folds=1, num_replicas=None, rank=None, shuffle=True) + +**Parameters** +- **dataset** —— An instance of the torch.utils.data.Dataset class. +- **num_folds** —— Integer, indicating the number of times data is repeated. +- **num_replicas** —— The number of data partitions, generally consistent with the world size. +- **rank** —— The rank representing the current process number. +- **shuffle** —— Whether to shuffle the data or not. + +#### function **MultiFoldDistributedSampler.__iter__** +() +Iterator, each iteration yields the index of a sample. + +#### function **MultiFoldDistributedSampler.set_epoch** +(epoch: int) +Sets the current epoch. +**Parameters** +- **epoch** —— The current epoch number. + +## EvalDistributedSampler +A sampler for evaluation during testing, where you may find that the last rank has less data than other ranks when not using padding mode. +#### function **EvalDistributedSampler.__init__** +(dataset: torch.utils.data.Dataset, num_replicas: Optional[int] = None, rank: Optional[int] = None, padding: bool = False) + +**Parameters** +- **dataset** —— Instance of the torch.utils.data.Dataset class. +- **num_replicas** —— The number of data partitions, typically consistent with the world size. +- **rank** —— The rank representing the current process number. +- **padding** —— Whether the data needs padding. If true, it can ensure the last rank has the same amount of data as other ranks. + +#### function **EvalDistributedSampler.__iter__** +() +Iterator, each iteration yields the index of a sample. + +#### function **EvalDistributedSampler.set_epoch** +(epoch: int) +Sets the current epoch. + +**Parameters** +- **epoch** —— The current epoch number. + +## MultiLevelBatchSampler +A sampler for large-scale data with multi-level indexing. +#### function **MultiLevelBatchSampler.__init__** +(index_file: str, batch_size: int, rank: int = 0, seed: int = 8888) + +**Parameters** +- **index_file** —— The index file for multi-level data indexing. +- **batch_size** —— The size of a batch. +- **rank** —— The rank representing the current process number. +- **seed** —— The seed for random sampling, obtained globally from data.registry. + +#### function **MultiLevelBatchSampler.__iter__** +() +Iterator, each iteration yields the index of a sample. diff --git a/docs/en/scepter/solvers.md b/docs/en/scepter/solvers.md new file mode 100644 index 0000000..e3ae16f --- /dev/null +++ b/docs/en/scepter/solvers.md @@ -0,0 +1,488 @@ +# Solvers + +## Overview +The Solver defines a workflow for model training, validation, and testing processes. +Within the Solver, based on the settings in the configuration YAML file, modules required such as data, model, optimizer, and scheduler are initialized one by one. +Each specific task's custom Solver must inherit from BaseSolver. + +In some special scenarios, it is also necessary to initialize some custom modules. For example, Hooks are needed to record and save intermediate results during the training process; Metrics need to be defined for validation during training, and so on. + + +
+ +## Basic Usage + +Usage of subclass Solver upon inheritance: + +```python +from scepter.modules.solver.registry import SOLVERS +from scepter.modules.solver import BaseSolver + + +@SOLVERS.register_class() +class XxxSolver(BaseSolver): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) +``` + +The actual usage of starting the Solver can be found in ***scepter.modules/task/cate_recognition*** within ***run_train.py*** and ***run_inference.py***: + +```python +solver = SOLVERS.build(cfg.SOLVER, logger=std_logger) +``` + +
+ +## **scepter.modules.solver.BaseSolver** +The base class for Solvers is an abstract base class defined through the ABCMeta metaclass, which supports registration operations. Custom solvers should all be subclasses of this class and be registered accordingly. + +BaseSolver is a concrete example of a Solver implementation, showcasing how to write a Solver with and without the use of pytorch_lightning. + +In practice, usage varies widely, so the application of Solvers is quite flexible. Most of its member functions can be overwritten in subclasses or enhanced with additional useful features to meet specific needs. Alternatively, completely new functions can be written to replace existing ones. +
+ +### function **\_\_init\_\_** + +(cfg: *scepter.modules.utils.config.Config*, logger = None) -> None + +Initialize the Solver class, obtaining and defining some necessary parameters from the configuration (***cfg***). +Details of some parameters initialized in the **\_\_init\_\_** method can be seen in the code: + +**Configs** + +- **FILE_SYSTEM** —— (cfg) File system configuration, default None. +- **WORK_DIR** —— (str) Definition of the work directory. +- **LOG_FILE** —— (str) Location of the log file. +- **RESUME_FROM** —— (str) Path to intermediate results of the model from a previous training session. +- **MAX_EPOCHS** —— (int) Maximum number of training epochs. +- **ACCU_STEP** —— (int) When using ddp (distributed data parallel), the gradient accumulation steps for each process, default 1. +- **NUM_FOLDS** —— (int) Number of folds for training, default 1. +- **EVAL_INTERVAL** —— (int) Interval for evaluating the model, default 1. +- **EXTRA_KEYS** —— (List) The extra keys for metrics, default []. +- **TRAIN_DATA** —— (cfg) Training data configuration. +- **EVAL_DATA** —— (cfg) Validation data configuration. +- **TEST_DATA** —— (cfg) Test data configuration. +- **TRAIN_HOOKS** —— (List) Training hooks. +- **EVAL_HOOKS** —— (List) Validation hooks. +- **TEST_HOOKS** —— (List) Test hooks. +- **MODEL** —— (cfg) Model configuration. +- **OPTIMIZER** —— (cfg) Optimizer configuration. +- **LR_SCHEDULER** —— (cfg) Learning rate scheduler configuration. +- **METRICS** —— (List) Test Metrics. + +**Parameters** + +- **cfg** —— The Config used to build solver. +- **logger** —— Instantiated Logger to print or save log. + +
+ +### function **set_up_pre** + +() -> None + +Configure the environment, log path, and call construct_hook to initialize hooks, among other things. +Choose between this and pytorch_lightning; it is called when not using pytorch_lightning (use_pl=False). It needs to be called before initiating any other operations. + +
+ +### function **set_up** + +() -> None + +Configure data, model, metrics, optimizer, and the pytorch_lightning environment (if used). + + +
+ +### function **construct_data** + +() -> None + +The actual method for constructing data, which by default is called within self.set_up, including TRAIN_DATA, EVAL_DATA, TEST_DATA. The instantiated results are written into self.datas. + +
+ +### function **construct_hook** + +() -> None + +The actual method for constructing hooks, which by default is called within self.set_up_pre, includes TRAIN_HOOKS, EVAL_HOOKS, TEST_HOOKS. The instantiated results are stored in self.hooks_dict. + +
+ +### function **construct_model** + +() -> None + +The actual method for constructing the Model, which by default is called within self.set_up, resulting in the instantiated model being assigned to self.model. + +
+ +### function **model_to_device** + +() -> None + +The actual method for constructing Metrics, which by default is called within self.set_up, and the instantiated results are stored in self.metrics. + +
+ +### function **model_to_device** + +() -> None or MODELS + +The method for configuring the model, including model sharding and other distributed settings, is typically called within self.set_up by default. + +**Parameters** + +- **tg_model_ins** —— The model to be configured. If it is None, self.model will be used by default. + +**Returns** + +- **tg_model_ins** —— The configured model. If tg_model_ins is None, there is no return value. + +
+ +### function **init_opti** + +() -> None + +The method for configuring the optimizer, which by default is called within self.set_up. The instantiated optimizer and learning rate scheduler are assigned to self.optimizer and self.lr_scheduler, respectively. + +
+ +### function **solve** + +(epoch = None, every_epoch = False) -> None + +Execute the actual training, validation, or testing operations by calling self.solve_train, self.solve_eval, self.solve_test, etc. Additionally, call self.before_solve and self.after_solve before and after execution to perform ***Hook*** logging. + +**Parameters** + +- **epoch** —— The number of epochs to set. +- **every_epoch** —— Related to the implementation of self.solve_train, self.solve_eval, self.solve_test, and the configuration of the data. It indicates whether it is necessary to call them again for each epoch. For example, with self.solve_train, if the implementation involves calling it once to execute one epoch, then every_epoch should be True. If a single call will execute until all epochs are finished, then it should be False. + +
+ +### function **solve_train** + +() -> None + +To execute the training process, you would call self.run_train. Additionally, you would invoke self.before_epoch and self.after_epoch before and after the training execution to carry out ***Hook*** logging. + +
+ +### function **solve_eval** + +() -> None + +Invoke self.run_eval to execute evaluation. +Additionally, call self.before_epoch and self.after_epoch respectively before and after the execution to log ***Hook***. + +
+ +### function **solve_test** + +() -> None + +Invoke self.run_test to execute testing. +Additionally, call self.before_epoch and self.after_epoch respectively before and after the execution to log ***Hook***. + +
+ +### function **before_solve** + +() -> None + +Before actually executing run_xxx, perform ***Hook*** logging. + +
+ +### function **after_solve** + +() -> None + +After executing run_xxx, perform ***Hook*** logging again. + +
+ +### function **run_train** + +() -> None + +Iteratively call self.run_step_train to execute the training process for one epoch or all epochs. +Additionally, call self.before_all_iter and self.after_all_iter respectively at the beginning and end of the loop to log ***Hook***. +Moreover, call self.before_iter and self.after_iter respectively before and after each iteration when calling self.run_step_train to log ***Hook***. + +
+ +### function **run_eval** + +() -> None + +Iteratively call self.run_step_eval to execute the validating process for one epoch or all epochs. +Additionally, call self.before_all_iter and self.after_all_iter respectively at the beginning and end of the loop to log ***Hook***. +Moreover, call self.before_iter and self.after_iter respectively before and after each iteration when calling self.run_step_eval to log ***Hook***. + +
+ +### function **run_test** + +() -> None + +Iteratively call self.run_step_test to execute the testing process for one epoch or all epochs. +Additionally, call self.before_all_iter and self.after_all_iter respectively at the beginning and end of the loop to log ***Hook***. +Moreover, call self.before_iter and self.after_iter respectively before and after each iteration when calling self.run_step_test to log ***Hook***. + +
+ +### function **run_step_train** + +(batch_data, batch_idx = 0, step = None, rank = None) -> None + +Perform model training for a single batch. + +
+ +### function **run_step_eval** + +(batch_data, batch_idx = 0, step = None, rank = None) -> None + +Perform model inference for a single batch. + +
+ +### function **run_step_test** + +(batch_data, batch_idx = 0, step = None, rank = None) -> None + +Perform model testing for a single batch. + +
+ +### function **register_flops** + +(Dict data, keys = []) -> None + +Use fvcore to calculate the FLOPs (floating-point operations per second), and save the result in self._model_flops. + +**Parameters** + +- **data** —— The input data for the model, formatted as a key-value structure, where the value is either a Tensor or a List. By default, only the first element along the batch dimension is taken to ensure batch_size = 1. +- **keys** —— Specifies the use of the values corresponding to the elements in keys as inputs to the model. If keys are empty, then all values in data are taken. + +
+ +### function **before_epoch** + +() -> None + +Execute before the start of each epoch. Perform before running run_xxx. + +
+ +### function **before_all_iter** + +() -> None + +Execute before the start of each epoch. Perform before the loop that iteratively calls run_step_xxx. + +
+ +### function **before_iter** + +() -> None + +Execute before the start of each step. Perform before running run_step_xxx. + +
+ +### function **after_epoch** + +() -> None + +Execute after the start of each epoch. Perform after running run_xxx. + +
+ +### function **after_all_iter** + +() -> None + +Execute after the end of each epoch. Perform after the loop that iteratively calls run_step_xxx. + + +
+ +### function **after_iter** + +() -> None + +Execute after the start of each step. Perform after running run_step_xxx. + +
+ +### function **collect_log_vars** + +() -> OrderedDict + +Obtain the variables you need to save in logs + +**Returns** + +- **ret** —— Return the required variables. + +
+ + +# Hooks + +## Overview + +During the execution of the Solver, it is necessary to perform tasks such as printing logs, recording to Tensorboard, computing and updating gradients, saving intermediate model parameters, and saving test results, all of which require Hooks to execute. + +
+ +## Basic Usage + +Method for creating a new Hook: + +```python +from scepter.modules.solver.hooks.hook import Hook +from scepter.modules.solver.hooks.registry import HOOKS + + +@HOOKS.register_class() +class XxxHook(Hook): + def __init__(self, cfg, logger=None): + super(XxxHook, self).__init__(cfg, logger=logger) +``` + +
+ + +## **scepter.modules.solvers.hooks.Hook** +A standard Hook base class is defined with a detailed list of member functions. Subclasses that inherit from this class need to select and implement some of these member functions. + +The member functions include: + +### function **__init__** +(cfg: *scepter.modules.utils.config.Config*, logger = None) -> None +initialize logger。 + +### function **before_solve** +(solver) -> None + +Execute before starting to solve. + +### function **after_solve** +(solver) -> None + +Execute after starting to solve. + +### function **before_epoch** +(solver) -> None + +Execute before each epoch. + +### function **after_epoch** +(solver) -> None + +Execute after each epoch. + +### function **before_all_iter** +(solver) -> None + +Execute before each iteration. + +### function **after_all_iter** +(solver) -> None + +Execute after each iteration. + +### function **before_iter** +(solver) -> None + +Execute before each step. + +### function **after_iter** +(solver) -> None + +Execute after each step. + +## **scepter.modules.solver.hooks.CheckpointHook** +Before starting the solve, load the checkpoint. The path to load the model comes from the Solver's RESUME_FROM parameter. It depends on the load_checkpoint member function implemented in the Solver. + +After the end of each epoch, save the checkpoint. + +**Configs** + +- **PRIORITY** —— (int)Default is _DEFAULT_CHECKPOINT_PRIORITY = 300 +- **INTERVAL** —— (int)The interval of epochs between saving checkpoints, default is 1 +- **SAVE_NAME_PREFIX** —— (str)The prefix name for saving checkpoints, default is 'ldm_step' +- **SAVE_LAST** —— (bool)Whether to save the checkpoint of the last iteration, default is False, applied in after_iter. +- **SAVE_BEST** —— (bool)Whether to save the best checkpoint, default is True, applied in after_epoch. SAVE_BEST_BY must also be set, otherwise it defaults back to False. +- **SAVE_BEST_BY** —— (str)The metric used to judge the best checkpoint, by default, the larger the better. + +## **scepter.modules.solver.hooks.BackwardHook** +The actions executed after each iteration step. This includes backpropagation of the loss, configuration of the optimizer's steps, and so on. + +**Configs** + +- **PRIORITY** —— (int)_DEFAULT_BACKWARD_PRIORITY=0 +- **GRADIENT_CLIP** —— (int)In torch.nn.utils.clip_grad_norm_, the max_norm parameter defaults to -1. It does not take effect when it is less than or equal to 0. +- **ACCUMULATE_STEP** —— (int)The gradient accumulation step count is used to specify how many forward/backward passes to accumulate gradients over before performing an optimizer step (gradient descent update). +- **EMPTY_CACHE_STEP** —— (int)The max_steps parameter for torch.cuda.empty_cache() defaults to -1. When set to a value less than or equal to 0, the cache clearing operation will not be performed. +- +## **scepter.modules.solver.hooks.LogHook** +Log hook。 + +**Configs** + +- **PRIORITY** —— (int)__DEFAULT_LOG_PRIORITY=100 +- **SHOW_GPU_MEM** —— (bool)To check whether to print memory usage information. +- **LOG_INTERVAL** —— (int)The interval of steps for printing logs, default is 10. It does not take effect when it is less than or equal to 0. + +## **scepter.modules.solver.hooks.TensorboardLogHook** +Tensorboard log hook。 + +**Configs** + +- **PRIORITY** —— (int)__DEFAULT_LOG_PRIORITY=100 +- **LOG_DIR** —— (str)The path to store Tensorboard logs. +- **LOG_INTERVAL** —— (int)The interval of steps for printing logs, default is 1000. + +## **scepter.modules.solver.hooks.LrHook** +The hook for changing the learning rate. + +**Configs** + +- **PRIORITY** —— (int)_DEFAULT_LR_PRIORITY=200 +- **WARMUP_FUNC** —— (str)The method of warmup: only "linear" is supported, default is "linear". +- **WARMUP_EPOCHS** —— (int)The number of warmup epochs, default is 1. +- **WARMUP_START_LR** —— (float)the initial learning rate for warmup, default is 0.0001. +- **SET_BY_EPOCH** —— (bool)whether to set the learning rate once per epoch, default is True. If False, the learning rate is set at each step. + +## **scepter.modules.solver.hooks.DistSamplerHook** +Before each epoch begins, sample data for one epoch. + +**Configs** + +- **PRIORITY** —— (int)__DEFAULT_SAMPLER_PRIORITY=400 + +## **scepter.modules.solver.hooks.ProbeDataHook** +Used to print the intermediate (visualization) results of train/eval stored by the probe. + +**Configs** + +- **PRIORITY** —— (int)_DEFAULT_LR_PRIORITY=200 +- **PROB_INTERVAL** —— (int)The interval of steps for printing the probe data, default is 1000. + +## **scepter.modules.solver.hooks.SafetensorsHook** +The ***.safetensors*** format is used to store model files + +**Configs** + +- **PRIORITY** —— (int)_DEFAULT_LR_PRIORITY=200 +- **INTERVAL** —— (int)The interval of steps for saving .safetensors file, default is 1000. +- **SAVE_NAME_PREFIX** —— (str)The prefix name for saved files。 diff --git a/docs/en/scepter/tools.md b/docs/en/scepter/tools.md new file mode 100644 index 0000000..c7f3ad2 --- /dev/null +++ b/docs/en/scepter/tools.md @@ -0,0 +1,54 @@ +# Tool Components(Tools) + +Supports querying components registered with the framework and obtaining parameter templates. + +## Overview +1、Query module types:scepter.module_list +2、Query objects by module:scepter.objects_by_module +3、Query parameter configurations by object:scepter.configures_by_objects + +
+ +## Basic Usage + +```python +from scepter import module_list, objects_by_module, configures_by_objects + +# Query module types +module_list() +# Query object list by module +objects_by_module("BACKBONES") +# Query parameter configurations by object +configures_by_objects("BACKBONES", "ResNet3D_TAda") +``` +
+ +### function **module_list** +() + +**Returns** + +- **list** —— A list of module names. + +### function **objects_by_module** +(module_name: str) + +**Parameters** + +- **module_name** —— Module name + +**Returns** + +- **list** —— A list of module names. + +### function **get_module_object_config** +(module_name: str, object_name: str) + +**Parameters** + +- **module_name** —— Module name +- **object_name** —— Object name + +**Returns** + +- **str** —— Parameter template. diff --git a/docs/en/scepter/transforms.md b/docs/en/scepter/transforms.md new file mode 100644 index 0000000..f7cc653 --- /dev/null +++ b/docs/en/scepter/transforms.md @@ -0,0 +1,419 @@ +# Transforms + +Data pre-processing module + +## Overview + +Supports various data pre-processing methods: + +1. ***scepter.transforms.image*** +2. ***scepter.transforms.io*** +3. ***scepter.transforms.io_video*** +4. ***scepter.transforms.tensor*** +5. ***scepter.transforms.augmention*** +6. ***scepter.transforms.video*** +7. ***scepter.transforms.transform_xl*** +8. ***scepter.transforms.identity*** +9. ***scepter.transforms.compose*** + +
+ +## Basic Usage + +```python +# 以 scepter.modules.transform.image.RandomResizedCrop 为例 +from scepter.modules.transform.image import RandomResizedCrop +from scepter.modules.utils.config import Config +import PIL + +cfg = Config(load=False, + cfg_dict={"SIZE": 224, "RATIO": [3. / 4., 4. / 3.], "SCALE": [0.08, 1.0], "INTERPOLATION": "bilinear"}) +transform = RandomResizedCrop(cfg) +input_img = {"img": PIL.Image} +output_img = transform(input_img) +``` + +
+ +## **scepter.modules.transform.image** + +Some pre-processing methods used for images. + +
+ +### scepter.modules.transform.image.ImageTransform + +Initialize the ***ImageTransform*** class, obtain and define some necessary parameters from ***cfg***. + +**Parameters** + +- **INPUT_KEY** —— (str) input key or key list +- **OUTPUT_KEY** —— (str) output key or key list +- **BACKEND** —— (str) backend, choose from pillow, cv2, torchvision + +
+ +### scepter.modules.transform.image.RandomResizedCrop + +Randomly crop the image to a specified size. + +**Parameters** + +- **SIZE** —— (int) crop size +- **RATIO** —— (list) ratio +- **SCALE** —— (list) scale +- **INTERPOLATION** —— (str) interpolation + +
+ +### scepter.modules.transform.image.RandomHorizontalFlip + +Randomly horizontally flip the given image with a given probability(***P***). + +**Parameters** + +- **P** —— (float) probability + +
+ +### scepter.modules.transform.image.Normalize + +Normalize the image using ***mean*** and standard deviation(***std***). + +**Parameters** + +- **MEAN** —— (list) mean +- **STD** —— (list) std + +
+ +### scepter.modules.transform.image.ImageToTensor + +transform ***PIL.Image / numpy.ndarray / unit8*** to ***float32 tensor***. + +**Parameters** + +
+ +### scepter.modules.transform.image.Resize + +Resize the given image to the given ***Size***. + +**Parameters** + +- **INTERPOLATION** —— (str) interpolation +- **SIZE** —— (int) resized size + +
+ +### scepter.modules.transform.image.CenterCrop + +Crop the given image from the center. + +**Parameters** + +- **SIZE** —— (int) crop size + +
+ +### scepter.modules.transform.image.FlexibleResize + +Resize the given image to the given ***Size***. + +**Parameters** + +- **INTERPOLATION** —— (str) interpolation + +
+ +### scepter.modules.transform.image.FlexibleCenterCrop + +Center crop the given image to the given ***Size***. + +**Parameters** + +
+ +## **scepter.modules.transform.io** + +Some methods for reading images from local disk. + +
+ +### scepter.modules.transform.io.LoadPILImageFromFile + +Read a local image file into ***PIL.Image*** format. + +**Parameters** + +- **RGB_ORDER** —— (str) "RGB" or "BGR" + +
+ +### scepter.modules.transform.io.LoadCvImageFromFile + +Read a local image file into ***cv2*** format. + +**Parameters** + +- **RGB_ORDER** —— (str) "RGB" or "BGR" + +
+ +### scepter.modules.transform.io.LoadImageFromFile + +Read a local image file into a specific format. + +**Parameters** + +- **RGB_ORDER** —— (str) "RGB" or "BGR" +- **BACKEND** —— (str) "pillow" or "cv2" or "torchvision" + +
+ +### scepter.modules.transform.io.LoadImageFromFileList + +Read a set of input images into a specific format. + +**Parameters** + +- **RGB_ORDER** —— (str) "RGB" or "BGR" +- **BACKEND** —— (str) "pillow" or "cv2" or "torchvision" +- **FILE_KEYS** —— (list) The file keys for input + +
+ +## **scepter.modules.transform.io_video** + +Some methods for reading videos from local disk. + +
+ +### scepter.modules.transform.io_video.DecodeVideoToTensor + +Decode local video files into tensors. + +**Parameters** + +- **NUM_FRAMES** —— (int) decode frames number +- **TARGET_FPS** —— (int) decode frames fps, default is 30. +- **SAMPLE_MODE** —— (str) interval or segment sampling, default is interval +- **SAMPLE_INTERVAL** —— (int) sample interval between output frames for interval sample mode +- **SAMPLE_MINUS_INTERVAL** —— (float) wheather minus interval for interval sample mode +- **REPEAT** —— (str) number of clips to be decoded from each video + +
+ +### scepter.modules.transform.io_video.LoadVideoFromFile + +Decode and read local video files into a sequence of frames. + +**Parameters** + +- **NUM_FRAMES** —— (int) decode frames number +- **SAMPLE_TYPE** —— (str) sample type +- **CLIP_DURATION** —— (float) needed for 'interval' sampling type +- **DECODER** —— (str) video decoder name + +
+ +## **scepter.modules.transform.tensor** + +Some methods for processing tensors. + +
+ +### scepter.modules.transform.tensor.ToTensor + +Convert input data from other formats into tensors. + +**Parameters** + +- **KEYS** —— (list) keys of input data + +
+ +### scepter.modules.transform.tensor.Select + +Select some keys from the input data and output them. + +**Parameters** + +- **META_KEYS** —— (list) chosen keys of input data + +
+ +### scepter.modules.transform.tensor.Rename + +Rename the keys of the input data. + +**Parameters** + +- **IN_KEYS** —— (list) input data keys +- **OUT_KEYS** —— (list) output data keys + +
+ +## **scepter.modules.transform.augmention** + +Some methods for enhancing the colors in images. + +
+ +### scepter.modules.transform.augmention.ColorJitterGeneral + +Randomly adjust the brightness, contrast, and saturation of an image. + +**Parameters** + +- **BRIGHTNESS** —— (float) (float or tuple of float (min, max)): How much to jitter brightness +- **CONTRAST** —— (float) (float or tuple of float (min, max)): How much to jitter contrast +- **SATURATION** —— (float) (float or tuple of float (min, max)): How much to jitter saturation +- **HUE** —— (float) (float or tuple of float (min, max)): How much to jitter hue +- **GRAYSCALE** —— (float) probablitities for rgb-to-gray 0~1 +- **CONSISTENT** —— (bool) for video input whether the augment scale is consistent or not +- **SHUFFLE** —— (bool) shuffle the transform's order when there are multiple transforms +- **GRAY_FIRST** —— (bool) whether use grayscale or not +- **IS_SPLIT** —— (bool) whether randomly chance the channel as the gray results + +
+ +## **scepter.modules.transform.video** + +Some methods for processing videos. + +
+ +### scepter.modules.transform.video.VideoTransform + +To initialize a ***VideoTransform*** class and define the necessary parameters from a ***cfg***. + +**Parameters** + +- **INPUT_KEY** —— (str) input key or key list +- **OUTPUT_KEY** —— (str) output key or key list +- **BACKEND** —— (str) backend, choose from pillow, cv2, torchvision + +
+ +### scepter.modules.transform.video.RandomResizedCropVideo + +Perform random cropping of the video to a specified size. + +**Parameters** + +- **META_KEYS** —— (list) chosen keys of input data + +
+ +### scepter.modules.transform.video.CenterCropVideo + +Renaming the keys of the input data. + +**Parameters** + +- **SIZE** —— (int) crop size +- **RATIO** —— (list) ratio +- **SCALE** —— (list) scale +- **INTERPOLATION** —— (str) interpolation + +
+ +### scepter.modules.transform.video.RandomHorizontalFlipVideo + +Randomly flip the given video horizontally with a given probability. + +**Parameters** + +- **P** —— (float) probability + +
+ +### scepter.modules.transform.video.NormalizeVideo + +Normalize the video using the ***mean*** and standard deviation(***std***). + +**Parameters** + +- **MEAN** —— (list) mean +- **STD** —— (list) std + +
+ +### scepter.modules.transform.video.VideoToTensor + +transform ***PIL.Image / numpy.ndarray / unit8*** to ***float32 tensor***. + +**Parameters** + +
+ +### scepter.modules.transform.video.AutoResizedCropVideo + +Crop the given video from the center. + +**Parameters** + +- **SCALE** —— (list) scale + +
+ +### scepter.modules.transform.video.ResizeVideo + +Resize the given video to the specified dimensions. + +**Parameters** + +- **SCALE** —— (list) scale +- **INTERPOLATION** —— (str) interpolation + +
+ +## **scepter.modules.transform.transform_xl** + +Some methods for image processing in SDXL to obtain the desired coordinates. + +
+ +### scepter.modules.transform.transform_xl.FlexibleCropXL + +Crop an image and obtain its original size, target size, and cropping coordinates (top/left). + +**Parameters** + +- **INPUT_KEY** —— (str) input key or key list +- **OUTPUT_KEY** —— (str) output key or key list +- **BACKEND** —— (str) backend, choose from pillow, cv2, torchvision +- **SIZE** —— (int or list) crop size if 'image_size' not in meta + +
+ +## **scepter.modules.transform.identity** + +Some methods to process images and obtain the required coordinates in SDXL. + +
+ +### scepter.modules.transform.identity.Identity + +Return the image itself. + +**Parameters** + +
+ +## **scepter.modules.transform.compose** + +Combine various transform methods. + +
+ +### scepter.modules.transform.compose.Compose + +Combine the various transform objects from ***scepter.transforms*** into a pipeline. + +**Parameters** + +- **TRANSFORMS** —— (list) transform config list + +
diff --git a/docs/en/scepter/utils/file_clients.md b/docs/en/scepter/utils/file_clients.md new file mode 100644 index 0000000..78a5d87 --- /dev/null +++ b/docs/en/scepter/utils/file_clients.md @@ -0,0 +1,526 @@ +# File System + +This is the File System Module, designed to handle file transfer functionalities. + +## Overview + +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 + +
+ +## Basic Usage + +```python +from scepter.utils.file_system import FS +from scepter.utils.config import Config + +fs_cfg = Config(load=False, cfg_dict={ + "NAME": "AliyunOssFs", + # ENDPOINT DESCRIPTION: the oss endpoint TYPE: str default: '' + "ENDPOINT": "xxxxx", + # BUCKET DESCRIPTION: the oss bucket TYPE: str default: '' + "BUCKET": "xxxxx", + # OSS_AK DESCRIPTION: the oss ak TYPE: str default: '' + "OSS_AK": "xxxxx", + # OSS_SK DESCRIPTION: the oss sk TYPE: str default: '' + "OSS_SK": "xxxxx", + # TEMP_DIR DESCRIPTION: default is None, means using system cache dir and auto remove! If you set dir, the data will be saved in this temp dir without autoremoving default. TYPE: NoneType default: None + "TEMP_DIR": "cache", + # AUTO_CLEAN DESCRIPTION: when TEMP_DIR is not None, if you set AUTO_CLEAN to True, the data will be clean automatics. TYPE: bool default: False + "AUTO_CLEAN": False +}) + +fs_prefix = FS.init_fs_client(fs_cfg, logger=None) +with FS.get_from("xxxxx", wait_finish=True) as local_object: +# do sth. using local_object here. +# Download multiple files at once. +generator = FS.get_batch_objects_from(["xxx", "xxxx"]) +for local_path in generator: + print(local_path) +``` +
+ +## **scepter.modules.utils.file_system.FileSystem** + +By building various File IO Handlers, it supports read and write operations for different types of files. +
+ +### function **\_\_init\_\_** + +() + +**Parameters** + +
+ +### function **init_fs_client** + +( cfg: *scepter.modules.utils.config.Config* = None, logger = None ) -> str + +The fs_client is instantiated through the cfg parameter and stored in the self._prefix_to_clients attribute, allowing access to the corresponding fs_client via the prefix. + +**Parameters** + +- **cfg** —— The Config used to build fs_client. If None, use LocalFs as default. +- **logger** —— Instantiated Logger to print or save log. + +**Returns** + +- *str* —— The prefix of instantiated fs_client + +
+ +### function **get_fs_client** + +( target_path: *str*, safe: *bool* = False ) + +Retrieve the corresponding fs_client based on the prefix of the target_path. + +**Parameters** + +- **target_path** —— Target file path. +- **safe** —— In safe mode, return a copy of the client; otherwise, return the client itself. + +**Returns** + +- *BaseFs* —— Instantiated fs_client. + +
+ +### function **get_from** + +( target_path: *str*, local_path: *str* = None, wait_finish: *bool* = False ) -> str + +Download remote files to the local system. + +**Parameters** + +- **target_path** —— Remote file path. +- **local_path** —— Local file path; if None, use the cache path. +- **wait_finish** —— if True, only the card 0 of each machine will download the data, and the other cards will wait for the download by card 0 to finish. + +**Returns** + +- *str* —— Local save file path. + +
+ +### function **get_dir_to_local_dir** + +( target_path: *str*, local_path: *str* = None, wait_finish: *bool* = False, timeout: *int* = 3600, worker_id: *int* = 0 ) -> str + +Download a folder from a remote path to the local system. + +**Parameters** + +- **target_path** —— Remote folder path. +- **local_path** —— Local folder path; if None, use the cache path. +- **wait_finish** —— if True,only the 0-card of each machine downloads the data, while the other cards wait for the 0-card to finish downloading. +- **timeout** —— Download timeout duration. +- **worker_id** —— Deprecated + +**Returns** + +- *str* —— 本地文件夹路径 + +
+ +### function **get_object** + +(target_path: *str*) -> bytes + +Read a remote file into memory + +**Parameters** + +- **target_path** — Target file path + +**Returns** + +- *bytes* — Binary data of the target file + +
+ +### function **put_object** + +(local_data: *bytes*, target_path: *str*) -> bool + +Upload a data stream to a specified file + +**Parameters** + +- **local_data** — Local data stream + +- **target_path** — Target file path + +**Returns** + +- *bool* — Whether the upload was successful + +
+ + +### function **delete_object** + +(target_path: *str*) -> bool + +Delete the target file + +**Parameters** + +- **target_path** — Target file path + +**Returns** + +- *bool* — Whether the deletion was successful + +
+ +### function **get_batch_objects_from** + +(target_path_list: *list[str]*, wait_finish: *bool*) -> Iterator[str] + +Batch download files + +**Parameters** + +- **target_path_list** — List of files to download + +**Returns** + +- *Iterator[str]* — Iterator for local file paths + +-
+ +### function **put_batch_objects_to** + +(local_path_list: *list[str]*, target_path_list: *list[str]*, wait_finish: *bool*) -> Iterator[tuple[str, str]] + +Batch upload files + +**Parameters** + +- **local_path_list** — List of files to upload + +- **target_path_list** — List of target file paths + +**Returns** + +- *Iterator[tuple[str, str]]* — Returns pairs of local file and target file paths + +
+ +### function **get_object_stream** + +(target_path: *str*, start: *int*, size: *int*, delimiter: *str*) -> bytes, int + +Batch upload files + +**Parameters** + +- **target_path** — Target file + +- **start** — Starting character position of the target stream + +- **size** — Size of the target stream starting from the character position + +- **delimiter** — Delimiter character for the end of the target stream + +**Returns** + +- *bytes, int* — Returns the data stream bytes and the end character position + +
+ +### function **get_object_chunk_list** + +(target_path: *str*, chunk_num: *int* = 1, delimiter: *str* = None) -> list[bytes] + +Get a remote file and divide it into chunks + +**Parameters** + +- **target_path** — Target file path + +-**chunk_num** — Number of chunks + +- **delimiter** — Delimiter to ensure the data downloaded is a complete record and not truncated in the middle + +**Returns** + +- *list[bytes]* — Chunked data + +
+ +### function **get_url** + +(target_path: *str*, set_public=False, lifecycle: *int* = 360000) -> str + +Get the URL of a remote file (only supports AliyunOssFs) + +**Parameters** + +- **target_path** — Target file path + +- **lifecycle** — Valid duration + +- **set_public** — Whether to provide a public link + +**Returns** + +- *str* — URL of the target file + +
+ +### function **put_to** + +(target_path: *str*) + +Supports uploading a local file to a remote path + +**Parameters** + +- **target_path** — Remote file path + +**Returns** + +- **None** + +
+ +```python +# Used as a context manager +with FS.put_to(target_path) as local_path: + # some operations on local_path. +``` + +
+ +### function **put_object_from_local_file** + +(local_path: *str*, target_path: *str*) -> bool + +Push a local file to a remote path + +**Parameters** + +- **local_path** — Local file path + +- **target_path** — Remote file path + +**Returns** + +- *bool* — Whether the upload was successful + +
+ +### function **put_dir_from_local_dir** + +(local_dir: *str*, target_dir: *str*) -> bool + +Push a local directory to a remote path + +**Parameters** + +- **local_dir** — Local directory path + +- **target_dir** — Remote directory path + +**Returns** + +- *bool* — Whether the upload was successful + +
+ +### function **add_target_local_map** + +(target_dir: *str*, local_dir: *str*) -> None + +Save the mapping relationship between the remote directory path and the local directory path as key-value pairs to self._target_local_mapper + +**Parameters** + +- **target_dir** — Remote directory path + +- **local_dir** — Local directory path + +**Returns** + +- *None* + +
+ +### function **make_dir** + +(target_dir: *str*) -> bool + +Create a remote directory + +**Parameters** + +- **target_dir** — Remote directory path + +**Returns** + +- *bool* — Whether the creation was successful + +
+ +### function **exists** + +(target_path: *str*) -> bool + +Check if the target path exists + +**Parameters** + +- **target_path** — Remote file path + +**Returns** + +- *bool* — Whether it exists + +
+ +### function **map_to_local** + +(target_path: *str*) -> str, bool + +Map the remote file path to a local path + +**Parameters** + +- **target_path** — Remote file path + +**Returns** + +- *str* — Local file path + +- *bool* — Whether the local file is a temporary file + +
+ +### function **walk_dir** + +(target_dir: *str*, recurse=True) -> Iterator + +Get the file list under the remote directory + +**Parameters** + +- **target_dir** — Remote directory path + +- **recurse** — Whether to traverse subdirectories, default is True + +**Returns** + +- *Iterator* — List of subfile paths + +
+ +### function **is_local_client** + +(target_path: *str*) -> bool + +Determine if the target file client is LocalFs + +**Parameters** + +- **target_path** — Target file path + +**Returns** + +- *bool* — Whether the client is LocalFs + +
+ +### function **size** + +(target_path: *str*) -> int + +Determine the size of the target file + +**Parameters** + +- **target_path** — Target file path + +**Returns** + +- *int* — Size of the target file + +
+ +### function **isfile** + +(target_path: *str*) -> bool + +Determine if the target path is an object + +**Parameters** + +- **target_path** — Target file path + +**Returns** + +- *bool* — Whether the target path is an object + +
+ +### function **isdir** + +(target_path: *str*) -> bool + +Determine if the target path is a directory + +**Parameters** + +- **target_path** — Target file path + +**Returns** + +- *bool* — Whether the target path is a directory + +
+ +## **scepter.modules.utils.file_clients.AliyunOssFs** + +```yaml +NAME: AliyunOssFs +TEMP_DIR: None +AUTO_CLEAN: False +ENDPOINT: +BUCKET: +OSS_AK: +OSS_SK: +PREFIX: "" +WRITABLE: True +CHECK_WRITABLE: False +RETRY_TIMES: 10 +``` + +
+ +## **scepter.modules.utils.file_clients.LocalFs** + +```yaml +NAME: LocalFs +TEMP_DIR: None +AUTO_CLEAN: False +``` + +
+ +## **scepter.modules.utils.file_clients.HttpFs** + +```yaml +NAME: HttpFs +TEMP_DIR: None +AUTO_CLEAN: False +RETRY_TIMES: 10 +``` + +
diff --git a/docs/en/scepter/utils/utils.md b/docs/en/scepter/utils/utils.md new file mode 100644 index 0000000..96e0cec --- /dev/null +++ b/docs/en/scepter/utils/utils.md @@ -0,0 +1,939 @@ +# Dependency Components (Utils) + +Relies on SDKs, which are used to organize modules and SDKs that are frequently reused throughout the framework and to aggregate them based on functional relevance. + +## 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) + +### Basic Usage + +```python +from scepter.utils.config import Config +# Initialize Config object from a dict +fs_cfg = Config(load=False, cfg_dict={"NAME": "LocalFs"}) +print(fs_cfg.NAME) +# Initialize Config object from a json file +import json +json.dump({"NAME": "LocalFs"}, open("examples.json", "w")) +fs_cfg = Config(load=True, cfg_file="examples.json") +print(fs_cfg.NAME) +# Initialize Config object from a yaml file +import yaml +yaml.dump({"NAME": "LocalFs"}, open("examples.yaml", "w")) +fs_cfg = Config(load=True, cfg_file="examples.yaml") +print(fs_cfg.NAME) +# Initialize Config object from an argparse object, in this mode cfg parameters are required, otherwise an error will be thrown. +import argparse +parser = argparse.ArgumentParser( + description="Argparser for Cate process:\n" +) +parser.add_argument( + "--stage", + dest="stage", + help="Running stage!", + default="train" +) +fs_cfg = Config(load=True, parser_ins=parser) +print(fs_cfg.args) +``` + +
+ +### function **__init__** + +(cfg_dict: dict = {}, load = True, cfg_file = None, logger = None, parser_ins: argparse.ArgumentParser = None) + +**Parameters** + +- **cfg_dict** — A dict containing parameters, default is {}. + +- **load** — When True, it indicates parameters need to be loaded from a file or argparse. + +- **cfg_file** — Supports loading parameters from json or yaml files. + +- **logger** — Logging instance, if None, a default logging instance to stdio will be initialized. + +- **parser_ins** — An argparse instance, default includes cfg parameter for passing in a parameter file. + +-- parser_ins will by default include the following system parameters: + + - cfg(--cfg) used to specify the parameter file location + + - local_rank(--local_rank) the default parameter read by torchrun, default is 0, can be ignored + + - launcher(-l) the method for starting the code, default is spawn, alternative is torchrun + + - data_online(-d) set global data not to be persisted to disk, should be set on pai clusters + + - share_storage(-s) set whether global data download is on a shared file system, like nas. When set, it implies the file system is shared across nodes, and only needs downloading at rank=0; when not set, it means data is downloaded on different nodes, and should only be downloaded when device_id=0. + +### function **dict_to_yaml** + +(module_name: str, name: str, json_config: dict, set_name: bool = False) + +**Parameters** + +- **module_name** — The module name, used at the start of the template to explain which module's template it is. + +- **name** — The default name for the Name field. + +- **json_config** — Parameter description, needs to satisfy {} (indicating dependency on a sub-module), [] (dependency on multiple sub-modules), {"value":"", "description":""} (leaf parameter value). + +- **set_name** — Whether to set the Name field. + +**Returns** + +- **str** — Template text + +## 2. Path sdk (scepter.modules.utils.directory) +Some commonly used path functions +### Basic Usage +```python +from scepter.utils.directory import osp_path +# Automatically join paths based on the path prefix +prefix = "xxxx" +data_file = "example_videos/1.mp4" +# Outputs as xxxx/example_videos/1.mp4 +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 +# 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 +# Get the md5 code of the text/path 34a447fb46d0b786a3999c9dad01d470 +print(get_md5(data_file)) +``` +
+ +### function **osp_path** + +( prefix: str, data_file: str ) -> str + +Automatically join paths based on the path prefix + +**Parameters** + +- **prefix** —— Path prefix. +- **data_file** —— File path. + +**Returns** + +- **str** —— Joined path after concatenation + +### function **get_relative_folder** + +( abs_path: str, keep_index: int = -1 ) -> str + +Get the folder path at a specified level according to the path + +**Parameters** + +- **abs_path** —— File path. +- **keep_index** —— Level to keep, -1 for the last level, -2 for the second last level. + +**Returns** + +- **str** —— Parsed path after resolution + +### function **get_md5** + +( ori_str: str) -> str + +Get the md5 code based on a string/path + +**Parameters** + +- **ori_str** —— File path or string. + +**Returns** + +- **str** —— md5 code + +## 3. PyTorch Distributed(scepter.modules.utils.distribute) +PyTorch distributed initialization SDK. By using this SDK, users can avoid focusing on the implementation details of PyTorch's distributed initialization. +### Basic Usage + +```python +from scepter.utils.distribute import we +from scepter.utils.config import Config + +cfg = Config(cfg_dict={}, load=False) + + +def fn(): + pass + + +print(we) +# Launch task +we.init_env(cfg, fn, logger=None) +``` +
+ +### class **Workenv** + +This is a class used to uniformly manage the running environment. It is usually not necessary to initialize this class. In scepter.modules.utils.distribute, +a global instance, 'we', will be initialized to manage some key flag variables. +- Specific explanations of some parameters of we are as follows: + - initialized marks whether the PyTorch process group has been initialized, default is False. + - is_distributed marks whether it is currently running in distributed mode, default is False. + - sync_bn marks whether to use sync_bn, default is False. + - rank marks the current process's rank, default is 0. + - world_size marks the total number of processes, default is 1. + - device_id marks the current device ID being used, default is 0. + - device_count marks the total number of devices in the current environment, default is 1. + - use_pl marks whether the pytorch_lighting engine is used in the current environment, default is False. + - launcher marks the method of starting the environment, default is spawn. + - data_online marks whether the io part of the data in the current environment is persisted to disk, default is False. + - share_storage marks whether different nodes in the current environment use the same file system, such as nas, default is False. +### function **we.init_env** + +( config: scepter.modules.utils.config.Config, fn: function, logger: logging.Logger = None ) + +As the entry point for executing any task. + +**Parameters** + +- **config** —— The passed instance of parameters. +- **fn** —— The function that needs to be executed. +- **logger** —— A standard logging instance. + +### function **we.get_env** +() -> dict + +Retrieve all class-internal parameters of we, stored in the form of a dict. + + +### function **we.set_env** +(we_env: dict) + +Reset all class-internal parameters of we, using a dict as input. + +**Parameters** + +- **we_env** —— A dict, each key represents an internal class variable. + +### function **get_dist_info** +() -> int, int + +Obtain the environment's rank/world size, which is directly acquired through torch's methods, typically used when the environment is not initialized with we.init_env. + +**Returns** + +- **rank** —— The rank of the current process, default is 0. +- **world_size** —— The total number of processes in the current environment, when in single-process mode it is 1. + +### function **gather_data** +(data: [list, dict, tensor, object] ) -> data + +Using scepter.distributed.all_gather to collect any instance, and merge it into a summarized instance on rank=0 process. + +**Parameters** + - **data** —— Supports dict/list, where elements can be any instance or tensor. + +**Returns** + - **data** —— A summarized data with the same structure as the input data. + +### function **gather_list** +(data: [list] ) -> data + +Using scepter.distributed.all_gather to collect any list instance, and merge it into a summarized instance on rank=0 process. + +**Parameters** + - **data** —— Supports list, where elements can be any instance or tensor. + +**Returns** + - **data** —— A summarized data with the same structure as the input data. + +### function **gather_picklable** +(data: [object] ) -> data + +Using scepter.distributed.all_gather to collect any picklable instance, and merge it into a summarized instance on rank=0 process. + +**Parameters** + - **data** —— A serializable instance. + +**Returns** + - **data** —— A summarized data with the same structure as the input data. + +### function **broadcast** +(tensor: **torch.Tensor**, src: **str**, group: **list** ) + +An optimized version of torch.distributed.broadcast, automatically checks if it is a distributed environment. + +**Parameters** + - **tensor** —— The tensor to be broadcast. + - **src** —— The source device for broadcasting. + - **group** —— The group for broadcasting. + +**Returns** + - **data** —— A summarized data with the same structure as the input data. +* Other functions such as barrier, all_reduce, reduce, send, recv, isend, irecv, scatter have also been adapted for this operation. + + +### function **gather_gpu_tensors** +(tensor: torch.Tensor ) -> tensor: torch.Tensor + +Using torch.distributed.all_gather to collect GPU tensors, and merge then transfer them to the CPU on rank=0. +Since cloning is involved, this may cause additional GPU memory waste. + +**Parameters** + - **tensor** —— The GPU tensor input. + +**Returns** + - **tensor** —— The output tensor on the CPU for process rank=0. + +## 4. 模型导出sdk(scepter.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 + +save_develop_model_multi_io( + model, + input_size, + input_type, + input_name, + output_name, + limit, + save_onnx_path=None, + save_pt_path=None +) +``` +
+ +### function **save_develop_model_multi_io** +(model: torch.nn.Module, input_size: list, input_type: list, input_name: list, +output_name: list, limit: list, save_onnx_path: str = None, save_pt_path: str = None) -> pt_module, onnx_module + +Supports importing and exporting models with multiple inputs and outputs + +**Parameters** + - **model** —— The model instance to be exported. + - **input_size** —— A list where each tuple contains the shape information of the data, such as [[1, 3, 224, 224]]. + - **input_type** —— A list where each tuple contains the type information of the data, corresponding to input_size, with possible values ("float32", "float16", "int8", "int16", "int32", "int64"). For example, ["float32"]. + - **input_name** —— A list used to name each input variable for ONNX, such as ["image"], corresponding to the above input_size and input_type. + - **output_name** —— A list used to name each output variable for ONNX, such as ["output"] + - **limit** —— A list where each tuple defines the upper and lower bounds for that input, such as [[-1, 1]], representing that the input tensor for the image is between -1 and 1. + - **save_onnx_path** —— If not None, the ONNX model will be exported and stored at this location. + - **save_pt_path** —— If not None, the TorchScript model will be exported and stored at this location. + + + +**Returns** + - **tensor** —— The output tensor on the CPU for process rank=0. + +## 5. 文件系统sdk(scepter.utils.file_system) +Refer to [file_clients](file_clients.md) + +## 6. Logging SDK(scepter.utils.logger) +Used to instantiate a standard logging instance for printing information. + +### Basic Usage + +```python +from scepter.utils.logger import get_logger, init_logger + +std_logger = get_logger(name="std_torch") +init_logger(std_logger, log_file="", dist_launcher="pytorch") +``` +
+ +### function **get_logger** +(name: str) -> logger + +Retrieve a logging instance. + +**Parameters** + - **name** —— The log prefix; it will be printed first every time the logger prints. + +**Returns** + - **logger** —— Returns a logging instance. + +### function **init_logger** +(in_logger: logger, log_file: str) -> logger + +Re-initialize a logging instance, which can assign a file for output storage. + +**Parameters** + - **in_logger** —— The existing logging instance. + - **log_file** —— The desired file location for storage. + - **dist_launcher** —— No longer important, deprecated + +### function **as_time** +(s: int) -> str + +Convert time in seconds s to the standard format of xxx days xxx hours xxx mins xxx secs + +**Parameters** + - **s** —— Represents the number of seconds s. + +**Returns** + - **str** —— Formatted output. + +### function **time_since** +(since: int, percent: float) -> str + +Calculate the time remaining until completion based on the current usage time and percentage. + +**Parameters** + - **since** —— Represents the current elapsed time. + - **percent** —— Represents the percentage of completion. + +**Returns** + - **str** —— Formatted output. + +## 7. Video Processing SDK (scepter.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 ( + VideoReaderWrapper, EasyVideoReader, FramesReaderWrapper +) +``` +
+ +### function **do_frame_sample** +(sampling_type: str, vid_len: int, vid_fps: int, num_frames: int, kwargs) -> list + +Get a frame sampler for videos. + +**Parameters** + - **sampling_type** —— The type of sampler, currently supports UniformSampler (uniform sampler), IntervalSampler (interval sampler), SegmentSampler (segment sampler). + - **vid_len** —— Video length. + - **vid_fps** —— Frame rate of the video. + - **num_frames** —— Number of frames in the video. + - **kwargs** —— Required parameters for the corresponding sampler, refer to the source code of the corresponding sampler. + +**Returns** + - **list** —— Sampled frames result. + +### class **VideoReaderWrapper** +A standard class for reading videos, with the underlying decoder being decord. + +#### function **VideoReaderWrapper.__init__** +(video_path: str) + +Initialize a video instance. + +**Parameters** + - **video_path** —— Video link. + +#### function **VideoReaderWrapper.len** +() -> int + +Get the total number of video frames. + +**Returns** + - **int** —— Number of video frames. + +#### function **VideoReaderWrapper.fps** +() -> float + +Get video frame rate. + +**Returns** + - **float** —— Video frame rate. + +#### function **VideoReaderWrapper.duration** +() -> float + +Get video duration. + +**Returns** + - **float** —— Video duration. + +#### function **VideoReaderWrapper.sample_frames** +(decode_list: torch.Tensor) -> torch.Tensor + +Tensor Get frame data based on frame numbers. + +**Parameters** + - **decode_list** —— List of sampled frame numbers. +**Returns** + - **tensor** —— Data tensor. + +### class **FramesReaderWrapper** +Reads frame data in order from a given fully decoded frame folder. + +#### function **FramesReaderWrapper.__init__** +(frame_dir: str, extract_fps: float, suffix: str) + +Initialize a video instance. + +**Parameters** + - **frame_dir** —— The frame folder. + - **extract_fps** —— FPS for extracting frames. + - **suffix** —— Suffix for the frame files, default is jpg. + +#### function **FramesReaderWrapper.len** +() -> int + +Get the total number of video frames. + +**Returns** + - **int** —— Number of video frames. + +#### function **FramesReaderWrapper.fps** +() -> float + +Get video frame rate. + +**Returns** + - **float** —— Video frame rate. + +#### function **FramesReaderWrapper.duration** +() -> float + +Get video duration. + +**Returns** + - **float** —— Video duration. + +#### function **FramesReaderWrapper.sample_frames** +(decode_list: torch.Tensor) -> torch.Tensor + +Get frame data based on frame numbers. + +**Parameters** + - **decode_list** —— List of sampled frame numbers. +**Returns** + - **tensor** —— Data tensor. + +### class **EasyVideoReader** +Used for reading, sampling, and preprocessing long videos. + +#### function **EasyVideoReader.__init__** +(video_path: str, num_frames: int, clip_duration: Union[float, Fraction, str], +overlap: Union[float, Fraction, str] = Fraction(0), transforms: Optional[Callable] = None) + +Initialize a video instance. + +**Parameters** + - **video_path** —— Video link. + - **num_frames** —— Number of video frames. + - **clip_duration** —— Length of each clip. + - **overlap** —— Proportion of overlap between clips. + - **transforms** —— Preprocessing operators. + +#### function **EasyVideoReader.__iter__** +() -> int + +Iterator + +#### function **EasyVideoReader.__next__** +() -> float + +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) +Used for managing various registered classes. + +### Basic Usage + +```python +from scepter.utils.registry import Registry +from scepter.utils.config import Config + +MODELS = Registry('MODELS') + + +@MODELS.register_class() +class ResNet(object): + pass + + +config = Config(load=False, cfg_dict={"NAME": "ResNet"}) +resnet = MODELS.build(config) +``` +
+ +### class **Registry** +Registry + +#### function **Registry.__init__** +(name: str, build_func: function = None, common_para: Config = None, allow_types: tuple = ("class", "function")) + +Initialize the registry module instance + +**Parameters** + - **name** —— Module name. + - **build_func** —— The function called when building the module. + - **common_para** —— Common parameters under this module. + - **allow_types** —— The types of classes or functions allowed to be registered in this module, by default, registration of both is allowed. + +#### function **Registry.build** +(cfg: Config, logger: logger = None, kwargs) -> cls_obj + +Build an instance of the target class + +**Returns** + - **cls_obj** —— An instance of a specific class. + +#### function **Registry.register_class** +(name: str) + +Register a class + +**Returns** + - **name** —— Registration name. + +#### function **Registry.register_function** +(name: str) + +Register a function + +**Returns** + - **name** —— Registration name. + +## 9. Data SDK(scepter.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 + +data = {"a": torch.Tensor([0])} +transfer_data_to_numpy(data) +transfer_data_to_cpu(data) +transfer_data_to_cuda(data) +``` +
+ + +#### function **transfer_data_to_numpy** +(data: list/dict of torch.Tensor) -> (data: list/dict of numpy.ndarray) + +Transfer data to numpy + +**Parameters** + - **data** —— Stored as a list/dict of torch.Tensor. +**Returns** + - **data** —— Stored as a list/dict of numpy.ndarray, consistent with the input format. + +#### function **transfer_data_to_cpu** +(data: list/dict of torch.Tensor(cuda)) -> (data: list/dict of torch.Tensor(cpu)) + +Transfer data from GPU to CPU + +**Parameters** + - **data** —— Stored as a list/dict of torch.Tensor[CUDA]. +**Returns** + - **data** —— Stored as a list/dict of torch.Tensor[CPU], consistent with the input format. + +#### function **transfer_data_to_cuda** +(data: list/dict of torch.Tensor(cpu)) -> (data: list/dict of torch.Tensor(cuda)) + +Transfer data from CPU to GPU + +**Parameters** + - **data** —— Stored as a list/dict of torch.Tensor[CPU]. +**Returns** + - **data** —— Stored as a list/dict of torch.Tensor[CUDA], consistent with the input format. + +## 10. Model SDK(torch.utils.model) +Used for operations such as loading and evaluating models + +### Basic Usage + +```python +import torch +from scepter.utils.model import move_model_to_cpu, load_pretrained, + count_params, init_weights +``` +
+ + +#### function **move_model_to_cpu** +(params: list/dict of torch.Tensor[cuda]) -> (data: torch.Tensor[cpu]) + +Move parameter data from GPU to CPU. + +**Parameters** + - **params** —— Stored as OrderedDict of torch.Tensor[cuda]. +**Returns** + - **params** —— Stored as torch.Tensor[cpu], consistent with the input format. + +#### function **load_pretrained** +(model: torch.nn.Module, path: str, map_location="cpu", logger=None, + sub_level=None) + +Load parameters into the model. + +**Parameters** + - **model** —— The torch.nn.Module model instance. + - **path** —— Pretrained model parameters. + - **map_location** —— cpu/cuda。 + - **logger** —— Standard logging instance. + - **sub_level** —— For example, when using DDP, sub-level indexing might be needed. + + +#### function **count_params** +(model: torch.nn.Module) -> (float) + +Count the total parameters of the model. + +**Parameters** + - **model** —— The torch.nn.Module model instance. +**Returns** + - **float** —— The quantity of model parameters (number of floating-point values). + +#### function **init_weights** +(model: torch.nn.Module) + +Initialize the parameters of the model modules. + +**Parameters** + - **module** —— The torch.nn.Module model instance. + +## 11. Sampler SDK(scepter.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, + EvalDistributedSampler, MultiLevelBatchSampler, MixtureOfSamplers +``` +
+ + +#### class **MultiFoldDistributedSampler** + +Multi-fold sampler, supports repeating multiple rounds of data within one epoch. + +#### function **MultiFoldDistributedSampler.__init__** + +( dataset: torch.data.dataset, num_folds=1, num_replicas=None, rank=None, shuffle=True) + +**Parameters** + +- **dataset** —— An instance of torch.data.dataset class. +- **num_folds** —— Int, indicates the number of times the data is repeated. +- **num_replicas** —— Indicates the number of data partitions, usually consistent with world-size. +- **rank** —— Indicates the current process number. +- **shuffle** —— Whether to shuffle the data. + +#### function **MultiFoldDistributedSampler.__iter__** +() + +Iterator, each iteration returns an index of a sample. + +#### function **MultiFoldDistributedSampler.set_epoch** +(epoch: int) + +Set the current epoch. + +**Parameters** + +- **epoch** —— The current epoch. + + +#### class **EvalDistributedSampler** + +A sampler for testing, when not using padding mode, it will be observed that the last rank has fewer data than other ranks. + +#### function **EvalDistributedSampler.__init__** + +( dataset: torch.data.dataset, num_replicas: Optional[int] =None, rank: Optional[int] =None, padding: bool =False) + +**Parameters** + +- **dataset** —— An instance of torch.data.dataset class. +- **num_replicas** —— Indicates the number of data partitions, usually consistent with world-size. +- **rank** —— Rank indicates the current process number. +- **padding** —— Whether the data needs to be padded, if padded it can ensure the last rank has the same amount of data as the other ranks. + +#### function **EvalDistributedSampler.__iter__** +() + +Iterator, each iteration returns an index of a sample. + +#### function **EvalDistributedSampler.set_epoch** +(epoch: int) + +Set the current epoch. + +**Parameters** + +- **epoch** —— The current epoch. + +#### class **MultiLevelBatchSampler** + +A sampler for multi-level indexing of large-scale data. + +#### function **MultiLevelBatchSampler.__init__** + +(index_file: str, batch_size: int, rank: int =0, seed: int = 8888) + +**Parameters** + +- **index_file** —— Index file for multi-level data indexing. +- **batch_size** —— The size of a batch. +- **rank** —— Rank indicates the current process number. +- **seed** —— Random sampling seed, obtained globally from data.registry. + +#### function **MultiLevelBatchSampler.__iter__** +() + +Iterator, each iteration returns an index of a sample. + + +#### class **MixtureOfSamplers** + +A sampler for multi-level indexing of large-scale data. + +#### function **MixtureOfSamplers.__init__** + +(samplers: list(sampler), probabilities: list(float), rank: int =0, seed: int = 8888) + +**Parameters** + +- **samplers** —— A list of samplers used for mixing. +- **probabilities** —— The probability of each sampler. +- **rank** —— Rank indicates the current process number. +- **seed** —— Random sampling seed, obtained globally from data.registry. + +#### function **MixtureOfSamplers.__iter__** +() + +Iterator, each iteration returns an index of a sample. + +## 12. Prober SDK(scepter.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 + + +class TestModel(BaseModel): + def forward(self, data): + self.register_probe(data) + # Test ProbeData example + view_distribute + self.register_probe( + {"data_key_dist": ProbeData(data["data_key"], view_distribute=True), + "data_folder": ProbeData(data["data_folder"], view_distribute=True)} + ) + + +class TestModel2(BaseModel): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.test_model = TestModel(cfg, logger=logger) + + def forward(self, data): + # Test list of str done + # Test dict of str done + # Test dict of number done + # Test list of number done + # Test number done + # Test str done + # Test np.array done + # Test list of np.ndarray must manually create ProbeData done + # Test 2D image must manually create ProbeData + # Test 3D image must manually create ProbeData + # Test 3D multiple 2D images must manually create ProbeData + # Test 3D list image must manually create ProbeData + # Test 4D Array image done + # Test 4D Array image save_html + # Test 4D List image save_html + self.register_probe(data) + self.register_probe({ + "test_np_list": ProbeData([np.zeros([40, 40, 3]).astype(dtype=np.int8) for _ in range(5)]), + "test_2d_img": ProbeData(np.zeros([40, 40]).astype(dtype=np.uint8), is_image=True), + "test_3d_n2d_img": ProbeData(np.zeros([10, 40, 40]).astype(dtype=np.uint8), is_image=True), + "test_3d_img": ProbeData(np.zeros([40, 40, 3]).astype(dtype=np.uint8), is_image=True), + "test_list_3d_img": ProbeData([np.zeros([40, 40, 3]).astype(dtype=np.uint8) for _ in range(5)], + is_image=True), + "test_4d_img": ProbeData(np.zeros([10, 40, 40, 3]).astype(dtype=np.uint8), is_image=True), + "test_4d_img_html": ProbeData(np.zeros([10, 40, 40, 3]).astype(dtype=np.uint8), is_image=True, + build_html=True, build_label="4d_data"), + "test_4d_img_list_html": ProbeData([np.zeros([10, 40, 40, 3]).astype(dtype=np.uint8) for _ in range(5)], + is_image=True, + build_html=True, build_label=[f"4d_data_{i}" for i in range(5)]), + }) + # Test nested types + self.test_model(data) + + +cfg = Config(cfg_file="./config/general_config.yaml") +if cfg.have("FILE_SYSTEMS"): + for file_sys in cfg.FILE_SYSTEMS: + fs_prefix = FS.init_fs_client(file_sys) +else: + fs_prefix = FS.init_fs_client(cfg) + +_model = TestModel2(cfg) + +data = { + "data_key": [1, 1], + "data_folder": {"mj": 1., "mj_square": 2.}, + "timestamp": 1, + "valid_str": "right", + "test_np": np.zeros([40, 40, 3]).astype(dtype=np.int8) +} +_model(data) +probe = _model.probe_data() +for key in probe: + print(key, probe[key].to_log(prefix=f"xxx/{key}")) +``` +
+ +Use in conjunction with Hooks as follows (where PROB_INTERVAL is the probe storage interval, i.e., the number of calls to probe_data()): + +```yaml +- + NAME: ProbeDataHook + PROB_INTERVAL: 100 +``` +#### class **ProbeData** + +Instance of probe data. + +#### function **ProbeData.__init__** + +(data, is_image = False, build_html = False, build_label = None, view_distribute = False) + +**Parameters** +- **data** —— The probe data passed in, currently supports str, Number, list, dict, tensor. +- **is_image** —— Whether to store as an image. +- **build_html** —— Whether to store as html. +- **build_label** —— The label for saving html. +- **view_distribute** —— To count the frequency of some values. diff --git a/docs/zh_cn/Makefile b/docs/zh_cn/Makefile new file mode 100644 index 0000000..ed88099 --- /dev/null +++ b/docs/zh_cn/Makefile @@ -0,0 +1,20 @@ +# Minimal makefile for Sphinx documentation +# + +# You can set these variables from the command line, and also +# from the environment for the first two. +SPHINXOPTS ?= +SPHINXBUILD ?= sphinx-build +SOURCEDIR = . +BUILDDIR = build + +# Put it first so that "make" without argument is like "make help". +help: + @$(SPHINXBUILD) -M help "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O) + +.PHONY: help Makefile + +# Catch-all target: route all unknown targets to Sphinx using the new +# "make mode" option. $(O) is meant as a shortcut for $(SPHINXOPTS). +%: Makefile + @$(SPHINXBUILD) -M $@ "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O) diff --git a/docs/zh_cn/conf.py b/docs/zh_cn/conf.py new file mode 100644 index 0000000..a713930 --- /dev/null +++ b/docs/zh_cn/conf.py @@ -0,0 +1,83 @@ +# -*- coding: utf-8 -*- +# Configuration file for the Sphinx documentation builder. +# +# This file only contains a selection of the most common options. For a full +# list see the documentation: +# https://www.sphinx-doc.org/en/master/usage/configuration.html + +# -- Path setup -------------------------------------------------------------- + +# If extensions (or modules to document with autodoc) are in another directory, +# add these directories to sys.path here. If the directory is relative to the +# documentation root, use os.path.abspath to make it absolute, like shown here. +# +# import os +# import sys +# sys.path.insert(0, os.path.abspath('.')) + +# -- Project information ----------------------------------------------------- + +project = 'scepter' +copyright = '2023, scepter' +author = 'scepter' + + +def get_version(): + version_path = '../../scepter/version.py' + with open(version_path) as f: + exec(compile(f.read(), version_path, 'exec')) + return locals()['__version__'] + + +release = get_version() + +# -- General configuration --------------------------------------------------- + +# Add any Sphinx extension module names here, as strings. They can be +# extensions coming with Sphinx (named 'sphinx.ext.*') or your custom +# ones. +extensions = [ + 'sphinx.ext.autodoc', + 'sphinx.ext.napoleon', + 'sphinx.ext.viewcode', + 'sphinx_copybutton', + 'recommonmark', + 'sphinx_markdown_tables', +] + +# Add any paths that contain templates here, relative to this directory. +templates_path = ['_templates'] + +# The language for content autogenerated by Sphinx. Refer to documentation +# for a list of supported languages. +# +# This is also used if you do content translation via gettext catalogs. +# Usually you set "language" from the command line for these cases. +language = 'zh_CN' + +master_doc = 'index' + +# List of patterns, relative to source directory, that match files and +# directories to ignore when looking for source files. +# This pattern also affects html_static_path and html_extra_path. +exclude_patterns = ['build'] + +# -- Options for HTML output ------------------------------------------------- + +source_suffix = { + '.rst': 'restructuredtext', + '.md': 'markdown', +} + +# The theme to use for HTML and HTML Help pages. See the documentation for +# a list of builtin themes. +html_theme = 'press' + +# Add any paths that contain custom static files (such as style sheets) here, +# relative to this directory. They are copied after the builtin static files, +# so a file named "default.css" will overwrite the builtin "default.css". +html_static_path = ['_static'] + +# import sphinx_rtd_theme +# html_theme = "sphinx_rtd_theme" +# html_theme_path = [sphinx_rtd_theme.get_html_theme_path()] diff --git a/docs/zh_cn/index.rst b/docs/zh_cn/index.rst new file mode 100644 index 0000000..f12f8cd --- /dev/null +++ b/docs/zh_cn/index.rst @@ -0,0 +1,32 @@ + +========================================== + +.. toctree:: + :maxdepth: 1 + :caption: 开始使用 + + scepter/quick_start.md + +.. toctree:: + :maxdepth: 2 + :caption: 核心模块 + + scepter/data.md + scepter/model.md + scepter/opt.md + scepter/solvers.md + scepter/tools.md + scepter/transforms.md + +.. toctree:: + :maxdepth: 2 + :caption: 工具模块 + + scepter/utils/file_clients.md + scepter/utils/utils.md + +.. toctree:: + :maxdepth: 2 + :caption: 变更日志 + + scepter/changelog.md diff --git a/docs/zh_cn/make.bat b/docs/zh_cn/make.bat new file mode 100644 index 0000000..061f32f --- /dev/null +++ b/docs/zh_cn/make.bat @@ -0,0 +1,35 @@ +@ECHO OFF + +pushd %~dp0 + +REM Command file for Sphinx documentation + +if "%SPHINXBUILD%" == "" ( + set SPHINXBUILD=sphinx-build +) +set SOURCEDIR=source +set BUILDDIR=build + +if "%1" == "" goto help + +%SPHINXBUILD% >NUL 2>NUL +if errorlevel 9009 ( + echo. + echo.The 'sphinx-build' command was not found. Make sure you have Sphinx + echo.installed, then set the SPHINXBUILD environment variable to point + echo.to the full path of the 'sphinx-build' executable. Alternatively you + echo.may add the Sphinx directory to PATH. + echo. + echo.If you don't have Sphinx installed, grab it from + echo.https://www.sphinx-doc.org/ + exit /b 1 +) + +%SPHINXBUILD% -M %1 %SOURCEDIR% %BUILDDIR% %SPHINXOPTS% %O% +goto end + +:help +%SPHINXBUILD% -M help %SOURCEDIR% %BUILDDIR% %SPHINXOPTS% %O% + +:end +popd diff --git a/docs/zh_cn/scepter/dataset.md b/docs/zh_cn/scepter/dataset.md new file mode 100644 index 0000000..893d3ea --- /dev/null +++ b/docs/zh_cn/scepter/dataset.md @@ -0,0 +1,65 @@ +# 数据集模块 (Dataset) + +## 总览 +在继承BaseDataset基础上注册每个task各自的数据集读取模块,BaseDataset中封装有File System以及transform pipeline; + +
+ +## **scepter.modules.data.dataset.BaseDataset** +### function **\_\_init\_\_()** +#### Parameters(输入参数) +(cfg: scepter.modules.utils.config.Config, logger = None) -> None +#### config: +* MODE —— [train, test, eval]; +* TRANSFORMS —— 用于构造pipeline处理每个样本,详见docs.transforms; +* FILE_SYSTEM —— File IO Handler,支持不同类型文件的读写操作,详见docs.utils.file_clients; + +#### object(内部) +* self.pipeline —— transform pipeline +* self.fs_prefix —— The prefix of instantiated fs_client +* self.local_we —— local rank info when is_distributed +### function **\_\_getitem\_\_()** +#### Parameters(输入参数):index +index的类别由sampler确定,详见data/registry.py; + +一般默认的torch dataloader传入的index为int型作为dataset下标; + +自定义的sampler则可以传入自定义item,sampler定义参照scepter.modules.utils.sampler; + +* 具体读取单条数据由_get()方法传入index参数实现; +* 输出经由pipeline转换后的数据(如果有pipeline); + +### function **worker_init_fn()** +#### Parameters(输入参数): +(worker_id, num_workers = 1) + +worker_id为分布式训练中工作节点id; + +dataloader一次性创建num_workers个工作进程; +* 用于初始化文件读取系统和设置多卡worker对应参数; + +### function **\_get()** +#### Parameters(输入参数):index(由__getitem__()传入) +抽象方法,需要由自定义dataset继承具体实现,用于根据index读取batch中每条数据; +
+ +## 基础用法 +子类注册: + +```python +from scepter.modules.data.dataset import BaseDataset +from scepter.modules.data.dataset import DATASETS + + +@DATASETS.register_class() +class XxxxxDataset(BaseDataset): + def __init__(self, cfg, logger=None): + super(XxxxxDataset, self).__init__(cfg, logger=logger) +``` +实际启动dataset用法: + +```python +if self.cfg.have("TRAIN_DATA"): + train_data = DATASETS.build(self.cfg.TRAIN_DATA, logger=self.logger) +``` +详见scepter/modules/solver/diffusion_solver.py,具体参数定义在配置yaml中TRAIN_DATA/EVAL_DATA/TEST_DATA模块下; diff --git a/docs/zh_cn/scepter/model.md b/docs/zh_cn/scepter/model.md new file mode 100644 index 0000000..213453f --- /dev/null +++ b/docs/zh_cn/scepter/model.md @@ -0,0 +1,184 @@ +# 模型模块 (Model) +## Overview +模型模块分为backbone、neck、head、loss、metric、network、tokenizer、tuner; +* backbone/neck:一般为提取feature主要模块(necks非必有); +* head:根据不同任务类型,输入backbone提取的feature,输出下游任务所需logit; +* loss:用于计算不同类型loss; +* metric:用于计算各类评测指标; +* tokenizer:用于分词; +* tuner:用于创建微调模块; +* network:train和test模块,对数据集输入的batch整合上述模块进行最终loss和指标计算; +
+ +## **backbone/neck/head/loss/tuner** +### Basic Usage +子类注册: + +```python +from scepter.model.registry import BACKBONES +from scepter.model.base_model import BaseModel + + +@BACKBONES.register_class("ResNet") +class ResNet(BaseModel): + def __init__(self, cfg, logger=None): + super(ResNet, self).__init__(cfg, logger=logger) +``` + +```python +from scepter.model.registry import NECKS +from scepter.model.base_model import BaseModel + + +@NECKS.register_class() +class GlobalAveragePooling(BaseModel): + def __init__(self, cfg, logger=None): + super(GlobalAveragePooling, self).__init__(cfg, logger=logger) +``` + +```python +from scepter.model.registry import HEADS +from scepter.model.base_model import BaseModel + + +@HEADS.register_class() +class ClassifierHead(BaseModel): + def __init__(self, cfg, logger=None): + super(ClassifierHead, self).__init__(cfg, logger=logger) +``` + +```python +from scepter.model.registry import LOSSES +import torch.nn as nn + + +@LOSSES.register_class() +class CrossEntropy(nn.Module): + def __init__(self, cfg, logger=None): + super(CrossEntropy, self).__init__(cfg, logger=logger) +``` +实际调用: + +```python +from scepter.model.registry import BACKBONES, NECKS, HEADS, LOSSES, TUNERS + +backbone = BACKBONES.build(cfg.BACKBONE, logger=logger) +neck = NECKS.build(cfg.NECK, logger=logger) +head = HEADS.build(cfg.HEAD, logger=logger) +loss = LOSSES.build(cfg.LOSS, logger=logger) +tuner = TUNERS.build(cfg.TUNER, logger=logger) +``` + +### function **\_\_init\_\_()** +#### Parameters(输入参数) +(cfg: scepter.modules.utils.config.Config, logger = None) -> None + +主要用于初始化模型各layer; + +### function **forward()** +根据需要具体实现; +
+ + +## **metric** +### Basic Usage +子类注册: + +```python +from scepter.model.metrics.registry import METRICS +from scepter.model.metrics.base_metric import BaseMetric + + +@METRICS.register_class("AccuracyMetric") +class AccuracyMetric(BaseMetric): + def __init__(self, cfg, logger=None): + super(CrossEntropy, self).__init__(cfg, logger=logger) +``` +实际用法: + +```python +from scepter.model.metrics.registry import METRICS + +metric = METRICS.build(cfgs, logger) +``` + +### function **\_\_init\_\_()** +#### Parameters(输入参数) +(cfg: scepter.modules.utils.config.Config, logger = None) -> None + +初始化计算metric所需超参,例如topk等系数; + +### function **\_\_call\_\_()** +@torch.no_grad() + +通常输入logit和label以及其他所需要的变量,输出计算指标; +
+ +## **tokenizer** +### Basic Usage +子类注册: + +```python +from scepter.model.registry import TOKENIZERS +from scepter.model.tokenizers import BaseTokenizer + + +@TOKENIZERS.register_class() +class BaseBertTokenizer(BaseTokenizer): + def __init__(self, cfg, logger=None): + super(BaseBertTokenizer, self).__init__(cfg, logger=logger) +``` +实际用法: + +```python +from scepter.model.registry import TOKENIZERS + +tokenizer = TOKENIZERS.build(cfgs, logger) +``` +### function **\_\_init\_\_()** +#### Parameters(输入参数) +(cfg: scepter.modules.utils.config.Config, logger = None) -> None +用于初始化加载tokenizer对象,例如BertTokenizer; + +### function **tokenize()** +输入需要分词的text list,输出分词后转换的token id sequence以及其他所需的attention mask/tpye id list/position id list等; +
+ +## **network** +### Basic Usage +子类注册: + +```python +from scepter.model.registry import MODELS +from scepter.model.networks.train_module import TrainModule + + +@MODELS.register_class() +class Classifier(TrainModule): + def __init__(self, cfg, logger=None): + super(Classifier, self).__init__(cfg, logger=logger) +``` +实际用法: + +```python +from scepter.model.registry import MODELS + +model = MODELS.build(self.cfg.MODEL, logger=self.logger) +``` + +### function **\_\_init\_\_()** +#### Parameters(输入参数) +(cfg: scepter.modules.utils.config.Config, logger = None) -> None + +结合上述build方法,初始化训练所需的backbone、neck、head、loss、metric、tokenizer模块; + +### function **forward_train()** +输入训练batch的数据,经backbone、neck、head、loss计算相关loss; + +### function **forward_test()** +输入测试batch的数据,经backbone、neck、head、metrics计算相关指标; + +### function **forward** +实际调用接口,用于分发任务至forward_train()/forward_test(); + +其他训练/测试所需函数可在network下自定义; diff --git a/docs/zh_cn/scepter/opt.md b/docs/zh_cn/scepter/opt.md new file mode 100644 index 0000000..c245b85 --- /dev/null +++ b/docs/zh_cn/scepter/opt.md @@ -0,0 +1,83 @@ +# 优化器 (Optimizer) +## 总览 +1. lr_schedulers +2. optimizers +
+ +## lr_schedulers +### 基础用法 +子lr_schedulers继承时用法: + +```python +from scepter.opt.lr_schedulers import LR_SCHEDULERS +from scepter.opt.lr_schedulers.base_scheduler import BaseScheduler + + +@LR_SCHEDULERS.register_class() +class XxxLR(BaseScheduler): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) +``` +实际启动lr_scheduler用法(optimizer必要),参考task/stable_diffusion/impls/solvers/diffusion_solver.py: +```python +if self.cfg.have("LR_SCHEDULER") and not self.optimizer is None: + self.lr_scheduler = LR_SCHEDULERS.build(self.cfg.LR_SCHEDULER, logger=self.logger, + optimizer=self.optimizer) +``` + +## **scepter.modules.opt.lr_schedulers.base_scheduler.BaseScheduler** +lr_schedulers的基类,支持注册操作,可根据需要自定义; + +### function **\_\_init\_\_()** +#### Parameters(输入参数) +(cfg: scepter.modules.utils.config.Config, logger = None) -> None +#### config(常用参数,实际需要根据不同schedulers设置,以StepLR为例): +* STEP_SIZE +* GAMMA +* LAST_EPOCH + +### function **\_\_call\_\_()** +#### Parameters(输入参数) +(optimizer: scepter.modules.opt.optimizers.OPTIMIZERS) -> None + +具体对传入的optimizer对象进行schedule设置; +
+ +## optimizers +### 基础用法 +子optimizers继承时用法: + +```python +from scepter.opt.optimizers.base_optimizer import BaseOptimize +from scepter.opt.optimizers.registry import OPTIMIZERS + + +@OPTIMIZERS.register_class() +class Xxx(BaseOptimize): + def __init__(self, cfg, logger=None): + super(Xxx, self).__init__(cfg, logger=logger) +``` +实际启动optimizers用法,参考task/stable_diffusion/impls/solvers/diffusion_solver.py,需要传入train_parameters: +```python +if self.cfg.have("OPTIMIZER"): + self.optimizer = OPTIMIZERS.build(self.cfg.OPTIMIZER, logger=self.logger, + parameters=self.train_parameters()) +``` + +## **scepter.modules.opt.optimizers.base_optimizer.BaseOptimize** +optimizers的基类,支持注册操作,可根据需要自定义; + +### function **\_\_init\_\_()** +#### Parameters(输入参数) +(cfg: scepter.modules.utils.config.Config, logger = None) -> None +#### config(常用参数,实际需要根据不同optimizer设置,以SGD为例): +* LEARNING_RATE +* MOMENTUM +* DAMPENING +* WEIGHT_DECAY +* NESTEROV + +### function **\_\_call\_\_()** +#### Parameters(输入参数) +(parameters:dict()) -> None +输入需要梯度更新的train parameters,格式为dict; diff --git a/docs/zh_cn/scepter/quick_start.md b/docs/zh_cn/scepter/quick_start.md new file mode 100644 index 0000000..9af7f39 --- /dev/null +++ b/docs/zh_cn/scepter/quick_start.md @@ -0,0 +1,251 @@ +# 快速开始 (Quick Start) + +本章节以stable diffusion v1.5为例,介绍如何从零开始构建一个网络结构,以及基于该结构训练一个模型和测试该模型。 + +# 1. 使用network类定义模型结构 + +network类包含了定义、训练和测试模型的方法,我们首先初始化一个network类,然后在里面定义所需的autoencoder, unet, embedder, 以及loss子模块。 + +```python +@MODELS.register_class() +class LatentDiffusion(TrainModule): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.model_config = cfg.DIFFUSION_MODEL + self.first_stage_config = cfg.FIRST_STAGE_MODEL + self.cond_stage_config = cfg.COND_STAGE_MODEL + self.loss_config = cfg.get('LOSS', None) + + self.model = BACKBONES.build(self.model_config, logger=self.logger) + self.first_stage_model = MODELS.build(self.first_stage_config, + logger=self.logger) + self.cond_stage_model = EMBEDDERS.build(self.cond_stage_config, + logger=self.logger) + if self.loss_config: + self.loss = LOSSES.build(self.loss_config, logger=self.logger) + + # 其他变量和模块定义 +``` + +# 2. 实现自定义network类的训练和测试方法 + +每个network类依赖forward_train和forward_test方法定义自己的训练和测试流程,在sd1.5中,forward_train对采样时刻t进行噪声预测以及进行loss计算 + +```python + def forward_train(self, image=None, noise=None, prompt=None, **kwargs): + x_start = self.encode_first_stage(image, **kwargs) + t = torch.randint(0, + self.num_timesteps, (x_start.shape[0], ), + device=x_start.device).long() + context = {} + if prompt and self.cond_stage_model: + zeros = (torch.rand(len(prompt)) < self.p_zero).numpy().tolist() + prompt = [ + self.train_n_prompt if zeros[idx] else p + for idx, p in enumerate(prompt) + ] + + with torch.autocast(device_type='cuda', enabled=False): + context = self.encode_condition( + self.tokenizer(prompt).to(we.device_id)) + + loss = self.diffusion.loss(x0=x_start, + t=t, + model=self.model, + model_kwargs={'cond': context}, + noise=noise) + loss = loss.mean() + ret = {'loss': loss, 'probe_data': {'prompt': prompt}} + return ret +``` + +forward_test函数用于推理阶段执行完整的图像去噪过程 + +```python +@torch.no_grad() +@torch.autocast('cuda', dtype=torch.float16) +def forward_test(self, + prompt=None, + n_prompt=None, + sampler='ddim', + sample_steps=50, + seed=2023, + guide_scale=7.5, + guide_rescale=0.5, + discretization='trailing', + run_train_n=True, + **kwargs): + g = torch.Generator(device=we.device_id) + seed = seed if seed >= 0 else random.randint(0, 2**32 - 1) + g.manual_seed(seed) + num_samples = len(prompt) + + n_prompt = default(n_prompt, [self.default_n_prompt] * len(prompt)) + assert isinstance(prompt, list) and \ + isinstance(n_prompt, list) and \ + len(prompt) == len(n_prompt) + + context = self.encode_condition(self.tokenizer(prompt).to( + we.device_id), method='encode_text') + null_context = self.encode_condition(self.tokenizer(n_prompt).to( + we.device_id), method='encode_text') + + width, height = 512, 512 + noise = self.noise_sample(num_samples, width // self.size_factor, + height // self.size_factor, g) + # UNet use input n_prompt + samples = self.diffusion.sample(solver=sampler, + noise=noise, + model=self.model, + model_kwargs=[{ + 'cond': context + }, { + 'cond': null_context + }], + steps=sample_steps, + guide_scale=guide_scale, + guide_rescale=guide_rescale, + discretization=discretization, + show_progress=True, + seed=seed, + condition_fn=None, + clamp=None, + percentile=None, + t_max=None, + t_min=None, + discard_penultimate_step=None, + return_intermediate=None, + **kwargs) + x_samples = self.decode_first_stage(samples).float() + x_samples = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0) + + outputs = list() + for p, np, img in zip(prompt, n_prompt, x_samples): + one_tup = {'prompt': p, 'n_prompt': np, 'image': img} + outputs.append(one_tup) + + return outputs +``` + +# 3. 子模块注册 + +在实现完network类之后,需要确保network类中用到的所有子模块都已完成注册。以sd1.5中的embedder为例,为了能在network的初始化方法中实例化该embedder,我们需要先实现该embedder类,并注册到scepter中 + +```python +@EMBEDDERS.register_class() +class FrozenCLIPEmbedder(BaseEmbedder): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + + with FS.get_dir_to_local_dir(cfg.PRETRAINED_MODEL, + wait_finish=True) as local_path: + self.tokenizer = CLIPTokenizer.from_pretrained(local_path) + self.transformer = CLIPTextModel.from_pretrained(local_path) + + self.use_grad = cfg.get('USE_GRAD', False) + self.freeze_flag = cfg.get('FREEZE', True) + if self.freeze_flag: + self.freeze() + + self.max_length = cfg.get('MAX_LENGTH', 77) + self.layer = cfg.get('LAYER', 'last') + self.layer_idx = cfg.get('LAYER_IDX', None) + self.use_final_layer_norm = cfg.get('USE_FINAL_LAYER_NORM', False) + assert self.layer in self.LAYERS + if self.layer == 'hidden': + assert self.layer_idx is not None + assert 0 <= abs(self.layer_idx) <= 12 + + + def encode_text(self, + tokens, + tokenizer=None, + append_sentence_embedding=False): + # 定义一些需要的方法 + pass + + @staticmethod + def get_config_template(): + return dict_to_yaml('MODELS', + __class__.__name__, + FrozenCLIPEmbedder.para_dict, + set_name=True) +``` + +# 4. Solver注册 + +solver类中封装了训练和测试一个network所需要的完整流程。以sd1.5为例,注册一个训练sd1.5模型的solver需要: +1. 创建一个或多个数据加载器(data_loader), 对应solver中的construct_data方法; +2. 实例化一个模型,对应solver中的construct_model方法 +3. (可选)定义度量指标,对应solver中的construct_metrics方法 +4. (可选)定义训练和测试用到的一些钩子(HOOKS),比如模型保存,预训练参数加载,日志打印,等等。 + +```python +@SOLVERS.register_class() +class LatentDiffusionSolver(BaseSolver): + def set_up(self): + self.construct_data() + self.construct_model() + self.construct_metrics() + self.model_to_device() + self.init_opti() + + def load_checkpoint(self, checkpoint): + # 这里定义加载模型的指令 + + def save_checkpoint(self): + # 这里定义保存模型的指令 + + def solve(self): + # 入口函数,根据数据类型选择执行训练或测试 + self.before_solve() + if 'train' in self._mode_set: + self.run_train() + if 'test' in self._mode_set: + self.run_test() + self.after_solve() + + def run_train(self): + # 模型训练 + + def run_eval(self): + # 模型验证 + + def run_test(self): + # 模型测试 +``` + +# 5. 训练/测试超参数定义 + +在注册完所需的各个组件(包括但不限于BACKBONE, NETWORK, EMBEDDER, SOLVER, METRIC)后,需要对其中用到的一些超参数进行设置,scepter使用yaml文件定义各模块超参数,具体参考scepter/modules/examples/sd15/sd15_512_full.yaml + +# 6. 模型训练 + +通过指定--cfg参数加载所需的yaml文件,完成训练或批量测试的操作 +```shell +多机多卡训练 +# 基于spawn方式,为默认模式 +export CUDA_VISIBLE_DEVICES=0,1,2,3 +export WORLD_SIZE=1 +python -W ignore scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml +# 基于原生的pytorch引擎 torchrun模式 +torchrun --nproc_per_node 4 scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml --launcher torchrun +# 基于pytorch_lightning引擎, ENV.USE_PL需要设置为true +python -W ignore scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml +# 单卡训练 +# 基于spawn方式,为默认模式 +export CUDA_VISIBLE_DEVICES=0 +export WORLD_SIZE=1 +python -W ignore scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml +# 基于原生的pytorch引擎 +python -W ignore scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml --launcher torchrun +# 基于pytorch_lightning引擎, ENV.USE_PL需要设置为true +python scepter/run_train.py --cfg scepter/examples/sd15/sd15_512_full.yaml +``` + +# 7. 模型推理 + +单次推理可以通过自定义run_inference.py文件来实现,参照sd1.5的推理方法 +```shell +python -W ignore scepter/run_inference.py --prompt "a woman" --n_prompt "" --num_samples 4 --pretrained_model "path/to/your/pretrained/model" +``` diff --git a/docs/zh_cn/scepter/sampler.md b/docs/zh_cn/scepter/sampler.md new file mode 100644 index 0000000..e7104f2 --- /dev/null +++ b/docs/zh_cn/scepter/sampler.md @@ -0,0 +1,165 @@ + +# 采样器模块(Sampler) + +## Overview +采样器定义了选择训练、验证和测试所需数据的采样方式。 +采样器比较具有通用性,大多数情况下不会进行定制开发,这里提供了几类常用的sampler采样器。 +其中部分采样器例如MixtureOfSamplers封装有SubSampler。 + +当然,采样器也支持定制开发并注册,开发者只需要在继承BaseSampler基础上对实现的sampler类进行注册。 + +## Basic Usage + +预定义samplers: +- TorchDefault +- LoopSampler, +- MixtureOfSamplers, +- MultiFoldDistributedSampler, +- EvalDistributedSampler, +- MultiLevelBatchSampler, +- MultiLevelBatchSamplerMultiSource +```python +from scepter.modules.data.sampler import ( + LoopSampler, + MixtureOfSamplers, + MultiFoldDistributedSampler, + EvalDistributedSampler, + MultiLevelBatchSampler, + MultiLevelBatchSamplerMultiSource) +``` +自定义sampler,以LoopSampler为例: +```python +@SAMPLERS.register_class() +class LoopSampler(BaseSampler): + para_dict = {} + + def __init__(self, cfg, logger): + super().__init__(cfg, logger) + rank = we.rank + self.rng = np.random.default_rng(self.seed + rank) + + def __iter__(self): + while True: + yield self.rng.choice(sys.maxsize) + + def __len__(self): + return sys.maxsize + + @staticmethod + def get_config_template(): + return dict_to_yaml('SAMPLERS', + __class__.__name__, + LoopSampler.para_dict, + set_name=True) +``` +实现__iter__和__len__方法实现采样功能,在__init__的入参cfg中可以读取相关配置参数。 +
+ +## LoopSampler +简单数据的无限循环sampler。 + +#### function **LoopSampler.__init__** +(cfg: scepter.modules.utils.config.Config, logger = None) -> None + +#### function **LoopSampler.__iter__** +迭代器,每迭代一次得到一个样本的index + + +## MixtureOfSamplers + +用于大规模数据的多级索引的sampler + +#### function **MixtureOfSamplers.__init__** + +(samplers: list(sampler), probabilities: list(float), rank: int =0, seed: int = 8888) + +**Parameters** + +- **samplers** —— 采样器列表,用于混合采样器。 +- **probabilities** —— 每个采样器的概率。 +- **rank** —— rank表示当前进程号。 +- **seed** —— 随机采样的seed,在data.registry中获取全局seed。 + +#### function **MultiLevelBatchSampler.__iter__** +() + +迭代器,每迭代一次得到一个样本的index + +## MultiFoldDistributedSampler + +多fold采样器,支持在一个epoch中重复多轮数据 + +#### function **MultiFoldDistributedSampler.__init__** + +( dataset: torch.data.dataset, num_folds=1, num_replicas=None, rank=None, shuffle=True) + +**Parameters** + +- **dataset** —— torch.data.dataset类实例 +- **num_folds** —— int,表示数据重复的轮数。 +- **num_replicas** —— 表示数据分割的片数,一般和world-size保持一致。 +- **rank** —— rank表示当前进程号。 +- **shuffle** —— 数据是否要打乱。 + +#### function **MultiFoldDistributedSampler.__iter__** +() + +迭代器,每迭代一次得到一个样本的index + +#### function **MultiFoldDistributedSampler.set_epoch** +(epoch: int) + +设置当前的epoch + +**Parameters** + +- **epoch** —— 当前的epoch。 + + +## EvalDistributedSampler + +用于测试时的采样器,当不用padding模式的时候,会发现最后一个rank的数据会少于其他rank。 + +#### function **EvalDistributedSampler.__init__** + +( dataset: torch.data.dataset, num_replicas: Optional[int] =None, rank: Optional[int] =None, padding: bool =False) + +**Parameters** + +- **dataset** —— torch.data.dataset类实例 +- **num_replicas** —— 表示数据分割的片数,一般和world-size保持一致。 +- **rank** —— rank表示当前进程号。 +- **padding** —— 数据是否需要padding,如果padding则能保证最后一个rank和其他rank数据量一致。 + +#### function **EvalDistributedSampler.__iter__** +() + +迭代器,每迭代一次得到一个样本的index + +#### function **EvalDistributedSampler.set_epoch** +(epoch: int) + +设置当前的epoch + +**Parameters** + +- **epoch** —— 当前的epoch。 + +## MultiLevelBatchSampler +用于大规模数据的多级索引的sampler + +#### function **MultiLevelBatchSampler.__init__** + +(index_file: str, batch_size: int, rank: int =0, seed: int = 8888) + +**Parameters** + +- **index_file** —— 多级数据索引的索引文件。 +- **batch_size** —— 一个batch的大小。 +- **rank** —— rank表示当前进程号。 +- **seed** —— 随机采样的seed,在data.registry中获取全局seed。 + +#### function **MultiLevelBatchSampler.__iter__** +() + +迭代器,每迭代一次得到一个样本的index diff --git a/docs/zh_cn/scepter/solvers.md b/docs/zh_cn/scepter/solvers.md new file mode 100644 index 0000000..e724d47 --- /dev/null +++ b/docs/zh_cn/scepter/solvers.md @@ -0,0 +1,489 @@ +# 训练器(Solvers) + +## 总览 +Solver是对模型训练、验证和测试过程的一个流程定义。 +在Solver中,会根据配置yaml文件的设置,对数据(data),模型(model),优化器(optimizer)和调度(scheduler)等需要的模块进行逐一初始化。 +每一个具体的任务的自定义Solver都要继承自BaseSolver。 + +在某些特殊场景下,还需要初始化一些自定义的模块。例如,在训练过程中记录保存中间结果需要用到Hooks;在训练过程中验证需要定义Metrics等等。 + +
+ +## 基础用法 + +子Solver继承时的用法: + +```python +from scepter.modules.solver.registry import SOLVERS +from scepter.modules.solver import BaseSolver + + +@SOLVERS.register_class() +class XxxSolver(BaseSolver): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) +``` + +实际启动Solver的使用方式见scepter.modules/task/cate_recognition中run_task.py和run_inference.py: + +```python +solver = SOLVERS.build(cfg.SOLVER, logger=std_logger) +``` + +
+ +## **scepter.modules.solver.BaseSolver** +Solver的基类,是通过元类ABCMeta定义的抽象基类,支持注册操作。自定义的solver均应该是该类的子类,并且进行注册。 + +BaseSolver是一个具体实现Solver的案例,展示了使用pytorch_lightning和不使用时两种Solver的写法。 +在实际使用中情况各异,因此Solver的使用也比较灵活,其中大部分的成员函数均可以按照需求在子类中,被复写或者新加好用的功能,甚至直接写新的函数代替。 + +
+ +### function **\_\_init\_\_** + +(cfg: *scepter.modules.utils.config.Config*, logger = None) -> None + +初始化Solver类,从cfg中获取并定义一些必要的参数。 + +部分在__init__中初始化的参数,详情见代码: + +**Configs** + +- **FILE_SYSTEM** —— (cfg) 文件系统配置,默认None. +- **WORK_DIR** —— (str) 工作路径定义. +- **LOG_FILE** —— (str) Log文件的位置. +- **RESUME_FROM** —— (str) 恢复训练模型保存的中间结果. +- **MAX_EPOCHS** —— (int) 训练最大epochs数. +- **ACCU_STEP** —— (int) When use ddp, the grad accumulate steps for each process,默认1. +- **NUM_FOLDS** —— (int) Num folds for training,默认1. +- **EVAL_INTERVAL** —— (int) Eval the model interval,默认1. +- **EXTRA_KEYS** —— (List) The extra keys for metrics,默认[]. +- **TRAIN_DATA** —— (cfg) 训练数据配置. +- **EVAL_DATA** —— (cfg) 验证数据配置. +- **TEST_DATA** —— (cfg) 测试数据配置. +- **TRAIN_HOOKS** —— (List) 训练HOOKS. +- **EVAL_HOOKS** —— (List) 验证HOOKS. +- **TEST_HOOKS** —— (List) 测试HOOKS. +- **MODEL** —— (cfg) 模型配置. +- **OPTIMIZER** —— (cfg) 优化器配置. +- **LR_SCHEDULER** —— (cfg) 学习率Scheduler配置. +- **METRICS** —— (List) Metrics. + +**Parameters** + +- **cfg** —— The Config used to build solver. +- **logger** —— Instantiated Logger to print or save log. + +
+ +### function **set_up_pre** + +() -> None + +配置环境、日志路径,调用construct_hook来初始化hook等等。 + +与pytorch_lightning二选一,在不使用pytorch_lightning(use_pl=Flase)时调用。需要在启动其他所有操作之前调用。 + +
+ +### function **set_up** + +() -> None + +配置数据、模型、metrics、优化器及pytorch_lightning环境(如果使用的话)。 + +
+ +### function **construct_data** + +() -> None + +实际数据的构建方法,默认会在self.set_up中被调用,包括TRAIN_DATA、EVAL_DATA、TEST_DATA。将实例化的结果写入self.datas中。 + +
+ +### function **construct_hook** + +() -> None + +实际Hook的构建方法,默认会在self.set_up_pre中被调用,包括TRAIN_HOOKS、EVAL_HOOKS、TEST_HOOKS。将实例化的结果写入self.hooks_dict中。 + +
+ +### function **construct_model** + +() -> None + +实际Hook的构建方法,默认会在self.set_up中被调用,将实例化的结果作为self.model。 + +
+ +### function **model_to_device** + +() -> None + +实际Metrics的构建方法,默认会在self.set_up中被调用,将实例化的结果写入self.metrics。 + +
+ +### function **model_to_device** + +(tg_model_ins=None) -> None or MODELS + +模型的配置方法,包括模型分片等分布式配置,默认会在self.set_up中被调用。 + +**Parameters** + +- **tg_model_ins** —— 待配置的model,如果为None,默认使用self.model。 + +**Returns** + +- **tg_model_ins** —— 配置好的model。如果tg_model_ins为None,则无返回值。 + +
+ +### function **init_opti** + +() -> None + +优化器的配置方法,默认会在self.set_up中被调用。将实例化的optimizer和lr_scheduler分别作为self.optimizer和self.lr_scheduler。 + +
+ +### function **solve** + +(epoch = None, every_epoch = False) -> None + +执行实际的训练、验证或者测试操作,执行self.solve_train、self.solve_eval、self.solve_test等。 +并且在执行前后分别调用self.before_solve、self.after_solve来进行Hook的记录。 + +**Parameters** + +- **epoch** —— 设定epoch数量。 +- **every_epoch** —— 与self.solve_train、self.solve_eval、self.solve_test的实现方式及Data的配置有关,标记是否每个epoch都需要重新调用一遍。 +以self.solve_train为例,如果实现方式是调用一次执行一个epoch,则every_epoch应为True;如果调用一次调用会执行到所有epoch都结束,则应为False。 + +
+ +### function **solve_train** + +() -> None + +调用self.run_train,执行train。 +并且在执行前后分别调用self.before_epoch、self.after_epoch来进行Hook的记录。 + +
+ +### function **solve_eval** + +() -> None + +调用self.run_eval,执行eval。 +并且在执行前后分别调用self.before_epoch、self.after_epoch来进行Hook的记录。 + +
+ +### function **solve_test** + +() -> None + +调用self.run_test,执行test。 +并且在执行前后分别调用self.before_epoch、self.after_epoch来进行Hook的记录。 + +
+ +### function **before_solve** + +() -> None + +在实际执行run_xxx之前,进行Hook的记录。 + +
+ +### function **after_solve** + +() -> None + +在执行run_xxx之后,再次进行Hook的记录。 + +
+ +### function **run_train** + +() -> None + +循环调用self.run_step_train,执行一个epoch或者所有epoch的训练过程。 +并且在循环开始及结束后分别调用self.before_all_iter、self.after_all_iter来进行Hook的记录。 +并且在每次循环调用self.run_step_train前后,分别调用self.before_iter、self.after_iter来进行Hook的记录。 + +
+ +### function **run_eval** + +() -> None + +循环调用self.run_step_eval,执行一个epoch或者所有epoch的验证过程。 +并且在循环开始及结束后分别调用self.before_all_iter、self.after_all_iter来进行Hook的记录。 +并且在每次循环调用self.run_step_eval前后,分别调用self.before_iter、self.after_iter来进行Hook的记录。 + +
+ +### function **run_test** + +() -> None + +循环调用self.run_step_test,执行一个epoch或者所有epoch的验证过程。 +并且在循环开始及结束后分别调用self.before_all_iter、self.after_all_iter来进行Hook的记录。 +并且在每次循环调用self.run_step_test前后,分别调用self.before_iter、self.after_iter来进行Hook的记录。 + +
+ +### function **run_step_train** + +(batch_data, batch_idx = 0, step = None, rank = None) -> None + +执行一个batch的模型推理。 + +
+ +### function **run_step_eval** + +(batch_data, batch_idx = 0, step = None, rank = None) -> None + +执行一个batch的模型推理。 + +
+ +### function **run_step_test** + +(batch_data, batch_idx = 0, step = None, rank = None) -> None + +执行一个batch的模型推理。 + +
+ +### function **register_flops** + +(Dict data, keys = []) -> None + +使用fvcore进行flops计算,结果保存在self._model_flops中。 + +**Parameters** + +- **data** —— 输入模型的data,格式为key-value结构,value为Tensor或者List,默认只取batch维度的第一个元素来保证batchsize=1。 +- **keys** —— 指定使用keys中的元素所对应的value作为模型输入。如果keys为空,则去data中所有的value。 + +
+ +### function **before_epoch** + +() -> None + +在每个epoch开始前执行。run_xxx前执行。 + +
+ +### function **before_all_iter** + +() -> None + +在每个epoch开始前执行。循环调用run_step_xxx的循环前执行。 + +
+ +### function **before_iter** + +() -> None + +在每个step开始前执行。run_step_xxx前执行。 + +
+ +### function **after_epoch** + +() -> None + +在每个epoch开始后执行。run_xxx后执行。 + +
+ +### function **after_all_iter** + +() -> None + +在每个epoch开始后执行。循环调用run_step_xxx的循环后执行。 + +
+ +### function **after_iter** + +() -> None + +在每个step开始后执行。run_step_xxx后执行。 + +
+ +### function **collect_log_vars** + +() -> OrderedDict + +获取需要在log中保存的变量。 + +**Returns** + +- **ret** —— 返回需要的变量。 + +
+ + +# Hooks + +## 总览 + +在Solver执行过程中,需要打印日志、记录Tensorboard、梯度计算和更新、保存中间模型参数、保存测试结果等,这些都需要Hook去执行。 + +
+ +## 基础用法 + +新建Hook的方法: + +```python +from scepter.modules.solver.hooks.hook import Hook +from scepter.modules.solver.hooks.registry import HOOKS + + +@HOOKS.register_class() +class XxxHook(Hook): + def __init__(self, cfg, logger=None): + super(XxxHook, self).__init__(cfg, logger=logger) +``` + +
+ + +## **scepter.modules.solvers.hooks.Hook** +定义了一个标准的Hook类的基类,定义了详细的成员函数列表。继承自该类的子类均需要在这些成员函数中挑选部分进行实现。 +成员函数包括: + +### function **__init__** +(cfg: *scepter.modules.utils.config.Config*, logger = None) -> None +初始化logger。 + +### function **before_solve** +(solver) -> None + +开始solve之前执行。 + +### function **after_solve** +(solver) -> None + +结束solve之后执行。 + +### function **before_epoch** +(solver) -> None + +每个epoch之前执行。 + +### function **after_epoch** +(solver) -> None + +每个epoch之后执行。 + +### function **before_all_iter** +(solver) -> None + +迭代开始之前执行。 + +### function **after_all_iter** +(solver) -> None + +迭代结束之后执行。 + +### function **before_iter** +(solver) -> None + +每个step之前执行。 + +### function **after_iter** +(solver) -> None + +每个step之后执行。 + +## **scepter.modules.solver.hooks.CheckpointHook** +在solve开始之前,加载checkpoint。加载模型的路径来自Solver的RESUME_FROM参数。依赖于Solver中实现的load_checkpoint成员函数。 + +每个epoch结束之后,保存checkpoint。 + +**Configs** + +- **PRIORITY** —— (int)默认为_DEFAULT_CHECKPOINT_PRIORITY=300 +- **INTERVAL** —— (int)保存checkpoint的epoch间隔,默认为1 +- **SAVE_NAME_PREFIX** —— (str)保存checkpoint的前缀名,默认为'ldm_step' +- **SAVE_LAST** —— (bool)是否保存最后iter的checkpoint,默认为False,作用于after_iter中。 +- **SAVE_BEST** —— (bool)是否保存最好的checkpoint,默认为True,作用于after_epoch中。需要同时设置SAVE_BEST_BY,否则退回默认False。 +- **SAVE_BEST_BY** —— (str)判断最好的指标,默认越大越好 + +## **scepter.modules.solver.hooks.BackwardHook** +在每步迭代之后执行的内容。包括loss的反传,optimizer的步数配置等等。 + +**Configs** + +- **PRIORITY** —— (int)_DEFAULT_BACKWARD_PRIORITY=0 +- **GRADIENT_CLIP** —— (int)torch.nn.utils.clip_grad_norm_中max_norm参数,默认为-1。小于等于0时不生效。 +- **ACCUMULATE_STEP** —— (int)用于设置梯度累计的步数 +- **EMPTY_CACHE_STEP** —— (int)torch.cuda.empty_cache每多少步清除一下memory,默认为-1。小于等于0时不生效。 +- +## **scepter.modules.solver.hooks.LogHook** +日志Hook。 + +**Configs** + +- **PRIORITY** —— (int)__DEFAULT_LOG_PRIORITY=100 +- **SHOW_GPU_MEM** —— (bool)判断是否打印内存使用情况 +- **LOG_INTERVAL** —— (int)打印log的步数间隔,默认为10。小于等于0时不生效。 + +## **scepter.modules.solver.hooks.TensorboardLogHook** +Tensorboard日志Hook。 + +**Configs** + +- **PRIORITY** —— (int)__DEFAULT_LOG_PRIORITY=100 +- **LOG_DIR** —— (str)存储Tensorboard log的路径。 +- **LOG_INTERVAL** —— (int)打印log的步数间隔,默认为1000。 + +## **scepter.modules.solver.hooks.LrHook** +学习率变化Hook。 + +**Configs** + +- **PRIORITY** —— (int)_DEFAULT_LR_PRIORITY=200 +- **WARMUP_FUNC** —— (str)warmup的方式:仅支持"linear",默认为"linear"。 +- **WARMUP_EPOCHS** —— (int)warmup epoch数,默认为1。 +- **WARMUP_START_LR** —— (float)warmup初始学习率,默认为0.0001。 +- **SET_BY_EPOCH** —— (bool)是否每个epoch设置一次学习率,默认为True。False则每个step设置一次。 + +## **scepter.modules.solver.hooks.DistSamplerHook** +每个epoch开始之前,采样一个epoch的数据。 + +**Configs** + +- **PRIORITY** —— (int)__DEFAULT_SAMPLER_PRIORITY=400 + +## **scepter.modules.solver.hooks.ProbeDataHook** +用于打印probe存储的train/eval中间(可视化)结果。 + +**Configs** + +- **PRIORITY** —— (int)_DEFAULT_LR_PRIORITY=200 +- **PROB_INTERVAL** —— (int)打印probe的步数间隔,默认为1000。 + +## **scepter.modules.solver.hooks.SafetensorsHook** +用于存储.safetensors格式的模型文件。 + +**Configs** + +- **PRIORITY** —— (int)_DEFAULT_LR_PRIORITY=200 +- **INTERVAL** —— (int)存储的步数间隔,默认为1000。 +- **SAVE_NAME_PREFIX** —— (str)保存文件的前缀名。 diff --git a/docs/zh_cn/scepter/tools.md b/docs/zh_cn/scepter/tools.md new file mode 100644 index 0000000..333eb1e --- /dev/null +++ b/docs/zh_cn/scepter/tools.md @@ -0,0 +1,54 @@ +# 工具组件(Tools) + +支持对框架注册的组件进行查询、获取参数模版。 + +## 总览 +1、查询模块类型:scepter.module_list +2、按模块查询对象:scepter.objects_by_module +3、按对象查询参数配置:scepter.configures_by_objects + +
+ +## 基础用法 + +```python +from scepter import module_list, objects_by_module, configures_by_objects + +# 查询模块类型 +module_list() +# 按照模块查询对象列表 +objects_by_module("BACKBONES") +# 按照对象查询参数配置 +configures_by_objects("BACKBONES", "ResNet3D_TAda") +``` +
+ +### function **module_list** +() + +**Returns** + +- **list** —— 模块名列表。 + +### function **objects_by_module** +(module_name: str) + +**Parameters** + +- **module_name** —— 模块名 + +**Returns** + +- **list** —— 模块名列表。 + +### function **get_module_object_config** +(module_name: str, object_name: str) + +**Parameters** + +- **module_name** —— 模块名 +- **object_name** —— 对象名 + +**Returns** + +- **str** —— 参数模版。 diff --git a/docs/zh_cn/scepter/transforms.md b/docs/zh_cn/scepter/transforms.md new file mode 100644 index 0000000..f6c34ce --- /dev/null +++ b/docs/zh_cn/scepter/transforms.md @@ -0,0 +1,419 @@ +# 数据转换(Transforms) + +数据预处理模块 + +## 总览 + +支持多种数据预处理方式: + +1. scepter.transforms.image +2. scepter.transforms.io +3. scepter.transforms.io_video +4. scepter.transforms.tensor +5. scepter.transforms.augmention +6. scepter.transforms.video +7. scepter.transforms.transform_xl +8. scepter.transforms.identity +9. scepter.transforms.compose + +
+ +## 基础用法 + +```python +# 以 scepter.modules.transform.image.RandomResizedCrop 为例 +from scepter.modules.transform.image import RandomResizedCrop +from scepter.modules.utils.config import Config +import PIL + +cfg = Config(load=False, + cfg_dict={"SIZE": 224, "RATIO": [3. / 4., 4. / 3.], "SCALE": [0.08, 1.0], "INTERPOLATION": "bilinear"}) +transform = RandomResizedCrop(cfg) +input_img = {"img": PIL.Image} +output_img = transform(input_img) +``` + +
+ +## **scepter.modules.transform.image** + +一些用于图像的预处理方法. + +
+ +### scepter.modules.transform.image.ImageTransform + +初始化ImageTransform类, 从cfg中获取并定义一些必要的参数. + +**Parameters** + +- **INPUT_KEY** —— (str) input key or key list +- **OUTPUT_KEY** —— (str) output key or key list +- **BACKEND** —— (str) backend, choose from pillow, cv2, torchvision + +
+ +### scepter.modules.transform.image.RandomResizedCrop + +对图像进行随机crop到指定大小. + +**Parameters** + +- **SIZE** —— (int) crop size +- **RATIO** —— (list) ratio +- **SCALE** —— (list) scale +- **INTERPOLATION** —— (str) interpolation + +
+ +### scepter.modules.transform.image.RandomHorizontalFlip + +以给定的概率随机水平翻转给定的图像. + +**Parameters** + +- **P** —— (float) probability + +
+ +### scepter.modules.transform.image.Normalize + +利用平均值和标准差对图像进行归一化. + +**Parameters** + +- **MEAN** —— (list) mean +- **STD** —— (list) std + +
+ +### scepter.modules.transform.image.ImageToTensor + +把PIL.Image / numpy.ndarray / unit8 转成float32 tensor. + +**Parameters** + +
+ +### scepter.modules.transform.image.Resize + +把给定的图像按照给定的尺寸进行resize. + +**Parameters** + +- **INTERPOLATION** —— (str) interpolation +- **SIZE** —— (int) resized size + +
+ +### scepter.modules.transform.image.CenterCrop + +对给定的图像从中心进行crop. + +**Parameters** + +- **SIZE** —— (int) crop size + +
+ +### scepter.modules.transform.image.FlexibleResize + +对给定的图像按照给定的尺寸进行resize. + +**Parameters** + +- **INTERPOLATION** —— (str) interpolation + +
+ +### scepter.modules.transform.image.FlexibleCenterCrop + +对给定的图像按照给定的尺寸进行center crop. + +**Parameters** + +
+ +## **scepter.modules.transform.io** + +一些用于图像的本地磁盘读取方法. + +
+ +### scepter.modules.transform.io.LoadPILImageFromFile + +将本地图片文件读取成PIL.Image的形式. + +**Parameters** + +- **RGB_ORDER** —— (str) "RGB" or "BGR" + +
+ +### scepter.modules.transform.io.LoadCvImageFromFile + +将本地图片文件读取成cv2的形式. + +**Parameters** + +- **RGB_ORDER** —— (str) "RGB" or "BGR" + +
+ +### scepter.modules.transform.io.LoadImageFromFile + +将本地图片文件读取成指定格式. + +**Parameters** + +- **RGB_ORDER** —— (str) "RGB" or "BGR" +- **BACKEND** —— (str) "pillow" or "cv2" or "torchvision" + +
+ +### scepter.modules.transform.io.LoadImageFromFileList + +将输入的一组图片读取成指定格式. + +**Parameters** + +- **RGB_ORDER** —— (str) "RGB" or "BGR" +- **BACKEND** —— (str) "pillow" or "cv2" or "torchvision" +- **FILE_KEYS** —— (list) The file keys for input + +
+ +## **scepter.modules.transform.io_video** + +一些用于视频的本地磁盘读取方法. + +
+ +### scepter.modules.transform.io_video.DecodeVideoToTensor + +将本地视频文件解码成tensor. + +**Parameters** + +- **NUM_FRAMES** —— (int) decode frames number +- **TARGET_FPS** —— (int) decode frames fps, default is 30. +- **SAMPLE_MODE** —— (str) interval or segment sampling, default is interval +- **SAMPLE_INTERVAL** —— (int) sample interval between output frames for interval sample mode +- **SAMPLE_MINUS_INTERVAL** —— (float) wheather minus interval for interval sample mode +- **REPEAT** —— (str) number of clips to be decoded from each video + +
+ +### scepter.modules.transform.io_video.LoadVideoFromFile + +将本地视频文件解码读取成帧序列的形式. + +**Parameters** + +- **NUM_FRAMES** —— (int) decode frames number +- **SAMPLE_TYPE** —— (str) sample type +- **CLIP_DURATION** —— (float) needed for 'interval' sampling type +- **DECODER** —— (str) video decoder name + +
+ +## **scepter.modules.transform.tensor** + +一些处理tensor的方法. + +
+ +### scepter.modules.transform.tensor.ToTensor + +将输入的其他形式的data转成tensor. + +**Parameters** + +- **KEYS** —— (list) keys of input data + +
+ +### scepter.modules.transform.tensor.Select + +选择输入data中的一些key并输出. + +**Parameters** + +- **META_KEYS** —— (list) chosen keys of input data + +
+ +### scepter.modules.transform.tensor.Rename + +将输入data的keys重新命名. + +**Parameters** + +- **IN_KEYS** —— (list) input data keys +- **OUT_KEYS** —— (list) output data keys + +
+ +## **scepter.modules.transform.augmention** + +一些图片颜色增强的方法. + +
+ +### scepter.modules.transform.augmention.ColorJitterGeneral + +随机改变图像的亮度、对比度和饱和度. + +**Parameters** + +- **BRIGHTNESS** —— (float) (float or tuple of float (min, max)): How much to jitter brightness +- **CONTRAST** —— (float) (float or tuple of float (min, max)): How much to jitter contrast +- **SATURATION** —— (float) (float or tuple of float (min, max)): How much to jitter saturation +- **HUE** —— (float) (float or tuple of float (min, max)): How much to jitter hue +- **GRAYSCALE** —— (float) probablitities for rgb-to-gray 0~1 +- **CONSISTENT** —— (bool) for video input whether the augment scale is consistent or not +- **SHUFFLE** —— (bool) shuffle the transform's order when there are multiple transforms +- **GRAY_FIRST** —— (bool) whether use grayscale or not +- **IS_SPLIT** —— (bool) whether randomly chance the channel as the gray results + +
+ +## **scepter.modules.transform.video** + +一些处理video的方法. + +
+ +### scepter.modules.transform.video.VideoTransform + +初始化VideoTransform类, 从cfg中获取并定义一些必要的参数. + +**Parameters** + +- **INPUT_KEY** —— (str) input key or key list +- **OUTPUT_KEY** —— (str) output key or key list +- **BACKEND** —— (str) backend, choose from pillow, cv2, torchvision + +
+ +### scepter.modules.transform.video.RandomResizedCropVideo + +对视频进行随机crop到指定大小. + +**Parameters** + +- **META_KEYS** —— (list) chosen keys of input data + +
+ +### scepter.modules.transform.video.CenterCropVideo + +将输入data的keys重新命名. + +**Parameters** + +- **SIZE** —— (int) crop size +- **RATIO** —— (list) ratio +- **SCALE** —— (list) scale +- **INTERPOLATION** —— (str) interpolation + +
+ +### scepter.modules.transform.video.RandomHorizontalFlipVideo + +以给定的概率随机水平翻转给定的视频. + +**Parameters** + +- **P** —— (float) probability + +
+ +### scepter.modules.transform.video.NormalizeVideo + +利用平均值和标准差对视频进行归一化. + +**Parameters** + +- **MEAN** —— (list) mean +- **STD** —— (list) std + +
+ +### scepter.modules.transform.video.VideoToTensor + +把PIL.Image / numpy.ndarray / unit8 转成float32 tensor. + +**Parameters** + +
+ +### scepter.modules.transform.video.AutoResizedCropVideo + +对给定的视频从中心进行crop. + +**Parameters** + +- **SCALE** —— (list) scale + +
+ +### scepter.modules.transform.video.ResizeVideo + +把给定的视频按照给定的尺寸进行resize. + +**Parameters** + +- **SCALE** —— (list) scale +- **INTERPOLATION** —— (str) interpolation + +
+ +## **scepter.modules.transform.transform_xl** + +sdxl中进行图像处理得到所需坐标的一些方法. + +
+ +### scepter.modules.transform.transform_xl.FlexibleCropXL + +对图像进行裁剪,并获取其原始尺寸、目标尺寸和裁剪坐标(top/left). + +**Parameters** + +- **INPUT_KEY** —— (str) input key or key list +- **OUTPUT_KEY** —— (str) output key or key list +- **BACKEND** —— (str) backend, choose from pillow, cv2, torchvision +- **SIZE** —— (int or list) crop size if 'image_size' not in meta + +
+ +## **scepter.modules.transform.identity** + +sdxl中进行图像处理得到所需坐标的一些方法. + +
+ +### scepter.modules.transform.identity.Identity + +返回图像本身. + +**Parameters** + +
+ +## **scepter.modules.transform.compose** + +组合各类transform方法. + +
+ +### scepter.modules.transform.compose.Compose + +将scepter.transforms中的各个transform对象组合为pipeline. + +**Parameters** + +- **TRANSFORMS** —— (list) transform config list + +
diff --git a/docs/zh_cn/scepter/utils/file_clients.md b/docs/zh_cn/scepter/utils/file_clients.md new file mode 100644 index 0000000..b05cbc8 --- /dev/null +++ b/docs/zh_cn/scepter/utils/file_clients.md @@ -0,0 +1,547 @@ +# 文件系统(File System) + +文件系统模块 + +## 总览 + +支持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 + + +
+ +## 基础用法 + +```python +from scepter.utils.file_system import FS +from scepter.utils.config import Config + +fs_cfg = Config(load=False, cfg_dict={ + "NAME": "AliyunOssFs", + # ENDPOINT DESCRIPTION: the oss endpoint TYPE: str default: '' + "ENDPOINT": "xxxxx", + # BUCKET DESCRIPTION: the oss bucket TYPE: str default: '' + "BUCKET": "xxxxx", + # OSS_AK DESCRIPTION: the oss ak TYPE: str default: '' + "OSS_AK": "xxxxx", + # OSS_SK DESCRIPTION: the oss sk TYPE: str default: '' + "OSS_SK": "xxxxx", + # TEMP_DIR DESCRIPTION: default is None, means using system cache dir and auto remove! If you set dir, the data will be saved in this temp dir without autoremoving default. TYPE: NoneType default: None + "TEMP_DIR": "cache", + # AUTO_CLEAN DESCRIPTION: when TEMP_DIR is not None, if you set AUTO_CLEAN to True, the data will be clean automatics. TYPE: bool default: False + "AUTO_CLEAN": False +}) + +fs_prefix = FS.init_fs_client(fs_cfg, logger=None) +with FS.get_from("xxxxx", wait_finish=True) as local_object: +# do sth. using local_object here. +# 一次下载多个文件 +generator = FS.get_batch_objects_from(["xxx", "xxxx"]) +for local_path in generator: + print(local_path) +``` +
+ +## **scepter.modules.utils.file_system.FileSystem** + +通过build多类File IO Handler,支持不同类型文件的读写操作。 +
+ +### function **\_\_init\_\_** + +() + +**Parameters** + +
+ +### function **init_fs_client** + +( cfg: *scepter.modules.utils.config.Config* = None, logger = None ) -> str + +通过cfg参数来实例化fs_client,存储在self._prefix_to_clients属性中,可通过prefix来access对应的fs_client. + +**Parameters** + +- **cfg** —— The Config used to build fs_client. If None, use LocalFs as default. +- **logger** —— Instantiated Logger to print or save log. + +**Returns** + +- *str* —— The prefix of instantiated fs_client + +
+ +### function **get_fs_client** + +( target_path: *str*, safe: *bool* = False ) + +通过target_path的前缀来获取对应的fs_client + +**Parameters** + +- **target_path** —— 目标文件路径 +- **safe** —— 安全模式,返回client的copy,否则返回client本身 + +**Returns** + +- *BaseFs* —— 实例化的fs_client + +
+ +### function **get_from** + +( target_path: *str*, local_path: *str* = None, wait_finish: *bool* = False ) -> str + +将远程文件下载到本地 + +**Parameters** + +- **target_path** —— 远程文件路径 +- **local_path** —— 本地文件路径,如果为None则使用cache路径 +- **wait_finish** —— if True,则只有每台机器的0卡下载数据,其他卡等待0卡下载结束 + +**Returns** + +- *str* —— 本地保存文件路径 + +
+ +### function **get_dir_to_local_dir** + +( target_path: *str*, local_path: *str* = None, wait_finish: *bool* = False, timeout: *int* = 3600, worker_id: *int* = 0 ) -> str + +将远程路径的文件夹下载到本地 + +**Parameters** + +- **target_path** —— 远程文件夹路径 +- **local_path** —— 本地文件夹路径,如果为None则使用cache路径 +- **wait_finish** —— if True,则只有每台机器的0卡下载数据,其他卡等待0卡下载结束 +- **timeout** —— 下载超时时间 +- **worker_id** —— Deprecated + +**Returns** + +- *str* —— 本地文件夹路径 + +
+ +### function **get_object** + +( target_path: *str* ) -> byte + +读取远程文件到内存中 + +**Parameters** + +- **target_path** —— 目标文件路径 + +**Returns** + +- *bytes* —— 目标文件的二进制数据 + +
+ +### function **get_object** + +( target_path: *str* ) -> byte + +读取远程文件到内存中 + +**Parameters** + +- **target_path** —— 目标文件路径 + +**Returns** + +- *bytes* —— 目标文件的二进制数据 + +
+ +### function **put_object** + +( local_data: *byte*, target_path: *str* ) -> bool + +上传数据流到指定文件 + +**Parameters** + +- **local_data** —— 本地数据流 +- **target_path** —— 目标文件路径 + +**Returns** + +- *bool* —— 是否上传成功 + +
+ +### function **delete_object** + +( target_path: *str* ) -> bool + +删除目标文件 + +**Parameters** + +- **target_path** —— + +**Returns** + +- *bool* —— 是否删除成功 + +
+ +### function **get_batch_objects_from** + +( target_path_list: *str*, wait_finish: *bool* ) -> *str* + +批量下载文件 + +**Parameters** + +- **target_path_list** —— 下载文件列表 + +**Returns** + +- *local_path* —— 本地文件generator + +
+ +### function **put_batch_objects_to** + +(local_path_list: *str*, target_path_list: *str*, wait_finish: *bool* ) -> *str* + +批量上传文件 + +**Parameters** + +- **local_path_list** —— 上传文件列表 +- **target_path_list** —— 目标文件列表 + +**Returns** + +- *local_path, target_path* —— 返回本地文件和目标文件的pair。 + +
+ +### function **get_object_stream** + +(target_path: *str*, start: *int*, size: *int*, delimiter: *str*) -> *byte*, *int* + +批量上传文件 + +**Parameters** + +- **target_path** —— 目标文件 +- **start** —— 目标流的开始字符位置 +- **size** —— 目标流的开始字符流大小 +- **delimiter** —— 目标流的结束字符 + +**Returns** + +- *local_data, end* —— 返回数据流字节和结束字符位置。 + +
+ +### function **get_object_chunk_list** + +( target_path: *str*, chunk_num: *int* = 1, delimiter: *str* = None ) -> list[bytes] + +获取远程文件且分块 + +**Parameters** + +- **target_path** —— 目标文件路径 +- **chunk_num** —— 分块个数 +- **delimiter** —— 分隔符,确保下载数据是完整的一条,不会从中间截断 + +**Returns** + +- *list[bytes]* —— 分块数据 + +
+ +### function **get_url** + +( target_path: *str*, set_public = False, lifecycle: *int* = 360000 ) -> str + +获取远程文件的url(仅支持AliyunOssFs) + +**Parameters** + +- **target_path** —— 目标文件路径 +- **lifecycle** —— 有效时间 +- **set_public** —— 反馈公开链接 + +**Returns** + +- *str* —— 目标文件的url + +
+ +### function **put_to** + +( target_path: *str* ) + +支持将本地文件上传到远程 + +**Parameters** + +- **target_path** —— 远程文件路径 + +**Returns** + +- **None** + +```python +# 作为上下文管理器使用 +with FS.put_to(target_path) as local_path: + # some operations on local_path. +``` + +
+ +### function **put_object_from_local_file** + +( local_path: *str*, target_path: *str* ) -> bool + +将本地文件push到远程路径 + +**Parameters** + +- **local_path** —— 本地文件路径 +- **target_path** —— 远程文件路径 + +**Returns** + +- *bool* —— 是否上传成功 + +
+ +### function **put_dir_from_local_dir** + +( local_dir: *str*, target_dir: *str* ) -> bool + +将本地文件夹push到远程路径 + +**Parameters** + +- **local_dir** —— 本地文件夹路径 +- **target_dir** —— 远程文件夹路径 + +**Returns** + +- *bool* —— 是否上传成功 + +
+ +### function **add_target_local_map** + +( target_dir: *str*, local_dir: *str* ) -> None + +将远程文件夹和本地文件夹路径的映射关系以key-value对的形式保存到self._target_local_mapper + +**Parameters** + +- **target_dir** —— 远程文件夹路径 +- **local_dir** —— 本地文件夹路径 + +**Returns** + +- *None* + +
+ +### function **make_dir** + +( target_dir: *str* ) -> bool + +创建远程文件夹 + +**Parameters** + +- **target_dir** —— 远程文件夹路径 + +**Returns** + +- *bool* —— 是否创建成功 + +
+ +### function **exists** + +( target_path: *str* ) -> bool + +判断目标路径是否存在 + +**Parameters** + +- **target_path** —— 远程文件路径 + +**Returns** + +- *bool* —— 是否存在 + +
+ +### function **map_to_local** + +( target_path: *str* ) -> str, bool + +将远程文件路径映射到本地路径 + +**Parameters** + +- **target_path** —— 远程文件路径 + +**Returns** + +- *str* —— 本地文件路径 +- *bool* —— 本地文件是否tmp文件 + +
+ +### function **walk_dir** + +( target_dir: *str*, recurse = True) -> Iterator + +获取远程文件夹下的文件列表 + +**Parameters** + +- **target_dir** —— 远程文件夹路径 +- **recurse** —— 是否遍历子文件夹,默认遍历为True + +**Returns** + +- *Iterator* —— 子文件路径列表 + +
+ +### function **is_local_client** + +( target_path: *str* ) -> bool + +判断目标文件client是不是LocalFs + +**Parameters** + +- **target_path** —— 目标文件路径 + +**Returns** + +- *bool* —— client是否LocalFs + +
+ +### function **size** + +( target_path: *str* ) -> int + +判断目标文件大小 + +**Parameters** + +- **target_path** —— 目标文件路径 + +**Returns** + +- *int* —— 目标文件size + +
+ +### function **isfile** + +( target_path: *str* ) -> bool + +判断目标路径是不是object + +**Parameters** + +- **target_path** —— 目标文件路径 + +**Returns** + +- *bool* —— 目标路径是不是object + +
+ +### function **isdir** + +( target_path: *str* ) -> bool + +判断目标路径是不是文件夹 + +**Parameters** + +- **target_path** —— 目标文件路径 + +**Returns** + +- *bool* —— 目标路径是不是文件夹 + +
+ +## **scepter.modules.utils.file_clients.AliyunOssFs** + +- target_path 格式: `oss://{bucket_name}/xxx/yy` + +```yaml +NAME: AliyunOssFs +TEMP_DIR: None +AUTO_CLEAN: False +ENDPOINT: +BUCKET: +OSS_AK: +OSS_SK: +PREFIX: "" +WRITABLE: True +CHECK_WRITABLE: False +RETRY_TIMES: 10 +``` + +
+ +## **scepter.modules.utils.file_clients.LocalFs** + +- target_path 格式: 本地路径 + +```yaml +NAME: LocalFs +TEMP_DIR: None +AUTO_CLEAN: False +``` + +
+ +## **scepter.modules.utils.file_clients.HttpFs** + +- target_path 格式: `http://xx/yy/zz` + +```yaml +NAME: HttpFs +TEMP_DIR: None +AUTO_CLEAN: False +RETRY_TIMES: 10 +``` + +
+ +## **scepter.modules.utils.file_clients.ModelscopeFs** + +- target_path 单文件格式: `ms://{group}/{name}/:{revision}@{filename}` +- target_path 全文件格式: `ms://{group}/{name}/:{revision}` + +```yaml +NAME: ModelscopeFs +TEMP_DIR: None +AUTO_CLEAN: False +RETRY_TIMES: 10 +``` + +
diff --git a/docs/zh_cn/scepter/utils/utils.md b/docs/zh_cn/scepter/utils/utils.md new file mode 100644 index 0000000..1efa161 --- /dev/null +++ b/docs/zh_cn/scepter/utils/utils.md @@ -0,0 +1,941 @@ +# 依赖组件(Utils) + +依赖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) + +### 基础用法 + +```python +from scepter.utils.config import Config + +# 从一个dict对象 初始化 Config对象 +fs_cfg = Config(load=False, cfg_dict={"NAME": "LocalFs"}) +print(fs_cfg.NAME) +# 从一个json文件中初始化 Config对象 +import json + +json.dump({"NAME": "LocalFs"}, open("examples.json", "w")) +fs_cfg = Config(load=True, cfg_file="examples.json") +print(fs_cfg.NAME) +# 从一个yaml文件中初始化 Config对象 +import yaml + +yaml.dump({"NAME": "LocalFs"}, open("examples.yaml", "w")) +fs_cfg = Config(load=True, cfg_file="examples.yaml") +print(fs_cfg.NAME) +# 从 argparse 对象中初始化 Config对象,该模式下cfg的参数为必需参数,否则会报错。 +import argparse + +parser = argparse.ArgumentParser( + description="Argparser for Cate process:\n" +) +parser.add_argument( + "--stage", + dest="stage", + help="Running stage!", + default="train" +) +fs_cfg = Config(load=True, parser_ins=parser) +print(fs_cfg.args) +``` +
+ +### function **__init__** + +( cfg_dict: dict = {}, load = True, cfg_file = None, logger = None, parser_ins: argparse.ArgumentParser = None ) + +**Parameters** + +- **cfg_dict** —— 包含参数的dict,默认为{}。 +- **load** —— 为True时说明需要从文件或argparse中载入参数。 +- **cfg_file** —— 支持从json文件或者yaml文件中载入参数。 +- **logger** —— 日志示例,如果为None,则会默认初始化一个stdio的日志实例。 +- **parser_ins** —— argparse实例,默认有cfg参数,用于传入参数文件。 +-- parser_ins 默认会加入系统参数,说明如下: + - cfg(--cfg) 用于指定参数文件位置 + - local_rank(--local_rank) torchrun默认读取参数,默认为0,可不管 + - launcher(-l) 启动代码的方式,默认为spawn,可选值为 torchrun + - data_online(-d) 设置全局下载数据不落盘,在pai集群上应设置该值 + - share_storage(-s) 设置全局下载数据是否共享文件系统,如nas。当设置时,说明文件系统不同节点互通,此时只需在rank=0时下载即可;当不设置时,说明 +是在不同的节点进行数据下载,此时应该只在device_id=0时下载。 + +### function **dict_to_yaml** +( module_name: str, name: str, json_config: dict, set_name: bool = False ) + +**Parameters** + +- **module_name** —— 模块名称,用于在模版开始说明是哪个模块的模版。 +- **name** —— Name字段的默认名称。 +- **json_config** —— 参数说明,需要满足{}(表示依赖一个子模块), [](依赖多个子模块), {"value":"", "description":""} (叶子参数值)。 +- **set_name** —— 是否设置Name字段。 + +**Returns** + +- **str** —— 模版文本 + + + +## 2. 路径sdk(scepter.modules.utils.directory) +一些常用的路径函数 +### 基础用法 + +```python +from scepter.utils.directory import osp_path + +# 根据路径前缀进行自动化路径拼接 +prefix = "xxxx" +data_file = "example_videos/1.mp4" +# 输出为 xxxx/example_videos/1.mp4 +print(osp_path(prefix, data_file)) +# 输出也为 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 + +# 根据路径获取指定层级的文件夹路径 +# 默认最后一级 xxxx/example_videos/ +print(get_relative_folder(data_file)) +# 倒数第二级 xxxx/ +print(get_relative_folder(data_file, keep_index=-2)) + +from scepter.utils.directory import get_md5 + +# 获取文本/路径的md5码 34a447fb46d0b786a3999c9dad01d470 +print(get_md5(data_file)) +``` +
+ +### function **osp_path** + +( prefix: str, data_file: str ) -> str + +根据路径前缀进行自动化路径拼接 + +**Parameters** + +- **prefix** —— 路径前缀。 +- **data_file** —— 文件路径。 + +**Returns** + +- **str** —— 拼接以后的路径 + +### function **get_relative_folder** + +( abs_path: str, keep_index: int = -1 ) -> str + +根据路径获取指定层级的文件夹路径 + +**Parameters** + +- **abs_path** —— 文件路径。 +- **keep_index** —— 保留层级,-1代表倒数第一级,-2 为倒数第二级。 + +**Returns** + +- **str** —— 解析以后的路径 + +### function **get_md5** + +( ori_str: str) -> str + +根据字符串/路径获取md5码 + +**Parameters** + +- **ori_str** —— 文件路径或字符串。 + +**Returns** + +- **str** —— md5码 + +## 3. torch分布式sdk(scepter.modules.utils.distribute) +torch分布式初始化sdk,使用该sdk,可以让用户不要关注torch的分布式初始化的实现。 +### 基础用法 + +```python +from scepter.utils.distribute import we +from scepter.utils.config import Config + +cfg = Config(cfg_dict={}, load=False) + + +def fn(): + pass + + +print(we) +# 启动任务 +we.init_env(cfg, fn, logger=None) +``` +
+ +### class **Workenv** + +这是一个用于统一管理运行环境的类,通常不需要使用该类做初始化,在scepter.modules.utils.distribute +中会初始化一个全局的实例we,用于管理一些关键性的标志变量。 + +- 关于we的一些参数,具体说明如下: + - initialized 标记是否初始化torch的process group,默认为False。 + - is_distributed 标记当前是否为分布式运行,默认为False。 + - sync_bn 标记是否使用sync_bn,默认为False。 + - rank 标记当前process的rank,默认为0。 + - world_size 标记当前所有的进程数,默认为1。 + - device_id 标记当前使用的设备ID,默认为0。 + - device_count 标记当前环境下的所有设备数,默认为1。 + - use_pl 标记当前环境是否使用pytorch_lighting引擎,默认为False + - launcher 标记当前环境的启动方式,默认为spawn。 + - data_online 标记当前环境下io部分的数据是否落盘,默认为False。 + - share_storage 标记当前环境下不同节点是否使用相同的文件系统,如nas,默认为False。 + +### function **we.init_env** + +( config: scepter.modules.utils.config.Config, fn: function, logger: logging.Logger = None ) + +作为启动任何任务的执行入口。 + +**Parameters** + +- **config** —— 传入的参数实例。 +- **fn** —— 需要执行的函数。 +- **logger** —— 标准的日志实例。 + +### function **we.get_env** +() -> dict + +获取we的所有类内参数,以dict的形式存储。 + + +### function **we.set_env** +(we_env: dict) + +重新设置we的所有类内参数,以dict的形式作为输入。 + +**Parameters** + +- **we_env** —— dict,每个key代表一个类内变量。 + +### function **get_dist_info** +() -> int, int + +获取环境的rank/world size,这个是直接通过torch的方法来获取的,一般用于当初始化环境的方式 +不是we.init_env的时候使用。 + +**Returns** + +- **rank** —— 当前进程的rank值,默认为0 +- **world_size** —— 当前环境的总进程数, 当单进程时为1。 + +### function **gather_data** +(data: [list, dict, tensor, object] ) -> data + +通过scepter.distributed.all_gather将任意实例收集起来,并在rank=0进程合并为一个汇总后的实例。 + +**Parameters** + - **data** —— 支持dict/list,其中元素支持任意实例或者tensor。 + +**Returns** + - **data** —— 一个和输入data相同结构的汇总过的数据。 + +### function **gather_list** +(data: [list] ) -> data + +通过scepter.distributed.all_gather将任意实例收集起来,并在rank=0进程合并为一个汇总后的实例。 + +**Parameters** + - **data** —— 支持list,其中元素支持任意实例或者tensor。 + +**Returns** + - **data** —— 一个和输入data相同结构的汇总过的数据。 + +### function **gather_picklable** +(data: [object] ) -> data + +通过scepter.distributed.all_gather将任意实例收集起来,并在rank=0进程合并为一个汇总后的实例。 + +**Parameters** + - **data** —— 为一个可序列化的实例。 + +**Returns** + - **data** —— 一个和输入data相同结构的汇总过的数据。 + +### function **broadcast** +(tensor: **torch.Tensor**, src: **str**, group: **list** ) + +为torch.distributed.broadcast的优化版本,自动确认是否为分布式环境。 + +**Parameters** + - **tensor** —— 需要广播的tensor。 + - **src** —— 需要广播的源设备。 + - **group** —— 需要广播的组。 + +**Returns** + - **data** —— 一个和输入data相同结构的汇总过的数据。 +* 其他函数如barrier、all_reduce、 reduce、send、recv、isend、irecv、scatter也做了此操作。 + + +### function **gather_gpu_tensors** +(tensor: torch.Tensor ) -> tensor: torch.Tensor + +通过torch.distributed.all_gather将gpu tensor收集起来,并在rank=0合并并转到cpu上。 +因为涉及到clone,因此有可能造成额外的显存浪费。 + +**Parameters** + - **tensor** —— 输入的gpu上的tensor。 + +**Returns** + - **tensor** —— 输出的在进程rank=0上的cpu的tensor。 + +## 4. 模型导出sdk(scepter.utils.export_model) + 用于模型导出为torchscript/Onnx格式的api。 +### 基础用法 + +```python +from scepter.utils.export_model import save_develop_model_multi_io + +save_develop_model_multi_io( + model, + input_size, + input_type, + input_name, + output_name, + limit, + save_onnx_path=None, + save_pt_path=None +) +``` +
+ +### function **save_develop_model_multi_io** +(model: torch.nn.Module, input_size: list, input_type: list, input_name: list, +output_name: list, limit: list, save_onnx_path: str = None, save_pt_path: str = None) -> pt_module, onnx_module + +支持多输入多输出的模型导入和导出 + +**Parameters** + - **model** —— 待导出的模型实例 + - **input_size** —— 为一个list,每个元组包含数据的shape信息,如[[1, 3, 224, 224]]。 + - **input_type** —— 为一个list,每个元组包含数据的type信息,与input_size一一对应,可选值为("float32", +"float16","int8","int16","int32","int64")。如["float32"]。 + - **input_name** —— 为一个list,为onnx的每个输入变量命名,如["image"],与上述input_size、 +input_type 一一对应。 + - **output_name** —— 为一个list,为onnx的每个输出变量命名,如["output"] + - **limit** —— 为一个list,每个元祖定义了该输入的上下届,如[[-1, 1]],代表了image的输入张量在-1~1之间。 + - **save_onnx_path** —— 不为None时,会导出onnx模型,存储在该位置。 + - **save_pt_path** —— 不为None时,会导出torchscript模型,存储在该位置。 + + + +**Returns** + - **tensor** —— 输出的在进程rank=0上的cpu的tensor。 + +## 5. 文件系统sdk(scepter.utils.file_system) +参考[file_clients](file_clients.md) + +## 6. 日志sdk(scepter.utils.logger) +用于实例化一个标准的日志实例,用于打印信息。 + +### 基础用法 + +```python +from scepter.utils.logger import get_logger, init_logger + +std_logger = get_logger(name="std_torch") +init_logger(std_logger, log_file="", dist_launcher="pytorch") +``` +
+ +### function **get_logger** +(name: str) -> logger + +获取日志实例。 + +**Parameters** + - **name** —— 日志前缀,每次打印会首先打印该前缀。 + +**Returns** + - **logger** —— 返回一个logging实例。 + +### function **init_logger** +(in_logger: logger, log_file: str) -> logger + +二次初始化日志实例,可以为该实例分配一个文件落盘。 + +**Parameters** + - **in_logger** —— 已有的日志实例。 + - **log_file** —— 希望存储的文件位置。 + - **dist_launcher** —— 已经不重要了,deprecated + +### function **as_time** +(s: int) -> str + +时间s转换为标准的xxx days xxx hours xxx mins xxx secs + +**Parameters** + - **s** —— 代表秒数s。 + +**Returns** + - **str** —— 格式化的输出。 + +### function **time_since** +(since: int, percent: float) -> str + +根据当前用时和百分比计算距离结束的时间。 + +**Parameters** + - **since** —— 代表当前已经用的时间。 + - **percent** —— 代表当前已经执行的百分比。 + +**Returns** + - **str** —— 格式化的输出。 + +## 7. 视频处理sdk(scepter.utils.video_reader) +用于处理视频读取的api。 + +### 基础用法 + +```python +from scepter.utils.video_reader.frame_sampler import do_frame_sample +from scepter.utils.video_reader.video_reader import ( + VideoReaderWrapper, EasyVideoReader, FramesReaderWrapper +) +``` +
+ +### function **do_frame_sample** +(sampling_type: str, vid_len: int, vid_fps: int, num_frames: int, kwargs) -> list + +获取针对于视频的帧采样器。 + +**Parameters** + - **sampling_type** —— 采样器类型,目前支持的UniformSampler(均匀采样器)、IntervalSampler(等间隔采样器)、SegmentSampler(切片采样器)。 + - **vid_len** —— 视频长度。 + - **vid_fps** —— 视频的帧率。 + - **num_frames** —— 视频包含的帧数。 + - **kwargs** —— 对应采样器的需要参数,需要参考对应采样器的源码。 + +**Returns** + - **list** —— 采样帧结果。 + +### class **VideoReaderWrapper** +读取视频的标准类,底层解码器为decord + +#### function **VideoReaderWrapper.__init__** +(video_path: str) + +初始化视频实例 + +**Parameters** + - **video_path** —— 视频链接。 + +#### function **VideoReaderWrapper.len** +() -> int + +获取视频帧总数。 + +**Returns** + - **int** —— 视频帧数。 + +#### function **VideoReaderWrapper.fps** +() -> float + +获取视频帧率 + +**Returns** + - **float** —— 视频帧率。 + +#### function **VideoReaderWrapper.duration** +() -> float + +获取视频时长 + +**Returns** + - **float** —— 视频时长。 + +#### function **VideoReaderWrapper.sample_frames** +(decode_list: torch.Tensor) -> torch.Tensor + +根据帧号,获取帧数据 + +**Parameters** + - **decode_list** —— 采样帧号列表。 +**Returns** + - **tensor** —— 数据张量。 + +### class **FramesReaderWrapper** +给定解完帧的文件夹,按顺序读取帧数据 + +#### function **FramesReaderWrapper.__init__** +(frame_dir: str, extract_fps: float, suffix: str) + +初始化视频实例 + +**Parameters** + - **frame_dir** —— 帧文件夹。 + - **extract_fps** —— 提取帧的fps。 + - **suffix** —— 帧文件的后缀,默认为jpg。 + +#### function **FramesReaderWrapper.len** +() -> int + +获取视频帧总数。 + +**Returns** + - **int** —— 视频帧数。 + +#### function **FramesReaderWrapper.fps** +() -> float + +获取视频帧率 + +**Returns** + - **float** —— 视频帧率。 + +#### function **FramesReaderWrapper.duration** +() -> float + +获取视频时长 + +**Returns** + - **float** —— 视频时长。 + +#### function **FramesReaderWrapper.sample_frames** +(decode_list: torch.Tensor) -> torch.Tensor + +根据帧号,获取帧数据 + +**Parameters** + - **decode_list** —— 采样帧号列表。 +**Returns** + - **tensor** —— 数据张量。 + +### class **EasyVideoReader** +用于长视频读取、采样和预处理的类。 + +#### function **EasyVideoReader.__init__** +(video_path: str, num_frames: int, clip_duration: Union[float, Fraction, str], +overlap: Union[float, Fraction, str] = Fraction(0), transforms: Optional[Callable] = None) + +初始化视频实例 + +**Parameters** + - **video_path** —— 视频链接。 + - **num_frames** —— 视频帧数。 + - **clip_duration** —— 单片段长度。 + - **overlap** —— 片段间重合比例。 + - **transforms** —— 预处理算子。 + +#### function **EasyVideoReader.__iter__** +() -> int + +迭代器 + +#### function **EasyVideoReader.__next__** +() -> float + +迭代器,每迭代一次,返回一个片段的tensor。 + +**Returns** + - **tensor** —— 视频片段的tensor。 + +## 8. 模块注册sdk(scepter.utils.registry) +用于管理各种注册的类。 + +### 基础用法 + +```python +from scepter.utils.registry import Registry +from scepter.utils.config import Config + +MODELS = Registry('MODELS') + + +@MODELS.register_class() +class ResNet(object): + pass + + +config = Config(load=False, cfg_dict={"NAME": "ResNet"}) +resnet = MODELS.build(config) +``` +
+ +### class **Registry** +注册器 + +#### function **Registry.__init__** +(name: str, build_func: function = None, common_para: Config = None, allow_types: tuple = ("class", "function")) + +初始化注册模块实例 + +**Parameters** + - **name** —— 模块名。 + - **build_func** —— build模块的时候调用的function。 + - **common_para** —— 该模块下的公共参数。 + - **allow_types** —— 该模块允许注册的类或者函数,默认都允许注册。 + +#### function **Registry.build** +(cfg: Config, logger: logger = None, kwargs) -> cls_obj + +build目标类的实例 + +**Returns** + - **cls_obj** —— 特定类的实例。 + +#### function **Registry.register_class** +(name: str) + +注册一个类 + +**Returns** + - **name** —— 注册名称。 + +#### function **Registry.register_function** +(name: str) + +注册一个函数 + +**Returns** + - **name** —— 注册名称。 + +## 9. 数据sdk(scepter.utils.data) +用于数据在设备间转移 + +### 基础用法 + +```python +import torch +from scepter.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) +transfer_data_to_cpu(data) +transfer_data_to_cuda(data) +``` +
+ + +#### function **transfer_data_to_numpy** +(data: list/dict of torch.Tensor) -> (data: list/dict of numpy.ndarray) + +将数据转移到numpy + +**Parameters** + - **data** —— torch.Tensor并以list/dict形式存储。 +**Returns** + - **data** —— numpy.ndarray并与输入一致的形式存储。 + +#### function **transfer_data_to_cpu** +(data: list/dict of torch.Tensor(cuda)) -> (data: list/dict of torch.Tensor(cpu)) + +将gpu数据转移到cpu + +**Parameters** + - **data** —— torch.Tensor[CUDA]并以list/dict形式存储。 +**Returns** + - **data** —— torch.Tensor[CPU]并与输入一致的形式存储。 + +#### function **transfer_data_to_cuda** +(data: list/dict of torch.Tensor(cpu)) -> (data: list/dict of torch.Tensor(cuda)) + +将cpu数据转移到gpu + +**Parameters** + - **data** —— torch.Tensor[CPU]并以list/dict形式存储。 +**Returns** + - **data** —— torch.Tensor[CUDA]并与输入一致的形式存储。 + +## 10. 模型sdk(torch.utils.model) +用于对模型进行加载、评估等操作 + +### 基础用法 + +```python +import torch +from scepter.utils.model import move_model_to_cpu, load_pretrained, + count_params, init_weights +``` +
+ + +#### function **move_model_to_cpu** +(params: list/dict of torch.Tensor[cuda]) -> (data: torch.Tensor[cpu]) + +将参数数据从gpu转移到cpu上。 + +**Parameters** + - **params** —— torch.Tensor[cuda]并以OrderedDict形式存储。 +**Returns** + - **params** —— torch.Tensor[cpu]并与输入一致的形式存储。 + +#### function **load_pretrained** +(model: torch.nn.Module, path: str, map_location="cpu", logger=None, + sub_level=None) + +加载参数到模型。 + +**Parameters** + - **model** —— torch.nn.Module模型实例。 + - **path** —— 预训练模型参数。 + - **map_location** —— cpu/cuda。 + - **logger** —— 标准日志实例。 + - **sub_level** —— 比如ddp时需要索引子层级。 + + +#### function **count_params** +(model: torch.nn.Module) -> (float) + +统计模型的总参数。 + +**Parameters** + - **model** —— torch.nn.Module模型实例。 +**Returns** + - **float** —— 模型参数量(浮点数个数)。 + +#### function **init_weights** +(model: torch.nn.Module) + +对模型模块进行参数初始化。 + +**Parameters** + - **module** —— torch.nn.Module模型实例。 + +## 11. 采样器sdk(scepter.utils.sampler) +采样器比较具有通用性,大多数情况下不会进行定制开发,这里提供了几类常用的sampler采样器。 + +### 基础用法 + +```python +import torch +from scepter.utils.sampler import MultiFoldDistributedSampler, + EvalDistributedSampler, MultiLevelBatchSampler, MixtureOfSamplers +``` +
+ + +#### class **MultiFoldDistributedSampler** + +多fold采样器,支持在一个epoch中重复多轮数据 + +#### function **MultiFoldDistributedSampler.__init__** + +( dataset: torch.data.dataset, num_folds=1, num_replicas=None, rank=None, shuffle=True) + +**Parameters** + +- **dataset** —— torch.data.dataset类实例 +- **num_folds** —— int,表示数据重复的轮数。 +- **num_replicas** —— 表示数据分割的片数,一般和world-size保持一致。 +- **rank** —— rank表示当前进程号。 +- **shuffle** —— 数据是否要打乱。 + +#### function **MultiFoldDistributedSampler.__iter__** +() + +迭代器,每迭代一次得到一个样本的index + +#### function **MultiFoldDistributedSampler.set_epoch** +(epoch: int) + +设置当前的epoch + +**Parameters** + +- **epoch** —— 当前的epoch。 + + +#### class **EvalDistributedSampler** + +用于测试时的采样器,当不用padding模式的时候,会发现最后一个rank的数据会少于其他rank。 + +#### function **EvalDistributedSampler.__init__** + +( dataset: torch.data.dataset, num_replicas: Optional[int] =None, rank: Optional[int] =None, padding: bool =False) + +**Parameters** + +- **dataset** —— torch.data.dataset类实例 +- **num_replicas** —— 表示数据分割的片数,一般和world-size保持一致。 +- **rank** —— rank表示当前进程号。 +- **padding** —— 数据是否需要padding,如果padding则能保证最后一个rank和其他rank数据量一致。 + +#### function **EvalDistributedSampler.__iter__** +() + +迭代器,每迭代一次得到一个样本的index + +#### function **EvalDistributedSampler.set_epoch** +(epoch: int) + +设置当前的epoch + +**Parameters** + +- **epoch** —— 当前的epoch。 + +#### class **MultiLevelBatchSampler** + +用于大规模数据的多级索引的sampler + +#### function **MultiLevelBatchSampler.__init__** + +(index_file: str, batch_size: int, rank: int =0, seed: int = 8888) + +**Parameters** + +- **index_file** —— 多级数据索引的索引文件。 +- **batch_size** —— 一个batch的大小。 +- **rank** —— rank表示当前进程号。 +- **seed** —— 随机采样的seed,在data.registry中获取全局seed。 + +#### function **MultiLevelBatchSampler.__iter__** +() + +迭代器,每迭代一次得到一个样本的index + + +#### class **MixtureOfSamplers** + +用于大规模数据的多级索引的sampler + +#### function **MixtureOfSamplers.__init__** + +(samplers: list(sampler), probabilities: list(float), rank: int =0, seed: int = 8888) + +**Parameters** + +- **samplers** —— 采样器列表,用于混合采样器。 +- **probabilities** —— 每个采样器的概率。 +- **rank** —— rank表示当前进程号。 +- **seed** —— 随机采样的seed,在data.registry中获取全局seed。 + +#### function **MixtureOfSamplers.__iter__** +() + +迭代器,每迭代一次得到一个样本的index + +## 12. 探针器sdk(scepter.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 + + +class TestModel(BaseModel): + def forward(self, data): + self.register_probe(data) + # 测试ProbeData示例 + view_distribute + self.register_probe( + {"data_key_dist": ProbeData(data["data_key"], view_distribute=True), + "data_folder": ProbeData(data["data_folder"], view_distribute=True)} + ) + + +class TestModel2(BaseModel): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.test_model = TestModel(cfg, logger=logger) + + def forward(self, data): + # 测试list of str done + # 测试dict of str done + # 测试dict of number done + # 测试list of number done + # 测试number done + # 测试str done + # 测试np.array done + # 测试list of np.ndarray 必须手动建立ProbeData done + # 测试2D 图 必须手动建立ProbeData + # 测试3D 图 必须手动建立ProbeData + # 测试3D 多个2维图 必须手动建立ProbeData + # 测试3D list 图 必须手动建立ProbeData + # 测试4D Array 图 done + # 测试4D Array 图 save_html + # 测试4D List 图 save_html + self.register_probe(data) + self.register_probe({ + "test_np_list": ProbeData([np.zeros([40, 40, 3]).astype(dtype=np.int8) for _ in range(5)]), + "test_2d_img": ProbeData(np.zeros([40, 40]).astype(dtype=np.uint8), is_image=True), + "test_3d_n2d_img": ProbeData(np.zeros([10, 40, 40]).astype(dtype=np.uint8), is_image=True), + "test_3d_img": ProbeData(np.zeros([40, 40, 3]).astype(dtype=np.uint8), is_image=True), + "test_list_3d_img": ProbeData([np.zeros([40, 40, 3]).astype(dtype=np.uint8) for _ in range(5)], + is_image=True), + "test_4d_img": ProbeData(np.zeros([10, 40, 40, 3]).astype(dtype=np.uint8), is_image=True), + "test_4d_img_html": ProbeData(np.zeros([10, 40, 40, 3]).astype(dtype=np.uint8), is_image=True, + build_html=True, build_label="4d_data"), + "test_4d_img_list_html": ProbeData([np.zeros([10, 40, 40, 3]).astype(dtype=np.uint8) for _ in range(5)], + is_image=True, + build_html=True, build_label=[f"4d_data_{i}" for i in range(5)]), + }) + # 测试嵌套类型 + self.test_model(data) + + +cfg = Config(cfg_file="./config/general_config.yaml") +if cfg.have("FILE_SYSTEMS"): + for file_sys in cfg.FILE_SYSTEMS: + fs_prefix = FS.init_fs_client(file_sys) +else: + fs_prefix = FS.init_fs_client(cfg) + +_model = TestModel2(cfg) + +data = { + "data_key": [1, 1], + "data_folder": {"mj": 1., "mj_square": 2.}, + "timestamp": 1, + "valid_str": "right", + "test_np": np.zeros([40, 40, 3]).astype(dtype=np.int8) +} +_model(data) +probe = _model.probe_data() +for key in probe: + print(key, probe[key].to_log(prefix=f"xxx/dev_easytorch/{key}")) +``` +
+ +配合Hook使用如下(其中PROB_INTERVAL探针存储间隔,即调用probe_data()的次数): + +```yaml +- + NAME: ProbeDataHook + PROB_INTERVAL: 100 +``` +#### class **ProbeData** + +探针数据的实例。 + +#### function **ProbeData.__init__** + +(data, is_image = False, build_html = False, build_label = None, view_distribute = False) + +**Parameters** +- **data** —— 传入的探针数据,目前支持str、Number、list、dict、tensor。 +- **is_image** —— 是否要存为图像。 +- **build_html** —— 是否存为html。 +- **build_label** —— 填入保存html的label html。 +- **view_distribute** —— 针对一些值统计频率。 diff --git a/environment.yaml b/environment.yaml new file mode 100644 index 0000000..f3555a6 --- /dev/null +++ b/environment.yaml @@ -0,0 +1,9 @@ +name: scepter +channels: + - defaults +dependencies: + - python==3.8 + - pip>=20.3 + - numpy>=1.23.1 + - pip: + - -r requirements.txt diff --git a/readme.md b/readme.md new file mode 100644 index 0000000..fc5417e --- /dev/null +++ b/readme.md @@ -0,0 +1,122 @@ +# 🪄SCPTER + +

+ + + +

+ +## 📖 Table of Contents +- [Introduction](#-introduction) +- [News](#-news) +- [Installation](#-installation) +- [Getting Started](#-getting-started) +- [Learn More](#-learn-more) +- [License](#license) + +## 📝 Introduction + +SCEPTER is an open-source code repository dedicated to generative training, fine-tuning, and inference, encompassing a suite of downstream tasks such as image generation, transfer, editing. It integrates popular community-driven implementations as well as proprietary methods by Tongyi Lab of Alibaba Group, offering a comprehensive toolkit for researchers and practitioners in the field of AIGC. This versatile library is designed to facilitate innovation and accelerate development in the rapidly evolving domain of generative models. + +Main Feature: + +- Training: + - distribute: DDP / FSDP / FairScale +- Inference + - text-to-image generation + - controllable image synthesis (TODO) +- Deploy-Gradio (TODO) + - fine-tuning + - inference + +Currently supported approches (and counting): + +1. SD Series: [Stable Diffusion v1.5](https://huggingface.co/runwayml/stable-diffusion-v1-5) / [Stable Diffusion v2.1](https://huggingface.co/runwayml/stable-diffusion-v1-5) / [Stable Diffusion XL](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0) +2. SCEdit: [SCEdit: Efficient and Controllable Image Diffusion Generation via Skip Connection Editing](https://arxiv.org/abs/2312.11392) [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=SCEdit&color=red&logo=arxiv)](https://arxiv.org/abs/2312.11392) [![Page link](https://img.shields.io/badge/Page-SCEdit-Gree)](https://scedit.github.io/) +3. Res-Tuning(TODO): [Res-Tuning: A Flexible and Efficient Tuning Paradigm via Unbinding Tuner from Backbone](https://arxiv.org/abs/2310.19859) [![Arxiv link](https://img.shields.io/static/v1?label=arXiv&message=ResTuning&color=red&logo=arxiv)](https://arxiv.org/abs/2310.19859) [![Page link](https://img.shields.io/badge/Page-ResTuning-Gree)](https://res-tuning.github.io/) + +## 🎉 News +- [2023.12]: We propose [SCEdit](https://arxiv.org/abs/2312.11392), an efficient and controllable generation framework. +- [2023.12]: We release [🪄SCEPTER](https://github.com/modelscope/scepter/) library. + +## 🛠️ Installation + +- Create new environment + +```shell +conda env create -f environment.yaml +conda activate scepter +``` + +- Install SCEPTER by the `pip` command: + +```shell +pip install scepter +``` + +## 🚀 Getting Started + +### Dataset + +#### Text-to-Image generation + +We use a [custom-stylized dataset](https://modelscope.cn/datasets/damo/style_custom_dataset/summary), which included classes 3D, anime, flat illustration, oil painting, sketch, and watercolor, each with 30 image-text pairs. + +```python +# pip install modelscope +from modelscope.msdatasets import MsDataset +ms_train_dataset = MsDataset.load('style_custom_dataset', namespace='damo', subset_name='3D', split='train_short') +print(next(iter(ms_train_dataset))) +``` + +### Training + +#### Text-to-Image generation + +- SCEdit + +```python +# SD v1.5 +python scepter/tools/run_train.py --cfg scepter/methods/SCEdit/t2i_sd15_512_sce.yaml +# SD v2.1 +python scepter/tools/run_train.py --cfg scepter/methods/SCEdit/t2i_sd21_768_sce.yaml +# SD XL +python scepter/tools/run_train.py --cfg scepter/methods/SCEdit/t2i_sdxl_1024_sce.yaml +``` + +- Existing strategies +```python +# fully-tuning on SD v1.5 +python scepter/tools/run_train.py --cfg scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml +# lora-tuning on SD v2.1 +python scepter/tools/run_train.py --cfg scepter/methods/examples/generation/stable_diffusion_2.1_768_lora.yaml +``` +#### Controllable Image Synthesis + +TODO + +### Inference + +```python +# generation on SD v1.5 +python scepter/tools/run_inference.py --cfg scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml --prompt 'a cute dog' --save_folder 'inference' +# generation on SD v2.1 +python scepter/tools/run_inference.py --cfg scepter/methods/examples/generation/stable_diffusion_2.1_768.yaml --prompt 'a cute dog' --save_folder 'inference' +# generation on SD XL +python scepter/tools/run_inference.py --cfg scepter/methods/examples/generation/stable_diffusion_xl_1024.yaml --prompt 'a cute dog' --save_folder 'inference' +``` + + +## 🔍 Learn More + +- [ModelScope library](https://github.com/modelscope/modelscope/) + + ModelScope Library is the model library of ModelScope project, which contains a large number of popular models. + +- [Alibaba TongYi Vision Intelligence Lab](https://github.com/damo-vilab) + + Discover more about open-source projects on image generation, video generation, and editing tasks. + +## License + +This project is licensed under the [Apache License (Version 2.0)](https://github.com/modelscope/modelscope/blob/master/LICENSE). diff --git a/scepter/__init__.py b/scepter/__init__.py new file mode 100644 index 0000000..3613881 --- /dev/null +++ b/scepter/__init__.py @@ -0,0 +1,18 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import os + +import scepter +from scepter.modules import data, model, opt, solver, transform, utils +from scepter.tools.helper import get_module_list as module_list +from scepter.tools.helper import \ + get_module_object_config as configures_by_objects +from scepter.tools.helper import get_module_objects as objects_by_module +from scepter.version import __version__, version_info + +dirname = os.path.dirname(scepter.__file__) + +__all__ = [ + utils, transform, data, model, solver, version_info, opt, '__version__', + 'dirname' +] diff --git a/scepter/methods/SCEdit/t2i_sd15_512_sce.yaml b/scepter/methods/SCEdit/t2i_sd15_512_sce.yaml new file mode 100644 index 0000000..defb3b4 --- /dev/null +++ b/scepter/methods/SCEdit/t2i_sd15_512_sce.yaml @@ -0,0 +1,226 @@ +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/t2i_sd15_512_sce + 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 + # + 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 + # + LOSS: + NAME: ReconstructLoss + LOSS_TYPE: l2 + # + SAMPLE_ARGS: + SAMPLER: ddim + SAMPLE_STEPS: 50 + SEED: 2023 + GUIDE_SCALE: 7.5 + GUIDE_RESCALE: + DISCRETIZATION: trailing + IMAGE_SIZE: [512, 512] + RUN_TRAIN_N: False + # + OPTIMIZER: + NAME: AdamW + LEARNING_RATE: 0.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 diff --git a/scepter/methods/SCEdit/t2i_sd21_768_sce.yaml b/scepter/methods/SCEdit/t2i_sd21_768_sce.yaml new file mode 100644 index 0000000..96b15c8 --- /dev/null +++ b/scepter/methods/SCEdit/t2i_sd21_768_sce.yaml @@ -0,0 +1,222 @@ +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/t2i_sd21_768_sce + 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 + # + MODEL: + NAME: LatentDiffusion + PARAMETERIZATION: v + TIMESTEPS: 1000 + MIN_SNR_GAMMA: + ZERO_TERMINAL_SNR: False + PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-2-1@v2-1_768-ema-pruned.safetensors + IGNORE_KEYS: [ ] + SCALE_FACTOR: 0.18215 + SIZE_FACTOR: 8 + DEFAULT_N_PROMPT: + SCHEDULE_ARGS: + "NAME": "scaled_linear" + "BETA_MIN": 0.00085 + "BETA_MAX": 0.012 + USE_EMA: False + # + DIFFUSION_MODEL: + NAME: DiffusionUNet + IN_CHANNELS: 4 + OUT_CHANNELS: 4 + MODEL_CHANNELS: 320 + NUM_HEADS_CHANNELS: 64 + NUM_RES_BLOCKS: 2 + ATTENTION_RESOLUTIONS: [ 4, 2, 1 ] + CHANNEL_MULT: [ 1, 2, 4, 4 ] + CONV_RESAMPLE: True + DIMS: 2 + USE_CHECKPOINT: False + USE_SCALE_SHIFT_NORM: False + RESBLOCK_UPDOWN: False + USE_SPATIAL_TRANSFORMER: True + TRANSFORMER_DEPTH: 1 + CONTEXT_DIM: 1024 + DISABLE_MIDDLE_SELF_ATTN: False + USE_LINEAR_IN_TRANSFORMER: True + PRETRAINED_MODEL: + # + FIRST_STAGE_MODEL: + NAME: AutoencoderKL + EMBED_DIM: 4 + PRETRAINED_MODEL: + IGNORE_KEYS: [ ] + BATCH_SIZE: 4 + # + ENCODER: + NAME: Encoder + CH: 128 + OUT_CH: 3 + NUM_RES_BLOCKS: 2 + IN_CHANNELS: 3 + ATTN_RESOLUTIONS: [ ] + CH_MULT: [ 1, 2, 4, 4 ] + Z_CHANNELS: 4 + DOUBLE_Z: True + DROPOUT: 0.0 + RESAMP_WITH_CONV: True + # + DECODER: + NAME: Decoder + CH: 128 + OUT_CH: 3 + NUM_RES_BLOCKS: 2 + IN_CHANNELS: 3 + ATTN_RESOLUTIONS: [ ] + CH_MULT: [ 1, 2, 4, 4 ] + Z_CHANNELS: 4 + DROPOUT: 0.0 + RESAMP_WITH_CONV: True + GIVE_PRE_END: False + TANH_OUT: False + # + TOKENIZER: + NAME: OpenClipTokenizer + LENGTH: 77 + # + COND_STAGE_MODEL: + NAME: FrozenOpenCLIPEmbedder + ARCH: ViT-H-14 + PRETRAINED_MODEL: + LAYER: penultimate + # + LOSS: + NAME: ReconstructLoss + LOSS_TYPE: l2 + # + SAMPLE_ARGS: + SAMPLER: ddim + SAMPLE_STEPS: 50 + SEED: 2023 + GUIDE_SCALE: 7.5 + GUIDE_RESCALE: + DISCRETIZATION: trailing + IMAGE_SIZE: [768, 768] + 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: 768 + INTERPOLATION: bilinear + INPUT_KEY: [ 'img' ] + OUTPUT_KEY: [ 'img' ] + BACKEND: pillow + - NAME: CenterCrop + SIZE: 768 + 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: [768, 768] + 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 diff --git a/scepter/methods/SCEdit/t2i_sdxl_1024_sce.yaml b/scepter/methods/SCEdit/t2i_sdxl_1024_sce.yaml new file mode 100644 index 0000000..b167c5e --- /dev/null +++ b/scepter/methods/SCEdit/t2i_sdxl_1024_sce.yaml @@ -0,0 +1,335 @@ +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/t2i_sdxl_1024_sce + 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 + # + 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: + 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 + IN_KEYS: [ 'img', 'img_original_size_as_tuple', 'img_target_size_as_tuple', 'img_crop_coords_top_left' ] + OUT_KEYS: [ '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 \ No newline at end of file diff --git a/scepter/methods/examples/classification/example.yaml b/scepter/methods/examples/classification/example.yaml new file mode 100644 index 0000000..3f4bc6c --- /dev/null +++ b/scepter/methods/examples/classification/example.yaml @@ -0,0 +1,260 @@ +ENV: + USE_PL: False +# SET GLOBAL SYSTEM +SOLVER: + # NAME DESCRIPTION: TYPE: default: 'TrainValSolver' + NAME: TrainValSolver + # RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: '' + RESUME_FROM: + # MAX_EPOCHS DESCRIPTION: Max epochs for training. TYPE: int default: 10 + MAX_EPOCHS: 200 + # NUM_FOLDS DESCRIPTION: Num folds for training. TYPE: int default: 0 + NUM_FOLDS: 1 + # WORK_DIR DESCRIPTION: Save dir of the training log or model. TYPE: str default: '' + WORK_DIR: ./exp12/ + LOG_FILE: std_log.txt + # EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1 + EVAL_INTERVAL: 1 + ACCU_STEP: 1 + # DO_FINAL_EVAL DESCRIPTION: If do final evaluation or not. TYPE: bool default: False + DO_FINAL_EVAL: True + # SAVE_EVAL_DATA DESCRIPTION: If save the evaluation data or not. TYPE: bool default: False + SAVE_EVAL_DATA: True + # EXTRA_KEYS DESCRIPTION: The extra keys for metric. TYPE: list default: [] + EXTRA_KEYS: [] + # TRAIN_DATA DESCRIPTION: Train data config. TYPE: default: '' + TRAIN_DATA: + # NAME DESCRIPTION: TYPE: default: 'ImageClassifyPublicDataset' + NAME: ImageClassifyPublicDataset + # DATASET DESCRIPTION: the public dataset name TYPE: str default: 'cifar10' + DATASET: cifar10 + # DATA_ROOT DESCRIPTION: the download data save path TYPE: str default: '' + DATA_ROOT: cifar10 + # MODE DESCRIPTION: test TYPE: str default: test + MODE: train + # PIN_MEMORY DESCRIPTION: pin_memory for data loader TYPE: bool default: False + PIN_MEMORY: True + # BATCH_SIZE DESCRIPTION: batch size for data TYPE: int default: 4 + BATCH_SIZE: 96 + # NUM_WORKERS DESCRIPTION: num workers for fetching data! TYPE: int default: 1 + NUM_WORKERS: 4 + # TRANSFORMS DESCRIPTION: TYPE: default: + TRANSFORMS: + # - DESCRIPTION: TYPE: default: + - # NAME DESCRIPTION: TYPE: default: 'RandomResizedCrop' + NAME: RandomResizedCrop + SIZE: 32 + # RATIO DESCRIPTION: ratio TYPE: list default: [0.75, 1.3333333333333333] + RATIO: [0.75, 1.33] + # SCALE DESCRIPTION: scale TYPE: list default: [0.08, 1.0] + SCALE: [0.8, 1.0] + # INTERPOLATION DESCRIPTION: interpolation TYPE: str default: 'blilinear' + INTERPOLATION: bilinear + # INPUT_KEY DESCRIPTION: input key TYPE: str default: 'img' + INPUT_KEY: img + # OUTPUT_KEY DESCRIPTION: input key TYPE: str default: 'img' + OUTPUT_KEY: img + # BACKEND DESCRIPTION: backend, choose from pillow, cv2, torchvision TYPE: str default: 'pillow' + BACKEND: pillow + - # NAME DESCRIPTION: TYPE: default: 'RandomHorizontalFlip' + NAME: RandomHorizontalFlip + # P DESCRIPTION: P TYPE: float default: 0.5 + P: 0.5 + # INPUT_KEY DESCRIPTION: input key TYPE: str default: 'img' + INPUT_KEY: img + # OUTPUT_KEY DESCRIPTION: input key TYPE: str default: 'img' + OUTPUT_KEY: img + # BACKEND DESCRIPTION: backend, choose from pillow, cv2, torchvision TYPE: str default: 'pillow' + BACKEND: pillow + - # NAME DESCRIPTION: TYPE: default: 'ImageToTensor' + NAME: ImageToTensor + # INPUT_KEY DESCRIPTION: input key TYPE: str default: 'img' + INPUT_KEY: img + # OUTPUT_KEY DESCRIPTION: input key TYPE: str default: 'img' + OUTPUT_KEY: img + # BACKEND DESCRIPTION: backend, choose from pillow, cv2, torchvision TYPE: str default: 'pillow' + BACKEND: pillow + - # NAME DESCRIPTION: TYPE: default: 'Normalize' + NAME: Normalize + # MEAN DESCRIPTION: mean TYPE: list default: [] + MEAN: [0.4914, 0.4822, 0.4465] + # STD DESCRIPTION: std TYPE: list default: [] + STD: [0.2023, 0.1994, 0.2010] + # INPUT_KEY DESCRIPTION: input key TYPE: str default: 'img' + INPUT_KEY: img + # OUTPUT_KEY DESCRIPTION: input key TYPE: str default: 'img' + OUTPUT_KEY: img + # BACKEND DESCRIPTION: backend, choose from pillow, cv2, torchvision TYPE: str default: 'pillow' + BACKEND: pillow + - NAME: ToTensor + # KEYS DESCRIPTION: keys TYPE: list default: [] + KEYS: ["img", "label"] + - # NAME DESCRIPTION: TYPE: default: 'Select' + NAME: Select + # KEYS DESCRIPTION: keys TYPE: list default: [] + KEYS: ["img", "label"] + # META_KEYS DESCRIPTION: meta keys TYPE: list default: [] + META_KEYS: [] + # EVAL_DATA DESCRIPTION: Eval data config. TYPE: default: '' + EVAL_DATA: + # NAME DESCRIPTION: TYPE: default: 'ImageClassifyPublicDataset' + NAME: ImageClassifyPublicDataset + # DATASET DESCRIPTION: the public dataset name TYPE: str default: 'cifar10' + DATASET: cifar10 + # DATA_ROOT DESCRIPTION: the download data save path TYPE: str default: '' + DATA_ROOT: ./local_data/cifar10 + # MODE DESCRIPTION: test TYPE: str default: test + MODE: test + # PIN_MEMORY DESCRIPTION: pin_memory for data loader TYPE: bool default: False + PIN_MEMORY: True + # BATCH_SIZE DESCRIPTION: batch size for data TYPE: int default: 4 + BATCH_SIZE: 96 + # NUM_WORKERS DESCRIPTION: num workers for fetching data! TYPE: int default: 1 + NUM_WORKERS: 4 + # TRANSFORMS DESCRIPTION: TYPE: default: + TRANSFORMS: + # - DESCRIPTION: TYPE: default: + - # NAME DESCRIPTION: TYPE: default: 'Resize' + NAME: Resize + SIZE: 32 + # INTERPOLATION DESCRIPTION: interpolation TYPE: str default: 'blilinear' + INTERPOLATION: bilinear + # INPUT_KEY DESCRIPTION: input key TYPE: str default: 'img' + INPUT_KEY: img + # OUTPUT_KEY DESCRIPTION: input key TYPE: str default: 'img' + OUTPUT_KEY: img + # BACKEND DESCRIPTION: backend, choose from pillow, cv2, torchvision TYPE: str default: 'pillow' + BACKEND: pillow + - # NAME DESCRIPTION: TYPE: default: 'ImageToTensor' + NAME: ImageToTensor + # INPUT_KEY DESCRIPTION: input key TYPE: str default: 'img' + INPUT_KEY: img + # OUTPUT_KEY DESCRIPTION: input key TYPE: str default: 'img' + OUTPUT_KEY: img + # BACKEND DESCRIPTION: backend, choose from pillow, cv2, torchvision TYPE: str default: 'pillow' + BACKEND: pillow + - # NAME DESCRIPTION: TYPE: default: 'Normalize' + NAME: Normalize + # MEAN DESCRIPTION: mean TYPE: list default: [] + MEAN: [0.4914, 0.4822, 0.4465] + # STD DESCRIPTION: std TYPE: list default: [] + STD: [0.2023, 0.1994, 0.2010] + # INPUT_KEY DESCRIPTION: input key TYPE: str default: 'img' + INPUT_KEY: img + # OUTPUT_KEY DESCRIPTION: input key TYPE: str default: 'img' + OUTPUT_KEY: img + # BACKEND DESCRIPTION: backend, choose from pillow, cv2, torchvision TYPE: str default: 'pillow' + BACKEND: pillow + - NAME: ToTensor + # KEYS DESCRIPTION: keys TYPE: list default: [] + KEYS: ["img", "label"] + - # NAME DESCRIPTION: TYPE: default: 'Select' + NAME: Select + # KEYS DESCRIPTION: keys TYPE: list default: [] + KEYS: ["img", "label"] + # META_KEYS DESCRIPTION: meta keys TYPE: list default: [] + META_KEYS: [] + # TRAIN_HOOKS DESCRIPTION: TYPE: default: '' + TRAIN_HOOKS: + - # NAME DESCRIPTION: TYPE: default: 'LogHook' + NAME: LogHook + # LOG_INTERVAL DESCRIPTION: the interval for log print! TYPE: int default: 10 + LOG_INTERVAL: 10 + # EVAL_HOOKS DESCRIPTION: TYPE: default: '' + EVAL_HOOKS: + - # NAME DESCRIPTION: TYPE: default: 'LogHook' + NAME: LogHook + # LOG_INTERVAL DESCRIPTION: the interval for log print! TYPE: int default: 10 + LOG_INTERVAL: 10 + # TEST_HOOKS DESCRIPTION: TYPE: default: '' + MODEL: + # NAME DESCRIPTION: TYPE: default: 'Classifier' + NAME: Classifier + # ACT_NAME DESCRIPTION: the activation function for logits, select from [softmax, sigmoid]! TYPE: str default: 'softmax' + ACT_NAME: softmax + # FREEZE_BN DESCRIPTION: if freeze bn of not TYPE: bool default: False + FREEZE_BN: False + # BACKBONE DESCRIPTION: TYPE: default: '' + BACKBONE: + # NAME DESCRIPTION: TYPE: default: 'ResNet' + NAME: ResNet + # DEPTH DESCRIPTION: the depth of network for resnet! TYPE: int default: 18 + DEPTH: 18 + # PRETRAINED DESCRIPTION: if load the official pretrained model or not. TYPE: bool default: False + PRETRAINED: false + # + KERNEL_SIZE: 3 + # USE_RELU DESCRIPTION: use relu or not! TYPE: bool default: True + USE_RELU: True + # USE_MAXPOOL DESCRIPTION: use maxpool or not! TYPE: bool default: True + USE_MAXPOOL: false + # FIRST_CONV_STRIDE DESCRIPTION: first conv stride 1 or 2! TYPE: int default: 1 + FIRST_CONV_STRIDE: 1 + # FIRST_MAX_POOL_STRIDE DESCRIPTION: first max pool stride 1 or 2! TYPE: int default: 1 + FIRST_MAX_POOL_STRIDE: 1 + # NECK DESCRIPTION: TYPE: default: '' + NECK: + # NAME DESCRIPTION: TYPE: default: 'GlobalAveragePooling' + NAME: GlobalAveragePooling + # DIM DESCRIPTION: GlobalAveragePooling dim! TYPE: int default: 2 + DIM: 2 + # HEAD DESCRIPTION: TYPE: default: '' + HEAD: + # NAME DESCRIPTION: TYPE: default: 'ClassifierHead' + NAME: ClassifierHead + # DIM DESCRIPTION: representation dim! TYPE: int default: 512 + DIM: 512 + # NUM_CLASSES DESCRIPTION: number of classes. TYPE: int default: 10 + NUM_CLASSES: 10 + # DROPOUT_RATE DESCRIPTION: dropout rate, default 0. TYPE: float default: 0.0 + DROPOUT_RATE: 0.0 + METRIC: + # NAME DESCRIPTION: TYPE: default: 'AccuracyMetric' + NAME: AccuracyMetric + # TOPK DESCRIPTION: topk accuracy! TYPE: int default: 1 + TOPK: 1 + # LOSS DESCRIPTION: TYPE: default: '' + LOSS: + # NAME DESCRIPTION: TYPE: default: 'CrossEntropy' + NAME: CrossEntropy + # REDUCE DESCRIPTION: reduce is False, returns a loss per batch element instead and ignores :attr: size_average. Default: True TYPE: NoneType default: None + # REDUCE: None + # SIZE_AVERAGE DESCRIPTION: Deprecated (see :attr: reduction). By default,the losses are averaged over each loss element in the batch. Note that forsome losses, there are multiple elements per sample. If the field :attr: size_averageis set to False, the losses are instead summed for each minibatch. Ignoredwhen :attr: reduce is False. Default: True TYPE: NoneType default: None + # SIZE_AVERAGE: None + # IGNORE_INDEX DESCRIPTION: Specifies a target value that is ignoredand does not contribute to the input gradient. When :attr: size_average isTrue, the loss is averaged over non-ignored targets. Note that:attr: ignore_index is only applicable when the target contains class indices. TYPE: int default: -100 + # IGNORE_INDEX: -100 + # REDUCTION DESCRIPTION: Specifies the reduction to apply to the output:'none' | 'mean' | 'sum'. 'none': no reduction willbe applied, 'mean': the weighted mean of the output is taken,'sum': the output will be summed. Note: :attr: size_averageand :attr:`reduce` are in the process of being deprecated, and inthe meantime, specifying either of those two args will override:attr:`reduction`. Default: 'mean' TYPE: str default: 'mean' + # REDUCTION: mean + # LABEL_SMOOTHING DESCRIPTION: A float in [0.0, 1.0]. Specifies the amountof smoothing when computing the loss, where 0.0 means no smoothing. TYPE: float default: 0.0 + # LABEL_SMOOTHING: 0.0 + # OPTIMIZER DESCRIPTION: TYPE: default: '' + OPTIMIZER: + # NAME DESCRIPTION: TYPE: default: 'SGD' + NAME: SGD + # LEARNING_RATE DESCRIPTION: the initial learning rate! TYPE: float default: 0.1 + LEARNING_RATE: 0.01 + # MOMENTUM DESCRIPTION: the momentum! TYPE: int default: 0 + MOMENTUM: 0.9 + # DAMPENING DESCRIPTION: the dampening! TYPE: int default: 0 + DAMPENING: 0 + # WEIGHT_DECAY DESCRIPTION: the weight decay! TYPE: int default: 0 + WEIGHT_DECAY: 5e-4 + # NESTEROV DESCRIPTION: the nesterov! TYPE: bool default: False + NESTEROV: False + # LR_SCHEDULER DESCRIPTION: TYPE: default: '' + LR_SCHEDULER: + # NAME DESCRIPTION: TYPE: default: 'CosineAnnealingLR' + NAME: CosineAnnealingLR + # T_MAX DESCRIPTION: the T max! TYPE: float default: 1.0 + T_MAX: 200.0 + # ETA_MIN DESCRIPTION: the eta min! TYPE: int default: 0 + ETA_MIN: 0 + # LAST_EPOCH DESCRIPTION: the last epoch! TYPE: int default: -1 + LAST_EPOCH: -1 + # METRICS DESCRIPTION: TYPE: default: '' + METRICS: + - # NAME DESCRIPTION: TYPE: default: 'AccuracyMetric' + NAME: AccuracyMetric + # TOPK DESCRIPTION: topk accuracy! TYPE: int default: 1 + TOPK: 1 + KEYS: ["logits", "label"] diff --git a/scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml b/scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml new file mode 100644 index 0000000..0d70b10 --- /dev/null +++ b/scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml @@ -0,0 +1,218 @@ +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/sd15_512_full + LOG_FILE: std_log.txt + # + FILE_SYSTEM: + NAME: "ModelscopeFs" + TEMP_DIR: "./cache/data" + # + MODEL: + NAME: LatentDiffusion + PARAMETERIZATION: eps + TIMESTEPS: 1000 + MIN_SNR_GAMMA: + ZERO_TERMINAL_SNR: False + PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-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 + # + LOSS: + NAME: ReconstructLoss + LOSS_TYPE: l2 + # + SAMPLE_ARGS: + SAMPLER: ddim + SAMPLE_STEPS: 50 + SEED: 2023 + GUIDE_SCALE: 7.5 + GUIDE_RESCALE: + DISCRETIZATION: trailing + IMAGE_SIZE: [512, 512] + RUN_TRAIN_N: False + # + OPTIMIZER: + NAME: AdamW + LEARNING_RATE: 0.0064 + BETAS: [ 0.9, 0.999 ] + EPS: 1e-8 + WEIGHT_DECAY: 1e-2 + AMSGRAD: False + # + TRAIN_DATA: + NAME: ImageTextPairMSDataset + MODE: train + MS_DATASET_NAME: style_custom_dataset + MS_DATASET_NAMESPACE: damo + MS_DATASET_SUBNAME: 3D + PROMPT_PREFIX: "" + MS_DATASET_SPLIT: train_short + MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' } + REPLACE_STYLE: False + PIN_MEMORY: True + BATCH_SIZE: 1 + NUM_WORKERS: 4 + SAMPLER: + NAME: LoopSampler + TRANSFORMS: + - NAME: LoadImageFromFile + RGB_ORDER: RGB + BACKEND: pillow + - NAME: Resize + SIZE: 512 + INTERPOLATION: bilinear + INPUT_KEY: [ 'img' ] + OUTPUT_KEY: [ 'img' ] + BACKEND: pillow + - NAME: CenterCrop + SIZE: 512 + INPUT_KEY: [ 'img' ] + OUTPUT_KEY: [ 'img' ] + BACKEND: pillow + - NAME: ImageToTensor + INPUT_KEY: [ 'img' ] + OUTPUT_KEY: [ 'img' ] + BACKEND: pillow + - NAME: Normalize + MEAN: [ 0.5, 0.5, 0.5 ] + STD: [ 0.5, 0.5, 0.5 ] + INPUT_KEY: [ 'img' ] + OUTPUT_KEY: [ 'image' ] + BACKEND: torchvision + - NAME: Select + KEYS: [ 'image', 'prompt' ] + META_KEYS: [ 'data_key' ] + # + EVAL_DATA: + NAME: ImageTextPairMSDataset + MODE: eval + MS_DATASET_NAME: style_custom_dataset + MS_DATASET_NAMESPACE: damo + MS_DATASET_SUBNAME: 3D + PROMPT_PREFIX: "" + MS_REMAP_KEYS: { 'Image': 'Target:FILE' } + MS_DATASET_SPLIT: test_short + OUTPUT_SIZE: [512, 512] + REPLACE_STYLE: False + PIN_MEMORY: True + BATCH_SIZE: 4 + NUM_WORKERS: 4 + FILE_SYSTEM: + NAME: "ModelscopeFs" + TEMP_DIR: "./cache/data" + # + TRANSFORMS: + - + NAME: Select + KEYS: ['prompt'] + META_KEYS: ['image_size'] + # + TRAIN_HOOKS: + - + NAME: BackwardHook + PRIORITY: 0 + - + NAME: LogHook + LOG_INTERVAL: 50 + - + NAME: CheckpointHook + INTERVAL: 1000 + - + NAME: ProbeDataHook + PROB_INTERVAL: 100 + # + EVAL_HOOKS: + - + NAME: ProbeDataHook + PROB_INTERVAL: 100 diff --git a/scepter/methods/examples/generation/stable_diffusion_1.5_512_lora.yaml b/scepter/methods/examples/generation/stable_diffusion_1.5_512_lora.yaml new file mode 100644 index 0000000..57b23f8 --- /dev/null +++ b/scepter/methods/examples/generation/stable_diffusion_1.5_512_lora.yaml @@ -0,0 +1,226 @@ +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/sd15_512_lora + LOG_FILE: std_log.txt + # + FILE_SYSTEM: + NAME: "ModelscopeFs" + TEMP_DIR: "./cache/data" + TUNER: + - + NAME: SwiftLoRA + R: 64 + LORA_ALPHA: 64 + LORA_DROPOUT: 0.0 + BIAS: "none" + TARGET_MODULES: model.*(to_q|to_k|to_v|to_out.0|net.0.proj|net.2)$ + # + MODEL: + NAME: LatentDiffusion + PARAMETERIZATION: eps + TIMESTEPS: 1000 + MIN_SNR_GAMMA: + ZERO_TERMINAL_SNR: False + PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-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 + # + LOSS: + NAME: ReconstructLoss + LOSS_TYPE: l2 + # + SAMPLE_ARGS: + SAMPLER: ddim + SAMPLE_STEPS: 50 + SEED: 2023 + GUIDE_SCALE: 7.5 + GUIDE_RESCALE: + DISCRETIZATION: trailing + IMAGE_SIZE: [512, 512] + RUN_TRAIN_N: False + # + OPTIMIZER: + NAME: AdamW + LEARNING_RATE: 0.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 diff --git a/scepter/methods/examples/generation/stable_diffusion_2.1_768.yaml b/scepter/methods/examples/generation/stable_diffusion_2.1_768.yaml new file mode 100644 index 0000000..ad0fc9f --- /dev/null +++ b/scepter/methods/examples/generation/stable_diffusion_2.1_768.yaml @@ -0,0 +1,214 @@ +ENV: + BACKEND: nccl +SOLVER: + NAME: LatentDiffusionSolver + RESUME_FROM: + LOAD_MODEL_ONLY: True + USE_FSDP: False + SHARDING_STRATEGY: + USE_AMP: True + DTYPE: float16 + CHANNELS_LAST: True + MAX_STEPS: 2000 + MAX_EPOCHS: -1 + NUM_FOLDS: 1 + ACCU_STEP: 1 + EVAL_INTERVAL: 100 + # + WORK_DIR: ./cache/sd21_768_full + LOG_FILE: std_log.txt + # + FILE_SYSTEM: + NAME: "ModelscopeFs" + TEMP_DIR: "./cache/data" + # + MODEL: + NAME: LatentDiffusion + PARAMETERIZATION: v + TIMESTEPS: 1000 + MIN_SNR_GAMMA: + ZERO_TERMINAL_SNR: False + PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-2-1@v2-1_768-ema-pruned.safetensors + IGNORE_KEYS: [ ] + SCALE_FACTOR: 0.18215 + SIZE_FACTOR: 8 + DEFAULT_N_PROMPT: + SCHEDULE_ARGS: + "NAME": "scaled_linear" + "BETA_MIN": 0.00085 + "BETA_MAX": 0.012 + USE_EMA: False + # + DIFFUSION_MODEL: + NAME: DiffusionUNet + IN_CHANNELS: 4 + OUT_CHANNELS: 4 + MODEL_CHANNELS: 320 + NUM_HEADS_CHANNELS: 64 + NUM_RES_BLOCKS: 2 + ATTENTION_RESOLUTIONS: [ 4, 2, 1 ] + CHANNEL_MULT: [ 1, 2, 4, 4 ] + CONV_RESAMPLE: True + DIMS: 2 + USE_CHECKPOINT: False + USE_SCALE_SHIFT_NORM: False + RESBLOCK_UPDOWN: False + USE_SPATIAL_TRANSFORMER: True + TRANSFORMER_DEPTH: 1 + CONTEXT_DIM: 1024 + DISABLE_MIDDLE_SELF_ATTN: False + USE_LINEAR_IN_TRANSFORMER: True + PRETRAINED_MODEL: + # + FIRST_STAGE_MODEL: + NAME: AutoencoderKL + EMBED_DIM: 4 + PRETRAINED_MODEL: + IGNORE_KEYS: [ ] + BATCH_SIZE: 4 + # + ENCODER: + NAME: Encoder + CH: 128 + OUT_CH: 3 + NUM_RES_BLOCKS: 2 + IN_CHANNELS: 3 + ATTN_RESOLUTIONS: [ ] + CH_MULT: [ 1, 2, 4, 4 ] + Z_CHANNELS: 4 + DOUBLE_Z: True + DROPOUT: 0.0 + RESAMP_WITH_CONV: True + # + DECODER: + NAME: Decoder + CH: 128 + OUT_CH: 3 + NUM_RES_BLOCKS: 2 + IN_CHANNELS: 3 + ATTN_RESOLUTIONS: [ ] + CH_MULT: [ 1, 2, 4, 4 ] + Z_CHANNELS: 4 + DROPOUT: 0.0 + RESAMP_WITH_CONV: True + GIVE_PRE_END: False + TANH_OUT: False + # + TOKENIZER: + NAME: OpenClipTokenizer + LENGTH: 77 + # + COND_STAGE_MODEL: + NAME: FrozenOpenCLIPEmbedder + ARCH: ViT-H-14 + PRETRAINED_MODEL: + LAYER: penultimate + # + LOSS: + NAME: ReconstructLoss + LOSS_TYPE: l2 + # + SAMPLE_ARGS: + SAMPLER: ddim + SAMPLE_STEPS: 50 + SEED: 2023 + GUIDE_SCALE: 7.5 + GUIDE_RESCALE: + DISCRETIZATION: trailing + IMAGE_SIZE: [768, 768] + RUN_TRAIN_N: False + # + OPTIMIZER: + NAME: AdamW + LEARNING_RATE: 0.0064 + BETAS: [ 0.9, 0.999 ] + EPS: 1e-8 + WEIGHT_DECAY: 1e-2 + AMSGRAD: False + # + TRAIN_DATA: + NAME: ImageTextPairMSDataset + MODE: train + MS_DATASET_NAME: style_custom_dataset + MS_DATASET_NAMESPACE: damo + MS_DATASET_SUBNAME: 3D + PROMPT_PREFIX: "" + MS_DATASET_SPLIT: train_short + MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' } + REPLACE_STYLE: False + PIN_MEMORY: True + BATCH_SIZE: 1 + NUM_WORKERS: 4 + SAMPLER: + NAME: LoopSampler + TRANSFORMS: + - NAME: LoadImageFromFile + RGB_ORDER: RGB + BACKEND: pillow + - NAME: Resize + SIZE: 768 + INTERPOLATION: bilinear + INPUT_KEY: [ 'img' ] + OUTPUT_KEY: [ 'img' ] + BACKEND: pillow + - NAME: CenterCrop + SIZE: 768 + 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: [768, 768] + 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 diff --git a/scepter/methods/examples/generation/stable_diffusion_2.1_768_lora.yaml b/scepter/methods/examples/generation/stable_diffusion_2.1_768_lora.yaml new file mode 100644 index 0000000..4dd1cb5 --- /dev/null +++ b/scepter/methods/examples/generation/stable_diffusion_2.1_768_lora.yaml @@ -0,0 +1,223 @@ +ENV: + BACKEND: nccl +SOLVER: + NAME: LatentDiffusionSolver + RESUME_FROM: + LOAD_MODEL_ONLY: True + USE_FSDP: False + SHARDING_STRATEGY: + USE_AMP: True + DTYPE: float16 + CHANNELS_LAST: True + MAX_STEPS: 2000 + MAX_EPOCHS: -1 + NUM_FOLDS: 1 + ACCU_STEP: 1 + EVAL_INTERVAL: 100 + # + WORK_DIR: ./cache/sd21_768_lora + LOG_FILE: std_log.txt + # + FILE_SYSTEM: + NAME: "ModelscopeFs" + TEMP_DIR: "./cache/data" + # + TUNER: + - + NAME: SwiftLoRA + R: 64 + LORA_ALPHA: 64 + LORA_DROPOUT: 0.0 + BIAS: "none" + TARGET_MODULES: model.*(to_q|to_k|to_v|to_out.0|net.0.proj|net.2)$ + # + MODEL: + NAME: LatentDiffusion + PARAMETERIZATION: v + TIMESTEPS: 1000 + MIN_SNR_GAMMA: + ZERO_TERMINAL_SNR: False + PRETRAINED_MODEL: ms://AI-ModelScope/stable-diffusion-2-1@v2-1_768-ema-pruned.safetensors + IGNORE_KEYS: [ ] + SCALE_FACTOR: 0.18215 + SIZE_FACTOR: 8 + DEFAULT_N_PROMPT: + SCHEDULE_ARGS: + "NAME": "scaled_linear" + "BETA_MIN": 0.00085 + "BETA_MAX": 0.012 + USE_EMA: False + # + DIFFUSION_MODEL: + NAME: DiffusionUNet + IN_CHANNELS: 4 + OUT_CHANNELS: 4 + MODEL_CHANNELS: 320 + NUM_HEADS_CHANNELS: 64 + NUM_RES_BLOCKS: 2 + ATTENTION_RESOLUTIONS: [ 4, 2, 1 ] + CHANNEL_MULT: [ 1, 2, 4, 4 ] + CONV_RESAMPLE: True + DIMS: 2 + USE_CHECKPOINT: False + USE_SCALE_SHIFT_NORM: False + RESBLOCK_UPDOWN: False + USE_SPATIAL_TRANSFORMER: True + TRANSFORMER_DEPTH: 1 + CONTEXT_DIM: 1024 + DISABLE_MIDDLE_SELF_ATTN: False + USE_LINEAR_IN_TRANSFORMER: True + PRETRAINED_MODEL: + # + FIRST_STAGE_MODEL: + NAME: AutoencoderKL + EMBED_DIM: 4 + PRETRAINED_MODEL: + IGNORE_KEYS: [ ] + BATCH_SIZE: 4 + # + ENCODER: + NAME: Encoder + CH: 128 + OUT_CH: 3 + NUM_RES_BLOCKS: 2 + IN_CHANNELS: 3 + ATTN_RESOLUTIONS: [ ] + CH_MULT: [ 1, 2, 4, 4 ] + Z_CHANNELS: 4 + DOUBLE_Z: True + DROPOUT: 0.0 + RESAMP_WITH_CONV: True + # + DECODER: + NAME: Decoder + CH: 128 + OUT_CH: 3 + NUM_RES_BLOCKS: 2 + IN_CHANNELS: 3 + ATTN_RESOLUTIONS: [ ] + CH_MULT: [ 1, 2, 4, 4 ] + Z_CHANNELS: 4 + DROPOUT: 0.0 + RESAMP_WITH_CONV: True + GIVE_PRE_END: False + TANH_OUT: False + # + TOKENIZER: + NAME: OpenClipTokenizer + LENGTH: 77 + # + COND_STAGE_MODEL: + NAME: FrozenOpenCLIPEmbedder + ARCH: ViT-H-14 + PRETRAINED_MODEL: + LAYER: penultimate + # + LOSS: + NAME: ReconstructLoss + LOSS_TYPE: l2 + # + SAMPLE_ARGS: + SAMPLER: ddim + SAMPLE_STEPS: 50 + SEED: 2023 + GUIDE_SCALE: 7.5 + GUIDE_RESCALE: + DISCRETIZATION: trailing + IMAGE_SIZE: [768, 768] + 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: 768 + INTERPOLATION: bilinear + INPUT_KEY: [ 'img' ] + OUTPUT_KEY: [ 'img' ] + BACKEND: pillow + - NAME: CenterCrop + SIZE: 768 + 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: [768, 768] + 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 diff --git a/scepter/methods/examples/generation/stable_diffusion_xl_1024.yaml b/scepter/methods/examples/generation/stable_diffusion_xl_1024.yaml new file mode 100644 index 0000000..6ac3f4d --- /dev/null +++ b/scepter/methods/examples/generation/stable_diffusion_xl_1024.yaml @@ -0,0 +1,328 @@ +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/sdxl_1024_full + LOG_FILE: std_log.txt + # + FILE_SYSTEM: + NAME: "ModelscopeFs" + TEMP_DIR: "./cache/data" + # + 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: + 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.0064 + BETAS: [ 0.9, 0.999 ] + EPS: 1e-8 + WEIGHT_DECAY: 1e-2 + AMSGRAD: False + # + TRAIN_DATA: + NAME: ImageTextPairMSDataset + MODE: train + MS_DATASET_NAME: style_custom_dataset + MS_DATASET_NAMESPACE: damo + MS_DATASET_SUBNAME: 3D + PROMPT_PREFIX: "" + MS_DATASET_SPLIT: train_short + MS_REMAP_KEYS: { 'Image:FILE': 'Target:FILE' } + REPLACE_STYLE: False + PIN_MEMORY: True + BATCH_SIZE: 1 + NUM_WORKERS: 4 + SAMPLER: + NAME: LoopSampler + TRANSFORMS: + - NAME: LoadImageFromFile + RGB_ORDER: RGB + BACKEND: pillow + - NAME: 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 + IN_KEYS: [ 'img', 'img_original_size_as_tuple', 'img_target_size_as_tuple', 'img_crop_coords_top_left' ] + OUT_KEYS: [ '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: 2 + 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 diff --git a/scepter/methods/examples/generation/stable_diffusion_xl_1024_lora.yaml b/scepter/methods/examples/generation/stable_diffusion_xl_1024_lora.yaml new file mode 100644 index 0000000..2f1d315 --- /dev/null +++ b/scepter/methods/examples/generation/stable_diffusion_xl_1024_lora.yaml @@ -0,0 +1,337 @@ +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/sdxl_1024_lora + LOG_FILE: std_log.txt + # + FILE_SYSTEM: + NAME: "ModelscopeFs" + TEMP_DIR: "./cache/data" + # + TUNER: + - + NAME: SwiftLoRA + R: 64 + LORA_ALPHA: 64 + LORA_DROPOUT: 0.0 + BIAS: "none" + TARGET_MODULES: model.*(to_q|to_k|to_v|to_out.0|net.0.proj|net.2)$ + # + MODEL: + NAME: 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: + 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 + IN_KEYS: [ 'img', 'img_original_size_as_tuple', 'img_target_size_as_tuple', 'img_crop_coords_top_left' ] + OUT_KEYS: [ '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 \ No newline at end of file diff --git a/scepter/modules/__init__.py b/scepter/modules/__init__.py new file mode 100644 index 0000000..43a95d8 --- /dev/null +++ b/scepter/modules/__init__.py @@ -0,0 +1,3 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules import data, model, opt, solver, transform, utils diff --git a/scepter/modules/data/__init__.py b/scepter/modules/data/__init__.py new file mode 100644 index 0000000..198e24f --- /dev/null +++ b/scepter/modules/data/__init__.py @@ -0,0 +1,4 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.data import dataset, sampler diff --git a/scepter/modules/data/dataset/__init__.py b/scepter/modules/data/dataset/__init__.py new file mode 100644 index 0000000..f513bb1 --- /dev/null +++ b/scepter/modules/data/dataset/__init__.py @@ -0,0 +1,10 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.data.dataset.base_dataset import BaseDataset +from scepter.modules.data.dataset.dataset import (Image2ImageDataset, + ImageClassifyPublicDataset, + ImageTextPairDataset, + Text2ImageDataset) +from scepter.modules.data.dataset.ms_dataset import ImageTextPairMSDataset +from scepter.modules.data.dataset.registry import DATASETS diff --git a/scepter/modules/data/dataset/base_dataset.py b/scepter/modules/data/dataset/base_dataset.py new file mode 100644 index 0000000..0211d6d --- /dev/null +++ b/scepter/modules/data/dataset/base_dataset.py @@ -0,0 +1,113 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import os +from abc import ABCMeta, abstractmethod + +from torch.utils.data import Dataset + +from scepter.modules.transform.registry import TRANSFORMS, build_pipeline +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import set_random_seed, we +from scepter.modules.utils.file_system import FS +from scepter.modules.utils.logger import get_logger +from scepter.modules.utils.registry import old_python_version + + +class BaseDataset(Dataset, metaclass=ABCMeta): + para_dict = { + 'MODE': { + 'value': 'train', + 'description': 'solver phase, select from [train, test, eval]' + }, + 'FILE_SYSTEM': {}, + 'TRANSFORMS': [{ + 'ImageToTensor': { + 'PADDING': { + 'value': None, + 'description': 'padding' + }, + 'PAD_IF_NEEDED': { + 'value': False, + 'description': 'pad if needed' + }, + 'FILL': { + 'value': 0, + 'description': 'fill' + }, + 'PADDING_MODE': { + 'value': 'constant', + 'description': 'padding mode' + } + } + }] + } + + def __init__(self, cfg, logger=None): + mode = cfg.get('MODE', 'train') + pipeline = cfg.get('TRANSFORMS', []) + super(BaseDataset, self).__init__() + self.mode = mode + self.logger = logger + self.worker_logger = get_logger(name='datasets') + self.pipeline = build_pipeline(pipeline, + TRANSFORMS, + logger=self.worker_logger) + self.file_systems = cfg.get('FILE_SYSTEM', None) + + if isinstance(self.file_systems, list): + for file_sys in self.file_systems: + self.fs_prefix = FS.init_fs_client(file_sys, + logger=self.logger, + overwrite=False) + elif self.file_systems is not None: + self.fs_prefix = FS.init_fs_client(self.file_systems, + logger=self.logger, + overwrite=False) + self.local_we = we.get_env() + if old_python_version: + self.file_systems.logger = None + + def __getitem__(self, index: int): + item = self._get(index) + return self.pipeline(item) + + def worker_init_fn(self, worker_id, num_workers=1): + if isinstance(self.file_systems, list): + for file_sys in self.file_systems: + self.fs_prefix = FS.init_fs_client(file_sys, + logger=self.logger, + overwrite=False) + elif self.file_systems is not None: + self.fs_prefix = FS.init_fs_client(self.file_systems, + logger=self.logger, + overwrite=False) + self.worker_id = worker_id + self.logger = self.worker_logger + set_random_seed(int(os.environ.get('ES_SEED', 2023))) + we.set_env(self.local_we) + + @abstractmethod + def _get(self, index: int): + pass + + def __repr__(self) -> str: + return f'{self.__class__.__name__}: mode={self.mode}, len={len(self)}' + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('DATASETS', + __class__.__name__, + BaseDataset.para_dict, + set_name=True) diff --git a/scepter/modules/data/dataset/dataset.py b/scepter/modules/data/dataset/dataset.py new file mode 100644 index 0000000..c82c50b --- /dev/null +++ b/scepter/modules/data/dataset/dataset.py @@ -0,0 +1,284 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import numbers +import sys +from collections.abc import Iterable + +import numpy as np +import torchvision + +from scepter.modules.data.dataset.base_dataset import BaseDataset +from scepter.modules.data.dataset.registry import DATASETS +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we +from scepter.modules.utils.file_system import DATA_FS as FS + + +@DATASETS.register_class() +class ImageClassifyPublicDataset(BaseDataset): + """ + Dataset for image classification wrapper + + Args: + json_path (str): json file which contains all instances, should be a list of dict + which contains img_path and gt_label + image_dir (str or None): image directory, if None, img_path in json_path will be considered as absolute path + classes (list[str] or None): image class description + """ + para_dict = { + 'DATASET': { + 'value': 'cifar10', + 'description': 'the public dataset name' + }, + 'DATA_ROOT': { + 'value': '', + 'description': 'the download data save path' + } + } + + para_dict.update(BaseDataset.para_dict) + + def __init__(self, cfg, logger=None): + + super(ImageClassifyPublicDataset, self).__init__(cfg, logger=logger) + + self.dataset_name = cfg.DATASET + self.data_root = cfg.DATA_ROOT + self.phase = cfg.MODE + if self.dataset_name == 'cifar10': + self.dataset = torchvision.datasets.CIFAR10( + root=self.data_root, + train=self.phase == 'train', + download=True) + + def __len__(self) -> int: + return len(self.dataset) + + def _get(self, index: int): + img, target = self.dataset.__getitem__(index) + ret = { + 'meta': {}, + 'label': np.asarray(target, dtype=np.int64), + 'img': img + } + return ret + + def worker_init_fn(self, worker_id, num_workers=1): + super(ImageClassifyPublicDataset, + self).worker_init_fn(worker_id, num_workers=num_workers) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('modename_DATA', + __class__.__name__, + ImageClassifyPublicDataset.para_dict, + set_name=True) + + +@DATASETS.register_class() +class ImageTextPairDataset(BaseDataset): + """ + Dataset for diffusion model training + """ + para_dict = { + 'P_ZERO': { + 'value': 0.0, + 'description': '', + }, + 'NEGTIVE_PROMPT': { + 'value': '', + 'description': 'The default negtive prompt', + } + } + para_dict.update(BaseDataset.para_dict) + + def __init__(self, cfg, logger=None): + super(ImageTextPairDataset, self).__init__(cfg, logger=logger) + self.p_zero = cfg.get('P_ZERO', 0.0) + self._default_item = { + 'meta': {}, + 'prompt': + 'Plants in the Water, Nature, Lake, Horizontal, Reflection, Photography, Backgrounds, Swamp, No People' + } + + def _get(self, index): + meta = dict() + # the last item is field_keys + for key, value in zip(index[-1], index[:-1]): + if key in ['oss_key', 'path', 'img_path', 'target_img_path']: + meta['img_path'] = value + elif key in ['prompt', 'caption', 'text']: + meta['ori_prompt'] = value + elif key in ['width', 'height']: + meta[key] = int(value) + else: + meta[key] = value + + prompt = meta.get('prompt_prefix', '') + meta.get('ori_prompt', '') + if self.mode == 'train' and np.random.uniform() < self.p_zero: + prompt = '' + + item = { + 'meta': meta, + 'prompt': prompt, + } + return item + + def __getitem__(self, index): + item = self._get(index) + item = self.pipeline(item) + return item + + def __len__(self) -> int: + return sys.maxsize + + @staticmethod + def get_config_template(): + return dict_to_yaml('DATASETS', + __class__.__name__, + ImageTextPairDataset.para_dict, + set_name=True) + + +@DATASETS.register_class() +class Image2ImageDataset(BaseDataset): + """ + Dataset for diffusion model training + """ + para_dict = {} + para_dict.update(BaseDataset.para_dict) + + def __init__(self, cfg, logger=None): + super(Image2ImageDataset, self).__init__(cfg, logger=logger) + self._default_item = { + 'meta': {}, + 'prompt': + 'Plants in the Water, Nature, Lake, Horizontal, Reflection, Photography, Backgrounds, Swamp, No People' + } + + def _get(self, index): + meta = dict() + # the last item is field_keys + for key, value in zip(index[-1], index[:-1]): + if key in ['oss_key', 'path', 'img_path']: + meta['img_path'] = value + elif key in ['prompt', 'caption', 'text']: + meta['ori_prompt'] = value + elif key in ['width', 'height']: + meta[key] = int(value) + else: + meta[key] = value + item = {'meta': meta} + return item + + def __getitem__(self, index): + item = self._get(index) + item = self.pipeline(item) + return item + + def __len__(self) -> int: + return sys.maxsize + + @staticmethod + def get_config_template(): + return dict_to_yaml('DATASETS', + __class__.__name__, + Image2ImageDataset.para_dict, + set_name=True) + + +@DATASETS.register_class() +class Text2ImageDataset(BaseDataset): + para_dict = { + 'PROMPT_FILE': { + 'value': '', + 'description': '' + }, + 'FIELDS': { + 'value': '', + 'description': '' + }, + 'DELIMITER': { + 'value': ',', + 'description': '' + }, + 'PROMPT_PREFIX': { + 'value': '', + 'description': '' + }, + 'IMAGE_SIZE': { + 'value': 512, + 'description': '' + }, + 'USE_NUM': { + 'value': -1, + 'description': '' + }, + } + para_dict.update(BaseDataset.para_dict) + + def __init__(self, cfg, logger=None): + super(Text2ImageDataset, self).__init__(cfg, logger=logger) + + delimiter = cfg.get('DELIMITER', ',') + fields = cfg.get('FIELDS', ['row_key', 'prompt']) + prompt_prefix = cfg.get('PROMPT_PREFIX', '') + use_num = cfg.get('USE_NUM', -1) + + image_size = cfg.get('IMAGE_SIZE', 1024) + if isinstance(image_size, numbers.Number): + image_size = [image_size, image_size] + assert isinstance(image_size, Iterable) and len(image_size) == 2 + + prompt_file = cfg.PROMPT_FILE + with FS.get_object(prompt_file) as local_data: + rows = [ + i.split(delimiter, + len(fields) - 1) + for i in local_data.decode('utf-8').strip().split('\n') + ] + + self.items = list() + for i, row in enumerate(rows): + item = {'index': i, 'meta': {'image_size': image_size}} + for key, value in zip(fields, row): + if key in ['prompt', 'caption', 'text']: + item['ori_prompt'] = value + item['prompt'] = prompt_prefix + value + elif key != 'meta': + item[key] = value + else: + continue + + self.items.append(item) + if use_num > 0: + self.items = self.items[:use_num] + if we.rank == 0: + logger.info(f'eval prompt num: {len(self.items)}') + logger.info('eval prompts: {}'.format( + [k['prompt'] for k in self.items])) + + def _get(self, index: int): + return self.items[index] + + def __len__(self) -> int: + return len(self.items) + + @staticmethod + def get_config_template(): + return dict_to_yaml('DATASETS', + __class__.__name__, + Text2ImageDataset.para_dict, + set_name=True) diff --git a/scepter/modules/data/dataset/ms_dataset.py b/scepter/modules/data/dataset/ms_dataset.py new file mode 100644 index 0000000..190135c --- /dev/null +++ b/scepter/modules/data/dataset/ms_dataset.py @@ -0,0 +1,179 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import numbers +import os +import sys + +from scepter.modules.data.dataset.base_dataset import BaseDataset +from scepter.modules.data.dataset.registry import DATASETS +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we + + +@DATASETS.register_class() +class ImageTextPairMSDataset(BaseDataset): + para_dict = { + 'MS_DATASET_NAME': { + 'value': '', + 'description': 'Modelscope dataset name.' + }, + 'MS_DATASET_NAMESPACE': { + 'value': '', + 'description': 'Modelscope dataset namespace.' + }, + 'MS_DATASET_SUBNAME': { + 'value': '', + 'description': 'Modelscope dataset subname.' + }, + 'MS_DATASET_SPLIT': { + 'value': '', + 'description': + 'Modelscope dataset split set name, default is train.' + }, + 'MS_REMAP_KEYS': { + 'value': + None, + 'description': + 'Modelscope dataset header of list file, the default is Target:FILE; ' + 'If your file is not this header, please set this field, which is a map dict.' + "For example, { 'Image:FILE': 'Target:FILE' } will replace the filed Image:FILE to Target:FILE" + }, + 'MS_REMAP_PATH': { + 'value': + None, + 'description': + 'When modelscope dataset name is not None, that means you use the dataset from modelscope,' + ' default is None. But if you want to use the datalist from modelscope and the file from ' + 'local device, you can use this field to set the root path of your images. ' + }, + 'TRIGGER_WORDS': { + 'value': + '', + 'description': + 'The words used to describe the common features of your data, especially when you customize a ' + 'tuner. Use these words you can get what you want.' + }, + 'REPLACE_STYLE': { + 'value': + False, + 'description': + 'Whether use the MS_DATASET_SUBNAME to replace the word in your description, default is False.' + }, + 'HIGHLIGHT_KEYWORDS': { + 'value': + '', + 'description': + 'The keywords you want to highlight in prompt, which will be replace by .' + }, + 'KEYWORDS_SIGN': { + 'value': + '', + 'description': + 'The keywords sign you want to add, which is like <{HIGHLIGHT_KEYWORDS}{KEYWORDS_SIGN}>' + }, + 'OUTPUT_SIZE': { + 'value': + None, + 'description': + 'If you use the FlexibleResize transforms, this filed will output the image_size as [h, w],' + 'which will be used to set the output size of images used to train the model.' + }, + } + + def __init__(self, cfg, logger=None): + super().__init__(cfg=cfg, logger=logger) + from modelscope import MsDataset + from modelscope.utils.constant import DownloadMode + ms_dataset_name = cfg.get('MS_DATASET_NAME', None) + ms_dataset_namespace = cfg.get('MS_DATASET_NAMESPACE', None) + ms_dataset_subname = cfg.get('MS_DATASET_SUBNAME', None) + ms_dataset_split = cfg.get('MS_DATASET_SPLIT', 'train') + ms_remap_keys = cfg.get('MS_REMAP_KEYS', None) + ms_remap_path = cfg.get('MS_REMAP_PATH', None) + self.replace_style = cfg.get('REPLACE_STYLE', False) + self.trigger_words = cfg.get('TRIGGER_WORDS', '') + self.replace_keywords = cfg.get('HIGHLIGHT_KEYWORDS', '') + self.keywords_sign = cfg.get('KEYWORDS_SIGN', '') + self.output_size = cfg.get('OUTPUT_SIZE', None) + if self.output_size is not None: + if isinstance(self.output_size, numbers.Number): + self.output_size = [self.output_size, self.output_size] + # Use modelscope dataset + + if not ms_dataset_name: + raise ( + 'Your must set MS_DATASET_NAME as modelscope dataset or your local dataset orignized ' + 'as modelscope dataset.') + try: + self.data = MsDataset.load(str(ms_dataset_name), + namespace=ms_dataset_namespace, + subset_name=ms_dataset_subname, + split=ms_dataset_split) + except Exception as e: + self.logger.info( + f"Load Modelscope dataset failed with {e}, retry with download_mode='force_redownload'." + ) + try: + self.data = MsDataset.load( + str(ms_dataset_name), + namespace=ms_dataset_namespace, + subset_name=ms_dataset_subname, + split=ms_dataset_split, + download_mode=DownloadMode.FORCE_REDOWNLOAD) + except Exception as sec_e: + raise f'Load Modelscope dataset failed {sec_e}.' + if ms_remap_keys: + self.data = self.data.remap_columns(ms_remap_keys.get_dict()) + + if ms_remap_path: + + def map_func(example): + example['Target:FILE'] = os.path.join(ms_remap_path, + example['Target:FILE']) + return example + + self.data = self.data.ds_instance.map(map_func) + self.real_number = len(self.data) + + def __len__(self): + if self.mode == 'train': + return sys.maxsize + else: + return len(self.data) + + def _get(self, index: int): + current_data = self.data[index % len(self.data)] + # print(current_data.keys()) + image_path = current_data['Target:FILE'] + prompt = current_data['Prompt'] + style = current_data['Style'] if 'Style' in current_data else '' + # print(prompt, style) + if self.replace_style and not style == '': + prompt = prompt.replace(style, f'<{self.keywords_sign}>') + elif not self.replace_keywords.strip() == '': + prompt = prompt.replace( + self.replace_keywords, + '<' + self.replace_keywords + f'{self.keywords_sign}>') + if not self.trigger_words == '': + prompt = self.trigger_words.strip() + ' ' + prompt + if we.debug: + print(prompt, self.replace_keywords.strip()) + ret_item = { + 'meta': { + 'img_path': image_path, + 'data_key': style, + 'data_num': self.real_number + }, + 'prompt': prompt + } + if self.output_size is not None: + ret_item['meta']['image_size'] = self.output_size + return ret_item + + @staticmethod + def get_config_template(): + return dict_to_yaml('DATASet', + __class__.__name__, + ImageTextPairMSDataset.para_dict, + set_name=True) diff --git a/scepter/modules/data/dataset/registry.py b/scepter/modules/data/dataset/registry.py new file mode 100644 index 0000000..a1e935e --- /dev/null +++ b/scepter/modules/data/dataset/registry.py @@ -0,0 +1,361 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import collections.abc as container_abcs +import inspect +import re +from functools import partial + +import torch +from torch.utils.data import DataLoader, DistributedSampler + +from scepter.modules.data.sampler import (SAMPLERS, MixtureOfSamplers, + MultiFoldDistributedSampler, + MultiLevelBatchSampler) +from scepter.modules.utils.registry import (Registry, deep_copy, + old_python_version) + +string_classes = (str, bytes) +int_classes = (int, ) + +np_str_obj_array_pattern = re.compile(r'[SaUO]') +default_collate_err_msg_format = ( + 'default_collate: batch must contain tensors, numpy arrays, numbers, ' + 'dicts or lists; found {}') + + +def gpu_batch_collate(batch, device_id=0): + """ Modified from pytorch default collate function. + When using gpu operation in preprocess pipelines, tensor could be on different devices. + While storage._new_shared will use only cuda:0, it will crash at elem.new(storage). + + Args: + batch (list): List of contents to be collated. + device_id (int): GPU device id where cuda type tensor will be on, default is 0. + + Returns: + Inputs in batch. + + Raises: + TypeError: Only support tensors, numpy arrays, numbers, dicts, lists, strings. + """ + elem: object = batch[0] + elem_type = type(elem) + if isinstance(elem, torch.Tensor): + if not elem.is_cuda: + out = None + if torch.utils.data.get_worker_info() is not None: + # If we're in a background process, concatenate directly into a + # shared memory tensor to avoid an extra copy + numel = sum([x.numel() for x in batch]) + storage = elem.storage()._new_shared(numel) + out = elem.new(storage) + return torch.stack(batch, 0, out=out) + else: + return torch.stack(batch, 0).cuda(device_id) + elif elem_type.__module__ == 'numpy' and elem_type.__name__ != 'str_' \ + and elem_type.__name__ != 'string_': + elem = batch[0] + if elem_type.__name__ == 'ndarray': + # array of string classes and object + if np_str_obj_array_pattern.search(elem.dtype.str) is not None: + raise TypeError( + default_collate_err_msg_format.format(elem.dtype)) + + return gpu_batch_collate([torch.as_tensor(b) for b in batch], + device_id=device_id) + elif elem.shape == (): # scalars + return torch.as_tensor(batch) + elif isinstance(elem, float): + return torch.tensor(batch, dtype=torch.float64) + elif isinstance(elem, int_classes): + return torch.tensor(batch) + elif isinstance(elem, string_classes): + return batch + elif isinstance(elem, container_abcs.Mapping): + return { + key: gpu_batch_collate([d[key] for d in batch], + device_id=device_id) + for key in elem + } + elif isinstance(elem, tuple) and hasattr(elem, '_fields'): # namedtuple + return elem_type(*(gpu_batch_collate(samples, device_id=device_id) + for samples in zip(*batch))) + elif isinstance(elem, container_abcs.Sequence): + transposed = zip(*batch) + return [ + gpu_batch_collate(samples, device_id=device_id) + for samples in transposed + ] + raise TypeError(default_collate_err_msg_format.format(elem_type)) + + +class DataObject(object): + para_dict = { + 'PIN_MEMORY': { + 'value': False, + 'description': 'pin_memory for data loader' + }, + 'SHUFFLE': { + 'value': False, + 'description': 'SHUFFLE the list or not!' + }, + 'BATCH_SIZE': { + 'value': 4, + 'description': 'batch size for data' + }, + 'NUM_WORKERS': { + 'value': 1, + 'description': 'num workers for fetching data!' + }, + 'NUM_FOLDS': { + 'value': + 0, + 'description': + 'if set, use MultiFoldDistributedSampler for distribute training!' + }, + 'SAMPLER': { + 'NAME': { + 'value': + None, + 'description': + 'If set None, system will choose the sampler according to world size. This means ' + 'when world size > 1 and num fold > 1, choose MultiFoldDistributedSampler,' + 'when world size > 1 and num fold <= 1, choose DistributedSampler,' + 'when world size = 1, choose default torch dataloader sampler.' + 'If set TorchDefault, choose default torch dataloader, which means every process will use all dataset.' + 'If set MultiLevelBatchSampler, choose MultiLevelBatchSampler, which loads data from index file.' + 'If set MixtureOfSamplers, choose MixtureOfSamplers, which means you can use different datasets from ' + 'different sources with multi MultiLevelBatchSampler by setting SUB_SAMPLERS.' + }, + 'INDEX_FILE': { + 'value': + None, + 'description': + 'Set your index file when you choose MultiLevelBatchSampler' + }, + 'SUB_SAMPLERS': [] + } + } + + def __init__(self, cfg, dataset, logger=None): + self.dataset = dataset + + self.logger = logger + self.pin_memory = cfg.get('PIN_MEMORY', False) + self.batch_size = cfg.get('BATCH_SIZE', 4) + self.num_workers = cfg.get('NUM_WORKERS', 1) + self.data_sampler_config = cfg.get('SAMPLER', None) + self.shuffle = cfg.get('MODE', 'test') == 'train' + self.cfg = cfg + worker_init_fn = dataset.worker_init_fn if hasattr( + dataset, 'worker_init_fn') else None + if worker_init_fn is not None: + worker_init_fn = partial(worker_init_fn, + num_workers=self.num_workers) + + self.registry_sampler(dataset) + + collate_fn = dataset.collate_fn if hasattr(dataset, + 'collate_fn') else None + if self.batch_sampler: + self.dataloader = DataLoader( + dataset, + batch_sampler=self.batch_sampler, + num_workers=self.num_workers, + collate_fn=collate_fn, + pin_memory=self.pin_memory, + prefetch_factor=None if self.num_workers == 0 else 2, + worker_init_fn=worker_init_fn, + persistent_workers=self.num_workers > 0, + timeout=2400 if self.num_workers > 0 else 0) + else: + self.dataloader = DataLoader( + dataset, + batch_size=self.batch_size, + shuffle=self.shuffle, + sampler=self.sampler, + num_workers=self.num_workers, + collate_fn=collate_fn, + pin_memory=self.pin_memory, + drop_last=self.drop_last, + worker_init_fn=worker_init_fn, + persistent_workers=self.num_workers > 0, + timeout=2400 if self.num_workers > 0 else 0) + if self.logger is not None: + self.logger.info( + f'Built dataloader with len {len(self.dataloader)}') + + def reload(self, dataset): + worker_init_fn = dataset.worker_init_fn if hasattr( + dataset, 'worker_init_fn') else None + if worker_init_fn is not None: + worker_init_fn = partial(worker_init_fn, + num_workers=self.num_workers) + + self.registry_sampler(dataset) + + collate_fn = dataset.collate_fn if hasattr(dataset, + 'collate_fn') else None + + if self.batch_sampler: + dataloader = DataLoader( + dataset, + batch_sampler=self.batch_sampler, + num_workers=self.num_workers, + collate_fn=collate_fn, + pin_memory=self.pin_memory, + prefetch_factor=None if self.num_workers == 0 else 2, + worker_init_fn=worker_init_fn, + persistent_workers=self.num_workers > 0, + timeout=2400 if self.num_workers > 0 else 0) + else: + dataloader = DataLoader( + dataset, + batch_size=self.batch_size, + shuffle=self.shuffle, + sampler=self.sampler, + num_workers=self.num_workers, + collate_fn=collate_fn, + pin_memory=self.pin_memory, + drop_last=self.drop_last, + worker_init_fn=worker_init_fn, + persistent_workers=self.num_workers > 0, + timeout=2400 if self.num_workers > 0 else 0) + + if self.logger is not None: + self.logger.info(f'Reload dataloader with len {len(dataloader)}') + return dataloader + + def registry_sampler(self, dataset): + self.sampler, self.batch_sampler, self.drop_last = None, None, False + from scepter.modules.utils.distribute import we + seed, world_size, rank = we.seed, we.world_size, we.rank + if not self.data_sampler_config: + if world_size > 1: + if self.cfg.have('NUM_FOLDS') and self.cfg.NUM_FOLDS > 1: + num_folds = self.cfg.NUM_FOLDS + print('num folds', num_folds) + self.sampler = MultiFoldDistributedSampler( + dataset, + num_folds, + world_size, + rank, + shuffle=self.shuffle) + else: + self.sampler = DistributedSampler(dataset, + world_size, + rank, + shuffle=self.shuffle) + # collate_fn = partial(gpu_batch_collate, device_id=device_id) + self.drop_last = False + self.shuffle = False + else: + self.sampler = None + # collate_fn = None + self.drop_last = False + else: + sampler_name = self.data_sampler_config.get('NAME', None) + if sampler_name == 'MixtureOfSamplers': + subsampler_configs = self.data_sampler_config.get( + 'SUB_SAMPLERS', []) + subsamplers = list() + subsampler_probs = list() + for ssconfig in subsampler_configs: + name = ssconfig.get('NAME') + prob = ssconfig.get('PROB', 0.0) + if name == 'MultiLevelBatchSampler': + subsampler = self._instantiate_multi_level_batch_sampler( + ssconfig, self.batch_size, rank, seed) + subsamplers.append(subsampler) + else: + # surpport register + ssconfig.SEED = seed + ssconfig.BATCH_SIZE = self.batch_size + subsampler = SAMPLERS.build(ssconfig, + logger=self.logger) + subsamplers.append(subsampler) + subsampler_probs.append(prob) + self.batch_sampler = MixtureOfSamplers(subsamplers, + subsampler_probs, rank, + seed) + elif sampler_name == 'MultiLevelBatchSampler': + self.batch_sampler = self._instantiate_multi_level_batch_sampler( + self.data_sampler_config, self.batch_size, rank, seed) + elif sampler_name == 'TorchDefault': + self.sampler = None + else: + 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) + + def _instantiate_multi_level_batch_sampler(self, sampler_config, + batch_size, rank, seed): + index_file = sampler_config.INDEX_FILE + image_size = sampler_config.get('IMAGE_SIZE', 512) + fields = sampler_config.get('FIELDS', ['img_path', 'prompt']) + delimiter = sampler_config.get('DELIMITER', ',') + path_prefix = sampler_config.get('PATH_PREFIX', '') + prompt_prefix = sampler_config.get('PROMPT_PREFIX', '') + return MultiLevelBatchSampler(batch_size, index_file, image_size, + fields, delimiter, path_prefix, + prompt_prefix, rank, seed) + + +def build_dataset_config(cfg, registry, logger=None, *args, **kwargs): + """ Default builder function. + + Args: + cfg (objective attribution): A set of objective attirbutions which contain + parameters passes to target class or function. + Must contains key 'type', indicates the target class or function name. + registry (Registry): An registry to search target class or function. + kwargs (dict, optional): Other params not in config dict. + + Returns: + Target class object or object returned by invoking function. + + Raises: + TypeError: + KeyError: + Exception: + """ + from scepter.modules.utils.config import Config + if not isinstance(cfg, Config): + raise TypeError(f'config must be type dict, got {type(cfg)}') + if not cfg.have('NAME'): + raise KeyError(f'config must contain key NAME, got {cfg}') + if not isinstance(registry, Registry): + raise TypeError( + f'registry must be type Registry, got {type(registry)}') + + cfg = deep_copy(cfg) + + req_type = cfg.get('NAME') + if isinstance(req_type, str): + req_type_entry = registry.get(req_type) + if req_type_entry is None: + raise KeyError(f'{req_type} not found in {registry.name} registry') + + if kwargs is not None: + cfg._update_dict(kwargs) + + if old_python_version: + logger = None + + if inspect.isclass(req_type_entry): + try: + dataset = req_type_entry(cfg, logger=logger, *args, **kwargs) + do = DataObject(cfg, dataset, logger=logger) + return do + except Exception as e: + raise Exception(f'Failed to init class {req_type_entry}, with {e}') + else: + raise TypeError(f'type must be class, got {type(req_type_entry)}') + + +DATASETS = Registry('DATASETS', + common_para=DataObject.para_dict, + build_func=build_dataset_config) diff --git a/scepter/modules/data/sampler/__init__.py b/scepter/modules/data/sampler/__init__.py new file mode 100644 index 0000000..40ff88a --- /dev/null +++ b/scepter/modules/data/sampler/__init__.py @@ -0,0 +1,9 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.data.sampler.base_sampler import BaseSampler +from scepter.modules.data.sampler.registry import SAMPLERS +from scepter.modules.data.sampler.sampler import ( + EvalDistributedSampler, LoopSampler, MixtureOfSamplers, + MultiFoldDistributedSampler, MultiLevelBatchSampler, + MultiLevelBatchSamplerMultiSource) diff --git a/scepter/modules/data/sampler/base_sampler.py b/scepter/modules/data/sampler/base_sampler.py new file mode 100644 index 0000000..017f359 --- /dev/null +++ b/scepter/modules/data/sampler/base_sampler.py @@ -0,0 +1,41 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from torch.utils.data.sampler import Sampler + +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we + + +class BaseSampler(Sampler): + para_dict = {} + + def __init__(self, cfg, logger=None): + self.cfg = cfg + self.logger = logger + self.seed = cfg.get('SEED', we.seed) + self.batch_size = cfg.get('BATCH_SIZE', 1) + + def __iter__(self): + pass + + def __len__(self): + return 1 + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('SAMPLERS', + __class__.__name__, + BaseSampler.para_dict, + set_name=True) diff --git a/scepter/modules/data/sampler/registry.py b/scepter/modules/data/sampler/registry.py new file mode 100644 index 0000000..4a3324b --- /dev/null +++ b/scepter/modules/data/sampler/registry.py @@ -0,0 +1,61 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import inspect + +from scepter.modules.utils.registry import (Registry, deep_copy, + old_python_version) + + +def build_sampler_config(cfg, registry, logger=None, **kwargs): + """ Default builder function. + + Args: + cfg (objective attribution): A set of objective attirbutions which contain + parameters passes to target class or function. + Must contains key 'type', indicates the target class or function name. + registry (Registry): An registry to search target class or function. + kwargs (dict, optional): Other params not in config dict. + + Returns: + Target class object or object returned by invoking function. + + Raises: + TypeError: + KeyError: + Exception: + """ + from scepter.modules.utils.config import Config + if not isinstance(cfg, Config): + raise TypeError(f'config must be type dict, got {type(cfg)}') + if not cfg.have('NAME'): + raise KeyError(f'config must contain key NAME, got {cfg}') + if not isinstance(registry, Registry): + raise TypeError( + f'registry must be type Registry, got {type(registry)}') + + cfg = deep_copy(cfg) + + req_type = cfg.get('NAME') + if isinstance(req_type, str): + req_type_entry = registry.get(req_type) + if req_type_entry is None: + raise KeyError(f'{req_type} not found in {registry.name} registry') + + if kwargs is not None: + cfg._update_dict(kwargs) + + if old_python_version: + logger = None + + if inspect.isclass(req_type_entry): + try: + sampler = req_type_entry(cfg, logger=logger) + return sampler + except Exception as e: + raise Exception(f'Failed to init class {req_type_entry}, with {e}') + else: + raise TypeError(f'type must be class, got {type(req_type_entry)}') + + +SAMPLERS = Registry('SAMPLERS', build_func=build_sampler_config) diff --git a/scepter/modules/data/sampler/sampler.py b/scepter/modules/data/sampler/sampler.py new file mode 100644 index 0000000..2f34563 --- /dev/null +++ b/scepter/modules/data/sampler/sampler.py @@ -0,0 +1,591 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import json +import math +import numbers +import os +import sys +from collections.abc import Iterable +from typing import Optional + +import numpy as np +import torch +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.utils.config import dict_to_yaml +from scepter.modules.utils.directory import osp_path +from scepter.modules.utils.distribute import we +from scepter.modules.utils.file_system import FS + + +@SAMPLERS.register_class() +class MultiLevelBatchSamplerMultiSource(BaseSampler): + """Sampler for database with multi-level indexing. + """ + para_dict = { + 'FIELDS': { + 'value': ['img_path', 'width', 'height', 'prompt'], + 'description': 'The fields list for input record.' + }, + 'DELIMITER': { + 'value': ',', + 'description': 'The fields delimiter for input record.' + }, + 'PATH_PREFIX': { + 'value': 'datasets', + 'description': 'The path prefix for input oss key.' + }, + 'INDEX_FILE': { + 'value': '', + 'description': 'The index file.' + }, + 'PROB': { + 'value': 1.0, + 'description': 'The prob for current sampler.' + }, + 'SELECT_SOURCES': { + 'value': + None, + 'description': + 'Select data name from index, default is None, which means use all data; ' + 'use as a [] of data key to select the used name' + }, + 'SUB_DATA_WEIGHTS': { + 'value': {}, + 'description': + 'The prob for sub data weights, default is 1, ' + 'which means computing the prob according to data num ratio of total num.' + }, + 'SUB_RESOLUTION_MAP': { + 'value': {}, + 'description': + 'The resolution map for im_type, if the resolution is in im_type, please ignore this para.' + }, + 'KARGS': { + 'value': {}, + 'description': + 'The extended parameters for transfering to downstream.' + } + } + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.rng = np.random.default_rng(self.seed + we.rank) + self.fields = cfg.get('FIELDS', []) + self.num_fields = len(self.fields) + self.delimiter = cfg.get('DELIMITER', ',') + self.path_prefix = cfg.get('PATH_PREFIX', '') + common_prob = cfg.get('PROB', 1) + sub_data_weights = cfg.get('SUB_DATA_WEIGHTS', None) + sub_data_weights = {} if sub_data_weights is None else sub_data_weights.get_dict( + ) + sub_resolution_map = cfg.get('SUB_RESOLUTION_MAP', None) + sub_resolution_map = {} if sub_resolution_map is None else sub_resolution_map.get_dict( + ) + sub_resolution_map = { + k.lower(): v + for k, v in sub_resolution_map.items() + } + kargs = cfg.get('KARGS', None) + kargs = {} if kargs is None else kargs.get_dict() + select_sources = cfg.get('SELECT_SOURCES', None) + if isinstance(select_sources, list) and len(select_sources) < 1: + raise 'SELECT_SOURCES must be None or non-empty list.' + self.kargs = {} + + for k, v in kargs.items(): + if isinstance(v, list) and len(v) == 0: + continue + if isinstance(v, dict): + v = {k.lower(): vv for k, vv in v.items()} + if v is None: + continue + self.kargs[k.lower()] = v + # read dataset according to the source fields. + index_file = cfg.INDEX_FILE + assert index_file.endswith('.json') + with FS.get_object(index_file) as local_data: + index = json.loads(local_data.decode('utf-8')) + self.sub_data_list = [] + self.key_args = {} + sub_data_num = [] + for key in index: + im_type = index[key]['image_type'] + if 'image_size' not in index[key]: + assert im_type in sub_resolution_map + index[key]['image_size'] = sub_resolution_map[im_type] + index[key]['data_key'] = key + data_name = index[key]['data_name'] + if select_sources is not None and data_name not in select_sources: + continue + sub_data_num.append(index[key]['total'] * + sub_data_weights.get(data_name, 1)) + self.sub_data_list.append(index[key]) + if data_name in self.kargs: + self.key_args[data_name] = self.kargs.pop(data_name) + + self.probabilities = np.array(sub_data_num) / np.sum( + np.array(sub_data_num)) + self.probabilities = self.probabilities.tolist() + for sub_data, p in zip(self.sub_data_list, self.probabilities): + logger.info( + f"{sub_data['data_key']}'s sample prob: {p} * {common_prob} = " + f"{p * common_prob} and samples'num: {sub_data['total']} in this cluster." + ) + self.rng = np.random.default_rng(self.seed + we.rank) + self.oss_prefix = '/'.join(index_file.split('/')[:3]) + self.index_dir = os.path.dirname(index_file) + + def __iter__(self): + while True: + index_id = self.rng.choice(len(self.sub_data_list), + p=self.probabilities) + index = self.sub_data_list[index_id] + image_size = index['image_size'] + data_name = index['data_name'] + data_key = index['data_key'] + batch = [] + while len(batch) < self.batch_size: + n = self.batch_size - len(batch) + # read items + items = index['list'] + for _ in range(index['index_level'] - 1): + list_file = self.rng.choice(items) + list_file = osp_path(self.oss_prefix, list_file) + if not list_file.startswith(self.index_dir): + list_file = os.path.join( + self.index_dir, + '/'.join(list_file.split('/')[-2:])) + with FS.get_object(list_file) as f: + items = f.decode('utf-8').strip().split('\n') + + # sample into batch + m = min(n, len(items)) + batch += [ + i + for i in self.rng.choice(items, m, replace=False).tolist() + ] + + # check batch size + if len(batch) == self.batch_size: + break + assert len(batch) == self.batch_size + + ret_data = [] + for u in batch: + one_data = {} + res_list = u.split(self.delimiter, self.num_fields - 1) + for idx, res in enumerate(res_list): + one_data[self.fields[idx]] = res + one_data.update(self.kargs) + if data_name in self.key_args: + one_data.update(self.key_args[data_name]) + one_data['image_size'] = image_size + one_data['data_key'] = data_key + one_data['prefix'] = osp_path(self.oss_prefix, + self.path_prefix) + ret_data.append(one_data) + yield ret_data + + def __len__(self): + return sys.maxsize + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('SAMPLERS', + __class__.__name__, + MultiLevelBatchSamplerMultiSource.para_dict, + set_name=True) + + +@SAMPLERS.register_class() +class MultiFoldDistributedSampler(BaseSampler): + """Modified from DistributedSampler, which performs multi fold training for + accelerating distributed training with large batches. + + Sampler that restricts data loading to a subset of the dataset. + + It is especially useful in conjunction with + :class:`torch.nn.parallel.DistributedDataParallel`. In such case, each + process can pass a DistributedSampler instance as a DataLoader sampler, + and load a subset of the original dataset that is exclusive to it. + + .. note:: + Dataset is assumed to be of constant size. + + Arguments: + dataset: Dataset used for sampling. + num_folds (optional): Number of folds, if 1, will act same as DistributeSampler + num_replicas (optional): Number of processes participating in + distributed training. + rank (optional): Rank of the current process within num_replicas. + shuffle (optional): If true (default), sampler will shuffle the indices + + .. warning:: + In distributed mode, calling the ``set_epoch`` method is needed to + make shuffling work; each process will use the same random seed + otherwise. + """ + para_dict = {} + + def __init__(self, + dataset, + num_folds=1, + num_replicas=None, + rank=None, + shuffle=True): + """ + When num_folds = 1, MultiFoldDistributedSampler degenerates to DistributedSampler. + """ + if num_replicas is None: + if not dist.is_available(): + raise RuntimeError( + 'Requires distributed package to be available') + num_replicas = dist.get_world_size() + if rank is None: + if not dist.is_available(): + raise RuntimeError( + 'Requires distributed package to be available') + rank = dist.get_rank() + self.dataset = dataset + self.num_folds = num_folds + self.num_replicas = num_replicas + self.rank = rank + self.epoch = 0 + self.num_samples = int( + math.ceil( + len(self.dataset) * self.num_folds * 1.0 / self.num_replicas)) + self.total_size = self.num_samples * self.num_replicas + self.shuffle = shuffle + + def __iter__(self): + # deterministically shuffle based on epoch + indices = [] + for fold_idx in range(self.num_folds): + g = torch.Generator() + g.manual_seed(self.epoch + fold_idx) + if self.shuffle: + indices += torch.randperm(len(self.dataset), + generator=g).tolist() + else: + indices += list(range(len(self.dataset))) + + # add extra samples to make it evenly divisible + indices += indices[:(self.total_size - len(indices))] + assert len(indices) == self.total_size + + # subsample + indices = indices[self.rank:self.total_size:self.num_replicas] + + assert len(indices) == self.num_samples + + return iter(indices) + + def __len__(self): + return self.num_samples + + def set_epoch(self, epoch): + self.epoch = epoch + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('SAMPLERS', + __class__.__name__, + MultiFoldDistributedSampler.para_dict, + set_name=True) + + +@SAMPLERS.register_class() +class EvalDistributedSampler(BaseSampler): + """Modified from DistributedSampler. + + Notice! + 1. This sampler should only be used in test mode. + 2. This sampler will pad indices or not pad, according to `padding` flag. + In no padding mode, the last rank device may get samples less than given batch_size. + The last rank device may have less iteration number than other rank. + By the way, __len__ function may return a fake number. + """ + para_dict = {} + + def __init__(self, + dataset, + num_replicas: Optional[int] = None, + rank: Optional[int] = None, + padding: bool = False) -> None: + if num_replicas is None: + if not dist.is_available(): + raise RuntimeError( + 'Requires distributed package to be available') + num_replicas = dist.get_world_size() + if rank is None: + if not dist.is_available(): + raise RuntimeError( + 'Requires distributed package to be available') + rank = dist.get_rank() + if rank >= num_replicas or rank < 0: + raise ValueError('Invalid rank {}, rank should be in the interval' + ' [0, {}]'.format(rank, num_replicas - 1)) + self.dataset = dataset + self.num_replicas = num_replicas + self.rank = rank + self.padding = padding + + self.perfect_num_samples = math.ceil( + len(self.dataset) / self.num_replicas) + self.perfect_total_size = self.perfect_num_samples * self.num_replicas + + def __iter__(self): + indices = list(range(len(self.dataset))) + + if self.padding and len(indices) < self.perfect_total_size: + padding_size = self.perfect_total_size - len(indices) + indices += indices[:padding_size] + + return iter(indices) + + def __len__(self) -> int: + return self.perfect_num_samples + + def set_epoch(self, epoch: int) -> None: + pass + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('SAMPLERS', + __class__.__name__, + EvalDistributedSampler.para_dict, + set_name=True) + + +@SAMPLERS.register_class() +class MultiLevelBatchSampler(BaseSampler): + """Sampler for database with multi-level indexing. + """ + para_dict = { + 'IMAGE_SIZE': { + 'value': [1024, 1024], + 'description': 'The image size for input image.' + }, + 'FIELDS': { + 'value': ['img_path', 'width', 'height', 'prompt'], + 'description': 'The fields list for input record.' + }, + 'DELIMITER': { + 'value': ',', + 'description': 'The fields delimiter for input record.' + }, + 'PATH_PREFIX': { + 'value': 'datasets', + 'description': 'The path prefix for input oss key.' + }, + 'PROMPT_PREFIX': { + 'value': '', + 'description': 'The prompt prefix.' + }, + 'INDEX_FILE': { + 'value': '', + 'description': 'The index file.' + }, + } + + def __init__(self, + batch_size, + index_file, + image_size=[1024, 1024], + fields=['oss_key', 'prompt'], + delimiter=',', + path_prefix='', + prompt_prefix='', + rank=0, + seed=8888): + self.batch_size = batch_size + self.seed = seed + self.rng = np.random.default_rng(seed + rank) + if isinstance(image_size, numbers.Number): + image_size = [image_size, image_size] + assert isinstance(image_size, Iterable) and len(image_size) == 2 + self.image_size = image_size + self.fields = fields + self.num_fields = len(fields) + self.delimiter = delimiter + self.path_prefix = path_prefix + self.prompt_prefix = prompt_prefix + with FS.get_object(index_file) as local_data: + if index_file.endswith('.json'): + self.index = json.loads(local_data.decode('utf-8')) + else: + self.index = { + 'list': local_data.decode('utf-8').strip().split('\n'), + 'index_level': 1, + 'num_fields': self.num_fields + } + self.oss_prefix = '/'.join(index_file.split('/')[:3]) + self.index_dir = os.path.dirname(index_file) + + def __iter__(self): + while True: + batch = [] + while len(batch) < self.batch_size: + n = self.batch_size - len(batch) + + # read items + items = self.index['list'] + for _ in range(self.index['index_level'] - 1): + list_file = self.rng.choice(items) + list_file = osp_path(self.oss_prefix, list_file) + if not list_file.startswith(self.index_dir): + list_file = os.path.join( + self.index_dir, + '/'.join(list_file.split('/')[-2:])) + with FS.get_object(list_file) as f: + items = f.decode('utf-8').strip().split('\n') + + # sample into batch + m = min(n, len(items)) + batch += [ + osp_path(self.oss_prefix, + os.path.join(self.path_prefix, i)) + for i in self.rng.choice(items, m, replace=False).tolist() + ] + + # check batch size + if len(batch) == self.batch_size: + break + assert len(batch) == self.batch_size + + fields = self.fields + ['image_size', 'prompt_prefix'] + yield [ + u.split(self.delimiter, self.num_fields - 1) + + [self.image_size, self.prompt_prefix, fields] for u in batch + ] + + def __len__(self): + return sys.maxsize + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('SAMPLERS', + __class__.__name__, + MultiLevelBatchSampler.para_dict, + set_name=True) + + +@SAMPLERS.register_class() +class MixtureOfSamplers(BaseSampler): + para_dict = {'SUB_SAMPLERS': []} + + def __init__(self, samplers, probabilities, rank=0, seed=8888): + self.samplers = samplers + self.iterators = [iter(u) for u in samplers] + self.probabilities = probabilities + self.seed = seed + self.rng = np.random.default_rng(seed + rank) + + def __iter__(self): + while True: + index = self.rng.choice(len(self.iterators), p=self.probabilities) + yield next(self.iterators[index]) + + def __len__(self): + return sys.maxsize + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('SAMPLERS', + __class__.__name__, + MixtureOfSamplers.para_dict, + set_name=True) + + +@SAMPLERS.register_class() +class LoopSampler(BaseSampler): + para_dict = {} + + def __init__(self, cfg, logger): + super().__init__(cfg, logger) + rank = we.rank + self.rng = np.random.default_rng(self.seed + rank) + + def __iter__(self): + while True: + yield self.rng.choice(sys.maxsize) + + def __len__(self): + return sys.maxsize + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('SAMPLERS', + __class__.__name__, + LoopSampler.para_dict, + set_name=True) diff --git a/scepter/modules/model/__init__.py b/scepter/modules/model/__init__.py new file mode 100644 index 0000000..66e7062 --- /dev/null +++ b/scepter/modules/model/__init__.py @@ -0,0 +1,5 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.model import (backbone, embedder, head, loss, metric, + neck, network, tokenizer, tuner) diff --git a/scepter/modules/model/backbone/__init__.py b/scepter/modules/model/backbone/__init__.py new file mode 100644 index 0000000..03879b0 --- /dev/null +++ b/scepter/modules/model/backbone/__init__.py @@ -0,0 +1,4 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.backbone import (autoencoder, image, unet, utils, + video) diff --git a/scepter/modules/model/backbone/autoencoder/__init__.py b/scepter/modules/model/backbone/autoencoder/__init__.py new file mode 100644 index 0000000..cfa8365 --- /dev/null +++ b/scepter/modules/model/backbone/autoencoder/__init__.py @@ -0,0 +1,4 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.backbone.autoencoder.ae_module import (Decoder, + Encoder) diff --git a/scepter/modules/model/backbone/autoencoder/ae_module.py b/scepter/modules/model/backbone/autoencoder/ae_module.py new file mode 100644 index 0000000..7c011b9 --- /dev/null +++ b/scepter/modules/model/backbone/autoencoder/ae_module.py @@ -0,0 +1,342 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import torch +import torch.nn as nn +from torch.utils.checkpoint import checkpoint + +from scepter.modules.model.backbone.autoencoder.ae_utils import ( + XFORMERS_IS_AVAILBLE, AttnBlock, Downsample, MemoryEfficientAttention, + Normalize, ResnetBlock, Upsample, nonlinearity) +from scepter.modules.model.base_model import BaseModel +from scepter.modules.model.registry import BACKBONES +from scepter.modules.utils.config import dict_to_yaml + + +@BACKBONES.register_class() +class Encoder(BaseModel): + para_dict = { + 'CH': { + 'value': 128, + 'description': '' + }, + 'NUM_RES_BLOCKS': { + 'value': 2, + 'description': '' + }, + 'IN_CHANNELS': { + 'value': 3, + 'description': '' + }, + 'ATTN_RESOLUTIONS': { + 'value': [], + 'description': '' + }, + 'CH_MULT': { + 'value': [1, 2, 4, 4], + 'description': '' + }, + 'Z_CHANNELS': { + 'value': 4, + 'description': '' + }, + 'DOUBLE_Z': { + 'value': True, + 'description': '' + }, + 'DROPOUT': { + 'value': 0.0, + 'description': '' + }, + 'RESAMP_WITH_CONV': { + 'value': True, + 'description': '' + } + } + para_dict.update(BaseModel.para_dict) + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.use_checkpoint = cfg.get('USE_CHECKPOINT', False) + self.ch = cfg.CH + self.out_ch = cfg.OUT_CH + self.num_res_blocks = cfg.NUM_RES_BLOCKS + self.in_channels = cfg.IN_CHANNELS + self.attn_resolutions = cfg.ATTN_RESOLUTIONS + self.ch_mult = tuple(cfg.get('CH_MULT', [1, 2, 4, 8])) + self.z_channels = cfg.Z_CHANNELS + self.double_z = cfg.get('DOUBLE_Z', True) + self.dropout = cfg.get('DROPOUT', 0.0) + self.resamp_with_conv = cfg.get('RESAMP_WITH_CONV', True) + self.temb_ch = 0 + self.construct_model() + + def construct_model(self): + self.num_resolutions = len(self.ch_mult) + self.logger.info( + f'AE Module XFORMERS_IS_AVAILBLE : {XFORMERS_IS_AVAILBLE}') + AttentionBuilder = MemoryEfficientAttention if XFORMERS_IS_AVAILBLE else AttnBlock + + self.conv_in = torch.nn.Conv2d(self.in_channels, + self.ch, + kernel_size=3, + stride=1, + padding=1) + + curr_res = 2**(self.num_resolutions - 1) + in_ch_mult = (1, ) + tuple(self.ch_mult) + self.in_ch_mult = in_ch_mult + self.down = nn.ModuleList() + for i_level in range(self.num_resolutions): + block = nn.ModuleList() + attn = nn.ModuleList() + block_in = self.ch * in_ch_mult[i_level] + block_out = self.ch * self.ch_mult[i_level] + for i_block in range(self.num_res_blocks): + block.append( + ResnetBlock(in_channels=block_in, + out_channels=block_out, + temb_channels=self.temb_ch, + dropout=self.dropout)) + block_in = block_out + if curr_res in self.attn_resolutions: + attn.append(AttentionBuilder(block_in)) + down = nn.Module() + down.block = block + down.attn = attn + if i_level != self.num_resolutions - 1: + down.downsample = Downsample(block_in, self.resamp_with_conv) + curr_res = curr_res // 2 + self.down.append(down) + + # middle + self.mid = nn.Module() + self.mid.block_1 = ResnetBlock(in_channels=block_in, + out_channels=block_in, + temb_channels=self.temb_ch, + dropout=self.dropout) + self.mid.attn_1 = AttentionBuilder(block_in) + self.mid.block_2 = ResnetBlock(in_channels=block_in, + out_channels=block_in, + temb_channels=self.temb_ch, + dropout=self.dropout) + + # end + self.norm_out = Normalize(block_in) + self.conv_out = torch.nn.Conv2d( + block_in, + 2 * self.z_channels if self.double_z else self.z_channels, + kernel_size=3, + stride=1, + padding=1) + + def forward_ori(self, x): + # timestep embedding + temb = None + # downsampling + hs = [self.conv_in(x)] + for i_level in range(self.num_resolutions): + for i_block in range(self.num_res_blocks): + h = self.down[i_level].block[i_block](hs[-1], temb) + if len(self.down[i_level].attn) > 0: + h = self.down[i_level].attn[i_block](h) + hs.append(h) + if i_level != self.num_resolutions - 1: + hs.append(self.down[i_level].downsample(hs[-1])) + + # middle + h = hs[-1] + h = self.mid.block_1(h, temb) + h = self.mid.attn_1(h) + h = self.mid.block_2(h, temb) + + # end + h = self.norm_out(h) + h = nonlinearity(h) + h = self.conv_out(h) + + return h + + def forward(self, x): + if self.use_checkpoint: + return checkpoint(self.forward_ori, x) + else: + return self.forward_ori(x) + + @staticmethod + def get_config_template(): + return dict_to_yaml('BACKBONE', + __class__.__name__, + Encoder.para_dict, + set_name=True) + + +@BACKBONES.register_class() +class Decoder(BaseModel): + para_dict = { + 'CH': { + 'value': 128, + 'description': '' + }, + 'OUT_CH': { + 'value': 3, + 'description': '' + }, + 'NUM_RES_BLOCKS': { + 'value': 2, + 'description': '' + }, + 'IN_CHANNELS': { + 'value': 3, + 'description': '' + }, + 'ATTN_RESOLUTIONS': { + 'value': [], + 'description': '' + }, + 'CH_MULT': { + 'value': [1, 2, 4, 4], + 'description': '' + }, + 'Z_CHANNELS': { + 'value': 4, + 'description': '' + }, + 'DROPOUT': { + 'value': 0.0, + 'description': '' + }, + 'RESAMP_WITH_CONV': { + 'value': True, + 'description': '' + }, + 'GIVE_PRE_END': { + 'value': False, + 'description': '' + }, + 'TANH_OUT': { + 'value': False, + 'description': '' + } + } + para_dict.update(BaseModel.para_dict) + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.use_checkpoint = cfg.get('USE_CHECKPOINT', False) + self.ch = cfg.CH + self.out_ch = cfg.OUT_CH + self.num_res_blocks = cfg.NUM_RES_BLOCKS + self.in_channels = cfg.IN_CHANNELS + self.attn_resolutions = cfg.ATTN_RESOLUTIONS + self.ch_mult = tuple(cfg.get('CH_MULT', [1, 2, 4, 8])) + self.z_channels = cfg.Z_CHANNELS + self.dropout = cfg.get('DROPOUT', 0.0) + self.resamp_with_conv = cfg.get('RESAMP_WITH_CONV', True) + self.give_pre_end = cfg.get('GIVE_PRE_END', False) + self.tanh_out = cfg.get('TANH_OUT', False) + self.temb_ch = 0 + self.construct_model() + + def construct_model(self): + self.num_resolutions = len(self.ch_mult) + AttentionBuilder = MemoryEfficientAttention if XFORMERS_IS_AVAILBLE else AttnBlock + + # compute in_ch_mult, block_in and curr_res at lowest res + block_in = self.ch * self.ch_mult[self.num_resolutions - 1] + curr_res = 1 + # z to block_in + self.conv_in = torch.nn.Conv2d(self.z_channels, + block_in, + kernel_size=3, + stride=1, + padding=1) + + # middle + self.mid = nn.Module() + self.mid.block_1 = ResnetBlock(in_channels=block_in, + out_channels=block_in, + temb_channels=self.temb_ch, + dropout=self.dropout) + self.mid.attn_1 = AttentionBuilder(block_in) + self.mid.block_2 = ResnetBlock(in_channels=block_in, + out_channels=block_in, + temb_channels=self.temb_ch, + dropout=self.dropout) + + # upsampling + self.up = nn.ModuleList() + for i_level in reversed(range(self.num_resolutions)): + block = nn.ModuleList() + attn = nn.ModuleList() + block_out = self.ch * self.ch_mult[i_level] + for i_block in range(self.num_res_blocks + 1): + block.append( + ResnetBlock(in_channels=block_in, + out_channels=block_out, + temb_channels=self.temb_ch, + dropout=self.dropout)) + block_in = block_out + if curr_res in self.attn_resolutions: + attn.append(AttentionBuilder(block_in)) + up = nn.Module() + up.block = block + up.attn = attn + if i_level != 0: + up.upsample = Upsample(block_in, self.resamp_with_conv) + curr_res = curr_res * 2 + self.up.insert(0, up) # prepend to get consistent order + + # end + self.norm_out = Normalize(block_in) + self.conv_out = torch.nn.Conv2d(block_in, + self.out_ch, + kernel_size=3, + stride=1, + padding=1) + + def mid_upsclae_transform(self, h, temb): + h = self.mid.block_1(h, temb) + h = self.mid.attn_1(h) + h = self.mid.block_2(h, temb) + + # upsampling + for i_level in reversed(range(self.num_resolutions)): + for i_block in range(self.num_res_blocks + 1): + h = self.up[i_level].block[i_block](h, temb) + if len(self.up[i_level].attn) > 0: + h = self.up[i_level].attn[i_block](h) + if i_level != 0: + h = self.up[i_level].upsample(h) + + return h + + def forward(self, z, cond=None): + # timestep embedding + temb = None + + h = self.conv_in(z) + + # middle + if not self.use_checkpoint: + h = self.mid_upsclae_transform(h, temb) + else: + h = checkpoint(self.mid_upsclae_transform, h, temb) + + # end + if self.give_pre_end: + return h + + h = self.norm_out(h) + h = nonlinearity(h) + h = self.conv_out(h) + if self.tanh_out: + h = torch.tanh(h) + return h + + @staticmethod + def get_config_template(): + return dict_to_yaml('BACKBONE', + __class__.__name__, + Decoder.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/autoencoder/ae_utils.py b/scepter/modules/model/backbone/autoencoder/ae_utils.py new file mode 100644 index 0000000..d0c96fc --- /dev/null +++ b/scepter/modules/model/backbone/autoencoder/ae_utils.py @@ -0,0 +1,271 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import math +import warnings +from typing import Any, Optional + +import torch +import torch.nn as nn +from torch.utils.checkpoint import checkpoint + +try: + import xformers + import xformers.ops + XFORMERS_IS_AVAILBLE = True +except Exception as e: + XFORMERS_IS_AVAILBLE = False + warnings.warn(f'{e}') + + +def get_timestep_embedding(timesteps, embedding_dim): + """ + This matches the implementation in Denoising Diffusion Probabilistic Models: + From Fairseq. + Build sinusoidal embeddings. + This matches the implementation in tensor2tensor, but differs slightly + from the description in Section 3.5 of "Attention Is All You Need". + """ + assert len(timesteps.shape) == 1 + + half_dim = embedding_dim // 2 + emb = math.log(10000) / (half_dim - 1) + emb = torch.exp(torch.arange(half_dim, dtype=torch.float32) * -emb) + emb = emb.to(device=timesteps.device) + emb = timesteps.float()[:, None] * emb[None, :] + emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1) + if embedding_dim % 2 == 1: # zero pad + emb = torch.nn.functional.pad(emb, (0, 1, 0, 0)) + return emb + + +def nonlinearity(x): + # swish + return x * torch.sigmoid(x) + + +def Normalize(in_channels, num_groups=32): + return torch.nn.GroupNorm(num_groups=num_groups, + num_channels=in_channels, + eps=1e-6, + affine=True) + + +class Upsample(nn.Module): + def __init__(self, in_channels, with_conv): + super().__init__() + self.with_conv = with_conv + if self.with_conv: + self.conv = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=3, + stride=1, + padding=1) + + def forward(self, x): + x = torch.nn.functional.interpolate(x.float(), + scale_factor=2.0, + mode='nearest').type_as(x) + if self.with_conv: + x = self.conv(x) + return x + + +class Downsample(nn.Module): + def __init__(self, in_channels, with_conv): + super().__init__() + self.with_conv = with_conv + if self.with_conv: + # no asymmetric padding in torch conv, must do it ourselves + self.conv = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=3, + stride=2, + padding=0) + + def forward(self, x): + if self.with_conv: + pad = (0, 1, 0, 1) + x = torch.nn.functional.pad(x, pad, mode='constant', value=0) + x = self.conv(x) + else: + x = torch.nn.functional.avg_pool2d(x, kernel_size=2, stride=2) + return x + + +class ResnetBlock(nn.Module): + def __init__(self, + *, + in_channels, + out_channels=None, + conv_shortcut=False, + dropout, + temb_channels=512): + super().__init__() + self.in_channels = in_channels + out_channels = in_channels if out_channels is None else out_channels + self.out_channels = out_channels + self.use_conv_shortcut = conv_shortcut + + self.norm1 = Normalize(in_channels) + self.conv1 = torch.nn.Conv2d(in_channels, + out_channels, + kernel_size=3, + stride=1, + padding=1) + if temb_channels > 0: + self.temb_proj = torch.nn.Linear(temb_channels, out_channels) + self.norm2 = Normalize(out_channels) + self.dropout = torch.nn.Dropout(dropout) + self.conv2 = torch.nn.Conv2d(out_channels, + out_channels, + kernel_size=3, + stride=1, + padding=1) + if self.in_channels != self.out_channels: + if self.use_conv_shortcut: + self.conv_shortcut = torch.nn.Conv2d(in_channels, + out_channels, + kernel_size=3, + stride=1, + padding=1) + else: + self.nin_shortcut = torch.nn.Conv2d(in_channels, + out_channels, + kernel_size=1, + stride=1, + padding=0) + + def forward(self, x, temb): + h = x + h = self.norm1(h) + h = nonlinearity(h) + h = self.conv1(h) + + if temb is not None: + h = h + self.temb_proj(nonlinearity(temb))[:, :, None, None] + + h = self.norm2(h) + h = nonlinearity(h) + h = self.dropout(h) + h = self.conv2(h) + + if self.in_channels != self.out_channels: + if self.use_conv_shortcut: + x = self.conv_shortcut(x) + else: + x = self.nin_shortcut(x) + + return x + h + + +class AttnBlock(nn.Module): + def __init__(self, in_channels): + super().__init__() + self.in_channels = in_channels + + self.norm = Normalize(in_channels) + self.q = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + self.k = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + self.v = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + self.proj_out = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + + def forward(self, x): + h_ = x + h_ = self.norm(h_) + q = self.q(h_) + k = self.k(h_) + v = self.v(h_) + + # compute attention + b, c, h, w = q.shape + q = q.reshape(b, c, h * w) + q = q.permute(0, 2, 1) # b,hw,c + k = k.reshape(b, c, h * w) # b,c,hw + w_ = torch.bmm(q, k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j] + w_ = w_ * (int(c)**(-0.5)) + w_ = torch.nn.functional.softmax(w_, dim=2) + + # attend to values + v = v.reshape(b, c, h * w) + w_ = w_.permute(0, 2, 1) # b,hw,hw (first hw of k, second of q) + h_ = torch.bmm( + v, w_) # b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j] + h_ = h_.reshape(b, c, h, w) + + h_ = self.proj_out(h_) + + return x + h_ + + +class MemoryEfficientAttention(nn.Module): + def __init__(self, in_channels, use_checkpoint=False): + super().__init__() + self.use_checkpoint = use_checkpoint + self.in_channels = in_channels + + self.norm = Normalize(in_channels) + self.q = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + self.k = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + self.v = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + self.proj_out = torch.nn.Conv2d(in_channels, + in_channels, + kernel_size=1, + stride=1, + padding=0) + self.attention_op: Optional[Any] = None + + def forward_ori(self, x): + h_ = x + h_ = self.norm(h_) + q = self.q(h_) + k = self.k(h_) + v = self.v(h_) + + b, c, h, w = q.shape + q, k, v = map( + lambda t: t.view(b, c, h * w).permute(0, 2, 1).contiguous(), + (q, k, v)) + + out = xformers.ops.memory_efficient_attention(q, + k, + v, + attn_bias=None, + op=self.attention_op) + out = out.permute(0, 2, 1).view(b, c, h, w) + + return x + self.proj_out(out) + + def forward(self, x): + if self.use_checkpoint: + return checkpoint(self.forward_ori, x) + else: + return self.forward_ori(x) diff --git a/scepter/modules/model/backbone/image/__init__.py b/scepter/modules/model/backbone/image/__init__.py new file mode 100644 index 0000000..0264a75 --- /dev/null +++ b/scepter/modules/model/backbone/image/__init__.py @@ -0,0 +1,8 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.backbone.image.mlp import MLP +from scepter.modules.model.backbone.image.resnet import ResNet +from scepter.modules.model.backbone.image.timm_model import TIMM_MODEL +from scepter.modules.model.backbone.image.vit_modify import ( + MultiHeadSomeFTVisualTransformer, SomeFTVisualTransformer, + SomeFTVisualTransformerTwoPart, VisualTransformer) diff --git a/scepter/modules/model/backbone/image/mlp.py b/scepter/modules/model/backbone/image/mlp.py new file mode 100644 index 0000000..58f10df --- /dev/null +++ b/scepter/modules/model/backbone/image/mlp.py @@ -0,0 +1,94 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import torch.nn as nn + +from scepter.modules.model.base_model import BaseModel +from scepter.modules.model.registry import BACKBONES +from scepter.modules.utils.config import dict_to_yaml + + +def MLP_unit(input_dim, output_dim, use_bn=False, use_relu=False): + layers = [] + layers.append(nn.Linear(input_dim, output_dim)) + if use_bn: + layers.append(nn.BatchNorm1d(output_dim)) + if use_relu: + layers.append(nn.ReLU(inplace=True)) + + return nn.Sequential(*layers) + + +@BACKBONES.register_class() +class MLP(BaseModel): + para_dict = { + 'IN_DIM': { + 'value': [10, 10], + 'description': + "the input dim for each linear, which also the previous' out dim!" + }, + 'OUT_DIM': { + 'value': + 10, + 'description': + 'The output dim for head, often this value is the classes number!' + }, + 'USE_BN': { + 'value': [True], + 'description': + 'The MLP before proj use bn or not for each layer! len() = len(IN_DIM) - 1' + }, + 'USE_RELU': { + 'value': [False], + 'description': + 'The MLP before proj use relu or not for each layer!' + } + } + + def __init__(self, cfg, logger=None): + super(MLP, self).__init__(cfg, logger=logger) + self.dim_list = cfg.IN_DIM + self.use_bn = cfg.USE_BN + self.use_relu = cfg.USE_RELU + self.out_feature = cfg.OUT_DIM + assert len(self.dim_list) >= 1 + layers = [] + for idx, dim in enumerate(self.dim_list): + if idx == 0: + in_feature = dim + else: + out_feature = dim + layers.append( + MLP_unit(in_feature, + out_feature, + use_bn=self.use_bn[idx - 1], + use_relu=self.use_relu[idx - 1])) + in_feature = dim + + self.mlp = nn.Sequential(*layers) + self.fc = nn.Linear(in_feature, self.out_feature) + self.bn = nn.BatchNorm1d(self.out_feature) + + def forward(self, x): + x = x.type(self.fc.weight.dtype) + x = self.mlp(x) + x = self.fc(x) + x = self.bn(x) + return x + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('HEADS', + __class__.__name__, + MLP.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/image/resnet.py b/scepter/modules/model/backbone/image/resnet.py new file mode 100644 index 0000000..04f42c4 --- /dev/null +++ b/scepter/modules/model/backbone/image/resnet.py @@ -0,0 +1,94 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.model.backbone.image.resnet_impl import (resnet18, + resnet34, + resnet50, + resnet101, + resnet152) +from scepter.modules.model.base_model import BaseModel +from scepter.modules.model.registry import BACKBONES +from scepter.modules.utils.config import dict_to_yaml + + +@BACKBONES.register_class('ResNet') +class ResNet(BaseModel): + para_dict = { + 'DEPTH': { + 'value': 18, + 'description': 'the depth of network for resnet!' + }, + 'KERNEL_SIZE': { + 'value': + 7, + 'description': + 'first conv kernel size, 7 or 3 (without stride and maxpooling)!' + }, + 'USE_RELU': { + 'value': True, + 'description': 'use relu or not!' + }, + 'USE_MAXPOOL': { + 'value': True, + 'description': 'use maxpool or not!' + }, + 'FIRST_CONV_STRIDE': { + 'value': 1, + 'description': 'first conv stride 1 or 2!' + }, + 'FIRST_MAX_POOL_STRIDE': { + 'value': 1, + 'description': 'first max pool stride 1 or 2!' + }, + 'PRETRAINED': { + 'value': False, + 'description': 'if load the official pretrained model or not.' + } + } + + def __init__(self, cfg, logger=None): + super(ResNet, self).__init__(cfg, logger=logger) + depth = cfg.get('DEPTH', 18) + pretrained = cfg.get('PRETRAINED', False) + kernel_size = cfg.get('KERNEL_SIZE', 7) + use_relu = cfg.get('USE_RELU', True) + use_maxpool = cfg.get('USE_MAXPOOL', True) + first_conv_stride = cfg.get('FIRST_CONV_STRIDE', 1) + first_max_pool_stride = cfg.get('FIRST_MAX_POOL_STRIDE', 1) + depth_mapper = { + 18: resnet18, + 34: resnet34, + 50: resnet50, + 101: resnet101, + 152: resnet152 + } + cons_func = depth_mapper.get(depth) + if cons_func is None: + raise KeyError(f'Unsupported depth for resnet, {depth}') + self.model = cons_func(pretrained=pretrained, + kernel_size=kernel_size, + use_relu=use_relu, + use_maxpool=use_maxpool, + first_conv_stride=first_conv_stride, + first_max_pool_stride=first_max_pool_stride) + + def forward(self, x): + return self.model.forward(x) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BACKBONES', + __class__.__name__, + ResNet.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/image/resnet_impl.py b/scepter/modules/model/backbone/image/resnet_impl.py new file mode 100644 index 0000000..2e3c660 --- /dev/null +++ b/scepter/modules/model/backbone/image/resnet_impl.py @@ -0,0 +1,572 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +# Modified from https://github.com/pytorch/vision/blob/main/torchvision/models/resnet.py + +# BSD 3-Clause License +# +# Copyright (c) Soumith Chintala 2016, +# All rights reserved. +# +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: +# +# * Redistributions of source code must retain the above copyright notice, this +# list of conditions and the following disclaimer. +# +# * Redistributions in binary form must reproduce the above copyright notice, +# this list of conditions and the following disclaimer in the documentation +# and/or other materials provided with the distribution. +# +# * Neither the name of the copyright holder nor the names of its +# contributors may be used to endorse or promote products derived from +# this software without specific prior written permission. +# +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE +# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +from typing import Any, Callable, List, Optional, Type, Union + +import torch.nn as nn +from torch import Tensor + +try: + from torch.hub import load_state_dict_from_url +except ImportError: + from torch.utils.model_zoo import load_url as load_state_dict_from_url + +__all__ = [ + 'ResNet', 'resnet18', 'resnet34', 'resnet50', 'resnet101', 'resnet152', + 'resnext50_32x4d', 'resnext101_32x8d', 'wide_resnet50_2', + 'wide_resnet101_2' +] + +model_urls = { + 'resnet18': + 'https://download.pytorch.org/models/resnet18-f37072fd.pth', + 'resnet34': + 'https://download.pytorch.org/models/resnet34-b627a593.pth', + 'resnet50': + 'https://download.pytorch.org/models/resnet50-0676ba61.pth', + 'resnet101': + 'https://download.pytorch.org/models/resnet101-63fe2227.pth', + 'resnet152': + 'https://download.pytorch.org/models/resnet152-394f9c45.pth', + 'resnext50_32x4d': + 'https://download.pytorch.org/models/resnext50_32x4d-7cdf4587.pth', + 'resnext101_32x8d': + 'https://download.pytorch.org/models/resnext101_32x8d-8ba56ff5.pth', + 'wide_resnet50_2': + 'https://download.pytorch.org/models/wide_resnet50_2-95faca4d.pth', + 'wide_resnet101_2': + 'https://download.pytorch.org/models/wide_resnet101_2-32ee1156.pth', +} + + +def conv3x3(in_planes: int, + out_planes: int, + stride: int = 1, + groups: int = 1, + dilation: int = 1) -> nn.Conv2d: + """3x3 convolution with padding""" + return nn.Conv2d(in_planes, + out_planes, + kernel_size=3, + stride=stride, + padding=dilation, + groups=groups, + bias=False, + dilation=dilation) + + +def conv1x1(in_planes: int, out_planes: int, stride: int = 1) -> nn.Conv2d: + """1x1 convolution""" + return nn.Conv2d(in_planes, + out_planes, + kernel_size=1, + stride=stride, + bias=False) + + +class BasicBlock(nn.Module): + expansion: int = 1 + + def __init__( + self, + inplanes: int, + planes: int, + stride: int = 1, + downsample: Optional[nn.Module] = None, + groups: int = 1, + base_width: int = 64, + dilation: int = 1, + norm_layer: Optional[Callable[..., nn.Module]] = None) -> None: + super(BasicBlock, self).__init__() + if norm_layer is None: + norm_layer = nn.BatchNorm2d + if groups != 1 or base_width != 64: + raise ValueError( + 'BasicBlock only supports groups=1 and base_width=64') + if dilation > 1: + raise NotImplementedError( + 'Dilation > 1 not supported in BasicBlock') + # Both self.conv1 and self.downsample layers downsample the input when stride != 1 + self.conv1 = conv3x3(inplanes, planes, stride) + self.bn1 = norm_layer(planes) + self.relu = nn.ReLU(inplace=True) + self.conv2 = conv3x3(planes, planes) + self.bn2 = norm_layer(planes) + self.downsample = downsample + self.stride = stride + + def forward(self, x: Tensor) -> Tensor: + identity = x + + out = self.conv1(x) + out = self.bn1(out) + out = self.relu(out) + + out = self.conv2(out) + out = self.bn2(out) + + if self.downsample is not None: + identity = self.downsample(x) + + out += identity + out = self.relu(out) + + return out + + +class Bottleneck(nn.Module): + # Bottleneck in torchvision places the stride for downsampling at 3x3 convolution(self.conv2) + # while original implementation places the stride at the first 1x1 convolution(self.conv1) + # according to "Deep residual learning for image recognition"https://arxiv.org/abs/1512.03385. + # This variant is also known as ResNet V1.5 and improves accuracy according to + # https://ngc.nvidia.com/catalog/model-scripts/nvidia:resnet_50_v1_5_for_pytorch. + + expansion: int = 4 + + def __init__( + self, + inplanes: int, + planes: int, + stride: int = 1, + downsample: Optional[nn.Module] = None, + groups: int = 1, + base_width: int = 64, + dilation: int = 1, + norm_layer: Optional[Callable[..., nn.Module]] = None) -> None: + super(Bottleneck, self).__init__() + if norm_layer is None: + norm_layer = nn.BatchNorm2d + width = int(planes * (base_width / 64.)) * groups + # Both self.conv2 and self.downsample layers downsample the input when stride != 1 + self.conv1 = conv1x1(inplanes, width) + self.bn1 = norm_layer(width) + self.conv2 = conv3x3(width, width, stride, groups, dilation) + self.bn2 = norm_layer(width) + self.conv3 = conv1x1(width, planes * self.expansion) + self.bn3 = norm_layer(planes * self.expansion) + self.relu = nn.ReLU(inplace=True) + self.downsample = downsample + self.stride = stride + + def forward(self, x: Tensor) -> Tensor: + identity = x + + out = self.conv1(x) + out = self.bn1(out) + out = self.relu(out) + + out = self.conv2(out) + out = self.bn2(out) + out = self.relu(out) + + out = self.conv3(out) + out = self.bn3(out) + + if self.downsample is not None: + identity = self.downsample(x) + + out += identity + out = self.relu(out) + + return out + + +class ResNet(nn.Module): + def __init__( + self, + block: Type[Union[BasicBlock, Bottleneck]], + layers: List[int], + kernel_size=7, + use_relu=True, + use_maxpool=True, + first_conv_stride=1, + first_max_pool_stride=1, + num_classes: int = 1000, + zero_init_residual: bool = False, + groups: int = 1, + width_per_group: int = 64, + replace_stride_with_dilation: Optional[List[bool]] = None, + norm_layer: Optional[Callable[..., nn.Module]] = None) -> None: + super(ResNet, self).__init__() + if norm_layer is None: + norm_layer = nn.BatchNorm2d + self._norm_layer = norm_layer + + self.inplanes = 64 + self.dilation = 1 + if replace_stride_with_dilation is None: + # each element in the tuple indicates if we should replace + # the 2x2 stride with a dilated convolution instead + replace_stride_with_dilation = [False, False, False] + if len(replace_stride_with_dilation) != 3: + raise ValueError('replace_stride_with_dilation should be None ' + 'or a 3-element tuple, got {}'.format( + replace_stride_with_dilation)) + self.groups = groups + self.base_width = width_per_group + if kernel_size == 7: + self.conv1 = nn.Conv2d(3, + self.inplanes, + kernel_size=7, + stride=2, + padding=3, + bias=False) + self.bn1 = norm_layer(self.inplanes) + self.relu = nn.ReLU(inplace=True) + self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1) + elif kernel_size == 3: + self.conv1 = nn.Conv2d(3, + self.inplanes, + kernel_size=3, + stride=first_conv_stride, + padding=1, + bias=False) + self.bn1 = norm_layer(self.inplanes) + if use_relu: + self.relu = nn.ReLU(inplace=True) + if use_maxpool: + self.maxpool = nn.MaxPool2d(kernel_size=3, + stride=first_max_pool_stride, + padding=1) + self.layer1 = self._make_layer(block, 64, layers[0]) + self.layer2 = self._make_layer(block, + 128, + layers[1], + stride=2, + dilate=replace_stride_with_dilation[0]) + self.layer3 = self._make_layer(block, + 256, + layers[2], + stride=2, + dilate=replace_stride_with_dilation[1]) + self.layer4 = self._make_layer(block, + 512, + layers[3], + stride=2, + dilate=replace_stride_with_dilation[2]) + # self.avgpool = nn.AdaptiveAvgPool2d((1, 1)) + # self.fc = nn.Linear(512 * block.expansion, num_classes) + + for m in self.modules(): + if isinstance(m, nn.Conv2d): + nn.init.kaiming_normal_(m.weight, + mode='fan_out', + nonlinearity='relu') + elif isinstance(m, (nn.BatchNorm2d, nn.GroupNorm)): + nn.init.constant_(m.weight, 1) + nn.init.constant_(m.bias, 0) + + # Zero-initialize the last BN in each residual branch, + # so that the residual branch starts with zeros, and each residual block behaves like an identity. + # This improves the model by 0.2~0.3% according to https://arxiv.org/abs/1706.02677 + if zero_init_residual: + for m in self.modules(): + if isinstance(m, Bottleneck): + nn.init.constant_(m.bn3.weight, + 0) # type: ignore[arg-type] + elif isinstance(m, BasicBlock): + nn.init.constant_(m.bn2.weight, + 0) # type: ignore[arg-type] + + def _make_layer(self, + block: Type[Union[BasicBlock, Bottleneck]], + planes: int, + blocks: int, + stride: int = 1, + dilate: bool = False) -> nn.Sequential: + norm_layer = self._norm_layer + downsample = None + previous_dilation = self.dilation + if dilate: + self.dilation *= stride + stride = 1 + if stride != 1 or self.inplanes != planes * block.expansion: + downsample = nn.Sequential( + conv1x1(self.inplanes, planes * block.expansion, stride), + norm_layer(planes * block.expansion), + ) + + layers = [] + layers.append( + block(self.inplanes, planes, stride, downsample, self.groups, + self.base_width, previous_dilation, norm_layer)) + self.inplanes = planes * block.expansion + for _ in range(1, blocks): + layers.append( + block(self.inplanes, + planes, + groups=self.groups, + base_width=self.base_width, + dilation=self.dilation, + norm_layer=norm_layer)) + + return nn.Sequential(*layers) + + def _forward_impl(self, x: Tensor) -> Tensor: + # See note [TorchScript super()] + x = self.conv1(x) + x = self.bn1(x) + if hasattr(self, 'relu'): + x = self.relu(x) + if hasattr(self, 'maxpool'): + x = self.maxpool(x) + + x = self.layer1(x) + x = self.layer2(x) + x = self.layer3(x) + x = self.layer4(x) + + # x = self.avgpool(x) + # x = torch.flatten(x, 1) + # x = self.fc(x) + + return x + + def forward(self, x: Tensor) -> Tensor: + return self._forward_impl(x) + + +def _resnet(arch: str, + block: Type[Union[BasicBlock, Bottleneck]], + layers: List[int], + pretrained: bool, + progress: bool, + kernel_size: int, + use_relu=True, + use_maxpool=True, + first_conv_stride=1, + first_max_pool_stride=1, + **kwargs: Any) -> ResNet: + model = ResNet(block, + layers, + kernel_size=kernel_size, + use_relu=use_relu, + use_maxpool=use_maxpool, + first_conv_stride=first_conv_stride, + first_max_pool_stride=first_max_pool_stride, + **kwargs) + if pretrained: + state_dict = load_state_dict_from_url(model_urls[arch], + progress=progress) + state_dict.move_to_end('fc.weight', last=True) + state_dict.popitem(last=True) + state_dict.move_to_end('fc.bias', last=True) + state_dict.popitem(last=True) + model.load_state_dict(state_dict) + return model + + +def resnet18(pretrained: bool = False, + progress: bool = True, + kernel_size=7, + use_relu=True, + use_maxpool=True, + first_conv_stride=1, + first_max_pool_stride=1, + **kwargs: Any) -> ResNet: + r"""ResNet-18 model from + `"Deep Residual Learning for Image Recognition" `_. + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + return _resnet('resnet18', + BasicBlock, [2, 2, 2, 2], + pretrained, + progress, + kernel_size=kernel_size, + use_relu=use_relu, + use_maxpool=use_maxpool, + first_conv_stride=first_conv_stride, + first_max_pool_stride=first_max_pool_stride, + **kwargs) + + +def resnet34(pretrained: bool = False, + progress: bool = True, + kernel_size=7, + **kwargs: Any) -> ResNet: + r"""ResNet-34 model from + `"Deep Residual Learning for Image Recognition" `_. + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + return _resnet('resnet34', + BasicBlock, [3, 4, 6, 3], + pretrained, + progress, + kernel_size=kernel_size, + **kwargs) + + +def resnet50(pretrained: bool = False, + progress: bool = True, + kernel_size=7, + **kwargs: Any) -> ResNet: + r"""ResNet-50 model from + `"Deep Residual Learning for Image Recognition" `_. + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + return _resnet('resnet50', + Bottleneck, [3, 4, 6, 3], + pretrained, + progress, + kernel_size=kernel_size, + **kwargs) + + +def resnet101(pretrained: bool = False, + progress: bool = True, + kernel_size=7, + **kwargs: Any) -> ResNet: + r"""ResNet-101 model from + `"Deep Residual Learning for Image Recognition" `_. + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + return _resnet('resnet101', + Bottleneck, [3, 4, 23, 3], + pretrained, + progress, + kernel_size=kernel_size, + **kwargs) + + +def resnet152(pretrained: bool = False, + progress: bool = True, + kernel_size=7, + **kwargs: Any) -> ResNet: + r"""ResNet-152 model from + `"Deep Residual Learning for Image Recognition" `_. + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + return _resnet('resnet152', + Bottleneck, [3, 8, 36, 3], + pretrained, + progress, + kernel_size=kernel_size, + **kwargs) + + +def resnext50_32x4d(pretrained: bool = False, + progress: bool = True, + kernel_size=7, + **kwargs: Any) -> ResNet: + r"""ResNeXt-50 32x4d model from + `"Aggregated Residual Transformation for Deep Neural Networks" `_. + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + kwargs['groups'] = 32 + kwargs['width_per_group'] = 4 + return _resnet('resnext50_32x4d', + Bottleneck, [3, 4, 6, 3], + pretrained, + progress, + kernel_size=kernel_size, + **kwargs) + + +def resnext101_32x8d(pretrained: bool = False, + progress: bool = True, + kernel_size=7, + **kwargs: Any) -> ResNet: + r"""ResNeXt-101 32x8d model from + `"Aggregated Residual Transformation for Deep Neural Networks" `_. + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + kwargs['groups'] = 32 + kwargs['width_per_group'] = 8 + return _resnet('resnext101_32x8d', + Bottleneck, [3, 4, 23, 3], + pretrained, + progress, + kernel_size=7, + **kwargs) + + +def wide_resnet50_2(pretrained: bool = False, + progress: bool = True, + kernel_size=7, + **kwargs: Any) -> ResNet: + r"""Wide ResNet-50-2 model from + `"Wide Residual Networks" `_. + The model is the same as ResNet except for the bottleneck number of channels + which is twice larger in every block. The number of channels in outer 1x1 + convolutions is the same, e.g. last block in ResNet-50 has 2048-512-2048 + channels, and in Wide ResNet-50-2 has 2048-1024-2048. + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + kwargs['width_per_group'] = 64 * 2 + return _resnet('wide_resnet50_2', + Bottleneck, [3, 4, 6, 3], + pretrained, + progress, + kernel_size=kernel_size, + **kwargs) + + +def wide_resnet101_2(pretrained: bool = False, + progress: bool = True, + kernel_size=7, + **kwargs: Any) -> ResNet: + r"""Wide ResNet-101-2 model from + `"Wide Residual Networks" `_. + The model is the same as ResNet except for the bottleneck number of channels + which is twice larger in every block. The number of channels in outer 1x1 + convolutions is the same, e.g. last block in ResNet-50 has 2048-512-2048 + channels, and in Wide ResNet-50-2 has 2048-1024-2048. + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + kwargs['width_per_group'] = 64 * 2 + return _resnet('wide_resnet101_2', + Bottleneck, [3, 4, 23, 3], + pretrained, + progress, + kernel_size=kernel_size, + **kwargs) diff --git a/scepter/modules/model/backbone/image/timm_model.py b/scepter/modules/model/backbone/image/timm_model.py new file mode 100644 index 0000000..0a2b526 --- /dev/null +++ b/scepter/modules/model/backbone/image/timm_model.py @@ -0,0 +1,55 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.base_model import BaseModel +from scepter.modules.model.registry import BACKBONES +from scepter.modules.utils.config import dict_to_yaml + + +@BACKBONES.register_class() +class TIMM_MODEL(BaseModel): + para_dict = { + 'MODEL_NAME': { + 'value': '', + 'description': 'The name of timm!' + }, + 'NUM_CLASSES': { + 'value': 1000, + 'description': 'The num class for your task!' + }, + 'PRETRAINED': { + 'value': True, + 'description': 'Use the pretrained model or not!' + } + } + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + import timm + num_classes = cfg.NUM_CLASSES + model_name = cfg.MODEL_NAME + pretrained = cfg.get('PRETRAINED', False) + self.visual = timm.create_model(model_name, + pretrained=pretrained, + num_classes=num_classes) + + def forward(self, x): + out = self.visual.forward(x) + return out + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BACKBONES', + __class__.__name__, + TIMM_MODEL.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/image/utils/__init__.py b/scepter/modules/model/backbone/image/utils/__init__.py new file mode 100644 index 0000000..5943878 --- /dev/null +++ b/scepter/modules/model/backbone/image/utils/__init__.py @@ -0,0 +1,3 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.backbone.image.utils import vit diff --git a/scepter/modules/model/backbone/image/utils/clip.py b/scepter/modules/model/backbone/image/utils/clip.py new file mode 100644 index 0000000..702be1d --- /dev/null +++ b/scepter/modules/model/backbone/image/utils/clip.py @@ -0,0 +1,565 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from collections import OrderedDict +from typing import List, Tuple, Union + +import numpy as np +import torch +import torch.nn.functional as F +from pkg_resources import packaging +from torch import nn + +from scepter.modules.model.backbone.image.utils.simple_tokenizer import \ + SimpleTokenizer as _Tokenizer + +_tokenizer = _Tokenizer() + + +class Bottleneck(nn.Module): + expansion = 4 + + def __init__(self, inplanes, planes, stride=1): + super().__init__() + + # all conv layers have stride 1. an avgpool is performed after the second convolution when stride > 1 + self.conv1 = nn.Conv2d(inplanes, planes, 1, bias=False) + self.bn1 = nn.BatchNorm2d(planes) + self.relu1 = nn.ReLU(inplace=True) + + self.conv2 = nn.Conv2d(planes, planes, 3, padding=1, bias=False) + self.bn2 = nn.BatchNorm2d(planes) + self.relu2 = nn.ReLU(inplace=True) + + self.avgpool = nn.AvgPool2d(stride) if stride > 1 else nn.Identity() + + self.conv3 = nn.Conv2d(planes, planes * self.expansion, 1, bias=False) + self.bn3 = nn.BatchNorm2d(planes * self.expansion) + self.relu3 = nn.ReLU(inplace=True) + + self.downsample = None + self.stride = stride + + if stride > 1 or inplanes != planes * Bottleneck.expansion: + # downsampling layer is prepended with an avgpool, and the subsequent convolution has stride 1 + self.downsample = nn.Sequential( + OrderedDict([('-1', nn.AvgPool2d(stride)), + ('0', + nn.Conv2d(inplanes, + planes * self.expansion, + 1, + stride=1, + bias=False)), + ('1', nn.BatchNorm2d(planes * self.expansion))])) + + def forward(self, x: torch.Tensor): + identity = x + + out = self.relu1(self.bn1(self.conv1(x))) + out = self.relu2(self.bn2(self.conv2(out))) + out = self.avgpool(out) + out = self.bn3(self.conv3(out)) + + if self.downsample is not None: + identity = self.downsample(x) + + out += identity + out = self.relu3(out) + return out + + +class AttentionPool2d(nn.Module): + def __init__(self, + spacial_dim: int, + embed_dim: int, + num_heads: int, + output_dim: int = None): + super().__init__() + self.positional_embedding = nn.Parameter( + torch.randn(spacial_dim**2 + 1, embed_dim) / embed_dim**0.5) + self.k_proj = nn.Linear(embed_dim, embed_dim) + self.q_proj = nn.Linear(embed_dim, embed_dim) + self.v_proj = nn.Linear(embed_dim, embed_dim) + self.c_proj = nn.Linear(embed_dim, output_dim or embed_dim) + self.num_heads = num_heads + + def forward(self, x): + x = x.flatten(start_dim=2).permute(2, 0, 1) # NCHW -> (HW)NC + x = torch.cat([x.mean(dim=0, keepdim=True), x], dim=0) # (HW+1)NC + x = x + self.positional_embedding[:, None, :].to(x.dtype) # (HW+1)NC + x, _ = F.multi_head_attention_forward( + query=x[:1], + key=x, + value=x, + embed_dim_to_check=x.shape[-1], + num_heads=self.num_heads, + q_proj_weight=self.q_proj.weight, + k_proj_weight=self.k_proj.weight, + v_proj_weight=self.v_proj.weight, + in_proj_weight=None, + in_proj_bias=torch.cat( + [self.q_proj.bias, self.k_proj.bias, self.v_proj.bias]), + bias_k=None, + bias_v=None, + add_zero_attn=False, + dropout_p=0, + out_proj_weight=self.c_proj.weight, + out_proj_bias=self.c_proj.bias, + use_separate_proj_weight=True, + training=self.training, + need_weights=False) + return x.squeeze(0) + + +class ModifiedResNet(nn.Module): + """ + A ResNet class that is similar to torchvision's but contains the following changes: + - There are now 3 "stem" convolutions as opposed to 1, with an average pool instead of a max pool. + - Performs anti-aliasing strided convolutions, where an avgpool is prepended to convolutions with stride > 1 + - The final pooling layer is a QKV attention instead of an average pool + """ + def __init__(self, + layers, + output_dim, + heads, + input_resolution=224, + width=64): + super().__init__() + self.output_dim = output_dim + self.input_resolution = input_resolution + + # the 3-layer stem + self.conv1 = nn.Conv2d(3, + width // 2, + kernel_size=3, + stride=2, + padding=1, + bias=False) + self.bn1 = nn.BatchNorm2d(width // 2) + self.relu1 = nn.ReLU(inplace=True) + self.conv2 = nn.Conv2d(width // 2, + width // 2, + kernel_size=3, + padding=1, + bias=False) + self.bn2 = nn.BatchNorm2d(width // 2) + self.relu2 = nn.ReLU(inplace=True) + self.conv3 = nn.Conv2d(width // 2, + width, + kernel_size=3, + padding=1, + bias=False) + self.bn3 = nn.BatchNorm2d(width) + self.relu3 = nn.ReLU(inplace=True) + self.avgpool = nn.AvgPool2d(2) + + # residual layers + self._inplanes = width # this is a *mutable* variable used during construction + self.layer1 = self._make_layer(width, layers[0]) + self.layer2 = self._make_layer(width * 2, layers[1], stride=2) + self.layer3 = self._make_layer(width * 4, layers[2], stride=2) + self.layer4 = self._make_layer(width * 8, layers[3], stride=2) + + embed_dim = width * 32 # the ResNet feature dimension + self.attnpool = AttentionPool2d(input_resolution // 32, embed_dim, + heads, output_dim) + + def _make_layer(self, planes, blocks, stride=1): + layers = [Bottleneck(self._inplanes, planes, stride)] + + self._inplanes = planes * Bottleneck.expansion + for _ in range(1, blocks): + layers.append(Bottleneck(self._inplanes, planes)) + + return nn.Sequential(*layers) + + def forward(self, x): + def stem(x): + x = self.relu1(self.bn1(self.conv1(x))) + x = self.relu2(self.bn2(self.conv2(x))) + x = self.relu3(self.bn3(self.conv3(x))) + x = self.avgpool(x) + return x + + x = x.type(self.conv1.weight.dtype) + x = stem(x) + x = self.layer1(x) + x = self.layer2(x) + x = self.layer3(x) + x = self.layer4(x) + x = self.attnpool(x) + + return x + + +class LayerNorm(nn.LayerNorm): + """Subclass torch's LayerNorm to handle fp16.""" + def forward(self, x: torch.Tensor): + orig_type = x.dtype + ret = super().forward(x.type(torch.float32)) + return ret.type(orig_type) + + +class QuickGELU(nn.Module): + def forward(self, x: torch.Tensor): + return x * torch.sigmoid(1.702 * x) + + +class ResidualAttentionBlock(nn.Module): + def __init__(self, + d_model: int, + n_head: int, + attn_mask: torch.Tensor = None): + super().__init__() + + self.attn = nn.MultiheadAttention(d_model, n_head) + self.ln_1 = LayerNorm(d_model) + self.mlp = nn.Sequential( + OrderedDict([('c_fc', nn.Linear(d_model, d_model * 4)), + ('gelu', QuickGELU()), + ('c_proj', nn.Linear(d_model * 4, d_model))])) + self.ln_2 = LayerNorm(d_model) + self.attn_mask = attn_mask + + def attention(self, x: torch.Tensor): + self.attn_mask = self.attn_mask.to( + dtype=x.dtype, + device=x.device) if self.attn_mask is not None else None + return self.attn(x, x, x, need_weights=False, + attn_mask=self.attn_mask)[0] + + def forward(self, x: torch.Tensor): + x = x + self.attention(self.ln_1(x)) + x = x + self.mlp(self.ln_2(x)) + return x + + +class Transformer(nn.Module): + def __init__(self, + width: int, + layers: int, + heads: int, + attn_mask: torch.Tensor = None): + super().__init__() + self.width = width + self.layers = layers + self.resblocks = nn.Sequential(*[ + ResidualAttentionBlock(width, heads, attn_mask) + for _ in range(layers) + ]) + + def forward(self, x: torch.Tensor): + return self.resblocks(x) + + +class VisionTransformer(nn.Module): + def __init__(self, input_resolution: int, patch_size: int, width: int, + layers: int, heads: int, output_dim: int): + super().__init__() + self.input_resolution = input_resolution + self.output_dim = output_dim + self.conv1 = nn.Conv2d(in_channels=3, + out_channels=width, + kernel_size=patch_size, + stride=patch_size, + bias=False) + + scale = width**-0.5 + self.class_embedding = nn.Parameter(scale * torch.randn(width)) + self.positional_embedding = nn.Parameter(scale * torch.randn( + (input_resolution // patch_size)**2 + 1, width)) + self.ln_pre = LayerNorm(width) + + self.transformer = Transformer(width, layers, heads) + + self.ln_post = LayerNorm(width) + self.proj = nn.Parameter(scale * torch.randn(width, output_dim)) + + def forward(self, x: torch.Tensor): + x = self.conv1(x) # shape = [*, width, grid, grid] + x = x.reshape(x.shape[0], x.shape[1], + -1) # shape = [*, width, grid ** 2] + x = x.permute(0, 2, 1) # shape = [*, grid ** 2, width] + x = torch.cat([ + self.class_embedding.to(x.dtype) + torch.zeros( + x.shape[0], 1, x.shape[-1], dtype=x.dtype, device=x.device), x + ], + dim=1) # shape = [*, grid ** 2 + 1, width] + x = x + self.positional_embedding.to(x.dtype) + x = self.ln_pre(x) + + x = x.permute(1, 0, 2) # NLD -> LND + x = self.transformer(x) + x = x.permute(1, 0, 2) # LND -> NLD + + x = self.ln_post(x[:, 0, :]) + + if self.proj is not None: + x = x @ self.proj + + return x + + +class CLIP(nn.Module): + def __init__( + self, + embed_dim: int, + # vision + image_resolution: int, + vision_layers: Union[Tuple[int, int, int, int], int], + vision_width: int, + vision_patch_size: int, + # text + context_length: int, + vocab_size: int, + transformer_width: int, + transformer_heads: int, + transformer_layers: int): + super().__init__() + + self.context_length = context_length + + if isinstance(vision_layers, (tuple, list)): + vision_heads = vision_width * 32 // 64 + self.visual = ModifiedResNet(layers=vision_layers, + output_dim=embed_dim, + heads=vision_heads, + input_resolution=image_resolution, + width=vision_width) + else: + vision_heads = vision_width // 64 + self.visual = VisionTransformer(input_resolution=image_resolution, + patch_size=vision_patch_size, + width=vision_width, + layers=vision_layers, + heads=vision_heads, + output_dim=embed_dim) + + self.transformer = Transformer(width=transformer_width, + layers=transformer_layers, + heads=transformer_heads, + attn_mask=self.build_attention_mask()) + + self.vocab_size = vocab_size + self.token_embedding = nn.Embedding(vocab_size, transformer_width) + self.positional_embedding = nn.Parameter( + torch.empty(self.context_length, transformer_width)) + self.ln_final = LayerNorm(transformer_width) + + self.text_projection = nn.Parameter( + torch.empty(transformer_width, embed_dim)) + self.logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07)) + + self.initialize_parameters() + + def initialize_parameters(self): + nn.init.normal_(self.token_embedding.weight, std=0.02) + nn.init.normal_(self.positional_embedding, std=0.01) + + if isinstance(self.visual, ModifiedResNet): + if self.visual.attnpool is not None: + std = self.visual.attnpool.c_proj.in_features**-0.5 + nn.init.normal_(self.visual.attnpool.q_proj.weight, std=std) + nn.init.normal_(self.visual.attnpool.k_proj.weight, std=std) + nn.init.normal_(self.visual.attnpool.v_proj.weight, std=std) + nn.init.normal_(self.visual.attnpool.c_proj.weight, std=std) + + for resnet_block in [ + self.visual.layer1, self.visual.layer2, self.visual.layer3, + self.visual.layer4 + ]: + for name, param in resnet_block.named_parameters(): + if name.endswith('bn3.weight'): + nn.init.zeros_(param) + + proj_std = (self.transformer.width**-0.5) * ( + (2 * self.transformer.layers)**-0.5) + attn_std = self.transformer.width**-0.5 + fc_std = (2 * self.transformer.width)**-0.5 + for block in self.transformer.resblocks: + nn.init.normal_(block.attn.in_proj_weight, std=attn_std) + nn.init.normal_(block.attn.out_proj.weight, std=proj_std) + nn.init.normal_(block.mlp.c_fc.weight, std=fc_std) + nn.init.normal_(block.mlp.c_proj.weight, std=proj_std) + + if self.text_projection is not None: + nn.init.normal_(self.text_projection, + std=self.transformer.width**-0.5) + + def build_attention_mask(self): + # lazily create causal attention mask, with full attention between the vision tokens + # pytorch uses additive attention mask; fill with -inf + mask = torch.empty(self.context_length, self.context_length) + mask.fill_(float('-inf')) + mask.triu_(1) # zero out the lower diagonal + return mask + + @property + def dtype(self): + return self.visual.conv1.weight.dtype + + def encode_image(self, image): + return self.visual(image.type(self.dtype)) + + def encode_text(self, text): + x = self.token_embedding(text).type( + self.dtype) # [batch_size, n_ctx, d_model] + + x = x + self.positional_embedding.type(self.dtype) + x = x.permute(1, 0, 2) # NLD -> LND + x = self.transformer(x) + x = x.permute(1, 0, 2) # LND -> NLD + x = self.ln_final(x).type(self.dtype) + + # x.shape = [batch_size, n_ctx, transformer.width] + # take features from the eot embedding (eot_token is the highest number in each sequence) + x = x[torch.arange(x.shape[0]), + text.argmax(dim=-1)] @ self.text_projection + + return x + + def forward(self, image, text): + image_features = self.encode_image(image) + text_features = self.encode_text(text) + + # normalized features + image_features = image_features / image_features.norm(dim=1, + keepdim=True) + text_features = text_features / text_features.norm(dim=1, keepdim=True) + + # cosine similarity as logits + logit_scale = self.logit_scale.exp() + logits_per_image = logit_scale * image_features @ text_features.t() + logits_per_text = logits_per_image.t() + + # shape = [global_batch_size, global_batch_size] + return logits_per_image, logits_per_text + + +def tokenize( + texts: Union[str, List[str]], + context_length: int = 77, + truncate: bool = False) -> Union[torch.IntTensor, torch.LongTensor]: + """ + Returns the tokenized representation of given input string(s) + + Parameters + ---------- + texts : Union[str, List[str]] + An input string or a list of input strings to tokenize + + context_length : int + The context length to use; all CLIP model use 77 as the context length + + truncate: bool + Whether to truncate the text in case its encoding is longer than the context length + + Returns + ------- + A two-dimensional tensor containing the resulting tokens, shape = [number of input strings, context_length]. + We return LongTensor when torch version is <1.8.0, since older index_select requires indices to be long. + """ + if isinstance(texts, str): + texts = [texts] + + sot_token = _tokenizer.encoder['<|startoftext|>'] + eot_token = _tokenizer.encoder['<|endoftext|>'] + all_tokens = [[sot_token] + _tokenizer.encode(text) + [eot_token] + for text in texts] + if packaging.version.parse( + torch.__version__) < packaging.version.parse('1.8.0'): + result = torch.zeros(len(all_tokens), context_length, dtype=torch.long) + else: + result = torch.zeros(len(all_tokens), context_length, dtype=torch.int) + + for i, tokens in enumerate(all_tokens): + if len(tokens) > context_length: + if truncate: + tokens = tokens[:context_length] + tokens[-1] = eot_token + else: + raise RuntimeError( + f'Input {texts[i]} is too long for context length {context_length}' + ) + result[i, :len(tokens)] = torch.tensor(tokens) + + return result + + +def convert_weights(model: nn.Module): + """Convert applicable model parameters to fp16""" + def _convert_weights_to_fp16(layer): + if isinstance(layer, (nn.Conv1d, nn.Conv2d, nn.Linear)): + layer.weight.data = layer.weight.data.half() + if layer.bias is not None: + layer.bias.data = layer.bias.data.half() + + if isinstance(layer, nn.MultiheadAttention): + for attr in [ + *[f'{s}_proj_weight' for s in ['in', 'q', 'k', 'v']], + 'in_proj_bias', 'bias_k', 'bias_v' + ]: + tensor = getattr(layer, attr) + if tensor is not None: + tensor.data = tensor.data.half() + + for name in ['text_projection', 'proj']: + if hasattr(layer, name): + attr = getattr(layer, name) + if attr is not None: + attr.data = attr.data.half() + + model.apply(_convert_weights_to_fp16) + + +def build_model(state_dict: dict): + vit = 'visual.proj' in state_dict + + if vit: + vision_width = state_dict['visual.conv1.weight'].shape[0] + vision_layers = len([ + k for k in state_dict.keys() + if k.startswith('visual.') and k.endswith('.attn.in_proj_weight') + ]) + vision_patch_size = state_dict['visual.conv1.weight'].shape[-1] + grid_size = round( + (state_dict['visual.positional_embedding'].shape[0] - 1)**0.5) + image_resolution = vision_patch_size * grid_size + else: + counts: list = [ + len( + set( + k.split('.')[2] for k in state_dict + if k.startswith(f'visual.layer{b}'))) + for b in [1, 2, 3, 4] + ] + vision_layers = tuple(counts) + vision_width = state_dict['visual.layer1.0.conv1.weight'].shape[0] + output_width = round( + (state_dict['visual.attnpool.positional_embedding'].shape[0] - + 1)**0.5) + vision_patch_size = None + assert output_width**2 + 1 == state_dict[ + 'visual.attnpool.positional_embedding'].shape[0] + image_resolution = output_width * 32 + + embed_dim = state_dict['text_projection'].shape[1] + context_length = state_dict['positional_embedding'].shape[0] + vocab_size = state_dict['token_embedding.weight'].shape[0] + transformer_width = state_dict['ln_final.weight'].shape[0] + transformer_heads = transformer_width // 64 + transformer_layers = len( + set( + k.split('.')[2] for k in state_dict + if k.startswith('transformer.resblocks'))) + + model = CLIP(embed_dim, image_resolution, vision_layers, vision_width, + vision_patch_size, context_length, vocab_size, + transformer_width, transformer_heads, transformer_layers) + + for key in ['input_resolution', 'context_length', 'vocab_size']: + if key in state_dict: + del state_dict[key] + + convert_weights(model) + model.load_state_dict(state_dict) + return model.eval() diff --git a/scepter/modules/model/backbone/image/utils/simple_tokenizer.py b/scepter/modules/model/backbone/image/utils/simple_tokenizer.py new file mode 100644 index 0000000..3804d47 --- /dev/null +++ b/scepter/modules/model/backbone/image/utils/simple_tokenizer.py @@ -0,0 +1,150 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import gzip +import html +import os +from functools import lru_cache + +import ftfy +import regex as re + + +@lru_cache() +def default_bpe(): + return os.path.join(os.path.dirname(os.path.abspath(__file__)), + 'bpe_simple_vocab_16e6.txt.gz') + + +@lru_cache() +def bytes_to_unicode(): + """ + Returns list of utf-8 byte and a corresponding list of unicode strings. + The reversible bpe codes work on unicode strings. + This means you need a large # of unicode characters in your vocab if you want to avoid UNKs. + When you're at something like a 10B token dataset you end up needing around 5K for decent coverage. + This is a signficant percentage of your normal, say, 32K bpe vocab. + To avoid that, we want lookup tables between utf-8 bytes and unicode strings. + And avoids mapping to whitespace/control characters the bpe code barfs on. + """ + bs = list(range(ord('!'), + ord('~') + 1)) + list(range( + ord('¡'), + ord('¬') + 1)) + list(range(ord('®'), + ord('ÿ') + 1)) + cs = bs[:] + n = 0 + for b in range(2**8): + if b not in bs: + bs.append(b) + cs.append(2**8 + n) + n += 1 + cs = [chr(n) for n in cs] + return dict(zip(bs, cs)) + + +def get_pairs(word): + """Return set of symbol pairs in a word. + Word is represented as tuple of symbols (symbols being variable-length strings). + """ + pairs = set() + prev_char = word[0] + for char in word[1:]: + pairs.add((prev_char, char)) + prev_char = char + return pairs + + +def basic_clean(text): + text = ftfy.fix_text(text) + text = html.unescape(html.unescape(text)) + return text.strip() + + +def whitespace_clean(text): + text = re.sub(r'\s+', ' ', text) + text = text.strip() + return text + + +class SimpleTokenizer(object): + def __init__(self, bpe_path: str = default_bpe()): + self.byte_encoder = bytes_to_unicode() + self.byte_decoder = {v: k for k, v in self.byte_encoder.items()} + merges = gzip.open(bpe_path).read().decode('utf-8').split('\n') + merges = merges[1:49152 - 256 - 2 + 1] + merges = [tuple(merge.split()) for merge in merges] + vocab = list(bytes_to_unicode().values()) + vocab = vocab + [v + '' for v in vocab] + for merge in merges: + vocab.append(''.join(merge)) + vocab.extend(['<|startoftext|>', '<|endoftext|>']) + self.encoder = dict(zip(vocab, range(len(vocab)))) + self.decoder = {v: k for k, v in self.encoder.items()} + self.bpe_ranks = dict(zip(merges, range(len(merges)))) + self.cache = { + '<|startoftext|>': '<|startoftext|>', + '<|endoftext|>': '<|endoftext|>' + } + self.pat = re.compile( + r"""<\|startoftext\|>|<\|endoftext\|>|'s|'t|'re|'ve|'m|'ll|'d|[\p{L}]+|[\p{N}]|[^\s\p{L}\p{N}]+""", + re.IGNORECASE) + + def bpe(self, token): + if token in self.cache: + return self.cache[token] + word = tuple(token[:-1]) + (token[-1] + '', ) + pairs = get_pairs(word) + + if not pairs: + return token + '' + + while True: + bigram = min( + pairs, key=lambda pair: self.bpe_ranks.get(pair, float('inf'))) + if bigram not in self.bpe_ranks: + break + first, second = bigram + new_word = [] + i = 0 + while i < len(word): + try: + j = word.index(first, i) + new_word.extend(word[i:j]) + i = j + except Exception: + new_word.extend(word[i:]) + break + + if word[i] == first and i < len(word) - 1 and word[ + i + 1] == second: + new_word.append(first + second) + i += 2 + else: + new_word.append(word[i]) + i += 1 + new_word = tuple(new_word) + word = new_word + if len(word) == 1: + break + else: + pairs = get_pairs(word) + word = ' '.join(word) + self.cache[token] = word + return word + + def encode(self, text): + bpe_tokens = [] + text = whitespace_clean(basic_clean(text)).lower() + for token in re.findall(self.pat, text): + token = ''.join(self.byte_encoder[b] + for b in token.encode('utf-8')) + bpe_tokens.extend(self.encoder[bpe_token] + for bpe_token in self.bpe(token).split(' ')) + return bpe_tokens + + def decode(self, tokens): + text = ''.join([self.decoder[token] for token in tokens]) + text = bytearray([self.byte_decoder[c] for c in text + ]).decode('utf-8', + errors='replace').replace('', ' ') + return text diff --git a/scepter/modules/model/backbone/image/utils/vit.py b/scepter/modules/model/backbone/image/utils/vit.py new file mode 100644 index 0000000..f2ff923 --- /dev/null +++ b/scepter/modules/model/backbone/image/utils/vit.py @@ -0,0 +1,590 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from collections import OrderedDict + +import torch +from torch import nn + + +class LayerNorm(nn.LayerNorm): + """Subclass torch's LayerNorm to handle fp16.""" + def forward(self, x: torch.Tensor): + orig_type = x.dtype + ret = super().forward(x.type(torch.float32)) + return ret.type(orig_type) + + +class QuickGELU(nn.Module): + def forward(self, x: torch.Tensor): + return x * torch.sigmoid(1.702 * x) + + +class ResidualAttentionBlock(nn.Module): + def __init__(self, + d_model: int, + n_head: int, + attn_mask: torch.Tensor = None): + super().__init__() + + self.attn = nn.MultiheadAttention(d_model, n_head) + self.ln_1 = LayerNorm(d_model) + self.mlp = nn.Sequential( + OrderedDict([('c_fc', nn.Linear(d_model, d_model * 4)), + ('gelu', QuickGELU()), + ('c_proj', nn.Linear(d_model * 4, d_model))])) + self.ln_2 = LayerNorm(d_model) + self.attn_mask = attn_mask + + def attention(self, x: torch.Tensor): + self.attn_mask = self.attn_mask.to( + dtype=x.dtype, + device=x.device) if self.attn_mask is not None else None + return self.attn(x, x, x, need_weights=False, + attn_mask=self.attn_mask)[0] + + def forward(self, x: torch.Tensor): + x = x + self.attention(self.ln_1(x)) + x = x + self.mlp(self.ln_2(x)) + return x + + +class FrozenTransformer(nn.Module): + def __init__(self, + width: int, + layers: int, + heads: int, + attn_mask: torch.Tensor = None): + super().__init__() + self.width = width + self.layers = layers + self.resblocks = nn.Sequential(*[ + ResidualAttentionBlock(width, heads, attn_mask) + for _ in range(layers) + ]) + + def forward(self, x: torch.Tensor): + return self.resblocks(x) + + def train(self, mode: bool = True): + self.training = False + for module in self.children(): + module.train(False) + return self + + +class Transformer(nn.Module): + def __init__(self, + width: int, + layers: int, + heads: int, + attn_mask: torch.Tensor = None): + super().__init__() + self.width = width + self.layers = layers + self.resblocks = nn.Sequential(*[ + ResidualAttentionBlock(width, heads, attn_mask) + for _ in range(layers) + ]) + + def forward(self, x: torch.Tensor): + return self.resblocks(x) + + +class VIT(nn.Module): + para_dict = { + 'INPUT_RESOLUTION': { + 'value': 224, + 'description': 'The input resolution of vit model!' + }, + 'PATCH_SIZE': { + 'value': 32, + 'description': 'The patch size of vit model!' + }, + 'WIDTH': { + 'value': 768, + 'description': 'The input embbeding dimention!' + }, + 'OUTPUT_DIM': { + 'value': 512, + 'description': 'The output embbeding dimention!' + }, + 'LAYERS': { + 'value': 12, + 'description': "Model's all layers num!" + }, + 'HEADS': { + 'value': 12, + 'description': 'The head number of transformer!' + }, + 'EXPORT': { + 'value': False, + 'description': 'Whether export model or not!' + }, + 'TOKEN_WISE': { + 'value': False, + 'description': 'Whether output token wise feature or not!' + } + } + + def __init__(self, cfg, logger=None): + super().__init__() + input_resolution = cfg.INPUT_RESOLUTION + width = cfg.WIDTH + patch_size = cfg.PATCH_SIZE + layers = cfg.LAYERS + heads = cfg.HEADS + output_dim = cfg.OUTPUT_DIM + use_proj = cfg.get('USE_PROJ', True) + self.export = cfg.get('EXPORT', False) + self.token_wise = cfg.get('TOKEN_WISE', False) + self.conv1 = nn.Conv2d(in_channels=3, + out_channels=width, + kernel_size=patch_size, + stride=patch_size, + bias=False) + scale = width**-0.5 + self.class_embedding = nn.Parameter(scale * torch.randn(width)) + self.positional_embedding = nn.Parameter(scale * torch.randn( + (input_resolution // patch_size)**2 + 1, width)) + self.ln_pre = LayerNorm(width) + self.transformer = Transformer(width, layers, heads) + + self.ln_post = LayerNorm(width) + if use_proj: + self.proj = nn.Parameter(scale * torch.randn(width, output_dim)) + else: + self.proj = None + + @property + def dtype(self): + return self.conv1.weight.dtype + + def forward(self, x: torch.Tensor): + x = self.conv1(x.type(self.dtype)) # shape = [*, width, grid, grid] + x = x.reshape(x.shape[0], x.shape[1], + -1) # shape = [*, width, grid ** 2] + x = x.permute(0, 2, 1) # shape = [*, grid ** 2, width] + # x = torch.cat([self.class_embedding.to(x.dtype) + torch.zeros(x.shape[0], + # 1, x.shape[-1], dtype=x.dtype, device=x.device), x], dim=1) # shape = [*, grid ** 2 + 1, width] + if self.export: + x = torch.cat([ + self.class_embedding.to(x.dtype).view(x.shape[0], 1, + x.shape[-1]), x + ], + dim=1) # shape = [*, grid ** 2 + 1, width] + else: + x = torch.cat([ + self.class_embedding.to(x.dtype) + torch.zeros( + x.shape[0], 1, x.shape[-1], dtype=x.dtype, + device=x.device), x + ], + dim=1) # shape = [*, grid ** 2 + 1, width] + x = x + self.positional_embedding.to(x.dtype) + x = self.ln_pre(x) + + x = x.permute(1, 0, 2) # NLD -> LND + x = self.transformer(x) + x = x.permute(1, 0, 2) # LND -> NLD + if self.token_wise: + return self.ln_post(x) + x = self.ln_post(x[:, 0, :]) + if not self.export: + if self.proj is not None: + x = x @ self.proj + return x + else: + before_proj = x + if self.proj is not None: + x = before_proj @ self.proj + return before_proj, x + + +class VIT_MODEL(nn.Module): + ''' + INPUT_RESOLUTION: 224 + PATCH_SIZE: 32 + WIDTH: 768 + OUTPUT_DIM: 512 + LAYERS: 12 + HEADS: 12 + ''' + para_dict = { + 'INPUT_RESOLUTION': { + 'value': 224, + 'description': 'The input resolution of vit model!' + }, + 'PATCH_SIZE': { + 'value': 32, + 'description': 'The patch size of vit model!' + }, + 'WIDTH': { + 'value': 768, + 'description': 'The input embbeding dimention!' + }, + 'OUTPUT_DIM': { + 'value': 512, + 'description': 'The output embbeding dimention!' + }, + 'FROZEN_LAYERS': { + 'value': 6, + 'description': "Frozen model's layers num!" + }, + 'FT_LAYERS': { + 'value': 6, + 'description': "Finetune model's layers num!" + }, + 'HEADS': { + 'value': 12, + 'description': 'The head number of transformer!' + } + } + + def __init__(self, cfg): + super().__init__() + input_resolution = cfg.INPUT_RESOLUTION + patch_size = cfg.PATCH_SIZE + width = cfg.WIDTH + output_dim = cfg.OUTPUT_DIM + frozen_layers = cfg.FROZEN_LAYERS + ft_layers = cfg.FT_LAYERS + heads = cfg.HEADS + + self.input_resolution = input_resolution + self.output_dim = output_dim + self.conv1 = nn.Conv2d(in_channels=3, + out_channels=width, + kernel_size=patch_size, + stride=patch_size, + bias=False) + + scale = width**-0.5 + self.class_embedding = nn.Parameter(scale * + torch.randn(width)) # [768] + self.positional_embedding = nn.Parameter(scale * torch.randn( + (input_resolution // patch_size)**2 + 1, width)) # [50, 768] + self.ln_pre = LayerNorm(width) + + self.frozen_transformer = FrozenTransformer(width, frozen_layers, + heads) + self.frozen_transformer.eval() + + self.ft_transformer = Transformer(width, ft_layers, heads) + + self.ln_post = LayerNorm(width) + self.proj = nn.Parameter(scale * torch.randn(width, output_dim)) + + def forward(self, x: torch.Tensor): + with torch.no_grad(): + x = self.conv1( + x) # shape = [*, width, grid, grid] -> [1, 768, 7, 7] + x = x.reshape( + x.shape[0], x.shape[1], + -1) # shape = [*, width, grid ** 2] -> [1, 768, 49] 49token + x = x.permute(0, 2, + 1) # shape = [*, grid ** 2, width] -> [1, 49, 768] + x = torch.cat([ + self.class_embedding.to(x.dtype) + torch.zeros( + x.shape[0], 1, x.shape[-1], dtype=x.dtype, + device=x.device), x + ], + dim=1 + ) # shape = [*, grid ** 2 + 1, width]-> [1, 50, 768] + x = x + self.positional_embedding.to(x.dtype) # [1, 50, 768] + x = self.ln_pre(x) + x = x.permute(1, 0, 2) # NLD -> LND [50, 1, 768] + x = self.frozen_transformer(x) + x = self.ft_transformer(x) + x = x.permute(1, 0, 2) # LND -> NLD + + x = self.ln_post(x[:, 0, :]) + + if self.proj is not None: + x = x @ self.proj + + return x + + +class MULTI_HEAD_VIT_MODEL(nn.Module): + para_dict = { + 'INPUT_RESOLUTION': { + 'value': 224, + 'description': 'The input resolution of vit model!' + }, + 'PATCH_SIZE': { + 'value': 32, + 'description': 'The patch size of vit model!' + }, + 'WIDTH': { + 'value': 768, + 'description': 'The input embbeding dimention!' + }, + 'OUTPUT_DIM': { + 'value': 512, + 'description': 'The output embbeding dimention!' + }, + 'LAYERS': { + 'value': 12, + 'description': "All of the vit model's layers!" + }, + 'FROZEN_LAYERS': { + 'value': 6, + 'description': 'The frozen layers number!' + }, + 'FT_LAYERS': { + 'value': 6, + 'description': 'The finetune layers number!' + }, + 'MULTI_HEAD': { + 'value': 2, + 'description': 'The head number of vit!' + }, + 'HEADS': { + 'value': 12, + 'description': 'The head number of transformer!' + } + } + + def __init__(self, cfg): + super().__init__() + input_resolution = cfg.INPUT_RESOLUTION + patch_size = cfg.PATCH_SIZE + width = cfg.WIDTH + output_dim = cfg.OUTPUT_DIM + frozen_layers = cfg.FROZEN_LAYERS + ft_layers = cfg.FT_LAYERS + self.multi_head = cfg.MULTI_HEAD + heads = cfg.HEADS + + self.input_resolution = input_resolution + self.output_dim = output_dim + self.conv1 = nn.Conv2d(in_channels=3, + out_channels=width, + kernel_size=patch_size, + stride=patch_size, + bias=False) + + scale = width**-0.5 + self.class_embedding = nn.Parameter(scale * + torch.randn(width)) # [768] + self.positional_embedding = nn.Parameter(scale * torch.randn( + (input_resolution // patch_size)**2 + 1, width)) # [50, 768] + self.ln_pre = LayerNorm(width) + + self.frozen_transformer = FrozenTransformer(width, frozen_layers, + heads) + self.frozen_transformer.eval() + + if self.multi_head == 2: + self.ft_transformer_1 = Transformer(width, ft_layers, heads) + self.ln_post_1 = LayerNorm(width) + self.proj_1 = nn.Parameter(scale * torch.randn(width, output_dim)) + + self.ft_transformer_2 = Transformer(width, ft_layers, heads) + self.ln_post_2 = LayerNorm(width) + self.proj_2 = nn.Parameter(scale * torch.randn(width, output_dim)) + elif self.multi_head == 3: + self.ft_transformer_1 = Transformer(width, ft_layers, heads) + self.ln_post_1 = LayerNorm(width) + self.proj_1 = nn.Parameter(scale * torch.randn(width, output_dim)) + + self.ft_transformer_2 = Transformer(width, ft_layers, heads) + self.ln_post_2 = LayerNorm(width) + self.proj_2 = nn.Parameter(scale * torch.randn(width, output_dim)) + + self.ft_transformer_3 = Transformer(width, ft_layers, heads) + self.ln_post_3 = LayerNorm(width) + self.proj_3 = nn.Parameter(scale * torch.randn(width, output_dim)) + elif self.multi_head == 4: + self.ft_transformer_1 = Transformer(width, ft_layers, heads) + self.ln_post_1 = LayerNorm(width) + self.proj_1 = nn.Parameter(scale * torch.randn(width, output_dim)) + + self.ft_transformer_2 = Transformer(width, ft_layers, heads) + self.ln_post_2 = LayerNorm(width) + self.proj_2 = nn.Parameter(scale * torch.randn(width, output_dim)) + + self.ft_transformer_3 = Transformer(width, ft_layers, heads) + self.ln_post_3 = LayerNorm(width) + self.proj_3 = nn.Parameter(scale * torch.randn(width, output_dim)) + + self.ft_transformer_4 = Transformer(width, ft_layers, heads) + self.ln_post_4 = LayerNorm(width) + self.proj_4 = nn.Parameter(scale * torch.randn(width, output_dim)) + + def forward(self, x: torch.Tensor): + with torch.no_grad(): + x = self.conv1( + x) # shape = [*, width, grid, grid] -> [1, 768, 7, 7] + x = x.reshape( + x.shape[0], x.shape[1], + -1) # shape = [*, width, grid ** 2] -> [1, 768, 49] 49token + x = x.permute(0, 2, + 1) # shape = [*, grid ** 2, width] -> [1, 49, 768] + x = torch.cat([ + self.class_embedding.to(x.dtype) + torch.zeros( + x.shape[0], 1, x.shape[-1], dtype=x.dtype, + device=x.device), x + ], + dim=1 + ) # shape = [*, grid ** 2 + 1, width]-> [1, 50, 768] + x = x + self.positional_embedding.to(x.dtype) # [1, 50, 768] + x = self.ln_pre(x) + x = x.permute(1, 0, 2) # NLD -> LND [50, 1, 768] + x = self.frozen_transformer(x) + + if self.multi_head == 2: + sub_x1 = self.ft_transformer_1(x) + sub_x1 = sub_x1.permute(1, 0, 2) # LND -> NLD + sub_x1 = self.ln_post_1(sub_x1[:, 0, :]) + sub_x1 = sub_x1 @ self.proj_1 + + sub_x2 = self.ft_transformer_2(x) + sub_x2 = sub_x2.permute(1, 0, 2) # LND -> NLD + sub_x2 = self.ln_post_2(sub_x2[:, 0, :]) + sub_x2 = sub_x2 @ self.proj_2 + + return sub_x1, sub_x2 + elif self.multi_head == 3: + sub_x1 = self.ft_transformer_1(x) + sub_x1 = sub_x1.permute(1, 0, 2) # LND -> NLD + sub_x1 = self.ln_post_1(sub_x1[:, 0, :]) + sub_x1 = sub_x1 @ self.proj_1 + + sub_x2 = self.ft_transformer_2(x) + sub_x2 = sub_x2.permute(1, 0, 2) # LND -> NLD + sub_x2 = self.ln_post_2(sub_x2[:, 0, :]) + sub_x2 = sub_x2 @ self.proj_2 + + sub_x3 = self.ft_transformer_3(x) + sub_x3 = sub_x3.permute(1, 0, 2) # LND -> NLD + sub_x3 = self.ln_post_3(sub_x3[:, 0, :]) + sub_x3 = sub_x3 @ self.proj_3 + + return sub_x1, sub_x2, sub_x3 + elif self.multi_head == 4: + sub_x1 = self.ft_transformer_1(x) + sub_x1 = sub_x1.permute(1, 0, 2) # LND -> NLD + sub_x1 = self.ln_post_1(sub_x1[:, 0, :]) + sub_x1 = sub_x1 @ self.proj_1 + + sub_x2 = self.ft_transformer_2(x) + sub_x2 = sub_x2.permute(1, 0, 2) # LND -> NLD + sub_x2 = self.ln_post_2(sub_x2[:, 0, :]) + sub_x2 = sub_x2 @ self.proj_2 + + sub_x3 = self.ft_transformer_3(x) + sub_x3 = sub_x3.permute(1, 0, 2) # LND -> NLD + sub_x3 = self.ln_post_3(sub_x3[:, 0, :]) + sub_x3 = sub_x3 @ self.proj_3 + + sub_x4 = self.ft_transformer_4(x) + sub_x4 = sub_x4.permute(1, 0, 2) # LND -> NLD + sub_x4 = self.ln_post_4(sub_x4[:, 0, :]) + sub_x4 = sub_x4 @ self.proj_4 + + return sub_x1, sub_x2, sub_x3, sub_x4 + + +class MULTI_HEAD_VIT_MODEL_Split(nn.Module): + para_dict = { + 'INPUT_RESOLUTION': { + 'value': 224, + 'description': 'The input resolution of vit model!' + }, + 'PATCH_SIZE': { + 'value': 32, + 'description': 'The patch size of vit model!' + }, + 'WIDTH': { + 'value': 768, + 'description': 'The input embbeding dimention!' + }, + 'OUTPUT_DIM': { + 'value': 512, + 'description': 'The output embbeding dimention!' + }, + 'LAYERS': { + 'value': 12, + 'description': "All of the vit model's layers!" + }, + 'FROZEN_LAYERS': { + 'value': 6, + 'description': 'The frozen layers number!' + }, + 'FT_LAYERS': { + 'value': 6, + 'description': 'The finetune layers number!' + }, + 'PART': { + 'value': 'backbone', + 'description': 'The part name of vit!' + }, + 'HEADS': { + 'value': 12, + 'description': 'The head number of transformer!' + } + } + + def __init__(self, cfg): + super().__init__() + input_resolution = cfg.INPUT_RESOLUTION + patch_size = cfg.PATCH_SIZE + width = cfg.WIDTH + output_dim = cfg.OUTPUT_DIM + frozen_layers = cfg.FROZEN_LAYERS + ft_layers = cfg.FT_LAYERS + heads = cfg.HEADS + self.PART = cfg.PART + scale = width**-0.5 + if self.PART == 'backbone': + self.input_resolution = input_resolution + self.output_dim = output_dim + self.conv1 = nn.Conv2d(in_channels=3, + out_channels=width, + kernel_size=patch_size, + stride=patch_size, + bias=False) + self.class_embedding = nn.Parameter(scale * + torch.randn(width)) # [768] + self.positional_embedding = nn.Parameter(scale * torch.randn( + (input_resolution // patch_size)**2 + 1, width)) # [50, 768] + self.ln_pre = LayerNorm(width) + self.frozen_transformer = FrozenTransformer( + width, frozen_layers, heads) + self.frozen_transformer.eval() + else: + self.ft_transformer = Transformer(width, ft_layers, heads) + self.ln_post = LayerNorm(width) + self.proj = nn.Parameter(scale * torch.randn(width, output_dim)) + + def forward(self, x: torch.Tensor): + if self.PART == 'backbone': + with torch.no_grad(): + x = self.conv1( + x) # shape = [*, width, grid, grid] -> [1, 768, 7, 7] + x = x.reshape( + x.shape[0], x.shape[1], -1 + ) # shape = [*, width, grid ** 2] -> [1, 768, 49] 49token + x = x.permute( + 0, 2, 1) # shape = [*, grid ** 2, width] -> [1, 49, 768] + x = torch.cat( + [ + self.class_embedding.to(x.dtype) + + torch.zeros(x.shape[0], + 1, + x.shape[-1], + dtype=x.dtype, + device=x.device), x + ], + dim=1) # shape = [*, grid ** 2 + 1, width]-> [1, 50, 768] + x = x + self.positional_embedding.to(x.dtype) # [1, 50, 768] + x = self.ln_pre(x) + x = x.permute(1, 0, 2) # NLD -> LND [50, 1, 768] + x = self.frozen_transformer(x) + return x + else: + sub_x = self.ft_transformer(x) + sub_x = sub_x.permute(1, 0, 2) # LND -> NLD + sub_x = self.ln_post(sub_x[:, 0, :]) + sub_x = sub_x @ self.proj + return sub_x diff --git a/scepter/modules/model/backbone/image/vit_modify.py b/scepter/modules/model/backbone/image/vit_modify.py new file mode 100644 index 0000000..4e5d84b --- /dev/null +++ b/scepter/modules/model/backbone/image/vit_modify.py @@ -0,0 +1,277 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import torch +import torch.nn as nn + +from scepter.modules.model.backbone.image.utils.vit import ( + MULTI_HEAD_VIT_MODEL, VIT, VIT_MODEL, MULTI_HEAD_VIT_MODEL_Split) +from scepter.modules.model.base_model import BaseModel +from scepter.modules.model.registry import BACKBONES +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.file_system import FS + + +def convert_weights(model: nn.Module): + """Convert applicable model parameters to fp16""" + def _convert_weights_to_fp16(layer): + if isinstance(layer, (nn.Conv1d, nn.Conv2d, nn.Linear)): + layer.weight.data = layer.weight.data.half() + if layer.bias is not None: + layer.bias.data = layer.bias.data.half() + + if isinstance(layer, nn.MultiheadAttention): + for attr in [ + *[f'{s}_proj_weight' for s in ['in', 'q', 'k', 'v']], + 'in_proj_bias', 'bias_k', 'bias_v' + ]: + tensor = getattr(layer, attr) + if tensor is not None: + tensor.data = tensor.data.half() + + for name in ['text_projection', 'proj']: + if hasattr(layer, name): + attr = getattr(layer, name) + if attr is not None: + attr.data = attr.data.half() + + model.apply(_convert_weights_to_fp16) + + +@BACKBONES.register_class() +class VisualTransformer(BaseModel): + ''' + B/16: Input 224 Patch-size 16 Layers 12 Heads 12 WIDTH 768 + B/32: Input 224 Patch-size 32 Layers 12 Heads 12 WIDTH 768 + L/16: Input 224/336 Patch-size 16 Layers 24 Heads 16 WIDTH 1024 + L/14: Input 224/336 Patch-size 14 Layers 24 Heads 16 WIDTH 1024 + L/32: Input 224 Patch-size 32 Layers 24 Heads 16 WIDTH 1024 + H/14: Input ... + INPUT_RESOLUTION: 224 + PATCH_SIZE: 32 + WIDTH: 768 + OUTPUT_DIM: 512 + LAYERS: 12 + HEADS: 12 + ''' + para_dict = { + 'PRETRAIN_PATH': { + 'value': '', + 'description': 'The file path of pretrained model!' + }, + 'PRETRAINED': { + 'value': True, + 'description': 'Use the pretrained model or not!' + } + } + para_dict.update(VIT.para_dict) + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.pretrain_path = cfg.PRETRAIN_PATH + self.pretrained = cfg.PRETRAINED + self.visual = VIT(cfg) + use_proj = cfg.get('USE_PROJ', True) + if self.pretrained: + with FS.get_from(self.pretrain_path, + wait_finish=True) as local_file: + logger.info(f'Loading checkpoint from {self.pretrain_path}') + visual_pre = torch.load(local_file, map_location='cpu') + if not use_proj: + visual_pre.pop('proj') + if visual_pre['conv1.weight'].dtype == torch.float16: + convert_weights(self.visual) + self.visual.load_state_dict(visual_pre, strict=True) + + def forward(self, x): + out = self.visual.forward(x) + return out + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BACKBONES', + __class__.__name__, + VisualTransformer.para_dict, + set_name=True) + + +@BACKBONES.register_class() +class SomeFTVisualTransformer(BaseModel): + ''' + INPUT_RESOLUTION: 224 + PATCH_SIZE: 32 + WIDTH: 768 + OUTPUT_DIM: 512 + LAYERS: 12 + HEADS: 12 + ''' + para_dict = { + 'PRETRAIN_PATH': { + 'value': '', + 'description': 'The file path of pretrained model!' + }, + 'PRETRAINED': { + 'value': True, + 'description': 'Use the pretrained model or not!' + }, + 'FROZEN_LAYERS': { + 'value': 6, + 'description': 'The frozen layers number!' + }, + 'FT_LAYERS': { + 'value': 6, + 'description': 'The finetune layers number!' + } + } + para_dict.update(VIT_MODEL.para_dict) + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.pretrain_path = cfg.PRETRAIN_PATH + self.pretrained = cfg.PRETRAINED + self.visual = VIT_MODEL(cfg) + self.frozen_layers = cfg.FROZEN_LAYERS + self.ft_layers = cfg.FT_LAYERS + if self.pretrained: + with FS.get_from(self.pretrain_path, + wait_finish=True) as local_file: + logger.info(f'Loading checkpoint from {self.pretrain_path}') + visual_pre = torch.load(local_file, map_location='cpu') + state_dict_update = self.reformat_state_dict(visual_pre) + self.visual.load_state_dict(state_dict_update, strict=True) + + def reformat_state_dict(self, state_dict): + state_dict_update = {} + for k, v in state_dict.items(): + if 'transformer.resblocks.' in k: + if int(k.split('.')[2]) < self.frozen_layers: + state_dict_update[k.replace( + 'transformer.resblocks', + 'frozen_transformer.resblocks')] = v + else: + new_k = k.replace('transformer.resblocks', + 'ft_transformer.resblocks') + k_tups = new_k.split('.') + k_tups[2] = str(int(k_tups[2]) - self.frozen_layers) + new_k = '.'.join(k_tups) + state_dict_update[new_k] = v + else: + state_dict_update[k] = v + return state_dict_update + + def forward(self, x): + out = self.visual.forward(x) + return out + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BACKBONES', + __class__.__name__, + SomeFTVisualTransformer.para_dict, + set_name=True) + + +@BACKBONES.register_class() +class MultiHeadSomeFTVisualTransformer(BaseModel): + ''' + INPUT_RESOLUTION: 224 + PATCH_SIZE: 32 + WIDTH: 768 + OUTPUT_DIM: 512 + LAYERS: 12 + HEADS: 12 + ''' + para_dict = {} + para_dict.update(MULTI_HEAD_VIT_MODEL.para_dict) + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.visual = MULTI_HEAD_VIT_MODEL(cfg) + self.multi_head = cfg.MULTI_HEAD + self.frozen_layers = cfg.FROZEN_LAYERS + self.ft_layers = cfg.FT_LAYERS + + def forward(self, x): + out = self.visual.forward(x) + return out + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BACKBONES', + __class__.__name__, + MultiHeadSomeFTVisualTransformer.para_dict, + set_name=True) + + +@BACKBONES.register_class() +class SomeFTVisualTransformerTwoPart(BaseModel): + ''' + INPUT_RESOLUTION: 224 + PATCH_SIZE: 32 + WIDTH: 768 + OUTPUT_DIM: 512 + LAYERS: 12 + HEADS: 12 + ''' + para_dict = {} + para_dict.update(MULTI_HEAD_VIT_MODEL_Split.para_dict) + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.visual = MULTI_HEAD_VIT_MODEL_Split(cfg) + self.frozen_layers = cfg.FROZEN_LAYERS + self.ft_layers = cfg.FT_LAYERS + + def forward(self, x): + out = self.visual.forward(x) + return out + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BACKBONES', + __class__.__name__, + SomeFTVisualTransformerTwoPart.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/unet/__init__.py b/scepter/modules/model/backbone/unet/__init__.py new file mode 100644 index 0000000..99d7eae --- /dev/null +++ b/scepter/modules/model/backbone/unet/__init__.py @@ -0,0 +1,3 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.backbone.unet.unet_module import DiffusionUNet diff --git a/scepter/modules/model/backbone/unet/unet_module.py b/scepter/modules/model/backbone/unet/unet_module.py new file mode 100644 index 0000000..b7e399e --- /dev/null +++ b/scepter/modules/model/backbone/unet/unet_module.py @@ -0,0 +1,816 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import copy +from collections import OrderedDict + +import torch +import torch.nn as nn + +from scepter.modules.model.backbone.unet.unet_utils import ( + Downsample, ResBlock, SpatialTransformer, Timestep, + TimestepEmbedSequential, Upsample, conv_nd, linear, normalization, + timestep_embedding, zero_module) +from scepter.modules.model.base_model import BaseModel +from scepter.modules.model.registry import BACKBONES +from scepter.modules.model.utils.basic_utils import exists +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we +from scepter.modules.utils.file_system import FS + + +def convert_module_to_f16(x): + pass + + +def convert_module_to_f32(x): + pass + + +@BACKBONES.register_class() +class DiffusionUNet(BaseModel): + para_dict = { + 'IN_CHANNELS': { + 'value': + 4, + 'description': + "Unet channels for input, considering the input image's channels." + }, + 'OUT_CHANNELS': { + 'value': + 4, + 'description': + "Unet channels for output, considering the input image's channels." + }, + 'NUM_RES_BLOCKS': { + 'value': 2, + 'description': "The blocks's number of res." + }, + 'MODEL_CHANNELS': { + 'value': 320, + 'description': 'base channel count for the model.' + }, + 'ATTENTION_RESOLUTIONS': { + 'value': [4, 2, 1], + 'description': + 'A collection of downsample rates at which ' + 'attention will take place. May be a set, list,' + ' or tuple. For example, if this contains 4, ' + 'then at 4x downsampling, attentio will be used.' + }, + 'DROPOUT': { + 'value': 0, + 'description': 'The dropout rate.' + }, + 'CHANNEL_MULT': { + 'value': [1, 2, 4, 4], + 'description': 'channel multiplier for each level of the UNet.' + }, + 'CONV_RESAMPLE': { + 'value': True, + 'description': 'Use conv to resample when downsample.' + }, + 'DIMS': { + 'value': 2, + 'description': 'The Conv dims which 2 represent Conv2D.' + }, + 'NUM_CLASSES': { + 'value': + None, + 'description': + 'The class num for class guided setting, also can be set as continuous.' + }, + 'USE_CHECKPOINT': { + 'value': True, + 'description': 'Use gradient checkpointing to reduce memory usage.' + }, + 'USE_FP16': { + 'value': False, + 'description': + 'Set the inference precision whether use FP16 or not.' + }, + 'NUM_HEADS': { + 'value': 8, + 'description': + 'The number of attention head in each attention layer.' + }, + 'NUM_HEADS_CHANNELS': { + 'value': + -1, + 'description': + 'If specified, ignore num_heads and instead use ' + 'a fixed channel width per attention head.' + }, + 'NUM_HEADS_UPSAMPLE': { + 'value': + -1, + 'description': + 'Works with num_heads to set a different number ' + 'of head for upsampling. Deprecated.' + }, + 'USE_SCALE_SHIFT_NORM': { + 'value': + False, + 'description': + 'The scale and shift for the outnorm of RESBLOCK, ' + 'use a FiLM-like conditioning mechanism.' + }, + 'RESBLOCK_UPDOWN': { + 'value': + False, + 'description': + 'Use residual blocks for up/downsampling, if False use Conv.' + }, + 'USE_NEW_ATTENTION_ORDER': { + 'value': + True, + 'description': + 'Whether use new attention(qkv before split head or not) or not.' + }, + 'USE_SPATIAL_TRANSFORMER': { + 'value': + True, + 'description': + 'Custom transformer which support the context, ' + 'if context_dim is not None, the parameter must set True' + }, + 'TRANSFORMER_DEPTH': { + 'value': + 1, + 'description': + "Custom transformer's depth, valid when USE_SPATIAL_TRANSFORMER is True." + }, + 'CONTEXT_DIM': { + 'value': + 768, + 'description': + 'Custom context info, if set, USE_SPATIAL_TRANSFORMER also set True.' + }, + 'N_EMBED': { + 'value': + None, + 'description': + 'Whether support predict_codebook_ids or not, which is the scale of codebook.' + }, + 'LEGACY': { + 'value': + False, + 'description': + 'Whether auto-compute dim_heads according to USE_SPATIAL_TRANSFORMER.' + }, + 'DISABLE_SELF_ATTENTIONS': { + 'value': + None, + 'description': + 'Whether disable the self-attentions on some level, should be a list, [False, True, ...]' + }, + 'NUM_ATTENTION_BLOCKS': { + 'value': None, + 'description': + 'The number of attention blocks for attention layer.' + }, + 'DISABLE_MIDDLE_SELF_ATTN': { + 'value': False, + 'description': + 'Whether disable the self-attentions in middle blocks.' + }, + 'USE_LINEAR_IN_TRANSFORMER': { + 'value': + False, + 'description': + "Custom transformer's parameter, valid when USE_SPATIAL_TRANSFORMER is True." + }, + 'ADM_IN_CHANNELS': { + 'value': 2048, + 'description': "Used when num_classes == 'sequential'." + }, + } + + def __init__(self, cfg, logger): + super().__init__(cfg, logger=logger) + self.init_params(cfg) + self.construct_network() + + def init_params(self, cfg): + self.in_channels = cfg.IN_CHANNELS + self.model_channels = cfg.MODEL_CHANNELS + self.out_channels = cfg.OUT_CHANNELS + self.num_res_blocks = cfg.NUM_RES_BLOCKS + self.attention_resolutions = cfg.ATTENTION_RESOLUTIONS + + self.num_heads = cfg.get('NUM_HEADS', -1) + self.num_head_channels = cfg.get('NUM_HEADS_CHANNELS', -1) + self.context_dim = cfg.CONTEXT_DIM + self.dropout = cfg.get('DROPOUT', 0) + self.channel_mult = tuple(cfg.get('CHANNEL_MULT', [1, 2, 4, 4])) + self.conv_resample = cfg.get('CONV_RESAMPLE', True) + self.dims = cfg.get('DIMS', 2) + self.num_classes = cfg.get('NUM_CLASSES', None) + self.use_checkpoint = cfg.get('USE_CHECKPOINT', False) + self.use_scale_shift_norm = cfg.get('USE_SCALE_SHIFT_NORM', False) + self.resblock_updown = cfg.get('RESBLOCK_UPDOWN', False) + self.use_new_attention_order = cfg.get('USE_NEW_ATTENTION_ORDER', True) + self.use_spatial_transformer = cfg.get('USE_SPATIAL_TRANSFORMER', True) + self.transformer_depth = cfg.get('TRANSFORMER_DEPTH', 1) + self.use_linear_in_transformer = cfg.get('USE_LINEAR_IN_TRANSFORMER', + False) + self.disable_self_attentions = cfg.get('DISABLE_SELF_ATTENTIONS', None) + self.disable_middle_self_attn = cfg.get('DISABLE_MIDDLE_SELF_ATTN', + False) + self.adm_in_channels = cfg.get('ADM_IN_CHANNELS', None) + self.pretrained_model = cfg.get('PRETRAINED_MODEL', None) + self.ignore_keys = cfg.get('IGNORE_KEYS', []) + + assert (self.num_heads > 0 or self.num_head_channels > 0) and \ + (self.num_heads == -1 or self.num_head_channels == -1) + + if isinstance(self.num_res_blocks, int): + self.num_res_blocks = len( + self.channel_mult) * [self.num_res_blocks] + elif len(self.num_res_blocks) != len(self.channel_mult): + raise ValueError( + 'provide num_res_blocks either as an int (globally constant) or ' + 'as a list/tuple (per-level) with the same length as channel_mult' + ) + + def construct_network(self): + in_channels = self.in_channels + model_channels = self.model_channels + out_channels = self.out_channels + attention_resolutions = self.attention_resolutions + channel_mult = self.channel_mult + num_classes = self.num_classes + num_heads = self.num_heads + num_head_channels = self.num_head_channels + dims = self.dims + dropout = self.dropout + use_checkpoint = self.use_checkpoint + use_scale_shift_norm = self.use_scale_shift_norm + disable_self_attentions = self.disable_self_attentions + disable_middle_self_attn = self.disable_middle_self_attn + transformer_depth = self.transformer_depth + context_dim = self.context_dim + use_linear_in_transformer = self.use_linear_in_transformer + resblock_updown = self.resblock_updown + conv_resample = self.conv_resample + adm_in_channels = self.adm_in_channels + + time_embed_dim = model_channels * 4 + self.time_embed = nn.Sequential( + linear(model_channels, time_embed_dim), + nn.SiLU(), + linear(time_embed_dim, time_embed_dim), + ) + + if self.num_classes is not None: + if isinstance(num_classes, int): + self.label_emb = nn.Embedding(num_classes, time_embed_dim) + elif self.num_classes == 'continuous': + print('setting up linear c_adm embedding layer') + self.label_emb = nn.Linear(1, time_embed_dim) + elif self.num_classes == 'sequential': + assert adm_in_channels is not None + self.label_emb = nn.Sequential( + nn.Sequential( + linear(adm_in_channels, time_embed_dim), + nn.SiLU(), + linear(time_embed_dim, time_embed_dim), + )) + else: + raise ValueError() + + self.input_blocks = nn.ModuleList([ + TimestepEmbedSequential( + conv_nd(dims, in_channels, model_channels, 3, padding=1)) + ]) + self._feature_size = model_channels + input_block_chans = [model_channels] + ch = model_channels + ds = 1 + for level, mult in enumerate(channel_mult): + for nr in range(self.num_res_blocks[level]): + layers = [ + ResBlock( + ch, + time_embed_dim, + dropout, + out_channels=mult * model_channels, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + ) + ] + ch = mult * model_channels + if ds in attention_resolutions: + if num_head_channels == -1: + dim_head = ch // num_heads + else: + num_heads = ch // num_head_channels + dim_head = num_head_channels + disabled_sa = disable_self_attentions[level] if exists( + disable_self_attentions) else False + + layers.append( + SpatialTransformer( + ch, + num_heads, + dim_head, + depth=transformer_depth, + context_dim=context_dim, + disable_self_attn=disabled_sa, + use_linear=use_linear_in_transformer, + use_checkpoint=use_checkpoint)) + self.input_blocks.append(TimestepEmbedSequential(*layers)) + self._feature_size += ch + input_block_chans.append(ch) + if level != len(channel_mult) - 1: + out_ch = ch + self.input_blocks.append( + TimestepEmbedSequential( + ResBlock( + ch, + time_embed_dim, + dropout, + out_channels=out_ch, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + down=True, + ) if resblock_updown else Downsample( + ch, conv_resample, dims=dims, out_channels=out_ch)) + ) + ch = out_ch + input_block_chans.append(ch) + ds *= 2 + self._feature_size += ch + self._input_block_chans = copy.deepcopy(input_block_chans) + + if num_head_channels == -1: + dim_head = ch // num_heads + else: + num_heads = ch // num_head_channels + dim_head = num_head_channels + self.middle_block = TimestepEmbedSequential( + ResBlock( + ch, + time_embed_dim, + dropout, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + ), + SpatialTransformer(ch, + num_heads, + dim_head, + depth=transformer_depth, + context_dim=context_dim, + disable_self_attn=disable_middle_self_attn, + use_linear=use_linear_in_transformer, + use_checkpoint=use_checkpoint), + ResBlock( + ch, + time_embed_dim, + dropout, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + ), + ) + self._feature_size += ch + self._middle_block_chans = [ch] + + self._output_block_chans = [] + self.output_blocks = nn.ModuleList([]) + self.lsc_identity = nn.ModuleList() + for level, mult in list(enumerate(channel_mult))[::-1]: + for i in range(self.num_res_blocks[level] + 1): + ich = input_block_chans.pop() + layers = [ + ResBlock( + ch + ich, + time_embed_dim, + dropout, + out_channels=model_channels * mult, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + ) + ] + ch = model_channels * mult + if ds in attention_resolutions: + if num_head_channels == -1: + dim_head = ch // num_heads + else: + num_heads = ch // num_head_channels + dim_head = num_head_channels + disabled_sa = disable_self_attentions[level] if exists( + disable_self_attentions) else False + layers.append( + SpatialTransformer( + ch, + num_heads, + dim_head, + depth=transformer_depth, + context_dim=context_dim, + disable_self_attn=disabled_sa, + use_linear=use_linear_in_transformer, + use_checkpoint=use_checkpoint)) + if level and i == self.num_res_blocks[level]: + out_ch = ch + layers.append( + ResBlock( + ch, + time_embed_dim, + dropout, + out_channels=out_ch, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + up=True, + ) if resblock_updown else Upsample( + ch, conv_resample, dims=dims, out_channels=out_ch)) + ds //= 2 + + self.output_blocks.append(TimestepEmbedSequential(*layers)) + self.lsc_identity.append(nn.Identity()) + self._feature_size += ch + self._output_block_chans.append(ch) + + self.out = nn.Sequential( + normalization(ch), + nn.SiLU(), + zero_module( + conv_nd(dims, model_channels, out_channels, 3, padding=1)), + ) + + def load_pretrained_model(self, pretrained_model): + if pretrained_model is not None: + with FS.get_from(pretrained_model, + wait_finish=True) as local_model: + self.init_from_ckpt(local_model, ignore_keys=self.ignore_keys) + + def init_from_ckpt(self, path, ignore_keys=list()): + if path.endswith('safetensors'): + from safetensors.torch import load_file as load_safetensors + sd = load_safetensors(path) + else: + sd = torch.load(path, map_location='cpu') + + new_sd = OrderedDict() + for k, v in sd.items(): + ignored = False + for ik in ignore_keys: + if ik in k: + if we.rank == 0: + self.logger.info( + 'Ignore key {} from state_dict.'.format(k)) + ignored = True + break + if not ignored: + new_sd[k] = v + + missing, unexpected = self.load_state_dict(new_sd, strict=False) + if we.rank == 0: + self.logger.info( + f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys' + ) + if len(missing) > 0: + self.logger.info(f'Missing Keys:\n {missing}') + if len(unexpected) > 0: + self.logger.info(f'\nUnexpected Keys:\n {unexpected}') + + def forward(self, x, t=None, cond=dict()): + t_emb = timestep_embedding(t, self.model_channels, repeat_only=False) + emb = self.time_embed(t_emb) + + if isinstance(cond, dict): + if 'y' in cond and cond['y'] is not None: + assert self.num_classes is not None + emb = emb + self.label_emb(cond['y']) + if 'concat' in cond: + c = cond['concat'] + x = torch.cat([x, c], dim=1) + + context = cond.get('crossattn', None) + else: + context = cond + + hs = [] + h = x + for module in self.input_blocks: + h = module(h, emb, context) + hs.append(h) + h = self.middle_block(h, emb, context) + for m_id, module in enumerate(self.output_blocks): + h = torch.cat([h, self.lsc_identity[m_id](hs.pop())], dim=1) + target_size = hs[-1].shape[-2:] if len(hs) > 0 else None + h = module(h, emb, context, target_size) + + return self.out(h) + + @staticmethod + def get_config_template(): + return dict_to_yaml('BACKBONE', + __class__.__name__, + DiffusionUNet.para_dict, + set_name=True) + + +@BACKBONES.register_class() +class DiffusionUNetXL(DiffusionUNet): + para_dict = { + 'TRANSFORMER_DEPTH_MIDDLE': { + 'value': + None, + 'description': + "Custom transformer's depth of middle block, If set None, use TRANSFORMER_DEPTH last value." + }, + } + para_dict.update(DiffusionUNet.para_dict) + + def init_params(self, cfg): + super().init_params(cfg) + + if isinstance(self.transformer_depth, int): + self.transformer_depth = len( + self.channel_mult) * [self.transformer_depth] + elif isinstance(self.transformer_depth, list): + assert len(self.transformer_depth) == len(self.channel_mult) + + self.transformer_depth_middle = cfg.get('TRANSFORMER_DEPTH_MIDDLE', + self.transformer_depth[-1]) + + def construct_network(self): + in_channels = self.in_channels + model_channels = self.model_channels + out_channels = self.out_channels + attention_resolutions = self.attention_resolutions + channel_mult = self.channel_mult + num_classes = self.num_classes + num_heads = self.num_heads + num_head_channels = self.num_head_channels + dims = self.dims + dropout = self.dropout + use_checkpoint = self.use_checkpoint + use_scale_shift_norm = self.use_scale_shift_norm + disable_self_attentions = self.disable_self_attentions + disable_middle_self_attn = self.disable_middle_self_attn + transformer_depth = self.transformer_depth + transformer_depth_middle = self.transformer_depth_middle + context_dim = self.context_dim + use_linear_in_transformer = self.use_linear_in_transformer + resblock_updown = self.resblock_updown + conv_resample = self.conv_resample + adm_in_channels = self.adm_in_channels + + time_embed_dim = model_channels * 4 + self.time_embed = nn.Sequential( + linear(model_channels, time_embed_dim), + nn.SiLU(), + linear(time_embed_dim, time_embed_dim), + ) + + if self.num_classes is not None: + if isinstance(self.num_classes, int): + self.label_emb = nn.Embedding(num_classes, time_embed_dim) + elif self.num_classes == 'continuous': + print('setting up linear c_adm embedding layer') + self.label_emb = nn.Linear(1, time_embed_dim) + elif self.num_classes == 'timestep': + self.label_emb = nn.Sequential( + Timestep(model_channels), + nn.Sequential( + linear(model_channels, time_embed_dim), + nn.SiLU(), + linear(time_embed_dim, time_embed_dim), + ), + ) + elif self.num_classes == 'sequential': + assert adm_in_channels is not None + self.label_emb = nn.Sequential( + nn.Sequential( + linear(adm_in_channels, time_embed_dim), + nn.SiLU(), + linear(time_embed_dim, time_embed_dim), + )) + else: + raise ValueError() + + self.input_blocks = nn.ModuleList([ + TimestepEmbedSequential( + conv_nd(dims, in_channels, model_channels, 3, padding=1)) + ]) + self._feature_size = model_channels + input_block_chans = [model_channels] + ch = model_channels + ds = 1 + for level, mult in enumerate(channel_mult): + for nr in range(self.num_res_blocks[level]): + layers = [ + ResBlock( + ch, + time_embed_dim, + dropout, + out_channels=mult * model_channels, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + ) + ] + ch = mult * model_channels + if ds in attention_resolutions: + if num_head_channels == -1: + dim_head = ch // num_heads + else: + num_heads = ch // num_head_channels + dim_head = num_head_channels + disabled_sa = disable_self_attentions[level] if exists( + disable_self_attentions) else False + + layers.append( + SpatialTransformer( + ch, + num_heads, + dim_head, + depth=transformer_depth[level], + context_dim=context_dim, + disable_self_attn=disabled_sa, + use_linear=use_linear_in_transformer, + use_checkpoint=use_checkpoint)) + self.input_blocks.append(TimestepEmbedSequential(*layers)) + self._feature_size += ch + input_block_chans.append(ch) + if level != len(channel_mult) - 1: + out_ch = ch + self.input_blocks.append( + TimestepEmbedSequential( + ResBlock( + ch, + time_embed_dim, + dropout, + out_channels=out_ch, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + down=True, + ) if resblock_updown else Downsample( + ch, conv_resample, dims=dims, out_channels=out_ch)) + ) + ch = out_ch + input_block_chans.append(ch) + ds *= 2 + self._feature_size += ch + self._input_block_chans = copy.deepcopy(input_block_chans) + + if num_head_channels == -1: + dim_head = ch // num_heads + else: + num_heads = ch // num_head_channels + dim_head = num_head_channels + self.middle_block = TimestepEmbedSequential( + ResBlock( + ch, + time_embed_dim, + dropout, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + ), + SpatialTransformer(ch, + num_heads, + dim_head, + depth=transformer_depth_middle, + context_dim=context_dim, + disable_self_attn=disable_middle_self_attn, + use_linear=use_linear_in_transformer, + use_checkpoint=use_checkpoint), + ResBlock( + ch, + time_embed_dim, + dropout, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + ), + ) + self._feature_size += ch + self._middle_block_chans = [ch] + + self._output_block_chans = [] + self.output_blocks = nn.ModuleList([]) + self.lsc_identity = nn.ModuleList() + for level, mult in list(enumerate(channel_mult))[::-1]: + for i in range(self.num_res_blocks[level] + 1): + ich = input_block_chans.pop() + layers = [ + ResBlock( + ch + ich, + time_embed_dim, + dropout, + out_channels=model_channels * mult, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + ) + ] + ch = model_channels * mult + if ds in attention_resolutions: + if num_head_channels == -1: + dim_head = ch // num_heads + else: + num_heads = ch // num_head_channels + dim_head = num_head_channels + disabled_sa = disable_self_attentions[level] if exists( + disable_self_attentions) else False + layers.append( + SpatialTransformer( + ch, + num_heads, + dim_head, + depth=transformer_depth[level], + context_dim=context_dim, + disable_self_attn=disabled_sa, + use_linear=use_linear_in_transformer, + use_checkpoint=use_checkpoint)) + if level and i == self.num_res_blocks[level]: + out_ch = ch + layers.append( + ResBlock( + ch, + time_embed_dim, + dropout, + out_channels=out_ch, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + up=True, + ) if resblock_updown else Upsample( + ch, conv_resample, dims=dims, out_channels=out_ch)) + ds //= 2 + + self.output_blocks.append(TimestepEmbedSequential(*layers)) + self.lsc_identity.append(nn.Identity()) + self._feature_size += ch + self._output_block_chans.append(ch) + + self.out = nn.Sequential( + normalization(ch), + nn.SiLU(), + zero_module( + conv_nd(dims, model_channels, out_channels, 3, padding=1)), + ) + + def forward(self, x, t=None, cond=dict()): + t_emb = timestep_embedding(t, + self.model_channels, + repeat_only=False, + legacy=True) + emb = self.time_embed(t_emb) + + if isinstance(cond, dict): + if 'y' in cond: + assert self.num_classes is not None + emb = emb + self.label_emb(cond['y']) + if 'concat' in cond: + c = cond['concat'] + x = torch.cat([x, c], dim=1) + + context = cond.get('crossattn', None) + else: + context = cond + + hs = [] + h = x + for module in self.input_blocks: + h = module(h, emb, context) + hs.append(h) + h = self.middle_block(h, emb, context) + for m_id, module in enumerate(self.output_blocks): + h = torch.cat([h, self.lsc_identity[m_id](hs.pop())], dim=1) + target_size = hs[-1].shape[-2:] if len(hs) > 0 else None + h = module(h, emb, context, target_size) + + return self.out(h) + + def convert_to_fp16(self): + """ + Convert the torso of the model to float16. + """ + self.input_blocks.apply(convert_module_to_f16) + self.middle_block.apply(convert_module_to_f16) + self.output_blocks.apply(convert_module_to_f16) + + def convert_to_fp32(self): + """ + Convert the torso of the model to float32. + """ + self.input_blocks.apply(convert_module_to_f32) + self.middle_block.apply(convert_module_to_f32) + self.output_blocks.apply(convert_module_to_f32) + + @staticmethod + def get_config_template(): + return dict_to_yaml('BACKBONE', + __class__.__name__, + DiffusionUNetXL.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/unet/unet_utils.py b/scepter/modules/model/backbone/unet/unet_utils.py new file mode 100644 index 0000000..be0f5c1 --- /dev/null +++ b/scepter/modules/model/backbone/unet/unet_utils.py @@ -0,0 +1,1000 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import math +import warnings +from abc import abstractmethod +from importlib import find_loader + +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F +from einops import rearrange, repeat + +from scepter.modules.model.utils.basic_utils import checkpoint, default, exists + +try: + import xformers + import xformers.ops + XFORMERS_IS_AVAILBLE = True +except Exception as e: + XFORMERS_IS_AVAILBLE = False + warnings.warn(f'{e}') + +if find_loader('flash_attn'): + FLASH_ATTN_IS_AVAILABLE = True +else: + FLASH_ATTN_IS_AVAILABLE = False + + +def normalization(channels): + """ + Make a standard normalization layer. + :param channels: number of input channels. + :return: an nn.Module for normalization. + """ + return GroupNorm32(32, channels) + + +class GroupNorm32(nn.GroupNorm): + def forward(self, x): + return super().forward(x.float()).type(x.dtype) + + +def count_flops_attn(model, _x, y): + """ + A counter for the `thop` package to count the operations in an + attention operation. + Meant to be used like: + macs, params = thop.profile( + model, + inputs=(inputs, timestamps), + custom_ops={QKVAttention: QKVAttention.count_flops}, + ) + """ + b, c, *spatial = y[0].shape + num_spatial = int(np.prod(spatial)) + # We perform two matmuls with the same number of ops. + # The first computes the weight matrix, the second computes + # the combination of the value vectors. + matmul_ops = 2 * b * (num_spatial**2) * c + model.total_ops += torch.DoubleTensor([matmul_ops]) + + +def conv_nd(dims, *args, **kwargs): + """ + Create a 1D, 2D, or 3D convolution module. + """ + if dims == 1: + return nn.Conv1d(*args, **kwargs) + elif dims == 2: + return nn.Conv2d(*args, **kwargs) + elif dims == 3: + return nn.Conv3d(*args, **kwargs) + raise ValueError(f'unsupported dimensions: {dims}') + + +def linear(*args, **kwargs): + """ + Create a linear module. + """ + return nn.Linear(*args, **kwargs) + + +def avg_pool_nd(dims, *args, **kwargs): + """ + Create a 1D, 2D, or 3D average pooling module. + """ + if dims == 1: + return nn.AvgPool1d(*args, **kwargs) + elif dims == 2: + return nn.AvgPool2d(*args, **kwargs) + elif dims == 3: + return nn.AvgPool3d(*args, **kwargs) + raise ValueError(f'unsupported dimensions: {dims}') + + +def zero_module(module): + """ + Zero out the parameters of a module and return it. + """ + for p in module.parameters(): + p.detach().zero_() + return module + + +def timestep_embedding(timesteps, + dim, + max_period=10000, + repeat_only=False, + legacy=False): + """ + Create sinusoidal timestep embeddings. + :param timesteps: a 1-D Tensor of N indices, one per batch element. + These may be fractional. + :param dim: the dimension of the output. + :param max_period: controls the minimum frequency of the embeddings. + :return: an [N x dim] Tensor of positional embeddings. + """ + if not repeat_only: + half = dim // 2 + freqs = torch.exp( + -math.log(max_period) * + torch.arange(start=0, end=half, dtype=torch.float32) / + half).to(device=timesteps.device) + if legacy: + args = timesteps[:, None].float() * freqs[None] + else: + args = torch.mm(timesteps.float().unsqueeze(1), + freqs.unsqueeze(0)).view(timesteps.shape[0], + len(freqs)) + embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) + if dim % 2: + embedding = torch.cat( + [embedding, torch.zeros_like(embedding[:, :1])], dim=-1) + else: + embedding = repeat(timesteps, 'b -> b d', d=dim) + return embedding + + +class Timestep(nn.Module): + def __init__(self, dim, legacy=False): + super().__init__() + self.dim = dim + self.legacy = legacy + + def forward(self, t): + return timestep_embedding(t, self.dim, legacy=self.legacy) + + +class TimestepBlock(nn.Module): + """ + Any module where forward() takes timestep embeddings as a second argument. + """ + @abstractmethod + def forward(self, x, emb): + """ + Apply the module to `x` given `emb` timestep embeddings. + """ + + +class TimestepEmbedSequential(nn.Sequential, TimestepBlock): + """ + A sequential module that passes timestep embeddings to the children that + support it as an extra input. + """ + def forward(self, x, emb, context=None, target_size=None): + for layer in self: + if isinstance(layer, TimestepBlock): + x = layer(x, emb) + elif isinstance(layer, SpatialTransformer): + x = layer(x, context) + elif isinstance(layer, Upsample): + x = layer(x, target_size) + else: + x = layer(x) + return x + + +class Upsample(nn.Module): + """ + An upsampling layer with an optional convolution. + :param channels: channels in the inputs and outputs. + :param use_conv: a bool determining if a convolution is applied. + :param dims: determines if the signal is 1D, 2D, or 3D. If 3D, then + upsampling occurs in the inner-two dimensions. + """ + def __init__(self, + channels, + use_conv, + dims=2, + out_channels=None, + padding=1): + super().__init__() + self.channels = channels + self.out_channels = out_channels or channels + self.use_conv = use_conv + self.dims = dims + if use_conv: + self.conv = conv_nd(dims, + self.channels, + self.out_channels, + 3, + padding=padding) + + def forward(self, x, target_size=None): + assert x.shape[1] == self.channels + if self.dims == 3: + x = F.interpolate(x.float(), + (x.shape[2], x.shape[3] * 2, x.shape[4] * 2), + mode='nearest').type_as(x) + else: + if target_size is None: + x = F.interpolate(x.float(), scale_factor=2, + mode='nearest').type_as(x) + else: + x = F.interpolate(x.float(), target_size, + mode='nearest').type_as(x) + if self.use_conv: + x = self.conv(x) + return x + + +class Downsample(nn.Module): + """ + A downsampling layer with an optional convolution. + :param channels: channels in the inputs and outputs. + :param use_conv: a bool determining if a convolution is applied. + :param dims: determines if the signal is 1D, 2D, or 3D. If 3D, then + downsampling occurs in the inner-two dimensions. + """ + def __init__(self, + channels, + use_conv, + dims=2, + out_channels=None, + padding=1): + super().__init__() + self.channels = channels + self.out_channels = out_channels or channels + self.use_conv = use_conv + self.dims = dims + stride = 2 if dims != 3 else (1, 2, 2) + if use_conv: + self.op = conv_nd(dims, + self.channels, + self.out_channels, + 3, + stride=stride, + padding=padding) + else: + assert self.channels == self.out_channels + self.op = avg_pool_nd(dims, kernel_size=stride, stride=stride) + + def forward(self, x): + assert x.shape[1] == self.channels + return self.op(x) + + +class ResBlock(TimestepBlock): + """ + A residual block that can optionally change the number of channels. + :param channels: the number of input channels. + :param emb_channels: the number of timestep embedding channels. + :param dropout: the rate of dropout. + :param out_channels: if specified, the number of out channels. + :param use_conv: if True and out_channels is specified, use a spatial + convolution instead of a smaller 1x1 convolution to change the + channels in the skip connection. + :param dims: determines if the signal is 1D, 2D, or 3D. + :param use_checkpoint: if True, use gradient checkpointing on this module. + :param up: if True, use this block for upsampling. + :param down: if True, use this block for downsampling. + """ + def __init__( + self, + channels, + emb_channels, + dropout, + out_channels=None, + use_conv=False, + use_scale_shift_norm=False, + dims=2, + use_checkpoint=False, + up=False, + down=False, + ): + super().__init__() + self.channels = channels + self.emb_channels = emb_channels + self.dropout = dropout + self.out_channels = out_channels or channels + self.use_conv = use_conv + self.use_checkpoint = use_checkpoint + self.use_scale_shift_norm = use_scale_shift_norm + + self.in_layers = nn.Sequential( + normalization(channels), + nn.SiLU(), + conv_nd(dims, channels, self.out_channels, 3, padding=1), + ) + + self.updown = up or down + + if up: + self.h_upd = Upsample(channels, False, dims) + self.x_upd = Upsample(channels, False, dims) + elif down: + self.h_upd = Downsample(channels, False, dims) + self.x_upd = Downsample(channels, False, dims) + else: + self.h_upd = self.x_upd = nn.Identity() + + self.emb_layers = nn.Sequential( + nn.SiLU(), + linear( + emb_channels, + 2 * self.out_channels + if use_scale_shift_norm else self.out_channels, + ), + ) + self.out_layers = nn.Sequential( + normalization(self.out_channels), + nn.SiLU(), + nn.Dropout(p=dropout), + zero_module( + conv_nd(dims, + self.out_channels, + self.out_channels, + 3, + padding=1)), + ) + + if self.out_channels == channels: + self.skip_connection = nn.Identity() + elif use_conv: + self.skip_connection = conv_nd(dims, + channels, + self.out_channels, + 3, + padding=1) + else: + self.skip_connection = conv_nd(dims, channels, self.out_channels, + 1) + + def forward(self, x, emb): + """ + Apply the block to a Tensor, conditioned on a timestep embedding. + :param x: an [N x C x ...] Tensor of features. + :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) + + def _forward(self, x, emb): + if self.updown: + in_rest, in_conv = self.in_layers[:-1], self.in_layers[-1] + h = in_rest(x) + h = self.h_upd(h) + x = self.x_upd(x) + h = in_conv(h) + else: + h = self.in_layers(x) + emb_out = self.emb_layers(emb).type(h.dtype) + while len(emb_out.shape) < len(h.shape): + emb_out = emb_out[..., None] + if self.use_scale_shift_norm: + out_norm, out_rest = self.out_layers[0], self.out_layers[1:] + scale, shift = torch.chunk(emb_out, 2, dim=1) + h = out_norm(h) * (1 + scale) + shift + h = out_rest(h) + else: + h = h + emb_out + h = self.out_layers(h) + return self.skip_connection(x) + h + + +class AttentionBlock(nn.Module): + """ + An attention block that allows spatial positions to attend to each other. + Originally ported from here, but adapted to the N-d case. + https://github.com/hojonathanho/diffusion/blob/1e0dceb3b3495bbe19116a5e1b3596cd0706c543/diffusion_tf/models/unet.py#L66. + """ + def __init__( + self, + channels, + num_heads=1, + num_head_channels=-1, + use_checkpoint=False, + use_new_attention_order=False, + ): + super().__init__() + self.channels = channels + if num_head_channels == -1: + self.num_heads = num_heads + else: + assert channels % num_head_channels == 0, \ + f'q,k,v channels {channels} is not divisible by num_head_channels {num_head_channels}' + self.num_heads = channels // num_head_channels + + self.use_checkpoint = use_checkpoint + self.norm = normalization(channels) + self.qkv = conv_nd(1, channels, channels * 3, 1) + if use_new_attention_order: + # split qkv before split head + self.attention = QKVAttention(self.num_heads) + else: + # split head before split qkv + self.attention = QKVAttentionLegacy(self.num_heads) + + 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) + + def _forward(self, x): + b, c, *spatial = x.shape + x = x.reshape(b, c, -1) + qkv = self.qkv(self.norm(x)) + h = self.attention(qkv) + h = self.proj_out(h) + return (x + h).reshape(b, c, *spatial) + + +class QKVAttentionLegacy(nn.Module): + """ + A module which performs QKV attention. Matches legacy QKVAttention + input/ouput head shaping + """ + def __init__(self, n_heads): + super().__init__() + self.n_heads = n_heads + + def forward(self, qkv): + """ + Apply QKV attention. + + :param qkv: an [N x (H * 3 * C) x T] tensor of Qs, Ks, and Vs. + :return: an [N x (H * C) x T] tensor after attention. + """ + bs, width, length = qkv.shape + assert width % (3 * self.n_heads) == 0 + ch = width // (3 * self.n_heads) + q, k, v = qkv.reshape(bs * self.n_heads, ch * 3, length).split(ch, + dim=1) + scale = 1 / math.sqrt(math.sqrt(ch)) + weight = torch.einsum( + 'bct,bcs->bts', q * scale, + k * scale) # More stable with f16 than dividing afterwards + weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype) + a = torch.einsum('bts,bcs->bct', weight, v) + return a.reshape(bs, -1, length) + + @staticmethod + def count_flops(model, _x, y): + return count_flops_attn(model, _x, y) + + +class QKVAttention(nn.Module): + """ + A module which performs QKV attention. Matches legacy QKVAttention + input/ouput head shaping + """ + def __init__(self, n_heads): + super().__init__() + self.n_heads = n_heads + + def forward(self, qkv): + """ + Apply QKV attention. + :param qkv: an [N x (H * 3 * C) x T] tensor of Qs, Ks, and Vs. + :return: an [N x (H * C) x T] tensor after attention. + """ + bs, width, length = qkv.shape + assert width % (3 * self.n_heads) == 0 + ch = width // (3 * self.n_heads) + q, k, v = qkv.reshape(bs * self.n_heads, ch * 3, length).split(ch, + dim=1) + scale = 1 / math.sqrt(math.sqrt(ch)) + weight = torch.einsum( + 'bct,bcs->bts', q * scale, + k * scale) # More stable with f16 than dividing afterwards + weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype) + a = torch.einsum('bts,bcs->bct', weight, v) + return a.reshape(bs, -1, length) + + +class GEGLU(nn.Module): + def __init__(self, dim_in, dim_out): + super().__init__() + self.proj = nn.Linear(dim_in, dim_out * 2) + + def forward(self, x): + x, gate = self.proj(x).chunk(2, dim=-1) + return x * F.gelu(gate) + + +class FeedForward(nn.Module): + def __init__(self, dim, dim_out=None, mult=4, glu=False, dropout=0.): + super().__init__() + inner_dim = int(dim * mult) + dim_out = default(dim_out, dim) + project_in = nn.Sequential(nn.Linear( + dim, inner_dim), nn.GELU()) if not glu else GEGLU(dim, inner_dim) + + self.net = nn.Sequential(project_in, nn.Dropout(dropout), + nn.Linear(inner_dim, dim_out)) + + def forward(self, x): + return self.net(x) + + +class MultiHeadAttention(nn.Module): + def __init__(self, + dim, + context_dim=None, + num_heads=None, + head_dim=None, + dropout=0.0, + flash_dtype=torch.float16): + # consider head_dim first, then num_heads + num_heads = dim // head_dim if head_dim else num_heads + head_dim = dim // num_heads + assert num_heads * head_dim == dim + context_dim = context_dim or dim + assert flash_dtype in (None, torch.float16, torch.bfloat16) + super().__init__() + self.dim = dim + self.context_dim = context_dim + self.num_heads = num_heads + self.head_dim = head_dim + self.scale = math.pow(head_dim, -0.25) + self.flash_dtype = flash_dtype + + # layers + self.q = nn.Linear(dim, dim, bias=False) + self.k = nn.Linear(context_dim, dim, bias=False) + self.v = nn.Linear(context_dim, dim, bias=False) + self.o = nn.Linear(dim, dim) + self.dropout = nn.Dropout(dropout) + + def forward(self, x, context=None): + """x: [B, L, C]. + context: [B, L', C'] or None. + """ + context = x if context is None else context + b, n, d = x.size(0), self.num_heads, self.head_dim + + # compute query, key, value + q = self.q(x).view(b, -1, n, d) + k = self.k(context).view(b, -1, n, d) + v = self.v(context).view(b, -1, n, d) + + attn = torch.einsum('binc,bjnc->bnij', q * self.scale, k * self.scale) + attn = F.softmax(attn.float(), dim=-1).type_as(attn) + x = torch.einsum('bnij,bjnc->binc', attn, v.float()) + # output + x = x.reshape(b, -1, n * d) + x = self.o(x) + x = self.dropout(x) + return x + + +class FlashattnMultiHeadAttention(nn.Module): + def __init__(self, + dim, + context_dim=None, + num_heads=None, + head_dim=None, + dropout=0.0, + flash_dtype=torch.float16): + # consider head_dim first, then num_heads + num_heads = dim // head_dim if head_dim else num_heads + head_dim = dim // num_heads + assert num_heads * head_dim == dim + context_dim = context_dim or dim + assert flash_dtype in (None, torch.float16, torch.bfloat16) + super().__init__() + self.dim = dim + self.context_dim = context_dim + self.num_heads = num_heads + self.head_dim = head_dim + self.scale = math.pow(head_dim, -0.25) + self.flash_dtype = flash_dtype + + # layers + self.q = nn.Linear(dim, dim, bias=False) + self.k = nn.Linear(context_dim, dim, bias=False) + self.v = nn.Linear(context_dim, dim, bias=False) + self.o = nn.Linear(dim, dim) + self.dropout = nn.Dropout(dropout) + + def forward(self, x, context=None): + """x: [B, L, C]. + context: [B, L', C'] or None. + """ + context = x if context is None else context + b, n, d = x.size(0), self.num_heads, self.head_dim + + # compute query, key, value + q = self.q(x).view(b, -1, n, d) + k = self.k(context).view(b, -1, n, d) + v = self.v(context).view(b, -1, n, d) + + # compute attention + if (x.device.type != 'cpu' and find_loader('flash_attn') + and self.head_dim % 8 == 0 and self.head_dim <= 128 + and self.flash_dtype is not None): + # flash implementation + from flash_attn.flash_attn_interface import \ + flash_attn_unpadded_kvpacked_func + dtype = q.dtype + if dtype != self.flash_dtype: + q = q.type(self.flash_dtype) + k = k.type(self.flash_dtype) + v = v.type(self.flash_dtype) + cu_seqlens_q = torch.arange(0, + b * q.size(1) + 1, + q.size(1), + dtype=torch.int32, + device=x.device) + cu_seqlens_k = torch.arange(0, + b * k.size(1) + 1, + k.size(1), + dtype=torch.int32, + device=x.device) + x = flash_attn_unpadded_kvpacked_func( + q=q.reshape(-1, n, d).contiguous(), + kv=torch.stack([k.reshape(-1, n, d), + v.reshape(-1, n, d)], + dim=1).contiguous(), + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_k=cu_seqlens_k, + max_seqlen_q=q.size(1), + max_seqlen_k=k.size(1), + dropout_p=self.dropout.p if self.training else 0.0, + return_attn_probs=False).reshape(b, -1, n, d).type(dtype) + else: + # attn = torch.einsum('binc,bjnc->bnij', q * self.scale, k * self.scale) + # attn = F.softmax(attn.float(), dim=-1).type_as(attn) + # x = torch.einsum('bnij,bjnc->binc', attn, v.float()) + # torch implementation + q = q.permute(0, 2, 1, 3) * self.scale + q = torch.clamp(q, min=-65504, max=66504) + k = k.permute(0, 2, 3, 1) * self.scale + k = torch.clamp(k, min=-65504, max=66504) + v = v.permute(0, 2, 1, 3) + if q.shape[1] == 10 and k.shape[ + 1] == 10 and q.shape[2] >= 8192 and k.shape[3] >= 8192: + qkv = zip(q.chunk(10, dim=1), k.chunk(10, dim=1), + v.chunk(10, dim=1)) + tmp = [] + for q, k, v in qkv: + attn = torch.matmul(q, k) + attn = torch.clamp(attn, min=-65504, max=65504) + # print(f"attn1 has no inf: {torch.all(torch.isinf(attn) == False)}, attn1 dtype: {attn.dtype}") + # print(f"attn1 has no nan: {torch.all(torch.isnan(attn) == False)}") + attn = F.softmax(attn.float(), dim=-1).type_as(attn) + tmp.append(torch.matmul(attn, v)) + x = torch.cat(tmp, 1).permute(0, 2, 1, 3) + else: + attn = torch.matmul(q, k) + attn = torch.clamp(attn, min=-65504, max=65504) + # print(f"attn1 has no inf: {torch.all(torch.isinf(attn) == False)}, attn1 dtype: {attn.dtype}") + # print(f"attn1 has no nan: {torch.all(torch.isnan(attn) == False)}") + attn = F.softmax(attn.float(), dim=-1).type_as(attn) + x = torch.matmul(attn, v).permute(0, 2, 1, 3) + + # output + x = x.reshape(b, -1, n * d) + x = self.o(x) + x = self.dropout(x) + return x + + +class XFormerMultiHeadAttention(nn.Module): + def __init__(self, + dim, + context_dim=None, + num_heads=None, + head_dim=None, + dropout=0.0): + super().__init__() + # consider head_dim first, then num_heads + num_heads = dim // head_dim if head_dim else num_heads + head_dim = dim // num_heads + assert num_heads * head_dim == dim + context_dim = context_dim or dim + self.dim = dim + self.context_dim = context_dim + self.num_heads = num_heads + self.head_dim = head_dim + self.scale = math.pow(head_dim, -0.25) + + # layers + self.q = nn.Linear(dim, dim, bias=False) + self.k = nn.Linear(context_dim, dim, bias=False) + self.v = nn.Linear(context_dim, dim, bias=False) + self.o = nn.Linear(dim, dim) + self.dropout = nn.Dropout(dropout) + self.attention_op = None + + def x_form(self, x, context=None): + context = x if context is None else context + b, n, d = x.size(0), self.num_heads, self.head_dim + # compute query, key, value + q = self.q(x).view(b, -1, n, d) + k = self.k(context).view(b, -1, n, d) + v = self.v(context).view(b, -1, n, d) + + x = xformers.ops.memory_efficient_attention(q, + k, + v, + attn_bias=None, + op=self.attention_op) + x = x.reshape(b, -1, n * d) + x = self.o(x) + x = self.dropout(x) + return x + + def x_ori(self, x, context=None): + """x: [B, L, C]. + context: [B, L', C'] or None. + """ + context = x if context is None else context + b, n, d = x.size(0), self.num_heads, self.head_dim + + # compute query, key, value + q = self.q(x).view(b, -1, n, d) + k = self.k(context).view(b, -1, n, d) + v = self.v(context).view(b, -1, n, d) + + attn = torch.einsum('binc,bjnc->bnij', q * self.scale, k * self.scale) + attn = F.softmax(attn.float(), dim=-1).type_as(attn) + x = torch.einsum('bnij,bjnc->binc', attn, v.float()) + # output + x = x.reshape(b, -1, n * d) + x = self.o(x) + x = self.dropout(x) + return x + + def forward(self, x, context=None): + """x: [B, L, C]. + context: [B, L', C'] or None. + """ + if XFORMERS_IS_AVAILBLE: + x = self.x_form(x, context=context) + else: + x = self.x_ori(x, context=context) + return x + + +class CrossAttention(nn.Module): + def __init__(self, + query_dim, + context_dim=None, + heads=8, + dim_head=64, + dropout=0.): + super().__init__() + inner_dim = dim_head * heads + context_dim = default(context_dim, query_dim) + + self.scale = dim_head**-0.5 + self.heads = heads + + self.to_q = nn.Linear(query_dim, inner_dim, bias=False) + self.to_k = nn.Linear(context_dim, inner_dim, bias=False) + self.to_v = nn.Linear(context_dim, inner_dim, bias=False) + + self.to_out = nn.Sequential(nn.Linear(inner_dim, query_dim), + nn.Dropout(dropout)) + + def forward(self, x, context=None, mask=None): + h = self.heads + + q = self.to_q(x) + context = default(context, x) + k = self.to_k(context) + v = self.to_v(context) + + q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), + (q, k, v)) + sim = torch.einsum('b i d, b j d -> b i j', q, k) * self.scale + + if exists(mask): + mask = rearrange(mask, 'b ... -> b (...)') + max_neg_value = -torch.finfo(sim.dtype).max + mask = repeat(mask, 'b j -> (b h) () j', h=h) + sim.masked_fill_(~mask, max_neg_value) + + # attention, what we cannot get enough of + sim = sim.softmax(dim=-1) + + out = torch.einsum('b i j, b j d -> b i d', sim, v) + out = rearrange(out, '(b h) n d -> b n (h d)', h=h) + return self.to_out(out) + + +class MemoryEfficientCrossAttention(nn.Module): + def __init__(self, + query_dim, + context_dim=None, + heads=8, + dim_head=64, + dropout=0.0): + super().__init__() + inner_dim = dim_head * heads + context_dim = default(context_dim, query_dim) + + self.heads = heads + self.dim_head = dim_head + + self.to_q = nn.Linear(query_dim, inner_dim, bias=False) + self.to_k = nn.Linear(context_dim, inner_dim, bias=False) + self.to_v = nn.Linear(context_dim, inner_dim, bias=False) + + self.to_out = nn.Sequential(nn.Linear(inner_dim, query_dim), + nn.Dropout(dropout)) + self.attention_op = None + + def forward(self, x, context=None, mask=None): + + if x.shape[-1] < 8: + with torch.autocast(enabled=False, device_type='cuda'): + q = self.to_q(x) + context = default(context, x) + k = self.to_k(context) + v = self.to_v(context) + else: + q = self.to_q(x) + context = default(context, x) + k = self.to_k(context) + v = self.to_v(context) + + b, _, _ = q.shape + q, k, v = map( + lambda t: t.unsqueeze(3).reshape(b, t.shape[ + 1], self.heads, self.dim_head).permute(0, 2, 1, 3).reshape( + b * self.heads, t.shape[1], self.dim_head).contiguous(), + (q, k, v), + ) + + # actually compute the attention, what we cannot get enough of + out = xformers.ops.memory_efficient_attention(q, + k, + v, + attn_bias=None, + op=self.attention_op) + + # TODO: Use this directly in the attention operation, as a bias + if exists(mask): + raise NotImplementedError + out = (out.unsqueeze(0).reshape( + b, self.heads, out.shape[1], + self.dim_head).permute(0, 2, 1, + 3).reshape(b, out.shape[1], + self.heads * self.dim_head)) + return self.to_out(out) + + +class BasicTransformerBlock(nn.Module): + def __init__(self, + dim, + n_heads, + d_head, + dropout=0., + context_dim=None, + gated_ff=True, + use_checkpoint=True, + disable_self_attn=False): + super().__init__() + self.disable_self_attn = disable_self_attn + AttentionBuilder = MemoryEfficientCrossAttention if XFORMERS_IS_AVAILBLE else CrossAttention + self.attn1 = AttentionBuilder( + query_dim=dim, + heads=n_heads, + dim_head=d_head, + dropout=dropout, + context_dim=context_dim + if self.disable_self_attn else None) # is a self-attention + self.ff = FeedForward(dim, dropout=dropout, glu=gated_ff) + self.attn2 = AttentionBuilder( + query_dim=dim, + context_dim=context_dim, + heads=n_heads, + dim_head=d_head, + dropout=dropout) # is self-attn if context is none + self.norm1 = nn.LayerNorm(dim) + self.norm2 = nn.LayerNorm(dim) + self.norm3 = nn.LayerNorm(dim) + self.use_checkpoint = use_checkpoint + + def forward(self, x, context=None): + return checkpoint(self._forward, (x, context), self.parameters(), + self.use_checkpoint) + + def _forward(self, x, context=None): + x = self.attn1(self.norm1(x), + context=context if self.disable_self_attn else None) + x + x = self.attn2(self.norm2(x), context=context) + x + x = self.ff(self.norm3(x)) + x + return x + + +class SpatialTransformer(nn.Module): + """ + Transformer block for image-like data. + First, project the input (aka embedding) + and reshape to b, t, d. + Then apply standard transformer action. + Finally, reshape to image + NEW: use_linear for more efficiency instead of the 1x1 convs + """ + def __init__(self, + in_channels, + n_heads, + d_head, + depth=1, + dropout=0., + context_dim=None, + disable_self_attn=False, + use_linear=False, + use_checkpoint=True): + super().__init__() + if exists(context_dim) and not isinstance(context_dim, list): + context_dim = [context_dim] + + if exists(context_dim) and not isinstance(context_dim, (list)): + context_dim = [context_dim] + if exists(context_dim) and isinstance(context_dim, list): + if depth != len(context_dim): + print( + f'WARNING: {self.__class__.__name__}: Found context dims {context_dim} of' + f" depth {len(context_dim)}, which does not match the specified 'depth' of" + f' {depth}. Setting context_dim to {depth * [context_dim[0]]} now.' + ) + # depth does not match context dims. + assert all( + map(lambda x: x == context_dim[0], context_dim) + ), 'need homogenous context_dim to match depth automatically' + context_dim = depth * [context_dim[0]] + elif context_dim is None: + context_dim = [None] * depth + + self.in_channels = in_channels + inner_dim = n_heads * d_head + self.norm = normalization(in_channels) + if not use_linear: + self.proj_in = nn.Conv2d(in_channels, + inner_dim, + kernel_size=1, + stride=1, + padding=0) + else: + self.proj_in = nn.Linear(in_channels, inner_dim) + + self.transformer_blocks = nn.ModuleList([ + BasicTransformerBlock(inner_dim, + n_heads, + d_head, + dropout=dropout, + context_dim=context_dim[d], + disable_self_attn=disable_self_attn, + use_checkpoint=use_checkpoint) + for d in range(depth) + ]) + if not use_linear: + self.proj_out = zero_module( + nn.Conv2d(inner_dim, + in_channels, + kernel_size=1, + stride=1, + padding=0)) + else: + self.proj_out = zero_module(nn.Linear(in_channels, inner_dim)) + self.use_linear = use_linear + + def forward(self, x, context=None): + # note: if no context is given, cross-attention defaults to self-attention + if not isinstance(context, list): + context = [context] + b, c, h, w = x.shape + x_in = x + x = self.norm(x) + if not self.use_linear: + x = self.proj_in(x) + x = rearrange(x, 'b c h w -> b (h w) c').contiguous() + if self.use_linear: + x = self.proj_in(x) + for i, block in enumerate(self.transformer_blocks): + if i > 0 and len(context) == 1: + i = 0 # use same context for each block + x = block(x, context=context[i]) + if self.use_linear: + x = self.proj_out(x) + x = rearrange(x, 'b (h w) c -> b c h w', h=h, w=w).contiguous() + if not self.use_linear: + x = self.proj_out(x) + return x + x_in diff --git a/scepter/modules/model/backbone/utils/__init__.py b/scepter/modules/model/backbone/utils/__init__.py new file mode 100644 index 0000000..cc26a06 --- /dev/null +++ b/scepter/modules/model/backbone/utils/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/modules/model/backbone/utils/transformer.py b/scepter/modules/model/backbone/utils/transformer.py new file mode 100644 index 0000000..5a6e452 --- /dev/null +++ b/scepter/modules/model/backbone/utils/transformer.py @@ -0,0 +1,66 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from collections import OrderedDict + +import torch +from torch import nn + + +class LayerNorm(nn.LayerNorm): + """Subclass torch's LayerNorm to handle fp16.""" + def forward(self, x: torch.Tensor): + orig_type = x.dtype + ret = super().forward(x.type(torch.float32)) + return ret.type(orig_type) + + +class QuickGELU(nn.Module): + def forward(self, x: torch.Tensor): + return x * torch.sigmoid(1.702 * x) + + +class ResidualAttentionBlock(nn.Module): + def __init__(self, + d_model: int, + n_head: int, + attn_mask: torch.Tensor = None): + super().__init__() + + self.attn = nn.MultiheadAttention(d_model, n_head) + self.ln_1 = LayerNorm(d_model) + self.mlp = nn.Sequential( + OrderedDict([('c_fc', nn.Linear(d_model, d_model * 4)), + ('gelu', QuickGELU()), + ('c_proj', nn.Linear(d_model * 4, d_model))])) + self.ln_2 = LayerNorm(d_model) + self.attn_mask = attn_mask + + def attention(self, x: torch.Tensor): + self.attn_mask = self.attn_mask.to( + dtype=x.dtype, + device=x.device) if self.attn_mask is not None else None + return self.attn(x, x, x, need_weights=False, + attn_mask=self.attn_mask)[0] + + def forward(self, x: torch.Tensor): + x = x + self.attention(self.ln_1(x)) + x = x + self.mlp(self.ln_2(x)) + return x + + +class Transformer(nn.Module): + def __init__(self, + width: int, + layers: int, + heads: int, + attn_mask: torch.Tensor = None): + super().__init__() + self.width = width + self.layers = layers + self.resblocks = nn.Sequential(*[ + ResidualAttentionBlock(width, heads, attn_mask) + for _ in range(layers) + ]) + + def forward(self, x: torch.Tensor): + return self.resblocks(x) diff --git a/scepter/modules/model/backbone/video/__init__.py b/scepter/modules/model/backbone/video/__init__.py new file mode 100644 index 0000000..5c3d613 --- /dev/null +++ b/scepter/modules/model/backbone/video/__init__.py @@ -0,0 +1,7 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.model.backbone.video import bricks +from scepter.modules.model.backbone.video.resnet_3d import ResNet3D +from scepter.modules.model.backbone.video.video_transformer import ( + FactorizedVideoTransformer, VideoTransformer) diff --git a/scepter/modules/model/backbone/video/bricks/__init__.py b/scepter/modules/model/backbone/video/bricks/__init__.py new file mode 100644 index 0000000..0f5eda1 --- /dev/null +++ b/scepter/modules/model/backbone/video/bricks/__init__.py @@ -0,0 +1,14 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.model.backbone.video.bricks import stems +from scepter.modules.model.backbone.video.bricks.csn_branch import CSNBranch +from scepter.modules.model.backbone.video.bricks.non_local import NonLocal +from scepter.modules.model.backbone.video.bricks.r2d3d_branch import \ + R2D3DBranch +from scepter.modules.model.backbone.video.bricks.r2plus1d_branch import \ + R2Plus1DBranch +from scepter.modules.model.backbone.video.bricks.tada_conv import \ + TAdaConvBlockAvgPool +from scepter.modules.model.backbone.video.bricks.transformer_branch import ( + BaseTransformerLayer, TimesformerLayer) diff --git a/scepter/modules/model/backbone/video/bricks/base_branch.py b/scepter/modules/model/backbone/video/bricks/base_branch.py new file mode 100644 index 0000000..f65cc33 --- /dev/null +++ b/scepter/modules/model/backbone/video/bricks/base_branch.py @@ -0,0 +1,64 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from abc import ABCMeta, abstractmethod + +from scepter.modules.model.backbone.video.bricks.visualize_3d_module import \ + Visualize3DModule +from scepter.modules.utils.config import dict_to_yaml + + +class BaseBranch(Visualize3DModule, metaclass=ABCMeta): + para_dict = { + 'BRANCH_STYLE': { + 'value': 'simple_block', + 'description': 'the branch style, default: simple_block!' + }, + 'CONSTRUCT_BRANCH': { + 'value': True, + 'description': 'construct branch or not!' + } + } + + def __init__(self, cfg, logger=None): + super(BaseBranch, self).__init__(cfg, logger=logger) + self.branch_style = cfg.get('BRANCH_STYLE', 'simple_block') + construct_branch = cfg.get('CONSTRUCT_BRANCH', True) + if construct_branch: + self._construct_branch() + + def _construct_branch(self): + if self.branch_style == 'simple_block': + self._construct_simple_block() + elif self.branch_style == 'bottleneck': + self._construct_bottleneck() + + @abstractmethod + def _construct_simple_block(self): + return + + @abstractmethod + def _construct_bottleneck(self): + return + + @abstractmethod + def forward(self, x): + return + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BRANCH', + __class__.__name__, + BaseBranch.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/video/bricks/csn_branch.py b/scepter/modules/model/backbone/video/bricks/csn_branch.py new file mode 100644 index 0000000..57bef51 --- /dev/null +++ b/scepter/modules/model/backbone/video/bricks/csn_branch.py @@ -0,0 +1,132 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import torch.nn as nn + +from scepter.modules.model.backbone.video.bricks.base_branch import BaseBranch +from scepter.modules.model.registry import BRICKS +from scepter.modules.utils.config import Config, dict_to_yaml + + +@BRICKS.register_class() +class CSNBranch(BaseBranch): + para_dict = { + 'DIM_IN': { + 'value': 64, + 'description': "the branch's dim in!" + }, + 'NUM_FILTERS': { + 'value': 64, + 'description': 'the num of filter!' + }, + 'DOWNSAMPLING': { + 'value': True, + 'description': 'downsample spatial data or not!' + }, + 'DOWNSAMPLING_TEMPORAL': { + 'value': True, + 'description': 'downsample temporal data or not!' + }, + 'EXPANISION_RATIO': { + 'value': 2, + 'description': 'expanision ratio for this branch!' + }, + 'BN_PARAMS': { + 'value': + None, + 'description': + 'bn params data, key/value align with torch.BatchNorm3d/2d/1d!' + } + } + para_dict.update(BaseBranch.para_dict) + + def __init__(self, cfg, logger=None): + self.dim_in = cfg.DIM_IN + self.num_filters = cfg.NUM_FILTERS + self.kernel_size = cfg.KERNEL_SIZE + self.downsampling = cfg.get('DOWNSAMPLING', True) + self.downsampling_temporal = cfg.get('DOWNSAMPLING_TEMPORAL', True) + self.expansion_ratio = cfg.get('EXPANISION_RATIO', 2) + # bn_params or {} + self.bn_params = cfg.get('BN_PARAMS', None) or dict() + if isinstance(self.bn_params, Config): + self.bn_params = self.bn_params.__dict__ + if self.downsampling: + if self.downsampling_temporal: + self.stride = (2, 2, 2) + else: + self.stride = (1, 2, 2) + else: + self.stride = (1, 1, 1) + super(CSNBranch, self).__init__(cfg, logger=logger) + + def _construct_simple_block(self): + raise NotImplementedError + + def _construct_bottleneck(self): + self.a = nn.Conv3d(in_channels=self.dim_in, + out_channels=self.num_filters // + self.expansion_ratio, + kernel_size=(1, 1, 1), + stride=(1, 1, 1), + padding=0, + bias=False) + self.a_bn = nn.BatchNorm3d(self.num_filters // self.expansion_ratio, + **self.bn_params) + self.a_relu = nn.ReLU(inplace=True) + + self.b = nn.Conv3d( + in_channels=self.num_filters // self.expansion_ratio, + out_channels=self.num_filters // self.expansion_ratio, + kernel_size=self.kernel_size, + stride=self.stride, + padding=[ + self.kernel_size[0] // 2, self.kernel_size[1] // 2, + self.kernel_size[2] // 2 + ], + bias=False, + groups=self.num_filters // self.expansion_ratio) + self.b_bn = nn.BatchNorm3d(self.num_filters // self.expansion_ratio, + **self.bn_params) + self.b_relu = nn.ReLU(inplace=True) + + self.c = nn.Conv3d(in_channels=self.num_filters // + self.expansion_ratio, + out_channels=self.num_filters, + kernel_size=(1, 1, 1), + stride=(1, 1, 1), + padding=0, + bias=False) + self.c_bn = nn.BatchNorm3d(self.num_filters, **self.bn_params) + + def forward(self, x): + if self.branch_style == 'bottleneck': + x = self.a(x) + x = self.a_bn(x) + x = self.a_relu(x) + + x = self.b(x) + x = self.b_bn(x) + x = self.b_relu(x) + + x = self.c(x) + x = self.c_bn(x) + return x + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BRANCH', + __class__.__name__, + CSNBranch.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/video/bricks/non_local.py b/scepter/modules/model/backbone/video/bricks/non_local.py new file mode 100644 index 0000000..61cf92f --- /dev/null +++ b/scepter/modules/model/backbone/video/bricks/non_local.py @@ -0,0 +1,112 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +""" NonLocal block. """ + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from scepter.modules.model.backbone.video.bricks.visualize_3d_module import \ + Visualize3DModule +from scepter.modules.model.registry import BRICKS +from scepter.modules.utils.config import Config, dict_to_yaml + + +@BRICKS.register_class() +class NonLocal(Visualize3DModule): + """ + Non-local block. + + See Xiaolong Wang et al. + Non-local Neural Networks. + """ + para_dict = { + 'DIM_IN': { + 'value': 64, + 'description': "the branch's dim in!" + }, + 'NUM_FILTERS': { + 'value': 64, + 'description': 'the num of filter!' + }, + 'BN_PARAMS': { + 'value': + None, + 'description': + 'bn params data, key/value align with torch.BatchNorm3d/2d/1d!' + } + } + para_dict.update(Visualize3DModule.para_dict) + + def __init__(self, cfg, logger=None): + super(NonLocal, self).__init__(cfg, logger=logger) + self.dim_in = cfg.DIM_IN + self.num_filters = cfg.NUM_FILTERS + bn_params = cfg.get('BN_PARAMS', None) or dict() + if isinstance(bn_params, Config): + bn_params = bn_params.__dict__ + bn_params['eps'] = 1e-5 + self.dim_middle = self.dim_in // 2 + + self.qconv = nn.Conv3d(self.dim_in, + self.dim_middle, + kernel_size=(1, 1, 1), + stride=(1, 1, 1), + padding=0) + + self.kconv = nn.Conv3d(self.dim_in, + self.dim_middle, + kernel_size=(1, 1, 1), + stride=(1, 1, 1), + padding=0) + + self.vconv = nn.Conv3d(self.dim_in, + self.dim_middle, + kernel_size=(1, 1, 1), + stride=(1, 1, 1), + padding=0) + + self.out_conv = nn.Conv3d( + self.dim_middle, + self.num_filters, + kernel_size=(1, 1, 1), + stride=(1, 1, 1), + padding=0, + ) + + self.out_bn = nn.BatchNorm3d(self.num_filters, **bn_params) + + def forward(self, x): + n, c, t, h, w = x.shape + + query = self.qconv(x).view(n, self.dim_middle, -1) + key = self.kconv(x).view(n, self.dim_middle, -1) + value = self.vconv(x).view(n, self.dim_middle, -1) + + attn = torch.einsum('nct,ncp->ntp', (query, key)) + attn = attn * (self.dim_middle**-0.5) + attn = F.softmax(attn, dim=2) + + out = torch.einsum('ntg,ncg->nct', (attn, value)) + out = out.view(n, self.dim_middle, t, h, w) + out = self.out_conv(out) + out = self.out_bn(out) + return x + out + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BRANCH', + __class__.__name__, + NonLocal.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/video/bricks/r2d3d_branch.py b/scepter/modules/model/backbone/video/bricks/r2d3d_branch.py new file mode 100644 index 0000000..d0386f8 --- /dev/null +++ b/scepter/modules/model/backbone/video/bricks/r2d3d_branch.py @@ -0,0 +1,167 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import torch.nn as nn + +from scepter.modules.model.backbone.video.bricks.base_branch import BaseBranch +from scepter.modules.model.registry import BRICKS +from scepter.modules.utils.config import Config, dict_to_yaml + + +@BRICKS.register_class() +class R2D3DBranch(BaseBranch): + para_dict = { + 'DIM_IN': { + 'value': 64, + 'description': "the branch's dim in!" + }, + 'NUM_FILTERS': { + 'value': 64, + 'description': 'the num of filter!' + }, + 'KERNEL_SIZE': { + 'value': [1, 7, 7], + 'description': 'the kernel size!' + }, + 'DOWNSAMPLING': { + 'value': True, + 'description': 'downsample spatial data or not!' + }, + 'DOWNSAMPLING_TEMPORAL': { + 'value': True, + 'description': 'downsample temporal data or not!' + }, + 'EXPANISION_RATIO': { + 'value': 2, + 'description': 'expanision ratio for this branch!' + }, + 'BN_PARAMS': { + 'value': + None, + 'description': + 'bn params data, key/value align with torch.BatchNorm3d/2d/1d!' + } + } + para_dict.update(BaseBranch.para_dict) + + def __init__(self, cfg, logger=None): + self.dim_in = cfg.DIM_IN + self.num_filters = cfg.NUM_FILTERS + self.kernel_size = cfg.KERNEL_SIZE + self.downsampling = cfg.get('DOWNSAMPLING', True) + self.downsampling_temporal = cfg.get('DOWNSAMPLING_TEMPORAL', True) + self.expansion_ratio = cfg.get('EXPANISION_RATIO', 2) + self.branch_style = cfg.get('BRANCH_STYLE', 'simple_block') + # bn_params or {} + self.bn_params = cfg.get('BN_PARAMS', None) or dict() + if isinstance(self.bn_params, Config): + self.bn_params = self.bn_params.__dict__ + if self.downsampling: + if self.downsampling_temporal: + self.stride = (2, 2, 2) + else: + self.stride = (1, 2, 2) + else: + self.stride = (1, 1, 1) + super(R2D3DBranch, self).__init__(cfg, logger=logger) + + def _construct_simple_block(self): + self.a = nn.Conv3d(in_channels=self.dim_in, + out_channels=self.num_filters, + kernel_size=self.kernel_size, + stride=self.stride, + padding=[ + self.kernel_size[0] // 2, + self.kernel_size[1] // 2, + self.kernel_size[2] // 2 + ], + bias=False) + self.a_bn = nn.BatchNorm3d(self.num_filters, **self.bn_params) + self.a_relu = nn.ReLU(inplace=True) + + self.b = nn.Conv3d(in_channels=self.num_filters, + out_channels=self.num_filters, + kernel_size=self.kernel_size, + stride=(1, 1, 1), + padding=[ + self.kernel_size[0] // 2, + self.kernel_size[1] // 2, + self.kernel_size[2] // 2 + ], + bias=False) + self.b_bn = nn.BatchNorm3d(self.num_filters, **self.bn_params) + + def _construct_bottleneck(self): + self.a = nn.Conv3d(in_channels=self.dim_in, + out_channels=self.num_filters // + self.expansion_ratio, + kernel_size=(1, 1, 1), + stride=(1, 1, 1), + padding=0, + bias=False) + self.a_bn = nn.BatchNorm3d(self.num_filters // self.expansion_ratio, + **self.bn_params) + self.a_relu = nn.ReLU(inplace=True) + + self.b = nn.Conv3d( + in_channels=self.num_filters // self.expansion_ratio, + out_channels=self.num_filters // self.expansion_ratio, + kernel_size=self.kernel_size, + stride=self.stride, + padding=[ + self.kernel_size[0] // 2, self.kernel_size[1] // 2, + self.kernel_size[2] // 2 + ], + bias=False) + self.b_bn = nn.BatchNorm3d(self.num_filters // self.expansion_ratio, + **self.bn_params) + self.b_relu = nn.ReLU(inplace=True) + + self.c = nn.Conv3d(in_channels=self.num_filters // + self.expansion_ratio, + out_channels=self.num_filters, + kernel_size=(1, 1, 1), + stride=(1, 1, 1), + padding=0, + bias=False) + self.c_bn = nn.BatchNorm3d(self.num_filters, **self.bn_params) + + def forward(self, x): + if self.branch_style == 'simple_block': + x = self.a(x) + x = self.a_bn(x) + x = self.a_relu(x) + + x = self.b(x) + x = self.b_bn(x) + return x + elif self.branch_style == 'bottleneck': + x = self.a(x) + x = self.a_bn(x) + x = self.a_relu(x) + + x = self.b(x) + x = self.b_bn(x) + x = self.b_relu(x) + + x = self.c(x) + x = self.c_bn(x) + return x + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BRANCH', + __class__.__name__, + R2D3DBranch.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/video/bricks/r2plus1d_branch.py b/scepter/modules/model/backbone/video/bricks/r2plus1d_branch.py new file mode 100644 index 0000000..e028ae8 --- /dev/null +++ b/scepter/modules/model/backbone/video/bricks/r2plus1d_branch.py @@ -0,0 +1,212 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import math + +import torch.nn as nn + +from scepter.modules.model.backbone.video.bricks.base_branch import BaseBranch +from scepter.modules.model.registry import BRICKS +from scepter.modules.utils.config import Config, dict_to_yaml + + +@BRICKS.register_class() +class R2Plus1DBranch(BaseBranch): + para_dict = { + 'DIM_IN': { + 'value': 64, + 'description': "the branch's dim in!" + }, + 'NUM_FILTERS': { + 'value': 64, + 'description': 'the num of filter!' + }, + 'DOWNSAMPLING': { + 'value': True, + 'description': 'downsample spatial data or not!' + }, + 'DOWNSAMPLING_TEMPORAL': { + 'value': True, + 'description': 'downsample temporal data or not!' + }, + 'EXPANISION_RATIO': { + 'value': 2, + 'description': 'expanision ratio for this branch!' + }, + 'BN_PARAMS': { + 'value': + None, + 'description': + 'bn params data, key/value align with torch.BatchNorm3d/2d/1d!' + } + } + para_dict.update(BaseBranch.para_dict) + + def __init__(self, cfg, logger=None): + self.dim_in = cfg.DIM_IN + self.num_filters = cfg.NUM_FILTERS + self.kernel_size = cfg.KERNEL_SIZE + self.downsampling = cfg.get('DOWNSAMPLING', True) + self.downsampling_temporal = cfg.get('DOWNSAMPLING_TEMPORAL', True) + self.expansion_ratio = cfg.get('EXPANISION_RATIO', 2) + # bn_params or {} + self.bn_params = cfg.get('BN_PARAMS', None) or dict() + if isinstance(self.bn_params, Config): + self.bn_params = self.bn_params.__dict__ + if self.downsampling: + if self.downsampling_temporal: + self.stride = (2, 2, 2) + else: + self.stride = (1, 2, 2) + else: + self.stride = (1, 1, 1) + super(R2Plus1DBranch, self).__init__(cfg, logger=logger) + + def _construct_simple_block(self): + mid_dim = int( + math.floor( + (self.kernel_size[0] * self.kernel_size[1] * + self.kernel_size[2] * self.dim_in * self.num_filters) / + (self.kernel_size[1] * self.kernel_size[2] * self.dim_in + + self.kernel_size[0] * self.num_filters))) + + self.a1 = nn.Conv3d(in_channels=self.dim_in, + out_channels=mid_dim, + kernel_size=(1, self.kernel_size[1], + self.kernel_size[2]), + stride=(1, self.stride[1], self.stride[2]), + padding=(0, self.kernel_size[1] // 2, + self.kernel_size[2] // 2), + bias=False) + self.a1_bn = nn.BatchNorm3d(mid_dim, **self.bn_params) + self.a1_relu = nn.ReLU(inplace=True) + + self.a2 = nn.Conv3d(in_channels=mid_dim, + out_channels=self.num_filters, + kernel_size=(self.kernel_size[0], 1, 1), + stride=(self.stride[0], 1, 1), + padding=(self.kernel_size[0] // 2, 0, 0), + bias=False) + self.a2_bn = nn.BatchNorm3d(self.num_filters, **self.bn_params) + self.a2_relu = nn.ReLU(inplace=True) + + mid_dim = int( + math.floor( + (self.kernel_size[0] * self.kernel_size[1] * + self.kernel_size[2] * self.num_filters * self.num_filters) / + (self.kernel_size[1] * self.kernel_size[2] * self.num_filters + + self.kernel_size[0] * self.num_filters))) + + self.b1 = nn.Conv3d(in_channels=self.num_filters, + out_channels=mid_dim, + kernel_size=(1, self.kernel_size[1], + self.kernel_size[2]), + stride=(1, 1, 1), + padding=(0, self.kernel_size[1] // 2, + self.kernel_size[2] // 2), + bias=False) + self.b1_bn = nn.BatchNorm3d(mid_dim, **self.bn_params) + self.b1_relu = nn.ReLU(inplace=True) + + self.b2 = nn.Conv3d(in_channels=mid_dim, + out_channels=self.num_filters, + kernel_size=(self.kernel_size[0], 1, 1), + stride=(1, 1, 1), + padding=(self.kernel_size[0] // 2, 0, 0), + bias=False) + self.b2_bn = nn.BatchNorm3d(self.num_filters, **self.bn_params) + + def _construct_bottleneck(self): + self.a = nn.Conv3d(in_channels=self.dim_in, + out_channels=self.num_filters // + self.expansion_ratio, + kernel_size=(1, 1, 1), + stride=(1, 1, 1), + padding=0, + bias=False) + self.a_bn = nn.BatchNorm3d(self.num_filters // self.expansion_ratio, + **self.bn_params) + self.a_relu = nn.ReLU(inplace=True) + + self.b1 = nn.Conv3d( + in_channels=self.num_filters // self.expansion_ratio, + out_channels=self.num_filters // self.expansion_ratio, + kernel_size=(1, self.kernel_size[1], self.kernel_size[2]), + stride=(1, self.stride[1], self.stride[2]), + padding=(0, self.kernel_size[1] // 2, self.kernel_size[2] // 2), + bias=False) + self.b1_bn = nn.BatchNorm3d(self.num_filters // self.expansion_ratio, + **self.bn_params) + self.b1_relu = nn.ReLU(inplace=True) + + self.b2 = nn.Conv3d( + in_channels=self.num_filters // self.expansion_ratio, + out_channels=self.num_filters // self.expansion_ratio, + kernel_size=(self.kernel_size[0], 1, 1), + stride=(self.stride[0], 1, 1), + padding=(self.kernel_size[0] // 2, 0, 0), + bias=False) + self.b2_bn = nn.BatchNorm3d(self.num_filters // self.expansion_ratio, + **self.bn_params) + self.b2_relu = nn.ReLU(inplace=True) + + self.c = nn.Conv3d(in_channels=self.num_filters // + self.expansion_ratio, + out_channels=self.num_filters, + kernel_size=(1, 1, 1), + stride=(1, 1, 1), + padding=0, + bias=False) + self.c_bn = nn.BatchNorm3d(self.num_filters, **self.bn_params) + + def forward(self, x): + if self.branch_style == 'simple_block': + x = self.a1(x) + x = self.a1_bn(x) + x = self.a1_relu(x) + + x = self.a2(x) + x = self.a2_bn(x) + x = self.a2_relu(x) + + x = self.b1(x) + x = self.b1_bn(x) + x = self.b1_relu(x) + + x = self.b2(x) + x = self.b2_bn(x) + return x + elif self.branch_style == 'bottleneck': + x = self.a(x) + x = self.a_bn(x) + x = self.a_relu(x) + + x = self.b1(x) + x = self.b1_bn(x) + x = self.b1_relu(x) + + x = self.b2(x) + x = self.b2_bn(x) + x = self.b2_relu(x) + + x = self.c(x) + x = self.c_bn(x) + return x + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BRANCH', + __class__.__name__, + R2Plus1DBranch.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/video/bricks/stems/__init__.py b/scepter/modules/model/backbone/video/bricks/stems/__init__.py new file mode 100644 index 0000000..3081326 --- /dev/null +++ b/scepter/modules/model/backbone/video/bricks/stems/__init__.py @@ -0,0 +1,13 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.model.backbone.video.bricks.stems.base_2d_stem import \ + Base2DStem +from scepter.modules.model.backbone.video.bricks.stems.base_3d_stem import \ + Base3DStem +from scepter.modules.model.backbone.video.bricks.stems.down_sample_stem import \ + DownSampleStem +from scepter.modules.model.backbone.video.bricks.stems.embedding_stem import ( + PatchEmbedStem, TubeletEmbeddingStem) +from scepter.modules.model.backbone.video.bricks.stems.r2plus1d_stem import \ + R2Plus1DStem diff --git a/scepter/modules/model/backbone/video/bricks/stems/base_2d_stem.py b/scepter/modules/model/backbone/video/bricks/stems/base_2d_stem.py new file mode 100644 index 0000000..cfecd0b --- /dev/null +++ b/scepter/modules/model/backbone/video/bricks/stems/base_2d_stem.py @@ -0,0 +1,96 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import torch.nn as nn + +from scepter.modules.model.backbone.video.bricks.visualize_3d_module import \ + Visualize3DModule +from scepter.modules.model.registry import STEMS +from scepter.modules.utils.config import Config, dict_to_yaml + + +@STEMS.register_class() +class Base2DStem(Visualize3DModule): + para_dict = { + 'DIM_IN': { + 'value': 3, + 'description': "the stem's dim in!" + }, + 'NUM_FILTERS': { + 'value': 64, + 'description': 'the num of filter!' + }, + 'KERNEL_SIZE': { + 'value': [1, 7, 1], + 'description': 'the kernel size!' + }, + 'DOWNSAMPLING': { + 'value': True, + 'description': 'downsample spatial data or not!' + }, + 'DOWNSAMPLING_TEMPORAL': { + 'value': False, + 'description': 'downsample temporal data or not!' + }, + 'BN_PARAMS': { + 'value': + None, + 'description': + 'bn params data, key/value align with torch.BatchNorm3d/2d/1d!' + } + } + para_dict.update(Visualize3DModule.para_dict) + + def __init__(self, cfg, logger=None): + super(Base2DStem, self).__init__(cfg, logger=logger) + + self.dim_in = cfg.get('DIM_IN', 3) + self.num_filters = cfg.get('NUM_FILTERS', 64) + self.kernel_size = cfg.get('KERNEL_SIZE', [1, 7, 1]) + downsampling = cfg.get('DOWNSAMPLING', True) + downsampling_temporal = cfg.get('DOWNSAMPLING_TEMPORAL', False) + # bn_params or {} + self.bn_params = cfg.get('BN_PARAMS', None) or dict() + if isinstance(self.bn_params, Config): + self.bn_params = self.bn_params.__dict__ + if downsampling: + if downsampling_temporal: + self.stride = [2, 2, 2] + else: + self.stride = [1, 2, 2] + else: + self.stride = [1, 1, 1] + + self._construct() + + def _construct(self): + self.a = nn.Conv3d( + self.dim_in, + self.num_filters, + kernel_size=(1, self.kernel_size[1], self.kernel_size[2]), + stride=(1, self.stride[1], self.stride[2]), + padding=[0, self.kernel_size[1] // 2, self.kernel_size[2] // 2], + bias=False) + self.a_bn = nn.BatchNorm3d(self.num_filters, **self.bn_params) + self.a_relu = nn.ReLU(inplace=True) + + def forward(self, x): + return self.a_relu(self.a_bn(self.a(x))) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('STEM', + __class__.__name__, + Base2DStem.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/video/bricks/stems/base_3d_stem.py b/scepter/modules/model/backbone/video/bricks/stems/base_3d_stem.py new file mode 100644 index 0000000..be7e64c --- /dev/null +++ b/scepter/modules/model/backbone/video/bricks/stems/base_3d_stem.py @@ -0,0 +1,99 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import torch.nn as nn + +from scepter.modules.model.backbone.video.bricks.visualize_3d_module import \ + Visualize3DModule +from scepter.modules.model.registry import STEMS +from scepter.modules.utils.config import Config, dict_to_yaml + + +@STEMS.register_class() +class Base3DStem(Visualize3DModule): + para_dict = { + 'DIM_IN': { + 'value': 3, + 'description': "the stem's dim in!" + }, + 'NUM_FILTERS': { + 'value': 64, + 'description': 'the num of filter!' + }, + 'KERNEL_SIZE': { + 'value': [1, 7, 7], + 'description': 'the kernel size!' + }, + 'DOWNSAMPLING': { + 'value': True, + 'description': 'downsample spatial data or not!' + }, + 'DOWNSAMPLING_TEMPORAL': { + 'value': False, + 'description': 'downsample temporal data or not!' + }, + 'BN_PARAMS': { + 'value': + None, + 'description': + 'bn params data, key/value align with torch.BatchNorm3d/2d/1d!' + } + } + para_dict.update(Visualize3DModule.para_dict) + + def __init__(self, cfg, logger=None): + super(Base3DStem, self).__init__(cfg, logger=logger) + + self.dim_in = cfg.get('DIM_IN', 3) + self.num_filters = cfg.get('NUM_FILTERS', 64) + self.kernel_size = cfg.get('KERNEL_SIZE', [1, 7, 1]) + downsampling = cfg.get('DOWNSAMPLING', True) + downsampling_temporal = cfg.get('DOWNSAMPLING_TEMPORAL', False) + # bn_params or {} + self.bn_params = cfg.get('BN_PARAMS', None) or dict() + if isinstance(self.bn_params, Config): + self.bn_params = self.bn_params.__dict__ + if downsampling: + if downsampling_temporal: + self.stride = (2, 2, 2) + else: + self.stride = (1, 2, 2) + else: + self.stride = (1, 1, 1) + + self._construct() + + def _construct(self): + self.a = nn.Conv3d(self.dim_in, + self.num_filters, + kernel_size=self.kernel_size, + stride=self.stride, + padding=[ + self.kernel_size[0] // 2, + self.kernel_size[1] // 2, + self.kernel_size[2] // 2 + ], + bias=False) + self.a_bn = nn.BatchNorm3d(self.num_filters, **self.bn_params) + self.a_relu = nn.ReLU(inplace=True) + + def forward(self, x): + return self.a_relu(self.a_bn(self.a(x))) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('STEM', + __class__.__name__, + Base3DStem.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/video/bricks/stems/down_sample_stem.py b/scepter/modules/model/backbone/video/bricks/stems/down_sample_stem.py new file mode 100644 index 0000000..cf3d332 --- /dev/null +++ b/scepter/modules/model/backbone/video/bricks/stems/down_sample_stem.py @@ -0,0 +1,42 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import torch.nn as nn + +from scepter.modules.model.backbone.video.bricks.stems.base_3d_stem import \ + Base3DStem +from scepter.modules.model.registry import STEMS +from scepter.modules.utils.config import dict_to_yaml + + +@STEMS.register_class() +class DownSampleStem(Base3DStem): + para_dict = {} + para_dict.update(Base3DStem.para_dict) + + def __init__(self, cfg, logger=None): + super(DownSampleStem, self).__init__(cfg, logger=logger) + self.maxpool = nn.MaxPool3d(kernel_size=(1, 3, 3), + stride=(1, 2, 2), + padding=(0, 1, 1)) + + def forward(self, x): + return self.maxpool(self.a_relu(self.a_bn(self.a(x)))) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('STEM', + __class__.__name__, + DownSampleStem.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/video/bricks/stems/embedding_stem.py b/scepter/modules/model/backbone/video/bricks/stems/embedding_stem.py new file mode 100644 index 0000000..725a308 --- /dev/null +++ b/scepter/modules/model/backbone/video/bricks/stems/embedding_stem.py @@ -0,0 +1,167 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import torch.nn as nn + +from scepter.modules.model.backbone.video.bricks.visualize_3d_module import \ + Visualize3DModule +from scepter.modules.model.registry import STEMS +from scepter.modules.utils.config import dict_to_yaml + + +@STEMS.register_class() +class PatchEmbedStem(Visualize3DModule): + para_dict = { + 'IMAGE_SIZE': { + 'value': 224, + 'description': "the stem's input frame size!" + }, + 'PATCH_SIZE': { + 'value': 16, + 'description': "the stem's input patch size!" + }, + 'NUM_FRAMES': { + 'value': 16, + 'description': "the stem's input frame num!" + }, + 'NUM_INPUT_CHANNELS': { + 'value': 3, + 'description': "the stem's input channels num!" + }, + 'DIM': { + 'value': 768, + 'description': "the stem's input dim!" + } + } + para_dict.update(Visualize3DModule.para_dict) + + def __init__(self, cfg, logger=None): + super(PatchEmbedStem, self).__init__(cfg, logger=logger) + image_size = cfg.get('IMAGE_SIZE', 224) + patch_size = cfg.get('PATCH_SIZE', 16) + num_frames = cfg.get('NUM_FRAMES', 16) + num_input_channels = cfg.get('NUM_INPUT_CHANNELS', 3) + dim = cfg.get('DIM', 768) + num_patches_per_image = (image_size // patch_size)**2 + num_patches = num_patches_per_image * num_frames + + self.image_size = image_size + self.patch_size = patch_size + self.num_frames = num_frames + self.num_patches = num_patches + + self.conv1 = nn.Conv3d(in_channels=num_input_channels, + out_channels=dim, + kernel_size=(1, patch_size, patch_size), + stride=(1, patch_size, patch_size), + bias=False) + + def forward(self, x): + h, w, p = x.shape[3], x.shape[4], self.patch_size + assert h % p == 0 and w % p == 0, f'height {h} and width {w} of video must be divisible by the patch size {p}' + x = self.conv1(x) + # b, c, t, h, w -> b, c, p (p: num patches) + x = x.reshape(x.shape[0], x.shape[1], -1) + # b, c, p -> b, p, c + x = x.permute(0, 2, 1) + return x + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('STEM', + __class__.__name__, + PatchEmbedStem.para_dict, + set_name=True) + + +@STEMS.register_class() +class TubeletEmbeddingStem(Visualize3DModule): + para_dict = { + 'IMAGE_SIZE': { + 'value': 224, + 'description': "the stem's input frame size!" + }, + 'PATCH_SIZE': { + 'value': 16, + 'description': "the stem's input patch size!" + }, + 'NUM_FRAMES': { + 'value': 16, + 'description': "the stem's input frame num!" + }, + 'NUM_INPUT_CHANNELS': { + 'value': 3, + 'description': "the stem's input channels num!" + }, + 'TUBELET_SIZE': { + 'value': 2, + 'description': "the stem's tubelet size!" + }, + 'DIM': { + 'value': 768, + 'description': "the stem's input dim!" + } + } + para_dict.update(Visualize3DModule.para_dict) + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + image_size = cfg.get('IMAGE_SIZE', 224) + patch_size = cfg.get('PATCH_SIZE', 16) + num_frames = cfg.get('NUM_FRAMES', 16) + num_input_channels = cfg.get('NUM_INPUT_CHANNELS', 3) + tubelet_size = cfg.get('TUBELET_SIZE', 2) + dim = cfg.get('DIM', 768) + num_patches_per_image = (image_size // patch_size)**2 + num_patches = num_patches_per_image * num_frames + + self.image_size = image_size + self.patch_size = patch_size + self.num_frames = num_frames + self.num_patches = num_patches + + self.conv1 = nn.Conv3d(in_channels=num_input_channels, + out_channels=dim, + kernel_size=(tubelet_size, patch_size, + patch_size), + stride=(tubelet_size, patch_size, patch_size), + bias=False) + + def forward(self, x): + h, w, p = x.shape[3], x.shape[4], self.patch_size + assert h % p == 0 and w % p == 0, f'height {h} and width {w} of video must be divisible by the patch size {p}' + x = self.conv1(x) + # b, c, t, h, w -> b, c, p (p: num patches) + x = x.reshape(x.shape[0], x.shape[1], -1) + # b, c, p -> b, p, c + x = x.permute(0, 2, 1) + return x + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('STEM', + __class__.__name__, + TubeletEmbeddingStem.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/video/bricks/stems/r2plus1d_stem.py b/scepter/modules/model/backbone/video/bricks/stems/r2plus1d_stem.py new file mode 100644 index 0000000..c1a5c69 --- /dev/null +++ b/scepter/modules/model/backbone/video/bricks/stems/r2plus1d_stem.py @@ -0,0 +1,76 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import math + +import torch.nn as nn + +from scepter.modules.model.backbone.video.bricks.stems.base_3d_stem import \ + Base3DStem +from scepter.modules.model.registry import STEMS +from scepter.modules.utils.config import dict_to_yaml + + +@STEMS.register_class() +class R2Plus1DStem(Base3DStem): + para_dict = {} + para_dict.update(Base3DStem.para_dict) + + def __init__(self, cfg, logger=None): + super(R2Plus1DStem, self).__init__(cfg, logger=logger) + + def _construct(self): + mid_dim = int( + math.floor( + (self.kernel_size[0] * self.kernel_size[1] * + self.kernel_size[2] * self.dim_in * self.num_filters) / + (self.kernel_size[1] * self.kernel_size[2] * self.dim_in + + self.kernel_size[0] * self.num_filters))) + + self.a1 = nn.Conv3d(in_channels=self.dim_in, + out_channels=mid_dim, + kernel_size=(1, self.kernel_size[1], + self.kernel_size[2]), + stride=(1, self.stride[1], self.stride[2]), + padding=(0, self.kernel_size[1] // 2, + self.kernel_size[2] // 2), + bias=False) + self.a1_bn = nn.BatchNorm3d(mid_dim, **self.bn_params) + self.a1_relu = nn.ReLU(inplace=True) + + self.a2 = nn.Conv3d(in_channels=mid_dim, + out_channels=self.num_filters, + kernel_size=(self.kernel_size[0], 1, 1), + stride=(self.stride[0], 1, 1), + padding=(self.kernel_size[0] // 2, 0, 0), + bias=False) + self.a2_bn = nn.BatchNorm3d(self.num_filters, **self.bn_params) + self.a2_relu = nn.ReLU(inplace=True) + + def forward(self, x): + x = self.a1(x) + x = self.a1_bn(x) + x = self.a1_relu(x) + + x = self.a2(x) + x = self.a2_bn(x) + x = self.a2_relu(x) + return x + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('STEM', + __class__.__name__, + R2Plus1DStem.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/video/bricks/tada_conv.py b/scepter/modules/model/backbone/video/bricks/tada_conv.py new file mode 100644 index 0000000..c673ab1 --- /dev/null +++ b/scepter/modules/model/backbone/video/bricks/tada_conv.py @@ -0,0 +1,335 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import collections +import math +from itertools import repeat + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from scepter.modules.model.backbone.video.bricks.base_branch import BaseBranch +from scepter.modules.model.registry import BRICKS +from scepter.modules.utils.config import Config, dict_to_yaml + + +def _ntuple(n): + def parse(x): + if isinstance(x, collections.abc.Iterable): + return tuple(x) + return tuple(repeat(x, n)) + + return parse + + +_single = _ntuple(1) +_pair = _ntuple(2) +_triple = _ntuple(3) +_quadruple = _ntuple(4) + + +class TAdaConv2d(nn.Module): + """ Performs temporally adaptive 2D convolution. + Currently, only application on 5D tensors is supported, which makes TAdaConv2d + essentially a 3D convolution with temporal kernel size of 1. + """ + def __init__(self, + in_channels, + out_channels, + kernel_size, + stride=1, + padding=0, + dilation=1, + groups=1, + bias=True): + """ + Args: + in_channels (int): number of input channels. + out_channels (int): number of output channels. + kernel_size (Tuple[int]): kernel size of TAdaConv2d. + stride (int, Tuple[int]): stride for the convolution in TAdaConv2d. + padding (int, Tuple[int]): padding for the convolution in TAdaConv2d. + dilation (Tuple[int]): dilation of the convolution in TAdaConv2d. + groups (int): number of groups for TAdaConv2d. + bias (bool): whether to use bias in TAdaConv2d. + """ + super(TAdaConv2d, self).__init__() + + kernel_size = _triple(kernel_size) + stride = _triple(stride) + padding = _triple(padding) + dilation = _triple(dilation) + + assert kernel_size[0] == 1 + assert stride[0] == 1 + assert padding[0] == 0 + assert dilation[0] == 1 + + self.in_channels = in_channels + self.out_channels = out_channels + self.kernel_size = kernel_size + self.stride = stride + self.padding = padding + self.dilation = dilation + self.groups = groups + + # base weights (W_b) + self.weight = nn.Parameter( + torch.Tensor(1, 1, out_channels, in_channels // groups, + kernel_size[1], kernel_size[2])) + if bias: + self.bias = nn.Parameter(torch.Tensor(out_channels)) + else: + self.register_parameter('bias', None) + + nn.init.kaiming_uniform_(self.weight, a=math.sqrt(5)) + if self.bias is not None: + fan_in, _ = nn.init._calculate_fan_in_and_fan_out(self.weight) + bound = 1 / math.sqrt(fan_in) + nn.init.uniform_(self.bias, -bound, bound) + + def forward(self, x, alpha): + """ + Args: + x (tensor): feature to perform convolution on. + alpha (tensor): calibration weight for the base weights. + W_t = alpha_t * W_b + """ + _, _, c_out, c_in, kh, kw = self.weight.size() + b, c_in, t, h, w = x.size() + x = x.permute(0, 2, 1, 3, 4).reshape(1, -1, h, w) + + # alpha: B, C, T, H(1), W(1) -> B, T, C, H(1), W(1) -> B, T, 1, C, H(1), W(1) + # corresponding to calibrating the input channel + weight = (alpha.permute(0, 2, 1, 3, 4).unsqueeze(2) * + self.weight).reshape(-1, c_in, kh, kw) + + bias = None + if self.bias is not None: + raise NotImplementedError + else: + output = F.conv2d(x, + weight=weight, + bias=bias, + stride=self.stride[1:], + padding=self.padding[1:], + dilation=self.dilation[1:], + groups=self.groups * b * t) + + output = output.view(b, t, c_out, output.size(-2), + output.size(-1)).permute(0, 2, 1, 3, 4) + + return output + + def extra_repr(self) -> str: + return f'{self.in_channels}, {self.out_channels}, ' \ + f'kernel_size: {self.kernel_size}, stride: {self.stride}, ' \ + f'padding: {self.padding}, dilation: {self.dilation}, ' \ + f'groups: {self.groups}' + + +class RouteFuncMLP(nn.Module): + """ The routing function for generating the calibration weights. + """ + def __init__(self, c_in, ratio, kernels, bn_eps=1e-5, bn_mmt=0.1): + """ + Args: + c_in (int): number of input channels. + ratio (int): reduction ratio for the routing function. + kernels (list): temporal kernel size of the stacked 1D convolutions + """ + super(RouteFuncMLP, self).__init__() + self.c_in = c_in + self.avgpool = nn.AdaptiveAvgPool3d((None, 1, 1)) + self.globalpool = nn.AdaptiveAvgPool3d(1) + self.g = nn.Conv3d( + in_channels=c_in, + out_channels=c_in, + kernel_size=(1, 1, 1), + padding=0, + ) + self.a = nn.Conv3d( + in_channels=c_in, + out_channels=int(c_in // ratio), + kernel_size=(kernels[0], 1, 1), + padding=(kernels[0] // 2, 0, 0), + ) + self.bn = nn.BatchNorm3d(int(c_in // ratio), + eps=bn_eps, + momentum=bn_mmt) + self.relu = nn.ReLU(inplace=True) + self.b = nn.Conv3d(in_channels=int(c_in // ratio), + out_channels=c_in, + kernel_size=(kernels[1], 1, 1), + padding=(kernels[1] // 2, 0, 0), + bias=False) + self.b.skip_init = True + self.b.weight.data.zero_() # to make sure the initial values + # for the output is 1. + + def forward(self, x): + g = self.globalpool(x) + x = self.avgpool(x) + x = self.a(x + self.g(g)) + x = self.bn(x) + x = self.relu(x) + x = self.b(x) + 1 + return x + + +@BRICKS.register_class() +class TAdaConvBlockAvgPool(BaseBranch): + """ The TAdaConv branch with average pooling as the feature aggregation scheme. + + For details, see + Ziyuan Huang, Shiwei Zhang, Liang Pan, Zhiwu Qing, Mingqian Tang, Ziwei Liu, and Marcelo H. Ang Jr. + "TAda! Temporally-Adaptive Convolutions for Video Understanding." + """ + para_dict = { + 'DIM_IN': { + 'value': 64, + 'description': "the branch's dim in!" + }, + 'NUM_FILTERS': { + 'value': 64, + 'description': 'the num of filter!' + }, + 'DOWNSAMPLING': { + 'value': True, + 'description': 'downsample spatial data or not!' + }, + 'DOWNSAMPLING_TEMPORAL': { + 'value': True, + 'description': 'downsample temporal data or not!' + }, + 'EXPANISION_RATIO': { + 'value': 2, + 'description': 'expanision ratio for this branch!' + }, + 'ROUTE_FUNC_R': { + 'value': 4, + 'description': 'the route func r!' + }, + 'ROUTE_FUNC_K': { + 'value': [3, 3], + 'description': 'the route func k!' + }, + 'POOL_K': { + 'value': [3, 1, 1], + 'description': 'the pool k!' + }, + 'BN_PARAMS': { + 'value': + None, + 'description': + 'bn params data, key/value align with torch.BatchNorm3d/2d/1d!' + } + } + para_dict.update(BaseBranch.para_dict) + + def __init__(self, cfg, logger=None): + self.dim_in = cfg.DIM_IN + self.num_filters = cfg.NUM_FILTERS + self.kernel_size = cfg.KERNEL_SIZE + self.downsampling = cfg.get('DOWNSAMPLING', True) + self.downsampling_temporal = cfg.get('DOWNSAMPLING_TEMPORAL', True) + self.expansion_ratio = cfg.get('EXPANISION_RATIO', 2) + self.route_func_r = cfg.get('ROUTE_FUNC_R', 4) + self.route_func_k = cfg.get('ROUTE_FUNC_K', [3, 3]) + self.pool_k = cfg.get('POOL_K', [3, 1, 1]) + # bn_params or {} + self.bn_params = cfg.get('BN_PARAMS', None) or dict() + if isinstance(self.bn_params, Config): + self.bn_params = self.bn_params.__dict__ + if self.downsampling: + if self.downsampling_temporal: + self.stride = [2, 2, 2] + else: + self.stride = [1, 2, 2] + else: + self.stride = [1, 1, 1] + + super(TAdaConvBlockAvgPool, self).__init__(cfg, logger=logger) + + def _construct_simple_block(self): + raise NotImplementedError + + def _construct_bottleneck(self): + self.a = nn.Conv3d(in_channels=self.dim_in, + out_channels=self.num_filters // + self.expansion_ratio, + kernel_size=(1, 1, 1), + stride=(1, 1, 1), + padding=0, + bias=False) + self.a_bn = nn.BatchNorm3d(self.num_filters // self.expansion_ratio, + **self.bn_params) + self.a_relu = nn.ReLU(inplace=True) + + self.b = TAdaConv2d( + in_channels=self.num_filters // self.expansion_ratio, + out_channels=self.num_filters // self.expansion_ratio, + kernel_size=(1, self.kernel_size[1], self.kernel_size[2]), + stride=(1, self.stride[1], self.stride[2]), + padding=(0, self.kernel_size[1] // 2, self.kernel_size[2] // 2), + bias=False) + self.b_rf = RouteFuncMLP(c_in=self.num_filters // self.expansion_ratio, + ratio=self.route_func_r, + kernels=self.route_func_k) + self.b_bn = nn.BatchNorm3d(self.num_filters // self.expansion_ratio, + **self.bn_params) + self.b_avgpool = nn.AvgPool3d(kernel_size=self.pool_k, + stride=1, + padding=(self.pool_k[0] // 2, + self.pool_k[1] // 2, + self.pool_k[2] // 2)) + self.b_avgpool_bn = nn.BatchNorm3d( + self.num_filters // self.expansion_ratio, **self.bn_params) + self.b_avgpool_bn.skip_init = True + self.b_avgpool_bn.weight.data.zero_() + self.b_avgpool_bn.bias.data.zero_() + + self.b_relu = nn.ReLU(inplace=True) + + self.c = nn.Conv3d(in_channels=self.num_filters // + self.expansion_ratio, + out_channels=self.num_filters, + kernel_size=(1, 1, 1), + stride=(1, 1, 1), + padding=0, + bias=False) + self.c_bn = nn.BatchNorm3d(self.num_filters, **self.bn_params) + + def forward(self, x): + if self.branch_style == 'simple_block': + raise NotImplementedError + + x = self.a(x) + x = self.a_bn(x) + x = self.a_relu(x) + x = self.b(x, self.b_rf(x)) + x = self.b_bn(x) + self.b_avgpool_bn(self.b_avgpool(x)) + x = self.b_relu(x) + + x = self.c(x) + x = self.c_bn(x) + return x + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BRANCH', + __class__.__name__, + TAdaConvBlockAvgPool.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/video/bricks/transformer_branch.py b/scepter/modules/model/backbone/video/bricks/transformer_branch.py new file mode 100644 index 0000000..31314db --- /dev/null +++ b/scepter/modules/model/backbone/video/bricks/transformer_branch.py @@ -0,0 +1,363 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +# drop_path function & DropPath class & Attention class +# Modified from https://github.com/rwightman/pytorch-image-models/blob/master/timm/models/layers/drop.py +# +# Copyright 2019, Facebook, Inc +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import torch +import torch.nn as nn +import torch.nn.functional as F +from einops import rearrange, repeat + +from scepter.modules.model.backbone.video.bricks.visualize_3d_module import \ + Visualize3DModule +from scepter.modules.model.registry import BRICKS +from scepter.modules.utils.config import dict_to_yaml + + +def drop_path(x, drop_prob: float = 0., training: bool = False): + if drop_prob == 0. or not training: + return x + keep_prob = 1 - drop_prob + shape = (x.shape[0], ) + (1, ) * ( + x.ndim - 1) # work with diff dim tensors, not just 2D ConvNets + random_tensor = keep_prob + torch.rand( + shape, dtype=x.dtype, device=x.device) + random_tensor.floor_() # binarize + output = x.div(keep_prob) * random_tensor + return output + + +class DropPath(nn.Module): + """ + From https://github.com/rwightman/pytorch-image-models/blob/master/timm/models/layers/drop.py. + Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks). + """ + def __init__(self, drop_prob=None): + super(DropPath, self).__init__() + self.drop_prob = drop_prob + + def forward(self, x): + return drop_path(x, self.drop_prob, self.training) + + def extra_repr(self) -> str: + return f'drop_prob={self.drop_prob}' + + +class GEGLU(nn.Module): + def forward(self, x): + x, gates = x.chunk(2, dim=-1) + return x * F.gelu(gates) + + +class FeedForward(nn.Module): + def __init__(self, dim, mult=4, ff_dropout=0.): + super().__init__() + self.net = nn.Sequential( + nn.Linear(dim, dim * mult), + nn.GELU(), + nn.Dropout(ff_dropout), + nn.Linear(dim * mult, dim), + nn.Dropout(ff_dropout), + ) + + def forward(self, x): + return self.net(x) + + +class Attention(nn.Module): + """ + Self-attention module. + Currently supports both full self-attention on all the input tokens, + or only-spatial/only-temporal self-attention. + + See Anurag Arnab et al. + ViVIT: A Video Vision Transformer. + and + Gedas Bertasius, Heng Wang, Lorenzo Torresani. + Is Space-Time Attention All You Need for Video Understanding? + """ + def __init__( + self, + dim, + num_heads=12, + attn_dropout=0., + ff_dropout=0., + einops_from=None, + einops_to=None, + **einops_dims, + ): + super(Attention, self).__init__() + self.num_heads = num_heads + dim_head = dim // num_heads + self.scale = dim_head**-0.5 + + self.to_qkv = nn.Linear(dim, dim * 3) + self.attn_dropout = nn.Dropout(attn_dropout) + self.proj = nn.Linear(dim, dim) + self.ff_dropout = nn.Dropout(ff_dropout) + + if einops_from is not None and einops_to is not None: + self.partial = True + self.einops_from = einops_from + self.einops_to = einops_to + self.einops_dims = einops_dims + else: + self.partial = False + + def forward(self, x): + if self.partial: + return self.forward_partial( + x, + self.einops_from, + self.einops_to, + **self.einops_dims, + ) + B, N, C = x.shape + qkv = self.to_qkv(x).reshape(B, N, 3, self.num_heads, + C // self.num_heads).permute( + 2, 0, 3, 1, 4) + q, k, v = qkv[0], qkv[1], qkv[ + 2] # make torchscript happy (cannot use tensor as tuple) + + attn = (q @ k.transpose(-2, -1)) * self.scale + attn = attn.softmax(dim=-1) + attn = self.attn_dropout(attn) + + x = (attn @ v).transpose(1, 2).reshape(B, N, C) + x = self.proj(x) + x = self.ff_dropout(x) + return x + + def forward_partial(self, x, einops_from, einops_to, **einops_dims): + h = self.num_heads + q, k, v = self.to_qkv(x).chunk(3, dim=-1) + q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), + (q, k, v)) + + q *= self.scale + + # splice out classification token at index 1 + (cls_q, q_), (cls_k, k_), (cls_v, + v_) = map(lambda t: (t[:, 0:1], t[:, 1:]), + (q, k, v)) + + # let classification token attend to key / values of all patches across time and space + cls_attn = (cls_q @ k.transpose(1, 2)).softmax(-1) + cls_attn = self.attn_dropout(cls_attn) + cls_out = cls_attn @ v + + # rearrange across time or space + q_, k_, v_ = map( + lambda t: rearrange(t, f'{einops_from} -> {einops_to}', ** + einops_dims), (q_, k_, v_)) + + # expand cls token keys and values across time or space and concat + r = q_.shape[0] // cls_k.shape[0] + cls_k, cls_v = map(lambda t: repeat(t, 'b () d -> (b r) () d', r=r), + (cls_k, cls_v)) + + k_ = torch.cat((cls_k, k_), dim=1) + v_ = torch.cat((cls_v, v_), dim=1) + + # attention + attn = (q_ @ k_.transpose(1, 2)).softmax(-1) + attn = self.attn_dropout(attn) + x = attn @ v_ + + # merge back time or space + x = rearrange(x, f'{einops_to} -> {einops_from}', **einops_dims) + + # concat back the cls token + x = torch.cat((cls_out, x), dim=1) + + # merge back the head + x = rearrange(x, '(b h) n d -> b n (h d)', h=h) + + # combine head out + x = self.proj(x) + x = self.ff_dropout(x) + return x + + def extra_repr(self) -> str: + return f'partial={self.partial}, ' + \ + '' if not self.partial \ + else f'einops_from={self.einops_from}, einops_to={self.einops_to}, einops_dims={self.einops_dims}' + + +@BRICKS.register_class() +class BaseTransformerLayer(Visualize3DModule): + para_dict = { + 'DIM': { + 'value': 768, + 'description': "the num of transformer's input dim!" + }, + 'NUM_HEADS': { + 'value': 12, + 'description': "the num of transformer's head!" + }, + 'ATTN_DROPOUT': { + 'value': 0.0, + 'description': 'the attention dropout of transformer!' + }, + 'FF_DROPOUT': { + 'value': 0.0, + 'description': 'the ff dropout of transformer!' + }, + 'MLP_MULT': { + 'value': 4, + 'description': "the mlp's mult!" + }, + 'DROP_PATH_PROB': { + 'value': 0.0, + 'description': 'the drop path prob!' + } + } + para_dict.update(Visualize3DModule.para_dict) + + def __init__(self, cfg, logger=None): + super(BaseTransformerLayer, self).__init__(cfg, logger=logger) + dim = cfg.DIM + num_heads = cfg.get('NUM_HEADS', 12) + attn_dropout = cfg.get('ATTN_DROPOUT', 0.0) + ff_dropout = cfg.get('FF_DROPOUT', 0.0) + mlp_mult = cfg.get('MLP_MULT', 4) + drop_path_prob = cfg.get('DROP_PATH_PROB', 0.0) + self.norm = nn.LayerNorm(dim, eps=1e-6) + self.attn = Attention(dim, + num_heads=num_heads, + attn_dropout=attn_dropout, + ff_dropout=ff_dropout) + self.norm_ffn = nn.LayerNorm(dim, eps=1e-6) + self.ffn = FeedForward(dim, mult=mlp_mult, ff_dropout=ff_dropout) + self.drop_path = DropPath(drop_prob=drop_path_prob + ) if drop_path_prob > 0. else nn.Identity() + + def forward(self, x): + x = x + self.drop_path(self.attn(self.norm(x))) + x = x + self.drop_path(self.ffn(self.norm_ffn(x))) + return x + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BRANCH', + __class__.__name__, + BaseTransformerLayer.para_dict, + set_name=True) + + +@BRICKS.register_class() +class TimesformerLayer(Visualize3DModule): + para_dict = { + 'NUM_PATCHES': { + 'value': 49, + 'description': "the num of transformer's input patches!" + }, + 'NUM_FRAMES': { + 'value': 30, + 'description': "the num of transformer's input frames!" + }, + 'DIM': { + 'value': 768, + 'description': "the num of transformer's input dim!" + }, + 'NUM_HEADS': { + 'value': 12, + 'description': "the num of transformer's head!" + }, + 'ATTN_DROPOUT': { + 'value': 0.0, + 'description': 'the attention dropout of transformer!' + }, + 'FF_DROPOUT': { + 'value': 0.0, + 'description': 'the ff dropout of transformer!' + }, + 'DROP_PATH_PROB': { + 'value': 0.0, + 'description': 'the drop path prob!' + } + } + para_dict.update(Visualize3DModule.para_dict) + + def __init__(self, cfg, logger=None): + super(TimesformerLayer, self).__init__(cfg, logger=logger) + + num_patches = cfg.get('NUM_PATCHES', 49) + num_frames = cfg.get('NUM_FRAMES', 30) + dim = cfg.DIM + num_heads = cfg.get('NUM_HEADS', 12) + attn_dropout = cfg.get('ATTN_DROPOUT', 0.0) + ff_dropout = cfg.get('FF_DROPOUT', 0.0) + drop_path_prob = cfg.get('DROP_PATH_PROB', 0.0) + + self.norm_temporal = nn.LayerNorm(dim, eps=1e-6) + self.attn_temporal = Attention(dim, + num_heads=num_heads, + attn_dropout=attn_dropout, + ff_dropout=ff_dropout, + einops_from='b (f n) d', + einops_to='(b n) f d', + n=num_patches) + self.norm = nn.LayerNorm(dim, eps=1e-6) + self.attn = Attention(dim, + num_heads=num_heads, + attn_dropout=attn_dropout, + ff_dropout=ff_dropout, + einops_from='b (f n) d', + einops_to='(b f) n d', + f=num_frames) + self.norm_ffn = nn.LayerNorm(dim, eps=1e-6) + self.ffn = FeedForward(dim=dim, ff_dropout=ff_dropout) + + self.drop_path = DropPath( + drop_path_prob) if drop_path_prob > 0. else nn.Identity() + + def forward(self, x): + x = x + self.drop_path(self.attn_temporal(self.norm_temporal(x))) + x = x + self.drop_path(self.attn(self.norm(x))) + x = x + self.drop_path(self.ffn(self.norm_ffn(x))) + return x + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BRANCH', + __class__.__name__, + TimesformerLayer.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/video/bricks/visualize_3d_module.py b/scepter/modules/model/backbone/video/bricks/visualize_3d_module.py new file mode 100644 index 0000000..c22f63e --- /dev/null +++ b/scepter/modules/model/backbone/video/bricks/visualize_3d_module.py @@ -0,0 +1,79 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import os + +import torch.nn as nn + +from scepter.modules.utils.config import dict_to_yaml + + +class Visualize3DModule(nn.Module): + para_dict = { + 'VISUALIZE': { + 'value': False, + 'description': 'visualize the layer output or not!' + }, + 'VISUALIZE_OUTPUT_DIR': { + 'value': '', + 'description': 'the visualize output saved dir!' + }, + } + + def __init__(self, cfg, logger=None): + super(Visualize3DModule, self).__init__() + # visualize=False, visualize_output_dir="" + self.logger = logger + self.visualize = cfg.get('VISUALIZE', False) + self.visualize_output_dir = cfg.get('VISUALIZE_OUTPUT_DIR', '') + self.id = 0 + + def visualize_features(self, module, input_x, output_x): + """ + Visualizes and saves the normalized output features for the module. + """ + import matplotlib.pyplot as plt + if not self.visualize: + return + b, c, t, h, w = output_x.shape + xmin, xmax = output_x.min(1).values.unsqueeze(1), output_x.max( + 1).values.unsqueeze(1) + x_vis = ((output_x.detach() - xmin) / (xmax - xmin)).permute(0, 1, 3, 2, 4) \ + .reshape(b, c * h, t * w).detach().cpu().numpy() + if hasattr(self, 'stage_id'): + stage_id = self.stage_id + block_id = self.block_id + else: + stage_id = 0 + block_id = 0 + for i in range(b): + if not os.path.exists( + f'{self.visualize_output_dir}/im_{self.id + i}/'): + os.makedirs(f'{self.visualize_output_dir}/im_{self.id + i}/') + plt.imsave( + f'{self.base_output_dir}/' + f'im_{self.id + i}/layer_{stage_id}_{block_id}_feature.jpg', + x_vis[i]) + self.id += b + + def set_stage_block_id(self, stage_id, block_id): + setattr(self, 'stage_id', stage_id) + setattr(self, 'block_id', block_id) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BRANCH', + __class__.__name__, + Visualize3DModule.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/video/init_helper.py b/scepter/modules/model/backbone/video/init_helper.py new file mode 100644 index 0000000..ff69bfb --- /dev/null +++ b/scepter/modules/model/backbone/video/init_helper.py @@ -0,0 +1,186 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +# _no_grad_trunc_normal_ & trunc_normal_ & variance_scaling_ & lecun_normal_ +# Modified from https://github.com/rwightman/pytorch-image-models/blob/master/timm/models/layers/weight_init.py. + +# Copyright 2019 Ross Wightman +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# _init_transformer_weights & c2_msra_fill & _init_convnet_weights +# Modified from https://github.com/facebookresearch/SlowFast/blob/main/slowfast/utils/weight_init_helper.py. + +# Copyright 2019, Facebook, Inc +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import math +import warnings + +import torch +import torch.nn as nn +from torch.nn.init import _calculate_fan_in_and_fan_out + + +def _no_grad_trunc_normal_(tensor, mean, std, a, b): + # Cut & paste from PyTorch official master until it's in a few official releases - RW + # Method based on https://people.sc.fsu.edu/~jburkardt/presentations/truncated_normal.pdf + def norm_cdf(x): + # Computes standard normal cumulative distribution function + return (1. + math.erf(x / math.sqrt(2.))) / 2. + + if (mean < a - 2 * std) or (mean > b + 2 * std): + warnings.warn( + 'mean is more than 2 std from [a, b] in nn.init.trunc_normal_. ' + 'The distribution of values may be incorrect.', + stacklevel=2) + + with torch.no_grad(): + # Values are generated by using a truncated uniform distribution and + # then using the inverse CDF for the normal distribution. + # Get upper and lower cdf values + le = norm_cdf((a - mean) / std) + u = norm_cdf((b - mean) / std) + + # Uniformly fill tensor with values from [l, u], then translate to + # [2l-1, 2u-1]. + tensor.uniform_(2 * le - 1, 2 * u - 1) + + # Use inverse cdf transform for normal distribution to get truncated + # standard normal + tensor.erfinv_() + + # Transform to proper mean, std + tensor.mul_(std * math.sqrt(2.)) + tensor.add_(mean) + + # Clamp to ensure it's in the proper range + tensor.clamp_(min=a, max=b) + return tensor + + +def trunc_normal_(tensor, mean=0., std=1., a=-2., b=2.): + # type: (torch.Tensor, float, float, float, float) -> torch.Tensor + r"""Fills the input Tensor with values drawn from a truncated + normal distribution. The values are effectively drawn from the + normal distribution :math:`\mathcal{N}(\text{mean}, \text{std}^2)` + with values outside :math:`[a, b]` redrawn until they are within + the bounds. The method used for generating the random values works + best when :math:`a \leq \text{mean} \leq b`. + Args: + tensor: an n-dimensional `torch.Tensor` + mean: the mean of the normal distribution + std: the standard deviation of the normal distribution + a: the minimum cutoff value + b: the maximum cutoff value + Examples: + >>> w = torch.empty(3, 5) + >>> nn.init.trunc_normal_(w) + """ + return _no_grad_trunc_normal_(tensor, mean, std, a, b) + + +def variance_scaling_(tensor, scale=1.0, mode='fan_in', distribution='normal'): + fan_in, fan_out = _calculate_fan_in_and_fan_out(tensor) + if mode == 'fan_in': + denom = fan_in + elif mode == 'fan_out': + denom = fan_out + elif mode == 'fan_avg': + denom = (fan_in + fan_out) / 2 + + variance = scale / denom + + if distribution == 'truncated_normal': + # constant is stddev of standard normal truncated to (-2, 2) + trunc_normal_(tensor, std=math.sqrt(variance) / .87962566103423978) + elif distribution == 'normal': + tensor.normal_(std=math.sqrt(variance)) + elif distribution == 'uniform': + bound = math.sqrt(3 * variance) + tensor.uniform_(-bound, bound) + else: + raise ValueError(f'invalid distribution {distribution}') + + +def lecun_normal_(tensor): + variance_scaling_(tensor, mode='fan_in', distribution='truncated_normal') + + +def _init_transformer_weights(m): + if isinstance(m, nn.Linear): + trunc_normal_(m.weight, std=.02) + if m.bias is not None: + nn.init.zeros_(m.bias) + elif isinstance(m, nn.LayerNorm): + nn.init.zeros_(m.bias) + nn.init.ones_(m.weight) + + +def c2_msra_fill(module: nn.Module) -> None: + """ + Initialize `module.weight` using the "MSRAFill" implemented in Caffe2. + Also initializes `module.bias` to 0. + Args: + module (torch.nn.Module): module to initialize. + """ + # pyre-ignore + nn.init.kaiming_normal_(module.weight, mode='fan_out', nonlinearity='relu') + if module.bias is not None: # pyre-ignore + nn.init.constant_(module.bias, 0) + + +def _init_convnet_weights(model, fc_init_std=0.01, zero_init_final_bn=True): + """ + Performs ResNet style weight initialization. + Args: + fc_init_std (float): the expected standard deviation for fc layer. + zero_init_final_bn (bool): if True, zero initialize the final bn for + every bottleneck. + """ + for m in model.modules(): + if hasattr(m, 'skip_init'): + continue + if isinstance(m, nn.Conv3d) and not hasattr(m, 'linear'): + """ + Follow the initialization method proposed in: + {He, Kaiming, et al. + "Delving deep into rectifiers: Surpassing human-level + performance on imagenet classification." + arXiv preprint arXiv:1502.01852 (2015)} + """ + c2_msra_fill(m) + elif isinstance(m, nn.BatchNorm3d): + if (hasattr(m, 'transform_final_bn') and m.transform_final_bn + and zero_init_final_bn): + batchnorm_weight = 0.0 + else: + batchnorm_weight = 1.0 + if m.weight is not None: + m.weight.data.fill_(batchnorm_weight) + if m.bias is not None: + m.bias.data.zero_() + if isinstance(m, nn.Linear) or hasattr(m, 'linear'): + m.weight.data.normal_(mean=0.0, std=fc_init_std) + if m.bias is not None: + m.bias.data.zero_() diff --git a/scepter/modules/model/backbone/video/resnet_3d.py b/scepter/modules/model/backbone/video/resnet_3d.py new file mode 100644 index 0000000..e5194b7 --- /dev/null +++ b/scepter/modules/model/backbone/video/resnet_3d.py @@ -0,0 +1,841 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import torch.nn as nn + +from scepter.modules.model.backbone.video.bricks.non_local import NonLocal +from scepter.modules.model.backbone.video.init_helper import \ + _init_convnet_weights +from scepter.modules.model.registry import BACKBONES, BRICKS, STEMS +from scepter.modules.utils.config import Config, dict_to_yaml + +_n_conv_resnet = { + 10: (1, 1, 1, 1), + 16: (2, 2, 2, 1), + 18: (2, 2, 2, 2), + 26: (2, 2, 2, 2), + 34: (3, 4, 6, 3), + 50: (3, 4, 6, 3), + 101: (3, 4, 23, 3), + 152: (3, 8, 36, 3), +} + + +@BRICKS.register_class() +class Base3DBlock(nn.Module): + para_dict = { + 'DIM_IN': { + 'value': 64, + 'description': 'this block input dim!' + }, + 'NUM_FILTERS': { + 'value': 64, + 'description': 'this block input num filter!' + }, + 'KERNEL_SIZE': { + 'value': [1, 7, 7], + 'description': 'the kernel size!' + }, + 'DOWNSAMPLING': { + 'value': True, + 'description': "this block's downsampling!" + }, + 'DOWNSAMPLING_TEMPORAL': { + 'value': False, + 'description': "this block's downsampling temporal!" + }, + 'BN_PARAMS': { + 'value': + None, + 'description': + 'bn params data, key/value align with torch.BatchNorm3d/2d/1d!' + }, + 'BRANCH': { + 'NAME': { + 'value': + '', + 'description': + "the branch's params, which shared the parameters DIM_IN, " + 'NUM_FILTERS, DOWNSAMPLING, DOWNSAMPLING_TEMPORAL, BN_PARAMS!' + } + } + } + + def __init__(self, cfg, logger=None): + super(Base3DBlock, self).__init__() + dim_in = cfg.DIM_IN + num_filters = cfg.NUM_FILTERS + kernel_size = cfg.KERNEL_SIZE + downsampling = cfg.get('DOWNSAMPLING', True) + downsampling_temporal = cfg.get('DOWNSAMPLING_TEMPORAL', True) + # bn_params or {} + bn_params = cfg.get('BN_PARAMS', None) or dict() + if isinstance(bn_params, Config): + bn_params = bn_params.__dict__ + if dim_in != num_filters or downsampling: + if downsampling: + if downsampling_temporal: + _stride = (2, 2, 2) + else: + _stride = (1, 2, 2) + else: + _stride = (1, 1, 1) + self.short_cut = nn.Conv3d(dim_in, + num_filters, + kernel_size=(1, 1, 1), + stride=_stride, + padding=0, + bias=False) + self.short_cut_bn = nn.BatchNorm3d(num_filters, **(bn_params + or {})) + branch_cfg = cfg.get('BRANCH', None) + assert branch_cfg is not None + branch_cfg.BN_PARAMS = bn_params + branch_cfg.DOWNSAMPLING = downsampling + branch_cfg.DOWNSAMPLING_TEMPORAL = downsampling_temporal + branch_cfg.DIM_IN = dim_in + branch_cfg.NUM_FILTERS = num_filters + branch_cfg.KERNEL_SIZE = kernel_size + self.conv_branch = BRICKS.build(branch_cfg, logger=logger) + self.relu = nn.ReLU(inplace=True) + + def forward(self, x): + short_cut = x + if hasattr(self, 'short_cut'): + short_cut = self.short_cut_bn(self.short_cut(short_cut)) + x = self.relu(short_cut + self.conv_branch(x)) + return x + + def set_stage_block(self, stage_id, block_id): + if hasattr(self.conv_branch, 'set_stage_block_id'): + self.conv_branch.set_stage_block_id(stage_id, block_id) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BRANCH', + __class__.__name__, + Base3DBlock.para_dict, + set_name=True) + + +@BRICKS.register_class() +class Base3DResStage(nn.Module): + """ + ResNet Stage containing several blocks. + """ + para_dict = { + 'NUM_BLOCKS': { + 'value': 5, + 'description': 'this stage contains num of block!' + }, + 'USE_NON_LOCAL': { + 'value': False, + 'description': 'use non local or not!' + }, + 'NUM_FILTERS': { + 'value': 64, + 'description': 'this block input num filter!' + }, + 'NON_LOCAL': { + 'NAME': { + 'value': 'NonLocal', + 'description': 'the non local config!' + } + } + } + para_dict.update(Base3DBlock.para_dict) + + def __init__(self, cfg, logger=None): + super(Base3DResStage, self).__init__() + self.num_blocks = cfg.NUM_BLOCKS + use_non_local = cfg.get('USE_NON_LOCAL', False) + non_local_cfg = cfg.get('NON_LOCAL', None) + res_block = Base3DBlock(cfg, logger=logger) + self.add_module('res_{}'.format(1), res_block) + for i in range(self.num_blocks - 1): + dim_in = cfg.NUM_FILTERS + downsampling = False + cfg.DIM_IN = dim_in + cfg.DOWNSAMPLING = downsampling + res_block = Base3DBlock(cfg, logger=logger) + self.add_module('res_{}'.format(i + 2), res_block) + if use_non_local: + non_local = NonLocal(non_local_cfg, logger=logger) + self.add_module('nonlocal', non_local) + + def forward(self, x): + # performs computation on the convolutions + for i in range(self.num_blocks): + res_block = getattr(self, 'res_{}'.format(i + 1)) + x = res_block(x) + + # performs non-local operations if specified. + if hasattr(self, 'nonlocal'): + non_local = getattr(self, 'nonlocal') + x = non_local(x) + return x + + def set_stage_id(self, stage_id): + for i in range(self.num_blocks): + res_block = getattr(self, 'res_{}'.format(i + 1)) + res_block.set_stage_block_id(stage_id, i) + if hasattr(self, 'nonlocal'): + non_local = getattr(self, 'nonlocal') + non_local.set_stage_block_id(stage_id, self.num_blocks) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BRANCH', + __class__.__name__, + Base3DResStage.para_dict, + set_name=True) + + +@BACKBONES.register_class() +class ResNet3D(nn.Module): + para_dict = { + 'DEPTH': { + 'value': 18, + 'description': "resnet model's depth!" + }, + 'NUM_INPUT_CHANNELS': { + 'value': 3, + 'description': 'the input channels num for the model!' + }, + 'NUM_FILTERS': { + 'value': [64, 64, 128, 256, 256], + 'description': 'this num filter for each layer!' + }, + 'KERNEL_SIZE': { + 'value': [[1, 7, 7], [1, 3, 3], [1, 3, 3], [3, 3, 3], [3, 3, 3]], + 'description': 'the kernel size for each layer!' + }, + 'DOWNSAMPLING': { + 'value': [True, False, True, True, True], + 'description': 'the downsample status for each layer!' + }, + 'DOWNSAMPLING_TEMPORAL': { + 'value': [False, False, False, True, True], + 'description': 'the temporal downsample status for each layer!' + }, + 'USE_NON_LOCAL': { + 'value': [False, False, False, False, False], + 'description': 'use non local or not for each layer!' + }, + 'NON_LOCAL': { + 'value': None, + 'description': 'the non local paramters!' + }, + 'STEM': { + 'NAME': { + 'value': + 'DownSampleStem', + 'description': + 'use the shared parameters as DIM_IN, NUM_FILTERS, VISUAL_CFG,' + 'KERNEL_SIZE, DOWNSAMPLING, DOWNSAMPLING_TEMPORAL, BN_PARAMS' + } + }, + 'BRANCH': { + 'NAME': { + 'value': 'R2D3DBranch', + 'description': 'use the shared parameters as BRANCH_STYLE' + }, + 'EXPANSION_RATIO': { + 'value': 2, + 'description': 'the expansion ratio value' + } + }, + 'BN_PARAMS': { + 'value': + None, + 'description': + 'bn params data, key/value align with torch.BatchNorm3d/2d/1d!' + }, + 'INIT_CFG': { + 'value': + None, + 'description': + 'the parameters init config, including name key and default as kaiming!' + }, + 'VISUAL_CFG': { + 'value': None, + 'description': 'the visualize config' + } + } + + def __init__(self, cfg, logger=None): + super(ResNet3D, self).__init__() + depth = cfg.get('DEPTH', 18) + num_input_channels = cfg.get('NUM_INPUT_CHANNELS', 3) + num_filters = cfg.get('NUM_FILTERS', [64, 64, 128, 256, 256]) + kernel_size = cfg.get( + 'KERNEL_SIZE', + [[1, 7, 7], [1, 3, 3], [1, 3, 3], [3, 3, 3], [3, 3, 3]]) + downsampling = cfg.get('DOWNSAMPLING', [True, False, True, True, True]) + downsampling_temporal = cfg.get('DOWNSAMPLING_TEMPORAL', + [False, False, False, True, True]) + use_non_local = cfg.get('USE_NON_LOCAL', + [False, False, False, False, False]) + non_local_cfg = cfg.get('NON_LOCAL', None) + stem_cfg = cfg.get('STEM', None) + branch_cfg = cfg.get('BRANCH', None) + + init_cfg = cfg.get('INIT_CFG', None) or dict() + visual_cfg = cfg.get('VISUAL_CFG', None) or dict() + # bn_params or {} + bn_params = cfg.get('BN_PARAMS', None) or dict() + if isinstance(bn_params, Config): + bn_params = bn_params.__dict__ + if len(bn_params) < 1: + bn_params = dict(eps=1e-3, momentum=0.1) + # Build stem cfg + stem_cfg.DIM_IN = num_input_channels + stem_cfg.NUM_FILTERS = num_filters[0] + stem_cfg.KERNEL_SIZE = tuple(kernel_size[0]) + stem_cfg.DOWNSAMPLING = downsampling[0] + stem_cfg.DOWNSAMPLING_TEMPORAL = downsampling_temporal[0] + stem_cfg.BN_PARAMS = bn_params + stem_cfg.VISUAL_CFG = visual_cfg + self.conv1 = STEMS.build(stem_cfg, logger=logger) + self.conv1.set_stage_block_id(0, 0) + # ------------------- Main arch ------------------- + branch_style = 'simple_block' if depth <= 34 else 'bottleneck' + blocks_list = _n_conv_resnet[depth] + for stage_id, num_blocks in enumerate(blocks_list): + stage_id = stage_id + 1 + # Build branch cfg + block_cfg = Config(cfg_dict={}, load=False, logger=logger) + block_cfg.NUM_BLOCKS = num_blocks + block_cfg.USE_NON_LOCAL = use_non_local[stage_id] + block_cfg.NUM_FILTERS = num_filters[stage_id] + block_cfg.NON_LOCAL = non_local_cfg + block_cfg.DIM_IN = num_filters[stage_id - 1] + block_cfg.DOWNSAMPLING = downsampling[stage_id] + block_cfg.DOWNSAMPLING_TEMPORAL = downsampling_temporal[stage_id] + block_cfg.BRANCH = branch_cfg + block_cfg.BRANCH.BRANCH_STYLE = branch_style + block_cfg.KERNEL_SIZE = tuple(kernel_size[stage_id]) + conv = Base3DResStage(block_cfg, logger=None) + setattr(self, f'conv{stage_id + 1}', conv) + # perform initialization + if isinstance(init_cfg, Config): + init_cfg = init_cfg.__dict__ + if init_cfg.get('name') == 'kaiming': + _init_convnet_weights(self) + + def forward(self, video): + x = self.conv1(video) + for i in range(2, 6): + x = getattr(self, f'conv{i}')(x) + return x + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BACKBONE', + __class__.__name__, + ResNet3D.para_dict, + set_name=True) + + +@BACKBONES.register_class() +class ResNet3D_2plus1d(nn.Module): + para_dict = { + 'DEPTH': { + 'value': 18, + 'description': "resnet model's depth!" + }, + 'NUM_INPUT_CHANNELS': { + 'value': 3, + 'description': 'the input channels num for the model!' + }, + 'NUM_FILTERS': { + 'value': [64, 64, 128, 256, 512], + 'description': 'this num filter for each layer!' + }, + 'KERNEL_SIZE': { + 'value': [[3, 7, 7], [3, 3, 3], [3, 3, 3], [3, 3, 3], [3, 3, 3]], + 'description': 'the kernel size for each layer!' + }, + 'DOWNSAMPLING': { + 'value': [True, False, True, True, True], + 'description': 'the downsample status for each layer!' + }, + 'DOWNSAMPLING_TEMPORAL': { + 'value': [False, False, True, True, True], + 'description': 'the temporal downsample status for each layer!' + }, + 'USE_NON_LOCAL': { + 'value': [False, False, False, False, False], + 'description': 'use non local or not for each layer!' + }, + 'NON_LOCAL': { + 'value': None, + 'description': 'the non local paramters!' + }, + 'STEM': { + 'NAME': { + 'value': + 'R2Plus1DStem', + 'description': + 'use the shared parameters as DIM_IN, NUM_FILTERS, VISUAL_CFG,' + 'KERNEL_SIZE, DOWNSAMPLING, DOWNSAMPLING_TEMPORAL, BN_PARAMS' + } + }, + 'BRANCH': { + 'NAME': { + 'value': 'R2Plus1DBranch', + 'description': 'use the shared parameters as BRANCH_STYLE' + }, + 'EXPANSION_RATIO': { + 'value': 2, + 'description': 'the expansion ratio value' + } + }, + 'BN_PARAMS': { + 'value': + None, + 'description': + 'bn params data, key/value align with torch.BatchNorm3d/2d/1d!' + }, + 'INIT_CFG': { + 'value': + None, + 'description': + 'the parameters init config, including name key and default as kaiming!' + }, + 'VISUAL_CFG': { + 'value': None, + 'description': 'the visualize config' + } + } + + def __init__(self, cfg, logger=None): + super(ResNet3D_2plus1d, self).__init__() + depth = cfg.get('DEPTH', 18) + num_input_channels = cfg.get('NUM_INPUT_CHANNELS', 3) + num_filters = cfg.get('NUM_FILTERS', [64, 64, 128, 256, 512]) + kernel_size = cfg.get( + 'KERNEL_SIZE', + [[3, 7, 7], [3, 3, 3], [3, 3, 3], [3, 3, 3], [3, 3, 3]]) + downsampling = cfg.get('DOWNSAMPLING', [True, False, True, True, True]) + downsampling_temporal = cfg.get('DOWNSAMPLING_TEMPORAL', + [False, False, False, True, True]) + use_non_local = cfg.get('USE_NON_LOCAL', + [False, False, False, False, False]) + non_local_cfg = cfg.get('NON_LOCAL', None) + stem_cfg = cfg.get('STEM', None) + branch_cfg = cfg.get('BRANCH', None) + + init_cfg = cfg.get('INIT_CFG', None) or dict() + visual_cfg = cfg.get('VISUAL_CFG', None) or dict() + # bn_params or {} + bn_params = cfg.get('BN_PARAMS', None) or dict() + if isinstance(bn_params, Config): + bn_params = bn_params.__dict__ + if len(bn_params) < 1: + bn_params = dict(eps=1e-3, momentum=0.1) + # Build stem cfg + stem_cfg.DIM_IN = num_input_channels + stem_cfg.NUM_FILTERS = num_filters[0] + stem_cfg.KERNEL_SIZE = tuple(kernel_size[0]) + stem_cfg.DOWNSAMPLING = downsampling[0] + stem_cfg.DOWNSAMPLING_TEMPORAL = downsampling_temporal[0] + stem_cfg.BN_PARAMS = bn_params + stem_cfg.VISUAL_CFG = visual_cfg + self.conv1 = STEMS.build(stem_cfg, logger=logger) + self.conv1.set_stage_block_id(0, 0) + # ------------------- Main arch ------------------- + branch_style = 'simple_block' if depth <= 34 else 'bottleneck' + blocks_list = _n_conv_resnet[depth] + for stage_id, num_blocks in enumerate(blocks_list): + stage_id = stage_id + 1 + # Build branch cfg + block_cfg = Config(cfg_dict={}, load=False, logger=logger) + block_cfg.NUM_BLOCKS = num_blocks + block_cfg.USE_NON_LOCAL = use_non_local[stage_id] + block_cfg.NUM_FILTERS = num_filters[stage_id] + block_cfg.NON_LOCAL = non_local_cfg + block_cfg.DIM_IN = num_filters[stage_id - 1] + block_cfg.DOWNSAMPLING = downsampling[stage_id] + block_cfg.DOWNSAMPLING_TEMPORAL = downsampling_temporal[stage_id] + block_cfg.BRANCH = branch_cfg + block_cfg.BRANCH.BRANCH_STYLE = branch_style + block_cfg.KERNEL_SIZE = tuple(kernel_size[stage_id]) + conv = Base3DResStage(block_cfg, logger=None) + setattr(self, f'conv{stage_id + 1}', conv) + # perform initialization + if isinstance(init_cfg, Config): + init_cfg = init_cfg.__dict__ + if init_cfg.get('name') == 'kaiming': + _init_convnet_weights(self) + + def forward(self, video): + x = self.conv1(video) + for i in range(2, 6): + x = getattr(self, f'conv{i}')(x) + return x + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BACKBONE', + __class__.__name__, + ResNet3D_2plus1d.para_dict, + set_name=True) + + +@BACKBONES.register_class() +class ResNet3D_CSN(nn.Module): + para_dict = { + 'DEPTH': { + 'value': 18, + 'description': "resnet model's depth!" + }, + 'NUM_INPUT_CHANNELS': { + 'value': 3, + 'description': 'the input channels num for the model!' + }, + 'NUM_FILTERS': { + 'value': [64, 256, 512, 1024, 2048], + 'description': 'this num filter for each layer!' + }, + 'KERNEL_SIZE': { + 'value': [[3, 7, 7], [3, 3, 3], [3, 3, 3], [3, 3, 3], [3, 3, 3]], + 'description': 'the kernel size for each layer!' + }, + 'DOWNSAMPLING': { + 'value': [True, False, True, True, True], + 'description': 'the downsample status for each layer!' + }, + 'DOWNSAMPLING_TEMPORAL': { + 'value': [False, False, True, True, True], + 'description': 'the temporal downsample status for each layer!' + }, + 'USE_NON_LOCAL': { + 'value': [False, False, False, False, False], + 'description': 'use non local or not for each layer!' + }, + 'NON_LOCAL': { + 'value': None, + 'description': 'the non local paramters!' + }, + 'STEM': { + 'NAME': { + 'value': + 'DownSampleStem', + 'description': + 'use the shared parameters as DIM_IN, NUM_FILTERS, VISUAL_CFG,' + 'KERNEL_SIZE, DOWNSAMPLING, DOWNSAMPLING_TEMPORAL, BN_PARAMS' + } + }, + 'BRANCH': { + 'NAME': { + 'value': 'CSNBranch', + 'description': 'use the shared parameters as BRANCH_STYLE' + }, + 'EXPANSION_RATIO': { + 'value': 4, + 'description': 'the expansion ratio value' + } + }, + 'BN_PARAMS': { + 'value': + None, + 'description': + 'bn params data, key/value align with torch.BatchNorm3d/2d/1d!' + }, + 'INIT_CFG': { + 'value': + None, + 'description': + 'the parameters init config, including name key and default as kaiming!' + }, + 'VISUAL_CFG': { + 'value': None, + 'description': 'the visualize config' + } + } + + def __init__(self, cfg, logger=None): + super(ResNet3D_CSN, self).__init__() + depth = cfg.get('DEPTH', 18) + num_input_channels = cfg.get('NUM_INPUT_CHANNELS', 3) + num_filters = cfg.get('NUM_FILTERS', [64, 64, 128, 256, 256]) + kernel_size = cfg.get( + 'KERNEL_SIZE', + [[1, 7, 7], [1, 3, 3], [1, 3, 3], [3, 3, 3], [3, 3, 3]]) + downsampling = cfg.get('DOWNSAMPLING', [True, False, True, True, True]) + downsampling_temporal = cfg.get('DOWNSAMPLING_TEMPORAL', + [False, False, False, True, True]) + use_non_local = cfg.get('USE_NON_LOCAL', + [False, False, False, False, False]) + non_local_cfg = cfg.get('NON_LOCAL', None) + stem_cfg = cfg.get('STEM', None) + branch_cfg = cfg.get('BRANCH', None) + + init_cfg = cfg.get('INIT_CFG', None) or dict() + visual_cfg = cfg.get('VISUAL_CFG', None) or dict() + # bn_params or {} + bn_params = cfg.get('BN_PARAMS', None) or dict() + if isinstance(bn_params, Config): + bn_params = bn_params.__dict__ + if len(bn_params) < 1: + bn_params = dict(eps=1e-3, momentum=0.1) + # Build stem cfg + stem_cfg.DIM_IN = num_input_channels + stem_cfg.NUM_FILTERS = num_filters[0] + stem_cfg.KERNEL_SIZE = tuple(kernel_size[0]) + stem_cfg.DOWNSAMPLING = downsampling[0] + stem_cfg.DOWNSAMPLING_TEMPORAL = downsampling_temporal[0] + stem_cfg.BN_PARAMS = bn_params + stem_cfg.VISUAL_CFG = visual_cfg + self.conv1 = STEMS.build(stem_cfg, logger=logger) + self.conv1.set_stage_block_id(0, 0) + # ------------------- Main arch ------------------- + branch_style = 'simple_block' if depth <= 34 else 'bottleneck' + blocks_list = _n_conv_resnet[depth] + for stage_id, num_blocks in enumerate(blocks_list): + stage_id = stage_id + 1 + # Build branch cfg + block_cfg = Config(cfg_dict={}, load=False, logger=logger) + block_cfg.NUM_BLOCKS = num_blocks + block_cfg.USE_NON_LOCAL = use_non_local[stage_id] + block_cfg.NUM_FILTERS = num_filters[stage_id] + block_cfg.NON_LOCAL = non_local_cfg + block_cfg.DIM_IN = num_filters[stage_id - 1] + block_cfg.DOWNSAMPLING = downsampling[stage_id] + block_cfg.DOWNSAMPLING_TEMPORAL = downsampling_temporal[stage_id] + block_cfg.BRANCH = branch_cfg + block_cfg.BRANCH.BRANCH_STYLE = branch_style + block_cfg.KERNEL_SIZE = tuple(kernel_size[stage_id]) + conv = Base3DResStage(block_cfg, logger=None) + setattr(self, f'conv{stage_id + 1}', conv) + # perform initialization + if isinstance(init_cfg, Config): + init_cfg = init_cfg.__dict__ + if init_cfg.get('name') == 'kaiming': + _init_convnet_weights(self) + + def forward(self, video): + x = self.conv1(video) + for i in range(2, 6): + x = getattr(self, f'conv{i}')(x) + return x + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BACKBONE', + __class__.__name__, + ResNet3D_CSN.para_dict, + set_name=True) + + +@BACKBONES.register_class() +class ResNet3D_TAda(nn.Module): + para_dict = { + 'DEPTH': { + 'value': 18, + 'description': "resnet model's depth!" + }, + 'NUM_INPUT_CHANNELS': { + 'value': 3, + 'description': 'the input channels num for the model!' + }, + 'NUM_FILTERS': { + 'value': [64, 256, 512, 1024, 2048], + 'description': 'this num filter for each layer!' + }, + 'KERNEL_SIZE': { + 'value': [[1, 7, 7], [1, 3, 3], [1, 3, 3], [1, 3, 3], [1, 3, 3]], + 'description': 'the kernel size for each layer!' + }, + 'DOWNSAMPLING': { + 'value': [True, True, True, True, True], + 'description': 'the downsample status for each layer!' + }, + 'DOWNSAMPLING_TEMPORAL': { + 'value': [False, False, False, False, False], + 'description': 'the temporal downsample status for each layer!' + }, + 'USE_NON_LOCAL': { + 'value': [False, False, False, False, False], + 'description': 'use non local or not for each layer!' + }, + 'NON_LOCAL': { + 'value': None, + 'description': 'the non local paramters!' + }, + 'STEM': { + 'NAME': { + 'value': + 'Base2DStem', + 'description': + 'use the shared parameters as DIM_IN, NUM_FILTERS, VISUAL_CFG,' + 'KERNEL_SIZE, DOWNSAMPLING, DOWNSAMPLING_TEMPORAL, BN_PARAMS' + } + }, + 'BRANCH': { + 'NAME': { + 'value': 'TAdaConvBlockAvgPool', + 'description': 'use the shared parameters as BRANCH_STYLE' + }, + 'EXPANSION_RATIO': { + 'value': 4, + 'description': 'the expansion ratio value' + } + }, + 'BN_PARAMS': { + 'value': + None, + 'description': + 'bn params data, key/value align with torch.BatchNorm3d/2d/1d!' + }, + 'INIT_CFG': { + 'name': { + 'value': 'kaiming', + 'description': 'init fn!' + } + }, + 'VISUAL_CFG': { + 'value': None, + 'description': 'the visualize config' + } + } + + def __init__(self, cfg, logger=None): + super(ResNet3D_TAda, self).__init__() + depth = cfg.get('DEPTH', 18) + num_input_channels = cfg.get('NUM_INPUT_CHANNELS', 3) + num_filters = cfg.get('NUM_FILTERS', [64, 64, 128, 256, 256]) + kernel_size = cfg.get( + 'KERNEL_SIZE', + [[1, 7, 7], [1, 3, 3], [1, 3, 3], [3, 3, 3], [3, 3, 3]]) + downsampling = cfg.get('DOWNSAMPLING', [True, False, True, True, True]) + downsampling_temporal = cfg.get('DOWNSAMPLING_TEMPORAL', + [False, False, False, True, True]) + use_non_local = cfg.get('USE_NON_LOCAL', + [False, False, False, False, False]) + non_local_cfg = cfg.get('NON_LOCAL', None) + stem_cfg = cfg.get('STEM', None) + branch_cfg = cfg.get('BRANCH', None) + + init_cfg = cfg.get('INIT_CFG', None) or dict() + visual_cfg = cfg.get('VISUAL_CFG', None) or dict() + # bn_params or {} + bn_params = cfg.get('BN_PARAMS', None) or dict() + if isinstance(bn_params, Config): + bn_params = bn_params.__dict__ + if len(bn_params) < 1: + bn_params = dict(eps=1e-3, momentum=0.1) + # Build stem cfg + stem_cfg.DIM_IN = num_input_channels + stem_cfg.NUM_FILTERS = num_filters[0] + stem_cfg.KERNEL_SIZE = tuple(kernel_size[0]) + stem_cfg.DOWNSAMPLING = downsampling[0] + stem_cfg.DOWNSAMPLING_TEMPORAL = downsampling_temporal[0] + stem_cfg.BN_PARAMS = bn_params + stem_cfg.VISUAL_CFG = visual_cfg + self.conv1 = STEMS.build(stem_cfg, logger=logger) + self.conv1.set_stage_block_id(0, 0) + # ------------------- Main arch ------------------- + branch_style = 'simple_block' if depth <= 34 else 'bottleneck' + blocks_list = _n_conv_resnet[depth] + for stage_id, num_blocks in enumerate(blocks_list): + stage_id = stage_id + 1 + # Build branch cfg + block_cfg = Config(cfg_dict={}, load=False, logger=logger) + block_cfg.NUM_BLOCKS = num_blocks + block_cfg.USE_NON_LOCAL = use_non_local[stage_id] + block_cfg.NUM_FILTERS = num_filters[stage_id] + block_cfg.NON_LOCAL = non_local_cfg + block_cfg.DIM_IN = num_filters[stage_id - 1] + block_cfg.DOWNSAMPLING = downsampling[stage_id] + block_cfg.DOWNSAMPLING_TEMPORAL = downsampling_temporal[stage_id] + block_cfg.BRANCH = branch_cfg + block_cfg.KERNEL_SIZE = tuple(kernel_size[stage_id]) + block_cfg.BRANCH.BRANCH_STYLE = branch_style + conv = Base3DResStage(block_cfg, logger=None) + setattr(self, f'conv{stage_id + 1}', conv) + # perform initialization + if isinstance(init_cfg, Config): + init_cfg = init_cfg.__dict__ + if init_cfg.get('name') == 'kaiming': + _init_convnet_weights(self) + + def forward(self, video): + x = self.conv1(video) + for i in range(2, 6): + x = getattr(self, f'conv{i}')(x) + return x + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BACKBONE', + __class__.__name__, + ResNet3D_2plus1d.para_dict, + set_name=True) diff --git a/scepter/modules/model/backbone/video/video_transformer.py b/scepter/modules/model/backbone/video/video_transformer.py new file mode 100644 index 0000000..7dca97f --- /dev/null +++ b/scepter/modules/model/backbone/video/video_transformer.py @@ -0,0 +1,383 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import math + +import torch +import torch.nn as nn +import torch.nn.functional +from einops import rearrange + +from scepter.modules.model.backbone.video.init_helper import ( + _init_transformer_weights, trunc_normal_) +from scepter.modules.model.registry import BACKBONES, BRICKS, STEMS +from scepter.modules.utils.config import dict_to_yaml +''' +The implementations of vivit as https://arxiv.org/abs/2103.15691. +The following setting alined the proposed model in the paper above. +Model1: Spatio-temporal attention. + class: VideoTransformer + stem: PatchEmbedStem/TubeletEmbeddingStem + branch: BaseTransformerLayer + complexity: (n_t * n_h * n_w) ** 2 +Model2: Factorised encoder + class: FactorizedVideoTransformer + stem: PatchEmbedStem/TubeletEmbeddingStem + branch: BaseTransformerLayer [drop_path=0.1] + complexity: (n_h * n_w) ** 2 + n_t ** 2 [attn_dropout = 0.0 ff_dropout=0.0] +Model3: Factorised self-attention + class: VideoTransformer + stem: TubeletEmbeddingStem + branch: TimesformerLayer + complexity: (n_h * n_w) ** 2 + O(attn_temp) +Model4: Factorised dot-product attention + coming soon... +TimesFormer: + class: VideoTransformer + stem: PatchEmbedStem [drop_path=0.0] + branch: TimesformerLayer [attn_dropout = 0.1 ff_dropout=0.1] + complexity: (n_h * n_w) ** 2 + O(attn_temp) +''' + + +@BACKBONES.register_class() +class VideoTransformer(nn.Module): + para_dict = { + 'NUM_INPUT_CHANNELS': { + 'value': 3, + 'description': "the input frames's channel!" + }, + 'NUM_FRAMES': { + 'value': 30, + 'description': "the num of transformer's input frames!" + }, + 'IMAGE_SIZE': { + 'value': 224, + 'description': "the input frame's size!" + }, + 'DIM': { + 'value': 768, + 'description': 'the patch embedding size!' + }, + 'PATCH_SIZE': { + 'value': 16, + 'description': 'the patch size!' + }, + 'DEPTH': { + 'value': 12, + 'description': 'the transformer network depth!' + }, + 'STEM': { + 'NAME': { + 'value': + 'PatchEmbedStem', + 'description': + 'also use TubeletEmbeddingStem, use the shared parameters as IMAGE_SIZE, PATCH_SIZE, NUM_FRAMES,' + 'NUM_INPUT_CHANNELS, DIM' + } + }, + 'BRANCH': { + 'NAME': { + 'value': + 'BaseTransformerLayer', + 'description': + 'also use TimesformerLayer, use the shared parameters as NUM_PATCHES, NUM_FRAMES, DIM' + } + } + } + + def __init__(self, cfg, logger=None): + super(VideoTransformer, self).__init__() + num_input_channels = cfg.get('NUM_INPUT_CHANNELS', 3) + num_frames = cfg.get('NUM_FRAMES', 8) + image_size = cfg.get('IMAGE_SIZE', 224) + num_features = cfg.get('DIM', 768) + patch_size = cfg.get('PATCH_SIZE', 16) + depth = cfg.get('DEPTH', 12) + drop_path = cfg.get('DROP_PATH', 0.1) + stem = cfg.STEM + branch = cfg.BRANCH + + assert image_size % patch_size == 0, 'Image dimensions must be divided by patch size.' + + self.num_patches_per_frame = (image_size // patch_size)**2 + if stem.NAME == 'TubeletEmbeddingStem': + self.num_patches = num_frames * self.num_patches_per_frame // stem.TUBELET_SIZE + else: + self.num_patches = num_frames * self.num_patches_per_frame + assert stem.NAME in ('PatchEmbedStem', 'TubeletEmbeddingStem') + stem.IMAGE_SIZE = image_size + stem.PATCH_SIZE = patch_size + stem.NUM_FRAMES = num_frames + stem.NUM_INPUT_CHANNELS = num_input_channels + stem.DIM = num_features + + self.stem = STEMS.build(stem, logger=logger) + + self.pos_embd = nn.Parameter( + torch.zeros(1, self.num_patches + 1, num_features)) + self.cls_token = nn.Parameter(torch.randn(1, 1, num_features)) + + assert branch.NAME in ('BaseTransformerLayer', 'TimesformerLayer') + if branch.NAME == 'TimesformerLayer': + branch.NUM_PATCHES = image_size // patch_size**2 + branch.NUM_FRAMES = num_frames + branch.DIM = num_features + + # construct spatial transformer layers + dpr = [x.item() for x in torch.linspace(0, drop_path, depth) + ] # stochastic depth decay rule + + layers = [] + for i in range(depth): + branch.DROP_PATH_PROB = dpr[i] + layers.append(BRICKS.build(branch, logger=logger)) + self.layers = nn.Sequential(*layers) + self.norm = nn.LayerNorm(num_features, eps=1e-6) + + # initialization + trunc_normal_(self.pos_embd, std=.02) + trunc_normal_(self.cls_token, std=.02) + self.apply(_init_transformer_weights) + + def forward(self, video): + x = video + x = self.stem(x) + + cls_token = self.cls_token.repeat((x.shape[0], 1, 1)) + x = torch.cat((cls_token, x), dim=1) + + x += self.pos_embd + + x = self.layers(x) + x = self.norm(x) + + return x[:, 0] + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BACKBONE', + __class__.__name__, + VideoTransformer.para_dict, + set_name=True) + + +@BACKBONES.register_class() +class FactorizedVideoTransformer(nn.Module): + para_dict = { + 'NUM_INPUT_CHANNELS': { + 'value': 3, + 'description': "the input frames's channel!" + }, + 'NUM_FRAMES': { + 'value': 30, + 'description': "the num of transformer's input frames!" + }, + 'IMAGE_SIZE': { + 'value': 224, + 'description': "the input frame's size!" + }, + 'DIM': { + 'value': 768, + 'description': 'the patch embedding size!' + }, + 'PATCH_SIZE': { + 'value': 16, + 'description': 'the patch size!' + }, + 'TUBELET_SIZE': { + 'value': 2, + 'description': 'the tubelet size, also means the temporal stride!' + }, + 'DEPTH': { + 'value': 12, + 'description': 'the transformer network depth!' + }, + 'DEPTH_TEMPORAL': { + 'value': 4, + 'description': 'the temporal network depth!' + }, + 'DROP_PATH': { + 'value': 0.1, + 'description': 'the drop path module value!' + }, + 'STEM': { + 'NAME': { + 'value': + 'PatchEmbedStem', + 'description': + 'use the shared parameters as IMAGE_SIZE, PATCH_SIZE, NUM_FRAMES,' + 'NUM_INPUT_CHANNELS, DIM' + } + }, + 'BRANCH': { + 'NAME': { + 'value': + 'BaseTransformerLayer', + 'description': + 'use the shared parameters as NUM_PATCHES, NUM_FRAMES, DIM, ' + 'TUBELET_SIZE (if use TubeletEmbeddingStem)' + } + } + } + + def __init__(self, cfg, logger=None): + super(FactorizedVideoTransformer, self).__init__() + num_input_channels = cfg.get('NUM_INPUT_CHANNELS', 3) + num_frames = cfg.get('NUM_FRAMES', 8) + image_size = cfg.get('IMAGE_SIZE', 224) + num_features = cfg.get('DIM', 768) + self.patch_size = cfg.get('PATCH_SIZE', 16) + tubelet_size = cfg.get('TUBELET_SIZE', 2) + depth = cfg.get('DEPTH', 12) + depth_temp = cfg.get('DEPTH_TEMPORAL', 4) + drop_path = cfg.get('DROP_PATH', 0.1) + stem = cfg.STEM + branch = cfg.BRANCH + + assert image_size % self.patch_size == 0, 'Image dimensions must be divided by patch size.' + + self.num_patches_per_frame = (image_size // self.patch_size)**2 + self.num_patches = num_frames * self.num_patches_per_frame // tubelet_size + + assert stem.NAME in ('PatchEmbedStem', 'TubeletEmbeddingStem') + stem.IMAGE_SIZE = image_size + stem.PATCH_SIZE = self.patch_size + stem.NUM_FRAMES = num_frames + stem.NUM_INPUT_CHANNELS = num_input_channels + stem.DIM = num_features + if stem.NAME == 'TubeletEmbeddingStem': + stem.TUBELET_SIZE = tubelet_size + + self.stem = STEMS.build(stem, logger=logger) + + self.pos_embd = nn.Parameter( + torch.zeros(1, self.num_patches_per_frame + 1, num_features)) + self.temp_embd = nn.Parameter( + torch.zeros(1, num_frames // tubelet_size + 1, num_features)) + self.cls_token = nn.Parameter(torch.randn(1, 1, num_features)) + self.cls_token_out = nn.Parameter(torch.randn(1, 1, num_features)) + + assert branch.NAME in ('BaseTransformerLayer', 'TimesformerLayer') + if branch.NAME == 'TimesformerLayer': + branch.NUM_PATCHES = image_size // self.patch_size**2 + branch.NUM_FRAMES = num_frames + branch.DIM = num_features + + # construct spatial transformer layers + dpr = [ + x.item() for x in torch.linspace(0, drop_path, depth + depth_temp) + ] # stochastic depth decay rule + + layers = [] + for i in range(depth): + branch.DROP_PATH_PROB = dpr[i] + layers.append(BRICKS.build(branch, logger=logger)) + self.layers = nn.Sequential(*layers) + self.norm = nn.LayerNorm(num_features, eps=1e-6) + + # construct temporal transformer layers + layers_temporal = [] + for i in range(depth_temp): + branch.DROP_PATH_PROB = dpr[i + depth] + layers_temporal.append(BRICKS.build(branch, logger=logger)) + self.layers_temporal = nn.Sequential(*layers_temporal) + + self.norm_out = nn.LayerNorm(num_features, eps=1e-6) + + # initialization + trunc_normal_(self.pos_embd, std=.02) + trunc_normal_(self.temp_embd, std=.02) + trunc_normal_(self.cls_token, std=.02) + trunc_normal_(self.cls_token_out, std=.02) + self.apply(_init_transformer_weights) + + def forward(self, video): + x = video + h, w = x.shape[-2:] + actual_num_patches_per_frame = (h // self.patch_size) * ( + w // self.patch_size) + x = self.stem(x) + + if actual_num_patches_per_frame != self.num_patches_per_frame: + assert not self.training + x = rearrange(x, + 'b (t n) c -> (b t) n c', + n=actual_num_patches_per_frame) + else: + x = rearrange(x, + 'b (t n) c -> (b t) n c', + n=self.num_patches_per_frame) + + cls_token = self.cls_token.repeat((x.shape[0], 1, 1)) + x = torch.cat((cls_token, x), dim=1) + + # to make the input video size changable + if actual_num_patches_per_frame != self.num_patches_per_frame: + actual_num_pathces_per_side = int( + math.sqrt(actual_num_patches_per_frame)) + if not hasattr(self, + 'new_pos_embd') or self.new_pos_embd.shape[1] != ( + actual_num_pathces_per_side**2 + 1): + cls_pos_embd = self.pos_embd[:, 0, :].unsqueeze(1) + pos_embd = self.pos_embd[:, 1:, :] + num_patches_per_side = int( + math.sqrt(self.num_patches_per_frame)) + pos_embd = pos_embd.reshape(1, num_patches_per_side, + num_patches_per_side, + -1).permute(0, 3, 1, 2) + pos_embd = torch.nn.functional.interpolate( + pos_embd, + size=(actual_num_pathces_per_side, + actual_num_pathces_per_side), + mode='bilinear').permute(0, 2, 3, 1).reshape( + 1, actual_num_pathces_per_side**2, -1) + self.new_pos_embd = torch.cat((cls_pos_embd, pos_embd), dim=1) + x += self.new_pos_embd + else: + x += self.pos_embd + + x = self.layers(x) + x = self.norm(x)[:, 0] + + x = rearrange(x, + '(b t) c -> b t c', + t=self.num_patches // self.num_patches_per_frame) + + cls_token_out = self.cls_token_out.repeat((x.shape[0], 1, 1)) + x = torch.cat((cls_token_out, x), dim=1) + + x += self.temp_embd + x = self.layers_temporal(x) + x = self.norm_out(x) + + return x[:, 0] + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('BACKBONE', + __class__.__name__, + FactorizedVideoTransformer.para_dict, + set_name=True) diff --git a/scepter/modules/model/base_model.py b/scepter/modules/model/base_model.py new file mode 100644 index 0000000..6adf8c1 --- /dev/null +++ b/scepter/modules/model/base_model.py @@ -0,0 +1,114 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import copy + +import torch.nn as nn + +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import gather_data, we +from scepter.modules.utils.probe import (ProbeData, merge_gathered_probe, + register_data) + + +class BaseModel(nn.Module): + para_dict = { + 'PRETRAINED_MODEL': { + 'value': None, + 'description': 'Pretrained model path.' + } + } + + def __init__(self, cfg, logger=None): + super(BaseModel, self).__init__() + self.logger = logger + self.cfg = cfg + self._probe_data = {} + self._dist_data = {} + + def __repr__(self) -> str: + return f'{self.__class__.__name__}' + ' ' + super().__repr__() + + def load_pretrained_model(self, pretrained_model): + pass + + def register_probe(self, probe_data: dict): + probe_da, dist_da = register_data(probe_data, + key_prefix=__class__.__name__) + self._probe_data.update(probe_da) + for key in dist_da: + if key not in self._dist_data: + self._dist_data[key] = dist_da[key] + else: + for k, v in dist_da[key].items(): + if k in self._dist_data[key]: + self._dist_data[key][k] += v + else: + self._dist_data[key][k] = v + + def probe_data(self): + gather_probe_data = gather_data(self._probe_data) + _dist_data_list = gather_data([self._dist_data]) + if not we.rank == 0: + self._probe_data = {} + self._dist_data = {} + # Iterate recurse the sub class's probe data for time-aware data. + for k, v in self._modules.items(): + if isinstance(getattr(self, k), BaseModel): + for kk, vv in getattr(self, k).probe_data().items(): + self._probe_data[f'{k}/{kk}'] = vv + + if gather_probe_data is not None: + # Before processing, just merge the data. + self._probe_data = merge_gathered_probe(gather_probe_data) + reduce_dist_data = {} + if _dist_data_list is not None: + reduce_dist_data = {} + for one_data in _dist_data_list: + for k, v in one_data.items(): + if k in reduce_dist_data: + for kk, vv in v.items(): + if kk in reduce_dist_data[k]: + reduce_dist_data[k][kk] += vv + else: + reduce_dist_data[k][kk] = vv + else: + reduce_dist_data[k] = v + self._dist_data = reduce_dist_data + # Iterate recurse the sub class's probe data for reduce data. + self._probe_data[f'{__class__.__name__}_distribute'] = ProbeData( + self._dist_data) + norm_dist_data = {} + for key, value in self._dist_data.items(): + total = 0 + for k, v in value.items(): + total += v + norm_v = {} + for k, v in value.items(): + norm_v[k] = v / total + norm_dist_data[key] = norm_v + self._probe_data[f'{__class__.__name__}_norm_distribute'] = ProbeData( + norm_dist_data) + ret_data = copy.deepcopy(self._probe_data) + self._probe_data = {} + return ret_data + + def clear_probe(self): + self._probe_data.clear() + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('MODELS', + __class__.__name__, + BaseModel.para_dict, + set_name=True) diff --git a/scepter/modules/model/embedder/__init__.py b/scepter/modules/model/embedder/__init__.py new file mode 100644 index 0000000..20e2f21 --- /dev/null +++ b/scepter/modules/model/embedder/__init__.py @@ -0,0 +1,7 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.embedder.embedder import (ConcatTimestepEmbedderND, + FrozenCLIPEmbedder, + FrozenOpenCLIPEmbedder, + FrozenOpenCLIPEmbedder2, + GeneralConditioner) diff --git a/scepter/modules/model/embedder/base_embedder.py b/scepter/modules/model/embedder/base_embedder.py new file mode 100644 index 0000000..d4a5c88 --- /dev/null +++ b/scepter/modules/model/embedder/base_embedder.py @@ -0,0 +1,31 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from abc import ABCMeta + +from scepter.modules.model.base_model import BaseModel +from scepter.modules.model.registry import EMBEDDERS +from scepter.modules.utils.config import dict_to_yaml + + +@EMBEDDERS.register_class() +class BaseEmbedder(BaseModel, metaclass=ABCMeta): + para_dict = {} + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + + def encode(self, *args, **kwargs): + raise NotImplementedError + + def encode_text(self, *args, **kwargs): + raise NotImplementedError + + def encode_image(self, *args, **kwargs): + raise NotImplementedError + + @staticmethod + def get_config_template(): + return dict_to_yaml('EMBEDDERS', + __class__.__name__, + BaseEmbedder.para_dict, + set_name=True) diff --git a/scepter/modules/model/embedder/embedder.py b/scepter/modules/model/embedder/embedder.py new file mode 100644 index 0000000..e710cef --- /dev/null +++ b/scepter/modules/model/embedder/embedder.py @@ -0,0 +1,675 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from collections import OrderedDict +from contextlib import nullcontext +from typing import Dict + +import numpy as np +import open_clip +import torch +import torch.nn as nn +import torch.utils.dlpack +from einops import rearrange +from torch.utils.checkpoint import checkpoint +from transformers import CLIPTextModel, CLIPTokenizer + +# to check +from scepter.modules.model.backbone.unet.unet_utils import Timestep +from scepter.modules.model.registry import EMBEDDERS +from scepter.modules.model.utils.basic_utils import expand_dims_like +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we +from scepter.modules.utils.file_system import FS + +from .base_embedder import BaseEmbedder + + +def autocast(f, enabled=True): + def do_autocast(*args, **kwargs): + with torch.cuda.amp.autocast( + enabled=enabled, + dtype=torch.get_autocast_gpu_dtype(), + cache_enabled=torch.is_autocast_cache_enabled(), + ): + return f(*args, **kwargs) + + return do_autocast + + +@EMBEDDERS.register_class() +class FrozenCLIPEmbedder(BaseEmbedder): + """Uses the CLIP transformer encoder for text (from huggingface)""" + para_dict = { + 'PRETRAINED_MODEL': { + 'value': None, + 'description': 'You should set pretrained_model: modelcard.' + }, + 'TOKENIZER_PATH': { + 'value': + None, + 'description': + 'If you want to use tokenizer in embedder, you should set this field.' + }, + 'MAX_LENGTH': { + 'value': 77, + 'description': '' + }, + 'FREEZE': { + 'value': True, + 'description': '' + }, + 'USE_GRAD': { + 'value': False, + 'description': 'Compute grad or not.' + }, + 'LAYER': { + 'value': 'last', + 'description': '' + }, + 'LAYER_IDX': { + 'value': None, + 'description': '' + }, + 'USE_FINAL_LAYER_NORM': { + 'value': False, + 'description': '' + }, + } + LAYERS = ['last', 'pooled', 'hidden'] + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + tokenizer_path = cfg.get('TOKENIZER_PATH', None) + if tokenizer_path is not None: + with FS.get_dir_to_local_dir(tokenizer_path, + wait_finish=True) as local_path: + self.tokenizer = CLIPTokenizer.from_pretrained(local_path) + + pretrained_model = cfg.get('PRETRAINED_MODEL', None) + if pretrained_model is None: + raise 'You should set pretrained_model: modelcard.' + with FS.get_dir_to_local_dir(cfg.PRETRAINED_MODEL, + wait_finish=True) as local_path: + self.transformer = CLIPTextModel.from_pretrained(local_path) + + self.use_grad = cfg.get('USE_GRAD', False) + self.freeze_flag = cfg.get('FREEZE', True) + if self.freeze_flag: + self.freeze() + + self.max_length = cfg.get('MAX_LENGTH', 77) + self.layer = cfg.get('LAYER', 'last') + self.layer_idx = cfg.get('LAYER_IDX', None) + self.use_final_layer_norm = cfg.get('USE_FINAL_LAYER_NORM', False) + assert self.layer in self.LAYERS + if self.layer == 'hidden': + assert self.layer_idx is not None + assert 0 <= abs(self.layer_idx) <= 12 + + def freeze(self): + self.transformer = self.transformer.eval() + for param in self.parameters(): + param.requires_grad = False + + # @torch.no_grad() + def _forward(self, text): + batch_encoding = self.tokenizer(text, + truncation=True, + max_length=self.max_length, + return_length=True, + return_overflowing_tokens=False, + padding='max_length', + return_tensors='pt') + tokens = batch_encoding['input_ids'].to(we.device_id) + outputs = self.transformer(input_ids=tokens, + output_hidden_states=self.layer == 'hidden') + if self.layer == 'last': + z = outputs.last_hidden_state + elif self.layer == 'pooled': + z = outputs.pooler_output[:, None, :] + if self.use_final_layer_norm: + z = self.transformer.text_model.final_layer_norm(z) + else: + z = outputs.hidden_states[self.layer_idx] + if self.use_final_layer_norm: + z = self.transformer.text_model.final_layer_norm(z) + return z + + @autocast + def forward(self, text): + if not self.use_grad: + with torch.no_grad(): + output = self._forward(text) + else: + output = self._forward(text) + return output + + def encode(self, text): + return self(text) + + # @torch.no_grad() + def _encode_text(self, + tokens, + tokenizer=None, + append_sentence_embedding=False): + outputs = self.transformer(input_ids=tokens, + output_hidden_states=self.layer == 'hidden') + if self.layer == 'last': + z = outputs.last_hidden_state + elif self.layer == 'pooled': + z = outputs.pooler_output[:, None, :] + if self.use_final_layer_norm: + z = self.transformer.text_model.final_layer_norm(z) + else: + z = outputs.hidden_states[self.layer_idx] + if self.use_final_layer_norm: + z = self.transformer.text_model.final_layer_norm(z) + return z + + def encode_text(self, + tokens, + tokenizer=None, + append_sentence_embedding=False): + if not self.use_grad: + with torch.no_grad(): + output = self._encode_text(tokens, tokenizer, + append_sentence_embedding) + else: + output = self._encode_text(tokens, tokenizer, + append_sentence_embedding) + return output + + @staticmethod + def get_config_template(): + return dict_to_yaml('MODELS', + __class__.__name__, + FrozenCLIPEmbedder.para_dict, + set_name=True) + + +@EMBEDDERS.register_class() +class FrozenOpenCLIPEmbedder(BaseEmbedder): + """ + Uses the OpenCLIP transformer encoder for text + """ + para_dict = { + 'ARCH': { + 'value': 'ViT-H-14', + 'description': '' + }, + 'PRETRAINED_MODEL': { + 'value': '', + 'description': '' + }, + 'MAX_LENGTH': { + 'value': 77, + 'description': '' + }, + 'FREEZE': { + 'value': True, + 'description': '' + }, + 'USE_GRAD': { + 'value': False, + 'description': 'Compute grad or not.' + }, + 'LAYER': { + 'value': 'last', + 'description': '' + }, + } + LAYERS = ['last', 'penultimate'] + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + + arch = cfg.get('ARCH', 'ViT-H-14') + if cfg.PRETRAINED_MODEL is None: + model, _, _ = open_clip.create_model_and_transforms( + arch, device=torch.device('cpu'), pretrained=None) + del model.visual + else: + with FS.get_from(cfg.PRETRAINED_MODEL, + wait_finish=True) as local_path: + model, _, _ = open_clip.create_model_and_transforms( + arch, device=torch.device('cpu'), pretrained=local_path) + del model.visual + self.model = model + + self.use_grad = cfg.get('USE_GRAD', False) + self.freeze_flag = cfg.get('FREEZE', True) + if self.freeze_flag: + self.freeze() + + self.max_length = cfg.get('MAX_LENGTH', 77) + self.layer = cfg.get('LAYER', 'penultimate') + assert self.layer in self.LAYERS + if self.layer == 'last': + self.layer_idx = 0 + elif self.layer == 'penultimate': + self.layer_idx = 1 + else: + raise NotImplementedError() + + def freeze(self): + self.model = self.model.eval() + for param in self.parameters(): + param.requires_grad = False + + @autocast + def forward(self, text): + embedding_context = nullcontext if self.use_grad else torch.no_grad + with embedding_context(): + tokens = open_clip.tokenize(text) + z = self.encode_with_transformer(tokens.to(we.device_id)) + return z + + def encode_with_transformer(self, text): + embedding_context = nullcontext if self.use_grad else torch.no_grad + with embedding_context(): + x = self.model.token_embedding( + text) # [batch_size, n_ctx, d_model] + x = x + self.model.positional_embedding + x = x.permute(1, 0, 2) # NLD -> LND + x = self.text_transformer_forward(x, + attn_mask=self.model.attn_mask) + x = x.permute(1, 0, 2) # LND -> NLD + x = self.model.ln_final(x) + return x + + def text_transformer_forward(self, x: torch.Tensor, attn_mask=None): + embedding_context = nullcontext if self.use_grad else torch.no_grad + with embedding_context(): + for i, r in enumerate(self.model.transformer.resblocks): + if i == len(self.model.transformer.resblocks) - self.layer_idx: + break + if self.model.transformer.grad_checkpointing and not torch.jit.is_scripting( + ): + x = checkpoint(r, x, attn_mask) + else: + x = r(x, attn_mask=attn_mask) + return x + + def encode_text(self, + tokens, + tokenizer=None, + append_sentence_embedding=False): + embedding_context = nullcontext if self.use_grad else torch.no_grad + with embedding_context(): + z = self.encode_with_transformer(tokens.to(we.device_id)) + return z + + def encode(self, text): + return self(text) + + @staticmethod + def get_config_template(): + return dict_to_yaml('MODELS', + __class__.__name__, + FrozenOpenCLIPEmbedder.para_dict, + set_name=True) + + +@EMBEDDERS.register_class() +class FrozenOpenCLIPEmbedder2(BaseEmbedder): + """ + Uses the OpenCLIP transformer encoder for text + """ + """ + Uses the OpenCLIP transformer encoder for text + """ + para_dict = { + 'ARCH': { + 'value': 'ViT-H-14', + 'description': '' + }, + 'PRETRAINED_MODEL': { + 'value': 'laion2b_s32b_b79k', + 'description': '' + }, + 'MAX_LENGTH': { + 'value': 77, + 'description': '' + }, + 'FREEZE': { + 'value': True, + 'description': '' + }, + 'USE_GRAD': { + 'value': False, + 'description': 'Compute grad or not.' + }, + 'ALWAYS_RETURN_POOLED': { + 'value': + False, + 'description': + 'Whether always return pooled results or not ,default False.' + }, + 'LEGACY': { + 'value': + True, + 'description': + 'Whether use legacy returnd feature or not ,default True.' + }, + 'LAYER': { + 'value': 'last', + 'description': '' + }, + } + + LAYERS = ['pooled', 'last', 'penultimate'] + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + arch = cfg.get('ARCH', 'ViT-H-14') + if cfg.PRETRAINED_MODEL is None: + model, _, _ = open_clip.create_model_and_transforms( + arch, device=torch.device('cpu'), pretrained=None) + del model.visual + else: + with FS.get_from(cfg.PRETRAINED_MODEL, + wait_finish=True) as local_path: + model, _, _ = open_clip.create_model_and_transforms( + arch, device=torch.device('cpu'), pretrained=local_path) + del model.visual + self.model = model + + self.max_length = cfg.get('MAX_LENGTH', 77) + self.layer = cfg.get('LAYER', 'last') + self.return_pooled = cfg.get('ALWAYS_RETURN_POOLED', False) + self.use_grad = cfg.get('USE_GRAD', False) + self.freeze_flag = cfg.get('FREEZE', True) + if self.freeze_flag: + self.freeze() + if self.layer == 'last': + self.layer_idx = 0 + elif self.layer == 'penultimate': + self.layer_idx = 1 + else: + raise NotImplementedError() + self.legacy = cfg.get('LEGACY', True) + + def freeze(self): + self.model = self.model.eval() + for param in self.parameters(): + param.requires_grad = False + + @autocast + def forward(self, text): + embedding_context = nullcontext if self.use_grad else torch.no_grad + with embedding_context(): + tokens = open_clip.tokenize(text) + z = self.encode_with_transformer(tokens.to(we.device_id)) + if not self.return_pooled and self.legacy: + return z + if self.return_pooled: + assert not self.legacy + return z[self.layer], z['pooled'] + return z[self.layer] + + def encode_with_transformer(self, text): + embedding_context = nullcontext if self.use_grad else torch.no_grad + with embedding_context(): + x = self.model.token_embedding( + text) # [batch_size, n_ctx, d_model] + x = x + self.model.positional_embedding + x = x.permute(1, 0, 2) # NLD -> LND + x = self.text_transformer_forward(x, + attn_mask=self.model.attn_mask) + if self.legacy: + x = x[self.layer] + x = self.model.ln_final(x) + return x + else: + # x is a dict and will stay a dict + o = x['last'] + o = self.model.ln_final(o) + pooled = self.pool(o, text) + x['pooled'] = pooled + return x + + def pool(self, x, text): + # take features from the eot embedding (eot_token is the highest number in each sequence) + x = (x[torch.arange(x.shape[0]), + text.argmax(dim=-1)] @ self.model.text_projection) + return x + + def text_transformer_forward(self, x: torch.Tensor, attn_mask=None): + outputs = {} + embedding_context = nullcontext if self.use_grad else torch.no_grad + with embedding_context(): + for i, r in enumerate(self.model.transformer.resblocks): + if i == len(self.model.transformer.resblocks) - 1: + outputs['penultimate'] = x.permute(1, 0, 2) # LND -> NLD + if (self.model.transformer.grad_checkpointing + and not torch.jit.is_scripting()): + x = checkpoint(r, x, attn_mask) + else: + x = r(x, attn_mask=attn_mask) + outputs['last'] = x.permute(1, 0, 2) # LND -> NLD + return outputs + + def encode(self, text): + return self(text) + + def encode_text(self, + tokens, + tokenizer=None, + append_sentence_embedding=False): + z = self.encode_with_transformer(tokens.to(we.device_id)) + if not self.return_pooled and self.legacy: + return z + if self.return_pooled: + assert not self.legacy + return z[self.layer], z['pooled'] + return z[self.layer] + + @staticmethod + def get_config_template(): + return dict_to_yaml('MODELS', + __class__.__name__, + FrozenOpenCLIPEmbedder2.para_dict, + set_name=True) + + +@EMBEDDERS.register_class() +class ConcatTimestepEmbedderND(BaseEmbedder): + """embeds each dimension independently and concatenates them""" + para_dict = { + 'OUT_DIM': { + 'value': 256, + 'description': 'Output dim' + }, + } + + def __init__(self, cfg, logger=None): + super().__init__(cfg=cfg, logger=logger) + outdim = cfg.get('OUT_DIM', 256) + self.timestep = Timestep(outdim, legacy=True) + self.outdim = outdim + + def forward(self, x): + if x.ndim == 1: + x = x[:, None] + assert len(x.shape) == 2 + b, dims = x.shape[0], x.shape[1] + x = rearrange(x, 'b d -> (b d)') + emb = self.timestep(x) + emb = rearrange(emb, + '(b d) d2 -> b (d d2)', + b=b, + d=dims, + d2=self.outdim) + return emb + + @staticmethod + def get_config_template(): + return dict_to_yaml('MODELS', + __class__.__name__, + ConcatTimestepEmbedderND.para_dict, + set_name=True) + + +@EMBEDDERS.register_class() +class GeneralConditioner(BaseEmbedder): + OUTPUT_DIM2KEYS = {2: 'y', 3: 'crossattn', 4: 'concat', 5: 'concat'} + KEY2CATDIM = {'y': 1, 'crossattn': 2, 'concat': 1} + para_dict = { + 'EMBEDDERS': [], + 'USE_GRAD': { + 'value': False, + 'description': 'Compute grad or not.' + }, + } + para_dict.update(para_dict) + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + emb_models = cfg.get('EMBEDDERS', []) + use_grad = cfg.get('USE_GRAD', False) + self.embedders = nn.ModuleList([]) + for n, embconfig in enumerate(emb_models): + embconfig.USE_GRAD = use_grad if not embconfig.have( + 'USE_GRAD' + ) or embconfig.USE_GRAD is None else embconfig.USE_GRAD + embedder = EMBEDDERS.build(embconfig, logger=logger) + embedder.ucg_rate = embconfig.get('UCG_RATE', 0.0) + embedder.input_keys = embconfig.get('INPUT_KEYS', []) + embedder.legacy_ucg_val = embconfig.get('LEGACY_UCG_VALUE', None) + if embedder.legacy_ucg_val is not None: + embedder.ucg_prng = np.random.RandomState() + + self.embedders.append(embedder) + + def load_pretrained_model(self, pretrained_model): + if pretrained_model is not None: + with FS.get_from(pretrained_model, + wait_finish=True) as local_model: + self.init_from_ckpt(local_model) + + def init_from_ckpt(self, path, ignore_keys=list()): + if path.endswith('safetensors'): + from safetensors.torch import load_file as load_safetensors + sd = load_safetensors(path) + else: + sd = torch.load(path, map_location='cpu') + new_sd = OrderedDict() + for k, v in sd.items(): + ignored = False + for ik in ignore_keys: + if ik in k: + if we.rank == 0: + self.logger.info( + 'Ignore key {} from state_dict.'.format(k)) + ignored = True + break + if not ignored: + new_sd[k] = v + + missing, unexpected = self.load_state_dict(new_sd, strict=False) + if we.rank == 0: + self.logger.info( + f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys' + ) + if len(missing) > 0: + self.logger.info(f'Missing Keys:\n {missing}') + if len(unexpected) > 0: + self.logger.info(f'\nUnexpected Keys:\n {unexpected}') + + def possibly_get_ucg_val(self, embedder, batch: Dict) -> Dict: + assert embedder.legacy_ucg_val is not None + p = embedder.ucg_rate + val = embedder.legacy_ucg_val + for i in range(len(batch[embedder.input_key])): + if embedder.ucg_prng.choice(2, p=[1 - p, p]): + batch[embedder.input_key][i] = val + return batch + + def forward(self, batch: Dict, force_zero_embeddings=None) -> Dict: + output = dict() + if force_zero_embeddings is None: + force_zero_embeddings = [] + for embedder in self.embedders: + embedding_context = nullcontext if hasattr( + embedder, 'use_grad') and embedder.use_grad else torch.no_grad + with embedding_context(): + if hasattr(embedder, 'input_key') and (embedder.input_key + is not None): + if embedder.legacy_ucg_val is not None: + batch = self.possibly_get_ucg_val(embedder, batch) + emb_out = embedder(batch[embedder.input_key]) + elif hasattr(embedder, 'input_keys'): + emb_out = embedder( + *[batch[k] for k in embedder.input_keys]) + assert isinstance( + emb_out, (torch.Tensor, list, tuple) + ), f'encoder outputs must be tensors or a sequence, but got {type(emb_out)}' + if not isinstance(emb_out, (list, tuple)): + emb_out = [emb_out] + for emb in emb_out: + # print("emb.shape", emb.shape) + # print("emb.input_keys", embedder.input_keys) + out_key = self.OUTPUT_DIM2KEYS[emb.dim()] + if embedder.ucg_rate > 0.0 and embedder.legacy_ucg_val is None: + emb = (expand_dims_like( + torch.bernoulli( + (1.0 - embedder.ucg_rate) * + torch.ones(emb.shape[0], device=emb.device)), + emb, + ) * emb) + if (hasattr(embedder, 'input_keys')): + if np.sum( + np.array([ + key in force_zero_embeddings + for key in embedder.input_keys + ])) > 0: + emb = torch.zeros_like(emb) + if out_key in output: + output[out_key] = torch.cat((output[out_key], emb), + self.KEY2CATDIM[out_key]) + else: + output[out_key] = emb + # if "y" in output: + # print("out.shape", output["y"].shape) + return output + + def get_unconditional_conditioning(self, + batch_c, + batch_uc=None, + force_uc_zero_embeddings=None): + if force_uc_zero_embeddings is None: + force_uc_zero_embeddings = [] + ucg_rates = list() + for embedder in self.embedders: + ucg_rates.append(embedder.ucg_rate) + embedder.ucg_rate = 0.0 + c = self(batch_c) + uc = self(batch_c if batch_uc is None else batch_uc, + force_uc_zero_embeddings) + + for embedder, rate in zip(self.embedders, ucg_rates): + embedder.ucg_rate = rate + return c, uc + + def encode(self, + batch_dict, + is_unconditional=False, + force_uc_zero_embeddings=None): + ucg_rates = list() + for embedder in self.embedders: + ucg_rates.append(embedder.ucg_rate) + embedder.ucg_rate = 0.0 + if is_unconditional: + c = self(batch_dict, force_uc_zero_embeddings) + else: + c = self(batch_dict) + for embedder, rate in zip(self.embedders, ucg_rates): + embedder.ucg_rate = rate + return c + + @staticmethod + def get_config_template(): + return dict_to_yaml('MODELS', + __class__.__name__, + GeneralConditioner.para_dict, + set_name=True) diff --git a/scepter/modules/model/head/__init__.py b/scepter/modules/model/head/__init__.py new file mode 100644 index 0000000..9f70332 --- /dev/null +++ b/scepter/modules/model/head/__init__.py @@ -0,0 +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) diff --git a/scepter/modules/model/head/classifier_head.py b/scepter/modules/model/head/classifier_head.py new file mode 100644 index 0000000..00cdb85 --- /dev/null +++ b/scepter/modules/model/head/classifier_head.py @@ -0,0 +1,406 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import math +from collections import OrderedDict + +import torch +import torch.nn as nn +import torch.nn.functional as F +from torch.nn.parameter import Parameter + +from scepter.modules.model.base_model import BaseModel +from scepter.modules.model.registry import HEADS +from scepter.modules.utils.config import dict_to_yaml + + +@HEADS.register_class() +class ClassifierHead(BaseModel): + para_dict = { + 'DIM': { + 'value': 512, + 'description': 'representation dim!' + }, + 'NUM_CLASSES': { + 'value': 10, + 'description': 'number of classes.' + }, + 'DROPOUT_RATE': { + 'value': 0.0, + 'description': 'dropout rate, default 0.' + } + } + + def __init__(self, cfg, logger=None): + super(ClassifierHead, self).__init__(cfg, logger=logger) + self.dim = cfg.DIM + self.num_classes = cfg.NUM_CLASSES + self.dropout_rate = cfg.DROPOUT_RATE + if self.dropout_rate > 0.0: + self.dropout = nn.Dropout(self.dropout_rate) + + self.fc = nn.Linear(self.dim, self.num_classes) + + def forward(self, x, label=None): + x = x.type(self.fc.weight.dtype) + if hasattr(self, 'dropout'): + x = self.dropout(x) + return self.fc(x) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('HEADS', + __class__.__name__, + ClassifierHead.para_dict, + set_name=True) + + +class CosineLinear(nn.Module): + def __init__(self, + in_features: int, + out_features: int, + sigma: bool = True): + super(CosineLinear, self).__init__() + self.in_features = in_features + self.out_features = out_features + self.weight = Parameter(torch.Tensor(out_features, in_features)) + if sigma: + self.sigma = Parameter(torch.Tensor(1)) + else: + self.register_parameter('sigma', None) + self.reset_parameters() + + def reset_parameters(self): + stdv = 1. / math.sqrt(self.weight.size(1)) + self.weight.data.uniform_(-stdv, stdv) + if self.sigma is not None: + self.sigma.data.fill_(1) # for initializaiton of sigma + + def forward(self, x, label=None): + out = F.linear(F.normalize(x, p=2, dim=1), + F.normalize(self.weight, p=2, dim=1)) + if self.sigma is not None: + out = self.sigma * out + return out + + +@HEADS.register_class() +class CosineLinearHead(BaseModel): + para_dict = { + 'IN_DIM': { + 'value': 64, + 'description': 'the input dim for head!' + }, + 'NUM_CLASSES': { + 'value': + 10, + 'description': + 'The output dim for head, often this value is the classes number!' + }, + 'SIGMA': { + 'value': True, + 'description': 'The cosine scale which is learned by the model!' + } + } + + def __init__(self, cfg, logger=None): + super(CosineLinearHead, self).__init__(cfg, logger=logger) + self.in_features = cfg.IN_DIM + self.out_features = cfg.NUM_CLASSES + sigma = cfg.get('SIGMA', True) + self.fc = CosineLinear(self.in_features, self.out_features, sigma) + + def forward(self, x, label=None): + x = x.type(self.fc.weight.dtype) + x = self.fc(x) + return x + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('HEADS', + __class__.__name__, + CosineLinearHead.para_dict, + set_name=True) + + +@HEADS.register_class() +class VideoClassifierHead(BaseModel): + para_dict = { + 'DIM': { + 'value': 512, + 'description': 'representation dim!' + }, + 'NUM_CLASSES': { + 'value': 10, + 'description': 'number of classes.' + }, + 'DROPOUT_RATE': { + 'value': 0.0, + 'description': 'dropout rate, default 0.' + } + } + + def __init__(self, cfg, logger=None): + super(VideoClassifierHead, self).__init__(cfg, logger=logger) + self.dim = cfg.DIM + self.num_classes = cfg.NUM_CLASSES + self.dropout_rate = cfg.get('DROPOUT_RATE', 0.5) + + if self.dropout_rate > 0.0: + self.dropout = nn.Dropout(self.dropout_rate) + + self.out = nn.Linear(self.dim, self.num_classes, bias=True) + + def forward(self, x, label=None): + if hasattr(self, 'dropout'): + x = self.dropout(x) + out = self.out(x) + + return out + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('HEADS', + __class__.__name__, + VideoClassifierHead.para_dict, + set_name=True) + + +@HEADS.register_class() +class VideoClassifierHeadx2(BaseModel): + para_dict = { + 'DIM': { + 'value': 512, + 'description': 'representation dim!' + }, + 'NUM_CLASSES': { + 'value': [10, 12], + 'description': 'number of classes for two head.' + }, + 'DROPOUT_RATE': { + 'value': 0.0, + 'description': 'dropout rate, default 0.' + } + } + + def __init__(self, cfg, logger=None): + super(VideoClassifierHeadx2, self).__init__(cfg, logger=logger) + self.dim = cfg.DIM + self.num_classes = cfg.NUM_CLASSES + self.dropout_rate = cfg.get('DROPOUT_RATE', 0.5) + assert type(self.num_classes) is list + assert len(self.num_classes) == 2 + + if self.dropout_rate > 0.0: + self.dropout = nn.Dropout(self.dropout_rate) + + self.linear1 = nn.Linear(self.dim, self.num_classes[0], bias=True) + self.linear2 = nn.Linear(self.dim, self.num_classes[1], bias=True) + + def forward(self, x, label=None): + if hasattr(self, 'dropout'): + out = self.dropout(x) + else: + out = x + + out1 = self.linear1(out) + out2 = self.linear2(out) + + return out1, out2 + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('HEADS', + __class__.__name__, + VideoClassifierHeadx2.para_dict, + set_name=True) + + +@HEADS.register_class() +class TransformerHead(BaseModel): + para_dict = { + 'DIM': { + 'value': 512, + 'description': 'representation dim!' + }, + 'NUM_CLASSES': { + 'value': 10, + 'description': 'number of classes.' + }, + 'DROPOUT_RATE': { + 'value': 0.0, + 'description': 'dropout rate, default 0.' + }, + 'PRE_LOGITS': { + 'value': False, + 'description': 'pre logits default False.' + } + } + + def __init__(self, cfg, logger=None): + super(TransformerHead, self).__init__(cfg, logger=logger) + self.dim = cfg.DIM + self.num_classes = cfg.NUM_CLASSES + self.dropout_rate = cfg.get('DROPOUT_RATE', 0.5) + self.pre_logits = cfg.get('PRE_LOGITS', False) + if self.pre_logits: + self.pre_logits = nn.Sequential( + OrderedDict([('fc', nn.Linear(self.dim, self.dim)), + ('act', nn.Tanh())])) + + if self.dropout_rate > 0.0: + self.dropout = nn.Dropout(self.dropout_rate) + + self.linear = nn.Linear(self.dim, self.num_classes, bias=True) + + def forward(self, x, label=None): + if hasattr(self, 'dropout'): + out = self.dropout(x) + else: + out = x + if hasattr(self, 'pre_logits'): + out = self.pre_logits(out) + out = self.linear(out) + + return out + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('HEADS', + __class__.__name__, + TransformerHead.para_dict, + set_name=True) + + +@HEADS.register_class() +class TransformerHeadx2(BaseModel): + para_dict = { + 'DIM': { + 'value': 512, + 'description': 'representation dim!' + }, + 'NUM_CLASSES': { + 'value': [10, 12], + 'description': 'number of classes for two head.' + }, + 'DROPOUT_RATE': { + 'value': 0.0, + 'description': 'dropout rate, default 0.' + }, + 'PRE_LOGITS': { + 'value': False, + 'description': 'pre logits default False.' + } + } + + def __init__(self, cfg, logger=None): + super(TransformerHeadx2, self).__init__(cfg, logger=logger) + self.dim = cfg.DIM + self.num_classes = cfg.NUM_CLASSES + self.dropout_rate = cfg.get('DROPOUT_RATE', 0.5) + self.pre_logits = cfg.get('PRE_LOGITS', False) + assert type(self.num_classes) is list + assert len(self.num_classes) == 2 + if self.pre_logits: + self.pre_logits1 = nn.Sequential( + OrderedDict([('fc', nn.Linear(self.dim, self.dim)), + ('act', nn.Tanh())])) + self.pre_logits2 = nn.Sequential( + OrderedDict([('fc', nn.Linear(self.dim, self.dim)), + ('act', nn.Tanh())])) + + if self.dropout_rate > 0.0: + self.dropout = nn.Dropout(self.dropout_rate) + + self.linear1 = nn.Linear(self.dim, self.num_classes[0], bias=True) + self.linear2 = nn.Linear(self.dim, self.num_classes[1], bias=True) + + def forward(self, x, label=None): + if hasattr(self, 'dropout'): + out = self.dropout(x) + else: + out = x + + if hasattr(self, 'pre_logits1'): + out1 = self.pre_logits1(out) + out2 = self.pre_logits2(out) + else: + out1, out2 = out, out + + out1 = self.linear1(out1) + out2 = self.linear2(out2) + + return out1, out2 + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('HEADS', + __class__.__name__, + TransformerHeadx2.para_dict, + set_name=True) diff --git a/scepter/modules/model/loss/__init__.py b/scepter/modules/model/loss/__init__.py new file mode 100644 index 0000000..bdc7778 --- /dev/null +++ b/scepter/modules/model/loss/__init__.py @@ -0,0 +1,4 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.loss.base_losses import CrossEntropy +from scepter.modules.model.loss.rec_loss import MinSNRLoss, ReconstructLoss diff --git a/scepter/modules/model/loss/base_losses.py b/scepter/modules/model/loss/base_losses.py new file mode 100644 index 0000000..181f72b --- /dev/null +++ b/scepter/modules/model/loss/base_losses.py @@ -0,0 +1,106 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import torch.nn as nn +from packaging import version +from torch.version import __version__ as torch_version + +from scepter.modules.model.registry import LOSSES +from scepter.modules.utils.config import dict_to_yaml + + +@LOSSES.register_class() +class CrossEntropy(nn.Module): + para_dict = { + 'REDUCE': { + 'value': + None, + 'description': + 'reduce is False, returns a loss per batch element instead ' + 'and ignores :attr: size_average. Default: True' + }, + 'SIZE_AVERAGE': { + 'value': + None, + 'description': + 'Deprecated (see :attr: reduction). By default,' + 'the losses are averaged over each loss element in the batch. Note that for' + 'some losses, there are multiple elements per sample. If the field :attr: size_average' + 'is set to False, the losses are instead summed for each minibatch. Ignored' + 'when :attr: reduce is False. Default: True' + }, + 'IGNORE_INDEX': { + 'value': + -100, + 'description': + 'Specifies a target value that is ignored' + 'and does not contribute to the input gradient. When :attr: size_average is' + 'True, the loss is averaged over non-ignored targets. Note that' + ':attr: ignore_index is only applicable when the target contains class indices.' + }, + 'REDUCTION': { + 'value': + 'mean', + 'description': + 'Specifies the reduction to apply to the output:' + "'none' | 'mean' | 'sum'. 'none': no reduction will" + "be applied, 'mean': the weighted mean of the output is taken," + "'sum': the output will be summed. Note: :attr: size_average" + 'and :attr:`reduce` are in the process of being deprecated, and in' + 'the meantime, specifying either of those two args will override' + ":attr:`reduction`. Default: 'mean'" + }, + 'LABEL_SMOOTHING': { + 'value': + 0.0, + 'description': + 'A float in [0.0, 1.0]. Specifies the amount' + 'of smoothing when computing the loss, where 0.0 means no smoothing. ' + } + } + + def __init__(self, cfg, logger=None): + super(CrossEntropy, self).__init__() + self.logger = logger + size_average = cfg.get('SIZE_AVERAGE', None) + ignore_index = cfg.get('IGNORE_INDEX', -100) + reduce = cfg.get('REDUCE', None) + reduction = cfg.get('REDUCTION', 'mean') + + if version.parse(torch_version) >= version.parse('1.10.0'): + label_smoothing = cfg.get('LABEL_SMOOTHING', 0.0) + self.loss_obj = nn.CrossEntropyLoss( + size_average=size_average, + ignore_index=ignore_index, + reduce=reduce, + reduction=reduction, + label_smoothing=label_smoothing) + else: + self.loss_obj = nn.CrossEntropyLoss(size_average=size_average, + ignore_index=ignore_index, + reduce=reduce, + reduction=reduction) + + def forward(self, input, target): + return self.loss_obj(input, target) + + def __repr__(self) -> str: + return f'{self.__class__.__name__}' + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('LOSS', + __class__.__name__, + CrossEntropy.para_dict, + set_name=True) diff --git a/scepter/modules/model/loss/rec_loss.py b/scepter/modules/model/loss/rec_loss.py new file mode 100644 index 0000000..75a9452 --- /dev/null +++ b/scepter/modules/model/loss/rec_loss.py @@ -0,0 +1,65 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import torch +from torch import nn + +from scepter.modules.model.registry import LOSSES +from scepter.modules.utils.config import dict_to_yaml + + +@LOSSES.register_class() +class ReconstructLoss(nn.Module): + para_dict = { + 'LOSS_TYPE': { + 'value': 'l1', + 'description': 'Used loss type l1 or l2.' + } + } + + def __init__(self, cfg, logger=None): + super(ReconstructLoss, self).__init__() + self.loss_type = cfg.get('LOSS_TYPE', 'l2') + + def forward(self, pred, target, mean=True): + if self.loss_type == 'l1': + loss = (target - pred).abs() + if mean: + loss = loss.mean() + elif self.loss_type == 'l2': + if mean: + loss = torch.nn.functional.mse_loss(target, pred) + else: + loss = torch.nn.functional.mse_loss(target, + pred, + reduction='none') + else: + raise NotImplementedError("unknown loss type '{loss_type}'") + + return loss + + @staticmethod + def get_config_template(): + return dict_to_yaml('LOSS', + __class__.__name__, + ReconstructLoss.para_dict, + set_name=True) + + +@LOSSES.register_class() +class MinSNRLoss(ReconstructLoss): + """Only used when parameterization=='eps'""" + para_dict = {'GAMMA': {'value': 5, 'description': 'max value of snr.'}} + para_dict.update(ReconstructLoss.para_dict) + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.gamma = cfg.get('GAMMA', 5) + + def forward(self, pred, target, alphas_cumprod, timesteps, mean=False): + loss = super().forward(pred, target, mean=mean) + + alpha = torch.sqrt(alphas_cumprod) + sigma = torch.sqrt(1.0 - alphas_cumprod) + all_snr = ((alpha / sigma)**2).to(pred.device) + snr_weight = (self.gamma / all_snr[timesteps]).clip(max=1.).float() + return loss * snr_weight.view(-1, 1, 1, 1) diff --git a/scepter/modules/model/metric/__init__.py b/scepter/modules/model/metric/__init__.py new file mode 100644 index 0000000..155ba54 --- /dev/null +++ b/scepter/modules/model/metric/__init__.py @@ -0,0 +1,5 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.metric.classification import (AccuracyMetric, + EnsembleAccuracyMetric + ) diff --git a/scepter/modules/model/metric/base_metric.py b/scepter/modules/model/metric/base_metric.py new file mode 100644 index 0000000..bad62bb --- /dev/null +++ b/scepter/modules/model/metric/base_metric.py @@ -0,0 +1,31 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.utils.config import dict_to_yaml + + +class BaseMetric(object): + para_dict = [{}] + + def __init__(self, cfg, logger=None): + self.logger = logger + + def __repr__(self) -> str: + return f'{self.__class__.__name__}' + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('METRICS', + __class__.__name__, + BaseMetric.para_dict, + set_name=True) diff --git a/scepter/modules/model/metric/classification.py b/scepter/modules/model/metric/classification.py new file mode 100644 index 0000000..fb83ebc --- /dev/null +++ b/scepter/modules/model/metric/classification.py @@ -0,0 +1,165 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from collections import OrderedDict + +import numpy as np +import torch + +from scepter.modules.model.metric.base_metric import BaseMetric +from scepter.modules.model.metric.registry import METRICS +from scepter.modules.utils.config import dict_to_yaml + + +@METRICS.register_class('AccuracyMetric') +class AccuracyMetric(BaseMetric): + para_dict = [{'TOPK': {'value': 1, 'description': 'topk accuracy!'}}] + + def __init__(self, cfg, logger=None): + super(AccuracyMetric, self).__init__(cfg, logger=logger) + topk = cfg.get('TOPK', 1) + if isinstance(topk, int): + topk = (topk, ) + self.topk = topk + self.maxk = max(self.topk) + + @torch.no_grad() + def __call__(self, logits, labels, label_map=None, prefix='acc'): + """ Compute Accuracy + Args: + logits (torch.Tensor or numpy.ndarray): + labels (torch.Tensor or numpy.ndarray): + prefix (str): Prefix string of ret key, default is acc. + + Returns: + A OrderedDict, contains accuracy tensors according to topk. + + """ + assert self.maxk <= logits.shape[-1] + + if isinstance(logits, np.ndarray): + logits = torch.from_numpy(logits) + if isinstance(labels, np.ndarray): + labels = torch.from_numpy(labels) + + batch_size = logits.size(0) + + _, pred = logits.topk(self.maxk, 1, True, True) + if label_map is not None: + pred = torch.gather(label_map, 1, pred) + # print(labels) + # print(pred) + + pred = pred.t() + corrects = pred.eq(labels.view(1, -1).expand_as(pred)) + + res = OrderedDict() + for k in self.topk: + correct_k = corrects[:k].contiguous().view(-1).float().sum(0) + res[f'{prefix}@{k}'] = correct_k.mul_(1.0 / batch_size) + return res + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('METRICS', + __class__.__name__, + AccuracyMetric.para_dict, + set_name=True) + + +@METRICS.register_class('EnsembleAccuracyMetric') +class EnsembleAccuracyMetric(object): + para_dict = [{ + 'TOPK': { + 'value': 1, + 'description': 'topk accuracy!' + }, + 'ENSEMBLE_METHOD': { + 'value': 'avg', + 'description': 'ensemble method from (avg, max)!' + } + }] + + def __init__(self, cfg, logger=None): + topk = cfg.get('TOPK', 1) + ensemble_method = cfg.get('ENSEMBLE_METHOD', 'avg') + if isinstance(topk, int): + topk = (topk, ) + self.topk = topk + self.maxk = max(self.topk) + assert ensemble_method in ( + 'avg', 'max' + ), f"Expected ensemble_method in ('avg', 'max'), got {ensemble_method}" + self.ensemble_method = ensemble_method + + @torch.no_grad() + def __call__(self, logits, labels, keys, prefix='acc'): + """ Compute Accuracy + Args: + logits (torch.Tensor or numpy.ndarray): + labels (torch.Tensor or numpy.ndarray): + keys (List[str]): Keys to accumulate logits. + prefix (str): Prefix string of ret key, default is acc. + + Returns: + A OrderedDict, contains accuracy tensors according to topk. + + """ + if isinstance(logits, np.ndarray): + logits = torch.from_numpy(logits) + if isinstance(labels, np.ndarray): + labels = torch.from_numpy(labels) + + agg_keys = list(set(keys)) + keys = np.asarray(keys) + + agg_logits = [] # N * Tensor([C]) + agg_labels = [] # N * Tensor(scalar) + + for key in agg_keys: + key_index = np.where(keys == key)[0] + key_index = torch.from_numpy(key_index) + key_logits = logits[key_index] + + if self.ensemble_method == 'avg': + key_logit = torch.mean(key_logits, dim=0) + else: + key_logit, _ = torch.max(key_logit, dim=0) + key_label = labels[key_index[0]] + + agg_logits.append(key_logit) + agg_labels.append(key_label) + + agg_logits = torch.vstack(agg_logits) + agg_labels = torch.hstack(agg_labels) + + return AccuracyMetric(self.topk)(agg_logits, agg_labels, prefix) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('METRICS', + __class__.__name__, + EnsembleAccuracyMetric.para_dict, + set_name=True) diff --git a/scepter/modules/model/metric/registry.py b/scepter/modules/model/metric/registry.py new file mode 100644 index 0000000..a6821ba --- /dev/null +++ b/scepter/modules/model/metric/registry.py @@ -0,0 +1,6 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.utils.registry import Registry + +METRICS = Registry('METRICS', allow_types=('class', )) diff --git a/scepter/modules/model/neck/__init__.py b/scepter/modules/model/neck/__init__.py new file mode 100644 index 0000000..258155d --- /dev/null +++ b/scepter/modules/model/neck/__init__.py @@ -0,0 +1,5 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.neck.global_average_pooling import \ + GlobalAveragePooling +from scepter.modules.model.neck.identity import Identity diff --git a/scepter/modules/model/neck/global_average_pooling.py b/scepter/modules/model/neck/global_average_pooling.py new file mode 100644 index 0000000..41da23c --- /dev/null +++ b/scepter/modules/model/neck/global_average_pooling.py @@ -0,0 +1,64 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import torch.nn as nn + +from scepter.modules.model.base_model import BaseModel +from scepter.modules.model.registry import NECKS +from scepter.modules.utils.config import dict_to_yaml + + +@NECKS.register_class() +class GlobalAveragePooling(BaseModel): + """Global Average Pooling neck. + Args: + dim (int): Dimensions of each sample channel, can be one of {1, 2, 3}. + Default: 2 + """ + para_dict = { + 'DIM': { + 'value': 2, + 'description': 'GlobalAveragePooling dim!' + } + } + + def __init__(self, cfg, logger=None): + super(GlobalAveragePooling, self).__init__(cfg, logger=logger) + dim = cfg.get('DIM', 2) + assert dim in [1, 2, 3], 'GlobalAveragePooling dim only support ' \ + f'{1, 2, 3}, get {dim} instead.' + if dim == 1: + self.gap = nn.AdaptiveAvgPool1d(1) + elif dim == 2: + self.gap = nn.AdaptiveAvgPool2d((1, 1)) + else: + self.gap = nn.AdaptiveAvgPool3d((1, 1, 1)) + + def infer(self, x): + if x.ndim == 2: + return x + return self.gap(x).view(x.size(0), -1) + + def forward(self, inputs): + if isinstance(inputs, tuple): + return tuple([self.infer(x) for x in inputs]) + else: + return self.infer(inputs) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('NECKS', + __class__.__name__, + GlobalAveragePooling.para_dict, + set_name=True) diff --git a/scepter/modules/model/neck/identity.py b/scepter/modules/model/neck/identity.py new file mode 100644 index 0000000..eebd452 --- /dev/null +++ b/scepter/modules/model/neck/identity.py @@ -0,0 +1,35 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.model.base_model import BaseModel +from scepter.modules.model.registry import NECKS +from scepter.modules.utils.config import dict_to_yaml + + +@NECKS.register_class() +class Identity(BaseModel): + para_dict = {} + + def __init__(self, cfg, logger=None): + super(Identity, self).__init__(cfg, logger=logger) + + def forward(self, inputs): + return inputs + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('neckname', + __class__.__name__, + Identity.para_dict, + set_name=True) diff --git a/scepter/modules/model/network/__init__.py b/scepter/modules/model/network/__init__.py new file mode 100644 index 0000000..dda2aba --- /dev/null +++ b/scepter/modules/model/network/__init__.py @@ -0,0 +1,7 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.network.autoencoder import ae_kl +from scepter.modules.model.network.classifier import Classifier +from scepter.modules.model.network.diffusion import (diffusion, schedules, + solvers) +from scepter.modules.model.network.ldm import ldm, ldm_xl diff --git a/scepter/modules/model/network/autoencoder/__init__.py b/scepter/modules/model/network/autoencoder/__init__.py new file mode 100644 index 0000000..c0708a2 --- /dev/null +++ b/scepter/modules/model/network/autoencoder/__init__.py @@ -0,0 +1,3 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.network.autoencoder.ae_kl import AutoencoderKL diff --git a/scepter/modules/model/network/autoencoder/ae_kl.py b/scepter/modules/model/network/autoencoder/ae_kl.py new file mode 100644 index 0000000..77bd9e3 --- /dev/null +++ b/scepter/modules/model/network/autoencoder/ae_kl.py @@ -0,0 +1,268 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from collections import OrderedDict + +import numpy as np +import torch + +from scepter.modules.model.network.train_module import TrainModule +from scepter.modules.model.registry import BACKBONES, LOSSES, MODELS +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we +from scepter.modules.utils.file_system import FS + + +class DiagonalGaussianDistribution(object): + def __init__(self, mean, logvar, deterministic=False): + self.mean = mean + self.logvar = torch.clamp(logvar, -30.0, 20.0) + self.deterministic = deterministic + self.std = torch.exp(0.5 * self.logvar) + self.var = torch.exp(self.logvar) + if self.deterministic: + self.var = self.std = torch.zeros_like( + self.mean).to(device=self.mean.device) + + def sample(self): + x = self.mean + self.std * torch.randn( + self.mean.shape).to(device=self.mean.device) + return x + + def kl(self, other=None): + if self.deterministic: + return torch.Tensor([0.]) + else: + if other is None: + return 0.5 * torch.sum( + torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar, + dim=[1, 2, 3]) + else: + return 0.5 * torch.sum( + torch.pow(self.mean - other.mean, 2) / other.var + + self.var / other.var - 1.0 - self.logvar + other.logvar, + dim=[1, 2, 3]) + + def nll(self, sample, dims=[1, 2, 3]): + if self.deterministic: + return torch.Tensor([0.]) + logtwopi = np.log(2.0 * np.pi) + return 0.5 * torch.sum(logtwopi + self.logvar + + torch.pow(sample - self.mean, 2) / self.var, + dim=dims) + + def mode(self): + print('*** use DiagonalGaussianDistribution.mode() ***') + return self.mean + + +@MODELS.register_class() +class AutoencoderKL(TrainModule): + para_dict = { + 'ENCODER': {}, + 'DECODER': {}, + 'LOSS': {}, + 'EMBED_DIM': { + 'value': 4, + 'description': '' + }, + 'PRETRAINED_MODEL': { + 'value': None, + 'description': '' + }, + 'IGNORE_KEYS': { + 'value': [], + 'description': '' + }, + 'BATCH_SIZE': { + 'value': 16, + 'description': '' + }, + } + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.encoder_cfg = self.cfg.ENCODER + self.decoder_cfg = self.cfg.DECODER + self.loss_cfg = self.cfg.get('LOSS', None) + self.embed_dim = self.cfg.get('EMBED_DIM', 4) + self.pretrained_model = self.cfg.get('PRETRAINED_MODEL', None) + self.ignore_keys = self.cfg.get('IGNORE_KEYS', []) + self.batch_size = self.cfg.get('BATCH_SIZE', 16) + + self.construct_network() + self.init_network() + + def construct_network(self): + z_channels = self.encoder_cfg.Z_CHANNELS + self.encoder = BACKBONES.build(self.encoder_cfg, logger=self.logger) + self.decoder = BACKBONES.build(self.decoder_cfg, logger=self.logger) + self.conv1 = torch.nn.Conv2d(2 * z_channels, 2 * self.embed_dim, 1) + self.conv2 = torch.nn.Conv2d(self.embed_dim, z_channels, 1) + + if self.loss_cfg is not None: + self.loss = LOSSES.build(self.loss_cfg, logger=self.logger) + + def init_network(self): + if self.pretrained_model is not None: + with FS.get_from(self.pretrained_model, + wait_finish=True) as local_model: + self.init_from_ckpt(local_model, ignore_keys=self.ignore_keys) + + def init_from_ckpt(self, path, ignore_keys=list()): + if path.find('.safetensors') > -1: + from safetensors import safe_open + sd = OrderedDict() + with safe_open(path, framework='pt', device='cpu') as f: + for k in f.keys(): + sd[k] = f.get_tensor(k) + else: + sd = torch.load(path, map_location='cpu') + if path.find('.pt') > -1 and 'state_dict' in sd: + sd = sd['state_dict'] + elif path.find('.ckpt') > -1 and 'state_dict' in sd: + sd = sd['state_dict'] + + 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 + k = k.replace('post_quant_conv', + 'conv2') if 'post_quant_conv' in k else k + k = k.replace('quant_conv', 'conv1') if 'quant_conv' in k else k + if not ignored: + new_sd[k] = v + + missing, unexpected = self.load_state_dict(new_sd, strict=False) + if we.rank == 0: + self.logger.info( + f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys' + ) + if len(missing) > 0: + self.logger.info(f'Missing Keys:\n {missing}') + if len(unexpected) > 0: + self.logger.info(f'\nUnexpected Keys:\n {unexpected}') + + def encode(self, x, return_mom=False): + return torch.cat([ + self._encode(batch, return_mom=return_mom) + for batch in x.split(self.batch_size, dim=0) + ], + dim=0) + + def decode(self, z): + return torch.cat( + [self._decode(batch) for batch in z.split(self.batch_size, dim=0)], + dim=0) + + def sample(self, moments): + mean, logvar = torch.chunk(moments, 2, dim=1) + posterior = DiagonalGaussianDistribution(mean, logvar) + z = posterior.sample() + return z + + def _encode(self, x, return_mom=False): + h = self.encoder(x) + moments = self.conv1(h) + if return_mom: + return moments + mean, logvar = torch.chunk(moments, 2, dim=1) + posterior = DiagonalGaussianDistribution(mean, logvar) + z = posterior.sample() + return z + + def _decode(self, z): + z = self.conv2(z) + dec = self.decoder(z) + return dec + + def share_forward(self, image, sample_posterior=True): + posterior = self.encode(image) + if sample_posterior: + z = posterior.sample() + else: + z = posterior.mode() + dec = self.decode(z) + return dec, posterior + + def forward(self, **kwargs): + if self.training: + ret = self.forward_train(**kwargs) + else: + ret = self.forward_test(**kwargs) + return ret + + def forward_train(self, + image=None, + sample_posterior=True, + optimizer_idx=0, + **kwargs): + reconstructions, posterior = self.share_forward( + image, sample_posterior) + ret = {} + if optimizer_idx == 0: + # train encoder+decoder+logvar + aeloss, log_dict_ae = self.loss(image, + reconstructions, + posterior, + optimizer_idx, + self.global_step, + last_layer=self.get_last_layer(), + split='train') + ret['loss'] = aeloss + ret.update(log_dict_ae) + + if optimizer_idx == 1: + # train the discriminator + discloss, log_dict_disc = self.loss( + image, + reconstructions, + posterior, + optimizer_idx, + self.global_step, + last_layer=self.get_last_layer(), + split='train') + ret['loss'] = discloss + ret.update(log_dict_disc) + + return ret + + def forward_test(self, image=None, sample_posterior=True, **kwargs): + reconstructions, posterior = self.share_forward( + image, sample_posterior) + ret = {} + aeloss, log_dict_ae = self.loss(image, + reconstructions, + posterior, + 0, + self.global_step, + last_layer=self.get_last_layer(), + split='val') + discloss, log_dict_disc = self.loss(image, + reconstructions, + posterior, + 1, + self.global_step, + last_layer=self.get_last_layer(), + split='val') + + ret.update(log_dict_ae) + ret.update(log_dict_disc) + + return ret + + def get_last_layer(self): + return self.decoder.conv_out.weight + + @staticmethod + def get_config_template(): + return dict_to_yaml('MODEL', + __class__.__name__, + AutoencoderKL.para_dict, + set_name=True) diff --git a/scepter/modules/model/network/classifier.py b/scepter/modules/model/network/classifier.py new file mode 100644 index 0000000..979bd6a --- /dev/null +++ b/scepter/modules/model/network/classifier.py @@ -0,0 +1,208 @@ +# -*- coding: utf-8 -*- +# Copyright 2021 Alibaba Group Holding Limited. All Rights Reserved. + +from collections import OrderedDict +from functools import partial + +import torch.nn as nn +from torch.nn.functional import sigmoid, softmax + +from scepter.modules.model.metric.registry import METRICS +from scepter.modules.model.network.train_module import TrainModule +from scepter.modules.model.registry import (BACKBONES, HEADS, LOSSES, MODELS, + NECKS) +from scepter.modules.utils.config import Config, dict_to_yaml + +_ACTIVATE_MAPPER = {'softmax': partial(softmax, dim=1), 'sigmoid': sigmoid} + + +@MODELS.register_class() +class Classifier(TrainModule): + """ Base classifier implementation. + + Args: + backbones (dict): Defines backbones. + neck (dict, optional): Defines neck. Use Identity if none. + head (dict): Defines head. + act_name (str): Defines activate function, 'softmax' or 'sigmoid'. + topk (Sequence[int]): Defines how to calculate accuracy metrics. + freeze_bn (bool): If True, freeze all BatchNorm layers including LayerNorm. + """ + para_dict = { + 'ACT_NAME': { + 'value': + 'softmax', + 'description': + 'the activation function for logits, select from [softmax, sigmoid]!' + }, + 'FREEZE_BN': { + 'value': False, + 'description': 'if freeze bn of not' + } + } + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + # Construct model + self.backbone = BACKBONES.build(cfg.BACKBONE, logger=logger) + necks_cfg = cfg.get('NECK', + Config(cfg_dict={'NAME': 'Identity'}, load=False)) + self.neck = NECKS.build(necks_cfg, logger=logger) + self.head = HEADS.build(cfg.HEAD, logger=logger) + freeze_bn = cfg.get('FREEZE_BN', False) + # Construct loss + loss = cfg.get('LOSS', + Config(cfg_dict={'NAME': 'CrossEntropy'}, load=False)) + self.loss = LOSSES.build(loss, logger=logger) + act_name = cfg.get('ACT_NAME', 'softmax') + # Construct activate function + self.act_fn = _ACTIVATE_MAPPER[act_name] + self.metric = METRICS.build(cfg.METRIC, logger=logger) + self.freeze_bn = freeze_bn + + def train(self, mode=True): + self.training = mode + super(Classifier, self).train(mode=mode) + if self.freeze_bn: + for module in self.modules(): + if isinstance(module, + (nn.BatchNorm2d, nn.BatchNorm3d, nn.LayerNorm)): + module.train(False) + return self + + def forward(self, img, label=None, **kwargs): + return self.forward_train( + img, label=label) if self.training else self.forward_test( + img, label=label) # noqa + + def forward_train(self, img, label=None): + probs = self.head(self.neck(self.backbone(img))) + if label is None: + return probs + + ret = OrderedDict() + loss = self.loss(probs, label) + ret['loss'] = loss + ret['batch_size'] = img.size(0) + ret.update(self.metric(probs, label)) + return ret + + def forward_test(self, img, label=None): + logits = self.act_fn(self.head(self.neck(self.backbone(img)))) + if label is not None: + ret = OrderedDict() + ret['logits'] = logits + ret['batch_size'] = img.size(0) + ret.update(self.metric(logits, label)) + return ret + return logits + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('MODEL', + __class__.__name__, + Classifier.para_dict, + set_name=True) + + +@MODELS.register_class() +class VideoClassifier(Classifier): + """ Classifier for video. + Default input tensor is video. + + """ + def forward(self, video, label=None, **kwargs): + return self.forward_train(video, label=label) \ + if self.training else self.forward_test(video, label=label) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('MODEL', + __class__.__name__, + VideoClassifier.para_dict, + set_name=True) + + +@MODELS.register_class() +class VideoClassifier2x(VideoClassifier): + """ A 2-way classifier for video. + + """ + def forward_train(self, video, label=None): + probs0, probs1 = self.head(self.neck(self.backbone(video))) + if label is not None: + ret = OrderedDict() + loss = self.loss(probs0, label[:, 0]) + self.loss( + probs1, label[:, 1]) + ret['loss'] = loss + ret['batch_size'] = video.size(0) + acc_0 = self.metric(probs0, label[:, 0]) + acc_0 = { + key.relace('@', '_0@'): value + for key, value in acc_0.items() + } + acc_1 = self.metric(probs1, label[:, 1]) + acc_1 = { + key.relace('@', '_1@'): value + for key, value in acc_1.items() + } + ret.update(acc_0) + ret.update(acc_1) + return ret + return {'logits0': self.act_fn(probs0), 'logits1': self.act_fn(probs1)} + + def forward_test(self, video, label=None): + probs0, probs1 = self.head(self.neck(self.backbone(video))) + logits0, logits1 = self.act_fn(probs0), self.act_fn(probs1) + if label is None: + return {'logits0': logits0, 'logits1': logits1} + ret = OrderedDict() + ret['logits0'] = logits0 + ret['logits1'] = logits1 + acc_0 = self.metric(probs0, label[:, 0]) + acc_0 = {key.relace('@', '_0@'): value for key, value in acc_0.items()} + acc_1 = self.metric(probs1, label[:, 1]) + acc_1 = {key.relace('@', '_1@'): value for key, value in acc_1.items()} + ret.update(acc_0) + ret.update(acc_1) + return ret + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('MODEL', + __class__.__name__, + VideoClassifier2x.para_dict, + set_name=True) diff --git a/scepter/modules/model/network/diffusion/__init__.py b/scepter/modules/model/network/diffusion/__init__.py new file mode 100644 index 0000000..f00e199 --- /dev/null +++ b/scepter/modules/model/network/diffusion/__init__.py @@ -0,0 +1,4 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.network.diffusion import (diffusion, schedules, + solvers) diff --git a/scepter/modules/model/network/diffusion/diffusion.py b/scepter/modules/model/network/diffusion/diffusion.py new file mode 100644 index 0000000..ac0240d --- /dev/null +++ b/scepter/modules/model/network/diffusion/diffusion.py @@ -0,0 +1,555 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +""" +GaussianDiffusion wraps operators for denoising diffusion models, including the +diffusion and denoising processes, as well as the loss evaluation. +""" +import copy +import random + +import torch + +from .schedules import karras_schedule +from .solvers import (sample_ddim, sample_dpm_2, sample_dpm_2_ancestral, + sample_dpmpp_2m, sample_dpmpp_2m_sde, + sample_dpmpp_2s_ancestral, sample_dpmpp_sde, + sample_euler, sample_euler_ancestral, sample_heun, + sample_img2img_euler, sample_img2img_euler_ancestral) + +__all__ = ['GaussianDiffusion'] + + +def _i(tensor, t, x): + """ + Index tensor using t and format the output according to x. + """ + shape = (x.size(0), ) + (1, ) * (x.ndim - 1) + return tensor[t.to(tensor.device)].view(shape).to(x.device) + + +class GaussianDiffusion(object): + def __init__(self, sigmas, prediction_type='eps'): + assert prediction_type in {'x0', 'eps', 'v'} + self.sigmas = sigmas # noise coefficients + self.alphas = torch.sqrt(1 - sigmas**2) # signal coefficients + self.num_timesteps = len(sigmas) + self.prediction_type = prediction_type + + def diffuse(self, x0, t, noise=None): + """ + Add Gaussian noise to signal x0 according to: + q(x_t | x_0) = N(x_t | alpha_t x_0, sigma_t^2 I). + """ + noise = torch.randn_like(x0) if noise is None else noise + xt = _i(self.alphas, t, x0) * x0 + _i(self.sigmas, t, x0) * noise + return xt + + def denoise(self, + xt, + t, + s, + model, + model_kwargs={}, + guide_scale=None, + guide_rescale=None, + clamp=None, + percentile=None, + cat_uc=False): + """ + Apply one step of denoising from the posterior distribution q(x_s | x_t, x0). + Since x0 is not available, estimate the denoising results using the learned + distribution p(x_s | x_t, \hat{x}_0 == f(x_t)). # noqa + """ + s = t - 1 if s is None else s + + # hyperparams + sigmas = _i(self.sigmas, t, xt) + alphas = _i(self.alphas, t, xt) + alphas_s = _i(self.alphas, s.clamp(0), xt) + alphas_s[s < 0] = 1. + sigmas_s = torch.sqrt(1 - alphas_s**2) + + # precompute variables + betas = 1 - (alphas / alphas_s)**2 + coef1 = betas * alphas_s / sigmas**2 + coef2 = (alphas * sigmas_s**2) / (alphas_s * sigmas**2) + var = betas * (sigmas_s / sigmas)**2 + log_var = torch.log(var).clamp_(-20, 20) + + # prediction + if guide_scale is None: + assert isinstance(model_kwargs, dict) + out = model(xt, t=t, **model_kwargs) + else: + # classifier-free guidance (arXiv:2207.12598) + # model_kwargs[0]: conditional kwargs + # model_kwargs[1]: non-conditional kwargs + assert isinstance(model_kwargs, list) and len(model_kwargs) == 2 + + if guide_scale == 1.: + out = model(xt, t=t, **model_kwargs[0]) + else: + if cat_uc: + + def parse_model_kwargs(prev_value, value): + if isinstance(value, torch.Tensor): + prev_value = torch.cat([prev_value, value], dim=0) + elif isinstance(value, dict): + for k, v in value.items(): + prev_value[k] = parse_model_kwargs( + prev_value[k], v) + elif isinstance(value, list): + for idx, v in enumerate(value): + prev_value[idx] = parse_model_kwargs( + prev_value[idx], v) + return prev_value + + all_model_kwargs = copy.deepcopy(model_kwargs[0]) + for model_kwarg in model_kwargs[1:]: + for key, value in model_kwarg.items(): + all_model_kwargs[key] = parse_model_kwargs( + all_model_kwargs[key], value) + all_out = model(xt.repeat(2, 1, 1, 1), + t=t.repeat(2), + **all_model_kwargs) + y_out, u_out = all_out.chunk(2) + else: + y_out = model(xt, t=t, **model_kwargs[0]) + u_out = model(xt, t=t, **model_kwargs[1]) + out = u_out + guide_scale * (y_out - u_out) + + # rescale the output according to arXiv:2305.08891 + if guide_rescale is not None: + assert guide_rescale >= 0 and guide_rescale <= 1 + ratio = (y_out.flatten(1).std(dim=1) / + (out.flatten(1).std(dim=1) + + 1e-12)).view((-1, ) + (1, ) * (y_out.ndim - 1)) + out *= guide_rescale * ratio + (1 - guide_rescale) * 1.0 + # compute x0 + if self.prediction_type == 'x0': + x0 = out + elif self.prediction_type == 'eps': + x0 = (xt - sigmas * out) / alphas + elif self.prediction_type == 'v': + x0 = alphas * xt - sigmas * out + else: + raise NotImplementedError( + f'prediction_type {self.prediction_type} not implemented') + + # restrict the range of x0 + if percentile is not None: + # NOTE: percentile should only be used when data is within range [-1, 1] + assert percentile > 0 and percentile <= 1 + s = torch.quantile(x0.flatten(1).abs(), percentile, dim=1) + s = s.clamp_(1.0).view((-1, ) + (1, ) * (xt.ndim - 1)) + x0 = torch.min(s, torch.max(-s, x0)) / s + elif clamp is not None: + x0 = x0.clamp(-clamp, clamp) + + # recompute eps using the restricted x0 + eps = (xt - alphas * x0) / sigmas + + # compute mu (mean of posterior distribution) using the restricted x0 + mu = coef1 * x0 + coef2 * xt + return mu, var, log_var, x0, eps + + def loss(self, + x0, + t, + model, + model_kwargs={}, + reduction='mean', + noise=None): + # hyperparams + sigmas = _i(self.sigmas, t, x0) + alphas = _i(self.alphas, t, x0) + + # diffuse and denoise + if noise is None: + noise = torch.randn_like(x0) + xt = self.diffuse(x0, t, noise) + out = model(xt, t=t, **model_kwargs) + + # mse loss + target = { + 'eps': noise, + 'x0': x0, + 'v': alphas * noise - sigmas * x0 + }[self.prediction_type] + loss = (out - target).pow(2) + if reduction == 'mean': + loss = loss.flatten(1).mean(dim=1) + return loss + + @torch.no_grad() + def sample(self, + noise, + model, + x=None, + denoising_strength=1.0, + refine_stage=False, + refine_strength=0.0, + model_kwargs={}, + condition_fn=None, + guide_scale=None, + guide_rescale=None, + clamp=None, + percentile=None, + solver='euler_a', + steps=20, + t_max=None, + t_min=None, + discretization=None, + discard_penultimate_step=None, + return_intermediate=None, + show_progress=False, + seed=-1, + intermediate_callback=None, + cat_uc=False, + **kwargs): + # sanity check + assert isinstance(steps, (int, torch.LongTensor)) + assert t_max is None or (t_max > 0 and t_max <= self.num_timesteps - 1) + assert t_min is None or (t_min >= 0 and t_min < self.num_timesteps - 1) + assert discretization in (None, 'leading', 'linspace', 'trailing') + assert discard_penultimate_step in (None, True, False) + assert return_intermediate in (None, 'x0', 'xt') + + # function of diffusion solver + solver_fn = { + 'ddim': sample_ddim, + 'euler_ancestral': sample_euler_ancestral, + 'euler': sample_euler, + 'heun': sample_heun, + 'dpm2': sample_dpm_2, + 'dpm2_ancestral': sample_dpm_2_ancestral, + 'dpmpp_2s_ancestral': sample_dpmpp_2s_ancestral, + 'dpmpp_2m': sample_dpmpp_2m, + 'dpmpp_sde': sample_dpmpp_sde, + 'dpmpp_2m_sde': sample_dpmpp_2m_sde, + 'dpm2_karras': sample_dpm_2, + 'dpm2_ancestral_karras': sample_dpm_2_ancestral, + 'dpmpp_2s_ancestral_karras': sample_dpmpp_2s_ancestral, + 'dpmpp_2m_karras': sample_dpmpp_2m, + 'dpmpp_sde_karras': sample_dpmpp_sde, + 'dpmpp_2m_sde_karras': sample_dpmpp_2m_sde + }[solver] + + # options + schedule = 'karras' if 'karras' in solver else None + discretization = discretization or 'linspace' + seed = seed if seed >= 0 else random.randint(0, 2**31) + if isinstance(steps, torch.LongTensor): + discard_penultimate_step = False + if discard_penultimate_step is None: + discard_penultimate_step = True if solver in ( + 'dpm2', 'dpm2_ancestral', 'dpmpp_2m_sde', 'dpm2_karras', + 'dpm2_ancestral_karras', 'dpmpp_2m_sde_karras') else False + + # function for denoising xt to get x0 + intermediates = [] + + def model_fn(xt, sigma): + # denoising + t = self._sigma_to_t(sigma).repeat(len(xt)).round().long() + x0 = self.denoise(xt, + t, + None, + model, + model_kwargs, + guide_scale, + guide_rescale, + clamp, + percentile, + cat_uc=cat_uc)[-2] + + # collect intermediate outputs + if return_intermediate == 'xt': + intermediates.append(xt) + elif return_intermediate == 'x0': + intermediates.append(x0) + if intermediate_callback is not None: + intermediate_callback(intermediates[-1]) + return x0 + + # get timesteps + if isinstance(steps, int): + steps += 1 if discard_penultimate_step else 0 + t_max = self.num_timesteps - 1 if t_max is None else t_max + t_min = 0 if t_min is None else t_min + + # discretize timesteps + if discretization == 'leading': + steps = torch.arange(t_min, t_max + 1, + (t_max - t_min + 1) / steps).flip(0) + elif discretization == 'linspace': + steps = torch.linspace(t_max, t_min, steps) + elif discretization == 'trailing': + steps = torch.arange(t_max, t_min - 1, + -((t_max - t_min + 1) / steps)) + else: + raise NotImplementedError( + f'{discretization} discretization not implemented') + steps = steps.clamp_(t_min, t_max) + steps = torch.as_tensor(steps, + dtype=torch.float32, + device=noise.device) + + # get sigmas + sigmas = self._t_to_sigma(steps) + sigmas = torch.cat([sigmas, sigmas.new_zeros([1])]) + t_enc = int(min(denoising_strength, 0.999) * len(steps)) + sigmas = sigmas[len(steps) - t_enc - 1:] + if refine_strength > 0: + t_refine = int(min(refine_strength, 0.999) * len(steps)) + if refine_stage: + sigmas = sigmas[-t_refine:] + else: + sigmas = sigmas[:-t_refine + 1] + # print(sigmas) + if x is not None: + noise = (x + noise * sigmas[0]) / torch.sqrt(1.0 + sigmas[0]**2.0) + + if schedule == 'karras': + if sigmas[0] == float('inf'): + sigmas = karras_schedule( + n=len(steps) - 1, + sigma_min=sigmas[sigmas > 0].min().item(), + sigma_max=sigmas[sigmas < float('inf')].max().item(), + rho=7.).to(sigmas) + sigmas = torch.cat([ + sigmas.new_tensor([float('inf')]), sigmas, + sigmas.new_zeros([1]) + ]) + else: + sigmas = karras_schedule( + n=len(steps), + sigma_min=sigmas[sigmas > 0].min().item(), + sigma_max=sigmas.max().item(), + rho=7.).to(sigmas) + sigmas = torch.cat([sigmas, sigmas.new_zeros([1])]) + if discard_penultimate_step: + sigmas = torch.cat([sigmas[:-2], sigmas[-1:]]) + kwargs['seed'] = seed + # sampling + x0 = solver_fn(noise, + model_fn, + sigmas, + show_progress=show_progress, + **kwargs) + return (x0, intermediates) if return_intermediate is not None else x0 + + def _sigma_to_t(self, sigma): + if sigma == float('inf'): + t = torch.full_like(sigma, len(self.sigmas) - 1) + else: + log_sigmas = torch.sqrt(self.sigmas**2 / + (1 - self.sigmas**2)).log().to(sigma) + log_sigma = sigma.log() + dists = log_sigma - log_sigmas[:, None] + low_idx = dists.ge(0).cumsum(dim=0).argmax(dim=0).clamp( + max=log_sigmas.shape[0] - 2) + high_idx = low_idx + 1 + low, high = log_sigmas[low_idx], log_sigmas[high_idx] + w = (low - log_sigma) / (low - high) + w = w.clamp(0, 1) + t = (1 - w) * low_idx + w * high_idx + t = t.view(sigma.shape) + if t.ndim == 0: + t = t.unsqueeze(0) + return t + + def _t_to_sigma(self, t): + t = t.float() + low_idx, high_idx, w = t.floor().long(), t.ceil().long(), t.frac() + log_sigmas = torch.sqrt(self.sigmas**2 / + (1 - self.sigmas**2)).log().to(t) + log_sigma = (1 - w) * log_sigmas[low_idx] + w * log_sigmas[high_idx] + log_sigma[torch.isnan(log_sigma) + | torch.isinf(log_sigma)] = float('inf') + return log_sigma.exp() + + @torch.no_grad() + def stochastic_encode(self, x0, t, steps): + # fast, but does not allow for exact reconstruction + # t serves as an index to gather the correct alphas + + t_max = None + t_min = None + + # discretization method + discretization = 'trailing' if self.prediction_type == 'v' else 'leading' + + # timesteps + if isinstance(steps, int): + t_max = self.num_timesteps - 1 if t_max is None else t_max + t_min = 0 if t_min is None else t_min + steps = discretize_timesteps(t_max, t_min, steps, discretization) + steps = torch.as_tensor(steps).round().long().flip(0).to(x0.device) + # steps = torch.as_tensor(steps).round().long().to(x0.device) + + # self.alphas_bar = torch.cumprod(1 - self.sigmas ** 2, dim=0) + # print('sigma: ', self.sigmas, len(self.sigmas)) + # print('alpha_bar: ', self.alphas_bar, len(self.alphas_bar)) + # print('steps: ', steps, len(steps)) + # sqrt_alphas_cumprod = torch.sqrt(self.alphas_bar).to(x0.device)[steps] + # sqrt_one_minus_alphas_cumprod = torch.sqrt(1 - self.alphas_bar).to(x0.device)[steps] + + sqrt_alphas_cumprod = self.alphas.to(x0.device)[steps] + sqrt_one_minus_alphas_cumprod = self.sigmas.to(x0.device)[steps] + # print('sigma: ', self.sigmas, len(self.sigmas)) + # print('alpha: ', self.alphas, len(self.alphas)) + # print('steps: ', steps, len(steps)) + + noise = torch.randn_like(x0) + return ( + extract_into_tensor(sqrt_alphas_cumprod, t, x0.shape) * x0 + + extract_into_tensor(sqrt_one_minus_alphas_cumprod, t, x0.shape) * + noise) + + @torch.no_grad() + def sample_img2img(self, + x, + noise, + model, + denoising_strength=1, + model_kwargs={}, + condition_fn=None, + guide_scale=None, + guide_rescale=None, + clamp=None, + percentile=None, + solver='euler_a', + steps=20, + t_max=None, + t_min=None, + discretization=None, + discard_penultimate_step=None, + return_intermediate=None, + show_progress=False, + seed=-1, + **kwargs): + # sanity check + assert isinstance(steps, (int, torch.LongTensor)) + assert t_max is None or (t_max > 0 and t_max <= self.num_timesteps - 1) + assert t_min is None or (t_min >= 0 and t_min < self.num_timesteps - 1) + assert discretization in (None, 'leading', 'linspace', 'trailing') + assert discard_penultimate_step in (None, True, False) + assert return_intermediate in (None, 'x0', 'xt') + # function of diffusion solver + solver_fn = { + 'euler_ancestral': sample_img2img_euler_ancestral, + 'euler': sample_img2img_euler, + }[solver] + # options + schedule = 'karras' if 'karras' in solver else None + discretization = discretization or 'linspace' + seed = seed if seed >= 0 else random.randint(0, 2**31) + if isinstance(steps, torch.LongTensor): + discard_penultimate_step = False + if discard_penultimate_step is None: + discard_penultimate_step = True if solver in ( + 'dpm2', 'dpm2_ancestral', 'dpmpp_2m_sde', 'dpm2_karras', + 'dpm2_ancestral_karras', 'dpmpp_2m_sde_karras') else False + + # function for denoising xt to get x0 + intermediates = [] + + def get_scalings(sigma): + c_out = -sigma + c_in = 1 / (sigma**2 + 1.**2)**0.5 + return c_out, c_in + + def model_fn(xt, sigma): + # denoising + c_out, c_in = get_scalings(sigma) + t = self._sigma_to_t(sigma).repeat(len(xt)).round().long() + + x0 = self.denoise(xt * c_in, t, None, model, model_kwargs, + guide_scale, guide_rescale, clamp, + percentile)[-2] + # collect intermediate outputs + if return_intermediate == 'xt': + intermediates.append(xt) + elif return_intermediate == 'x0': + intermediates.append(x0) + return xt + x0 * c_out + + # get timesteps + if isinstance(steps, int): + steps += 1 if discard_penultimate_step else 0 + t_max = self.num_timesteps - 1 if t_max is None else t_max + t_min = 0 if t_min is None else t_min + # discretize timesteps + if discretization == 'leading': + steps = torch.arange(t_min, t_max + 1, + (t_max - t_min + 1) / steps).flip(0) + elif discretization == 'linspace': + steps = torch.linspace(t_max, t_min, steps) + elif discretization == 'trailing': + steps = torch.arange(t_max, t_min - 1, + -((t_max - t_min + 1) / steps)) + else: + raise NotImplementedError( + f'{discretization} discretization not implemented') + steps = steps.clamp_(t_min, t_max) + steps = torch.as_tensor(steps, dtype=torch.float32, device=x.device) + # get sigmas + sigmas = self._t_to_sigma(steps) + sigmas = torch.cat([sigmas, sigmas.new_zeros([1])]) + t_enc = int(min(denoising_strength, 0.999) * len(steps)) + sigmas = sigmas[len(steps) - t_enc - 1:] + noise = x + noise * sigmas[0] + + if schedule == 'karras': + if sigmas[0] == float('inf'): + sigmas = karras_schedule( + n=len(steps) - 1, + sigma_min=sigmas[sigmas > 0].min().item(), + sigma_max=sigmas[sigmas < float('inf')].max().item(), + rho=7.).to(sigmas) + sigmas = torch.cat([ + sigmas.new_tensor([float('inf')]), sigmas, + sigmas.new_zeros([1]) + ]) + else: + sigmas = karras_schedule( + n=len(steps), + sigma_min=sigmas[sigmas > 0].min().item(), + sigma_max=sigmas.max().item(), + rho=7.).to(sigmas) + sigmas = torch.cat([sigmas, sigmas.new_zeros([1])]) + if discard_penultimate_step: + sigmas = torch.cat([sigmas[:-2], sigmas[-1:]]) + + # sampling + x0 = solver_fn(noise, + model_fn, + sigmas, + seed=seed, + show_progress=show_progress, + **kwargs) + return (x0, intermediates) if return_intermediate is not None else x0 + + +def extract_into_tensor(a, t, x_shape): + b, *_ = t.shape + out = a.gather(-1, t) + return out.reshape(b, *((1, ) * (len(x_shape) - 1))) + + +def discretize_timesteps(t_max, t_min, steps, discretization): + """ + Implementation of timestep discretization methods. + """ + if discretization == 'leading': + steps = torch.arange(t_min, t_max + 1, + (t_max - t_min + 1) / steps).flip(0) + elif discretization == 'linspace': + steps = torch.linspace(t_max, t_min, steps) + elif discretization == 'trailing': + steps = torch.arange(t_max, t_min - 1, -((t_max - t_min + 1) / steps)) + else: + raise NotImplementedError( + f'{discretization} discretization not implemented') + return steps.clamp_(t_min, t_max) diff --git a/scepter/modules/model/network/diffusion/schedules.py b/scepter/modules/model/network/diffusion/schedules.py new file mode 100644 index 0000000..72a1fb3 --- /dev/null +++ b/scepter/modules/model/network/diffusion/schedules.py @@ -0,0 +1,181 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +""" +Noise schedules of denoising diffusion probabilistic models. + +We consider a variance preserving (VP) process, and we use the standard deviation +sigma_t of the noise added to the signal at time t to represent the noise schedule. The +corresponding diffusion process is: + +q(x_t | x_0) = N(x_t | alpha_t x_0, sigma_t^2 I), + +where alpha_t^2 = 1 - sigma_t^2. +""" +import math + +import torch + +__all__ = [ + 'betas_to_sigmas', 'sigmas_to_betas', 'logsnrs_to_sigmas', + 'sigmas_to_logsnrs', 'linear_schedule', 'quadratic_schedule', + 'scaled_linear_schedule', 'cosine_schedule', 'sigmoid_schedule', + 'karras_schedule', 'exponential_schedule', 'polyexponential_schedule', + 'vp_schedule', 'logsnr_cosine_schedule', 'logsnr_cosine_shifted_schedule', + 'logsnr_cosine_interp_schedule', 'noise_schedule' +] + + +def betas_to_sigmas(betas): + return torch.sqrt(1 - torch.cumprod(1 - betas, dim=0)) + + +def sigmas_to_betas(sigmas): + square_alphas = 1 - sigmas**2 + betas = 1 - torch.cat( + [square_alphas[:1], square_alphas[1:] / square_alphas[:-1]]) + return betas + + +def logsnrs_to_sigmas(logsnrs): + return torch.sqrt(torch.sigmoid(-logsnrs)) + + +def sigmas_to_logsnrs(sigmas): + square_sigmas = sigmas**2 + return torch.log(square_sigmas / (1 - square_sigmas)) + + +def linear_schedule(n, beta_min=0.0001, beta_max=0.02): + betas = torch.linspace(beta_min, beta_max, n, dtype=torch.float32) + return betas_to_sigmas(betas) + + +def scaled_linear_schedule(n, beta_min=0.00085, beta_max=0.012): + betas = torch.linspace(beta_min**0.5, + beta_max**0.5, + n, + dtype=torch.float32)**2 + return betas_to_sigmas(betas) + + +def quadratic_schedule(n=1000, init_beta=0.00085, last_beta=0.012): + betas = torch.linspace(init_beta**0.5, + last_beta**0.5, + n, + dtype=torch.float32)**2 + return betas_to_sigmas(betas) + + +def cosine_schedule(n, cosine_s=0.008): + ramp = torch.linspace(0, 1, n + 1) + square_alphas = torch.cos( + (ramp + cosine_s) / (1 + cosine_s) * torch.pi / 2)**2 + betas = (1 - square_alphas[1:] / square_alphas[:-1]).clamp(max=0.999) + return betas_to_sigmas(betas) + + +def sigmoid_schedule(n, beta_min=0.0001, beta_max=0.02): + betas = torch.sigmoid(torch.linspace(-6, 6, + n)) * (beta_max - beta_min) + beta_min + return betas_to_sigmas(betas) + + +def karras_schedule(n, sigma_min=0.002, sigma_max=80.0, rho=7.0): + ramp = torch.linspace(1, 0, n) + min_inv_rho = sigma_min**(1 / rho) + max_inv_rho = sigma_max**(1 / rho) + sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho))**rho + sigmas = torch.sqrt(sigmas**2 / (1 + sigmas**2)) # VE -> VP + return sigmas + + +def exponential_schedule(n, sigma_min=0.002, sigma_max=80.0): + sigmas = torch.linspace(math.log(sigma_min), math.log(sigma_max), n).exp() + sigmas = torch.sqrt(sigmas**2 / (1 + sigmas**2)) # VE -> VP + return sigmas + + +def polyexponential_schedule(n, sigma_min=0.002, sigma_max=80.0): + ramp = torch.linspace(0, 1, n) + sigmas = torch.exp(ramp * (math.log(sigma_max) - math.log(sigma_min)) + + math.log(sigma_min)) + sigmas = torch.sqrt(sigmas**2 / (1 + sigmas**2)) # VE -> VP + return sigmas + + +def vp_schedule(n, beta_d=19.9, beta_min=0.1, eps_s=1e-3): + t = torch.linspace(eps_s, 1, n) + sigmas = torch.sqrt(torch.exp(beta_d * t**2 / 2 + beta_min * t) - 1) + sigmas = torch.sqrt(sigmas**2 / (1 + sigmas**2)) # VE -> VP + return sigmas + + +def _logsnr_cosine(n, logsnr_min=-15, logsnr_max=15): + t_min = math.atan(math.exp(-0.5 * logsnr_min)) + t_max = math.atan(math.exp(-0.5 * logsnr_max)) + t = torch.linspace(1, 0, n) + logsnrs = -2 * torch.log(torch.tan(t_min + t * (t_max - t_min))) + return logsnrs + + +def _logsnr_cosine_shifted(n, logsnr_min=-15, logsnr_max=15, scale=2): + logsnrs = _logsnr_cosine(n, logsnr_min, logsnr_max) + logsnrs += 2 * math.log(1 / scale) + return logsnrs + + +def _logsnr_cosine_interp(n, + logsnr_min=-15, + logsnr_max=15, + scale_min=2, + scale_max=4): + t = torch.linspace(1, 0, n) + logsnrs_min = _logsnr_cosine_shifted(n, logsnr_min, logsnr_max, scale_min) + logsnrs_max = _logsnr_cosine_shifted(n, logsnr_min, logsnr_max, scale_max) + logsnrs = t * logsnrs_min + (1 - t) * logsnrs_max + return logsnrs + + +def logsnr_cosine_schedule(n, logsnr_min=-15, logsnr_max=15): + return logsnrs_to_sigmas(_logsnr_cosine(n, logsnr_min, logsnr_max)) + + +def logsnr_cosine_shifted_schedule(n, logsnr_min=-15, logsnr_max=15, scale=2): + return logsnrs_to_sigmas( + _logsnr_cosine_shifted(n, logsnr_min, logsnr_max, scale)) + + +def logsnr_cosine_interp_schedule(n, + logsnr_min=-15, + logsnr_max=15, + scale_min=2, + scale_max=4): + return logsnrs_to_sigmas( + _logsnr_cosine_interp(n, logsnr_min, logsnr_max, scale_min, scale_max)) + + +def noise_schedule(schedule='logsnr_cosine_interp', + n=1000, + zero_terminal_snr=False, + **kwargs): + # compute sigmas + sigmas = { + 'linear': linear_schedule, + 'scaled_linear': scaled_linear_schedule, + 'quadratic': quadratic_schedule, + 'cosine': cosine_schedule, + 'sigmoid': sigmoid_schedule, + 'karras': karras_schedule, + 'exponential': exponential_schedule, + 'polyexponential': polyexponential_schedule, + 'vp': vp_schedule, + 'logsnr_cosine': logsnr_cosine_schedule, + 'logsnr_cosine_shifted': logsnr_cosine_shifted_schedule, + 'logsnr_cosine_interp': logsnr_cosine_interp_schedule + }[schedule](n, **kwargs) + + # post-processing + if zero_terminal_snr and sigmas.max() != 1.0: + scale = (1.0 - sigmas.min()) / (sigmas.max() - sigmas.min()) + sigmas = sigmas.min() + scale * (sigmas - sigmas.min()) + return sigmas diff --git a/scepter/modules/model/network/diffusion/solvers.py b/scepter/modules/model/network/diffusion/solvers.py new file mode 100644 index 0000000..a8fcb7c --- /dev/null +++ b/scepter/modules/model/network/diffusion/solvers.py @@ -0,0 +1,611 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +""" +ODE/SDE solver for denoising diffusion models under either variation preserving (VP) or +variation exploding (VE) settings. Under the VE setting, the diffusion process is: + +q(x_t | x_0) = N(x_t | x_0, sigma_t^2 I), + +where 0 <= sigma_t <= inf; while under the VP setting, the diffusion process is: + +q(x_t | x_0) = N(x_t | alpha_t x_0, sigma_t^2 I), + +where 0 <= sigma_t <= 1 and alpha_t^2 = 1 - sigma_t^2. +""" +import torch +from tqdm.auto import trange + +__all__ = [ + 'sample_euler', 'sample_euler_ancestral', 'sample_heun', 'sample_dpm_2', + 'sample_dpm_2_ancestral', 'sample_dpmpp_2s_ancestral', 'sample_dpmpp_sde', + 'sample_dpmpp_2m', 'sample_dpmpp_2m_sde', 'sample_ddim' +] + +# -------------------- variation exploding (VE) solver --------------------# + + +def get_ancestral_step(sigma_from, sigma_to, eta=1.): + """ + Calculates the noise level (sigma_down) to step down to and the amount + of noise to add (sigma_up) when doing an ancestral sampling step. + """ + if not eta: + return sigma_to, 0. + sigma_up = min( + sigma_to, + eta * (sigma_to**2 * + (sigma_from**2 - sigma_to**2) / sigma_from**2)**0.5) + sigma_down = (sigma_to**2 - sigma_up**2)**0.5 + return sigma_down, sigma_up + + +def get_scalings(sigma): + c_out = -sigma + c_in = 1 / (sigma**2 + 1.**2)**0.5 + return c_out, c_in + + +@torch.no_grad() +def sample_euler(noise, + model, + sigmas, + s_churn=0., + s_tmin=0., + s_tmax=float('inf'), + s_noise=1., + seed=None, + show_progress=True): + """ + Implements Algorithm 2 (Euler steps) from Karras et al. (2022). + """ + x = noise * sigmas[0] + for i in trange(len(sigmas) - 1, disable=not show_progress): + gamma = 0. + if s_tmin <= sigmas[i] <= s_tmax and sigmas[i] < float('inf'): + gamma = min(s_churn / (len(sigmas) - 1), 2**0.5 - 1) + eps = torch.randn_like(x) * s_noise + sigma_hat = sigmas[i] * (gamma + 1) + if gamma > 0: + x = x + eps * (sigma_hat**2 - sigmas[i]**2)**0.5 + # Euler method + if sigmas[i] == float('inf'): + denoised = model(noise, sigma_hat) + x = denoised + sigmas[i + 1] * (gamma + 1) * noise + else: + _, c_in = get_scalings(sigma_hat) + denoised = model(x * c_in, sigma_hat) + d = (x - denoised) / sigma_hat + dt = sigmas[i + 1] - sigma_hat + x = x + d * dt + return x + + +@torch.no_grad() +def sample_euler_ancestral(noise, + model, + sigmas, + eta=1., + s_noise=1., + seed=None, + show_progress=True): + """ + Ancestral sampling with Euler method steps. + """ + x = noise * sigmas[0] + for i in trange(len(sigmas) - 1, disable=not show_progress): + sigma_down, sigma_up = get_ancestral_step(sigmas[i], + sigmas[i + 1], + eta=eta) + # Euler method + if sigmas[i] == float('inf'): + denoised = model(noise, sigmas[i]) + x = denoised + sigmas[i + 1] * noise + else: + _, c_in = get_scalings(sigmas[i]) + denoised = model(x * c_in, sigmas[i]) + d = (x - denoised) / sigmas[i] + dt = sigma_down - sigmas[i] + x = x + d * dt + if sigmas[i + 1] > 0: + x = x + torch.randn_like(x) * s_noise * sigma_up + return x + + +@torch.no_grad() +def sample_heun(noise, + model, + sigmas, + s_churn=0., + s_tmin=0., + s_tmax=float('inf'), + s_noise=1., + seed=None, + show_progress=True): + """ + Implements Algorithm 2 (Heun steps) from Karras et al. (2022). + """ + x = noise * sigmas[0] + for i in trange(len(sigmas) - 1, disable=not show_progress): + gamma = 0. + if s_tmin <= sigmas[i] <= s_tmax and sigmas[i] < float('inf'): + gamma = min(s_churn / (len(sigmas) - 1), 2**0.5 - 1) + eps = torch.randn_like(x) * s_noise + sigma_hat = sigmas[i] * (gamma + 1) + if gamma > 0: + x = x + eps * (sigma_hat**2 - sigmas[i]**2)**0.5 + if sigmas[i] == float('inf'): + # Euler method + denoised = model(noise, sigma_hat) + x = denoised + sigmas[i + 1] * (gamma + 1) * noise + else: + _, c_in = get_scalings(sigma_hat) + denoised = model(x * c_in, sigma_hat) + d = (x - denoised) / sigma_hat + dt = sigmas[i + 1] - sigma_hat + if sigmas[i + 1] == 0: + # Euler method + x = x + d * dt + else: + # Heun's method + x_2 = x + d * dt + _, c_in = get_scalings(sigmas[i + 1]) + denoised_2 = model(x_2 * c_in, sigmas[i + 1]) + d_2 = (x_2 - denoised_2) / sigmas[i + 1] + d_prime = (d + d_2) / 2 + x = x + d_prime * dt + return x + + +@torch.no_grad() +def sample_dpm_2(noise, + model, + sigmas, + s_churn=0., + s_tmin=0., + s_tmax=float('inf'), + s_noise=1., + seed=None, + show_progress=True): + """ + A sampler inspired by DPM-Solver-2 and Algorithm 2 from Karras et al. (2022). + """ + x = noise * sigmas[0] + for i in trange(len(sigmas) - 1, disable=not show_progress): + gamma = 0. + if s_tmin <= sigmas[i] <= s_tmax and sigmas[i] < float('inf'): + gamma = min(s_churn / (len(sigmas) - 1), 2**0.5 - 1) + eps = torch.randn_like(x) * s_noise + sigma_hat = sigmas[i] * (gamma + 1) + if gamma > 0: + x = x + eps * (sigma_hat**2 - sigmas[i]**2)**0.5 + if sigmas[i] == float('inf'): + # Euler method + denoised = model(noise, sigma_hat) + x = denoised + sigmas[i + 1] * (gamma + 1) * noise + else: + _, c_in = get_scalings(sigma_hat) + denoised = model(x * c_in, sigma_hat) + d = (x - denoised) / sigma_hat + if sigmas[i + 1] == 0: + # Euler method + dt = sigmas[i + 1] - sigma_hat + x = x + d * dt + else: + # DPM-Solver-2 + sigma_mid = sigma_hat.log().lerp(sigmas[i + 1].log(), + 0.5).exp() + dt_1 = sigma_mid - sigma_hat + dt_2 = sigmas[i + 1] - sigma_hat + x_2 = x + d * dt_1 + _, c_in = get_scalings(sigma_mid) + denoised_2 = model(x_2 * c_in, sigma_mid) + d_2 = (x_2 - denoised_2) / sigma_mid + x = x + d_2 * dt_2 + return x + + +@torch.no_grad() +def sample_dpm_2_ancestral(noise, + model, + sigmas, + eta=1., + s_noise=1., + seed=None, + show_progress=True): + """ + Ancestral sampling with DPM-Solver second-order steps. + """ + x = noise * sigmas[0] + for i in trange(len(sigmas) - 1, disable=not show_progress): + sigma_down, sigma_up = get_ancestral_step(sigmas[i], + sigmas[i + 1], + eta=eta) + if sigmas[i] == float('inf'): + # Euler method + denoised = model(noise, sigmas[i]) + x = denoised + sigmas[i + 1] * noise + else: + _, c_in = get_scalings(sigmas[i]) + denoised = model(x * c_in, sigmas[i]) + d = (x - denoised) / sigmas[i] + if sigma_down == 0: + # Euler method + dt = sigma_down - sigmas[i] + x = x + d * dt + else: + # DPM-Solver-2 + sigma_mid = sigmas[i].log().lerp(sigma_down.log(), 0.5).exp() + dt_1 = sigma_mid - sigmas[i] + dt_2 = sigma_down - sigmas[i] + x_2 = x + d * dt_1 + _, c_in = get_scalings(sigma_mid) + denoised_2 = model(x_2 * c_in, sigma_mid) + d_2 = (x_2 - denoised_2) / sigma_mid + x = x + d_2 * dt_2 + x = x + torch.randn_like(x) * s_noise * sigma_up + return x + + +@torch.no_grad() +def sample_dpmpp_2s_ancestral(noise, + model, + sigmas, + eta=1., + s_noise=1., + seed=None, + show_progress=True): + """ + Ancestral sampling with DPM-Solver++ (2S) second-order steps. + """ + def t_to_sigma(t): + return t.neg().exp() + + def sigma_to_t(sigma): + return sigma.log().neg() + + # x = noise * sigmas[0] + x = noise * torch.sqrt(1.0 + sigmas[0]**2.0) + for i in trange(len(sigmas) - 1, disable=not show_progress): + sigma_down, sigma_up = get_ancestral_step(sigmas[i], + sigmas[i + 1], + eta=eta) + if sigmas[i] == float('inf'): + # Euler method + denoised = model(noise, sigmas[i]) + x = denoised + sigma_down * noise + else: + _, c_in = get_scalings(sigmas[i]) + denoised = model(x * c_in, sigmas[i]) + if sigma_down == 0: + # Euler method + d = (x - denoised) / sigmas[i] + dt = sigma_down - sigmas[i] + x = x + d * dt + else: + # DPM-Solver++(2S) + t, t_next = sigma_to_t(sigmas[i]), sigma_to_t(sigma_down) + r = 1 / 2 + h = t_next - t + s = t + r * h + x_2 = (t_to_sigma(s) / + t_to_sigma(t)) * x - (-h * r).expm1() * denoised + _, c_in = get_scalings(t_to_sigma(s)) + denoised_2 = model(x_2 * c_in, t_to_sigma(s)) + x = (t_to_sigma(t_next) / + t_to_sigma(t)) * x - (-h).expm1() * denoised_2 + # Noise addition + if sigmas[i + 1] > 0: + x = x + torch.randn_like(x) * s_noise * sigma_up + return x + + +class BatchedBrownianTree: + """ + A wrapper around torchsde.BrownianTree that enables batches of entropy. + """ + def __init__(self, x, t0, t1, seed=None, **kwargs): + import torchsde + t0, t1, self.sign = self.sort(t0, t1) + w0 = kwargs.get('w0', torch.zeros_like(x)) + if seed is None: + seed = torch.randint(0, 2**63 - 1, []).item() + self.batched = True + try: + assert len(seed) == x.shape[0] + w0 = w0[0] + except TypeError: + seed = [seed] + self.batched = False + self.trees = [ + torchsde.BrownianTree(t0, w0, t1, entropy=s, **kwargs) + for s in seed + ] + + @staticmethod + def sort(a, b): + return (a, b, 1) if a < b else (b, a, -1) + + def __call__(self, t0, t1): + t0, t1, sign = self.sort(t0, t1) + w = torch.stack([tree(t0, t1) + for tree in self.trees]) * (self.sign * sign) + return w if self.batched else w[0] + + +class BrownianTreeNoiseSampler: + """ + A noise sampler backed by a torchsde.BrownianTree. + + Args: + x (Tensor): The tensor whose shape, device and dtype to use to generate + random samples. + sigma_min (float): The low end of the valid interval. + sigma_max (float): The high end of the valid interval. + seed (int or List[int]): The random seed. If a list of seeds is + supplied instead of a single integer, then the noise sampler will + use one BrownianTree per batch item, each with its own seed. + transform (callable): A function that maps sigma to the sampler's + internal timestep. + """ + def __init__(self, + x, + sigma_min, + sigma_max, + seed=None, + transform=lambda x: x): + self.transform = transform + t0 = self.transform(torch.as_tensor(sigma_min)) + t1 = self.transform(torch.as_tensor(sigma_max)) + self.tree = BatchedBrownianTree(x, t0, t1, seed) + + def __call__(self, sigma, sigma_next): + t0 = self.transform(torch.as_tensor(sigma)) + t1 = self.transform(torch.as_tensor(sigma_next)) + return self.tree(t0, t1) / (t1 - t0).abs().sqrt() + + +@torch.no_grad() +def sample_dpmpp_sde(noise, + model, + sigmas, + eta=1., + s_noise=1., + r=1 / 2, + seed=None, + show_progress=True): + """ + DPM-Solver++ (stochastic). + """ + def t_to_sigma(t): + return t.neg().exp() + + def sigma_to_t(sigma): + return sigma.log().neg() + + x = noise * sigmas[0] + sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas[ + sigmas < float('inf')].max() + noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed) + for i in trange(len(sigmas) - 1, disable=not show_progress): + if sigmas[i] == float('inf'): + # Euler method + denoised = model(noise, sigmas[i]) + x = denoised + sigmas[i + 1] * noise + else: + _, c_in = get_scalings(sigmas[i]) + denoised = model(x * c_in, sigmas[i]) + if sigmas[i + 1] == 0: + # Euler method + d = (x - denoised) / sigmas[i] + dt = sigmas[i + 1] - sigmas[i] + x = x + d * dt + else: + # DPM-Solver++ + t, t_next = sigma_to_t(sigmas[i]), sigma_to_t(sigmas[i + 1]) + h = t_next - t + s = t + h * r + fac = 1 / (2 * r) + + # Step 1 + sd, su = get_ancestral_step(t_to_sigma(t), t_to_sigma(s), eta) + s_ = sigma_to_t(sd) + x_2 = (t_to_sigma(s_) / + t_to_sigma(t)) * x - (t - s_).expm1() * denoised + x_2 = x_2 + noise_sampler(t_to_sigma(t), + t_to_sigma(s)) * s_noise * su + _, c_in = get_scalings(t_to_sigma(s)) + denoised_2 = model(x_2 * c_in, t_to_sigma(s)) + + # Step 2 + sd, su = get_ancestral_step(t_to_sigma(t), t_to_sigma(t_next), + eta) + t_next_ = sigma_to_t(sd) + denoised_d = (1 - fac) * denoised + fac * denoised_2 + x = (t_to_sigma(t_next_) / t_to_sigma(t)) * x - \ + (t - t_next_).expm1() * denoised_d + x = x + noise_sampler(t_to_sigma(t), + t_to_sigma(t_next)) * s_noise * su + return x + + +@torch.no_grad() +def sample_dpmpp_2m(noise, model, sigmas, seed=None, show_progress=True): + """ + DPM-Solver++ (2M). + """ + def t_to_sigma(t): + return t.neg().exp() + + def sigma_to_t(sigma): + return sigma.log().neg() + + x = noise * sigmas[0] + old_denoised = None + for i in trange(len(sigmas) - 1, disable=not show_progress): + if sigmas[i] == float('inf'): + # Euler method + denoised = model(noise, sigmas[i]) + x = denoised + sigmas[i + 1] * noise + else: + _, c_in = get_scalings(sigmas[i]) + denoised = model(x * c_in, sigmas[i]) + t, t_next = sigma_to_t(sigmas[i]), sigma_to_t(sigmas[i + 1]) + h = t_next - t + if (old_denoised is None or sigmas[i - 1] == float('inf') + or sigmas[i + 1] == 0): + x = (t_to_sigma(t_next) / + t_to_sigma(t)) * x - (-h).expm1() * denoised + else: + h_last = t - sigma_to_t(sigmas[i - 1]) + r = h_last / h + denoised_d = (1 + 1 / + (2 * r)) * denoised - (1 / + (2 * r)) * old_denoised + x = (t_to_sigma(t_next) / + t_to_sigma(t)) * x - (-h).expm1() * denoised_d + old_denoised = denoised + return x + + +@torch.no_grad() +def sample_dpmpp_2m_sde(noise, + model, + sigmas, + eta=1., + s_noise=1., + solver_type='midpoint', + seed=None, + show_progress=True): + """ + DPM-Solver++ (2M) SDE. + """ + assert solver_type in {'heun', 'midpoint'} + + x = noise * sigmas[0] + sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas[ + sigmas < float('inf')].max() + noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed) + old_denoised = None + h_last = None + + for i in trange(len(sigmas) - 1, disable=not show_progress): + if sigmas[i] == float('inf'): + # Euler method + denoised = model(noise, sigmas[i]) + x = denoised + sigmas[i + 1] * noise + else: + _, c_in = get_scalings(sigmas[i]) + denoised = model(x * c_in, sigmas[i]) + if sigmas[i + 1] == 0: + # Denoising step + x = denoised + else: + # DPM-Solver++(2M) SDE + t, s = -sigmas[i].log(), -sigmas[i + 1].log() + h = s - t + eta_h = eta * h + + x = sigmas[i + 1] / sigmas[i] * (-eta_h).exp() * x + \ + (-h - eta_h).expm1().neg() * denoised + + if old_denoised is not None: + r = h_last / h + if solver_type == 'heun': + x = x + ((-h - eta_h).expm1().neg() / (-h - eta_h) + 1) * \ + (1 / r) * (denoised - old_denoised) + elif solver_type == 'midpoint': + x = x + 0.5 * (-h - eta_h).expm1().neg() * \ + (1 / r) * (denoised - old_denoised) + + x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[ + i + 1] * (-2 * eta_h).expm1().neg().sqrt() * s_noise + + old_denoised = denoised + h_last = h + return x + + +# -------------------- variation preserving (VP) solver --------------------# +@torch.no_grad() +def sample_ddim(noise, model, sigmas, eta=0., seed=None, show_progress=True): + """ + DDIM solver steps. + """ + x = noise + sigmas_vp = (sigmas**2 / (1 + sigmas**2))**0.5 + sigmas_vp[sigmas == float('inf')] = 1. + for i in trange(len(sigmas) - 1, disable=not show_progress): + denoised = model(x, sigmas[i]) + noise_factor = eta * (sigmas_vp[i + 1]**2 / sigmas_vp[i]**2 * + (1 - (1 - sigmas_vp[i]**2) / + (1 - sigmas_vp[i + 1]**2))) + d = (x - (1 - sigmas_vp[i]**2)**0.5 * denoised) / sigmas_vp[i] + x = (1 - sigmas_vp[i + 1] ** 2) ** 0.5 * denoised + \ + (sigmas_vp[i + 1] ** 2 - noise_factor ** 2) ** 0.5 * d + if sigmas_vp[i + 1] > 0: + x += noise_factor * torch.randn_like(x) + return x + + +@torch.no_grad() +def sample_img2img_euler(noise, + model, + sigmas, + s_churn=0., + s_tmin=0., + s_tmax=float('inf'), + s_noise=1., + seed=None, + show_progress=True): + """ + Implements Algorithm 2 (Euler steps) from Karras et al. (2022). + """ + x = noise + for i in trange(len(sigmas) - 1, disable=not show_progress): + gamma = 0. + if s_tmin <= sigmas[i] <= s_tmax and sigmas[i] < float('inf'): + gamma = min(s_churn / (len(sigmas) - 1), 2**0.5 - 1) + eps = torch.randn_like(x) * s_noise + sigma_hat = sigmas[i] * (gamma + 1) + if gamma > 0: + x = x + eps * (sigma_hat**2 - sigmas[i]**2)**0.5 + # Euler method + if sigmas[i] == float('inf'): + denoised = model(noise, sigma_hat) + x = denoised + sigmas[i + 1] * (gamma + 1) * noise + else: + denoised = model(x, sigma_hat) + d = (x - denoised) / sigma_hat + dt = sigmas[i + 1] - sigma_hat + x = x + d * dt + return x + + +@torch.no_grad() +def sample_img2img_euler_ancestral(noise, + model, + sigmas, + eta=1., + s_noise=1., + seed=None, + show_progress=True): + """ + Ancestral sampling with Euler method steps. + """ + x = noise + for i in trange(len(sigmas) - 1, disable=not show_progress): + sigma_down, sigma_up = get_ancestral_step(sigmas[i], + sigmas[i + 1], + eta=eta) + # Euler method + if sigmas[i] == float('inf'): + denoised = model(noise, sigmas[i]) + x = denoised + sigmas[i + 1] * noise + else: + denoised = model(x, sigmas[i]) + d = (x - denoised) / sigmas[i] + dt = sigma_down - sigmas[i] + x = x + d * dt + if sigmas[i + 1] > 0: + x = x + torch.randn_like(x) * s_noise * sigma_up + return x diff --git a/scepter/modules/model/network/ldm/__init__.py b/scepter/modules/model/network/ldm/__init__.py new file mode 100644 index 0000000..30f2253 --- /dev/null +++ b/scepter/modules/model/network/ldm/__init__.py @@ -0,0 +1,4 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.network.ldm.ldm import LatentDiffusion +from scepter.modules.model.network.ldm.ldm_xl import LatentDiffusionXL diff --git a/scepter/modules/model/network/ldm/ldm.py b/scepter/modules/model/network/ldm/ldm.py new file mode 100644 index 0000000..5e0ffff --- /dev/null +++ b/scepter/modules/model/network/ldm/ldm.py @@ -0,0 +1,432 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import numbers +import random +from collections import OrderedDict + +import torch + +from scepter.modules.model.network.diffusion.diffusion import GaussianDiffusion +from scepter.modules.model.network.diffusion.schedules import noise_schedule +from scepter.modules.model.network.train_module import TrainModule +from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, LOSSES, + MODELS, TOKENIZERS) +from scepter.modules.model.utils.basic_utils import count_params, default +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we +from scepter.modules.utils.file_system import FS + + +def disabled_train(self, mode=True): + """Overwrite model.train with this function to make sure train/eval mode + does not change anymore.""" + return self + + +@MODELS.register_class() +class LatentDiffusion(TrainModule): + para_dict = { + 'PARAMETERIZATION': { + 'value': + 'v', + 'description': + "The prediction type, you can choose from 'eps' and 'x0' and 'v'", + }, + 'TIMESTEPS': { + 'value': 1000, + 'description': 'The schedule steps for diffusion.', + }, + 'SCHEDULE_ARGS': {}, + 'MIN_SNR_GAMMA': { + 'value': None, + 'description': 'The minimum snr gamma, default is None.', + }, + 'ZERO_TERMINAL_SNR': { + 'value': False, + 'description': 'Whether zero terminal snr, default is False.', + }, + 'PRETRAINED_MODEL': { + 'value': None, + 'description': "Whole model's pretrained model path.", + }, + 'IGNORE_KEYS': { + 'value': [], + 'description': 'The ignore keys for pretrain model loaded.', + }, + 'SCALE_FACTOR': { + 'value': 0.18215, + 'description': 'The vae embeding scale.', + }, + 'SIZE_FACTOR': { + 'value': 8, + 'description': 'The vae size factor.', + }, + 'DEFAULT_N_PROMPT': { + 'value': '', + 'description': 'The default negtive prompt.', + }, + 'TRAIN_N_PROMPT': { + 'value': '', + 'description': 'The negtive prompt used in train phase.', + }, + 'P_ZERO': { + 'value': 0.0, + 'description': 'The prob for zero or negtive prompt.', + }, + 'USE_EMA': { + 'value': True, + 'description': 'Use Ema or not. Default True', + }, + 'DIFFUSION_MODEL': {}, + 'DIFFUSION_MODEL_EMA': {}, + 'FIRST_STAGE_MODEL': {}, + 'COND_STAGE_MODEL': {}, + 'TOKENIZER': {} + } + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.init_params() + self.construct_network() + + def init_params(self): + self.parameterization = self.cfg.get('PARAMETERIZATION', 'eps') + assert self.parameterization in [ + 'eps', 'x0', 'v' + ], 'currently only supporting "eps" and "x0" and "v"' + self.num_timesteps = self.cfg.get('TIMESTEPS', 1000) + + self.schedule_args = { + k.lower(): v + for k, v in self.cfg.get('SCHEDULE_ARGS', { + 'NAME': 'logsnr_cosine_interp', + 'SCALE_MIN': 2.0, + 'SCALE_MAX': 4.0 + }).items() + } + + self.min_snr_gamma = self.cfg.get('MIN_SNR_GAMMA', None) + + self.zero_terminal_snr = self.cfg.get('ZERO_TERMINAL_SNR', False) + if self.zero_terminal_snr: + assert self.parameterization == 'v', 'Now zero_terminal_snr only support v-prediction mode.' + + self.sigmas = noise_schedule(schedule=self.schedule_args.pop('name'), + n=self.num_timesteps, + zero_terminal_snr=self.zero_terminal_snr, + **self.schedule_args) + + self.diffusion = GaussianDiffusion( + sigmas=self.sigmas, prediction_type=self.parameterization) + + self.pretrained_model = self.cfg.get('PRETRAINED_MODEL', None) + self.ignore_keys = self.cfg.get('IGNORE_KEYS', []) + + self.model_config = self.cfg.DIFFUSION_MODEL + self.first_stage_config = self.cfg.FIRST_STAGE_MODEL + self.cond_stage_config = self.cfg.COND_STAGE_MODEL + self.tokenizer_config = self.cfg.get('TOKENIZER', None) + self.loss_config = self.cfg.get('LOSS', None) + + self.scale_factor = self.cfg.get('SCALE_FACTOR', 0.18215) + self.size_factor = self.cfg.get('SIZE_FACTOR', 8) + self.default_n_prompt = self.cfg.get('DEFAULT_N_PROMPT', '') + self.default_n_prompt = '' if self.default_n_prompt is None else self.default_n_prompt + self.p_zero = self.cfg.get('P_ZERO', 0.0) + self.train_n_prompt = self.cfg.get('TRAIN_N_PROMPT', '') + if self.default_n_prompt is None: + self.default_n_prompt = '' + if self.train_n_prompt is None: + self.train_n_prompt = '' + self.use_ema = self.cfg.get('USE_EMA', True) + self.model_ema_config = self.cfg.get('DIFFUSION_MODEL_EMA', None) + + def construct_network(self): + self.model = BACKBONES.build(self.model_config, logger=self.logger) + self.logger.info('all parameters:{}'.format(count_params(self.model))) + if self.use_ema and self.model_ema_config: + self.model_ema = BACKBONES.build(self.model_ema_config, + logger=self.logger) + self.model_ema = self.model_ema.eval() + for param in self.model_ema.parameters(): + param.requires_grad = False + if self.loss_config: + self.loss = LOSSES.build(self.loss_config, logger=self.logger) + if self.tokenizer_config is not None: + self.tokenizer = TOKENIZERS.build(self.tokenizer_config, + logger=self.logger) + + self.first_stage_model = MODELS.build(self.first_stage_config, + logger=self.logger) + self.first_stage_model = self.first_stage_model.eval() + self.first_stage_model.train = disabled_train + for param in self.first_stage_model.parameters(): + param.requires_grad = False + if self.tokenizer_config is not None: + self.cond_stage_config.KWARGS = { + 'vocab_size': self.tokenizer.vocab_size + } + if self.cond_stage_config == '__is_unconditional__': + print( + f'Training {self.__class__.__name__} as an unconditional model.' + ) + self.cond_stage_model = None + else: + model = EMBEDDERS.build(self.cond_stage_config, logger=self.logger) + self.cond_stage_model = model.eval().requires_grad_(False) + self.cond_stage_model.train = disabled_train + + def load_pretrained_model(self, pretrained_model): + if pretrained_model is not None: + with FS.get_from(pretrained_model, + wait_finish=True) as local_model: + self.init_from_ckpt(local_model, ignore_keys=self.ignore_keys) + + def init_from_ckpt(self, path, ignore_keys=list()): + if path.endswith('safetensors'): + from safetensors.torch import load_file as load_safetensors + sd = load_safetensors(path) + else: + sd = torch.load(path, map_location='cpu') + new_sd = OrderedDict() + for k, v in sd.items(): + ignored = False + for ik in ignore_keys: + if ik in k: + if we.rank == 0: + self.logger.info( + 'Ignore key {} from state_dict.'.format(k)) + ignored = True + break + if not ignored: + if k.startswith('model.diffusion_model.'): + k = k.replace('model.diffusion_model.', 'model.') + k = k.replace('post_quant_conv', + 'conv2') if 'post_quant_conv' in k else k + k = k.replace('quant_conv', + 'conv1') if 'quant_conv' in k else k + new_sd[k] = v + + missing, unexpected = self.load_state_dict(new_sd, strict=False) + if we.rank == 0: + self.logger.info( + f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys' + ) + if len(missing) > 0: + self.logger.info(f'Missing Keys:\n {missing}') + if len(unexpected) > 0: + self.logger.info(f'\nUnexpected Keys:\n {unexpected}') + + def encode_condition(self, input, method='encode_text'): + if hasattr(self.cond_stage_model, method): + return getattr(self.cond_stage_model, + method)(input, tokenizer=self.tokenizer) + else: + return self.cond_stage_model(input) + + def forward_train(self, image=None, noise=None, prompt=None, **kwargs): + + x_start = self.encode_first_stage(image, **kwargs) + t = torch.randint(0, + self.num_timesteps, (x_start.shape[0], ), + device=x_start.device).long() + context = {} + if prompt and self.cond_stage_model: + zeros = (torch.rand(len(prompt)) < self.p_zero).numpy().tolist() + prompt = [ + self.train_n_prompt if zeros[idx] else p + for idx, p in enumerate(prompt) + ] + self.register_probe({'after_prompt': prompt}) + with torch.autocast(device_type='cuda', enabled=False): + context = self.encode_condition( + self.tokenizer(prompt).to(we.device_id)) + if self.min_snr_gamma is not None: + alphas = self.diffusion.alphas.to(we.device_id)[t] + sigmas = self.diffusion.sigmas.pow(2).to(we.device_id)[t] + snrs = (alphas / sigmas).clamp(min=1e-20) + min_snrs = snrs.clamp(max=self.min_snr_gamma) + weights = min_snrs / snrs + else: + weights = 1 + self.register_probe({'snrs_weights': weights}) + loss = self.diffusion.loss(x0=x_start, + t=t, + model=self.model, + model_kwargs={'cond': context}, + noise=noise) + loss = loss * weights + loss = loss.mean() + ret = {'loss': loss, 'probe_data': {'prompt': prompt}} + return ret + + def noise_sample(self, batch_size, h, w, g): + noise = torch.empty(batch_size, 4, h, w, + device=we.device_id).normal_(generator=g) + return noise + + def forward(self, **kwargs): + if self.training: + return self.forward_train(**kwargs) + else: + return self.forward_test(**kwargs) + + @torch.no_grad() + @torch.autocast('cuda', dtype=torch.float16) + def forward_test(self, + prompt=None, + n_prompt=None, + sampler='ddim', + sample_steps=50, + seed=2023, + guide_scale=7.5, + guide_rescale=0.5, + discretization='trailing', + run_train_n=True, + **kwargs): + g = torch.Generator(device=we.device_id) + seed = seed if seed >= 0 else random.randint(0, 2**32 - 1) + g.manual_seed(seed) + num_samples = len(prompt) + if 'dynamic_encode_text' in kwargs and kwargs.pop( + 'dynamic_encode_text'): + method = 'dynamic_encode_text' + else: + method = 'encode_text' + + n_prompt = default(n_prompt, [self.default_n_prompt] * len(prompt)) + assert isinstance(prompt, list) and \ + isinstance(n_prompt, list) and \ + len(prompt) == len(n_prompt) + # with torch.autocast(device_type="cuda", enabled=False): + context = self.encode_condition(self.tokenizer(prompt).to( + we.device_id), + method=method) + null_context = self.encode_condition(self.tokenizer(n_prompt).to( + we.device_id), + method=method) + + if 'index' in kwargs: + kwargs.pop('index') + image_size = None + if 'meta' in kwargs: + meta = kwargs.pop('meta') + if 'image_size' in meta: + h = int(meta['image_size'][0][0]) + w = int(meta['image_size'][1][0]) + image_size = [h, w] + if 'image_size' in kwargs: + image_size = kwargs.pop('image_size') + if image_size is None or isinstance(image_size, numbers.Number): + image_size = [1024, 1024] + height, width = image_size + noise = self.noise_sample(num_samples, height // self.size_factor, + width // self.size_factor, g) + # UNet use input n_prompt + samples = self.diffusion.sample(solver=sampler, + noise=noise, + model=self.model, + model_kwargs=[{ + 'cond': context + }, { + 'cond': null_context + }], + steps=sample_steps, + guide_scale=guide_scale, + guide_rescale=guide_rescale, + discretization=discretization, + show_progress=True, + seed=seed, + condition_fn=None, + clamp=None, + percentile=None, + t_max=None, + t_min=None, + discard_penultimate_step=None, + return_intermediate=None, + **kwargs) + x_samples = self.decode_first_stage(samples).float() + x_samples = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0) + # UNet use train n_prompt + if not self.default_n_prompt == self.train_n_prompt and run_train_n: + train_n_prompt = [self.train_n_prompt] * len(prompt) + null_train_context = self.encode_condition( + self.tokenizer(train_n_prompt).to(we.device_id), method=method) + + tn_samples = self.diffusion.sample(solver=sampler, + noise=noise, + model=self.model, + model_kwargs=[{ + 'cond': context + }, { + 'cond': + null_train_context + }], + steps=sample_steps, + guide_scale=guide_scale, + guide_rescale=guide_rescale, + discretization=discretization, + show_progress=we.rank == 0, + seed=seed, + condition_fn=None, + clamp=None, + percentile=None, + t_max=None, + t_min=None, + discard_penultimate_step=None, + return_intermediate=None, + **kwargs) + + t_x_samples = self.decode_first_stage(tn_samples).float() + + t_x_samples = torch.clamp((t_x_samples + 1.0) / 2.0, + min=0.0, + max=1.0) + else: + train_n_prompt = ['' for _ in prompt] + t_x_samples = [None for _ in prompt] + + outputs = list() + for p, np, tnp, img, t_img in zip(prompt, n_prompt, train_n_prompt, + x_samples, t_x_samples): + one_tup = {'prompt': p, 'n_prompt': np, 'image': img} + if t_img is not None: + one_tup['train_n_prompt'] = tnp + one_tup['train_n_image'] = t_img + outputs.append(one_tup) + + return outputs + + @torch.no_grad() + def log_images(self, image=None, prompt=None, n_prompt=None, **kwargs): + results = self.forward_test(prompt=prompt, n_prompt=n_prompt, **kwargs) + outputs = list() + for img, res in zip(image, results): + one_tup = { + 'orig': torch.clamp((img + 1.0) / 2.0, min=0.0, max=1.0), + 'recon': res['image'], + 'prompt': res['prompt'], + 'n_prompt': res['n_prompt'] + } + if 'train_n_prompt' in res: + one_tup['train_n_prompt'] = res['train_n_prompt'] + one_tup['train_n_image'] = res['train_n_image'] + outputs.append(one_tup) + return outputs + + @torch.no_grad() + def encode_first_stage(self, x, **kwargs): + z = self.first_stage_model.encode(x) + return self.scale_factor * z + + @torch.no_grad() + def decode_first_stage(self, z): + z = 1. / self.scale_factor * z + return self.first_stage_model.decode(z) + + @staticmethod + def get_config_template(): + return dict_to_yaml('MODEL', + __class__.__name__, + LatentDiffusion.para_dict, + set_name=True) diff --git a/scepter/modules/model/network/ldm/ldm_xl.py b/scepter/modules/model/network/ldm/ldm_xl.py new file mode 100644 index 0000000..a02edf2 --- /dev/null +++ b/scepter/modules/model/network/ldm/ldm_xl.py @@ -0,0 +1,507 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import numbers +import random +from collections import OrderedDict + +import torch +import torch.nn.functional as F + +from scepter.modules.model.network.ldm import LatentDiffusion +from scepter.modules.model.registry import BACKBONES, MODELS +from scepter.modules.model.utils.basic_utils import default +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we + + +@MODELS.register_class() +class LatentDiffusionXL(LatentDiffusion): + para_dict = { + 'LOAD_REFINER': { + 'value': False, + 'description': 'Whether load REFINER or Not.' + } + } + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.load_refiner = cfg.get('LOAD_REFINER', False) + self.latent_cache_data = {} + self.SURPPORT_RATIOS = { + '0.5': (704, 1408), + '0.52': (704, 1344), + '0.57': (768, 1344), + '0.6': (768, 1280), + '0.68': (832, 1216), + '0.72': (832, 1152), + '0.78': (896, 1152), + '0.82': (896, 1088), + '0.88': (960, 1088), + '0.94': (960, 1024), + '1.0': (1024, 1024), + '1.07': (1024, 960), + '1.13': (1088, 960), + '1.21': (1088, 896), + '1.29': (1152, 896), + '1.38': (1152, 832), + '1.46': (1216, 832), + '1.67': (1280, 768), + '1.75': (1344, 768), + '1.91': (1344, 704), + '2.0': (1408, 704), + '2.09': (1472, 704), + '2.4': (1536, 640), + '2.5': (1600, 640), + '2.89': (1664, 576), + '3.0': (1728, 576), + } + + def construct_network(self): + super().construct_network() + self.refiner_cfg = self.cfg.get('REFINER_MODEL', None) + self.refiner_cond_cfg = self.cfg.get('REFINER_COND_MODEL', None) + if self.refiner_cfg and self.load_refiner: + self.refiner_model = BACKBONES.build(self.refiner_cfg, + logger=self.logger) + self.refiner_cond_model = BACKBONES.build(self.refiner_cond_cfg, + logger=self.logger) + else: + self.refiner_model = None + self.refiner_cond_model = None + self.input_keys = self.get_unique_embedder_keys_from_conditioner( + self.cond_stage_model) + if self.refiner_cond_model: + self.input_refiner_keys = self.get_unique_embedder_keys_from_conditioner( + self.refiner_cond_model) + else: + self.input_refiner_keys = [] + + def init_from_ckpt(self, path, ignore_keys=list()): + if path.endswith('safetensors'): + from safetensors.torch import load_file as load_safetensors + sd = load_safetensors(path) + else: + sd = torch.load(path, map_location='cpu') + new_sd = OrderedDict() + for k, v in sd.items(): + ignored = False + for ik in ignore_keys: + if ik in k: + if we.rank == 0: + self.logger.info( + 'Ignore key {} from state_dict.'.format(k)) + ignored = True + break + if not ignored: + if k.startswith('model.diffusion_model.'): + k = k.replace('model.diffusion_model.', 'model.') + if k.startswith('conditioner.'): + k = k.replace('conditioner.', 'cond_stage_model.') + k = k.replace('post_quant_conv', + 'conv2') if 'post_quant_conv' in k else k + k = k.replace('quant_conv', + 'conv1') if 'quant_conv' in k else k + new_sd[k] = v + + missing, unexpected = self.load_state_dict(new_sd, strict=False) + if we.rank == 0: + self.logger.info( + f'Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys' + ) + if len(missing) > 0: + self.logger.info(f'Missing Keys:\n {missing}') + if len(unexpected) > 0: + self.logger.info(f'\nUnexpected Keys:\n {unexpected}') + + def get_unique_embedder_keys_from_conditioner(self, conditioner): + input_keys = [] + for x in conditioner.embedders: + input_keys.extend(x.input_keys) + return list(set(input_keys)) + + def get_batch(self, keys, value_dict, num_samples=1): + batch = {} + batch_uc = {} + N = num_samples + device = we.device_id + for key in keys: + 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)) + 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 forward_train(self, image=None, noise=None, prompt=None, **kwargs): + with torch.autocast('cuda', enabled=False): + x_start = self.encode_first_stage(image, **kwargs) + + t = torch.randint(0, + self.num_timesteps, (x_start.shape[0], ), + device=x_start.device).long() + + if prompt and self.cond_stage_model: + zeros = (torch.rand(len(prompt)) < self.p_zero).numpy().tolist() + prompt = [ + self.train_n_prompt if zeros[idx] else p + for idx, p in enumerate(prompt) + ] + self.register_probe({'after_prompt': prompt}) + batch = {'prompt': prompt} + for key in self.input_keys: + if key not in kwargs: + continue + batch[key] = kwargs[key].to(we.device_id) + context = getattr(self.cond_stage_model, 'encode')(batch) + + if self.min_snr_gamma is not None: + alphas = self.diffusion.alphas.to(we.device_id)[t] + sigmas = self.diffusion.sigmas.pow(2).to(we.device_id)[t] + snrs = (alphas / sigmas).clamp(min=1e-20) + min_snrs = snrs.clamp(max=self.min_snr_gamma) + weights = min_snrs / snrs + else: + weights = 1 + self.register_probe({'snrs_weights': weights}) + loss = self.diffusion.loss(x0=x_start, + t=t, + model=self.model, + model_kwargs={'cond': context}, + noise=noise) + loss = loss * weights + loss = loss.mean() + ret = {'loss': loss, 'probe_data': {'prompt': prompt}} + return ret + + def check_valid_inputs(self, kwargs): + batch_data = {} + all_keys = set(self.input_keys + self.input_refiner_keys) + for key in all_keys: + if key in kwargs: + batch_data[key] = kwargs.pop(key) + return batch_data + + @torch.no_grad() + def forward_test(self, + prompt=None, + n_prompt=None, + image=None, + sampler='ddim', + sample_steps=50, + seed=2023, + guide_scale=7.5, + guide_rescale=0.5, + discretization='trailing', + img_to_img_strength=0.0, + run_train_n=True, + refine_strength=0.0, + refine_sampler='ddim', + **kwargs): + g = torch.Generator(device=we.device_id) + seed = seed if seed >= 0 else random.randint(0, 2**32 - 1) + g.manual_seed(seed) + num_samples = len(prompt) + n_prompt = default(n_prompt, [self.default_n_prompt] * len(prompt)) + assert isinstance(prompt, list) and \ + isinstance(n_prompt, list) and \ + len(prompt) == len(n_prompt) + image_size = None + if 'meta' in kwargs: + meta = kwargs.pop('meta') + if 'image_size' in meta: + h = int(meta['image_size'][0][0]) + w = int(meta['image_size'][1][0]) + image_size = [h, w] + if 'image_size' in kwargs: + image_size = kwargs.pop('image_size') + if image_size is None or isinstance(image_size, numbers.Number): + image_size = [1024, 1024] + pre_batch = self.check_valid_inputs(kwargs) + if len(pre_batch) > 0: + batch = {'prompt': prompt} + batch.update(pre_batch) + batch_uc = {'prompt': n_prompt} + batch_uc.update(pre_batch) + else: + height, width = image_size + if image is None: + ori_width = width + ori_height = height + else: + ori_height, ori_width = image.shape[-2:] + + value_dict = { + 'original_size_as_tuple': [ori_height, ori_width], + 'target_size_as_tuple': [height, width], + 'prompt': prompt, + 'negative_prompt': n_prompt, + 'crop_coords_top_left': [0, 0] + } + if refine_strength > 0: + assert 'aesthetic_score' in kwargs and 'negative_aesthetic_score' in kwargs + value_dict['aesthetic_score'] = kwargs.pop('aesthetic_score') + value_dict['negative_aesthetic_score'] = kwargs.pop( + 'negative_aesthetic_score') + + batch, batch_uc = self.get_batch(self.input_keys, + value_dict, + num_samples=num_samples) + + context = getattr(self.cond_stage_model, 'encode')(batch) + null_context = getattr(self.cond_stage_model, 'encode')(batch_uc) + + if 'index' in kwargs: + kwargs.pop('index') + height, width = batch['target_size_as_tuple'][0].cpu().numpy().tolist() + noise = self.noise_sample(num_samples, height // self.size_factor, + width // self.size_factor, g) + if image is not None and img_to_img_strength > 0: + # run image2image + if not (ori_width == width and ori_height == height): + image = F.interpolate(image, (height, width), mode='bicubic') + with torch.autocast('cuda', enabled=False): + z = self.encode_first_stage(image, **kwargs) + else: + z = None + + # UNet use input n_prompt + samples = self.diffusion.sample( + noise=noise, + x=z, + denoising_strength=img_to_img_strength if z is not None else 1.0, + refine_strength=refine_strength, + solver=sampler, + model=self.model, + model_kwargs=[{ + 'cond': context + }, { + 'cond': null_context + }], + steps=sample_steps, + guide_scale=guide_scale, + guide_rescale=guide_rescale, + discretization=discretization, + show_progress=True, + seed=seed, + condition_fn=None, + clamp=None, + percentile=None, + t_max=None, + t_min=None, + discard_penultimate_step=None, + return_intermediate=None, + **kwargs) + + # apply refiner + if refine_strength > 0: + assert self.refiner_model is not None + assert self.refiner_cond_model is not None + with torch.autocast('cuda', enabled=False): + before_refiner_samples = self.decode_first_stage( + samples).float() + before_refiner_samples = torch.clamp( + (before_refiner_samples + 1.0) / 2.0, min=0.0, max=1.0) + + if len(pre_batch) > 0: + batch = {'prompt': prompt} + batch.update(pre_batch) + batch_uc = {'prompt': n_prompt} + batch_uc.update(pre_batch) + else: + batch, batch_uc = self.get_batch(self.input_refiner_keys, + value_dict, + num_samples=num_samples) + + context = getattr(self.refiner_cond_model, 'encode')(batch) + null_context = getattr(self.refiner_cond_model, 'encode')(batch_uc) + + samples = self.diffusion.sample( + noise=noise, + x=samples, + denoising_strength=img_to_img_strength + if z is not None else 1.0, + refine_strength=refine_strength, + refine_stage=True, + solver=sampler, + model=self.refiner_model, + model_kwargs=[{ + 'cond': context + }, { + 'cond': null_context + }], + steps=sample_steps, + guide_scale=guide_scale, + guide_rescale=guide_rescale, + discretization=discretization, + show_progress=True, + seed=seed, + condition_fn=None, + clamp=None, + percentile=None, + t_max=None, + t_min=None, + discard_penultimate_step=None, + return_intermediate=None, + **kwargs) + else: + before_refiner_samples = [None for _ in prompt] + + with torch.autocast('cuda', enabled=False): + x_samples = self.decode_first_stage(samples).float() + x_samples = torch.clamp((x_samples + 1.0) / 2.0, min=0.0, max=1.0) + + # UNet use train n_prompt + if not self.default_n_prompt == self.train_n_prompt and run_train_n: + train_n_prompt = [self.train_n_prompt] * len(prompt) + if len(pre_batch) > 0: + pre_batch = {'prompt': prompt} + batch.update(pre_batch) + batch_uc = {'prompt': train_n_prompt} + batch_uc.update(pre_batch) + else: + value_dict['negative_prompt'] = train_n_prompt + batch, batch_uc = self.get_batch(self.input_keys, + value_dict, + num_samples=num_samples) + + context = getattr(self.cond_stage_model, 'encode')(batch) + null_context = getattr(self.cond_stage_model, 'encode')(batch_uc) + + tn_samples = self.diffusion.sample( + noise=noise, + x=z, + denoising_strength=img_to_img_strength + if z is not None else 1.0, + refine_strength=refine_strength, + solver=sampler, + model=self.model, + model_kwargs=[{ + 'cond': context + }, { + 'cond': null_context + }], + steps=sample_steps, + guide_scale=guide_scale, + guide_rescale=guide_rescale, + discretization=discretization, + show_progress=we.rank == 0, + seed=seed, + condition_fn=None, + clamp=None, + percentile=None, + t_max=None, + t_min=None, + discard_penultimate_step=None, + return_intermediate=None, + **kwargs) + + if refine_strength > 0: + assert self.refiner_model is not None + assert self.refiner_cond_model is not None + with torch.autocast('cuda', enabled=False): + before_refiner_t_samples = self.decode_first_stage( + samples).float() + before_refiner_t_samples = torch.clamp( + (before_refiner_t_samples + 1.0) / 2.0, min=0.0, max=1.0) + + if len(pre_batch) > 0: + pre_batch = {'prompt': prompt} + batch.update(pre_batch) + batch_uc = {'prompt': train_n_prompt} + batch_uc.update(pre_batch) + else: + batch, batch_uc = self.get_batch(self.input_refiner_keys, + value_dict, + num_samples=num_samples) + + context = getattr(self.refiner_cond_model, 'encode')(batch) + null_context = getattr(self.refiner_cond_model, + 'encode')(batch_uc) + tn_samples = self.diffusion.sample( + noise=noise, + x=tn_samples, + denoising_strength=img_to_img_strength + if z is not None else 1.0, + refine_strength=refine_strength, + refine_stage=True, + solver=sampler, + model=self.refiner_model, + model_kwargs=[{ + 'cond': context + }, { + 'cond': null_context + }], + steps=sample_steps, + guide_scale=guide_scale, + guide_rescale=guide_rescale, + discretization=discretization, + show_progress=True, + seed=seed, + condition_fn=None, + clamp=None, + percentile=None, + t_max=None, + t_min=None, + discard_penultimate_step=None, + return_intermediate=None, + **kwargs) + else: + before_refiner_t_samples = [None for _ in prompt] + + t_x_samples = self.decode_first_stage(tn_samples).float() + + t_x_samples = torch.clamp((t_x_samples + 1.0) / 2.0, + min=0.0, + max=1.0) + else: + train_n_prompt = ['' for _ in prompt] + t_x_samples = [None for _ in prompt] + before_refiner_t_samples = [None for _ in prompt] + + outputs = list() + for p, np, tnp, img, r_img, t_img, r_t_img in zip( + prompt, n_prompt, train_n_prompt, x_samples, + before_refiner_samples, t_x_samples, before_refiner_t_samples): + one_tup = { + 'prompt': p, + 'n_prompt': np, + 'image': img, + 'before_refiner_image': r_img + } + if t_img is not None: + one_tup['train_n_prompt'] = tnp + one_tup['train_n_image'] = t_img + one_tup['train_n_before_refiner_image'] = r_t_img + outputs.append(one_tup) + + return outputs + + @staticmethod + def get_config_template(): + return dict_to_yaml('MODEL', + __class__.__name__, + LatentDiffusionXL.para_dict, + set_name=True) diff --git a/scepter/modules/model/network/train_module.py b/scepter/modules/model/network/train_module.py new file mode 100644 index 0000000..b13439a --- /dev/null +++ b/scepter/modules/model/network/train_module.py @@ -0,0 +1,49 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from abc import ABCMeta, abstractmethod + +import torch + +from scepter.modules.model.base_model import BaseModel +from scepter.modules.utils.config import dict_to_yaml + + +class TrainModule(BaseModel, metaclass=ABCMeta): + para_dict = {} + + def __init__(self, cfg, logger=None): + super(TrainModule, self).__init__(cfg, logger=logger) + self.logger = logger + self.cfg = cfg + + @abstractmethod + def forward(self, *inputs, **kwargs): + pass + + @abstractmethod + def forward_train(self, *inputs, **kwargs): + pass + + @abstractmethod + @torch.no_grad() + def forward_test(self, *inputs, **kwargs): + pass + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('networkname', + __class__.__name__, + TrainModule.para_dict, + set_name=True) diff --git a/scepter/modules/model/registry.py b/scepter/modules/model/registry.py new file mode 100644 index 0000000..b1293b0 --- /dev/null +++ b/scepter/modules/model/registry.py @@ -0,0 +1,39 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.utils.config import Config +from scepter.modules.utils.registry import Registry, build_from_config + + +def build_model(cfg, registry, logger=None, *args, **kwargs): + """ After build model, load pretrained model if exists key `pretrain`. + + pretrain (str, dict): Describes how to load pretrained model. + str, treat pretrain as model path; + dict: should contains key `path`, and other parameters token by function load_pretrained(); + """ + if not isinstance(cfg, Config): + raise TypeError(f'Config must be type dict, got {type(cfg)}') + if cfg.have('PRETRAINED_MODEL'): + pretrain_cfg = cfg.PRETRAINED_MODEL + if pretrain_cfg is not None and not isinstance(pretrain_cfg, (str)): + raise TypeError('Pretrain parameter must be a string') + else: + pretrain_cfg = None + + model = build_from_config(cfg, registry, logger=logger, *args, **kwargs) + if pretrain_cfg is not None: + if hasattr(model, 'load_pretrained_model'): + model.load_pretrained_model(pretrain_cfg) + return model + + +MODELS = Registry('MODELS', build_func=build_model) +TOKENIZERS = Registry('TOKENIZER', build_func=build_model) +EMBEDDERS = Registry('EMBEDDERS', build_func=build_model) +BACKBONES = Registry('BACKBONES', build_func=build_model) +NECKS = Registry('NECKS', build_func=build_model) +HEADS = Registry('HEADS', build_func=build_model) +BRICKS = Registry('BRICKS', build_func=build_model) +STEMS = BRICKS +LOSSES = Registry('LOSSES', build_func=build_model) +TUNERS = Registry('TUNERS', build_func=build_model) diff --git a/scepter/modules/model/tokenizer/__init__.py b/scepter/modules/model/tokenizer/__init__.py new file mode 100644 index 0000000..a7a9df6 --- /dev/null +++ b/scepter/modules/model/tokenizer/__init__.py @@ -0,0 +1,6 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.tokenizer.base_tokenizer import BaseTokenizer +from scepter.modules.model.tokenizer.tokenizer import (ClipTokenizer, + HuggingfaceTokenizer, + OpenClipTokenizer) diff --git a/scepter/modules/model/tokenizer/base_tokenizer.py b/scepter/modules/model/tokenizer/base_tokenizer.py new file mode 100644 index 0000000..bf6e7be --- /dev/null +++ b/scepter/modules/model/tokenizer/base_tokenizer.py @@ -0,0 +1,28 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from abc import ABCMeta, abstractmethod + +from scepter.modules.model.registry import TOKENIZERS +from scepter.modules.utils.config import dict_to_yaml + + +@TOKENIZERS.register_class() +class BaseTokenizer(metaclass=ABCMeta): + para_dict = {} + + def __init__(self, cfg, logger=None): + pass + + def tokenize(self, x, **kwargs): + raise NotImplementedError + + @abstractmethod + def __call__(self, x, **kwargs): + self.tokenize(x, **kwargs) + + @staticmethod + def get_config_template(): + return dict_to_yaml('TOKENIZERS', + __class__.__name__, + BaseTokenizer.para_dict, + set_name=True) diff --git a/scepter/modules/model/tokenizer/tokenizer.py b/scepter/modules/model/tokenizer/tokenizer.py new file mode 100644 index 0000000..9a03a08 --- /dev/null +++ b/scepter/modules/model/tokenizer/tokenizer.py @@ -0,0 +1,160 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import open_clip +from transformers import CLIPTokenizer as transformer_clip_tokenizer + +from scepter.modules.model.registry import TOKENIZERS +from scepter.modules.model.tokenizer import BaseTokenizer +from scepter.modules.model.tokenizer.tokenizer_component import ( + basic_clean, whitespace_clean) +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.file_system import FS + + +@TOKENIZERS.register_class() +class HuggingfaceTokenizer(BaseTokenizer): + para_dict = { + 'PRETRAINED_PATH': { + 'value': '', + 'description': "Huggingface tokenizer's pretrained path." + }, + 'LENGTH': { + 'value': 77, + 'description': "The input prompt's length." + }, + 'CLEAN': { + 'value': True, + 'description': 'Clean the special chars or not.' + } + } + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.pretrained_path = cfg.get('PRETRAINED_PATH', 'xlm-roberta-large') + self.length = cfg.get('LENGTH', 77) + self.clean = cfg.get('CLEAN', True) + + # init tokenizer + from transformers import AutoTokenizer + with FS.get_dir_to_local_dir(self.pretrained_path) as local_path: + self.tokenizer = AutoTokenizer.from_pretrained(local_path) + self.vocab_size = len(self.tokenizer) + + # special tokens + self.comma_token = self.tokenizer(',')['input_ids'][ + 2] # same as CN comma + self.sos_token = self.tokenizer( + self.tokenizer.bos_token)['input_ids'][1] + self.eos_token = self.tokenizer( + self.tokenizer.eos_token)['input_ids'][1] + self.pad_token = self.tokenizer( + self.tokenizer.pad_token)['input_ids'][1] + + def __call__(self, sequence, **kwargs): + # arguments + _kwargs = {'return_tensors': 'pt'} + if self.length is not None: + _kwargs.update({ + 'padding': 'max_length', + 'truncation': True, + 'max_length': self.length + }) + _kwargs.update(**kwargs) + + # tokenization + if isinstance(sequence, str): + sequence = [sequence] + if self.clean: + sequence = [whitespace_clean(basic_clean(u)) for u in sequence] + tokens = self.tokenizer(sequence, **_kwargs) + return tokens.input_ids + + @staticmethod + def get_config_template(): + return dict_to_yaml('TOKENIZERS', + __class__.__name__, + HuggingfaceTokenizer.para_dict, + set_name=True) + + +@TOKENIZERS.register_class() +class ClipTokenizer(BaseTokenizer): + para_dict = { + 'PRETRAINED_PATH': { + 'value': '', + 'description': "Huggingface tokenizer's pretrained path." + }, + 'LENGTH': { + 'value': 77, + 'description': "The input prompt's length." + }, + 'CLEAN': { + 'value': True, + 'description': 'Clean the special chars or not.' + } + } + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.pretrained_path = cfg.get('PRETRAINED_PATH', 'xlm-roberta-large') + self.length = cfg.get('LENGTH', 77) + self.clean = cfg.get('CLEAN', True) + with FS.get_dir_to_local_dir(self.pretrained_path, + wait_finish=True) as local_path: + self.tokenizer = transformer_clip_tokenizer.from_pretrained( + local_path) + self.vocab_size = len(self.tokenizer) + # special tokens + self.comma_token = self.tokenizer(',')['input_ids'][ + 2] # same as CN comma + self.sos_token = self.tokenizer( + self.tokenizer.bos_token)['input_ids'][1] + self.eos_token = self.tokenizer( + self.tokenizer.eos_token)['input_ids'][1] + self.pad_token = self.tokenizer( + self.tokenizer.pad_token)['input_ids'][1] + + def __call__(self, sequence, **kwargs): + # arguments + batch_encoding = self.tokenizer(sequence, + truncation=True, + max_length=self.length, + return_length=True, + return_overflowing_tokens=False, + padding='max_length', + return_tensors='pt') + return batch_encoding['input_ids'] + + @staticmethod + def get_config_template(): + return dict_to_yaml('TOKENIZERS', + __class__.__name__, + ClipTokenizer.para_dict, + set_name=True) + + +@TOKENIZERS.register_class() +class OpenClipTokenizer(BaseTokenizer): + para_dict = { + 'LENGTH': { + 'value': 77, + 'description': "The input prompt's length." + } + } + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.length = cfg.get('LENGTH', 77) + self.vocab_size = open_clip.tokenizer._tokenizer.vocab_size + + def __call__(self, sequence, **kwargs): + # arguments + tokens = open_clip.tokenize(sequence) + return tokens + + @staticmethod + def get_config_template(): + return dict_to_yaml('TOKENIZERS', + __class__.__name__, + OpenClipTokenizer.para_dict, + set_name=True) diff --git a/scepter/modules/model/tokenizer/tokenizer_component.py b/scepter/modules/model/tokenizer/tokenizer_component.py new file mode 100644 index 0000000..75c26eb --- /dev/null +++ b/scepter/modules/model/tokenizer/tokenizer_component.py @@ -0,0 +1,61 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import html +from functools import lru_cache + +import ftfy +import regex as re + + +@lru_cache() +def bytes_to_unicode(): + """ + Returns list of utf-8 byte and a corresponding list of unicode strings. + The reversible bpe codes work on unicode strings. + This means you need a large # of unicode characters in your vocab if you want to + avoid UNKs. + When you're at something like a 10B token dataset you end up needing around 5K for + decent coverage. + This is a signficant percentage of your normal, say, 32K bpe vocab. + To avoid that, we want lookup tables between utf-8 bytes and unicode strings. + And avoids mapping to whitespace/control characters the bpe code barfs on. + """ + bs = (list(range(ord('!'), + ord('~') + 1)) + list(range(ord('¡'), + ord('¬') + 1)) + + list(range(ord('®'), + ord('ÿ') + 1))) + cs = bs[:] + n = 0 + for b in range(2**8): + if b not in bs: + bs.append(b) + cs.append(2**8 + n) + n += 1 + cs = [chr(n) for n in cs] + return dict(zip(bs, cs)) + + +def get_pairs(word): + """Return set of symbol pairs in a word. + + Word is represented as tuple of symbols (symbols being variable-length strings). + """ + pairs = set() + prev_char = word[0] + for char in word[1:]: + pairs.add((prev_char, char)) + prev_char = char + return pairs + + +def basic_clean(text): + text = ftfy.fix_text(text) + text = html.unescape(html.unescape(text)) + return text.strip() + + +def whitespace_clean(text): + text = re.sub(r'\s+', ' ', text) + text = text.strip() + return text diff --git a/scepter/modules/model/tuner/__init__.py b/scepter/modules/model/tuner/__init__.py new file mode 100644 index 0000000..3ecebea --- /dev/null +++ b/scepter/modules/model/tuner/__init__.py @@ -0,0 +1,4 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.tuner.swift_tuner import (SwiftAdapter, SwiftFull, + SwiftLoRA) diff --git a/scepter/modules/model/tuner/base_tuner.py b/scepter/modules/model/tuner/base_tuner.py new file mode 100644 index 0000000..560682d --- /dev/null +++ b/scepter/modules/model/tuner/base_tuner.py @@ -0,0 +1,25 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from abc import ABCMeta + +from scepter.modules.model.base_model import BaseModel +from scepter.modules.model.registry import TUNERS +from scepter.modules.utils.config import dict_to_yaml + + +@TUNERS.register_class() +class BaseTuner(BaseModel, metaclass=ABCMeta): + para_dict = {} + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + + def forward(self, *args, **kwargs): + raise NotImplementedError + + @staticmethod + def get_config_template(): + return dict_to_yaml('TUNERS', + __class__.__name__, + BaseTuner.para_dict, + set_name=True) diff --git a/scepter/modules/model/tuner/swift_tuner.py b/scepter/modules/model/tuner/swift_tuner.py new file mode 100644 index 0000000..df877ff --- /dev/null +++ b/scepter/modules/model/tuner/swift_tuner.py @@ -0,0 +1,144 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.registry import TUNERS +from scepter.modules.utils.config import dict_to_yaml + +from .base_tuner import BaseTuner + + +@TUNERS.register_class() +class SwiftFull(BaseTuner): + para_dict = {} + + def __init__(self, cfg, logger=None): + self.logger = logger + + def __call__(self, *args, **kwargs): + return None + + @staticmethod + def get_config_template(): + return dict_to_yaml('TUNERS', + __class__.__name__, + SwiftFull.para_dict, + set_name=True) + + +@TUNERS.register_class() +class SwiftLoRA(): + para_dict = { + 'R': { + 'value': 64, + 'description': 'Rank of lora.' + }, + 'LORA_ALPHA': { + 'value': 64, + 'description': 'Lora alpha of lora, lora_alpha/rank=weight.' + }, + 'LORA_DROPOUT': { + 'value': 0.0, + 'description': 'Lora dropout, default is 0.0.' + }, + 'BIAS': { + 'value': None, + 'description': "Linear's bias for lora." + }, + 'TARGET_MODULES': { + 'value': '', + 'description': 'The norm expression of target modules.' + } + } + + def __init__(self, cfg, logger=None): + from swift import LoRAConfig + self.logger = logger + self.init_config = LoRAConfig(r=cfg.R, + lora_alpha=cfg.LORA_ALPHA, + lora_dropout=cfg.LORA_DROPOUT, + bias=cfg.BIAS, + target_modules=cfg.TARGET_MODULES) + + def __call__(self, *args, **kwargs): + return self.init_config + + @staticmethod + def get_config_template(): + return dict_to_yaml('TUNERS', + __class__.__name__, + SwiftLoRA.para_dict, + set_name=True) + + +@TUNERS.register_class() +class SwiftAdapter(BaseTuner): + para_dict = { + 'DIMS': { + 'value': [], + 'description': 'DIMS.' + }, + 'TARGET_MODULES': { + 'value': '', + 'description': 'The norm expression of target modules.' + }, + 'ADAPTER_LENGTH': { + 'value': '', + 'description': 'The length of adapter.' + } + } + + def __init__(self, cfg, logger=None): + from swift import AdapterConfig + self.logger = logger + self.init_config = AdapterConfig(dim=cfg.DIMS, + hidden_pos=0, + target_modules=cfg.TARGET_MODULES, + adapter_length=cfg.ADAPTER_LENGTH) + + def __call__(self, *args, **kwargs): + return self.init_config + + @staticmethod + def get_config_template(): + return dict_to_yaml('TUNERS', + __class__.__name__, + SwiftAdapter.para_dict, + set_name=True) + +@TUNERS.register_class() +class SwiftSCETuning(BaseTuner): + para_dict = { + 'DIMS': { + 'value': [], + 'description': 'DIMS.' + }, + 'TARGET_MODULES': { + 'value': '', + 'description': 'The norm expression of target modules.' + }, + 'DOWN_RATIO': { + 'value': 1.0, + 'description': 'The dim down ratio of tuner hidden state.' + }, + 'TUNER_MODE': { + 'value': 'identity', + 'description': 'Location of tuner operation' + } + } + + def __init__(self, cfg, logger=None): + from swift import SCETuningConfig + self.logger = logger + self.init_config = SCETuningConfig(dims=cfg.DIMS, + target_modules=cfg.TARGET_MODULES, + down_ratio=cfg.DOWN_RATIO, + tuner_mode=cfg.TUNER_MODE) + + def __call__(self, *args, **kwargs): + return self.init_config + + @staticmethod + def get_config_template(): + return dict_to_yaml('TUNERS', + __class__.__name__, + SwiftSCETuning.para_dict, + set_name=True) diff --git a/scepter/modules/model/tuner/tuner_component.py b/scepter/modules/model/tuner/tuner_component.py new file mode 100644 index 0000000..bd2183f --- /dev/null +++ b/scepter/modules/model/tuner/tuner_component.py @@ -0,0 +1,177 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import math + +import torch +import torch.nn as nn + + +class Prompt(nn.Module): + """The implementation of vision prompt tuning method. + + Visual prompt tuning (VPT) is proposed to initialize tunable prompt tokens + and prepend to the original tokens in the first layer or multiple layers. + 'Visual Prompt Tuning' by Jia et al.(2022) + See https://arxiv.org/abs/2203.12119 + + Attributes: + dim: An integer indicating the embedding dimension. + layer_num: An integer indicating number of layers. + prompt_length: An integer indicating the length of vision prompt tuning. + prompt_type: A string indicating the type of vision prompt tuning. + """ + def __init__(self, dim, layer_num, prompt_length=None, prompt_type=None): + super(Prompt, self).__init__() + self.dim = dim + self.layer_num = layer_num + self.prompt_length = prompt_length + self.prompt_type = prompt_type + + self.prompt_token = nn.Parameter(torch.zeros(1, prompt_length, dim)) + nn.init.xavier_uniform_(self.prompt_token) + + def forward(self, x): + B, N, C = x.shape + prompt_token = self.prompt_token.expand(B, -1, -1) + + if self.layer_num == 0: + x = torch.cat((x, prompt_token), dim=1) + else: + x = torch.cat((x[:, :-self.prompt_length, :], prompt_token), dim=1) + return x + + def extract(self, x): + return x[:, :-self.prompt_length, :] + + +class Adapter(nn.Module): + """The implementation of adapter tuning method. + + Adapters project input tokens by an MLP layer. + 'Parameter-Efficient Transfer Learning for NLP' by Houlsby et al.(2019) + See http://arxiv.org/abs/1902.00751 + + Attributes: + dim: An integer indicating the embedding dimension. + adapter_length: An integer indicating the length of adapter tuning. + adapter_type: A string indicating the type of adapter tuning. + """ + def __init__( + self, + dim, + adapter_length=None, + adapter_type=None, + act_layer=nn.GELU, + ): + super(Adapter, self).__init__() + self.dim = dim + self.adapter_length = adapter_length + self.adapter_type = adapter_type + self.ln1 = nn.Linear(dim, adapter_length) + self.activate = act_layer() + self.ln2 = nn.Linear(adapter_length, dim) + self.init_weights() + + def init_weights(self): + def _init_weights(m): + if isinstance(m, nn.Linear): + nn.init.xavier_uniform_(m.weight) + nn.init.normal_(m.bias, std=1e-6) + + self.apply(_init_weights) + + def forward(self, x, identity=None): + out = self.ln2(self.activate(self.ln1(x))) + if identity is None: + identity = x + out = identity + out + return out + + +class LoRA(nn.Module): + """The implementation of LoRA tuning method. + + LoRA constructs an additional layer with low-rank decomposition matrices of the weights in the network. + 'LoRA: Low-Rank Adaptation of Large Language Models' by Hu et al.(2021) + See https://arxiv.org/abs/2106.09685 + + Attributes: + dim: An integer indicating the embedding dimension. + num_heads: An integer indicating number of attention head. + lora_length: An integer indicating the length of LoRA tuning. + lora_type: A string indicating the type of LoRA tuning. + """ + def __init__( + self, + dim, + num_heads, + lora_length=None, + lora_type=None, + ): + super(LoRA, self).__init__() + self.dim = dim + self.num_heads = num_heads + if isinstance(dim, int): + self.lora_a = nn.Linear(dim, lora_length, bias=False) + self.lora_b = nn.Linear(lora_length, dim * 3, bias=False) + else: + self.lora_a = nn.Linear(dim[0], lora_length, bias=False) + self.lora_b = nn.Linear(lora_length, dim[1] * 3, bias=False) + nn.init.kaiming_uniform_(self.lora_a.weight, a=math.sqrt(5)) + nn.init.zeros_(self.lora_b.weight) + + self.lora_length = lora_length + self.lora_type = lora_type + + def forward(self, x, q, k, v): + B, N, C = x.shape + qkv_delta = self.lora_b(self.lora_a(x)) + qkv_delta = qkv_delta.reshape(B, N, 3, self.num_heads, + -1).permute(2, 0, 3, 1, 4) + q_delta, k_delta, v_delta = qkv_delta.unbind(0) + # q, k, v = q + q_delta, k + k_delta, v + v_delta + k, v = k + k_delta, v + v_delta + return q, k, v + + +class Prefix(nn.Module): + """The implementation of prefix tuning method. + + Prefix tuning optimizes the task-specific vector in the multi-head attention layer. + 'Prefix-tuning: Optimizing continuous prompts for generation' by Li & Liang(2021) + See https://arxiv.org/abs/2101.00190 + + Attributes: + dim: An integer indicating the embedding dimension. + num_heads: An integer indicating number of attention head. + prefix_length: An integer indicating the length of prefix tuning. + prefix_type: A string indicating the type of prefix tuning. + """ + def __init__( + self, + dim, + num_heads, + prefix_length=None, + prefix_type=None, + ): + super(Prefix, self).__init__() + self.dim = dim + self.num_heads = num_heads + self.prefix_length = prefix_length + self.prefix_type = prefix_type + self.prefix_key = nn.Parameter(torch.zeros(1, prefix_length, dim)) + self.prefix_value = nn.Parameter(torch.zeros(1, prefix_length, dim)) + nn.init.xavier_uniform_(self.prefix_key) + nn.init.xavier_uniform_(self.prefix_value) + + def forward(self, x, q, k, v): + B, N, C = x.shape + prefix_key = self.prefix_key.expand(B, -1, -1).reshape( + B, self.prefix_length, self.num_heads, + self.dim // self.num_heads).permute(0, 2, 1, 3) + prefix_value = self.prefix_value.expand(B, -1, -1).reshape( + B, self.prefix_length, self.num_heads, + self.dim // self.num_heads).permute(0, 2, 1, 3) + k, v = torch.cat((k, prefix_key), dim=2), torch.cat((v, prefix_value), + dim=2) + return q, k, v diff --git a/scepter/modules/model/tuner/tuner_utils.py b/scepter/modules/model/tuner/tuner_utils.py new file mode 100644 index 0000000..6eedf49 --- /dev/null +++ b/scepter/modules/model/tuner/tuner_utils.py @@ -0,0 +1,31 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import torch +import torch.nn as nn + + +def choose_weight_type(weight_type, dim): + if weight_type == 'gate': + scaling = nn.Linear(dim, 1) + elif weight_type == 'scale': + scaling = nn.Parameter(torch.Tensor(1)) + scaling.data.fill_(1) + elif weight_type == 'scale_channel': + scaling = nn.Parameter(torch.Tensor(dim)) + scaling.data.fill_(1) + elif weight_type and weight_type.startswith('scalar'): + scaling = float(weight_type.split('_')[-1]) + else: + scaling = None + return scaling + + +def get_weight_value(weight_type, scaling, x): + if weight_type in ['gate']: + scaling = torch.mean(torch.sigmoid(scaling(x)), dim=1).view(-1, 1, 1) + elif weight_type in ['scale', 'scale_channel' + ] or weight_type.startswith('scalar'): + scaling = scaling + else: + scaling = None + return scaling diff --git a/scepter/modules/model/utils/__init__.py b/scepter/modules/model/utils/__init__.py new file mode 100644 index 0000000..cc26a06 --- /dev/null +++ b/scepter/modules/model/utils/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/modules/model/utils/basic_utils.py b/scepter/modules/model/utils/basic_utils.py new file mode 100644 index 0000000..932739e --- /dev/null +++ b/scepter/modules/model/utils/basic_utils.py @@ -0,0 +1,105 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from inspect import isfunction + +import torch + + +def exists(x): + return x is not None + + +def default(val, d): + if exists(val): + return val + 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.""" + return self + + +def transfer_size(para_num): + if para_num > 1000 * 1000 * 1000 * 1000: + bill = para_num / (1000 * 1000 * 1000 * 1000) + return '{:.2f}T'.format(bill) + elif para_num > 1000 * 1000 * 1000: + gyte = para_num / (1000 * 1000 * 1000) + return '{:.2f}B'.format(gyte) + elif para_num > (1000 * 1000): + meta = para_num / (1000 * 1000) + return '{:.2f}M'.format(meta) + elif para_num > 1000: + kelo = para_num / 1000 + return '{:.2f}K'.format(kelo) + else: + return para_num + + +def count_params(model): + total_params = sum(p.numel() for p in model.parameters()) + return transfer_size(total_params) + + +def expand_dims_like(x, y): + while x.dim() != y.dim(): + x = x.unsqueeze(-1) + return x diff --git a/scepter/modules/opt/__init__.py b/scepter/modules/opt/__init__.py new file mode 100644 index 0000000..7779fc5 --- /dev/null +++ b/scepter/modules/opt/__init__.py @@ -0,0 +1,4 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.opt import lr_schedulers, optimizers diff --git a/scepter/modules/opt/lr_schedulers/__init__.py b/scepter/modules/opt/lr_schedulers/__init__.py new file mode 100644 index 0000000..36a5138 --- /dev/null +++ b/scepter/modules/opt/lr_schedulers/__init__.py @@ -0,0 +1,6 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.opt.lr_schedulers.define_schedulers import LinoPolyLR +from scepter.modules.opt.lr_schedulers.official_schedulers import * # noqa +from scepter.modules.opt.lr_schedulers.warmup import WarmupToConstantLR diff --git a/scepter/modules/opt/lr_schedulers/base_scheduler.py b/scepter/modules/opt/lr_schedulers/base_scheduler.py new file mode 100644 index 0000000..a8a85f5 --- /dev/null +++ b/scepter/modules/opt/lr_schedulers/base_scheduler.py @@ -0,0 +1,10 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + + +class BaseScheduler(): + def __init__(self, cfg, logger=None): + self.logger = logger + + def __call__(self, parameters): + return self diff --git a/scepter/modules/opt/lr_schedulers/define_schedulers.py b/scepter/modules/opt/lr_schedulers/define_schedulers.py new file mode 100644 index 0000000..b03bab5 --- /dev/null +++ b/scepter/modules/opt/lr_schedulers/define_schedulers.py @@ -0,0 +1,91 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from torch.optim.lr_scheduler import _LRScheduler + +from scepter.modules.opt.lr_schedulers.base_scheduler import BaseScheduler +from scepter.modules.opt.lr_schedulers.registry import LR_SCHEDULERS +from scepter.modules.utils.config import dict_to_yaml + + +class PolyLR(_LRScheduler): + """Decays the learning rate of each parameter group by gamma every epoch. + When last_epoch=-1, sets initial lr as lr. + + Args: + optimizer (Optimizer): Wrapped optimizer. + power (float): Multiplicative factor of learning rate decay. + end_epoch (int): Total epoches for the task. + last_epoch (int): The index of last epoch. Default: -1. + verbose (bool): If ``True``, prints a message to stdout for + each update. Default: ``False``. + """ + def __init__(self, + optimizer, + power, + end_epoch, + last_epoch=-1, + verbose=False): + self.power = power + self.end_epoch = end_epoch + super(PolyLR, self).__init__(optimizer, last_epoch, verbose) + + def get_lr(self): + assert self.last_epoch >= 0 + current_epoch = self.last_epoch + if self.last_epoch > self.end_epoch: + current_epoch = self.end_epoch + factor = (1 - current_epoch / self.end_epoch)**self.power + return [base_lr * factor for base_lr in self.base_lrs] + + def _get_closed_form_lr(self): + assert self.last_epoch >= 0 + current_epoch = self.last_epoch + if self.last_epoch > self.end_epoch: + current_epoch = self.end_epoch + factor = (1 - current_epoch / self.end_epoch)**self.power + return [base_lr * factor for base_lr in self.base_lrs] + + +@LR_SCHEDULERS.register_class() +class LinoPolyLR(BaseScheduler): + para_dict = { + 'END_EPOCH': { + 'value': 1, + 'description': 'The total epoches!' + }, + 'POWER': { + 'value': 1, + 'description': 'the lr decay rate!' + }, + 'LAST_EPOCH': { + 'value': -1, + 'description': 'the last epoch!' + } + } + + def __init__(self, cfg, logger=None): + super(LinoPolyLR, self).__init__(cfg, logger=logger) + self.power = cfg.get('POWER', 1) + self.end_epoch = cfg.END_EPOCH + self.last_epoch = cfg.get('LAST_EPOCH', -1) + + def __call__(self, optimizer): + return PolyLR(optimizer, self.power, self.end_epoch, last_epoch=-1) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('LR_SCHEDULERS', + __class__.__name__, + LinoPolyLR.para_dict, + set_name=True) diff --git a/scepter/modules/opt/lr_schedulers/official_schedulers.py b/scepter/modules/opt/lr_schedulers/official_schedulers.py new file mode 100644 index 0000000..8b304db --- /dev/null +++ b/scepter/modules/opt/lr_schedulers/official_schedulers.py @@ -0,0 +1,473 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import torch.optim.lr_scheduler as lr_sch + +from scepter.modules.opt.lr_schedulers.base_scheduler import BaseScheduler +from scepter.modules.opt.lr_schedulers.registry import LR_SCHEDULERS +from scepter.modules.utils.config import dict_to_yaml + +SUPPORT_TYPES = ('StepLR', 'CyclicLR', 'LambdaLR', 'MultiStepLR', + 'ExponentialLR', 'CosineAnnealingLR', + 'CosineAnnealingWarmRestarts', 'ReduceLROnPlateau') + + +@LR_SCHEDULERS.register_class() +class StepLR(BaseScheduler): + para_dict = { + 'STEP_SIZE': { + 'value': 1, + 'description': 'the epoch step size!' + }, + 'GAMMA': { + 'value': 0.1, + 'description': 'the gamma!' + }, + 'LAST_EPOCH': { + 'value': -1, + 'description': 'the last epoch!' + } + } + + def __init__(self, cfg, logger=None): + super(StepLR, self).__init__(cfg, logger=logger) + self.step_size = cfg.STEP_SIZE + self.gamma = cfg.get('GAMMA', 0.1) + self.last_epoch = cfg.get('LAST_EPOCH', -1) + + def __call__(self, optimizer): + return lr_sch.StepLR(optimizer, + step_size=self.step_size, + gamma=0.1, + last_epoch=-1) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('LR_SCHEDULERS', + __class__.__name__, + StepLR.para_dict, + set_name=True) + + +@LR_SCHEDULERS.register_class() +class CyclicLR(BaseScheduler): + para_dict = { + 'BASE_LR': { + 'value': 0.1, + 'description': 'the base lr!' + }, + 'MAX_LR': { + 'value': 0.5, + 'description': 'the max lr!' + }, + 'STEP_SIZE_UP': { + 'value': 2000, + 'description': 'the step size up!' + }, + 'STEP_SIZE_DOWN': { + 'value': None, + 'description': 'the step size down!' + }, + 'MODE': { + 'value': 'triangular', + 'description': 'the mode triangular!' + }, + 'GAMMA': { + 'value': 1, + 'description': 'the gamma!' + }, + 'SCALE_FN': { + 'value': None, + 'description': 'the scale fn!' + }, + 'SCALE_MODE': { + 'value': 'cycle', + 'description': 'the scale mode!' + }, + 'CYCLE_MOMENTUM': { + 'value': True, + 'description': 'the cycle momentum!' + }, + 'BASE_MOMENTUM': { + 'value': 0.8, + 'description': 'the base momentum!' + }, + 'MAX_MOMENTUM': { + 'value': 0.9, + 'description': 'the max momentum!' + }, + 'LAST_EPOCH': { + 'value': -1, + 'description': 'the last epoch!' + } + } + + def __init__(self, cfg, logger=None): + super(CyclicLR, self).__init__(cfg, logger=logger) + self.base_lr = cfg.BASE_LR + self.max_lr = cfg.MAX_LR + self.step_size_up = cfg.get('STEP_SIZE_UP', 2000) + self.step_size_down = cfg.get('STEP_SIZE_DOWN', None) + self.mode = cfg.get('MODE', 'triangular') + self.gamma = cfg.get('GAMMA', 1) + self.scale_fn = cfg.get('SCALE_FN', None) + self.scale_mode = cfg.get('SCALE_MODE', 'cycle') + self.cycle_momentum = cfg.get('CYCLE_MOMENTUM', True) + self.base_momentum = cfg.get('BASE_MOMENTUM', 0.8) + self.max_momentum = cfg.get('MAX_MOMENTUM', 0.9) + self.last_epoch = cfg.get('LAST_EPOCH', -1) + + def __call__(self, optimizer): + return lr_sch.CyclicLR(optimizer, + base_lr=self.base_lr, + max_lr=self.max_lr, + step_size_up=self.step_size_up, + step_size_down=self.step_size_down, + mode=self.mode, + gamma=self.gamma, + scale_fn=self.scale_fn, + scale_mode=self.scale_mode, + cycle_momentum=self.cycle_momentum, + base_momentum=self.base_momentum, + max_momentum=self.max_momentum, + last_epoch=self.last_epoch) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('LR_SCHEDULERS', + __class__.__name__, + CyclicLR.para_dict, + set_name=True) + + +@LR_SCHEDULERS.register_class() +class LambdaLR(BaseScheduler): + para_dict = { + 'LR_LAMBDA': { + 'value': 0.1, + 'description': 'the lr lambda!' + }, + 'LAST_EPOCH': { + 'value': -1, + 'description': 'the last epoch!' + } + } + + def __init__(self, cfg, logger=None): + super(LambdaLR, self).__init__(cfg, logger=logger) + self.lr_lambda = cfg.LR_LAMBDA + self.last_epoch = cfg.get('LAST_EPOCH', -1) + + def __call__(self, optimizer): + return lr_sch.LambdaLR(optimizer, + lr_lambda=self.lr_lambda, + last_epoch=self.last_epoch) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('LR_SCHEDULERS', + __class__.__name__, + LambdaLR.para_dict, + set_name=True) + + +@LR_SCHEDULERS.register_class() +class MultiStepLR(BaseScheduler): + para_dict = { + 'MILESTONES': { + 'value': [10000, 20000], + 'description': 'the lr lambda!' + }, + 'GAMMA': { + 'value': 0.1, + 'description': 'the gamma!' + }, + 'LAST_EPOCH': { + 'value': -1, + 'description': 'the last epoch!' + } + } + + def __init__(self, cfg, logger=None): + super(MultiStepLR, self).__init__(cfg, logger=logger) + self.milestones = cfg.MILESTONES + self.gamma = cfg.get('GAMMA', 0.1) + self.last_epoch = cfg.get('LAST_EPOCH', -1) + + def __call__(self, optimizer): + return lr_sch.MultiStepLR(optimizer, + milestones=self.milestones, + gamma=self.gamma, + last_epoch=self.last_epoch) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('LR_SCHEDULERS', + __class__.__name__, + MultiStepLR.para_dict, + set_name=True) + + +@LR_SCHEDULERS.register_class() +class ExponentialLR(BaseScheduler): + para_dict = { + 'GAMMA': { + 'value': 0.1, + 'description': 'the gamma!' + }, + 'LAST_EPOCH': { + 'value': -1, + 'description': 'the last epoch!' + } + } + + def __init__(self, cfg, logger=None): + super(ExponentialLR, self).__init__(cfg, logger=logger) + self.gamma = cfg.get('GAMMA', 0.1) + self.last_epoch = cfg.get('LAST_EPOCH', -1) + + def __call__(self, optimizer): + return lr_sch.ExponentialLR(optimizer, + gamma=self.gamma, + last_epoch=self.last_epoch) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('LR_SCHEDULERS', + __class__.__name__, + ExponentialLR.para_dict, + set_name=True) + + +@LR_SCHEDULERS.register_class() +class CosineAnnealingLR(BaseScheduler): + para_dict = { + 'T_MAX': { + 'value': 1.0, + 'description': 'the T max!' + }, + 'ETA_MIN': { + 'value': 0, + 'description': 'the eta min!' + }, + 'LAST_EPOCH': { + 'value': -1, + 'description': 'the last epoch!' + } + } + + def __init__(self, cfg, logger=None): + super(CosineAnnealingLR, self).__init__(cfg, logger=logger) + self.T_max = cfg.T_MAX + self.eta_min = cfg.get('ETA_MIN', 0) + self.last_epoch = cfg.get('LAST_EPOCH', -1) + + def __call__(self, optimizer): + return lr_sch.CosineAnnealingLR(optimizer, + T_max=self.T_max, + eta_min=self.eta_min, + last_epoch=self.last_epoch) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('LR_SCHEDULERS', + __class__.__name__, + CosineAnnealingLR.para_dict, + set_name=True) + + +@LR_SCHEDULERS.register_class() +class CosineAnnealingWarmRestarts(BaseScheduler): + para_dict = { + 'T_0': { + 'value': 1.0, + 'description': 'the T 0!' + }, + 'T_MULT': { + 'value': 1.0, + 'description': 'the T mult!' + }, + 'ETA_MIN': { + 'value': 1.0, + 'description': 'the eta min!' + }, + 'LAST_EPOCH': { + 'value': -1, + 'description': 'the last epoch!' + } + } + + def __init__(self, cfg, logger=None): + super(CosineAnnealingWarmRestarts, self).__init__(cfg, logger=logger) + self.T_0 = cfg.T_0 + self.T_mult = cfg.get('T_MULT', 1.0) + self.eta_min = cfg.get('ETA_MIN', 1.0) + self.last_epoch = cfg.get('LAST_EPOCH', -1) + + def __call__(self, optimizer): + return lr_sch.CosineAnnealingWarmRestarts(optimizer, + T_0=self.T_0, + T_mult=self.T_mult, + eta_min=self.eta_min, + last_epoch=self.last_epoch) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('LR_SCHEDULERS', + __class__.__name__, + CosineAnnealingWarmRestarts.para_dict, + set_name=True) + + +@LR_SCHEDULERS.register_class() +class ReduceLROnPlateau(BaseScheduler): + para_dict = { + 'MODE': { + 'value': 'min', + 'description': 'the mode!' + }, + 'FACTOR': { + 'value': 0.1, + 'description': 'the factor!' + }, + 'PATIENCE': { + 'value': 10, + 'description': 'the patience!' + }, + 'THRESHOLD': { + 'value': 1e-4, + 'description': 'the threshold!' + }, + 'THRESHOLD_MODE': { + 'value': 'rel', + 'description': 'the threshold mode!' + }, + 'COOLDOWN': { + 'value': 0, + 'description': 'the cooldown!' + }, + 'MIN_LR': { + 'value': 0, + 'description': 'the min lr!' + }, + 'EPS': { + 'value': 1e-8, + 'description': 'the eps!' + } + } + + def __init__(self, cfg, logger=None): + super(ReduceLROnPlateau, self).__init__(cfg, logger=logger) + self.model = cfg.get('MODE', 'min') + self.factor = cfg.get('FACTOR', 0.1) + self.patience = cfg.get('PATIENCE', 10) + self.threshold = cfg.get('THRESHOLD', 1e-4) + self.threshold_mode = cfg.get('THRESHOLD_MODE', 'rel') + self.cooldown = cfg.get('COOLDOWN', 0) + self.min_lr = cfg.get('MIN_LR', 0) + self.eps = cfg.get('EPS', 1e-8) + + def __call__(self, optimizer): + return lr_sch.ReduceLROnPlateau(optimizer, + mode=self.mode, + factor=self.factor, + patience=self.patience, + threshold=self.threshold, + threshold_mode=self.threshold_mode, + cooldown=self.cooldown, + min_lr=self.min_lr, + eps=self.eps) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('LR_SCHEDULERS', + __class__.__name__, + ReduceLROnPlateau.para_dict, + set_name=True) diff --git a/scepter/modules/opt/lr_schedulers/registry.py b/scepter/modules/opt/lr_schedulers/registry.py new file mode 100644 index 0000000..e1bc28a --- /dev/null +++ b/scepter/modules/opt/lr_schedulers/registry.py @@ -0,0 +1,38 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import inspect + +from scepter.modules.utils.registry import Registry, deep_copy + + +def build_lr_scheduler(cfg, registry, logger=None, *args, **kwargs): + from scepter.modules.utils.config import Config + if not isinstance(cfg, Config): + raise TypeError(f'config must be type dict, got {type(cfg)}') + if not cfg.have('NAME'): + raise KeyError(f'config must contain key NAME, got {cfg}') + if not isinstance(registry, Registry): + raise TypeError( + f'registry must be type Registry, got {type(registry)}') + + cfg = deep_copy(cfg) + assert kwargs is not None and 'optimizer' in kwargs + optimizer = kwargs['optimizer'] + + req_type = cfg.get('NAME') + if isinstance(req_type, str): + req_type_entry = registry.get(req_type) + if req_type_entry is None: + raise KeyError(f'{req_type} not found in {registry.name} registry') + if inspect.isclass(req_type_entry): + try: + Scheduler = req_type_entry(cfg, logger=logger) + return Scheduler(optimizer) + except Exception as e: + raise Exception(f'Failed to init class {req_type_entry}, with {e}') + else: + raise TypeError( + f'type must be str or class, got {type(req_type_entry)}') + + +LR_SCHEDULERS = Registry('LR_SCHEDULERS', build_func=build_lr_scheduler) diff --git a/scepter/modules/opt/lr_schedulers/warmup.py b/scepter/modules/opt/lr_schedulers/warmup.py new file mode 100644 index 0000000..97f30d0 --- /dev/null +++ b/scepter/modules/opt/lr_schedulers/warmup.py @@ -0,0 +1,154 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import math + +import torch.optim.lr_scheduler as lr_scheduler +from torch.optim.lr_scheduler import _LRScheduler + +from scepter.modules.opt.lr_schedulers import LR_SCHEDULERS +from scepter.modules.opt.lr_schedulers.base_scheduler import BaseScheduler +from scepter.modules.utils.config import dict_to_yaml + + +@LR_SCHEDULERS.register_class() +class WarmupToConstantLR(BaseScheduler): + para_dict = { + 'WARMUP_STEPS': { + 'value': 10000, + 'description': 'warmup steps' + } + } + + def __init__(self, cfg, logger=None): + super(WarmupToConstantLR, self).__init__(cfg, logger=logger) + warmup_steps = cfg.get('WARMUP_STEPS', 10000) + self.warmup_func = lambda step: min(1.0, step / warmup_steps) + + def __call__(self, optimizer): + return lr_scheduler.LambdaLR(optimizer, lr_lambda=[self.warmup_func]) + + @staticmethod + def get_config_template(): + return dict_to_yaml('LR_SCHEDULERS', + __class__.__name__, + WarmupToConstantLR.para_dict, + set_name=True) + + +class AnnealingLR(_LRScheduler): + def __init__(self, + optimizer, + warmup_steps, + total_steps, + decay_mode='cosine', + min_lr=0.0, + last_step=-1): + assert decay_mode in ['linear', 'cosine', 'none'] + self.optimizer = optimizer + for group in optimizer.param_groups: + if 'initial_lr' not in group: + group.setdefault('initial_lr', group['lr']) + self.base_lrs = [ + group['initial_lr'] for group in optimizer.param_groups + ] + self.warmup_steps = warmup_steps + self.total_steps = total_steps + self.decay_mode = decay_mode + self.min_lr = min_lr + self.current_step = last_step + 1 + self.step(self.current_step) + + def get_lr(self): + if self.warmup_steps > 0 and self.current_step <= self.warmup_steps: + return [ + base_lr * self.current_step / self.warmup_steps + for base_lr in self.base_lrs + ] + else: + ratio = (self.current_step - self.warmup_steps) / ( + self.total_steps - self.warmup_steps) + ratio = min(1.0, max(0.0, ratio)) + if self.decay_mode == 'linear': + return [base_lr * (1 - ratio) for base_lr in self.base_lrs] + elif self.decay_mode == 'cosine': + return [ + base_lr * (math.cos(math.pi * ratio) + 1.0) / 2.0 + for base_lr in self.base_lrs + ] + else: + return self.base_lrs + + def step(self, current_step=None): + if current_step is None: + current_step = self.current_step + 1 + self.current_step = current_step + new_lrs = self.get_lr() + new_lrs = [max(self.min_lr, new_lr) for new_lr in new_lrs] + for new_lr, group in zip(new_lrs, self.optimizer.param_groups): + group['lr'] = new_lr + + def state_dict(self): + return { + 'base_lrs': self.base_lrs, + 'warmup_steps': self.warmup_steps, + 'total_steps': self.total_steps, + 'decay_mode': self.decay_mode, + 'current_step': self.current_step + } + + def load_state_dict(self, state_dict): + self.base_lrs = state_dict['base_lrs'] + self.warmup_steps = state_dict['warmup_steps'] + self.total_steps = state_dict['total_steps'] + self.decay_mode = state_dict['decay_mode'] + self.current_step = state_dict['current_step'] + + +@LR_SCHEDULERS.register_class() +class StepAnnealingLR(BaseScheduler): + para_dict = { + 'WARMUP_STEPS': { + 'value': 0, + 'description': 'Setting warmup steps.' + }, + 'TOTAL_STEPS': { + 'value': 10000000, + 'description': 'The total training steps.' + }, + 'DECAY_MODE': { + 'value': 'cosine', + 'description': 'The lr decay mode, default is cosine.' + }, + 'MIN_LR': { + 'value': 0.0, + 'description': 'The minimum learning rate, default is 0.0.' + }, + 'LAST_STEP': { + 'value': -1, + 'description': + 'The the last step before last runing, default is -1.' + } + } + + def __init__(self, cfg, logger=None): + super(StepAnnealingLR, self).__init__(cfg, logger=logger) + self.warmup_steps = cfg.WARMUP_STEPS + self.total_steps = cfg.TOTAL_STEPS + self.decay_mode = cfg.get('DECAY_MODE', 'cosine') + self.min_lr = cfg.get('MIN_LR', 0.0) + self.last_step = cfg.get('LAST_STEP', -1) + + def __call__(self, optimizer): + return AnnealingLR(optimizer, + warmup_steps=self.warmup_steps, + total_steps=self.total_steps, + decay_mode=self.decay_mode, + min_lr=self.min_lr, + last_step=self.last_step) + + @staticmethod + def get_config_template(): + return dict_to_yaml('LR_SCHEDULERS', + __class__.__name__, + StepAnnealingLR.para_dict, + set_name=True) diff --git a/scepter/modules/opt/optimizers/__init__.py b/scepter/modules/opt/optimizers/__init__.py new file mode 100644 index 0000000..675bcbb --- /dev/null +++ b/scepter/modules/opt/optimizers/__init__.py @@ -0,0 +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.registry import OPTIMIZERS diff --git a/scepter/modules/opt/optimizers/base_optimizer.py b/scepter/modules/opt/optimizers/base_optimizer.py new file mode 100644 index 0000000..5071f28 --- /dev/null +++ b/scepter/modules/opt/optimizers/base_optimizer.py @@ -0,0 +1,10 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + + +class BaseOptimize(): + def __init__(self, cfg, logger=None): + self.logger = logger + + def __call__(self, parameters): + return self diff --git a/scepter/modules/opt/optimizers/official_optimizers.py b/scepter/modules/opt/optimizers/official_optimizers.py new file mode 100644 index 0000000..45909a4 --- /dev/null +++ b/scepter/modules/opt/optimizers/official_optimizers.py @@ -0,0 +1,655 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import torch.optim as optim + +from scepter.modules.opt.optimizers.base_optimizer import BaseOptimize +from scepter.modules.opt.optimizers.registry import OPTIMIZERS +from scepter.modules.utils.config import dict_to_yaml + +SUPPORT_TYPES = ('Adadelta', 'Adagrad', 'Adam', 'Adamax', 'AdamW', 'ASGD', + 'LBFGS', 'RMSprop', 'Rprop', 'SGD', 'SparseAdam') + + +@OPTIMIZERS.register_class() +class SGD(BaseOptimize): + para_dict = { + 'LEARNING_RATE': { + 'value': 0.1, + 'description': 'the initial learning rate!' + }, + 'MOMENTUM': { + 'value': 0, + 'description': 'the momentum!' + }, + 'DAMPENING': { + 'value': 0, + 'description': 'the dampening!' + }, + 'WEIGHT_DECAY': { + 'value': 0, + 'description': 'the weight decay!' + }, + 'NESTEROV': { + 'value': False, + 'description': 'the nesterov!' + } + } + + def __init__(self, cfg, logger=None): + super(SGD, self).__init__(cfg, logger=logger) + self.lr = cfg.LEARNING_RATE + self.momentum = cfg.get('MOMENTUM', 0) + self.dampening = cfg.get('DAMPENING', 0) + self.weight_decay = cfg.get('WEIGHT_DECAY', 0) + self.nesterov = cfg.get('NESTEROV', False) + + def __call__(self, parameters): + return optim.SGD(parameters, + lr=self.lr, + momentum=self.momentum, + dampening=self.dampening, + weight_decay=self.weight_decay, + nesterov=self.nesterov) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('OPTIMIZER', + __class__.__name__, + SGD.para_dict, + set_name=True) + + +@OPTIMIZERS.register_class() +class Adadelta(BaseOptimize): + para_dict = { + 'LEARNING_RATE': { + 'value': 1.0, + 'description': 'the initial learning rate!' + }, + 'RHO': { + 'value': 0, + 'description': 'the rho!' + }, + 'EPS': { + 'value': 1e-6, + 'description': 'the eps!' + }, + 'WEIGHT_DECAY': { + 'value': 0, + 'description': 'the weight decay!' + } + } + + def __init__(self, cfg, logger=None): + super(Adadelta, self).__init__(cfg, logger=logger) + self.lr = cfg.get('LEARNING_RATE', 1.0) + self.rho = cfg.get('RHO', 0.0) + self.eps = cfg.get('EPS', 1e-6) + self.weight_decay = cfg.get('WEIGHT_DECAY', 0) + + def __call__(self, parameters): + return optim.Adadelta(parameters, + lr=self.lr, + rho=self.rho, + eps=self.eps, + weight_decay=self.weight_decay) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('OPTIMIZER', + __class__.__name__, + Adadelta.para_dict, + set_name=True) + + +@OPTIMIZERS.register_class() +class Adagrad(BaseOptimize): + para_dict = { + 'LEARNING_RATE': { + 'value': 1e-2, + 'description': 'the initial learning rate!' + }, + 'LEARNING_RATE_DECAY': { + 'value': 0, + 'description': 'the lr decay!' + }, + 'EPS': { + 'value': 1e-10, + 'description': 'the eps!' + }, + 'WEIGHT_DECAY': { + 'value': 0, + 'description': 'the weight decay!' + }, + 'INITIAL_ACCUMULATOR_VALUE': { + 'value': 0, + 'description': 'the initial accumulator value!' + } + } + + def __init__(self, cfg, logger=None): + super(Adagrad, self).__init__(cfg, logger=logger) + self.lr = cfg.get('LEARNING_RATE', 1e-2) + self.lr_decay = cfg.get('LEARNING_RATE_DECAY', 0.0) + self.weight_decay = cfg.get('WEIGHT_DECAY', 0) + self.initial_accumulator_value = cfg.get('INITIAL_ACCUMULATOR_VALUE', + 0) + self.eps = cfg.get('EPS', 1e-10) + + def __call__(self, parameters): + return optim.Adagrad( + parameters, + lr=self.lr, + lr_decay=self.lr_decay, + weight_decay=self.weight_decay, + initial_accumulator_value=self.initial_accumulator_value, + eps=self.eps) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('OPTIMIZER', + __class__.__name__, + Adagrad.para_dict, + set_name=True) + + +@OPTIMIZERS.register_class() +class Adam(BaseOptimize): + para_dict = { + 'LEARNING_RATE': { + 'value': 1e-2, + 'description': 'the initial learning rate!' + }, + 'BETAS': { + 'value': [0.9, 0.999], + 'description': 'the rho!' + }, + 'EPS': { + 'value': 1e-6, + 'description': 'the eps!' + }, + 'WEIGHT_DECAY': { + 'value': 0, + 'description': 'the weight decay!' + }, + 'AMSGRAD': { + 'value': False, + 'description': 'the amsgrad!' + } + } + + def __init__(self, cfg, logger=None): + super(Adam, self).__init__(cfg, logger=logger) + self.lr = cfg.get('LEARNING_RATE', 1e-2) + self.betas = cfg.get('BETAS', [0.9, 0.999]) + self.weight_decay = cfg.get('WEIGHT_DECAY', 0) + self.amsgrad = cfg.get('AMSGRAD', False) + self.eps = cfg.get('EPS', 1e-10) + + def __call__(self, parameters): + return optim.Adam(parameters, + lr=self.lr, + betas=tuple(self.betas), + eps=self.eps, + weight_decay=self.weight_decay, + amsgrad=self.amsgrad) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('OPTIMIZER', + __class__.__name__, + Adam.para_dict, + set_name=True) + + +@OPTIMIZERS.register_class() +class Adamax(BaseOptimize): + para_dict = { + 'LEARNING_RATE': { + 'value': 2e-3, + 'description': 'the initial learning rate!' + }, + 'BETAS': { + 'value': [0.9, 0.999], + 'description': 'the rho!' + }, + 'EPS': { + 'value': 1e-8, + 'description': 'the eps!' + }, + 'WEIGHT_DECAY': { + 'value': 0, + 'description': 'the weight decay!' + } + } + + def __init__(self, cfg, logger=None): + super(Adamax, self).__init__(cfg, logger=logger) + self.lr = cfg.get('LEARNING_RATE', 2e-3) + self.betas = cfg.get('BETAS', [0.9, 0.999]) + self.weight_decay = cfg.get('WEIGHT_DECAY', 0) + self.eps = cfg.get('EPS', 1e-8) + + def __call__(self, parameters): + return optim.Adamax(parameters, + lr=self.lr, + betas=tuple(self.betas), + eps=self.eps, + weight_decay=self.weight_decay) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('OPTIMIZER', + __class__.__name__, + Adamax.para_dict, + set_name=True) + + +@OPTIMIZERS.register_class() +class AdamW(BaseOptimize): + para_dict = { + 'LEARNING_RATE': { + 'value': 1e-3, + 'description': 'the initial learning rate!' + }, + 'BETAS': { + 'value': [0.9, 0.999], + 'description': 'the rho!' + }, + 'EPS': { + 'value': 1e-8, + 'description': 'the eps!' + }, + 'WEIGHT_DECAY': { + 'value': 0, + 'description': 'the weight decay!' + }, + 'AMSGRAD': { + 'value': False, + 'description': 'the amsgrad!' + } + } + + def __init__(self, cfg, logger=None): + super(AdamW, self).__init__(cfg, logger=logger) + self.lr = cfg.get('LEARNING_RATE', 1e-3) + self.betas = cfg.get('BETAS', [0.9, 0.999]) + self.weight_decay = cfg.get('WEIGHT_DECAY', 0) + self.eps = cfg.get('EPS', 1e-8) + self.amsgrad = cfg.get('AMSGRAD', False) + + def __call__(self, parameters): + return optim.AdamW(parameters, + lr=self.lr, + betas=tuple(self.betas), + eps=self.eps, + weight_decay=self.weight_decay, + amsgrad=self.amsgrad) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('OPTIMIZER', + __class__.__name__, + AdamW.para_dict, + set_name=True) + + +@OPTIMIZERS.register_class() +class ASGD(BaseOptimize): + para_dict = { + 'LEARNING_RATE': { + 'value': 1e-2, + 'description': 'the initial learning rate!' + }, + 'LAMBD': { + 'value': 1e-4, + 'description': 'the rho!' + }, + 'ALPHA': { + 'value': 0.75, + 'description': 'the alpha!' + }, + 'T0': { + 'value': 1e6, + 'description': 'the t0!' + }, + 'WEIGHT_DECAY': { + 'value': 0, + 'description': 'the weight decay!' + } + } + + def __init__(self, cfg, logger=None): + super(ASGD, self).__init__(cfg, logger=logger) + self.lr = cfg.get('LEARNING_RATE', 1e-2) + self.lambd = cfg.get('LAMBD', 1e-4) + self.alpha = cfg.get('ALPHA', 0.75) + self.t0 = cfg.get('T0', 1e6) + self.weight_decay = cfg.get('WEIGHT_DECAY', 0) + + def __call__(self, parameters): + return optim.ASGD(parameters, + lr=self.lr, + lambd=self.lambd, + alpha=self.alpha, + t0=self.t0, + weight_decay=self.weight_decay) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('OPTIMIZER', + __class__.__name__, + ASGD.para_dict, + set_name=True) + + +@OPTIMIZERS.register_class() +class LBFGS(BaseOptimize): + para_dict = { + 'LEARNING_RATE': { + 'value': 1.0, + 'description': 'the initial learning rate!' + }, + 'MAX_ITER': { + 'value': 20, + 'description': 'the max iter!' + }, + 'MAX_EVAL': { + 'value': None, + 'description': 'the max eval!' + }, + 'TOLERANCE_GRAD': { + 'value': 1e-7, + 'description': 'the tolerance grad!' + }, + 'TOLERANCE_CHANGE': { + 'value': 1e-9, + 'description': 'the tolerance change!' + }, + 'HISTORY_SIZE': { + 'value': 100, + 'description': 'the history size!' + }, + 'LINE_SEARCH_FN': { + 'value': None, + 'description': 'the line search fn!' + } + } + + def __init__(self, cfg, logger=None): + super(LBFGS, self).__init__(cfg, logger=logger) + self.lr = cfg.get('LEARNING_RATE', 1) + self.max_iter = cfg.get('MAX_ITER', 20) + self.max_eval = cfg.get('MAX_EVAL', None) + self.tolerance_grad = cfg.get('TOLERANCE_GRAD', 1e-7) + self.tolerance_change = cfg.get('TOLERANCE_CHANGE', 1e-9) + self.history_size = cfg.get('HISTORY_SIZE', 100) + self.line_search_fn = cfg.get('LINE_SEARCH_FN', None) + + def __call__(self, parameters): + return optim.LBFGS(parameters, + lr=self.lr, + max_iter=self.max_iter, + max_eval=self.max_eval, + tolerance_grad=self.tolerance_grad, + tolerance_change=self.tolerance_change, + history_size=self.history_size, + line_search_fn=self.line_search_fn) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('OPTIMIZER', + __class__.__name__, + LBFGS.para_dict, + set_name=True) + + +@OPTIMIZERS.register_class() +class RMSprop(BaseOptimize): + para_dict = { + 'LEARNING_RATE': { + 'value': 1e-2, + 'description': 'the initial learning rate!' + }, + 'ALPHA': { + 'value': 0.99, + 'description': 'the alpha!' + }, + 'EPS': { + 'value': 1e-8, + 'description': 'the eps!' + }, + 'WEIGHT_DECAY': { + 'value': 0, + 'description': 'the weight decay!' + }, + 'MOMENTUM': { + 'value': 0, + 'description': 'the momentum!' + }, + 'CENTERED': { + 'value': False, + 'description': 'the centered!' + } + } + + def __init__(self, cfg, logger=None): + super(RMSprop, self).__init__(cfg, logger=logger) + self.lr = cfg.get('LEARNING_RATE', 1e-2) + self.alpha = cfg.get('ALPHA', 0.99) + self.eps = cfg.get('EPS', 1e-8) + self.weight_decay = cfg.get('WEIGHT_DECAY', False) + self.momentum = cfg.get('MOMENTUM', 0) + self.centered = cfg.get('CENTERED', False) + + def __call__(self, parameters): + return optim.RMSprop(parameters, + lr=self.lr, + alpha=self.alpha, + eps=self.eps, + weight_decay=self.weight_decay, + momentum=self.momentum, + centered=self.centered) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('OPTIMIZER', + __class__.__name__, + RMSprop.para_dict, + set_name=True) + + +@OPTIMIZERS.register_class() +class Rprop(BaseOptimize): + para_dict = { + 'LEARNING_RATE': { + 'value': 1.0, + 'description': 'the initial learning rate!' + }, + 'ETAS': { + 'value': [0.5, 1.2], + 'description': 'the etas!' + }, + 'STEP_SIZES': { + 'value': [1e-6, 50], + 'description': 'the step sizes!' + } + } + + def __init__(self, cfg, logger=None): + super(Rprop, self).__init__(cfg, logger=logger) + self.lr = cfg.get('LEARNING_RATE', 1e-2) + self.etas = cfg.get('ETAS', [0.5, 1.2]) + self.step_sizes = cfg.get('STEP_SIZES', [1e-6, 50]) + + def __call__(self, parameters): + return optim.Rprop(parameters, + lr=self.lr, + etas=tuple(self.etas), + step_sizes=tuple(self.step_sizes)) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('OPTIMIZER', + __class__.__name__, + Rprop.para_dict, + set_name=True) + + +@OPTIMIZERS.register_class() +class SparseAdam(BaseOptimize): + para_dict = { + 'LEARNING_RATE': { + 'value': 1.0, + 'description': 'the initial learning rate!' + }, + 'BETAS': { + 'value': [0.9, 0.999], + 'description': 'the betas!' + }, + 'EPS': { + 'value': 1e-8, + 'description': 'the eps!' + } + } + + def __init__(self, cfg, logger=None): + super(SparseAdam, self).__init__(cfg, logger=logger) + self.lr = cfg.get('LEARNING_RATE', 1e-3) + self.betas = cfg.get('BETAS', [0.9, 0.999]) + self.eps = cfg.get('EPS', 1e-8) + + def __call__(self, parameters): + return optim.SparseAdam(parameters, + lr=self.lr, + betas=tuple(self.betas), + eps=self.eps) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('OPTIMIZER', + __class__.__name__, + SparseAdam.para_dict, + set_name=True) diff --git a/scepter/modules/opt/optimizers/registry.py b/scepter/modules/opt/optimizers/registry.py new file mode 100644 index 0000000..b3bc53a --- /dev/null +++ b/scepter/modules/opt/optimizers/registry.py @@ -0,0 +1,39 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import inspect + +from scepter.modules.utils.registry import Registry, deep_copy + + +def build_optimizer(cfg, registry, logger=None, *args, **kwargs): + from scepter.modules.utils.config import Config + if not isinstance(cfg, Config): + raise TypeError(f'config must be type dict, got {type(cfg)}') + if not cfg.have('NAME'): + raise KeyError(f'config must contain key NAME, got {cfg}') + if not isinstance(registry, Registry): + raise TypeError( + f'registry must be type Registry, got {type(registry)}') + + assert kwargs is not None and 'parameters' in kwargs + parameters = kwargs['parameters'] + + cfg = deep_copy(cfg) + + req_type = cfg.get('NAME') + if isinstance(req_type, str): + req_type_entry = registry.get(req_type) + if req_type_entry is None: + raise KeyError(f'{req_type} not found in {registry.name} registry') + if inspect.isclass(req_type_entry): + try: + Opti = req_type_entry(cfg, logger=logger) + return Opti(parameters) + except Exception as e: + raise Exception(f'Failed to init class {req_type_entry}, with {e}') + else: + raise TypeError( + f'type must be str or class, got {type(req_type_entry)}') + + +OPTIMIZERS = Registry('OPTIMIZERS', build_func=build_optimizer) diff --git a/scepter/modules/solver/__init__.py b/scepter/modules/solver/__init__.py new file mode 100644 index 0000000..11eff4c --- /dev/null +++ b/scepter/modules/solver/__init__.py @@ -0,0 +1,6 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.solver import hooks +from scepter.modules.solver.base_solver import BaseSolver +from scepter.modules.solver.diffusion_solver import LatentDiffusionSolver +from scepter.modules.solver.train_val_solver import TrainValSolver diff --git a/scepter/modules/solver/base_solver.py b/scepter/modules/solver/base_solver.py new file mode 100644 index 0000000..e62d59c --- /dev/null +++ b/scepter/modules/solver/base_solver.py @@ -0,0 +1,918 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import copy +import numbers +import os +import warnings +from abc import ABCMeta +from collections import OrderedDict, defaultdict + +import torch +from torch.nn.parallel import DistributedDataParallel + +from scepter.modules.data.dataset import DATASETS +from scepter.modules.model.base_model import BaseModel +from scepter.modules.model.metric.registry import METRICS +from scepter.modules.model.registry import MODELS +from scepter.modules.opt.lr_schedulers import LR_SCHEDULERS +from scepter.modules.opt.optimizers import OPTIMIZERS +from scepter.modules.solver.hooks import HOOKS +from scepter.modules.utils.config import Config, dict_to_yaml +from scepter.modules.utils.data import transfer_data_to_cuda +from scepter.modules.utils.directory import get_relative_folder, osp_path +from scepter.modules.utils.distribute import dist, gather_data, we +from scepter.modules.utils.file_system import FS +from scepter.modules.utils.logger import get_logger, init_logger +from scepter.modules.utils.probe import (ProbeData, merge_gathered_probe, + register_data) + +try: + import pytorch_lightning as pl + + class PyLightningWrapper(pl.LightningModule): + def __init__(self, solver): + super().__init__() + self.solver = solver + self.model = self.solver.model + + def configure_optimizers(self): + optimizer, lr_scheduler = self.solver.optimizer, self.solver.lr_scheduler + if lr_scheduler is None: + return [optimizer] + return [optimizer], [lr_scheduler] + + def forward(self, x): + return self.model(x) + + def training_step(self, batch, batch_idx): + mode = 'train' + self.solver._mode = mode + self.solver._epoch = self.current_epoch + self.solver._total_iter[mode] = self.global_step + self.solver._iter[mode] = batch_idx + if self.current_epoch >= 1: + self.solver._epoch_max_iter[mode] = ( + self.global_step - batch_idx) // self.current_epoch + self.solver.before_iter(self.solver.hooks_dict[mode]) + results = self.solver.run_step_train(batch, + batch_idx, + step=self.global_step, + rank=self.global_rank) + self.solver._iter_outputs[mode] = self._reduce_scalar(results) + self.set_log(mode) + self.solver.after_iter(self.solver.hooks_dict[mode]) + return results['loss'] + + @torch.no_grad() + def validation_step(self, batch, batch_idx): + mode = 'eval' + self.solver._mode = mode + self.solver._iter[mode] = batch_idx + self.solver.before_iter(self.solver.hooks_dict[mode]) + results = self.solver.run_step_eval(batch, + batch_idx, + step=self.global_step, + rank=self.global_rank) + self.solver._iter_outputs[mode] = self._reduce_scalar(results) + self.set_log(mode) + self.solver.after_iter(self.solver.hooks_dict[mode]) + + @torch.no_grad() + def test_step(self, batch, batch_idx): + mode = 'test' + self.solver._mode = mode + self.solver.before_iter(self.hooks_dict[mode]) + results = self.solver.run_step_test(batch, + batch_idx, + step=self.global_step, + rank=self.global_rank) + self.solver._iter_outputs[mode] = self._reduce_scalar(results) + self.set_log(mode) + self.solver.after_iter(self.solver.hooks_dict[mode]) + + # redefine the folowing function + def on_train_epoch_start(self) -> None: + self.solver.before_epoch(self.solver.hooks_dict['train']) + + def on_train_epoch_end(self) -> None: + self.solver.after_epoch(self.solver.hooks_dict['train']) + + def on_validation_epoch_start(self) -> None: + self.solver.before_epoch(self.solver.hooks_dict['eval']) + + def on_validation_epoch_end(self) -> None: + self.solver.after_epoch(self.solver.hooks_dict['eval']) + + def on_test_epoch_start(self) -> None: + self.solver.before_epoch(self.solver.hooks_dict['test']) + + def on_test_epoch_end(self) -> None: + self.solver.after_epoch(self.solver.hooks_dict['test']) + + def setup(self, stage: str) -> None: + self.solver.logger = get_logger(name='std_torch') + self.solver._prefix = FS.init_fs_client(self.solver.file_system, + logger=self.solver.logger) + self.solver._local_rank = self.global_rank + we.rank = self.global_rank + local_devices = os.environ.get( + 'LOCAL_WORLD_SIZE') or torch.cuda.device_count() + local_devices = int(local_devices) + we.device_count = local_devices + we.device_id = self.global_rank % local_devices + self.solver.set_up_pre() + super().setup(stage) + + def on_fit_start(self): + super().on_fit_start() + self.solver.before_solve() + + def on_fit_end(self): + self.solver.after_solve() + super().on_fit_end() + + def _reduce_scalar(self, data_dict: dict): + """ Only reduce all scalar tensor values if distributed. + Any way, loss tensor will be specially processed just in case. + + Args: + data_dict: Dict result returned by model. + + Returns: + A new data dict whose tensor scalar values is all-reduced. + + """ + if isinstance(data_dict, OrderedDict): + keys = data_dict.keys() + else: + keys = sorted(list(data_dict.keys())) + + ret = OrderedDict() + for key in keys: + value = data_dict[key] + if isinstance(value, torch.Tensor) and value.ndim == 0: + ret[key] = value.data.clone() + else: + ret[key] = value + return ret + + def set_log(self, mode): + now_outputs = self.solver._iter_outputs[mode] + extra_vars = self.solver.collect_log_vars() + now_outputs.update(extra_vars) + for k, v in now_outputs.items(): + if k == 'batch_size': + continue + if isinstance(v, torch.Tensor) and v.ndim == 0 or isinstance( + v, numbers.Number): + if mode == 'train': + self.log(f'{mode}_{k}', + v, + prog_bar=True, + logger=True, + on_step=True, + rank_zero_only=True, + sync_dist=True) + else: + self.log(f'{mode}_{k}', + v, + prog_bar=True, + logger=True, + on_epoch=True, + rank_zero_only=True, + sync_dist=True) + +except Exception as e: + warnings.warn(f'{e}') + + +class BaseSolver(object, metaclass=ABCMeta): + """ Base Solver. + To initialize the solver. + We have to initialize the data, model, optimizer and schedule. + To process the common processing we also have to initialize the hooks. + How to support Pytorch_lightning framework? Take a simple task as an examples. + """ + para_dict = { + 'TRAIN_PRECISION': { + 'value': 32, + 'description': 'The precision for train process.' + }, + 'FILE_SYSTEM': {}, + 'ACCU_STEP': { + 'value': + 1, + 'description': + 'When use ddp, the grad accumulate steps for each process.' + }, + 'RESUME_FROM': { + 'value': '', + 'description': 'Resume from some state of training!' + }, + 'MAX_EPOCHS': { + 'value': 10, + 'description': 'Max epochs for training.' + }, + 'NUM_FOLDS': { + 'value': 1, + 'description': 'Num folds for training.' + }, + 'WORK_DIR': { + 'value': '', + 'description': 'Save dir of the training log or model.' + }, + 'LOG_FILE': { + 'value': '', + 'description': 'Save log path.' + }, + 'EVAL_INTERVAL': { + 'value': 1, + 'description': 'Eval the model interval.' + }, + 'EXTRA_KEYS': { + 'value': [], + 'description': 'The extra keys for metric.' + }, + 'TRAIN_DATA': { + 'description': 'Train data config.' + }, + 'EVAL_DATA': { + 'description': 'Eval data config.' + }, + 'TEST_DATA': { + 'description': 'Test data config.' + }, + 'TRAIN_HOOKS': [], + 'EVAL_HOOKS': [], + 'TEST_HOOKS': [], + 'MODEL': {}, + 'OPTIMIZER': {}, + 'LR_SCHEDULER': {}, + 'METRICS': [] + } + + def __init__(self, cfg, logger=None): + # initialize some hyperparameters + self.file_system = cfg.get('FILE_SYSTEM', None) + self.work_dir: str = cfg.WORK_DIR + self.pl_dir = self.work_dir + self.log_file = osp_path(self.work_dir, cfg.LOG_FILE) + self.optimizer, self.lr_scheduler = None, None + self.cfg = cfg + self.logger = logger + self.resume_from: str = cfg.RESUME_FROM + self.max_epochs: int = cfg.MAX_EPOCHS + self.use_pl = we.use_pl + self.train_precision = self.cfg.get('TRAIN_PRECISION', 32) + self._mode_set = set() + self._mode = 'train' + self.probe_ins = {} + self.clear_probe_ins = {} + self._num_folds: int = 1 + if not self.use_pl: + world_size = we.world_size + if world_size > 1: + self._num_folds: int = cfg.NUM_FOLDS + if cfg.have('MODE'): + self._mode_set.add(cfg.MODE) + self._mode = cfg.MODE + if we.is_distributed: + self.accu_step = cfg.get('ACCU_STEP', 1) + + self.do_step = True + self.hooks_dict = {'train': [], 'eval': [], 'test': []} + self.datas = {} + # Other initialized parameters + self._epoch: int = 0 + # epoch_max_iter, iter, total_iter, iter_outputs, epoch_outputs + # values is different according to self._mode + self._epoch_max_iter: defaultdict = defaultdict(int) + self._iter: defaultdict = defaultdict(int) + self._total_iter: defaultdict = defaultdict(int) + self._iter_outputs = defaultdict(dict) + self._agg_iter_outputs = defaultdict(dict) + self._epoch_outputs = defaultdict(dict) + self._probe_data = defaultdict(dict) + self._dist_data = defaultdict(dict) + self._model_parameters = 0 + self._model_flops = 0 + self._loss = None # loss tensor + self._local_rank = we.rank + if isinstance(self.file_system, list): + for file_sys in self.file_system: + FS.init_fs_client(file_sys, logger=self.logger) + elif self.file_system is not None: + FS.init_fs_client(self.file_system, logger=self.logger) + self._prefix = FS.get_fs_client(self.work_dir).get_prefix() + if not FS.exists(self.work_dir): + FS.make_dir(self.work_dir) + self.logger.info( + f"Parse work dir {self.work_dir}'s prefix is {self._prefix}") + + def set_up_pre(self): + # initialize Enviranment + if self._local_rank == 0: + if self.log_file.startswith('file://'): + save_folder = get_relative_folder(self.log_file, -1) + if not os.path.exists(save_folder): + os.makedirs(save_folder) + elif not self.log_file.startswith(self._prefix): + self.log_file = os.path.join(self._prefix, self.log_file) + init_logger(self.logger, + log_file=self.log_file, + dist_launcher='pytorch') + self.construct_hook() + + def __setattr__(self, key, value): + if isinstance(value, BaseModel): + self.probe_ins[key] = value.probe_data + self.clear_probe_ins[key] = value.clear_probe + super().__setattr__(key, value) + + def set_up(self): + self.construct_data() + self.construct_model() + self.construct_metrics() + if not self.use_pl: + self.model_to_device() + self.init_opti() + if self.use_pl: + self.local_work_dir, _ = FS.map_to_local(self.work_dir) + os.makedirs(self.local_work_dir, exist_ok=True) + # resume + resume_local_file = None + if self.resume_from is not None and FS.exists(self.resume_from): + with FS.get_from(self.resume_from, + wait_finish=True) as local_file: + self.logger.info( + f'Loading checkpoint from {self.resume_from}') + resume_local_file = local_file + self.pl_ins = PyLightningWrapper(self) + self.pl_trainer = pl.Trainer( + default_root_dir=self.local_work_dir, + max_epochs=self.max_epochs, + precision=self.train_precision, + accelerator='auto', + devices='auto', + check_val_every_n_epoch=self.eval_interval, + resume_from_checkpoint=resume_local_file) + self.pl_dir = os.path.join( + self.work_dir, + '/'.join(self.pl_trainer.log_dir.split('/')[-2:])) + + def construct_data(self): + def one_device_init(): + # initialize data + assert self.cfg.have('TRAIN_DATA') or self.cfg.have( + 'EVAL_DATA') or self.cfg.have('TEST_DATA') + if self.cfg.have('TRAIN_DATA') and ('train' in self._mode_set + or len(self._mode_set) < 1): + self.cfg.TRAIN_DATA.NUM_FOLDS = self.num_folds + train_data = DATASETS.build(self.cfg.TRAIN_DATA, + logger=self.logger) + self.datas['train'] = train_data + self._mode_set.add('train') + if not self.use_pl: + self._epoch_max_iter['train'] = len( + train_data.dataloader) // self.num_folds + 1 + else: + self._epoch_max_iter['train'] = -1 + if self.cfg.have('EVAL_DATA'): + eval_data = DATASETS.build(self.cfg.EVAL_DATA, + logger=self.logger) + self.datas['eval'] = eval_data + self._mode_set.add('eval') + if not self.use_pl: + self._epoch_max_iter['eval'] = len(eval_data.dataloader) + else: + self._epoch_max_iter['eval'] = -1 + if self.cfg.have('TEST_DATA'): + test_data = DATASETS.build(self.cfg.TEST_DATA, + logger=self.logger) + self.datas['test'] = test_data + if not self.use_pl: + self._epoch_max_iter['test'] = len(test_data.dataloader) + else: + self._epoch_max_iter['test'] = -1 + self._mode_set.add('test') + + one_device_init() + + def construct_hook(self): + # initialize data + assert self.use_pl or self.cfg.have('TRAIN_HOOKS') or self.cfg.have( + 'EVAL_HOOKS') or self.cfg.have('TEST_HOOKS') + if self.cfg.have('TRAIN_HOOKS') and ('train' in self._mode_set + or len(self._mode_set) < 1): + self.hooks_dict['train'] = self._load_hook(self.cfg.TRAIN_HOOKS) + if self.cfg.have('EVAL_HOOKS'): + assert self.cfg.have('EVAL_HOOKS') + self.hooks_dict['eval'] = self._load_hook(self.cfg.EVAL_HOOKS) + if self.cfg.have('TEST_HOOKS') and ('test' in self._mode_set + or 'eval' not in self._mode_set): + self.hooks_dict['test'] = self._load_hook(self.cfg.TEST_HOOKS) + + def construct_model(self): + # initialize Model + assert self.cfg.have('MODEL') + self.model = MODELS.build(self.cfg.MODEL, logger=self.logger) + + def construct_metrics(self): + # Initial metric + self.metrics = [] + self.eval_interval = self.cfg.get('EVAL_INTERVAL', 1) + + def model_to_device(self, tg_model_ins=None): + # Initialize distributed model + if tg_model_ins is None: + tg_model = self.model + else: + tg_model = tg_model_ins + if we.is_distributed and we.sync_bn is True: + self.logger.info('Convert BatchNorm to Synchronized BatchNorm...') + tg_model = torch.nn.SyncBatchNorm.convert_sync_batchnorm(tg_model) + tg_model = tg_model.to(we.device_id) + if we.is_distributed: + tg_model = DistributedDataParallel( + tg_model, + device_ids=[torch.cuda.current_device()], + output_device=torch.cuda.current_device(), + broadcast_buffers=True) + self.logger.info('Transfer to ddp ...') + if tg_model_ins is None: + self.model = tg_model + else: + return tg_model + + def init_opti(self): + if self.cfg.have('OPTIMIZER'): + self.optimizer = OPTIMIZERS.build( + self.cfg.OPTIMIZER, + logger=self.logger, + parameters=self.model.parameters()) + if self.cfg.have('LR_SCHEDULER') and self.optimizer is not None: + self.lr_scheduler = LR_SCHEDULERS.build(self.cfg.LR_SCHEDULER, + logger=self.logger, + optimizer=self.optimizer) + + def solve(self, epoch=None, every_epoch=False): + if not self.use_pl: + if epoch is not None: + self.epoch = epoch + self.before_solve() + if self.epoch >= self.max_epochs: + self.logger.info( + f'Nothing to do because current epoch {self.epoch} greater max epoches {self.epoch}' + ) + while self.epoch < self.max_epochs: + self.solve_train() + self.solve_eval() + self.solve_test() + if 'train' not in self._mode_set and not every_epoch: + break + self.after_solve() + else: + train_dataloader = None + if 'train' in self.datas: + train_dataloader = self.datas['train'].dataloader + val_dataloader = None + if 'eval' in self.datas: + val_dataloader = self.datas['eval'].dataloader + self.pl_trainer.fit(self.pl_ins, + train_dataloaders=train_dataloader, + val_dataloaders=val_dataloader) + if 'test' in self.datas: + self.pl_trainer.test(self.pl_ins, + dataloaders=self.datas['test'].dataloader) + + def solve_train(self): + current_mode = 'train' + if current_mode in self._mode_set: + self.logger.info( + f'Begin to solve {current_mode} at Epoch [{self.epoch}/{self.max_epochs}]...' + ) + self.before_epoch(self.hooks_dict[current_mode]) + self.run_train() + self.after_epoch(self.hooks_dict[current_mode]) + + def solve_eval(self): + current_mode = 'eval' + if current_mode in self._mode_set and self.epoch % self.eval_interval == 0: + self.logger.info( + f'Begin to solve {current_mode} at Epoch [{self.epoch}/{self.max_epochs}]...' + ) + self.before_epoch(self.hooks_dict[current_mode]) + self.run_eval() + self.after_epoch(self.hooks_dict[current_mode]) + + def solve_test(self): + current_mode = 'test' + if current_mode in self._mode_set: + self.logger.info( + f'Begin to solve {current_mode} at Epoch [{self.epoch}/{self.max_epochs}]...' + ) + self.before_epoch(self.hooks_dict[current_mode]) + self.run_test() + self.after_epoch(self.hooks_dict[current_mode]) + + def before_solve(self): + for k, hooks in self.hooks_dict.items(): + [t.before_solve(self) for t in hooks] + + def after_solve(self): + for k, hooks in self.hooks_dict.items(): + [t.after_solve(self) for t in hooks] + + def run_train(self): + self.train_mode() + self.before_all_iter(self.hooks_dict[self._mode]) + for batch_idx, batch_data in enumerate( + self.datas[self._mode].dataloader): + self.before_iter(self.hooks_dict[self._mode]) + results = self.run_step_train(transfer_data_to_cuda(batch_data), + batch_idx, + step=self.total_iter, + rank=we.rank) + self._iter_outputs[self._mode] = self._reduce_scalar(results) + self.after_iter(self.hooks_dict[self._mode]) + self.after_all_iter(self.hooks_dict[self._mode]) + + def run_step_train(self, batch_data, batch_idx=0, step=None, rank=None): + results = self.model(**batch_data) + return results + + @torch.no_grad() + def run_eval(self): + self.eval_mode() + self.before_all_iter(self.hooks_dict[self._mode]) + for batch_idx, batch_data in enumerate( + self.datas[self._mode].dataloader): + self.before_iter(self.hooks_dict[self._mode]) + results = self.run_step_eval(transfer_data_to_cuda(batch_data), + batch_idx, + step=self.total_iter, + rank=we.rank) + self._iter_outputs[self._mode] = self._reduce_scalar(results) + self.after_iter(self.hooks_dict[self._mode]) + self.after_all_iter(self.hooks_dict[self._mode]) + + def run_step_eval(self, batch_data, batch_idx=0, step=None, rank=None): + results = self.model(**batch_data) + return results + + @torch.no_grad() + def run_test(self): + self.test_mode() + self.before_all_iter(self.hooks_dict[self._mode]) + for batch_idx, batch_data in enumerate( + self.datas[self._mode].dataloader): + self.before_iter(self.hooks_dict[self._mode]) + results = self.run_step_test(transfer_data_to_cuda(batch_data), + batch_idx, + step=self.total_iter, + rank=we.rank) + self._iter_outputs[self._mode] = self._reduce_scalar(results) + self.after_iter(self.hooks_dict[self._mode]) + self.after_all_iter(self.hooks_dict[self._mode]) + + def run_step_test(self, batch_data, batch_idx=0, step=None, rank=None): + results = self.model(**batch_data) + return results + + @torch.no_grad() + def register_flops(self, data, keys=[]): + from fvcore.nn import FlopCountAnalysis + if len(keys) < 1: + keys = list(data.keys()) + for key in data: + if isinstance(data[key], torch.Tensor): + batch_one_data = data[key][0, ...] + batch_one_data = torch.unsqueeze(batch_one_data, dim=0) + data[key] = batch_one_data + elif isinstance(data[key], list): + data[key] = data[key][0] + + tensor = [data[k] for k in keys] + flops = FlopCountAnalysis(self.model, tuple(tensor)) + self._model_flops = flops.total() + + def before_epoch(self, hooks): + [t.before_epoch(self) for t in hooks] + + def before_all_iter(self, hooks): + [t.before_all_iter(self) for t in hooks] + + def before_iter(self, hooks): + if not self.use_pl and self.is_train_mode: + self._epoch = self._total_iter[self._mode] // self._epoch_max_iter[ + self._mode] + 1 + if self._iter[self._mode] % self._epoch_max_iter[self._mode] == 0: + self._iter[self._mode] = 0 + [t.before_iter(self) for t in hooks] + + def after_iter(self, hooks): + [t.after_iter(self) for t in hooks] + if not self.use_pl: + self._total_iter[self._mode] += 1 + self._iter[self._mode] += 1 + self.clear_probe() + + def after_all_iter(self, hooks): + [t.after_all_iter(self) for t in hooks] + + def after_epoch(self, hooks): + [t.after_epoch(self) for t in hooks] + self._iter.clear() + self._iter_outputs.clear() + self._epoch_outputs.clear() + if self.use_pl: + FS.put_dir_from_local_dir(self.local_work_dir, self.work_dir) + + def collect_log_vars(self) -> OrderedDict: + ret = OrderedDict() + if self.is_train_mode and self.optimizer is not None: + for idx, pg in enumerate(self.optimizer.param_groups): + ret[f'pg{idx}_lr'] = pg['lr'] + return ret + + def load_checkpoint(self, checkpoint: dict): + """ + Load checkpoint function + :param checkpoint: all tensors are on cpu, you need to transfer to gpu by hand + :return: + """ + pass + + def save_checkpoint(self) -> dict: + """ + Save checkpoint function, you need to transfer all tensors to cpu by hand + :return: + """ + pass + + @property + def num_folds(self) -> int: + return self._num_folds + + @property + def epoch(self) -> int: + return self._epoch + + @epoch.setter + def epoch(self, new_epoch): + self._epoch = new_epoch + + @property + def iter(self) -> int: + return self._iter[self._mode] + + @property + def probe_data(self): + return self._probe_data[self._mode] + + @property + def total_iter(self) -> int: + return self._total_iter[self._mode] + + @property + def epoch_max_iter(self) -> int: + return self._epoch_max_iter[self._mode] + + @property + def mode(self) -> str: + return self._mode + + @property + def iter_outputs(self) -> dict: + return self._iter_outputs[self._mode] + + @property + def agg_iter_outputs(self) -> dict: + return self._agg_iter_outputs + + @agg_iter_outputs.setter + def agg_iter_outputs(self, new_outputs): + assert type(new_outputs) is dict + self._agg_iter_outputs[self._mode] = new_outputs + + @property + def epoch_outputs(self) -> dict: + return self._epoch_outputs + + @property + def is_train_mode(self): + return self._mode == 'train' + + @property + def is_eval_mode(self): + return self._mode == 'eval' + + @property + def is_test_mode(self): + return self._mode == 'test' + + def train_mode(self): + self.model.train() + self._mode = 'train' + + def eval_mode(self): + self.model.eval() + self._mode = 'eval' + + def test_mode(self): + self.model.eval() + self._mode = 'test' + + def register_probe(self, probe_data: dict): + probe_da, dist_da = register_data(probe_data, + key_prefix=__class__.__name__) + self._probe_data[self.mode].update(probe_da) + for key in dist_da: + if key not in self._dist_data[self.mode]: + self._dist_data[self.mode][key] = dist_da[key] + else: + for k, v in dist_da[key].items(): + if k in self._dist_data[self.mode][key]: + self._dist_data[self.mode][key][k] += v + else: + self._dist_data[self.mode][key][k] = v + + @property + def probe_data(self): # noqa + gather_probe_data = gather_data(self._probe_data[self.mode]) + _dist_data_list = gather_data([self._dist_data[self.mode] or {}]) + if not we.rank == 0: + self._probe_data[self.mode] = {} + self._dist_data[self.mode] = {} + # Iterate recurse the sub class's probe data. + for k, func in self.probe_ins.items(): + for kk, vv in func().items(): + self._probe_data[self.mode][f'{k}/{kk}'] = vv + if gather_probe_data is not None: + # Before processing, just merge the data. + self._probe_data[self.mode] = merge_gathered_probe( + gather_probe_data) + if _dist_data_list is not None and len(_dist_data_list) > 0: + reduce_dist_data = {} + for one_data in _dist_data_list: + for k, v in one_data.items(): + if k in reduce_dist_data: + for kk, vv in v.items(): + if kk in reduce_dist_data[k]: + reduce_dist_data[k][kk] += vv + else: + reduce_dist_data[k][kk] = vv + else: + reduce_dist_data[k] = v + self._dist_data[self.mode] = reduce_dist_data + self._probe_data[ + self.mode][f'{__class__.__name__}_distribute'] = ProbeData( + self._dist_data[self.mode]) + norm_dist_data = {} + for key, value in self._dist_data[self.mode].items(): + total = 0 + for k, v in value.items(): + total += v + norm_v = {} + for k, v in value.items(): + norm_v[k] = v / total + norm_dist_data[key] = norm_v + self._probe_data[ + self.mode][f'{__class__.__name__}_norm_distribute'] = ProbeData( + norm_dist_data) + ret_data = copy.deepcopy(self._probe_data[self.mode]) + self._probe_data[self.mode] = {} + return ret_data + + def clear_probe(self): + self._probe_data[self.mode].clear() + # Iterate recurse the sub class's probe data. + for k, func in self.clear_probe_ins.items(): + func() + + def _load_hook(self, hooks): + ret_hooks = [] + if hooks is not None and len(hooks) > 0: + for hook_cfg in hooks: + if self.use_pl: + if 'backward' in hook_cfg.NAME.lower( + ) or 'lrhook' in hook_cfg.NAME.lower( + ) or 'samplerhook' in hook_cfg.NAME.lower(): + self.logger.info( + f'Hook {hook_cfg.NAME} is not useful when use PytorchLightning!' + ) + continue + ret_hooks.append(HOOKS.build(hook_cfg, logger=self.logger)) + ret_hooks.sort(key=lambda a: a.priority) + return ret_hooks + + def get_optim_parameters(self): + return self.model.parameters() + + def __repr__(self) -> str: + return f'{self.__class__.__name__}' + + def _reduce_scalar(self, data_dict: dict): + """ Only reduce all scalar tensor values if distributed. + Any way, loss tensor will be specially processed just in case. + + Args: + data_dict: Dict result returned by model. + + Returns: + A new data dict whose tensor scalar values is all-reduced. + + """ + if 'loss' in data_dict: + self.loss = data_dict['loss'] + data_dict['loss'] = self.loss.data.clone() + + if isinstance(data_dict, OrderedDict): + keys = data_dict.keys() + else: + keys = sorted(list(data_dict.keys())) + + ret = OrderedDict() + # print([(key, type(data_dict[key])) for key in keys], f"{dist.get_rank()}", f"{self.iter}") + for key in keys: + value = data_dict[key] + if isinstance(value, torch.Tensor) and value.ndim == 0: + if dist.is_available() and dist.is_initialized(): + value = value.data.clone() + dist.all_reduce(value.div_(dist.get_world_size())) + ret[key] = value + else: + ret[key] = value + + return ret + + def _build_metrics(self, cfgs, logger=None): + if isinstance(cfgs, (list, tuple)): + for cfg in cfgs: + self._build_metrics(cfg, logger=logger) + elif isinstance(cfgs, Config): + fn = METRICS.build(cfgs, logger) + keys = cfgs.KEYS + self.metrics.append({'fn': fn, 'keys': keys}) + self._collect_keys.update(keys) + + def print_memory_status(self): + if torch.cuda.is_available(): + nvi_info = os.popen('nvidia-smi').read() + gpu_mem = nvi_info.split('\n')[9].split('|')[2].split( + '/')[0].strip() + else: + gpu_mem = '' + return gpu_mem + + def print_model_params_status(self, model=None, logger=None): + """Print the status and parameters of the model""" + if model is None: + model = self.model + if logger is None: + logger = self.logger + train_param_dict = {} + forzen_param_dict = {} + all_param_numel = 0 + for key, val in model.named_parameters(): + if val.requires_grad: + sub_key = '.'.join(key.split('.', 1)[-1].split('.', 2)[:2]) + if sub_key in train_param_dict: + train_param_dict[sub_key] += val.numel() + else: + train_param_dict[sub_key] = val.numel() + else: + sub_key = '.'.join(key.split('.', 1)[-1].split('.', 1)[:1]) + if sub_key in forzen_param_dict: + forzen_param_dict[sub_key] += val.numel() + else: + forzen_param_dict[sub_key] = val.numel() + all_param_numel += val.numel() + train_param_numel = sum(train_param_dict.values()) + forzen_param_numel = sum(forzen_param_dict.values()) + logger.info( + f'Load trainable params {train_param_numel} / {all_param_numel} = ' + f'{train_param_numel / all_param_numel:.2%}, ' + f'train part: {train_param_dict}.') + logger.info( + f'Load forzen params {forzen_param_numel} / {all_param_numel} = ' + f'{forzen_param_numel / all_param_numel:.2%}, ' + f'forzen part: {forzen_param_dict}.') + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('solvername', + __class__.__name__, + BaseSolver.para_dict, + set_name=True) diff --git a/scepter/modules/solver/diffusion_solver.py b/scepter/modules/solver/diffusion_solver.py new file mode 100644 index 0000000..f2705ae --- /dev/null +++ b/scepter/modules/solver/diffusion_solver.py @@ -0,0 +1,687 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import os +from collections import OrderedDict, defaultdict + +import numpy as np +import torch +import torch.cuda.amp as amp +from torch.distributed.fsdp import (BackwardPrefetch, CPUOffload, + FullStateDictConfig, + FullyShardedDataParallel, MixedPrecision, + ShardingStrategy, StateDictType) +from torch.nn.parallel import DistributedDataParallel +from tqdm import tqdm + +from scepter.modules.data.dataset import DATASETS +from scepter.modules.opt.lr_schedulers import LR_SCHEDULERS +from scepter.modules.opt.optimizers import OPTIMIZERS +from scepter.modules.solver import BaseSolver +from scepter.modules.solver.registry import SOLVERS +from scepter.modules.utils.config import Config, dict_to_yaml +from scepter.modules.utils.data import transfer_data_to_cuda +from scepter.modules.utils.distribute import we +from scepter.modules.utils.probe import ProbeData + +sharding_strategy_map = { + 'full_shard': ShardingStrategy.FULL_SHARD, + 'shard_grad_op': ShardingStrategy.SHARD_GRAD_OP +} + + +@SOLVERS.register_class() +class LatentDiffusionSolver(BaseSolver): + para_dict = { + 'MAX_STEPS': { + 'value': 100000, + 'description': 'The total steps for training.', + }, + 'USE_AMP': { + 'value': + False, + 'description': + 'Use amp to surpport mix precision or not, default is False.', + }, + 'DTYPE': { + 'value': 'float32', + 'description': 'The precision for training.', + }, + 'USE_FAIRSCALE': { + 'value': False, + 'description': + 'Use fairscale as the backend of ddp, default False.', + }, + 'USE_FSDP': { + 'value': False, + 'description': 'Use fsdp as the backend of ddp, default False.', + }, + 'SHARDING_STRATEGY': { + 'value': + 'shard_grad_op', + 'description': + f'The shard strategy for fsdp, select from {list(sharding_strategy_map.keys())}', + }, + 'IMAGE_LOG_STEP': { + 'value': 2000, + 'description': 'The interval for image log.', + }, + 'LOAD_MODEL_ONLY': { + 'value': + False, + 'description': + 'Only load the model rather than the optimizer and schedule, default is False.', + }, + 'CHANNELS_LAST': { + 'value': False, + 'description': 'The channels last, default is False.', + }, + 'SAMPLE_ARGS': { + 'value': + None, + 'description': + 'Sampling related parameters, default is None( use default sample args ).', + }, + 'TUNER': { + 'value': None, + 'description': 'Tuner config, default is None.', + }, + 'FREEZE': { + 'value': + None, + 'description': + 'Specify freezing and training parameters, default is None.', + } + } + para_dict.update(BaseSolver.para_dict) + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.max_steps = cfg.MAX_STEPS + self.use_amp = cfg.get('USE_AMP', False) + self.dtype = getattr(torch, cfg.DTYPE) + self.use_fairscale = cfg.get('USE_FAIRSCALE', False) + self.use_fsdp = cfg.get('USE_FSDP', False) + if self.use_fairscale and self.use_fsdp: + raise 'fairscale and fsdp is not allowed used meanwhile.' + elif self.use_fairscale: + self.logger.info('Use fairscale as the backend of ddp.') + elif self.use_fsdp: + self.logger.info('Use fsdp as the backend of ddp.') + else: + self.logger.info('Use default backend.') + self.model_shard = cfg.get('SHARDING_STRATEGY', 'full_shard') + self.image_log_step = cfg.get('IMAGE_LOG_STEP', 2000) + self._image_out = defaultdict(list) + self.load_model_only = cfg.get('LOAD_MODEL_ONLY', False) + self.channels_last = cfg.get('CHANNELS_LAST', False) + self.current_batch_data = defaultdict(dict) + self.sample_args = cfg.get('SAMPLE_ARGS', None) + self.tuner_cfg = cfg.get('TUNER', None) + self.freeze_cfg = cfg.get('FREEZE', None) + + def set_up(self): + self.construct_data() + self.construct_model() + self.construct_metrics() + self.model_to_device() + if 'train' in self.datas and self.cfg.have('OPTIMIZER'): + if we.world_size > 1: + all_batch_size = self.datas['train'].batch_size * we.world_size + else: + all_batch_size = self.datas['train'].batch_size + self.cfg.OPTIMIZER.LEARNING_RATE *= all_batch_size + self.cfg.OPTIMIZER.LEARNING_RATE /= 640 + self.init_opti() + + def construct_hook(self): + # initialize data + assert self.cfg.have('TRAIN_HOOKS') or self.cfg.have( + 'EVAL_HOOKS') or self.cfg.have('TEST_HOOKS') + if self.cfg.have('TRAIN_HOOKS'): + self.hooks_dict['train'] = self._load_hook(self.cfg.TRAIN_HOOKS) + if self.cfg.have('EVAL_HOOKS'): + self.hooks_dict['eval'] = self._load_hook(self.cfg.EVAL_HOOKS) + if self.cfg.have('TEST_HOOKS'): + self.hooks_dict['test'] = self._load_hook(self.cfg.TEST_HOOKS) + + def construct_data(self): + # assert self.cfg.have("TRAIN_DATA") or self.cfg.have("EVAL_DATA") or self.cfg.have("TEST_DATA") + if self.cfg.have('TRAIN_DATA'): + train_data = DATASETS.build(self.cfg.TRAIN_DATA, + logger=self.logger) + self.datas['train'] = train_data + self._epoch_max_iter['train'] = len(train_data.dataloader) + self._mode_set.add('train') + if self.cfg.have('EVAL_DATA') and 'train' in self._mode_set: + eval_data = DATASETS.build(self.cfg.EVAL_DATA, logger=self.logger) + self.datas['eval'] = eval_data + self._epoch_max_iter['eval'] = len(eval_data.dataloader) + self._mode_set.add('eval') + if self.cfg.have('TEST_DATA'): + test_data = DATASETS.build(self.cfg.TEST_DATA, logger=self.logger) + self.datas['test'] = test_data + self._epoch_max_iter['test'] = len(test_data.dataloader) + self._mode_set.add('test') + + def construct_model(self): + super().construct_model() + if self.tuner_cfg: + self.model = self.add_tuner(self.tuner_cfg, self.model) + if self.freeze_cfg: + freeze_cfg = Config.get_plain_cfg(self.freeze_cfg) + self.model = self.freeze(freeze_cfg, self.model) + if self.channels_last: + self.model = self.model.to(memory_format=torch.channels_last) + self.print_model_params_status() + + if we.debug: + module_keys = [key for key, _ in self.model.named_modules()] + self.logger.info(module_keys) + + def model_to_device(self): + self.model = self.model.to(we.device_id) + + def init_opti(self): + if hasattr(self.model, 'ignored_parameters'): + train_params, ignored_params = self.model.parameters( + ), self.model.ignored_parameters() + else: + train_params, ignored_params = self.model.parameters(), None + if we.is_distributed: + if self.use_fairscale: + from fairscale.nn.data_parallel import ShardedDataParallel + from fairscale.optim.oss import OSS + self.optimizer = OSS(params=train_params, + optim=torch.optim.AdamW, + lr=self.cfg.OPTIMIZER.LEARNING_RATE) + self.model = ShardedDataParallel(self.model, self.optimizer) + elif self.use_fsdp: + mixed_precision = MixedPrecision(param_dtype=self.dtype, + reduce_dtype=self.dtype, + buffer_dtype=self.dtype) + sharding_strategy = sharding_strategy_map[self.model_shard] + self.model = FullyShardedDataParallel( + self.model, + mixed_precision=mixed_precision, + cpu_offload=CPUOffload(offload_params=False), + sharding_strategy=sharding_strategy, + backward_prefetch=BackwardPrefetch.BACKWARD_PRE, + device_id=torch.cuda.current_device(), + ignored_parameters=ignored_params) + self.optimizer = OPTIMIZERS.build(self.cfg.OPTIMIZER, + logger=self.logger, + parameters=train_params) + else: + assert not self.model.use_ema + self.model = DistributedDataParallel( + self.model, + device_ids=[torch.cuda.current_device()], + output_device=torch.cuda.current_device(), + find_unused_parameters=True) + self.optimizer = OPTIMIZERS.build(self.cfg.OPTIMIZER, + logger=self.logger, + parameters=train_params) + else: + self.optimizer = OPTIMIZERS.build(self.cfg.OPTIMIZER, + logger=self.logger, + parameters=train_params) + + if self.cfg.have('LR_SCHEDULER') and self.optimizer is not None: + self.cfg.LR_SCHEDULER.TOTAL_STEPS = self.max_steps + self.lr_scheduler = LR_SCHEDULERS.build(self.cfg.LR_SCHEDULER, + logger=self.logger, + optimizer=self.optimizer) + + if self.cfg.DTYPE == 'float16': + if we.is_distributed: + if self.use_fairscale: + from fairscale.optim.grad_scaler import ShardedGradScaler + self.scaler = ShardedGradScaler(enabled=True) + elif self.use_fsdp: + from torch.distributed.fsdp.sharded_grad_scaler import \ + ShardedGradScaler + self.scaler = ShardedGradScaler() + else: + self.scaler = amp.GradScaler() + else: + self.scaler = amp.GradScaler() + else: + self.scaler = None + + def load_checkpoint(self, checkpoint: dict): + """ + Load checkpoint function + :param checkpoint: all tensors are on cpu, you need to transfer to gpu by hand + :return: + """ + if 'model' in checkpoint: + if hasattr(self.model, 'module'): + self.model.module.load_state_dict(checkpoint['model']) + else: + self.model.load_state_dict(checkpoint['model']) + else: + if hasattr(self.model, 'module'): + self.model.module.load_state_dict(checkpoint) + else: + self.model.load_state_dict(checkpoint) + if not self.load_model_only: + if 'optimizer' in checkpoint and self.optimizer: + self.optimizer.load_state_dict(checkpoint['optimizer']) + if 'scaler' in checkpoint and self.scaler: + self.scaler.load_state_dict(checkpoint['scaler']) + + def save_checkpoint(self) -> dict: + """ + Save checkpoint function, you need to transfer all tensors to cpu by hand + :return: + """ + ckpt = dict() + if we.is_distributed: + if self.use_fsdp: + save_policy = FullStateDictConfig(offload_to_cpu=True, + rank0_only=True) + with FullyShardedDataParallel.state_dict_type( + self.model, StateDictType.FULL_STATE_DICT, + save_policy): + ckpt['model'] = self.model.state_dict() + else: + if hasattr(self.model, 'module'): + ckpt['model'] = self.model.module.state_dict() + else: + ckpt['model'] = self.model.state_dict() + else: + ckpt['model'] = self.model.state_dict() + if self.optimizer and not self.use_fairscale: + ckpt['optimizer'] = self.optimizer.state_dict() + if self.scaler: + ckpt['scaler'] = self.scaler.state_dict() + return ckpt + + def solve(self): + self.before_solve() + if 'train' in self._mode_set: + self.run_train() + if 'test' in self._mode_set: + self.run_test() + self.after_solve() + + def run_train(self): + self.train_mode() + self.before_all_iter(self.hooks_dict[self._mode]) + data_iter = iter(self.datas[self._mode].dataloader) + self.print_memory_status() + for step in range(self.max_steps): + if 'eval' in self._mode_set and (step % self.eval_interval == 0 + or step == self.max_steps - 1): + self.run_eval() + self.train_mode() + self.before_iter(self.hooks_dict[self._mode]) + batch_data = next(data_iter) + if self.sample_args: + batch_data.update(self.sample_args.get_lowercase_dict()) + if 'meta' in batch_data: + self.register_probe({ + 'data_key': + ProbeData(batch_data['meta'].get('data_key', []), + view_distribute=True) + }) + self.register_probe({ + 'prompt': batch_data['prompt'], + 'batch_size': len(batch_data['prompt']) + }) + self.current_batch_data[self.mode] = batch_data + with torch.autocast(device_type='cuda', + enabled=self.use_amp, + dtype=self.dtype): + results = self.run_step_train( + transfer_data_to_cuda(batch_data), + step, + step=self.total_iter, + rank=we.rank) + self._iter_outputs[self._mode] = self._reduce_scalar(results) + self.after_iter(self.hooks_dict[self._mode]) + if we.debug: + self.print_trainable_params_status(prefix='model.') + self.after_all_iter(self.hooks_dict[self._mode]) + + @torch.no_grad() + def run_eval(self): + self.eval_mode() + self.before_all_iter(self.hooks_dict[self._mode]) + all_results = [] + for batch_idx, batch_data in tqdm( + enumerate(self.datas[self._mode].dataloader)): + self.before_iter(self.hooks_dict[self._mode]) + if self.sample_args: + batch_data.update(self.sample_args.get_lowercase_dict()) + with torch.autocast(device_type='cuda', + enabled=self.use_amp, + dtype=self.dtype): + results = self.run_step_eval(transfer_data_to_cuda(batch_data), + batch_idx, + step=self.total_iter, + rank=we.rank) + all_results.extend(results) + self.after_iter(self.hooks_dict[self._mode]) + log_data, log_label, ori_label = [], [], [] + for result in all_results: + # the inference image use + log_data.append((result['image'].permute(1, 2, 0).cpu().numpy() * + 255).astype(np.uint8)) + log_label.append(result['prompt'] + ' NegPrompt: ' + + result['n_prompt']) + ori_label.append(result['prompt']) + + self.register_probe({'test_label': log_label}) + self.register_probe({ + 'test_image': + ProbeData(log_data, + is_image=True, + build_html=True, + build_label=log_label) + }) + + log_data, log_label, ori_label = [], [], [] + for result in all_results: + # the inference image use + if 'train_n_image' in result: + log_data.append( + (result['train_n_image'].permute(1, 2, 0).cpu().numpy() * + 255).astype(np.uint8)) + log_label.append(result['prompt'] + 'NegPrompt' + + result['train_n_prompt']) + ori_label.append(result['prompt']) + if len(log_data) > 0: + self.register_probe({'test_train_n_label': log_label}) + self.register_probe({ + 'test_train_n_image': + ProbeData(log_data, + is_image=True, + build_html=True, + build_label=log_label) + }) + self.after_all_iter(self.hooks_dict[self._mode]) + + @torch.no_grad() + def run_test(self): + self.test_mode() + self.before_all_iter(self.hooks_dict[self._mode]) + all_results = [] + for batch_idx, batch_data in tqdm( + enumerate(self.datas[self._mode].dataloader)): + self.before_iter(self.hooks_dict[self._mode]) + if self.sample_args: + batch_data.update(self.sample_args.get_lowercase_dict()) + with torch.autocast(device_type='cuda', + enabled=self.use_amp, + dtype=self.dtype): + results = self.run_step_eval(transfer_data_to_cuda(batch_data), + batch_idx, + step=self.total_iter, + rank=we.rank) + all_results.extend(results) + self.after_iter(self.hooks_dict[self._mode]) + log_data, log_label = [], [] + for result in all_results: + # the inference image use + log_data.append((result['image'].permute(1, 2, 0).cpu().numpy() * + 255).astype(np.uint8)) + log_label.append(result['prompt'] + + " |NegPrompt| " + + result['n_prompt']) + + self.register_probe({'test_label': log_label}) + self.register_probe({ + 'test_image': + ProbeData(log_data, + is_image=True, + build_html=True, + build_label=log_label) + }) + + log_data, log_label = [], [] + for result in all_results: + # the inference image use + if 'train_n_image' in result: + log_data.append( + (result['train_n_image'].permute(1, 2, 0).cpu().numpy() * + 255).astype(np.uint8)) + log_label.append(result['prompt'] + + " |NegPrompt| " + + result['train_n_prompt']) + if len(log_data) > 0: + self.register_probe({'test_train_n_label': log_label}) + self.register_probe({ + 'test_train_n_image': + ProbeData(log_data, + is_image=True, + build_html=True, + build_label=log_label) + }) + + self.after_all_iter(self.hooks_dict[self._mode]) + + def add_tuner(self, tuner_cfg, model=None): + from scepter.modules.model.registry import TUNERS + + if model is None: + model = self.model + swift_cfg_dict = {} + for t_id, t_cfg in enumerate(tuner_cfg): + cfg_name = t_cfg['NAME'] + init_config = TUNERS.build(t_cfg, logger=self.logger)() + if init_config is None: + continue + swift_cfg_dict[f'{t_id}_{cfg_name}'] = init_config + if len(swift_cfg_dict) > 0: + from swift import Swift + model = Swift.prepare_model(self.model, config=swift_cfg_dict) + return model + + def freeze(self, freeze_cfg, model=None): + """ Freeze or train the model based on the config. + """ + if model is None: + model = self.model + freeze_part = freeze_cfg[ + 'FREEZE_PART'] if 'FREEZE_PART' in freeze_cfg else [] + train_part = freeze_cfg[ + 'TRAIN_PART'] if 'TRAIN_PART' in freeze_cfg else [] + + if hasattr(model, 'module'): + freeze_model = model.module + else: + freeze_model = model + + if freeze_part: + if isinstance(freeze_part, dict): + if 'BACKBONE' in freeze_part: + part = freeze_part['BACKBONE'] + for name, param in freeze_model.backbone.named_parameters( + ): + freeze_flag = sum([p in name for p in part]) > 0 + if freeze_flag: + param.requires_grad = False + elif 'HEAD' in freeze_part: + part = freeze_part['HEAD'] + for name, param in freeze_model.head.named_parameters(): + freeze_flag = sum([p in name for p in part]) > 0 + if freeze_flag: + param.requires_grad = False + elif isinstance(freeze_part, list): + for name, param in freeze_model.named_parameters(): + freeze_flag = sum([p in name for p in freeze_part]) > 0 + if freeze_flag: + param.requires_grad = False + if train_part: + if isinstance(train_part, dict): + if 'BACKBONE' in train_part: + part = train_part['BACKBONE'] + for name, param in freeze_model.backbone.named_parameters( + ): + freeze_flag = sum([p in name for p in part]) > 0 + if freeze_flag: + param.requires_grad = True + elif 'HEAD' in train_part: + part = train_part['HEAD'] + for name, param in freeze_model.head.named_parameters(): + freeze_flag = sum([p in name for p in part]) > 0 + if freeze_flag: + param.requires_grad = True + elif isinstance(train_part, list): + for name, param in freeze_model.named_parameters(): + freeze_flag = sum([p in name for p in train_part]) > 0 + if freeze_flag: + param.requires_grad = True + return model + + @torch.no_grad() + def log_image(self, batch_data): + self.eval_mode() + if we.is_distributed: + if hasattr(self.model, 'module'): + images = self.model.module.log_images(**batch_data) + else: + images = self.model.log_images(**batch_data) + else: + images = self.model.log_images(**batch_data) + self.train_mode() + return images + + @staticmethod + def get_config_template(): + return dict_to_yaml('solvername', + __class__.__name__, + LatentDiffusionSolver.para_dict, + set_name=True) + + @property + def image_out(self): + return self._image_out[self._mode] + + def collect_log_vars(self) -> OrderedDict: + ret = OrderedDict() + if self.is_train_mode and self.optimizer is not None: + for idx, pg in enumerate(self.optimizer.param_groups): + ret[f'pg{idx}_lr'] = pg['lr'] + if self.is_train_mode and self.scaler is not None: + ret['scale'] = self.scaler.get_scale() + return ret + + @property + def probe_data(self): + if not we.debug and self.mode == 'train': + with torch.autocast(device_type='cuda', + enabled=self.use_amp, + dtype=self.dtype): + outputs = self.log_image( + transfer_data_to_cuda(self.current_batch_data[self.mode])) + log_data, log_label = [], [] + for result in outputs: + merge_image = torch.cat([result['orig'], result['recon']], + dim=2) + log_data.append((merge_image.permute(1, 2, 0).cpu().numpy() * + 255).astype(np.uint8)) + log_label.append('recon image: ' + result['prompt'] + + " |NegPrompt| " + + result['n_prompt']) + + self.register_probe({ + 'train_image': + ProbeData(log_data, + is_image=True, + build_html=True, + build_label=log_label) + }) + self.register_probe({'train_label': log_label}) + + # the inference image use + log_data, log_label = [], [] + for result in outputs: + if 'train_n_image' in result: + merge_image = torch.cat( + [result['orig'], result['train_n_image']], dim=2) + log_data.append( + (merge_image.permute(1, 2, 0).cpu().numpy() * + 255).astype(np.uint8)) + log_label.append( + 'recon image: ' + result['prompt'] + + " |NegPrompt| " + + result['train_n_prompt']) + + if len(log_data) > 0: + self.register_probe({'train_n_label': log_label}) + self.register_probe({ + 'train_n_image': + ProbeData(log_data, + is_image=True, + build_html=True, + build_label=log_label) + }) + return super().probe_data + + def print_memory_status(self): + """Print the memory usage status of the model""" + if torch.cuda.is_available(): + nvi_info = os.popen('nvidia-smi').read() + gpu_mem = nvi_info.split('\n')[9].split('|')[2].split( + '/')[0].strip() + else: + gpu_mem = '' + return gpu_mem + + def print_trainable_params_status(self, + model=None, + logger=None, + prefix=''): + """Print the status and parameters of the model""" + if model is None: + model = self.model + if logger is None: + logger = self.logger + + for key, val in model.named_parameters(): + if val.requires_grad: + if prefix in key: + logger.info( + f"param {key} value'sum {torch.sum(val)} with shape {val.shape}." + ) + + def print_model_params_status(self, model=None, logger=None): + """Print the status and parameters of the model""" + if model is None: + model = self.model + if logger is None: + logger = self.logger + train_param_dict = {} + forzen_param_dict = {} + all_param_numel = 0 + if we.debug: + for key, _ in model.named_modules(): + logger.info(f'sub modules {key}.') + for key, val in model.named_parameters(): + if val.requires_grad: + sub_key = '.'.join(key.split('.', 1)[-1].split('.', 2)[:2]) + if sub_key in train_param_dict: + train_param_dict[sub_key] += val.numel() + else: + train_param_dict[sub_key] = val.numel() + else: + sub_key = '.'.join(key.split('.', 1)[-1].split('.', 1)[:1]) + if sub_key in forzen_param_dict: + forzen_param_dict[sub_key] += val.numel() + else: + forzen_param_dict[sub_key] = val.numel() + all_param_numel += val.numel() + if we.debug: + logger.info(key) + train_param_numel = sum(train_param_dict.values()) + forzen_param_numel = sum(forzen_param_dict.values()) + logger.info( + f'Load trainable params {train_param_numel} / {all_param_numel} = ' + f'{train_param_numel / all_param_numel:.2%}, ' + f'train part: {train_param_dict}.') + logger.info( + f'Load forzen params {forzen_param_numel} / {all_param_numel} = ' + f'{forzen_param_numel / all_param_numel:.2%}, ' + f'forzen part: {forzen_param_dict}.') diff --git a/scepter/modules/solver/hooks/__init__.py b/scepter/modules/solver/hooks/__init__.py new file mode 100644 index 0000000..f2fca74 --- /dev/null +++ b/scepter/modules/solver/hooks/__init__.py @@ -0,0 +1,51 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.solver.hooks.backward import BackwardHook +from scepter.modules.solver.hooks.checkpoint import CheckpointHook +from scepter.modules.solver.hooks.data_probe import ProbeDataHook +from scepter.modules.solver.hooks.hook import Hook +from scepter.modules.solver.hooks.log import LogHook, TensorboardLogHook +from scepter.modules.solver.hooks.lr import LrHook +from scepter.modules.solver.hooks.registry import HOOKS +from scepter.modules.solver.hooks.safetensors import SafetensorsHook +from scepter.modules.solver.hooks.sampler import DistSamplerHook +""" +Normally, hooks have priorities, below we recommend priority that runs fine (low score MEANS high priority) +BackwardHook: 0 +LogHook: 100 +LrHook: 200 +CheckpointHook: 300 +SamplerHook: 400 + +Recommend sequences in training are: +before solve: + TensorboardLogHook: prepare file handler + CheckpointHook: resume checkpoint + +before epoch: + LogHook: clear epoch variables + DistSamplerHook: change sampler seed + +before iter: + LogHook: record data time + +after iter: + BackwardHook: network backward + LogHook: log + TensorboardLogHook: log + CheckpointHook: save checkpoint + SafetensorsHook: save checkpoint + +after epoch: + LrHook: reset learning rate + CheckpointHook: save checkpoint + +after solve: + TensorboardLogHook: close file handler +""" + +__all__ = [ + 'HOOKS', 'BackwardHook', 'CheckpointHook', 'Hook', 'LrHook', 'LogHook', + 'TensorboardLogHook', 'DistSamplerHook', 'ProbeDataHook', 'SafetensorsHook' +] diff --git a/scepter/modules/solver/hooks/backward.py b/scepter/modules/solver/hooks/backward.py new file mode 100644 index 0000000..08e5e8e --- /dev/null +++ b/scepter/modules/solver/hooks/backward.py @@ -0,0 +1,93 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import warnings + +import torch + +from scepter.modules.solver.hooks.hook import Hook +from scepter.modules.solver.hooks.registry import HOOKS +from scepter.modules.utils.config import dict_to_yaml + +_DEFAULT_BACKWARD_PRIORITY = 0 + + +@HOOKS.register_class() +class BackwardHook(Hook): + para_dict = [{ + 'PRIORITY': { + 'value': _DEFAULT_BACKWARD_PRIORITY, + 'description': 'the priority for processing!' + }, + 'GRADIENT_CLIP': { + 'value': -1, + 'description': 'the gradient clip max_norm value for parameters!' + }, + 'ACCUMULATE_STEP': { + 'value': + 1, + 'description': + 'the gradient accumulate steps for backward step, default is 1' + }, + 'EMPTY_CACHE_STEP': { + 'value': -1, + 'description': 'the memory empty step!' + } + }] + + def __init__(self, cfg, logger=None): + super(BackwardHook, self).__init__(cfg, logger=logger) + self.priority = cfg.get('PRIORITY', _DEFAULT_BACKWARD_PRIORITY) + self.gradient_clip = cfg.get('GRADIENT_CLIP', -1) + self.empty_cache_step = cfg.get('EMPTY_CACHE_STEP', -1) + self.accumulate_step = cfg.get('ACCUMULATE_STEP', 1) + self.current_step = 0 + + def grad_clip(self, parameters): + torch.nn.utils.clip_grad_norm_(parameters=parameters, + max_norm=self.gradient_clip, + norm_type=2) + + def after_iter(self, solver): + if (hasattr(solver, 'use_fsdp') + and solver.use_fsdp) and self.accumulate_step > 1: + self.logger.info("Fsdp don't surpport gradient accumulate.") + self.accumulate_step = 1 + if solver.optimizer is not None and solver.is_train_mode: + if solver.loss is None: + warnings.warn( + 'solver.loss should not be None in train mode, remember to call solver._reduce_scalar()!' + ) + return + if solver.scaler is not None: + solver.scaler.scale(solver.loss).backward() + if self.gradient_clip > 0: + solver.scaler.unscale_(solver.optimizer) + self.grad_clip(solver.train_parameters()) + self.current_step += 1 + if self.current_step % self.accumulate_step == 0: + solver.scaler.step(solver.optimizer) + solver.scaler.update() + solver.optimizer.zero_grad() + else: + solver.loss.backward() + if self.gradient_clip > 0: + self.grad_clip(solver.train_parameters()) + self.current_step += 1 + if self.current_step % self.accumulate_step == 0: + solver.optimizer.step() + solver.optimizer.zero_grad() + if solver.lr_scheduler: + if self.current_step % self.accumulate_step == 0: + solver.lr_scheduler.step() + if self.current_step % self.accumulate_step == 0: + self.current_step = 0 + solver.loss = None + if self.empty_cache_step > 0 and solver.total_iter % self.empty_cache_step == 0: + torch.cuda.empty_cache() + + @staticmethod + def get_config_template(): + return dict_to_yaml('hook', + __class__.__name__, + BackwardHook.para_dict, + set_name=True) diff --git a/scepter/modules/solver/hooks/checkpoint.py b/scepter/modules/solver/hooks/checkpoint.py new file mode 100644 index 0000000..ec38d81 --- /dev/null +++ b/scepter/modules/solver/hooks/checkpoint.py @@ -0,0 +1,181 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import os.path as osp +import sys +import warnings + +import torch +import torch.distributed as du + +from scepter.modules.solver.hooks.hook import Hook +from scepter.modules.solver.hooks.registry import HOOKS +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we +from scepter.modules.utils.file_system import FS + +_DEFAULT_CHECKPOINT_PRIORITY = 300 + + +@HOOKS.register_class() +class CheckpointHook(Hook): + """ Checkpoint resume or save hook. + Args: + interval (int): Save interval, by epoch. + save_best (bool): Save the best checkpoint by a metric key, default is False. + save_best_by (str): How to get the best the checkpoint by the metric key, default is ''. + + means the higher the best (default). + - means the lower the best. + E.g. +acc@1, -err@1, acc@5(same as +acc@5) + """ + para_dict = [{ + 'PRIORITY': { + 'value': _DEFAULT_CHECKPOINT_PRIORITY, + 'description': 'the priority for processing!' + }, + 'INTERVAL': { + 'value': 1, + 'description': 'the interval of saving checkpoint!' + }, + 'SAVE_BEST': { + 'value': False, + 'description': 'If save the best model or not!' + }, + 'SAVE_BEST_BY': { + 'value': + '', + 'description': + 'If save the best model, which order should be sorted, +/-!' + } + }] + + def __init__(self, cfg, logger=None): + super(CheckpointHook, self).__init__(cfg, logger=logger) + self.priority = cfg.get('PRIORITY', _DEFAULT_CHECKPOINT_PRIORITY) + self.interval = cfg.get('INTERVAL', 1) + self.save_name_prefix = cfg.get('SAVE_NAME_PREFIX', 'ldm_step') + self.save_last = cfg.get('SAVE_LAST', False) + self.save_best = cfg.get('SAVE_BEST', False) + self.save_best_by = cfg.get('SAVE_BEST_BY', '') + if self.save_best and not self.save_best_by: + warnings.warn( + "CheckpointHook: Parameter 'save_best_by' is not set, turn off save_best function." + ) + self.save_best = False + self.higher_the_best = True + if self.save_best: + if self.save_best_by.startswith('+'): + self.save_best_by = self.save_best_by[1:] + elif self.save_best_by.startswith('-'): + self.save_best_by = self.save_best_by[1:] + self.higher_the_best = False + if self.save_best and not self.save_best_by: + warnings.warn( + "CheckpointHook: Parameter 'save_best_by' is not valid, turn off save_best function." + ) + self.save_best = False + self._last_best = None if not self.save_best else ( + sys.float_info.min if self.higher_the_best else sys.float_info.max) + + def before_solve(self, solver): + if solver.resume_from is None: + return + if not FS.exists(solver.resume_from): + solver.logger.error(f'File not exists {solver.resume_from}') + return + + with FS.get_from(solver.resume_from, wait_finish=True) as local_file: + solver.logger.info(f'Loading checkpoint from {solver.resume_from}') + checkpoint = torch.load(local_file, + map_location=torch.device('cpu')) + + solver.load_checkpoint(checkpoint) + if self.save_best and '_CheckpointHook_best' in checkpoint: + self._last_best = checkpoint['_CheckpointHook_best'] + + def after_iter(self, solver): + if solver.total_iter != 0 and ( + (solver.total_iter + 1) % self.interval == 0 + or solver.total_iter == solver.max_steps - 1): + checkpoint = solver.save_checkpoint() + solver.logger.info( + f'Saving checkpoint after {solver.total_iter + 1} steps') + if we.rank == 0: + save_path = osp.join( + solver.work_dir, + 'checkpoints/{}-{}.pth'.format(self.save_name_prefix, + solver.total_iter + 1)) + with FS.put_to(save_path) as local_path: + with open(local_path, 'wb') as f: + torch.save(checkpoint, f) + if self.save_last and solver.total_iter == solver.max_steps - 1: + with FS.get_fs_client(save_path) as client: + last_path = osp.join(solver.work_dir, 'checkpoint.pth') + client.make_link(last_path, save_path) + + torch.cuda.synchronize() + if we.is_distributed: + torch.distributed.barrier() + + def after_epoch(self, solver): + if du.is_available() and du.is_initialized() and du.get_rank() != 0: + return + if (solver.epoch + 1) % self.interval == 0: + solver.logger.info( + f'Saving checkpoint after {solver.epoch} epochs') + checkpoint = solver.save_checkpoint() + if checkpoint is None or len(checkpoint) == 0: + return + cur_is_best = False + if self.save_best: + # Try to get current state from epoch_outputs["eval"] + cur_state = None \ + if self.save_best_by not in solver.epoch_outputs['eval'] \ + else solver.epoch_outputs['eval'][self.save_best_by] + # Try to get current state from agg_iter_outputs["eval"] if do_final_eval is False + if cur_state is None: + cur_state = None \ + if self.save_best_by not in solver.agg_iter_outputs['eval'] \ + else solver.agg_iter_outputs['eval'][self.save_best_by] + # Try to get current state from agg_iter_outputs["train"] if no evaluation + if cur_state is None: + cur_state = None \ + if self.save_best_by not in solver.agg_iter_outputs['train'] \ + else solver.agg_iter_outputs['train'][self.save_best_by] + if cur_state is not None: + if self.higher_the_best and cur_state > self._last_best: + self._last_best = cur_state + cur_is_best = True + elif not self.higher_the_best and cur_state < self._last_best: + self._last_best = cur_state + cur_is_best = True + checkpoint['_CheckpointHook_best'] = self._last_best + # minus 1, means index + save_path = osp.join(solver.work_dir, + 'epoch-{:05d}.pth'.format(solver.epoch)) + + with FS.get_fs_client(save_path) as client: + local_file = client.convert_to_local_path(save_path) + with open(local_file, 'wb') as f: + torch.save(checkpoint, f) + client.put_object_from_local_file(local_file, save_path) + + if cur_is_best: + best_path = osp.join(solver.work_dir, 'best.pth') + client.make_link(best_path, save_path) + # save pretrain checkout + if 'pre_state_dict' in checkpoint: + save_path = osp.join( + solver.work_dir, + 'epoch-{:05d}_pretrain.pth'.format(solver.epoch)) + with FS.get_fs_client(save_path) as client: + local_file = client.convert_to_local_path(save_path) + with open(local_file, 'wb') as f: + torch.save(checkpoint['pre_state_dict'], f) + client.put_object_from_local_file(local_file, save_path) + + @staticmethod + def get_config_template(): + return dict_to_yaml('hook', + __class__.__name__, + CheckpointHook.para_dict, + set_name=True) diff --git a/scepter/modules/solver/hooks/data_probe.py b/scepter/modules/solver/hooks/data_probe.py new file mode 100644 index 0000000..a5b9657 --- /dev/null +++ b/scepter/modules/solver/hooks/data_probe.py @@ -0,0 +1,99 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import json +import os + +import torch + +from scepter.modules.solver.hooks.hook import Hook +from scepter.modules.solver.hooks.registry import HOOKS +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import barrier, we +from scepter.modules.utils.file_system import FS + +_DEFAULT_PROBE_PRIORITY = 1000 + + +@HOOKS.register_class() +class ProbeDataHook(Hook): + para_dict = [{ + 'PRIORITY': { + 'value': _DEFAULT_PROBE_PRIORITY, + 'description': 'The priority for processing!' + }, + 'PROB_INTERVAL': { + 'value': 1000, + 'description': 'the interval for log print!' + } + }] + + def __init__(self, cfg, logger=None): + super(ProbeDataHook, self).__init__(cfg, logger=logger) + self.priority = cfg.get('PRIORITY', _DEFAULT_PROBE_PRIORITY) + self.log_interval = cfg.get('PROB_INTERVAL', 1000) + + def before_all_iter(self, solver): + pass + + def before_iter(self, solver): + pass + + def after_iter(self, solver): + if solver.mode == 'train' and solver.total_iter % self.log_interval == 0: + probe_dict = solver.probe_data + if we.rank == 0: + save_folder = os.path.join( + solver.work_dir, + f'{solver.mode}_probe/step_{solver.total_iter}') + ret_data = {} + for k, v in probe_dict.items(): + ret_one = v.to_log( + os.path.join( + save_folder, + k.replace('/', '_') + + f'_step_{solver.total_iter}')) + if (isinstance(ret_one, list) + or isinstance(ret_one, dict)) and len(ret_one) < 1: + continue + ret_data[k] = ret_one + with FS.put_to(os.path.join(save_folder, + 'meta.json')) as local_path: + json.dump(ret_data, + open(local_path, 'w'), + ensure_ascii=False) + solver.clear_probe() + torch.cuda.synchronize() + barrier() + + def after_all_iter(self, solver): + if not solver.mode == 'train': + probe_dict = solver.probe_data + if we.rank == 0: + step = solver._total_iter[ + 'train'] if 'train' in solver._total_iter else 0 + save_folder = os.path.join(solver.work_dir, + f'{solver.mode}_probe/step_{step}') + ret_data = {} + for k, v in probe_dict.items(): + ret_one = v.to_log( + os.path.join(save_folder, + k.replace('/', '_') + f'_step_{step}')) + if (isinstance(ret_one, list) + or isinstance(ret_one, dict)) and len(ret_one) < 1: + continue + ret_data[k] = ret_one + with FS.put_to(os.path.join(save_folder, + 'meta.json')) as local_path: + json.dump(ret_data, + open(local_path, 'w'), + ensure_ascii=False) + solver.clear_probe() + torch.cuda.synchronize() + barrier() + + @staticmethod + def get_config_template(): + return dict_to_yaml('HOOK', + __class__.__name__, + ProbeDataHook.para_dict, + set_name=True) diff --git a/scepter/modules/solver/hooks/hook.py b/scepter/modules/solver/hooks/hook.py new file mode 100644 index 0000000..ef68ab4 --- /dev/null +++ b/scepter/modules/solver/hooks/hook.py @@ -0,0 +1,33 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from abc import ABCMeta + + +class Hook(object, metaclass=ABCMeta): + def __init__(self, cfg, logger=None): + self.logger = logger + + def before_solve(self, solver): + pass + + def after_solve(self, solver): + pass + + def before_epoch(self, solver): + pass + + def after_epoch(self, solver): + pass + + def before_all_iter(self, solver): + pass + + def before_iter(self, solver): + pass + + def after_iter(self, solver): + pass + + def after_all_iter(self, solver): + pass diff --git a/scepter/modules/solver/hooks/log.py b/scepter/modules/solver/hooks/log.py new file mode 100644 index 0000000..4ea3b4c --- /dev/null +++ b/scepter/modules/solver/hooks/log.py @@ -0,0 +1,312 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import numbers +import os +import os.path as osp +import time +import warnings +from collections import defaultdict +from typing import Optional + +import numpy as np +import torch + +from scepter.modules.solver.hooks.hook import Hook +from scepter.modules.solver.hooks.registry import HOOKS +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we +from scepter.modules.utils.file_system import FS +from scepter.modules.utils.logger import LogAgg, time_since + +try: + from torch.utils.tensorboard import SummaryWriter +except Exception as e: + warnings.warn(f'Runing without tensorboard! {e}') + +_DEFAULT_LOG_PRIORITY = 100 + + +def _format_float(x): + try: + if abs(x) - int(abs(x)) < 0.01: + return '{:.6f}'.format(x) + else: + return '{:.4f}'.format(x) + except Exception: + return 'NaN' + + +def _print_v(x): + if isinstance(x, float): + return _format_float(x) + elif isinstance(x, torch.Tensor) and x.ndim == 0: + return _print_v(x.item()) + else: + return f'{x}' + + +def _print_iter_log(solver, outputs, final=False, start_time=0, mode=None): + extra_vars = solver.collect_log_vars() + outputs.update(extra_vars) + s = [] + for k, v in outputs.items(): + if k in ('data_time', 'time'): + continue + if isinstance(v, (list, tuple)) and len(v) == 2: + s.append(f'{k}: ' + _print_v(v[0]) + f'({_print_v(v[1])})') + else: + s.append(f'{k}: ' + _print_v(v)) + if 'time' in outputs: + v = outputs['time'] + s.insert(0, 'time: ' + _print_v(v[0]) + f'({_print_v(v[1])})') + if 'data_time' in outputs: + v = outputs['data_time'] + s.insert(0, 'data_time: ' + _print_v(v[0]) + f'({_print_v(v[1])})') + + if solver.max_epochs == -1: + assert solver.max_steps > 0 + percent = (solver.total_iter + + 1 if not final else solver.total_iter) / solver.max_steps + now_status = time_since(start_time, percent) + solver.logger.info( + f'Stage [{mode}] ' + f'iter: [{solver.total_iter + 1 if not final else solver.total_iter}/{solver.max_steps}], ' + f"{', '.join(s)}, " + f'[{now_status}]') + else: + assert solver.max_epochs > 0 and solver.epoch_max_iter > 0 + percent = (solver.total_iter + 1 if not final else solver.total_iter + ) / (solver.epoch_max_iter * solver.max_epochs) + now_status = time_since(start_time, percent) + solver.logger.info( + f'Epoch [{solver.epoch}/{solver.max_epochs}], stage [{mode}] ' + f'iter: [{solver.total_iter + 1 if not final else solver.total_iter}/{solver.epoch_max_iter * solver.max_epochs}], ' # noqa + f'iter: [{solver.iter + 1 if not final else solver.iter}/{solver.epoch_max_iter}], ' + f"{', '.join(s)}, " + f'[{now_status}]') + + +def print_memory_status(): + if torch.cuda.is_available(): + nvi_info = os.popen('nvidia-smi').read() + gpu_mem = nvi_info.split('\n')[9].split('|')[2].split('/')[0].strip() + gpu_mem = int(gpu_mem.replace('MiB', '')) + else: + gpu_mem = 0 + return gpu_mem + + +@HOOKS.register_class() +class LogHook(Hook): + para_dict = [{ + 'PRIORITY': { + 'value': _DEFAULT_LOG_PRIORITY, + 'description': 'the priority for processing!' + }, + 'LOG_INTERVAL': { + 'value': 10, + 'description': 'the interval for log print!' + }, + 'SHOW_GPU_MEM': { + 'value': False, + 'description': 'to show the gpu memory' + } + }] + + def __init__(self, cfg, logger=None): + super(LogHook, self).__init__(cfg, logger=logger) + self.priority = cfg.get('PRIORITY', _DEFAULT_LOG_PRIORITY) + self.log_interval = cfg.get('LOG_INTERVAL', 10) + self.show_gpu_mem = cfg.get('SHOW_GPU_MEM', False) + self.log_agg_dict = defaultdict(LogAgg) + + self.last_log_step = ('train', 0) + + self.time = time.time() + self.start_time = time.time() + self.data_time = 0 + + def before_all_iter(self, solver): + self.time = time.time() + self.last_log_step = (solver.mode, 0) + + def before_iter(self, solver): + data_time = time.time() - self.time + self.data_time = data_time + + def after_iter(self, solver): + log_agg = self.log_agg_dict[solver.mode] + iter_time = time.time() - self.time + self.time = time.time() + outputs = solver.iter_outputs.copy() + outputs['time'] = iter_time + outputs['data_time'] = self.data_time + if 'batch_size' in outputs: + batch_size = outputs.pop('batch_size') + else: + batch_size = 1 + if self.show_gpu_mem: + outputs['nvidia-smi'] = print_memory_status() + log_agg.update(outputs, batch_size) + if (solver.iter + 1) % self.log_interval == 0: + _print_iter_log(solver, + log_agg.aggregate(self.log_interval), + start_time=self.start_time, + mode=solver.mode) + self.last_log_step = (solver.mode, solver.iter + 1) + + def after_all_iter(self, solver): + outputs = self.log_agg_dict[solver.mode].aggregate( + solver.iter - self.last_log_step[1]) + solver.agg_iter_outputs = { + key: value[1] + for key, value in outputs.items() + } + current_log_step = (solver.mode, solver.iter) + if current_log_step != self.last_log_step: + _print_iter_log(solver, + outputs, + final=True, + start_time=self.start_time, + mode=solver.mode) + self.last_log_step = current_log_step + + for _, value in self.log_agg_dict.items(): + value.reset() + + def after_epoch(self, solver): + outputs = solver.epoch_outputs + mode_s = [] + for mode_name, kvs in outputs.items(): + if len(kvs) == 0: + return + s = [f'{k}: ' + _print_v(v) for k, v in kvs.items()] + mode_s.append(f"{mode_name} -> {', '.join(s)}") + if len(mode_s) > 1: + states = '\n\t'.join(mode_s) + solver.logger.info( + f'Epoch [{solver.epoch}/{solver.max_epochs}], \n\t' + f'{states}') + elif len(mode_s) == 1: + solver.logger.info( + f'Epoch [{solver.epoch}/{solver.max_epochs}], {mode_s[0]}') + # summary + + for mode in self.log_agg_dict: + solver.logger.info(f'Current Epoch {mode} Summary:') + log_agg = self.log_agg_dict[mode] + _print_iter_log(solver, + log_agg.aggregate(self.log_interval), + start_time=self.start_time, + mode=mode) + if not mode == 'train': + self.log_agg_dict[mode].reset() + + @staticmethod + def get_config_template(): + return dict_to_yaml('HOOK', + __class__.__name__, + LogHook.para_dict, + set_name=True) + + +@HOOKS.register_class() +class TensorboardLogHook(Hook): + para_dict = [{ + 'PRIORITY': { + 'value': _DEFAULT_LOG_PRIORITY, + 'description': 'the priority for processing!' + }, + 'LOG_DIR': { + 'value': None, + 'description': 'the dir for tensorboard log!' + }, + 'LOG_INTERVAL': { + 'value': 10000, + 'description': 'the interval for log upload!' + } + }] + + def __init__(self, cfg, logger=None): + super(TensorboardLogHook, self).__init__(cfg, logger=logger) + self.priority = cfg.get('PRIORITY', _DEFAULT_LOG_PRIORITY) + self.log_dir = cfg.get('LOG_DIR', None) + self.log_interval = cfg.get('LOG_INTERVAL', 1000) + self._local_log_dir = None + self.writer: Optional[SummaryWriter] = None + + def before_solve(self, solver): + if we.rank != 0: + return + + if self.log_dir is None: + self.log_dir = osp.join(solver.work_dir, 'tensorboard') + + self._local_log_dir, _ = FS.map_to_local(self.log_dir) + os.makedirs(self._local_log_dir, exist_ok=True) + self.writer = SummaryWriter(self._local_log_dir) + solver.logger.info(f'Tensorboard: save to {self.log_dir}') + + def after_iter(self, solver): + if self.writer is None: + return + outputs = solver.iter_outputs.copy() + extra_vars = solver.collect_log_vars() + outputs.update(extra_vars) + mode = solver.mode + for key, value in outputs.items(): + if key == 'batch_size': + continue + if isinstance(value, torch.Tensor): + # Must be scalar + if not value.ndim == 0: + continue + value = value.item() + elif isinstance(value, np.ndarray): + # Must be scalar + if not value.ndim == 0: + continue + value = float(value) + elif isinstance(value, numbers.Number): + # Must be number + pass + else: + continue + + self.writer.add_scalar(f'{mode}/iter/{key}', + value, + global_step=solver.total_iter) + if solver.total_iter % self.log_interval: + self.writer.flush() + # Put to remote file systems every epoch + FS.put_dir_from_local_dir(self._local_log_dir, self.log_dir) + + def after_epoch(self, solver): + if self.writer is None: + return + outputs = solver.epoch_outputs.copy() + for mode, kvs in outputs.items(): + for key, value in kvs.items(): + self.writer.add_scalar(f'{mode}/epoch/{key}', + value, + global_step=solver.epoch) + + self.writer.flush() + # Put to remote file systems every epoch + FS.put_dir_from_local_dir(self._local_log_dir, self.log_dir) + + def after_solve(self, solver): + if self.writer is None: + return + if self.writer: + self.writer.close() + + FS.put_dir_from_local_dir(self._local_log_dir, self.log_dir) + + @staticmethod + def get_config_template(): + return dict_to_yaml('HOOK', + __class__.__name__, + TensorboardLogHook.para_dict, + set_name=True) diff --git a/scepter/modules/solver/hooks/lr.py b/scepter/modules/solver/hooks/lr.py new file mode 100644 index 0000000..52d1b16 --- /dev/null +++ b/scepter/modules/solver/hooks/lr.py @@ -0,0 +1,131 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from ...utils.config import dict_to_yaml +from .hook import Hook +from .registry import HOOKS + +_DEFAULT_LR_PRIORITY = 200 + + +def _get_lr_from_scheduler(lr_scheduler, cur_epoch): + """Ugly solution to get lr by epoch. + PyTorch lr scheduler get_lr() function is recommended to call in step() + Here we mock the environment. + + Args: + lr_scheduler (torch.optim.lr_scheduler._LRScheduler): + cur_epoch (number): int or float (when num_folds > 1) + + Returns: + Learning rate at cur_epoch. + """ + lr_scheduler._get_lr_called_within_step = True + last_epoch_bk = lr_scheduler.last_epoch + lr_scheduler.last_epoch = cur_epoch + if hasattr(lr_scheduler, '_get_closed_form_lr'): + lr = lr_scheduler._get_closed_form_lr()[0] + else: + lr = lr_scheduler.get_lr()[0] + lr_scheduler._get_lr_called_within_step = False + lr_scheduler.last_epoch = last_epoch_bk + return lr + + +@HOOKS.register_class() +class LrHook(Hook): + """ Learning rate updater hook. + If warmup, warmup_end_lr will be calculated by lr_scheduler at warmup_epochs. + Lr in warmup period is set based on warmup_func. + If set_by_epoch, lr is set at end of epoch. Otherwise, lr is set before training iteration. + + Args: + set_by_epoch (bool): Reset learning rate by epoch, we recommend true if solver.num_folds == 1 + warmup_func (str, None): Do not warm up if None, currently support linear warmup + warmup_epochs (int): + warmup_start_lr (float): + """ + para_dict = [{ + 'PRIORITY': { + 'value': _DEFAULT_LR_PRIORITY, + 'description': 'the priority for processing!' + }, + 'WARMUP_FUNC': { + 'value': 'linear', + 'description': 'Only linear warmup supported!' + }, + 'WARMUP_EPOCHS': { + 'value': 1, + 'description': 'The warmup epochs!' + }, + 'WARMUP_START_LR': { + 'value': 0.0001, + 'description': 'The warmup start learning rate!' + }, + 'SET_BY_EPOCH': { + 'value': True, + 'description': 'Set the learning rate by epoch!' + } + }] + + def __init__(self, cfg, logger=None): + super(LrHook, self).__init__(cfg, logger=logger) + self.priority = cfg.get('PRIORITY', _DEFAULT_LR_PRIORITY) + self.warmup_func = cfg.get('WARMUP_FUNC', 'linear') + if self.warmup_func is not None: + assert self.warmup_func in ( + 'linear', ), 'Only linear warmup supported' + self.warmup_epochs = cfg.get('WARMUP_EPOCHS', 1) + self.warmup_start_lr = cfg.get('WARMUP_START_LR', 0.0001) + self.warmup_end_lr = 0 + self.set_by_epoch = cfg.get('SET_BY_EPOCH', True) + + def before_solve(self, solver): + if self.warmup_func is not None and self.warmup_epochs > 0: + self.warmup_end_lr = _get_lr_from_scheduler( + solver.lr_scheduler, self.warmup_epochs) + for param_group in solver.optimizer.param_groups: + param_group['lr'] = self.warmup_start_lr + + def _get_warmup_lr(self, cur_epoch): + # if self.warmup_func == "linear": + alpha = (self.warmup_end_lr - + self.warmup_start_lr) / self.warmup_epochs + return self.warmup_start_lr + alpha * cur_epoch + + def after_epoch(self, solver): + if solver.lr_scheduler is not None and solver.is_train_mode: + if self.set_by_epoch: + last_lr = solver.optimizer.param_groups[0]['lr'] + for _ in range(solver.num_folds): + solver.lr_scheduler.step() + new_lr = solver.optimizer.param_groups[0]['lr'] + print(f'now {new_lr}') + if self.warmup_func is not None and solver.epoch < self.warmup_epochs: + new_lr = self._get_warmup_lr(solver.epoch) + for param_group in solver.optimizer.param_groups: + param_group['lr'] = new_lr + if last_lr != new_lr: + solver.logger.info( + f'Change learning rate from {last_lr} to {new_lr}') + else: + solver.logger.info(f'Keep learning rate = {last_lr}') + + def before_iter(self, solver): + if not self.set_by_epoch and solver.is_train_mode and solver.lr_scheduler is not None: + cur_epoch_float = solver.epoch + solver.iter / solver.epoch_max_iter - 1 + # solver.logger.info(cur_epoch_float) + if self.warmup_func is not None and cur_epoch_float < self.warmup_epochs: + new_lr = self._get_warmup_lr(cur_epoch_float) + else: + new_lr = _get_lr_from_scheduler(solver.lr_scheduler, + cur_epoch_float) + for param_group in solver.optimizer.param_groups: + param_group['lr'] = new_lr + + @staticmethod + def get_config_template(): + return dict_to_yaml('HOOK', + __class__.__name__, + LrHook.para_dict, + set_name=True) diff --git a/scepter/modules/solver/hooks/registry.py b/scepter/modules/solver/hooks/registry.py new file mode 100644 index 0000000..a7802ff --- /dev/null +++ b/scepter/modules/solver/hooks/registry.py @@ -0,0 +1,6 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.utils.registry import Registry + +HOOKS = Registry('HOOKS') diff --git a/scepter/modules/solver/hooks/safetensors.py b/scepter/modules/solver/hooks/safetensors.py new file mode 100644 index 0000000..5abfb4d --- /dev/null +++ b/scepter/modules/solver/hooks/safetensors.py @@ -0,0 +1,59 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import os.path as osp + +import torch +from safetensors.torch import save_file + +from scepter.modules.solver.hooks.hook import Hook +from scepter.modules.solver.hooks.registry import HOOKS +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we +from scepter.modules.utils.file_system import FS + +_DEFAULT_CHECKPOINT_PRIORITY = 300 + + +@HOOKS.register_class() +class SafetensorsHook(Hook): + para_dict = [{ + 'PRIORITY': { + 'value': _DEFAULT_CHECKPOINT_PRIORITY, + 'description': 'the priority for processing!' + }, + 'INTERVAL': { + 'value': 1, + 'description': 'the interval of saving checkpoint!' + } + }] + + def __init__(self, cfg, logger=None): + super(SafetensorsHook, self).__init__(cfg, logger=logger) + self.priority = cfg.get('PRIORITY', _DEFAULT_CHECKPOINT_PRIORITY) + self.interval = cfg.get('INTERVAL', 5000) + self.save_name_prefix = cfg.get('SAVE_NAME_PREFIX', 'ldm_step') + + def after_iter(self, solver): + if solver.total_iter != 0 and ( + (solver.total_iter + 1) % self.interval == 0 + or solver.total_iter == solver.max_steps - 1): + state_dict, metadata = solver.save_safetensors() + solver.logger.info( + f'Saving safetensors after {solver.total_iter + 1} steps') + if we.rank == 0: + save_path = osp.join( + solver.work_dir, 'safetensors/{}-{}.safetensors'.format( + self.save_name_prefix, solver.total_iter + 1)) + with FS.put_to(save_path) as local_path: + with open(local_path, 'wb'): + save_file(state_dict, local_path, metadata) + torch.cuda.synchronize() + if we.is_distributed: + torch.distributed.barrier() + + @staticmethod + def get_config_template(): + return dict_to_yaml('hook', + __class__.__name__, + SafetensorsHook.para_dict, + set_name=True) diff --git a/scepter/modules/solver/hooks/sampler.py b/scepter/modules/solver/hooks/sampler.py new file mode 100644 index 0000000..0f6537b --- /dev/null +++ b/scepter/modules/solver/hooks/sampler.py @@ -0,0 +1,42 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.solver.hooks.hook import Hook +from scepter.modules.solver.hooks.registry import HOOKS +from scepter.modules.utils.config import dict_to_yaml + +_DEFAULT_SAMPLER_PRIORITY = 400 + + +@HOOKS.register_class() +class DistSamplerHook(Hook): + """ DistributedDataSampler needs to set_epoch to shuffle sample indexes + """ + para_dict = [{ + 'PRIORITY': { + 'value': _DEFAULT_SAMPLER_PRIORITY, + 'description': 'the priority for processing!' + } + }] + + def __init__(self, cfg, logger=None): + super(DistSamplerHook, self).__init__(cfg, logger=logger) + self.priority = cfg.get('PRIORITY', _DEFAULT_SAMPLER_PRIORITY) + + def before_epoch(self, solver): + for name, data_ins in solver.datas.items(): + if name == 'train': + data_loader = data_ins.dataloader + solver.logger.info( + f'distribute sampler set_epoch to {solver.epoch}') + if hasattr(data_loader.sampler, 'set_epoch'): + data_loader.sampler.set_epoch(solver.epoch) + elif hasattr(data_loader.batch_sampler.sampler, 'set_epoch'): + data_loader.batch_sampler.sampler.set_epoch(solver.epoch) + + @staticmethod + def get_config_template(): + return dict_to_yaml('HOOK', + __class__.__name__, + DistSamplerHook.para_dict, + set_name=True) diff --git a/scepter/modules/solver/registry.py b/scepter/modules/solver/registry.py new file mode 100644 index 0000000..82919fd --- /dev/null +++ b/scepter/modules/solver/registry.py @@ -0,0 +1,40 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import inspect + +from scepter.modules.utils.registry import Registry, deep_copy + + +def build_solver(cfg, registry, logger=None, *args, **kwargs): + from scepter.modules.utils.config import Config + if not isinstance(cfg, Config): + raise TypeError(f'config must be type Config, got {type(cfg)}') + if not cfg.have('NAME'): + raise KeyError(f'config must contain key NAME, got {cfg}') + if not isinstance(registry, Registry): + raise TypeError( + f'registry must be type Registry, got {type(registry)}') + + cfg = deep_copy(cfg) + + req_type = cfg.get('NAME') + if isinstance(req_type, str): + req_type_entry = registry.get(req_type) + if req_type_entry is None: + raise KeyError(f'{req_type} not found in {registry.name} registry') + + if kwargs is not None: + cfg._update_dict(kwargs) + + if inspect.isclass(req_type_entry): + try: + return req_type_entry(cfg, logger=logger, *args, **kwargs) + except Exception as e: + raise Exception(f'Failed to init class {req_type_entry}, with {e}') + else: + raise TypeError( + f'type must be str or class, got {type(req_type_entry)}') + + +SOLVERS = Registry('SOLVERS', build_func=build_solver, allow_types=('class', )) diff --git a/scepter/modules/solver/train_val_solver.py b/scepter/modules/solver/train_val_solver.py new file mode 100644 index 0000000..25e65b6 --- /dev/null +++ b/scepter/modules/solver/train_val_solver.py @@ -0,0 +1,200 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import os.path as osp +from collections import OrderedDict, defaultdict + +import torch + +from scepter.modules.solver.base_solver import BaseSolver +from scepter.modules.solver.registry import SOLVERS +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.data import (transfer_data_to_cpu, + transfer_data_to_cuda) +from scepter.modules.utils.distribute import gather_data, we +from scepter.modules.utils.file_system import FS + + +def _get_value(data: dict, key: str): + """ Recursively get value from data by a multi-level key. + + Args: + data (dict): + key (str): 'data', 'meta.path', 'a.b.c' + + Returns: + Value. + + """ + if not isinstance(data, dict): + return None + if key in data: + return data[key] + elif '.' in key: + par_key = key.split('.')[0] + sub_key = '.'.join(key.split('.')[1:]) + if par_key in data: + return _get_value(data[par_key], sub_key) + return None + + +@SOLVERS.register_class() +class TrainValSolver(BaseSolver): + """ Standard train and eval steps solver + + Args: + model (torch.nn.Module): Model to train or eval. + + """ + para_dict = { + 'DO_FINAL_EVAL': { + 'value': False, + 'description': 'If do final evaluation or not.' + }, + 'SAVE_EVAL_DATA': { + 'value': False, + 'description': 'If save the evaluation data or not.' + } + } + para_dict.update(BaseSolver.para_dict) + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + if not self.use_pl: + if 'train' in self.datas and self.cfg.have('OPTIMIZER'): + self.cfg.OPTIMIZER.LEARNING_RATE *= self.datas[ + 'train'].batch_size + if we.world_size > 1: + self.cfg.OPTIMIZER.LEARNING_RATE *= we.world_size + self.cfg.OPTIMIZER.LEARNING_RATE *= self.accu_step + self.cfg.OPTIMIZER.LEARNING_RATE /= 96 + + def construct_metrics(self): + # Initial metric + super().construct_metrics() + self.metrics = [] + if self.cfg.have('METRICS'): + self.extra_keys = self.cfg.get('EXTRA_KEYS', []) + self._collect_keys = set() + self.do_final_eval = self.cfg.get('DO_FINAL_EVAL', False) + self.save_eval_data = self.cfg.get('SAVE_EVAL_DATA', False) + if self.do_final_eval or self.save_eval_data: + self._build_metrics(self.cfg.METRICS) + self._collect_keys.update(list(self.extra_keys or [])) + self._collect_keys = sorted(list(self._collect_keys)) + if len(self._collect_keys) > 0: + self.logger.info( + f"{', '.join(self._collect_keys)} will be collected during eval epoch" + ) + + @torch.no_grad() + def run_eval(self): + self.eval_mode() + collect_data = defaultdict(list) + rank, world_size = we.rank, we.world_size + self.before_all_iter(self.hooks_dict[self._mode]) + for data in self.datas[self._mode].dataloader: + self.before_iter(self.hooks_dict[self._mode]) + data_gpu = transfer_data_to_cuda(data) + result = self.model(**data_gpu) + self._iter_outputs[self._mode] = self._reduce_scalar(result) + if self.do_final_eval or self.save_eval_data: + # Collect data + if isinstance(result, torch.Tensor): + data_gpu['result'] = result + elif isinstance(result, dict): + data_gpu.update(result) + + step_data = OrderedDict() + for key in self._collect_keys: + value = _get_value(data_gpu, key) + if value is None: + raise ValueError( + f'Cannot get valid value from model input or output data with key {key}' + ) + step_data[key] = value + + step_data = transfer_data_to_cpu(step_data) + + for key, value in step_data.items(): + if isinstance(value, torch.Tensor): + collect_data[key].append(value.clone()) + else: + collect_data[key].append(value) + + self.after_iter(self.hooks_dict[self._mode]) + self.after_all_iter(self.hooks_dict[self._mode]) + + if self.do_final_eval or self.save_eval_data: + # Concat collect_data + concat_collect_data = OrderedDict() + for key, tensors in collect_data.items(): + if isinstance(tensors[0], torch.Tensor): + concat_collect_data[key] = torch.cat(tensors) + elif isinstance(tensors[0], list): + concat_collect_data[key] = sum(tensors, []) + else: + concat_collect_data[key] = tensors + + # If distributed and use DistributedSampler + # Gather all collect data to rank 0 + if world_size > 1 and type( + self.datas[self._mode].sampler + ) is torch.utils.data.DistributedSampler: + concat_collect_data = { + key: gather_data(concat_collect_data[key]) + for key in self._collect_keys + } + + # Do final evaluate + if self.do_final_eval and rank == 0: + for metric in self.metrics: + self._epoch_outputs[self._mode].update(metric['fn']( + *[concat_collect_data[key] for key in metric['keys']])) + + # Save all data + if self.save_eval_data and rank == 0: + # minus 1, means index + save_path = osp.join( + self.work_dir, + 'eval_{:05d}.pth'.format(self.epoch + self.num_folds)) + with FS.put_to(save_path) as local_file: + torch.save(concat_collect_data, local_file) + + def load_checkpoint(self, checkpoint: dict): + self._epoch = checkpoint['epoch'] + for mode_name, total_iter in checkpoint['total_iters'].items(): + self._total_iter[mode_name] = total_iter + self.model.load_state_dict(checkpoint['state_dict']) + self.optimizer.load_state_dict(checkpoint['checkpoint']) + if self.lr_scheduler is not None: + self.lr_scheduler.load_state_dict(checkpoint['lr_scheduler']) + self._epoch += 1 # Move to next epoch + + def save_checkpoint(self) -> dict: + checkpoint = { + 'epoch': self._epoch, + 'total_iters': self._total_iter, + 'state_dict': self.model.state_dict(), + 'checkpoint': self.optimizer.state_dict(), + } + if self.lr_scheduler is not None: + checkpoint['lr_scheduler'] = self.lr_scheduler.state_dict() + return checkpoint + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('solvername', + __class__.__name__, + TrainValSolver.para_dict, + set_name=True) diff --git a/scepter/modules/transform/__init__.py b/scepter/modules/transform/__init__.py new file mode 100644 index 0000000..b57b57e --- /dev/null +++ b/scepter/modules/transform/__init__.py @@ -0,0 +1,26 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.transform.augmention import ColorJitterGeneral +from scepter.modules.transform.compose import Compose +from scepter.modules.transform.identity import Identity +from scepter.modules.transform.image import (CenterCrop, FlexibleCenterCrop, + FlexibleResize, ImageToTensor, + ImageTransform, Normalize, + RandomHorizontalFlip, + RandomResizedCrop, Resize) +from scepter.modules.transform.io import (LoadCvImageFromFile, + LoadImageFromFile, + LoadImageFromFileList, + LoadPILImageFromFile) +from scepter.modules.transform.io_video import (DecodeVideoToTensor, + LoadVideoFromFile) +from scepter.modules.transform.registry import TRANSFORMS, build_pipeline +from scepter.modules.transform.tensor import Rename, Select, ToTensor +from scepter.modules.transform.transform_xl import FlexibleCropXL +from scepter.modules.transform.video import (AutoResizedCropVideo, + CenterCropVideo, NormalizeVideo, + RandomHorizontalFlipVideo, + RandomResizedCropVideo, + ResizeVideo, VideoToTensor, + VideoTransform) diff --git a/scepter/modules/transform/augmention.py b/scepter/modules/transform/augmention.py new file mode 100644 index 0000000..c419443 --- /dev/null +++ b/scepter/modules/transform/augmention.py @@ -0,0 +1,506 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import numbers +import random + +# https://github.com/TengdaHan/DPC/blob/master/utils/augmentation.py +import torch +from torchvision.transforms import Compose, Lambda + +from scepter.modules.transform.image import ImageTransform +from scepter.modules.transform.registry import TRANSFORMS +from scepter.modules.transform.utils import (BACKEND_CV2, BACKEND_PILLOW, + BACKEND_TORCHVISION, + TORCHVISION_CAPABILITY) +from scepter.modules.utils.config import dict_to_yaml + +if TORCHVISION_CAPABILITY: + BACKENDS = (BACKEND_PILLOW, BACKEND_CV2, BACKEND_TORCHVISION) +else: + BACKENDS = (BACKEND_PILLOW, BACKEND_CV2) + + +def _is_tensor_a_torch_image(input): + return input.ndim >= 2 + + +def _blend(img1, img2, ratio): + # type: (torch.Tensor, torch.Tensor, float) -> torch.Tensor + bound = 1 if img1.dtype in [torch.half, torch.float32, torch.float64 + ] else 255 + return (ratio * img1 + (1 - ratio) * img2).clamp(0, bound).to(img1.dtype) + + +def rgb_to_grayscale(img, split=False): + # type: (torch.Tensor) -> torch.Tensor + """Convert the given RGB Image Tensor to Grayscale. + For RGB to Grayscale conversion, ITU-R 601-2 luma transform is performed which + is L = R * 0.2989 + G * 0.5870 + B * 0.1140 + Args: + img (Tensor): Image to be converted to Grayscale in the form [C, H, W]. + Returns: + Tensor: Grayscale image. + Args: + clip (torch.tensor): Size is (T, H, W, C) + Return: + clip (torch.tensor): Size is (T, H, W, C) + """ + orig_dtype = img.dtype + rgb_convert = torch.tensor([0.299, 0.587, 0.114]) + if split: + rgb_convert *= 0 + channel = random.randint(0, 2) + rgb_convert[channel] = 1 + + assert img.shape[0] == 3, 'First dimension need to be 3 Channels' + if img.is_cuda: + rgb_convert = rgb_convert.to(img.device) + + img = img.float().permute(1, 2, 3, 0).matmul(rgb_convert).to(orig_dtype) + return torch.stack([img, img, img], 0) + + +def _rgb2hsv(img): + r, g, b = img.unbind(0) + + maxc, _ = torch.max(img, dim=0) + minc, _ = torch.min(img, dim=0) + + eqc = maxc == minc + cr = maxc - minc + s = cr / torch.where(eqc, maxc.new_ones(()), maxc) + cr_divisor = torch.where(eqc, maxc.new_ones(()), cr) + rc = (maxc - r) / cr_divisor + gc = (maxc - g) / cr_divisor + bc = (maxc - b) / cr_divisor + + hr = (maxc == r) * (bc - gc) + hg = ((maxc == g) & (maxc != r)) * (2.0 + rc - bc) + hb = ((maxc != g) & (maxc != r)) * (4.0 + gc - rc) + h = (hr + hg + hb) + h = torch.fmod((h / 6.0 + 1.0), 1.0) + return torch.stack((h, s, maxc)) + + +def _hsv2rgb(img): + l = len(img.shape) # noqa + h, s, v = img.unbind(0) + i = torch.floor(h * 6.0) + f = (h * 6.0) - i + i = i.to(dtype=torch.int32) + + p = torch.clamp((v * (1.0 - s)), 0.0, 1.0) + q = torch.clamp((v * (1.0 - s * f)), 0.0, 1.0) + t = torch.clamp((v * (1.0 - s * (1.0 - f))), 0.0, 1.0) + i = i % 6 + + if l == 3: # noqa + tmp = torch.arange(6)[:, None, None] + elif l == 4: # noqa + tmp = torch.arange(6)[:, None, None, None] + + if img.is_cuda: + tmp = tmp.to(img.device) + + mask = i == tmp # (H, W) == (6, H, W) + + a1 = torch.stack((v, q, p, p, t, v)) + a2 = torch.stack((t, v, v, q, p, p)) + a3 = torch.stack((p, p, t, v, v, q)) + a4 = torch.stack((a1, a2, a3)) # (3, 6, H, W) + + if l == 3: # noqa + return torch.einsum('ijk, xijk -> xjk', mask.to(dtype=img.dtype), + a4) # (C, H, W) + elif l == 4: # noqa + return torch.einsum('itjk, xitjk -> xtjk', mask.to(dtype=img.dtype), + a4) # (C, T, H, W) + + +def adjust_brightness(img, brightness_factor): + # type: (torch.Tensor, float) -> torch.Tensor + if not _is_tensor_a_torch_image(img): + raise TypeError('tensor is not a torch image.') + + return _blend(img, torch.zeros_like(img), brightness_factor) + + +def adjust_contrast(img, contrast_factor): + # type: (torch.Tensor, float) -> torch.Tensor + if not _is_tensor_a_torch_image(img): + raise TypeError('tensor is not a torch image.') + + mean = torch.mean(rgb_to_grayscale(img).to(torch.float), + dim=(-4, -2, -1), + keepdim=True) + + return _blend(img, mean, contrast_factor) + + +def adjust_saturation(img, saturation_factor): + # type: (torch.Tensor, float) -> torch.Tensor + if not _is_tensor_a_torch_image(img): + raise TypeError('tensor is not a torch image.') + + return _blend(img, rgb_to_grayscale(img), saturation_factor) + + +def adjust_hue(img, hue_factor): + """Adjust hue of an image. + The image hue is adjusted by converting the image to HSV and + cyclically shifting the intensities in the hue channel (H). + The image is then converted back to original image mode. + `hue_factor` is the amount of shift in H channel and must be in the + interval `[-0.5, 0.5]`. + See `Hue`_ for more details. + .. _Hue: https://en.wikipedia.org/wiki/Hue + Args: + img (Tensor): Image to be adjusted. Image type is either uint8 or float. + hue_factor (float): How much to shift the hue channel. Should be in + [-0.5, 0.5]. 0.5 and -0.5 give complete reversal of hue channel in + HSV space in positive and negative direction respectively. + 0 means no shift. Therefore, both -0.5 and 0.5 will give an image + with complementary colors while 0 gives the original image. + Returns: + Tensor: Hue adjusted image. + """ + if isinstance(hue_factor, float) and not (-0.5 <= hue_factor <= 0.5): + raise ValueError( + 'hue_factor ({}) is not in [-0.5, 0.5].'.format(hue_factor)) + elif (isinstance(hue_factor, torch.Tensor) + and not ((-0.5 <= hue_factor).sum() == hue_factor.shape[0] and + (hue_factor <= 0.5).sum() == hue_factor.shape[0])): + raise ValueError( + 'hue_factor ({}) is not in [-0.5, 0.5].'.format(hue_factor)) + + if not _is_tensor_a_torch_image(img): + raise TypeError('tensor is not a torch image.') + + orig_dtype = img.dtype + if img.dtype == torch.uint8: + img = img.to(dtype=torch.float32) / 255.0 + + img = _rgb2hsv(img) + h, s, v = img.unbind(0) + h += hue_factor + h = h % 1.0 + img = torch.stack((h, s, v)) + img_hue_adj = _hsv2rgb(img) + + if orig_dtype == torch.uint8: + img_hue_adj = (img_hue_adj * 255.0).to(dtype=orig_dtype) + return img_hue_adj + + +# https://github.com/TengdaHan/DPC/blob/master/utils/augmentation.py +class ColorJitter(object): + """Randomly change the brightness, contrast and saturation of an image. + Args: + brightness (float or tuple of float (min, max)): How much to jitter brightness. + brightness_factor is chosen uniformly from [max(0, 1 - brightness), 1 + brightness] + or the given [min, max]. Should be non negative numbers. + contrast (float or tuple of float (min, max)): How much to jitter contrast. + contrast_factor is chosen uniformly from [max(0, 1 - contrast), 1 + contrast] + or the given [min, max]. Should be non negative numbers. + saturation (float or tuple of float (min, max)): How much to jitter saturation. + saturation_factor is chosen uniformly from [max(0, 1 - saturation), 1 + saturation] + or the given [min, max]. Should be non negative numbers. + hue (float or tuple of float (min, max)): How much to jitter hue. + hue_factor is chosen uniformly from [-hue, hue] or the given [min, max]. + Should have 0<= hue <= 0.5 or -0.5 <= min <= max <= 0.5. + grayscale (probablitities for rgb-to-gray 0~1) + consistent (for video input whether the augment scale is consistent or not!) + shuffle (shuffle the transform's order when there are multiple transform) + gray_first (whether use grayscale or not) + is_split (whether randomly chance the channel as the gray results.) + """ + def __init__(self, + brightness=0, + contrast=0, + saturation=0, + hue=0, + grayscale=0, + consistent=False, + shuffle=True, + gray_first=True, + is_split=False): + self.brightness = self._check_input(brightness, 'brightness') + self.contrast = self._check_input(contrast, 'contrast') + self.saturation = self._check_input(saturation, 'saturation') + self.hue = self._check_input(hue, + 'hue', + center=0, + bound=(-0.5, 0.5), + clip_first_on_zero=False) + self.grayscale = grayscale + self.consistent = consistent + self.shuffle = shuffle + self.gray_first = gray_first + self.is_split = is_split + + def _check_input(self, + value, + name, + center=1, + bound=(0, float('inf')), + clip_first_on_zero=True): + if isinstance(value, numbers.Number): + if value < 0: + raise ValueError( + 'If {} is a single number, it must be non negative.'. + format(name)) + value = [center - float(value), center + float(value)] + if clip_first_on_zero: + value[0] = max(value[0], 0.0) + elif isinstance(value, (tuple, list)) and len(value) == 2: + if not bound[0] <= value[0] <= value[1] <= bound[1]: + raise ValueError('{} values should be between {}'.format( + name, bound)) + else: + raise TypeError( + '{} should be a single number or a list/tuple with lenght 2.'. + format(name)) + + # if value is 0 or (1., 1.) for brightness/contrast/saturation + # or (0., 0.) for hue, do nothing + if value[0] == value[1] == center: + value = None + return value + + def _get_transform(self, T, device): + """Get a randomized transform to be applied on image. + Arguments are same as that of __init__. + Arg: + T (int): number of frames. Used when consistent = False. + Returns: + Transform which randomly adjusts brightness, contrast and + saturation in a random order. + """ + transforms = [] + if self.brightness is not None: + if self.consistent: + brightness_factor = random.uniform(self.brightness[0], + self.brightness[1]) + else: + brightness_factor = torch.empty([1, T, 1, 1], + device=device).uniform_( + self.brightness[0], + self.brightness[1]) + transforms.append( + Lambda( + lambda frame: adjust_brightness(frame, brightness_factor))) + + if self.contrast is not None: + if self.consistent: + contrast_factor = random.uniform(self.contrast[0], + self.contrast[1]) + else: + contrast_factor = torch.empty([1, T, 1, 1], + device=device).uniform_( + self.contrast[0], + self.contrast[1]) + transforms.append( + Lambda(lambda frame: adjust_contrast(frame, contrast_factor))) + + if self.saturation is not None: + if self.consistent: + saturation_factor = random.uniform(self.saturation[0], + self.saturation[1]) + else: + saturation_factor = torch.empty([1, T, 1, 1], + device=device).uniform_( + self.saturation[0], + self.saturation[1]) + transforms.append( + Lambda( + lambda frame: adjust_saturation(frame, saturation_factor))) + + if self.hue is not None: + if self.consistent: + hue_factor = random.uniform(self.hue[0], self.hue[1]) + else: + hue_factor = torch.empty([T, 1, 1], device=device).uniform_( + self.hue[0], self.hue[1]) + transforms.append( + Lambda(lambda frame: adjust_hue(frame, hue_factor))) + + if self.shuffle: + random.shuffle(transforms) + + if random.uniform(0, 1) < self.grayscale: + gray_transform = Lambda( + lambda frame: rgb_to_grayscale(frame, split=self.is_split)) + if self.gray_first: + transforms.insert(0, gray_transform) + else: + transforms.append(gray_transform) + + transform = Compose(transforms) + return transform + + def __call__(self, clip): + """ + Args: + clip (torch.tensor): Size is (C, T, H, W) + Return: + clip (torch.tensor): Size is (C, T, H, W) + """ + + is_frame = False + + if len(clip.shape) == 3: + # frame data, transfer to 4dim + clip = torch.unsqueeze(clip, dim=1) + is_frame = True + # (C, T, H, W) + raw_shape = clip.shape + device = clip.device + T = raw_shape[1] + transform = self._get_transform(T, device) + clip = transform(clip) + assert clip.shape == raw_shape + if is_frame: + clip = torch.squeeze(clip) + return clip # (C, T, H, W) + + def __repr__(self): + format_string = self.__class__.__name__ + '(' + format_string += 'brightness={0}'.format(self.brightness) + format_string += ', contrast={0}'.format(self.contrast) + format_string += ', saturation={0}'.format(self.saturation) + format_string += ', hue={0})'.format(self.hue) + format_string += ', grayscale={0})'.format(self.grayscale) + return format_string + + +@TRANSFORMS.register_class() +class ColorJitterGeneral(ImageTransform): + ''' + brightness (float or tuple of float (min, max)): How much to jitter brightness. + brightness_factor is chosen uniformly from [max(0, 1 - brightness), 1 + brightness] + or the given [min, max]. Should be non negative numbers. + contrast (float or tuple of float (min, max)): How much to jitter contrast. + contrast_factor is chosen uniformly from [max(0, 1 - contrast), 1 + contrast] + or the given [min, max]. Should be non negative numbers. + saturation (float or tuple of float (min, max)): How much to jitter saturation. + saturation_factor is chosen uniformly from [max(0, 1 - saturation), 1 + saturation] + or the given [min, max]. Should be non negative numbers. + hue (float or tuple of float (min, max)): How much to jitter hue. + hue_factor is chosen uniformly from [-hue, hue] or the given [min, max]. + Should have 0<= hue <= 0.5 or -0.5 <= min <= max <= 0.5. + grayscale (probablitities for rgb-to-gray 0~1) + consistent (for video input whether the augment scale is consistent or not!) + shuffle (shuffle the transform's order when there are multiple transform) + gray_first (whether use grayscale or not) + is_split (whether randomly chance the channel as the gray results.) + ''' + para_dict = [{ + 'BRIGHTNESS': { + 'value': + 0, + 'description': + '(float or tuple of float (min, max)): How much to jitter brightness.' + 'brightness_factor is chosen uniformly from [max(0, 1 - brightness), 1 + brightness]' + 'or the given [min, max]. Should be non negative numbers.' + }, + 'CONTRAST': { + 'value': + 0, + 'description': + '(float or tuple of float (min, max)): How much to jitter contrast.' + 'contrast_factor is chosen uniformly from [max(0, 1 - contrast), 1 + contrast]' + 'or the given [min, max]. Should be non negative numbers.' + }, + 'SATURATION': { + 'value': + 0, + 'description': + '(float or tuple of float (min, max)): How much to jitter saturation.' + 'saturation_factor is chosen uniformly from [max(0, 1 - saturation), 1 + saturation]' + 'or the given [min, max]. Should be non negative numbers.' + }, + 'HUE': { + 'value': + 0, + 'description': + '(float or tuple of float (min, max)): How much to jitter hue.' + 'hue_factor is chosen uniformly from [-hue, hue] or the given [min, max].' + 'Should have 0<= hue <= 0.5 or -0.5 <= min <= max <= 0.5.' + }, + 'GRAYSCALE': { + 'value': 0, + 'description': '(probablitities for rgb-to-gray 0~1)' + }, + 'CONSISTENT': { + 'value': + False, + 'description': + '(for video input whether the augment scale is consistent or not!)' + }, + 'SHUFFLE': { + 'value': + False, + 'description': + "(shuffle the transform's order when there are multiple transform)" + }, + 'GRAY_FIRST': { + 'value': False, + 'description': '(whether use grayscale or not)' + }, + 'IS_SPLIT': { + 'value': + False, + 'description': + '(whether randomly chance the channel as the gray results.)' + } + }] + para_dict[0].update(ImageTransform.para_dict[0]) + para_dict[0].pop('BACKEND') + + def __init__(self, cfg, logger=None): + super(ColorJitterGeneral, self).__init__(cfg, logger=None) + # brightness=0, contrast=0, saturation=0, hue=0, grayscale=0, + # consistent=False, shuffle=True, gray_first=True, is_split=False + brightness = cfg.get('BRIGHTNESS', 0) + contrast = cfg.get('CONTRAST', 0) + saturation = cfg.get('SATURATION', 0) + hue = cfg.get('HUE', 0) + grayscale = cfg.get('GRAYSCALE', 0) + consistent = cfg.get('CONSISTENT', False) + shuffle = cfg.get('SHUFFLE', False) + gray_first = cfg.get('GRAY_FIRST', False) + is_split = cfg.get('IS_SPLIT', False) + + cj_ins = ColorJitter(brightness=brightness, + contrast=contrast, + saturation=saturation, + hue=hue, + grayscale=grayscale, + consistent=consistent, + shuffle=shuffle, + gray_first=gray_first, + is_split=is_split) + self.callable = cj_ins + + def __call__(self, item): + item[self.output_key] = self.callable(item[self.input_key]) + # item["meta"]["normalize_params"] = dict(mean=self.mean, std=self.std) + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + ColorJitterGeneral.para_dict, + set_name=True) diff --git a/scepter/modules/transform/compose.py b/scepter/modules/transform/compose.py new file mode 100644 index 0000000..ac78ce6 --- /dev/null +++ b/scepter/modules/transform/compose.py @@ -0,0 +1,39 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.transform.registry import TRANSFORMS +from scepter.modules.utils.config import dict_to_yaml + + +@TRANSFORMS.register_class() +class Compose(object): + """ Compose all transform function into one. + + Args: + transform (List[dict]): List of transform configs. + + """ + def __init__(self, cfg, logger=None): + self.transforms = [TRANSFORMS.build(t) for t in cfg.TRANSFORMS] + + def __call__(self, item): + for t in self.transforms: + item = t(item) + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, [{}], + set_name=True) diff --git a/scepter/modules/transform/identity.py b/scepter/modules/transform/identity.py new file mode 100644 index 0000000..fa1f9b5 --- /dev/null +++ b/scepter/modules/transform/identity.py @@ -0,0 +1,31 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from ..utils.config import dict_to_yaml +from .registry import TRANSFORMS + + +@TRANSFORMS.register_class() +class Identity(object): + def __init__(self, cfg, logger=None): + pass + + def __call__(self, item): + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, [{}], + set_name=True) diff --git a/scepter/modules/transform/image.py b/scepter/modules/transform/image.py new file mode 100644 index 0000000..fea07d0 --- /dev/null +++ b/scepter/modules/transform/image.py @@ -0,0 +1,626 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import numbers + +import numpy as np +import opencv_transforms.functional as cv2_TF +import opencv_transforms.transforms as cv2_transforms +import torch +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.utils.config import dict_to_yaml + +if TORCHVISION_CAPABILITY: + BACKENDS = (BACKEND_PILLOW, BACKEND_CV2, BACKEND_TORCHVISION) +else: + BACKENDS = (BACKEND_PILLOW, BACKEND_CV2) + + +class ImageTransform(object): + para_dict = [{ + 'INPUT_KEY': { + 'value': 'img', + 'description': 'input key or key list.' + }, + 'OUTPUT_KEY': { + 'value': 'img', + 'description': 'input key or key list.' + }, + 'BACKEND': { + 'value': 'pillow', + 'description': 'backend, choose from pillow, cv2, torchvision' + } + }] + + def __init__(self, cfg, logger=None): + self.input_key = cfg.get('INPUT_KEY', 'img') + self.output_key = cfg.get('OUTPUT_KEY', 'img') + self.backend = cfg.get('BACKEND', BACKEND_PILLOW) + + def check_image_type(self, input_img): + if self.backend == BACKEND_PILLOW: + assert is_pil_image(input_img), INPUT_PIL_TYPE_WARNING + w, h = input_img.size + return h, w + elif self.backend == BACKEND_CV2: + assert is_cv2_image(input_img), INPUT_CV2_TYPE_WARNING + h, w, c = input_img.shape + return h, w + elif TORCHVISION_CAPABILITY: + if self.backend == BACKEND_TORCHVISION: + assert is_tensor(input_img), INPUT_TENSOR_TYPE_WARNING + c, h, w = input_img.shape + return h, w + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + ImageTransform.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class RandomCrop(ImageTransform): + """ Crop a random portion of image. + If the image is torch Tensor, it is expected to have [..., H, W] shape. + + Args: + size (sequence or int): Desired output size. + If size is a sequence like (h, w), the output size will be matched to this. + If size is an int, the output size will be matched to (size, size). + padding (sequence or int): Optional padding on each border of the image. Default is None. + pad_if_needed (bool): It will pad the image if smaller than the desired size to avoid raising an exception. + fill (number or str or tuple): Pixel fill value for constant fill. Default is 0. + padding_mode (str): Type of padding. Should be: constant, edge, reflect or symmetric. + Default is constant. + """ + para_dict = [{ + 'SIZE': { + 'value': 224, + 'description': 'crop size' + }, + 'PADDING': { + 'value': None, + 'description': 'padding' + }, + 'PAD_IF_NEEDED': { + 'value': False, + 'description': 'pad if needed' + }, + 'FILL': { + 'value': 0, + 'description': 'fill' + }, + 'PADDING_MODE': { + 'value': 'constant', + 'description': 'padding mode' + } + }] + para_dict[0].update(ImageTransform.para_dict[0]) + + def __init__(self, cfg, logger=None): + size = cfg.SIZE + padding = cfg.get('PADDING', None) + pad_if_needed = cfg.get('PAD_IF_NEEDED', False) + fill = cfg.get('FILL', 0) + padding_mode = cfg.get('PADDING_MODE', 'constant') + super(RandomCrop, self).__init__(cfg) + assert self.backend in BACKENDS + if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION): + self.callable = transforms.RandomCrop(size, + padding=padding, + pad_if_needed=pad_if_needed, + fill=fill, + padding_mode=padding_mode) + else: + self.callable = cv2_transforms.RandomCrop( + size, + padding=padding, + pad_if_needed=pad_if_needed, + fill=fill, + padding_mode=padding_mode) + + def __call__(self, item): + if isinstance(self.input_key, str): + self.input_key = [self.input_key] + if isinstance(self.output_key, str): + self.output_key = [self.output_key] + for idx, key in enumerate(self.input_key): + self.check_image_type(item[key]) + item[self.output_key[idx]] = self.callable(item[key]) + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + RandomCrop.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class RandomResizedCrop(ImageTransform): + """Crop a random portion of image and resize it to a given size. + + If the image is torch Tensor, it is expected to have [..., H, W] shape. + + Args: + size (int or sequence): Desired output size. + If size is a sequence like (h, w), the output size will be matched to this. + If size is an int, the output size will be matched to (size, size). + scale (tuple of float): Specifies the lower and upper bounds for the random area of the crop, + before resizing. The scale is defined with respect to the area of the original image. + ratio (tuple of float): lower and upper bounds for the random aspect ratio of the crop, before + resizing. + interpolation (str): Desired interpolation string, 'bilinear', 'nearest', 'bicubic' are supported. + """ + para_dict = [{ + 'SIZE': { + 'value': 224, + 'description': 'crop size' + }, + 'RATIO': { + 'value': [3. / 4., 4. / 3.], + 'description': 'ratio' + }, + 'SCALE': { + 'value': [0.08, 1.0], + 'description': 'scale' + }, + 'INTERPOLATION': { + 'value': 'bilinear', + 'description': 'interpolation' + } + }] + para_dict[0].update(ImageTransform.para_dict[0]) + + def __init__(self, cfg, logger=None): + super(RandomResizedCrop, self).__init__(cfg, logger=logger) + assert self.backend in BACKENDS + self.interpolation = cfg.get('INTERPOLATION', 'bilinear') + self.size = cfg.SIZE + self.scale = tuple(cfg.get('SCALE', [0.08, 1.0])) + self.ratio = tuple(cfg.get('RATIO', [3. / 4., 4. / 3.])) + if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION): + assert self.interpolation in INTERPOLATION_STYLE + else: + assert self.interpolation in INTERPOLATION_STYLE_CV2 + self.callable = transforms.RandomResizedCrop(self.size, + self.scale, self.ratio, INTERPOLATION_STYLE[self.interpolation]) \ + if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION) \ + else cv2_transforms.RandomResizedCrop(self.size, self.scale, self.ratio, INTERPOLATION_STYLE_CV2[self.interpolation]) # noqa + + def __call__(self, item): + if isinstance(self.input_key, str): + self.input_key = [self.input_key] + if isinstance(self.output_key, str): + self.output_key = [self.output_key] + for idx, key in enumerate(self.input_key): + self.check_image_type(item[key]) + item[self.output_key[idx]] = self.callable(item[key]) + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + RandomResizedCrop.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class Resize(ImageTransform): + """Resize image to a given size. + + If the image is torch Tensor, it is expected to have [..., H, W] shape. + + Args: + size (int or sequence): Desired output size. + If size is a sequence like (h, w), the output size will be matched to this. + If size is an int, the smaller edge of the image will be matched to this number + maintaining the aspect ratio. + interpolation (str): Desired interpolation string, 'bilinear', 'nearest', 'bicubic' are supported. + """ + para_dict = [{ + 'INTERPOLATION': { + 'value': 'bilinear', + 'description': 'interpolation' + }, + 'SIZE': { + 'value': 224, + 'description': 'resize to size 224' + } + }] + para_dict[0].update(ImageTransform.para_dict[0]) + + def __init__(self, cfg, logger=None): + super(Resize, self).__init__(cfg, logger=logger) + assert self.backend in BACKENDS + self.size = cfg.SIZE + self.interpolation = cfg.get('INTERPOLATION', 'bilinear') + if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION): + assert self.interpolation in INTERPOLATION_STYLE + else: + assert self.interpolation in INTERPOLATION_STYLE_CV2 + self.callable = transforms.Resize(self.size, INTERPOLATION_STYLE[self.interpolation]) \ + if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION) \ + else cv2_transforms.Resize(self.size, INTERPOLATION_STYLE_CV2[self.interpolation]) + + def __call__(self, item): + if isinstance(self.input_key, str): + self.input_key = [self.input_key] + if isinstance(self.output_key, str): + self.output_key = [self.output_key] + for idx, key in enumerate(self.input_key): + self.check_image_type(item[key]) + item[self.output_key[idx]] = self.callable(item[key]) + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + Resize.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class CenterCrop(ImageTransform): + """ Crops the given image at the center. + + If the image is torch Tensor, it is expected to have [..., H, W] shape. + + Args: + size (sequence or int): Desired output size. + If size is a sequence like (h, w), the output size will be matched to this. + If size is an int, the output size will be matched to (size, size). + """ + para_dict = [{'SIZE': {'value': 224, 'description': 'resize to size 224'}}] + para_dict[0].update(ImageTransform.para_dict[0]) + + def __init__(self, cfg, logger=None): + super(CenterCrop, self).__init__(cfg, logger=None) + assert self.backend in BACKENDS + self.size = cfg.SIZE + self.callable = transforms.CenterCrop(self.size) \ + if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION) else cv2_transforms.CenterCrop(self.size) + + def __call__(self, item): + if isinstance(self.input_key, str): + self.input_key = [self.input_key] + if isinstance(self.output_key, str): + self.output_key = [self.output_key] + for idx, key in enumerate(self.input_key): + self.check_image_type(item[key]) + item[self.output_key[idx]] = self.callable(item[key]) + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + CenterCrop.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class RandomHorizontalFlip(ImageTransform): + """ Horizontally flip the given image randomly with a given probability. + + If the image is torch Tensor, it is expected to have [..., H, W] shape. + + Args: + p (float): probability of the image being flipped. Default value is 0.5 + """ + para_dict = [{'P': {'value': 0.5, 'description': 'P'}}] + para_dict[0].update(ImageTransform.para_dict[0]) + + def __init__(self, cfg, logger=None): + super(RandomHorizontalFlip, self).__init__(cfg, logger=None) + p = cfg.get('P', 0.5) + assert self.backend in BACKENDS + self.callable = transforms.RandomHorizontalFlip(p) \ + if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION) else cv2_transforms.RandomHorizontalFlip(p) + + def __call__(self, item): + if isinstance(self.input_key, str): + self.input_key = [self.input_key] + if isinstance(self.output_key, str): + self.output_key = [self.output_key] + for idx, key in enumerate(self.input_key): + self.check_image_type(item[key]) + item[self.output_key[idx]] = self.callable(item[key]) + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + RandomHorizontalFlip.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class Normalize(ImageTransform): + """ Normalize a tensor image with mean and standard deviation. + This transform only support tensor image. + + Args: + mean (sequence): Sequence of means for each channel. + std (sequence): Sequence of standard deviations for each channel. + """ + para_dict = [{ + 'MEAN': { + 'value': [], + 'description': 'mean' + }, + 'STD': { + 'value': [], + 'description': 'std' + }, + }] + para_dict[0].update(ImageTransform.para_dict[0]) + + def __init__(self, cfg, logger=None): + super(Normalize, self).__init__(cfg, logger=None) + assert self.backend in BACKENDS + mean = cfg.MEAN + std = cfg.STD + self.mean = np.array(mean, dtype=np.float32) + self.std = np.array(std, dtype=np.float32) + self.callable = transforms.Normalize(self.mean, self.std) \ + if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION) else cv2_transforms.Normalize(self.mean, self.std) + + def __call__(self, item): + if isinstance(self.input_key, str): + self.input_key = [self.input_key] + if isinstance(self.output_key, str): + self.output_key = [self.output_key] + for idx, key in enumerate(self.input_key): + item[self.output_key[idx]] = self.callable(item[key]) + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + Normalize.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class ImageToTensor(ImageTransform): + """ Convert a ``PIL Image`` or ``numpy.ndarray`` or uint8 type tensor to a float32 tensor, + and scale output to [0.0, 1.0]. + """ + para_dict = [{}] + para_dict[0].update(ImageTransform.para_dict[0]) + + def __init__(self, cfg, logger=None): + super(ImageToTensor, self).__init__(cfg, logger) + assert self.backend in BACKENDS + + if self.backend == BACKEND_PILLOW: + self.callable = transforms.ToTensor() + elif self.backend == BACKEND_CV2: + self.callable = cv2_transforms.ToTensor() + else: + self.callable = transforms.ConvertImageDtype(torch.float) + + def __call__(self, item): + if isinstance(self.input_key, str): + self.input_key = [self.input_key] + if isinstance(self.output_key, str): + self.output_key = [self.output_key] + for idx, key in enumerate(self.input_key): + item[self.output_key[idx]] = self.callable(item[key]) + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + ImageToTensor.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class FlexibleResize(ImageTransform): + para_dict = [{ + 'INTERPOLATION': { + 'value': 'bilinear', + 'description': 'interpolation' + }, + }] + para_dict[0].update(ImageTransform.para_dict[0]) + + def __init__(self, cfg, logger=None): + super(FlexibleResize, self).__init__(cfg, logger=logger) + assert self.backend in BACKENDS + self.size = cfg.get('SIZE', None) + if self.size is not None: + if isinstance(self.size, numbers.Number): + self.size = [self.size, self.size] + + interpolation = cfg.get('INTERPOLATION', 'bilinear') + if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION): + assert interpolation in INTERPOLATION_STYLE + else: + assert interpolation in INTERPOLATION_STYLE_CV2 + + if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION): + self.callable = TF.resize + self.interpolation = INTERPOLATION_STYLE[interpolation] + else: + self.callable = cv2_TF.resize + self.interpolation = INTERPOLATION_STYLE_CV2[interpolation] + + def __call__(self, item): + if isinstance(self.input_key, str): + self.input_key = [self.input_key] + if isinstance(self.output_key, str): + self.output_key = [self.output_key] + + for idx, key in enumerate(self.input_key): + ih, iw = self.check_image_type(item[key]) + meta = item.get('meta', {}) + if 'image_size' in meta: + iw, ih, ow, oh = iw, ih, meta['image_size'][1], meta[ + 'image_size'][0] + elif self.size is not None: + iw, ih, ow, oh = iw, ih, self.size[1], self.size[0] + meta['image_size'] = [oh, ow] + else: + raise KeyError( + 'The meta of input item must consists of ' + "['width', 'height', 'image_size'], and at least one key is missing." + ) + scale = max(ow / iw, oh / ih) + new_size = (round(scale * ih), round(scale * iw)) + item[self.output_key[idx]] = self.callable(item[key], new_size, + self.interpolation) + return item + + @staticmethod + def get_config_template(): + return dict_to_yaml('TRANSFORM', + __class__.__name__, + FlexibleResize.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class FlexibleCenterCrop(ImageTransform): + para_dict = [{}] + para_dict[0].update(ImageTransform.para_dict[0]) + + def __init__(self, cfg, logger=None): + super(FlexibleCenterCrop, self).__init__(cfg, logger=None) + assert self.backend in BACKENDS + self.size = cfg.get('SIZE', None) + if self.size is not None: + if isinstance(self.size, numbers.Number): + self.size = [self.size, self.size] + self.callable = TF.center_crop if self.backend in ( + BACKEND_PILLOW, BACKEND_TORCHVISION) else cv2_TF.center_crop + + def __call__(self, item): + if isinstance(self.input_key, str): + self.input_key = [self.input_key] + if isinstance(self.output_key, str): + self.output_key = [self.output_key] + + meta = item.get('meta', {}) + if 'image_size' in meta: + oh, ow = meta['image_size'] + out_size = (oh, ow) + else: + out_size = self.size + + for idx, key in enumerate(self.input_key): + self.check_image_type(item[key]) + item[self.output_key[idx]] = self.callable(item[key], out_size) + return item + + @staticmethod + def get_config_template(): + return dict_to_yaml('TRANSFORM', + __class__.__name__, + FlexibleCenterCrop.para_dict, + set_name=True) diff --git a/scepter/modules/transform/io.py b/scepter/modules/transform/io.py new file mode 100644 index 0000000..b0b4151 --- /dev/null +++ b/scepter/modules/transform/io.py @@ -0,0 +1,341 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import io +import os + +import cv2 +import numpy as np +import torch +from PIL import Image + +from scepter.modules.transform.registry import TRANSFORMS +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we +from scepter.modules.utils.file_system import DATA_FS as FS + + +def pillow_convert(image, rgb_order): + if image.mode != rgb_order: + if image.mode == 'P': + image = image.convert(f'{rgb_order}A') + if image.mode == f'{rgb_order}A': + bg = Image.new(rgb_order, + size=(image.width, image.height), + color=(255, 255, 255)) + bg.paste(image, (0, 0), mask=image) + image = bg + else: + image = image.convert('RGB') + return image + + +@TRANSFORMS.register_class() +class LoadImageFromFile(object): + """ Load Image from file. We have multi ways to load image. Here we compose them into one transform. + + Args: + rgb_order (str): 'RGB' or 'BGR'. + backend (str): 'pillow', 'cv2' or 'torchvision'. Image should be read as uint8 dtype. + - 'pillow': Read image file as PIL.Image object. + - 'cv2': Read image file as numpy.ndarray object. + - 'torchvision': Read image file as tensor object. + """ + def __init__(self, cfg, logger=None): + rgb_order = cfg.get('RGB_ORDER', 'RGB') + backend = cfg.get('BACKEND', 'pillow') + assert rgb_order in ('RGB', 'BGR') + assert backend in ('pillow', 'cv2', 'torchvision') + self.rgb_order = rgb_order + self.backend = backend + + def read_file(self, img_path): + if not we.data_online: + with FS.get_from(img_path) as img_path: + if self.backend == 'pillow': + try: + image = Image.open(img_path) + image = pillow_convert(image, self.rgb_order) + except Exception as e: + print(img_path, e) + elif self.backend == 'cv2': + image = cv2.imread(img_path, cv2.IMREAD_COLOR) + if self.rgb_order == 'RGB': + try: + cv2.cvtColor(image, cv2.COLOR_BGR2RGB, image) + except Exception as e: + print(img_path, e) + else: + image = Image.open(img_path).convert(self.rgb_order) + image_np = np.asarray(image).transpose( + (2, 0, 1)) # Tensor type needs shape to be (C, H, W) + image = torch.from_numpy(image_np) + return image + else: + with FS.get_object(img_path) as image_data: + if self.backend == 'pillow': + try: + image = Image.open(io.BytesIO(image_data)) + image = pillow_convert(image, self.rgb_order) + except Exception as e: + print(img_path, e) + elif self.backend == 'cv2': + image = cv2.imdecode( + np.array(bytearray(image_data), dtype='uint8'), + cv2.IMREAD_COLOR) + if self.rgb_order == 'RGB': + try: + cv2.cvtColor(image, cv2.COLOR_BGR2RGB, image) + except Exception as e: + print(img_path, e) + else: + image = Image.open(img_path).convert(self.rgb_order) + image_np = np.asarray(image).transpose( + (2, 0, 1)) # Tensor type needs shape to be (C, H, W) + image = torch.from_numpy(image_np) + return image + + def __call__(self, item): + if 'prefix' in item['meta']: + img_path = os.path.join(item['meta']['prefix'], + item['meta']['img_path']) + else: + img_path = item['meta']['img_path'] + item['img'] = self.read_file(img_path) + item['meta']['rgb_order'] = self.rgb_order + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + para_dict = [{ + 'RGB_ORDER': { + 'value': 'RGB', + 'description': 'rgb order' + }, + 'BACKEND': { + 'value': 'pillow', + 'description': 'input backend' + } + }] + return dict_to_yaml('TRANSFORM', + __class__.__name__, + para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class LoadImageFromFileList(object): + """ Load Image from file. We have multi ways to load image. Here we compose them into one transform. + + Args: + rgb_order (str): 'RGB' or 'BGR'. + backend (str): 'pillow', 'cv2' or 'torchvision'. Image should be read as uint8 dtype. + - 'pillow': Read image file as PIL.Image object. + - 'cv2': Read image file as numpy.ndarray object. + - 'torchvision': Read image file as tensor object. + """ + para_dict = [{ + 'RGB_ORDER': { + 'value': 'RGB', + 'description': 'Rgb order!' + }, + 'BACKEND': { + 'value': 'pillow', + 'description': 'Input backend!' + }, + 'FILE_KEYS': { + 'value': [], + 'description': + "The file keys for input, if key include '_path', " + "the return results will be saved with key as key.replace('_path', '')!" + } + }] + + def __init__(self, cfg, logger=None): + rgb_order = cfg.get('RGB_ORDER', 'RGB') + backend = cfg.get('BACKEND', 'pillow') + self.file_keys = cfg.get('FILE_KEYS', ['img_path']) + if isinstance(self.file_keys, str): + self.file_keys = [self.file_keys] + assert rgb_order in ('RGB', 'BGR') + assert backend in ('pillow', 'cv2', 'torchvision') + self.rgb_order = rgb_order + self.backend = backend + + def read_file(self, img_path): + if not we.data_online: + with FS.get_from(img_path) as img_path: + if self.backend == 'pillow': + try: + image = Image.open(img_path) + image = pillow_convert(image, self.rgb_order) + except Exception as e: + print(img_path, e) + elif self.backend == 'cv2': + image = cv2.imread(img_path, cv2.IMREAD_COLOR) + if self.rgb_order == 'RGB': + try: + cv2.cvtColor(image, cv2.COLOR_BGR2RGB, image) + except Exception as e: + print(img_path, e) + else: + image = Image.open(img_path).convert(self.rgb_order) + image_np = np.asarray(image).transpose( + (2, 0, 1)) # Tensor type needs shape to be (C, H, W) + image = torch.from_numpy(image_np) + return image + else: + with FS.get_object(img_path) as image_data: + if self.backend == 'pillow': + try: + image = Image.open(io.BytesIO(image_data)) + image = pillow_convert(image, self.rgb_order) + except Exception as e: + print(img_path, e) + elif self.backend == 'cv2': + image = cv2.imdecode( + np.array(bytearray(image_data), dtype='uint8'), + cv2.IMREAD_COLOR) + if self.rgb_order == 'RGB': + try: + cv2.cvtColor(image, cv2.COLOR_BGR2RGB, image) + except Exception as e: + print(img_path, e) + else: + image = Image.open(img_path).convert(self.rgb_order) + image_np = np.asarray(image).transpose( + (2, 0, 1)) # Tensor type needs shape to be (C, H, W) + image = torch.from_numpy(image_np) + return image + + def __call__(self, item): + for key in self.file_keys: + if 'prefix' in item['meta']: + img_path = os.path.join(item['meta']['prefix'], + item['meta'][key]) + else: + img_path = item['meta'][key] + item[key.replace('_path', '')] = self.read_file(img_path) + item['meta']['rgb_order'] = self.rgb_order + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + + return dict_to_yaml('TRANSFORM', + __class__.__name__, + LoadImageFromFileList.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class LoadPILImageFromFile(object): + def __init__(self, cfg, logger=None): + rgb_order = cfg.get('RGB_ORDER', 'RGB') + assert rgb_order in ('RGB', 'BGR') + self.rgb_order = rgb_order + + def __call__(self, item): + if 'prefix' in item['meta']: + img_path = os.path.join(item['meta']['prefix'], + item['meta']['img_path']) + else: + img_path = item['meta']['img_path'] + + with FS.get_from(img_path) as img_path: + image = Image.open(img_path).convert(self.rgb_order) + item['img'] = image + item['meta']['rgb_order'] = self.rgb_order + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + para_dict = [{ + 'RGB_ORDER': { + 'value': 'RGB', + 'description': 'rgb order' + } + }] + return dict_to_yaml('TRANSFORM', + __class__.__name__, + para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class LoadCvImageFromFile(object): + def __init__(self, cfg, logger=None): + rgb_order = cfg.get('RGB_ORDER', 'RGB') + assert rgb_order in ('RGB', 'BGR') + self.rgb_order = rgb_order + + def __call__(self, item): + if 'prefix' in item['meta']: + img_path = os.path.join(item['meta']['prefix'], + item['meta']['img_path']) + else: + img_path = item['meta']['img_path'] + + with FS.get_from(img_path) as img_path: + image = cv2.imread(img_path, cv2.IMREAD_COLOR) + if self.rgb_order == 'RGB': + cv2.cvtColor(image, cv2.COLOR_BGR2RGB, image) + item['img'] = image + item['meta']['rgb_order'] = self.rgb_order + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + para_dict = [{ + 'RGB_ORDER': { + 'value': 'RGB', + 'description': 'rgb order' + } + }] + return dict_to_yaml('TRANSFORM', + __class__.__name__, + para_dict, + set_name=True) diff --git a/scepter/modules/transform/io_video.py b/scepter/modules/transform/io_video.py new file mode 100644 index 0000000..b0a7170 --- /dev/null +++ b/scepter/modules/transform/io_video.py @@ -0,0 +1,486 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import numbers +import os.path as osp +import queue +import random +import threading + +import numpy as np +import torch + +from scepter.modules.transform import LoadImageFromFile +from scepter.modules.transform.registry import TRANSFORMS +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.file_system import DATA_FS as FS +from scepter.modules.utils.video_reader.frame_sampler import do_frame_sample +from scepter.modules.utils.video_reader.video_reader import VideoReaderWrapper + + +def _interval_based_sampling(vid_length, + vid_fps, + target_fps, + clip_idx, + num_clips, + num_frames, + interval, + minus_interval=False): + """ Generates the frame index list using interval based sampling. + + Args: + vid_length (int): The length of the whole video (valid selection range). + vid_fps (float): The original video fps. + target_fps (int): The target decode fps. + clip_idx (int): -1 for random temporal sampling, and positive values for sampling specific clip from the video. + num_clips (int): The total clips to be sampled from each video. + Combined with clip_idx, the sampled video is the "clip_idx-th" video from "num_clips" videos. + num_frames (int): Number of frames in each sampled clips. + interval (int): The interval to sample each frame. + minus_interval (bool): + + Returns: + index (torch.Tensor): The sampled frame indexes. + """ + if num_frames == 1: + index = [random.randint(0, vid_length - 1)] + else: + # transform FPS + clip_length = num_frames * interval * vid_fps / target_fps + + max_idx = max(vid_length - clip_length, 0) + if clip_idx == -1: # random sampling + start_idx = random.uniform(0, max_idx) + else: + if num_clips == 1: + start_idx = max_idx / 2 + else: + start_idx = max_idx * clip_idx / num_clips + if minus_interval: + end_idx = start_idx + clip_length - interval + else: + end_idx = start_idx + clip_length - 1 + + index = torch.linspace(start_idx, end_idx, num_frames) + index = torch.clamp(index, 0, vid_length - 1).long() + + return index + + +def _segment_based_sampling(vid_length, clip_idx, num_clips, num_frames, + random_sample): + """ Generates the frame index list using segment based sampling. + + Args: + vid_length (int): The length of the whole video (valid selection range). + clip_idx (int): -1 for random temporal sampling, and positive values for sampling specific clip from the video. + num_clips (int): The total clips to be sampled from each video. + Combined with clip_idx, the sampled video is the "clip_idx-th" video from "num_clips" videos. + num_frames (int): Number of frames in each sampled clips. + random_sample (bool): Whether or not to randomly sample from each segment. True for train and False for test. + + Returns: + index (torch.Tensor): The sampled frame indexes. + """ + index = torch.zeros(num_frames) + index_range = torch.linspace(0, vid_length, num_frames + 1) + for idx in range(num_frames): + if random_sample: + index[idx] = random.uniform(index_range[idx], index_range[idx + 1]) + else: + if num_clips == 1: + index[idx] = (index_range[idx] + index_range[idx + 1]) / 2 + else: + index[idx] = index_range[idx] + (index_range[ + idx + 1] - index_range[idx]) * (clip_idx + 1) / num_clips + index = torch.round(torch.clamp(index, 0, vid_length - 1)).long() + + return index + + +R = threading.Lock() + + +@TRANSFORMS.register_class() +class LoadVideoFromFile(object): + """ Open video file, extract frames, convert to tensor. + + Args: + num_frames (int): T dimension value. + sample_type (str): See + `from essmc2.metric.video_reader import FRAME_SAMPLERS; print(FRAME_SAMPLERS)` to get candidates, + default is 'interval'. + clip_duration (Optional[float]): Needed for 'interval' sampling type. + decoder (str): Video decoder name, default is decord. + + """ + def __init__(self, cfg, logger=None): + self.num_frames = cfg.NUM_FRAMES + self.sample_type = cfg.get('SAMPLE_TYPE', 'interval') + self.clip_duration = cfg.get('CLIP_DURATION', None) + + assert self.sample_type in ('uniform', 'interval', 'segment'), \ + f'Expected sample type in (uniform, interval, segment), got {self.sample_type}' + if self.sample_type == 'interval': + assert isinstance(self.clip_duration, numbers.Number), \ + 'Interval style sampling needs clip_duration not None' + self.decoder = cfg.get('DECODER', 'decord') + + def __call__(self, item): + """ + Args: + item (dict): + item['meta']['prefix'] (Optional[str]): Prefix of video_path. + item['meta']['video_path'] (str): Required. + item['meta']['clip_id'] (Optional[int]): Multi-view test needs it, default is 0. + item['meta']['num_clips'] (Optional[int]): Multi-view test needs it, default is 1. + item['meta']['start_sec'] (Optional[float]): Uniform sampling needs it. + item['meta']['end_sec'] (Optional[float]): Uniform sampling needs it. + + Returns: + item(dict): + item['video'] (torch.Tensor): a THWC tensor. + """ + meta = item['meta'] + video_path = meta['video_path'] if 'prefix' not in meta else osp.join( + meta['prefix'], meta['video_path']) + + with FS.get_from(video_path) as local_path: + vr = VideoReaderWrapper(local_path, decoder=self.decoder) + + params = dict() + clip_id = meta.get('clip_id') or 0 + num_clips = meta.get('num_clips') or 1 + if self.sample_type == 'interval': + # default is test mode for interval and segment + params.update(clip_duration=self.clip_duration, + clip_id=clip_id, + num_clips=num_clips) + elif self.sample_type == 'segment': + # default is test mode for interval and segment + params.update(clip_id=clip_id, num_clips=num_clips) + else: + # uniform, needs start_sec, clip_duration or end_sec + start_sec = meta['start_sec'] - meta['start_sec'] + if 'end_sec' in meta: + end_sec = meta['end_sec'] - meta['start_sec'] + elif self.clip_duration is not None: + end_sec = start_sec + self.clip_duration + else: + raise ValueError( + 'Uniform sampling needs start_sec & end_sec / start_sec & clip_duration' + ) + params.update(start_sec=start_sec, end_sec=end_sec) + + decode_list = do_frame_sample(self.sample_type, vr.len, vr.fps, + self.num_frames, **params) + item['video'] = vr.sample_frames(decode_list) + + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + para_dict = [{ + 'SAMPLE_TYPE': { + 'value': 'interval', + 'description': 'sample type' + }, + 'CLIP_DURATION': { + 'value': None, + 'description': 'clip duration' + } + }] + return dict_to_yaml('TRANSFORM', + __class__.__name__, + para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class LoadVideoFromFrameList(object): + """ extract frames, convert to tensor. + + Args: + num_frames (int): T dimension value. + sample_type (str): See + `from essmc2.metric.video_reader import FRAME_SAMPLERS; print(FRAME_SAMPLERS)` to get candidates, + default is 'interval'. + clip_duration (Optional[float]): Needed for 'interval' sampling type. + decoder (str): Video decoder name, default is decord. + """ + para_dict = [{ + 'NUM_FRAMES': { + 'value': 30, + 'description': 'clip length!' + }, + 'SAMPLE_TYPE': { + 'value': 'interval', + 'description': 'sample type' + }, + 'CLIP_DURATION': { + 'value': None, + 'description': 'clip duration' + }, + 'RGB_ORDER': { + 'value': 'RGB', + 'description': 'rgb order' + }, + 'BACKEND': { + 'value': 'pillow', + 'description': 'input backend' + } + }] + + def __init__(self, cfg, logger=None): + self.load_ins = LoadImageFromFile(cfg, logger=logger) + self.num_frames = cfg.NUM_FRAMES + self.sample_type = cfg.get('SAMPLE_TYPE', 'interval') + self.clip_duration = cfg.get('CLIP_DURATION', None) + + assert self.sample_type in ('uniform', 'interval', 'segment'), \ + f'Expected sample type in (uniform, interval, segment), got {self.sample_type}' + if self.sample_type == 'interval': + assert isinstance(self.clip_duration, numbers.Number), \ + 'Interval style sampling needs clip_duration not None' + + def __call__(self, item): + """ + Args: + item (dict): + item['meta']['video_path'] (str): Required. + item['meta']['clip_id'] (Optional[int]): Multi-view test needs it, default is 0. + item['meta']['num_clips'] (Optional[int]): Multi-view test needs it, default is 1. + item['meta']['start_sec'] (Optional[float]): Uniform sampling needs it. + item['meta']['end_sec'] (Optional[float]): Uniform sampling needs it. + item['meta']['frames'](list): all frames file path list + + Returns: + item(dict): + item['video'] (torch.Tensor): a THWC tensor. + """ + meta = item['meta'] + params = dict() + clip_id = meta.get('clip_id') or 0 + num_clips = meta.get('num_clips') or 1 + if self.sample_type == 'interval': + # default is test mode for interval and segment + params.update(clip_duration=self.clip_duration, + clip_id=clip_id, + num_clips=num_clips) + elif self.sample_type == 'segment': + # default is test mode for interval and segment + params.update(clip_id=clip_id, num_clips=num_clips) + else: + # uniform, needs start_sec, clip_duration or end_sec + start_sec = meta['start_sec'] - meta['start_sec'] + if 'end_sec' in meta: + end_sec = meta['end_sec'] - meta['start_sec'] + elif self.clip_duration is not None: + end_sec = start_sec + self.clip_duration + else: + raise ValueError( + 'Uniform sampling needs start_sec & end_sec / start_sec & clip_duration' + ) + params.update(start_sec=start_sec, end_sec=end_sec) + frames = meta['frames'] + fps = meta['fps'] + decode_list = do_frame_sample(self.sample_type, len(frames), fps, + self.num_frames, **params) + sample_frames = [{'meta': {'img_path': frame}} for frame in frames] + img_path_queue = queue.Queue() + [ + img_path_queue.put_nowait([idx, item]) + for idx, item in enumerate(sample_frames) + ] + img_queue = queue.Queue() + + def download_file(): + while not img_path_queue.empty(): + R.acquire() + try: + idx, item = img_path_queue.get_nowait() + except Exception: + R.release() + continue + R.release() + img_queue.put_nowait([idx, self.load_ins(item)]) + + threading_list = [] + for _ in range(8): + t = threading.Thread(target=download_file) + t.daemon = True + t.start() + threading_list.append(t) + [th.join() for th in threading_list] + # print(f"one video download time {time.time() - st}") + sample_frames = [] + while not img_queue.empty(): + sample_frames.append(img_queue.get_nowait()) + sample_frames.sort(key=lambda x: x[0]) + # sample_frames = [self.load_ins(item) for item in sample_frames] + item['video'] = np.array( + [sample_frames[frame_id][1]['img'] for frame_id in decode_list]) + # item['video'] = item['video'].transpose([0, 3, 1, 2]) + item['meta'].pop('frames') + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + LoadVideoFromFrameList.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class DecodeVideoToTensor(object): + def __init__(self, cfg, logger=None): + """ DecodeVideoToTensor + Args: + num_frames (int): Decode frames number. + target_fps (int): Decode frames fps, default is 30. + sample_mode (str): Interval or segment sampling, default is interval. + sample_interval (int): Sample interval between output frames for interval sample mode, default is 4. + sample_minus_interval (bool): If minus interval for interval sample mode, default is False. + repeat (int): Number of clips to be decoded from each video, if repeat > 1, outputs will be named like + 'video-0', 'video-1'. Normally, 1 for classification task, 2 for contrastive learning. + """ + import decord + from decord import VideoReader + self.VideoReader = VideoReader + decord.bridge.set_bridge('torch') + self.num_frames = cfg.NUM_FRAMES + self.target_fps = cfg.get('TARGET_FPS', 30) + self.sample_mode = cfg.get('SAMPLE_MODE', 'interval') + self.sample_interval = cfg.get('SAMPLE_INTERVAL', 4) + self.sample_minus_interval = cfg.get('SAMPLE_MINUS_INTERVAL', False) + self.repeat = cfg.get('REPEAT', 1) + + def __call__(self, item): + """ Call to invoke decode + Args: + item (dict): A dict contains which file to decode and how to decode. + Normally, it has structure like + { + "meta": { + "prefix" (str, None): if not None, prefix will be added to video_path. + "video_path" (str): Absolute (prefix is None) or relative path. + "clip_idx" (int): -1 means random sampling, >=0 means do temporal crop. + "num_clips" (int): if clip_idx >= 0, clip_idx must < num_clips + } + } + + Returns: + A dict contains original input item and "video" tensor. + """ + meta = item['meta'] + video_path = meta['video_path'] \ + if 'prefix' not in meta else osp.join(meta['prefix'], meta['video_path']) + + with FS.get_from(video_path) as local_path: + vr = self.VideoReader(local_path) + # default is test mode + clip_id = meta.get('clip_id') or 0 + num_clips = meta.get('num_clips') or 1 + + vid_len = len(vr) + vid_fps = vr.get_avg_fps() + + frame_list = [] + for _ in range(self.repeat): + if self.sample_mode == 'interval': + decode_list = _interval_based_sampling( + vid_len, vid_fps, self.target_fps, clip_id, num_clips, + self.num_frames, self.sample_interval, + self.sample_minus_interval) + else: + decode_list = _segment_based_sampling( + vid_len, clip_id, num_clips, self.num_frames, + clip_id == -1) + + # Decord gives inconsistent result for avi files. Getting full frames will fix it, although slower. + # See https://github.com/dmlc/decord/issues/195 + if video_path.lower().endswith('avi'): + full_decode_list = list( + range(0, + torch.max(decode_list).item() + 1)) + full_frames = vr.get_batch(full_decode_list) + frames = full_frames[decode_list].clone() + else: + frames = vr.get_batch(decode_list).clone() + + frame_list.append(frames) + + if self.repeat == 1: + item['video'] = frame_list[0] + else: + for idx, frame_tensor in zip(range(self.repeat), frame_list): + item[f'video-{idx}'] = frame_tensor + + del vr + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + para_dict = [{ + 'NUM_FRAMES': { + 'value': 30, + 'description': 'num frame' + }, + 'TARGET_FPS': { + 'value': 30, + 'description': 'target fps' + }, + 'SAMPLE_MODE': { + 'value': 'interval', + 'description': 'sample mode' + }, + 'SAMPLE_INTERVAL': { + 'value': 4, + 'description': 'sample interval' + }, + 'SAMPLE_MINUS_INTERVAL': { + 'value': False, + 'description': 'sample minus interval' + }, + 'REPEAT': { + 'value': 1, + 'description': 'repeat' + } + }] + return dict_to_yaml('TRANSFORM', + __class__.__name__, + para_dict, + set_name=True) diff --git a/scepter/modules/transform/registry.py b/scepter/modules/transform/registry.py new file mode 100644 index 0000000..2788b56 --- /dev/null +++ b/scepter/modules/transform/registry.py @@ -0,0 +1,50 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.utils.config import Config +from scepter.modules.utils.registry import Registry, build_from_config + + +def build_pipeline(pipeline, registry, logger=None, *args, **kwargs): + + if isinstance(pipeline, list): + if len(pipeline) == 0: + return build_from_config(Config(cfg_dict={'NAME': 'Identity'}, + load=False), + registry, + logger=logger, + *args, + **kwargs) + elif len(pipeline) == 1: + return build_pipeline(pipeline[0], registry, logger, *args, + **kwargs) + else: + return build_from_config(Config(cfg_dict={ + 'NAME': 'Compose', + 'TRANSFORMS': pipeline + }, + load=False), + registry, + logger=logger, + *args, + **kwargs) + elif isinstance(pipeline, Config): + return build_from_config(pipeline, + registry, + logger=logger, + *args, + **kwargs) + elif pipeline is None: + return build_from_config(Config(cfg_dict={'NAME': 'Identity'}, + load=False), + registry, + logger=logger, + *args, + **kwargs) + else: + raise TypeError( + f'Expect pipeline_cfg to be dict or list or None, got {type(pipeline)}' + ) + + +TRANSFORMS = Registry('TRANSFORMS', build_func=build_pipeline) diff --git a/scepter/modules/transform/tensor.py b/scepter/modules/transform/tensor.py new file mode 100644 index 0000000..4612b2f --- /dev/null +++ b/scepter/modules/transform/tensor.py @@ -0,0 +1,198 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import numpy as np +import torch + +from scepter.modules.transform.registry import TRANSFORMS +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we + + +def to_tensor(data): + if isinstance(data, torch.Tensor): + return data + elif isinstance(data, np.ndarray): + return torch.from_numpy(data) + elif isinstance(data, list): + return torch.tensor(data) + elif isinstance(data, int): + return torch.LongTensor([data]) + elif isinstance(data, float): + return torch.FloatTensor([data]) + else: + raise TypeError(f'Unsupported type {type(data)}') + + +@TRANSFORMS.register_class() +class ToTensor(object): + def __init__(self, cfg, logger=None): + self.keys = cfg.KEYS + + def __call__(self, item): + for key in self.keys: + item[key] = to_tensor(item[key]) + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + para_dict = [{'KEYS': {'value': [], 'description': 'keys'}}] + return dict_to_yaml('TRANSFORM', + __class__.__name__, + para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class Select(object): + def __init__(self, cfg, logger=None): + self.keys = cfg.KEYS + meta_keys = cfg.get('META_KEYS', []) + if not isinstance(meta_keys, (list, tuple)): + raise TypeError( + f'Expected meta_keys to be list or tuple, got {type(meta_keys)}' + ) + self.meta_keys = meta_keys + + def __call__(self, item): + data = {} + for key in self.keys: + data[key] = item[key] + if 'meta' in item and len(self.meta_keys) > 0: + data['meta'] = {} + for key in self.meta_keys: + data['meta'][key] = item['meta'][key] + return data + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + para_dict = [{ + 'KEYS': { + 'value': [], + 'description': 'keys' + }, + 'META_KEYS': { + 'value': [], + 'description': 'meta keys' + } + }] + return dict_to_yaml('TRANSFORM', + __class__.__name__, + para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class Rename(object): + def __init__(self, cfg, logger=None): + self.in_keys = cfg.IN_KEYS + self.out_keys = cfg.OUT_KEYS + + def __call__(self, item): + data = {} + for idx, key in enumerate(self.in_keys): + data[self.out_keys[idx]] = item[key] + have_key_set = set(self.in_keys) + for k, v in item.items(): + if k not in have_key_set: + data[k] = v + return data + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + para_dict = [{ + 'IN_KEYS': { + 'value': [], + 'description': + 'The keys need to rename, the other keys are outputed by default.' + }, + 'OUT_KEYS': { + 'value': [], + 'description': + 'The keys need to rename, the other keys are outputed by default.' + } + }] + return dict_to_yaml('TRANSFORM', + __class__.__name__, + para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class TensorToGPU(object): + def __init__(self, cfg, logger=None): + self.keys = cfg.KEYS + self.device_id = we.rank + + def __call__(self, item): + ret = {} + for key, value in item.items(): + if key in self.keys and isinstance( + value, torch.Tensor) and torch.cuda.is_available(): + ret[key] = value.cuda(self.device_id, non_blocking=True) + else: + ret[key] = value + return ret + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + para_dict = [{ + 'KEYS': { + 'value': [], + 'description': 'keys' + }, + 'DEVICE_ID': { + 'value': + 0, + 'description': + "device id, which should be set according to current GPU's rank" + } + }] + return dict_to_yaml('TRANSFORM', + __class__.__name__, + para_dict, + set_name=True) diff --git a/scepter/modules/transform/transform_xl.py b/scepter/modules/transform/transform_xl.py new file mode 100644 index 0000000..b3ba963 --- /dev/null +++ b/scepter/modules/transform/transform_xl.py @@ -0,0 +1,84 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import numbers + +import numpy as np +import opencv_transforms.functional as cv2_TF +import torch +import torchvision.transforms.functional as TF +from PIL.Image import Image + +from scepter.modules.transform import TRANSFORMS, ImageTransform +from scepter.modules.transform.image import BACKENDS +from scepter.modules.transform.utils import BACKEND_PILLOW, BACKEND_TORCHVISION +from scepter.modules.utils.config import dict_to_yaml + + +@TRANSFORMS.register_class() +class FlexibleCropXL(ImageTransform): + para_dict = [{ + 'IS_CENTER': { + 'value': False, + 'description': 'Use center crop or not.' + } + }] + para_dict[0].update(ImageTransform.para_dict[0]) + + def __init__(self, cfg, logger=None): + super(FlexibleCropXL, self).__init__(cfg, logger=logger) + assert self.backend in BACKENDS + self.size = cfg.get('SIZE', None) + self.is_center = cfg.get('IS_CENTER', False) + if self.size is not None: + if isinstance(self.size, numbers.Number): + self.size = [self.size, self.size] + self.callable = TF.crop if self.backend in ( + BACKEND_PILLOW, BACKEND_TORCHVISION) else cv2_TF.crop + + def __call__(self, item): + if isinstance(self.input_key, str): + self.input_key = [self.input_key] + if isinstance(self.output_key, str): + self.output_key = [self.output_key] + meta = item.get('meta', {}) + for idx, key in enumerate(self.input_key): + self.check_image_type(item[key]) + if isinstance(item[key], (torch.Tensor, np.ndarray)): + if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION): + w, h, c = item[key].shape + else: + h, w, c = item[key].shape + elif isinstance(item[key], Image): + w, h = item[key].size + if 'image_size' in meta: + oh, ow = meta['image_size'] + else: + assert self.size is not None + oh, ow = self.size + meta['image_size'] = [oh, ow] + + delta_h = h - oh + delta_w = w - ow + if not self.is_center: + top = np.random.randint(0, delta_h + 1) + left = np.random.randint(0, delta_w + 1) + else: + top = delta_h // 2 + left = delta_w // 2 + + item[self.output_key[idx]] = self.callable(item[key], top, left, + oh, ow) + item[self.output_key[idx] + '_' + + 'original_size_as_tuple'] = torch.tensor([h, w]) + item[self.output_key[idx] + '_' + + 'target_size_as_tuple'] = torch.tensor([oh, ow]) + item[self.output_key[idx] + '_' + + 'crop_coords_top_left'] = torch.tensor([top, left]) + return item + + @staticmethod + def get_config_template(): + return dict_to_yaml('TRANSFORM', + __class__.__name__, + FlexibleCropXL.para_dict, + set_name=True) diff --git a/scepter/modules/transform/utils.py b/scepter/modules/transform/utils.py new file mode 100644 index 0000000..f842915 --- /dev/null +++ b/scepter/modules/transform/utils.py @@ -0,0 +1,66 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import cv2 +import numpy as np +import torch +from packaging import version +from PIL import Image +from torchvision.version import __version__ as tv_version + +try: + import accimage +except ImportError: + accimage = None + + +def is_pil_image(img): + if accimage is not None: + return isinstance(img, (Image.Image, accimage.Image)) + else: + return isinstance(img, Image.Image) + + +def is_cv2_image(img): + return isinstance(img, np.ndarray) and img.dtype == np.uint8 + + +def is_tensor(t): + return isinstance(t, torch.Tensor) + + +INPUT_PIL_TYPE_WARNING = 'input should be PIL Image' +INPUT_CV2_TYPE_WARNING = 'input should be cv2 image(uint8 np.ndarray)' +INPUT_TENSOR_TYPE_WARNING = 'input should be tensor(uint8 np.ndarray)' + +# Recommend to use nn.Module backend to transform +TORCHVISION_CAPABILITY = version.parse(tv_version) >= version.parse('0.8.0') + +BACKEND_TORCHVISION = 'torchvision' +BACKEND_PILLOW = 'pillow' +BACKEND_CV2 = 'cv2' + +# Recommend to use InterpolationMode since torchvision 0.9.0 +INTERPOLATION_MODE_CAPABILITY = version.parse(tv_version) >= version.parse( + '0.9.0') +if INTERPOLATION_MODE_CAPABILITY: + from torchvision.transforms.functional import InterpolationMode +else: + import warnings + + warnings.filterwarnings('ignore', message='Default upsampling behavior.*') +INTERPOLATION_STYLE = { + 'bilinear': + Image.BILINEAR + if not INTERPOLATION_MODE_CAPABILITY else InterpolationMode('bilinear'), + 'nearest': + Image.NEAREST + if not INTERPOLATION_MODE_CAPABILITY else InterpolationMode('nearest'), + 'bicubic': + Image.BICUBIC + if not INTERPOLATION_MODE_CAPABILITY else InterpolationMode('bicubic'), +} +INTERPOLATION_STYLE_CV2 = { + 'bilinear': cv2.INTER_LINEAR, + 'nearest': cv2.INTER_NEAREST, + 'bicubic': cv2.INTER_CUBIC, +} diff --git a/scepter/modules/transform/video.py b/scepter/modules/transform/video.py new file mode 100644 index 0000000..af5d00a --- /dev/null +++ b/scepter/modules/transform/video.py @@ -0,0 +1,561 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import random + +import numpy as np +import torch +import torchvision.transforms.functional as functional +import torchvision.transforms.transforms as transforms +from packaging import version +from torchvision.version import __version__ as tv_version + +from scepter.modules.transform.registry import TRANSFORMS +from scepter.modules.transform.utils import (BACKEND_TORCHVISION, + INTERPOLATION_STYLE, is_tensor) +# torchvision.transform._transforms_video is deprecated since torchvision 0.10.0, use transform instead +from scepter.modules.utils.config import dict_to_yaml + +use_video_transforms = version.parse(tv_version) < version.parse('0.10.0') + +BACKENDS = (BACKEND_TORCHVISION, ) + + +class VideoTransform(object): + para_dict = [{ + 'INPUT_KEY': { + 'value': 'img', + 'description': 'input key' + }, + 'OUTPUT_KEY': { + 'value': 'img', + 'description': 'input key' + }, + 'BACKEND': { + 'value': 'pillow', + 'description': 'backend' + } + }] + + def __init__(self, cfg, logger=None): + backend = cfg.get('BACKEND', BACKEND_TORCHVISION) + self.input_key = cfg.get('INPUT_KEY', 'video') + self.output_key = cfg.get('OUTPUT_KEY', 'video') + self.backend = backend + + def check_video_type(self, input_video): + if self.backend == BACKEND_TORCHVISION: + assert is_tensor(input_video) + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + VideoTransform.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class RandomResizedCropVideo(VideoTransform): + """Crop a random portion of video and resize it to a given size. + + Expect the video is a torch tensor with shape [..., H, W] + + Args: + size (int or sequence): Desired output size. + If size is a sequence like (h, w), the output size will be matched to this. + If size is an int, the output size will be matched to (size, size). + scale (tuple of float): Specifies the lower and upper bounds for the random area of the crop, + before resizing. The scale is defined with respect to the area of the original image. + ratio (tuple of float): lower and upper bounds for the random aspect ratio of the crop, before + resizing. + interpolation (str): Desired interpolation string, 'bilinear', 'nearest', 'bicubic' are supported. + """ + para_dict = [{ + 'SIZE': { + 'value': 0, + 'description': 'size' + }, + 'SCALE': { + 'value': [0.08, 1.0], + 'description': 'scale' + }, + 'RATIO': { + 'value': [3. / 4., 4. / 3.], + 'description': 'ratio' + }, + 'INTERPOLATION': { + 'value': 'bilinear', + 'description': 'interpolation' + } + }] + para_dict[0].update(VideoTransform.para_dict[0]) + + def __init__(self, cfg, logger=None): + size, scale = cfg.SIZE, cfg.get('SCALE', [0.08, 1.0]) + ratio = cfg.get('RATIO', [3. / 4., 4. / 3.]) + interpolation = cfg.get('INTERPOLATION', 'bilinear') + super(RandomResizedCropVideo, self).__init__(cfg, logger=logger) + assert self.backend in BACKENDS + self.interpolation = interpolation + if isinstance(size, (tuple, list)): + assert len(size) == 2 + size = tuple(size) + elif isinstance(size, int): + size = (size, size) + else: + raise ValueError( + f'Unexpected type {type(size)}, expected int or tuple or list') + + if use_video_transforms: + from torchvision.transforms._transforms_video import \ + RandomResizedCropVideo as RandomResizedCropVideoOp + self.callable = RandomResizedCropVideoOp(size, scale, ratio, + self.interpolation) + else: + self.callable = transforms.RandomResizedCrop( + size, scale, ratio, INTERPOLATION_STYLE[self.interpolation]) + + def __call__(self, item): + self.check_video_type(item[self.input_key]) + item[self.output_key] = self.callable(item[self.input_key]) + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + RandomResizedCropVideo.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class CenterCropVideo(VideoTransform): + """ Crops the given video at the center. + + Expect the video is a torch tensor with shape [..., H, W] + + Args: + size (sequence or int): Desired output size. + If size is a sequence like (h, w), the output size will be matched to this. + If size is an int, the output size will be matched to (size, size). + """ + para_dict = [{'SIZE': {'value': 0, 'description': 'size'}}] + para_dict[0].update(VideoTransform.para_dict[0]) + + def __init__(self, cfg, logger=None): + size = cfg.SIZE + super(CenterCropVideo, self).__init__(cfg, logger=logger) + assert self.backend in BACKENDS + self.size = size + + if use_video_transforms: + from torchvision.transforms._transforms_video import \ + CenterCropVideo as CenterCropVideoOp + self.callable = CenterCropVideoOp(size) + else: + self.callable = transforms.CenterCrop(size) + + def __call__(self, item): + self.check_video_type(item[self.input_key]) + item[self.output_key] = self.callable(item[self.input_key]) + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + CenterCropVideo.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class RandomHorizontalFlipVideo(VideoTransform): + """ Horizontally flip the given video randomly with a given probability. + + Expect the video is a torch tensor with shape [..., H, W] + + Args: + p (float): probability of the image being flipped. Default value is 0.5 + """ + para_dict = [{'P': {'value': 0.5, 'description': 'P'}}] + para_dict[0].update(VideoTransform.para_dict[0]) + + def __init__(self, cfg, logger=None): + p = cfg.get('P', 0.5) + super(RandomHorizontalFlipVideo, self).__init__(cfg, logger=logger) + assert self.backend in BACKENDS + + if use_video_transforms: + from torchvision.transforms._transforms_video import \ + RandomHorizontalFlipVideo as RandomHorizontalFlipVideoOp + self.callable = RandomHorizontalFlipVideoOp(p) + else: + self.callable = transforms.RandomHorizontalFlip(p) + + def __call__(self, item): + self.check_video_type(item[self.input_key]) + item[self.output_key] = self.callable(item[self.input_key]) + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + RandomHorizontalFlipVideo.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class NormalizeVideo(VideoTransform): + """ Normalize a tensor video with mean and standard deviation. + Expect the video is a torch tensor with shape [..., H, W] + + Args: + mean (sequence): Sequence of means for each channel. + std (sequence): Sequence of standard deviations for each channel. + """ + para_dict = [{ + 'MEAN': { + 'value': [], + 'description': 'mean' + }, + 'STD': { + 'value': [], + 'description': 'std' + } + }] + para_dict[0].update(VideoTransform.para_dict[0]) + + def __init__(self, cfg, logger=None): + mean, std = cfg.MEAN, cfg.STD + super(NormalizeVideo, self).__init__(cfg, logger=logger) + assert self.backend in BACKENDS + self.mean = np.array(mean, dtype=np.float32) + self.std = np.array(std, dtype=np.float32) + + if use_video_transforms: + from torchvision.transforms._transforms_video import \ + NormalizeVideo as NormalizeVideoOp + self.callable = NormalizeVideoOp(self.mean, self.std) + else: + self.callable = transforms.Normalize(self.mean, self.std) + + def __call__(self, item): + video = item[self.input_key] + if not use_video_transforms: + video = video.permute(1, 0, 2, 3) + video = self.callable(video) + if not use_video_transforms: + video = video.permute(1, 0, 2, 3) + item[self.output_key] = video + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + NormalizeVideo.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class VideoToTensor(VideoTransform): + """ Convert a uint8 type tensor to a float32 tensor, permute it and scale output to [0.0, 1.0]. + + Expect the video is a uint8 torch tensor with shape [T, H, W, C] + """ + para_dict = [{}] + para_dict[0].update(VideoTransform.para_dict[0]) + + def __init__(self, cfg, logger=None): + super(VideoToTensor, self).__init__(cfg, logger=logger) + assert self.backend in BACKENDS + + def __call__(self, item): + video = item[self.input_key] + if isinstance(video, np.ndarray): + video = torch.tensor(video) + + if not torch.is_tensor(video): + raise TypeError('video should be Tensor. Got %s' % type(video)) + + if not video.ndimension() == 4: + raise ValueError('video should be 4D. Got %dD' % video.dim()) + + if not video.dtype == torch.uint8: + raise TypeError( + 'video tensor should have data type uint8. Got %s' % + str(video.dtype)) + + item[self.output_key] = video.float().permute(3, 0, 1, 2) / 255.0 + + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + VideoToTensor.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class AutoResizedCropVideo(VideoTransform): + """ Crop the video with a given position and resize it to a given size. + + Expect the video is a torch tensor with shape [..., H, W]. + + Input ``crop_mode`` supports values: + - `cc`: center-center + - `cl`: left-center + - `cr`: right-center + - `tl`: left-top + - `tr`: right-top + - `bl`: left-bottom + - `br`: right-bottom + + Args: + size (int or sequence): Desired output size. + If size is a sequence like (h, w), the output size will be matched to this. + If size is an int, the output size will be matched to (size, size). + scale (tuple of float): Specifies the lower and upper bounds for the random area of the crop, + before resizing. The scale is defined with respect to the area of the original image. + interpolation (str): Desired interpolation string, 'bilinear', 'nearest', 'bicubic' are supported. + """ + para_dict = [{'SCALE': {'value': [0.08, 1.0], 'description': 'scale'}}] + para_dict[0].update(VideoTransform.para_dict[0]) + + def __init__(self, cfg, logger=None): + size, scale = cfg.SIZE, cfg.get('SCALE', [0.08, 1.0]) + interpolation = cfg.get('INTERPOLATION', 'bilinear') + super(AutoResizedCropVideo, self).__init__(cfg, logger=logger) + if isinstance(size, (tuple, list)): + assert len(size) == 2 + size = tuple(size) + elif isinstance(size, int): + size = (size, size) + else: + raise ValueError( + f'Unexpected type {type(size)}, expected int or tuple or list') + self.size = size + self.scale = scale + self.interpolation_mode = interpolation + + def get_crop(self, clip, crop_mode='cc'): + scale = random.uniform(*self.scale) + _, _, video_height, video_width = clip.shape + min_length = min(video_height, video_width) + crop_size = int(min_length * scale) + center_x = video_width // 2 + center_y = video_height // 2 + box_half = crop_size // 2 + + # default is cc + x0 = center_x - box_half + y0 = center_y - box_half + if crop_mode == 'cl': + x0 = 0 + y0 = center_y - box_half + elif crop_mode == 'cr': + x0 = video_width - crop_size + y0 = center_y - box_half + elif crop_mode == 'tl': + x0 = 0 + y0 = 0 + elif crop_mode == 'tr': + x0 = video_width - crop_size + y0 = 0 + elif crop_mode == 'bl': + x0 = 0 + y0 = video_height - crop_size + elif crop_mode == 'br': + x0 = video_width - crop_size + y0 = video_height - crop_size + + if use_video_transforms: + from torchvision.transforms.functional import resized_crop + return resized_crop(clip, y0, x0, crop_size, crop_size, self.size, + self.interpolation_mode) + else: + return functional.resized_crop( + clip, y0, x0, crop_size, crop_size, self.size, + INTERPOLATION_STYLE[self.interpolation_mode]) + + def __call__(self, item): + self.check_video_type(item[self.input_key]) + crop_mode = item['meta'].get('crop_mode') or 'cc' + item[self.output_key] = self.get_crop(item[self.input_key], crop_mode) + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + AutoResizedCropVideo.para_dict, + set_name=True) + + +@TRANSFORMS.register_class() +class ResizeVideo(VideoTransform): + """Resize video to a given size. + + Expect the video is a torch tensor with shape [..., H, W]. + + Args: + size (int or sequence): Desired output size. + If size is a sequence like (h, w), the output size will be matched to this. + If size is an int, the smaller edge of the image will be matched to this + number maintaining the aspect ratio. + interpolation (str): Desired interpolation string, 'bilinear', 'nearest', 'bicubic' are supported. + """ + para_dict = [{ + 'SCALE': { + 'value': [0.08, 1.0], + 'description': 'scale' + }, + 'INTERPOLATION': { + 'value': 'bilinear', + 'description': 'interpolation' + } + }] + para_dict[0].update(VideoTransform.para_dict[0]) + + def __init__(self, cfg, logger=None): + size = cfg.SIZE + interpolation = cfg.get('INTERPOLATION', 'bilinear') + super(ResizeVideo, self).__init__(cfg, logger=logger) + self.size = size + if isinstance(self.size, (tuple, list)): + self.size = tuple(self.size) + assert len(self.size) == 2 + else: + if not isinstance(self.size, int): + raise ValueError( + f'Expected size to be tuple or list or int, got {type(self.size)}' + ) + self.interpolation_mode = interpolation + + def resize(self, clip): + if use_video_transforms: + from torchvision.transforms.functional import resize + + # resize function only takes a tuple size + # so we need to compute scaled target size here + if isinstance(self.size, int): + h, w = clip.shape[-2], clip.shape[-1] + if (w <= h and w == self.size) or (h <= w and h == self.size): + return clip + if w < h: + ow = self.size + oh = int(self.size * h / w) + else: + oh = self.size + ow = int(self.size * w / h) + size = (oh, ow) + else: + size = self.size + return resize(clip, size, self.interpolation_mode) + else: + return functional.resize( + clip, self.size, INTERPOLATION_STYLE[self.interpolation_mode]) + + def __call__(self, item): + self.check_video_type(item[self.input_key]) + item[self.output_key] = self.resize(item[self.input_key]) + return item + + @staticmethod + def get_config_template(): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + :return: + ''' + return dict_to_yaml('TRANSFORM', + __class__.__name__, + ResizeVideo.para_dict, + set_name=True) diff --git a/scepter/modules/utils/__init__.py b/scepter/modules/utils/__init__.py new file mode 100644 index 0000000..c243788 --- /dev/null +++ b/scepter/modules/utils/__init__.py @@ -0,0 +1,3 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.utils import config, distribute, file_clients, file_system diff --git a/scepter/modules/utils/config.py b/scepter/modules/utils/config.py new file mode 100644 index 0000000..947c232 --- /dev/null +++ b/scepter/modules/utils/config.py @@ -0,0 +1,604 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import argparse +import copy +import json +import os +import sys + +import yaml + +from scepter.modules.utils.model import StdMsg + + +def dict_to_yaml(module_name, name, json_config, set_name=False): + ''' + { "ENV" : + { "description" : "", + "A" : { + "value": 1.0, + "description": "" + } + } + } + convert std dict to yaml + :param module_name: + :param json_config: + :return: + ''' + def convert_yaml_style(level=1, + name='ENV', + description='ENV PARA', + default='', + type_name='', + is_sys=False): + new_line = '' + new_line += '{}# {} DESCRIPTION: {} TYPE: {} default: {}\n'.format( + '\t' * (level - 1), name.upper(), description, type_name, + f'\'{default}\'' if isinstance(default, str) else default) + if is_sys: + if name == '-': + new_line += '{}{}\n'.format('\t' * (level - 1), name.upper()) + else: + new_line += '{}{}:\n'.format('\t' * (level - 1), name.upper()) + else: + # if isinstance(default, str): + # default = f'\'{default}\'' + if default is None: + new_line += '{}# {}: {}\n'.format('\t' * (level - 1), + name.upper(), default) + else: + new_line += '{}{}: {}\n'.format('\t' * (level - 1), + name.upper(), default) + return new_line + + def parse_dict(json_config, + level_num, + parent_key, + set_name=False, + name='', + parent_type='dict'): + yaml_str = '' + # print(level_num, json_config) + if isinstance(json_config, dict): + if 'value' in json_config: + value = json_config['value'] + if isinstance(value, dict): + assert len(value) < 1 + value = None + description = json_config.get('description', '') + yaml_str += convert_yaml_style(level=level_num - 1, + name=parent_key, + description=description, + default=value, + type_name=type(value).__name__) + return True, yaml_str + else: + if len(json_config) < 1: + yaml_str += convert_yaml_style(level=level_num, + name='NAME', + description='', + default='', + type_name='') + level_num += 1 + for k, v in json_config.items(): + if k == 'description': + continue + if isinstance(v, dict): + is_final, new_yaml_str = parse_dict(v, + level_num, + k, + parent_type='dict') + if not is_final and parent_type == 'dict': + description = v.get('description', '') + yaml_str += convert_yaml_style( + level=level_num - 1, + name=k, + description=description, + default='', + type_name='', + is_sys=True) + if not is_final and parent_type == 'list': + yaml_str += convert_yaml_style(level=level_num, + name='NAME', + description='', + default=k, + type_name='') + yaml_str += new_yaml_str + elif isinstance(v, list): + base_yaml_str = convert_yaml_style(level=level_num - 1, + name=k, + description='', + default='', + type_name='', + is_sys=True) + yaml_str += base_yaml_str + for tup in v: + is_final, new_yaml_str = parse_dict( + tup, level_num, '-', parent_type='list') + if not is_final: + yaml_str += convert_yaml_style(level=level_num, + name='-', + description='', + default='', + type_name='', + is_sys=True) + yaml_str += new_yaml_str + else: + raise KeyError( + f'json config {json_config} must be a dict of list' + ) + + elif isinstance(json_config, list): + level_num += 1 + for tup in json_config: + is_final, new_yaml_str = parse_dict(tup, level_num, '-') + if not is_final: + + yaml_str += convert_yaml_style(level=level_num - 1, + name='-', + description='', + default='', + type_name='', + is_sys=True) + if set_name: + yaml_str += convert_yaml_style(level=level_num, + name='NAME', + description='', + default=name, + type_name='') + yaml_str += new_yaml_str + else: + raise KeyError(f'json config {json_config} must be a dict') + return False, yaml_str + + if isinstance(json_config, dict): + first_dict, sec_dict, third_dict = {}, {}, {} + for key, value in json_config.items(): + if isinstance(value, dict) and len(value) > 0: + first_dict[key] = value + elif isinstance(value, dict) and len(value) == 0: + sec_dict[key] = value + elif isinstance(value, list): + third_dict[key] = value + else: + raise f'Config {json_config} is illegal' + json_config = {} + json_config.update(first_dict) + json_config.update(sec_dict) + json_config.update(third_dict) + + yaml_str = f'[{module_name}] module yaml examples:\n' + level_num = 1 + base_yaml_str = convert_yaml_style(level=level_num, + name=module_name, + description='', + default='', + type_name='', + is_sys=True) + level_num += 1 + + is_final, new_yaml_str = parse_dict(json_config, + level_num, + module_name, + set_name=isinstance(json_config, list) + and set_name, + name=name) + if not is_final: + yaml_str += base_yaml_str + if set_name and not isinstance(json_config, list): + yaml_str += convert_yaml_style(level=level_num, + name='NAME', + description='', + default=name, + type_name='') + yaml_str += new_yaml_str + else: + yaml_str += new_yaml_str[1:] + + return yaml_str + + +def _parse_args(parser): + if parser is None: + parser = argparse.ArgumentParser( + description='Argparser for My codebase:\n') + else: + assert isinstance(parser, argparse.ArgumentParser) + parser.add_argument('--cfg', + dest='cfg_file', + help='Path to the configuration file', + required=False, + default=None) + parser.add_argument('--local_rank', + dest='local_rank', + help='torch distributed launch args!', + default=0) + parser.add_argument( + '-l', + '--launcher', + dest='launcher', + help='spawn launcher is using python scripts, torchrun launcher is ' + 'using torchrun module, default is spawn!', + default='spawn') + + parser.add_argument('-o', + '--data_online', + dest='data_online', + action='store_false', + help='Read data from online or save local as cache. ' + 'Default is from online.') + + parser.add_argument('-s', + '--share_storage', + dest='share_storage', + action='store_true', + help='If use nas as the common cache folder, ' + 'set True to avoid download conflict.') + parser.add_argument('--debug', + dest='debug', + action='store_true', + help='Swich debug mode.') + + return parser.parse_args() + + +class Config(object): + def __init__(self, + cfg_dict={}, + load=True, + cfg_file=None, + logger=None, + parser_ins=None): + ''' + support to parse json/dict/yaml_file of parameters. + :param load: whether load parameters or not. + :param cfg_dict: default None. + :param cfg_level: default None, means the current cfg-level for recurrent cfg presentation. + :param logger: logger instance for print the cfg log. + one examples: + import argparse + parser = argparse.ArgumentParser( + description="Argparser for Cate process:\n" + ) + parser.add_argument( + "--stage", + dest="stage", + help="Running stage!", + default="train", + choices=["train"] + ) + + cfg = Config(load=True, parser_ins=parser) + ''' + # checking that the logger exists or not + if logger is None: + self.logger = StdMsg(name='Config') + else: + self.logger = logger + self.cfg_dict = cfg_dict + if load: + if cfg_file is None: + assert parser_ins is not None + self.args = _parse_args(parser_ins) + self.load_from_file(self.args.cfg_file) + # os.environ["LAUNCHER"] = self.args.launcher + os.environ['DATA_ONLINE'] = str(self.args.data_online).lower() + os.environ['SHARE_STORAGE'] = str( + self.args.share_storage).lower() + os.environ['ES_DEBUG'] = str(self.args.debug).lower() + else: + self.load_from_file(cfg_file) + if 'ENV' not in self.cfg_dict: + self.cfg_dict['ENV'] = { + 'SEED': 2023, + 'USE_PL': False, + 'BACKEND': 'nccl', + 'SYNC_BN': False, + 'CUDNN_DETERMINISTIC': True, + 'CUDNN_BENCHMARK': False + } + self.logger.info( + f"ENV is not set and will use default ENV as {self.cfg_dict['ENV']}; " + f'If want to change this value, please set them in your config.' + ) + else: + if 'SEED' not in self.cfg_dict['ENV']: + self.cfg_dict['ENV']['SEED'] = 2023 + self.logger.info( + f"SEED is not set and will use default SEED as {self.cfg_dict['ENV']['SEED']}; " + f'If want to change this value, please set it in your config.' + ) + os.environ['ES_SEED'] = str(self.cfg_dict['ENV']['SEED']) + self._update_dict(self.cfg_dict) + if load: + self.logger.info(f'Parse cfg file as \n {self.dump()}') + + def load_from_file(self, file_name): + self.logger.info(f'Loading config from {file_name}') + if file_name is None or not os.path.exists(file_name): + self.logger.info(f'File {file_name} does not exist!') + self.logger.warning( + f"Cfg file is None or doesn't exist, Skip loading config from {file_name}." + ) + return + if file_name.endswith('.json'): + self.cfg_dict = self._load_json(file_name) + self.logger.info( + f'System take {file_name} as json, because we find json in this file' + ) + elif file_name.endswith('.yaml'): + self.cfg_dict = self._load_yaml(file_name) + self.logger.info( + f'System take {file_name} as yaml, because we find yaml in this file' + ) + else: + self.logger.info( + f'No config file found! Because we do not find json or yaml in --cfg {file_name}' + ) + + def _update_dict(self, cfg_dict): + def recur(key, elem): + if type(elem) is dict: + return key, Config(load=False, + cfg_dict=elem, + logger=self.logger) + elif type(elem) is list: + config_list = [] + for idx, ele in enumerate(elem): + if type(ele) is str and ele[1:3] == 'e-': + ele = float(ele) + config_list.append(ele) + elif type(ele) is str: + config_list.append(ele) + elif type(ele) is dict: + config_list.append( + Config(load=False, + cfg_dict=ele, + logger=self.logger)) + elif type(ele) is list: + config_list.append(ele) + else: + config_list.append(ele) + return key, config_list + else: + if type(elem) is str and elem[1:3] == 'e-': + elem = float(elem) + return key, elem + + dic = dict(recur(k, v) for k, v in cfg_dict.items()) + self.__dict__.update(dic) + + def _load_json(self, cfg_file): + ''' + :param cfg_file: + :return: + ''' + if cfg_file is None: + self.logger.warning( + f'Cfg file is None, Skip loading config from {cfg_file}.') + return {} + file_name = cfg_file + try: + cfg = json.load(open(file_name, 'r')) + except Exception as e: + self.logger.error(f'Load json from {cfg_file} error. Message: {e}') + sys.exit() + return cfg + + def _load_yaml(self, cfg_file): + ''' + if replace some parameters from Base, You can reference the base parameters use Base. + + :param cfg_file: + :return: + ''' + if cfg_file is None: + self.logger.warning( + f'Cfg file is None, Skip loading config from {cfg_file}.') + return {} + file_name = cfg_file + try: + with open(cfg_file, 'r') as f: + cfg = yaml.load(f.read(), Loader=yaml.SafeLoader) + except Exception as e: + self.logger.error(f'Load yaml from {cfg_file} error. Message: {e}') + sys.exit() + if '_BASE_RUN' not in cfg.keys() and '_BASE_MODEL' not in cfg.keys( + ) and '_BASE' not in cfg.keys(): + return cfg + + if '_BASE' in cfg.keys(): + if cfg['_BASE'][1] == '.': + prev_count = cfg['_BASE'].count('..') + cfg_base_file = self._path_join( + file_name.split('/')[:(-1 - cfg['_BASE'].count('..'))] + + cfg['_BASE'].split('/')[prev_count:]) + else: + cfg_base_file = cfg['_BASE'].replace( + './', file_name.replace(file_name.split('/')[-1], '')) + cfg_base = self._load_yaml(cfg_base_file) + cfg = self._merge_cfg_from_base(cfg_base, cfg) + else: + if '_BASE_RUN' in cfg.keys(): + if cfg['_BASE_RUN'][1] == '.': + prev_count = cfg['_BASE_RUN'].count('..') + cfg_base_file = self._path_join( + file_name.split('/')[:(-1 - prev_count)] + + cfg['_BASE_RUN'].split('/')[prev_count:]) + else: + cfg_base_file = cfg['_BASE_RUN'].replace( + './', file_name.replace(file_name.split('/')[-1], '')) + cfg_base = self._load_yaml(cfg_base_file) + cfg = self._merge_cfg_from_base(cfg_base, + cfg, + preserve_base=True) + if '_BASE_MODEL' in cfg.keys(): + if cfg['_BASE_MODEL'][1] == '.': + prev_count = cfg['_BASE_MODEL'].count('..') + cfg_base_file = self._path_join( + file_name.split('/')[:( + -1 - cfg['_BASE_MODEL'].count('..'))] + + cfg['_BASE_MODEL'].split('/')[prev_count:]) + else: + cfg_base_file = cfg['_BASE_MODEL'].replace( + './', file_name.replace(file_name.split('/')[-1], '')) + cfg_base = self._load_yaml(cfg_base_file) + cfg = self._merge_cfg_from_base(cfg_base, cfg) + return cfg + + def _path_join(self, path_list): + path = '' + for p in path_list: + path += p + '/' + return path[:-1] + + def items(self): + return self.cfg_dict.items() + + def _merge_cfg_from_base(self, cfg_base, cfg, preserve_base=False): + for k, v in cfg.items(): + if k in cfg_base.keys(): + if isinstance(v, dict): + self._merge_cfg_from_base(cfg_base[k], v) + else: + cfg_base[k] = v + else: + if 'BASE' not in k or preserve_base: + cfg_base[k] = v + return cfg_base + + def _merge_cfg_from_command(self, args, cfg): + assert len( + args.opts + ) % 2 == 0, f'Override list {args.opts} has odd length: {len(args.opts)}' + + keys = args.opts[0::2] + vals = args.opts[1::2] + + # maximum supported depth 3 + for idx, key in enumerate(keys): + key_split = key.split('.') + assert len( + key_split + ) <= 4, 'Key depth error. \n Maximum depth: 3\n Get depth: {}'.format( + len(key_split)) + assert key_split[0] in cfg.keys(), 'Non-existant key: {}.'.format( + key_split[0]) + if len(key_split) == 2: + assert key_split[1] in cfg[ + key_split[0]].keys(), 'Non-existant key: {}'.format(key) + elif len(key_split) == 3: + assert key_split[1] in cfg[ + key_split[0]].keys(), 'Non-existant key: {}'.format(key) + assert key_split[2] in cfg[key_split[0]][ + key_split[1]].keys(), 'Non-existant key: {}'.format(key) + elif len(key_split) == 4: + assert key_split[1] in cfg[ + key_split[0]].keys(), 'Non-existant key: {}'.format(key) + assert key_split[2] in cfg[key_split[0]][ + key_split[1]].keys(), 'Non-existant key: {}'.format(key) + assert key_split[3] in cfg[key_split[0]][key_split[1]][ + key_split[2]].keys(), 'Non-existant key: {}'.format(key) + + if len(key_split) == 1: + cfg[key_split[0]] = vals[idx] + elif len(key_split) == 2: + cfg[key_split[0]][key_split[1]] = vals[idx] + elif len(key_split) == 3: + cfg[key_split[0]][key_split[1]][key_split[2]] = vals[idx] + elif len(key_split) == 4: + cfg[key_split[0]][key_split[1]][key_split[2]][ + key_split[3]] = vals[idx] + + return cfg + + def __repr__(self): + return '{}\n'.format(self.dump()) + + def dump(self): + return json.dumps(self.cfg_dict, indent=2) + + def deep_copy(self): + return copy.deepcopy(self) + + def have(self, name): + if name in self.__dict__: + return True + return False + + def get(self, name, default=None): + if name in self.__dict__: + return self.__dict__[name] + return default + + def __getitem__(self, key): + return self.__dict__.__getitem__(key) + + def __setattr__(self, key, value): + super().__setattr__(key, value) + if hasattr(self, 'cfg_dict') and key in self.cfg_dict: + if isinstance(value, Config): + value = value.cfg_dict + self.cfg_dict[key] = value + + def __setitem__(self, key, value): + self.__dict__[key] = value + self.__setattr__(key, value) + + def __iter__(self): + return iter(self.__dict__) + + def set(self, name, value): + new_dict = {name: value} + self.__dict__.update(new_dict) + self.__setattr__(name, value) + + def get_dict(self): + return self.cfg_dict + + def get_lowercase_dict(self, cfg_dict=None): + if cfg_dict is None: + cfg_dict = self.get_dict() + config_new = {} + for key, val in cfg_dict.items(): + if isinstance(key, str): + if isinstance(val, dict): + config_new[key.lower()] = self.get_lowercase_dict(val) + else: + config_new[key.lower()] = val + else: + config_new[key] = val + return config_new + + @staticmethod + def get_plain_cfg(cfg=None): + if isinstance(cfg, Config): + cfg_new = {} + cfg_dict = cfg.get_dict() + for key, val in cfg_dict.items(): + if isinstance(val, (Config, dict, list)): + cfg_new[key] = Config.get_plain_cfg(val) + else: + cfg_new[key] = val + return cfg_new + elif isinstance(cfg, dict): + cfg_new = {} + cfg_dict = cfg + for key, val in cfg_dict.items(): + if isinstance(val, (Config, dict, list)): + cfg_new[key] = Config.get_plain_cfg(val) + else: + cfg_new[key] = val + return cfg_new + elif isinstance(cfg, list): + cfg_new = [] + cfg_list = cfg + for val in cfg_list: + if isinstance(val, (Config, dict, list)): + cfg_new.append(Config.get_plain_cfg(val)) + else: + cfg_new.append(val) + return cfg_new + else: + return cfg diff --git a/scepter/modules/utils/data.py b/scepter/modules/utils/data.py new file mode 100644 index 0000000..b4e734a --- /dev/null +++ b/scepter/modules/utils/data.py @@ -0,0 +1,88 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from collections import OrderedDict + +import torch + + +def transfer_data_to_numpy(data_map: dict) -> dict: + """ Transfer tensors in data_map to numpy type. + Will recursively walk through inner list, tuple and dict values. + + Args: + data_map (dict): a dictionary which contains tensors to be transferred + + Returns: + A dict which has same structure with input `data_map`. + """ + if not isinstance(data_map, dict): + return data_map + ret = OrderedDict() + for key, value in data_map.items(): + if isinstance(value, torch.Tensor): + ret[key] = value.detach().cpu().numpy() + elif isinstance(value, dict): + ret[key] = transfer_data_to_numpy(value) + elif isinstance(value, (list, tuple)): + ret[key] = type(value)([transfer_data_to_numpy(t) for t in value]) + else: + ret[key] = value + return ret + + +def transfer_data_to_cpu(data_map: dict) -> dict: + """ Transfer tensors in data_map to cpu device. + Will recursively walk through inner list, tuple and dict values. + + Args: + data_map (dict): a dictionary which contains tensors to be transferred + + Returns: + A dict which has same structure with input `data_map`. + """ + if not isinstance(data_map, dict): + return data_map + ret = OrderedDict() + for key, value in data_map.items(): + if isinstance(value, torch.Tensor): + ret[key] = value.detach().cpu() + elif isinstance(value, dict): + ret[key] = transfer_data_to_cpu(value) + elif isinstance(value, (list, tuple)): + ret[key] = type(value)([transfer_data_to_cpu(t) for t in value]) + else: + ret[key] = value + torch.cuda.empty_cache() + return ret + + +def transfer_data_to_cuda(data_map: dict) -> dict: + """ Transfer tensors in data_map to current default gpu device. + Will recursively walk through inner list, tuple and dict values. + + Args: + data_map (dict): a dictionary which contains tensors to be transferred + + Returns: + A dict which has same structure with input `data_map`. + """ + import platform + if platform.system() == 'Darwin': + return data_map + if not isinstance(data_map, dict): + return data_map + ret = OrderedDict() + for key, value in data_map.items(): + if isinstance(value, torch.Tensor): + if value.is_cuda: + ret[key] = value + else: + ret[key] = value.cuda(non_blocking=True) + elif isinstance(value, dict): + ret[key] = transfer_data_to_cuda(value) + elif isinstance(value, (list, tuple)): + ret[key] = type(value)([transfer_data_to_cuda(t) for t in value]) + else: + ret[key] = value + return ret diff --git a/scepter/modules/utils/directory.py b/scepter/modules/utils/directory.py new file mode 100644 index 0000000..16392f5 --- /dev/null +++ b/scepter/modules/utils/directory.py @@ -0,0 +1,21 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import hashlib +import os.path as osp + + +def osp_path(prefix, data_file): + if data_file.startswith(prefix): + return data_file + else: + return osp.join(prefix, data_file) + + +def get_relative_folder(abs_path, keep_index=-1): + path_tup = abs_path.split('/')[:keep_index] + return '/'.join(path_tup) + + +def get_md5(ori_str): + md5 = hashlib.md5(ori_str.encode()).hexdigest() + return md5 diff --git a/scepter/modules/utils/distribute.py b/scepter/modules/utils/distribute.py new file mode 100644 index 0000000..abdefb4 --- /dev/null +++ b/scepter/modules/utils/distribute.py @@ -0,0 +1,455 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import functools +import os +import pickle +import random +import warnings +from collections import OrderedDict + +import numpy as np +import torch +import torch.distributed as dist + +from scepter.modules.utils.model import StdMsg + +__all__ = [ + 'gather_data', 'we', 'broadcast', 'barrier', 'reduce_scatter', 'reduce', + 'all_reduce', 'send', 'recv', 'isend', 'irecv', 'scatter', + 'shared_random_seed' +] + +try: + from onnxruntime.transformers.benchmark_helper import set_random_seed +except Exception: + + def set_random_seed(seed): + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + np.random.seed(seed) + random.seed(seed) + + +def get_dist_info(): + if dist.is_available() and dist.is_initialized(): + return dist.get_rank(), dist.get_world_size() + else: + return 0, 1 + + +def gather_data(data): + """ Gather tensors and other picklable objects to rank 0. + Will recursively walk through inner list and dict values. + + Args: + data (any): Anything. + + Returns: + A object has same structure with input `data`. + """ + if not we.is_distributed: + return data + if isinstance(data, torch.Tensor): + return gather_gpu_tensors(data) + elif isinstance(data, dict): + # Keep in order, dict type DO NOT guarantee a fixed key order + keys = sorted(list(data.keys())) + ret = OrderedDict() + for key in keys: + ret[key] = gather_data(data[key]) + return ret + elif isinstance(data, list): + return gather_list(data) + else: + return gather_picklable(data) + + +def gather_list(data): + """ Gather list of picklable objects to a new list on rank 0. + Will NOT recursively walk through. + + Args: + data (list): List of picklable things. + + Returns: + A new flat list. + """ + rank, _ = get_dist_info() + list_of_list = gather_picklable(data) + if rank == 0: + return sum(list_of_list, []) + + +def gather_picklable(data): + """ Gather picklable object to a list on rank 0. + Will NOT recursively walk through. + + Args: + data (picklable): Picklable data. + + Returns: + A list contains data collected. + """ + from packaging import version + from torch.version import __version__ + if version.parse(__version__) < version.parse('1.8.0'): + return _gather_picklable_custom(data) + else: + rank, world_size = we.rank, we.world_size + obj_list = [None for _ in range(world_size)] + dist.all_gather_object(obj_list, data) + if rank == 0: + return obj_list + + +def _gather_picklable_custom(data): + """ Custom implementation function to gather picklable object to a list on rank 0. + If torch version is lower than 1.8.0, use this. + + Args: + data (picklable): Picklable data. + + Returns: + A list contains data collected. + """ + import pickle + byte_tensor = torch.tensor(bytearray(pickle.dumps(data)), + dtype=torch.uint8, + device='cuda') + rank, world_size = we.rank, we.world_size + shape_tensor = torch.tensor(byte_tensor.shape, device='cuda') + shape_list = [shape_tensor.clone() for _ in range(world_size)] + dist.all_gather(shape_list, shape_tensor) + shape_max = torch.tensor(shape_list).max() + + tensor_send = torch.zeros(shape_max, + dtype=byte_tensor.dtype, + device='cuda') + tensor_send[0:shape_tensor[0]] = byte_tensor + tensor_list = [torch.zeros_like(tensor_send) for _ in range(world_size)] + dist.all_gather(tensor_list, tensor_send) + + if rank == 0: + data_out = [] + for tensor_recv, shape_recv in zip(tensor_list, shape_list): + new_data = pickle.loads( + tensor_recv[:shape_recv[0]].cpu().numpy().tobytes()) + data_out.append(new_data) + return data_out + + +def gather_gpu_tensors(tensor, all_recv=False, is_cat=True): + """ + Args: + tensor (torch.Tensor): + all_recv: Gather tensor to rank 0 and concat it. + + Returns: + A new tensor. + """ + assert dist.get_backend() == 'nccl' + + device = tensor.device + if device.type == 'cpu': + tensor = tensor.to(we.device_id) + + rank, world_size = we.rank, we.world_size + + shape_tensor = torch.tensor(tensor.shape[0], device='cuda') + shape_list = [shape_tensor.clone() for _ in range(world_size)] + dist.all_gather(shape_list, shape_tensor) + shape_max = torch.tensor(shape_list).max() + + tensor_send = torch.zeros((shape_max, *tensor.shape[1:]), + dtype=tensor.dtype, + device='cuda') + tensor_send[0:tensor.shape[0]] = tensor + tensor_list = [torch.zeros_like(tensor_send) for _ in range(world_size)] + dist.all_gather(tensor_list, tensor_send) + if not all_recv: + if rank == 0: + if not is_cat: + return tensor_list, shape_list + tensors_out = [] + for tensor_recv, shape_recv in zip(tensor_list, shape_list): + tensors_out.append(tensor_recv[0:shape_recv]) + tensor_out = torch.cat(tensors_out).contiguous() + if device.type == 'cpu': + tensor_out = tensor_out.cpu() + del tensor_list, shape_list + return tensor_out + else: + del tensor_list, shape_list + else: + if not is_cat: + return tensor_list, shape_list + tensors_out = [] + for tensor_recv, shape_recv in zip(tensor_list, shape_list): + tensors_out.append(tensor_recv[0:shape_recv]) + tensor_out = torch.cat(tensors_out).contiguous() + if device.type == 'cpu': + tensor_out = tensor_out.cpu() + del tensor_list, shape_list + return tensor_out + + +def broadcast(tensor, src, group=None, **kwargs): + if we.is_distributed: + return dist.broadcast(tensor, src, group, **kwargs) + + +def barrier(): + if we.is_distributed: + dist.barrier() + + +@functools.lru_cache() +def get_global_gloo_group(): + backend = dist.get_backend() + assert backend in ['gloo', 'nccl'] + if backend == 'nccl': + return dist.new_group(backend='gloo') + else: + return dist.group.WORLD + + +def reduce_scatter(output, + input_list, + op=dist.ReduceOp.SUM, + group=None, + **kwargs): + if we.is_distributed: + return dist.reduce_scatter(output, input_list, op, group, **kwargs) + + +def all_reduce(tensor, op=dist.ReduceOp.SUM, group=None, **kwargs): + if we.is_distributed: + return dist.all_reduce(tensor, op, group, **kwargs) + + +def reduce(tensor, dst, op=dist.ReduceOp.SUM, group=None, **kwargs): + if we.is_distributed: + return dist.reduce(tensor, dst, op, group, **kwargs) + + +def _serialize_to_tensor(data): + buffer = pickle.dumps(data) + storage = torch.ByteStorage.from_buffer(buffer) + tensor = torch.ByteTensor(storage) + return tensor + + +def _unserialize_from_tensor(recv_data): + buffer = recv_data.cpu().numpy().tobytes() + return pickle.loads(buffer) + + +def send(tensor, dst, group=None, **kwargs): + if we.is_distributed: + assert tensor.is_contiguous( + ), 'ops.send requires the tensor to be contiguous()' + return dist.send(tensor, dst, group, **kwargs) + + +def recv(tensor, src=None, group=None, **kwargs): + if we.is_distributed: + assert tensor.is_contiguous( + ), 'ops.recv requires the tensor to be contiguous()' + return dist.recv(tensor, src, group, **kwargs) + + +def isend(tensor, dst, group=None, **kwargs): + if we.is_distributed: + assert tensor.is_contiguous( + ), 'ops.isend requires the tensor to be contiguous()' + return dist.isend(tensor, dst, group, **kwargs) + + +def irecv(tensor, src=None, group=None, **kwargs): + if we.is_distributed: + assert tensor.is_contiguous( + ), 'ops.irecv requires the tensor to be contiguous()' + return dist.irecv(tensor, src, group, **kwargs) + + +def scatter(data, scatter_list=None, src=0, group=None, **kwargs): + r"""NOTE: only supports CPU tensor communication. + """ + world_size = we.world_size + if world_size == 1: + data.copy_(scatter_list[0]) + if group is None: + group = get_global_gloo_group() + return dist.scatter(data, scatter_list, src, group, **kwargs) + + +def shared_random_seed(): + seed = np.random.randint(2**31) + all_seeds, _ = gather_gpu_tensors(seed, all_recv=True, is_cat=False) + return all_seeds[0] + + +global we + + +def mp_worker(gpu, ngpus_per_node, cfg, fn, pmi_rank, world_size, work_env): + rank = pmi_rank * ngpus_per_node + gpu + work_env.device_id = gpu % ngpus_per_node + work_env.rank = rank + dist.init_process_group(backend='nccl', world_size=world_size, rank=rank) + torch.backends.cudnn.deterministic = cfg.ENV.get('CUDNN_DETERMINISTIC', + True) + torch.backends.cudnn.benchmark = cfg.ENV.get('CUDNN_BENCHMARK', False) + torch.cuda.set_device(work_env.device_id) + if work_env.logger is not None: + work_env.logger.info( + f'Now running in the distributed environment with world size {work_env.world_size}!' + ) + work_env.logger.info(f'PMI rank {pmi_rank}!') + work_env.logger.info(f'Nums of gpu {ngpus_per_node}!') + work_env.logger.info( + f'Current rank {work_env.rank} current devices num {ngpus_per_node} ' + f'current machine rank {pmi_rank} and all world size {world_size}') + + we.set_env(work_env.get_env()) + fn(cfg) + + +class Workenv(object): + def __init__(self): + self.initialized = False + self.is_distributed = False + self.sync_bn = False + self.rank = 0 + self.world_size = 1 + self.device_id = 0 + self.device_count = 1 + self.seed = 2023 + self.debug = False + self.use_pl = False + self.launcher = 'spawn' + self.data_online = False + self.share_storage = False + + def init_env(self, config, fn, logger=None): + # if use pytorch_lightning: then direct use pytorch_lightning. + config.ENV = config.get('ENV', {}) + self.seed = config.ENV.get('SEED', 2023) + self.debug = os.environ.get('ES_DEBUG', None) == 'true' + set_random_seed(self.seed) + if logger is not None: + logger.info(f'And running with seed {self.seed}!') + if config.ENV.get('USE_PL', False): + self.use_pl = config.ENV.USE_PL + fn(config) + return + if hasattr(config, 'args') and hasattr(config.args, 'launcher'): + self.launcher = config.args.launcher + if logger is None: + self.logger = StdMsg(name='env') + else: + self.logger = logger + + self.data_online = os.environ.get('DATA_ONLINE', None) == 'true' + self.share_storage = os.environ.get('SHARE_STORAGE', None) == 'true' + + if not torch.cuda.is_available(): + self.device_id = 'cpu' + fn(config) + return + + if (os.environ.get('WORLD_SIZE') is None or os.environ.get('WORLD_SIZE') == 1) \ + and torch.cuda.device_count() == 1 and not self.launcher == 'dist': + self.device_id = 0 + fn(config) + return + + if self.launcher == 'torchrun': + try: + torch.multiprocessing.set_start_method('spawn') + except Exception as e: + warnings.warn(f'{e}') + # checking train mode is distributed or not + if not os.environ.get('WORLD_SIZE') is None: + if self.logger is not None: + self.logger.info( + f"Now running in the distributed environment with {os.environ.get('WORLD_SIZE')}!" + ) + self.is_distributed = True + if not self.initialized: + if self.is_distributed: + self.backend = config.ENV.get('BACKEND', 'nccl') + self.sync_bn = config.ENV.get('SYNC_BN', False) + dist.init_process_group(backend=self.backend) + # dist.barrier() + self.initialized = True + if dist.is_initialized(): + self.rank, self.world_size = dist.get_rank( + ), dist.get_world_size() + if self.logger is not None: + self.logger.info(f'And running in rank {self.rank}!') + self.logger.info( + f"And cuda visible devices {os.environ.get('CUDA_VISIBLE_DEVICES')}!" + ) + else: + self.rank, self.world_size = 0, 1 + local_devices = os.environ.get( + 'LOCAL_WORLD_SIZE') or torch.cuda.device_count() + local_devices = int(local_devices) + self.device_count = local_devices + self.device_id = self.rank % local_devices + self.logger.info(f"We's attributes: \n" + f' launcher {self.launcher} \n' + f' rank {self.rank} \n' + f' world size {self.world_size} \n' + f' device_id {self.device_id}') + torch.cuda.set_device(self.device_id) + torch.backends.cudnn.deterministic = config.ENV.get( + 'CUDNN_DETERMINISTIC', True) + torch.backends.cudnn.benchmark = config.ENV.get( + 'CUDNN_BENCHMARK', False) + fn(config) + else: + import torch.multiprocessing as mp + if 'MASTER_ADDR' not in os.environ: + os.environ['MASTER_ADDR'] = 'localhost' + if 'MASTER_PORT' not in os.environ: + os.environ['MASTER_PORT'] = '14567' + pmi_rank = int(os.environ.get('RANK', 0)) + pmi_world_size = int(os.environ.get('WORLD_SIZE', 1)) + ngpus_per_node = os.environ.get( + 'LOCAL_WORLD_SIZE') or torch.cuda.device_count() + ngpus_per_node = int(ngpus_per_node) + self.device_count = ngpus_per_node + world_size = ngpus_per_node * pmi_world_size + self.world_size = world_size + if self.world_size > 1: + self.is_distributed = True + self.initialized = True + if self.is_distributed: + self.backend = config.ENV.get('BACKEND', 'nccl') + self.sync_bn = config.ENV.get('SYNC_BN', False) + mp.spawn(mp_worker, + nprocs=ngpus_per_node, + args=(ngpus_per_node, config, fn, pmi_rank, world_size, + self)) + + def get_env(self): + return self.__dict__ + + def set_env(self, we_env): + for k, v in we_env.items(): + setattr(self, k, v) + + def __str__(self): + environ_str = f'Now running in the distributed environment with world size {self.world_size}\n!' + environ_str += f'Current pod have {self.device_count} devices!\n' + environ_str += f'Current task executes on device {self.device_id}!\n' + environ_str += f"Current task's global rank is {self.rank} \n" + environ_str += f"Current task's data online is set {self.data_online}" + environ_str += f"Current task's share storage is set {self.share_storage}" + environ_str += f"Current task's global seed is set {self.seed}" + return environ_str + + +we = Workenv() diff --git a/scepter/modules/utils/export_model.py b/scepter/modules/utils/export_model.py new file mode 100644 index 0000000..679e3ef --- /dev/null +++ b/scepter/modules/utils/export_model.py @@ -0,0 +1,113 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import io +from io import BytesIO + +import onnx +import onnxruntime +import torch +from torch.onnx import OperatorExportTypes + +from scepter.modules.utils.distribute import we + +type_map = { + 'float32': torch.float32, + 'float16': torch.float16, + 'int64': torch.int64, + 'int32': torch.int32, + 'int16': torch.int16, + 'int8': torch.int8 +} + + +@torch.no_grad() +def save_develop_model_multi_io(model, + input_size, + input_type, + input_name, + output_name, + limit, + save_onnx_path=None, + save_pt_path=None): + + # save aggregation + rank, word_size = we.rank, we.world_size + assert isinstance(input_type, list) + example = [] + dynamic_axes = {} + + for idx, type_name in enumerate(input_type): + assert type_name in type_map + torch_type = type_map[type_name] + size = input_size[idx] + if 'float' in type_name: + input_ex = torch.rand(tuple(size)).type(torch_type).to(rank) + elif 'int' in type_name: + input_ex = torch.randint(limit[idx][0], limit[idx][1], + tuple(size)).type(torch_type).to(rank) + example.append(input_ex) + dynamic_axes[input_name[idx]] = {0: 'batch_size'} + + if word_size > 0: + save_module = model.module + else: + save_module = model + + def _check_eval(module): + assert not module.training + + save_module.apply(_check_eval) + + if len(example) == 1: + input_example = example[0] + else: + input_example = tuple(example) + traced_script_module = torch.jit.trace(save_module, input_example) + + for p in traced_script_module.parameters(): + p.requires_grad = False + if len(example) == 1: + output = save_module(input_example) + else: + output = save_module(*input_example) + print('Ori output:', output) + + module = None + if save_pt_path is not None: + traced_script_module.save(save_pt_path) + module = torch.jit.load(io.BytesIO(open(save_pt_path, 'rb').read()), + map_location=torch.device(rank)) + if len(example) == 1: + output = module(input_example) + else: + output = module(*input_example) + print('PT output:', output) + + onnx_module = None + if save_onnx_path is not None: + # export the model to ONNX + with torch.autocast(device_type='cpu', + enabled=True, + dtype=torch.bfloat16): + with BytesIO() as f: + torch.onnx.export( + save_module, + input_example, + f, + operator_export_type=OperatorExportTypes.ONNX, + opset_version=11, + input_names=input_name, + output_names=output_name, + dynamic_axes=dynamic_axes, + export_params=True, + do_constant_folding=True) + onnx_model = onnx.load_from_string(f.getvalue()) + onnx.save(onnx_model, save_onnx_path) + onnx_module = onnxruntime.InferenceSession( + save_onnx_path, providers=['CUDAExecutionProvider']) + input_data = {} + for idx, ex in enumerate(example): + input_data[input_name[idx]] = ex.detach().cpu().numpy() + output_tensor = onnx_module.run(output_name, input_data) + print('ONNX_OUTPUT', output_tensor, output_tensor[0].shape) + return module, onnx_module diff --git a/scepter/modules/utils/file_clients/__init__.py b/scepter/modules/utils/file_clients/__init__.py new file mode 100644 index 0000000..84656dc --- /dev/null +++ b/scepter/modules/utils/file_clients/__init__.py @@ -0,0 +1,6 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.utils.file_clients.aliyun_oss_fs import AliyunOssFs +from scepter.modules.utils.file_clients.http_fs import HttpFs +from scepter.modules.utils.file_clients.local_fs import LocalFs +from scepter.modules.utils.file_clients.modelscope_fs import ModelscopeFs diff --git a/scepter/modules/utils/file_clients/aliyun_oss_fs.py b/scepter/modules/utils/file_clients/aliyun_oss_fs.py new file mode 100644 index 0000000..27de807 --- /dev/null +++ b/scepter/modules/utils/file_clients/aliyun_oss_fs.py @@ -0,0 +1,1003 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import copy +import json +import logging +import os +import os.path as osp +import queue +import random +import sys +import tempfile +import threading +import time +import warnings +from typing import Optional + +import oss2 +from oss2 import determine_part_size +from oss2.models import PartInfo +from tqdm import tqdm + +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.directory import get_md5 +from scepter.modules.utils.file_clients.base_fs import BaseFs +from scepter.modules.utils.file_clients.registry import FILE_SYSTEMS + + +def compute_size(bytes): + if bytes > 1024 * 1024 * 1024: + gyte = bytes / (1024 * 1024 * 1024) + return '{:.2f}G'.format(gyte) + elif bytes > 1024 * 1024: + gyte = bytes / (1024 * 1024) + return '{:.2f}M'.format(gyte) + elif bytes > 1024: + gyte = bytes / 1024 + return '{:.2f}K'.format(gyte) + else: + return bytes + + +def upload_process_bar(consumed_bytes, total_bytes): + sys.stdout.flush() + percent = 100 * consumed_bytes / total_bytes + sys.stdout.write('upload {:.2f}% [{}/{}]'.format( + percent, compute_size(consumed_bytes), compute_size(total_bytes))) + sys.stdout.flush() + sys.stdout.write('\r') + + +def download_process_bar(consumed_bytes, total_bytes): + sys.stdout.flush() + percent = 100 * consumed_bytes / total_bytes + sys.stdout.write('Download {:.2f}% [{}/{}]'.format( + percent, compute_size(consumed_bytes), compute_size(total_bytes))) + sys.stdout.flush() + sys.stdout.write('\r') + + +def process_msg(msg): + sys.stdout.write(' ' * 100) + sys.stdout.flush() + sys.stdout.write(msg) + sys.stdout.flush() + sys.stdout.write('\r') + sys.stdout.flush() + + +class OssLoggingHandler(logging.StreamHandler): + def __init__(self, log_file, cfg): + super(OssLoggingHandler, self).__init__() + self.cfg = cfg + self._log_file = log_file + self._sessions = {} + _bucket = self._init_bucket() + if _bucket.object_exists(self._log_file): + size = _bucket.get_object_meta(self._log_file).content_length + else: + size = 0 + self._pos = _bucket.append_object(self._log_file, size, '') + + def _init_bucket(self): + endpoint = self.cfg.ENDPOINT + bucket = self.cfg.BUCKET + ak = self.cfg.OSS_AK + sk = self.cfg.OSS_SK + # session + session = self._sessions.setdefault(f'{bucket}@{os.getpid()}', + oss2.Session()) + _bucket: oss2.Bucket = oss2.Bucket(oss2.Auth(ak, sk), + endpoint, + bucket, + session=session) + return _bucket + + def emit(self, record): + msg = self.format(record) + '\n' + for _ in range(5): + _bucket = self._init_bucket() + try: + self._pos = _bucket.append_object(self._log_file, + self._pos.next_position, msg) + break + except oss2.exceptions.PositionNotEqualToLength: + self._pos = _bucket.get_object_meta( + self._log_file).content_length + self._pos = _bucket.append_object(self._log_file, + self._pos.next_position, msg) + break + except Exception: + continue + + +@FILE_SYSTEMS.register_class() +class AliyunOssFs(BaseFs): + para_dict = { + 'ENDPOINT': { + 'value': '', + 'description': 'the oss endpoint' + }, + 'BUCKET': { + 'value': '', + 'description': 'the oss bucket' + }, + 'OSS_AK': { + 'value': '', + 'description': 'the oss ak' + }, + 'OSS_SK': { + 'value': '', + 'description': 'the oss sk' + }, + 'PREFIX': { + 'value': '', + 'description': 'the file system prefix!' + }, + 'WRITABLE': { + 'value': True, + 'description': 'this file system is writable or not!' + }, + 'CHECK_WRITABLE': { + 'value': False, + 'decription': 'check fs is writable or not!' + }, + 'RETRY_TIMES': { + 'value': 10, + 'description': 'for one file retry download or upload times!' + } + } + para_dict.update(BaseFs.para_dict) + + def __init__(self, cfg, logger=None): + super(AliyunOssFs, self).__init__(cfg, logger=logger) + prefix = cfg.get('PREFIX', None) + writable = cfg.get('WRITABLE', True) + check_writable = cfg.get('CHECK_WRITABLE', False) + retry_times = cfg.get('RETRY_TIMES', 10) + bucket = self.cfg.BUCKET + self._sessions = {} + _bucket = self._init_bucket() + self._fs_prefix = f'oss://{bucket}/' + ('' + if prefix is None else prefix) + self._prefix = f'oss://{bucket}/' + try: + _bucket.list_objects(max_keys=1) + except Exception as e: + warnings.warn( + f'Cannot list objects in {self._prefix}, please check auth information. \n{e}' + ) + self._retry_times = retry_times + + self._writable = writable + if check_writable: + self._writable = self._test_write(bucket, prefix) + + def _init_bucket(self): + endpoint = self.cfg.ENDPOINT + bucket = self.cfg.BUCKET + ak = self.cfg.OSS_AK + sk = self.cfg.OSS_SK + # session + session = self._sessions.setdefault(f'{bucket}@{os.getpid()}', + oss2.Session()) + _bucket: oss2.Bucket = oss2.Bucket(oss2.Auth(ak, sk), + endpoint, + bucket, + session=session) + return _bucket + + def _test_write(self, bucket, prefix) -> bool: + local_tmp_file = osp.join( + tempfile.gettempdir(), + f"oss_{bucket}_{'' if prefix is None else prefix}_try_test_write" + + ''.join([str(random.randint(1, 10) for _ in range(5))])) + with open(local_tmp_file, 'w') as f: + f.write('Try to write') + target_tmp_file = osp.join(self._prefix, osp.basename(local_tmp_file)) + status = self.put_object_from_local_file(local_tmp_file, + target_tmp_file) + if status: + self.remove(target_tmp_file) + return status + + def get_prefix(self) -> str: + return self._fs_prefix + + def support_write(self) -> bool: + return self._writable + + def support_link(self) -> bool: + return self._writable + + def get_meta(self, target_path): + key = osp.relpath(target_path, self._prefix) + retry, wait_retry = 0, 0 # noqa + try: + while retry < self._retry_times: + _bucket = self._init_bucket() + try: + meta = _bucket.get_object_meta(key) + etag = meta.etag + size = meta.content_length + break + except oss2.exceptions.NoSuchKey as e: + warnings.warn(f'Get file meta {e}') + return None, None + except Exception as e: + warnings.warn(f'Get file meta {e}') + retry += 1 + if retry >= self._retry_times: + return None, None + except Exception: + etag = '' + size = 100 + return etag, size + + def get_object_to_local_file(self, + target_path, + local_path=None, + wait_finish=False, + worker_id=-1) -> Optional[str]: + key = osp.relpath(target_path, self._prefix) + etag, size = self.get_meta(target_path) + if local_path is None: + local_path, is_tmp = self.map_to_local(target_path, etag=etag) + else: + is_tmp = False + + if wait_finish: + from scepter.modules.utils.distribute import we + wait_times = size / 1024 / 1024 + cur_time = 0 + interval = 30 + while not worker_id == 0 and not osp.exists( + local_path) and wait_times > 0: + if we.share_storage and we.rank == 0: + break + if not we.share_storage and (we.device_id == 0 + or we.device_id == 'cpu'): + break + if self.logger is not None: + self.logger.info( + f'GPU {we.device_id} have waited for {cur_time}s. Status: ' + f'the data {target_path} is downloaded to {local_path},' + f'share storage {we.share_storage}, device {we.device_id}, rank {we.rank}!' + ) + time.sleep(interval) + cur_time += interval + wait_times -= interval + if osp.exists(local_path) and osp.getsize(local_path) == size: + return local_path + + os.makedirs(osp.dirname(local_path), exist_ok=True) + retry, _ = 0, 0 + temp_file = local_path + '.{}_temp'.format(time.time()) + + _bucket = self._init_bucket() + if not (os.path.exists(local_path) + and osp.getsize(local_path) == size): + while retry < self._retry_times: + if size < 100 * 1024 * 1024: + try: + if wait_finish: + _bucket.get_object_to_file( + key, + temp_file, + progress_callback=download_process_bar) + else: + _bucket.get_object_to_file(key, temp_file) + break + except oss2.exceptions.NoSuchKey as e: + warnings.warn(f'Download {key} error {e}') + return None + except Exception as e: + warnings.warn(f'{e}') + retry += 1 + else: + try: + _ = self._download_object_multi_part(target_path, + temp_file, + chunk_size=50 * + 1024 * 1024) + break + except Exception as e: + retry += 1 + self.logger.info( + 'Download file {} error {} retry {} times!'.format( + target_path, e, retry)) + + if retry >= self._retry_times: + return None + + try: + if not os.path.exists(local_path): + os.rename(temp_file, local_path) + elif not osp.getsize(local_path) == size: + os.remove(local_path) + os.rename(temp_file, local_path) + except Exception as e: + warnings.warn(f'Download local path {e}') + if os.path.exists(temp_file): + try: + os.remove(temp_file) + except Exception: + warnings.warn('Remove temp file error!') + + if is_tmp: + self.add_temp_file(local_path) + return local_path + + def _download_object_multi_part(self, + target_path, + local_path, + thread_num=20, + chunk_size=10 * 1024) -> bool: + ''' + Args: + target_path: + local_path: + Returns: + ''' + key = osp.relpath(target_path, self._prefix) + chunk_list = self.get_object_chunk_list(target_path, + chunk_num=-1, + chunk_size=chunk_size) + if chunk_list is None or len(chunk_list) == 0: + self.logger.info("File {} isn't exists!".format(target_path)) + return False + object_size = chunk_list[-1][0] + chunk_list[-1][1] + slice_queue = queue.Queue() + ret_slice_queue = queue.Queue() + R = threading.Lock() + for chunk_id, chunk in enumerate(chunk_list): + slice_queue.put_nowait([chunk_id, chunk]) + process_msg( + f'Split {target_path} to {len(chunk_list)} parts to download!') + all_part_number = len(chunk_list) + cache_folder = f'{local_path}_cache' + os.makedirs(cache_folder, exist_ok=True) + _bucket = self._init_bucket() + + def download_one_part(key): + while not slice_queue.empty(): + R.acquire() + try: + if not slice_queue.empty(): + part_number, chunk = slice_queue.get_nowait() + data_status = True + else: + data_status = False + except Exception as e: + self.logger.info('Get chunk {} error {}'.format( + target_path, e)) + data_status = False + R.release() + if not data_status: + continue + temp_part_file = os.path.join(cache_folder, f'{part_number}') + retry = 0 + while retry < self._retry_times: + try: + data, end = self.get_object_stream( + target_path, chunk[0], chunk[1]) + with open(temp_part_file, 'wb') as f: + f.write(data) + R.acquire() + try: + ret_slice_queue.put_nowait({ + 'part_number': + part_number, + 'status': + True, + 'file_path': + temp_part_file + }) + except Exception as e: + self.logger.info('Ret status error {}'.format(e)) + R.release() + break + except Exception as e: + retry += 1 + self.logger.info( + 'Download part {} for {} error {} retry {} times!'. + format(part_number, key, e, retry)) + if retry >= self._retry_times: + R.acquire() + try: + ret_slice_queue.put_nowait({ + 'part_number': part_number, + 'status': False, + 'file_path': None + }) + except Exception as e: + self.logger.info('Ret status error {}'.format(e)) + R.release() + if ret_slice_queue.qsize() < all_part_number: + download_process_bar(chunk_size * ret_slice_queue.qsize(), + object_size) + else: + download_process_bar(object_size, object_size) + + threading_list = [] + for i in range(thread_num): + t = threading.Thread(target=download_one_part, args=(key, )) + t.daemon = True + t.start() + threading_list.append(t) + for thread in threading_list: + thread.join() + parts = [] + while not ret_slice_queue.empty(): + upload_status = ret_slice_queue.get_nowait() + if upload_status['status']: + parts.append( + (upload_status['part_number'], upload_status['file_path'])) + else: + return False + + parts.sort(key=lambda x: x[0]) + + with open(local_path, 'wb') as fw: + for part in tqdm(parts, desc='merge parts to file...'): + with open(part[1], 'rb') as f: + fw.write(f.read()) + try: + os.system('rm -rf {}'.format(cache_folder)) + except Exception as e: + self.logger.info('Remove cache error {}'.format(e)) + if _bucket.object_exists(key): + return True + else: + return False + + def check_folder(self, check_file): + meta_dict = json.load(open(check_file, 'r')) + for key, v in meta_dict.items(): + etag, size = self.get_meta(key) + if not etag == v: + return False + return True + + def _get_dir(self, + target_path, + local_path, + wait_finish=False, + meta_dict={}): + local_path = local_path.replace('/./', '/') + os.makedirs(local_path, exist_ok=True) + generator = self.walk_dir(target_path) + for file_name in generator: + if file_name == target_path or file_name == target_path + '/': + continue + local_file_name = os.path.join( + local_path, + file_name.split(target_path)[-1]).replace('/./', '/') + if not self.isdir(file_name): + etag, size = self.get_meta(file_name) + if local_file_name in meta_dict and meta_dict[ + local_file_name] == etag: + continue + process_msg(f'Download {file_name} to {local_file_name}....') + local_file_name = self.get_object_to_local_file( + file_name, local_file_name, wait_finish=wait_finish) + assert local_file_name is not None + meta_dict[file_name] = etag + else: + meta_dict.update( + self._get_dir(file_name, + local_file_name, + meta_dict=copy.deepcopy(meta_dict))) + return meta_dict + + def get_dir_to_local_dir(self, + target_path, + local_path=None, + wait_finish=False, + timeout=3600, + worker_id=-1) -> Optional[str]: + if not self.isdir(target_path): + self.logger.info( + f"{target_path} is not directory or doesn't exist.") + if not target_path.endswith('/'): + target_path += '/' + if local_path is None: + local_path, is_tmp = self.map_to_local(target_path) + else: + is_tmp = False + local_path = local_path.replace('/./', '/') + os.makedirs(local_path, exist_ok=True) + check_file = os.path.join(local_path, + f'{get_md5(target_path)}_data_meta.json') + if wait_finish: + from scepter.modules.utils.distribute import we + wait_times = timeout + while wait_times > 0 and not (osp.exists(check_file) + and self.check_folder(check_file)): + if we.share_storage and we.rank == 0: + break + if not we.share_storage and (we.device_id == 0 + or we.device_id == 'cpu'): + break + if self.logger is not None: + self.logger.info( + f'GPU {we.device_id} is waiting that ' + f'the data {target_path} is downloaded to {local_path}!' + ) + time.sleep(5) + wait_times -= 1 + if osp.exists(check_file) and self.check_folder(check_file): + return local_path + + if osp.exists(check_file): + meta_dict = json.load(open(check_file, 'r')) + else: + meta_dict = {} + meta_dict = self._get_dir(target_path, + local_path=local_path, + meta_dict=copy.deepcopy(meta_dict)) + json.dump(meta_dict, open(check_file, 'w')) + if is_tmp: + self.add_temp_file(local_path) + return local_path + + def put_object_from_local_file(self, local_path, target_path) -> bool: + key = osp.relpath(target_path, self._prefix) + retry = 0 + object_size = os.path.getsize(local_path) + while retry < self._retry_times: + _bucket = self._init_bucket() + if object_size <= 100 * 1024 * 1024: + try: + _bucket.put_object_from_file(key, local_path) + if _bucket.object_exists(key): + break + except Exception as e: + retry += 1 + self.logger.info( + 'Upload file {} error {} retry {} times!'.format( + target_path, e, retry)) + else: + try: + flag = self._put_object_multi_part(local_path, + key, + object_size=object_size) + return flag + except Exception as e: + retry += 1 + self.logger.info( + 'Upload file {} error {} retry {} times!'.format( + target_path, e, retry)) + + if retry >= self._retry_times: + return False + + return True + + def _put_object_multi_part(self, + local_path, + key, + object_size=-1, + thread_num=20) -> bool: + ''' + Args: + local_path: + key: + object_size: + + Returns: + + ''' + if object_size < 0: + object_size = os.path.getsize(local_path) + _bucket = self._init_bucket() + part_size = determine_part_size(object_size, + preferred_size=10 * 1024 * 1024) + upload_id = _bucket.init_multipart_upload(key).upload_id + parts = [] + fileobj = open(local_path, 'rb') + slice_queue = queue.Queue() + ret_slice_queue = queue.Queue() + R = threading.Lock() + part_number = 1 + offset = 0 + while offset < object_size: + num_to_upload = min(part_size, object_size - offset) + slice_queue.put_nowait([part_number, offset, num_to_upload]) + offset += num_to_upload + part_number += 1 + process_msg(f'Split {key} to {part_number - 1} parts to upload!') + all_part_number = part_number + + def upload_one_part(key, upload_id): + while not slice_queue.empty(): + R.acquire() + try: + if not slice_queue.empty(): + part_number, offset, num_to_upload = slice_queue.get_nowait( + ) + fileobj.seek(offset) + raw_data = fileobj.read(num_to_upload) + data_status = True + else: + data_status = False + except Exception as e: + self.logger.info('Seek file {} error {}'.format( + local_path, e)) + data_status = False + R.release() + if not data_status: + continue + retry = 0 + while retry < self._retry_times: + _bucket = self._init_bucket() + try: + result = _bucket.upload_part(key, upload_id, + part_number, raw_data) + R.acquire() + try: + ret_slice_queue.put_nowait({ + 'part_number': part_number, + 'status': True, + 'etag': result.etag + }) + except Exception as e: + self.logger.info('Ret status error {}'.format(e)) + R.release() + break + except Exception as e: + retry += 1 + self.logger.info( + 'Upload part {} for {} error {} retry {} times!'. + format(part_number, key, e, retry)) + if retry >= self._retry_times: + R.acquire() + try: + ret_slice_queue.put_nowait({ + 'part_number': part_number, + 'status': False, + 'etag': None + }) + except Exception as e: + self.logger.info('Ret status error {}'.format(e)) + R.release() + if ret_slice_queue.qsize() < all_part_number - 1: + upload_process_bar(part_size * ret_slice_queue.qsize(), + object_size) + else: + upload_process_bar(object_size, object_size) + + threading_list = [] + for i in range(thread_num): + t = threading.Thread(target=upload_one_part, args=(key, upload_id)) + t.daemon = True + t.start() + threading_list.append(t) + for thread in threading_list: + thread.join() + + while not ret_slice_queue.empty(): + upload_status = ret_slice_queue.get_nowait() + if upload_status['status']: + parts.append( + PartInfo(upload_status['part_number'], + upload_status['etag'])) + else: + return False + _ = _bucket.complete_multipart_upload(key, upload_id, parts) + try: + fileobj.close() + except Exception as e: + self.logger.info(f'{e}') + if _bucket.object_exists(key): + return True + else: + return False + + def get_object(self, target_path): + try: + retry = 0 + key = osp.relpath(target_path, self._prefix) + local_data = None + while retry < self._retry_times: + _bucket = self._init_bucket() + try: + local_data = _bucket.get_object(key).read() + break + except oss2.exceptions.NoSuchKey as e: + self.logger.info(f'Download local path {e}') + return None + except Exception as e: + self.logger.info(f'{e}') + retry += 1 + except Exception as e: + self.logger.error(f'Read {target_path} error {e}') + local_data = None + return local_data + + def get_object_stream(self, + target_path, + start, + size=10000, + delimiter=None): + end = start + size - 1 + retry = 0 + key = osp.relpath(target_path, self._prefix) + local_data = None + _bucket = self._init_bucket() + if not _bucket.object_exists(key): + return local_data, end + meta_data = _bucket.get_object_meta(key) + content_length = meta_data.content_length + end = min(end, content_length - 1) + + if start >= content_length - 1: + return local_data, None + + while retry < self._retry_times: + _bucket = self._init_bucket() + try: + local_data = _bucket.get_object(key, byte_range=(start, + end)).read() + break + except oss2.exceptions.NoSuchKey as e: + self.logger.info(f'Download local path {e}') + return None, end + except Exception as e: + self.logger.info(f'{e}') + retry += 1 + if delimiter is not None and not end == content_length - 1: + try: + total_len = len(local_data) + offset = 0 + try: + sp_data = local_data.split(bytes(delimiter, 'utf-8')) + # if failed, suppose the bytes is splited error. + except Exception as e: + self.logger.info( + f'Return data split error,please check your delimiter {e}' + ) + return None, end + if not len(sp_data[-1]) == len(local_data): + cur_offset = len(sp_data[-1]) + offset += cur_offset + local_data = local_data[:total_len - cur_offset] + local_data = local_data[len(sp_data[0]):] + end = end - offset + except Exception as e: + self.logger.info(f'Local data decode error {e}') + return local_data, end + 1 + + def get_object_chunk_list(self, + target_path, + chunk_num=-1, + chunk_size=-1, + delimiter=None): + chunk_st_et = [] + key = osp.relpath(target_path, self._prefix) + _bucket = self._init_bucket() + if not _bucket.object_exists(key): + self.logger.info(f'{target_path} is not exists!') + return chunk_st_et + if chunk_size < 0 and chunk_num < 0: + self.logger.error( + 'Suppose chunk size > 0 or chunk num > 0, instead of both < 0') + return None + meta_data = _bucket.get_object_meta(key) + content_length = meta_data.content_length + if chunk_size < 0 and chunk_num > 0: + chunk_size = content_length // chunk_num + 1 + # Don't care of the lines info + if delimiter is None: + chunk_start = 0 + # Iterate over all chunks and construct arguments for `process_chunk` + while chunk_start < content_length - 1: + chunk_end = min(content_length - 1, chunk_start + chunk_size) + chunk_st_et.append([chunk_start, chunk_end - chunk_start + 1]) + chunk_start = chunk_end + 1 + else: + chunk_start = 0 + offset = 0 + while chunk_start < content_length - 1: + chunk_end = min(content_length - 1, + chunk_start + chunk_size + offset) + retry = 0 + while retry < self._retry_times: + _bucket = self._init_bucket() + try: + quota_st = max(0, chunk_end - 20000) + quota_st = max(chunk_start, quota_st) + local_data = _bucket.get_object( + key, byte_range=(quota_st, chunk_end)).read() + break + except oss2.exceptions.NoSuchKey as e: + warnings.warn(f'Download local path {e}') + return None + except Exception as e: + warnings.warn(f'{e}') + retry += 1 + if not chunk_end == content_length - 1: + try: + offset = 0 + try: + sp_data = local_data.split( + bytes(delimiter, 'utf-8')) + # if failed, suppose the bytes is splited error. + except Exception as e: + self.logger.info( + f'Return data split error,please check your delimiter {e}' + ) + return None + if not len(sp_data[-1]) == len(local_data): + cur_offset = len(sp_data[-1]) + offset += cur_offset + chunk_end = chunk_end - offset + except Exception as e: + self.logger.info( + f'Local data decode error {e}, check your data is supported by str.decode().' + ) + chunk_st_et.append([chunk_start, chunk_end - chunk_start + 1]) + chunk_start = chunk_end + 1 + return chunk_st_et + + def put_object(self, local_data, target_path): + key = osp.relpath(target_path, self._prefix) + retry = 0 + while retry < self._retry_times: + _bucket = self._init_bucket() + try: + _bucket.put_object(key, local_data) + if _bucket.object_exists(key): + break + except Exception as e: + retry += 1 + self.logger.info('Upload file {} error {}'.format( + target_path, e)) + + if retry >= self._retry_times: + return False + + return True + + def get_url(self, + target_path, + set_public=False, + lifecycle=3600 * 100, + slash_safe=True): + key = osp.relpath(target_path, self._prefix) + _bucket = self._init_bucket() + if not _bucket.object_exists(key): + self.logger.info(f'{target_path} is not exists!') + return None + retry = 0 + while retry < self._retry_times: + _bucket = self._init_bucket() + try: + output_url = _bucket.sign_url('GET', + key, + lifecycle, + slash_safe=slash_safe) + _bucket.put_object_acl(key, oss2.OBJECT_ACL_PUBLIC_READ) + if set_public: + output_url = output_url.replace('%2F', '/').split('?')[0] + return output_url + except Exception as e: + retry += 1 + self.logger.info('Upload file {} error {}'.format( + target_path, e)) + return None + + def make_link(self, target_link_path, target_path) -> bool: + if not self.support_link(): + return False + link_key = osp.relpath(target_link_path, self._prefix) + target_key = osp.relpath(target_path, self._prefix) + _bucket = self._init_bucket() + try: + _bucket.put_symlink(target_key, link_key) + except Exception as e: + print(e) + return False + return True + + def make_dir(self, target_dir) -> bool: + # OSS treat file path as a key, it will create directory automatically when putting a file. + return True + + def remove(self, target_path) -> bool: + key = osp.relpath(target_path, self._prefix) + _bucket = self._init_bucket() + try: + _bucket.delete_object(key) + return True + except Exception as e: + print(e) + return False + + def get_logging_handler(self, target_logging_path): + oss_key = osp.relpath(target_logging_path, self._prefix) + _ = self._init_bucket() + return OssLoggingHandler(oss_key, self.cfg) + + def walk_dir(self, file_dir, recurse=True): + key = file_dir.replace(self._prefix, '') + if self.isdir(file_dir) and not file_dir.endswith('/'): + key += '/' + if recurse: + delimiter = '' + else: + delimiter = '/' + _bucket = self._init_bucket() + for obj in oss2.ObjectIteratorV2(_bucket, + prefix=key, + delimiter=delimiter, + max_keys=1000): + # if obj.is_prefix(): + # continue + if obj.key == key: + continue + yield osp.join(self._prefix, obj.key) + + def put_dir_from_local_dir(self, local_dir, target_dir) -> bool: + for folder, sub_folders, files in os.walk(local_dir): + for file in files: + file_abs_path = osp.join(folder, file) + file_rel_path = osp.relpath(file_abs_path, local_dir) + target_path = osp.join(target_dir, file_rel_path) + status = self.put_object_from_local_file( + file_abs_path, target_path) + if not status: + return False + return True + + def size(self, target_path) -> Optional[int]: + key = osp.relpath(target_path, self._prefix) + _bucket = self._init_bucket() + if not _bucket.object_exists(key): + self.logger.info(f"File {key} doesn't exist.") + return -1 + meta_data = _bucket.get_object_meta(key) + content_length = meta_data.content_length + return content_length + + def exists(self, target_path) -> bool: + # if is folder,try list all objects + _bucket = self._init_bucket() + target_path = target_path.replace(self._prefix, '') + try: + object_exists = _bucket.object_exists(target_path) + except Exception as e: + print(e) + return False + # if object doesnot exist, suppose it's a folder + if not object_exists: + if not target_path.endswith('/'): + target_path += '/' + ret_object_list = _bucket.list_objects(target_path, + max_keys=10).object_list + return len(ret_object_list) > 0 + return object_exists + + def isfile(self, target_path) -> bool: + if target_path.endswith('/'): + return False + target_path = osp.relpath(target_path, self._prefix) + _bucket = self._init_bucket() + try: + object_exists = _bucket.object_exists(target_path) + except Exception as e: + print(e) + return False + return object_exists + + def isdir(self, target_path) -> bool: + if not target_path.endswith('/'): + target_path += '/' + return self.exists(target_path) + + @staticmethod + def get_config_template(): + return dict_to_yaml('FILE_SYSTEMS', + __class__.__name__, + AliyunOssFs.para_dict, + set_name=True) diff --git a/scepter/modules/utils/file_clients/base_fs.py b/scepter/modules/utils/file_clients/base_fs.py new file mode 100644 index 0000000..a7a43db --- /dev/null +++ b/scepter/modules/utils/file_clients/base_fs.py @@ -0,0 +1,383 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import datetime +import os +import os.path as osp +import random +import tempfile +import warnings +from abc import ABCMeta, abstractmethod +from copy import copy +from typing import Optional + +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.directory import get_md5 +from scepter.modules.utils.file_clients.utils import remove_temp_path +from scepter.modules.utils.logger import get_logger + + +class BaseFs(object, metaclass=ABCMeta): + para_dict = { + 'TEMP_DIR': { + 'value': + None, + 'description': + 'default is None, means using system cache dir and auto remove! If you set dir, the data will' + ' be saved in this temp dir without autoremoving default.' + }, + 'AUTO_CLEAN': { + 'value': + False, + 'description': + 'when TEMP_DIR is not None, if you set AUTO_CLEAN to True, the data will be clean automatics.' + } + } + + def __init__(self, cfg, logger=None): + self._target_local_mapper = {} + self._temp_files = set() + self.cfg = cfg + self.tmp_dir = cfg.get('TEMP_DIR', None) + self.auto_clean = cfg.get('AUTO_CLEAN', False) + if self.tmp_dir is None: + self.auto_clean = True + if self.tmp_dir is not None: + if not os.path.exists(self.tmp_dir): + try: + os.makedirs(self.tmp_dir, exist_ok=True) + except Exception as e: + warnings.warn( + f'Create cache folder failed use default cache file{e}!' + .format(self.tmp_dir)) + self.tmp_dir = None + # checking that the logger exists or not + if logger is None: + self.logger = get_logger(name='File System') + else: + self.logger = logger + + # Functions without io + @abstractmethod + def get_prefix(self) -> str: + """ Get supported path prefix to determine which handler to use. + + Returns: + A prefix. + """ + pass + + @abstractmethod + def support_write(self) -> bool: + """ Return flag if this file system supports write operation. + + Returns: + Bool. + """ + pass + + @abstractmethod + def support_link(self) -> bool: + """ Return if this file system supports create a soft link. + + Returns: + Bool. + """ + pass + + def add_target_local_map(self, target_dir, local_dir): + """ Map target directory to local file system directory + + Args: + target_dir (str): Target directory. + local_dir (str): Directory in local file system. + """ + self._target_local_mapper[target_dir] = local_dir + + def map_to_local(self, target_path, etag='') -> (str, bool): + """ Map target path to local file path. (NO IO HERE). + + Args: + target_path (str): Target file path. + + Returns: + A local path and a flag indicates if the local path is a temporary file. + """ + for target_dir, local_dir in self._target_local_mapper.items(): + if target_path.startswith(target_dir): + return osp.join(local_dir, osp.relpath(target_path, + target_dir)), False + else: + return self._make_temporary_file(target_path, etag=etag), True + + def convert_to_local_path(self, target_path, etag='') -> str: + """ Deprecated. Use map_to_local() function instead. + """ + warnings.warn( + 'Function convert_to_local_path is deprecated, use map_to_local() function instead.' + ) + local_path, _ = self.map_to_local(target_path, etag=etag) + return local_path + + def basename(self, target_path) -> str: + """ Get file name from target_path + + Args: + target_path (str): Target file path. + + Returns: + A file name. + """ + return osp.basename(target_path) + + # Functions with heavy io + @abstractmethod + def get_object_to_local_file(self, + target_path, + local_path=None, + wait_finish=False) -> Optional[str]: + """ Transfer file object to local file. + If local_path is not None, + if path can be searched in local_mapper, download it as a persistent file + else, download it as a temporary file + else + download it as a persistent file + + wait_finish when multi-processing download the same data, set wait_finish as True to avoid conflict + + Args: + target_path (str): path of object in different file systems + local_path (Optional[str]): If not None, will write path to local_path. + + Returns: + Local file path of the object, none means a failure happened. + """ + pass + + # Functions with heavy io + @abstractmethod + def get_object(self, target_path): + """ Transfer file object to local file. + If local_path is not None, + if path can be searched in local_mapper, download it as a persistent file + else, download it as a temporary file + else + download it as a persistent file + + Args: + target_path (str): path of object in different file systems + local_path (Optional[str]): If not None, will write path to local_path. + + Returns: + Local file path of the object, none means a failure happened. + """ + pass + + @abstractmethod + def get_object_stream(self, target_path, start, size, delimiter=None): + """ Transfer file object to local file. + If local_path is not None, + if path can be searched in local_mapper, download it as a persistent file + else, download it as a temporary file + else + download it as a persistent file + + Args: + target_path (str): path of object in different file systems + start (int): object's start position. + size (int): object's bytes size. + delimiter (str): records's delimiter. + + Returns: + Local file path of the object, none means a failure happened. + """ + pass + + @abstractmethod + def put_object_from_local_file(self, local_path, target_path) -> bool: + """ Put local file to target file system path. + + Args: + local_path (str): local file path of the object + target_path (str): target file path of the object + + Returns: + Bool. + """ + pass + + @abstractmethod + def put_object(self, local_data, target_path) -> bool: + """ Put local file to target file system path. + + Args: + local_path (binary): local data of the object + target_path (str): target file path of the object + + Returns: + Bool. + """ + pass + + @abstractmethod + def make_link(self, target_link_path, target_path) -> bool: + """ Make soft link to target_path. + + Args: + target_link_path (str): + target_path (str) + + Returns: + Bool. + """ + pass + + @abstractmethod + def make_dir(self, target_dir) -> bool: + """ Make a directory. + If target_dir is already exists, return True. + + Args: + target_dir (str): + + Returns: + True if target_dir exists or created. + """ + pass + + @abstractmethod + def remove(self, target_path) -> bool: + """ Remove target file. + + Args: + target_path (str): + + Returns: + Bool. + """ + pass + + @abstractmethod + def get_logging_handler(self, target_logging_path): + """ Get logging handler to target logging path. + + Args: + target_logging_path: + + Returns: + A handler which has a type of subclass of logging.Handler. + """ + pass + + @abstractmethod + def walk_dir(self, file_dir, recurse=True): + pass + + @abstractmethod + def put_dir_from_local_dir(self, local_dir, target_dir) -> bool: + """ Upload all contents in local_dir to target_dir, keep the file tree. + + Args: + local_dir (str): + target_dir (str): + + Returns: + Bool. + """ + pass + + def _make_temporary_file(self, target_path, etag=''): + """ Make a temporary file for target_path, which should have the same suffix. + + Args: + target_path (str): + + Returns: + A path (str). + """ + file_name = self.basename(target_path) + _, suffix = osp.splitext(file_name) + if self.tmp_dir is None: + rand_name = '{0:%Y%m%d%H%M%S%f}'.format( + datetime.datetime.now()) + '_' + ''.join( + [str(random.randint(1, 10)) for _ in range(5)]) + # rand_name = get_md5(target_path) + if suffix: + rand_name += f'{suffix}' + tmp_file = osp.join(tempfile.gettempdir(), rand_name) + else: + cache_name = '{}{}{}'.format(etag, get_md5(target_path), suffix) + tmp_file = osp.join(self.tmp_dir, cache_name) + return tmp_file + + # Functions only for status, light io + @abstractmethod + def exists(self, target_path) -> bool: + """ Check if target_path exists. + + Args: + target_path (str): + + Returns: + Bool. + """ + pass + + @abstractmethod + def isfile(self, target_path) -> bool: + """ Check if target_path is a file. + + Args: + target_path (str): + + Returns: + Bool. + """ + pass + + @abstractmethod + def isdir(self, target_path) -> bool: + """ Check if target_path is a directory. + + Args: + target_path (str): + + Returns: + Bool. + """ + + def add_temp_file(self, tmp_file): + self._temp_files.add(tmp_file) + + def clear(self): + """Delete all temp files + """ + if self.auto_clean: + for temp_local_file in self._temp_files: + remove_temp_path(temp_local_file) + + # Functions for context + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + # if self.tmp_dir is None or self.auto_clean: + # for temp_local_file in self._temp_files: + # remove_temp_path(temp_local_file) + pass + + def __del__(self): + pass + + def copy(self): + obj = copy(self) + obj._temp_files = set( + ) # A new obj to avoid confusing in multi-thread context. + return obj + + @staticmethod + def get_config_template(): + return dict_to_yaml('FILE_SYSTEMS', + __class__.__name__, + BaseFs.para_dict, + set_name=True) diff --git a/scepter/modules/utils/file_clients/http_fs.py b/scepter/modules/utils/file_clients/http_fs.py new file mode 100644 index 0000000..3418334 --- /dev/null +++ b/scepter/modules/utils/file_clients/http_fs.py @@ -0,0 +1,140 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import os +import os.path as osp +import urllib.parse as parse +import urllib.request +from typing import Optional, Union + +from scepter.modules.utils.file_clients.base_fs import BaseFs +from scepter.modules.utils.file_clients.registry import FILE_SYSTEMS + + +@FILE_SYSTEMS.register_class() +class HttpFs(BaseFs): + para_dict = { + 'RETRY_TIMES': { + 'value': 10, + 'description': 'Retry get object times.' + } + } + para_dict.update(BaseFs.para_dict) + + def __init__(self, cfg, logger): + super(HttpFs, self).__init__(cfg, logger=logger) + retry_times = cfg.get('RETRY_TIMES', 10) + self._retry_times = retry_times + + def get_prefix(self) -> str: + return 'http' + + def support_write(self) -> bool: + return False + + def support_link(self) -> bool: + return False + + def basename(self, target_path) -> str: + url = parse.unquote(target_path) + url = url.split('?')[0] + return osp.basename(url) + + def get_object_to_local_file(self, + target_path, + local_path=None, + wait_finish=False) -> Optional[str]: + if local_path is None: + local_path, is_tmp = self.map_to_local(target_path) + else: + is_tmp = False + + os.makedirs(osp.dirname(local_path), exist_ok=True) + + retry = 0 + while retry < self._retry_times: + try: + target_url = urllib.parse.quote(target_path, + safe=":/?#[]@!$&'()*+,;=%") + urllib.request.urlretrieve(target_url, local_path) + if osp.exists(local_path): + break + except Exception: + retry += 1 + + if retry >= self._retry_times: + return None + + if is_tmp: + self.add_temp_file(local_path) + return local_path + + def get_object(self, target_path): + try: + local_data = open(self.get_object_to_local_file(target_path), + 'rb').read() + except Exception as e: + self.logger.error(f'Read {target_path} error {e}') + local_data = None + return local_data + + def put_object(self, local_data, target_path): + raise NotImplementedError + + def put_object_from_local_file(self, local_path, target_path) -> bool: + raise NotImplementedError + + def make_link(self, target_link_path, target_path) -> bool: + raise NotImplementedError + + def make_dir(self, target_dir) -> bool: + raise NotImplementedError + + def remove(self, target_path) -> bool: + raise NotImplementedError + + def get_logging_handler(self, target_logging_path): + raise NotImplementedError + + def walk_dir(self, file_dir, recurse=True): + raise NotImplementedError + + def put_dir_from_local_dir(self, local_dir, target_dir) -> bool: + raise NotImplementedError + + def size(self, target_path) -> Optional[int]: + raise NotImplementedError + + def get_object_chunk_list(self, + target_path, + chunk_num=1, + delimiter=None) -> Optional[list]: + raise NotImplementedError + + def get_object_stream( + self, + target_path, + start, + size=10000, + delimiter=None) -> (Union[bytes, str, None], Optional[int]): + raise NotImplementedError + + def get_url(self, target_path, lifecycle=3600 * 100): + return target_path + + def exists(self, target_path) -> bool: + req = urllib.request.Request(target_path) + req.get_method = lambda: 'HEAD' + + try: + urllib.request.urlopen(req) + return True + except Exception: + return False + + def isfile(self, target_path) -> bool: + # Well for a http url, it should only be a file. + return True + + def isdir(self, target_path) -> bool: + return False diff --git a/scepter/modules/utils/file_clients/local_fs.py b/scepter/modules/utils/file_clients/local_fs.py new file mode 100644 index 0000000..b568503 --- /dev/null +++ b/scepter/modules/utils/file_clients/local_fs.py @@ -0,0 +1,337 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import logging +import os +import os.path as osp +import shutil +from typing import Optional, Union + +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.file_clients.base_fs import BaseFs +from scepter.modules.utils.file_clients.registry import FILE_SYSTEMS + + +def is_start_of_line(f, position, delimiter='\n'): + if position == 0: + return True + # Check whether the previous character is EOL + f.seek(position - 1) + return f.read(1) == delimiter + + +def get_next_line_position(f, position): + # Read the current line till the end + f.seek(position) + f.readline() + # Return a position after reading the line + return f.tell() + + +@FILE_SYSTEMS.register_class() +class LocalFs(BaseFs): + def __init__(self, cfg, logger=None): + super(LocalFs, self).__init__(cfg, logger=logger) + self._fs_prefix = os.path.abspath(os.curdir) + + def get_prefix(self) -> str: + return self._fs_prefix + + def reconstruct_path(self, target_path) -> str: + if target_path.startswith(self.get_prefix()): + return target_path + if target_path.startswith('./') or target_path.startswith('../'): + return os.path.join(self.get_prefix(), + target_path).replace('/./', + '/').replace('/../', '/') + if target_path.startswith('/'): + return target_path + if target_path.startswith('file://'): + return os.path.join(self.get_prefix(), + target_path[len('file://'):]) + return os.path.join(self.get_prefix(), target_path) + + def support_write(self) -> bool: + return True + + def support_link(self) -> bool: + return True + + def map_to_local(self, target_path) -> (str, bool): + target_path = self.reconstruct_path(target_path) + return target_path, False + + def get_object_to_local_file(self, + target_path, + local_path=None, + wait_finish=False) -> Optional[str]: + target_path = self.reconstruct_path(target_path) + if local_path is not None: + local_path = self.reconstruct_path(local_path) + local_path = osp.abspath(local_path) + if local_path != target_path: + # copy target_path to local_path + os.makedirs(osp.dirname(local_path), exist_ok=True) + try: + shutil.copy(target_path, local_path) + except Exception as e: + self.logger.info(f'Copy file failed {e}') + return None + + return local_path + return target_path + + def get_dir_to_local_dir(self, + target_path, + local_path=None, + wait_finish=False, + timeout=3600, + worker_id=0) -> Optional[str]: + if not self.isdir(target_path): + self.logger.info( + f"{target_path} is not directory or doesn't exist.") + if not target_path.endswith('/'): + target_path += '/' + if local_path is None: + local_path, is_tmp = self.map_to_local(target_path) + else: + is_tmp = False + local_path = local_path.replace('/./', '/') + os.makedirs(local_path, exist_ok=True) + generator = self.walk_dir(target_path) + for file_name in generator: + if file_name == target_path or file_name == target_path + '/': + continue + local_file_name = os.path.join( + local_path, + file_name.split(target_path)[-1]).replace('/./', '/') + if not self.isdir(file_name): + self.get_object_to_local_file(file_name, + local_file_name, + wait_finish=wait_finish) + else: + self.get_dir_to_local_dir(file_name, + local_file_name, + wait_finish=wait_finish) + if is_tmp: + self.add_temp_file(local_path) + return local_path + + def get_object(self, target_path) -> Optional[bytes]: + target_path = self.reconstruct_path(target_path) + try: + local_data = open(target_path, 'rb').read() + except Exception as e: + self.logger.error(f'Read {target_path} error {e}') + local_data = None + return local_data + + def get_object_stream( + self, + target_path, + start, + size=10000, + delimiter=None) -> (Union[bytes, str, None], Optional[int]): + target_path = self.reconstruct_path(target_path) + if not osp.exists(target_path): + self.logger.error(f'Read {target_path} error: file not exist.') + return None, None + + file_size = os.path.getsize(target_path) + start, end = start, min(file_size, start + size) + if start >= end - 1: + return None, None + + with open(target_path, 'rb') as f: + if delimiter is None or end == file_size: + f.seek(start) + local_data = f.read(end - start) + return local_data, end + f.seek(start) + local_data = f.read(end - start) + try: + total_len = len(local_data) + offset = 0 + try: + sp_data = local_data.split(bytes(delimiter, 'utf-8')) + # if failed, suppose the bytes is splited error. + except Exception as e: + self.logger.info( + f'Return data split error,please check your delimiter {e}' + ) + return None, end + if not len(sp_data[-1]) == len(local_data): + cur_offset = len(sp_data[-1]) + offset += cur_offset + local_data = local_data[:total_len - cur_offset] + local_data = local_data[len(sp_data[0]):] + end = end - offset + except Exception as e: + self.logger.info(f'Local data decode error {e}') + return local_data, end + + def get_object_chunk_list(self, + target_path, + chunk_num=-1, + chunk_size=-1, + delimiter=None) -> Optional[list]: + target_path = self.reconstruct_path(target_path) + if not osp.exists(target_path): + self.logger.error(f'Read {target_path} error: file not exist.') + return None + file_size = os.path.getsize(target_path) + if chunk_size < 0 and chunk_num < 0: + self.logger.error( + 'Suppose chunk size > 0 or chunk num > 0, instead of both < 0') + return None + if chunk_size < 0 and chunk_num > 0: + chunk_size = file_size // chunk_num + 1 + chunk_st_et = [] + # Don't care of the lines info + if delimiter is None: + chunk_start = 0 + # Iterate over all chunks and construct arguments for `process_chunk` + while chunk_start < file_size: + chunk_end = min(file_size, chunk_start + chunk_size) + chunk_st_et.append([chunk_start, chunk_end - chunk_start]) + chunk_start = chunk_end + else: + with open(target_path, 'rb') as f: + chunk_start = 0 + offset = 0 + # Iterate over all chunks and construct arguments for `process_chunk` + while chunk_start < file_size: + chunk_end = min(file_size, + chunk_start + chunk_size + offset) + quota_st = max(0, chunk_end - 20000) + quota_st = max(chunk_start, quota_st) + f.seek(quota_st) + local_data = f.read(chunk_end - quota_st) + offset = 0 + if not chunk_end == file_size: + try: + try: + sp_data = local_data.split( + bytes(delimiter, 'utf-8')) + # if failed, suppose the bytes is splited error. + except Exception as e: + self.logger.info( + f'Return data split error,please check your delimiter {e}' + ) + return None + if not len(sp_data[-1]) == len(local_data): + cur_offset = len(sp_data[-1]) + offset += cur_offset + chunk_end = chunk_end - offset + except Exception as e: + self.logger.info( + f'Local data decode error {e}, check your data is supported by str.decode().' + ) + chunk_end = chunk_end - offset + chunk_st_et.append([chunk_start, chunk_end - chunk_start]) + chunk_start = chunk_end + return chunk_st_et + + def put_object(self, local_data, target_path) -> bool: + target_path = self.reconstruct_path(target_path) + with open(target_path, 'w') as f: + f.write(local_data) + return True + + def walk_dir(self, file_dir, recurse=True): + for root, dirs, files in os.walk(file_dir, topdown=True): + sub_files = files + dirs + for name in sub_files: + yield os.path.join(root, name) + + def put_object_from_local_file(self, local_path, target_path) -> bool: + target_path = self.reconstruct_path(target_path) + local_path = self.reconstruct_path(local_path) + if local_path != target_path: + try: + shutil.copy(local_path, target_path) + except Exception: + return False + return True + + def get_url(self, target_path, set_public=False, lifecycle=3600 * 100): + return target_path + + def make_dir(self, target_dir) -> bool: + target_dir = self.reconstruct_path(target_dir) + if osp.exists(target_dir): + if osp.isfile(target_dir): + self.logger.error(f'{target_dir} already exists as a file!') + return False + return True + try: + os.makedirs(target_dir) + except Exception as e: + self.logger.error(e) + return False + return True + + def make_link(self, target_link_path, target_path) -> bool: + target_path = self.reconstruct_path(target_path) + target_link_path = self.reconstruct_path(target_link_path) + try: + if osp.lexists(target_link_path): + os.remove(target_link_path) + os.symlink(target_path, target_link_path) + return True + except Exception: + return False + + def remove(self, target_path) -> bool: + target_path = self.reconstruct_path(target_path) + if osp.exists(target_path): + try: + os.remove(target_path) + except Exception: + return False + return True + + def get_logging_handler(self, target_logging_path): + return logging.FileHandler(target_logging_path) + + def put_dir_from_local_dir(self, local_dir, target_dir) -> bool: + local_dir = self.reconstruct_path(local_dir) + target_dir = self.reconstruct_path(target_dir) + if local_dir == target_dir: + return True + # cp -f local_dir/* target_dir/* + if not osp.exists(target_dir): + status = os.system(f'mkdir -p {target_dir}') + if status != 0: + return False + try: + shutil.copytree(local_dir, target_dir, symlinks=True) + except Exception: + return False + return True + + def size(self, target_path) -> Optional[int]: + target_path = self.reconstruct_path(target_path) + if not osp.exists(target_path): + self.logger.info(f"File {target_path} doesn't exist.") + return -1 + file_size = os.path.getsize(target_path) + return file_size + + def exists(self, target_path) -> bool: + target_path = self.reconstruct_path(target_path) + return osp.exists(target_path) + + def isfile(self, target_path) -> bool: + target_path = self.reconstruct_path(target_path) + return osp.isfile(target_path) + + def isdir(self, target_path) -> bool: + target_path = self.reconstruct_path(target_path) + return osp.isdir(target_path) + + @staticmethod + def get_config_template(): + return dict_to_yaml('FILE_SYSTEMS', + __class__.__name__, {}, + set_name=True) diff --git a/scepter/modules/utils/file_clients/modelscope_fs.py b/scepter/modules/utils/file_clients/modelscope_fs.py new file mode 100644 index 0000000..7d20b98 --- /dev/null +++ b/scepter/modules/utils/file_clients/modelscope_fs.py @@ -0,0 +1,202 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import os +import os.path as osp +import urllib.parse as parse +import urllib.request +from typing import Optional, Union + +from scepter.modules.utils.file_clients.base_fs import BaseFs +from scepter.modules.utils.file_clients.registry import FILE_SYSTEMS + + +@FILE_SYSTEMS.register_class() +class ModelscopeFs(BaseFs): + para_dict = { + 'RETRY_TIMES': { + 'value': 10, + 'description': 'Retry get object times.' + } + } + para_dict.update(BaseFs.para_dict) + + def __init__(self, cfg, logger): + super(ModelscopeFs, self).__init__(cfg, logger=logger) + retry_times = cfg.get('RETRY_TIMES', 10) + self._retry_times = retry_times + + def get_prefix(self) -> str: + return 'ms://' + + def support_write(self) -> bool: + return False + + def support_link(self) -> bool: + return False + + def basename(self, target_path) -> str: + url = parse.unquote(target_path) + url = url.split('?')[0] + return osp.basename(url) + + def get_object_to_local_file(self, + target_path, + local_path=None, + wait_finish=False) -> Optional[str]: + from modelscope.hub.file_download import model_file_download + + key = osp.relpath(target_path, self.get_prefix()) + key, file_path = key.split('@', 1) + + if ':' in key: + key, revision = key.split(':', 1) + else: + revision = None + + if local_path is None: + local_path, is_tmp = self.map_to_local(key) + else: + is_tmp = False + + if revision is not None: + local_path = local_path + '_' + str(revision) + + retry = 0 + while retry < self._retry_times: + try: + local_path = model_file_download(model_id=key, + revision=revision, + file_path=file_path, + cache_dir=local_path) + if osp.exists(local_path): + break + except Exception: + retry += 1 + + if retry >= self._retry_times: + return None + + if is_tmp: + self.add_temp_file(local_path) + return local_path + + def get_dir_to_local_dir(self, + target_path, + local_path=None, + wait_finish=False, + timeout=3600, + worker_id=-1) -> Optional[str]: + from modelscope.hub.snapshot_download import snapshot_download + assert target_path.startswith(self.get_prefix()) + + key = osp.relpath(target_path, self.get_prefix()) + if '@' not in key: + key, ret_folder = key.split('@', 1)[0], '' + else: + at_level_folder = key.split('@') + if len(at_level_folder) > 2: + raise f'Target path should include only one @, but you give {len(at_level_folder)} @.' + key, ret_folder = at_level_folder + + if ':' in key: + key, revision = key.split(':', 1) + else: + revision = None + + if local_path is None: + local_path, is_tmp = self.map_to_local(key) + else: + is_tmp = False + + if revision is not None: + local_path = local_path + '_' + str(revision) + + retry = 0 + while retry < self._retry_times: + try: + local_path = snapshot_download(key, + revision=revision, + cache_dir=local_path) + if osp.exists(local_path): + break + except Exception: + retry += 1 + + if retry >= self._retry_times: + return None + + if is_tmp: + self.add_temp_file(local_path) + if not ret_folder == '': + local_path = os.path.join(local_path, ret_folder) + return local_path + + def get_object(self, target_path): + try: + local_data = open(self.get_object_to_local_file(target_path), + 'rb').read() + except Exception as e: + self.logger.error(f'Read {target_path} error {e}') + local_data = None + return local_data + + def put_object(self, local_data, target_path): + raise NotImplementedError + + def put_object_from_local_file(self, local_path, target_path) -> bool: + raise NotImplementedError + + def make_link(self, target_link_path, target_path) -> bool: + raise NotImplementedError + + def make_dir(self, target_dir) -> bool: + raise NotImplementedError + + def remove(self, target_path) -> bool: + raise NotImplementedError + + def get_logging_handler(self, target_logging_path): + raise NotImplementedError + + def walk_dir(self, file_dir, recurse=True): + raise NotImplementedError + + def put_dir_from_local_dir(self, local_dir, target_dir) -> bool: + raise NotImplementedError + + def size(self, target_path) -> Optional[int]: + raise NotImplementedError + + def get_object_chunk_list(self, + target_path, + chunk_num=1, + delimiter=None) -> Optional[list]: + raise NotImplementedError + + def get_object_stream( + self, + target_path, + start, + size=10000, + delimiter=None) -> (Union[bytes, str, None], Optional[int]): + raise NotImplementedError + + def get_url(self, target_path, lifecycle=3600 * 100): + return target_path + + def exists(self, target_path) -> bool: + req = urllib.request.Request(target_path) + req.get_method = lambda: 'HEAD' + + try: + urllib.request.urlopen(req) + return True + except Exception: + return False + + def isfile(self, target_path) -> bool: + # Well for a http url, it should only be a file. + return True + + def isdir(self, target_path) -> bool: + return False diff --git a/scepter/modules/utils/file_clients/registry.py b/scepter/modules/utils/file_clients/registry.py new file mode 100644 index 0000000..68bb8f9 --- /dev/null +++ b/scepter/modules/utils/file_clients/registry.py @@ -0,0 +1,6 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +from scepter.modules.utils.registry import Registry + +FILE_SYSTEMS = Registry('FILE_SYSTEMS') diff --git a/scepter/modules/utils/file_clients/utils.py b/scepter/modules/utils/file_clients/utils.py new file mode 100644 index 0000000..3250c15 --- /dev/null +++ b/scepter/modules/utils/file_clients/utils.py @@ -0,0 +1,52 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import os +import os.path as osp + + +def check_if_local_path(path): + """ + Check if path is a local path, no matter file or directory. + + True: + file:///home/admin/a.txt (standard path) + /home/admin/a.txt (standard unix path) + C:\\Users\\a.txt (standard windows path) + C:/Users/a.txt (works as well) + ./a.txt (relative path) + a.txt (relative path) + False: + http://www.aliyun.com/a.txt + http://www.aliyun.com/a.txt + oss://aliyun/a.txt + + Args: + path (str): + + Returns: + True if path is a local path. + """ + if path.startswith('file://'): + return True + return '://' not in path + + +def remove_temp_path(path): + """ + Delete local temp path. + + Args: + path (str): + + Returns: + """ + if not osp.exists(path): + return + if not osp.isfile(path): + return + try: + os.remove(path) + except Exception: + pass + # warnings.warn(f"remove {path}") diff --git a/scepter/modules/utils/file_system.py b/scepter/modules/utils/file_system.py new file mode 100644 index 0000000..b4b3c10 --- /dev/null +++ b/scepter/modules/utils/file_system.py @@ -0,0 +1,445 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import os +import threading +import time +import warnings +from contextlib import contextmanager +from queue import Queue + +from scepter.modules.utils.config import Config +from scepter.modules.utils.file_clients.base_fs import BaseFs +from scepter.modules.utils.file_clients.local_fs import LocalFs +from scepter.modules.utils.file_clients.registry import FILE_SYSTEMS +from scepter.modules.utils.file_clients.utils import check_if_local_path + + +class IoString(str): + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + pass + + +class IoBytes(bytes): + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + pass + + +class ReadException(Exception): + pass + + +class WriteException(Exception): + pass + + +class FileSystem(object): + def __init__(self): + self._prefix_to_clients = {} + self._default_client = None + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + pass + + def __del__(self): + for k, client in self._prefix_to_clients.items(): + client.clear() + + @property + def support_prefix(self): + return self._prefix_to_clients + + def init_fs_client(self, cfg=None, logger=None, overwrite=True): + """ Initialize file system backend + Supported backend: + 1. Local file system, e.g. /home/admin/work_dir, work_dir_bk/imagenet_pretrain + 2. Aliyun Oss, e.g. oss://bucket_name/work_dir + 3. Http, only support to read content, e.g. + https://www.google.com.hk/images/branding/googlelogo/2x/googlelogo_color_272x92dp.png + 4. other fs backend... + + Args: + cfg (list, dict, optional): + list: list of file system configs to be initialized + dict: a dict contains file system configs as values or a file system config dict + optional: Will only use default LocalFs + """ + + fs_cfg = cfg or Config(load=False) + if not isinstance(fs_cfg, Config): + raise '{} is not a Config Instance!'.format(fs_cfg) + + if not fs_cfg.have('NAME'): + raise KeyError(f'{fs_cfg} does not contain key NAME!') + + fs_client = FILE_SYSTEMS.build(fs_cfg, logger=logger) + _prefix = fs_client.get_prefix() + if _prefix in self._prefix_to_clients and not overwrite: + return _prefix + if _prefix in self._prefix_to_clients: + warnings.warn( + 'File client {} has already been set, will be replaced by newer config.' + .format(_prefix)) + self._prefix_to_clients[_prefix] = fs_client + return _prefix + + def get_fs_client(self, target_path, safe=False) -> BaseFs: + """ Get the client by input path. + Every file system has its own identifier, default will use local file system to have a try. + If copy needed, only do shallow copy. + + Args: + target_path (str): + safe (bool): In safe mode, get the copy of the client. + """ + obj = None + + for prefix in sorted(list(self._prefix_to_clients.keys()), + key=lambda a: -len(a)): + if target_path.startswith(prefix): + obj = self._prefix_to_clients[prefix] + break + if obj is not None: + if safe: + return obj.copy() + else: + return obj + + if not check_if_local_path(target_path): + warnings.warn( + f'{target_path} is not a local path, use LocalFs may cause an error.' + ) + if self._default_client is None: + self._default_client = LocalFs(Config(load=False)) + if safe: + return self._default_client.copy() + else: + return self._default_client + + def get_dir_to_local_dir(self, + target_path, + local_path=None, + wait_finish=False, + timeout=3600, + worker_id=0): + with self.get_fs_client(target_path) as client: + local_path = client.get_dir_to_local_dir(target_path, + local_path=local_path, + wait_finish=wait_finish, + timeout=timeout, + worker_id=worker_id) + if local_path is None: + raise ReadException( + f'Failed to fetch {target_path} to {local_path}') + return IoString(local_path) + + def add_target_local_map(self, target_dir, local_dir): + """ Map target directory to local file system directory + + Args: + target_dir (str): Target directory. + local_dir (str): Directory in local file system. + """ + with self.get_fs_client(target_dir, safe=False) as client: + client.add_target_local_map(target_dir, local_dir) + + def make_dir(self, target_dir): + """ Make a directory. + If target_dir is already exists, return True. + + Args: + target_dir (str): + + Returns: + True if target_dir exists or created. + """ + with self.get_fs_client(target_dir) as client: + return client.make_dir(target_dir) + + def exists(self, target_path): + """ Check if target_path exists. + + Args: + target_path (str): + + Returns: + Bool. + """ + with self.get_fs_client(target_path) as client: + return client.exists(target_path) + + def map_to_local(self, target_path): + """ Map target path to local file path. (NO IO HERE). + + Args: + target_path (str): Target file path. + + Returns: + A local path and a flag indicates if the local path is a temporary file. + """ + with self.get_fs_client(target_path) as client: + local_path, is_tmp = client.map_to_local(target_path) + return local_path, is_tmp + + def put_dir_from_local_dir(self, local_dir, target_dir): + """ Upload all contents in local_dir to target_dir, keep the file tree. + + Args: + local_dir (str): + target_dir (str): + + Returns: + Bool. + """ + with self.get_fs_client(target_dir) as client: + return client.put_dir_from_local_dir(local_dir, target_dir) + + def walk_dir(self, target_dir, recurse=True): + """ Iterator to access the files of target dir. + Args: + target_dir (str): + + Returns: + Generator. + """ + with self.get_fs_client(target_dir) as client: + return client.walk_dir(target_dir, recurse=recurse) + + def is_local_client(self, target_path) -> bool: + """ Check if the client support read or write to target_path is a LocalFs. + + Args: + target_path (str): + + Returns: + Bool. + """ + with self.get_fs_client(target_path) as client: + return type(client) is LocalFs + + def put_object_from_local_file(self, local_path, target_path) -> bool: + with self.get_fs_client(target_path) as client: + flag = client.put_object_from_local_file(local_path, target_path) + return flag + + def get_from(self, target_path, local_path=None, wait_finish=False): + with self.get_fs_client(target_path) as client: + local_path = client.get_object_to_local_file( + target_path, local_path=local_path, wait_finish=wait_finish) + if local_path is None: + raise ReadException( + f'Failed to fetch {target_path} to {local_path}') + return IoString(local_path) + + def get_url(self, target_path, set_public=False, lifecycle=3600 * 100): + with self.get_fs_client(target_path) as client: + output_url = client.get_url(target_path, + set_public=set_public, + lifecycle=lifecycle) + return output_url + + def get_object(self, target_path): + with self.get_fs_client(target_path) as client: + local_data = client.get_object(target_path) + if local_data is None: + return IoBytes(None) + return IoBytes(local_data) + + def put_object(self, local_data, target_path): + with self.get_fs_client(target_path) as client: + flg = client.put_object(local_data, target_path) + return flg + + def delete_object(self, target_path): + with self.get_fs_client(target_path) as client: + if self.isfile(target_path): + flg = client.remove(target_path) + return flg + else: + return False + + def get_batch_objects_from(self, target_path_list, wait_finish=False): + data_quene = Queue() + batch_size = 20 + R = threading.Lock() + + def get_one_object(target_path_list): + for target_path in target_path_list: + if self.exists(target_path): + local_path = self.get_from(target_path, + wait_finish=wait_finish) + else: + local_path = None + R.acquire() + try: + data_quene.put_nowait([target_path, local_path]) + except Exception: + R.release() + R.release() + + while True: + batch_list = target_path_list[:4 * batch_size] + if len(batch_list) < 1: + break + target_path_list = target_path_list[4 * batch_size:] + threading_list = [] + for i in range(batch_size): + t = threading.Thread(target=get_one_object, + args=(batch_list[i::batch_size], )) + t.daemon = True + t.start() + threading_list.append(t) + [threading_t.join() for threading_t in threading_list] + file_dict = {} + while not data_quene.empty(): + target_path, local_path = data_quene.get_nowait() + file_dict[target_path] = local_path + + for target_path in batch_list: + local_path = file_dict.get(target_path, None) + yield local_path + + def put_batch_objects_to(self, + local_path_list, + target_path_list, + batch_size=20, + wait_finish=False): + data_quene = Queue() + R = threading.Lock() + + def put_one_object(local_path_list, target_path_list): + for local_path, target_path in zip(local_path_list, + target_path_list): + if local_path is None or target_path is None: + flg = False + elif self.exists(local_path): + local_cache = self.get_from(local_path, + local_path + f'{time.time()}', + wait_finish=wait_finish) + flg = self.put_object_from_local_file( + local_cache, target_path) + try: + if os.path.exists(local_cache): + os.remove(local_cache) + except Exception: + pass + else: + flg = False + R.acquire() + try: + data_quene.put_nowait([local_path, target_path, flg]) + except Exception: + R.release() + R.release() + + while True: + batch_local_list = local_path_list[:4 * batch_size] + batch_target_list = target_path_list[:4 * batch_size] + if len(batch_local_list) < 1: + break + local_path_list = local_path_list[4 * batch_size:] + target_path_list = target_path_list[4 * batch_size:] + threading_list = [] + for i in range(batch_size): + t = threading.Thread(target=put_one_object, + args=( + batch_local_list[i::batch_size], + batch_target_list[i::batch_size], + )) + t.daemon = True + t.start() + threading_list.append(t) + [threading_t.join() for threading_t in threading_list] + file_dict = {} + while not data_quene.empty(): + local_path, target_path, flg = data_quene.get_nowait() + file_dict[local_path] = [local_path, target_path, flg] + + for idx, local_path in enumerate(batch_local_list): + local_path, target_path, flg = file_dict.get( + local_path, + [batch_local_list[idx], batch_target_list[idx], False]) + yield local_path, target_path, flg + + def get_object_stream(self, + target_path, + start, + size=10000, + delimiter=None): + with self.get_fs_client(target_path) as client: + local_data, end = client.get_object_stream(target_path, + start, + size=size, + delimiter=delimiter) + return local_data, end + + def get_object_chunk_list(self, + target_path, + chunk_num=1, + chunk_size=-1, + delimiter=None): + with self.get_fs_client(target_path) as client: + chunk_list = client.get_object_chunk_list(target_path, + chunk_num=chunk_num, + chunk_size=chunk_size, + delimiter=delimiter) + if chunk_list is None: + raise ReadException(f'Failed to fetch {target_path}') + return chunk_list + + def size(self, target_path): + with self.get_fs_client(target_path) as client: + size = client.size(target_path) + return size + + def isfile(self, target_path): + with self.get_fs_client(target_path) as client: + is_file = client.isfile(target_path) + return is_file + + def isdir(self, target_path): + with self.get_fs_client(target_path) as client: + is_dir = client.isdir(target_path) + return is_dir + + @contextmanager + def put_to(self, target_path): + with self.get_fs_client(target_path) as client: + local_path, is_tmp = client.map_to_local(target_path) + if is_tmp: + client.add_temp_file(local_path) + if not os.path.exists(os.path.dirname(local_path)): + os.makedirs(os.path.dirname(local_path)) + yield local_path + status = client.put_object_from_local_file(local_path, target_path) + if not status: + raise WriteException( + f'Failed to upload from {local_path} to {target_path}') + if not isinstance(client, LocalFs): + try: + if os.path.exists(local_path): + os.remove(local_path) + except Exception: + pass + + def __repr__(self) -> str: + s = 'Support prefix list:\n' + for prefix in sorted(list(self._prefix_to_clients.keys()), + key=lambda a: -len(a)): + s += f'\t{prefix} -> {self._prefix_to_clients[prefix]}\n' + return s + + +global FS, DATA_FS, MODEL_FS +# global instance, easy to use +FS = FileSystem() +DATA_FS = FS +MODEL_FS = FS diff --git a/scepter/modules/utils/logger.py b/scepter/modules/utils/logger.py new file mode 100644 index 0000000..7a5648d --- /dev/null +++ b/scepter/modules/utils/logger.py @@ -0,0 +1,177 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import logging +import numbers +import sys +import time +from collections import OrderedDict + +import numpy as np +import torch + +from scepter.modules.utils.distribute import get_dist_info + + +def as_time(s): + s = int(s) + one_day, one_hour, one_min = 3600 * 24, 3600, 60 + day, hour, min = 0, 0, 0 + # compute day 3600 * 24 + if s >= one_day: + day = int(s // one_day) + s = s % one_day + # compute hour 3600 + if s >= one_hour: + hour = int(s // one_hour) + s = s % one_hour + # compute min 60 + if s >= one_min: + min = int(s // one_min) + s = s % one_min + output_str = [] + if day > 0: + output_str.append('{}days'.format(day)) + if hour > 0: + output_str.append('{}hours'.format(hour)) + if min > 0: + output_str.append('{}mins'.format(min)) + output_str.append('{}secs'.format(int(s))) + return ' '.join(output_str) + + +def time_since(since, percent): + now = time.time() + s = now - since + es = s / (percent) + rs = es - s + return '{} {:.2f}%({})'.format(as_time(s), 100 * percent, as_time(rs)) + + +def get_logger(name='torch dist'): + logger = logging.getLogger(name) + logger.propagate = False + if len(logger.handlers) == 0: + std_handler = logging.StreamHandler(sys.stdout) + formatter = logging.Formatter( + '%(name)s [%(levelname)s] %(asctime)s ' + '[File: %(filename)s Function: %(funcName)s at line %(lineno)d] %(message)s' + ) + std_handler.setFormatter(formatter) + std_handler.setLevel(logging.INFO) + logger.setLevel(logging.INFO) + logger.addHandler(std_handler) + return logger + + +def init_logger(in_logger, log_file=None, dist_launcher='pytorch'): + """ Add file handler to logger on rank 0 and set log level by dist_launcher + + Args: + in_logger (logging.Logger): + log_file (str, None): if not None, a file handler will be add to in_logger + dist_launcher (str, None): + """ + rank, _ = get_dist_info() + if rank == 0: + if log_file is not None: + from scepter.modules.utils.file_system import FS + file_handler = FS.get_fs_client(log_file).get_logging_handler( + log_file) + formatter = logging.Formatter( + '%(name)s [%(levelname)s] %(asctime)s [File: %(filename)s ' + 'Function: %(funcName)s at line %(lineno)d] %(message)s') + file_handler.setFormatter(formatter) + file_handler.setLevel(logging.INFO) + in_logger.addHandler(file_handler) + in_logger.info(f'Running task with log file: {log_file}') + in_logger.setLevel(logging.INFO) + else: + if dist_launcher == 'pytorch': + in_logger.setLevel(logging.ERROR) + else: + # Distribute Training with more than one machine, we'd like to show logs on every machine. + in_logger.setLevel(logging.INFO) + + +class LogAgg(object): + """ Log variable aggregate tool. Recommend to invoke clear() function after one epoch. + In distributed training environment, tensor variable will be all reduced to get an average. + + Example: + >>> agg = LogAgg() + >>> agg.update(dict(loss=0.1, accuracy=0.5)) + >>> agg.update(dict(loss=0.2, accuracy=0.6)) + >>> agg.update(dict(loss=0.3, accuracy=0.7)) + >>> agg.aggregate() + OrderedDict([('loss', (0.3, 0.20000000000000004)), ('accuracy', (0.7, 0.6))]) + """ + def __init__(self): + self.buffer = OrderedDict() + self.counter = [] + + def update(self, kv: dict, count=1): + """ Update variables + + Args: + kv (dict): a dict with value type in (torch.Tensor, numbers) + count (int): divider, default is 1 + """ + for k, v in kv.items(): + if isinstance(v, torch.Tensor): + # Must be scalar + if not v.ndim == 0: + continue + v = v.item() + elif isinstance(v, np.ndarray): + # Must be scalar + if not v.ndim == 0: + continue + elif isinstance(v, numbers.Number): + # Must be number + pass + else: + continue + + if k not in self.buffer: + self.buffer[k] = [] + self.buffer[k].append(v) + self.counter.append(count) + + def _aggregate(self, n=0): + """ Do aggregation. + + Args: + n (int): recent n numbers, if 0, start from 0 + + Returns: + A dict contains aggregate values. + """ + + ret = OrderedDict() + for key in self.buffer: + values = np.array(self.buffer[key][-n:]) + nums = np.array(self.counter[-n:]) + avg = np.sum(values * nums) / np.sum(nums) + ret[key] = avg + return ret + + def aggregate(self, log_interval=1): + """ Do aggregation with current step values and all mean values. + + Args: + log_interval (int): Steps to aggregate current state, default is 1. + + Returns: + A dict contains current step and all step mean values. + """ + cur = self._aggregate(log_interval) + all_mean = self._aggregate(0) + ret = OrderedDict() + for key in cur: + ret[key] = (cur[key], all_mean[key]) + return ret + + def reset(self): + self.buffer.clear() + self.counter.clear() diff --git a/scepter/modules/utils/math_plot.py b/scepter/modules/utils/math_plot.py new file mode 100644 index 0000000..0c7da1b --- /dev/null +++ b/scepter/modules/utils/math_plot.py @@ -0,0 +1,103 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import warnings + +import numpy as np + +try: + import matplotlib.pyplot as plt +except Exception as e: + warnings.warn(f'Runing without matplotlib {e}') + +color_list = ['b', 'g', 'r', 'c', 'm', 'y', 'k'] +line_list = ['-', '--', '-.', ':'] + + +def plot_multi_curves(x, + y, + show=False, + title=None, + save_path=None, + x_label=None, + y_label=None): + ''' + Args: + x: the x-axis data + y: the y-axis data dict + like: [{"data": np.ndarrays, "label": ""}] + title: None + show: False + save_path: None + x_label: None + y_label: None + Returns: + ''' + if save_path is not None: + plt.figure() + + x_max, x_min = np.max(x), np.min(x) + + max_num, min_num = 0, 0 + for y_id, data in enumerate(y): + max_n = np.max(data['data']) + min_n = np.min(data['data']) + max_num = max_n if max_n > max_num else max_num + min_num = min_n if min_n < min_num else min_num + plt.plot(x, + data['data'], + linestyle=line_list[y_id % len(line_list)], + linewidth=2, + color=color_list[y_id % len(color_list)], + label=data['label'], + alpha=1.00) + plt.title(title, loc='center') + + plt.legend(loc='upper right') + if x_label is not None: + plt.xlabel(x_label) + if y_label is not None: + plt.ylabel(y_label) + + x_step = (x_max - x_min) / 5 + y_step = (max_num - min_num) / 5 + plt.xticks(np.arange(x_min - x_step / 2, x_max + x_step / 2, x_step)) + plt.yticks(np.arange(min_num - y_step / 2, max_num + y_step / 2, y_step)) + plt.grid() + if save_path is not None: + plt.savefig(save_path) + if show: + plt.show() + plt.clf() + plt.cla() + plt.close() + return True + + +def plt_curve(x, + y, + show=False, + title=None, + save_path=None, + x_label=None, + y_label=None): + ''' + Args: + x: the x-axis data + y: the y-axis data dict + like: [{"data": np.ndarrays, "label": ""}] + title: None + show: False + save_path: None + x_label: None + y_label: None + Returns: + ''' + return plot_multi_curves(x, [{ + 'data': y, + 'label': 'y' + }], + show=show, + title=title, + save_path=save_path, + x_label=x_label, + y_label=y_label) diff --git a/scepter/modules/utils/model.py b/scepter/modules/utils/model.py new file mode 100644 index 0000000..053955d --- /dev/null +++ b/scepter/modules/utils/model.py @@ -0,0 +1,164 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import os +import re +import sys +from collections import OrderedDict + +import torch +import torch.nn as nn +from torch.utils.model_zoo import load_url as load_state_dict_from_url + + +class StdMsg(): + def __init__(self, name='msg'): + self.name = name + + def info(self, msg): + sys.stdout.write('[Info]: ' + msg + '\n') + + def error(self, msg): + sys.stdout.write('[Error]: ' + msg + '\n') + + def warning(self, msg): + sys.stdout.write('[Warning]: ' + msg + '\n') + + +def move_model_to_cpu(params): + cpu_params = OrderedDict() + for key, val in params.items(): + cpu_params[key] = val.cpu() + return cpu_params + + +def load_pretrained(model: torch.nn.Module, + path: str, + map_location='cpu', + logger=None, + sub_level=None): + if logger: + logger.info( + f'Load pretrained model [{model.__class__.__name__}] from {path}') + if os.path.exists(path): + # From local + state_dict = torch.load(path, map_location) + elif path.startswith('http'): + # From url + state_dict = load_state_dict_from_url(path, + map_location=map_location, + check_hash=False) + else: + raise Exception(f'Cannot find {path} when load pretrained') + + return load_pretrained_dict(model, state_dict, logger, sub_level=sub_level) + + +def _auto_drop_invalid(model: torch.nn.Module, state_dict: dict, logger=None): + """ Strip unmatched parameters in state_dict, e.g. shape not matched, type not matched. + + Args: + model (torch.nn.Module): + state_dict (dict): + logger (logging.Logger, None): + + Returns: + A new state dict. + """ + ret_dict = state_dict.copy() + invalid_msgs = [] + for key, value in model.state_dict().items(): + if key in state_dict: + # Check shape + new_value = state_dict[key] + if value.shape != new_value.shape: + invalid_msgs.append( + f'{key}: invalid shape, dst {value.shape} vs. src {new_value.shape}' + ) + ret_dict.pop(key) + elif value.dtype != new_value.dtype: + invalid_msgs.append( + f'{key}: invalid dtype, dst {value.dtype} vs. src {new_value.dtype}' + ) + ret_dict.pop(key) + if len(invalid_msgs) > 0: + warning_msg = 'ignore keys from source: \n' + '\n'.join(invalid_msgs) + if logger: + logger.warning(warning_msg) + else: + import warnings + warnings.warn(warning_msg) + return ret_dict + + +def load_pretrained_dict(model: torch.nn.Module, + state_dict: dict, + logger=None, + sub_level=None): + """ Load parameters to model with + 1. Sub name by revise_keys For DataParallelModel or DistributeParallelModel. + 2. Load 'state_dict' again if possible by key 'state_dict' or 'model_state'. + 3. Take sub level keys from source, e.g. load 'backbone' part from a classifier into a backbone model. + 4. Auto remove invalid parameters from source. + 5. Log or warning if unexpected key exists or key misses. + + Args: + model (torch.nn.Module): + state_dict (dict): dict of parameters + logger (logging.Logger, None): + sub_level (str, optional): If not None, parameters with key startswith sub_level will remove the prefix + to fit actual model keys. This action happens if user want to load sub module parameters + into a sub module model. + """ + revise_keys = [(r'^module\.', '')] + + if 'state_dict' in state_dict: + state_dict = state_dict['state_dict'] + if 'model_state' in state_dict: + state_dict = state_dict['model_state'] + + for p, r in revise_keys: + state_dict = {re.sub(p, r, k): v for k, v in state_dict.items()} + + if sub_level: + sub_level = sub_level if sub_level.endswith('.') else (sub_level + '.') + sub_level_len = len(sub_level) + state_dict = { + key[sub_level_len:]: value + for key, value in state_dict.items() if key.startswith(sub_level) + } + + state_dict = _auto_drop_invalid(model, state_dict, logger=logger) + + load_status = model.load_state_dict(state_dict, strict=False) + unexpected_keys = load_status.unexpected_keys + missing_keys = load_status.missing_keys + err_msgs = [] + if unexpected_keys: + err_msgs.append('unexpected key in source ' + f'state_dict: {", ".join(unexpected_keys)}\n') + if missing_keys: + err_msgs.append('missing key in source ' + f'state_dict: {", ".join(missing_keys)}\n') + err_msgs = '\n'.join(err_msgs) + + if len(err_msgs) > 0: + if logger: + logger.warning(err_msgs) + else: + import warnings + warnings.warn(err_msgs) + + +def count_params(model): + total_params = sum(p.numel() for p in model.parameters()) + return total_params + + +def init_weights(module): + if isinstance(module, (nn.Linear, nn.Embedding)): + module.weight.data.normal_(mean=0.0, std=0.02) + elif isinstance(module, nn.LayerNorm): + module.bias.data.zero_() + module.weight.data.fill_(1.0) + if isinstance(module, nn.Linear) and module.bias is not None: + module.bias.data.zero_() diff --git a/scepter/modules/utils/probe.py b/scepter/modules/utils/probe.py new file mode 100644 index 0000000..c4577a6 --- /dev/null +++ b/scepter/modules/utils/probe.py @@ -0,0 +1,391 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import copy +import os.path +from numbers import Number + +import numpy as np +import torch +from PIL import Image + +from scepter.modules.utils.file_system import FS + + +def check_legal_type(data): + if isinstance(data, str) or isinstance(data, Number): + return True + elif isinstance(data, dict): + for k, v in data.items(): + if not check_legal_type(v): + return False + return True + elif isinstance(data, list): + for v in data: + if not check_legal_type(v): + return False + return True + else: + return False + + +def register_data(probe_data: dict, key_prefix=''): + ret_data = {} + dist_data = {} + for k, v in probe_data.items(): + key = f'{key_prefix}_{k}' + if isinstance(v, torch.Tensor) or isinstance(v, np.ndarray): + ret_data[key] = ProbeData(v) + elif isinstance(v, ProbeData): + ret_data[key] = v + else: + if not check_legal_type(v): + raise f'The datatype of {key} should be included in [array, tensor, number, str] or the dict or ' \ + f'list of (number, str); if you want register the list of image, please use ProbeData instance.' + ret_data[key] = ProbeData(v) + if ret_data[key].view_distribute: + dist_data[key] = ret_data[key].distribute + + return ret_data, dist_data + + +def merge_gathered_probe(all_gathered_data): + ''' + Merge the gathered data on rank_0. + Returns: + The merged data. + + ''' + for key, gathered_data in all_gathered_data.items(): + # Must be the list of ProbeData. + if isinstance(gathered_data, list): + for v in gathered_data: + if not isinstance(v, ProbeData): + all_gathered_data[key] = gathered_data + # Must be the gathered data. + ret_data = gathered_data[0] + if not isinstance(ret_data.data, + list) and (isinstance(ret_data.data, np.ndarray) + or isinstance(ret_data.data, dict) + or check_legal_type(ret_data.data)): + new_data = [v.data for v in gathered_data] + if ret_data.build_label is not None: + ret_data.build_label = [ + v.build_label for v in gathered_data + ] + all_gathered_data[key] = ProbeData( + new_data, + is_image=ret_data.is_image, + build_html=ret_data.build_html, + build_label=ret_data.build_label, + view_distribute=ret_data.view_distribute) + elif isinstance(ret_data.data, list): + if ret_data.build_label is not None: + if isinstance(ret_data.build_label, str): + ret_data.build_label = [ + ret_data.build_label for _ in ret_data.data + ] + for v in gathered_data[1:]: + ret_data.data += v.data + if ret_data.build_label is not None: + if isinstance(v.build_label, str): + ret_data.build_label.extend( + [v.build_label for _ in v.data]) + ret_data.build_label.extend(v.build_label) + all_gathered_data[key] = ProbeData( + ret_data.data, + is_image=ret_data.is_image, + build_html=ret_data.build_html, + build_label=ret_data.build_label, + view_distribute=ret_data.view_distribute) + else: + all_gathered_data[key] = gathered_data + return all_gathered_data + + +class ProbeData(): + def __init__(self, + data, + is_image=False, + build_html=False, + build_label=None, + view_distribute=False): + ''' Probe Data Initialize. + We only support basic types such as [torch.Tensor, numpy.ndarray, number, str], + or [dict, list] of [dict, list, + number, str], or [dict, list] of [tensor, array] + ''' + data = copy.deepcopy(data) + self.basic_type = True + self._distribute_dict = {} + if view_distribute: + is_legal = True + if isinstance(data, str) or isinstance(data, Number): + is_legal = True + if data in self._distribute_dict: + self._distribute_dict[data] += 1 + else: + self._distribute_dict[data] = 1 + elif isinstance(data, list): + for v in data: + if isinstance(v, str) or isinstance(v, Number): + is_legal = True + if v in self._distribute_dict: + self._distribute_dict[v] += 1 + else: + self._distribute_dict[v] = 1 + else: + is_legal = False + elif isinstance(data, dict): + for k, v in data.items(): + if isinstance(v, str) or isinstance(v, Number): + is_legal = True + n_k = f'{k}_{v}' + if n_k in self._distribute_dict: + self._distribute_dict[n_k] += 1 + else: + self._distribute_dict[n_k] = 1 + else: + is_legal = False + else: + is_legal = False + if not is_legal: + print('Unsurpport data type', data) + assert is_legal + self.view_distribute = view_distribute + if isinstance(data, torch.Tensor): + self.data = data.detach().cpu().numpy() + elif isinstance(data, np.ndarray): + self.data = data + elif isinstance(data, dict): + for k, v in data.items(): + if not check_legal_type(v): + if isinstance(v, torch.Tensor): + data[k] = v.detach().cpu().numpy() + self.basic_type = False + elif isinstance(v, np.ndarray): + data[k] = v + self.basic_type = False + else: + raise f'Unsupport data type for {v}' + self.data = data + elif isinstance(data, list): + for idx, v in enumerate(data): + if not check_legal_type(v): + if isinstance(v, torch.Tensor): + data[idx] = v.detach().cpu().numpy() + self.basic_type = False + elif isinstance(v, np.ndarray): + data[idx] = v + self.basic_type = False + else: + raise f'Unsupport data type for {v}' + self.data = data + elif check_legal_type(data): + self.data = data + else: + raise f'Unsupport data type for {data}' + + self.is_image = is_image + self.build_html = build_html + + if self.build_html: + assert build_label is not None + if isinstance(self.data, str): + assert isinstance(build_label, str) + if isinstance(self.data, list): + assert isinstance(build_label, str) or isinstance( + build_label, list) + if isinstance(self.data, dict): + assert isinstance(build_label, str) or isinstance( + build_label, dict) + self.build_label = build_label + + def save_image(self, file_prefix, images): + np_shape = images.shape + # 4D + shape_str = '_'.join([str(v) for v in np_shape]) + if len(np_shape) == 4: + # channel is 1 or 3 + if np_shape[-1] == 1 or np_shape[-1] == 3: + if np_shape[-1] == 1: + images = images.reshape(images.shape[:-1]) + file_list = [] + for idx in range(np_shape[0]): + file_path = file_prefix + f'_probe_{idx}_[{shape_str}].jpg' + with FS.put_to(file_path) as local_path: + Image.fromarray(images[idx, ...]).save(local_path) + file_list.append(file_path) + return file_list + else: + raise f"Ensure your data's dim is BWHC, and channel is 1 or 3 for {file_prefix}" + elif len(np_shape) == 3: + if np_shape[-1] == 1 or np_shape[-1] == 3: + if np_shape[-1] == 1: + images = images.reshape(images.shape[:-1]) + file_path = file_prefix + f'_probe_[{shape_str}].jpg' + with FS.put_to(file_path) as local_path: + Image.fromarray(images).save(local_path) + return file_path + else: + images = images.reshape(list(images.shape) + [1]) + return self.save_image(file_prefix, images) + elif len(np_shape) == 2: + file_path = file_prefix + f'_probe_[{shape_str}].jpg' + with FS.put_to(file_path) as local_path: + Image.fromarray(images).save(local_path) + return file_path + else: + raise f"Ensure your data's dim is BWHC or WHC or WH, and channel is 1 or 3 for {file_prefix}" + + def save_npy(self, file_prefix, data): + shape_str = '_'.join([str(v) for v in data.shape]) + file_path = file_prefix + f'_{shape_str}.npy' + with FS.put_to(file_path) as local_path: + np.save(local_path, data) + return file_path + + def save_html(self, html_prefix, ret_data, ret_label): + height = 600 + with FS.put_to(html_prefix) as local_path: + with open(local_path, 'w') as f: + f.writelines('\n') + f.writelines('\n') + f.writelines('

\n') + all_ranks = list() + for save_id, save_data in enumerate(zip(ret_data, ret_label)): + save_path, save_label = save_data + one_rank = '' + for idx, one_data in enumerate(zip(save_path, save_label)): + one_path, one_label = one_data + one_label = one_label.replace('<', '<').replace( + '>', '>') + url = FS.get_url(one_path, + lifecycle=3600 * 365 * 24).replace( + '.oss-internal.aliyun-inc.', + '.oss.aliyuncs.') + one_rank += ( + f'' + ) + one_rank += '
' + f'
{save_id}-{idx}|{one_label}

' + all_ranks.append(one_rank) + f.writelines('\n'.join(all_ranks)) + return html_prefix + + @property + def distribute(self): + return self._distribute_dict + + def to_log(self, prefix=None): + if isinstance(self.data, np.ndarray): + if prefix is None: + raise 'You should provide the save prefix for array sample.' + # save jpg + if self.is_image: + ret_data = self.save_image(prefix, self.data) + if isinstance(ret_data, list): + ret_data = [ret_data] + if self.build_html: + ret_label = [] + if isinstance(self.build_label, str): + ret_label.append( + [self.build_label for _ in ret_data[0]]) + else: + ret_label.append(self.build_label) + if not len(ret_data[0]) == len(ret_label[0]): + raise f"The {prefix} label's length should be equal with 1st dim {self.data.shape[0]}." + html_prefix = prefix + '_probe.html' + html_file = self.save_html(html_prefix, ret_data, + ret_label) + return {'ori_file': ret_data, 'html': html_file} + else: + return {'ori_file': ret_data} + else: + return ret_data + else: + ret_data = self.save_npy(prefix, self.data) + return ret_data + elif isinstance(self.data, list): + if not self.basic_type: + ret_data = [] + ret_label = [] + for idx, v in enumerate(self.data): + prefix_path = os.path.join(prefix, f'{idx}') + if self.is_image: + ret_images = self.save_image(prefix_path, v) + ret_data.append(ret_images if isinstance( + ret_images, list) else [ret_images]) + if self.build_html: + if isinstance(ret_images, list): + if isinstance(self.build_label, str): + ret_label.append( + [self.build_label for _ in ret_images]) + elif isinstance(self.build_label[idx], list): + assert len(self.build_label[idx]) == len( + ret_images) + ret_label.append(self.build_label[idx]) + else: + ret_label.append([ + self.build_label[idx] + for _ in ret_images + ]) + else: + if isinstance(self.build_label, str): + ret_label.append([self.build_label]) + else: + ret_label.append([self.build_label[idx]]) + else: + ret_data.append(self.save_npy(prefix_path, v)) + if self.is_image and self.build_html: + html_prefix = prefix + '_probe.html' + html_file = self.save_html(html_prefix, ret_data, + ret_label) + return {'ori_file': ret_data, 'html': html_file} + else: + return {'ori_file': ret_data} + else: + return self.data + elif isinstance(self.data, dict): + if not self.basic_type: + ret_data = [] + ret_label = [] + for k, v in self.data: + prefix_path = os.path.join(prefix, f'{k}_') + if self.is_image: + ret_images = self.save_image(prefix_path, v) + if isinstance(ret_images, list): + ret_data.append(ret_images) + else: + ret_data.append([ret_images]) + if self.build_html: + if isinstance(ret_images, list): + if isinstance(self.build_label, str): + ret_label.append( + [self.build_label for _ in ret_images]) + elif isinstance(self.build_label[k], list): + assert len( + self.build_label[k]) == len(ret_images) + ret_label.append(self.build_label[k]) + else: + ret_label.append([ + self.build_label[k] for _ in ret_images + ]) + else: + if isinstance(self.build_label, str): + ret_label.append([self.build_label]) + else: + ret_label.append([self.build_label[k]]) + else: + ret_data.append(self.save_npy(prefix_path, v)) + if self.is_image and self.build_html: + html_prefix = prefix + '_probe.html' + html_file = self.save_html(html_prefix, ret_data, + ret_label) + return {'ori_file': ret_data, 'html': html_file} + else: + return {'ori_file': ret_data} + else: + return self.data + else: + return self.data diff --git a/scepter/modules/utils/registry.py b/scepter/modules/utils/registry.py new file mode 100644 index 0000000..39cae95 --- /dev/null +++ b/scepter/modules/utils/registry.py @@ -0,0 +1,212 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +# Modified Based on the following original code. + +# Registry class & build_from_config function partially modified from +# https://github.com/open-mmlab/mmcv/blob/master/mmcv/utils/registry.py +# Copyright 2018-2020 Open-MMLab. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import inspect +import sys +import warnings + +from scepter.modules.utils.config import dict_to_yaml + +old_python_version = '3.6' in sys.version +if old_python_version: + + def deep_copy(obj): + return obj +else: + import copy + + def deep_copy(obj): + return copy.deepcopy(obj) + + +def build_from_config(cfg, registry, logger=None, *args, **kwargs): + """ Default builder function. + + Args: + cfg (objective attribution): A set of objective attirbutions + which contain parameters passes to target class or function. + Must contains key 'type', indicates the target class or function name. + registry (Registry): An registry to search target class or function. + kwargs (dict, optional): Other params not in config dict. + + Returns: + Target class object or object returned by invoking function. + + Raises: + TypeError: + KeyError: + Exception: + """ + from scepter.modules.utils.config import Config + if not isinstance(cfg, Config): + raise TypeError(f'config must be type dict, got {type(cfg)}') + if not cfg.have('NAME'): + raise KeyError(f'config must contain key NAME, got {cfg}') + if not isinstance(registry, Registry): + raise TypeError( + f'registry must be type Registry, got {type(registry)}') + + cfg = deep_copy(cfg) + req_type = cfg.get('NAME') + if isinstance(req_type, str): + req_type_entry = registry.get(req_type) + if req_type_entry is None: + raise KeyError(f'{req_type} not found in {registry.name} registry') + + if kwargs is not None: + cfg._update_dict(kwargs) + + if inspect.isclass(req_type_entry): + try: + return req_type_entry(cfg, logger=logger, *args, **kwargs) + except Exception as e: + raise Exception(f'Failed to init class {req_type_entry}, with {e}') + elif inspect.isfunction(req_type_entry): + try: + return req_type_entry(cfg, logger=logger, *args, **kwargs) + except Exception as e: + raise Exception( + f'Failed to invoke function {req_type_entry}, with {e}') + else: + raise TypeError( + f'type must be str or class, got {type(req_type_entry)}') + + +REGISTRY_LIST = [] + + +class Registry(object): + """ A registry maps key to classes or functions. + + Example: + # >>> MODELS = Registry('MODELS') + # >>> @MODELS.register_class() + # >>> class ResNet(object): + # >>> pass + # >>> config = Config(cfg_dict = {"NAME":"ResNet"}) + # >>> resnet = MODELS.build(config) + # >>> + # >>> import torchvision + # >>> @MODELS.register_function("InceptionV3") + # >>> def get_inception_v3(pretrained=False, progress=True): + # >>> return torchvision.model.inception_v3(pretrained=pretrained, progress=progress) + # >>> config = Config(cfg_dict = {"NAME":"InceptionV3"}) + # >>> inception_v3 = MODELS.build(config) + + Args: + name (str): Registry name. + build_func (func, None): Instance construct function. Default is build_from_config. + allow_types (tuple): Indicates how to construct the instance, by constructing class or invoking function. + """ + def __init__(self, + name, + build_func=None, + common_para=None, + allow_types=('class', 'function')): + self.name = name + self.allow_types = allow_types + self.class_map = {} + self.func_map = {} + self.common_para = common_para + self.build_func = build_func or build_from_config + REGISTRY_LIST.append(self) + + def get(self, req_type): + return self.class_map.get(req_type) or self.func_map.get(req_type) + + def build(self, cfg, logger=None, *args, **kwargs): + return self.build_func(cfg, + registry=self, + logger=logger, + *args, + **kwargs) + + def register_class(self, name=None): + def _register(cls): + if not inspect.isclass(cls): + raise TypeError(f'Module must be type class, got {type(cls)}') + if 'class' not in self.allow_types: + raise TypeError( + f'Register {self.name} only allows type {self.allow_types}, got class' + ) + module_name = name or cls.__name__ + if module_name in self.class_map: + warnings.warn( + f'Class {module_name} already registered by {self.class_map[module_name]}, ' + f'will be replaced by {cls}') + self.class_map[module_name] = cls + return cls + + return _register + + def register_function(self, name=None): + def _register(func): + if not inspect.isfunction(func): + raise TypeError( + f'Registry must be type function, got {type(func)}') + if 'function' not in self.allow_types: + raise TypeError( + f'Registry {self.name} only allows type {self.allow_types}, got function' + ) + func_name = name or func.__name__ + if func_name in self.class_map: + warnings.warn( + f'Function {func_name} already registered by {self.func_map[func_name]}, ' + f'will be replaced by {func}') + self.func_map[func_name] = func + return func + + return _register + + def _list(self): + keys = sorted(list(self.class_map.keys()) + list(self.func_map.keys())) + descriptions = [] + for key in keys: + if key in self.class_map: + descriptions.append(f'{key}: {self.class_map[key]}') + else: + descriptions.append( + f"{key}: " + ) + return '\n'.join(descriptions) + + def __repr__(self): + description = self._list() + description = '\n'.join(['\t' + s for s in description.split('\n')]) + return f'{self.__class__.__name__} [{self.name}], \n' + description + + def get_config_template(self, name): + common_yaml_str = '' + if self.common_para is not None: + common_yaml_str += 'The following para are used for this class.\n' + common_yaml_str += dict_to_yaml('common_parameter', + __class__.__name__, + self.common_para, + set_name=False) + + req_type_entry = self.get(name) + if req_type_entry is None: + raise KeyError(f'{name} not found in {self.name} registry') + if inspect.isclass(req_type_entry): + return req_type_entry.get_config_template() + common_yaml_str + elif inspect.isfunction(req_type_entry): + return '{} is a function!'.format(name) + else: + return 'Unsurport object type!' diff --git a/scepter/modules/utils/video_reader/__init__.py b/scepter/modules/utils/video_reader/__init__.py new file mode 100644 index 0000000..5392d67 --- /dev/null +++ b/scepter/modules/utils/video_reader/__init__.py @@ -0,0 +1,6 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from .frame_sampler import (FRAME_SAMPLERS, IntervalSampler, SegmentSampler, + UniformSampler, do_frame_sample) +from .video_reader import (EasyVideoReader, FramesReaderWrapper, + VideoReaderWrapper) diff --git a/scepter/modules/utils/video_reader/frame_sampler.py b/scepter/modules/utils/video_reader/frame_sampler.py new file mode 100644 index 0000000..943eb12 --- /dev/null +++ b/scepter/modules/utils/video_reader/frame_sampler.py @@ -0,0 +1,165 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +""" +FrameSampler. + +Sample: + 1. give start & end time, num_frames, e.g. 16 frames from [1.0s ,3.0s], + usually used in real applications. + fixed args: + `sample_type`='uniform'; + `vid_len` (int): valid total frame numbers in video; + `vid_fps` (float): video fps; + `num_frames` (int): number of frames to be extracted; + extra args: + + 2. give a fixed clip duration, num_frames, e.g. 16 frames from a 2s clip. + In train mode (`clip_id`=-1), this clip will be randomly sampled from video. + In test mode (`clip_id`>=0), this clip is the center part of the video + which acts the same as `DecodeVideoToTensor` op. + fixed args: `sample_mode`='interval', `vid_len`, `vid_fps`, `num_frames` + extra args: `clip_duration`, `clip_id`, `num_clips`=1 + + 3. give a fixed clip duration, constant total clips, current clip index, num_frames, + e.g. three 2-s clips will be sampled from the video, and 16 frames from the first clip. + Usually used in multi-view test. + In train mode (`clip_id`=-1), constant total clips will be ignored, so this will act the same b. + In test mode (`clip_id`>=0), video is splitted into constant clips (uniformly and allow overlap, + a 3s video splits into three 2s clips, [0, 2), [0.5, 2.5), [1.0, 3.0) ), + then sample frames from one clip. + fixed args: `sample_mode`='interval', `vid_len`, `vid_fps`, `num_frames` + extra args: `clip_duration`, `clip_id`, `num_clips` + + 4. give num_frames, do segment-sampling, e.g. 16 frames from whole video, then splits the video into 16 segments, + and sample one frame from each segment. + In train mode (`clip_id`=-1), sample a frame randomly from a segment. + In test mode (`clip_id`>=0), the center frame in each part will be chosen. + fixed args: `sample_mode`='segment', `vid_len`, `vid_fps`, `num_frames` + call args: `clip_id`, `num_clips`=1 + + 5. give constant total clips, current clip index, num_frames, + e.g. splits the video into 16 segments, split one segment into 3 parts, + if clip_index=0, sample one frame from the first part, and loop 16 times. + In train mode, sample a frame randomly from a segment. + In test mode, int(`clip_id`/`num_clips` * segment_frames) will be chosen. + fixed args: `sample_mode`='segment', `vid_len`, `vid_fps`, `num_frames` + call args: `clip_id`, `num_clips` + +Output: + A list of frame indices (torch.Tensor) + +""" +import math +import random + +import torch + +from scepter.modules.utils.config import Config +from scepter.modules.utils.registry import Registry + +FRAME_SAMPLERS = Registry('FRAME_SAMPLERS') + + +def do_frame_sample(sampling_type: str, vid_len: int, vid_fps: float, + num_frames: int, **kwargs) -> torch.Tensor: + params = dict(vid_len=vid_len, + vid_fps=vid_fps, + num_frames=num_frames, + **kwargs) + return FRAME_SAMPLERS.build( + Config(cfg_dict={'NAME': sampling_type}, load=False))(**params) + + +@FRAME_SAMPLERS.register_class('uniform') +class UniformSampler(object): + def __init__(self, cfg, logger=None): + self.cfg = cfg + self.logger = logger + + def __call__(self, + vid_len: int, + vid_fps: float, + num_frames: int, + start_sec: float = 0, + end_sec: float = -1) -> torch.Tensor: + start_sec = max(start_sec, 0) + if end_sec < 0: + new_end_sec = vid_len / vid_fps + else: + new_end_sec = min(end_sec, vid_len / vid_fps) + assert new_end_sec > start_sec, ( + f'end_sec should be greater then start_sec, ' + f'got end_sec={new_end_sec}, start_sec={start_sec}') + end_sec = new_end_sec + + start_idx = math.floor(start_sec / vid_fps) + end_idx_exc = min(vid_len, math.ceil(end_sec / vid_fps)) + + index = torch.linspace(start_idx, end_idx_exc, num_frames) + index = torch.clamp(index, 0, vid_len - 1).long() + + return index + + +@FRAME_SAMPLERS.register_class('interval') +class IntervalSampler(object): + def __init__(self, cfg, logger=None): + self.cfg = cfg + self.logger = logger + + def __call__(self, + vid_len: int, + vid_fps: float, + num_frames: int, + clip_duration: float, + clip_id: int = 0, + num_clips: int = 1) -> torch.Tensor: + if num_frames == 1: + return torch.randint(0, vid_len, (1, )) + + clip_len = int(clip_duration / vid_fps) + max_idx = max(vid_len, clip_len, 0) + + if clip_id == -1: + start_idx = random.uniform(0, max_idx) + else: + if num_clips == 1: + start_idx = max_idx / 2 + else: + start_idx = max_idx * clip_id / num_clips + + end_idx = start_idx + clip_len - 1 + index = torch.linspace(start_idx, end_idx, num_frames) + index = torch.clamp(index, 0, vid_len - 1).long() + return index + + +@FRAME_SAMPLERS.register_class('segment') +class SegmentSampler(object): + def __init__(self, cfg, logger=None): + self.cfg = cfg + self.logger = logger + + def __call__(self, + vid_len: int, + vid_fps: float, + num_frames: int, + clip_id: int = 0, + num_clips: int = 1) -> torch.Tensor: + index = torch.zeros(num_frames) + index_range = torch.linspace(0, vid_len, num_frames + 1) + for idx in range(num_frames): + if clip_id == -1: + index[idx] = random.uniform(index_range[idx], + index_range[idx + 1]) + else: + if num_clips == 1: + index[idx] = (index_range[idx] + index_range[idx + 1]) / 2 + else: + index[idx] = index_range[idx] + ( + index_range[idx + 1] - + index_range[idx]) * (clip_id + 1) / num_clips + + index = torch.round(torch.clamp(index, 0, vid_len - 1)).long() + + return index diff --git a/scepter/modules/utils/video_reader/video_reader.py b/scepter/modules/utils/video_reader/video_reader.py new file mode 100644 index 0000000..ff6684e --- /dev/null +++ b/scepter/modules/utils/video_reader/video_reader.py @@ -0,0 +1,166 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import os +from fractions import Fraction +from typing import Callable, Optional, Union + +import cv2 +import numpy as np +import torch +import torch.utils.dlpack as dlpack + +from scepter.modules.utils.file_system import FS +from scepter.modules.utils.video_reader.frame_sampler import do_frame_sample + + +class _Wrapper(object): + @property + def len(self) -> int: + raise NotImplementedError + + @property + def fps(self) -> float: + raise NotImplementedError + + @property + def duration(self) -> float: + raise NotImplementedError + + def sample_frames(self, decode_list: torch.Tensor) -> torch.Tensor: + raise NotImplementedError + + +class VideoReaderWrapper(_Wrapper): + def __init__(self, video_path): + import decord + self._video_path = video_path + self._decoder_type = 'decord' + self._vr = decord.VideoReader(self._video_path) + + @property + def len(self): + return len(self._vr) + + @property + def fps(self): + return self._vr.get_avg_fps() + + @property + def duration(self): + return float(self.len) / self.fps + + def __del__(self): + if self._vr is not None: + del self._vr + self._vr = None + + def sample_frames(self, decode_list: torch.Tensor) -> torch.Tensor: + frames = dlpack.from_dlpack( + self._vr.get_batch(decode_list).to_dlpack()).clone() + return frames + + +class FramesReaderWrapper(_Wrapper): + def __init__(self, frame_dir: str, extract_fps: float, suffix='.jpg'): + self._frame_dir = frame_dir + self._extract_fps = extract_fps + self._suffix = suffix + self._frame_list = sorted([ + os.path.join(self._frame_dir, t) + for t in os.listdir(self._frame_dir) if t.endswith(self._suffix) + ]) + self._frames = [None] * len(self._frame_list) + + @property + def len(self) -> int: + return len(self._frame_list) + + @property + def fps(self): + return self._extract_fps + + @property + def duration(self): + return float(self.len) / self.fps + + def _load_frame(self, idx): + path = self._frame_list[idx] + img = cv2.imread(path, cv2.IMREAD_COLOR) + img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) + self._frames[idx] = img + + def sample_frames(self, decode_list: torch.Tensor) -> torch.Tensor: + ret = [] + for idx in decode_list.numpy(): + if self._frames[idx] is None: + self._load_frame(idx) + ret.append(self._frames[idx].copy()) + ret = np.asarray(ret) + return torch.from_numpy(ret) + + +class EasyVideoReader(object): + """ A video reader which is easy to use in real applications. + + Args: + video_path (str): Path of video file. + num_frames (int): Extract frames for one sample. + clip_duration (Union[float, Fraction, str]): Clip duration to be extracted uniformly. + overlap (Union[float, Fraction, str]): The offset (in secs) of + the next clip overlaps the last clip, default is 0 no overlap. + transforms (Optional[Callable]): Do transform operations, default is None. + + """ + def __init__(self, + video_path: str, + num_frames: int, + clip_duration: Union[float, Fraction, str], + overlap: Union[float, Fraction, str] = Fraction(0), + transforms: Optional[Callable] = None): + self._video_path: str = video_path + self._num_frames: int = num_frames + self._clip_duration: Fraction = Fraction(clip_duration) + self._overlap: Fraction = Fraction(overlap) + assert self._overlap < self._clip_duration, 'Overlap must be smaller than clip_duration!' + self._transforms = transforms + + self._last_end: Fraction = Fraction(0) + + client = FS.get_fs_client(self._video_path) + local_path = client.get_object_to_local_file(self._video_path) + self._vr = VideoReaderWrapper(local_path) + + def __iter__(self): + return self + + def __next__(self): + start_sec = max(Fraction(0), self._last_end - self._overlap) + end_sec = start_sec + self._clip_duration + + if end_sec > self._vr.duration: + del self._vr + raise StopIteration + + decode_list = do_frame_sample('uniform', + self._vr.len, + self._vr.fps, + self._num_frames, + start_sec=float(start_sec), + end_sec=float(end_sec)) + output_tensor = self._vr.sample_frames(decode_list) + + self._last_end = end_sec + + output = { + 'video': output_tensor, + 'meta': { + 'video_path': self._video_path, + 'start_sec': float(start_sec), + 'end_sec': float(end_sec) + } + } + + if self._transforms is not None: + return self._transforms(output) + return output diff --git a/scepter/tools/__init__.py b/scepter/tools/__init__.py new file mode 100644 index 0000000..cc26a06 --- /dev/null +++ b/scepter/tools/__init__.py @@ -0,0 +1,2 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. diff --git a/scepter/tools/helper.py b/scepter/tools/helper.py new file mode 100644 index 0000000..4dadeb1 --- /dev/null +++ b/scepter/tools/helper.py @@ -0,0 +1,98 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import argparse +import os +import sys + +# +from scepter.modules.utils.registry import REGISTRY_LIST + +sys.path.insert(0, os.path.abspath(os.curdir)) + + +def get_module_list(): + ret_msg = [v.name for v in REGISTRY_LIST] + print(ret_msg) + return None + + +def get_module_objects(module_name): + ret_msg = 'Not surpport module!' + for v in REGISTRY_LIST: + if v.name == module_name: + class_map = v.class_map + func_map = v.func_map + ret_msg = 'The {} module surpport the following object: \n'.format( + module_name) + index = 0 + for k in class_map: + ret_msg += '{}: class {}\n'.format(index, k) + index += 1 + for k in func_map: + ret_msg += '{}: function {}\n'.format(index, k) + index += 1 + print(ret_msg) + return None + + +def get_module_object_config(module_name, object_name): + ret_msg = 'Not surpport module or object!' + for v in REGISTRY_LIST: + if v.name == module_name: + class_map = v.class_map + func_map = v.func_map + + for k in class_map: + if object_name == k: + ret_msg = 'The {module_name}:{object_name} object need follow config: \n' + ret_msg += v.get_config_template(k) + for k in func_map: + if object_name == k: + ret_msg = 'The {module_name}:{object_name} object need follow config: \n' + ret_msg += v.get_config_template(k) + print(ret_msg) + return None + + +if __name__ == '__main__': + # initialize the data manager instance + + usage_string = 'usage for the torch dist: \n' \ + '1. view mudule list \n' \ + ' examples: \n' \ + ' ' + + parser = argparse.ArgumentParser(usage=usage_string) + + parser.add_argument('-t', + '--tool', + dest='tool_type', + type=str, + choices=['lm', 'lo', 'config'], + default='data_info', + help='choose your operation for torch dist!') + + parser.add_argument('-m', + '--module', + dest='module', + type=str, + choices=get_module_list(), + default='', + help='choose the module from {}!'.format( + get_module_list())) + + parser.add_argument('-o', + '--object', + dest='object', + type=str, + default='', + help='choose the object!') + + args = parser.parse_args() + + if args.tool_type == 'lm': + get_module_list() + if args.tool_type == 'lo': + get_module_objects(args.module) + if args.tool_type == 'config': + get_module_object_config(args.module, args.object) diff --git a/scepter/tools/local.sh b/scepter/tools/local.sh new file mode 100644 index 0000000..6a0ef5b --- /dev/null +++ b/scepter/tools/local.sh @@ -0,0 +1,2 @@ +export CUDA_VISIBLE_DEVICES=0 +python -W ignore scepter/tools/run_inference.py --cfg scepter/methods/examples/generation/stable_diffusion_xl_1024_align.yaml diff --git a/scepter/tools/run_inference.py b/scepter/tools/run_inference.py new file mode 100644 index 0000000..2d0e98f --- /dev/null +++ b/scepter/tools/run_inference.py @@ -0,0 +1,132 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import argparse +import os + +import cv2 +import numpy as np +import torch +import torch.cuda.amp as amp + +from scepter.modules.solver.registry import SOLVERS +from scepter.modules.utils.config import Config +from scepter.modules.utils.data import transfer_data_to_cuda +from scepter.modules.utils.distribute import we +from scepter.modules.utils.file_system import FS +from scepter.modules.utils.logger import get_logger + + +def run_task(cfg): + std_logger = get_logger(name='scepter') + solver = SOLVERS.build(cfg.SOLVER, logger=std_logger) + solver.set_up() + if not cfg.args.pretrained_model == '': + with FS.get_from(cfg.args.pretrained_model, + wait_finish=True) as local_path: + solver.model.load_state_dict( + torch.load(local_path, map_location='cuda')['model']) + solver.test_mode() + num_samples = cfg.args.num_samples + prompt = [cfg.args.prompt] * num_samples + n_prompt = [cfg.args.n_prompt] * num_samples + sampler = cfg.args.sampler + sample_steps = cfg.args.sample_steps + seed = cfg.args.seed + guide_scale = cfg.args.guide_scale + guide_rescale = cfg.args.guide_rescale + image_size = cfg.args.image_size + if image_size is not None: + if ',' in image_size: + h, w = image_size.split(',') + image_size = [int(h), int(w)] + else: + image_size = [int(image_size), int(image_size)] + + batch_data = {} + if solver.sample_args: + batch_data.update(solver.sample_args.get_lowercase_dict()) + if image_size is not None: + batch_data.update({'image_size': image_size}) + batch_data.update({ + 'prompt': prompt, + 'n_prompt': n_prompt, + 'sampler': sampler, + 'sample_steps': sample_steps, + 'seed': seed, + 'guide_scale': guide_scale, + 'guide_rescale': guide_rescale, + }) + + dtype = getattr(torch, cfg.SOLVER.DTYPE) + with amp.autocast(enabled=True, dtype=dtype): + batch_data = transfer_data_to_cuda(batch_data) + ret = solver.run_step_test(batch_data) + save_folder = os.path.join(solver.work_dir, cfg.args.save_folder) + for idx, out in enumerate(ret): + img = out['image'] + img = img.permute(1, 2, 0).cpu().numpy() + img = (img * 255).astype(np.uint8) + filename = '{}_{}.png'.format('inference', idx) + save_file = os.path.join(save_folder, filename) + with FS.put_to(save_file) as local_path: + image = img.copy() + cv2.cvtColor(image, cv2.COLOR_RGB2BGR, image) + cv2.imwrite(local_path, image) + + +if __name__ == '__main__': + parser = argparse.ArgumentParser(description='Argparser for Scepter:\n') + parser.add_argument( + '--prompt', + dest='prompt', + help='Prompt sentence!', + default='a woman is walking on the street in a rainy day.') + parser.add_argument('--n_prompt', + dest='n_prompt', + help='Add Prompt sentence!', + default='') + parser.add_argument('--num_samples', + dest='num_samples', + help="Output image's number!", + default=4, + type=int) + parser.add_argument('--sampler', + dest='sampler', + help='Sampler method!', + default='ddim', + type=str) + parser.add_argument('--sample_steps', + dest='sample_steps', + help='Sample steps!', + default=50, + type=int) + parser.add_argument('--seed', + dest='seed', + help='Random seed!', + default=2023, + type=int) + parser.add_argument('--guide_scale', + dest='guide_scale', + help='Guidance scale!', + default=7.5, + type=float) + parser.add_argument('--guide_rescale', + dest='guide_rescale', + help='Guidance rescale!', + default=0.5, + type=float) + parser.add_argument('--image_size', + dest='image_size', + help='Output image size! (h, w)', + default=None, + type=str) + parser.add_argument('--save_folder', + dest='save_folder', + help="Output image's save folder!", + default='test_images') + parser.add_argument('--pretrained_model', + dest='pretrained_model', + help='The pretrained model for our network!', + default='') + cfg = Config(load=True, parser_ins=parser) + we.init_env(cfg, logger=None, fn=run_task) diff --git a/scepter/tools/run_train.py b/scepter/tools/run_train.py new file mode 100644 index 0000000..0761a9e --- /dev/null +++ b/scepter/tools/run_train.py @@ -0,0 +1,23 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import argparse + +from scepter.modules.solver.registry import SOLVERS +from scepter.modules.utils.config import Config +from scepter.modules.utils.distribute import we +from scepter.modules.utils.logger import get_logger + + +def run_task(cfg): + std_logger = get_logger(name='scepter') + solver = SOLVERS.build(cfg.SOLVER, logger=std_logger) + solver.set_up_pre() + solver.set_up() + solver.solve() + + +if __name__ == '__main__': + parser = argparse.ArgumentParser(description='Argparser for Scepter:\n') + + cfg = Config(load=True, parser_ins=parser) + we.init_env(cfg, logger=None, fn=run_task) diff --git a/scepter/version.py b/scepter/version.py new file mode 100644 index 0000000..737ae04 --- /dev/null +++ b/scepter/version.py @@ -0,0 +1,8 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +__version__ = '1.0.0' + +version_info = tuple(int(x) for x in __version__.split('.')[0:3]) + +__all__ = ['__version__', version_info] diff --git a/setup.py b/setup.py new file mode 100644 index 0000000..a65c893 --- /dev/null +++ b/setup.py @@ -0,0 +1,162 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import glob +import os + +import setuptools + + +def get_long_description(): + with open('readme.md') as f: + long_description = f.read() + return long_description + + +def get_version(): + version_path = 'scepter/version.py' + with open(version_path) as f: + exec(compile(f.read(), version_path, 'exec')) + return locals()['__version__'] + + +def parse_requirements(fname='requirements.txt', with_version=True): + """ + Parse the package dependencies listed in a requirements file but strips + specific versioning information. + + Args: + fname (str): path to requirements file + with_version (bool, default=False): if True include version specs + + Returns: + List[str]: list of requirements items + + CommandLine: + python -c "import setup; print(setup.parse_requirements())" + """ + import re + import sys + from os.path import exists + require_fpath = fname + + def parse_line(line): + """ + Parse information from a line in a requirements text file + """ + if line.startswith('-r '): + # Allow specifying requirements in other files + target = line.split(' ')[1] + relative_base = os.path.dirname(fname) + absolute_target = os.path.join(relative_base, target) + for info in parse_require_file(absolute_target): + yield info + else: + info = {'line': line} + if line.startswith('-e '): + info['package'] = line.split('#egg=')[1] + else: + # Remove versioning from the package + pat = '(' + '|'.join(['>=', '==', '>']) + ')' + parts = re.split(pat, line, maxsplit=1) + parts = [p.strip() for p in parts] + + info['package'] = parts[0] + if len(parts) > 1: + op, rest = parts[1:] + if ';' in rest: + # Handle platform specific dependencies + # http://setuptools.readthedocs.io/en/latest/setuptools.html#declaring-platform-specific-dependencies + version, platform_deps = map(str.strip, + rest.split(';')) + info['platform_deps'] = platform_deps + else: + version = rest # NOQA + info['version'] = (op, version) + yield info + + def parse_require_file(fpath): + with open(fpath, 'r', encoding='utf-8') as f: + for line in f.readlines(): + line = line.strip() + if line.startswith('http'): + print('skip http requirements %s' % line) + continue + if line and not line.startswith('#') and not line.startswith( + '--'): + for info in parse_line(line): + yield info + elif line and line.startswith('--find-links'): + eles = line.split() + for e in eles: + e = e.strip() + if 'http' in e: + info = dict(dependency_links=e) + yield info + + def gen_packages_items(): + items = [] + deps_link = [] + if exists(require_fpath): + for info in parse_require_file(require_fpath): + if 'dependency_links' not in info: + parts = [info['package']] + if with_version and 'version' in info: + parts.extend(info['version']) + if not sys.version.startswith('3.4'): + # apparently package_deps are broken in 3.4 + platform_deps = info.get('platform_deps') + if platform_deps is not None: + parts.append(';' + platform_deps) + item = ''.join(parts) + items.append(item) + else: + deps_link.append(info['dependency_links']) + return items, deps_link + + return gen_packages_items() + + +def backupfile(): + contents = open('scepter/__init__.py', 'r').read() + with open('scepter/__init__.py', 'w') as f: + f.write(contents.replace('from scepter import task', '')) + return contents + + +def restorefile(contents): + with open('scepter/__init__.py', 'w') as f: + f.write(contents) + + +required = parse_requirements() + +contents = backupfile() + +setuptools.setup( + name='scepter', + version=get_version(), + author='Tongyi Lab', + author_email='', + description='', + keywords='compute vision, framework, generation, image edition.', + long_description=get_long_description(), + long_description_content_type='text/markdown', + url='', + packages=[ + pkg for pkg in setuptools.find_packages() + if '__pycache__' not in pkg and 'scepter' in pkg + ], + data_files=[('lib/docs', glob.glob('docs/*.md') + glob.glob('docs/*/*.md')) + ], + classifiers=[ + 'Programming Language :: Python :: 3.8', + 'Programming Language :: Python :: 3.9', + 'Programming Language :: Python :: 3.10', + 'License :: OSI Approved :: Apache License (Version 2.0)', + 'Operating System :: OS Independent', + ], + python_requires='>=3.8', + install_requires=required, +) + +restorefile(contents) diff --git a/tests/tools/test_inference.py b/tests/tools/test_inference.py new file mode 100644 index 0000000..1462e49 --- /dev/null +++ b/tests/tools/test_inference.py @@ -0,0 +1,65 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import os +import unittest + + +class InferenceTest(unittest.TestCase): + def setUp(self): + print(('Testing %s.%s' % (type(self).__name__, self._testMethodName))) + + def tearDown(self): + super().tearDown() + + # @unittest.skip('') + def test_infer_args(self): + os.system("python scepter/tools/run_inference.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml " + "--prompt 'a cute dog' --save_folder 'test_prompt_a_cute_dog_2023'") + os.system("python scepter/tools/run_inference.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml " + "--prompt 'a cute dog' --save_folder 'test_prompt_a_cute_dog_2024' --seed 2024") + os.system("python scepter/tools/run_inference.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml " + "--prompt 'a cute dog' --save_folder 'test_prompt_a_cute_dog_size768' " + "--image_size '768'") + os.system("python scepter/tools/run_inference.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml " + "--prompt 'a cute dog' --save_folder 'test_prompt_a_cute_dog_1280_720' " + "--image_size '1280,720'") + os.system("python scepter/tools/run_inference.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml " + "--prompt 'a cute dog' --save_folder 'test_prompt_a_cute_dog_720_1280_step10' " + "--image_size '720,1280' --sample_steps 10") + os.system("python scepter/tools/run_inference.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml " + "--prompt 'a cute dog' --save_folder 'test_prompt_a_cute_dog_num2' --num_samples 2") + os.system("python scepter/tools/run_inference.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml " + "--prompt 'a cute dog' --save_folder 'test_prompt_a_cute_dog_dpmpp_2s_ancestral' " + "--sampler 'dpmpp_2s_ancestral'") + os.system("python scepter/tools/run_inference.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml " + "--prompt 'a cute dog' --save_folder 'test_prompt_a_cute_dog_scale5' " + "--guide_scale 5.0") + os.system("python scepter/tools/run_inference.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml " + "--prompt 'a cute dog' --save_folder 'test_prompt_a_cute_dog_rescale0_1' " + "--guide_rescale 0.1") + + # @unittest.skip('') + def test_example_infer(self): + os.system("python scepter/tools/run_inference.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml " + "--prompt 'a cute dog' --save_folder 'test_prompt_a_cute_dog'") + os.system("python scepter/tools/run_inference.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_2.1_768.yaml " + "--prompt 'a cute dog' --save_folder 'test_prompt_a_cute_dog'") + os.system("python scepter/tools/run_inference.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_xl_1024.yaml " + "--prompt 'a cute dog' --save_folder 'test_prompt_a_cute_dog'") + + +if __name__ == '__main__': + unittest.main() diff --git a/tests/tools/test_train.py b/tests/tools/test_train.py new file mode 100644 index 0000000..7f0216c --- /dev/null +++ b/tests/tools/test_train.py @@ -0,0 +1,45 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import os +import unittest + + +class TrainTest(unittest.TestCase): + def setUp(self): + print(('Testing %s.%s' % (type(self).__name__, self._testMethodName))) + + def tearDown(self): + super().tearDown() + + # @unittest.skip('') + def test_generation_example_full(self): + os.system("python scepter/tools/run_train.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_1.5_512.yaml ") + os.system("python scepter/tools/run_train.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_2.1_768.yaml ") + os.system("python scepter/tools/run_train.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_xl_1024.yaml ") + + # @unittest.skip('') + def test_generation_example_lora(self): + os.system("python scepter/tools/run_train.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_1.5_512_lora.yaml ") + os.system("python scepter/tools/run_train.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_2.1_768_lora.yaml ") + os.system("python scepter/tools/run_train.py " + "--cfg scepter/methods/examples/generation/stable_diffusion_xl_1024_lora.yaml ") + + # @unittest.skip('') + def test_generation_example_scedit(self): + os.system("python scepter/tools/run_train.py " + "--cfg scepter/methods/SCEdit/t2i_sd15_512_sce.yaml ") + os.system("python scepter/tools/run_train.py " + "--cfg scepter/methods/SCEdit/t2i_sd21_768_sce.yaml ") + os.system("python scepter/tools/run_train.py " + "--cfg scepter/methods/SCEdit/t2i_sdxl_1024_sce.yaml ") + + + +if __name__ == '__main__': + unittest.main() diff --git a/tests/utils/test_fs.py b/tests/utils/test_fs.py new file mode 100644 index 0000000..969ba93 --- /dev/null +++ b/tests/utils/test_fs.py @@ -0,0 +1,50 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. + +import os +import unittest + +from scepter.modules.utils.config import Config +from scepter.modules.utils.file_system import FS + + +class FSTest(unittest.TestCase): + def setUp(self): + print(('Testing %s.%s' % (type(self).__name__, self._testMethodName))) + + def tearDown(self): + super().tearDown() + + def test_modelscope(self): + fs_info = {'NAME': 'ModelscopeFs', 'TEMP_DIR': 'cache/data'} + config = Config(load=False, cfg_dict=fs_info) + FS.init_fs_client(config) + + path = 'ms://AI-ModelScope/stable-diffusion-v1-5' + with FS.get_dir_to_local_dir(path, wait_finish=True) as local_path: + print(f'Download from {path} to {local_path}') + self.assertTrue(os.path.exists(local_path)) + + path = 'ms://AI-ModelScope/stable-diffusion-v1-5:v1.0.8' + with FS.get_dir_to_local_dir(path, wait_finish=True) as local_path: + print(f'Download from {path} to {local_path}') + self.assertTrue(os.path.exists(local_path)) + + path = 'ms://AI-ModelScope/stable-diffusion-v1-5@configuration.json' + with FS.get_from(path, wait_finish=True) as local_path: + print(f'Download from {path} to {local_path}') + self.assertTrue(os.path.exists(local_path)) + + path = 'ms://AI-ModelScope/stable-diffusion-v1-5:v1.0.8@v1-5-pruned-emaonly.ckpt' + with FS.get_from(path, wait_finish=True) as local_path: + print(f'Download from {path} to {local_path}') + self.assertTrue(os.path.exists(local_path)) + + path = 'ms://AI-ModelScope/stable-diffusion-v1-5:v1.0.8@text_encoder' + with FS.get_dir_to_local_dir(path, wait_finish=True) as local_path: + print(f'Download from {path} to {local_path}') + self.assertTrue(os.path.exists(local_path)) + + +if __name__ == '__main__': + unittest.main()