Improve WebUI and add loudness normalization

This commit is contained in:
litagin02
2023-12-29 18:43:14 +09:00
parent b0fbcc4667
commit fc2bde673a
7 changed files with 245 additions and 148 deletions

View File

@@ -46,7 +46,7 @@ python initialize.py # 必要なモデルとデフォルトTTSモデルをダ
### 音声合成 ### 音声合成
`App.bat`をダブルクリックするとWebUIが起動します。インストール時にデフォルトのモデルがダウンロードされているので、学習していなくてもそれを使うことができます。 `App.bat`をダブルクリックか、`python app.py`するとWebUIが起動します。インストール時にデフォルトのモデルがダウンロードされているので、学習していなくてもそれを使うことができます。
音声合成に必要なモデルファイルたちの構造は以下の通りです(手動で配置する必要はありません)。 音声合成に必要なモデルファイルたちの構造は以下の通りです(手動で配置する必要はありません)。
``` ```
@@ -64,18 +64,20 @@ model_assets
### 学習 ### 学習
`Train.bat`をダブルクリックするとWebUIが起動します。 `Train.bat`をダブルクリック`python webui_train.py`するとWebUIが起動します。
### スタイルの生成 ### スタイルの生成
- `Style.bat`をダブルクリックするとWebUIが起動します。 - `Style.bat`をダブルクリック`python webui_style_vectors.py`するとWebUIが起動します。
- この手順は、音声ファイルたちからスタイルを作るのに必要な手順です。 - この手順は、音声ファイルたちからスタイルを作るのに必要な手順です。
- 学習とは独立しているので、学習中でもできるし、学習が終わっても何度もやりなおせます。 - 学習とは独立しているので、学習中でもできるし、学習が終わっても何度もやりなおせます。
- スタイルについての詳細は[clustering.ipynb](clustering.ipynb)を参照してください。 - スタイルについての詳細は[clustering.ipynb](clustering.ipynb)を参照してください。
### データセット作り ### データセット作り
- `Dataset.bat`をダブルクリックすると、音声ファイルからデータセットを作るためのWebUIが起動します。音声ファイルのみからでもこれを使って学習できます。 - `Dataset.bat`をダブルクリック`python webui_dataset.py`すると、音声ファイルからデータセットを作るためのWebUIが起動します。音声ファイルのみからでもこれを使って学習できます。
注意: データセットの手動修正やノイズ除去や、より高品質なデータセットを作りたい場合は、[Aivis](https://github.com/tsukumijima/Aivis)や、そのデータセット部分のWindows対応版 [Aivis Dataset](https://github.com/litagin02/Aivis-Dataset) を使うのをおすすめします。
## Bert-VITS2 v2.1との関係 ## Bert-VITS2 v2.1との関係

View File

@@ -16,6 +16,7 @@ numba
numpy numpy
psutil psutil
pyannote.audio>=3.1.0 pyannote.audio>=3.1.0
pyloudnorm
pyopenjtalk-prebuilt pyopenjtalk-prebuilt
pypinyin pypinyin
PyYAML PyYAML

View File

@@ -4,10 +4,20 @@ import sys
from multiprocessing import Pool, cpu_count from multiprocessing import Pool, cpu_count
import librosa import librosa
import pyloudnorm as pyln
import soundfile import soundfile
from tqdm import tqdm from tqdm import tqdm
from config import config from config import config
from tools.log import logger
def normalize_audio(data, sr):
meter = pyln.Meter(sr) # create BS.1770 meter
loudness = meter.integrated_loudness(data)
# logger.info(f"loudness: {loudness}")
data = pyln.normalize.loudness(data, loudness, -23.0)
return data
def process(item): def process(item):
@@ -15,6 +25,8 @@ def process(item):
wav_path = os.path.join(args.in_dir, spkdir, wav_name) wav_path = os.path.join(args.in_dir, spkdir, wav_name)
if os.path.exists(wav_path) and wav_path.lower().endswith(".wav"): if os.path.exists(wav_path) and wav_path.lower().endswith(".wav"):
wav, sr = librosa.load(wav_path, sr=args.sr) wav, sr = librosa.load(wav_path, sr=args.sr)
if args.normalize:
wav = normalize_audio(wav, sr)
soundfile.write(os.path.join(args.out_dir, spkdir, wav_name), wav, sr) soundfile.write(os.path.join(args.out_dir, spkdir, wav_name), wav, sr)
@@ -44,13 +56,18 @@ if __name__ == "__main__":
default=4, default=4,
help="cpu_processes", help="cpu_processes",
) )
parser.add_argument(
"--normalize",
action="store_true",
default=True,
help="loudness normalize audio",
)
args, _ = parser.parse_known_args() args, _ = parser.parse_known_args()
# autodl 无卡模式会识别出46个cpu # autodl 无卡模式会识别出46个cpu
if args.num_processes == 0: if args.num_processes == 0:
processes = cpu_count() - 2 if cpu_count() > 4 else 1 processes = cpu_count() - 2 if cpu_count() > 4 else 1
else: else:
processes = args.num_processes processes = args.num_processes
pool = Pool(processes=processes)
tasks = [] tasks = []
@@ -66,7 +83,10 @@ if __name__ == "__main__":
tasks.append(twople) tasks.append(twople)
if len(tasks) == 0: if len(tasks) == 0:
logger.error(f"No wav files found in {args.in_dir}")
raise ValueError(f"No wav files found in {args.in_dir}") raise ValueError(f"No wav files found in {args.in_dir}")
pool = Pool(processes=processes)
for _ in tqdm( for _ in tqdm(
pool.imap_unordered(process, tasks), file=sys.stdout, total=len(tasks) pool.imap_unordered(process, tasks), file=sys.stdout, total=len(tasks)
): ):

