CoolFace
Apppublic

kwau/sovits-isla

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
train.py329 linesDownload Raw Back to root
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()