CoolFace
Apppublic

DD0101/VITS

sourceHugging Faceupdated 2y agoView on Hugging Face
1likes
train_ms.py295 linesDownload Raw Back to root
1import os2import json3import argparse4import itertools5import math6import torch7from torch import nn, optim8from torch.nn import functional as F9from torch.utils.data import DataLoader10from torch.utils.tensorboard import SummaryWriter11import torch.multiprocessing as mp12import torch.distributed as dist13from torch.nn.parallel import DistributedDataParallel as DDP14from torch.cuda.amp import autocast, GradScaler15 16import commons17import utils18from data_utils import (19  TextAudioSpeakerLoader,20  TextAudioSpeakerCollate,21  DistributedBucketSampler22)23from models import (24  SynthesizerTrn,25  MultiPeriodDiscriminator,26)27from losses import (28  generator_loss,29  discriminator_loss,30  feature_loss,31  kl_loss32)33from mel_processing import mel_spectrogram_torch, spec_to_mel_torch34from text.symbols import symbols35 36 37torch.backends.cudnn.benchmark = True38global_step = 039 40 41def main():42  """Assume Single Node Multi GPUs Training Only"""43  assert torch.cuda.is_available(), "CPU training is not allowed."44 45  n_gpus = torch.cuda.device_count()46  os.environ['MASTER_ADDR'] = 'localhost'47  os.environ['MASTER_PORT'] = '80000'48 49  hps = utils.get_hparams()50  mp.spawn(run, nprocs=n_gpus, args=(n_gpus, hps,))51 52 53def run(rank, n_gpus, hps):54  global global_step55  if rank == 0:56    logger = utils.get_logger(hps.model_dir)57    logger.info(hps)58    utils.check_git_hash(hps.model_dir)59    writer = SummaryWriter(log_dir=hps.model_dir)60    writer_eval = SummaryWriter(log_dir=os.path.join(hps.model_dir, "eval"))61 62  dist.init_process_group(backend='nccl', init_method='env://', world_size=n_gpus, rank=rank)63  torch.manual_seed(hps.train.seed)64  torch.cuda.set_device(rank)65 66  train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data)67  train_sampler = DistributedBucketSampler(68      train_dataset,69      hps.train.batch_size,70      [32,300,400,500,600,700,800,900,1000],71      num_replicas=n_gpus,72      rank=rank,73      shuffle=True)74  collate_fn = TextAudioSpeakerCollate()75  train_loader = DataLoader(train_dataset, num_workers=8, shuffle=False, pin_memory=True,76      collate_fn=collate_fn, batch_sampler=train_sampler)77  if rank == 0:78    eval_dataset = TextAudioSpeakerLoader(hps.data.validation_files, hps.data)79    eval_loader = DataLoader(eval_dataset, num_workers=8, shuffle=False,80        batch_size=hps.train.batch_size, pin_memory=True,81        drop_last=False, collate_fn=collate_fn)82 83  net_g = SynthesizerTrn(84      len(symbols),85      hps.data.filter_length // 2 + 1,86      hps.train.segment_size // hps.data.hop_length,87      n_speakers=hps.data.n_speakers,88      **hps.model).cuda(rank)89  net_d = MultiPeriodDiscriminator(hps.model.use_spectral_norm).cuda(rank)90  optim_g = torch.optim.AdamW(91      net_g.parameters(), 92      hps.train.learning_rate, 93      betas=hps.train.betas, 94      eps=hps.train.eps)95  optim_d = torch.optim.AdamW(96      net_d.parameters(),97      hps.train.learning_rate, 98      betas=hps.train.betas, 99      eps=hps.train.eps)100  net_g = DDP(net_g, device_ids=[rank])101  net_d = DDP(net_d, device_ids=[rank])102 103  try:104    _, _, _, epoch_str = utils.load_checkpoint(utils.latest_checkpoint_path(hps.model_dir, "G_*.pth"), net_g, optim_g)105    _, _, _, epoch_str = utils.load_checkpoint(utils.latest_checkpoint_path(hps.model_dir, "D_*.pth"), net_d, optim_d)106    global_step = (epoch_str - 1) * len(train_loader)107  except:108    epoch_str = 1109    global_step = 0110 111  scheduler_g = torch.optim.lr_scheduler.ExponentialLR(optim_g, gamma=hps.train.lr_decay, last_epoch=epoch_str-2)112  scheduler_d = torch.optim.lr_scheduler.ExponentialLR(optim_d, gamma=hps.train.lr_decay, last_epoch=epoch_str-2)113 114  scaler = GradScaler(enabled=hps.train.fp16_run)115 116  for epoch in range(epoch_str, hps.train.epochs + 1):117    if rank==0:118      train_and_evaluate(rank, epoch, hps, [net_g, net_d], [optim_g, optim_d], [scheduler_g, scheduler_d], scaler, [train_loader, eval_loader], logger, [writer, writer_eval])119    else:120      train_and_evaluate(rank, epoch, hps, [net_g, net_d], [optim_g, optim_d], [scheduler_g, scheduler_d], scaler, [train_loader, None], None, None)121    scheduler_g.step()122    scheduler_d.step()123 124 125def train_and_evaluate(rank, epoch, hps, nets, optims, schedulers, scaler, loaders, logger, writers):126  net_g, net_d = nets127  optim_g, optim_d = optims128  scheduler_g, scheduler_d = schedulers129  train_loader, eval_loader = loaders130  if writers is not None:131    writer, writer_eval = writers132 133  train_loader.batch_sampler.set_epoch(epoch)134  global global_step135 136  net_g.train()137  net_d.train()138  for batch_idx, (x, x_lengths, spec, spec_lengths, y, y_lengths, speakers) in enumerate(train_loader):139    x, x_lengths = x.cuda(rank, non_blocking=True), x_lengths.cuda(rank, non_blocking=True)140    spec, spec_lengths = spec.cuda(rank, non_blocking=True), spec_lengths.cuda(rank, non_blocking=True)141    y, y_lengths = y.cuda(rank, non_blocking=True), y_lengths.cuda(rank, non_blocking=True)142    speakers = speakers.cuda(rank, non_blocking=True)143 144    with autocast(enabled=hps.train.fp16_run):145      y_hat, l_length, attn, ids_slice, x_mask, z_mask,\146      (z, z_p, m_p, logs_p, m_q, logs_q) = net_g(x, x_lengths, spec, spec_lengths, speakers)147 148      mel = spec_to_mel_torch(149          spec, 150          hps.data.filter_length, 151          hps.data.n_mel_channels, 152          hps.data.sampling_rate,153          hps.data.mel_fmin, 154          hps.data.mel_fmax)155      y_mel = commons.slice_segments(mel, ids_slice, hps.train.segment_size // hps.data.hop_length)156      y_hat_mel = mel_spectrogram_torch(157          y_hat.squeeze(1), 158          hps.data.filter_length, 159          hps.data.n_mel_channels, 160          hps.data.sampling_rate, 161          hps.data.hop_length, 162          hps.data.win_length, 163          hps.data.mel_fmin, 164          hps.data.mel_fmax165      )166 167      y = commons.slice_segments(y, ids_slice * hps.data.hop_length, hps.train.segment_size) # slice 168 169      # Discriminator170      y_d_hat_r, y_d_hat_g, _, _ = net_d(y, y_hat.detach())171      with autocast(enabled=False):172        loss_disc, losses_disc_r, losses_disc_g = discriminator_loss(y_d_hat_r, y_d_hat_g)173        loss_disc_all = loss_disc174    optim_d.zero_grad()175    scaler.scale(loss_disc_all).backward()176    scaler.unscale_(optim_d)177    grad_norm_d = commons.clip_grad_value_(net_d.parameters(), None)178    scaler.step(optim_d)179 180    with autocast(enabled=hps.train.fp16_run):181      # Generator182      y_d_hat_r, y_d_hat_g, fmap_r, fmap_g = net_d(y, y_hat)183      with autocast(enabled=False):184        loss_dur = torch.sum(l_length.float())185        loss_mel = F.l1_loss(y_mel, y_hat_mel) * hps.train.c_mel186        loss_kl = kl_loss(z_p, logs_q, m_p, logs_p, z_mask) * hps.train.c_kl187 188        loss_fm = feature_loss(fmap_r, fmap_g)189        loss_gen, losses_gen = generator_loss(y_d_hat_g)190        loss_gen_all = loss_gen + loss_fm + loss_mel + loss_dur + loss_kl191    optim_g.zero_grad()192    scaler.scale(loss_gen_all).backward()193    scaler.unscale_(optim_g)194    grad_norm_g = commons.clip_grad_value_(net_g.parameters(), None)195    scaler.step(optim_g)196    scaler.update()197 198    if rank==0:199      if global_step % hps.train.log_interval == 0:200        lr = optim_g.param_groups[0]['lr']201        losses = [loss_disc, loss_gen, loss_fm, loss_mel, loss_dur, loss_kl]202        logger.info('Train Epoch: {} [{:.0f}%]'.format(203          epoch,204          100. * batch_idx / len(train_loader)))205        logger.info([x.item() for x in losses] + [global_step, lr])206        207        scalar_dict = {"loss/g/total": loss_gen_all, "loss/d/total": loss_disc_all, "learning_rate": lr, "grad_norm_d": grad_norm_d, "grad_norm_g": grad_norm_g}208        scalar_dict.update({"loss/g/fm": loss_fm, "loss/g/mel": loss_mel, "loss/g/dur": loss_dur, "loss/g/kl": loss_kl})209 210        scalar_dict.update({"loss/g/{}".format(i): v for i, v in enumerate(losses_gen)})211        scalar_dict.update({"loss/d_r/{}".format(i): v for i, v in enumerate(losses_disc_r)})212        scalar_dict.update({"loss/d_g/{}".format(i): v for i, v in enumerate(losses_disc_g)})213        image_dict = { 214            "slice/mel_org": utils.plot_spectrogram_to_numpy(y_mel[0].data.cpu().numpy()),215            "slice/mel_gen": utils.plot_spectrogram_to_numpy(y_hat_mel[0].data.cpu().numpy()), 216            "all/mel": utils.plot_spectrogram_to_numpy(mel[0].data.cpu().numpy()),217            "all/attn": utils.plot_alignment_to_numpy(attn[0,0].data.cpu().numpy())218        }219        utils.summarize(220          writer=writer,221          global_step=global_step, 222          images=image_dict,223          scalars=scalar_dict)224 225      if global_step % hps.train.eval_interval == 0:226        evaluate(hps, net_g, eval_loader, writer_eval)227        utils.save_checkpoint(net_g, optim_g, hps.train.learning_rate, epoch, os.path.join(hps.model_dir, "G_{}.pth".format(global_step)))228        utils.save_checkpoint(net_d, optim_d, hps.train.learning_rate, epoch, os.path.join(hps.model_dir, "D_{}.pth".format(global_step)))229    global_step += 1230  231  if rank == 0:232    logger.info('====> Epoch: {}'.format(epoch))233 234 235def evaluate(hps, generator, eval_loader, writer_eval):236    generator.eval()237    with torch.no_grad():238      for batch_idx, (x, x_lengths, spec, spec_lengths, y, y_lengths, speakers) in enumerate(eval_loader):239        x, x_lengths = x.cuda(0), x_lengths.cuda(0)240        spec, spec_lengths = spec.cuda(0), spec_lengths.cuda(0)241        y, y_lengths = y.cuda(0), y_lengths.cuda(0)242        speakers = speakers.cuda(0)243 244        # remove else245        x = x[:1]246        x_lengths = x_lengths[:1]247        spec = spec[:1]248        spec_lengths = spec_lengths[:1]249        y = y[:1]250        y_lengths = y_lengths[:1]251        speakers = speakers[:1]252        break253      y_hat, attn, mask, *_ = generator.module.infer(x, x_lengths, speakers, max_len=1000)254      y_hat_lengths = mask.sum([1,2]).long() * hps.data.hop_length255 256      mel = spec_to_mel_torch(257        spec, 258        hps.data.filter_length, 259        hps.data.n_mel_channels, 260        hps.data.sampling_rate,261        hps.data.mel_fmin, 262        hps.data.mel_fmax)263      y_hat_mel = mel_spectrogram_torch(264        y_hat.squeeze(1).float(),265        hps.data.filter_length,266        hps.data.n_mel_channels,267        hps.data.sampling_rate,268        hps.data.hop_length,269        hps.data.win_length,270        hps.data.mel_fmin,271        hps.data.mel_fmax272      )273    image_dict = {274      "gen/mel": utils.plot_spectrogram_to_numpy(y_hat_mel[0].cpu().numpy())275    }276    audio_dict = {277      "gen/audio": y_hat[0,:,:y_hat_lengths[0]]278    }279    if global_step == 0:280      image_dict.update({"gt/mel": utils.plot_spectrogram_to_numpy(mel[0].cpu().numpy())})281      audio_dict.update({"gt/audio": y[0,:,:y_lengths[0]]})282 283    utils.summarize(284      writer=writer_eval,285      global_step=global_step, 286      images=image_dict,287      audios=audio_dict,288      audio_sampling_rate=hps.data.sampling_rate289    )290    generator.train()291 292                           293if __name__ == "__main__":294  main()295