diff --git a/infer.py b/infer.py index 7c2f7d3..8497998 100644 --- a/infer.py +++ b/infer.py @@ -201,6 +201,5 @@ def infer( .float() .numpy() ) - del x_tst, tones, lang_ids, bert, x_tst_lengths, speakers - torch.cuda.empty_cache() + del x_tst, tones, lang_ids, bert, x_tst_lengths, speakers, ja_bert, en_bert return audio diff --git a/server_fastapi.py b/server_fastapi.py index 6c231c7..1f22860 100644 --- a/server_fastapi.py +++ b/server_fastapi.py @@ -5,6 +5,7 @@ import logging import gc import random +import numpy as np import utils from fastapi import FastAPI, Query, Request from fastapi.responses import Response, FileResponse @@ -23,6 +24,8 @@ from urllib.parse import unquote from infer import infer, get_net_g, latest_version import tools.translate as trans +from re_matching import cut_sent + from config import config @@ -183,6 +186,7 @@ if __name__ == "__main__": length: float = Query(1, description="语速"), language: str = Query(None, description="语言"), # 若不指定使用语言则使用默认值 auto_translate: bool = Query(False, description="自动翻译"), + auto_split: bool = Query(False, description="自动切分"), ): """语音接口""" logger.info( @@ -206,19 +210,41 @@ if __name__ == "__main__": language = loaded_models.models[model_id].language if auto_translate: text = trans.translate(Sentence=text, to_Language=language.lower()) - with torch.no_grad(): - audio = infer( - text=text, - sdp_ratio=sdp_ratio, - noise_scale=noise, - noise_scale_w=noisew, - length_scale=length, - sid=speaker_name, - language=language, - hps=loaded_models.models[model_id].hps, - net_g=loaded_models.models[model_id].net_g, - device=loaded_models.models[model_id].device, - ) + if not auto_split: + with torch.no_grad(): + audio = infer( + text=text, + sdp_ratio=sdp_ratio, + noise_scale=noise, + noise_scale_w=noisew, + length_scale=length, + sid=speaker_name, + language=language, + hps=loaded_models.models[model_id].hps, + net_g=loaded_models.models[model_id].net_g, + device=loaded_models.models[model_id].device, + ) + else: + texts = cut_sent(text) + audios = [] + with torch.no_grad(): + for t in texts: + audios.append( + infer( + text=t, + sdp_ratio=sdp_ratio, + noise_scale=noise, + noise_scale_w=noisew, + length_scale=length, + sid=speaker_name, + language=language, + hps=loaded_models.models[model_id].hps, + net_g=loaded_models.models[model_id].net_g, + device=loaded_models.models[model_id].device, + ) + ) + audios.append(np.zeros((int)(44100 * 0.3))) + audio = np.concatenate(audios) wavContent = BytesIO() wavfile.write( wavContent, loaded_models.models[model_id].hps.data.sampling_rate, audio @@ -308,7 +334,7 @@ if __name__ == "__main__": model_files = [] for sub_file in sub_files: relpath = os.path.realpath(os.path.join(sub_dir, sub_file)) - if only_unloaded and relpath in loaded_models.path2count.keys(): + if only_unloaded and relpath in loaded_models.path2ids.keys(): continue if sub_file.endswith(".pth") and sub_file.startswith("G_"): if os.path.isfile(relpath): @@ -327,7 +353,7 @@ if __name__ == "__main__": sub_files = os.listdir(models_dir) for sub_file in sub_files: relpath = os.path.realpath(os.path.join(models_dir, sub_file)) - if only_unloaded and relpath in loaded_models.path2count.keys(): + if only_unloaded and relpath in loaded_models.path2ids.keys(): continue if sub_file.endswith(".pth") and sub_file.startswith("G_"): if os.path.isfile(os.path.join(models_dir, sub_file)): diff --git a/tools/translate.py b/tools/translate.py index d2c5aa8..9368b5f 100644 --- a/tools/translate.py +++ b/tools/translate.py @@ -22,13 +22,13 @@ def translate(Sentence: str, to_Language: str = "jp", from_Language: str = ""): if appid == "" or key == "": return "请开发者在config.yml中配置app_key与secret_key" url = "https://fanyi-api.baidu.com/api/trans/vip/translate" - texts = Sentence.split("\n") + texts = Sentence.splitlines() outTexts = [] for t in texts: if t != "": # 签名计算 参考文档 https://api.fanyi.baidu.com/product/113 salt = str(random.randint(1, 100000)) - signString = appid + Sentence + salt + key + signString = appid + t + salt + key hs = hashlib.md5() hs.update(signString.encode("utf-8")) signString = hs.hexdigest() @@ -36,7 +36,7 @@ def translate(Sentence: str, to_Language: str = "jp", from_Language: str = ""): from_Language = "auto" headers = {"Content-Type": "application/x-www-form-urlencoded"} payload = { - "q": Sentence, + "q": t, "from": from_Language, "to": to_Language, "appid": appid,