new project from v0.0.1
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.utils import config, distribute, file_clients, file_system
|
||||
@@ -0,0 +1,604 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import argparse
|
||||
import copy
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
|
||||
import yaml
|
||||
|
||||
from scepter.modules.utils.model import StdMsg
|
||||
|
||||
|
||||
def dict_to_yaml(module_name, name, json_config, set_name=False):
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
convert std dict to yaml
|
||||
:param module_name:
|
||||
:param json_config:
|
||||
:return:
|
||||
'''
|
||||
def convert_yaml_style(level=1,
|
||||
name='ENV',
|
||||
description='ENV PARA',
|
||||
default='',
|
||||
type_name='',
|
||||
is_sys=False):
|
||||
new_line = ''
|
||||
new_line += '{}# {} DESCRIPTION: {} TYPE: {} default: {}\n'.format(
|
||||
'\t' * (level - 1), name.upper(), description, type_name,
|
||||
f'\'{default}\'' if isinstance(default, str) else default)
|
||||
if is_sys:
|
||||
if name == '-':
|
||||
new_line += '{}{}\n'.format('\t' * (level - 1), name.upper())
|
||||
else:
|
||||
new_line += '{}{}:\n'.format('\t' * (level - 1), name.upper())
|
||||
else:
|
||||
# if isinstance(default, str):
|
||||
# default = f'\'{default}\''
|
||||
if default is None:
|
||||
new_line += '{}# {}: {}\n'.format('\t' * (level - 1),
|
||||
name.upper(), default)
|
||||
else:
|
||||
new_line += '{}{}: {}\n'.format('\t' * (level - 1),
|
||||
name.upper(), default)
|
||||
return new_line
|
||||
|
||||
def parse_dict(json_config,
|
||||
level_num,
|
||||
parent_key,
|
||||
set_name=False,
|
||||
name='',
|
||||
parent_type='dict'):
|
||||
yaml_str = ''
|
||||
# print(level_num, json_config)
|
||||
if isinstance(json_config, dict):
|
||||
if 'value' in json_config:
|
||||
value = json_config['value']
|
||||
if isinstance(value, dict):
|
||||
assert len(value) < 1
|
||||
value = None
|
||||
description = json_config.get('description', '')
|
||||
yaml_str += convert_yaml_style(level=level_num - 1,
|
||||
name=parent_key,
|
||||
description=description,
|
||||
default=value,
|
||||
type_name=type(value).__name__)
|
||||
return True, yaml_str
|
||||
else:
|
||||
if len(json_config) < 1:
|
||||
yaml_str += convert_yaml_style(level=level_num,
|
||||
name='NAME',
|
||||
description='',
|
||||
default='',
|
||||
type_name='')
|
||||
level_num += 1
|
||||
for k, v in json_config.items():
|
||||
if k == 'description':
|
||||
continue
|
||||
if isinstance(v, dict):
|
||||
is_final, new_yaml_str = parse_dict(v,
|
||||
level_num,
|
||||
k,
|
||||
parent_type='dict')
|
||||
if not is_final and parent_type == 'dict':
|
||||
description = v.get('description', '')
|
||||
yaml_str += convert_yaml_style(
|
||||
level=level_num - 1,
|
||||
name=k,
|
||||
description=description,
|
||||
default='',
|
||||
type_name='',
|
||||
is_sys=True)
|
||||
if not is_final and parent_type == 'list':
|
||||
yaml_str += convert_yaml_style(level=level_num,
|
||||
name='NAME',
|
||||
description='',
|
||||
default=k,
|
||||
type_name='')
|
||||
yaml_str += new_yaml_str
|
||||
elif isinstance(v, list):
|
||||
base_yaml_str = convert_yaml_style(level=level_num - 1,
|
||||
name=k,
|
||||
description='',
|
||||
default='',
|
||||
type_name='',
|
||||
is_sys=True)
|
||||
yaml_str += base_yaml_str
|
||||
for tup in v:
|
||||
is_final, new_yaml_str = parse_dict(
|
||||
tup, level_num, '-', parent_type='list')
|
||||
if not is_final:
|
||||
yaml_str += convert_yaml_style(level=level_num,
|
||||
name='-',
|
||||
description='',
|
||||
default='',
|
||||
type_name='',
|
||||
is_sys=True)
|
||||
yaml_str += new_yaml_str
|
||||
else:
|
||||
raise KeyError(
|
||||
f'json config {json_config} must be a dict of list'
|
||||
)
|
||||
|
||||
elif isinstance(json_config, list):
|
||||
level_num += 1
|
||||
for tup in json_config:
|
||||
is_final, new_yaml_str = parse_dict(tup, level_num, '-')
|
||||
if not is_final:
|
||||
|
||||
yaml_str += convert_yaml_style(level=level_num - 1,
|
||||
name='-',
|
||||
description='',
|
||||
default='',
|
||||
type_name='',
|
||||
is_sys=True)
|
||||
if set_name:
|
||||
yaml_str += convert_yaml_style(level=level_num,
|
||||
name='NAME',
|
||||
description='',
|
||||
default=name,
|
||||
type_name='')
|
||||
yaml_str += new_yaml_str
|
||||
else:
|
||||
raise KeyError(f'json config {json_config} must be a dict')
|
||||
return False, yaml_str
|
||||
|
||||
if isinstance(json_config, dict):
|
||||
first_dict, sec_dict, third_dict = {}, {}, {}
|
||||
for key, value in json_config.items():
|
||||
if isinstance(value, dict) and len(value) > 0:
|
||||
first_dict[key] = value
|
||||
elif isinstance(value, dict) and len(value) == 0:
|
||||
sec_dict[key] = value
|
||||
elif isinstance(value, list):
|
||||
third_dict[key] = value
|
||||
else:
|
||||
raise f'Config {json_config} is illegal'
|
||||
json_config = {}
|
||||
json_config.update(first_dict)
|
||||
json_config.update(sec_dict)
|
||||
json_config.update(third_dict)
|
||||
|
||||
yaml_str = f'[{module_name}] module yaml examples:\n'
|
||||
level_num = 1
|
||||
base_yaml_str = convert_yaml_style(level=level_num,
|
||||
name=module_name,
|
||||
description='',
|
||||
default='',
|
||||
type_name='',
|
||||
is_sys=True)
|
||||
level_num += 1
|
||||
|
||||
is_final, new_yaml_str = parse_dict(json_config,
|
||||
level_num,
|
||||
module_name,
|
||||
set_name=isinstance(json_config, list)
|
||||
and set_name,
|
||||
name=name)
|
||||
if not is_final:
|
||||
yaml_str += base_yaml_str
|
||||
if set_name and not isinstance(json_config, list):
|
||||
yaml_str += convert_yaml_style(level=level_num,
|
||||
name='NAME',
|
||||
description='',
|
||||
default=name,
|
||||
type_name='')
|
||||
yaml_str += new_yaml_str
|
||||
else:
|
||||
yaml_str += new_yaml_str[1:]
|
||||
|
||||
return yaml_str
|
||||
|
||||
|
||||
def _parse_args(parser):
|
||||
if parser is None:
|
||||
parser = argparse.ArgumentParser(
|
||||
description='Argparser for My codebase:\n')
|
||||
else:
|
||||
assert isinstance(parser, argparse.ArgumentParser)
|
||||
parser.add_argument('--cfg',
|
||||
dest='cfg_file',
|
||||
help='Path to the configuration file',
|
||||
required=False,
|
||||
default=None)
|
||||
parser.add_argument('--local_rank',
|
||||
dest='local_rank',
|
||||
help='torch distributed launch args!',
|
||||
default=0)
|
||||
parser.add_argument(
|
||||
'-l',
|
||||
'--launcher',
|
||||
dest='launcher',
|
||||
help='spawn launcher is using python scripts, torchrun launcher is '
|
||||
'using torchrun module, default is spawn!',
|
||||
default='spawn')
|
||||
|
||||
parser.add_argument('-o',
|
||||
'--data_online',
|
||||
dest='data_online',
|
||||
action='store_false',
|
||||
help='Read data from online or save local as cache. '
|
||||
'Default is from online.')
|
||||
|
||||
parser.add_argument('-s',
|
||||
'--share_storage',
|
||||
dest='share_storage',
|
||||
action='store_true',
|
||||
help='If use nas as the common cache folder, '
|
||||
'set True to avoid download conflict.')
|
||||
parser.add_argument('--debug',
|
||||
dest='debug',
|
||||
action='store_true',
|
||||
help='Swich debug mode.')
|
||||
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
class Config(object):
|
||||
def __init__(self,
|
||||
cfg_dict={},
|
||||
load=True,
|
||||
cfg_file=None,
|
||||
logger=None,
|
||||
parser_ins=None):
|
||||
'''
|
||||
support to parse json/dict/yaml_file of parameters.
|
||||
:param load: whether load parameters or not.
|
||||
:param cfg_dict: default None.
|
||||
:param cfg_level: default None, means the current cfg-level for recurrent cfg presentation.
|
||||
:param logger: logger instance for print the cfg log.
|
||||
one examples:
|
||||
import argparse
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Argparser for Cate process:\n"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--stage",
|
||||
dest="stage",
|
||||
help="Running stage!",
|
||||
default="train",
|
||||
choices=["train"]
|
||||
)
|
||||
|
||||
cfg = Config(load=True, parser_ins=parser)
|
||||
'''
|
||||
# checking that the logger exists or not
|
||||
if logger is None:
|
||||
self.logger = StdMsg(name='Config')
|
||||
else:
|
||||
self.logger = logger
|
||||
self.cfg_dict = cfg_dict
|
||||
if load:
|
||||
if cfg_file is None:
|
||||
assert parser_ins is not None
|
||||
self.args = _parse_args(parser_ins)
|
||||
self.load_from_file(self.args.cfg_file)
|
||||
# os.environ["LAUNCHER"] = self.args.launcher
|
||||
os.environ['DATA_ONLINE'] = str(self.args.data_online).lower()
|
||||
os.environ['SHARE_STORAGE'] = str(
|
||||
self.args.share_storage).lower()
|
||||
os.environ['ES_DEBUG'] = str(self.args.debug).lower()
|
||||
else:
|
||||
self.load_from_file(cfg_file)
|
||||
if 'ENV' not in self.cfg_dict:
|
||||
self.cfg_dict['ENV'] = {
|
||||
'SEED': 2023,
|
||||
'USE_PL': False,
|
||||
'BACKEND': 'nccl',
|
||||
'SYNC_BN': False,
|
||||
'CUDNN_DETERMINISTIC': True,
|
||||
'CUDNN_BENCHMARK': False
|
||||
}
|
||||
self.logger.info(
|
||||
f"ENV is not set and will use default ENV as {self.cfg_dict['ENV']}; "
|
||||
f'If want to change this value, please set them in your config.'
|
||||
)
|
||||
else:
|
||||
if 'SEED' not in self.cfg_dict['ENV']:
|
||||
self.cfg_dict['ENV']['SEED'] = 2023
|
||||
self.logger.info(
|
||||
f"SEED is not set and will use default SEED as {self.cfg_dict['ENV']['SEED']}; "
|
||||
f'If want to change this value, please set it in your config.'
|
||||
)
|
||||
os.environ['ES_SEED'] = str(self.cfg_dict['ENV']['SEED'])
|
||||
self._update_dict(self.cfg_dict)
|
||||
if load:
|
||||
self.logger.info(f'Parse cfg file as \n {self.dump()}')
|
||||
|
||||
def load_from_file(self, file_name):
|
||||
self.logger.info(f'Loading config from {file_name}')
|
||||
if file_name is None or not os.path.exists(file_name):
|
||||
self.logger.info(f'File {file_name} does not exist!')
|
||||
self.logger.warning(
|
||||
f"Cfg file is None or doesn't exist, Skip loading config from {file_name}."
|
||||
)
|
||||
return
|
||||
if file_name.endswith('.json'):
|
||||
self.cfg_dict = self._load_json(file_name)
|
||||
self.logger.info(
|
||||
f'System take {file_name} as json, because we find json in this file'
|
||||
)
|
||||
elif file_name.endswith('.yaml'):
|
||||
self.cfg_dict = self._load_yaml(file_name)
|
||||
self.logger.info(
|
||||
f'System take {file_name} as yaml, because we find yaml in this file'
|
||||
)
|
||||
else:
|
||||
self.logger.info(
|
||||
f'No config file found! Because we do not find json or yaml in --cfg {file_name}'
|
||||
)
|
||||
|
||||
def _update_dict(self, cfg_dict):
|
||||
def recur(key, elem):
|
||||
if type(elem) is dict:
|
||||
return key, Config(load=False,
|
||||
cfg_dict=elem,
|
||||
logger=self.logger)
|
||||
elif type(elem) is list:
|
||||
config_list = []
|
||||
for idx, ele in enumerate(elem):
|
||||
if type(ele) is str and ele[1:3] == 'e-':
|
||||
ele = float(ele)
|
||||
config_list.append(ele)
|
||||
elif type(ele) is str:
|
||||
config_list.append(ele)
|
||||
elif type(ele) is dict:
|
||||
config_list.append(
|
||||
Config(load=False,
|
||||
cfg_dict=ele,
|
||||
logger=self.logger))
|
||||
elif type(ele) is list:
|
||||
config_list.append(ele)
|
||||
else:
|
||||
config_list.append(ele)
|
||||
return key, config_list
|
||||
else:
|
||||
if type(elem) is str and elem[1:3] == 'e-':
|
||||
elem = float(elem)
|
||||
return key, elem
|
||||
|
||||
dic = dict(recur(k, v) for k, v in cfg_dict.items())
|
||||
self.__dict__.update(dic)
|
||||
|
||||
def _load_json(self, cfg_file):
|
||||
'''
|
||||
:param cfg_file:
|
||||
:return:
|
||||
'''
|
||||
if cfg_file is None:
|
||||
self.logger.warning(
|
||||
f'Cfg file is None, Skip loading config from {cfg_file}.')
|
||||
return {}
|
||||
file_name = cfg_file
|
||||
try:
|
||||
cfg = json.load(open(file_name, 'r'))
|
||||
except Exception as e:
|
||||
self.logger.error(f'Load json from {cfg_file} error. Message: {e}')
|
||||
sys.exit()
|
||||
return cfg
|
||||
|
||||
def _load_yaml(self, cfg_file):
|
||||
'''
|
||||
if replace some parameters from Base, You can reference the base parameters use Base.
|
||||
|
||||
:param cfg_file:
|
||||
:return:
|
||||
'''
|
||||
if cfg_file is None:
|
||||
self.logger.warning(
|
||||
f'Cfg file is None, Skip loading config from {cfg_file}.')
|
||||
return {}
|
||||
file_name = cfg_file
|
||||
try:
|
||||
with open(cfg_file, 'r') as f:
|
||||
cfg = yaml.load(f.read(), Loader=yaml.SafeLoader)
|
||||
except Exception as e:
|
||||
self.logger.error(f'Load yaml from {cfg_file} error. Message: {e}')
|
||||
sys.exit()
|
||||
if '_BASE_RUN' not in cfg.keys() and '_BASE_MODEL' not in cfg.keys(
|
||||
) and '_BASE' not in cfg.keys():
|
||||
return cfg
|
||||
|
||||
if '_BASE' in cfg.keys():
|
||||
if cfg['_BASE'][1] == '.':
|
||||
prev_count = cfg['_BASE'].count('..')
|
||||
cfg_base_file = self._path_join(
|
||||
file_name.split('/')[:(-1 - cfg['_BASE'].count('..'))] +
|
||||
cfg['_BASE'].split('/')[prev_count:])
|
||||
else:
|
||||
cfg_base_file = cfg['_BASE'].replace(
|
||||
'./', file_name.replace(file_name.split('/')[-1], ''))
|
||||
cfg_base = self._load_yaml(cfg_base_file)
|
||||
cfg = self._merge_cfg_from_base(cfg_base, cfg)
|
||||
else:
|
||||
if '_BASE_RUN' in cfg.keys():
|
||||
if cfg['_BASE_RUN'][1] == '.':
|
||||
prev_count = cfg['_BASE_RUN'].count('..')
|
||||
cfg_base_file = self._path_join(
|
||||
file_name.split('/')[:(-1 - prev_count)] +
|
||||
cfg['_BASE_RUN'].split('/')[prev_count:])
|
||||
else:
|
||||
cfg_base_file = cfg['_BASE_RUN'].replace(
|
||||
'./', file_name.replace(file_name.split('/')[-1], ''))
|
||||
cfg_base = self._load_yaml(cfg_base_file)
|
||||
cfg = self._merge_cfg_from_base(cfg_base,
|
||||
cfg,
|
||||
preserve_base=True)
|
||||
if '_BASE_MODEL' in cfg.keys():
|
||||
if cfg['_BASE_MODEL'][1] == '.':
|
||||
prev_count = cfg['_BASE_MODEL'].count('..')
|
||||
cfg_base_file = self._path_join(
|
||||
file_name.split('/')[:(
|
||||
-1 - cfg['_BASE_MODEL'].count('..'))] +
|
||||
cfg['_BASE_MODEL'].split('/')[prev_count:])
|
||||
else:
|
||||
cfg_base_file = cfg['_BASE_MODEL'].replace(
|
||||
'./', file_name.replace(file_name.split('/')[-1], ''))
|
||||
cfg_base = self._load_yaml(cfg_base_file)
|
||||
cfg = self._merge_cfg_from_base(cfg_base, cfg)
|
||||
return cfg
|
||||
|
||||
def _path_join(self, path_list):
|
||||
path = ''
|
||||
for p in path_list:
|
||||
path += p + '/'
|
||||
return path[:-1]
|
||||
|
||||
def items(self):
|
||||
return self.cfg_dict.items()
|
||||
|
||||
def _merge_cfg_from_base(self, cfg_base, cfg, preserve_base=False):
|
||||
for k, v in cfg.items():
|
||||
if k in cfg_base.keys():
|
||||
if isinstance(v, dict):
|
||||
self._merge_cfg_from_base(cfg_base[k], v)
|
||||
else:
|
||||
cfg_base[k] = v
|
||||
else:
|
||||
if 'BASE' not in k or preserve_base:
|
||||
cfg_base[k] = v
|
||||
return cfg_base
|
||||
|
||||
def _merge_cfg_from_command(self, args, cfg):
|
||||
assert len(
|
||||
args.opts
|
||||
) % 2 == 0, f'Override list {args.opts} has odd length: {len(args.opts)}'
|
||||
|
||||
keys = args.opts[0::2]
|
||||
vals = args.opts[1::2]
|
||||
|
||||
# maximum supported depth 3
|
||||
for idx, key in enumerate(keys):
|
||||
key_split = key.split('.')
|
||||
assert len(
|
||||
key_split
|
||||
) <= 4, 'Key depth error. \n Maximum depth: 3\n Get depth: {}'.format(
|
||||
len(key_split))
|
||||
assert key_split[0] in cfg.keys(), 'Non-existant key: {}.'.format(
|
||||
key_split[0])
|
||||
if len(key_split) == 2:
|
||||
assert key_split[1] in cfg[
|
||||
key_split[0]].keys(), 'Non-existant key: {}'.format(key)
|
||||
elif len(key_split) == 3:
|
||||
assert key_split[1] in cfg[
|
||||
key_split[0]].keys(), 'Non-existant key: {}'.format(key)
|
||||
assert key_split[2] in cfg[key_split[0]][
|
||||
key_split[1]].keys(), 'Non-existant key: {}'.format(key)
|
||||
elif len(key_split) == 4:
|
||||
assert key_split[1] in cfg[
|
||||
key_split[0]].keys(), 'Non-existant key: {}'.format(key)
|
||||
assert key_split[2] in cfg[key_split[0]][
|
||||
key_split[1]].keys(), 'Non-existant key: {}'.format(key)
|
||||
assert key_split[3] in cfg[key_split[0]][key_split[1]][
|
||||
key_split[2]].keys(), 'Non-existant key: {}'.format(key)
|
||||
|
||||
if len(key_split) == 1:
|
||||
cfg[key_split[0]] = vals[idx]
|
||||
elif len(key_split) == 2:
|
||||
cfg[key_split[0]][key_split[1]] = vals[idx]
|
||||
elif len(key_split) == 3:
|
||||
cfg[key_split[0]][key_split[1]][key_split[2]] = vals[idx]
|
||||
elif len(key_split) == 4:
|
||||
cfg[key_split[0]][key_split[1]][key_split[2]][
|
||||
key_split[3]] = vals[idx]
|
||||
|
||||
return cfg
|
||||
|
||||
def __repr__(self):
|
||||
return '{}\n'.format(self.dump())
|
||||
|
||||
def dump(self):
|
||||
return json.dumps(self.cfg_dict, indent=2)
|
||||
|
||||
def deep_copy(self):
|
||||
return copy.deepcopy(self)
|
||||
|
||||
def have(self, name):
|
||||
if name in self.__dict__:
|
||||
return True
|
||||
return False
|
||||
|
||||
def get(self, name, default=None):
|
||||
if name in self.__dict__:
|
||||
return self.__dict__[name]
|
||||
return default
|
||||
|
||||
def __getitem__(self, key):
|
||||
return self.__dict__.__getitem__(key)
|
||||
|
||||
def __setattr__(self, key, value):
|
||||
super().__setattr__(key, value)
|
||||
if hasattr(self, 'cfg_dict') and key in self.cfg_dict:
|
||||
if isinstance(value, Config):
|
||||
value = value.cfg_dict
|
||||
self.cfg_dict[key] = value
|
||||
|
||||
def __setitem__(self, key, value):
|
||||
self.__dict__[key] = value
|
||||
self.__setattr__(key, value)
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self.__dict__)
|
||||
|
||||
def set(self, name, value):
|
||||
new_dict = {name: value}
|
||||
self.__dict__.update(new_dict)
|
||||
self.__setattr__(name, value)
|
||||
|
||||
def get_dict(self):
|
||||
return self.cfg_dict
|
||||
|
||||
def get_lowercase_dict(self, cfg_dict=None):
|
||||
if cfg_dict is None:
|
||||
cfg_dict = self.get_dict()
|
||||
config_new = {}
|
||||
for key, val in cfg_dict.items():
|
||||
if isinstance(key, str):
|
||||
if isinstance(val, dict):
|
||||
config_new[key.lower()] = self.get_lowercase_dict(val)
|
||||
else:
|
||||
config_new[key.lower()] = val
|
||||
else:
|
||||
config_new[key] = val
|
||||
return config_new
|
||||
|
||||
@staticmethod
|
||||
def get_plain_cfg(cfg=None):
|
||||
if isinstance(cfg, Config):
|
||||
cfg_new = {}
|
||||
cfg_dict = cfg.get_dict()
|
||||
for key, val in cfg_dict.items():
|
||||
if isinstance(val, (Config, dict, list)):
|
||||
cfg_new[key] = Config.get_plain_cfg(val)
|
||||
else:
|
||||
cfg_new[key] = val
|
||||
return cfg_new
|
||||
elif isinstance(cfg, dict):
|
||||
cfg_new = {}
|
||||
cfg_dict = cfg
|
||||
for key, val in cfg_dict.items():
|
||||
if isinstance(val, (Config, dict, list)):
|
||||
cfg_new[key] = Config.get_plain_cfg(val)
|
||||
else:
|
||||
cfg_new[key] = val
|
||||
return cfg_new
|
||||
elif isinstance(cfg, list):
|
||||
cfg_new = []
|
||||
cfg_list = cfg
|
||||
for val in cfg_list:
|
||||
if isinstance(val, (Config, dict, list)):
|
||||
cfg_new.append(Config.get_plain_cfg(val))
|
||||
else:
|
||||
cfg_new.append(val)
|
||||
return cfg_new
|
||||
else:
|
||||
return cfg
|
||||
@@ -0,0 +1,88 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def transfer_data_to_numpy(data_map: dict) -> dict:
|
||||
""" Transfer tensors in data_map to numpy type.
|
||||
Will recursively walk through inner list, tuple and dict values.
|
||||
|
||||
Args:
|
||||
data_map (dict): a dictionary which contains tensors to be transferred
|
||||
|
||||
Returns:
|
||||
A dict which has same structure with input `data_map`.
|
||||
"""
|
||||
if not isinstance(data_map, dict):
|
||||
return data_map
|
||||
ret = OrderedDict()
|
||||
for key, value in data_map.items():
|
||||
if isinstance(value, torch.Tensor):
|
||||
ret[key] = value.detach().cpu().numpy()
|
||||
elif isinstance(value, dict):
|
||||
ret[key] = transfer_data_to_numpy(value)
|
||||
elif isinstance(value, (list, tuple)):
|
||||
ret[key] = type(value)([transfer_data_to_numpy(t) for t in value])
|
||||
else:
|
||||
ret[key] = value
|
||||
return ret
|
||||
|
||||
|
||||
def transfer_data_to_cpu(data_map: dict) -> dict:
|
||||
""" Transfer tensors in data_map to cpu device.
|
||||
Will recursively walk through inner list, tuple and dict values.
|
||||
|
||||
Args:
|
||||
data_map (dict): a dictionary which contains tensors to be transferred
|
||||
|
||||
Returns:
|
||||
A dict which has same structure with input `data_map`.
|
||||
"""
|
||||
if not isinstance(data_map, dict):
|
||||
return data_map
|
||||
ret = OrderedDict()
|
||||
for key, value in data_map.items():
|
||||
if isinstance(value, torch.Tensor):
|
||||
ret[key] = value.detach().cpu()
|
||||
elif isinstance(value, dict):
|
||||
ret[key] = transfer_data_to_cpu(value)
|
||||
elif isinstance(value, (list, tuple)):
|
||||
ret[key] = type(value)([transfer_data_to_cpu(t) for t in value])
|
||||
else:
|
||||
ret[key] = value
|
||||
torch.cuda.empty_cache()
|
||||
return ret
|
||||
|
||||
|
||||
def transfer_data_to_cuda(data_map: dict) -> dict:
|
||||
""" Transfer tensors in data_map to current default gpu device.
|
||||
Will recursively walk through inner list, tuple and dict values.
|
||||
|
||||
Args:
|
||||
data_map (dict): a dictionary which contains tensors to be transferred
|
||||
|
||||
Returns:
|
||||
A dict which has same structure with input `data_map`.
|
||||
"""
|
||||
import platform
|
||||
if platform.system() == 'Darwin':
|
||||
return data_map
|
||||
if not isinstance(data_map, dict):
|
||||
return data_map
|
||||
ret = OrderedDict()
|
||||
for key, value in data_map.items():
|
||||
if isinstance(value, torch.Tensor):
|
||||
if value.is_cuda:
|
||||
ret[key] = value
|
||||
else:
|
||||
ret[key] = value.cuda(non_blocking=True)
|
||||
elif isinstance(value, dict):
|
||||
ret[key] = transfer_data_to_cuda(value)
|
||||
elif isinstance(value, (list, tuple)):
|
||||
ret[key] = type(value)([transfer_data_to_cuda(t) for t in value])
|
||||
else:
|
||||
ret[key] = value
|
||||
return ret
|
||||
@@ -0,0 +1,21 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import hashlib
|
||||
import os.path as osp
|
||||
|
||||
|
||||
def osp_path(prefix, data_file):
|
||||
if data_file.startswith(prefix):
|
||||
return data_file
|
||||
else:
|
||||
return osp.join(prefix, data_file)
|
||||
|
||||
|
||||
def get_relative_folder(abs_path, keep_index=-1):
|
||||
path_tup = abs_path.split('/')[:keep_index]
|
||||
return '/'.join(path_tup)
|
||||
|
||||
|
||||
def get_md5(ori_str):
|
||||
md5 = hashlib.md5(ori_str.encode()).hexdigest()
|
||||
return md5
|
||||
@@ -0,0 +1,455 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import functools
|
||||
import os
|
||||
import pickle
|
||||
import random
|
||||
import warnings
|
||||
from collections import OrderedDict
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from scepter.modules.utils.model import StdMsg
|
||||
|
||||
__all__ = [
|
||||
'gather_data', 'we', 'broadcast', 'barrier', 'reduce_scatter', 'reduce',
|
||||
'all_reduce', 'send', 'recv', 'isend', 'irecv', 'scatter',
|
||||
'shared_random_seed'
|
||||
]
|
||||
|
||||
try:
|
||||
from onnxruntime.transformers.benchmark_helper import set_random_seed
|
||||
except Exception:
|
||||
|
||||
def set_random_seed(seed):
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
np.random.seed(seed)
|
||||
random.seed(seed)
|
||||
|
||||
|
||||
def get_dist_info():
|
||||
if dist.is_available() and dist.is_initialized():
|
||||
return dist.get_rank(), dist.get_world_size()
|
||||
else:
|
||||
return 0, 1
|
||||
|
||||
|
||||
def gather_data(data):
|
||||
""" Gather tensors and other picklable objects to rank 0.
|
||||
Will recursively walk through inner list and dict values.
|
||||
|
||||
Args:
|
||||
data (any): Anything.
|
||||
|
||||
Returns:
|
||||
A object has same structure with input `data`.
|
||||
"""
|
||||
if not we.is_distributed:
|
||||
return data
|
||||
if isinstance(data, torch.Tensor):
|
||||
return gather_gpu_tensors(data)
|
||||
elif isinstance(data, dict):
|
||||
# Keep in order, dict type DO NOT guarantee a fixed key order
|
||||
keys = sorted(list(data.keys()))
|
||||
ret = OrderedDict()
|
||||
for key in keys:
|
||||
ret[key] = gather_data(data[key])
|
||||
return ret
|
||||
elif isinstance(data, list):
|
||||
return gather_list(data)
|
||||
else:
|
||||
return gather_picklable(data)
|
||||
|
||||
|
||||
def gather_list(data):
|
||||
""" Gather list of picklable objects to a new list on rank 0.
|
||||
Will NOT recursively walk through.
|
||||
|
||||
Args:
|
||||
data (list): List of picklable things.
|
||||
|
||||
Returns:
|
||||
A new flat list.
|
||||
"""
|
||||
rank, _ = get_dist_info()
|
||||
list_of_list = gather_picklable(data)
|
||||
if rank == 0:
|
||||
return sum(list_of_list, [])
|
||||
|
||||
|
||||
def gather_picklable(data):
|
||||
""" Gather picklable object to a list on rank 0.
|
||||
Will NOT recursively walk through.
|
||||
|
||||
Args:
|
||||
data (picklable): Picklable data.
|
||||
|
||||
Returns:
|
||||
A list contains data collected.
|
||||
"""
|
||||
from packaging import version
|
||||
from torch.version import __version__
|
||||
if version.parse(__version__) < version.parse('1.8.0'):
|
||||
return _gather_picklable_custom(data)
|
||||
else:
|
||||
rank, world_size = we.rank, we.world_size
|
||||
obj_list = [None for _ in range(world_size)]
|
||||
dist.all_gather_object(obj_list, data)
|
||||
if rank == 0:
|
||||
return obj_list
|
||||
|
||||
|
||||
def _gather_picklable_custom(data):
|
||||
""" Custom implementation function to gather picklable object to a list on rank 0.
|
||||
If torch version is lower than 1.8.0, use this.
|
||||
|
||||
Args:
|
||||
data (picklable): Picklable data.
|
||||
|
||||
Returns:
|
||||
A list contains data collected.
|
||||
"""
|
||||
import pickle
|
||||
byte_tensor = torch.tensor(bytearray(pickle.dumps(data)),
|
||||
dtype=torch.uint8,
|
||||
device='cuda')
|
||||
rank, world_size = we.rank, we.world_size
|
||||
shape_tensor = torch.tensor(byte_tensor.shape, device='cuda')
|
||||
shape_list = [shape_tensor.clone() for _ in range(world_size)]
|
||||
dist.all_gather(shape_list, shape_tensor)
|
||||
shape_max = torch.tensor(shape_list).max()
|
||||
|
||||
tensor_send = torch.zeros(shape_max,
|
||||
dtype=byte_tensor.dtype,
|
||||
device='cuda')
|
||||
tensor_send[0:shape_tensor[0]] = byte_tensor
|
||||
tensor_list = [torch.zeros_like(tensor_send) for _ in range(world_size)]
|
||||
dist.all_gather(tensor_list, tensor_send)
|
||||
|
||||
if rank == 0:
|
||||
data_out = []
|
||||
for tensor_recv, shape_recv in zip(tensor_list, shape_list):
|
||||
new_data = pickle.loads(
|
||||
tensor_recv[:shape_recv[0]].cpu().numpy().tobytes())
|
||||
data_out.append(new_data)
|
||||
return data_out
|
||||
|
||||
|
||||
def gather_gpu_tensors(tensor, all_recv=False, is_cat=True):
|
||||
"""
|
||||
Args:
|
||||
tensor (torch.Tensor):
|
||||
all_recv: Gather tensor to rank 0 and concat it.
|
||||
|
||||
Returns:
|
||||
A new tensor.
|
||||
"""
|
||||
assert dist.get_backend() == 'nccl'
|
||||
|
||||
device = tensor.device
|
||||
if device.type == 'cpu':
|
||||
tensor = tensor.to(we.device_id)
|
||||
|
||||
rank, world_size = we.rank, we.world_size
|
||||
|
||||
shape_tensor = torch.tensor(tensor.shape[0], device='cuda')
|
||||
shape_list = [shape_tensor.clone() for _ in range(world_size)]
|
||||
dist.all_gather(shape_list, shape_tensor)
|
||||
shape_max = torch.tensor(shape_list).max()
|
||||
|
||||
tensor_send = torch.zeros((shape_max, *tensor.shape[1:]),
|
||||
dtype=tensor.dtype,
|
||||
device='cuda')
|
||||
tensor_send[0:tensor.shape[0]] = tensor
|
||||
tensor_list = [torch.zeros_like(tensor_send) for _ in range(world_size)]
|
||||
dist.all_gather(tensor_list, tensor_send)
|
||||
if not all_recv:
|
||||
if rank == 0:
|
||||
if not is_cat:
|
||||
return tensor_list, shape_list
|
||||
tensors_out = []
|
||||
for tensor_recv, shape_recv in zip(tensor_list, shape_list):
|
||||
tensors_out.append(tensor_recv[0:shape_recv])
|
||||
tensor_out = torch.cat(tensors_out).contiguous()
|
||||
if device.type == 'cpu':
|
||||
tensor_out = tensor_out.cpu()
|
||||
del tensor_list, shape_list
|
||||
return tensor_out
|
||||
else:
|
||||
del tensor_list, shape_list
|
||||
else:
|
||||
if not is_cat:
|
||||
return tensor_list, shape_list
|
||||
tensors_out = []
|
||||
for tensor_recv, shape_recv in zip(tensor_list, shape_list):
|
||||
tensors_out.append(tensor_recv[0:shape_recv])
|
||||
tensor_out = torch.cat(tensors_out).contiguous()
|
||||
if device.type == 'cpu':
|
||||
tensor_out = tensor_out.cpu()
|
||||
del tensor_list, shape_list
|
||||
return tensor_out
|
||||
|
||||
|
||||
def broadcast(tensor, src, group=None, **kwargs):
|
||||
if we.is_distributed:
|
||||
return dist.broadcast(tensor, src, group, **kwargs)
|
||||
|
||||
|
||||
def barrier():
|
||||
if we.is_distributed:
|
||||
dist.barrier()
|
||||
|
||||
|
||||
@functools.lru_cache()
|
||||
def get_global_gloo_group():
|
||||
backend = dist.get_backend()
|
||||
assert backend in ['gloo', 'nccl']
|
||||
if backend == 'nccl':
|
||||
return dist.new_group(backend='gloo')
|
||||
else:
|
||||
return dist.group.WORLD
|
||||
|
||||
|
||||
def reduce_scatter(output,
|
||||
input_list,
|
||||
op=dist.ReduceOp.SUM,
|
||||
group=None,
|
||||
**kwargs):
|
||||
if we.is_distributed:
|
||||
return dist.reduce_scatter(output, input_list, op, group, **kwargs)
|
||||
|
||||
|
||||
def all_reduce(tensor, op=dist.ReduceOp.SUM, group=None, **kwargs):
|
||||
if we.is_distributed:
|
||||
return dist.all_reduce(tensor, op, group, **kwargs)
|
||||
|
||||
|
||||
def reduce(tensor, dst, op=dist.ReduceOp.SUM, group=None, **kwargs):
|
||||
if we.is_distributed:
|
||||
return dist.reduce(tensor, dst, op, group, **kwargs)
|
||||
|
||||
|
||||
def _serialize_to_tensor(data):
|
||||
buffer = pickle.dumps(data)
|
||||
storage = torch.ByteStorage.from_buffer(buffer)
|
||||
tensor = torch.ByteTensor(storage)
|
||||
return tensor
|
||||
|
||||
|
||||
def _unserialize_from_tensor(recv_data):
|
||||
buffer = recv_data.cpu().numpy().tobytes()
|
||||
return pickle.loads(buffer)
|
||||
|
||||
|
||||
def send(tensor, dst, group=None, **kwargs):
|
||||
if we.is_distributed:
|
||||
assert tensor.is_contiguous(
|
||||
), 'ops.send requires the tensor to be contiguous()'
|
||||
return dist.send(tensor, dst, group, **kwargs)
|
||||
|
||||
|
||||
def recv(tensor, src=None, group=None, **kwargs):
|
||||
if we.is_distributed:
|
||||
assert tensor.is_contiguous(
|
||||
), 'ops.recv requires the tensor to be contiguous()'
|
||||
return dist.recv(tensor, src, group, **kwargs)
|
||||
|
||||
|
||||
def isend(tensor, dst, group=None, **kwargs):
|
||||
if we.is_distributed:
|
||||
assert tensor.is_contiguous(
|
||||
), 'ops.isend requires the tensor to be contiguous()'
|
||||
return dist.isend(tensor, dst, group, **kwargs)
|
||||
|
||||
|
||||
def irecv(tensor, src=None, group=None, **kwargs):
|
||||
if we.is_distributed:
|
||||
assert tensor.is_contiguous(
|
||||
), 'ops.irecv requires the tensor to be contiguous()'
|
||||
return dist.irecv(tensor, src, group, **kwargs)
|
||||
|
||||
|
||||
def scatter(data, scatter_list=None, src=0, group=None, **kwargs):
|
||||
r"""NOTE: only supports CPU tensor communication.
|
||||
"""
|
||||
world_size = we.world_size
|
||||
if world_size == 1:
|
||||
data.copy_(scatter_list[0])
|
||||
if group is None:
|
||||
group = get_global_gloo_group()
|
||||
return dist.scatter(data, scatter_list, src, group, **kwargs)
|
||||
|
||||
|
||||
def shared_random_seed():
|
||||
seed = np.random.randint(2**31)
|
||||
all_seeds, _ = gather_gpu_tensors(seed, all_recv=True, is_cat=False)
|
||||
return all_seeds[0]
|
||||
|
||||
|
||||
global we
|
||||
|
||||
|
||||
def mp_worker(gpu, ngpus_per_node, cfg, fn, pmi_rank, world_size, work_env):
|
||||
rank = pmi_rank * ngpus_per_node + gpu
|
||||
work_env.device_id = gpu % ngpus_per_node
|
||||
work_env.rank = rank
|
||||
dist.init_process_group(backend='nccl', world_size=world_size, rank=rank)
|
||||
torch.backends.cudnn.deterministic = cfg.ENV.get('CUDNN_DETERMINISTIC',
|
||||
True)
|
||||
torch.backends.cudnn.benchmark = cfg.ENV.get('CUDNN_BENCHMARK', False)
|
||||
torch.cuda.set_device(work_env.device_id)
|
||||
if work_env.logger is not None:
|
||||
work_env.logger.info(
|
||||
f'Now running in the distributed environment with world size {work_env.world_size}!'
|
||||
)
|
||||
work_env.logger.info(f'PMI rank {pmi_rank}!')
|
||||
work_env.logger.info(f'Nums of gpu {ngpus_per_node}!')
|
||||
work_env.logger.info(
|
||||
f'Current rank {work_env.rank} current devices num {ngpus_per_node} '
|
||||
f'current machine rank {pmi_rank} and all world size {world_size}')
|
||||
|
||||
we.set_env(work_env.get_env())
|
||||
fn(cfg)
|
||||
|
||||
|
||||
class Workenv(object):
|
||||
def __init__(self):
|
||||
self.initialized = False
|
||||
self.is_distributed = False
|
||||
self.sync_bn = False
|
||||
self.rank = 0
|
||||
self.world_size = 1
|
||||
self.device_id = 0
|
||||
self.device_count = 1
|
||||
self.seed = 2023
|
||||
self.debug = False
|
||||
self.use_pl = False
|
||||
self.launcher = 'spawn'
|
||||
self.data_online = False
|
||||
self.share_storage = False
|
||||
|
||||
def init_env(self, config, fn, logger=None):
|
||||
# if use pytorch_lightning: then direct use pytorch_lightning.
|
||||
config.ENV = config.get('ENV', {})
|
||||
self.seed = config.ENV.get('SEED', 2023)
|
||||
self.debug = os.environ.get('ES_DEBUG', None) == 'true'
|
||||
set_random_seed(self.seed)
|
||||
if logger is not None:
|
||||
logger.info(f'And running with seed {self.seed}!')
|
||||
if config.ENV.get('USE_PL', False):
|
||||
self.use_pl = config.ENV.USE_PL
|
||||
fn(config)
|
||||
return
|
||||
if hasattr(config, 'args') and hasattr(config.args, 'launcher'):
|
||||
self.launcher = config.args.launcher
|
||||
if logger is None:
|
||||
self.logger = StdMsg(name='env')
|
||||
else:
|
||||
self.logger = logger
|
||||
|
||||
self.data_online = os.environ.get('DATA_ONLINE', None) == 'true'
|
||||
self.share_storage = os.environ.get('SHARE_STORAGE', None) == 'true'
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
self.device_id = 'cpu'
|
||||
fn(config)
|
||||
return
|
||||
|
||||
if (os.environ.get('WORLD_SIZE') is None or os.environ.get('WORLD_SIZE') == 1) \
|
||||
and torch.cuda.device_count() == 1 and not self.launcher == 'dist':
|
||||
self.device_id = 0
|
||||
fn(config)
|
||||
return
|
||||
|
||||
if self.launcher == 'torchrun':
|
||||
try:
|
||||
torch.multiprocessing.set_start_method('spawn')
|
||||
except Exception as e:
|
||||
warnings.warn(f'{e}')
|
||||
# checking train mode is distributed or not
|
||||
if not os.environ.get('WORLD_SIZE') is None:
|
||||
if self.logger is not None:
|
||||
self.logger.info(
|
||||
f"Now running in the distributed environment with {os.environ.get('WORLD_SIZE')}!"
|
||||
)
|
||||
self.is_distributed = True
|
||||
if not self.initialized:
|
||||
if self.is_distributed:
|
||||
self.backend = config.ENV.get('BACKEND', 'nccl')
|
||||
self.sync_bn = config.ENV.get('SYNC_BN', False)
|
||||
dist.init_process_group(backend=self.backend)
|
||||
# dist.barrier()
|
||||
self.initialized = True
|
||||
if dist.is_initialized():
|
||||
self.rank, self.world_size = dist.get_rank(
|
||||
), dist.get_world_size()
|
||||
if self.logger is not None:
|
||||
self.logger.info(f'And running in rank {self.rank}!')
|
||||
self.logger.info(
|
||||
f"And cuda visible devices {os.environ.get('CUDA_VISIBLE_DEVICES')}!"
|
||||
)
|
||||
else:
|
||||
self.rank, self.world_size = 0, 1
|
||||
local_devices = os.environ.get(
|
||||
'LOCAL_WORLD_SIZE') or torch.cuda.device_count()
|
||||
local_devices = int(local_devices)
|
||||
self.device_count = local_devices
|
||||
self.device_id = self.rank % local_devices
|
||||
self.logger.info(f"We's attributes: \n"
|
||||
f' launcher {self.launcher} \n'
|
||||
f' rank {self.rank} \n'
|
||||
f' world size {self.world_size} \n'
|
||||
f' device_id {self.device_id}')
|
||||
torch.cuda.set_device(self.device_id)
|
||||
torch.backends.cudnn.deterministic = config.ENV.get(
|
||||
'CUDNN_DETERMINISTIC', True)
|
||||
torch.backends.cudnn.benchmark = config.ENV.get(
|
||||
'CUDNN_BENCHMARK', False)
|
||||
fn(config)
|
||||
else:
|
||||
import torch.multiprocessing as mp
|
||||
if 'MASTER_ADDR' not in os.environ:
|
||||
os.environ['MASTER_ADDR'] = 'localhost'
|
||||
if 'MASTER_PORT' not in os.environ:
|
||||
os.environ['MASTER_PORT'] = '14567'
|
||||
pmi_rank = int(os.environ.get('RANK', 0))
|
||||
pmi_world_size = int(os.environ.get('WORLD_SIZE', 1))
|
||||
ngpus_per_node = os.environ.get(
|
||||
'LOCAL_WORLD_SIZE') or torch.cuda.device_count()
|
||||
ngpus_per_node = int(ngpus_per_node)
|
||||
self.device_count = ngpus_per_node
|
||||
world_size = ngpus_per_node * pmi_world_size
|
||||
self.world_size = world_size
|
||||
if self.world_size > 1:
|
||||
self.is_distributed = True
|
||||
self.initialized = True
|
||||
if self.is_distributed:
|
||||
self.backend = config.ENV.get('BACKEND', 'nccl')
|
||||
self.sync_bn = config.ENV.get('SYNC_BN', False)
|
||||
mp.spawn(mp_worker,
|
||||
nprocs=ngpus_per_node,
|
||||
args=(ngpus_per_node, config, fn, pmi_rank, world_size,
|
||||
self))
|
||||
|
||||
def get_env(self):
|
||||
return self.__dict__
|
||||
|
||||
def set_env(self, we_env):
|
||||
for k, v in we_env.items():
|
||||
setattr(self, k, v)
|
||||
|
||||
def __str__(self):
|
||||
environ_str = f'Now running in the distributed environment with world size {self.world_size}\n!'
|
||||
environ_str += f'Current pod have {self.device_count} devices!\n'
|
||||
environ_str += f'Current task executes on device {self.device_id}!\n'
|
||||
environ_str += f"Current task's global rank is {self.rank} \n"
|
||||
environ_str += f"Current task's data online is set {self.data_online}"
|
||||
environ_str += f"Current task's share storage is set {self.share_storage}"
|
||||
environ_str += f"Current task's global seed is set {self.seed}"
|
||||
return environ_str
|
||||
|
||||
|
||||
we = Workenv()
|
||||
@@ -0,0 +1,113 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import io
|
||||
from io import BytesIO
|
||||
|
||||
import onnx
|
||||
import onnxruntime
|
||||
import torch
|
||||
from torch.onnx import OperatorExportTypes
|
||||
|
||||
from scepter.modules.utils.distribute import we
|
||||
|
||||
type_map = {
|
||||
'float32': torch.float32,
|
||||
'float16': torch.float16,
|
||||
'int64': torch.int64,
|
||||
'int32': torch.int32,
|
||||
'int16': torch.int16,
|
||||
'int8': torch.int8
|
||||
}
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def save_develop_model_multi_io(model,
|
||||
input_size,
|
||||
input_type,
|
||||
input_name,
|
||||
output_name,
|
||||
limit,
|
||||
save_onnx_path=None,
|
||||
save_pt_path=None):
|
||||
|
||||
# save aggregation
|
||||
rank, word_size = we.rank, we.world_size
|
||||
assert isinstance(input_type, list)
|
||||
example = []
|
||||
dynamic_axes = {}
|
||||
|
||||
for idx, type_name in enumerate(input_type):
|
||||
assert type_name in type_map
|
||||
torch_type = type_map[type_name]
|
||||
size = input_size[idx]
|
||||
if 'float' in type_name:
|
||||
input_ex = torch.rand(tuple(size)).type(torch_type).to(rank)
|
||||
elif 'int' in type_name:
|
||||
input_ex = torch.randint(limit[idx][0], limit[idx][1],
|
||||
tuple(size)).type(torch_type).to(rank)
|
||||
example.append(input_ex)
|
||||
dynamic_axes[input_name[idx]] = {0: 'batch_size'}
|
||||
|
||||
if word_size > 0:
|
||||
save_module = model.module
|
||||
else:
|
||||
save_module = model
|
||||
|
||||
def _check_eval(module):
|
||||
assert not module.training
|
||||
|
||||
save_module.apply(_check_eval)
|
||||
|
||||
if len(example) == 1:
|
||||
input_example = example[0]
|
||||
else:
|
||||
input_example = tuple(example)
|
||||
traced_script_module = torch.jit.trace(save_module, input_example)
|
||||
|
||||
for p in traced_script_module.parameters():
|
||||
p.requires_grad = False
|
||||
if len(example) == 1:
|
||||
output = save_module(input_example)
|
||||
else:
|
||||
output = save_module(*input_example)
|
||||
print('Ori output:', output)
|
||||
|
||||
module = None
|
||||
if save_pt_path is not None:
|
||||
traced_script_module.save(save_pt_path)
|
||||
module = torch.jit.load(io.BytesIO(open(save_pt_path, 'rb').read()),
|
||||
map_location=torch.device(rank))
|
||||
if len(example) == 1:
|
||||
output = module(input_example)
|
||||
else:
|
||||
output = module(*input_example)
|
||||
print('PT output:', output)
|
||||
|
||||
onnx_module = None
|
||||
if save_onnx_path is not None:
|
||||
# export the model to ONNX
|
||||
with torch.autocast(device_type='cpu',
|
||||
enabled=True,
|
||||
dtype=torch.bfloat16):
|
||||
with BytesIO() as f:
|
||||
torch.onnx.export(
|
||||
save_module,
|
||||
input_example,
|
||||
f,
|
||||
operator_export_type=OperatorExportTypes.ONNX,
|
||||
opset_version=11,
|
||||
input_names=input_name,
|
||||
output_names=output_name,
|
||||
dynamic_axes=dynamic_axes,
|
||||
export_params=True,
|
||||
do_constant_folding=True)
|
||||
onnx_model = onnx.load_from_string(f.getvalue())
|
||||
onnx.save(onnx_model, save_onnx_path)
|
||||
onnx_module = onnxruntime.InferenceSession(
|
||||
save_onnx_path, providers=['CUDAExecutionProvider'])
|
||||
input_data = {}
|
||||
for idx, ex in enumerate(example):
|
||||
input_data[input_name[idx]] = ex.detach().cpu().numpy()
|
||||
output_tensor = onnx_module.run(output_name, input_data)
|
||||
print('ONNX_OUTPUT', output_tensor, output_tensor[0].shape)
|
||||
return module, onnx_module
|
||||
@@ -0,0 +1,6 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.utils.file_clients.aliyun_oss_fs import AliyunOssFs
|
||||
from scepter.modules.utils.file_clients.http_fs import HttpFs
|
||||
from scepter.modules.utils.file_clients.local_fs import LocalFs
|
||||
from scepter.modules.utils.file_clients.modelscope_fs import ModelscopeFs
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,383 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import datetime
|
||||
import os
|
||||
import os.path as osp
|
||||
import random
|
||||
import tempfile
|
||||
import warnings
|
||||
from abc import ABCMeta, abstractmethod
|
||||
from copy import copy
|
||||
from typing import Optional
|
||||
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.directory import get_md5
|
||||
from scepter.modules.utils.file_clients.utils import remove_temp_path
|
||||
from scepter.modules.utils.logger import get_logger
|
||||
|
||||
|
||||
class BaseFs(object, metaclass=ABCMeta):
|
||||
para_dict = {
|
||||
'TEMP_DIR': {
|
||||
'value':
|
||||
None,
|
||||
'description':
|
||||
'default is None, means using system cache dir and auto remove! If you set dir, the data will'
|
||||
' be saved in this temp dir without autoremoving default.'
|
||||
},
|
||||
'AUTO_CLEAN': {
|
||||
'value':
|
||||
False,
|
||||
'description':
|
||||
'when TEMP_DIR is not None, if you set AUTO_CLEAN to True, the data will be clean automatics.'
|
||||
}
|
||||
}
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
self._target_local_mapper = {}
|
||||
self._temp_files = set()
|
||||
self.cfg = cfg
|
||||
self.tmp_dir = cfg.get('TEMP_DIR', None)
|
||||
self.auto_clean = cfg.get('AUTO_CLEAN', False)
|
||||
if self.tmp_dir is None:
|
||||
self.auto_clean = True
|
||||
if self.tmp_dir is not None:
|
||||
if not os.path.exists(self.tmp_dir):
|
||||
try:
|
||||
os.makedirs(self.tmp_dir, exist_ok=True)
|
||||
except Exception as e:
|
||||
warnings.warn(
|
||||
f'Create cache folder failed use default cache file{e}!'
|
||||
.format(self.tmp_dir))
|
||||
self.tmp_dir = None
|
||||
# checking that the logger exists or not
|
||||
if logger is None:
|
||||
self.logger = get_logger(name='File System')
|
||||
else:
|
||||
self.logger = logger
|
||||
|
||||
# Functions without io
|
||||
@abstractmethod
|
||||
def get_prefix(self) -> str:
|
||||
""" Get supported path prefix to determine which handler to use.
|
||||
|
||||
Returns:
|
||||
A prefix.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def support_write(self) -> bool:
|
||||
""" Return flag if this file system supports write operation.
|
||||
|
||||
Returns:
|
||||
Bool.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def support_link(self) -> bool:
|
||||
""" Return if this file system supports create a soft link.
|
||||
|
||||
Returns:
|
||||
Bool.
|
||||
"""
|
||||
pass
|
||||
|
||||
def add_target_local_map(self, target_dir, local_dir):
|
||||
""" Map target directory to local file system directory
|
||||
|
||||
Args:
|
||||
target_dir (str): Target directory.
|
||||
local_dir (str): Directory in local file system.
|
||||
"""
|
||||
self._target_local_mapper[target_dir] = local_dir
|
||||
|
||||
def map_to_local(self, target_path, etag='') -> (str, bool):
|
||||
""" Map target path to local file path. (NO IO HERE).
|
||||
|
||||
Args:
|
||||
target_path (str): Target file path.
|
||||
|
||||
Returns:
|
||||
A local path and a flag indicates if the local path is a temporary file.
|
||||
"""
|
||||
for target_dir, local_dir in self._target_local_mapper.items():
|
||||
if target_path.startswith(target_dir):
|
||||
return osp.join(local_dir, osp.relpath(target_path,
|
||||
target_dir)), False
|
||||
else:
|
||||
return self._make_temporary_file(target_path, etag=etag), True
|
||||
|
||||
def convert_to_local_path(self, target_path, etag='') -> str:
|
||||
""" Deprecated. Use map_to_local() function instead.
|
||||
"""
|
||||
warnings.warn(
|
||||
'Function convert_to_local_path is deprecated, use map_to_local() function instead.'
|
||||
)
|
||||
local_path, _ = self.map_to_local(target_path, etag=etag)
|
||||
return local_path
|
||||
|
||||
def basename(self, target_path) -> str:
|
||||
""" Get file name from target_path
|
||||
|
||||
Args:
|
||||
target_path (str): Target file path.
|
||||
|
||||
Returns:
|
||||
A file name.
|
||||
"""
|
||||
return osp.basename(target_path)
|
||||
|
||||
# Functions with heavy io
|
||||
@abstractmethod
|
||||
def get_object_to_local_file(self,
|
||||
target_path,
|
||||
local_path=None,
|
||||
wait_finish=False) -> Optional[str]:
|
||||
""" Transfer file object to local file.
|
||||
If local_path is not None,
|
||||
if path can be searched in local_mapper, download it as a persistent file
|
||||
else, download it as a temporary file
|
||||
else
|
||||
download it as a persistent file
|
||||
|
||||
wait_finish when multi-processing download the same data, set wait_finish as True to avoid conflict
|
||||
|
||||
Args:
|
||||
target_path (str): path of object in different file systems
|
||||
local_path (Optional[str]): If not None, will write path to local_path.
|
||||
|
||||
Returns:
|
||||
Local file path of the object, none means a failure happened.
|
||||
"""
|
||||
pass
|
||||
|
||||
# Functions with heavy io
|
||||
@abstractmethod
|
||||
def get_object(self, target_path):
|
||||
""" Transfer file object to local file.
|
||||
If local_path is not None,
|
||||
if path can be searched in local_mapper, download it as a persistent file
|
||||
else, download it as a temporary file
|
||||
else
|
||||
download it as a persistent file
|
||||
|
||||
Args:
|
||||
target_path (str): path of object in different file systems
|
||||
local_path (Optional[str]): If not None, will write path to local_path.
|
||||
|
||||
Returns:
|
||||
Local file path of the object, none means a failure happened.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_object_stream(self, target_path, start, size, delimiter=None):
|
||||
""" Transfer file object to local file.
|
||||
If local_path is not None,
|
||||
if path can be searched in local_mapper, download it as a persistent file
|
||||
else, download it as a temporary file
|
||||
else
|
||||
download it as a persistent file
|
||||
|
||||
Args:
|
||||
target_path (str): path of object in different file systems
|
||||
start (int): object's start position.
|
||||
size (int): object's bytes size.
|
||||
delimiter (str): records's delimiter.
|
||||
|
||||
Returns:
|
||||
Local file path of the object, none means a failure happened.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def put_object_from_local_file(self, local_path, target_path) -> bool:
|
||||
""" Put local file to target file system path.
|
||||
|
||||
Args:
|
||||
local_path (str): local file path of the object
|
||||
target_path (str): target file path of the object
|
||||
|
||||
Returns:
|
||||
Bool.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def put_object(self, local_data, target_path) -> bool:
|
||||
""" Put local file to target file system path.
|
||||
|
||||
Args:
|
||||
local_path (binary): local data of the object
|
||||
target_path (str): target file path of the object
|
||||
|
||||
Returns:
|
||||
Bool.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def make_link(self, target_link_path, target_path) -> bool:
|
||||
""" Make soft link to target_path.
|
||||
|
||||
Args:
|
||||
target_link_path (str):
|
||||
target_path (str)
|
||||
|
||||
Returns:
|
||||
Bool.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def make_dir(self, target_dir) -> bool:
|
||||
""" Make a directory.
|
||||
If target_dir is already exists, return True.
|
||||
|
||||
Args:
|
||||
target_dir (str):
|
||||
|
||||
Returns:
|
||||
True if target_dir exists or created.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def remove(self, target_path) -> bool:
|
||||
""" Remove target file.
|
||||
|
||||
Args:
|
||||
target_path (str):
|
||||
|
||||
Returns:
|
||||
Bool.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_logging_handler(self, target_logging_path):
|
||||
""" Get logging handler to target logging path.
|
||||
|
||||
Args:
|
||||
target_logging_path:
|
||||
|
||||
Returns:
|
||||
A handler which has a type of subclass of logging.Handler.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def walk_dir(self, file_dir, recurse=True):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def put_dir_from_local_dir(self, local_dir, target_dir) -> bool:
|
||||
""" Upload all contents in local_dir to target_dir, keep the file tree.
|
||||
|
||||
Args:
|
||||
local_dir (str):
|
||||
target_dir (str):
|
||||
|
||||
Returns:
|
||||
Bool.
|
||||
"""
|
||||
pass
|
||||
|
||||
def _make_temporary_file(self, target_path, etag=''):
|
||||
""" Make a temporary file for target_path, which should have the same suffix.
|
||||
|
||||
Args:
|
||||
target_path (str):
|
||||
|
||||
Returns:
|
||||
A path (str).
|
||||
"""
|
||||
file_name = self.basename(target_path)
|
||||
_, suffix = osp.splitext(file_name)
|
||||
if self.tmp_dir is None:
|
||||
rand_name = '{0:%Y%m%d%H%M%S%f}'.format(
|
||||
datetime.datetime.now()) + '_' + ''.join(
|
||||
[str(random.randint(1, 10)) for _ in range(5)])
|
||||
# rand_name = get_md5(target_path)
|
||||
if suffix:
|
||||
rand_name += f'{suffix}'
|
||||
tmp_file = osp.join(tempfile.gettempdir(), rand_name)
|
||||
else:
|
||||
cache_name = '{}{}{}'.format(etag, get_md5(target_path), suffix)
|
||||
tmp_file = osp.join(self.tmp_dir, cache_name)
|
||||
return tmp_file
|
||||
|
||||
# Functions only for status, light io
|
||||
@abstractmethod
|
||||
def exists(self, target_path) -> bool:
|
||||
""" Check if target_path exists.
|
||||
|
||||
Args:
|
||||
target_path (str):
|
||||
|
||||
Returns:
|
||||
Bool.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def isfile(self, target_path) -> bool:
|
||||
""" Check if target_path is a file.
|
||||
|
||||
Args:
|
||||
target_path (str):
|
||||
|
||||
Returns:
|
||||
Bool.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def isdir(self, target_path) -> bool:
|
||||
""" Check if target_path is a directory.
|
||||
|
||||
Args:
|
||||
target_path (str):
|
||||
|
||||
Returns:
|
||||
Bool.
|
||||
"""
|
||||
|
||||
def add_temp_file(self, tmp_file):
|
||||
self._temp_files.add(tmp_file)
|
||||
|
||||
def clear(self):
|
||||
"""Delete all temp files
|
||||
"""
|
||||
if self.auto_clean:
|
||||
for temp_local_file in self._temp_files:
|
||||
remove_temp_path(temp_local_file)
|
||||
|
||||
# Functions for context
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
# if self.tmp_dir is None or self.auto_clean:
|
||||
# for temp_local_file in self._temp_files:
|
||||
# remove_temp_path(temp_local_file)
|
||||
pass
|
||||
|
||||
def __del__(self):
|
||||
pass
|
||||
|
||||
def copy(self):
|
||||
obj = copy(self)
|
||||
obj._temp_files = set(
|
||||
) # A new obj to avoid confusing in multi-thread context.
|
||||
return obj
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('FILE_SYSTEMS',
|
||||
__class__.__name__,
|
||||
BaseFs.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,140 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import os
|
||||
import os.path as osp
|
||||
import urllib.parse as parse
|
||||
import urllib.request
|
||||
from typing import Optional, Union
|
||||
|
||||
from scepter.modules.utils.file_clients.base_fs import BaseFs
|
||||
from scepter.modules.utils.file_clients.registry import FILE_SYSTEMS
|
||||
|
||||
|
||||
@FILE_SYSTEMS.register_class()
|
||||
class HttpFs(BaseFs):
|
||||
para_dict = {
|
||||
'RETRY_TIMES': {
|
||||
'value': 10,
|
||||
'description': 'Retry get object times.'
|
||||
}
|
||||
}
|
||||
para_dict.update(BaseFs.para_dict)
|
||||
|
||||
def __init__(self, cfg, logger):
|
||||
super(HttpFs, self).__init__(cfg, logger=logger)
|
||||
retry_times = cfg.get('RETRY_TIMES', 10)
|
||||
self._retry_times = retry_times
|
||||
|
||||
def get_prefix(self) -> str:
|
||||
return 'http'
|
||||
|
||||
def support_write(self) -> bool:
|
||||
return False
|
||||
|
||||
def support_link(self) -> bool:
|
||||
return False
|
||||
|
||||
def basename(self, target_path) -> str:
|
||||
url = parse.unquote(target_path)
|
||||
url = url.split('?')[0]
|
||||
return osp.basename(url)
|
||||
|
||||
def get_object_to_local_file(self,
|
||||
target_path,
|
||||
local_path=None,
|
||||
wait_finish=False) -> Optional[str]:
|
||||
if local_path is None:
|
||||
local_path, is_tmp = self.map_to_local(target_path)
|
||||
else:
|
||||
is_tmp = False
|
||||
|
||||
os.makedirs(osp.dirname(local_path), exist_ok=True)
|
||||
|
||||
retry = 0
|
||||
while retry < self._retry_times:
|
||||
try:
|
||||
target_url = urllib.parse.quote(target_path,
|
||||
safe=":/?#[]@!$&'()*+,;=%")
|
||||
urllib.request.urlretrieve(target_url, local_path)
|
||||
if osp.exists(local_path):
|
||||
break
|
||||
except Exception:
|
||||
retry += 1
|
||||
|
||||
if retry >= self._retry_times:
|
||||
return None
|
||||
|
||||
if is_tmp:
|
||||
self.add_temp_file(local_path)
|
||||
return local_path
|
||||
|
||||
def get_object(self, target_path):
|
||||
try:
|
||||
local_data = open(self.get_object_to_local_file(target_path),
|
||||
'rb').read()
|
||||
except Exception as e:
|
||||
self.logger.error(f'Read {target_path} error {e}')
|
||||
local_data = None
|
||||
return local_data
|
||||
|
||||
def put_object(self, local_data, target_path):
|
||||
raise NotImplementedError
|
||||
|
||||
def put_object_from_local_file(self, local_path, target_path) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
def make_link(self, target_link_path, target_path) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
def make_dir(self, target_dir) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
def remove(self, target_path) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
def get_logging_handler(self, target_logging_path):
|
||||
raise NotImplementedError
|
||||
|
||||
def walk_dir(self, file_dir, recurse=True):
|
||||
raise NotImplementedError
|
||||
|
||||
def put_dir_from_local_dir(self, local_dir, target_dir) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
def size(self, target_path) -> Optional[int]:
|
||||
raise NotImplementedError
|
||||
|
||||
def get_object_chunk_list(self,
|
||||
target_path,
|
||||
chunk_num=1,
|
||||
delimiter=None) -> Optional[list]:
|
||||
raise NotImplementedError
|
||||
|
||||
def get_object_stream(
|
||||
self,
|
||||
target_path,
|
||||
start,
|
||||
size=10000,
|
||||
delimiter=None) -> (Union[bytes, str, None], Optional[int]):
|
||||
raise NotImplementedError
|
||||
|
||||
def get_url(self, target_path, lifecycle=3600 * 100):
|
||||
return target_path
|
||||
|
||||
def exists(self, target_path) -> bool:
|
||||
req = urllib.request.Request(target_path)
|
||||
req.get_method = lambda: 'HEAD'
|
||||
|
||||
try:
|
||||
urllib.request.urlopen(req)
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def isfile(self, target_path) -> bool:
|
||||
# Well for a http url, it should only be a file.
|
||||
return True
|
||||
|
||||
def isdir(self, target_path) -> bool:
|
||||
return False
|
||||
@@ -0,0 +1,337 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import logging
|
||||
import os
|
||||
import os.path as osp
|
||||
import shutil
|
||||
from typing import Optional, Union
|
||||
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.file_clients.base_fs import BaseFs
|
||||
from scepter.modules.utils.file_clients.registry import FILE_SYSTEMS
|
||||
|
||||
|
||||
def is_start_of_line(f, position, delimiter='\n'):
|
||||
if position == 0:
|
||||
return True
|
||||
# Check whether the previous character is EOL
|
||||
f.seek(position - 1)
|
||||
return f.read(1) == delimiter
|
||||
|
||||
|
||||
def get_next_line_position(f, position):
|
||||
# Read the current line till the end
|
||||
f.seek(position)
|
||||
f.readline()
|
||||
# Return a position after reading the line
|
||||
return f.tell()
|
||||
|
||||
|
||||
@FILE_SYSTEMS.register_class()
|
||||
class LocalFs(BaseFs):
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(LocalFs, self).__init__(cfg, logger=logger)
|
||||
self._fs_prefix = os.path.abspath(os.curdir)
|
||||
|
||||
def get_prefix(self) -> str:
|
||||
return self._fs_prefix
|
||||
|
||||
def reconstruct_path(self, target_path) -> str:
|
||||
if target_path.startswith(self.get_prefix()):
|
||||
return target_path
|
||||
if target_path.startswith('./') or target_path.startswith('../'):
|
||||
return os.path.join(self.get_prefix(),
|
||||
target_path).replace('/./',
|
||||
'/').replace('/../', '/')
|
||||
if target_path.startswith('/'):
|
||||
return target_path
|
||||
if target_path.startswith('file://'):
|
||||
return os.path.join(self.get_prefix(),
|
||||
target_path[len('file://'):])
|
||||
return os.path.join(self.get_prefix(), target_path)
|
||||
|
||||
def support_write(self) -> bool:
|
||||
return True
|
||||
|
||||
def support_link(self) -> bool:
|
||||
return True
|
||||
|
||||
def map_to_local(self, target_path) -> (str, bool):
|
||||
target_path = self.reconstruct_path(target_path)
|
||||
return target_path, False
|
||||
|
||||
def get_object_to_local_file(self,
|
||||
target_path,
|
||||
local_path=None,
|
||||
wait_finish=False) -> Optional[str]:
|
||||
target_path = self.reconstruct_path(target_path)
|
||||
if local_path is not None:
|
||||
local_path = self.reconstruct_path(local_path)
|
||||
local_path = osp.abspath(local_path)
|
||||
if local_path != target_path:
|
||||
# copy target_path to local_path
|
||||
os.makedirs(osp.dirname(local_path), exist_ok=True)
|
||||
try:
|
||||
shutil.copy(target_path, local_path)
|
||||
except Exception as e:
|
||||
self.logger.info(f'Copy file failed {e}')
|
||||
return None
|
||||
|
||||
return local_path
|
||||
return target_path
|
||||
|
||||
def get_dir_to_local_dir(self,
|
||||
target_path,
|
||||
local_path=None,
|
||||
wait_finish=False,
|
||||
timeout=3600,
|
||||
worker_id=0) -> Optional[str]:
|
||||
if not self.isdir(target_path):
|
||||
self.logger.info(
|
||||
f"{target_path} is not directory or doesn't exist.")
|
||||
if not target_path.endswith('/'):
|
||||
target_path += '/'
|
||||
if local_path is None:
|
||||
local_path, is_tmp = self.map_to_local(target_path)
|
||||
else:
|
||||
is_tmp = False
|
||||
local_path = local_path.replace('/./', '/')
|
||||
os.makedirs(local_path, exist_ok=True)
|
||||
generator = self.walk_dir(target_path)
|
||||
for file_name in generator:
|
||||
if file_name == target_path or file_name == target_path + '/':
|
||||
continue
|
||||
local_file_name = os.path.join(
|
||||
local_path,
|
||||
file_name.split(target_path)[-1]).replace('/./', '/')
|
||||
if not self.isdir(file_name):
|
||||
self.get_object_to_local_file(file_name,
|
||||
local_file_name,
|
||||
wait_finish=wait_finish)
|
||||
else:
|
||||
self.get_dir_to_local_dir(file_name,
|
||||
local_file_name,
|
||||
wait_finish=wait_finish)
|
||||
if is_tmp:
|
||||
self.add_temp_file(local_path)
|
||||
return local_path
|
||||
|
||||
def get_object(self, target_path) -> Optional[bytes]:
|
||||
target_path = self.reconstruct_path(target_path)
|
||||
try:
|
||||
local_data = open(target_path, 'rb').read()
|
||||
except Exception as e:
|
||||
self.logger.error(f'Read {target_path} error {e}')
|
||||
local_data = None
|
||||
return local_data
|
||||
|
||||
def get_object_stream(
|
||||
self,
|
||||
target_path,
|
||||
start,
|
||||
size=10000,
|
||||
delimiter=None) -> (Union[bytes, str, None], Optional[int]):
|
||||
target_path = self.reconstruct_path(target_path)
|
||||
if not osp.exists(target_path):
|
||||
self.logger.error(f'Read {target_path} error: file not exist.')
|
||||
return None, None
|
||||
|
||||
file_size = os.path.getsize(target_path)
|
||||
start, end = start, min(file_size, start + size)
|
||||
if start >= end - 1:
|
||||
return None, None
|
||||
|
||||
with open(target_path, 'rb') as f:
|
||||
if delimiter is None or end == file_size:
|
||||
f.seek(start)
|
||||
local_data = f.read(end - start)
|
||||
return local_data, end
|
||||
f.seek(start)
|
||||
local_data = f.read(end - start)
|
||||
try:
|
||||
total_len = len(local_data)
|
||||
offset = 0
|
||||
try:
|
||||
sp_data = local_data.split(bytes(delimiter, 'utf-8'))
|
||||
# if failed, suppose the bytes is splited error.
|
||||
except Exception as e:
|
||||
self.logger.info(
|
||||
f'Return data split error,please check your delimiter {e}'
|
||||
)
|
||||
return None, end
|
||||
if not len(sp_data[-1]) == len(local_data):
|
||||
cur_offset = len(sp_data[-1])
|
||||
offset += cur_offset
|
||||
local_data = local_data[:total_len - cur_offset]
|
||||
local_data = local_data[len(sp_data[0]):]
|
||||
end = end - offset
|
||||
except Exception as e:
|
||||
self.logger.info(f'Local data decode error {e}')
|
||||
return local_data, end
|
||||
|
||||
def get_object_chunk_list(self,
|
||||
target_path,
|
||||
chunk_num=-1,
|
||||
chunk_size=-1,
|
||||
delimiter=None) -> Optional[list]:
|
||||
target_path = self.reconstruct_path(target_path)
|
||||
if not osp.exists(target_path):
|
||||
self.logger.error(f'Read {target_path} error: file not exist.')
|
||||
return None
|
||||
file_size = os.path.getsize(target_path)
|
||||
if chunk_size < 0 and chunk_num < 0:
|
||||
self.logger.error(
|
||||
'Suppose chunk size > 0 or chunk num > 0, instead of both < 0')
|
||||
return None
|
||||
if chunk_size < 0 and chunk_num > 0:
|
||||
chunk_size = file_size // chunk_num + 1
|
||||
chunk_st_et = []
|
||||
# Don't care of the lines info
|
||||
if delimiter is None:
|
||||
chunk_start = 0
|
||||
# Iterate over all chunks and construct arguments for `process_chunk`
|
||||
while chunk_start < file_size:
|
||||
chunk_end = min(file_size, chunk_start + chunk_size)
|
||||
chunk_st_et.append([chunk_start, chunk_end - chunk_start])
|
||||
chunk_start = chunk_end
|
||||
else:
|
||||
with open(target_path, 'rb') as f:
|
||||
chunk_start = 0
|
||||
offset = 0
|
||||
# Iterate over all chunks and construct arguments for `process_chunk`
|
||||
while chunk_start < file_size:
|
||||
chunk_end = min(file_size,
|
||||
chunk_start + chunk_size + offset)
|
||||
quota_st = max(0, chunk_end - 20000)
|
||||
quota_st = max(chunk_start, quota_st)
|
||||
f.seek(quota_st)
|
||||
local_data = f.read(chunk_end - quota_st)
|
||||
offset = 0
|
||||
if not chunk_end == file_size:
|
||||
try:
|
||||
try:
|
||||
sp_data = local_data.split(
|
||||
bytes(delimiter, 'utf-8'))
|
||||
# if failed, suppose the bytes is splited error.
|
||||
except Exception as e:
|
||||
self.logger.info(
|
||||
f'Return data split error,please check your delimiter {e}'
|
||||
)
|
||||
return None
|
||||
if not len(sp_data[-1]) == len(local_data):
|
||||
cur_offset = len(sp_data[-1])
|
||||
offset += cur_offset
|
||||
chunk_end = chunk_end - offset
|
||||
except Exception as e:
|
||||
self.logger.info(
|
||||
f'Local data decode error {e}, check your data is supported by str.decode().'
|
||||
)
|
||||
chunk_end = chunk_end - offset
|
||||
chunk_st_et.append([chunk_start, chunk_end - chunk_start])
|
||||
chunk_start = chunk_end
|
||||
return chunk_st_et
|
||||
|
||||
def put_object(self, local_data, target_path) -> bool:
|
||||
target_path = self.reconstruct_path(target_path)
|
||||
with open(target_path, 'w') as f:
|
||||
f.write(local_data)
|
||||
return True
|
||||
|
||||
def walk_dir(self, file_dir, recurse=True):
|
||||
for root, dirs, files in os.walk(file_dir, topdown=True):
|
||||
sub_files = files + dirs
|
||||
for name in sub_files:
|
||||
yield os.path.join(root, name)
|
||||
|
||||
def put_object_from_local_file(self, local_path, target_path) -> bool:
|
||||
target_path = self.reconstruct_path(target_path)
|
||||
local_path = self.reconstruct_path(local_path)
|
||||
if local_path != target_path:
|
||||
try:
|
||||
shutil.copy(local_path, target_path)
|
||||
except Exception:
|
||||
return False
|
||||
return True
|
||||
|
||||
def get_url(self, target_path, set_public=False, lifecycle=3600 * 100):
|
||||
return target_path
|
||||
|
||||
def make_dir(self, target_dir) -> bool:
|
||||
target_dir = self.reconstruct_path(target_dir)
|
||||
if osp.exists(target_dir):
|
||||
if osp.isfile(target_dir):
|
||||
self.logger.error(f'{target_dir} already exists as a file!')
|
||||
return False
|
||||
return True
|
||||
try:
|
||||
os.makedirs(target_dir)
|
||||
except Exception as e:
|
||||
self.logger.error(e)
|
||||
return False
|
||||
return True
|
||||
|
||||
def make_link(self, target_link_path, target_path) -> bool:
|
||||
target_path = self.reconstruct_path(target_path)
|
||||
target_link_path = self.reconstruct_path(target_link_path)
|
||||
try:
|
||||
if osp.lexists(target_link_path):
|
||||
os.remove(target_link_path)
|
||||
os.symlink(target_path, target_link_path)
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def remove(self, target_path) -> bool:
|
||||
target_path = self.reconstruct_path(target_path)
|
||||
if osp.exists(target_path):
|
||||
try:
|
||||
os.remove(target_path)
|
||||
except Exception:
|
||||
return False
|
||||
return True
|
||||
|
||||
def get_logging_handler(self, target_logging_path):
|
||||
return logging.FileHandler(target_logging_path)
|
||||
|
||||
def put_dir_from_local_dir(self, local_dir, target_dir) -> bool:
|
||||
local_dir = self.reconstruct_path(local_dir)
|
||||
target_dir = self.reconstruct_path(target_dir)
|
||||
if local_dir == target_dir:
|
||||
return True
|
||||
# cp -f local_dir/* target_dir/*
|
||||
if not osp.exists(target_dir):
|
||||
status = os.system(f'mkdir -p {target_dir}')
|
||||
if status != 0:
|
||||
return False
|
||||
try:
|
||||
shutil.copytree(local_dir, target_dir, symlinks=True)
|
||||
except Exception:
|
||||
return False
|
||||
return True
|
||||
|
||||
def size(self, target_path) -> Optional[int]:
|
||||
target_path = self.reconstruct_path(target_path)
|
||||
if not osp.exists(target_path):
|
||||
self.logger.info(f"File {target_path} doesn't exist.")
|
||||
return -1
|
||||
file_size = os.path.getsize(target_path)
|
||||
return file_size
|
||||
|
||||
def exists(self, target_path) -> bool:
|
||||
target_path = self.reconstruct_path(target_path)
|
||||
return osp.exists(target_path)
|
||||
|
||||
def isfile(self, target_path) -> bool:
|
||||
target_path = self.reconstruct_path(target_path)
|
||||
return osp.isfile(target_path)
|
||||
|
||||
def isdir(self, target_path) -> bool:
|
||||
target_path = self.reconstruct_path(target_path)
|
||||
return osp.isdir(target_path)
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('FILE_SYSTEMS',
|
||||
__class__.__name__, {},
|
||||
set_name=True)
|
||||
@@ -0,0 +1,202 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import os
|
||||
import os.path as osp
|
||||
import urllib.parse as parse
|
||||
import urllib.request
|
||||
from typing import Optional, Union
|
||||
|
||||
from scepter.modules.utils.file_clients.base_fs import BaseFs
|
||||
from scepter.modules.utils.file_clients.registry import FILE_SYSTEMS
|
||||
|
||||
|
||||
@FILE_SYSTEMS.register_class()
|
||||
class ModelscopeFs(BaseFs):
|
||||
para_dict = {
|
||||
'RETRY_TIMES': {
|
||||
'value': 10,
|
||||
'description': 'Retry get object times.'
|
||||
}
|
||||
}
|
||||
para_dict.update(BaseFs.para_dict)
|
||||
|
||||
def __init__(self, cfg, logger):
|
||||
super(ModelscopeFs, self).__init__(cfg, logger=logger)
|
||||
retry_times = cfg.get('RETRY_TIMES', 10)
|
||||
self._retry_times = retry_times
|
||||
|
||||
def get_prefix(self) -> str:
|
||||
return 'ms://'
|
||||
|
||||
def support_write(self) -> bool:
|
||||
return False
|
||||
|
||||
def support_link(self) -> bool:
|
||||
return False
|
||||
|
||||
def basename(self, target_path) -> str:
|
||||
url = parse.unquote(target_path)
|
||||
url = url.split('?')[0]
|
||||
return osp.basename(url)
|
||||
|
||||
def get_object_to_local_file(self,
|
||||
target_path,
|
||||
local_path=None,
|
||||
wait_finish=False) -> Optional[str]:
|
||||
from modelscope.hub.file_download import model_file_download
|
||||
|
||||
key = osp.relpath(target_path, self.get_prefix())
|
||||
key, file_path = key.split('@', 1)
|
||||
|
||||
if ':' in key:
|
||||
key, revision = key.split(':', 1)
|
||||
else:
|
||||
revision = None
|
||||
|
||||
if local_path is None:
|
||||
local_path, is_tmp = self.map_to_local(key)
|
||||
else:
|
||||
is_tmp = False
|
||||
|
||||
if revision is not None:
|
||||
local_path = local_path + '_' + str(revision)
|
||||
|
||||
retry = 0
|
||||
while retry < self._retry_times:
|
||||
try:
|
||||
local_path = model_file_download(model_id=key,
|
||||
revision=revision,
|
||||
file_path=file_path,
|
||||
cache_dir=local_path)
|
||||
if osp.exists(local_path):
|
||||
break
|
||||
except Exception:
|
||||
retry += 1
|
||||
|
||||
if retry >= self._retry_times:
|
||||
return None
|
||||
|
||||
if is_tmp:
|
||||
self.add_temp_file(local_path)
|
||||
return local_path
|
||||
|
||||
def get_dir_to_local_dir(self,
|
||||
target_path,
|
||||
local_path=None,
|
||||
wait_finish=False,
|
||||
timeout=3600,
|
||||
worker_id=-1) -> Optional[str]:
|
||||
from modelscope.hub.snapshot_download import snapshot_download
|
||||
assert target_path.startswith(self.get_prefix())
|
||||
|
||||
key = osp.relpath(target_path, self.get_prefix())
|
||||
if '@' not in key:
|
||||
key, ret_folder = key.split('@', 1)[0], ''
|
||||
else:
|
||||
at_level_folder = key.split('@')
|
||||
if len(at_level_folder) > 2:
|
||||
raise f'Target path should include only one @, but you give {len(at_level_folder)} @.'
|
||||
key, ret_folder = at_level_folder
|
||||
|
||||
if ':' in key:
|
||||
key, revision = key.split(':', 1)
|
||||
else:
|
||||
revision = None
|
||||
|
||||
if local_path is None:
|
||||
local_path, is_tmp = self.map_to_local(key)
|
||||
else:
|
||||
is_tmp = False
|
||||
|
||||
if revision is not None:
|
||||
local_path = local_path + '_' + str(revision)
|
||||
|
||||
retry = 0
|
||||
while retry < self._retry_times:
|
||||
try:
|
||||
local_path = snapshot_download(key,
|
||||
revision=revision,
|
||||
cache_dir=local_path)
|
||||
if osp.exists(local_path):
|
||||
break
|
||||
except Exception:
|
||||
retry += 1
|
||||
|
||||
if retry >= self._retry_times:
|
||||
return None
|
||||
|
||||
if is_tmp:
|
||||
self.add_temp_file(local_path)
|
||||
if not ret_folder == '':
|
||||
local_path = os.path.join(local_path, ret_folder)
|
||||
return local_path
|
||||
|
||||
def get_object(self, target_path):
|
||||
try:
|
||||
local_data = open(self.get_object_to_local_file(target_path),
|
||||
'rb').read()
|
||||
except Exception as e:
|
||||
self.logger.error(f'Read {target_path} error {e}')
|
||||
local_data = None
|
||||
return local_data
|
||||
|
||||
def put_object(self, local_data, target_path):
|
||||
raise NotImplementedError
|
||||
|
||||
def put_object_from_local_file(self, local_path, target_path) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
def make_link(self, target_link_path, target_path) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
def make_dir(self, target_dir) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
def remove(self, target_path) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
def get_logging_handler(self, target_logging_path):
|
||||
raise NotImplementedError
|
||||
|
||||
def walk_dir(self, file_dir, recurse=True):
|
||||
raise NotImplementedError
|
||||
|
||||
def put_dir_from_local_dir(self, local_dir, target_dir) -> bool:
|
||||
raise NotImplementedError
|
||||
|
||||
def size(self, target_path) -> Optional[int]:
|
||||
raise NotImplementedError
|
||||
|
||||
def get_object_chunk_list(self,
|
||||
target_path,
|
||||
chunk_num=1,
|
||||
delimiter=None) -> Optional[list]:
|
||||
raise NotImplementedError
|
||||
|
||||
def get_object_stream(
|
||||
self,
|
||||
target_path,
|
||||
start,
|
||||
size=10000,
|
||||
delimiter=None) -> (Union[bytes, str, None], Optional[int]):
|
||||
raise NotImplementedError
|
||||
|
||||
def get_url(self, target_path, lifecycle=3600 * 100):
|
||||
return target_path
|
||||
|
||||
def exists(self, target_path) -> bool:
|
||||
req = urllib.request.Request(target_path)
|
||||
req.get_method = lambda: 'HEAD'
|
||||
|
||||
try:
|
||||
urllib.request.urlopen(req)
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def isfile(self, target_path) -> bool:
|
||||
# Well for a http url, it should only be a file.
|
||||
return True
|
||||
|
||||
def isdir(self, target_path) -> bool:
|
||||
return False
|
||||
@@ -0,0 +1,6 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
from scepter.modules.utils.registry import Registry
|
||||
|
||||
FILE_SYSTEMS = Registry('FILE_SYSTEMS')
|
||||
@@ -0,0 +1,52 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import os
|
||||
import os.path as osp
|
||||
|
||||
|
||||
def check_if_local_path(path):
|
||||
"""
|
||||
Check if path is a local path, no matter file or directory.
|
||||
|
||||
True:
|
||||
file:///home/admin/a.txt (standard path)
|
||||
/home/admin/a.txt (standard unix path)
|
||||
C:\\Users\\a.txt (standard windows path)
|
||||
C:/Users/a.txt (works as well)
|
||||
./a.txt (relative path)
|
||||
a.txt (relative path)
|
||||
False:
|
||||
http://www.aliyun.com/a.txt
|
||||
http://www.aliyun.com/a.txt
|
||||
oss://aliyun/a.txt
|
||||
|
||||
Args:
|
||||
path (str):
|
||||
|
||||
Returns:
|
||||
True if path is a local path.
|
||||
"""
|
||||
if path.startswith('file://'):
|
||||
return True
|
||||
return '://' not in path
|
||||
|
||||
|
||||
def remove_temp_path(path):
|
||||
"""
|
||||
Delete local temp path.
|
||||
|
||||
Args:
|
||||
path (str):
|
||||
|
||||
Returns:
|
||||
"""
|
||||
if not osp.exists(path):
|
||||
return
|
||||
if not osp.isfile(path):
|
||||
return
|
||||
try:
|
||||
os.remove(path)
|
||||
except Exception:
|
||||
pass
|
||||
# warnings.warn(f"remove {path}")
|
||||
@@ -0,0 +1,445 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import os
|
||||
import threading
|
||||
import time
|
||||
import warnings
|
||||
from contextlib import contextmanager
|
||||
from queue import Queue
|
||||
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.file_clients.base_fs import BaseFs
|
||||
from scepter.modules.utils.file_clients.local_fs import LocalFs
|
||||
from scepter.modules.utils.file_clients.registry import FILE_SYSTEMS
|
||||
from scepter.modules.utils.file_clients.utils import check_if_local_path
|
||||
|
||||
|
||||
class IoString(str):
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
pass
|
||||
|
||||
|
||||
class IoBytes(bytes):
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
pass
|
||||
|
||||
|
||||
class ReadException(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class WriteException(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class FileSystem(object):
|
||||
def __init__(self):
|
||||
self._prefix_to_clients = {}
|
||||
self._default_client = None
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
pass
|
||||
|
||||
def __del__(self):
|
||||
for k, client in self._prefix_to_clients.items():
|
||||
client.clear()
|
||||
|
||||
@property
|
||||
def support_prefix(self):
|
||||
return self._prefix_to_clients
|
||||
|
||||
def init_fs_client(self, cfg=None, logger=None, overwrite=True):
|
||||
""" Initialize file system backend
|
||||
Supported backend:
|
||||
1. Local file system, e.g. /home/admin/work_dir, work_dir_bk/imagenet_pretrain
|
||||
2. Aliyun Oss, e.g. oss://bucket_name/work_dir
|
||||
3. Http, only support to read content, e.g.
|
||||
https://www.google.com.hk/images/branding/googlelogo/2x/googlelogo_color_272x92dp.png
|
||||
4. other fs backend...
|
||||
|
||||
Args:
|
||||
cfg (list, dict, optional):
|
||||
list: list of file system configs to be initialized
|
||||
dict: a dict contains file system configs as values or a file system config dict
|
||||
optional: Will only use default LocalFs
|
||||
"""
|
||||
|
||||
fs_cfg = cfg or Config(load=False)
|
||||
if not isinstance(fs_cfg, Config):
|
||||
raise '{} is not a Config Instance!'.format(fs_cfg)
|
||||
|
||||
if not fs_cfg.have('NAME'):
|
||||
raise KeyError(f'{fs_cfg} does not contain key NAME!')
|
||||
|
||||
fs_client = FILE_SYSTEMS.build(fs_cfg, logger=logger)
|
||||
_prefix = fs_client.get_prefix()
|
||||
if _prefix in self._prefix_to_clients and not overwrite:
|
||||
return _prefix
|
||||
if _prefix in self._prefix_to_clients:
|
||||
warnings.warn(
|
||||
'File client {} has already been set, will be replaced by newer config.'
|
||||
.format(_prefix))
|
||||
self._prefix_to_clients[_prefix] = fs_client
|
||||
return _prefix
|
||||
|
||||
def get_fs_client(self, target_path, safe=False) -> BaseFs:
|
||||
""" Get the client by input path.
|
||||
Every file system has its own identifier, default will use local file system to have a try.
|
||||
If copy needed, only do shallow copy.
|
||||
|
||||
Args:
|
||||
target_path (str):
|
||||
safe (bool): In safe mode, get the copy of the client.
|
||||
"""
|
||||
obj = None
|
||||
|
||||
for prefix in sorted(list(self._prefix_to_clients.keys()),
|
||||
key=lambda a: -len(a)):
|
||||
if target_path.startswith(prefix):
|
||||
obj = self._prefix_to_clients[prefix]
|
||||
break
|
||||
if obj is not None:
|
||||
if safe:
|
||||
return obj.copy()
|
||||
else:
|
||||
return obj
|
||||
|
||||
if not check_if_local_path(target_path):
|
||||
warnings.warn(
|
||||
f'{target_path} is not a local path, use LocalFs may cause an error.'
|
||||
)
|
||||
if self._default_client is None:
|
||||
self._default_client = LocalFs(Config(load=False))
|
||||
if safe:
|
||||
return self._default_client.copy()
|
||||
else:
|
||||
return self._default_client
|
||||
|
||||
def get_dir_to_local_dir(self,
|
||||
target_path,
|
||||
local_path=None,
|
||||
wait_finish=False,
|
||||
timeout=3600,
|
||||
worker_id=0):
|
||||
with self.get_fs_client(target_path) as client:
|
||||
local_path = client.get_dir_to_local_dir(target_path,
|
||||
local_path=local_path,
|
||||
wait_finish=wait_finish,
|
||||
timeout=timeout,
|
||||
worker_id=worker_id)
|
||||
if local_path is None:
|
||||
raise ReadException(
|
||||
f'Failed to fetch {target_path} to {local_path}')
|
||||
return IoString(local_path)
|
||||
|
||||
def add_target_local_map(self, target_dir, local_dir):
|
||||
""" Map target directory to local file system directory
|
||||
|
||||
Args:
|
||||
target_dir (str): Target directory.
|
||||
local_dir (str): Directory in local file system.
|
||||
"""
|
||||
with self.get_fs_client(target_dir, safe=False) as client:
|
||||
client.add_target_local_map(target_dir, local_dir)
|
||||
|
||||
def make_dir(self, target_dir):
|
||||
""" Make a directory.
|
||||
If target_dir is already exists, return True.
|
||||
|
||||
Args:
|
||||
target_dir (str):
|
||||
|
||||
Returns:
|
||||
True if target_dir exists or created.
|
||||
"""
|
||||
with self.get_fs_client(target_dir) as client:
|
||||
return client.make_dir(target_dir)
|
||||
|
||||
def exists(self, target_path):
|
||||
""" Check if target_path exists.
|
||||
|
||||
Args:
|
||||
target_path (str):
|
||||
|
||||
Returns:
|
||||
Bool.
|
||||
"""
|
||||
with self.get_fs_client(target_path) as client:
|
||||
return client.exists(target_path)
|
||||
|
||||
def map_to_local(self, target_path):
|
||||
""" Map target path to local file path. (NO IO HERE).
|
||||
|
||||
Args:
|
||||
target_path (str): Target file path.
|
||||
|
||||
Returns:
|
||||
A local path and a flag indicates if the local path is a temporary file.
|
||||
"""
|
||||
with self.get_fs_client(target_path) as client:
|
||||
local_path, is_tmp = client.map_to_local(target_path)
|
||||
return local_path, is_tmp
|
||||
|
||||
def put_dir_from_local_dir(self, local_dir, target_dir):
|
||||
""" Upload all contents in local_dir to target_dir, keep the file tree.
|
||||
|
||||
Args:
|
||||
local_dir (str):
|
||||
target_dir (str):
|
||||
|
||||
Returns:
|
||||
Bool.
|
||||
"""
|
||||
with self.get_fs_client(target_dir) as client:
|
||||
return client.put_dir_from_local_dir(local_dir, target_dir)
|
||||
|
||||
def walk_dir(self, target_dir, recurse=True):
|
||||
""" Iterator to access the files of target dir.
|
||||
Args:
|
||||
target_dir (str):
|
||||
|
||||
Returns:
|
||||
Generator.
|
||||
"""
|
||||
with self.get_fs_client(target_dir) as client:
|
||||
return client.walk_dir(target_dir, recurse=recurse)
|
||||
|
||||
def is_local_client(self, target_path) -> bool:
|
||||
""" Check if the client support read or write to target_path is a LocalFs.
|
||||
|
||||
Args:
|
||||
target_path (str):
|
||||
|
||||
Returns:
|
||||
Bool.
|
||||
"""
|
||||
with self.get_fs_client(target_path) as client:
|
||||
return type(client) is LocalFs
|
||||
|
||||
def put_object_from_local_file(self, local_path, target_path) -> bool:
|
||||
with self.get_fs_client(target_path) as client:
|
||||
flag = client.put_object_from_local_file(local_path, target_path)
|
||||
return flag
|
||||
|
||||
def get_from(self, target_path, local_path=None, wait_finish=False):
|
||||
with self.get_fs_client(target_path) as client:
|
||||
local_path = client.get_object_to_local_file(
|
||||
target_path, local_path=local_path, wait_finish=wait_finish)
|
||||
if local_path is None:
|
||||
raise ReadException(
|
||||
f'Failed to fetch {target_path} to {local_path}')
|
||||
return IoString(local_path)
|
||||
|
||||
def get_url(self, target_path, set_public=False, lifecycle=3600 * 100):
|
||||
with self.get_fs_client(target_path) as client:
|
||||
output_url = client.get_url(target_path,
|
||||
set_public=set_public,
|
||||
lifecycle=lifecycle)
|
||||
return output_url
|
||||
|
||||
def get_object(self, target_path):
|
||||
with self.get_fs_client(target_path) as client:
|
||||
local_data = client.get_object(target_path)
|
||||
if local_data is None:
|
||||
return IoBytes(None)
|
||||
return IoBytes(local_data)
|
||||
|
||||
def put_object(self, local_data, target_path):
|
||||
with self.get_fs_client(target_path) as client:
|
||||
flg = client.put_object(local_data, target_path)
|
||||
return flg
|
||||
|
||||
def delete_object(self, target_path):
|
||||
with self.get_fs_client(target_path) as client:
|
||||
if self.isfile(target_path):
|
||||
flg = client.remove(target_path)
|
||||
return flg
|
||||
else:
|
||||
return False
|
||||
|
||||
def get_batch_objects_from(self, target_path_list, wait_finish=False):
|
||||
data_quene = Queue()
|
||||
batch_size = 20
|
||||
R = threading.Lock()
|
||||
|
||||
def get_one_object(target_path_list):
|
||||
for target_path in target_path_list:
|
||||
if self.exists(target_path):
|
||||
local_path = self.get_from(target_path,
|
||||
wait_finish=wait_finish)
|
||||
else:
|
||||
local_path = None
|
||||
R.acquire()
|
||||
try:
|
||||
data_quene.put_nowait([target_path, local_path])
|
||||
except Exception:
|
||||
R.release()
|
||||
R.release()
|
||||
|
||||
while True:
|
||||
batch_list = target_path_list[:4 * batch_size]
|
||||
if len(batch_list) < 1:
|
||||
break
|
||||
target_path_list = target_path_list[4 * batch_size:]
|
||||
threading_list = []
|
||||
for i in range(batch_size):
|
||||
t = threading.Thread(target=get_one_object,
|
||||
args=(batch_list[i::batch_size], ))
|
||||
t.daemon = True
|
||||
t.start()
|
||||
threading_list.append(t)
|
||||
[threading_t.join() for threading_t in threading_list]
|
||||
file_dict = {}
|
||||
while not data_quene.empty():
|
||||
target_path, local_path = data_quene.get_nowait()
|
||||
file_dict[target_path] = local_path
|
||||
|
||||
for target_path in batch_list:
|
||||
local_path = file_dict.get(target_path, None)
|
||||
yield local_path
|
||||
|
||||
def put_batch_objects_to(self,
|
||||
local_path_list,
|
||||
target_path_list,
|
||||
batch_size=20,
|
||||
wait_finish=False):
|
||||
data_quene = Queue()
|
||||
R = threading.Lock()
|
||||
|
||||
def put_one_object(local_path_list, target_path_list):
|
||||
for local_path, target_path in zip(local_path_list,
|
||||
target_path_list):
|
||||
if local_path is None or target_path is None:
|
||||
flg = False
|
||||
elif self.exists(local_path):
|
||||
local_cache = self.get_from(local_path,
|
||||
local_path + f'{time.time()}',
|
||||
wait_finish=wait_finish)
|
||||
flg = self.put_object_from_local_file(
|
||||
local_cache, target_path)
|
||||
try:
|
||||
if os.path.exists(local_cache):
|
||||
os.remove(local_cache)
|
||||
except Exception:
|
||||
pass
|
||||
else:
|
||||
flg = False
|
||||
R.acquire()
|
||||
try:
|
||||
data_quene.put_nowait([local_path, target_path, flg])
|
||||
except Exception:
|
||||
R.release()
|
||||
R.release()
|
||||
|
||||
while True:
|
||||
batch_local_list = local_path_list[:4 * batch_size]
|
||||
batch_target_list = target_path_list[:4 * batch_size]
|
||||
if len(batch_local_list) < 1:
|
||||
break
|
||||
local_path_list = local_path_list[4 * batch_size:]
|
||||
target_path_list = target_path_list[4 * batch_size:]
|
||||
threading_list = []
|
||||
for i in range(batch_size):
|
||||
t = threading.Thread(target=put_one_object,
|
||||
args=(
|
||||
batch_local_list[i::batch_size],
|
||||
batch_target_list[i::batch_size],
|
||||
))
|
||||
t.daemon = True
|
||||
t.start()
|
||||
threading_list.append(t)
|
||||
[threading_t.join() for threading_t in threading_list]
|
||||
file_dict = {}
|
||||
while not data_quene.empty():
|
||||
local_path, target_path, flg = data_quene.get_nowait()
|
||||
file_dict[local_path] = [local_path, target_path, flg]
|
||||
|
||||
for idx, local_path in enumerate(batch_local_list):
|
||||
local_path, target_path, flg = file_dict.get(
|
||||
local_path,
|
||||
[batch_local_list[idx], batch_target_list[idx], False])
|
||||
yield local_path, target_path, flg
|
||||
|
||||
def get_object_stream(self,
|
||||
target_path,
|
||||
start,
|
||||
size=10000,
|
||||
delimiter=None):
|
||||
with self.get_fs_client(target_path) as client:
|
||||
local_data, end = client.get_object_stream(target_path,
|
||||
start,
|
||||
size=size,
|
||||
delimiter=delimiter)
|
||||
return local_data, end
|
||||
|
||||
def get_object_chunk_list(self,
|
||||
target_path,
|
||||
chunk_num=1,
|
||||
chunk_size=-1,
|
||||
delimiter=None):
|
||||
with self.get_fs_client(target_path) as client:
|
||||
chunk_list = client.get_object_chunk_list(target_path,
|
||||
chunk_num=chunk_num,
|
||||
chunk_size=chunk_size,
|
||||
delimiter=delimiter)
|
||||
if chunk_list is None:
|
||||
raise ReadException(f'Failed to fetch {target_path}')
|
||||
return chunk_list
|
||||
|
||||
def size(self, target_path):
|
||||
with self.get_fs_client(target_path) as client:
|
||||
size = client.size(target_path)
|
||||
return size
|
||||
|
||||
def isfile(self, target_path):
|
||||
with self.get_fs_client(target_path) as client:
|
||||
is_file = client.isfile(target_path)
|
||||
return is_file
|
||||
|
||||
def isdir(self, target_path):
|
||||
with self.get_fs_client(target_path) as client:
|
||||
is_dir = client.isdir(target_path)
|
||||
return is_dir
|
||||
|
||||
@contextmanager
|
||||
def put_to(self, target_path):
|
||||
with self.get_fs_client(target_path) as client:
|
||||
local_path, is_tmp = client.map_to_local(target_path)
|
||||
if is_tmp:
|
||||
client.add_temp_file(local_path)
|
||||
if not os.path.exists(os.path.dirname(local_path)):
|
||||
os.makedirs(os.path.dirname(local_path))
|
||||
yield local_path
|
||||
status = client.put_object_from_local_file(local_path, target_path)
|
||||
if not status:
|
||||
raise WriteException(
|
||||
f'Failed to upload from {local_path} to {target_path}')
|
||||
if not isinstance(client, LocalFs):
|
||||
try:
|
||||
if os.path.exists(local_path):
|
||||
os.remove(local_path)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def __repr__(self) -> str:
|
||||
s = 'Support prefix list:\n'
|
||||
for prefix in sorted(list(self._prefix_to_clients.keys()),
|
||||
key=lambda a: -len(a)):
|
||||
s += f'\t{prefix} -> {self._prefix_to_clients[prefix]}\n'
|
||||
return s
|
||||
|
||||
|
||||
global FS, DATA_FS, MODEL_FS
|
||||
# global instance, easy to use
|
||||
FS = FileSystem()
|
||||
DATA_FS = FS
|
||||
MODEL_FS = FS
|
||||
@@ -0,0 +1,177 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import logging
|
||||
import numbers
|
||||
import sys
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from scepter.modules.utils.distribute import get_dist_info
|
||||
|
||||
|
||||
def as_time(s):
|
||||
s = int(s)
|
||||
one_day, one_hour, one_min = 3600 * 24, 3600, 60
|
||||
day, hour, min = 0, 0, 0
|
||||
# compute day 3600 * 24
|
||||
if s >= one_day:
|
||||
day = int(s // one_day)
|
||||
s = s % one_day
|
||||
# compute hour 3600
|
||||
if s >= one_hour:
|
||||
hour = int(s // one_hour)
|
||||
s = s % one_hour
|
||||
# compute min 60
|
||||
if s >= one_min:
|
||||
min = int(s // one_min)
|
||||
s = s % one_min
|
||||
output_str = []
|
||||
if day > 0:
|
||||
output_str.append('{}days'.format(day))
|
||||
if hour > 0:
|
||||
output_str.append('{}hours'.format(hour))
|
||||
if min > 0:
|
||||
output_str.append('{}mins'.format(min))
|
||||
output_str.append('{}secs'.format(int(s)))
|
||||
return ' '.join(output_str)
|
||||
|
||||
|
||||
def time_since(since, percent):
|
||||
now = time.time()
|
||||
s = now - since
|
||||
es = s / (percent)
|
||||
rs = es - s
|
||||
return '{} {:.2f}%({})'.format(as_time(s), 100 * percent, as_time(rs))
|
||||
|
||||
|
||||
def get_logger(name='torch dist'):
|
||||
logger = logging.getLogger(name)
|
||||
logger.propagate = False
|
||||
if len(logger.handlers) == 0:
|
||||
std_handler = logging.StreamHandler(sys.stdout)
|
||||
formatter = logging.Formatter(
|
||||
'%(name)s [%(levelname)s] %(asctime)s '
|
||||
'[File: %(filename)s Function: %(funcName)s at line %(lineno)d] %(message)s'
|
||||
)
|
||||
std_handler.setFormatter(formatter)
|
||||
std_handler.setLevel(logging.INFO)
|
||||
logger.setLevel(logging.INFO)
|
||||
logger.addHandler(std_handler)
|
||||
return logger
|
||||
|
||||
|
||||
def init_logger(in_logger, log_file=None, dist_launcher='pytorch'):
|
||||
""" Add file handler to logger on rank 0 and set log level by dist_launcher
|
||||
|
||||
Args:
|
||||
in_logger (logging.Logger):
|
||||
log_file (str, None): if not None, a file handler will be add to in_logger
|
||||
dist_launcher (str, None):
|
||||
"""
|
||||
rank, _ = get_dist_info()
|
||||
if rank == 0:
|
||||
if log_file is not None:
|
||||
from scepter.modules.utils.file_system import FS
|
||||
file_handler = FS.get_fs_client(log_file).get_logging_handler(
|
||||
log_file)
|
||||
formatter = logging.Formatter(
|
||||
'%(name)s [%(levelname)s] %(asctime)s [File: %(filename)s '
|
||||
'Function: %(funcName)s at line %(lineno)d] %(message)s')
|
||||
file_handler.setFormatter(formatter)
|
||||
file_handler.setLevel(logging.INFO)
|
||||
in_logger.addHandler(file_handler)
|
||||
in_logger.info(f'Running task with log file: {log_file}')
|
||||
in_logger.setLevel(logging.INFO)
|
||||
else:
|
||||
if dist_launcher == 'pytorch':
|
||||
in_logger.setLevel(logging.ERROR)
|
||||
else:
|
||||
# Distribute Training with more than one machine, we'd like to show logs on every machine.
|
||||
in_logger.setLevel(logging.INFO)
|
||||
|
||||
|
||||
class LogAgg(object):
|
||||
""" Log variable aggregate tool. Recommend to invoke clear() function after one epoch.
|
||||
In distributed training environment, tensor variable will be all reduced to get an average.
|
||||
|
||||
Example:
|
||||
>>> agg = LogAgg()
|
||||
>>> agg.update(dict(loss=0.1, accuracy=0.5))
|
||||
>>> agg.update(dict(loss=0.2, accuracy=0.6))
|
||||
>>> agg.update(dict(loss=0.3, accuracy=0.7))
|
||||
>>> agg.aggregate()
|
||||
OrderedDict([('loss', (0.3, 0.20000000000000004)), ('accuracy', (0.7, 0.6))])
|
||||
"""
|
||||
def __init__(self):
|
||||
self.buffer = OrderedDict()
|
||||
self.counter = []
|
||||
|
||||
def update(self, kv: dict, count=1):
|
||||
""" Update variables
|
||||
|
||||
Args:
|
||||
kv (dict): a dict with value type in (torch.Tensor, numbers)
|
||||
count (int): divider, default is 1
|
||||
"""
|
||||
for k, v in kv.items():
|
||||
if isinstance(v, torch.Tensor):
|
||||
# Must be scalar
|
||||
if not v.ndim == 0:
|
||||
continue
|
||||
v = v.item()
|
||||
elif isinstance(v, np.ndarray):
|
||||
# Must be scalar
|
||||
if not v.ndim == 0:
|
||||
continue
|
||||
elif isinstance(v, numbers.Number):
|
||||
# Must be number
|
||||
pass
|
||||
else:
|
||||
continue
|
||||
|
||||
if k not in self.buffer:
|
||||
self.buffer[k] = []
|
||||
self.buffer[k].append(v)
|
||||
self.counter.append(count)
|
||||
|
||||
def _aggregate(self, n=0):
|
||||
""" Do aggregation.
|
||||
|
||||
Args:
|
||||
n (int): recent n numbers, if 0, start from 0
|
||||
|
||||
Returns:
|
||||
A dict contains aggregate values.
|
||||
"""
|
||||
|
||||
ret = OrderedDict()
|
||||
for key in self.buffer:
|
||||
values = np.array(self.buffer[key][-n:])
|
||||
nums = np.array(self.counter[-n:])
|
||||
avg = np.sum(values * nums) / np.sum(nums)
|
||||
ret[key] = avg
|
||||
return ret
|
||||
|
||||
def aggregate(self, log_interval=1):
|
||||
""" Do aggregation with current step values and all mean values.
|
||||
|
||||
Args:
|
||||
log_interval (int): Steps to aggregate current state, default is 1.
|
||||
|
||||
Returns:
|
||||
A dict contains current step and all step mean values.
|
||||
"""
|
||||
cur = self._aggregate(log_interval)
|
||||
all_mean = self._aggregate(0)
|
||||
ret = OrderedDict()
|
||||
for key in cur:
|
||||
ret[key] = (cur[key], all_mean[key])
|
||||
return ret
|
||||
|
||||
def reset(self):
|
||||
self.buffer.clear()
|
||||
self.counter.clear()
|
||||
@@ -0,0 +1,103 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import warnings
|
||||
|
||||
import numpy as np
|
||||
|
||||
try:
|
||||
import matplotlib.pyplot as plt
|
||||
except Exception as e:
|
||||
warnings.warn(f'Runing without matplotlib {e}')
|
||||
|
||||
color_list = ['b', 'g', 'r', 'c', 'm', 'y', 'k']
|
||||
line_list = ['-', '--', '-.', ':']
|
||||
|
||||
|
||||
def plot_multi_curves(x,
|
||||
y,
|
||||
show=False,
|
||||
title=None,
|
||||
save_path=None,
|
||||
x_label=None,
|
||||
y_label=None):
|
||||
'''
|
||||
Args:
|
||||
x: the x-axis data
|
||||
y: the y-axis data dict
|
||||
like: [{"data": np.ndarrays, "label": ""}]
|
||||
title: None
|
||||
show: False
|
||||
save_path: None
|
||||
x_label: None
|
||||
y_label: None
|
||||
Returns:
|
||||
'''
|
||||
if save_path is not None:
|
||||
plt.figure()
|
||||
|
||||
x_max, x_min = np.max(x), np.min(x)
|
||||
|
||||
max_num, min_num = 0, 0
|
||||
for y_id, data in enumerate(y):
|
||||
max_n = np.max(data['data'])
|
||||
min_n = np.min(data['data'])
|
||||
max_num = max_n if max_n > max_num else max_num
|
||||
min_num = min_n if min_n < min_num else min_num
|
||||
plt.plot(x,
|
||||
data['data'],
|
||||
linestyle=line_list[y_id % len(line_list)],
|
||||
linewidth=2,
|
||||
color=color_list[y_id % len(color_list)],
|
||||
label=data['label'],
|
||||
alpha=1.00)
|
||||
plt.title(title, loc='center')
|
||||
|
||||
plt.legend(loc='upper right')
|
||||
if x_label is not None:
|
||||
plt.xlabel(x_label)
|
||||
if y_label is not None:
|
||||
plt.ylabel(y_label)
|
||||
|
||||
x_step = (x_max - x_min) / 5
|
||||
y_step = (max_num - min_num) / 5
|
||||
plt.xticks(np.arange(x_min - x_step / 2, x_max + x_step / 2, x_step))
|
||||
plt.yticks(np.arange(min_num - y_step / 2, max_num + y_step / 2, y_step))
|
||||
plt.grid()
|
||||
if save_path is not None:
|
||||
plt.savefig(save_path)
|
||||
if show:
|
||||
plt.show()
|
||||
plt.clf()
|
||||
plt.cla()
|
||||
plt.close()
|
||||
return True
|
||||
|
||||
|
||||
def plt_curve(x,
|
||||
y,
|
||||
show=False,
|
||||
title=None,
|
||||
save_path=None,
|
||||
x_label=None,
|
||||
y_label=None):
|
||||
'''
|
||||
Args:
|
||||
x: the x-axis data
|
||||
y: the y-axis data dict
|
||||
like: [{"data": np.ndarrays, "label": ""}]
|
||||
title: None
|
||||
show: False
|
||||
save_path: None
|
||||
x_label: None
|
||||
y_label: None
|
||||
Returns:
|
||||
'''
|
||||
return plot_multi_curves(x, [{
|
||||
'data': y,
|
||||
'label': 'y'
|
||||
}],
|
||||
show=show,
|
||||
title=title,
|
||||
save_path=save_path,
|
||||
x_label=x_label,
|
||||
y_label=y_label)
|
||||
@@ -0,0 +1,164 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.utils.model_zoo import load_url as load_state_dict_from_url
|
||||
|
||||
|
||||
class StdMsg():
|
||||
def __init__(self, name='msg'):
|
||||
self.name = name
|
||||
|
||||
def info(self, msg):
|
||||
sys.stdout.write('[Info]: ' + msg + '\n')
|
||||
|
||||
def error(self, msg):
|
||||
sys.stdout.write('[Error]: ' + msg + '\n')
|
||||
|
||||
def warning(self, msg):
|
||||
sys.stdout.write('[Warning]: ' + msg + '\n')
|
||||
|
||||
|
||||
def move_model_to_cpu(params):
|
||||
cpu_params = OrderedDict()
|
||||
for key, val in params.items():
|
||||
cpu_params[key] = val.cpu()
|
||||
return cpu_params
|
||||
|
||||
|
||||
def load_pretrained(model: torch.nn.Module,
|
||||
path: str,
|
||||
map_location='cpu',
|
||||
logger=None,
|
||||
sub_level=None):
|
||||
if logger:
|
||||
logger.info(
|
||||
f'Load pretrained model [{model.__class__.__name__}] from {path}')
|
||||
if os.path.exists(path):
|
||||
# From local
|
||||
state_dict = torch.load(path, map_location)
|
||||
elif path.startswith('http'):
|
||||
# From url
|
||||
state_dict = load_state_dict_from_url(path,
|
||||
map_location=map_location,
|
||||
check_hash=False)
|
||||
else:
|
||||
raise Exception(f'Cannot find {path} when load pretrained')
|
||||
|
||||
return load_pretrained_dict(model, state_dict, logger, sub_level=sub_level)
|
||||
|
||||
|
||||
def _auto_drop_invalid(model: torch.nn.Module, state_dict: dict, logger=None):
|
||||
""" Strip unmatched parameters in state_dict, e.g. shape not matched, type not matched.
|
||||
|
||||
Args:
|
||||
model (torch.nn.Module):
|
||||
state_dict (dict):
|
||||
logger (logging.Logger, None):
|
||||
|
||||
Returns:
|
||||
A new state dict.
|
||||
"""
|
||||
ret_dict = state_dict.copy()
|
||||
invalid_msgs = []
|
||||
for key, value in model.state_dict().items():
|
||||
if key in state_dict:
|
||||
# Check shape
|
||||
new_value = state_dict[key]
|
||||
if value.shape != new_value.shape:
|
||||
invalid_msgs.append(
|
||||
f'{key}: invalid shape, dst {value.shape} vs. src {new_value.shape}'
|
||||
)
|
||||
ret_dict.pop(key)
|
||||
elif value.dtype != new_value.dtype:
|
||||
invalid_msgs.append(
|
||||
f'{key}: invalid dtype, dst {value.dtype} vs. src {new_value.dtype}'
|
||||
)
|
||||
ret_dict.pop(key)
|
||||
if len(invalid_msgs) > 0:
|
||||
warning_msg = 'ignore keys from source: \n' + '\n'.join(invalid_msgs)
|
||||
if logger:
|
||||
logger.warning(warning_msg)
|
||||
else:
|
||||
import warnings
|
||||
warnings.warn(warning_msg)
|
||||
return ret_dict
|
||||
|
||||
|
||||
def load_pretrained_dict(model: torch.nn.Module,
|
||||
state_dict: dict,
|
||||
logger=None,
|
||||
sub_level=None):
|
||||
""" Load parameters to model with
|
||||
1. Sub name by revise_keys For DataParallelModel or DistributeParallelModel.
|
||||
2. Load 'state_dict' again if possible by key 'state_dict' or 'model_state'.
|
||||
3. Take sub level keys from source, e.g. load 'backbone' part from a classifier into a backbone model.
|
||||
4. Auto remove invalid parameters from source.
|
||||
5. Log or warning if unexpected key exists or key misses.
|
||||
|
||||
Args:
|
||||
model (torch.nn.Module):
|
||||
state_dict (dict): dict of parameters
|
||||
logger (logging.Logger, None):
|
||||
sub_level (str, optional): If not None, parameters with key startswith sub_level will remove the prefix
|
||||
to fit actual model keys. This action happens if user want to load sub module parameters
|
||||
into a sub module model.
|
||||
"""
|
||||
revise_keys = [(r'^module\.', '')]
|
||||
|
||||
if 'state_dict' in state_dict:
|
||||
state_dict = state_dict['state_dict']
|
||||
if 'model_state' in state_dict:
|
||||
state_dict = state_dict['model_state']
|
||||
|
||||
for p, r in revise_keys:
|
||||
state_dict = {re.sub(p, r, k): v for k, v in state_dict.items()}
|
||||
|
||||
if sub_level:
|
||||
sub_level = sub_level if sub_level.endswith('.') else (sub_level + '.')
|
||||
sub_level_len = len(sub_level)
|
||||
state_dict = {
|
||||
key[sub_level_len:]: value
|
||||
for key, value in state_dict.items() if key.startswith(sub_level)
|
||||
}
|
||||
|
||||
state_dict = _auto_drop_invalid(model, state_dict, logger=logger)
|
||||
|
||||
load_status = model.load_state_dict(state_dict, strict=False)
|
||||
unexpected_keys = load_status.unexpected_keys
|
||||
missing_keys = load_status.missing_keys
|
||||
err_msgs = []
|
||||
if unexpected_keys:
|
||||
err_msgs.append('unexpected key in source '
|
||||
f'state_dict: {", ".join(unexpected_keys)}\n')
|
||||
if missing_keys:
|
||||
err_msgs.append('missing key in source '
|
||||
f'state_dict: {", ".join(missing_keys)}\n')
|
||||
err_msgs = '\n'.join(err_msgs)
|
||||
|
||||
if len(err_msgs) > 0:
|
||||
if logger:
|
||||
logger.warning(err_msgs)
|
||||
else:
|
||||
import warnings
|
||||
warnings.warn(err_msgs)
|
||||
|
||||
|
||||
def count_params(model):
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
return total_params
|
||||
|
||||
|
||||
def init_weights(module):
|
||||
if isinstance(module, (nn.Linear, nn.Embedding)):
|
||||
module.weight.data.normal_(mean=0.0, std=0.02)
|
||||
elif isinstance(module, nn.LayerNorm):
|
||||
module.bias.data.zero_()
|
||||
module.weight.data.fill_(1.0)
|
||||
if isinstance(module, nn.Linear) and module.bias is not None:
|
||||
module.bias.data.zero_()
|
||||
@@ -0,0 +1,391 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
import os.path
|
||||
from numbers import Number
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
|
||||
def check_legal_type(data):
|
||||
if isinstance(data, str) or isinstance(data, Number):
|
||||
return True
|
||||
elif isinstance(data, dict):
|
||||
for k, v in data.items():
|
||||
if not check_legal_type(v):
|
||||
return False
|
||||
return True
|
||||
elif isinstance(data, list):
|
||||
for v in data:
|
||||
if not check_legal_type(v):
|
||||
return False
|
||||
return True
|
||||
else:
|
||||
return False
|
||||
|
||||
|
||||
def register_data(probe_data: dict, key_prefix=''):
|
||||
ret_data = {}
|
||||
dist_data = {}
|
||||
for k, v in probe_data.items():
|
||||
key = f'{key_prefix}_{k}'
|
||||
if isinstance(v, torch.Tensor) or isinstance(v, np.ndarray):
|
||||
ret_data[key] = ProbeData(v)
|
||||
elif isinstance(v, ProbeData):
|
||||
ret_data[key] = v
|
||||
else:
|
||||
if not check_legal_type(v):
|
||||
raise f'The datatype of {key} should be included in [array, tensor, number, str] or the dict or ' \
|
||||
f'list of (number, str); if you want register the list of image, please use ProbeData instance.'
|
||||
ret_data[key] = ProbeData(v)
|
||||
if ret_data[key].view_distribute:
|
||||
dist_data[key] = ret_data[key].distribute
|
||||
|
||||
return ret_data, dist_data
|
||||
|
||||
|
||||
def merge_gathered_probe(all_gathered_data):
|
||||
'''
|
||||
Merge the gathered data on rank_0.
|
||||
Returns:
|
||||
The merged data.
|
||||
|
||||
'''
|
||||
for key, gathered_data in all_gathered_data.items():
|
||||
# Must be the list of ProbeData.
|
||||
if isinstance(gathered_data, list):
|
||||
for v in gathered_data:
|
||||
if not isinstance(v, ProbeData):
|
||||
all_gathered_data[key] = gathered_data
|
||||
# Must be the gathered data.
|
||||
ret_data = gathered_data[0]
|
||||
if not isinstance(ret_data.data,
|
||||
list) and (isinstance(ret_data.data, np.ndarray)
|
||||
or isinstance(ret_data.data, dict)
|
||||
or check_legal_type(ret_data.data)):
|
||||
new_data = [v.data for v in gathered_data]
|
||||
if ret_data.build_label is not None:
|
||||
ret_data.build_label = [
|
||||
v.build_label for v in gathered_data
|
||||
]
|
||||
all_gathered_data[key] = ProbeData(
|
||||
new_data,
|
||||
is_image=ret_data.is_image,
|
||||
build_html=ret_data.build_html,
|
||||
build_label=ret_data.build_label,
|
||||
view_distribute=ret_data.view_distribute)
|
||||
elif isinstance(ret_data.data, list):
|
||||
if ret_data.build_label is not None:
|
||||
if isinstance(ret_data.build_label, str):
|
||||
ret_data.build_label = [
|
||||
ret_data.build_label for _ in ret_data.data
|
||||
]
|
||||
for v in gathered_data[1:]:
|
||||
ret_data.data += v.data
|
||||
if ret_data.build_label is not None:
|
||||
if isinstance(v.build_label, str):
|
||||
ret_data.build_label.extend(
|
||||
[v.build_label for _ in v.data])
|
||||
ret_data.build_label.extend(v.build_label)
|
||||
all_gathered_data[key] = ProbeData(
|
||||
ret_data.data,
|
||||
is_image=ret_data.is_image,
|
||||
build_html=ret_data.build_html,
|
||||
build_label=ret_data.build_label,
|
||||
view_distribute=ret_data.view_distribute)
|
||||
else:
|
||||
all_gathered_data[key] = gathered_data
|
||||
return all_gathered_data
|
||||
|
||||
|
||||
class ProbeData():
|
||||
def __init__(self,
|
||||
data,
|
||||
is_image=False,
|
||||
build_html=False,
|
||||
build_label=None,
|
||||
view_distribute=False):
|
||||
''' Probe Data Initialize.
|
||||
We only support basic types such as [torch.Tensor, numpy.ndarray, number, str],
|
||||
or [dict, list] of [dict, list,
|
||||
number, str], or [dict, list] of [tensor, array]
|
||||
'''
|
||||
data = copy.deepcopy(data)
|
||||
self.basic_type = True
|
||||
self._distribute_dict = {}
|
||||
if view_distribute:
|
||||
is_legal = True
|
||||
if isinstance(data, str) or isinstance(data, Number):
|
||||
is_legal = True
|
||||
if data in self._distribute_dict:
|
||||
self._distribute_dict[data] += 1
|
||||
else:
|
||||
self._distribute_dict[data] = 1
|
||||
elif isinstance(data, list):
|
||||
for v in data:
|
||||
if isinstance(v, str) or isinstance(v, Number):
|
||||
is_legal = True
|
||||
if v in self._distribute_dict:
|
||||
self._distribute_dict[v] += 1
|
||||
else:
|
||||
self._distribute_dict[v] = 1
|
||||
else:
|
||||
is_legal = False
|
||||
elif isinstance(data, dict):
|
||||
for k, v in data.items():
|
||||
if isinstance(v, str) or isinstance(v, Number):
|
||||
is_legal = True
|
||||
n_k = f'{k}_{v}'
|
||||
if n_k in self._distribute_dict:
|
||||
self._distribute_dict[n_k] += 1
|
||||
else:
|
||||
self._distribute_dict[n_k] = 1
|
||||
else:
|
||||
is_legal = False
|
||||
else:
|
||||
is_legal = False
|
||||
if not is_legal:
|
||||
print('Unsurpport data type', data)
|
||||
assert is_legal
|
||||
self.view_distribute = view_distribute
|
||||
if isinstance(data, torch.Tensor):
|
||||
self.data = data.detach().cpu().numpy()
|
||||
elif isinstance(data, np.ndarray):
|
||||
self.data = data
|
||||
elif isinstance(data, dict):
|
||||
for k, v in data.items():
|
||||
if not check_legal_type(v):
|
||||
if isinstance(v, torch.Tensor):
|
||||
data[k] = v.detach().cpu().numpy()
|
||||
self.basic_type = False
|
||||
elif isinstance(v, np.ndarray):
|
||||
data[k] = v
|
||||
self.basic_type = False
|
||||
else:
|
||||
raise f'Unsupport data type for {v}'
|
||||
self.data = data
|
||||
elif isinstance(data, list):
|
||||
for idx, v in enumerate(data):
|
||||
if not check_legal_type(v):
|
||||
if isinstance(v, torch.Tensor):
|
||||
data[idx] = v.detach().cpu().numpy()
|
||||
self.basic_type = False
|
||||
elif isinstance(v, np.ndarray):
|
||||
data[idx] = v
|
||||
self.basic_type = False
|
||||
else:
|
||||
raise f'Unsupport data type for {v}'
|
||||
self.data = data
|
||||
elif check_legal_type(data):
|
||||
self.data = data
|
||||
else:
|
||||
raise f'Unsupport data type for {data}'
|
||||
|
||||
self.is_image = is_image
|
||||
self.build_html = build_html
|
||||
|
||||
if self.build_html:
|
||||
assert build_label is not None
|
||||
if isinstance(self.data, str):
|
||||
assert isinstance(build_label, str)
|
||||
if isinstance(self.data, list):
|
||||
assert isinstance(build_label, str) or isinstance(
|
||||
build_label, list)
|
||||
if isinstance(self.data, dict):
|
||||
assert isinstance(build_label, str) or isinstance(
|
||||
build_label, dict)
|
||||
self.build_label = build_label
|
||||
|
||||
def save_image(self, file_prefix, images):
|
||||
np_shape = images.shape
|
||||
# 4D
|
||||
shape_str = '_'.join([str(v) for v in np_shape])
|
||||
if len(np_shape) == 4:
|
||||
# channel is 1 or 3
|
||||
if np_shape[-1] == 1 or np_shape[-1] == 3:
|
||||
if np_shape[-1] == 1:
|
||||
images = images.reshape(images.shape[:-1])
|
||||
file_list = []
|
||||
for idx in range(np_shape[0]):
|
||||
file_path = file_prefix + f'_probe_{idx}_[{shape_str}].jpg'
|
||||
with FS.put_to(file_path) as local_path:
|
||||
Image.fromarray(images[idx, ...]).save(local_path)
|
||||
file_list.append(file_path)
|
||||
return file_list
|
||||
else:
|
||||
raise f"Ensure your data's dim is BWHC, and channel is 1 or 3 for {file_prefix}"
|
||||
elif len(np_shape) == 3:
|
||||
if np_shape[-1] == 1 or np_shape[-1] == 3:
|
||||
if np_shape[-1] == 1:
|
||||
images = images.reshape(images.shape[:-1])
|
||||
file_path = file_prefix + f'_probe_[{shape_str}].jpg'
|
||||
with FS.put_to(file_path) as local_path:
|
||||
Image.fromarray(images).save(local_path)
|
||||
return file_path
|
||||
else:
|
||||
images = images.reshape(list(images.shape) + [1])
|
||||
return self.save_image(file_prefix, images)
|
||||
elif len(np_shape) == 2:
|
||||
file_path = file_prefix + f'_probe_[{shape_str}].jpg'
|
||||
with FS.put_to(file_path) as local_path:
|
||||
Image.fromarray(images).save(local_path)
|
||||
return file_path
|
||||
else:
|
||||
raise f"Ensure your data's dim is BWHC or WHC or WH, and channel is 1 or 3 for {file_prefix}"
|
||||
|
||||
def save_npy(self, file_prefix, data):
|
||||
shape_str = '_'.join([str(v) for v in data.shape])
|
||||
file_path = file_prefix + f'_{shape_str}.npy'
|
||||
with FS.put_to(file_path) as local_path:
|
||||
np.save(local_path, data)
|
||||
return file_path
|
||||
|
||||
def save_html(self, html_prefix, ret_data, ret_label):
|
||||
height = 600
|
||||
with FS.put_to(html_prefix) as local_path:
|
||||
with open(local_path, 'w') as f:
|
||||
f.writelines('<meta charset="utf-8">\n')
|
||||
f.writelines('<style>input{height:' + f'{height}px;' +
|
||||
'opacity:1.0;}</style>\n')
|
||||
f.writelines('<br><hr/>\n')
|
||||
all_ranks = list()
|
||||
for save_id, save_data in enumerate(zip(ret_data, ret_label)):
|
||||
save_path, save_label = save_data
|
||||
one_rank = '<table><tr>'
|
||||
for idx, one_data in enumerate(zip(save_path, save_label)):
|
||||
one_path, one_label = one_data
|
||||
one_label = one_label.replace('<', '<').replace(
|
||||
'>', '>')
|
||||
url = FS.get_url(one_path,
|
||||
lifecycle=3600 * 365 * 24).replace(
|
||||
'.oss-internal.aliyun-inc.',
|
||||
'.oss.aliyuncs.')
|
||||
one_rank += (
|
||||
f'<td align="center"><input type="image" src="{url}" >'
|
||||
f'<br><font size="4"><strong>{save_id}-{idx}|{one_label}<strong></font><br/></td>'
|
||||
)
|
||||
one_rank += '</tr></table><hr/>'
|
||||
all_ranks.append(one_rank)
|
||||
f.writelines('\n'.join(all_ranks))
|
||||
return html_prefix
|
||||
|
||||
@property
|
||||
def distribute(self):
|
||||
return self._distribute_dict
|
||||
|
||||
def to_log(self, prefix=None):
|
||||
if isinstance(self.data, np.ndarray):
|
||||
if prefix is None:
|
||||
raise 'You should provide the save prefix for array sample.'
|
||||
# save jpg
|
||||
if self.is_image:
|
||||
ret_data = self.save_image(prefix, self.data)
|
||||
if isinstance(ret_data, list):
|
||||
ret_data = [ret_data]
|
||||
if self.build_html:
|
||||
ret_label = []
|
||||
if isinstance(self.build_label, str):
|
||||
ret_label.append(
|
||||
[self.build_label for _ in ret_data[0]])
|
||||
else:
|
||||
ret_label.append(self.build_label)
|
||||
if not len(ret_data[0]) == len(ret_label[0]):
|
||||
raise f"The {prefix} label's length should be equal with 1st dim {self.data.shape[0]}."
|
||||
html_prefix = prefix + '_probe.html'
|
||||
html_file = self.save_html(html_prefix, ret_data,
|
||||
ret_label)
|
||||
return {'ori_file': ret_data, 'html': html_file}
|
||||
else:
|
||||
return {'ori_file': ret_data}
|
||||
else:
|
||||
return ret_data
|
||||
else:
|
||||
ret_data = self.save_npy(prefix, self.data)
|
||||
return ret_data
|
||||
elif isinstance(self.data, list):
|
||||
if not self.basic_type:
|
||||
ret_data = []
|
||||
ret_label = []
|
||||
for idx, v in enumerate(self.data):
|
||||
prefix_path = os.path.join(prefix, f'{idx}')
|
||||
if self.is_image:
|
||||
ret_images = self.save_image(prefix_path, v)
|
||||
ret_data.append(ret_images if isinstance(
|
||||
ret_images, list) else [ret_images])
|
||||
if self.build_html:
|
||||
if isinstance(ret_images, list):
|
||||
if isinstance(self.build_label, str):
|
||||
ret_label.append(
|
||||
[self.build_label for _ in ret_images])
|
||||
elif isinstance(self.build_label[idx], list):
|
||||
assert len(self.build_label[idx]) == len(
|
||||
ret_images)
|
||||
ret_label.append(self.build_label[idx])
|
||||
else:
|
||||
ret_label.append([
|
||||
self.build_label[idx]
|
||||
for _ in ret_images
|
||||
])
|
||||
else:
|
||||
if isinstance(self.build_label, str):
|
||||
ret_label.append([self.build_label])
|
||||
else:
|
||||
ret_label.append([self.build_label[idx]])
|
||||
else:
|
||||
ret_data.append(self.save_npy(prefix_path, v))
|
||||
if self.is_image and self.build_html:
|
||||
html_prefix = prefix + '_probe.html'
|
||||
html_file = self.save_html(html_prefix, ret_data,
|
||||
ret_label)
|
||||
return {'ori_file': ret_data, 'html': html_file}
|
||||
else:
|
||||
return {'ori_file': ret_data}
|
||||
else:
|
||||
return self.data
|
||||
elif isinstance(self.data, dict):
|
||||
if not self.basic_type:
|
||||
ret_data = []
|
||||
ret_label = []
|
||||
for k, v in self.data:
|
||||
prefix_path = os.path.join(prefix, f'{k}_')
|
||||
if self.is_image:
|
||||
ret_images = self.save_image(prefix_path, v)
|
||||
if isinstance(ret_images, list):
|
||||
ret_data.append(ret_images)
|
||||
else:
|
||||
ret_data.append([ret_images])
|
||||
if self.build_html:
|
||||
if isinstance(ret_images, list):
|
||||
if isinstance(self.build_label, str):
|
||||
ret_label.append(
|
||||
[self.build_label for _ in ret_images])
|
||||
elif isinstance(self.build_label[k], list):
|
||||
assert len(
|
||||
self.build_label[k]) == len(ret_images)
|
||||
ret_label.append(self.build_label[k])
|
||||
else:
|
||||
ret_label.append([
|
||||
self.build_label[k] for _ in ret_images
|
||||
])
|
||||
else:
|
||||
if isinstance(self.build_label, str):
|
||||
ret_label.append([self.build_label])
|
||||
else:
|
||||
ret_label.append([self.build_label[k]])
|
||||
else:
|
||||
ret_data.append(self.save_npy(prefix_path, v))
|
||||
if self.is_image and self.build_html:
|
||||
html_prefix = prefix + '_probe.html'
|
||||
html_file = self.save_html(html_prefix, ret_data,
|
||||
ret_label)
|
||||
return {'ori_file': ret_data, 'html': html_file}
|
||||
else:
|
||||
return {'ori_file': ret_data}
|
||||
else:
|
||||
return self.data
|
||||
else:
|
||||
return self.data
|
||||
@@ -0,0 +1,212 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
# Modified Based on the following original code.
|
||||
|
||||
# Registry class & build_from_config function partially modified from
|
||||
# https://github.com/open-mmlab/mmcv/blob/master/mmcv/utils/registry.py
|
||||
# Copyright 2018-2020 Open-MMLab. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
import inspect
|
||||
import sys
|
||||
import warnings
|
||||
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
|
||||
old_python_version = '3.6' in sys.version
|
||||
if old_python_version:
|
||||
|
||||
def deep_copy(obj):
|
||||
return obj
|
||||
else:
|
||||
import copy
|
||||
|
||||
def deep_copy(obj):
|
||||
return copy.deepcopy(obj)
|
||||
|
||||
|
||||
def build_from_config(cfg, registry, logger=None, *args, **kwargs):
|
||||
""" Default builder function.
|
||||
|
||||
Args:
|
||||
cfg (objective attribution): A set of objective attirbutions
|
||||
which contain parameters passes to target class or function.
|
||||
Must contains key 'type', indicates the target class or function name.
|
||||
registry (Registry): An registry to search target class or function.
|
||||
kwargs (dict, optional): Other params not in config dict.
|
||||
|
||||
Returns:
|
||||
Target class object or object returned by invoking function.
|
||||
|
||||
Raises:
|
||||
TypeError:
|
||||
KeyError:
|
||||
Exception:
|
||||
"""
|
||||
from scepter.modules.utils.config import Config
|
||||
if not isinstance(cfg, Config):
|
||||
raise TypeError(f'config must be type dict, got {type(cfg)}')
|
||||
if not cfg.have('NAME'):
|
||||
raise KeyError(f'config must contain key NAME, got {cfg}')
|
||||
if not isinstance(registry, Registry):
|
||||
raise TypeError(
|
||||
f'registry must be type Registry, got {type(registry)}')
|
||||
|
||||
cfg = deep_copy(cfg)
|
||||
req_type = cfg.get('NAME')
|
||||
if isinstance(req_type, str):
|
||||
req_type_entry = registry.get(req_type)
|
||||
if req_type_entry is None:
|
||||
raise KeyError(f'{req_type} not found in {registry.name} registry')
|
||||
|
||||
if kwargs is not None:
|
||||
cfg._update_dict(kwargs)
|
||||
|
||||
if inspect.isclass(req_type_entry):
|
||||
try:
|
||||
return req_type_entry(cfg, logger=logger, *args, **kwargs)
|
||||
except Exception as e:
|
||||
raise Exception(f'Failed to init class {req_type_entry}, with {e}')
|
||||
elif inspect.isfunction(req_type_entry):
|
||||
try:
|
||||
return req_type_entry(cfg, logger=logger, *args, **kwargs)
|
||||
except Exception as e:
|
||||
raise Exception(
|
||||
f'Failed to invoke function {req_type_entry}, with {e}')
|
||||
else:
|
||||
raise TypeError(
|
||||
f'type must be str or class, got {type(req_type_entry)}')
|
||||
|
||||
|
||||
REGISTRY_LIST = []
|
||||
|
||||
|
||||
class Registry(object):
|
||||
""" A registry maps key to classes or functions.
|
||||
|
||||
Example:
|
||||
# >>> MODELS = Registry('MODELS')
|
||||
# >>> @MODELS.register_class()
|
||||
# >>> class ResNet(object):
|
||||
# >>> pass
|
||||
# >>> config = Config(cfg_dict = {"NAME":"ResNet"})
|
||||
# >>> resnet = MODELS.build(config)
|
||||
# >>>
|
||||
# >>> import torchvision
|
||||
# >>> @MODELS.register_function("InceptionV3")
|
||||
# >>> def get_inception_v3(pretrained=False, progress=True):
|
||||
# >>> return torchvision.model.inception_v3(pretrained=pretrained, progress=progress)
|
||||
# >>> config = Config(cfg_dict = {"NAME":"InceptionV3"})
|
||||
# >>> inception_v3 = MODELS.build(config)
|
||||
|
||||
Args:
|
||||
name (str): Registry name.
|
||||
build_func (func, None): Instance construct function. Default is build_from_config.
|
||||
allow_types (tuple): Indicates how to construct the instance, by constructing class or invoking function.
|
||||
"""
|
||||
def __init__(self,
|
||||
name,
|
||||
build_func=None,
|
||||
common_para=None,
|
||||
allow_types=('class', 'function')):
|
||||
self.name = name
|
||||
self.allow_types = allow_types
|
||||
self.class_map = {}
|
||||
self.func_map = {}
|
||||
self.common_para = common_para
|
||||
self.build_func = build_func or build_from_config
|
||||
REGISTRY_LIST.append(self)
|
||||
|
||||
def get(self, req_type):
|
||||
return self.class_map.get(req_type) or self.func_map.get(req_type)
|
||||
|
||||
def build(self, cfg, logger=None, *args, **kwargs):
|
||||
return self.build_func(cfg,
|
||||
registry=self,
|
||||
logger=logger,
|
||||
*args,
|
||||
**kwargs)
|
||||
|
||||
def register_class(self, name=None):
|
||||
def _register(cls):
|
||||
if not inspect.isclass(cls):
|
||||
raise TypeError(f'Module must be type class, got {type(cls)}')
|
||||
if 'class' not in self.allow_types:
|
||||
raise TypeError(
|
||||
f'Register {self.name} only allows type {self.allow_types}, got class'
|
||||
)
|
||||
module_name = name or cls.__name__
|
||||
if module_name in self.class_map:
|
||||
warnings.warn(
|
||||
f'Class {module_name} already registered by {self.class_map[module_name]}, '
|
||||
f'will be replaced by {cls}')
|
||||
self.class_map[module_name] = cls
|
||||
return cls
|
||||
|
||||
return _register
|
||||
|
||||
def register_function(self, name=None):
|
||||
def _register(func):
|
||||
if not inspect.isfunction(func):
|
||||
raise TypeError(
|
||||
f'Registry must be type function, got {type(func)}')
|
||||
if 'function' not in self.allow_types:
|
||||
raise TypeError(
|
||||
f'Registry {self.name} only allows type {self.allow_types}, got function'
|
||||
)
|
||||
func_name = name or func.__name__
|
||||
if func_name in self.class_map:
|
||||
warnings.warn(
|
||||
f'Function {func_name} already registered by {self.func_map[func_name]}, '
|
||||
f'will be replaced by {func}')
|
||||
self.func_map[func_name] = func
|
||||
return func
|
||||
|
||||
return _register
|
||||
|
||||
def _list(self):
|
||||
keys = sorted(list(self.class_map.keys()) + list(self.func_map.keys()))
|
||||
descriptions = []
|
||||
for key in keys:
|
||||
if key in self.class_map:
|
||||
descriptions.append(f'{key}: {self.class_map[key]}')
|
||||
else:
|
||||
descriptions.append(
|
||||
f"{key}: <function '{self.func_map[key].__module__}.{self.func_map[key].__name__}'>"
|
||||
)
|
||||
return '\n'.join(descriptions)
|
||||
|
||||
def __repr__(self):
|
||||
description = self._list()
|
||||
description = '\n'.join(['\t' + s for s in description.split('\n')])
|
||||
return f'{self.__class__.__name__} [{self.name}], \n' + description
|
||||
|
||||
def get_config_template(self, name):
|
||||
common_yaml_str = ''
|
||||
if self.common_para is not None:
|
||||
common_yaml_str += 'The following para are used for this class.\n'
|
||||
common_yaml_str += dict_to_yaml('common_parameter',
|
||||
__class__.__name__,
|
||||
self.common_para,
|
||||
set_name=False)
|
||||
|
||||
req_type_entry = self.get(name)
|
||||
if req_type_entry is None:
|
||||
raise KeyError(f'{name} not found in {self.name} registry')
|
||||
if inspect.isclass(req_type_entry):
|
||||
return req_type_entry.get_config_template() + common_yaml_str
|
||||
elif inspect.isfunction(req_type_entry):
|
||||
return '{} is a function!'.format(name)
|
||||
else:
|
||||
return 'Unsurport object type!'
|
||||
@@ -0,0 +1,6 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from .frame_sampler import (FRAME_SAMPLERS, IntervalSampler, SegmentSampler,
|
||||
UniformSampler, do_frame_sample)
|
||||
from .video_reader import (EasyVideoReader, FramesReaderWrapper,
|
||||
VideoReaderWrapper)
|
||||
@@ -0,0 +1,165 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
"""
|
||||
FrameSampler.
|
||||
|
||||
Sample:
|
||||
1. give start & end time, num_frames, e.g. 16 frames from [1.0s ,3.0s],
|
||||
usually used in real applications.
|
||||
fixed args:
|
||||
`sample_type`='uniform';
|
||||
`vid_len` (int): valid total frame numbers in video;
|
||||
`vid_fps` (float): video fps;
|
||||
`num_frames` (int): number of frames to be extracted;
|
||||
extra args:
|
||||
|
||||
2. give a fixed clip duration, num_frames, e.g. 16 frames from a 2s clip.
|
||||
In train mode (`clip_id`=-1), this clip will be randomly sampled from video.
|
||||
In test mode (`clip_id`>=0), this clip is the center part of the video
|
||||
which acts the same as `DecodeVideoToTensor` op.
|
||||
fixed args: `sample_mode`='interval', `vid_len`, `vid_fps`, `num_frames`
|
||||
extra args: `clip_duration`, `clip_id`, `num_clips`=1
|
||||
|
||||
3. give a fixed clip duration, constant total clips, current clip index, num_frames,
|
||||
e.g. three 2-s clips will be sampled from the video, and 16 frames from the first clip.
|
||||
Usually used in multi-view test.
|
||||
In train mode (`clip_id`=-1), constant total clips will be ignored, so this will act the same b.
|
||||
In test mode (`clip_id`>=0), video is splitted into constant clips (uniformly and allow overlap,
|
||||
a 3s video splits into three 2s clips, [0, 2), [0.5, 2.5), [1.0, 3.0) ),
|
||||
then sample frames from one clip.
|
||||
fixed args: `sample_mode`='interval', `vid_len`, `vid_fps`, `num_frames`
|
||||
extra args: `clip_duration`, `clip_id`, `num_clips`
|
||||
|
||||
4. give num_frames, do segment-sampling, e.g. 16 frames from whole video, then splits the video into 16 segments,
|
||||
and sample one frame from each segment.
|
||||
In train mode (`clip_id`=-1), sample a frame randomly from a segment.
|
||||
In test mode (`clip_id`>=0), the center frame in each part will be chosen.
|
||||
fixed args: `sample_mode`='segment', `vid_len`, `vid_fps`, `num_frames`
|
||||
call args: `clip_id`, `num_clips`=1
|
||||
|
||||
5. give constant total clips, current clip index, num_frames,
|
||||
e.g. splits the video into 16 segments, split one segment into 3 parts,
|
||||
if clip_index=0, sample one frame from the first part, and loop 16 times.
|
||||
In train mode, sample a frame randomly from a segment.
|
||||
In test mode, int(`clip_id`/`num_clips` * segment_frames) will be chosen.
|
||||
fixed args: `sample_mode`='segment', `vid_len`, `vid_fps`, `num_frames`
|
||||
call args: `clip_id`, `num_clips`
|
||||
|
||||
Output:
|
||||
A list of frame indices (torch.Tensor)
|
||||
|
||||
"""
|
||||
import math
|
||||
import random
|
||||
|
||||
import torch
|
||||
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.registry import Registry
|
||||
|
||||
FRAME_SAMPLERS = Registry('FRAME_SAMPLERS')
|
||||
|
||||
|
||||
def do_frame_sample(sampling_type: str, vid_len: int, vid_fps: float,
|
||||
num_frames: int, **kwargs) -> torch.Tensor:
|
||||
params = dict(vid_len=vid_len,
|
||||
vid_fps=vid_fps,
|
||||
num_frames=num_frames,
|
||||
**kwargs)
|
||||
return FRAME_SAMPLERS.build(
|
||||
Config(cfg_dict={'NAME': sampling_type}, load=False))(**params)
|
||||
|
||||
|
||||
@FRAME_SAMPLERS.register_class('uniform')
|
||||
class UniformSampler(object):
|
||||
def __init__(self, cfg, logger=None):
|
||||
self.cfg = cfg
|
||||
self.logger = logger
|
||||
|
||||
def __call__(self,
|
||||
vid_len: int,
|
||||
vid_fps: float,
|
||||
num_frames: int,
|
||||
start_sec: float = 0,
|
||||
end_sec: float = -1) -> torch.Tensor:
|
||||
start_sec = max(start_sec, 0)
|
||||
if end_sec < 0:
|
||||
new_end_sec = vid_len / vid_fps
|
||||
else:
|
||||
new_end_sec = min(end_sec, vid_len / vid_fps)
|
||||
assert new_end_sec > start_sec, (
|
||||
f'end_sec should be greater then start_sec, '
|
||||
f'got end_sec={new_end_sec}, start_sec={start_sec}')
|
||||
end_sec = new_end_sec
|
||||
|
||||
start_idx = math.floor(start_sec / vid_fps)
|
||||
end_idx_exc = min(vid_len, math.ceil(end_sec / vid_fps))
|
||||
|
||||
index = torch.linspace(start_idx, end_idx_exc, num_frames)
|
||||
index = torch.clamp(index, 0, vid_len - 1).long()
|
||||
|
||||
return index
|
||||
|
||||
|
||||
@FRAME_SAMPLERS.register_class('interval')
|
||||
class IntervalSampler(object):
|
||||
def __init__(self, cfg, logger=None):
|
||||
self.cfg = cfg
|
||||
self.logger = logger
|
||||
|
||||
def __call__(self,
|
||||
vid_len: int,
|
||||
vid_fps: float,
|
||||
num_frames: int,
|
||||
clip_duration: float,
|
||||
clip_id: int = 0,
|
||||
num_clips: int = 1) -> torch.Tensor:
|
||||
if num_frames == 1:
|
||||
return torch.randint(0, vid_len, (1, ))
|
||||
|
||||
clip_len = int(clip_duration / vid_fps)
|
||||
max_idx = max(vid_len, clip_len, 0)
|
||||
|
||||
if clip_id == -1:
|
||||
start_idx = random.uniform(0, max_idx)
|
||||
else:
|
||||
if num_clips == 1:
|
||||
start_idx = max_idx / 2
|
||||
else:
|
||||
start_idx = max_idx * clip_id / num_clips
|
||||
|
||||
end_idx = start_idx + clip_len - 1
|
||||
index = torch.linspace(start_idx, end_idx, num_frames)
|
||||
index = torch.clamp(index, 0, vid_len - 1).long()
|
||||
return index
|
||||
|
||||
|
||||
@FRAME_SAMPLERS.register_class('segment')
|
||||
class SegmentSampler(object):
|
||||
def __init__(self, cfg, logger=None):
|
||||
self.cfg = cfg
|
||||
self.logger = logger
|
||||
|
||||
def __call__(self,
|
||||
vid_len: int,
|
||||
vid_fps: float,
|
||||
num_frames: int,
|
||||
clip_id: int = 0,
|
||||
num_clips: int = 1) -> torch.Tensor:
|
||||
index = torch.zeros(num_frames)
|
||||
index_range = torch.linspace(0, vid_len, num_frames + 1)
|
||||
for idx in range(num_frames):
|
||||
if clip_id == -1:
|
||||
index[idx] = random.uniform(index_range[idx],
|
||||
index_range[idx + 1])
|
||||
else:
|
||||
if num_clips == 1:
|
||||
index[idx] = (index_range[idx] + index_range[idx + 1]) / 2
|
||||
else:
|
||||
index[idx] = index_range[idx] + (
|
||||
index_range[idx + 1] -
|
||||
index_range[idx]) * (clip_id + 1) / num_clips
|
||||
|
||||
index = torch.round(torch.clamp(index, 0, vid_len - 1)).long()
|
||||
|
||||
return index
|
||||
@@ -0,0 +1,166 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import os
|
||||
from fractions import Fraction
|
||||
from typing import Callable, Optional, Union
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.utils.dlpack as dlpack
|
||||
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.modules.utils.video_reader.frame_sampler import do_frame_sample
|
||||
|
||||
|
||||
class _Wrapper(object):
|
||||
@property
|
||||
def len(self) -> int:
|
||||
raise NotImplementedError
|
||||
|
||||
@property
|
||||
def fps(self) -> float:
|
||||
raise NotImplementedError
|
||||
|
||||
@property
|
||||
def duration(self) -> float:
|
||||
raise NotImplementedError
|
||||
|
||||
def sample_frames(self, decode_list: torch.Tensor) -> torch.Tensor:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class VideoReaderWrapper(_Wrapper):
|
||||
def __init__(self, video_path):
|
||||
import decord
|
||||
self._video_path = video_path
|
||||
self._decoder_type = 'decord'
|
||||
self._vr = decord.VideoReader(self._video_path)
|
||||
|
||||
@property
|
||||
def len(self):
|
||||
return len(self._vr)
|
||||
|
||||
@property
|
||||
def fps(self):
|
||||
return self._vr.get_avg_fps()
|
||||
|
||||
@property
|
||||
def duration(self):
|
||||
return float(self.len) / self.fps
|
||||
|
||||
def __del__(self):
|
||||
if self._vr is not None:
|
||||
del self._vr
|
||||
self._vr = None
|
||||
|
||||
def sample_frames(self, decode_list: torch.Tensor) -> torch.Tensor:
|
||||
frames = dlpack.from_dlpack(
|
||||
self._vr.get_batch(decode_list).to_dlpack()).clone()
|
||||
return frames
|
||||
|
||||
|
||||
class FramesReaderWrapper(_Wrapper):
|
||||
def __init__(self, frame_dir: str, extract_fps: float, suffix='.jpg'):
|
||||
self._frame_dir = frame_dir
|
||||
self._extract_fps = extract_fps
|
||||
self._suffix = suffix
|
||||
self._frame_list = sorted([
|
||||
os.path.join(self._frame_dir, t)
|
||||
for t in os.listdir(self._frame_dir) if t.endswith(self._suffix)
|
||||
])
|
||||
self._frames = [None] * len(self._frame_list)
|
||||
|
||||
@property
|
||||
def len(self) -> int:
|
||||
return len(self._frame_list)
|
||||
|
||||
@property
|
||||
def fps(self):
|
||||
return self._extract_fps
|
||||
|
||||
@property
|
||||
def duration(self):
|
||||
return float(self.len) / self.fps
|
||||
|
||||
def _load_frame(self, idx):
|
||||
path = self._frame_list[idx]
|
||||
img = cv2.imread(path, cv2.IMREAD_COLOR)
|
||||
img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
|
||||
self._frames[idx] = img
|
||||
|
||||
def sample_frames(self, decode_list: torch.Tensor) -> torch.Tensor:
|
||||
ret = []
|
||||
for idx in decode_list.numpy():
|
||||
if self._frames[idx] is None:
|
||||
self._load_frame(idx)
|
||||
ret.append(self._frames[idx].copy())
|
||||
ret = np.asarray(ret)
|
||||
return torch.from_numpy(ret)
|
||||
|
||||
|
||||
class EasyVideoReader(object):
|
||||
""" A video reader which is easy to use in real applications.
|
||||
|
||||
Args:
|
||||
video_path (str): Path of video file.
|
||||
num_frames (int): Extract frames for one sample.
|
||||
clip_duration (Union[float, Fraction, str]): Clip duration to be extracted uniformly.
|
||||
overlap (Union[float, Fraction, str]): The offset (in secs) of
|
||||
the next clip overlaps the last clip, default is 0 no overlap.
|
||||
transforms (Optional[Callable]): Do transform operations, default is None.
|
||||
|
||||
"""
|
||||
def __init__(self,
|
||||
video_path: str,
|
||||
num_frames: int,
|
||||
clip_duration: Union[float, Fraction, str],
|
||||
overlap: Union[float, Fraction, str] = Fraction(0),
|
||||
transforms: Optional[Callable] = None):
|
||||
self._video_path: str = video_path
|
||||
self._num_frames: int = num_frames
|
||||
self._clip_duration: Fraction = Fraction(clip_duration)
|
||||
self._overlap: Fraction = Fraction(overlap)
|
||||
assert self._overlap < self._clip_duration, 'Overlap must be smaller than clip_duration!'
|
||||
self._transforms = transforms
|
||||
|
||||
self._last_end: Fraction = Fraction(0)
|
||||
|
||||
client = FS.get_fs_client(self._video_path)
|
||||
local_path = client.get_object_to_local_file(self._video_path)
|
||||
self._vr = VideoReaderWrapper(local_path)
|
||||
|
||||
def __iter__(self):
|
||||
return self
|
||||
|
||||
def __next__(self):
|
||||
start_sec = max(Fraction(0), self._last_end - self._overlap)
|
||||
end_sec = start_sec + self._clip_duration
|
||||
|
||||
if end_sec > self._vr.duration:
|
||||
del self._vr
|
||||
raise StopIteration
|
||||
|
||||
decode_list = do_frame_sample('uniform',
|
||||
self._vr.len,
|
||||
self._vr.fps,
|
||||
self._num_frames,
|
||||
start_sec=float(start_sec),
|
||||
end_sec=float(end_sec))
|
||||
output_tensor = self._vr.sample_frames(decode_list)
|
||||
|
||||
self._last_end = end_sec
|
||||
|
||||
output = {
|
||||
'video': output_tensor,
|
||||
'meta': {
|
||||
'video_path': self._video_path,
|
||||
'start_sec': float(start_sec),
|
||||
'end_sec': float(end_sec)
|
||||
}
|
||||
}
|
||||
|
||||
if self._transforms is not None:
|
||||
return self._transforms(output)
|
||||
return output
|
||||
Reference in New Issue
Block a user