Improve: Models can now be loaded directly onto NVIDIA GPUs

This commit is contained in:
tsukumi
2024-08-11 08:40:50 +09:00
parent 0e71d7ef7c
commit cfd6bd5bb6
4 changed files with 8 additions and 6 deletions

View File

@@ -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(

View File

@@ -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()