View File

@@ -3,8 +3,8 @@ import os
import shutil import shutil
import sys import sys
import soundfile as sf
import torch import torch
from pydub import AudioSegment
from tqdm import tqdm from tqdm import tqdm
vad_model, utils = torch.hub.load( vad_model, utils = torch.hub.load(
@@ -52,16 +52,9 @@ def split_wav(
audio_file, min_silence_dur_ms=min_silence_dur_ms, min_sec=min_sec audio_file, min_silence_dur_ms=min_silence_dur_ms, min_sec=min_sec
) )
# WAVファイルを読み込む data, sr = sf.read(audio_file)
audio = AudioSegment.from_wav(audio_file)
# リサンプリング44100Hz total_ms = len(data) / sr * 1000
audio = audio.set_frame_rate(44100)
# ステレオをモノラルに変換
audio = audio.set_channels(1)
total_ms = len(audio)
file_name = os.path.basename(audio_file).split(".")[0] file_name = os.path.basename(audio_file).split(".")[0]
os.makedirs(target_dir, exist_ok=True) os.makedirs(target_dir, exist_ok=True)
@@ -74,8 +67,15 @@ def split_wav(
end_ms = min(ts["end"] / 16 + margin, total_ms) end_ms = min(ts["end"] / 16 + margin, total_ms)
if end_ms - start_ms > upper_bound_ms: if end_ms - start_ms > upper_bound_ms:
continue continue
segment = audio[start_ms:end_ms]
segment.export(os.path.join(target_dir, f"{file_name}-{i}.wav"), format="wav") start_sample = int(start_ms / 1000 * sr)
end_sample = int(end_ms / 1000 * sr)
segment = data[start_sample:end_sample]
if normalize:
segment = normalize_audio(segment, sr)
sf.write(os.path.join(target_dir, f"{file_name}-{i}.wav"), segment, sr)
total_time_ms += end_ms - start_ms total_time_ms += end_ms - start_ms
return total_time_ms / 1000 return total_time_ms / 1000

View File

@@ -658,8 +658,9 @@ def train_and_evaluate(
) )
global_step += 1 global_step += 1
gc.collect() # 本家ではこれをスピードアップのために消すと書かれていたので、一応消してみる
torch.cuda.empty_cache() # gc.collect()
# torch.cuda.empty_cache()
if rank == 0: if rank == 0:
logger.info(f"====> Epoch: {epoch}, step: {global_step}") logger.info(f"====> Epoch: {epoch}, step: {global_step}")

View File

@@ -1,43 +1,35 @@
import os import os
import subprocess
import sys
import gradio as gr import gradio as gr
python = sys.executable from tools.log import logger
from tools.subprocess_utils import run_script_with_log, second_elem_of
def subprocess_wrapper(cmd): def do_slice(model_name, normalize):
return subprocess.run( logger.info("Start slicing...")
cmd,
stdout=sys.stdout,
stderr=subprocess.PIPE,
text=True,
)
def do_slice(model_name):
input_dir = "inputs" input_dir = "inputs"
output_dir = os.path.join("Data", model_name, "raw") output_dir = os.path.join("Data", model_name, "raw")
result = subprocess_wrapper( cmd = [
[ "slice.py",
python, "--input_dir",
"slice.py", input_dir,
"--input_dir", "--output_dir",
input_dir, output_dir,
"--output_dir", ]
output_dir, if normalize:
] cmd.append("--normalize")
) success, message = run_script_with_log(cmd)
return "ターミナルを見て結果を確認してください。" if not success:
return f"Error: {message}"
return "音声のスライスが完了しました。"
def do_transcribe(model_name): def do_transcribe(model_name):
input_dir = os.path.join("Data", model_name, "raw") input_dir = os.path.join("Data", model_name, "raw")
output_file = os.path.join("Data", model_name, "esd.list") output_file = os.path.join("Data", model_name, "esd.list")
result = subprocess_wrapper( result = run_script_with_log(
[ [
python,
"transcribe.py", "transcribe.py",
"--input_dir", "--input_dir",
input_dir, input_dir,
@@ -47,13 +39,13 @@ def do_transcribe(model_name):
model_name, model_name,
] ]
) )
if result.stderr:
return f"{result.stderr}"
return "音声の文字起こしが完了しました。" return "音声の文字起こしが完了しました。"
initial_md = """ initial_md = """
# 学習用データセット作成ツール # 簡易学習用データセット作成ツール
**注意**:より精密で高品質なデータセットを作成したい・書き起こしをいろいろ修正したい場合は、[Aivis Dataset](https://github.com/litagin02/Aivis-Dataset)をおすすめします。書き起こし部分もかなり工夫されています。このツールはあくまでスライスして書き起こすという簡易的なことしかしていません。
Style-Bert-VITS2の学習用データセットを作成するためのツールです。与えられた音声からちょうどいい長さの発話区間を切り取りスライスし、それぞれの音声に対して文字起こしを行います。 Style-Bert-VITS2の学習用データセットを作成するためのツールです。与えられた音声からちょうどいい長さの発話区間を切り取りスライスし、それぞれの音声に対して文字起こしを行います。
@@ -69,23 +61,24 @@ Style-Bert-VITS2の学習用データセットを作成するためのツール
細かいパラメータ調整とかがしたい人は、`slice.py`と`transcribe.py`を眺めて直接実行してください。 細かいパラメータ調整とかがしたい人は、`slice.py`と`transcribe.py`を眺めて直接実行してください。
また、出来上がった音声ファイルたちは`Data/{モデル名}/raw`に、書き起こしファイルは`Data/{モデル名}/esd.list`に保存されます。 また、出来上がった音声ファイルたちは`Data/{モデル名}/raw`に、書き起こしファイルは`Data/{モデル名}/esd.list`に保存されます。
書き起こしの結果は、**そこまで正確に誤字や誤りを修正しなくても、それなりの質になる**ので、あまり修正は必要ないかもしれません(私は手動修正したことないです 書き起こしの結果をどれだけ修正すればいいかはデータセットに依存しそうです。
**ffmpeg のインストールが別途必要のよう**です、「Couldn't find ffmpeg」とか怒られたら、「Windows ffmpeg インストール」等でググって別途インストールしてください。
""" """
with gr.Blocks(theme="NoCrypt/miku") as app: with gr.Blocks(theme="NoCrypt/miku") as app:
gr.Markdown(initial_md) gr.Markdown(initial_md)
model_name = gr.Textbox(label="モデル名を入力してください(話者名としても使われます)。") model_name = gr.Textbox(label="モデル名を入力してください(話者名としても使われます)。")
with gr.Row(): with gr.Accordion("音声のスライス"):
slice_button = gr.Button("音声のスライス") with gr.Row():
result1 = gr.Textbox(label="結果") with gr.Column():
normalize = gr.Checkbox(label="スライスされた音声の音量を正規化する", value=True)
slice_button = gr.Button("スライスを実行")
result1 = gr.Textbox(label="結果")
with gr.Row(): with gr.Row():
transcribe_button = gr.Button("音声の文字起こし") transcribe_button = gr.Button("音声の文字起こし")
result2 = gr.Textbox(label="結果") result2 = gr.Textbox(label="結果")
slice_button.click( slice_button.click(
do_slice, do_slice,
inputs=[model_name], inputs=[model_name, normalize],
outputs=[result1], outputs=[result1],
) )
transcribe_button.click( transcribe_button.click(

View File

@@ -2,7 +2,7 @@ import json
import os import os
import shutil import shutil
import sys import sys
from multiprocessing import Pool, cpu_count from multiprocessing import cpu_count
import gradio as gr import gradio as gr
import yaml import yaml
@@ -21,7 +21,7 @@ def get_path(model_name):
return dataset_path, lbl_path, train_path, val_path, config_path return dataset_path, lbl_path, train_path, val_path, config_path
def initialize(model_name, batch_size, epochs, bf16_run): def initialize(model_name, batch_size, epochs, save_every_steps, bf16_run):
logger.info("Step 1: start initialization...") logger.info("Step 1: start initialization...")
dataset_path, _, train_path, val_path, config_path = get_path(model_name) dataset_path, _, train_path, val_path, config_path = get_path(model_name)
if os.path.isfile(config_path): if os.path.isfile(config_path):
@@ -35,6 +35,7 @@ def initialize(model_name, batch_size, epochs, bf16_run):
config["train"]["batch_size"] = batch_size config["train"]["batch_size"] = batch_size
config["train"]["epochs"] = epochs config["train"]["epochs"] = epochs
config["train"]["bf16_run"] = bf16_run config["train"]["bf16_run"] = bf16_run
config["train"]["eval_interval"] = save_every_steps
model_path = os.path.join(dataset_path, "models") model_path = os.path.join(dataset_path, "models")
try: try:
@@ -61,22 +62,25 @@ def initialize(model_name, batch_size, epochs, bf16_run):
return True, "Step 1, Success: 初期設定が完了しました" return True, "Step 1, Success: 初期設定が完了しました"
def resample(model_name): def resample(model_name, normalize, num_processes):
logger.info("Step 2: start resampling...") logger.info("Step 2: start resampling...")
dataset_path, _, _, _, _ = get_path(model_name) dataset_path, _, _, _, _ = get_path(model_name)
in_dir = os.path.join(dataset_path, "raw") in_dir = os.path.join(dataset_path, "raw")
out_dir = os.path.join(dataset_path, "wavs") out_dir = os.path.join(dataset_path, "wavs")
success, message = run_script_with_log( cmd = [
[ "resample.py",
"resample.py", "--in_dir",
"--in_dir", in_dir,
in_dir, "--out_dir",
"--out_dir", out_dir,
out_dir, "--num_processes",
"--sr", str(num_processes),
"44100", "--sr",
] "44100",
) ]
if normalize:
cmd.append("--normalize")
success, message = run_script_with_log(cmd)
if not success: if not success:
logger.error(f"Step 2: resampling failed.") logger.error(f"Step 2: resampling failed.")
return False, f"Step 2, Error: 音声ファイルの前処理に失敗しました:\n{message}" return False, f"Step 2, Error: 音声ファイルの前処理に失敗しました:\n{message}"
@@ -126,8 +130,17 @@ def preprocess_text(model_name):
def bert_gen(model_name): def bert_gen(model_name):
logger.info("Step 4: start bert_gen...")
_, _, _, _, config_path = get_path(model_name) _, _, _, _, config_path = get_path(model_name)
success, message = run_script_with_log(["bert_gen.py", "--config", config_path]) success, message = run_script_with_log(
[
"bert_gen.py",
"--config",
config_path,
# "--num_processes", # bert_genは重いのでプロセス数いじらない
# str(num_processes),
]
)
if not success: if not success:
logger.error(f"Step 4: bert_gen failed.") logger.error(f"Step 4: bert_gen failed.")
return False, f"Step 4, Error: BERT特徴ファイルの生成に失敗しました:\n{message}" return False, f"Step 4, Error: BERT特徴ファイルの生成に失敗しました:\n{message}"
@@ -139,6 +152,7 @@ def bert_gen(model_name):
def style_gen(model_name, num_processes): def style_gen(model_name, num_processes):
logger.info("Step 5: start style_gen...")
_, _, _, _, config_path = get_path(model_name) _, _, _, _, config_path = get_path(model_name)
success, message = run_script_with_log( success, message = run_script_with_log(
[ [
@@ -159,23 +173,29 @@ def style_gen(model_name, num_processes):
return True, "Step 5, Success: スタイル特徴ファイルの生成が完了しました" return True, "Step 5, Success: スタイル特徴ファイルの生成が完了しました"
def preprocess_all(model_name, batch_size, epochs, bf16_run, num_processes): def preprocess_all(
success, message = initialize(model_name, batch_size, epochs, bf16_run) model_name, batch_size, epochs, save_every_steps, bf16_run, num_processes, normalize
):
if model_name == "":
return False, "Error: モデル名を入力してください"
success, message = initialize(
model_name, batch_size, epochs, save_every_steps, bf16_run
)
if not success: if not success:
return success, message return False, message
success, message = resample(model_name) success, message = resample(model_name, normalize, num_processes)
if not success: if not success:
return success, message return False, message
success, message = preprocess_text(model_name) success, message = preprocess_text(model_name)
if not success: if not success:
return success, message return False, message
success, message = bert_gen(model_name) success, message = bert_gen(model_name) # bert_genは重いのでプロセス数いじらない
if not success: if not success:
return success, message return False, message
success, message = style_gen(model_name, num_processes) success, message = style_gen(model_name, num_processes)
if not success: if not success:
return success, message return False, message
logger.success("Success: all preprocess finished.") logger.success("Success: All preprocess finished!")
return True, "Success: 全ての前処理が完了しました。ターミナルを確認しておかしいところがないか確認するのをおすすめします。" return True, "Success: 全ての前処理が完了しました。ターミナルを確認しておかしいところがないか確認するのをおすすめします。"
@@ -199,16 +219,20 @@ initial_md = """
## 使い方 ## 使い方
- データを準備して、各ステップを順に実行してください。進捗状況等はターミナルに表示されます。 - データを準備して、モデル名を入力して、必要なら設定を調整してから、「自動前処理を実行」ボタンを押してください。進捗状況等はターミナルに表示されます。
- 途中から学習を再開する場合は、モデル名を入力してFinal Stepだけ実行すればよいです - 各ステップごとに実行する場合は「手動前処理」を使ってください(基本的には自動でいいはず)
- 前処理が終わったら、「学習を開始する」ボタンを押すと学習が開始されます。
- 途中から学習を再開する場合は、モデル名を入力してから「学習を開始する」を押せばよいです。
注意: 音声合成で使うには、スタイルベクトルファイル`style_vectors.npy`を作る必要があります。これは、`Style.bat`を実行してそこで作成してください。 注意: 音声合成で使うには、スタイルベクトルファイル`style_vectors.npy`を作る必要があります。これは、`Style.bat`を実行してそこで作成してください。
動作は軽いはずなので、学習中でも実行でき、何度でも繰り返して試せます。 動作は軽いはずなので、学習中でも実行でき、何度でも繰り返して試せます。
""" """
prepare_md = """ prepare_md = """
まず音声データwavファイルで1ファイルが2-15秒程度の、長すぎず短すぎない発話のものをいくつか)と、書き起こしテキストを用意してください。 まず音声データwavファイルで1ファイルが2-12秒程度の、長すぎず短すぎない発話のものをいくつか)と、書き起こしテキストを用意してください。
それを次のように配置します。 それを次のように配置します。
``` ```
@@ -237,56 +261,64 @@ english_teacher.wav|Mary|EN|How are you? I'm fine, thank you, and you?
日本語話者の単一話者データセットでも構いません。 日本語話者の単一話者データセットでも構いません。
""" """
css = """
div.panel {
justify-content: space-between;
}
"""
if __name__ == "__main__": if __name__ == "__main__":
with gr.Blocks(theme="NoCrypt/miku", css=css) as app: with gr.Blocks(theme="NoCrypt/miku") as app:
gr.Markdown(initial_md) gr.Markdown(initial_md)
with gr.Accordion(label="データの前準備", open=False): with gr.Accordion(label="データの前準備", open=False):
gr.Markdown(prepare_md) gr.Markdown(prepare_md)
model_name = gr.Textbox( model_name = gr.Textbox(
label="モデル名", label="モデル名",
) )
info = gr.Textbox(label="状況")
gr.Markdown("### 自動前処理") gr.Markdown("### 自動前処理")
with gr.Row(variant="panel"): with gr.Row(variant="panel"):
batch_size = gr.Slider( with gr.Column():
label="バッチサイズ", batch_size = gr.Slider(
info="VRAM 12GBで4くらい", label="バッチサイズ",
value=4, info="VRAM 12GBで4くらい",
minimum=1, value=4,
maximum=64, minimum=1,
step=1, maximum=64,
) step=1,
epochs = gr.Slider( )
label="エポック数", epochs = gr.Slider(
info="100もあれば十分そう", label="エポック数",
value=100, info="100もあれば十分そう",
minimum=1, value=100,
maximum=1000, minimum=1,
step=1, maximum=1000,
) step=1,
bf16_run = gr.Checkbox( )
label="bfloat16を使う", save_every_steps = gr.Slider(
info="bfloat16を使うかどうか。新しめのグラボだと学習が早くなるかも、古いグラボだと動かないかも。", label="何ステップごとに結果を保存するか",
value=True, info="エポック数とは違うことに注意",
) value=1000,
num_processes = gr.Slider( minimum=100,
label="プロセス数", maximum=10000,
info="前処理時の並列処理プロセス数、大きすぎるとフリーズするかも", step=100,
value=cpu_count() // 2, )
minimum=1, bf16_run = gr.Checkbox(
maximum=cpu_count(), label="bfloat16を使う",
step=1, info="bfloat16を使うかどうか。新しめのグラボだと学習が早くなるかも、古いグラボだと動かないかも。",
) value=True,
preprocess_button = gr.Button(value="実行", variant="primary") )
num_processes = gr.Slider(
label="プロセス数",
info="前処理時の並列処理プロセス数、大きすぎるとフリーズするかも",
value=cpu_count() // 2,
minimum=1,
maximum=cpu_count(),
step=1,
)
normalize = gr.Checkbox(
label="音声の音量を正規化する",
value=True,
)
with gr.Column():
preprocess_button = gr.Button(value="自動前処理を実行", variant="primary")
info_all = gr.Textbox(label="状況")
with gr.Accordion(open=False, label="手動前処理"): with gr.Accordion(open=False, label="手動前処理"):
with gr.Row(variant="panel"): with gr.Row(variant="panel"):
with gr.Column(variant="panel", min_width=160): with gr.Column():
gr.Markdown(value="#### Step 1: 設定ファイルの生成") gr.Markdown(value="#### Step 1: 設定ファイルの生成")
batch_size_manual = gr.Slider( batch_size_manual = gr.Slider(
label="バッチサイズ", label="バッチサイズ",
@@ -302,12 +334,22 @@ if __name__ == "__main__":
maximum=1000, maximum=1000,
step=1, step=1,
) )
save_every_steps_manual = gr.Slider(
label="何ステップごとに結果を保存するか",
value=1000,
minimum=100,
maximum=10000,
step=100,
)
bf16_run_manual = gr.Checkbox( bf16_run_manual = gr.Checkbox(
label="bfloat16を使う", label="bfloat16を使う",
value=True, value=True,
) )
with gr.Column():
generate_config_btn = gr.Button(value="実行", variant="primary") generate_config_btn = gr.Button(value="実行", variant="primary")
with gr.Column(variant="panel", min_width=160): info_init = gr.Textbox(label="状況")
with gr.Row(variant="panel"):
with gr.Column():
gr.Markdown(value="#### Step 2: 音声ファイルの前処理") gr.Markdown(value="#### Step 2: 音声ファイルの前処理")
num_processes_resample = gr.Slider( num_processes_resample = gr.Slider(
label="プロセス数", label="プロセス数",
@@ -316,21 +358,35 @@ if __name__ == "__main__":
maximum=cpu_count(), maximum=cpu_count(),
step=1, step=1,
) )
resample_btn = gr.Button(value="実行", variant="primary") normalize_resample = gr.Checkbox(
with gr.Column(variant="panel", min_width=160): label="音声の音量を正規化する",
gr.Markdown(value="#### Step 3: 書き起こしファイルの前処理") value=True,
preprocess_text_btn = gr.Button(value="実行", variant="primary")
with gr.Column(variant="panel", min_width=160):
gr.Markdown(value="#### Step 4: BERT特徴ファイルの生成")
num_processes_bert = gr.Slider(
label="プロセス数",
value=cpu_count() // 2,
minimum=1,
maximum=cpu_count(),
step=1,
) )
with gr.Column():
resample_btn = gr.Button(value="実行", variant="primary")
info_resample = gr.Textbox(label="状況")
with gr.Row(variant="panel"):
with gr.Column():
gr.Markdown(value="#### Step 3: 書き起こしファイルの前処理")
with gr.Column():
preprocess_text_btn = gr.Button(value="実行", variant="primary")
info_preprocess_text = gr.Textbox(label="状況")
with gr.Row(variant="panel"):
with gr.Column():
gr.Markdown(value="#### Step 4: BERT特徴ファイルの生成")
# num_processes_bert = gr.Slider(
# label="プロセス数",
# value=cpu_count() // 2,
# minimum=1,
# maximum=cpu_count(),
# step=1,
# )
# bert_genは重いようで、初期の4がよいみたい
with gr.Column():
bert_gen_btn = gr.Button(value="実行", variant="primary") bert_gen_btn = gr.Button(value="実行", variant="primary")
with gr.Column(variant="panel", min_width=160): info_bert = gr.Textbox(label="状況")
with gr.Row(variant="panel"):
with gr.Column():
gr.Markdown(value="#### Step 5: スタイル特徴ファイルの生成") gr.Markdown(value="#### Step 5: スタイル特徴ファイルの生成")
num_processes_style = gr.Slider( num_processes_style = gr.Slider(
label="プロセス数", label="プロセス数",
@@ -339,36 +395,60 @@ if __name__ == "__main__":
maximum=cpu_count(), maximum=cpu_count(),
step=1, step=1,
) )
with gr.Column():
style_gen_btn = gr.Button(value="実行", variant="primary") style_gen_btn = gr.Button(value="実行", variant="primary")
info_style = gr.Textbox(label="状況")
gr.Markdown("## 学習")
with gr.Row(variant="panel"): with gr.Row(variant="panel"):
with gr.Column(): train_btn = gr.Button(value="学習を開始する", variant="primary")
gr.Markdown("## 学習") info_train = gr.Textbox(label="状況")
train_btn = gr.Button(value="学習を開始する", variant="primary")
preprocess_button.click( preprocess_button.click(
second_elem_of(preprocess_all), second_elem_of(preprocess_all),
inputs=[model_name, batch_size, epochs, bf16_run, num_processes], inputs=[
outputs=[info], model_name,
batch_size,
epochs,
save_every_steps,
bf16_run,
num_processes,
normalize,
],
outputs=[info_all],
) )
generate_config_btn.click( generate_config_btn.click(
second_elem_of(initialize), second_elem_of(initialize),
inputs=[model_name, batch_size_manual, epochs_manual, bf16_run_manual], inputs=[
outputs=[info], model_name,
batch_size_manual,
epochs_manual,
save_every_steps_manual,
bf16_run_manual,
],
outputs=[info_init],
) )
resample_btn.click( resample_btn.click(
second_elem_of(resample), inputs=[model_name], outputs=[info] second_elem_of(resample),
inputs=[model_name, normalize_resample, num_processes_resample],
outputs=[info_resample],
) )
preprocess_text_btn.click( preprocess_text_btn.click(
second_elem_of(preprocess_text), inputs=[model_name], outputs=[info] second_elem_of(preprocess_text),
inputs=[model_name],
outputs=[info_preprocess_text],
) )
bert_gen_btn.click( bert_gen_btn.click(
second_elem_of(bert_gen), inputs=[model_name], outputs=[info] second_elem_of(bert_gen),
inputs=[model_name],
outputs=[info_bert],
) )
style_gen_btn.click( style_gen_btn.click(
second_elem_of(style_gen), second_elem_of(style_gen),
inputs=[model_name, num_processes_style], inputs=[model_name, num_processes_style],
outputs=[info], outputs=[info_style],
)
train_btn.click(
second_elem_of(train), inputs=[model_name], outputs=[info_train]
) )
train_btn.click(second_elem_of(train), inputs=[model_name], outputs=[info])
app.launch(inbrowser=True) app.launch(inbrowser=True)