new project from v0.0.1

This commit is contained in:
duanmurs@163.com
2023-12-28 18:23:39 +08:00
parent fda66e7ca6
commit 66f979f7e9
204 changed files with 37941 additions and 0 deletions
+20
View File
@@ -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)
+83
View File
@@ -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()]
+32
View File
@@ -0,0 +1,32 @@
<scepter>
==========================================
.. 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
+35
View File
@@ -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
+53
View File
@@ -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**
### <font color="#0FB0E4">function **\_\_init\_\_()**</font>
#### 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
### <font color="#0FB0E4">function **\_\_getitem\_\_()**</font>
#### 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);
### <font color="#0FB0E4">function **worker_init_fn()**</font>
#### 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;
### <font color="#0FB0E4">function **\_get()**</font>
#### 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.
+181
View File
@@ -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;
* <font color="#0FB0E4">network</font>:train和test模块,对数据集输入的batch整合上述模块进行最终loss和指标计算;
<hr/>
## **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)
```
### <font color="#0FB0E4">function **\_\_init\_\_()**</font>
#### Parameters(input parameters)
(cfg: scepter.modules.utils.config.Config, logger = None) -> None
Mainly used for initializing various layers of the model;
### <font color="#0FB0E4">function **forward()**</font>
To be implemented specifically as needed;
<hr/>
## **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)
```
### <font color="#0FB0E4">function **\_\_init\_\_()**</font>
#### Parameters(input parameters)
(cfg: scepter.modules.utils.config.Config, logger = None) -> None
Initializes the hyperparameters needed for calculating metrics, such as coefficients like topk;
### <font color="#0FB0E4">function **\_\_call\_\_()**</font>
@torch.no_grad()
Typically takes logits and labels as well as other necessary variables as inputs and outputs calculated metrics;
<hr/>
## **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)
```
### <font color="#0FB0E4">function **\_\_init\_\_()**</font>
#### Parameters(input parameters)
(cfg: scepter.modules.utils.config.Config, logger = None) -> None
None Used for initializing and loading the tokenizer object, such as BertTokenizer;
### <font color="#0FB0E4">function **tokenize()**</font>
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;
<hr/>
## **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)
```
### <font color="#0FB0E4">function **\_\_init\_\_()**</font>
#### 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;
### <font color="#0FB0E4">function **forward_train()**</font>
Takes data from a training batch, processes it through the backbone, neck, head, and loss to calculate the relevant loss;
### <font color="#0FB0E4">function **forward_test()**</font>
Takes data from a test batch, processes it through the backbone, neck, head, and metrics to calculate the relevant metrics;
### <font color="#0FB0E4">function **forward**</font>
The actual calling interface, used to dispatch tasks to forward_train()/forward_test();
Other functions needed for training/testing can be customized under network;"
+83
View File
@@ -0,0 +1,83 @@
# Optimizer (Optimizer)
## Overview
1. lr_schedulers
2. optimizers
<hr/>
## 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;
### <font color="#0FB0E4">function **\_\_init\_\_()**</font>
#### 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
### <font color="#0FB0E4">function **\_\_call\_\_()**</font>
#### Parameters(input parameters)
(optimizer: scepter.modules.opt.optimizers.OPTIMIZERS) -> None
Sets up the schedule for the passed-in optimizer object;
<hr/>
## 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;
### <font color="#0FB0E4">function **\_\_init\_\_()**</font>
#### 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
### <font color="#0FB0E4">function **\_\_call\_\_()**</font>
#### Parameters(input parameters)
(parameters:dict()) -> None
Inputs the train parameters that need gradient updates, in dict format;
+282
View File
@@ -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"
```
+130
View File
@@ -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.
<hr/>
## LoopSampler
An infinite looping sampler for simple data.
#### <font color="#0FB0E4">function **LoopSampler.__init__**</font>
(cfg: scepter.modules.utils.config.Config, logger=None) -> None
#### <font color="#0FB0E4">function **LoopSampler.__iter__**</font>
Iterator, each iteration yields the index of a sample.
## MixtureOfSamplers
A sampler for large-scale data with multi-level indexing.
#### <font color="#0FB0E4">function **MixtureOfSamplers.__init__**</font>
(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.
#### <font color="#0FB0E4">function **MultiLevelBatchSampler.__iter__**</font>
()
Iterator, each iteration yields the index of a sample.
## MultiFoldDistributedSampler
A multi-fold sampler that supports repeating data several times within one epoch.
#### <font color="#0FB0E4">function **MultiFoldDistributedSampler.__init__**</font>
(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.
#### <font color="#0FB0E4">function **MultiFoldDistributedSampler.__iter__**</font>
()
Iterator, each iteration yields the index of a sample.
#### <font color="#0FB0E4">function **MultiFoldDistributedSampler.set_epoch**</font>
(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.
#### <font color="#0FB0E4">function **EvalDistributedSampler.__init__**</font>
(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.
#### <font color="#0FB0E4">function **EvalDistributedSampler.__iter__**</font>
()
Iterator, each iteration yields the index of a sample.
#### <font color="#0FB0E4">function **EvalDistributedSampler.set_epoch**</font>
(epoch: int)
Sets the current epoch.
**Parameters**
- **epoch** —— The current epoch number.
## MultiLevelBatchSampler
A sampler for large-scale data with multi-level indexing.
#### <font color="#0FB0E4">function **MultiLevelBatchSampler.__init__**</font>
(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.
#### <font color="#0FB0E4">function **MultiLevelBatchSampler.__iter__**</font>
()
Iterator, each iteration yields the index of a sample.
+488
View File
@@ -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.
<hr/>
## 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)
```
<hr/>
## **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.
<br>
### <font color="#0FB0E4">function **\_\_init\_\_**</font>
(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.
<br>
### <font color="#0FB0E4">function **set_up_pre**</font>
() -> 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.
<br>
### <font color="#0FB0E4">function **set_up**</font>
() -> None
Configure data, model, metrics, optimizer, and the pytorch_lightning environment (if used).
<br>
### <font color="#0FB0E4">function **construct_data**</font>
() -> 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.
<br>
### <font color="#0FB0E4">function **construct_hook**</font>
() -> 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.
<br>
### <font color="#0FB0E4">function **construct_model**</font>
() -> 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.
<br>
### <font color="#0FB0E4">function **model_to_device**</font>
() -> 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.
<br>
### <font color="#0FB0E4">function **model_to_device**</font>
() -> 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.
<br>
### <font color="#0FB0E4">function **init_opti**</font>
() -> 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.
<br>
### <font color="#0FB0E4">function **solve**</font>
(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.
<br>
### <font color="#0FB0E4">function **solve_train**</font>
() -> 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.
<br>
### <font color="#0FB0E4">function **solve_eval**</font>
() -> 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***.
<br>
### <font color="#0FB0E4">function **solve_test**</font>
() -> 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***.
<br>
### <font color="#0FB0E4">function **before_solve**</font>
() -> None
Before actually executing run_xxx, perform ***Hook*** logging.
<br>
### <font color="#0FB0E4">function **after_solve**</font>
() -> None
After executing run_xxx, perform ***Hook*** logging again.
<br>
### <font color="#0FB0E4">function **run_train**</font>
() -> 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***.
<br>
### <font color="#0FB0E4">function **run_eval**</font>
() -> 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***.
<br>
### <font color="#0FB0E4">function **run_test**</font>
() -> 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***.
<br>
### <font color="#0FB0E4">function **run_step_train**</font>
(batch_data, batch_idx = 0, step = None, rank = None) -> None
Perform model training for a single batch.
<br>
### <font color="#0FB0E4">function **run_step_eval**</font>
(batch_data, batch_idx = 0, step = None, rank = None) -> None
Perform model inference for a single batch.
<br>
### <font color="#0FB0E4">function **run_step_test**</font>
(batch_data, batch_idx = 0, step = None, rank = None) -> None
Perform model testing for a single batch.
<br>
### <font color="#0FB0E4">function **register_flops**</font>
(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.
<br>
### <font color="#0FB0E4">function **before_epoch**</font>
() -> None
Execute before the start of each epoch. Perform before running run_xxx.
<br>
### <font color="#0FB0E4">function **before_all_iter**</font>
() -> None
Execute before the start of each epoch. Perform before the loop that iteratively calls run_step_xxx.
<br>
### <font color="#0FB0E4">function **before_iter**</font>
() -> None
Execute before the start of each step. Perform before running run_step_xxx.
<br>
### <font color="#0FB0E4">function **after_epoch**</font>
() -> None
Execute after the start of each epoch. Perform after running run_xxx.
<br>
### <font color="#0FB0E4">function **after_all_iter**</font>
() -> None
Execute after the end of each epoch. Perform after the loop that iteratively calls run_step_xxx.
<br>
### <font color="#0FB0E4">function **after_iter**</font>
() -> None
Execute after the start of each step. Perform after running run_step_xxx.
<br>
### <font color="#0FB0E4">function **collect_log_vars**</font>
() -> OrderedDict
Obtain the variables you need to save in logs
**Returns**
- **ret** —— Return the required variables.
<br>
# 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.
<hr/>
## 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)
```
<hr/>
## **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:
### <font color="#0FB0E4">function **__init__**</font>
(cfg: *scepter.modules.utils.config.Config*, logger = None) -> None
initialize logger。
### <font color="#0FB0E4">function **before_solve**</font>
(solver) -> None
Execute before starting to solve.
### <font color="#0FB0E4">function **after_solve**</font>
(solver) -> None
Execute after starting to solve.
### <font color="#0FB0E4">function **before_epoch**</font>
(solver) -> None
Execute before each epoch.
### <font color="#0FB0E4">function **after_epoch**</font>
(solver) -> None
Execute after each epoch.
### <font color="#0FB0E4">function **before_all_iter**</font>
(solver) -> None
Execute before each iteration.
### <font color="#0FB0E4">function **after_all_iter**</font>
(solver) -> None
Execute after each iteration.
### <font color="#0FB0E4">function **before_iter**</font>
(solver) -> None
Execute before each step.
### <font color="#0FB0E4">function **after_iter**</font>
(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。
+54
View File
@@ -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
<hr/>
## 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")
```
<hr/>
### <font color="#0FB0E4">function **module_list**</font>
()
**Returns**
- **list** —— A list of module names.
### <font color="#0FB0E4">function **objects_by_module**</font>
(module_name: str)
**Parameters**
- **module_name** —— Module name
**Returns**
- **list** —— A list of module names.
### <font color="#0FB0E4">function **get_module_object_config**</font>
(module_name: str, object_name: str)
**Parameters**
- **module_name** —— Module name
- **object_name** —— Object name
**Returns**
- **str** —— Parameter template.
+419
View File
@@ -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***
<hr/>
## 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)
```
<hr/>
## **scepter.modules.transform.image**
Some pre-processing methods used for images.
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.ImageTransform</font>
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
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.RandomResizedCrop</font>
Randomly crop the image to a specified size.
**Parameters**
- **SIZE** —— (int) crop size
- **RATIO** —— (list) ratio
- **SCALE** —— (list) scale
- **INTERPOLATION** —— (str) interpolation
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.RandomHorizontalFlip</font>
Randomly horizontally flip the given image with a given probability(***P***).
**Parameters**
- **P** —— (float) probability
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.Normalize</font>
Normalize the image using ***mean*** and standard deviation(***std***).
**Parameters**
- **MEAN** —— (list) mean
- **STD** —— (list) std
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.ImageToTensor</font>
transform ***PIL.Image / numpy.ndarray / unit8*** to ***float32 tensor***.
**Parameters**
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.Resize</font>
Resize the given image to the given ***Size***.
**Parameters**
- **INTERPOLATION** —— (str) interpolation
- **SIZE** —— (int) resized size
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.CenterCrop</font>
Crop the given image from the center.
**Parameters**
- **SIZE** —— (int) crop size
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.FlexibleResize</font>
Resize the given image to the given ***Size***.
**Parameters**
- **INTERPOLATION** —— (str) interpolation
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.FlexibleCenterCrop</font>
Center crop the given image to the given ***Size***.
**Parameters**
<br>
## **scepter.modules.transform.io**
Some methods for reading images from local disk.
<br>
### <font color="#0FB0E4">scepter.modules.transform.io.LoadPILImageFromFile</font>
Read a local image file into ***PIL.Image*** format.
**Parameters**
- **RGB_ORDER** —— (str) "RGB" or "BGR"
<br>
### <font color="#0FB0E4">scepter.modules.transform.io.LoadCvImageFromFile</font>
Read a local image file into ***cv2*** format.
**Parameters**
- **RGB_ORDER** —— (str) "RGB" or "BGR"
<br>
### <font color="#0FB0E4">scepter.modules.transform.io.LoadImageFromFile</font>
Read a local image file into a specific format.
**Parameters**
- **RGB_ORDER** —— (str) "RGB" or "BGR"
- **BACKEND** —— (str) "pillow" or "cv2" or "torchvision"
<br>
### <font color="#0FB0E4">scepter.modules.transform.io.LoadImageFromFileList</font>
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
<br>
## **scepter.modules.transform.io_video**
Some methods for reading videos from local disk.
<br>
### <font color="#0FB0E4">scepter.modules.transform.io_video.DecodeVideoToTensor</font>
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
<br>
### <font color="#0FB0E4">scepter.modules.transform.io_video.LoadVideoFromFile</font>
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
<br>
## **scepter.modules.transform.tensor**
Some methods for processing tensors.
<br>
### <font color="#0FB0E4">scepter.modules.transform.tensor.ToTensor</font>
Convert input data from other formats into tensors.
**Parameters**
- **KEYS** —— (list) keys of input data
<br>
### <font color="#0FB0E4">scepter.modules.transform.tensor.Select</font>
Select some keys from the input data and output them.
**Parameters**
- **META_KEYS** —— (list) chosen keys of input data
<br>
### <font color="#0FB0E4">scepter.modules.transform.tensor.Rename</font>
Rename the keys of the input data.
**Parameters**
- **IN_KEYS** —— (list) input data keys
- **OUT_KEYS** —— (list) output data keys
<br>
## **scepter.modules.transform.augmention**
Some methods for enhancing the colors in images.
<br>
### <font color="#0FB0E4">scepter.modules.transform.augmention.ColorJitterGeneral</font>
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
<br>
## **scepter.modules.transform.video**
Some methods for processing videos.
<br>
### <font color="#0FB0E4">scepter.modules.transform.video.VideoTransform</font>
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
<br>
### <font color="#0FB0E4">scepter.modules.transform.video.RandomResizedCropVideo</font>
Perform random cropping of the video to a specified size.
**Parameters**
- **META_KEYS** —— (list) chosen keys of input data
<br>
### <font color="#0FB0E4">scepter.modules.transform.video.CenterCropVideo</font>
Renaming the keys of the input data.
**Parameters**
- **SIZE** —— (int) crop size
- **RATIO** —— (list) ratio
- **SCALE** —— (list) scale
- **INTERPOLATION** —— (str) interpolation
<br>
### <font color="#0FB0E4">scepter.modules.transform.video.RandomHorizontalFlipVideo</font>
Randomly flip the given video horizontally with a given probability.
**Parameters**
- **P** —— (float) probability
<br>
### <font color="#0FB0E4">scepter.modules.transform.video.NormalizeVideo</font>
Normalize the video using the ***mean*** and standard deviation(***std***).
**Parameters**
- **MEAN** —— (list) mean
- **STD** —— (list) std
<br>
### <font color="#0FB0E4">scepter.modules.transform.video.VideoToTensor</font>
transform ***PIL.Image / numpy.ndarray / unit8*** to ***float32 tensor***.
**Parameters**
<br>
### <font color="#0FB0E4">scepter.modules.transform.video.AutoResizedCropVideo</font>
Crop the given video from the center.
**Parameters**
- **SCALE** —— (list) scale
<br>
### <font color="#0FB0E4">scepter.modules.transform.video.ResizeVideo</font>
Resize the given video to the specified dimensions.
**Parameters**
- **SCALE** —— (list) scale
- **INTERPOLATION** —— (str) interpolation
<br>
## **scepter.modules.transform.transform_xl**
Some methods for image processing in SDXL to obtain the desired coordinates.
<br>
### <font color="#0FB0E4">scepter.modules.transform.transform_xl.FlexibleCropXL</font>
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
<br>
## **scepter.modules.transform.identity**
Some methods to process images and obtain the required coordinates in SDXL.
<br>
### <font color="#0FB0E4">scepter.modules.transform.identity.Identity</font>
Return the image itself.
**Parameters**
<br>
## **scepter.modules.transform.compose**
Combine various transform methods.
<br>
### <font color="#0FB0E4">scepter.modules.transform.compose.Compose</font>
Combine the various transform objects from ***scepter.transforms*** into a pipeline.
**Parameters**
- **TRANSFORMS** —— (list) transform config list
<br>
+526
View File
@@ -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
<hr/>
## 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)
```
<hr/>
## **scepter.modules.utils.file_system.FileSystem**
By building various File IO Handlers, it supports read and write operations for different types of files.
<br>
### <font color="#0FB0E4">function **\_\_init\_\_**</font>
()
**Parameters**
<br>
### <font color="#0FB0E4">function **init_fs_client**</font>
( 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
<br>
### <font color="#0FB0E4">function **get_fs_client**</font>
( 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.
<br>
### <font color="#0FB0E4">function **get_from**</font>
( 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.
<br>
### <font color="#0FB0E4">function **get_dir_to_local_dir**</font>
( 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* —— 本地文件夹路径
<br>
### <font color="#0FB0E4">function **get_object**</font>
(target_path: *str*) -> bytes
Read a remote file into memory
**Parameters**
- **target_path** — Target file path
**Returns**
- *bytes* — Binary data of the target file
<br>
### <font color="#0FB0E4">function **put_object**</font>
(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
<br>
### <font color="#0FB0E4">function **delete_object**</font>
(target_path: *str*) -> bool
Delete the target file
**Parameters**
- **target_path** — Target file path
**Returns**
- *bool* — Whether the deletion was successful
<br>
### <font color="#0FB0E4">function **get_batch_objects_from**</font>
(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
- <br>
### <font color="#0FB0E4">function **put_batch_objects_to**</font>
(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
<br>
### <font color="#0FB0E4">function **get_object_stream**</font>
(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
<br>
### <font color="#0FB0E4">function **get_object_chunk_list**</font>
(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
<br>
### <font color="#0FB0E4">function **get_url**</font>
(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
<br>
### <font color="#0FB0E4">function **put_to**</font>
(target_path: *str*)
Supports uploading a local file to a remote path
**Parameters**
- **target_path** — Remote file path
**Returns**
- **None**
<br>
```python
# Used as a context manager
with FS.put_to(target_path) as local_path:
# some operations on local_path.
```
<br>
### <font color="#0FB0E4">function **put_object_from_local_file**</font>
(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
<br>
### <font color="#0FB0E4">function **put_dir_from_local_dir**</font>
(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
<br>
### <font color="#0FB0E4">function **add_target_local_map**</font>
(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*
<br>
### <font color="#0FB0E4">function **make_dir**</font>
(target_dir: *str*) -> bool
Create a remote directory
**Parameters**
- **target_dir** — Remote directory path
**Returns**
- *bool* — Whether the creation was successful
<br>
### <font color="#0FB0E4">function **exists**</font>
(target_path: *str*) -> bool
Check if the target path exists
**Parameters**
- **target_path** — Remote file path
**Returns**
- *bool* — Whether it exists
<br>
### <font color="#0FB0E4">function **map_to_local**</font>
(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
<br>
### <font color="#0FB0E4">function **walk_dir**</font>
(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
<br>
### <font color="#0FB0E4">function **is_local_client**</font>
(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
<br>
### <font color="#0FB0E4">function **size**</font>
(target_path: *str*) -> int
Determine the size of the target file
**Parameters**
- **target_path** — Target file path
**Returns**
- *int* — Size of the target file
<br>
### <font color="#0FB0E4">function **isfile**</font>
(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
<br>
### <font color="#0FB0E4">function **isdir**</font>
(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
<hr/>
## **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
```
<hr/>
## **scepter.modules.utils.file_clients.LocalFs**
```yaml
NAME: LocalFs
TEMP_DIR: None
AUTO_CLEAN: False
```
<hr/>
## **scepter.modules.utils.file_clients.HttpFs**
```yaml
NAME: HttpFs
TEMP_DIR: None
AUTO_CLEAN: False
RETRY_TIMES: 10
```
<hr/>
+939
View File
@@ -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)
<hr/>
## 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)
```
<hr/>
### <font color="#0FB0E4">function **__init__**</font>
(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.
### <font color="#0FB0E4">function **dict_to_yaml**</font>
(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))
```
<hr/>
### <font color="#0FB0E4">function **osp_path**</font>
( 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
### <font color="#0FB0E4">function **get_relative_folder**</font>
( 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
### <font color="#0FB0E4">function **get_md5**</font>
( 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)
```
<hr/>
### <font color="#0FB0E4">class **Workenv**</font>
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.
### <font color="#0FB0E4">function **we.init_env**</font>
( 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.
### <font color="#0FB0E4">function **we.get_env**</font>
() -> dict
Retrieve all class-internal parameters of we, stored in the form of a dict.
### <font color="#0FB0E4">function **we.set_env**</font>
(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.
### <font color="#0FB0E4">function **get_dist_info**</font>
() -> 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.
### <font color="#0FB0E4">function **gather_data**</font>
(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.
### <font color="#0FB0E4">function **gather_list**</font>
(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.
### <font color="#0FB0E4">function **gather_picklable**</font>
(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.
### <font color="#0FB0E4">function **broadcast**</font>
(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.
### <font color="#0FB0E4">function **gather_gpu_tensors**</font>
(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
)
```
<hr/>
### <font color="#0FB0E4">function **save_develop_model_multi_io**</font>
(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")
```
<hr/>
### <font color="#0FB0E4">function **get_logger**</font>
(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.
### <font color="#0FB0E4">function **init_logger**</font>
(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
### <font color="#0FB0E4">function **as_time**</font>
(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.
### <font color="#0FB0E4">function **time_since**</font>
(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
)
```
<hr/>
### <font color="#0FB0E4">function **do_frame_sample**</font>
(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.
### <font color="#0FB0E4">class **VideoReaderWrapper**</font>
A standard class for reading videos, with the underlying decoder being decord.
#### <font color="#0FB0E4">function **VideoReaderWrapper.__init__**</font>
(video_path: str)
Initialize a video instance.
**Parameters**
- **video_path** —— Video link.
#### <font color="#0FB0E4">function **VideoReaderWrapper.len**</font>
() -> int
Get the total number of video frames.
**Returns**
- **int** —— Number of video frames.
#### <font color="#0FB0E4">function **VideoReaderWrapper.fps**</font>
() -> float
Get video frame rate.
**Returns**
- **float** —— Video frame rate.
#### <font color="#0FB0E4">function **VideoReaderWrapper.duration**</font>
() -> float
Get video duration.
**Returns**
- **float** —— Video duration.
#### <font color="#0FB0E4">function **VideoReaderWrapper.sample_frames**</font>
(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.
### <font color="#0FB0E4">class **FramesReaderWrapper**</font>
Reads frame data in order from a given fully decoded frame folder.
#### <font color="#0FB0E4">function **FramesReaderWrapper.__init__**</font>
(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.
#### <font color="#0FB0E4">function **FramesReaderWrapper.len**</font>
() -> int
Get the total number of video frames.
**Returns**
- **int** —— Number of video frames.
#### <font color="#0FB0E4">function **FramesReaderWrapper.fps**</font>
() -> float
Get video frame rate.
**Returns**
- **float** —— Video frame rate.
#### <font color="#0FB0E4">function **FramesReaderWrapper.duration**</font>
() -> float
Get video duration.
**Returns**
- **float** —— Video duration.
#### <font color="#0FB0E4">function **FramesReaderWrapper.sample_frames**</font>
(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.
### <font color="#0FB0E4">class **EasyVideoReader**</font>
Used for reading, sampling, and preprocessing long videos.
#### <font color="#0FB0E4">function **EasyVideoReader.__init__**</font>
(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.
#### <font color="#0FB0E4">function **EasyVideoReader.__iter__**</font>
() -> int
Iterator
#### <font color="#0FB0E4">function **EasyVideoReader.__next__**</font>
() -> 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)
```
<hr/>
### <font color="#0FB0E4">class **Registry**</font>
Registry
#### <font color="#0FB0E4">function **Registry.__init__**</font>
(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.
#### <font color="#0FB0E4">function **Registry.build**</font>
(cfg: Config, logger: logger = None, kwargs) -> cls_obj
Build an instance of the target class
**Returns**
- **cls_obj** —— An instance of a specific class.
#### <font color="#0FB0E4">function **Registry.register_class**</font>
(name: str)
Register a class
**Returns**
- **name** —— Registration name.
#### <font color="#0FB0E4">function **Registry.register_function**</font>
(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)
```
<hr/>
#### <font color="#0FB0E4">function **transfer_data_to_numpy**</font>
(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.
#### <font color="#0FB0E4">function **transfer_data_to_cpu**</font>
(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.
#### <font color="#0FB0E4">function **transfer_data_to_cuda**</font>
(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
```
<hr/>
#### <font color="#0FB0E4">function **move_model_to_cpu**</font>
(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.
#### <font color="#0FB0E4">function **load_pretrained**</font>
(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.
#### <font color="#0FB0E4">function **count_params**</font>
(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).
#### <font color="#0FB0E4">function **init_weights**</font>
(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
```
<hr/>
#### <font color="#0FB0E4">class **MultiFoldDistributedSampler**</font>
Multi-fold sampler, supports repeating multiple rounds of data within one epoch.
#### <font color="#0FB0E4">function **MultiFoldDistributedSampler.__init__**</font>
( 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.
#### <font color="#0FB0E4">function **MultiFoldDistributedSampler.__iter__**</font>
()
Iterator, each iteration returns an index of a sample.
#### <font color="#0FB0E4">function **MultiFoldDistributedSampler.set_epoch**</font>
(epoch: int)
Set the current epoch.
**Parameters**
- **epoch** —— The current epoch.
#### <font color="#0FB0E4">class **EvalDistributedSampler**</font>
A sampler for testing, when not using padding mode, it will be observed that the last rank has fewer data than other ranks.
#### <font color="#0FB0E4">function **EvalDistributedSampler.__init__**</font>
( 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.
#### <font color="#0FB0E4">function **EvalDistributedSampler.__iter__**</font>
()
Iterator, each iteration returns an index of a sample.
#### <font color="#0FB0E4">function **EvalDistributedSampler.set_epoch**</font>
(epoch: int)
Set the current epoch.
**Parameters**
- **epoch** —— The current epoch.
#### <font color="#0FB0E4">class **MultiLevelBatchSampler**</font>
A sampler for multi-level indexing of large-scale data.
#### <font color="#0FB0E4">function **MultiLevelBatchSampler.__init__**</font>
(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.
#### <font color="#0FB0E4">function **MultiLevelBatchSampler.__iter__**</font>
()
Iterator, each iteration returns an index of a sample.
#### <font color="#0FB0E4">class **MixtureOfSamplers**</font>
A sampler for multi-level indexing of large-scale data.
#### <font color="#0FB0E4">function **MixtureOfSamplers.__init__**</font>
(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.
#### <font color="#0FB0E4">function **MixtureOfSamplers.__iter__**</font>
()
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}"))
```
<hr/>
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
```
#### <font color="#0FB0E4">class **ProbeData**</font>
Instance of probe data.
#### <font color="#0FB0E4">function **ProbeData.__init__**</font>
(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.
+20
View File
@@ -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)
+83
View File
@@ -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()]
+32
View File
@@ -0,0 +1,32 @@
<scepter>
==========================================
.. 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
+35
View File
@@ -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
+65
View File
@@ -0,0 +1,65 @@
# 数据集模块 (Dataset)
## 总览
在继承BaseDataset基础上注册每个task各自的数据集读取模块,BaseDataset中封装有File System以及transform pipeline;
<hr/>
## **scepter.modules.data.dataset.BaseDataset**
### <font color="#0FB0E4">function **\_\_init\_\_()**</font>
#### 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
### <font color="#0FB0E4">function **\_\_getitem\_\_()**</font>
#### Parameters(输入参数):index
index的类别由sampler确定,详见data/registry.py;
一般默认的torch dataloader传入的index为int型作为dataset下标;
自定义的sampler则可以传入自定义item,sampler定义参照scepter.modules.utils.sampler;
* 具体读取单条数据由_get()方法传入index参数实现;
* 输出经由pipeline转换后的数据(如果有pipeline);
### <font color="#0FB0E4">function **worker_init_fn()**</font>
#### Parameters(输入参数):
(worker_id, num_workers = 1)
worker_id为分布式训练中工作节点id;
dataloader一次性创建num_workers个工作进程;
* 用于初始化文件读取系统和设置多卡worker对应参数;
### <font color="#0FB0E4">function **\_get()**</font>
#### Parameters(输入参数):index(由__getitem__()传入)
抽象方法,需要由自定义dataset继承具体实现,用于根据index读取batch中每条数据;
<hr/>
## 基础用法
子类注册:
```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模块下;
+184
View File
@@ -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:用于创建微调模块;
* <font color="#0FB0E4">network</font>:train和test模块,对数据集输入的batch整合上述模块进行最终loss和指标计算;
<hr/>
## **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)
```
### <font color="#0FB0E4">function **\_\_init\_\_()**</font>
#### Parameters(输入参数)
(cfg: scepter.modules.utils.config.Config, logger = None) -> None
主要用于初始化模型各layer;
### <font color="#0FB0E4">function **forward()**</font>
根据需要具体实现;
<hr/>
## **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)
```
### <font color="#0FB0E4">function **\_\_init\_\_()**</font>
#### Parameters(输入参数)
(cfg: scepter.modules.utils.config.Config, logger = None) -> None
初始化计算metric所需超参,例如topk等系数;
### <font color="#0FB0E4">function **\_\_call\_\_()**</font>
@torch.no_grad()
通常输入logit和label以及其他所需要的变量,输出计算指标;
<hr/>
## **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)
```
### <font color="#0FB0E4">function **\_\_init\_\_()**</font>
#### Parameters(输入参数)
(cfg: scepter.modules.utils.config.Config, logger = None) -> None
用于初始化加载tokenizer对象,例如BertTokenizer;
### <font color="#0FB0E4">function **tokenize()**</font>
输入需要分词的text list,输出分词后转换的token id sequence以及其他所需的attention mask/tpye id list/position id list等;
<hr/>
## **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)
```
### <font color="#0FB0E4">function **\_\_init\_\_()**</font>
#### Parameters(输入参数)
(cfg: scepter.modules.utils.config.Config, logger = None) -> None
结合上述build方法,初始化训练所需的backbone、neck、head、loss、metric、tokenizer模块;
### <font color="#0FB0E4">function **forward_train()**</font>
输入训练batch的数据,经backbone、neck、head、loss计算相关loss;
### <font color="#0FB0E4">function **forward_test()**</font>
输入测试batch的数据,经backbone、neck、head、metrics计算相关指标;
### <font color="#0FB0E4">function **forward**</font>
实际调用接口,用于分发任务至forward_train()/forward_test();
其他训练/测试所需函数可在network下自定义;
+83
View File
@@ -0,0 +1,83 @@
# 优化器 (Optimizer)
## 总览
1. lr_schedulers
2. optimizers
<hr/>
## 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的基类,支持注册操作,可根据需要自定义;
### <font color="#0FB0E4">function **\_\_init\_\_()**</font>
#### Parameters(输入参数)
(cfg: scepter.modules.utils.config.Config, logger = None) -> None
#### config(常用参数,实际需要根据不同schedulers设置,以StepLR为例):
* STEP_SIZE
* GAMMA
* LAST_EPOCH
### <font color="#0FB0E4">function **\_\_call\_\_()**</font>
#### Parameters(输入参数)
(optimizer: scepter.modules.opt.optimizers.OPTIMIZERS) -> None
具体对传入的optimizer对象进行schedule设置;
<hr/>
## 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的基类,支持注册操作,可根据需要自定义;
### <font color="#0FB0E4">function **\_\_init\_\_()**</font>
#### Parameters(输入参数)
(cfg: scepter.modules.utils.config.Config, logger = None) -> None
#### config(常用参数,实际需要根据不同optimizer设置,以SGD为例):
* LEARNING_RATE
* MOMENTUM
* DAMPENING
* WEIGHT_DECAY
* NESTEROV
### <font color="#0FB0E4">function **\_\_call\_\_()**</font>
#### Parameters(输入参数)
(parameters:dict()) -> None
输入需要梯度更新的train parameters,格式为dict;
+251
View File
@@ -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"
```
+165
View File
@@ -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中可以读取相关配置参数。
<hr/>
## LoopSampler
简单数据的无限循环sampler。
#### <font color="#0FB0E4">function **LoopSampler.__init__**</font>
(cfg: scepter.modules.utils.config.Config, logger = None) -> None
#### <font color="#0FB0E4">function **LoopSampler.__iter__**</font>
迭代器,每迭代一次得到一个样本的index
## MixtureOfSamplers
用于大规模数据的多级索引的sampler
#### <font color="#0FB0E4">function **MixtureOfSamplers.__init__**</font>
(samplers: list(sampler), probabilities: list(float), rank: int =0, seed: int = 8888)
**Parameters**
- **samplers** —— 采样器列表,用于混合采样器。
- **probabilities** —— 每个采样器的概率。
- **rank** —— rank表示当前进程号。
- **seed** —— 随机采样的seed,在data.registry中获取全局seed。
#### <font color="#0FB0E4">function **MultiLevelBatchSampler.__iter__**</font>
()
迭代器,每迭代一次得到一个样本的index
## MultiFoldDistributedSampler
多fold采样器,支持在一个epoch中重复多轮数据
#### <font color="#0FB0E4">function **MultiFoldDistributedSampler.__init__**</font>
( 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** —— 数据是否要打乱。
#### <font color="#0FB0E4">function **MultiFoldDistributedSampler.__iter__**</font>
()
迭代器,每迭代一次得到一个样本的index
#### <font color="#0FB0E4">function **MultiFoldDistributedSampler.set_epoch**</font>
(epoch: int)
设置当前的epoch
**Parameters**
- **epoch** —— 当前的epoch。
## EvalDistributedSampler
用于测试时的采样器,当不用padding模式的时候,会发现最后一个rank的数据会少于其他rank。
#### <font color="#0FB0E4">function **EvalDistributedSampler.__init__**</font>
( 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数据量一致。
#### <font color="#0FB0E4">function **EvalDistributedSampler.__iter__**</font>
()
迭代器,每迭代一次得到一个样本的index
#### <font color="#0FB0E4">function **EvalDistributedSampler.set_epoch**</font>
(epoch: int)
设置当前的epoch
**Parameters**
- **epoch** —— 当前的epoch。
## MultiLevelBatchSampler
用于大规模数据的多级索引的sampler
#### <font color="#0FB0E4">function **MultiLevelBatchSampler.__init__**</font>
(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。
#### <font color="#0FB0E4">function **MultiLevelBatchSampler.__iter__**</font>
()
迭代器,每迭代一次得到一个样本的index
+489
View File
@@ -0,0 +1,489 @@
# 训练器(Solvers)
## 总览
Solver是对模型训练、验证和测试过程的一个流程定义。
在Solver中,会根据配置yaml文件的设置,对数据(data),模型(model),优化器(optimizer)和调度(scheduler)等需要的模块进行逐一初始化。
每一个具体的任务的自定义Solver都要继承自BaseSolver。
在某些特殊场景下,还需要初始化一些自定义的模块。例如,在训练过程中记录保存中间结果需要用到Hooks;在训练过程中验证需要定义Metrics等等。
<hr/>
## 基础用法
子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)
```
<hr/>
## **scepter.modules.solver.BaseSolver**
Solver的基类,是通过元类ABCMeta定义的抽象基类,支持注册操作。自定义的solver均应该是该类的子类,并且进行注册。
BaseSolver是一个具体实现Solver的案例,展示了使用pytorch_lightning和不使用时两种Solver的写法。
在实际使用中情况各异,因此Solver的使用也比较灵活,其中大部分的成员函数均可以按照需求在子类中,被复写或者新加好用的功能,甚至直接写新的函数代替。
<br>
### <font color="#0FB0E4">function **\_\_init\_\_**</font>
(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.
<br>
### <font color="#0FB0E4">function **set_up_pre**</font>
() -> None
配置环境、日志路径,调用construct_hook来初始化hook等等。
与pytorch_lightning二选一,在不使用pytorch_lightning(use_pl=Flase)时调用。需要在启动其他所有操作之前调用。
<br>
### <font color="#0FB0E4">function **set_up**</font>
() -> None
配置数据、模型、metrics、优化器及pytorch_lightning环境(如果使用的话)。
<br>
### <font color="#0FB0E4">function **construct_data**</font>
() -> None
实际数据的构建方法,默认会在self.set_up中被调用,包括TRAIN_DATA、EVAL_DATA、TEST_DATA。将实例化的结果写入self.datas中。
<br>
### <font color="#0FB0E4">function **construct_hook**</font>
() -> None
实际Hook的构建方法,默认会在self.set_up_pre中被调用,包括TRAIN_HOOKS、EVAL_HOOKS、TEST_HOOKS。将实例化的结果写入self.hooks_dict中。
<br>
### <font color="#0FB0E4">function **construct_model**</font>
() -> None
实际Hook的构建方法,默认会在self.set_up中被调用,将实例化的结果作为self.model。
<br>
### <font color="#0FB0E4">function **model_to_device**</font>
() -> None
实际Metrics的构建方法,默认会在self.set_up中被调用,将实例化的结果写入self.metrics。
<br>
### <font color="#0FB0E4">function **model_to_device**</font>
(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,则无返回值。
<br>
### <font color="#0FB0E4">function **init_opti**</font>
() -> None
优化器的配置方法,默认会在self.set_up中被调用。将实例化的optimizer和lr_scheduler分别作为self.optimizer和self.lr_scheduler。
<br>
### <font color="#0FB0E4">function **solve**</font>
(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。
<br>
### <font color="#0FB0E4">function **solve_train**</font>
() -> None
调用self.run_train,执行train。
并且在执行前后分别调用self.before_epoch、self.after_epoch来进行Hook的记录。
<br>
### <font color="#0FB0E4">function **solve_eval**</font>
() -> None
调用self.run_eval,执行eval。
并且在执行前后分别调用self.before_epoch、self.after_epoch来进行Hook的记录。
<br>
### <font color="#0FB0E4">function **solve_test**</font>
() -> None
调用self.run_test,执行test。
并且在执行前后分别调用self.before_epoch、self.after_epoch来进行Hook的记录。
<br>
### <font color="#0FB0E4">function **before_solve**</font>
() -> None
在实际执行run_xxx之前,进行Hook的记录。
<br>
### <font color="#0FB0E4">function **after_solve**</font>
() -> None
在执行run_xxx之后,再次进行Hook的记录。
<br>
### <font color="#0FB0E4">function **run_train**</font>
() -> 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的记录。
<br>
### <font color="#0FB0E4">function **run_eval**</font>
() -> 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的记录。
<br>
### <font color="#0FB0E4">function **run_test**</font>
() -> 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的记录。
<br>
### <font color="#0FB0E4">function **run_step_train**</font>
(batch_data, batch_idx = 0, step = None, rank = None) -> None
执行一个batch的模型推理。
<br>
### <font color="#0FB0E4">function **run_step_eval**</font>
(batch_data, batch_idx = 0, step = None, rank = None) -> None
执行一个batch的模型推理。
<br>
### <font color="#0FB0E4">function **run_step_test**</font>
(batch_data, batch_idx = 0, step = None, rank = None) -> None
执行一个batch的模型推理。
<br>
### <font color="#0FB0E4">function **register_flops**</font>
(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。
<br>
### <font color="#0FB0E4">function **before_epoch**</font>
() -> None
在每个epoch开始前执行。run_xxx前执行。
<br>
### <font color="#0FB0E4">function **before_all_iter**</font>
() -> None
在每个epoch开始前执行。循环调用run_step_xxx的循环前执行。
<br>
### <font color="#0FB0E4">function **before_iter**</font>
() -> None
在每个step开始前执行。run_step_xxx前执行。
<br>
### <font color="#0FB0E4">function **after_epoch**</font>
() -> None
在每个epoch开始后执行。run_xxx后执行。
<br>
### <font color="#0FB0E4">function **after_all_iter**</font>
() -> None
在每个epoch开始后执行。循环调用run_step_xxx的循环后执行。
<br>
### <font color="#0FB0E4">function **after_iter**</font>
() -> None
在每个step开始后执行。run_step_xxx后执行。
<br>
### <font color="#0FB0E4">function **collect_log_vars**</font>
() -> OrderedDict
获取需要在log中保存的变量。
**Returns**
- **ret** —— 返回需要的变量。
<br>
# Hooks
## 总览
在Solver执行过程中,需要打印日志、记录Tensorboard、梯度计算和更新、保存中间模型参数、保存测试结果等,这些都需要Hook去执行。
<hr/>
## 基础用法
新建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)
```
<hr/>
## **scepter.modules.solvers.hooks.Hook**
定义了一个标准的Hook类的基类,定义了详细的成员函数列表。继承自该类的子类均需要在这些成员函数中挑选部分进行实现。
成员函数包括:
### <font color="#0FB0E4">function **__init__**</font>
(cfg: *scepter.modules.utils.config.Config*, logger = None) -> None
初始化logger。
### <font color="#0FB0E4">function **before_solve**</font>
(solver) -> None
开始solve之前执行。
### <font color="#0FB0E4">function **after_solve**</font>
(solver) -> None
结束solve之后执行。
### <font color="#0FB0E4">function **before_epoch**</font>
(solver) -> None
每个epoch之前执行。
### <font color="#0FB0E4">function **after_epoch**</font>
(solver) -> None
每个epoch之后执行。
### <font color="#0FB0E4">function **before_all_iter**</font>
(solver) -> None
迭代开始之前执行。
### <font color="#0FB0E4">function **after_all_iter**</font>
(solver) -> None
迭代结束之后执行。
### <font color="#0FB0E4">function **before_iter**</font>
(solver) -> None
每个step之前执行。
### <font color="#0FB0E4">function **after_iter**</font>
(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)保存文件的前缀名。
+54
View File
@@ -0,0 +1,54 @@
# 工具组件(Tools)
支持对框架注册的组件进行查询、获取参数模版。
## 总览
1、查询模块类型:scepter.module_list
2、按模块查询对象:scepter.objects_by_module
3、按对象查询参数配置:scepter.configures_by_objects
<hr/>
## 基础用法
```python
from scepter import module_list, objects_by_module, configures_by_objects
# 查询模块类型
module_list()
# 按照模块查询对象列表
objects_by_module("BACKBONES")
# 按照对象查询参数配置
configures_by_objects("BACKBONES", "ResNet3D_TAda")
```
<hr/>
### <font color="#0FB0E4">function **module_list**</font>
()
**Returns**
- **list** —— 模块名列表。
### <font color="#0FB0E4">function **objects_by_module**</font>
(module_name: str)
**Parameters**
- **module_name** —— 模块名
**Returns**
- **list** —— 模块名列表。
### <font color="#0FB0E4">function **get_module_object_config**</font>
(module_name: str, object_name: str)
**Parameters**
- **module_name** —— 模块名
- **object_name** —— 对象名
**Returns**
- **str** —— 参数模版。
+419
View File
@@ -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
<hr/>
## 基础用法
```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)
```
<hr/>
## **scepter.modules.transform.image**
一些用于图像的预处理方法.
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.ImageTransform</font>
初始化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
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.RandomResizedCrop</font>
对图像进行随机crop到指定大小.
**Parameters**
- **SIZE** —— (int) crop size
- **RATIO** —— (list) ratio
- **SCALE** —— (list) scale
- **INTERPOLATION** —— (str) interpolation
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.RandomHorizontalFlip</font>
以给定的概率随机水平翻转给定的图像.
**Parameters**
- **P** —— (float) probability
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.Normalize</font>
利用平均值和标准差对图像进行归一化.
**Parameters**
- **MEAN** —— (list) mean
- **STD** —— (list) std
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.ImageToTensor</font>
把PIL.Image / numpy.ndarray / unit8 转成float32 tensor.
**Parameters**
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.Resize</font>
把给定的图像按照给定的尺寸进行resize.
**Parameters**
- **INTERPOLATION** —— (str) interpolation
- **SIZE** —— (int) resized size
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.CenterCrop</font>
对给定的图像从中心进行crop.
**Parameters**
- **SIZE** —— (int) crop size
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.FlexibleResize</font>
对给定的图像按照给定的尺寸进行resize.
**Parameters**
- **INTERPOLATION** —— (str) interpolation
<br>
### <font color="#0FB0E4">scepter.modules.transform.image.FlexibleCenterCrop</font>
对给定的图像按照给定的尺寸进行center crop.
**Parameters**
<br>
## **scepter.modules.transform.io**
一些用于图像的本地磁盘读取方法.
<br>
### <font color="#0FB0E4">scepter.modules.transform.io.LoadPILImageFromFile</font>
将本地图片文件读取成PIL.Image的形式.
**Parameters**
- **RGB_ORDER** —— (str) "RGB" or "BGR"
<br>
### <font color="#0FB0E4">scepter.modules.transform.io.LoadCvImageFromFile</font>
将本地图片文件读取成cv2的形式.
**Parameters**
- **RGB_ORDER** —— (str) "RGB" or "BGR"
<br>
### <font color="#0FB0E4">scepter.modules.transform.io.LoadImageFromFile</font>
将本地图片文件读取成指定格式.
**Parameters**
- **RGB_ORDER** —— (str) "RGB" or "BGR"
- **BACKEND** —— (str) "pillow" or "cv2" or "torchvision"
<br>
### <font color="#0FB0E4">scepter.modules.transform.io.LoadImageFromFileList</font>
将输入的一组图片读取成指定格式.
**Parameters**
- **RGB_ORDER** —— (str) "RGB" or "BGR"
- **BACKEND** —— (str) "pillow" or "cv2" or "torchvision"
- **FILE_KEYS** —— (list) The file keys for input
<br>
## **scepter.modules.transform.io_video**
一些用于视频的本地磁盘读取方法.
<br>
### <font color="#0FB0E4">scepter.modules.transform.io_video.DecodeVideoToTensor</font>
将本地视频文件解码成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
<br>
### <font color="#0FB0E4">scepter.modules.transform.io_video.LoadVideoFromFile</font>
将本地视频文件解码读取成帧序列的形式.
**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
<br>
## **scepter.modules.transform.tensor**
一些处理tensor的方法.
<br>
### <font color="#0FB0E4">scepter.modules.transform.tensor.ToTensor</font>
将输入的其他形式的data转成tensor.
**Parameters**
- **KEYS** —— (list) keys of input data
<br>
### <font color="#0FB0E4">scepter.modules.transform.tensor.Select</font>
选择输入data中的一些key并输出.
**Parameters**
- **META_KEYS** —— (list) chosen keys of input data
<br>
### <font color="#0FB0E4">scepter.modules.transform.tensor.Rename</font>
将输入data的keys重新命名.
**Parameters**
- **IN_KEYS** —— (list) input data keys
- **OUT_KEYS** —— (list) output data keys
<br>
## **scepter.modules.transform.augmention**
一些图片颜色增强的方法.
<br>
### <font color="#0FB0E4">scepter.modules.transform.augmention.ColorJitterGeneral</font>
随机改变图像的亮度、对比度和饱和度.
**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
<br>
## **scepter.modules.transform.video**
一些处理video的方法.
<br>
### <font color="#0FB0E4">scepter.modules.transform.video.VideoTransform</font>
初始化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
<br>
### <font color="#0FB0E4">scepter.modules.transform.video.RandomResizedCropVideo</font>
对视频进行随机crop到指定大小.
**Parameters**
- **META_KEYS** —— (list) chosen keys of input data
<br>
### <font color="#0FB0E4">scepter.modules.transform.video.CenterCropVideo</font>
将输入data的keys重新命名.
**Parameters**
- **SIZE** —— (int) crop size
- **RATIO** —— (list) ratio
- **SCALE** —— (list) scale
- **INTERPOLATION** —— (str) interpolation
<br>
### <font color="#0FB0E4">scepter.modules.transform.video.RandomHorizontalFlipVideo</font>
以给定的概率随机水平翻转给定的视频.
**Parameters**
- **P** —— (float) probability
<br>
### <font color="#0FB0E4">scepter.modules.transform.video.NormalizeVideo</font>
利用平均值和标准差对视频进行归一化.
**Parameters**
- **MEAN** —— (list) mean
- **STD** —— (list) std
<br>
### <font color="#0FB0E4">scepter.modules.transform.video.VideoToTensor</font>
把PIL.Image / numpy.ndarray / unit8 转成float32 tensor.
**Parameters**
<br>
### <font color="#0FB0E4">scepter.modules.transform.video.AutoResizedCropVideo</font>
对给定的视频从中心进行crop.
**Parameters**
- **SCALE** —— (list) scale
<br>
### <font color="#0FB0E4">scepter.modules.transform.video.ResizeVideo</font>
把给定的视频按照给定的尺寸进行resize.
**Parameters**
- **SCALE** —— (list) scale
- **INTERPOLATION** —— (str) interpolation
<br>
## **scepter.modules.transform.transform_xl**
sdxl中进行图像处理得到所需坐标的一些方法.
<br>
### <font color="#0FB0E4">scepter.modules.transform.transform_xl.FlexibleCropXL</font>
对图像进行裁剪,并获取其原始尺寸、目标尺寸和裁剪坐标(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
<br>
## **scepter.modules.transform.identity**
sdxl中进行图像处理得到所需坐标的一些方法.
<br>
### <font color="#0FB0E4">scepter.modules.transform.identity.Identity</font>
返回图像本身.
**Parameters**
<br>
## **scepter.modules.transform.compose**
组合各类transform方法.
<br>
### <font color="#0FB0E4">scepter.modules.transform.compose.Compose</font>
将scepter.transforms中的各个transform对象组合为pipeline.
**Parameters**
- **TRANSFORMS** —— (list) transform config list
<br>
+547
View File
@@ -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
<hr/>
## 基础用法
```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)
```
<hr/>
## **scepter.modules.utils.file_system.FileSystem**
通过build多类File IO Handler,支持不同类型文件的读写操作。
<br>
### <font color="#0FB0E4">function **\_\_init\_\_**</font>
()
**Parameters**
<br>
### <font color="#0FB0E4">function **init_fs_client**</font>
( 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
<br>
### <font color="#0FB0E4">function **get_fs_client**</font>
( target_path: *str*, safe: *bool* = False )
通过target_path的前缀来获取对应的fs_client
**Parameters**
- **target_path** —— 目标文件路径
- **safe** —— 安全模式,返回client的copy,否则返回client本身
**Returns**
- *BaseFs* —— 实例化的fs_client
<br>
### <font color="#0FB0E4">function **get_from**</font>
( 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* —— 本地保存文件路径
<br>
### <font color="#0FB0E4">function **get_dir_to_local_dir**</font>
( 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* —— 本地文件夹路径
<br>
### <font color="#0FB0E4">function **get_object**</font>
( target_path: *str* ) -> byte
读取远程文件到内存中
**Parameters**
- **target_path** —— 目标文件路径
**Returns**
- *bytes* —— 目标文件的二进制数据
<br>
### <font color="#0FB0E4">function **get_object**</font>
( target_path: *str* ) -> byte
读取远程文件到内存中
**Parameters**
- **target_path** —— 目标文件路径
**Returns**
- *bytes* —— 目标文件的二进制数据
<br>
### <font color="#0FB0E4">function **put_object**</font>
( local_data: *byte*, target_path: *str* ) -> bool
上传数据流到指定文件
**Parameters**
- **local_data** —— 本地数据流
- **target_path** —— 目标文件路径
**Returns**
- *bool* —— 是否上传成功
<br>
### <font color="#0FB0E4">function **delete_object**</font>
( target_path: *str* ) -> bool
删除目标文件
**Parameters**
- **target_path** ——
**Returns**
- *bool* —— 是否删除成功
<br>
### <font color="#0FB0E4">function **get_batch_objects_from**</font>
( target_path_list: *str*, wait_finish: *bool* ) -> *str*
批量下载文件
**Parameters**
- **target_path_list** —— 下载文件列表
**Returns**
- *local_path* —— 本地文件generator
<br>
### <font color="#0FB0E4">function **put_batch_objects_to**</font>
(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。
<br>
### <font color="#0FB0E4">function **get_object_stream**</font>
(target_path: *str*, start: *int*, size: *int*, delimiter: *str*) -> *byte*, *int*
批量上传文件
**Parameters**
- **target_path** —— 目标文件
- **start** —— 目标流的开始字符位置
- **size** —— 目标流的开始字符流大小
- **delimiter** —— 目标流的结束字符
**Returns**
- *local_data, end* —— 返回数据流字节和结束字符位置。
<br>
### <font color="#0FB0E4">function **get_object_chunk_list**</font>
( target_path: *str*, chunk_num: *int* = 1, delimiter: *str* = None ) -> list[bytes]
获取远程文件且分块
**Parameters**
- **target_path** —— 目标文件路径
- **chunk_num** —— 分块个数
- **delimiter** —— 分隔符,确保下载数据是完整的一条,不会从中间截断
**Returns**
- *list[bytes]* —— 分块数据
<br>
### <font color="#0FB0E4">function **get_url**</font>
( target_path: *str*, set_public = False, lifecycle: *int* = 360000 ) -> str
获取远程文件的url(仅支持AliyunOssFs)
**Parameters**
- **target_path** —— 目标文件路径
- **lifecycle** —— 有效时间
- **set_public** —— 反馈公开链接
**Returns**
- *str* —— 目标文件的url
<br>
### <font color="#0FB0E4">function **put_to**</font>
( target_path: *str* )
支持将本地文件上传到远程
**Parameters**
- **target_path** —— 远程文件路径
**Returns**
- **None**
```python
# 作为上下文管理器使用
with FS.put_to(target_path) as local_path:
# some operations on local_path.
```
<br>
### <font color="#0FB0E4">function **put_object_from_local_file**</font>
( local_path: *str*, target_path: *str* ) -> bool
将本地文件push到远程路径
**Parameters**
- **local_path** —— 本地文件路径
- **target_path** —— 远程文件路径
**Returns**
- *bool* —— 是否上传成功
<br>
### <font color="#0FB0E4">function **put_dir_from_local_dir**</font>
( local_dir: *str*, target_dir: *str* ) -> bool
将本地文件夹push到远程路径
**Parameters**
- **local_dir** —— 本地文件夹路径
- **target_dir** —— 远程文件夹路径
**Returns**
- *bool* —— 是否上传成功
<br>
### <font color="#0FB0E4">function **add_target_local_map**</font>
( target_dir: *str*, local_dir: *str* ) -> None
将远程文件夹和本地文件夹路径的映射关系以key-value对的形式保存到self._target_local_mapper
**Parameters**
- **target_dir** —— 远程文件夹路径
- **local_dir** —— 本地文件夹路径
**Returns**
- *None*
<br>
### <font color="#0FB0E4">function **make_dir**</font>
( target_dir: *str* ) -> bool
创建远程文件夹
**Parameters**
- **target_dir** —— 远程文件夹路径
**Returns**
- *bool* —— 是否创建成功
<br>
### <font color="#0FB0E4">function **exists**</font>
( target_path: *str* ) -> bool
判断目标路径是否存在
**Parameters**
- **target_path** —— 远程文件路径
**Returns**
- *bool* —— 是否存在
<br>
### <font color="#0FB0E4">function **map_to_local**</font>
( target_path: *str* ) -> str, bool
将远程文件路径映射到本地路径
**Parameters**
- **target_path** —— 远程文件路径
**Returns**
- *str* —— 本地文件路径
- *bool* —— 本地文件是否tmp文件
<br>
### <font color="#0FB0E4">function **walk_dir**</font>
( target_dir: *str*, recurse = True) -> Iterator
获取远程文件夹下的文件列表
**Parameters**
- **target_dir** —— 远程文件夹路径
- **recurse** —— 是否遍历子文件夹,默认遍历为True
**Returns**
- *Iterator* —— 子文件路径列表
<br>
### <font color="#0FB0E4">function **is_local_client**</font>
( target_path: *str* ) -> bool
判断目标文件client是不是LocalFs
**Parameters**
- **target_path** —— 目标文件路径
**Returns**
- *bool* —— client是否LocalFs
<br>
### <font color="#0FB0E4">function **size**</font>
( target_path: *str* ) -> int
判断目标文件大小
**Parameters**
- **target_path** —— 目标文件路径
**Returns**
- *int* —— 目标文件size
<br>
### <font color="#0FB0E4">function **isfile**</font>
( target_path: *str* ) -> bool
判断目标路径是不是object
**Parameters**
- **target_path** —— 目标文件路径
**Returns**
- *bool* —— 目标路径是不是object
<br>
### <font color="#0FB0E4">function **isdir**</font>
( target_path: *str* ) -> bool
判断目标路径是不是文件夹
**Parameters**
- **target_path** —— 目标文件路径
**Returns**
- *bool* —— 目标路径是不是文件夹
<hr/>
## **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
```
<hr/>
## **scepter.modules.utils.file_clients.LocalFs**
- target_path 格式: 本地路径
```yaml
NAME: LocalFs
TEMP_DIR: None
AUTO_CLEAN: False
```
<hr/>
## **scepter.modules.utils.file_clients.HttpFs**
- target_path 格式: `http://xx/yy/zz`
```yaml
NAME: HttpFs
TEMP_DIR: None
AUTO_CLEAN: False
RETRY_TIMES: 10
```
<hr/>
## **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
```
<hr/>
+941
View File
@@ -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)
<hr/>
## 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)
```
<hr/>
### <font color="#0FB0E4">function **__init__**</font>
( 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时下载。
### <font color="#0FB0E4">function **dict_to_yaml**</font>
( 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))
```
<hr/>
### <font color="#0FB0E4">function **osp_path**</font>
( prefix: str, data_file: str ) -> str
根据路径前缀进行自动化路径拼接
**Parameters**
- **prefix** —— 路径前缀。
- **data_file** —— 文件路径。
**Returns**
- **str** —— 拼接以后的路径
### <font color="#0FB0E4">function **get_relative_folder**</font>
( abs_path: str, keep_index: int = -1 ) -> str
根据路径获取指定层级的文件夹路径
**Parameters**
- **abs_path** —— 文件路径。
- **keep_index** —— 保留层级,-1代表倒数第一级,-2 为倒数第二级。
**Returns**
- **str** —— 解析以后的路径
### <font color="#0FB0E4">function **get_md5**</font>
( 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)
```
<hr/>
### <font color="#0FB0E4">class **Workenv**</font>
这是一个用于统一管理运行环境的类,通常不需要使用该类做初始化,在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。
### <font color="#0FB0E4">function **we.init_env**</font>
( config: scepter.modules.utils.config.Config, fn: function, logger: logging.Logger = None )
作为启动任何任务的执行入口。
**Parameters**
- **config** —— 传入的参数实例。
- **fn** —— 需要执行的函数。
- **logger** —— 标准的日志实例。
### <font color="#0FB0E4">function **we.get_env**</font>
() -> dict
获取we的所有类内参数,以dict的形式存储。
### <font color="#0FB0E4">function **we.set_env**</font>
(we_env: dict)
重新设置we的所有类内参数,以dict的形式作为输入。
**Parameters**
- **we_env** —— dict,每个key代表一个类内变量。
### <font color="#0FB0E4">function **get_dist_info**</font>
() -> int, int
获取环境的rank/world size,这个是直接通过torch的方法来获取的,一般用于当初始化环境的方式
不是we.init_env的时候使用。
**Returns**
- **rank** —— 当前进程的rank值,默认为0
- **world_size** —— 当前环境的总进程数, 当单进程时为1。
### <font color="#0FB0E4">function **gather_data**</font>
(data: [list, dict, tensor, object] ) -> data
通过scepter.distributed.all_gather将任意实例收集起来,并在rank=0进程合并为一个汇总后的实例。
**Parameters**
- **data** —— 支持dict/list,其中元素支持任意实例或者tensor。
**Returns**
- **data** —— 一个和输入data相同结构的汇总过的数据。
### <font color="#0FB0E4">function **gather_list**</font>
(data: [list] ) -> data
通过scepter.distributed.all_gather将任意实例收集起来,并在rank=0进程合并为一个汇总后的实例。
**Parameters**
- **data** —— 支持list,其中元素支持任意实例或者tensor。
**Returns**
- **data** —— 一个和输入data相同结构的汇总过的数据。
### <font color="#0FB0E4">function **gather_picklable**</font>
(data: [object] ) -> data
通过scepter.distributed.all_gather将任意实例收集起来,并在rank=0进程合并为一个汇总后的实例。
**Parameters**
- **data** —— 为一个可序列化的实例。
**Returns**
- **data** —— 一个和输入data相同结构的汇总过的数据。
### <font color="#0FB0E4">function **broadcast**</font>
(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也做了此操作。
### <font color="#0FB0E4">function **gather_gpu_tensors**</font>
(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
)
```
<hr/>
### <font color="#0FB0E4">function **save_develop_model_multi_io**</font>
(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")
```
<hr/>
### <font color="#0FB0E4">function **get_logger**</font>
(name: str) -> logger
获取日志实例。
**Parameters**
- **name** —— 日志前缀,每次打印会首先打印该前缀。
**Returns**
- **logger** —— 返回一个logging实例。
### <font color="#0FB0E4">function **init_logger**</font>
(in_logger: logger, log_file: str) -> logger
二次初始化日志实例,可以为该实例分配一个文件落盘。
**Parameters**
- **in_logger** —— 已有的日志实例。
- **log_file** —— 希望存储的文件位置。
- **dist_launcher** —— 已经不重要了,deprecated
### <font color="#0FB0E4">function **as_time**</font>
(s: int) -> str
时间s转换为标准的xxx days xxx hours xxx mins xxx secs
**Parameters**
- **s** —— 代表秒数s。
**Returns**
- **str** —— 格式化的输出。
### <font color="#0FB0E4">function **time_since**</font>
(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
)
```
<hr/>
### <font color="#0FB0E4">function **do_frame_sample**</font>
(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** —— 采样帧结果。
### <font color="#0FB0E4">class **VideoReaderWrapper**</font>
读取视频的标准类,底层解码器为decord
#### <font color="#0FB0E4">function **VideoReaderWrapper.__init__**</font>
(video_path: str)
初始化视频实例
**Parameters**
- **video_path** —— 视频链接。
#### <font color="#0FB0E4">function **VideoReaderWrapper.len**</font>
() -> int
获取视频帧总数。
**Returns**
- **int** —— 视频帧数。
#### <font color="#0FB0E4">function **VideoReaderWrapper.fps**</font>
() -> float
获取视频帧率
**Returns**
- **float** —— 视频帧率。
#### <font color="#0FB0E4">function **VideoReaderWrapper.duration**</font>
() -> float
获取视频时长
**Returns**
- **float** —— 视频时长。
#### <font color="#0FB0E4">function **VideoReaderWrapper.sample_frames**</font>
(decode_list: torch.Tensor) -> torch.Tensor
根据帧号,获取帧数据
**Parameters**
- **decode_list** —— 采样帧号列表。
**Returns**
- **tensor** —— 数据张量。
### <font color="#0FB0E4">class **FramesReaderWrapper**</font>
给定解完帧的文件夹,按顺序读取帧数据
#### <font color="#0FB0E4">function **FramesReaderWrapper.__init__**</font>
(frame_dir: str, extract_fps: float, suffix: str)
初始化视频实例
**Parameters**
- **frame_dir** —— 帧文件夹。
- **extract_fps** —— 提取帧的fps。
- **suffix** —— 帧文件的后缀,默认为jpg。
#### <font color="#0FB0E4">function **FramesReaderWrapper.len**</font>
() -> int
获取视频帧总数。
**Returns**
- **int** —— 视频帧数。
#### <font color="#0FB0E4">function **FramesReaderWrapper.fps**</font>
() -> float
获取视频帧率
**Returns**
- **float** —— 视频帧率。
#### <font color="#0FB0E4">function **FramesReaderWrapper.duration**</font>
() -> float
获取视频时长
**Returns**
- **float** —— 视频时长。
#### <font color="#0FB0E4">function **FramesReaderWrapper.sample_frames**</font>
(decode_list: torch.Tensor) -> torch.Tensor
根据帧号,获取帧数据
**Parameters**
- **decode_list** —— 采样帧号列表。
**Returns**
- **tensor** —— 数据张量。
### <font color="#0FB0E4">class **EasyVideoReader**</font>
用于长视频读取、采样和预处理的类。
#### <font color="#0FB0E4">function **EasyVideoReader.__init__**</font>
(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** —— 预处理算子。
#### <font color="#0FB0E4">function **EasyVideoReader.__iter__**</font>
() -> int
迭代器
#### <font color="#0FB0E4">function **EasyVideoReader.__next__**</font>
() -> 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)
```
<hr/>
### <font color="#0FB0E4">class **Registry**</font>
注册器
#### <font color="#0FB0E4">function **Registry.__init__**</font>
(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** —— 该模块允许注册的类或者函数,默认都允许注册。
#### <font color="#0FB0E4">function **Registry.build**</font>
(cfg: Config, logger: logger = None, kwargs) -> cls_obj
build目标类的实例
**Returns**
- **cls_obj** —— 特定类的实例。
#### <font color="#0FB0E4">function **Registry.register_class**</font>
(name: str)
注册一个类
**Returns**
- **name** —— 注册名称。
#### <font color="#0FB0E4">function **Registry.register_function**</font>
(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)
```
<hr/>
#### <font color="#0FB0E4">function **transfer_data_to_numpy**</font>
(data: list/dict of torch.Tensor) -> (data: list/dict of numpy.ndarray)
将数据转移到numpy
**Parameters**
- **data** —— torch.Tensor并以list/dict形式存储。
**Returns**
- **data** —— numpy.ndarray并与输入一致的形式存储。
#### <font color="#0FB0E4">function **transfer_data_to_cpu**</font>
(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]并与输入一致的形式存储。
#### <font color="#0FB0E4">function **transfer_data_to_cuda**</font>
(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
```
<hr/>
#### <font color="#0FB0E4">function **move_model_to_cpu**</font>
(params: list/dict of torch.Tensor[cuda]) -> (data: torch.Tensor[cpu])
将参数数据从gpu转移到cpu上。
**Parameters**
- **params** —— torch.Tensor[cuda]并以OrderedDict形式存储。
**Returns**
- **params** —— torch.Tensor[cpu]并与输入一致的形式存储。
#### <font color="#0FB0E4">function **load_pretrained**</font>
(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时需要索引子层级。
#### <font color="#0FB0E4">function **count_params**</font>
(model: torch.nn.Module) -> (float)
统计模型的总参数。
**Parameters**
- **model** —— torch.nn.Module模型实例。
**Returns**
- **float** —— 模型参数量(浮点数个数)。
#### <font color="#0FB0E4">function **init_weights**</font>
(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
```
<hr/>
#### <font color="#0FB0E4">class **MultiFoldDistributedSampler**</font>
多fold采样器,支持在一个epoch中重复多轮数据
#### <font color="#0FB0E4">function **MultiFoldDistributedSampler.__init__**</font>
( 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** —— 数据是否要打乱。
#### <font color="#0FB0E4">function **MultiFoldDistributedSampler.__iter__**</font>
()
迭代器,每迭代一次得到一个样本的index
#### <font color="#0FB0E4">function **MultiFoldDistributedSampler.set_epoch**</font>
(epoch: int)
设置当前的epoch
**Parameters**
- **epoch** —— 当前的epoch。
#### <font color="#0FB0E4">class **EvalDistributedSampler**</font>
用于测试时的采样器,当不用padding模式的时候,会发现最后一个rank的数据会少于其他rank。
#### <font color="#0FB0E4">function **EvalDistributedSampler.__init__**</font>
( 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数据量一致。
#### <font color="#0FB0E4">function **EvalDistributedSampler.__iter__**</font>
()
迭代器,每迭代一次得到一个样本的index
#### <font color="#0FB0E4">function **EvalDistributedSampler.set_epoch**</font>
(epoch: int)
设置当前的epoch
**Parameters**
- **epoch** —— 当前的epoch。
#### <font color="#0FB0E4">class **MultiLevelBatchSampler**</font>
用于大规模数据的多级索引的sampler
#### <font color="#0FB0E4">function **MultiLevelBatchSampler.__init__**</font>
(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。
#### <font color="#0FB0E4">function **MultiLevelBatchSampler.__iter__**</font>
()
迭代器,每迭代一次得到一个样本的index
#### <font color="#0FB0E4">class **MixtureOfSamplers**</font>
用于大规模数据的多级索引的sampler
#### <font color="#0FB0E4">function **MixtureOfSamplers.__init__**</font>
(samplers: list(sampler), probabilities: list(float), rank: int =0, seed: int = 8888)
**Parameters**
- **samplers** —— 采样器列表,用于混合采样器。
- **probabilities** —— 每个采样器的概率。
- **rank** —— rank表示当前进程号。
- **seed** —— 随机采样的seed,在data.registry中获取全局seed。
#### <font color="#0FB0E4">function **MixtureOfSamplers.__iter__**</font>
()
迭代器,每迭代一次得到一个样本的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}"))
```
<hr/>
配合Hook使用如下(其中PROB_INTERVAL探针存储间隔,即调用probe_data()的次数):
```yaml
-
NAME: ProbeDataHook
PROB_INTERVAL: 100
```
#### <font color="#0FB0E4">class **ProbeData**</font>
探针数据的实例。
#### <font color="#0FB0E4">function **ProbeData.__init__**</font>
(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** —— 针对一些值统计频率。