Apply black formatter
This commit is contained in:
@@ -74,16 +74,19 @@ def clean_text(
|
||||
if language == Languages.JP:
|
||||
from style_bert_vits2.nlp.japanese.g2p import g2p
|
||||
from style_bert_vits2.nlp.japanese.normalizer import normalize_text
|
||||
|
||||
norm_text = normalize_text(text)
|
||||
phones, tones, word2ph = g2p(norm_text, use_jp_extra, raise_yomi_error)
|
||||
elif language == Languages.EN:
|
||||
from style_bert_vits2.nlp.english.g2p import g2p
|
||||
from style_bert_vits2.nlp.english.normalizer import normalize_text
|
||||
|
||||
norm_text = normalize_text(text)
|
||||
phones, tones, word2ph = g2p(norm_text)
|
||||
elif language == Languages.ZH:
|
||||
from style_bert_vits2.nlp.chinese.g2p import g2p
|
||||
from style_bert_vits2.nlp.chinese.normalizer import normalize_text
|
||||
|
||||
norm_text = normalize_text(text)
|
||||
phones, tones, word2ph = g2p(norm_text)
|
||||
else:
|
||||
@@ -92,7 +95,9 @@ def clean_text(
|
||||
return norm_text, phones, tones, word2ph
|
||||
|
||||
|
||||
def cleaned_text_to_sequence(cleaned_phones: list[str], tones: list[int], language: Languages) -> tuple[list[int], list[int], list[int]]:
|
||||
def cleaned_text_to_sequence(
|
||||
cleaned_phones: list[str], tones: list[int], language: Languages
|
||||
) -> tuple[list[int], list[int], list[int]]:
|
||||
"""
|
||||
テキスト文字列を、テキスト内の記号に対応する一連の ID に変換する
|
||||
|
||||
|
||||
@@ -30,7 +30,9 @@ from style_bert_vits2.logging import logger
|
||||
__loaded_models: dict[Languages, Union[PreTrainedModel, DebertaV2Model]] = {}
|
||||
|
||||
# 各言語ごとのロード済みの BERT トークナイザーを格納する辞書
|
||||
__loaded_tokenizers: dict[Languages, Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]] = {}
|
||||
__loaded_tokenizers: dict[
|
||||
Languages, Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]
|
||||
] = {}
|
||||
|
||||
|
||||
def load_model(
|
||||
@@ -63,18 +65,24 @@ def load_model(
|
||||
|
||||
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
|
||||
if pretrained_model_name_or_path is None:
|
||||
assert DEFAULT_BERT_TOKENIZER_PATHS[language].exists(), \
|
||||
f"The default {language} BERT model does not exist on the file system. Please specify the path to the pre-trained model."
|
||||
assert DEFAULT_BERT_TOKENIZER_PATHS[
|
||||
language
|
||||
].exists(), f"The default {language} BERT model does not exist on the file system. Please specify the path to the pre-trained model."
|
||||
pretrained_model_name_or_path = str(DEFAULT_BERT_TOKENIZER_PATHS[language])
|
||||
|
||||
# BERT モデルをロードし、辞書に格納して返す
|
||||
## 英語のみ DebertaV2Model でロードする必要がある
|
||||
if language == Languages.EN:
|
||||
model = cast(DebertaV2Model, DebertaV2Model.from_pretrained(pretrained_model_name_or_path))
|
||||
model = cast(
|
||||
DebertaV2Model,
|
||||
DebertaV2Model.from_pretrained(pretrained_model_name_or_path),
|
||||
)
|
||||
else:
|
||||
model = AutoModelForMaskedLM.from_pretrained(pretrained_model_name_or_path)
|
||||
__loaded_models[language] = model
|
||||
logger.info(f"Loaded the {language} BERT model from {pretrained_model_name_or_path}")
|
||||
logger.info(
|
||||
f"Loaded the {language} BERT model from {pretrained_model_name_or_path}"
|
||||
)
|
||||
|
||||
return model
|
||||
|
||||
@@ -109,8 +117,9 @@ def load_tokenizer(
|
||||
|
||||
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
|
||||
if pretrained_model_name_or_path is None:
|
||||
assert DEFAULT_BERT_TOKENIZER_PATHS[language].exists(), \
|
||||
f"The default {language} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model."
|
||||
assert DEFAULT_BERT_TOKENIZER_PATHS[
|
||||
language
|
||||
].exists(), f"The default {language} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model."
|
||||
pretrained_model_name_or_path = str(DEFAULT_BERT_TOKENIZER_PATHS[language])
|
||||
|
||||
# BERT トークナイザーをロードし、辞書に格納して返す
|
||||
@@ -120,7 +129,9 @@ def load_tokenizer(
|
||||
else:
|
||||
tokenizer = AutoTokenizer.from_pretrained(pretrained_model_name_or_path)
|
||||
__loaded_tokenizers[language] = tokenizer
|
||||
logger.info(f"Loaded the {language} BERT tokenizer from {pretrained_model_name_or_path}")
|
||||
logger.info(
|
||||
f"Loaded the {language} BERT tokenizer from {pretrained_model_name_or_path}"
|
||||
)
|
||||
|
||||
return tokenizer
|
||||
|
||||
|
||||
@@ -95,7 +95,11 @@ def __g2p(segments: list[str]) -> tuple[list[str], list[int], list[int]]:
|
||||
if pinyin[0] in single_rep_map.keys():
|
||||
pinyin = single_rep_map[pinyin[0]] + pinyin[1:]
|
||||
|
||||
assert pinyin in __PINYIN_TO_SYMBOL_MAP.keys(), (pinyin, seg, raw_pinyin)
|
||||
assert pinyin in __PINYIN_TO_SYMBOL_MAP.keys(), (
|
||||
pinyin,
|
||||
seg,
|
||||
raw_pinyin,
|
||||
)
|
||||
phone = __PINYIN_TO_SYMBOL_MAP[pinyin].split(" ")
|
||||
word2ph.append(len(phone))
|
||||
|
||||
@@ -125,7 +129,7 @@ if __name__ == "__main__":
|
||||
text = normalize_text(text)
|
||||
print(text)
|
||||
phones, tones, word2ph = g2p(text)
|
||||
bert = extract_bert_feature(text, word2ph, 'cuda')
|
||||
bert = extract_bert_feature(text, word2ph, "cuda")
|
||||
|
||||
print(phones, tones, word2ph, bert.shape)
|
||||
|
||||
|
||||
@@ -121,7 +121,9 @@ def __expand_number(m: re.Match[str]) -> str:
|
||||
else:
|
||||
return __INFLECT.number_to_words(
|
||||
num, andword="", zero="oh", group=2 # type: ignore
|
||||
).replace(", ", " ") # type: ignore
|
||||
).replace(
|
||||
", ", " "
|
||||
) # type: ignore
|
||||
else:
|
||||
return __INFLECT.number_to_words(num, andword="") # type: ignore
|
||||
|
||||
|
||||
@@ -10,9 +10,7 @@ from style_bert_vits2.nlp.symbols import PUNCTUATIONS
|
||||
|
||||
|
||||
def g2p(
|
||||
norm_text: str,
|
||||
use_jp_extra: bool = True,
|
||||
raise_yomi_error: bool = False
|
||||
norm_text: str, use_jp_extra: bool = True, raise_yomi_error: bool = False
|
||||
) -> tuple[list[str], list[int], list[int]]:
|
||||
"""
|
||||
他で使われるメインの関数。`normalize_text()` で正規化された `norm_text` を受け取り、
|
||||
@@ -93,8 +91,7 @@ def g2p(
|
||||
|
||||
|
||||
def text_to_sep_kata(
|
||||
norm_text: str,
|
||||
raise_yomi_error: bool = False
|
||||
norm_text: str, raise_yomi_error: bool = False
|
||||
) -> tuple[list[str], list[str]]:
|
||||
"""
|
||||
`normalize_text` で正規化済みの `norm_text` を受け取り、それを単語分割し、
|
||||
@@ -212,7 +209,9 @@ def __g2phone_tone_wo_punct(text: str) -> list[tuple[str, int]]:
|
||||
return result
|
||||
|
||||
|
||||
def __pyopenjtalk_g2p_prosody(text: str, drop_unvoiced_vowels: bool = True) -> list[str]:
|
||||
def __pyopenjtalk_g2p_prosody(
|
||||
text: str, drop_unvoiced_vowels: bool = True
|
||||
) -> list[str]:
|
||||
"""
|
||||
ESPnet の実装から引用、変更点無し。「ん」は「N」なことに注意。
|
||||
ref: https://github.com/espnet/espnet/blob/master/espnet2/text/phoneme_tokenizer.py
|
||||
@@ -414,8 +413,7 @@ def __kata_to_phoneme_list(text: str) -> list[str]:
|
||||
|
||||
|
||||
def __align_tones(
|
||||
phones_with_punct: list[str],
|
||||
phone_tone_list: list[tuple[str, int]]
|
||||
phones_with_punct: list[str], phone_tone_list: list[tuple[str, int]]
|
||||
) -> list[tuple[str, int]]:
|
||||
"""
|
||||
例: …私は、、そう思う。
|
||||
|
||||
@@ -34,11 +34,13 @@ def phone_tone2kata_tone(phone_tone: list[tuple[str, int]]) -> list[tuple[str, i
|
||||
"""
|
||||
|
||||
# 子音の集合
|
||||
CONSONANTS = set([
|
||||
consonant
|
||||
for consonant, _ in MORA_KATA_TO_MORA_PHONEMES.values()
|
||||
if consonant is not None
|
||||
])
|
||||
CONSONANTS = set(
|
||||
[
|
||||
consonant
|
||||
for consonant, _ in MORA_KATA_TO_MORA_PHONEMES.values()
|
||||
if consonant is not None
|
||||
]
|
||||
)
|
||||
|
||||
phone_tone = phone_tone[1:] # 最初の("_", 0)を無視
|
||||
phones = [phone for phone, _ in phone_tone]
|
||||
|
||||
@@ -25,6 +25,7 @@ def run_frontend(text: str) -> list[dict[str, Any]]:
|
||||
else:
|
||||
# without worker
|
||||
import pyopenjtalk
|
||||
|
||||
return pyopenjtalk.run_frontend(text)
|
||||
|
||||
|
||||
@@ -36,6 +37,7 @@ def make_label(njd_features: Any) -> list[str]:
|
||||
else:
|
||||
# without worker
|
||||
import pyopenjtalk
|
||||
|
||||
return pyopenjtalk.make_label(njd_features)
|
||||
|
||||
|
||||
@@ -45,6 +47,7 @@ def mecab_dict_index(path: str, out_path: str, dn_mecab: Optional[str] = None) -
|
||||
else:
|
||||
# without worker
|
||||
import pyopenjtalk
|
||||
|
||||
pyopenjtalk.mecab_dict_index(path, out_path, dn_mecab)
|
||||
|
||||
|
||||
@@ -54,6 +57,7 @@ def update_global_jtalk_with_user_dict(path: str) -> None:
|
||||
else:
|
||||
# without worker
|
||||
import pyopenjtalk
|
||||
|
||||
pyopenjtalk.update_global_jtalk_with_user_dict(path)
|
||||
|
||||
|
||||
@@ -63,6 +67,7 @@ def unset_user_dict() -> None:
|
||||
else:
|
||||
# without worker
|
||||
import pyopenjtalk
|
||||
|
||||
pyopenjtalk.unset_user_dict()
|
||||
|
||||
|
||||
@@ -102,7 +107,12 @@ def initialize_worker(port: int = WORKER_PORT) -> None:
|
||||
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)
|
||||
subprocess.Popen(
|
||||
args,
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
start_new_session=True,
|
||||
)
|
||||
|
||||
# wait until server listening
|
||||
count = 0
|
||||
|
||||
@@ -2,12 +2,15 @@ import socket
|
||||
from typing import Any, cast
|
||||
|
||||
from style_bert_vits2.logging import logger
|
||||
from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_common import RequestType, receive_data, send_data
|
||||
from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_common import (
|
||||
RequestType,
|
||||
receive_data,
|
||||
send_data,
|
||||
)
|
||||
|
||||
|
||||
class WorkerClient:
|
||||
""" pyopenjtalk worker client """
|
||||
|
||||
"""pyopenjtalk worker client"""
|
||||
|
||||
def __init__(self, port: int) -> None:
|
||||
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||
@@ -16,19 +19,15 @@ class WorkerClient:
|
||||
sock.connect((socket.gethostname(), port))
|
||||
self.sock = sock
|
||||
|
||||
|
||||
def __enter__(self) -> "WorkerClient":
|
||||
return self
|
||||
|
||||
|
||||
def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> None:
|
||||
self.close()
|
||||
|
||||
|
||||
def close(self) -> None:
|
||||
self.sock.close()
|
||||
|
||||
|
||||
def dispatch_pyopenjtalk(self, func: str, *args: Any, **kwargs: Any) -> Any:
|
||||
data = {
|
||||
"request-type": RequestType.PYOPENJTALK,
|
||||
@@ -43,7 +42,6 @@ class WorkerClient:
|
||||
logger.trace(f"client received response: {response}")
|
||||
return response.get("return")
|
||||
|
||||
|
||||
def status(self) -> int:
|
||||
data = {"request-type": RequestType.STATUS}
|
||||
logger.trace(f"client sends request: {data}")
|
||||
@@ -53,7 +51,6 @@ class WorkerClient:
|
||||
logger.trace(f"client received response: {response}")
|
||||
return cast(int, response.get("client-count"))
|
||||
|
||||
|
||||
def quit_server(self) -> None:
|
||||
data = {"request-type": RequestType.QUIT_SERVER}
|
||||
logger.trace(f"client sends request: {data}")
|
||||
|
||||
@@ -26,14 +26,12 @@ PYOPENJTALK_FUNC_DICT = {
|
||||
|
||||
|
||||
class WorkerServer:
|
||||
""" pyopenjtalk worker server """
|
||||
|
||||
"""pyopenjtalk worker server"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.client_count: int = 0
|
||||
self.quit: bool = False
|
||||
|
||||
|
||||
def handle_request(self, request: dict[str, Any]) -> dict[str, Any]:
|
||||
request_type = None
|
||||
try:
|
||||
@@ -70,7 +68,6 @@ class WorkerServer:
|
||||
|
||||
return response
|
||||
|
||||
|
||||
def start_server(self, port: int, no_client_timeout: int = 30) -> None:
|
||||
logger.info("start pyopenjtalk worker server")
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as server_socket:
|
||||
|
||||
@@ -18,7 +18,11 @@ from fastapi import HTTPException
|
||||
from style_bert_vits2.constants import DEFAULT_USER_DICT_DIR
|
||||
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
||||
from style_bert_vits2.nlp.japanese.user_dict.word_model import UserDictWord, WordTypes
|
||||
from style_bert_vits2.nlp.japanese.user_dict.part_of_speech_data import MAX_PRIORITY, MIN_PRIORITY, part_of_speech_data
|
||||
from style_bert_vits2.nlp.japanese.user_dict.part_of_speech_data import (
|
||||
MAX_PRIORITY,
|
||||
MIN_PRIORITY,
|
||||
part_of_speech_data,
|
||||
)
|
||||
|
||||
# root_dir = engine_root()
|
||||
# save_dir = get_save_dir()
|
||||
@@ -26,9 +30,13 @@ from style_bert_vits2.nlp.japanese.user_dict.part_of_speech_data import MAX_PRIO
|
||||
# if not save_dir.is_dir():
|
||||
# save_dir.mkdir(parents=True)
|
||||
|
||||
default_dict_path = DEFAULT_USER_DICT_DIR / "default.csv" # VOICEVOXデフォルト辞書ファイルのパス
|
||||
default_dict_path = (
|
||||
DEFAULT_USER_DICT_DIR / "default.csv"
|
||||
) # VOICEVOXデフォルト辞書ファイルのパス
|
||||
user_dict_path = DEFAULT_USER_DICT_DIR / "user_dict.json" # ユーザー辞書ファイルのパス
|
||||
compiled_dict_path = DEFAULT_USER_DICT_DIR / "user.dic" # コンパイル済み辞書ファイルのパス
|
||||
compiled_dict_path = (
|
||||
DEFAULT_USER_DICT_DIR / "user.dic"
|
||||
) # コンパイル済み辞書ファイルのパス
|
||||
|
||||
|
||||
# # 同時書き込みの制御
|
||||
|
||||
Reference in New Issue
Block a user