CoolFace
Apppublic

kkvc-hf/Style-Bert-VITS2-AS2

sourceHugging Faceapache-2.0updated 11mo agoView on Hugging Face
1likes
server_fastapi.py343 linesDownload Raw Back to root
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