update 1.4.0
This commit is contained in:
@@ -1,4 +1,23 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.utils import (config, distribute, file_clients,
|
||||
file_system, module_transform)
|
||||
from typing import TYPE_CHECKING
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from scepter.modules.utils import (config, distribute, file_clients,
|
||||
file_system, module_transform)
|
||||
else:
|
||||
_import_structure = {
|
||||
'utils': ['config', 'distribute', 'file_clients',
|
||||
'file_system', 'module_transform']
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
@@ -0,0 +1,486 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import ast
|
||||
import logging
|
||||
import os
|
||||
import os.path as osp
|
||||
import time
|
||||
import traceback
|
||||
from pathlib import Path
|
||||
from typing import Union, Any
|
||||
|
||||
|
||||
p = Path(__file__)
|
||||
|
||||
|
||||
SKIP_FUNCTION_SCANNING = True
|
||||
SCEPTER_PATH = p.resolve().parents[2]
|
||||
REGISTER_CLASS = 'register_class'
|
||||
IGNORED_PACKAGES = ['.']
|
||||
SCAN_SUB_FOLDERS = [
|
||||
'modules', 'studio', 'tools', 'workflow'
|
||||
]
|
||||
INDEXER_FILE = 'ast_indexer'
|
||||
DECORATOR_KEY = 'decorators'
|
||||
EXPRESS_KEY = 'express'
|
||||
FROM_IMPORT_KEY = 'from_imports'
|
||||
IMPORT_KEY = 'imports'
|
||||
FILE_NAME_KEY = 'filepath'
|
||||
INDEX_KEY = 'index'
|
||||
REQUIREMENT_KEY = 'requirements'
|
||||
MODULE_KEY = 'module'
|
||||
CLASS_NAME = 'class_name'
|
||||
|
||||
|
||||
def get_ast_logger():
|
||||
ast_logger = logging.getLogger('scepter.ast')
|
||||
ast_logger.setLevel(logging.INFO)
|
||||
return ast_logger
|
||||
|
||||
|
||||
logger = get_ast_logger()
|
||||
|
||||
|
||||
class AstScanning(object):
|
||||
def __init__(self) -> None:
|
||||
self.result_import = dict()
|
||||
self.result_from_import = dict()
|
||||
self.result_decorator = []
|
||||
self.express = []
|
||||
|
||||
def _is_sub_node(self, node: object) -> bool:
|
||||
return isinstance(node,
|
||||
ast.AST) and not isinstance(node, ast.expr_context)
|
||||
|
||||
def _is_leaf(self, node: ast.AST) -> bool:
|
||||
for field in node._fields:
|
||||
attr = getattr(node, field)
|
||||
if self._is_sub_node(attr):
|
||||
return False
|
||||
elif isinstance(attr, (list, tuple)):
|
||||
for val in attr:
|
||||
if self._is_sub_node(val):
|
||||
return False
|
||||
else:
|
||||
return True
|
||||
|
||||
def _skip_function(self, node: Union[ast.AST, 'str']) -> bool:
|
||||
if SKIP_FUNCTION_SCANNING:
|
||||
if type(node).__name__ == 'FunctionDef' or node == 'FunctionDef':
|
||||
return True
|
||||
return False
|
||||
|
||||
def _fields(self, n: ast.AST, show_offsets: bool = True) -> tuple:
|
||||
if show_offsets:
|
||||
return n._attributes + n._fields
|
||||
else:
|
||||
return n._fields
|
||||
|
||||
def _leaf(self, node: ast.AST, show_offsets: bool = True) -> str:
|
||||
output = dict()
|
||||
if isinstance(node, ast.AST):
|
||||
local_dict = dict()
|
||||
for field in self._fields(node, show_offsets=show_offsets):
|
||||
field_output = self._leaf(
|
||||
getattr(node, field), show_offsets=show_offsets)
|
||||
local_dict[field] = field_output
|
||||
output[type(node).__name__] = local_dict
|
||||
return output
|
||||
else:
|
||||
return node
|
||||
|
||||
def _refresh(self):
|
||||
self.result_import = dict()
|
||||
self.result_from_import = dict()
|
||||
self.result_decorator = []
|
||||
self.result_express = []
|
||||
|
||||
def scan_ast(self, node: Union[ast.AST, None, str]):
|
||||
self._setup_global()
|
||||
self.scan_import(node, indent=' ', show_offsets=False)
|
||||
|
||||
def scan_import(
|
||||
self,
|
||||
node: Union[ast.AST, None, str],
|
||||
show_offsets: bool = True,
|
||||
parent_node_name: str = '',
|
||||
) -> None | str | dict[Any, Any]:
|
||||
if node is None:
|
||||
return node
|
||||
elif self._is_leaf(node):
|
||||
return self._leaf(node, show_offsets=show_offsets)
|
||||
else:
|
||||
def _scan_import(el: Union[ast.AST, None, str],
|
||||
parent_node_name: str = '') -> str:
|
||||
return self.scan_import(
|
||||
el,
|
||||
show_offsets=show_offsets,
|
||||
parent_node_name=parent_node_name)
|
||||
|
||||
outputs = dict()
|
||||
# add relative path expression
|
||||
if type(node).__name__ == 'ImportFrom':
|
||||
level = getattr(node, 'level')
|
||||
if level >= 1:
|
||||
path_level = ''.join(['.'] * level)
|
||||
setattr(node, 'level', 0)
|
||||
module_name = getattr(node, 'module')
|
||||
if module_name is None:
|
||||
setattr(node, 'module', path_level)
|
||||
else:
|
||||
setattr(node, 'module', path_level + module_name)
|
||||
|
||||
for field in self._fields(node, show_offsets=show_offsets):
|
||||
attr = getattr(node, field)
|
||||
if not attr:
|
||||
outputs[field] = []
|
||||
elif self._skip_function(parent_node_name):
|
||||
continue
|
||||
elif (isinstance(attr, list) and len(attr) == 1
|
||||
and isinstance(attr[0], ast.AST)
|
||||
and self._is_leaf(attr[0])):
|
||||
local_out = _scan_import(attr[0])
|
||||
outputs[field] = local_out
|
||||
elif isinstance(attr, list):
|
||||
el_dict = dict()
|
||||
for el in attr:
|
||||
local_out = _scan_import(el, type(el).__name__)
|
||||
name = type(el).__name__
|
||||
if (name == 'Import' or name == 'ImportFrom'
|
||||
or parent_node_name == 'ImportFrom'
|
||||
or parent_node_name == 'Import'):
|
||||
if name not in el_dict:
|
||||
el_dict[name] = []
|
||||
el_dict[name].append(local_out)
|
||||
outputs[field] = el_dict
|
||||
elif isinstance(attr, ast.AST):
|
||||
output = _scan_import(attr)
|
||||
outputs[field] = output
|
||||
else:
|
||||
outputs[field] = attr
|
||||
|
||||
if (type(node).__name__ == 'Import'
|
||||
or type(node).__name__ == 'ImportFrom'):
|
||||
if type(node).__name__ == 'ImportFrom':
|
||||
if field == 'module':
|
||||
self.result_from_import[outputs[field]] = dict()
|
||||
if field == 'names':
|
||||
if isinstance(outputs[field]['alias'], list):
|
||||
item_name = []
|
||||
for item in outputs[field]['alias']:
|
||||
local_name = item['alias']['name']
|
||||
item_name.append(local_name)
|
||||
self.result_from_import[
|
||||
outputs['module']] = item_name
|
||||
else:
|
||||
local_name = outputs[field]['alias']['name']
|
||||
self.result_from_import[outputs['module']] = [
|
||||
local_name
|
||||
]
|
||||
|
||||
if type(node).__name__ == 'Import':
|
||||
final_dict = outputs[field]['alias']
|
||||
if isinstance(final_dict, list):
|
||||
for item in final_dict:
|
||||
self.result_import[item['alias']
|
||||
['name']] = item['alias']
|
||||
else:
|
||||
self.result_import[outputs[field]['alias']
|
||||
['name']] = final_dict
|
||||
|
||||
if 'decorator_list' == field and attr != []:
|
||||
for item in attr:
|
||||
setattr(item, CLASS_NAME, node.name)
|
||||
self.result_decorator.extend(attr)
|
||||
|
||||
if attr != [] and type(
|
||||
attr
|
||||
).__name__ == 'Call' and parent_node_name == 'Expr':
|
||||
self.result_express.append(attr)
|
||||
return {IMPORT_KEY: self.result_import,
|
||||
FROM_IMPORT_KEY: self.result_from_import,
|
||||
DECORATOR_KEY: self.result_decorator,
|
||||
EXPRESS_KEY: self.result_express}
|
||||
|
||||
def _parse_decorator(self, node: ast.AST) -> tuple:
|
||||
def _get_attribute_item(node: ast.AST) -> tuple:
|
||||
value, id, attr = None, None, None
|
||||
if type(node).__name__ == 'Attribute':
|
||||
value = getattr(node, 'value')
|
||||
id = getattr(value, 'id', None)
|
||||
attr = getattr(node, 'attr')
|
||||
if type(node).__name__ == 'Name':
|
||||
id = getattr(node, 'id')
|
||||
return id, attr
|
||||
|
||||
def _get_args_name(nodes: list) -> list:
|
||||
result = []
|
||||
for node in nodes:
|
||||
if type(node).__name__ == 'Str':
|
||||
result.append((node.s, None))
|
||||
elif type(node).__name__ == 'Constant':
|
||||
result.append((node.value, None))
|
||||
else:
|
||||
result.append(_get_attribute_item(node))
|
||||
return result
|
||||
|
||||
def _get_keyword_name(nodes: ast.AST) -> list:
|
||||
result = []
|
||||
for node in nodes:
|
||||
if type(node).__name__ == 'keyword':
|
||||
attribute_node = getattr(node, 'value')
|
||||
if type(attribute_node).__name__ == 'Str':
|
||||
result.append((getattr(node,
|
||||
'arg'), attribute_node.s, None))
|
||||
elif type(attribute_node).__name__ == 'Constant':
|
||||
result.append(
|
||||
(getattr(node, 'arg'), attribute_node.value, None))
|
||||
else:
|
||||
result.append((getattr(node, 'arg'), )
|
||||
+ _get_attribute_item(attribute_node))
|
||||
return result
|
||||
|
||||
functions = _get_attribute_item(node.func)
|
||||
args_list = _get_args_name(node.args)
|
||||
keyword_list = _get_keyword_name(node.keywords)
|
||||
return functions, args_list, keyword_list
|
||||
|
||||
def _registry_indexer(self, parsed_input: tuple, class_name: str) -> tuple:
|
||||
"""format registry information to a tuple indexer
|
||||
|
||||
Return:
|
||||
tuple: (MODELS, ClassName, RegisterName)
|
||||
"""
|
||||
functions, args_list, keyword_list = parsed_input
|
||||
|
||||
if REGISTER_CLASS != functions[1]:
|
||||
return None
|
||||
output = [functions[0]]
|
||||
return (output[0], class_name, args_list)
|
||||
|
||||
def parse_decorators(self, nodes: list) -> list:
|
||||
"""parse the AST nodes of decorators object to registry indexer
|
||||
|
||||
Args:
|
||||
nodes (list): list of AST decorator nodes
|
||||
|
||||
Returns:
|
||||
list: list of registry indexer
|
||||
"""
|
||||
results = []
|
||||
for node in nodes:
|
||||
if type(node).__name__ != 'Call':
|
||||
continue
|
||||
class_name = getattr(node, CLASS_NAME, None)
|
||||
func = getattr(node, 'func')
|
||||
if getattr(func, 'attr', None) != REGISTER_CLASS:
|
||||
continue
|
||||
|
||||
parse_output = self._parse_decorator(node)
|
||||
index = self._registry_indexer(parse_output, class_name)
|
||||
if None is not index:
|
||||
results.append(index)
|
||||
return results
|
||||
|
||||
def generate_ast(self, file):
|
||||
self._refresh()
|
||||
with open(file, 'r', encoding='utf8') as code:
|
||||
data = code.readlines()
|
||||
data = ''.join(data)
|
||||
node = ast.parse(data)
|
||||
output = self.scan_import(node, show_offsets=False)
|
||||
output[DECORATOR_KEY] = self.parse_decorators(output[DECORATOR_KEY])
|
||||
output[EXPRESS_KEY] = self.parse_decorators(output[EXPRESS_KEY])
|
||||
output[DECORATOR_KEY].extend(output[EXPRESS_KEY])
|
||||
return output
|
||||
|
||||
|
||||
class FilesAstScanning(object):
|
||||
def __init__(self) -> None:
|
||||
self.astScaner = AstScanning()
|
||||
self.file_dirs = []
|
||||
self.requirement_dirs = []
|
||||
|
||||
def _parse_import_path(self,
|
||||
import_package: str,
|
||||
current_path: str = None) -> str:
|
||||
"""
|
||||
Args:
|
||||
import_package (str): relative import or abs import
|
||||
current_path (str): path/to/current/file
|
||||
"""
|
||||
if import_package.startswith(IGNORED_PACKAGES[0]):
|
||||
return SCEPTER_PATH + '/' + '/'.join(
|
||||
import_package.split('.')[1:]) + '.py'
|
||||
elif import_package.startswith(IGNORED_PACKAGES[1]):
|
||||
current_path_list = current_path.split('/')
|
||||
import_package_list = import_package.split('.')
|
||||
level = 0
|
||||
for index, item in enumerate(import_package_list):
|
||||
if item != '':
|
||||
level = index
|
||||
break
|
||||
|
||||
abs_path_list = current_path_list[0:-level]
|
||||
abs_path_list.extend(import_package_list[index:])
|
||||
return '/' + '/'.join(abs_path_list) + '.py'
|
||||
else:
|
||||
return current_path
|
||||
|
||||
def parse_import(self, scan_result: dict) -> list:
|
||||
"""parse import and from import dicts to a third party package list
|
||||
|
||||
Args:
|
||||
scan_result (dict): including the import and from import result
|
||||
|
||||
Returns:
|
||||
list: a list of package ignored 'scepter' and relative path import
|
||||
"""
|
||||
output = []
|
||||
output.extend(list(scan_result[IMPORT_KEY].keys()))
|
||||
output.extend(list(scan_result[FROM_IMPORT_KEY].keys()))
|
||||
|
||||
# get the package name
|
||||
for index, item in enumerate(output):
|
||||
if '' == item.split('.')[0]:
|
||||
output[index] = '.'
|
||||
else:
|
||||
output[index] = item.split('.')[0]
|
||||
|
||||
ignored = set()
|
||||
for item in output:
|
||||
for ignored_package in IGNORED_PACKAGES:
|
||||
if item.startswith(ignored_package):
|
||||
ignored.add(item)
|
||||
return list(set(output) - set(ignored))
|
||||
|
||||
def traversal_files(self, path, check_sub_dir=None, include_init=False):
|
||||
self.file_dirs = []
|
||||
if check_sub_dir is None or len(check_sub_dir) == 0:
|
||||
self._traversal_files(path, include_init=include_init)
|
||||
else:
|
||||
for item in check_sub_dir:
|
||||
sub_dir = os.path.join(path, item)
|
||||
if os.path.isdir(sub_dir):
|
||||
self._traversal_files(sub_dir, include_init=include_init)
|
||||
|
||||
def _traversal_files(self, path, include_init=False):
|
||||
dir_list = os.scandir(path)
|
||||
for item in dir_list:
|
||||
if item.name == '__init__.py' and not include_init:
|
||||
continue
|
||||
elif (item.name.startswith('__')
|
||||
and item.name != '__init__.py') or item.name.endswith(
|
||||
'.json') or item.name.endswith('.md'):
|
||||
continue
|
||||
if item.is_dir():
|
||||
self._traversal_files(item.path, include_init=include_init)
|
||||
elif item.is_file() and item.name.endswith('.py'):
|
||||
self.file_dirs.append(item.path)
|
||||
elif item.is_file() and 'requirement' in item.name:
|
||||
self.requirement_dirs.append(item.path)
|
||||
|
||||
def _get_single_file_scan_result(self, file):
|
||||
try:
|
||||
output = self.astScaner.generate_ast(file)
|
||||
except Exception as e:
|
||||
detail = traceback.extract_tb(e.__traceback__)
|
||||
raise Exception(
|
||||
f'During ast indexing the file {file}, a related error excepted '
|
||||
f'in the file {detail[-1].filename} at line: '
|
||||
f'{detail[-1].lineno}: "{detail[-1].line}" with error msg: '
|
||||
f'"{type(e).__name__}: {e}", please double check the origin file {file} '
|
||||
f'to see whether the file is correctly edited.')
|
||||
|
||||
import_list = self.parse_import(output)
|
||||
return output[DECORATOR_KEY], import_list
|
||||
|
||||
def _inverted_index(self, forward_index):
|
||||
inverted_index = dict()
|
||||
for index in forward_index:
|
||||
for item in forward_index[index][DECORATOR_KEY]:
|
||||
inverted_index[item[:2]] = {
|
||||
FILE_NAME_KEY: index,
|
||||
IMPORT_KEY: forward_index[index][IMPORT_KEY],
|
||||
MODULE_KEY: forward_index[index][MODULE_KEY],
|
||||
}
|
||||
if item[-1]:
|
||||
for register_name in item[-1]:
|
||||
inverted_index[(item[0], register_name[0])] = {
|
||||
FILE_NAME_KEY: index,
|
||||
IMPORT_KEY: forward_index[index][IMPORT_KEY],
|
||||
MODULE_KEY: forward_index[index][MODULE_KEY],
|
||||
}
|
||||
return inverted_index
|
||||
|
||||
def _module_import(self, forward_index):
|
||||
module_import = dict()
|
||||
for index, value_dict in forward_index.items():
|
||||
module_import[value_dict[MODULE_KEY]] = value_dict[IMPORT_KEY]
|
||||
return module_import
|
||||
|
||||
def get_files_scan_results(self,
|
||||
target_file_list=None,
|
||||
target_dir=SCEPTER_PATH,
|
||||
target_folders=SCAN_SUB_FOLDERS):
|
||||
"""the entry method of the ast scan method
|
||||
|
||||
Args:
|
||||
target_file_list can override the dir and folders combine
|
||||
target_dir (str, optional): the absolute path of the target directory to be scanned. Defaults to None.
|
||||
target_folder (list, optional): the list of
|
||||
sub-folders to be scanned in the target folder.
|
||||
Defaults to SCAN_SUB_FOLDERS.
|
||||
|
||||
Returns:
|
||||
dict: indexer of registry
|
||||
"""
|
||||
start = time.time()
|
||||
if target_file_list is not None:
|
||||
self.file_dirs = target_file_list
|
||||
else:
|
||||
self.traversal_files(target_dir, target_folders)
|
||||
logger.info(
|
||||
f'AST-Scanning the path "{target_dir}" with the following sub folders {target_folders}'
|
||||
)
|
||||
|
||||
result = dict()
|
||||
for file in self.file_dirs:
|
||||
filepath = file[file.rfind('scepter'):]
|
||||
module_name = filepath.replace(osp.sep, '.').replace('.py', '')
|
||||
decorator_list, import_list = self._get_single_file_scan_result(
|
||||
file)
|
||||
result[file] = {
|
||||
DECORATOR_KEY: decorator_list,
|
||||
IMPORT_KEY: import_list,
|
||||
MODULE_KEY: module_name
|
||||
}
|
||||
|
||||
inverted_index_with_results = self._inverted_index(result)
|
||||
module_import = self._module_import(result)
|
||||
index = {
|
||||
INDEX_KEY: inverted_index_with_results,
|
||||
REQUIREMENT_KEY: module_import
|
||||
}
|
||||
logger.info(
|
||||
f'Scanning done! A number of {len(inverted_index_with_results)} '
|
||||
f'components indexed or updated! Time consumed {time.time()-start}s'
|
||||
)
|
||||
return index
|
||||
|
||||
|
||||
file_scanner = FilesAstScanning()
|
||||
file_index = None
|
||||
|
||||
|
||||
def load_index(file_list=None):
|
||||
global file_index
|
||||
if file_index is None:
|
||||
logger.info('Building ast index from scanning every file!')
|
||||
file_index = file_scanner.get_files_scan_results(file_list)
|
||||
return file_index
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
index = load_index()
|
||||
print(index)
|
||||
@@ -10,7 +10,8 @@ import sys
|
||||
|
||||
import yaml
|
||||
|
||||
from scepter.modules.utils.model import StdMsg
|
||||
from scepter.modules.utils.logger import StdMsg
|
||||
|
||||
|
||||
_SECURE_KEYWORDS = [
|
||||
'ENDPOINT', 'BUCKET', 'OSS_AK', 'OSS_SK', 'OSS', 'TOKEN', 'APPKEY'
|
||||
@@ -211,7 +212,7 @@ def dict_to_yaml(module_name, name, json_config, set_name=False, exclude_keys=[]
|
||||
return yaml_str
|
||||
|
||||
|
||||
pattern = re.compile('.*?(\${\w+}).*?') # noqa
|
||||
pattern = re.compile(r'.*?(\${\w+}).*?') # noqa
|
||||
|
||||
|
||||
def env_var_constructor(loader, node):
|
||||
@@ -669,4 +670,4 @@ class Config(object):
|
||||
return len(self.cfg_dict)
|
||||
|
||||
def pop(self, name):
|
||||
self.cfg_dict.pop(name)
|
||||
return self.cfg_dict.pop(name)
|
||||
|
||||
@@ -13,8 +13,8 @@ import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.autograd import Function
|
||||
|
||||
from scepter.modules.utils.model import StdMsg
|
||||
import platform
|
||||
from scepter.modules.utils.logger import StdMsg
|
||||
|
||||
__all__ = [
|
||||
'gather_data', 'we', 'broadcast', 'barrier', 'reduce_scatter', 'reduce',
|
||||
@@ -620,7 +620,11 @@ class Workenv(object):
|
||||
self.sync_bn = False
|
||||
self.rank = 0
|
||||
self.world_size = 1
|
||||
self.device_id = 0
|
||||
if torch.cuda.is_available():
|
||||
self.device_id = 0
|
||||
else:
|
||||
self.device_id = 'mps' if platform.system() == "Darwin" else 'cpu'
|
||||
|
||||
self.backend = ''
|
||||
self.device_count = 1
|
||||
self.seed = 2023
|
||||
@@ -665,7 +669,7 @@ class Workenv(object):
|
||||
self.share_storage = os.environ.get('SHARE_STORAGE', None) == 'true'
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
self.device_id = 'cpu'
|
||||
self.device_id = 'mps' if platform.system() == "Darwin" else 'cpu'
|
||||
fn(config)
|
||||
return
|
||||
|
||||
|
||||
@@ -0,0 +1,208 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
# docstyle-ignore
|
||||
ALBUMENTATIONS_IMPORT_ERROR = """
|
||||
{0} requires the albumentations library but it was not found in your environment. You can install it with pip:
|
||||
`pip install albumentations`
|
||||
"""
|
||||
|
||||
# docstyle-ignore
|
||||
SENTENCEPIECE_IMPORT_ERROR = """
|
||||
{0} requires the SentencePiece library but it was not found in your environment. Checkout the instructions on the
|
||||
installation page of its repo: https://github.com/google/sentencepiece#installation and follow the ones
|
||||
that match your environment.
|
||||
"""
|
||||
|
||||
# docstyle-ignore
|
||||
SKLEARN_IMPORT_ERROR = """
|
||||
{0} requires the scikit-learn library but it was not found in your environment. You can install it with:
|
||||
```
|
||||
pip install -U scikit-learn
|
||||
```
|
||||
In a notebook or a colab, you can install it by executing a cell with
|
||||
```
|
||||
!pip install -U scikit-learn
|
||||
```
|
||||
"""
|
||||
|
||||
# docstyle-ignore
|
||||
TIMM_IMPORT_ERROR = """
|
||||
{0} requires the timm library but it was not found in your environment. You can install it with pip:
|
||||
`pip install timm`
|
||||
"""
|
||||
|
||||
# docstyle-ignore
|
||||
SCEPTER_IMPORT_ERROR = """
|
||||
{0} requires the scepter library but it was not found in your environment. You can install it with pip:
|
||||
`pip install scepter`
|
||||
"""
|
||||
|
||||
# docstyle-ignore
|
||||
PYTORCH_IMPORT_ERROR = """
|
||||
{0} requires the PyTorch library but it was not found in your environment. Checkout the instructions on the
|
||||
installation page: https://pytorch.org/get-started/locally/ and follow the ones that match your environment.
|
||||
"""
|
||||
|
||||
WENETRUNTIME_IMPORT_ERROR = """
|
||||
{0} requires the wenetruntime library but it was not found in your environment. You can install it with pip:
|
||||
`pip install wenetruntime==TORCH_VER`
|
||||
"""
|
||||
|
||||
# docstyle-ignore
|
||||
TORCHVISION_IMPORT_ERROR = """
|
||||
{0} requires the scipy library but it was not found in your environment. You can install it with pip:
|
||||
`pip install torchvision`
|
||||
"""
|
||||
|
||||
# docstyle-ignore
|
||||
OPENCV_IMPORT_ERROR = """
|
||||
{0} requires the opencv library but it was not found in your environment. You can install it with pip:
|
||||
`pip install opencv-python`
|
||||
"""
|
||||
|
||||
PILLOW_IMPORT_ERROR = """
|
||||
{0} requires the Pillow library but it was not found in your environment. You can install it with pip:
|
||||
`pip install Pillow`
|
||||
"""
|
||||
|
||||
MODELSCOPE_IMPORT_ERROR = """
|
||||
{0} requires the modelscope library but it was not found in your environment. You can install it with pip:
|
||||
`pip install modelscope`
|
||||
"""
|
||||
|
||||
FLASH_ATTN_IMPORT_ERROR = """
|
||||
{0} requires the flash_attn library but it was not found in your environment. You can install it with pip:
|
||||
`pip install flash_attn==2.5.8`
|
||||
"""
|
||||
|
||||
XFORMERS_IMPORT_ERROR = """
|
||||
{0} requires the xformers library but it was not found in your environment. You can install it with pip:
|
||||
`pip install xformers`
|
||||
"""
|
||||
|
||||
DECORD_IMPORT_ERROR = """
|
||||
{0} requires the decord library but it was not found in your environment. You can install it with pip:
|
||||
`pip install decord>=0.6.0`
|
||||
"""
|
||||
|
||||
# docstyle-ignore
|
||||
BEAUTIFULSOUP4_IMPORT_ERROR = """
|
||||
{0} requires the decord library but it was not found in your environment. You can install it with pip:
|
||||
`pip install beautifulsoup4`
|
||||
"""
|
||||
|
||||
# docstyle-ignore
|
||||
BEZIER_IMPORT_ERROR = """
|
||||
{0} requires the beizer library but it was not found in your environment. You can install it with pip:
|
||||
`pip install beizer`
|
||||
"""
|
||||
|
||||
# docstyle-ignore
|
||||
EINOPS_IMPORT_ERROR = """
|
||||
{0} requires the einops library but it was not found in your environment. You can install it with pip:
|
||||
`pip install einops`
|
||||
"""
|
||||
|
||||
# docstyle-ignore
|
||||
EASYNLP_IMPORT_ERROR = """
|
||||
{0} requires the easynlp library but it was not found in your environment.
|
||||
You can install it with pip on linux or mac:
|
||||
`pip install pai-easynlp -f https://modelscope.oss-cn-beijing.aliyuncs.com/releases/repo.html`
|
||||
Or you can checkout the instructions on the
|
||||
installation page: https://github.com/alibaba/EasyNLP and follow the ones that match your environment.
|
||||
"""
|
||||
|
||||
# docstyle-ignore
|
||||
NUMPY_IMPORT_ERROR = """
|
||||
{0} requires the megatron_util library but it was not found in your environment. You can install it with pip:
|
||||
`pip install numpy`
|
||||
"""
|
||||
|
||||
# docstyle-ignore
|
||||
OSS2_IMPORT_ERROR = """
|
||||
{0} requires the oss2 library but it was not found in your environment. You can install it with pip:
|
||||
`pip install oss2`
|
||||
"""
|
||||
|
||||
# docstyle-ignore
|
||||
PYCOCOTOOLS_IMPORT_ERROR = """
|
||||
{0} requires the pycocotools library but it was not found in your environment. You can install it with pip:
|
||||
`pip install pycocotools`
|
||||
"""
|
||||
|
||||
# docstyle-ignore
|
||||
OPENCLIP_IMPORT_ERROR = """
|
||||
{0} requires the fasttext library but it was not found in your environment.
|
||||
You can install it with pip on linux or mac:
|
||||
`pip install open_clip_torch`
|
||||
Or you can checkout the instructions on the
|
||||
installation page: https://github.com/mlfoundations/open_clip and follow the ones that match your environment.
|
||||
"""
|
||||
|
||||
# docstyle-ignore
|
||||
PYYAML_IMPORT_ERROR = """
|
||||
{0} requires the pyyaml library but it was not found in your environment. You can install it with pip:
|
||||
`pip install pyyaml`
|
||||
"""
|
||||
|
||||
# docstyle-ignore
|
||||
SWIFT_IMPORT_ERROR = """
|
||||
{0} requires the ms-swift library but it was not found in your environment. You can install it with pip:
|
||||
`pip install ms-swift`
|
||||
"""
|
||||
|
||||
SCIKIT_IMAGE_IMPORT_ERROR = """
|
||||
{0} requires the ms-swift library but it was not found in your environment. You can install it with pip:
|
||||
`pip install scikit-image`
|
||||
"""
|
||||
|
||||
SCIKIT_LEARN_IMPORT_ERROR = """
|
||||
{0} requires the ms-swift library but it was not found in your environment. You can install it with pip:
|
||||
`pip install scikit-learn`
|
||||
"""
|
||||
|
||||
TORCHSDE_IMPORT_ERROR = """
|
||||
{0} requires the ms-swift library but it was not found in your environment. You can install it with pip:
|
||||
`pip install torchsde
|
||||
"""
|
||||
|
||||
BITSANDBYTES_IMPORT_ERROR = """
|
||||
{0} requires the ms-swift library but it was not found in your environment. You can install it with pip:
|
||||
`pip install bitsandbytes
|
||||
"""
|
||||
|
||||
GRADIO_IMAGESLIDER_IMPORT_ERROR = """
|
||||
{0} requires the ms-swift library but it was not found in your environment. You can install it with pip:
|
||||
`pip install gradio_imageslider
|
||||
"""
|
||||
|
||||
IMAGEHASH_IMPORT_ERROR = """
|
||||
{0} requires the ms-swift library but it was not found in your environment. You can install it with pip:
|
||||
`pip install imagehash
|
||||
"""
|
||||
|
||||
PSUTIL_IMPORT_ERROR = """
|
||||
{0} requires the ms-swift library but it was not found in your environment. You can install it with pip:
|
||||
`pip install psutil
|
||||
"""
|
||||
|
||||
TIKTOKEN_IMPORT_ERROR = """
|
||||
{0} requires the ms-swift library but it was not found in your environment. You can install it with pip:
|
||||
`pip install tiktoken
|
||||
"""
|
||||
|
||||
TRANSFORMERS_IMPORT_ERROR = """
|
||||
{0} requires the transformers library but it was not found in your environment. You can install it with pip:
|
||||
`pip install transformers`
|
||||
"""
|
||||
|
||||
GENERAL_IMPORT_ERROR = """
|
||||
{0} requires the REQ library but it was not found in your environment. You can install it with pip:
|
||||
`pip install REQ`
|
||||
"""
|
||||
|
||||
GRADIO_IMPORT_ERROR = """
|
||||
{0} requires the gradio library but it was not found in your environment. You can install it with pip:
|
||||
`pip install gradio`
|
||||
"""
|
||||
@@ -1,7 +1,29 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.utils.file_clients.aliyun_oss_fs import AliyunOssFs
|
||||
from scepter.modules.utils.file_clients.http_fs import HttpFs
|
||||
from scepter.modules.utils.file_clients.huggingface_fs import HuggingfaceFs
|
||||
from scepter.modules.utils.file_clients.local_fs import LocalFs
|
||||
from scepter.modules.utils.file_clients.modelscope_fs import ModelscopeFs
|
||||
from typing import TYPE_CHECKING
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from scepter.modules.utils.file_clients.aliyun_oss_fs import AliyunOssFs
|
||||
from scepter.modules.utils.file_clients.http_fs import HttpFs
|
||||
from scepter.modules.utils.file_clients.huggingface_fs import HuggingfaceFs
|
||||
from scepter.modules.utils.file_clients.local_fs import LocalFs
|
||||
from scepter.modules.utils.file_clients.modelscope_fs import ModelscopeFs
|
||||
else:
|
||||
_import_structure = {
|
||||
'aliyun_oss_fs': ['AliyunOssFs'],
|
||||
'http_fs': ['HttpFs'],
|
||||
'huggingface_fs': ['HuggingfaceFs'],
|
||||
'local_fs': ['LocalFs'],
|
||||
'modelscope_fs': ['ModelscopeFs']
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
@@ -0,0 +1,333 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import functools
|
||||
import importlib
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from collections import OrderedDict
|
||||
from importlib import import_module
|
||||
from itertools import chain
|
||||
from types import ModuleType
|
||||
from typing import Any
|
||||
import scepter
|
||||
from scepter.modules.utils.ast_utils import (INDEX_KEY,
|
||||
MODULE_KEY,
|
||||
REQUIREMENT_KEY,
|
||||
load_index)
|
||||
from scepter.modules.utils.error import *
|
||||
from scepter.modules.utils.logger import get_logger
|
||||
|
||||
|
||||
if sys.version_info < (3, 8):
|
||||
import importlib_metadata
|
||||
else:
|
||||
import importlib.metadata as importlib_metadata
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
def get_dirname():
|
||||
return os.path.dirname(scepter.__file__)
|
||||
|
||||
|
||||
def import_modules(imports, allow_failed_imports=False):
|
||||
"""Import modules from the given list of strings.
|
||||
|
||||
Args:
|
||||
imports (list | str | None): The given module names to be imported.
|
||||
allow_failed_imports (bool): If True, the failed imports will return
|
||||
None. Otherwise, an ImportError is raise. Default: False.
|
||||
|
||||
Returns:
|
||||
list[module] | module | None: The imported modules.
|
||||
|
||||
Examples:
|
||||
>>> osp, sys = import_modules(
|
||||
... ['os.path', 'sys'])
|
||||
>>> import os.path as osp_
|
||||
>>> import sys as sys_
|
||||
>>> assert osp == osp_
|
||||
>>> assert sys == sys_
|
||||
"""
|
||||
if not imports:
|
||||
return
|
||||
single_import = False
|
||||
if isinstance(imports, str):
|
||||
single_import = True
|
||||
imports = [imports]
|
||||
if not isinstance(imports, list):
|
||||
raise TypeError(
|
||||
f'custom_imports must be a list but got type {type(imports)}')
|
||||
imported = []
|
||||
for imp in imports:
|
||||
if not isinstance(imp, str):
|
||||
raise TypeError(
|
||||
f'{imp} is of type {type(imp)} and cannot be imported.')
|
||||
try:
|
||||
imported_tmp = import_module(imp)
|
||||
except ImportError:
|
||||
if allow_failed_imports:
|
||||
logger.warning(f'{imp} failed to import and is ignored.')
|
||||
imported_tmp = None
|
||||
else:
|
||||
raise ImportError
|
||||
imported.append(imported_tmp)
|
||||
if single_import:
|
||||
imported = imported[0]
|
||||
return imported
|
||||
|
||||
|
||||
# following code borrows implementation from huggingface/transformers
|
||||
ENV_VARS_TRUE_VALUES = {'1', 'ON', 'YES', 'TRUE'}
|
||||
ENV_VARS_TRUE_AND_AUTO_VALUES = ENV_VARS_TRUE_VALUES.union({'AUTO'})
|
||||
USE_TORCH = os.environ.get('USE_TORCH', 'AUTO').upper()
|
||||
|
||||
_torch_version = 'N/A'
|
||||
if USE_TORCH in ENV_VARS_TRUE_AND_AUTO_VALUES:
|
||||
_torch_available = importlib.util.find_spec('torch') is not None
|
||||
if _torch_available:
|
||||
try:
|
||||
_torch_version = importlib_metadata.version('torch')
|
||||
logger.info(f'PyTorch version {_torch_version} Found.')
|
||||
except importlib_metadata.PackageNotFoundError:
|
||||
_torch_available = False
|
||||
else:
|
||||
logger.info('Disabling PyTorch because USE_TF is set')
|
||||
_torch_available = False
|
||||
|
||||
|
||||
def is_torchvision_available():
|
||||
return importlib.util.find_spec('torchvision') is not None
|
||||
|
||||
|
||||
def is_sentencepiece_available():
|
||||
return importlib.util.find_spec('sentencepiece') is not None
|
||||
|
||||
|
||||
def is_scepter_available():
|
||||
return importlib.util.find_spec('scepter') is not None
|
||||
|
||||
|
||||
def is_torch_available():
|
||||
return _torch_available
|
||||
|
||||
|
||||
def is_torch_cuda_available():
|
||||
if is_torch_available():
|
||||
import torch
|
||||
return torch.cuda.is_available()
|
||||
else:
|
||||
return False
|
||||
|
||||
|
||||
def is_swift_available():
|
||||
return importlib.util.find_spec('swift') is not None
|
||||
|
||||
|
||||
def is_opencv_available():
|
||||
return importlib.util.find_spec('cv2') is not None
|
||||
|
||||
|
||||
def is_pillow_available():
|
||||
return importlib.util.find_spec('PIL.Image') is not None
|
||||
|
||||
|
||||
def _is_package_available_fn(pkg_name):
|
||||
return importlib.util.find_spec(pkg_name) is not None
|
||||
|
||||
|
||||
def is_package_available(pkg_name):
|
||||
return functools.partial(_is_package_available_fn, pkg_name)
|
||||
|
||||
|
||||
def is_flash_attn_available():
|
||||
return importlib.util.find_spec('flash-attn') is not None
|
||||
|
||||
|
||||
def is_transformers_available():
|
||||
return importlib.util.find_spec('transformers') is not None
|
||||
|
||||
|
||||
REQUIREMENTS_MAAPING = OrderedDict([
|
||||
('scepter', (is_scepter_available(), SCEPTER_IMPORT_ERROR)),
|
||||
('torch', (is_torch_available, PYTORCH_IMPORT_ERROR)),
|
||||
('torchvision', (is_torchvision_available(), TORCHVISION_IMPORT_ERROR)),
|
||||
('cv2', (is_opencv_available, OPENCV_IMPORT_ERROR)),
|
||||
('PIL', (is_pillow_available, PILLOW_IMPORT_ERROR)),
|
||||
('modelscope', (is_package_available('modelscope'), MODELSCOPE_IMPORT_ERROR)),
|
||||
('flash-attn', (is_flash_attn_available, FLASH_ATTN_IMPORT_ERROR)),
|
||||
('xformers', (is_package_available('funasr'), XFORMERS_IMPORT_ERROR)),
|
||||
('albumentations', (is_package_available('albumentations'), ALBUMENTATIONS_IMPORT_ERROR)),
|
||||
('decord', (is_package_available('decord'), DECORD_IMPORT_ERROR)),
|
||||
('beautifulsoup4', (is_package_available('beautifulsoup4'), BEAUTIFULSOUP4_IMPORT_ERROR)),
|
||||
('bezier', (is_package_available('bezier'), BEZIER_IMPORT_ERROR)),
|
||||
('einops', (is_package_available('einops'), EINOPS_IMPORT_ERROR)),
|
||||
('numpy', (is_package_available('numpy'), NUMPY_IMPORT_ERROR)),
|
||||
('oss2', (is_package_available('oss2'), OSS2_IMPORT_ERROR)),
|
||||
('pycocotools', (is_package_available('pycocotools'), PYCOCOTOOLS_IMPORT_ERROR)),
|
||||
('open_clip', (is_package_available('open_clip'), OPENCLIP_IMPORT_ERROR)),
|
||||
('pyyaml', (is_package_available('pyyaml'), PYYAML_IMPORT_ERROR)),
|
||||
('transformers', (is_package_available('transformers'), TRANSFORMERS_IMPORT_ERROR)),
|
||||
('ms-swift', (is_package_available('ms-swift'), SWIFT_IMPORT_ERROR)),
|
||||
('gradio', (is_package_available('gradio'), SWIFT_IMPORT_ERROR)),
|
||||
('scikit-image', (is_package_available('scikit-image'), SCIKIT_IMAGE_IMPORT_ERROR)),
|
||||
('scikit-learn', (is_package_available('scikit-learn'), SCIKIT_LEARN_IMPORT_ERROR)),
|
||||
('sentencepiece', (is_package_available('sentencepiece'), SENTENCEPIECE_IMPORT_ERROR)),
|
||||
('torchsde', (is_package_available('torchsde'), TORCHSDE_IMPORT_ERROR)),
|
||||
('bitsandbytes', (is_package_available('bitsandbytes'), BITSANDBYTES_IMPORT_ERROR)),
|
||||
('gradio_imageslider', (is_package_available('gradio_imageslider'), GRADIO_IMAGESLIDER_IMPORT_ERROR)),
|
||||
('imagehash', (is_package_available('imagehash'), IMAGEHASH_IMPORT_ERROR)),
|
||||
('psutil', (is_package_available('psutil'), PSUTIL_IMPORT_ERROR)),
|
||||
('tiktoken', (is_package_available('tiktoken'), TIKTOKEN_IMPORT_ERROR))
|
||||
])
|
||||
|
||||
SYSTEM_PACKAGE = set(['os', 'sys', 'typing'])
|
||||
|
||||
|
||||
def requires(obj, requirements):
|
||||
if not isinstance(requirements, (list, tuple)):
|
||||
requirements = [requirements]
|
||||
if isinstance(obj, str):
|
||||
name = obj
|
||||
else:
|
||||
name = obj.__name__ if hasattr(obj,
|
||||
'__name__') else obj.__class__.__name__
|
||||
|
||||
checks = []
|
||||
for req in requirements:
|
||||
if req == '' or req in SYSTEM_PACKAGE:
|
||||
continue
|
||||
if req in REQUIREMENTS_MAAPING:
|
||||
check = REQUIREMENTS_MAAPING[req]
|
||||
else:
|
||||
check_fn = is_package_available(req)
|
||||
err_msg = GENERAL_IMPORT_ERROR.replace('REQ', req)
|
||||
check = (check_fn, err_msg)
|
||||
checks.append(check)
|
||||
|
||||
failed = [msg.format(name) for available, msg in checks if not available]
|
||||
if failed:
|
||||
raise ImportError(''.join(failed))
|
||||
|
||||
|
||||
def torch_required(func):
|
||||
# Chose a different decorator name than in tests so it's clear they are not the same.
|
||||
@functools.wraps(func)
|
||||
def wrapper(*args, **kwargs):
|
||||
if is_torch_available():
|
||||
return func(*args, **kwargs)
|
||||
else:
|
||||
raise ImportError(f'Method `{func.__name__}` requires PyTorch.')
|
||||
return wrapper
|
||||
|
||||
|
||||
class LazyImportModule(ModuleType):
|
||||
_AST_INDEX = None
|
||||
|
||||
def __init__(self,
|
||||
name,
|
||||
module_file,
|
||||
import_structure,
|
||||
module_spec=None,
|
||||
extra_objects=None,
|
||||
try_to_pre_import=False):
|
||||
super().__init__(name)
|
||||
self._modules = set(import_structure.keys())
|
||||
self._class_to_module = {}
|
||||
for key, values in import_structure.items():
|
||||
for value in values:
|
||||
self._class_to_module[value] = key
|
||||
# Needed for autocompletion in an IDE
|
||||
self.__all__ = list(import_structure.keys()) + list(
|
||||
chain(*import_structure.values()))
|
||||
self.__file__ = module_file
|
||||
self.__spec__ = module_spec
|
||||
self.__path__ = [os.path.dirname(module_file)]
|
||||
self._objects = {} if extra_objects is None else extra_objects
|
||||
self._name = name
|
||||
self._import_structure = import_structure
|
||||
if try_to_pre_import:
|
||||
self._try_to_import()
|
||||
|
||||
def _try_to_import(self):
|
||||
for sub_module in self._class_to_module.keys():
|
||||
try:
|
||||
getattr(self, sub_module)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f'pre load module {sub_module} error, please check {e}')
|
||||
|
||||
def __dir__(self):
|
||||
result = super().__dir__()
|
||||
for attr in self.__all__:
|
||||
if attr not in result:
|
||||
result.append(attr)
|
||||
return result
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
if name in self._objects:
|
||||
return self._objects[name]
|
||||
if name in self._modules:
|
||||
value = self._get_module(name)
|
||||
elif name in self._class_to_module.keys():
|
||||
module = self._get_module(self._class_to_module[name])
|
||||
value = getattr(module, name)
|
||||
else:
|
||||
raise AttributeError(
|
||||
f'module {self.__name__} has no attribute {name}')
|
||||
|
||||
setattr(self, name, value)
|
||||
return value
|
||||
|
||||
def _get_module(self, module_name: str):
|
||||
try:
|
||||
module_name_full = self.__name__ + '.' + module_name
|
||||
if not any(
|
||||
module_name_full.startswith(f'scepter.{prefix}')
|
||||
for prefix in ['modules', 'studio', 'version', 'tools', 'workflow']):
|
||||
# check requirements before module import
|
||||
requirements = self.get_requirements()
|
||||
if module_name_full in requirements:
|
||||
requires(module_name_full, requirements)
|
||||
return importlib.import_module('.' + module_name, self.__name__)
|
||||
except Exception as e:
|
||||
raise RuntimeError(
|
||||
f'Failed to import {self.__name__}.{module_name} because of the following error '
|
||||
f'(look up to see its traceback):\n{e}') from e
|
||||
|
||||
def __reduce__(self):
|
||||
return self.__class__, (self._name, self.__file__,
|
||||
self._import_structure)
|
||||
|
||||
@staticmethod
|
||||
def get_ast_index():
|
||||
if LazyImportModule._AST_INDEX is None:
|
||||
LazyImportModule._AST_INDEX = load_index()
|
||||
return LazyImportModule._AST_INDEX
|
||||
|
||||
@staticmethod
|
||||
def import_module(signature):
|
||||
""" import a lazy import module using signature
|
||||
|
||||
Args:
|
||||
signature (tuple): a tuple of str, (registry_name, class_name)
|
||||
"""
|
||||
ast_index = LazyImportModule.get_ast_index()
|
||||
if signature in ast_index[INDEX_KEY]:
|
||||
mod_index = ast_index[INDEX_KEY][signature]
|
||||
module_name = mod_index[MODULE_KEY]
|
||||
if module_name in ast_index[REQUIREMENT_KEY]:
|
||||
requirements = ast_index[REQUIREMENT_KEY][module_name]
|
||||
requires(module_name, requirements)
|
||||
importlib.import_module(module_name)
|
||||
else:
|
||||
logger.warning(f'{signature} not found in ast index file')
|
||||
|
||||
@staticmethod
|
||||
def get_module_type(module):
|
||||
ast_index = LazyImportModule.get_ast_index()
|
||||
if module in ast_index[INDEX_KEY]:
|
||||
return True
|
||||
else:
|
||||
return False
|
||||
@@ -1,15 +1,13 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import logging
|
||||
import numbers
|
||||
import sys
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from scepter.modules.utils.distribute import get_dist_info
|
||||
import numbers
|
||||
|
||||
|
||||
def as_time(s):
|
||||
@@ -71,6 +69,7 @@ def init_logger(in_logger, log_file=None, dist_launcher='pytorch'):
|
||||
log_file (str, None): if not None, a file handler will be add to in_logger
|
||||
dist_launcher (str, None):
|
||||
"""
|
||||
from scepter.modules.utils.distribute import get_dist_info
|
||||
rank, _ = get_dist_info()
|
||||
if rank == 0:
|
||||
if log_file is not None:
|
||||
@@ -93,6 +92,20 @@ def init_logger(in_logger, log_file=None, dist_launcher='pytorch'):
|
||||
in_logger.setLevel(logging.INFO)
|
||||
|
||||
|
||||
class StdMsg():
|
||||
def __init__(self, name='msg'):
|
||||
self.name = name
|
||||
|
||||
def info(self, msg):
|
||||
sys.stdout.write('[Info]: ' + msg + '\n')
|
||||
|
||||
def error(self, msg):
|
||||
sys.stdout.write('[Error]: ' + msg + '\n')
|
||||
|
||||
def warning(self, msg):
|
||||
sys.stdout.write('[Warning]: ' + msg + '\n')
|
||||
|
||||
|
||||
class LogAgg(object):
|
||||
""" Log variable aggregate tool. Recommend to invoke clear() function after one epoch.
|
||||
In distributed training environment, tensor variable will be all reduced to get an average.
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
@@ -10,20 +9,6 @@ import torch.nn as nn
|
||||
from torch.utils.model_zoo import load_url as load_state_dict_from_url
|
||||
|
||||
|
||||
class StdMsg():
|
||||
def __init__(self, name='msg'):
|
||||
self.name = name
|
||||
|
||||
def info(self, msg):
|
||||
sys.stdout.write('[Info]: ' + msg + '\n')
|
||||
|
||||
def error(self, msg):
|
||||
sys.stdout.write('[Error]: ' + msg + '\n')
|
||||
|
||||
def warning(self, msg):
|
||||
sys.stdout.write('[Warning]: ' + msg + '\n')
|
||||
|
||||
|
||||
def move_model_to_cpu(params):
|
||||
cpu_params = OrderedDict()
|
||||
for key, val in params.items():
|
||||
@@ -41,7 +26,7 @@ def load_pretrained(model: torch.nn.Module,
|
||||
f'Load pretrained model [{model.__class__.__name__}] from {path}')
|
||||
if os.path.exists(path):
|
||||
# From local
|
||||
state_dict = torch.load(path, map_location)
|
||||
state_dict = torch.load(path, map_location, weights_only=True)
|
||||
elif path.startswith('http'):
|
||||
# From url
|
||||
state_dict = load_state_dict_from_url(path,
|
||||
|
||||
@@ -65,6 +65,13 @@ def build_from_config(cfg, registry, logger=None, *args, **kwargs):
|
||||
|
||||
cfg = deep_copy(cfg)
|
||||
req_type = cfg.get('NAME')
|
||||
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
sig = (registry.name.upper(), req_type)
|
||||
if (LazyImportModule.get_module_type(sig)
|
||||
and req_type not in registry.class_map.keys()):
|
||||
LazyImportModule.import_module(sig)
|
||||
|
||||
if isinstance(req_type, str):
|
||||
req_type_entry = registry.get(req_type)
|
||||
if req_type_entry is None:
|
||||
|
||||
@@ -1,6 +1,26 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from .frame_sampler import (FRAME_SAMPLERS, IntervalSampler, SegmentSampler,
|
||||
UniformSampler, do_frame_sample)
|
||||
from .video_reader import (EasyVideoReader, FramesReaderWrapper,
|
||||
VideoReaderWrapper)
|
||||
from typing import TYPE_CHECKING
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .frame_sampler import (FRAME_SAMPLERS, IntervalSampler, SegmentSampler,
|
||||
UniformSampler, do_frame_sample)
|
||||
from .video_reader import (EasyVideoReader, FramesReaderWrapper,
|
||||
VideoReaderWrapper)
|
||||
else:
|
||||
_import_structure = {
|
||||
'frame_sampler': ['FRAME_SAMPLERS', 'IntervalSampler', 'SegmentSampler',
|
||||
'UniformSampler', 'do_frame_sample'],
|
||||
'video_reader': ['EasyVideoReader', 'FramesReaderWrapper', 'VideoReaderWrapper']
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
from enum import Enum
|
||||
import os
|
||||
from enum import Enum
|
||||
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
@@ -17,17 +17,15 @@ class Media(Enum):
|
||||
|
||||
|
||||
class HtmlVisualization(object):
|
||||
def __init__(
|
||||
self,
|
||||
allow_annotation=False,
|
||||
slice_size=1000,
|
||||
align='center',
|
||||
width_scale='60%',
|
||||
title='Visualization',
|
||||
height=600,
|
||||
width=None,
|
||||
text_cols=40
|
||||
):
|
||||
def __init__(self,
|
||||
allow_annotation=False,
|
||||
slice_size=1000,
|
||||
align='center',
|
||||
width_scale='60%',
|
||||
title='Visualization',
|
||||
height=600,
|
||||
width=None,
|
||||
text_cols=40):
|
||||
self.content_list = []
|
||||
self.rows_meta = []
|
||||
self.allow_annotation = allow_annotation
|
||||
@@ -37,9 +35,9 @@ class HtmlVisualization(object):
|
||||
self.title = title
|
||||
self.html_start = '<html>'
|
||||
self.html_head = f'<head><meta charset="utf-8"><title>{title}</title></head>'
|
||||
self.height = height if height is not None else "600"
|
||||
self.width = width if width is not None else "auto"
|
||||
self.text_cols = text_cols if text_cols is not None else "auto"
|
||||
self.height = height if height is not None else '600'
|
||||
self.width = width if width is not None else 'auto'
|
||||
self.text_cols = text_cols if text_cols is not None else 'auto'
|
||||
self.html_style = ('''
|
||||
<style> \n
|
||||
.container {
|
||||
@@ -88,12 +86,11 @@ class HtmlVisualization(object):
|
||||
resize: none; \n
|
||||
border: 1px solid #ccc; \n
|
||||
} \n
|
||||
.large-checkbox {transform: scale(2.5); margin-left: 20px; margin-bottom: 20px; vertical-align: middle;} \n
|
||||
.large-checkbox {transform: scale(2.5); margin-left: 20px; margin-bottom: 20px; vertical-align: middle;} \n # noqa
|
||||
</style> \n
|
||||
\n
|
||||
'''.replace('{width_scale}',
|
||||
self.width_scale).replace('{align}', self.align)
|
||||
.replace('{pair_height}', f'{self.height}'))
|
||||
'''.replace('{width_scale}', self.width_scale).replace(
|
||||
'{align}', self.align).replace('{pair_height}', f'{self.height}'))
|
||||
|
||||
self.html_body_script = '''
|
||||
<script>\n
|
||||
@@ -120,9 +117,6 @@ class HtmlVisualization(object):
|
||||
|
||||
let percentage = (clientX - left) / width * 100;\n
|
||||
|
||||
|
||||
// 限制百分比在0到100之间\n
|
||||
|
||||
percentage = Math.max(0, Math.min(100, percentage));\n
|
||||
|
||||
media2.style.clipPath = `inset(0 ${100 - percentage}% 0 0)`;\n
|
||||
@@ -132,7 +126,6 @@ class HtmlVisualization(object):
|
||||
console.info(slider.style.left);\n
|
||||
|
||||
});\n
|
||||
// 初始化滑块位置\n
|
||||
slider.style.left = '50%';\n
|
||||
});\n
|
||||
</script>\n
|
||||
@@ -165,17 +158,16 @@ class HtmlVisualization(object):
|
||||
|
||||
'''
|
||||
self.label_button = (
|
||||
'<table><tr><td>' +
|
||||
"<button style='height: 50px;' type=\"button\" onclick=\"saveSamples()\">Save Samples</button>"
|
||||
+ '</td></tr></table>')
|
||||
'<table><tr><td>' +
|
||||
"<button style='height: 50px;' type=\"button\" onclick=\"saveSamples()\">Save Samples</button>"
|
||||
+ '</td></tr></table>')
|
||||
|
||||
def format_col(self,
|
||||
content='',
|
||||
label='',
|
||||
type=Media.TEXT,
|
||||
show_label=True,
|
||||
cols_span=1
|
||||
):
|
||||
cols_span=1):
|
||||
if type == Media.TEXT:
|
||||
ret_str = '<textarea' # noqa: E501
|
||||
if self.height is not None:
|
||||
@@ -185,7 +177,7 @@ class HtmlVisualization(object):
|
||||
cols = f"cols={self.text_cols * cols_span}"
|
||||
ret_str += f" {cols}"
|
||||
ret_str += f'>"{content}"</textarea>'
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ''
|
||||
elif type == Media.IMAGE:
|
||||
ret_str = f'<img src="{content}"'
|
||||
if self.height is not None:
|
||||
@@ -195,7 +187,7 @@ class HtmlVisualization(object):
|
||||
width = f'width="{self.width}"'
|
||||
ret_str += f" {width}"
|
||||
ret_str += ' >'
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ''
|
||||
elif type == Media.VIDEO:
|
||||
ret_str = '<video' # noqa
|
||||
if self.height is not None:
|
||||
@@ -206,34 +198,36 @@ class HtmlVisualization(object):
|
||||
ret_str += f" {width}"
|
||||
ret_str += ' preload="none" autoplay muted loop>'
|
||||
ret_str += f'<source src="{content}" type="video/mp4"></video>'
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ''
|
||||
elif type == Media.AUDIO:
|
||||
ret_str = f'<audio src="{content}" controls>'
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ''
|
||||
elif type == Media.IMAGE_PAIR:
|
||||
assert isinstance(content, (list, tuple)) and len(content) == 2
|
||||
ret_str = f'\n'
|
||||
ret_str += f' <div class="container"'
|
||||
ret_str += (f'> \n'
|
||||
f' <div class="image" id="media1">'
|
||||
f' <img src="{content[1]}" alt="before">\n'
|
||||
f' </div>\n'
|
||||
f' <div class="image" id="media2" style="clip-path: inset(0 50% 0 0);">\n'
|
||||
f' <img src="{content[0]}" alt="after">\n'
|
||||
f' </div>\n'
|
||||
f' <div class="slider" id="slider"></div>\n'
|
||||
f'')
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
ret_str = f'\n' # noqa
|
||||
ret_str += f' <div class="container"' # noqa
|
||||
ret_str += (
|
||||
f'> \n'
|
||||
f' <div class="image" id="media1">' # noqa
|
||||
f' <img src="{content[1]}" alt="before">\n' # noqa
|
||||
f' </div>\n' # noqa
|
||||
f' <div class="image" id="media2" style="clip-path: inset(0 50% 0 0);">\n' # noqa
|
||||
f' <img src="{content[0]}" alt="after">\n' # noqa
|
||||
f' </div>\n' # noqa
|
||||
f' <div class="slider" id="slider"></div>\n' # noqa
|
||||
f'')
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ''
|
||||
elif type == Media.VIDEO_PAIR:
|
||||
assert isinstance(content, (list, tuple)) and len(content) == 2
|
||||
ret_str = f'\n'
|
||||
ret_str += f' <div class="container"'
|
||||
ret_str += (f'> \n'
|
||||
f' <video autoplay muted loop class="video" id="media1"><source src="{content[1]}" type="video/mp4"></video>\n'
|
||||
f' <video autoplay muted loop class="video" id="media2" style="clip-path: inset(0 50% 0 0);"><source src="{content[0]}" type="video/mp4"></video>\n'
|
||||
f' <div class="slider" id="slider"></div>\n'
|
||||
f'</div>')
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
ret_str = f'\n' # noqa
|
||||
ret_str += f' <div class="container"' # noqa
|
||||
ret_str += (
|
||||
f'> \n'
|
||||
f' <video autoplay muted loop class="video" id="media1"><source src="{content[1]}" type="video/mp4"></video>\n' # noqa
|
||||
f' <video autoplay muted loop class="video" id="media2" style="clip-path: inset(0 50% 0 0);"><source src="{content[0]}" type="video/mp4"></video>\n' # noqa
|
||||
f' <div class="slider" id="slider"></div>\n' # noqa
|
||||
f'</div>')
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ''
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -242,10 +236,10 @@ class HtmlVisualization(object):
|
||||
|
||||
if cols_span > 1:
|
||||
ret_str = f'<th colspan="{cols_span}">{ret_str}</th>\n'
|
||||
sec_ret_str = f'<th colspan="{cols_span}">{sec_ret_str}</th>\n' if not sec_ret_str == "" else sec_ret_str
|
||||
sec_ret_str = f'<th colspan="{cols_span}">{sec_ret_str}</th>\n' if not sec_ret_str == '' else sec_ret_str
|
||||
else:
|
||||
ret_str = f'<td>{ret_str}</td>\n'
|
||||
sec_ret_str = f'<td align="center">{sec_ret_str}</td>\n' if not sec_ret_str == "" else sec_ret_str
|
||||
sec_ret_str = f'<td align="center">{sec_ret_str}</td>\n' if not sec_ret_str == '' else sec_ret_str
|
||||
return [ret_str, sec_ret_str]
|
||||
|
||||
def format_row(self):
|
||||
@@ -267,7 +261,10 @@ class HtmlVisualization(object):
|
||||
if not self.allow_annotation:
|
||||
one_row_str += '\n'.join([v[0] for v in one_content])
|
||||
else:
|
||||
one_row_str += '\n'.join([v[0].replace('#sample_id#', f'{sample_id}') for v in one_content])
|
||||
one_row_str += '\n'.join([
|
||||
v[0].replace('#sample_id#', f'{sample_id}')
|
||||
for v in one_content
|
||||
])
|
||||
row_meta = '#;#'.join(one_row_meta)
|
||||
one_row_str += (
|
||||
f'<td><input type="checkbox" class="large-checkbox" '
|
||||
@@ -282,7 +279,8 @@ class HtmlVisualization(object):
|
||||
# one_row_str = f'<label for="sample{sample_id}">{one_row_str}</label>'
|
||||
current_sample_html.append(one_row_str)
|
||||
sample_id += 1
|
||||
all_sample_html.append("<table>" + '\n'.join(current_sample_html) + "</table>")
|
||||
all_sample_html.append('<table>' + '\n'.join(current_sample_html) +
|
||||
'</table>')
|
||||
|
||||
return all_sample_html
|
||||
|
||||
@@ -304,8 +302,11 @@ class HtmlVisualization(object):
|
||||
if col_id > len(self.content_list[row_id]):
|
||||
raise RuntimeError(
|
||||
'col_id should be next number of the last col_id.')
|
||||
format_col = self.format_col(content, f"{row_id}-{col_id}: {label}",
|
||||
type, show_label=show_label, cols_span=cols_span)
|
||||
format_col = self.format_col(content,
|
||||
f"{row_id}-{col_id}: {label}",
|
||||
type,
|
||||
show_label=show_label,
|
||||
cols_span=cols_span)
|
||||
|
||||
annotation_meta = annotation_meta if annotation_meta else ''
|
||||
if col_id == len(self.content_list[row_id]):
|
||||
@@ -320,8 +321,8 @@ class HtmlVisualization(object):
|
||||
if isinstance(html_body, list) and len(html_body) > 1:
|
||||
try:
|
||||
os.makedirs(path, exist_ok=True)
|
||||
except:
|
||||
print("Create folder path failed.")
|
||||
except: # noqa
|
||||
print('Create folder path failed.')
|
||||
for html_id, one_html in enumerate(html_body):
|
||||
ret_html_list = [
|
||||
self.html_start, self.html_head, self.html_style,
|
||||
@@ -332,7 +333,8 @@ class HtmlVisualization(object):
|
||||
ret_html_list.append(self.html_script)
|
||||
ret_html_list.append(self.html_end)
|
||||
ret_html = '\n'.join(ret_html_list)
|
||||
FS.put_object(ret_html.encode(), os.path.join(path, f"{html_id}.html"))
|
||||
FS.put_object(ret_html.encode(),
|
||||
os.path.join(path, f"{html_id}.html"))
|
||||
else:
|
||||
ret_html_list = [
|
||||
self.html_start, self.html_head, self.html_style,
|
||||
|
||||
Reference in New Issue
Block a user