File size: 4,131 Bytes
4da87a3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1a38d44
4da87a3
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
"""Convert official BF16 checkpoints to ComfyUI's native mixed NVFP4 format."""
import argparse
import hashlib
import json
from pathlib import Path
import re
import sys
import time

import torch
from safetensors import safe_open
from safetensors.torch import save_file


def digest(path):
    with open(path, "rb") as stream:
        return hashlib.file_digest(stream, "sha256").hexdigest()


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("source", type=Path)
    parser.add_argument("output", type=Path)
    parser.add_argument("--component", choices=["dit", "encoder"], required=True)
    parser.add_argument("--comfy-root", type=Path, required=True)
    args = parser.parse_args()
    sys.path.insert(0, str(args.comfy_root.resolve()))
    # ComfyUI parses process arguments when its operation registry is imported.
    sys.argv = [sys.argv[0]]
    from comfy.quant_ops import TensorCoreNVFP4Layout

    patterns = {
        "dit": r"transformer_blocks\.\d+\.(attn\.(to_q|to_k|to_v|to_out\.0)|img_mlp\.(gate_up|out))\.weight",
        "encoder": r"model\.layers\.\d+\.(self_attn\.(q_proj|k_proj|v_proj|o_proj)|mlp\.(gate_proj|up_proj|down_proj))\.weight",
    }
    output, records = {}, []
    start = time.perf_counter()
    with safe_open(args.source, framework="pt", device="cpu") as source:
        metadata = dict(source.metadata() or {})
        for key in source.keys():
            tensor = source.get_tensor(key)
            record = {"name": key, "shape": list(tensor.shape), "source_dtype": str(tensor.dtype)}
            if re.fullmatch(patterns[args.component], key):
                if tensor.shape[1] % 32 or tensor.shape[0] % 8:
                    raise ValueError(f"Unsupported NVFP4 GEMM alignment: {key} {tensor.shape}")
                original = tensor.to("cuda")
                scale = original.float().abs().amax().clamp_min(1e-12) / (448 * 6)
                packed, params = TensorCoreNVFP4Layout.quantize(original, scale=scale)
                reconstructed = TensorCoreNVFP4Layout.dequantize(packed, params)
                error = (original.float() - reconstructed.float()).square().mean()
                record.update(storage="nvfp4", relative_rmse=float((error / original.float().square().mean().clamp_min(1e-30)).sqrt()))
                for suffix, value in TensorCoreNVFP4Layout.state_dict_tensors(packed, params).items():
                    output[key + suffix] = value.cpu().contiguous()
                marker = json.dumps({"format": "nvfp4"}, separators=(",", ":")).encode()
                output[key.removesuffix("weight") + "comfy_quant"] = torch.tensor(list(marker), dtype=torch.uint8)
                del original, reconstructed, packed, params
                print(f"NVFP4 {key} relative_rmse={record['relative_rmse']:.5f}", flush=True)
            else:
                output[key] = tensor.clone()
                record.update(storage=str(tensor.dtype), sha256=hashlib.sha256(tensor.view(torch.uint8).numpy().tobytes()).hexdigest())
            records.append(record)
    count = sum(r["storage"] == "nvfp4" for r in records)
    expected = 192 if args.component == "dit" else 252
    if count != expected:
        raise ValueError(f"Recipe matched {count} matrices, expected {expected}")
    metadata.update(format="pt", conversion="Qwen Image 2.1 native NVFP4; protected tensors unchanged; Built with Qwen")
    args.output.parent.mkdir(parents=True, exist_ok=True)
    save_file(output, str(args.output), metadata=metadata)
    report = {"source": args.source.name, "source_sha256": digest(args.source), "output": args.output.name,
              "output_sha256": digest(args.output), "output_bytes": args.output.stat().st_size,
              "nvfp4_matrices": count, "elapsed_seconds": time.perf_counter() - start, "tensors": records}
    args.output.with_suffix(".manifest.json").write_text(json.dumps(report, indent=2), encoding="utf-8")
    print(json.dumps({k: v for k, v in report.items() if k != "tensors"}), flush=True)


if __name__ == "__main__":
    main()