BennyDaBall's picture
Attribute Qwen Image 2.1 and update release name and workflows
4da87a3 verified
Raw History Blame
2 kB
"""Verify saved native NVFP4 shapes, metadata, scales, and protected tensors."""
import argparse
import hashlib
import json
from pathlib import Path
import torch
from safetensors import safe_open
parser = argparse.ArgumentParser()
parser.add_argument('source', type=Path)
parser.add_argument('converted', type=Path)
args = parser.parse_args()
manifest = json.loads(args.converted.with_suffix('.manifest.json').read_text())
with args.converted.open('rb') as stream:
assert hashlib.file_digest(stream, 'sha256').hexdigest() == manifest['output_sha256']
with safe_open(args.source, framework='pt') as original, safe_open(args.converted, framework='pt') as result:
expected_keys = set(original.keys())
for record in manifest['tensors']:
key = record['name']
value = result.get_tensor(key)
if record['storage'] == 'nvfp4':
marker_key = key.removesuffix('weight') + 'comfy_quant'
assert json.loads(result.get_tensor(marker_key).numpy().tobytes()) == {'format':'nvfp4'}
assert value.dtype == torch.uint8
assert list(value.shape) == [record['shape'][0], record['shape'][1] // 2]
block = result.get_tensor(key + '_scale')
scale = result.get_tensor(key + '_scale_2')
assert block.dtype == torch.float8_e4m3fn and scale.dtype == torch.float32
assert bool(torch.isfinite(block.float()).all()) and bool(torch.isfinite(scale).all())
assert bool((scale > 0).all())
expected_keys.update([key + '_scale', key + '_scale_2', marker_key])
else:
source = original.get_tensor(key)
assert source.dtype == value.dtype and torch.equal(source, value), key
assert set(result.keys()) == expected_keys
print(json.dumps({'file':args.converted.name, 'verified_tensors':len(manifest['tensors']), 'nvfp4_matrices':manifest['nvfp4_matrices'], 'protected_exact':True, 'sha256':manifest['output_sha256']}))