Support multilang generation (#185)

* Support multilang generation

* Update

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
OedoSoldier
2023-11-15 20:22:48 +08:00
committed by GitHub
parent ad694ee050
commit 1fbddf4202
2 changed files with 213 additions and 22 deletions

View File

@@ -221,3 +221,93 @@ def infer(
if torch.cuda.is_available():
torch.cuda.empty_cache()
return audio
def infer_multilang(
text,
sdp_ratio,
noise_scale,
noise_scale_w,
length_scale,
sid,
language,
hps,
net_g,
device,
skip_start=False,
skip_end=False,
):
bert, ja_bert, en_bert, phones, tones, lang_ids = [], [], [], [], [], []
# bert, ja_bert, en_bert, phones, tones, lang_ids = get_text(
# text, language, hps, device
# )
for idx, (t, l) in enumerate(zip(text, language)):
skip_start = (idx != 0) or (skip_start and idx == 0)
skip_end = (idx != len(text) - 1) or (skip_end and idx == len(text) - 1)
(
temp_bert,
temp_ja_bert,
temp_en_bert,
temp_phones,
temp_tones,
temp_lang_ids,
) = get_text(t, l, hps, device)
if skip_start:
temp_bert = temp_bert[:, 1:]
temp_ja_bert = temp_ja_bert[:, 1:]
temp_en_bert = temp_en_bert[:, 1:]
temp_phones = temp_phones[1:]
temp_tones = temp_tones[1:]
temp_lang_ids = temp_lang_ids[1:]
if skip_end:
temp_bert = temp_bert[:, :-1]
temp_ja_bert = temp_ja_bert[:, :-1]
temp_en_bert = temp_en_bert[:, :-1]
temp_phones = temp_phones[:-1]
temp_tones = temp_tones[:-1]
temp_lang_ids = temp_lang_ids[:-1]
bert.append(temp_bert)
ja_bert.append(temp_ja_bert)
en_bert.append(temp_en_bert)
phones.append(temp_phones)
tones.append(temp_tones)
lang_ids.append(temp_lang_ids)
bert = torch.concatenate(bert, dim=1)
ja_bert = torch.concatenate(ja_bert, dim=1)
en_bert = torch.concatenate(en_bert, dim=1)
phones = torch.concatenate(phones, dim=0)
tones = torch.concatenate(tones, dim=0)
lang_ids = torch.concatenate(lang_ids, dim=0)
with torch.no_grad():
x_tst = phones.to(device).unsqueeze(0)
tones = tones.to(device).unsqueeze(0)
lang_ids = lang_ids.to(device).unsqueeze(0)
bert = bert.to(device).unsqueeze(0)
ja_bert = ja_bert.to(device).unsqueeze(0)
en_bert = en_bert.to(device).unsqueeze(0)
x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device)
del phones
speakers = torch.LongTensor([hps.data.spk2id[sid]]).to(device)
audio = (
net_g.infer(
x_tst,
x_tst_lengths,
speakers,
tones,
lang_ids,
bert,
ja_bert,
en_bert,
sdp_ratio=sdp_ratio,
noise_scale=noise_scale,
noise_scale_w=noise_scale_w,
length_scale=length_scale,
)[0][0, 0]
.data.cpu()
.float()
.numpy()
)
del x_tst, tones, lang_ids, bert, x_tst_lengths, speakers, ja_bert, en_bert
if torch.cuda.is_available():
torch.cuda.empty_cache()
return audio

145
webui.py
View File

@@ -3,7 +3,7 @@ import os
import logging
import re_matching
from tools.sentence import split_by_language, sentence_split
from tools.sentence import split_by_language
logging.getLogger("numba").setLevel(logging.WARNING)
logging.getLogger("markdown_it").setLevel(logging.WARNING)
@@ -18,7 +18,7 @@ logger = logging.getLogger(__name__)
import torch
import utils
from infer import infer, latest_version, get_net_g
from infer import infer, latest_version, get_net_g, infer_multilang
import gradio as gr
import webbrowser
import numpy as np
@@ -69,6 +69,43 @@ def generate_audio(
return audio_list
def generate_audio_multilang(
slices,
sdp_ratio,
noise_scale,
noise_scale_w,
length_scale,
speaker,
language,
skip_start=False,
skip_end=False,
):
audio_list = []
# silence = np.zeros(hps.data.sampling_rate // 2, dtype=np.int16)
with torch.no_grad():
for idx, piece in enumerate(slices):
skip_start = (idx != 0) and skip_start
skip_end = (idx != len(slices) - 1) and skip_end
audio = infer_multilang(
piece,
sdp_ratio=sdp_ratio,
noise_scale=noise_scale,
noise_scale_w=noise_scale_w,
length_scale=length_scale,
sid=speaker,
language=language[idx],
hps=hps,
net_g=net_g,
device=device,
skip_start=skip_start,
skip_end=skip_end,
)
audio16bit = gr.processing_utils.convert_to_16_bit_wav(audio)
audio_list.append(audio16bit)
# audio_list.append(silence) # 将静音添加到列表中
return audio_list
def tts_split(
text: str,
speaker,
@@ -165,52 +202,116 @@ def tts_fn(
hps.data.sampling_rate,
np.concatenate([np.zeros(hps.data.sampling_rate // 2)]),
)
result = re_matching.text_matching(text)
for idx, one in enumerate(result):
skip_start = idx != 0
skip_end = idx != len(result) - 1
result = []
for slice in re_matching.text_matching(text):
_speaker = slice.pop()
temp_contant = []
temp_lang = []
for lang, content in slice:
if "|" in content:
temp = []
temp_ = []
for i in content.split("|"):
if i != "":
temp.append([i])
temp_.append([lang])
else:
temp.append([])
temp_.append([])
temp_contant += temp
temp_lang += temp_
else:
if len(temp_contant) == 0:
temp_contant.append([])
temp_lang.append([])
temp_contant[-1].append(content)
temp_lang[-1].append(lang)
for i, j in zip(temp_lang, temp_contant):
result.append([*zip(i, j), _speaker])
for i, one in enumerate(result):
skip_start = i != 0
skip_end = i != len(result) - 1
_speaker = one.pop()
for idx, (lang, content) in enumerate(one):
idx = 0
while idx < len(one):
text_to_generate = []
lang_to_generate = []
while True:
lang, content = one[idx]
temp_text = [content]
if len(text_to_generate) > 0:
text_to_generate[-1] += [temp_text.pop(0)]
lang_to_generate[-1] += [lang]
if len(temp_text) > 0:
text_to_generate += [[i] for i in temp_text]
lang_to_generate += [[lang]] * len(temp_text)
if idx + 1 < len(one):
idx += 1
else:
break
skip_start = (idx != 0) and skip_start
skip_end = (idx != len(one) - 1) and skip_end
print(text_to_generate, lang_to_generate)
audio_list.extend(
generate_audio(
content.split("|"),
generate_audio_multilang(
text_to_generate,
sdp_ratio,
noise_scale,
noise_scale_w,
length_scale,
_speaker,
lang,
lang_to_generate,
skip_start,
skip_end,
)
)
idx += 1
elif language.lower() == "auto":
sentences_list = split_by_language(text, target_languages=["zh", "ja", "en"])
for idx, (sentences, lang) in enumerate(sentences_list):
for idx, slice in enumerate(text.split("|")):
if slice == "":
continue
skip_start = idx != 0
skip_end = idx != len(sentences_list) - 1
lang = lang.upper()
if lang == "JA":
lang = "JP"
sentences = sentence_split(sentences, max=250)
for idx, content in enumerate(sentences):
skip_end = idx != len(text.split("|")) - 1
sentences_list = split_by_language(
slice, target_languages=["zh", "ja", "en"]
)
idx = 0
while idx < len(sentences_list):
text_to_generate = []
lang_to_generate = []
while True:
content, lang = sentences_list[idx]
temp_text = [content]
lang = lang.upper()
if lang == "JA":
lang = "JP"
if len(text_to_generate) > 0:
text_to_generate[-1] += [temp_text.pop(0)]
lang_to_generate[-1] += [lang]
if len(temp_text) > 0:
text_to_generate += [[i] for i in temp_text]
lang_to_generate += [[lang]] * len(temp_text)
if idx + 1 < len(sentences_list):
idx += 1
else:
break
skip_start = (idx != 0) and skip_start
skip_end = (idx != len(sentences) - 1) and skip_end
skip_end = (idx != len(sentences_list) - 1) and skip_end
print(text_to_generate, lang_to_generate)
audio_list.extend(
generate_audio(
content.split("|"),
generate_audio_multilang(
text_to_generate,
sdp_ratio,
noise_scale,
noise_scale_w,
length_scale,
speaker,
lang,
lang_to_generate,
skip_start,
skip_end,
)
)
idx += 1
else:
audio_list.extend(
generate_audio(