mae-remote-sensing / utils /checkpoint.py
ANMOL1140W's picture
app done
78627b9
Raw
History Blame Contribute Delete
569 Bytes
import torch
import torch.nn as nn
def load_checkpoint(model: nn.Module, filename: str, device: torch.device | None = None) -> dict:
"""Load model weights from a training checkpoint."""
device = device or torch.device("cpu")
checkpoint = torch.load(filename, map_location=device, weights_only=False)
if not isinstance(checkpoint, dict) or "state_dict" not in checkpoint:
raise ValueError(f"Checkpoint at {filename} does not contain a 'state_dict' key.")
model.load_state_dict(checkpoint["state_dict"], strict=True)
return checkpoint