diff --git a/infer.py b/infer.py index eec7698..0f3eb9f 100644 --- a/infer.py +++ b/infer.py @@ -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 diff --git a/webui.py b/webui.py index 32140da..bc80690 100644 --- a/webui.py +++ b/webui.py @@ -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(