Merge pull request #12 from medioqrity/master

修复一些辅助脚本的bug;增加resume训练的选项
This commit is contained in:
Stardust·减
2023-09-02 14:21:57 +08:00
committed by GitHub
4 changed files with 18 additions and 10 deletions

View File

@@ -39,8 +39,14 @@ def process_line(line):
assert bert.shape[-1] == len(phone)
torch.save(bert, bert_path)
with open(hps.data.training_files, encoding='utf-8') as f:
lines = f.readlines()
lines = []
with open(hps.data.training_files) as f:
lines.extend(f.readlines())
with open(hps.data.validation_files) as f:
lines.extend(f.readlines())
with Pool(processes=12) as pool: #A100 suitable config,if coom,please decrease the processess number.
for _ in tqdm(pool.imap_unordered(process_line, lines)):

View File

@@ -14,10 +14,10 @@ def process(item):
speaker = spkdir.replace("\\", "/").split("/")[-1]
wav_path = os.path.join(args.in_dir, speaker, wav_name)
if os.path.exists(wav_path) and '.wav' in wav_path:
os.makedirs(os.path.join(args.out_dir2, speaker), exist_ok=True)
wav, sr = librosa.load(wav_path, sr=args.sr2)
os.makedirs(os.path.join(args.out_dir, speaker), exist_ok=True)
wav, sr = librosa.load(wav_path, sr=args.sr)
soundfile.write(
os.path.join(args.out_dir2, speaker, wav_name),
os.path.join(args.out_dir, speaker, wav_name),
wav,
sr
)

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