Skip SOS and EOS when processing sliced sentances (#179)

This commit is contained in:
OedoSoldier
2023-11-14 08:45:00 +08:00
committed by GitHub
parent 52f895aa9f
commit ad694ee050
2 changed files with 54 additions and 10 deletions

View File

@@ -125,6 +125,8 @@ def infer(
hps, hps,
net_g, net_g,
device, device,
skip_start=False,
skip_end=False,
): ):
# 支持中日双语版本 # 支持中日双语版本
inferMap_V2 = { inferMap_V2 = {
@@ -172,6 +174,20 @@ def infer(
bert, ja_bert, en_bert, phones, tones, lang_ids = get_text( bert, ja_bert, en_bert, phones, tones, lang_ids = get_text(
text, language, hps, device text, language, hps, device
) )
if skip_start:
phones = phones[1:]
tones = tones[1:]
lang_ids = lang_ids[1:]
bert = bert[:, 1:]
ja_bert = ja_bert[:, 1:]
en_bert = en_bert[:, 1:]
if skip_end:
phones = phones[:-1]
tones = tones[:-1]
lang_ids = lang_ids[:-1]
bert = bert[:, :-1]
ja_bert = ja_bert[:, :-1]
en_bert = en_bert[:, :-1]
with torch.no_grad(): with torch.no_grad():
x_tst = phones.to(device).unsqueeze(0) x_tst = phones.to(device).unsqueeze(0)
tones = tones.to(device).unsqueeze(0) tones = tones.to(device).unsqueeze(0)

View File

@@ -40,11 +40,15 @@ def generate_audio(
length_scale, length_scale,
speaker, speaker,
language, language,
skip_start=False,
skip_end=False,
): ):
audio_list = [] audio_list = []
silence = np.zeros(hps.data.sampling_rate // 2, dtype=np.int16) # silence = np.zeros(hps.data.sampling_rate // 2, dtype=np.int16)
with torch.no_grad(): with torch.no_grad():
for piece in slices: for idx, piece in enumerate(slices):
skip_start = (idx != 0) and skip_start
skip_end = (idx != len(slices) - 1) and skip_end
audio = infer( audio = infer(
piece, piece,
sdp_ratio=sdp_ratio, sdp_ratio=sdp_ratio,
@@ -56,10 +60,12 @@ def generate_audio(
hps=hps, hps=hps,
net_g=net_g, net_g=net_g,
device=device, device=device,
skip_start=skip_start,
skip_end=skip_end,
) )
audio16bit = gr.processing_utils.convert_to_16_bit_wav(audio) audio16bit = gr.processing_utils.convert_to_16_bit_wav(audio)
audio_list.append(audio16bit) audio_list.append(audio16bit)
audio_list.append(silence) # 将静音添加到列表中 # audio_list.append(silence) # 将静音添加到列表中
return audio_list return audio_list
@@ -82,7 +88,9 @@ def tts_split(
para_list = re_matching.cut_para(text) para_list = re_matching.cut_para(text)
audio_list = [] audio_list = []
if not cut_by_sent: if not cut_by_sent:
for p in para_list: for idx, p in enumerate(para_list):
skip_start = idx != 0
skip_end = idx != len(para_list) - 1
audio = infer( audio = infer(
p, p,
sdp_ratio=sdp_ratio, sdp_ratio=sdp_ratio,
@@ -94,16 +102,22 @@ def tts_split(
hps=hps, hps=hps,
net_g=net_g, net_g=net_g,
device=device, device=device,
skip_start=skip_start,
skip_end=skip_end,
) )
audio16bit = gr.processing_utils.convert_to_16_bit_wav(audio) audio16bit = gr.processing_utils.convert_to_16_bit_wav(audio)
audio_list.append(audio16bit) audio_list.append(audio16bit)
silence = np.zeros((int)(44100 * interval_between_para), dtype=np.int16) silence = np.zeros((int)(44100 * interval_between_para), dtype=np.int16)
audio_list.append(silence) audio_list.append(silence)
else: else:
for p in para_list: for idx, p in enumerate(para_list):
skip_start = idx != 0
skip_end = idx != len(para_list) - 1
audio_list_sent = [] audio_list_sent = []
sent_list = re_matching.cut_sent(p) sent_list = re_matching.cut_sent(p)
for s in sent_list: for idx, s in enumerate(sent_list):
skip_start = (idx != 0) and skip_start
skip_end = (idx != len(sent_list) - 1) and skip_end
audio = infer( audio = infer(
s, s,
sdp_ratio=sdp_ratio, sdp_ratio=sdp_ratio,
@@ -115,6 +129,8 @@ def tts_split(
hps=hps, hps=hps,
net_g=net_g, net_g=net_g,
device=device, device=device,
skip_start=skip_start,
skip_end=skip_end,
) )
audio_list_sent.append(audio) audio_list_sent.append(audio)
silence = np.zeros((int)(44100 * interval_between_sent)) silence = np.zeros((int)(44100 * interval_between_sent))
@@ -150,9 +166,13 @@ def tts_fn(
np.concatenate([np.zeros(hps.data.sampling_rate // 2)]), np.concatenate([np.zeros(hps.data.sampling_rate // 2)]),
) )
result = re_matching.text_matching(text) result = re_matching.text_matching(text)
for one in result: for idx, one in enumerate(result):
skip_start = idx != 0
skip_end = idx != len(result) - 1
_speaker = one.pop() _speaker = one.pop()
for lang, content in one: for idx, (lang, content) in enumerate(one):
skip_start = (idx != 0) and skip_start
skip_end = (idx != len(one) - 1) and skip_end
audio_list.extend( audio_list.extend(
generate_audio( generate_audio(
content.split("|"), content.split("|"),
@@ -162,16 +182,22 @@ def tts_fn(
length_scale, length_scale,
_speaker, _speaker,
lang, lang,
skip_start,
skip_end,
) )
) )
elif language.lower() == "auto": elif language.lower() == "auto":
sentences_list = split_by_language(text, target_languages=["zh", "ja", "en"]) sentences_list = split_by_language(text, target_languages=["zh", "ja", "en"])
for sentences, lang in sentences_list: for idx, (sentences, lang) in enumerate(sentences_list):
skip_start = idx != 0
skip_end = idx != len(sentences_list) - 1
lang = lang.upper() lang = lang.upper()
if lang == "JA": if lang == "JA":
lang = "JP" lang = "JP"
sentences = sentence_split(sentences, max=250) sentences = sentence_split(sentences, max=250)
for content in sentences: for idx, content in enumerate(sentences):
skip_start = (idx != 0) and skip_start
skip_end = (idx != len(sentences) - 1) and skip_end
audio_list.extend( audio_list.extend(
generate_audio( generate_audio(
content.split("|"), content.split("|"),
@@ -181,6 +207,8 @@ def tts_fn(
length_scale, length_scale,
speaker, speaker,
lang, lang,
skip_start,
skip_end,
) )
) )
else: else: