Improve: Use the device_map option in transformers to load BERT models directly to the GPU
This commit is contained in:
@@ -22,6 +22,7 @@ classifiers = [
|
|||||||
"Programming Language :: Python :: Implementation :: CPython",
|
"Programming Language :: Python :: Implementation :: CPython",
|
||||||
]
|
]
|
||||||
dependencies = [
|
dependencies = [
|
||||||
|
"accelerate",
|
||||||
"cmudict",
|
"cmudict",
|
||||||
"cn2an",
|
"cn2an",
|
||||||
"g2p_en",
|
"g2p_en",
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
accelerate
|
||||||
cmudict
|
cmudict
|
||||||
cn2an
|
cn2an
|
||||||
g2p_en
|
g2p_en
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
accelerate
|
||||||
cmudict
|
cmudict
|
||||||
cn2an
|
cn2an
|
||||||
# faster-whisper==0.10.1
|
# faster-whisper==0.10.1
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
accelerate
|
||||||
cmudict
|
cmudict
|
||||||
cn2an
|
cn2an
|
||||||
faster-whisper==0.10.1
|
faster-whisper==0.10.1
|
||||||
|
|||||||
@@ -156,12 +156,6 @@ pyopenjtalk.initialize_worker()
|
|||||||
# pyopenjtalk の辞書を更新
|
# pyopenjtalk の辞書を更新
|
||||||
update_dict()
|
update_dict()
|
||||||
|
|
||||||
# 事前に BERT モデル/トークナイザーをロードしておく
|
|
||||||
## ここでロードしなくても必要になった際に自動ロードされるが、時間がかかるため事前にロードしておいた方が体験が良い
|
|
||||||
## server_editor.py は日本語にしか対応していないため、日本語の BERT モデル/トークナイザーのみロードする
|
|
||||||
bert_models.load_model(Languages.JP)
|
|
||||||
bert_models.load_tokenizer(Languages.JP)
|
|
||||||
|
|
||||||
|
|
||||||
class AudioResponse(Response):
|
class AudioResponse(Response):
|
||||||
media_type = "audio/wav"
|
media_type = "audio/wav"
|
||||||
@@ -194,6 +188,12 @@ port = int(args.port)
|
|||||||
# download_default_models()
|
# download_default_models()
|
||||||
skip_static_files = bool(args.skip_static_files)
|
skip_static_files = bool(args.skip_static_files)
|
||||||
|
|
||||||
|
# 事前に BERT モデル/トークナイザーをロードしておく
|
||||||
|
## ここでロードしなくても必要になった際に自動ロードされるが、時間がかかるため事前にロードしておいた方が体験が良い
|
||||||
|
## server_editor.py は日本語にしか対応していないため、日本語の BERT モデル/トークナイザーのみロードする
|
||||||
|
bert_models.load_model(Languages.JP, device_map=device)
|
||||||
|
bert_models.load_tokenizer(Languages.JP)
|
||||||
|
|
||||||
model_holder = TTSModelHolder(model_dir, device)
|
model_holder = TTSModelHolder(model_dir, device)
|
||||||
if len(model_holder.model_names) == 0:
|
if len(model_holder.model_names) == 0:
|
||||||
logger.error(f"Models not found in {model_dir}.")
|
logger.error(f"Models not found in {model_dir}.")
|
||||||
|
|||||||
@@ -51,15 +51,6 @@ pyopenjtalk.initialize_worker()
|
|||||||
# dict_data/ 以下の辞書データを pyopenjtalk に適用
|
# dict_data/ 以下の辞書データを pyopenjtalk に適用
|
||||||
update_dict()
|
update_dict()
|
||||||
|
|
||||||
# 事前に BERT モデル/トークナイザーをロードしておく
|
|
||||||
## ここでロードしなくても必要になった際に自動ロードされるが、時間がかかるため事前にロードしておいた方が体験が良い
|
|
||||||
bert_models.load_model(Languages.JP)
|
|
||||||
bert_models.load_tokenizer(Languages.JP)
|
|
||||||
bert_models.load_model(Languages.EN)
|
|
||||||
bert_models.load_tokenizer(Languages.EN)
|
|
||||||
bert_models.load_model(Languages.ZH)
|
|
||||||
bert_models.load_tokenizer(Languages.ZH)
|
|
||||||
|
|
||||||
|
|
||||||
def raise_validation_error(msg: str, param: str):
|
def raise_validation_error(msg: str, param: str):
|
||||||
logger.warning(f"Validation error: {msg}")
|
logger.warning(f"Validation error: {msg}")
|
||||||
@@ -104,6 +95,15 @@ if __name__ == "__main__":
|
|||||||
else:
|
else:
|
||||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
|
||||||
|
# 事前に BERT モデル/トークナイザーをロードしておく
|
||||||
|
## ここでロードしなくても必要になった際に自動ロードされるが、時間がかかるため事前にロードしておいた方が体験が良い
|
||||||
|
bert_models.load_model(Languages.JP, device_map=device)
|
||||||
|
bert_models.load_tokenizer(Languages.JP)
|
||||||
|
bert_models.load_model(Languages.EN, device_map=device)
|
||||||
|
bert_models.load_tokenizer(Languages.EN)
|
||||||
|
bert_models.load_model(Languages.ZH, device_map=device)
|
||||||
|
bert_models.load_tokenizer(Languages.ZH)
|
||||||
|
|
||||||
model_dir = Path(args.dir)
|
model_dir = Path(args.dir)
|
||||||
model_holder = TTSModelHolder(model_dir, device)
|
model_holder = TTSModelHolder(model_dir, device)
|
||||||
if len(model_holder.model_names) == 0:
|
if len(model_holder.model_names) == 0:
|
||||||
|
|||||||
@@ -38,6 +38,9 @@ __loaded_tokenizers: dict[
|
|||||||
def load_model(
|
def load_model(
|
||||||
language: Languages,
|
language: Languages,
|
||||||
pretrained_model_name_or_path: Optional[str] = None,
|
pretrained_model_name_or_path: Optional[str] = None,
|
||||||
|
device_map: Optional[
|
||||||
|
Union[str, dict[str, Union[int, str, torch.device]], int, torch.device]
|
||||||
|
] = None,
|
||||||
cache_dir: Optional[str] = None,
|
cache_dir: Optional[str] = None,
|
||||||
revision: str = "main",
|
revision: str = "main",
|
||||||
) -> Union[PreTrainedModel, DebertaV2Model]:
|
) -> Union[PreTrainedModel, DebertaV2Model]:
|
||||||
@@ -46,6 +49,7 @@ def load_model(
|
|||||||
一度ロードされていれば、ロード済みの BERT モデルを即座に返す。
|
一度ロードされていれば、ロード済みの BERT モデルを即座に返す。
|
||||||
ライブラリ利用時は常に必ず pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある。
|
ライブラリ利用時は常に必ず pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある。
|
||||||
ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき。
|
ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき。
|
||||||
|
device_map は既に指定された言語の BERT モデルがロードされている場合は効果がない。
|
||||||
cache_dir と revision は pretrain_model_name_or_path がリポジトリ名の場合のみ有効。
|
cache_dir と revision は pretrain_model_name_or_path がリポジトリ名の場合のみ有効。
|
||||||
|
|
||||||
Style-Bert-VITS2 では、BERT モデルに下記の 3 つが利用されている。
|
Style-Bert-VITS2 では、BERT モデルに下記の 3 つが利用されている。
|
||||||
@@ -57,6 +61,9 @@ def load_model(
|
|||||||
Args:
|
Args:
|
||||||
language (Languages): ロードする学習済みモデルの対象言語
|
language (Languages): ロードする学習済みモデルの対象言語
|
||||||
pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None)
|
pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None)
|
||||||
|
device_map (Optional[str]): accelerate を使用して高速にデバイスにモデルをロードするためのデバイスマップ。
|
||||||
|
指定しない場合は通常のモデルロード処理になる (デフォルト: None)
|
||||||
|
ref: https://huggingface.co/docs/accelerate/usage_guides/big_modeling
|
||||||
cache_dir (Optional[str]): モデルのキャッシュディレクトリ。指定しない場合はデフォルトのキャッシュディレクトリが利用される (デフォルト: None)
|
cache_dir (Optional[str]): モデルのキャッシュディレクトリ。指定しない場合はデフォルトのキャッシュディレクトリが利用される (デフォルト: None)
|
||||||
revision (str): モデルの Hugging Face 上の Git リビジョン。指定しない場合は最新の main ブランチの内容が利用される (デフォルト: None)
|
revision (str): モデルの Hugging Face 上の Git リビジョン。指定しない場合は最新の main ブランチの内容が利用される (デフォルト: None)
|
||||||
|
|
||||||
@@ -81,12 +88,18 @@ def load_model(
|
|||||||
__loaded_models[language] = cast(
|
__loaded_models[language] = cast(
|
||||||
DebertaV2Model,
|
DebertaV2Model,
|
||||||
DebertaV2Model.from_pretrained(
|
DebertaV2Model.from_pretrained(
|
||||||
pretrained_model_name_or_path, cache_dir=cache_dir, revision=revision
|
pretrained_model_name_or_path,
|
||||||
|
device_map=device_map,
|
||||||
|
cache_dir=cache_dir,
|
||||||
|
revision=revision,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
__loaded_models[language] = AutoModelForMaskedLM.from_pretrained(
|
__loaded_models[language] = AutoModelForMaskedLM.from_pretrained(
|
||||||
pretrained_model_name_or_path, cache_dir=cache_dir, revision=revision
|
pretrained_model_name_or_path,
|
||||||
|
device_map=device_map,
|
||||||
|
cache_dir=cache_dir,
|
||||||
|
revision=revision,
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Loaded the {language} BERT model from {pretrained_model_name_or_path}"
|
f"Loaded the {language} BERT model from {pretrained_model_name_or_path}"
|
||||||
@@ -170,10 +183,15 @@ 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} is not loaded.")
|
||||||
|
|
||||||
|
# 既に指定されたデバイスにモデルがロードされている場合は何もしない
|
||||||
|
# ex: current_device="cuda:0", device="cuda" → 何もしない
|
||||||
|
# ex: current_device="cuda:0", device="cpu" → モデルを CPU に移動
|
||||||
current_device = str(__loaded_models[language].device)
|
current_device = str(__loaded_models[language].device)
|
||||||
if current_device != device:
|
if current_device.startswith(device):
|
||||||
__loaded_models[language].to(device) # type: ignore
|
return
|
||||||
logger.info(
|
|
||||||
|
__loaded_models[language].to(device) # type: ignore
|
||||||
|
logger.info(
|
||||||
f"Transferred the {language} BERT model from {current_device} to {device}"
|
f"Transferred the {language} BERT model from {current_device} to {device}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user