Fix: server_fastapi.py, support auto split (#152)
* Fix: server_fastapi.py, support auto split * Fix: translate.py
This commit is contained in:
3
infer.py
3
infer.py
@@ -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
|
||||||
|
|||||||
@@ -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)):
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
Reference in New Issue
Block a user