HandEdit-LoRA / scripts /infer_longcat_lora.py
HandEdit's picture
Add files using upload-large-folder tool
ce47bc4 verified
Raw
History Blame Contribute Delete
6.08 kB
#!/usr/bin/env python3
"""Batch image editing with LongCat-Image-Edit and a PEFT LoRA adapter."""
from __future__ import annotations
import argparse
from pathlib import Path
import torch
from diffusers import LongCatImageEditPipeline
from peft import PeftModel
from PIL import Image
IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
DEFAULT_BASE = "meituan-longcat/LongCat-Image-Edit"
DEFAULT_PROMPT = (
"Edit only the human hand region. Replace the human hand with a realistic Inspire robotic "
"hand with correct robotic finger structure and joints. Preserve the original wrist pose, "
"palm orientation, finger articulation, grasp geometry, and contact points with the object. "
"The robot hand must be kinematically feasible and physically plausible, without penetrating "
"the object. Keep the object pose, shape, texture, background, lighting, camera viewpoint, "
"and all non-hand regions unchanged."
)
def list_images(input_dir: Path, recursive: bool) -> list[Path]:
iterator = input_dir.rglob("*") if recursive else input_dir.iterdir()
return sorted(
path
for path in iterator
if path.is_file() and path.suffix.lower() in IMAGE_EXTS
)
def apply_lora_scale(model: torch.nn.Module, scale: float) -> int:
"""Multiply every active PEFT LoRA layer scale by ``scale``."""
if scale == 1.0:
return 0
changed = 0
for module in model.modules():
scaling = getattr(module, "scaling", None)
if isinstance(scaling, dict):
for adapter_name in list(scaling):
scaling[adapter_name] *= scale
changed += 1
return changed
def main() -> None:
parser = argparse.ArgumentParser(
description="Batch inference for LongCat-Image-Edit with the HandEdit LoRA."
)
parser.add_argument(
"--base",
default=DEFAULT_BASE,
help=f"Base model directory or model ID (default: {DEFAULT_BASE}).",
)
parser.add_argument(
"--lora",
default="./checkpoints/longcat",
help="PEFT LoRA directory (default: ./checkpoints/longcat).",
)
parser.add_argument("--input_dir", required=True)
parser.add_argument("--output_dir", required=True)
parser.add_argument("--prompt", default=DEFAULT_PROMPT)
parser.add_argument("--negative_prompt", default="")
parser.add_argument("--guidance_scale", type=float, default=4.5)
parser.add_argument("--steps", type=int, default=50)
parser.add_argument("--seed", type=int, default=43)
parser.add_argument(
"--seed_mode", choices=["fixed", "increment"], default="fixed"
)
parser.add_argument("--lora_scale", type=float, default=1.0)
parser.add_argument(
"--offload", choices=["model", "sequential", "none"], default="model"
)
parser.add_argument("--local_files_only", action="store_true")
parser.add_argument("--suffix", default="")
parser.add_argument("--recursive", action="store_true")
parser.add_argument("--skip_existing", action="store_true")
args = parser.parse_args()
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required for practical LongCat inference.")
if args.steps < 2:
raise ValueError("--steps must be at least 2")
if args.guidance_scale < 0:
raise ValueError("--guidance_scale must be non-negative")
input_root = Path(args.input_dir).expanduser().resolve()
output_root = Path(args.output_dir).expanduser().resolve()
lora_root = Path(args.lora).expanduser().resolve()
if not input_root.is_dir():
raise FileNotFoundError(f"Input directory not found: {input_root}")
if not lora_root.is_dir():
raise FileNotFoundError(f"LoRA directory not found: {lora_root}")
paths = list_images(input_root, args.recursive)
if not paths:
raise RuntimeError(f"No images found under: {input_root}")
output_root.mkdir(parents=True, exist_ok=True)
print("[1/3] Loading LongCat-Image-Edit...")
pipe = LongCatImageEditPipeline.from_pretrained(
args.base,
torch_dtype=torch.bfloat16,
local_files_only=args.local_files_only,
)
print("[2/3] Loading HandEdit LoRA...")
pipe.transformer = PeftModel.from_pretrained(
pipe.transformer,
str(lora_root),
is_trainable=False,
)
changed = apply_lora_scale(pipe.transformer, args.lora_scale)
if args.lora_scale != 1.0:
print(
f"[INFO] LoRA scale={args.lora_scale}; "
f"adjusted {changed} active LoRA layers."
)
if args.offload == "model":
pipe.enable_model_cpu_offload()
elif args.offload == "sequential":
pipe.enable_sequential_cpu_offload()
else:
pipe.to("cuda")
print(f"[3/3] Processing {len(paths)} images...")
current_seed = args.seed
for index, input_path in enumerate(paths, start=1):
relative = input_path.relative_to(input_root)
output_path = output_root / relative.with_name(
relative.stem + args.suffix + ".png"
)
output_path.parent.mkdir(parents=True, exist_ok=True)
if args.skip_existing and output_path.exists():
print(f"[{index}/{len(paths)}] SKIP {output_path}")
else:
with Image.open(input_path) as source:
image = source.convert("RGB")
generator = torch.Generator("cpu").manual_seed(current_seed)
result = pipe(
image,
args.prompt,
negative_prompt=args.negative_prompt,
guidance_scale=args.guidance_scale,
num_inference_steps=args.steps,
num_images_per_prompt=1,
generator=generator,
).images[0]
result.save(output_path)
print(f"[{index}/{len(paths)}] OK {input_path.name} -> {output_path}")
if args.seed_mode == "increment":
current_seed += 1
print(f"[DONE] Results saved under: {output_root}")
if __name__ == "__main__":
main()