Fix: ONNX BERT model/tokenizer is not preloaded by default to avoid wasting VRAM
This commit is contained in:
@@ -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,6 +195,8 @@ 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)
|
||||
# VRAM を浪費しないように、既定では ONNX 版 BERT モデル/トークナイザーは事前ロードしない
|
||||
if args.preload_onnx_bert:
|
||||
onnx_bert_models.load_model(
|
||||
Languages.JP, onnx_providers=torch_device_to_onnx_providers(device)
|
||||
)
|
||||
|
||||
@@ -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,6 +105,8 @@ if __name__ == "__main__":
|
||||
bert_models.load_tokenizer(Languages.EN)
|
||||
bert_models.load_model(Languages.ZH, device_map=device)
|
||||
bert_models.load_tokenizer(Languages.ZH)
|
||||
# VRAM を浪費しないように、既定では ONNX 版 BERT モデル/トークナイザーは事前ロードしない
|
||||
if args.preload_onnx_bert:
|
||||
onnx_bert_models.load_model(
|
||||
Languages.JP, onnx_providers=torch_device_to_onnx_providers(device)
|
||||
)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user