justyoung/DiffSinger
1
1import glob2import re3import subprocess4from datetime import datetime5 6import matplotlib7 8matplotlib.use('Agg')9 10from utils.hparams import hparams, set_hparams11import random12import sys13import numpy as np14import torch.distributed as dist15from pytorch_lightning.loggers import TensorBoardLogger16from utils.pl_utils import LatestModelCheckpoint, BaseTrainer, data_loader, DDP17from torch import nn18import torch.utils.data19import utils20import logging21import os22 23torch.multiprocessing.set_sharing_strategy(os.getenv('TORCH_SHARE_STRATEGY', 'file_system'))24 25log_format = '%(asctime)s %(message)s'26logging.basicConfig(stream=sys.stdout, level=logging.INFO,27 format=log_format, datefmt='%m/%d %I:%M:%S %p')28 29 30class BaseDataset(torch.utils.data.Dataset):31 def __init__(self, shuffle):32 super().__init__()33 self.hparams = hparams34 self.shuffle = shuffle35 self.sort_by_len = hparams['sort_by_len']36 self.sizes = None37 38 @property39 def _sizes(self):40 return self.sizes41 42 def __getitem__(self, index):43 raise NotImplementedError44 45 def collater(self, samples):46 raise NotImplementedError47 48 def __len__(self):49 return len(self._sizes)50 51 def num_tokens(self, index):52 return self.size(index)53 54 def size(self, index):55 """Return an example's size as a float or tuple. This value is used when56 filtering a dataset with ``--max-positions``."""57 size = min(self._sizes[index], hparams['max_frames'])58 return size59 60 def ordered_indices(self):61 """Return an ordered list of indices. Batches will be constructed based62 on this order."""63 if self.shuffle:64 indices = np.random.permutation(len(self))65 if self.sort_by_len:66 indices = indices[np.argsort(np.array(self._sizes)[indices], kind='mergesort')]67 # 先random, 然后稳定排序, 保证排序后同长度的数据顺序是依照random permutation的 (被其随机打乱).68 else:69 indices = np.arange(len(self))70 return indices71 72 @property73 def num_workers(self):74 return int(os.getenv('NUM_WORKERS', hparams['ds_workers']))75 76 77class BaseTask(nn.Module):78 def __init__(self, *args, **kwargs):79 # dataset configs80 super(BaseTask, self).__init__(*args, **kwargs)81 self.current_epoch = 082 self.global_step = 083 self.loaded_optimizer_states_dict = {}84 self.trainer = None85 self.logger = None86 self.on_gpu = False87 self.use_dp = False88 self.use_ddp = False89 self.example_input_array = None90 91 self.max_tokens = hparams['max_tokens']92 self.max_sentences = hparams['max_sentences']93 self.max_eval_tokens = hparams['max_eval_tokens']94 if self.max_eval_tokens == -1:95 hparams['max_eval_tokens'] = self.max_eval_tokens = self.max_tokens96 self.max_eval_sentences = hparams['max_eval_sentences']97 if self.max_eval_sentences == -1:98 hparams['max_eval_sentences'] = self.max_eval_sentences = self.max_sentences99 100 self.model = None101 self.training_losses_meter = None102 103 ###########104 # Training, validation and testing105 ###########106 def build_model(self):107 raise NotImplementedError108 109 def load_ckpt(self, ckpt_base_dir, current_model_name=None, model_name='model', force=True, strict=True):110 # This function is updated on 2021.12.13111 if current_model_name is None:112 current_model_name = model_name113 utils.load_ckpt(self.__getattr__(current_model_name), ckpt_base_dir, current_model_name, force, strict)114 115 def on_epoch_start(self):116 self.training_losses_meter = {'total_loss': utils.AvgrageMeter()}117 118 def _training_step(self, sample, batch_idx, optimizer_idx):119 """120 121 :param sample:122 :param batch_idx:123 :return: total loss: torch.Tensor, loss_log: dict124 """125 raise NotImplementedError126 127 def training_step(self, sample, batch_idx, optimizer_idx=-1):128 loss_ret = self._training_step(sample, batch_idx, optimizer_idx)129 self.opt_idx = optimizer_idx130 if loss_ret is None:131 return {'loss': None}132 total_loss, log_outputs = loss_ret133 log_outputs = utils.tensors_to_scalars(log_outputs)134 for k, v in log_outputs.items():135 if k not in self.training_losses_meter:136 self.training_losses_meter[k] = utils.AvgrageMeter()137 if not np.isnan(v):138 self.training_losses_meter[k].update(v)139 self.training_losses_meter['total_loss'].update(total_loss.item())140 141 try:142 log_outputs['lr'] = self.scheduler.get_lr()143 if isinstance(log_outputs['lr'], list):144 log_outputs['lr'] = log_outputs['lr'][0]145 except:146 pass147 148 # log_outputs['all_loss'] = total_loss.item()149 progress_bar_log = log_outputs150 tb_log = {f'tr/{k}': v for k, v in log_outputs.items()}151 return {152 'loss': total_loss,153 'progress_bar': progress_bar_log,154 'log': tb_log155 }156 157 def optimizer_step(self, epoch, batch_idx, optimizer, optimizer_idx):158 optimizer.step()159 optimizer.zero_grad()160 if self.scheduler is not None:161 self.scheduler.step(self.global_step // hparams['accumulate_grad_batches'])162 163 def on_epoch_end(self):164 loss_outputs = {k: round(v.avg, 4) for k, v in self.training_losses_meter.items()}165 print(f"\n==============\n "166 f"Epoch {self.current_epoch} ended. Steps: {self.global_step}. {loss_outputs}"167 f"\n==============\n")168 169 def validation_step(self, sample, batch_idx):170 """171 172 :param sample:173 :param batch_idx:174 :return: output: dict175 """176 raise NotImplementedError177 178 def _validation_end(self, outputs):179 """180 181 :param outputs:182 :return: loss_output: dict183 """184 raise NotImplementedError185 186 def validation_end(self, outputs):187 loss_output = self._validation_end(outputs)188 print(f"\n==============\n "189 f"valid results: {loss_output}"190 f"\n==============\n")191 return {192 'log': {f'val/{k}': v for k, v in loss_output.items()},193 'val_loss': loss_output['total_loss']194 }195 196 def build_scheduler(self, optimizer):197 raise NotImplementedError198 199 def build_optimizer(self, model):200 raise NotImplementedError201 202 def configure_optimizers(self):203 optm = self.build_optimizer(self.model)204 self.scheduler = self.build_scheduler(optm)205 return [optm]206 207 def test_start(self):208 pass209 210 def test_step(self, sample, batch_idx):211 return self.validation_step(sample, batch_idx)212 213 def test_end(self, outputs):214 return self.validation_end(outputs)215 216 ###########217 # Running configuration218 ###########219 220 @classmethod221 def start(cls):222 set_hparams()223 os.environ['MASTER_PORT'] = str(random.randint(15000, 30000))224 random.seed(hparams['seed'])225 np.random.seed(hparams['seed'])226 task = cls()227 work_dir = hparams['work_dir']228 trainer = BaseTrainer(checkpoint_callback=LatestModelCheckpoint(229 filepath=work_dir,230 verbose=True,231 monitor='val_loss',232 mode='min',233 num_ckpt_keep=hparams['num_ckpt_keep'],234 save_best=hparams['save_best'],235 period=1 if hparams['save_ckpt'] else 100000236 ),237 logger=TensorBoardLogger(238 save_dir=work_dir,239 name='lightning_logs',240 version='lastest'241 ),242 gradient_clip_val=hparams['clip_grad_norm'],243 val_check_interval=hparams['val_check_interval'],244 row_log_interval=hparams['log_interval'],245 max_updates=hparams['max_updates'],246 num_sanity_val_steps=hparams['num_sanity_val_steps'] if not hparams[247 'validate'] else 10000,248 accumulate_grad_batches=hparams['accumulate_grad_batches'])249 if not hparams['infer']: # train250 t = datetime.now().strftime('%Y%m%d%H%M%S')251 code_dir = f'{work_dir}/codes/{t}'252 subprocess.check_call(f'mkdir -p "{code_dir}"', shell=True)253 for c in hparams['save_codes']:254 subprocess.check_call(f'cp -r "{c}" "{code_dir}/"', shell=True)255 print(f"| Copied codes to {code_dir}.")256 trainer.checkpoint_callback.task = task257 trainer.fit(task)258 else:259 trainer.test(task)260 261 def configure_ddp(self, model, device_ids):262 model = DDP(263 model,264 device_ids=device_ids,265 find_unused_parameters=True266 )267 if dist.get_rank() != 0 and not hparams['debug']:268 sys.stdout = open(os.devnull, "w")269 sys.stderr = open(os.devnull, "w")270 random.seed(hparams['seed'])271 np.random.seed(hparams['seed'])272 return model273 274 def training_end(self, *args, **kwargs):275 return None276 277 def init_ddp_connection(self, proc_rank, world_size):278 set_hparams(print_hparams=False)279 # guarantees unique ports across jobs from same grid search280 default_port = 12910281 # if user gave a port number, use that one instead282 try:283 default_port = os.environ['MASTER_PORT']284 except Exception:285 os.environ['MASTER_PORT'] = str(default_port)286 287 # figure out the root node addr288 root_node = '127.0.0.2'289 root_node = self.trainer.resolve_root_node_address(root_node)290 os.environ['MASTER_ADDR'] = root_node291 dist.init_process_group('nccl', rank=proc_rank, world_size=world_size)292 293 @data_loader294 def train_dataloader(self):295 return None296 297 @data_loader298 def test_dataloader(self):299 return None300 301 @data_loader302 def val_dataloader(self):303 return None304 305 def on_load_checkpoint(self, checkpoint):306 pass307 308 def on_save_checkpoint(self, checkpoint):309 pass310 311 def on_sanity_check_start(self):312 pass313 314 def on_train_start(self):315 pass316 317 def on_train_end(self):318 pass319 320 def on_batch_start(self, batch):321 pass322 323 def on_batch_end(self):324 pass325 326 def on_pre_performance_check(self):327 pass328 329 def on_post_performance_check(self):330 pass331 332 def on_before_zero_grad(self, optimizer):333 pass334 335 def on_after_backward(self):336 pass337 338 def backward(self, loss, optimizer):339 loss.backward()340 341 def grad_norm(self, norm_type):342 results = {}343 total_norm = 0344 for name, p in self.named_parameters():345 if p.requires_grad:346 try:347 param_norm = p.grad.data.norm(norm_type)348 total_norm += param_norm ** norm_type349 norm = param_norm ** (1 / norm_type)350 351 grad = round(norm.data.cpu().numpy().flatten()[0], 3)352 results['grad_{}_norm_{}'.format(norm_type, name)] = grad353 except Exception:354 # this param had no grad355 pass356 357 total_norm = total_norm ** (1. / norm_type)358 grad = round(total_norm.data.cpu().numpy().flatten()[0], 3)359 results['grad_{}_norm_total'.format(norm_type)] = grad360 return results361 