修改加载latest model的方式,修改global_step计算,增加preserved参数,增加train_with_pretrained_model参数
This commit is contained in:
+22
-11
@@ -100,18 +100,26 @@ def run(rank, n_gpus, hps):
|
|||||||
# load existing model
|
# load existing model
|
||||||
if hps.cont:
|
if hps.cont:
|
||||||
try:
|
try:
|
||||||
_, _, _, epoch_str = utils.load_checkpoint("./OUTPUT_MODEL/G_latest.pth", net_g, None)
|
_, _, _, epoch_str = utils.load_checkpoint(utils.latest_checkpoint_path(hps.model_dir, "G_[0-9]*.pth"), net_g, None)
|
||||||
_, _, _, epoch_str = utils.load_checkpoint("./OUTPUT_MODEL/D_latest.pth", net_d, None)
|
_, _, _, epoch_str = utils.load_checkpoint(utils.latest_checkpoint_path(hps.model_dir, "D_[0-9]*.pth"), net_d, None)
|
||||||
global_step = epoch_str * hps.train.batch_size
|
global_step = (epoch_str - 1) * len(train_loader)
|
||||||
except:
|
except:
|
||||||
print("Failed to find latest checkpoint, loading G_0.pth...")
|
print("Failed to find latest checkpoint, loading G_0.pth...")
|
||||||
_, _, _, epoch_str = utils.load_checkpoint("./pretrained_models/G_0.pth", net_g, None)
|
if hps.train_with_pretrained_model:
|
||||||
_, _, _, epoch_str = utils.load_checkpoint("./pretrained_models/D_0.pth", net_d, None)
|
print("Train with pretrained model...")
|
||||||
|
_, _, _, epoch_str = utils.load_checkpoint("./pretrained_models/G_0.pth", net_g, None)
|
||||||
|
_, _, _, epoch_str = utils.load_checkpoint("./pretrained_models/D_0.pth", net_d, None)
|
||||||
|
else:
|
||||||
|
print("Train without pretrained model...")
|
||||||
epoch_str = 1
|
epoch_str = 1
|
||||||
global_step = 0
|
global_step = 0
|
||||||
else:
|
else:
|
||||||
_, _, _, epoch_str = utils.load_checkpoint("./pretrained_models/G_0.pth", net_g, None)
|
if hps.train_with_pretrained_model:
|
||||||
_, _, _, epoch_str = utils.load_checkpoint("./pretrained_models/D_0.pth", net_d, None)
|
print("Train with pretrained model...")
|
||||||
|
_, _, _, epoch_str = utils.load_checkpoint("./pretrained_models/G_0.pth", net_g, None)
|
||||||
|
_, _, _, epoch_str = utils.load_checkpoint("./pretrained_models/D_0.pth", net_d, None)
|
||||||
|
else:
|
||||||
|
print("Train without pretrained model...")
|
||||||
epoch_str = 1
|
epoch_str = 1
|
||||||
global_step = 0
|
global_step = 0
|
||||||
# freeze all other layers except speaker embedding
|
# freeze all other layers except speaker embedding
|
||||||
@@ -256,13 +264,16 @@ def train_and_evaluate(rank, epoch, hps, nets, optims, schedulers, scaler, loade
|
|||||||
utils.save_checkpoint(net_g, None, hps.train.learning_rate, epoch,
|
utils.save_checkpoint(net_g, None, hps.train.learning_rate, epoch,
|
||||||
os.path.join(hps.model_dir, "G_latest.pth".format(global_step)))
|
os.path.join(hps.model_dir, "G_latest.pth".format(global_step)))
|
||||||
utils.save_checkpoint(net_d, None, hps.train.learning_rate, epoch, os.path.join(hps.model_dir, "D_{}.pth".format(global_step)))
|
utils.save_checkpoint(net_d, None, hps.train.learning_rate, epoch, os.path.join(hps.model_dir, "D_{}.pth".format(global_step)))
|
||||||
utils.save_checkpoint(net_d, None, hps.train.learning_rate, epoch,
|
# utils.save_checkpoint(net_d, None, hps.train.learning_rate, epoch,
|
||||||
os.path.join(hps.model_dir, "D_latest.pth".format(global_step)))
|
# os.path.join(hps.model_dir, "D_latest.pth".format(global_step)))
|
||||||
old_g=os.path.join(hps.model_dir, "G_{}.pth".format(global_step-4000))
|
old_g = utils.oldest_checkpoint_path(hps.model_dir, "G_[0-9]*.pth",
|
||||||
old_d=os.path.join(hps.model_dir, "D_{}.pth".format(global_step-4000))
|
preserved=hps.preserved) # Preserve 4 (default) historical checkpoints.
|
||||||
|
old_d = utils.oldest_checkpoint_path(hps.model_dir, "D_[0-9]*.pth", preserved=hps.preserved)
|
||||||
if os.path.exists(old_g):
|
if os.path.exists(old_g):
|
||||||
|
print(f"remove {old_g}")
|
||||||
os.remove(old_g)
|
os.remove(old_g)
|
||||||
if os.path.exists(old_d):
|
if os.path.exists(old_d):
|
||||||
|
print(f"remove {old_d}")
|
||||||
os.remove(old_d)
|
os.remove(old_d)
|
||||||
global_step += 1
|
global_step += 1
|
||||||
if epoch > hps.max_epochs:
|
if epoch > hps.max_epochs:
|
||||||
|
|||||||
@@ -204,14 +204,29 @@ def summarize(writer, global_step, scalars={}, histograms={}, images={}, audios=
|
|||||||
writer.add_audio(k, v, global_step, audio_sampling_rate)
|
writer.add_audio(k, v, global_step, audio_sampling_rate)
|
||||||
|
|
||||||
|
|
||||||
def latest_checkpoint_path(dir_path, regex="G_*.pth"):
|
def extract_digits(f):
|
||||||
|
digits = "".join(filter(str.isdigit, f))
|
||||||
|
return int(digits) if digits else -1
|
||||||
|
|
||||||
|
|
||||||
|
def latest_checkpoint_path(dir_path, regex="G_[0-9]*.pth"):
|
||||||
f_list = glob.glob(os.path.join(dir_path, regex))
|
f_list = glob.glob(os.path.join(dir_path, regex))
|
||||||
f_list.sort(key=lambda f: int("".join(filter(str.isdigit, f))))
|
f_list.sort(key=lambda f: extract_digits(f))
|
||||||
x = f_list[-1]
|
x = f_list[-1]
|
||||||
print(x)
|
print(f"latest_checkpoint_path:{x}")
|
||||||
return x
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
def oldest_checkpoint_path(dir_path, regex="G_[0-9]*.pth", preserved=4):
|
||||||
|
f_list = glob.glob(os.path.join(dir_path, regex))
|
||||||
|
f_list.sort(key=lambda f: extract_digits(f))
|
||||||
|
if len(f_list) > preserved:
|
||||||
|
x = f_list[0]
|
||||||
|
print(f"oldest_checkpoint_path:{x}")
|
||||||
|
return x
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
def plot_spectrogram_to_numpy(spectrogram):
|
def plot_spectrogram_to_numpy(spectrogram):
|
||||||
global MATPLOTLIB_FLAG
|
global MATPLOTLIB_FLAG
|
||||||
if not MATPLOTLIB_FLAG:
|
if not MATPLOTLIB_FLAG:
|
||||||
@@ -288,6 +303,10 @@ def get_hparams(init=True):
|
|||||||
help='finetune epochs')
|
help='finetune epochs')
|
||||||
parser.add_argument('--cont', type=bool, default=False, help='whether to continue training on the latest checkpoint')
|
parser.add_argument('--cont', type=bool, default=False, help='whether to continue training on the latest checkpoint')
|
||||||
parser.add_argument('--drop_speaker_embed', type=bool, default=False, help='whether to drop existing characters')
|
parser.add_argument('--drop_speaker_embed', type=bool, default=False, help='whether to drop existing characters')
|
||||||
|
parser.add_argument('--train_with_pretrained_model', type=bool, default=True,
|
||||||
|
help='whether to train with pretrained model')
|
||||||
|
parser.add_argument('--preserved', type=int, default=4,
|
||||||
|
help='Number of preserved models')
|
||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
model_dir = os.path.join("./", args.model)
|
model_dir = os.path.join("./", args.model)
|
||||||
@@ -312,6 +331,8 @@ def get_hparams(init=True):
|
|||||||
hparams.max_epochs = args.max_epochs
|
hparams.max_epochs = args.max_epochs
|
||||||
hparams.cont = args.cont
|
hparams.cont = args.cont
|
||||||
hparams.drop_speaker_embed = args.drop_speaker_embed
|
hparams.drop_speaker_embed = args.drop_speaker_embed
|
||||||
|
hparams.train_with_pretrained_model = args.train_with_pretrained_model
|
||||||
|
hparams.preserved = args.preserved
|
||||||
return hparams
|
return hparams
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user