Improve: Use the device_map option in transformers to load BERT models directly to the GPU
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user