From 87144dd321ea6cc3c65d76772ca0d01e0699c0cf Mon Sep 17 00:00:00 2001 From: tsukumi Date: Mon, 6 May 2024 20:56:07 +0900 Subject: [PATCH] Improve: automatically generate configs/paths.yml by copying it from configs/default_paths.yml when running initialize.py If configs/paths.yml itself is included in version control, differences will occur when it is changed in each environment, which is troublesome. --- .gitignore | 2 ++ configs/{paths.yml => default_paths.yml} | 0 gradio_tabs/style_vectors.py | 1 + initialize.py | 8 +++++++- style_bert_vits2/tts_model.py | 5 +++-- 5 files changed, 13 insertions(+), 3 deletions(-) rename configs/{paths.yml => default_paths.yml} (100%) diff --git a/.gitignore b/.gitignore index aa114e3..cc5fed4 100644 --- a/.gitignore +++ b/.gitignore @@ -14,6 +14,8 @@ dist/ /bert/*/*.safetensors /bert/*/*.msgpack +/configs/paths.yml + /pretrained/*.safetensors /pretrained/*.pth diff --git a/configs/paths.yml b/configs/default_paths.yml similarity index 100% rename from configs/paths.yml rename to configs/default_paths.yml diff --git a/gradio_tabs/style_vectors.py b/gradio_tabs/style_vectors.py index 9234cff..af43258 100644 --- a/gradio_tabs/style_vectors.py +++ b/gradio_tabs/style_vectors.py @@ -16,6 +16,7 @@ from config import config from style_bert_vits2.constants import DEFAULT_STYLE, GRADIO_THEME from style_bert_vits2.logging import logger + # Get path settings with open(os.path.join("configs", "paths.yml"), "r", encoding="utf-8") as f: path_config: dict[str, str] = yaml.safe_load(f.read()) diff --git a/initialize.py b/initialize.py index bfe59e4..664d570 100644 --- a/initialize.py +++ b/initialize.py @@ -1,5 +1,6 @@ import argparse import json +import shutil from pathlib import Path import yaml @@ -102,11 +103,16 @@ def main(): download_pretrained_models() download_jp_extra_pretrained_models() + # If configs/paths.yml not exists, create it + default_paths_yml = Path("configs/default_paths.yml") + paths_yml = Path("configs/paths.yml") + if not paths_yml.exists(): + shutil.copy(default_paths_yml, paths_yml) + 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: diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index c60a632..092957c 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -69,8 +69,9 @@ class TTSModel: # ハイパーパラメータのパスが指定された else: self.config_path: Path = config_path - self.hyper_parameters: HyperParameters = \ - HyperParameters.load_from_json(self.config_path) + self.hyper_parameters: HyperParameters = HyperParameters.load_from_json( + self.config_path + ) # スタイルベクトルの NDArray が直接指定された if isinstance(style_vec_path, np.ndarray):