Image21-INT8 / scripts /provenance.py
ixim's picture
Release verified Image21-INT8 conversion
6f8c80a verified
Raw History Blame Contribute Delete
4.5 kB
"""Validate complete model structure and bind evaluation to exact model files."""
import hashlib
import json
from pathlib import Path
from safetensors import safe_open
from scripts.integrity import sha256
REVISION = 'b3179ad355be050328e483a9dfdd9e60cd62adfa'
def validate_model_structure(root):
root = Path(root)
required = ['model_index.json', 'processor/tokenizer.json',
'processor/tokenizer_config.json', 'processor/preprocessor_config.json',
'scheduler/scheduler_config.json']
required += [f'{c}/config.json' for c in ('transformer', 'text_encoder', 'vae')]
missing = [p for p in required if not (root / p).is_file()]
if missing:
raise ValueError(f'Missing model components/configs: {missing}')
model_index = json.loads((root / 'model_index.json').read_text())
if model_index.get('_class_name') != 'QwenImage21Pipeline':
raise ValueError('Unexpected pipeline class')
quantized = []
for component in ('transformer', 'text_encoder', 'vae'):
folder = root / component
config = json.loads((folder / 'config.json').read_text())
files = sorted(folder.glob('*.safetensors'))
indexes = list(folder.glob('*.safetensors.index.json'))
if not files or len(indexes) > 1 or (len(files) > 1 and not indexes):
raise ValueError(f'Missing or ambiguous shards: {component}')
index = json.loads(indexes[0].read_text())['weight_map'] if indexes else None
if index is not None and set(index.values()) != {f.name for f in files}:
raise ValueError(f'Incomplete shard index: {component}')
tensor_names, integer_count = set(), 0
for file in files:
with safe_open(str(file), framework='np') as f:
keys = set(f.keys())
if not keys or tensor_names & keys:
raise ValueError(f'Empty/duplicate tensors: {file}')
if index is not None and any(index.get(k) != file.name for k in keys):
raise ValueError(f'Shard tensor mapping mismatch: {file}')
tensor_names |= keys
integer_count += sum(f.get_slice(k).get_dtype() == 'I8' for k in keys)
if index is not None and tensor_names != set(index):
raise ValueError(f'Missing indexed tensors: {component}')
is_int8 = bool(config.get('quantization_config', {}).get('load_in_8bit'))
if is_int8 != bool(integer_count):
raise ValueError(f'Quantization config/storage mismatch: {component}')
if component != 'vae':
quantized.append(is_int8)
if quantized not in ([False, False], [True, True]):
raise ValueError('Expected both main components to use the same release precision')
return 'int8' if all(quantized) else 'bf16'
def inference_files(root):
root = Path(root)
result = [root / 'model_index.json']
for component in ('transformer', 'text_encoder', 'vae', 'processor', 'scheduler'):
result.extend(p for p in (root / component).rglob('*')
if p.is_file() and not any(part.startswith('.') for part in p.relative_to(root).parts))
return sorted(result)
def model_identity(root, verified_rows=None):
root = Path(root)
kind = validate_model_structure(root)
known = {r['path']: r for r in verified_rows} if verified_rows is not None else None
rows = []
for file in inference_files(root):
name = file.relative_to(root).as_posix()
if known is not None:
if name not in known or file.stat().st_size != known[name]['size']:
raise ValueError(f'Preverified inventory mismatch: {name}')
row = known[name]
else:
row = {'path': name, 'size': file.stat().st_size, 'sha256': sha256(file)}
rows.append(row)
encoded = json.dumps(rows, sort_keys=True, separators=(',', ':')).encode()
return {'kind': kind, 'base_revision': REVISION, 'files': rows,
'fingerprint': hashlib.sha256(encoded).hexdigest()}
def validate_roles(baseline, quantized):
if baseline['kind'] != 'bf16' or quantized['kind'] != 'int8':
raise ValueError('Comparison requires BF16 baseline and INT8 candidate')
if baseline['fingerprint'] == quantized['fingerprint']:
raise ValueError('Baseline and candidate are identical')
if baseline['base_revision'] != quantized['base_revision']:
raise ValueError('Different upstream revisions')