Improve: Use the device_map option in transformers to load BERT models directly to the GPU

This commit is contained in:
tsukumi
2024-08-10 16:56:19 +09:00
parent bdf20e8911
commit 80dd6dbc22
7 changed files with 42 additions and 20 deletions

View File

@@ -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",

View File

@@ -1,3 +1,4 @@
accelerate
cmudict cmudict
cn2an cn2an
g2p_en g2p_en

View File

@@ -1,3 +1,4 @@
accelerate
cmudict cmudict
cn2an cn2an
# faster-whisper==0.10.1 # faster-whisper==0.10.1

View File

@@ -1,3 +1,4 @@
accelerate
cmudict cmudict
cn2an cn2an
faster-whisper==0.10.1 faster-whisper==0.10.1

View File

@@ -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}.")

View File

@@ -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:

View File

@@ -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}"
) )