Wan2.2-NVFP4-Sparse / convert_lightx2v_nvfp4_to_comfy.py
charles2530's picture
Add files using upload-large-folder tool
e2634b7 verified
Raw History Blame Contribute Delete
5.94 kB
#!/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()