Refactor: moved text/cleaner.py to style_bert_vits2/text_processing/

This commit is contained in:
tsukumi
2024-03-07 00:24:28 +00:00
parent 89825e68d8
commit e826faf62e
7 changed files with 376 additions and 351 deletions

View File

@@ -7,10 +7,10 @@ from typing import Optional
import click import click
from tqdm import tqdm from tqdm import tqdm
from style_bert_vits2.logging import logger
from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT
from config import config from config import config
from text.cleaner import clean_text from style_bert_vits2.logging import logger
from style_bert_vits2.text_processing.cleaner import clean_text
from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT
preprocess_text_config = config.preprocess_text_config preprocess_text_config = config.preprocess_text_config
@@ -72,7 +72,7 @@ def preprocess(
utt, spk, language, text = line.strip().split("|") utt, spk, language, text = line.strip().split("|")
norm_text, phones, tones, word2ph = clean_text( norm_text, phones, tones, word2ph = clean_text(
text=text, text=text,
language=language, language=language, # type: ignore
use_jp_extra=use_jp_extra, use_jp_extra=use_jp_extra,
raise_yomi_error=(yomi_error != "use"), raise_yomi_error=(yomi_error != "use"),
) )

View File

@@ -1,12 +1,14 @@
from typing import Literal
import torch import torch
import utils import utils
from text import cleaned_text_to_sequence, get_bert from text import cleaned_text_to_sequence, get_bert
from text.cleaner import clean_text
from style_bert_vits2.logging import logger from style_bert_vits2.logging import logger
from style_bert_vits2.models import commons from style_bert_vits2.models import commons
from style_bert_vits2.models.models import SynthesizerTrn from style_bert_vits2.models.models import SynthesizerTrn
from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra
from style_bert_vits2.text_processing.cleaner import clean_text
from style_bert_vits2.text_processing.symbols import SYMBOLS from style_bert_vits2.text_processing.symbols import SYMBOLS
@@ -45,18 +47,21 @@ def get_net_g(model_path: str, version: str, device: str, hps):
def get_text( def get_text(
text, text: str,
language_str, language_str: Literal["JP", "EN", "ZH"],
hps, hps,
device, device: str,
assist_text=None, assist_text: str | None = None,
assist_text_weight=0.7, assist_text_weight: float = 0.7,
given_tone=None, given_tone: list[int] | None = None,
): ):
use_jp_extra = hps.version.endswith("JP-Extra") use_jp_extra = hps.version.endswith("JP-Extra")
# 推論のときにのみ呼び出されるので、raise_yomi_errorFalseに設定 # 推論のみ呼び出されるので、raise_yomi_errorFalse に設定
norm_text, phone, tone, word2ph = clean_text( norm_text, phone, tone, word2ph = clean_text(
text, language_str, use_jp_extra, raise_yomi_error=False text,
language_str,
use_jp_extra = use_jp_extra,
raise_yomi_error = False,
) )
if given_tone is not None: if given_tone is not None:
if len(given_tone) != len(phone): if len(given_tone) != len(phone):
@@ -110,22 +115,22 @@ def get_text(
def infer( def infer(
text, text: str,
style_vec, style_vec,
sdp_ratio, sdp_ratio: float,
noise_scale, noise_scale: float,
noise_scale_w, noise_scale_w: float,
length_scale, length_scale: float,
sid: int, # In the original Bert-VITS2, its speaker_name: str, but here it's id sid: int, # In the original Bert-VITS2, its speaker_name: str, but here it's id
language, language: Literal["JP", "EN", "ZH"],
hps, hps,
net_g, net_g,
device, device: str,
skip_start=False, skip_start: bool = False,
skip_end=False, skip_end: bool = False,
assist_text=None, assist_text: str | None = None,
assist_text_weight=0.7, assist_text_weight: float = 0.7,
given_tone=None, given_tone: list[int] | None = None,
): ):
is_jp_extra = hps.version.endswith("JP-Extra") is_jp_extra = hps.version.endswith("JP-Extra")
bert, ja_bert, en_bert, phones, tones, lang_ids = get_text( bert, ja_bert, en_bert, phones, tones, lang_ids = get_text(
@@ -210,19 +215,19 @@ def infer(
def infer_multilang( def infer_multilang(
text, text: str,
style_vec, style_vec,
sdp_ratio, sdp_ratio: float,
noise_scale, noise_scale: float,
noise_scale_w, noise_scale_w: float,
length_scale, length_scale: float,
sid, sid: int,
language, language: Literal["JP", "EN", "ZH"],
hps, hps,
net_g, net_g,
device, device: str,
skip_start=False, skip_start: bool = False,
skip_end=False, skip_end: bool = False,
): ):
bert, ja_bert, en_bert, phones, tones, lang_ids = [], [], [], [], [], [] bert, ja_bert, en_bert, phones, tones, lang_ids = [], [], [], [], [], []
# emo = get_emo_(reference_audio, emotion, sid) # emo = get_emo_(reference_audio, emotion, sid)
@@ -241,7 +246,7 @@ def infer_multilang(
temp_phones, temp_phones,
temp_tones, temp_tones,
temp_lang_ids, temp_lang_ids,
) = get_text(txt, lang, hps, device) ) = get_text(txt, lang, hps, device) # type: ignore
if _skip_start: if _skip_start:
temp_bert = temp_bert[:, 3:] temp_bert = temp_bert[:, 3:]
temp_ja_bert = temp_ja_bert[:, 3:] temp_ja_bert = temp_ja_bert[:, 3:]

View File

@@ -0,0 +1,46 @@
from typing import Literal
def clean_text(
text: str,
language: Literal["JP", "EN", "ZH"],
use_jp_extra: bool = True,
raise_yomi_error: bool = False,
) -> tuple[str, list[str], list[int], list[int]]:
"""
テキストをクリーニングし、音素に変換する
Args:
text (str): クリーニングするテキスト
language (Literal["JP", "EN", "ZH"]): テキストの言語
use_jp_extra (bool, optional): テキストが日本語の場合に JP-Extra モデルを利用するかどうか。Defaults to True.
raise_yomi_error (bool, optional): False の場合、読めない文字が消えたような扱いとして処理される。Defaults to False.
Returns:
tuple[str, list[str], list[int], list[int]]: クリーニングされたテキストと、音素・アクセント・元のテキストの各文字に音素が何個割り当てられるかのリスト
"""
# Changed to import inside if condition to avoid unnecessary import
if language == "JP":
from transformers import AutoTokenizer
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":
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":
from ...text import chinese as language_module
norm_text = language_module.normalize_text(text)
phones, tones, word2ph = language_module.g2p(norm_text)
else:
raise ValueError(f"Language {language} not supported")
return norm_text, phones, tones, word2ph

View File

@@ -168,7 +168,7 @@ def _g2p(segments):
return phones_list, tones_list, word2ph return phones_list, tones_list, word2ph
def text_normalize(text): def normalize_text(text):
numbers = re.findall(r"\d+(?:\.?\d+)?", text) numbers = re.findall(r"\d+(?:\.?\d+)?", text)
for number in numbers: for number in numbers:
text = text.replace(number, cn2an.an2cn(number), 1) text = text.replace(number, cn2an.an2cn(number), 1)
@@ -186,7 +186,7 @@ if __name__ == "__main__":
from text.chinese_bert import get_bert_feature from text.chinese_bert import get_bert_feature
text = "啊!但是《原神》是由,米哈\游自主, [研发]的一款全.新开放世界.冒险游戏" text = "啊!但是《原神》是由,米哈\游自主, [研发]的一款全.新开放世界.冒险游戏"
text = text_normalize(text) text = normalize_text(text)
print(text) print(text)
phones, tones, word2ph = g2p(text) phones, tones, word2ph = g2p(text)
bert = get_bert_feature(text, word2ph) bert = get_bert_feature(text, word2ph)

View File

@@ -1,26 +0,0 @@
def clean_text(text, language, use_jp_extra=True, raise_yomi_error=False):
# Changed to import inside if condition to avoid unnecessary import
if language == "ZH":
from . import chinese as language_module
norm_text = language_module.text_normalize(text)
phones, tones, word2ph = language_module.g2p(norm_text)
elif language == "EN":
from . import english as language_module
norm_text = language_module.text_normalize(text)
phones, tones, word2ph = language_module.g2p(norm_text)
elif language == "JP":
from . import japanese as language_module
norm_text = language_module.text_normalize(text)
phones, tones, word2ph = language_module.g2p(
norm_text, use_jp_extra, raise_yomi_error=raise_yomi_error
)
else:
raise ValueError(f"Language {language} not supported")
return norm_text, phones, tones, word2ph
if __name__ == "__main__":
pass

View File

@@ -369,7 +369,7 @@ def normalize_numbers(text):
return text return text
def text_normalize(text): def normalize_text(text):
text = normalize_numbers(text) text = normalize_numbers(text)
text = replace_punctuation(text) text = replace_punctuation(text)
text = re.sub(r"([,;.\?\!])([\w])", r"\1 \2", text) text = re.sub(r"([,;.\?\!])([\w])", r"\1 \2", text)

View File

@@ -96,7 +96,7 @@ rep_map = {
} }
def text_normalize(text): def normalize_text(text):
""" """
日本語のテキストを正規化する。 日本語のテキストを正規化する。
結果は、ちょうど次の文字のみからなる: 結果は、ちょうど次の文字のみからなる:
@@ -177,7 +177,7 @@ def g2p(
norm_text: str, use_jp_extra: bool = True, raise_yomi_error: bool = False norm_text: str, use_jp_extra: bool = True, raise_yomi_error: bool = False
) -> tuple[list[str], list[int], list[int]]: ) -> tuple[list[str], list[int], list[int]]:
""" """
他で使われるメインの関数。`text_normalize()`で正規化された`norm_text`を受け取り、 他で使われるメインの関数。`normalize_text()`で正規化された`norm_text`を受け取り、
- phones: 音素のリスト(ただし`!`や`,`や`.`等punctuationが含まれうる - phones: 音素のリスト(ただし`!`や`,`や`.`等punctuationが含まれうる
- tones: アクセントのリスト、0と1からなり、phonesと同じ長さ - tones: アクセントのリスト、0と1からなり、phonesと同じ長さ
- word2ph: 元のテキストの各文字に音素が何個割り当てられるかを表すリスト - word2ph: 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
@@ -350,7 +350,7 @@ def text2sep_kata(
norm_text: str, raise_yomi_error: bool = False norm_text: str, raise_yomi_error: bool = False
) -> tuple[list[str], list[str]]: ) -> tuple[list[str], list[str]]:
""" """
`text_normalize`で正規化済みの`norm_text`を受け取り、それを単語分割し、 `normalize_text()`で正規化済みの`norm_text`を受け取り、それを単語分割し、
分割された単語リストとその読みカタカナor記号1文字のリストのタプルを返す。 分割された単語リストとその読みカタカナor記号1文字のリストのタプルを返す。
単語分割結果は、`g2p()`の`word2ph`で1文字あたりに割り振る音素記号の数を決めるために使う。 単語分割結果は、`g2p()`の`word2ph`で1文字あたりに割り振る音素記号の数を決めるために使う。
例: 例:
@@ -634,7 +634,7 @@ if __name__ == "__main__":
text = "こんにちは、世界。" text = "こんにちは、世界。"
from text.japanese_bert import get_bert_feature from text.japanese_bert import get_bert_feature
text = text_normalize(text) text = normalize_text(text)
phones, tones, word2ph = g2p(text) phones, tones, word2ph = g2p(text)
bert = get_bert_feature(text, word2ph) bert = get_bert_feature(text, word2ph)