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,314 +1,319 @@
import torch from typing import Literal
import utils import torch
from text import cleaned_text_to_sequence, get_bert
from text.cleaner import clean_text import utils
from style_bert_vits2.logging import logger from text import cleaned_text_to_sequence, get_bert
from style_bert_vits2.models import commons from style_bert_vits2.logging import logger
from style_bert_vits2.models.models import SynthesizerTrn from style_bert_vits2.models import commons
from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra from style_bert_vits2.models.models import SynthesizerTrn
from style_bert_vits2.text_processing.symbols import SYMBOLS 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
class InvalidToneError(ValueError):
pass
class InvalidToneError(ValueError):
pass
def get_net_g(model_path: str, version: str, device: str, hps):
if version.endswith("JP-Extra"):
logger.info("Using JP-Extra model") def get_net_g(model_path: str, version: str, device: str, hps):
net_g = SynthesizerTrnJPExtra( if version.endswith("JP-Extra"):
len(SYMBOLS), logger.info("Using JP-Extra model")
hps.data.filter_length // 2 + 1, net_g = SynthesizerTrnJPExtra(
hps.train.segment_size // hps.data.hop_length, len(SYMBOLS),
n_speakers=hps.data.n_speakers, hps.data.filter_length // 2 + 1,
**hps.model, hps.train.segment_size // hps.data.hop_length,
).to(device) n_speakers=hps.data.n_speakers,
else: **hps.model,
logger.info("Using normal model") ).to(device)
net_g = SynthesizerTrn( else:
len(SYMBOLS), logger.info("Using normal model")
hps.data.filter_length // 2 + 1, net_g = SynthesizerTrn(
hps.train.segment_size // hps.data.hop_length, len(SYMBOLS),
n_speakers=hps.data.n_speakers, hps.data.filter_length // 2 + 1,
**hps.model, hps.train.segment_size // hps.data.hop_length,
).to(device) n_speakers=hps.data.n_speakers,
net_g.state_dict() **hps.model,
_ = net_g.eval() ).to(device)
if model_path.endswith(".pth") or model_path.endswith(".pt"): net_g.state_dict()
_ = utils.load_checkpoint(model_path, net_g, None, skip_optimizer=True) _ = net_g.eval()
elif model_path.endswith(".safetensors"): if model_path.endswith(".pth") or model_path.endswith(".pt"):
_ = utils.load_safetensors(model_path, net_g, True) _ = utils.load_checkpoint(model_path, net_g, None, skip_optimizer=True)
else: elif model_path.endswith(".safetensors"):
raise ValueError(f"Unknown model format: {model_path}") _ = utils.load_safetensors(model_path, net_g, True)
return net_g else:
raise ValueError(f"Unknown model format: {model_path}")
return net_g
def get_text(
text,
language_str, def get_text(
hps, text: str,
device, language_str: Literal["JP", "EN", "ZH"],
assist_text=None, hps,
assist_text_weight=0.7, device: str,
given_tone=None, assist_text: str | None = None,
): assist_text_weight: float = 0.7,
use_jp_extra = hps.version.endswith("JP-Extra") given_tone: list[int] | None = None,
# 推論のときにのみ呼び出されるので、raise_yomi_errorはFalseに設定 ):
norm_text, phone, tone, word2ph = clean_text( use_jp_extra = hps.version.endswith("JP-Extra")
text, language_str, use_jp_extra, raise_yomi_error=False # 推論時のみ呼び出されるので、raise_yomi_errorFalse に設定
) norm_text, phone, tone, word2ph = clean_text(
if given_tone is not None: text,
if len(given_tone) != len(phone): language_str,
raise InvalidToneError( use_jp_extra = use_jp_extra,
f"Length of given_tone ({len(given_tone)}) != length of phone ({len(phone)})" raise_yomi_error = False,
) )
tone = given_tone if given_tone is not None:
phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str) if len(given_tone) != len(phone):
raise InvalidToneError(
if hps.data.add_blank: f"Length of given_tone ({len(given_tone)}) != length of phone ({len(phone)})"
phone = commons.intersperse(phone, 0) )
tone = commons.intersperse(tone, 0) tone = given_tone
language = commons.intersperse(language, 0) phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str)
for i in range(len(word2ph)):
word2ph[i] = word2ph[i] * 2 if hps.data.add_blank:
word2ph[0] += 1 phone = commons.intersperse(phone, 0)
bert_ori = get_bert( tone = commons.intersperse(tone, 0)
norm_text, language = commons.intersperse(language, 0)
word2ph, for i in range(len(word2ph)):
language_str, word2ph[i] = word2ph[i] * 2
device, word2ph[0] += 1
assist_text, bert_ori = get_bert(
assist_text_weight, norm_text,
) word2ph,
del word2ph language_str,
assert bert_ori.shape[-1] == len(phone), phone device,
assist_text,
if language_str == "ZH": assist_text_weight,
bert = bert_ori )
ja_bert = torch.zeros(1024, len(phone)) del word2ph
en_bert = torch.zeros(1024, len(phone)) assert bert_ori.shape[-1] == len(phone), phone
elif language_str == "JP":
bert = torch.zeros(1024, len(phone)) if language_str == "ZH":
ja_bert = bert_ori bert = bert_ori
en_bert = torch.zeros(1024, len(phone)) ja_bert = torch.zeros(1024, len(phone))
elif language_str == "EN": en_bert = torch.zeros(1024, len(phone))
bert = torch.zeros(1024, len(phone)) elif language_str == "JP":
ja_bert = torch.zeros(1024, len(phone)) bert = torch.zeros(1024, len(phone))
en_bert = bert_ori ja_bert = bert_ori
else: en_bert = torch.zeros(1024, len(phone))
raise ValueError("language_str should be ZH, JP or EN") elif language_str == "EN":
bert = torch.zeros(1024, len(phone))
assert bert.shape[-1] == len( ja_bert = torch.zeros(1024, len(phone))
phone en_bert = bert_ori
), f"Bert seq len {bert.shape[-1]} != {len(phone)}" else:
raise ValueError("language_str should be ZH, JP or EN")
phone = torch.LongTensor(phone)
tone = torch.LongTensor(tone) assert bert.shape[-1] == len(
language = torch.LongTensor(language) phone
return bert, ja_bert, en_bert, phone, tone, language ), f"Bert seq len {bert.shape[-1]} != {len(phone)}"
phone = torch.LongTensor(phone)
def infer( tone = torch.LongTensor(tone)
text, language = torch.LongTensor(language)
style_vec, return bert, ja_bert, en_bert, phone, tone, language
sdp_ratio,
noise_scale,
noise_scale_w, def infer(
length_scale, text: str,
sid: int, # In the original Bert-VITS2, its speaker_name: str, but here it's id style_vec,
language, sdp_ratio: float,
hps, noise_scale: float,
net_g, noise_scale_w: float,
device, length_scale: float,
skip_start=False, sid: int, # In the original Bert-VITS2, its speaker_name: str, but here it's id
skip_end=False, language: Literal["JP", "EN", "ZH"],
assist_text=None, hps,
assist_text_weight=0.7, net_g,
given_tone=None, device: str,
): skip_start: bool = False,
is_jp_extra = hps.version.endswith("JP-Extra") skip_end: bool = False,
bert, ja_bert, en_bert, phones, tones, lang_ids = get_text( assist_text: str | None = None,
text, assist_text_weight: float = 0.7,
language, given_tone: list[int] | None = None,
hps, ):
device, is_jp_extra = hps.version.endswith("JP-Extra")
assist_text=assist_text, bert, ja_bert, en_bert, phones, tones, lang_ids = get_text(
assist_text_weight=assist_text_weight, text,
given_tone=given_tone, language,
) hps,
if skip_start: device,
phones = phones[3:] assist_text=assist_text,
tones = tones[3:] assist_text_weight=assist_text_weight,
lang_ids = lang_ids[3:] given_tone=given_tone,
bert = bert[:, 3:] )
ja_bert = ja_bert[:, 3:] if skip_start:
en_bert = en_bert[:, 3:] phones = phones[3:]
if skip_end: tones = tones[3:]
phones = phones[:-2] lang_ids = lang_ids[3:]
tones = tones[:-2] bert = bert[:, 3:]
lang_ids = lang_ids[:-2] ja_bert = ja_bert[:, 3:]
bert = bert[:, :-2] en_bert = en_bert[:, 3:]
ja_bert = ja_bert[:, :-2] if skip_end:
en_bert = en_bert[:, :-2] phones = phones[:-2]
with torch.no_grad(): tones = tones[:-2]
x_tst = phones.to(device).unsqueeze(0) lang_ids = lang_ids[:-2]
tones = tones.to(device).unsqueeze(0) bert = bert[:, :-2]
lang_ids = lang_ids.to(device).unsqueeze(0) ja_bert = ja_bert[:, :-2]
bert = bert.to(device).unsqueeze(0) en_bert = en_bert[:, :-2]
ja_bert = ja_bert.to(device).unsqueeze(0) with torch.no_grad():
en_bert = en_bert.to(device).unsqueeze(0) x_tst = phones.to(device).unsqueeze(0)
x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device) tones = tones.to(device).unsqueeze(0)
style_vec = torch.from_numpy(style_vec).to(device).unsqueeze(0) lang_ids = lang_ids.to(device).unsqueeze(0)
del phones bert = bert.to(device).unsqueeze(0)
sid_tensor = torch.LongTensor([sid]).to(device) ja_bert = ja_bert.to(device).unsqueeze(0)
if is_jp_extra: en_bert = en_bert.to(device).unsqueeze(0)
output = net_g.infer( x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device)
x_tst, style_vec = torch.from_numpy(style_vec).to(device).unsqueeze(0)
x_tst_lengths, del phones
sid_tensor, sid_tensor = torch.LongTensor([sid]).to(device)
tones, if is_jp_extra:
lang_ids, output = net_g.infer(
ja_bert, x_tst,
style_vec=style_vec, x_tst_lengths,
sdp_ratio=sdp_ratio, sid_tensor,
noise_scale=noise_scale, tones,
noise_scale_w=noise_scale_w, lang_ids,
length_scale=length_scale, ja_bert,
) style_vec=style_vec,
else: sdp_ratio=sdp_ratio,
output = net_g.infer( noise_scale=noise_scale,
x_tst, noise_scale_w=noise_scale_w,
x_tst_lengths, length_scale=length_scale,
sid_tensor, )
tones, else:
lang_ids, output = net_g.infer(
bert, x_tst,
ja_bert, x_tst_lengths,
en_bert, sid_tensor,
style_vec=style_vec, tones,
sdp_ratio=sdp_ratio, lang_ids,
noise_scale=noise_scale, bert,
noise_scale_w=noise_scale_w, ja_bert,
length_scale=length_scale, en_bert,
) style_vec=style_vec,
audio = output[0][0, 0].data.cpu().float().numpy() sdp_ratio=sdp_ratio,
del ( noise_scale=noise_scale,
x_tst, noise_scale_w=noise_scale_w,
tones, length_scale=length_scale,
lang_ids, )
bert, audio = output[0][0, 0].data.cpu().float().numpy()
x_tst_lengths, del (
sid_tensor, x_tst,
ja_bert, tones,
en_bert, lang_ids,
style_vec, bert,
) # , emo x_tst_lengths,
if torch.cuda.is_available(): sid_tensor,
torch.cuda.empty_cache() ja_bert,
return audio en_bert,
style_vec,
) # , emo
def infer_multilang( if torch.cuda.is_available():
text, torch.cuda.empty_cache()
style_vec, return audio
sdp_ratio,
noise_scale,
noise_scale_w, def infer_multilang(
length_scale, text: str,
sid, style_vec,
language, sdp_ratio: float,
hps, noise_scale: float,
net_g, noise_scale_w: float,
device, length_scale: float,
skip_start=False, sid: int,
skip_end=False, language: Literal["JP", "EN", "ZH"],
): hps,
bert, ja_bert, en_bert, phones, tones, lang_ids = [], [], [], [], [], [] net_g,
# emo = get_emo_(reference_audio, emotion, sid) device: str,
# if isinstance(reference_audio, np.ndarray): skip_start: bool = False,
# emo = get_clap_audio_feature(reference_audio, device) skip_end: bool = False,
# else: ):
# emo = get_clap_text_feature(emotion, device) bert, ja_bert, en_bert, phones, tones, lang_ids = [], [], [], [], [], []
# emo = torch.squeeze(emo, dim=1) # emo = get_emo_(reference_audio, emotion, sid)
for idx, (txt, lang) in enumerate(zip(text, language)): # if isinstance(reference_audio, np.ndarray):
_skip_start = (idx != 0) or (skip_start and idx == 0) # emo = get_clap_audio_feature(reference_audio, device)
_skip_end = (idx != len(language) - 1) or skip_end # else:
( # emo = get_clap_text_feature(emotion, device)
temp_bert, # emo = torch.squeeze(emo, dim=1)
temp_ja_bert, for idx, (txt, lang) in enumerate(zip(text, language)):
temp_en_bert, _skip_start = (idx != 0) or (skip_start and idx == 0)
temp_phones, _skip_end = (idx != len(language) - 1) or skip_end
temp_tones, (
temp_lang_ids, temp_bert,
) = get_text(txt, lang, hps, device) temp_ja_bert,
if _skip_start: temp_en_bert,
temp_bert = temp_bert[:, 3:] temp_phones,
temp_ja_bert = temp_ja_bert[:, 3:] temp_tones,
temp_en_bert = temp_en_bert[:, 3:] temp_lang_ids,
temp_phones = temp_phones[3:] ) = get_text(txt, lang, hps, device) # type: ignore
temp_tones = temp_tones[3:] if _skip_start:
temp_lang_ids = temp_lang_ids[3:] temp_bert = temp_bert[:, 3:]
if _skip_end: temp_ja_bert = temp_ja_bert[:, 3:]
temp_bert = temp_bert[:, :-2] temp_en_bert = temp_en_bert[:, 3:]
temp_ja_bert = temp_ja_bert[:, :-2] temp_phones = temp_phones[3:]
temp_en_bert = temp_en_bert[:, :-2] temp_tones = temp_tones[3:]
temp_phones = temp_phones[:-2] temp_lang_ids = temp_lang_ids[3:]
temp_tones = temp_tones[:-2] if _skip_end:
temp_lang_ids = temp_lang_ids[:-2] temp_bert = temp_bert[:, :-2]
bert.append(temp_bert) temp_ja_bert = temp_ja_bert[:, :-2]
ja_bert.append(temp_ja_bert) temp_en_bert = temp_en_bert[:, :-2]
en_bert.append(temp_en_bert) temp_phones = temp_phones[:-2]
phones.append(temp_phones) temp_tones = temp_tones[:-2]
tones.append(temp_tones) temp_lang_ids = temp_lang_ids[:-2]
lang_ids.append(temp_lang_ids) bert.append(temp_bert)
bert = torch.concatenate(bert, dim=1) ja_bert.append(temp_ja_bert)
ja_bert = torch.concatenate(ja_bert, dim=1) en_bert.append(temp_en_bert)
en_bert = torch.concatenate(en_bert, dim=1) phones.append(temp_phones)
phones = torch.concatenate(phones, dim=0) tones.append(temp_tones)
tones = torch.concatenate(tones, dim=0) lang_ids.append(temp_lang_ids)
lang_ids = torch.concatenate(lang_ids, dim=0) bert = torch.concatenate(bert, dim=1)
with torch.no_grad(): ja_bert = torch.concatenate(ja_bert, dim=1)
x_tst = phones.to(device).unsqueeze(0) en_bert = torch.concatenate(en_bert, dim=1)
tones = tones.to(device).unsqueeze(0) phones = torch.concatenate(phones, dim=0)
lang_ids = lang_ids.to(device).unsqueeze(0) tones = torch.concatenate(tones, dim=0)
bert = bert.to(device).unsqueeze(0) lang_ids = torch.concatenate(lang_ids, dim=0)
ja_bert = ja_bert.to(device).unsqueeze(0) with torch.no_grad():
en_bert = en_bert.to(device).unsqueeze(0) x_tst = phones.to(device).unsqueeze(0)
# emo = emo.to(device).unsqueeze(0) tones = tones.to(device).unsqueeze(0)
x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device) lang_ids = lang_ids.to(device).unsqueeze(0)
del phones bert = bert.to(device).unsqueeze(0)
speakers = torch.LongTensor([hps.data.spk2id[sid]]).to(device) ja_bert = ja_bert.to(device).unsqueeze(0)
audio = ( en_bert = en_bert.to(device).unsqueeze(0)
net_g.infer( # emo = emo.to(device).unsqueeze(0)
x_tst, x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device)
x_tst_lengths, del phones
speakers, speakers = torch.LongTensor([hps.data.spk2id[sid]]).to(device)
tones, audio = (
lang_ids, net_g.infer(
bert, x_tst,
ja_bert, x_tst_lengths,
en_bert, speakers,
style_vec=style_vec, tones,
sdp_ratio=sdp_ratio, lang_ids,
noise_scale=noise_scale, bert,
noise_scale_w=noise_scale_w, ja_bert,
length_scale=length_scale, en_bert,
)[0][0, 0] style_vec=style_vec,
.data.cpu() sdp_ratio=sdp_ratio,
.float() noise_scale=noise_scale,
.numpy() noise_scale_w=noise_scale_w,
) length_scale=length_scale,
del ( )[0][0, 0]
x_tst, .data.cpu()
tones, .float()
lang_ids, .numpy()
bert, )
x_tst_lengths, del (
speakers, x_tst,
ja_bert, tones,
en_bert, lang_ids,
) # , emo bert,
if torch.cuda.is_available(): x_tst_lengths,
torch.cuda.empty_cache() speakers,
return audio ja_bert,
en_bert,
) # , emo
if torch.cuda.is_available():
torch.cuda.empty_cache()
return audio

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)