Add preprocess_all.py and optimize cli settings

This commit is contained in:
litagin02
2024-02-19 17:22:40 +09:00
parent d26336bf0a
commit 036cea17f3
3 changed files with 135 additions and 10 deletions

View File

@@ -3,22 +3,20 @@ import shutil
import yaml
import argparse
parser = argparse.ArgumentParser(description="config.ymlの生成。あらかじめ前準備をしたデータをバッチファイルなどで連続で学習する時にtrain_ms.pyより前に使用する。")
# そうしないと最後の前準備したデータで学習してしまう
parser.add_argument(
"--model_name",
type=str,
help="Model name",
required=True
parser = argparse.ArgumentParser(
description="config.ymlの生成。あらかじめ前準備したデータをバッチファイルなどで連続で学習する時にtrain_ms.pyより前に使用する。"
)
# そうしないと最後の前準備したデータで学習してしまう
parser.add_argument("--model_name", type=str, help="Model name", required=True)
parser.add_argument(
"--dataset_path",
type=str,
help="Dataset path(example: Data\\your_model_name)",
required=True
required=True,
)
args = parser.parse_args()
def gen_yaml(model_name, dataset_path):
if not os.path.exists("config.yml"):
shutil.copy(src="default_config.yml", dst="config.yml")
@@ -29,6 +27,6 @@ def gen_yaml(model_name, dataset_path):
with open("config.yml", "w", encoding="utf-8") as f:
yaml.dump(yml_data, f, allow_unicode=True)
if __name__ == "__main__":
gen_yaml(args.model_name, args.dataset_path)

View File

@@ -2,6 +2,7 @@ import argparse
import json
from pathlib import Path
import yaml
from huggingface_hub import hf_hub_download
from common.log import logger
@@ -90,9 +91,21 @@ def download_jvnv_models():
)
if __name__ == "__main__":
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--skip_jvnv", action="store_true")
parser.add_argument(
"--dataset_root",
type=str,
help="Dataset root path (default: Data)",
default=None,
)
parser.add_argument(
"--assets_root",
type=str,
help="Assets root path (default: model_assets)",
default=None,
)
args = parser.parse_args()
download_bert_models()
@@ -105,3 +118,21 @@ if __name__ == "__main__":
if not args.skip_jvnv:
download_jvnv_models()
if args.dataset_root is None and args.assets_root is None:
return
# Change default paths if necessary
paths_yml = Path("configs/paths.yml")
with open(paths_yml, "r", encoding="utf-8") as f:
yml_data = yaml.safe_load(f)
if args.assets_root is not None:
yml_data["assets_root"] = args.assets_root
if args.dataset_root is not None:
yml_data["dataset_root"] = args.dataset_root
with open(paths_yml, "w", encoding="utf-8") as f:
yaml.dump(yml_data, f, allow_unicode=True)
if __name__ == "__main__":
main()

96
preprocess_all.py Normal file
View File

@@ -0,0 +1,96 @@
import argparse
from webui_train import preprocess_all
from multiprocessing import cpu_count
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument(
"--model_name", "-m", type=str, help="Model name", required=True
)
parser.add_argument("--batch_size", "-b", type=int, help="Batch size", default=2)
parser.add_argument("--epochs", "-e", type=int, help="Epochs", default=100)
parser.add_argument(
"--save_every_steps",
"-s",
type=int,
help="Save every steps",
default=1000,
)
parser.add_argument(
"--num_processes",
type=int,
help="Number of processes",
default=cpu_count() // 2,
)
parser.add_argument(
"--normalize",
action="store_true",
help="Loudness normalize audio",
)
parser.add_argument(
"--trim",
action="store_true",
help="Trim silence",
)
parser.add_argument(
"--freeze_EN_bert",
action="store_true",
help="Freeze English BERT",
)
parser.add_argument(
"--freeze_JP_bert",
action="store_true",
help="Freeze Japanese BERT",
)
parser.add_argument(
"--freeze_ZH_bert",
action="store_true",
help="Freeze Chinese BERT",
)
parser.add_argument(
"--freeze_style",
action="store_true",
help="Freeze style vector",
)
parser.add_argument(
"--use_jp_extra",
action="store_true",
help="Use JP-Extra pretrained model",
)
parser.add_argument(
"--val_per_lang",
type=int,
help="Validation per language",
default=0,
)
parser.add_argument(
"--log_interval",
type=int,
help="Log interval",
default=200,
)
parser.add_argument(
"--skip_invalid",
action="store_true",
help="Skip invalid",
)
args = parser.parse_args()
preprocess_all(
model_name=args.model_name,
batch_size=args.batch_size,
epochs=args.epochs,
save_every_steps=args.save_every_steps,
num_processes=args.num_processes,
normalize=args.normalize,
trim=args.trim,
freeze_EN_bert=args.freeze_EN_bert,
freeze_JP_bert=args.freeze_JP_bert,
freeze_ZH_bert=args.freeze_ZH_bert,
freeze_style=args.freeze_style,
use_jp_extra=args.use_jp_extra,
val_per_lang=args.val_per_lang,
log_interval=args.log_interval,
skip_invalid=args.skip_invalid,
)