This commit is contained in:
tuna2134
2026-07-26 22:35:32 +09:00
parent aa70540aef
commit a283bec776

View File

@@ -16,6 +16,7 @@ from typing import Any
import torch import torch
from torch.nn import functional as F from torch.nn import functional as F
from torch.utils.data import DataLoader from torch.utils.data import DataLoader
from tqdm import tqdm
from data_utils import TextAudioSpeakerCollate, TextAudioSpeakerLoader from data_utils import TextAudioSpeakerCollate, TextAudioSpeakerLoader
from losses import discriminator_loss, feature_loss, generator_loss from losses import discriminator_loss, feature_loss, generator_loss
@@ -25,6 +26,7 @@ from style_bert_vits2.models.hyper_parameters import HyperParameters
from style_bert_vits2.models.mel_synthesizer import JointMelSynthesizer from style_bert_vits2.models.mel_synthesizer import JointMelSynthesizer
from style_bert_vits2.models.models import MultiPeriodDiscriminator from style_bert_vits2.models.models import MultiPeriodDiscriminator
from style_bert_vits2.nlp.symbols import SYMBOLS from style_bert_vits2.nlp.symbols import SYMBOLS
from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT
def _checkpoint_state(path: str) -> dict[str, torch.Tensor]: def _checkpoint_state(path: str) -> dict[str, torch.Tensor]:
@@ -217,6 +219,11 @@ def run() -> None:
parser.add_argument("--save-every", type=int, default=1000) parser.add_argument("--save-every", type=int, default=1000)
parser.add_argument("--num-workers", type=int, default=1) parser.add_argument("--num-workers", type=int, default=1)
parser.add_argument("--joint-timesteps", type=int, default=2) parser.add_argument("--joint-timesteps", type=int, default=2)
parser.add_argument(
"--no_progress_bar",
action="store_true",
help="Disable the tqdm training progress bar.",
)
args = parser.parse_args() args = parser.parse_args()
if args.stage == "joint" and ( if args.stage == "joint" and (
@@ -287,10 +294,19 @@ def run() -> None:
output_dir = Path(args.output_dir or Path(args.config).parent / "models") output_dir = Path(args.output_dir or Path(args.config).parent / "models")
segment_frames = hps.train.segment_size // hps.data.hop_length segment_frames = hps.train.segment_size // hps.data.hop_length
global_step = 0 global_step = 0
pbar = None
if not args.no_progress_bar:
pbar = tqdm(
total=hps.train.epochs * len(loader),
initial=global_step,
smoothing=0.05,
file=SAFE_STDOUT,
dynamic_ncols=True,
)
for epoch in range(1, hps.train.epochs + 1): for epoch in range(1, hps.train.epochs + 1):
model.train() model.train()
for batch in loader: for batch_idx, batch in enumerate(loader):
inputs, waveform = unpack_batch(batch, hps, device) inputs, waveform = unpack_batch(batch, hps, device)
if args.stage == "acoustic": if args.stage == "acoustic":
outputs = acoustic( outputs = acoustic(
@@ -391,11 +407,31 @@ def run() -> None:
torch.nn.utils.clip_grad_norm_(model.parameters(), 500) torch.nn.utils.clip_grad_norm_(model.parameters(), 500)
optimizer.step() optimizer.step()
global_step += 1 global_step += 1
if pbar is not None:
epoch_progress = 100.0 * (batch_idx + 1) / len(loader)
pbar.set_description(
f"{args.stage.capitalize()} "
f"{epoch}({epoch_progress:.0f}%)/{hps.train.epochs}"
)
if args.stage == "acoustic":
pbar.set_postfix(
loss=f"{loss.item():.4f}",
dur=f"{outputs['duration_loss'].item():.4f}",
prior=f"{outputs['prior_loss'].item():.4f}",
flow=f"{outputs['flow_loss'].item():.4f}",
)
else:
pbar.set_postfix(loss=f"{loss.item():.4f}")
pbar.update()
if global_step % hps.train.log_interval == 0: if global_step % hps.train.log_interval == 0:
print( message = (
f"epoch={epoch} step={global_step} " f"epoch={epoch} step={global_step} "
f"stage={args.stage} loss={loss.item():.5f}" f"stage={args.stage} loss={loss.item():.5f}"
) )
if pbar is not None:
pbar.write(message)
else:
print(message)
if global_step % args.save_every == 0: if global_step % args.save_every == 0:
name = args.stage.upper() name = args.stage.upper()
save_training_checkpoint( save_training_checkpoint(
@@ -422,6 +458,8 @@ def run() -> None:
hps.train.epochs, hps.train.epochs,
global_step, global_step,
) )
if pbar is not None:
pbar.close()
if __name__ == "__main__": if __name__ == "__main__":