CoolFace
Apppublic

justyoung/DiffSinger

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
base_task.py361 linesDownload Raw Back to tasks
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