← back to Local Model Leaderboard Watch
verify-model-weights.py
108 lines
#!/usr/bin/env python3
"""Offline structural preflight: a config directory is not a downloaded model.
Read safetensors headers, never model payloads. This is completeness validation,
not publisher authentication or a guarantee that the inference engine supports it.
"""
import json
import math
from pathlib import Path
import struct
import sys
# Official safetensors Dtype::bitsize; unknown future types fail closed.
# https://github.com/safetensors/safetensors/blob/main/safetensors/src/tensor.rs
DTYPE_BITS = {
'BOOL': 8, 'U8': 8, 'I8': 8, 'F8_E5M2': 8, 'F8_E4M3': 8,
'F8_E8M0': 8, 'F8_E4M3FNUZ': 8, 'F8_E5M2FNUZ': 8,
'I16': 16, 'U16': 16, 'F16': 16, 'BF16': 16,
'I32': 32, 'U32': 32, 'F32': 32,
'I64': 64, 'U64': 64, 'F64': 64, 'C64': 64,
'F4': 4, 'F6_E2M3': 6, 'F6_E3M2': 6,
}
def verify(root):
root = Path(root).resolve(strict=True)
config = json.loads((root / 'config.json').read_text())
if not isinstance(config, dict) or not config:
raise ValueError('missing model configuration')
index = root / 'model.safetensors.index.json'
if index.exists():
index_data = json.loads(index.read_text())
if not isinstance(index_data, dict):
raise ValueError('index must be a JSON object')
mapping = index_data.get('weight_map', {})
if not isinstance(mapping, dict) or not mapping:
raise ValueError('empty weight map')
else:
mapping = None
names = list(mapping.values()) if mapping else ['model.safetensors']
if any(not isinstance(name, str) for name in names):
raise ValueError('shard filenames must be strings')
names = set(names)
tensors = {}
for name in names:
if not isinstance(name, str) or Path(name).name != name or not name.endswith('.safetensors'):
raise ValueError('invalid shard filename')
shard = root / name
if shard.resolve().parent != root or not shard.is_file():
raise ValueError('missing or outside-root shard: ' + name)
size = shard.stat().st_size
with shard.open('rb') as stream:
length = stream.read(8)
if len(length) != 8:
raise ValueError('truncated shard: ' + name)
header_size = struct.unpack('<Q', length)[0]
if header_size < 2 or header_size > 100_000_000 or header_size + 8 > size:
raise ValueError('invalid shard header: ' + name)
header = json.loads(stream.read(header_size))
if not isinstance(header, dict):
raise ValueError('shard header must be a JSON object: ' + name)
payload_size = size - 8 - header_size
entries = {key: value for key, value in header.items() if key != '__metadata__'}
if not entries:
raise ValueError('empty shard: ' + name)
ranges = []
for tensor, info in entries.items():
if not isinstance(info, dict):
raise ValueError('tensor metadata must be a JSON object: ' + tensor)
dtype = info.get('dtype')
if not isinstance(dtype, str) or dtype not in DTYPE_BITS:
raise ValueError('unsupported tensor dtype: ' + str(dtype))
shape = info.get('shape')
if not isinstance(shape, list) or any(type(d) is not int or d < 0 for d in shape):
raise ValueError('tensor shape must contain nonnegative integers: ' + tensor)
offsets = info.get('data_offsets')
if not isinstance(offsets, list) or len(offsets) != 2:
raise ValueError('tensor offsets must contain two integers: ' + tensor)
start, end = offsets
if type(start) is not int or type(end) is not int or not 0 <= start <= end <= payload_size:
raise ValueError('invalid tensor offsets: ' + name)
expected_bits = math.prod(shape) * DTYPE_BITS[dtype]
if expected_bits % 8 or expected_bits // 8 != end - start:
raise ValueError('tensor shape/dtype byte-length mismatch: ' + tensor)
ranges.append((start, end))
cursor = 0
for start, end in sorted(ranges):
if start != cursor:
raise ValueError('noncontiguous tensor payload: ' + name)
cursor = end
if cursor != payload_size:
raise ValueError('incomplete or extra tensor payload: ' + name)
tensors[name] = entries
if mapping:
for tensor, name in mapping.items():
if tensor not in tensors[name]:
raise ValueError('index tensor missing from shard: ' + tensor)
return len(names)
if __name__ == '__main__':
try:
print('weight preflight PASS:', verify(sys.argv[1]), 'complete shard(s)')
except (OSError, ValueError, KeyError, TypeError, IndexError, struct.error) as error:
print('weight preflight REFUSED:', error)
sys.exit(1)