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 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 