Merge pull request #14 from RedRayz/dev

学習中に全体の進捗を表示
This commit is contained in:
litagin02
2024-01-06 11:46:25 +09:00
committed by GitHub

View File

@@ -107,6 +107,11 @@ def run():
action="store_true", action="store_true",
help="Skip saving default style config and mean vector.", help="Skip saving default style config and mean vector.",
) )
parser.add_argument(
'--no_progress_bar',
action='store_true',
help='Do not show the progress bar while training.'
)
args = parser.parse_args() args = parser.parse_args()
model_dir = os.path.join(args.model, config.train_ms_config.model_dir) model_dir = os.path.join(args.model, config.train_ms_config.model_dir)
if not os.path.exists(model_dir): if not os.path.exists(model_dir):
@@ -327,7 +332,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:
@@ -364,6 +369,13 @@ def run():
else: else:
scheduler_dur_disc = None scheduler_dur_disc = None
scaler = GradScaler(enabled=hps.train.bf16_run) scaler = GradScaler(enabled=hps.train.bf16_run)
logger.info("Start training.")
diff = abs(epoch_str * len(train_loader) - (hps.train.epochs + 1) * len(train_loader))
pbar = None
if not args.no_progress_bar:
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:
@@ -379,6 +391,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 +407,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 +449,9 @@ def run():
for_infer=True, for_infer=True,
) )
if pbar is not None:
pbar.close()
def train_and_evaluate( def train_and_evaluate(
rank, rank,
@@ -446,6 +465,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 +496,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 +684,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,13 +727,15 @@ 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: if pbar is None and rank == 0:
logger.info(f"====> Epoch: {epoch}, step: {global_step}") logger.info(f"====> Epoch: {epoch}, step: {global_step}")
def evaluate(hps, generator, eval_loader, writer_eval): def evaluate(hps, generator, eval_loader, writer_eval):
generator.eval() generator.eval()
image_dict = {} image_dict = {}