From e4cd4f84239a6121d767f6d082a80ea2c4305d41 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Wed, 6 Mar 2024 07:38:51 +0000 Subject: [PATCH 01/64] Fix: ignore .venv/ --- .gitignore | 1 + 1 file changed, 1 insertion(+) diff --git a/.gitignore b/.gitignore index b556dff..88048d7 100644 --- a/.gitignore +++ b/.gitignore @@ -2,6 +2,7 @@ __pycache__/ venv/ +.venv/ .ipynb_checkpoints/ /*.yml From f26def436972c13d8a01cf57e2ee874970cbca90 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Wed, 6 Mar 2024 12:39:24 +0000 Subject: [PATCH 02/64] Remove: code that is not referenced anywhere --- spec_gen.py | 87 -------------------------------------------- update_status.py | 93 ------------------------------------------------ 2 files changed, 180 deletions(-) delete mode 100644 spec_gen.py delete mode 100644 update_status.py diff --git a/spec_gen.py b/spec_gen.py deleted file mode 100644 index b6715fa..0000000 --- a/spec_gen.py +++ /dev/null @@ -1,87 +0,0 @@ -import torch -from tqdm import tqdm -from multiprocessing import Pool -from mel_processing import spectrogram_torch, mel_spectrogram_torch -from utils import load_wav_to_torch - - -class AudioProcessor: - def __init__( - self, - max_wav_value, - use_mel_spec_posterior, - filter_length, - n_mel_channels, - sampling_rate, - hop_length, - win_length, - mel_fmin, - mel_fmax, - ): - self.max_wav_value = max_wav_value - self.use_mel_spec_posterior = use_mel_spec_posterior - self.filter_length = filter_length - self.n_mel_channels = n_mel_channels - self.sampling_rate = sampling_rate - self.hop_length = hop_length - self.win_length = win_length - self.mel_fmin = mel_fmin - self.mel_fmax = mel_fmax - - def process_audio(self, filename): - audio, sampling_rate = load_wav_to_torch(filename) - audio_norm = audio / self.max_wav_value - audio_norm = audio_norm.unsqueeze(0) - spec_filename = filename.replace(".wav", ".spec.pt") - if self.use_mel_spec_posterior: - spec_filename = spec_filename.replace(".spec.pt", ".mel.pt") - try: - spec = torch.load(spec_filename) - except: - if self.use_mel_spec_posterior: - spec = mel_spectrogram_torch( - audio_norm, - self.filter_length, - self.n_mel_channels, - self.sampling_rate, - self.hop_length, - self.win_length, - self.mel_fmin, - self.mel_fmax, - center=False, - ) - else: - spec = spectrogram_torch( - audio_norm, - self.filter_length, - self.sampling_rate, - self.hop_length, - self.win_length, - center=False, - ) - spec = torch.squeeze(spec, 0) - torch.save(spec, spec_filename) - return spec, audio_norm - - -# 使用示例 -processor = AudioProcessor( - max_wav_value=32768.0, - use_mel_spec_posterior=False, - filter_length=2048, - n_mel_channels=128, - sampling_rate=44100, - hop_length=512, - win_length=2048, - mel_fmin=0.0, - mel_fmax="null", -) - -with open("filelists/train.list", "r") as f: - filepaths = [line.split("|")[0] for line in f] # 取每一行的第一部分作为audiopath - -# 使用多进程处理 -with Pool(processes=32) as pool: # 使用4个进程 - with tqdm(total=len(filepaths)) as pbar: - for i, _ in enumerate(pool.imap_unordered(processor.process_audio, filepaths)): - pbar.update() diff --git a/update_status.py b/update_status.py deleted file mode 100644 index 7d768c6..0000000 --- a/update_status.py +++ /dev/null @@ -1,93 +0,0 @@ -import os -import gradio as gr - -lang_dict = {"EN(英文)": "_en", "ZH(中文)": "_zh", "JP(日语)": "_jp"} - - -def raw_dir_convert_to_path(target_dir: str, lang): - res = target_dir.rstrip("/").rstrip("\\") - if (not target_dir.startswith("raw")) and (not target_dir.startswith("./raw")): - res = os.path.join("./raw", res) - if ( - (not res.endswith("_zh")) - and (not res.endswith("_jp")) - and (not res.endswith("_en")) - ): - res += lang_dict[lang] - return res - - -def update_g_files(): - g_files = [] - cnt = 0 - for root, dirs, files in os.walk(os.path.abspath("./logs")): - for file in files: - if file.startswith("G_") and file.endswith(".pth"): - g_files.append(os.path.join(root, file)) - cnt += 1 - print(g_files) - return f"更新模型列表完成, 共找到{cnt}个模型", gr.Dropdown.update(choices=g_files) - - -def update_c_files(): - c_files = [] - cnt = 0 - for root, dirs, files in os.walk(os.path.abspath("./logs")): - for file in files: - if file.startswith("config.json"): - c_files.append(os.path.join(root, file)) - cnt += 1 - print(c_files) - return f"更新模型列表完成, 共找到{cnt}个配置文件", gr.Dropdown.update( - choices=c_files - ) - - -def update_model_folders(): - subdirs = [] - cnt = 0 - for root, dirs, files in os.walk(os.path.abspath("./logs")): - for dir_name in dirs: - if os.path.basename(dir_name) != "eval": - subdirs.append(os.path.join(root, dir_name)) - cnt += 1 - print(subdirs) - return f"更新模型文件夹列表完成, 共找到{cnt}个文件夹", gr.Dropdown.update( - choices=subdirs - ) - - -def update_wav_lab_pairs(): - wav_count = tot_count = 0 - for root, _, files in os.walk("./raw"): - for file in files: - # print(file) - file_path = os.path.join(root, file) - if file.lower().endswith(".wav"): - lab_file = os.path.splitext(file_path)[0] + ".lab" - if os.path.exists(lab_file): - wav_count += 1 - tot_count += 1 - return f"{wav_count} / {tot_count}" - - -def update_raw_folders(): - subdirs = [] - cnt = 0 - script_path = os.path.dirname(os.path.abspath(__file__)) # 获取当前脚本的绝对路径 - raw_path = os.path.join(script_path, "raw") - print(raw_path) - os.makedirs(raw_path, exist_ok=True) - for root, dirs, files in os.walk(raw_path): - for dir_name in dirs: - relative_path = os.path.relpath( - os.path.join(root, dir_name), script_path - ) # 获取相对路径 - subdirs.append(relative_path) - cnt += 1 - print(subdirs) - return ( - f"更新raw音频文件夹列表完成, 共找到{cnt}个文件夹", - gr.Dropdown.update(choices=subdirs), - gr.Textbox.update(value=update_wav_lab_pairs()), - ) From 918d168ae707c464a0582b186f5ffdfdad8fbf94 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Wed, 6 Mar 2024 20:56:21 +0000 Subject: [PATCH 03/64] Refactor: rewrote Japanese natural language processing code imported from server_editor.py The logic has not been changed, only renaming, splitting and moving modules on a per-function basis. Existing code will be left in place for the time being to avoid breaking the training code, which is not subject to refactoring this time. --- server_editor.py | 40 +- style_bert_vits2/.editorconfig | 15 + style_bert_vits2/__init__.py | 0 style_bert_vits2/constants.py | 32 ++ style_bert_vits2/logging.py | 15 + .../text_processing/japanese/g2p.py | 493 ++++++++++++++++++ .../text_processing/japanese/g2p_utils.py | 94 ++++ .../text_processing/japanese/mora_list.py | 236 +++++++++ .../text_processing/japanese/normalizer.py | 161 ++++++ style_bert_vits2/text_processing/symbols.py | 192 +++++++ style_bert_vits2/utils/stdout_wrapper.py | 47 ++ 11 files changed, 1312 insertions(+), 13 deletions(-) create mode 100644 style_bert_vits2/.editorconfig create mode 100644 style_bert_vits2/__init__.py create mode 100644 style_bert_vits2/constants.py create mode 100644 style_bert_vits2/logging.py create mode 100644 style_bert_vits2/text_processing/japanese/g2p.py create mode 100644 style_bert_vits2/text_processing/japanese/g2p_utils.py create mode 100644 style_bert_vits2/text_processing/japanese/mora_list.py create mode 100644 style_bert_vits2/text_processing/japanese/normalizer.py create mode 100644 style_bert_vits2/text_processing/symbols.py create mode 100644 style_bert_vits2/utils/stdout_wrapper.py diff --git a/server_editor.py b/server_editor.py index afb9321..70f7c9b 100644 --- a/server_editor.py +++ b/server_editor.py @@ -12,14 +12,13 @@ import io import shutil import sys import webbrowser +import yaml import zipfile from datetime import datetime from io import BytesIO from pathlib import Path -import yaml import numpy as np -import pyopenjtalk import requests import torch import uvicorn @@ -29,21 +28,29 @@ from fastapi.responses import JSONResponse, Response from fastapi.staticfiles import StaticFiles from pydantic import BaseModel from scipy.io import wavfile +from transformers import AutoTokenizer -from common.constants import ( +from common.tts_model import ModelHolder +from style_bert_vits2.constants import ( DEFAULT_ASSIST_TEXT_WEIGHT, DEFAULT_NOISE, DEFAULT_NOISEW, DEFAULT_SDP_RATIO, DEFAULT_STYLE, DEFAULT_STYLE_WEIGHT, - LATEST_VERSION, + VERSION, Languages, ) -from common.log import logger -from common.tts_model import ModelHolder -from text.japanese import g2kata_tone, kata_tone2phone_tone, text_normalize -from text.user_dict import apply_word, update_dict, read_dict, rewrite_word, delete_word +from style_bert_vits2.logging import logger +from style_bert_vits2.text_processing.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone +from style_bert_vits2.text_processing.japanese.normalizer import normalize_text +from text.user_dict import ( + apply_word, + update_dict, + read_dict, + rewrite_word, + delete_word, +) # ---フロントエンド部分に関する処理--- @@ -140,6 +147,12 @@ def save_last_download(latest_release): # ---フロントエンド部分に関する処理ここまで--- # 以降はAPIの設定 +# 最初に pyopenjtalk の辞書を更新 +update_dict() + +# 単語分割に使う BERT トークナイザーをロード +tokenizer = AutoTokenizer.from_pretrained("./bert/deberta-v2-large-japanese-char-wwm") + class AudioResponse(Response): media_type = "audio/wav" @@ -197,7 +210,7 @@ router = APIRouter() @router.get("/version") def version() -> str: - return LATEST_VERSION + return VERSION class MoraTone(BaseModel): @@ -213,8 +226,8 @@ class TextRequest(BaseModel): async def read_item(item: TextRequest): try: # 最初に正規化しないと整合性がとれない - text = text_normalize(item.text) - kata_tone_list = g2kata_tone(text) + text = normalize_text(item.text) + kata_tone_list = g2kata_tone(text, tokenizer) except Exception as e: raise HTTPException( status_code=400, @@ -224,8 +237,8 @@ async def read_item(item: TextRequest): @router.post("/normalize") -async def normalize_text(item: TextRequest): - return text_normalize(item.text) +async def normalize(item: TextRequest): + return normalize_text(item.text) @router.get("/models_info") @@ -311,6 +324,7 @@ def multi_synthesis(request: MultiSynthesisRequest): detail=f"行数は{args.line_count}行以下にしてください。", ) audios = [] + sr = None for i, req in enumerate(lines): if args.line_length is not None and len(req.text) > args.line_length: raise HTTPException( diff --git a/style_bert_vits2/.editorconfig b/style_bert_vits2/.editorconfig new file mode 100644 index 0000000..bbc8a70 --- /dev/null +++ b/style_bert_vits2/.editorconfig @@ -0,0 +1,15 @@ +root = true + +[*] +charset = utf-8 +end_of_line = lf +insert_final_newline = true +indent_size = 4 +indent_style = space +trim_trailing_whitespace = true + +[*.md] +trim_trailing_whitespace = false + +[*.yml] +indent_size = 2 diff --git a/style_bert_vits2/__init__.py b/style_bert_vits2/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/style_bert_vits2/constants.py b/style_bert_vits2/constants.py new file mode 100644 index 0000000..a90bc01 --- /dev/null +++ b/style_bert_vits2/constants.py @@ -0,0 +1,32 @@ +from enum import Enum +from pathlib import Path + + +# Style-Bert-VITS2 のバージョン +VERSION = "2.3.1" + +# ユーザー辞書ディレクトリ +USER_DICT_DIR = Path("dict_data") + +# Gradio のテーマ +## Built-in theme: "default", "base", "monochrome", "soft", "glass" +## See https://huggingface.co/spaces/gradio/theme-gallery for more themes +GRADIO_THEME = "NoCrypt/miku" + +# 利用可能な言語 +class Languages(str, Enum): + JP = "JP" + EN = "EN" + ZH = "ZH" + +# 推論パラメータのデフォルト値 +DEFAULT_STYLE = "Neutral" +DEFAULT_STYLE_WEIGHT = 5.0 +DEFAULT_SDP_RATIO = 0.2 +DEFAULT_NOISE = 0.6 +DEFAULT_NOISEW = 0.8 +DEFAULT_LENGTH = 1.0 +DEFAULT_LINE_SPLIT = True +DEFAULT_SPLIT_INTERVAL = 0.5 +DEFAULT_ASSIST_TEXT_WEIGHT = 0.7 +DEFAULT_ASSIST_TEXT_WEIGHT = 1.0 diff --git a/style_bert_vits2/logging.py b/style_bert_vits2/logging.py new file mode 100644 index 0000000..eec887c --- /dev/null +++ b/style_bert_vits2/logging.py @@ -0,0 +1,15 @@ +from loguru import logger + +from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT + + +# Remove all default handlers +logger.remove() + +# Add a new handler +logger.add( + SAFE_STDOUT, + format = "{time:MM-DD HH:mm:ss} |{level:^8}| {file}:{line} | {message}", + backtrace = True, + diagnose = True, +) diff --git a/style_bert_vits2/text_processing/japanese/g2p.py b/style_bert_vits2/text_processing/japanese/g2p.py new file mode 100644 index 0000000..1b8bd69 --- /dev/null +++ b/style_bert_vits2/text_processing/japanese/g2p.py @@ -0,0 +1,493 @@ +import pyopenjtalk +import re +from transformers import PreTrainedTokenizer, PreTrainedTokenizerFast + +from style_bert_vits2.logging import logger +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 + + +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]]: + """ + 他で使われるメインの関数。`normalize_text()` で正規化された `norm_text` を受け取り、 + - phones: 音素のリスト(ただし `!` や `,` や `.` など punctuation が含まれうる) + - tones: アクセントのリスト、0(低)と1(高)からなり、phones と同じ長さ + - 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. + + Returns: + tuple[list[str], list[int], list[int]]: 音素のリスト、アクセントのリスト、word2ph のリスト + """ + + # pyopenjtalk のフルコンテキストラベルを使ってアクセントを取り出すと、punctuation の位置が消えてしまい情報が失われてしまう: + # 「こんにちは、世界。」と「こんにちは!世界。」と「こんにちは!!!???世界……。」は全て同じになる。 + # よって、まず punctuation 無しの音素とアクセントのリストを作り、 + # それとは別に pyopenjtalk.run_frontend() で得られる音素リスト(こちらは punctuation が保持される)を使い、 + # アクセント割当をしなおすことによって punctuation を含めた音素とアクセントのリストを作る。 + + # punctuation がすべて消えた、音素とアクセントのタプルのリスト(「ん」は「N」) + phone_tone_list_wo_punct = __g2phone_tone_wo_punct(norm_text) + + # sep_text: 単語単位の単語のリスト、読めない文字があったら raise_yomi_error なら例外、そうでないなら読めない文字が消えて返ってくる + # sep_kata: 単語単位の単語のカタカナ読みのリスト + sep_text, sep_kata = text_to_sep_kata(norm_text, raise_yomi_error=raise_yomi_error) + + # sep_phonemes: 各単語ごとの音素のリストのリスト + sep_phonemes = __handle_long([__kata_to_phoneme_list(i) for i in sep_kata]) + + # phone_w_punct: sep_phonemes を結合した、punctuation を元のまま保持した音素列 + phone_w_punct: list[str] = [] + for i in sep_phonemes: + phone_w_punct += i + + # punctuation 無しのアクセント情報を使って、punctuation を含めたアクセント情報を作る + phone_tone_list = __align_tones(phone_w_punct, phone_tone_list_wo_punct) + # logger.debug(f"phone_tone_list:\n{phone_tone_list}") + + # word2ph は厳密な解答は不可能なので(「今日」「眼鏡」等の熟字訓が存在)、 + # Bert-VITS2 では、単語単位の分割を使って、単語の文字ごとにだいたい均等に音素を分配する + + # sep_text から、各単語を1文字1文字分割して、文字のリスト(のリスト)を作る + sep_tokenized: list[list[str]] = [] + for i in sep_text: + if i not in PUNCTUATIONS: + sep_tokenized.append( + tokenizer.tokenize(i) + ) # ここでおそらく`i`が文字単位に分割される + else: + sep_tokenized.append([i]) + + # 各単語について、音素の数と文字の数を比較して、均等っぽく分配する + word2ph = [] + for token, phoneme in zip(sep_tokenized, sep_phonemes): + phone_len = len(phoneme) + word_len = len(token) + word2ph += __distribute_phone(phone_len, word_len) + + # 最初と最後に `_` 記号を追加、アクセントは 0(低)、word2ph もそれに合わせて追加 + phone_tone_list = [("_", 0)] + phone_tone_list + [("_", 0)] + word2ph = [1] + word2ph + [1] + + phones = [phone for phone, _ in phone_tone_list] + tones = [tone for _, tone in phone_tone_list] + + assert len(phones) == sum(word2ph), f"{len(phones)} != {sum(word2ph)}" + + # use_jp_extra でない場合は「N」を「n」に変換 + if not use_jp_extra: + phones = [phone if phone != "N" else "n" for phone in phones] + + return phones, tones, word2ph + + +def text_to_sep_kata( + norm_text: str, + raise_yomi_error: bool = False +) -> tuple[list[str], list[str]]: + """ + `normalize_text` で正規化済みの `norm_text` を受け取り、それを単語分割し、 + 分割された単語リストとその読み(カタカナ or 記号1文字)のリストのタプルを返す。 + 単語分割結果は、`g2p()` の `word2ph` で1文字あたりに割り振る音素記号の数を決めるために使う。 + 例: + `私はそう思う!って感じ?` → + ["私", "は", "そう", "思う", "!", "って", "感じ", "?"], ["ワタシ", "ワ", "ソー", "オモウ", "!", "ッテ", "カンジ", "?"] + + Args: + norm_text (str): 正規化されたテキスト + raise_yomi_error (bool, optional): False の場合、読めない文字が消えたような扱いとして処理される。Defaults to False. + + Returns: + tuple[list[str], list[str]]: 分割された単語リストと、その読み(カタカナ or 記号1文字)のリスト + """ + + # parsed: OpenJTalkの解析結果 + parsed = pyopenjtalk.run_frontend(norm_text) + sep_text: list[str] = [] + sep_kata: list[str] = [] + + for parts in parsed: + # word: 実際の単語の文字列 + # yomi: その読み、但し無声化サインの`’`は除去 + word, yomi = replace_punctuation(parts["string"]), parts["pron"].replace( + "’", "" + ) + """ + ここで `yomi` の取りうる値は以下の通りのはず。 + - `word` が通常単語 → 通常の読み(カタカナ) + (カタカナからなり、長音記号も含みうる、`アー` 等) + - `word` が `ー` から始まる → `ーラー` や `ーーー` など + - `word` が句読点や空白等 → `、` + - `word` が punctuation の繰り返し → 全角にしたもの + 基本的に punctuation は1文字ずつ分かれるが、何故かある程度連続すると1つにまとまる。 + 他にも `word` が読めないキリル文字アラビア文字等が来ると `、` になるが、正規化でこの場合は起きないはず。 + また元のコードでは `yomi` が空白の場合の処理があったが、これは起きないはず。 + 処理すべきは `yomi` が `、` の場合のみのはず。 + """ + assert yomi != "", f"Empty yomi: {word}" + if yomi == "、": + # word は正規化されているので、`.`, `,`, `!`, `'`, `-`, `--` のいずれか + if not set(word).issubset(set(PUNCTUATIONS)): # 記号繰り返しか判定 + # ここは pyopenjtalk が読めない文字等のときに起こる + if raise_yomi_error: + raise YomiError(f"Cannot read: {word} in:\n{norm_text}") + logger.warning(f"Ignoring unknown: {word} in:\n{norm_text}") + continue + # yomi は元の記号のままに変更 + yomi = word + elif yomi == "?": + assert word == "?", f"yomi `?` comes from: {word}" + yomi = "?" + sep_text.append(word) + sep_kata.append(yomi) + + return sep_text, sep_kata + + +def __g2phone_tone_wo_punct(text: str) -> list[tuple[str, int]]: + """ + テキストに対して、音素とアクセント(0か1)のペアのリストを返す。 + ただし「!」「.」「?」等の非音素記号 (punctuation) は全て消える(ポーズ記号も残さない)。 + 非音素記号を含める処理は `align_tones()` で行われる。 + また「っ」は「q」に、「ん」は「N」に変換される。 + 例: "こんにちは、世界ー。。元気?!" → + [('k', 0), ('o', 0), ('N', 1), ('n', 1), ('i', 1), ('ch', 1), ('i', 1), ('w', 1), ('a', 1), ('s', 1), ('e', 1), ('k', 0), ('a', 0), ('i', 0), ('i', 0), ('g', 1), ('e', 1), ('N', 0), ('k', 0), ('i', 0)] + + Args: + text (str): テキスト + + Returns: + list[tuple[str, int]]: 音素とアクセントのペアのリスト + """ + + prosodies = pyopenjtalk_g2p_prosody(text, drop_unvoiced_vowels=True) + # logger.debug(f"prosodies: {prosodies}") + result: list[tuple[str, int]] = [] + current_phrase: list[tuple[str, int]] = [] + current_tone = 0 + + for i, letter in enumerate(prosodies): + # 特殊記号の処理 + + # 文頭記号、無視する + if letter == "^": + assert i == 0, "Unexpected ^" + # アクセント句の終わりに来る記号 + elif letter in ("$", "?", "_", "#"): + # 保持しているフレーズを、アクセント数値を 0-1 に修正し結果に追加 + result.extend(__fix_phone_tone(current_phrase)) + # 末尾に来る終了記号、無視(文中の疑問文は `_` になる) + if letter in ("$", "?"): + assert i == len(prosodies) - 1, f"Unexpected {letter}" + # あとは "_"(ポーズ)と "#"(アクセント句の境界)のみ + # これらは残さず、次のアクセント句に備える。 + current_phrase = [] + # 0 を基準点にしてそこから上昇・下降する(負の場合は上の `fix_phone_tone` で直る) + current_tone = 0 + # アクセント上昇記号 + elif letter == "[": + current_tone = current_tone + 1 + # アクセント下降記号 + elif letter == "]": + current_tone = current_tone - 1 + # それ以外は通常の音素 + else: + if letter == "cl": # 「っ」の処理 + letter = "q" + # elif letter == "N": # 「ん」の処理 + # letter = "n" + current_phrase.append((letter, current_tone)) + + return result + + +def pyopenjtalk_g2p_prosody(text: str, drop_unvoiced_vowels: bool = True) -> list[str]: + """ + ESPnet の実装から引用、変更点無し。「ん」は「N」なことに注意。 + ref: https://github.com/espnet/espnet/blob/master/espnet2/text/phoneme_tokenizer.py + ------------------------------------------------------------------------------------------ + + Extract phoneme + prosoody symbol sequence from input full-context labels. + + The algorithm is based on `Prosodic features control by symbols as input of + sequence-to-sequence acoustic modeling for neural TTS`_ with some r9y9's tweaks. + + Args: + text (str): Input text. + drop_unvoiced_vowels (bool): whether to drop unvoiced vowels. + + Returns: + List[str]: List of phoneme + prosody symbols. + + Examples: + >>> from espnet2.text.phoneme_tokenizer import pyopenjtalk_g2p_prosody + >>> pyopenjtalk_g2p_prosody("こんにちは。") + ['^', 'k', 'o', '[', 'N', 'n', 'i', 'ch', 'i', 'w', 'a', '$'] + + .. _`Prosodic features control by symbols as input of sequence-to-sequence acoustic + modeling for neural TTS`: https://doi.org/10.1587/transinf.2020EDP7104 + """ + + def _numeric_feature_by_regex(regex: str, s: str) -> int: + match = re.search(regex, s) + if match is None: + return -50 + return int(match.group(1)) + + labels = pyopenjtalk.make_label(pyopenjtalk.run_frontend(text)) + N = len(labels) + + phones = [] + for n in range(N): + lab_curr = labels[n] + + # current phoneme + p3 = re.search(r"\-(.*?)\+", lab_curr).group(1) # type: ignore + # deal unvoiced vowels as normal vowels + if drop_unvoiced_vowels and p3 in "AEIOU": + p3 = p3.lower() + + # deal with sil at the beginning and the end of text + if p3 == "sil": + assert n == 0 or n == N - 1 + if n == 0: + phones.append("^") + elif n == N - 1: + # check question form or not + e3 = _numeric_feature_by_regex(r"!(\d+)_", lab_curr) + if e3 == 0: + phones.append("$") + elif e3 == 1: + phones.append("?") + continue + elif p3 == "pau": + phones.append("_") + continue + else: + phones.append(p3) + + # accent type and position info (forward or backward) + a1 = _numeric_feature_by_regex(r"/A:([0-9\-]+)\+", lab_curr) + a2 = _numeric_feature_by_regex(r"\+(\d+)\+", lab_curr) + a3 = _numeric_feature_by_regex(r"\+(\d+)/", lab_curr) + + # number of mora in accent phrase + f1 = _numeric_feature_by_regex(r"/F:(\d+)_", lab_curr) + + a2_next = _numeric_feature_by_regex(r"\+(\d+)\+", labels[n + 1]) + # accent phrase border + if a3 == 1 and a2_next == 1 and p3 in "aeiouAEIOUNcl": + phones.append("#") + # pitch falling + elif a1 == 0 and a2_next == a2 + 1 and a2 != f1: + phones.append("]") + # pitch rising + elif a2 == 1 and a2_next == 2: + phones.append("[") + + return phones + + +def __fix_phone_tone(phone_tone_list: list[tuple[str, int]]) -> list[tuple[str, int]]: + """ + `phone_tone_list` の tone(アクセントの値)を 0 か 1 の範囲に修正する。 + 例: [(a, 0), (i, -1), (u, -1)] → [(a, 1), (i, 0), (u, 0)] + + Args: + phone_tone_list (list[tuple[str, int]]): 音素とアクセントのペアのリスト + + Returns: + list[tuple[str, int]]: 修正された音素とアクセントのペアのリスト + """ + + tone_values = set(tone for _, tone in phone_tone_list) + if len(tone_values) == 1: + assert tone_values == {0}, tone_values + return phone_tone_list + elif len(tone_values) == 2: + if tone_values == {0, 1}: + return phone_tone_list + elif tone_values == {-1, 0}: + return [ + (letter, 0 if tone == -1 else 1) for letter, tone in phone_tone_list + ] + else: + raise ValueError(f"Unexpected tone values: {tone_values}") + else: + raise ValueError(f"Unexpected tone values: {tone_values}") + + +def __handle_long(sep_phonemes: list[list[str]]) -> list[list[str]]: + """ + フレーズごとに分かれた音素(長音記号がそのまま)のリストのリスト `sep_phonemes` を受け取り、 + その長音記号を処理して、音素のリストのリストを返す。 + 基本的には直前の音素を伸ばすが、直前の音素が母音でない場合もしくは冒頭の場合は、 + おそらく長音記号とダッシュを勘違いしていると思われるので、ダッシュに対応する音素 `-` に変換する。 + + Args: + sep_phonemes (list[list[str]]): フレーズごとに分かれた音素のリストのリスト + + Returns: + list[list[str]]: 長音記号を処理した音素のリストのリスト + """ + + # 母音の集合 (便宜上「ん」を含める) + VOWELS = {"a", "i", "u", "e", "o", "N"} + + for i in range(len(sep_phonemes)): + if len(sep_phonemes[i]) == 0: + # 空白文字等でリストが空の場合 + continue + if sep_phonemes[i][0] == "ー": + if i != 0: + prev_phoneme = sep_phonemes[i - 1][-1] + if prev_phoneme in VOWELS: + # 母音と「ん」のあとの伸ばし棒なので、その母音に変換 + sep_phonemes[i][0] = sep_phonemes[i - 1][-1] + else: + # 「。ーー」等おそらく予期しない長音記号 + # ダッシュの勘違いだと思われる + sep_phonemes[i][0] = "-" + else: + # 冒頭に長音記号が来ていおり、これはダッシュの勘違いと思われる + sep_phonemes[i][0] = "-" + if "ー" in sep_phonemes[i]: + for j in range(len(sep_phonemes[i])): + if sep_phonemes[i][j] == "ー": + sep_phonemes[i][j] = sep_phonemes[i][j - 1][-1] + + return sep_phonemes + + +def __kata_to_phoneme_list(text: str) -> list[str]: + """ + 原則カタカナの `text` を受け取り、それをそのままいじらずに音素記号のリストに変換。 + 注意点: + - punctuation かその繰り返しが来た場合、punctuation たちをそのままリストにして返す。 + - 冒頭に続く「ー」はそのまま「ー」のままにする(`handle_long()` で処理される) + - 文中の「ー」は前の音素記号の最後の音素記号に変換される。 + 例: + `ーーソーナノカーー` → ["ー", "ー", "s", "o", "o", "n", "a", "n", "o", "k", "a", "a", "a"] + `?` → ["?"] + `!?!?!?!?!` → ["!", "?", "!", "?", "!", "?", "!", "?", "!"] + + Args: + text (str): カタカナのテキスト + + Returns: + list[str]: 音素記号のリスト + """ + + if set(text).issubset(set(PUNCTUATIONS)): + return list(text) + # `text`がカタカナ(`ー`含む)のみからなるかどうかをチェック + if re.fullmatch(r"[\u30A0-\u30FF]+", text) is None: + raise ValueError(f"Input must be katakana only: {text}") + sorted_keys = sorted(MORA_KATA_TO_MORA_PHONEMES.keys(), key=len, reverse=True) + pattern = "|".join(map(re.escape, sorted_keys)) + + def mora2phonemes(mora: str) -> str: + cosonant, vowel = MORA_KATA_TO_MORA_PHONEMES[mora] + if cosonant is None: + return f" {vowel}" + return f" {cosonant} {vowel}" + + spaced_phonemes = re.sub(pattern, lambda m: mora2phonemes(m.group()), text) + + # 長音記号「ー」の処理 + long_pattern = r"(\w)(ー*)" + long_replacement = lambda m: m.group(1) + (" " + m.group(1)) * len(m.group(2)) # type: ignore + spaced_phonemes = re.sub(long_pattern, long_replacement, spaced_phonemes) + + return spaced_phonemes.strip().split(" ") + + +def __align_tones( + phones_with_punct: list[str], + phone_tone_list: list[tuple[str, int]] +) -> list[tuple[str, int]]: + """ + 例: …私は、、そう思う。 + phones_with_punct: + [".", ".", ".", "w", "a", "t", "a", "sh", "i", "w", "a", ",", ",", "s", "o", "o", "o", "m", "o", "u", "."] + phone_tone_list: + [("w", 0), ("a", 0), ("t", 1), ("a", 1), ("sh", 1), ("i", 1), ("w", 1), ("a", 1), ("_", 0), ("s", 0), ("o", 0), ("o", 1), ("o", 1), ("m", 1), ("o", 1), ("u", 0))] + Return: + [(".", 0), (".", 0), (".", 0), ("w", 0), ("a", 0), ("t", 1), ("a", 1), ("sh", 1), ("i", 1), ("w", 1), ("a", 1), (",", 0), (",", 0), ("s", 0), ("o", 0), ("o", 1), ("o", 1), ("m", 1), ("o", 1), ("u", 0), (".", 0)] + + Args: + phones_with_punct (list[str]): punctuation を含む音素のリスト + phone_tone_list (list[tuple[str, int]]): punctuation を含まない音素とアクセントのペアのリスト + + Returns: + list[tuple[str, int]]: punctuation を含む音素とアクセントのペアのリスト + """ + + result: list[tuple[str, int]] = [] + tone_index = 0 + for phone in phones_with_punct: + if tone_index >= len(phone_tone_list): + # 余ったpunctuationがある場合 → (punctuation, 0)を追加 + result.append((phone, 0)) + elif phone == phone_tone_list[tone_index][0]: + # phone_tone_listの現在の音素と一致する場合 → toneをそこから取得、(phone, tone)を追加 + result.append((phone, phone_tone_list[tone_index][1])) + # 探すindexを1つ進める + tone_index += 1 + elif phone in PUNCTUATIONS: + # phoneがpunctuationの場合 → (phone, 0)を追加 + result.append((phone, 0)) + else: + logger.debug(f"phones: {phones_with_punct}") + logger.debug(f"phone_tone_list: {phone_tone_list}") + logger.debug(f"result: {result}") + logger.debug(f"tone_index: {tone_index}") + logger.debug(f"phone: {phone}") + raise ValueError(f"Unexpected phone: {phone}") + + return result + + +def __distribute_phone(n_phone: int, n_word: int) -> list[int]: + """ + 左から右に 1 ずつ振り分け、次にまた左から右に1ずつ増やし、というふうに、 + 音素の数 `n_phone` を単語の数 `n_word` に分配する。 + + Args: + n_phone (int): 音素の数 + n_word (int): 単語の数 + + Returns: + list[int]: 単語ごとの音素の数のリスト + """ + + phones_per_word = [0] * n_word + for _ in range(n_phone): + min_tasks = min(phones_per_word) + min_index = phones_per_word.index(min_tasks) + phones_per_word[min_index] += 1 + + return phones_per_word + + +class YomiError(Exception): + """ + OpenJTalk で、読みが正しく取得できない箇所があるときに発生する例外。 + 基本的に「学習の前処理のテキスト処理時」には発生させ、そうでない場合は、 + ignore_yomi_error=True にしておいて、この例外を発生させないようにする。 + """ + + pass diff --git a/style_bert_vits2/text_processing/japanese/g2p_utils.py b/style_bert_vits2/text_processing/japanese/g2p_utils.py new file mode 100644 index 0000000..e095602 --- /dev/null +++ b/style_bert_vits2/text_processing/japanese/g2p_utils.py @@ -0,0 +1,94 @@ +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, + MORA_PHONEMES_TO_MORA_KATA, +) +from style_bert_vits2.text_processing.symbols import PUNCTUATIONS + + +def g2kata_tone(norm_text: str, tokenizer: PreTrainedTokenizer | PreTrainedTokenizerFast) -> list[tuple[str, int]]: + """ + テキストからカタカナとアクセントのペアのリストを返す。 + 推論時のみに使われるので、常に`raise_yomi_error=False`でg2pを呼ぶ。 + tokenizer には deberta-v2-large-japanese-char-wwm を AutoTokenizer.from_pretrained() でロードしたものを指定する。 + + Args: + norm_text: 正規化されたテキスト。 + tokenizer (PreTrainedTokenizer | PreTrainedTokenizerFast): 単語分割に使うロード済みの BERT Tokenizer インスタンス + + Returns: + カタカナと音高のリスト。 + """ + + phones, tones, _ = g2p(norm_text, tokenizer, use_jp_extra=True, raise_yomi_error=False) + return phone_tone2kata_tone(list(zip(phones, tones))) + + +def phone_tone2kata_tone(phone_tone: list[tuple[str, int]]) -> list[tuple[str, int]]: + """ + phone_tone の phone 部分をカタカナに変換する。ただし最初と最後の ("_", 0) は無視する。 + + Args: + phone_tone: 音素と音高のリスト。 + + Returns: + カタカナと音高のリスト。 + """ + + # 子音の集合 + CONSONANTS = set([ + consonant + for consonant, _ in MORA_KATA_TO_MORA_PHONEMES.values() + if consonant is not None + ]) + + phone_tone = phone_tone[1:] # 最初の("_", 0)を無視 + phones = [phone for phone, _ in phone_tone] + tones = [tone for _, tone in phone_tone] + result: list[tuple[str, int]] = [] + current_mora = "" + for phone, next_phone, tone, next_tone in zip(phones, phones[1:], tones, tones[1:]): + # zip の関係で最後の ("_", 0) は無視されている + if phone in PUNCTUATIONS: + result.append((phone, tone)) + continue + if phone in CONSONANTS: # n以外の子音の場合 + assert current_mora == "", f"Unexpected {phone} after {current_mora}" + assert tone == next_tone, f"Unexpected {phone} tone {tone} != {next_tone}" + current_mora = phone + else: + # phoneが母音もしくは「N」 + current_mora += phone + result.append((MORA_PHONEMES_TO_MORA_KATA[current_mora], tone)) + current_mora = "" + + return result + + +def kata_tone2phone_tone(kata_tone: list[tuple[str, int]]) -> list[tuple[str, int]]: + """ + `phone_tone2kata_tone()` の逆の変換を行う。 + + Args: + kata_tone: カタカナと音高のリスト。 + + Returns: + 音素と音高のリスト。 + """ + + result: list[tuple[str, int]] = [("_", 0)] + for mora, tone in kata_tone: + if mora in PUNCTUATIONS: + result.append((mora, tone)) + else: + consonant, vowel = MORA_KATA_TO_MORA_PHONEMES[mora] + if consonant is None: + result.append((vowel, tone)) + else: + result.append((consonant, tone)) + result.append((vowel, tone)) + result.append(("_", 0)) + + return result diff --git a/style_bert_vits2/text_processing/japanese/mora_list.py b/style_bert_vits2/text_processing/japanese/mora_list.py new file mode 100644 index 0000000..a0dab2f --- /dev/null +++ b/style_bert_vits2/text_processing/japanese/mora_list.py @@ -0,0 +1,236 @@ +""" +以下のコードは VOICEVOX のソースコードからお借りし最低限の改造を行ったもの。 +https://github.com/VOICEVOX/voicevox_engine/blob/master/voicevox_engine/tts_pipeline/mora_list.py +""" + +""" +以下のモーラ対応表は OpenJTalk のソースコードから取得し、 +カタカナ表記とモーラが一対一対応するように改造した。 +ライセンス表記: +----------------------------------------------------------------- + The Japanese TTS System "Open JTalk" + developed by HTS Working Group + http://open-jtalk.sourceforge.net/ +----------------------------------------------------------------- + + Copyright (c) 2008-2014 Nagoya Institute of Technology + Department of Computer Science + +All rights reserved. + +Redistribution and use in source and binary forms, with or +without modification, are permitted provided that the following +conditions are met: + +- Redistributions of source code must retain the above copyright + notice, this list of conditions and the following disclaimer. +- Redistributions in binary form must reproduce the above + copyright notice, this list of conditions and the following + disclaimer in the documentation and/or other materials provided + with the distribution. +- Neither the name of the HTS working group nor the names of its + contributors may be used to endorse or promote products derived + from this software without specific prior written permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND +CONTRIBUTORS "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, +INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF +MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT OWNER OR CONTRIBUTORS +BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, +EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED +TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, +DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON +ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY +OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE +POSSIBILITY OF SUCH DAMAGE. +""" + +from typing import Optional + + +# (カタカナ, 子音, 母音)の順。子音がない場合は None を入れる。 +# 但し「ン」と「ッ」は母音のみという扱いで、「ン」は「N」、「ッ」は「q」とする。 +# (元々「ッ」は「cl」) +# また「デェ = dy e」は pyopenjtalk の出力(de e)と合わないため削除 +__MORA_LIST_MINIMUM: list[tuple[str, Optional[str], str]] = [ + ("ヴォ", "v", "o"), + ("ヴェ", "v", "e"), + ("ヴィ", "v", "i"), + ("ヴァ", "v", "a"), + ("ヴ", "v", "u"), + ("ン", None, "N"), + ("ワ", "w", "a"), + ("ロ", "r", "o"), + ("レ", "r", "e"), + ("ル", "r", "u"), + ("リョ", "ry", "o"), + ("リュ", "ry", "u"), + ("リャ", "ry", "a"), + ("リェ", "ry", "e"), + ("リ", "r", "i"), + ("ラ", "r", "a"), + ("ヨ", "y", "o"), + ("ユ", "y", "u"), + ("ヤ", "y", "a"), + ("モ", "m", "o"), + ("メ", "m", "e"), + ("ム", "m", "u"), + ("ミョ", "my", "o"), + ("ミュ", "my", "u"), + ("ミャ", "my", "a"), + ("ミェ", "my", "e"), + ("ミ", "m", "i"), + ("マ", "m", "a"), + ("ポ", "p", "o"), + ("ボ", "b", "o"), + ("ホ", "h", "o"), + ("ペ", "p", "e"), + ("ベ", "b", "e"), + ("ヘ", "h", "e"), + ("プ", "p", "u"), + ("ブ", "b", "u"), + ("フォ", "f", "o"), + ("フェ", "f", "e"), + ("フィ", "f", "i"), + ("ファ", "f", "a"), + ("フ", "f", "u"), + ("ピョ", "py", "o"), + ("ピュ", "py", "u"), + ("ピャ", "py", "a"), + ("ピェ", "py", "e"), + ("ピ", "p", "i"), + ("ビョ", "by", "o"), + ("ビュ", "by", "u"), + ("ビャ", "by", "a"), + ("ビェ", "by", "e"), + ("ビ", "b", "i"), + ("ヒョ", "hy", "o"), + ("ヒュ", "hy", "u"), + ("ヒャ", "hy", "a"), + ("ヒェ", "hy", "e"), + ("ヒ", "h", "i"), + ("パ", "p", "a"), + ("バ", "b", "a"), + ("ハ", "h", "a"), + ("ノ", "n", "o"), + ("ネ", "n", "e"), + ("ヌ", "n", "u"), + ("ニョ", "ny", "o"), + ("ニュ", "ny", "u"), + ("ニャ", "ny", "a"), + ("ニェ", "ny", "e"), + ("ニ", "n", "i"), + ("ナ", "n", "a"), + ("ドゥ", "d", "u"), + ("ド", "d", "o"), + ("トゥ", "t", "u"), + ("ト", "t", "o"), + ("デョ", "dy", "o"), + ("デュ", "dy", "u"), + ("デャ", "dy", "a"), + # ("デェ", "dy", "e"), + ("ディ", "d", "i"), + ("デ", "d", "e"), + ("テョ", "ty", "o"), + ("テュ", "ty", "u"), + ("テャ", "ty", "a"), + ("ティ", "t", "i"), + ("テ", "t", "e"), + ("ツォ", "ts", "o"), + ("ツェ", "ts", "e"), + ("ツィ", "ts", "i"), + ("ツァ", "ts", "a"), + ("ツ", "ts", "u"), + ("ッ", None, "q"), # 「cl」から「q」に変更 + ("チョ", "ch", "o"), + ("チュ", "ch", "u"), + ("チャ", "ch", "a"), + ("チェ", "ch", "e"), + ("チ", "ch", "i"), + ("ダ", "d", "a"), + ("タ", "t", "a"), + ("ゾ", "z", "o"), + ("ソ", "s", "o"), + ("ゼ", "z", "e"), + ("セ", "s", "e"), + ("ズィ", "z", "i"), + ("ズ", "z", "u"), + ("スィ", "s", "i"), + ("ス", "s", "u"), + ("ジョ", "j", "o"), + ("ジュ", "j", "u"), + ("ジャ", "j", "a"), + ("ジェ", "j", "e"), + ("ジ", "j", "i"), + ("ショ", "sh", "o"), + ("シュ", "sh", "u"), + ("シャ", "sh", "a"), + ("シェ", "sh", "e"), + ("シ", "sh", "i"), + ("ザ", "z", "a"), + ("サ", "s", "a"), + ("ゴ", "g", "o"), + ("コ", "k", "o"), + ("ゲ", "g", "e"), + ("ケ", "k", "e"), + ("グヮ", "gw", "a"), + ("グ", "g", "u"), + ("クヮ", "kw", "a"), + ("ク", "k", "u"), + ("ギョ", "gy", "o"), + ("ギュ", "gy", "u"), + ("ギャ", "gy", "a"), + ("ギェ", "gy", "e"), + ("ギ", "g", "i"), + ("キョ", "ky", "o"), + ("キュ", "ky", "u"), + ("キャ", "ky", "a"), + ("キェ", "ky", "e"), + ("キ", "k", "i"), + ("ガ", "g", "a"), + ("カ", "k", "a"), + ("オ", None, "o"), + ("エ", None, "e"), + ("ウォ", "w", "o"), + ("ウェ", "w", "e"), + ("ウィ", "w", "i"), + ("ウ", None, "u"), + ("イェ", "y", "e"), + ("イ", None, "i"), + ("ア", None, "a"), +] +__MORA_LIST_ADDITIONAL: list[tuple[str, Optional[str], str]] = [ + ("ヴョ", "by", "o"), + ("ヴュ", "by", "u"), + ("ヴャ", "by", "a"), + ("ヲ", None, "o"), + ("ヱ", None, "e"), + ("ヰ", None, "i"), + ("ヮ", "w", "a"), + ("ョ", "y", "o"), + ("ュ", "y", "u"), + ("ヅ", "z", "u"), + ("ヂ", "j", "i"), + ("ヶ", "k", "e"), + ("ャ", "y", "a"), + ("ォ", None, "o"), + ("ェ", None, "e"), + ("ゥ", None, "u"), + ("ィ", None, "i"), + ("ァ", None, "a"), +] + +# モーラの音素表記とカタカナの対応表 +# 例: "vo" -> "ヴォ", "a" -> "ア" +MORA_PHONEMES_TO_MORA_KATA: dict[str, str] = { + (consonant or "") + vowel: kana for [kana, consonant, vowel] in __MORA_LIST_MINIMUM +} + +# モーラのカタカナ表記と音素の対応表 +# 例: "ヴォ" -> ("v", "o"), "ア" -> (None, "a") +MORA_KATA_TO_MORA_PHONEMES: dict[str, tuple[Optional[str], str]] = { + kana: (consonant, vowel) + for [kana, consonant, vowel] in __MORA_LIST_MINIMUM + __MORA_LIST_ADDITIONAL +} diff --git a/style_bert_vits2/text_processing/japanese/normalizer.py b/style_bert_vits2/text_processing/japanese/normalizer.py new file mode 100644 index 0000000..92c2e87 --- /dev/null +++ b/style_bert_vits2/text_processing/japanese/normalizer.py @@ -0,0 +1,161 @@ +import re +import unicodedata +from num2words import num2words + +from style_bert_vits2.text_processing.symbols import PUNCTUATIONS + + +def normalize_text(text: str) -> str: + """ + 日本語のテキストを正規化する。 + 結果は、ちょうど次の文字のみからなる: + - ひらがな + - カタカナ(全角長音記号「ー」が入る!) + - 漢字 + - 半角アルファベット(大文字と小文字) + - ギリシャ文字 + - `.` (句点`。`や`…`の一部や改行等) + - `,` (読点`、`や`:`等) + - `?` (疑問符`?`) + - `!` (感嘆符`!`) + - `'` (`「`や`」`等) + - `-` (`―`(ダッシュ、長音記号ではない)や`-`等) + + 注意点: + - 三点リーダー`…`は`...`に変換される(`なるほど…。` → `なるほど....`) + - 数字は漢字に変換される(`1,100円` → `千百円`、`52.34` → `五十二点三四`) + - 読点や疑問符等の位置・個数等は保持される(`??あ、、!!!` → `??あ,,!!!`) + + Args: + text (str): 正規化するテキスト + + Returns: + str: 正規化されたテキスト + """ + + res = unicodedata.normalize("NFKC", text) # ここでアルファベットは半角になる + res = __convert_numbers_to_words(res) # 「100円」→「百円」等 + # 「~」と「~」も長音記号として扱う + res = res.replace("~", "ー") + res = res.replace("~", "ー") + + res = replace_punctuation(res) # 句読点等正規化、読めない文字を削除 + + # 結合文字の濁点・半濁点を削除 + # 通常の「ば」等はそのままのこされる、「あ゛」は上で「あ゙」になりここで「あ」になる + res = res.replace("\u3099", "") # 結合文字の濁点を削除、る゙ → る + res = res.replace("\u309A", "") # 結合文字の半濁点を削除、な゚ → な + return res + + +def __convert_numbers_to_words(text: str) -> str: + """ + 記号や数字を日本語の文字表現に変換する。 + + Args: + text (str): 変換するテキスト + + Returns: + str: 変換されたテキスト + """ + + NUMBER_WITH_SEPARATOR_PATTERN = re.compile("[0-9]{1,3}(,[0-9]{3})+") + CURRENCY_MAP = {"$": "ドル", "¥": "円", "£": "ポンド", "€": "ユーロ"} + CURRENCY_PATTERN = re.compile(r"([$¥£€])([0-9.]*[0-9])") + NUMBER_PATTERN = re.compile(r"[0-9]+(\.[0-9]+)?") + + res = NUMBER_WITH_SEPARATOR_PATTERN.sub(lambda m: m[0].replace(",", ""), text) + res = CURRENCY_PATTERN.sub(lambda m: m[2] + CURRENCY_MAP.get(m[1], m[1]), res) + res = NUMBER_PATTERN.sub(lambda m: num2words(m[0], lang="ja"), res) + + return res + + +def replace_punctuation(text: str) -> str: + """ + 句読点等を「.」「,」「!」「?」「'」「-」に正規化し、OpenJTalk で読みが取得できるもののみ残す: + 漢字・平仮名・カタカナ、アルファベット、ギリシャ文字 + + Args: + text (str): 正規化するテキスト + + Returns: + str: 正規化されたテキスト + """ + + # 記号類の正規化変換マップ + REPLACE_MAP = { + ":": ",", + ";": ",", + ",": ",", + "。": ".", + "!": "!", + "?": "?", + "\n": ".", + ".": ".", + "…": "...", + "···": "...", + "・・・": "...", + "·": ",", + "・": ",", + "、": ",", + "$": ".", + "“": "'", + "”": "'", + '"': "'", + "‘": "'", + "’": "'", + "(": "'", + ")": "'", + "(": "'", + ")": "'", + "《": "'", + "》": "'", + "【": "'", + "】": "'", + "[": "'", + "]": "'", + # NFKC 正規化後のハイフン・ダッシュの変種を全て通常半角ハイフン - \u002d に変換 + "\u02d7": "\u002d", # ˗, Modifier Letter Minus Sign + "\u2010": "\u002d", # ‐, Hyphen, + # "\u2011": "\u002d", # ‑, Non-Breaking Hyphen, NFKC により \u2010 に変換される + "\u2012": "\u002d", # ‒, Figure Dash + "\u2013": "\u002d", # –, En Dash + "\u2014": "\u002d", # —, Em Dash + "\u2015": "\u002d", # ―, Horizontal Bar + "\u2043": "\u002d", # ⁃, Hyphen Bullet + "\u2212": "\u002d", # −, Minus Sign + "\u23af": "\u002d", # ⎯, Horizontal Line Extension + "\u23e4": "\u002d", # ⏤, Straightness + "\u2500": "\u002d", # ─, Box Drawings Light Horizontal + "\u2501": "\u002d", # ━, Box Drawings Heavy Horizontal + "\u2e3a": "\u002d", # ⸺, Two-Em Dash + "\u2e3b": "\u002d", # ⸻, Three-Em Dash + # "~": "-", # これは長音記号「ー」として扱うよう変更 + # "~": "-", # これも長音記号「ー」として扱うよう変更 + "「": "'", + "」": "'", + } + + pattern = re.compile("|".join(re.escape(p) for p in REPLACE_MAP.keys())) + + # 句読点を辞書で置換 + replaced_text = pattern.sub(lambda x: REPLACE_MAP[x.group()], text) + + replaced_text = re.sub( + # ↓ ひらがな、カタカナ、漢字 + r"[^\u3040-\u309F\u30A0-\u30FF\u4E00-\u9FFF\u3400-\u4DBF\u3005" + # ↓ 半角アルファベット(大文字と小文字) + + r"\u0041-\u005A\u0061-\u007A" + # ↓ 全角アルファベット(大文字と小文字) + + r"\uFF21-\uFF3A\uFF41-\uFF5A" + # ↓ ギリシャ文字 + + r"\u0370-\u03FF\u1F00-\u1FFF" + # ↓ "!", "?", "…", ",", ".", "'", "-", 但し`…`はすでに`...`に変換されている + + "".join(PUNCTUATIONS) + r"]+", + # 上述以外の文字を削除 + "", + replaced_text, + ) + + return replaced_text diff --git a/style_bert_vits2/text_processing/symbols.py b/style_bert_vits2/text_processing/symbols.py new file mode 100644 index 0000000..628edc6 --- /dev/null +++ b/style_bert_vits2/text_processing/symbols.py @@ -0,0 +1,192 @@ +# Punctuations +PUNCTUATIONS = ["!", "?", "…", ",", ".", "'", "-"] + +# Punctuations and special tokens +PUNCTUATION_SYMBOLS = PUNCTUATIONS + ["SP", "UNK"] + +# Padding +PAD = "_" + +# Chinese symbols +ZH_SYMBOLS = [ + "E", + "En", + "a", + "ai", + "an", + "ang", + "ao", + "b", + "c", + "ch", + "d", + "e", + "ei", + "en", + "eng", + "er", + "f", + "g", + "h", + "i", + "i0", + "ia", + "ian", + "iang", + "iao", + "ie", + "in", + "ing", + "iong", + "ir", + "iu", + "j", + "k", + "l", + "m", + "n", + "o", + "ong", + "ou", + "p", + "q", + "r", + "s", + "sh", + "t", + "u", + "ua", + "uai", + "uan", + "uang", + "ui", + "un", + "uo", + "v", + "van", + "ve", + "vn", + "w", + "x", + "y", + "z", + "zh", + "AA", + "EE", + "OO", +] +NUM_ZH_TONES = 6 + +# japanese +JA_SYMBOLS = [ + "N", + "a", + "a:", + "b", + "by", + "ch", + "d", + "dy", + "e", + "e:", + "f", + "g", + "gy", + "h", + "hy", + "i", + "i:", + "j", + "k", + "ky", + "m", + "my", + "n", + "ny", + "o", + "o:", + "p", + "py", + "q", + "r", + "ry", + "s", + "sh", + "t", + "ts", + "ty", + "u", + "u:", + "w", + "y", + "z", + "zy", +] +NUM_JA_TONES = 2 + +# English +EN_SYMBOLS = [ + "aa", + "ae", + "ah", + "ao", + "aw", + "ay", + "b", + "ch", + "d", + "dh", + "eh", + "er", + "ey", + "f", + "g", + "hh", + "ih", + "iy", + "jh", + "k", + "l", + "m", + "n", + "ng", + "ow", + "oy", + "p", + "r", + "s", + "sh", + "t", + "th", + "uh", + "uw", + "V", + "w", + "y", + "z", + "zh", +] +NUM_EN_TONES = 4 + +# Combine all symbols +NORMAL_SYMBOLS = sorted(set(ZH_SYMBOLS + JA_SYMBOLS + EN_SYMBOLS)) +SYMBOLS = [PAD] + NORMAL_SYMBOLS + PUNCTUATION_SYMBOLS +SIL_PHONEMES_IDS = [SYMBOLS.index(i) for i in PUNCTUATION_SYMBOLS] + +# Combine all tones +num_tones = NUM_ZH_TONES + NUM_JA_TONES + NUM_EN_TONES + +# Language maps +LANGUAGE_ID_MAP = {"ZH": 0, "JP": 1, "EN": 2} +NUM_LANGUAGES = len(LANGUAGE_ID_MAP.keys()) + +LANGUAGE_TONE_START_MAP = { + "ZH": 0, + "JP": NUM_ZH_TONES, + "EN": NUM_ZH_TONES + NUM_JA_TONES, +} + +if __name__ == "__main__": + a = set(ZH_SYMBOLS) + b = set(EN_SYMBOLS) + print(sorted(a & b)) diff --git a/style_bert_vits2/utils/stdout_wrapper.py b/style_bert_vits2/utils/stdout_wrapper.py new file mode 100644 index 0000000..09254ad --- /dev/null +++ b/style_bert_vits2/utils/stdout_wrapper.py @@ -0,0 +1,47 @@ +import sys +import tempfile +from typing import TextIO + + +class StdoutWrapper(TextIO): + """ + `sys.stdout` wrapper for both Google Colab and local environment. + """ + + + def __init__(self) -> None: + self.temp_file = tempfile.NamedTemporaryFile( + mode="w+", delete=False, encoding="utf-8" + ) + self.original_stdout = sys.stdout + + + def write(self, message: str) -> int: + result = self.temp_file.write(message) + self.temp_file.flush() + print(message, end="", file=self.original_stdout) + return result + + + def flush(self) -> None: + self.temp_file.flush() + + + def read(self, n: int = -1) -> str: + self.temp_file.seek(0) + return self.temp_file.read(n) + + + def close(self) -> None: + self.temp_file.close() + + + def fileno(self) -> int: + return self.temp_file.fileno() + + +try: + import google.colab # type: ignore + SAFE_STDOUT = StdoutWrapper() +except ImportError: + SAFE_STDOUT = sys.stdout From 46c83cf89a1b298fff5b956912108ad6aec80023 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Wed, 6 Mar 2024 21:35:47 +0000 Subject: [PATCH 04/64] Refactor: moved the user dictionary implementation ported from VOICEVOX to style_bert_vits2/text_processing/japanese/user_dict/ --- server_editor.py | 2 +- style_bert_vits2/constants.py | 21 ++++++----- .../japanese/user_dict/README.md | 27 ++++++++++++++ .../japanese}/user_dict/__init__.py | 35 +++++++++---------- .../user_dict/part_of_speech_data.py | 13 +++---- .../japanese}/user_dict/word_model.py | 12 ++++--- text/japanese.py | 2 +- text/user_dict/README.md | 19 ---------- 8 files changed, 71 insertions(+), 60 deletions(-) create mode 100644 style_bert_vits2/text_processing/japanese/user_dict/README.md rename {text => style_bert_vits2/text_processing/japanese}/user_dict/__init__.py (93%) rename {text => style_bert_vits2/text_processing/japanese}/user_dict/part_of_speech_data.py (87%) rename {text => style_bert_vits2/text_processing/japanese}/user_dict/word_model.py (93%) delete mode 100644 text/user_dict/README.md diff --git a/server_editor.py b/server_editor.py index 70f7c9b..7577637 100644 --- a/server_editor.py +++ b/server_editor.py @@ -44,7 +44,7 @@ from style_bert_vits2.constants import ( from style_bert_vits2.logging import logger from style_bert_vits2.text_processing.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone from style_bert_vits2.text_processing.japanese.normalizer import normalize_text -from text.user_dict import ( +from style_bert_vits2.text_processing.japanese.user_dict import ( apply_word, update_dict, read_dict, diff --git a/style_bert_vits2/constants.py b/style_bert_vits2/constants.py index a90bc01..2d6e9ad 100644 --- a/style_bert_vits2/constants.py +++ b/style_bert_vits2/constants.py @@ -5,21 +5,17 @@ from pathlib import Path # Style-Bert-VITS2 のバージョン VERSION = "2.3.1" -# ユーザー辞書ディレクトリ -USER_DICT_DIR = Path("dict_data") - # Gradio のテーマ ## Built-in theme: "default", "base", "monochrome", "soft", "glass" ## See https://huggingface.co/spaces/gradio/theme-gallery for more themes GRADIO_THEME = "NoCrypt/miku" -# 利用可能な言語 -class Languages(str, Enum): - JP = "JP" - EN = "EN" - ZH = "ZH" +# デフォルトのユーザー辞書ディレクトリ +## style_bert_vits2.text_processing.japanese.user_dict モジュールのデフォルト値として利用される +## ライブラリとしての利用などで外部のユーザー辞書を指定したい場合は、user_dict 以下の各関数の実行時、引数に辞書データファイルのパスを指定する +DEFAULT_USER_DICT_DIR = Path(__file__).parent.parent / "dict_data" -# 推論パラメータのデフォルト値 +# デフォルトの推論パラメータ DEFAULT_STYLE = "Neutral" DEFAULT_STYLE_WEIGHT = 5.0 DEFAULT_SDP_RATIO = 0.2 @@ -30,3 +26,10 @@ DEFAULT_LINE_SPLIT = True 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" diff --git a/style_bert_vits2/text_processing/japanese/user_dict/README.md b/style_bert_vits2/text_processing/japanese/user_dict/README.md new file mode 100644 index 0000000..3e7a0c0 --- /dev/null +++ b/style_bert_vits2/text_processing/japanese/user_dict/README.md @@ -0,0 +1,27 @@ + +## ユーザー辞書関連のコードについて + +このフォルダに含まれるユーザー辞書関連のコードは、[VOICEVOX ENGINE](https://github.com/VOICEVOX/voicevox_engine) プロジェクトのコードを改変したものを使用しています。 +VOICEVOX プロジェクトのチームに深く感謝し、その貢献を尊重します。 + +### 元のコード + +- [voicevox_engine/user_dict/](https://github.com/VOICEVOX/voicevox_engine/tree/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/user_dict) +- [voicevox_engine/model.py](https://github.com/VOICEVOX/voicevox_engine/blob/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/model.py#L207) + +### 改変の詳細 + +- ファイル名の書き換えおよびそれに伴う import 文の書き換え。 +- VOICEVOX 固有の部分をコメントアウト。 +- mutex を使用している部分をコメントアウト。 +- 参照している pyopenjtalk の違いによるメソッド名の書き換え。 +- UserDictWord の mora_count のデフォルト値を None に指定。 +- `model.py` のうち、必要な Pydantic モデルのみを抽出。 + +### ライセンス + +元の VOICEVOX ENGINE のリポジトリのコードは、LGPL v3 と、ソースコードの公開が不要な別ライセンスのデュアルライセンスの下で使用されています。 +当プロジェクトにおけるこのモジュールも LGPL ライセンスの下にあります。 + +詳細については、プロジェクトのルートディレクトリにある [LGPL_LICENSE](/LGPL_LICENSE) ファイルをご参照ください。 +また、元の VOICEVOX ENGINE プロジェクトのライセンスについては、[こちら](https://github.com/VOICEVOX/voicevox_engine/blob/master/LICENSE) をご覧ください。 diff --git a/text/user_dict/__init__.py b/style_bert_vits2/text_processing/japanese/user_dict/__init__.py similarity index 93% rename from text/user_dict/__init__.py rename to style_bert_vits2/text_processing/japanese/user_dict/__init__.py index c12b3d1..10ea02f 100644 --- a/text/user_dict/__init__.py +++ b/style_bert_vits2/text_processing/japanese/user_dict/__init__.py @@ -1,11 +1,12 @@ -# このファイルは、VOICEVOXプロジェクトのVOICEVOX engineからお借りしています。 -# 引用元: -# https://github.com/VOICEVOX/voicevox_engine/blob/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/user_dict/user_dict.py -# ライセンス: LGPL-3.0 -# 詳しくは、このファイルと同じフォルダにあるREADME.mdを参照してください。 +""" +このファイルは、VOICEVOX プロジェクトの VOICEVOX ENGINE からお借りしています。 +引用元: https://github.com/VOICEVOX/voicevox_engine/blob/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/user_dict/user_dict.py +ライセンス: LGPL-3.0 +詳しくは、このファイルと同じフォルダにある README.md を参照してください。 +""" + import json import sys -import threading import traceback from pathlib import Path from typing import Dict, List, Optional @@ -15,25 +16,21 @@ import numpy as np import pyopenjtalk from fastapi import HTTPException -from .word_model import UserDictWord, WordTypes - +from style_bert_vits2.constants import DEFAULT_USER_DICT_DIR +from style_bert_vits2.text_processing.japanese.user_dict.word_model import UserDictWord, WordTypes # from ..utility.mutex_utility import mutex_wrapper # from ..utility.path_utility import engine_root, get_save_dir -from .part_of_speech_data import MAX_PRIORITY, MIN_PRIORITY, part_of_speech_data -from common.constants import USER_DICT_DIR +from style_bert_vits2.text_processing.japanese.user_dict.part_of_speech_data import MAX_PRIORITY, MIN_PRIORITY, part_of_speech_data # root_dir = engine_root() # save_dir = get_save_dir() -root_dir = Path(USER_DICT_DIR) -save_dir = Path(USER_DICT_DIR) +# if not save_dir.is_dir(): +# save_dir.mkdir(parents=True) -if not save_dir.is_dir(): - save_dir.mkdir(parents=True) - -default_dict_path = root_dir / "default.csv" # VOICEVOXデフォルト辞書ファイルのパス -user_dict_path = save_dir / "user_dict.json" # ユーザー辞書ファイルのパス -compiled_dict_path = save_dir / "user.dic" # コンパイル済み辞書ファイルのパス +default_dict_path = DEFAULT_USER_DICT_DIR / "default.csv" # VOICEVOXデフォルト辞書ファイルのパス +user_dict_path = DEFAULT_USER_DICT_DIR / "user_dict.json" # ユーザー辞書ファイルのパス +compiled_dict_path = DEFAULT_USER_DICT_DIR / "user.dic" # コンパイル済み辞書ファイルのパス # # 同時書き込みの制御 @@ -54,7 +51,7 @@ def _write_to_json(user_dict: Dict[str, UserDictWord], user_dict_path: Path) -> """ converted_user_dict = {} for word_uuid, word in user_dict.items(): - word_dict = word.dict() + word_dict = word.model_dump() word_dict["cost"] = _priority2cost( word_dict["context_id"], word_dict["priority"] ) diff --git a/text/user_dict/part_of_speech_data.py b/style_bert_vits2/text_processing/japanese/user_dict/part_of_speech_data.py similarity index 87% rename from text/user_dict/part_of_speech_data.py rename to style_bert_vits2/text_processing/japanese/user_dict/part_of_speech_data.py index 7e22699..db42d38 100644 --- a/text/user_dict/part_of_speech_data.py +++ b/style_bert_vits2/text_processing/japanese/user_dict/part_of_speech_data.py @@ -1,12 +1,13 @@ -# このファイルは、VOICEVOXプロジェクトのVOICEVOX engineからお借りしています。 -# 引用元: -# https://github.com/VOICEVOX/voicevox_engine/blob/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/user_dict/part_of_speech_data.py -# ライセンス: LGPL-3.0 -# 詳しくは、このファイルと同じフォルダにあるREADME.mdを参照してください。 +""" +このファイルは、VOICEVOX プロジェクトの VOICEVOX ENGINE からお借りしています。 +引用元: https://github.com/VOICEVOX/voicevox_engine/blob/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/user_dict/part_of_speech_data.py +ライセンス: LGPL-3.0 +詳しくは、このファイルと同じフォルダにある README.md を参照してください。 +""" from typing import Dict -from .word_model import ( +from style_bert_vits2.text_processing.japanese.user_dict.word_model import ( USER_DICT_MAX_PRIORITY, USER_DICT_MIN_PRIORITY, PartOfSpeechDetail, diff --git a/text/user_dict/word_model.py b/style_bert_vits2/text_processing/japanese/user_dict/word_model.py similarity index 93% rename from text/user_dict/word_model.py rename to style_bert_vits2/text_processing/japanese/user_dict/word_model.py index f05d8dc..bcd4d37 100644 --- a/text/user_dict/word_model.py +++ b/style_bert_vits2/text_processing/japanese/user_dict/word_model.py @@ -1,8 +1,10 @@ -# このファイルは、VOICEVOXプロジェクトのVOICEVOX engineからお借りしています。 -# 引用元: -# https://github.com/VOICEVOX/voicevox_engine/blob/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/model.py#L207 -# ライセンス: LGPL-3.0 -# 詳しくは、このファイルと同じフォルダにあるREADME.mdを参照してください。 +""" +このファイルは、VOICEVOX プロジェクトの VOICEVOX ENGINE からお借りしています。 +引用元: https://github.com/VOICEVOX/voicevox_engine/blob/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/model.py#L207 +ライセンス: LGPL-3.0 +詳しくは、このファイルと同じフォルダにある README.md を参照してください。 +""" + from enum import Enum from re import findall, fullmatch from typing import List, Optional diff --git a/text/japanese.py b/text/japanese.py index b18bc68..47b21c5 100644 --- a/text/japanese.py +++ b/text/japanese.py @@ -15,7 +15,7 @@ from text.japanese_mora_list import ( mora_phonemes_to_mora_kata, ) -from text.user_dict import update_dict +from style_bert_vits2.text_processing.japanese.user_dict import update_dict # 最初にpyopenjtalkの辞書を更新 update_dict() diff --git a/text/user_dict/README.md b/text/user_dict/README.md deleted file mode 100644 index 6f5618e..0000000 --- a/text/user_dict/README.md +++ /dev/null @@ -1,19 +0,0 @@ -このフォルダに含まれるユーザー辞書関連のコードは、[VOICEVOX engine](https://github.com/VOICEVOX/voicevox_engine)プロジェクトのコードを改変したものを使用しています。VOICEVOXプロジェクトのチームに深く感謝し、その貢献を尊重します。 - -**元のコード**: - -- [voicevox_engine/user_dict/](https://github.com/VOICEVOX/voicevox_engine/tree/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/user_dict) -- [voicevox_engine/model.py](https://github.com/VOICEVOX/voicevox_engine/blob/f181411ec69812296989d9cc583826c22eec87ae/voicevox_engine/model.py#L207) - -**改変の詳細**: - -- ファイル名の書き換えおよびそれに伴うimport文の書き換え。 -- VOICEVOX固有の部分をコメントアウト。 -- mutexを使用している部分をコメントアウト。 -- 参照しているpyopenjtalkの違いによるメソッド名の書き換え。 -- UserDictWordのmora_countのデフォルト値をNoneに指定。 -- Pydanticのモデルで必要な箇所のみを抽出。 - -**ライセンス**: - -元のVOICEVOX engineのリポジトリのコードは、LGPL v3 と、ソースコードの公開が不要な別ライセンスのデュアルライセンスの下で使用されています。当プロジェクトにおけるこのモジュールもLGPLライセンスの下にあります。詳細については、プロジェクトのルートディレクトリにある[LGPL_LICENSE](/LGPL_LICENSE)ファイルをご参照ください。また、元のVOICEVOX engineプロジェクトのライセンスについては、[こちら](https://github.com/VOICEVOX/voicevox_engine/blob/master/LICENSE)をご覧ください。 From 95464954349441f06dcc344be98bcbc845ea1d87 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Wed, 6 Mar 2024 22:17:03 +0000 Subject: [PATCH 05/64] Refactor: moved commons.py to style_bert_vits2/models/ and added type definitions and comments --- attentions.py | 2 +- bert_gen.py | 2 +- commons.py | 152 ------------- data_utils.py | 2 +- infer.py | 2 +- models.py | 4 +- models_jp_extra.py | 4 +- modules.py | 4 +- style_bert_vits2/models/commons.py | 336 +++++++++++++++++++++++++++++ train_ms.py | 2 +- train_ms_jp_extra.py | 2 +- 11 files changed, 348 insertions(+), 164 deletions(-) delete mode 100644 commons.py create mode 100644 style_bert_vits2/models/commons.py diff --git a/attentions.py b/attentions.py index 9a4bba9..87a7f08 100644 --- a/attentions.py +++ b/attentions.py @@ -3,7 +3,7 @@ import torch from torch import nn from torch.nn import functional as F -import commons +from style_bert_vits2.models import commons from common.log import logger as logging diff --git a/bert_gen.py b/bert_gen.py index a5f7c25..70af9fd 100644 --- a/bert_gen.py +++ b/bert_gen.py @@ -5,7 +5,7 @@ import torch import torch.multiprocessing as mp from tqdm import tqdm -import commons +from style_bert_vits2.models import commons import utils from common.log import logger from common.stdout_wrapper import SAFE_STDOUT diff --git a/commons.py b/commons.py deleted file mode 100644 index 081b8a0..0000000 --- a/commons.py +++ /dev/null @@ -1,152 +0,0 @@ -import math -import torch -from torch.nn import functional as F - - -def init_weights(m, mean=0.0, std=0.01): - classname = m.__class__.__name__ - if classname.find("Conv") != -1: - m.weight.data.normal_(mean, std) - - -def get_padding(kernel_size, dilation=1): - return int((kernel_size * dilation - dilation) / 2) - - -def convert_pad_shape(pad_shape): - layer = pad_shape[::-1] - pad_shape = [item for sublist in layer for item in sublist] - return pad_shape - - -def intersperse(lst, item): - result = [item] * (len(lst) * 2 + 1) - result[1::2] = lst - return result - - -def kl_divergence(m_p, logs_p, m_q, logs_q): - """KL(P||Q)""" - kl = (logs_q - logs_p) - 0.5 - kl += ( - 0.5 * (torch.exp(2.0 * logs_p) + ((m_p - m_q) ** 2)) * torch.exp(-2.0 * logs_q) - ) - return kl - - -def rand_gumbel(shape): - """Sample from the Gumbel distribution, protect from overflows.""" - uniform_samples = torch.rand(shape) * 0.99998 + 0.00001 - return -torch.log(-torch.log(uniform_samples)) - - -def rand_gumbel_like(x): - g = rand_gumbel(x.size()).to(dtype=x.dtype, device=x.device) - return g - - -def slice_segments(x, ids_str, segment_size=4): - gather_indices = ids_str.view(x.size(0), 1, 1).repeat( - 1, x.size(1), 1 - ) + torch.arange(segment_size, device=x.device) - return torch.gather(x, 2, gather_indices) - - -def rand_slice_segments(x, x_lengths=None, segment_size=4): - b, d, t = x.size() - if x_lengths is None: - x_lengths = t - ids_str_max = torch.clamp(x_lengths - segment_size + 1, min=0) - ids_str = (torch.rand([b], device=x.device) * ids_str_max).to(dtype=torch.long) - ret = slice_segments(x, ids_str, segment_size) - return ret, ids_str - - -def get_timing_signal_1d(length, channels, min_timescale=1.0, max_timescale=1.0e4): - position = torch.arange(length, dtype=torch.float) - num_timescales = channels // 2 - log_timescale_increment = math.log(float(max_timescale) / float(min_timescale)) / ( - num_timescales - 1 - ) - inv_timescales = min_timescale * torch.exp( - torch.arange(num_timescales, dtype=torch.float) * -log_timescale_increment - ) - scaled_time = position.unsqueeze(0) * inv_timescales.unsqueeze(1) - signal = torch.cat([torch.sin(scaled_time), torch.cos(scaled_time)], 0) - signal = F.pad(signal, [0, 0, 0, channels % 2]) - signal = signal.view(1, channels, length) - return signal - - -def add_timing_signal_1d(x, min_timescale=1.0, max_timescale=1.0e4): - b, channels, length = x.size() - signal = get_timing_signal_1d(length, channels, min_timescale, max_timescale) - return x + signal.to(dtype=x.dtype, device=x.device) - - -def cat_timing_signal_1d(x, min_timescale=1.0, max_timescale=1.0e4, axis=1): - b, channels, length = x.size() - signal = get_timing_signal_1d(length, channels, min_timescale, max_timescale) - return torch.cat([x, signal.to(dtype=x.dtype, device=x.device)], axis) - - -def subsequent_mask(length): - mask = torch.tril(torch.ones(length, length)).unsqueeze(0).unsqueeze(0) - return mask - - -@torch.jit.script -def fused_add_tanh_sigmoid_multiply(input_a, input_b, n_channels): - n_channels_int = n_channels[0] - in_act = input_a + input_b - t_act = torch.tanh(in_act[:, :n_channels_int, :]) - s_act = torch.sigmoid(in_act[:, n_channels_int:, :]) - acts = t_act * s_act - return acts - - -def shift_1d(x): - x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [1, 0]]))[:, :, :-1] - return x - - -def sequence_mask(length, max_length=None): - if max_length is None: - max_length = length.max() - x = torch.arange(max_length, dtype=length.dtype, device=length.device) - return x.unsqueeze(0) < length.unsqueeze(1) - - -def generate_path(duration, mask): - """ - duration: [b, 1, t_x] - mask: [b, 1, t_y, t_x] - """ - - b, _, t_y, t_x = mask.shape - cum_duration = torch.cumsum(duration, -1) - - cum_duration_flat = cum_duration.view(b * t_x) - path = sequence_mask(cum_duration_flat, t_y).to(mask.dtype) - path = path.view(b, t_x, t_y) - path = path - F.pad(path, convert_pad_shape([[0, 0], [1, 0], [0, 0]]))[:, :-1] - path = path.unsqueeze(1).transpose(2, 3) * mask - return path - - -def clip_grad_value_(parameters, clip_value, norm_type=2): - if isinstance(parameters, torch.Tensor): - parameters = [parameters] - parameters = list(filter(lambda p: p.grad is not None, parameters)) - norm_type = float(norm_type) - if clip_value is not None: - clip_value = float(clip_value) - - total_norm = 0 - for p in parameters: - param_norm = p.grad.data.norm(norm_type) - total_norm += param_norm.item() ** norm_type - if clip_value is not None: - p.grad.data.clamp_(min=-clip_value, max=clip_value) - total_norm = total_norm ** (1.0 / norm_type) - return total_norm diff --git a/data_utils.py b/data_utils.py index ac038c2..81c0ade 100644 --- a/data_utils.py +++ b/data_utils.py @@ -7,7 +7,7 @@ import torch import torch.utils.data from tqdm import tqdm -import commons +from style_bert_vits2.models import commons from config import config from mel_processing import mel_spectrogram_torch, spectrogram_torch from text import cleaned_text_to_sequence diff --git a/infer.py b/infer.py index 3707df1..b976486 100644 --- a/infer.py +++ b/infer.py @@ -1,6 +1,6 @@ import torch -import commons +from style_bert_vits2.models import commons import utils from models import SynthesizerTrn from models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra diff --git a/models.py b/models.py index eb706fd..501fcaf 100644 --- a/models.py +++ b/models.py @@ -8,10 +8,10 @@ from torch.nn import functional as F from torch.nn.utils import remove_weight_norm, spectral_norm, weight_norm import attentions -import commons +from style_bert_vits2.models import commons import modules import monotonic_align -from commons import get_padding, init_weights +from style_bert_vits2.models.commons import get_padding, init_weights from text import num_languages, num_tones, symbols diff --git a/models_jp_extra.py b/models_jp_extra.py index 1bb2dd2..3e87ced 100644 --- a/models_jp_extra.py +++ b/models_jp_extra.py @@ -3,7 +3,7 @@ import torch from torch import nn from torch.nn import functional as F -import commons +from style_bert_vits2.models import commons import modules import attentions import monotonic_align @@ -11,7 +11,7 @@ import monotonic_align from torch.nn import Conv1d, ConvTranspose1d, Conv2d from torch.nn.utils import weight_norm, remove_weight_norm, spectral_norm -from commons import init_weights, get_padding +from style_bert_vits2.models.commons import init_weights, get_padding from text import symbols, num_tones, num_languages diff --git a/modules.py b/modules.py index 86b93b5..68b0b9a 100644 --- a/modules.py +++ b/modules.py @@ -7,9 +7,9 @@ from torch.nn import Conv1d from torch.nn import functional as F from torch.nn.utils import remove_weight_norm, weight_norm -import commons +from style_bert_vits2.models import commons from attentions import Encoder -from commons import get_padding, init_weights +from style_bert_vits2.models.commons import get_padding, init_weights from transforms import piecewise_rational_quadratic_transform LRELU_SLOPE = 0.1 diff --git a/style_bert_vits2/models/commons.py b/style_bert_vits2/models/commons.py new file mode 100644 index 0000000..064ef5f --- /dev/null +++ b/style_bert_vits2/models/commons.py @@ -0,0 +1,336 @@ +""" +以下に記述されている関数のコメントはリファクタリング時に GPT-4 に生成させたもので、 +コードと完全に一致している保証はない。あくまで参考程度とすること。 +""" + +import math +import torch +from torch.nn import functional as F +from typing import Any + + +def init_weights(m: torch.nn.Module, mean: float = 0.0, std: float = 0.01) -> None: + """ + モジュールの重みを初期化する + + Args: + m (torch.nn.Module): 重みを初期化する対象のモジュール + mean (float): 正規分布の平均 + std (float): 正規分布の標準偏差 + """ + classname = m.__class__.__name__ + if classname.find("Conv") != -1: + m.weight.data.normal_(mean, std) + + +def get_padding(kernel_size: int, dilation: int = 1) -> int: + """ + カーネルサイズと膨張率からパディングの大きさを計算する + + Args: + kernel_size (int): カーネルのサイズ + dilation (int): 膨張率 + + Returns: + int: 計算されたパディングの大きさ + """ + return int((kernel_size * dilation - dilation) / 2) + + +def convert_pad_shape(pad_shape: list[list[Any]]) -> list[Any]: + """ + パディングの形状を変換する + + Args: + pad_shape (list[list[Any]]): 変換前のパディングの形状 + + Returns: + list[Any]: 変換後のパディングの形状 + """ + layer = pad_shape[::-1] + new_pad_shape = [item for sublist in layer for item in sublist] + return new_pad_shape + + +def intersperse(lst: list[Any], item: Any) -> list[Any]: + """ + リストの要素の間に特定のアイテムを挿入する + + Args: + lst (list[Any]): 元のリスト + item (Any): 挿入するアイテム + + Returns: + list[Any]: 新しいリスト + """ + result = [item] * (len(lst) * 2 + 1) + result[1::2] = lst + return result + + +def kl_divergence(m_p: torch.Tensor, logs_p: torch.Tensor, m_q: torch.Tensor, logs_q: torch.Tensor) -> torch.Tensor: + """ + 2つの正規分布間の KL ダイバージェンスを計算する + + Args: + m_p (torch.Tensor): P の平均 + logs_p (torch.Tensor): P の対数標準偏差 + m_q (torch.Tensor): Q の平均 + logs_q (torch.Tensor): Q の対数標準偏差 + + Returns: + torch.Tensor: KL ダイバージェンスの値。 + """ + kl = (logs_q - logs_p) - 0.5 + kl += ( + 0.5 * (torch.exp(2.0 * logs_p) + ((m_p - m_q) ** 2)) * torch.exp(-2.0 * logs_q) + ) + return kl + + +def rand_gumbel(shape: torch.Size) -> torch.Tensor: + """ + Gumbel 分布からサンプリングし、オーバーフローを防ぐ + + Args: + shape (torch.Size): サンプルの形状 + + Returns: + torch.Tensor: Gumbel 分布からのサンプル + """ + uniform_samples = torch.rand(shape) * 0.99998 + 0.00001 + return -torch.log(-torch.log(uniform_samples)) + + +def rand_gumbel_like(x: torch.Tensor) -> torch.Tensor: + """ + 引数と同じ形状のテンソルで、Gumbel 分布からサンプリングする + + Args: + x (torch.Tensor): 形状を基にするテンソル + + Returns: + torch.Tensor: Gumbel 分布からのサンプル + """ + g = rand_gumbel(x.size()).to(dtype=x.dtype, device=x.device) + return g + + +def slice_segments(x: torch.Tensor, ids_str: torch.Tensor, segment_size: int = 4) -> torch.Tensor: + """ + テンソルからセグメントをスライスする + + Args: + x (torch.Tensor): 入力テンソル + ids_str (torch.Tensor): スライスを開始するインデックス + segment_size (int, optional): スライスのサイズ (デフォルト: 4) + + Returns: + torch.Tensor: スライスされたセグメント + """ + gather_indices = ids_str.view(x.size(0), 1, 1).repeat( + 1, x.size(1), 1 + ) + torch.arange(segment_size, device=x.device) + return torch.gather(x, 2, gather_indices) + + +def rand_slice_segments(x: torch.Tensor, x_lengths: torch.Tensor | None = None, segment_size: int = 4) -> tuple[torch.Tensor, torch.Tensor]: + """ + ランダムなセグメントをスライスする + + Args: + x (torch.Tensor): 入力テンソル + x_lengths (torch.Tensor, optional): 各バッチの長さ (デフォルト: None) + segment_size (int, optional): スライスのサイズ (デフォルト: 4) + + Returns: + tuple[torch.Tensor, torch.Tensor]: スライスされたセグメントと開始インデックス + """ + b, d, t = x.size() + if x_lengths is None: + x_lengths = t # type: ignore + ids_str_max = torch.clamp(x_lengths - segment_size + 1, min=0) # type: ignore + ids_str = (torch.rand([b], device=x.device) * ids_str_max).to(dtype=torch.long) + ret = slice_segments(x, ids_str, segment_size) + return ret, ids_str + + +def get_timing_signal_1d(length: int, channels: int, min_timescale: float = 1.0, max_timescale: float = 1.0e4) -> torch.Tensor: + """ + 1D タイミング信号を取得する + + Args: + length (int): シグナルの長さ + channels (int): シグナルのチャネル数 + min_timescale (float, optional): 最小のタイムスケール (デフォルト: 1.0) + max_timescale (float, optional): 最大のタイムスケール (デフォルト: 1.0e4) + + Returns: + torch.Tensor: タイミング信号 + """ + position = torch.arange(length, dtype=torch.float) + num_timescales = channels // 2 + log_timescale_increment = math.log(float(max_timescale) / float(min_timescale)) / ( + num_timescales - 1 + ) + inv_timescales = min_timescale * torch.exp( + torch.arange(num_timescales, dtype=torch.float) * -log_timescale_increment + ) + scaled_time = position.unsqueeze(0) * inv_timescales.unsqueeze(1) + signal = torch.cat([torch.sin(scaled_time), torch.cos(scaled_time)], 0) + signal = F.pad(signal, [0, 0, 0, channels % 2]) + signal = signal.view(1, channels, length) + return signal + + +def add_timing_signal_1d(x: torch.Tensor, min_timescale: float = 1.0, max_timescale: float = 1.0e4) -> torch.Tensor: + """ + 1D タイミング信号をテンソルに追加する + + Args: + x (torch.Tensor): 入力テンソル + min_timescale (float, optional): 最小のタイムスケール (デフォルト: 1.0) + max_timescale (float, optional): 最大のタイムスケール (デフォルト: 1.0e4) + + Returns: + torch.Tensor: タイミング信号が追加されたテンソル + """ + b, channels, length = x.size() + signal = get_timing_signal_1d(length, channels, min_timescale, max_timescale) + return x + signal.to(dtype=x.dtype, device=x.device) + + +def cat_timing_signal_1d(x: torch.Tensor, min_timescale: float = 1.0, max_timescale: float = 1.0e4, axis: int = 1) -> torch.Tensor: + """ + 1D タイミング信号をテンソルに連結する + + Args: + x (torch.Tensor): 入力テンソル + min_timescale (float, optional): 最小のタイムスケール (デフォルト: 1.0) + max_timescale (float, optional): 最大のタイムスケール (デフォルト: 1.0e4) + axis (int, optional): 連結する軸 (デフォルト: 1) + + Returns: + torch.Tensor: タイミング信号が連結されたテンソル + """ + b, channels, length = x.size() + signal = get_timing_signal_1d(length, channels, min_timescale, max_timescale) + return torch.cat([x, signal.to(dtype=x.dtype, device=x.device)], axis) + + +def subsequent_mask(length: int) -> torch.Tensor: + """ + 後続のマスクを生成する + + Args: + length (int): マスクのサイズ + + Returns: + torch.Tensor: 生成されたマスク + """ + mask = torch.tril(torch.ones(length, length)).unsqueeze(0).unsqueeze(0) + return mask + + +@torch.jit.script # type: ignore +def fused_add_tanh_sigmoid_multiply(input_a: torch.Tensor, input_b: torch.Tensor, n_channels: torch.Tensor) -> torch.Tensor: + """ + 加算、tanh、sigmoid の活性化関数を組み合わせた演算を行う + + Args: + input_a (torch.Tensor): 入力テンソル A + input_b (torch.Tensor): 入力テンソル B + n_channels (torch.Tensor): チャネル数 + + Returns: + torch.Tensor: 演算結果 + """ + n_channels_int = n_channels[0] + in_act = input_a + input_b + t_act = torch.tanh(in_act[:, :n_channels_int, :]) + s_act = torch.sigmoid(in_act[:, n_channels_int:, :]) + acts = t_act * s_act + return acts + + +def shift_1d(x: torch.Tensor) -> torch.Tensor: + """ + 与えられたテンソルを 1D でシフトする + + Args: + x (torch.Tensor): シフトするテンソル + + Returns: + torch.Tensor: シフトされたテンソル + """ + x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [1, 0]]))[:, :, :-1] + return x + + +def sequence_mask(length: torch.Tensor, max_length: int | None = None) -> torch.Tensor: + """ + シーケンスマスクを生成する + + Args: + length (torch.Tensor): 各シーケンスの長さ + max_length (int | None): 最大のシーケンス長さ。指定されていない場合は length の最大値を使用 + + Returns: + torch.Tensor: 生成されたシーケンスマスク + """ + if max_length is None: + max_length = length.max() # type: ignore + x = torch.arange(max_length, dtype=length.dtype, device=length.device) # type: ignore + return x.unsqueeze(0) < length.unsqueeze(1) + + +def generate_path(duration: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: + """ + パスを生成する + + Args: + duration (torch.Tensor): 各時間ステップの持続時間 + mask (torch.Tensor): マスクテンソル + + Returns: + torch.Tensor: 生成されたパス + """ + b, _, t_y, t_x = mask.shape + cum_duration = torch.cumsum(duration, -1) + + cum_duration_flat = cum_duration.view(b * t_x) + path = sequence_mask(cum_duration_flat, t_y).to(mask.dtype) + path = path.view(b, t_x, t_y) + path = path - F.pad(path, convert_pad_shape([[0, 0], [1, 0], [0, 0]]))[:, :-1] + path = path.unsqueeze(1).transpose(2, 3) * mask + return path + + +def clip_grad_value_(parameters: torch.Tensor | list[torch.Tensor], clip_value: float | None, norm_type: float = 2.0) -> float: + """ + 勾配の値をクリップする + + Args: + parameters (torch.Tensor | list[torch.Tensor]): クリップするパラメータ + clip_value (float | None): クリップする値。None の場合はクリップしない + norm_type (float): ノルムの種類 + + Returns: + float: 総ノルム + """ + if isinstance(parameters, torch.Tensor): + parameters = [parameters] + parameters = list(filter(lambda p: p.grad is not None, parameters)) + norm_type = float(norm_type) + if clip_value is not None: + clip_value = float(clip_value) + + total_norm = 0.0 + for p in parameters: + assert p.grad is not None + param_norm = p.grad.data.norm(norm_type) + total_norm += param_norm.item() ** norm_type + if clip_value is not None: + p.grad.data.clamp_(min=-clip_value, max=clip_value) + total_norm = total_norm ** (1.0 / norm_type) + return total_norm diff --git a/train_ms.py b/train_ms.py index 07b5100..0631ffe 100644 --- a/train_ms.py +++ b/train_ms.py @@ -15,7 +15,7 @@ from torch.utils.tensorboard import SummaryWriter from tqdm import tqdm # logging.getLogger("numba").setLevel(logging.WARNING) -import commons +from style_bert_vits2.models import commons import default_style import utils from common.log import logger diff --git a/train_ms_jp_extra.py b/train_ms_jp_extra.py index bae922d..aa5925b 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -15,7 +15,7 @@ from tqdm import tqdm from huggingface_hub import HfApi # logging.getLogger("numba").setLevel(logging.WARNING) -import commons +from style_bert_vits2.models import commons import default_style import utils from common.log import logger From f880641eb5e82508b2a801014e859bf71c52e0c8 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Wed, 6 Mar 2024 22:29:12 +0000 Subject: [PATCH 06/64] Remove: modules under common/ that have been rewritten --- app.py | 8 ++++---- attentions.py | 2 +- bert_gen.py | 4 ++-- common/constants.py | 28 ---------------------------- common/log.py | 17 ----------------- common/stdout_wrapper.py | 40 ---------------------------------------- config.py | 2 +- data_utils.py | 2 +- default_style.py | 4 ++-- infer.py | 2 +- initialize.py | 2 +- losses.py | 2 +- preprocess_text.py | 4 ++-- resample.py | 4 ++-- server_fastapi.py | 4 ++-- slice.py | 4 ++-- speech_mos.py | 2 +- style_gen.py | 4 ++-- text/japanese.py | 2 +- train_ms.py | 4 ++-- train_ms_jp_extra.py | 4 ++-- transcribe.py | 6 +++--- utils.py | 2 +- webui_dataset.py | 4 ++-- webui_merge.py | 4 ++-- webui_style_vectors.py | 4 ++-- webui_train.py | 8 ++++---- 27 files changed, 44 insertions(+), 129 deletions(-) delete mode 100644 common/constants.py delete mode 100644 common/log.py delete mode 100644 common/stdout_wrapper.py diff --git a/app.py b/app.py index 1514444..acdb646 100644 --- a/app.py +++ b/app.py @@ -10,7 +10,7 @@ import gradio as gr import torch import yaml -from common.constants import ( +from style_bert_vits2.constants import ( DEFAULT_ASSIST_TEXT_WEIGHT, DEFAULT_LENGTH, DEFAULT_LINE_SPLIT, @@ -21,10 +21,10 @@ from common.constants import ( DEFAULT_STYLE, DEFAULT_STYLE_WEIGHT, GRADIO_THEME, - LATEST_VERSION, + VERSION, Languages, ) -from common.log import logger +from style_bert_vits2.logging import logger from common.tts_model import ModelHolder from infer import InvalidToneError from text.japanese import g2kata_tone, kata_tone2phone_tone, text_normalize @@ -202,7 +202,7 @@ examples = [ ] initial_md = f""" -# Style-Bert-VITS2 ver {LATEST_VERSION} 音声合成 +# Style-Bert-VITS2 ver {VERSION} 音声合成 - Ver 2.3で追加されたエディターのほうが実際に読み上げさせるには使いやすいかもしれません。`Editor.bat`か`python server_editor.py`で起動できます。 diff --git a/attentions.py b/attentions.py index 87a7f08..1cca086 100644 --- a/attentions.py +++ b/attentions.py @@ -4,7 +4,7 @@ from torch import nn from torch.nn import functional as F from style_bert_vits2.models import commons -from common.log import logger as logging +from style_bert_vits2.logging import logger as logging class LayerNorm(nn.Module): diff --git a/bert_gen.py b/bert_gen.py index 70af9fd..fd0b54e 100644 --- a/bert_gen.py +++ b/bert_gen.py @@ -7,8 +7,8 @@ from tqdm import tqdm from style_bert_vits2.models import commons import utils -from common.log import logger -from common.stdout_wrapper import SAFE_STDOUT +from style_bert_vits2.logging import logger +from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from config import config from text import cleaned_text_to_sequence, get_bert diff --git a/common/constants.py b/common/constants.py deleted file mode 100644 index fe62019..0000000 --- a/common/constants.py +++ /dev/null @@ -1,28 +0,0 @@ -import enum - -# Built-in theme: "default", "base", "monochrome", "soft", "glass" -# See https://huggingface.co/spaces/gradio/theme-gallery for more themes -GRADIO_THEME: str = "NoCrypt/miku" - -LATEST_VERSION: str = "2.3.1" - -USER_DICT_DIR = "dict_data" - -DEFAULT_STYLE: str = "Neutral" -DEFAULT_STYLE_WEIGHT: float = 5.0 - - -class Languages(str, enum.Enum): - JP = "JP" - EN = "EN" - ZH = "ZH" - - -DEFAULT_SDP_RATIO: float = 0.2 -DEFAULT_NOISE: float = 0.6 -DEFAULT_NOISEW: float = 0.8 -DEFAULT_LENGTH: float = 1.0 -DEFAULT_LINE_SPLIT: bool = True -DEFAULT_SPLIT_INTERVAL: float = 0.5 -DEFAULT_ASSIST_TEXT_WEIGHT: float = 0.7 -DEFAULT_ASSIST_TEXT_WEIGHT: float = 1.0 diff --git a/common/log.py b/common/log.py deleted file mode 100644 index 679bb2c..0000000 --- a/common/log.py +++ /dev/null @@ -1,17 +0,0 @@ -""" -logger封装 -""" - -from loguru import logger - -from .stdout_wrapper import SAFE_STDOUT - -# 移除所有默认的处理器 -logger.remove() - -# 自定义格式并添加到标准输出 -log_format = ( - "{time:MM-DD HH:mm:ss} |{level:^8}| {file}:{line} | {message}" -) - -logger.add(SAFE_STDOUT, format=log_format, backtrace=True, diagnose=True) diff --git a/common/stdout_wrapper.py b/common/stdout_wrapper.py deleted file mode 100644 index 192f908..0000000 --- a/common/stdout_wrapper.py +++ /dev/null @@ -1,40 +0,0 @@ -""" -`sys.stdout` wrapper for both Google Colab and local environment. -""" - -import sys -import tempfile - - -class StdoutWrapper: - def __init__(self): - self.temp_file = tempfile.NamedTemporaryFile( - mode="w+", delete=False, encoding="utf-8" - ) - self.original_stdout = sys.stdout - - def write(self, message: str): - self.temp_file.write(message) - self.temp_file.flush() - print(message, end="", file=self.original_stdout) - - def flush(self): - self.temp_file.flush() - - def read(self): - self.temp_file.seek(0) - return self.temp_file.read() - - def close(self): - self.temp_file.close() - - def fileno(self): - return self.temp_file.fileno() - - -try: - import google.colab - - SAFE_STDOUT = StdoutWrapper() -except ImportError: - SAFE_STDOUT = sys.stdout diff --git a/config.py b/config.py index 056e7f1..6369e6b 100644 --- a/config.py +++ b/config.py @@ -9,7 +9,7 @@ from typing import Dict, List import torch import yaml -from common.log import logger +from style_bert_vits2.logging import logger # If not cuda available, set possible devices to cpu cuda_available = torch.cuda.is_available() diff --git a/data_utils.py b/data_utils.py index 81c0ade..1118102 100644 --- a/data_utils.py +++ b/data_utils.py @@ -11,7 +11,7 @@ from style_bert_vits2.models import commons from config import config from mel_processing import mel_spectrogram_torch, spectrogram_torch from text import cleaned_text_to_sequence -from common.log import logger +from style_bert_vits2.logging import logger from utils import load_filepaths_and_text, load_wav_to_torch """Multi speaker version""" diff --git a/default_style.py b/default_style.py index 9198ca8..763e291 100644 --- a/default_style.py +++ b/default_style.py @@ -1,6 +1,6 @@ import os -from common.log import logger -from common.constants import DEFAULT_STYLE +from style_bert_vits2.logging import logger +from style_bert_vits2.constants import DEFAULT_STYLE import numpy as np import json diff --git a/infer.py b/infer.py index b976486..525219d 100644 --- a/infer.py +++ b/infer.py @@ -7,7 +7,7 @@ from models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra from text import cleaned_text_to_sequence, get_bert from text.cleaner import clean_text from text.symbols import symbols -from common.log import logger +from style_bert_vits2.logging import logger class InvalidToneError(ValueError): diff --git a/initialize.py b/initialize.py index 5e35061..927a91a 100644 --- a/initialize.py +++ b/initialize.py @@ -5,7 +5,7 @@ from pathlib import Path import yaml from huggingface_hub import hf_hub_download -from common.log import logger +from style_bert_vits2.logging import logger def download_bert_models(): diff --git a/losses.py b/losses.py index 763cc02..4a890ba 100644 --- a/losses.py +++ b/losses.py @@ -2,7 +2,7 @@ import torch import torchaudio from transformers import AutoModel -from common.log import logger +from style_bert_vits2.logging import logger def feature_loss(fmap_r, fmap_g): diff --git a/preprocess_text.py b/preprocess_text.py index b3aaa17..126ba2c 100644 --- a/preprocess_text.py +++ b/preprocess_text.py @@ -7,8 +7,8 @@ from typing import Optional import click from tqdm import tqdm -from common.log import logger -from common.stdout_wrapper import SAFE_STDOUT +from style_bert_vits2.logging import logger +from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from config import config from text.cleaner import clean_text diff --git a/resample.py b/resample.py index 5aff5ef..7001af6 100644 --- a/resample.py +++ b/resample.py @@ -7,8 +7,8 @@ import pyloudnorm as pyln import soundfile from tqdm import tqdm -from common.log import logger -from common.stdout_wrapper import SAFE_STDOUT +from style_bert_vits2.logging import logger +from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from config import config DEFAULT_BLOCK_SIZE: float = 0.400 # seconds diff --git a/server_fastapi.py b/server_fastapi.py index ce6ed04..132cc17 100644 --- a/server_fastapi.py +++ b/server_fastapi.py @@ -20,7 +20,7 @@ from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import FileResponse, Response from scipy.io import wavfile -from common.constants import ( +from style_bert_vits2.constants import ( DEFAULT_ASSIST_TEXT_WEIGHT, DEFAULT_LENGTH, DEFAULT_LINE_SPLIT, @@ -32,7 +32,7 @@ from common.constants import ( DEFAULT_STYLE_WEIGHT, Languages, ) -from common.log import logger +from style_bert_vits2.logging import logger from common.tts_model import Model, ModelHolder from config import config diff --git a/slice.py b/slice.py index 2d56427..c69f8bf 100644 --- a/slice.py +++ b/slice.py @@ -8,8 +8,8 @@ import torch import yaml from tqdm import tqdm -from common.log import logger -from common.stdout_wrapper import SAFE_STDOUT +from style_bert_vits2.logging import logger +from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT vad_model, utils = torch.hub.load( repo_or_dir="snakers4/silero-vad", diff --git a/speech_mos.py b/speech_mos.py index d69a23a..15cccef 100644 --- a/speech_mos.py +++ b/speech_mos.py @@ -10,7 +10,7 @@ import pandas as pd import torch from tqdm import tqdm -from common.log import logger +from style_bert_vits2.logging import logger from common.tts_model import Model from config import config diff --git a/style_gen.py b/style_gen.py index 97a0aee..1c1f034 100644 --- a/style_gen.py +++ b/style_gen.py @@ -7,8 +7,8 @@ import torch from tqdm import tqdm import utils -from common.log import logger -from common.stdout_wrapper import SAFE_STDOUT +from style_bert_vits2.logging import logger +from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from config import config warnings.filterwarnings("ignore", category=UserWarning) diff --git a/text/japanese.py b/text/japanese.py index 47b21c5..12dc349 100644 --- a/text/japanese.py +++ b/text/japanese.py @@ -8,7 +8,7 @@ import pyopenjtalk from num2words import num2words from transformers import AutoTokenizer -from common.log import logger +from style_bert_vits2.logging import logger from text import punctuation from text.japanese_mora_list import ( mora_kata_to_mora_phonemes, diff --git a/train_ms.py b/train_ms.py index 0631ffe..3e9e707 100644 --- a/train_ms.py +++ b/train_ms.py @@ -18,8 +18,8 @@ from tqdm import tqdm from style_bert_vits2.models import commons import default_style import utils -from common.log import logger -from common.stdout_wrapper import SAFE_STDOUT +from style_bert_vits2.logging import logger +from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from config import config from data_utils import ( DistributedBucketSampler, diff --git a/train_ms_jp_extra.py b/train_ms_jp_extra.py index aa5925b..9a964c3 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -18,8 +18,8 @@ from huggingface_hub import HfApi from style_bert_vits2.models import commons import default_style import utils -from common.log import logger -from common.stdout_wrapper import SAFE_STDOUT +from style_bert_vits2.logging import logger +from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from config import config from data_utils import ( DistributedBucketSampler, diff --git a/transcribe.py b/transcribe.py index b3f9014..18509c9 100644 --- a/transcribe.py +++ b/transcribe.py @@ -7,9 +7,9 @@ import yaml from faster_whisper import WhisperModel from tqdm import tqdm -from common.constants import Languages -from common.log import logger -from common.stdout_wrapper import SAFE_STDOUT +from style_bert_vits2.constants import Languages +from style_bert_vits2.logging import logger +from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT def transcribe(wav_path: Path, initial_prompt=None, language="ja"): diff --git a/utils.py b/utils.py index ca194bf..80dfa66 100644 --- a/utils.py +++ b/utils.py @@ -13,7 +13,7 @@ from safetensors import safe_open from safetensors.torch import save_file from scipy.io.wavfile import read -from common.log import logger +from style_bert_vits2.logging import logger MATPLOTLIB_FLAG = False diff --git a/webui_dataset.py b/webui_dataset.py index fec7a9a..169cc4d 100644 --- a/webui_dataset.py +++ b/webui_dataset.py @@ -4,8 +4,8 @@ import os import gradio as gr import yaml -from common.constants import GRADIO_THEME -from common.log import logger +from style_bert_vits2.constants import GRADIO_THEME +from style_bert_vits2.logging import logger from common.subprocess_utils import run_script_with_log # Get path settings diff --git a/webui_merge.py b/webui_merge.py index a58471a..0a39902 100644 --- a/webui_merge.py +++ b/webui_merge.py @@ -11,8 +11,8 @@ import yaml from safetensors import safe_open from safetensors.torch import save_file -from common.constants import DEFAULT_STYLE, GRADIO_THEME -from common.log import logger +from style_bert_vits2.constants import DEFAULT_STYLE, GRADIO_THEME +from style_bert_vits2.logging import logger from common.tts_model import Model, ModelHolder voice_keys = ["dec"] diff --git a/webui_style_vectors.py b/webui_style_vectors.py index b89c9c9..cf53ca2 100644 --- a/webui_style_vectors.py +++ b/webui_style_vectors.py @@ -12,8 +12,8 @@ from sklearn.cluster import DBSCAN, AgglomerativeClustering, KMeans from sklearn.manifold import TSNE from umap import UMAP -from common.constants import DEFAULT_STYLE, GRADIO_THEME -from common.log import logger +from style_bert_vits2.constants import DEFAULT_STYLE, GRADIO_THEME +from style_bert_vits2.logging import logger from config import config # Get path settings diff --git a/webui_train.py b/webui_train.py index 59cc9f5..0f272e8 100644 --- a/webui_train.py +++ b/webui_train.py @@ -14,9 +14,9 @@ from pathlib import Path import gradio as gr import yaml -from common.constants import GRADIO_THEME, LATEST_VERSION -from common.log import logger -from common.stdout_wrapper import SAFE_STDOUT +from style_bert_vits2.constants import GRADIO_THEME, VERSION +from style_bert_vits2.logging import logger +from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from common.subprocess_utils import run_script_with_log, second_elem_of logger_handler = None @@ -399,7 +399,7 @@ def run_tensorboard(model_name): initial_md = f""" -# Style-Bert-VITS2 ver {LATEST_VERSION} 学習用WebUI +# Style-Bert-VITS2 ver {VERSION} 学習用WebUI ## 使い方 From 1936344c0c6af8213b4327cb964b32dbd20130ba Mon Sep 17 00:00:00 2001 From: tsukumi Date: Wed, 6 Mar 2024 22:51:25 +0000 Subject: [PATCH 07/64] Refactor: remove old code that can be deleted and update where modules are imported --- app.py | 5 +- infer.py | 6 +- models.py | 8 +- models_jp_extra.py | 8 +- style_bert_vits2/text_processing/symbols.py | 2 +- text/__init__.py | 8 +- text/chinese.py | 8 +- text/english.py | 13 +- text/japanese.py | 34 ++- text/japanese_bert.py | 6 +- text/japanese_mora_list.py | 232 -------------------- text/symbols.py | 187 ---------------- train_ms.py | 4 +- train_ms_jp_extra.py | 4 +- 14 files changed, 52 insertions(+), 473 deletions(-) delete mode 100644 text/japanese_mora_list.py delete mode 100644 text/symbols.py diff --git a/app.py b/app.py index acdb646..c70a89d 100644 --- a/app.py +++ b/app.py @@ -27,7 +27,8 @@ from style_bert_vits2.constants import ( from style_bert_vits2.logging import logger from common.tts_model import ModelHolder from infer import InvalidToneError -from text.japanese import g2kata_tone, kata_tone2phone_tone, text_normalize +from style_bert_vits2.text_processing.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone +from style_bert_vits2.text_processing.japanese.normalizer import normalize_text # Get path settings with open(os.path.join("configs", "paths.yml"), "r", encoding="utf-8") as f: @@ -131,7 +132,7 @@ def tts_fn( if tone is None and language == "JP": # アクセント指定に使えるようにアクセント情報を返す - norm_text = text_normalize(text) + norm_text = normalize_text(text) kata_tone = g2kata_tone(norm_text) kata_tone_json_str = json.dumps(kata_tone, ensure_ascii=False) elif tone is None: diff --git a/infer.py b/infer.py index 525219d..4afd048 100644 --- a/infer.py +++ b/infer.py @@ -6,7 +6,7 @@ from models import SynthesizerTrn from models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra from text import cleaned_text_to_sequence, get_bert from text.cleaner import clean_text -from text.symbols import symbols +from style_bert_vits2.text_processing.symbols import SYMBOLS from style_bert_vits2.logging import logger @@ -18,7 +18,7 @@ def get_net_g(model_path: str, version: str, device: str, hps): if version.endswith("JP-Extra"): logger.info("Using JP-Extra model") net_g = SynthesizerTrnJPExtra( - len(symbols), + len(SYMBOLS), hps.data.filter_length // 2 + 1, hps.train.segment_size // hps.data.hop_length, n_speakers=hps.data.n_speakers, @@ -27,7 +27,7 @@ def get_net_g(model_path: str, version: str, device: str, hps): else: logger.info("Using normal model") net_g = SynthesizerTrn( - len(symbols), + len(SYMBOLS), hps.data.filter_length // 2 + 1, hps.train.segment_size // hps.data.hop_length, n_speakers=hps.data.n_speakers, diff --git a/models.py b/models.py index 501fcaf..ef581f2 100644 --- a/models.py +++ b/models.py @@ -12,7 +12,7 @@ from style_bert_vits2.models import commons import modules import monotonic_align from style_bert_vits2.models.commons import get_padding, init_weights -from text import num_languages, num_tones, symbols +from style_bert_vits2.text_processing.symbols import NUM_LANGUAGES, NUM_TONES, SYMBOLS class DurationDiscriminator(nn.Module): # vits2 @@ -334,11 +334,11 @@ class TextEncoder(nn.Module): self.kernel_size = kernel_size self.p_dropout = p_dropout self.gin_channels = gin_channels - self.emb = nn.Embedding(len(symbols), hidden_channels) + self.emb = nn.Embedding(len(SYMBOLS), hidden_channels) nn.init.normal_(self.emb.weight, 0.0, hidden_channels**-0.5) - self.tone_emb = nn.Embedding(num_tones, hidden_channels) + self.tone_emb = nn.Embedding(NUM_TONES, hidden_channels) nn.init.normal_(self.tone_emb.weight, 0.0, hidden_channels**-0.5) - self.language_emb = nn.Embedding(num_languages, hidden_channels) + self.language_emb = nn.Embedding(NUM_LANGUAGES, hidden_channels) nn.init.normal_(self.language_emb.weight, 0.0, hidden_channels**-0.5) self.bert_proj = nn.Conv1d(1024, hidden_channels, 1) self.ja_bert_proj = nn.Conv1d(1024, hidden_channels, 1) diff --git a/models_jp_extra.py b/models_jp_extra.py index 3e87ced..16cacc7 100644 --- a/models_jp_extra.py +++ b/models_jp_extra.py @@ -12,7 +12,7 @@ from torch.nn import Conv1d, ConvTranspose1d, Conv2d from torch.nn.utils import weight_norm, remove_weight_norm, spectral_norm from style_bert_vits2.models.commons import init_weights, get_padding -from text import symbols, num_tones, num_languages +from style_bert_vits2.text_processing.symbols import SYMBOLS, NUM_TONES, NUM_LANGUAGES class DurationDiscriminator(nn.Module): # vits2 @@ -353,11 +353,11 @@ class TextEncoder(nn.Module): self.kernel_size = kernel_size self.p_dropout = p_dropout self.gin_channels = gin_channels - self.emb = nn.Embedding(len(symbols), hidden_channels) + self.emb = nn.Embedding(len(SYMBOLS), hidden_channels) nn.init.normal_(self.emb.weight, 0.0, hidden_channels**-0.5) - self.tone_emb = nn.Embedding(num_tones, hidden_channels) + self.tone_emb = nn.Embedding(NUM_TONES, hidden_channels) nn.init.normal_(self.tone_emb.weight, 0.0, hidden_channels**-0.5) - self.language_emb = nn.Embedding(num_languages, hidden_channels) + self.language_emb = nn.Embedding(NUM_LANGUAGES, hidden_channels) nn.init.normal_(self.language_emb.weight, 0.0, hidden_channels**-0.5) self.bert_proj = nn.Conv1d(1024, hidden_channels, 1) diff --git a/style_bert_vits2/text_processing/symbols.py b/style_bert_vits2/text_processing/symbols.py index 628edc6..d69bc1c 100644 --- a/style_bert_vits2/text_processing/symbols.py +++ b/style_bert_vits2/text_processing/symbols.py @@ -174,7 +174,7 @@ SYMBOLS = [PAD] + NORMAL_SYMBOLS + PUNCTUATION_SYMBOLS SIL_PHONEMES_IDS = [SYMBOLS.index(i) for i in PUNCTUATION_SYMBOLS] # Combine all tones -num_tones = NUM_ZH_TONES + NUM_JA_TONES + NUM_EN_TONES +NUM_TONES = NUM_ZH_TONES + NUM_JA_TONES + NUM_EN_TONES # Language maps LANGUAGE_ID_MAP = {"ZH": 0, "JP": 1, "EN": 2} diff --git a/text/__init__.py b/text/__init__.py index d8ae88d..ce4c008 100644 --- a/text/__init__.py +++ b/text/__init__.py @@ -1,6 +1,6 @@ -from text.symbols import * +from style_bert_vits2.text_processing.symbols import * -_symbol_to_id = {s: i for i, s in enumerate(symbols)} +_symbol_to_id = {s: i for i, s in enumerate(SYMBOLS)} def cleaned_text_to_sequence(cleaned_text, tones, language): @@ -11,9 +11,9 @@ def cleaned_text_to_sequence(cleaned_text, tones, language): List of integers corresponding to the symbols in the text """ phones = [_symbol_to_id[symbol] for symbol in cleaned_text] - tone_start = language_tone_start_map[language] + tone_start = LANGUAGE_TONE_START_MAP[language] tones = [i + tone_start for i in tones] - lang_id = language_id_map[language] + lang_id = LANGUAGE_ID_MAP[language] lang_ids = [lang_id for i in phones] return phones, tones, lang_ids diff --git a/text/chinese.py b/text/chinese.py index d9174ee..56dc4f3 100644 --- a/text/chinese.py +++ b/text/chinese.py @@ -4,7 +4,7 @@ import re import cn2an from pypinyin import lazy_pinyin, Style -from text.symbols import punctuation +from style_bert_vits2.text_processing.symbols import PUNCTUATIONS from text.tone_sandhi import ToneSandhi current_file_path = os.path.dirname(__file__) @@ -60,14 +60,14 @@ def replace_punctuation(text): replaced_text = pattern.sub(lambda x: rep_map[x.group()], text) replaced_text = re.sub( - r"[^\u4e00-\u9fa5" + "".join(punctuation) + r"]+", "", replaced_text + r"[^\u4e00-\u9fa5" + "".join(PUNCTUATIONS) + r"]+", "", replaced_text ) return replaced_text def g2p(text): - pattern = r"(?<=[{0}])\s*".format("".join(punctuation)) + pattern = r"(?<=[{0}])\s*".format("".join(PUNCTUATIONS)) sentences = [i for i in re.split(pattern, text) if i.strip() != ""] phones, tones, word2ph = _g2p(sentences) assert sum(word2ph) == len(phones) @@ -119,7 +119,7 @@ def _g2p(segments): # NOTE: post process for pypinyin outputs # we discriminate i, ii and iii if c == v: - assert c in punctuation + assert c in PUNCTUATIONS phone = [c] tone = "0" word2ph.append(1) diff --git a/text/english.py b/text/english.py index 4a2af95..f38ee84 100644 --- a/text/english.py +++ b/text/english.py @@ -4,8 +4,7 @@ import re from g2p_en import G2p from transformers import DebertaV2Tokenizer -from text import symbols -from text.symbols import punctuation +from style_bert_vits2.text_processing.symbols import PUNCTUATIONS, SYMBOLS current_file_path = os.path.dirname(__file__) CMU_DICT_PATH = os.path.join(current_file_path, "cmudict.rep") @@ -107,9 +106,9 @@ def post_replace_ph(ph): } if ph in rep_map.keys(): ph = rep_map[ph] - if ph in symbols: + if ph in SYMBOLS: return ph - if ph not in symbols: + if ph not in SYMBOLS: ph = "UNK" return ph @@ -399,13 +398,13 @@ def text_to_words(text): if t.startswith("▁"): words.append([t[1:]]) else: - if t in punctuation: + if t in PUNCTUATIONS: if idx == len(tokens) - 1: words.append([f"{t}"]) else: if ( not tokens[idx + 1].startswith("▁") - and tokens[idx + 1] not in punctuation + and tokens[idx + 1] not in PUNCTUATIONS ): if idx == 0: words.append([]) @@ -433,7 +432,7 @@ def g2p(text): if "'" in word: word = ["".join(word)] for w in word: - if w in punctuation: + if w in PUNCTUATIONS: temp_phones.append(w) temp_tones.append(0) continue diff --git a/text/japanese.py b/text/japanese.py index 12dc349..0dc6aa8 100644 --- a/text/japanese.py +++ b/text/japanese.py @@ -2,20 +2,18 @@ # compatible with Julius https://github.com/julius-speech/segmentation-kit import re import unicodedata -from pathlib import Path import pyopenjtalk from num2words import num2words from transformers import AutoTokenizer from style_bert_vits2.logging import logger -from text import punctuation -from text.japanese_mora_list import ( - mora_kata_to_mora_phonemes, - mora_phonemes_to_mora_kata, +from style_bert_vits2.text_processing.japanese.mora_list import ( + MORA_KATA_TO_MORA_PHONEMES, + MORA_PHONEMES_TO_MORA_KATA, ) - from style_bert_vits2.text_processing.japanese.user_dict import update_dict +from style_bert_vits2.text_processing.symbols import PUNCTUATIONS # 最初にpyopenjtalkの辞書を更新 update_dict() @@ -24,7 +22,7 @@ update_dict() COSONANTS = set( [ cosonant - for cosonant, _ in mora_kata_to_mora_phonemes.values() + for cosonant, _ in MORA_KATA_TO_MORA_PHONEMES.values() if cosonant is not None ] ) @@ -153,7 +151,7 @@ def replace_punctuation(text: str) -> str: # ↓ ギリシャ文字 + r"\u0370-\u03FF\u1F00-\u1FFF" # ↓ "!", "?", "…", ",", ".", "'", "-", 但し`…`はすでに`...`に変換されている - + "".join(punctuation) + r"]+", + + "".join(PUNCTUATIONS) + r"]+", # 上述以外の文字を削除 "", replaced_text, @@ -220,7 +218,7 @@ def g2p( # sep_textから、各単語を1文字1文字分割して、文字のリスト(のリスト)を作る sep_tokenized: list[list[str]] = [] for i in sep_text: - if i not in punctuation: + if i not in PUNCTUATIONS: sep_tokenized.append( tokenizer.tokenize(i) ) # ここでおそらく`i`が文字単位に分割される @@ -268,7 +266,7 @@ def phone_tone2kata_tone(phone_tone: list[tuple[str, int]]) -> list[tuple[str, i current_mora = "" for phone, next_phone, tone, next_tone in zip(phones, phones[1:], tones, tones[1:]): # zipの関係で最後の("_", 0)は無視されている - if phone in punctuation: + if phone in PUNCTUATIONS: result.append((phone, tone)) continue if phone in COSONANTS: # n以外の子音の場合 @@ -278,7 +276,7 @@ def phone_tone2kata_tone(phone_tone: list[tuple[str, int]]) -> list[tuple[str, i else: # phoneが母音もしくは「N」 current_mora += phone - result.append((mora_phonemes_to_mora_kata[current_mora], tone)) + result.append((MORA_PHONEMES_TO_MORA_KATA[current_mora], tone)) current_mora = "" return result @@ -287,10 +285,10 @@ def kata_tone2phone_tone(kata_tone: list[tuple[str, int]]) -> list[tuple[str, in """`phone_tone2kata_tone()`の逆。""" result: list[tuple[str, int]] = [("_", 0)] for mora, tone in kata_tone: - if mora in punctuation: + if mora in PUNCTUATIONS: result.append((mora, tone)) else: - cosonant, vowel = mora_kata_to_mora_phonemes[mora] + cosonant, vowel = MORA_KATA_TO_MORA_PHONEMES[mora] if cosonant is None: result.append((vowel, tone)) else: @@ -387,7 +385,7 @@ def text2sep_kata( assert yomi != "", f"Empty yomi: {word}" if yomi == "、": # wordは正規化されているので、`.`, `,`, `!`, `'`, `-`, `--` のいずれか - if not set(word).issubset(set(punctuation)): # 記号繰り返しか判定 + if not set(word).issubset(set(PUNCTUATIONS)): # 記号繰り返しか判定 # ここはpyopenjtalkが読めない文字等のときに起こる if raise_yomi_error: raise YomiError(f"Cannot read: {word} in:\n{norm_text}") @@ -581,7 +579,7 @@ def align_tones( result.append((phone, phone_tone_list[tone_index][1])) # 探すindexを1つ進める tone_index += 1 - elif phone in punctuation: + elif phone in PUNCTUATIONS: # phoneがpunctuationの場合 → (phone, 0)を追加 result.append((phone, 0)) else: @@ -606,16 +604,16 @@ def kata2phoneme_list(text: str) -> list[str]: `?` → ["?"] `!?!?!?!?!` → ["!", "?", "!", "?", "!", "?", "!", "?", "!"] """ - if set(text).issubset(set(punctuation)): + if set(text).issubset(set(PUNCTUATIONS)): return list(text) # `text`がカタカナ(`ー`含む)のみからなるかどうかをチェック if re.fullmatch(r"[\u30A0-\u30FF]+", text) is None: raise ValueError(f"Input must be katakana only: {text}") - sorted_keys = sorted(mora_kata_to_mora_phonemes.keys(), key=len, reverse=True) + sorted_keys = sorted(MORA_KATA_TO_MORA_PHONEMES.keys(), key=len, reverse=True) pattern = "|".join(map(re.escape, sorted_keys)) def mora2phonemes(mora: str) -> str: - cosonant, vowel = mora_kata_to_mora_phonemes[mora] + cosonant, vowel = MORA_KATA_TO_MORA_PHONEMES[mora] if cosonant is None: return f" {vowel}" return f" {cosonant} {vowel}" diff --git a/text/japanese_bert.py b/text/japanese_bert.py index dcee0f3..fbeb94d 100644 --- a/text/japanese_bert.py +++ b/text/japanese_bert.py @@ -4,7 +4,7 @@ import torch from transformers import AutoModelForMaskedLM, AutoTokenizer from config import config -from text.japanese import text2sep_kata, text_normalize +from style_bert_vits2.text_processing.japanese.g2p import text_to_sep_kata LOCAL_PATH = "./bert/deberta-v2-large-japanese-char-wwm" @@ -22,10 +22,10 @@ def get_bert_feature( ): # 各単語が何文字かを作る`word2ph`を使う必要があるので、読めない文字は必ず無視する # でないと`word2ph`の結果とテキストの文字数結果が整合性が取れない - text = "".join(text2sep_kata(text, raise_yomi_error=False)[0]) + text = "".join(text_to_sep_kata(text, raise_yomi_error=False)[0]) if assist_text: - assist_text = "".join(text2sep_kata(assist_text, raise_yomi_error=False)[0]) + assist_text = "".join(text_to_sep_kata(assist_text, raise_yomi_error=False)[0]) if ( sys.platform == "darwin" and torch.backends.mps.is_available() diff --git a/text/japanese_mora_list.py b/text/japanese_mora_list.py deleted file mode 100644 index b43e54d..0000000 --- a/text/japanese_mora_list.py +++ /dev/null @@ -1,232 +0,0 @@ -""" -VOICEVOXのソースコードからお借りして最低限に改造したコード。 -https://github.com/VOICEVOX/voicevox_engine/blob/master/voicevox_engine/tts_pipeline/mora_list.py -""" - -""" -以下のモーラ対応表はOpenJTalkのソースコードから取得し、 -カタカナ表記とモーラが一対一対応するように改造した。 -ライセンス表記: ------------------------------------------------------------------ - The Japanese TTS System "Open JTalk" - developed by HTS Working Group - http://open-jtalk.sourceforge.net/ ------------------------------------------------------------------ - - Copyright (c) 2008-2014 Nagoya Institute of Technology - Department of Computer Science - -All rights reserved. - -Redistribution and use in source and binary forms, with or -without modification, are permitted provided that the following -conditions are met: - -- Redistributions of source code must retain the above copyright - notice, this list of conditions and the following disclaimer. -- Redistributions in binary form must reproduce the above - copyright notice, this list of conditions and the following - disclaimer in the documentation and/or other materials provided - with the distribution. -- Neither the name of the HTS working group nor the names of its - contributors may be used to endorse or promote products derived - from this software without specific prior written permission. - -THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND -CONTRIBUTORS "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, -INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF -MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE -DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT OWNER OR CONTRIBUTORS -BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, -EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED -TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, -DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON -ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, -OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY -OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE -POSSIBILITY OF SUCH DAMAGE. -""" -from typing import Optional - -# (カタカナ, 子音, 母音)の順。子音がない場合はNoneを入れる。 -# 但し「ン」と「ッ」は母音のみという扱いで、「ン」は「N」、「ッ」は「q」とする。 -# (元々「ッ」は「cl」) -# また「デェ = dy e」はpyopenjtalkの出力(de e)と合わないため削除 -_mora_list_minimum: list[tuple[str, Optional[str], str]] = [ - ("ヴォ", "v", "o"), - ("ヴェ", "v", "e"), - ("ヴィ", "v", "i"), - ("ヴァ", "v", "a"), - ("ヴ", "v", "u"), - ("ン", None, "N"), - ("ワ", "w", "a"), - ("ロ", "r", "o"), - ("レ", "r", "e"), - ("ル", "r", "u"), - ("リョ", "ry", "o"), - ("リュ", "ry", "u"), - ("リャ", "ry", "a"), - ("リェ", "ry", "e"), - ("リ", "r", "i"), - ("ラ", "r", "a"), - ("ヨ", "y", "o"), - ("ユ", "y", "u"), - ("ヤ", "y", "a"), - ("モ", "m", "o"), - ("メ", "m", "e"), - ("ム", "m", "u"), - ("ミョ", "my", "o"), - ("ミュ", "my", "u"), - ("ミャ", "my", "a"), - ("ミェ", "my", "e"), - ("ミ", "m", "i"), - ("マ", "m", "a"), - ("ポ", "p", "o"), - ("ボ", "b", "o"), - ("ホ", "h", "o"), - ("ペ", "p", "e"), - ("ベ", "b", "e"), - ("ヘ", "h", "e"), - ("プ", "p", "u"), - ("ブ", "b", "u"), - ("フォ", "f", "o"), - ("フェ", "f", "e"), - ("フィ", "f", "i"), - ("ファ", "f", "a"), - ("フ", "f", "u"), - ("ピョ", "py", "o"), - ("ピュ", "py", "u"), - ("ピャ", "py", "a"), - ("ピェ", "py", "e"), - ("ピ", "p", "i"), - ("ビョ", "by", "o"), - ("ビュ", "by", "u"), - ("ビャ", "by", "a"), - ("ビェ", "by", "e"), - ("ビ", "b", "i"), - ("ヒョ", "hy", "o"), - ("ヒュ", "hy", "u"), - ("ヒャ", "hy", "a"), - ("ヒェ", "hy", "e"), - ("ヒ", "h", "i"), - ("パ", "p", "a"), - ("バ", "b", "a"), - ("ハ", "h", "a"), - ("ノ", "n", "o"), - ("ネ", "n", "e"), - ("ヌ", "n", "u"), - ("ニョ", "ny", "o"), - ("ニュ", "ny", "u"), - ("ニャ", "ny", "a"), - ("ニェ", "ny", "e"), - ("ニ", "n", "i"), - ("ナ", "n", "a"), - ("ドゥ", "d", "u"), - ("ド", "d", "o"), - ("トゥ", "t", "u"), - ("ト", "t", "o"), - ("デョ", "dy", "o"), - ("デュ", "dy", "u"), - ("デャ", "dy", "a"), - # ("デェ", "dy", "e"), - ("ディ", "d", "i"), - ("デ", "d", "e"), - ("テョ", "ty", "o"), - ("テュ", "ty", "u"), - ("テャ", "ty", "a"), - ("ティ", "t", "i"), - ("テ", "t", "e"), - ("ツォ", "ts", "o"), - ("ツェ", "ts", "e"), - ("ツィ", "ts", "i"), - ("ツァ", "ts", "a"), - ("ツ", "ts", "u"), - ("ッ", None, "q"), # 「cl」から「q」に変更 - ("チョ", "ch", "o"), - ("チュ", "ch", "u"), - ("チャ", "ch", "a"), - ("チェ", "ch", "e"), - ("チ", "ch", "i"), - ("ダ", "d", "a"), - ("タ", "t", "a"), - ("ゾ", "z", "o"), - ("ソ", "s", "o"), - ("ゼ", "z", "e"), - ("セ", "s", "e"), - ("ズィ", "z", "i"), - ("ズ", "z", "u"), - ("スィ", "s", "i"), - ("ス", "s", "u"), - ("ジョ", "j", "o"), - ("ジュ", "j", "u"), - ("ジャ", "j", "a"), - ("ジェ", "j", "e"), - ("ジ", "j", "i"), - ("ショ", "sh", "o"), - ("シュ", "sh", "u"), - ("シャ", "sh", "a"), - ("シェ", "sh", "e"), - ("シ", "sh", "i"), - ("ザ", "z", "a"), - ("サ", "s", "a"), - ("ゴ", "g", "o"), - ("コ", "k", "o"), - ("ゲ", "g", "e"), - ("ケ", "k", "e"), - ("グヮ", "gw", "a"), - ("グ", "g", "u"), - ("クヮ", "kw", "a"), - ("ク", "k", "u"), - ("ギョ", "gy", "o"), - ("ギュ", "gy", "u"), - ("ギャ", "gy", "a"), - ("ギェ", "gy", "e"), - ("ギ", "g", "i"), - ("キョ", "ky", "o"), - ("キュ", "ky", "u"), - ("キャ", "ky", "a"), - ("キェ", "ky", "e"), - ("キ", "k", "i"), - ("ガ", "g", "a"), - ("カ", "k", "a"), - ("オ", None, "o"), - ("エ", None, "e"), - ("ウォ", "w", "o"), - ("ウェ", "w", "e"), - ("ウィ", "w", "i"), - ("ウ", None, "u"), - ("イェ", "y", "e"), - ("イ", None, "i"), - ("ア", None, "a"), -] -_mora_list_additional: list[tuple[str, Optional[str], str]] = [ - ("ヴョ", "by", "o"), - ("ヴュ", "by", "u"), - ("ヴャ", "by", "a"), - ("ヲ", None, "o"), - ("ヱ", None, "e"), - ("ヰ", None, "i"), - ("ヮ", "w", "a"), - ("ョ", "y", "o"), - ("ュ", "y", "u"), - ("ヅ", "z", "u"), - ("ヂ", "j", "i"), - ("ヶ", "k", "e"), - ("ャ", "y", "a"), - ("ォ", None, "o"), - ("ェ", None, "e"), - ("ゥ", None, "u"), - ("ィ", None, "i"), - ("ァ", None, "a"), -] - -# 例: "vo" -> "ヴォ", "a" -> "ア" -mora_phonemes_to_mora_kata: dict[str, str] = { - (consonant or "") + vowel: kana for [kana, consonant, vowel] in _mora_list_minimum -} - -# 例: "ヴォ" -> ("v", "o"), "ア" -> (None, "a") -mora_kata_to_mora_phonemes: dict[str, tuple[Optional[str], str]] = { - kana: (consonant, vowel) - for [kana, consonant, vowel] in _mora_list_minimum + _mora_list_additional -} diff --git a/text/symbols.py b/text/symbols.py deleted file mode 100644 index 846de64..0000000 --- a/text/symbols.py +++ /dev/null @@ -1,187 +0,0 @@ -punctuation = ["!", "?", "…", ",", ".", "'", "-"] -pu_symbols = punctuation + ["SP", "UNK"] -pad = "_" - -# chinese -zh_symbols = [ - "E", - "En", - "a", - "ai", - "an", - "ang", - "ao", - "b", - "c", - "ch", - "d", - "e", - "ei", - "en", - "eng", - "er", - "f", - "g", - "h", - "i", - "i0", - "ia", - "ian", - "iang", - "iao", - "ie", - "in", - "ing", - "iong", - "ir", - "iu", - "j", - "k", - "l", - "m", - "n", - "o", - "ong", - "ou", - "p", - "q", - "r", - "s", - "sh", - "t", - "u", - "ua", - "uai", - "uan", - "uang", - "ui", - "un", - "uo", - "v", - "van", - "ve", - "vn", - "w", - "x", - "y", - "z", - "zh", - "AA", - "EE", - "OO", -] -num_zh_tones = 6 - -# japanese -ja_symbols = [ - "N", - "a", - "a:", - "b", - "by", - "ch", - "d", - "dy", - "e", - "e:", - "f", - "g", - "gy", - "h", - "hy", - "i", - "i:", - "j", - "k", - "ky", - "m", - "my", - "n", - "ny", - "o", - "o:", - "p", - "py", - "q", - "r", - "ry", - "s", - "sh", - "t", - "ts", - "ty", - "u", - "u:", - "w", - "y", - "z", - "zy", -] -num_ja_tones = 2 - -# English -en_symbols = [ - "aa", - "ae", - "ah", - "ao", - "aw", - "ay", - "b", - "ch", - "d", - "dh", - "eh", - "er", - "ey", - "f", - "g", - "hh", - "ih", - "iy", - "jh", - "k", - "l", - "m", - "n", - "ng", - "ow", - "oy", - "p", - "r", - "s", - "sh", - "t", - "th", - "uh", - "uw", - "V", - "w", - "y", - "z", - "zh", -] -num_en_tones = 4 - -# combine all symbols -normal_symbols = sorted(set(zh_symbols + ja_symbols + en_symbols)) -symbols = [pad] + normal_symbols + pu_symbols -sil_phonemes_ids = [symbols.index(i) for i in pu_symbols] - -# combine all tones -num_tones = num_zh_tones + num_ja_tones + num_en_tones - -# language maps -language_id_map = {"ZH": 0, "JP": 1, "EN": 2} -num_languages = len(language_id_map.keys()) - -language_tone_start_map = { - "ZH": 0, - "JP": num_zh_tones, - "EN": num_zh_tones + num_ja_tones, -} - -if __name__ == "__main__": - a = set(zh_symbols) - b = set(en_symbols) - print(sorted(a & b)) diff --git a/train_ms.py b/train_ms.py index 3e9e707..783b157 100644 --- a/train_ms.py +++ b/train_ms.py @@ -29,7 +29,7 @@ from data_utils import ( from losses import discriminator_loss, feature_loss, generator_loss, kl_loss from mel_processing import mel_spectrogram_torch, spec_to_mel_torch from models import DurationDiscriminator, MultiPeriodDiscriminator, SynthesizerTrn -from text.symbols import symbols +from style_bert_vits2.text_processing.symbols import SYMBOLS torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = ( @@ -279,7 +279,7 @@ def run(): logger.info("Using normal encoder for VITS1") net_g = SynthesizerTrn( - len(symbols), + len(SYMBOLS), hps.data.filter_length // 2 + 1, hps.train.segment_size // hps.data.hop_length, n_speakers=hps.data.n_speakers, diff --git a/train_ms_jp_extra.py b/train_ms_jp_extra.py index 9a964c3..e4e4e56 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -34,7 +34,7 @@ from models_jp_extra import ( SynthesizerTrn, WavLMDiscriminator, ) -from text.symbols import symbols +from style_bert_vits2.text_processing.symbols import SYMBOLS torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = ( @@ -293,7 +293,7 @@ def run(): logger.info("Using normal encoder for VITS1") net_g = SynthesizerTrn( - len(symbols), + len(SYMBOLS), hps.data.filter_length // 2 + 1, hps.train.segment_size // hps.data.hop_length, n_speakers=hps.data.n_speakers, From ca4c03c67bbf838eec27ca6811c90a5b9d9a5ec2 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Wed, 6 Mar 2024 23:03:06 +0000 Subject: [PATCH 08/64] Fix: import error --- common/subprocess_utils.py | 4 ++-- common/tts_model.py | 6 ++---- 2 files changed, 4 insertions(+), 6 deletions(-) diff --git a/common/subprocess_utils.py b/common/subprocess_utils.py index 40426f7..d7f3fc4 100644 --- a/common/subprocess_utils.py +++ b/common/subprocess_utils.py @@ -1,8 +1,8 @@ import subprocess import sys -from .log import logger -from .stdout_wrapper import SAFE_STDOUT +from style_bert_vits2.logging import logger +from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT python = sys.executable diff --git a/common/tts_model.py b/common/tts_model.py index e09787e..3f17d15 100644 --- a/common/tts_model.py +++ b/common/tts_model.py @@ -1,11 +1,9 @@ -import os import warnings from pathlib import Path from typing import Optional, Union import gradio as gr import numpy as np - import torch from gradio.processing_utils import convert_to_16_bit_wav @@ -14,7 +12,7 @@ from infer import get_net_g, infer from models import SynthesizerTrn from models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra -from .constants import ( +from style_bert_vits2.constants import ( DEFAULT_ASSIST_TEXT_WEIGHT, DEFAULT_LENGTH, DEFAULT_LINE_SPLIT, @@ -25,7 +23,7 @@ from .constants import ( DEFAULT_STYLE, DEFAULT_STYLE_WEIGHT, ) -from .log import logger +from style_bert_vits2.logging import logger def adjust_voice(fs, wave, pitch_scale, intonation_scale): From a52fda7a88ad296831419993364f261559147939 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Wed, 6 Mar 2024 23:11:27 +0000 Subject: [PATCH 09/64] Refactor: moved common/subprocess_utils.py to style_bert_vits2/utils/subprocess.py --- common/subprocess_utils.py | 33 ---------------- style_bert_vits2/utils/subprocess.py | 56 ++++++++++++++++++++++++++++ webui_dataset.py | 2 +- webui_train.py | 2 +- 4 files changed, 58 insertions(+), 35 deletions(-) delete mode 100644 common/subprocess_utils.py create mode 100644 style_bert_vits2/utils/subprocess.py diff --git a/common/subprocess_utils.py b/common/subprocess_utils.py deleted file mode 100644 index d7f3fc4..0000000 --- a/common/subprocess_utils.py +++ /dev/null @@ -1,33 +0,0 @@ -import subprocess -import sys - -from style_bert_vits2.logging import logger -from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT - -python = sys.executable - - -def run_script_with_log(cmd: list[str], ignore_warning=False) -> tuple[bool, str]: - logger.info(f"Running: {' '.join(cmd)}") - result = subprocess.run( - [python] + cmd, - stdout=SAFE_STDOUT, # type: ignore - stderr=subprocess.PIPE, - text=True, - encoding="utf-8", - ) - if result.returncode != 0: - logger.error(f"Error: {' '.join(cmd)}\n{result.stderr}") - return False, result.stderr - elif result.stderr and not ignore_warning: - logger.warning(f"Warning: {' '.join(cmd)}\n{result.stderr}") - return True, result.stderr - logger.success(f"Success: {' '.join(cmd)}") - return True, "" - - -def second_elem_of(original_function): - def inner_function(*args, **kwargs): - return original_function(*args, **kwargs)[1] - - return inner_function diff --git a/style_bert_vits2/utils/subprocess.py b/style_bert_vits2/utils/subprocess.py new file mode 100644 index 0000000..b152702 --- /dev/null +++ b/style_bert_vits2/utils/subprocess.py @@ -0,0 +1,56 @@ +import subprocess +import sys +from typing import Any, Callable + +from style_bert_vits2.logging import logger +from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT + +PYTHON = sys.executable + + +def run_script_with_log(cmd: list[str], ignore_warning: bool = False) -> tuple[bool, str]: + """ + 指定されたコマンドを実行し、そのログを記録する + + Args: + cmd: 実行するコマンドのリスト + ignore_warning: 警告を無視するかどうかのフラグ + + Returns: + tuple[bool, str]: 実行が成功したかどうかのブール値と、エラーまたは警告のメッセージ(ある場合) + """ + + logger.info(f"Running: {' '.join(cmd)}") + result = subprocess.run( + [PYTHON] + cmd, + stdout = SAFE_STDOUT, + stderr = subprocess.PIPE, + text = True, + encoding = "utf-8", + ) + if result.returncode != 0: + logger.error(f"Error: {' '.join(cmd)}\n{result.stderr}") + return False, result.stderr + elif result.stderr and not ignore_warning: + logger.warning(f"Warning: {' '.join(cmd)}\n{result.stderr}") + return True, result.stderr + logger.success(f"Success: {' '.join(cmd)}") + + return True, "" + + +def second_elem_of(original_function: Callable[..., tuple[Any, Any]]) -> Callable[..., Any]: + """ + 与えられた関数をラップし、その戻り値の 2 番目の要素のみを返す関数を生成する + + Args: + original_function (Callable[..., tuple[Any, Any]])): ラップする元の関数 + + Returns: + Callable[..., Any]: 元の関数の戻り値の 2 番目の要素のみを返す関数 + """ + + def inner_function(*args, **kwargs) -> Any: # type: ignore + return original_function(*args, **kwargs)[1] + + return inner_function diff --git a/webui_dataset.py b/webui_dataset.py index 169cc4d..3ad63c1 100644 --- a/webui_dataset.py +++ b/webui_dataset.py @@ -6,7 +6,7 @@ import yaml from style_bert_vits2.constants import GRADIO_THEME from style_bert_vits2.logging import logger -from common.subprocess_utils import run_script_with_log +from style_bert_vits2.utils.subprocess import run_script_with_log # Get path settings with open(os.path.join("configs", "paths.yml"), "r", encoding="utf-8") as f: diff --git a/webui_train.py b/webui_train.py index 0f272e8..6c89b26 100644 --- a/webui_train.py +++ b/webui_train.py @@ -17,7 +17,7 @@ import yaml from style_bert_vits2.constants import GRADIO_THEME, VERSION from style_bert_vits2.logging import logger from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT -from common.subprocess_utils import run_script_with_log, second_elem_of +from style_bert_vits2.utils.subprocess import run_script_with_log, second_elem_of logger_handler = None tensorboard_executed = False From 89825e68d88586225d93b7486d866578f1e4990c Mon Sep 17 00:00:00 2001 From: tsukumi Date: Wed, 6 Mar 2024 23:43:25 +0000 Subject: [PATCH 10/64] Refactor: moved model, attentions definitions and inference code to style_bert_vits2/models/ The code has not yet been cleaned up, just moved. --- app.py | 2 +- common/tts_model.py | 7 ++--- .../models/attentions.py | 3 +- infer.py => style_bert_vits2/models/infer.py | 8 +++--- .../models/models.py | 7 ++--- .../models/models_jp_extra.py | 26 ++++++++--------- .../models/modules.py | 28 +++++++++---------- train_ms.py | 13 +++++---- train_ms_jp_extra.py | 10 +++---- webui.py | 15 +++++----- 10 files changed, 58 insertions(+), 61 deletions(-) rename attentions.py => style_bert_vits2/models/attentions.py (99%) rename infer.py => style_bert_vits2/models/infer.py (95%) rename models.py => style_bert_vits2/models/models.py (99%) rename models_jp_extra.py => style_bert_vits2/models/models_jp_extra.py (98%) rename modules.py => style_bert_vits2/models/modules.py (95%) diff --git a/app.py b/app.py index c70a89d..a0a215c 100644 --- a/app.py +++ b/app.py @@ -26,7 +26,7 @@ from style_bert_vits2.constants import ( ) from style_bert_vits2.logging import logger from common.tts_model import ModelHolder -from infer import InvalidToneError +from style_bert_vits2.models.infer import InvalidToneError from style_bert_vits2.text_processing.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone from style_bert_vits2.text_processing.japanese.normalizer import normalize_text diff --git a/common/tts_model.py b/common/tts_model.py index 3f17d15..c924686 100644 --- a/common/tts_model.py +++ b/common/tts_model.py @@ -8,10 +8,6 @@ import torch from gradio.processing_utils import convert_to_16_bit_wav import utils -from infer import get_net_g, infer -from models import SynthesizerTrn -from models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra - from style_bert_vits2.constants import ( DEFAULT_ASSIST_TEXT_WEIGHT, DEFAULT_LENGTH, @@ -23,6 +19,9 @@ from style_bert_vits2.constants import ( DEFAULT_STYLE, DEFAULT_STYLE_WEIGHT, ) +from style_bert_vits2.models.infer import get_net_g, infer +from style_bert_vits2.models.models import SynthesizerTrn +from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra from style_bert_vits2.logging import logger diff --git a/attentions.py b/style_bert_vits2/models/attentions.py similarity index 99% rename from attentions.py rename to style_bert_vits2/models/attentions.py index 1cca086..6d43e08 100644 --- a/attentions.py +++ b/style_bert_vits2/models/attentions.py @@ -4,7 +4,6 @@ from torch import nn from torch.nn import functional as F from style_bert_vits2.models import commons -from style_bert_vits2.logging import logger as logging class LayerNorm(nn.Module): @@ -67,7 +66,7 @@ class Encoder(nn.Module): self.cond_layer_idx = ( kwargs["cond_layer_idx"] if "cond_layer_idx" in kwargs else 2 ) - # logging.debug(self.gin_channels, self.cond_layer_idx) + # logger.debug(self.gin_channels, self.cond_layer_idx) assert ( self.cond_layer_idx < self.n_layers ), "cond_layer_idx should be less than n_layers" diff --git a/infer.py b/style_bert_vits2/models/infer.py similarity index 95% rename from infer.py rename to style_bert_vits2/models/infer.py index 4afd048..9abd378 100644 --- a/infer.py +++ b/style_bert_vits2/models/infer.py @@ -1,13 +1,13 @@ import torch -from style_bert_vits2.models import commons import utils -from models import SynthesizerTrn -from models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra from text import cleaned_text_to_sequence, get_bert from text.cleaner import clean_text -from style_bert_vits2.text_processing.symbols import SYMBOLS from style_bert_vits2.logging import logger +from style_bert_vits2.models import commons +from style_bert_vits2.models.models import SynthesizerTrn +from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra +from style_bert_vits2.text_processing.symbols import SYMBOLS class InvalidToneError(ValueError): diff --git a/models.py b/style_bert_vits2/models/models.py similarity index 99% rename from models.py rename to style_bert_vits2/models/models.py index ef581f2..a8c6695 100644 --- a/models.py +++ b/style_bert_vits2/models/models.py @@ -1,5 +1,4 @@ import math -import warnings import torch from torch import nn @@ -7,10 +6,10 @@ from torch.nn import Conv1d, Conv2d, ConvTranspose1d from torch.nn import functional as F from torch.nn.utils import remove_weight_norm, spectral_norm, weight_norm -import attentions -from style_bert_vits2.models import commons -import modules import monotonic_align +from style_bert_vits2.models import attentions +from style_bert_vits2.models import commons +from style_bert_vits2.models import modules from style_bert_vits2.models.commons import get_padding, init_weights from style_bert_vits2.text_processing.symbols import NUM_LANGUAGES, NUM_TONES, SYMBOLS diff --git a/models_jp_extra.py b/style_bert_vits2/models/models_jp_extra.py similarity index 98% rename from models_jp_extra.py rename to style_bert_vits2/models/models_jp_extra.py index 16cacc7..8c4a4ec 100644 --- a/models_jp_extra.py +++ b/style_bert_vits2/models/models_jp_extra.py @@ -1,17 +1,15 @@ import math + import torch from torch import nn +from torch.nn import Conv1d, Conv2d, ConvTranspose1d from torch.nn import functional as F +from torch.nn.utils import remove_weight_norm, spectral_norm, weight_norm -from style_bert_vits2.models import commons -import modules -import attentions import monotonic_align - -from torch.nn import Conv1d, ConvTranspose1d, Conv2d -from torch.nn.utils import weight_norm, remove_weight_norm, spectral_norm - -from style_bert_vits2.models.commons import init_weights, get_padding +from style_bert_vits2.models import attentions +from style_bert_vits2.models import commons +from style_bert_vits2.models import modules from style_bert_vits2.text_processing.symbols import SYMBOLS, NUM_TONES, NUM_LANGUAGES @@ -529,7 +527,7 @@ class Generator(torch.nn.Module): self.resblocks.append(resblock(ch, k, d)) self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False) - self.ups.apply(init_weights) + self.ups.apply(commons.init_weights) if gin_channels != 0: self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1) @@ -577,7 +575,7 @@ class DiscriminatorP(torch.nn.Module): 32, (kernel_size, 1), (stride, 1), - padding=(get_padding(kernel_size, 1), 0), + padding=(commons.get_padding(kernel_size, 1), 0), ) ), norm_f( @@ -586,7 +584,7 @@ class DiscriminatorP(torch.nn.Module): 128, (kernel_size, 1), (stride, 1), - padding=(get_padding(kernel_size, 1), 0), + padding=(commons.get_padding(kernel_size, 1), 0), ) ), norm_f( @@ -595,7 +593,7 @@ class DiscriminatorP(torch.nn.Module): 512, (kernel_size, 1), (stride, 1), - padding=(get_padding(kernel_size, 1), 0), + padding=(commons.get_padding(kernel_size, 1), 0), ) ), norm_f( @@ -604,7 +602,7 @@ class DiscriminatorP(torch.nn.Module): 1024, (kernel_size, 1), (stride, 1), - padding=(get_padding(kernel_size, 1), 0), + padding=(commons.get_padding(kernel_size, 1), 0), ) ), norm_f( @@ -613,7 +611,7 @@ class DiscriminatorP(torch.nn.Module): 1024, (kernel_size, 1), 1, - padding=(get_padding(kernel_size, 1), 0), + padding=(commons.get_padding(kernel_size, 1), 0), ) ), ] diff --git a/modules.py b/style_bert_vits2/models/modules.py similarity index 95% rename from modules.py rename to style_bert_vits2/models/modules.py index 68b0b9a..e0885c4 100644 --- a/modules.py +++ b/style_bert_vits2/models/modules.py @@ -1,16 +1,14 @@ import math -import warnings import torch from torch import nn from torch.nn import Conv1d from torch.nn import functional as F from torch.nn.utils import remove_weight_norm, weight_norm +from transforms import piecewise_rational_quadratic_transform from style_bert_vits2.models import commons -from attentions import Encoder -from style_bert_vits2.models.commons import get_padding, init_weights -from transforms import piecewise_rational_quadratic_transform +from style_bert_vits2.models.attentions import Encoder LRELU_SLOPE = 0.1 @@ -231,7 +229,7 @@ class ResBlock1(torch.nn.Module): kernel_size, 1, dilation=dilation[0], - padding=get_padding(kernel_size, dilation[0]), + padding=commons.get_padding(kernel_size, dilation[0]), ) ), weight_norm( @@ -241,7 +239,7 @@ class ResBlock1(torch.nn.Module): kernel_size, 1, dilation=dilation[1], - padding=get_padding(kernel_size, dilation[1]), + padding=commons.get_padding(kernel_size, dilation[1]), ) ), weight_norm( @@ -251,12 +249,12 @@ class ResBlock1(torch.nn.Module): kernel_size, 1, dilation=dilation[2], - padding=get_padding(kernel_size, dilation[2]), + padding=commons.get_padding(kernel_size, dilation[2]), ) ), ] ) - self.convs1.apply(init_weights) + self.convs1.apply(commons.init_weights) self.convs2 = nn.ModuleList( [ @@ -267,7 +265,7 @@ class ResBlock1(torch.nn.Module): kernel_size, 1, dilation=1, - padding=get_padding(kernel_size, 1), + padding=commons.get_padding(kernel_size, 1), ) ), weight_norm( @@ -277,7 +275,7 @@ class ResBlock1(torch.nn.Module): kernel_size, 1, dilation=1, - padding=get_padding(kernel_size, 1), + padding=commons.get_padding(kernel_size, 1), ) ), weight_norm( @@ -287,12 +285,12 @@ class ResBlock1(torch.nn.Module): kernel_size, 1, dilation=1, - padding=get_padding(kernel_size, 1), + padding=commons.get_padding(kernel_size, 1), ) ), ] ) - self.convs2.apply(init_weights) + self.convs2.apply(commons.init_weights) def forward(self, x, x_mask=None): for c1, c2 in zip(self.convs1, self.convs2): @@ -328,7 +326,7 @@ class ResBlock2(torch.nn.Module): kernel_size, 1, dilation=dilation[0], - padding=get_padding(kernel_size, dilation[0]), + padding=commons.get_padding(kernel_size, dilation[0]), ) ), weight_norm( @@ -338,12 +336,12 @@ class ResBlock2(torch.nn.Module): kernel_size, 1, dilation=dilation[1], - padding=get_padding(kernel_size, dilation[1]), + padding=commons.get_padding(kernel_size, dilation[1]), ) ), ] ) - self.convs.apply(init_weights) + self.convs.apply(commons.init_weights) def forward(self, x, x_mask=None): for c in self.convs: diff --git a/train_ms.py b/train_ms.py index 783b157..05cd4cb 100644 --- a/train_ms.py +++ b/train_ms.py @@ -1,6 +1,5 @@ import argparse import datetime -import gc import os import platform @@ -15,11 +14,8 @@ from torch.utils.tensorboard import SummaryWriter from tqdm import tqdm # logging.getLogger("numba").setLevel(logging.WARNING) -from style_bert_vits2.models import commons import default_style import utils -from style_bert_vits2.logging import logger -from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from config import config from data_utils import ( DistributedBucketSampler, @@ -28,8 +24,15 @@ from data_utils import ( ) from losses import discriminator_loss, feature_loss, generator_loss, kl_loss from mel_processing import mel_spectrogram_torch, spec_to_mel_torch -from models import DurationDiscriminator, MultiPeriodDiscriminator, SynthesizerTrn +from style_bert_vits2.logging import logger +from style_bert_vits2.models import commons +from style_bert_vits2.models.models import ( + DurationDiscriminator, + MultiPeriodDiscriminator, + SynthesizerTrn, +) from style_bert_vits2.text_processing.symbols import SYMBOLS +from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = ( diff --git a/train_ms_jp_extra.py b/train_ms_jp_extra.py index e4e4e56..1a287d9 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -6,20 +6,17 @@ import platform import torch import torch.distributed as dist +from huggingface_hub import HfApi from torch.cuda.amp import GradScaler, autocast from torch.nn import functional as F from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data import DataLoader from torch.utils.tensorboard import SummaryWriter from tqdm import tqdm -from huggingface_hub import HfApi # logging.getLogger("numba").setLevel(logging.WARNING) -from style_bert_vits2.models import commons import default_style import utils -from style_bert_vits2.logging import logger -from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from config import config from data_utils import ( DistributedBucketSampler, @@ -28,13 +25,16 @@ from data_utils import ( ) from losses import WavLMLoss, discriminator_loss, feature_loss, generator_loss, kl_loss from mel_processing import mel_spectrogram_torch, spec_to_mel_torch -from models_jp_extra import ( +from style_bert_vits2.logging import logger +from style_bert_vits2.models import commons +from style_bert_vits2.models.models_jp_extra import ( DurationDiscriminator, MultiPeriodDiscriminator, SynthesizerTrn, WavLMDiscriminator, ) from style_bert_vits2.text_processing.symbols import SYMBOLS +from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = ( diff --git a/webui.py b/webui.py index 90318a1..1f31e94 100644 --- a/webui.py +++ b/webui.py @@ -19,15 +19,16 @@ logging.basicConfig( logger = logging.getLogger(__name__) -import torch -import utils -from infer import infer, latest_version, get_net_g, infer_multilang import gradio as gr -import webbrowser -import numpy as np -from config import config -from tools.translate import translate import librosa +import numpy as np +import torch +import webbrowser + +import utils +from config import config +from style_bert_vits2.models.infer import infer, latest_version, get_net_g, infer_multilang +from tools.translate import translate net_g = None From e826faf62e69552d5cbcbab079ab0f4ffe5d8972 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Thu, 7 Mar 2024 00:24:28 +0000 Subject: [PATCH 11/64] Refactor: moved text/cleaner.py to style_bert_vits2/text_processing/ --- preprocess_text.py | 8 +- style_bert_vits2/models/infer.py | 633 ++++++++++---------- style_bert_vits2/text_processing/cleaner.py | 46 ++ text/chinese.py | 4 +- text/cleaner.py | 26 - text/english.py | 2 +- text/japanese.py | 8 +- 7 files changed, 376 insertions(+), 351 deletions(-) create mode 100644 style_bert_vits2/text_processing/cleaner.py delete mode 100644 text/cleaner.py diff --git a/preprocess_text.py b/preprocess_text.py index 126ba2c..92e00b9 100644 --- a/preprocess_text.py +++ b/preprocess_text.py @@ -7,10 +7,10 @@ from typing import Optional import click 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 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 @@ -72,7 +72,7 @@ def preprocess( utt, spk, language, text = line.strip().split("|") norm_text, phones, tones, word2ph = clean_text( text=text, - language=language, + language=language, # type: ignore use_jp_extra=use_jp_extra, raise_yomi_error=(yomi_error != "use"), ) diff --git a/style_bert_vits2/models/infer.py b/style_bert_vits2/models/infer.py index 9abd378..e0d8691 100644 --- a/style_bert_vits2/models/infer.py +++ b/style_bert_vits2/models/infer.py @@ -1,314 +1,319 @@ -import torch - -import utils -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.models import commons -from style_bert_vits2.models.models import SynthesizerTrn -from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra -from style_bert_vits2.text_processing.symbols import SYMBOLS - - -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") - net_g = SynthesizerTrnJPExtra( - len(SYMBOLS), - hps.data.filter_length // 2 + 1, - hps.train.segment_size // hps.data.hop_length, - n_speakers=hps.data.n_speakers, - **hps.model, - ).to(device) - else: - logger.info("Using normal model") - net_g = SynthesizerTrn( - len(SYMBOLS), - hps.data.filter_length // 2 + 1, - hps.train.segment_size // hps.data.hop_length, - n_speakers=hps.data.n_speakers, - **hps.model, - ).to(device) - net_g.state_dict() - _ = net_g.eval() - if model_path.endswith(".pth") or model_path.endswith(".pt"): - _ = utils.load_checkpoint(model_path, net_g, None, skip_optimizer=True) - elif model_path.endswith(".safetensors"): - _ = utils.load_safetensors(model_path, net_g, True) - else: - raise ValueError(f"Unknown model format: {model_path}") - return net_g - - -def get_text( - text, - language_str, - hps, - device, - assist_text=None, - assist_text_weight=0.7, - given_tone=None, -): - use_jp_extra = hps.version.endswith("JP-Extra") - # 推論のときにのみ呼び出されるので、raise_yomi_errorはFalseに設定 - norm_text, phone, tone, word2ph = clean_text( - text, language_str, use_jp_extra, raise_yomi_error=False - ) - if given_tone is not None: - if len(given_tone) != len(phone): - raise InvalidToneError( - f"Length of given_tone ({len(given_tone)}) != length of phone ({len(phone)})" - ) - tone = given_tone - phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str) - - if hps.data.add_blank: - phone = commons.intersperse(phone, 0) - tone = commons.intersperse(tone, 0) - language = commons.intersperse(language, 0) - for i in range(len(word2ph)): - word2ph[i] = word2ph[i] * 2 - word2ph[0] += 1 - bert_ori = get_bert( - norm_text, - word2ph, - language_str, - device, - assist_text, - assist_text_weight, - ) - del word2ph - assert bert_ori.shape[-1] == len(phone), phone - - if language_str == "ZH": - bert = bert_ori - ja_bert = torch.zeros(1024, len(phone)) - en_bert = torch.zeros(1024, len(phone)) - elif language_str == "JP": - bert = torch.zeros(1024, len(phone)) - ja_bert = bert_ori - en_bert = torch.zeros(1024, len(phone)) - elif language_str == "EN": - bert = torch.zeros(1024, len(phone)) - ja_bert = torch.zeros(1024, len(phone)) - en_bert = bert_ori - else: - raise ValueError("language_str should be ZH, JP or EN") - - assert bert.shape[-1] == len( - phone - ), f"Bert seq len {bert.shape[-1]} != {len(phone)}" - - phone = torch.LongTensor(phone) - tone = torch.LongTensor(tone) - language = torch.LongTensor(language) - return bert, ja_bert, en_bert, phone, tone, language - - -def infer( - text, - style_vec, - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - sid: int, # In the original Bert-VITS2, its speaker_name: str, but here it's id - language, - hps, - net_g, - device, - skip_start=False, - skip_end=False, - assist_text=None, - assist_text_weight=0.7, - given_tone=None, -): - is_jp_extra = hps.version.endswith("JP-Extra") - bert, ja_bert, en_bert, phones, tones, lang_ids = get_text( - text, - language, - hps, - device, - assist_text=assist_text, - assist_text_weight=assist_text_weight, - given_tone=given_tone, - ) - if skip_start: - phones = phones[3:] - tones = tones[3:] - lang_ids = lang_ids[3:] - bert = bert[:, 3:] - ja_bert = ja_bert[:, 3:] - en_bert = en_bert[:, 3:] - if skip_end: - phones = phones[:-2] - tones = tones[:-2] - lang_ids = lang_ids[:-2] - bert = bert[:, :-2] - ja_bert = ja_bert[:, :-2] - en_bert = en_bert[:, :-2] - with torch.no_grad(): - x_tst = phones.to(device).unsqueeze(0) - tones = tones.to(device).unsqueeze(0) - lang_ids = lang_ids.to(device).unsqueeze(0) - bert = bert.to(device).unsqueeze(0) - ja_bert = ja_bert.to(device).unsqueeze(0) - en_bert = en_bert.to(device).unsqueeze(0) - x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device) - style_vec = torch.from_numpy(style_vec).to(device).unsqueeze(0) - del phones - sid_tensor = torch.LongTensor([sid]).to(device) - if is_jp_extra: - output = net_g.infer( - x_tst, - x_tst_lengths, - sid_tensor, - tones, - lang_ids, - ja_bert, - style_vec=style_vec, - sdp_ratio=sdp_ratio, - noise_scale=noise_scale, - noise_scale_w=noise_scale_w, - length_scale=length_scale, - ) - else: - output = net_g.infer( - x_tst, - x_tst_lengths, - sid_tensor, - tones, - lang_ids, - bert, - ja_bert, - en_bert, - style_vec=style_vec, - sdp_ratio=sdp_ratio, - noise_scale=noise_scale, - noise_scale_w=noise_scale_w, - length_scale=length_scale, - ) - audio = output[0][0, 0].data.cpu().float().numpy() - del ( - x_tst, - tones, - lang_ids, - bert, - x_tst_lengths, - sid_tensor, - ja_bert, - en_bert, - style_vec, - ) # , emo - if torch.cuda.is_available(): - torch.cuda.empty_cache() - return audio - - -def infer_multilang( - text, - style_vec, - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - sid, - language, - hps, - net_g, - device, - skip_start=False, - skip_end=False, -): - bert, ja_bert, en_bert, phones, tones, lang_ids = [], [], [], [], [], [] - # emo = get_emo_(reference_audio, emotion, sid) - # if isinstance(reference_audio, np.ndarray): - # emo = get_clap_audio_feature(reference_audio, device) - # else: - # emo = get_clap_text_feature(emotion, device) - # emo = torch.squeeze(emo, dim=1) - for idx, (txt, lang) in enumerate(zip(text, language)): - _skip_start = (idx != 0) or (skip_start and idx == 0) - _skip_end = (idx != len(language) - 1) or skip_end - ( - temp_bert, - temp_ja_bert, - temp_en_bert, - temp_phones, - temp_tones, - temp_lang_ids, - ) = get_text(txt, lang, hps, device) - if _skip_start: - temp_bert = temp_bert[:, 3:] - temp_ja_bert = temp_ja_bert[:, 3:] - temp_en_bert = temp_en_bert[:, 3:] - temp_phones = temp_phones[3:] - temp_tones = temp_tones[3:] - temp_lang_ids = temp_lang_ids[3:] - if _skip_end: - temp_bert = temp_bert[:, :-2] - temp_ja_bert = temp_ja_bert[:, :-2] - temp_en_bert = temp_en_bert[:, :-2] - temp_phones = temp_phones[:-2] - temp_tones = temp_tones[:-2] - temp_lang_ids = temp_lang_ids[:-2] - bert.append(temp_bert) - ja_bert.append(temp_ja_bert) - en_bert.append(temp_en_bert) - phones.append(temp_phones) - tones.append(temp_tones) - lang_ids.append(temp_lang_ids) - bert = torch.concatenate(bert, dim=1) - ja_bert = torch.concatenate(ja_bert, dim=1) - en_bert = torch.concatenate(en_bert, dim=1) - phones = torch.concatenate(phones, dim=0) - tones = torch.concatenate(tones, dim=0) - lang_ids = torch.concatenate(lang_ids, dim=0) - with torch.no_grad(): - x_tst = phones.to(device).unsqueeze(0) - tones = tones.to(device).unsqueeze(0) - lang_ids = lang_ids.to(device).unsqueeze(0) - bert = bert.to(device).unsqueeze(0) - ja_bert = ja_bert.to(device).unsqueeze(0) - en_bert = en_bert.to(device).unsqueeze(0) - # emo = emo.to(device).unsqueeze(0) - x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device) - del phones - speakers = torch.LongTensor([hps.data.spk2id[sid]]).to(device) - audio = ( - net_g.infer( - x_tst, - x_tst_lengths, - speakers, - tones, - lang_ids, - bert, - ja_bert, - en_bert, - style_vec=style_vec, - sdp_ratio=sdp_ratio, - noise_scale=noise_scale, - noise_scale_w=noise_scale_w, - length_scale=length_scale, - )[0][0, 0] - .data.cpu() - .float() - .numpy() - ) - del ( - x_tst, - tones, - lang_ids, - bert, - x_tst_lengths, - speakers, - ja_bert, - en_bert, - ) # , emo - if torch.cuda.is_available(): - torch.cuda.empty_cache() - return audio +from typing import Literal + +import torch + +import utils +from text import cleaned_text_to_sequence, get_bert +from style_bert_vits2.logging import logger +from style_bert_vits2.models import commons +from style_bert_vits2.models.models import SynthesizerTrn +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 + + +def get_net_g(model_path: str, version: str, device: str, hps): + if version.endswith("JP-Extra"): + logger.info("Using JP-Extra model") + net_g = SynthesizerTrnJPExtra( + len(SYMBOLS), + hps.data.filter_length // 2 + 1, + hps.train.segment_size // hps.data.hop_length, + n_speakers=hps.data.n_speakers, + **hps.model, + ).to(device) + else: + logger.info("Using normal model") + net_g = SynthesizerTrn( + len(SYMBOLS), + hps.data.filter_length // 2 + 1, + hps.train.segment_size // hps.data.hop_length, + n_speakers=hps.data.n_speakers, + **hps.model, + ).to(device) + net_g.state_dict() + _ = net_g.eval() + if model_path.endswith(".pth") or model_path.endswith(".pt"): + _ = utils.load_checkpoint(model_path, net_g, None, skip_optimizer=True) + elif model_path.endswith(".safetensors"): + _ = utils.load_safetensors(model_path, net_g, True) + else: + raise ValueError(f"Unknown model format: {model_path}") + return net_g + + +def get_text( + text: str, + language_str: Literal["JP", "EN", "ZH"], + hps, + device: str, + assist_text: str | None = None, + assist_text_weight: float = 0.7, + given_tone: list[int] | None = None, +): + use_jp_extra = hps.version.endswith("JP-Extra") + # 推論時のみ呼び出されるので、raise_yomi_error は False に設定 + norm_text, phone, tone, word2ph = clean_text( + text, + language_str, + use_jp_extra = use_jp_extra, + raise_yomi_error = False, + ) + if given_tone is not None: + if len(given_tone) != len(phone): + raise InvalidToneError( + f"Length of given_tone ({len(given_tone)}) != length of phone ({len(phone)})" + ) + tone = given_tone + phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str) + + if hps.data.add_blank: + phone = commons.intersperse(phone, 0) + tone = commons.intersperse(tone, 0) + language = commons.intersperse(language, 0) + for i in range(len(word2ph)): + word2ph[i] = word2ph[i] * 2 + word2ph[0] += 1 + bert_ori = get_bert( + norm_text, + word2ph, + language_str, + device, + assist_text, + assist_text_weight, + ) + del word2ph + assert bert_ori.shape[-1] == len(phone), phone + + if language_str == "ZH": + bert = bert_ori + ja_bert = torch.zeros(1024, len(phone)) + en_bert = torch.zeros(1024, len(phone)) + elif language_str == "JP": + bert = torch.zeros(1024, len(phone)) + ja_bert = bert_ori + en_bert = torch.zeros(1024, len(phone)) + elif language_str == "EN": + bert = torch.zeros(1024, len(phone)) + ja_bert = torch.zeros(1024, len(phone)) + en_bert = bert_ori + else: + raise ValueError("language_str should be ZH, JP or EN") + + assert bert.shape[-1] == len( + phone + ), f"Bert seq len {bert.shape[-1]} != {len(phone)}" + + phone = torch.LongTensor(phone) + tone = torch.LongTensor(tone) + language = torch.LongTensor(language) + return bert, ja_bert, en_bert, phone, tone, language + + +def infer( + text: str, + style_vec, + sdp_ratio: float, + noise_scale: float, + 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"], + hps, + net_g, + device: str, + skip_start: bool = False, + skip_end: bool = False, + assist_text: str | None = None, + assist_text_weight: float = 0.7, + given_tone: list[int] | None = None, +): + is_jp_extra = hps.version.endswith("JP-Extra") + bert, ja_bert, en_bert, phones, tones, lang_ids = get_text( + text, + language, + hps, + device, + assist_text=assist_text, + assist_text_weight=assist_text_weight, + given_tone=given_tone, + ) + if skip_start: + phones = phones[3:] + tones = tones[3:] + lang_ids = lang_ids[3:] + bert = bert[:, 3:] + ja_bert = ja_bert[:, 3:] + en_bert = en_bert[:, 3:] + if skip_end: + phones = phones[:-2] + tones = tones[:-2] + lang_ids = lang_ids[:-2] + bert = bert[:, :-2] + ja_bert = ja_bert[:, :-2] + en_bert = en_bert[:, :-2] + with torch.no_grad(): + x_tst = phones.to(device).unsqueeze(0) + tones = tones.to(device).unsqueeze(0) + lang_ids = lang_ids.to(device).unsqueeze(0) + bert = bert.to(device).unsqueeze(0) + ja_bert = ja_bert.to(device).unsqueeze(0) + en_bert = en_bert.to(device).unsqueeze(0) + x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device) + style_vec = torch.from_numpy(style_vec).to(device).unsqueeze(0) + del phones + sid_tensor = torch.LongTensor([sid]).to(device) + if is_jp_extra: + output = net_g.infer( + x_tst, + x_tst_lengths, + sid_tensor, + tones, + lang_ids, + ja_bert, + style_vec=style_vec, + sdp_ratio=sdp_ratio, + noise_scale=noise_scale, + noise_scale_w=noise_scale_w, + length_scale=length_scale, + ) + else: + output = net_g.infer( + x_tst, + x_tst_lengths, + sid_tensor, + tones, + lang_ids, + bert, + ja_bert, + en_bert, + style_vec=style_vec, + sdp_ratio=sdp_ratio, + noise_scale=noise_scale, + noise_scale_w=noise_scale_w, + length_scale=length_scale, + ) + audio = output[0][0, 0].data.cpu().float().numpy() + del ( + x_tst, + tones, + lang_ids, + bert, + x_tst_lengths, + sid_tensor, + ja_bert, + en_bert, + style_vec, + ) # , emo + if torch.cuda.is_available(): + torch.cuda.empty_cache() + return audio + + +def infer_multilang( + text: str, + style_vec, + sdp_ratio: float, + noise_scale: float, + noise_scale_w: float, + length_scale: float, + sid: int, + language: Literal["JP", "EN", "ZH"], + hps, + net_g, + device: str, + skip_start: bool = False, + skip_end: bool = False, +): + bert, ja_bert, en_bert, phones, tones, lang_ids = [], [], [], [], [], [] + # emo = get_emo_(reference_audio, emotion, sid) + # if isinstance(reference_audio, np.ndarray): + # emo = get_clap_audio_feature(reference_audio, device) + # else: + # emo = get_clap_text_feature(emotion, device) + # emo = torch.squeeze(emo, dim=1) + for idx, (txt, lang) in enumerate(zip(text, language)): + _skip_start = (idx != 0) or (skip_start and idx == 0) + _skip_end = (idx != len(language) - 1) or skip_end + ( + temp_bert, + temp_ja_bert, + temp_en_bert, + temp_phones, + temp_tones, + temp_lang_ids, + ) = get_text(txt, lang, hps, device) # type: ignore + if _skip_start: + temp_bert = temp_bert[:, 3:] + temp_ja_bert = temp_ja_bert[:, 3:] + temp_en_bert = temp_en_bert[:, 3:] + temp_phones = temp_phones[3:] + temp_tones = temp_tones[3:] + temp_lang_ids = temp_lang_ids[3:] + if _skip_end: + temp_bert = temp_bert[:, :-2] + temp_ja_bert = temp_ja_bert[:, :-2] + temp_en_bert = temp_en_bert[:, :-2] + temp_phones = temp_phones[:-2] + temp_tones = temp_tones[:-2] + temp_lang_ids = temp_lang_ids[:-2] + bert.append(temp_bert) + ja_bert.append(temp_ja_bert) + en_bert.append(temp_en_bert) + phones.append(temp_phones) + tones.append(temp_tones) + lang_ids.append(temp_lang_ids) + bert = torch.concatenate(bert, dim=1) + ja_bert = torch.concatenate(ja_bert, dim=1) + en_bert = torch.concatenate(en_bert, dim=1) + phones = torch.concatenate(phones, dim=0) + tones = torch.concatenate(tones, dim=0) + lang_ids = torch.concatenate(lang_ids, dim=0) + with torch.no_grad(): + x_tst = phones.to(device).unsqueeze(0) + tones = tones.to(device).unsqueeze(0) + lang_ids = lang_ids.to(device).unsqueeze(0) + bert = bert.to(device).unsqueeze(0) + ja_bert = ja_bert.to(device).unsqueeze(0) + en_bert = en_bert.to(device).unsqueeze(0) + # emo = emo.to(device).unsqueeze(0) + x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device) + del phones + speakers = torch.LongTensor([hps.data.spk2id[sid]]).to(device) + audio = ( + net_g.infer( + x_tst, + x_tst_lengths, + speakers, + tones, + lang_ids, + bert, + ja_bert, + en_bert, + style_vec=style_vec, + sdp_ratio=sdp_ratio, + noise_scale=noise_scale, + noise_scale_w=noise_scale_w, + length_scale=length_scale, + )[0][0, 0] + .data.cpu() + .float() + .numpy() + ) + del ( + x_tst, + tones, + lang_ids, + bert, + x_tst_lengths, + speakers, + ja_bert, + en_bert, + ) # , emo + if torch.cuda.is_available(): + torch.cuda.empty_cache() + return audio diff --git a/style_bert_vits2/text_processing/cleaner.py b/style_bert_vits2/text_processing/cleaner.py new file mode 100644 index 0000000..400d198 --- /dev/null +++ b/style_bert_vits2/text_processing/cleaner.py @@ -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 diff --git a/text/chinese.py b/text/chinese.py index 56dc4f3..3d9c392 100644 --- a/text/chinese.py +++ b/text/chinese.py @@ -168,7 +168,7 @@ def _g2p(segments): return phones_list, tones_list, word2ph -def text_normalize(text): +def normalize_text(text): numbers = re.findall(r"\d+(?:\.?\d+)?", text) for number in numbers: text = text.replace(number, cn2an.an2cn(number), 1) @@ -186,7 +186,7 @@ if __name__ == "__main__": from text.chinese_bert import get_bert_feature text = "啊!但是《原神》是由,米哈\游自主, [研发]的一款全.新开放世界.冒险游戏" - text = text_normalize(text) + text = normalize_text(text) print(text) phones, tones, word2ph = g2p(text) bert = get_bert_feature(text, word2ph) diff --git a/text/cleaner.py b/text/cleaner.py deleted file mode 100644 index d805b51..0000000 --- a/text/cleaner.py +++ /dev/null @@ -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 diff --git a/text/english.py b/text/english.py index f38ee84..3dcfdec 100644 --- a/text/english.py +++ b/text/english.py @@ -369,7 +369,7 @@ def normalize_numbers(text): return text -def text_normalize(text): +def normalize_text(text): text = normalize_numbers(text) text = replace_punctuation(text) text = re.sub(r"([,;.\?\!])([\w])", r"\1 \2", text) diff --git a/text/japanese.py b/text/japanese.py index 0dc6aa8..03682e3 100644 --- a/text/japanese.py +++ b/text/japanese.py @@ -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 ) -> tuple[list[str], list[int], list[int]]: """ - 他で使われるメインの関数。`text_normalize()`で正規化された`norm_text`を受け取り、 + 他で使われるメインの関数。`normalize_text()`で正規化された`norm_text`を受け取り、 - phones: 音素のリスト(ただし`!`や`,`や`.`等punctuationが含まれうる) - tones: アクセントのリスト、0(低)と1(高)からなり、phonesと同じ長さ - word2ph: 元のテキストの各文字に音素が何個割り当てられるかを表すリスト @@ -350,7 +350,7 @@ def text2sep_kata( norm_text: str, raise_yomi_error: bool = False ) -> tuple[list[str], list[str]]: """ - `text_normalize`で正規化済みの`norm_text`を受け取り、それを単語分割し、 + `normalize_text()`で正規化済みの`norm_text`を受け取り、それを単語分割し、 分割された単語リストとその読み(カタカナor記号1文字)のリストのタプルを返す。 単語分割結果は、`g2p()`の`word2ph`で1文字あたりに割り振る音素記号の数を決めるために使う。 例: @@ -634,7 +634,7 @@ if __name__ == "__main__": text = "こんにちは、世界。" from text.japanese_bert import get_bert_feature - text = text_normalize(text) + text = normalize_text(text) phones, tones, word2ph = g2p(text) bert = get_bert_feature(text, word2ph) From c3c0dd8b32db40b385995a1389402d3683a21762 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Thu, 7 Mar 2024 02:31:30 +0000 Subject: [PATCH 12/64] 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 --- server_editor.py | 10 +- style_bert_vits2/constants.py | 34 +- style_bert_vits2/models/infer.py | 15 +- .../text_processing/bert_models.py | 123 ++++ style_bert_vits2/text_processing/cleaner.py | 20 +- .../text_processing/japanese/g2p.py | 8 +- .../text_processing/japanese/g2p_utils.py | 10 +- text/__init__.py | 24 +- text/chinese_bert.py | 17 +- text/english.py | 7 +- text/english_bert_mock.py | 18 +- text/japanese.py | 642 ------------------ text/japanese_bert.py | 17 +- 13 files changed, 217 insertions(+), 728 deletions(-) create mode 100644 style_bert_vits2/text_processing/bert_models.py delete mode 100644 text/japanese.py diff --git a/server_editor.py b/server_editor.py index 7577637..eb1cffd 100644 --- a/server_editor.py +++ b/server_editor.py @@ -28,7 +28,6 @@ from fastapi.responses import JSONResponse, Response from fastapi.staticfiles import StaticFiles from pydantic import BaseModel from scipy.io import wavfile -from transformers import AutoTokenizer from common.tts_model import ModelHolder from style_bert_vits2.constants import ( @@ -42,6 +41,7 @@ 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.g2p_utils import g2kata_tone, kata_tone2phone_tone from style_bert_vits2.text_processing.japanese.normalizer import normalize_text from style_bert_vits2.text_processing.japanese.user_dict import ( @@ -150,8 +150,10 @@ def save_last_download(latest_release): # 最初に pyopenjtalk の辞書を更新 update_dict() -# 単語分割に使う BERT トークナイザーをロード -tokenizer = AutoTokenizer.from_pretrained("./bert/deberta-v2-large-japanese-char-wwm") +# 単語分割に使う BERT モデル/トークナイザーを事前にロードしておく +## server_editor.py は日本語にしか対応していないため、日本語の BERT モデル/トークナイザーのみロードする +bert_models.load_model(Languages.JP) +bert_models.load_tokenizer(Languages.JP) class AudioResponse(Response): @@ -227,7 +229,7 @@ async def read_item(item: TextRequest): try: # 最初に正規化しないと整合性がとれない text = normalize_text(item.text) - kata_tone_list = g2kata_tone(text, tokenizer) + kata_tone_list = g2kata_tone(text) except Exception as e: raise HTTPException( status_code=400, diff --git a/style_bert_vits2/constants.py b/style_bert_vits2/constants.py index 2d6e9ad..d6b54c8 100644 --- a/style_bert_vits2/constants.py +++ b/style_bert_vits2/constants.py @@ -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" diff --git a/style_bert_vits2/models/infer.py b/style_bert_vits2/models/infer.py index e0d8691..99999a9 100644 --- a/style_bert_vits2/models/infer.py +++ b/style_bert_vits2/models/infer.py @@ -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, diff --git a/style_bert_vits2/text_processing/bert_models.py b/style_bert_vits2/text_processing/bert_models.py new file mode 100644 index 0000000..9132082 --- /dev/null +++ b/style_bert_vits2/text_processing/bert_models.py @@ -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 diff --git a/style_bert_vits2/text_processing/cleaner.py b/style_bert_vits2/text_processing/cleaner.py index 400d198..0420099 100644 --- a/style_bert_vits2/text_processing/cleaner.py +++ b/style_bert_vits2/text_processing/cleaner.py @@ -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) diff --git a/style_bert_vits2/text_processing/japanese/g2p.py b/style_bert_vits2/text_processing/japanese/g2p.py index 1b8bd69..8ef4d3e 100644 --- a/style_bert_vits2/text_processing/japanese/g2p.py +++ b/style_bert_vits2/text_processing/japanese/g2p.py @@ -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]) diff --git a/style_bert_vits2/text_processing/japanese/g2p_utils.py b/style_bert_vits2/text_processing/japanese/g2p_utils.py index e095602..3a91a00 100644 --- a/style_bert_vits2/text_processing/japanese/g2p_utils.py +++ b/style_bert_vits2/text_processing/japanese/g2p_utils.py @@ -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))) diff --git a/text/__init__.py b/text/__init__.py index ce4c008..efff830 100644 --- a/text/__init__.py +++ b/text/__init__.py @@ -1,12 +1,17 @@ +from style_bert_vits2.constants import Languages from style_bert_vits2.text_processing.symbols import * + _symbol_to_id = {s: i for i, s in enumerate(SYMBOLS)} -def cleaned_text_to_sequence(cleaned_text, tones, language): - """Converts a string of text to a sequence of IDs corresponding to the symbols in the text. +def cleaned_text_to_sequence(cleaned_text: str, tones: list[int], language: Languages): + """ + Converts a string of text to a sequence of IDs corresponding to the symbols in the text. + Args: text: string to convert to a sequence + Returns: List of integers corresponding to the symbols in the text """ @@ -18,12 +23,19 @@ def cleaned_text_to_sequence(cleaned_text, tones, language): return phones, tones, lang_ids -def get_bert(text, word2ph, language, device, assist_text=None, assist_text_weight=0.7): - if language == "ZH": +def get_bert( + text: str, + word2ph, + language: Languages, + device: str, + assist_text: str | None = None, + assist_text_weight: float = 0.7, +): + if language == Languages.ZH: from .chinese_bert import get_bert_feature - elif language == "EN": + elif language == Languages.EN: from .english_bert_mock import get_bert_feature - elif language == "JP": + elif language == Languages.JP: from .japanese_bert import get_bert_feature else: raise ValueError(f"Language {language} not supported") diff --git a/text/chinese_bert.py b/text/chinese_bert.py index 94e1408..2ee6e9d 100644 --- a/text/chinese_bert.py +++ b/text/chinese_bert.py @@ -1,23 +1,21 @@ import sys import torch -from transformers import AutoModelForMaskedLM, AutoTokenizer from config import config +from style_bert_vits2.constants import Languages +from style_bert_vits2.text_processing import bert_models -LOCAL_PATH = "./bert/chinese-roberta-wwm-ext-large" - -tokenizer = AutoTokenizer.from_pretrained(LOCAL_PATH) models = dict() def get_bert_feature( - text, + text: str, word2ph, - device=config.bert_gen_config.device, - assist_text=None, - assist_text_weight=0.7, + device = config.bert_gen_config.device, + assist_text: str | None = None, + assist_text_weight: float = 0.7, ): if ( sys.platform == "darwin" @@ -30,8 +28,9 @@ def get_bert_feature( if device == "cuda" and not torch.cuda.is_available(): device = "cpu" if device not in models.keys(): - models[device] = AutoModelForMaskedLM.from_pretrained(LOCAL_PATH).to(device) + models[device] = bert_models.load_model(Languages.ZH).to(device) with torch.no_grad(): + tokenizer = bert_models.load_tokenizer(Languages.ZH) inputs = tokenizer(text, return_tensors="pt") for i in inputs: inputs[i] = inputs[i].to(device) diff --git a/text/english.py b/text/english.py index 3dcfdec..a6d71e9 100644 --- a/text/english.py +++ b/text/english.py @@ -2,16 +2,16 @@ import pickle import os import re from g2p_en import G2p -from transformers import DebertaV2Tokenizer +from style_bert_vits2.constants import Languages +from style_bert_vits2.text_processing import bert_models from style_bert_vits2.text_processing.symbols import PUNCTUATIONS, SYMBOLS + current_file_path = os.path.dirname(__file__) CMU_DICT_PATH = os.path.join(current_file_path, "cmudict.rep") CACHE_PATH = os.path.join(current_file_path, "cmudict_cache.pickle") _g2p = G2p() -LOCAL_PATH = "./bert/deberta-v3-large" -tokenizer = DebertaV2Tokenizer.from_pretrained(LOCAL_PATH) arpa = { "AH0", @@ -392,6 +392,7 @@ def sep_text(text): def text_to_words(text): + tokenizer = bert_models.load_tokenizer(Languages.EN) tokens = tokenizer.tokenize(text) words = [] for idx, t in enumerate(tokens): diff --git a/text/english_bert_mock.py b/text/english_bert_mock.py index 782b65d..8e57df1 100644 --- a/text/english_bert_mock.py +++ b/text/english_bert_mock.py @@ -1,24 +1,21 @@ import sys import torch -from transformers import DebertaV2Model, DebertaV2Tokenizer from config import config +from style_bert_vits2.constants import Languages +from style_bert_vits2.text_processing import bert_models -LOCAL_PATH = "./bert/deberta-v3-large" - -tokenizer = DebertaV2Tokenizer.from_pretrained(LOCAL_PATH) - models = dict() def get_bert_feature( - text, + text: str, word2ph, - device=config.bert_gen_config.device, - assist_text=None, - assist_text_weight=0.7, + device = config.bert_gen_config.device, + assist_text: str | None = None, + assist_text_weight: float = 0.7, ): if ( sys.platform == "darwin" @@ -31,8 +28,9 @@ def get_bert_feature( if device == "cuda" and not torch.cuda.is_available(): device = "cpu" if device not in models.keys(): - models[device] = DebertaV2Model.from_pretrained(LOCAL_PATH).to(device) + models[device] = bert_models.load_model(Languages.EN).to(device) with torch.no_grad(): + tokenizer = bert_models.load_tokenizer(Languages.EN) inputs = tokenizer(text, return_tensors="pt") for i in inputs: inputs[i] = inputs[i].to(device) diff --git a/text/japanese.py b/text/japanese.py deleted file mode 100644 index 03682e3..0000000 --- a/text/japanese.py +++ /dev/null @@ -1,642 +0,0 @@ -# Convert Japanese text to phonemes which is -# compatible with Julius https://github.com/julius-speech/segmentation-kit -import re -import unicodedata - -import pyopenjtalk -from num2words import num2words -from transformers import AutoTokenizer - -from style_bert_vits2.logging import logger -from style_bert_vits2.text_processing.japanese.mora_list import ( - MORA_KATA_TO_MORA_PHONEMES, - MORA_PHONEMES_TO_MORA_KATA, -) -from style_bert_vits2.text_processing.japanese.user_dict import update_dict -from style_bert_vits2.text_processing.symbols import PUNCTUATIONS - -# 最初にpyopenjtalkの辞書を更新 -update_dict() - -# 子音の集合 -COSONANTS = set( - [ - cosonant - for cosonant, _ in MORA_KATA_TO_MORA_PHONEMES.values() - if cosonant is not None - ] -) - -# 母音の集合、便宜上「ん」を含める -VOWELS = {"a", "i", "u", "e", "o", "N"} - - -class YomiError(Exception): - """ - OpenJTalkで、読みが正しく取得できない箇所があるときに発生する例外。 - 基本的に「学習の前処理のテキスト処理時」には発生させ、そうでない場合は、 - ignore_yomi_error=Trueにしておいて、この例外を発生させないようにする。 - """ - - pass - - -# 正規化で記号を変換するための辞書 -rep_map = { - ":": ",", - ";": ",", - ",": ",", - "。": ".", - "!": "!", - "?": "?", - "\n": ".", - ".": ".", - "…": "...", - "···": "...", - "・・・": "...", - "·": ",", - "・": ",", - "、": ",", - "$": ".", - "“": "'", - "”": "'", - '"': "'", - "‘": "'", - "’": "'", - "(": "'", - ")": "'", - "(": "'", - ")": "'", - "《": "'", - "》": "'", - "【": "'", - "】": "'", - "[": "'", - "]": "'", - # NFKC正規化後のハイフン・ダッシュの変種を全て通常半角ハイフン - \u002d に変換 - "\u02d7": "\u002d", # ˗, Modifier Letter Minus Sign - "\u2010": "\u002d", # ‐, Hyphen, - # "\u2011": "\u002d", # ‑, Non-Breaking Hyphen, NFKCにより\u2010に変換される - "\u2012": "\u002d", # ‒, Figure Dash - "\u2013": "\u002d", # –, En Dash - "\u2014": "\u002d", # —, Em Dash - "\u2015": "\u002d", # ―, Horizontal Bar - "\u2043": "\u002d", # ⁃, Hyphen Bullet - "\u2212": "\u002d", # −, Minus Sign - "\u23af": "\u002d", # ⎯, Horizontal Line Extension - "\u23e4": "\u002d", # ⏤, Straightness - "\u2500": "\u002d", # ─, Box Drawings Light Horizontal - "\u2501": "\u002d", # ━, Box Drawings Heavy Horizontal - "\u2e3a": "\u002d", # ⸺, Two-Em Dash - "\u2e3b": "\u002d", # ⸻, Three-Em Dash - # "~": "-", # これは長音記号「ー」として扱うよう変更 - # "~": "-", # これも長音記号「ー」として扱うよう変更 - "「": "'", - "」": "'", -} - - -def normalize_text(text): - """ - 日本語のテキストを正規化する。 - 結果は、ちょうど次の文字のみからなる: - - ひらがな - - カタカナ(全角長音記号「ー」が入る!) - - 漢字 - - 半角アルファベット(大文字と小文字) - - ギリシャ文字 - - `.` (句点`。`や`…`の一部や改行等) - - `,` (読点`、`や`:`等) - - `?` (疑問符`?`) - - `!` (感嘆符`!`) - - `'` (`「`や`」`等) - - `-` (`―`(ダッシュ、長音記号ではない)や`-`等) - - 注意点: - - 三点リーダー`…`は`...`に変換される(`なるほど…。` → `なるほど....`) - - 数字は漢字に変換される(`1,100円` → `千百円`、`52.34` → `五十二点三四`) - - 読点や疑問符等の位置・個数等は保持される(`??あ、、!!!` → `??あ,,!!!`) - """ - res = unicodedata.normalize("NFKC", text) # ここでアルファベットは半角になる - res = japanese_convert_numbers_to_words(res) # 「100円」→「百円」等 - # 「~」と「~」も長音記号として扱う - res = res.replace("~", "ー") - res = res.replace("~", "ー") - - res = replace_punctuation(res) # 句読点等正規化、読めない文字を削除 - - # 結合文字の濁点・半濁点を削除 - # 通常の「ば」等はそのままのこされる、「あ゛」は上で「あ゙」になりここで「あ」になる - res = res.replace("\u3099", "") # 結合文字の濁点を削除、る゙ → る - res = res.replace("\u309A", "") # 結合文字の半濁点を削除、な゚ → な - return res - - -def replace_punctuation(text: str) -> str: - """句読点等を「.」「,」「!」「?」「'」「-」に正規化し、OpenJTalkで読みが取得できるもののみ残す: - 漢字・平仮名・カタカナ、アルファベット、ギリシャ文字 - """ - pattern = re.compile("|".join(re.escape(p) for p in rep_map.keys())) - - # 句読点を辞書で置換 - replaced_text = pattern.sub(lambda x: rep_map[x.group()], text) - - replaced_text = re.sub( - # ↓ ひらがな、カタカナ、漢字 - r"[^\u3040-\u309F\u30A0-\u30FF\u4E00-\u9FFF\u3400-\u4DBF\u3005" - # ↓ 半角アルファベット(大文字と小文字) - + r"\u0041-\u005A\u0061-\u007A" - # ↓ 全角アルファベット(大文字と小文字) - + r"\uFF21-\uFF3A\uFF41-\uFF5A" - # ↓ ギリシャ文字 - + r"\u0370-\u03FF\u1F00-\u1FFF" - # ↓ "!", "?", "…", ",", ".", "'", "-", 但し`…`はすでに`...`に変換されている - + "".join(PUNCTUATIONS) + r"]+", - # 上述以外の文字を削除 - "", - replaced_text, - ) - - return replaced_text - - -_NUMBER_WITH_SEPARATOR_RX = re.compile("[0-9]{1,3}(,[0-9]{3})+") -_CURRENCY_MAP = {"$": "ドル", "¥": "円", "£": "ポンド", "€": "ユーロ"} -_CURRENCY_RX = re.compile(r"([$¥£€])([0-9.]*[0-9])") -_NUMBER_RX = re.compile(r"[0-9]+(\.[0-9]+)?") - - -def japanese_convert_numbers_to_words(text: str) -> str: - res = _NUMBER_WITH_SEPARATOR_RX.sub(lambda m: m[0].replace(",", ""), text) - res = _CURRENCY_RX.sub(lambda m: m[2] + _CURRENCY_MAP.get(m[1], m[1]), res) - res = _NUMBER_RX.sub(lambda m: num2words(m[0], lang="ja"), res) - return res - - -def g2p( - norm_text: str, use_jp_extra: bool = True, raise_yomi_error: bool = False -) -> tuple[list[str], list[int], list[int]]: - """ - 他で使われるメインの関数。`normalize_text()`で正規化された`norm_text`を受け取り、 - - phones: 音素のリスト(ただし`!`や`,`や`.`等punctuationが含まれうる) - - tones: アクセントのリスト、0(低)と1(高)からなり、phonesと同じ長さ - - word2ph: 元のテキストの各文字に音素が何個割り当てられるかを表すリスト - のタプルを返す。 - ただし`phones`と`tones`の最初と終わりに`_`が入り、応じて`word2ph`の最初と最後に1が追加される。 - - use_jp_extra: Falseの場合、「ん」の音素を「N」ではなく「n」とする。 - raise_yomi_error: Trueの場合、読めない文字があるときに例外を発生させる。 - Falseの場合は読めない文字が消えたような扱いとして処理される。 - """ - # pyopenjtalkのフルコンテキストラベルを使ってアクセントを取り出すと、punctuationの位置が消えてしまい情報が失われてしまう: - # 「こんにちは、世界。」と「こんにちは!世界。」と「こんにちは!!!???世界……。」は全て同じになる。 - # よって、まずpunctuation無しの音素とアクセントのリストを作り、 - # それとは別にpyopenjtalk.run_frontend()で得られる音素リスト(こちらはpunctuationが保持される)を使い、 - # アクセント割当をしなおすことによってpunctuationを含めた音素とアクセントのリストを作る。 - - # punctuationがすべて消えた、音素とアクセントのタプルのリスト(「ん」は「N」) - phone_tone_list_wo_punct = g2phone_tone_wo_punct(norm_text) - - # sep_text: 単語単位の単語のリスト、読めない文字があったらraise_yomi_errorなら例外、そうでないなら読めない文字が消えて返ってくる - # sep_kata: 単語単位の単語のカタカナ読みのリスト - sep_text, sep_kata = text2sep_kata(norm_text, raise_yomi_error=raise_yomi_error) - - # sep_phonemes: 各単語ごとの音素のリストのリスト - sep_phonemes = handle_long([kata2phoneme_list(i) for i in sep_kata]) - - # phone_w_punct: sep_phonemesを結合した、punctuationを元のまま保持した音素列 - phone_w_punct: list[str] = [] - for i in sep_phonemes: - phone_w_punct += i - - # punctuation無しのアクセント情報を使って、punctuationを含めたアクセント情報を作る - phone_tone_list = align_tones(phone_w_punct, phone_tone_list_wo_punct) - # logger.debug(f"phone_tone_list:\n{phone_tone_list}") - # word2phは厳密な解答は不可能なので(「今日」「眼鏡」等の熟字訓が存在)、 - # Bert-VITS2では、単語単位の分割を使って、単語の文字ごとにだいたい均等に音素を分配する - - # sep_textから、各単語を1文字1文字分割して、文字のリスト(のリスト)を作る - sep_tokenized: list[list[str]] = [] - for i in sep_text: - if i not in PUNCTUATIONS: - sep_tokenized.append( - tokenizer.tokenize(i) - ) # ここでおそらく`i`が文字単位に分割される - else: - sep_tokenized.append([i]) - - # 各単語について、音素の数と文字の数を比較して、均等っぽく分配する - word2ph = [] - for token, phoneme in zip(sep_tokenized, sep_phonemes): - phone_len = len(phoneme) - word_len = len(token) - word2ph += distribute_phone(phone_len, word_len) - - # 最初と最後に`_`記号を追加、アクセントは0(低)、word2phもそれに合わせて追加 - phone_tone_list = [("_", 0)] + phone_tone_list + [("_", 0)] - word2ph = [1] + word2ph + [1] - - phones = [phone for phone, _ in phone_tone_list] - tones = [tone for _, tone in phone_tone_list] - - assert len(phones) == sum(word2ph), f"{len(phones)} != {sum(word2ph)}" - - # use_jp_extraでない場合は「N」を「n」に変換 - if not use_jp_extra: - phones = [phone if phone != "N" else "n" for phone in phones] - - return phones, tones, word2ph - - -def g2kata_tone(norm_text: str) -> list[tuple[str, int]]: - """ - テキストからカタカナとアクセントのペアのリストを返す。 - 推論時のみに使われるので、常に`raise_yomi_error=False`でg2pを呼ぶ。 - """ - phones, tones, _ = g2p(norm_text, use_jp_extra=True, raise_yomi_error=False) - return phone_tone2kata_tone(list(zip(phones, tones))) - - -def phone_tone2kata_tone(phone_tone: list[tuple[str, int]]) -> list[tuple[str, int]]: - """phone_toneをのphone部分をカタカナに変換する。ただし最初と最後の("_", 0)は無視""" - phone_tone = phone_tone[1:] # 最初の("_", 0)を無視 - phones = [phone for phone, _ in phone_tone] - tones = [tone for _, tone in phone_tone] - result: list[tuple[str, int]] = [] - current_mora = "" - for phone, next_phone, tone, next_tone in zip(phones, phones[1:], tones, tones[1:]): - # zipの関係で最後の("_", 0)は無視されている - if phone in PUNCTUATIONS: - result.append((phone, tone)) - continue - if phone in COSONANTS: # n以外の子音の場合 - assert current_mora == "", f"Unexpected {phone} after {current_mora}" - assert tone == next_tone, f"Unexpected {phone} tone {tone} != {next_tone}" - current_mora = phone - else: - # phoneが母音もしくは「N」 - current_mora += phone - result.append((MORA_PHONEMES_TO_MORA_KATA[current_mora], tone)) - current_mora = "" - return result - - -def kata_tone2phone_tone(kata_tone: list[tuple[str, int]]) -> list[tuple[str, int]]: - """`phone_tone2kata_tone()`の逆。""" - result: list[tuple[str, int]] = [("_", 0)] - for mora, tone in kata_tone: - if mora in PUNCTUATIONS: - result.append((mora, tone)) - else: - cosonant, vowel = MORA_KATA_TO_MORA_PHONEMES[mora] - if cosonant is None: - result.append((vowel, tone)) - else: - result.append((cosonant, tone)) - result.append((vowel, tone)) - result.append(("_", 0)) - return result - - -def g2phone_tone_wo_punct(text: str) -> list[tuple[str, int]]: - """ - テキストに対して、音素とアクセント(0か1)のペアのリストを返す。 - ただし「!」「.」「?」等の非音素記号(punctuation)は全て消える(ポーズ記号も残さない)。 - 非音素記号を含める処理は`align_tones()`で行われる。 - また「っ」は「q」に、「ん」は「N」に変換される。 - 例: "こんにちは、世界ー。。元気?!" → - [('k', 0), ('o', 0), ('N', 1), ('n', 1), ('i', 1), ('ch', 1), ('i', 1), ('w', 1), ('a', 1), ('s', 1), ('e', 1), ('k', 0), ('a', 0), ('i', 0), ('i', 0), ('g', 1), ('e', 1), ('N', 0), ('k', 0), ('i', 0)] - """ - prosodies = pyopenjtalk_g2p_prosody(text, drop_unvoiced_vowels=True) - # logger.debug(f"prosodies: {prosodies}") - result: list[tuple[str, int]] = [] - current_phrase: list[tuple[str, int]] = [] - current_tone = 0 - for i, letter in enumerate(prosodies): - # 特殊記号の処理 - - # 文頭記号、無視する - if letter == "^": - assert i == 0, "Unexpected ^" - # アクセント句の終わりに来る記号 - elif letter in ("$", "?", "_", "#"): - # 保持しているフレーズを、アクセント数値を0-1に修正し結果に追加 - result.extend(fix_phone_tone(current_phrase)) - # 末尾に来る終了記号、無視(文中の疑問文は`_`になる) - if letter in ("$", "?"): - assert i == len(prosodies) - 1, f"Unexpected {letter}" - # あとは"_"(ポーズ)と"#"(アクセント句の境界)のみ - # これらは残さず、次のアクセント句に備える。 - current_phrase = [] - # 0を基準点にしてそこから上昇・下降する(負の場合は上の`fix_phone_tone`で直る) - current_tone = 0 - # アクセント上昇記号 - elif letter == "[": - current_tone = current_tone + 1 - # アクセント下降記号 - elif letter == "]": - current_tone = current_tone - 1 - # それ以外は通常の音素 - else: - if letter == "cl": # 「っ」の処理 - letter = "q" - # elif letter == "N": # 「ん」の処理 - # letter = "n" - current_phrase.append((letter, current_tone)) - return result - - -def text2sep_kata( - norm_text: str, raise_yomi_error: bool = False -) -> tuple[list[str], list[str]]: - """ - `normalize_text()`で正規化済みの`norm_text`を受け取り、それを単語分割し、 - 分割された単語リストとその読み(カタカナor記号1文字)のリストのタプルを返す。 - 単語分割結果は、`g2p()`の`word2ph`で1文字あたりに割り振る音素記号の数を決めるために使う。 - 例: - `私はそう思う!って感じ?` → - ["私", "は", "そう", "思う", "!", "って", "感じ", "?"], ["ワタシ", "ワ", "ソー", "オモウ", "!", "ッテ", "カンジ", "?"] - - raise_yomi_error: Trueの場合、読めない文字があるときに例外を発生させる。 - Falseの場合は読めない文字が消えたような扱いとして処理される。 - """ - # parsed: OpenJTalkの解析結果 - parsed = pyopenjtalk.run_frontend(norm_text) - sep_text: list[str] = [] - sep_kata: list[str] = [] - for parts in parsed: - # word: 実際の単語の文字列 - # yomi: その読み、但し無声化サインの`’`は除去 - word, yomi = replace_punctuation(parts["string"]), parts["pron"].replace( - "’", "" - ) - """ - ここで`yomi`の取りうる値は以下の通りのはず。 - - `word`が通常単語 → 通常の読み(カタカナ) - (カタカナからなり、長音記号も含みうる、`アー` 等) - - `word`が`ー` から始まる → `ーラー` や `ーーー` など - - `word`が句読点や空白等 → `、` - - `word`がpunctuationの繰り返し → 全角にしたもの - 基本的にpunctuationは1文字ずつ分かれるが、何故かある程度連続すると1つにまとまる。 - 他にも`word`が読めないキリル文字アラビア文字等が来ると`、`になるが、正規化でこの場合は起きないはず。 - また元のコードでは`yomi`が空白の場合の処理があったが、これは起きないはず。 - 処理すべきは`yomi`が`、`の場合のみのはず。 - """ - assert yomi != "", f"Empty yomi: {word}" - if yomi == "、": - # wordは正規化されているので、`.`, `,`, `!`, `'`, `-`, `--` のいずれか - if not set(word).issubset(set(PUNCTUATIONS)): # 記号繰り返しか判定 - # ここはpyopenjtalkが読めない文字等のときに起こる - if raise_yomi_error: - raise YomiError(f"Cannot read: {word} in:\n{norm_text}") - logger.warning(f"Ignoring unknown: {word} in:\n{norm_text}") - continue - # yomiは元の記号のままに変更 - yomi = word - elif yomi == "?": - assert word == "?", f"yomi `?` comes from: {word}" - yomi = "?" - sep_text.append(word) - sep_kata.append(yomi) - return sep_text, sep_kata - - -# ESPnetの実装から引用、変更点無し。「ん」は「N」なことに注意。 -# https://github.com/espnet/espnet/blob/master/espnet2/text/phoneme_tokenizer.py -def pyopenjtalk_g2p_prosody(text: str, drop_unvoiced_vowels: bool = True) -> list[str]: - """Extract phoneme + prosoody symbol sequence from input full-context labels. - - The algorithm is based on `Prosodic features control by symbols as input of - sequence-to-sequence acoustic modeling for neural TTS`_ with some r9y9's tweaks. - - Args: - text (str): Input text. - drop_unvoiced_vowels (bool): whether to drop unvoiced vowels. - - Returns: - List[str]: List of phoneme + prosody symbols. - - Examples: - >>> from espnet2.text.phoneme_tokenizer import pyopenjtalk_g2p_prosody - >>> pyopenjtalk_g2p_prosody("こんにちは。") - ['^', 'k', 'o', '[', 'N', 'n', 'i', 'ch', 'i', 'w', 'a', '$'] - - .. _`Prosodic features control by symbols as input of sequence-to-sequence acoustic - modeling for neural TTS`: https://doi.org/10.1587/transinf.2020EDP7104 - - """ - labels = pyopenjtalk.make_label(pyopenjtalk.run_frontend(text)) - N = len(labels) - - phones = [] - for n in range(N): - lab_curr = labels[n] - - # current phoneme - p3 = re.search(r"\-(.*?)\+", lab_curr).group(1) - # deal unvoiced vowels as normal vowels - if drop_unvoiced_vowels and p3 in "AEIOU": - p3 = p3.lower() - - # deal with sil at the beginning and the end of text - if p3 == "sil": - assert n == 0 or n == N - 1 - if n == 0: - phones.append("^") - elif n == N - 1: - # check question form or not - e3 = _numeric_feature_by_regex(r"!(\d+)_", lab_curr) - if e3 == 0: - phones.append("$") - elif e3 == 1: - phones.append("?") - continue - elif p3 == "pau": - phones.append("_") - continue - else: - phones.append(p3) - - # accent type and position info (forward or backward) - a1 = _numeric_feature_by_regex(r"/A:([0-9\-]+)\+", lab_curr) - a2 = _numeric_feature_by_regex(r"\+(\d+)\+", lab_curr) - a3 = _numeric_feature_by_regex(r"\+(\d+)/", lab_curr) - - # number of mora in accent phrase - f1 = _numeric_feature_by_regex(r"/F:(\d+)_", lab_curr) - - a2_next = _numeric_feature_by_regex(r"\+(\d+)\+", labels[n + 1]) - # accent phrase border - if a3 == 1 and a2_next == 1 and p3 in "aeiouAEIOUNcl": - phones.append("#") - # pitch falling - elif a1 == 0 and a2_next == a2 + 1 and a2 != f1: - phones.append("]") - # pitch rising - elif a2 == 1 and a2_next == 2: - phones.append("[") - - return phones - - -def _numeric_feature_by_regex(regex, s): - match = re.search(regex, s) - if match is None: - return -50 - return int(match.group(1)) - - -def fix_phone_tone(phone_tone_list: list[tuple[str, int]]) -> list[tuple[str, int]]: - """ - `phone_tone_list`のtone(アクセントの値)を0か1の範囲に修正する。 - 例: [(a, 0), (i, -1), (u, -1)] → [(a, 1), (i, 0), (u, 0)] - """ - tone_values = set(tone for _, tone in phone_tone_list) - if len(tone_values) == 1: - assert tone_values == {0}, tone_values - return phone_tone_list - elif len(tone_values) == 2: - if tone_values == {0, 1}: - return phone_tone_list - elif tone_values == {-1, 0}: - return [ - (letter, 0 if tone == -1 else 1) for letter, tone in phone_tone_list - ] - else: - raise ValueError(f"Unexpected tone values: {tone_values}") - else: - raise ValueError(f"Unexpected tone values: {tone_values}") - - -def distribute_phone(n_phone: int, n_word: int) -> list[int]: - """ - 左から右に1ずつ振り分け、次にまた左から右に1ずつ増やし、というふうに、 - 音素の数`n_phone`を単語の数`n_word`に分配する。 - """ - phones_per_word = [0] * n_word - for _ in range(n_phone): - min_tasks = min(phones_per_word) - min_index = phones_per_word.index(min_tasks) - phones_per_word[min_index] += 1 - return phones_per_word - - -def handle_long(sep_phonemes: list[list[str]]) -> list[list[str]]: - """ - フレーズごとに分かれた音素(長音記号がそのまま)のリストのリスト`sep_phonemes`を受け取り、 - その長音記号を処理して、音素のリストのリストを返す。 - 基本的には直前の音素を伸ばすが、直前の音素が母音でない場合もしくは冒頭の場合は、 - おそらく長音記号とダッシュを勘違いしていると思われるので、ダッシュに対応する音素`-`に変換する。 - """ - for i in range(len(sep_phonemes)): - if len(sep_phonemes[i]) == 0: - # 空白文字等でリストが空の場合 - continue - if sep_phonemes[i][0] == "ー": - if i != 0: - prev_phoneme = sep_phonemes[i - 1][-1] - if prev_phoneme in VOWELS: - # 母音と「ん」のあとの伸ばし棒なので、その母音に変換 - sep_phonemes[i][0] = sep_phonemes[i - 1][-1] - else: - # 「。ーー」等おそらく予期しない長音記号 - # ダッシュの勘違いだと思われる - sep_phonemes[i][0] = "-" - else: - # 冒頭に長音記号が来ていおり、これはダッシュの勘違いと思われる - sep_phonemes[i][0] = "-" - if "ー" in sep_phonemes[i]: - for j in range(len(sep_phonemes[i])): - if sep_phonemes[i][j] == "ー": - sep_phonemes[i][j] = sep_phonemes[i][j - 1][-1] - return sep_phonemes - - -tokenizer = AutoTokenizer.from_pretrained("./bert/deberta-v2-large-japanese-char-wwm") - - -def align_tones( - phones_with_punct: list[str], phone_tone_list: list[tuple[str, int]] -) -> list[tuple[str, int]]: - """ - 例: - …私は、、そう思う。 - phones_with_punct: - [".", ".", ".", "w", "a", "t", "a", "sh", "i", "w", "a", ",", ",", "s", "o", "o", "o", "m", "o", "u", "."] - phone_tone_list: - [("w", 0), ("a", 0), ("t", 1), ("a", 1), ("sh", 1), ("i", 1), ("w", 1), ("a", 1), ("_", 0), ("s", 0), ("o", 0), ("o", 1), ("o", 1), ("m", 1), ("o", 1), ("u", 0))] - Return: - [(".", 0), (".", 0), (".", 0), ("w", 0), ("a", 0), ("t", 1), ("a", 1), ("sh", 1), ("i", 1), ("w", 1), ("a", 1), (",", 0), (",", 0), ("s", 0), ("o", 0), ("o", 1), ("o", 1), ("m", 1), ("o", 1), ("u", 0), (".", 0)] - """ - result: list[tuple[str, int]] = [] - tone_index = 0 - for phone in phones_with_punct: - if tone_index >= len(phone_tone_list): - # 余ったpunctuationがある場合 → (punctuation, 0)を追加 - result.append((phone, 0)) - elif phone == phone_tone_list[tone_index][0]: - # phone_tone_listの現在の音素と一致する場合 → toneをそこから取得、(phone, tone)を追加 - result.append((phone, phone_tone_list[tone_index][1])) - # 探すindexを1つ進める - tone_index += 1 - elif phone in PUNCTUATIONS: - # phoneがpunctuationの場合 → (phone, 0)を追加 - result.append((phone, 0)) - else: - logger.debug(f"phones: {phones_with_punct}") - logger.debug(f"phone_tone_list: {phone_tone_list}") - logger.debug(f"result: {result}") - logger.debug(f"tone_index: {tone_index}") - logger.debug(f"phone: {phone}") - raise ValueError(f"Unexpected phone: {phone}") - return result - - -def kata2phoneme_list(text: str) -> list[str]: - """ - 原則カタカナの`text`を受け取り、それをそのままいじらずに音素記号のリストに変換。 - 注意点: - - punctuationかその繰り返しが来た場合、punctuationたちをそのままリストにして返す。 - - 冒頭に続く「ー」はそのまま「ー」のままにする(`handle_long()`で処理される) - - 文中の「ー」は前の音素記号の最後の音素記号に変換される。 - 例: - `ーーソーナノカーー` → ["ー", "ー", "s", "o", "o", "n", "a", "n", "o", "k", "a", "a", "a"] - `?` → ["?"] - `!?!?!?!?!` → ["!", "?", "!", "?", "!", "?", "!", "?", "!"] - """ - if set(text).issubset(set(PUNCTUATIONS)): - return list(text) - # `text`がカタカナ(`ー`含む)のみからなるかどうかをチェック - if re.fullmatch(r"[\u30A0-\u30FF]+", text) is None: - raise ValueError(f"Input must be katakana only: {text}") - sorted_keys = sorted(MORA_KATA_TO_MORA_PHONEMES.keys(), key=len, reverse=True) - pattern = "|".join(map(re.escape, sorted_keys)) - - def mora2phonemes(mora: str) -> str: - cosonant, vowel = MORA_KATA_TO_MORA_PHONEMES[mora] - if cosonant is None: - return f" {vowel}" - return f" {cosonant} {vowel}" - - spaced_phonemes = re.sub(pattern, lambda m: mora2phonemes(m.group()), text) - - # 長音記号「ー」の処理 - long_pattern = r"(\w)(ー*)" - long_replacement = lambda m: m.group(1) + (" " + m.group(1)) * len(m.group(2)) - spaced_phonemes = re.sub(long_pattern, long_replacement, spaced_phonemes) - return spaced_phonemes.strip().split(" ") - - -if __name__ == "__main__": - tokenizer = AutoTokenizer.from_pretrained( - "./bert/deberta-v2-large-japanese-char-wwm" - ) - text = "こんにちは、世界。" - from text.japanese_bert import get_bert_feature - - text = normalize_text(text) - - phones, tones, word2ph = g2p(text) - bert = get_bert_feature(text, word2ph) - - print(phones, tones, word2ph, bert.shape) diff --git a/text/japanese_bert.py b/text/japanese_bert.py index fbeb94d..4c1bccc 100644 --- a/text/japanese_bert.py +++ b/text/japanese_bert.py @@ -1,24 +1,22 @@ import sys import torch -from transformers import AutoModelForMaskedLM, AutoTokenizer from config import config +from style_bert_vits2.constants import Languages +from style_bert_vits2.text_processing import bert_models from style_bert_vits2.text_processing.japanese.g2p import text_to_sep_kata -LOCAL_PATH = "./bert/deberta-v2-large-japanese-char-wwm" - -tokenizer = AutoTokenizer.from_pretrained(LOCAL_PATH) models = dict() def get_bert_feature( - text, + text: str, word2ph, - device=config.bert_gen_config.device, - assist_text=None, - assist_text_weight=0.7, + device = config.bert_gen_config.device, + assist_text: str | None = None, + assist_text_weight: float = 0.7, ): # 各単語が何文字かを作る`word2ph`を使う必要があるので、読めない文字は必ず無視する # でないと`word2ph`の結果とテキストの文字数結果が整合性が取れない @@ -37,8 +35,9 @@ def get_bert_feature( if device == "cuda" and not torch.cuda.is_available(): device = "cpu" if device not in models.keys(): - models[device] = AutoModelForMaskedLM.from_pretrained(LOCAL_PATH).to(device) + models[device] = bert_models.load_model(Languages.JP).to(device) with torch.no_grad(): + tokenizer = bert_models.load_tokenizer(Languages.JP) inputs = tokenizer(text, return_tensors="pt") for i in inputs: inputs[i] = inputs[i].to(device) From 62919e904e643d0d72ef18fd5424260a38157cb8 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Thu, 7 Mar 2024 03:32:07 +0000 Subject: [PATCH 13/64] Refactor: moved the module for extracting BERT features from text in each language to style_bert_vits2/text_processing/(language)/bert_feature.py --- app.py | 3 +- bert_gen.py | 10 +-- data_utils.py | 6 +- server_fastapi.py | 4 +- style_bert_vits2/models/infer.py | 4 +- style_bert_vits2/text_processing/__init__.py | 68 +++++++++++++++++++ .../text_processing/chinese/bert_feature.py | 36 +++++++--- .../text_processing/english/bert_feature.py | 36 +++++++--- .../text_processing/japanese/bert_feature.py | 39 ++++++++--- text/__init__.py | 43 ------------ text/chinese.py | 10 +-- text/english.py | 6 -- webui_dataset.py | 1 + webui_merge.py | 3 +- webui_style_vectors.py | 3 +- webui_train.py | 1 + 16 files changed, 172 insertions(+), 101 deletions(-) create mode 100644 style_bert_vits2/text_processing/__init__.py rename text/chinese_bert.py => style_bert_vits2/text_processing/chinese/bert_feature.py (74%) rename text/english_bert_mock.py => style_bert_vits2/text_processing/english/bert_feature.py (65%) rename text/japanese_bert.py => style_bert_vits2/text_processing/japanese/bert_feature.py (63%) delete mode 100644 text/__init__.py diff --git a/app.py b/app.py index a0a215c..02c01cd 100644 --- a/app.py +++ b/app.py @@ -10,6 +10,7 @@ import gradio as gr import torch import yaml +from common.tts_model import ModelHolder from style_bert_vits2.constants import ( DEFAULT_ASSIST_TEXT_WEIGHT, DEFAULT_LENGTH, @@ -25,11 +26,11 @@ from style_bert_vits2.constants import ( Languages, ) from style_bert_vits2.logging import logger -from common.tts_model import ModelHolder from style_bert_vits2.models.infer import InvalidToneError from style_bert_vits2.text_processing.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone from style_bert_vits2.text_processing.japanese.normalizer import normalize_text + # Get path settings with open(os.path.join("configs", "paths.yml"), "r", encoding="utf-8") as f: path_config: dict[str, str] = yaml.safe_load(f.read()) diff --git a/bert_gen.py b/bert_gen.py index fd0b54e..2b512d2 100644 --- a/bert_gen.py +++ b/bert_gen.py @@ -5,12 +5,12 @@ import torch import torch.multiprocessing as mp from tqdm import tqdm -from style_bert_vits2.models import commons import utils -from style_bert_vits2.logging import logger -from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from config import config -from text import cleaned_text_to_sequence, get_bert +from style_bert_vits2.logging import logger +from style_bert_vits2.models import commons +from style_bert_vits2.text_processing import cleaned_text_to_sequence, extract_bert_feature +from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT def process_line(x): @@ -45,7 +45,7 @@ def process_line(x): bert = torch.load(bert_path) assert bert.shape[-1] == len(phone) except Exception: - bert = get_bert(text, word2ph, language_str, device) + bert = extract_bert_feature(text, word2ph, language_str, device) assert bert.shape[-1] == len(phone) torch.save(bert, bert_path) diff --git a/data_utils.py b/data_utils.py index 1118102..7738247 100644 --- a/data_utils.py +++ b/data_utils.py @@ -7,12 +7,12 @@ import torch import torch.utils.data from tqdm import tqdm -from style_bert_vits2.models import commons from config import config from mel_processing import mel_spectrogram_torch, spectrogram_torch -from text import cleaned_text_to_sequence -from style_bert_vits2.logging import logger from utils import load_filepaths_and_text, load_wav_to_torch +from style_bert_vits2.logging import logger +from style_bert_vits2.models import commons +from style_bert_vits2.text_processing import cleaned_text_to_sequence """Multi speaker version""" diff --git a/server_fastapi.py b/server_fastapi.py index 132cc17..ca9520c 100644 --- a/server_fastapi.py +++ b/server_fastapi.py @@ -20,6 +20,8 @@ from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import FileResponse, Response from scipy.io import wavfile +from common.tts_model import Model, ModelHolder +from config import config from style_bert_vits2.constants import ( DEFAULT_ASSIST_TEXT_WEIGHT, DEFAULT_LENGTH, @@ -33,8 +35,6 @@ from style_bert_vits2.constants import ( Languages, ) from style_bert_vits2.logging import logger -from common.tts_model import Model, ModelHolder -from config import config ln = config.server_config.language diff --git a/style_bert_vits2/models/infer.py b/style_bert_vits2/models/infer.py index 99999a9..c4f6aed 100644 --- a/style_bert_vits2/models/infer.py +++ b/style_bert_vits2/models/infer.py @@ -1,12 +1,12 @@ 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 from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra +from style_bert_vits2.text_processing import cleaned_text_to_sequence, extract_bert_feature from style_bert_vits2.text_processing.cleaner import clean_text from style_bert_vits2.text_processing.symbols import SYMBOLS @@ -77,7 +77,7 @@ def get_text( for i in range(len(word2ph)): word2ph[i] = word2ph[i] * 2 word2ph[0] += 1 - bert_ori = get_bert( + bert_ori = extract_bert_feature( norm_text, word2ph, language_str, diff --git a/style_bert_vits2/text_processing/__init__.py b/style_bert_vits2/text_processing/__init__.py new file mode 100644 index 0000000..42bb3a5 --- /dev/null +++ b/style_bert_vits2/text_processing/__init__.py @@ -0,0 +1,68 @@ +import torch + +from style_bert_vits2.constants import Languages +from style_bert_vits2.text_processing.symbols import ( + LANGUAGE_ID_MAP, + LANGUAGE_TONE_START_MAP, + SYMBOLS, +) + + +_symbol_to_id = {s: i for i, s in enumerate(SYMBOLS)} + + +def cleaned_text_to_sequence(cleaned_text: str, tones: list[int], language: Languages) -> tuple[list[int], list[int], list[int]]: + """ + Converts a string of text to a sequence of IDs corresponding to the symbols in the text. + + Args: + cleaned_text (str): string to convert to a sequence + tones (list[int]): List of tones + language (Languages): Language of the text + + Returns: + tuple[list[int], list[int], list[int]]: List of integers corresponding to the symbols in the text + """ + + phones = [_symbol_to_id[symbol] for symbol in cleaned_text] + tone_start = LANGUAGE_TONE_START_MAP[language] + tones = [i + tone_start for i in tones] + lang_id = LANGUAGE_ID_MAP[language] + lang_ids = [lang_id for i in phones] + + return phones, tones, lang_ids + + +def extract_bert_feature( + text: str, + word2ph: list[int], + language: Languages, + device: torch.device | str, + assist_text: str | None = None, + assist_text_weight: float = 0.7, +) -> torch.Tensor: + """ + テキストから BERT の特徴量を抽出する + + Args: + text (str): テキスト + word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト + language (Languages): テキストの言語 + device (torch.device | str): 推論に利用するデバイス + assist_text (str | None, optional): 補助テキスト (デフォルト: None) + assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7) + + Returns: + torch.Tensor: BERT の特徴量 + """ + + if language == Languages.JP: + from style_bert_vits2.text_processing.japanese.bert_feature import extract_bert_feature + elif language == Languages.EN: + from style_bert_vits2.text_processing.english.bert_feature import extract_bert_feature + elif language == Languages.ZH: + from style_bert_vits2.text_processing.chinese.bert_feature import extract_bert_feature + else: + raise ValueError(f"Language {language} not supported") + + return extract_bert_feature(text, word2ph, device, assist_text, assist_text_weight) diff --git a/text/chinese_bert.py b/style_bert_vits2/text_processing/chinese/bert_feature.py similarity index 74% rename from text/chinese_bert.py rename to style_bert_vits2/text_processing/chinese/bert_feature.py index 2ee6e9d..c2085a4 100644 --- a/text/chinese_bert.py +++ b/style_bert_vits2/text_processing/chinese/bert_feature.py @@ -1,22 +1,36 @@ import sys import torch +from transformers import PreTrainedModel -from config import config from style_bert_vits2.constants import Languages from style_bert_vits2.text_processing import bert_models -models = dict() +models: dict[str, PreTrainedModel] = {} -def get_bert_feature( +def extract_bert_feature( text: str, - word2ph, - device = config.bert_gen_config.device, + word2ph: list[int], + device: torch.device | str, assist_text: str | None = None, assist_text_weight: float = 0.7, -): +) -> torch.Tensor: + """ + 中国語のテキストから BERT の特徴量を抽出する + + Args: + text (str): 中国語のテキスト + word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト + device (torch.device | str): 推論に利用するデバイス + assist_text (str | None, optional): 補助テキスト (デフォルト: None) + assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7) + + Returns: + torch.Tensor: BERT の特徴量 + """ + if ( sys.platform == "darwin" and torch.backends.mps.is_available() @@ -28,26 +42,30 @@ def get_bert_feature( if device == "cuda" and not torch.cuda.is_available(): device = "cpu" if device not in models.keys(): - models[device] = bert_models.load_model(Languages.ZH).to(device) + models[device] = bert_models.load_model(Languages.ZH).to(device) # type: ignore + + style_res_mean = None with torch.no_grad(): tokenizer = bert_models.load_tokenizer(Languages.ZH) inputs = tokenizer(text, return_tensors="pt") for i in inputs: - inputs[i] = inputs[i].to(device) + inputs[i] = inputs[i].to(device) # type: ignore res = models[device](**inputs, output_hidden_states=True) res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu() if assist_text: style_inputs = tokenizer(assist_text, return_tensors="pt") for i in style_inputs: - style_inputs[i] = style_inputs[i].to(device) + style_inputs[i] = style_inputs[i].to(device) # type: ignore style_res = models[device](**style_inputs, output_hidden_states=True) style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu() style_res_mean = style_res.mean(0) + assert len(word2ph) == len(text) + 2 word2phone = word2ph phone_level_feature = [] for i in range(len(word2phone)): if assist_text: + assert style_res_mean is not None repeat_feature = ( res[i].repeat(word2phone[i], 1) * (1 - assist_text_weight) + style_res_mean.repeat(word2phone[i], 1) * assist_text_weight diff --git a/text/english_bert_mock.py b/style_bert_vits2/text_processing/english/bert_feature.py similarity index 65% rename from text/english_bert_mock.py rename to style_bert_vits2/text_processing/english/bert_feature.py index 8e57df1..d9c1025 100644 --- a/text/english_bert_mock.py +++ b/style_bert_vits2/text_processing/english/bert_feature.py @@ -1,22 +1,36 @@ import sys import torch +from transformers import PreTrainedModel -from config import config from style_bert_vits2.constants import Languages from style_bert_vits2.text_processing import bert_models -models = dict() +models: dict[str, PreTrainedModel] = {} -def get_bert_feature( +def extract_bert_feature( text: str, - word2ph, - device = config.bert_gen_config.device, + word2ph: list[int], + device: torch.device | str, assist_text: str | None = None, assist_text_weight: float = 0.7, -): +) -> torch.Tensor: + """ + 英語のテキストから BERT の特徴量を抽出する + + Args: + text (str): 英語のテキスト + word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト + device (torch.device | str): 推論に利用するデバイス + assist_text (str | None, optional): 補助テキスト (デフォルト: None) + assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7) + + Returns: + torch.Tensor: BERT の特徴量 + """ + if ( sys.platform == "darwin" and torch.backends.mps.is_available() @@ -28,26 +42,30 @@ def get_bert_feature( if device == "cuda" and not torch.cuda.is_available(): device = "cpu" if device not in models.keys(): - models[device] = bert_models.load_model(Languages.EN).to(device) + models[device] = bert_models.load_model(Languages.EN).to(device) # type: ignore + + style_res_mean = None with torch.no_grad(): tokenizer = bert_models.load_tokenizer(Languages.EN) inputs = tokenizer(text, return_tensors="pt") for i in inputs: - inputs[i] = inputs[i].to(device) + inputs[i] = inputs[i].to(device) # type: ignore res = models[device](**inputs, output_hidden_states=True) res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu() if assist_text: style_inputs = tokenizer(assist_text, return_tensors="pt") for i in style_inputs: - style_inputs[i] = style_inputs[i].to(device) + style_inputs[i] = style_inputs[i].to(device) # type: ignore style_res = models[device](**style_inputs, output_hidden_states=True) style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu() style_res_mean = style_res.mean(0) + assert len(word2ph) == res.shape[0], (text, res.shape[0], len(word2ph)) word2phone = word2ph phone_level_feature = [] for i in range(len(word2phone)): if assist_text: + assert style_res_mean is not None repeat_feature = ( res[i].repeat(word2phone[i], 1) * (1 - assist_text_weight) + style_res_mean.repeat(word2phone[i], 1) * assist_text_weight diff --git a/text/japanese_bert.py b/style_bert_vits2/text_processing/japanese/bert_feature.py similarity index 63% rename from text/japanese_bert.py rename to style_bert_vits2/text_processing/japanese/bert_feature.py index 4c1bccc..078bb5c 100644 --- a/text/japanese_bert.py +++ b/style_bert_vits2/text_processing/japanese/bert_feature.py @@ -1,25 +1,39 @@ import sys import torch +from transformers import PreTrainedModel -from config import config from style_bert_vits2.constants import Languages from style_bert_vits2.text_processing import bert_models from style_bert_vits2.text_processing.japanese.g2p import text_to_sep_kata -models = dict() +models: dict[str, PreTrainedModel] = {} -def get_bert_feature( +def extract_bert_feature( text: str, - word2ph, - device = config.bert_gen_config.device, + word2ph: list[int], + device: torch.device | str, assist_text: str | None = None, assist_text_weight: float = 0.7, -): - # 各単語が何文字かを作る`word2ph`を使う必要があるので、読めない文字は必ず無視する - # でないと`word2ph`の結果とテキストの文字数結果が整合性が取れない +) -> torch.Tensor: + """ + 日本語のテキストから BERT の特徴量を抽出する + + Args: + text (str): 日本語のテキスト + word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト + device (torch.device | str): 推論に利用するデバイス + assist_text (str | None, optional): 補助テキスト (デフォルト: None) + assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7) + + Returns: + torch.Tensor: BERT の特徴量 + """ + + # 各単語が何文字かを作る `word2ph` を使う必要があるので、読めない文字は必ず無視する + # でないと `word2ph` の結果とテキストの文字数結果が整合性が取れない text = "".join(text_to_sep_kata(text, raise_yomi_error=False)[0]) if assist_text: @@ -35,18 +49,20 @@ def get_bert_feature( if device == "cuda" and not torch.cuda.is_available(): device = "cpu" if device not in models.keys(): - models[device] = bert_models.load_model(Languages.JP).to(device) + models[device] = bert_models.load_model(Languages.JP).to(device) # type: ignore + + style_res_mean = None with torch.no_grad(): tokenizer = bert_models.load_tokenizer(Languages.JP) inputs = tokenizer(text, return_tensors="pt") for i in inputs: - inputs[i] = inputs[i].to(device) + inputs[i] = inputs[i].to(device) # type: ignore res = models[device](**inputs, output_hidden_states=True) res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu() if assist_text: style_inputs = tokenizer(assist_text, return_tensors="pt") for i in style_inputs: - style_inputs[i] = style_inputs[i].to(device) + style_inputs[i] = style_inputs[i].to(device) # type: ignore style_res = models[device](**style_inputs, output_hidden_states=True) style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu() style_res_mean = style_res.mean(0) @@ -56,6 +72,7 @@ def get_bert_feature( phone_level_feature = [] for i in range(len(word2phone)): if assist_text: + assert style_res_mean is not None repeat_feature = ( res[i].repeat(word2phone[i], 1) * (1 - assist_text_weight) + style_res_mean.repeat(word2phone[i], 1) * assist_text_weight diff --git a/text/__init__.py b/text/__init__.py deleted file mode 100644 index efff830..0000000 --- a/text/__init__.py +++ /dev/null @@ -1,43 +0,0 @@ -from style_bert_vits2.constants import Languages -from style_bert_vits2.text_processing.symbols import * - - -_symbol_to_id = {s: i for i, s in enumerate(SYMBOLS)} - - -def cleaned_text_to_sequence(cleaned_text: str, tones: list[int], language: Languages): - """ - Converts a string of text to a sequence of IDs corresponding to the symbols in the text. - - Args: - text: string to convert to a sequence - - Returns: - List of integers corresponding to the symbols in the text - """ - phones = [_symbol_to_id[symbol] for symbol in cleaned_text] - tone_start = LANGUAGE_TONE_START_MAP[language] - tones = [i + tone_start for i in tones] - lang_id = LANGUAGE_ID_MAP[language] - lang_ids = [lang_id for i in phones] - return phones, tones, lang_ids - - -def get_bert( - text: str, - word2ph, - language: Languages, - device: str, - assist_text: str | None = None, - assist_text_weight: float = 0.7, -): - if language == Languages.ZH: - from .chinese_bert import get_bert_feature - elif language == Languages.EN: - from .english_bert_mock import get_bert_feature - elif language == Languages.JP: - from .japanese_bert import get_bert_feature - else: - raise ValueError(f"Language {language} not supported") - - return get_bert_feature(text, word2ph, device, assist_text, assist_text_weight) diff --git a/text/chinese.py b/text/chinese.py index 3d9c392..94266c2 100644 --- a/text/chinese.py +++ b/text/chinese.py @@ -176,20 +176,14 @@ def normalize_text(text): return text -def get_bert_feature(text, word2ph): - from text import chinese_bert - - return chinese_bert.get_bert_feature(text, word2ph) - - if __name__ == "__main__": - from text.chinese_bert import get_bert_feature + from style_bert_vits2.text_processing.chinese.bert_feature import extract_bert_feature text = "啊!但是《原神》是由,米哈\游自主, [研发]的一款全.新开放世界.冒险游戏" text = normalize_text(text) print(text) phones, tones, word2ph = g2p(text) - bert = get_bert_feature(text, word2ph) + bert = extract_bert_feature(text, word2ph, 'cuda') print(phones, tones, word2ph, bert.shape) diff --git a/text/english.py b/text/english.py index a6d71e9..419be3b 100644 --- a/text/english.py +++ b/text/english.py @@ -477,12 +477,6 @@ def g2p(text): return phones, tones, word2ph -def get_bert_feature(text, word2ph): - from text import english_bert_mock - - return english_bert_mock.get_bert_feature(text, word2ph) - - if __name__ == "__main__": # print(get_dict()) # print(eng_word_to_phoneme("hello")) diff --git a/webui_dataset.py b/webui_dataset.py index 3ad63c1..1796864 100644 --- a/webui_dataset.py +++ b/webui_dataset.py @@ -8,6 +8,7 @@ from style_bert_vits2.constants import GRADIO_THEME from style_bert_vits2.logging import logger from style_bert_vits2.utils.subprocess import run_script_with_log + # Get path settings with open(os.path.join("configs", "paths.yml"), "r", encoding="utf-8") as f: path_config: dict[str, str] = yaml.safe_load(f.read()) diff --git a/webui_merge.py b/webui_merge.py index 0a39902..5bf0568 100644 --- a/webui_merge.py +++ b/webui_merge.py @@ -11,9 +11,10 @@ import yaml from safetensors import safe_open from safetensors.torch import save_file +from common.tts_model import Model, ModelHolder from style_bert_vits2.constants import DEFAULT_STYLE, GRADIO_THEME from style_bert_vits2.logging import logger -from common.tts_model import Model, ModelHolder + voice_keys = ["dec"] voice_pitch_keys = ["flow"] diff --git a/webui_style_vectors.py b/webui_style_vectors.py index cf53ca2..1cbabbd 100644 --- a/webui_style_vectors.py +++ b/webui_style_vectors.py @@ -12,9 +12,10 @@ from sklearn.cluster import DBSCAN, AgglomerativeClustering, KMeans from sklearn.manifold import TSNE from umap import UMAP +from config import config from style_bert_vits2.constants import DEFAULT_STYLE, GRADIO_THEME from style_bert_vits2.logging import logger -from config import config + # Get path settings with open(os.path.join("configs", "paths.yml"), "r", encoding="utf-8") as f: diff --git a/webui_train.py b/webui_train.py index 6c89b26..fda31ca 100644 --- a/webui_train.py +++ b/webui_train.py @@ -19,6 +19,7 @@ from style_bert_vits2.logging import logger from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from style_bert_vits2.utils.subprocess import run_script_with_log, second_elem_of + logger_handler = None tensorboard_executed = False From 4f11b011fdca5a860ba83f9bdc4062b7667780fd Mon Sep 17 00:00:00 2001 From: tsukumi Date: Thu, 7 Mar 2024 03:54:02 +0000 Subject: [PATCH 14/64] Refactor: minor adjustments --- .../text_processing/chinese/bert_feature.py | 2 +- .../text_processing/english/bert_feature.py | 2 +- .../text_processing/japanese/bert_feature.py | 2 +- style_bert_vits2/text_processing/japanese/g2p.py | 10 +++++----- .../text_processing/japanese/g2p_utils.py | 2 +- style_bert_vits2/text_processing/symbols.py | 14 ++++++++------ style_bert_vits2/utils/subprocess.py | 4 +--- utils.py | 15 --------------- 8 files changed, 18 insertions(+), 33 deletions(-) diff --git a/style_bert_vits2/text_processing/chinese/bert_feature.py b/style_bert_vits2/text_processing/chinese/bert_feature.py index c2085a4..25024cb 100644 --- a/style_bert_vits2/text_processing/chinese/bert_feature.py +++ b/style_bert_vits2/text_processing/chinese/bert_feature.py @@ -7,7 +7,7 @@ from style_bert_vits2.constants import Languages from style_bert_vits2.text_processing import bert_models -models: dict[str, PreTrainedModel] = {} +models: dict[torch.device | str, PreTrainedModel] = {} def extract_bert_feature( diff --git a/style_bert_vits2/text_processing/english/bert_feature.py b/style_bert_vits2/text_processing/english/bert_feature.py index d9c1025..ec556c2 100644 --- a/style_bert_vits2/text_processing/english/bert_feature.py +++ b/style_bert_vits2/text_processing/english/bert_feature.py @@ -7,7 +7,7 @@ from style_bert_vits2.constants import Languages from style_bert_vits2.text_processing import bert_models -models: dict[str, PreTrainedModel] = {} +models: dict[torch.device | str, PreTrainedModel] = {} def extract_bert_feature( diff --git a/style_bert_vits2/text_processing/japanese/bert_feature.py b/style_bert_vits2/text_processing/japanese/bert_feature.py index 078bb5c..3ff9d7b 100644 --- a/style_bert_vits2/text_processing/japanese/bert_feature.py +++ b/style_bert_vits2/text_processing/japanese/bert_feature.py @@ -8,7 +8,7 @@ from style_bert_vits2.text_processing import bert_models from style_bert_vits2.text_processing.japanese.g2p import text_to_sep_kata -models: dict[str, PreTrainedModel] = {} +models: dict[torch.device | str, PreTrainedModel] = {} def extract_bert_feature( diff --git a/style_bert_vits2/text_processing/japanese/g2p.py b/style_bert_vits2/text_processing/japanese/g2p.py index 8ef4d3e..7968751 100644 --- a/style_bert_vits2/text_processing/japanese/g2p.py +++ b/style_bert_vits2/text_processing/japanese/g2p.py @@ -391,7 +391,7 @@ def __kata_to_phoneme_list(text: str) -> list[str]: if set(text).issubset(set(PUNCTUATIONS)): return list(text) - # `text`がカタカナ(`ー`含む)のみからなるかどうかをチェック + # `text` がカタカナ(`ー`含む)のみからなるかどうかをチェック if re.fullmatch(r"[\u30A0-\u30FF]+", text) is None: raise ValueError(f"Input must be katakana only: {text}") sorted_keys = sorted(MORA_KATA_TO_MORA_PHONEMES.keys(), key=len, reverse=True) @@ -438,15 +438,15 @@ def __align_tones( tone_index = 0 for phone in phones_with_punct: if tone_index >= len(phone_tone_list): - # 余ったpunctuationがある場合 → (punctuation, 0)を追加 + # 余った punctuation がある場合 → (punctuation, 0) を追加 result.append((phone, 0)) elif phone == phone_tone_list[tone_index][0]: - # phone_tone_listの現在の音素と一致する場合 → toneをそこから取得、(phone, tone)を追加 + # phone_tone_list の現在の音素と一致する場合 → tone をそこから取得、(phone, tone) を追加 result.append((phone, phone_tone_list[tone_index][1])) - # 探すindexを1つ進める + # 探す index を1つ進める tone_index += 1 elif phone in PUNCTUATIONS: - # phoneがpunctuationの場合 → (phone, 0)を追加 + # phone が punctuation の場合 → (phone, 0) を追加 result.append((phone, 0)) else: logger.debug(f"phones: {phones_with_punct}") diff --git a/style_bert_vits2/text_processing/japanese/g2p_utils.py b/style_bert_vits2/text_processing/japanese/g2p_utils.py index 3a91a00..4ea56e8 100644 --- a/style_bert_vits2/text_processing/japanese/g2p_utils.py +++ b/style_bert_vits2/text_processing/japanese/g2p_utils.py @@ -9,7 +9,7 @@ from style_bert_vits2.text_processing.symbols import PUNCTUATIONS def g2kata_tone(norm_text: str) -> list[tuple[str, int]]: """ テキストからカタカナとアクセントのペアのリストを返す。 - 推論時のみに使われるので、常に `raise_yomi_error=False` で g2p を呼ぶ。 + 推論時のみに使われる関数のため、常に `raise_yomi_error=False` を指定して g2p() を呼ぶ仕様になっている。 Args: norm_text: 正規化されたテキスト。 diff --git a/style_bert_vits2/text_processing/symbols.py b/style_bert_vits2/text_processing/symbols.py index d69bc1c..e3a650f 100644 --- a/style_bert_vits2/text_processing/symbols.py +++ b/style_bert_vits2/text_processing/symbols.py @@ -77,8 +77,8 @@ ZH_SYMBOLS = [ ] NUM_ZH_TONES = 6 -# japanese -JA_SYMBOLS = [ +# Japanese +JP_SYMBOLS = [ "N", "a", "a:", @@ -122,7 +122,7 @@ JA_SYMBOLS = [ "z", "zy", ] -NUM_JA_TONES = 2 +NUM_JP_TONES = 2 # English EN_SYMBOLS = [ @@ -169,23 +169,25 @@ EN_SYMBOLS = [ NUM_EN_TONES = 4 # Combine all symbols -NORMAL_SYMBOLS = sorted(set(ZH_SYMBOLS + JA_SYMBOLS + EN_SYMBOLS)) +NORMAL_SYMBOLS = sorted(set(ZH_SYMBOLS + JP_SYMBOLS + EN_SYMBOLS)) SYMBOLS = [PAD] + NORMAL_SYMBOLS + PUNCTUATION_SYMBOLS SIL_PHONEMES_IDS = [SYMBOLS.index(i) for i in PUNCTUATION_SYMBOLS] # Combine all tones -NUM_TONES = NUM_ZH_TONES + NUM_JA_TONES + NUM_EN_TONES +NUM_TONES = NUM_ZH_TONES + NUM_JP_TONES + NUM_EN_TONES # Language maps LANGUAGE_ID_MAP = {"ZH": 0, "JP": 1, "EN": 2} NUM_LANGUAGES = len(LANGUAGE_ID_MAP.keys()) +# Language tone start map LANGUAGE_TONE_START_MAP = { "ZH": 0, "JP": NUM_ZH_TONES, - "EN": NUM_ZH_TONES + NUM_JA_TONES, + "EN": NUM_ZH_TONES + NUM_JP_TONES, } + if __name__ == "__main__": a = set(ZH_SYMBOLS) b = set(EN_SYMBOLS) diff --git a/style_bert_vits2/utils/subprocess.py b/style_bert_vits2/utils/subprocess.py index b152702..5ff267b 100644 --- a/style_bert_vits2/utils/subprocess.py +++ b/style_bert_vits2/utils/subprocess.py @@ -5,8 +5,6 @@ from typing import Any, Callable from style_bert_vits2.logging import logger from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT -PYTHON = sys.executable - def run_script_with_log(cmd: list[str], ignore_warning: bool = False) -> tuple[bool, str]: """ @@ -22,7 +20,7 @@ def run_script_with_log(cmd: list[str], ignore_warning: bool = False) -> tuple[b logger.info(f"Running: {' '.join(cmd)}") result = subprocess.run( - [PYTHON] + cmd, + [sys.executable] + cmd, stdout = SAFE_STDOUT, stderr = subprocess.PIPE, text = True, diff --git a/utils.py b/utils.py index 80dfa66..d4ff765 100644 --- a/utils.py +++ b/utils.py @@ -450,21 +450,6 @@ class HParams: return self.__dict__.__repr__() -def load_model(model_path, config_path): - hps = get_hparams_from_file(config_path) - net = SynthesizerTrn( - # len(symbols), - 108, - hps.data.filter_length // 2 + 1, - hps.train.segment_size // hps.data.hop_length, - n_speakers=hps.data.n_speakers, - **hps.model, - ).to("cpu") - _ = net.eval() - _ = load_checkpoint(model_path, net, None, skip_optimizer=True) - return net - - def mix_model( network1, network2, output_path, voice_ratio=(0.5, 0.5), tone_ratio=(0.5, 0.5) ): From b01168309d20296de273222052a537871e758678 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Thu, 7 Mar 2024 04:08:56 +0000 Subject: [PATCH 15/64] Remove: remove currently unused code in utils.py --- utils.py | 288 ++++++++++++++++++++++--------------------------------- 1 file changed, 117 insertions(+), 171 deletions(-) diff --git a/utils.py b/utils.py index d4ff765..8c7e842 100644 --- a/utils.py +++ b/utils.py @@ -8,28 +8,16 @@ import subprocess import numpy as np import torch -from huggingface_hub import hf_hub_download from safetensors import safe_open from safetensors.torch import save_file from scipy.io.wavfile import read from style_bert_vits2.logging import logger + MATPLOTLIB_FLAG = False -def download_checkpoint( - dir_path, repo_config, token=None, regex="G_*.pth", mirror="openi" -): - repo_id = repo_config["repo_id"] - f_list = glob.glob(os.path.join(dir_path, regex)) - if f_list: - print("Use existed model, skip downloading.") - return - for file in ["DUR_0.pth", "D_0.pth", "G_0.pth"]: - hf_hub_download(repo_id, file, local_dir=dir_path, local_dir_use_symlinks=False) - - def load_checkpoint( checkpoint_path, model, optimizer=None, skip_optimizer=False, for_infer=False ): @@ -114,28 +102,54 @@ def save_checkpoint(model, optimizer, learning_rate, iteration, checkpoint_path) ) -def save_safetensors(model, iteration, checkpoint_path, is_half=False, for_infer=False): - """ - Save model with safetensors. - """ - if hasattr(model, "module"): - state_dict = model.module.state_dict() - else: - state_dict = model.state_dict() - keys = [] - for k in state_dict: - if "enc_q" in k and for_infer: - continue # noqa: E701 - keys.append(k) +def clean_checkpoints(path_to_models="logs/44k/", n_ckpts_to_keep=2, sort_by_time=True): + """Freeing up space by deleting saved ckpts - new_dict = ( - {k: state_dict[k].half() for k in keys} - if is_half - else {k: state_dict[k] for k in keys} - ) - new_dict["iteration"] = torch.LongTensor([iteration]) - logger.info(f"Saved safetensors to {checkpoint_path}") - save_file(new_dict, checkpoint_path) + Arguments: + path_to_models -- Path to the model directory + n_ckpts_to_keep -- Number of ckpts to keep, excluding G_0.pth and D_0.pth + sort_by_time -- True -> chronologically delete ckpts + False -> lexicographically delete ckpts + """ + import re + + ckpts_files = [ + f + for f in os.listdir(path_to_models) + if os.path.isfile(os.path.join(path_to_models, f)) + ] + + def name_key(_f): + return int(re.compile("._(\\d+)\\.pth").match(_f).group(1)) + + def time_key(_f): + return os.path.getmtime(os.path.join(path_to_models, _f)) + + sort_key = time_key if sort_by_time else name_key + + def x_sorted(_x): + return sorted( + [f for f in ckpts_files if f.startswith(_x) and not f.endswith("_0.pth")], + key=sort_key, + ) + + to_del = [ + os.path.join(path_to_models, fn) + for fn in ( + x_sorted("G_")[:-n_ckpts_to_keep] + + x_sorted("D_")[:-n_ckpts_to_keep] + + x_sorted("WD_")[:-n_ckpts_to_keep] + + x_sorted("DUR_")[:-n_ckpts_to_keep] + ) + ] + + def del_info(fn): + return logger.info(f"Free up space by deleting ckpt {fn}") + + def del_routine(x): + return [os.remove(x), del_info(x)] + + [del_routine(fn) for fn in to_del] def load_safetensors(checkpoint_path, model, for_infer=False): @@ -169,6 +183,30 @@ def load_safetensors(checkpoint_path, model, for_infer=False): return model, iteration +def save_safetensors(model, iteration, checkpoint_path, is_half=False, for_infer=False): + """ + Save model with safetensors. + """ + if hasattr(model, "module"): + state_dict = model.module.state_dict() + else: + state_dict = model.state_dict() + keys = [] + for k in state_dict: + if "enc_q" in k and for_infer: + continue # noqa: E701 + keys.append(k) + + new_dict = ( + {k: state_dict[k].half() for k in keys} + if is_half + else {k: state_dict[k] for k in keys} + ) + new_dict["iteration"] = torch.LongTensor([iteration]) + logger.info(f"Saved safetensors to {checkpoint_path}") + save_file(new_dict, checkpoint_path) + + def summarize( writer, global_step, @@ -274,6 +312,51 @@ def load_filepaths_and_text(filename, split="|"): return filepaths_and_text +def get_logger(model_dir, filename="train.log"): + global logger + logger = logging.getLogger(os.path.basename(model_dir)) + logger.setLevel(logging.DEBUG) + + formatter = logging.Formatter("%(asctime)s\t%(name)s\t%(levelname)s\t%(message)s") + if not os.path.exists(model_dir): + os.makedirs(model_dir) + h = logging.FileHandler(os.path.join(model_dir, filename)) + h.setLevel(logging.DEBUG) + h.setFormatter(formatter) + logger.addHandler(h) + return logger + + +def get_steps(model_path): + matches = re.findall(r"\d+", model_path) + return matches[-1] if matches else None + + +def check_git_hash(model_dir): + source_dir = os.path.dirname(os.path.realpath(__file__)) + if not os.path.exists(os.path.join(source_dir, ".git")): + logger.warning( + "{} is not a git repository, therefore hash value comparison will be ignored.".format( + source_dir + ) + ) + return + + cur_hash = subprocess.getoutput("git rev-parse HEAD") + + path = os.path.join(model_dir, "githash") + if os.path.exists(path): + saved_hash = open(path).read() + if saved_hash != cur_hash: + logger.warning( + "git hash values are different. {}(saved) != {}(current)".format( + saved_hash[:8], cur_hash[:8] + ) + ) + else: + open(path, "w").write(cur_hash) + + def get_hparams(init=True): parser = argparse.ArgumentParser() parser.add_argument( @@ -307,67 +390,6 @@ def get_hparams(init=True): return hparams -def clean_checkpoints(path_to_models="logs/44k/", n_ckpts_to_keep=2, sort_by_time=True): - """Freeing up space by deleting saved ckpts - - Arguments: - path_to_models -- Path to the model directory - n_ckpts_to_keep -- Number of ckpts to keep, excluding G_0.pth and D_0.pth - sort_by_time -- True -> chronologically delete ckpts - False -> lexicographically delete ckpts - """ - import re - - ckpts_files = [ - f - for f in os.listdir(path_to_models) - if os.path.isfile(os.path.join(path_to_models, f)) - ] - - def name_key(_f): - return int(re.compile("._(\\d+)\\.pth").match(_f).group(1)) - - def time_key(_f): - return os.path.getmtime(os.path.join(path_to_models, _f)) - - sort_key = time_key if sort_by_time else name_key - - def x_sorted(_x): - return sorted( - [f for f in ckpts_files if f.startswith(_x) and not f.endswith("_0.pth")], - key=sort_key, - ) - - to_del = [ - os.path.join(path_to_models, fn) - for fn in ( - x_sorted("G_")[:-n_ckpts_to_keep] - + x_sorted("D_")[:-n_ckpts_to_keep] - + x_sorted("WD_")[:-n_ckpts_to_keep] - + x_sorted("DUR_")[:-n_ckpts_to_keep] - ) - ] - - def del_info(fn): - return logger.info(f"Free up space by deleting ckpt {fn}") - - def del_routine(x): - return [os.remove(x), del_info(x)] - - [del_routine(fn) for fn in to_del] - - -def get_hparams_from_dir(model_dir): - config_save_path = os.path.join(model_dir, "config.json") - with open(config_save_path, "r", encoding="utf-8") as f: - data = f.read() - config = json.loads(data) - - hparams = HParams(**config) - hparams.model_dir = model_dir - return hparams - - def get_hparams_from_file(config_path): # print("config_path: ", config_path) with open(config_path, "r", encoding="utf-8") as f: @@ -378,46 +400,6 @@ def get_hparams_from_file(config_path): return hparams -def check_git_hash(model_dir): - source_dir = os.path.dirname(os.path.realpath(__file__)) - if not os.path.exists(os.path.join(source_dir, ".git")): - logger.warning( - "{} is not a git repository, therefore hash value comparison will be ignored.".format( - source_dir - ) - ) - return - - cur_hash = subprocess.getoutput("git rev-parse HEAD") - - path = os.path.join(model_dir, "githash") - if os.path.exists(path): - saved_hash = open(path).read() - if saved_hash != cur_hash: - logger.warning( - "git hash values are different. {}(saved) != {}(current)".format( - saved_hash[:8], cur_hash[:8] - ) - ) - else: - open(path, "w").write(cur_hash) - - -def get_logger(model_dir, filename="train.log"): - global logger - logger = logging.getLogger(os.path.basename(model_dir)) - logger.setLevel(logging.DEBUG) - - formatter = logging.Formatter("%(asctime)s\t%(name)s\t%(levelname)s\t%(message)s") - if not os.path.exists(model_dir): - os.makedirs(model_dir) - h = logging.FileHandler(os.path.join(model_dir, filename)) - h.setLevel(logging.DEBUG) - h.setFormatter(formatter) - logger.addHandler(h) - return logger - - class HParams: def __init__(self, **kwargs): for k, v in kwargs.items(): @@ -448,39 +430,3 @@ class HParams: def __repr__(self): return self.__dict__.__repr__() - - -def mix_model( - network1, network2, output_path, voice_ratio=(0.5, 0.5), tone_ratio=(0.5, 0.5) -): - if hasattr(network1, "module"): - state_dict1 = network1.module.state_dict() - state_dict2 = network2.module.state_dict() - else: - state_dict1 = network1.state_dict() - state_dict2 = network2.state_dict() - for k in state_dict1.keys(): - if k not in state_dict2.keys(): - continue - if "enc_p" in k: - state_dict1[k] = ( - state_dict1[k].clone() * tone_ratio[0] - + state_dict2[k].clone() * tone_ratio[1] - ) - else: - state_dict1[k] = ( - state_dict1[k].clone() * voice_ratio[0] - + state_dict2[k].clone() * voice_ratio[1] - ) - for k in state_dict2.keys(): - if k not in state_dict1.keys(): - state_dict1[k] = state_dict2[k].clone() - torch.save( - {"model": state_dict1, "iteration": 0, "optimizer": None, "learning_rate": 0}, - output_path, - ) - - -def get_steps(model_path): - matches = re.findall(r"\d+", model_path) - return matches[-1] if matches else None From def6d88425d296457c3b1e04f3bc969c6e5d60cd Mon Sep 17 00:00:00 2001 From: tsukumi Date: Thu, 7 Mar 2024 04:19:40 +0000 Subject: [PATCH 16/64] Refactor: style_bert_vits2/text_processing/cleaner.py integrated into style_bert_vits2/text_processing/__init__.py This was often used in 3 function sets and felt like a wasteful division with few lines. --- preprocess_text.py | 2 +- style_bert_vits2/models/infer.py | 3 +- style_bert_vits2/text_processing/__init__.py | 83 ++++++++++++++------ style_bert_vits2/text_processing/cleaner.py | 40 ---------- 4 files changed, 63 insertions(+), 65 deletions(-) delete mode 100644 style_bert_vits2/text_processing/cleaner.py diff --git a/preprocess_text.py b/preprocess_text.py index 92e00b9..03e1232 100644 --- a/preprocess_text.py +++ b/preprocess_text.py @@ -9,7 +9,7 @@ from tqdm import tqdm from config import config from style_bert_vits2.logging import logger -from style_bert_vits2.text_processing.cleaner import clean_text +from style_bert_vits2.text_processing import clean_text from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT preprocess_text_config = config.preprocess_text_config diff --git a/style_bert_vits2/models/infer.py b/style_bert_vits2/models/infer.py index c4f6aed..0e556d0 100644 --- a/style_bert_vits2/models/infer.py +++ b/style_bert_vits2/models/infer.py @@ -6,8 +6,7 @@ from style_bert_vits2.logging import logger from style_bert_vits2.models import commons from style_bert_vits2.models.models import SynthesizerTrn from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra -from style_bert_vits2.text_processing import cleaned_text_to_sequence, extract_bert_feature -from style_bert_vits2.text_processing.cleaner import clean_text +from style_bert_vits2.text_processing import clean_text, cleaned_text_to_sequence, extract_bert_feature from style_bert_vits2.text_processing.symbols import SYMBOLS diff --git a/style_bert_vits2/text_processing/__init__.py b/style_bert_vits2/text_processing/__init__.py index 42bb3a5..cd56ee7 100644 --- a/style_bert_vits2/text_processing/__init__.py +++ b/style_bert_vits2/text_processing/__init__.py @@ -11,28 +11,6 @@ from style_bert_vits2.text_processing.symbols import ( _symbol_to_id = {s: i for i, s in enumerate(SYMBOLS)} -def cleaned_text_to_sequence(cleaned_text: str, tones: list[int], language: Languages) -> tuple[list[int], list[int], list[int]]: - """ - Converts a string of text to a sequence of IDs corresponding to the symbols in the text. - - Args: - cleaned_text (str): string to convert to a sequence - tones (list[int]): List of tones - language (Languages): Language of the text - - Returns: - tuple[list[int], list[int], list[int]]: List of integers corresponding to the symbols in the text - """ - - phones = [_symbol_to_id[symbol] for symbol in cleaned_text] - tone_start = LANGUAGE_TONE_START_MAP[language] - tones = [i + tone_start for i in tones] - lang_id = LANGUAGE_ID_MAP[language] - lang_ids = [lang_id for i in phones] - - return phones, tones, lang_ids - - def extract_bert_feature( text: str, word2ph: list[int], @@ -66,3 +44,64 @@ def extract_bert_feature( raise ValueError(f"Language {language} not supported") return extract_bert_feature(text, word2ph, device, assist_text, assist_text_weight) + + +def clean_text( + text: str, + language: Languages, + use_jp_extra: bool = True, + raise_yomi_error: bool = False, +) -> tuple[str, list[str], list[int], list[int]]: + """ + テキストをクリーニングし、音素に変換する + + Args: + text (str): クリーニングするテキスト + language (Languages): テキストの言語 + 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 == 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, 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 == Languages.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 + + +def cleaned_text_to_sequence(cleaned_phones: list[str], tones: list[int], language: Languages) -> tuple[list[int], list[int], list[int]]: + """ + テキスト文字列を、テキスト内の記号に対応する一連の ID に変換する + + Args: + cleaned_phones (list[str]): clean_text() でクリーニングされた音素のリスト (?) + tones (list[int]): 各音素のアクセント + language (Languages): テキストの言語 + + Returns: + tuple[list[int], list[int], list[int]]: List of integers corresponding to the symbols in the text + """ + + phones = [_symbol_to_id[symbol] for symbol in cleaned_phones] + tone_start = LANGUAGE_TONE_START_MAP[language] + tones = [i + tone_start for i in tones] + lang_id = LANGUAGE_ID_MAP[language] + lang_ids = [lang_id for i in phones] + + return phones, tones, lang_ids diff --git a/style_bert_vits2/text_processing/cleaner.py b/style_bert_vits2/text_processing/cleaner.py deleted file mode 100644 index 0420099..0000000 --- a/style_bert_vits2/text_processing/cleaner.py +++ /dev/null @@ -1,40 +0,0 @@ -from style_bert_vits2.constants import Languages - - -def clean_text( - text: str, - language: Languages, - use_jp_extra: bool = True, - raise_yomi_error: bool = False, -) -> tuple[str, list[str], list[int], list[int]]: - """ - テキストをクリーニングし、音素に変換する - - Args: - text (str): クリーニングするテキスト - language (Languages): テキストの言語 - 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 == 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, 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 == Languages.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 From 1450bfd06f6510a3003a4613f234cb63d932b0e1 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Thu, 7 Mar 2024 04:21:48 +0000 Subject: [PATCH 17/64] Remove: remove webui.py, which is no longer maintained in Style-Bert-VITS2 Since app.py and server_editor.py already exist as alternative Web UI, there is no need to revive webui.py in the future. --- webui.py | 559 ------------------------------------------------------- 1 file changed, 559 deletions(-) delete mode 100644 webui.py diff --git a/webui.py b/webui.py deleted file mode 100644 index 1f31e94..0000000 --- a/webui.py +++ /dev/null @@ -1,559 +0,0 @@ -""" -Original `webui.py` for Bert-VITS2, not working with Style-Bert-VITS2 yet. -""" - -# flake8: noqa: E402 -import os -import logging -import re_matching -from tools.sentence import split_by_language - -logging.getLogger("numba").setLevel(logging.WARNING) -logging.getLogger("markdown_it").setLevel(logging.WARNING) -logging.getLogger("urllib3").setLevel(logging.WARNING) -logging.getLogger("matplotlib").setLevel(logging.WARNING) - -logging.basicConfig( - level=logging.INFO, format="| %(name)s | %(levelname)s | %(message)s" -) - -logger = logging.getLogger(__name__) - -import gradio as gr -import librosa -import numpy as np -import torch -import webbrowser - -import utils -from config import config -from style_bert_vits2.models.infer import infer, latest_version, get_net_g, infer_multilang -from tools.translate import translate - -net_g = None - -device = config.webui_config.device -if device == "mps": - os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = "1" - - -def generate_audio( - slices, - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - speaker, - language, - reference_audio, - emotion, - style_text, - style_weight, - skip_start=False, - skip_end=False, -): - audio_list = [] - # silence = np.zeros(hps.data.sampling_rate // 2, dtype=np.int16) - with torch.no_grad(): - for idx, piece in enumerate(slices): - skip_start = idx != 0 - skip_end = idx != len(slices) - 1 - audio = infer( - piece, - reference_audio=reference_audio, - emotion=emotion, - sdp_ratio=sdp_ratio, - noise_scale=noise_scale, - noise_scale_w=noise_scale_w, - length_scale=length_scale, - sid=speaker, - language=language, - hps=hps, - net_g=net_g, - device=device, - skip_start=skip_start, - skip_end=skip_end, - assist_text=style_text, - assist_text_weight=style_weight, - ) - audio16bit = gr.processing_utils.convert_to_16_bit_wav(audio) - audio_list.append(audio16bit) - return audio_list - - -def generate_audio_multilang( - slices, - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - speaker, - language, - reference_audio, - emotion, - skip_start=False, - skip_end=False, -): - audio_list = [] - # silence = np.zeros(hps.data.sampling_rate // 2, dtype=np.int16) - with torch.no_grad(): - for idx, piece in enumerate(slices): - skip_start = idx != 0 - skip_end = idx != len(slices) - 1 - audio = infer_multilang( - piece, - reference_audio=reference_audio, - emotion=emotion, - sdp_ratio=sdp_ratio, - noise_scale=noise_scale, - noise_scale_w=noise_scale_w, - length_scale=length_scale, - sid=speaker, - language=language[idx], - hps=hps, - net_g=net_g, - device=device, - skip_start=skip_start, - skip_end=skip_end, - ) - audio16bit = gr.processing_utils.convert_to_16_bit_wav(audio) - audio_list.append(audio16bit) - return audio_list - - -def tts_split( - text: str, - speaker, - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - language, - cut_by_sent, - interval_between_para, - interval_between_sent, - reference_audio, - emotion, - style_text, - style_weight, -): - while text.find("\n\n") != -1: - text = text.replace("\n\n", "\n") - text = text.replace("|", "") - para_list = re_matching.cut_para(text) - para_list = [p for p in para_list if p != ""] - audio_list = [] - for p in para_list: - if not cut_by_sent: - audio_list += process_text( - p, - speaker, - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - language, - reference_audio, - emotion, - style_text, - style_weight, - ) - silence = np.zeros((int)(44100 * interval_between_para), dtype=np.int16) - audio_list.append(silence) - else: - audio_list_sent = [] - sent_list = re_matching.cut_sent(p) - sent_list = [s for s in sent_list if s != ""] - for s in sent_list: - audio_list_sent += process_text( - s, - speaker, - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - language, - reference_audio, - emotion, - style_text, - style_weight, - ) - silence = np.zeros((int)(44100 * interval_between_sent)) - audio_list_sent.append(silence) - if (interval_between_para - interval_between_sent) > 0: - silence = np.zeros( - (int)(44100 * (interval_between_para - interval_between_sent)) - ) - audio_list_sent.append(silence) - audio16bit = gr.processing_utils.convert_to_16_bit_wav( - np.concatenate(audio_list_sent) - ) # 对完整句子做音量归一 - audio_list.append(audio16bit) - audio_concat = np.concatenate(audio_list) - return ("Success", (hps.data.sampling_rate, audio_concat)) - - -def process_mix(slice): - _speaker = slice.pop() - _text, _lang = [], [] - for lang, content in slice: - content = content.split("|") - content = [part for part in content if part != ""] - if len(content) == 0: - continue - if len(_text) == 0: - _text = [[part] for part in content] - _lang = [[lang] for part in content] - else: - _text[-1].append(content[0]) - _lang[-1].append(lang) - if len(content) > 1: - _text += [[part] for part in content[1:]] - _lang += [[lang] for part in content[1:]] - return _text, _lang, _speaker - - -def process_auto(text): - _text, _lang = [], [] - for slice in text.split("|"): - if slice == "": - continue - temp_text, temp_lang = [], [] - sentences_list = split_by_language(slice, target_languages=["zh", "ja", "en"]) - for sentence, lang in sentences_list: - if sentence == "": - continue - temp_text.append(sentence) - if lang == "ja": - lang = "jp" - temp_lang.append(lang.upper()) - _text.append(temp_text) - _lang.append(temp_lang) - return _text, _lang - - -def process_text( - text: str, - speaker, - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - language, - reference_audio, - emotion, - style_text=None, - style_weight=0, -): - audio_list = [] - if language == "mix": - bool_valid, str_valid = re_matching.validate_text(text) - if not bool_valid: - return str_valid, ( - hps.data.sampling_rate, - np.concatenate([np.zeros(hps.data.sampling_rate // 2)]), - ) - for slice in re_matching.text_matching(text): - _text, _lang, _speaker = process_mix(slice) - if _speaker is None: - continue - print(f"Text: {_text}\nLang: {_lang}") - audio_list.extend( - generate_audio_multilang( - _text, - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - _speaker, - _lang, - reference_audio, - emotion, - ) - ) - elif language.lower() == "auto": - _text, _lang = process_auto(text) - print(f"Text: {_text}\nLang: {_lang}") - audio_list.extend( - generate_audio_multilang( - _text, - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - speaker, - _lang, - reference_audio, - emotion, - ) - ) - else: - audio_list.extend( - generate_audio( - text.split("|"), - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - speaker, - language, - reference_audio, - emotion, - style_text, - style_weight, - ) - ) - return audio_list - - -def tts_fn( - text: str, - speaker, - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - language, - reference_audio, - emotion, - prompt_mode, - style_text=None, - style_weight=0, -): - if style_text == "": - style_text = None - if prompt_mode == "Audio prompt": - if reference_audio == None: - return ("Invalid audio prompt", None) - else: - reference_audio = load_audio(reference_audio)[1] - else: - reference_audio = None - - audio_list = process_text( - text, - speaker, - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - language, - reference_audio, - emotion, - style_text, - style_weight, - ) - - audio_concat = np.concatenate(audio_list) - return "Success", (hps.data.sampling_rate, audio_concat) - - -def format_utils(text, speaker): - _text, _lang = process_auto(text) - res = f"[{speaker}]" - for lang_s, content_s in zip(_lang, _text): - for lang, content in zip(lang_s, content_s): - res += f"<{lang.lower()}>{content}" - res += "|" - return "mix", res[:-1] - - -def load_audio(path): - audio, sr = librosa.load(path, 48000) - # audio = librosa.resample(audio, 44100, 48000) - return sr, audio - - -def gr_util(item): - if item == "Text prompt": - return {"visible": True, "__type__": "update"}, { - "visible": False, - "__type__": "update", - } - else: - return {"visible": False, "__type__": "update"}, { - "visible": True, - "__type__": "update", - } - - -if __name__ == "__main__": - if config.webui_config.debug: - logger.info("Enable DEBUG-LEVEL log") - logging.basicConfig(level=logging.DEBUG) - hps = utils.get_hparams_from_file(config.webui_config.config_path) - # 若config.json中未指定版本则默认为最新版本 - version = hps.version if hasattr(hps, "version") else latest_version - net_g = get_net_g( - model_path=config.webui_config.model, version=version, device=device, hps=hps - ) - speaker_ids = hps.data.spk2id - speakers = list(speaker_ids.keys()) - languages = ["ZH", "JP", "EN", "mix", "auto"] - with gr.Blocks() as app: - with gr.Row(): - with gr.Column(): - text = gr.TextArea( - label="输入文本内容", - placeholder=""" - 如果你选择语言为\'mix\',必须按照格式输入,否则报错: - 格式举例(zh是中文,jp是日语,不区分大小写;说话人举例:gongzi): - [说话人1]你好,こんにちは! こんにちは,世界。 - [说话人2]你好吗?元気ですか? - [说话人3]谢谢。どういたしまして。 - ... - 另外,所有的语言选项都可以用'|'分割长段实现分句生成。 - """, - ) - trans = gr.Button("中翻日", variant="primary") - slicer = gr.Button("快速切分", variant="primary") - formatter = gr.Button("检测语言,并整理为 MIX 格式", variant="primary") - speaker = gr.Dropdown( - choices=speakers, value=speakers[0], label="Speaker" - ) - _ = gr.Markdown( - value="提示模式(Prompt mode):可选文字提示或音频提示,用于生成文字或音频指定风格的声音。\n", - visible=False, - ) - prompt_mode = gr.Radio( - ["Text prompt", "Audio prompt"], - label="Prompt Mode", - value="Text prompt", - visible=False, - ) - text_prompt = gr.Textbox( - label="Text prompt", - placeholder="用文字描述生成风格。如:Happy", - value="Happy", - visible=False, - ) - audio_prompt = gr.Audio( - label="Audio prompt", type="filepath", visible=False - ) - sdp_ratio = gr.Slider( - minimum=0, maximum=1, value=0.5, step=0.1, label="SDP Ratio" - ) - noise_scale = gr.Slider( - minimum=0.1, maximum=2, value=0.6, step=0.1, label="Noise" - ) - noise_scale_w = gr.Slider( - minimum=0.1, maximum=2, value=0.9, step=0.1, label="Noise_W" - ) - length_scale = gr.Slider( - minimum=0.1, maximum=2, value=1.0, step=0.1, label="Length" - ) - language = gr.Dropdown( - choices=languages, value=languages[0], label="Language" - ) - btn = gr.Button("生成音频!", variant="primary") - with gr.Column(): - with gr.Accordion("融合文本语义", open=False): - gr.Markdown( - value="使用辅助文本的语意来辅助生成对话(语言保持与主文本相同)\n\n" - "**注意**:不要使用**指令式文本**(如:开心),要使用**带有强烈情感的文本**(如:我好快乐!!!)\n\n" - "效果较不明确,留空即为不使用该功能" - ) - style_text = gr.Textbox(label="辅助文本") - style_weight = gr.Slider( - minimum=0, - maximum=1, - value=0.7, - step=0.1, - label="Weight", - info="主文本和辅助文本的bert混合比率,0表示仅主文本,1表示仅辅助文本", - ) - with gr.Row(): - with gr.Column(): - interval_between_sent = gr.Slider( - minimum=0, - maximum=5, - value=0.2, - step=0.1, - label="句间停顿(秒),勾选按句切分才生效", - ) - interval_between_para = gr.Slider( - minimum=0, - maximum=10, - value=1, - step=0.1, - label="段间停顿(秒),需要大于句间停顿才有效", - ) - opt_cut_by_sent = gr.Checkbox( - label="按句切分 在按段落切分的基础上再按句子切分文本" - ) - slicer = gr.Button("切分生成", variant="primary") - text_output = gr.Textbox(label="状态信息") - audio_output = gr.Audio(label="输出音频") - # explain_image = gr.Image( - # label="参数解释信息", - # show_label=True, - # show_share_button=False, - # show_download_button=False, - # value=os.path.abspath("./img/参数说明.png"), - # ) - btn.click( - tts_fn, - inputs=[ - text, - speaker, - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - language, - audio_prompt, - text_prompt, - prompt_mode, - style_text, - style_weight, - ], - outputs=[text_output, audio_output], - ) - - trans.click( - translate, - inputs=[text], - outputs=[text], - ) - slicer.click( - tts_split, - inputs=[ - text, - speaker, - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - language, - opt_cut_by_sent, - interval_between_para, - interval_between_sent, - audio_prompt, - text_prompt, - style_text, - style_weight, - ], - outputs=[text_output, audio_output], - ) - - prompt_mode.change( - lambda x: gr_util(x), - inputs=[prompt_mode], - outputs=[text_prompt, audio_prompt], - ) - - audio_prompt.upload( - lambda x: load_audio(x), - inputs=[audio_prompt], - outputs=[audio_prompt], - ) - - formatter.click( - format_utils, - inputs=[text, speaker], - outputs=[language, text], - ) - - print("推理页面已开启!") - webbrowser.open(f"http://127.0.0.1:{config.webui_config.port}") - app.launch(share=config.webui_config.share, server_port=config.webui_config.port) From d36401849b74e155ac0972569d6fa4e040eb6928 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Thu, 7 Mar 2024 04:36:05 +0000 Subject: [PATCH 18/64] Refactor: moved monotonic_align/ to style_bert_vits2/models/monotonic_alignment.py --- monotonic_align/__init__.py | 16 ---- monotonic_align/core.py | 46 ---------- style_bert_vits2/models/infer.py | 8 +- style_bert_vits2/models/models.py | 4 +- style_bert_vits2/models/models_jp_extra.py | 4 +- style_bert_vits2/models/modules.py | 1 + .../models/monotonic_alignment.py | 88 +++++++++++++++++++ 7 files changed, 97 insertions(+), 70 deletions(-) delete mode 100644 monotonic_align/__init__.py delete mode 100644 monotonic_align/core.py create mode 100644 style_bert_vits2/models/monotonic_alignment.py diff --git a/monotonic_align/__init__.py b/monotonic_align/__init__.py deleted file mode 100644 index 15d8e60..0000000 --- a/monotonic_align/__init__.py +++ /dev/null @@ -1,16 +0,0 @@ -from numpy import zeros, int32, float32 -from torch import from_numpy - -from .core import maximum_path_jit - - -def maximum_path(neg_cent, mask): - device = neg_cent.device - dtype = neg_cent.dtype - neg_cent = neg_cent.data.cpu().numpy().astype(float32) - path = zeros(neg_cent.shape, dtype=int32) - - t_t_max = mask.sum(1)[:, 0].data.cpu().numpy().astype(int32) - t_s_max = mask.sum(2)[:, 0].data.cpu().numpy().astype(int32) - maximum_path_jit(path, neg_cent, t_t_max, t_s_max) - return from_numpy(path).to(device=device, dtype=dtype) diff --git a/monotonic_align/core.py b/monotonic_align/core.py deleted file mode 100644 index ffa489d..0000000 --- a/monotonic_align/core.py +++ /dev/null @@ -1,46 +0,0 @@ -import numba - - -@numba.jit( - numba.void( - numba.int32[:, :, ::1], - numba.float32[:, :, ::1], - numba.int32[::1], - numba.int32[::1], - ), - nopython=True, - nogil=True, -) -def maximum_path_jit(paths, values, t_ys, t_xs): - b = paths.shape[0] - max_neg_val = -1e9 - for i in range(int(b)): - path = paths[i] - value = values[i] - t_y = t_ys[i] - t_x = t_xs[i] - - v_prev = v_cur = 0.0 - index = t_x - 1 - - for y in range(t_y): - for x in range(max(0, t_x + y - t_y), min(t_x, y + 1)): - if x == y: - v_cur = max_neg_val - else: - v_cur = value[y - 1, x] - if x == 0: - if y == 0: - v_prev = 0.0 - else: - v_prev = max_neg_val - else: - v_prev = value[y - 1, x - 1] - value[y, x] += max(v_prev, v_cur) - - for y in range(t_y - 1, -1, -1): - path[y, index] = 1 - if index != 0 and ( - index == y or value[y - 1, index] < value[y - 1, index - 1] - ): - index = index - 1 diff --git a/style_bert_vits2/models/infer.py b/style_bert_vits2/models/infer.py index 0e556d0..5eb9ebd 100644 --- a/style_bert_vits2/models/infer.py +++ b/style_bert_vits2/models/infer.py @@ -10,10 +10,6 @@ from style_bert_vits2.text_processing import clean_text, cleaned_text_to_sequenc from style_bert_vits2.text_processing.symbols import SYMBOLS -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") @@ -315,3 +311,7 @@ def infer_multilang( if torch.cuda.is_available(): torch.cuda.empty_cache() return audio + + +class InvalidToneError(ValueError): + pass diff --git a/style_bert_vits2/models/models.py b/style_bert_vits2/models/models.py index a8c6695..7f14be4 100644 --- a/style_bert_vits2/models/models.py +++ b/style_bert_vits2/models/models.py @@ -6,10 +6,10 @@ from torch.nn import Conv1d, Conv2d, ConvTranspose1d from torch.nn import functional as F from torch.nn.utils import remove_weight_norm, spectral_norm, weight_norm -import monotonic_align from style_bert_vits2.models import attentions from style_bert_vits2.models import commons from style_bert_vits2.models import modules +from style_bert_vits2.models import monotonic_alignment from style_bert_vits2.models.commons import get_padding, init_weights from style_bert_vits2.text_processing.symbols import NUM_LANGUAGES, NUM_TONES, SYMBOLS @@ -932,7 +932,7 @@ class SynthesizerTrn(nn.Module): attn_mask = torch.unsqueeze(x_mask, 2) * torch.unsqueeze(y_mask, -1) attn = ( - monotonic_align.maximum_path(neg_cent, attn_mask.squeeze(1)) + monotonic_alignment.maximum_path(neg_cent, attn_mask.squeeze(1)) .unsqueeze(1) .detach() ) diff --git a/style_bert_vits2/models/models_jp_extra.py b/style_bert_vits2/models/models_jp_extra.py index 8c4a4ec..54591f6 100644 --- a/style_bert_vits2/models/models_jp_extra.py +++ b/style_bert_vits2/models/models_jp_extra.py @@ -6,10 +6,10 @@ from torch.nn import Conv1d, Conv2d, ConvTranspose1d from torch.nn import functional as F from torch.nn.utils import remove_weight_norm, spectral_norm, weight_norm -import monotonic_align from style_bert_vits2.models import attentions from style_bert_vits2.models import commons from style_bert_vits2.models import modules +from style_bert_vits2.models import monotonic_alignment from style_bert_vits2.text_processing.symbols import SYMBOLS, NUM_TONES, NUM_LANGUAGES @@ -979,7 +979,7 @@ class SynthesizerTrn(nn.Module): attn_mask = torch.unsqueeze(x_mask, 2) * torch.unsqueeze(y_mask, -1) attn = ( - monotonic_align.maximum_path(neg_cent, attn_mask.squeeze(1)) + monotonic_alignment.maximum_path(neg_cent, attn_mask.squeeze(1)) .unsqueeze(1) .detach() ) diff --git a/style_bert_vits2/models/modules.py b/style_bert_vits2/models/modules.py index e0885c4..c076e21 100644 --- a/style_bert_vits2/models/modules.py +++ b/style_bert_vits2/models/modules.py @@ -10,6 +10,7 @@ from transforms import piecewise_rational_quadratic_transform from style_bert_vits2.models import commons from style_bert_vits2.models.attentions import Encoder + LRELU_SLOPE = 0.1 diff --git a/style_bert_vits2/models/monotonic_alignment.py b/style_bert_vits2/models/monotonic_alignment.py new file mode 100644 index 0000000..0f393c1 --- /dev/null +++ b/style_bert_vits2/models/monotonic_alignment.py @@ -0,0 +1,88 @@ +""" +以下に記述されている関数のコメントはリファクタリング時に GPT-4 に生成させたもので、 +コードと完全に一致している保証はない。あくまで参考程度とすること。 +""" + +import numba +import torch +from numpy import int32, float32, zeros +from typing import Any + + +def maximum_path(neg_cent: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: + """ + 与えられた負の中心とマスクを使用して最大パスを計算する + + Args: + neg_cent (torch.Tensor): 負の中心を表すテンソル + mask (torch.Tensor): マスクを表すテンソル + + Returns: + Tensor: 計算された最大パスを表すテンソル + """ + + device = neg_cent.device + dtype = neg_cent.dtype + neg_cent = neg_cent.data.cpu().numpy().astype(float32) + path = zeros(neg_cent.shape, dtype=int32) + + t_t_max = mask.sum(1)[:, 0].data.cpu().numpy().astype(int32) + t_s_max = mask.sum(2)[:, 0].data.cpu().numpy().astype(int32) + maximum_path_jit(path, neg_cent, t_t_max, t_s_max) + + return torch.from_numpy(path).to(device=device, dtype=dtype) + + +@numba.jit( + numba.void( + numba.int32[:, :, ::1], + numba.float32[:, :, ::1], + numba.int32[::1], + numba.int32[::1], + ), + nopython = True, + nogil = True, +) # type: ignore +def maximum_path_jit(paths: Any, values: Any, t_ys: Any, t_xs: Any) -> None: + """ + 与えられたパス、値、およびターゲットの y と x 座標を使用して JIT で最大パスを計算する + + Args: + paths: 計算されたパスを格納するための整数型の 3 次元配列 + values: 値を格納するための浮動小数点型の 3 次元配列 + t_ys: ターゲットの y 座標を格納するための整数型の 1 次元配列 + t_xs: ターゲットの x 座標を格納するための整数型の 1 次元配列 + """ + + b = paths.shape[0] + max_neg_val = -1e9 + for i in range(int(b)): + path = paths[i] + value = values[i] + t_y = t_ys[i] + t_x = t_xs[i] + + v_prev = v_cur = 0.0 + index = t_x - 1 + + for y in range(t_y): + for x in range(max(0, t_x + y - t_y), min(t_x, y + 1)): + if x == y: + v_cur = max_neg_val + else: + v_cur = value[y - 1, x] + if x == 0: + if y == 0: + v_prev = 0.0 + else: + v_prev = max_neg_val + else: + v_prev = value[y - 1, x - 1] + value[y, x] += max(v_prev, v_cur) + + for y in range(t_y - 1, -1, -1): + path[y, index] = 1 + if index != 0 and ( + index == y or value[y - 1, index] < value[y - 1, index - 1] + ): + index = index - 1 From f8f798d10a63739c9329d386e49657715b26c7b3 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Thu, 7 Mar 2024 04:48:11 +0000 Subject: [PATCH 19/64] Refactor: moved text/ to style_bert_vits2/text_processing/(language)/ --- app.py | 2 +- server_editor.py | 2 +- style_bert_vits2/text_processing/__init__.py | 15 +++++++-------- .../text_processing/chinese/__init__.py | 4 ++-- .../text_processing/chinese}/tone_sandhi.py | 0 .../text_processing/english/__init__.py | 4 ++-- .../text_processing/english}/cmudict.rep | 0 .../english}/cmudict_cache.pickle | Bin .../text_processing/english}/opencpop-strict.txt | 0 .../text_processing/japanese/__init__.py | 2 ++ 10 files changed, 15 insertions(+), 14 deletions(-) rename text/chinese.py => style_bert_vits2/text_processing/chinese/__init__.py (98%) rename {text => style_bert_vits2/text_processing/chinese}/tone_sandhi.py (100%) rename text/english.py => style_bert_vits2/text_processing/english/__init__.py (99%) rename {text => style_bert_vits2/text_processing/english}/cmudict.rep (100%) rename {text => style_bert_vits2/text_processing/english}/cmudict_cache.pickle (100%) rename {text => style_bert_vits2/text_processing/english}/opencpop-strict.txt (100%) create mode 100644 style_bert_vits2/text_processing/japanese/__init__.py diff --git a/app.py b/app.py index 02c01cd..f03e91b 100644 --- a/app.py +++ b/app.py @@ -27,8 +27,8 @@ from style_bert_vits2.constants import ( ) from style_bert_vits2.logging import logger from style_bert_vits2.models.infer import InvalidToneError +from style_bert_vits2.text_processing.japanese import normalize_text from style_bert_vits2.text_processing.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone -from style_bert_vits2.text_processing.japanese.normalizer import normalize_text # Get path settings diff --git a/server_editor.py b/server_editor.py index eb1cffd..a1e249d 100644 --- a/server_editor.py +++ b/server_editor.py @@ -42,8 +42,8 @@ from style_bert_vits2.constants import ( ) from style_bert_vits2.logging import logger from style_bert_vits2.text_processing import bert_models +from style_bert_vits2.text_processing.japanese import normalize_text from style_bert_vits2.text_processing.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone -from style_bert_vits2.text_processing.japanese.normalizer import normalize_text from style_bert_vits2.text_processing.japanese.user_dict import ( apply_word, update_dict, diff --git a/style_bert_vits2/text_processing/__init__.py b/style_bert_vits2/text_processing/__init__.py index cd56ee7..5e77df6 100644 --- a/style_bert_vits2/text_processing/__init__.py +++ b/style_bert_vits2/text_processing/__init__.py @@ -67,18 +67,17 @@ def clean_text( # Changed to import inside if condition to avoid unnecessary import 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 + from style_bert_vits2.text_processing.japanese import g2p, normalize_text norm_text = normalize_text(text) 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) + from style_bert_vits2.text_processing.english import g2p, normalize_text + norm_text = normalize_text(text) + phones, tones, word2ph = g2p(norm_text) 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) + from style_bert_vits2.text_processing.chinese import g2p, normalize_text + norm_text = normalize_text(text) + phones, tones, word2ph = g2p(norm_text) else: raise ValueError(f"Language {language} not supported") diff --git a/text/chinese.py b/style_bert_vits2/text_processing/chinese/__init__.py similarity index 98% rename from text/chinese.py rename to style_bert_vits2/text_processing/chinese/__init__.py index 94266c2..92c8214 100644 --- a/text/chinese.py +++ b/style_bert_vits2/text_processing/chinese/__init__.py @@ -66,7 +66,7 @@ def replace_punctuation(text): return replaced_text -def g2p(text): +def g2p(text: str) -> tuple[list[str], list[int], list[int]]: pattern = r"(?<=[{0}])\s*".format("".join(PUNCTUATIONS)) sentences = [i for i in re.split(pattern, text) if i.strip() != ""] phones, tones, word2ph = _g2p(sentences) @@ -168,7 +168,7 @@ def _g2p(segments): return phones_list, tones_list, word2ph -def normalize_text(text): +def normalize_text(text: str) -> str: numbers = re.findall(r"\d+(?:\.?\d+)?", text) for number in numbers: text = text.replace(number, cn2an.an2cn(number), 1) diff --git a/text/tone_sandhi.py b/style_bert_vits2/text_processing/chinese/tone_sandhi.py similarity index 100% rename from text/tone_sandhi.py rename to style_bert_vits2/text_processing/chinese/tone_sandhi.py diff --git a/text/english.py b/style_bert_vits2/text_processing/english/__init__.py similarity index 99% rename from text/english.py rename to style_bert_vits2/text_processing/english/__init__.py index 419be3b..852431a 100644 --- a/text/english.py +++ b/style_bert_vits2/text_processing/english/__init__.py @@ -369,7 +369,7 @@ def normalize_numbers(text): return text -def normalize_text(text): +def normalize_text(text: str) -> str: text = normalize_numbers(text) text = replace_punctuation(text) text = re.sub(r"([,;.\?\!])([\w])", r"\1 \2", text) @@ -419,7 +419,7 @@ def text_to_words(text): return words -def g2p(text): +def g2p(text: str) -> tuple[list[str], list[int], list[int]]: phones = [] tones = [] phone_len = [] diff --git a/text/cmudict.rep b/style_bert_vits2/text_processing/english/cmudict.rep similarity index 100% rename from text/cmudict.rep rename to style_bert_vits2/text_processing/english/cmudict.rep diff --git a/text/cmudict_cache.pickle b/style_bert_vits2/text_processing/english/cmudict_cache.pickle similarity index 100% rename from text/cmudict_cache.pickle rename to style_bert_vits2/text_processing/english/cmudict_cache.pickle diff --git a/text/opencpop-strict.txt b/style_bert_vits2/text_processing/english/opencpop-strict.txt similarity index 100% rename from text/opencpop-strict.txt rename to style_bert_vits2/text_processing/english/opencpop-strict.txt diff --git a/style_bert_vits2/text_processing/japanese/__init__.py b/style_bert_vits2/text_processing/japanese/__init__.py new file mode 100644 index 0000000..17e1785 --- /dev/null +++ b/style_bert_vits2/text_processing/japanese/__init__.py @@ -0,0 +1,2 @@ +from style_bert_vits2.text_processing.japanese.g2p import g2p # type: ignore +from style_bert_vits2.text_processing.japanese.normalizer import normalize_text # type: ignore From bffd5a67bb93174869db2f1164211970cec67f63 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Thu, 7 Mar 2024 08:34:54 +0000 Subject: [PATCH 20/64] Fix: import error --- .../text_processing/chinese/__init__.py | 6 ++-- .../text_processing/chinese/tone_sandhi.py | 36 +++++++++---------- 2 files changed, 20 insertions(+), 22 deletions(-) diff --git a/style_bert_vits2/text_processing/chinese/__init__.py b/style_bert_vits2/text_processing/chinese/__init__.py index 92c8214..23acbf3 100644 --- a/style_bert_vits2/text_processing/chinese/__init__.py +++ b/style_bert_vits2/text_processing/chinese/__init__.py @@ -2,10 +2,12 @@ import os import re import cn2an +import jieba.posseg as psg from pypinyin import lazy_pinyin, Style +from style_bert_vits2.text_processing.chinese.tone_sandhi import ToneSandhi from style_bert_vits2.text_processing.symbols import PUNCTUATIONS -from text.tone_sandhi import ToneSandhi + current_file_path = os.path.dirname(__file__) pinyin_to_symbol_map = { @@ -13,8 +15,6 @@ pinyin_to_symbol_map = { for line in open(os.path.join(current_file_path, "opencpop-strict.txt")).readlines() } -import jieba.posseg as psg - rep_map = { ":": ",", diff --git a/style_bert_vits2/text_processing/chinese/tone_sandhi.py b/style_bert_vits2/text_processing/chinese/tone_sandhi.py index 38f3137..5832434 100644 --- a/style_bert_vits2/text_processing/chinese/tone_sandhi.py +++ b/style_bert_vits2/text_processing/chinese/tone_sandhi.py @@ -11,8 +11,6 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -from typing import List -from typing import Tuple import jieba from pypinyin import lazy_pinyin @@ -463,7 +461,7 @@ class ToneSandhi: # word: "家里" # pos: "s" # finals: ['ia1', 'i3'] - def _neural_sandhi(self, word: str, pos: str, finals: List[str]) -> List[str]: + def _neural_sandhi(self, word: str, pos: str, finals: list[str]) -> list[str]: # reduplication words for n. and v. e.g. 奶奶, 试试, 旺旺 for j, item in enumerate(word): if ( @@ -522,7 +520,7 @@ class ToneSandhi: finals = sum(finals_list, []) return finals - def _bu_sandhi(self, word: str, finals: List[str]) -> List[str]: + def _bu_sandhi(self, word: str, finals: list[str]) -> list[str]: # e.g. 看不懂 if len(word) == 3 and word[1] == "不": finals[1] = finals[1][:-1] + "5" @@ -533,7 +531,7 @@ class ToneSandhi: finals[i] = finals[i][:-1] + "2" return finals - def _yi_sandhi(self, word: str, finals: List[str]) -> List[str]: + def _yi_sandhi(self, word: str, finals: list[str]) -> list[str]: # "一" in number sequences, e.g. 一零零, 二一零 if word.find("一") != -1 and all( [item.isnumeric() for item in word if item != "一"] @@ -558,9 +556,9 @@ class ToneSandhi: finals[i] = finals[i][:-1] + "4" return finals - def _split_word(self, word: str) -> List[str]: + def _split_word(self, word: str) -> list[str]: word_list = jieba.cut_for_search(word) - word_list = sorted(word_list, key=lambda i: len(i), reverse=False) + word_list = sorted(word_list, key=lambda i: len(i), reverse=False) # type: ignore first_subword = word_list[0] first_begin_idx = word.find(first_subword) if first_begin_idx == 0: @@ -571,7 +569,7 @@ class ToneSandhi: new_word_list = [second_subword, first_subword] return new_word_list - def _three_sandhi(self, word: str, finals: List[str]) -> List[str]: + def _three_sandhi(self, word: str, finals: list[str]) -> list[str]: if len(word) == 2 and self._all_tone_three(finals): finals[0] = finals[0][:-1] + "2" elif len(word) == 3: @@ -611,12 +609,12 @@ class ToneSandhi: return finals - def _all_tone_three(self, finals: List[str]) -> bool: + def _all_tone_three(self, finals: list[str]) -> bool: return all(x[-1] == "3" for x in finals) # merge "不" and the word behind it # if don't merge, "不" sometimes appears alone according to jieba, which may occur sandhi error - def _merge_bu(self, seg: List[Tuple[str, str]]) -> List[Tuple[str, str]]: + def _merge_bu(self, seg: list[tuple[str, str]]) -> list[tuple[str, str]]: new_seg = [] last_word = "" for word, pos in seg: @@ -636,7 +634,7 @@ class ToneSandhi: # e.g. # input seg: [('听', 'v'), ('一', 'm'), ('听', 'v')] # output seg: [['听一听', 'v']] - def _merge_yi(self, seg: List[Tuple[str, str]]) -> List[Tuple[str, str]]: + def _merge_yi(self, seg: list[tuple[str, str]]) -> list[tuple[str, str]]: new_seg = [] * len(seg) # function 1 i = 0 @@ -674,8 +672,8 @@ class ToneSandhi: # the first and the second words are all_tone_three def _merge_continuous_three_tones( - self, seg: List[Tuple[str, str]] - ) -> List[Tuple[str, str]]: + self, seg: list[tuple[str, str]] + ) -> list[tuple[str, str]]: new_seg = [] sub_finals_list = [ lazy_pinyin(word, neutral_tone_with_five=True, style=Style.FINALS_TONE3) @@ -709,8 +707,8 @@ class ToneSandhi: # the last char of first word and the first char of second word is tone_three def _merge_continuous_three_tones_2( - self, seg: List[Tuple[str, str]] - ) -> List[Tuple[str, str]]: + self, seg: list[tuple[str, str]] + ) -> list[tuple[str, str]]: new_seg = [] sub_finals_list = [ lazy_pinyin(word, neutral_tone_with_five=True, style=Style.FINALS_TONE3) @@ -738,7 +736,7 @@ class ToneSandhi: new_seg.append([word, pos]) return new_seg - def _merge_er(self, seg: List[Tuple[str, str]]) -> List[Tuple[str, str]]: + def _merge_er(self, seg: list[tuple[str, str]]) -> list[tuple[str, str]]: new_seg = [] for i, (word, pos) in enumerate(seg): if i - 1 >= 0 and word == "儿" and seg[i - 1][0] != "#": @@ -747,7 +745,7 @@ class ToneSandhi: new_seg.append([word, pos]) return new_seg - def _merge_reduplication(self, seg: List[Tuple[str, str]]) -> List[Tuple[str, str]]: + def _merge_reduplication(self, seg: list[tuple[str, str]]) -> list[tuple[str, str]]: new_seg = [] for i, (word, pos) in enumerate(seg): if new_seg and word == new_seg[-1][0]: @@ -756,7 +754,7 @@ class ToneSandhi: new_seg.append([word, pos]) return new_seg - def pre_merge_for_modify(self, seg: List[Tuple[str, str]]) -> List[Tuple[str, str]]: + def pre_merge_for_modify(self, seg: list[tuple[str, str]]) -> list[tuple[str, str]]: seg = self._merge_bu(seg) try: seg = self._merge_yi(seg) @@ -768,7 +766,7 @@ class ToneSandhi: seg = self._merge_er(seg) return seg - def modified_tone(self, word: str, pos: str, finals: List[str]) -> List[str]: + def modified_tone(self, word: str, pos: str, finals: list[str]) -> list[str]: finals = self._bu_sandhi(word, finals) finals = self._yi_sandhi(word, finals) finals = self._neural_sandhi(word, pos, finals) From e57cfbf072c2a34dcfd22201ba9c8c67632f9122 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Thu, 7 Mar 2024 08:46:13 +0000 Subject: [PATCH 21/64] Remove: remove currently unused code in style_bert_vits2/models/commons.py --- style_bert_vits2/models/commons.py | 126 ------------------ .../text_processing/japanese/g2p.py | 4 +- .../text_processing/japanese/normalizer.py | 46 +++---- 3 files changed, 25 insertions(+), 151 deletions(-) diff --git a/style_bert_vits2/models/commons.py b/style_bert_vits2/models/commons.py index 064ef5f..969ed36 100644 --- a/style_bert_vits2/models/commons.py +++ b/style_bert_vits2/models/commons.py @@ -3,7 +3,6 @@ コードと完全に一致している保証はない。あくまで参考程度とすること。 """ -import math import torch from torch.nn import functional as F from typing import Any @@ -68,54 +67,6 @@ def intersperse(lst: list[Any], item: Any) -> list[Any]: return result -def kl_divergence(m_p: torch.Tensor, logs_p: torch.Tensor, m_q: torch.Tensor, logs_q: torch.Tensor) -> torch.Tensor: - """ - 2つの正規分布間の KL ダイバージェンスを計算する - - Args: - m_p (torch.Tensor): P の平均 - logs_p (torch.Tensor): P の対数標準偏差 - m_q (torch.Tensor): Q の平均 - logs_q (torch.Tensor): Q の対数標準偏差 - - Returns: - torch.Tensor: KL ダイバージェンスの値。 - """ - kl = (logs_q - logs_p) - 0.5 - kl += ( - 0.5 * (torch.exp(2.0 * logs_p) + ((m_p - m_q) ** 2)) * torch.exp(-2.0 * logs_q) - ) - return kl - - -def rand_gumbel(shape: torch.Size) -> torch.Tensor: - """ - Gumbel 分布からサンプリングし、オーバーフローを防ぐ - - Args: - shape (torch.Size): サンプルの形状 - - Returns: - torch.Tensor: Gumbel 分布からのサンプル - """ - uniform_samples = torch.rand(shape) * 0.99998 + 0.00001 - return -torch.log(-torch.log(uniform_samples)) - - -def rand_gumbel_like(x: torch.Tensor) -> torch.Tensor: - """ - 引数と同じ形状のテンソルで、Gumbel 分布からサンプリングする - - Args: - x (torch.Tensor): 形状を基にするテンソル - - Returns: - torch.Tensor: Gumbel 分布からのサンプル - """ - g = rand_gumbel(x.size()).to(dtype=x.dtype, device=x.device) - return g - - def slice_segments(x: torch.Tensor, ids_str: torch.Tensor, segment_size: int = 4) -> torch.Tensor: """ テンソルからセグメントをスライスする @@ -155,69 +106,6 @@ def rand_slice_segments(x: torch.Tensor, x_lengths: torch.Tensor | None = None, return ret, ids_str -def get_timing_signal_1d(length: int, channels: int, min_timescale: float = 1.0, max_timescale: float = 1.0e4) -> torch.Tensor: - """ - 1D タイミング信号を取得する - - Args: - length (int): シグナルの長さ - channels (int): シグナルのチャネル数 - min_timescale (float, optional): 最小のタイムスケール (デフォルト: 1.0) - max_timescale (float, optional): 最大のタイムスケール (デフォルト: 1.0e4) - - Returns: - torch.Tensor: タイミング信号 - """ - position = torch.arange(length, dtype=torch.float) - num_timescales = channels // 2 - log_timescale_increment = math.log(float(max_timescale) / float(min_timescale)) / ( - num_timescales - 1 - ) - inv_timescales = min_timescale * torch.exp( - torch.arange(num_timescales, dtype=torch.float) * -log_timescale_increment - ) - scaled_time = position.unsqueeze(0) * inv_timescales.unsqueeze(1) - signal = torch.cat([torch.sin(scaled_time), torch.cos(scaled_time)], 0) - signal = F.pad(signal, [0, 0, 0, channels % 2]) - signal = signal.view(1, channels, length) - return signal - - -def add_timing_signal_1d(x: torch.Tensor, min_timescale: float = 1.0, max_timescale: float = 1.0e4) -> torch.Tensor: - """ - 1D タイミング信号をテンソルに追加する - - Args: - x (torch.Tensor): 入力テンソル - min_timescale (float, optional): 最小のタイムスケール (デフォルト: 1.0) - max_timescale (float, optional): 最大のタイムスケール (デフォルト: 1.0e4) - - Returns: - torch.Tensor: タイミング信号が追加されたテンソル - """ - b, channels, length = x.size() - signal = get_timing_signal_1d(length, channels, min_timescale, max_timescale) - return x + signal.to(dtype=x.dtype, device=x.device) - - -def cat_timing_signal_1d(x: torch.Tensor, min_timescale: float = 1.0, max_timescale: float = 1.0e4, axis: int = 1) -> torch.Tensor: - """ - 1D タイミング信号をテンソルに連結する - - Args: - x (torch.Tensor): 入力テンソル - min_timescale (float, optional): 最小のタイムスケール (デフォルト: 1.0) - max_timescale (float, optional): 最大のタイムスケール (デフォルト: 1.0e4) - axis (int, optional): 連結する軸 (デフォルト: 1) - - Returns: - torch.Tensor: タイミング信号が連結されたテンソル - """ - b, channels, length = x.size() - signal = get_timing_signal_1d(length, channels, min_timescale, max_timescale) - return torch.cat([x, signal.to(dtype=x.dtype, device=x.device)], axis) - - def subsequent_mask(length: int) -> torch.Tensor: """ 後続のマスクを生成する @@ -253,20 +141,6 @@ def fused_add_tanh_sigmoid_multiply(input_a: torch.Tensor, input_b: torch.Tensor return acts -def shift_1d(x: torch.Tensor) -> torch.Tensor: - """ - 与えられたテンソルを 1D でシフトする - - Args: - x (torch.Tensor): シフトするテンソル - - Returns: - torch.Tensor: シフトされたテンソル - """ - x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [1, 0]]))[:, :, :-1] - return x - - def sequence_mask(length: torch.Tensor, max_length: int | None = None) -> torch.Tensor: """ シーケンスマスクを生成する diff --git a/style_bert_vits2/text_processing/japanese/g2p.py b/style_bert_vits2/text_processing/japanese/g2p.py index 7968751..c04d078 100644 --- a/style_bert_vits2/text_processing/japanese/g2p.py +++ b/style_bert_vits2/text_processing/japanese/g2p.py @@ -171,7 +171,7 @@ def __g2phone_tone_wo_punct(text: str) -> list[tuple[str, int]]: list[tuple[str, int]]: 音素とアクセントのペアのリスト """ - prosodies = pyopenjtalk_g2p_prosody(text, drop_unvoiced_vowels=True) + prosodies = __pyopenjtalk_g2p_prosody(text, drop_unvoiced_vowels=True) # logger.debug(f"prosodies: {prosodies}") result: list[tuple[str, int]] = [] current_phrase: list[tuple[str, int]] = [] @@ -212,7 +212,7 @@ def __g2phone_tone_wo_punct(text: str) -> list[tuple[str, int]]: return result -def pyopenjtalk_g2p_prosody(text: str, drop_unvoiced_vowels: bool = True) -> list[str]: +def __pyopenjtalk_g2p_prosody(text: str, drop_unvoiced_vowels: bool = True) -> list[str]: """ ESPnet の実装から引用、変更点無し。「ん」は「N」なことに注意。 ref: https://github.com/espnet/espnet/blob/master/espnet2/text/phoneme_tokenizer.py diff --git a/style_bert_vits2/text_processing/japanese/normalizer.py b/style_bert_vits2/text_processing/japanese/normalizer.py index 92c2e87..8276338 100644 --- a/style_bert_vits2/text_processing/japanese/normalizer.py +++ b/style_bert_vits2/text_processing/japanese/normalizer.py @@ -48,29 +48,6 @@ def normalize_text(text: str) -> str: return res -def __convert_numbers_to_words(text: str) -> str: - """ - 記号や数字を日本語の文字表現に変換する。 - - Args: - text (str): 変換するテキスト - - Returns: - str: 変換されたテキスト - """ - - NUMBER_WITH_SEPARATOR_PATTERN = re.compile("[0-9]{1,3}(,[0-9]{3})+") - CURRENCY_MAP = {"$": "ドル", "¥": "円", "£": "ポンド", "€": "ユーロ"} - CURRENCY_PATTERN = re.compile(r"([$¥£€])([0-9.]*[0-9])") - NUMBER_PATTERN = re.compile(r"[0-9]+(\.[0-9]+)?") - - res = NUMBER_WITH_SEPARATOR_PATTERN.sub(lambda m: m[0].replace(",", ""), text) - res = CURRENCY_PATTERN.sub(lambda m: m[2] + CURRENCY_MAP.get(m[1], m[1]), res) - res = NUMBER_PATTERN.sub(lambda m: num2words(m[0], lang="ja"), res) - - return res - - def replace_punctuation(text: str) -> str: """ 句読点等を「.」「,」「!」「?」「'」「-」に正規化し、OpenJTalk で読みが取得できるもののみ残す: @@ -159,3 +136,26 @@ def replace_punctuation(text: str) -> str: ) return replaced_text + + +def __convert_numbers_to_words(text: str) -> str: + """ + 記号や数字を日本語の文字表現に変換する。 + + Args: + text (str): 変換するテキスト + + Returns: + str: 変換されたテキスト + """ + + NUMBER_WITH_SEPARATOR_PATTERN = re.compile("[0-9]{1,3}(,[0-9]{3})+") + CURRENCY_MAP = {"$": "ドル", "¥": "円", "£": "ポンド", "€": "ユーロ"} + CURRENCY_PATTERN = re.compile(r"([$¥£€])([0-9.]*[0-9])") + NUMBER_PATTERN = re.compile(r"[0-9]+(\.[0-9]+)?") + + res = NUMBER_WITH_SEPARATOR_PATTERN.sub(lambda m: m[0].replace(",", ""), text) + res = CURRENCY_PATTERN.sub(lambda m: m[2] + CURRENCY_MAP.get(m[1], m[1]), res) + res = NUMBER_PATTERN.sub(lambda m: num2words(m[0], lang="ja"), res) + + return res From 70f8d53a1e41f83ef33f20062bd45b4ff193cb64 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Thu, 7 Mar 2024 09:09:24 +0000 Subject: [PATCH 22/64] Add: empty __init__.py --- style_bert_vits2/models/__init__.py | 0 style_bert_vits2/utils/__init__.py | 0 2 files changed, 0 insertions(+), 0 deletions(-) create mode 100644 style_bert_vits2/models/__init__.py create mode 100644 style_bert_vits2/utils/__init__.py diff --git a/style_bert_vits2/models/__init__.py b/style_bert_vits2/models/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/style_bert_vits2/utils/__init__.py b/style_bert_vits2/utils/__init__.py new file mode 100644 index 0000000..e69de29 From 3f07c256e300b4111201fee7fb2b3b0ddd538f04 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Thu, 7 Mar 2024 19:33:21 +0000 Subject: [PATCH 23/64] Refactor: make variables private that are not used externally --- default_style.py | 2 +- gen_yaml.py | 3 +- .../models/monotonic_alignment.py | 4 +- style_bert_vits2/text_processing/__init__.py | 4 +- .../text_processing/bert_models.py | 72 ++++++++++++++++--- .../text_processing/chinese/bert_feature.py | 10 +-- .../text_processing/english/bert_feature.py | 10 +-- .../text_processing/japanese/bert_feature.py | 10 +-- 8 files changed, 85 insertions(+), 30 deletions(-) diff --git a/default_style.py b/default_style.py index 763e291..67b6fc3 100644 --- a/default_style.py +++ b/default_style.py @@ -1,6 +1,6 @@ import os -from style_bert_vits2.logging import logger from style_bert_vits2.constants import DEFAULT_STYLE +from style_bert_vits2.logging import logger import numpy as np import json diff --git a/gen_yaml.py b/gen_yaml.py index 91301ac..76df206 100644 --- a/gen_yaml.py +++ b/gen_yaml.py @@ -1,7 +1,8 @@ +import argparse import os import shutil import yaml -import argparse + parser = argparse.ArgumentParser( description="config.ymlの生成。あらかじめ前準備をしたデータをバッチファイルなどで連続で学習する時にtrain_ms.pyより前に使用する。" diff --git a/style_bert_vits2/models/monotonic_alignment.py b/style_bert_vits2/models/monotonic_alignment.py index 0f393c1..b499ad0 100644 --- a/style_bert_vits2/models/monotonic_alignment.py +++ b/style_bert_vits2/models/monotonic_alignment.py @@ -28,7 +28,7 @@ def maximum_path(neg_cent: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: t_t_max = mask.sum(1)[:, 0].data.cpu().numpy().astype(int32) t_s_max = mask.sum(2)[:, 0].data.cpu().numpy().astype(int32) - maximum_path_jit(path, neg_cent, t_t_max, t_s_max) + __maximum_path_jit(path, neg_cent, t_t_max, t_s_max) return torch.from_numpy(path).to(device=device, dtype=dtype) @@ -43,7 +43,7 @@ def maximum_path(neg_cent: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: nopython = True, nogil = True, ) # type: ignore -def maximum_path_jit(paths: Any, values: Any, t_ys: Any, t_xs: Any) -> None: +def __maximum_path_jit(paths: Any, values: Any, t_ys: Any, t_xs: Any) -> None: """ 与えられたパス、値、およびターゲットの y と x 座標を使用して JIT で最大パスを計算する diff --git a/style_bert_vits2/text_processing/__init__.py b/style_bert_vits2/text_processing/__init__.py index 5e77df6..4719525 100644 --- a/style_bert_vits2/text_processing/__init__.py +++ b/style_bert_vits2/text_processing/__init__.py @@ -8,7 +8,7 @@ from style_bert_vits2.text_processing.symbols import ( ) -_symbol_to_id = {s: i for i, s in enumerate(SYMBOLS)} +__symbol_to_id = {s: i for i, s in enumerate(SYMBOLS)} def extract_bert_feature( @@ -97,7 +97,7 @@ def cleaned_text_to_sequence(cleaned_phones: list[str], tones: list[int], langua tuple[list[int], list[int], list[int]]: List of integers corresponding to the symbols in the text """ - phones = [_symbol_to_id[symbol] for symbol in cleaned_phones] + phones = [__symbol_to_id[symbol] for symbol in cleaned_phones] tone_start = LANGUAGE_TONE_START_MAP[language] tones = [i + tone_start for i in tones] lang_id = LANGUAGE_ID_MAP[language] diff --git a/style_bert_vits2/text_processing/bert_models.py b/style_bert_vits2/text_processing/bert_models.py index 9132082..e8ef4b4 100644 --- a/style_bert_vits2/text_processing/bert_models.py +++ b/style_bert_vits2/text_processing/bert_models.py @@ -5,11 +5,13 @@ Style-Bert-VITS2 の学習・推論に必要な各言語ごとの BERT モデル 場合によっては多重にロードされて非効率なほか、BERT モデルのロード元のパスがハードコードされているためライブラリ化ができない。 そこで、ライブラリの利用前に、音声合成に利用する言語の BERT モデルだけを「明示的に」ロードできるようにした。 -一度 load_tokenizer() で当該言語の BERT モデルがロードされていれば、ライブラリ内部のどこからでもロード済みのモデル/トークナイザーを取得できる。 +一度 load_model/tokenizer() で当該言語の BERT モデルがロードされていれば、ライブラリ内部のどこからでもロード済みのモデル/トークナイザーを取得できる。 """ +import gc from typing import cast +import torch from transformers import ( AutoModelForMaskedLM, AutoTokenizer, @@ -25,10 +27,10 @@ from style_bert_vits2.logging import logger # 各言語ごとのロード済みの BERT モデルを格納する辞書 -loaded_models: dict[Languages, PreTrainedModel | DebertaV2Model] = {} +__loaded_models: dict[Languages, PreTrainedModel | DebertaV2Model] = {} # 各言語ごとのロード済みの BERT トークナイザーを格納する辞書 -loaded_tokenizers: dict[Languages, PreTrainedTokenizer | PreTrainedTokenizerFast | DebertaV2Tokenizer] = {} +__loaded_tokenizers: dict[Languages, PreTrainedTokenizer | PreTrainedTokenizerFast | DebertaV2Tokenizer] = {} def load_model( @@ -56,8 +58,8 @@ def load_model( """ # すでにロード済みの場合はそのまま返す - if language in loaded_models: - return loaded_models[language] + if language in __loaded_models: + return __loaded_models[language] # pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用 if pretrained_model_name_or_path is None: @@ -71,7 +73,7 @@ def load_model( 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 + __loaded_models[language] = model logger.info(f"Loaded the {language} BERT model from {pretrained_model_name_or_path}") return model @@ -102,8 +104,8 @@ def load_tokenizer( """ # すでにロード済みの場合はそのまま返す - if language in loaded_tokenizers: - return loaded_tokenizers[language] + if language in __loaded_tokenizers: + return __loaded_tokenizers[language] # pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用 if pretrained_model_name_or_path is None: @@ -117,7 +119,59 @@ def load_tokenizer( tokenizer = DebertaV2Tokenizer.from_pretrained(pretrained_model_name_or_path) else: tokenizer = AutoTokenizer.from_pretrained(pretrained_model_name_or_path) - loaded_tokenizers[language] = tokenizer + __loaded_tokenizers[language] = tokenizer logger.info(f"Loaded the {language} BERT tokenizer from {pretrained_model_name_or_path}") return tokenizer + + +def unload_model(language: Languages) -> None: + """ + 指定された言語の BERT モデルをアンロードする + + Args: + language (Languages): アンロードする BERT モデルの言語 + """ + + if language in __loaded_models: + del __loaded_models[language] + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + logger.info(f"Unloaded the {language} BERT model") + + +def unload_tokenizer(language: Languages) -> None: + """ + 指定された言語の BERT トークナイザーをアンロードする + + Args: + language (Languages): アンロードする BERT トークナイザーの言語 + """ + + if language in __loaded_tokenizers: + del __loaded_tokenizers[language] + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + logger.info(f"Unloaded the {language} BERT tokenizer") + + +def unload_all_models() -> None: + """ + すべての BERT モデルをアンロードする + """ + + for language in list(__loaded_models.keys()): + unload_model(language) + logger.info("Unloaded all BERT models") + + +def unload_all_tokenizers() -> None: + """ + すべての BERT トークナイザーをアンロードする + """ + + for language in list(__loaded_tokenizers.keys()): + unload_tokenizer(language) + logger.info("Unloaded all BERT tokenizers") diff --git a/style_bert_vits2/text_processing/chinese/bert_feature.py b/style_bert_vits2/text_processing/chinese/bert_feature.py index 25024cb..3178565 100644 --- a/style_bert_vits2/text_processing/chinese/bert_feature.py +++ b/style_bert_vits2/text_processing/chinese/bert_feature.py @@ -7,7 +7,7 @@ from style_bert_vits2.constants import Languages from style_bert_vits2.text_processing import bert_models -models: dict[torch.device | str, PreTrainedModel] = {} +__models: dict[torch.device | str, PreTrainedModel] = {} def extract_bert_feature( @@ -41,8 +41,8 @@ def extract_bert_feature( device = "cuda" if device == "cuda" and not torch.cuda.is_available(): device = "cpu" - if device not in models.keys(): - models[device] = bert_models.load_model(Languages.ZH).to(device) # type: ignore + if device not in __models.keys(): + __models[device] = bert_models.load_model(Languages.ZH).to(device) # type: ignore style_res_mean = None with torch.no_grad(): @@ -50,13 +50,13 @@ def extract_bert_feature( inputs = tokenizer(text, return_tensors="pt") for i in inputs: inputs[i] = inputs[i].to(device) # type: ignore - res = models[device](**inputs, output_hidden_states=True) + res = __models[device](**inputs, output_hidden_states=True) res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu() if assist_text: style_inputs = tokenizer(assist_text, return_tensors="pt") for i in style_inputs: style_inputs[i] = style_inputs[i].to(device) # type: ignore - style_res = models[device](**style_inputs, output_hidden_states=True) + style_res = __models[device](**style_inputs, output_hidden_states=True) style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu() style_res_mean = style_res.mean(0) diff --git a/style_bert_vits2/text_processing/english/bert_feature.py b/style_bert_vits2/text_processing/english/bert_feature.py index ec556c2..b29d531 100644 --- a/style_bert_vits2/text_processing/english/bert_feature.py +++ b/style_bert_vits2/text_processing/english/bert_feature.py @@ -7,7 +7,7 @@ from style_bert_vits2.constants import Languages from style_bert_vits2.text_processing import bert_models -models: dict[torch.device | str, PreTrainedModel] = {} +__models: dict[torch.device | str, PreTrainedModel] = {} def extract_bert_feature( @@ -41,8 +41,8 @@ def extract_bert_feature( device = "cuda" if device == "cuda" and not torch.cuda.is_available(): device = "cpu" - if device not in models.keys(): - models[device] = bert_models.load_model(Languages.EN).to(device) # type: ignore + if device not in __models.keys(): + __models[device] = bert_models.load_model(Languages.EN).to(device) # type: ignore style_res_mean = None with torch.no_grad(): @@ -50,13 +50,13 @@ def extract_bert_feature( inputs = tokenizer(text, return_tensors="pt") for i in inputs: inputs[i] = inputs[i].to(device) # type: ignore - res = models[device](**inputs, output_hidden_states=True) + res = __models[device](**inputs, output_hidden_states=True) res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu() if assist_text: style_inputs = tokenizer(assist_text, return_tensors="pt") for i in style_inputs: style_inputs[i] = style_inputs[i].to(device) # type: ignore - style_res = models[device](**style_inputs, output_hidden_states=True) + style_res = __models[device](**style_inputs, output_hidden_states=True) style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu() style_res_mean = style_res.mean(0) diff --git a/style_bert_vits2/text_processing/japanese/bert_feature.py b/style_bert_vits2/text_processing/japanese/bert_feature.py index 3ff9d7b..d1809fe 100644 --- a/style_bert_vits2/text_processing/japanese/bert_feature.py +++ b/style_bert_vits2/text_processing/japanese/bert_feature.py @@ -8,7 +8,7 @@ from style_bert_vits2.text_processing import bert_models from style_bert_vits2.text_processing.japanese.g2p import text_to_sep_kata -models: dict[torch.device | str, PreTrainedModel] = {} +__models: dict[torch.device | str, PreTrainedModel] = {} def extract_bert_feature( @@ -48,8 +48,8 @@ def extract_bert_feature( device = "cuda" if device == "cuda" and not torch.cuda.is_available(): device = "cpu" - if device not in models.keys(): - models[device] = bert_models.load_model(Languages.JP).to(device) # type: ignore + if device not in __models.keys(): + __models[device] = bert_models.load_model(Languages.JP).to(device) # type: ignore style_res_mean = None with torch.no_grad(): @@ -57,13 +57,13 @@ def extract_bert_feature( inputs = tokenizer(text, return_tensors="pt") for i in inputs: inputs[i] = inputs[i].to(device) # type: ignore - res = models[device](**inputs, output_hidden_states=True) + res = __models[device](**inputs, output_hidden_states=True) res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu() if assist_text: style_inputs = tokenizer(assist_text, return_tensors="pt") for i in style_inputs: style_inputs[i] = style_inputs[i].to(device) # type: ignore - style_res = models[device](**style_inputs, output_hidden_states=True) + style_res = __models[device](**style_inputs, output_hidden_states=True) style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu() style_res_mean = style_res.mean(0) From 4d5c537f959a18ee978a0ce7cc9ebaea5e754057 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Thu, 7 Mar 2024 19:52:34 +0000 Subject: [PATCH 24/64] Refactor: introducing Ruff --- losses.py | 2 -- pyproject.toml | 9 +++++++++ server_fastapi.py | 4 ++-- .../text_processing/english/__init__.py | 1 - .../text_processing/japanese/__init__.py | 4 ++-- train_ms_jp_extra.py | 1 - webui_style_vectors.py | 6 +++--- webui_train.py | 20 +++++++++---------- 8 files changed, 26 insertions(+), 21 deletions(-) create mode 100644 pyproject.toml diff --git a/losses.py b/losses.py index 4a890ba..9bb50af 100644 --- a/losses.py +++ b/losses.py @@ -2,8 +2,6 @@ import torch import torchaudio from transformers import AutoModel -from style_bert_vits2.logging import logger - def feature_loss(fmap_r, fmap_g): loss = 0 diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..0249363 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,9 @@ +[tool.ruff] +# インデント幅を 4 に設定 +indent-width = 4 + +# 行の長さを 100 文字に設定 +line-length = 100 + +# Python 3.10 向けにフォーマット +target-version = "py310" diff --git a/server_fastapi.py b/server_fastapi.py index ca9520c..b7ebb77 100644 --- a/server_fastapi.py +++ b/server_fastapi.py @@ -104,7 +104,7 @@ if __name__ == "__main__": @app.get("/voice", response_class=AudioResponse) async def voice( request: Request, - text: str = Query(..., min_length=1, max_length=limit, description=f"セリフ"), + text: str = Query(..., min_length=1, max_length=limit, description="セリフ"), encoding: str = Query(None, description="textをURLデコードする(ex, `utf-8`)"), model_id: int = Query( 0, description="モデルID。`GET /models/info`のkeyの値を指定ください" @@ -132,7 +132,7 @@ if __name__ == "__main__": DEFAULT_LENGTH, description="話速。基準は1で大きくするほど音声は長くなり読み上げが遅まる", ), - language: Languages = Query(ln, description=f"textの言語"), + language: Languages = Query(ln, description="textの言語"), auto_split: bool = Query(DEFAULT_LINE_SPLIT, description="改行で分けて生成"), split_interval: float = Query( DEFAULT_SPLIT_INTERVAL, description="分けた場合に挟む無音の長さ(秒)" diff --git a/style_bert_vits2/text_processing/english/__init__.py b/style_bert_vits2/text_processing/english/__init__.py index 852431a..b57e250 100644 --- a/style_bert_vits2/text_processing/english/__init__.py +++ b/style_bert_vits2/text_processing/english/__init__.py @@ -234,7 +234,6 @@ def refine_syllables(syllables): return phonemes, tones -import re import inflect _inflect = inflect.engine() diff --git a/style_bert_vits2/text_processing/japanese/__init__.py b/style_bert_vits2/text_processing/japanese/__init__.py index 17e1785..9123377 100644 --- a/style_bert_vits2/text_processing/japanese/__init__.py +++ b/style_bert_vits2/text_processing/japanese/__init__.py @@ -1,2 +1,2 @@ -from style_bert_vits2.text_processing.japanese.g2p import g2p # type: ignore -from style_bert_vits2.text_processing.japanese.normalizer import normalize_text # type: ignore +from style_bert_vits2.text_processing.japanese.g2p import g2p # noqa: F401 +from style_bert_vits2.text_processing.japanese.normalizer import normalize_text # noqa: F401 diff --git a/train_ms_jp_extra.py b/train_ms_jp_extra.py index 1a287d9..f04d2ae 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -1,6 +1,5 @@ import argparse import datetime -import gc import os import platform diff --git a/webui_style_vectors.py b/webui_style_vectors.py index 1cbabbd..b149e2b 100644 --- a/webui_style_vectors.py +++ b/webui_style_vectors.py @@ -153,7 +153,7 @@ def do_dbscan_gradio(eps=2.5, min_samples=15): return [ plt, gr.Slider(maximum=MAX_CLUSTER_NUM), - f"クラスタが数が0です。パラメータを変えてみてください。", + "クラスタが数が0です。パラメータを変えてみてください。", ] + [gr.Audio(visible=False)] * MAX_AUDIO_NUM return [plt, gr.Slider(maximum=n_clusters, value=1), n_clusters] + [ @@ -212,7 +212,7 @@ def save_style_vectors_from_clustering(model_name, style_names_str: str): if len(style_name_list) != len(centroids) + 1: return f"スタイルの数が合いません。`,`で正しく{len(centroids)}個に区切られているか確認してください: {style_names_str}" if len(set(style_names)) != len(style_names): - return f"スタイル名が重複しています。" + return "スタイル名が重複しています。" logger.info(f"Backup {config_path} to {config_path}.bak") shutil.copy(config_path, f"{config_path}.bak") @@ -243,7 +243,7 @@ def save_style_vectors_from_files( return f"音声ファイルとスタイル名の数が合いません。`,`で正しく{len(style_names)}個に区切られているか確認してください: {audio_files_str}と{style_names_str}" style_name_list = [DEFAULT_STYLE] + style_names if len(set(style_names)) != len(style_names): - return f"スタイル名が重複しています。" + return "スタイル名が重複しています。" style_vectors = [mean] wavs_dir = os.path.join(dataset_root, model_name, "wavs") diff --git a/webui_train.py b/webui_train.py index fda31ca..0a7000c 100644 --- a/webui_train.py +++ b/webui_train.py @@ -147,10 +147,10 @@ def resample(model_name, normalize, trim, num_processes): cmd.append("--trim") success, message = run_script_with_log(cmd) if not success: - logger.error(f"Step 2: resampling failed.") + logger.error("Step 2: resampling failed.") return False, f"Step 2, Error: 音声ファイルの前処理に失敗しました:\n{message}" elif message: - logger.warning(f"Step 2: resampling finished with stderr.") + logger.warning("Step 2: resampling finished with stderr.") return True, f"Step 2, Success: 音声ファイルの前処理が完了しました:\n{message}" logger.success("Step 2: resampling finished.") return True, "Step 2, Success: 音声ファイルの前処理が完了しました" @@ -197,13 +197,13 @@ def preprocess_text(model_name, use_jp_extra, val_per_lang, yomi_error): cmd.append("--use_jp_extra") success, message = run_script_with_log(cmd) if not success: - logger.error(f"Step 3: preprocessing text failed.") + logger.error("Step 3: preprocessing text failed.") return ( False, f"Step 3, Error: 書き起こしファイルの前処理に失敗しました:\n{message}", ) elif message: - logger.warning(f"Step 3: preprocessing text finished with stderr.") + logger.warning("Step 3: preprocessing text finished with stderr.") return ( True, f"Step 3, Success: 書き起こしファイルの前処理が完了しました:\n{message}", @@ -225,10 +225,10 @@ def bert_gen(model_name): ] ) if not success: - logger.error(f"Step 4: bert_gen failed.") + logger.error("Step 4: bert_gen failed.") return False, f"Step 4, Error: BERT特徴ファイルの生成に失敗しました:\n{message}" elif message: - logger.warning(f"Step 4: bert_gen finished with stderr.") + logger.warning("Step 4: bert_gen finished with stderr.") return ( True, f"Step 4, Success: BERT特徴ファイルの生成が完了しました:\n{message}", @@ -250,13 +250,13 @@ def style_gen(model_name, num_processes): ] ) if not success: - logger.error(f"Step 5: style_gen failed.") + logger.error("Step 5: style_gen failed.") return ( False, f"Step 5, Error: スタイル特徴ファイルの生成に失敗しました:\n{message}", ) elif message: - logger.warning(f"Step 5: style_gen finished with stderr.") + logger.warning("Step 5: style_gen finished with stderr.") return ( True, f"Step 5, Success: スタイル特徴ファイルの生成が完了しました:\n{message}", @@ -350,10 +350,10 @@ def train(model_name, skip_style=False, use_jp_extra=True, speedup=False): cmd.append("--speedup") success, message = run_script_with_log(cmd, ignore_warning=True) if not success: - logger.error(f"Train failed.") + logger.error("Train failed.") return False, f"Error: 学習に失敗しました:\n{message}" elif message: - logger.warning(f"Train finished with stderr.") + logger.warning("Train finished with stderr.") return True, f"Success: 学習が完了しました:\n{message}" logger.success("Train finished.") return True, "Success: 学習が完了しました" From 4a3519c4b934b820ddb46f451b830c461c04962a Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 05:51:54 +0000 Subject: [PATCH 25/64] Remove: Ruff I have determined that this is excessive for this project at this time. --- pyproject.toml | 9 --------- 1 file changed, 9 deletions(-) delete mode 100644 pyproject.toml diff --git a/pyproject.toml b/pyproject.toml deleted file mode 100644 index 0249363..0000000 --- a/pyproject.toml +++ /dev/null @@ -1,9 +0,0 @@ -[tool.ruff] -# インデント幅を 4 に設定 -indent-width = 4 - -# 行の長さを 100 文字に設定 -line-length = 100 - -# Python 3.10 向けにフォーマット -target-version = "py310" From 8add1b42023f23b82efb69d981dd9dbe7fbb14bf Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 06:14:50 +0000 Subject: [PATCH 26/64] Fix: maintain compatibility with Python 3.9 --- app.py | 4 +- server_editor.py | 7 ++-- style_bert_vits2/constants.py | 2 +- style_bert_vits2/models/commons.py | 14 +++---- style_bert_vits2/models/infer.py | 9 ++-- style_bert_vits2/text_processing/__init__.py | 5 ++- .../text_processing/bert_models.py | 42 +++++++++---------- .../text_processing/chinese/bert_feature.py | 5 ++- .../text_processing/english/bert_feature.py | 5 ++- .../text_processing/japanese/bert_feature.py | 5 ++- 10 files changed, 52 insertions(+), 46 deletions(-) diff --git a/app.py b/app.py index 56f0cd8..13929cd 100644 --- a/app.py +++ b/app.py @@ -5,7 +5,7 @@ import gradio as gr import torch import yaml -from style_bert_vits2.constants import GRADIO_THEME, VERSION +from style_bert_vits2.constants import GRADIO_THEME, LATEST_VERSION from common.tts_model import ModelHolder from webui import ( create_dataset_app, @@ -34,7 +34,7 @@ if device == "cuda" and not torch.cuda.is_available(): model_holder = ModelHolder(Path(assets_root), device) with gr.Blocks(theme=GRADIO_THEME) as app: - gr.Markdown(f"# Style-Bert-VITS2 WebUI (version {VERSION})") + gr.Markdown(f"# Style-Bert-VITS2 WebUI (version {LATEST_VERSION})") with gr.Tabs(): with gr.Tab("音声合成"): create_inference_app(model_holder=model_holder) diff --git a/server_editor.py b/server_editor.py index 73f76e8..b74d74f 100644 --- a/server_editor.py +++ b/server_editor.py @@ -16,6 +16,7 @@ import zipfile from datetime import datetime from io import BytesIO from pathlib import Path +from typing import Optional import numpy as np import requests @@ -37,7 +38,7 @@ from style_bert_vits2.constants import ( DEFAULT_SDP_RATIO, DEFAULT_STYLE, DEFAULT_STYLE_WEIGHT, - VERSION, + LATEST_VERSION, Languages, ) from style_bert_vits2.logging import logger @@ -212,7 +213,7 @@ router = APIRouter() @router.get("/version") def version() -> str: - return VERSION + return LATEST_VERSION class MoraTone(BaseModel): @@ -265,7 +266,7 @@ class SynthesisRequest(BaseModel): silenceAfter: float = 0.5 pitchScale: float = 1.0 intonationScale: float = 1.0 - speaker: str | None = None + speaker: Optional[str] = None @router.post("/synthesis", response_class=AudioResponse) diff --git a/style_bert_vits2/constants.py b/style_bert_vits2/constants.py index 7b83595..064221e 100644 --- a/style_bert_vits2/constants.py +++ b/style_bert_vits2/constants.py @@ -3,7 +3,7 @@ from pathlib import Path # Style-Bert-VITS2 のバージョン -VERSION = "2.4" +LATEST_VERSION = "2.4" # Style-Bert-VITS2 のベースディレクトリ BASE_DIR = Path(__file__).parent.parent diff --git a/style_bert_vits2/models/commons.py b/style_bert_vits2/models/commons.py index 969ed36..89d07d5 100644 --- a/style_bert_vits2/models/commons.py +++ b/style_bert_vits2/models/commons.py @@ -5,7 +5,7 @@ import torch from torch.nn import functional as F -from typing import Any +from typing import Any, Optional def init_weights(m: torch.nn.Module, mean: float = 0.0, std: float = 0.01) -> None: @@ -85,13 +85,13 @@ def slice_segments(x: torch.Tensor, ids_str: torch.Tensor, segment_size: int = 4 return torch.gather(x, 2, gather_indices) -def rand_slice_segments(x: torch.Tensor, x_lengths: torch.Tensor | None = None, segment_size: int = 4) -> tuple[torch.Tensor, torch.Tensor]: +def rand_slice_segments(x: torch.Tensor, x_lengths: Optional[torch.Tensor] = None, segment_size: int = 4) -> tuple[torch.Tensor, torch.Tensor]: """ ランダムなセグメントをスライスする Args: x (torch.Tensor): 入力テンソル - x_lengths (torch.Tensor, optional): 各バッチの長さ (デフォルト: None) + x_lengths (Optional[torch.Tensor], optional): 各バッチの長さ (デフォルト: None) segment_size (int, optional): スライスのサイズ (デフォルト: 4) Returns: @@ -141,13 +141,13 @@ def fused_add_tanh_sigmoid_multiply(input_a: torch.Tensor, input_b: torch.Tensor return acts -def sequence_mask(length: torch.Tensor, max_length: int | None = None) -> torch.Tensor: +def sequence_mask(length: torch.Tensor, max_length: Optional[int] = None) -> torch.Tensor: """ シーケンスマスクを生成する Args: length (torch.Tensor): 各シーケンスの長さ - max_length (int | None): 最大のシーケンス長さ。指定されていない場合は length の最大値を使用 + max_length (Optional[int]): 最大のシーケンス長さ。指定されていない場合は length の最大値を使用 Returns: torch.Tensor: 生成されたシーケンスマスク @@ -180,13 +180,13 @@ def generate_path(duration: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: return path -def clip_grad_value_(parameters: torch.Tensor | list[torch.Tensor], clip_value: float | None, norm_type: float = 2.0) -> float: +def clip_grad_value_(parameters: torch.Tensor | list[torch.Tensor], clip_value: Optional[float], norm_type: float = 2.0) -> float: """ 勾配の値をクリップする Args: parameters (torch.Tensor | list[torch.Tensor]): クリップするパラメータ - clip_value (float | None): クリップする値。None の場合はクリップしない + clip_value (Optional[float]): クリップする値。None の場合はクリップしない norm_type (float): ノルムの種類 Returns: diff --git a/style_bert_vits2/models/infer.py b/style_bert_vits2/models/infer.py index 5eb9ebd..a5ce70f 100644 --- a/style_bert_vits2/models/infer.py +++ b/style_bert_vits2/models/infer.py @@ -1,4 +1,5 @@ import torch +from typing import Optional import utils from style_bert_vits2.constants import Languages @@ -45,9 +46,9 @@ def get_text( language_str: Languages, hps, device: str, - assist_text: str | None = None, + assist_text: Optional[str] = None, assist_text_weight: float = 0.7, - given_tone: list[int] | None = None, + given_tone: Optional[list[int]] = None, ): use_jp_extra = hps.version.endswith("JP-Extra") # 推論時のみ呼び出されるので、raise_yomi_error は False に設定 @@ -122,9 +123,9 @@ def infer( device: str, skip_start: bool = False, skip_end: bool = False, - assist_text: str | None = None, + assist_text: Optional[str] = None, assist_text_weight: float = 0.7, - given_tone: list[int] | None = None, + given_tone: Optional[list[int]] = None, ): is_jp_extra = hps.version.endswith("JP-Extra") bert, ja_bert, en_bert, phones, tones, lang_ids = get_text( diff --git a/style_bert_vits2/text_processing/__init__.py b/style_bert_vits2/text_processing/__init__.py index 4719525..4dadaed 100644 --- a/style_bert_vits2/text_processing/__init__.py +++ b/style_bert_vits2/text_processing/__init__.py @@ -1,4 +1,5 @@ import torch +from typing import Optional from style_bert_vits2.constants import Languages from style_bert_vits2.text_processing.symbols import ( @@ -16,7 +17,7 @@ def extract_bert_feature( word2ph: list[int], language: Languages, device: torch.device | str, - assist_text: str | None = None, + assist_text: Optional[str] = None, assist_text_weight: float = 0.7, ) -> torch.Tensor: """ @@ -27,7 +28,7 @@ def extract_bert_feature( word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト language (Languages): テキストの言語 device (torch.device | str): 推論に利用するデバイス - assist_text (str | None, optional): 補助テキスト (デフォルト: None) + assist_text (Optional[str], optional): 補助テキスト (デフォルト: None) assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7) Returns: diff --git a/style_bert_vits2/text_processing/bert_models.py b/style_bert_vits2/text_processing/bert_models.py index e8ef4b4..df7a016 100644 --- a/style_bert_vits2/text_processing/bert_models.py +++ b/style_bert_vits2/text_processing/bert_models.py @@ -9,7 +9,7 @@ Style-Bert-VITS2 の学習・推論に必要な各言語ごとの BERT モデル """ import gc -from typing import cast +from typing import cast, Optional import torch from transformers import ( @@ -35,23 +35,23 @@ __loaded_tokenizers: dict[Languages, PreTrainedTokenizer | PreTrainedTokenizerFa def load_model( language: Languages, - pretrained_model_name_or_path: str | None = None, + pretrained_model_name_or_path: Optional[str] = None, ) -> PreTrainedModel | DebertaV2Model: """ - 指定された言語の BERT モデルをロードし、ロード済みの BERT モデルを返す - 一度ロードされていれば、ロード済みの BERT モデルを即座に返す - ライブラリ利用時は常に pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある - ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき + 指定された言語の BERT モデルをロードし、ロード済みの BERT モデルを返す。 + 一度ロードされていれば、ロード済みの BERT モデルを即座に返す。 + ライブラリ利用時は常に pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある。 + ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき。 - Style-Bert-VITS2 では、BERT モデルに下記の 3 つが利用されている - これ以外の BERT モデルを指定した場合は正常に動作しない可能性が高い + 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) + pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None) Returns: PreTrainedModel | DebertaV2Model: ロード済みの BERT モデル @@ -81,23 +81,23 @@ def load_model( def load_tokenizer( language: Languages, - pretrained_model_name_or_path: str | None = None, + pretrained_model_name_or_path: Optional[str] = None, ) -> PreTrainedTokenizer | PreTrainedTokenizerFast | DebertaV2Tokenizer: """ - 指定された言語の BERT モデルをロードし、ロード済みの BERT トークナイザーを返す - 一度ロードされていれば、ロード済みの BERT トークナイザーを即座に返す - ライブラリ利用時は常に pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある - ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき + 指定された言語の BERT モデルをロードし、ロード済みの BERT トークナイザーを返す。 + 一度ロードされていれば、ロード済みの BERT トークナイザーを即座に返す。 + ライブラリ利用時は常に pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある。 + ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき。 - Style-Bert-VITS2 では、BERT モデルに下記の 3 つが利用されている - これ以外の BERT モデルを指定した場合は正常に動作しない可能性が高い + 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) + pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None) Returns: PreTrainedTokenizer | PreTrainedTokenizerFast | DebertaV2Tokenizer: ロード済みの BERT トークナイザー @@ -127,7 +127,7 @@ def load_tokenizer( def unload_model(language: Languages) -> None: """ - 指定された言語の BERT モデルをアンロードする + 指定された言語の BERT モデルをアンロードする。 Args: language (Languages): アンロードする BERT モデルの言語 @@ -143,7 +143,7 @@ def unload_model(language: Languages) -> None: def unload_tokenizer(language: Languages) -> None: """ - 指定された言語の BERT トークナイザーをアンロードする + 指定された言語の BERT トークナイザーをアンロードする。 Args: language (Languages): アンロードする BERT トークナイザーの言語 @@ -159,7 +159,7 @@ def unload_tokenizer(language: Languages) -> None: def unload_all_models() -> None: """ - すべての BERT モデルをアンロードする + すべての BERT モデルをアンロードする。 """ for language in list(__loaded_models.keys()): @@ -169,7 +169,7 @@ def unload_all_models() -> None: def unload_all_tokenizers() -> None: """ - すべての BERT トークナイザーをアンロードする + すべての BERT トークナイザーをアンロードする。 """ for language in list(__loaded_tokenizers.keys()): diff --git a/style_bert_vits2/text_processing/chinese/bert_feature.py b/style_bert_vits2/text_processing/chinese/bert_feature.py index 3178565..f8bce04 100644 --- a/style_bert_vits2/text_processing/chinese/bert_feature.py +++ b/style_bert_vits2/text_processing/chinese/bert_feature.py @@ -1,4 +1,5 @@ import sys +from typing import Optional import torch from transformers import PreTrainedModel @@ -14,7 +15,7 @@ def extract_bert_feature( text: str, word2ph: list[int], device: torch.device | str, - assist_text: str | None = None, + assist_text: Optional[str] = None, assist_text_weight: float = 0.7, ) -> torch.Tensor: """ @@ -24,7 +25,7 @@ def extract_bert_feature( text (str): 中国語のテキスト word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト device (torch.device | str): 推論に利用するデバイス - assist_text (str | None, optional): 補助テキスト (デフォルト: None) + assist_text (Optional[str], optional): 補助テキスト (デフォルト: None) assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7) Returns: diff --git a/style_bert_vits2/text_processing/english/bert_feature.py b/style_bert_vits2/text_processing/english/bert_feature.py index b29d531..79f24c9 100644 --- a/style_bert_vits2/text_processing/english/bert_feature.py +++ b/style_bert_vits2/text_processing/english/bert_feature.py @@ -1,4 +1,5 @@ import sys +from typing import Optional import torch from transformers import PreTrainedModel @@ -14,7 +15,7 @@ def extract_bert_feature( text: str, word2ph: list[int], device: torch.device | str, - assist_text: str | None = None, + assist_text: Optional[str] = None, assist_text_weight: float = 0.7, ) -> torch.Tensor: """ @@ -24,7 +25,7 @@ def extract_bert_feature( text (str): 英語のテキスト word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト device (torch.device | str): 推論に利用するデバイス - assist_text (str | None, optional): 補助テキスト (デフォルト: None) + assist_text (Optional[str], optional): 補助テキスト (デフォルト: None) assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7) Returns: diff --git a/style_bert_vits2/text_processing/japanese/bert_feature.py b/style_bert_vits2/text_processing/japanese/bert_feature.py index d1809fe..3a44a8d 100644 --- a/style_bert_vits2/text_processing/japanese/bert_feature.py +++ b/style_bert_vits2/text_processing/japanese/bert_feature.py @@ -1,4 +1,5 @@ import sys +from typing import Optional import torch from transformers import PreTrainedModel @@ -15,7 +16,7 @@ def extract_bert_feature( text: str, word2ph: list[int], device: torch.device | str, - assist_text: str | None = None, + assist_text: Optional[str] = None, assist_text_weight: float = 0.7, ) -> torch.Tensor: """ @@ -25,7 +26,7 @@ def extract_bert_feature( text (str): 日本語のテキスト word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト device (torch.device | str): 推論に利用するデバイス - assist_text (str | None, optional): 補助テキスト (デフォルト: None) + assist_text (Optional[str], optional): 補助テキスト (デフォルト: None) assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7) Returns: From fac4f9a8ab5f8bdf5eb714cc149658210925163f Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 06:20:44 +0000 Subject: [PATCH 27/64] Refactor: rename text_processing to nlp "text_processing" is clearer, but the import statement is longer. "nlp" is shorter and makes it clear that it is natural language processing. --- bert_gen.py | 4 ++-- data_utils.py | 2 +- preprocess_text.py | 2 +- server_editor.py | 8 ++++---- style_bert_vits2/constants.py | 2 +- style_bert_vits2/models/infer.py | 4 ++-- style_bert_vits2/models/models.py | 2 +- style_bert_vits2/models/models_jp_extra.py | 2 +- .../{text_processing => nlp}/__init__.py | 14 +++++++------- .../{text_processing => nlp}/bert_models.py | 0 .../{text_processing => nlp}/chinese/__init__.py | 6 +++--- .../chinese/bert_feature.py | 2 +- .../chinese/tone_sandhi.py | 0 .../{text_processing => nlp}/english/__init__.py | 4 ++-- .../english/bert_feature.py | 2 +- .../{text_processing => nlp}/english/cmudict.rep | 0 .../english/cmudict_cache.pickle | Bin .../english/opencpop-strict.txt | 0 style_bert_vits2/nlp/japanese/__init__.py | 2 ++ .../japanese/bert_feature.py | 4 ++-- .../{text_processing => nlp}/japanese/g2p.py | 10 +++++----- .../japanese/g2p_utils.py | 6 +++--- .../japanese/mora_list.py | 0 .../japanese/normalizer.py | 2 +- .../japanese/pyopenjtalk_worker/__init__.py | 4 ++-- .../japanese/pyopenjtalk_worker/__main__.py | 4 ++-- .../japanese/pyopenjtalk_worker/worker_client.py | 2 +- .../japanese/pyopenjtalk_worker/worker_common.py | 0 .../japanese/pyopenjtalk_worker/worker_server.py | 2 +- .../japanese/user_dict/README.md | 0 .../japanese/user_dict/__init__.py | 6 +++--- .../japanese/user_dict/part_of_speech_data.py | 2 +- .../japanese/user_dict/word_model.py | 0 .../{text_processing => nlp}/symbols.py | 0 .../text_processing/japanese/__init__.py | 2 -- train_ms.py | 2 +- train_ms_jp_extra.py | 2 +- webui/inference.py | 4 ++-- 38 files changed, 54 insertions(+), 54 deletions(-) rename style_bert_vits2/{text_processing => nlp}/__init__.py (86%) rename style_bert_vits2/{text_processing => nlp}/bert_models.py (100%) rename style_bert_vits2/{text_processing => nlp}/chinese/__init__.py (95%) rename style_bert_vits2/{text_processing => nlp}/chinese/bert_feature.py (98%) rename style_bert_vits2/{text_processing => nlp}/chinese/tone_sandhi.py (100%) rename style_bert_vits2/{text_processing => nlp}/english/__init__.py (98%) rename style_bert_vits2/{text_processing => nlp}/english/bert_feature.py (98%) rename style_bert_vits2/{text_processing => nlp}/english/cmudict.rep (100%) rename style_bert_vits2/{text_processing => nlp}/english/cmudict_cache.pickle (100%) rename style_bert_vits2/{text_processing => nlp}/english/opencpop-strict.txt (100%) create mode 100644 style_bert_vits2/nlp/japanese/__init__.py rename style_bert_vits2/{text_processing => nlp}/japanese/bert_feature.py (96%) rename style_bert_vits2/{text_processing => nlp}/japanese/g2p.py (98%) rename style_bert_vits2/{text_processing => nlp}/japanese/g2p_utils.py (93%) rename style_bert_vits2/{text_processing => nlp}/japanese/mora_list.py (100%) rename style_bert_vits2/{text_processing => nlp}/japanese/normalizer.py (98%) rename style_bert_vits2/{text_processing => nlp}/japanese/pyopenjtalk_worker/__init__.py (94%) rename style_bert_vits2/{text_processing => nlp}/japanese/pyopenjtalk_worker/__main__.py (58%) rename style_bert_vits2/{text_processing => nlp}/japanese/pyopenjtalk_worker/worker_client.py (93%) rename style_bert_vits2/{text_processing => nlp}/japanese/pyopenjtalk_worker/worker_common.py (100%) rename style_bert_vits2/{text_processing => nlp}/japanese/pyopenjtalk_worker/worker_server.py (98%) rename style_bert_vits2/{text_processing => nlp}/japanese/user_dict/README.md (100%) rename style_bert_vits2/{text_processing => nlp}/japanese/user_dict/__init__.py (98%) rename style_bert_vits2/{text_processing => nlp}/japanese/user_dict/part_of_speech_data.py (97%) rename style_bert_vits2/{text_processing => nlp}/japanese/user_dict/word_model.py (100%) rename style_bert_vits2/{text_processing => nlp}/symbols.py (100%) delete mode 100644 style_bert_vits2/text_processing/japanese/__init__.py diff --git a/bert_gen.py b/bert_gen.py index f22c2e8..79ea65f 100644 --- a/bert_gen.py +++ b/bert_gen.py @@ -9,8 +9,8 @@ import utils from config import config from style_bert_vits2.logging import logger from style_bert_vits2.models import commons -from style_bert_vits2.text_processing import cleaned_text_to_sequence, extract_bert_feature -from style_bert_vits2.text_processing.japanese import pyopenjtalk_worker as pyopenjtalk +from style_bert_vits2.nlp import cleaned_text_to_sequence, extract_bert_feature +from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT pyopenjtalk.initialize() diff --git a/data_utils.py b/data_utils.py index 7738247..460c15f 100644 --- a/data_utils.py +++ b/data_utils.py @@ -12,7 +12,7 @@ from mel_processing import mel_spectrogram_torch, spectrogram_torch from utils import load_filepaths_and_text, load_wav_to_torch from style_bert_vits2.logging import logger from style_bert_vits2.models import commons -from style_bert_vits2.text_processing import cleaned_text_to_sequence +from style_bert_vits2.nlp import cleaned_text_to_sequence """Multi speaker version""" diff --git a/preprocess_text.py b/preprocess_text.py index 03e1232..4966305 100644 --- a/preprocess_text.py +++ b/preprocess_text.py @@ -9,7 +9,7 @@ from tqdm import tqdm from config import config from style_bert_vits2.logging import logger -from style_bert_vits2.text_processing import clean_text +from style_bert_vits2.nlp import clean_text from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT preprocess_text_config = config.preprocess_text_config diff --git a/server_editor.py b/server_editor.py index b74d74f..f72800d 100644 --- a/server_editor.py +++ b/server_editor.py @@ -42,10 +42,10 @@ 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 import normalize_text -from style_bert_vits2.text_processing.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone -from style_bert_vits2.text_processing.japanese.user_dict import ( +from style_bert_vits2.nlp import bert_models +from style_bert_vits2.nlp.japanese import normalize_text +from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone +from style_bert_vits2.nlp.japanese.user_dict import ( apply_word, delete_word, read_dict, diff --git a/style_bert_vits2/constants.py b/style_bert_vits2/constants.py index 064221e..d5a46bf 100644 --- a/style_bert_vits2/constants.py +++ b/style_bert_vits2/constants.py @@ -23,7 +23,7 @@ DEFAULT_BERT_TOKENIZER_PATHS = { } # デフォルトのユーザー辞書ディレクトリ -## style_bert_vits2.text_processing.japanese.user_dict モジュールのデフォルト値として利用される +## style_bert_vits2.nlp.japanese.user_dict モジュールのデフォルト値として利用される ## ライブラリとしての利用などで外部のユーザー辞書を指定したい場合は、user_dict 以下の各関数の実行時、引数に辞書データファイルのパスを指定する DEFAULT_USER_DICT_DIR = BASE_DIR / "dict_data" diff --git a/style_bert_vits2/models/infer.py b/style_bert_vits2/models/infer.py index a5ce70f..9265eee 100644 --- a/style_bert_vits2/models/infer.py +++ b/style_bert_vits2/models/infer.py @@ -7,8 +7,8 @@ from style_bert_vits2.logging import logger from style_bert_vits2.models import commons from style_bert_vits2.models.models import SynthesizerTrn from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra -from style_bert_vits2.text_processing import clean_text, cleaned_text_to_sequence, extract_bert_feature -from style_bert_vits2.text_processing.symbols import SYMBOLS +from style_bert_vits2.nlp import clean_text, cleaned_text_to_sequence, extract_bert_feature +from style_bert_vits2.nlp.symbols import SYMBOLS def get_net_g(model_path: str, version: str, device: str, hps): diff --git a/style_bert_vits2/models/models.py b/style_bert_vits2/models/models.py index 7f14be4..b4e1eb7 100644 --- a/style_bert_vits2/models/models.py +++ b/style_bert_vits2/models/models.py @@ -11,7 +11,7 @@ from style_bert_vits2.models import commons from style_bert_vits2.models import modules from style_bert_vits2.models import monotonic_alignment from style_bert_vits2.models.commons import get_padding, init_weights -from style_bert_vits2.text_processing.symbols import NUM_LANGUAGES, NUM_TONES, SYMBOLS +from style_bert_vits2.nlp.symbols import NUM_LANGUAGES, NUM_TONES, SYMBOLS class DurationDiscriminator(nn.Module): # vits2 diff --git a/style_bert_vits2/models/models_jp_extra.py b/style_bert_vits2/models/models_jp_extra.py index 54591f6..7bae8b4 100644 --- a/style_bert_vits2/models/models_jp_extra.py +++ b/style_bert_vits2/models/models_jp_extra.py @@ -10,7 +10,7 @@ from style_bert_vits2.models import attentions from style_bert_vits2.models import commons from style_bert_vits2.models import modules from style_bert_vits2.models import monotonic_alignment -from style_bert_vits2.text_processing.symbols import SYMBOLS, NUM_TONES, NUM_LANGUAGES +from style_bert_vits2.nlp.symbols import SYMBOLS, NUM_TONES, NUM_LANGUAGES class DurationDiscriminator(nn.Module): # vits2 diff --git a/style_bert_vits2/text_processing/__init__.py b/style_bert_vits2/nlp/__init__.py similarity index 86% rename from style_bert_vits2/text_processing/__init__.py rename to style_bert_vits2/nlp/__init__.py index 4dadaed..47786d1 100644 --- a/style_bert_vits2/text_processing/__init__.py +++ b/style_bert_vits2/nlp/__init__.py @@ -2,7 +2,7 @@ import torch from typing import Optional from style_bert_vits2.constants import Languages -from style_bert_vits2.text_processing.symbols import ( +from style_bert_vits2.nlp.symbols import ( LANGUAGE_ID_MAP, LANGUAGE_TONE_START_MAP, SYMBOLS, @@ -36,11 +36,11 @@ def extract_bert_feature( """ if language == Languages.JP: - from style_bert_vits2.text_processing.japanese.bert_feature import extract_bert_feature + from style_bert_vits2.nlp.japanese.bert_feature import extract_bert_feature elif language == Languages.EN: - from style_bert_vits2.text_processing.english.bert_feature import extract_bert_feature + from style_bert_vits2.nlp.english.bert_feature import extract_bert_feature elif language == Languages.ZH: - from style_bert_vits2.text_processing.chinese.bert_feature import extract_bert_feature + from style_bert_vits2.nlp.chinese.bert_feature import extract_bert_feature else: raise ValueError(f"Language {language} not supported") @@ -68,15 +68,15 @@ def clean_text( # Changed to import inside if condition to avoid unnecessary import if language == Languages.JP: - from style_bert_vits2.text_processing.japanese import g2p, normalize_text + from style_bert_vits2.nlp.japanese import g2p, normalize_text norm_text = normalize_text(text) phones, tones, word2ph = g2p(norm_text, use_jp_extra, raise_yomi_error) elif language == Languages.EN: - from style_bert_vits2.text_processing.english import g2p, normalize_text + from style_bert_vits2.nlp.english import g2p, normalize_text norm_text = normalize_text(text) phones, tones, word2ph = g2p(norm_text) elif language == Languages.ZH: - from style_bert_vits2.text_processing.chinese import g2p, normalize_text + from style_bert_vits2.nlp.chinese import g2p, normalize_text norm_text = normalize_text(text) phones, tones, word2ph = g2p(norm_text) else: diff --git a/style_bert_vits2/text_processing/bert_models.py b/style_bert_vits2/nlp/bert_models.py similarity index 100% rename from style_bert_vits2/text_processing/bert_models.py rename to style_bert_vits2/nlp/bert_models.py diff --git a/style_bert_vits2/text_processing/chinese/__init__.py b/style_bert_vits2/nlp/chinese/__init__.py similarity index 95% rename from style_bert_vits2/text_processing/chinese/__init__.py rename to style_bert_vits2/nlp/chinese/__init__.py index 23acbf3..5b70863 100644 --- a/style_bert_vits2/text_processing/chinese/__init__.py +++ b/style_bert_vits2/nlp/chinese/__init__.py @@ -5,8 +5,8 @@ import cn2an import jieba.posseg as psg from pypinyin import lazy_pinyin, Style -from style_bert_vits2.text_processing.chinese.tone_sandhi import ToneSandhi -from style_bert_vits2.text_processing.symbols import PUNCTUATIONS +from style_bert_vits2.nlp.chinese.tone_sandhi import ToneSandhi +from style_bert_vits2.nlp.symbols import PUNCTUATIONS current_file_path = os.path.dirname(__file__) @@ -177,7 +177,7 @@ def normalize_text(text: str) -> str: if __name__ == "__main__": - from style_bert_vits2.text_processing.chinese.bert_feature import extract_bert_feature + from style_bert_vits2.nlp.chinese.bert_feature import extract_bert_feature text = "啊!但是《原神》是由,米哈\游自主, [研发]的一款全.新开放世界.冒险游戏" text = normalize_text(text) diff --git a/style_bert_vits2/text_processing/chinese/bert_feature.py b/style_bert_vits2/nlp/chinese/bert_feature.py similarity index 98% rename from style_bert_vits2/text_processing/chinese/bert_feature.py rename to style_bert_vits2/nlp/chinese/bert_feature.py index f8bce04..b97950b 100644 --- a/style_bert_vits2/text_processing/chinese/bert_feature.py +++ b/style_bert_vits2/nlp/chinese/bert_feature.py @@ -5,7 +5,7 @@ import torch from transformers import PreTrainedModel from style_bert_vits2.constants import Languages -from style_bert_vits2.text_processing import bert_models +from style_bert_vits2.nlp import bert_models __models: dict[torch.device | str, PreTrainedModel] = {} diff --git a/style_bert_vits2/text_processing/chinese/tone_sandhi.py b/style_bert_vits2/nlp/chinese/tone_sandhi.py similarity index 100% rename from style_bert_vits2/text_processing/chinese/tone_sandhi.py rename to style_bert_vits2/nlp/chinese/tone_sandhi.py diff --git a/style_bert_vits2/text_processing/english/__init__.py b/style_bert_vits2/nlp/english/__init__.py similarity index 98% rename from style_bert_vits2/text_processing/english/__init__.py rename to style_bert_vits2/nlp/english/__init__.py index b57e250..3610b13 100644 --- a/style_bert_vits2/text_processing/english/__init__.py +++ b/style_bert_vits2/nlp/english/__init__.py @@ -4,8 +4,8 @@ import re from g2p_en import G2p from style_bert_vits2.constants import Languages -from style_bert_vits2.text_processing import bert_models -from style_bert_vits2.text_processing.symbols import PUNCTUATIONS, SYMBOLS +from style_bert_vits2.nlp import bert_models +from style_bert_vits2.nlp.symbols import PUNCTUATIONS, SYMBOLS current_file_path = os.path.dirname(__file__) diff --git a/style_bert_vits2/text_processing/english/bert_feature.py b/style_bert_vits2/nlp/english/bert_feature.py similarity index 98% rename from style_bert_vits2/text_processing/english/bert_feature.py rename to style_bert_vits2/nlp/english/bert_feature.py index 79f24c9..647920d 100644 --- a/style_bert_vits2/text_processing/english/bert_feature.py +++ b/style_bert_vits2/nlp/english/bert_feature.py @@ -5,7 +5,7 @@ import torch from transformers import PreTrainedModel from style_bert_vits2.constants import Languages -from style_bert_vits2.text_processing import bert_models +from style_bert_vits2.nlp import bert_models __models: dict[torch.device | str, PreTrainedModel] = {} diff --git a/style_bert_vits2/text_processing/english/cmudict.rep b/style_bert_vits2/nlp/english/cmudict.rep similarity index 100% rename from style_bert_vits2/text_processing/english/cmudict.rep rename to style_bert_vits2/nlp/english/cmudict.rep diff --git a/style_bert_vits2/text_processing/english/cmudict_cache.pickle b/style_bert_vits2/nlp/english/cmudict_cache.pickle similarity index 100% rename from style_bert_vits2/text_processing/english/cmudict_cache.pickle rename to style_bert_vits2/nlp/english/cmudict_cache.pickle diff --git a/style_bert_vits2/text_processing/english/opencpop-strict.txt b/style_bert_vits2/nlp/english/opencpop-strict.txt similarity index 100% rename from style_bert_vits2/text_processing/english/opencpop-strict.txt rename to style_bert_vits2/nlp/english/opencpop-strict.txt diff --git a/style_bert_vits2/nlp/japanese/__init__.py b/style_bert_vits2/nlp/japanese/__init__.py new file mode 100644 index 0000000..5c7f19f --- /dev/null +++ b/style_bert_vits2/nlp/japanese/__init__.py @@ -0,0 +1,2 @@ +from style_bert_vits2.nlp.japanese.g2p import g2p # noqa: F401 +from style_bert_vits2.nlp.japanese.normalizer import normalize_text # noqa: F401 diff --git a/style_bert_vits2/text_processing/japanese/bert_feature.py b/style_bert_vits2/nlp/japanese/bert_feature.py similarity index 96% rename from style_bert_vits2/text_processing/japanese/bert_feature.py rename to style_bert_vits2/nlp/japanese/bert_feature.py index 3a44a8d..ede1f83 100644 --- a/style_bert_vits2/text_processing/japanese/bert_feature.py +++ b/style_bert_vits2/nlp/japanese/bert_feature.py @@ -5,8 +5,8 @@ import torch from transformers import PreTrainedModel from style_bert_vits2.constants import Languages -from style_bert_vits2.text_processing import bert_models -from style_bert_vits2.text_processing.japanese.g2p import text_to_sep_kata +from style_bert_vits2.nlp import bert_models +from style_bert_vits2.nlp.japanese.g2p import text_to_sep_kata __models: dict[torch.device | str, PreTrainedModel] = {} diff --git a/style_bert_vits2/text_processing/japanese/g2p.py b/style_bert_vits2/nlp/japanese/g2p.py similarity index 98% rename from style_bert_vits2/text_processing/japanese/g2p.py rename to style_bert_vits2/nlp/japanese/g2p.py index 2f9746b..45e7817 100644 --- a/style_bert_vits2/text_processing/japanese/g2p.py +++ b/style_bert_vits2/nlp/japanese/g2p.py @@ -2,11 +2,11 @@ import re 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 import pyopenjtalk_worker as pyopenjtalk -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 +from style_bert_vits2.nlp import bert_models +from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk +from style_bert_vits2.nlp.japanese.mora_list import MORA_KATA_TO_MORA_PHONEMES +from style_bert_vits2.nlp.japanese.normalizer import replace_punctuation +from style_bert_vits2.nlp.symbols import PUNCTUATIONS pyopenjtalk.initialize() diff --git a/style_bert_vits2/text_processing/japanese/g2p_utils.py b/style_bert_vits2/nlp/japanese/g2p_utils.py similarity index 93% rename from style_bert_vits2/text_processing/japanese/g2p_utils.py rename to style_bert_vits2/nlp/japanese/g2p_utils.py index 4ea56e8..893d3b5 100644 --- a/style_bert_vits2/text_processing/japanese/g2p_utils.py +++ b/style_bert_vits2/nlp/japanese/g2p_utils.py @@ -1,9 +1,9 @@ -from style_bert_vits2.text_processing.japanese.g2p import g2p -from style_bert_vits2.text_processing.japanese.mora_list import ( +from style_bert_vits2.nlp.japanese.g2p import g2p +from style_bert_vits2.nlp.japanese.mora_list import ( MORA_KATA_TO_MORA_PHONEMES, MORA_PHONEMES_TO_MORA_KATA, ) -from style_bert_vits2.text_processing.symbols import PUNCTUATIONS +from style_bert_vits2.nlp.symbols import PUNCTUATIONS def g2kata_tone(norm_text: str) -> list[tuple[str, int]]: diff --git a/style_bert_vits2/text_processing/japanese/mora_list.py b/style_bert_vits2/nlp/japanese/mora_list.py similarity index 100% rename from style_bert_vits2/text_processing/japanese/mora_list.py rename to style_bert_vits2/nlp/japanese/mora_list.py diff --git a/style_bert_vits2/text_processing/japanese/normalizer.py b/style_bert_vits2/nlp/japanese/normalizer.py similarity index 98% rename from style_bert_vits2/text_processing/japanese/normalizer.py rename to style_bert_vits2/nlp/japanese/normalizer.py index 8276338..b8cad90 100644 --- a/style_bert_vits2/text_processing/japanese/normalizer.py +++ b/style_bert_vits2/nlp/japanese/normalizer.py @@ -2,7 +2,7 @@ import re import unicodedata from num2words import num2words -from style_bert_vits2.text_processing.symbols import PUNCTUATIONS +from style_bert_vits2.nlp.symbols import PUNCTUATIONS def normalize_text(text: str) -> str: diff --git a/style_bert_vits2/text_processing/japanese/pyopenjtalk_worker/__init__.py b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py similarity index 94% rename from style_bert_vits2/text_processing/japanese/pyopenjtalk_worker/__init__.py rename to style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py index c5fc697..fc8f6da 100644 --- a/style_bert_vits2/text_processing/japanese/pyopenjtalk_worker/__init__.py +++ b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py @@ -6,8 +6,8 @@ to avoid user dictionary access error from typing import Any, Optional from style_bert_vits2.logging import logger -from style_bert_vits2.text_processing.japanese.pyopenjtalk_worker.worker_client import WorkerClient -from style_bert_vits2.text_processing.japanese.pyopenjtalk_worker.worker_common import WORKER_PORT +from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_client import WorkerClient +from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_common import WORKER_PORT WORKER_CLIENT: Optional[WorkerClient] = None diff --git a/style_bert_vits2/text_processing/japanese/pyopenjtalk_worker/__main__.py b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__main__.py similarity index 58% rename from style_bert_vits2/text_processing/japanese/pyopenjtalk_worker/__main__.py rename to style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__main__.py index f022901..2452a16 100644 --- a/style_bert_vits2/text_processing/japanese/pyopenjtalk_worker/__main__.py +++ b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__main__.py @@ -1,7 +1,7 @@ import argparse -from style_bert_vits2.text_processing.japanese.pyopenjtalk_worker.worker_common import WORKER_PORT -from style_bert_vits2.text_processing.japanese.pyopenjtalk_worker.worker_server import WorkerServer +from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_common import WORKER_PORT +from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_server import WorkerServer def main() -> None: diff --git a/style_bert_vits2/text_processing/japanese/pyopenjtalk_worker/worker_client.py b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_client.py similarity index 93% rename from style_bert_vits2/text_processing/japanese/pyopenjtalk_worker/worker_client.py rename to style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_client.py index 0bfe87a..b87a937 100644 --- a/style_bert_vits2/text_processing/japanese/pyopenjtalk_worker/worker_client.py +++ b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_client.py @@ -2,7 +2,7 @@ import socket from typing import Any, cast from style_bert_vits2.logging import logger -from style_bert_vits2.text_processing.japanese.pyopenjtalk_worker.worker_common import RequestType, receive_data, send_data +from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_common import RequestType, receive_data, send_data class WorkerClient: diff --git a/style_bert_vits2/text_processing/japanese/pyopenjtalk_worker/worker_common.py b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_common.py similarity index 100% rename from style_bert_vits2/text_processing/japanese/pyopenjtalk_worker/worker_common.py rename to style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_common.py diff --git a/style_bert_vits2/text_processing/japanese/pyopenjtalk_worker/worker_server.py b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_server.py similarity index 98% rename from style_bert_vits2/text_processing/japanese/pyopenjtalk_worker/worker_server.py rename to style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_server.py index 4282c78..e9a3dc1 100644 --- a/style_bert_vits2/text_processing/japanese/pyopenjtalk_worker/worker_server.py +++ b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_server.py @@ -6,7 +6,7 @@ from typing import Any, cast import pyopenjtalk from style_bert_vits2.logging import logger -from style_bert_vits2.text_processing.japanese.pyopenjtalk_worker.worker_common import ( +from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_common import ( ConnectionClosedException, RequestType, receive_data, diff --git a/style_bert_vits2/text_processing/japanese/user_dict/README.md b/style_bert_vits2/nlp/japanese/user_dict/README.md similarity index 100% rename from style_bert_vits2/text_processing/japanese/user_dict/README.md rename to style_bert_vits2/nlp/japanese/user_dict/README.md diff --git a/style_bert_vits2/text_processing/japanese/user_dict/__init__.py b/style_bert_vits2/nlp/japanese/user_dict/__init__.py similarity index 98% rename from style_bert_vits2/text_processing/japanese/user_dict/__init__.py rename to style_bert_vits2/nlp/japanese/user_dict/__init__.py index 041f37c..097cc6a 100644 --- a/style_bert_vits2/text_processing/japanese/user_dict/__init__.py +++ b/style_bert_vits2/nlp/japanese/user_dict/__init__.py @@ -16,9 +16,9 @@ import numpy as np from fastapi import HTTPException from style_bert_vits2.constants import DEFAULT_USER_DICT_DIR -from style_bert_vits2.text_processing.japanese import pyopenjtalk_worker as pyopenjtalk -from style_bert_vits2.text_processing.japanese.user_dict.word_model import UserDictWord, WordTypes -from style_bert_vits2.text_processing.japanese.user_dict.part_of_speech_data import MAX_PRIORITY, MIN_PRIORITY, part_of_speech_data +from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk +from style_bert_vits2.nlp.japanese.user_dict.word_model import UserDictWord, WordTypes +from style_bert_vits2.nlp.japanese.user_dict.part_of_speech_data import MAX_PRIORITY, MIN_PRIORITY, part_of_speech_data pyopenjtalk.initialize() diff --git a/style_bert_vits2/text_processing/japanese/user_dict/part_of_speech_data.py b/style_bert_vits2/nlp/japanese/user_dict/part_of_speech_data.py similarity index 97% rename from style_bert_vits2/text_processing/japanese/user_dict/part_of_speech_data.py rename to style_bert_vits2/nlp/japanese/user_dict/part_of_speech_data.py index db42d38..443bdc5 100644 --- a/style_bert_vits2/text_processing/japanese/user_dict/part_of_speech_data.py +++ b/style_bert_vits2/nlp/japanese/user_dict/part_of_speech_data.py @@ -7,7 +7,7 @@ from typing import Dict -from style_bert_vits2.text_processing.japanese.user_dict.word_model import ( +from style_bert_vits2.nlp.japanese.user_dict.word_model import ( USER_DICT_MAX_PRIORITY, USER_DICT_MIN_PRIORITY, PartOfSpeechDetail, diff --git a/style_bert_vits2/text_processing/japanese/user_dict/word_model.py b/style_bert_vits2/nlp/japanese/user_dict/word_model.py similarity index 100% rename from style_bert_vits2/text_processing/japanese/user_dict/word_model.py rename to style_bert_vits2/nlp/japanese/user_dict/word_model.py diff --git a/style_bert_vits2/text_processing/symbols.py b/style_bert_vits2/nlp/symbols.py similarity index 100% rename from style_bert_vits2/text_processing/symbols.py rename to style_bert_vits2/nlp/symbols.py diff --git a/style_bert_vits2/text_processing/japanese/__init__.py b/style_bert_vits2/text_processing/japanese/__init__.py deleted file mode 100644 index 9123377..0000000 --- a/style_bert_vits2/text_processing/japanese/__init__.py +++ /dev/null @@ -1,2 +0,0 @@ -from style_bert_vits2.text_processing.japanese.g2p import g2p # noqa: F401 -from style_bert_vits2.text_processing.japanese.normalizer import normalize_text # noqa: F401 diff --git a/train_ms.py b/train_ms.py index 05cd4cb..0cb7a82 100644 --- a/train_ms.py +++ b/train_ms.py @@ -31,7 +31,7 @@ from style_bert_vits2.models.models import ( MultiPeriodDiscriminator, SynthesizerTrn, ) -from style_bert_vits2.text_processing.symbols import SYMBOLS +from style_bert_vits2.nlp.symbols import SYMBOLS from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT torch.backends.cuda.matmul.allow_tf32 = True diff --git a/train_ms_jp_extra.py b/train_ms_jp_extra.py index f04d2ae..7d0636b 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -32,7 +32,7 @@ from style_bert_vits2.models.models_jp_extra import ( SynthesizerTrn, WavLMDiscriminator, ) -from style_bert_vits2.text_processing.symbols import SYMBOLS +from style_bert_vits2.nlp.symbols import SYMBOLS from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT torch.backends.cuda.matmul.allow_tf32 = True diff --git a/webui/inference.py b/webui/inference.py index f8de8c3..026481f 100644 --- a/webui/inference.py +++ b/webui/inference.py @@ -20,8 +20,8 @@ from style_bert_vits2.constants import ( ) from style_bert_vits2.logging import logger from style_bert_vits2.models.infer import InvalidToneError -from style_bert_vits2.text_processing.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone -from style_bert_vits2.text_processing.japanese.normalizer import normalize_text +from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone +from style_bert_vits2.nlp.japanese.normalizer import normalize_text languages = [l.value for l in Languages] From 75467936d96e4396fcb60eeaeaea48ce379a9027 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 06:39:51 +0000 Subject: [PATCH 28/64] Refactor: cleanup style_bert_vits2/nlp/english/__init__.py --- style_bert_vits2/nlp/english/__init__.py | 558 ++++++++++------------- 1 file changed, 242 insertions(+), 316 deletions(-) diff --git a/style_bert_vits2/nlp/english/__init__.py b/style_bert_vits2/nlp/english/__init__.py index 3610b13..e5041d1 100644 --- a/style_bert_vits2/nlp/english/__init__.py +++ b/style_bert_vits2/nlp/english/__init__.py @@ -1,6 +1,9 @@ import pickle import os import re +from pathlib import Path + +import inflect from g2p_en import G2p from style_bert_vits2.constants import Languages @@ -8,88 +11,215 @@ from style_bert_vits2.nlp import bert_models from style_bert_vits2.nlp.symbols import PUNCTUATIONS, SYMBOLS -current_file_path = os.path.dirname(__file__) -CMU_DICT_PATH = os.path.join(current_file_path, "cmudict.rep") -CACHE_PATH = os.path.join(current_file_path, "cmudict_cache.pickle") -_g2p = G2p() - -arpa = { - "AH0", - "S", - "AH1", - "EY2", - "AE2", - "EH0", - "OW2", - "UH0", - "NG", - "B", - "G", - "AY0", - "M", - "AA0", - "F", - "AO0", - "ER2", - "UH1", - "IY1", - "AH2", - "DH", - "IY0", - "EY1", - "IH0", - "K", - "N", - "W", - "IY2", - "T", - "AA1", - "ER1", - "EH2", - "OY0", - "UH2", - "UW1", - "Z", - "AW2", - "AW1", - "V", - "UW2", - "AA2", - "ER", - "AW0", - "UW0", - "R", - "OW1", - "EH1", - "ZH", - "AE0", - "IH2", - "IH", - "Y", - "JH", - "P", - "AY1", - "EY0", - "OY2", - "TH", - "HH", - "D", - "ER0", - "CH", - "AO1", - "AE1", - "AO2", - "OY1", - "AY2", - "IH1", - "OW0", - "L", - "SH", -} +CMU_DICT_PATH = Path(__file__).parent / "cmudict.rep" +CACHE_PATH = Path(__file__).parent / "cmudict_cache.pickle" -def post_replace_ph(ph): - rep_map = { +def g2p(text: str) -> tuple[list[str], list[int], list[int]]: + + ARPA = { + "AH0", + "S", + "AH1", + "EY2", + "AE2", + "EH0", + "OW2", + "UH0", + "NG", + "B", + "G", + "AY0", + "M", + "AA0", + "F", + "AO0", + "ER2", + "UH1", + "IY1", + "AH2", + "DH", + "IY0", + "EY1", + "IH0", + "K", + "N", + "W", + "IY2", + "T", + "AA1", + "ER1", + "EH2", + "OY0", + "UH2", + "UW1", + "Z", + "AW2", + "AW1", + "V", + "UW2", + "AA2", + "ER", + "AW0", + "UW0", + "R", + "OW1", + "EH1", + "ZH", + "AE0", + "IH2", + "IH", + "Y", + "JH", + "P", + "AY1", + "EY0", + "OY2", + "TH", + "HH", + "D", + "ER0", + "CH", + "AO1", + "AE1", + "AO2", + "OY1", + "AY2", + "IH1", + "OW0", + "L", + "SH", + } + + _g2p = G2p() + + phones = [] + tones = [] + phone_len = [] + # tokens = [tokenizer.tokenize(i) for i in words] + words = __text_to_words(text) + eng_dict = __get_dict() + + for word in words: + temp_phones, temp_tones = [], [] + if len(word) > 1: + if "'" in word: + word = ["".join(word)] + for w in word: + if w in PUNCTUATIONS: + temp_phones.append(w) + temp_tones.append(0) + continue + if w.upper() in eng_dict: + phns, tns = __refine_syllables(eng_dict[w.upper()]) + temp_phones += [__post_replace_ph(i) for i in phns] + temp_tones += tns + # w2ph.append(len(phns)) + else: + phone_list = list(filter(lambda p: p != " ", _g2p(w))) # type: ignore + phns = [] + tns = [] + for ph in phone_list: + if ph in ARPA: + ph, tn = __refine_ph(ph) + phns.append(ph) + tns.append(tn) + else: + phns.append(ph) + tns.append(0) + temp_phones += [__post_replace_ph(i) for i in phns] + temp_tones += tns + phones += temp_phones + tones += temp_tones + phone_len.append(len(temp_phones)) + # phones = [post_replace_ph(i) for i in phones] + + word2ph = [] + for token, pl in zip(words, phone_len): + word_len = len(token) + + aaa = __distribute_phone(pl, word_len) + word2ph += aaa + + phones = ["_"] + phones + ["_"] + tones = [0] + tones + [0] + word2ph = [1] + word2ph + [1] + assert len(phones) == len(tones), text + assert len(phones) == sum(word2ph), text + + return phones, tones, word2ph + + +def normalize_text(text: str) -> str: + text = __normalize_numbers(text) + text = __replace_punctuation(text) + text = re.sub(r"([,;.\?\!])([\w])", r"\1 \2", text) + return text + + +def __normalize_numbers(text: str) -> str: + text = re.sub(__comma_number_re, __remove_commas, text) + text = re.sub(__pounds_re, r"\1 pounds", text) + text = re.sub(__dollars_re, __expand_dollars, text) + text = re.sub(__decimal_number_re, __expand_decimal_point, text) + text = re.sub(__ordinal_re, __expand_ordinal, text) + text = re.sub(__number_re, __expand_number, text) + return text + + +def __replace_punctuation(text: str) -> str: + REPLACE_MAP = { + ":": ",", + ";": ",", + ",": ",", + "。": ".", + "!": "!", + "?": "?", + "\n": ".", + ".": ".", + "…": "...", + "···": "...", + "・・・": "...", + "·": ",", + "・": ",", + "、": ",", + "$": ".", + "“": "'", + "”": "'", + '"': "'", + "‘": "'", + "’": "'", + "(": "'", + ")": "'", + "(": "'", + ")": "'", + "《": "'", + "》": "'", + "【": "'", + "】": "'", + "[": "'", + "]": "'", + "—": "-", + "−": "-", + "~": "-", + "~": "-", + "「": "'", + "」": "'", + } + pattern = re.compile("|".join(re.escape(p) for p in REPLACE_MAP.keys())) + replaced_text = pattern.sub(lambda x: REPLACE_MAP[x.group()], text) + # replaced_text = re.sub( + # r"[^\u3040-\u309F\u30A0-\u30FF\u4E00-\u9FFF\u3400-\u4DBF\u3005" + # + "".join(punctuation) + # + r"]+", + # "", + # replaced_text, + # ) + return replaced_text + + +def __post_replace_ph(ph: str) -> str: + REPLACE_MAP = { ":": ",", ";": ",", ",": ",", @@ -104,8 +234,8 @@ def post_replace_ph(ph): "・・・": "...", "v": "V", } - if ph in rep_map.keys(): - ph = rep_map[ph] + if ph in REPLACE_MAP.keys(): + ph = REPLACE_MAP[ph] if ph in SYMBOLS: return ph if ph not in SYMBOLS: @@ -113,63 +243,7 @@ def post_replace_ph(ph): return ph -rep_map = { - ":": ",", - ";": ",", - ",": ",", - "。": ".", - "!": "!", - "?": "?", - "\n": ".", - ".": ".", - "…": "...", - "···": "...", - "・・・": "...", - "·": ",", - "・": ",", - "、": ",", - "$": ".", - "“": "'", - "”": "'", - '"': "'", - "‘": "'", - "’": "'", - "(": "'", - ")": "'", - "(": "'", - ")": "'", - "《": "'", - "》": "'", - "【": "'", - "】": "'", - "[": "'", - "]": "'", - "—": "-", - "−": "-", - "~": "-", - "~": "-", - "「": "'", - "」": "'", -} - - -def replace_punctuation(text): - pattern = re.compile("|".join(re.escape(p) for p in rep_map.keys())) - - replaced_text = pattern.sub(lambda x: rep_map[x.group()], text) - - # replaced_text = re.sub( - # r"[^\u3040-\u309F\u30A0-\u30FF\u4E00-\u9FFF\u3400-\u4DBF\u3005" - # + "".join(punctuation) - # + r"]+", - # "", - # replaced_text, - # ) - - return replaced_text - - -def read_dict(): +def __read_dict() -> dict[str, list[list[str]]]: g2p_dict = {} start_line = 49 with open(CMU_DICT_PATH) as f: @@ -193,26 +267,23 @@ def read_dict(): return g2p_dict -def cache_dict(g2p_dict, file_path): +def __cache_dict(g2p_dict: dict[str, list[list[str]]], file_path: Path) -> None: with open(file_path, "wb") as pickle_file: pickle.dump(g2p_dict, pickle_file) -def get_dict(): +def __get_dict() -> dict[str, list[list[str]]]: if os.path.exists(CACHE_PATH): with open(CACHE_PATH, "rb") as pickle_file: g2p_dict = pickle.load(pickle_file) else: - g2p_dict = read_dict() - cache_dict(g2p_dict, CACHE_PATH) + g2p_dict = __read_dict() + __cache_dict(g2p_dict, CACHE_PATH) return g2p_dict -eng_dict = get_dict() - - -def refine_ph(phn): +def __refine_ph(phn: str) -> tuple[str, int]: tone = 0 if re.search(r"\d$", phn): tone = int(phn[-1]) + 1 @@ -222,93 +293,28 @@ def refine_ph(phn): return phn.lower(), tone -def refine_syllables(syllables): +def __refine_syllables(syllables: list[list[str]]) -> tuple[list[str], list[int]]: tones = [] phonemes = [] for phn_list in syllables: for i in range(len(phn_list)): phn = phn_list[i] - phn, tone = refine_ph(phn) + phn, tone = __refine_ph(phn) phonemes.append(phn) tones.append(tone) return phonemes, tones -import inflect - -_inflect = inflect.engine() -_comma_number_re = re.compile(r"([0-9][0-9\,]+[0-9])") -_decimal_number_re = re.compile(r"([0-9]+\.[0-9]+)") -_pounds_re = re.compile(r"£([0-9\,]*[0-9]+)") -_dollars_re = re.compile(r"\$([0-9\.\,]*[0-9]+)") -_ordinal_re = re.compile(r"[0-9]+(st|nd|rd|th)") -_number_re = re.compile(r"[0-9]+") - -# List of (regular expression, replacement) pairs for abbreviations: -_abbreviations = [ - (re.compile("\\b%s\\." % x[0], re.IGNORECASE), x[1]) - for x in [ - ("mrs", "misess"), - ("mr", "mister"), - ("dr", "doctor"), - ("st", "saint"), - ("co", "company"), - ("jr", "junior"), - ("maj", "major"), - ("gen", "general"), - ("drs", "doctors"), - ("rev", "reverend"), - ("lt", "lieutenant"), - ("hon", "honorable"), - ("sgt", "sergeant"), - ("capt", "captain"), - ("esq", "esquire"), - ("ltd", "limited"), - ("col", "colonel"), - ("ft", "fort"), - ] -] +__inflect = inflect.engine() +__comma_number_re = re.compile(r"([0-9][0-9\,]+[0-9])") +__decimal_number_re = re.compile(r"([0-9]+\.[0-9]+)") +__pounds_re = re.compile(r"£([0-9\,]*[0-9]+)") +__dollars_re = re.compile(r"\$([0-9\.\,]*[0-9]+)") +__ordinal_re = re.compile(r"[0-9]+(st|nd|rd|th)") +__number_re = re.compile(r"[0-9]+") -# List of (ipa, lazy ipa) pairs: -_lazy_ipa = [ - (re.compile("%s" % x[0]), x[1]) - for x in [ - ("r", "ɹ"), - ("æ", "e"), - ("ɑ", "a"), - ("ɔ", "o"), - ("ð", "z"), - ("θ", "s"), - ("ɛ", "e"), - ("ɪ", "i"), - ("ʊ", "u"), - ("ʒ", "ʥ"), - ("ʤ", "ʥ"), - ("ˈ", "↓"), - ] -] - -# List of (ipa, lazy ipa2) pairs: -_lazy_ipa2 = [ - (re.compile("%s" % x[0]), x[1]) - for x in [ - ("r", "ɹ"), - ("ð", "z"), - ("θ", "s"), - ("ʒ", "ʑ"), - ("ʤ", "dʑ"), - ("ˈ", "↓"), - ] -] - -# List of (ipa, ipa2) pairs -_ipa_to_ipa2 = [ - (re.compile("%s" % x[0]), x[1]) for x in [("r", "ɹ"), ("ʤ", "dʒ"), ("ʧ", "tʃ")] -] - - -def _expand_dollars(m): +def __expand_dollars(m: re.Match[str]) -> str: match = m.group(1) parts = match.split(".") if len(parts) > 2: @@ -329,53 +335,36 @@ def _expand_dollars(m): return "zero dollars" -def _remove_commas(m): +def __remove_commas(m: re.Match[str]) -> str: return m.group(1).replace(",", "") -def _expand_ordinal(m): - return _inflect.number_to_words(m.group(0)) +def __expand_ordinal(m: re.Match[str]) -> str: + return __inflect.number_to_words(m.group(0)) # type: ignore -def _expand_number(m): +def __expand_number(m: re.Match[str]) -> str: num = int(m.group(0)) if num > 1000 and num < 3000: if num == 2000: return "two thousand" elif num > 2000 and num < 2010: - return "two thousand " + _inflect.number_to_words(num % 100) + return "two thousand " + __inflect.number_to_words(num % 100) # type: ignore elif num % 100 == 0: - return _inflect.number_to_words(num // 100) + " hundred" + return __inflect.number_to_words(num // 100) + " hundred" # type: ignore else: - return _inflect.number_to_words( - num, andword="", zero="oh", group=2 - ).replace(", ", " ") + return __inflect.number_to_words( + num, andword="", zero="oh", group=2 # type: ignore + ).replace(", ", " ") # type: ignore else: - return _inflect.number_to_words(num, andword="") + return __inflect.number_to_words(num, andword="") # type: ignore -def _expand_decimal_point(m): +def __expand_decimal_point(m: re.Match[str]) -> str: return m.group(1).replace(".", " point ") -def normalize_numbers(text): - text = re.sub(_comma_number_re, _remove_commas, text) - text = re.sub(_pounds_re, r"\1 pounds", text) - text = re.sub(_dollars_re, _expand_dollars, text) - text = re.sub(_decimal_number_re, _expand_decimal_point, text) - text = re.sub(_ordinal_re, _expand_ordinal, text) - text = re.sub(_number_re, _expand_number, text) - return text - - -def normalize_text(text: str) -> str: - text = normalize_numbers(text) - text = replace_punctuation(text) - text = re.sub(r"([,;.\?\!])([\w])", r"\1 \2", text) - return text - - -def distribute_phone(n_phone, n_word): +def __distribute_phone(n_phone: int, n_word: int) -> list[int]: phones_per_word = [0] * n_word for task in range(n_phone): min_tasks = min(phones_per_word) @@ -384,13 +373,7 @@ def distribute_phone(n_phone, n_word): return phones_per_word -def sep_text(text): - words = re.split(r"([,;.\?\!\s+])", text) - words = [word for word in words if word.strip() != ""] - return words - - -def text_to_words(text): +def __text_to_words(text: str) -> list[list[str]]: tokenizer = bert_models.load_tokenizer(Languages.EN) tokens = tokenizer.tokenize(text) words = [] @@ -418,69 +401,12 @@ def text_to_words(text): return words -def g2p(text: str) -> tuple[list[str], list[int], list[int]]: - phones = [] - tones = [] - phone_len = [] - # words = sep_text(text) - # tokens = [tokenizer.tokenize(i) for i in words] - words = text_to_words(text) - - for word in words: - temp_phones, temp_tones = [], [] - if len(word) > 1: - if "'" in word: - word = ["".join(word)] - for w in word: - if w in PUNCTUATIONS: - temp_phones.append(w) - temp_tones.append(0) - continue - if w.upper() in eng_dict: - phns, tns = refine_syllables(eng_dict[w.upper()]) - temp_phones += [post_replace_ph(i) for i in phns] - temp_tones += tns - # w2ph.append(len(phns)) - else: - phone_list = list(filter(lambda p: p != " ", _g2p(w))) - phns = [] - tns = [] - for ph in phone_list: - if ph in arpa: - ph, tn = refine_ph(ph) - phns.append(ph) - tns.append(tn) - else: - phns.append(ph) - tns.append(0) - temp_phones += [post_replace_ph(i) for i in phns] - temp_tones += tns - phones += temp_phones - tones += temp_tones - phone_len.append(len(temp_phones)) - # phones = [post_replace_ph(i) for i in phones] - - word2ph = [] - for token, pl in zip(words, phone_len): - word_len = len(token) - - aaa = distribute_phone(pl, word_len) - word2ph += aaa - - phones = ["_"] + phones + ["_"] - tones = [0] + tones + [0] - word2ph = [1] + word2ph + [1] - assert len(phones) == len(tones), text - assert len(phones) == sum(word2ph), text - - return phones, tones, word2ph - - if __name__ == "__main__": # print(get_dict()) # print(eng_word_to_phoneme("hello")) print(g2p("In this paper, we propose 1 DSPGAN, a GAN-based universal vocoder.")) # all_phones = set() + # eng_dict = get_dict() # for k, syllables in eng_dict.items(): # for group in syllables: # for ph in group: From 5de4884075a4b40c1b20c95d39397726f31083df Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 07:31:33 +0000 Subject: [PATCH 29/64] Fix: import error pyopenjtalk_worker.initialize() has the side effect of starting another process and should not be executed automatically on import. --- bert_gen.py | 3 --- style_bert_vits2/nlp/english/__init__.py | 2 +- style_bert_vits2/nlp/japanese/g2p.py | 5 +++-- .../nlp/japanese/pyopenjtalk_worker/__init__.py | 14 +++++++------- .../nlp/japanese/user_dict/__init__.py | 7 +++++-- 5 files changed, 16 insertions(+), 15 deletions(-) diff --git a/bert_gen.py b/bert_gen.py index 79ea65f..26df64d 100644 --- a/bert_gen.py +++ b/bert_gen.py @@ -10,11 +10,8 @@ from config import config from style_bert_vits2.logging import logger from style_bert_vits2.models import commons from style_bert_vits2.nlp import cleaned_text_to_sequence, extract_bert_feature -from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT -pyopenjtalk.initialize() - def process_line(x): line, add_blank = x diff --git a/style_bert_vits2/nlp/english/__init__.py b/style_bert_vits2/nlp/english/__init__.py index e5041d1..b2067c3 100644 --- a/style_bert_vits2/nlp/english/__init__.py +++ b/style_bert_vits2/nlp/english/__init__.py @@ -273,7 +273,7 @@ def __cache_dict(g2p_dict: dict[str, list[list[str]]], file_path: Path) -> None: def __get_dict() -> dict[str, list[list[str]]]: - if os.path.exists(CACHE_PATH): + if CACHE_PATH.exists(): with open(CACHE_PATH, "rb") as pickle_file: g2p_dict = pickle.load(pickle_file) else: diff --git a/style_bert_vits2/nlp/japanese/g2p.py b/style_bert_vits2/nlp/japanese/g2p.py index 45e7817..0c44563 100644 --- a/style_bert_vits2/nlp/japanese/g2p.py +++ b/style_bert_vits2/nlp/japanese/g2p.py @@ -8,8 +8,6 @@ from style_bert_vits2.nlp.japanese.mora_list import MORA_KATA_TO_MORA_PHONEMES from style_bert_vits2.nlp.japanese.normalizer import replace_punctuation from style_bert_vits2.nlp.symbols import PUNCTUATIONS -pyopenjtalk.initialize() - def g2p( norm_text: str, @@ -114,6 +112,9 @@ def text_to_sep_kata( tuple[list[str], list[str]]: 分割された単語リストと、その読み(カタカナ or 記号1文字)のリスト """ + # pyopenjtalk_worker を初期化 + pyopenjtalk.initialize() + # parsed: OpenJTalkの解析結果 parsed = pyopenjtalk.run_frontend(norm_text) sep_text: list[str] = [] diff --git a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py index fc8f6da..90419f5 100644 --- a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py +++ b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py @@ -50,11 +50,11 @@ def unset_user_dict(): def initialize(port: int = WORKER_PORT) -> None: - import time - import socket - import sys import atexit import signal + import socket + import sys + import time logger.debug("initialize") global WORKER_CLIENT @@ -83,7 +83,7 @@ def initialize(port: int = WORKER_PORT) -> None: else: # align with Windows behavior # start_new_session is same as specifying setsid in preexec_fn - subprocess.Popen(args, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, start_new_session=True) # type: ignore + subprocess.Popen(args, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, start_new_session=True) # wait until server listening count = 0 @@ -92,10 +92,10 @@ def initialize(port: int = WORKER_PORT) -> None: client = WorkerClient(port) break except socket.error: - time.sleep(1) + time.sleep(0.5) count += 1 - # 10: max number of retries - if count == 10: + # 20: max number of retries + if count == 20: raise TimeoutError("サーバーに接続できませんでした") WORKER_CLIENT = client diff --git a/style_bert_vits2/nlp/japanese/user_dict/__init__.py b/style_bert_vits2/nlp/japanese/user_dict/__init__.py index 097cc6a..3032c01 100644 --- a/style_bert_vits2/nlp/japanese/user_dict/__init__.py +++ b/style_bert_vits2/nlp/japanese/user_dict/__init__.py @@ -20,8 +20,6 @@ from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk from style_bert_vits2.nlp.japanese.user_dict.word_model import UserDictWord, WordTypes from style_bert_vits2.nlp.japanese.user_dict.part_of_speech_data import MAX_PRIORITY, MIN_PRIORITY, part_of_speech_data -pyopenjtalk.initialize() - # root_dir = engine_root() # save_dir = get_save_dir() @@ -81,6 +79,11 @@ def update_dict( compiled_dict_path : Path コンパイル済み辞書ファイルのパス """ + + # pyopenjtalk_worker を初期化 + # ファイルを開く前に実行する必要がある + pyopenjtalk.initialize() + random_string = uuid4() tmp_csv_path = compiled_dict_path.with_suffix( f".dict_csv-{random_string}.tmp" From e2daa550002a13f578d2b09cb7dcd07ccae6d751 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 08:43:54 +0000 Subject: [PATCH 30/64] Refactor: split style_bert_vits2.nlp.english package --- style_bert_vits2/nlp/__init__.py | 20 +- style_bert_vits2/nlp/english/__init__.py | 414 ------------------- style_bert_vits2/nlp/english/bert_feature.py | 6 +- style_bert_vits2/nlp/english/cmudict.py | 46 +++ style_bert_vits2/nlp/english/g2p.py | 240 +++++++++++ style_bert_vits2/nlp/english/normalizer.py | 130 ++++++ 6 files changed, 432 insertions(+), 424 deletions(-) create mode 100644 style_bert_vits2/nlp/english/cmudict.py create mode 100644 style_bert_vits2/nlp/english/g2p.py create mode 100644 style_bert_vits2/nlp/english/normalizer.py diff --git a/style_bert_vits2/nlp/__init__.py b/style_bert_vits2/nlp/__init__.py index 47786d1..99f56a9 100644 --- a/style_bert_vits2/nlp/__init__.py +++ b/style_bert_vits2/nlp/__init__.py @@ -1,5 +1,4 @@ -import torch -from typing import Optional +from typing import Optional, TYPE_CHECKING from style_bert_vits2.constants import Languages from style_bert_vits2.nlp.symbols import ( @@ -8,6 +7,11 @@ from style_bert_vits2.nlp.symbols import ( SYMBOLS, ) +# __init__.py は配下のモジュールをインポートした時点で実行される +# Pytorch のインポートは重いので、型チェック時以外はインポートしない +if TYPE_CHECKING: + import torch + __symbol_to_id = {s: i for i, s in enumerate(SYMBOLS)} @@ -16,10 +20,10 @@ def extract_bert_feature( text: str, word2ph: list[int], language: Languages, - device: torch.device | str, + device: str, assist_text: Optional[str] = None, assist_text_weight: float = 0.7, -) -> torch.Tensor: +) -> "torch.Tensor": """ テキストから BERT の特徴量を抽出する @@ -27,7 +31,7 @@ def extract_bert_feature( text (str): テキスト word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト language (Languages): テキストの言語 - device (torch.device | str): 推論に利用するデバイス + device (str): 推論に利用するデバイス assist_text (Optional[str], optional): 補助テキスト (デフォルト: None) assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7) @@ -68,11 +72,13 @@ def clean_text( # Changed to import inside if condition to avoid unnecessary import if language == Languages.JP: - from style_bert_vits2.nlp.japanese import g2p, normalize_text + from style_bert_vits2.nlp.japanese.g2p import g2p + from style_bert_vits2.nlp.japanese.normalizer import normalize_text norm_text = normalize_text(text) phones, tones, word2ph = g2p(norm_text, use_jp_extra, raise_yomi_error) elif language == Languages.EN: - from style_bert_vits2.nlp.english import g2p, normalize_text + from style_bert_vits2.nlp.english.g2p import g2p + from style_bert_vits2.nlp.english.normalizer import normalize_text norm_text = normalize_text(text) phones, tones, word2ph = g2p(norm_text) elif language == Languages.ZH: diff --git a/style_bert_vits2/nlp/english/__init__.py b/style_bert_vits2/nlp/english/__init__.py index b2067c3..e69de29 100644 --- a/style_bert_vits2/nlp/english/__init__.py +++ b/style_bert_vits2/nlp/english/__init__.py @@ -1,414 +0,0 @@ -import pickle -import os -import re -from pathlib import Path - -import inflect -from g2p_en import G2p - -from style_bert_vits2.constants import Languages -from style_bert_vits2.nlp import bert_models -from style_bert_vits2.nlp.symbols import PUNCTUATIONS, SYMBOLS - - -CMU_DICT_PATH = Path(__file__).parent / "cmudict.rep" -CACHE_PATH = Path(__file__).parent / "cmudict_cache.pickle" - - -def g2p(text: str) -> tuple[list[str], list[int], list[int]]: - - ARPA = { - "AH0", - "S", - "AH1", - "EY2", - "AE2", - "EH0", - "OW2", - "UH0", - "NG", - "B", - "G", - "AY0", - "M", - "AA0", - "F", - "AO0", - "ER2", - "UH1", - "IY1", - "AH2", - "DH", - "IY0", - "EY1", - "IH0", - "K", - "N", - "W", - "IY2", - "T", - "AA1", - "ER1", - "EH2", - "OY0", - "UH2", - "UW1", - "Z", - "AW2", - "AW1", - "V", - "UW2", - "AA2", - "ER", - "AW0", - "UW0", - "R", - "OW1", - "EH1", - "ZH", - "AE0", - "IH2", - "IH", - "Y", - "JH", - "P", - "AY1", - "EY0", - "OY2", - "TH", - "HH", - "D", - "ER0", - "CH", - "AO1", - "AE1", - "AO2", - "OY1", - "AY2", - "IH1", - "OW0", - "L", - "SH", - } - - _g2p = G2p() - - phones = [] - tones = [] - phone_len = [] - # tokens = [tokenizer.tokenize(i) for i in words] - words = __text_to_words(text) - eng_dict = __get_dict() - - for word in words: - temp_phones, temp_tones = [], [] - if len(word) > 1: - if "'" in word: - word = ["".join(word)] - for w in word: - if w in PUNCTUATIONS: - temp_phones.append(w) - temp_tones.append(0) - continue - if w.upper() in eng_dict: - phns, tns = __refine_syllables(eng_dict[w.upper()]) - temp_phones += [__post_replace_ph(i) for i in phns] - temp_tones += tns - # w2ph.append(len(phns)) - else: - phone_list = list(filter(lambda p: p != " ", _g2p(w))) # type: ignore - phns = [] - tns = [] - for ph in phone_list: - if ph in ARPA: - ph, tn = __refine_ph(ph) - phns.append(ph) - tns.append(tn) - else: - phns.append(ph) - tns.append(0) - temp_phones += [__post_replace_ph(i) for i in phns] - temp_tones += tns - phones += temp_phones - tones += temp_tones - phone_len.append(len(temp_phones)) - # phones = [post_replace_ph(i) for i in phones] - - word2ph = [] - for token, pl in zip(words, phone_len): - word_len = len(token) - - aaa = __distribute_phone(pl, word_len) - word2ph += aaa - - phones = ["_"] + phones + ["_"] - tones = [0] + tones + [0] - word2ph = [1] + word2ph + [1] - assert len(phones) == len(tones), text - assert len(phones) == sum(word2ph), text - - return phones, tones, word2ph - - -def normalize_text(text: str) -> str: - text = __normalize_numbers(text) - text = __replace_punctuation(text) - text = re.sub(r"([,;.\?\!])([\w])", r"\1 \2", text) - return text - - -def __normalize_numbers(text: str) -> str: - text = re.sub(__comma_number_re, __remove_commas, text) - text = re.sub(__pounds_re, r"\1 pounds", text) - text = re.sub(__dollars_re, __expand_dollars, text) - text = re.sub(__decimal_number_re, __expand_decimal_point, text) - text = re.sub(__ordinal_re, __expand_ordinal, text) - text = re.sub(__number_re, __expand_number, text) - return text - - -def __replace_punctuation(text: str) -> str: - REPLACE_MAP = { - ":": ",", - ";": ",", - ",": ",", - "。": ".", - "!": "!", - "?": "?", - "\n": ".", - ".": ".", - "…": "...", - "···": "...", - "・・・": "...", - "·": ",", - "・": ",", - "、": ",", - "$": ".", - "“": "'", - "”": "'", - '"': "'", - "‘": "'", - "’": "'", - "(": "'", - ")": "'", - "(": "'", - ")": "'", - "《": "'", - "》": "'", - "【": "'", - "】": "'", - "[": "'", - "]": "'", - "—": "-", - "−": "-", - "~": "-", - "~": "-", - "「": "'", - "」": "'", - } - pattern = re.compile("|".join(re.escape(p) for p in REPLACE_MAP.keys())) - replaced_text = pattern.sub(lambda x: REPLACE_MAP[x.group()], text) - # replaced_text = re.sub( - # r"[^\u3040-\u309F\u30A0-\u30FF\u4E00-\u9FFF\u3400-\u4DBF\u3005" - # + "".join(punctuation) - # + r"]+", - # "", - # replaced_text, - # ) - return replaced_text - - -def __post_replace_ph(ph: str) -> str: - REPLACE_MAP = { - ":": ",", - ";": ",", - ",": ",", - "。": ".", - "!": "!", - "?": "?", - "\n": ".", - "·": ",", - "、": ",", - "…": "...", - "···": "...", - "・・・": "...", - "v": "V", - } - if ph in REPLACE_MAP.keys(): - ph = REPLACE_MAP[ph] - if ph in SYMBOLS: - return ph - if ph not in SYMBOLS: - ph = "UNK" - return ph - - -def __read_dict() -> dict[str, list[list[str]]]: - g2p_dict = {} - start_line = 49 - with open(CMU_DICT_PATH) as f: - line = f.readline() - line_index = 1 - while line: - if line_index >= start_line: - line = line.strip() - word_split = line.split(" ") - word = word_split[0] - - syllable_split = word_split[1].split(" - ") - g2p_dict[word] = [] - for syllable in syllable_split: - phone_split = syllable.split(" ") - g2p_dict[word].append(phone_split) - - line_index = line_index + 1 - line = f.readline() - - return g2p_dict - - -def __cache_dict(g2p_dict: dict[str, list[list[str]]], file_path: Path) -> None: - with open(file_path, "wb") as pickle_file: - pickle.dump(g2p_dict, pickle_file) - - -def __get_dict() -> dict[str, list[list[str]]]: - if CACHE_PATH.exists(): - with open(CACHE_PATH, "rb") as pickle_file: - g2p_dict = pickle.load(pickle_file) - else: - g2p_dict = __read_dict() - __cache_dict(g2p_dict, CACHE_PATH) - - return g2p_dict - - -def __refine_ph(phn: str) -> tuple[str, int]: - tone = 0 - if re.search(r"\d$", phn): - tone = int(phn[-1]) + 1 - phn = phn[:-1] - else: - tone = 3 - return phn.lower(), tone - - -def __refine_syllables(syllables: list[list[str]]) -> tuple[list[str], list[int]]: - tones = [] - phonemes = [] - for phn_list in syllables: - for i in range(len(phn_list)): - phn = phn_list[i] - phn, tone = __refine_ph(phn) - phonemes.append(phn) - tones.append(tone) - return phonemes, tones - - -__inflect = inflect.engine() -__comma_number_re = re.compile(r"([0-9][0-9\,]+[0-9])") -__decimal_number_re = re.compile(r"([0-9]+\.[0-9]+)") -__pounds_re = re.compile(r"£([0-9\,]*[0-9]+)") -__dollars_re = re.compile(r"\$([0-9\.\,]*[0-9]+)") -__ordinal_re = re.compile(r"[0-9]+(st|nd|rd|th)") -__number_re = re.compile(r"[0-9]+") - - -def __expand_dollars(m: re.Match[str]) -> str: - match = m.group(1) - parts = match.split(".") - if len(parts) > 2: - return match + " dollars" # Unexpected format - dollars = int(parts[0]) if parts[0] else 0 - cents = int(parts[1]) if len(parts) > 1 and parts[1] else 0 - if dollars and cents: - dollar_unit = "dollar" if dollars == 1 else "dollars" - cent_unit = "cent" if cents == 1 else "cents" - return "%s %s, %s %s" % (dollars, dollar_unit, cents, cent_unit) - elif dollars: - dollar_unit = "dollar" if dollars == 1 else "dollars" - return "%s %s" % (dollars, dollar_unit) - elif cents: - cent_unit = "cent" if cents == 1 else "cents" - return "%s %s" % (cents, cent_unit) - else: - return "zero dollars" - - -def __remove_commas(m: re.Match[str]) -> str: - return m.group(1).replace(",", "") - - -def __expand_ordinal(m: re.Match[str]) -> str: - return __inflect.number_to_words(m.group(0)) # type: ignore - - -def __expand_number(m: re.Match[str]) -> str: - num = int(m.group(0)) - if num > 1000 and num < 3000: - if num == 2000: - return "two thousand" - elif num > 2000 and num < 2010: - return "two thousand " + __inflect.number_to_words(num % 100) # type: ignore - elif num % 100 == 0: - return __inflect.number_to_words(num // 100) + " hundred" # type: ignore - else: - return __inflect.number_to_words( - num, andword="", zero="oh", group=2 # type: ignore - ).replace(", ", " ") # type: ignore - else: - return __inflect.number_to_words(num, andword="") # type: ignore - - -def __expand_decimal_point(m: re.Match[str]) -> str: - return m.group(1).replace(".", " point ") - - -def __distribute_phone(n_phone: int, n_word: int) -> list[int]: - phones_per_word = [0] * n_word - for task in range(n_phone): - min_tasks = min(phones_per_word) - min_index = phones_per_word.index(min_tasks) - phones_per_word[min_index] += 1 - return phones_per_word - - -def __text_to_words(text: str) -> list[list[str]]: - tokenizer = bert_models.load_tokenizer(Languages.EN) - tokens = tokenizer.tokenize(text) - words = [] - for idx, t in enumerate(tokens): - if t.startswith("▁"): - words.append([t[1:]]) - else: - if t in PUNCTUATIONS: - if idx == len(tokens) - 1: - words.append([f"{t}"]) - else: - if ( - not tokens[idx + 1].startswith("▁") - and tokens[idx + 1] not in PUNCTUATIONS - ): - if idx == 0: - words.append([]) - words[-1].append(f"{t}") - else: - words.append([f"{t}"]) - else: - if idx == 0: - words.append([]) - words[-1].append(f"{t}") - return words - - -if __name__ == "__main__": - # print(get_dict()) - # print(eng_word_to_phoneme("hello")) - print(g2p("In this paper, we propose 1 DSPGAN, a GAN-based universal vocoder.")) - # all_phones = set() - # eng_dict = get_dict() - # for k, syllables in eng_dict.items(): - # for group in syllables: - # for ph in group: - # all_phones.add(ph) - # print(all_phones) diff --git a/style_bert_vits2/nlp/english/bert_feature.py b/style_bert_vits2/nlp/english/bert_feature.py index 647920d..27fd501 100644 --- a/style_bert_vits2/nlp/english/bert_feature.py +++ b/style_bert_vits2/nlp/english/bert_feature.py @@ -8,13 +8,13 @@ from style_bert_vits2.constants import Languages from style_bert_vits2.nlp import bert_models -__models: dict[torch.device | str, PreTrainedModel] = {} +__models: dict[str, PreTrainedModel] = {} def extract_bert_feature( text: str, word2ph: list[int], - device: torch.device | str, + device: str, assist_text: Optional[str] = None, assist_text_weight: float = 0.7, ) -> torch.Tensor: @@ -24,7 +24,7 @@ def extract_bert_feature( Args: text (str): 英語のテキスト word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト - device (torch.device | str): 推論に利用するデバイス + device (str): 推論に利用するデバイス assist_text (Optional[str], optional): 補助テキスト (デフォルト: None) assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7) diff --git a/style_bert_vits2/nlp/english/cmudict.py b/style_bert_vits2/nlp/english/cmudict.py new file mode 100644 index 0000000..e6afb89 --- /dev/null +++ b/style_bert_vits2/nlp/english/cmudict.py @@ -0,0 +1,46 @@ +import pickle +from pathlib import Path + + +CMU_DICT_PATH = Path(__file__).parent / "cmudict.rep" +CACHE_PATH = Path(__file__).parent / "cmudict_cache.pickle" + + +def get_dict() -> dict[str, list[list[str]]]: + if CACHE_PATH.exists(): + with open(CACHE_PATH, "rb") as pickle_file: + g2p_dict = pickle.load(pickle_file) + else: + g2p_dict = read_dict() + cache_dict(g2p_dict, CACHE_PATH) + + return g2p_dict + + +def read_dict() -> dict[str, list[list[str]]]: + g2p_dict = {} + start_line = 49 + with open(CMU_DICT_PATH) as f: + line = f.readline() + line_index = 1 + while line: + if line_index >= start_line: + line = line.strip() + word_split = line.split(" ") + word = word_split[0] + + syllable_split = word_split[1].split(" - ") + g2p_dict[word] = [] + for syllable in syllable_split: + phone_split = syllable.split(" ") + g2p_dict[word].append(phone_split) + + line_index = line_index + 1 + line = f.readline() + + return g2p_dict + + +def cache_dict(g2p_dict: dict[str, list[list[str]]], file_path: Path) -> None: + with open(file_path, "wb") as pickle_file: + pickle.dump(g2p_dict, pickle_file) diff --git a/style_bert_vits2/nlp/english/g2p.py b/style_bert_vits2/nlp/english/g2p.py new file mode 100644 index 0000000..db2a87f --- /dev/null +++ b/style_bert_vits2/nlp/english/g2p.py @@ -0,0 +1,240 @@ +import re + +from g2p_en import G2p + +from style_bert_vits2.constants import Languages +from style_bert_vits2.nlp import bert_models +from style_bert_vits2.nlp.english.cmudict import get_dict +from style_bert_vits2.nlp.symbols import PUNCTUATIONS, SYMBOLS + + +def g2p(text: str) -> tuple[list[str], list[int], list[int]]: + + ARPA = { + "AH0", + "S", + "AH1", + "EY2", + "AE2", + "EH0", + "OW2", + "UH0", + "NG", + "B", + "G", + "AY0", + "M", + "AA0", + "F", + "AO0", + "ER2", + "UH1", + "IY1", + "AH2", + "DH", + "IY0", + "EY1", + "IH0", + "K", + "N", + "W", + "IY2", + "T", + "AA1", + "ER1", + "EH2", + "OY0", + "UH2", + "UW1", + "Z", + "AW2", + "AW1", + "V", + "UW2", + "AA2", + "ER", + "AW0", + "UW0", + "R", + "OW1", + "EH1", + "ZH", + "AE0", + "IH2", + "IH", + "Y", + "JH", + "P", + "AY1", + "EY0", + "OY2", + "TH", + "HH", + "D", + "ER0", + "CH", + "AO1", + "AE1", + "AO2", + "OY1", + "AY2", + "IH1", + "OW0", + "L", + "SH", + } + + _g2p = G2p() + + phones = [] + tones = [] + phone_len = [] + # tokens = [tokenizer.tokenize(i) for i in words] + words = __text_to_words(text) + eng_dict = get_dict() + + for word in words: + temp_phones, temp_tones = [], [] + if len(word) > 1: + if "'" in word: + word = ["".join(word)] + for w in word: + if w in PUNCTUATIONS: + temp_phones.append(w) + temp_tones.append(0) + continue + if w.upper() in eng_dict: + phns, tns = __refine_syllables(eng_dict[w.upper()]) + temp_phones += [__post_replace_ph(i) for i in phns] + temp_tones += tns + # w2ph.append(len(phns)) + else: + phone_list = list(filter(lambda p: p != " ", _g2p(w))) # type: ignore + phns = [] + tns = [] + for ph in phone_list: + if ph in ARPA: + ph, tn = __refine_ph(ph) + phns.append(ph) + tns.append(tn) + else: + phns.append(ph) + tns.append(0) + temp_phones += [__post_replace_ph(i) for i in phns] + temp_tones += tns + phones += temp_phones + tones += temp_tones + phone_len.append(len(temp_phones)) + # phones = [post_replace_ph(i) for i in phones] + + word2ph = [] + for token, pl in zip(words, phone_len): + word_len = len(token) + + aaa = __distribute_phone(pl, word_len) + word2ph += aaa + + phones = ["_"] + phones + ["_"] + tones = [0] + tones + [0] + word2ph = [1] + word2ph + [1] + assert len(phones) == len(tones), text + assert len(phones) == sum(word2ph), text + + return phones, tones, word2ph + + +def __post_replace_ph(ph: str) -> str: + REPLACE_MAP = { + ":": ",", + ";": ",", + ",": ",", + "。": ".", + "!": "!", + "?": "?", + "\n": ".", + "·": ",", + "、": ",", + "…": "...", + "···": "...", + "・・・": "...", + "v": "V", + } + if ph in REPLACE_MAP.keys(): + ph = REPLACE_MAP[ph] + if ph in SYMBOLS: + return ph + if ph not in SYMBOLS: + ph = "UNK" + return ph + + +def __refine_ph(phn: str) -> tuple[str, int]: + tone = 0 + if re.search(r"\d$", phn): + tone = int(phn[-1]) + 1 + phn = phn[:-1] + else: + tone = 3 + return phn.lower(), tone + + +def __refine_syllables(syllables: list[list[str]]) -> tuple[list[str], list[int]]: + tones = [] + phonemes = [] + for phn_list in syllables: + for i in range(len(phn_list)): + phn = phn_list[i] + phn, tone = __refine_ph(phn) + phonemes.append(phn) + tones.append(tone) + return phonemes, tones + + +def __distribute_phone(n_phone: int, n_word: int) -> list[int]: + phones_per_word = [0] * n_word + for task in range(n_phone): + min_tasks = min(phones_per_word) + min_index = phones_per_word.index(min_tasks) + phones_per_word[min_index] += 1 + return phones_per_word + + +def __text_to_words(text: str) -> list[list[str]]: + tokenizer = bert_models.load_tokenizer(Languages.EN) + tokens = tokenizer.tokenize(text) + words = [] + for idx, t in enumerate(tokens): + if t.startswith("▁"): + words.append([t[1:]]) + else: + if t in PUNCTUATIONS: + if idx == len(tokens) - 1: + words.append([f"{t}"]) + else: + if ( + not tokens[idx + 1].startswith("▁") + and tokens[idx + 1] not in PUNCTUATIONS + ): + if idx == 0: + words.append([]) + words[-1].append(f"{t}") + else: + words.append([f"{t}"]) + else: + if idx == 0: + words.append([]) + words[-1].append(f"{t}") + return words + + +if __name__ == "__main__": + # print(get_dict()) + # print(eng_word_to_phoneme("hello")) + print(g2p("In this paper, we propose 1 DSPGAN, a GAN-based universal vocoder.")) + # all_phones = set() + # eng_dict = get_dict() + # for k, syllables in eng_dict.items(): + # for group in syllables: + # for ph in group: + # all_phones.add(ph) + # print(all_phones) diff --git a/style_bert_vits2/nlp/english/normalizer.py b/style_bert_vits2/nlp/english/normalizer.py new file mode 100644 index 0000000..0886581 --- /dev/null +++ b/style_bert_vits2/nlp/english/normalizer.py @@ -0,0 +1,130 @@ +import re + +import inflect + + +__INFLECT = inflect.engine() +__COMMA_NUMBER_PATTERN = re.compile(r"([0-9][0-9\,]+[0-9])") +__DECIMAL_NUMBER_PATTERN = re.compile(r"([0-9]+\.[0-9]+)") +__POUNDS_PATTERN = re.compile(r"£([0-9\,]*[0-9]+)") +__DOLLARS_PATTERN = re.compile(r"\$([0-9\.\,]*[0-9]+)") +__ORDINAL_PATTERN = re.compile(r"[0-9]+(st|nd|rd|th)") +__NUMBER_PATTERN = re.compile(r"[0-9]+") + + +def normalize_text(text: str) -> str: + text = __normalize_numbers(text) + text = __replace_punctuation(text) + text = re.sub(r"([,;.\?\!])([\w])", r"\1 \2", text) + return text + + +def __normalize_numbers(text: str) -> str: + text = re.sub(__COMMA_NUMBER_PATTERN, __remove_commas, text) + text = re.sub(__POUNDS_PATTERN, r"\1 pounds", text) + text = re.sub(__DOLLARS_PATTERN, __expand_dollars, text) + text = re.sub(__DECIMAL_NUMBER_PATTERN, __expand_decimal_point, text) + text = re.sub(__ORDINAL_PATTERN, __expand_ordinal, text) + text = re.sub(__NUMBER_PATTERN, __expand_number, text) + return text + + +def __replace_punctuation(text: str) -> str: + REPLACE_MAP = { + ":": ",", + ";": ",", + ",": ",", + "。": ".", + "!": "!", + "?": "?", + "\n": ".", + ".": ".", + "…": "...", + "···": "...", + "・・・": "...", + "·": ",", + "・": ",", + "、": ",", + "$": ".", + "“": "'", + "”": "'", + '"': "'", + "‘": "'", + "’": "'", + "(": "'", + ")": "'", + "(": "'", + ")": "'", + "《": "'", + "》": "'", + "【": "'", + "】": "'", + "[": "'", + "]": "'", + "—": "-", + "−": "-", + "~": "-", + "~": "-", + "「": "'", + "」": "'", + } + pattern = re.compile("|".join(re.escape(p) for p in REPLACE_MAP.keys())) + replaced_text = pattern.sub(lambda x: REPLACE_MAP[x.group()], text) + # replaced_text = re.sub( + # r"[^\u3040-\u309F\u30A0-\u30FF\u4E00-\u9FFF\u3400-\u4DBF\u3005" + # + "".join(punctuation) + # + r"]+", + # "", + # replaced_text, + # ) + return replaced_text + + +def __expand_dollars(m: re.Match[str]) -> str: + match = m.group(1) + parts = match.split(".") + if len(parts) > 2: + return match + " dollars" # Unexpected format + dollars = int(parts[0]) if parts[0] else 0 + cents = int(parts[1]) if len(parts) > 1 and parts[1] else 0 + if dollars and cents: + dollar_unit = "dollar" if dollars == 1 else "dollars" + cent_unit = "cent" if cents == 1 else "cents" + return "%s %s, %s %s" % (dollars, dollar_unit, cents, cent_unit) + elif dollars: + dollar_unit = "dollar" if dollars == 1 else "dollars" + return "%s %s" % (dollars, dollar_unit) + elif cents: + cent_unit = "cent" if cents == 1 else "cents" + return "%s %s" % (cents, cent_unit) + else: + return "zero dollars" + + +def __remove_commas(m: re.Match[str]) -> str: + return m.group(1).replace(",", "") + + +def __expand_ordinal(m: re.Match[str]) -> str: + return __INFLECT.number_to_words(m.group(0)) # type: ignore + + +def __expand_number(m: re.Match[str]) -> str: + num = int(m.group(0)) + if num > 1000 and num < 3000: + if num == 2000: + return "two thousand" + elif num > 2000 and num < 2010: + return "two thousand " + __INFLECT.number_to_words(num % 100) # type: ignore + elif num % 100 == 0: + return __INFLECT.number_to_words(num // 100) + " hundred" # type: ignore + else: + return __INFLECT.number_to_words( + num, andword="", zero="oh", group=2 # type: ignore + ).replace(", ", " ") # type: ignore + else: + return __INFLECT.number_to_words(num, andword="") # type: ignore + + +def __expand_decimal_point(m: re.Match[str]) -> str: + return m.group(1).replace(".", " point ") From a672aeefd93188de76e9978cd9487bed0f91d91d Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 08:45:56 +0000 Subject: [PATCH 31/64] Fix: maintain compatibility with Python 3.9 --- server_editor.py | 2 +- style_bert_vits2/models/commons.py | 6 +++--- style_bert_vits2/nlp/bert_models.py | 14 +++++++------- style_bert_vits2/nlp/chinese/bert_feature.py | 6 +++--- style_bert_vits2/nlp/japanese/__init__.py | 2 -- style_bert_vits2/nlp/japanese/bert_feature.py | 6 +++--- 6 files changed, 17 insertions(+), 19 deletions(-) diff --git a/server_editor.py b/server_editor.py index f72800d..890a074 100644 --- a/server_editor.py +++ b/server_editor.py @@ -43,8 +43,8 @@ from style_bert_vits2.constants import ( ) from style_bert_vits2.logging import logger from style_bert_vits2.nlp import bert_models -from style_bert_vits2.nlp.japanese import normalize_text from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone +from style_bert_vits2.nlp.japanese.normalizer import normalize_text from style_bert_vits2.nlp.japanese.user_dict import ( apply_word, delete_word, diff --git a/style_bert_vits2/models/commons.py b/style_bert_vits2/models/commons.py index 89d07d5..1106b21 100644 --- a/style_bert_vits2/models/commons.py +++ b/style_bert_vits2/models/commons.py @@ -5,7 +5,7 @@ import torch from torch.nn import functional as F -from typing import Any, Optional +from typing import Any, Optional, Union def init_weights(m: torch.nn.Module, mean: float = 0.0, std: float = 0.01) -> None: @@ -180,12 +180,12 @@ def generate_path(duration: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: return path -def clip_grad_value_(parameters: torch.Tensor | list[torch.Tensor], clip_value: Optional[float], norm_type: float = 2.0) -> float: +def clip_grad_value_(parameters: Union[torch.Tensor, list[torch.Tensor]], clip_value: Optional[float], norm_type: float = 2.0) -> float: """ 勾配の値をクリップする Args: - parameters (torch.Tensor | list[torch.Tensor]): クリップするパラメータ + parameters (Union[torch.Tensor, list[torch.Tensor]]): クリップするパラメータ clip_value (Optional[float]): クリップする値。None の場合はクリップしない norm_type (float): ノルムの種類 diff --git a/style_bert_vits2/nlp/bert_models.py b/style_bert_vits2/nlp/bert_models.py index df7a016..4385d64 100644 --- a/style_bert_vits2/nlp/bert_models.py +++ b/style_bert_vits2/nlp/bert_models.py @@ -9,7 +9,7 @@ Style-Bert-VITS2 の学習・推論に必要な各言語ごとの BERT モデル """ import gc -from typing import cast, Optional +from typing import cast, Optional, Union import torch from transformers import ( @@ -27,16 +27,16 @@ from style_bert_vits2.logging import logger # 各言語ごとのロード済みの BERT モデルを格納する辞書 -__loaded_models: dict[Languages, PreTrainedModel | DebertaV2Model] = {} +__loaded_models: dict[Languages, Union[PreTrainedModel, DebertaV2Model]] = {} # 各言語ごとのロード済みの BERT トークナイザーを格納する辞書 -__loaded_tokenizers: dict[Languages, PreTrainedTokenizer | PreTrainedTokenizerFast | DebertaV2Tokenizer] = {} +__loaded_tokenizers: dict[Languages, Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]] = {} def load_model( language: Languages, pretrained_model_name_or_path: Optional[str] = None, -) -> PreTrainedModel | DebertaV2Model: +) -> Union[PreTrainedModel, DebertaV2Model]: """ 指定された言語の BERT モデルをロードし、ロード済みの BERT モデルを返す。 一度ロードされていれば、ロード済みの BERT モデルを即座に返す。 @@ -54,7 +54,7 @@ def load_model( pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None) Returns: - PreTrainedModel | DebertaV2Model: ロード済みの BERT モデル + Union[PreTrainedModel, DebertaV2Model]: ロード済みの BERT モデル """ # すでにロード済みの場合はそのまま返す @@ -82,7 +82,7 @@ def load_model( def load_tokenizer( language: Languages, pretrained_model_name_or_path: Optional[str] = None, -) -> PreTrainedTokenizer | PreTrainedTokenizerFast | DebertaV2Tokenizer: +) -> Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]: """ 指定された言語の BERT モデルをロードし、ロード済みの BERT トークナイザーを返す。 一度ロードされていれば、ロード済みの BERT トークナイザーを即座に返す。 @@ -100,7 +100,7 @@ def load_tokenizer( pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None) Returns: - PreTrainedTokenizer | PreTrainedTokenizerFast | DebertaV2Tokenizer: ロード済みの BERT トークナイザー + Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]: ロード済みの BERT トークナイザー """ # すでにロード済みの場合はそのまま返す diff --git a/style_bert_vits2/nlp/chinese/bert_feature.py b/style_bert_vits2/nlp/chinese/bert_feature.py index b97950b..f448b30 100644 --- a/style_bert_vits2/nlp/chinese/bert_feature.py +++ b/style_bert_vits2/nlp/chinese/bert_feature.py @@ -8,13 +8,13 @@ from style_bert_vits2.constants import Languages from style_bert_vits2.nlp import bert_models -__models: dict[torch.device | str, PreTrainedModel] = {} +__models: dict[str, PreTrainedModel] = {} def extract_bert_feature( text: str, word2ph: list[int], - device: torch.device | str, + device: str, assist_text: Optional[str] = None, assist_text_weight: float = 0.7, ) -> torch.Tensor: @@ -24,7 +24,7 @@ def extract_bert_feature( Args: text (str): 中国語のテキスト word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト - device (torch.device | str): 推論に利用するデバイス + device (str): 推論に利用するデバイス assist_text (Optional[str], optional): 補助テキスト (デフォルト: None) assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7) diff --git a/style_bert_vits2/nlp/japanese/__init__.py b/style_bert_vits2/nlp/japanese/__init__.py index 5c7f19f..e69de29 100644 --- a/style_bert_vits2/nlp/japanese/__init__.py +++ b/style_bert_vits2/nlp/japanese/__init__.py @@ -1,2 +0,0 @@ -from style_bert_vits2.nlp.japanese.g2p import g2p # noqa: F401 -from style_bert_vits2.nlp.japanese.normalizer import normalize_text # noqa: F401 diff --git a/style_bert_vits2/nlp/japanese/bert_feature.py b/style_bert_vits2/nlp/japanese/bert_feature.py index ede1f83..0d70014 100644 --- a/style_bert_vits2/nlp/japanese/bert_feature.py +++ b/style_bert_vits2/nlp/japanese/bert_feature.py @@ -9,13 +9,13 @@ from style_bert_vits2.nlp import bert_models from style_bert_vits2.nlp.japanese.g2p import text_to_sep_kata -__models: dict[torch.device | str, PreTrainedModel] = {} +__models: dict[str, PreTrainedModel] = {} def extract_bert_feature( text: str, word2ph: list[int], - device: torch.device | str, + device: str, assist_text: Optional[str] = None, assist_text_weight: float = 0.7, ) -> torch.Tensor: @@ -25,7 +25,7 @@ def extract_bert_feature( Args: text (str): 日本語のテキスト word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト - device (torch.device | str): 推論に利用するデバイス + device (str): 推論に利用するデバイス assist_text (Optional[str], optional): 補助テキスト (デフォルト: None) assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7) From df687716512a8e5f0c89e9dda959c75348f6fb0e Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 09:04:40 +0000 Subject: [PATCH 32/64] Refactor: split style_bert_vits2.nlp.chinese package Configured so that the same public function is exported from the module with the same name for each language. --- style_bert_vits2/nlp/__init__.py | 3 +- style_bert_vits2/nlp/chinese/__init__.py | 193 ------------------ style_bert_vits2/nlp/chinese/g2p.py | 135 ++++++++++++ style_bert_vits2/nlp/chinese/normalizer.py | 61 ++++++ .../{english => chinese}/opencpop-strict.txt | 0 style_bert_vits2/nlp/english/normalizer.py | 24 +-- 6 files changed, 210 insertions(+), 206 deletions(-) create mode 100644 style_bert_vits2/nlp/chinese/g2p.py create mode 100644 style_bert_vits2/nlp/chinese/normalizer.py rename style_bert_vits2/nlp/{english => chinese}/opencpop-strict.txt (100%) diff --git a/style_bert_vits2/nlp/__init__.py b/style_bert_vits2/nlp/__init__.py index 99f56a9..afc6cb0 100644 --- a/style_bert_vits2/nlp/__init__.py +++ b/style_bert_vits2/nlp/__init__.py @@ -82,7 +82,8 @@ def clean_text( norm_text = normalize_text(text) phones, tones, word2ph = g2p(norm_text) elif language == Languages.ZH: - from style_bert_vits2.nlp.chinese import g2p, normalize_text + from style_bert_vits2.nlp.chinese.g2p import g2p + from style_bert_vits2.nlp.chinese.normalizer import normalize_text norm_text = normalize_text(text) phones, tones, word2ph = g2p(norm_text) else: diff --git a/style_bert_vits2/nlp/chinese/__init__.py b/style_bert_vits2/nlp/chinese/__init__.py index 5b70863..e69de29 100644 --- a/style_bert_vits2/nlp/chinese/__init__.py +++ b/style_bert_vits2/nlp/chinese/__init__.py @@ -1,193 +0,0 @@ -import os -import re - -import cn2an -import jieba.posseg as psg -from pypinyin import lazy_pinyin, Style - -from style_bert_vits2.nlp.chinese.tone_sandhi import ToneSandhi -from style_bert_vits2.nlp.symbols import PUNCTUATIONS - - -current_file_path = os.path.dirname(__file__) -pinyin_to_symbol_map = { - line.split("\t")[0]: line.strip().split("\t")[1] - for line in open(os.path.join(current_file_path, "opencpop-strict.txt")).readlines() -} - - -rep_map = { - ":": ",", - ";": ",", - ",": ",", - "。": ".", - "!": "!", - "?": "?", - "\n": ".", - "·": ",", - "、": ",", - "...": "…", - "$": ".", - "“": "'", - "”": "'", - '"': "'", - "‘": "'", - "’": "'", - "(": "'", - ")": "'", - "(": "'", - ")": "'", - "《": "'", - "》": "'", - "【": "'", - "】": "'", - "[": "'", - "]": "'", - "—": "-", - "~": "-", - "~": "-", - "「": "'", - "」": "'", -} - -tone_modifier = ToneSandhi() - - -def replace_punctuation(text): - text = text.replace("嗯", "恩").replace("呣", "母") - pattern = re.compile("|".join(re.escape(p) for p in rep_map.keys())) - - replaced_text = pattern.sub(lambda x: rep_map[x.group()], text) - - replaced_text = re.sub( - r"[^\u4e00-\u9fa5" + "".join(PUNCTUATIONS) + r"]+", "", replaced_text - ) - - return replaced_text - - -def g2p(text: str) -> tuple[list[str], list[int], list[int]]: - pattern = r"(?<=[{0}])\s*".format("".join(PUNCTUATIONS)) - sentences = [i for i in re.split(pattern, text) if i.strip() != ""] - phones, tones, word2ph = _g2p(sentences) - assert sum(word2ph) == len(phones) - assert len(word2ph) == len(text) # Sometimes it will crash,you can add a try-catch. - phones = ["_"] + phones + ["_"] - tones = [0] + tones + [0] - word2ph = [1] + word2ph + [1] - return phones, tones, word2ph - - -def _get_initials_finals(word): - initials = [] - finals = [] - orig_initials = lazy_pinyin(word, neutral_tone_with_five=True, style=Style.INITIALS) - orig_finals = lazy_pinyin( - word, neutral_tone_with_five=True, style=Style.FINALS_TONE3 - ) - for c, v in zip(orig_initials, orig_finals): - initials.append(c) - finals.append(v) - return initials, finals - - -def _g2p(segments): - phones_list = [] - tones_list = [] - word2ph = [] - for seg in segments: - # Replace all English words in the sentence - seg = re.sub("[a-zA-Z]+", "", seg) - seg_cut = psg.lcut(seg) - initials = [] - finals = [] - seg_cut = tone_modifier.pre_merge_for_modify(seg_cut) - for word, pos in seg_cut: - if pos == "eng": - continue - sub_initials, sub_finals = _get_initials_finals(word) - sub_finals = tone_modifier.modified_tone(word, pos, sub_finals) - initials.append(sub_initials) - finals.append(sub_finals) - - # assert len(sub_initials) == len(sub_finals) == len(word) - initials = sum(initials, []) - finals = sum(finals, []) - # - for c, v in zip(initials, finals): - raw_pinyin = c + v - # NOTE: post process for pypinyin outputs - # we discriminate i, ii and iii - if c == v: - assert c in PUNCTUATIONS - phone = [c] - tone = "0" - word2ph.append(1) - else: - v_without_tone = v[:-1] - tone = v[-1] - - pinyin = c + v_without_tone - assert tone in "12345" - - if c: - # 多音节 - v_rep_map = { - "uei": "ui", - "iou": "iu", - "uen": "un", - } - if v_without_tone in v_rep_map.keys(): - pinyin = c + v_rep_map[v_without_tone] - else: - # 单音节 - pinyin_rep_map = { - "ing": "ying", - "i": "yi", - "in": "yin", - "u": "wu", - } - if pinyin in pinyin_rep_map.keys(): - pinyin = pinyin_rep_map[pinyin] - else: - single_rep_map = { - "v": "yu", - "e": "e", - "i": "y", - "u": "w", - } - if pinyin[0] in single_rep_map.keys(): - pinyin = single_rep_map[pinyin[0]] + pinyin[1:] - - assert pinyin in pinyin_to_symbol_map.keys(), (pinyin, seg, raw_pinyin) - phone = pinyin_to_symbol_map[pinyin].split(" ") - word2ph.append(len(phone)) - - phones_list += phone - tones_list += [int(tone)] * len(phone) - return phones_list, tones_list, word2ph - - -def normalize_text(text: str) -> str: - numbers = re.findall(r"\d+(?:\.?\d+)?", text) - for number in numbers: - text = text.replace(number, cn2an.an2cn(number), 1) - text = replace_punctuation(text) - return text - - -if __name__ == "__main__": - from style_bert_vits2.nlp.chinese.bert_feature import extract_bert_feature - - text = "啊!但是《原神》是由,米哈\游自主, [研发]的一款全.新开放世界.冒险游戏" - text = normalize_text(text) - print(text) - phones, tones, word2ph = g2p(text) - bert = extract_bert_feature(text, word2ph, 'cuda') - - print(phones, tones, word2ph, bert.shape) - - -# # 示例用法 -# text = "这是一个示例文本:,你好!这是一个测试...." -# print(g2p_paddle(text)) # 输出: 这是一个示例文本你好这是一个测试 diff --git a/style_bert_vits2/nlp/chinese/g2p.py b/style_bert_vits2/nlp/chinese/g2p.py new file mode 100644 index 0000000..1cb3839 --- /dev/null +++ b/style_bert_vits2/nlp/chinese/g2p.py @@ -0,0 +1,135 @@ +import re +from pathlib import Path + +import jieba.posseg as psg +from pypinyin import lazy_pinyin, Style + +from style_bert_vits2.nlp.chinese.tone_sandhi import ToneSandhi +from style_bert_vits2.nlp.symbols import PUNCTUATIONS + + +__PINYIN_TO_SYMBOL_MAP = { + line.split("\t")[0]: line.strip().split("\t")[1] + for line in open(Path(__file__).parent / "opencpop-strict.txt").readlines() +} + + +def g2p(text: str) -> tuple[list[str], list[int], list[int]]: + pattern = r"(?<=[{0}])\s*".format("".join(PUNCTUATIONS)) + sentences = [i for i in re.split(pattern, text) if i.strip() != ""] + phones, tones, word2ph = __g2p(sentences) + assert sum(word2ph) == len(phones) + assert len(word2ph) == len(text) # Sometimes it will crash,you can add a try-catch. + phones = ["_"] + phones + ["_"] + tones = [0] + tones + [0] + word2ph = [1] + word2ph + [1] + return phones, tones, word2ph + + +def __g2p(segments: list[str]) -> tuple[list[str], list[int], list[int]]: + phones_list = [] + tones_list = [] + word2ph = [] + tone_modifier = ToneSandhi() + for seg in segments: + # Replace all English words in the sentence + seg = re.sub("[a-zA-Z]+", "", seg) + seg_cut = psg.lcut(seg) + initials = [] + finals = [] + seg_cut = tone_modifier.pre_merge_for_modify(seg_cut) # type: ignore + for word, pos in seg_cut: + if pos == "eng": + continue + sub_initials, sub_finals = __get_initials_finals(word) + sub_finals = tone_modifier.modified_tone(word, pos, sub_finals) + initials.append(sub_initials) + finals.append(sub_finals) + + # assert len(sub_initials) == len(sub_finals) == len(word) + initials = sum(initials, []) + finals = sum(finals, []) + # + for c, v in zip(initials, finals): + raw_pinyin = c + v + # NOTE: post process for pypinyin outputs + # we discriminate i, ii and iii + if c == v: + assert c in PUNCTUATIONS + phone = [c] + tone = "0" + word2ph.append(1) + else: + v_without_tone = v[:-1] + tone = v[-1] + + pinyin = c + v_without_tone + assert tone in "12345" + + if c: + # 多音节 + v_rep_map = { + "uei": "ui", + "iou": "iu", + "uen": "un", + } + if v_without_tone in v_rep_map.keys(): + pinyin = c + v_rep_map[v_without_tone] + else: + # 单音节 + pinyin_rep_map = { + "ing": "ying", + "i": "yi", + "in": "yin", + "u": "wu", + } + if pinyin in pinyin_rep_map.keys(): + pinyin = pinyin_rep_map[pinyin] + else: + single_rep_map = { + "v": "yu", + "e": "e", + "i": "y", + "u": "w", + } + if pinyin[0] in single_rep_map.keys(): + pinyin = single_rep_map[pinyin[0]] + pinyin[1:] + + assert pinyin in __PINYIN_TO_SYMBOL_MAP.keys(), (pinyin, seg, raw_pinyin) + phone = __PINYIN_TO_SYMBOL_MAP[pinyin].split(" ") + word2ph.append(len(phone)) + + phones_list += phone + tones_list += [int(tone)] * len(phone) + return phones_list, tones_list, word2ph + + +def __get_initials_finals(word: str) -> tuple[list[str], list[str]]: + initials = [] + finals = [] + orig_initials = lazy_pinyin(word, neutral_tone_with_five=True, style=Style.INITIALS) + orig_finals = lazy_pinyin( + word, neutral_tone_with_five=True, style=Style.FINALS_TONE3 + ) + for c, v in zip(orig_initials, orig_finals): + initials.append(c) + finals.append(v) + return initials, finals + + +if __name__ == "__main__": + from style_bert_vits2.nlp.chinese.bert_feature import extract_bert_feature + from style_bert_vits2.nlp.chinese.normalizer import normalize_text + + text = "啊!但是《原神》是由,米哈游自主, [研发]的一款全.新开放世界.冒险游戏" + text = normalize_text(text) + print(text) + phones, tones, word2ph = g2p(text) + bert = extract_bert_feature(text, word2ph, 'cuda') + + print(phones, tones, word2ph, bert.shape) + + +# 示例用法 +# text = "这是一个示例文本:,你好!这是一个测试...." +# print(g2p_paddle(text)) # 输出: 这是一个示例文本你好这是一个测试 diff --git a/style_bert_vits2/nlp/chinese/normalizer.py b/style_bert_vits2/nlp/chinese/normalizer.py new file mode 100644 index 0000000..c56c636 --- /dev/null +++ b/style_bert_vits2/nlp/chinese/normalizer.py @@ -0,0 +1,61 @@ +import re + +import cn2an + +from style_bert_vits2.nlp.symbols import PUNCTUATIONS + + +def normalize_text(text: str) -> str: + numbers = re.findall(r"\d+(?:\.?\d+)?", text) + for number in numbers: + text = text.replace(number, cn2an.an2cn(number), 1) + text = replace_punctuation(text) + return text + + +def replace_punctuation(text: str) -> str: + + REPLACE_MAP = { + ":": ",", + ";": ",", + ",": ",", + "。": ".", + "!": "!", + "?": "?", + "\n": ".", + "·": ",", + "、": ",", + "...": "…", + "$": ".", + "“": "'", + "”": "'", + '"': "'", + "‘": "'", + "’": "'", + "(": "'", + ")": "'", + "(": "'", + ")": "'", + "《": "'", + "》": "'", + "【": "'", + "】": "'", + "[": "'", + "]": "'", + "—": "-", + "~": "-", + "~": "-", + "「": "'", + "」": "'", + } + + text = text.replace("嗯", "恩").replace("呣", "母") + pattern = re.compile("|".join(re.escape(p) for p in REPLACE_MAP.keys())) + + replaced_text = pattern.sub(lambda x: REPLACE_MAP[x.group()], text) + + replaced_text = re.sub( + r"[^\u4e00-\u9fa5" + "".join(PUNCTUATIONS) + r"]+", "", replaced_text + ) + + return replaced_text diff --git a/style_bert_vits2/nlp/english/opencpop-strict.txt b/style_bert_vits2/nlp/chinese/opencpop-strict.txt similarity index 100% rename from style_bert_vits2/nlp/english/opencpop-strict.txt rename to style_bert_vits2/nlp/chinese/opencpop-strict.txt diff --git a/style_bert_vits2/nlp/english/normalizer.py b/style_bert_vits2/nlp/english/normalizer.py index 0886581..81b71d7 100644 --- a/style_bert_vits2/nlp/english/normalizer.py +++ b/style_bert_vits2/nlp/english/normalizer.py @@ -14,22 +14,12 @@ __NUMBER_PATTERN = re.compile(r"[0-9]+") def normalize_text(text: str) -> str: text = __normalize_numbers(text) - text = __replace_punctuation(text) + text = replace_punctuation(text) text = re.sub(r"([,;.\?\!])([\w])", r"\1 \2", text) return text -def __normalize_numbers(text: str) -> str: - text = re.sub(__COMMA_NUMBER_PATTERN, __remove_commas, text) - text = re.sub(__POUNDS_PATTERN, r"\1 pounds", text) - text = re.sub(__DOLLARS_PATTERN, __expand_dollars, text) - text = re.sub(__DECIMAL_NUMBER_PATTERN, __expand_decimal_point, text) - text = re.sub(__ORDINAL_PATTERN, __expand_ordinal, text) - text = re.sub(__NUMBER_PATTERN, __expand_number, text) - return text - - -def __replace_punctuation(text: str) -> str: +def replace_punctuation(text: str) -> str: REPLACE_MAP = { ":": ",", ";": ",", @@ -80,6 +70,16 @@ def __replace_punctuation(text: str) -> str: return replaced_text +def __normalize_numbers(text: str) -> str: + text = re.sub(__COMMA_NUMBER_PATTERN, __remove_commas, text) + text = re.sub(__POUNDS_PATTERN, r"\1 pounds", text) + text = re.sub(__DOLLARS_PATTERN, __expand_dollars, text) + text = re.sub(__DECIMAL_NUMBER_PATTERN, __expand_decimal_point, text) + text = re.sub(__ORDINAL_PATTERN, __expand_ordinal, text) + text = re.sub(__NUMBER_PATTERN, __expand_number, text) + return text + + def __expand_dollars(m: re.Match[str]) -> str: match = m.group(1) parts = match.split(".") From 766699e812084322822027dee76b1003dbe02918 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 09:21:09 +0000 Subject: [PATCH 33/64] Fix: pyopenjtalk_worker not working --- style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py | 4 ++-- .../nlp/japanese/pyopenjtalk_worker/worker_server.py | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py index 90419f5..670461d 100644 --- a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py +++ b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py @@ -56,7 +56,6 @@ def initialize(port: int = WORKER_PORT) -> None: import sys import time - logger.debug("initialize") global WORKER_CLIENT if WORKER_CLIENT: return @@ -98,6 +97,7 @@ def initialize(port: int = WORKER_PORT) -> None: if count == 20: raise TimeoutError("サーバーに接続できませんでした") + logger.debug("pyopenjtalk worker server started") WORKER_CLIENT = client atexit.register(terminate) @@ -110,7 +110,7 @@ def initialize(port: int = WORKER_PORT) -> None: # top-level declaration def terminate() -> None: - logger.debug("terminate") + logger.debug("pyopenjtalk worker server terminated") global WORKER_CLIENT if not WORKER_CLIENT: return diff --git a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_server.py b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_server.py index e9a3dc1..149323a 100644 --- a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_server.py +++ b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_server.py @@ -16,7 +16,7 @@ from style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_common import ( # To make it as fast as possible # Probably faster than calling getattr every time -__PYOPENJTALK_FUNC_DICT = { +PYOPENJTALK_FUNC_DICT = { "run_frontend": pyopenjtalk.run_frontend, "make_label": pyopenjtalk.make_label, "mecab_dict_index": pyopenjtalk.mecab_dict_index, @@ -57,7 +57,7 @@ class WorkerServer: elif request_type == RequestType.PYOPENJTALK: func_name = request.get("func") assert isinstance(func_name, str) - func = __PYOPENJTALK_FUNC_DICT[func_name] + func = PYOPENJTALK_FUNC_DICT[func_name] args = request.get("args") kwargs = request.get("kwargs") assert isinstance(args, list) From fe7e31e0806a82972c44569bbb9ac9abc2be4c87 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 09:34:44 +0000 Subject: [PATCH 34/64] Refactor: moved common/tts_model.py to style_bert_vits2/ --- app.py | 2 +- server_editor.py | 2 +- server_fastapi.py | 2 +- speech_mos.py | 4 ++-- {common => style_bert_vits2}/tts_model.py | 0 webui/inference.py | 2 +- webui/merge.py | 2 +- 7 files changed, 7 insertions(+), 7 deletions(-) rename {common => style_bert_vits2}/tts_model.py (100%) diff --git a/app.py b/app.py index 13929cd..5d05d7c 100644 --- a/app.py +++ b/app.py @@ -6,7 +6,7 @@ import torch import yaml from style_bert_vits2.constants import GRADIO_THEME, LATEST_VERSION -from common.tts_model import ModelHolder +from style_bert_vits2.tts_model import ModelHolder from webui import ( create_dataset_app, create_inference_app, diff --git a/server_editor.py b/server_editor.py index 890a074..022a95a 100644 --- a/server_editor.py +++ b/server_editor.py @@ -30,7 +30,6 @@ from fastapi.staticfiles import StaticFiles from pydantic import BaseModel from scipy.io import wavfile -from common.tts_model import ModelHolder from style_bert_vits2.constants import ( DEFAULT_ASSIST_TEXT_WEIGHT, DEFAULT_NOISE, @@ -52,6 +51,7 @@ from style_bert_vits2.nlp.japanese.user_dict import ( rewrite_word, update_dict, ) +from style_bert_vits2.tts_model import ModelHolder # ---フロントエンド部分に関する処理--- diff --git a/server_fastapi.py b/server_fastapi.py index b7ebb77..833af4a 100644 --- a/server_fastapi.py +++ b/server_fastapi.py @@ -20,7 +20,6 @@ from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import FileResponse, Response from scipy.io import wavfile -from common.tts_model import Model, ModelHolder from config import config from style_bert_vits2.constants import ( DEFAULT_ASSIST_TEXT_WEIGHT, @@ -35,6 +34,7 @@ from style_bert_vits2.constants import ( Languages, ) from style_bert_vits2.logging import logger +from style_bert_vits2.tts_model import Model, ModelHolder ln = config.server_config.language diff --git a/speech_mos.py b/speech_mos.py index 15cccef..79221ad 100644 --- a/speech_mos.py +++ b/speech_mos.py @@ -10,9 +10,9 @@ import pandas as pd import torch from tqdm import tqdm -from style_bert_vits2.logging import logger -from common.tts_model import Model from config import config +from style_bert_vits2.logging import logger +from style_bert_vits2.tts_model import Model warnings.filterwarnings("ignore") diff --git a/common/tts_model.py b/style_bert_vits2/tts_model.py similarity index 100% rename from common/tts_model.py rename to style_bert_vits2/tts_model.py diff --git a/webui/inference.py b/webui/inference.py index 026481f..02d6f3a 100644 --- a/webui/inference.py +++ b/webui/inference.py @@ -4,7 +4,6 @@ from typing import Optional import gradio as gr -from common.tts_model import ModelHolder from style_bert_vits2.constants import ( DEFAULT_ASSIST_TEXT_WEIGHT, DEFAULT_LENGTH, @@ -22,6 +21,7 @@ from style_bert_vits2.logging import logger from style_bert_vits2.models.infer import InvalidToneError from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone from style_bert_vits2.nlp.japanese.normalizer import normalize_text +from style_bert_vits2.tts_model import ModelHolder languages = [l.value for l in Languages] diff --git a/webui/merge.py b/webui/merge.py index 1fc3cee..e9dd1f0 100644 --- a/webui/merge.py +++ b/webui/merge.py @@ -9,9 +9,9 @@ import yaml from safetensors import safe_open from safetensors.torch import save_file -from common.tts_model import Model, ModelHolder from style_bert_vits2.constants import DEFAULT_STYLE, GRADIO_THEME from style_bert_vits2.logging import logger +from style_bert_vits2.tts_model import Model, ModelHolder voice_keys = ["dec"] From 67ff3105c1a7065971319bf71cebe1fef4478b75 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 09:40:27 +0000 Subject: [PATCH 35/64] Refactor: moved utils.py to style_bert_vits2/models/ --- bert_gen.py | 2 +- data_utils.py | 2 +- style_bert_vits2/models/infer.py | 2 +- utils.py => style_bert_vits2/models/utils.py | 0 style_bert_vits2/tts_model.py | 2 +- style_gen.py | 2 +- train_ms.py | 2 +- train_ms_jp_extra.py | 2 +- 8 files changed, 7 insertions(+), 7 deletions(-) rename utils.py => style_bert_vits2/models/utils.py (100%) diff --git a/bert_gen.py b/bert_gen.py index 26df64d..0935929 100644 --- a/bert_gen.py +++ b/bert_gen.py @@ -5,10 +5,10 @@ import torch import torch.multiprocessing as mp from tqdm import tqdm -import utils from config import config from style_bert_vits2.logging import logger from style_bert_vits2.models import commons +from style_bert_vits2.models import utils from style_bert_vits2.nlp import cleaned_text_to_sequence, extract_bert_feature from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT diff --git a/data_utils.py b/data_utils.py index 460c15f..99da2e4 100644 --- a/data_utils.py +++ b/data_utils.py @@ -9,9 +9,9 @@ from tqdm import tqdm from config import config from mel_processing import mel_spectrogram_torch, spectrogram_torch -from utils import load_filepaths_and_text, load_wav_to_torch from style_bert_vits2.logging import logger from style_bert_vits2.models import commons +from style_bert_vits2.models.utils import load_filepaths_and_text, load_wav_to_torch from style_bert_vits2.nlp import cleaned_text_to_sequence """Multi speaker version""" diff --git a/style_bert_vits2/models/infer.py b/style_bert_vits2/models/infer.py index 9265eee..7394d0b 100644 --- a/style_bert_vits2/models/infer.py +++ b/style_bert_vits2/models/infer.py @@ -1,10 +1,10 @@ import torch from typing import Optional -import utils 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 import utils from style_bert_vits2.models.models import SynthesizerTrn from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra from style_bert_vits2.nlp import clean_text, cleaned_text_to_sequence, extract_bert_feature diff --git a/utils.py b/style_bert_vits2/models/utils.py similarity index 100% rename from utils.py rename to style_bert_vits2/models/utils.py diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index a5901e1..bf04fe2 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -7,7 +7,6 @@ import numpy as np import torch from gradio.processing_utils import convert_to_16_bit_wav -import utils from style_bert_vits2.constants import ( DEFAULT_ASSIST_TEXT_WEIGHT, DEFAULT_LENGTH, @@ -19,6 +18,7 @@ from style_bert_vits2.constants import ( DEFAULT_STYLE, DEFAULT_STYLE_WEIGHT, ) +from style_bert_vits2.models import utils from style_bert_vits2.models.infer import get_net_g, infer from style_bert_vits2.models.models import SynthesizerTrn from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra diff --git a/style_gen.py b/style_gen.py index 1c1f034..d7f692f 100644 --- a/style_gen.py +++ b/style_gen.py @@ -6,8 +6,8 @@ import numpy as np import torch from tqdm import tqdm -import utils from style_bert_vits2.logging import logger +from style_bert_vits2.models import utils from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from config import config diff --git a/train_ms.py b/train_ms.py index 0cb7a82..977b393 100644 --- a/train_ms.py +++ b/train_ms.py @@ -15,7 +15,6 @@ from tqdm import tqdm # logging.getLogger("numba").setLevel(logging.WARNING) import default_style -import utils from config import config from data_utils import ( DistributedBucketSampler, @@ -26,6 +25,7 @@ from losses import discriminator_loss, feature_loss, generator_loss, kl_loss from mel_processing import mel_spectrogram_torch, spec_to_mel_torch from style_bert_vits2.logging import logger from style_bert_vits2.models import commons +from style_bert_vits2.models import utils from style_bert_vits2.models.models import ( DurationDiscriminator, MultiPeriodDiscriminator, diff --git a/train_ms_jp_extra.py b/train_ms_jp_extra.py index 7d0636b..3b1c01a 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -15,7 +15,6 @@ from tqdm import tqdm # logging.getLogger("numba").setLevel(logging.WARNING) import default_style -import utils from config import config from data_utils import ( DistributedBucketSampler, @@ -26,6 +25,7 @@ from losses import WavLMLoss, discriminator_loss, feature_loss, generator_loss, from mel_processing import mel_spectrogram_torch, spec_to_mel_torch from style_bert_vits2.logging import logger from style_bert_vits2.models import commons +from style_bert_vits2.models import utils from style_bert_vits2.models.models_jp_extra import ( DurationDiscriminator, MultiPeriodDiscriminator, From c915215ad2c3d99b01f0d7843dc8e35efcd4cf1c Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 09:43:31 +0000 Subject: [PATCH 36/64] Refactor: moved transforms.py to style_bert_vits2/models/ --- style_bert_vits2/models/modules.py | 2 +- transforms.py => style_bert_vits2/models/transforms.py | 0 2 files changed, 1 insertion(+), 1 deletion(-) rename transforms.py => style_bert_vits2/models/transforms.py (100%) diff --git a/style_bert_vits2/models/modules.py b/style_bert_vits2/models/modules.py index c076e21..eede771 100644 --- a/style_bert_vits2/models/modules.py +++ b/style_bert_vits2/models/modules.py @@ -5,10 +5,10 @@ from torch import nn from torch.nn import Conv1d from torch.nn import functional as F from torch.nn.utils import remove_weight_norm, weight_norm -from transforms import piecewise_rational_quadratic_transform from style_bert_vits2.models import commons from style_bert_vits2.models.attentions import Encoder +from style_bert_vits2.models.transforms import piecewise_rational_quadratic_transform LRELU_SLOPE = 0.1 diff --git a/transforms.py b/style_bert_vits2/models/transforms.py similarity index 100% rename from transforms.py rename to style_bert_vits2/models/transforms.py From 8fad591ac0894a7cae11ceb9afa617ecf15a1231 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 11:59:14 +0000 Subject: [PATCH 37/64] Fix: backported StrEnum from Python 3.11 --- style_bert_vits2/constants.py | 3 ++- style_bert_vits2/utils/strenum.py | 36 +++++++++++++++++++++++++++++++ 2 files changed, 38 insertions(+), 1 deletion(-) create mode 100644 style_bert_vits2/utils/strenum.py diff --git a/style_bert_vits2/constants.py b/style_bert_vits2/constants.py index d5a46bf..5f32a15 100644 --- a/style_bert_vits2/constants.py +++ b/style_bert_vits2/constants.py @@ -1,6 +1,7 @@ -from enum import StrEnum from pathlib import Path +from style_bert_vits2.utils.strenum import StrEnum + # Style-Bert-VITS2 のバージョン LATEST_VERSION = "2.4" diff --git a/style_bert_vits2/utils/strenum.py b/style_bert_vits2/utils/strenum.py new file mode 100644 index 0000000..40d3b0c --- /dev/null +++ b/style_bert_vits2/utils/strenum.py @@ -0,0 +1,36 @@ +import enum + + +class StrEnum(str, enum.Enum): + """ + Enum where members are also (and must be) strings (backported from Python 3.11). + """ + + def __new__(cls, *values: str) -> "StrEnum": + "values must already be of type `str`" + if len(values) > 3: + raise TypeError('too many arguments for str(): %r' % (values, )) + if len(values) == 1: + # it must be a string + if not isinstance(values[0], str): # type: ignore + raise TypeError('%r is not a string' % (values[0], )) + if len(values) >= 2: + # check that encoding argument is a string + if not isinstance(values[1], str): # type: ignore + raise TypeError('encoding must be a string, not %r' % (values[1], )) + if len(values) == 3: + # check that errors argument is a string + if not isinstance(values[2], str): # type: ignore + raise TypeError('errors must be a string, not %r' % (values[2])) + value = str(*values) + member = str.__new__(cls, value) + member._value_ = value + return member + + + @staticmethod + def _generate_next_value_(name: str, start: int, count: int, last_values: list[str]) -> str: + """ + Return the lower-cased version of the member name. + """ + return name.lower() From 7f0b2528066e3b7e02bfed100dcef3ca5806dca3 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 12:49:11 +0000 Subject: [PATCH 38/64] Refactor: add Pydantic model representing hyper-parameters of Style-Bert-VITS2 model --- requirements.txt | 1 + style_bert_vits2/models/hyper_parameters.py | 116 ++++++++++++++++++++ style_bert_vits2/models/utils.py | 33 ------ 3 files changed, 117 insertions(+), 33 deletions(-) create mode 100644 style_bert_vits2/models/hyper_parameters.py diff --git a/requirements.txt b/requirements.txt index 6bf329f..669af2f 100644 --- a/requirements.txt +++ b/requirements.txt @@ -14,6 +14,7 @@ numba numpy psutil pyannote.audio>=3.1.0 +pydantic pyloudnorm # pyopenjtalk-prebuilt # Should be manually uninstalled pyopenjtalk-dict diff --git a/style_bert_vits2/models/hyper_parameters.py b/style_bert_vits2/models/hyper_parameters.py new file mode 100644 index 0000000..cee7924 --- /dev/null +++ b/style_bert_vits2/models/hyper_parameters.py @@ -0,0 +1,116 @@ +""" +Style-Bert-VITS2 モデルのハイパーパラメータを表す Pydantic モデル。 +デフォルト値は configs/configs_jp_extra.json 内の定義と同一で、 +万が一ロードした config.json に存在しないキーがあった際のフェイルセーフとして適用される。 +""" + +from pathlib import Path +from typing import Optional, Union + +from pydantic import BaseModel + + +class __HyperParametersTrain(BaseModel): + log_interval: int = 200 + eval_interval: int = 1000 + seed: int = 42 + epochs: int = 1000 + learning_rate: float = 0.0001 + betas: list[float] = [0.8, 0.99] + eps: float = 1e-9 + batch_size: int = 2 + bf16_run: bool = False + fp16_run: bool = False + lr_decay: float = 0.99996 + segment_size: int = 16384 + init_lr_ratio: int = 1 + warmup_epochs: int = 0 + c_mel: int = 45 + c_kl: float = 1.0 + c_commit: int = 100 + skip_optimizer: bool = False + freeze_ZH_bert: bool = False + freeze_JP_bert: bool = False + freeze_EN_bert: bool = False + freeze_emo: bool = False + freeze_style: bool = False + freeze_decoder: bool = False + +class __HyperParametersData(BaseModel): + use_jp_extra: bool = True + training_files: str = "Data/dummy/train.list" + validation_files: str = "Data/dummy/val.list" + max_wav_value: float = 32768.0 + sampling_rate: int = 44100 + filter_length: int = 2048 + hop_length: int = 512 + win_length: int = 2048 + n_mel_channels: int = 128 + mel_fmin: float = 0.0 + mel_fmax: Optional[float] = None + add_blank: bool = True + n_speakers: int = 512 + cleaned_text: bool = True + spk2id: dict[str, int] = { + "dummy": 0 + } + num_styles: int = 1 + style2id: dict[str, int] = { + "Neutral": 0, + } + +class __HyperParametersModel(BaseModel): + use_spk_conditioned_encoder: bool = True + use_noise_scaled_mas: bool = True + use_mel_posterior_encoder: bool = False + use_duration_discriminator: bool = False + use_wavlm_discriminator: bool = True + inter_channels: int = 192 + hidden_channels: int = 192 + filter_channels: int = 768 + n_heads: int = 2 + n_layers: int = 6 + kernel_size: int = 3 + p_dropout: float = 0.1 + resblock: str = "1" + resblock_kernel_sizes: list[int] = [3, 7, 11] + resblock_dilation_sizes: list[list[int]] = [ + [1, 3, 5], + [1, 3, 5], + [1, 3, 5] + ] + upsample_rates: list[int] = [8, 8, 2, 2, 2] + upsample_initial_channel: int = 512 + upsample_kernel_sizes: list[int] = [16, 16, 8, 2, 2] + n_layers_q: int = 3 + use_spectral_norm: bool = False + gin_channels: int = 512 + slm: dict[str, Union[int, str]] = { + "model": "./slm/wavlm-base-plus", + "sr": 16000, + "hidden": 768, + "nlayers": 13, + "initial_channel": 64 + } + +class HyperParameters(BaseModel): + version: str = "2.0-JP-Extra" + model_name: str = 'dummy' + train: __HyperParametersTrain + data: __HyperParametersData + model: __HyperParametersModel + + + @staticmethod + def load_from_json(json_path: Union[str, Path]) -> "HyperParameters": + """ + 与えられた JSON ファイルからハイパーパラメータを読み込む。 + + Args: + json_path (Union[str, Path]): JSON ファイルのパス + + Returns: + HyperParameters: ハイパーパラメータ + """ + with open(json_path, "r") as f: + return HyperParameters.model_validate_json(f.read()) diff --git a/style_bert_vits2/models/utils.py b/style_bert_vits2/models/utils.py index 8c7e842..9a7947f 100644 --- a/style_bert_vits2/models/utils.py +++ b/style_bert_vits2/models/utils.py @@ -357,39 +357,6 @@ def check_git_hash(model_dir): open(path, "w").write(cur_hash) -def get_hparams(init=True): - parser = argparse.ArgumentParser() - parser.add_argument( - "-c", - "--config", - type=str, - default="./configs/base.json", - help="JSON file for configuration", - ) - parser.add_argument("-m", "--model", type=str, required=True, help="Model name") - - args = parser.parse_args() - model_dir = os.path.join("./logs", args.model) - - if not os.path.exists(model_dir): - os.makedirs(model_dir) - - config_path = args.config - config_save_path = os.path.join(model_dir, "config.json") - if init: - with open(config_path, "r", encoding="utf-8") as f: - data = f.read() - with open(config_save_path, "w", encoding="utf-8") as f: - f.write(data) - else: - with open(config_save_path, "r", vencoding="utf-8") as f: - data = f.read() - config = json.loads(data) - hparams = HParams(**config) - hparams.model_dir = model_dir - return hparams - - def get_hparams_from_file(config_path): # print("config_path: ", config_path) with open(config_path, "r", encoding="utf-8") as f: From a0b5cd3d1b0f0b78aa3af3931dfba78916639479 Mon Sep 17 00:00:00 2001 From: kale4eat Date: Sat, 9 Mar 2024 00:23:23 +0900 Subject: [PATCH 39/64] Extended timeout in client side --- text/pyopenjtalk_worker/worker_client.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/text/pyopenjtalk_worker/worker_client.py b/text/pyopenjtalk_worker/worker_client.py index 86d8969..710243d 100644 --- a/text/pyopenjtalk_worker/worker_client.py +++ b/text/pyopenjtalk_worker/worker_client.py @@ -9,8 +9,8 @@ from common.log import logger class WorkerClient: def __init__(self, port: int) -> None: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - # 5: timeout - sock.settimeout(5) + # 60: timeout + sock.settimeout(60) sock.connect((socket.gethostname(), port)) self.sock = sock From a84783a6cc75ef17c8f6773728368c1b1642285c Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 15:52:37 +0000 Subject: [PATCH 40/64] Refactor: replace utils.HParams with HyperParameters Pydantic model HyperParameters is largely a drop-in replacement for utils.HParams, which ensures type safety for hyper-parameters. --- bert_gen.py | 4 +- data_utils.py | 3 +- style_bert_vits2/models/hyper_parameters.py | 30 ++- style_bert_vits2/models/infer.py | 192 +++++++------------- style_bert_vits2/models/models.py | 4 +- style_bert_vits2/models/models_jp_extra.py | 4 +- style_bert_vits2/models/utils.py | 42 ----- style_bert_vits2/tts_model.py | 73 ++++---- style_gen.py | 3 +- train_ms.py | 29 ++- train_ms_jp_extra.py | 27 ++- 11 files changed, 190 insertions(+), 221 deletions(-) diff --git a/bert_gen.py b/bert_gen.py index 0935929..5a16af7 100644 --- a/bert_gen.py +++ b/bert_gen.py @@ -8,7 +8,7 @@ from tqdm import tqdm from config import config from style_bert_vits2.logging import logger from style_bert_vits2.models import commons -from style_bert_vits2.models import utils +from style_bert_vits2.models.hyper_parameters import HyperParameters from style_bert_vits2.nlp import cleaned_text_to_sequence, extract_bert_feature from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT @@ -62,7 +62,7 @@ if __name__ == "__main__": ) args, _ = parser.parse_known_args() config_path = args.config - hps = utils.get_hparams_from_file(config_path) + hps = HyperParameters.load_from_json(config_path) lines = [] with open(hps.data.training_files, encoding="utf-8") as f: lines.extend(f.readlines()) diff --git a/data_utils.py b/data_utils.py index 99da2e4..04047e2 100644 --- a/data_utils.py +++ b/data_utils.py @@ -11,6 +11,7 @@ from config import config from mel_processing import mel_spectrogram_torch, spectrogram_torch from style_bert_vits2.logging import logger from style_bert_vits2.models import commons +from style_bert_vits2.models.hyper_parameters import HyperParametersData from style_bert_vits2.models.utils import load_filepaths_and_text, load_wav_to_torch from style_bert_vits2.nlp import cleaned_text_to_sequence @@ -24,7 +25,7 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset): 3) computes spectrograms from audio files. """ - def __init__(self, audiopaths_sid_text, hparams): + def __init__(self, audiopaths_sid_text: str, hparams: HyperParametersData): self.audiopaths_sid_text = load_filepaths_and_text(audiopaths_sid_text) self.max_wav_value = hparams.max_wav_value self.sampling_rate = hparams.sampling_rate diff --git a/style_bert_vits2/models/hyper_parameters.py b/style_bert_vits2/models/hyper_parameters.py index cee7924..9dc5afb 100644 --- a/style_bert_vits2/models/hyper_parameters.py +++ b/style_bert_vits2/models/hyper_parameters.py @@ -1,16 +1,16 @@ """ Style-Bert-VITS2 モデルのハイパーパラメータを表す Pydantic モデル。 -デフォルト値は configs/configs_jp_extra.json 内の定義と同一で、 +デフォルト値は configs/configs_jp_extra.json 内の定義と概ね同一で、 万が一ロードした config.json に存在しないキーがあった際のフェイルセーフとして適用される。 """ from pathlib import Path from typing import Optional, Union -from pydantic import BaseModel +from pydantic import BaseModel, ConfigDict -class __HyperParametersTrain(BaseModel): +class HyperParametersTrain(BaseModel): log_interval: int = 200 eval_interval: int = 1000 seed: int = 42 @@ -36,7 +36,8 @@ class __HyperParametersTrain(BaseModel): freeze_style: bool = False freeze_decoder: bool = False -class __HyperParametersData(BaseModel): + +class HyperParametersData(BaseModel): use_jp_extra: bool = True training_files: str = "Data/dummy/train.list" validation_files: str = "Data/dummy/val.list" @@ -59,7 +60,8 @@ class __HyperParametersData(BaseModel): "Neutral": 0, } -class __HyperParametersModel(BaseModel): + +class HyperParametersModel(BaseModel): use_spk_conditioned_encoder: bool = True use_noise_scaled_mas: bool = True use_mel_posterior_encoder: bool = False @@ -93,12 +95,21 @@ class __HyperParametersModel(BaseModel): "initial_channel": 64 } + class HyperParameters(BaseModel): - version: str = "2.0-JP-Extra" model_name: str = 'dummy' - train: __HyperParametersTrain - data: __HyperParametersData - model: __HyperParametersModel + version: str = "2.0-JP-Extra" + train: HyperParametersTrain + data: HyperParametersData + model: HyperParametersModel + + # 以下は学習時にのみ動的に設定されるパラメータ (通常 config.json には存在しない) + model_dir: Optional[str] = None + speedup: bool = False + repo_id: Optional[str] = None + + # model_ 以下を Pydantic の保護対象から除外する + model_config = ConfigDict(protected_namespaces=()) @staticmethod @@ -112,5 +123,6 @@ class HyperParameters(BaseModel): Returns: HyperParameters: ハイパーパラメータ """ + with open(json_path, "r") as f: return HyperParameters.model_validate_json(f.read()) diff --git a/style_bert_vits2/models/infer.py b/style_bert_vits2/models/infer.py index 7394d0b..93bf8c2 100644 --- a/style_bert_vits2/models/infer.py +++ b/style_bert_vits2/models/infer.py @@ -1,34 +1,81 @@ +from typing import Any, cast, Optional, Union + import torch -from typing import Optional +from numpy.typing import NDArray 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 import utils +from style_bert_vits2.models.hyper_parameters import HyperParameters from style_bert_vits2.models.models import SynthesizerTrn from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra from style_bert_vits2.nlp import clean_text, cleaned_text_to_sequence, extract_bert_feature from style_bert_vits2.nlp.symbols import SYMBOLS -def get_net_g(model_path: str, version: str, device: str, hps): +def get_net_g(model_path: str, version: str, device: str, hps: HyperParameters): if version.endswith("JP-Extra"): logger.info("Using JP-Extra model") net_g = SynthesizerTrnJPExtra( - len(SYMBOLS), - hps.data.filter_length // 2 + 1, - hps.train.segment_size // hps.data.hop_length, - n_speakers=hps.data.n_speakers, - **hps.model, + n_vocab = len(SYMBOLS), + spec_channels = hps.data.filter_length // 2 + 1, + segment_size = hps.train.segment_size // hps.data.hop_length, + n_speakers = hps.data.n_speakers, + # hps.model 以下のすべての値を引数に渡す + use_spk_conditioned_encoder = hps.model.use_spk_conditioned_encoder, + use_noise_scaled_mas = hps.model.use_noise_scaled_mas, + use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder, + use_duration_discriminator = hps.model.use_duration_discriminator, + use_wavlm_discriminator = hps.model.use_wavlm_discriminator, + inter_channels = hps.model.inter_channels, + hidden_channels = hps.model.hidden_channels, + filter_channels = hps.model.filter_channels, + n_heads = hps.model.n_heads, + n_layers = hps.model.n_layers, + kernel_size = hps.model.kernel_size, + p_dropout = hps.model.p_dropout, + resblock = hps.model.resblock, + resblock_kernel_sizes = hps.model.resblock_kernel_sizes, + resblock_dilation_sizes = hps.model.resblock_dilation_sizes, + upsample_rates = hps.model.upsample_rates, + upsample_initial_channel = hps.model.upsample_initial_channel, + upsample_kernel_sizes = hps.model.upsample_kernel_sizes, + n_layers_q = hps.model.n_layers_q, + use_spectral_norm = hps.model.use_spectral_norm, + gin_channels = hps.model.gin_channels, + slm = hps.model.slm, ).to(device) else: logger.info("Using normal model") net_g = SynthesizerTrn( - len(SYMBOLS), - hps.data.filter_length // 2 + 1, - hps.train.segment_size // hps.data.hop_length, + n_vocab = len(SYMBOLS), + spec_channels = hps.data.filter_length // 2 + 1, + segment_size = hps.train.segment_size // hps.data.hop_length, n_speakers=hps.data.n_speakers, - **hps.model, + # hps.model 以下のすべての値を引数に渡す + use_spk_conditioned_encoder = hps.model.use_spk_conditioned_encoder, + use_noise_scaled_mas = hps.model.use_noise_scaled_mas, + use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder, + use_duration_discriminator = hps.model.use_duration_discriminator, + use_wavlm_discriminator = hps.model.use_wavlm_discriminator, + inter_channels = hps.model.inter_channels, + hidden_channels = hps.model.hidden_channels, + filter_channels = hps.model.filter_channels, + n_heads = hps.model.n_heads, + n_layers = hps.model.n_layers, + kernel_size = hps.model.kernel_size, + p_dropout = hps.model.p_dropout, + resblock = hps.model.resblock, + resblock_kernel_sizes = hps.model.resblock_kernel_sizes, + resblock_dilation_sizes = hps.model.resblock_dilation_sizes, + upsample_rates = hps.model.upsample_rates, + upsample_initial_channel = hps.model.upsample_initial_channel, + upsample_kernel_sizes = hps.model.upsample_kernel_sizes, + n_layers_q = hps.model.n_layers_q, + use_spectral_norm = hps.model.use_spectral_norm, + gin_channels = hps.model.gin_channels, + slm = hps.model.slm, ).to(device) net_g.state_dict() _ = net_g.eval() @@ -44,7 +91,7 @@ def get_net_g(model_path: str, version: str, device: str, hps): def get_text( text: str, language_str: Languages, - hps, + hps: HyperParameters, device: str, assist_text: Optional[str] = None, assist_text_weight: float = 0.7, @@ -111,15 +158,15 @@ def get_text( def infer( text: str, - style_vec, + style_vec: NDArray[Any], sdp_ratio: float, noise_scale: float, noise_scale_w: float, length_scale: float, sid: int, # In the original Bert-VITS2, its speaker_name: str, but here it's id language: Languages, - hps, - net_g, + hps: HyperParameters, + net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra], device: str, skip_start: bool = False, skip_end: bool = False, @@ -159,25 +206,25 @@ def infer( ja_bert = ja_bert.to(device).unsqueeze(0) en_bert = en_bert.to(device).unsqueeze(0) x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device) - style_vec = torch.from_numpy(style_vec).to(device).unsqueeze(0) + style_vec_tensor = torch.from_numpy(style_vec).to(device).unsqueeze(0) del phones sid_tensor = torch.LongTensor([sid]).to(device) if is_jp_extra: - output = net_g.infer( + output = cast(SynthesizerTrnJPExtra, net_g).infer( x_tst, x_tst_lengths, sid_tensor, tones, lang_ids, ja_bert, - style_vec=style_vec, + style_vec=style_vec_tensor, sdp_ratio=sdp_ratio, noise_scale=noise_scale, noise_scale_w=noise_scale_w, length_scale=length_scale, ) else: - output = net_g.infer( + output = cast(SynthesizerTrn, net_g).infer( x_tst, x_tst_lengths, sid_tensor, @@ -186,7 +233,7 @@ def infer( bert, ja_bert, en_bert, - style_vec=style_vec, + style_vec=style_vec_tensor, sdp_ratio=sdp_ratio, noise_scale=noise_scale, noise_scale_w=noise_scale_w, @@ -209,110 +256,5 @@ def infer( return audio -def infer_multilang( - text: str, - style_vec, - sdp_ratio: float, - noise_scale: float, - noise_scale_w: float, - length_scale: float, - sid: int, - language: Languages, - hps, - net_g, - device: str, - skip_start: bool = False, - skip_end: bool = False, -): - bert, ja_bert, en_bert, phones, tones, lang_ids = [], [], [], [], [], [] - # emo = get_emo_(reference_audio, emotion, sid) - # if isinstance(reference_audio, np.ndarray): - # emo = get_clap_audio_feature(reference_audio, device) - # else: - # emo = get_clap_text_feature(emotion, device) - # emo = torch.squeeze(emo, dim=1) - for idx, (txt, lang) in enumerate(zip(text, language)): - _skip_start = (idx != 0) or (skip_start and idx == 0) - _skip_end = (idx != len(language) - 1) or skip_end - ( - temp_bert, - temp_ja_bert, - temp_en_bert, - temp_phones, - temp_tones, - temp_lang_ids, - ) = get_text(txt, lang, hps, device) # type: ignore - if _skip_start: - temp_bert = temp_bert[:, 3:] - temp_ja_bert = temp_ja_bert[:, 3:] - temp_en_bert = temp_en_bert[:, 3:] - temp_phones = temp_phones[3:] - temp_tones = temp_tones[3:] - temp_lang_ids = temp_lang_ids[3:] - if _skip_end: - temp_bert = temp_bert[:, :-2] - temp_ja_bert = temp_ja_bert[:, :-2] - temp_en_bert = temp_en_bert[:, :-2] - temp_phones = temp_phones[:-2] - temp_tones = temp_tones[:-2] - temp_lang_ids = temp_lang_ids[:-2] - bert.append(temp_bert) - ja_bert.append(temp_ja_bert) - en_bert.append(temp_en_bert) - phones.append(temp_phones) - tones.append(temp_tones) - lang_ids.append(temp_lang_ids) - bert = torch.concatenate(bert, dim=1) - ja_bert = torch.concatenate(ja_bert, dim=1) - en_bert = torch.concatenate(en_bert, dim=1) - phones = torch.concatenate(phones, dim=0) - tones = torch.concatenate(tones, dim=0) - lang_ids = torch.concatenate(lang_ids, dim=0) - with torch.no_grad(): - x_tst = phones.to(device).unsqueeze(0) - tones = tones.to(device).unsqueeze(0) - lang_ids = lang_ids.to(device).unsqueeze(0) - bert = bert.to(device).unsqueeze(0) - ja_bert = ja_bert.to(device).unsqueeze(0) - en_bert = en_bert.to(device).unsqueeze(0) - # emo = emo.to(device).unsqueeze(0) - x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device) - del phones - speakers = torch.LongTensor([hps.data.spk2id[sid]]).to(device) - audio = ( - net_g.infer( - x_tst, - x_tst_lengths, - speakers, - tones, - lang_ids, - bert, - ja_bert, - en_bert, - style_vec=style_vec, - sdp_ratio=sdp_ratio, - noise_scale=noise_scale, - noise_scale_w=noise_scale_w, - length_scale=length_scale, - )[0][0, 0] - .data.cpu() - .float() - .numpy() - ) - del ( - x_tst, - tones, - lang_ids, - bert, - x_tst_lengths, - speakers, - ja_bert, - en_bert, - ) # , emo - if torch.cuda.is_available(): - torch.cuda.empty_cache() - return audio - - class InvalidToneError(ValueError): pass diff --git a/style_bert_vits2/models/models.py b/style_bert_vits2/models/models.py index b4e1eb7..4efd100 100644 --- a/style_bert_vits2/models/models.py +++ b/style_bert_vits2/models/models.py @@ -983,10 +983,10 @@ class SynthesizerTrn(nn.Module): en_bert, style_vec, noise_scale=0.667, - length_scale=1, + length_scale=1.0, noise_scale_w=0.8, max_len=None, - sdp_ratio=0, + sdp_ratio=0.0, y=None, ): # x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert) diff --git a/style_bert_vits2/models/models_jp_extra.py b/style_bert_vits2/models/models_jp_extra.py index 7bae8b4..3a43d51 100644 --- a/style_bert_vits2/models/models_jp_extra.py +++ b/style_bert_vits2/models/models_jp_extra.py @@ -1029,10 +1029,10 @@ class SynthesizerTrn(nn.Module): bert, style_vec, noise_scale=0.667, - length_scale=1, + length_scale=1.0, noise_scale_w=0.8, max_len=None, - sdp_ratio=0, + sdp_ratio=0.0, y=None, ): # x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert) diff --git a/style_bert_vits2/models/utils.py b/style_bert_vits2/models/utils.py index 9a7947f..37ea269 100644 --- a/style_bert_vits2/models/utils.py +++ b/style_bert_vits2/models/utils.py @@ -355,45 +355,3 @@ def check_git_hash(model_dir): ) else: open(path, "w").write(cur_hash) - - -def get_hparams_from_file(config_path): - # print("config_path: ", config_path) - with open(config_path, "r", encoding="utf-8") as f: - data = f.read() - config = json.loads(data) - - hparams = HParams(**config) - return hparams - - -class HParams: - def __init__(self, **kwargs): - for k, v in kwargs.items(): - if type(v) == dict: - v = HParams(**v) - self[k] = v - - def keys(self): - return self.__dict__.keys() - - def items(self): - return self.__dict__.items() - - def values(self): - return self.__dict__.values() - - def __len__(self): - return len(self.__dict__) - - def __getitem__(self, key): - return getattr(self, key) - - def __setitem__(self, key, value): - return setattr(self, key, value) - - def __contains__(self, key): - return key in self.__dict__ - - def __repr__(self): - return self.__dict__.__repr__() diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index bf04fe2..86c8425 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -1,11 +1,12 @@ import warnings from pathlib import Path -from typing import Optional, Union +from typing import Any, Optional, Union import gradio as gr import numpy as np import torch from gradio.processing_utils import convert_to_16_bit_wav +from numpy.typing import NDArray from style_bert_vits2.constants import ( DEFAULT_ASSIST_TEXT_WEIGHT, @@ -17,15 +18,22 @@ from style_bert_vits2.constants import ( DEFAULT_SPLIT_INTERVAL, DEFAULT_STYLE, DEFAULT_STYLE_WEIGHT, + Languages, ) -from style_bert_vits2.models import utils +from style_bert_vits2.models.hyper_parameters import HyperParameters from style_bert_vits2.models.infer import get_net_g, infer from style_bert_vits2.models.models import SynthesizerTrn from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra from style_bert_vits2.logging import logger -def adjust_voice(fs, wave, pitch_scale, intonation_scale): +def adjust_voice( + fs: int, + wave: NDArray[Any], + pitch_scale: float, + intonation_scale: float, +) -> tuple[int, NDArray[Any]]: + if pitch_scale == 1.0 and intonation_scale == 1.0: # 初期値の場合は、音質劣化を避けるためにそのまま返す return fs, wave @@ -37,15 +45,17 @@ def adjust_voice(fs, wave, pitch_scale, intonation_scale): "pyworld is not installed. Please install it by `pip install pyworld`" ) - # pyworldでf0を加工して合成 - # pyworldよりもよいのがあるかもしれないが…… + # pyworld で f0 を加工して合成 + # pyworld よりもよいのがあるかもしれないが…… + ## pyworld は Cython で書かれているが、スタブファイルがないため型補完が全く効かない… wave = wave.astype(np.double) - f0, t = pyworld.harvest(wave, fs) - # 質が高そうだしとりあえずharvestにしておく - sp = pyworld.cheaptrick(wave, f0, t, fs) - ap = pyworld.d4c(wave, f0, t, fs) + # 質が高そうだしとりあえずharvestにしておく + f0, t = pyworld.harvest(wave, fs) # type: ignore + + sp = pyworld.cheaptrick(wave, f0, t, fs) # type: ignore + ap = pyworld.d4c(wave, f0, t, fs) # type: ignore non_zero_f0 = [f for f in f0 if f != 0] f0_mean = sum(non_zero_f0) / len(non_zero_f0) @@ -55,7 +65,7 @@ def adjust_voice(fs, wave, pitch_scale, intonation_scale): continue f0[i] = pitch_scale * f0_mean + intonation_scale * (f - f0_mean) - wave = pyworld.synthesize(f0, sp, ap, fs) + wave = pyworld.synthesize(f0, sp, ap, fs) # type: ignore return fs, wave @@ -67,7 +77,7 @@ class Model: self.config_path: Path = config_path self.style_vec_path: Path = style_vec_path self.device: str = device - self.hps: utils.HParams = utils.get_hparams_from_file(self.config_path) + self.hps: HyperParameters = HyperParameters.load_from_json(self.config_path) self.spk2id: dict[str, int] = self.hps.data.spk2id self.id2spk: dict[int, str] = {v: k for k, v in self.spk2id.items()} @@ -81,7 +91,7 @@ class Model: f"Number of styles ({self.num_styles}) does not match the number of style2id ({len(self.style2id)})" ) - self.style_vectors: np.ndarray = np.load(self.style_vec_path) + self.style_vectors: NDArray[Any] = np.load(self.style_vec_path) if self.style_vectors.shape[0] != self.num_styles: raise ValueError( f"The number of styles ({self.num_styles}) does not match the number of style vectors ({self.style_vectors.shape[0]})" @@ -97,7 +107,7 @@ class Model: hps=self.hps, ) - def get_style_vector(self, style_id: int, weight: float = 1.0) -> np.ndarray: + def get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]: mean = self.style_vectors[0] style_vec = self.style_vectors[style_id] style_vec = mean + (style_vec - mean) * weight @@ -105,7 +115,7 @@ class Model: def get_style_vector_from_audio( self, audio_path: str, weight: float = 1.0 - ) -> np.ndarray: + ) -> NDArray[Any]: from style_gen import get_style_vector xvec = get_style_vector(audio_path) @@ -116,7 +126,7 @@ class Model: def infer( self, text: str, - language: str = "JP", + language: Languages = Languages.JP, sid: int = 0, reference_audio_path: Optional[str] = None, sdp_ratio: float = DEFAULT_SDP_RATIO, @@ -133,7 +143,7 @@ class Model: given_tone: Optional[list[int]] = None, pitch_scale: float = 1.0, intonation_scale: float = 1.0, - ) -> tuple[int, np.ndarray]: + ) -> tuple[int, NDArray[Any]]: logger.info(f"Start generating audio data from text:\n{text}") if language != "JP" and self.hps.version.endswith("JP-Extra"): raise ValueError( @@ -146,6 +156,7 @@ class Model: if self.net_g is None: self.load_net_g() + assert self.net_g is not None if reference_audio_path is None: style_id = self.style2id[style] style_vector = self.get_style_vector(style_id, style_weight) @@ -246,19 +257,17 @@ class ModelHolder: continue self.model_files_dict[model_dir.name] = model_files self.model_names.append(model_dir.name) - hps = utils.get_hparams_from_file(config_path) + hps = HyperParameters.load_from_json(config_path) style2id: dict[str, int] = hps.data.style2id styles = list(style2id.keys()) spk2id: dict[str, int] = hps.data.spk2id speakers = list(spk2id.keys()) - self.models_info.append( - { - "name": model_dir.name, - "files": [str(f) for f in model_files], - "styles": styles, - "speakers": speakers, - } - ) + self.models_info.append({ + "name": model_dir.name, + "files": [str(f) for f in model_files], + "styles": styles, + "speakers": speakers, + }) def load_model(self, model_name: str, model_path_str: str): model_path = Path(model_path_str) @@ -291,9 +300,9 @@ class ModelHolder: speakers = list(self.current_model.spk2id.keys()) styles = list(self.current_model.style2id.keys()) return ( - gr.Dropdown(choices=styles, value=styles[0]), + gr.Dropdown(choices=styles, value=styles[0]), # type: ignore gr.Button(interactive=True, value="音声合成"), - gr.Dropdown(choices=speakers, value=speakers[0]), + gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore ) self.current_model = Model( model_path=model_path, @@ -304,21 +313,21 @@ class ModelHolder: speakers = list(self.current_model.spk2id.keys()) styles = list(self.current_model.style2id.keys()) return ( - gr.Dropdown(choices=styles, value=styles[0]), + gr.Dropdown(choices=styles, value=styles[0]), # type: ignore gr.Button(interactive=True, value="音声合成"), - gr.Dropdown(choices=speakers, value=speakers[0]), + gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore ) def update_model_files_gr(self, model_name: str) -> gr.Dropdown: model_files = self.model_files_dict[model_name] - return gr.Dropdown(choices=model_files, value=model_files[0]) + return gr.Dropdown(choices=model_files, value=model_files[0]) # type: ignore def update_model_names_gr(self) -> tuple[gr.Dropdown, gr.Dropdown, gr.Button]: self.refresh() initial_model_name = self.model_names[0] initial_model_files = self.model_files_dict[initial_model_name] return ( - gr.Dropdown(choices=self.model_names, value=initial_model_name), - gr.Dropdown(choices=initial_model_files, value=initial_model_files[0]), + gr.Dropdown(choices=self.model_names, value=initial_model_name), # type: ignore + gr.Dropdown(choices=initial_model_files, value=initial_model_files[0]), # type: ignore gr.Button(interactive=False), # For tts_button ) diff --git a/style_gen.py b/style_gen.py index d7f692f..ec0b507 100644 --- a/style_gen.py +++ b/style_gen.py @@ -8,6 +8,7 @@ from tqdm import tqdm from style_bert_vits2.logging import logger from style_bert_vits2.models import utils +from style_bert_vits2.models.hyper_parameters import HyperParameters from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from config import config @@ -72,7 +73,7 @@ if __name__ == "__main__": config_path = args.config num_processes = args.num_processes - hps = utils.get_hparams_from_file(config_path) + hps = HyperParameters.load_from_json(config_path) device = config.style_gen_config.device diff --git a/train_ms.py b/train_ms.py index 977b393..b2cb02b 100644 --- a/train_ms.py +++ b/train_ms.py @@ -26,6 +26,7 @@ from mel_processing import mel_spectrogram_torch, spec_to_mel_torch from style_bert_vits2.logging import logger from style_bert_vits2.models import commons from style_bert_vits2.models import utils +from style_bert_vits2.models.hyper_parameters import HyperParameters from style_bert_vits2.models.models import ( DurationDiscriminator, MultiPeriodDiscriminator, @@ -130,7 +131,7 @@ def run(): local_rank = int(os.environ["LOCAL_RANK"]) n_gpus = dist.get_world_size() - hps = utils.get_hparams_from_file(args.config) + hps = HyperParameters.load_from_json(args.config) # This is needed because we have to pass values to `train_and_evaluate()` hps.model_dir = model_dir hps.speedup = args.speedup @@ -288,7 +289,29 @@ def run(): n_speakers=hps.data.n_speakers, mas_noise_scale_initial=mas_noise_scale_initial, noise_scale_delta=noise_scale_delta, - **hps.model, + # hps.model 以下のすべての値を引数に渡す + use_spk_conditioned_encoder = hps.model.use_spk_conditioned_encoder, + use_noise_scaled_mas = hps.model.use_noise_scaled_mas, + use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder, + use_duration_discriminator = hps.model.use_duration_discriminator, + use_wavlm_discriminator = hps.model.use_wavlm_discriminator, + inter_channels = hps.model.inter_channels, + hidden_channels = hps.model.hidden_channels, + filter_channels = hps.model.filter_channels, + n_heads = hps.model.n_heads, + n_layers = hps.model.n_layers, + kernel_size = hps.model.kernel_size, + p_dropout = hps.model.p_dropout, + resblock = hps.model.resblock, + resblock_kernel_sizes = hps.model.resblock_kernel_sizes, + resblock_dilation_sizes = hps.model.resblock_dilation_sizes, + upsample_rates = hps.model.upsample_rates, + upsample_initial_channel = hps.model.upsample_initial_channel, + upsample_kernel_sizes = hps.model.upsample_kernel_sizes, + n_layers_q = hps.model.n_layers_q, + use_spectral_norm = hps.model.use_spectral_norm, + gin_channels = hps.model.gin_channels, + slm = hps.model.slm, ).cuda(local_rank) if getattr(hps.train, "freeze_ZH_bert", False): @@ -547,7 +570,7 @@ def train_and_evaluate( rank, local_rank, epoch, - hps, + hps: HyperParameters, nets, optims, schedulers, diff --git a/train_ms_jp_extra.py b/train_ms_jp_extra.py index 3b1c01a..a464d29 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -26,6 +26,7 @@ from mel_processing import mel_spectrogram_torch, spec_to_mel_torch from style_bert_vits2.logging import logger from style_bert_vits2.models import commons from style_bert_vits2.models import utils +from style_bert_vits2.models.hyper_parameters import HyperParameters from style_bert_vits2.models.models_jp_extra import ( DurationDiscriminator, MultiPeriodDiscriminator, @@ -129,7 +130,7 @@ def run(): local_rank = int(os.environ["LOCAL_RANK"]) n_gpus = dist.get_world_size() - hps = utils.get_hparams_from_file(args.config) + hps = HyperParameters.load_from_json(args.config) # This is needed because we have to pass values to `train_and_evaluate() hps.model_dir = model_dir hps.speedup = args.speedup @@ -298,7 +299,29 @@ def run(): n_speakers=hps.data.n_speakers, mas_noise_scale_initial=mas_noise_scale_initial, noise_scale_delta=noise_scale_delta, - **hps.model, + # hps.model 以下のすべての値を引数に渡す + use_spk_conditioned_encoder = hps.model.use_spk_conditioned_encoder, + use_noise_scaled_mas = hps.model.use_noise_scaled_mas, + use_mel_posterior_encoder = hps.model.use_mel_posterior_encoder, + use_duration_discriminator = hps.model.use_duration_discriminator, + use_wavlm_discriminator = hps.model.use_wavlm_discriminator, + inter_channels = hps.model.inter_channels, + hidden_channels = hps.model.hidden_channels, + filter_channels = hps.model.filter_channels, + n_heads = hps.model.n_heads, + n_layers = hps.model.n_layers, + kernel_size = hps.model.kernel_size, + p_dropout = hps.model.p_dropout, + resblock = hps.model.resblock, + resblock_kernel_sizes = hps.model.resblock_kernel_sizes, + resblock_dilation_sizes = hps.model.resblock_dilation_sizes, + upsample_rates = hps.model.upsample_rates, + upsample_initial_channel = hps.model.upsample_initial_channel, + upsample_kernel_sizes = hps.model.upsample_kernel_sizes, + n_layers_q = hps.model.n_layers_q, + use_spectral_norm = hps.model.use_spectral_norm, + gin_channels = hps.model.gin_channels, + slm = hps.model.slm, ).cuda(local_rank) if getattr(hps.train, "freeze_JP_bert", False): logger.info("Freezing (JP) bert encoder !!!") From b9e486e72a3fedcbe91f60f32ffdfc213d11c2ec Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 16:44:56 +0000 Subject: [PATCH 41/64] Fix: extend timeout for style_bert_vits2.nlp.japanese.pyopenjtalk_worker.worker_client ref: https://github.com/litagin02/Style-Bert-VITS2/pull/91 --- .../nlp/japanese/pyopenjtalk_worker/worker_client.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_client.py b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_client.py index b87a937..4507af4 100644 --- a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_client.py +++ b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/worker_client.py @@ -11,8 +11,8 @@ class WorkerClient: def __init__(self, port: int) -> None: sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - # 5: timeout - sock.settimeout(5) + # 60: timeout + sock.settimeout(60) sock.connect((socket.gethostname(), port)) self.sock = sock From 30ea08d6ea662b26f9f57698465027eda687f186 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 18:22:40 +0000 Subject: [PATCH 42/64] Fix: forgot to write pyopenjtalk.initialize() --- style_bert_vits2/nlp/japanese/g2p.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/style_bert_vits2/nlp/japanese/g2p.py b/style_bert_vits2/nlp/japanese/g2p.py index 839fb16..70b4e81 100644 --- a/style_bert_vits2/nlp/japanese/g2p.py +++ b/style_bert_vits2/nlp/japanese/g2p.py @@ -249,6 +249,10 @@ def __pyopenjtalk_g2p_prosody(text: str, drop_unvoiced_vowels: bool = True) -> l return -50 return int(match.group(1)) + # pyopenjtalk_worker を初期化 + ## 一度 worker を起動すれば、明示的に終了するかプロセス終了まで同一の worker に接続される + pyopenjtalk.initialize() + labels = pyopenjtalk.make_label(pyopenjtalk.run_frontend(text)) N = len(labels) From d22a11ebb2fb71bde61155305fe7842aa0201a9c Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 22:09:47 +0000 Subject: [PATCH 43/64] Fix: a bug that prevented speech synthesis in app.py --- app.py | 1 + server_editor.py | 8 +++++--- server_fastapi.py | 17 +++++++++++++++++ .../japanese/pyopenjtalk_worker/__init__.py | 7 ++++++- webui/inference.py | 18 +++++++++++++++++- 5 files changed, 46 insertions(+), 5 deletions(-) diff --git a/app.py b/app.py index 5d05d7c..c2237b5 100644 --- a/app.py +++ b/app.py @@ -15,6 +15,7 @@ from webui import ( create_train_app, ) + # Get path settings with Path("configs/paths.yml").open("r", encoding="utf-8") as f: path_config: dict[str, str] = yaml.safe_load(f.read()) diff --git a/server_editor.py b/server_editor.py index 022a95a..75fa472 100644 --- a/server_editor.py +++ b/server_editor.py @@ -149,9 +149,11 @@ def save_last_download(latest_release): # 以降はAPIの設定 # 最初に pyopenjtalk の辞書を更新 +## pyopenjtalk_worker の起動も同時に行われる update_dict() -# 単語分割に使う BERT モデル/トークナイザーを事前にロードしておく +# 事前に BERT モデル/トークナイザーをロードしておく +## ここでロードしなくても必要になった際に自動ロードされるが、時間がかかるため事前にロードしておいた方が体験が良い ## server_editor.py は日本語にしか対応していないため、日本語の BERT モデル/トークナイザーのみロードする bert_models.load_model(Languages.JP) bert_models.load_tokenizer(Languages.JP) @@ -301,7 +303,7 @@ def synthesis(request: SynthesisRequest): ) sr, audio = model.infer( text=text, - language=request.language.value, + language=request.language, sdp_ratio=request.sdpRatio, noise=request.noise, noisew=request.noisew, @@ -361,7 +363,7 @@ def multi_synthesis(request: MultiSynthesisRequest): tone = [t for _, t in phone_tone] sr, audio = model.infer( text=text, - language=req.language.value, + language=req.language, sdp_ratio=req.sdpRatio, noise=req.noise, noisew=req.noisew, diff --git a/server_fastapi.py b/server_fastapi.py index 833af4a..d8da956 100644 --- a/server_fastapi.py +++ b/server_fastapi.py @@ -34,11 +34,28 @@ from style_bert_vits2.constants import ( Languages, ) from style_bert_vits2.logging import logger +from style_bert_vits2.nlp import bert_models +from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk from style_bert_vits2.tts_model import Model, ModelHolder ln = config.server_config.language +# pyopenjtalk_worker を起動 +## Gradio はマルチスレッドだが、initialize() 内部で利用されている signal はマルチスレッドから設定できない +## さらに起動には若干時間がかかるため、事前に起動しておいた方が体験が良い +pyopenjtalk.initialize() + +# 事前に BERT モデル/トークナイザーをロードしておく +## ここでロードしなくても必要になった際に自動ロードされるが、時間がかかるため事前にロードしておいた方が体験が良い +bert_models.load_model(Languages.JP) +bert_models.load_tokenizer(Languages.JP) +bert_models.load_model(Languages.EN) +bert_models.load_tokenizer(Languages.EN) +bert_models.load_model(Languages.ZH) +bert_models.load_tokenizer(Languages.ZH) + + def raise_validation_error(msg: str, param: str): logger.warning(f"Validation error: {msg}") raise HTTPException( diff --git a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py index 670461d..6c87123 100644 --- a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py +++ b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py @@ -105,7 +105,12 @@ def initialize(port: int = WORKER_PORT) -> None: def signal_handler(signum: int, frame: Any): terminate() - signal.signal(signal.SIGTERM, signal_handler) + try: + signal.signal(signal.SIGINT, signal_handler) + signal.signal(signal.SIGTERM, signal_handler) + except ValueError: + # signal only works in main thread + pass # top-level declaration diff --git a/webui/inference.py b/webui/inference.py index 02d6f3a..ef1f97f 100644 --- a/webui/inference.py +++ b/webui/inference.py @@ -19,12 +19,28 @@ from style_bert_vits2.constants import ( ) from style_bert_vits2.logging import logger from style_bert_vits2.models.infer import InvalidToneError +from style_bert_vits2.nlp import bert_models +from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone from style_bert_vits2.nlp.japanese.normalizer import normalize_text from style_bert_vits2.tts_model import ModelHolder -languages = [l.value for l in Languages] +# pyopenjtalk_worker を起動 +## Gradio はマルチスレッドだが、initialize() 内部で利用されている signal はマルチスレッドから設定できない +## さらに起動には若干時間がかかるため、事前に起動しておいた方が体験が良い +pyopenjtalk.initialize() + +# 事前に BERT モデル/トークナイザーをロードしておく +## ここでロードしなくても必要になった際に自動ロードされるが、時間がかかるため事前にロードしておいた方が体験が良い +bert_models.load_model(Languages.JP) +bert_models.load_tokenizer(Languages.JP) +bert_models.load_model(Languages.EN) +bert_models.load_tokenizer(Languages.EN) +bert_models.load_model(Languages.ZH) +bert_models.load_tokenizer(Languages.ZH) + +languages = [l.value for l in Languages] initial_text = "こんにちは、初めまして。あなたの名前はなんていうの?" From b98fecb46d18108ea1c5fdd36a256dcf7e8f655d Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 22:14:39 +0000 Subject: [PATCH 44/64] Add: --host/--port option to app.py to allow specifying listening host/port --- app.py | 4 +++- webui/train.py | 1 - 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/app.py b/app.py index c2237b5..f0f94d0 100644 --- a/app.py +++ b/app.py @@ -24,6 +24,8 @@ with Path("configs/paths.yml").open("r", encoding="utf-8") as f: parser = argparse.ArgumentParser() parser.add_argument("--device", type=str, default="cuda") +parser.add_argument("--host", type=str, default="127.0.0.1") +parser.add_argument("--port", type=int, default=7860) parser.add_argument("--no_autolaunch", action="store_true") parser.add_argument("--share", action="store_true") @@ -49,4 +51,4 @@ with gr.Blocks(theme=GRADIO_THEME) as app: create_merge_app(model_holder=model_holder) -app.launch(inbrowser=not args.no_autolaunch, share=args.share) +app.launch(server_name=args.host, server_port=args.port, inbrowser=not args.no_autolaunch, share=args.share) diff --git a/webui/train.py b/webui/train.py index fa29719..d837471 100644 --- a/webui/train.py +++ b/webui/train.py @@ -792,5 +792,4 @@ def create_train_app(): outputs=[use_jp_extra_train], ) - # app.launch(inbrowser=not args.no_autolaunch, server_name=args.server_name) return app From 717ba7925f5fc76debbedfaf14a99c0226df88c9 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 22:27:21 +0000 Subject: [PATCH 45/64] Fix: app.py cannot be closed with Ctrl+C --- style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py | 1 - 1 file changed, 1 deletion(-) diff --git a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py index 6c87123..3593e60 100644 --- a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py +++ b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py @@ -106,7 +106,6 @@ def initialize(port: int = WORKER_PORT) -> None: terminate() try: - signal.signal(signal.SIGINT, signal_handler) signal.signal(signal.SIGTERM, signal_handler) except ValueError: # signal only works in main thread From e1fad54c9950e001df4ef4a9439702c493c4fdde Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 22:58:40 +0000 Subject: [PATCH 46/64] Refactor: add type hints to attentions.py / modules.py / transforms.py I didn't add docstring because it is very technical code and I don't understand what is being implemented. --- style_bert_vits2/models/attentions.py | 110 +++++++++++---------- style_bert_vits2/models/modules.py | 137 +++++++++++++++----------- style_bert_vits2/models/transforms.py | 84 ++++++++-------- 3 files changed, 182 insertions(+), 149 deletions(-) diff --git a/style_bert_vits2/models/attentions.py b/style_bert_vits2/models/attentions.py index 6d43e08..b262681 100644 --- a/style_bert_vits2/models/attentions.py +++ b/style_bert_vits2/models/attentions.py @@ -1,3 +1,5 @@ +from typing import Any, Optional + import math import torch from torch import nn @@ -7,7 +9,7 @@ from style_bert_vits2.models import commons class LayerNorm(nn.Module): - def __init__(self, channels, eps=1e-5): + def __init__(self, channels: int, eps: float = 1e-5): super().__init__() self.channels = channels self.eps = eps @@ -15,14 +17,14 @@ class LayerNorm(nn.Module): self.gamma = nn.Parameter(torch.ones(channels)) self.beta = nn.Parameter(torch.zeros(channels)) - def forward(self, x): + def forward(self, x: torch.Tensor) -> torch.Tensor: x = x.transpose(1, -1) x = F.layer_norm(x, (self.channels,), self.gamma, self.beta, self.eps) return x.transpose(1, -1) -@torch.jit.script -def fused_add_tanh_sigmoid_multiply(input_a, input_b, n_channels): +@torch.jit.script # type: ignore +def fused_add_tanh_sigmoid_multiply(input_a: torch.Tensor, input_b: torch.Tensor, n_channels: list[int]) -> torch.Tensor: n_channels_int = n_channels[0] in_act = input_a + input_b t_act = torch.tanh(in_act[:, :n_channels_int, :]) @@ -34,15 +36,15 @@ def fused_add_tanh_sigmoid_multiply(input_a, input_b, n_channels): class Encoder(nn.Module): def __init__( self, - hidden_channels, - filter_channels, - n_heads, - n_layers, - kernel_size=1, - p_dropout=0.0, - window_size=4, - isflow=True, - **kwargs + hidden_channels: int, + filter_channels: int, + n_heads: int, + n_layers: int, + kernel_size: int = 1, + p_dropout: float = 0.0, + window_size: int = 4, + isflow: bool = True, + **kwargs: Any ): super().__init__() self.hidden_channels = hidden_channels @@ -97,12 +99,13 @@ class Encoder(nn.Module): ) self.norm_layers_2.append(LayerNorm(hidden_channels)) - def forward(self, x, x_mask, g=None): + def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor: attn_mask = x_mask.unsqueeze(2) * x_mask.unsqueeze(-1) x = x * x_mask for i in range(self.n_layers): if i == self.cond_layer_idx and g is not None: g = self.spk_emb_linear(g.transpose(1, 2)) + assert g is not None g = g.transpose(1, 2) x = x + g x = x * x_mask @@ -120,15 +123,15 @@ class Encoder(nn.Module): class Decoder(nn.Module): def __init__( self, - hidden_channels, - filter_channels, - n_heads, - n_layers, - kernel_size=1, - p_dropout=0.0, - proximal_bias=False, - proximal_init=True, - **kwargs + hidden_channels: int, + filter_channels: int, + n_heads: int, + n_layers: int, + kernel_size: int = 1, + p_dropout: float = 0.0, + proximal_bias: bool = False, + proximal_init: bool = True, + **kwargs: Any ): super().__init__() self.hidden_channels = hidden_channels @@ -177,7 +180,7 @@ class Decoder(nn.Module): ) self.norm_layers_2.append(LayerNorm(hidden_channels)) - def forward(self, x, x_mask, h, h_mask): + def forward(self, x: torch.Tensor, x_mask: torch.Tensor, h: torch.Tensor, h_mask: torch.Tensor): """ x: decoder input h: encoder output @@ -206,15 +209,15 @@ class Decoder(nn.Module): class MultiHeadAttention(nn.Module): def __init__( self, - channels, - out_channels, - n_heads, - p_dropout=0.0, - window_size=None, - heads_share=True, - block_length=None, - proximal_bias=False, - proximal_init=False, + channels: int, + out_channels: int, + n_heads: int, + p_dropout: float = 0.0, + window_size: Optional[int] = None, + heads_share: bool = True, + block_length: Optional[int] = None, + proximal_bias: bool = False, + proximal_init: bool = False, ): super().__init__() assert channels % n_heads == 0 @@ -255,9 +258,11 @@ class MultiHeadAttention(nn.Module): if proximal_init: with torch.no_grad(): self.conv_k.weight.copy_(self.conv_q.weight) + assert self.conv_k.bias is not None + assert self.conv_q.bias is not None self.conv_k.bias.copy_(self.conv_q.bias) - def forward(self, x, c, attn_mask=None): + def forward(self, x: torch.Tensor, c: torch.Tensor, attn_mask: Optional[torch.Tensor] = None) -> torch.Tensor: q = self.conv_q(x) k = self.conv_k(c) v = self.conv_v(c) @@ -267,7 +272,7 @@ class MultiHeadAttention(nn.Module): x = self.conv_o(x) return x - def attention(self, query, key, value, mask=None): + def attention(self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, mask: Optional[torch.Tensor] = None) -> tuple[torch.Tensor, torch.Tensor]: # reshape [b, d, t] -> [b, n_h, t, d_k] b, d, t_s, t_t = (*key.size(), query.size(2)) query = query.view(b, self.n_heads, self.k_channels, t_t).transpose(2, 3) @@ -318,7 +323,7 @@ class MultiHeadAttention(nn.Module): ) # [b, n_h, t_t, d_k] -> [b, d, t_t] return output, p_attn - def _matmul_with_relative_values(self, x, y): + def _matmul_with_relative_values(self, x: torch.Tensor, y: torch.Tensor) -> torch.Tensor: """ x: [b, h, l, m] y: [h or 1, m, d] @@ -327,7 +332,7 @@ class MultiHeadAttention(nn.Module): ret = torch.matmul(x, y.unsqueeze(0)) return ret - def _matmul_with_relative_keys(self, x, y): + def _matmul_with_relative_keys(self, x: torch.Tensor, y: torch.Tensor) -> torch.Tensor: """ x: [b, h, l, d] y: [h or 1, m, d] @@ -336,8 +341,9 @@ class MultiHeadAttention(nn.Module): ret = torch.matmul(x, y.unsqueeze(0).transpose(-2, -1)) return ret - def _get_relative_embeddings(self, relative_embeddings, length): - 2 * self.window_size + 1 + def _get_relative_embeddings(self, relative_embeddings: torch.Tensor, length: int) -> torch.Tensor: + assert self.window_size is not None + 2 * self.window_size + 1 # type: ignore # Pad first before slice to avoid using cond ops. pad_length = max(length - (self.window_size + 1), 0) slice_start_position = max((self.window_size + 1) - length, 0) @@ -354,7 +360,7 @@ class MultiHeadAttention(nn.Module): ] return used_relative_embeddings - def _relative_position_to_absolute_position(self, x): + def _relative_position_to_absolute_position(self, x: torch.Tensor) -> torch.Tensor: """ x: [b, h, l, 2*l-1] ret: [b, h, l, l] @@ -375,7 +381,7 @@ class MultiHeadAttention(nn.Module): ] return x_final - def _absolute_position_to_relative_position(self, x): + def _absolute_position_to_relative_position(self, x: torch.Tensor) -> torch.Tensor: """ x: [b, h, l, l] ret: [b, h, l, 2*l-1] @@ -391,7 +397,7 @@ class MultiHeadAttention(nn.Module): x_final = x_flat.view([batch, heads, length, 2 * length])[:, :, :, 1:] return x_final - def _attention_bias_proximal(self, length): + def _attention_bias_proximal(self, length: int) -> torch.Tensor: """Bias for self-attention to encourage attention to close positions. Args: length: an integer scalar. @@ -406,13 +412,13 @@ class MultiHeadAttention(nn.Module): class FFN(nn.Module): def __init__( self, - in_channels, - out_channels, - filter_channels, - kernel_size, - p_dropout=0.0, - activation=None, - causal=False, + in_channels: int, + out_channels: int, + filter_channels: int, + kernel_size: int, + p_dropout: float = 0.0, + activation: Optional[str] = None, + causal: bool = False, ): super().__init__() self.in_channels = in_channels @@ -432,7 +438,7 @@ class FFN(nn.Module): self.conv_2 = nn.Conv1d(filter_channels, out_channels, kernel_size) self.drop = nn.Dropout(p_dropout) - def forward(self, x, x_mask): + def forward(self, x: torch.Tensor, x_mask: torch.Tensor) -> torch.Tensor: x = self.conv_1(self.padding(x * x_mask)) if self.activation == "gelu": x = x * torch.sigmoid(1.702 * x) @@ -442,7 +448,7 @@ class FFN(nn.Module): x = self.conv_2(self.padding(x * x_mask)) return x * x_mask - def _causal_padding(self, x): + def _causal_padding(self, x: torch.Tensor) -> torch.Tensor: if self.kernel_size == 1: return x pad_l = self.kernel_size - 1 @@ -451,7 +457,7 @@ class FFN(nn.Module): x = F.pad(x, commons.convert_pad_shape(padding)) return x - def _same_padding(self, x): + def _same_padding(self, x: torch.Tensor) -> torch.Tensor: if self.kernel_size == 1: return x pad_l = (self.kernel_size - 1) // 2 diff --git a/style_bert_vits2/models/modules.py b/style_bert_vits2/models/modules.py index eede771..df38071 100644 --- a/style_bert_vits2/models/modules.py +++ b/style_bert_vits2/models/modules.py @@ -1,4 +1,5 @@ import math +from typing import Any, Optional, Union import torch from torch import nn @@ -15,7 +16,7 @@ LRELU_SLOPE = 0.1 class LayerNorm(nn.Module): - def __init__(self, channels, eps=1e-5): + def __init__(self, channels: int, eps: float = 1e-5): super().__init__() self.channels = channels self.eps = eps @@ -23,7 +24,7 @@ class LayerNorm(nn.Module): self.gamma = nn.Parameter(torch.ones(channels)) self.beta = nn.Parameter(torch.zeros(channels)) - def forward(self, x): + def forward(self, x: torch.Tensor) -> torch.Tensor: x = x.transpose(1, -1) x = F.layer_norm(x, (self.channels,), self.gamma, self.beta, self.eps) return x.transpose(1, -1) @@ -32,12 +33,12 @@ class LayerNorm(nn.Module): class ConvReluNorm(nn.Module): def __init__( self, - in_channels, - hidden_channels, - out_channels, - kernel_size, - n_layers, - p_dropout, + in_channels: int, + hidden_channels: int, + out_channels: int, + kernel_size: int, + n_layers: int, + p_dropout: float, ): super().__init__() self.in_channels = in_channels @@ -69,9 +70,10 @@ class ConvReluNorm(nn.Module): self.norm_layers.append(LayerNorm(hidden_channels)) self.proj = nn.Conv1d(hidden_channels, out_channels, 1) self.proj.weight.data.zero_() + assert self.proj.bias is not None self.proj.bias.data.zero_() - def forward(self, x, x_mask): + def forward(self, x: torch.Tensor, x_mask: torch.Tensor) -> torch.Tensor: x_org = x for i in range(self.n_layers): x = self.conv_layers[i](x * x_mask) @@ -86,7 +88,7 @@ class DDSConv(nn.Module): Dialted and Depth-Separable Convolution """ - def __init__(self, channels, kernel_size, n_layers, p_dropout=0.0): + def __init__(self, channels: int, kernel_size: int, n_layers: int, p_dropout: float = 0.0): super().__init__() self.channels = channels self.kernel_size = kernel_size @@ -115,7 +117,7 @@ class DDSConv(nn.Module): self.norms_1.append(LayerNorm(channels)) self.norms_2.append(LayerNorm(channels)) - def forward(self, x, x_mask, g=None): + def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor: if g is not None: x = x + g for i in range(self.n_layers): @@ -133,12 +135,12 @@ class DDSConv(nn.Module): class WN(torch.nn.Module): def __init__( self, - hidden_channels, - kernel_size, - dilation_rate, - n_layers, - gin_channels=0, - p_dropout=0, + hidden_channels: int, + kernel_size: int, + dilation_rate: int, + n_layers: int, + gin_channels: int = 0, + p_dropout: float = 0, ): super(WN, self).__init__() assert kernel_size % 2 == 1 @@ -182,7 +184,7 @@ class WN(torch.nn.Module): res_skip_layer = torch.nn.utils.weight_norm(res_skip_layer, name="weight") self.res_skip_layers.append(res_skip_layer) - def forward(self, x, x_mask, g=None, **kwargs): + def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None, **kwargs: Any) -> torch.Tensor: output = torch.zeros_like(x) n_channels_tensor = torch.IntTensor([self.hidden_channels]) @@ -209,7 +211,7 @@ class WN(torch.nn.Module): output = output + res_skip_acts return output * x_mask - def remove_weight_norm(self): + def remove_weight_norm(self) -> None: if self.gin_channels != 0: torch.nn.utils.remove_weight_norm(self.cond_layer) for l in self.in_layers: @@ -219,7 +221,7 @@ class WN(torch.nn.Module): class ResBlock1(torch.nn.Module): - def __init__(self, channels, kernel_size=3, dilation=(1, 3, 5)): + def __init__(self, channels: int, kernel_size: int = 3, dilation: tuple[int, int, int] = (1, 3, 5)): super(ResBlock1, self).__init__() self.convs1 = nn.ModuleList( [ @@ -293,7 +295,7 @@ class ResBlock1(torch.nn.Module): ) self.convs2.apply(commons.init_weights) - def forward(self, x, x_mask=None): + def forward(self, x: torch.Tensor, x_mask: Optional[torch.Tensor] = None) -> torch.Tensor: for c1, c2 in zip(self.convs1, self.convs2): xt = F.leaky_relu(x, LRELU_SLOPE) if x_mask is not None: @@ -308,7 +310,7 @@ class ResBlock1(torch.nn.Module): x = x * x_mask return x - def remove_weight_norm(self): + def remove_weight_norm(self) -> None: for l in self.convs1: remove_weight_norm(l) for l in self.convs2: @@ -316,7 +318,7 @@ class ResBlock1(torch.nn.Module): class ResBlock2(torch.nn.Module): - def __init__(self, channels, kernel_size=3, dilation=(1, 3)): + def __init__(self, channels: int, kernel_size: int = 3, dilation: tuple[int, int] = (1, 3)): super(ResBlock2, self).__init__() self.convs = nn.ModuleList( [ @@ -344,7 +346,7 @@ class ResBlock2(torch.nn.Module): ) self.convs.apply(commons.init_weights) - def forward(self, x, x_mask=None): + def forward(self, x: torch.Tensor, x_mask: Optional[torch.Tensor] = None) -> torch.Tensor: for c in self.convs: xt = F.leaky_relu(x, LRELU_SLOPE) if x_mask is not None: @@ -355,13 +357,13 @@ class ResBlock2(torch.nn.Module): x = x * x_mask return x - def remove_weight_norm(self): + def remove_weight_norm(self) -> None: for l in self.convs: remove_weight_norm(l) class Log(nn.Module): - def forward(self, x, x_mask, reverse=False, **kwargs): + def forward(self, x: torch.Tensor, x_mask: torch.Tensor, reverse: bool = False, **kwargs: Any): if not reverse: y = torch.log(torch.clamp_min(x, 1e-5)) * x_mask logdet = torch.sum(-y, [1, 2]) @@ -372,7 +374,13 @@ class Log(nn.Module): class Flip(nn.Module): - def forward(self, x, *args, reverse=False, **kwargs): + def forward( + self, + x: torch.Tensor, + *args: Any, + reverse: bool = False, + **kwargs: Any, + ) -> Union[tuple[torch.Tensor, torch.Tensor], torch.Tensor]: x = torch.flip(x, [1]) if not reverse: logdet = torch.zeros(x.size(0)).to(dtype=x.dtype, device=x.device) @@ -382,13 +390,19 @@ class Flip(nn.Module): class ElementwiseAffine(nn.Module): - def __init__(self, channels): + def __init__(self, channels: int): super().__init__() self.channels = channels self.m = nn.Parameter(torch.zeros(channels, 1)) self.logs = nn.Parameter(torch.zeros(channels, 1)) - def forward(self, x, x_mask, reverse=False, **kwargs): + def forward( + self, + x: torch.Tensor, + x_mask: torch.Tensor, + reverse: bool = False, + **kwargs: Any, + ) -> Union[tuple[torch.Tensor, torch.Tensor], torch.Tensor]: if not reverse: y = self.m + torch.exp(self.logs) * x y = y * x_mask @@ -402,14 +416,14 @@ class ElementwiseAffine(nn.Module): class ResidualCouplingLayer(nn.Module): def __init__( self, - channels, - hidden_channels, - kernel_size, - dilation_rate, - n_layers, - p_dropout=0, - gin_channels=0, - mean_only=False, + channels: int, + hidden_channels: int, + kernel_size: int, + dilation_rate: int, + n_layers: int, + p_dropout: float = 0, + gin_channels: int = 0, + mean_only: bool = False, ): assert channels % 2 == 0, "channels should be divisible by 2" super().__init__() @@ -432,9 +446,10 @@ class ResidualCouplingLayer(nn.Module): ) self.post = nn.Conv1d(hidden_channels, self.half_channels * (2 - mean_only), 1) self.post.weight.data.zero_() + assert self.post.bias is not None self.post.bias.data.zero_() - def forward(self, x, x_mask, g=None, reverse=False): + def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None, reverse: bool = False): x0, x1 = torch.split(x, [self.half_channels] * 2, 1) h = self.pre(x0) * x_mask h = self.enc(h, x_mask, g=g) @@ -459,12 +474,12 @@ class ResidualCouplingLayer(nn.Module): class ConvFlow(nn.Module): def __init__( self, - in_channels, - filter_channels, - kernel_size, - n_layers, - num_bins=10, - tail_bound=5.0, + in_channels: int, + filter_channels: int, + kernel_size: int, + n_layers: int, + num_bins: int = 10, + tail_bound: float = 5.0, ): super().__init__() self.in_channels = in_channels @@ -481,9 +496,10 @@ class ConvFlow(nn.Module): filter_channels, self.half_channels * (num_bins * 3 - 1), 1 ) self.proj.weight.data.zero_() + assert self.proj.bias is not None self.proj.bias.data.zero_() - def forward(self, x, x_mask, g=None, reverse=False): + def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None, reverse: bool = False): x0, x1 = torch.split(x, [self.half_channels] * 2, 1) h = self.pre(x0) h = self.convs(h, x_mask, g=g) @@ -519,17 +535,17 @@ class ConvFlow(nn.Module): class TransformerCouplingLayer(nn.Module): def __init__( self, - channels, - hidden_channels, - kernel_size, - n_layers, - n_heads, - p_dropout=0, - filter_channels=0, - mean_only=False, - wn_sharing_parameter=None, - gin_channels=0, - ): + channels: int, + hidden_channels: int, + kernel_size: int, + n_layers: int, + n_heads: int, + p_dropout: float = 0, + filter_channels: int = 0, + mean_only: bool = False, + wn_sharing_parameter: Optional[nn.Module] = None, + gin_channels: int = 0, + ) -> None: assert channels % 2 == 0, "channels should be divisible by 2" super().__init__() self.channels = channels @@ -556,9 +572,16 @@ class TransformerCouplingLayer(nn.Module): ) self.post = nn.Conv1d(hidden_channels, self.half_channels * (2 - mean_only), 1) self.post.weight.data.zero_() + assert self.post.bias is not None self.post.bias.data.zero_() - def forward(self, x, x_mask, g=None, reverse=False): + def forward( + self, + x: torch.Tensor, + x_mask: torch.Tensor, + g: Optional[torch.Tensor] = None, + reverse: bool = False, + ) -> Union[tuple[torch.Tensor, torch.Tensor], torch.Tensor]: x0, x1 = torch.split(x, [self.half_channels] * 2, 1) h = self.pre(x0) * x_mask h = self.enc(h, x_mask, g=g) diff --git a/style_bert_vits2/models/transforms.py b/style_bert_vits2/models/transforms.py index a11f799..61306ad 100644 --- a/style_bert_vits2/models/transforms.py +++ b/style_bert_vits2/models/transforms.py @@ -1,7 +1,8 @@ -import torch -from torch.nn import functional as F +from typing import Optional import numpy as np +import torch +from torch.nn import functional as F DEFAULT_MIN_BIN_WIDTH = 1e-3 @@ -10,17 +11,18 @@ DEFAULT_MIN_DERIVATIVE = 1e-3 def piecewise_rational_quadratic_transform( - inputs, - unnormalized_widths, - unnormalized_heights, - unnormalized_derivatives, - inverse=False, - tails=None, - tail_bound=1.0, - min_bin_width=DEFAULT_MIN_BIN_WIDTH, - min_bin_height=DEFAULT_MIN_BIN_HEIGHT, - min_derivative=DEFAULT_MIN_DERIVATIVE, -): + inputs: torch.Tensor, + unnormalized_widths: torch.Tensor, + unnormalized_heights: torch.Tensor, + unnormalized_derivatives: torch.Tensor, + inverse: bool = False, + tails: Optional[str] = None, + tail_bound: float = 1.0, + min_bin_width: float = DEFAULT_MIN_BIN_WIDTH, + min_bin_height: float = DEFAULT_MIN_BIN_HEIGHT, + min_derivative: float = DEFAULT_MIN_DERIVATIVE, +) -> tuple[torch.Tensor, torch.Tensor]: + if tails is None: spline_fn = rational_quadratic_spline spline_kwargs = {} @@ -37,28 +39,29 @@ def piecewise_rational_quadratic_transform( min_bin_width=min_bin_width, min_bin_height=min_bin_height, min_derivative=min_derivative, - **spline_kwargs + **spline_kwargs # type: ignore ) return outputs, logabsdet -def searchsorted(bin_locations, inputs, eps=1e-6): +def searchsorted(bin_locations: torch.Tensor, inputs: torch.Tensor, eps: float = 1e-6) -> torch.Tensor: bin_locations[..., -1] += eps return torch.sum(inputs[..., None] >= bin_locations, dim=-1) - 1 def unconstrained_rational_quadratic_spline( - inputs, - unnormalized_widths, - unnormalized_heights, - unnormalized_derivatives, - inverse=False, - tails="linear", - tail_bound=1.0, - min_bin_width=DEFAULT_MIN_BIN_WIDTH, - min_bin_height=DEFAULT_MIN_BIN_HEIGHT, - min_derivative=DEFAULT_MIN_DERIVATIVE, -): + inputs: torch.Tensor, + unnormalized_widths: torch.Tensor, + unnormalized_heights: torch.Tensor, + unnormalized_derivatives: torch.Tensor, + inverse: bool = False, + tails: str = "linear", + tail_bound: float = 1.0, + min_bin_width: float = DEFAULT_MIN_BIN_WIDTH, + min_bin_height: float = DEFAULT_MIN_BIN_HEIGHT, + min_derivative: float = DEFAULT_MIN_DERIVATIVE, +) -> tuple[torch.Tensor, torch.Tensor]: + inside_interval_mask = (inputs >= -tail_bound) & (inputs <= tail_bound) outside_interval_mask = ~inside_interval_mask @@ -74,7 +77,7 @@ def unconstrained_rational_quadratic_spline( outputs[outside_interval_mask] = inputs[outside_interval_mask] logabsdet[outside_interval_mask] = 0 else: - raise RuntimeError("{} tails are not implemented.".format(tails)) + raise RuntimeError(f"{tails} tails are not implemented.") ( outputs[inside_interval_mask], @@ -98,19 +101,20 @@ def unconstrained_rational_quadratic_spline( def rational_quadratic_spline( - inputs, - unnormalized_widths, - unnormalized_heights, - unnormalized_derivatives, - inverse=False, - left=0.0, - right=1.0, - bottom=0.0, - top=1.0, - min_bin_width=DEFAULT_MIN_BIN_WIDTH, - min_bin_height=DEFAULT_MIN_BIN_HEIGHT, - min_derivative=DEFAULT_MIN_DERIVATIVE, -): + inputs: torch.Tensor, + unnormalized_widths: torch.Tensor, + unnormalized_heights: torch.Tensor, + unnormalized_derivatives: torch.Tensor, + inverse: bool = False, + left: float = 0.0, + right: float = 1.0, + bottom: float = 0.0, + top: float = 1.0, + min_bin_width: float = DEFAULT_MIN_BIN_WIDTH, + min_bin_height: float = DEFAULT_MIN_BIN_HEIGHT, + min_derivative: float = DEFAULT_MIN_DERIVATIVE, +) -> tuple[torch.Tensor, torch.Tensor]: + if torch.min(inputs) < left or torch.max(inputs) > right: raise ValueError("Input to a transform is not within its domain") From 8feef04cef8e49bbb895d2c954d112ad43c8d063 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 8 Mar 2024 23:52:44 +0000 Subject: [PATCH 47/64] Refactor: add type hints to models.py / models_jp_extra.py I didn't add docstring because it is very technical code and I don't understand what is being implemented. --- style_bert_vits2/models/attentions.py | 20 +- style_bert_vits2/models/models.py | 384 ++++++++++++--------- style_bert_vits2/models/models_jp_extra.py | 369 ++++++++++++-------- style_bert_vits2/models/modules.py | 42 ++- 4 files changed, 490 insertions(+), 325 deletions(-) diff --git a/style_bert_vits2/models/attentions.py b/style_bert_vits2/models/attentions.py index b262681..03b238d 100644 --- a/style_bert_vits2/models/attentions.py +++ b/style_bert_vits2/models/attentions.py @@ -9,7 +9,7 @@ from style_bert_vits2.models import commons class LayerNorm(nn.Module): - def __init__(self, channels: int, eps: float = 1e-5): + def __init__(self, channels: int, eps: float = 1e-5) -> None: super().__init__() self.channels = channels self.eps = eps @@ -45,7 +45,7 @@ class Encoder(nn.Module): window_size: int = 4, isflow: bool = True, **kwargs: Any - ): + ) -> None: super().__init__() self.hidden_channels = hidden_channels self.filter_channels = filter_channels @@ -132,7 +132,7 @@ class Decoder(nn.Module): proximal_bias: bool = False, proximal_init: bool = True, **kwargs: Any - ): + ) -> None: super().__init__() self.hidden_channels = hidden_channels self.filter_channels = filter_channels @@ -180,7 +180,7 @@ class Decoder(nn.Module): ) self.norm_layers_2.append(LayerNorm(hidden_channels)) - def forward(self, x: torch.Tensor, x_mask: torch.Tensor, h: torch.Tensor, h_mask: torch.Tensor): + def forward(self, x: torch.Tensor, x_mask: torch.Tensor, h: torch.Tensor, h_mask: torch.Tensor) -> torch.Tensor: """ x: decoder input h: encoder output @@ -218,7 +218,7 @@ class MultiHeadAttention(nn.Module): block_length: Optional[int] = None, proximal_bias: bool = False, proximal_init: bool = False, - ): + ) -> None: super().__init__() assert channels % n_heads == 0 @@ -272,7 +272,13 @@ class MultiHeadAttention(nn.Module): x = self.conv_o(x) return x - def attention(self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, mask: Optional[torch.Tensor] = None) -> tuple[torch.Tensor, torch.Tensor]: + def attention( + self, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + mask: Optional[torch.Tensor] = None, + ) -> tuple[torch.Tensor, torch.Tensor]: # reshape [b, d, t] -> [b, n_h, t, d_k] b, d, t_s, t_t = (*key.size(), query.size(2)) query = query.view(b, self.n_heads, self.k_channels, t_t).transpose(2, 3) @@ -419,7 +425,7 @@ class FFN(nn.Module): p_dropout: float = 0.0, activation: Optional[str] = None, causal: bool = False, - ): + ) -> None: super().__init__() self.in_channels = in_channels self.out_channels = out_channels diff --git a/style_bert_vits2/models/models.py b/style_bert_vits2/models/models.py index 4efd100..829ca4a 100644 --- a/style_bert_vits2/models/models.py +++ b/style_bert_vits2/models/models.py @@ -1,4 +1,5 @@ import math +from typing import Any, Optional import torch from torch import nn @@ -10,14 +11,18 @@ from style_bert_vits2.models import attentions from style_bert_vits2.models import commons from style_bert_vits2.models import modules from style_bert_vits2.models import monotonic_alignment -from style_bert_vits2.models.commons import get_padding, init_weights from style_bert_vits2.nlp.symbols import NUM_LANGUAGES, NUM_TONES, SYMBOLS class DurationDiscriminator(nn.Module): # vits2 def __init__( - self, in_channels, filter_channels, kernel_size, p_dropout, gin_channels=0 - ): + self, + in_channels: int, + filter_channels: int, + kernel_size: int, + p_dropout: float, + gin_channels: int = 0 + ) -> None: super().__init__() self.in_channels = in_channels @@ -51,7 +56,13 @@ class DurationDiscriminator(nn.Module): # vits2 self.output_layer = nn.Sequential(nn.Linear(filter_channels, 1), nn.Sigmoid()) - def forward_probability(self, x, x_mask, dur, g=None): + def forward_probability( + self, + x: torch.Tensor, + x_mask: torch.Tensor, + dur: torch.Tensor, + g: Optional[torch.Tensor] = None, + ) -> torch.Tensor: dur = self.dur_proj(dur) x = torch.cat([x, dur], dim=1) x = self.pre_out_conv_1(x * x_mask) @@ -67,7 +78,14 @@ class DurationDiscriminator(nn.Module): # vits2 output_prob = self.output_layer(x) return output_prob - def forward(self, x, x_mask, dur_r, dur_hat, g=None): + def forward( + self, + x: torch.Tensor, + x_mask: torch.Tensor, + dur_r: torch.Tensor, + dur_hat: torch.Tensor, + g: Optional[torch.Tensor] = None, + ) -> list[torch.Tensor]: x = torch.detach(x) if g is not None: g = torch.detach(g) @@ -92,17 +110,17 @@ class DurationDiscriminator(nn.Module): # vits2 class TransformerCouplingBlock(nn.Module): def __init__( self, - channels, - hidden_channels, - filter_channels, - n_heads, - n_layers, - kernel_size, - p_dropout, - n_flows=4, - gin_channels=0, - share_parameter=False, - ): + channels: int, + hidden_channels: int, + filter_channels: int, + n_heads: int, + n_layers: int, + kernel_size: int, + p_dropout: float, + n_flows: int = 4, + gin_channels: int = 0, + share_parameter: bool = False, + ) -> None: super().__init__() self.channels = channels self.hidden_channels = hidden_channels @@ -114,16 +132,17 @@ class TransformerCouplingBlock(nn.Module): self.flows = nn.ModuleList() self.wn = ( - attentions.FFT( - hidden_channels, - filter_channels, - n_heads, - n_layers, - kernel_size, - p_dropout, - isflow=True, - gin_channels=self.gin_channels, - ) + # attentions.FFT( + # hidden_channels, + # filter_channels, + # n_heads, + # n_layers, + # kernel_size, + # p_dropout, + # isflow=True, + # gin_channels=self.gin_channels, + # ) + None if share_parameter else None ) @@ -145,7 +164,13 @@ class TransformerCouplingBlock(nn.Module): ) self.flows.append(modules.Flip()) - def forward(self, x, x_mask, g=None, reverse=False): + def forward( + self, + x: torch.Tensor, + x_mask: torch.Tensor, + g: Optional[torch.Tensor] = None, + reverse: bool = False, + ) -> torch.Tensor: if not reverse: for flow in self.flows: x, _ = flow(x, x_mask, g=g, reverse=reverse) @@ -158,13 +183,13 @@ class TransformerCouplingBlock(nn.Module): class StochasticDurationPredictor(nn.Module): def __init__( self, - in_channels, - filter_channels, - kernel_size, - p_dropout, - n_flows=4, - gin_channels=0, - ): + in_channels: int, + filter_channels: int, + kernel_size: int, + p_dropout: float, + n_flows: int = 4, + gin_channels: int = 0, + ) -> None: super().__init__() filter_channels = in_channels # it needs to be removed from future version. self.in_channels = in_channels @@ -204,7 +229,15 @@ class StochasticDurationPredictor(nn.Module): if gin_channels != 0: self.cond = nn.Conv1d(gin_channels, filter_channels, 1) - def forward(self, x, x_mask, w=None, g=None, reverse=False, noise_scale=1.0): + def forward( + self, + x: torch.Tensor, + x_mask: torch.Tensor, + w: Optional[torch.Tensor] = None, + g: Optional[torch.Tensor] = None, + reverse: bool = False, + noise_scale: float = 1.0, + ) -> torch.Tensor: x = torch.detach(x) x = self.pre(x) if g is not None: @@ -268,8 +301,13 @@ class StochasticDurationPredictor(nn.Module): class DurationPredictor(nn.Module): def __init__( - self, in_channels, filter_channels, kernel_size, p_dropout, gin_channels=0 - ): + self, + in_channels: int, + filter_channels: int, + kernel_size: int, + p_dropout: float, + gin_channels: int = 0, + ) -> None: super().__init__() self.in_channels = in_channels @@ -292,7 +330,7 @@ class DurationPredictor(nn.Module): if gin_channels != 0: self.cond = nn.Conv1d(gin_channels, in_channels, 1) - def forward(self, x, x_mask, g=None): + def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor: x = torch.detach(x) if g is not None: g = torch.detach(g) @@ -312,17 +350,17 @@ class DurationPredictor(nn.Module): class TextEncoder(nn.Module): def __init__( self, - n_vocab, - out_channels, - hidden_channels, - filter_channels, - n_heads, - n_layers, - kernel_size, - p_dropout, - n_speakers, - gin_channels=0, - ): + n_vocab: int, + out_channels: int, + hidden_channels: int, + filter_channels: int, + n_heads: int, + n_layers: int, + kernel_size: int, + p_dropout: float, + n_speakers: int, + gin_channels: int = 0, + ) -> None: super().__init__() self.n_vocab = n_vocab self.out_channels = out_channels @@ -357,17 +395,17 @@ class TextEncoder(nn.Module): def forward( self, - x, - x_lengths, - tone, - language, - bert, - ja_bert, - en_bert, - style_vec, - sid, - g=None, - ): + x: torch.Tensor, + x_lengths: torch.Tensor, + tone: torch.Tensor, + language: torch.Tensor, + bert: torch.Tensor, + ja_bert: torch.Tensor, + en_bert: torch.Tensor, + style_vec: torch.Tensor, + sid: torch.Tensor, + g: Optional[torch.Tensor] = None, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: bert_emb = self.bert_proj(bert).transpose(1, 2) ja_bert_emb = self.ja_bert_proj(ja_bert).transpose(1, 2) en_bert_emb = self.en_bert_proj(en_bert).transpose(1, 2) @@ -399,14 +437,14 @@ class TextEncoder(nn.Module): class ResidualCouplingBlock(nn.Module): def __init__( self, - channels, - hidden_channels, - kernel_size, - dilation_rate, - n_layers, - n_flows=4, - gin_channels=0, - ): + channels: int, + hidden_channels: int, + kernel_size: int, + dilation_rate: int, + n_layers: int, + n_flows: int = 4, + gin_channels: int = 0, + ) -> None: super().__init__() self.channels = channels self.hidden_channels = hidden_channels @@ -431,7 +469,13 @@ class ResidualCouplingBlock(nn.Module): ) self.flows.append(modules.Flip()) - def forward(self, x, x_mask, g=None, reverse=False): + def forward( + self, + x: torch.Tensor, + x_mask: torch.Tensor, + g: Optional[torch.Tensor] = None, + reverse: bool = False, + ) -> torch.Tensor: if not reverse: for flow in self.flows: x, _ = flow(x, x_mask, g=g, reverse=reverse) @@ -444,14 +488,14 @@ class ResidualCouplingBlock(nn.Module): class PosteriorEncoder(nn.Module): def __init__( self, - in_channels, - out_channels, - hidden_channels, - kernel_size, - dilation_rate, - n_layers, - gin_channels=0, - ): + in_channels: int, + out_channels: int, + hidden_channels: int, + kernel_size: int, + dilation_rate: int, + n_layers: int, + gin_channels: int = 0, + ) -> None: super().__init__() self.in_channels = in_channels self.out_channels = out_channels @@ -471,7 +515,12 @@ class PosteriorEncoder(nn.Module): ) self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1) - def forward(self, x, x_lengths, g=None): + def forward( + self, + x: torch.Tensor, + x_lengths: torch.Tensor, + g: Optional[torch.Tensor] = None, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to( x.dtype ) @@ -486,22 +535,22 @@ class PosteriorEncoder(nn.Module): class Generator(torch.nn.Module): def __init__( self, - initial_channel, - resblock, - resblock_kernel_sizes, - resblock_dilation_sizes, - upsample_rates, - upsample_initial_channel, - upsample_kernel_sizes, - gin_channels=0, - ): + initial_channel: int, + resblock_str: str, + resblock_kernel_sizes: list[int], + resblock_dilation_sizes: list[list[int]], + upsample_rates: list[int], + upsample_initial_channel: int, + upsample_kernel_sizes: list[int], + gin_channels: int = 0, + ) -> None: super(Generator, self).__init__() self.num_kernels = len(resblock_kernel_sizes) self.num_upsamples = len(upsample_rates) self.conv_pre = Conv1d( initial_channel, upsample_initial_channel, 7, 1, padding=3 ) - resblock = modules.ResBlock1 if resblock == "1" else modules.ResBlock2 + resblock = modules.ResBlock1 if resblock_str == "1" else modules.ResBlock2 self.ups = nn.ModuleList() for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)): @@ -518,20 +567,22 @@ class Generator(torch.nn.Module): ) self.resblocks = nn.ModuleList() + ch = None for i in range(len(self.ups)): ch = upsample_initial_channel // (2 ** (i + 1)) for j, (k, d) in enumerate( zip(resblock_kernel_sizes, resblock_dilation_sizes) ): - self.resblocks.append(resblock(ch, k, d)) + self.resblocks.append(resblock(ch, k, d)) # type: ignore + assert ch is not None self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False) - self.ups.apply(init_weights) + self.ups.apply(commons.init_weights) if gin_channels != 0: self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1) - def forward(self, x, g=None): + def forward(self, x: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor: x = self.conv_pre(x) if g is not None: x = x + self.cond(g) @@ -545,6 +596,7 @@ class Generator(torch.nn.Module): xs = self.resblocks[i * self.num_kernels + j](x) else: xs += self.resblocks[i * self.num_kernels + j](x) + assert xs is not None x = xs / self.num_kernels x = F.leaky_relu(x) x = self.conv_post(x) @@ -552,7 +604,7 @@ class Generator(torch.nn.Module): return x - def remove_weight_norm(self): + def remove_weight_norm(self) -> None: print("Removing weight norm...") for layer in self.ups: remove_weight_norm(layer) @@ -561,7 +613,7 @@ class Generator(torch.nn.Module): class DiscriminatorP(torch.nn.Module): - def __init__(self, period, kernel_size=5, stride=3, use_spectral_norm=False): + def __init__(self, period: int, kernel_size: int = 5, stride: int = 3, use_spectral_norm: bool = False) -> None: super(DiscriminatorP, self).__init__() self.period = period self.use_spectral_norm = use_spectral_norm @@ -574,7 +626,7 @@ class DiscriminatorP(torch.nn.Module): 32, (kernel_size, 1), (stride, 1), - padding=(get_padding(kernel_size, 1), 0), + padding=(commons.get_padding(kernel_size, 1), 0), ) ), norm_f( @@ -583,7 +635,7 @@ class DiscriminatorP(torch.nn.Module): 128, (kernel_size, 1), (stride, 1), - padding=(get_padding(kernel_size, 1), 0), + padding=(commons.get_padding(kernel_size, 1), 0), ) ), norm_f( @@ -592,7 +644,7 @@ class DiscriminatorP(torch.nn.Module): 512, (kernel_size, 1), (stride, 1), - padding=(get_padding(kernel_size, 1), 0), + padding=(commons.get_padding(kernel_size, 1), 0), ) ), norm_f( @@ -601,7 +653,7 @@ class DiscriminatorP(torch.nn.Module): 1024, (kernel_size, 1), (stride, 1), - padding=(get_padding(kernel_size, 1), 0), + padding=(commons.get_padding(kernel_size, 1), 0), ) ), norm_f( @@ -610,14 +662,14 @@ class DiscriminatorP(torch.nn.Module): 1024, (kernel_size, 1), 1, - padding=(get_padding(kernel_size, 1), 0), + padding=(commons.get_padding(kernel_size, 1), 0), ) ), ] ) self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0))) - def forward(self, x): + def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, list[torch.Tensor]]: fmap = [] # 1d to 2d @@ -640,7 +692,7 @@ class DiscriminatorP(torch.nn.Module): class DiscriminatorS(torch.nn.Module): - def __init__(self, use_spectral_norm=False): + def __init__(self, use_spectral_norm: bool = False) -> None: super(DiscriminatorS, self).__init__() norm_f = weight_norm if use_spectral_norm is False else spectral_norm self.convs = nn.ModuleList( @@ -655,7 +707,7 @@ class DiscriminatorS(torch.nn.Module): ) self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1)) - def forward(self, x): + def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, list[torch.Tensor]]: fmap = [] for layer in self.convs: @@ -670,7 +722,7 @@ class DiscriminatorS(torch.nn.Module): class MultiPeriodDiscriminator(torch.nn.Module): - def __init__(self, use_spectral_norm=False): + def __init__(self, use_spectral_norm: bool = False) -> None: super(MultiPeriodDiscriminator, self).__init__() periods = [2, 3, 5, 7, 11] @@ -680,7 +732,11 @@ class MultiPeriodDiscriminator(torch.nn.Module): ] self.discriminators = nn.ModuleList(discs) - def forward(self, y, y_hat): + def forward( + self, + y: torch.Tensor, + y_hat: torch.Tensor, + ) -> tuple[list[torch.Tensor], list[torch.Tensor], list[torch.Tensor], list[torch.Tensor]]: y_d_rs = [] y_d_gs = [] fmap_rs = [] @@ -702,7 +758,7 @@ class ReferenceEncoder(nn.Module): outputs --- [N, ref_enc_gru_size] """ - def __init__(self, spec_channels, gin_channels=0): + def __init__(self, spec_channels: int, gin_channels: int = 0) -> None: super().__init__() self.spec_channels = spec_channels ref_enc_filters = [32, 32, 64, 64, 128, 128] @@ -731,7 +787,7 @@ class ReferenceEncoder(nn.Module): ) self.proj = nn.Linear(128, gin_channels) - def forward(self, inputs, mask=None): + def forward(self, inputs: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor: N = inputs.size(0) out = inputs.view(N, 1, -1, self.spec_channels) # [N, 1, Ty, n_freqs] for conv in self.convs: @@ -749,7 +805,7 @@ class ReferenceEncoder(nn.Module): return self.proj(out.squeeze(0)) - def calculate_channels(self, L, kernel_size, stride, pad, n_convs): + def calculate_channels(self, L: int, kernel_size: int, stride: int, pad: int, n_convs: int) -> int: for i in range(n_convs): L = (L - kernel_size + 2 * pad) // stride + 1 return L @@ -762,31 +818,31 @@ class SynthesizerTrn(nn.Module): def __init__( self, - n_vocab, - spec_channels, - segment_size, - inter_channels, - hidden_channels, - filter_channels, - n_heads, - n_layers, - kernel_size, - p_dropout, - resblock, - resblock_kernel_sizes, - resblock_dilation_sizes, - upsample_rates, - upsample_initial_channel, - upsample_kernel_sizes, - n_speakers=256, - gin_channels=256, - use_sdp=True, - n_flow_layer=4, - n_layers_trans_flow=4, - flow_share_parameter=False, - use_transformer_flow=True, - **kwargs, - ): + n_vocab: int, + spec_channels: int, + segment_size: int, + inter_channels: int, + hidden_channels: int, + filter_channels: int, + n_heads: int, + n_layers: int, + kernel_size: int, + p_dropout: float, + resblock: str, + resblock_kernel_sizes: list[int], + resblock_dilation_sizes: list[list[int]], + upsample_rates: list[int], + upsample_initial_channel: int, + upsample_kernel_sizes: list[int], + n_speakers: int = 256, + gin_channels: int = 256, + use_sdp: bool = True, + n_flow_layer: int = 4, + n_layers_trans_flow: int = 4, + flow_share_parameter: bool = False, + use_transformer_flow: bool = True, + **kwargs: Any, + ) -> None: super().__init__() self.n_vocab = n_vocab self.spec_channels = spec_channels @@ -884,18 +940,27 @@ class SynthesizerTrn(nn.Module): def forward( self, - x, - x_lengths, - y, - y_lengths, - sid, - tone, - language, - bert, - ja_bert, - en_bert, - style_vec, - ): + x: torch.Tensor, + x_lengths: torch.Tensor, + y: torch.Tensor, + y_lengths: torch.Tensor, + sid: torch.Tensor, + tone: torch.Tensor, + language: torch.Tensor, + bert: torch.Tensor, + ja_bert: torch.Tensor, + en_bert: torch.Tensor, + style_vec: torch.Tensor, + ) -> tuple[ + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, + tuple[torch.Tensor, ...], + tuple[torch.Tensor, ...], + ]: if self.n_speakers > 0: g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1] else: @@ -973,27 +1038,28 @@ class SynthesizerTrn(nn.Module): def infer( self, - x, - x_lengths, - sid, - tone, - language, - bert, - ja_bert, - en_bert, - style_vec, - noise_scale=0.667, - length_scale=1.0, - noise_scale_w=0.8, - max_len=None, - sdp_ratio=0.0, - y=None, - ): + x: torch.Tensor, + x_lengths: torch.Tensor, + sid: torch.Tensor, + tone: torch.Tensor, + language: torch.Tensor, + bert: torch.Tensor, + ja_bert: torch.Tensor, + en_bert: torch.Tensor, + style_vec: torch.Tensor, + noise_scale: float = 0.667, + length_scale: float = 1.0, + noise_scale_w: float = 0.8, + max_len: Optional[int] = None, + sdp_ratio: float = 0.0, + y: Optional[torch.Tensor] = None, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, tuple[torch.Tensor, ...]]: # x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert) # g = self.gst(y) if self.n_speakers > 0: g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1] else: + assert y is not None g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1) x, m_p, logs_p, x_mask = self.enc_p( x, x_lengths, tone, language, bert, ja_bert, en_bert, style_vec, sid, g=g diff --git a/style_bert_vits2/models/models_jp_extra.py b/style_bert_vits2/models/models_jp_extra.py index 3a43d51..a43a715 100644 --- a/style_bert_vits2/models/models_jp_extra.py +++ b/style_bert_vits2/models/models_jp_extra.py @@ -1,4 +1,5 @@ import math +from typing import Any, Optional import torch from torch import nn @@ -10,13 +11,18 @@ from style_bert_vits2.models import attentions from style_bert_vits2.models import commons from style_bert_vits2.models import modules from style_bert_vits2.models import monotonic_alignment -from style_bert_vits2.nlp.symbols import SYMBOLS, NUM_TONES, NUM_LANGUAGES +from style_bert_vits2.nlp.symbols import NUM_LANGUAGES, NUM_TONES, SYMBOLS class DurationDiscriminator(nn.Module): # vits2 def __init__( - self, in_channels, filter_channels, kernel_size, p_dropout, gin_channels=0 - ): + self, + in_channels: int, + filter_channels: int, + kernel_size: int, + p_dropout: float, + gin_channels: int = 0 + ) -> None: super().__init__() self.in_channels = in_channels @@ -47,7 +53,7 @@ class DurationDiscriminator(nn.Module): # vits2 nn.Linear(2 * filter_channels, 1), nn.Sigmoid() ) - def forward_probability(self, x, dur): + def forward_probability(self, x: torch.Tensor, dur: torch.Tensor) -> torch.Tensor: dur = self.dur_proj(dur) x = torch.cat([x, dur], dim=1) x = x.transpose(1, 2) @@ -55,7 +61,14 @@ class DurationDiscriminator(nn.Module): # vits2 output_prob = self.output_layer(x) return output_prob - def forward(self, x, x_mask, dur_r, dur_hat, g=None): + def forward( + self, + x: torch.Tensor, + x_mask: torch.Tensor, + dur_r: torch.Tensor, + dur_hat: torch.Tensor, + g: Optional[torch.Tensor] = None, + ) -> list[torch.Tensor]: x = torch.detach(x) if g is not None: g = torch.detach(g) @@ -80,17 +93,17 @@ class DurationDiscriminator(nn.Module): # vits2 class TransformerCouplingBlock(nn.Module): def __init__( self, - channels, - hidden_channels, - filter_channels, - n_heads, - n_layers, - kernel_size, - p_dropout, - n_flows=4, - gin_channels=0, - share_parameter=False, - ): + channels: int, + hidden_channels: int, + filter_channels: int, + n_heads: int, + n_layers: int, + kernel_size: int, + p_dropout: float, + n_flows: int = 4, + gin_channels: int = 0, + share_parameter: bool = False, + ) -> None: super().__init__() self.channels = channels self.hidden_channels = hidden_channels @@ -102,16 +115,17 @@ class TransformerCouplingBlock(nn.Module): self.flows = nn.ModuleList() self.wn = ( - attentions.FFT( - hidden_channels, - filter_channels, - n_heads, - n_layers, - kernel_size, - p_dropout, - isflow=True, - gin_channels=self.gin_channels, - ) + # attentions.FFT( + # hidden_channels, + # filter_channels, + # n_heads, + # n_layers, + # kernel_size, + # p_dropout, + # isflow=True, + # gin_channels=self.gin_channels, + # ) + None if share_parameter else None ) @@ -133,7 +147,13 @@ class TransformerCouplingBlock(nn.Module): ) self.flows.append(modules.Flip()) - def forward(self, x, x_mask, g=None, reverse=False): + def forward( + self, + x: torch.Tensor, + x_mask: torch.Tensor, + g: Optional[torch.Tensor] = None, + reverse: bool = False, + ) -> torch.Tensor: if not reverse: for flow in self.flows: x, _ = flow(x, x_mask, g=g, reverse=reverse) @@ -146,13 +166,13 @@ class TransformerCouplingBlock(nn.Module): class StochasticDurationPredictor(nn.Module): def __init__( self, - in_channels, - filter_channels, - kernel_size, - p_dropout, - n_flows=4, - gin_channels=0, - ): + in_channels: int, + filter_channels: int, + kernel_size: int, + p_dropout: float, + n_flows: int = 4, + gin_channels: int = 0, + ) -> None: super().__init__() filter_channels = in_channels # it needs to be removed from future version. self.in_channels = in_channels @@ -192,7 +212,15 @@ class StochasticDurationPredictor(nn.Module): if gin_channels != 0: self.cond = nn.Conv1d(gin_channels, filter_channels, 1) - def forward(self, x, x_mask, w=None, g=None, reverse=False, noise_scale=1.0): + def forward( + self, + x: torch.Tensor, + x_mask: torch.Tensor, + w: Optional[torch.Tensor] = None, + g: Optional[torch.Tensor] = None, + reverse: bool = False, + noise_scale: float = 1.0, + ) -> torch.Tensor: x = torch.detach(x) x = self.pre(x) if g is not None: @@ -256,8 +284,13 @@ class StochasticDurationPredictor(nn.Module): class DurationPredictor(nn.Module): def __init__( - self, in_channels, filter_channels, kernel_size, p_dropout, gin_channels=0 - ): + self, + in_channels: int, + filter_channels: int, + kernel_size: int, + p_dropout: float, + gin_channels: int = 0, + ) -> None: super().__init__() self.in_channels = in_channels @@ -280,7 +313,7 @@ class DurationPredictor(nn.Module): if gin_channels != 0: self.cond = nn.Conv1d(gin_channels, in_channels, 1) - def forward(self, x, x_mask, g=None): + def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor: x = torch.detach(x) if g is not None: g = torch.detach(g) @@ -298,14 +331,14 @@ class DurationPredictor(nn.Module): class Bottleneck(nn.Sequential): - def __init__(self, in_dim, hidden_dim): + def __init__(self, in_dim: int, hidden_dim: int) -> None: c_fc1 = nn.Linear(in_dim, hidden_dim, bias=False) c_fc2 = nn.Linear(in_dim, hidden_dim, bias=False) - super().__init__(*[c_fc1, c_fc2]) + super().__init__(c_fc1, c_fc2) class Block(nn.Module): - def __init__(self, in_dim, hidden_dim) -> None: + def __init__(self, in_dim: int, hidden_dim: int) -> None: super().__init__() self.norm = nn.LayerNorm(in_dim) self.mlp = MLP(in_dim, hidden_dim) @@ -316,13 +349,13 @@ class Block(nn.Module): class MLP(nn.Module): - def __init__(self, in_dim, hidden_dim): + def __init__(self, in_dim: int, hidden_dim: int) -> None: super().__init__() self.c_fc1 = nn.Linear(in_dim, hidden_dim, bias=False) self.c_fc2 = nn.Linear(in_dim, hidden_dim, bias=False) self.c_proj = nn.Linear(hidden_dim, in_dim, bias=False) - def forward(self, x: torch.Tensor): + def forward(self, x: torch.Tensor) -> torch.Tensor: x = F.silu(self.c_fc1(x)) * self.c_fc2(x) x = self.c_proj(x) return x @@ -331,16 +364,16 @@ class MLP(nn.Module): class TextEncoder(nn.Module): def __init__( self, - n_vocab, - out_channels, - hidden_channels, - filter_channels, - n_heads, - n_layers, - kernel_size, - p_dropout, - gin_channels=0, - ): + n_vocab: int, + out_channels: int, + hidden_channels: int, + filter_channels: int, + n_heads: int, + n_layers: int, + kernel_size: int, + p_dropout: float, + gin_channels: int = 0, + ) -> None: super().__init__() self.n_vocab = n_vocab self.out_channels = out_channels @@ -373,7 +406,16 @@ class TextEncoder(nn.Module): ) self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1) - def forward(self, x, x_lengths, tone, language, bert, style_vec, g=None): + def forward( + self, + x: torch.Tensor, + x_lengths: torch.Tensor, + tone: torch.Tensor, + language: torch.Tensor, + bert: torch.Tensor, + style_vec: torch.Tensor, + g: Optional[torch.Tensor] = None, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: bert_emb = self.bert_proj(bert).transpose(1, 2) style_emb = self.style_proj(style_vec.unsqueeze(1)) x = ( @@ -400,14 +442,14 @@ class TextEncoder(nn.Module): class ResidualCouplingBlock(nn.Module): def __init__( self, - channels, - hidden_channels, - kernel_size, - dilation_rate, - n_layers, - n_flows=4, - gin_channels=0, - ): + channels: int, + hidden_channels: int, + kernel_size: int, + dilation_rate: int, + n_layers: int, + n_flows: int = 4, + gin_channels: int = 0, + ) -> None: super().__init__() self.channels = channels self.hidden_channels = hidden_channels @@ -432,7 +474,13 @@ class ResidualCouplingBlock(nn.Module): ) self.flows.append(modules.Flip()) - def forward(self, x, x_mask, g=None, reverse=False): + def forward( + self, + x: torch.Tensor, + x_mask: torch.Tensor, + g: Optional[torch.Tensor] = None, + reverse: bool = False, + ) -> torch.Tensor: if not reverse: for flow in self.flows: x, _ = flow(x, x_mask, g=g, reverse=reverse) @@ -445,14 +493,14 @@ class ResidualCouplingBlock(nn.Module): class PosteriorEncoder(nn.Module): def __init__( self, - in_channels, - out_channels, - hidden_channels, - kernel_size, - dilation_rate, - n_layers, - gin_channels=0, - ): + in_channels: int, + out_channels: int, + hidden_channels: int, + kernel_size: int, + dilation_rate: int, + n_layers: int, + gin_channels: int = 0, + ) -> None: super().__init__() self.in_channels = in_channels self.out_channels = out_channels @@ -472,7 +520,12 @@ class PosteriorEncoder(nn.Module): ) self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1) - def forward(self, x, x_lengths, g=None): + def forward( + self, + x: torch.Tensor, + x_lengths: torch.Tensor, + g: Optional[torch.Tensor] = None, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to( x.dtype ) @@ -487,22 +540,22 @@ class PosteriorEncoder(nn.Module): class Generator(torch.nn.Module): def __init__( self, - initial_channel, - resblock, - resblock_kernel_sizes, - resblock_dilation_sizes, - upsample_rates, - upsample_initial_channel, - upsample_kernel_sizes, - gin_channels=0, - ): + initial_channel: int, + resblock_str: str, + resblock_kernel_sizes: list[int], + resblock_dilation_sizes: list[list[int]], + upsample_rates: list[int], + upsample_initial_channel: int, + upsample_kernel_sizes: list[int], + gin_channels: int = 0, + ) -> None: super(Generator, self).__init__() self.num_kernels = len(resblock_kernel_sizes) self.num_upsamples = len(upsample_rates) self.conv_pre = Conv1d( initial_channel, upsample_initial_channel, 7, 1, padding=3 ) - resblock = modules.ResBlock1 if resblock == "1" else modules.ResBlock2 + resblock = modules.ResBlock1 if resblock_str == "1" else modules.ResBlock2 self.ups = nn.ModuleList() for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)): @@ -519,20 +572,22 @@ class Generator(torch.nn.Module): ) self.resblocks = nn.ModuleList() + ch = None for i in range(len(self.ups)): ch = upsample_initial_channel // (2 ** (i + 1)) for j, (k, d) in enumerate( zip(resblock_kernel_sizes, resblock_dilation_sizes) ): - self.resblocks.append(resblock(ch, k, d)) + self.resblocks.append(resblock(ch, k, d)) # type: ignore + assert ch is not None self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False) self.ups.apply(commons.init_weights) if gin_channels != 0: self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1) - def forward(self, x, g=None): + def forward(self, x: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor: x = self.conv_pre(x) if g is not None: x = x + self.cond(g) @@ -546,6 +601,7 @@ class Generator(torch.nn.Module): xs = self.resblocks[i * self.num_kernels + j](x) else: xs += self.resblocks[i * self.num_kernels + j](x) + assert xs is not None x = xs / self.num_kernels x = F.leaky_relu(x) x = self.conv_post(x) @@ -553,7 +609,7 @@ class Generator(torch.nn.Module): return x - def remove_weight_norm(self): + def remove_weight_norm(self) -> None: print("Removing weight norm...") for layer in self.ups: remove_weight_norm(layer) @@ -562,7 +618,7 @@ class Generator(torch.nn.Module): class DiscriminatorP(torch.nn.Module): - def __init__(self, period, kernel_size=5, stride=3, use_spectral_norm=False): + def __init__(self, period: int, kernel_size: int = 5, stride: int = 3, use_spectral_norm: bool = False) -> None: super(DiscriminatorP, self).__init__() self.period = period self.use_spectral_norm = use_spectral_norm @@ -618,7 +674,7 @@ class DiscriminatorP(torch.nn.Module): ) self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0))) - def forward(self, x): + def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, list[torch.Tensor]]: fmap = [] # 1d to 2d @@ -641,7 +697,7 @@ class DiscriminatorP(torch.nn.Module): class DiscriminatorS(torch.nn.Module): - def __init__(self, use_spectral_norm=False): + def __init__(self, use_spectral_norm: bool = False) -> None: super(DiscriminatorS, self).__init__() norm_f = weight_norm if use_spectral_norm is False else spectral_norm self.convs = nn.ModuleList( @@ -656,7 +712,7 @@ class DiscriminatorS(torch.nn.Module): ) self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1)) - def forward(self, x): + def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, list[torch.Tensor]]: fmap = [] for layer in self.convs: @@ -671,7 +727,7 @@ class DiscriminatorS(torch.nn.Module): class MultiPeriodDiscriminator(torch.nn.Module): - def __init__(self, use_spectral_norm=False): + def __init__(self, use_spectral_norm: bool = False) -> None: super(MultiPeriodDiscriminator, self).__init__() periods = [2, 3, 5, 7, 11] @@ -681,7 +737,11 @@ class MultiPeriodDiscriminator(torch.nn.Module): ] self.discriminators = nn.ModuleList(discs) - def forward(self, y, y_hat): + def forward( + self, + y: torch.Tensor, + y_hat: torch.Tensor, + ) -> tuple[list[torch.Tensor], list[torch.Tensor], list[torch.Tensor], list[torch.Tensor]]: y_d_rs = [] y_d_gs = [] fmap_rs = [] @@ -701,8 +761,12 @@ class WavLMDiscriminator(nn.Module): """docstring for Discriminator.""" def __init__( - self, slm_hidden=768, slm_layers=13, initial_channel=64, use_spectral_norm=False - ): + self, + slm_hidden: int = 768, + slm_layers: int = 13, + initial_channel: int = 64, + use_spectral_norm: bool = False, + ) -> None: super(WavLMDiscriminator, self).__init__() norm_f = weight_norm if use_spectral_norm == False else spectral_norm self.pre = norm_f( @@ -732,7 +796,7 @@ class WavLMDiscriminator(nn.Module): self.conv_post = norm_f(Conv1d(initial_channel * 4, 1, 3, 1, padding=1)) - def forward(self, x): + def forward(self, x: torch.Tensor) -> torch.Tensor: x = self.pre(x) fmap = [] @@ -752,7 +816,7 @@ class ReferenceEncoder(nn.Module): outputs --- [N, ref_enc_gru_size] """ - def __init__(self, spec_channels, gin_channels=0): + def __init__(self, spec_channels: int, gin_channels: int = 0) -> None: super().__init__() self.spec_channels = spec_channels ref_enc_filters = [32, 32, 64, 64, 128, 128] @@ -781,7 +845,7 @@ class ReferenceEncoder(nn.Module): ) self.proj = nn.Linear(128, gin_channels) - def forward(self, inputs, mask=None): + def forward(self, inputs: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor: N = inputs.size(0) out = inputs.view(N, 1, -1, self.spec_channels) # [N, 1, Ty, n_freqs] for conv in self.convs: @@ -799,7 +863,7 @@ class ReferenceEncoder(nn.Module): return self.proj(out.squeeze(0)) - def calculate_channels(self, L, kernel_size, stride, pad, n_convs): + def calculate_channels(self, L: int, kernel_size: int, stride: int, pad: int, n_convs: int) -> int: for i in range(n_convs): L = (L - kernel_size + 2 * pad) // stride + 1 return L @@ -812,31 +876,31 @@ class SynthesizerTrn(nn.Module): def __init__( self, - n_vocab, - spec_channels, - segment_size, - inter_channels, - hidden_channels, - filter_channels, - n_heads, - n_layers, - kernel_size, - p_dropout, - resblock, - resblock_kernel_sizes, - resblock_dilation_sizes, - upsample_rates, - upsample_initial_channel, - upsample_kernel_sizes, - n_speakers=256, - gin_channels=256, - use_sdp=True, - n_flow_layer=4, - n_layers_trans_flow=6, - flow_share_parameter=False, - use_transformer_flow=True, - **kwargs - ): + n_vocab: int, + spec_channels: int, + segment_size: int, + inter_channels: int, + hidden_channels: int, + filter_channels: int, + n_heads: int, + n_layers: int, + kernel_size: int, + p_dropout: float, + resblock: str, + resblock_kernel_sizes: list[int], + resblock_dilation_sizes: list[list[int]], + upsample_rates: list[int], + upsample_initial_channel: int, + upsample_kernel_sizes: list[int], + n_speakers: int = 256, + gin_channels: int = 256, + use_sdp: bool = True, + n_flow_layer: int = 4, + n_layers_trans_flow: int = 6, + flow_share_parameter: bool = False, + use_transformer_flow: bool = True, + **kwargs: Any, + ) -> None: super().__init__() self.n_vocab = n_vocab self.spec_channels = spec_channels @@ -933,16 +997,26 @@ class SynthesizerTrn(nn.Module): def forward( self, - x, - x_lengths, - y, - y_lengths, - sid, - tone, - language, - bert, - style_vec, - ): + x: torch.Tensor, + x_lengths: torch.Tensor, + y: torch.Tensor, + y_lengths: torch.Tensor, + sid: torch.Tensor, + tone: torch.Tensor, + language: torch.Tensor, + bert: torch.Tensor, + style_vec: torch.Tensor, + ) -> tuple[ + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, + tuple[torch.Tensor, ...], + tuple[torch.Tensor, ...], + ]: if self.n_speakers > 0: g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1] else: @@ -1014,32 +1088,33 @@ class SynthesizerTrn(nn.Module): ids_slice, x_mask, y_mask, - (z, z_p, m_p, logs_p, m_q, logs_q), + (z, z_p, m_p, logs_p, m_q, logs_q), # type: ignore (x, logw, logw_), # , logw_sdp), g, ) def infer( self, - x, - x_lengths, - sid, - tone, - language, - bert, - style_vec, - noise_scale=0.667, - length_scale=1.0, - noise_scale_w=0.8, - max_len=None, - sdp_ratio=0.0, - y=None, - ): + x: torch.Tensor, + x_lengths: torch.Tensor, + sid: torch.Tensor, + tone: torch.Tensor, + language: torch.Tensor, + bert: torch.Tensor, + style_vec: torch.Tensor, + noise_scale: float = 0.667, + length_scale: float = 1.0, + noise_scale_w: float = 0.8, + max_len: Optional[int] = None, + sdp_ratio: float = 0.0, + y: Optional[torch.Tensor] = None, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, tuple[torch.Tensor, ...]]: # x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert) # g = self.gst(y) if self.n_speakers > 0: g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1] else: + assert y is not None g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1) x, m_p, logs_p, x_mask = self.enc_p( x, x_lengths, tone, language, bert, style_vec, g=g diff --git a/style_bert_vits2/models/modules.py b/style_bert_vits2/models/modules.py index df38071..8eed963 100644 --- a/style_bert_vits2/models/modules.py +++ b/style_bert_vits2/models/modules.py @@ -16,7 +16,7 @@ LRELU_SLOPE = 0.1 class LayerNorm(nn.Module): - def __init__(self, channels: int, eps: float = 1e-5): + def __init__(self, channels: int, eps: float = 1e-5) -> None: super().__init__() self.channels = channels self.eps = eps @@ -39,7 +39,7 @@ class ConvReluNorm(nn.Module): kernel_size: int, n_layers: int, p_dropout: float, - ): + ) -> None: super().__init__() self.in_channels = in_channels self.hidden_channels = hidden_channels @@ -88,7 +88,7 @@ class DDSConv(nn.Module): Dialted and Depth-Separable Convolution """ - def __init__(self, channels: int, kernel_size: int, n_layers: int, p_dropout: float = 0.0): + def __init__(self, channels: int, kernel_size: int, n_layers: int, p_dropout: float = 0.0) -> None: super().__init__() self.channels = channels self.kernel_size = kernel_size @@ -141,7 +141,7 @@ class WN(torch.nn.Module): n_layers: int, gin_channels: int = 0, p_dropout: float = 0, - ): + ) -> None: super(WN, self).__init__() assert kernel_size % 2 == 1 self.hidden_channels = hidden_channels @@ -221,7 +221,7 @@ class WN(torch.nn.Module): class ResBlock1(torch.nn.Module): - def __init__(self, channels: int, kernel_size: int = 3, dilation: tuple[int, int, int] = (1, 3, 5)): + def __init__(self, channels: int, kernel_size: int = 3, dilation: tuple[int, int, int] = (1, 3, 5)) -> None: super(ResBlock1, self).__init__() self.convs1 = nn.ModuleList( [ @@ -318,7 +318,7 @@ class ResBlock1(torch.nn.Module): class ResBlock2(torch.nn.Module): - def __init__(self, channels: int, kernel_size: int = 3, dilation: tuple[int, int] = (1, 3)): + def __init__(self, channels: int, kernel_size: int = 3, dilation: tuple[int, int] = (1, 3)) -> None: super(ResBlock2, self).__init__() self.convs = nn.ModuleList( [ @@ -363,7 +363,13 @@ class ResBlock2(torch.nn.Module): class Log(nn.Module): - def forward(self, x: torch.Tensor, x_mask: torch.Tensor, reverse: bool = False, **kwargs: Any): + def forward( + self, + x: torch.Tensor, + x_mask: torch.Tensor, + reverse: bool = False, + **kwargs: Any, + ) -> Union[tuple[torch.Tensor, torch.Tensor], torch.Tensor]: if not reverse: y = torch.log(torch.clamp_min(x, 1e-5)) * x_mask logdet = torch.sum(-y, [1, 2]) @@ -390,7 +396,7 @@ class Flip(nn.Module): class ElementwiseAffine(nn.Module): - def __init__(self, channels: int): + def __init__(self, channels: int) -> None: super().__init__() self.channels = channels self.m = nn.Parameter(torch.zeros(channels, 1)) @@ -424,7 +430,7 @@ class ResidualCouplingLayer(nn.Module): p_dropout: float = 0, gin_channels: int = 0, mean_only: bool = False, - ): + ) -> None: assert channels % 2 == 0, "channels should be divisible by 2" super().__init__() self.channels = channels @@ -449,7 +455,13 @@ class ResidualCouplingLayer(nn.Module): assert self.post.bias is not None self.post.bias.data.zero_() - def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None, reverse: bool = False): + def forward( + self, + x: torch.Tensor, + x_mask: torch.Tensor, + g: Optional[torch.Tensor] = None, + reverse: bool = False, + ) -> Union[tuple[torch.Tensor, torch.Tensor], torch.Tensor]: x0, x1 = torch.split(x, [self.half_channels] * 2, 1) h = self.pre(x0) * x_mask h = self.enc(h, x_mask, g=g) @@ -480,7 +492,7 @@ class ConvFlow(nn.Module): n_layers: int, num_bins: int = 10, tail_bound: float = 5.0, - ): + ) -> None: super().__init__() self.in_channels = in_channels self.filter_channels = filter_channels @@ -499,7 +511,13 @@ class ConvFlow(nn.Module): assert self.proj.bias is not None self.proj.bias.data.zero_() - def forward(self, x: torch.Tensor, x_mask: torch.Tensor, g: Optional[torch.Tensor] = None, reverse: bool = False): + def forward( + self, + x: torch.Tensor, + x_mask: torch.Tensor, + g: Optional[torch.Tensor] = None, + reverse: bool = False, + ) -> Union[tuple[torch.Tensor, torch.Tensor], torch.Tensor]: x0, x1 = torch.split(x, [self.half_channels] * 2, 1) h = self.pre(x0) h = self.convs(h, x_mask, g=g) From c594f7ea7a51e225e5bb980b8d1a947cac07f32e Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sat, 9 Mar 2024 00:26:51 +0000 Subject: [PATCH 48/64] Refactor: change execution location of pyopenjtalk.initialize() Considering library design, this function with many side effects should not be executed in a library. --- server_editor.py | 8 ++++++-- server_fastapi.py | 3 +-- style_bert_vits2/nlp/japanese/g2p.py | 8 -------- style_bert_vits2/nlp/japanese/user_dict/__init__.py | 4 ---- webui/inference.py | 3 +-- 5 files changed, 8 insertions(+), 18 deletions(-) diff --git a/server_editor.py b/server_editor.py index 75fa472..9a6a08a 100644 --- a/server_editor.py +++ b/server_editor.py @@ -42,6 +42,7 @@ from style_bert_vits2.constants import ( ) from style_bert_vits2.logging import logger from style_bert_vits2.nlp import bert_models +from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone from style_bert_vits2.nlp.japanese.normalizer import normalize_text from style_bert_vits2.nlp.japanese.user_dict import ( @@ -148,8 +149,11 @@ def save_last_download(latest_release): # ---フロントエンド部分に関する処理ここまで--- # 以降はAPIの設定 -# 最初に pyopenjtalk の辞書を更新 -## pyopenjtalk_worker の起動も同時に行われる +# pyopenjtalk_worker を起動 +## pyopenjtalk_worker は TCP ソケットサーバーのため、ここで起動する +pyopenjtalk.initialize() + +# pyopenjtalk の辞書を更新 update_dict() # 事前に BERT モデル/トークナイザーをロードしておく diff --git a/server_fastapi.py b/server_fastapi.py index d8da956..e8309da 100644 --- a/server_fastapi.py +++ b/server_fastapi.py @@ -42,8 +42,7 @@ ln = config.server_config.language # pyopenjtalk_worker を起動 -## Gradio はマルチスレッドだが、initialize() 内部で利用されている signal はマルチスレッドから設定できない -## さらに起動には若干時間がかかるため、事前に起動しておいた方が体験が良い +## pyopenjtalk_worker は TCP ソケットサーバーのため、ここで起動する pyopenjtalk.initialize() # 事前に BERT モデル/トークナイザーをロードしておく diff --git a/style_bert_vits2/nlp/japanese/g2p.py b/style_bert_vits2/nlp/japanese/g2p.py index 70b4e81..7fc97f2 100644 --- a/style_bert_vits2/nlp/japanese/g2p.py +++ b/style_bert_vits2/nlp/japanese/g2p.py @@ -112,10 +112,6 @@ def text_to_sep_kata( tuple[list[str], list[str]]: 分割された単語リストと、その読み(カタカナ or 記号1文字)のリスト """ - # pyopenjtalk_worker を初期化 - ## 一度 worker を起動すれば、明示的に終了するかプロセス終了まで同一の worker に接続される - pyopenjtalk.initialize() - # parsed: OpenJTalkの解析結果 parsed = pyopenjtalk.run_frontend(norm_text) sep_text: list[str] = [] @@ -249,10 +245,6 @@ def __pyopenjtalk_g2p_prosody(text: str, drop_unvoiced_vowels: bool = True) -> l return -50 return int(match.group(1)) - # pyopenjtalk_worker を初期化 - ## 一度 worker を起動すれば、明示的に終了するかプロセス終了まで同一の worker に接続される - pyopenjtalk.initialize() - labels = pyopenjtalk.make_label(pyopenjtalk.run_frontend(text)) N = len(labels) diff --git a/style_bert_vits2/nlp/japanese/user_dict/__init__.py b/style_bert_vits2/nlp/japanese/user_dict/__init__.py index 1887c13..a2cc43e 100644 --- a/style_bert_vits2/nlp/japanese/user_dict/__init__.py +++ b/style_bert_vits2/nlp/japanese/user_dict/__init__.py @@ -80,10 +80,6 @@ def update_dict( コンパイル済み辞書ファイルのパス """ - # pyopenjtalk_worker を初期化 - ## 一度 worker を起動すれば、明示的に終了するかプロセス終了まで同一の worker に接続される - pyopenjtalk.initialize() - random_string = uuid4() tmp_csv_path = compiled_dict_path.with_suffix( f".dict_csv-{random_string}.tmp" diff --git a/webui/inference.py b/webui/inference.py index ef1f97f..536b38b 100644 --- a/webui/inference.py +++ b/webui/inference.py @@ -27,8 +27,7 @@ from style_bert_vits2.tts_model import ModelHolder # pyopenjtalk_worker を起動 -## Gradio はマルチスレッドだが、initialize() 内部で利用されている signal はマルチスレッドから設定できない -## さらに起動には若干時間がかかるため、事前に起動しておいた方が体験が良い +## pyopenjtalk_worker は TCP ソケットサーバーのため、ここで起動する pyopenjtalk.initialize() # 事前に BERT モデル/トークナイザーをロードしておく From 98ab8e79789cdf6f8f95d92a769d391994b8ead5 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sat, 9 Mar 2024 15:58:57 +0000 Subject: [PATCH 49/64] Refactor: separate module for utilities related to loading/saving checkpoints and safetensors --- style_bert_vits2/models/infer.py | 4 +- style_bert_vits2/models/utils.py | 357 ------------------- style_bert_vits2/models/utils/__init__.py | 156 ++++++++ style_bert_vits2/models/utils/checkpoints.py | 194 ++++++++++ style_bert_vits2/models/utils/safetensors.py | 91 +++++ train_ms.py | 71 ++-- train_ms_jp_extra.py | 94 ++--- 7 files changed, 510 insertions(+), 457 deletions(-) delete mode 100644 style_bert_vits2/models/utils.py create mode 100644 style_bert_vits2/models/utils/__init__.py create mode 100644 style_bert_vits2/models/utils/checkpoints.py create mode 100644 style_bert_vits2/models/utils/safetensors.py diff --git a/style_bert_vits2/models/infer.py b/style_bert_vits2/models/infer.py index 93bf8c2..d3ec963 100644 --- a/style_bert_vits2/models/infer.py +++ b/style_bert_vits2/models/infer.py @@ -80,9 +80,9 @@ def get_net_g(model_path: str, version: str, device: str, hps: HyperParameters): net_g.state_dict() _ = net_g.eval() if model_path.endswith(".pth") or model_path.endswith(".pt"): - _ = utils.load_checkpoint(model_path, net_g, None, skip_optimizer=True) + _ = utils.checkpoints.load_checkpoint(model_path, net_g, None, skip_optimizer=True) elif model_path.endswith(".safetensors"): - _ = utils.load_safetensors(model_path, net_g, True) + _ = utils.safetensors.load_safetensors(model_path, net_g, True) else: raise ValueError(f"Unknown model format: {model_path}") return net_g diff --git a/style_bert_vits2/models/utils.py b/style_bert_vits2/models/utils.py deleted file mode 100644 index 37ea269..0000000 --- a/style_bert_vits2/models/utils.py +++ /dev/null @@ -1,357 +0,0 @@ -import argparse -import glob -import json -import logging -import os -import re -import subprocess - -import numpy as np -import torch -from safetensors import safe_open -from safetensors.torch import save_file -from scipy.io.wavfile import read - -from style_bert_vits2.logging import logger - - -MATPLOTLIB_FLAG = False - - -def load_checkpoint( - checkpoint_path, model, optimizer=None, skip_optimizer=False, for_infer=False -): - assert os.path.isfile(checkpoint_path) - checkpoint_dict = torch.load(checkpoint_path, map_location="cpu") - iteration = checkpoint_dict["iteration"] - learning_rate = checkpoint_dict["learning_rate"] - logger.info( - f"Loading model and optimizer at iteration {iteration} from {checkpoint_path}" - ) - if ( - optimizer is not None - and not skip_optimizer - and checkpoint_dict["optimizer"] is not None - ): - optimizer.load_state_dict(checkpoint_dict["optimizer"]) - elif optimizer is None and not skip_optimizer: - # else: Disable this line if Infer and resume checkpoint,then enable the line upper - new_opt_dict = optimizer.state_dict() - new_opt_dict_params = new_opt_dict["param_groups"][0]["params"] - new_opt_dict["param_groups"] = checkpoint_dict["optimizer"]["param_groups"] - new_opt_dict["param_groups"][0]["params"] = new_opt_dict_params - optimizer.load_state_dict(new_opt_dict) - - saved_state_dict = checkpoint_dict["model"] - if hasattr(model, "module"): - state_dict = model.module.state_dict() - else: - state_dict = model.state_dict() - - new_state_dict = {} - for k, v in state_dict.items(): - try: - # assert "emb_g" not in k - new_state_dict[k] = saved_state_dict[k] - assert saved_state_dict[k].shape == v.shape, ( - saved_state_dict[k].shape, - v.shape, - ) - except: - # For upgrading from the old version - if "ja_bert_proj" in k: - v = torch.zeros_like(v) - logger.warning( - f"Seems you are using the old version of the model, the {k} is automatically set to zero for backward compatibility" - ) - elif "enc_q" in k and for_infer: - continue - else: - logger.error(f"{k} is not in the checkpoint {checkpoint_path}") - - new_state_dict[k] = v - - if hasattr(model, "module"): - model.module.load_state_dict(new_state_dict, strict=False) - else: - model.load_state_dict(new_state_dict, strict=False) - - logger.info("Loaded '{}' (iteration {})".format(checkpoint_path, iteration)) - - return model, optimizer, learning_rate, iteration - - -def save_checkpoint(model, optimizer, learning_rate, iteration, checkpoint_path): - logger.info( - "Saving model and optimizer state at iteration {} to {}".format( - iteration, checkpoint_path - ) - ) - if hasattr(model, "module"): - state_dict = model.module.state_dict() - else: - state_dict = model.state_dict() - torch.save( - { - "model": state_dict, - "iteration": iteration, - "optimizer": optimizer.state_dict(), - "learning_rate": learning_rate, - }, - checkpoint_path, - ) - - -def clean_checkpoints(path_to_models="logs/44k/", n_ckpts_to_keep=2, sort_by_time=True): - """Freeing up space by deleting saved ckpts - - Arguments: - path_to_models -- Path to the model directory - n_ckpts_to_keep -- Number of ckpts to keep, excluding G_0.pth and D_0.pth - sort_by_time -- True -> chronologically delete ckpts - False -> lexicographically delete ckpts - """ - import re - - ckpts_files = [ - f - for f in os.listdir(path_to_models) - if os.path.isfile(os.path.join(path_to_models, f)) - ] - - def name_key(_f): - return int(re.compile("._(\\d+)\\.pth").match(_f).group(1)) - - def time_key(_f): - return os.path.getmtime(os.path.join(path_to_models, _f)) - - sort_key = time_key if sort_by_time else name_key - - def x_sorted(_x): - return sorted( - [f for f in ckpts_files if f.startswith(_x) and not f.endswith("_0.pth")], - key=sort_key, - ) - - to_del = [ - os.path.join(path_to_models, fn) - for fn in ( - x_sorted("G_")[:-n_ckpts_to_keep] - + x_sorted("D_")[:-n_ckpts_to_keep] - + x_sorted("WD_")[:-n_ckpts_to_keep] - + x_sorted("DUR_")[:-n_ckpts_to_keep] - ) - ] - - def del_info(fn): - return logger.info(f"Free up space by deleting ckpt {fn}") - - def del_routine(x): - return [os.remove(x), del_info(x)] - - [del_routine(fn) for fn in to_del] - - -def load_safetensors(checkpoint_path, model, for_infer=False): - """ - Load safetensors model. - """ - - tensors = {} - iteration = None - with safe_open(checkpoint_path, framework="pt", device="cpu") as f: - for key in f.keys(): - if key == "iteration": - iteration = f.get_tensor(key).item() - tensors[key] = f.get_tensor(key) - if hasattr(model, "module"): - result = model.module.load_state_dict(tensors, strict=False) - else: - result = model.load_state_dict(tensors, strict=False) - for key in result.missing_keys: - if key.startswith("enc_q") and for_infer: - continue - logger.warning(f"Missing key: {key}") - for key in result.unexpected_keys: - if key == "iteration": - continue - logger.warning(f"Unexpected key: {key}") - if iteration is None: - logger.info(f"Loaded '{checkpoint_path}'") - else: - logger.info(f"Loaded '{checkpoint_path}' (iteration {iteration})") - return model, iteration - - -def save_safetensors(model, iteration, checkpoint_path, is_half=False, for_infer=False): - """ - Save model with safetensors. - """ - if hasattr(model, "module"): - state_dict = model.module.state_dict() - else: - state_dict = model.state_dict() - keys = [] - for k in state_dict: - if "enc_q" in k and for_infer: - continue # noqa: E701 - keys.append(k) - - new_dict = ( - {k: state_dict[k].half() for k in keys} - if is_half - else {k: state_dict[k] for k in keys} - ) - new_dict["iteration"] = torch.LongTensor([iteration]) - logger.info(f"Saved safetensors to {checkpoint_path}") - save_file(new_dict, checkpoint_path) - - -def summarize( - writer, - global_step, - scalars={}, - histograms={}, - images={}, - audios={}, - audio_sampling_rate=22050, -): - for k, v in scalars.items(): - writer.add_scalar(k, v, global_step) - for k, v in histograms.items(): - writer.add_histogram(k, v, global_step) - for k, v in images.items(): - writer.add_image(k, v, global_step, dataformats="HWC") - for k, v in audios.items(): - writer.add_audio(k, v, global_step, audio_sampling_rate) - - -def is_resuming(dir_path): - # JP-ExtraバージョンではDURがなくWDがあったり変わるため、Gのみで判断する - g_list = glob.glob(os.path.join(dir_path, "G_*.pth")) - # d_list = glob.glob(os.path.join(dir_path, "D_*.pth")) - # dur_list = glob.glob(os.path.join(dir_path, "DUR_*.pth")) - return len(g_list) > 0 - - -def latest_checkpoint_path(dir_path, regex="G_*.pth"): - f_list = glob.glob(os.path.join(dir_path, regex)) - f_list.sort(key=lambda f: int("".join(filter(str.isdigit, f)))) - try: - x = f_list[-1] - except IndexError: - raise ValueError(f"No checkpoint found in {dir_path} with regex {regex}") - return x - - -def plot_spectrogram_to_numpy(spectrogram): - global MATPLOTLIB_FLAG - if not MATPLOTLIB_FLAG: - import matplotlib - - matplotlib.use("Agg") - MATPLOTLIB_FLAG = True - mpl_logger = logging.getLogger("matplotlib") - mpl_logger.setLevel(logging.WARNING) - import matplotlib.pylab as plt - import numpy as np - - fig, ax = plt.subplots(figsize=(10, 2)) - im = ax.imshow(spectrogram, aspect="auto", origin="lower", interpolation="none") - plt.colorbar(im, ax=ax) - plt.xlabel("Frames") - plt.ylabel("Channels") - plt.tight_layout() - - fig.canvas.draw() - data = np.fromstring(fig.canvas.tostring_rgb(), dtype=np.uint8, sep="") - data = data.reshape(fig.canvas.get_width_height()[::-1] + (3,)) - plt.close() - return data - - -def plot_alignment_to_numpy(alignment, info=None): - global MATPLOTLIB_FLAG - if not MATPLOTLIB_FLAG: - import matplotlib - - matplotlib.use("Agg") - MATPLOTLIB_FLAG = True - mpl_logger = logging.getLogger("matplotlib") - mpl_logger.setLevel(logging.WARNING) - import matplotlib.pylab as plt - import numpy as np - - fig, ax = plt.subplots(figsize=(6, 4)) - im = ax.imshow( - alignment.transpose(), aspect="auto", origin="lower", interpolation="none" - ) - fig.colorbar(im, ax=ax) - xlabel = "Decoder timestep" - if info is not None: - xlabel += "\n\n" + info - plt.xlabel(xlabel) - plt.ylabel("Encoder timestep") - plt.tight_layout() - - fig.canvas.draw() - data = np.fromstring(fig.canvas.tostring_rgb(), dtype=np.uint8, sep="") - data = data.reshape(fig.canvas.get_width_height()[::-1] + (3,)) - plt.close() - return data - - -def load_wav_to_torch(full_path): - sampling_rate, data = read(full_path) - return torch.FloatTensor(data.astype(np.float32)), sampling_rate - - -def load_filepaths_and_text(filename, split="|"): - with open(filename, encoding="utf-8") as f: - filepaths_and_text = [line.strip().split(split) for line in f] - return filepaths_and_text - - -def get_logger(model_dir, filename="train.log"): - global logger - logger = logging.getLogger(os.path.basename(model_dir)) - logger.setLevel(logging.DEBUG) - - formatter = logging.Formatter("%(asctime)s\t%(name)s\t%(levelname)s\t%(message)s") - if not os.path.exists(model_dir): - os.makedirs(model_dir) - h = logging.FileHandler(os.path.join(model_dir, filename)) - h.setLevel(logging.DEBUG) - h.setFormatter(formatter) - logger.addHandler(h) - return logger - - -def get_steps(model_path): - matches = re.findall(r"\d+", model_path) - return matches[-1] if matches else None - - -def check_git_hash(model_dir): - source_dir = os.path.dirname(os.path.realpath(__file__)) - if not os.path.exists(os.path.join(source_dir, ".git")): - logger.warning( - "{} is not a git repository, therefore hash value comparison will be ignored.".format( - source_dir - ) - ) - return - - cur_hash = subprocess.getoutput("git rev-parse HEAD") - - path = os.path.join(model_dir, "githash") - if os.path.exists(path): - saved_hash = open(path).read() - if saved_hash != cur_hash: - logger.warning( - "git hash values are different. {}(saved) != {}(current)".format( - saved_hash[:8], cur_hash[:8] - ) - ) - else: - open(path, "w").write(cur_hash) diff --git a/style_bert_vits2/models/utils/__init__.py b/style_bert_vits2/models/utils/__init__.py new file mode 100644 index 0000000..f488e09 --- /dev/null +++ b/style_bert_vits2/models/utils/__init__.py @@ -0,0 +1,156 @@ +import glob +import logging +import os +import re +import subprocess + +import numpy as np +import torch +from scipy.io.wavfile import read + +from style_bert_vits2.logging import logger +from style_bert_vits2.models.utils import checkpoints # type: ignore +from style_bert_vits2.models.utils import safetensors # type: ignore + + +MATPLOTLIB_FLAG = False + + +def summarize( + writer, + global_step, + scalars={}, + histograms={}, + images={}, + audios={}, + audio_sampling_rate=22050, +): + for k, v in scalars.items(): + writer.add_scalar(k, v, global_step) + for k, v in histograms.items(): + writer.add_histogram(k, v, global_step) + for k, v in images.items(): + writer.add_image(k, v, global_step, dataformats="HWC") + for k, v in audios.items(): + writer.add_audio(k, v, global_step, audio_sampling_rate) + + +def is_resuming(dir_path): + # JP-ExtraバージョンではDURがなくWDがあったり変わるため、Gのみで判断する + g_list = glob.glob(os.path.join(dir_path, "G_*.pth")) + # d_list = glob.glob(os.path.join(dir_path, "D_*.pth")) + # dur_list = glob.glob(os.path.join(dir_path, "DUR_*.pth")) + return len(g_list) > 0 + + +def plot_spectrogram_to_numpy(spectrogram): + global MATPLOTLIB_FLAG + if not MATPLOTLIB_FLAG: + import matplotlib + + matplotlib.use("Agg") + MATPLOTLIB_FLAG = True + mpl_logger = logging.getLogger("matplotlib") + mpl_logger.setLevel(logging.WARNING) + import matplotlib.pylab as plt + import numpy as np + + fig, ax = plt.subplots(figsize=(10, 2)) + im = ax.imshow(spectrogram, aspect="auto", origin="lower", interpolation="none") + plt.colorbar(im, ax=ax) + plt.xlabel("Frames") + plt.ylabel("Channels") + plt.tight_layout() + + fig.canvas.draw() + data = np.fromstring(fig.canvas.tostring_rgb(), dtype=np.uint8, sep="") + data = data.reshape(fig.canvas.get_width_height()[::-1] + (3,)) + plt.close() + return data + + +def plot_alignment_to_numpy(alignment, info=None): + global MATPLOTLIB_FLAG + if not MATPLOTLIB_FLAG: + import matplotlib + + matplotlib.use("Agg") + MATPLOTLIB_FLAG = True + mpl_logger = logging.getLogger("matplotlib") + mpl_logger.setLevel(logging.WARNING) + import matplotlib.pylab as plt + import numpy as np + + fig, ax = plt.subplots(figsize=(6, 4)) + im = ax.imshow( + alignment.transpose(), aspect="auto", origin="lower", interpolation="none" + ) + fig.colorbar(im, ax=ax) + xlabel = "Decoder timestep" + if info is not None: + xlabel += "\n\n" + info + plt.xlabel(xlabel) + plt.ylabel("Encoder timestep") + plt.tight_layout() + + fig.canvas.draw() + data = np.fromstring(fig.canvas.tostring_rgb(), dtype=np.uint8, sep="") + data = data.reshape(fig.canvas.get_width_height()[::-1] + (3,)) + plt.close() + return data + + +def load_wav_to_torch(full_path): + sampling_rate, data = read(full_path) + return torch.FloatTensor(data.astype(np.float32)), sampling_rate + + +def load_filepaths_and_text(filename, split="|"): + with open(filename, encoding="utf-8") as f: + filepaths_and_text = [line.strip().split(split) for line in f] + return filepaths_and_text + + +def get_logger(model_dir, filename="train.log"): + global logger + logger = logging.getLogger(os.path.basename(model_dir)) + logger.setLevel(logging.DEBUG) + + formatter = logging.Formatter("%(asctime)s\t%(name)s\t%(levelname)s\t%(message)s") + if not os.path.exists(model_dir): + os.makedirs(model_dir) + h = logging.FileHandler(os.path.join(model_dir, filename)) + h.setLevel(logging.DEBUG) + h.setFormatter(formatter) + logger.addHandler(h) + return logger + + +def get_steps(model_path): + matches = re.findall(r"\d+", model_path) + return matches[-1] if matches else None + + +def check_git_hash(model_dir): + source_dir = os.path.dirname(os.path.realpath(__file__)) + if not os.path.exists(os.path.join(source_dir, ".git")): + logger.warning( + "{} is not a git repository, therefore hash value comparison will be ignored.".format( + source_dir + ) + ) + return + + cur_hash = subprocess.getoutput("git rev-parse HEAD") + + path = os.path.join(model_dir, "githash") + if os.path.exists(path): + saved_hash = open(path).read() + if saved_hash != cur_hash: + logger.warning( + "git hash values are different. {}(saved) != {}(current)".format( + saved_hash[:8], cur_hash[:8] + ) + ) + else: + open(path, "w").write(cur_hash) diff --git a/style_bert_vits2/models/utils/checkpoints.py b/style_bert_vits2/models/utils/checkpoints.py new file mode 100644 index 0000000..f26f8fc --- /dev/null +++ b/style_bert_vits2/models/utils/checkpoints.py @@ -0,0 +1,194 @@ +import glob +import os +import re +from pathlib import Path +from typing import Any, Optional, Union + +import torch + +from style_bert_vits2.logging import logger + + +def load_checkpoint( + checkpoint_path: Union[str, Path], + model: torch.nn.Module, + optimizer: Optional[torch.optim.Optimizer] = None, + skip_optimizer: bool = False, + for_infer: bool = False +) -> tuple[torch.nn.Module, Optional[torch.optim.Optimizer], float, int]: + """ + 指定されたパスからチェックポイントを読み込み、モデルとオプティマイザーを更新する。 + + Args: + checkpoint_path (Union[str, Path]): チェックポイントファイルのパス + model (torch.nn.Module): 更新するモデル + optimizer (Optional[torch.optim.Optimizer]): 更新するオプティマイザー。None の場合は更新しない + skip_optimizer (bool): オプティマイザーの更新をスキップするかどうかのフラグ + for_infer (bool): 推論用に読み込むかどうかのフラグ + + Returns: + tuple[torch.nn.Module, Optional[torch.optim.Optimizer], float, int]: 更新されたモデルとオプティマイザー、学習率、イテレーション番号 + """ + + assert os.path.isfile(checkpoint_path) + checkpoint_dict = torch.load(checkpoint_path, map_location="cpu") + iteration = checkpoint_dict["iteration"] + learning_rate = checkpoint_dict["learning_rate"] + logger.info( + f"Loading model and optimizer at iteration {iteration} from {checkpoint_path}" + ) + if ( + optimizer is not None + and not skip_optimizer + and checkpoint_dict["optimizer"] is not None + ): + optimizer.load_state_dict(checkpoint_dict["optimizer"]) + elif optimizer is None and not skip_optimizer: + # else: Disable this line if Infer and resume checkpoint,then enable the line upper + new_opt_dict = optimizer.state_dict() # type: ignore + new_opt_dict_params = new_opt_dict["param_groups"][0]["params"] + new_opt_dict["param_groups"] = checkpoint_dict["optimizer"]["param_groups"] + new_opt_dict["param_groups"][0]["params"] = new_opt_dict_params + optimizer.load_state_dict(new_opt_dict) # type: ignore + + saved_state_dict = checkpoint_dict["model"] + if hasattr(model, "module"): + state_dict = model.module.state_dict() + else: + state_dict = model.state_dict() + + new_state_dict = {} + for k, v in state_dict.items(): + try: + # assert "emb_g" not in k + new_state_dict[k] = saved_state_dict[k] + assert saved_state_dict[k].shape == v.shape, ( + saved_state_dict[k].shape, + v.shape, + ) + except: + # For upgrading from the old version + if "ja_bert_proj" in k: + v = torch.zeros_like(v) + logger.warning( + f"Seems you are using the old version of the model, the {k} is automatically set to zero for backward compatibility" + ) + elif "enc_q" in k and for_infer: + continue + else: + logger.error(f"{k} is not in the checkpoint {checkpoint_path}") + + new_state_dict[k] = v + + if hasattr(model, "module"): + model.module.load_state_dict(new_state_dict, strict=False) + else: + model.load_state_dict(new_state_dict, strict=False) + + logger.info("Loaded '{}' (iteration {})".format(checkpoint_path, iteration)) + + return model, optimizer, learning_rate, iteration + + +def save_checkpoint( + model: torch.nn.Module, + optimizer: Union[torch.optim.Optimizer, torch.optim.AdamW], + learning_rate: float, + iteration: int, + checkpoint_path: Union[str, Path], +) -> None: + """ + モデルとオプティマイザーの状態を指定されたパスに保存する。 + + Args: + model (torch.nn.Module): 保存するモデル + optimizer (Union[torch.optim.Optimizer, torch.optim.AdamW]): 保存するオプティマイザー + learning_rate (float): 学習率 + iteration (int): イテレーション数 + checkpoint_path (Union[str, Path]): 保存先のパス + """ + logger.info(f"Saving model and optimizer state at iteration {iteration} to {checkpoint_path}") + if hasattr(model, "module"): + state_dict = model.module.state_dict() + else: + state_dict = model.state_dict() + torch.save( + { + "model": state_dict, + "iteration": iteration, + "optimizer": optimizer.state_dict(), + "learning_rate": learning_rate, + }, + checkpoint_path, + ) + + +def clean_checkpoints(model_dir_path: Union[str, Path] = "logs/44k/", n_ckpts_to_keep: int = 2, sort_by_time: bool = True) -> None: + """ + 指定されたディレクトリから古いチェックポイントを削除して空き容量を確保する + + Args: + model_dir_path (Union[str, Path]): モデルが保存されているディレクトリのパス + n_ckpts_to_keep (int): 保持するチェックポイントの数(G_0.pth と D_0.pth を除く) + sort_by_time (bool): True の場合、時間順に削除。False の場合、名前順に削除 + """ + + ckpts_files = [ + f + for f in os.listdir(model_dir_path) + if os.path.isfile(os.path.join(model_dir_path, f)) + ] + + def name_key(_f: str) -> int: + return int(re.compile("._(\\d+)\\.pth").match(_f).group(1)) # type: ignore + + def time_key(_f: str) -> float: + return os.path.getmtime(os.path.join(model_dir_path, _f)) + + sort_key = time_key if sort_by_time else name_key + + def x_sorted(_x: str) -> list[str]: + return sorted( + [f for f in ckpts_files if f.startswith(_x) and not f.endswith("_0.pth")], + key=sort_key, + ) + + to_del = [ + os.path.join(model_dir_path, fn) + for fn in ( + x_sorted("G_")[:-n_ckpts_to_keep] + + x_sorted("D_")[:-n_ckpts_to_keep] + + x_sorted("WD_")[:-n_ckpts_to_keep] + + x_sorted("DUR_")[:-n_ckpts_to_keep] + ) + ] + + def del_info(fn: str) -> None: + return logger.info(f"Free up space by deleting ckpt {fn}") + + def del_routine(x: str) -> list[Any]: + return [os.remove(x), del_info(x)] + + [del_routine(fn) for fn in to_del] + + +def get_latest_checkpoint_path(model_dir_path: Union[str, Path], regex: str = "G_*.pth") -> str: + """ + 指定されたディレクトリから最新のチェックポイントのパスを取得する + + Args: + model_dir_path (Union[str, Path]): モデルが保存されているディレクトリのパス + regex (str): チェックポイントのファイル名の正規表現 + + Returns: + str: 最新のチェックポイントのパス + """ + + f_list = glob.glob(os.path.join(str(model_dir_path), regex)) + f_list.sort(key=lambda f: int("".join(filter(str.isdigit, f)))) + try: + x = f_list[-1] + except IndexError: + raise ValueError(f"No checkpoint found in {model_dir_path} with regex {regex}") + + return x diff --git a/style_bert_vits2/models/utils/safetensors.py b/style_bert_vits2/models/utils/safetensors.py new file mode 100644 index 0000000..8917c77 --- /dev/null +++ b/style_bert_vits2/models/utils/safetensors.py @@ -0,0 +1,91 @@ +from pathlib import Path +from typing import Any, Optional, Union + +import torch +from safetensors import safe_open +from safetensors.torch import save_file + +from style_bert_vits2.logging import logger + + +def load_safetensors( + checkpoint_path: Union[str, Path], + model: torch.nn.Module, + for_infer: bool = False, +) -> tuple[torch.nn.Module, Optional[int]]: + """ + 指定されたパスから safetensors モデルを読み込み、モデルとイテレーションを返す。 + + Args: + checkpoint_path (Union[str, Path]): モデルのチェックポイントファイルのパス + model (torch.nn.Module): 読み込む対象のモデル + for_infer (bool): 推論用に読み込むかどうかのフラグ + + Returns: + tuple[torch.nn.Module, Optional[int]]: 読み込まれたモデルとイテレーション番号(存在する場合) + """ + + tensors: dict[str, Any] = {} + iteration: Optional[int] = None + with safe_open(str(checkpoint_path), framework="pt", device="cpu") as f: # type: ignore + for key in f.keys(): + if key == "iteration": + iteration = f.get_tensor(key).item() + tensors[key] = f.get_tensor(key) + if hasattr(model, "module"): + result = model.module.load_state_dict(tensors, strict=False) + else: + result = model.load_state_dict(tensors, strict=False) + for key in result.missing_keys: + if key.startswith("enc_q") and for_infer: + continue + logger.warning(f"Missing key: {key}") + for key in result.unexpected_keys: + if key == "iteration": + continue + logger.warning(f"Unexpected key: {key}") + if iteration is None: + logger.info(f"Loaded '{checkpoint_path}'") + else: + logger.info(f"Loaded '{checkpoint_path}' (iteration {iteration})") + + return model, iteration + + +def save_safetensors( + model: torch.nn.Module, + iteration: int, + checkpoint_path: Union[str, Path], + is_half: bool = False, + for_infer: bool = False, +) -> None: + """ + モデルを safetensors 形式で保存する。 + + Args: + model (torch.nn.Module): 保存するモデル + iteration (int): イテレーション番号 + checkpoint_path (Union[str, Path]): 保存先のパス + is_half (bool): モデルを半精度で保存するかどうかのフラグ + for_infer (bool): 推論用に保存するかどうかのフラグ + """ + + if hasattr(model, "module"): + state_dict = model.module.state_dict() + else: + state_dict = model.state_dict() + keys = [] + for k in state_dict: + if "enc_q" in k and for_infer: + continue # noqa: E701 + keys.append(k) + + new_dict = ( + {k: state_dict[k].half() for k in keys} + if is_half + else {k: state_dict[k] for k in keys} + ) + new_dict["iteration"] = torch.LongTensor([iteration]) + logger.info(f"Saved safetensors to {checkpoint_path}") + + save_file(new_dict, checkpoint_path) diff --git a/train_ms.py b/train_ms.py index b2cb02b..3d50c03 100644 --- a/train_ms.py +++ b/train_ms.py @@ -248,10 +248,7 @@ def run(): drop_last=False, collate_fn=collate_fn, ) - if ( - "use_noise_scaled_mas" in hps.model.keys() - and hps.model.use_noise_scaled_mas is True - ): + if hps.model.use_noise_scaled_mas is True: logger.info("Using noise scaled MAS for VITS2") mas_noise_scale_initial = 0.01 noise_scale_delta = 2e-6 @@ -259,10 +256,7 @@ def run(): logger.info("Using normal MAS for VITS1") mas_noise_scale_initial = 0.0 noise_scale_delta = 0.0 - if ( - "use_duration_discriminator" in hps.model.keys() - and hps.model.use_duration_discriminator is True - ): + if hps.model.use_duration_discriminator is True: logger.info("Using duration discriminator for VITS2") net_dur_disc = DurationDiscriminator( hps.model.hidden_channels, @@ -271,10 +265,7 @@ def run(): 0.1, gin_channels=hps.model.gin_channels if hps.data.n_speakers != 0 else 0, ).cuda(local_rank) - if ( - "use_spk_conditioned_encoder" in hps.model.keys() - and hps.model.use_spk_conditioned_encoder is True - ): + if hps.model.use_spk_conditioned_encoder is True: if hps.data.n_speakers == 0: raise ValueError( "n_speakers must be > 0 when using spk conditioned encoder to train multi-speaker model" @@ -370,31 +361,25 @@ def run(): if utils.is_resuming(model_dir): if net_dur_disc is not None: - _, _, dur_resume_lr, epoch_str = utils.load_checkpoint( - utils.latest_checkpoint_path(model_dir, "DUR_*.pth"), + _, _, dur_resume_lr, epoch_str = utils.checkpoints.load_checkpoint( + utils.checkpoints.get_latest_checkpoint_path(model_dir, "DUR_*.pth"), net_dur_disc, optim_dur_disc, - skip_optimizer=( - hps.train.skip_optimizer if "skip_optimizer" in hps.train else True - ), + skip_optimizer=hps.train.skip_optimizer, ) if not optim_dur_disc.param_groups[0].get("initial_lr"): optim_dur_disc.param_groups[0]["initial_lr"] = dur_resume_lr - _, optim_g, g_resume_lr, epoch_str = utils.load_checkpoint( - utils.latest_checkpoint_path(model_dir, "G_*.pth"), + _, optim_g, g_resume_lr, epoch_str = utils.checkpoints.load_checkpoint( + utils.checkpoints.get_latest_checkpoint_path(model_dir, "G_*.pth"), net_g, optim_g, - skip_optimizer=( - hps.train.skip_optimizer if "skip_optimizer" in hps.train else True - ), + skip_optimizer=hps.train.skip_optimizer, ) - _, optim_d, d_resume_lr, epoch_str = utils.load_checkpoint( - utils.latest_checkpoint_path(model_dir, "D_*.pth"), + _, optim_d, d_resume_lr, epoch_str = utils.checkpoints.load_checkpoint( + utils.checkpoints.get_latest_checkpoint_path(model_dir, "D_*.pth"), net_d, optim_d, - skip_optimizer=( - hps.train.skip_optimizer if "skip_optimizer" in hps.train else True - ), + skip_optimizer=hps.train.skip_optimizer, ) if not optim_g.param_groups[0].get("initial_lr"): optim_g.param_groups[0]["initial_lr"] = g_resume_lr @@ -404,21 +389,21 @@ def run(): epoch_str = max(epoch_str, 1) # global_step = (epoch_str - 1) * len(train_loader) global_step = int( - utils.get_steps(utils.latest_checkpoint_path(model_dir, "G_*.pth")) + utils.get_steps(utils.checkpoints.get_latest_checkpoint_path(model_dir, "G_*.pth")) ) logger.info( f"******************Found the model. Current epoch is {epoch_str}, gloabl step is {global_step}*********************" ) else: try: - _ = utils.load_safetensors( + _ = utils.safetensors.load_safetensors( os.path.join(model_dir, "G_0.safetensors"), net_g ) - _ = utils.load_safetensors( + _ = utils.safetensors.load_safetensors( os.path.join(model_dir, "D_0.safetensors"), net_d ) if net_dur_disc is not None: - _ = utils.load_safetensors( + _ = utils.safetensors.load_safetensors( os.path.join(model_dir, "DUR_0.safetensors"), net_dur_disc ) logger.info("Loaded the pretrained models.") @@ -511,14 +496,16 @@ def run(): if epoch == hps.train.epochs: # Save the final models - utils.save_checkpoint( + assert optim_g is not None + utils.checkpoints.save_checkpoint( net_g, optim_g, hps.train.learning_rate, epoch, os.path.join(model_dir, "G_{}.pth".format(global_step)), ) - utils.save_checkpoint( + assert optim_d is not None + utils.checkpoints.save_checkpoint( net_d, optim_d, hps.train.learning_rate, @@ -526,14 +513,15 @@ def run(): os.path.join(model_dir, "D_{}.pth".format(global_step)), ) if net_dur_disc is not None: - utils.save_checkpoint( + assert optim_dur_disc is not None + utils.checkpoints.save_checkpoint( net_dur_disc, optim_dur_disc, hps.train.learning_rate, epoch, os.path.join(model_dir, "DUR_{}.pth".format(global_step)), ) - utils.save_safetensors( + utils.safetensors.save_safetensors( net_g, epoch, os.path.join( @@ -804,14 +792,15 @@ def train_and_evaluate( ): if not hps.speedup: evaluate(hps, net_g, eval_loader, writer_eval) - utils.save_checkpoint( + assert hps.model_dir is not None + utils.checkpoints.save_checkpoint( net_g, optim_g, hps.train.learning_rate, epoch, os.path.join(hps.model_dir, "G_{}.pth".format(global_step)), ) - utils.save_checkpoint( + utils.checkpoints.save_checkpoint( net_d, optim_d, hps.train.learning_rate, @@ -819,7 +808,7 @@ def train_and_evaluate( os.path.join(hps.model_dir, "D_{}.pth".format(global_step)), ) if net_dur_disc is not None: - utils.save_checkpoint( + utils.checkpoints.save_checkpoint( net_dur_disc, optim_dur_disc, hps.train.learning_rate, @@ -828,13 +817,13 @@ def train_and_evaluate( ) keep_ckpts = config.train_ms_config.keep_ckpts if keep_ckpts > 0: - utils.clean_checkpoints( - path_to_models=hps.model_dir, + utils.checkpoints.clean_checkpoints( + model_dir_path=hps.model_dir, n_ckpts_to_keep=keep_ckpts, sort_by_time=True, ) # Save safetensors (for inference) to `model_assets/{model_name}` - utils.save_safetensors( + utils.safetensors.save_safetensors( net_g, epoch, os.path.join( diff --git a/train_ms_jp_extra.py b/train_ms_jp_extra.py index a464d29..5ed68a1 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -247,10 +247,7 @@ def run(): drop_last=False, collate_fn=collate_fn, ) - if ( - "use_noise_scaled_mas" in hps.model.keys() - and hps.model.use_noise_scaled_mas is True - ): + if hps.model.use_noise_scaled_mas is True: logger.info("Using noise scaled MAS for VITS2") mas_noise_scale_initial = 0.01 noise_scale_delta = 2e-6 @@ -258,10 +255,7 @@ def run(): logger.info("Using normal MAS for VITS1") mas_noise_scale_initial = 0.0 noise_scale_delta = 0.0 - if ( - "use_duration_discriminator" in hps.model.keys() - and hps.model.use_duration_discriminator is True - ): + if hps.model.use_duration_discriminator is True: logger.info("Using duration discriminator for VITS2") net_dur_disc = DurationDiscriminator( hps.model.hidden_channels, @@ -272,19 +266,13 @@ def run(): ).cuda(local_rank) else: net_dur_disc = None - if ( - "use_wavlm_discriminator" in hps.model.keys() - and hps.model.use_wavlm_discriminator is True - ): + if hps.model.use_wavlm_discriminator is True: net_wd = WavLMDiscriminator( hps.model.slm.hidden, hps.model.slm.nlayers, hps.model.slm.initial_channel ).cuda(local_rank) else: net_wd = None - if ( - "use_spk_conditioned_encoder" in hps.model.keys() - and hps.model.use_spk_conditioned_encoder is True - ): + if hps.model.use_spk_conditioned_encoder is True: if hps.data.n_speakers == 0: raise ValueError( "n_speakers must be > 0 when using spk conditioned encoder to train multi-speaker model" @@ -394,15 +382,11 @@ def run(): if utils.is_resuming(model_dir): if net_dur_disc is not None: try: - _, _, dur_resume_lr, epoch_str = utils.load_checkpoint( - utils.latest_checkpoint_path(model_dir, "DUR_*.pth"), + _, _, dur_resume_lr, epoch_str = utils.checkpoints.load_checkpoint( + utils.checkpoints.get_latest_checkpoint_path(model_dir, "DUR_*.pth"), net_dur_disc, optim_dur_disc, - skip_optimizer=( - hps.train.skip_optimizer - if "skip_optimizer" in hps.train - else True - ), + skip_optimizer=hps.train.skip_optimizer, ) if not optim_dur_disc.param_groups[0].get("initial_lr"): optim_dur_disc.param_groups[0]["initial_lr"] = dur_resume_lr @@ -412,15 +396,11 @@ def run(): print("Initialize dur_disc") if net_wd is not None: try: - _, optim_wd, wd_resume_lr, epoch_str = utils.load_checkpoint( - utils.latest_checkpoint_path(model_dir, "WD_*.pth"), + _, optim_wd, wd_resume_lr, epoch_str = utils.checkpoints.load_checkpoint( + utils.checkpoints.get_latest_checkpoint_path(model_dir, "WD_*.pth"), net_wd, optim_wd, - skip_optimizer=( - hps.train.skip_optimizer - if "skip_optimizer" in hps.train - else True - ), + skip_optimizer=hps.train.skip_optimizer, ) if not optim_wd.param_groups[0].get("initial_lr"): optim_wd.param_groups[0]["initial_lr"] = wd_resume_lr @@ -430,21 +410,17 @@ def run(): logger.info("Initialize wavlm") try: - _, optim_g, g_resume_lr, epoch_str = utils.load_checkpoint( - utils.latest_checkpoint_path(model_dir, "G_*.pth"), + _, optim_g, g_resume_lr, epoch_str = utils.checkpoints.load_checkpoint( + utils.checkpoints.get_latest_checkpoint_path(model_dir, "G_*.pth"), net_g, optim_g, - skip_optimizer=( - hps.train.skip_optimizer if "skip_optimizer" in hps.train else True - ), + skip_optimizer=hps.train.skip_optimizer, ) - _, optim_d, d_resume_lr, epoch_str = utils.load_checkpoint( - utils.latest_checkpoint_path(model_dir, "D_*.pth"), + _, optim_d, d_resume_lr, epoch_str = utils.checkpoints.load_checkpoint( + utils.checkpoints.get_latest_checkpoint_path(model_dir, "D_*.pth"), net_d, optim_d, - skip_optimizer=( - hps.train.skip_optimizer if "skip_optimizer" in hps.train else True - ), + skip_optimizer=hps.train.skip_optimizer, ) if not optim_g.param_groups[0].get("initial_lr"): optim_g.param_groups[0]["initial_lr"] = g_resume_lr @@ -454,7 +430,7 @@ def run(): epoch_str = max(epoch_str, 1) # global_step = (epoch_str - 1) * len(train_loader) global_step = int( - utils.get_steps(utils.latest_checkpoint_path(model_dir, "G_*.pth")) + utils.get_steps(utils.checkpoints.get_latest_checkpoint_path(model_dir, "G_*.pth")) ) logger.info( f"******************Found the model. Current epoch is {epoch_str}, gloabl step is {global_step}*********************" @@ -468,18 +444,18 @@ def run(): global_step = 0 else: try: - _ = utils.load_safetensors( + _ = utils.safetensors.load_safetensors( os.path.join(model_dir, "G_0.safetensors"), net_g ) - _ = utils.load_safetensors( + _ = utils.safetensors.load_safetensors( os.path.join(model_dir, "D_0.safetensors"), net_d ) if net_dur_disc is not None: - _ = utils.load_safetensors( + _ = utils.safetensors.load_safetensors( os.path.join(model_dir, "DUR_0.safetensors"), net_dur_disc ) if net_wd is not None: - _ = utils.load_safetensors( + _ = utils.safetensors.load_safetensors( os.path.join(model_dir, "WD_0.safetensors"), net_wd ) logger.info("Loaded the pretrained models.") @@ -586,14 +562,16 @@ def run(): scheduler_wd.step() if epoch == hps.train.epochs: # Save the final models - utils.save_checkpoint( + assert optim_g is not None + utils.checkpoints.save_checkpoint( net_g, optim_g, hps.train.learning_rate, epoch, os.path.join(model_dir, "G_{}.pth".format(global_step)), ) - utils.save_checkpoint( + assert optim_d is not None + utils.checkpoints.save_checkpoint( net_d, optim_d, hps.train.learning_rate, @@ -601,7 +579,8 @@ def run(): os.path.join(model_dir, "D_{}.pth".format(global_step)), ) if net_dur_disc is not None: - utils.save_checkpoint( + assert optim_dur_disc is not None + utils.checkpoints.save_checkpoint( net_dur_disc, optim_dur_disc, hps.train.learning_rate, @@ -609,14 +588,15 @@ def run(): os.path.join(model_dir, "DUR_{}.pth".format(global_step)), ) if net_wd is not None: - utils.save_checkpoint( + assert optim_wd is not None + utils.checkpoints.save_checkpoint( net_wd, optim_wd, hps.train.learning_rate, epoch, os.path.join(model_dir, "WD_{}.pth".format(global_step)), ) - utils.save_safetensors( + utils.safetensors.save_safetensors( net_g, epoch, os.path.join( @@ -949,14 +929,14 @@ def train_and_evaluate( ): if not hps.speedup: evaluate(hps, net_g, eval_loader, writer_eval) - utils.save_checkpoint( + utils.checkpoints.save_checkpoint( net_g, optim_g, hps.train.learning_rate, epoch, os.path.join(hps.model_dir, "G_{}.pth".format(global_step)), ) - utils.save_checkpoint( + utils.checkpoints.save_checkpoint( net_d, optim_d, hps.train.learning_rate, @@ -964,7 +944,7 @@ def train_and_evaluate( os.path.join(hps.model_dir, "D_{}.pth".format(global_step)), ) if net_dur_disc is not None: - utils.save_checkpoint( + utils.checkpoints.save_checkpoint( net_dur_disc, optim_dur_disc, hps.train.learning_rate, @@ -972,7 +952,7 @@ def train_and_evaluate( os.path.join(hps.model_dir, "DUR_{}.pth".format(global_step)), ) if net_wd is not None: - utils.save_checkpoint( + utils.checkpoints.save_checkpoint( net_wd, optim_wd, hps.train.learning_rate, @@ -981,13 +961,13 @@ def train_and_evaluate( ) keep_ckpts = config.train_ms_config.keep_ckpts if keep_ckpts > 0: - utils.clean_checkpoints( - path_to_models=hps.model_dir, + utils.checkpoints.clean_checkpoints( + model_dir_path=hps.model_dir, n_ckpts_to_keep=keep_ckpts, sort_by_time=True, ) # Save safetensors (for inference) to `model_assets/{model_name}` - utils.save_safetensors( + utils.safetensors.save_safetensors( net_g, epoch, os.path.join( From 61e2a1deae543e2403ffa56e97570e11fad4a609 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sat, 9 Mar 2024 17:26:42 +0000 Subject: [PATCH 50/64] Refactor: add type hints to style_bert_vits2.models.utils --- style_bert_vits2/models/utils/__init__.py | 158 +++++++++++++++++----- style_bert_vits2/nlp/__init__.py | 2 +- 2 files changed, 127 insertions(+), 33 deletions(-) diff --git a/style_bert_vits2/models/utils/__init__.py b/style_bert_vits2/models/utils/__init__.py index f488e09..51e19d9 100644 --- a/style_bert_vits2/models/utils/__init__.py +++ b/style_bert_vits2/models/utils/__init__.py @@ -3,28 +3,44 @@ import logging import os import re import subprocess +from pathlib import Path +from typing import Any, Optional, Union import numpy as np import torch +from numpy.typing import NDArray from scipy.io.wavfile import read +from torch.utils.tensorboard import SummaryWriter from style_bert_vits2.logging import logger from style_bert_vits2.models.utils import checkpoints # type: ignore from style_bert_vits2.models.utils import safetensors # type: ignore -MATPLOTLIB_FLAG = False +__is_matplotlib_imported = False def summarize( - writer, - global_step, - scalars={}, - histograms={}, - images={}, - audios={}, - audio_sampling_rate=22050, -): + writer: SummaryWriter, + global_step: int, + scalars: dict[str, float] = {}, + histograms: dict[str, Any] = {}, + images: dict[str, Any] = {}, + audios: dict[str, Any] = {}, + audio_sampling_rate: int = 22050, +) -> None: + """ + 指定されたデータを TensorBoard にまとめて追加する + + Args: + writer (SummaryWriter): TensorBoard への書き込みを行うオブジェクト + global_step (int): グローバルステップ数 + scalars (dict[str, float]): スカラー値の辞書 + histograms (dict[str, Any]): ヒストグラムの辞書 + images (dict[str, Any]): 画像データの辞書 + audios (dict[str, Any]): 音声データの辞書 + audio_sampling_rate (int): 音声データのサンプリングレート + """ for k, v in scalars.items(): writer.add_scalar(k, v, global_step) for k, v in histograms.items(): @@ -35,7 +51,16 @@ def summarize( writer.add_audio(k, v, global_step, audio_sampling_rate) -def is_resuming(dir_path): +def is_resuming(dir_path: Union[str, Path]) -> bool: + """ + 指定されたディレクトリパスに再開可能なモデルが存在するかどうかを返す + + Args: + dir_path: チェックするディレクトリのパス + + Returns: + bool: 再開可能なモデルが存在するかどうか + """ # JP-ExtraバージョンではDURがなくWDがあったり変わるため、Gのみで判断する g_list = glob.glob(os.path.join(dir_path, "G_*.pth")) # d_list = glob.glob(os.path.join(dir_path, "D_*.pth")) @@ -43,13 +68,23 @@ def is_resuming(dir_path): return len(g_list) > 0 -def plot_spectrogram_to_numpy(spectrogram): - global MATPLOTLIB_FLAG - if not MATPLOTLIB_FLAG: +def plot_spectrogram_to_numpy(spectrogram: NDArray[Any]) -> NDArray[Any]: + """ + 指定されたスペクトログラムを画像データに変換する + + Args: + spectrogram (NDArray[Any]): スペクトログラム + + Returns: + NDArray[Any]: 画像データ + """ + + global __is_matplotlib_imported + if not __is_matplotlib_imported: import matplotlib matplotlib.use("Agg") - MATPLOTLIB_FLAG = True + __is_matplotlib_imported = True mpl_logger = logging.getLogger("matplotlib") mpl_logger.setLevel(logging.WARNING) import matplotlib.pylab as plt @@ -63,23 +98,33 @@ def plot_spectrogram_to_numpy(spectrogram): plt.tight_layout() fig.canvas.draw() - data = np.fromstring(fig.canvas.tostring_rgb(), dtype=np.uint8, sep="") + data = np.fromstring(fig.canvas.tostring_rgb(), dtype=np.uint8, sep="") # type: ignore data = data.reshape(fig.canvas.get_width_height()[::-1] + (3,)) plt.close() return data -def plot_alignment_to_numpy(alignment, info=None): - global MATPLOTLIB_FLAG - if not MATPLOTLIB_FLAG: +def plot_alignment_to_numpy(alignment: NDArray[Any], info: Optional[str] = None) -> NDArray[Any]: + """ + 指定されたアライメントを画像データに変換する + + Args: + alignment (NDArray[Any]): アライメント + info (Optional[str]): 画像に追加する情報 + + Returns: + NDArray[Any]: 画像データ + """ + + global __is_matplotlib_imported + if not __is_matplotlib_imported: import matplotlib matplotlib.use("Agg") - MATPLOTLIB_FLAG = True + __is_matplotlib_imported = True mpl_logger = logging.getLogger("matplotlib") mpl_logger.setLevel(logging.WARNING) import matplotlib.pylab as plt - import numpy as np fig, ax = plt.subplots(figsize=(6, 4)) im = ax.imshow( @@ -94,44 +139,93 @@ def plot_alignment_to_numpy(alignment, info=None): plt.tight_layout() fig.canvas.draw() - data = np.fromstring(fig.canvas.tostring_rgb(), dtype=np.uint8, sep="") + data = np.fromstring(fig.canvas.tostring_rgb(), dtype=np.uint8, sep="") # type: ignore data = data.reshape(fig.canvas.get_width_height()[::-1] + (3,)) plt.close() return data -def load_wav_to_torch(full_path): +def load_wav_to_torch(full_path: Union[str, Path]) -> tuple[torch.FloatTensor, int]: + """ + 指定された音声ファイルを読み込み、PyTorch のテンソルに変換して返す + + Args: + full_path (Union[str, Path]): 音声ファイルのパス + + Returns: + tuple[torch.FloatTensor, int]: 音声データのテンソルとサンプリングレート + """ + sampling_rate, data = read(full_path) return torch.FloatTensor(data.astype(np.float32)), sampling_rate -def load_filepaths_and_text(filename, split="|"): +def load_filepaths_and_text(filename: Union[str, Path], split: str = "|") -> list[list[str]]: + """ + 指定されたファイルからファイルパスとテキストを読み込む + + Args: + filename (Union[str, Path]): ファイルのパス + split (str): ファイルの区切り文字 (デフォルト: "|") + + Returns: + list[list[str]]: ファイルパスとテキストのリスト + """ + with open(filename, encoding="utf-8") as f: filepaths_and_text = [line.strip().split(split) for line in f] return filepaths_and_text -def get_logger(model_dir, filename="train.log"): +def get_logger(model_dir_path: Union[str, Path], filename: str = "train.log") -> logging.Logger: + """ + ロガーを取得する + + Args: + model_dir_path (Union[str, Path]): ログを保存するディレクトリのパス + filename (str): ログファイルの名前 (デフォルト: "train.log") + + Returns: + logging.Logger: ロガー + """ + global logger - logger = logging.getLogger(os.path.basename(model_dir)) + logger = logging.getLogger(os.path.basename(model_dir_path)) logger.setLevel(logging.DEBUG) formatter = logging.Formatter("%(asctime)s\t%(name)s\t%(levelname)s\t%(message)s") - if not os.path.exists(model_dir): - os.makedirs(model_dir) - h = logging.FileHandler(os.path.join(model_dir, filename)) + if not os.path.exists(model_dir_path): + os.makedirs(model_dir_path) + h = logging.FileHandler(os.path.join(model_dir_path, filename)) h.setLevel(logging.DEBUG) h.setFormatter(formatter) logger.addHandler(h) return logger -def get_steps(model_path): - matches = re.findall(r"\d+", model_path) +def get_steps(model_path: Union[str, Path]) -> Optional[int]: + """ + モデルのパスからイテレーション番号を取得する + + Args: + model_path (Union[str, Path]): モデルのパス + + Returns: + Optional[int]: イテレーション番号 + """ + + matches = re.findall(r"\d+", model_path) # type: ignore return matches[-1] if matches else None -def check_git_hash(model_dir): +def check_git_hash(model_dir_path: Union[str, Path]) -> None: + """ + モデルのディレクトリに .git ディレクトリが存在する場合、ハッシュ値を比較する + + Args: + model_dir_path (Union[str, Path]): モデルのディレクトリのパス + """ + source_dir = os.path.dirname(os.path.realpath(__file__)) if not os.path.exists(os.path.join(source_dir, ".git")): logger.warning( @@ -143,7 +237,7 @@ def check_git_hash(model_dir): cur_hash = subprocess.getoutput("git rev-parse HEAD") - path = os.path.join(model_dir, "githash") + path = os.path.join(model_dir_path, "githash") if os.path.exists(path): saved_hash = open(path).read() if saved_hash != cur_hash: diff --git a/style_bert_vits2/nlp/__init__.py b/style_bert_vits2/nlp/__init__.py index afc6cb0..683d6d4 100644 --- a/style_bert_vits2/nlp/__init__.py +++ b/style_bert_vits2/nlp/__init__.py @@ -8,7 +8,7 @@ from style_bert_vits2.nlp.symbols import ( ) # __init__.py は配下のモジュールをインポートした時点で実行される -# Pytorch のインポートは重いので、型チェック時以外はインポートしない +# PyTorch のインポートは重いので、型チェック時以外はインポートしない if TYPE_CHECKING: import torch From 96d22102f3f881964c047bc99507369bde62ac83 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sat, 9 Mar 2024 17:45:56 +0000 Subject: [PATCH 51/64] Refactor: separate adjust_voice() function from tts_model.py --- style_bert_vits2/tts_model.py | 171 +++++++++++++++------------------- style_bert_vits2/voice.py | 46 +++++++++ webui/inference.py | 6 +- webui/merge.py | 6 +- 4 files changed, 127 insertions(+), 102 deletions(-) create mode 100644 style_bert_vits2/voice.py diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index 86c8425..a3a83c3 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -25,54 +25,23 @@ from style_bert_vits2.models.infer import get_net_g, infer from style_bert_vits2.models.models import SynthesizerTrn from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra from style_bert_vits2.logging import logger - - -def adjust_voice( - fs: int, - wave: NDArray[Any], - pitch_scale: float, - intonation_scale: float, -) -> tuple[int, NDArray[Any]]: - - if pitch_scale == 1.0 and intonation_scale == 1.0: - # 初期値の場合は、音質劣化を避けるためにそのまま返す - return fs, wave - - try: - import pyworld - except ImportError: - raise ImportError( - "pyworld is not installed. Please install it by `pip install pyworld`" - ) - - # pyworld で f0 を加工して合成 - # pyworld よりもよいのがあるかもしれないが…… - ## pyworld は Cython で書かれているが、スタブファイルがないため型補完が全く効かない… - - wave = wave.astype(np.double) - - # 質が高そうだしとりあえずharvestにしておく - f0, t = pyworld.harvest(wave, fs) # type: ignore - - sp = pyworld.cheaptrick(wave, f0, t, fs) # type: ignore - ap = pyworld.d4c(wave, f0, t, fs) # type: ignore - - non_zero_f0 = [f for f in f0 if f != 0] - f0_mean = sum(non_zero_f0) / len(non_zero_f0) - - for i, f in enumerate(f0): - if f == 0: - continue - f0[i] = pitch_scale * f0_mean + intonation_scale * (f - f0_mean) - - wave = pyworld.synthesize(f0, sp, ap, fs) # type: ignore - return fs, wave +from style_bert_vits2.voice import adjust_voice class Model: + """ + Style-Bert-Vits2 の音声合成モデルを操作するためのクラス + モデル/ハイパーパラメータ/スタイルベクトルのパスとデバイスを指定して初期化し、model.infer() メソッドを呼び出すと音声合成を行える + """ + + def __init__( - self, model_path: Path, config_path: Path, style_vec_path: Path, device: str - ): + self, + model_path: Path, + config_path: Path, + style_vec_path: Path, + device: str, + ) -> None: self.model_path: Path = model_path self.config_path: Path = config_path self.style_vec_path: Path = style_vec_path @@ -99,7 +68,8 @@ class Model: self.net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None - def load_net_g(self): + + def load_net_g(self) -> None: self.net_g = get_net_g( model_path=str(self.model_path), version=self.hps.version, @@ -107,15 +77,15 @@ class Model: hps=self.hps, ) + def get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]: mean = self.style_vectors[0] style_vec = self.style_vectors[style_id] style_vec = mean + (style_vec - mean) * weight return style_vec - def get_style_vector_from_audio( - self, audio_path: str, weight: float = 1.0 - ) -> NDArray[Any]: + + def get_style_vector_from_audio(self, audio_path: str, weight: float = 1.0) -> NDArray[Any]: from style_gen import get_style_vector xvec = get_style_vector(audio_path) @@ -123,6 +93,7 @@ class Model: xvec = mean + (xvec - mean) * weight return xvec + def infer( self, text: str, @@ -167,20 +138,20 @@ class Model: if not line_split: with torch.no_grad(): audio = infer( - text=text, - sdp_ratio=sdp_ratio, - noise_scale=noise, - noise_scale_w=noisew, - length_scale=length, - sid=sid, - language=language, - hps=self.hps, - net_g=self.net_g, - device=self.device, - assist_text=assist_text, - assist_text_weight=assist_text_weight, - style_vec=style_vector, - given_tone=given_tone, + text = text, + sdp_ratio = sdp_ratio, + noise_scale = noise, + noise_scale_w = noisew, + length_scale = length, + sid = sid, + language = language, + hps = self.hps, + net_g = self.net_g, + device = self.device, + assist_text = assist_text, + assist_text_weight = assist_text_weight, + style_vec = style_vector, + given_tone = given_tone, ) else: texts = text.split("\n") @@ -190,19 +161,19 @@ class Model: for i, t in enumerate(texts): audios.append( infer( - text=t, - sdp_ratio=sdp_ratio, - noise_scale=noise, - noise_scale_w=noisew, - length_scale=length, - sid=sid, - language=language, - hps=self.hps, - net_g=self.net_g, - device=self.device, - assist_text=assist_text, - assist_text_weight=assist_text_weight, - style_vec=style_vector, + text = t, + sdp_ratio = sdp_ratio, + noise_scale = noise, + noise_scale_w = noisew, + length_scale = length, + sid = sid, + language = language, + hps = self.hps, + net_g = self.net_g, + device = self.device, + assist_text = assist_text, + assist_text_weight = assist_text_weight, + style_vec = style_vector, ) ) if i != len(texts) - 1: @@ -211,10 +182,10 @@ class Model: logger.info("Audio data generated successfully") if not (pitch_scale == 1.0 and intonation_scale == 1.0): _, audio = adjust_voice( - fs=self.hps.data.sampling_rate, - wave=audio, - pitch_scale=pitch_scale, - intonation_scale=intonation_scale, + fs = self.hps.data.sampling_rate, + wave = audio, + pitch_scale = pitch_scale, + intonation_scale = intonation_scale, ) with warnings.catch_warnings(): warnings.simplefilter("ignore") @@ -223,8 +194,13 @@ class Model: class ModelHolder: - def __init__(self, root_dir: Path, device: str): - self.root_dir: Path = root_dir + """ + Style-Bert-Vits2 の音声合成モデルを管理するためのクラス + """ + + + def __init__(self, model_root_dir: Path, device: str) -> None: + self.root_dir: Path = model_root_dir self.device: str = device self.model_files_dict: dict[str, list[Path]] = {} self.current_model: Optional[Model] = None @@ -233,7 +209,8 @@ class ModelHolder: self.models_info: list[dict[str, Union[str, list[str]]]] = [] self.refresh() - def refresh(self): + + def refresh(self) -> None: self.model_files_dict = {} self.model_names = [] self.current_model = None @@ -269,7 +246,8 @@ class ModelHolder: "speakers": speakers, }) - def load_model(self, model_name: str, model_path_str: str): + + def load_model(self, model_name: str, model_path_str: str) -> Model: model_path = Path(model_path_str) if model_name not in self.model_files_dict: raise ValueError(f"Model `{model_name}` is not found") @@ -277,16 +255,15 @@ class ModelHolder: raise ValueError(f"Model file `{model_path}` is not found") if self.current_model is None or self.current_model.model_path != model_path: self.current_model = Model( - model_path=model_path, - config_path=self.root_dir / model_name / "config.json", - style_vec_path=self.root_dir / model_name / "style_vectors.npy", - device=self.device, + model_path = model_path, + config_path = self.root_dir / model_name / "config.json", + style_vec_path = self.root_dir / model_name / "style_vectors.npy", + device = self.device, ) return self.current_model - def load_model_gr( - self, model_name: str, model_path_str: str - ) -> tuple[gr.Dropdown, gr.Button, gr.Dropdown]: + + def load_model_for_gradio(self, model_name: str, model_path_str: str) -> tuple[gr.Dropdown, gr.Button, gr.Dropdown]: model_path = Path(model_path_str) if model_name not in self.model_files_dict: raise ValueError(f"Model `{model_name}` is not found") @@ -305,10 +282,10 @@ class ModelHolder: gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore ) self.current_model = Model( - model_path=model_path, - config_path=self.root_dir / model_name / "config.json", - style_vec_path=self.root_dir / model_name / "style_vectors.npy", - device=self.device, + model_path = model_path, + config_path = self.root_dir / model_name / "config.json", + style_vec_path = self.root_dir / model_name / "style_vectors.npy", + device = self.device, ) speakers = list(self.current_model.spk2id.keys()) styles = list(self.current_model.style2id.keys()) @@ -318,11 +295,13 @@ class ModelHolder: gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore ) - def update_model_files_gr(self, model_name: str) -> gr.Dropdown: + + def update_model_files_for_gradio(self, model_name: str) -> gr.Dropdown: model_files = self.model_files_dict[model_name] return gr.Dropdown(choices=model_files, value=model_files[0]) # type: ignore - def update_model_names_gr(self) -> tuple[gr.Dropdown, gr.Dropdown, gr.Button]: + + def update_model_names_for_gradio(self) -> tuple[gr.Dropdown, gr.Dropdown, gr.Button]: self.refresh() initial_model_name = self.model_names[0] initial_model_files = self.model_files_dict[initial_model_name] diff --git a/style_bert_vits2/voice.py b/style_bert_vits2/voice.py new file mode 100644 index 0000000..f1cfcac --- /dev/null +++ b/style_bert_vits2/voice.py @@ -0,0 +1,46 @@ +from typing import Any + +import numpy as np +from numpy.typing import NDArray + + +def adjust_voice( + fs: int, + wave: NDArray[Any], + pitch_scale: float = 1.0, + intonation_scale: float = 1.0, +) -> tuple[int, NDArray[Any]]: + + if pitch_scale == 1.0 and intonation_scale == 1.0: + # 初期値の場合は、音質劣化を避けるためにそのまま返す + return fs, wave + + try: + import pyworld + except ImportError: + raise ImportError( + "pyworld is not installed. Please install it by `pip install pyworld`" + ) + + # pyworld で f0 を加工して合成 + # pyworld よりもよいのがあるかもしれないが…… + ## pyworld は Cython で書かれているが、スタブファイルがないため型補完が全く効かない… + + wave = wave.astype(np.double) + + # 質が高そうだしとりあえずharvestにしておく + f0, t = pyworld.harvest(wave, fs) # type: ignore + + sp = pyworld.cheaptrick(wave, f0, t, fs) # type: ignore + ap = pyworld.d4c(wave, f0, t, fs) # type: ignore + + non_zero_f0 = [f for f in f0 if f != 0] + f0_mean = sum(non_zero_f0) / len(non_zero_f0) + + for i, f in enumerate(f0): + if f == 0: + continue + f0[i] = pitch_scale * f0_mean + intonation_scale * (f - f0_mean) + + wave = pyworld.synthesize(f0, sp, ap, fs) # type: ignore + return fs, wave diff --git a/webui/inference.py b/webui/inference.py index 536b38b..9c2bf63 100644 --- a/webui/inference.py +++ b/webui/inference.py @@ -446,7 +446,7 @@ def create_inference_app(model_holder: ModelHolder) -> gr.Blocks: ) model_name.change( - model_holder.update_model_files_gr, + model_holder.update_model_files_for_gradio, inputs=[model_name], outputs=[model_path], ) @@ -454,12 +454,12 @@ def create_inference_app(model_holder: ModelHolder) -> gr.Blocks: model_path.change(make_non_interactive, outputs=[tts_button]) refresh_button.click( - model_holder.update_model_names_gr, + model_holder.update_model_names_for_gradio, outputs=[model_name, model_path, tts_button], ) load_button.click( - model_holder.load_model_gr, + model_holder.load_model_for_gradio, inputs=[model_name, model_path], outputs=[style, tts_button, speaker], ) diff --git a/webui/merge.py b/webui/merge.py index e9dd1f0..85806f4 100644 --- a/webui/merge.py +++ b/webui/merge.py @@ -255,7 +255,7 @@ def simple_tts(model_name, text, style=DEFAULT_STYLE, style_weight=1.0): def update_two_model_names_dropdown(model_holder: ModelHolder): - new_names, new_files, _ = model_holder.update_model_names_gr() + new_names, new_files, _ = model_holder.update_model_names_for_gradio() return new_names, new_files, new_names, new_files @@ -444,12 +444,12 @@ def create_merge_app(model_holder: ModelHolder) -> gr.Blocks: audio_output = gr.Audio(label="結果") model_name_a.change( - model_holder.update_model_files_gr, + model_holder.update_model_files_for_gradio, inputs=[model_name_a], outputs=[model_path_a], ) model_name_b.change( - model_holder.update_model_files_gr, + model_holder.update_model_files_for_gradio, inputs=[model_name_b], outputs=[model_path_b], ) From 1d320915e0704fbaabd49d117b3ba8b8ccf363c2 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sat, 9 Mar 2024 18:10:25 +0000 Subject: [PATCH 52/64] Refactor: make tts_model.py independent of style_gen.py --- style_bert_vits2/tts_model.py | 102 ++++++++++++++++++--------- style_bert_vits2/utils/subprocess.py | 4 +- style_gen.py | 3 +- 3 files changed, 70 insertions(+), 39 deletions(-) diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index a3a83c3..42f6b84 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -4,6 +4,7 @@ from typing import Any, Optional, Union import gradio as gr import numpy as np +import pyannote.audio import torch from gradio.processing_utils import convert_to_16_bit_wav from numpy.typing import NDArray @@ -30,8 +31,8 @@ from style_bert_vits2.voice import adjust_voice class Model: """ - Style-Bert-Vits2 の音声合成モデルを操作するためのクラス - モデル/ハイパーパラメータ/スタイルベクトルのパスとデバイスを指定して初期化し、model.infer() メソッドを呼び出すと音声合成を行える + Style-Bert-Vits2 の音声合成モデルを操作するためのクラス。 + モデル/ハイパーパラメータ/スタイルベクトルのパスとデバイスを指定して初期化し、model.infer() メソッドを呼び出すと音声合成を行える。 """ @@ -46,50 +47,81 @@ class Model: self.config_path: Path = config_path self.style_vec_path: Path = style_vec_path self.device: str = device - self.hps: HyperParameters = HyperParameters.load_from_json(self.config_path) - self.spk2id: dict[str, int] = self.hps.data.spk2id + self.hyper_parameters: HyperParameters = HyperParameters.load_from_json(self.config_path) + self.spk2id: dict[str, int] = self.hyper_parameters.data.spk2id self.id2spk: dict[int, str] = {v: k for k, v in self.spk2id.items()} - self.num_styles: int = self.hps.data.num_styles - if hasattr(self.hps.data, "style2id"): - self.style2id: dict[str, int] = self.hps.data.style2id + num_styles: int = self.hyper_parameters.data.num_styles + if hasattr(self.hyper_parameters.data, "style2id"): + self.style2id: dict[str, int] = self.hyper_parameters.data.style2id else: - self.style2id: dict[str, int] = {str(i): i for i in range(self.num_styles)} - if len(self.style2id) != self.num_styles: + self.style2id: dict[str, int] = {str(i): i for i in range(num_styles)} + if len(self.style2id) != num_styles: raise ValueError( - f"Number of styles ({self.num_styles}) does not match the number of style2id ({len(self.style2id)})" + f"Number of styles ({num_styles}) does not match the number of style2id ({len(self.style2id)})" ) - self.style_vectors: NDArray[Any] = np.load(self.style_vec_path) - if self.style_vectors.shape[0] != self.num_styles: + self.__style_vector_inference: Optional[pyannote.audio.Inference] = None + self.__style_vectors: NDArray[Any] = np.load(self.style_vec_path) + if self.__style_vectors.shape[0] != num_styles: raise ValueError( - f"The number of styles ({self.num_styles}) does not match the number of style vectors ({self.style_vectors.shape[0]})" + f"The number of styles ({num_styles}) does not match the number of style vectors ({self.__style_vectors.shape[0]})" ) - self.net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None + self.__net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None def load_net_g(self) -> None: - self.net_g = get_net_g( + """ + net_g をロードする。 + """ + self.__net_g = get_net_g( model_path=str(self.model_path), - version=self.hps.version, + version=self.hyper_parameters.version, device=self.device, - hps=self.hps, + hps=self.hyper_parameters, ) def get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]: - mean = self.style_vectors[0] - style_vec = self.style_vectors[style_id] + """ + スタイルベクトルを取得する。 + + Args: + style_id (int): スタイル ID + weight (float, optional): スタイルベクトルの重み. Defaults to 1.0. + + Returns: + NDArray[Any]: スタイルベクトル + """ + mean = self.__style_vectors[0] + style_vec = self.__style_vectors[style_id] style_vec = mean + (style_vec - mean) * weight return style_vec def get_style_vector_from_audio(self, audio_path: str, weight: float = 1.0) -> NDArray[Any]: - from style_gen import get_style_vector + """ + 音声からスタイルベクトルを推論する。 - xvec = get_style_vector(audio_path) - mean = self.style_vectors[0] + Args: + audio_path (str): 音声ファイルのパス + weight (float, optional): スタイルベクトルの重み. Defaults to 1.0. + Returns: + NDArray[Any]: スタイルベクトル + """ + + # スタイルベクトルを取得するための推論モデルを初期化 + if self.__style_vector_inference is None: + self.__style_vector_inference = pyannote.audio.Inference( + model = pyannote.audio.Model.from_pretrained("pyannote/wespeaker-voxceleb-resnet34-LM"), + window = "whole", + ) + self.__style_vector_inference.to(torch.device(self.device)) + + # 音声からスタイルベクトルを推論 + xvec = self.__style_vector_inference(audio_path) + mean = self.__style_vectors[0] xvec = mean + (xvec - mean) * weight return xvec @@ -116,7 +148,7 @@ class Model: intonation_scale: float = 1.0, ) -> tuple[int, NDArray[Any]]: logger.info(f"Start generating audio data from text:\n{text}") - if language != "JP" and self.hps.version.endswith("JP-Extra"): + if language != "JP" and self.hyper_parameters.version.endswith("JP-Extra"): raise ValueError( "The model is trained with JP-Extra, but the language is not JP" ) @@ -125,9 +157,9 @@ class Model: if assist_text == "" or not use_assist_text: assist_text = None - if self.net_g is None: + if self.__net_g is None: self.load_net_g() - assert self.net_g is not None + assert self.__net_g is not None if reference_audio_path is None: style_id = self.style2id[style] style_vector = self.get_style_vector(style_id, style_weight) @@ -145,8 +177,8 @@ class Model: length_scale = length, sid = sid, language = language, - hps = self.hps, - net_g = self.net_g, + hps = self.hyper_parameters, + net_g = self.__net_g, device = self.device, assist_text = assist_text, assist_text_weight = assist_text_weight, @@ -168,8 +200,8 @@ class Model: length_scale = length, sid = sid, language = language, - hps = self.hps, - net_g = self.net_g, + hps = self.hyper_parameters, + net_g = self.__net_g, device = self.device, assist_text = assist_text, assist_text_weight = assist_text_weight, @@ -182,7 +214,7 @@ class Model: logger.info("Audio data generated successfully") if not (pitch_scale == 1.0 and intonation_scale == 1.0): _, audio = adjust_voice( - fs = self.hps.data.sampling_rate, + fs = self.hyper_parameters.data.sampling_rate, wave = audio, pitch_scale = pitch_scale, intonation_scale = intonation_scale, @@ -190,12 +222,12 @@ class Model: with warnings.catch_warnings(): warnings.simplefilter("ignore") audio = convert_to_16_bit_wav(audio) - return (self.hps.data.sampling_rate, audio) + return (self.hyper_parameters.data.sampling_rate, audio) class ModelHolder: """ - Style-Bert-Vits2 の音声合成モデルを管理するためのクラス + Style-Bert-Vits2 の音声合成モデルを管理するためのクラス。 """ @@ -234,10 +266,10 @@ class ModelHolder: continue self.model_files_dict[model_dir.name] = model_files self.model_names.append(model_dir.name) - hps = HyperParameters.load_from_json(config_path) - style2id: dict[str, int] = hps.data.style2id + hyper_parameters = HyperParameters.load_from_json(config_path) + style2id: dict[str, int] = hyper_parameters.data.style2id styles = list(style2id.keys()) - spk2id: dict[str, int] = hps.data.spk2id + spk2id: dict[str, int] = hyper_parameters.data.spk2id speakers = list(spk2id.keys()) self.models_info.append({ "name": model_dir.name, diff --git a/style_bert_vits2/utils/subprocess.py b/style_bert_vits2/utils/subprocess.py index 5ff267b..542f94b 100644 --- a/style_bert_vits2/utils/subprocess.py +++ b/style_bert_vits2/utils/subprocess.py @@ -8,7 +8,7 @@ from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT def run_script_with_log(cmd: list[str], ignore_warning: bool = False) -> tuple[bool, str]: """ - 指定されたコマンドを実行し、そのログを記録する + 指定されたコマンドを実行し、そのログを記録する。 Args: cmd: 実行するコマンドのリスト @@ -39,7 +39,7 @@ def run_script_with_log(cmd: list[str], ignore_warning: bool = False) -> tuple[b def second_elem_of(original_function: Callable[..., tuple[Any, Any]]) -> Callable[..., Any]: """ - 与えられた関数をラップし、その戻り値の 2 番目の要素のみを返す関数を生成する + 与えられた関数をラップし、その戻り値の 2 番目の要素のみを返す関数を生成する。 Args: original_function (Callable[..., tuple[Any, Any]])): ラップする元の関数 diff --git a/style_gen.py b/style_gen.py index ec0b507..5190575 100644 --- a/style_gen.py +++ b/style_gen.py @@ -6,11 +6,10 @@ import numpy as np import torch from tqdm import tqdm +from config import config from style_bert_vits2.logging import logger -from style_bert_vits2.models import utils from style_bert_vits2.models.hyper_parameters import HyperParameters from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT -from config import config warnings.filterwarnings("ignore", category=UserWarning) from pyannote.audio import Inference, Model From a79b1910fb1945151903458e4ddca938d008d096 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sat, 9 Mar 2024 18:38:16 +0000 Subject: [PATCH 53/64] Add: use hatch to build style-bert-vits2 as a library --- .gitignore | 1 + app.py | 4 +- pyproject.toml | 92 +++++++++++++++++++++++++++++++++++ server_editor.py | 4 +- style_bert_vits2/constants.py | 2 +- tests/__init__.py | 0 6 files changed, 98 insertions(+), 5 deletions(-) create mode 100644 pyproject.toml create mode 100644 tests/__init__.py diff --git a/.gitignore b/.gitignore index 88048d7..b8a19a4 100644 --- a/.gitignore +++ b/.gitignore @@ -3,6 +3,7 @@ __pycache__/ venv/ .venv/ +dist/ .ipynb_checkpoints/ /*.yml diff --git a/app.py b/app.py index f0f94d0..eaf07b7 100644 --- a/app.py +++ b/app.py @@ -5,7 +5,7 @@ import gradio as gr import torch import yaml -from style_bert_vits2.constants import GRADIO_THEME, LATEST_VERSION +from style_bert_vits2.constants import GRADIO_THEME, VERSION from style_bert_vits2.tts_model import ModelHolder from webui import ( create_dataset_app, @@ -37,7 +37,7 @@ if device == "cuda" and not torch.cuda.is_available(): model_holder = ModelHolder(Path(assets_root), device) with gr.Blocks(theme=GRADIO_THEME) as app: - gr.Markdown(f"# Style-Bert-VITS2 WebUI (version {LATEST_VERSION})") + gr.Markdown(f"# Style-Bert-VITS2 WebUI (version {VERSION})") with gr.Tabs(): with gr.Tab("音声合成"): create_inference_app(model_holder=model_holder) diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..4db5fe9 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,92 @@ +[build-system] +requires = ["hatchling"] +build-backend = "hatchling.build" + +[project] +name = "style-bert-vits2" +dynamic = ["version"] +description = 'Style-Bert-VITS2: Bert-VITS2 with more controllable voice styles.' +readme = "README.md" +requires-python = ">=3.9" +license = "AGPL-3.0" +keywords = [] +authors = [ + { name = "litagin02", email = "139731664+litagin02@users.noreply.github.com" }, +] +classifiers = [ + "Development Status :: 4 - Beta", + "Programming Language :: Python", + "Programming Language :: Python :: 3.9", + "Programming Language :: Python :: 3.10", + "Programming Language :: Python :: 3.11", + "Programming Language :: Python :: 3.12", + "Programming Language :: Python :: Implementation :: CPython", +] +dependencies = [ + 'cmudict', + 'cn2an', + 'g2p_en', + 'gradio', + 'jieba', + 'librosa==0.9.2', + 'loguru', + 'num2words', + 'numba', + 'numpy', + 'pyannote.audio>=3.1.0', + 'pydantic', + 'pyopenjtalk-dict', + 'pypinyin', + 'pyworld', + 'safetensors', + 'scipy', + 'torch>=2.1,<2.2', + 'transformers', +] + +[project.urls] +Documentation = "https://github.com/litagin02/Style-Bert-VITS2#readme" +Issues = "https://github.com/litagin02/Style-Bert-VITS2/issues" +Source = "https://github.com/litagin02/Style-Bert-VITS2" + +[tool.hatch.version] +path = "style_bert_vits2/constants.py" + +[tool.hatch.envs.default] +dependencies = [ + "coverage[toml]>=6.5", + "pytest", +] +[tool.hatch.envs.default.scripts] +test = "pytest {args:tests}" +test-cov = "coverage run -m pytest {args:tests}" +cov-report = [ + "- coverage combine", + "coverage report", +] +cov = [ + "test-cov", + "cov-report", +] + +[[tool.hatch.envs.all.matrix]] +python = ["3.9", "3.10", "3.11", "3.12"] + +[tool.coverage.run] +source_pkgs = ["style_bert_vits2", "tests"] +branch = true +parallel = true +omit = [ + "style_bert_vits2/constants.py", +] + +[tool.coverage.paths] +style_bert_vits2 = ["style_bert_vits2", "*/style-bert-vits2/style_bert_vits2"] +tests = ["tests", "*/style-bert-vits2/tests"] + +[tool.coverage.report] +exclude_lines = [ + "no cov", + "if __name__ == .__main__.:", + "if TYPE_CHECKING:", +] diff --git a/server_editor.py b/server_editor.py index 9a6a08a..b12f304 100644 --- a/server_editor.py +++ b/server_editor.py @@ -37,7 +37,7 @@ from style_bert_vits2.constants import ( DEFAULT_SDP_RATIO, DEFAULT_STYLE, DEFAULT_STYLE_WEIGHT, - LATEST_VERSION, + VERSION, Languages, ) from style_bert_vits2.logging import logger @@ -219,7 +219,7 @@ router = APIRouter() @router.get("/version") def version() -> str: - return LATEST_VERSION + return VERSION class MoraTone(BaseModel): diff --git a/style_bert_vits2/constants.py b/style_bert_vits2/constants.py index 5f32a15..735aebc 100644 --- a/style_bert_vits2/constants.py +++ b/style_bert_vits2/constants.py @@ -4,7 +4,7 @@ from style_bert_vits2.utils.strenum import StrEnum # Style-Bert-VITS2 のバージョン -LATEST_VERSION = "2.4" +VERSION = "2.4" # Style-Bert-VITS2 のベースディレクトリ BASE_DIR = Path(__file__).parent.parent diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..e69de29 From b7d7c7820364c7d17f8bedad04b8aa5441bec6ba Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sun, 10 Mar 2024 03:04:27 +0000 Subject: [PATCH 54/64] Refactor: when pyopenjtalk_worker is called without initialization, continue processing without a worker When using style-bert-vits2 as a library, the requirement to be able to launch it in multiple processes may not be necessary. Also, if the library is embedded and exe-ed using PyInstaller or similar, it is difficult to make pyopenjtalk_worker run in a separate process. Therefore, we changed it so that the worker is used only when it is explicitly initialized. --- server_editor.py | 2 +- server_fastapi.py | 2 +- .../japanese/pyopenjtalk_worker/__init__.py | 64 ++++++++++++------- webui/inference.py | 2 +- 4 files changed, 45 insertions(+), 25 deletions(-) diff --git a/server_editor.py b/server_editor.py index b12f304..3c0d2c8 100644 --- a/server_editor.py +++ b/server_editor.py @@ -151,7 +151,7 @@ def save_last_download(latest_release): # pyopenjtalk_worker を起動 ## pyopenjtalk_worker は TCP ソケットサーバーのため、ここで起動する -pyopenjtalk.initialize() +pyopenjtalk.initialize_worker() # pyopenjtalk の辞書を更新 update_dict() diff --git a/server_fastapi.py b/server_fastapi.py index e8309da..30f253f 100644 --- a/server_fastapi.py +++ b/server_fastapi.py @@ -43,7 +43,7 @@ ln = config.server_config.language # pyopenjtalk_worker を起動 ## pyopenjtalk_worker は TCP ソケットサーバーのため、ここで起動する -pyopenjtalk.initialize() +pyopenjtalk.initialize_worker() # 事前に BERT モデル/トークナイザーをロードしておく ## ここでロードしなくても必要になった際に自動ロードされるが、時間がかかるため事前にロードしておいた方が体験が良い diff --git a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py index 3593e60..212d21a 100644 --- a/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py +++ b/style_bert_vits2/nlp/japanese/pyopenjtalk_worker/__init__.py @@ -18,38 +18,58 @@ WORKER_CLIENT: Optional[WorkerClient] = None def run_frontend(text: str) -> list[dict[str, Any]]: - assert WORKER_CLIENT - ret = WORKER_CLIENT.dispatch_pyopenjtalk("run_frontend", text) - assert isinstance(ret, list) - return ret + if WORKER_CLIENT is not None: + ret = WORKER_CLIENT.dispatch_pyopenjtalk("run_frontend", text) + assert isinstance(ret, list) + return ret + else: + # without worker + import pyopenjtalk + return pyopenjtalk.run_frontend(text) def make_label(njd_features: Any) -> list[str]: - assert WORKER_CLIENT - ret = WORKER_CLIENT.dispatch_pyopenjtalk("make_label", njd_features) - assert isinstance(ret, list) - return ret + if WORKER_CLIENT is not None: + ret = WORKER_CLIENT.dispatch_pyopenjtalk("make_label", njd_features) + assert isinstance(ret, list) + return ret + else: + # without worker + import pyopenjtalk + return pyopenjtalk.make_label(njd_features) -def mecab_dict_index(path: str, out_path: str, dn_mecab: Optional[str] = None): - assert WORKER_CLIENT - WORKER_CLIENT.dispatch_pyopenjtalk("mecab_dict_index", path, out_path, dn_mecab) +def mecab_dict_index(path: str, out_path: str, dn_mecab: Optional[str] = None) -> None: + if WORKER_CLIENT is not None: + WORKER_CLIENT.dispatch_pyopenjtalk("mecab_dict_index", path, out_path, dn_mecab) + else: + # without worker + import pyopenjtalk + pyopenjtalk.mecab_dict_index(path, out_path, dn_mecab) -def update_global_jtalk_with_user_dict(path: str): - assert WORKER_CLIENT - WORKER_CLIENT.dispatch_pyopenjtalk("update_global_jtalk_with_user_dict", path) +def update_global_jtalk_with_user_dict(path: str) -> None: + if WORKER_CLIENT is not None: + WORKER_CLIENT.dispatch_pyopenjtalk("update_global_jtalk_with_user_dict", path) + else: + # without worker + import pyopenjtalk + pyopenjtalk.update_global_jtalk_with_user_dict(path) -def unset_user_dict(): - assert WORKER_CLIENT - WORKER_CLIENT.dispatch_pyopenjtalk("unset_user_dict") +def unset_user_dict() -> None: + if WORKER_CLIENT is not None: + WORKER_CLIENT.dispatch_pyopenjtalk("unset_user_dict") + else: + # without worker + import pyopenjtalk + pyopenjtalk.unset_user_dict() # initialize module when imported -def initialize(port: int = WORKER_PORT) -> None: +def initialize_worker(port: int = WORKER_PORT) -> None: import atexit import signal import socket @@ -99,11 +119,11 @@ def initialize(port: int = WORKER_PORT) -> None: logger.debug("pyopenjtalk worker server started") WORKER_CLIENT = client - atexit.register(terminate) + atexit.register(terminate_worker) # when the process is killed def signal_handler(signum: int, frame: Any): - terminate() + terminate_worker() try: signal.signal(signal.SIGTERM, signal_handler) @@ -113,13 +133,13 @@ def initialize(port: int = WORKER_PORT) -> None: # top-level declaration -def terminate() -> None: +def terminate_worker() -> None: logger.debug("pyopenjtalk worker server terminated") global WORKER_CLIENT if not WORKER_CLIENT: return - # repare for unexpected errors + # prepare for unexpected errors try: if WORKER_CLIENT.status() == 1: WORKER_CLIENT.quit_server() diff --git a/webui/inference.py b/webui/inference.py index 9c2bf63..6d6eae6 100644 --- a/webui/inference.py +++ b/webui/inference.py @@ -28,7 +28,7 @@ from style_bert_vits2.tts_model import ModelHolder # pyopenjtalk_worker を起動 ## pyopenjtalk_worker は TCP ソケットサーバーのため、ここで起動する -pyopenjtalk.initialize() +pyopenjtalk.initialize_worker() # 事前に BERT モデル/トークナイザーをロードしておく ## ここでロードしなくても必要になった際に自動ロードされるが、時間がかかるため事前にロードしておいた方が体験が良い From d2fd378b565239fd272520a011f2df84d1916a3b Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sun, 10 Mar 2024 03:47:17 +0000 Subject: [PATCH 55/64] Refactor: rename Model / ModelHolder to TTSModel / TTSModelHolder for clarification and add comments to each method --- app.py | 4 +- server_editor.py | 14 ++-- server_fastapi.py | 14 ++-- speech_mos.py | 4 +- style_bert_vits2/tts_model.py | 133 +++++++++++++++++++++++++++------- webui/inference.py | 12 +-- webui/merge.py | 8 +- 7 files changed, 133 insertions(+), 56 deletions(-) diff --git a/app.py b/app.py index eaf07b7..b4fbc3e 100644 --- a/app.py +++ b/app.py @@ -6,7 +6,7 @@ import torch import yaml from style_bert_vits2.constants import GRADIO_THEME, VERSION -from style_bert_vits2.tts_model import ModelHolder +from style_bert_vits2.tts_model import TTSModelHolder from webui import ( create_dataset_app, create_inference_app, @@ -34,7 +34,7 @@ device = args.device if device == "cuda" and not torch.cuda.is_available(): device = "cpu" -model_holder = ModelHolder(Path(assets_root), device) +model_holder = TTSModelHolder(Path(assets_root), device) with gr.Blocks(theme=GRADIO_THEME) as app: gr.Markdown(f"# Style-Bert-VITS2 WebUI (version {VERSION})") diff --git a/server_editor.py b/server_editor.py index 3c0d2c8..aa85ef0 100644 --- a/server_editor.py +++ b/server_editor.py @@ -52,7 +52,7 @@ from style_bert_vits2.nlp.japanese.user_dict import ( rewrite_word, update_dict, ) -from style_bert_vits2.tts_model import ModelHolder +from style_bert_vits2.tts_model import TTSModelHolder # ---フロントエンド部分に関する処理--- @@ -198,7 +198,7 @@ if device == "cuda" and not torch.cuda.is_available(): model_dir = Path(args.model_dir) port = int(args.port) -model_holder = ModelHolder(model_dir, device) +model_holder = TTSModelHolder(model_dir, device) if len(model_holder.model_names) == 0: logger.error(f"Models not found in {model_dir}.") sys.exit(1) @@ -283,7 +283,7 @@ def synthesis(request: SynthesisRequest): detail=f"1行の文字数は{args.line_length}文字以下にしてください。", ) try: - model = model_holder.load_model( + model = model_holder.get_model( model_name=request.model, model_path_str=request.modelFile ) except Exception as e: @@ -310,7 +310,7 @@ def synthesis(request: SynthesisRequest): language=request.language, sdp_ratio=request.sdpRatio, noise=request.noise, - noisew=request.noisew, + noise_w=request.noisew, length=1 / request.speed, given_tone=tone, style=request.style, @@ -321,7 +321,7 @@ def synthesis(request: SynthesisRequest): line_split=False, pitch_scale=request.pitchScale, intonation_scale=request.intonationScale, - sid=sid, + speaker_id=sid, ) with BytesIO() as wavContent: @@ -350,7 +350,7 @@ def multi_synthesis(request: MultiSynthesisRequest): detail=f"1行の文字数は{args.line_length}文字以下にしてください。", ) try: - model = model_holder.load_model( + model = model_holder.get_model( model_name=req.model, model_path_str=req.modelFile ) except Exception as e: @@ -370,7 +370,7 @@ def multi_synthesis(request: MultiSynthesisRequest): language=req.language, sdp_ratio=req.sdpRatio, noise=req.noise, - noisew=req.noisew, + noise_w=req.noisew, length=1 / req.speed, given_tone=tone, style=req.style, diff --git a/server_fastapi.py b/server_fastapi.py index 30f253f..1f60b3c 100644 --- a/server_fastapi.py +++ b/server_fastapi.py @@ -36,7 +36,7 @@ from style_bert_vits2.constants import ( from style_bert_vits2.logging import logger from style_bert_vits2.nlp import bert_models from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk -from style_bert_vits2.tts_model import Model, ModelHolder +from style_bert_vits2.tts_model import TTSModel, TTSModelHolder ln = config.server_config.language @@ -67,16 +67,16 @@ class AudioResponse(Response): media_type = "audio/wav" -def load_models(model_holder: ModelHolder): +def load_models(model_holder: TTSModelHolder): model_holder.models = [] for model_name, model_paths in model_holder.model_files_dict.items(): - model = Model( + model = TTSModel( model_path=model_paths[0], config_path=model_holder.root_dir / model_name / "config.json", style_vec_path=model_holder.root_dir / model_name / "style_vectors.npy", device=model_holder.device, ) - model.load_net_g() + model.load() model_holder.models.append(model) @@ -94,7 +94,7 @@ if __name__ == "__main__": device = "cuda" if torch.cuda.is_available() else "cpu" model_dir = Path(args.dir) - model_holder = ModelHolder(model_dir, device) + model_holder = TTSModelHolder(model_dir, device) if len(model_holder.model_names) == 0: logger.error(f"Models not found in {model_dir}.") sys.exit(1) @@ -194,11 +194,11 @@ if __name__ == "__main__": sr, audio = model.infer( text=text, language=language, - sid=speaker_id, + speaker_id=speaker_id, reference_audio_path=reference_audio_path, sdp_ratio=sdp_ratio, noise=noise, - noisew=noisew, + noise_w=noisew, length=length, line_split=auto_split, split_interval=split_interval, diff --git a/speech_mos.py b/speech_mos.py index 79221ad..6dd2caa 100644 --- a/speech_mos.py +++ b/speech_mos.py @@ -12,7 +12,7 @@ from tqdm import tqdm from config import config from style_bert_vits2.logging import logger -from style_bert_vits2.tts_model import Model +from style_bert_vits2.tts_model import TTSModel warnings.filterwarnings("ignore") @@ -54,7 +54,7 @@ safetensors_files = model_path.glob("*.safetensors") def get_model(model_file: Path): - return Model( + return TTSModel( model_path=str(model_file), config_path=str(model_file.parent / "config.json"), style_vec_path=str(model_file.parent / "style_vectors.npy"), diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index 42f6b84..109ef0e 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -29,9 +29,9 @@ from style_bert_vits2.logging import logger from style_bert_vits2.voice import adjust_voice -class Model: +class TTSModel: """ - Style-Bert-Vits2 の音声合成モデルを操作するためのクラス。 + Style-Bert-Vits2 の音声合成モデルを操作するクラス。 モデル/ハイパーパラメータ/スタイルベクトルのパスとデバイスを指定して初期化し、model.infer() メソッドを呼び出すと音声合成を行える。 """ @@ -43,6 +43,17 @@ class Model: style_vec_path: Path, device: str, ) -> None: + """ + Style-Bert-Vits2 の音声合成モデルを初期化する。 + この時点ではモデルはロードされていない (明示的にロードしたい場合は model.load() を呼び出す)。 + + Args: + model_path (Path): モデル (.safetensors) のパス + config_path (Path): ハイパーパラメータ (config.json) のパス + style_vec_path (Path): スタイルベクトル (style_vectors.npy) のパス + device (str): 音声合成時に利用するデバイス (cpu, cuda, mps など) + """ + self.model_path: Path = model_path self.config_path: Path = config_path self.style_vec_path: Path = style_vec_path @@ -71,24 +82,24 @@ class Model: self.__net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None - def load_net_g(self) -> None: + def load(self) -> None: """ - net_g をロードする。 + 音声合成モデルをデバイスにロードする。 """ self.__net_g = get_net_g( - model_path=str(self.model_path), - version=self.hyper_parameters.version, - device=self.device, - hps=self.hyper_parameters, + model_path = str(self.model_path), + version = self.hyper_parameters.version, + device = self.device, + hps = self.hyper_parameters, ) - def get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]: + def __get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]: """ スタイルベクトルを取得する。 Args: - style_id (int): スタイル ID + style_id (int): スタイル ID (0 から始まるインデックス) weight (float, optional): スタイルベクトルの重み. Defaults to 1.0. Returns: @@ -100,7 +111,7 @@ class Model: return style_vec - def get_style_vector_from_audio(self, audio_path: str, weight: float = 1.0) -> NDArray[Any]: + def __get_style_vector_from_audio(self, audio_path: str, weight: float = 1.0) -> NDArray[Any]: """ 音声からスタイルベクトルを推論する。 @@ -130,11 +141,11 @@ class Model: self, text: str, language: Languages = Languages.JP, - sid: int = 0, + speaker_id: int = 0, reference_audio_path: Optional[str] = None, sdp_ratio: float = DEFAULT_SDP_RATIO, noise: float = DEFAULT_NOISE, - noisew: float = DEFAULT_NOISEW, + noise_w: float = DEFAULT_NOISEW, length: float = DEFAULT_LENGTH, line_split: bool = DEFAULT_LINE_SPLIT, split_interval: float = DEFAULT_SPLIT_INTERVAL, @@ -147,6 +158,33 @@ class Model: pitch_scale: float = 1.0, intonation_scale: float = 1.0, ) -> tuple[int, NDArray[Any]]: + """ + テキストから音声を合成する。 + + Args: + text (str): 読み上げるテキスト + language (Languages, optional): 言語. Defaults to Languages.JP. + speaker_id (int, optional): 話者 ID. Defaults to 0. + reference_audio_path (Optional[str], optional): 音声スタイルの参照元の音声ファイルのパス. Defaults to None. + sdp_ratio (float, optional): SDP レシオ (値を大きくするとより感情豊かになる傾向がある). Defaults to DEFAULT_SDP_RATIO. + noise (float, optional): ノイズの大きさ. Defaults to DEFAULT_NOISE. + noise_w (float, optional): ノイズの大きさの重み. Defaults to DEFAULT_NOISEW. + length (float, optional): 長さ. Defaults to DEFAULT_LENGTH. + line_split (bool, optional): テキストを改行ごとに分割して生成するかどうか. Defaults to DEFAULT_LINE_SPLIT. + split_interval (float, optional): 改行ごとに分割する場合の無音 (秒). Defaults to DEFAULT_SPLIT_INTERVAL. + assist_text (Optional[str], optional): 感情表現の参照元の補助テキスト. Defaults to None. + assist_text_weight (float, optional): 感情表現の補助テキストを適用する強さ. Defaults to DEFAULT_ASSIST_TEXT_WEIGHT. + use_assist_text (bool, optional): 音声合成時に感情表現の補助テキストを使用するかどうか. Defaults to False. + style (str, optional): 音声スタイル (Neutral, Happy など). Defaults to DEFAULT_STYLE. + style_weight (float, optional): 音声スタイルを適用する強さ. Defaults to DEFAULT_STYLE_WEIGHT. + given_tone (Optional[list[int]], optional): アクセントのトーンのリスト. Defaults to None. + pitch_scale (float, optional): ピッチの高さ (1.0 から変更すると若干音質が低下する). Defaults to 1.0. + intonation_scale (float, optional): イントネーションの高さ (1.0 から変更すると若干音質が低下する). Defaults to 1.0. + + Returns: + tuple[int, NDArray[Any]]: サンプリングレートと音声データ (16bit PCM) + """ + logger.info(f"Start generating audio data from text:\n{text}") if language != "JP" and self.hyper_parameters.version.endswith("JP-Extra"): raise ValueError( @@ -158,13 +196,13 @@ class Model: assist_text = None if self.__net_g is None: - self.load_net_g() + self.load() assert self.__net_g is not None if reference_audio_path is None: style_id = self.style2id[style] - style_vector = self.get_style_vector(style_id, style_weight) + style_vector = self.__get_style_vector(style_id, style_weight) else: - style_vector = self.get_style_vector_from_audio( + style_vector = self.__get_style_vector_from_audio( reference_audio_path, style_weight ) if not line_split: @@ -173,9 +211,9 @@ class Model: text = text, sdp_ratio = sdp_ratio, noise_scale = noise, - noise_scale_w = noisew, + noise_scale_w = noise_w, length_scale = length, - sid = sid, + sid = speaker_id, language = language, hps = self.hyper_parameters, net_g = self.__net_g, @@ -196,9 +234,9 @@ class Model: text = t, sdp_ratio = sdp_ratio, noise_scale = noise, - noise_scale_w = noisew, + noise_scale_w = noise_w, length_scale = length, - sid = sid, + sid = speaker_id, language = language, hps = self.hyper_parameters, net_g = self.__net_g, @@ -225,24 +263,50 @@ class Model: return (self.hyper_parameters.data.sampling_rate, audio) -class ModelHolder: +class TTSModelHolder: """ - Style-Bert-Vits2 の音声合成モデルを管理するためのクラス。 + Style-Bert-Vits2 の音声合成モデルを管理するクラス。 + model_holder.models_info から指定されたディレクトリ内にある音声合成モデルの一覧を取得できる。 """ def __init__(self, model_root_dir: Path, device: str) -> None: + """ + Style-Bert-Vits2 の音声合成モデルを管理するクラスを初期化する。 + 音声合成モデルは下記のように配置されていることを前提とする (.safetensors のファイル名は自由) 。 + ``` + model_root_dir + ├── model-name-1 + │ ├── config.json + │ ├── model-name-1_e160_s14000.safetensors + │ └── style_vectors.npy + ├── model-name-2 + │ ├── config.json + │ ├── model-name-2_e160_s14000.safetensors + │ └── style_vectors.npy + └── ... + ``` + + Args: + model_root_dir (Path): 音声合成モデルが配置されているディレクトリのパス + device (str): 音声合成時に利用するデバイス (cpu, cuda, mps など) + """ + self.root_dir: Path = model_root_dir self.device: str = device self.model_files_dict: dict[str, list[Path]] = {} - self.current_model: Optional[Model] = None + self.current_model: Optional[TTSModel] = None self.model_names: list[str] = [] - self.models: list[Model] = [] + self.models: list[TTSModel] = [] self.models_info: list[dict[str, Union[str, list[str]]]] = [] self.refresh() def refresh(self) -> None: + """ + 音声合成モデルの一覧を更新する。 + """ + self.model_files_dict = {} self.model_names = [] self.current_model = None @@ -279,23 +343,36 @@ class ModelHolder: }) - def load_model(self, model_name: str, model_path_str: str) -> Model: + def get_model(self, model_name: str, model_path_str: str) -> TTSModel: + """ + 指定された音声合成モデルのインスタンスを取得する。 + この時点ではモデルはロードされていない (明示的にロードしたい場合は model.load() を呼び出す)。 + + Args: + model_name (str): 音声合成モデルの名前 + model_path_str (str): 音声合成モデルのファイルパス (.safetensors) + + Returns: + TTSModel: 音声合成モデルのインスタンス + """ + model_path = Path(model_path_str) if model_name not in self.model_files_dict: raise ValueError(f"Model `{model_name}` is not found") if model_path not in self.model_files_dict[model_name]: raise ValueError(f"Model file `{model_path}` is not found") if self.current_model is None or self.current_model.model_path != model_path: - self.current_model = Model( + self.current_model = TTSModel( model_path = model_path, config_path = self.root_dir / model_name / "config.json", style_vec_path = self.root_dir / model_name / "style_vectors.npy", device = self.device, ) + return self.current_model - def load_model_for_gradio(self, model_name: str, model_path_str: str) -> tuple[gr.Dropdown, gr.Button, gr.Dropdown]: + def get_model_for_gradio(self, model_name: str, model_path_str: str) -> tuple[gr.Dropdown, gr.Button, gr.Dropdown]: model_path = Path(model_path_str) if model_name not in self.model_files_dict: raise ValueError(f"Model `{model_name}` is not found") @@ -313,7 +390,7 @@ class ModelHolder: gr.Button(interactive=True, value="音声合成"), gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore ) - self.current_model = Model( + self.current_model = TTSModel( model_path = model_path, config_path = self.root_dir / model_name / "config.json", style_vec_path = self.root_dir / model_name / "style_vectors.npy", diff --git a/webui/inference.py b/webui/inference.py index 6d6eae6..e32907c 100644 --- a/webui/inference.py +++ b/webui/inference.py @@ -23,7 +23,7 @@ from style_bert_vits2.nlp import bert_models from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone from style_bert_vits2.nlp.japanese.normalizer import normalize_text -from style_bert_vits2.tts_model import ModelHolder +from style_bert_vits2.tts_model import TTSModelHolder # pyopenjtalk_worker を起動 @@ -151,7 +151,7 @@ def gr_util(item): return (gr.update(visible=False), gr.update(visible=True)) -def create_inference_app(model_holder: ModelHolder) -> gr.Blocks: +def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: def tts_fn( model_name, model_path, @@ -175,7 +175,7 @@ def create_inference_app(model_holder: ModelHolder) -> gr.Blocks: pitch_scale, intonation_scale, ): - model_holder.load_model(model_name, model_path) + model_holder.get_model(model_name, model_path) assert model_holder.current_model is not None wrong_tone_message = "" @@ -218,7 +218,7 @@ def create_inference_app(model_holder: ModelHolder) -> gr.Blocks: reference_audio_path=reference_audio_path, sdp_ratio=sdp_ratio, noise=noise_scale, - noisew=noise_scale_w, + noise_w=noise_scale_w, length=length_scale, line_split=line_split, split_interval=split_interval, @@ -228,7 +228,7 @@ def create_inference_app(model_holder: ModelHolder) -> gr.Blocks: style=style, style_weight=style_weight, given_tone=tone, - sid=speaker_id, + speaker_id=speaker_id, pitch_scale=pitch_scale, intonation_scale=intonation_scale, ) @@ -459,7 +459,7 @@ def create_inference_app(model_holder: ModelHolder) -> gr.Blocks: ) load_button.click( - model_holder.load_model_for_gradio, + model_holder.get_model_for_gradio, inputs=[model_name, model_path], outputs=[style, tts_button, speaker], ) diff --git a/webui/merge.py b/webui/merge.py index 85806f4..f692b67 100644 --- a/webui/merge.py +++ b/webui/merge.py @@ -11,7 +11,7 @@ from safetensors.torch import save_file from style_bert_vits2.constants import DEFAULT_STYLE, GRADIO_THEME from style_bert_vits2.logging import logger -from style_bert_vits2.tts_model import Model, ModelHolder +from style_bert_vits2.tts_model import TTSModel, TTSModelHolder voice_keys = ["dec"] @@ -250,11 +250,11 @@ def simple_tts(model_name, text, style=DEFAULT_STYLE, style_weight=1.0): config_path = os.path.join(assets_root, model_name, "config.json") style_vec_path = os.path.join(assets_root, model_name, "style_vectors.npy") - model = Model(Path(model_path), Path(config_path), Path(style_vec_path), device) + model = TTSModel(Path(model_path), Path(config_path), Path(style_vec_path), device) return model.infer(text, style=style, style_weight=style_weight) -def update_two_model_names_dropdown(model_holder: ModelHolder): +def update_two_model_names_dropdown(model_holder: TTSModelHolder): new_names, new_files, _ = model_holder.update_model_names_for_gradio() return new_names, new_files, new_names, new_files @@ -328,7 +328,7 @@ Happy, Surprise, HappySurprise """ -def create_merge_app(model_holder: ModelHolder) -> gr.Blocks: +def create_merge_app(model_holder: TTSModelHolder) -> gr.Blocks: model_names = model_holder.model_names if len(model_names) == 0: logger.error( From afff154da4dd9eb9870e895cb143b2a2f67f19a2 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sun, 10 Mar 2024 04:18:39 +0000 Subject: [PATCH 56/64] Add: test code for style-bert-vits2 as a library By executing "hatch run test:test", you can check whether the test passes in all Python 3.9 to 3.12 environments. --- pyproject.toml | 10 ++++++--- style_bert_vits2/tts_model.py | 12 ++++++++--- tests/.gitignore | 1 + tests/test_main.py | 39 +++++++++++++++++++++++++++++++++++ 4 files changed, 56 insertions(+), 6 deletions(-) create mode 100644 tests/.gitignore create mode 100644 tests/test_main.py diff --git a/pyproject.toml b/pyproject.toml index 4db5fe9..5995fc9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -52,24 +52,28 @@ Source = "https://github.com/litagin02/Style-Bert-VITS2" [tool.hatch.version] path = "style_bert_vits2/constants.py" -[tool.hatch.envs.default] +[tool.hatch.envs.test] dependencies = [ "coverage[toml]>=6.5", "pytest", ] -[tool.hatch.envs.default.scripts] +[tool.hatch.envs.test.scripts] +# Usage: `hatch run test:test` test = "pytest {args:tests}" +# Usage: `hatch run test:coverage` test-cov = "coverage run -m pytest {args:tests}" +# Usage: `hatch run test:cov-report` cov-report = [ "- coverage combine", "coverage report", ] +# Usage: `hatch run test:cov` cov = [ "test-cov", "cov-report", ] -[[tool.hatch.envs.all.matrix]] +[[tool.hatch.envs.test.matrix]] python = ["3.9", "3.10", "3.11", "3.12"] [tool.coverage.run] diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index 109ef0e..fe4dcb5 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -1,6 +1,6 @@ import warnings from pathlib import Path -from typing import Any, Optional, Union +from typing import Any, Optional, Union, TypedDict import gradio as gr import numpy as np @@ -263,6 +263,13 @@ class TTSModel: return (self.hyper_parameters.data.sampling_rate, audio) +class TTSModelInfo(TypedDict): + name: str + files: list[str] + styles: list[str] + speakers: list[str] + + class TTSModelHolder: """ Style-Bert-Vits2 の音声合成モデルを管理するクラス。 @@ -297,8 +304,7 @@ class TTSModelHolder: self.model_files_dict: dict[str, list[Path]] = {} self.current_model: Optional[TTSModel] = None self.model_names: list[str] = [] - self.models: list[TTSModel] = [] - self.models_info: list[dict[str, Union[str, list[str]]]] = [] + self.models_info: list[TTSModelInfo] = [] self.refresh() diff --git a/tests/.gitignore b/tests/.gitignore new file mode 100644 index 0000000..697e56f --- /dev/null +++ b/tests/.gitignore @@ -0,0 +1 @@ +*.wav \ No newline at end of file diff --git a/tests/test_main.py b/tests/test_main.py new file mode 100644 index 0000000..597c232 --- /dev/null +++ b/tests/test_main.py @@ -0,0 +1,39 @@ +import pytest +from scipy.io import wavfile + +from style_bert_vits2.constants import BASE_DIR +from style_bert_vits2.tts_model import TTSModelHolder + + +def synthesize(device: str = 'cpu'): + + # モデル一覧を取得 + model_holder = TTSModelHolder(BASE_DIR / 'model_assets', device) + + # モデルが存在する場合、音声合成を実行 + if len(model_holder.models_info) > 0: + + # jvnv-F1-jp モデルを探す + for model_info in model_holder.models_info: + if model_info['name'] == 'jvnv-F1-jp': + + # 音声合成を実行 + model = model_holder.get_model(model_info['name'], model_info['files'][0]) + model.load() + sample_rate, audio_data = model.infer("あらゆる現実を、すべて自分のほうへねじ曲げたのだ。") + + # 音声データを保存 + with open(BASE_DIR / 'tests/test.wav', mode='wb') as f: + wavfile.write(f, sample_rate, audio_data) + else: + pytest.skip("音声合成モデルが見つかりませんでした。") + + +def test_synthesize_cpu(): + synthesize(device='cpu') + assert (BASE_DIR / 'tests/test.wav').exists() + + +def test_synthesize_cuda(): + synthesize(device='cuda') + assert (BASE_DIR / 'tests/test.wav').exists() From 84b1dbe1b5cd91e70286cb006df0b79dfd760838 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sun, 10 Mar 2024 10:48:21 +0000 Subject: [PATCH 57/64] Fix: problem with test failures Style-Bert-VITS2 has been reported to not work with some PyTorch 2.2 series, but Python 3.12 is only supported in Torch >= 2.2, so Python 3.12 support is not provided for the time being --- pyproject.toml | 5 ++--- style_bert_vits2/models/utils/__init__.py | 9 ++++++--- 2 files changed, 8 insertions(+), 6 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 5995fc9..9e68d6d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -19,7 +19,6 @@ classifiers = [ "Programming Language :: Python :: 3.9", "Programming Language :: Python :: 3.10", "Programming Language :: Python :: 3.11", - "Programming Language :: Python :: 3.12", "Programming Language :: Python :: Implementation :: CPython", ] dependencies = [ @@ -37,7 +36,7 @@ dependencies = [ 'pydantic', 'pyopenjtalk-dict', 'pypinyin', - 'pyworld', + # 'pyworld', 'safetensors', 'scipy', 'torch>=2.1,<2.2', @@ -74,7 +73,7 @@ cov = [ ] [[tool.hatch.envs.test.matrix]] -python = ["3.9", "3.10", "3.11", "3.12"] +python = ["3.9", "3.10", "3.11"] [tool.coverage.run] source_pkgs = ["style_bert_vits2", "tests"] diff --git a/style_bert_vits2/models/utils/__init__.py b/style_bert_vits2/models/utils/__init__.py index 51e19d9..e0d922f 100644 --- a/style_bert_vits2/models/utils/__init__.py +++ b/style_bert_vits2/models/utils/__init__.py @@ -4,24 +4,27 @@ import os import re import subprocess from pathlib import Path -from typing import Any, Optional, Union +from typing import Any, Optional, Union, TYPE_CHECKING import numpy as np import torch from numpy.typing import NDArray from scipy.io.wavfile import read -from torch.utils.tensorboard import SummaryWriter from style_bert_vits2.logging import logger from style_bert_vits2.models.utils import checkpoints # type: ignore from style_bert_vits2.models.utils import safetensors # type: ignore +if TYPE_CHECKING: + # tensorboard はライブラリとしてインストールされている場合は依存関係に含まれないため、型チェック時のみインポートする + from torch.utils.tensorboard import SummaryWriter + __is_matplotlib_imported = False def summarize( - writer: SummaryWriter, + writer: "SummaryWriter", global_step: int, scalars: dict[str, float] = {}, histograms: dict[str, Any] = {}, From cdc47a98cebef125b7a33d8b24a5970887f6d9a8 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sun, 10 Mar 2024 11:10:06 +0000 Subject: [PATCH 58/64] Improve: test code --- .gitignore | 1 + tests/test_main.py | 46 +++++++++++++++++++++++++++++++--------------- 2 files changed, 32 insertions(+), 15 deletions(-) diff --git a/.gitignore b/.gitignore index b8a19a4..3160231 100644 --- a/.gitignore +++ b/.gitignore @@ -4,6 +4,7 @@ __pycache__/ venv/ .venv/ dist/ +.coverage .ipynb_checkpoints/ /*.yml diff --git a/tests/test_main.py b/tests/test_main.py index 597c232..c5d786f 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -1,39 +1,55 @@ import pytest from scipy.io import wavfile -from style_bert_vits2.constants import BASE_DIR +from style_bert_vits2.constants import BASE_DIR, Languages from style_bert_vits2.tts_model import TTSModelHolder def synthesize(device: str = 'cpu'): - # モデル一覧を取得 + # 音声合成モデルが配置されていれば、音声合成を実行 model_holder = TTSModelHolder(BASE_DIR / 'model_assets', device) - - # モデルが存在する場合、音声合成を実行 if len(model_holder.models_info) > 0: - # jvnv-F1-jp モデルを探す + # jvnv-F2-jp モデルを探す for model_info in model_holder.models_info: - if model_info['name'] == 'jvnv-F1-jp': + if model_info['name'] == 'jvnv-F2-jp': + # すべてのスタイルに対して音声合成を実行 + for style in model_info['styles']: - # 音声合成を実行 - model = model_holder.get_model(model_info['name'], model_info['files'][0]) - model.load() - sample_rate, audio_data = model.infer("あらゆる現実を、すべて自分のほうへねじ曲げたのだ。") + # 音声合成を実行 + model = model_holder.get_model(model_info['name'], model_info['files'][0]) + model.load() + sample_rate, audio_data = model.infer( + "あらゆる現実を、すべて自分のほうへねじ曲げたのだ。", + # 言語 (JP, EN, ZH / JP-Extra モデルの場合は JP のみ) + language = Languages.JP, + # 話者 ID (音声合成モデルに複数の話者が含まれる場合のみ必須、単一話者のみの場合は 0) + speaker_id = 0, + # 感情表現の強さ (0.0 〜 1.0) + sdp_ratio = 0.4, + # スタイル (Neutral, Happy など) + style = style, + # スタイルの強さ (0.0 〜 100.0) + style_weight = 6.0, + ) - # 音声データを保存 - with open(BASE_DIR / 'tests/test.wav', mode='wb') as f: - wavfile.write(f, sample_rate, audio_data) + # 音声データを保存 + (BASE_DIR / 'tests/wavs').mkdir(exist_ok=True, parents=True) + wav_file_path = BASE_DIR / f'tests/wavs/{style}.wav' + with open(wav_file_path, 'wb') as f: + wavfile.write(f, sample_rate, audio_data) + + # 音声データが保存されたことを確認 + assert wav_file_path.exists() + # wav_file_path.unlink() else: pytest.skip("音声合成モデルが見つかりませんでした。") def test_synthesize_cpu(): synthesize(device='cpu') - assert (BASE_DIR / 'tests/test.wav').exists() def test_synthesize_cuda(): synthesize(device='cuda') - assert (BASE_DIR / 'tests/test.wav').exists() From 00bf496325ebe3ddaf2ed673d7a82db34e7a2d43 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sun, 10 Mar 2024 13:51:56 +0000 Subject: [PATCH 59/64] Add: VSCode settings Enabling type checking with Pylance. --- .gitignore | 2 -- .vscode/extensions.json | 6 ++++++ .vscode/settings.json | 22 ++++++++++++++++++++++ 3 files changed, 28 insertions(+), 2 deletions(-) create mode 100644 .vscode/extensions.json create mode 100644 .vscode/settings.json diff --git a/.gitignore b/.gitignore index 3160231..ae22e64 100644 --- a/.gitignore +++ b/.gitignore @@ -1,5 +1,3 @@ -.vscode/ - __pycache__/ venv/ .venv/ diff --git a/.vscode/extensions.json b/.vscode/extensions.json new file mode 100644 index 0000000..7478fbd --- /dev/null +++ b/.vscode/extensions.json @@ -0,0 +1,6 @@ +{ + "recommendations": [ + "ms-python.python", + "ms-python.vscode-pylance" + ] +} \ No newline at end of file diff --git a/.vscode/settings.json b/.vscode/settings.json new file mode 100644 index 0000000..342c024 --- /dev/null +++ b/.vscode/settings.json @@ -0,0 +1,22 @@ +{ + // Pylance の Type Checking を有効化 + "python.languageServer": "Pylance", + "python.analysis.typeCheckingMode": "strict", + // Pylance の Type Checking のうち、いくつかのエラー報告を抑制する + "python.analysis.diagnosticSeverityOverrides": { + "reportConstantRedefinition": "none", + "reportGeneralTypeIssues": "warning", + "reportMissingParameterType": "warning", + "reportMissingTypeStubs": "none", + "reportPrivateImportUsage": "none", + "reportPrivateUsage": "warning", + "reportShadowedImports": "none", + "reportUnnecessaryComparison": "none", + "reportUnknownArgumentType": "none", + "reportUnknownMemberType": "none", + "reportUnknownParameterType": "warning", + "reportUnknownVariableType": "none", + "reportUnusedFunction": "none", + "reportUnusedVariable": "information", + }, +} \ No newline at end of file From 9c233630efbe48840dbff1e4d4774cf37413de42 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sun, 10 Mar 2024 14:18:57 +0000 Subject: [PATCH 60/64] Refactor: don't keep models in bert_feature module for each language BERT models and tokenizers are already stored and managed in the bert_models module and should not be stored here. In addition, since there may be situations where the user would like to use cpu instead of mps for inference when using it as a library, the automatic switching process to mps was removed. --- style_bert_vits2/nlp/chinese/bert_feature.py | 20 +++-------------- style_bert_vits2/nlp/english/bert_feature.py | 20 +++-------------- style_bert_vits2/nlp/japanese/bert_feature.py | 22 ++++--------------- 3 files changed, 10 insertions(+), 52 deletions(-) diff --git a/style_bert_vits2/nlp/chinese/bert_feature.py b/style_bert_vits2/nlp/chinese/bert_feature.py index f448b30..adc0788 100644 --- a/style_bert_vits2/nlp/chinese/bert_feature.py +++ b/style_bert_vits2/nlp/chinese/bert_feature.py @@ -1,16 +1,11 @@ -import sys from typing import Optional import torch -from transformers import PreTrainedModel from style_bert_vits2.constants import Languages from style_bert_vits2.nlp import bert_models -__models: dict[str, PreTrainedModel] = {} - - def extract_bert_feature( text: str, word2ph: list[int], @@ -32,18 +27,9 @@ def extract_bert_feature( torch.Tensor: BERT の特徴量 """ - if ( - sys.platform == "darwin" - and torch.backends.mps.is_available() - and device == "cpu" - ): - device = "mps" - if not device: - device = "cuda" if device == "cuda" and not torch.cuda.is_available(): device = "cpu" - if device not in __models.keys(): - __models[device] = bert_models.load_model(Languages.ZH).to(device) # type: ignore + model = bert_models.load_model(Languages.ZH).to(device) # type: ignore style_res_mean = None with torch.no_grad(): @@ -51,13 +37,13 @@ def extract_bert_feature( inputs = tokenizer(text, return_tensors="pt") for i in inputs: inputs[i] = inputs[i].to(device) # type: ignore - res = __models[device](**inputs, output_hidden_states=True) + res = model(**inputs, output_hidden_states=True) res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu() if assist_text: style_inputs = tokenizer(assist_text, return_tensors="pt") for i in style_inputs: style_inputs[i] = style_inputs[i].to(device) # type: ignore - style_res = __models[device](**style_inputs, output_hidden_states=True) + style_res = model(**style_inputs, output_hidden_states=True) style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu() style_res_mean = style_res.mean(0) diff --git a/style_bert_vits2/nlp/english/bert_feature.py b/style_bert_vits2/nlp/english/bert_feature.py index 27fd501..7692c9e 100644 --- a/style_bert_vits2/nlp/english/bert_feature.py +++ b/style_bert_vits2/nlp/english/bert_feature.py @@ -1,16 +1,11 @@ -import sys from typing import Optional import torch -from transformers import PreTrainedModel from style_bert_vits2.constants import Languages from style_bert_vits2.nlp import bert_models -__models: dict[str, PreTrainedModel] = {} - - def extract_bert_feature( text: str, word2ph: list[int], @@ -32,18 +27,9 @@ def extract_bert_feature( torch.Tensor: BERT の特徴量 """ - if ( - sys.platform == "darwin" - and torch.backends.mps.is_available() - and device == "cpu" - ): - device = "mps" - if not device: - device = "cuda" if device == "cuda" and not torch.cuda.is_available(): device = "cpu" - if device not in __models.keys(): - __models[device] = bert_models.load_model(Languages.EN).to(device) # type: ignore + model = bert_models.load_model(Languages.EN).to(device) # type: ignore style_res_mean = None with torch.no_grad(): @@ -51,13 +37,13 @@ def extract_bert_feature( inputs = tokenizer(text, return_tensors="pt") for i in inputs: inputs[i] = inputs[i].to(device) # type: ignore - res = __models[device](**inputs, output_hidden_states=True) + res = model(**inputs, output_hidden_states=True) res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu() if assist_text: style_inputs = tokenizer(assist_text, return_tensors="pt") for i in style_inputs: style_inputs[i] = style_inputs[i].to(device) # type: ignore - style_res = __models[device](**style_inputs, output_hidden_states=True) + style_res = model(**style_inputs, output_hidden_states=True) style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu() style_res_mean = style_res.mean(0) diff --git a/style_bert_vits2/nlp/japanese/bert_feature.py b/style_bert_vits2/nlp/japanese/bert_feature.py index 0d70014..d6e214b 100644 --- a/style_bert_vits2/nlp/japanese/bert_feature.py +++ b/style_bert_vits2/nlp/japanese/bert_feature.py @@ -1,17 +1,12 @@ -import sys from typing import Optional import torch -from transformers import PreTrainedModel from style_bert_vits2.constants import Languages from style_bert_vits2.nlp import bert_models from style_bert_vits2.nlp.japanese.g2p import text_to_sep_kata -__models: dict[str, PreTrainedModel] = {} - - def extract_bert_feature( text: str, word2ph: list[int], @@ -36,21 +31,12 @@ def extract_bert_feature( # 各単語が何文字かを作る `word2ph` を使う必要があるので、読めない文字は必ず無視する # でないと `word2ph` の結果とテキストの文字数結果が整合性が取れない text = "".join(text_to_sep_kata(text, raise_yomi_error=False)[0]) - if assist_text: assist_text = "".join(text_to_sep_kata(assist_text, raise_yomi_error=False)[0]) - if ( - sys.platform == "darwin" - and torch.backends.mps.is_available() - and device == "cpu" - ): - device = "mps" - if not device: - device = "cuda" + if device == "cuda" and not torch.cuda.is_available(): device = "cpu" - if device not in __models.keys(): - __models[device] = bert_models.load_model(Languages.JP).to(device) # type: ignore + model = bert_models.load_model(Languages.JP).to(device) # type: ignore style_res_mean = None with torch.no_grad(): @@ -58,13 +44,13 @@ def extract_bert_feature( inputs = tokenizer(text, return_tensors="pt") for i in inputs: inputs[i] = inputs[i].to(device) # type: ignore - res = __models[device](**inputs, output_hidden_states=True) + res = model(**inputs, output_hidden_states=True) res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu() if assist_text: style_inputs = tokenizer(assist_text, return_tensors="pt") for i in style_inputs: style_inputs[i] = style_inputs[i].to(device) # type: ignore - style_res = __models[device](**style_inputs, output_hidden_states=True) + style_res = model(**style_inputs, output_hidden_states=True) style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu() style_res_mean = style_res.mean(0) From 733a9d838d10f3accdc4151c1e14988538ff1ff9 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sun, 10 Mar 2024 15:38:08 +0000 Subject: [PATCH 61/64] Fix: clearly include Pydantic v2 in the dependencies The Pydantic models in the library are written for Pydantic v2 and will not work with Pydantic v1. --- pyproject.toml | 2 +- requirements.txt | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 9e68d6d..0a1bd99 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -33,7 +33,7 @@ dependencies = [ 'numba', 'numpy', 'pyannote.audio>=3.1.0', - 'pydantic', + 'pydantic>=2.0', 'pyopenjtalk-dict', 'pypinyin', # 'pyworld', diff --git a/requirements.txt b/requirements.txt index 669af2f..75dbdcd 100644 --- a/requirements.txt +++ b/requirements.txt @@ -14,7 +14,7 @@ numba numpy psutil pyannote.audio>=3.1.0 -pydantic +pydantic>=2.0 pyloudnorm # pyopenjtalk-prebuilt # Should be manually uninstalled pyopenjtalk-dict From be265d42ed1a946bbb4bcc12d76fe9ab01da59b2 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sun, 10 Mar 2024 16:51:58 +0000 Subject: [PATCH 62/64] Improve: switch pyworld to pyworld-prebuilt and enable it by default Prebuilt wheels for almost every OS/architecture (except musl) are now available on PyPI, eliminating the need for a build environment. ref: https://pypi.org/project/pyworld-prebuilt/ --- pyproject.toml | 2 +- requirements.txt | 2 +- style_bert_vits2/voice.py | 30 ++++++++++++++++++------------ webui/inference.py | 2 -- 4 files changed, 20 insertions(+), 16 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 0a1bd99..870c3b2 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -36,7 +36,7 @@ dependencies = [ 'pydantic>=2.0', 'pyopenjtalk-dict', 'pypinyin', - # 'pyworld', + 'pyworld-prebuilt', 'safetensors', 'scipy', 'torch>=2.1,<2.2', diff --git a/requirements.txt b/requirements.txt index 75dbdcd..bf47ec0 100644 --- a/requirements.txt +++ b/requirements.txt @@ -19,7 +19,7 @@ pyloudnorm # pyopenjtalk-prebuilt # Should be manually uninstalled pyopenjtalk-dict pypinyin -# pyworld # Not supported on Windows without Cython... +pyworld-prebuilt PyYAML requests safetensors diff --git a/style_bert_vits2/voice.py b/style_bert_vits2/voice.py index f1cfcac..ed7843f 100644 --- a/style_bert_vits2/voice.py +++ b/style_bert_vits2/voice.py @@ -1,6 +1,7 @@ from typing import Any import numpy as np +import pyworld from numpy.typing import NDArray @@ -10,29 +11,34 @@ def adjust_voice( pitch_scale: float = 1.0, intonation_scale: float = 1.0, ) -> tuple[int, NDArray[Any]]: + """ + 音声のピッチとイントネーションを調整する。 + 変更すると若干音質が劣化するので、どちらも初期値のままならそのまま返す。 + + Args: + fs (int): 音声のサンプリング周波数 + wave (NDArray[Any]): 音声データ + pitch_scale (float, optional): ピッチの高さ. Defaults to 1.0. + intonation_scale (float, optional): イントネーションの高さ. Defaults to 1.0. + + Returns: + tuple[int, NDArray[Any]]: 調整後の音声データのサンプリング周波数と音声データ + """ if pitch_scale == 1.0 and intonation_scale == 1.0: # 初期値の場合は、音質劣化を避けるためにそのまま返す return fs, wave - try: - import pyworld - except ImportError: - raise ImportError( - "pyworld is not installed. Please install it by `pip install pyworld`" - ) - # pyworld で f0 を加工して合成 # pyworld よりもよいのがあるかもしれないが…… - ## pyworld は Cython で書かれているが、スタブファイルがないため型補完が全く効かない… wave = wave.astype(np.double) # 質が高そうだしとりあえずharvestにしておく - f0, t = pyworld.harvest(wave, fs) # type: ignore + f0, t = pyworld.harvest(wave, fs) - sp = pyworld.cheaptrick(wave, f0, t, fs) # type: ignore - ap = pyworld.d4c(wave, f0, t, fs) # type: ignore + sp = pyworld.cheaptrick(wave, f0, t, fs) + ap = pyworld.d4c(wave, f0, t, fs) non_zero_f0 = [f for f in f0 if f != 0] f0_mean = sum(non_zero_f0) / len(non_zero_f0) @@ -42,5 +48,5 @@ def adjust_voice( continue f0[i] = pitch_scale * f0_mean + intonation_scale * (f - f0_mean) - wave = pyworld.synthesize(f0, sp, ap, fs) # type: ignore + wave = pyworld.synthesize(f0, sp, ap, fs) return fs, wave diff --git a/webui/inference.py b/webui/inference.py index e32907c..91baed5 100644 --- a/webui/inference.py +++ b/webui/inference.py @@ -294,7 +294,6 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: value=1, step=0.05, label="音程(1以外では音質劣化)", - visible=False, # pyworldが必要 ) intonation_scale = gr.Slider( minimum=0, @@ -302,7 +301,6 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: value=1, step=0.1, label="抑揚(1以外では音質劣化)", - visible=False, # pyworldが必要 ) line_split = gr.Checkbox( From 859d940916e253da5680f59c1091ced23d41804d Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sun, 10 Mar 2024 19:15:11 +0000 Subject: [PATCH 63/64] Fix: failed to start training --- style_bert_vits2/models/hyper_parameters.py | 31 +++++++++++---------- 1 file changed, 16 insertions(+), 15 deletions(-) diff --git a/style_bert_vits2/models/hyper_parameters.py b/style_bert_vits2/models/hyper_parameters.py index 9dc5afb..53dc7f3 100644 --- a/style_bert_vits2/models/hyper_parameters.py +++ b/style_bert_vits2/models/hyper_parameters.py @@ -39,8 +39,8 @@ class HyperParametersTrain(BaseModel): class HyperParametersData(BaseModel): use_jp_extra: bool = True - training_files: str = "Data/dummy/train.list" - validation_files: str = "Data/dummy/val.list" + training_files: str = "Data/Dummy/train.list" + validation_files: str = "Data/Dummy/val.list" max_wav_value: float = 32768.0 sampling_rate: int = 44100 filter_length: int = 2048 @@ -53,7 +53,7 @@ class HyperParametersData(BaseModel): n_speakers: int = 512 cleaned_text: bool = True spk2id: dict[str, int] = { - "dummy": 0 + "Dummy": 0, } num_styles: int = 1 style2id: dict[str, int] = { @@ -61,6 +61,13 @@ class HyperParametersData(BaseModel): } +class HyperParametersModelSLM(BaseModel): + model: str = "./slm/wavlm-base-plus" + sr: int = 16000 + hidden: int = 768 + nlayers: int = 13 + initial_channel: int = 64 + class HyperParametersModel(BaseModel): use_spk_conditioned_encoder: bool = True use_noise_scaled_mas: bool = True @@ -79,7 +86,7 @@ class HyperParametersModel(BaseModel): resblock_dilation_sizes: list[list[int]] = [ [1, 3, 5], [1, 3, 5], - [1, 3, 5] + [1, 3, 5], ] upsample_rates: list[int] = [8, 8, 2, 2, 2] upsample_initial_channel: int = 512 @@ -87,21 +94,15 @@ class HyperParametersModel(BaseModel): n_layers_q: int = 3 use_spectral_norm: bool = False gin_channels: int = 512 - slm: dict[str, Union[int, str]] = { - "model": "./slm/wavlm-base-plus", - "sr": 16000, - "hidden": 768, - "nlayers": 13, - "initial_channel": 64 - } + slm: HyperParametersModelSLM = HyperParametersModelSLM() class HyperParameters(BaseModel): - model_name: str = 'dummy' + model_name: str = 'Dummy' version: str = "2.0-JP-Extra" - train: HyperParametersTrain - data: HyperParametersData - model: HyperParametersModel + train: HyperParametersTrain = HyperParametersTrain() + data: HyperParametersData = HyperParametersData() + model: HyperParametersModel = HyperParametersModel() # 以下は学習時にのみ動的に設定されるパラメータ (通常 config.json には存在しない) model_dir: Optional[str] = None From 7f02b0f1d5e6452dba368178d95037aa69acfaae Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sun, 10 Mar 2024 19:21:22 +0000 Subject: [PATCH 64/64] Refactor: TTSModelInfo changed from TypedDict to Pydantic model Pydantic models are more robust and properties can be accessed by dots. --- server_editor.py | 4 ++-- style_bert_vits2/tts_model.py | 17 +++++++++-------- tests/test_main.py | 6 +++--- 3 files changed, 14 insertions(+), 13 deletions(-) diff --git a/server_editor.py b/server_editor.py index aa85ef0..c6f856e 100644 --- a/server_editor.py +++ b/server_editor.py @@ -52,7 +52,7 @@ from style_bert_vits2.nlp.japanese.user_dict import ( rewrite_word, update_dict, ) -from style_bert_vits2.tts_model import TTSModelHolder +from style_bert_vits2.tts_model import TTSModelHolder, TTSModelInfo # ---フロントエンド部分に関する処理--- @@ -250,7 +250,7 @@ async def normalize(item: TextRequest): return normalize_text(item.text) -@router.get("/models_info") +@router.get("/models_info", response_model=list[TTSModelInfo]) def models_info(): return model_holder.models_info diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index fe4dcb5..527b4ab 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -1,6 +1,6 @@ import warnings from pathlib import Path -from typing import Any, Optional, Union, TypedDict +from typing import Any, Optional, Union import gradio as gr import numpy as np @@ -8,6 +8,7 @@ import pyannote.audio import torch from gradio.processing_utils import convert_to_16_bit_wav from numpy.typing import NDArray +from pydantic import BaseModel from style_bert_vits2.constants import ( DEFAULT_ASSIST_TEXT_WEIGHT, @@ -263,7 +264,7 @@ class TTSModel: return (self.hyper_parameters.data.sampling_rate, audio) -class TTSModelInfo(TypedDict): +class TTSModelInfo(BaseModel): name: str files: list[str] styles: list[str] @@ -341,12 +342,12 @@ class TTSModelHolder: styles = list(style2id.keys()) spk2id: dict[str, int] = hyper_parameters.data.spk2id speakers = list(spk2id.keys()) - self.models_info.append({ - "name": model_dir.name, - "files": [str(f) for f in model_files], - "styles": styles, - "speakers": speakers, - }) + self.models_info.append(TTSModelInfo( + name = model_dir.name, + files = [str(f) for f in model_files], + styles = styles, + speakers = speakers, + )) def get_model(self, model_name: str, model_path_str: str) -> TTSModel: diff --git a/tests/test_main.py b/tests/test_main.py index c5d786f..0c0fe77 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -13,12 +13,12 @@ def synthesize(device: str = 'cpu'): # jvnv-F2-jp モデルを探す for model_info in model_holder.models_info: - if model_info['name'] == 'jvnv-F2-jp': + if model_info.name == 'jvnv-F2-jp': # すべてのスタイルに対して音声合成を実行 - for style in model_info['styles']: + for style in model_info.styles: # 音声合成を実行 - model = model_holder.get_model(model_info['name'], model_info['files'][0]) + model = model_holder.get_model(model_info.name, model_info.files[0]) model.load() sample_rate, audio_data = model.infer( "あらゆる現実を、すべて自分のほうへねじ曲げたのだ。",