This commit is contained in:
tuna2134
2026-07-24 16:24:50 +09:00
parent 1785e81509
commit 8125666e22
7 changed files with 1053 additions and 2 deletions

407
train_ms_mel.py Normal file
View 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()