| import os | |
| from typing import Any, Dict, Optional | |
| import torch | |
| import torch.nn as nn | |
| def save_checkpoint( | |
| path: str, | |
| model: nn.Module, | |
| optimizer: Optional[torch.optim.Optimizer] = None, | |
| epoch: int = 0, | |
| metrics: Optional[Dict[str, Any]] = None, | |
| config: Optional[Dict[str, Any]] = None, | |
| ) -> None: | |
| os.makedirs(os.path.dirname(path) or ".", exist_ok=True) | |
| state = { | |
| "epoch": epoch, | |
| "model_state_dict": model.state_dict(), | |
| "metrics": metrics or {}, | |
| } | |
| if optimizer is not None: | |
| state["optimizer_state_dict"] = optimizer.state_dict() | |
| if config is not None: | |
| state["config"] = config | |
| torch.save(state, path) | |
| def load_checkpoint( | |
| path: str, | |
| model: nn.Module, | |
| optimizer: Optional[torch.optim.Optimizer] = None, | |
| device: Optional[torch.device] = None, | |
| ) -> Dict[str, Any]: | |
| state = torch.load(path, map_location=device or "cpu", weights_only=False) | |
| model.load_state_dict(state["model_state_dict"]) | |
| if optimizer is not None and "optimizer_state_dict" in state: | |
| optimizer.load_state_dict(state["optimizer_state_dict"]) | |
| return state | |