mirror of
https://github.com/storytold/vits-finetuning.git
synced 2026-10-09 00:09:52 +00:00
Make trainer generate sample sentences
This commit is contained in:
@@ -0,0 +1,81 @@
|
||||
{
|
||||
"train": {
|
||||
"log_interval": 200,
|
||||
"eval_interval": 2000,
|
||||
"save_interval": 5000,
|
||||
"seed": 1234,
|
||||
"epochs": 20000,
|
||||
"learning_rate": 2e-4,
|
||||
"betas": [0.8, 0.99],
|
||||
"eps": 1e-9,
|
||||
"batch_size": 32,
|
||||
"grad_acc_steps": 2,
|
||||
"fp16_run": true,
|
||||
"lr_decay": 0.999875,
|
||||
"segment_size": 8192,
|
||||
"init_lr_ratio": 1,
|
||||
"warmup_epochs": 0,
|
||||
"c_mel": 45,
|
||||
"c_kl": 1.0,
|
||||
"use_8bit": true
|
||||
},
|
||||
"data": {
|
||||
"training_files":"filelists/train_filelist.txt.cleaned",
|
||||
"validation_files":"filelists/val_filelist.txt.cleaned",
|
||||
"text_cleaners":["arpa_cleaners"],
|
||||
"max_wav_value": 32768.0,
|
||||
"sampling_rate": 44100,
|
||||
"filter_length": 2048,
|
||||
"hop_length": 512,
|
||||
"win_length": 2048,
|
||||
"n_mel_channels": 80,
|
||||
"mel_fmin": 20,
|
||||
"mel_fmax": 11025,
|
||||
"add_blank": true,
|
||||
"n_speakers": 0,
|
||||
"cleaned_text": true,
|
||||
"bert_model_name": "huawei-noah/TinyBERT_General_4L_312D"
|
||||
},
|
||||
"model": {
|
||||
"inter_channels": 192,
|
||||
"hidden_channels": 192,
|
||||
"filter_channels": 768,
|
||||
"n_heads": 2,
|
||||
"n_layers": 6,
|
||||
"kernel_size": 3,
|
||||
"p_dropout": 0.1,
|
||||
"resblock": "1",
|
||||
"resblock_kernel_sizes": [3,7,11],
|
||||
"resblock_dilation_sizes": [[1,3,5], [1,3,5], [1,3,5]],
|
||||
"upsample_rates": [8,8],
|
||||
"upsample_initial_channel": 512,
|
||||
"upsample_kernel_sizes": [16,16],
|
||||
"n_layers_q": 3,
|
||||
"use_spectral_norm": false,
|
||||
"gen_istft_n_fft" : 32,
|
||||
"gen_istft_hop_size": 8,
|
||||
"moji_start_size": 2304,
|
||||
"moji_enc_sizes": [1024,512,128,32],
|
||||
"bert_size": 312,
|
||||
"bert_final": 32
|
||||
},
|
||||
"disc": {
|
||||
"combd_channels": [16, 64, 256, 1024, 1024, 1024],
|
||||
"combd_kernels" : [[7, 11, 11, 11, 11, 5], [11, 21, 21, 21, 21, 5], [15, 41, 41, 41, 41, 5]],
|
||||
"combd_groups" : [1, 4, 16, 64, 256, 1],
|
||||
"combd_strides": [1, 1, 4, 4, 4, 1],
|
||||
|
||||
"tkernels" : [7, 5, 3],
|
||||
"fkernel" : 5,
|
||||
"tchannels" : [64, 128, 256, 256, 256],
|
||||
"fchannels" : [32, 64, 128, 128, 128],
|
||||
"tstrides" : [[1, 1, 3, 3, 1], [1, 1, 3, 3, 1], [1, 1, 3, 3, 1]],
|
||||
"fstride" : [1, 1, 3, 3, 1],
|
||||
"tdilations" : [[[5, 7, 11], [5, 7, 11], [5, 7, 11], [5, 7, 11], [5, 7, 11], [5, 7, 11]], [[3, 5, 7], [3, 5, 7], [3, 5, 7], [3, 5, 7], [3, 5, 7]], [[1, 2, 3], [1, 2, 3], [1, 2, 3], [1, 2, 3], [1, 2, 3]]],
|
||||
"fdilations" : [[1, 2, 3], [1, 2, 3], [1, 2, 3], [2, 3, 5], [2, 3, 5]],
|
||||
"pqmf_n" : 16,
|
||||
"pqmf_m" : 64,
|
||||
"freq_init_ch" : 128,
|
||||
"tsubband" : [6, 11, 16]
|
||||
}
|
||||
}
|
||||
@@ -6,6 +6,7 @@ import os
|
||||
from moji.moji import TorchMoji
|
||||
import torch
|
||||
from bertfe import BERTFrontEnd
|
||||
import commons
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser()
|
||||
@@ -14,11 +15,21 @@ if __name__ == '__main__':
|
||||
parser.add_argument("--filelists", nargs="+", default=["filelists/ljs_audio_text_val_filelist.txt", "filelists/ljs_audio_text_test_filelist.txt"])
|
||||
parser.add_argument("--text_cleaners", nargs="+", default=["english_cleaners2"])
|
||||
parser.add_argument("--bert", default="huawei-noah/TinyBERT_General_4L_312D")
|
||||
parser.add_argument("--test", default="")
|
||||
|
||||
|
||||
args = parser.parse_args()
|
||||
moji = TorchMoji(verbose=True)
|
||||
bert_f = BERTFrontEnd(model_name=args.bert)
|
||||
test_sentences = ["The quick brown fox jumps over the lazy dog",
|
||||
"In a galaxy far, far away, a young hero embarks on an epic adventure",
|
||||
"Ladies and gentlemen, welcome to the annual science fair!",
|
||||
"The chef skillfully prepares a delicious gourmet meal with fresh ingredients and exquisite flavors",
|
||||
"The crowd erupted in cheers as the team scored the winning goal in the final seconds of the match"]
|
||||
|
||||
if len(args.test) > 1:
|
||||
with open(args.test) as f:
|
||||
test_sentences = f.readlines()
|
||||
|
||||
|
||||
if "arpa_cleaners" in args.text_cleaners:
|
||||
@@ -61,4 +72,28 @@ if __name__ == '__main__':
|
||||
with open(new_filelist, "w", encoding="utf-8") as f:
|
||||
for x in new_fp_text:
|
||||
f.write("|".join(x) + "\n")
|
||||
|
||||
print("Writing test data...")
|
||||
test_path = "test_preproc"
|
||||
for i, test_sent in tqdm(enumerate(test_sentences)):
|
||||
moji_t = moji(test_sent)
|
||||
bert_t, _ = bert_f.infer(test_sent)
|
||||
|
||||
test_cleaned_text = text._clean_text(test_sent, args.text_cleaners)
|
||||
cleaned_text_indices = text.cleaned_text_to_sequence(test_cleaned_text)
|
||||
text_norm = commons.intersperse(cleaned_text_indices, 0)
|
||||
text_norm = torch.LongTensor(text_norm)
|
||||
|
||||
if not os.path.exists(test_path):
|
||||
os.makedirs(test_path)
|
||||
|
||||
test_rawfn = f"test{i}.pt"
|
||||
test_fullfn = os.path.join(test_path, test_rawfn)
|
||||
|
||||
torch.save([text_norm, moji_t, bert_t], test_fullfn)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -355,7 +355,9 @@ def evaluate(hps, generator, eval_loader, writer_eval):
|
||||
bert = bert[:1]
|
||||
bert_lens = bert_lens[:1]
|
||||
break
|
||||
|
||||
|
||||
image_dict = {}
|
||||
audio_dict = {}
|
||||
y_hat, attn, mask, *_ = generator.module.infer(x, x_lengths, tm_hidden, bert, bert_lens, max_len=1000)
|
||||
y_hat_lengths = mask.sum([1,2]).long() * hps.data.hop_length
|
||||
|
||||
@@ -376,12 +378,38 @@ def evaluate(hps, generator, eval_loader, writer_eval):
|
||||
hps.data.mel_fmin,
|
||||
hps.data.mel_fmax
|
||||
)
|
||||
image_dict = {
|
||||
"gen/mel": utils.plot_spectrogram_to_numpy(y_hat_mel[0].cpu().numpy())
|
||||
}
|
||||
audio_dict = {
|
||||
"gen/audio": y_hat[0,:,:y_hat_lengths[0]]
|
||||
}
|
||||
image_dict["gen/mel"] = utils.plot_spectrogram_to_numpy(y_hat_mel[0].cpu().numpy())
|
||||
audio_dict["gen/audio"] = y_hat[0,:,:y_hat_lengths[0]]
|
||||
print("Inferring test sentences...")
|
||||
for t_idx, test_file in enumerate(sorted(os.listdir(hps.test_dp))):
|
||||
if not ".pt" in test_file:
|
||||
continue
|
||||
|
||||
t_text_norm, t_moji, t_bert = torch.load(os.path.join(hps.test_dp, test_file))
|
||||
|
||||
t_text_norm = t_text_norm.unsqueeze(0).cuda(0)
|
||||
t_text_lengths = torch.LongTensor([t_text_norm.size(1)]).cuda(0)
|
||||
|
||||
t_moji = t_moji.squeeze().unsqueeze(0).cuda(0)
|
||||
t_bert_lens = torch.LongTensor([t_bert.size(1)]).cuda(0)
|
||||
t_bert = t_bert.cuda(0)
|
||||
|
||||
audio, attn, t_mask, _ = generator.module.infer(t_text_norm, t_text_lengths, t_moji, t_bert, t_bert_lens, noise_scale=.667, noise_scale_w=0.8, length_scale=1.0)
|
||||
test_audio_lengths = t_mask.sum([1,2]).long() * hps.data.hop_length
|
||||
y_test_mel = mel_spectrogram_torch(
|
||||
audio.squeeze(1).float(),
|
||||
hps.data.filter_length,
|
||||
hps.data.n_mel_channels,
|
||||
hps.data.sampling_rate,
|
||||
hps.data.hop_length,
|
||||
hps.data.win_length,
|
||||
hps.data.mel_fmin,
|
||||
hps.data.mel_fmax
|
||||
)
|
||||
image_dict[f"gen/mel_test{t_idx}"] = utils.plot_spectrogram_to_numpy(y_test_mel[0].cpu().numpy())
|
||||
audio_dict[f"gen/audio_test{t_idx}"] = audio[0,:,:test_audio_lengths[0]]
|
||||
|
||||
|
||||
if global_step == 0:
|
||||
image_dict.update({"gt/mel": utils.plot_spectrogram_to_numpy(mel[0].cpu().numpy())})
|
||||
audio_dict.update({"gt/audio": y[0,:,:y_lengths[0]]})
|
||||
|
||||
@@ -161,6 +161,8 @@ def get_hparams(init=True):
|
||||
parser.add_argument('-l','--use-latest', action='store_true', help='Whether to use latest-style checkpointing (saves to the same file every time) instead of regular. Saves disk space')
|
||||
parser.add_argument('-w', '--workers', type=int, default=8,
|
||||
help='Number of workers to use for dataloading. Should not be less than the amount of CPU threads')
|
||||
parser.add_argument('-t', '--test-dp', type=str, default="test_preproc",
|
||||
help='Path where preprocessed test data is')
|
||||
|
||||
|
||||
|
||||
@@ -169,6 +171,7 @@ def get_hparams(init=True):
|
||||
pt_path = args.pretrained
|
||||
use_latest = args.use_latest
|
||||
n_workers = args.workers
|
||||
test_data_path = args.test_dp
|
||||
|
||||
|
||||
if not os.path.exists(model_dir):
|
||||
@@ -191,6 +194,7 @@ def get_hparams(init=True):
|
||||
hparams.pt_path = pt_path
|
||||
hparams.use_latest = use_latest
|
||||
hparams.n_workers = n_workers
|
||||
hparams.test_dp = test_data_path
|
||||
return hparams
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user