Add: Preparation for ONNX inference support ④
This commit is contained in:
@@ -3,6 +3,10 @@
|
|||||||
"repo_id": "ku-nlp/deberta-v2-large-japanese-char-wwm",
|
"repo_id": "ku-nlp/deberta-v2-large-japanese-char-wwm",
|
||||||
"files": ["pytorch_model.bin"]
|
"files": ["pytorch_model.bin"]
|
||||||
},
|
},
|
||||||
|
"deberta-v2-large-japanese-char-wwm-onnx": {
|
||||||
|
"repo_id": "tsukumijima/deberta-v2-large-japanese-char-wwm-onnx",
|
||||||
|
"files": ["model.onnx"]
|
||||||
|
},
|
||||||
"chinese-roberta-wwm-ext-large": {
|
"chinese-roberta-wwm-ext-large": {
|
||||||
"repo_id": "hfl/chinese-roberta-wwm-ext-large",
|
"repo_id": "hfl/chinese-roberta-wwm-ext-large",
|
||||||
"files": ["pytorch_model.bin"]
|
"files": ["pytorch_model.bin"]
|
||||||
|
|||||||
@@ -18,13 +18,18 @@ class Languages(StrEnum):
|
|||||||
ZH = "ZH"
|
ZH = "ZH"
|
||||||
|
|
||||||
|
|
||||||
# 言語ごとのデフォルトの BERT トークナイザーのパス
|
# 言語ごとのデフォルトの BERT モデルのパス
|
||||||
DEFAULT_BERT_TOKENIZER_PATHS = {
|
DEFAULT_BERT_MODEL_PATHS = {
|
||||||
Languages.JP: BASE_DIR / "bert" / "deberta-v2-large-japanese-char-wwm",
|
Languages.JP: BASE_DIR / "bert" / "deberta-v2-large-japanese-char-wwm",
|
||||||
Languages.EN: BASE_DIR / "bert" / "deberta-v3-large",
|
Languages.EN: BASE_DIR / "bert" / "deberta-v3-large",
|
||||||
Languages.ZH: BASE_DIR / "bert" / "chinese-roberta-wwm-ext-large",
|
Languages.ZH: BASE_DIR / "bert" / "chinese-roberta-wwm-ext-large",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# 言語ごとのデフォルトの BERT モデル (ONNX 版) のパス
|
||||||
|
DEFAULT_ONNX_BERT_MODEL_PATHS = {
|
||||||
|
Languages.JP: BASE_DIR / "bert" / "deberta-v2-large-japanese-char-wwm-onnx",
|
||||||
|
}
|
||||||
|
|
||||||
# デフォルトのユーザー辞書ディレクトリ
|
# デフォルトのユーザー辞書ディレクトリ
|
||||||
## style_bert_vits2.nlp.japanese.user_dict モジュールのデフォルト値として利用される
|
## style_bert_vits2.nlp.japanese.user_dict モジュールのデフォルト値として利用される
|
||||||
## ライブラリとしての利用などで外部のユーザー辞書を指定したい場合は、user_dict 以下の各関数の実行時、引数に辞書データファイルのパスを指定する
|
## ライブラリとしての利用などで外部のユーザー辞書を指定したい場合は、user_dict 以下の各関数の実行時、引数に辞書データファイルのパスを指定する
|
||||||
|
|||||||
@@ -16,13 +16,13 @@ from transformers import (
|
|||||||
AutoModelForMaskedLM,
|
AutoModelForMaskedLM,
|
||||||
AutoTokenizer,
|
AutoTokenizer,
|
||||||
DebertaV2Model,
|
DebertaV2Model,
|
||||||
DebertaV2Tokenizer,
|
DebertaV2TokenizerFast,
|
||||||
PreTrainedModel,
|
PreTrainedModel,
|
||||||
PreTrainedTokenizer,
|
PreTrainedTokenizer,
|
||||||
PreTrainedTokenizerFast,
|
PreTrainedTokenizerFast,
|
||||||
)
|
)
|
||||||
|
|
||||||
from style_bert_vits2.constants import DEFAULT_BERT_TOKENIZER_PATHS, Languages
|
from style_bert_vits2.constants import DEFAULT_BERT_MODEL_PATHS, Languages
|
||||||
from style_bert_vits2.logging import logger
|
from style_bert_vits2.logging import logger
|
||||||
|
|
||||||
|
|
||||||
@@ -31,7 +31,8 @@ __loaded_models: dict[Languages, Union[PreTrainedModel, DebertaV2Model]] = {}
|
|||||||
|
|
||||||
# 各言語ごとのロード済みの BERT トークナイザーを格納する辞書
|
# 各言語ごとのロード済みの BERT トークナイザーを格納する辞書
|
||||||
__loaded_tokenizers: dict[
|
__loaded_tokenizers: dict[
|
||||||
Languages, Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]
|
Languages,
|
||||||
|
Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2TokenizerFast],
|
||||||
] = {}
|
] = {}
|
||||||
|
|
||||||
|
|
||||||
@@ -77,10 +78,10 @@ def load_model(
|
|||||||
|
|
||||||
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
|
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
|
||||||
if pretrained_model_name_or_path is None:
|
if pretrained_model_name_or_path is None:
|
||||||
assert DEFAULT_BERT_TOKENIZER_PATHS[
|
assert DEFAULT_BERT_MODEL_PATHS[
|
||||||
language
|
language
|
||||||
].exists(), f"The default {language} BERT model does not exist on the file system. Please specify the path to the pre-trained model."
|
].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])
|
pretrained_model_name_or_path = str(DEFAULT_BERT_MODEL_PATHS[language])
|
||||||
|
|
||||||
# BERT モデルをロードし、辞書に格納して返す
|
# BERT モデルをロードし、辞書に格納して返す
|
||||||
## 英語のみ DebertaV2Model でロードする必要がある
|
## 英語のみ DebertaV2Model でロードする必要がある
|
||||||
@@ -113,9 +114,9 @@ def load_tokenizer(
|
|||||||
pretrained_model_name_or_path: Optional[str] = None,
|
pretrained_model_name_or_path: Optional[str] = None,
|
||||||
cache_dir: Optional[str] = None,
|
cache_dir: Optional[str] = None,
|
||||||
revision: str = "main",
|
revision: str = "main",
|
||||||
) -> Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]:
|
) -> Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2TokenizerFast]:
|
||||||
"""
|
"""
|
||||||
指定された言語の BERT モデルをロードし、ロード済みの BERT トークナイザーを返す。
|
指定された言語の BERT トークナイザーをロードし、ロード済みの BERT トークナイザーを返す。
|
||||||
一度ロードされていれば、ロード済みの BERT トークナイザーを即座に返す。
|
一度ロードされていれば、ロード済みの BERT トークナイザーを即座に返す。
|
||||||
ライブラリ利用時は常に必ず pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある。
|
ライブラリ利用時は常に必ず pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある。
|
||||||
ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき。
|
ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき。
|
||||||
@@ -143,15 +144,15 @@ def load_tokenizer(
|
|||||||
|
|
||||||
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
|
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
|
||||||
if pretrained_model_name_or_path is None:
|
if pretrained_model_name_or_path is None:
|
||||||
assert DEFAULT_BERT_TOKENIZER_PATHS[
|
assert DEFAULT_BERT_MODEL_PATHS[
|
||||||
language
|
language
|
||||||
].exists(), f"The default {language} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model."
|
].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])
|
pretrained_model_name_or_path = str(DEFAULT_BERT_MODEL_PATHS[language])
|
||||||
|
|
||||||
# BERT トークナイザーをロードし、辞書に格納して返す
|
# BERT トークナイザーをロードし、辞書に格納して返す
|
||||||
## 英語のみ DebertaV2Tokenizer でロードする必要がある
|
## 英語のみ DebertaV2TokenizerFast でロードする必要がある
|
||||||
if language == Languages.EN:
|
if language == Languages.EN:
|
||||||
__loaded_tokenizers[language] = DebertaV2Tokenizer.from_pretrained(
|
__loaded_tokenizers[language] = DebertaV2TokenizerFast.from_pretrained(
|
||||||
pretrained_model_name_or_path,
|
pretrained_model_name_or_path,
|
||||||
cache_dir=cache_dir,
|
cache_dir=cache_dir,
|
||||||
revision=revision,
|
revision=revision,
|
||||||
@@ -161,6 +162,7 @@ def load_tokenizer(
|
|||||||
pretrained_model_name_or_path,
|
pretrained_model_name_or_path,
|
||||||
cache_dir=cache_dir,
|
cache_dir=cache_dir,
|
||||||
revision=revision,
|
revision=revision,
|
||||||
|
use_fast=True, # デフォルトで True だが念のため明示的に指定
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Loaded the {language} BERT tokenizer from {pretrained_model_name_or_path}"
|
f"Loaded the {language} BERT tokenizer from {pretrained_model_name_or_path}"
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ import numpy as np
|
|||||||
from numpy.typing import NDArray
|
from numpy.typing import NDArray
|
||||||
|
|
||||||
from style_bert_vits2.constants import Languages
|
from style_bert_vits2.constants import Languages
|
||||||
from style_bert_vits2.nlp import bert_models
|
from style_bert_vits2.nlp import bert_models, onnx_bert_models
|
||||||
from style_bert_vits2.nlp.japanese.g2p import text_to_sep_kata
|
from style_bert_vits2.nlp.japanese.g2p import text_to_sep_kata
|
||||||
|
|
||||||
|
|
||||||
@@ -112,31 +112,47 @@ def extract_bert_feature_onnx(
|
|||||||
if assist_text:
|
if assist_text:
|
||||||
assist_text = "".join(text_to_sep_kata(assist_text, raise_yomi_error=False)[0])
|
assist_text = "".join(text_to_sep_kata(assist_text, raise_yomi_error=False)[0])
|
||||||
|
|
||||||
tokenizer = Tokenizer.from_file("tokenizer.json")
|
tokenizer = onnx_bert_models.load_tokenizer(Languages.JP)
|
||||||
token_ids = [1]
|
inputs = tokenizer(text, return_tensors="pt")
|
||||||
attention_mask = [1]
|
|
||||||
for word in text:
|
|
||||||
encoded = tokenizer.encode(word)
|
|
||||||
token_ids.extend(encoded.ids[1:-1])
|
|
||||||
attention_mask.extend(encoded.attention_mask[1:-1])
|
|
||||||
|
|
||||||
token_ids.append(2)
|
session = onnx_bert_models.load_model(
|
||||||
attention_mask.append(1)
|
language=Languages.JP,
|
||||||
|
onnx_providers=onnx_providers,
|
||||||
bert_output_name = bert_session.get_outputs()[0].name
|
onnx_provider_options=onnx_provider_options,
|
||||||
res = bert_session.run(
|
)
|
||||||
[bert_output_name],
|
output_name = session.get_outputs()[0].name
|
||||||
|
res = session.run(
|
||||||
|
[output_name],
|
||||||
{
|
{
|
||||||
"input_ids": np.array(token_ids).reshape(1, -1),
|
"input_ids": inputs["input_ids"].detach().numpy(),
|
||||||
"attention_mask": np.array(attention_mask).reshape(1, -1),
|
"attention_mask": inputs["attention_mask"].detach().numpy(),
|
||||||
},
|
},
|
||||||
)[0]
|
)[0]
|
||||||
|
|
||||||
|
style_res_mean = None
|
||||||
|
if assist_text:
|
||||||
|
style_inputs = tokenizer(assist_text, return_tensors="pt")
|
||||||
|
style_res = session.run(
|
||||||
|
[output_name],
|
||||||
|
{
|
||||||
|
"input_ids": style_inputs["input_ids"].detach().numpy(),
|
||||||
|
"attention_mask": style_inputs["attention_mask"].detach().numpy(),
|
||||||
|
},
|
||||||
|
)[0]
|
||||||
|
style_res_mean = np.mean(style_res, axis=0)
|
||||||
|
|
||||||
assert len(word2ph) == len(text) + 2, text
|
assert len(word2ph) == len(text) + 2, text
|
||||||
word2phone = word2ph
|
word2phone = word2ph
|
||||||
phone_level_feature = []
|
phone_level_feature = []
|
||||||
for i in range(len(word2phone)):
|
for i in range(len(word2phone)):
|
||||||
repeat_feature = np.tile(res[i], (word2phone[i], 1))
|
if assist_text:
|
||||||
|
assert style_res_mean is not None
|
||||||
|
repeat_feature = (
|
||||||
|
np.tile(res[i], (word2phone[i], 1)) * (1 - assist_text_weight)
|
||||||
|
+ np.tile(style_res_mean, (word2phone[i], 1)) * assist_text_weight
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
repeat_feature = np.tile(res[i], (word2phone[i], 1))
|
||||||
phone_level_feature.append(repeat_feature)
|
phone_level_feature.append(repeat_feature)
|
||||||
|
|
||||||
phone_level_feature = np.concatenate(phone_level_feature, axis=0)
|
phone_level_feature = np.concatenate(phone_level_feature, axis=0)
|
||||||
|
|||||||
215
style_bert_vits2/nlp/onnx_bert_models.py
Normal file
215
style_bert_vits2/nlp/onnx_bert_models.py
Normal file
@@ -0,0 +1,215 @@
|
|||||||
|
"""
|
||||||
|
Style-Bert-VITS2 の ONNX 推論に必要な各言語ごとの ONNX 版 BERT モデルをロード/取得するためのモジュール。
|
||||||
|
このモジュールは style_bert_vits2.nlp.bert_models での実装を ONNX 推論向けに変更したもの。
|
||||||
|
|
||||||
|
オリジナルの Bert-VITS2 では各言語ごとの BERT モデルが初回インポート時にハードコードされたパスから「暗黙的に」ロードされているが、
|
||||||
|
場合によっては多重にロードされて非効率なほか、BERT モデルのロード元のパスがハードコードされているためライブラリ化ができない。
|
||||||
|
|
||||||
|
そこで、ライブラリの利用前に、音声合成に利用する言語の BERT モデルだけを「明示的に」ロードできるようにした。
|
||||||
|
一度 load_model/tokenizer() で当該言語の BERT モデルがロードされていれば、ライブラリ内部のどこからでもロード済みのモデル/トークナイザーを取得できる。
|
||||||
|
"""
|
||||||
|
|
||||||
|
import gc
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Optional, Sequence, Union
|
||||||
|
|
||||||
|
import onnxruntime
|
||||||
|
from huggingface_hub import hf_hub_download
|
||||||
|
from transformers import (
|
||||||
|
AutoTokenizer,
|
||||||
|
DebertaV2TokenizerFast,
|
||||||
|
PreTrainedTokenizer,
|
||||||
|
PreTrainedTokenizerFast,
|
||||||
|
)
|
||||||
|
|
||||||
|
from style_bert_vits2.constants import DEFAULT_ONNX_BERT_MODEL_PATHS, Languages
|
||||||
|
from style_bert_vits2.logging import logger
|
||||||
|
|
||||||
|
|
||||||
|
# 各言語ごとのロード済みの BERT モデルを格納する辞書
|
||||||
|
__loaded_models: dict[Languages, onnxruntime.InferenceSession] = {}
|
||||||
|
|
||||||
|
# 各言語ごとのロード済みの BERT トークナイザーを格納する辞書
|
||||||
|
__loaded_tokenizers: dict[
|
||||||
|
Languages,
|
||||||
|
Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2TokenizerFast],
|
||||||
|
] = {}
|
||||||
|
|
||||||
|
|
||||||
|
def load_model(
|
||||||
|
language: Languages,
|
||||||
|
pretrained_model_name_or_path: Optional[str] = None,
|
||||||
|
onnx_providers: list[str] = ["CPUExecutionProvider"],
|
||||||
|
onnx_provider_options: Optional[Sequence[dict[str, Any]]] = None,
|
||||||
|
cache_dir: Optional[str] = None,
|
||||||
|
revision: str = "main",
|
||||||
|
) -> onnxruntime.InferenceSession:
|
||||||
|
"""
|
||||||
|
指定された言語の ONNX 版 BERT モデルをロードし、ロード済みの ONNX 版 BERT モデルを返す。
|
||||||
|
一度ロードされていれば、ロード済みの ONNX 版 BERT モデルを即座に返す。
|
||||||
|
ライブラリ利用時は常に必ず pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある。
|
||||||
|
ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき。
|
||||||
|
cache_dir と revision は pretrain_model_name_or_path がリポジトリ名の場合のみ有効。
|
||||||
|
|
||||||
|
Style-Bert-VITS2 では、ONNX 版 BERT モデルに下記の 3 つが利用されている。
|
||||||
|
これ以外の ONNX 版 BERT モデルを指定した場合は正常に動作しない可能性が高い。
|
||||||
|
- 日本語: tsukumijima/deberta-v2-large-japanese-char-wwm-onnx
|
||||||
|
|
||||||
|
Args:
|
||||||
|
language (Languages): ロードする学習済みモデルの対象言語
|
||||||
|
pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None)
|
||||||
|
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
|
||||||
|
onnx_provider_options (Optional[dict[str, Any]]): ONNX 推論で利用する ExecutionProvider のオプション
|
||||||
|
cache_dir (Optional[str]): モデルのキャッシュディレクトリ。指定しない場合はデフォルトのキャッシュディレクトリが利用される (デフォルト: None)
|
||||||
|
revision (str): モデルの Hugging Face 上の Git リビジョン。指定しない場合は最新の main ブランチの内容が利用される (デフォルト: None)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
onnxruntime.InferenceSession: ロード済みの BERT モデル
|
||||||
|
"""
|
||||||
|
|
||||||
|
# すでにロード済みの場合はそのまま返す
|
||||||
|
if language in __loaded_models:
|
||||||
|
return __loaded_models[language]
|
||||||
|
|
||||||
|
# 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} ONNX 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_ONNX_BERT_MODEL_PATHS[language])
|
||||||
|
|
||||||
|
# pretrained_model_name_or_path に Hugging Face のリポジトリ名が指定された場合 (aaaa/bbbb のフォーマットを想定):
|
||||||
|
# 指定された revision の ONNX 版 BERT モデルを cache_dir にダウンロードする (既にダウンロード済みの場合は何も行われない)
|
||||||
|
if len(pretrained_model_name_or_path.split("/")) == 2:
|
||||||
|
model_path = Path(
|
||||||
|
hf_hub_download(
|
||||||
|
repo_id=pretrained_model_name_or_path,
|
||||||
|
filename="model.onnx",
|
||||||
|
cache_dir=cache_dir,
|
||||||
|
revision=revision,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
# pretrained_model_name_or_path にファイルパスが指定された場合:
|
||||||
|
# 既にダウンロード済みという前提のもと、モデルへのローカルパスを model_path に格納する
|
||||||
|
else:
|
||||||
|
model_path = Path(pretrained_model_name_or_path).resolve() / "model.onnx"
|
||||||
|
|
||||||
|
# BERT モデルをロードし、辞書に格納して返す
|
||||||
|
__loaded_models[language] = onnxruntime.InferenceSession(
|
||||||
|
model_path,
|
||||||
|
providers=onnx_providers,
|
||||||
|
provider_options=onnx_provider_options,
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
f"Loaded the {language} ONNX BERT model from {pretrained_model_name_or_path}"
|
||||||
|
)
|
||||||
|
|
||||||
|
return __loaded_models[language]
|
||||||
|
|
||||||
|
|
||||||
|
def load_tokenizer(
|
||||||
|
language: Languages,
|
||||||
|
pretrained_model_name_or_path: Optional[str] = None,
|
||||||
|
cache_dir: Optional[str] = None,
|
||||||
|
revision: str = "main",
|
||||||
|
) -> Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2TokenizerFast]:
|
||||||
|
"""
|
||||||
|
指定された言語の ONNX 版 BERT トークナイザーをロードし、ロード済みの ONNX 版 BERT トークナイザーを返す。
|
||||||
|
一度ロードされていれば、ロード済みの ONNX 版 BERT トークナイザーを即座に返す。
|
||||||
|
ライブラリ利用時は常に必ず pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある。
|
||||||
|
ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき。
|
||||||
|
cache_dir と revision は pretrain_model_name_or_path がリポジトリ名の場合のみ有効。
|
||||||
|
|
||||||
|
Style-Bert-VITS2 では、ONNX 版 BERT モデルに下記の 3 つが利用されている。
|
||||||
|
これ以外の ONNX 版 BERT モデルを指定した場合は正常に動作しない可能性が高い。
|
||||||
|
- 日本語: tsukumijima/deberta-v2-large-japanese-char-wwm-onnx
|
||||||
|
|
||||||
|
Args:
|
||||||
|
language (Languages): ロードする学習済みモデルの対象言語
|
||||||
|
pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None)
|
||||||
|
cache_dir (Optional[str]): モデルのキャッシュディレクトリ。指定しない場合はデフォルトのキャッシュディレクトリが利用される (デフォルト: None)
|
||||||
|
revision (str): モデルの Hugging Face 上の Git リビジョン。指定しない場合は最新の main ブランチの内容が利用される (デフォルト: None)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]: ロード済みの BERT トークナイザー
|
||||||
|
"""
|
||||||
|
|
||||||
|
# すでにロード済みの場合はそのまま返す
|
||||||
|
if language in __loaded_tokenizers:
|
||||||
|
return __loaded_tokenizers[language]
|
||||||
|
|
||||||
|
# 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} 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_ONNX_BERT_MODEL_PATHS[language])
|
||||||
|
|
||||||
|
# BERT トークナイザーをロードし、辞書に格納して返す
|
||||||
|
## 英語のみ DebertaV2TokenizerFast でロードする必要がある
|
||||||
|
if language == Languages.EN:
|
||||||
|
__loaded_tokenizers[language] = DebertaV2TokenizerFast.from_pretrained(
|
||||||
|
pretrained_model_name_or_path,
|
||||||
|
cache_dir=cache_dir,
|
||||||
|
revision=revision,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
__loaded_tokenizers[language] = AutoTokenizer.from_pretrained(
|
||||||
|
pretrained_model_name_or_path,
|
||||||
|
cache_dir=cache_dir,
|
||||||
|
revision=revision,
|
||||||
|
use_fast=True, # デフォルトで True だが念のため明示的に指定
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
f"Loaded the {language} ONNX BERT tokenizer from {pretrained_model_name_or_path}"
|
||||||
|
)
|
||||||
|
|
||||||
|
return __loaded_tokenizers[language]
|
||||||
|
|
||||||
|
|
||||||
|
def unload_model(language: Languages) -> None:
|
||||||
|
"""
|
||||||
|
指定された言語の BERT モデルをアンロードする。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
language (Languages): アンロードする BERT モデルの言語
|
||||||
|
"""
|
||||||
|
|
||||||
|
if language in __loaded_models:
|
||||||
|
del __loaded_models[language]
|
||||||
|
gc.collect()
|
||||||
|
logger.info(f"Unloaded the {language} ONNX BERT model")
|
||||||
|
|
||||||
|
|
||||||
|
def unload_tokenizer(language: Languages) -> None:
|
||||||
|
"""
|
||||||
|
指定された言語の BERT トークナイザーをアンロードする。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
language (Languages): アンロードする BERT トークナイザーの言語
|
||||||
|
"""
|
||||||
|
|
||||||
|
if language in __loaded_tokenizers:
|
||||||
|
del __loaded_tokenizers[language]
|
||||||
|
gc.collect()
|
||||||
|
logger.info(f"Unloaded the {language} ONNX BERT tokenizer")
|
||||||
|
|
||||||
|
|
||||||
|
def unload_all_models() -> None:
|
||||||
|
"""
|
||||||
|
すべての BERT モデルをアンロードする。
|
||||||
|
"""
|
||||||
|
|
||||||
|
for language in list(__loaded_models.keys()):
|
||||||
|
unload_model(language)
|
||||||
|
logger.info("Unloaded all ONNX BERT models")
|
||||||
|
|
||||||
|
|
||||||
|
def unload_all_tokenizers() -> None:
|
||||||
|
"""
|
||||||
|
すべての BERT トークナイザーをアンロードする。
|
||||||
|
"""
|
||||||
|
|
||||||
|
for language in list(__loaded_tokenizers.keys()):
|
||||||
|
unload_tokenizer(language)
|
||||||
|
logger.info("Unloaded all ONNX BERT tokenizers")
|
||||||
Reference in New Issue
Block a user