Image21-INT4 / scripts /provenance.py
ixim's picture
Release verified Image21-INT4 conversion
9116984 verified
Raw History Blame Contribute Delete
4.68 kB
"""Validate pipeline structure and bind an evaluation to exact files."""
import hashlib
import json
from pathlib import Path
from scripts.integrity import sha256
REVISION = 'b3179ad355be050328e483a9dfdd9e60cd62adfa'
def _quantization(config):
raw = config.get('quantization_config') or {}
method = str(raw.get('quant_method', '')).lower()
dtype = str(raw.get('weights_dtype', '')).lower()
return method == 'sdnq' and dtype == 'uint4' and not raw.get('use_quantized_matmul')
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'{component}/config.json' for component in ('transformer', 'text_encoder', 'vae')]
missing = [path for path in required if not (root / path).is_file()]
if missing:
raise ValueError(f'Missing model components/configs: {missing}')
model_index = json.loads((root / 'model_index.json').read_text(encoding='utf-8'))
if model_index.get('_class_name') != 'QwenImage21Pipeline':
raise ValueError('Unexpected pipeline class')
flags = []
for component in ('transformer', 'text_encoder', 'vae'):
folder = root / component
config = json.loads((folder / 'config.json').read_text(encoding='utf-8'))
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(encoding='utf-8'))['weight_map'] if indexes else None
if index is not None and set(index.values()) != {file.name for file in files}:
raise ValueError(f'Incomplete shard index: {component}')
names = set()
for file in files:
from safetensors import safe_open
with safe_open(str(file), framework='np') as handle:
keys = set(handle.keys())
if not keys or names & keys:
raise ValueError(f'Empty or duplicate tensors: {file}')
if index is not None and any(index.get(key) != file.name for key in keys):
raise ValueError(f'Shard tensor mapping mismatch: {file}')
names |= keys
if index is not None and names != set(index):
raise ValueError(f'Missing indexed tensors: {component}')
quantized = _quantization(config)
if component == 'vae' and quantized:
raise ValueError('VAE must remain floating point')
if component != 'vae':
flags.append(quantized)
if flags not in ([False, False], [True, True]):
raise ValueError('Transformer and text encoder must use the same precision')
return 'int4' if all(flags) 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(path for path in (root / component).rglob('*')
if path.is_file() and not any(part.startswith('.') for part in path.relative_to(root).parts))
return sorted(result)
def model_identity(root, verified_rows=None):
root = Path(root)
kind = validate_model_structure(root)
known = {row['path']: row for row 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'] != 'int4':
raise ValueError('Comparison requires a BF16 baseline and an INT4 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')