update v1.1.0
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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 == '':
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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('<', '<').replace(
|
||||
'>', '>')
|
||||
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}
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user