Refactor: add style_bert_vits2/text_processing/bert_models.py to hold loaded BERT models/tokenizer and replace all from_pretrained() to load_model/load_tokenizer

This commit is contained in:
tsukumi
2024-03-07 02:31:30 +00:00
parent e826faf62e
commit c3c0dd8b32
13 changed files with 217 additions and 728 deletions

View File

@@ -0,0 +1,123 @@
"""
Style-Bert-VITS2 の学習・推論に必要な各言語ごとの BERT モデルをロード/取得するためのモジュール。
オリジナルの Bert-VITS2 では各言語ごとの BERT モデルが初回インポート時にハードコードされたパスから「暗黙的に」ロードされているが、
場合によっては多重にロードされて非効率なほか、BERT モデルのロード元のパスがハードコードされているためライブラリ化ができない。
そこで、ライブラリの利用前に、音声合成に利用する言語の BERT モデルだけを「明示的に」ロードできるようにした。
一度 load_tokenizer() で当該言語の BERT モデルがロードされていれば、ライブラリ内部のどこからでもロード済みのモデル/トークナイザーを取得できる。
"""
from typing import cast
from transformers import (
AutoModelForMaskedLM,
AutoTokenizer,
DebertaV2Model,
DebertaV2Tokenizer,
PreTrainedModel,
PreTrainedTokenizer,
PreTrainedTokenizerFast,
)
from style_bert_vits2.constants import DEFAULT_BERT_TOKENIZER_PATHS, Languages
from style_bert_vits2.logging import logger
# 各言語ごとのロード済みの BERT モデルを格納する辞書
loaded_models: dict[Languages, PreTrainedModel | DebertaV2Model] = {}
# 各言語ごとのロード済みの BERT トークナイザーを格納する辞書
loaded_tokenizers: dict[Languages, PreTrainedTokenizer | PreTrainedTokenizerFast | DebertaV2Tokenizer] = {}
def load_model(
language: Languages,
pretrained_model_name_or_path: str | None = None,
) -> PreTrainedModel | DebertaV2Model:
"""
指定された言語の BERT モデルをロードし、ロード済みの BERT モデルを返す
一度ロードされていれば、ロード済みの BERT モデルを即座に返す
ライブラリ利用時は常に pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある
ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき
Style-Bert-VITS2 では、BERT モデルに下記の 3 つが利用されている
これ以外の BERT モデルを指定した場合は正常に動作しない可能性が高い
- 日本語: ku-nlp/deberta-v2-large-japanese-char-wwm
- 英語: microsoft/deberta-v3-large
- 中国語: hfl/chinese-roberta-wwm-ext-large
Args:
language (Languages): ロードする学習済みモデルの対象言語
pretrained_model_name_or_path (str | None): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None)
Returns:
PreTrainedModel | DebertaV2Model: ロード済みの 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_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))
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}")
return model
def load_tokenizer(
language: Languages,
pretrained_model_name_or_path: str | None = None,
) -> PreTrainedTokenizer | PreTrainedTokenizerFast | DebertaV2Tokenizer:
"""
指定された言語の BERT モデルをロードし、ロード済みの BERT トークナイザーを返す
一度ロードされていれば、ロード済みの BERT トークナイザーを即座に返す
ライブラリ利用時は常に pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある
ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき
Style-Bert-VITS2 では、BERT モデルに下記の 3 つが利用されている
これ以外の BERT モデルを指定した場合は正常に動作しない可能性が高い
- 日本語: ku-nlp/deberta-v2-large-japanese-char-wwm
- 英語: microsoft/deberta-v3-large
- 中国語: hfl/chinese-roberta-wwm-ext-large
Args:
language (Languages): ロードする学習済みモデルの対象言語
pretrained_model_name_or_path (str | None): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None)
Returns:
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_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 トークナイザーをロードし、辞書に格納して返す
## 英語のみ DebertaV2Tokenizer でロードする必要がある
if language == Languages.EN:
tokenizer = DebertaV2Tokenizer.from_pretrained(pretrained_model_name_or_path)
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}")
return tokenizer

View File

