""" 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", {}), }