CoolFace
Apppublic

Julius8888/XJP_Voice

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
server_fastapi.py681 linesDownload Raw Back to root
1"""2api服务 多版本多模型 fastapi实现3"""4import logging5import gc6import random7 8import librosa9import gradio10import numpy as np11import utils12from fastapi import FastAPI, Query, Request, File, UploadFile, Form13from fastapi.responses import Response, FileResponse14from fastapi.staticfiles import StaticFiles15from io import BytesIO16from scipy.io import wavfile17import uvicorn18import torch19import webbrowser20import psutil21import GPUtil22from typing import Dict, Optional, List, Set, Union23import os24from tools.log import logger25from urllib.parse import unquote26 27from infer import infer, get_net_g, latest_version28import tools.translate as trans29from re_matching import cut_sent30 31 32from config import config33 34os.environ["TOKENIZERS_PARALLELISM"] = "false"35 36 37class Model:38    """模型封装类"""39 40    def __init__(self, config_path: str, model_path: str, device: str, language: str):41        self.config_path: str = os.path.normpath(config_path)42        self.model_path: str = os.path.normpath(model_path)43        self.device: str = device44        self.language: str = language45        self.hps = utils.get_hparams_from_file(config_path)46        self.spk2id: Dict[str, int] = self.hps.data.spk2id  # spk - id 映射字典47        self.id2spk: Dict[int, str] = dict()  # id - spk 映射字典48        for speaker, speaker_id in self.hps.data.spk2id.items():49            self.id2spk[speaker_id] = speaker50        self.version: str = (51            self.hps.version if hasattr(self.hps, "version") else latest_version52        )53        self.net_g = get_net_g(54            model_path=model_path,55            version=self.version,56            device=device,57            hps=self.hps,58        )59 60    def to_dict(self) -> Dict[str, any]:61        return {62            "config_path": self.config_path,63            "model_path": self.model_path,64            "device": self.device,65            "language": self.language,66            "spk2id": self.spk2id,67            "id2spk": self.id2spk,68            "version": self.version,69        }70 71 72class Models:73    def __init__(self):74        self.models: Dict[int, Model] = dict()75        self.num = 076        # spkInfo[角色名][模型id] = 角色id77        self.spk_info: Dict[str, Dict[int, int]] = dict()78        self.path2ids: Dict[str, Set[int]] = dict()  # 路径指向的model的id79 80    def init_model(81        self, config_path: str, model_path: str, device: str, language: str82    ) -> int:83        """84        初始化并添加一个模型85 86        :param config_path: 模型config.json路径87        :param model_path: 模型路径88        :param device: 模型推理使用设备89        :param language: 模型推理默认语言90        """91        # 若文件不存在则不进行加载92        if not os.path.isfile(model_path):93            if model_path != "":94                logger.warning(f"模型文件{model_path} 不存在,不进行初始化")95            return self.num96        if not os.path.isfile(config_path):97            if config_path != "":98                logger.warning(f"配置文件{config_path} 不存在,不进行初始化")99            return self.num100 101        # 若路径中的模型已存在,则不添加模型,若不存在,则进行初始化。102        model_path = os.path.realpath(model_path)103        if model_path not in self.path2ids.keys():104            self.path2ids[model_path] = {self.num}105            self.models[self.num] = Model(106                config_path=config_path,107                model_path=model_path,108                device=device,109                language=language,110            )111            logger.success(f"添加模型{model_path},使用配置文件{os.path.realpath(config_path)}")112        else:113            # 获取一个指向id114            m_id = next(iter(self.path2ids[model_path]))115            self.models[self.num] = self.models[m_id]116            self.path2ids[model_path].add(self.num)117            logger.success("模型已存在,添加模型引用。")118        # 添加角色信息119        for speaker, speaker_id in self.models[self.num].spk2id.items():120            if speaker not in self.spk_info.keys():121                self.spk_info[speaker] = {self.num: speaker_id}122            else:123                self.spk_info[speaker][self.num] = speaker_id124        # 修改计数125        self.num += 1126        return self.num - 1127 128    def del_model(self, index: int) -> Optional[int]:129        """删除对应序号的模型,若不存在则返回None"""130        if index not in self.models.keys():131            return None132        # 删除角色信息133        for speaker, speaker_id in self.models[index].spk2id.items():134            self.spk_info[speaker].pop(index)135            if len(self.spk_info[speaker]) == 0:136                # 若对应角色的所有模型都被删除,则清除该角色信息137                self.spk_info.pop(speaker)138        # 删除路径信息139        model_path = os.path.realpath(self.models[index].model_path)140        self.path2ids[model_path].remove(index)141        if len(self.path2ids[model_path]) == 0:142            self.path2ids.pop(model_path)143            logger.success(f"删除模型{model_path}, id = {index}")144        else:145            logger.success(f"删除模型引用{model_path}, id = {index}")146        # 删除模型147        self.models.pop(index)148        gc.collect()149        if torch.cuda.is_available():150            torch.cuda.empty_cache()151        return index152 153    def get_models(self):154        """获取所有模型"""155        return self.models156 157 158if __name__ == "__main__":159    app = FastAPI()160    app.logger = logger161    # 挂载静态文件162    logger.info("开始挂载网页页面")163    StaticDir: str = "./Web"164    if not os.path.isdir(StaticDir):165        logger.warning(166            "缺少网页资源,无法开启网页页面,如有需要请在 https://github.com/jiangyuxiaoxiao/Bert-VITS2-UI 或者Bert-VITS对应版本的release页面下载"167        )168    else:169        dirs = [fir.name for fir in os.scandir(StaticDir) if fir.is_dir()]170        files = [fir.name for fir in os.scandir(StaticDir) if fir.is_dir()]171        for dirName in dirs:172            app.mount(173                f"/{dirName}",174                StaticFiles(directory=f"./{StaticDir}/{dirName}"),175                name=dirName,176            )177    loaded_models = Models()178    # 加载模型179    logger.info("开始加载模型")180    models_info = config.server_config.models181    for model_info in models_info:182        loaded_models.init_model(183            config_path=model_info["config"],184            model_path=model_info["model"],185            device=model_info["device"],186            language=model_info["language"],187        )188 189    @app.get("/")190    async def index():191        return FileResponse("./Web/index.html")192 193    async def _voice(194        text: str,195        model_id: int,196        speaker_name: str,197        speaker_id: int,198        sdp_ratio: float,199        noise: float,200        noisew: float,201        length: float,202        language: str,203        auto_translate: bool,204        auto_split: bool,205        emotion: Optional[Union[int, str]] = None,206        reference_audio=None,207        style_text: Optional[str] = None,208        style_weight: float = 0.7,209    ) -> Union[Response, Dict[str, any]]:210        """TTS实现函数"""211        # 检查模型是否存在212        if model_id not in loaded_models.models.keys():213            logger.error(f"/voice 请求错误:模型model_id={model_id}未加载")214            return {"status": 10, "detail": f"模型model_id={model_id}未加载"}215        # 检查是否提供speaker216        if speaker_name is None and speaker_id is None:217            logger.error("/voice 请求错误:推理请求未提供speaker_name或speaker_id")218            return {"status": 11, "detail": "请提供speaker_name或speaker_id"}219        elif speaker_name is None:220            # 检查speaker_id是否存在221            if speaker_id not in loaded_models.models[model_id].id2spk.keys():222                logger.error(f"/voice 请求错误:角色speaker_id={speaker_id}不存在")223                return {"status": 12, "detail": f"角色speaker_id={speaker_id}不存在"}224            speaker_name = loaded_models.models[model_id].id2spk[speaker_id]225        # 检查speaker_name是否存在226        if speaker_name not in loaded_models.models[model_id].spk2id.keys():227            logger.error(f"/voice 请求错误:角色speaker_name={speaker_name}不存在")228            return {"status": 13, "detail": f"角色speaker_name={speaker_name}不存在"}229        # 未传入则使用默认语言230        if language is None:231            language = loaded_models.models[model_id].language232        # 翻译会破坏mix结构,auto也会变得无意义。不要在这两个模式下使用233        if auto_translate:234            if language == "auto" or language == "mix":235                logger.error(236                    f"/voice 请求错误:请勿同时使用language = {language}与auto_translate模式"237                )238                return {239                    "status": 20,240                    "detail": f"请勿同时使用language = {language}与auto_translate模式",241                }242            text = trans.translate(Sentence=text, to_Language=language.lower())243        if reference_audio is not None:244            ref_audio = BytesIO(await reference_audio.read())245            # 2.2 适配246            if loaded_models.models[model_id].version == "2.2":247                ref_audio, _ = librosa.load(ref_audio, 48000)248 249        else:250            ref_audio = reference_audio251        if not auto_split:252            with torch.no_grad():253                audio = infer(254                    text=text,255                    sdp_ratio=sdp_ratio,256                    noise_scale=noise,257                    noise_scale_w=noisew,258                    length_scale=length,259                    sid=speaker_name,260                    language=language,261                    hps=loaded_models.models[model_id].hps,262                    net_g=loaded_models.models[model_id].net_g,263                    device=loaded_models.models[model_id].device,264                    emotion=emotion,265                    reference_audio=ref_audio,266                    style_text=style_text,267                    style_weight=style_weight,268                )269                audio = gradio.processing_utils.convert_to_16_bit_wav(audio)270        else:271            texts = cut_sent(text)272            audios = []273            with torch.no_grad():274                for t in texts:275                    audios.append(276                        infer(277                            text=t,278                            sdp_ratio=sdp_ratio,279                            noise_scale=noise,280                            noise_scale_w=noisew,281                            length_scale=length,282                            sid=speaker_name,283                            language=language,284                            hps=loaded_models.models[model_id].hps,285                            net_g=loaded_models.models[model_id].net_g,286                            device=loaded_models.models[model_id].device,287                            emotion=emotion,288                            reference_audio=ref_audio,289                            style_text=style_text,290                            style_weight=style_weight,291                        )292                    )293                    audios.append(np.zeros(int(44100 * 0.2)))294                audio = np.concatenate(audios)295                audio = gradio.processing_utils.convert_to_16_bit_wav(audio)296        with BytesIO() as wavContent:297            wavfile.write(298                wavContent, loaded_models.models[model_id].hps.data.sampling_rate, audio299            )300            response = Response(content=wavContent.getvalue(), media_type="audio/wav")301            return response302 303    @app.post("/voice")304    async def voice(305        request: Request,  # fastapi自动注入306        text: str = Form(...),307        model_id: int = Query(..., description="模型ID"),  # 模型序号308        speaker_name: str = Query(309            None, description="说话人名"310        ),  # speaker_name与 speaker_id二者选其一311        speaker_id: int = Query(None, description="说话人id,与speaker_name二选一"),312        sdp_ratio: float = Query(0.2, description="SDP/DP混合比"),313        noise: float = Query(0.2, description="感情"),314        noisew: float = Query(0.9, description="音素长度"),315        length: float = Query(1, description="语速"),316        language: str = Query(None, description="语言"),  # 若不指定使用语言则使用默认值317        auto_translate: bool = Query(False, description="自动翻译"),318        auto_split: bool = Query(False, description="自动切分"),319        emotion: Optional[Union[int, str]] = Query(None, description="emo"),320        reference_audio: UploadFile = File(None),321        style_text: Optional[str] = Form(None, description="风格文本"),322        style_weight: float = Query(0.7, description="风格权重"),323    ):324        """语音接口,若需要上传参考音频请仅使用post请求"""325        logger.info(326            f"{request.client.host}:{request.client.port}/voice  { unquote(str(request.query_params) )} text={text}"327        )328        return await _voice(329            text=text,330            model_id=model_id,331            speaker_name=speaker_name,332            speaker_id=speaker_id,333            sdp_ratio=sdp_ratio,334            noise=noise,335            noisew=noisew,336            length=length,337            language=language,338            auto_translate=auto_translate,339            auto_split=auto_split,340            emotion=emotion,341            reference_audio=reference_audio,342            style_text=style_text,343            style_weight=style_weight,344        )345 346    @app.get("/voice")347    async def voice(348        request: Request,  # fastapi自动注入349        text: str = Query(..., description="输入文字"),350        model_id: int = Query(..., description="模型ID"),  # 模型序号351        speaker_name: str = Query(352            None, description="说话人名"353        ),  # speaker_name与 speaker_id二者选其一354        speaker_id: int = Query(None, description="说话人id,与speaker_name二选一"),355        sdp_ratio: float = Query(0.2, description="SDP/DP混合比"),356        noise: float = Query(0.2, description="感情"),357        noisew: float = Query(0.9, description="音素长度"),358        length: float = Query(1, description="语速"),359        language: str = Query(None, description="语言"),  # 若不指定使用语言则使用默认值360        auto_translate: bool = Query(False, description="自动翻译"),361        auto_split: bool = Query(False, description="自动切分"),362        emotion: Optional[Union[int, str]] = Query(None, description="emo"),363        style_text: Optional[str] = Query(None, description="风格文本"),364        style_weight: float = Query(0.7, description="风格权重"),365    ):366        """语音接口"""367        logger.info(368            f"{request.client.host}:{request.client.port}/voice  { unquote(str(request.query_params) )}"369        )370        return await _voice(371            text=text,372            model_id=model_id,373            speaker_name=speaker_name,374            speaker_id=speaker_id,375            sdp_ratio=sdp_ratio,376            noise=noise,377            noisew=noisew,378            length=length,379            language=language,380            auto_translate=auto_translate,381            auto_split=auto_split,382            emotion=emotion,383            style_text=style_text,384            style_weight=style_weight,385        )386 387    @app.get("/models/info")388    def get_loaded_models_info(request: Request):389        """获取已加载模型信息"""390 391        result: Dict[str, Dict] = dict()392        for key, model in loaded_models.models.items():393            result[str(key)] = model.to_dict()394        return result395 396    @app.get("/models/delete")397    def delete_model(398        request: Request, model_id: int = Query(..., description="删除模型id")399    ):400        """删除指定模型"""401        logger.info(402            f"{request.client.host}:{request.client.port}/models/delete  { unquote(str(request.query_params) )}"403        )404        result = loaded_models.del_model(model_id)405        if result is None:406            logger.error(f"/models/delete 模型删除错误:模型{model_id}不存在,删除失败")407            return {"status": 14, "detail": f"模型{model_id}不存在,删除失败"}408 409        return {"status": 0, "detail": "删除成功"}410 411    @app.get("/models/add")412    def add_model(413        request: Request,414        model_path: str = Query(..., description="添加模型路径"),415        config_path: str = Query(416            None, description="添加模型配置文件路径,不填则使用./config.json或../config.json"417        ),418        device: str = Query("cuda", description="推理使用设备"),419        language: str = Query("ZH", description="模型默认语言"),420    ):421        """添加指定模型:允许重复添加相同路径模型,且不重复占用内存"""422        logger.info(423            f"{request.client.host}:{request.client.port}/models/add  { unquote(str(request.query_params) )}"424        )425        if config_path is None:426            model_dir = os.path.dirname(model_path)427            if os.path.isfile(os.path.join(model_dir, "config.json")):428                config_path = os.path.join(model_dir, "config.json")429            elif os.path.isfile(os.path.join(model_dir, "../config.json")):430                config_path = os.path.join(model_dir, "../config.json")431            else:432                logger.error("/models/add 模型添加失败:未在模型所在目录以及上级目录找到config.json文件")433                return {434                    "status": 15,435                    "detail": "查询未传入配置文件路径,同时默认路径./与../中不存在配置文件config.json。",436                }437        try:438            model_id = loaded_models.init_model(439                config_path=config_path,440                model_path=model_path,441                device=device,442                language=language,443            )444        except Exception:445            logging.exception("模型加载出错")446            return {447                "status": 16,448                "detail": "模型加载出错,详细查看日志",449            }450        return {451            "status": 0,452            "detail": "模型添加成功",453            "Data": {454                "model_id": model_id,455                "model_info": loaded_models.models[model_id].to_dict(),456            },457        }458 459    def _get_all_models(root_dir: str = "Data", only_unloaded: bool = False):460        """从root_dir搜索获取所有可用模型"""461        result: Dict[str, List[str]] = dict()462        files = os.listdir(root_dir) + ["."]463        for file in files:464            if os.path.isdir(os.path.join(root_dir, file)):465                sub_dir = os.path.join(root_dir, file)466                # 搜索 "sub_dir" 、 "sub_dir/models" 两个路径467                result[file] = list()468                sub_files = os.listdir(sub_dir)469                model_files = []470                for sub_file in sub_files:471                    relpath = os.path.realpath(os.path.join(sub_dir, sub_file))472                    if only_unloaded and relpath in loaded_models.path2ids.keys():473                        continue474                    if sub_file.endswith(".pth") and sub_file.startswith("G_"):475                        if os.path.isfile(relpath):476                            model_files.append(sub_file)477                # 对模型文件按步数排序478                model_files = sorted(479                    model_files,480                    key=lambda pth: int(pth.lstrip("G_").rstrip(".pth"))481                    if pth.lstrip("G_").rstrip(".pth").isdigit()482                    else 10**10,483                )484                result[file] = model_files485                models_dir = os.path.join(sub_dir, "models")486                model_files = []487                if os.path.isdir(models_dir):488                    sub_files = os.listdir(models_dir)489                    for sub_file in sub_files:490                        relpath = os.path.realpath(os.path.join(models_dir, sub_file))491                        if only_unloaded and relpath in loaded_models.path2ids.keys():492                            continue493                        if sub_file.endswith(".pth") and sub_file.startswith("G_"):494                            if os.path.isfile(os.path.join(models_dir, sub_file)):495                                model_files.append(f"models/{sub_file}")496                    # 对模型文件按步数排序497                    model_files = sorted(498                        model_files,499                        key=lambda pth: int(pth.lstrip("models/G_").rstrip(".pth"))500                        if pth.lstrip("models/G_").rstrip(".pth").isdigit()501                        else 10**10,502                    )503                    result[file] += model_files504                if len(result[file]) == 0:505                    result.pop(file)506 507        return result508 509    @app.get("/models/get_unloaded")510    def get_unloaded_models_info(511        request: Request, root_dir: str = Query("Data", description="搜索根目录")512    ):513        """获取未加载模型"""514        logger.info(515            f"{request.client.host}:{request.client.port}/models/get_unloaded  { unquote(str(request.query_params) )}"516        )517        return _get_all_models(root_dir, only_unloaded=True)518 519    @app.get("/models/get_local")520    def get_local_models_info(521        request: Request, root_dir: str = Query("Data", description="搜索根目录")522    ):523        """获取全部本地模型"""524        logger.info(525            f"{request.client.host}:{request.client.port}/models/get_local  { unquote(str(request.query_params) )}"526        )527        return _get_all_models(root_dir, only_unloaded=False)528 529    @app.get("/status")530    def get_status():531        """获取电脑运行状态"""532        cpu_percent = psutil.cpu_percent(interval=1)533        memory_info = psutil.virtual_memory()534        memory_total = memory_info.total535        memory_available = memory_info.available536        memory_used = memory_info.used537        memory_percent = memory_info.percent538        gpuInfo = []539        devices = ["cpu"]540        for i in range(torch.cuda.device_count()):541            devices.append(f"cuda:{i}")542        gpus = GPUtil.getGPUs()543        for gpu in gpus:544            gpuInfo.append(545                {546                    "gpu_id": gpu.id,547                    "gpu_load": gpu.load,548                    "gpu_memory": {549                        "total": gpu.memoryTotal,550                        "used": gpu.memoryUsed,551                        "free": gpu.memoryFree,552                    },553                }554            )555        return {556            "devices": devices,557            "cpu_percent": cpu_percent,558            "memory_total": memory_total,559            "memory_available": memory_available,560            "memory_used": memory_used,561            "memory_percent": memory_percent,562            "gpu": gpuInfo,563        }564 565    @app.get("/tools/translate")566    def translate(567        request: Request,568        texts: str = Query(..., description="待翻译文本"),569        to_language: str = Query(..., description="翻译目标语言"),570    ):571        """翻译"""572        logger.info(573            f"{request.client.host}:{request.client.port}/tools/translate  { unquote(str(request.query_params) )}"574        )575        return {"texts": trans.translate(Sentence=texts, to_Language=to_language)}576 577    all_examples: Dict[str, Dict[str, List]] = dict()  # 存放示例578 579    @app.get("/tools/random_example")580    def random_example(581        request: Request,582        language: str = Query(None, description="指定语言,未指定则随机返回"),583        root_dir: str = Query("Data", description="搜索根目录"),584    ):585        """586        获取一个随机音频+文本,用于对比,音频会从本地目录随机选择。587        """588        logger.info(589            f"{request.client.host}:{request.client.port}/tools/random_example  { unquote(str(request.query_params) )}"590        )591        global all_examples592        # 数据初始化593        if root_dir not in all_examples.keys():594            all_examples[root_dir] = {"ZH": [], "JP": [], "EN": []}595 596            examples = all_examples[root_dir]597 598            # 从项目Data目录中搜索train/val.list599            for root, directories, _files in os.walk(root_dir):600                for file in _files:601                    if file in ["train.list", "val.list"]:602                        with open(603                            os.path.join(root, file), mode="r", encoding="utf-8"604                        ) as f:605                            lines = f.readlines()606                            for line in lines:607                                data = line.split("|")608                                if len(data) != 7:609                                    continue610                                # 音频存在 且语言为ZH/EN/JP611                                if os.path.isfile(data[0]) and data[2] in [612                                    "ZH",613                                    "JP",614                                    "EN",615                                ]:616                                    examples[data[2]].append(617                                        {618                                            "text": data[3],619                                            "audio": data[0],620                                            "speaker": data[1],621                                        }622                                    )623 624        examples = all_examples[root_dir]625        if language is None:626            if len(examples["ZH"]) + len(examples["JP"]) + len(examples["EN"]) == 0:627                return {"status": 17, "detail": "没有加载任何示例数据"}628            else:629                # 随机选一个630                rand_num = random.randint(631                    0,632                    len(examples["ZH"]) + len(examples["JP"]) + len(examples["EN"]) - 1,633                )634                # ZH635                if rand_num < len(examples["ZH"]):636                    return {"status": 0, "Data": examples["ZH"][rand_num]}637                # JP638                if rand_num < len(examples["ZH"]) + len(examples["JP"]):639                    return {640                        "status": 0,641                        "Data": examples["JP"][rand_num - len(examples["ZH"])],642                    }643                # EN644                return {645                    "status": 0,646                    "Data": examples["EN"][647                        rand_num - len(examples["ZH"]) - len(examples["JP"])648                    ],649                }650 651        else:652            if len(examples[language]) == 0:653                return {"status": 17, "detail": f"没有加载任何{language}数据"}654            return {655                "status": 0,656                "Data": examples[language][657                    random.randint(0, len(examples[language]) - 1)658                ],659            }660 661    @app.get("/tools/get_audio")662    def get_audio(request: Request, path: str = Query(..., description="本地音频路径")):663        logger.info(664            f"{request.client.host}:{request.client.port}/tools/get_audio  { unquote(str(request.query_params) )}"665        )666        if not os.path.isfile(path):667            logger.error(f"/tools/get_audio 获取音频错误:指定音频{path}不存在")668            return {"status": 18, "detail": "指定音频不存在"}669        if not path.lower().endswith(".wav"):670            logger.error(f"/tools/get_audio 获取音频错误:音频{path}非wav文件")671            return {"status": 19, "detail": "非wav格式文件"}672        return FileResponse(path=path)673 674    logger.warning("本地服务,请勿将服务端口暴露于外网")675    logger.info(f"api文档地址 http://127.0.0.1:{config.server_config.port}/docs")676    if os.path.isdir(StaticDir):677        webbrowser.open(f"http://127.0.0.1:{config.server_config.port}")678    uvicorn.run(679        app, port=config.server_config.port, host="0.0.0.0", log_level="warning"680    )681