"""Train the separated text-to-mel model, then fine-tune it with a vocoder. Examples: python train_ms_mel.py -c Data/model/config.json --stage acoustic python train_ms_mel.py -c Data/model/config.json --stage vocoder python train_ms_mel.py -c Data/model/config.json --stage joint \ --acoustic-checkpoint Data/model/models/ACOUSTIC_10000.pth \ --vocoder-checkpoint Data/model/models/VOCODER_10000.pth """ import argparse import os from pathlib import Path from typing import Any import torch from torch.nn import functional as F from torch.utils.data import DataLoader from data_utils import TextAudioSpeakerCollate, TextAudioSpeakerLoader from losses import discriminator_loss, feature_loss, generator_loss from mel_processing import mel_spectrogram_torch, spec_to_mel_torch from style_bert_vits2.models import commons from style_bert_vits2.models.hyper_parameters import HyperParameters from style_bert_vits2.models.mel_synthesizer import JointMelSynthesizer from style_bert_vits2.models.models import MultiPeriodDiscriminator from style_bert_vits2.nlp.symbols import SYMBOLS def _checkpoint_state(path: str) -> dict[str, torch.Tensor]: if path.endswith(".safetensors"): from safetensors.torch import load_file return load_file(path) checkpoint = torch.load(path, map_location="cpu", weights_only=True) return checkpoint.get("model", checkpoint) def load_component( module: torch.nn.Module, path: str, prefixes: tuple[str, ...] = () ) -> None: """Load either a standalone or a prefixed joint/legacy component.""" saved = _checkpoint_state(path) target = module.state_dict() selected: dict[str, torch.Tensor] = {} for key, value in saved.items(): candidates = [key] if key.startswith("dec."): candidates.append("generator." + key[len("dec.") :]) if key.startswith("module.dec."): candidates.append("generator." + key[len("module.dec.") :]) for prefix in prefixes: if key.startswith(prefix): candidates.append(key[len(prefix) :]) for candidate in candidates: if candidate in target and target[candidate].shape == value.shape: selected[candidate] = value break if not selected: raise ValueError(f"No compatible parameters found in {path}") result = module.load_state_dict(selected, strict=False) print( f"Loaded {len(selected)}/{len(target)} tensors from {path} " f"({len(result.missing_keys)} parameters kept at initialization)" ) def save_training_checkpoint( path: Path, model: torch.nn.Module, optimizer: torch.optim.Optimizer, epoch: int, step: int, ) -> None: path.parent.mkdir(parents=True, exist_ok=True) torch.save( { "model": model.state_dict(), "optimizer": optimizer.state_dict(), "iteration": epoch, "global_step": step, "learning_rate": optimizer.param_groups[0]["lr"], }, path, ) def build_models(hps: HyperParameters) -> tuple[torch.nn.Module, torch.nn.Module]: module = ( __import__( "style_bert_vits2.models.models_jp_extra", fromlist=["build_mel_synthesizer"], ) if hps.data.use_jp_extra else __import__( "style_bert_vits2.models.models", fromlist=["build_mel_synthesizer"], ) ) acoustic = module.build_mel_synthesizer( len(SYMBOLS), hps.data.n_mel_channels, hps.model.hidden_channels, hps.model.filter_channels, hps.model.n_heads, hps.model.n_layers, hps.model.kernel_size, hps.model.p_dropout, hps.data.n_speakers, hps.model.gin_channels, hps.model.matcha_channels, hps.model.matcha_num_heads, hps.model.matcha_dropout, hps.model.matcha_sigma_min, hps.model.matcha_n_timesteps, hps.model.matcha_use_diff_attention, ) vocoder = module.build_mel_vocoder( hps.data.n_mel_channels, hps.model.resblock, hps.model.resblock_kernel_sizes, hps.model.resblock_dilation_sizes, hps.model.upsample_rates, hps.model.upsample_initial_channel, hps.model.upsample_kernel_sizes, hps.data.n_speakers, hps.model.gin_channels, ) return acoustic, vocoder def unpack_batch( batch: tuple[torch.Tensor, ...], hps: HyperParameters, device: torch.device, ) -> tuple[dict[str, Any], torch.Tensor]: ( x, x_lengths, spec, spec_lengths, waveform, _, speakers, tone, language, bert, ja_bert, en_bert, style_vec, ) = (item.to(device, non_blocking=True) for item in batch) mel = ( spec if hps.model.use_mel_posterior_encoder else spec_to_mel_torch( spec, hps.data.filter_length, hps.data.n_mel_channels, hps.data.sampling_rate, hps.data.mel_fmin, hps.data.mel_fmax, ) ) encoder_args = ( (tone, language, bert, style_vec) if hps.data.use_jp_extra else (tone, language, bert, ja_bert, en_bert, style_vec, speakers) ) return { "x": x, "x_lengths": x_lengths, "mel": mel, "mel_lengths": spec_lengths, "sid": speakers, "encoder_args": encoder_args, }, waveform def acoustic_loss(outputs: dict[str, torch.Tensor], hps: HyperParameters) -> torch.Tensor: return ( outputs["duration_loss"] + outputs["prior_loss"] + outputs["flow_loss"] * hps.train.c_matcha ) def run() -> None: parser = argparse.ArgumentParser() parser.add_argument("-c", "--config", required=True) parser.add_argument( "--stage", choices=("acoustic", "vocoder", "joint"), required=True ) parser.add_argument("--acoustic-checkpoint") parser.add_argument("--vocoder-checkpoint") parser.add_argument("--output-dir") parser.add_argument("--save-every", type=int, default=1000) parser.add_argument("--num-workers", type=int, default=1) parser.add_argument("--joint-timesteps", type=int, default=2) args = parser.parse_args() if args.stage == "joint" and ( not args.acoustic_checkpoint or not args.vocoder_checkpoint ): parser.error( "joint stage requires --acoustic-checkpoint and --vocoder-checkpoint" ) hps = HyperParameters.load_from_json(args.config) if hps.data.n_speakers < 1: raise ValueError("Separated mel training requires data.n_speakers >= 1") device = torch.device("cuda" if torch.cuda.is_available() else "cpu") acoustic, vocoder = build_models(hps) if args.acoustic_checkpoint: load_component( acoustic, args.acoustic_checkpoint, ("acoustic_model.", "module.acoustic_model.", "module."), ) if args.vocoder_checkpoint: load_component( vocoder, args.vocoder_checkpoint, ("vocoder.", "module.vocoder.", "dec.", "module.dec.", "module."), ) discriminator = None if args.stage == "acoustic": model: torch.nn.Module = acoustic.to(device) elif args.stage == "vocoder": model = vocoder.to(device) discriminator = MultiPeriodDiscriminator( hps.model.use_spectral_norm ).to(device) else: model = JointMelSynthesizer(acoustic, vocoder).to(device) discriminator = MultiPeriodDiscriminator( hps.model.use_spectral_norm ).to(device) optimizer = torch.optim.AdamW( model.parameters(), hps.train.learning_rate, betas=hps.train.betas, eps=hps.train.eps, ) optimizer_d = ( torch.optim.AdamW( discriminator.parameters(), hps.train.learning_rate, betas=hps.train.betas, eps=hps.train.eps, ) if discriminator is not None else None ) dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data) loader = DataLoader( dataset, batch_size=hps.train.batch_size, shuffle=True, num_workers=args.num_workers, pin_memory=device.type == "cuda", collate_fn=TextAudioSpeakerCollate(), drop_last=True, ) output_dir = Path(args.output_dir or Path(args.config).parent / "models") segment_frames = hps.train.segment_size // hps.data.hop_length global_step = 0 for epoch in range(1, hps.train.epochs + 1): model.train() for batch in loader: inputs, waveform = unpack_batch(batch, hps, device) if args.stage == "acoustic": outputs = acoustic( inputs["x"], inputs["x_lengths"], inputs["mel"], inputs["mel_lengths"], inputs["sid"], *inputs["encoder_args"], out_size=segment_frames, ) loss = acoustic_loss(outputs, hps) elif args.stage == "vocoder": assert discriminator is not None and optimizer_d is not None target_mel, ids = commons.rand_slice_segments( inputs["mel"], inputs["mel_lengths"], segment_frames ) real = commons.slice_segments( waveform, ids * hps.data.hop_length, hps.train.segment_size, ) generated = vocoder(target_mel, inputs["sid"]) real_scores, fake_scores, _, _ = discriminator( real, generated.detach() ) loss_d, _, _ = discriminator_loss(real_scores, fake_scores) optimizer_d.zero_grad(set_to_none=True) loss_d.backward() optimizer_d.step() _, fake_scores, fmap_r, fmap_g = discriminator(real, generated) generated_mel = mel_spectrogram_torch( generated.squeeze(1).float(), hps.data.filter_length, hps.data.n_mel_channels, hps.data.sampling_rate, hps.data.hop_length, hps.data.win_length, hps.data.mel_fmin, hps.data.mel_fmax, ) loss_g, _ = generator_loss(fake_scores) loss = ( loss_g + feature_loss(fmap_r, fmap_g) + F.l1_loss(target_mel, generated_mel) * hps.train.c_mel ) else: assert discriminator is not None and optimizer_d is not None generated, outputs = model( inputs["x"], inputs["x_lengths"], inputs["mel"], inputs["mel_lengths"], inputs["sid"], *inputs["encoder_args"], out_size=segment_frames, n_timesteps=args.joint_timesteps, ) ids = outputs["ids_slice"] real = commons.slice_segments( waveform, ids * hps.data.hop_length, hps.train.segment_size, ) real_scores, fake_scores, _, _ = discriminator( real, generated.detach() ) loss_d, _, _ = discriminator_loss(real_scores, fake_scores) optimizer_d.zero_grad(set_to_none=True) loss_d.backward() optimizer_d.step() real_scores, fake_scores, fmap_r, fmap_g = discriminator( real, generated ) generated_mel = mel_spectrogram_torch( generated.squeeze(1).float(), hps.data.filter_length, hps.data.n_mel_channels, hps.data.sampling_rate, hps.data.hop_length, hps.data.win_length, hps.data.mel_fmin, hps.data.mel_fmax, ) loss_g, _ = generator_loss(fake_scores) loss = ( acoustic_loss(outputs, hps) + loss_g + feature_loss(fmap_r, fmap_g) + F.l1_loss(outputs["target_mel"], generated_mel) * hps.train.c_mel ) optimizer.zero_grad(set_to_none=True) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 500) optimizer.step() global_step += 1 if global_step % hps.train.log_interval == 0: print( f"epoch={epoch} step={global_step} " f"stage={args.stage} loss={loss.item():.5f}" ) if global_step % args.save_every == 0: name = args.stage.upper() save_training_checkpoint( output_dir / f"{name}_{global_step}.pth", model, optimizer, epoch, global_step, ) if optimizer_d is not None and discriminator is not None: save_training_checkpoint( output_dir / f"D_{global_step}.pth", discriminator, optimizer_d, epoch, global_step, ) name = args.stage.upper() save_training_checkpoint( output_dir / f"{name}_{global_step}.pth", model, optimizer, hps.train.epochs, global_step, ) if __name__ == "__main__": run()