diff --git a/train_ms_mel.py b/train_ms_mel.py index a78e5bf..8c3db6f 100644 --- a/train_ms_mel.py +++ b/train_ms_mel.py @@ -16,6 +16,7 @@ from typing import Any import torch from torch.nn import functional as F from torch.utils.data import DataLoader +from tqdm import tqdm from data_utils import TextAudioSpeakerCollate, TextAudioSpeakerLoader 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.models import MultiPeriodDiscriminator 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]: @@ -217,6 +219,11 @@ def run() -> None: 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) + parser.add_argument( + "--no_progress_bar", + action="store_true", + help="Disable the tqdm training progress bar.", + ) args = parser.parse_args() if args.stage == "joint" and ( @@ -287,10 +294,19 @@ def run() -> None: 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 + 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): model.train() - for batch in loader: + for batch_idx, batch in enumerate(loader): inputs, waveform = unpack_batch(batch, hps, device) if args.stage == "acoustic": outputs = acoustic( @@ -391,11 +407,31 @@ def run() -> None: torch.nn.utils.clip_grad_norm_(model.parameters(), 500) optimizer.step() 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: - print( + message = ( f"epoch={epoch} step={global_step} " 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: name = args.stage.upper() save_training_checkpoint( @@ -422,6 +458,8 @@ def run() -> None: hps.train.epochs, global_step, ) + if pbar is not None: + pbar.close() if __name__ == "__main__":