From 1f080322dcc484dcc86fcb052771b88b0f2d30d7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stardust=C2=B7=E5=87=8F?= Date: Tue, 5 Sep 2023 19:40:15 +0800 Subject: [PATCH] =?UTF-8?q?=E6=B7=BB=E5=8A=A0device=E4=BC=A0=E5=8F=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- text/japanese_bert.py | 13 +++---------- 1 file changed, 3 insertions(+), 10 deletions(-) diff --git a/text/japanese_bert.py b/text/japanese_bert.py index bb83696..3bb8688 100644 --- a/text/japanese_bert.py +++ b/text/japanese_bert.py @@ -1,20 +1,13 @@ import torch from transformers import AutoTokenizer, AutoModelForMaskedLM -device = torch.device( - "cuda" - if torch.cuda.is_available() - else ( - "mps" - if sys.platform == "darwin" and torch.backends.mps.is_available() - else "cpu" - ) - ) tokenizer = AutoTokenizer.from_pretrained("./bert/bert-base-japanese-v3") -model = AutoModelForMaskedLM.from_pretrained("./bert/bert-base-japanese-v3").to(device) def get_bert_feature(text, word2ph): + if sys.platform == "darwin" and torch.backends.mps.is_available() and device == "cpu" + device == "mps" + model = AutoModelForMaskedLM.from_pretrained("./bert/bert-base-japanese-v3").to(device) with torch.no_grad(): inputs = tokenizer(text, return_tensors='pt') for i in inputs: