#!/usr/bin/env python3 """Migrate legacy ComfyUI int4_tensorwise+ConvRot safetensors metadata. This converter does NOT dequantize or requantize model weights. It rewrites: * safetensors __metadata__["_quantization_metadata"] * persisted .comfy_quant JSON tensors, when present * legacy [N, 1] weight_scale header shapes to the official [N] shape The packed INT4 weight payload and FP32 scale bytes are copied unchanged. Only legacy layers explicitly marked ConvRot are converted unless --assume-convrot is supplied. No third-party Python packages are required. """ from __future__ import annotations import argparse import copy import json import math import os import shutil import struct import sys import tempfile from collections import OrderedDict from dataclasses import dataclass, field from pathlib import Path from typing import Any, BinaryIO, Iterable, Mapping, MutableMapping OLD_FORMAT = "int4_tensorwise" NEW_FORMAT = "convrot_w4a4" DEFAULT_CONVROT_GROUPSIZE = 256 OFFICIAL_QUANT_GROUP_SIZE = 64 DEFAULT_LINEAR_DTYPE = "int4" HEADER_LEN_STRUCT = struct.Struct(" int: return len(self.converted_layers) @dataclass(frozen=True) class TensorEntry: name: str info: dict[str, Any] start: int end: int @property def size(self) -> int: return self.end - self.start def _json_dumps(value: Any) -> str: return json.dumps(value, ensure_ascii=False, separators=(",", ":")) def _read_exact(handle: BinaryIO, size: int) -> bytes: data = handle.read(size) if len(data) != size: raise ConversionError(f"Unexpected EOF: wanted {size} bytes, got {len(data)}") return data def read_safetensors_header(path: Path) -> tuple[OrderedDict[str, Any], int, int]: file_size = path.stat().st_size if file_size < HEADER_LEN_STRUCT.size: raise ConversionError("File is too small to be a safetensors file") with path.open("rb") as handle: header_len = HEADER_LEN_STRUCT.unpack(_read_exact(handle, HEADER_LEN_STRUCT.size))[0] data_start = HEADER_LEN_STRUCT.size + header_len if header_len <= 1 or data_start > file_size: raise ConversionError( f"Invalid safetensors header length {header_len} for file size {file_size}" ) raw_header = _read_exact(handle, header_len) try: decoded = raw_header.rstrip(b" \t\r\n\x00").decode("utf-8") header = json.loads(decoded, object_pairs_hook=OrderedDict) except (UnicodeDecodeError, json.JSONDecodeError) as exc: raise ConversionError(f"Invalid safetensors JSON header: {exc}") from exc if not isinstance(header, OrderedDict): raise ConversionError("Safetensors header root is not a JSON object") return header, data_start, file_size def collect_tensor_entries( header: Mapping[str, Any], data_size: int ) -> list[TensorEntry]: entries: list[TensorEntry] = [] for name, info in header.items(): if name == "__metadata__": continue if not isinstance(info, dict): raise ConversionError(f"Tensor header for {name!r} is not an object") offsets = info.get("data_offsets") if ( not isinstance(offsets, list) or len(offsets) != 2 or not all(isinstance(v, int) for v in offsets) ): raise ConversionError(f"Tensor {name!r} has invalid data_offsets: {offsets!r}") start, end = offsets if start < 0 or end < start or end > data_size: raise ConversionError( f"Tensor {name!r} has out-of-range offsets [{start}, {end}] " f"for data section size {data_size}" ) entries.append(TensorEntry(name=name, info=dict(info), start=start, end=end)) entries.sort(key=lambda item: (item.start, item.end, item.name)) previous_end = 0 for entry in entries: if entry.start < previous_end: raise ConversionError(f"Overlapping tensor data near {entry.name!r}") previous_end = entry.end return entries def _is_power_of_four(value: int) -> bool: if value < 4 or value & (value - 1): return False exponent = int(math.log2(value)) return exponent % 2 == 0 def _first_present(configs: Iterable[Mapping[str, Any]], key: str) -> Any: for config in configs: if key in config: return config[key] return None def convert_layer_config( config: Any, *, layer_name: str, assume_convrot: bool, linear_dtype: str, report: Report, ) -> tuple[Any, bool]: """Convert one layer's quantization config, preserving unrelated fields.""" if isinstance(config, str): if config != OLD_FORMAT: return config, False if not assume_convrot: raise ConversionError( f"Layer {layer_name!r} uses {OLD_FORMAT!r} but has no ConvRot marker. " "Refusing to relabel it; use --assume-convrot only when you know it was rotated." ) new_config: dict[str, Any] = { "format": NEW_FORMAT, "params": { "convrot_groupsize": DEFAULT_CONVROT_GROUPSIZE, "quant_group_size": OFFICIAL_QUANT_GROUP_SIZE, "linear_dtype": linear_dtype, }, } report.converted_layers.add(layer_name) return new_config, True if not isinstance(config, dict): return config, False if config.get("format") != OLD_FORMAT: return config, False params = config.get("params") if params is None: params = {} elif not isinstance(params, dict): raise ConversionError( f"Layer {layer_name!r} has a non-object params field: {params!r}" ) sources = (config, params) explicit_convrot = _first_present(sources, "convrot") group_size_value = _first_present(sources, "convrot_groupsize") inferred_convrot = group_size_value is not None if explicit_convrot is False: raise ConversionError( f"Layer {layer_name!r} explicitly has convrot=false. It cannot be safely " f"converted from {OLD_FORMAT!r} to {NEW_FORMAT!r}." ) if explicit_convrot is not True and not inferred_convrot and not assume_convrot: raise ConversionError( f"Layer {layer_name!r} uses {OLD_FORMAT!r} but is not marked ConvRot. " "Use --assume-convrot only if the checkpoint is known to be W4A4 ConvRot." ) convrot_groupsize = ( DEFAULT_CONVROT_GROUPSIZE if group_size_value is None else group_size_value ) if not isinstance(convrot_groupsize, int) or not _is_power_of_four(convrot_groupsize): raise ConversionError( f"Layer {layer_name!r} has invalid convrot_groupsize={convrot_groupsize!r}; " "the official regular Hadamard path requires a power of 4 (4, 16, 64, 256, ...)." ) old_quant_group_size = _first_present(sources, "quant_group_size") if old_quant_group_size not in (None, OFFICIAL_QUANT_GROUP_SIZE): raise ConversionError( f"Layer {layer_name!r} requests quant_group_size={old_quant_group_size!r}, " f"but the official INT4 MMA contract requires {OFFICIAL_QUANT_GROUP_SIZE}." ) new_config = copy.deepcopy(config) new_config["format"] = NEW_FORMAT # Normalize recipe options under `params`, which is the generic modern loader path. new_params = copy.deepcopy(params) for legacy_key in ("convrot", "per_channel", "is_weight"): new_params.pop(legacy_key, None) new_params["convrot_groupsize"] = convrot_groupsize new_params["quant_group_size"] = OFFICIAL_QUANT_GROUP_SIZE new_params["linear_dtype"] = linear_dtype new_config["params"] = new_params # Remove the legacy top-level knobs after moving their meaningful value to params. for legacy_key in ( "convrot", "convrot_groupsize", "quant_group_size", "linear_dtype", "per_channel", "is_weight", ): new_config.pop(legacy_key, None) report.converted_layers.add(layer_name) return new_config, True def convert_quantization_metadata_object( value: Any, *, assume_convrot: bool, linear_dtype: str, report: Report, ) -> tuple[Any, bool]: if not isinstance(value, dict): raise ConversionError("_quantization_metadata JSON must be an object") root = copy.deepcopy(value) target = root.get("_quantization_metadata", root) if not isinstance(target, dict): raise ConversionError("Nested _quantization_metadata must be an object") layers = target.get("layers") if layers is None: return root, False if not isinstance(layers, dict): raise ConversionError("_quantization_metadata.layers must be an object") changed = False for layer_name, layer_config in list(layers.items()): new_config, layer_changed = convert_layer_config( layer_config, layer_name=layer_name, assume_convrot=assume_convrot, linear_dtype=linear_dtype, report=report, ) if layer_changed: layers[layer_name] = new_config report.converted_metadata_layers.add(layer_name) changed = True return root, changed def convert_file_metadata( header: MutableMapping[str, Any], *, assume_convrot: bool, linear_dtype: str, report: Report, ) -> bool: metadata = header.get("__metadata__") if metadata is None: return False if not isinstance(metadata, dict): raise ConversionError("Safetensors __metadata__ must be an object") raw = metadata.get("_quantization_metadata") if raw is None: return False if not isinstance(raw, str): raise ConversionError( "Safetensors metadata values must be strings; _quantization_metadata is not a string" ) try: parsed = json.loads(raw) except json.JSONDecodeError as exc: raise ConversionError(f"Invalid _quantization_metadata JSON: {exc}") from exc converted, changed = convert_quantization_metadata_object( parsed, assume_convrot=assume_convrot, linear_dtype=linear_dtype, report=report, ) if changed: metadata["_quantization_metadata"] = _json_dumps(converted) return changed def read_tensor_bytes(handle: BinaryIO, data_start: int, entry: TensorEntry) -> bytes: handle.seek(data_start + entry.start) return _read_exact(handle, entry.size) def convert_comfy_quant_payload( raw: bytes, *, tensor_name: str, assume_convrot: bool, linear_dtype: str, report: Report, ) -> tuple[bytes, bool]: try: text = raw.rstrip(b"\x00").decode("utf-8") parsed = json.loads(text) except (UnicodeDecodeError, json.JSONDecodeError) as exc: raise ConversionError(f"Invalid JSON in tensor {tensor_name!r}: {exc}") from exc layer_name = tensor_name[: -len(".comfy_quant")] converted, changed = convert_layer_config( parsed, layer_name=layer_name, assume_convrot=assume_convrot, linear_dtype=linear_dtype, report=report, ) if not changed: return raw, False report.converted_comfy_quant_tensors.add(tensor_name) return _json_dumps(converted).encode("utf-8"), True def _numel(shape: Any) -> int | None: if not isinstance(shape, list) or not all(isinstance(v, int) and v >= 0 for v in shape): return None result = 1 for dim in shape: result *= dim return result def _find_matching_layer(tensor_name: str, layers: set[str]) -> str | None: suffix = ".weight_scale" if not tensor_name.endswith(suffix): return None base = tensor_name[: -len(suffix)] if base in layers: return base # Prefix changes can occur during ComfyUI model loading. Use an unambiguous # suffix match only; never guess if multiple metadata layer names fit. candidates = [layer for layer in layers if base.endswith(layer) or layer.endswith(base)] return candidates[0] if len(candidates) == 1 else None def normalize_scale_shape( tensor_name: str, info: MutableMapping[str, Any], *, converted_layers: set[str], keep_scale_shape: bool, report: Report, ) -> bool: if keep_scale_shape or _find_matching_layer(tensor_name, converted_layers) is None: return False shape = info.get("shape") if isinstance(shape, list) and len(shape) == 2 and shape[1] == 1: info["shape"] = [shape[0]] report.normalized_scale_tensors.add(tensor_name) return True return False def validate_converted_layer_tensors( tensor_infos: Mapping[str, Mapping[str, Any]], *, converted_layers: set[str], report: Report, ) -> None: for layer in sorted(converted_layers): weight_name = f"{layer}.weight" scale_name = f"{layer}.weight_scale" weight_info = tensor_infos.get(weight_name) scale_info = tensor_infos.get(scale_name) if weight_info is None or scale_info is None: report.warnings.append( f"Could not validate tensors for {layer!r} by exact name " f"(expected {weight_name!r} and {scale_name!r}); metadata prefix remapping may be in use." ) continue if weight_info.get("dtype") != "I8": raise ConversionError( f"{weight_name!r} dtype is {weight_info.get('dtype')!r}, expected I8 packed INT4" ) weight_shape = weight_info.get("shape") if not ( isinstance(weight_shape, list) and len(weight_shape) == 2 and all(isinstance(v, int) and v > 0 for v in weight_shape) ): raise ConversionError( f"{weight_name!r} must have packed 2D shape [N, K/2], got {weight_shape!r}" ) original_k = weight_shape[1] * 2 if original_k % OFFICIAL_QUANT_GROUP_SIZE != 0: raise ConversionError( f"{weight_name!r} reconstructs K={original_k}, not divisible by " f"official quant_group_size={OFFICIAL_QUANT_GROUP_SIZE}" ) if scale_info.get("dtype") != "F32": raise ConversionError( f"{scale_name!r} dtype is {scale_info.get('dtype')!r}, expected F32" ) scale_numel = _numel(scale_info.get("shape")) if scale_numel != weight_shape[0]: raise ConversionError( f"{scale_name!r} has {scale_numel} elements, expected one scale per output row " f"({weight_shape[0]})" ) def padded_header_bytes(header: Mapping[str, Any]) -> bytes: raw = _json_dumps(header).encode("utf-8") padding = (-len(raw)) % 8 return raw + (b" " * padding) def copy_region( source: BinaryIO, destination: BinaryIO, *, absolute_start: int, size: int, ) -> None: source.seek(absolute_start) remaining = size while remaining: chunk = source.read(min(COPY_CHUNK_SIZE, remaining)) if not chunk: raise ConversionError("Unexpected EOF while copying tensor payload") destination.write(chunk) remaining -= len(chunk) def build_conversion_plan( source_path: Path, *, assume_convrot: bool, linear_dtype: str, keep_scale_shape: bool, ) -> tuple[ OrderedDict[str, Any], int, list[TensorEntry], dict[str, bytes], Report, ]: header, data_start, file_size = read_safetensors_header(source_path) entries = collect_tensor_entries(header, file_size - data_start) report = Report() convert_file_metadata( header, assume_convrot=assume_convrot, linear_dtype=linear_dtype, report=report, ) replacements: dict[str, bytes] = {} with source_path.open("rb") as source: for entry in entries: if not entry.name.endswith(".comfy_quant"): continue if entry.info.get("dtype") not in ("U8", "I8"): raise ConversionError( f"{entry.name!r} must use U8/I8 storage, got {entry.info.get('dtype')!r}" ) raw = read_tensor_bytes(source, data_start, entry) replacement, changed = convert_comfy_quant_payload( raw, tensor_name=entry.name, assume_convrot=assume_convrot, linear_dtype=linear_dtype, report=report, ) if changed: replacements[entry.name] = replacement # Rebuild tensor entries in physical order and recalculate every data offset. output_header: OrderedDict[str, Any] = OrderedDict() if "__metadata__" in header: output_header["__metadata__"] = header["__metadata__"] cursor = 0 final_entries: list[TensorEntry] = [] tensor_infos: dict[str, dict[str, Any]] = {} for entry in entries: info = copy.deepcopy(entry.info) replacement = replacements.get(entry.name) size = len(replacement) if replacement is not None else entry.size if replacement is not None: info["shape"] = [size] normalize_scale_shape( entry.name, info, converted_layers=report.converted_layers, keep_scale_shape=keep_scale_shape, report=report, ) info["data_offsets"] = [cursor, cursor + size] output_header[entry.name] = info tensor_infos[entry.name] = info final_entries.append( TensorEntry(name=entry.name, info=info, start=entry.start, end=entry.end) ) cursor += size validate_converted_layer_tensors( tensor_infos, converted_layers=report.converted_layers, report=report, ) return output_header, data_start, final_entries, replacements, report def write_converted_file( source_path: Path, destination_path: Path, *, header: OrderedDict[str, Any], source_data_start: int, entries: list[TensorEntry], replacements: Mapping[str, bytes], ) -> None: destination_path.parent.mkdir(parents=True, exist_ok=True) header_bytes = padded_header_bytes(header) with source_path.open("rb") as source, destination_path.open("wb") as destination: destination.write(HEADER_LEN_STRUCT.pack(len(header_bytes))) destination.write(header_bytes) for entry in entries: replacement = replacements.get(entry.name) if replacement is not None: destination.write(replacement) else: copy_region( source, destination, absolute_start=source_data_start + entry.start, size=entry.size, ) destination.flush() os.fsync(destination.fileno()) def default_output_path(source: Path) -> Path: if source.suffix.lower() == ".safetensors": return source.with_name(f"{source.stem}.convrot_w4a4.safetensors") return source.with_name(f"{source.name}.convrot_w4a4.safetensors") def print_report(report: Report, source: Path, destination: Path | None) -> None: print(f"Source: {source}") if destination is not None: print(f"Output: {destination}") print(f"Converted layers: {report.converted_count}") print(f" _quantization_metadata entries: {len(report.converted_metadata_layers)}") print(f" persisted .comfy_quant tensors: {len(report.converted_comfy_quant_tensors)}") print(f" normalized weight_scale shapes: {len(report.normalized_scale_tensors)}") for layer in sorted(report.converted_layers): print(f" - {layer}") for warning in report.warnings: print(f"WARNING: {warning}", file=sys.stderr) def parse_args(argv: list[str] | None = None) -> argparse.Namespace: parser = argparse.ArgumentParser( description=( "Convert legacy int4_tensorwise+ConvRot safetensors metadata to the " "official convrot_w4a4 format without requantizing model weights." ) ) parser.add_argument("input", type=Path, help="Input .safetensors checkpoint") parser.add_argument( "output", nargs="?", type=Path, help="Output path (default: .convrot_w4a4.safetensors)", ) parser.add_argument( "--in-place", action="store_true", help="Atomically replace the input file after successful conversion", ) parser.add_argument( "--dry-run", action="store_true", help="Inspect and validate the conversion without writing a file", ) parser.add_argument( "--assume-convrot", action="store_true", help=( "Convert int4_tensorwise entries that lack an explicit convrot marker. " "Dangerous for non-ConvRot checkpoints." ), ) parser.add_argument( "--linear-dtype", choices=("int4", "int8"), default=DEFAULT_LINEAR_DTYPE, help="Official kernel path recorded in metadata (default: int4)", ) parser.add_argument( "--keep-scale-shape", action="store_true", help="Keep legacy weight_scale shape [N,1] instead of normalizing it to [N]", ) parser.add_argument( "--force", action="store_true", help="Overwrite an existing output file", ) parser.add_argument( "--allow-noop", action="store_true", help="Exit successfully even when no eligible legacy ConvRot layers are found", ) return parser.parse_args(argv) def main(argv: list[str] | None = None) -> int: args = parse_args(argv) source = args.input.expanduser().resolve() if not source.is_file(): raise ConversionError(f"Input file does not exist: {source}") if args.in_place and args.output is not None: raise ConversionError("Do not supply an output path together with --in-place") ( output_header, source_data_start, entries, replacements, report, ) = build_conversion_plan( source, assume_convrot=args.assume_convrot, linear_dtype=args.linear_dtype, keep_scale_shape=args.keep_scale_shape, ) if report.converted_count == 0 and not args.allow_noop: raise ConversionError( "No eligible int4_tensorwise ConvRot layers were found. " "The file may already use convrot_w4a4, may be non-ConvRot, or may not carry ComfyUI quantization metadata." ) if args.dry_run: print_report(report, source, None) print("Dry run only; no file was written.") return 0 if args.in_place: fd, temp_name = tempfile.mkstemp( prefix=f".{source.name}.", suffix=".tmp", dir=source.parent ) os.close(fd) temp_path = Path(temp_name) try: write_converted_file( source, temp_path, header=output_header, source_data_start=source_data_start, entries=entries, replacements=replacements, ) shutil.copystat(source, temp_path) os.replace(temp_path, source) finally: temp_path.unlink(missing_ok=True) print_report(report, source, source) else: destination = ( args.output.expanduser().resolve() if args.output is not None else default_output_path(source) ) if destination == source: raise ConversionError("Output equals input; use --in-place for atomic replacement") if destination.exists() and not args.force: raise ConversionError( f"Output already exists: {destination} (use --force to overwrite)" ) write_converted_file( source, destination, header=output_header, source_data_start=source_data_start, entries=entries, replacements=replacements, ) print_report(report, source, destination) print("Done. Packed INT4 weights and FP32 scale payloads were copied without requantization.") return 0 if __name__ == "__main__": try: raise SystemExit(main()) except ConversionError as exc: print(f"ERROR: {exc}", file=sys.stderr) raise SystemExit(2)