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("--line_count", type=int, default=None)
|
||||||
# parser.add_argument("--skip_default_models", action="store_true")
|
# parser.add_argument("--skip_default_models", action="store_true")
|
||||||
parser.add_argument("--skip_static_files", 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()
|
args = parser.parse_args()
|
||||||
device = args.device
|
device = args.device
|
||||||
if device == "cuda" and not torch.cuda.is_available():
|
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 モデル/トークナイザーのみロードする
|
## server_editor.py は日本語にしか対応していないため、日本語の BERT モデル/トークナイザーのみロードする
|
||||||
bert_models.load_model(Languages.JP, device_map=device)
|
bert_models.load_model(Languages.JP, device_map=device)
|
||||||
bert_models.load_tokenizer(Languages.JP)
|
bert_models.load_tokenizer(Languages.JP)
|
||||||
onnx_bert_models.load_model(
|
# VRAM を浪費しないように、既定では ONNX 版 BERT モデル/トークナイザーは事前ロードしない
|
||||||
Languages.JP, onnx_providers=torch_device_to_onnx_providers(device)
|
if args.preload_onnx_bert:
|
||||||
)
|
onnx_bert_models.load_model(
|
||||||
onnx_bert_models.load_tokenizer(Languages.JP)
|
Languages.JP, onnx_providers=torch_device_to_onnx_providers(device)
|
||||||
|
)
|
||||||
|
onnx_bert_models.load_tokenizer(Languages.JP)
|
||||||
|
|
||||||
model_holder = TTSModelHolder(
|
model_holder = TTSModelHolder(
|
||||||
model_dir, device, torch_device_to_onnx_providers(device)
|
model_dir, device, torch_device_to_onnx_providers(device)
|
||||||
|
|||||||
@@ -89,6 +89,7 @@ if __name__ == "__main__":
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--dir", "-d", type=str, help="Model directory", default=config.assets_root
|
"--dir", "-d", type=str, help="Model directory", default=config.assets_root
|
||||||
)
|
)
|
||||||
|
parser.add_argument("--preload_onnx_bert", action="store_true")
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
if args.cpu:
|
if args.cpu:
|
||||||
@@ -104,10 +105,12 @@ if __name__ == "__main__":
|
|||||||
bert_models.load_tokenizer(Languages.EN)
|
bert_models.load_tokenizer(Languages.EN)
|
||||||
bert_models.load_model(Languages.ZH, device_map=device)
|
bert_models.load_model(Languages.ZH, device_map=device)
|
||||||
bert_models.load_tokenizer(Languages.ZH)
|
bert_models.load_tokenizer(Languages.ZH)
|
||||||
onnx_bert_models.load_model(
|
# VRAM を浪費しないように、既定では ONNX 版 BERT モデル/トークナイザーは事前ロードしない
|
||||||
Languages.JP, onnx_providers=torch_device_to_onnx_providers(device)
|
if args.preload_onnx_bert:
|
||||||
)
|
onnx_bert_models.load_model(
|
||||||
onnx_bert_models.load_tokenizer(Languages.JP)
|
Languages.JP, onnx_providers=torch_device_to_onnx_providers(device)
|
||||||
|
)
|
||||||
|
onnx_bert_models.load_tokenizer(Languages.JP)
|
||||||
|
|
||||||
model_dir = Path(args.dir)
|
model_dir = Path(args.dir)
|
||||||
model_holder = TTSModelHolder(
|
model_holder = TTSModelHolder(
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ Style-Bert-VITS2 の学習・推論に必要な各言語ごとの BERT モデル
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import gc
|
import gc
|
||||||
|
import time
|
||||||
from typing import Optional, Union, cast
|
from typing import Optional, Union, cast
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -80,11 +81,12 @@ def load_model(
|
|||||||
if pretrained_model_name_or_path is None:
|
if pretrained_model_name_or_path is None:
|
||||||
assert DEFAULT_BERT_MODEL_PATHS[
|
assert DEFAULT_BERT_MODEL_PATHS[
|
||||||
language
|
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])
|
pretrained_model_name_or_path = str(DEFAULT_BERT_MODEL_PATHS[language])
|
||||||
|
|
||||||
# BERT モデルをロードし、辞書に格納して返す
|
# BERT モデルをロードし、辞書に格納して返す
|
||||||
## 英語のみ DebertaV2Model でロードする必要がある
|
## 英語のみ DebertaV2Model でロードする必要がある
|
||||||
|
start_time = time.time()
|
||||||
if language == Languages.EN:
|
if language == Languages.EN:
|
||||||
__loaded_models[language] = cast(
|
__loaded_models[language] = cast(
|
||||||
DebertaV2Model,
|
DebertaV2Model,
|
||||||
@@ -103,7 +105,7 @@ def load_model(
|
|||||||
revision=revision,
|
revision=revision,
|
||||||
)
|
)
|
||||||
logger.info(
|
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]
|
return __loaded_models[language]
|
||||||
@@ -146,7 +148,7 @@ def load_tokenizer(
|
|||||||
if pretrained_model_name_or_path is None:
|
if pretrained_model_name_or_path is None:
|
||||||
assert DEFAULT_BERT_MODEL_PATHS[
|
assert DEFAULT_BERT_MODEL_PATHS[
|
||||||
language
|
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])
|
pretrained_model_name_or_path = str(DEFAULT_BERT_MODEL_PATHS[language])
|
||||||
|
|
||||||
# BERT トークナイザーをロードし、辞書に格納して返す
|
# BERT トークナイザーをロードし、辞書に格納して返す
|
||||||
@@ -165,7 +167,7 @@ def load_tokenizer(
|
|||||||
use_fast=True, # デフォルトで True だが念のため明示的に指定
|
use_fast=True, # デフォルトで True だが念のため明示的に指定
|
||||||
)
|
)
|
||||||
logger.info(
|
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]
|
return __loaded_tokenizers[language]
|
||||||
@@ -183,7 +185,7 @@ def transfer_model(language: Languages, device: str) -> None:
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
if language not in __loaded_models:
|
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" → 何もしない
|
# 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
|
__loaded_models[language].to(device) # type: ignore
|
||||||
logger.info(
|
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()
|
gc.collect()
|
||||||
if torch.cuda.is_available():
|
if torch.cuda.is_available():
|
||||||
torch.cuda.empty_cache()
|
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:
|
def unload_tokenizer(language: Languages) -> None:
|
||||||
@@ -227,7 +229,7 @@ def unload_tokenizer(language: Languages) -> None:
|
|||||||
gc.collect()
|
gc.collect()
|
||||||
if torch.cuda.is_available():
|
if torch.cuda.is_available():
|
||||||
torch.cuda.empty_cache()
|
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:
|
def unload_all_models() -> None:
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ Style-Bert-VITS2 の ONNX 推論に必要な各言語ごとの ONNX 版 BERT モ
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import gc
|
import gc
|
||||||
|
import time
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Optional, Sequence, Union
|
from typing import Any, Optional, Sequence, Union
|
||||||
|
|
||||||
@@ -73,7 +74,7 @@ def load_model(
|
|||||||
if pretrained_model_name_or_path is None:
|
if pretrained_model_name_or_path is None:
|
||||||
assert DEFAULT_ONNX_BERT_MODEL_PATHS[
|
assert DEFAULT_ONNX_BERT_MODEL_PATHS[
|
||||||
language
|
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 = str(DEFAULT_ONNX_BERT_MODEL_PATHS[language])
|
||||||
|
|
||||||
# pretrained_model_name_or_path に Hugging Face のリポジトリ名が指定された場合 (aaaa/bbbb のフォーマットを想定):
|
# 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"
|
model_path = Path(pretrained_model_name_or_path).resolve() / "model.onnx"
|
||||||
|
|
||||||
# BERT モデルをロードし、辞書に格納して返す
|
# BERT モデルをロードし、辞書に格納して返す
|
||||||
|
start_time = time.time()
|
||||||
__loaded_models[language] = onnxruntime.InferenceSession(
|
__loaded_models[language] = onnxruntime.InferenceSession(
|
||||||
model_path,
|
model_path,
|
||||||
providers=onnx_providers,
|
providers=onnx_providers,
|
||||||
)
|
)
|
||||||
logger.info(
|
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]
|
return __loaded_models[language]
|
||||||
@@ -139,7 +141,7 @@ def load_tokenizer(
|
|||||||
if pretrained_model_name_or_path is None:
|
if pretrained_model_name_or_path is None:
|
||||||
assert DEFAULT_ONNX_BERT_MODEL_PATHS[
|
assert DEFAULT_ONNX_BERT_MODEL_PATHS[
|
||||||
language
|
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])
|
pretrained_model_name_or_path = str(DEFAULT_ONNX_BERT_MODEL_PATHS[language])
|
||||||
|
|
||||||
# BERT トークナイザーをロードし、辞書に格納して返す
|
# BERT トークナイザーをロードし、辞書に格納して返す
|
||||||
@@ -158,7 +160,7 @@ def load_tokenizer(
|
|||||||
use_fast=True, # デフォルトで True だが念のため明示的に指定
|
use_fast=True, # デフォルトで True だが念のため明示的に指定
|
||||||
)
|
)
|
||||||
logger.info(
|
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]
|
return __loaded_tokenizers[language]
|
||||||
@@ -175,7 +177,7 @@ def unload_model(language: Languages) -> None:
|
|||||||
if language in __loaded_models:
|
if language in __loaded_models:
|
||||||
del __loaded_models[language]
|
del __loaded_models[language]
|
||||||
gc.collect()
|
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:
|
def unload_tokenizer(language: Languages) -> None:
|
||||||
@@ -189,7 +191,7 @@ def unload_tokenizer(language: Languages) -> None:
|
|||||||
if language in __loaded_tokenizers:
|
if language in __loaded_tokenizers:
|
||||||
del __loaded_tokenizers[language]
|
del __loaded_tokenizers[language]
|
||||||
gc.collect()
|
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:
|
def unload_all_models() -> None:
|
||||||
|
|||||||
Reference in New Issue
Block a user