Add: Preparation for ONNX inference support ④

This commit is contained in:
tsukumi
2024-09-17 13:25:17 +09:00
parent ebd4551d1e
commit 9de6e04ad7
5 changed files with 272 additions and 30 deletions

View File

@@ -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}"