From 9723f75f9911e2921bd4c6c522c20f0029253de2 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Nov 2024 02:56:05 +0900 Subject: [PATCH] Fix: g2p process fails in PyTorch-independent environment --- style_bert_vits2/nlp/bert_models.py | 32 +++++++++++++++++----- style_bert_vits2/nlp/onnx_bert_models.py | 34 +++++++++++++++++------- 2 files changed, 50 insertions(+), 16 deletions(-) diff --git a/style_bert_vits2/nlp/bert_models.py b/style_bert_vits2/nlp/bert_models.py index 2775e29..9c05786 100644 --- a/style_bert_vits2/nlp/bert_models.py +++ b/style_bert_vits2/nlp/bert_models.py @@ -26,6 +26,7 @@ from transformers import ( from style_bert_vits2.constants import DEFAULT_BERT_MODEL_PATHS, Languages from style_bert_vits2.logging import logger +from style_bert_vits2.nlp import onnx_bert_models if TYPE_CHECKING: @@ -84,9 +85,8 @@ def load_model( # pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用 if pretrained_model_name_or_path is None: - assert DEFAULT_BERT_MODEL_PATHS[ - language - ].exists(), f"The default {language.name} BERT model does not exist on the file system. Please specify the path to the pre-trained model." + assert DEFAULT_BERT_MODEL_PATHS[language].exists(), \ + f"The default {language.name} BERT model does not exist on the file system. Please specify the path to the pre-trained model." # fmt: skip pretrained_model_name_or_path = str(DEFAULT_BERT_MODEL_PATHS[language]) # BERT モデルをロードし、辞書に格納して返す @@ -151,9 +151,13 @@ def load_tokenizer( # pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用 if pretrained_model_name_or_path is None: - assert DEFAULT_BERT_MODEL_PATHS[ - language - ].exists(), f"The default {language.name} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model." + # ライブラリ利用時、特例的にこの状況で ONNX 版 BERT トークナイザーがロードされている場合はそのまま返す + ## ONNX 版 BERT トークナイザー単独で g2p 処理を行うために必要 (各言語の g2p.py はこの関数に依存している) + ## 設計的には微妙だがこの方が差異を吸収できて手っ取り早い + if DEFAULT_BERT_MODEL_PATHS[language].exists() is False and onnx_bert_models.is_tokenizer_loaded(language): # fmt: skip + return onnx_bert_models.load_tokenizer(language) + assert DEFAULT_BERT_MODEL_PATHS[language].exists(), \ + f"The default {language.name} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model." # fmt: skip pretrained_model_name_or_path = str(DEFAULT_BERT_MODEL_PATHS[language]) # BERT トークナイザーをロードし、辞書に格納して返す @@ -205,6 +209,22 @@ def transfer_model(language: Languages, device: str) -> None: ) +def is_model_loaded(language: Languages) -> bool: + """ + 指定された言語の BERT モデルがロード済みかどうかを返す。 + """ + + return language in __loaded_models + + +def is_tokenizer_loaded(language: Languages) -> bool: + """ + 指定された言語の BERT トークナイザーがロード済みかどうかを返す。 + """ + + return language in __loaded_tokenizers + + def unload_model(language: Languages) -> None: """ 指定された言語の BERT モデルをアンロードする。 diff --git a/style_bert_vits2/nlp/onnx_bert_models.py b/style_bert_vits2/nlp/onnx_bert_models.py index 04c8d65..d4b095c 100644 --- a/style_bert_vits2/nlp/onnx_bert_models.py +++ b/style_bert_vits2/nlp/onnx_bert_models.py @@ -72,9 +72,8 @@ def load_model( # pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用 if pretrained_model_name_or_path is None: - assert DEFAULT_ONNX_BERT_MODEL_PATHS[ - language - ].exists(), f"The default {language.name} ONNX BERT model does not exist on the file system. Please specify the path to the pre-trained model." + assert DEFAULT_ONNX_BERT_MODEL_PATHS[language].exists(), \ + f"The default {language.name} ONNX BERT model does not exist on the file system. Please specify the path to the pre-trained model." # fmt: skip pretrained_model_name_or_path = str(DEFAULT_ONNX_BERT_MODEL_PATHS[language]) # pretrained_model_name_or_path に Hugging Face のリポジトリ名が指定された場合 (aaaa/bbbb のフォーマットを想定): @@ -148,9 +147,8 @@ def load_tokenizer( # pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用 if pretrained_model_name_or_path is None: - assert DEFAULT_ONNX_BERT_MODEL_PATHS[ - language - ].exists(), f"The default {language.name} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model." + assert DEFAULT_ONNX_BERT_MODEL_PATHS[language].exists(), \ + f"The default {language.name} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model." # fmt: skip pretrained_model_name_or_path = str(DEFAULT_ONNX_BERT_MODEL_PATHS[language]) # BERT トークナイザーをロードし、辞書に格納して返す @@ -175,9 +173,25 @@ def load_tokenizer( return __loaded_tokenizers[language] +def is_model_loaded(language: Languages) -> bool: + """ + 指定された言語の ONNX 版 BERT モデルがロード済みかどうかを返す。 + """ + + return language in __loaded_models + + +def is_tokenizer_loaded(language: Languages) -> bool: + """ + 指定された言語の ONNX 版 BERT トークナイザーがロード済みかどうかを返す。 + """ + + return language in __loaded_tokenizers + + def unload_model(language: Languages) -> None: """ - 指定された言語の BERT モデルをアンロードする。 + 指定された言語の ONNX 版 BERT モデルをアンロードする。 Args: language (Languages): アンロードする BERT モデルの言語 @@ -191,7 +205,7 @@ def unload_model(language: Languages) -> None: def unload_tokenizer(language: Languages) -> None: """ - 指定された言語の BERT トークナイザーをアンロードする。 + 指定された言語の ONNX 版 BERT トークナイザーをアンロードする。 Args: language (Languages): アンロードする BERT トークナイザーの言語 @@ -205,7 +219,7 @@ def unload_tokenizer(language: Languages) -> None: def unload_all_models() -> None: """ - すべての BERT モデルをアンロードする。 + すべての ONNX 版 BERT モデルをアンロードする。 """ for language in list(__loaded_models.keys()): @@ -215,7 +229,7 @@ def unload_all_models() -> None: def unload_all_tokenizers() -> None: """ - すべての BERT トークナイザーをアンロードする。 + すべての ONNX 版 BERT トークナイザーをアンロードする。 """ for language in list(__loaded_tokenizers.keys()):