Swift-1.5-5bit-MLX / verify_release.py
ukisai's picture
initial release
e476358
Raw History Blame
6.66 kB
"""Offline, bounded-memory file/index integrity check. Does not execute model code."""
import argparse
import hashlib
import json
import math
from pathlib import Path, PurePosixPath
import re
import struct
def unique(pairs):
result = {}
for key, value in pairs:
if key in result:
raise ValueError('duplicate JSON key')
result[key] = value
return result
def read_json(path):
return json.loads(path.read_text(encoding='utf-8'), object_pairs_hook=unique)
def safe_path(root, name):
if not isinstance(name, str) or not name or '\\' in name:
raise ValueError('unsafe file path')
p = PurePosixPath(name)
if p.is_absolute() or '..' in p.parts or str(p) != name:
raise ValueError('unsafe file path')
target = root.joinpath(*p.parts)
if any(parent.is_symlink() for parent in (target, *target.parents) if parent != root.parent):
raise ValueError('symlink path not allowed')
target.resolve().relative_to(root.resolve())
return target
def file_hashes(path):
size = path.stat().st_size
sha = hashlib.sha256()
blob = hashlib.sha1(f'blob {size}\0'.encode())
with path.open('rb') as stream:
for chunk in iter(lambda: stream.read(8 * 1024**2), b''):
sha.update(chunk)
blob.update(chunk)
return size, sha.hexdigest(), blob.hexdigest()
def tensor_header(path):
size = path.stat().st_size
with path.open('rb') as stream:
prefix = stream.read(8)
if len(prefix) != 8:
raise ValueError('short safetensors prefix')
n = struct.unpack('<Q', prefix)[0]
if n > 16 * 1024**2 or 8+n > size:
raise ValueError('invalid header boundary')
header = json.loads(stream.read(n), object_pairs_hook=unique)
widths = {'BF16': 2, 'F16': 2, 'F32': 4, 'F64': 8, 'U32': 4, 'I32': 4,
'U8': 1, 'I8': 1, 'U16': 2, 'I16': 2, 'U64': 8, 'I64': 8, 'BOOL': 1}
spans, names = [], set()
for name, value in header.items():
if name == '__metadata__':
continue
shape = value['shape']
if not isinstance(shape, list) or any(type(x) is not int or x < 0 for x in shape):
raise ValueError('invalid tensor shape')
start, end = value['data_offsets']
if type(start) is not int or type(end) is not int or not 0 <= start <= end <= size-8-n:
raise ValueError('invalid tensor payload boundary')
if end-start != math.prod(shape)*widths[value['dtype']]:
raise ValueError('dtype/shape byte count mismatch')
spans.append((start, end))
names.add(name)
cursor = 0
for start, end in sorted(spans):
if start != cursor:
raise ValueError('payload gap or overlap')
cursor = end
if cursor != size-8-n:
raise ValueError('unreferenced or truncated payload')
return names
def verify(root, expected_manifest_sha256=None):
root = Path(root).absolute()
errors, checked = [], []
try:
manifest_path = safe_path(root, 'UPLOAD_MANIFEST.json')
_, manifest_sha, _ = file_hashes(manifest_path)
if expected_manifest_sha256 and expected_manifest_sha256 != manifest_sha:
raise ValueError('trusted manifest hash mismatch')
manifest = read_json(manifest_path)
entries = manifest['files']
names = [e['path'] for e in entries]
if len(names) != len(set(names)):
raise ValueError('duplicate manifest file')
if any(type(e['bytes']) is not int or e['bytes'] < 0 for e in entries):
raise ValueError('invalid manifest size')
if len(entries) != manifest['file_count'] or sum(e['bytes'] for e in entries) != manifest['total_bytes']:
raise ValueError('manifest aggregates mismatch')
for e in entries:
if e['path'] == 'UPLOAD_MANIFEST.json' or not re.fullmatch('[a-f0-9]{64}', e['sha256']):
raise ValueError('invalid manifest entry')
p = safe_path(root, e['path'])
if not p.is_file():
errors.append({'file': e['path'], 'reason': 'missing file'})
continue
size, sha, blob = file_hashes(p)
if size != e['bytes'] or sha != e['sha256'] or (e.get('git_blob_sha1') and blob != e['git_blob_sha1']):
errors.append({'file': e['path'], 'reason': 'size or full-file hash mismatch'})
continue
if p.suffix == '.json':
read_json(p)
checked.append(e['path'])
index_name = 'model.safetensors.index.json'
if index_name not in names:
errors.append({'file': index_name, 'reason': 'required index not in manifest'})
elif index_name in checked:
index = read_json(root/index_name)['weight_map']
if not isinstance(index, dict) or not index:
raise ValueError('invalid weight map')
shards = set(index.values())
expected_shards = {n for n in names if n.endswith('.safetensors')}
if shards != expected_shards:
errors.append({'reason': 'index/manifest shard set mismatch'})
actual = {}
for shard in sorted(shards):
safe_path(root, shard)
if shard not in checked:
errors.append({'file': shard, 'reason': 'referenced shard missing or not hash-verified'})
continue
for tensor in tensor_header(root/shard):
if tensor in actual:
raise ValueError('duplicate tensor across shards')
actual[tensor] = shard
if actual != index:
errors.append({'reason': 'actual tensor inventory differs from weight map'})
except (OSError, ValueError, TypeError, KeyError, AttributeError) as exc:
errors.append({'reason': 'invalid or incomplete package', 'exception_type': type(exc).__name__})
return {'status': 'FAIL' if errors else 'PASS', 'errors': errors,
'fully_hashed_files': len(checked), 'manifest_pinned': bool(expected_manifest_sha256),
'scope': 'File integrity and tensor index only; not generation, finite-value or quality validation'}
if __name__ == '__main__':
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument('snapshot', type=Path)
parser.add_argument('--manifest-sha256')
args = parser.parse_args()
result = verify(args.snapshot, args.manifest_sha256)
print(json.dumps(result, indent=2))
raise SystemExit(0 if result['status'] == 'PASS' else 1)