408 lines
14 KiB
Python
408 lines
14 KiB
Python
"""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()
|