"""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']}))