Update webui

This commit is contained in:
litagin02
2024-01-16 18:38:36 +09:00
parent 1289681514
commit 714be0750b
2 changed files with 16 additions and 4 deletions

View File

@@ -198,6 +198,7 @@ def save_style_vectors_from_clustering(model_name, style_names_str: str):
style_vectors = np.stack([mean] + centroids) style_vectors = np.stack([mean] + centroids)
style_vector_path = os.path.join(result_dir, "style_vectors.npy") style_vector_path = os.path.join(result_dir, "style_vectors.npy")
if os.path.exists(style_vector_path): if os.path.exists(style_vector_path):
logger.info(f"Backup {style_vector_path} to {style_vector_path}.bak")
shutil.copy(style_vector_path, f"{style_vector_path}.bak") shutil.copy(style_vector_path, f"{style_vector_path}.bak")
np.save(style_vector_path, style_vectors) np.save(style_vector_path, style_vectors)
@@ -214,6 +215,7 @@ def save_style_vectors_from_clustering(model_name, style_names_str: str):
if len(set(style_names)) != len(style_names): if len(set(style_names)) != len(style_names):
return f"スタイル名が重複しています。" return f"スタイル名が重複しています。"
logger.info(f"Backup {config_path} to {config_path}.bak")
shutil.copy(config_path, f"{config_path}.bak") shutil.copy(config_path, f"{config_path}.bak")
with open(config_path, "r", encoding="utf-8") as f: with open(config_path, "r", encoding="utf-8") as f:
json_dict = json.load(f) json_dict = json.load(f)
@@ -255,6 +257,7 @@ def save_style_vectors_from_files(
assert len(style_name_list) == len(style_vectors) assert len(style_name_list) == len(style_vectors)
style_vector_path = os.path.join(result_dir, "style_vectors.npy") style_vector_path = os.path.join(result_dir, "style_vectors.npy")
if os.path.exists(style_vector_path): if os.path.exists(style_vector_path):
logger.info(f"Backup {style_vector_path} to {style_vector_path}.bak")
shutil.copy(style_vector_path, f"{style_vector_path}.bak") shutil.copy(style_vector_path, f"{style_vector_path}.bak")
np.save(style_vector_path, style_vectors) np.save(style_vector_path, style_vectors)
@@ -262,6 +265,7 @@ def save_style_vectors_from_files(
config_path = os.path.join(result_dir, "config.json") config_path = os.path.join(result_dir, "config.json")
if not os.path.exists(config_path): if not os.path.exists(config_path):
return f"{config_path}が存在しません。" return f"{config_path}が存在しません。"
logger.info(f"Backup {config_path} to {config_path}.bak")
shutil.copy(config_path, f"{config_path}.bak") shutil.copy(config_path, f"{config_path}.bak")
with open(config_path, "r", encoding="utf-8") as f: with open(config_path, "r", encoding="utf-8") as f:
@@ -291,7 +295,7 @@ Style-Bert-VITS2でこまかくスタイルを指定して音声合成するに
- 方法3: 自分でもっと頑張ってこだわって作るJVNVコーパスなど、もともとスタイルラベル等が利用可能な場合はこれがよいかも - 方法3: 自分でもっと頑張ってこだわって作るJVNVコーパスなど、もともとスタイルラベル等が利用可能な場合はこれがよいかも
""" """
method1 = """ method1 = f"""
学習の時に取り出したスタイルベクトルを読み込んで、可視化を見ながらスタイルを分けていきます。 学習の時に取り出したスタイルベクトルを読み込んで、可視化を見ながらスタイルを分けていきます。
手順: 手順:
@@ -303,7 +307,7 @@ method1 = """
詳細: スタイルベクトル(256次元)たちを適当なアルゴリズムでクラスタリングして、各クラスタの中心のベクトル(と全体の平均ベクトル)を保存します。 詳細: スタイルベクトル(256次元)たちを適当なアルゴリズムでクラスタリングして、各クラスタの中心のベクトル(と全体の平均ベクトル)を保存します。
平均スタイル({DEFAULT_EMOTION})は自動的に保存されます。 平均スタイル({DEFAULT_STYLE})は自動的に保存されます。
""" """
dbscan_md = """ dbscan_md = """

View File

@@ -238,7 +238,7 @@ def preprocess_all(
return True, "Success: 全ての前処理が完了しました。ターミナルを確認しておかしいところがないか確認するのをおすすめします。" return True, "Success: 全ての前処理が完了しました。ターミナルを確認しておかしいところがないか確認するのをおすすめします。"
def train(model_name): def train(model_name, skip_style=False):
dataset_path, _, _, _, config_path = get_path(model_name) dataset_path, _, _, _, config_path = get_path(model_name)
# 学習再開の場合は念のためconfig.ymlの名前等を更新 # 学習再開の場合は念のためconfig.ymlの名前等を更新
with open("config.yml", "r", encoding="utf-8") as f: with open("config.yml", "r", encoding="utf-8") as f:
@@ -247,6 +247,9 @@ def train(model_name):
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)
cmd = ["train_ms.py", "--config", config_path, "--model", dataset_path]
if skip_style:
cmd.append("--skip_default_style")
success, message = run_script_with_log( success, message = run_script_with_log(
["train_ms.py", "--config", config_path, "--model", dataset_path] ["train_ms.py", "--config", config_path, "--model", dataset_path]
) )
@@ -445,6 +448,11 @@ if __name__ == "__main__":
info_style = gr.Textbox(label="状況") info_style = gr.Textbox(label="状況")
gr.Markdown("## 学習") gr.Markdown("## 学習")
with gr.Row(variant="panel"): with gr.Row(variant="panel"):
skip_style = gr.Checkbox(
label="スタイルファイルの生成をスキップする",
info="学習再開の場合の場合はチェックしてください",
value=False,
)
train_btn = gr.Button(value="学習を開始する", variant="primary") train_btn = gr.Button(value="学習を開始する", variant="primary")
info_train = gr.Textbox(label="状況") info_train = gr.Textbox(label="状況")
@@ -499,7 +507,7 @@ if __name__ == "__main__":
outputs=[info_style], outputs=[info_style],
) )
train_btn.click( train_btn.click(
second_elem_of(train), inputs=[model_name], outputs=[info_train] second_elem_of(train), inputs=[model_name, skip_style], outputs=[info_train]
) )
parser = argparse.ArgumentParser() parser = argparse.ArgumentParser()