split
This commit is contained in:
407
train_ms_mel.py
Normal file
407
train_ms_mel.py
Normal file
@@ -0,0 +1,407 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user