kkvc-hf/Style-Bert-VITS2-AS2
1
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 