From fe904b9f5f0c4b5145c170775ec6018608302566 Mon Sep 17 00:00:00 2001 From: med Date: Sat, 2 Sep 2023 01:44:45 -0400 Subject: [PATCH] add --resume option --- train_ms.py | 10 +++++----- utils.py | 2 ++ 2 files changed, 7 insertions(+), 5 deletions(-) diff --git a/train_ms.py b/train_ms.py index 81048e3..93846c2 100644 --- a/train_ms.py +++ b/train_ms.py @@ -155,11 +155,11 @@ def run(rank, n_gpus, hps): if pretrain_dir is None: try: if net_dur_disc is not None: - _, _, _, epoch_str = utils.load_checkpoint(utils.latest_checkpoint_path(hps.model_dir, "DUR_*.pth"), net_dur_disc, optim_dur_disc, skip_optimizer=True) - _, _, _, epoch_str = utils.load_checkpoint(utils.latest_checkpoint_path(hps.model_dir, "G_*.pth"), net_g, - optim_g, skip_optimizer=True) - _, _, _, epoch_str = utils.load_checkpoint(utils.latest_checkpoint_path(hps.model_dir, "D_*.pth"), net_d, - optim_d, skip_optimizer=True) + _, optim_dur_disc, _, epoch_str = utils.load_checkpoint(utils.latest_checkpoint_path(hps.model_dir, "DUR_*.pth"), net_dur_disc, optim_dur_disc, skip_optimizer=not hps.resume) + _, optim_g, _, epoch_str = utils.load_checkpoint(utils.latest_checkpoint_path(hps.model_dir, "G_*.pth"), net_g, + optim_g, skip_optimizer=not hps.resume) + _, optim_d, _, epoch_str = utils.load_checkpoint(utils.latest_checkpoint_path(hps.model_dir, "D_*.pth"), net_d, + optim_d, skip_optimizer=not hps.resume) epoch_str = max(epoch_str, 1) global_step = (epoch_str - 1) * len(train_loader) diff --git a/utils.py b/utils.py index 0e00972..f07a682 100644 --- a/utils.py +++ b/utils.py @@ -158,6 +158,7 @@ def get_hparams(init=True): help='JSON file for configuration') parser.add_argument('-m', '--model', type=str, required=True, help='Model name') + parser.add_argument('--resume', dest='resume', action="store_true", default=False, help="resume training from latest checkpoint of the given model") args = parser.parse_args() model_dir = os.path.join("./logs", args.model) @@ -179,6 +180,7 @@ def get_hparams(init=True): hparams = HParams(**config) hparams.model_dir = model_dir + hparams.resume = args.resume return hparams