#!/usr/bin/env python3 """Portable, lossless OpenWAM checkpoint archive. Python standard library only. Weights are copied, never linked. Import is append-only, content-addressed shared components are immutable, and publication uses a filesystem lock and atomic JSON. """ from __future__ import annotations import argparse import base64 import concurrent.futures import contextlib import datetime import fcntl import hashlib import json import math import os from pathlib import Path, PurePosixPath import re import shutil import struct import subprocess import sys import tempfile import uuid VERSION = 1 BLOCK = 8 * 1024 * 1024 PROCESS_HASH_CACHE = {} DTYPE_BYTES = {'BOOL': 1, 'U8': 1, 'I8': 1, 'F8_E4M3': 1, 'F8_E5M2': 1, 'I16': 2, 'U16': 2, 'F16': 2, 'BF16': 2, 'I32': 4, 'U32': 4, 'F32': 4, 'I64': 8, 'U64': 8, 'F64': 8} def now(): return datetime.datetime.now(datetime.timezone.utc).isoformat() def json_bytes(value): return json.dumps(value, sort_keys=True, separators=(',', ':'), ensure_ascii=False).encode() def atomic_json(path, value): path = Path(path) path.parent.mkdir(parents=True, exist_ok=True) temp = path.with_name(path.name + '.tmp-' + uuid.uuid4().hex) try: with temp.open('xb') as f: f.write(json.dumps(value, indent=2, ensure_ascii=False).encode() + b'\n') f.flush() os.fsync(f.fileno()) os.replace(temp, path) finally: temp.unlink(missing_ok=True) def sha256_file(path): h = hashlib.sha256() with Path(path).open('rb') as f: while chunk := f.read(BLOCK): h.update(chunk) return h.hexdigest() def stable_stat(path): s = Path(path).stat() return {'size': s.st_size, 'mtime_ns': s.st_mtime_ns, 'inode': s.st_ino, 'device': s.st_dev} def checked_path(root, relative): """No absolute paths, traversal, symlinks, or references outside the archive.""" rel = PurePosixPath(relative) if rel.is_absolute() or not rel.parts or any(p in ('..', '.') for p in rel.parts): raise ValueError(f'Invalid archive-relative path: {relative!r}') root = Path(root).resolve() path = root for part in rel.parts: path = path / part if path.is_symlink(): raise ValueError(f'Archive contains a symlink: {path}') if not path.resolve().is_relative_to(root): raise ValueError(f'Archive path escaped root: {relative}') return path @contextlib.contextmanager def locked(path): path = Path(path) path.parent.mkdir(parents=True, exist_ok=True) with path.open('a') as f: fcntl.flock(f, fcntl.LOCK_EX) yield def canonical_run(run): if not re.fullmatch(r'run\d+', run): raise ValueError('run id must be run followed by digits, e.g. run41') return f'run{int(run[3:]):02d}' def read_header(path): with Path(path).open('rb') as f: raw_n = f.read(8) if len(raw_n) != 8: raise ValueError(f'Truncated safetensors: {path}') n = struct.unpack(' 100_000_000 or n < 2: raise ValueError(f'Invalid safetensors header: {path}') raw = f.read(n) if len(raw) != n: raise ValueError(f'Truncated safetensors header: {path}') def pairs_hook(pairs): result = {} for key, value in pairs: if key in result: raise ValueError(f'Duplicate JSON key: {key}') result[key] = value return result header = json.loads(raw, object_pairs_hook=pairs_hook) tensors = {k: v for k, v in header.items() if k != '__metadata__'} end = 0 for key, value in sorted(tensors.items(), key=lambda kv: kv[1]['data_offsets']): a, b = value['data_offsets'] dtype = value['dtype'] shape = value['shape'] if dtype not in DTYPE_BYTES or any(not isinstance(x, int) or x < 0 for x in shape): raise ValueError(f'Unsupported tensor layout: {path}: {key}') if a != end or b - a != math.prod(shape) * DTYPE_BYTES[dtype]: raise ValueError(f'Invalid tensor offsets: {path}: {key}') end = b if len(raw_n) + len(raw) + end != Path(path).stat().st_size: raise ValueError(f'Checkpoint length mismatch: {path}') return raw_n + raw, tensors def shard_layout(keys, tensors): offsets, header, end = {}, {}, 0 for key in keys: src = tensors[key] n = src['data_offsets'][1] - src['data_offsets'][0] offsets[key] = (end, end + n) header[key] = {'dtype': src['dtype'], 'shape': src['shape'], 'data_offsets': [end, end + n]} end += n raw = json_bytes(header) raw += b' ' * (-len(raw) % 8) return struct.pack('