From 3623eed89637407de784a19501e694e2113ff8b2 Mon Sep 17 00:00:00 2001 From: Sora <654163754@qq.com> Date: Sun, 3 Dec 2023 22:41:36 +0800 Subject: [PATCH 1/2] =?UTF-8?q?update=20server=5Ffastapi.py:=20=E6=B7=BB?= =?UTF-8?q?=E5=8A=A02.1=E6=A8=A1=E5=9E=8B=E6=8E=A8=E7=90=86=E6=94=AF?= =?UTF-8?q?=E6=8C=81=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- get_emo.py | 2 +- server_fastapi.py | 88 +++++++++++++++++++++++++++-------------------- 2 files changed, 51 insertions(+), 39 deletions(-) diff --git a/get_emo.py b/get_emo.py index d095980..054ebde 100644 --- a/get_emo.py +++ b/get_emo.py @@ -17,7 +17,7 @@ def get_emo(path): wav, sr = librosa.load(path, 16000) device = config.bert_gen_config.device return process_func( - np.expand_dims(wav, 0).astype(np.float), + np.expand_dims(wav, 0).astype(np.float64), sr, model, processor, diff --git a/server_fastapi.py b/server_fastapi.py index 571a4b8..b7c1a52 100644 --- a/server_fastapi.py +++ b/server_fastapi.py @@ -9,7 +9,7 @@ from pydantic import BaseModel import gradio import numpy as np import utils -from fastapi import FastAPI, Query, Request +from fastapi import FastAPI, Query, Request, File, UploadFile, Form from fastapi.responses import Response, FileResponse from fastapi.staticfiles import StaticFiles from io import BytesIO @@ -180,10 +180,7 @@ if __name__ == "__main__": async def index(): return FileResponse("./Web/index.html") - class Text(BaseModel): - text: str - - def _voice( + async def _voice( text: str, model_id: int, speaker_name: str, @@ -195,7 +192,10 @@ if __name__ == "__main__": language: str, auto_translate: bool, auto_split: bool, + emotion: Optional[int] = None, + reference_audio=None, ) -> Union[Response, Dict[str, any]]: + """TTS实现函数""" # 检查模型是否存在 if model_id not in loaded_models.models.keys(): return {"status": 10, "detail": f"模型model_id={model_id}未加载"} @@ -214,28 +214,12 @@ if __name__ == "__main__": language = loaded_models.models[model_id].language if auto_translate: text = trans.translate(Sentence=text, to_Language=language.lower()) - 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, + if reference_audio is not None: + with BytesIO(await reference_audio.read()) as ref_audio: + if not auto_split: + with torch.no_grad(): + audio = infer( + text=text, sdp_ratio=sdp_ratio, noise_scale=noise, noise_scale_w=noisew, @@ -245,11 +229,34 @@ if __name__ == "__main__": hps=loaded_models.models[model_id].hps, net_g=loaded_models.models[model_id].net_g, device=loaded_models.models[model_id].device, + emotion=emotion, + reference_audio=ref_audio, ) - ) - audios.append(np.zeros(int(44100 * 0.2))) - audio = np.concatenate(audios) - audio = gradio.processing_utils.convert_to_16_bit_wav(audio) + audio = gradio.processing_utils.convert_to_16_bit_wav(audio) + 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, + emotion=emotion, + reference_audio=ref_audio, + ) + ) + audios.append(np.zeros(int(44100 * 0.2))) + audio = np.concatenate(audios) + audio = gradio.processing_utils.convert_to_16_bit_wav(audio) with BytesIO() as wavContent: wavfile.write( wavContent, loaded_models.models[model_id].hps.data.sampling_rate, audio @@ -258,9 +265,9 @@ if __name__ == "__main__": return response @app.post("/voice") - def voice( + async def voice( request: Request, # fastapi自动注入 - text: Text, + text: str = Form(...), model_id: int = Query(..., description="模型ID"), # 模型序号 speaker_name: str = Query( None, description="说话人名" @@ -273,13 +280,14 @@ if __name__ == "__main__": language: str = Query(None, description="语言"), # 若不指定使用语言则使用默认值 auto_translate: bool = Query(False, description="自动翻译"), auto_split: bool = Query(False, description="自动切分"), + emotion: Optional[int] = Query(None, description="emo"), + reference_audio: UploadFile = File(None), ): - """语音接口""" - text = text.text + """语音接口,若需要上传参考音频请仅使用post请求""" logger.info( f"{request.client.host}:{request.client.port}/voice { unquote(str(request.query_params) )} text={text}" ) - return _voice( + return await _voice( text=text, model_id=model_id, speaker_name=speaker_name, @@ -291,10 +299,12 @@ if __name__ == "__main__": language=language, auto_translate=auto_translate, auto_split=auto_split, + emotion=emotion, + reference_audio=reference_audio, ) @app.get("/voice") - def voice( + async def voice( request: Request, # fastapi自动注入 text: str = Query(..., description="输入文字"), model_id: int = Query(..., description="模型ID"), # 模型序号 @@ -309,12 +319,13 @@ if __name__ == "__main__": language: str = Query(None, description="语言"), # 若不指定使用语言则使用默认值 auto_translate: bool = Query(False, description="自动翻译"), auto_split: bool = Query(False, description="自动切分"), + emotion: Optional[int] = Query(None, description="emo"), ): """语音接口""" logger.info( f"{request.client.host}:{request.client.port}/voice { unquote(str(request.query_params) )}" ) - return _voice( + return await _voice( text=text, model_id=model_id, speaker_name=speaker_name, @@ -326,6 +337,7 @@ if __name__ == "__main__": language=language, auto_translate=auto_translate, auto_split=auto_split, + emotion=emotion, ) @app.get("/models/info") From 872e8cb0e54e545162122c810995e15c3592c69a Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Sun, 3 Dec 2023 14:44:08 +0000 Subject: [PATCH 2/2] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- server_fastapi.py | 1 - 1 file changed, 1 deletion(-) diff --git a/server_fastapi.py b/server_fastapi.py index b7c1a52..e858f90 100644 --- a/server_fastapi.py +++ b/server_fastapi.py @@ -5,7 +5,6 @@ import logging import gc import random -from pydantic import BaseModel import gradio import numpy as np import utils