From 1fa7a8086b133d365e666a98e29ca9b665e40a5a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stardust=C2=B7=E5=87=8F?= <2225664821@qq.com> Date: Wed, 23 Aug 2023 17:06:45 +0800 Subject: [PATCH] Update train_ms.py --- train_ms.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/train_ms.py b/train_ms.py index 41ea352..7299116 100644 --- a/train_ms.py +++ b/train_ms.py @@ -331,7 +331,9 @@ def train_and_evaluate(rank, epoch, hps, nets, optims, schedulers, scaler, loade os.path.join(hps.model_dir, "G_{}.pth".format(global_step))) utils.save_checkpoint(net_d, optim_d, hps.train.learning_rate, epoch, os.path.join(hps.model_dir, "D_{}.pth".format(global_step))) - keep_ckpts = getattr(hps.train, 'keep_ckpts', 3) + if net_dur_disc is not None: + utils.save_checkpoint(net_dur_disc, optim_dur_disc, hps.train.learning_rate, epoch, os.path.join(hps.model_dir, "DUR_{}.pth".format(global_step))) + keep_ckpts = getattr(hps.train, 'keep_ckpts', 5) if keep_ckpts > 0: utils.clean_checkpoints(path_to_models=hps.model_dir, n_ckpts_to_keep=keep_ckpts, sort_by_time=True)