new project from v0.0.1

This commit is contained in:
duanmurs@163.com
2023-12-28 18:23:39 +08:00
parent fda66e7ca6
commit 66f979f7e9
204 changed files with 37941 additions and 0 deletions
+3
View File
@@ -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
+604
View File
@@ -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
+88
View File
@@ -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
+21
View File
@@ -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
+455
View File
@@ -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()
+113
View File
@@ -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}")
+445
View File
@@ -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
+177
View File
@@ -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()
+103
View File
@@ -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)
+164
View File
@@ -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_()
+391
View File
@@ -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('<', '&lt;').replace(
'>', '&gt;')
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
+212
View File
@@ -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