Fix: server_fastapi.py, support auto split (#152)

* Fix: server_fastapi.py, support auto split

* Fix: translate.py
This commit is contained in:
Sora
2023-11-06 19:40:58 +08:00
committed by GitHub
parent 855774213e
commit 096d661a4b
3 changed files with 45 additions and 20 deletions

View File

@@ -201,6 +201,5 @@ def infer(
.float() .float()
.numpy() .numpy()
) )
del x_tst, tones, lang_ids, bert, x_tst_lengths, speakers del x_tst, tones, lang_ids, bert, x_tst_lengths, speakers, ja_bert, en_bert
torch.cuda.empty_cache()
return audio return audio

View File

@@ -5,6 +5,7 @@ import logging
import gc import gc
import random import random
import numpy as np
import utils import utils
from fastapi import FastAPI, Query, Request from fastapi import FastAPI, Query, Request
from fastapi.responses import Response, FileResponse from fastapi.responses import Response, FileResponse
@@ -23,6 +24,8 @@ from urllib.parse import unquote
from infer import infer, get_net_g, latest_version from infer import infer, get_net_g, latest_version
import tools.translate as trans import tools.translate as trans
from re_matching import cut_sent
from config import config from config import config
@@ -183,6 +186,7 @@ if __name__ == "__main__":
length: float = Query(1, description="语速"), length: float = Query(1, description="语速"),
language: str = Query(None, description="语言"), # 若不指定使用语言则使用默认值 language: str = Query(None, description="语言"), # 若不指定使用语言则使用默认值
auto_translate: bool = Query(False, description="自动翻译"), auto_translate: bool = Query(False, description="自动翻译"),
auto_split: bool = Query(False, description="自动切分"),
): ):
"""语音接口""" """语音接口"""
logger.info( logger.info(
@@ -206,6 +210,7 @@ if __name__ == "__main__":
language = loaded_models.models[model_id].language language = loaded_models.models[model_id].language
if auto_translate: if auto_translate:
text = trans.translate(Sentence=text, to_Language=language.lower()) text = trans.translate(Sentence=text, to_Language=language.lower())
if not auto_split:
with torch.no_grad(): with torch.no_grad():
audio = infer( audio = infer(
text=text, text=text,
@@ -219,6 +224,27 @@ if __name__ == "__main__":
net_g=loaded_models.models[model_id].net_g, net_g=loaded_models.models[model_id].net_g,
device=loaded_models.models[model_id].device, 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() wavContent = BytesIO()
wavfile.write( wavfile.write(
wavContent, loaded_models.models[model_id].hps.data.sampling_rate, audio wavContent, loaded_models.models[model_id].hps.data.sampling_rate, audio
@@ -308,7 +334,7 @@ if __name__ == "__main__":
model_files = [] model_files = []
for sub_file in sub_files: for sub_file in sub_files:
relpath = os.path.realpath(os.path.join(sub_dir, sub_file)) 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 continue
if sub_file.endswith(".pth") and sub_file.startswith("G_"): if sub_file.endswith(".pth") and sub_file.startswith("G_"):
if os.path.isfile(relpath): if os.path.isfile(relpath):
@@ -327,7 +353,7 @@ if __name__ == "__main__":
sub_files = os.listdir(models_dir) sub_files = os.listdir(models_dir)
for sub_file in sub_files: for sub_file in sub_files:
relpath = os.path.realpath(os.path.join(models_dir, sub_file)) 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 continue
if sub_file.endswith(".pth") and sub_file.startswith("G_"): if sub_file.endswith(".pth") and sub_file.startswith("G_"):
if os.path.isfile(os.path.join(models_dir, sub_file)): if os.path.isfile(os.path.join(models_dir, sub_file)):

View File

@@ -22,13 +22,13 @@ def translate(Sentence: str, to_Language: str = "jp", from_Language: str = ""):
if appid == "" or key == "": if appid == "" or key == "":
return "请开发者在config.yml中配置app_key与secret_key" return "请开发者在config.yml中配置app_key与secret_key"
url = "https://fanyi-api.baidu.com/api/trans/vip/translate" url = "https://fanyi-api.baidu.com/api/trans/vip/translate"
texts = Sentence.split("\n") texts = Sentence.splitlines()
outTexts = [] outTexts = []
for t in texts: for t in texts:
if t != "": if t != "":
# 签名计算 参考文档 https://api.fanyi.baidu.com/product/113 # 签名计算 参考文档 https://api.fanyi.baidu.com/product/113
salt = str(random.randint(1, 100000)) salt = str(random.randint(1, 100000))
signString = appid + Sentence + salt + key signString = appid + t + salt + key
hs = hashlib.md5() hs = hashlib.md5()
hs.update(signString.encode("utf-8")) hs.update(signString.encode("utf-8"))
signString = hs.hexdigest() signString = hs.hexdigest()
@@ -36,7 +36,7 @@ def translate(Sentence: str, to_Language: str = "jp", from_Language: str = ""):
from_Language = "auto" from_Language = "auto"
headers = {"Content-Type": "application/x-www-form-urlencoded"} headers = {"Content-Type": "application/x-www-form-urlencoded"}
payload = { payload = {
"q": Sentence, "q": t,
"from": from_Language, "from": from_Language,
"to": to_Language, "to": to_Language,
"appid": appid, "appid": appid,