OpenWAM-3sys / tools /checkpoint_archive.py
qiuly's picture
Freeze OpenWAM-3sys archive with final Run41-47 policies and Run42 pretraining
a658404 verified
Raw History Blame Contribute Delete
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
@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('<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()