DD0101/VITS
1
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 TextAudioLoader,20 TextAudioCollate,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 = TextAudioLoader(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 = TextAudioCollate()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 = TextAudioLoader(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 **hps.model).cuda(rank)88 net_d = MultiPeriodDiscriminator(hps.model.use_spectral_norm).cuda(rank)89 optim_g = torch.optim.AdamW(90 net_g.parameters(), 91 hps.train.learning_rate, 92 betas=hps.train.betas, 93 eps=hps.train.eps)94 optim_d = torch.optim.AdamW(95 net_d.parameters(),96 hps.train.learning_rate, 97 betas=hps.train.betas, 98 eps=hps.train.eps)99 net_g = DDP(net_g, device_ids=[rank])100 net_d = DDP(net_d, device_ids=[rank])101 102 try:103 _, _, _, epoch_str = utils.load_checkpoint(utils.latest_checkpoint_path(hps.model_dir, "G_*.pth"), net_g, optim_g)104 _, _, _, epoch_str = utils.load_checkpoint(utils.latest_checkpoint_path(hps.model_dir, "D_*.pth"), net_d, optim_d)105 global_step = (epoch_str - 1) * len(train_loader)106 except:107 epoch_str = 1108 global_step = 0109 110 scheduler_g = torch.optim.lr_scheduler.ExponentialLR(optim_g, gamma=hps.train.lr_decay, last_epoch=epoch_str-2)111 scheduler_d = torch.optim.lr_scheduler.ExponentialLR(optim_d, gamma=hps.train.lr_decay, last_epoch=epoch_str-2)112 113 scaler = GradScaler(enabled=hps.train.fp16_run)114 115 for epoch in range(epoch_str, hps.train.epochs + 1):116 if rank==0:117 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])118 else:119 train_and_evaluate(rank, epoch, hps, [net_g, net_d], [optim_g, optim_d], [scheduler_g, scheduler_d], scaler, [train_loader, None], None, None)120 scheduler_g.step()121 scheduler_d.step()122 123 124def train_and_evaluate(rank, epoch, hps, nets, optims, schedulers, scaler, loaders, logger, writers):125 net_g, net_d = nets126 optim_g, optim_d = optims127 scheduler_g, scheduler_d = schedulers128 train_loader, eval_loader = loaders129 if writers is not None:130 writer, writer_eval = writers131 132 train_loader.batch_sampler.set_epoch(epoch)133 global global_step134 135 net_g.train()136 net_d.train()137 for batch_idx, (x, x_lengths, spec, spec_lengths, y, y_lengths) in enumerate(train_loader):138 x, x_lengths = x.cuda(rank, non_blocking=True), x_lengths.cuda(rank, non_blocking=True)139 spec, spec_lengths = spec.cuda(rank, non_blocking=True), spec_lengths.cuda(rank, non_blocking=True)140 y, y_lengths = y.cuda(rank, non_blocking=True), y_lengths.cuda(rank, non_blocking=True)141 142 with autocast(enabled=hps.train.fp16_run):143 y_hat, l_length, attn, ids_slice, x_mask, z_mask,\144 (z, z_p, m_p, logs_p, m_q, logs_q) = net_g(x, x_lengths, spec, spec_lengths)145 146 mel = spec_to_mel_torch(147 spec, 148 hps.data.filter_length, 149 hps.data.n_mel_channels, 150 hps.data.sampling_rate,151 hps.data.mel_fmin, 152 hps.data.mel_fmax)153 y_mel = commons.slice_segments(mel, ids_slice, hps.train.segment_size // hps.data.hop_length)154 y_hat_mel = mel_spectrogram_torch(155 y_hat.squeeze(1), 156 hps.data.filter_length, 157 hps.data.n_mel_channels, 158 hps.data.sampling_rate, 159 hps.data.hop_length, 160 hps.data.win_length, 161 hps.data.mel_fmin, 162 hps.data.mel_fmax163 )164 165 y = commons.slice_segments(y, ids_slice * hps.data.hop_length, hps.train.segment_size) # slice 166 167 # Discriminator168 y_d_hat_r, y_d_hat_g, _, _ = net_d(y, y_hat.detach())169 with autocast(enabled=False):170 loss_disc, losses_disc_r, losses_disc_g = discriminator_loss(y_d_hat_r, y_d_hat_g)171 loss_disc_all = loss_disc172 optim_d.zero_grad()173 scaler.scale(loss_disc_all).backward()174 scaler.unscale_(optim_d)175 grad_norm_d = commons.clip_grad_value_(net_d.parameters(), None)176 scaler.step(optim_d)177 178 with autocast(enabled=hps.train.fp16_run):179 # Generator180 y_d_hat_r, y_d_hat_g, fmap_r, fmap_g = net_d(y, y_hat)181 with autocast(enabled=False):182 loss_dur = torch.sum(l_length.float())183 loss_mel = F.l1_loss(y_mel, y_hat_mel) * hps.train.c_mel184 loss_kl = kl_loss(z_p, logs_q, m_p, logs_p, z_mask) * hps.train.c_kl185 186 loss_fm = feature_loss(fmap_r, fmap_g)187 loss_gen, losses_gen = generator_loss(y_d_hat_g)188 loss_gen_all = loss_gen + loss_fm + loss_mel + loss_dur + loss_kl189 optim_g.zero_grad()190 scaler.scale(loss_gen_all).backward()191 scaler.unscale_(optim_g)192 grad_norm_g = commons.clip_grad_value_(net_g.parameters(), None)193 scaler.step(optim_g)194 scaler.update()195 196 if rank==0:197 if global_step % hps.train.log_interval == 0:198 lr = optim_g.param_groups[0]['lr']199 losses = [loss_disc, loss_gen, loss_fm, loss_mel, loss_dur, loss_kl]200 logger.info('Train Epoch: {} [{:.0f}%]'.format(201 epoch,202 100. * batch_idx / len(train_loader)))203 logger.info([x.item() for x in losses] + [global_step, lr])204 205 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}206 scalar_dict.update({"loss/g/fm": loss_fm, "loss/g/mel": loss_mel, "loss/g/dur": loss_dur, "loss/g/kl": loss_kl})207 208 scalar_dict.update({"loss/g/{}".format(i): v for i, v in enumerate(losses_gen)})209 scalar_dict.update({"loss/d_r/{}".format(i): v for i, v in enumerate(losses_disc_r)})210 scalar_dict.update({"loss/d_g/{}".format(i): v for i, v in enumerate(losses_disc_g)})211 image_dict = { 212 "slice/mel_org": utils.plot_spectrogram_to_numpy(y_mel[0].data.cpu().numpy()),213 "slice/mel_gen": utils.plot_spectrogram_to_numpy(y_hat_mel[0].data.cpu().numpy()), 214 "all/mel": utils.plot_spectrogram_to_numpy(mel[0].data.cpu().numpy()),215 "all/attn": utils.plot_alignment_to_numpy(attn[0,0].data.cpu().numpy())216 }217 utils.summarize(218 writer=writer,219 global_step=global_step, 220 images=image_dict,221 scalars=scalar_dict)222 223 if global_step % hps.train.eval_interval == 0:224 evaluate(hps, net_g, eval_loader, writer_eval)225 utils.save_checkpoint(net_g, optim_g, hps.train.learning_rate, epoch, os.path.join(hps.model_dir, "G_{}.pth".format(global_step)))226 utils.save_checkpoint(net_d, optim_d, hps.train.learning_rate, epoch, os.path.join(hps.model_dir, "D_{}.pth".format(global_step)))227 global_step += 1228 229 if rank == 0:230 logger.info('====> Epoch: {}'.format(epoch))231 232 233def evaluate(hps, generator, eval_loader, writer_eval):234 generator.eval()235 with torch.no_grad():236 for batch_idx, (x, x_lengths, spec, spec_lengths, y, y_lengths) in enumerate(eval_loader):237 x, x_lengths = x.cuda(0), x_lengths.cuda(0)238 spec, spec_lengths = spec.cuda(0), spec_lengths.cuda(0)239 y, y_lengths = y.cuda(0), y_lengths.cuda(0)240 241 # remove else242 x = x[:1]243 x_lengths = x_lengths[:1]244 spec = spec[:1]245 spec_lengths = spec_lengths[:1]246 y = y[:1]247 y_lengths = y_lengths[:1]248 break249 y_hat, attn, mask, *_ = generator.module.infer(x, x_lengths, max_len=1000)250 y_hat_lengths = mask.sum([1,2]).long() * hps.data.hop_length251 252 mel = spec_to_mel_torch(253 spec, 254 hps.data.filter_length, 255 hps.data.n_mel_channels, 256 hps.data.sampling_rate,257 hps.data.mel_fmin, 258 hps.data.mel_fmax)259 y_hat_mel = mel_spectrogram_torch(260 y_hat.squeeze(1).float(),261 hps.data.filter_length,262 hps.data.n_mel_channels,263 hps.data.sampling_rate,264 hps.data.hop_length,265 hps.data.win_length,266 hps.data.mel_fmin,267 hps.data.mel_fmax268 )269 image_dict = {270 "gen/mel": utils.plot_spectrogram_to_numpy(y_hat_mel[0].cpu().numpy())271 }272 audio_dict = {273 "gen/audio": y_hat[0,:,:y_hat_lengths[0]]274 }275 if global_step == 0:276 image_dict.update({"gt/mel": utils.plot_spectrogram_to_numpy(mel[0].cpu().numpy())})277 audio_dict.update({"gt/audio": y[0,:,:y_lengths[0]]})278 279 utils.summarize(280 writer=writer_eval,281 global_step=global_step, 282 images=image_dict,283 audios=audio_dict,284 audio_sampling_rate=hps.data.sampling_rate285 )286 generator.train()287 288 289if __name__ == "__main__":290 main()291 