Add: Preparation for ONNX inference support ④
This commit is contained in:
@@ -16,13 +16,13 @@ from transformers import (
|
||||
AutoModelForMaskedLM,
|
||||
AutoTokenizer,
|
||||
DebertaV2Model,
|
||||
DebertaV2Tokenizer,
|
||||
DebertaV2TokenizerFast,
|
||||
PreTrainedModel,
|
||||
PreTrainedTokenizer,
|
||||
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
|
||||
|
||||
|
||||
@@ -31,7 +31,8 @@ __loaded_models: dict[Languages, Union[PreTrainedModel, DebertaV2Model]] = {}
|
||||
|
||||
# 各言語ごとのロード済みの BERT トークナイザーを格納する辞書
|
||||
__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 が指定されていない場合はデフォルトのパスを利用
|
||||
if pretrained_model_name_or_path is None:
|
||||
assert DEFAULT_BERT_TOKENIZER_PATHS[
|
||||
assert DEFAULT_BERT_MODEL_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])
|
||||
pretrained_model_name_or_path = str(DEFAULT_BERT_MODEL_PATHS[language])
|
||||
|
||||
# BERT モデルをロードし、辞書に格納して返す
|
||||
## 英語のみ DebertaV2Model でロードする必要がある
|
||||
@@ -113,9 +114,9 @@ def load_tokenizer(
|
||||
pretrained_model_name_or_path: Optional[str] = None,
|
||||
cache_dir: Optional[str] = None,
|
||||
revision: str = "main",
|
||||
) -> Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]:
|
||||
) -> Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2TokenizerFast]:
|
||||
"""
|
||||
指定された言語の BERT モデルをロードし、ロード済みの BERT トークナイザーを返す。
|
||||
指定された言語の BERT トークナイザーをロードし、ロード済みの BERT トークナイザーを返す。
|
||||
一度ロードされていれば、ロード済みの BERT トークナイザーを即座に返す。
|
||||
ライブラリ利用時は常に必ず pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある。
|
||||
ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき。
|
||||
@@ -143,15 +144,15 @@ def load_tokenizer(
|
||||
|
||||
# pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用
|
||||
if pretrained_model_name_or_path is None:
|
||||
assert DEFAULT_BERT_TOKENIZER_PATHS[
|
||||
assert DEFAULT_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_BERT_TOKENIZER_PATHS[language])
|
||||
pretrained_model_name_or_path = str(DEFAULT_BERT_MODEL_PATHS[language])
|
||||
|
||||
# BERT トークナイザーをロードし、辞書に格納して返す
|
||||
## 英語のみ DebertaV2Tokenizer でロードする必要がある
|
||||
## 英語のみ DebertaV2TokenizerFast でロードする必要がある
|
||||
if language == Languages.EN:
|
||||
__loaded_tokenizers[language] = DebertaV2Tokenizer.from_pretrained(
|
||||
__loaded_tokenizers[language] = DebertaV2TokenizerFast.from_pretrained(
|
||||
pretrained_model_name_or_path,
|
||||
cache_dir=cache_dir,
|
||||
revision=revision,
|
||||
@@ -161,6 +162,7 @@ def load_tokenizer(
|
||||
pretrained_model_name_or_path,
|
||||
cache_dir=cache_dir,
|
||||
revision=revision,
|
||||
use_fast=True, # デフォルトで True だが念のため明示的に指定
|
||||
)
|
||||
logger.info(
|
||||
f"Loaded the {language} BERT tokenizer from {pretrained_model_name_or_path}"
|
||||
|
||||
Reference in New Issue
Block a user