HandEdit-LoRA / scripts /infer_step1x_lora.py
HandEdit's picture
Add files using upload-large-folder tool
ce47bc4 verified
Raw
History Blame Contribute Delete
6.1 kB
#!/usr/bin/env python3
"""Batch image editing with the official Step1X-Edit code and a LoRA file."""
import argparse
import importlib.util
import sys
from pathlib import Path
import torch
from PIL import Image
IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
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):
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 import_official_inference(repo_dir: Path):
inference_file = repo_dir / "inference.py"
if not inference_file.is_file():
raise FileNotFoundError(f"Official inference.py not found: {inference_file}")
sys.path.insert(0, str(repo_dir))
spec = importlib.util.spec_from_file_location("step1x_official_inference", inference_file)
if spec is None or spec.loader is None:
raise RuntimeError(f"Unable to import: {inference_file}")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module
def main():
parser = argparse.ArgumentParser(
description="Batch inference for Step1X-Edit v1.0/v1.1 with a LoRA checkpoint."
)
parser.add_argument("--repo_dir", required=True, help="Official Step1X-Edit repository directory.")
parser.add_argument(
"--model_dir",
required=True,
help="Directory containing the DiT checkpoint, VAE, and Qwen2.5-VL directory.",
)
parser.add_argument(
"--lora",
default="./checkpoints/step1x/inspire_step1x_r32_a16_res512.safetensors",
help="Step1X LoRA .safetensors file.",
)
parser.add_argument("--input_dir", required=True)
parser.add_argument("--output_dir", required=True)
parser.add_argument("--prompt", default=DEFAULT_PROMPT)
parser.add_argument("--version", choices=["v1.0", "v1.1"], default="v1.0")
parser.add_argument("--steps", type=int, default=28)
parser.add_argument("--cfg_guidance", type=float, default=6.0)
parser.add_argument("--size_level", type=int, default=512)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--seed_mode", choices=["fixed", "increment"], default="fixed")
parser.add_argument("--quantized", action="store_true")
parser.add_argument("--offload", 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 Step1X inference.")
repo_dir = Path(args.repo_dir).expanduser().resolve()
model_dir = Path(args.model_dir).expanduser().resolve()
input_root = Path(args.input_dir).expanduser().resolve()
output_root = Path(args.output_dir).expanduser().resolve()
lora_path = Path(args.lora).expanduser().resolve()
if not input_root.is_dir():
raise FileNotFoundError(f"Input directory not found: {input_root}")
if not lora_path.is_file():
raise FileNotFoundError(f"LoRA file not found: {lora_path}")
ckpt_name = (
"step1x-edit-i1258.safetensors"
if args.version == "v1.0"
else "step1x-edit-v1p1-official.safetensors"
)
required = [
model_dir / ckpt_name,
model_dir / "vae.safetensors",
model_dir / "Qwen2.5-VL-7B-Instruct",
]
for path in required:
if not path.exists():
raise FileNotFoundError(f"Required Step1X component not found: {path}")
output_root.mkdir(parents=True, exist_ok=True)
paths = list_images(input_root, args.recursive)
if not paths:
raise RuntimeError(f"No images found under: {input_root}")
official = import_official_inference(repo_dir)
print("[1/2] Loading Step1X-Edit and LoRA...")
generator = official.ImageGenerator(
ae_path=str(model_dir / "vae.safetensors"),
dit_path=str(model_dir / ckpt_name),
qwen2vl_model_path=str(model_dir / "Qwen2.5-VL-7B-Instruct"),
max_length=640,
quantized=args.quantized,
offload=args.offload,
lora=str(lora_path),
mode="flash",
version=args.version,
)
print(f"[2/2] 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:
ref_image = Image.open(input_path).convert("RGB")
result = generator.generate_image(
prompt=args.prompt,
negative_prompt="",
ref_images=ref_image,
num_steps=args.steps,
cfg_guidance=args.cfg_guidance,
seed=current_seed,
num_samples=1,
show_progress=True,
size_level=args.size_level,
)[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()