← 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)