fix
This commit is contained in:
@@ -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__":
|
||||
|
||||
Reference in New Issue
Block a user