#!/usr/bin/env python3 """Export an SFP PEFT generator checkpoint for LightX2V inference.""" from __future__ import annotations import argparse import json import os from pathlib import Path import torch from safetensors.torch import save_file LORA_SUFFIXES = (".lora_A.weight", ".lora_B.weight") KNOWN_PREFIXES = ( "base_model.model.", "model.base_model.model.", "module.base_model.model.", ) def normalize_key(key: str) -> str: for prefix in KNOWN_PREFIXES: if key.startswith(prefix): key = key[len(prefix) :] break key = key.replace(".default.weight", ".weight") if not key.endswith(LORA_SUFFIXES): raise ValueError(f"Unexpected generator LoRA key: {key}") return key def base_weight_key(lora_key: str) -> str: for suffix in LORA_SUFFIXES: if lora_key.endswith(suffix): return f"{lora_key[:-len(suffix)]}.weight" raise ValueError(f"Not a LoRA key: {lora_key}") def export_generator_lora( checkpoint_path: Path, base_index_path: Path, output_path: Path, *, expected_rank: int, expected_pairs: int, ) -> None: checkpoint = torch.load( checkpoint_path, map_location="cpu", weights_only=True, mmap=True, ) if not isinstance(checkpoint, dict) or "generator_lora" not in checkpoint: raise ValueError(f"Missing generator_lora in {checkpoint_path}") with base_index_path.open("r", encoding="utf-8") as handle: base_keys = set(json.load(handle)["weight_map"]) exported: dict[str, torch.Tensor] = {} pair_members: dict[str, set[str]] = {} for source_key, tensor in checkpoint["generator_lora"].items(): key = normalize_key(source_key) target_key = base_weight_key(key) if target_key not in base_keys: raise ValueError( f"LoRA target {target_key} derived from {source_key} is absent " "from the Wan2.1 base checkpoint" ) if key in exported: raise ValueError(f"Duplicate normalized LoRA key: {key}") if tensor.ndim != 2: raise ValueError( f"Expected a matrix for {source_key}, got {tuple(tensor.shape)}" ) member = "A" if key.endswith(".lora_A.weight") else "B" rank = tensor.shape[0] if member == "A" else tensor.shape[1] if rank != expected_rank: raise ValueError( f"Unexpected rank for {source_key}: {rank}, expected {expected_rank}" ) base = key.rsplit(".lora_", 1)[0] pair_members.setdefault(base, set()).add(member) exported[key] = tensor.to(dtype=torch.bfloat16).contiguous() incomplete = sorted( base for base, members in pair_members.items() if members != {"A", "B"} ) if incomplete: raise ValueError(f"Incomplete LoRA pairs: {incomplete[:8]}") if len(pair_members) != expected_pairs: raise ValueError( f"Found {len(pair_members)} LoRA pairs, expected {expected_pairs}" ) if len(exported) != expected_pairs * 2: raise ValueError( f"Found {len(exported)} LoRA tensors, expected {expected_pairs * 2}" ) output_path.parent.mkdir(parents=True, exist_ok=True) temporary_path = output_path.with_suffix(f"{output_path.suffix}.tmp") save_file( exported, temporary_path, metadata={ "format": "pt", "source": str(checkpoint_path), "rank": str(expected_rank), "alpha": str(expected_rank), "pairs": str(expected_pairs), }, ) os.replace(temporary_path, output_path) print( f"Exported {len(exported)} generator LoRA tensors " f"({len(pair_members)} pairs, rank {expected_rank}) to {output_path}" ) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--checkpoint", type=Path, required=True) parser.add_argument("--base-index", type=Path, required=True) parser.add_argument("--output", type=Path, required=True) parser.add_argument("--expected-rank", type=int, default=128) parser.add_argument("--expected-pairs", type=int, default=400) return parser.parse_args() if __name__ == "__main__": args = parse_args() export_generator_lora( args.checkpoint, args.base_index, args.output, expected_rank=args.expected_rank, expected_pairs=args.expected_pairs, )