Download tools/checkpoint_archive.py from qiuly/OpenWAM-3sys: direct link, hf CLI and curl.
- Browser
- Download file 35.5 kB
-
https://huggingface.co/qiuly/OpenWAM-3sys/resolve/main/tools/checkpoint_archive.py
- Command line
-
hf download hf://qiuly/OpenWAM-3sys/tools/checkpoint_archive.py
-
curl -L -o checkpoint_archive.py https://huggingface.co/qiuly/OpenWAM-3sys/resolve/main/tools/checkpoint_archive.py
35.5 kB
| #!/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 | |
| 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('<Q', raw_n)[0] | |
| if n > 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('<Q', len(raw)) + raw, offsets | |
| def component(key, share_decoder): | |
| for prefix, name in [ | |
| ('video_backbone.reason1.', 'reason1'), | |
| ('video_backbone.vae.', 'wan-vae'), | |
| ('video_backbone.video_encoder._m.', 'visual-encoder'), | |
| ('video_backbone.video_encoder._svae.', 'svae'), | |
| ('video_backbone.video_encoder.teacher_model.', 'videorae-teacher'), | |
| ('video_backbone.video_encoder.latent_compressor.', 'videorae-compressor'), | |
| ('vlm_expert.vlm.visual.', 'qwen-vision'), | |
| ('vlm_expert.vlm.language_model.embed_tokens.', 'qwen-embeddings'), | |
| ]: | |
| if key.startswith(prefix): | |
| return name | |
| if share_decoder and key.startswith('vlm_expert.vlm.'): | |
| return 'qwen-frozen-decoder' | |
| return 'private' | |
| class Archive: | |
| def __init__(self, root): | |
| self.root = Path(root).resolve() | |
| self.hash_cache = PROCESS_HASH_CACHE | |
| def init(self): | |
| self.root.mkdir(parents=True, exist_ok=True) | |
| with locked(self.root / '.archive.lock'): | |
| path = self.root / 'catalog.json' | |
| if not path.exists(): | |
| atomic_json(path, {'format': 'openwam-checkpoint-archive', 'schema_version': VERSION, | |
| 'created_at': now(), 'updated_at': now(), 'entries': {}}) | |
| self.catalog() | |
| def catalog(self): | |
| c = json.loads((self.root / 'catalog.json').read_text()) | |
| if c.get('schema_version') != VERSION or c.get('format') != 'openwam-checkpoint-archive': | |
| raise ValueError('Unsupported catalog version/format') | |
| return c | |
| def check_file(self, record, full=True): | |
| path = checked_path(self.root, record['path']) | |
| if not path.is_file() or path.stat().st_size != record['bytes']: | |
| raise ValueError(f'Missing or truncated archive file: {path}') | |
| if full: | |
| stat = path.stat() | |
| cache_key = (str(path), stat.st_size, stat.st_mtime_ns, stat.st_ctime_ns) | |
| digest = self.hash_cache.get(cache_key) | |
| if digest is None: | |
| digest = sha256_file(path) | |
| self.hash_cache[cache_key] = digest | |
| if digest != record['sha256']: | |
| raise ValueError(f'Checksum mismatch: {path}') | |
| return path | |
| def resolve(self, run, stage='sft'): | |
| entries = self.catalog()['entries'] | |
| key = run if '/' in run else f'{stage}/{canonical_run(run)}' | |
| if key not in entries: | |
| raise KeyError(f'No archive entry {key}; use list to see available runs/stages') | |
| return entries[key] | |
| def index(self, entry): | |
| path = self.check_file(entry['index']) | |
| index = json.loads(path.read_text()) | |
| if index.get('schema_version') != VERSION or index['id'] != entry['id']: | |
| raise ValueError('Index/catalog identity mismatch') | |
| return index | |
| def _publish(self, stage_dir, destination, index, entry): | |
| final = checked_path(self.root, destination) | |
| with locked(self.root / '.archive.lock'): | |
| catalog = self.catalog() | |
| for dep in entry.get('dependencies', {}).values(): | |
| if dep not in catalog['entries']: | |
| raise ValueError(f'Missing catalog dependency: {dep}') | |
| if entry['id'] in catalog['entries'] or final.exists(): | |
| raise FileExistsError(f'Archive entry already exists: {entry["id"]}; immutable entries cannot be overwritten') | |
| atomic_json(stage_dir / 'checkpoint.index.json', index) | |
| entry['index'] = {'path': destination + '/checkpoint.index.json', | |
| 'bytes': (stage_dir / 'checkpoint.index.json').stat().st_size, | |
| 'sha256': sha256_file(stage_dir / 'checkpoint.index.json')} | |
| final.parent.mkdir(parents=True, exist_ok=True) | |
| os.rename(stage_dir, final) | |
| catalog['entries'][entry['id']] = entry | |
| catalog['updated_at'] = now() | |
| try: | |
| atomic_json(self.root / 'catalog.json', catalog) | |
| except BaseException: | |
| os.rename(final, stage_dir) | |
| raise | |
| return entry | |
| def _copy_file(self, source, dest, relative): | |
| source, dest = Path(source), Path(dest) | |
| before = stable_stat(source) | |
| dest.parent.mkdir(parents=True, exist_ok=True) | |
| h = hashlib.sha256() | |
| with source.open('rb') as src, dest.open('xb') as out: | |
| while chunk := src.read(BLOCK): | |
| out.write(chunk) | |
| h.update(chunk) | |
| out.flush() | |
| os.fsync(out.fileno()) | |
| if stable_stat(source) != before or sha256_file(dest) != h.hexdigest(): | |
| raise ValueError(f'Copy verification failed: {source}') | |
| if dest.stat().st_ino == source.stat().st_ino and dest.stat().st_dev == source.stat().st_dev: | |
| raise ValueError('Copy unexpectedly shares an inode with its source') | |
| return {'path': relative, 'bytes': before['size'], 'sha256': h.hexdigest(), | |
| 'source': str(source), 'source_stat': before} | |
| def _asset_sources(self, source_dir): | |
| for root, dirs, files in os.walk(source_dir, followlinks=False): | |
| dirs[:] = [d for d in dirs if not d.startswith(('accel_state_', '.')) and d not in ('wandb', 'logs', 'eval', 'features')] | |
| for d in dirs: | |
| if (Path(root) / d).is_symlink(): | |
| raise ValueError(f'Unexpected symlinked metadata directory: {Path(root) / d}') | |
| for name in sorted(files): | |
| if name.startswith(('checkpoint_step_', '.')) or name.endswith(('.log', '.mp4')): | |
| continue | |
| src = Path(root) / name | |
| rel = src.relative_to(source_dir).as_posix() | |
| yield src, rel | |
| def _assets(self, source_dir, stage_dir, destination): | |
| assets = [] | |
| for src, rel in self._asset_sources(source_dir): | |
| record = self._copy_file(src, stage_dir / 'private/assets' / rel, | |
| destination + '/private/assets/' + rel) | |
| record['restore_path'] = rel | |
| assets.append(record) | |
| return assets | |
| def _install_shared(self, source, prefix, tensors, keys, hashes, label, stage_dir): | |
| identity = hashlib.sha256(json_bytes([ | |
| [k, tensors[k]['dtype'], tensors[k]['shape'], hashes[k]] for k in sorted(keys) | |
| ])).hexdigest() | |
| base = ('svae/_components' if label == 'svae' else 'shared/' + label) + '/' + identity | |
| relative = base + '/model.safetensors' | |
| path = checked_path(self.root, relative) | |
| meta_path = checked_path(self.root, base + '/component.json') | |
| with locked(self.root / '.archive.lock'): | |
| if path.exists() and meta_path.exists(): | |
| record = json.loads(meta_path.read_text()) | |
| if record['identity'] != identity or record['path'] != relative: | |
| raise ValueError('Invalid shared component metadata') | |
| self.check_file(record) | |
| return record | |
| header, _ = shard_layout(keys, tensors) | |
| temp = stage_dir / (label + '-' + uuid.uuid4().hex + '.safetensors') | |
| h = hashlib.sha256(header) | |
| with source.open('rb') as src, temp.open('xb') as out: | |
| out.write(header) | |
| for key in keys: | |
| a, b = tensors[key]['data_offsets'] | |
| src.seek(len(prefix) + a) | |
| remaining = b - a | |
| tensor_hash = hashlib.sha256() | |
| while remaining: | |
| data = src.read(min(BLOCK, remaining)) | |
| if not data: | |
| raise ValueError('Source truncated during shared-component copy') | |
| out.write(data) | |
| h.update(data) | |
| tensor_hash.update(data) | |
| remaining -= len(data) | |
| if tensor_hash.hexdigest() != hashes[key]: | |
| raise ValueError(f'Source changed while importing: {source}: {key}') | |
| out.flush() | |
| os.fsync(out.fileno()) | |
| record = {'path': relative, 'bytes': temp.stat().st_size, 'sha256': h.hexdigest(), | |
| 'identity': identity, 'component': label, 'tensor_count': len(keys)} | |
| if sha256_file(temp) != record['sha256']: | |
| raise ValueError('Shared component readback verification failed') | |
| with locked(self.root / '.archive.lock'): | |
| if path.exists(): | |
| if not meta_path.exists(): | |
| raise ValueError(f'Incomplete component publication: {path}') | |
| self.check_file(record) | |
| temp.unlink() | |
| else: | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| os.rename(temp, path) | |
| atomic_json(meta_path, record) | |
| return record | |
| def add_policy(self, *, run, name, stage, source, share_decoder=False, dependencies=None): | |
| self.init() | |
| run = canonical_run(run) | |
| if stage not in ('sft', 'pretrain') or not re.fullmatch(r'[A-Za-z0-9_.-]+', name): | |
| raise ValueError('Invalid stage/name') | |
| if not re.match(r'^run0*' + str(int(run[3:])) + r'(?:_|-|$)', name): | |
| raise ValueError('Run name must start with the requested run number') | |
| source = Path(source).resolve() | |
| match = re.fullmatch(r'checkpoint_step_(\d+)\.safetensors', source.name) | |
| if not match or not source.is_file(): | |
| raise ValueError('Pass the explicit completed checkpoint_step_N.safetensors file') | |
| for required in ['config.yaml'] + (['normalization_stats.npy'] if stage == 'sft' else []): | |
| if not (source.parent / required).is_file(): | |
| raise ValueError(f'Missing deploy asset: {source.parent / required}') | |
| key = stage + '/' + run | |
| existing = self.catalog()['entries'].get(key) | |
| if existing: | |
| if existing['source']['sha256'] != sha256_file(source) or existing.get('dependencies', {}) != (dependencies or {}): | |
| raise FileExistsError(f'Different checkpoint/dependencies already registered for {key}') | |
| prior_assets = self.index(existing)['assets'] | |
| if {rel for _, rel in self._asset_sources(source.parent)} != {a['restore_path'] for a in prior_assets}: | |
| raise FileExistsError(f'Deploy metadata file set changed for existing entry {key}') | |
| for asset in prior_assets: | |
| current = checked_path(source.parent, asset['restore_path']) | |
| if not current.is_file() or sha256_file(current) != asset['sha256']: | |
| raise FileExistsError(f'Deploy metadata changed for existing entry {key}: {current}') | |
| self.verify_entry(existing) | |
| return existing | |
| name = re.sub(r'^run\d+', run, name) | |
| destination = stage + '/' + name | |
| (self.root / '.staging').mkdir(exist_ok=True) | |
| temp = Path(tempfile.mkdtemp(prefix=run + '-', dir=self.root / '.staging')) | |
| try: | |
| before = stable_stat(source) | |
| prefix, tensors = read_header(source) | |
| keys = sorted(tensors, key=lambda k: tensors[k]['data_offsets']) | |
| groups = {} | |
| for k in keys: | |
| groups.setdefault(component(k, share_decoder), []).append(k) | |
| private_keys = groups.pop('private', []) | |
| private_header, _ = shard_layout(private_keys, tensors) | |
| private = temp / 'private/weights.safetensors' | |
| private.parent.mkdir() | |
| private_hash = hashlib.sha256(private_header) | |
| original_hash = hashlib.sha256(prefix) | |
| hashes = {} | |
| with source.open('rb') as src, private.open('xb') as out: | |
| out.write(private_header) | |
| src.seek(len(prefix)) | |
| for k in keys: | |
| a, b = tensors[k]['data_offsets'] | |
| left = b - a | |
| h = hashlib.sha256() | |
| own = component(k, share_decoder) == 'private' | |
| while left: | |
| chunk = src.read(min(BLOCK, left)) | |
| if not chunk: | |
| raise ValueError(f'Truncated source: {source}') | |
| original_hash.update(chunk) | |
| h.update(chunk) | |
| if own: | |
| out.write(chunk) | |
| private_hash.update(chunk) | |
| left -= len(chunk) | |
| hashes[k] = h.hexdigest() | |
| out.flush() | |
| os.fsync(out.fileno()) | |
| if stable_stat(source) != before: | |
| raise ValueError(f'Source changed during import: {source}') | |
| private_record = {'path': destination + '/private/weights.safetensors', | |
| 'bytes': private.stat().st_size, 'sha256': private_hash.hexdigest()} | |
| if sha256_file(private) != private_record['sha256']: | |
| raise ValueError('Private weights readback verification failed') | |
| records = [private_record] | |
| mapping = {k: private_record['path'] for k in private_keys} | |
| for label, group_keys in groups.items(): | |
| record = self._install_shared(source, prefix, tensors, group_keys, hashes, label, temp) | |
| records.append(record) | |
| mapping.update({k: record['path'] for k in group_keys}) | |
| assets = self._assets(source.parent, temp, destination) | |
| if stable_stat(source) != before: | |
| raise ValueError('Source changed while metadata was copied') | |
| provenance = {'path': str(source), 'bytes': before['size'], 'sha256': original_hash.hexdigest(), | |
| 'stat': before, 'filename': source.name} | |
| index = {'schema_version': VERSION, 'kind': 'policy', 'id': key, | |
| 'original': provenance, 'original_header_base64': base64.b64encode(prefix).decode(), | |
| 'weight_map': mapping, 'files': records, 'assets': assets, 'tensor_sha256': hashes} | |
| entry = {'id': key, 'run_id': run, 'run_name': name, 'stage': stage, 'step': int(match[1]), | |
| 'kind': 'policy', 'directory': destination, 'source': provenance, | |
| 'dependencies': dependencies or {}, 'created_at': now(), 'tensor_count': len(keys)} | |
| self._publish(temp, destination, index, entry) | |
| print(json.dumps({'imported': key, 'step': entry['step'], 'original_bytes': before['size'], | |
| 'private_bytes': private_record['bytes']}), flush=True, file=sys.stderr) | |
| return entry | |
| finally: | |
| if temp.exists(): | |
| shutil.rmtree(temp) | |
| def add_model(self, *, key, name, source, files, metadata=None, extra_assets=None): | |
| self.init() | |
| if not re.fullmatch(r'(?:svae|pretrain)/[A-Za-z0-9_.-]+', key) or not re.fullmatch(r'[A-Za-z0-9_.-]+', name): | |
| raise ValueError('Invalid model key/name') | |
| stage, run = key.split('/') | |
| source = Path(source).resolve() | |
| destination = stage + '/' + name | |
| self.root.joinpath('.staging').mkdir(exist_ok=True) | |
| temp = Path(tempfile.mkdtemp(prefix='model-', dir=self.root / '.staging')) | |
| try: | |
| assets = [] | |
| sources = [(checked_path(source, rel), rel) for rel in sorted(files)] | |
| sources.extend((Path(src), rel) for rel, src in (extra_assets or {}).items()) | |
| if len({rel for _, rel in sources}) != len(sources): | |
| raise ValueError('Duplicate destination in model assets') | |
| for src, rel in sources: | |
| if not src.is_file(): | |
| raise FileNotFoundError(src) | |
| r = self._copy_file(src, checked_path(temp / 'private', rel), destination + '/private/' + rel) | |
| r['restore_path'] = rel | |
| assets.append(r) | |
| existing = self.catalog()['entries'].get(key) | |
| if existing: | |
| prior = self.index(existing) | |
| if {r['restore_path']: r['sha256'] for r in prior['assets']} != {r['restore_path']: r['sha256'] for r in assets}: | |
| raise FileExistsError(f'Different model already registered: {key}') | |
| self.verify_entry(existing) | |
| return existing | |
| index = {'schema_version': VERSION, 'kind': 'model', 'id': key, 'assets': assets, 'files': [], 'metadata': metadata or {}} | |
| entry = {'id': key, 'run_id': run, 'run_name': name, 'stage': stage, 'kind': 'model', | |
| 'directory': destination, 'created_at': now(), 'dependencies': {}, 'metadata': metadata or {}} | |
| return self._publish(temp, destination, index, entry) | |
| finally: | |
| if temp.exists(): | |
| shutil.rmtree(temp) | |
| def verify_entry(self, entry, reconstruct=False): | |
| index = self.index(entry) | |
| for record in index['files'] + index['assets']: | |
| self.check_file(record) | |
| for dependency in entry.get('dependencies', {}).values(): | |
| self.resolve(dependency) | |
| if index['kind'] == 'policy': | |
| self._validate_mapping(index) | |
| if reconstruct: | |
| self._reconstruct(index) | |
| return {'id': entry['id'], 'files': len(index['files']) + len(index['assets']), | |
| 'original_byte_equivalence': bool(reconstruct and index['kind'] == 'policy')} | |
| def _validate_mapping(self, index): | |
| prefix = base64.b64decode(index['original_header_base64'], validate=True) | |
| n = struct.unpack('<Q', prefix[:8])[0] | |
| if len(prefix) != n + 8: | |
| raise ValueError('Invalid original checkpoint header length') | |
| original = {k: v for k, v in json.loads(prefix[8:]).items() if k != '__metadata__'} | |
| if set(original) != set(index['weight_map']) or set(original) != set(index['tensor_sha256']): | |
| raise ValueError('Index must cover exactly every original tensor') | |
| files = {f['path']: f for f in index['files']} | |
| if len(files) != len(index['files']) or set(index['weight_map'].values()) - files.keys(): | |
| raise ValueError('Duplicate or undeclared shard reference') | |
| shards = {} | |
| for relative in files: | |
| p = self.check_file(files[relative], full=False) | |
| h, ts = read_header(p) | |
| assigned = {k for k, path in index['weight_map'].items() if path == relative} | |
| if set(ts) != assigned: | |
| raise ValueError('Shard/index keys differ; cannot silently load extra or missing parameters') | |
| for k in assigned: | |
| if any(ts[k][f] != original[k][f] for f in ['dtype', 'shape']): | |
| raise ValueError(f'Tensor schema mismatch: {k}') | |
| shards[relative] = (p, len(h), ts) | |
| return prefix, original, shards | |
| def _reconstruct(self, index, output=None): | |
| prefix, tensors, shards = self._validate_mapping(index) | |
| h = hashlib.sha256(prefix) | |
| if output: | |
| output.write(prefix) | |
| handles = {} | |
| try: | |
| for relative, (path, start, _) in shards.items(): | |
| handles[relative] = path.open('rb') | |
| written = len(prefix) | |
| for key in sorted(tensors, key=lambda k: tensors[k]['data_offsets']): | |
| rel = index['weight_map'][key] | |
| _, offset, table = shards[rel] | |
| a, b = table[key]['data_offsets'] | |
| f = handles[rel] | |
| f.seek(offset + a) | |
| left = b - a | |
| tensor_hash = hashlib.sha256() | |
| while left: | |
| chunk = f.read(min(BLOCK, left)) | |
| if not chunk: | |
| raise ValueError('Truncated archive shard') | |
| tensor_hash.update(chunk) | |
| h.update(chunk) | |
| if output: | |
| output.write(chunk) | |
| written += len(chunk) | |
| left -= len(chunk) | |
| if tensor_hash.hexdigest() != index['tensor_sha256'][key]: | |
| raise ValueError(f'Tensor byte mismatch: {key}') | |
| if h.hexdigest() != index['original']['sha256'] or written != index['original']['bytes']: | |
| raise ValueError('Reconstruction does not match the original checkpoint byte-for-byte') | |
| finally: | |
| for f in handles.values(): | |
| f.close() | |
| def materialize(self, run, output, stage='sft'): | |
| entry = self.resolve(run, stage) | |
| index = self.index(entry) | |
| output = Path(output).resolve() | |
| if output.is_relative_to(self.root) or self.root.is_relative_to(output): | |
| raise ValueError('Materialization must be outside the archive, not an ancestor of it') | |
| output.parent.mkdir(parents=True, exist_ok=True) | |
| with locked(output.parent / ('.' + output.name + '.lock')): | |
| marker = output / '.archive-materialized.json' | |
| if output.exists(): | |
| if not marker.is_file() or json.loads(marker.read_text()).get('index_sha256') != entry['index']['sha256']: | |
| raise FileExistsError(f'Output exists and is not a matching verified archive export: {output}') | |
| proof = json.loads(marker.read_text()) | |
| expected = [{'path': a['restore_path'], 'bytes': a['bytes'], 'sha256': a['sha256']} for a in index['assets']] | |
| if index['kind'] == 'policy': | |
| expected.append({'path': index['original']['filename'], 'bytes': index['original']['bytes'], 'sha256': index['original']['sha256']}) | |
| if sorted(proof.get('files', []), key=lambda r: r['path']) != sorted(expected, key=lambda r: r['path']): | |
| raise ValueError(f'Materialization marker does not match the archive index: {marker}') | |
| for record in expected: | |
| p = checked_path(output, record['path']) | |
| if not p.is_file() or p.stat().st_size != record['bytes'] or sha256_file(p) != record['sha256']: | |
| raise ValueError(f'Modified materialized file: {p}; remove the cache explicitly and recreate it') | |
| return output | |
| # Stage next to destination so final publication is an atomic rename. | |
| temp = Path(tempfile.mkdtemp(prefix='.' + output.name + '-', dir=output.parent)) | |
| try: | |
| proof = [] | |
| for r in index['files'] + index['assets']: | |
| self.check_file(r) | |
| for record in index['assets']: | |
| src = checked_path(self.root, record['path']) | |
| rel = record['restore_path'] | |
| dest = checked_path(temp, rel) | |
| copied = self._copy_file(src, dest, rel) | |
| proof.append({k: copied[k] for k in ['path', 'bytes', 'sha256']}) | |
| if index['kind'] == 'policy': | |
| dest = checked_path(temp, index['original']['filename']) | |
| with dest.open('xb') as f: | |
| self._reconstruct(index, f) | |
| f.flush() | |
| os.fsync(f.fileno()) | |
| if sha256_file(dest) != index['original']['sha256']: | |
| raise ValueError('Materialized checkpoint readback verification failed') | |
| proof.append({'path': dest.name, 'bytes': dest.stat().st_size, 'sha256': index['original']['sha256']}) | |
| atomic_json(temp / '.archive-materialized.json', {'id': entry['id'], 'index_sha256': entry['index']['sha256'], 'files': proof}) | |
| os.rename(temp, output) | |
| finally: | |
| if temp.exists(): | |
| shutil.rmtree(temp) | |
| return output | |
| def verify(self, reconstruct=False, report=None, run=None, stage='sft', workers=1): | |
| initial_catalog_hash = sha256_file(self.root / 'catalog.json') | |
| catalog = self.catalog() | |
| results = [] | |
| selected = set(catalog['entries']) | |
| if run: | |
| selected = {self.resolve(run, stage)['id']} | |
| pending = list(selected) | |
| while pending: | |
| entry = catalog['entries'][pending.pop()] | |
| for dep in entry.get('dependencies', {}).values(): | |
| if dep not in selected: | |
| self.resolve(dep) | |
| selected.add(dep) | |
| pending.append(dep) | |
| work = [] | |
| for key, entry in sorted(catalog['entries'].items()): | |
| if key != entry['id']: | |
| raise ValueError('Catalog key/entry identity mismatch') | |
| if key not in selected: | |
| continue | |
| work.append(entry) | |
| def verify_one(entry): | |
| result = self.verify_entry(entry, reconstruct=reconstruct) | |
| print(json.dumps(result), flush=True, file=sys.stderr) | |
| return result | |
| if workers < 1: | |
| raise ValueError('workers must be positive') | |
| with concurrent.futures.ThreadPoolExecutor(max_workers=workers) as pool: | |
| results = list(pool.map(verify_one, work)) | |
| for root, dirs, files in os.walk(self.root): | |
| dirs[:] = [d for d in dirs if d != '.staging'] | |
| for name in dirs + files: | |
| p = Path(root) / name | |
| if p.is_symlink() or (p.is_file() and p.stat().st_nlink != 1): | |
| raise ValueError(f'Archive files must be independent real files: {p}') | |
| if name == 'checkpoint.index.json' and p.relative_to(self.root).as_posix() not in { | |
| e['index']['path'] for e in catalog['entries'].values() | |
| }: | |
| raise ValueError(f'Unregistered checkpoint index (interrupted publication): {p}') | |
| if sha256_file(self.root / 'catalog.json') != initial_catalog_hash: | |
| raise ValueError('Catalog changed during verification; verify the new snapshot again') | |
| result = {'status': 'verified', 'schema_version': VERSION, 'time': now(), | |
| 'catalog_sha256': initial_catalog_hash, 'reconstructed': reconstruct, | |
| 'scope': 'all' if run is None else self.resolve(run, stage)['id'], | |
| 'entries': results} | |
| if report: | |
| atomic_json(report, result) | |
| return result | |
| def cli(): | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument('--root', type=Path, default=Path(__file__).resolve().parent.parent) | |
| sub = parser.add_subparsers(dest='cmd', required=True) | |
| sub.add_parser('init') | |
| sub.add_parser('list') | |
| p = sub.add_parser('resolve'); p.add_argument('run'); p.add_argument('--stage', default='sft') | |
| p = sub.add_parser('add-policy') | |
| p.add_argument('--run', required=True); p.add_argument('--name', required=True) | |
| p.add_argument('--stage', choices=['sft', 'pretrain'], required=True) | |
| p.add_argument('--checkpoint', type=Path, required=True) | |
| p.add_argument('--share-frozen-decoder', action='store_true') | |
| p.add_argument('--svae'); p.add_argument('--pretrain') | |
| p = sub.add_parser('add-model') | |
| p.add_argument('--id', required=True); p.add_argument('--name', required=True) | |
| p.add_argument('--source-dir', required=True); p.add_argument('--files', nargs='+', required=True) | |
| p.add_argument('--extra-file', action='append', default=[], metavar='RELATIVE=SOURCE') | |
| p = sub.add_parser('verify'); p.add_argument('--reconstruct', action='store_true'); p.add_argument('--report', type=Path) | |
| p.add_argument('--run'); p.add_argument('--stage', default='sft') | |
| p.add_argument('--workers', type=int, default=1) | |
| p = sub.add_parser('materialize'); p.add_argument('run'); p.add_argument('--stage', default='sft'); p.add_argument('--output', required=True) | |
| p = sub.add_parser('eval'); p.add_argument('run'); p.add_argument('--repo', default=os.environ.get('OPENWAM_REPO', '/mnt/data/limingleyang/OpenWAM-3sys')) | |
| p.add_argument('--cache-root', default=os.environ.get('OPENWAM_CHECKPOINT_EVAL_CACHE')) | |
| p.add_argument('--dry-run', action='store_true') | |
| args = parser.parse_args() | |
| archive = Archive(args.root) | |
| if args.cmd == 'init': | |
| archive.init(); return | |
| if args.cmd == 'list': | |
| for key, entry in sorted(archive.catalog()['entries'].items()): | |
| print(key, entry.get('step', '-'), entry['directory'], sep='\t') | |
| elif args.cmd == 'resolve': | |
| print(json.dumps(archive.resolve(args.run, args.stage), indent=2)) | |
| elif args.cmd == 'add-policy': | |
| deps = {k: v for k, v in [('svae', args.svae), ('pretrain', args.pretrain)] if v} | |
| result = archive.add_policy(run=args.run, name=args.name, stage=args.stage, source=args.checkpoint, | |
| share_decoder=args.share_frozen_decoder, dependencies=deps) | |
| print(json.dumps(result, indent=2)) | |
| elif args.cmd == 'add-model': | |
| extras = dict(item.split('=', 1) for item in args.extra_file) | |
| print(json.dumps(archive.add_model(key=args.id, name=args.name, source=args.source_dir, files=args.files, extra_assets=extras), indent=2)) | |
| elif args.cmd == 'verify': | |
| result = archive.verify(args.reconstruct, args.report, args.run, args.stage, args.workers) | |
| print(json.dumps({'status': result['status'], 'entries': len(result['entries']), 'reconstructed': args.reconstruct})) | |
| elif args.cmd == 'materialize': | |
| print(archive.materialize(args.run, args.output, args.stage)) | |
| elif args.cmd == 'eval': | |
| entry = archive.resolve(args.run) | |
| if entry['kind'] != 'policy' or entry['stage'] != 'sft': | |
| raise ValueError('LIBERO evaluation requires an SFT policy') | |
| repo = Path(args.repo).resolve() | |
| script = repo / 'scripts/libero/eval.sh' | |
| if not script.is_file(): | |
| raise FileNotFoundError(script) | |
| cache = Path(args.cache_root) if args.cache_root else archive.root.parent / 'checkpoint_eval_cache' | |
| output = cache / entry['run_id'] / entry['index']['sha256'][:16] | |
| materialized = archive.materialize(args.run, output) | |
| env = os.environ.copy() | |
| env.setdefault('EVAL_RUN_NAME', entry['run_name'] + '_step' + str(entry['step'])) | |
| if args.dry_run: | |
| env['EVAL_DRY_RUN'] = '1' | |
| raise SystemExit(subprocess.call(['bash', str(script), str(materialized), entry['source']['filename']], cwd=repo, env=env)) | |
| if __name__ == '__main__': | |
| cli() | |