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:
@@ -1,19 +1,31 @@
|
||||
from enum import Enum
|
||||
from enum import StrEnum
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
# Style-Bert-VITS2 のバージョン
|
||||
VERSION = "2.3.1"
|
||||
|
||||
# Gradio のテーマ
|
||||
## Built-in theme: "default", "base", "monochrome", "soft", "glass"
|
||||
## See https://huggingface.co/spaces/gradio/theme-gallery for more themes
|
||||
GRADIO_THEME = "NoCrypt/miku"
|
||||
# Style-Bert-VITS2 のベースディレクトリ
|
||||
BASE_DIR = Path(__file__).parent.parent
|
||||
|
||||
# 利用可能な言語
|
||||
## JP-Extra モデル利用時は JP 以外の言語の音声合成はできない
|
||||
class Languages(StrEnum):
|
||||
JP = "JP"
|
||||
EN = "EN"
|
||||
ZH = "ZH"
|
||||
|
||||
# 言語ごとのデフォルトの BERT トークナイザーのパス
|
||||
DEFAULT_BERT_TOKENIZER_PATHS = {
|
||||
Languages.JP: BASE_DIR / "bert" / "deberta-v2-large-japanese-char-wwm",
|
||||
Languages.EN: BASE_DIR / "bert" / "deberta-v3-large",
|
||||
Languages.ZH: BASE_DIR / "bert" / "chinese-roberta-wwm-ext-large",
|
||||
}
|
||||
|
||||
# デフォルトのユーザー辞書ディレクトリ
|
||||
## style_bert_vits2.text_processing.japanese.user_dict モジュールのデフォルト値として利用される
|
||||
## ライブラリとしての利用などで外部のユーザー辞書を指定したい場合は、user_dict 以下の各関数の実行時、引数に辞書データファイルのパスを指定する
|
||||
DEFAULT_USER_DICT_DIR = Path(__file__).parent.parent / "dict_data"
|
||||
DEFAULT_USER_DICT_DIR = BASE_DIR / "dict_data"
|
||||
|
||||
# デフォルトの推論パラメータ
|
||||
DEFAULT_STYLE = "Neutral"
|
||||
@@ -27,9 +39,7 @@ DEFAULT_SPLIT_INTERVAL = 0.5
|
||||
DEFAULT_ASSIST_TEXT_WEIGHT = 0.7
|
||||
DEFAULT_ASSIST_TEXT_WEIGHT = 1.0
|
||||
|
||||
# 利用可能な言語
|
||||
## JP-Extra モデル利用時は JP 以外の言語の音声合成はできない
|
||||
class Languages(str, Enum):
|
||||
JP = "JP"
|
||||
EN = "EN"
|
||||
ZH = "ZH"
|
||||
# Gradio のテーマ
|
||||
## Built-in theme: "default", "base", "monochrome", "soft", "glass"
|
||||
## See https://huggingface.co/spaces/gradio/theme-gallery for more themes
|
||||
GRADIO_THEME = "NoCrypt/miku"
|
||||
|
||||
@@ -1,9 +1,8 @@
|
||||
from typing import Literal
|
||||
|
||||
import torch
|
||||
|
||||
import utils
|
||||
from text import cleaned_text_to_sequence, get_bert
|
||||
from style_bert_vits2.constants import Languages
|
||||
from style_bert_vits2.logging import logger
|
||||
from style_bert_vits2.models import commons
|
||||
from style_bert_vits2.models.models import SynthesizerTrn
|
||||
@@ -48,7 +47,7 @@ def get_net_g(model_path: str, version: str, device: str, hps):
|
||||
|
||||
def get_text(
|
||||
text: str,
|
||||
language_str: Literal["JP", "EN", "ZH"],
|
||||
language_str: Languages,
|
||||
hps,
|
||||
device: str,
|
||||
assist_text: str | None = None,
|
||||
@@ -89,15 +88,15 @@ def get_text(
|
||||
del word2ph
|
||||
assert bert_ori.shape[-1] == len(phone), phone
|
||||
|
||||
if language_str == "ZH":
|
||||
if language_str == Languages.ZH:
|
||||
bert = bert_ori
|
||||
ja_bert = torch.zeros(1024, len(phone))
|
||||
en_bert = torch.zeros(1024, len(phone))
|
||||
elif language_str == "JP":
|
||||
elif language_str == Languages.JP:
|
||||
bert = torch.zeros(1024, len(phone))
|
||||
ja_bert = bert_ori
|
||||
en_bert = torch.zeros(1024, len(phone))
|
||||
elif language_str == "EN":
|
||||
elif language_str == Languages.EN:
|
||||
bert = torch.zeros(1024, len(phone))
|
||||
ja_bert = torch.zeros(1024, len(phone))
|
||||
en_bert = bert_ori
|
||||
@@ -122,7 +121,7 @@ def infer(
|
||||
noise_scale_w: float,
|
||||
length_scale: float,
|
||||
sid: int, # In the original Bert-VITS2, its speaker_name: str, but here it's id
|
||||
language: Literal["JP", "EN", "ZH"],
|
||||
language: Languages,
|
||||
hps,
|
||||
net_g,
|
||||
device: str,
|
||||
@@ -222,7 +221,7 @@ def infer_multilang(
|
||||
noise_scale_w: float,
|
||||
length_scale: float,
|
||||
sid: int,
|
||||
language: Literal["JP", "EN", "ZH"],
|
||||
language: Languages,
|
||||
hps,
|
||||
net_g,
|
||||
device: str,
|
||||
|
||||
123
style_bert_vits2/text_processing/bert_models.py
Normal file
123
style_bert_vits2/text_processing/bert_models.py
Normal 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
|
||||
@@ -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)
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -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)))
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user