From bc1058270c9a8648cae01bd301a136ef4ff92466 Mon Sep 17 00:00:00 2001 From: kale4eat Date: Tue, 5 Mar 2024 12:16:14 +0900 Subject: [PATCH 01/11] Delete pyopenjtalk import --- server_editor.py | 1 - 1 file changed, 1 deletion(-) diff --git a/server_editor.py b/server_editor.py index afb9321..1ae323f 100644 --- a/server_editor.py +++ b/server_editor.py @@ -19,7 +19,6 @@ from pathlib import Path import yaml import numpy as np -import pyopenjtalk import requests import torch import uvicorn From 1eb8fb4b08c5dba3c2d7bdd40bf56da46369d321 Mon Sep 17 00:00:00 2001 From: kale4eat Date: Tue, 5 Mar 2024 12:17:15 +0900 Subject: [PATCH 02/11] add openjtalk worker pkg --- text/pyopenjtalk_worker/__init__.py | 100 ++++++++++++++++++++++ text/pyopenjtalk_worker/__main__.py | 16 ++++ text/pyopenjtalk_worker/worker_client.py | 40 +++++++++ text/pyopenjtalk_worker/worker_common.py | 44 ++++++++++ text/pyopenjtalk_worker/worker_server.py | 103 +++++++++++++++++++++++ 5 files changed, 303 insertions(+) create mode 100644 text/pyopenjtalk_worker/__init__.py create mode 100644 text/pyopenjtalk_worker/__main__.py create mode 100644 text/pyopenjtalk_worker/worker_client.py create mode 100644 text/pyopenjtalk_worker/worker_common.py create mode 100644 text/pyopenjtalk_worker/worker_server.py diff --git a/text/pyopenjtalk_worker/__init__.py b/text/pyopenjtalk_worker/__init__.py new file mode 100644 index 0000000..6ab4ca1 --- /dev/null +++ b/text/pyopenjtalk_worker/__init__.py @@ -0,0 +1,100 @@ +""" +Run the pyopenjtalk worker in a separate process +to avoid user dictionary access error +""" + +from typing import Optional, Any + +from .worker_common import WOKER_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 = WOKER_PORT): + import time + import socket + import sys + import atexit + + global WORKER_CLIENT + logger.debug("initialize") + if WORKER_CLIENT: + return + + client = None + try: + client = WorkerClient(port) + except (socket.timeout, socket.error): + logger.debug("try starting worker server") + import os + import subprocess + + worker_pkg_path = os.path.relpath( + os.path.dirname(__file__), os.getcwd() + ).replace(os.sep, ".") + subprocess.Popen([sys.executable, "-m", worker_pkg_path, "--port", str(port)]) + # 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 + + def terminate(): + global WORKER_CLIENT + if not WORKER_CLIENT: + return + + if WORKER_CLIENT.status().get("client-count") == 1: + WORKER_CLIENT.quit_server() + WORKER_CLIENT.close() + WORKER_CLIENT = None + + atexit.register(terminate) diff --git a/text/pyopenjtalk_worker/__main__.py b/text/pyopenjtalk_worker/__main__.py new file mode 100644 index 0000000..8b67aa0 --- /dev/null +++ b/text/pyopenjtalk_worker/__main__.py @@ -0,0 +1,16 @@ +import argparse + +from .worker_server import WorkerServer +from .worker_common import WOKER_PORT + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--port", type=int, default=WOKER_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..bd9a32a --- /dev/null +++ b/text/pyopenjtalk_worker/worker_client.py @@ -0,0 +1,40 @@ +from typing import Any +import socket + +from .worker_common import RequestType, receive_data, send_data + + +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, + } + send_data(self.sock, data) + return receive_data(self.sock).get("return") + + def status(self): + send_data(self.sock, {"request-type": RequestType.STATUS}) + return receive_data(self.sock) + + def quit_server(self): + send_data(self.sock, {"request-type": RequestType.QUIT_SERVER}) + receive_data(self.sock) diff --git a/text/pyopenjtalk_worker/worker_common.py b/text/pyopenjtalk_worker/worker_common.py new file mode 100644 index 0000000..bea552e --- /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 + +WOKER_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..a2b9f2f --- /dev/null +++ b/text/pyopenjtalk_worker/worker_server.py @@ -0,0 +1,103 @@ +import pyopenjtalk +import socket +import select + + +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): + 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] + while True: + 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 ConnectionClosedException as e: + sock.close() + sockets.remove(sock) + self.client_count -= 1 + logger.info("close connection") + continue + + logger.debug(f"receive request: {request}") + + response = self.handle_request(request) + logger.debug(f"send response: {response}") + try: + send_data(sock, response) + except Exception: + logger.warning( + "an exception occurred during sending responce" + ) + if self.quit: + logger.info("quit pyopenjtalk worker server") + return From 85f5b9bd25e930dc8e0fd8e96614cd1eaa5127f7 Mon Sep 17 00:00:00 2001 From: kale4eat Date: Tue, 5 Mar 2024 12:18:05 +0900 Subject: [PATCH 03/11] replace pyopenjtalk import with worker --- text/japanese.py | 4 +++- text/user_dict/__init__.py | 4 +++- 2 files changed, 6 insertions(+), 2 deletions(-) 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/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 From 98ff976ec0e187bef23c6bf5d91bbe09829c9474 Mon Sep 17 00:00:00 2001 From: kale4eat Date: Tue, 5 Mar 2024 15:04:27 +0900 Subject: [PATCH 04/11] modify logging --- text/pyopenjtalk_worker/__init__.py | 2 +- text/pyopenjtalk_worker/worker_client.py | 8 +++++++- text/pyopenjtalk_worker/worker_server.py | 6 +++--- 3 files changed, 11 insertions(+), 5 deletions(-) diff --git a/text/pyopenjtalk_worker/__init__.py b/text/pyopenjtalk_worker/__init__.py index 6ab4ca1..ef5873c 100644 --- a/text/pyopenjtalk_worker/__init__.py +++ b/text/pyopenjtalk_worker/__init__.py @@ -64,7 +64,7 @@ def initialize(port: int = WOKER_PORT): try: client = WorkerClient(port) except (socket.timeout, socket.error): - logger.debug("try starting worker server") + logger.debug("try starting pyopenjtalk worker server") import os import subprocess diff --git a/text/pyopenjtalk_worker/worker_client.py b/text/pyopenjtalk_worker/worker_client.py index bd9a32a..23f7dbe 100644 --- a/text/pyopenjtalk_worker/worker_client.py +++ b/text/pyopenjtalk_worker/worker_client.py @@ -3,6 +3,8 @@ import socket from .worker_common import RequestType, receive_data, send_data +from common.log import logger + class WorkerClient: def __init__(self, port: int) -> None: @@ -28,8 +30,12 @@ class WorkerClient: "args": args, "kwargs": kwargs, } + logger.trace(f"client sends request: {data}") send_data(self.sock, data) - return receive_data(self.sock).get("return") + 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): send_data(self.sock, {"request-type": RequestType.STATUS}) diff --git a/text/pyopenjtalk_worker/worker_server.py b/text/pyopenjtalk_worker/worker_server.py index a2b9f2f..13bc771 100644 --- a/text/pyopenjtalk_worker/worker_server.py +++ b/text/pyopenjtalk_worker/worker_server.py @@ -2,7 +2,6 @@ import pyopenjtalk import socket import select - from .worker_common import ( ConnectionClosedException, RequestType, @@ -88,12 +87,13 @@ class WorkerServer: logger.info("close connection") continue - logger.debug(f"receive request: {request}") + logger.trace(f"server received request: {request}") response = self.handle_request(request) - logger.debug(f"send response: {response}") + 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" From 2ab025d5fe4a21522ffe7c9ba7d12a0a061cfa1d Mon Sep 17 00:00:00 2001 From: kale4eat Date: Wed, 6 Mar 2024 09:45:25 +0900 Subject: [PATCH 05/11] Run in a separate process group to avoid receiving signals by Ctrl + C Minor Correction: * logging * declaration of terminate --- text/pyopenjtalk_worker/__init__.py | 34 +++++++++++++++++++++-------- 1 file changed, 25 insertions(+), 9 deletions(-) diff --git a/text/pyopenjtalk_worker/__init__.py b/text/pyopenjtalk_worker/__init__.py index ef5873c..3fd9971 100644 --- a/text/pyopenjtalk_worker/__init__.py +++ b/text/pyopenjtalk_worker/__init__.py @@ -55,8 +55,8 @@ def initialize(port: int = WOKER_PORT): import sys import atexit - global WORKER_CLIENT logger.debug("initialize") + global WORKER_CLIENT if WORKER_CLIENT: return @@ -71,7 +71,16 @@ def initialize(port: int = WOKER_PORT): worker_pkg_path = os.path.relpath( os.path.dirname(__file__), os.getcwd() ).replace(os.sep, ".") - subprocess.Popen([sys.executable, "-m", worker_pkg_path, "--port", str(port)]) + args = [sys.executable, "-m", worker_pkg_path, "--port", str(port)] + # new session, new process group + if sys.platform.startswith("win"): + cf = subprocess.DETACHED_PROCESS | subprocess.CREATE_NEW_PROCESS_GROUP # type: ignore + subprocess.Popen(args, creationflags=cf) + 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: @@ -86,15 +95,22 @@ def initialize(port: int = WOKER_PORT): raise TimeoutError("サーバーに接続できませんでした") WORKER_CLIENT = client + atexit.register(terminate) - def terminate(): - global WORKER_CLIENT - if not WORKER_CLIENT: - return +# 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().get("client-count") == 1: WORKER_CLIENT.quit_server() - WORKER_CLIENT.close() - WORKER_CLIENT = None + except Exception as e: + logger.error(e) - atexit.register(terminate) + WORKER_CLIENT.close() + WORKER_CLIENT = None From 45c6bde2e75b3115ce32cb50059d783495d8168e Mon Sep 17 00:00:00 2001 From: kale4eat Date: Wed, 6 Mar 2024 11:19:55 +0900 Subject: [PATCH 06/11] In Windows, create new console and hide it Minor Correction: * logging * change status return type --- text/pyopenjtalk_worker/__init__.py | 9 ++++++--- text/pyopenjtalk_worker/worker_client.py | 17 +++++++++++++---- 2 files changed, 19 insertions(+), 7 deletions(-) diff --git a/text/pyopenjtalk_worker/__init__.py b/text/pyopenjtalk_worker/__init__.py index 3fd9971..8a10266 100644 --- a/text/pyopenjtalk_worker/__init__.py +++ b/text/pyopenjtalk_worker/__init__.py @@ -74,8 +74,11 @@ def initialize(port: int = WOKER_PORT): args = [sys.executable, "-m", worker_pkg_path, "--port", str(port)] # new session, new process group if sys.platform.startswith("win"): - cf = subprocess.DETACHED_PROCESS | subprocess.CREATE_NEW_PROCESS_GROUP # type: ignore - subprocess.Popen(args, creationflags=cf) + 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 @@ -107,7 +110,7 @@ def terminate(): # repare for unexpected errors try: - if WORKER_CLIENT.status().get("client-count") == 1: + if WORKER_CLIENT.status() == 1: WORKER_CLIENT.quit_server() except Exception as e: logger.error(e) diff --git a/text/pyopenjtalk_worker/worker_client.py b/text/pyopenjtalk_worker/worker_client.py index 23f7dbe..86d8969 100644 --- a/text/pyopenjtalk_worker/worker_client.py +++ b/text/pyopenjtalk_worker/worker_client.py @@ -38,9 +38,18 @@ class WorkerClient: return response.get("return") def status(self): - send_data(self.sock, {"request-type": RequestType.STATUS}) - return receive_data(self.sock) + 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): - send_data(self.sock, {"request-type": RequestType.QUIT_SERVER}) - receive_data(self.sock) + 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}") From ed90af1b8718fe31f7bf3e7565b8583a35f77113 Mon Sep 17 00:00:00 2001 From: kale4eat Date: Wed, 6 Mar 2024 12:13:27 +0900 Subject: [PATCH 07/11] Enhanced server error handling Add signal handling for when the process is killed --- text/pyopenjtalk_worker/__init__.py | 10 ++++++++++ text/pyopenjtalk_worker/worker_server.py | 6 ++++++ 2 files changed, 16 insertions(+) diff --git a/text/pyopenjtalk_worker/__init__.py b/text/pyopenjtalk_worker/__init__.py index 8a10266..e968aed 100644 --- a/text/pyopenjtalk_worker/__init__.py +++ b/text/pyopenjtalk_worker/__init__.py @@ -54,6 +54,7 @@ def initialize(port: int = WOKER_PORT): import socket import sys import atexit + import signal logger.debug("initialize") global WORKER_CLIENT @@ -100,6 +101,15 @@ def initialize(port: int = WOKER_PORT): WORKER_CLIENT = client atexit.register(terminate) + # when the process is killed + def signal_handler(signum, frame): + with open("signal_handler.txt", mode="w") as f: + + pass + terminate() + + signal.signal(signal.SIGTERM, signal_handler) + # top-level declaration def terminate(): diff --git a/text/pyopenjtalk_worker/worker_server.py b/text/pyopenjtalk_worker/worker_server.py index 13bc771..dc6d476 100644 --- a/text/pyopenjtalk_worker/worker_server.py +++ b/text/pyopenjtalk_worker/worker_server.py @@ -86,6 +86,12 @@ class WorkerServer: self.client_count -= 1 logger.info("close connection") continue + except Exception as e: + sock.close() + sockets.remove(sock) + self.client_count -= 1 + logger.error(e) + continue logger.trace(f"server received request: {request}") From 419d0e5bed6e91028c570dad1282f4cc027cf433 Mon Sep 17 00:00:00 2001 From: kale4eat Date: Wed, 6 Mar 2024 22:38:44 +0900 Subject: [PATCH 08/11] Delete debugging traces --- text/pyopenjtalk_worker/__init__.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/text/pyopenjtalk_worker/__init__.py b/text/pyopenjtalk_worker/__init__.py index e968aed..9677666 100644 --- a/text/pyopenjtalk_worker/__init__.py +++ b/text/pyopenjtalk_worker/__init__.py @@ -103,9 +103,6 @@ def initialize(port: int = WOKER_PORT): # when the process is killed def signal_handler(signum, frame): - with open("signal_handler.txt", mode="w") as f: - - pass terminate() signal.signal(signal.SIGTERM, signal_handler) From 2994873e3885fc1bae602dae8d311838697b788e Mon Sep 17 00:00:00 2001 From: kale4eat Date: Thu, 7 Mar 2024 09:40:21 +0900 Subject: [PATCH 09/11] add no client timeout summarize the except statement at disconnection --- text/pyopenjtalk_worker/worker_server.py | 25 ++++++++++++++++-------- 1 file changed, 17 insertions(+), 8 deletions(-) diff --git a/text/pyopenjtalk_worker/worker_server.py b/text/pyopenjtalk_worker/worker_server.py index dc6d476..6babd91 100644 --- a/text/pyopenjtalk_worker/worker_server.py +++ b/text/pyopenjtalk_worker/worker_server.py @@ -1,6 +1,7 @@ import pyopenjtalk import socket import select +import time from .worker_common import ( ConnectionClosedException, @@ -62,13 +63,23 @@ class WorkerServer: return response - def start_server(self, port: int): + 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: @@ -80,17 +91,15 @@ class WorkerServer: # client try: request = receive_data(sock) - except ConnectionClosedException as e: - sock.close() - sockets.remove(sock) - self.client_count -= 1 - logger.info("close connection") - continue except Exception as e: sock.close() sockets.remove(sock) self.client_count -= 1 - logger.error(e) + # unexpected disconnections + if not isinstance(e, ConnectionClosedException): + logger.error(e) + + logger.info("close connection") continue logger.trace(f"server received request: {request}") From d024b71340630d20a20f4ea4b8c4198e52e1dcc2 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Fri, 8 Mar 2024 10:10:53 +0900 Subject: [PATCH 10/11] Fix typo --- common/log.py | 1 + text/pyopenjtalk_worker/__init__.py | 4 ++-- text/pyopenjtalk_worker/__main__.py | 4 ++-- text/pyopenjtalk_worker/worker_common.py | 2 +- 4 files changed, 6 insertions(+), 5 deletions(-) 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/text/pyopenjtalk_worker/__init__.py b/text/pyopenjtalk_worker/__init__.py index 9677666..7aa0d0e 100644 --- a/text/pyopenjtalk_worker/__init__.py +++ b/text/pyopenjtalk_worker/__init__.py @@ -5,7 +5,7 @@ to avoid user dictionary access error from typing import Optional, Any -from .worker_common import WOKER_PORT +from .worker_common import WORKER_PORT from .worker_client import WorkerClient from common.log import logger @@ -49,7 +49,7 @@ def unset_user_dict(): # initialize module when imported -def initialize(port: int = WOKER_PORT): +def initialize(port: int = WORKER_PORT): import time import socket import sys diff --git a/text/pyopenjtalk_worker/__main__.py b/text/pyopenjtalk_worker/__main__.py index 8b67aa0..3bb6b53 100644 --- a/text/pyopenjtalk_worker/__main__.py +++ b/text/pyopenjtalk_worker/__main__.py @@ -1,12 +1,12 @@ import argparse from .worker_server import WorkerServer -from .worker_common import WOKER_PORT +from .worker_common import WORKER_PORT def main(): parser = argparse.ArgumentParser() - parser.add_argument("--port", type=int, default=WOKER_PORT) + parser.add_argument("--port", type=int, default=WORKER_PORT) args = parser.parse_args() server = WorkerServer() server.start_server(port=args.port) diff --git a/text/pyopenjtalk_worker/worker_common.py b/text/pyopenjtalk_worker/worker_common.py index bea552e..606d0c3 100644 --- a/text/pyopenjtalk_worker/worker_common.py +++ b/text/pyopenjtalk_worker/worker_common.py @@ -3,7 +3,7 @@ from enum import IntEnum, auto import socket import json -WOKER_PORT: Final[int] = 7861 +WORKER_PORT: Final[int] = 7861 HEADER_SIZE: Final[int] = 4 From 25ed226acbe13ba9425169c89bc1452c842f424e Mon Sep 17 00:00:00 2001 From: litagin02 Date: Fri, 8 Mar 2024 11:19:03 +0900 Subject: [PATCH 11/11] 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: