Improve: Models can now be loaded directly onto NVIDIA GPUs
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user