This commit is contained in:
rcell
2023-07-29 13:06:50 +08:00
parent 9fd6e9a259
commit 3dd30a204e
3 changed files with 52 additions and 52 deletions

View File

@@ -8,7 +8,7 @@ import torch.utils.data
import commons
from mel_processing import spectrogram_torch
from utils import load_wav_to_torch, load_filepaths_and_text
from text import cleaned_text_to_sequence, get_bert
from text import cleaned_text_to_sequence
"""Multi speaker version"""
@@ -113,14 +113,14 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
word2ph[i] = word2ph[i] * 2
word2ph[0] += 1
bert_path = wav_path.replace(".wav", ".bert.pt")
try:
bert = torch.load(bert_path)
assert bert.shape[-1] == len(phone)
except:
bert = get_bert(text, word2ph, language_str)
torch.save(bert, bert_path)
print(bert.shape[-1], bert_path, text, pold)
assert bert.shape[-1] == len(phone)
# try:
bert = torch.load(bert_path)
assert bert.shape[-1] == len(phone)
# except:
# bert = get_bert(text, word2ph, language_str)
# torch.save(bert, bert_path)
# print(bert.shape[-1], bert_path, text, pold)
# assert bert.shape[-1] == len(phone)
assert bert.shape[-1] == len(phone), (
bert.shape, len(phone), sum(word2ph), p1, p2, t1, t2, pold, pold2, word2ph, text, w2pho)

View File

@@ -11,47 +11,47 @@ from text import cleaned_text_to_sequence, get_bert
config_path = 'configs/config.json'
hps = utils.get_hparams_from_file(config_path)
train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data)
eval_dataset = TextAudioSpeakerLoader(hps.data.validation_files, hps.data)
# train_dataset = TextAudioSpeakerLoader(hps.data.training_files, hps.data)
# eval_dataset = TextAudioSpeakerLoader(hps.data.validation_files, hps.data)
#
# collate_fn = TextAudioSpeakerCollate()
# train_loader = DataLoader(train_dataset, num_workers=12, shuffle=False,
# batch_size=32, pin_memory=True,
# drop_last=False, collate_fn=collate_fn)
# eval_loader = DataLoader(eval_dataset, num_workers=12, shuffle=False,
# batch_size=32, pin_memory=True,
# drop_last=False, collate_fn=collate_fn)
# for _ in tqdm(train_loader):
# pass
# for _ in tqdm(eval_loader):
# pass
collate_fn = TextAudioSpeakerCollate()
train_loader = DataLoader(train_dataset, num_workers=12, shuffle=False,
batch_size=32, pin_memory=True,
drop_last=False, collate_fn=collate_fn)
eval_loader = DataLoader(eval_dataset, num_workers=12, shuffle=False,
batch_size=32, pin_memory=True,
drop_last=False, collate_fn=collate_fn)
for _ in tqdm(train_loader):
pass
for _ in tqdm(eval_loader):
pass
for line in tqdm( open(hps.data.training_files).readlines()):
_id, spk, language_str, text, phones, tone, word2ph = line.strip().split("|")
phone = phones.split(" ")
tone = [int(i) for i in tone.split(" ")]
word2ph = [int(i) for i in word2ph.split(" ")]
# print(text, word2ph,phone, tone, language_str)
w2pho = [i for i in word2ph]
word2ph = [i for i in word2ph]
phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str)
pold2 = phone
# for line in tqdm( open(hps.data.training_files).readlines()):
# _id, spk, language_str, text, phones, tone, word2ph = line.strip().split("|")
# phone = phones.split(" ")
# tone = [int(i) for i in tone.split(" ")]
# word2ph = [int(i) for i in word2ph.split(" ")]
# # print(text, word2ph,phone, tone, language_str)
# w2pho = [i for i in word2ph]
# word2ph = [i for i in word2ph]
# phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str)
# pold2 = phone
if hps.data.add_blank:
phone = commons.intersperse(phone, 0)
tone = commons.intersperse(tone, 0)
language = commons.intersperse(language, 0)
for i in range(len(word2ph)):
word2ph[i] = word2ph[i] * 2
word2ph[0] += 1
wav_path = f'dataset/{spk}/{_id}.wav'
# if hps.data.add_blank:
# phone = commons.intersperse(phone, 0)
# tone = commons.intersperse(tone, 0)
# language = commons.intersperse(language, 0)
# for i in range(len(word2ph)):
# word2ph[i] = word2ph[i] * 2
# word2ph[0] += 1
# wav_path = f'dataset/{spk}/{_id}.wav'
# bert_path = wav_path.replace(".wav", ".bert.pt")
# try:
# bert = torch.load(bert_path)
# assert bert.shape[-1] == len(phone)
# except:
# bert = get_bert(text, word2ph, language_str)
# assert bert.shape[-1] == len(phone)
# torch.save(bert, bert_path)
bert_path = wav_path.replace(".wav", ".bert.pt")
try:
bert = torch.load(bert_path)
assert bert.shape[-1] == len(phone)
except:
bert = get_bert(text, word2ph, language_str)
assert bert.shape[-1] == len(phone)
torch.save(bert, bert_path)

View File

@@ -82,13 +82,13 @@ class ToneSandhi():
} and word not in self.must_not_neural_tone_words:
finals[j] = finals[j][:-1] + "5"
ge_idx = word.find("")
if len(word) >= 1 and word[-1] in "吧呢啊呐噻嘛吖嗨呐哦哒额滴哩哟喽啰耶喔诶":
if len(word) >= 1 and word[-1] in "吧呢啊呐噻嘛吖嗨呐哦哒额滴哩哟喽啰耶喔诶":
finals[-1] = finals[-1][:-1] + "5"
elif len(word) >= 1 and word[-1] in "的地得":
finals[-1] = finals[-1][:-1] + "5"
# e.g. 走了, 看着, 去过
elif len(word) == 1 and word in "了着过" and pos in {"ul", "uz", "ug"}:
finals[-1] = finals[-1][:-1] + "5"
# elif len(word) == 1 and word in "了着过" and pos in {"ul", "uz", "ug"}:
# finals[-1] = finals[-1][:-1] + "5"
elif len(word) > 1 and word[-1] in "们子" and pos in {
"r", "n"
} and word not in self.must_not_neural_tone_words: