From 81baabbb30c66aaf32b09b216af882aba93ce410 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Mon, 4 Mar 2024 11:46:09 +0900 Subject: [PATCH 01/14] wip --- app.py | 520 +----------------- webui/__init__.py | 16 + webui_dataset.py => webui/dataset.py | 428 +++++++------- webui/inference.py | 504 +++++++++++++++++ webui_merge.py => webui/merge.py | 351 ++++++------ .../style_vectors.py | 295 +++++----- webui_train.py => webui/train.py | 6 +- 7 files changed, 1095 insertions(+), 1025 deletions(-) create mode 100644 webui/__init__.py rename webui_dataset.py => webui/dataset.py (50%) create mode 100644 webui/inference.py rename webui_merge.py => webui/merge.py (65%) rename webui_style_vectors.py => webui/style_vectors.py (65%) rename webui_train.py => webui/train.py (99%) diff --git a/app.py b/app.py index 1514444..7057f45 100644 --- a/app.py +++ b/app.py @@ -1,502 +1,30 @@ -import argparse -import datetime -import json -import os -import sys -from pathlib import Path -from typing import Optional - +import pyopenjtalk import gradio as gr -import torch -import yaml - -from common.constants import ( - DEFAULT_ASSIST_TEXT_WEIGHT, - DEFAULT_LENGTH, - DEFAULT_LINE_SPLIT, - DEFAULT_NOISE, - DEFAULT_NOISEW, - DEFAULT_SDP_RATIO, - DEFAULT_SPLIT_INTERVAL, - DEFAULT_STYLE, - DEFAULT_STYLE_WEIGHT, - GRADIO_THEME, - LATEST_VERSION, - Languages, +from webui import ( + create_dataset_app, + create_train_app, + create_merge_app, + create_style_vectors_app, ) -from common.log import logger -from common.tts_model import ModelHolder -from infer import InvalidToneError -from text.japanese import g2kata_tone, kata_tone2phone_tone, text_normalize +from pathlib import Path -# 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()) - # dataset_root = path_config["dataset_root"] - assets_root = path_config["assets_root"] +pyopenjtalk.unset_user_dict() -languages = [l.value for l in Languages] +setting_json = Path("webui/setting.json") + +with gr.Blocks() as app: + with gr.Tabs(): + with gr.Tab("Hello"): + gr.Markdown("## Hello, Gradio!") + gr.Textbox("input", label="Input Text") + with gr.Tab("Dataset"): + create_dataset_app() + with gr.Tab("Train"): + create_train_app() + with gr.Tab("Merge"): + create_merge_app() + with gr.Tab("Create Style Vectors"): + create_style_vectors_app() -def tts_fn( - model_name, - model_path, - text, - language, - reference_audio_path, - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - line_split, - split_interval, - assist_text, - assist_text_weight, - use_assist_text, - style, - style_weight, - kata_tone_json_str, - use_tone, - speaker, - pitch_scale, - intonation_scale, -): - model_holder.load_model_gr(model_name, model_path) - - wrong_tone_message = "" - kata_tone: Optional[list[tuple[str, int]]] = None - if use_tone and kata_tone_json_str != "": - if language != "JP": - logger.warning("Only Japanese is supported for tone generation.") - wrong_tone_message = "アクセント指定は現在日本語のみ対応しています。" - if line_split: - logger.warning("Tone generation is not supported for line split.") - wrong_tone_message = ( - "アクセント指定は改行で分けて生成を使わない場合のみ対応しています。" - ) - try: - kata_tone = [] - json_data = json.loads(kata_tone_json_str) - # tupleを使うように変換 - for kana, tone in json_data: - assert isinstance(kana, str) and tone in (0, 1), f"{kana}, {tone}" - kata_tone.append((kana, tone)) - except Exception as e: - logger.warning(f"Error occurred when parsing kana_tone_json: {e}") - wrong_tone_message = f"アクセント指定が不正です: {e}" - kata_tone = None - - # toneは実際に音声合成に代入される際のみnot Noneになる - tone: Optional[list[int]] = None - if kata_tone is not None: - phone_tone = kata_tone2phone_tone(kata_tone) - tone = [t for _, t in phone_tone] - - speaker_id = model_holder.current_model.spk2id[speaker] - - start_time = datetime.datetime.now() - - assert model_holder.current_model is not None - - try: - sr, audio = model_holder.current_model.infer( - text=text, - language=language, - reference_audio_path=reference_audio_path, - sdp_ratio=sdp_ratio, - noise=noise_scale, - noisew=noise_scale_w, - length=length_scale, - line_split=line_split, - split_interval=split_interval, - assist_text=assist_text, - assist_text_weight=assist_text_weight, - use_assist_text=use_assist_text, - style=style, - style_weight=style_weight, - given_tone=tone, - sid=speaker_id, - pitch_scale=pitch_scale, - intonation_scale=intonation_scale, - ) - except InvalidToneError as e: - logger.error(f"Tone error: {e}") - return f"Error: アクセント指定が不正です:\n{e}", None, kata_tone_json_str - except ValueError as e: - logger.error(f"Value error: {e}") - return f"Error: {e}", None, kata_tone_json_str - - end_time = datetime.datetime.now() - duration = (end_time - start_time).total_seconds() - - if tone is None and language == "JP": - # アクセント指定に使えるようにアクセント情報を返す - norm_text = text_normalize(text) - kata_tone = g2kata_tone(norm_text) - kata_tone_json_str = json.dumps(kata_tone, ensure_ascii=False) - elif tone is None: - kata_tone_json_str = "" - message = f"Success, time: {duration} seconds." - if wrong_tone_message != "": - message = wrong_tone_message + "\n" + message - return message, (sr, audio), kata_tone_json_str - - -initial_text = "こんにちは、初めまして。あなたの名前はなんていうの?" - -examples = [ - [initial_text, "JP"], - [ - """あなたがそんなこと言うなんて、私はとっても嬉しい。 -あなたがそんなこと言うなんて、私はとっても怒ってる。 -あなたがそんなこと言うなんて、私はとっても驚いてる。 -あなたがそんなこと言うなんて、私はとっても辛い。""", - "JP", - ], - [ # ChatGPTに考えてもらった告白セリフ - """私、ずっと前からあなたのことを見てきました。あなたの笑顔、優しさ、強さに、心惹かれていたんです。 -友達として過ごす中で、あなたのことがだんだんと特別な存在になっていくのがわかりました。 -えっと、私、あなたのことが好きです!もしよければ、私と付き合ってくれませんか?""", - "JP", - ], - [ # 夏目漱石『吾輩は猫である』 - """吾輩は猫である。名前はまだ無い。 -どこで生れたかとんと見当がつかぬ。なんでも薄暗いじめじめした所でニャーニャー泣いていた事だけは記憶している。 -吾輩はここで初めて人間というものを見た。しかもあとで聞くと、それは書生という、人間中で一番獰悪な種族であったそうだ。 -この書生というのは時々我々を捕まえて煮て食うという話である。""", - "JP", - ], - [ # 梶井基次郎『桜の樹の下には』 - """桜の樹の下には屍体が埋まっている!これは信じていいことなんだよ。 -何故って、桜の花があんなにも見事に咲くなんて信じられないことじゃないか。俺はあの美しさが信じられないので、このにさんにち不安だった。 -しかしいま、やっとわかるときが来た。桜の樹の下には屍体が埋まっている。これは信じていいことだ。""", - "JP", - ], - [ # ChatGPTと考えた、感情を表すセリフ - """やったー!テストで満点取れた!私とっても嬉しいな! -どうして私の意見を無視するの?許せない!ムカつく!あんたなんか死ねばいいのに。 -あはははっ!この漫画めっちゃ笑える、見てよこれ、ふふふ、あはは。 -あなたがいなくなって、私は一人になっちゃって、泣いちゃいそうなほど悲しい。""", - "JP", - ], - [ # 上の丁寧語バージョン - """やりました!テストで満点取れましたよ!私とっても嬉しいです! -どうして私の意見を無視するんですか?許せません!ムカつきます!あんたなんか死んでください。 -あはははっ!この漫画めっちゃ笑えます、見てくださいこれ、ふふふ、あはは。 -あなたがいなくなって、私は一人になっちゃって、泣いちゃいそうなほど悲しいです。""", - "JP", - ], - [ # ChatGPTに考えてもらった音声合成の説明文章 - """音声合成は、機械学習を活用して、テキストから人の声を再現する技術です。この技術は、言語の構造を解析し、それに基づいて音声を生成します。 -この分野の最新の研究成果を使うと、より自然で表現豊かな音声の生成が可能である。深層学習の応用により、感情やアクセントを含む声質の微妙な変化も再現することが出来る。""", - "JP", - ], - [ - "Speech synthesis is the artificial production of human speech. A computer system used for this purpose is called a speech synthesizer, and can be implemented in software or hardware products.", - "EN", - ], - [ - "语音合成是人工制造人类语音。用于此目的的计算机系统称为语音合成器,可以通过软件或硬件产品实现。", - "ZH", - ], -] - -initial_md = f""" -# Style-Bert-VITS2 ver {LATEST_VERSION} 音声合成 - -- Ver 2.3で追加されたエディターのほうが実際に読み上げさせるには使いやすいかもしれません。`Editor.bat`か`python server_editor.py`で起動できます。 - -- 初期からある[jvnvのモデル](https://huggingface.co/litagin/style_bert_vits2_jvnv)は、[JVNVコーパス(言語音声と非言語音声を持つ日本語感情音声コーパス)](https://sites.google.com/site/shinnosuketakamichi/research-topics/jvnv_corpus)で学習されたモデルです。ライセンスは[CC BY-SA 4.0](https://creativecommons.org/licenses/by-sa/4.0/deed.ja)です。 -""" - -how_to_md = """ -下のように`model_assets`ディレクトリの中にモデルファイルたちを置いてください。 -``` -model_assets -├── your_model -│ ├── config.json -│ ├── your_model_file1.safetensors -│ ├── your_model_file2.safetensors -│ ├── ... -│ └── style_vectors.npy -└── another_model - ├── ... -``` -各モデルにはファイルたちが必要です: -- `config.json`:学習時の設定ファイル -- `*.safetensors`:学習済みモデルファイル(1つ以上が必要、複数可) -- `style_vectors.npy`:スタイルベクトルファイル - -上2つは`Train.bat`による学習で自動的に正しい位置に保存されます。`style_vectors.npy`は`Style.bat`を実行して指示に従って生成してください。 -""" - -style_md = f""" -- プリセットまたは音声ファイルから読み上げの声音・感情・スタイルのようなものを制御できます。 -- デフォルトの{DEFAULT_STYLE}でも、十分に読み上げる文に応じた感情で感情豊かに読み上げられます。このスタイル制御は、それを重み付きで上書きするような感じです。 -- 強さを大きくしすぎると発音が変になったり声にならなかったりと崩壊することがあります。 -- どのくらいに強さがいいかはモデルやスタイルによって異なるようです。 -- 音声ファイルを入力する場合は、学習データと似た声音の話者(特に同じ性別)でないとよい効果が出ないかもしれません。 -""" - - -def make_interactive(): - return gr.update(interactive=True, value="音声合成") - - -def make_non_interactive(): - return gr.update(interactive=False, value="音声合成(モデルをロードしてください)") - - -def gr_util(item): - if item == "プリセットから選ぶ": - return (gr.update(visible=True), gr.Audio(visible=False, value=None)) - else: - return (gr.update(visible=False), gr.update(visible=True)) - - -if __name__ == "__main__": - parser = argparse.ArgumentParser() - parser.add_argument("--cpu", action="store_true", help="Use CPU instead of GPU") - parser.add_argument( - "--dir", "-d", type=str, help="Model directory", default=assets_root - ) - parser.add_argument( - "--share", action="store_true", help="Share this app publicly", default=False - ) - parser.add_argument( - "--server-name", - type=str, - default=None, - help="Server name for Gradio app", - ) - parser.add_argument( - "--no-autolaunch", - action="store_true", - default=False, - help="Do not launch app automatically", - ) - args = parser.parse_args() - model_dir = Path(args.dir) - - if args.cpu: - device = "cpu" - else: - device = "cuda" if torch.cuda.is_available() else "cpu" - - model_holder = ModelHolder(model_dir, device) - - model_names = model_holder.model_names - if len(model_names) == 0: - logger.error( - f"モデルが見つかりませんでした。{model_dir}にモデルを置いてください。" - ) - sys.exit(1) - initial_id = 0 - initial_pth_files = model_holder.model_files_dict[model_names[initial_id]] - - with gr.Blocks(theme=GRADIO_THEME) as app: - gr.Markdown(initial_md) - with gr.Accordion(label="使い方", open=False): - gr.Markdown(how_to_md) - with gr.Row(): - with gr.Column(): - with gr.Row(): - with gr.Column(scale=3): - model_name = gr.Dropdown( - label="モデル一覧", - choices=model_names, - value=model_names[initial_id], - ) - model_path = gr.Dropdown( - label="モデルファイル", - choices=initial_pth_files, - value=initial_pth_files[0], - ) - refresh_button = gr.Button("更新", scale=1, visible=True) - load_button = gr.Button("ロード", scale=1, variant="primary") - text_input = gr.TextArea(label="テキスト", value=initial_text) - pitch_scale = gr.Slider( - minimum=0.8, - maximum=1.5, - value=1, - step=0.05, - label="音程(1以外では音質劣化)", - visible=False, # pyworldが必要 - ) - intonation_scale = gr.Slider( - minimum=0, - maximum=2, - value=1, - step=0.1, - label="抑揚(1以外では音質劣化)", - visible=False, # pyworldが必要 - ) - - line_split = gr.Checkbox( - label="改行で分けて生成(分けたほうが感情が乗ります)", - value=DEFAULT_LINE_SPLIT, - ) - split_interval = gr.Slider( - minimum=0.0, - maximum=2, - value=DEFAULT_SPLIT_INTERVAL, - step=0.1, - label="改行ごとに挟む無音の長さ(秒)", - ) - line_split.change( - lambda x: (gr.Slider(visible=x)), - inputs=[line_split], - outputs=[split_interval], - ) - tone = gr.Textbox( - label="アクセント調整(数値は 0=低 か1=高 のみ)", - info="改行で分けない場合のみ使えます。万能ではありません。", - ) - use_tone = gr.Checkbox(label="アクセント調整を使う", value=False) - use_tone.change( - lambda x: (gr.Checkbox(value=False) if x else gr.Checkbox()), - inputs=[use_tone], - outputs=[line_split], - ) - language = gr.Dropdown(choices=languages, value="JP", label="Language") - speaker = gr.Dropdown(label="話者") - with gr.Accordion(label="詳細設定", open=False): - sdp_ratio = gr.Slider( - minimum=0, - maximum=1, - value=DEFAULT_SDP_RATIO, - step=0.1, - label="SDP Ratio", - ) - noise_scale = gr.Slider( - minimum=0.1, - maximum=2, - value=DEFAULT_NOISE, - step=0.1, - label="Noise", - ) - noise_scale_w = gr.Slider( - minimum=0.1, - maximum=2, - value=DEFAULT_NOISEW, - step=0.1, - label="Noise_W", - ) - length_scale = gr.Slider( - minimum=0.1, - maximum=2, - value=DEFAULT_LENGTH, - step=0.1, - label="Length", - ) - use_assist_text = gr.Checkbox( - label="Assist textを使う", value=False - ) - assist_text = gr.Textbox( - label="Assist text", - placeholder="どうして私の意見を無視するの?許せない、ムカつく!死ねばいいのに。", - info="このテキストの読み上げと似た声音・感情になりやすくなります。ただ抑揚やテンポ等が犠牲になる傾向があります。", - visible=False, - ) - assist_text_weight = gr.Slider( - minimum=0, - maximum=1, - value=DEFAULT_ASSIST_TEXT_WEIGHT, - step=0.1, - label="Assist textの強さ", - visible=False, - ) - use_assist_text.change( - lambda x: (gr.Textbox(visible=x), gr.Slider(visible=x)), - inputs=[use_assist_text], - outputs=[assist_text, assist_text_weight], - ) - with gr.Column(): - with gr.Accordion("スタイルについて詳細", open=False): - gr.Markdown(style_md) - style_mode = gr.Radio( - ["プリセットから選ぶ", "音声ファイルを入力"], - label="スタイルの指定方法", - value="プリセットから選ぶ", - ) - style = gr.Dropdown( - label=f"スタイル({DEFAULT_STYLE}が平均スタイル)", - choices=["モデルをロードしてください"], - value="モデルをロードしてください", - ) - style_weight = gr.Slider( - minimum=0, - maximum=50, - value=DEFAULT_STYLE_WEIGHT, - step=0.1, - label="スタイルの強さ", - ) - ref_audio_path = gr.Audio( - label="参照音声", type="filepath", visible=False - ) - tts_button = gr.Button( - "音声合成(モデルをロードしてください)", - variant="primary", - interactive=False, - ) - text_output = gr.Textbox(label="情報") - audio_output = gr.Audio(label="結果") - with gr.Accordion("テキスト例", open=False): - gr.Examples(examples, inputs=[text_input, language]) - - tts_button.click( - tts_fn, - inputs=[ - model_name, - model_path, - text_input, - language, - ref_audio_path, - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - line_split, - split_interval, - assist_text, - assist_text_weight, - use_assist_text, - style, - style_weight, - tone, - use_tone, - speaker, - pitch_scale, - intonation_scale, - ], - outputs=[text_output, audio_output, tone], - ) - - model_name.change( - model_holder.update_model_files_gr, - inputs=[model_name], - outputs=[model_path], - ) - - model_path.change(make_non_interactive, outputs=[tts_button]) - - refresh_button.click( - model_holder.update_model_names_gr, - outputs=[model_name, model_path, tts_button], - ) - - load_button.click( - model_holder.load_model_gr, - inputs=[model_name, model_path], - outputs=[style, tts_button, speaker], - ) - - style_mode.change( - gr_util, - inputs=[style_mode], - outputs=[style, ref_audio_path], - ) - - app.launch( - inbrowser=not args.no_autolaunch, share=args.share, server_name=args.server_name - ) +app.launch(inbrowser=True) diff --git a/webui/__init__.py b/webui/__init__.py new file mode 100644 index 0000000..4bc8efa --- /dev/null +++ b/webui/__init__.py @@ -0,0 +1,16 @@ +from .dataset import create_dataset_app +from .inference import create_inference_app +from .merge import create_merge_app +from .style_vectors import create_style_vectors_app +from .train import create_train_app + + +class TrainSettings: + def __init__(self, setting_json): + self.setting_json = setting_json + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_value, traceback): + pass diff --git a/webui_dataset.py b/webui/dataset.py similarity index 50% rename from webui_dataset.py rename to webui/dataset.py index fec7a9a..5ed656c 100644 --- a/webui_dataset.py +++ b/webui/dataset.py @@ -1,208 +1,220 @@ -import argparse -import os - -import gradio as gr -import yaml - -from common.constants import GRADIO_THEME -from common.log import logger -from common.subprocess_utils 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()) - dataset_root = path_config["dataset_root"] - # assets_root = path_config["assets_root"] - - -def do_slice( - model_name: str, - min_sec: float, - max_sec: float, - min_silence_dur_ms: int, - input_dir: str, -): - if model_name == "": - return "Error: モデル名を入力してください。" - logger.info("Start slicing...") - cmd = [ - "slice.py", - "--model_name", - model_name, - "--min_sec", - str(min_sec), - "--max_sec", - str(max_sec), - "--min_silence_dur_ms", - str(min_silence_dur_ms), - ] - if input_dir != "": - cmd += ["--input_dir", input_dir] - # onnxの警告が出るので無視する - success, message = run_script_with_log(cmd, ignore_warning=True) - if not success: - return f"Error: {message}" - return "音声のスライスが完了しました。" - - -def do_transcribe( - model_name, whisper_model, compute_type, language, initial_prompt, device -): - if model_name == "": - return "Error: モデル名を入力してください。" - - success, message = run_script_with_log( - [ - "transcribe.py", - "--model_name", - model_name, - "--model", - whisper_model, - "--compute_type", - compute_type, - "--device", - device, - "--language", - language, - "--initial_prompt", - f'"{initial_prompt}"', - ] - ) - if not success: - return f"Error: {message}" - return "音声の文字起こしが完了しました。" - - -initial_md = """ -# 簡易学習用データセット作成ツール - -Style-Bert-VITS2の学習用データセットを作成するためのツールです。以下の2つからなります。 - -- 与えられた音声からちょうどいい長さの発話区間を切り取りスライス -- 音声に対して文字起こし - -このうち両方を使ってもよいし、スライスする必要がない場合は後者のみを使ってもよいです。 - -## 必要なもの - -学習したい音声が入ったwavファイルいくつか。 -合計時間がある程度はあったほうがいいかも、10分とかでも大丈夫だったとの報告あり。単一ファイルでも良いし複数ファイルでもよい。 - -## スライス使い方 -1. `inputs`フォルダにwavファイルをすべて入れる -2. `モデル名`を入力して、設定を必要なら調整して`音声のスライス`ボタンを押す -3. 出来上がった音声ファイルたちは`Data/{モデル名}/raw`に保存される - -## 書き起こし使い方 - -1. 書き起こしたい音声ファイルのあるフォルダを指定(デフォルトは`Data/{モデル名}/raw`なのでスライス後に行う場合は省略してよい) -2. 設定を必要なら調整してボタンを押す -3. 書き起こしファイルは`Data/{モデル名}/esd.list`に保存される - -## 注意 - -- 長すぎる秒数(12-15秒くらいより長い?)のwavファイルは学習に用いられないようです。また短すぎてもあまりよくない可能性もあります。 -- 書き起こしの結果をどれだけ修正すればいいかはデータセットに依存しそうです。 -- 手動で書き起こしをいろいろ修正したり結果を細かく確認したい場合は、[Aivis Dataset](https://github.com/litagin02/Aivis-Dataset)もおすすめします。書き起こし部分もかなり工夫されています。ですがファイル数が多い場合などは、このツールで簡易的に切り出してデータセットを作るだけでも十分という気もしています。 -""" - -with gr.Blocks(theme=GRADIO_THEME) as app: - gr.Markdown(initial_md) - model_name = gr.Textbox( - label="モデル名を入力してください(話者名としても使われます)。" - ) - with gr.Accordion("音声のスライス"): - with gr.Row(): - with gr.Column(): - input_dir = gr.Textbox( - label="入力フォルダ名(デフォルトはinputs)", - placeholder="inputs", - info="下記フォルダにwavファイルを入れておいてください", - ) - min_sec = gr.Slider( - minimum=0, - maximum=10, - value=2, - step=0.5, - label="この秒数未満は切り捨てる", - ) - max_sec = gr.Slider( - minimum=0, - maximum=15, - value=12, - step=0.5, - label="この秒数以上は切り捨てる", - ) - min_silence_dur_ms = gr.Slider( - minimum=0, - maximum=2000, - value=700, - step=100, - label="無音とみなして区切る最小の無音の長さ(ms)", - ) - slice_button = gr.Button("スライスを実行") - result1 = gr.Textbox(label="結果") - with gr.Row(): - with gr.Column(): - whisper_model = gr.Dropdown( - ["tiny", "base", "small", "medium", "large", "large-v2", "large-v3"], - label="Whisperモデル", - value="large-v3", - ) - compute_type = gr.Dropdown( - [ - "int8", - "int8_float32", - "int8_float16", - "int8_bfloat16", - "int16", - "float16", - "bfloat16", - "float32", - ], - label="計算精度", - value="bfloat16", - ) - device = gr.Radio(["cuda", "cpu"], label="デバイス", value="cuda") - language = gr.Dropdown(["ja", "en", "zh"], value="ja", label="言語") - initial_prompt = gr.Textbox( - label="初期プロンプト", - value="こんにちは。元気、ですかー?ふふっ、私は……ちゃんと元気だよ!", - info="このように書き起こしてほしいという例文(句読点の入れ方・笑い方・固有名詞等)", - ) - transcribe_button = gr.Button("音声の文字起こし") - result2 = gr.Textbox(label="結果") - slice_button.click( - do_slice, - inputs=[model_name, min_sec, max_sec, min_silence_dur_ms, input_dir], - outputs=[result1], - ) - transcribe_button.click( - do_transcribe, - inputs=[ - model_name, - whisper_model, - compute_type, - language, - initial_prompt, - device, - ], - outputs=[result2], - ) - -parser = argparse.ArgumentParser() -parser.add_argument( - "--server-name", - type=str, - default=None, - help="Server name for Gradio app", -) -parser.add_argument( - "--no-autolaunch", - action="store_true", - default=False, - help="Do not launch app automatically", -) -args = parser.parse_args() - -app.launch(inbrowser=not args.no_autolaunch, server_name=args.server_name) +import argparse +import os + +import gradio as gr +import yaml + +from common.constants import GRADIO_THEME +from common.log import logger +from common.subprocess_utils 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()) + dataset_root = path_config["dataset_root"] + # assets_root = path_config["assets_root"] + + +def do_slice( + model_name: str, + min_sec: float, + max_sec: float, + min_silence_dur_ms: int, + input_dir: str, +): + if model_name == "": + return "Error: モデル名を入力してください。" + logger.info("Start slicing...") + cmd = [ + "slice.py", + "--model_name", + model_name, + "--min_sec", + str(min_sec), + "--max_sec", + str(max_sec), + "--min_silence_dur_ms", + str(min_silence_dur_ms), + ] + if input_dir != "": + cmd += ["--input_dir", input_dir] + # onnxの警告が出るので無視する + success, message = run_script_with_log(cmd, ignore_warning=True) + if not success: + return f"Error: {message}" + return "音声のスライスが完了しました。" + + +def do_transcribe( + model_name, whisper_model, compute_type, language, initial_prompt, device +): + if model_name == "": + return "Error: モデル名を入力してください。" + + success, message = run_script_with_log( + [ + "transcribe.py", + "--model_name", + model_name, + "--model", + whisper_model, + "--compute_type", + compute_type, + "--device", + device, + "--language", + language, + "--initial_prompt", + f'"{initial_prompt}"', + ] + ) + if not success: + return f"Error: {message}" + return "音声の文字起こしが完了しました。" + + +initial_md = """ +# 簡易学習用データセット作成ツール + +Style-Bert-VITS2の学習用データセットを作成するためのツールです。以下の2つからなります。 + +- 与えられた音声からちょうどいい長さの発話区間を切り取りスライス +- 音声に対して文字起こし + +このうち両方を使ってもよいし、スライスする必要がない場合は後者のみを使ってもよいです。 + +## 必要なもの + +学習したい音声が入ったwavファイルいくつか。 +合計時間がある程度はあったほうがいいかも、10分とかでも大丈夫だったとの報告あり。単一ファイルでも良いし複数ファイルでもよい。 + +## スライス使い方 +1. `inputs`フォルダにwavファイルをすべて入れる +2. `モデル名`を入力して、設定を必要なら調整して`音声のスライス`ボタンを押す +3. 出来上がった音声ファイルたちは`Data/{モデル名}/raw`に保存される + +## 書き起こし使い方 + +1. 書き起こしたい音声ファイルのあるフォルダを指定(デフォルトは`Data/{モデル名}/raw`なのでスライス後に行う場合は省略してよい) +2. 設定を必要なら調整してボタンを押す +3. 書き起こしファイルは`Data/{モデル名}/esd.list`に保存される + +## 注意 + +- 長すぎる秒数(12-15秒くらいより長い?)のwavファイルは学習に用いられないようです。また短すぎてもあまりよくない可能性もあります。 +- 書き起こしの結果をどれだけ修正すればいいかはデータセットに依存しそうです。 +- 手動で書き起こしをいろいろ修正したり結果を細かく確認したい場合は、[Aivis Dataset](https://github.com/litagin02/Aivis-Dataset)もおすすめします。書き起こし部分もかなり工夫されています。ですがファイル数が多い場合などは、このツールで簡易的に切り出してデータセットを作るだけでも十分という気もしています。 +""" + + +def create_dataset_app(): + + with gr.Blocks(theme=GRADIO_THEME) as app: + gr.Markdown(initial_md) + model_name = gr.Textbox( + label="モデル名を入力してください(話者名としても使われます)。" + ) + with gr.Accordion("音声のスライス"): + with gr.Row(): + with gr.Column(): + input_dir = gr.Textbox( + label="入力フォルダ名(デフォルトはinputs)", + placeholder="inputs", + info="下記フォルダにwavファイルを入れておいてください", + ) + min_sec = gr.Slider( + minimum=0, + maximum=10, + value=2, + step=0.5, + label="この秒数未満は切り捨てる", + ) + max_sec = gr.Slider( + minimum=0, + maximum=15, + value=12, + step=0.5, + label="この秒数以上は切り捨てる", + ) + min_silence_dur_ms = gr.Slider( + minimum=0, + maximum=2000, + value=700, + step=100, + label="無音とみなして区切る最小の無音の長さ(ms)", + ) + slice_button = gr.Button("スライスを実行") + result1 = gr.Textbox(label="結果") + with gr.Row(): + with gr.Column(): + whisper_model = gr.Dropdown( + [ + "tiny", + "base", + "small", + "medium", + "large", + "large-v2", + "large-v3", + ], + label="Whisperモデル", + value="large-v3", + ) + compute_type = gr.Dropdown( + [ + "int8", + "int8_float32", + "int8_float16", + "int8_bfloat16", + "int16", + "float16", + "bfloat16", + "float32", + ], + label="計算精度", + value="bfloat16", + ) + device = gr.Radio(["cuda", "cpu"], label="デバイス", value="cuda") + language = gr.Dropdown(["ja", "en", "zh"], value="ja", label="言語") + initial_prompt = gr.Textbox( + label="初期プロンプト", + value="こんにちは。元気、ですかー?ふふっ、私は……ちゃんと元気だよ!", + info="このように書き起こしてほしいという例文(句読点の入れ方・笑い方・固有名詞等)", + ) + transcribe_button = gr.Button("音声の文字起こし") + result2 = gr.Textbox(label="結果") + slice_button.click( + do_slice, + inputs=[model_name, min_sec, max_sec, min_silence_dur_ms, input_dir], + outputs=[result1], + ) + transcribe_button.click( + do_transcribe, + inputs=[ + model_name, + whisper_model, + compute_type, + language, + initial_prompt, + device, + ], + outputs=[result2], + ) + + parser = argparse.ArgumentParser() + parser.add_argument( + "--server-name", + type=str, + default=None, + help="Server name for Gradio app", + ) + parser.add_argument( + "--no-autolaunch", + action="store_true", + default=False, + help="Do not launch app automatically", + ) + args = parser.parse_args() + + # app.launch(inbrowser=not args.no_autolaunch, server_name=args.server_name) + return app diff --git a/webui/inference.py b/webui/inference.py new file mode 100644 index 0000000..663d5a6 --- /dev/null +++ b/webui/inference.py @@ -0,0 +1,504 @@ +import argparse +import datetime +import json +import os +import sys +from pathlib import Path +from typing import Optional + +import gradio as gr +import torch +import yaml + +from common.constants import ( + DEFAULT_ASSIST_TEXT_WEIGHT, + DEFAULT_LENGTH, + DEFAULT_LINE_SPLIT, + DEFAULT_NOISE, + DEFAULT_NOISEW, + DEFAULT_SDP_RATIO, + DEFAULT_SPLIT_INTERVAL, + DEFAULT_STYLE, + DEFAULT_STYLE_WEIGHT, + GRADIO_THEME, + LATEST_VERSION, + Languages, +) +from common.log import logger +from common.tts_model import ModelHolder +from infer import InvalidToneError +from text.japanese import g2kata_tone, kata_tone2phone_tone, text_normalize + +# 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()) + # dataset_root = path_config["dataset_root"] + assets_root = path_config["assets_root"] + +languages = [l.value for l in Languages] + + +def tts_fn( + model_name, + model_path, + text, + language, + reference_audio_path, + sdp_ratio, + noise_scale, + noise_scale_w, + length_scale, + line_split, + split_interval, + assist_text, + assist_text_weight, + use_assist_text, + style, + style_weight, + kata_tone_json_str, + use_tone, + speaker, + pitch_scale, + intonation_scale, +): + model_holder.load_model_gr(model_name, model_path) + + wrong_tone_message = "" + kata_tone: Optional[list[tuple[str, int]]] = None + if use_tone and kata_tone_json_str != "": + if language != "JP": + logger.warning("Only Japanese is supported for tone generation.") + wrong_tone_message = "アクセント指定は現在日本語のみ対応しています。" + if line_split: + logger.warning("Tone generation is not supported for line split.") + wrong_tone_message = ( + "アクセント指定は改行で分けて生成を使わない場合のみ対応しています。" + ) + try: + kata_tone = [] + json_data = json.loads(kata_tone_json_str) + # tupleを使うように変換 + for kana, tone in json_data: + assert isinstance(kana, str) and tone in (0, 1), f"{kana}, {tone}" + kata_tone.append((kana, tone)) + except Exception as e: + logger.warning(f"Error occurred when parsing kana_tone_json: {e}") + wrong_tone_message = f"アクセント指定が不正です: {e}" + kata_tone = None + + # toneは実際に音声合成に代入される際のみnot Noneになる + tone: Optional[list[int]] = None + if kata_tone is not None: + phone_tone = kata_tone2phone_tone(kata_tone) + tone = [t for _, t in phone_tone] + + speaker_id = model_holder.current_model.spk2id[speaker] + + start_time = datetime.datetime.now() + + assert model_holder.current_model is not None + + try: + sr, audio = model_holder.current_model.infer( + text=text, + language=language, + reference_audio_path=reference_audio_path, + sdp_ratio=sdp_ratio, + noise=noise_scale, + noisew=noise_scale_w, + length=length_scale, + line_split=line_split, + split_interval=split_interval, + assist_text=assist_text, + assist_text_weight=assist_text_weight, + use_assist_text=use_assist_text, + style=style, + style_weight=style_weight, + given_tone=tone, + sid=speaker_id, + pitch_scale=pitch_scale, + intonation_scale=intonation_scale, + ) + except InvalidToneError as e: + logger.error(f"Tone error: {e}") + return f"Error: アクセント指定が不正です:\n{e}", None, kata_tone_json_str + except ValueError as e: + logger.error(f"Value error: {e}") + return f"Error: {e}", None, kata_tone_json_str + + end_time = datetime.datetime.now() + duration = (end_time - start_time).total_seconds() + + if tone is None and language == "JP": + # アクセント指定に使えるようにアクセント情報を返す + norm_text = text_normalize(text) + kata_tone = g2kata_tone(norm_text) + kata_tone_json_str = json.dumps(kata_tone, ensure_ascii=False) + elif tone is None: + kata_tone_json_str = "" + message = f"Success, time: {duration} seconds." + if wrong_tone_message != "": + message = wrong_tone_message + "\n" + message + return message, (sr, audio), kata_tone_json_str + + +initial_text = "こんにちは、初めまして。あなたの名前はなんていうの?" + +examples = [ + [initial_text, "JP"], + [ + """あなたがそんなこと言うなんて、私はとっても嬉しい。 +あなたがそんなこと言うなんて、私はとっても怒ってる。 +あなたがそんなこと言うなんて、私はとっても驚いてる。 +あなたがそんなこと言うなんて、私はとっても辛い。""", + "JP", + ], + [ # ChatGPTに考えてもらった告白セリフ + """私、ずっと前からあなたのことを見てきました。あなたの笑顔、優しさ、強さに、心惹かれていたんです。 +友達として過ごす中で、あなたのことがだんだんと特別な存在になっていくのがわかりました。 +えっと、私、あなたのことが好きです!もしよければ、私と付き合ってくれませんか?""", + "JP", + ], + [ # 夏目漱石『吾輩は猫である』 + """吾輩は猫である。名前はまだ無い。 +どこで生れたかとんと見当がつかぬ。なんでも薄暗いじめじめした所でニャーニャー泣いていた事だけは記憶している。 +吾輩はここで初めて人間というものを見た。しかもあとで聞くと、それは書生という、人間中で一番獰悪な種族であったそうだ。 +この書生というのは時々我々を捕まえて煮て食うという話である。""", + "JP", + ], + [ # 梶井基次郎『桜の樹の下には』 + """桜の樹の下には屍体が埋まっている!これは信じていいことなんだよ。 +何故って、桜の花があんなにも見事に咲くなんて信じられないことじゃないか。俺はあの美しさが信じられないので、このにさんにち不安だった。 +しかしいま、やっとわかるときが来た。桜の樹の下には屍体が埋まっている。これは信じていいことだ。""", + "JP", + ], + [ # ChatGPTと考えた、感情を表すセリフ + """やったー!テストで満点取れた!私とっても嬉しいな! +どうして私の意見を無視するの?許せない!ムカつく!あんたなんか死ねばいいのに。 +あはははっ!この漫画めっちゃ笑える、見てよこれ、ふふふ、あはは。 +あなたがいなくなって、私は一人になっちゃって、泣いちゃいそうなほど悲しい。""", + "JP", + ], + [ # 上の丁寧語バージョン + """やりました!テストで満点取れましたよ!私とっても嬉しいです! +どうして私の意見を無視するんですか?許せません!ムカつきます!あんたなんか死んでください。 +あはははっ!この漫画めっちゃ笑えます、見てくださいこれ、ふふふ、あはは。 +あなたがいなくなって、私は一人になっちゃって、泣いちゃいそうなほど悲しいです。""", + "JP", + ], + [ # ChatGPTに考えてもらった音声合成の説明文章 + """音声合成は、機械学習を活用して、テキストから人の声を再現する技術です。この技術は、言語の構造を解析し、それに基づいて音声を生成します。 +この分野の最新の研究成果を使うと、より自然で表現豊かな音声の生成が可能である。深層学習の応用により、感情やアクセントを含む声質の微妙な変化も再現することが出来る。""", + "JP", + ], + [ + "Speech synthesis is the artificial production of human speech. A computer system used for this purpose is called a speech synthesizer, and can be implemented in software or hardware products.", + "EN", + ], + [ + "语音合成是人工制造人类语音。用于此目的的计算机系统称为语音合成器,可以通过软件或硬件产品实现。", + "ZH", + ], +] + +initial_md = f""" +# Style-Bert-VITS2 ver {LATEST_VERSION} 音声合成 + +- Ver 2.3で追加されたエディターのほうが実際に読み上げさせるには使いやすいかもしれません。`Editor.bat`か`python server_editor.py`で起動できます。 + +- 初期からある[jvnvのモデル](https://huggingface.co/litagin/style_bert_vits2_jvnv)は、[JVNVコーパス(言語音声と非言語音声を持つ日本語感情音声コーパス)](https://sites.google.com/site/shinnosuketakamichi/research-topics/jvnv_corpus)で学習されたモデルです。ライセンスは[CC BY-SA 4.0](https://creativecommons.org/licenses/by-sa/4.0/deed.ja)です。 +""" + +how_to_md = """ +下のように`model_assets`ディレクトリの中にモデルファイルたちを置いてください。 +``` +model_assets +├── your_model +│ ├── config.json +│ ├── your_model_file1.safetensors +│ ├── your_model_file2.safetensors +│ ├── ... +│ └── style_vectors.npy +└── another_model + ├── ... +``` +各モデルにはファイルたちが必要です: +- `config.json`:学習時の設定ファイル +- `*.safetensors`:学習済みモデルファイル(1つ以上が必要、複数可) +- `style_vectors.npy`:スタイルベクトルファイル + +上2つは`Train.bat`による学習で自動的に正しい位置に保存されます。`style_vectors.npy`は`Style.bat`を実行して指示に従って生成してください。 +""" + +style_md = f""" +- プリセットまたは音声ファイルから読み上げの声音・感情・スタイルのようなものを制御できます。 +- デフォルトの{DEFAULT_STYLE}でも、十分に読み上げる文に応じた感情で感情豊かに読み上げられます。このスタイル制御は、それを重み付きで上書きするような感じです。 +- 強さを大きくしすぎると発音が変になったり声にならなかったりと崩壊することがあります。 +- どのくらいに強さがいいかはモデルやスタイルによって異なるようです。 +- 音声ファイルを入力する場合は、学習データと似た声音の話者(特に同じ性別)でないとよい効果が出ないかもしれません。 +""" + + +def make_interactive(): + return gr.update(interactive=True, value="音声合成") + + +def make_non_interactive(): + return gr.update(interactive=False, value="音声合成(モデルをロードしてください)") + + +def gr_util(item): + if item == "プリセットから選ぶ": + return (gr.update(visible=True), gr.Audio(visible=False, value=None)) + else: + return (gr.update(visible=False), gr.update(visible=True)) + + +def create_inference_app(): + parser = argparse.ArgumentParser() + parser.add_argument("--cpu", action="store_true", help="Use CPU instead of GPU") + parser.add_argument( + "--dir", "-d", type=str, help="Model directory", default=assets_root + ) + parser.add_argument( + "--share", action="store_true", help="Share this app publicly", default=False + ) + parser.add_argument( + "--server-name", + type=str, + default=None, + help="Server name for Gradio app", + ) + parser.add_argument( + "--no-autolaunch", + action="store_true", + default=False, + help="Do not launch app automatically", + ) + args = parser.parse_args() + model_dir = Path(args.dir) + + if args.cpu: + device = "cpu" + else: + device = "cuda" if torch.cuda.is_available() else "cpu" + + model_holder = ModelHolder(model_dir, device) + + model_names = model_holder.model_names + if len(model_names) == 0: + logger.error( + f"モデルが見つかりませんでした。{model_dir}にモデルを置いてください。" + ) + sys.exit(1) + initial_id = 0 + initial_pth_files = model_holder.model_files_dict[model_names[initial_id]] + + with gr.Blocks(theme=GRADIO_THEME) as app: + gr.Markdown(initial_md) + with gr.Accordion(label="使い方", open=False): + gr.Markdown(how_to_md) + with gr.Row(): + with gr.Column(): + with gr.Row(): + with gr.Column(scale=3): + model_name = gr.Dropdown( + label="モデル一覧", + choices=model_names, + value=model_names[initial_id], + ) + model_path = gr.Dropdown( + label="モデルファイル", + choices=initial_pth_files, + value=initial_pth_files[0], + ) + refresh_button = gr.Button("更新", scale=1, visible=True) + load_button = gr.Button("ロード", scale=1, variant="primary") + text_input = gr.TextArea(label="テキスト", value=initial_text) + pitch_scale = gr.Slider( + minimum=0.8, + maximum=1.5, + value=1, + step=0.05, + label="音程(1以外では音質劣化)", + visible=False, # pyworldが必要 + ) + intonation_scale = gr.Slider( + minimum=0, + maximum=2, + value=1, + step=0.1, + label="抑揚(1以外では音質劣化)", + visible=False, # pyworldが必要 + ) + + line_split = gr.Checkbox( + label="改行で分けて生成(分けたほうが感情が乗ります)", + value=DEFAULT_LINE_SPLIT, + ) + split_interval = gr.Slider( + minimum=0.0, + maximum=2, + value=DEFAULT_SPLIT_INTERVAL, + step=0.1, + label="改行ごとに挟む無音の長さ(秒)", + ) + line_split.change( + lambda x: (gr.Slider(visible=x)), + inputs=[line_split], + outputs=[split_interval], + ) + tone = gr.Textbox( + label="アクセント調整(数値は 0=低 か1=高 のみ)", + info="改行で分けない場合のみ使えます。万能ではありません。", + ) + use_tone = gr.Checkbox(label="アクセント調整を使う", value=False) + use_tone.change( + lambda x: (gr.Checkbox(value=False) if x else gr.Checkbox()), + inputs=[use_tone], + outputs=[line_split], + ) + language = gr.Dropdown(choices=languages, value="JP", label="Language") + speaker = gr.Dropdown(label="話者") + with gr.Accordion(label="詳細設定", open=False): + sdp_ratio = gr.Slider( + minimum=0, + maximum=1, + value=DEFAULT_SDP_RATIO, + step=0.1, + label="SDP Ratio", + ) + noise_scale = gr.Slider( + minimum=0.1, + maximum=2, + value=DEFAULT_NOISE, + step=0.1, + label="Noise", + ) + noise_scale_w = gr.Slider( + minimum=0.1, + maximum=2, + value=DEFAULT_NOISEW, + step=0.1, + label="Noise_W", + ) + length_scale = gr.Slider( + minimum=0.1, + maximum=2, + value=DEFAULT_LENGTH, + step=0.1, + label="Length", + ) + use_assist_text = gr.Checkbox( + label="Assist textを使う", value=False + ) + assist_text = gr.Textbox( + label="Assist text", + placeholder="どうして私の意見を無視するの?許せない、ムカつく!死ねばいいのに。", + info="このテキストの読み上げと似た声音・感情になりやすくなります。ただ抑揚やテンポ等が犠牲になる傾向があります。", + visible=False, + ) + assist_text_weight = gr.Slider( + minimum=0, + maximum=1, + value=DEFAULT_ASSIST_TEXT_WEIGHT, + step=0.1, + label="Assist textの強さ", + visible=False, + ) + use_assist_text.change( + lambda x: (gr.Textbox(visible=x), gr.Slider(visible=x)), + inputs=[use_assist_text], + outputs=[assist_text, assist_text_weight], + ) + with gr.Column(): + with gr.Accordion("スタイルについて詳細", open=False): + gr.Markdown(style_md) + style_mode = gr.Radio( + ["プリセットから選ぶ", "音声ファイルを入力"], + label="スタイルの指定方法", + value="プリセットから選ぶ", + ) + style = gr.Dropdown( + label=f"スタイル({DEFAULT_STYLE}が平均スタイル)", + choices=["モデルをロードしてください"], + value="モデルをロードしてください", + ) + style_weight = gr.Slider( + minimum=0, + maximum=50, + value=DEFAULT_STYLE_WEIGHT, + step=0.1, + label="スタイルの強さ", + ) + ref_audio_path = gr.Audio( + label="参照音声", type="filepath", visible=False + ) + tts_button = gr.Button( + "音声合成(モデルをロードしてください)", + variant="primary", + interactive=False, + ) + text_output = gr.Textbox(label="情報") + audio_output = gr.Audio(label="結果") + with gr.Accordion("テキスト例", open=False): + gr.Examples(examples, inputs=[text_input, language]) + + tts_button.click( + tts_fn, + inputs=[ + model_name, + model_path, + text_input, + language, + ref_audio_path, + sdp_ratio, + noise_scale, + noise_scale_w, + length_scale, + line_split, + split_interval, + assist_text, + assist_text_weight, + use_assist_text, + style, + style_weight, + tone, + use_tone, + speaker, + pitch_scale, + intonation_scale, + ], + outputs=[text_output, audio_output, tone], + ) + + model_name.change( + model_holder.update_model_files_gr, + inputs=[model_name], + outputs=[model_path], + ) + + model_path.change(make_non_interactive, outputs=[tts_button]) + + refresh_button.click( + model_holder.update_model_names_gr, + outputs=[model_name, model_path, tts_button], + ) + + load_button.click( + model_holder.load_model_gr, + inputs=[model_name, model_path], + outputs=[style, tts_button, speaker], + ) + + style_mode.change( + gr_util, + inputs=[style_mode], + outputs=[style, ref_audio_path], + ) + + # app.launch( + # inbrowser=not args.no_autolaunch, share=args.share, server_name=args.server_name + # ) + + return app diff --git a/webui_merge.py b/webui/merge.py similarity index 65% rename from webui_merge.py rename to webui/merge.py index a58471a..c002e33 100644 --- a/webui_merge.py +++ b/webui/merge.py @@ -330,187 +330,192 @@ Happy, Surprise, HappySurprise - 構造上の相性の関係で、スタイルベクトルを混ぜる重みは、上の「話し方」と同じ比率で混ぜられます。例えば「話し方」が0のときはモデルAのみしか使われません。 """ -model_names = model_holder.model_names -if len(model_names) == 0: - logger.error( - f"モデルが見つかりませんでした。{assets_root}にモデルを置いてください。" - ) - sys.exit(1) -initial_id = 0 -initial_model_files = model_holder.model_files_dict[model_names[initial_id]] -with gr.Blocks(theme=GRADIO_THEME) as app: - gr.Markdown(initial_md) - with gr.Accordion(label="使い方", open=False): +def create_merge_app(): + model_names = model_holder.model_names + if len(model_names) == 0: + logger.error( + f"モデルが見つかりませんでした。{assets_root}にモデルを置いてください。" + ) + sys.exit(1) + initial_id = 0 + initial_model_files = model_holder.model_files_dict[model_names[initial_id]] + + with gr.Blocks(theme=GRADIO_THEME) as app: gr.Markdown(initial_md) - with gr.Row(): - with gr.Column(scale=3): - model_name_a = gr.Dropdown( - label="モデルA", - choices=model_names, - value=model_names[initial_id], - ) - model_path_a = gr.Dropdown( - label="モデルファイル", - choices=initial_model_files, - value=initial_model_files[0], - ) - with gr.Column(scale=3): - model_name_b = gr.Dropdown( - label="モデルB", - choices=model_names, - value=model_names[initial_id], - ) - model_path_b = gr.Dropdown( - label="モデルファイル", - choices=initial_model_files, - value=initial_model_files[0], - ) - refresh_button = gr.Button("更新", scale=1, visible=True) - with gr.Column(variant="panel"): - new_name = gr.Textbox(label="新しいモデル名", placeholder="new_model") + with gr.Accordion(label="使い方", open=False): + gr.Markdown(initial_md) with gr.Row(): - voice_slider = gr.Slider( - label="声質", - value=0, - minimum=0, - maximum=1, - step=0.1, - ) - voice_pitch_slider = gr.Slider( - label="声の高さ", - value=0, - minimum=0, - maximum=1, - step=0.1, - ) - speech_style_slider = gr.Slider( - label="話し方(抑揚・感情表現等)", - value=0, - minimum=0, - maximum=1, - step=0.1, - ) - tempo_slider = gr.Slider( - label="話す速さ・リズム・テンポ", - value=0, - minimum=0, - maximum=1, - step=0.1, - ) - use_slerp_instead_of_lerp = gr.Checkbox( - label="線形補完のかわりに球面線形補完を使う", - value=False, - ) + with gr.Column(scale=3): + model_name_a = gr.Dropdown( + label="モデルA", + choices=model_names, + value=model_names[initial_id], + ) + model_path_a = gr.Dropdown( + label="モデルファイル", + choices=initial_model_files, + value=initial_model_files[0], + ) + with gr.Column(scale=3): + model_name_b = gr.Dropdown( + label="モデルB", + choices=model_names, + value=model_names[initial_id], + ) + model_path_b = gr.Dropdown( + label="モデルファイル", + choices=initial_model_files, + value=initial_model_files[0], + ) + refresh_button = gr.Button("更新", scale=1, visible=True) with gr.Column(variant="panel"): - gr.Markdown("## モデルファイル(safetensors)のマージ") - model_merge_button = gr.Button("モデルファイルのマージ", variant="primary") - info_model_merge = gr.Textbox(label="情報") - with gr.Column(variant="panel"): - gr.Markdown(style_merge_md) + new_name = gr.Textbox(label="新しいモデル名", placeholder="new_model") with gr.Row(): - load_style_button = gr.Button("スタイル一覧をロード", scale=1) - styles_a = gr.Textbox(label="モデルAのスタイル一覧") - styles_b = gr.Textbox(label="モデルBのスタイル一覧") - style_triple_list = gr.TextArea( - label="スタイルのマージリスト", - placeholder=f"{DEFAULT_STYLE}, {DEFAULT_STYLE},{DEFAULT_STYLE}\nAngry, Angry, Angry", - value=f"{DEFAULT_STYLE}, {DEFAULT_STYLE}, {DEFAULT_STYLE}", - ) - style_merge_button = gr.Button("スタイルのマージ", variant="primary") - info_style_merge = gr.Textbox(label="情報") + voice_slider = gr.Slider( + label="声質", + value=0, + minimum=0, + maximum=1, + step=0.1, + ) + voice_pitch_slider = gr.Slider( + label="声の高さ", + value=0, + minimum=0, + maximum=1, + step=0.1, + ) + speech_style_slider = gr.Slider( + label="話し方(抑揚・感情表現等)", + value=0, + minimum=0, + maximum=1, + step=0.1, + ) + tempo_slider = gr.Slider( + label="話す速さ・リズム・テンポ", + value=0, + minimum=0, + maximum=1, + step=0.1, + ) + use_slerp_instead_of_lerp = gr.Checkbox( + label="線形補完のかわりに球面線形補完を使う", + value=False, + ) + with gr.Column(variant="panel"): + gr.Markdown("## モデルファイル(safetensors)のマージ") + model_merge_button = gr.Button( + "モデルファイルのマージ", variant="primary" + ) + info_model_merge = gr.Textbox(label="情報") + with gr.Column(variant="panel"): + gr.Markdown(style_merge_md) + with gr.Row(): + load_style_button = gr.Button("スタイル一覧をロード", scale=1) + styles_a = gr.Textbox(label="モデルAのスタイル一覧") + styles_b = gr.Textbox(label="モデルBのスタイル一覧") + style_triple_list = gr.TextArea( + label="スタイルのマージリスト", + placeholder=f"{DEFAULT_STYLE}, {DEFAULT_STYLE},{DEFAULT_STYLE}\nAngry, Angry, Angry", + value=f"{DEFAULT_STYLE}, {DEFAULT_STYLE}, {DEFAULT_STYLE}", + ) + style_merge_button = gr.Button("スタイルのマージ", variant="primary") + info_style_merge = gr.Textbox(label="情報") - text_input = gr.TextArea( - label="テキスト", value="これはテストです。聞こえていますか?" - ) - style = gr.Dropdown( - label="スタイル", - choices=["スタイルをマージしてください"], - value="スタイルをマージしてください", - ) - emotion_weight = gr.Slider( - minimum=0, - maximum=50, - value=1, - step=0.1, - label="スタイルの強さ", - ) - tts_button = gr.Button("音声合成", variant="primary") - audio_output = gr.Audio(label="結果") + text_input = gr.TextArea( + label="テキスト", value="これはテストです。聞こえていますか?" + ) + style = gr.Dropdown( + label="スタイル", + choices=["スタイルをマージしてください"], + value="スタイルをマージしてください", + ) + emotion_weight = gr.Slider( + minimum=0, + maximum=50, + value=1, + step=0.1, + label="スタイルの強さ", + ) + tts_button = gr.Button("音声合成", variant="primary") + audio_output = gr.Audio(label="結果") - model_name_a.change( - model_holder.update_model_files_gr, - inputs=[model_name_a], - outputs=[model_path_a], - ) - model_name_b.change( - model_holder.update_model_files_gr, - inputs=[model_name_b], - outputs=[model_path_b], - ) + model_name_a.change( + model_holder.update_model_files_gr, + inputs=[model_name_a], + outputs=[model_path_a], + ) + model_name_b.change( + model_holder.update_model_files_gr, + inputs=[model_name_b], + outputs=[model_path_b], + ) - refresh_button.click( - update_two_model_names_dropdown, - outputs=[model_name_a, model_path_a, model_name_b, model_path_b], + refresh_button.click( + update_two_model_names_dropdown, + outputs=[model_name_a, model_path_a, model_name_b, model_path_b], + ) + + load_style_button.click( + load_styles_gr, + inputs=[model_name_a, model_name_b], + outputs=[styles_a, styles_b, style_triple_list], + ) + + model_merge_button.click( + merge_models_gr, + inputs=[ + model_name_a, + model_path_a, + model_name_b, + model_path_b, + new_name, + voice_slider, + voice_pitch_slider, + speech_style_slider, + tempo_slider, + use_slerp_instead_of_lerp, + ], + outputs=[info_model_merge], + ) + + style_merge_button.click( + merge_style_gr, + inputs=[ + model_name_a, + model_name_b, + speech_style_slider, + new_name, + style_triple_list, + ], + outputs=[info_style_merge, style], + ) + + tts_button.click( + simple_tts, + inputs=[new_name, text_input, style, emotion_weight], + outputs=[audio_output], + ) + + parser = argparse.ArgumentParser() + parser.add_argument( + "--server-name", + type=str, + default=None, + help="Server name for Gradio app", ) - - load_style_button.click( - load_styles_gr, - inputs=[model_name_a, model_name_b], - outputs=[styles_a, styles_b, style_triple_list], + parser.add_argument( + "--no-autolaunch", + action="store_true", + default=False, + help="Do not launch app automatically", ) + parser.add_argument("--share", action="store_true", default=False) + args = parser.parse_args() - model_merge_button.click( - merge_models_gr, - inputs=[ - model_name_a, - model_path_a, - model_name_b, - model_path_b, - new_name, - voice_slider, - voice_pitch_slider, - speech_style_slider, - tempo_slider, - use_slerp_instead_of_lerp, - ], - outputs=[info_model_merge], - ) - - style_merge_button.click( - merge_style_gr, - inputs=[ - model_name_a, - model_name_b, - speech_style_slider, - new_name, - style_triple_list, - ], - outputs=[info_style_merge, style], - ) - - tts_button.click( - simple_tts, - inputs=[new_name, text_input, style, emotion_weight], - outputs=[audio_output], - ) - -parser = argparse.ArgumentParser() -parser.add_argument( - "--server-name", - type=str, - default=None, - help="Server name for Gradio app", -) -parser.add_argument( - "--no-autolaunch", - action="store_true", - default=False, - help="Do not launch app automatically", -) -parser.add_argument("--share", action="store_true", default=False) -args = parser.parse_args() - -app.launch( - inbrowser=not args.no_autolaunch, server_name=args.server_name, share=args.share -) + # app.launch( + # inbrowser=not args.no_autolaunch, server_name=args.server_name, share=args.share + # ) + return app diff --git a/webui_style_vectors.py b/webui/style_vectors.py similarity index 65% rename from webui_style_vectors.py rename to webui/style_vectors.py index b89c9c9..27b95bb 100644 --- a/webui_style_vectors.py +++ b/webui/style_vectors.py @@ -323,158 +323,161 @@ UMAPの場合はepsは0.3くらい、t-SNEの場合は2.5くらいがいいか https://ja.wikipedia.org/wiki/DBSCAN """ -with gr.Blocks(theme=GRADIO_THEME) as app: - gr.Markdown(initial_md) - with gr.Row(): - model_name = gr.Textbox(placeholder="your_model_name", label="モデル名") - reduction_method = gr.Radio( - choices=["UMAP", "t-SNE"], - label="次元削減方法", - info="v 1.3以前はt-SNEでしたがUMAPのほうがよい可能性もあります。", - value="UMAP", - ) - load_button = gr.Button("スタイルベクトルを読み込む", variant="primary") - output = gr.Plot(label="音声スタイルの可視化") - load_button.click(load, inputs=[model_name, reduction_method], outputs=[output]) - with gr.Tab("方法1: スタイル分けを自動で行う"): - with gr.Tab("スタイル分け1"): - n_clusters = gr.Slider( - minimum=2, - maximum=10, - step=1, - value=4, - label="作るスタイルの数(平均スタイルを除く)", - info="上の図を見ながらスタイルの数を試行錯誤してください。", + +def create_style_vectors_app(): + with gr.Blocks(theme=GRADIO_THEME) as app: + gr.Markdown(initial_md) + with gr.Row(): + model_name = gr.Textbox(placeholder="your_model_name", label="モデル名") + reduction_method = gr.Radio( + choices=["UMAP", "t-SNE"], + label="次元削減方法", + info="v 1.3以前はt-SNEでしたがUMAPのほうがよい可能性もあります。", + value="UMAP", ) - c_method = gr.Radio( - choices=[ - "Agglomerative after reduction", - "KMeans after reduction", - "Agglomerative", - "KMeans", - ], - label="アルゴリズム", - info="分類する(クラスタリング)アルゴリズムを選択します。いろいろ試してみてください。", - value="Agglomerative after reduction", - ) - c_button = gr.Button("スタイル分けを実行") - with gr.Tab("スタイル分け2: DBSCAN"): - gr.Markdown(dbscan_md) - eps = gr.Slider( - minimum=0.1, - maximum=10, - step=0.01, - value=0.3, - label="eps", - ) - min_samples = gr.Slider( - minimum=1, - maximum=50, - step=1, - value=15, - label="min_samples", + load_button = gr.Button("スタイルベクトルを読み込む", variant="primary") + output = gr.Plot(label="音声スタイルの可視化") + load_button.click(load, inputs=[model_name, reduction_method], outputs=[output]) + with gr.Tab("方法1: スタイル分けを自動で行う"): + with gr.Tab("スタイル分け1"): + n_clusters = gr.Slider( + minimum=2, + maximum=10, + step=1, + value=4, + label="作るスタイルの数(平均スタイルを除く)", + info="上の図を見ながらスタイルの数を試行錯誤してください。", + ) + c_method = gr.Radio( + choices=[ + "Agglomerative after reduction", + "KMeans after reduction", + "Agglomerative", + "KMeans", + ], + label="アルゴリズム", + info="分類する(クラスタリング)アルゴリズムを選択します。いろいろ試してみてください。", + value="Agglomerative after reduction", + ) + c_button = gr.Button("スタイル分けを実行") + with gr.Tab("スタイル分け2: DBSCAN"): + gr.Markdown(dbscan_md) + eps = gr.Slider( + minimum=0.1, + maximum=10, + step=0.01, + value=0.3, + label="eps", + ) + min_samples = gr.Slider( + minimum=1, + maximum=50, + step=1, + value=15, + label="min_samples", + ) + with gr.Row(): + dbscan_button = gr.Button("スタイル分けを実行") + num_styles_result = gr.Textbox(label="スタイル数") + gr.Markdown("スタイル分けの結果") + gr.Markdown( + "注意: もともと256次元なものをを2次元に落としているので、正確なベクトルの位置関係ではありません。" ) with gr.Row(): - dbscan_button = gr.Button("スタイル分けを実行") - num_styles_result = gr.Textbox(label="スタイル数") - gr.Markdown("スタイル分けの結果") - gr.Markdown( - "注意: もともと256次元なものをを2次元に落としているので、正確なベクトルの位置関係ではありません。" - ) - with gr.Row(): - gr_plot = gr.Plot() - with gr.Column(): - with gr.Row(): - cluster_index = gr.Slider( - minimum=1, - maximum=MAX_CLUSTER_NUM, - step=1, - value=1, - label="スタイル番号", - info="選択したスタイルの代表音声を表示します。", - ) - num_files = gr.Slider( - minimum=1, - maximum=MAX_AUDIO_NUM, - step=1, - value=5, - label="代表音声の数をいくつ表示するか", - ) - get_audios_button = gr.Button("代表音声を取得") - with gr.Row(): - audio_list = [] - for i in range(MAX_AUDIO_NUM): - audio_list.append(gr.Audio(visible=False, show_label=True)) - c_button.click( - do_clustering_gradio, - inputs=[n_clusters, c_method], - outputs=[gr_plot, cluster_index] + audio_list, + gr_plot = gr.Plot() + with gr.Column(): + with gr.Row(): + cluster_index = gr.Slider( + minimum=1, + maximum=MAX_CLUSTER_NUM, + step=1, + value=1, + label="スタイル番号", + info="選択したスタイルの代表音声を表示します。", + ) + num_files = gr.Slider( + minimum=1, + maximum=MAX_AUDIO_NUM, + step=1, + value=5, + label="代表音声の数をいくつ表示するか", + ) + get_audios_button = gr.Button("代表音声を取得") + with gr.Row(): + audio_list = [] + for i in range(MAX_AUDIO_NUM): + audio_list.append(gr.Audio(visible=False, show_label=True)) + c_button.click( + do_clustering_gradio, + inputs=[n_clusters, c_method], + outputs=[gr_plot, cluster_index] + audio_list, + ) + dbscan_button.click( + do_dbscan_gradio, + inputs=[eps, min_samples], + outputs=[gr_plot, cluster_index, num_styles_result] + audio_list, + ) + get_audios_button.click( + representative_wav_files_gradio, + inputs=[cluster_index, num_files], + outputs=audio_list, + ) + gr.Markdown("結果が良さそうなら、これを保存します。") + style_names = gr.Textbox( + "Angry, Sad, Happy", + label="スタイルの名前", + info=f"スタイルの名前を`,`で区切って入力してください(日本語可)。例: `Angry, Sad, Happy`や`怒り, 悲しみ, 喜び`など。平均音声は{DEFAULT_STYLE}として自動的に保存されます。", ) - dbscan_button.click( - do_dbscan_gradio, - inputs=[eps, min_samples], - outputs=[gr_plot, cluster_index, num_styles_result] + audio_list, - ) - get_audios_button.click( - representative_wav_files_gradio, - inputs=[cluster_index, num_files], - outputs=audio_list, - ) - gr.Markdown("結果が良さそうなら、これを保存します。") - style_names = gr.Textbox( - "Angry, Sad, Happy", - label="スタイルの名前", - info=f"スタイルの名前を`,`で区切って入力してください(日本語可)。例: `Angry, Sad, Happy`や`怒り, 悲しみ, 喜び`など。平均音声は{DEFAULT_STYLE}として自動的に保存されます。", - ) - with gr.Row(): - save_button1 = gr.Button("スタイルベクトルを保存", variant="primary") - info2 = gr.Textbox(label="保存結果") + with gr.Row(): + save_button1 = gr.Button("スタイルベクトルを保存", variant="primary") + info2 = gr.Textbox(label="保存結果") - save_button1.click( - save_style_vectors_from_clustering, - inputs=[model_name, style_names], - outputs=[info2], - ) - with gr.Tab("方法2: 手動でスタイルを選ぶ"): - gr.Markdown( - "下のテキスト欄に、各スタイルの代表音声のファイル名を`,`区切りで、その横に対応するスタイル名を`,`区切りで入力してください。" - ) - gr.Markdown("例: `angry.wav, sad.wav, happy.wav`と`Angry, Sad, Happy`") - gr.Markdown( - f"注意: {DEFAULT_STYLE}スタイルは自動的に保存されます、手動では{DEFAULT_STYLE}という名前のスタイルは指定しないでください。" - ) - with gr.Row(): - audio_files_text = gr.Textbox( - label="音声ファイル名", placeholder="angry.wav, sad.wav, happy.wav" - ) - style_names_text = gr.Textbox( - label="スタイル名", placeholder="Angry, Sad, Happy" - ) - with gr.Row(): - save_button2 = gr.Button("スタイルベクトルを保存", variant="primary") - info2 = gr.Textbox(label="保存結果") - save_button2.click( - save_style_vectors_from_files, - inputs=[model_name, audio_files_text, style_names_text], + save_button1.click( + save_style_vectors_from_clustering, + inputs=[model_name, style_names], outputs=[info2], ) + with gr.Tab("方法2: 手動でスタイルを選ぶ"): + gr.Markdown( + "下のテキスト欄に、各スタイルの代表音声のファイル名を`,`区切りで、その横に対応するスタイル名を`,`区切りで入力してください。" + ) + gr.Markdown("例: `angry.wav, sad.wav, happy.wav`と`Angry, Sad, Happy`") + gr.Markdown( + f"注意: {DEFAULT_STYLE}スタイルは自動的に保存されます、手動では{DEFAULT_STYLE}という名前のスタイルは指定しないでください。" + ) + with gr.Row(): + audio_files_text = gr.Textbox( + label="音声ファイル名", placeholder="angry.wav, sad.wav, happy.wav" + ) + style_names_text = gr.Textbox( + label="スタイル名", placeholder="Angry, Sad, Happy" + ) + with gr.Row(): + save_button2 = gr.Button("スタイルベクトルを保存", variant="primary") + info2 = gr.Textbox(label="保存結果") + save_button2.click( + save_style_vectors_from_files, + inputs=[model_name, audio_files_text, style_names_text], + outputs=[info2], + ) -parser = argparse.ArgumentParser() -parser.add_argument( - "--server-name", - type=str, - default=None, - help="Server name for Gradio app", -) -parser.add_argument( - "--no-autolaunch", - action="store_true", - default=False, - help="Do not launch app automatically", -) -parser.add_argument("--share", action="store_true", default=False) -args = parser.parse_args() + parser = argparse.ArgumentParser() + parser.add_argument( + "--server-name", + type=str, + default=None, + help="Server name for Gradio app", + ) + parser.add_argument( + "--no-autolaunch", + action="store_true", + default=False, + help="Do not launch app automatically", + ) + parser.add_argument("--share", action="store_true", default=False) + args = parser.parse_args() -app.launch( - inbrowser=not args.no_autolaunch, server_name=args.server_name, share=args.share -) + # app.launch( + # inbrowser=not args.no_autolaunch, server_name=args.server_name, share=args.share + # ) + return app diff --git a/webui_train.py b/webui/train.py similarity index 99% rename from webui_train.py rename to webui/train.py index 59cc9f5..1c9a214 100644 --- a/webui_train.py +++ b/webui/train.py @@ -450,7 +450,8 @@ english_teacher.wav|Mary|EN|How are you? I'm fine, thank you, and you? 日本語話者の単一話者データセットでも構いません。 """ -if __name__ == "__main__": + +def create_train_app(): with gr.Blocks(theme=GRADIO_THEME).queue() as app: gr.Markdown(initial_md) with gr.Accordion(label="データの前準備", open=False): @@ -808,4 +809,5 @@ if __name__ == "__main__": ) args = parser.parse_args() - app.launch(inbrowser=not args.no_autolaunch, server_name=args.server_name) + # app.launch(inbrowser=not args.no_autolaunch, server_name=args.server_name) + return app From bc1058270c9a8648cae01bd301a136ef4ff92466 Mon Sep 17 00:00:00 2001 From: kale4eat Date: Tue, 5 Mar 2024 12:16:14 +0900 Subject: [PATCH 02/14] Delete pyopenjtalk import --- server_editor.py | 1 - 1 file changed, 1 deletion(-) diff --git a/server_editor.py b/server_editor.py index afb9321..1ae323f 100644 --- a/server_editor.py +++ b/server_editor.py @@ -19,7 +19,6 @@ from pathlib import Path import yaml import numpy as np -import pyopenjtalk import requests import torch import uvicorn From 1eb8fb4b08c5dba3c2d7bdd40bf56da46369d321 Mon Sep 17 00:00:00 2001 From: kale4eat Date: Tue, 5 Mar 2024 12:17:15 +0900 Subject: [PATCH 03/14] add openjtalk worker pkg --- text/pyopenjtalk_worker/__init__.py | 100 ++++++++++++++++++++++ text/pyopenjtalk_worker/__main__.py | 16 ++++ text/pyopenjtalk_worker/worker_client.py | 40 +++++++++ text/pyopenjtalk_worker/worker_common.py | 44 ++++++++++ text/pyopenjtalk_worker/worker_server.py | 103 +++++++++++++++++++++++ 5 files changed, 303 insertions(+) create mode 100644 text/pyopenjtalk_worker/__init__.py create mode 100644 text/pyopenjtalk_worker/__main__.py create mode 100644 text/pyopenjtalk_worker/worker_client.py create mode 100644 text/pyopenjtalk_worker/worker_common.py create mode 100644 text/pyopenjtalk_worker/worker_server.py diff --git a/text/pyopenjtalk_worker/__init__.py b/text/pyopenjtalk_worker/__init__.py new file mode 100644 index 0000000..6ab4ca1 --- /dev/null +++ b/text/pyopenjtalk_worker/__init__.py @@ -0,0 +1,100 @@ +""" +Run the pyopenjtalk worker in a separate process +to avoid user dictionary access error +""" + +from typing import Optional, Any + +from .worker_common import WOKER_PORT +from .worker_client import WorkerClient + +from common.log import logger + +WORKER_CLIENT: Optional[WorkerClient] = None + +# pyopenjtalk interface + +# g2p: not used + + +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 + + +def make_label(njd_features) -> list[str]: + assert WORKER_CLIENT + ret = WORKER_CLIENT.dispatch_pyopenjtalk("make_label", njd_features) + assert isinstance(ret, list) + return ret + + +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 update_global_jtalk_with_user_dict(path: str): + assert WORKER_CLIENT + WORKER_CLIENT.dispatch_pyopenjtalk("update_global_jtalk_with_user_dict", path) + + +def unset_user_dict(): + assert WORKER_CLIENT + WORKER_CLIENT.dispatch_pyopenjtalk("unset_user_dict") + + +# initialize module when imported + + +def initialize(port: int = WOKER_PORT): + import time + import socket + import sys + import atexit + + global WORKER_CLIENT + logger.debug("initialize") + if WORKER_CLIENT: + return + + client = None + try: + client = WorkerClient(port) + except (socket.timeout, socket.error): + logger.debug("try starting worker server") + import os + import subprocess + + worker_pkg_path = os.path.relpath( + os.path.dirname(__file__), os.getcwd() + ).replace(os.sep, ".") + subprocess.Popen([sys.executable, "-m", worker_pkg_path, "--port", str(port)]) + # wait until server listening + count = 0 + while True: + try: + client = WorkerClient(port) + break + except socket.error: + time.sleep(1) + count += 1 + # 10: max number of retries + if count == 10: + raise TimeoutError("サーバーに接続できませんでした") + + WORKER_CLIENT = client + + def terminate(): + global WORKER_CLIENT + if not WORKER_CLIENT: + return + + if WORKER_CLIENT.status().get("client-count") == 1: + WORKER_CLIENT.quit_server() + WORKER_CLIENT.close() + WORKER_CLIENT = None + + atexit.register(terminate) diff --git a/text/pyopenjtalk_worker/__main__.py b/text/pyopenjtalk_worker/__main__.py new file mode 100644 index 0000000..8b67aa0 --- /dev/null +++ b/text/pyopenjtalk_worker/__main__.py @@ -0,0 +1,16 @@ +import argparse + +from .worker_server import WorkerServer +from .worker_common import WOKER_PORT + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--port", type=int, default=WOKER_PORT) + args = parser.parse_args() + server = WorkerServer() + server.start_server(port=args.port) + + +if __name__ == "__main__": + main() diff --git a/text/pyopenjtalk_worker/worker_client.py b/text/pyopenjtalk_worker/worker_client.py new file mode 100644 index 0000000..bd9a32a --- /dev/null +++ b/text/pyopenjtalk_worker/worker_client.py @@ -0,0 +1,40 @@ +from typing import Any +import socket + +from .worker_common import RequestType, receive_data, send_data + + +class WorkerClient: + def __init__(self, port: int) -> None: + sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + # 5: timeout + sock.settimeout(5) + sock.connect((socket.gethostname(), port)) + self.sock = sock + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_value, traceback): + self.close() + + def close(self): + self.sock.close() + + def dispatch_pyopenjtalk(self, func: str, *args, **kwargs): + data = { + "request-type": RequestType.PYOPENJTALK, + "func": func, + "args": args, + "kwargs": kwargs, + } + send_data(self.sock, data) + return receive_data(self.sock).get("return") + + def status(self): + send_data(self.sock, {"request-type": RequestType.STATUS}) + return receive_data(self.sock) + + def quit_server(self): + send_data(self.sock, {"request-type": RequestType.QUIT_SERVER}) + receive_data(self.sock) diff --git a/text/pyopenjtalk_worker/worker_common.py b/text/pyopenjtalk_worker/worker_common.py new file mode 100644 index 0000000..bea552e --- /dev/null +++ b/text/pyopenjtalk_worker/worker_common.py @@ -0,0 +1,44 @@ +from typing import Any, Optional, Final +from enum import IntEnum, auto +import socket +import json + +WOKER_PORT: Final[int] = 7861 +HEADER_SIZE: Final[int] = 4 + + +class RequestType(IntEnum): + STATUS = auto() + QUIT_SERVER = auto() + PYOPENJTALK = auto() + + +class ConnectionClosedException(Exception): + pass + + +# socket communication + + +def send_data(sock: socket.socket, data: dict[str, Any]): + json_data = json.dumps(data).encode() + header = len(json_data).to_bytes(HEADER_SIZE, byteorder="big") + sock.sendall(header + json_data) + + +def _receive_until(sock: socket.socket, size: int): + data = b"" + while len(data) < size: + part = sock.recv(size - len(data)) + if part == b"": + raise ConnectionClosedException("接続が閉じられました") + data += part + + return data + + +def receive_data(sock: socket.socket) -> dict[str, Any]: + header = _receive_until(sock, HEADER_SIZE) + data_length = int.from_bytes(header, byteorder="big") + body = _receive_until(sock, data_length) + return json.loads(body.decode()) diff --git a/text/pyopenjtalk_worker/worker_server.py b/text/pyopenjtalk_worker/worker_server.py new file mode 100644 index 0000000..a2b9f2f --- /dev/null +++ b/text/pyopenjtalk_worker/worker_server.py @@ -0,0 +1,103 @@ +import pyopenjtalk +import socket +import select + + +from .worker_common import ( + ConnectionClosedException, + RequestType, + receive_data, + send_data, +) + +from common.log import logger + +# To make it as fast as possible +# Probably faster than calling getattr every time +_PYOPENJTALK_FUNC_DICT = { + "run_frontend": pyopenjtalk.run_frontend, + "make_label": pyopenjtalk.make_label, + "mecab_dict_index": pyopenjtalk.mecab_dict_index, + "update_global_jtalk_with_user_dict": pyopenjtalk.update_global_jtalk_with_user_dict, + "unset_user_dict": pyopenjtalk.unset_user_dict, +} + + +class WorkerServer: + def __init__(self) -> None: + self.client_count: int = 0 + self.quit: bool = False + + def handle_request(self, request): + request_type = None + try: + request_type = RequestType(request.get("request-type")) + except Exception: + return { + "success": False, + "reason": "request-type is invalid", + } + + if request_type: + if request_type == RequestType.STATUS: + response = { + "success": True, + "client-count": self.client_count, + } + elif request_type == RequestType.QUIT_SERVER: + self.quit = True + response = {"success": True} + elif request_type == RequestType.PYOPENJTALK: + func_name = request.get("func") + assert isinstance(func_name, str) + func = _PYOPENJTALK_FUNC_DICT[func_name] + args = request.get("args") + kwargs = request.get("kwargs") + assert isinstance(args, list) + assert isinstance(kwargs, dict) + ret = func(*args, **kwargs) + response = {"success": True, "return": ret} + else: + # NOT REACHED + response = request + + return response + + def start_server(self, port: int): + logger.info("start pyopenjtalk worker server") + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as server_socket: + server_socket.bind((socket.gethostname(), port)) + server_socket.listen() + sockets = [server_socket] + while True: + ready_sockets, _, _ = select.select(sockets, [], [], 0.1) + for sock in ready_sockets: + if sock is server_socket: + logger.info("new client connected") + client_socket, _ = server_socket.accept() + sockets.append(client_socket) + self.client_count += 1 + else: + # client + try: + request = receive_data(sock) + except ConnectionClosedException as e: + sock.close() + sockets.remove(sock) + self.client_count -= 1 + logger.info("close connection") + continue + + logger.debug(f"receive request: {request}") + + response = self.handle_request(request) + logger.debug(f"send response: {response}") + try: + send_data(sock, response) + except Exception: + logger.warning( + "an exception occurred during sending responce" + ) + if self.quit: + logger.info("quit pyopenjtalk worker server") + return From 85f5b9bd25e930dc8e0fd8e96614cd1eaa5127f7 Mon Sep 17 00:00:00 2001 From: kale4eat Date: Tue, 5 Mar 2024 12:18:05 +0900 Subject: [PATCH 04/14] replace pyopenjtalk import with worker --- text/japanese.py | 4 +++- text/user_dict/__init__.py | 4 +++- 2 files changed, 6 insertions(+), 2 deletions(-) diff --git a/text/japanese.py b/text/japanese.py index b18bc68..fea0eaa 100644 --- a/text/japanese.py +++ b/text/japanese.py @@ -4,7 +4,9 @@ import re import unicodedata from pathlib import Path -import pyopenjtalk +from . import pyopenjtalk_worker as pyopenjtalk + +pyopenjtalk.initialize() from num2words import num2words from transformers import AutoTokenizer diff --git a/text/user_dict/__init__.py b/text/user_dict/__init__.py index c12b3d1..95515cb 100644 --- a/text/user_dict/__init__.py +++ b/text/user_dict/__init__.py @@ -12,7 +12,9 @@ from typing import Dict, List, Optional from uuid import UUID, uuid4 import numpy as np -import pyopenjtalk +from .. import pyopenjtalk_worker as pyopenjtalk + +pyopenjtalk.initialize() from fastapi import HTTPException from .word_model import UserDictWord, WordTypes From 98ff976ec0e187bef23c6bf5d91bbe09829c9474 Mon Sep 17 00:00:00 2001 From: kale4eat Date: Tue, 5 Mar 2024 15:04:27 +0900 Subject: [PATCH 05/14] modify logging --- text/pyopenjtalk_worker/__init__.py | 2 +- text/pyopenjtalk_worker/worker_client.py | 8 +++++++- text/pyopenjtalk_worker/worker_server.py | 6 +++--- 3 files changed, 11 insertions(+), 5 deletions(-) diff --git a/text/pyopenjtalk_worker/__init__.py b/text/pyopenjtalk_worker/__init__.py index 6ab4ca1..ef5873c 100644 --- a/text/pyopenjtalk_worker/__init__.py +++ b/text/pyopenjtalk_worker/__init__.py @@ -64,7 +64,7 @@ def initialize(port: int = WOKER_PORT): try: client = WorkerClient(port) except (socket.timeout, socket.error): - logger.debug("try starting worker server") + logger.debug("try starting pyopenjtalk worker server") import os import subprocess diff --git a/text/pyopenjtalk_worker/worker_client.py b/text/pyopenjtalk_worker/worker_client.py index bd9a32a..23f7dbe 100644 --- a/text/pyopenjtalk_worker/worker_client.py +++ b/text/pyopenjtalk_worker/worker_client.py @@ -3,6 +3,8 @@ import socket from .worker_common import RequestType, receive_data, send_data +from common.log import logger + class WorkerClient: def __init__(self, port: int) -> None: @@ -28,8 +30,12 @@ class WorkerClient: "args": args, "kwargs": kwargs, } + logger.trace(f"client sends request: {data}") send_data(self.sock, data) - return receive_data(self.sock).get("return") + logger.trace("client sent request successfully") + response = receive_data(self.sock) + logger.trace(f"client received response: {response}") + return response.get("return") def status(self): send_data(self.sock, {"request-type": RequestType.STATUS}) diff --git a/text/pyopenjtalk_worker/worker_server.py b/text/pyopenjtalk_worker/worker_server.py index a2b9f2f..13bc771 100644 --- a/text/pyopenjtalk_worker/worker_server.py +++ b/text/pyopenjtalk_worker/worker_server.py @@ -2,7 +2,6 @@ import pyopenjtalk import socket import select - from .worker_common import ( ConnectionClosedException, RequestType, @@ -88,12 +87,13 @@ class WorkerServer: logger.info("close connection") continue - logger.debug(f"receive request: {request}") + logger.trace(f"server received request: {request}") response = self.handle_request(request) - logger.debug(f"send response: {response}") + logger.trace(f"server sends response: {response}") try: send_data(sock, response) + logger.trace("server sent response successfully") except Exception: logger.warning( "an exception occurred during sending responce" From 2ab025d5fe4a21522ffe7c9ba7d12a0a061cfa1d Mon Sep 17 00:00:00 2001 From: kale4eat Date: Wed, 6 Mar 2024 09:45:25 +0900 Subject: [PATCH 06/14] Run in a separate process group to avoid receiving signals by Ctrl + C Minor Correction: * logging * declaration of terminate --- text/pyopenjtalk_worker/__init__.py | 34 +++++++++++++++++++++-------- 1 file changed, 25 insertions(+), 9 deletions(-) diff --git a/text/pyopenjtalk_worker/__init__.py b/text/pyopenjtalk_worker/__init__.py index ef5873c..3fd9971 100644 --- a/text/pyopenjtalk_worker/__init__.py +++ b/text/pyopenjtalk_worker/__init__.py @@ -55,8 +55,8 @@ def initialize(port: int = WOKER_PORT): import sys import atexit - global WORKER_CLIENT logger.debug("initialize") + global WORKER_CLIENT if WORKER_CLIENT: return @@ -71,7 +71,16 @@ def initialize(port: int = WOKER_PORT): worker_pkg_path = os.path.relpath( os.path.dirname(__file__), os.getcwd() ).replace(os.sep, ".") - subprocess.Popen([sys.executable, "-m", worker_pkg_path, "--port", str(port)]) + args = [sys.executable, "-m", worker_pkg_path, "--port", str(port)] + # new session, new process group + if sys.platform.startswith("win"): + cf = subprocess.DETACHED_PROCESS | subprocess.CREATE_NEW_PROCESS_GROUP # type: ignore + subprocess.Popen(args, creationflags=cf) + 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 + # wait until server listening count = 0 while True: @@ -86,15 +95,22 @@ def initialize(port: int = WOKER_PORT): raise TimeoutError("サーバーに接続できませんでした") WORKER_CLIENT = client + atexit.register(terminate) - def terminate(): - global WORKER_CLIENT - if not WORKER_CLIENT: - return +# top-level declaration +def terminate(): + logger.debug("terminate") + global WORKER_CLIENT + if not WORKER_CLIENT: + return + + # repare for unexpected errors + try: if WORKER_CLIENT.status().get("client-count") == 1: WORKER_CLIENT.quit_server() - WORKER_CLIENT.close() - WORKER_CLIENT = None + except Exception as e: + logger.error(e) - atexit.register(terminate) + WORKER_CLIENT.close() + WORKER_CLIENT = None From 45c6bde2e75b3115ce32cb50059d783495d8168e Mon Sep 17 00:00:00 2001 From: kale4eat Date: Wed, 6 Mar 2024 11:19:55 +0900 Subject: [PATCH 07/14] In Windows, create new console and hide it Minor Correction: * logging * change status return type --- text/pyopenjtalk_worker/__init__.py | 9 ++++++--- text/pyopenjtalk_worker/worker_client.py | 17 +++++++++++++---- 2 files changed, 19 insertions(+), 7 deletions(-) diff --git a/text/pyopenjtalk_worker/__init__.py b/text/pyopenjtalk_worker/__init__.py index 3fd9971..8a10266 100644 --- a/text/pyopenjtalk_worker/__init__.py +++ b/text/pyopenjtalk_worker/__init__.py @@ -74,8 +74,11 @@ def initialize(port: int = WOKER_PORT): args = [sys.executable, "-m", worker_pkg_path, "--port", str(port)] # new session, new process group if sys.platform.startswith("win"): - cf = subprocess.DETACHED_PROCESS | subprocess.CREATE_NEW_PROCESS_GROUP # type: ignore - subprocess.Popen(args, creationflags=cf) + cf = subprocess.CREATE_NEW_CONSOLE | subprocess.CREATE_NEW_PROCESS_GROUP # type: ignore + si = subprocess.STARTUPINFO() # type: ignore + si.dwFlags |= subprocess.STARTF_USESHOWWINDOW # type: ignore + si.wShowWindow = subprocess.SW_HIDE # type: ignore + subprocess.Popen(args, creationflags=cf, startupinfo=si) else: # align with Windows behavior # start_new_session is same as specifying setsid in preexec_fn @@ -107,7 +110,7 @@ def terminate(): # repare for unexpected errors try: - if WORKER_CLIENT.status().get("client-count") == 1: + if WORKER_CLIENT.status() == 1: WORKER_CLIENT.quit_server() except Exception as e: logger.error(e) diff --git a/text/pyopenjtalk_worker/worker_client.py b/text/pyopenjtalk_worker/worker_client.py index 23f7dbe..86d8969 100644 --- a/text/pyopenjtalk_worker/worker_client.py +++ b/text/pyopenjtalk_worker/worker_client.py @@ -38,9 +38,18 @@ class WorkerClient: return response.get("return") def status(self): - send_data(self.sock, {"request-type": RequestType.STATUS}) - return receive_data(self.sock) + data = {"request-type": RequestType.STATUS} + logger.trace(f"client sends request: {data}") + send_data(self.sock, data) + logger.trace("client sent request successfully") + response = receive_data(self.sock) + logger.trace(f"client received response: {response}") + return response.get("client-count") def quit_server(self): - send_data(self.sock, {"request-type": RequestType.QUIT_SERVER}) - receive_data(self.sock) + data = {"request-type": RequestType.QUIT_SERVER} + logger.trace(f"client sends request: {data}") + send_data(self.sock, data) + logger.trace("client sent request successfully") + response = receive_data(self.sock) + logger.trace(f"client received response: {response}") From ed90af1b8718fe31f7bf3e7565b8583a35f77113 Mon Sep 17 00:00:00 2001 From: kale4eat Date: Wed, 6 Mar 2024 12:13:27 +0900 Subject: [PATCH 08/14] Enhanced server error handling Add signal handling for when the process is killed --- text/pyopenjtalk_worker/__init__.py | 10 ++++++++++ text/pyopenjtalk_worker/worker_server.py | 6 ++++++ 2 files changed, 16 insertions(+) diff --git a/text/pyopenjtalk_worker/__init__.py b/text/pyopenjtalk_worker/__init__.py index 8a10266..e968aed 100644 --- a/text/pyopenjtalk_worker/__init__.py +++ b/text/pyopenjtalk_worker/__init__.py @@ -54,6 +54,7 @@ def initialize(port: int = WOKER_PORT): import socket import sys import atexit + import signal logger.debug("initialize") global WORKER_CLIENT @@ -100,6 +101,15 @@ def initialize(port: int = WOKER_PORT): WORKER_CLIENT = client atexit.register(terminate) + # when the process is killed + def signal_handler(signum, frame): + with open("signal_handler.txt", mode="w") as f: + + pass + terminate() + + signal.signal(signal.SIGTERM, signal_handler) + # top-level declaration def terminate(): diff --git a/text/pyopenjtalk_worker/worker_server.py b/text/pyopenjtalk_worker/worker_server.py index 13bc771..dc6d476 100644 --- a/text/pyopenjtalk_worker/worker_server.py +++ b/text/pyopenjtalk_worker/worker_server.py @@ -86,6 +86,12 @@ class WorkerServer: self.client_count -= 1 logger.info("close connection") continue + except Exception as e: + sock.close() + sockets.remove(sock) + self.client_count -= 1 + logger.error(e) + continue logger.trace(f"server received request: {request}") From 419d0e5bed6e91028c570dad1282f4cc027cf433 Mon Sep 17 00:00:00 2001 From: kale4eat Date: Wed, 6 Mar 2024 22:38:44 +0900 Subject: [PATCH 09/14] Delete debugging traces --- text/pyopenjtalk_worker/__init__.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/text/pyopenjtalk_worker/__init__.py b/text/pyopenjtalk_worker/__init__.py index e968aed..9677666 100644 --- a/text/pyopenjtalk_worker/__init__.py +++ b/text/pyopenjtalk_worker/__init__.py @@ -103,9 +103,6 @@ def initialize(port: int = WOKER_PORT): # when the process is killed def signal_handler(signum, frame): - with open("signal_handler.txt", mode="w") as f: - - pass terminate() signal.signal(signal.SIGTERM, signal_handler) From 2994873e3885fc1bae602dae8d311838697b788e Mon Sep 17 00:00:00 2001 From: kale4eat Date: Thu, 7 Mar 2024 09:40:21 +0900 Subject: [PATCH 10/14] add no client timeout summarize the except statement at disconnection --- text/pyopenjtalk_worker/worker_server.py | 25 ++++++++++++++++-------- 1 file changed, 17 insertions(+), 8 deletions(-) diff --git a/text/pyopenjtalk_worker/worker_server.py b/text/pyopenjtalk_worker/worker_server.py index dc6d476..6babd91 100644 --- a/text/pyopenjtalk_worker/worker_server.py +++ b/text/pyopenjtalk_worker/worker_server.py @@ -1,6 +1,7 @@ import pyopenjtalk import socket import select +import time from .worker_common import ( ConnectionClosedException, @@ -62,13 +63,23 @@ class WorkerServer: return response - def start_server(self, port: int): + def start_server(self, port: int, no_client_timeout: int = 30): logger.info("start pyopenjtalk worker server") with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as server_socket: server_socket.bind((socket.gethostname(), port)) server_socket.listen() sockets = [server_socket] + no_client_since = time.time() while True: + if self.client_count == 0: + if no_client_since is None: + no_client_since = time.time() + elif (time.time() - no_client_since) > no_client_timeout: + logger.info("quit because there is no client") + return + else: + no_client_since = None + ready_sockets, _, _ = select.select(sockets, [], [], 0.1) for sock in ready_sockets: if sock is server_socket: @@ -80,17 +91,15 @@ class WorkerServer: # client try: request = receive_data(sock) - except ConnectionClosedException as e: - sock.close() - sockets.remove(sock) - self.client_count -= 1 - logger.info("close connection") - continue except Exception as e: sock.close() sockets.remove(sock) self.client_count -= 1 - logger.error(e) + # unexpected disconnections + if not isinstance(e, ConnectionClosedException): + logger.error(e) + + logger.info("close connection") continue logger.trace(f"server received request: {request}") From d024b71340630d20a20f4ea4b8c4198e52e1dcc2 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Fri, 8 Mar 2024 10:10:53 +0900 Subject: [PATCH 11/14] Fix typo --- common/log.py | 1 + text/pyopenjtalk_worker/__init__.py | 4 ++-- text/pyopenjtalk_worker/__main__.py | 4 ++-- text/pyopenjtalk_worker/worker_common.py | 2 +- 4 files changed, 6 insertions(+), 5 deletions(-) diff --git a/common/log.py b/common/log.py index 679bb2c..71b5e63 100644 --- a/common/log.py +++ b/common/log.py @@ -14,4 +14,5 @@ log_format = ( "{time:MM-DD HH:mm:ss} |{level:^8}| {file}:{line} | {message}" ) +# logger.add(SAFE_STDOUT, format=log_format, backtrace=True, diagnose=True, level="TRACE") logger.add(SAFE_STDOUT, format=log_format, backtrace=True, diagnose=True) diff --git a/text/pyopenjtalk_worker/__init__.py b/text/pyopenjtalk_worker/__init__.py index 9677666..7aa0d0e 100644 --- a/text/pyopenjtalk_worker/__init__.py +++ b/text/pyopenjtalk_worker/__init__.py @@ -5,7 +5,7 @@ to avoid user dictionary access error from typing import Optional, Any -from .worker_common import WOKER_PORT +from .worker_common import WORKER_PORT from .worker_client import WorkerClient from common.log import logger @@ -49,7 +49,7 @@ def unset_user_dict(): # initialize module when imported -def initialize(port: int = WOKER_PORT): +def initialize(port: int = WORKER_PORT): import time import socket import sys diff --git a/text/pyopenjtalk_worker/__main__.py b/text/pyopenjtalk_worker/__main__.py index 8b67aa0..3bb6b53 100644 --- a/text/pyopenjtalk_worker/__main__.py +++ b/text/pyopenjtalk_worker/__main__.py @@ -1,12 +1,12 @@ import argparse from .worker_server import WorkerServer -from .worker_common import WOKER_PORT +from .worker_common import WORKER_PORT def main(): parser = argparse.ArgumentParser() - parser.add_argument("--port", type=int, default=WOKER_PORT) + parser.add_argument("--port", type=int, default=WORKER_PORT) args = parser.parse_args() server = WorkerServer() server.start_server(port=args.port) diff --git a/text/pyopenjtalk_worker/worker_common.py b/text/pyopenjtalk_worker/worker_common.py index bea552e..606d0c3 100644 --- a/text/pyopenjtalk_worker/worker_common.py +++ b/text/pyopenjtalk_worker/worker_common.py @@ -3,7 +3,7 @@ from enum import IntEnum, auto import socket import json -WOKER_PORT: Final[int] = 7861 +WORKER_PORT: Final[int] = 7861 HEADER_SIZE: Final[int] = 4 From 25ed226acbe13ba9425169c89bc1452c842f424e Mon Sep 17 00:00:00 2001 From: litagin02 Date: Fri, 8 Mar 2024 11:19:03 +0900 Subject: [PATCH 12/14] Add speaker list api --- common/tts_model.py | 21 ++++++++------------- server_editor.py | 17 +++++++++++++---- 2 files changed, 21 insertions(+), 17 deletions(-) diff --git a/common/tts_model.py b/common/tts_model.py index e09787e..9e2171d 100644 --- a/common/tts_model.py +++ b/common/tts_model.py @@ -222,12 +222,14 @@ class ModelHolder: self.current_model: Optional[Model] = None self.model_names: list[str] = [] self.models: list[Model] = [] + self.models_info: list[dict[str, Union[str, list[str]]]] = [] self.refresh() def refresh(self): self.model_files_dict = {} self.model_names = [] self.current_model = None + self.models_info = [] model_dirs = [d for d in self.root_dir.iterdir() if d.is_dir()] for model_dir in model_dirs: @@ -247,26 +249,19 @@ class ModelHolder: continue self.model_files_dict[model_dir.name] = model_files self.model_names.append(model_dir.name) - - def models_info(self): - if hasattr(self, "_models_info"): - return self._models_info - result = [] - for name, files in self.model_files_dict.items(): - # Get styles - config_path = self.root_dir / name / "config.json" hps = utils.get_hparams_from_file(config_path) style2id: dict[str, int] = hps.data.style2id styles = list(style2id.keys()) - result.append( + spk2id: dict[str, int] = hps.data.spk2id + speakers = list(spk2id.keys()) + self.models_info.append( { - "name": name, - "files": [str(f) for f in files], + "name": model_dir.name, + "files": [str(f) for f in model_files], "styles": styles, + "speakers": speakers, } ) - self._models_info = result - return result def load_model(self, model_name: str, model_path_str: str): model_path = Path(model_path_str) diff --git a/server_editor.py b/server_editor.py index 1ae323f..402116f 100644 --- a/server_editor.py +++ b/server_editor.py @@ -16,12 +16,13 @@ import zipfile from datetime import datetime from io import BytesIO from pathlib import Path -import yaml +from typing import Optional import numpy as np import requests import torch import uvicorn +import yaml from fastapi import APIRouter, FastAPI, HTTPException, status from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import JSONResponse, Response @@ -42,8 +43,7 @@ from common.constants import ( 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 text.user_dict import apply_word, delete_word, read_dict, rewrite_word, update_dict # ---フロントエンド部分に関する処理--- @@ -229,7 +229,7 @@ async def normalize_text(item: TextRequest): @router.get("/models_info") def models_info(): - return model_holder.models_info() + return model_holder.models_info class SynthesisRequest(BaseModel): @@ -249,6 +249,7 @@ class SynthesisRequest(BaseModel): silenceAfter: float = 0.5 pitchScale: float = 1.0 intonationScale: float = 1.0 + speaker: Optional[str] = None @router.post("/synthesis", response_class=AudioResponse) @@ -274,6 +275,13 @@ def synthesis(request: SynthesisRequest): ] phone_tone = kata_tone2phone_tone(kata_tone_list) tone = [t for _, t in phone_tone] + try: + sid = 0 if request.speaker is None else model.spk2id[request.speaker] + except KeyError: + raise HTTPException( + status_code=400, + detail=f"Speaker {request.speaker} not found in {model.spk2id}", + ) sr, audio = model.infer( text=text, language=request.language.value, @@ -290,6 +298,7 @@ def synthesis(request: SynthesisRequest): line_split=False, pitch_scale=request.pitchScale, intonation_scale=request.intonationScale, + sid=sid, ) with BytesIO() as wavContent: From 1cbe9648b0f103e803b88c1d3976e37571b45da3 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Fri, 8 Mar 2024 12:59:27 +0900 Subject: [PATCH 13/14] Initialize pyopenjtalk worker for multi-threading --- bert_gen.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/bert_gen.py b/bert_gen.py index a5f7c25..1e4fb61 100644 --- a/bert_gen.py +++ b/bert_gen.py @@ -6,12 +6,15 @@ import torch.multiprocessing as mp from tqdm import tqdm import commons +import text.pyopenjtalk_worker as pyopenjtalk import utils from common.log import logger from common.stdout_wrapper import SAFE_STDOUT from config import config from text import cleaned_text_to_sequence, get_bert +pyopenjtalk.initialize() + def process_line(x): line, add_blank = x From 9f01e54d2dc72f1c62aab504ec4795e3ef619158 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Fri, 8 Mar 2024 13:00:15 +0900 Subject: [PATCH 14/14] Unify webui to single webui --- app.py | 61 ++++++---- common/constants.py | 2 +- webui/dataset.py | 40 ++----- webui/inference.py | 260 +++++++++++++++++------------------------ webui/merge.py | 30 +---- webui/style_vectors.py | 26 +---- webui/train.py | 29 +---- 7 files changed, 172 insertions(+), 276 deletions(-) diff --git a/app.py b/app.py index 7057f45..997bd04 100644 --- a/app.py +++ b/app.py @@ -1,30 +1,51 @@ -import pyopenjtalk -import gradio as gr -from webui import ( - create_dataset_app, - create_train_app, - create_merge_app, - create_style_vectors_app, -) +import argparse from pathlib import Path -pyopenjtalk.unset_user_dict() +import gradio as gr +import torch +import yaml -setting_json = Path("webui/setting.json") +from common.constants import GRADIO_THEME, LATEST_VERSION +from common.tts_model import ModelHolder +from webui import ( + create_dataset_app, + create_inference_app, + create_merge_app, + create_style_vectors_app, + create_train_app, +) -with gr.Blocks() as 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()) + # dataset_root = path_config["dataset_root"] + assets_root = path_config["assets_root"] + +parser = argparse.ArgumentParser() +parser.add_argument("--device", type=str, default="cuda") +parser.add_argument("--no_autolaunch", action="store_true") +parser.add_argument("--share", action="store_true") + +args = parser.parse_args() +device = args.device +if device == "cuda" and not torch.cuda.is_available(): + device = "cpu" + +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})") with gr.Tabs(): - with gr.Tab("Hello"): - gr.Markdown("## Hello, Gradio!") - gr.Textbox("input", label="Input Text") - with gr.Tab("Dataset"): + with gr.Tab("音声合成"): + create_inference_app(model_holder=model_holder) + with gr.Tab("データセット作成"): create_dataset_app() - with gr.Tab("Train"): + with gr.Tab("学習"): create_train_app() - with gr.Tab("Merge"): - create_merge_app() - with gr.Tab("Create Style Vectors"): + with gr.Tab("スタイル作成"): create_style_vectors_app() + with gr.Tab("マージ"): + create_merge_app(model_holder=model_holder) -app.launch(inbrowser=True) +app.launch(inbrowser=not args.no_autolaunch, share=args.share) diff --git a/common/constants.py b/common/constants.py index fe62019..f3ccb69 100644 --- a/common/constants.py +++ b/common/constants.py @@ -4,7 +4,7 @@ import enum # See https://huggingface.co/spaces/gradio/theme-gallery for more themes GRADIO_THEME: str = "NoCrypt/miku" -LATEST_VERSION: str = "2.3.1" +LATEST_VERSION: str = "2.4" USER_DICT_DIR = "dict_data" diff --git a/webui/dataset.py b/webui/dataset.py index 5ed656c..3169fd2 100644 --- a/webui/dataset.py +++ b/webui/dataset.py @@ -8,12 +8,6 @@ from common.constants import GRADIO_THEME from common.log import logger from common.subprocess_utils 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()) - dataset_root = path_config["dataset_root"] - # assets_root = path_config["assets_root"] - def do_slice( model_name: str, @@ -69,13 +63,11 @@ def do_transcribe( ] ) if not success: - return f"Error: {message}" + return f"Error: {message}. しかし何故かエラーが起きても正常に終了している場合がほとんどなので、書き起こし結果を確認して問題なければ学習に使えます。" return "音声の文字起こしが完了しました。" -initial_md = """ -# 簡易学習用データセット作成ツール - +how_to_md = """ Style-Bert-VITS2の学習用データセットを作成するためのツールです。以下の2つからなります。 - 与えられた音声からちょうどいい長さの発話区間を切り取りスライス @@ -107,10 +99,10 @@ Style-Bert-VITS2の学習用データセットを作成するためのツール """ -def create_dataset_app(): - - with gr.Blocks(theme=GRADIO_THEME) as app: - gr.Markdown(initial_md) +def create_dataset_app() -> gr.Blocks: + with gr.Blocks() as app: + with gr.Accordion("使い方", open=False): + gr.Markdown(how_to_md) model_name = gr.Textbox( label="モデル名を入力してください(話者名としても使われます)。" ) @@ -118,8 +110,8 @@ def create_dataset_app(): with gr.Row(): with gr.Column(): input_dir = gr.Textbox( - label="入力フォルダ名(デフォルトはinputs)", - placeholder="inputs", + label="元音声の入っているフォルダパス", + value="inputs", info="下記フォルダにwavファイルを入れておいてください", ) min_sec = gr.Slider( @@ -201,20 +193,4 @@ def create_dataset_app(): outputs=[result2], ) - parser = argparse.ArgumentParser() - parser.add_argument( - "--server-name", - type=str, - default=None, - help="Server name for Gradio app", - ) - parser.add_argument( - "--no-autolaunch", - action="store_true", - default=False, - help="Do not launch app automatically", - ) - args = parser.parse_args() - - # app.launch(inbrowser=not args.no_autolaunch, server_name=args.server_name) return app diff --git a/webui/inference.py b/webui/inference.py index 663d5a6..94124b9 100644 --- a/webui/inference.py +++ b/webui/inference.py @@ -1,14 +1,9 @@ import argparse import datetime import json -import os -import sys -from pathlib import Path from typing import Optional import gradio as gr -import torch -import yaml from common.constants import ( DEFAULT_ASSIST_TEXT_WEIGHT, @@ -21,7 +16,6 @@ from common.constants import ( DEFAULT_STYLE, DEFAULT_STYLE_WEIGHT, GRADIO_THEME, - LATEST_VERSION, Languages, ) from common.log import logger @@ -29,119 +23,9 @@ from common.tts_model import ModelHolder from infer import InvalidToneError from text.japanese import g2kata_tone, kata_tone2phone_tone, text_normalize -# 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()) - # dataset_root = path_config["dataset_root"] - assets_root = path_config["assets_root"] - languages = [l.value for l in Languages] -def tts_fn( - model_name, - model_path, - text, - language, - reference_audio_path, - sdp_ratio, - noise_scale, - noise_scale_w, - length_scale, - line_split, - split_interval, - assist_text, - assist_text_weight, - use_assist_text, - style, - style_weight, - kata_tone_json_str, - use_tone, - speaker, - pitch_scale, - intonation_scale, -): - model_holder.load_model_gr(model_name, model_path) - - wrong_tone_message = "" - kata_tone: Optional[list[tuple[str, int]]] = None - if use_tone and kata_tone_json_str != "": - if language != "JP": - logger.warning("Only Japanese is supported for tone generation.") - wrong_tone_message = "アクセント指定は現在日本語のみ対応しています。" - if line_split: - logger.warning("Tone generation is not supported for line split.") - wrong_tone_message = ( - "アクセント指定は改行で分けて生成を使わない場合のみ対応しています。" - ) - try: - kata_tone = [] - json_data = json.loads(kata_tone_json_str) - # tupleを使うように変換 - for kana, tone in json_data: - assert isinstance(kana, str) and tone in (0, 1), f"{kana}, {tone}" - kata_tone.append((kana, tone)) - except Exception as e: - logger.warning(f"Error occurred when parsing kana_tone_json: {e}") - wrong_tone_message = f"アクセント指定が不正です: {e}" - kata_tone = None - - # toneは実際に音声合成に代入される際のみnot Noneになる - tone: Optional[list[int]] = None - if kata_tone is not None: - phone_tone = kata_tone2phone_tone(kata_tone) - tone = [t for _, t in phone_tone] - - speaker_id = model_holder.current_model.spk2id[speaker] - - start_time = datetime.datetime.now() - - assert model_holder.current_model is not None - - try: - sr, audio = model_holder.current_model.infer( - text=text, - language=language, - reference_audio_path=reference_audio_path, - sdp_ratio=sdp_ratio, - noise=noise_scale, - noisew=noise_scale_w, - length=length_scale, - line_split=line_split, - split_interval=split_interval, - assist_text=assist_text, - assist_text_weight=assist_text_weight, - use_assist_text=use_assist_text, - style=style, - style_weight=style_weight, - given_tone=tone, - sid=speaker_id, - pitch_scale=pitch_scale, - intonation_scale=intonation_scale, - ) - except InvalidToneError as e: - logger.error(f"Tone error: {e}") - return f"Error: アクセント指定が不正です:\n{e}", None, kata_tone_json_str - except ValueError as e: - logger.error(f"Value error: {e}") - return f"Error: {e}", None, kata_tone_json_str - - end_time = datetime.datetime.now() - duration = (end_time - start_time).total_seconds() - - if tone is None and language == "JP": - # アクセント指定に使えるようにアクセント情報を返す - norm_text = text_normalize(text) - kata_tone = g2kata_tone(norm_text) - kata_tone_json_str = json.dumps(kata_tone, ensure_ascii=False) - elif tone is None: - kata_tone_json_str = "" - message = f"Success, time: {duration} seconds." - if wrong_tone_message != "": - message = wrong_tone_message + "\n" + message - return message, (sr, audio), kata_tone_json_str - - initial_text = "こんにちは、初めまして。あなたの名前はなんていうの?" examples = [ @@ -202,9 +86,7 @@ examples = [ ] initial_md = f""" -# Style-Bert-VITS2 ver {LATEST_VERSION} 音声合成 - -- Ver 2.3で追加されたエディターのほうが実際に読み上げさせるには使いやすいかもしれません。`Editor.bat`か`python server_editor.py`で起動できます。 +- Ver 2.3で追加されたエディターのほうが実際に読み上げさせるには使いやすいかもしれません。`Editor.bat`か`python server_editor.py --inbrowser`で起動できます。 - 初期からある[jvnvのモデル](https://huggingface.co/litagin/style_bert_vits2_jvnv)は、[JVNVコーパス(言語音声と非言語音声を持つ日本語感情音声コーパス)](https://sites.google.com/site/shinnosuketakamichi/research-topics/jvnv_corpus)で学習されたモデルです。ライセンスは[CC BY-SA 4.0](https://creativecommons.org/licenses/by-sa/4.0/deed.ja)です。 """ @@ -254,43 +136,119 @@ def gr_util(item): return (gr.update(visible=False), gr.update(visible=True)) -def create_inference_app(): - parser = argparse.ArgumentParser() - parser.add_argument("--cpu", action="store_true", help="Use CPU instead of GPU") - parser.add_argument( - "--dir", "-d", type=str, help="Model directory", default=assets_root - ) - parser.add_argument( - "--share", action="store_true", help="Share this app publicly", default=False - ) - parser.add_argument( - "--server-name", - type=str, - default=None, - help="Server name for Gradio app", - ) - parser.add_argument( - "--no-autolaunch", - action="store_true", - default=False, - help="Do not launch app automatically", - ) - args = parser.parse_args() - model_dir = Path(args.dir) +def create_inference_app(model_holder: ModelHolder) -> gr.Blocks: + def tts_fn( + model_name, + model_path, + text, + language, + reference_audio_path, + sdp_ratio, + noise_scale, + noise_scale_w, + length_scale, + line_split, + split_interval, + assist_text, + assist_text_weight, + use_assist_text, + style, + style_weight, + kata_tone_json_str, + use_tone, + speaker, + pitch_scale, + intonation_scale, + ): + model_holder.load_model(model_name, model_path) + assert model_holder.current_model is not None - if args.cpu: - device = "cpu" - else: - device = "cuda" if torch.cuda.is_available() else "cpu" + wrong_tone_message = "" + kata_tone: Optional[list[tuple[str, int]]] = None + if use_tone and kata_tone_json_str != "": + if language != "JP": + logger.warning("Only Japanese is supported for tone generation.") + wrong_tone_message = "アクセント指定は現在日本語のみ対応しています。" + if line_split: + logger.warning("Tone generation is not supported for line split.") + wrong_tone_message = ( + "アクセント指定は改行で分けて生成を使わない場合のみ対応しています。" + ) + try: + kata_tone = [] + json_data = json.loads(kata_tone_json_str) + # tupleを使うように変換 + for kana, tone in json_data: + assert isinstance(kana, str) and tone in (0, 1), f"{kana}, {tone}" + kata_tone.append((kana, tone)) + except Exception as e: + logger.warning(f"Error occurred when parsing kana_tone_json: {e}") + wrong_tone_message = f"アクセント指定が不正です: {e}" + kata_tone = None - model_holder = ModelHolder(model_dir, device) + # toneは実際に音声合成に代入される際のみnot Noneになる + tone: Optional[list[int]] = None + if kata_tone is not None: + phone_tone = kata_tone2phone_tone(kata_tone) + tone = [t for _, t in phone_tone] + + speaker_id = model_holder.current_model.spk2id[speaker] + + start_time = datetime.datetime.now() + + try: + sr, audio = model_holder.current_model.infer( + text=text, + language=language, + reference_audio_path=reference_audio_path, + sdp_ratio=sdp_ratio, + noise=noise_scale, + noisew=noise_scale_w, + length=length_scale, + line_split=line_split, + split_interval=split_interval, + assist_text=assist_text, + assist_text_weight=assist_text_weight, + use_assist_text=use_assist_text, + style=style, + style_weight=style_weight, + given_tone=tone, + sid=speaker_id, + pitch_scale=pitch_scale, + intonation_scale=intonation_scale, + ) + except InvalidToneError as e: + logger.error(f"Tone error: {e}") + return f"Error: アクセント指定が不正です:\n{e}", None, kata_tone_json_str + except ValueError as e: + logger.error(f"Value error: {e}") + return f"Error: {e}", None, kata_tone_json_str + + end_time = datetime.datetime.now() + duration = (end_time - start_time).total_seconds() + + if tone is None and language == "JP": + # アクセント指定に使えるようにアクセント情報を返す + norm_text = text_normalize(text) + kata_tone = g2kata_tone(norm_text) + kata_tone_json_str = json.dumps(kata_tone, ensure_ascii=False) + elif tone is None: + kata_tone_json_str = "" + message = f"Success, time: {duration} seconds." + if wrong_tone_message != "": + message = wrong_tone_message + "\n" + message + return message, (sr, audio), kata_tone_json_str model_names = model_holder.model_names if len(model_names) == 0: logger.error( - f"モデルが見つかりませんでした。{model_dir}にモデルを置いてください。" + f"モデルが見つかりませんでした。{model_holder.root_dir}にモデルを置いてください。" ) - sys.exit(1) + with gr.Blocks() as app: + gr.Markdown( + f"Error: モデルが見つかりませんでした。{model_holder.root_dir}にモデルを置いてください。" + ) + return app initial_id = 0 initial_pth_files = model_holder.model_files_dict[model_names[initial_id]] @@ -497,8 +455,4 @@ def create_inference_app(): outputs=[style, ref_audio_path], ) - # app.launch( - # inbrowser=not args.no_autolaunch, share=args.share, server_name=args.server_name - # ) - return app diff --git a/webui/merge.py b/webui/merge.py index c002e33..c9386ef 100644 --- a/webui/merge.py +++ b/webui/merge.py @@ -1,7 +1,6 @@ import argparse import json import os -import sys from pathlib import Path import gradio as gr @@ -28,8 +27,6 @@ with open(os.path.join("configs", "paths.yml"), "r", encoding="utf-8") as f: # dataset_root = path_config["dataset_root"] assets_root = path_config["assets_root"] -model_holder = ModelHolder(Path(assets_root), device) - def merge_style(model_name_a, model_name_b, weight, output_name, style_triple_list): """ @@ -331,13 +328,17 @@ Happy, Surprise, HappySurprise """ -def create_merge_app(): +def create_merge_app(model_holder: ModelHolder) -> gr.Blocks: model_names = model_holder.model_names if len(model_names) == 0: logger.error( f"モデルが見つかりませんでした。{assets_root}にモデルを置いてください。" ) - sys.exit(1) + with gr.Blocks() as app: + gr.Markdown( + f"Error: モデルが見つかりませんでした。{assets_root}にモデルを置いてください。" + ) + return app initial_id = 0 initial_model_files = model_holder.model_files_dict[model_names[initial_id]] @@ -499,23 +500,4 @@ def create_merge_app(): outputs=[audio_output], ) - parser = argparse.ArgumentParser() - parser.add_argument( - "--server-name", - type=str, - default=None, - help="Server name for Gradio app", - ) - parser.add_argument( - "--no-autolaunch", - action="store_true", - default=False, - help="Do not launch app automatically", - ) - parser.add_argument("--share", action="store_true", default=False) - args = parser.parse_args() - - # app.launch( - # inbrowser=not args.no_autolaunch, server_name=args.server_name, share=args.share - # ) return app diff --git a/webui/style_vectors.py b/webui/style_vectors.py index 27b95bb..8effa37 100644 --- a/webui/style_vectors.py +++ b/webui/style_vectors.py @@ -277,9 +277,7 @@ def save_style_vectors_from_files( return f"成功!\n{style_vector_path}に保存し{config_path}を更新しました。" -initial_md = f""" -# Style Bert-VITS2 スタイルベクトルの作成 - +how_to_md = f""" Style-Bert-VITS2でこまかくスタイルを指定して音声合成するには、モデルごとにスタイルベクトルのファイル`style_vectors.npy`を手動で作成する必要があります。 ただし、学習の過程で自動的に平均スタイル「{DEFAULT_STYLE}」のみは作成されるので、それをそのまま使うこともできます(その場合はこのWebUIは使いません)。 @@ -326,7 +324,8 @@ https://ja.wikipedia.org/wiki/DBSCAN def create_style_vectors_app(): with gr.Blocks(theme=GRADIO_THEME) as app: - gr.Markdown(initial_md) + with gr.Accordion("使い方", open=False): + gr.Markdown(how_to_md) with gr.Row(): model_name = gr.Textbox(placeholder="your_model_name", label="モデル名") reduction_method = gr.Radio( @@ -461,23 +460,4 @@ def create_style_vectors_app(): outputs=[info2], ) - parser = argparse.ArgumentParser() - parser.add_argument( - "--server-name", - type=str, - default=None, - help="Server name for Gradio app", - ) - parser.add_argument( - "--no-autolaunch", - action="store_true", - default=False, - help="Do not launch app automatically", - ) - parser.add_argument("--share", action="store_true", default=False) - args = parser.parse_args() - - # app.launch( - # inbrowser=not args.no_autolaunch, server_name=args.server_name, share=args.share - # ) return app diff --git a/webui/train.py b/webui/train.py index 1c9a214..dee018e 100644 --- a/webui/train.py +++ b/webui/train.py @@ -14,7 +14,6 @@ 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 common.subprocess_utils import run_script_with_log, second_elem_of @@ -398,9 +397,7 @@ def run_tensorboard(model_name): yield gr.Button("Tensorboardを開く") -initial_md = f""" -# Style-Bert-VITS2 ver {LATEST_VERSION} 学習用WebUI - +how_to_md = f""" ## 使い方 - データを準備して、モデル名を入力して、必要なら設定を調整してから、「自動前処理を実行」ボタンを押してください。進捗状況等はターミナルに表示されます。 @@ -452,10 +449,11 @@ english_teacher.wav|Mary|EN|How are you? I'm fine, thank you, and you? def create_train_app(): - with gr.Blocks(theme=GRADIO_THEME).queue() as app: - gr.Markdown(initial_md) - with gr.Accordion(label="データの前準備", open=False): - gr.Markdown(prepare_md) + with gr.Blocks().queue() as app: + with gr.Accordion("使い方", open=False): + gr.Markdown(how_to_md) + with gr.Accordion(label="データの前準備", open=False): + gr.Markdown(prepare_md) model_name = gr.Textbox(label="モデル名") gr.Markdown("### 自動前処理") with gr.Row(variant="panel"): @@ -794,20 +792,5 @@ def create_train_app(): outputs=[use_jp_extra_train], ) - parser = argparse.ArgumentParser() - parser.add_argument( - "--server-name", - type=str, - default=None, - help="Server name for Gradio app", - ) - parser.add_argument( - "--no-autolaunch", - action="store_true", - default=False, - help="Do not launch app automatically", - ) - args = parser.parse_args() - # app.launch(inbrowser=not args.no_autolaunch, server_name=args.server_name) return app