Refactor: Use bert_models.transfer_model()
This commit is contained in:
@@ -80,10 +80,12 @@ cov = ["test-cov", "cov-report"]
|
|||||||
detached = true
|
detached = true
|
||||||
dependencies = ["black[jupyter]", "isort"]
|
dependencies = ["black[jupyter]", "isort"]
|
||||||
[tool.hatch.envs.style.scripts]
|
[tool.hatch.envs.style.scripts]
|
||||||
|
# Usage: `hatch run style:check`
|
||||||
check = [
|
check = [
|
||||||
"black --check --diff .",
|
"black --check --diff .",
|
||||||
"isort --check-only --diff --profile black --gitignore --lai 2 . --sg \"Data/*\" --sg \"inputs/*\" --sg \"model_assets/*\" --sg \"static/*\"",
|
"isort --check-only --diff --profile black --gitignore --lai 2 . --sg \"Data/*\" --sg \"inputs/*\" --sg \"model_assets/*\" --sg \"static/*\"",
|
||||||
]
|
]
|
||||||
|
# Usage: `hatch run style:fmt`
|
||||||
fmt = [
|
fmt = [
|
||||||
"black .",
|
"black .",
|
||||||
"isort --profile black --gitignore --lai 2 . --sg \"Data/*\" --sg \"inputs/*\" --sg \"model_assets/*\" --sg \"static/*\"",
|
"isort --profile black --gitignore --lai 2 . --sg \"Data/*\" --sg \"inputs/*\" --sg \"model_assets/*\" --sg \"static/*\"",
|
||||||
|
|||||||
@@ -159,6 +159,8 @@ def load_tokenizer(
|
|||||||
def transfer_model(language: Languages, device: str) -> None:
|
def transfer_model(language: Languages, device: str) -> None:
|
||||||
"""
|
"""
|
||||||
指定された言語の BERT モデルを、指定されたデバイスに移動する。
|
指定された言語の BERT モデルを、指定されたデバイスに移動する。
|
||||||
|
モデルのロード後に推論デバイスを変更したい場合に利用する。
|
||||||
|
既に指定されたデバイスにモデルがロードされている場合は何も行われない。
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
language (Languages): モデルを移動する言語
|
language (Languages): モデルを移動する言語
|
||||||
@@ -171,7 +173,9 @@ def transfer_model(language: Languages, device: str) -> None:
|
|||||||
current_device = str(__loaded_models[language].device)
|
current_device = str(__loaded_models[language].device)
|
||||||
if current_device != device:
|
if current_device != device:
|
||||||
__loaded_models[language].to(device) # type: ignore
|
__loaded_models[language].to(device) # type: ignore
|
||||||
logger.info(f"Transferred the {language} BERT model from {current_device} to {device}")
|
logger.info(
|
||||||
|
f"Transferred the {language} BERT model from {current_device} to {device}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def unload_model(language: Languages) -> None:
|
def unload_model(language: Languages) -> None:
|
||||||
|
|||||||
@@ -29,7 +29,8 @@ 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).to(device) # type: ignore
|
model = bert_models.load_model(Languages.ZH)
|
||||||
|
bert_models.transfer_model(Languages.ZH, device)
|
||||||
|
|
||||||
style_res_mean = None
|
style_res_mean = None
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
|
|||||||
@@ -29,7 +29,8 @@ 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).to(device) # type: ignore
|
model = bert_models.load_model(Languages.EN)
|
||||||
|
bert_models.transfer_model(Languages.EN, device)
|
||||||
|
|
||||||
style_res_mean = None
|
style_res_mean = None
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
|
|||||||
@@ -36,7 +36,8 @@ 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).to(device) # type: ignore
|
model = bert_models.load_model(Languages.JP)
|
||||||
|
bert_models.transfer_model(Languages.JP, device)
|
||||||
|
|
||||||
style_res_mean = None
|
style_res_mean = None
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
|
|||||||
Reference in New Issue
Block a user