CoolFace
Modelpublic

scrapegoat/Neural-Audio-Codec

sourceHugging Faceupdated 2y agoView on Hugging Face
2likes
utils.py260 linesDownload Raw Back to utils
1import dataclasses2import glob3import importlib4import random5import numpy as np6import torch7import warnings8import os9import time10import torch.utils.tensorboard as tensorboard11from torch import distributed as dist12import sys13import yaml14import json15import re16import pathlib17import matplotlib18matplotlib.use("Agg")19import matplotlib.pylab as plt20 21 22def plot_spectrogram(spectrogram):23    fig, ax = plt.subplots(figsize=(10, 2))24    im = ax.imshow(spectrogram, aspect="auto", origin="lower",25                   interpolation='none')26    plt.colorbar(im, ax=ax)27 28    fig.canvas.draw()29    plt.close()30 31    return fig32 33 34def seed_everything(seed, cudnn_deterministic=False):35    """36    Function that sets seed for pseudo-random number generators in:37    pytorch, numpy, python.random38    39    Args:40        seed: the integer value seed for global random state41    """42    if seed is not None:43        random.seed(seed)44        np.random.seed(seed)45        torch.manual_seed(seed)46        torch.cuda.manual_seed_all(seed)47 48    if cudnn_deterministic:49        torch.backends.cudnn.deterministic = True50        warnings.warn('You have chosen to seed training. '51                      'This will turn on the CUDNN deterministic setting, '52                      'which can slow down your training considerably! '53                      'You may see unexpected behavior when restarting '54                      'from checkpoints.')55 56def is_primary():57    return get_rank() == 058 59 60def get_rank():61    if not dist.is_available():62        return 063    if not dist.is_initialized():64        return 065 66    return dist.get_rank()67 68 69def load_yaml_config(path):70    with open(path) as f:71        config = yaml.full_load(f)72    return config73 74 75def save_config_to_yaml(config, path):76    assert path.endswith('.yaml')77    with open(path, 'w') as f:78        f.write(yaml.dump(config))79        f.close()80 81 82def save_dict_to_json(d, path, indent=None):83    json.dump(d, open(path, 'w'), indent=indent)84 85 86def load_dict_from_json(path):87    return json.load(open(path, 'r'))88 89 90def write_args(args, path):91    args_dict = dict((name, getattr(args, name)) for name in dir(args)if not name.startswith('_'))92    with open(path, 'a') as args_file:93        args_file.write('==> torch version: {}\n'.format(torch.__version__))94        args_file.write('==> cudnn version: {}\n'.format(torch.backends.cudnn.version()))95        args_file.write('==> Cmd:\n')96        args_file.write(str(sys.argv))97        args_file.write('\n==> args:\n')98        for k, v in sorted(args_dict.items()):99            args_file.write('  %s: %s\n' % (str(k), str(v)))100        args_file.close()101 102 103class Logger(object):104    def __init__(self, args):105        self.args = args106        self.save_dir = args.log_dir107        self.is_primary = is_primary()108        109        if self.is_primary:110            os.makedirs(self.save_dir, exist_ok=True)111            112            # save the args and config113            self.config_dir = os.path.join(self.save_dir, 'configs')114            os.makedirs(self.config_dir, exist_ok=True)115            file_name = os.path.join(self.config_dir, 'args.txt')116            write_args(args, file_name)117 118            log_dir = os.path.join(self.save_dir, 'logs')119            if not os.path.exists(log_dir):120                os.makedirs(log_dir, exist_ok=True)121            self.text_writer = open(os.path.join(log_dir, 'log.txt'), 'a') # 'w')122            if args.tensorboard:123                self.log_info('using tensorboard')124                self.tb_writer = torch.utils.tensorboard.SummaryWriter(log_dir=log_dir) # tensorboard.SummaryWriter(log_dir=log_dir)125            else:126                self.tb_writer = None127 128    def save_config(self, config):129        if self.is_primary:130            save_config_to_yaml(config, os.path.join(self.config_dir, 'config.yaml'))131 132    def log_info(self, info, check_primary=True):133        if self.is_primary or (not check_primary):134            print(info)135            if self.is_primary:136                info = str(info)137                time_str = time.strftime('%Y-%m-%d-%H-%M')138                info = '{}: {}'.format(time_str, info)139                if not info.endswith('\n'):140                    info += '\n'141                self.text_writer.write(info)142                self.text_writer.flush()143 144    def add_scalar(self, **kargs):145        """Log a scalar variable."""146        if self.is_primary:147            if self.tb_writer is not None:148                self.tb_writer.add_scalar(**kargs)149 150    def add_scalars(self, **kargs):151        """Log a scalar variable."""152        if self.is_primary:153            if self.tb_writer is not None:154                self.tb_writer.add_scalars(**kargs)155 156    def add_image(self, **kargs):157        """Log a scalar variable."""158        if self.is_primary:159            if self.tb_writer is not None:160                self.tb_writer.add_image(**kargs)161 162    def add_images(self, **kargs):163        """Log a scalar variable."""164        if self.is_primary:165            if self.tb_writer is not None:166                self.tb_writer.add_images(**kargs)167 168    def close(self):169        if self.is_primary:170            self.text_writer.close()171            self.tb_writer.close()172 173 174def cal_model_size(model, name=""):175 176    all_size = sum(p.numel() for p in model.parameters())/1024.0/1024.0177    return f'Model size of {name}: {all_size:.3f} MB'178 179    param_size = 0180    param_sum = 0181    for param in model.parameters():182        param_size += param.nelement() * param.element_size()183        param_sum += param.nelement()184    buffer_size = 0185    buffer_sum = 0186    for buffer in model.buffers():187        buffer_size += buffer.nelement() * buffer.element_size()188        buffer_sum += buffer.nelement()189    all_size = (param_size + buffer_size) / 1024 / 1024190 191    return f'Model size of {name}: {all_size:.3f} MB'192    # print(f'Model size of {name}: {all_size:.3f}MB')193    # return (param_size, param_sum, buffer_size, buffer_sum, all_size)194 195 196def load_obj(obj_path: str, default_obj_path: str = ''):197    """ Extract an object from a given path.198    Args:199        obj_path: Path to an object to be extracted, including the object name.200            e.g.: `src.trainers.meta_trainer.MetaTrainer`201                  `src.models.ada_style_speech.AdaStyleSpeechModel`202        default_obj_path: Default object path.203    204    Returns:205        Extracted object.206    Raises:207        AttributeError: When the object does not have the given named attribute.208    209    """210    obj_path_list = obj_path.rsplit('.', 1)211    obj_path = obj_path_list.pop(0) if len(obj_path_list) > 1 else default_obj_path212    obj_name = obj_path_list[0]213    module_obj = importlib.import_module(obj_path)214    if not hasattr(module_obj, obj_name):215        raise AttributeError(f'Object `{obj_name}` cannot be loaded from `{obj_path}`.')216    return getattr(module_obj, obj_name)217 218 219def to_device(data, device=None, dtype=None, non_blocking=False, copy=False):220    """Change the device of object recursively"""221    if isinstance(data, dict):222        return {223            k: to_device(v, device, dtype, non_blocking, copy) for k, v in data.items()224        }225    elif dataclasses.is_dataclass(data) and not isinstance(data, type):226        return type(data)(227            *[228                to_device(v, device, dtype, non_blocking, copy)229                for v in dataclasses.astuple(data)230            ]231        )232    # maybe namedtuple. I don't know the correct way to judge namedtuple.233    elif isinstance(data, tuple) and type(data) is not tuple:234        return type(data)(235            *[to_device(o, device, dtype, non_blocking, copy) for o in data]236        )237    elif isinstance(data, (list, tuple)):238        return type(data)(to_device(v, device, dtype, non_blocking, copy) for v in data)239    elif isinstance(data, np.ndarray):240        return to_device(torch.from_numpy(data), device, dtype, non_blocking, copy)241    elif isinstance(data, torch.Tensor):242        return data.to(device, dtype, non_blocking, copy)243    else:244        return data245 246 247def save_checkpoint(filepath, obj, ext='pth', num_ckpt_keep=10):248    ckpts = sorted(pathlib.Path(filepath).parent.glob(f'*.{ext}'))249    if len(ckpts) > num_ckpt_keep:250        [os.remove(c) for c in ckpts[:-num_ckpt_keep]]251    torch.save(obj, filepath)252 253 254def scan_checkpoint(cp_dir, prefix='ckpt_'):255    pattern = os.path.join(cp_dir, prefix + '????????.pth')256    cp_list = glob.glob(pattern)257    if len(cp_list) == 0:258        return None259    return sorted(cp_list)[-1]260