33
train_ms.py
33
train_ms.py
@@ -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 = {}
|
||||||
|
|||||||
Reference in New Issue
Block a user