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

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