From fa7123c422c76e441dad440be848fae83f2ca928 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Wed, 18 Sep 2024 05:11:15 +0900 Subject: [PATCH] Fix: ONNX BERT model/tokenizer is not preloaded by default to avoid wasting VRAM --- server_editor.py | 11 +++++++---- server_fastapi.py | 11 +++++++---- style_bert_vits2/nlp/bert_models.py | 18 ++++++++++-------- style_bert_vits2/nlp/onnx_bert_models.py | 14 ++++++++------ 4 files changed, 32 insertions(+), 22 deletions(-) diff --git a/server_editor.py b/server_editor.py index eee8754..50921c1 100644 --- a/server_editor.py +++ b/server_editor.py @@ -179,6 +179,7 @@ parser.add_argument("--line_length", type=int, default=None) parser.add_argument("--line_count", type=int, default=None) # parser.add_argument("--skip_default_models", action="store_true") parser.add_argument("--skip_static_files", action="store_true") +parser.add_argument("--preload_onnx_bert", action="store_true") args = parser.parse_args() device = args.device if device == "cuda" and not torch.cuda.is_available(): @@ -194,10 +195,12 @@ skip_static_files = bool(args.skip_static_files) ## server_editor.py は日本語にしか対応していないため、日本語の BERT モデル/トークナイザーのみロードする bert_models.load_model(Languages.JP, device_map=device) bert_models.load_tokenizer(Languages.JP) -onnx_bert_models.load_model( - Languages.JP, onnx_providers=torch_device_to_onnx_providers(device) -) -onnx_bert_models.load_tokenizer(Languages.JP) +# VRAM を浪費しないように、既定では ONNX 版 BERT モデル/トークナイザーは事前ロードしない +if args.preload_onnx_bert: + onnx_bert_models.load_model( + Languages.JP, onnx_providers=torch_device_to_onnx_providers(device) + ) + onnx_bert_models.load_tokenizer(Languages.JP) model_holder = TTSModelHolder( model_dir, device, torch_device_to_onnx_providers(device) diff --git a/server_fastapi.py b/server_fastapi.py index eca242e..a331578 100644 --- a/server_fastapi.py +++ b/server_fastapi.py @@ -89,6 +89,7 @@ if __name__ == "__main__": parser.add_argument( "--dir", "-d", type=str, help="Model directory", default=config.assets_root ) + parser.add_argument("--preload_onnx_bert", action="store_true") args = parser.parse_args() if args.cpu: @@ -104,10 +105,12 @@ if __name__ == "__main__": bert_models.load_tokenizer(Languages.EN) bert_models.load_model(Languages.ZH, device_map=device) bert_models.load_tokenizer(Languages.ZH) - onnx_bert_models.load_model( - Languages.JP, onnx_providers=torch_device_to_onnx_providers(device) - ) - onnx_bert_models.load_tokenizer(Languages.JP) + # VRAM を浪費しないように、既定では ONNX 版 BERT モデル/トークナイザーは事前ロードしない + if args.preload_onnx_bert: + onnx_bert_models.load_model( + Languages.JP, onnx_providers=torch_device_to_onnx_providers(device) + ) + onnx_bert_models.load_tokenizer(Languages.JP) model_dir = Path(args.dir) model_holder = TTSModelHolder( diff --git a/style_bert_vits2/nlp/bert_models.py b/style_bert_vits2/nlp/bert_models.py index 0c4e194..da74a4f 100644 --- a/style_bert_vits2/nlp/bert_models.py +++ b/style_bert_vits2/nlp/bert_models.py @@ -9,6 +9,7 @@ Style-Bert-VITS2 の学習・推論に必要な各言語ごとの BERT モデル """ import gc +import time from typing import Optional, Union, cast import torch @@ -80,11 +81,12 @@ def load_model( if pretrained_model_name_or_path is None: assert DEFAULT_BERT_MODEL_PATHS[ language - ].exists(), f"The default {language} BERT model does not exist on the file system. Please specify the path to the pre-trained model." + ].exists(), f"The default {language.name} BERT model does not exist on the file system. Please specify the path to the pre-trained model." pretrained_model_name_or_path = str(DEFAULT_BERT_MODEL_PATHS[language]) # BERT モデルをロードし、辞書に格納して返す ## 英語のみ DebertaV2Model でロードする必要がある + start_time = time.time() if language == Languages.EN: __loaded_models[language] = cast( DebertaV2Model, @@ -103,7 +105,7 @@ def load_model( revision=revision, ) logger.info( - f"Loaded the {language} BERT model from {pretrained_model_name_or_path}" + f"Loaded the {language.name} BERT model from {pretrained_model_name_or_path} ({time.time() - start_time:.2f}s)" ) return __loaded_models[language] @@ -146,7 +148,7 @@ def load_tokenizer( if pretrained_model_name_or_path is None: assert DEFAULT_BERT_MODEL_PATHS[ language - ].exists(), f"The default {language} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model." + ].exists(), f"The default {language.name} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model." pretrained_model_name_or_path = str(DEFAULT_BERT_MODEL_PATHS[language]) # BERT トークナイザーをロードし、辞書に格納して返す @@ -165,7 +167,7 @@ def load_tokenizer( use_fast=True, # デフォルトで True だが念のため明示的に指定 ) logger.info( - f"Loaded the {language} BERT tokenizer from {pretrained_model_name_or_path}" + f"Loaded the {language.name} BERT tokenizer from {pretrained_model_name_or_path}" ) return __loaded_tokenizers[language] @@ -183,7 +185,7 @@ def transfer_model(language: Languages, device: str) -> None: """ if language not in __loaded_models: - raise ValueError(f"BERT model for {language} is not loaded.") + raise ValueError(f"BERT model for {language.name} is not loaded.") # 既に指定されたデバイスにモデルがロードされている場合は何もしない # ex: current_device="cuda:0", device="cuda" → 何もしない @@ -194,7 +196,7 @@ def transfer_model(language: Languages, device: str) -> None: __loaded_models[language].to(device) # type: ignore logger.info( - f"Transferred the {language} BERT model from {current_device} to {device}" + f"Transferred the {language.name} BERT model from {current_device} to {device}" ) @@ -211,7 +213,7 @@ def unload_model(language: Languages) -> None: gc.collect() if torch.cuda.is_available(): torch.cuda.empty_cache() - logger.info(f"Unloaded the {language} BERT model") + logger.info(f"Unloaded the {language.name} BERT model") def unload_tokenizer(language: Languages) -> None: @@ -227,7 +229,7 @@ def unload_tokenizer(language: Languages) -> None: gc.collect() if torch.cuda.is_available(): torch.cuda.empty_cache() - logger.info(f"Unloaded the {language} BERT tokenizer") + logger.info(f"Unloaded the {language.name} BERT tokenizer") def unload_all_models() -> None: diff --git a/style_bert_vits2/nlp/onnx_bert_models.py b/style_bert_vits2/nlp/onnx_bert_models.py index 3339496..b0ad870 100644 --- a/style_bert_vits2/nlp/onnx_bert_models.py +++ b/style_bert_vits2/nlp/onnx_bert_models.py @@ -10,6 +10,7 @@ Style-Bert-VITS2 の ONNX 推論に必要な各言語ごとの ONNX 版 BERT モ """ import gc +import time from pathlib import Path from typing import Any, Optional, Sequence, Union @@ -73,7 +74,7 @@ def load_model( if pretrained_model_name_or_path is None: assert DEFAULT_ONNX_BERT_MODEL_PATHS[ language - ].exists(), f"The default {language} ONNX BERT model does not exist on the file system. Please specify the path to the pre-trained model." + ].exists(), f"The default {language.name} ONNX BERT model does not exist on the file system. Please specify the path to the pre-trained model." pretrained_model_name_or_path = str(DEFAULT_ONNX_BERT_MODEL_PATHS[language]) # pretrained_model_name_or_path に Hugging Face のリポジトリ名が指定された場合 (aaaa/bbbb のフォーマットを想定): @@ -93,12 +94,13 @@ def load_model( model_path = Path(pretrained_model_name_or_path).resolve() / "model.onnx" # BERT モデルをロードし、辞書に格納して返す + start_time = time.time() __loaded_models[language] = onnxruntime.InferenceSession( model_path, providers=onnx_providers, ) logger.info( - f"Loaded the {language} ONNX BERT model from {pretrained_model_name_or_path}" + f"Loaded the {language.name} ONNX BERT model from {pretrained_model_name_or_path} ({time.time() - start_time:.2f}s)" ) return __loaded_models[language] @@ -139,7 +141,7 @@ def load_tokenizer( if pretrained_model_name_or_path is None: assert DEFAULT_ONNX_BERT_MODEL_PATHS[ language - ].exists(), f"The default {language} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model." + ].exists(), f"The default {language.name} BERT tokenizer does not exist on the file system. Please specify the path to the pre-trained model." pretrained_model_name_or_path = str(DEFAULT_ONNX_BERT_MODEL_PATHS[language]) # BERT トークナイザーをロードし、辞書に格納して返す @@ -158,7 +160,7 @@ def load_tokenizer( use_fast=True, # デフォルトで True だが念のため明示的に指定 ) logger.info( - f"Loaded the {language} ONNX BERT tokenizer from {pretrained_model_name_or_path}" + f"Loaded the {language.name} ONNX BERT tokenizer from {pretrained_model_name_or_path}" ) return __loaded_tokenizers[language] @@ -175,7 +177,7 @@ def unload_model(language: Languages) -> None: if language in __loaded_models: del __loaded_models[language] gc.collect() - logger.info(f"Unloaded the {language} ONNX BERT model") + logger.info(f"Unloaded the {language.name} ONNX BERT model") def unload_tokenizer(language: Languages) -> None: @@ -189,7 +191,7 @@ def unload_tokenizer(language: Languages) -> None: if language in __loaded_tokenizers: del __loaded_tokenizers[language] gc.collect() - logger.info(f"Unloaded the {language} ONNX BERT tokenizer") + logger.info(f"Unloaded the {language.name} ONNX BERT tokenizer") def unload_all_models() -> None: