CoolFace
Apppublic

kkvc-hf/Style-Bert-VITS2-AS2

sourceHugging Faceapache-2.0updated 11mo agoView on Hugging Face
1likes
train_ms.py987 linesDownload Raw Back to root
1import argparse2import datetime3import gc4import os5import platform6 7import torch8import torch.distributed as dist9from huggingface_hub import HfApi10from torch.cuda.amp import GradScaler, autocast11from torch.nn import functional as F12from torch.nn.parallel import DistributedDataParallel as DDP13from torch.utils.data import DataLoader14from torch.utils.tensorboard import SummaryWriter15from tqdm import tqdm16 17# logging.getLogger("numba").setLevel(logging.WARNING)18import default_style19from config import get_config20from data_utils import (21    DistributedBucketSampler,22    TextAudioSpeakerCollate,23    TextAudioSpeakerLoader,24)25from losses import discriminator_loss, feature_loss, generator_loss, kl_loss26from mel_processing import mel_spectrogram_torch, spec_to_mel_torch27from style_bert_vits2.logging import logger28from style_bert_vits2.models import commons, utils29from style_bert_vits2.models.hyper_parameters import HyperParameters30from style_bert_vits2.models.models import (31    DurationDiscriminator,32    MultiPeriodDiscriminator,33    SynthesizerTrn,34)35from style_bert_vits2.nlp.symbols import SYMBOLS36from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT37 38 39torch.backends.cuda.matmul.allow_tf32 = True40torch.backends.cudnn.allow_tf32 = (41    True  # If encontered training problem,please try to disable TF32.42)43torch.set_float32_matmul_precision("medium")44torch.backends.cuda.sdp_kernel("flash")45torch.backends.cuda.enable_flash_sdp(True)46torch.backends.cuda.enable_mem_efficient_sdp(47    True48)  # Not available if torch version is lower than 2.049torch.backends.cuda.enable_math_sdp(True)50 51config = get_config()52global_step = 053 54api = HfApi()55 56 57def run():58    # Command line configuration is not recommended unless necessary, use config.yml59    parser = argparse.ArgumentParser()60    parser.add_argument(61        "-c",62        "--config",63        type=str,64        default=config.train_ms_config.config_path,65        help="JSON file for configuration",66    )67    parser.add_argument(68        "-m",69        "--model",70        type=str,71        help="数据集文件夹路径,请注意,数据不再默认放在/logs文件夹下。如果需要用命令行配置,请声明相对于根目录的路径",72        default=config.dataset_path,73    )74    parser.add_argument(75        "--assets_root",76        type=str,77        help="Root directory of model assets needed for inference.",78        default=config.assets_root,79    )80    parser.add_argument(81        "--skip_default_style",82        action="store_true",83        help="Skip saving default style config and mean vector.",84    )85    parser.add_argument(86        "--no_progress_bar",87        action="store_true",88        help="Do not show the progress bar while training.",89    )90    parser.add_argument(91        "--speedup",92        action="store_true",93        help="Speed up training by disabling logging and evaluation.",94    )95    parser.add_argument(96        "--repo_id",97        help="Huggingface model repo id to backup the model.",98        default=None,99    )100    parser.add_argument(101        "--not_use_custom_batch_sampler",102        help="Don't use custom batch sampler for training, which was used in the version < 2.5",103        action="store_true",104    )105    args = parser.parse_args()106 107    # Set log file108    model_dir = os.path.join(args.model, config.train_ms_config.model_dir)109    timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")110    logger.add(os.path.join(args.model, f"train_{timestamp}.log"))111 112    # Parsing environment variables113    envs = config.train_ms_config.env114    for env_name, env_value in envs.items():115        if env_name not in os.environ.keys():116            logger.info(f"Loading configuration from config {env_value!s}")117            os.environ[env_name] = str(env_value)118    logger.info(119        "Loading environment variables \nMASTER_ADDR: {},\nMASTER_PORT: {},\nWORLD_SIZE: {},\nRANK: {},\nLOCAL_RANK: {}".format(120            os.environ["MASTER_ADDR"],121            os.environ["MASTER_PORT"],122            os.environ["WORLD_SIZE"],123            os.environ["RANK"],124            os.environ["LOCAL_RANK"],125        )126    )127 128    backend = "nccl"129    if platform.system() == "Windows":130        backend = "gloo"  # If Windows,switch to gloo backend.131    dist.init_process_group(132        backend=backend,133        init_method="env://",134        timeout=datetime.timedelta(seconds=300),135    )  # Use torchrun instead of mp.spawn136    rank = dist.get_rank()137    local_rank = int(os.environ["LOCAL_RANK"])138    n_gpus = dist.get_world_size()139 140    hps = HyperParameters.load_from_json(args.config)141    # This is needed because we have to pass values to `train_and_evaluate()`142    hps.model_dir = model_dir143    hps.speedup = args.speedup144    hps.repo_id = args.repo_id145 146    # 比较路径是否相同147    if os.path.realpath(args.config) != os.path.realpath(148        config.train_ms_config.config_path149    ):150        with open(args.config, encoding="utf-8") as f:151            data = f.read()152        os.makedirs(os.path.dirname(config.train_ms_config.config_path), exist_ok=True)153        with open(config.train_ms_config.config_path, "w", encoding="utf-8") as f:154            f.write(data)155 156    """157    Path constants are a bit complicated...158    TODO: Refactor or rename these?159    (Both `config.yml` and `config.json` are used, which is confusing I think.)160 161    args.model: For saving all info needed for training.162        default: `Data/{model_name}`.163    hps.model_dir := model_dir: For saving checkpoints (for resuming training).164        default: `Data/{model_name}/models`.165        (Use `hps` since we have to pass `model_dir` to `train_and_evaluate()`.166 167    args.assets_root: The root directory of model assets needed for inference.168        default: config.assets_root == `model_assets`.169 170    config.out_dir: The directory for model assets of this model (for inference).171        default: `model_assets/{model_name}`.172    """173 174    if args.repo_id is not None:175        # First try to upload config.json to check if the repo exists176        try:177            api.upload_file(178                path_or_fileobj=args.config,179                path_in_repo=f"Data/{config.model_name}/config.json",180                repo_id=hps.repo_id,181            )182        except Exception as e:183            logger.error(e)184            logger.error(185                f"Failed to upload files to the repo {hps.repo_id}. Please check if the repo exists and you have logged in using `huggingface-cli login`."186            )187            raise e188        # Upload Data dir for resuming training189        api.upload_folder(190            repo_id=hps.repo_id,191            folder_path=config.dataset_path,192            path_in_repo=f"Data/{config.model_name}",193            delete_patterns="*.pth",  # Only keep the latest checkpoint194            run_as_future=True,195        )196 197    os.makedirs(config.out_dir, exist_ok=True)198 199    if not args.skip_default_style:200        default_style.save_styles_by_dirs(201            os.path.join(args.model, "wavs"),202            config.out_dir,203            config_path=args.config,204            config_output_path=os.path.join(config.out_dir, "config.json"),205        )206 207    torch.manual_seed(hps.train.seed)208    torch.cuda.set_device(local_rank)209 210    global global_step211    writer = None212    writer_eval = None213    if rank == 0 and not args.speedup:214        # logger = utils.get_logger(hps.model_dir)215        # logger.info(hps)216        utils.check_git_hash(model_dir)217        writer = SummaryWriter(log_dir=model_dir)218        writer_eval = SummaryWriter(log_dir=os.path.join(model_dir, "eval"))219    train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data)220    collate_fn = TextAudioSpeakerCollate()221    if not args.not_use_custom_batch_sampler:222        train_sampler = DistributedBucketSampler(223            train_dataset,224            hps.train.batch_size,225            [32, 300, 400, 500, 600, 700, 800, 900, 1000],226            num_replicas=n_gpus,227            rank=rank,228            shuffle=True,229        )230        train_loader = DataLoader(231            train_dataset,232            # メモリ消費量を減らそうとnum_workersを1にしてみる233            # num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2),234            num_workers=1,235            shuffle=False,236            pin_memory=True,237            collate_fn=collate_fn,238            batch_sampler=train_sampler,239            # batch_size=hps.train.batch_size,240            persistent_workers=True,241            # これもメモリ消費量を減らそうとしてコメントアウト242            # prefetch_factor=6,243        )244    else:245        train_loader = DataLoader(246            train_dataset,247            # メモリ消費量を減らそうとnum_workersを1にしてみる248            # num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2),249            num_workers=1,250            shuffle=True,251            pin_memory=True,252            collate_fn=collate_fn,253            # batch_sampler=train_sampler,254            batch_size=hps.train.batch_size,255            persistent_workers=True,256            # これもメモリ消費量を減らそうとしてコメントアウト257            # prefetch_factor=6,258        )259    eval_dataset = None260    eval_loader = None261    if rank == 0 and not args.speedup:262        eval_dataset = TextAudioSpeakerLoader(hps.data.validation_files, hps.data)263        eval_loader = DataLoader(264            eval_dataset,265            num_workers=0,266            shuffle=False,267            batch_size=1,268            pin_memory=True,269            drop_last=False,270            collate_fn=collate_fn,271        )272    if hps.model.use_noise_scaled_mas is True:273        logger.info("Using noise scaled MAS for VITS2")274        mas_noise_scale_initial = 0.01275        noise_scale_delta = 2e-6276    else:277        logger.info("Using normal MAS for VITS1")278        mas_noise_scale_initial = 0.0279        noise_scale_delta = 0.0280    if hps.model.use_duration_discriminator is True:281        logger.info("Using duration discriminator for VITS2")282        net_dur_disc = DurationDiscriminator(283            hps.model.hidden_channels,284            hps.model.hidden_channels,285            3,286            0.1,287            gin_channels=hps.model.gin_channels if hps.data.n_speakers != 0 else 0,288        ).cuda(local_rank)289    if hps.model.use_spk_conditioned_encoder is True:290        if hps.data.n_speakers == 0:291            raise ValueError(292                "n_speakers must be > 0 when using spk conditioned encoder to train multi-speaker model"293            )294    else:295        logger.info("Using normal encoder for VITS1")296 297    net_g = SynthesizerTrn(298        len(SYMBOLS),299        hps.data.filter_length // 2 + 1,300        hps.train.segment_size // hps.data.hop_length,301        n_speakers=hps.data.n_speakers,302        mas_noise_scale_initial=mas_noise_scale_initial,303        noise_scale_delta=noise_scale_delta,304        # hps.model 以下のすべての値を引数に渡す305        use_spk_conditioned_encoder=hps.model.use_spk_conditioned_encoder,306        use_noise_scaled_mas=hps.model.use_noise_scaled_mas,307        use_mel_posterior_encoder=hps.model.use_mel_posterior_encoder,308        use_duration_discriminator=hps.model.use_duration_discriminator,309        use_wavlm_discriminator=hps.model.use_wavlm_discriminator,310        inter_channels=hps.model.inter_channels,311        hidden_channels=hps.model.hidden_channels,312        filter_channels=hps.model.filter_channels,313        n_heads=hps.model.n_heads,314        n_layers=hps.model.n_layers,315        kernel_size=hps.model.kernel_size,316        p_dropout=hps.model.p_dropout,317        resblock=hps.model.resblock,318        resblock_kernel_sizes=hps.model.resblock_kernel_sizes,319        resblock_dilation_sizes=hps.model.resblock_dilation_sizes,320        upsample_rates=hps.model.upsample_rates,321        upsample_initial_channel=hps.model.upsample_initial_channel,322        upsample_kernel_sizes=hps.model.upsample_kernel_sizes,323        n_layers_q=hps.model.n_layers_q,324        use_spectral_norm=hps.model.use_spectral_norm,325        gin_channels=hps.model.gin_channels,326        slm=hps.model.slm,327    ).cuda(local_rank)328 329    if getattr(hps.train, "freeze_ZH_bert", False):330        logger.info("Freezing ZH bert encoder !!!")331        for param in net_g.enc_p.bert_proj.parameters():332            param.requires_grad = False333 334    if getattr(hps.train, "freeze_EN_bert", False):335        logger.info("Freezing EN bert encoder !!!")336        for param in net_g.enc_p.en_bert_proj.parameters():337            param.requires_grad = False338 339    if getattr(hps.train, "freeze_JP_bert", False):340        logger.info("Freezing JP bert encoder !!!")341        for param in net_g.enc_p.ja_bert_proj.parameters():342            param.requires_grad = False343    if getattr(hps.train, "freeze_style", False):344        logger.info("Freezing style encoder !!!")345        for param in net_g.enc_p.style_proj.parameters():346            param.requires_grad = False347 348    if getattr(hps.train, "freeze_decoder", False):349        logger.info("Freezing decoder !!!")350        for param in net_g.dec.parameters():351            param.requires_grad = False352 353    net_d = MultiPeriodDiscriminator(hps.model.use_spectral_norm).cuda(local_rank)354    optim_g = torch.optim.AdamW(355        filter(lambda p: p.requires_grad, net_g.parameters()),356        hps.train.learning_rate,357        betas=hps.train.betas,358        eps=hps.train.eps,359    )360    optim_d = torch.optim.AdamW(361        net_d.parameters(),362        hps.train.learning_rate,363        betas=hps.train.betas,364        eps=hps.train.eps,365    )366    if net_dur_disc is not None:367        optim_dur_disc = torch.optim.AdamW(368            net_dur_disc.parameters(),369            hps.train.learning_rate,370            betas=hps.train.betas,371            eps=hps.train.eps,372        )373    else:374        optim_dur_disc = None375    net_g = DDP(net_g, device_ids=[local_rank])376    net_d = DDP(net_d, device_ids=[local_rank])377    dur_resume_lr = None378    if net_dur_disc is not None:379        net_dur_disc = DDP(380            net_dur_disc, device_ids=[local_rank], find_unused_parameters=True381        )382 383    if utils.is_resuming(model_dir):384        if net_dur_disc is not None:385            _, _, dur_resume_lr, epoch_str = utils.checkpoints.load_checkpoint(386                utils.checkpoints.get_latest_checkpoint_path(model_dir, "DUR_*.pth"),387                net_dur_disc,388                optim_dur_disc,389                skip_optimizer=hps.train.skip_optimizer,390            )391            if not optim_dur_disc.param_groups[0].get("initial_lr"):392                optim_dur_disc.param_groups[0]["initial_lr"] = dur_resume_lr393        _, optim_g, g_resume_lr, epoch_str = utils.checkpoints.load_checkpoint(394            utils.checkpoints.get_latest_checkpoint_path(model_dir, "G_*.pth"),395            net_g,396            optim_g,397            skip_optimizer=hps.train.skip_optimizer,398        )399        _, optim_d, d_resume_lr, epoch_str = utils.checkpoints.load_checkpoint(400            utils.checkpoints.get_latest_checkpoint_path(model_dir, "D_*.pth"),401            net_d,402            optim_d,403            skip_optimizer=hps.train.skip_optimizer,404        )405        if not optim_g.param_groups[0].get("initial_lr"):406            optim_g.param_groups[0]["initial_lr"] = g_resume_lr407        if not optim_d.param_groups[0].get("initial_lr"):408            optim_d.param_groups[0]["initial_lr"] = d_resume_lr409 410        epoch_str = max(epoch_str, 1)411        # global_step = (epoch_str - 1) * len(train_loader)412        global_step = int(413            utils.get_steps(414                utils.checkpoints.get_latest_checkpoint_path(model_dir, "G_*.pth")415            )416        )417        logger.info(418            f"******************Found the model. Current epoch is {epoch_str}, gloabl step is {global_step}*********************"419        )420    else:421        try:422            _ = utils.safetensors.load_safetensors(423                os.path.join(model_dir, "G_0.safetensors"), net_g424            )425            _ = utils.safetensors.load_safetensors(426                os.path.join(model_dir, "D_0.safetensors"), net_d427            )428            if net_dur_disc is not None:429                _ = utils.safetensors.load_safetensors(430                    os.path.join(model_dir, "DUR_0.safetensors"), net_dur_disc431                )432            logger.info("Loaded the pretrained models.")433        except Exception as e:434            logger.warning(e)435            logger.warning(436                "It seems that you are not using the pretrained models, so we will train from scratch."437            )438        finally:439            epoch_str = 1440            global_step = 0441 442    def lr_lambda(epoch):443        """444        Learning rate scheduler for warmup and exponential decay.445        - During the warmup period, the learning rate increases linearly.446        - After the warmup period, the learning rate decreases exponentially.447        """448        if epoch < hps.train.warmup_epochs:449            return float(epoch) / float(max(1, hps.train.warmup_epochs))450        else:451            return hps.train.lr_decay ** (epoch - hps.train.warmup_epochs)452 453    scheduler_last_epoch = epoch_str - 2454    scheduler_g = torch.optim.lr_scheduler.LambdaLR(455        optim_g, lr_lambda=lr_lambda, last_epoch=scheduler_last_epoch456    )457    scheduler_d = torch.optim.lr_scheduler.LambdaLR(458        optim_d, lr_lambda=lr_lambda, last_epoch=scheduler_last_epoch459    )460    if net_dur_disc is not None:461        scheduler_dur_disc = torch.optim.lr_scheduler.LambdaLR(462            optim_dur_disc, lr_lambda=lr_lambda, last_epoch=scheduler_last_epoch463        )464    else:465        scheduler_dur_disc = None466    scaler = GradScaler(enabled=hps.train.bf16_run)467    logger.info("Start training.")468 469    diff = abs(470        epoch_str * len(train_loader) - (hps.train.epochs + 1) * len(train_loader)471    )472    pbar = None473    if not args.no_progress_bar:474        pbar = tqdm(475            total=global_step + diff,476            initial=global_step,477            smoothing=0.05,478            file=SAFE_STDOUT,479        )480    initial_step = global_step481 482    for epoch in range(epoch_str, hps.train.epochs + 1):483        if rank == 0:484            train_and_evaluate(485                rank,486                local_rank,487                epoch,488                hps,489                [net_g, net_d, net_dur_disc],490                [optim_g, optim_d, optim_dur_disc],491                [scheduler_g, scheduler_d, scheduler_dur_disc],492                scaler,493                [train_loader, eval_loader],494                logger,495                [writer, writer_eval],496                pbar,497                initial_step,498            )499        else:500            train_and_evaluate(501                rank,502                local_rank,503                epoch,504                hps,505                [net_g, net_d, net_dur_disc],506                [optim_g, optim_d, optim_dur_disc],507                [scheduler_g, scheduler_d, scheduler_dur_disc],508                scaler,509                [train_loader, None],510                None,511                None,512                pbar,513                initial_step,514            )515        scheduler_g.step()516        scheduler_d.step()517        if net_dur_disc is not None:518            scheduler_dur_disc.step()519 520        if epoch == hps.train.epochs:521            # Save the final models522            assert optim_g is not None523            utils.checkpoints.save_checkpoint(524                net_g,525                optim_g,526                hps.train.learning_rate,527                epoch,528                os.path.join(model_dir, f"G_{global_step}.pth"),529            )530            assert optim_d is not None531            utils.checkpoints.save_checkpoint(532                net_d,533                optim_d,534                hps.train.learning_rate,535                epoch,536                os.path.join(model_dir, f"D_{global_step}.pth"),537            )538            if net_dur_disc is not None:539                assert optim_dur_disc is not None540                utils.checkpoints.save_checkpoint(541                    net_dur_disc,542                    optim_dur_disc,543                    hps.train.learning_rate,544                    epoch,545                    os.path.join(model_dir, f"DUR_{global_step}.pth"),546                )547            utils.safetensors.save_safetensors(548                net_g,549                epoch,550                os.path.join(551                    config.out_dir,552                    f"{config.model_name}_e{epoch}_s{global_step}.safetensors",553                ),554                for_infer=True,555            )556            if hps.repo_id is not None:557                future1 = api.upload_folder(558                    repo_id=hps.repo_id,559                    folder_path=config.dataset_path,560                    path_in_repo=f"Data/{config.model_name}",561                    delete_patterns="*.pth",  # Only keep the latest checkpoint562                    run_as_future=True,563                )564                future2 = api.upload_folder(565                    repo_id=hps.repo_id,566                    folder_path=config.out_dir,567                    path_in_repo=f"model_assets/{config.model_name}",568                    run_as_future=True,569                )570                try:571                    future1.result()572                    future2.result()573                except Exception as e:574                    logger.error(e)575 576    if pbar is not None:577        pbar.close()578 579 580def train_and_evaluate(581    rank,582    local_rank,583    epoch,584    hps: HyperParameters,585    nets,586    optims,587    schedulers,588    scaler,589    loaders,590    logger,591    writers,592    pbar: tqdm,593    initial_step: int,594):595    net_g, net_d, net_dur_disc = nets596    optim_g, optim_d, optim_dur_disc = optims597    scheduler_g, scheduler_d, scheduler_dur_disc = schedulers598    train_loader, eval_loader = loaders599    if writers is not None:600        writer, writer_eval = writers601 602    train_loader.batch_sampler.set_epoch(epoch)603    global global_step604 605    net_g.train()606    net_d.train()607    if net_dur_disc is not None:608        net_dur_disc.train()609    for batch_idx, (610        x,611        x_lengths,612        spec,613        spec_lengths,614        y,615        y_lengths,616        speakers,617        tone,618        language,619        bert,620        ja_bert,621        en_bert,622        style_vec,623    ) in enumerate(train_loader):624        if net_g.module.use_noise_scaled_mas:625            current_mas_noise_scale = (626                net_g.module.mas_noise_scale_initial627                - net_g.module.noise_scale_delta * global_step628            )629            net_g.module.current_mas_noise_scale = max(current_mas_noise_scale, 0.0)630        x, x_lengths = x.cuda(local_rank, non_blocking=True), x_lengths.cuda(631            local_rank, non_blocking=True632        )633        spec, spec_lengths = spec.cuda(634            local_rank, non_blocking=True635        ), spec_lengths.cuda(local_rank, non_blocking=True)636        y, y_lengths = y.cuda(local_rank, non_blocking=True), y_lengths.cuda(637            local_rank, non_blocking=True638        )639        speakers = speakers.cuda(local_rank, non_blocking=True)640        tone = tone.cuda(local_rank, non_blocking=True)641        language = language.cuda(local_rank, non_blocking=True)642        bert = bert.cuda(local_rank, non_blocking=True)643        ja_bert = ja_bert.cuda(local_rank, non_blocking=True)644        en_bert = en_bert.cuda(local_rank, non_blocking=True)645        style_vec = style_vec.cuda(local_rank, non_blocking=True)646 647        with autocast(enabled=hps.train.bf16_run, dtype=torch.bfloat16):648            (649                y_hat,650                l_length,651                attn,652                ids_slice,653                x_mask,654                z_mask,655                (z, z_p, m_p, logs_p, m_q, logs_q),656                (hidden_x, logw, logw_),657            ) = net_g(658                x,659                x_lengths,660                spec,661                spec_lengths,662                speakers,663                tone,664                language,665                bert,666                ja_bert,667                en_bert,668                style_vec,669            )670            mel = spec_to_mel_torch(671                spec,672                hps.data.filter_length,673                hps.data.n_mel_channels,674                hps.data.sampling_rate,675                hps.data.mel_fmin,676                hps.data.mel_fmax,677            )678            y_mel = commons.slice_segments(679                mel, ids_slice, hps.train.segment_size // hps.data.hop_length680            )681            y_hat_mel = mel_spectrogram_torch(682                y_hat.squeeze(1).float(),683                hps.data.filter_length,684                hps.data.n_mel_channels,685                hps.data.sampling_rate,686                hps.data.hop_length,687                hps.data.win_length,688                hps.data.mel_fmin,689                hps.data.mel_fmax,690            )691 692            y = commons.slice_segments(693                y, ids_slice * hps.data.hop_length, hps.train.segment_size694            )  # slice695 696            # Discriminator697            y_d_hat_r, y_d_hat_g, _, _ = net_d(y, y_hat.detach())698            with autocast(enabled=hps.train.bf16_run, dtype=torch.bfloat16):699                loss_disc, losses_disc_r, losses_disc_g = discriminator_loss(700                    y_d_hat_r, y_d_hat_g701                )702                loss_disc_all = loss_disc703            if net_dur_disc is not None:704                y_dur_hat_r, y_dur_hat_g = net_dur_disc(705                    hidden_x.detach(), x_mask.detach(), logw.detach(), logw_.detach()706                )707                with autocast(enabled=hps.train.bf16_run, dtype=torch.bfloat16):708                    # TODO: I think need to mean using the mask, but for now, just mean all709                    (710                        loss_dur_disc,711                        losses_dur_disc_r,712                        losses_dur_disc_g,713                    ) = discriminator_loss(y_dur_hat_r, y_dur_hat_g)714                    loss_dur_disc_all = loss_dur_disc715                optim_dur_disc.zero_grad()716                scaler.scale(loss_dur_disc_all).backward()717                scaler.unscale_(optim_dur_disc)718                commons.clip_grad_value_(net_dur_disc.parameters(), None)719                scaler.step(optim_dur_disc)720 721        optim_d.zero_grad()722        scaler.scale(loss_disc_all).backward()723        scaler.unscale_(optim_d)724        if getattr(hps.train, "bf16_run", False):725            torch.nn.utils.clip_grad_norm_(parameters=net_d.parameters(), max_norm=200)726        grad_norm_d = commons.clip_grad_value_(net_d.parameters(), None)727        scaler.step(optim_d)728 729        with autocast(enabled=hps.train.bf16_run, dtype=torch.bfloat16):730            # Generator731            y_d_hat_r, y_d_hat_g, fmap_r, fmap_g = net_d(y, y_hat)732            if net_dur_disc is not None:733                y_dur_hat_r, y_dur_hat_g = net_dur_disc(hidden_x, x_mask, logw, logw_)734            with autocast(enabled=hps.train.bf16_run, dtype=torch.bfloat16):735                loss_dur = torch.sum(l_length.float())736                loss_mel = F.l1_loss(y_mel, y_hat_mel) * hps.train.c_mel737                loss_kl = kl_loss(z_p, logs_q, m_p, logs_p, z_mask) * hps.train.c_kl738 739                loss_fm = feature_loss(fmap_r, fmap_g)740                loss_gen, losses_gen = generator_loss(y_d_hat_g)741                loss_gen_all = loss_gen + loss_fm + loss_mel + loss_dur + loss_kl742                if net_dur_disc is not None:743                    loss_dur_gen, losses_dur_gen = generator_loss(y_dur_hat_g)744                    loss_gen_all += loss_dur_gen745        optim_g.zero_grad()746        scaler.scale(loss_gen_all).backward()747        scaler.unscale_(optim_g)748        if getattr(hps.train, "bf16_run", False):749            torch.nn.utils.clip_grad_norm_(parameters=net_g.parameters(), max_norm=500)750        grad_norm_g = commons.clip_grad_value_(net_g.parameters(), None)751        scaler.step(optim_g)752        scaler.update()753 754        if rank == 0:755            if global_step % hps.train.log_interval == 0 and not hps.speedup:756                lr = optim_g.param_groups[0]["lr"]757                losses = [loss_disc, loss_gen, loss_fm, loss_mel, loss_dur, loss_kl]758                # logger.info(759                #     "Train Epoch: {} [{:.0f}%]".format(760                #         epoch, 100.0 * batch_idx / len(train_loader)761                #     )762                # )763                # logger.info([x.item() for x in losses] + [global_step, lr])764 765                scalar_dict = {766                    "loss/g/total": loss_gen_all,767                    "loss/d/total": loss_disc_all,768                    "learning_rate": lr,769                    "grad_norm_d": grad_norm_d,770                    "grad_norm_g": grad_norm_g,771                }772                scalar_dict.update(773                    {774                        "loss/g/fm": loss_fm,775                        "loss/g/mel": loss_mel,776                        "loss/g/dur": loss_dur,777                        "loss/g/kl": loss_kl,778                    }779                )780                scalar_dict.update({f"loss/g/{i}": v for i, v in enumerate(losses_gen)})781                scalar_dict.update(782                    {f"loss/d_r/{i}": v for i, v in enumerate(losses_disc_r)}783                )784                scalar_dict.update(785                    {f"loss/d_g/{i}": v for i, v in enumerate(losses_disc_g)}786                )787                # 以降のログは計算が重い気がするし誰も見てない気がするのでコメントアウト788                # image_dict = {789                #     "slice/mel_org": utils.plot_spectrogram_to_numpy(790                #         y_mel[0].data.cpu().numpy()791                #     ),792                #     "slice/mel_gen": utils.plot_spectrogram_to_numpy(793                #         y_hat_mel[0].data.cpu().numpy()794                #     ),795                #     "all/mel": utils.plot_spectrogram_to_numpy(796                #         mel[0].data.cpu().numpy()797                #     ),798                #     "all/attn": utils.plot_alignment_to_numpy(799                #         attn[0, 0].data.cpu().numpy()800                #     ),801                # }802                utils.summarize(803                    writer=writer,804                    global_step=global_step,805                    # images=image_dict,806                    scalars=scalar_dict,807                )808 809            if (810                global_step % hps.train.eval_interval == 0811                and global_step != 0812                and initial_step != global_step813            ):814                if not hps.speedup:815                    evaluate(hps, net_g, eval_loader, writer_eval)816                assert hps.model_dir is not None817                utils.checkpoints.save_checkpoint(818                    net_g,819                    optim_g,820                    hps.train.learning_rate,821                    epoch,822                    os.path.join(hps.model_dir, f"G_{global_step}.pth"),823                )824                utils.checkpoints.save_checkpoint(825                    net_d,826                    optim_d,827                    hps.train.learning_rate,828                    epoch,829                    os.path.join(hps.model_dir, f"D_{global_step}.pth"),830                )831                if net_dur_disc is not None:832                    utils.checkpoints.save_checkpoint(833                        net_dur_disc,834                        optim_dur_disc,835                        hps.train.learning_rate,836                        epoch,837                        os.path.join(hps.model_dir, f"DUR_{global_step}.pth"),838                    )839                keep_ckpts = config.train_ms_config.keep_ckpts840                if keep_ckpts > 0:841                    utils.checkpoints.clean_checkpoints(842                        model_dir_path=hps.model_dir,843                        n_ckpts_to_keep=keep_ckpts,844                        sort_by_time=True,845                    )846                # Save safetensors (for inference) to `model_assets/{model_name}`847                utils.safetensors.save_safetensors(848                    net_g,849                    epoch,850                    os.path.join(851                        config.out_dir,852                        f"{config.model_name}_e{epoch}_s{global_step}.safetensors",853                    ),854                    for_infer=True,855                )856                if hps.repo_id is not None:857                    api.upload_folder(858                        repo_id=hps.repo_id,859                        folder_path=config.dataset_path,860                        path_in_repo=f"Data/{config.model_name}",861                        delete_patterns="*.pth",  # Only keep the latest checkpoint862                        run_as_future=True,863                    )864                    api.upload_folder(865                        repo_id=hps.repo_id,866                        folder_path=config.out_dir,867                        path_in_repo=f"model_assets/{config.model_name}",868                        run_as_future=True,869                    )870 871        global_step += 1872        if pbar is not None:873            pbar.set_description(874                f"Epoch {epoch}({100.0 * batch_idx / len(train_loader):.0f}%)/{hps.train.epochs}"875            )876            pbar.update()877    # 本家ではこれをスピードアップのために消すと書かれていたので、一応消してみる878    # と思ったけどメモリ使用量が減るかもしれないのでつけてみる879    gc.collect()880    torch.cuda.empty_cache()881    if pbar is None and rank == 0:882        logger.info(f"====> Epoch: {epoch}, step: {global_step}")883 884 885def evaluate(hps, generator, eval_loader, writer_eval):886    generator.eval()887    image_dict = {}888    audio_dict = {}889    print()890    logger.info("Evaluating ...")891    with torch.no_grad():892        for batch_idx, (893            x,894            x_lengths,895            spec,896            spec_lengths,897            y,898            y_lengths,899            speakers,900            tone,901            language,902            bert,903            ja_bert,904            en_bert,905            style_vec,906        ) in enumerate(eval_loader):907            x, x_lengths = x.cuda(), x_lengths.cuda()908            spec, spec_lengths = spec.cuda(), spec_lengths.cuda()909            y, y_lengths = y.cuda(), y_lengths.cuda()910            speakers = speakers.cuda()911            bert = bert.cuda()912            ja_bert = ja_bert.cuda()913            en_bert = en_bert.cuda()914            tone = tone.cuda()915            language = language.cuda()916            style_vec = style_vec.cuda()917            for use_sdp in [True, False]:918                y_hat, attn, mask, *_ = generator.module.infer(919                    x,920                    x_lengths,921                    speakers,922                    tone,923                    language,924                    bert,925                    ja_bert,926                    en_bert,927                    style_vec,928                    y=spec,929                    max_len=1000,930                    sdp_ratio=0.0 if not use_sdp else 1.0,931                )932                y_hat_lengths = mask.sum([1, 2]).long() * hps.data.hop_length933                # 以降のログは計算が重い気がするし誰も見てない気がするのでコメントアウト934                # mel = spec_to_mel_torch(935                #     spec,936                #     hps.data.filter_length,937                #     hps.data.n_mel_channels,938                #     hps.data.sampling_rate,939                #     hps.data.mel_fmin,940                #     hps.data.mel_fmax,941                # )942                # y_hat_mel = mel_spectrogram_torch(943                #     y_hat.squeeze(1).float(),944                #     hps.data.filter_length,945                #     hps.data.n_mel_channels,946                #     hps.data.sampling_rate,947                #     hps.data.hop_length,948                #     hps.data.win_length,949                #     hps.data.mel_fmin,950                #     hps.data.mel_fmax,951                # )952                # image_dict.update(953                #     {954                #         f"gen/mel_{batch_idx}": utils.plot_spectrogram_to_numpy(955                #             y_hat_mel[0].cpu().numpy()956                #         )957                #     }958                # )959                # image_dict.update(960                #     {961                #         f"gt/mel_{batch_idx}": utils.plot_spectrogram_to_numpy(962                #             mel[0].cpu().numpy()963                #         )964                #     }965                # )966                audio_dict.update(967                    {968                        f"gen/audio_{batch_idx}_{use_sdp}": y_hat[969                            0, :, : y_hat_lengths[0]970                        ]971                    }972                )973                audio_dict.update({f"gt/audio_{batch_idx}": y[0, :, : y_lengths[0]]})974 975    utils.summarize(976        writer=writer_eval,977        global_step=global_step,978        images=image_dict,979        audios=audio_dict,980        audio_sampling_rate=hps.data.sampling_rate,981    )982    generator.train()983 984 985if __name__ == "__main__":986    run()987