From 80dd6dbc22452cbf4e967e566c4a96036589026d Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sat, 10 Aug 2024 16:56:19 +0900 Subject: [PATCH] Improve: Use the device_map option in transformers to load BERT models directly to the GPU --- pyproject.toml | 1 + requirements-colab.txt | 1 + requirements-infer.txt | 1 + requirements.txt | 1 + server_editor.py | 12 ++++++------ server_fastapi.py | 18 +++++++++--------- style_bert_vits2/nlp/bert_models.py | 28 +++++++++++++++++++++++----- 7 files changed, 42 insertions(+), 20 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 795fd76..acaf305 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -22,6 +22,7 @@ classifiers = [ "Programming Language :: Python :: Implementation :: CPython", ] dependencies = [ + "accelerate", "cmudict", "cn2an", "g2p_en", diff --git a/requirements-colab.txt b/requirements-colab.txt index b34b256..39ac92b 100644 --- a/requirements-colab.txt +++ b/requirements-colab.txt @@ -1,3 +1,4 @@ +accelerate cmudict cn2an g2p_en diff --git a/requirements-infer.txt b/requirements-infer.txt index 336b64f..14f770c 100644 --- a/requirements-infer.txt +++ b/requirements-infer.txt @@ -1,3 +1,4 @@ +accelerate cmudict cn2an # faster-whisper==0.10.1 diff --git a/requirements.txt b/requirements.txt index 1d54159..253b8a5 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,3 +1,4 @@ +accelerate cmudict cn2an faster-whisper==0.10.1 diff --git a/server_editor.py b/server_editor.py index cde2c7d..21b53ca 100644 --- a/server_editor.py +++ b/server_editor.py @@ -156,12 +156,6 @@ pyopenjtalk.initialize_worker() # pyopenjtalk の辞書を更新 update_dict() -# 事前に BERT モデル/トークナイザーをロードしておく -## ここでロードしなくても必要になった際に自動ロードされるが、時間がかかるため事前にロードしておいた方が体験が良い -## server_editor.py は日本語にしか対応していないため、日本語の BERT モデル/トークナイザーのみロードする -bert_models.load_model(Languages.JP) -bert_models.load_tokenizer(Languages.JP) - class AudioResponse(Response): media_type = "audio/wav" @@ -194,6 +188,12 @@ port = int(args.port) # download_default_models() 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) if len(model_holder.model_names) == 0: logger.error(f"Models not found in {model_dir}.") diff --git a/server_fastapi.py b/server_fastapi.py index 38f58a2..0ce8624 100644 --- a/server_fastapi.py +++ b/server_fastapi.py @@ -51,15 +51,6 @@ pyopenjtalk.initialize_worker() # dict_data/ 以下の辞書データを pyopenjtalk に適用 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): logger.warning(f"Validation error: {msg}") @@ -104,6 +95,15 @@ if __name__ == "__main__": else: 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_holder = TTSModelHolder(model_dir, device) if len(model_holder.model_names) == 0: diff --git a/style_bert_vits2/nlp/bert_models.py b/style_bert_vits2/nlp/bert_models.py index 1dcadb5..336da9f 100644 --- a/style_bert_vits2/nlp/bert_models.py +++ b/style_bert_vits2/nlp/bert_models.py @@ -38,6 +38,9 @@ __loaded_tokenizers: dict[ def load_model( language: Languages, 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, revision: str = "main", ) -> Union[PreTrainedModel, DebertaV2Model]: @@ -46,6 +49,7 @@ def load_model( 一度ロードされていれば、ロード済みの BERT モデルを即座に返す。 ライブラリ利用時は常に必ず pretrain_model_name_or_path (Hugging Face のリポジトリ名 or ローカルのファイルパス) を指定する必要がある。 ロードにはそれなりに時間がかかるため、ライブラリ利用前に明示的に pretrained_model_name_or_path を指定してロードしておくべき。 + device_map は既に指定された言語の BERT モデルがロードされている場合は効果がない。 cache_dir と revision は pretrain_model_name_or_path がリポジトリ名の場合のみ有効。 Style-Bert-VITS2 では、BERT モデルに下記の 3 つが利用されている。 @@ -57,6 +61,9 @@ def load_model( Args: language (Languages): ロードする学習済みモデルの対象言語 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) revision (str): モデルの Hugging Face 上の Git リビジョン。指定しない場合は最新の main ブランチの内容が利用される (デフォルト: None) @@ -81,12 +88,18 @@ def load_model( __loaded_models[language] = cast( DebertaV2Model, 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: __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( 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: 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) - if current_device != device: - __loaded_models[language].to(device) # type: ignore - logger.info( + if current_device.startswith(device): + return + + __loaded_models[language].to(device) # type: ignore + logger.info( f"Transferred the {language} BERT model from {current_device} to {device}" )