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