Apex / checkpoint.py
yadavkapil7155's picture
Upload folder using huggingface_hub
1e4fb05 verified
Raw
History Blame Contribute Delete
1.95 kB
"""
Checkpointing utilities for saving and restoring training state
"""
import os
import shutil
import torch
def save_checkpoint(path: str, model, optimizer=None, scaler=None, step=0, val_loss=None, config=None):
os.makedirs(os.path.dirname(path), exist_ok=True)
raw_model = model.module if hasattr(model, "module") else model
state = {
"step": step,
"model_state_dict": raw_model.state_dict(),
"val_loss": val_loss,
"config": config.__dict__ if config is not None and hasattr(config, "__dict__") else config,
}
if optimizer is not None:
state["optimizer_state_dict"] = optimizer.state_dict()
if scaler is not None:
state["scaler_state_dict"] = scaler.state_dict()
torch.save(state, path)
def mirror_checkpoint(path: str, mirror_dir: str):
if not mirror_dir:
return None
os.makedirs(mirror_dir, exist_ok=True)
mirror_path = os.path.join(mirror_dir, os.path.basename(path))
shutil.copy2(path, mirror_path)
return mirror_path
def load_checkpoint(path: str, model, optimizer=None, scaler=None, map_location="cpu"):
if not os.path.exists(path):
raise FileNotFoundError(f"Checkpoint file not found: {path}")
print(f"[*] Loading checkpoint from {path}...")
checkpoint = torch.load(path, map_location=map_location)
raw_model = model.module if hasattr(model, "module") else model
raw_model.load_state_dict(checkpoint["model_state_dict"])
if optimizer is not None and "optimizer_state_dict" in checkpoint:
optimizer.load_state_dict(checkpoint["optimizer_state_dict"])
if scaler is not None and "scaler_state_dict" in checkpoint:
scaler.load_state_dict(checkpoint["scaler_state_dict"])
step = checkpoint.get("step", 0)
val_loss = checkpoint.get("val_loss", None)
print(f"[*] Successfully restored checkpoint at step {step} (val_loss: {val_loss})")
return step, val_loss