"""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()