From 036cea17f3181514c1dc5c5f269869a82ae19da5 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Mon, 19 Feb 2024 17:22:40 +0900 Subject: [PATCH] Add preprocess_all.py and optimize cli settings --- gen_yaml.py | 16 ++++---- initialize.py | 33 +++++++++++++++- preprocess_all.py | 96 +++++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 135 insertions(+), 10 deletions(-) create mode 100644 preprocess_all.py diff --git a/gen_yaml.py b/gen_yaml.py index 8650dc0..91301ac 100644 --- a/gen_yaml.py +++ b/gen_yaml.py @@ -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) - diff --git a/initialize.py b/initialize.py index c163ef9..5e35061 100644 --- a/initialize.py +++ b/initialize.py @@ -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() diff --git a/preprocess_all.py b/preprocess_all.py new file mode 100644 index 0000000..b83e93b --- /dev/null +++ b/preprocess_all.py @@ -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, + )