Make trainer generate sample sentences

This commit is contained in:
ZDisket
2023-05-25 23:48:49 -03:00
parent e0848f7715
commit 2a263444a1
4 changed files with 155 additions and 7 deletions
+81
View File
@@ -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]
}
}
+35
View File
@@ -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)
+35 -7
View File
@@ -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]]})
+4
View File
@@ -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