GID-Flow / PDGrapher /src /gidflow /utils /checkpoint.py
Boom5426's picture
Upload GID-Flow project snapshot (deduped: code + key artifacts)
07fcdfe verified
Raw
History Blame Contribute Delete
1.15 kB
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