Julius8888/XiJingPing_Voice_Clone
0
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 