update v1.1.0

This commit is contained in:
zeyinzi.jzyz
2024-10-21 00:35:53 +08:00
parent 7d6451efad
commit 0bba2c319d
148 changed files with 18476 additions and 1356 deletions
+32 -15
View File
@@ -20,7 +20,7 @@ _SECURE_KEYWORDS = [
_SECURE_VALUEWORDS = ['oss://', 'oss-'] # -> "#####"
def dict_to_yaml(module_name, name, json_config, set_name=False):
def dict_to_yaml(module_name, name, json_config, set_name=False, exclude_keys=[]):
'''
{ "ENV" :
{ "description" : "",
@@ -70,6 +70,9 @@ def dict_to_yaml(module_name, name, json_config, set_name=False):
yaml_str = ''
# print(level_num, json_config)
if isinstance(json_config, dict):
for key in exclude_keys:
if key in json_config:
json_config.pop(key)
if 'value' in json_config:
value = json_config['value']
if isinstance(value, dict):
@@ -322,39 +325,39 @@ class Config(object):
'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.'
)
# 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.'
)
# 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()}')
# 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}')
# 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}."
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'
f'Loading config from [{file_name}] as json 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'
f'Loading config from [{file_name}] as yaml file.'
)
else:
self.logger.info(
@@ -616,6 +619,20 @@ class Config(object):
config_new[key] = val
return config_new
def get_uppercase_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.upper()] = self.get_uppercase_dict(val)
else:
config_new[key.upper()] = val
else:
config_new[key] = val
return config_new
@staticmethod
def get_plain_cfg(cfg=None):
if isinstance(cfg, Config):
+7 -1
View File
@@ -83,7 +83,13 @@ def transfer_data_to_cuda(data_map: dict) -> dict:
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])
ret_data = []
for t in value:
if not isinstance(t, dict):
ret_data.append(transfer_data_to_cuda({'data': t})['data'])
else:
ret_data.append(transfer_data_to_cuda(t))
ret[key] = type(value)(ret_data)
else:
ret[key] = value
return ret
+362 -15
View File
@@ -4,12 +4,16 @@ import functools
import os
import pickle
import random
import socket
import warnings
from collections import OrderedDict
from datetime import timedelta
import numpy as np
import torch
import torch.distributed as dist
from torch.autograd import Function
from scepter.modules.utils.model import StdMsg
__all__ = [
@@ -202,6 +206,96 @@ def barrier():
dist.barrier()
def all_gather(tensor, uniform_size=True, group=None, **kwargs):
world_size = dist.get_world_size(group)
if world_size == 1:
return [tensor]
assert tensor.is_contiguous(), \
'ops.all_gather requires the tensor to be contiguous()'
if uniform_size:
tensor_list = [torch.empty_like(tensor) for _ in range(world_size)]
dist.all_gather(tensor_list, tensor, group, **kwargs)
return tensor_list
else:
# collect tensor shapes across GPUs
shape = tuple(tensor.shape)
shape_list = generalized_all_gather(shape, group)
# flatten the tensor
tensor = tensor.reshape(-1)
size = int(np.prod(shape))
size_list = [int(np.prod(u)) for u in shape_list]
max_size = max(size_list)
# pad to maximum size
if size != max_size:
padding = tensor.new_zeros(max_size - size)
tensor = torch.cat([tensor, padding], dim=0)
# all_gather
tensor_list = [torch.empty_like(tensor) for _ in range(world_size)]
dist.all_gather(tensor_list, tensor, group, **kwargs)
# reshape tensors
tensor_list = [
t[:n].view(s)
for t, n, s in zip(tensor_list, size_list, shape_list)
]
return tensor_list
def _pad_to_largest_tensor(tensor, group):
world_size = dist.get_world_size(group=group)
assert world_size >= 1, \
'gather/all_gather must be called from ranks within' \
'the give group!'
local_size = torch.tensor([tensor.numel()],
dtype=torch.int64,
device=tensor.device)
size_list = [
torch.zeros([1], dtype=torch.int64, device=tensor.device)
for _ in range(world_size)
]
# gather tensors and compute the maximum size
dist.all_gather(size_list, local_size, group=group)
size_list = [int(size.item()) for size in size_list]
max_size = max(size_list)
# pad tensors to the same size
if local_size != max_size:
padding = torch.zeros((max_size - local_size, ),
dtype=torch.uint8,
device=tensor.device)
tensor = torch.cat((tensor, padding), dim=0)
return size_list, tensor
def generalized_all_gather(data, group=None):
if dist.get_world_size(group) == 1:
return [data]
if group is None:
group = get_global_gloo_group()
tensor = _serialize_to_tensor(data, group)
size_list, tensor = _pad_to_largest_tensor(tensor, group)
max_size = max(size_list)
# receiving tensors from all ranks
tensor_list = [
torch.empty((max_size, ), dtype=torch.uint8, device=tensor.device)
for _ in size_list
]
dist.all_gather(tensor_list, tensor, group=group)
data_list = []
for size, tensor in zip(size_list, tensor_list):
buffer = tensor.cpu().numpy().tobytes()[:size]
data_list.append(pickle.loads(buffer))
return data_list
@functools.lru_cache()
def get_global_gloo_group():
backend = dist.get_backend()
@@ -223,12 +317,14 @@ def reduce_scatter(output,
def all_reduce(tensor, op=dist.ReduceOp.SUM, group=None, **kwargs):
if we.is_distributed:
return dist.all_reduce(tensor, op, group, **kwargs)
dist.all_reduce(tensor, op, group, **kwargs)
return tensor
def reduce(tensor, dst, op=dist.ReduceOp.SUM, group=None, **kwargs):
if we.is_distributed:
return dist.reduce(tensor, dst, op, group, **kwargs)
dist.reduce(tensor, dst, op, group, **kwargs)
return tensor
def _serialize_to_tensor(data):
@@ -243,6 +339,24 @@ def _unserialize_from_tensor(recv_data):
return pickle.loads(buffer)
def find_free_port():
# Copied from https://github.com/facebookresearch/detectron2/blob/main/detectron2/engine/launch.py # noqa: E501
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
# Binding to port 0 will cause the OS to find an available port for us
sock.bind(('', 0))
port = sock.getsockname()[1]
sock.close()
# NOTE: there is still a chance the port could be taken by other processes.
return port
def is_free_port(port):
ips = socket.gethostbyname_ex(socket.gethostname())[-1]
ips.append('localhost')
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
return all(s.connect_ex((ip, port)) != 0 for ip in ips)
def send(tensor, dst, group=None, **kwargs):
if we.is_distributed:
assert tensor.is_contiguous(
@@ -288,6 +402,144 @@ def shared_random_seed():
return all_seeds[0]
def all_to_all(x, scatter_dim, gather_dim, group=None, **kwargs):
"""
`scatter` along one dimension and `gather` along another.
"""
world_size = dist.get_world_size(group) if we.is_distributed else 1
if world_size > 1:
inputs = [u.contiguous() for u in x.chunk(world_size, dim=scatter_dim)]
outputs = [torch.empty_like(u) for u in inputs]
dist.all_to_all(outputs, inputs, group=group, **kwargs)
x = torch.cat(outputs, dim=gather_dim).contiguous()
return x
def _split(input, dim, group):
# skip if world_size == 1
rank = dist.get_rank(group=group)
world_size = dist.get_world_size(group=group)
if world_size == 1:
return input
# split sequence
assert input.size(dim) % world_size == 0
return input.chunk(world_size, dim=dim)[rank].contiguous()
def _gather(input, dim, group):
# skip if world_size == 1
world_size = dist.get_world_size(group=group)
if world_size == 1:
return input
# gather sequence
output = all_gather(input, uniform_size=True, group=group)
return torch.cat(output, dim=dim).contiguous()
class AllToAll(Function):
@staticmethod
def forward(ctx, input, scatter_dim, gather_dim, group):
ctx.scatter_dim = scatter_dim
ctx.gather_dim = gather_dim
ctx.group = group
return all_to_all(input, scatter_dim, gather_dim, group)
@staticmethod
def backward(ctx, grad_output):
return (all_to_all(grad_output, ctx.gather_dim, ctx.scatter_dim,
ctx.group), None, None, None)
class GradScaler(Function):
@staticmethod
def forward(ctx, input, scale):
ctx.scale = scale
return input
@staticmethod
def backward(ctx, grad_output):
if ctx.scale != 1:
grad_output = grad_output * ctx.scale
return grad_output, None
class AllGather(Function):
@staticmethod
def forward(ctx, input, dim, group=None):
ctx.dim = dim
ctx.group = group
output = all_gather(input, uniform_size=True, group=group)
return torch.cat(output, dim=dim).contiguous()
@staticmethod
def backward(ctx, grad_output):
rank = dist.get_rank(group=ctx.group)
world_size = dist.get_world_size(group=ctx.group)
return grad_output.chunk(world_size,
dim=ctx.dim)[rank].contiguous(), None, None
def diff_all_to_all(input, scatter_dim, gather_dim, group=None):
return AllToAll.apply(input, scatter_dim, gather_dim, group)
def diff_scatter_sequence(input, dim, group=None):
rank = dist.get_rank(group)
world_size = dist.get_world_size(group)
output = input.chunk(world_size, dim=dim)[rank].contiguous()
return GradScaler.apply(output, 1. / world_size)
def diff_gather_sequence(input, dim, group=None):
world_size = dist.get_world_size(group)
output = AllGather.apply(input, dim, group)
return GradScaler.apply(output, world_size)
class SplitForwardGatherBackward(Function):
@staticmethod
def forward(ctx, input, dim, group=None, grad_scale=None):
ctx.dim = dim
ctx.group = group
ctx.grad_scale = grad_scale
return _split(input, dim, group)
@staticmethod
def backward(ctx, grad_output):
if ctx.grad_scale == 'up':
grad_output = grad_output * dist.get_world_size(group=ctx.group)
elif ctx.grad_scale == 'down':
grad_output = grad_output / dist.get_world_size(group=ctx.group)
return _gather(grad_output, ctx.dim, ctx.group), None, None, None
class GatherForwardSplitBackward(Function):
@staticmethod
def forward(ctx, input, dim, group=None, grad_scale=None):
ctx.dim = dim
ctx.group = group
ctx.grad_scale = grad_scale
return _gather(input, dim, group)
@staticmethod
def backward(ctx, grad_output):
if ctx.grad_scale == 'up':
grad_output = grad_output * dist.get_world_size(group=ctx.group)
elif ctx.grad_scale == 'down':
grad_output = grad_output / dist.get_world_size(group=ctx.group)
return _split(grad_output, ctx.dim, ctx.group), None, None, None
def split_forward_gather_backward(input, dim, group=None, grad_scale=None):
return SplitForwardGatherBackward.apply(input, dim, group, grad_scale)
def gather_forward_split_backward(input, dim, group=None, grad_scale=None):
return GatherForwardSplitBackward.apply(input, dim, group, grad_scale)
global we
@@ -295,7 +547,10 @@ 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)
dist.init_process_group(backend='nccl',
world_size=world_size,
rank=rank,
timeout=timedelta(seconds=18000))
torch.backends.cudnn.deterministic = cfg.ENV.get('CUDNN_DETERMINISTIC',
True)
torch.backends.cudnn.benchmark = cfg.ENV.get('CUDNN_BENCHMARK', False)
@@ -307,11 +562,50 @@ def mp_worker(gpu, ngpus_per_node, cfg, fn, pmi_rank, world_size, work_env):
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 rank {work_env.rank} current devices num {ngpus_per_node} \n'
f'current machine rank {pmi_rank} and all world size {world_size}')
# model parallel
tensor_parallel_size = cfg.ENV.get('TENSOR_PARALLEL_SIZE', 1)
pipeline_parallel_size = cfg.ENV.get('PIPELINE_PARALLEL_SIZE', 1)
we.set_env(work_env.get_env())
env_dict = work_env.get_env()
if tensor_parallel_size * pipeline_parallel_size > 1:
'''
'''
assert world_size % tensor_parallel_size == 0
assert world_size % (tensor_parallel_size *
pipeline_parallel_size) == 0
data_parallel_size = world_size // (tensor_parallel_size *
pipeline_parallel_size)
mesh = torch.arange(world_size).view(data_parallel_size,
pipeline_parallel_size,
tensor_parallel_size)
index = torch.where(mesh == rank)
assert all(u.numel() == 1 for u in index)
index = [u.item() for u in index]
for j in range(pipeline_parallel_size):
for k in range(tensor_parallel_size):
group = dist.new_group(mesh[:, j, k].tolist())
if j == index[1] and k == index[2]:
env_dict['data_parallel_group'] = group
for i in range(data_parallel_size):
for j in range(pipeline_parallel_size):
group = dist.new_group(mesh[i, j, :].tolist())
if i == index[0] and j == index[1]:
env_dict['tensor_parallel_group'] = group
for i in range(data_parallel_size):
for k in range(tensor_parallel_size):
ranks = mesh[i, :, k].tolist()
group = dist.new_group(ranks)
if i == index[0] and k == index[2]:
env_dict['pipeline_parallel_group'] = group
env_dict['pipeline_parallel_ranks'] = ranks
we.set_env(env_dict)
work_env.logger.info(str(we))
fn(cfg)
torch.cuda.synchronize()
barrier()
class Workenv(object):
@@ -322,6 +616,7 @@ class Workenv(object):
self.rank = 0
self.world_size = 1
self.device_id = 0
self.backend = ''
self.device_count = 1
self.seed = 2023
self.debug = False
@@ -330,24 +625,36 @@ class Workenv(object):
self.data_online = False
self.share_storage = False
self.data_parallel_group = None
self.tensor_parallel_group = None
self.pipleline_parallel_group = None
self.pipeline_parallel_ranks = None
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'
if logger is None:
self.logger = StdMsg(name='env')
else:
self.logger = logger
self.sys_envs = config.ENV.get('SYS_ENVS', None)
if self.sys_envs:
for k, v in self.sys_envs.items():
os.environ[k] = v
self.logger.info(f'Set env variable {k}={v}')
set_random_seed(self.seed)
if logger is not None:
logger.info(f'And running with seed {self.seed}!')
self.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'
@@ -379,7 +686,8 @@ class Workenv(object):
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.init_process_group(backend=self.backend,
timeout=timedelta(seconds=18000))
# dist.barrier()
self.initialized = True
if dist.is_initialized():
@@ -422,7 +730,7 @@ class Workenv(object):
self.device_count = ngpus_per_node
world_size = ngpus_per_node * pmi_world_size
self.world_size = world_size
if self.world_size > 1:
if self.world_size >= 1:
self.is_distributed = True
self.initialized = True
if self.is_distributed:
@@ -434,13 +742,41 @@ class Workenv(object):
self))
def get_env(self):
return self.__dict__
ret_dict = {}
for k, v in self.__dict__.items():
if isinstance(v, (list, dict, int, float, str, bool)):
ret_dict[k] = v
return ret_dict
def set_env(self, we_env):
for k, v in we_env.items():
setattr(self, k, v)
set_random_seed(self.seed)
def group_info(self, group):
group_info = f'group size: {group.size()}\n'
group_info += f'group rank: {group.rank()}\n'
group_info += f'group name: {group.name()}\n'
return group_info
@property
def data_group_world_size(self):
if self.data_parallel_group is not None:
return self.data_parallel_group.size()
return self.world_size
@property
def tensor_group_world_size(self):
if self.tensor_parallel_group is not None:
return self.tensor_parallel_group.size()
return 1
@property
def pipeline_group_world_size(self):
if self.pipleline_parallel_group is not None:
return self.pipleline_parallel_group.size()
return 1
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'
@@ -448,8 +784,19 @@ class Workenv(object):
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} \n"
environ_str += f"Current task's share storage is set {self.share_storage} \n"
environ_str += f"Current task's global seed is set {self.seed}"
environ_str += f"Current task's global seed is set {self.seed} \n"
if self.data_parallel_group is not None:
environ_str += f"Current task's data parallel group: {self.group_info(self.data_parallel_group)} \n"
if self.pipleline_parallel_group is not None:
environ_str += f"Current task's pipeline parallel group: {self.group_info(self.pipleline_parallel_group)} \n"
if self.tensor_parallel_group is not None:
environ_str += f"Current task's tensor parallel group: {self.group_info(self.tensor_parallel_group)} \n"
environ_str += f"Current task's backend is set {self.backend} \n"
return environ_str
def __del__(self):
if we.is_distributed:
dist.destroy_process_group()
we = Workenv()
@@ -298,8 +298,7 @@ class AliyunOssFs(BaseFs):
try:
_ = self._download_object_multi_part(target_path,
temp_file,
chunk_size=50 *
1024 * 1024)
chunk_size=size // 100)
break
except Exception as e:
retry += 1
@@ -361,7 +360,7 @@ class AliyunOssFs(BaseFs):
def download_one_part(key):
while not slice_queue.empty():
R.acquire()
R.acquire(timeout=60)
try:
if not slice_queue.empty():
part_number, chunk = slice_queue.get_nowait()
@@ -383,7 +382,7 @@ class AliyunOssFs(BaseFs):
target_path, chunk[0], chunk[1])
with open(temp_part_file, 'wb') as f:
f.write(data)
R.acquire()
R.acquire(timeout=60)
try:
ret_slice_queue.put_nowait({
'part_number':
@@ -403,7 +402,7 @@ class AliyunOssFs(BaseFs):
'Download part {} for {} error {} retry {} times!'.
format(part_number, key, e, retry))
if retry >= self._retry_times:
R.acquire()
R.acquire(timeout=60)
try:
ret_slice_queue.put_nowait({
'part_number': part_number,
@@ -539,7 +538,7 @@ class AliyunOssFs(BaseFs):
meta_dict[target_path] = etag
else:
local_path = None
R.acquire()
R.acquire(timeout=60)
try:
data_quene.put_nowait([target_path, local_path])
except Exception:
@@ -623,7 +622,11 @@ class AliyunOssFs(BaseFs):
meta_dict = self._get_dir(target_path,
local_path=local_path,
meta_dict=copy.deepcopy(meta_dict))
json.dump(meta_dict, open(check_file, 'w'))
for _ in range(5):
try:
json.dump(meta_dict, open(check_file, 'w'))
except:
time.sleep(1)
if is_tmp:
self.add_temp_file(local_path)
return local_path
@@ -698,7 +701,7 @@ class AliyunOssFs(BaseFs):
def upload_one_part(key, upload_id):
while not slice_queue.empty():
R.acquire()
R.acquire(timeout=60)
try:
if not slice_queue.empty():
part_number, offset, num_to_upload = slice_queue.get_nowait(
@@ -721,7 +724,7 @@ class AliyunOssFs(BaseFs):
try:
result = _bucket.upload_part(key, upload_id,
part_number, raw_data)
R.acquire()
R.acquire(timeout=60)
try:
ret_slice_queue.put_nowait({
'part_number': part_number,
@@ -738,7 +741,7 @@ class AliyunOssFs(BaseFs):
'Upload part {} for {} error {} retry {} times!'.
format(part_number, key, e, retry))
if retry >= self._retry_times:
R.acquire()
R.acquire(timeout=60)
try:
ret_slice_queue.put_nowait({
'part_number': part_number,
@@ -965,8 +968,8 @@ class AliyunOssFs(BaseFs):
key,
lifecycle,
slash_safe=slash_safe)
_bucket.put_object_acl(key, oss2.OBJECT_ACL_PUBLIC_READ)
if set_public:
_bucket.put_object_acl(key, oss2.OBJECT_ACL_PUBLIC_READ)
output_url = output_url.replace('%2F', '/').split('?')[0]
return output_url
except Exception as e:
@@ -1060,7 +1063,7 @@ class AliyunOssFs(BaseFs):
local_path, target_path)
else:
flg = False
R.acquire()
R.acquire(timeout=60)
try:
data_quene.put_nowait([local_path, target_path, flg])
except Exception:
@@ -39,6 +39,7 @@ class BaseFs(object, metaclass=ABCMeta):
self._temp_files = set()
self.cfg = cfg
self.tmp_dir = cfg.get('TEMP_DIR', None)
self.enable_md5_path = cfg.get('ENABLE_MD5_PATH', True)
self.auto_clean = cfg.get('AUTO_CLEAN', False)
if self.tmp_dir is None:
self.auto_clean = True
@@ -306,7 +307,10 @@ class BaseFs(object, metaclass=ABCMeta):
rand_name += f'{suffix}'
tmp_file = osp.join(tempfile.gettempdir(), rand_name)
else:
cache_name = '{}{}{}'.format(etag, get_md5(target_path), suffix)
if self.enable_md5_path:
cache_name = '{}{}{}'.format(etag, get_md5(target_path), suffix)
else:
cache_name = ''
tmp_file = osp.join(self.tmp_dir, cache_name)
return tmp_file
@@ -84,6 +84,7 @@ class HuggingfaceFs(BaseFs):
target_path,
local_path=None,
wait_finish=False,
multi_thread=False,
timeout=3600,
sign_key=None,
worker_id=-1) -> Optional[str]:
+19 -6
View File
@@ -41,9 +41,7 @@ class LocalFs(BaseFs):
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('/../', '/')
return os.path.abspath(os.path.join(self.get_prefix(), target_path))
if target_path.startswith('/'):
return target_path
if target_path.startswith('file://'):
@@ -236,8 +234,17 @@ class LocalFs(BaseFs):
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)
dirname = os.path.dirname(target_path)
if not os.path.exists(dirname):
os.makedirs(dirname, exist_ok=True)
if isinstance(local_data, str):
with open(target_path, 'w') as f:
f.write(local_data)
elif isinstance(local_data, bytes):
with open(target_path, 'wb') as f:
f.write(local_data)
else:
raise NotImplementedError
return True
def walk_dir(self, file_dir, recurse=True):
@@ -249,6 +256,9 @@ class LocalFs(BaseFs):
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)
dirname = os.path.dirname(target_path)
if not os.path.exists(dirname):
os.makedirs(dirname, exist_ok=True)
if local_path != target_path:
try:
shutil.copy(local_path, target_path)
@@ -271,7 +281,7 @@ class LocalFs(BaseFs):
return False
return True
try:
os.makedirs(target_dir)
os.makedirs(target_dir, exist_ok=True)
except Exception as e:
self.logger.error(e)
return False
@@ -309,6 +319,9 @@ class LocalFs(BaseFs):
multi_thread=False) -> bool:
local_dir = self.reconstruct_path(local_dir)
target_dir = self.reconstruct_path(target_dir)
dirname = os.path.dirname(target_dir)
if not os.path.exists(dirname):
os.makedirs(dirname, exist_ok=True)
if local_dir == target_dir:
return True
# # cp -f local_dir/* target_dir/*
@@ -24,8 +24,8 @@ class ModelscopeFs(BaseFs):
super(ModelscopeFs, self).__init__(cfg, logger=logger)
retry_times = cfg.get('RETRY_TIMES', 10)
self._retry_times = retry_times
self._model_id_loaded = set()
self._model_file_loaded = set()
self._model_id_loaded = dict()
self._model_file_loaded = dict()
def get_prefix(self) -> str:
return 'ms://'
@@ -69,14 +69,14 @@ class ModelscopeFs(BaseFs):
if revision is not None:
local_path = local_path + '_' + str(revision)
model_file = os.path.join(key, file_path)
retry = 0
while retry < self._retry_times:
try:
model_file = os.path.join(key, file_path)
if model_file in self._model_file_loaded:
local_path = os.path.join(local_path, model_file)
local_path = self._model_file_loaded[model_file]
if not osp.exists(local_path):
self._model_file_loaded.remove(key)
self._model_file_loaded.pop(model_file)
else:
if token is not None:
cookies = self.get_modelscope_cookie(token)
@@ -94,7 +94,7 @@ class ModelscopeFs(BaseFs):
if retry >= self._retry_times:
return None
self._model_file_loaded.add(model_file)
self._model_file_loaded[model_file] = local_path
if is_tmp:
self.add_temp_file(local_path)
return local_path
@@ -142,9 +142,9 @@ class ModelscopeFs(BaseFs):
while retry < self._retry_times:
try:
if key in self._model_id_loaded:
local_path = os.path.join(local_path, key)
local_path = self._model_id_loaded[key]
if not osp.exists(local_path):
self._model_id_loaded.remove(key)
self._model_id_loaded.pop(key)
else:
if token is not None:
cookies = self.get_modelscope_cookie(token)
@@ -162,7 +162,7 @@ class ModelscopeFs(BaseFs):
if retry >= self._retry_times:
return None
self._model_id_loaded.add(key)
self._model_id_loaded[key] = local_path
if is_tmp:
self.add_temp_file(local_path)
if not ret_folder == '':
+7 -4
View File
@@ -281,7 +281,7 @@ class FileSystem(object):
else:
return False
def get_batch_objects_from(self, target_path_list, wait_finish=False):
def get_batch_objects_from(self, target_path_list, wait_finish=False, return_target_path=False):
data_quene = Queue()
batch_size = 20
R = threading.Lock()
@@ -293,7 +293,7 @@ class FileSystem(object):
wait_finish=wait_finish)
else:
local_path = None
R.acquire()
R.acquire(timeout = 2)
try:
data_quene.put_nowait([target_path, local_path])
except Exception:
@@ -320,7 +320,10 @@ class FileSystem(object):
for target_path in batch_list:
local_path = file_dict.get(target_path, None)
yield local_path
if return_target_path:
yield target_path, local_path
else:
yield local_path
def put_batch_objects_to(self,
local_path_list,
@@ -352,7 +355,7 @@ class FileSystem(object):
pass
else:
flg = False
R.acquire()
R.acquire(timeout=2)
try:
data_quene.put_nowait([local_path, target_path, flg])
except Exception:
+51 -1
View File
@@ -6,7 +6,6 @@ from tqdm import tqdm
from scepter.modules.utils.file_system import FS
def init_1level_llfs(list_file,
max_lines=1024,
index_name='index',
@@ -14,6 +13,57 @@ def init_1level_llfs(list_file,
r"""Construct large-list-file index.
"""
index_dir = osp.splitext(list_file)[0]
total_failed = 0
with FS.get_from(list_file, wait_finish=True) as local_path:
_num_split, index_stack, save_files = 0, [], []
cache_data_list, target_path_list = [], []
with open(local_path, 'r', buffering=1000000) as f:
for line in tqdm(f):
index_stack.append(line.strip())
if len(index_stack) >= max_lines:
save_file = f'{index_dir}/{index_name}/{_num_split + 1:09d}.txt'
cache_data_list.append('\n'.join(index_stack).encode())
target_path_list.append(save_file)
# with FS.put_to(save_file) as cache_path:
# with open(cache_path, 'w') as f_w:
# f_w.write('\n'.join(index_stack))
if len(target_path_list) >= 1000:
put_res = [(target_path, flg) for local_path, target_path, flg
in FS.put_batch_objects_to(cache_data_list, target_path_list, batch_size=40)]
for put_flg in put_res:
if not put_flg:
total_failed += 1
cache_data_list, target_path_list = [], []
index_stack = []
_num_split += 1
save_files.append(save_file)
put_res = [(target_path, flg) for local_path, target_path, flg
in FS.put_batch_objects_to(cache_data_list, target_path_list, batch_size=50)]
for put_flg in put_res:
if not put_flg:
total_failed += 1
print(f'Failed to put {total_failed} files.')
if len(index_stack) > 0:
save_file = f'{index_dir}/{index_name}/{_num_split + 1:06d}.txt'
with FS.put_to(save_file) as cache_path:
with open(cache_path, 'w') as f_w:
f_w.write('\n'.join(index_stack))
save_files.append(save_file)
# output meta-file
index_file = osp.join(index_dir, f'{index_name}.txt')
with FS.put_to(index_file) as cache_path:
with open(cache_path, 'w') as f_w:
f_w.write('\n'.join(save_files))
return index_file
def init_1level_llfs_single_threading(list_file,
max_lines=1024,
index_name='index',
delimiter='\n'):
r"""Construct large-list-file index.
"""
index_dir = osp.splitext(list_file)[0]
print(list_file)
with FS.get_from(list_file, wait_finish=True) as local_path:
print(local_path)
+3 -3
View File
@@ -47,7 +47,7 @@ def time_since(since, percent):
return '{} {:.2f}%({})'.format(as_time(s), 100 * percent, as_time(rs))
def get_logger(name='torch dist'):
def get_logger(name='scepter', level=logging.INFO):
logger = logging.getLogger(name)
logger.propagate = False
if len(logger.handlers) == 0:
@@ -57,8 +57,8 @@ def get_logger(name='torch dist'):
'[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)
std_handler.setLevel(level)
logger.setLevel(level)
logger.addHandler(std_handler)
return logger
+328 -123
View File
@@ -1,7 +1,9 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import json
import os.path
from io import BytesIO
from numbers import Number
import numpy as np
@@ -76,7 +78,8 @@ def merge_gathered_probe(all_gathered_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)
view_distribute=ret_data.view_distribute,
is_presave=ret_data.is_presave)
elif isinstance(ret_data.data, list):
if ret_data.build_label is not None:
if isinstance(ret_data.build_label, str):
@@ -95,19 +98,57 @@ def merge_gathered_probe(all_gathered_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)
view_distribute=ret_data.view_distribute,
is_presave=ret_data.is_presave)
else:
all_gathered_data[key] = gathered_data
return all_gathered_data
class MediaHandler():
def __init__(self, batch_size = 10):
self.file_list = []
self.target_path_list = []
self.target_status = {}
self.batch_size = batch_size
def append(self, source_file, target_path):
self.file_list.append(source_file)
self.target_path_list.append(target_path)
if len(self.file_list) > 2 * self.batch_size:
generator = FS.put_batch_objects_to(self.file_list, self.target_path_list, batch_size=self.batch_size)
for local_path, target_path, flg in generator:
self.target_status[target_path] = flg
self.file_list.clear()
self.target_path_list.clear()
def sync(self):
if len(self.file_list) > 0:
if len(self.file_list) > 4 * self.batch_size:
generator = FS.put_batch_objects_to(self.file_list, self.target_path_list, batch_size=self.batch_size)
for local_path, target_path, flg in generator:
self.target_status[target_path] = flg
else:
for file_, target_path in zip(self.file_list, self.target_path_list):
self.target_status[target_path] = FS.put_object(file_.getvalue(), target_path)
self.file_list.clear()
self.target_path_list.clear()
def clear(self):
self.file_list.clear()
self.target_path_list.clear()
self.target_status.clear()
class ProbeData():
def __init__(self,
data,
is_image=False,
is_video=False,
fps=8,
build_html=False,
build_label=None,
view_distribute=False):
view_distribute=False,
is_presave = False):
''' Probe Data Initialize.
We only support basic types such as [torch.Tensor, numpy.ndarray, number, str],
or [dict, list] of [dict, list,
@@ -164,6 +205,28 @@ class ProbeData():
elif isinstance(v, np.ndarray):
data[k] = v
self.basic_type = False
elif isinstance(v, list):
for v_idx, v_v in enumerate(v):
if not check_legal_type(v_v):
if isinstance(v_v, torch.Tensor):
data[k][v_idx] = v_v.detach().cpu().numpy()
self.basic_type = False
elif isinstance(v_v, np.ndarray):
data[k][v_idx] = v_v
self.basic_type = False
else:
raise f'Unsupport data type for {v_v}'
elif isinstance(v, dict):
for k_k, v_v in v.items():
if not check_legal_type(v_v):
if isinstance(v_v, torch.Tensor):
data[k][k_k] = v_v.detach().cpu().numpy()
self.basic_type = False
elif isinstance(v_v, np.ndarray):
data[k][k_k] = v_v
self.basic_type = False
else:
raise f'Unsupport data type for {v_v}'
else:
raise f'Unsupport data type for {v}'
self.data = data
@@ -176,6 +239,28 @@ class ProbeData():
elif isinstance(v, np.ndarray):
data[idx] = v
self.basic_type = False
elif isinstance(v, list):
for v_idx, v_v in enumerate(v):
if not check_legal_type(v_v):
if isinstance(v_v, torch.Tensor):
data[idx][v_idx] = v_v.detach().cpu().numpy()
self.basic_type = False
elif isinstance(v_v, np.ndarray):
data[idx][v_idx] = v_v
self.basic_type = False
else:
raise f'Unsupport data type for {v_v}'
elif isinstance(v, dict):
for k, v_v in v.items():
if not check_legal_type(v_v):
if isinstance(v_v, torch.Tensor):
data[idx][k] = v_v.detach().cpu().numpy()
self.basic_type = False
elif isinstance(v_v, np.ndarray):
data[idx][k] = v_v
self.basic_type = False
else:
raise f'Unsupport data type for {v_v}'
else:
raise f'Unsupport data type for {v}'
self.data = data
@@ -185,7 +270,13 @@ class ProbeData():
raise f'Unsupport data type for {data}'
self.is_image = is_image
self.is_video = is_video
self.image_postfix = 'jpg'
self.video_postfix = 'mp4'
self.fps = fps
self.build_html = build_html
self.media_handler = MediaHandler()
self.is_presave = is_presave
if self.build_html:
assert build_label is not None
@@ -199,7 +290,84 @@ class ProbeData():
build_label, dict)
self.build_label = build_label
def save_image(self, file_prefix, images, image_postfix):
def get_format(self, extension):
if extension.lower() in ['jpg', 'jpeg']:
return 'JPEG'
if extension.lower() in ['png']:
return 'PNG'
return 'JPEG'
def save_one_video(self, file_path, videos, fps = 8):
# write video
import imageio
try:
writer = imageio.get_writer(file_path, fps=fps, format=".mp4", codec='libx264', quality=8)
for frame in videos:
writer.append_data(frame)
writer.close()
return True
except:
return False
def save_video(self, file_prefix, videos, video_postfix, fps = 8, rank = 0):
if isinstance(videos, list):
for video in videos:
if isinstance(video, list):
raise f"Only surpport one layer nested list."
return [self.save_video(file_prefix + f'_{rank}_{idx}', v, video_postfix, fps) for idx, v in enumerate(videos)]
np_shape = videos.shape
# 4D
shape_str = '_'.join([str(v) for v in np_shape])
if len(np_shape) == 5:
# channel is 1 or 3
if np_shape[-1] == 1 or np_shape[-1] == 3 or np_shape[-1] == 4:
if np_shape[-1] == 1:
videos = videos.reshape(videos.shape[:-1])
file_list = []
for idx in range(np_shape[0]):
if videos[idx].shape[0] > 1:
file_path = os.path.join(file_prefix, f'probe_{rank}_{idx}_[{shape_str}].{video_postfix}')
byio = BytesIO()
is_suc = self.save_one_video(byio, videos[idx], fps)
if not is_suc:
byio.write(b"")
else:
file_path = os.path.join(file_prefix, f'probe_{rank}_{idx}_[{shape_str}].{self.image_postfix}')
byio = BytesIO()
Image.fromarray(videos[idx][0]).save(byio, self.get_format(self.image_postfix))
self.media_handler.append(byio, file_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) == 4:
if np_shape[-1] == 1 or np_shape[-1] == 3 or np_shape[-1] == 4:
if np_shape[-1] == 1:
videos = videos.reshape(videos.shape[:-1])
if videos.shape[0] > 1:
file_path = file_prefix + f'_probe_{rank}_[{shape_str}].{video_postfix}'
byio = BytesIO()
is_suc = self.save_one_video(byio, videos, fps)
if not is_suc:
byio.write(b"")
else:
file_path = file_prefix + f'_probe_{rank}_[{shape_str}].{self.image_postfix}'
byio = BytesIO()
Image.fromarray(videos[0]).save(byio, self.get_format(self.image_postfix))
self.media_handler.append(byio, file_path)
return file_path
else:
videos = videos.reshape(list(videos.shape) + [1])
return self.save_video(file_prefix, videos, video_postfix, fps = fps)
else:
raise f"Ensure your data's dim is BFWHC or FWHC, and channel is 1 or 3 for {file_prefix}"
def save_image(self, file_prefix, images, image_postfix, rank = 0):
if isinstance(images, list):
for image in images:
if isinstance(image, list):
raise f"Only surpport one layer nested list."
return [self.save_image(file_prefix + f'_{rank}_{idx}', v, image_postfix) for idx, v in enumerate(images)]
np_shape = images.shape
# 4D
shape_str = '_'.join([str(v) for v in np_shape])
@@ -210,9 +378,10 @@ class ProbeData():
images = images.reshape(images.shape[:-1])
file_list = []
for idx in range(np_shape[0]):
file_path = file_prefix + f'_probe_{idx}_[{shape_str}].{image_postfix}'
with FS.put_to(file_path) as local_path:
Image.fromarray(images[idx, ...]).save(local_path)
file_path = os.path.join(file_prefix, f'probe_{rank}_{idx}_[{shape_str}].{image_postfix}')
byio = BytesIO()
Image.fromarray(images[idx, ...]).save(byio, self.get_format(self.image_postfix))
self.media_handler.append(byio, file_path)
file_list.append(file_path)
return file_list
else:
@@ -221,26 +390,29 @@ class ProbeData():
if np_shape[-1] == 1 or np_shape[-1] == 3 or np_shape[-1] == 4:
if np_shape[-1] == 1:
images = images.reshape(images.shape[:-1])
file_path = file_prefix + f'_probe_[{shape_str}].{image_postfix}'
with FS.put_to(file_path) as local_path:
Image.fromarray(images).save(local_path)
file_path = file_prefix + f'_probe_{rank}_[{shape_str}].{image_postfix}'
byio = BytesIO()
Image.fromarray(images).save(byio, self.get_format(self.image_postfix))
self.media_handler.append(byio, file_path)
return file_path
else:
images = images.reshape(list(images.shape) + [1])
return self.save_image(file_prefix, images, image_postfix)
elif len(np_shape) == 2:
file_path = file_prefix + f'_probe_[{shape_str}].{image_postfix}'
with FS.put_to(file_path) as local_path:
Image.fromarray(images).save(local_path)
file_path = file_prefix + f'_probe_{rank}_[{shape_str}].{image_postfix}'
byio = BytesIO()
Image.fromarray(images).save(byio, self.get_format(self.image_postfix))
self.media_handler.append(byio, file_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):
def save_npy(self, file_prefix, data, rank = 0):
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)
file_path = file_prefix + f'_{rank}_{shape_str}.npy'
byio = BytesIO()
np.save(byio, data)
self.media_handler.append(byio, file_path)
return file_path
def save_html(self, html_prefix, ret_data, ret_label):
@@ -249,9 +421,10 @@ class ProbeData():
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')
'opacity:1.0;} textarea {font-size: 32px;}</style>\n')
f.writelines('<br><hr/>\n')
all_ranks = list()
is_textarea = False
for save_id, save_data in enumerate(zip(ret_data, ret_label)):
save_path, save_label = save_data
one_rank = '<table><tr>'
@@ -259,15 +432,32 @@ class ProbeData():
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.').replace(
'-internal', '')
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>'
)
try:
url = FS.get_url(one_path,
lifecycle=3600 * 365 * 24).replace(
'.oss-internal.aliyun-inc.',
'.oss.aliyuncs.').replace(
'-internal', '')
except:
url = one_path
if len(one_label) > 10 and idx == len(save_path) - 1:
is_textarea = True
if self.is_video and one_path.endswith(self.video_postfix):
one_rank += f'<td align="center"><video height="{height}" controls="">'
one_rank += f'<source src="{url}" type="video/mp4"></video>'
if idx == len(save_path) - 1 and is_textarea:
one_rank += f'<br><font size="4"><strong>{save_id}-{idx}<strong></font><br/></td>'
one_rank += f'<td align="center"><textarea rows="16" cols="40">{one_label}</textarea><br><font size="4"><strong>{save_id}-{idx}<strong></font><br/></td>'
else:
one_rank += f'<br><font size="4"><strong>{save_id}-{idx}|{one_label}<strong></font><br/></td>'
else:
one_rank += f'<td align="center"><input type="image" src="{url}" >'
if idx == len(save_path) - 1 and is_textarea:
one_rank += f'<br><font size="4"><strong>{save_id}-{idx}<strong></font><br/></td>'
one_rank += f'<td align="center"><textarea rows="16" cols="40">{one_label}</textarea><br><font size="4"><strong>{save_id}-{idx}<strong></font><br/></td>'
else:
one_rank += 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))
@@ -277,117 +467,132 @@ class ProbeData():
def distribute(self):
return self._distribute_dict
def to_log(self, prefix=None, image_postfix='jpg'):
def save_one_media(self, idx, v, prefix_path, image_postfix, video_postfix, rank = 0):
ret_label = None
if self.is_image:
ret_medias = self.save_image(prefix_path, v,
image_postfix, rank = rank)
elif self.is_video:
ret_medias = self.save_video(prefix_path, v,
video_postfix, fps=self.fps, rank = rank)
else:
ret_data = self.save_npy(prefix_path, v, rank = rank)
return ret_data, ret_label
ret_data = ret_medias if isinstance(ret_medias, list) else [ret_medias]
if self.build_html:
if isinstance(ret_medias, list):
if isinstance(self.build_label, str):
ret_label = [self.build_label for _ in ret_medias]
elif isinstance(self.build_label[idx], list):
assert len(self.build_label[idx]) == len(
ret_medias)
ret_label = self.build_label[idx]
else:
ret_label = [
self.build_label[idx]
for _ in ret_medias
]
else:
if isinstance(self.build_label, str):
ret_label = [self.build_label]
else:
ret_label = [self.build_label[idx]]
return ret_data, ret_label
def presave(self, prefix=None, image_postfix='jpg', video_postfix='mp4', rank = 0):
self.image_postfix = image_postfix
self.video_postfix = video_postfix
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, image_postfix)
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
ret_data = self.save_image(prefix, self.data, image_postfix, rank = rank)
elif self.is_video:
ret_data = self.save_video(prefix, self.data, video_postfix, fps=self.fps, rank = rank)
else:
ret_data = self.save_npy(prefix, self.data)
return ret_data
ret_data = self.save_npy(prefix, self.data, rank = rank)
self.media_handler.sync()
self.media_handler.clear()
if isinstance(ret_data, list):
ret_data = [ret_data]
ret_label = []
if self.build_html:
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]}."
self.data = {"ret_data": ret_data, "ret_label": ret_label}
else:
self.data = ret_data
self.is_presave = True
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,
image_postfix)
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
ret_one_data, ret_one_label = self.save_one_media(idx, v,
prefix_path,
image_postfix,
video_postfix,
rank=rank)
ret_data.append(ret_one_data)
ret_label.append(ret_one_label) if ret_one_label is not None else ret_label
self.media_handler.sync()
self.media_handler.clear()
self.data = {"ret_data": ret_data, "ret_label": ret_label}
self.is_presave = True
elif isinstance(self.data, dict):
if not self.basic_type:
ret_data = []
ret_label = []
for k, v in self.data:
for k, v in self.data.items():
prefix_path = os.path.join(prefix, f'{k}_')
if self.is_image:
ret_images = self.save_image(prefix_path, v,
image_postfix)
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:
ret_one_data, ret_one_label = self.save_one_media(k, v,
prefix_path,
image_postfix,
video_postfix,
rank = rank)
ret_data.append(ret_one_data)
ret_label.append(ret_one_label) if ret_one_label is not None else ret_label
self.media_handler.sync()
self.media_handler.clear()
self.data = {"ret_data": ret_data, "ret_label": ret_label}
self.is_presave = True
def to_log(self, prefix=None, image_postfix='jpg', video_postfix='mp4', rank = 0):
if not self.is_presave:
self.presave(prefix, image_postfix, video_postfix, rank = rank)
if not self.is_presave:
return self.data
if isinstance(self.data, str):
return self.data
elif isinstance(self.data, dict):
ret_data, ret_label = self.data["ret_data"], self.data["ret_label"]
if 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}
elif isinstance(self.data, list):
ret_data, ret_label = [], []
for one_data in self.data:
if isinstance(one_data, dict):
one_ret_data, one_ret_label = one_data["ret_data"], one_data["ret_label"]
ret_data.extend(one_ret_data)
ret_label.extend(one_ret_label)
elif isinstance(one_data, str):
ret_data.append(one_data)
if (self.is_image or self.is_video) 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}
+257
View File
@@ -0,0 +1,257 @@
# -*- coding: utf-8 -*-
from enum import Enum
class Media(Enum):
TEXT = 1
IMAGE = 2
VIDEO = 3
AUDIO = 4
class HtmlVisualization(object):
def __init__(
self,
allow_annotation=False,
slice_size=1000,
align='center',
width_scale='60%',
title='Visualization',
):
self.content_list = []
self.rows_meta = []
self.allow_annotation = allow_annotation
self.slice_size = slice_size
self.align = align
self.width_scale = width_scale
self.title = title
self.html_start = '<html>'
self.html_head = f'<head><meta charset="utf-8"><title>{title}</title></head>'
self.html_style = '''
<style>
table {
border-collapse: collapse;
}
td {
width: "{width_scale}";
align: "{align}";
margin: 0px;
border: 0px;
padding: 0px;
vertical-align: top;
}
video {
margin: 0px;
border: 0px solid #ccc;
padding: 0px;
}
textarea {
margin: 0px;
border: 0px;
padding: 0px;
resize: none;
border: 1px solid #ccc;
}
</style>
<script>
function adjustHeight() {
const textareas = document.querySelectorAll('textarea');
textareas.forEach(textarea => {
const td = textarea.parentNode;
const tdHeight = td.clientHeight;
textarea.style.height = tdHeight + 'px';
});
}
window.onload = adjustHeight;
window.onresize = adjustHeight;
</script>
'''.replace('{width_scale}',
self.width_scale).replace('{align}', self.align)
self.html_body = '<body>{BODY}</body>\n'
self.html_end = '</html>'
self.html_script = '''
<script>
function saveSamples() {
let selectedSamples = document.querySelectorAll('input[name="sample[]"]:checked');
let notSelectedSamples = document.querySelectorAll('input[name="sample[]"]:not(:checked)');
let sampleUrls = [];
for (let i=0; i<selectedSamples.length; i++) {
sampleUrls.push(selectedSamples[i].value + "#;#" + "1");
}
for (let i=0; i<notSelectedSamples.length; i++) {
sampleUrls.push(notSelectedSamples[i].value + "#;#" + "0");
}
let fileContent = sampleUrls.join('\\n');
let file = new Blob([fileContent], {type: 'text/plain'});
let a = document.createElement('a');
a.href = URL.createObjectURL(file);
a.download = 'result.txt';
a.click();
}
</script>
'''
self.label_button = (
'<table><tr><td>' +
"<button style='height: 50px;' type=\"button\" onclick=\"saveSamples()\">Save Samples</button>"
+ '</td></tr></table>')
def format_col(self,
content='',
label='',
type=Media.TEXT,
content_height=400,
content_width=600):
if type == Media.TEXT:
ret_str = '<td><textarea' # noqa: E501
# if content_height is not None:
# rows = f"rows={content_height//30}"
# ret_str += f" {rows}"
if content_width is not None:
cols = f"cols={content_width//15}"
ret_str += f" {cols}"
ret_str += f'>"{content}"</textarea></td>\n'
sec_ret_str = f'<td align="center"><font size="3"><strong>{label}<strong></font></td>\n'
return [ret_str, sec_ret_str]
elif type == Media.IMAGE:
ret_str = f'<td><img src="{content}"'
if content_height is not None:
height = f'height="{content_height}"'
ret_str += f" {height}"
if content_width is not None:
width = f'width="{content_width}"'
ret_str += f" {width}"
ret_str += ' ></td>\n'
sec_ret_str = f'<td align="center"><font size="3"><strong>{label}<strong></font></td>\n'
return [ret_str, sec_ret_str]
elif type == Media.VIDEO:
ret_str = '<td><video' # noqa
if content_height is not None:
height = f'height="{content_height}"'
ret_str += f" {height}"
if content_width is not None:
width = f'width="{content_width}"'
ret_str += f" {width}"
ret_str += ' controls>'
ret_str += f'<source src="{content}" type="video/mp4"></video></td>\n'
sec_ret_str = f'<td align="center"><font size="3"><strong>{label}<strong></font></td>\n'
return [ret_str, sec_ret_str]
elif type == Media.AUDIO:
ret_str = f'<td><audio src="{content}" controls></td>\n'
sec_ret_str = f'<td align="center"><font size="3"><strong>{label}<strong></font></td>\n'
return [ret_str, sec_ret_str]
else:
raise NotImplementedError
def format_row(self):
sample_id = 0
all_sample_html = []
for one_content, one_row_meta in zip(self.content_list,
self.rows_meta):
one_row_str = '<table><tr>'
one_row_str += '\n'.join([v[0] for v in one_content])
if self.allow_annotation:
row_meta = '#;#'.join(one_row_meta)
one_row_str += (
f'<td><input type="checkbox" class="large-checkbox" '
f'id="sample{sample_id}" name="sample[]" value="{row_meta}"></td>\n'
)
one_row_str += '</tr><tr>'
one_row_str += '\n'.join([v[1] for v in one_content]) # noqa
if self.allow_annotation: # noqa
one_row_str += f'<td></td>\n' # noqa
one_row_str += '</tr></table>'
if self.allow_annotation:
one_row_str = f'<label for="sample{sample_id}">{one_row_str}</label>'
all_sample_html.append(one_row_str)
sample_id += 1
return '\n'.join(all_sample_html)
def add_record(self,
content='',
label='',
type=Media.TEXT,
row_id=1,
col_id=1,
annotation_meta=None,
content_height=None,
content_width=None):
if row_id >= len(self.content_list):
self.content_list.append([])
self.rows_meta.append([])
if row_id != len(self.content_list) - 1:
raise RuntimeError(
'row_id should be next number of the last row_id.')
if col_id > len(self.content_list[row_id]):
raise RuntimeError(
'col_id should be next number of the last col_id.')
format_col = self.format_col(content, f"{row_id}-{col_id}: {label}",
type, content_height, content_width)
annotation_meta = annotation_meta if annotation_meta else ''
if col_id == len(self.content_list[row_id]):
self.content_list[row_id].append(format_col)
self.rows_meta[row_id].append(annotation_meta)
else:
self.content_list[row_id][col_id] = format_col
self.rows_meta[row_id][col_id] = annotation_meta
def save_html(self, path):
html_body = self.format_row()
ret_html_list = [
self.html_start, self.html_head, self.html_style,
self.html_body.replace('{BODY}', html_body)
]
if self.allow_annotation:
ret_html_list.append(self.label_button)
ret_html_list.append(self.html_script)
ret_html_list.append(self.html_end)
ret_html = '\n'.join(ret_html_list)
with open(path, 'w') as f:
f.write(ret_html)
if __name__ == '__main__':
from scepter.modules.utils.config import Config
from scepter.modules.utils.file_system import FS
FS.init_fs_client(Config(cfg_dict={}, load=False))
image_content_oss = '0_probe_0_[1024_2048_3].jpg'
content_oss = '6UTWGRG1lx08iRBx5REA01041200dzcb0E010.mp4'
caption = 'a little girl says hello.'
html_ins = HtmlVisualization(allow_annotation=True,
slice_size=1000,
title='Visualization',
width_scale='100%')
for i in range(4):
content_url = FS.get_url(content_oss, skip_check=True)
html_ins.add_record(content=content_url,
label='caption',
type=Media.VIDEO,
row_id=i,
col_id=0,
annotation_meta=None,
content_height=600,
content_width=None)
html_ins.add_record(content=caption,
label='caption',
type=Media.TEXT,
row_id=i,
col_id=1,
annotation_meta=None,
content_height=600,
content_width=750)
image_content_url = FS.get_url(image_content_oss, skip_check=True)
html_ins.add_record(content=image_content_url,
label='caption',
type=Media.IMAGE,
row_id=i,
col_id=2,
annotation_meta=None,
content_height=600,
content_width=None)
with FS.put_to('visualize.html') as local_path:
html_ins.save_html(local_path)