@@ -1,9 +1,9 @@
from typing import Literal
from style_bert_vits2.constants import Languages
def clean_text(
text: str,
language: Literal["JP", "EN", "ZH"],
language: Languages,
use_jp_extra: bool = True,
raise_yomi_error: bool = False,
) -> tuple[str, list[str], list[int], list[int]]:
@@ -12,7 +12,7 @@ def clean_text(
Args:
text (str): クリーニングするテキスト
language (Literal["JP", "EN", "ZH"]): テキストの言語
language (Languages): テキストの言語
use_jp_extra (bool, optional): テキストが日本語の場合に JP-Extra モデルを利用するかどうか。Defaults to True.
raise_yomi_error (bool, optional): False の場合、読めない文字が消えたような扱いとして処理される。Defaults to False.
@@ -21,22 +21,16 @@ def clean_text(
"""
# Changed to import inside if condition to avoid unnecessary import
if language == "JP":
from transformers import AutoTokenizer
if language == Languages.JP:
from style_bert_vits2.text_processing.japanese.g2p import g2p
from style_bert_vits2.text_processing.japanese.normalizer import normalize_text
norm_text = normalize_text(text)
phones, tones, word2ph = g2p(
norm_text,
tokenizer = AutoTokenizer.from_pretrained("./bert/deberta-v2-large-japanese-char-wwm"), # 暫定的にここで指定
use_jp_extra = use_jp_extra,
raise_yomi_error = raise_yomi_error,
)
elif language == "EN":
phones, tones, word2ph = g2p(norm_text, use_jp_extra, raise_yomi_error)
elif language == Languages.EN:
from ...text import english as language_module
norm_text = language_module.normalize_text(text)
phones, tones, word2ph = language_module.g2p(norm_text)
elif language == "ZH":
elif language == Languages.ZH:
from ...text import chinese as language_module
norm_text = language_module.normalize_text(text)
phones, tones, word2ph = language_module.g2p(norm_text)

View File

@@ -1,8 +1,9 @@
import pyopenjtalk
import re
from transformers import PreTrainedTokenizer, PreTrainedTokenizerFast
from style_bert_vits2.constants import Languages
from style_bert_vits2.logging import logger
from style_bert_vits2.text_processing import bert_models
from style_bert_vits2.text_processing.japanese.mora_list import MORA_KATA_TO_MORA_PHONEMES
from style_bert_vits2.text_processing.japanese.normalizer import replace_punctuation
from style_bert_vits2.text_processing.symbols import PUNCTUATIONS
@@ -10,7 +11,6 @@ from style_bert_vits2.text_processing.symbols import PUNCTUATIONS
def g2p(
norm_text: str,
tokenizer: PreTrainedTokenizer | PreTrainedTokenizerFast,
use_jp_extra: bool = True,
raise_yomi_error: bool = False
) -> tuple[list[str], list[int], list[int]]:
@@ -21,11 +21,9 @@ def g2p(
- word2ph: 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
のタプルを返す。
ただし `phones` と `tones` の最初と終わりに `_` が入り、応じて `word2ph` の最初と最後に 1 が追加される。
tokenizer には deberta-v2-large-japanese-char-wwm を AutoTokenizer.from_pretrained() でロードしたものを指定する。
Args:
norm_text (str): 正規化されたテキスト
tokenizer (PreTrainedTokenizer | PreTrainedTokenizerFast): 単語分割に使うロード済みの BERT Tokenizer インスタンス
use_jp_extra (bool, optional): False の場合、「ん」の音素を「N」ではなく「n」とする。Defaults to True.
raise_yomi_error (bool, optional): False の場合、読めない文字が消えたような扱いとして処理される。Defaults to False.
@@ -66,7 +64,7 @@ def g2p(
for i in sep_text:
if i not in PUNCTUATIONS:
sep_tokenized.append(
tokenizer.tokenize(i)
bert_models.load_tokenizer(Languages.JP).tokenize(i)
) # ここでおそらく`i`が文字単位に分割される
else:
sep_tokenized.append([i])

View File

@@ -1,5 +1,3 @@
from transformers import PreTrainedTokenizer, PreTrainedTokenizerFast
from style_bert_vits2.text_processing.japanese.g2p import g2p
from style_bert_vits2.text_processing.japanese.mora_list import (
MORA_KATA_TO_MORA_PHONEMES,
@@ -8,21 +6,19 @@ from style_bert_vits2.text_processing.japanese.mora_list import (
from style_bert_vits2.text_processing.symbols import PUNCTUATIONS
def g2kata_tone(norm_text: str, tokenizer: PreTrainedTokenizer | PreTrainedTokenizerFast) -> list[tuple[str, int]]:
def g2kata_tone(norm_text: str) -> list[tuple[str, int]]:
"""
テキストからカタカナとアクセントのペアのリストを返す。
推論時のみに使われるので、常に`raise_yomi_error=False`g2pを呼ぶ。
tokenizer には deberta-v2-large-japanese-char-wwm を AutoTokenizer.from_pretrained() でロードしたものを指定する。
推論時のみに使われるので、常に `raise_yomi_error=False`g2p を呼ぶ。
Args:
norm_text: 正規化されたテキスト。
tokenizer (PreTrainedTokenizer | PreTrainedTokenizerFast): 単語分割に使うロード済みの BERT Tokenizer インスタンス
Returns:
カタカナと音高のリスト。
"""
phones, tones, _ = g2p(norm_text, tokenizer, use_jp_extra=True, raise_yomi_error=False)
phones, tones, _ = g2p(norm_text, use_jp_extra=True, raise_yomi_error=False)
return phone_tone2kata_tone(list(zip(phones, tones)))