canopy-258m-r3 / canopy_r3 /checkpointing.py
psikosen's picture
Update to v5: Prefix Sliding KV-cache, SMELT scaling, CMA checkpoint, sPTC tool caller
2976bd9 verified
Raw History Blame Contribute Delete
11 kB
"""
Canopy-R3 Best-Practice Checkpointing, Resumption, Atomic Writes,
Storage Safety, and System Resource Monitoring.
"""
import os
import gc
import glob
import shutil
import random
import logging
from pathlib import Path
from typing import Optional, Dict, Any, Tuple, List, Union
import numpy as np
import psutil
import torch
import torch.nn as nn
from safetensors.torch import save_file
logger = logging.getLogger("canopy.checkpoint")
class ResourceMonitor:
"""Monitors CPU RAM, GPU VRAM, and Filesystem Disk Space."""
@staticmethod
def get_system_ram_gib() -> Dict[str, float]:
vm = psutil.virtual_memory()
return {
"total_gib": vm.total / (1024 ** 3),
"available_gib": vm.available / (1024 ** 3),
"used_gib": vm.used / (1024 ** 3),
"percent": vm.percent,
}
@staticmethod
def get_disk_free_gib(path: Union[str, Path] = ".") -> Dict[str, float]:
usage = shutil.disk_usage(path)
return {
"total_gib": usage.total / (1024 ** 3),
"used_gib": usage.used / (1024 ** 3),
"free_gib": usage.free / (1024 ** 3),
"percent": (usage.used / usage.total) * 100.0,
}
@staticmethod
def get_gpu_vram_gib(device: Union[int, str, torch.device] = "cuda") -> Dict[str, float]:
if not torch.cuda.is_available():
return {"allocated_gib": 0.0, "reserved_gib": 0.0, "max_gib": 0.0}
return {
"allocated_gib": torch.cuda.memory_allocated(device) / (1024 ** 3),
"reserved_gib": torch.cuda.memory_reserved(device) / (1024 ** 3),
"max_gib": torch.cuda.max_memory_allocated(device) / (1024 ** 3),
}
@classmethod
def get_status_string(cls, path: Union[str, Path] = ".", device: str = "cuda") -> str:
ram = cls.get_system_ram_gib()
disk = cls.get_disk_free_gib(path)
vram = cls.get_gpu_vram_gib(device)
return (
f"RAM: {ram['used_gib']:.1f}/{ram['total_gib']:.1f}G ({ram['percent']:.0f}%) | "
f"Disk: {disk['free_gib']:.1f}G free | "
f"VRAM: {vram['allocated_gib']:.1f}G"
)
class CheckpointManager:
"""
Manages robust, atomic checkpoint saving, automatic rotation/pruning of older checkpoints,
lightweight weights export (safetensors + pt), and exact state restoration.
"""
def __init__(
self,
output_dir: Union[str, Path],
keep_last_k: int = 3,
min_free_disk_gb: float = 10.0,
save_weights_only: bool = True,
save_safetensors: bool = True,
):
self.output_dir = Path(output_dir)
self.output_dir.mkdir(parents=True, exist_ok=True)
self.keep_last_k = max(1, keep_last_k)
self.min_free_disk_gb = min_free_disk_gb
self.save_weights_only = save_weights_only
self.save_safetensors = save_safetensors
def check_disk_space(self) -> float:
"""Verifies sufficient disk space remains on the storage volume."""
free_gb = ResourceMonitor.get_disk_free_gib(self.output_dir)["free_gib"]
if free_gb < self.min_free_disk_gb:
logger.warning(
f"WARNING: Low disk space detected on {self.output_dir}: "
f"{free_gb:.2f} GiB remaining (threshold: {self.min_free_disk_gb} GiB)."
)
return free_gb
def save_checkpoint(
self,
step: int,
model: nn.Module,
optimizer: Optional[torch.optim.Optimizer] = None,
scheduler: Optional[Any] = None,
sampler: Optional[Any] = None,
tokens_seen: int = 0,
metrics: Optional[Dict[str, Any]] = None,
is_best: bool = False,
config_dict: Optional[Dict[str, Any]] = None,
) -> Path:
"""
Atomically saves complete training state, exports weights, and prunes old checkpoints.
"""
self.check_disk_space()
# Capture RNG states for bit-for-bit reproducibility
rng_state = {
"python": random.getstate(),
"numpy": np.random.get_state(),
"torch_cpu": torch.get_rng_state(),
"torch_cuda": torch.cuda.get_rng_state() if torch.cuda.is_available() else None,
}
# Handle unwrapped models (e.g. torch.compile or DDP)
raw_model = getattr(model, "_orig_mod", model)
model_state = raw_model.state_dict()
payload = {
"step": step,
"tokens_seen": tokens_seen,
"model_state": model_state,
"optimizer_state": optimizer.state_dict() if optimizer else None,
"scheduler_state": scheduler.state_dict() if scheduler else None,
"sampler_state": sampler.state_dict() if (sampler and hasattr(sampler, "state_dict")) else None,
"rng_state": rng_state,
"metrics": metrics or {},
"config": config_dict or {},
}
# 1. Atomic write of full training checkpoint (via temporary file)
target_path = self.output_dir / f"checkpoint_step_{step:08d}.pt"
tmp_path = self.output_dir / f"checkpoint_step_{step:08d}.pt.tmp"
torch.save(payload, tmp_path)
os.replace(tmp_path, target_path)
# 2. Save standalone lightweight weights (for inference / evaluation without optimizer memory)
if self.save_weights_only:
weights_path = self.output_dir / f"weights_step_{step:08d}.pt"
weights_tmp = self.output_dir / f"weights_step_{step:08d}.pt.tmp"
torch.save({"model_state": model_state, "config": config_dict or {}}, weights_tmp)
os.replace(weights_tmp, weights_path)
# 3. Save Safetensors format
if self.save_safetensors:
st_path = self.output_dir / f"model_step_{step:08d}.safetensors"
cloned_state = {}
for k, v in model_state.items():
if isinstance(v, torch.Tensor):
cloned_state[k] = v.clone().cpu()
save_file(cloned_state, str(st_path))
# 4. Update latest checkpoint pointer
latest_record = self.output_dir / "latest_checkpoint.txt"
with open(latest_record, "w", encoding="utf-8") as f:
f.write(str(target_path.name) + "\n")
# 5. If best model, save dedicated copy
if is_best:
best_target = self.output_dir / "best_model.pt"
torch.save({"model_state": model_state, "step": step, "metrics": metrics}, best_target)
# 6. Prune old checkpoints to maintain bounded disk footprint
self.prune_old_checkpoints()
return target_path
def prune_old_checkpoints(self):
"""Keeps only the last keep_last_k checkpoints and removes older ones."""
ckpts = sorted(self.output_dir.glob("checkpoint_step_*.pt"), key=os.path.getmtime)
if len(ckpts) > self.keep_last_k:
to_delete = ckpts[: -self.keep_last_k]
for ckpt in to_delete:
try:
step_str = ckpt.stem.replace("checkpoint_step_", "")
ckpt.unlink(missing_ok=True)
# Also prune corresponding weights file if present
w_file = self.output_dir / f"weights_step_{step_str}.pt"
w_file.unlink(missing_ok=True)
st_file = self.output_dir / f"model_step_{step_str}.safetensors"
st_file.unlink(missing_ok=True)
except Exception as e:
logger.warning(f"Failed to delete old checkpoint {ckpt}: {e}")
@staticmethod
def find_latest_checkpoint(dir_path: Union[str, Path]) -> Optional[Path]:
"""Finds the most recent checkpoint in directory, using latest_checkpoint.txt or max step."""
p = Path(dir_path)
if not p.exists():
return None
# Check latest_checkpoint.txt first
pointer = p / "latest_checkpoint.txt"
if pointer.exists():
with open(pointer, "r", encoding="utf-8") as f:
name = f.read().strip()
target = p / name
if target.exists():
return target
# Fallback to scanning
ckpts = sorted(p.glob("checkpoint_step_*.pt"))
if ckpts:
return ckpts[-1]
return None
@staticmethod
def load_checkpoint(
checkpoint_path: Union[str, Path],
model: nn.Module,
optimizer: Optional[torch.optim.Optimizer] = None,
scheduler: Optional[Any] = None,
sampler: Optional[Any] = None,
device: str = "cpu",
strict: bool = True,
) -> Dict[str, Any]:
"""
Safely restores model weights, optimizer moments, scheduler, sampler, and RNG states.
"""
ckpt_path = Path(checkpoint_path)
if not ckpt_path.exists():
raise FileNotFoundError(f"Checkpoint not found at: {ckpt_path}")
print(f" Loading checkpoint from: {ckpt_path} (target device: {device})...")
ckpt = torch.load(ckpt_path, map_location=device, weights_only=False)
# Restore model weights
raw_model = getattr(model, "_orig_mod", model)
raw_model.load_state_dict(ckpt["model_state"], strict=strict)
# Restore optimizer
if optimizer and ckpt.get("optimizer_state"):
optimizer.load_state_dict(ckpt["optimizer_state"])
# Restore scheduler
if scheduler and ckpt.get("scheduler_state"):
scheduler.load_state_dict(ckpt["scheduler_state"])
# Restore sampler
if sampler and ckpt.get("sampler_state") and hasattr(sampler, "load_state_dict"):
sampler.load_state_dict(ckpt["sampler_state"])
# Restore RNG states
if "rng_state" in ckpt:
rng = ckpt["rng_state"]
if "python" in rng and rng["python"]:
random.setstate(rng["python"])
if "numpy" in rng and rng["numpy"]:
np.random.set_state(rng["numpy"])
if "torch_cpu" in rng and rng["torch_cpu"] is not None:
cpu_state = rng["torch_cpu"]
if isinstance(cpu_state, torch.Tensor):
cpu_state = cpu_state.cpu().to(torch.uint8)
torch.set_rng_state(cpu_state)
if "torch_cuda" in rng and rng["torch_cuda"] is not None and torch.cuda.is_available():
cuda_state = rng["torch_cuda"]
if isinstance(cuda_state, torch.Tensor):
cuda_state = cuda_state.cpu().to(torch.uint8)
torch.cuda.set_rng_state(cuda_state)
step = ckpt.get("step", 0)
tokens_seen = ckpt.get("tokens_seen", 0)
metrics = ckpt.get("metrics", {})
print(f" Resumed successfully from Step {step:,} ({tokens_seen:,} tokens seen).")
return {
"step": step,
"tokens_seen": tokens_seen,
"metrics": metrics,
"config": ckpt.get("config", {}),
}