Fix: Specifying the device_map option in extract_bert_feature()
This commit is contained in:
@@ -29,7 +29,7 @@ def extract_bert_feature(
|
|||||||
|
|
||||||
if device == "cuda" and not torch.cuda.is_available():
|
if device == "cuda" and not torch.cuda.is_available():
|
||||||
device = "cpu"
|
device = "cpu"
|
||||||
model = bert_models.load_model(Languages.ZH)
|
model = bert_models.load_model(Languages.ZH, device_map=device)
|
||||||
bert_models.transfer_model(Languages.ZH, device)
|
bert_models.transfer_model(Languages.ZH, device)
|
||||||
|
|
||||||
style_res_mean = None
|
style_res_mean = None
|
||||||
|
|||||||
@@ -29,7 +29,7 @@ def extract_bert_feature(
|
|||||||
|
|
||||||
if device == "cuda" and not torch.cuda.is_available():
|
if device == "cuda" and not torch.cuda.is_available():
|
||||||
device = "cpu"
|
device = "cpu"
|
||||||
model = bert_models.load_model(Languages.EN)
|
model = bert_models.load_model(Languages.EN, device_map=device)
|
||||||
bert_models.transfer_model(Languages.EN, device)
|
bert_models.transfer_model(Languages.EN, device)
|
||||||
|
|
||||||
style_res_mean = None
|
style_res_mean = None
|
||||||
|
|||||||
@@ -36,7 +36,7 @@ def extract_bert_feature(
|
|||||||
|
|
||||||
if device == "cuda" and not torch.cuda.is_available():
|
if device == "cuda" and not torch.cuda.is_available():
|
||||||
device = "cpu"
|
device = "cpu"
|
||||||
model = bert_models.load_model(Languages.JP)
|
model = bert_models.load_model(Languages.JP, device_map=device)
|
||||||
bert_models.transfer_model(Languages.JP, device)
|
bert_models.transfer_model(Languages.JP, device)
|
||||||
|
|
||||||
style_res_mean = None
|
style_res_mean = None
|
||||||
|
|||||||
Reference in New Issue
Block a user