Download code/state_io.py from laion/Humaneness-Voice-Small: direct link, hf CLI and curl.
- Browser
- Download file 3.81 kB
-
https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/state_io.py
- Command line
-
hf download hf://laion/Humaneness-Voice-Small/code/state_io.py
-
curl -L -o state_io.py https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/state_io.py
3.81 kB
| """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 | |