#!/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()