radames/Text2Human-API
1
1import logging2import os3import random4import sys5import time6from shutil import get_terminal_size7 8import numpy as np9import torch10 11logger = logging.getLogger('base')12 13 14def make_exp_dirs(opt):15 """Make dirs for experiments."""16 path_opt = opt['path'].copy()17 if opt['is_train']:18 overwrite = True if 'debug' in opt['name'] else False19 os.makedirs(path_opt.pop('experiments_root'), exist_ok=overwrite)20 os.makedirs(path_opt.pop('models'), exist_ok=overwrite)21 else:22 os.makedirs(path_opt.pop('results_root'))23 24 25def set_random_seed(seed):26 """Set random seeds."""27 random.seed(seed)28 np.random.seed(seed)29 torch.manual_seed(seed)30 torch.cuda.manual_seed(seed)31 torch.cuda.manual_seed_all(seed)32 33 34class ProgressBar(object):35 """A progress bar which can print the progress.36 37 Modified from:38 https://github.com/hellock/cvbase/blob/master/cvbase/progress.py39 """40 41 def __init__(self, task_num=0, bar_width=50, start=True):42 self.task_num = task_num43 max_bar_width = self._get_max_bar_width()44 self.bar_width = (45 bar_width if bar_width <= max_bar_width else max_bar_width)46 self.completed = 047 if start:48 self.start()49 50 def _get_max_bar_width(self):51 terminal_width, _ = get_terminal_size()52 max_bar_width = min(int(terminal_width * 0.6), terminal_width - 50)53 if max_bar_width < 10:54 print(f'terminal width is too small ({terminal_width}), '55 'please consider widen the terminal for better '56 'progressbar visualization')57 max_bar_width = 1058 return max_bar_width59 60 def start(self):61 if self.task_num > 0:62 sys.stdout.write(f"[{' ' * self.bar_width}] 0/{self.task_num}, "63 f'elapsed: 0s, ETA:\nStart...\n')64 else:65 sys.stdout.write('completed: 0, elapsed: 0s')66 sys.stdout.flush()67 self.start_time = time.time()68 69 def update(self, msg='In progress...'):70 self.completed += 171 elapsed = time.time() - self.start_time72 fps = self.completed / elapsed73 if self.task_num > 0:74 percentage = self.completed / float(self.task_num)75 eta = int(elapsed * (1 - percentage) / percentage + 0.5)76 mark_width = int(self.bar_width * percentage)77 bar_chars = '>' * mark_width + '-' * (self.bar_width - mark_width)78 sys.stdout.write('\033[2F') # cursor up 2 lines79 sys.stdout.write(80 '\033[J'81 ) # clean the output (remove extra chars since last display)82 sys.stdout.write(83 f'[{bar_chars}] {self.completed}/{self.task_num}, '84 f'{fps:.1f} task/s, elapsed: {int(elapsed + 0.5)}s, '85 f'ETA: {eta:5}s\n{msg}\n')86 else:87 sys.stdout.write(88 f'completed: {self.completed}, elapsed: {int(elapsed + 0.5)}s, '89 f'{fps:.1f} tasks/s')90 sys.stdout.flush()91 92 93class AverageMeter(object):94 """95 Computes and stores the average and current value96 Imported from97 https://github.com/pytorch/examples/blob/master/imagenet/main.py#L247-L26298 """99 100 def __init__(self):101 self.reset()102 103 def reset(self):104 self.val = 0105 self.avg = 0 # running average = running sum / running count106 self.sum = 0 # running sum107 self.count = 0 # running count108 109 def update(self, val, n=1):110 # n = batch_size111 112 # val = batch accuracy for an attribute113 # self.val = val114 115 # sum = 100 * accumulative correct predictions for this attribute116 self.sum += val * n117 118 # count = total samples so far119 self.count += n120 121 # avg = 100 * avg accuracy for this attribute122 # for all the batches so far123 self.avg = self.sum / self.count124 