From cfd6bd5bb65a9869e955a3c69c6a3ab5b0a2bf4a Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sun, 11 Aug 2024 08:40:50 +0900 Subject: [PATCH] Improve: Models can now be loaded directly onto NVIDIA GPUs --- style_bert_vits2/models/infer.py | 4 ++-- style_bert_vits2/models/utils/checkpoints.py | 3 ++- style_bert_vits2/models/utils/safetensors.py | 3 ++- style_bert_vits2/nlp/bert_models.py | 4 ++-- 4 files changed, 8 insertions(+), 6 deletions(-) diff --git a/style_bert_vits2/models/infer.py b/style_bert_vits2/models/infer.py index a9a4869..9b83803 100644 --- a/style_bert_vits2/models/infer.py +++ b/style_bert_vits2/models/infer.py @@ -86,10 +86,10 @@ def get_net_g(model_path: str, version: str, device: str, hps: HyperParameters): _ = net_g.eval() if model_path.endswith(".pth") or model_path.endswith(".pt"): _ = utils.checkpoints.load_checkpoint( - model_path, net_g, None, skip_optimizer=True + model_path, net_g, None, skip_optimizer=True, device=device ) elif model_path.endswith(".safetensors"): - _ = utils.safetensors.load_safetensors(model_path, net_g, True) + _ = utils.safetensors.load_safetensors(model_path, net_g, True, device=device) else: raise ValueError(f"Unknown model format: {model_path}") return net_g diff --git a/style_bert_vits2/models/utils/checkpoints.py b/style_bert_vits2/models/utils/checkpoints.py index c601dda..756586c 100644 --- a/style_bert_vits2/models/utils/checkpoints.py +++ b/style_bert_vits2/models/utils/checkpoints.py @@ -15,6 +15,7 @@ def load_checkpoint( optimizer: Optional[torch.optim.Optimizer] = None, skip_optimizer: bool = False, for_infer: bool = False, + device: Union[str, torch.device] = "cpu", ) -> tuple[torch.nn.Module, Optional[torch.optim.Optimizer], float, int]: """ 指定されたパスからチェックポイントを読み込み、モデルとオプティマイザーを更新する。 @@ -31,7 +32,7 @@ def load_checkpoint( """ assert os.path.isfile(checkpoint_path) - checkpoint_dict = torch.load(checkpoint_path, map_location="cpu") + checkpoint_dict = torch.load(checkpoint_path, map_location=device) iteration = checkpoint_dict["iteration"] learning_rate = checkpoint_dict["learning_rate"] logger.info( diff --git a/style_bert_vits2/models/utils/safetensors.py b/style_bert_vits2/models/utils/safetensors.py index 4b4ef3f..024fe06 100644 --- a/style_bert_vits2/models/utils/safetensors.py +++ b/style_bert_vits2/models/utils/safetensors.py @@ -12,6 +12,7 @@ def load_safetensors( checkpoint_path: Union[str, Path], model: torch.nn.Module, for_infer: bool = False, + device: Union[str, torch.device] = "cpu", ) -> tuple[torch.nn.Module, Optional[int]]: """ 指定されたパスから safetensors モデルを読み込み、モデルとイテレーションを返す。 @@ -27,7 +28,7 @@ def load_safetensors( tensors: dict[str, Any] = {} iteration: Optional[int] = None - with safe_open(str(checkpoint_path), framework="pt", device="cpu") as f: # type: ignore + with safe_open(str(checkpoint_path), framework="pt", device=device) as f: # type: ignore for key in f.keys(): if key == "iteration": iteration = f.get_tensor(key).item() diff --git a/style_bert_vits2/nlp/bert_models.py b/style_bert_vits2/nlp/bert_models.py index 336da9f..37d8cd5 100644 --- a/style_bert_vits2/nlp/bert_models.py +++ b/style_bert_vits2/nlp/bert_models.py @@ -192,8 +192,8 @@ def transfer_model(language: Languages, device: str) -> None: __loaded_models[language].to(device) # type: ignore logger.info( - f"Transferred the {language} BERT model from {current_device} to {device}" - ) + f"Transferred the {language} BERT model from {current_device} to {device}" + ) def unload_model(language: Languages) -> None: