kwau/sovits-isla
0
1import logging2import multiprocessing3import os4import time5 6import torch7import torch.distributed as dist8import torch.multiprocessing as mp9from torch.cuda.amp import GradScaler, autocast10from torch.nn import functional as F11from torch.nn.parallel import DistributedDataParallel as DDP12from torch.utils.data import DataLoader13from torch.utils.tensorboard import SummaryWriter14 15import modules.commons as commons16import utils17from data_utils import TextAudioCollate, TextAudioSpeakerLoader18from models import (19 MultiPeriodDiscriminator,20 SynthesizerTrn,21)22from modules.losses import discriminator_loss, feature_loss, generator_loss, kl_loss23from modules.mel_processing import mel_spectrogram_torch, spec_to_mel_torch24 25logging.getLogger('matplotlib').setLevel(logging.WARNING)26logging.getLogger('numba').setLevel(logging.WARNING)27 28torch.backends.cudnn.benchmark = True29global_step = 030start_time = time.time()31 32# os.environ['TORCH_DISTRIBUTED_DEBUG'] = 'INFO'33 34 35def main():36 """Assume Single Node Multi GPUs Training Only"""37 assert torch.cuda.is_available(), "CPU training is not allowed."38 hps = utils.get_hparams()39 40 n_gpus = torch.cuda.device_count()41 os.environ['MASTER_ADDR'] = 'localhost'42 os.environ['MASTER_PORT'] = hps.train.port43 44 mp.spawn(run, nprocs=n_gpus, args=(n_gpus, hps,))45 46 47def run(rank, n_gpus, hps):48 global global_step49 if rank == 0:50 logger = utils.get_logger(hps.model_dir)51 logger.info(hps)52 utils.check_git_hash(hps.model_dir)53 writer = SummaryWriter(log_dir=hps.model_dir)54 writer_eval = SummaryWriter(log_dir=os.path.join(hps.model_dir, "eval"))55 56 # for pytorch on win, backend use gloo 57 dist.init_process_group(backend= 'gloo' if os.name == 'nt' else 'nccl', init_method='env://', world_size=n_gpus, rank=rank)58 torch.manual_seed(hps.train.seed)59 torch.cuda.set_device(rank)60 collate_fn = TextAudioCollate()61 all_in_mem = hps.train.all_in_mem # If you have enough memory, turn on this option to avoid disk IO and speed up training.62 train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps, all_in_mem=all_in_mem)63 num_workers = 5 if multiprocessing.cpu_count() > 4 else multiprocessing.cpu_count()64 if all_in_mem:65 num_workers = 066 train_loader = DataLoader(train_dataset, num_workers=num_workers, shuffle=False, pin_memory=True,67 batch_size=hps.train.batch_size, collate_fn=collate_fn)68 if rank == 0:69 eval_dataset = TextAudioSpeakerLoader(hps.data.validation_files, hps, all_in_mem=all_in_mem,vol_aug = False)70 eval_loader = DataLoader(eval_dataset, num_workers=1, shuffle=False,71 batch_size=1, pin_memory=False,72 drop_last=False, collate_fn=collate_fn)73 74 net_g = SynthesizerTrn(75 hps.data.filter_length // 2 + 1,76 hps.train.segment_size // hps.data.hop_length,77 **hps.model).cuda(rank)78 net_d = MultiPeriodDiscriminator(hps.model.use_spectral_norm).cuda(rank)79 optim_g = torch.optim.AdamW(80 net_g.parameters(),81 hps.train.learning_rate,82 betas=hps.train.betas,83 eps=hps.train.eps)84 optim_d = torch.optim.AdamW(85 net_d.parameters(),86 hps.train.learning_rate,87 betas=hps.train.betas,88 eps=hps.train.eps)89 net_g = DDP(net_g, device_ids=[rank]) # , find_unused_parameters=True)90 net_d = DDP(net_d, device_ids=[rank])91 92 skip_optimizer = False93 try:94 _, _, _, epoch_str = utils.load_checkpoint(utils.latest_checkpoint_path(hps.model_dir, "G_*.pth"), net_g,95 optim_g, skip_optimizer)96 _, _, _, epoch_str = utils.load_checkpoint(utils.latest_checkpoint_path(hps.model_dir, "D_*.pth"), net_d,97 optim_d, skip_optimizer)98 epoch_str = max(epoch_str, 1)99 name=utils.latest_checkpoint_path(hps.model_dir, "D_*.pth")100 global_step=int(name[name.rfind("_")+1:name.rfind(".")])+1101 #global_step = (epoch_str - 1) * len(train_loader)102 except Exception:103 print("load old checkpoint failed...")104 epoch_str = 1105 global_step = 0106 if skip_optimizer:107 epoch_str = 1108 global_step = 0109 110 warmup_epoch = hps.train.warmup_epochs111 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 # set up warm-up learning rate118 if epoch <= warmup_epoch:119 for param_group in optim_g.param_groups:120 param_group['lr'] = hps.train.learning_rate / warmup_epoch * epoch121 for param_group in optim_d.param_groups:122 param_group['lr'] = hps.train.learning_rate / warmup_epoch * epoch123 # training124 if rank == 0:125 train_and_evaluate(rank, epoch, hps, [net_g, net_d], [optim_g, optim_d], [scheduler_g, scheduler_d], scaler,126 [train_loader, eval_loader], logger, [writer, writer_eval])127 else:128 train_and_evaluate(rank, epoch, hps, [net_g, net_d], [optim_g, optim_d], [scheduler_g, scheduler_d], scaler,129 [train_loader, None], None, None)130 # update learning rate131 scheduler_g.step()132 scheduler_d.step()133 134 135def train_and_evaluate(rank, epoch, hps, nets, optims, schedulers, scaler, loaders, logger, writers):136 net_g, net_d = nets137 optim_g, optim_d = optims138 scheduler_g, scheduler_d = schedulers139 train_loader, eval_loader = loaders140 if writers is not None:141 writer, writer_eval = writers142 143 half_type = torch.bfloat16 if hps.train.half_type=="bf16" else torch.float16144 145 # train_loader.batch_sampler.set_epoch(epoch)146 global global_step147 148 net_g.train()149 net_d.train()150 for batch_idx, items in enumerate(train_loader):151 c, f0, spec, y, spk, lengths, uv,volume = items152 g = spk.cuda(rank, non_blocking=True)153 spec, y = spec.cuda(rank, non_blocking=True), y.cuda(rank, non_blocking=True)154 c = c.cuda(rank, non_blocking=True)155 f0 = f0.cuda(rank, non_blocking=True)156 uv = uv.cuda(rank, non_blocking=True)157 lengths = lengths.cuda(rank, non_blocking=True)158 mel = spec_to_mel_torch(159 spec,160 hps.data.filter_length,161 hps.data.n_mel_channels,162 hps.data.sampling_rate,163 hps.data.mel_fmin,164 hps.data.mel_fmax)165 166 with autocast(enabled=hps.train.fp16_run, dtype=half_type):167 y_hat, ids_slice, z_mask, \168 (z, z_p, m_p, logs_p, m_q, logs_q), pred_lf0, norm_lf0, lf0 = net_g(c, f0, uv, spec, g=g, c_lengths=lengths,169 spec_lengths=lengths,vol = volume)170 171 y_mel = commons.slice_segments(mel, ids_slice, hps.train.segment_size // hps.data.hop_length)172 y_hat_mel = mel_spectrogram_torch(173 y_hat.squeeze(1),174 hps.data.filter_length,175 hps.data.n_mel_channels,176 hps.data.sampling_rate,177 hps.data.hop_length,178 hps.data.win_length,179 hps.data.mel_fmin,180 hps.data.mel_fmax181 )182 y = commons.slice_segments(y, ids_slice * hps.data.hop_length, hps.train.segment_size) # slice183 184 # Discriminator185 y_d_hat_r, y_d_hat_g, _, _ = net_d(y, y_hat.detach())186 187 with autocast(enabled=False, dtype=half_type):188 loss_disc, losses_disc_r, losses_disc_g = discriminator_loss(y_d_hat_r, y_d_hat_g)189 loss_disc_all = loss_disc190 191 optim_d.zero_grad()192 scaler.scale(loss_disc_all).backward()193 scaler.unscale_(optim_d)194 grad_norm_d = commons.clip_grad_value_(net_d.parameters(), None)195 scaler.step(optim_d)196 197 198 with autocast(enabled=hps.train.fp16_run, dtype=half_type):199 # Generator200 y_d_hat_r, y_d_hat_g, fmap_r, fmap_g = net_d(y, y_hat)201 with autocast(enabled=False, dtype=half_type):202 loss_mel = F.l1_loss(y_mel, y_hat_mel) * hps.train.c_mel203 loss_kl = kl_loss(z_p, logs_q, m_p, logs_p, z_mask) * hps.train.c_kl204 loss_fm = feature_loss(fmap_r, fmap_g)205 loss_gen, losses_gen = generator_loss(y_d_hat_g)206 loss_lf0 = F.mse_loss(pred_lf0, lf0) if net_g.module.use_automatic_f0_prediction else 0207 loss_gen_all = loss_gen + loss_fm + loss_mel + loss_kl + loss_lf0208 optim_g.zero_grad()209 scaler.scale(loss_gen_all).backward()210 scaler.unscale_(optim_g)211 grad_norm_g = commons.clip_grad_value_(net_g.parameters(), None)212 scaler.step(optim_g)213 scaler.update()214 215 if rank == 0:216 if global_step % hps.train.log_interval == 0:217 lr = optim_g.param_groups[0]['lr']218 losses = [loss_disc, loss_gen, loss_fm, loss_mel, loss_kl]219 reference_loss=0220 for i in losses:221 reference_loss += i222 logger.info('Train Epoch: {} [{:.0f}%]'.format(223 epoch,224 100. * batch_idx / len(train_loader)))225 logger.info(f"Losses: {[x.item() for x in losses]}, step: {global_step}, lr: {lr}, reference_loss: {reference_loss}")226 227 scalar_dict = {"loss/g/total": loss_gen_all, "loss/d/total": loss_disc_all, "learning_rate": lr,228 "grad_norm_d": grad_norm_d, "grad_norm_g": grad_norm_g}229 scalar_dict.update({"loss/g/fm": loss_fm, "loss/g/mel": loss_mel, "loss/g/kl": loss_kl,230 "loss/g/lf0": loss_lf0})231 232 # scalar_dict.update({"loss/g/{}".format(i): v for i, v in enumerate(losses_gen)})233 # scalar_dict.update({"loss/d_r/{}".format(i): v for i, v in enumerate(losses_disc_r)})234 # scalar_dict.update({"loss/d_g/{}".format(i): v for i, v in enumerate(losses_disc_g)})235 image_dict = {236 "slice/mel_org": utils.plot_spectrogram_to_numpy(y_mel[0].data.cpu().numpy()),237 "slice/mel_gen": utils.plot_spectrogram_to_numpy(y_hat_mel[0].data.cpu().numpy()),238 "all/mel": utils.plot_spectrogram_to_numpy(mel[0].data.cpu().numpy())239 }240 241 if net_g.module.use_automatic_f0_prediction:242 image_dict.update({243 "all/lf0": utils.plot_data_to_numpy(lf0[0, 0, :].cpu().numpy(),244 pred_lf0[0, 0, :].detach().cpu().numpy()),245 "all/norm_lf0": utils.plot_data_to_numpy(lf0[0, 0, :].cpu().numpy(),246 norm_lf0[0, 0, :].detach().cpu().numpy())247 })248 249 utils.summarize(250 writer=writer,251 global_step=global_step,252 images=image_dict,253 scalars=scalar_dict254 )255 256 if global_step % hps.train.eval_interval == 0:257 evaluate(hps, net_g, eval_loader, writer_eval)258 utils.save_checkpoint(net_g, optim_g, hps.train.learning_rate, epoch,259 os.path.join(hps.model_dir, "G_{}.pth".format(global_step)))260 utils.save_checkpoint(net_d, optim_d, hps.train.learning_rate, epoch,261 os.path.join(hps.model_dir, "D_{}.pth".format(global_step)))262 keep_ckpts = getattr(hps.train, 'keep_ckpts', 0)263 if keep_ckpts > 0:264 utils.clean_checkpoints(path_to_models=hps.model_dir, n_ckpts_to_keep=keep_ckpts, sort_by_time=True)265 266 global_step += 1267 268 if rank == 0:269 global start_time270 now = time.time()271 durtaion = format(now - start_time, '.2f')272 logger.info(f'====> Epoch: {epoch}, cost {durtaion} s')273 start_time = now274 275 276def evaluate(hps, generator, eval_loader, writer_eval):277 generator.eval()278 image_dict = {}279 audio_dict = {}280 with torch.no_grad():281 for batch_idx, items in enumerate(eval_loader):282 c, f0, spec, y, spk, _, uv,volume = items283 g = spk[:1].cuda(0)284 spec, y = spec[:1].cuda(0), y[:1].cuda(0)285 c = c[:1].cuda(0)286 f0 = f0[:1].cuda(0)287 uv= uv[:1].cuda(0)288 if volume is not None:289 volume = volume[:1].cuda(0)290 mel = spec_to_mel_torch(291 spec,292 hps.data.filter_length,293 hps.data.n_mel_channels,294 hps.data.sampling_rate,295 hps.data.mel_fmin,296 hps.data.mel_fmax)297 y_hat,_ = generator.module.infer(c, f0, uv, g=g,vol = volume)298 299 y_hat_mel = mel_spectrogram_torch(300 y_hat.squeeze(1).float(),301 hps.data.filter_length,302 hps.data.n_mel_channels,303 hps.data.sampling_rate,304 hps.data.hop_length,305 hps.data.win_length,306 hps.data.mel_fmin,307 hps.data.mel_fmax308 )309 310 audio_dict.update({311 f"gen/audio_{batch_idx}": y_hat[0],312 f"gt/audio_{batch_idx}": y[0]313 })314 image_dict.update({315 "gen/mel": utils.plot_spectrogram_to_numpy(y_hat_mel[0].cpu().numpy()),316 "gt/mel": utils.plot_spectrogram_to_numpy(mel[0].cpu().numpy())317 })318 utils.summarize(319 writer=writer_eval,320 global_step=global_step,321 images=image_dict,322 audios=audio_dict,323 audio_sampling_rate=hps.data.sampling_rate324 )325 generator.train()326 327 328if __name__ == "__main__":329 main()