Team Ai
Apppublic

radames/Text2Human-API

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
util.py124 linesDownload Raw Back to utils
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