From 25ed226acbe13ba9425169c89bc1452c842f424e Mon Sep 17 00:00:00 2001 From: litagin02 Date: Fri, 8 Mar 2024 11:19:03 +0900 Subject: [PATCH] Add speaker list api --- common/tts_model.py | 21 ++++++++------------- server_editor.py | 17 +++++++++++++---- 2 files changed, 21 insertions(+), 17 deletions(-) diff --git a/common/tts_model.py b/common/tts_model.py index e09787e..9e2171d 100644 --- a/common/tts_model.py +++ b/common/tts_model.py @@ -222,12 +222,14 @@ class ModelHolder: self.current_model: Optional[Model] = None self.model_names: list[str] = [] self.models: list[Model] = [] + self.models_info: list[dict[str, Union[str, list[str]]]] = [] self.refresh() def refresh(self): self.model_files_dict = {} self.model_names = [] self.current_model = None + self.models_info = [] model_dirs = [d for d in self.root_dir.iterdir() if d.is_dir()] for model_dir in model_dirs: @@ -247,26 +249,19 @@ class ModelHolder: continue self.model_files_dict[model_dir.name] = model_files self.model_names.append(model_dir.name) - - def models_info(self): - if hasattr(self, "_models_info"): - return self._models_info - result = [] - for name, files in self.model_files_dict.items(): - # Get styles - config_path = self.root_dir / name / "config.json" hps = utils.get_hparams_from_file(config_path) style2id: dict[str, int] = hps.data.style2id styles = list(style2id.keys()) - result.append( + spk2id: dict[str, int] = hps.data.spk2id + speakers = list(spk2id.keys()) + self.models_info.append( { - "name": name, - "files": [str(f) for f in files], + "name": model_dir.name, + "files": [str(f) for f in model_files], "styles": styles, + "speakers": speakers, } ) - self._models_info = result - return result def load_model(self, model_name: str, model_path_str: str): model_path = Path(model_path_str) diff --git a/server_editor.py b/server_editor.py index 1ae323f..402116f 100644 --- a/server_editor.py +++ b/server_editor.py @@ -16,12 +16,13 @@ import zipfile from datetime import datetime from io import BytesIO from pathlib import Path -import yaml +from typing import Optional import numpy as np import requests import torch import uvicorn +import yaml from fastapi import APIRouter, FastAPI, HTTPException, status from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import JSONResponse, Response @@ -42,8 +43,7 @@ from common.constants import ( from common.log import logger from common.tts_model import ModelHolder from text.japanese import g2kata_tone, kata_tone2phone_tone, text_normalize -from text.user_dict import apply_word, update_dict, read_dict, rewrite_word, delete_word - +from text.user_dict import apply_word, delete_word, read_dict, rewrite_word, update_dict # ---フロントエンド部分に関する処理--- @@ -229,7 +229,7 @@ async def normalize_text(item: TextRequest): @router.get("/models_info") def models_info(): - return model_holder.models_info() + return model_holder.models_info class SynthesisRequest(BaseModel): @@ -249,6 +249,7 @@ class SynthesisRequest(BaseModel): silenceAfter: float = 0.5 pitchScale: float = 1.0 intonationScale: float = 1.0 + speaker: Optional[str] = None @router.post("/synthesis", response_class=AudioResponse) @@ -274,6 +275,13 @@ def synthesis(request: SynthesisRequest): ] phone_tone = kata_tone2phone_tone(kata_tone_list) tone = [t for _, t in phone_tone] + try: + sid = 0 if request.speaker is None else model.spk2id[request.speaker] + except KeyError: + raise HTTPException( + status_code=400, + detail=f"Speaker {request.speaker} not found in {model.spk2id}", + ) sr, audio = model.infer( text=text, language=request.language.value, @@ -290,6 +298,7 @@ def synthesis(request: SynthesisRequest): line_split=False, pitch_scale=request.pitchScale, intonation_scale=request.intonationScale, + sid=sid, ) with BytesIO() as wavContent: