BennyDaBall's picture
Link original Qwen model and finish naming updates
1a38d44 verified
Raw
History Blame Contribute Delete
4.13 kB
"""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()