#!/usr/bin/env python3 """Convert LightX2V Wan2.2 NVFP4 safetensors to ComfyUI NVFP4 format.""" from __future__ import annotations import argparse import json from pathlib import Path import torch from safetensors import safe_open from safetensors.torch import save_file COMFY_QUANT_CONF = {"format": "nvfp4"} QUANT_LAYER_CONF = {"format": "nvfp4", "full_precision_matrix_mult": False} def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser( description=( "Convert LightX2V Wan2.2 NVFP4 Sparse safetensors into ComfyUI's " "native NVFP4 safetensors convention." ) ) parser.add_argument( "inputs", nargs="*", type=Path, help="Input .safetensors files. Defaults to non-_comfy safetensors in the current directory.", ) parser.add_argument( "--output-dir", type=Path, default=Path("."), help="Directory for converted files. Defaults to the current directory.", ) parser.add_argument( "--suffix", default="_comfy", help="Suffix appended before .safetensors. Defaults to _comfy.", ) parser.add_argument( "--overwrite", action="store_true", help="Overwrite existing converted files.", ) parser.add_argument( "--dry-run", action="store_true", help="Inspect inputs and print planned outputs without writing files.", ) return parser.parse_args() def default_inputs() -> list[Path]: return sorted( p for p in Path(".").glob("*.safetensors") if not p.name.endswith("_comfy.safetensors") ) def output_path(input_path: Path, output_dir: Path, suffix: str) -> Path: return output_dir / f"{input_path.stem}{suffix}{input_path.suffix}" def is_quant_layer(prefix: str, keys: set[str]) -> bool: return ( f"{prefix}.weight" in keys and f"{prefix}.weight_scale" in keys and f"{prefix}.alpha" in keys and f"{prefix}.input_global_scale" in keys ) def quant_layer_prefixes(keys: list[str]) -> list[str]: key_set = set(keys) prefixes: list[str] = [] for key in keys: if key.endswith(".weight"): prefix = key[: -len(".weight")] if is_quant_layer(prefix, key_set): prefixes.append(prefix) return prefixes def swap_fp4_nibbles(weight: torch.Tensor) -> torch.Tensor: if weight.dtype != torch.uint8: raise TypeError(f"Expected uint8 packed NVFP4 weight, got {weight.dtype}") low = torch.bitwise_left_shift(torch.bitwise_and(weight, 0x0F), 4) high = torch.bitwise_right_shift(weight, 4) return torch.bitwise_or(low, high).contiguous() def comfy_quant_tensor() -> torch.Tensor: payload = json.dumps(COMFY_QUANT_CONF).encode("utf-8") return torch.tensor(list(payload), dtype=torch.uint8) def convert_one(input_path: Path, output_path_: Path, overwrite: bool, dry_run: bool) -> None: if output_path_.exists() and not overwrite and not dry_run: raise FileExistsError(f"{output_path_} exists; pass --overwrite to replace it") with safe_open(input_path, framework="pt", device="cpu") as sf: keys = list(sf.keys()) quant_prefixes = quant_layer_prefixes(keys) quant_prefix_set = set(quant_prefixes) if dry_run: print( f"{input_path} -> {output_path_}: " f"{len(keys)} input tensors, {len(quant_prefixes)} quantized layers" ) return tensors: dict[str, torch.Tensor] = {} for key in keys: if key.endswith(".alpha") or key.endswith(".input_global_scale"): prefix = key.rsplit(".", 1)[0] if prefix in quant_prefix_set: continue if key.endswith(".weight"): prefix = key[: -len(".weight")] tensor = sf.get_tensor(key) if prefix in quant_prefix_set: tensors[key] = swap_fp4_nibbles(tensor) alpha = sf.get_tensor(f"{prefix}.alpha").to(torch.float32) input_global_scale = sf.get_tensor(f"{prefix}.input_global_scale").to(torch.float32) tensors[f"{prefix}.weight_scale_2"] = (alpha * input_global_scale).contiguous() tensors[f"{prefix}.input_scale"] = (1.0 / input_global_scale).contiguous() tensors[f"{prefix}.comfy_quant"] = comfy_quant_tensor() else: tensors[key] = tensor.contiguous() if not tensor.is_contiguous() else tensor continue tensor = sf.get_tensor(key) tensors[key] = tensor.contiguous() if not tensor.is_contiguous() else tensor metadata = { "format": "pt", "_quantization_metadata": json.dumps( {"layers": {prefix: QUANT_LAYER_CONF for prefix in quant_prefixes}}, sort_keys=True, ), } tmp_path = output_path_.with_name(f"{output_path_.name}.tmp") if tmp_path.exists(): tmp_path.unlink() save_file(tensors, tmp_path, metadata=metadata) tmp_path.replace(output_path_) print(f"Wrote {output_path_} ({len(tensors)} tensors, {len(quant_prefixes)} quantized layers)") def main() -> None: args = parse_args() inputs = args.inputs or default_inputs() if not inputs: raise SystemExit("No input .safetensors files found") args.output_dir.mkdir(parents=True, exist_ok=True) for input_path in inputs: input_path = input_path.resolve() if input_path.suffix != ".safetensors": raise ValueError(f"Expected a .safetensors input, got {input_path}") out = output_path(input_path, args.output_dir, args.suffix).resolve() convert_one(input_path, out, overwrite=args.overwrite, dry_run=args.dry_run) if __name__ == "__main__": main()