"""Atomic full-state DDP checkpoints with per-rank RNG and pinned contracts.""" from datetime import datetime import json import os from pathlib import Path import random import numpy as np import torch import torch.distributed as dist from continuation_control import sha def ranks(): return (dist.get_rank(), dist.get_world_size()) if dist.is_initialized() else (0, 1) def barrier(): if dist.is_initialized(): dist.barrier() def unwrap(model): return model.module if hasattr(model, 'module') else model def save(root, model, optimizer, scheduler, step, contract, export=False): rank, world = ranks() root = Path(root) rng = {'torch': torch.get_rng_state(), 'cuda': torch.cuda.get_rng_state(), 'numpy': np.random.get_state(), 'python': random.getstate()} all_rng = [None] * world if rank == 0 else None if world > 1: dist.gather_object(rng, all_rng, dst=0) else: all_rng = [rng] directory = root / f'step-{step:08d}' if rank == 0: temporary = root / f'.step-{step:08d}-{os.environ.get("SLURM_JOB_ID", "local")}.partial' temporary.mkdir(parents=True, exist_ok=True) torch.save({'model': unwrap(model).state_dict(), 'optimizer': optimizer.state_dict(), 'scheduler': scheduler.state_dict(), 'rng': all_rng, 'step': step, 'world': world, 'contract': contract}, temporary / 'training_state.pt') metadata = {'step': step, 'world': world, 'contract': contract, 'created_at': datetime.now().astimezone().isoformat(), 'state_bytes': (temporary / 'training_state.pt').stat().st_size, 'state_sha256': sha(temporary / 'training_state.pt')} if export: base = unwrap(model).base cache, state = {}, {} for name, value in base.state_dict().items(): key = (value.data_ptr(), tuple(value.shape), value.dtype) if key not in cache: cache[key] = value.detach().to(dtype=torch.bfloat16, device='cpu') state[name] = cache[key] torch.save(state, temporary / 'model_bf16.pt') metadata['export_bytes'] = (temporary / 'model_bf16.pt').stat().st_size metadata['export_sha256'] = sha(temporary / 'model_bf16.pt') (temporary / 'COMPLETE.json').write_text(json.dumps(metadata, indent=2) + '\n') if directory.exists(): raise RuntimeError(f'Refusing to overwrite existing checkpoint {directory}') temporary.replace(directory) pointer = root / 'current.json.tmp' pointer.write_text(json.dumps({'directory': directory.name, 'step': step}) + '\n') pointer.replace(root / 'current.json') barrier() return directory def load(directory, model, optimizer, scheduler, contract): directory = Path(directory) rank, world = ranks() metadata = json.loads((directory / 'COMPLETE.json').read_text()) assert metadata['contract'] == contract and metadata['world'] == world assert (directory / 'training_state.pt').stat().st_size == metadata['state_bytes'] state = torch.load(directory / 'training_state.pt', map_location='cpu', weights_only=False) assert state['contract'] == contract and state['world'] == world and len(state['rng']) == world unwrap(model).load_state_dict(state['model'], strict=True) optimizer.load_state_dict(state['optimizer']); scheduler.load_state_dict(state['scheduler']) rng = state['rng'][rank] torch.set_rng_state(rng['torch']); torch.cuda.set_rng_state(rng['cuda']) np.random.set_state(rng['numpy']); random.setstate(rng['python']) step = int(state['step']) assert step==metadata['step'] del state barrier() return step