Add preprocess_all.py and optimize cli settings
This commit is contained in:
16
gen_yaml.py
16
gen_yaml.py
@@ -3,22 +3,20 @@ import shutil
|
|||||||
import yaml
|
import yaml
|
||||||
import argparse
|
import argparse
|
||||||
|
|
||||||
parser = argparse.ArgumentParser(description="config.ymlの生成。あらかじめ前準備をしたデータをバッチファイルなどで連続で学習する時にtrain_ms.pyより前に使用する。")
|
parser = argparse.ArgumentParser(
|
||||||
# そうしないと最後の前準備したデータで学習してしまう
|
description="config.ymlの生成。あらかじめ前準備をしたデータをバッチファイルなどで連続で学習する時にtrain_ms.pyより前に使用する。"
|
||||||
parser.add_argument(
|
|
||||||
"--model_name",
|
|
||||||
type=str,
|
|
||||||
help="Model name",
|
|
||||||
required=True
|
|
||||||
)
|
)
|
||||||
|
# そうしないと最後の前準備したデータで学習してしまう
|
||||||
|
parser.add_argument("--model_name", type=str, help="Model name", required=True)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--dataset_path",
|
"--dataset_path",
|
||||||
type=str,
|
type=str,
|
||||||
help="Dataset path(example: Data\\your_model_name)",
|
help="Dataset path(example: Data\\your_model_name)",
|
||||||
required=True
|
required=True,
|
||||||
)
|
)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
def gen_yaml(model_name, dataset_path):
|
def gen_yaml(model_name, dataset_path):
|
||||||
if not os.path.exists("config.yml"):
|
if not os.path.exists("config.yml"):
|
||||||
shutil.copy(src="default_config.yml", dst="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:
|
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)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
gen_yaml(args.model_name, args.dataset_path)
|
gen_yaml(args.model_name, args.dataset_path)
|
||||||
|
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ import argparse
|
|||||||
import json
|
import json
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
import yaml
|
||||||
from huggingface_hub import hf_hub_download
|
from huggingface_hub import hf_hub_download
|
||||||
|
|
||||||
from common.log import logger
|
from common.log import logger
|
||||||
@@ -90,9 +91,21 @@ def download_jvnv_models():
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
def main():
|
||||||
parser = argparse.ArgumentParser()
|
parser = argparse.ArgumentParser()
|
||||||
parser.add_argument("--skip_jvnv", action="store_true")
|
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()
|
args = parser.parse_args()
|
||||||
|
|
||||||
download_bert_models()
|
download_bert_models()
|
||||||
@@ -105,3 +118,21 @@ if __name__ == "__main__":
|
|||||||
|
|
||||||
if not args.skip_jvnv:
|
if not args.skip_jvnv:
|
||||||
download_jvnv_models()
|
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
96
preprocess_all.py
Normal 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,
|
||||||
|
)
|
||||||
Reference in New Issue
Block a user