Merge branch 'dev-webui'
This commit is contained in:
10
README.md
10
README.md
@@ -48,7 +48,7 @@ python initialize.py # 必要なモデルとデフォルトTTSモデルをダ
|
|||||||
|
|
||||||
### 音声合成
|
### 音声合成
|
||||||
|
|
||||||
`App.bat`をダブルクリックするとWebUIが起動します。インストール時にデフォルトのモデルがダウンロードされているので、学習していなくてもそれを使うことができます。
|
`App.bat`をダブルクリックか、`python app.py`するとWebUIが起動します。インストール時にデフォルトのモデルがダウンロードされているので、学習していなくてもそれを使うことができます。
|
||||||
|
|
||||||
音声合成に必要なモデルファイルたちの構造は以下の通りです(手動で配置する必要はありません)。
|
音声合成に必要なモデルファイルたちの構造は以下の通りです(手動で配置する必要はありません)。
|
||||||
```
|
```
|
||||||
@@ -66,18 +66,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との関係
|
||||||
|
|
||||||
|
|||||||
25
app.py
25
app.py
@@ -14,9 +14,6 @@ from config import config
|
|||||||
from infer import get_net_g, infer
|
from infer import get_net_g, infer
|
||||||
from tools.log import logger
|
from tools.log import logger
|
||||||
|
|
||||||
is_hf_spaces = os.getenv("SYSTEM") == "spaces"
|
|
||||||
limit = 100
|
|
||||||
|
|
||||||
|
|
||||||
class Model:
|
class Model:
|
||||||
def __init__(self, model_path, config_path, style_vec_path, device):
|
def __init__(self, model_path, config_path, style_vec_path, device):
|
||||||
@@ -222,9 +219,6 @@ def tts_fn(
|
|||||||
emotion,
|
emotion,
|
||||||
emotion_weight,
|
emotion_weight,
|
||||||
):
|
):
|
||||||
if is_hf_spaces and len(text) > limit:
|
|
||||||
raise Exception(f"文字数が{limit}文字を超えています")
|
|
||||||
|
|
||||||
assert model_holder.current_model is not None
|
assert model_holder.current_model is not None
|
||||||
|
|
||||||
start_time = datetime.datetime.now()
|
start_time = datetime.datetime.now()
|
||||||
@@ -253,7 +247,7 @@ def tts_fn(
|
|||||||
|
|
||||||
initial_text = "こんにちは、初めまして。あなたの名前はなんていうの?"
|
initial_text = "こんにちは、初めまして。あなたの名前はなんていうの?"
|
||||||
|
|
||||||
example_local = [
|
examples = [
|
||||||
[initial_text, "JP"],
|
[initial_text, "JP"],
|
||||||
[
|
[
|
||||||
"""あなたがそんなこと言うなんて、私はとっても嬉しい。
|
"""あなたがそんなこと言うなんて、私はとっても嬉しい。
|
||||||
@@ -307,18 +301,6 @@ example_local = [
|
|||||||
["语音合成是人工制造人类语音。用于此目的的计算机系统称为语音合成器,可以通过软件或硬件产品实现。", "ZH"],
|
["语音合成是人工制造人类语音。用于此目的的计算机系统称为语音合成器,可以通过软件或硬件产品实现。", "ZH"],
|
||||||
]
|
]
|
||||||
|
|
||||||
example_hf_spaces = [
|
|
||||||
[initial_text, "JP"],
|
|
||||||
["えっと、私、あなたのことが好きです!もしよければ付き合ってくれませんか?", "JP"],
|
|
||||||
["吾輩は猫である。名前はまだ無い。", "JP"],
|
|
||||||
["どこで生れたかとんと見当がつかぬ。なんでも薄暗いじめじめした所でニャーニャー泣いていた事だけは記憶している。", "JP"],
|
|
||||||
["やったー!テストで満点取れたよ!私とっても嬉しいな!", "JP"],
|
|
||||||
["どうして私の意見を無視するの?許せない!ムカつく!あんたなんか死ねばいいのに。", "JP"],
|
|
||||||
["あはははっ!この漫画めっちゃ笑える、見てよこれ、ふふふ、あはは。", "JP"],
|
|
||||||
["あなたがいなくなって、私は一人になっちゃって、泣いちゃいそうなほど悲しい。", "JP"],
|
|
||||||
["深層学習の応用により、感情やアクセントを含む声質の微妙な変化も再現されている。", "JP"],
|
|
||||||
]
|
|
||||||
|
|
||||||
initial_md = """
|
initial_md = """
|
||||||
# Style-Bert-VITS2 音声合成
|
# Style-Bert-VITS2 音声合成
|
||||||
|
|
||||||
@@ -389,13 +371,12 @@ if __name__ == "__main__":
|
|||||||
model_holder = ModelHolder(model_dir, device)
|
model_holder = ModelHolder(model_dir, device)
|
||||||
|
|
||||||
languages = ["JP", "EN", "ZH"]
|
languages = ["JP", "EN", "ZH"]
|
||||||
examples = example_hf_spaces if is_hf_spaces else example_local
|
|
||||||
|
|
||||||
model_names = model_holder.model_names
|
model_names = model_holder.model_names
|
||||||
if len(model_names) == 0:
|
if len(model_names) == 0:
|
||||||
logger.error(f"モデルが見つかりませんでした。{model_dir}にモデルを置いてください。")
|
logger.error(f"モデルが見つかりませんでした。{model_dir}にモデルを置いてください。")
|
||||||
sys.exit(1)
|
sys.exit(1)
|
||||||
initial_id = 1 if is_hf_spaces else 0
|
initial_id = 0
|
||||||
initial_pth_files = model_holder.model_files_dict[model_names[initial_id]]
|
initial_pth_files = model_holder.model_files_dict[model_names[initial_id]]
|
||||||
|
|
||||||
with gr.Blocks(theme="NoCrypt/miku") as app:
|
with gr.Blocks(theme="NoCrypt/miku") as app:
|
||||||
@@ -416,7 +397,7 @@ if __name__ == "__main__":
|
|||||||
choices=initial_pth_files,
|
choices=initial_pth_files,
|
||||||
value=initial_pth_files[0],
|
value=initial_pth_files[0],
|
||||||
)
|
)
|
||||||
refresh_button = gr.Button("更新", scale=1, visible=not is_hf_spaces)
|
refresh_button = gr.Button("更新", scale=1, visible=True)
|
||||||
load_button = gr.Button("ロード", scale=1, variant="primary")
|
load_button = gr.Button("ロード", scale=1, variant="primary")
|
||||||
text_input = gr.TextArea(label="テキスト", value=initial_text)
|
text_input = gr.TextArea(label="テキスト", value=initial_text)
|
||||||
|
|
||||||
|
|||||||
@@ -97,7 +97,7 @@ class Style_gen_config:
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config_path: str,
|
config_path: str,
|
||||||
num_processes: int = 2,
|
num_processes: int = 4,
|
||||||
device: str = "cuda",
|
device: str = "cuda",
|
||||||
):
|
):
|
||||||
self.config_path = config_path
|
self.config_path = config_path
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
24
resample.py
24
resample.py
@@ -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 = []
|
||||||
|
|
||||||
@@ -65,6 +82,11 @@ if __name__ == "__main__":
|
|||||||
twople = (spk_dir, filename, args)
|
twople = (spk_dir, filename, args)
|
||||||
tasks.append(twople)
|
tasks.append(twople)
|
||||||
|
|
||||||
|
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}")
|
||||||
|
|
||||||
|
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)
|
||||||
):
|
):
|
||||||
|
|||||||
24
slice.py
24
slice.py
@@ -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
|
||||||
|
|||||||
33
tools/subprocess_utils.py
Normal file
33
tools/subprocess_utils.py
Normal file
@@ -0,0 +1,33 @@
|
|||||||
|
import subprocess
|
||||||
|
import sys
|
||||||
|
|
||||||
|
from .log import logger
|
||||||
|
|
||||||
|
python = sys.executable
|
||||||
|
|
||||||
|
|
||||||
|
def run_script_with_log(cmd: list[str]) -> tuple[bool, str]:
|
||||||
|
logger.info(f"Running: {' '.join(cmd)}")
|
||||||
|
result = subprocess.run(
|
||||||
|
[python] + cmd,
|
||||||
|
stdout=sys.stdout,
|
||||||
|
stderr=subprocess.PIPE,
|
||||||
|
text=True,
|
||||||
|
)
|
||||||
|
if result.returncode != 0:
|
||||||
|
logger.error(f"Error: {' '.join(cmd)}")
|
||||||
|
print(result.stderr)
|
||||||
|
return False, result.stderr
|
||||||
|
elif result.stderr:
|
||||||
|
logger.warning(f"Warning: {' '.join(cmd)}")
|
||||||
|
print(result.stderr)
|
||||||
|
return True, result.stderr
|
||||||
|
logger.success(f"Success: {' '.join(cmd)}")
|
||||||
|
return True, ""
|
||||||
|
|
||||||
|
|
||||||
|
def second_elem_of(original_function):
|
||||||
|
def inner_function(*args, **kwargs):
|
||||||
|
return original_function(*args, **kwargs)[1]
|
||||||
|
|
||||||
|
return inner_function
|
||||||
@@ -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}")
|
||||||
|
|
||||||
|
|||||||
@@ -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 = [
|
||||||
[
|
|
||||||
python,
|
|
||||||
"slice.py",
|
"slice.py",
|
||||||
"--input_dir",
|
"--input_dir",
|
||||||
input_dir,
|
input_dir,
|
||||||
"--output_dir",
|
"--output_dir",
|
||||||
output_dir,
|
output_dir,
|
||||||
]
|
]
|
||||||
)
|
if normalize:
|
||||||
return "ターミナルを見て結果を確認してください。"
|
cmd.append("--normalize")
|
||||||
|
success, message = run_script_with_log(cmd)
|
||||||
|
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.Accordion("音声のスライス"):
|
||||||
with gr.Row():
|
with gr.Row():
|
||||||
slice_button = gr.Button("音声のスライス")
|
with gr.Column():
|
||||||
|
normalize = gr.Checkbox(label="スライスされた音声の音量を正規化する", value=True)
|
||||||
|
slice_button = gr.Button("スライスを実行")
|
||||||
result1 = gr.Textbox(label="結果")
|
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(
|
||||||
|
|||||||
366
webui_train.py
366
webui_train.py
@@ -1,22 +1,14 @@
|
|||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
import shutil
|
import shutil
|
||||||
import subprocess
|
|
||||||
import sys
|
import sys
|
||||||
|
from multiprocessing import cpu_count
|
||||||
|
|
||||||
import gradio as gr
|
import gradio as gr
|
||||||
import yaml
|
import yaml
|
||||||
|
|
||||||
python = sys.executable
|
from tools.log import logger
|
||||||
|
from tools.subprocess_utils import run_script_with_log, second_elem_of
|
||||||
|
|
||||||
def subprocess_wrapper(cmd):
|
|
||||||
return subprocess.run(
|
|
||||||
cmd,
|
|
||||||
stdout=sys.stdout,
|
|
||||||
stderr=subprocess.PIPE,
|
|
||||||
text=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def get_path(model_name):
|
def get_path(model_name):
|
||||||
@@ -29,7 +21,8 @@ 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...")
|
||||||
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):
|
||||||
config = json.load(open(config_path, "r", encoding="utf-8"))
|
config = json.load(open(config_path, "r", encoding="utf-8"))
|
||||||
@@ -42,14 +35,17 @@ 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:
|
||||||
shutil.copytree(src="pretrained", dst=model_path)
|
shutil.copytree(src="pretrained", dst=model_path)
|
||||||
except FileExistsError:
|
except FileExistsError:
|
||||||
return f"Error: モデルフォルダ {model_path} が既に存在します。問題なければ削除してください。"
|
logger.error(f"Step 1: {model_path} already exists.")
|
||||||
|
return False, f"Step1, Error: モデルフォルダ {model_path} が既に存在します。問題なければ削除してください。"
|
||||||
except FileNotFoundError:
|
except FileNotFoundError:
|
||||||
return "Error: pretrainedフォルダが見つかりません。"
|
logger.error("Step 1: `pretrained` folder not found.")
|
||||||
|
return False, "Step 1, Error: pretrainedフォルダが見つかりません。"
|
||||||
|
|
||||||
with open(config_path, "w", encoding="utf-8") as f:
|
with open(config_path, "w", encoding="utf-8") as f:
|
||||||
json.dump(config, f, indent=2)
|
json.dump(config, f, indent=2)
|
||||||
@@ -62,33 +58,47 @@ def initialize(model_name, batch_size, epochs, bf16_run):
|
|||||||
yml_data["dataset_path"] = dataset_path
|
yml_data["dataset_path"] = dataset_path
|
||||||
with open("config.yml", "w", encoding="utf-8") as f:
|
with open("config.yml", "w", encoding="utf-8") as f:
|
||||||
yaml.dump(yml_data, f, allow_unicode=True)
|
yaml.dump(yml_data, f, allow_unicode=True)
|
||||||
return "Step 1: 初期設定が完了しました"
|
logger.success("Step 1: initialization finished.")
|
||||||
|
return True, "Step 1, Success: 初期設定が完了しました"
|
||||||
|
|
||||||
|
|
||||||
def resample(model_name):
|
def resample(model_name, normalize, num_processes):
|
||||||
|
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")
|
||||||
result = subprocess_wrapper(
|
cmd = [
|
||||||
[
|
|
||||||
python,
|
|
||||||
"resample.py",
|
"resample.py",
|
||||||
"--in_dir",
|
"--in_dir",
|
||||||
in_dir,
|
in_dir,
|
||||||
"--out_dir",
|
"--out_dir",
|
||||||
out_dir,
|
out_dir,
|
||||||
|
"--num_processes",
|
||||||
|
str(num_processes),
|
||||||
"--sr",
|
"--sr",
|
||||||
"44100",
|
"44100",
|
||||||
]
|
]
|
||||||
)
|
if normalize:
|
||||||
if result.stderr:
|
cmd.append("--normalize")
|
||||||
return f"{result.stderr}"
|
success, message = run_script_with_log(cmd)
|
||||||
return "Step 2: 音声ファイルの前処理が完了しました"
|
if not success:
|
||||||
|
logger.error(f"Step 2: resampling failed.")
|
||||||
|
return False, f"Step 2, Error: 音声ファイルの前処理に失敗しました:\n{message}"
|
||||||
|
elif message:
|
||||||
|
logger.warning(f"Step 2: resampling finished with stderr.")
|
||||||
|
return True, f"Step 2, Success: 音声ファイルの前処理が完了しました:\n{message}"
|
||||||
|
logger.success("Step 2: resampling finished.")
|
||||||
|
return True, "Step 2, Success: 音声ファイルの前処理が完了しました"
|
||||||
|
|
||||||
|
|
||||||
def preprocess_text(model_name):
|
def preprocess_text(model_name):
|
||||||
|
logger.info("Step 3: start preprocessing text...")
|
||||||
dataset_path, lbl_path, train_path, val_path, config_path = get_path(model_name)
|
dataset_path, lbl_path, train_path, val_path, config_path = get_path(model_name)
|
||||||
|
try:
|
||||||
lines = open(lbl_path, "r", encoding="utf-8").readlines()
|
lines = open(lbl_path, "r", encoding="utf-8").readlines()
|
||||||
|
except FileNotFoundError:
|
||||||
|
logger.error(f"Step 3: {lbl_path} not found.")
|
||||||
|
return False, f"Step 3, Error: 書き起こしファイル {lbl_path} が見つかりません。"
|
||||||
with open(lbl_path, "w", encoding="utf-8") as f:
|
with open(lbl_path, "w", encoding="utf-8") as f:
|
||||||
for line in lines:
|
for line in lines:
|
||||||
path, spk, language, text = line.strip().split("|")
|
path, spk, language, text = line.strip().split("|")
|
||||||
@@ -96,9 +106,8 @@ def preprocess_text(model_name):
|
|||||||
"\\", "/"
|
"\\", "/"
|
||||||
)
|
)
|
||||||
f.writelines(f"{path}|{spk}|{language}|{text}\n")
|
f.writelines(f"{path}|{spk}|{language}|{text}\n")
|
||||||
result = subprocess_wrapper(
|
success, message = run_script_with_log(
|
||||||
[
|
[
|
||||||
python,
|
|
||||||
"preprocess_text.py",
|
"preprocess_text.py",
|
||||||
"--config-path",
|
"--config-path",
|
||||||
config_path,
|
config_path,
|
||||||
@@ -110,38 +119,99 @@ def preprocess_text(model_name):
|
|||||||
val_path,
|
val_path,
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
if not success:
|
||||||
if result.stderr:
|
logger.error(f"Step 3: preprocessing text failed.")
|
||||||
return f"{result.stderr}"
|
return False, f"Step 3, Error: 書き起こしファイルの前処理に失敗しました:\n{message}"
|
||||||
return "Step 3: 書き起こしファイルの前処理が完了しました"
|
elif message:
|
||||||
|
logger.warning(f"Step 3: preprocessing text finished with stderr.")
|
||||||
|
return True, f"Step 3, Success: 書き起こしファイルの前処理が完了しました:\n{message}"
|
||||||
|
logger.success("Step 3: preprocessing text finished.")
|
||||||
|
return True, "Step 3, Success: 書き起こしファイルの前処理が完了しました"
|
||||||
|
|
||||||
|
|
||||||
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)
|
||||||
result = subprocess_wrapper([python, "bert_gen.py", "--config", config_path])
|
success, message = run_script_with_log(
|
||||||
if result.stderr:
|
[
|
||||||
return f"{result.stderr}"
|
"bert_gen.py",
|
||||||
return "Step 4: BERT特徴ファイルの生成が完了しました"
|
"--config",
|
||||||
|
config_path,
|
||||||
|
# "--num_processes", # bert_genは重いのでプロセス数いじらない
|
||||||
def style_gen(model_name):
|
# str(num_processes),
|
||||||
dataset_path, _, _, _, config_path = get_path(model_name)
|
]
|
||||||
result = subprocess_wrapper(
|
|
||||||
[python, "style_gen.py", "--config", config_path, "--model", dataset_path]
|
|
||||||
)
|
)
|
||||||
if result.stderr:
|
if not success:
|
||||||
return f"{result.stderr}"
|
logger.error(f"Step 4: bert_gen failed.")
|
||||||
return "Step 5: スタイル特徴ファイルの生成が完了しました"
|
return False, f"Step 4, Error: BERT特徴ファイルの生成に失敗しました:\n{message}"
|
||||||
|
elif message:
|
||||||
|
logger.warning(f"Step 4: bert_gen finished with stderr.")
|
||||||
|
return True, f"Step 4, Success: BERT特徴ファイルの生成が完了しました:\n{message}"
|
||||||
|
logger.success("Step 4: bert_gen finished.")
|
||||||
|
return True, "Step 4, Success: BERT特徴ファイルの生成が完了しました"
|
||||||
|
|
||||||
|
|
||||||
|
def style_gen(model_name, num_processes):
|
||||||
|
logger.info("Step 5: start style_gen...")
|
||||||
|
_, _, _, _, config_path = get_path(model_name)
|
||||||
|
success, message = run_script_with_log(
|
||||||
|
[
|
||||||
|
"style_gen.py",
|
||||||
|
"--config",
|
||||||
|
config_path,
|
||||||
|
"--num_processes",
|
||||||
|
str(num_processes),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
if not success:
|
||||||
|
logger.error(f"Step 5: style_gen failed.")
|
||||||
|
return False, f"Step 5, Error: スタイル特徴ファイルの生成に失敗しました:\n{message}"
|
||||||
|
elif message:
|
||||||
|
logger.warning(f"Step 5: style_gen finished with stderr.")
|
||||||
|
return True, f"Step 5, Success: スタイル特徴ファイルの生成が完了しました:\n{message}"
|
||||||
|
logger.success("Step 5: style_gen finished.")
|
||||||
|
return True, "Step 5, Success: スタイル特徴ファイルの生成が完了しました"
|
||||||
|
|
||||||
|
|
||||||
|
def preprocess_all(
|
||||||
|
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:
|
||||||
|
return False, message
|
||||||
|
success, message = resample(model_name, normalize, num_processes)
|
||||||
|
if not success:
|
||||||
|
return False, message
|
||||||
|
success, message = preprocess_text(model_name)
|
||||||
|
if not success:
|
||||||
|
return False, message
|
||||||
|
success, message = bert_gen(model_name) # bert_genは重いのでプロセス数いじらない
|
||||||
|
if not success:
|
||||||
|
return False, message
|
||||||
|
success, message = style_gen(model_name, num_processes)
|
||||||
|
if not success:
|
||||||
|
return False, message
|
||||||
|
logger.success("Success: All preprocess finished!")
|
||||||
|
return True, "Success: 全ての前処理が完了しました。ターミナルを確認しておかしいところがないか確認するのをおすすめします。"
|
||||||
|
|
||||||
|
|
||||||
def train(model_name):
|
def train(model_name):
|
||||||
dataset_path, _, _, _, config_path = get_path(model_name)
|
dataset_path, _, _, _, config_path = get_path(model_name)
|
||||||
result = subprocess_wrapper(
|
success, message = run_script_with_log(
|
||||||
[python, "train_ms.py", "--config", config_path, "--model", dataset_path]
|
["train_ms.py", "--config", config_path, "--model", dataset_path]
|
||||||
)
|
)
|
||||||
if result.stderr:
|
if not success:
|
||||||
return f"{result.stderr}"
|
logger.error(f"Final Step: train failed.")
|
||||||
return "Final Step: 学習が完了しました!"
|
return False, f"Final Step, Error: 学習に失敗しました:\n{message}"
|
||||||
|
elif message:
|
||||||
|
logger.warning(f"Final Step: train finished with stderr.")
|
||||||
|
return True, f"Final Step, Success: 学習が完了しました:\n{message}"
|
||||||
|
logger.success("Final Step: train finished.")
|
||||||
|
return True, "Final Step, Success: 学習が完了しました"
|
||||||
|
|
||||||
|
|
||||||
initial_md = """
|
initial_md = """
|
||||||
@@ -149,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秒程度の、長すぎず短すぎない発話のものをいくつか)と、書き起こしテキストを用意してください。
|
||||||
|
|
||||||
それを次のように配置します。
|
それを次のように配置します。
|
||||||
```
|
```
|
||||||
@@ -195,10 +269,9 @@ if __name__ == "__main__":
|
|||||||
model_name = gr.Textbox(
|
model_name = gr.Textbox(
|
||||||
label="モデル名",
|
label="モデル名",
|
||||||
)
|
)
|
||||||
info = gr.Textbox(label="状況")
|
gr.Markdown("### 自動前処理")
|
||||||
with gr.Row():
|
with gr.Row(variant="panel"):
|
||||||
with gr.Column():
|
with gr.Column():
|
||||||
gr.Markdown(value="### Step 1: 設定ファイルの生成")
|
|
||||||
batch_size = gr.Slider(
|
batch_size = gr.Slider(
|
||||||
label="バッチサイズ",
|
label="バッチサイズ",
|
||||||
info="VRAM 12GBで4くらい",
|
info="VRAM 12GBで4くらい",
|
||||||
@@ -215,38 +288,167 @@ if __name__ == "__main__":
|
|||||||
maximum=1000,
|
maximum=1000,
|
||||||
step=1,
|
step=1,
|
||||||
)
|
)
|
||||||
|
save_every_steps = gr.Slider(
|
||||||
|
label="何ステップごとに結果を保存するか",
|
||||||
|
info="エポック数とは違うことに注意",
|
||||||
|
value=1000,
|
||||||
|
minimum=100,
|
||||||
|
maximum=10000,
|
||||||
|
step=100,
|
||||||
|
)
|
||||||
bf16_run = gr.Checkbox(
|
bf16_run = gr.Checkbox(
|
||||||
label="bfloat16を使う",
|
label="bfloat16を使う",
|
||||||
info="新しめのグラボだと学習が早くなるかも、古いグラボだと動かないかも",
|
info="bfloat16を使うかどうか。新しめのグラボだと学習が早くなるかも、古いグラボだと動かないかも。",
|
||||||
value=True,
|
value=True,
|
||||||
)
|
)
|
||||||
generate_config_btn = gr.Button(value="実行", variant="primary")
|
num_processes = gr.Slider(
|
||||||
with gr.Column():
|
label="プロセス数",
|
||||||
gr.Markdown(value="### Step 2: 音声ファイルの前処理")
|
info="前処理時の並列処理プロセス数、大きすぎるとフリーズするかも",
|
||||||
resample_btn = gr.Button(value="実行", variant="primary")
|
value=cpu_count() // 2,
|
||||||
with gr.Column():
|
minimum=1,
|
||||||
gr.Markdown(value="### Step 3: 書き起こしファイルの前処理")
|
maximum=cpu_count(),
|
||||||
preprocess_text_btn = gr.Button(value="実行", variant="primary")
|
step=1,
|
||||||
with gr.Column():
|
|
||||||
gr.Markdown(value="### Step 4: BERT特徴ファイルの生成")
|
|
||||||
bert_gen_btn = gr.Button(value="実行", variant="primary")
|
|
||||||
with gr.Column():
|
|
||||||
gr.Markdown(value="### Step 5: スタイル特徴ファイルの生成")
|
|
||||||
style_gen_btn = gr.Button(value="実行", variant="primary")
|
|
||||||
with gr.Row():
|
|
||||||
with gr.Column():
|
|
||||||
gr.Markdown(value="### Final Step: 学習")
|
|
||||||
train_btn = gr.Button(value="学習", variant="primary")
|
|
||||||
|
|
||||||
generate_config_btn.click(
|
|
||||||
initialize,
|
|
||||||
inputs=[model_name, batch_size, epochs, bf16_run],
|
|
||||||
outputs=[info],
|
|
||||||
)
|
)
|
||||||
resample_btn.click(resample, inputs=[model_name], outputs=[info])
|
normalize = gr.Checkbox(
|
||||||
preprocess_text_btn.click(preprocess_text, inputs=[model_name], outputs=[info])
|
label="音声の音量を正規化する",
|
||||||
bert_gen_btn.click(bert_gen, inputs=[model_name], outputs=[info])
|
value=True,
|
||||||
style_gen_btn.click(style_gen, inputs=[model_name], outputs=[info])
|
)
|
||||||
train_btn.click(train, inputs=[model_name], outputs=[info])
|
with gr.Column():
|
||||||
|
preprocess_button = gr.Button(value="自動前処理を実行", variant="primary")
|
||||||
|
info_all = gr.Textbox(label="状況")
|
||||||
|
with gr.Accordion(open=False, label="手動前処理"):
|
||||||
|
with gr.Row(variant="panel"):
|
||||||
|
with gr.Column():
|
||||||
|
gr.Markdown(value="#### Step 1: 設定ファイルの生成")
|
||||||
|
batch_size_manual = gr.Slider(
|
||||||
|
label="バッチサイズ",
|
||||||
|
value=4,
|
||||||
|
minimum=1,
|
||||||
|
maximum=64,
|
||||||
|
step=1,
|
||||||
|
)
|
||||||
|
epochs_manual = gr.Slider(
|
||||||
|
label="エポック数",
|
||||||
|
value=100,
|
||||||
|
minimum=1,
|
||||||
|
maximum=1000,
|
||||||
|
step=1,
|
||||||
|
)
|
||||||
|
save_every_steps_manual = gr.Slider(
|
||||||
|
label="何ステップごとに結果を保存するか",
|
||||||
|
value=1000,
|
||||||
|
minimum=100,
|
||||||
|
maximum=10000,
|
||||||
|
step=100,
|
||||||
|
)
|
||||||
|
bf16_run_manual = gr.Checkbox(
|
||||||
|
label="bfloat16を使う",
|
||||||
|
value=True,
|
||||||
|
)
|
||||||
|
with gr.Column():
|
||||||
|
generate_config_btn = gr.Button(value="実行", variant="primary")
|
||||||
|
info_init = gr.Textbox(label="状況")
|
||||||
|
with gr.Row(variant="panel"):
|
||||||
|
with gr.Column():
|
||||||
|
gr.Markdown(value="#### Step 2: 音声ファイルの前処理")
|
||||||
|
num_processes_resample = gr.Slider(
|
||||||
|
label="プロセス数",
|
||||||
|
value=cpu_count() // 2,
|
||||||
|
minimum=1,
|
||||||
|
maximum=cpu_count(),
|
||||||
|
step=1,
|
||||||
|
)
|
||||||
|
normalize_resample = gr.Checkbox(
|
||||||
|
label="音声の音量を正規化する",
|
||||||
|
value=True,
|
||||||
|
)
|
||||||
|
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")
|
||||||
|
info_bert = gr.Textbox(label="状況")
|
||||||
|
with gr.Row(variant="panel"):
|
||||||
|
with gr.Column():
|
||||||
|
gr.Markdown(value="#### Step 5: スタイル特徴ファイルの生成")
|
||||||
|
num_processes_style = gr.Slider(
|
||||||
|
label="プロセス数",
|
||||||
|
value=cpu_count() // 2,
|
||||||
|
minimum=1,
|
||||||
|
maximum=cpu_count(),
|
||||||
|
step=1,
|
||||||
|
)
|
||||||
|
with gr.Column():
|
||||||
|
style_gen_btn = gr.Button(value="実行", variant="primary")
|
||||||
|
info_style = gr.Textbox(label="状況")
|
||||||
|
gr.Markdown("## 学習")
|
||||||
|
with gr.Row(variant="panel"):
|
||||||
|
train_btn = gr.Button(value="学習を開始する", variant="primary")
|
||||||
|
info_train = gr.Textbox(label="状況")
|
||||||
|
|
||||||
app.launch(share=False, inbrowser=True)
|
preprocess_button.click(
|
||||||
|
second_elem_of(preprocess_all),
|
||||||
|
inputs=[
|
||||||
|
model_name,
|
||||||
|
batch_size,
|
||||||
|
epochs,
|
||||||
|
save_every_steps,
|
||||||
|
bf16_run,
|
||||||
|
num_processes,
|
||||||
|
normalize,
|
||||||
|
],
|
||||||
|
outputs=[info_all],
|
||||||
|
)
|
||||||
|
generate_config_btn.click(
|
||||||
|
second_elem_of(initialize),
|
||||||
|
inputs=[
|
||||||
|
model_name,
|
||||||
|
batch_size_manual,
|
||||||
|
epochs_manual,
|
||||||
|
save_every_steps_manual,
|
||||||
|
bf16_run_manual,
|
||||||
|
],
|
||||||
|
outputs=[info_init],
|
||||||
|
)
|
||||||
|
resample_btn.click(
|
||||||
|
second_elem_of(resample),
|
||||||
|
inputs=[model_name, normalize_resample, num_processes_resample],
|
||||||
|
outputs=[info_resample],
|
||||||
|
)
|
||||||
|
preprocess_text_btn.click(
|
||||||
|
second_elem_of(preprocess_text),
|
||||||
|
inputs=[model_name],
|
||||||
|
outputs=[info_preprocess_text],
|
||||||
|
)
|
||||||
|
bert_gen_btn.click(
|
||||||
|
second_elem_of(bert_gen),
|
||||||
|
inputs=[model_name],
|
||||||
|
outputs=[info_bert],
|
||||||
|
)
|
||||||
|
style_gen_btn.click(
|
||||||
|
second_elem_of(style_gen),
|
||||||
|
inputs=[model_name, num_processes_style],
|
||||||
|
outputs=[info_style],
|
||||||
|
)
|
||||||
|
train_btn.click(
|
||||||
|
second_elem_of(train), inputs=[model_name], outputs=[info_train]
|
||||||
|
)
|
||||||
|
|
||||||
|
app.launch(inbrowser=True)
|
||||||
|
|||||||
Reference in New Issue
Block a user