Show total progress and don't evaluate when script starts training

This commit is contained in:
RedRayz
2024-01-05 21:32:00 +09:00
parent c7571a7c39
commit 701adfa2a5

View File

@@ -327,7 +327,7 @@ def run():
utils.get_steps(utils.latest_checkpoint_path(model_dir, "G_*.pth")) utils.get_steps(utils.latest_checkpoint_path(model_dir, "G_*.pth"))
) )
logger.info( logger.info(
f"******************检测到模型存在,epoch {epoch_str}gloabl step {global_step}*********************" f"******************Found the model. Current epoch is {epoch_str}, gloabl step is {global_step}*********************"
) )
else: else:
try: try:
@@ -365,6 +365,10 @@ def run():
scheduler_dur_disc = None scheduler_dur_disc = None
scaler = GradScaler(enabled=hps.train.bf16_run) scaler = GradScaler(enabled=hps.train.bf16_run)
diff = abs(epoch_str * len(train_loader) - (hps.train.epochs + 1) * len(train_loader))
pbar = tqdm(total=global_step + diff, initial=global_step, smoothing=0.05, file=SAFE_STDOUT)
initial_step = global_step
for epoch in range(epoch_str, hps.train.epochs + 1): for epoch in range(epoch_str, hps.train.epochs + 1):
if rank == 0: if rank == 0:
train_and_evaluate( train_and_evaluate(
@@ -379,6 +383,8 @@ def run():
[train_loader, eval_loader], [train_loader, eval_loader],
logger, logger,
[writer, writer_eval], [writer, writer_eval],
pbar,
initial_step,
) )
else: else:
train_and_evaluate( train_and_evaluate(
@@ -393,6 +399,8 @@ def run():
[train_loader, None], [train_loader, None],
None, None,
None, None,
pbar,
initial_step,
) )
scheduler_g.step() scheduler_g.step()
scheduler_d.step() scheduler_d.step()
@@ -433,6 +441,8 @@ def run():
for_infer=True, for_infer=True,
) )
pbar.close()
def train_and_evaluate( def train_and_evaluate(
rank, rank,
@@ -446,6 +456,8 @@ def train_and_evaluate(
loaders, loaders,
logger, logger,
writers, writers,
pbar: tqdm,
initial_step: int
): ):
net_g, net_d, net_dur_disc = nets net_g, net_d, net_dur_disc = nets
optim_g, optim_d, optim_dur_disc = optims optim_g, optim_d, optim_dur_disc = optims
@@ -475,7 +487,7 @@ def train_and_evaluate(
ja_bert, ja_bert,
en_bert, en_bert,
style_vec, style_vec,
) in enumerate(tqdm(train_loader, file=SAFE_STDOUT)): ) in enumerate(train_loader):
if net_g.module.use_noise_scaled_mas: if net_g.module.use_noise_scaled_mas:
current_mas_noise_scale = ( current_mas_noise_scale = (
net_g.module.mas_noise_scale_initial net_g.module.mas_noise_scale_initial
@@ -663,7 +675,7 @@ def train_and_evaluate(
scalars=scalar_dict, scalars=scalar_dict,
) )
if global_step % hps.train.eval_interval == 0 and global_step != 0: if global_step % hps.train.eval_interval == 0 and global_step != 0 and initial_step != global_step:
evaluate(hps, net_g, eval_loader, writer_eval) evaluate(hps, net_g, eval_loader, writer_eval)
utils.save_checkpoint( utils.save_checkpoint(
net_g, net_g,
@@ -706,11 +718,12 @@ def train_and_evaluate(
) )
global_step += 1 global_step += 1
if pbar is not None:
pbar.set_description("Epoch {}({:.0f}%)/{}".format(epoch, 100.0 * batch_idx / len(train_loader), hps.train.epochs))
pbar.update()
# 本家ではこれをスピードアップのために消すと書かれていたので、一応消してみる # 本家ではこれをスピードアップのために消すと書かれていたので、一応消してみる
# gc.collect() # gc.collect()
# torch.cuda.empty_cache() # torch.cuda.empty_cache()
if rank == 0:
logger.info(f"====> Epoch: {epoch}, step: {global_step}")
def evaluate(hps, generator, eval_loader, writer_eval): def evaluate(hps, generator, eval_loader, writer_eval):