kkvc-hf/Style-Bert-VITS2-AS2
1
1"""2API server for TTS3TODO: server_editor.pyと統合する?4"""5 6import argparse7import os8import sys9from io import BytesIO10from pathlib import Path11from typing import Any, Optional12from urllib.parse import unquote13 14import GPUtil15import psutil16import torch17import uvicorn18from fastapi import FastAPI, HTTPException, Query, Request, status19from fastapi.middleware.cors import CORSMiddleware20from fastapi.responses import FileResponse, Response21from scipy.io import wavfile22 23from config import get_config24from style_bert_vits2.constants import (25 DEFAULT_ASSIST_TEXT_WEIGHT,26 DEFAULT_LENGTH,27 DEFAULT_LINE_SPLIT,28 DEFAULT_NOISE,29 DEFAULT_NOISEW,30 DEFAULT_SDP_RATIO,31 DEFAULT_SPLIT_INTERVAL,32 DEFAULT_STYLE,33 DEFAULT_STYLE_WEIGHT,34 Languages,35)36from style_bert_vits2.logging import logger37from style_bert_vits2.nlp import bert_models38from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk39from style_bert_vits2.nlp.japanese.user_dict import update_dict40from style_bert_vits2.tts_model import TTSModel, TTSModelHolder41 42 43config = get_config()44ln = config.server_config.language45 46 47# pyopenjtalk_worker を起動48## pyopenjtalk_worker は TCP ソケットサーバーのため、ここで起動する49pyopenjtalk.initialize_worker()50 51# dict_data/ 以下の辞書データを pyopenjtalk に適用52update_dict()53 54# 事前に BERT モデル/トークナイザーをロードしておく55## ここでロードしなくても必要になった際に自動ロードされるが、時間がかかるため事前にロードしておいた方が体験が良い56bert_models.load_model(Languages.JP)57bert_models.load_tokenizer(Languages.JP)58bert_models.load_model(Languages.EN)59bert_models.load_tokenizer(Languages.EN)60bert_models.load_model(Languages.ZH)61bert_models.load_tokenizer(Languages.ZH)62 63 64def raise_validation_error(msg: str, param: str):65 logger.warning(f"Validation error: {msg}")66 raise HTTPException(67 status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,68 detail=[dict(type="invalid_params", msg=msg, loc=["query", param])],69 )70 71 72class AudioResponse(Response):73 media_type = "audio/wav"74 75 76loaded_models: list[TTSModel] = []77 78 79def load_models(model_holder: TTSModelHolder):80 global loaded_models81 loaded_models = []82 for model_name, model_paths in model_holder.model_files_dict.items():83 model = TTSModel(84 model_path=model_paths[0],85 config_path=model_holder.root_dir / model_name / "config.json",86 style_vec_path=model_holder.root_dir / model_name / "style_vectors.npy",87 device=model_holder.device,88 )89 # 起動時に全てのモデルを読み込むのは時間がかかりメモリを食うのでやめる90 # model.load()91 loaded_models.append(model)92 93 94if __name__ == "__main__":95 parser = argparse.ArgumentParser()96 parser.add_argument("--cpu", action="store_true", help="Use CPU instead of GPU")97 parser.add_argument(98 "--dir", "-d", type=str, help="Model directory", default=config.assets_root99 )100 args = parser.parse_args()101 102 if args.cpu:103 device = "cpu"104 else:105 device = "cuda" if torch.cuda.is_available() else "cpu"106 107 model_dir = Path(args.dir)108 model_holder = TTSModelHolder(model_dir, device)109 if len(model_holder.model_names) == 0:110 logger.error(f"Models not found in {model_dir}.")111 sys.exit(1)112 113 logger.info("Loading models...")114 load_models(model_holder)115 116 limit = config.server_config.limit117 if limit < 1:118 limit = None119 else:120 logger.info(121 f"The maximum length of the text is {limit}. If you want to change it, modify config.yml. Set limit to -1 to remove the limit."122 )123 app = FastAPI()124 allow_origins = config.server_config.origins125 if allow_origins:126 logger.warning(127 f"CORS allow_origins={config.server_config.origins}. If you don't want, modify config.yml"128 )129 app.add_middleware(130 CORSMiddleware,131 allow_origins=config.server_config.origins,132 allow_credentials=True,133 allow_methods=["*"],134 allow_headers=["*"],135 )136 # app.logger = logger137 # ↑効いていなさそう。loggerをどうやって上書きするかはよく分からなかった。138 139 @app.api_route("/voice", methods=["GET", "POST"], response_class=AudioResponse)140 async def voice(141 request: Request,142 text: str = Query(..., min_length=1, max_length=limit, description="セリフ"),143 encoding: str = Query(None, description="textをURLデコードする(ex, `utf-8`)"),144 model_name: str = Query(145 None,146 description="モデル名(model_idより優先)。model_assets内のディレクトリ名を指定",147 ),148 model_id: int = Query(149 0, description="モデルID。`GET /models/info`のkeyの値を指定ください"150 ),151 speaker_name: str = Query(152 None,153 description="話者名(speaker_idより優先)。esd.listの2列目の文字列を指定",154 ),155 speaker_id: int = Query(156 0, description="話者ID。model_assets>[model]>config.json内のspk2idを確認"157 ),158 sdp_ratio: float = Query(159 DEFAULT_SDP_RATIO,160 description="SDP(Stochastic Duration Predictor)/DP混合比。比率が高くなるほどトーンのばらつきが大きくなる",161 ),162 noise: float = Query(163 DEFAULT_NOISE,164 description="サンプルノイズの割合。大きくするほどランダム性が高まる",165 ),166 noisew: float = Query(167 DEFAULT_NOISEW,168 description="SDPノイズ。大きくするほど発音の間隔にばらつきが出やすくなる",169 ),170 length: float = Query(171 DEFAULT_LENGTH,172 description="話速。基準は1で大きくするほど音声は長くなり読み上げが遅まる",173 ),174 language: Languages = Query(ln, description="textの言語"),175 auto_split: bool = Query(DEFAULT_LINE_SPLIT, description="改行で分けて生成"),176 split_interval: float = Query(177 DEFAULT_SPLIT_INTERVAL, description="分けた場合に挟む無音の長さ(秒)"178 ),179 assist_text: Optional[str] = Query(180 None,181 description="このテキストの読み上げと似た声音・感情になりやすくなる。ただし抑揚やテンポ等が犠牲になる傾向がある",182 ),183 assist_text_weight: float = Query(184 DEFAULT_ASSIST_TEXT_WEIGHT, description="assist_textの強さ"185 ),186 style: Optional[str] = Query(DEFAULT_STYLE, description="スタイル"),187 style_weight: float = Query(DEFAULT_STYLE_WEIGHT, description="スタイルの強さ"),188 reference_audio_path: Optional[str] = Query(189 None, description="スタイルを音声ファイルで行う"190 ),191 ):192 """Infer text to speech(テキストから感情付き音声を生成する)"""193 logger.info(194 f"{request.client.host}:{request.client.port}/voice { unquote(str(request.query_params) )}"195 )196 if request.method == "GET":197 logger.warning(198 "The GET method is not recommended for this endpoint due to various restrictions. Please use the POST method."199 )200 if model_id >= len(201 model_holder.model_names202 ): # /models/refresh があるためQuery(le)で表現不可203 raise_validation_error(f"model_id={model_id} not found", "model_id")204 205 if model_name:206 # load_models() の 処理内容が i の正当性を担保していることに注意207 model_ids = [i for i, x in enumerate(model_holder.models_info) if x.name == model_name]208 if not model_ids:209 raise_validation_error(210 f"model_name={model_name} not found", "model_name"211 )212 # 今の実装ではディレクトリ名が重複することは無いはずだが...213 if len(model_ids) > 1:214 raise_validation_error(215 f"model_name={model_name} is ambiguous", "model_name"216 )217 model_id = model_ids[0]218 219 model = loaded_models[model_id]220 if speaker_name is None:221 if speaker_id not in model.id2spk.keys():222 raise_validation_error(223 f"speaker_id={speaker_id} not found", "speaker_id"224 )225 else:226 if speaker_name not in model.spk2id.keys():227 raise_validation_error(228 f"speaker_name={speaker_name} not found", "speaker_name"229 )230 speaker_id = model.spk2id[speaker_name]231 if style not in model.style2id.keys():232 raise_validation_error(f"style={style} not found", "style")233 assert style is not None234 if encoding is not None:235 text = unquote(text, encoding=encoding)236 sr, audio = model.infer(237 text=text,238 language=language,239 speaker_id=speaker_id,240 reference_audio_path=reference_audio_path,241 sdp_ratio=sdp_ratio,242 noise=noise,243 noise_w=noisew,244 length=length,245 line_split=auto_split,246 split_interval=split_interval,247 assist_text=assist_text,248 assist_text_weight=assist_text_weight,249 use_assist_text=bool(assist_text),250 style=style,251 style_weight=style_weight,252 )253 logger.success("Audio data generated and sent successfully")254 with BytesIO() as wavContent:255 wavfile.write(wavContent, sr, audio)256 return Response(content=wavContent.getvalue(), media_type="audio/wav")257 258 @app.post("/g2p")259 def g2p(text: str):260 return g2kata_tone(normalize_text(text))261 262 @app.get("/models/info")263 def get_loaded_models_info():264 """ロードされたモデル情報の取得"""265 266 result: dict[str, dict[str, Any]] = dict()267 for model_id, model in enumerate(loaded_models):268 result[str(model_id)] = {269 "config_path": model.config_path,270 "model_path": model.model_path,271 "device": model.device,272 "spk2id": model.spk2id,273 "id2spk": model.id2spk,274 "style2id": model.style2id,275 }276 return result277 278 @app.post("/models/refresh")279 def refresh():280 """モデルをパスに追加/削除した際などに読み込ませる"""281 model_holder.refresh()282 load_models(model_holder)283 return get_loaded_models_info()284 285 @app.get("/status")286 def get_status():287 """実行環境のステータスを取得"""288 cpu_percent = psutil.cpu_percent(interval=1)289 memory_info = psutil.virtual_memory()290 memory_total = memory_info.total291 memory_available = memory_info.available292 memory_used = memory_info.used293 memory_percent = memory_info.percent294 gpuInfo = []295 devices = ["cpu"]296 for i in range(torch.cuda.device_count()):297 devices.append(f"cuda:{i}")298 gpus = GPUtil.getGPUs()299 for gpu in gpus:300 gpuInfo.append(301 {302 "gpu_id": gpu.id,303 "gpu_load": gpu.load,304 "gpu_memory": {305 "total": gpu.memoryTotal,306 "used": gpu.memoryUsed,307 "free": gpu.memoryFree,308 },309 }310 )311 return {312 "devices": devices,313 "cpu_percent": cpu_percent,314 "memory_total": memory_total,315 "memory_available": memory_available,316 "memory_used": memory_used,317 "memory_percent": memory_percent,318 "gpu": gpuInfo,319 }320 321 @app.get("/tools/get_audio", response_class=AudioResponse)322 def get_audio(323 request: Request, path: str = Query(..., description="local wav path")324 ):325 """wavデータを取得する"""326 logger.info(327 f"{request.client.host}:{request.client.port}/tools/get_audio { unquote(str(request.query_params) )}"328 )329 if not os.path.isfile(path):330 raise_validation_error(f"path={path} not found", "path")331 if not path.lower().endswith(".wav"):332 raise_validation_error(f"wav file not found in {path}", "path")333 return FileResponse(path=path, media_type="audio/wav")334 335 logger.info(f"server listen: http://127.0.0.1:{config.server_config.port}")336 logger.info(f"API docs: http://127.0.0.1:{config.server_config.port}/docs")337 logger.info(338 f"Input text length limit: {limit}. You can change it in server.limit in config.yml"339 )340 uvicorn.run(341 app, port=config.server_config.port, host="0.0.0.0", log_level="warning"342 )343 