scrapegoat/Neural-Audio-Codec
2
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 