From 245bfe54aebd9de7e304c1591d0daa10e48d7a0d Mon Sep 17 00:00:00 2001 From: tsukumi Date: Mon, 23 Sep 2024 08:13:16 +0900 Subject: [PATCH] =?UTF-8?q?Fix:=20In=20the=20=E2=80=9CMerge=E2=80=9D=20tab?= =?UTF-8?q?=20of=20the=20Web=20UI,=20re-create=20the=20TTSModelHolder=20us?= =?UTF-8?q?ing=20the=20passed=20TTSModelHolder=20instance=20variable=20so?= =?UTF-8?q?=20that=20ONNX=20models=20do=20not=20get=20mixed=20up?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- gradio_tabs/merge.py | 7 ++++++- style_bert_vits2/tts_model.py | 8 +++++++- 2 files changed, 13 insertions(+), 2 deletions(-) diff --git a/gradio_tabs/merge.py b/gradio_tabs/merge.py index c5688e0..6082d1f 100644 --- a/gradio_tabs/merge.py +++ b/gradio_tabs/merge.py @@ -1003,6 +1003,11 @@ def method_change(x: str): def create_merge_app(model_holder: TTSModelHolder) -> gr.Blocks: + # ONNX モデルが混じらないよう、渡された TTSModelHolder のインスタンス変数を使って TTSModelHolder を作り直す + model_holder = TTSModelHolder( + model_holder.root_dir, model_holder.device, model_holder.onnx_providers, ignore_onnx=True + ) + model_names = model_holder.model_names if len(model_names) == 0: logger.error( @@ -1527,7 +1532,7 @@ def create_merge_app(model_holder: TTSModelHolder) -> gr.Blocks: if __name__ == "__main__": device = "cuda" if torch.cuda.is_available() else "cpu" model_holder = TTSModelHolder( - assets_root, device, torch_device_to_onnx_providers(device) + assets_root, device, torch_device_to_onnx_providers(device), ignore_onnx=True ) app = create_merge_app(model_holder) app.launch(inbrowser=True) diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index 8c9570c..47f9cc9 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -467,6 +467,7 @@ class TTSModelHolder: model_root_dir: Path, device: str, onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]], + ignore_onnx: bool = False, ) -> None: """ Style-Bert-VITS2 の音声合成モデルを管理するクラスを初期化する。 @@ -488,11 +489,13 @@ class TTSModelHolder: model_root_dir (Path): 音声合成モデルが配置されているディレクトリのパス device (str): PyTorch 推論での音声合成時に利用するデバイス (cpu, cuda, mps など) onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など) + ignore_onnx (bool, optional): ONNX モデルを除外するかどうか. Defaults to False. """ self.root_dir: Path = model_root_dir self.device: str = device self.onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]] = onnx_providers # fmt: skip + self.ignore_onnx: bool = ignore_onnx self.model_files_dict: dict[str, list[Path]] = {} self.current_model: Optional[TTSModel] = None self.model_names: list[str] = [] @@ -513,11 +516,14 @@ class TTSModelHolder: for model_dir in model_dirs: if model_dir.name.startswith("."): continue + suffixes = [".pth", ".pt", ".safetensors"] + if self.ignore_onnx is False: + suffixes.append(".onnx") model_files = sorted( [ f for f in model_dir.iterdir() - if f.suffix in [".pth", ".pt", ".safetensors", ".onnx"] + if f.suffix in suffixes ] ) if len(model_files) == 0: