Fix Qwen 2.1 GGUF load: promote Q8 norms to F32 and wire TextEncode latent.
abenzerps Q8_0 ships 1D RMSNorms as packed Q8 (136 vs 128), which breaks Comfy rms_rope; tagger now dequantizes small tensors and the graph uses TextEncodeQwenImage21's 64-ch latent plus AuraFlow shift. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -0,0 +1,132 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Make abenzerps Qwen Image 2.1 DiT GGUF loadable in city96 ComfyUI-GGUF.
|
||||
|
||||
abenzerps/Qwen-Image-2.1-Uncensored-GGUF ships Q8_0 with:
|
||||
- kv_count=0 (no general.architecture)
|
||||
- 1D RMSNorm weights quantized to Q8_0 (logical 128 -> packed 136),
|
||||
which breaks Comfy's fused rms_rope path
|
||||
|
||||
This rewrite:
|
||||
1. Adds general.architecture=qwen_image
|
||||
2. Dequantizes small / 1D tensors to F32 (city96 convert keeps them hiprec)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
# Match city96 ComfyUI-GGUF/tools/convert.py QUANTIZATION_THRESHOLD
|
||||
QUANTIZATION_THRESHOLD = 1024
|
||||
|
||||
|
||||
def main() -> int:
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("src", type=Path)
|
||||
ap.add_argument("dst", type=Path, nargs="?", default=None)
|
||||
ap.add_argument("--arch", default="qwen_image")
|
||||
ap.add_argument("--name", default="qwen-image-2.1")
|
||||
ap.add_argument("--inplace", action="store_true", help="Replace src after a successful tag")
|
||||
args = ap.parse_args()
|
||||
|
||||
import gguf
|
||||
import numpy as np
|
||||
|
||||
src = args.src.resolve()
|
||||
if not src.is_file():
|
||||
print(f"missing source: {src}", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
dst = (args.dst.resolve() if args.dst else src.with_name(src.stem + ".tagged.gguf"))
|
||||
|
||||
reader = gguf.GGUFReader(str(src))
|
||||
|
||||
def get_field(name: str):
|
||||
field = reader.fields.get(name)
|
||||
if field is None:
|
||||
return None
|
||||
try:
|
||||
return field.contents()
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
if dst == src:
|
||||
dst = src.with_name(src.stem + ".retag.gguf")
|
||||
|
||||
f32 = gguf.GGMLQuantizationType.F32
|
||||
compat = {gguf.GGMLQuantizationType.F32, gguf.GGMLQuantizationType.F16, gguf.GGMLQuantizationType.BF16}
|
||||
|
||||
print(f"rewriting {src} -> {dst} (arch={args.arch}, tensors={len(reader.tensors)})")
|
||||
writer = gguf.GGUFWriter(str(dst), arch=args.arch, use_temp_file=True)
|
||||
writer.add_name(args.name)
|
||||
try:
|
||||
writer.add_type("model")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
promoted = 0
|
||||
for tensor in reader.tensors:
|
||||
logical = tuple(int(x) for x in tensor.shape)
|
||||
# GGUF stores dims reversed vs torch; logical numel is what matters
|
||||
numel = 1
|
||||
for d in logical:
|
||||
numel *= d
|
||||
qtype = tensor.tensor_type
|
||||
data = tensor.data
|
||||
if hasattr(data, "copy"):
|
||||
data = data.copy()
|
||||
|
||||
needs_f32 = qtype not in compat and (len(logical) <= 1 or numel <= QUANTIZATION_THRESHOLD)
|
||||
if needs_f32:
|
||||
# Dequantize packed blocks -> float32 with logical shape (GGUF order)
|
||||
dequant = gguf.quants.dequantize(np.asarray(data), qtype).astype(np.float32, copy=False)
|
||||
expected = numel
|
||||
if dequant.size != expected:
|
||||
print(
|
||||
f"dequant size mismatch {tensor.name}: got {dequant.size} want {expected} "
|
||||
f"shape={logical} qtype={qtype}",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return 3
|
||||
data = dequant.reshape(logical)
|
||||
qtype = f32
|
||||
promoted += 1
|
||||
|
||||
writer.add_tensor(tensor.name, data, raw_dtype=qtype)
|
||||
|
||||
writer.write_header_to_file()
|
||||
writer.write_kv_data_to_file()
|
||||
writer.write_tensors_to_file(progress=True)
|
||||
writer.close()
|
||||
|
||||
check = gguf.GGUFReader(str(dst))
|
||||
field = check.fields.get("general.architecture")
|
||||
got = field.contents() if field is not None else None
|
||||
print(
|
||||
f"verified architecture={got!r} tensors={len(check.tensors)} "
|
||||
f"promoted_f32={promoted} size={dst.stat().st_size}"
|
||||
)
|
||||
if got != args.arch or len(check.tensors) != len(reader.tensors):
|
||||
print("tag failed", file=sys.stderr)
|
||||
return 2
|
||||
|
||||
# Spot-check a known norm weight is F32 @ 128
|
||||
sample = next((t for t in check.tensors if t.name.endswith("attn.norm_q.weight")), None)
|
||||
if sample is not None:
|
||||
print(
|
||||
f"sample {sample.name}: shape={tuple(int(x) for x in sample.shape)} "
|
||||
f"type={sample.tensor_type.name} data_shape={tuple(sample.data.shape)}"
|
||||
)
|
||||
|
||||
if args.inplace:
|
||||
bak = src.with_suffix(src.suffix + ".untagged.bak")
|
||||
if bak.exists():
|
||||
bak.unlink()
|
||||
src.replace(bak)
|
||||
dst.replace(src)
|
||||
print(f"inplace: {src} (backup {bak})")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Reference in New Issue
Block a user