diff --git a/common/log.py b/common/log.py index 679bb2c..71b5e63 100644 --- a/common/log.py +++ b/common/log.py @@ -14,4 +14,5 @@ log_format = ( "{time:MM-DD HH:mm:ss} |{level:^8}| {file}:{line} | {message}" ) +# logger.add(SAFE_STDOUT, format=log_format, backtrace=True, diagnose=True, level="TRACE") logger.add(SAFE_STDOUT, format=log_format, backtrace=True, diagnose=True) 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 afb9321..402116f 100644 --- a/server_editor.py +++ b/server_editor.py @@ -16,13 +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 pyopenjtalk 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 @@ -43,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 # ---フロントエンド部分に関する処理--- @@ -230,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): @@ -250,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) @@ -275,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, @@ -291,6 +298,7 @@ def synthesis(request: SynthesisRequest): line_split=False, pitch_scale=request.pitchScale, intonation_scale=request.intonationScale, + sid=sid, ) with BytesIO() as wavContent: diff --git a/text/japanese.py b/text/japanese.py index b18bc68..fea0eaa 100644 --- a/text/japanese.py +++ b/text/japanese.py @@ -4,7 +4,9 @@ import re import unicodedata from pathlib import Path -import pyopenjtalk +from . import pyopenjtalk_worker as pyopenjtalk + +pyopenjtalk.initialize() from num2words import num2words from transformers import AutoTokenizer diff --git a/text/pyopenjtalk_worker/__init__.py b/text/pyopenjtalk_worker/__init__.py new file mode 100644 index 0000000..7aa0d0e --- /dev/null +++ b/text/pyopenjtalk_worker/__init__.py @@ -0,0 +1,126 @@ +""" +Run the pyopenjtalk worker in a separate process +to avoid user dictionary access error +""" + +from typing import Optional, Any + +from .worker_common import WORKER_PORT +from .worker_client import WorkerClient + +from common.log import logger + +WORKER_CLIENT: Optional[WorkerClient] = None + +# pyopenjtalk interface + +# g2p: not used + + +def run_frontend(text: str) -> list[dict[str, Any]]: + assert WORKER_CLIENT + ret = WORKER_CLIENT.dispatch_pyopenjtalk("run_frontend", text) + assert isinstance(ret, list) + return ret + + +def make_label(njd_features) -> list[str]: + assert WORKER_CLIENT + ret = WORKER_CLIENT.dispatch_pyopenjtalk("make_label", njd_features) + assert isinstance(ret, list) + return ret + + +def mecab_dict_index(path: str, out_path: str, dn_mecab: Optional[str] = None): + assert WORKER_CLIENT + WORKER_CLIENT.dispatch_pyopenjtalk("mecab_dict_index", path, out_path, dn_mecab) + + +def update_global_jtalk_with_user_dict(path: str): + assert WORKER_CLIENT + WORKER_CLIENT.dispatch_pyopenjtalk("update_global_jtalk_with_user_dict", path) + + +def unset_user_dict(): + assert WORKER_CLIENT + WORKER_CLIENT.dispatch_pyopenjtalk("unset_user_dict") + + +# initialize module when imported + + +def initialize(port: int = WORKER_PORT): + import time + import socket + import sys + import atexit + import signal + + logger.debug("initialize") + global WORKER_CLIENT + if WORKER_CLIENT: + return + + client = None + try: + client = WorkerClient(port) + except (socket.timeout, socket.error): + logger.debug("try starting pyopenjtalk worker server") + import os + import subprocess + + worker_pkg_path = os.path.relpath( + os.path.dirname(__file__), os.getcwd() + ).replace(os.sep, ".") + args = [sys.executable, "-m", worker_pkg_path, "--port", str(port)] + # new session, new process group + if sys.platform.startswith("win"): + cf = subprocess.CREATE_NEW_CONSOLE | subprocess.CREATE_NEW_PROCESS_GROUP # type: ignore + si = subprocess.STARTUPINFO() # type: ignore + si.dwFlags |= subprocess.STARTF_USESHOWWINDOW # type: ignore + si.wShowWindow = subprocess.SW_HIDE # type: ignore + subprocess.Popen(args, creationflags=cf, startupinfo=si) + else: + # align with Windows behavior + # start_new_session is same as specifying setsid in preexec_fn + subprocess.Popen(args, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, start_new_session=True) # type: ignore + + # wait until server listening + count = 0 + while True: + try: + client = WorkerClient(port) + break + except socket.error: + time.sleep(1) + count += 1 + # 10: max number of retries + if count == 10: + raise TimeoutError("サーバーに接続できませんでした") + + WORKER_CLIENT = client + atexit.register(terminate) + + # when the process is killed + def signal_handler(signum, frame): + terminate() + + signal.signal(signal.SIGTERM, signal_handler) + + +# top-level declaration +def terminate(): + logger.debug("terminate") + global WORKER_CLIENT + if not WORKER_CLIENT: + return + + # repare for unexpected errors + try: + if WORKER_CLIENT.status() == 1: + WORKER_CLIENT.quit_server() + except Exception as e: + logger.error(e) + + WORKER_CLIENT.close() + WORKER_CLIENT = None diff --git a/text/pyopenjtalk_worker/__main__.py b/text/pyopenjtalk_worker/__main__.py new file mode 100644 index 0000000..3bb6b53 --- /dev/null +++ b/text/pyopenjtalk_worker/__main__.py @@ -0,0 +1,16 @@ +import argparse + +from .worker_server import WorkerServer +from .worker_common import WORKER_PORT + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--port", type=int, default=WORKER_PORT) + args = parser.parse_args() + server = WorkerServer() + server.start_server(port=args.port) + + +if __name__ == "__main__": + main() diff --git a/text/pyopenjtalk_worker/worker_client.py b/text/pyopenjtalk_worker/worker_client.py new file mode 100644 index 0000000..86d8969 --- /dev/null +++ b/text/pyopenjtalk_worker/worker_client.py @@ -0,0 +1,55 @@ +from typing import Any +import socket + +from .worker_common import RequestType, receive_data, send_data + +from common.log import logger + + +class WorkerClient: + def __init__(self, port: int) -> None: + sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + # 5: timeout + sock.settimeout(5) + sock.connect((socket.gethostname(), port)) + self.sock = sock + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_value, traceback): + self.close() + + def close(self): + self.sock.close() + + def dispatch_pyopenjtalk(self, func: str, *args, **kwargs): + data = { + "request-type": RequestType.PYOPENJTALK, + "func": func, + "args": args, + "kwargs": kwargs, + } + logger.trace(f"client sends request: {data}") + send_data(self.sock, data) + logger.trace("client sent request successfully") + response = receive_data(self.sock) + logger.trace(f"client received response: {response}") + return response.get("return") + + def status(self): + data = {"request-type": RequestType.STATUS} + logger.trace(f"client sends request: {data}") + send_data(self.sock, data) + logger.trace("client sent request successfully") + response = receive_data(self.sock) + logger.trace(f"client received response: {response}") + return response.get("client-count") + + def quit_server(self): + data = {"request-type": RequestType.QUIT_SERVER} + logger.trace(f"client sends request: {data}") + send_data(self.sock, data) + logger.trace("client sent request successfully") + response = receive_data(self.sock) + logger.trace(f"client received response: {response}") diff --git a/text/pyopenjtalk_worker/worker_common.py b/text/pyopenjtalk_worker/worker_common.py new file mode 100644 index 0000000..606d0c3 --- /dev/null +++ b/text/pyopenjtalk_worker/worker_common.py @@ -0,0 +1,44 @@ +from typing import Any, Optional, Final +from enum import IntEnum, auto +import socket +import json + +WORKER_PORT: Final[int] = 7861 +HEADER_SIZE: Final[int] = 4 + + +class RequestType(IntEnum): + STATUS = auto() + QUIT_SERVER = auto() + PYOPENJTALK = auto() + + +class ConnectionClosedException(Exception): + pass + + +# socket communication + + +def send_data(sock: socket.socket, data: dict[str, Any]): + json_data = json.dumps(data).encode() + header = len(json_data).to_bytes(HEADER_SIZE, byteorder="big") + sock.sendall(header + json_data) + + +def _receive_until(sock: socket.socket, size: int): + data = b"" + while len(data) < size: + part = sock.recv(size - len(data)) + if part == b"": + raise ConnectionClosedException("接続が閉じられました") + data += part + + return data + + +def receive_data(sock: socket.socket) -> dict[str, Any]: + header = _receive_until(sock, HEADER_SIZE) + data_length = int.from_bytes(header, byteorder="big") + body = _receive_until(sock, data_length) + return json.loads(body.decode()) diff --git a/text/pyopenjtalk_worker/worker_server.py b/text/pyopenjtalk_worker/worker_server.py new file mode 100644 index 0000000..6babd91 --- /dev/null +++ b/text/pyopenjtalk_worker/worker_server.py @@ -0,0 +1,118 @@ +import pyopenjtalk +import socket +import select +import time + +from .worker_common import ( + ConnectionClosedException, + RequestType, + receive_data, + send_data, +) + +from common.log import logger + +# To make it as fast as possible +# Probably faster than calling getattr every time +_PYOPENJTALK_FUNC_DICT = { + "run_frontend": pyopenjtalk.run_frontend, + "make_label": pyopenjtalk.make_label, + "mecab_dict_index": pyopenjtalk.mecab_dict_index, + "update_global_jtalk_with_user_dict": pyopenjtalk.update_global_jtalk_with_user_dict, + "unset_user_dict": pyopenjtalk.unset_user_dict, +} + + +class WorkerServer: + def __init__(self) -> None: + self.client_count: int = 0 + self.quit: bool = False + + def handle_request(self, request): + request_type = None + try: + request_type = RequestType(request.get("request-type")) + except Exception: + return { + "success": False, + "reason": "request-type is invalid", + } + + if request_type: + if request_type == RequestType.STATUS: + response = { + "success": True, + "client-count": self.client_count, + } + elif request_type == RequestType.QUIT_SERVER: + self.quit = True + response = {"success": True} + elif request_type == RequestType.PYOPENJTALK: + func_name = request.get("func") + assert isinstance(func_name, str) + func = _PYOPENJTALK_FUNC_DICT[func_name] + args = request.get("args") + kwargs = request.get("kwargs") + assert isinstance(args, list) + assert isinstance(kwargs, dict) + ret = func(*args, **kwargs) + response = {"success": True, "return": ret} + else: + # NOT REACHED + response = request + + return response + + def start_server(self, port: int, no_client_timeout: int = 30): + logger.info("start pyopenjtalk worker server") + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as server_socket: + server_socket.bind((socket.gethostname(), port)) + server_socket.listen() + sockets = [server_socket] + no_client_since = time.time() + while True: + if self.client_count == 0: + if no_client_since is None: + no_client_since = time.time() + elif (time.time() - no_client_since) > no_client_timeout: + logger.info("quit because there is no client") + return + else: + no_client_since = None + + ready_sockets, _, _ = select.select(sockets, [], [], 0.1) + for sock in ready_sockets: + if sock is server_socket: + logger.info("new client connected") + client_socket, _ = server_socket.accept() + sockets.append(client_socket) + self.client_count += 1 + else: + # client + try: + request = receive_data(sock) + except Exception as e: + sock.close() + sockets.remove(sock) + self.client_count -= 1 + # unexpected disconnections + if not isinstance(e, ConnectionClosedException): + logger.error(e) + + logger.info("close connection") + continue + + logger.trace(f"server received request: {request}") + + response = self.handle_request(request) + logger.trace(f"server sends response: {response}") + try: + send_data(sock, response) + logger.trace("server sent response successfully") + except Exception: + logger.warning( + "an exception occurred during sending responce" + ) + if self.quit: + logger.info("quit pyopenjtalk worker server") + return diff --git a/text/user_dict/__init__.py b/text/user_dict/__init__.py index c12b3d1..95515cb 100644 --- a/text/user_dict/__init__.py +++ b/text/user_dict/__init__.py @@ -12,7 +12,9 @@ from typing import Dict, List, Optional from uuid import UUID, uuid4 import numpy as np -import pyopenjtalk +from .. import pyopenjtalk_worker as pyopenjtalk + +pyopenjtalk.initialize() from fastapi import HTTPException from .word_model import UserDictWord, WordTypes