add --resume option

This commit is contained in:
med
2023-09-02 01:44:45 -04:00
parent 757ee212f5
commit fe904b9f5f
2 changed files with 7 additions and 5 deletions

View File

@@ -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)

View File

@@ -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