#!/usr/bin/env python3 """Batch image editing with OmniGen-v1 and a fine-tuned LoRA checkpoint.""" import argparse from pathlib import Path from PIL import Image from OmniGen import OmniGenPipeline IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"} DEFAULT_BASE = "Shitao/OmniGen-v1" DEFAULT_PROMPT = ( "<|image_1|> 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 target_size(image_path: Path, max_size: int): with Image.open(image_path) as image: width, height = image.size scale = min(1.0, max_size / max(width, height)) width = max(16, round(width * scale / 16) * 16) height = max(16, round(height * scale / 16) * 16) return width, height def main(): parser = argparse.ArgumentParser( description="Batch inference for OmniGen-v1 with a LoRA checkpoint." ) parser.add_argument( "--base", default=DEFAULT_BASE, help=f"OmniGen-v1 model directory or model ID (default: {DEFAULT_BASE}).", ) parser.add_argument( "--lora", default="./checkpoints/omnigen", help="LoRA checkpoint directory (default: ./checkpoints/omnigen).", ) parser.add_argument("--input_dir", required=True) parser.add_argument("--output_dir", required=True) parser.add_argument("--prompt", default=DEFAULT_PROMPT) parser.add_argument("--steps", type=int, default=50) parser.add_argument("--guidance_scale", type=float, default=2.5) parser.add_argument("--img_guidance_scale", type=float, default=1.6) parser.add_argument("--max_size", type=int, default=512) parser.add_argument("--seed", type=int, default=43) parser.add_argument("--seed_mode", choices=["fixed", "increment"], default="fixed") parser.add_argument("--suffix", default="") parser.add_argument("--recursive", action="store_true") parser.add_argument("--skip_existing", action="store_true") parser.add_argument("--offload_model", action="store_true") args = parser.parse_args() input_root = Path(args.input_dir).expanduser().resolve() output_root = Path(args.output_dir).expanduser().resolve() if not input_root.is_dir(): raise FileNotFoundError(f"Input directory not found: {input_root}") 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}") print("[1/3] Loading OmniGen-v1...") pipe = OmniGenPipeline.from_pretrained(args.base) print("[2/3] Merging LoRA checkpoint...") pipe.merge_lora(args.lora) 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: width, height = target_size(input_path, args.max_size) images = pipe( prompt=args.prompt, input_images=[str(input_path)], height=height, width=width, num_inference_steps=args.steps, guidance_scale=args.guidance_scale, img_guidance_scale=args.img_guidance_scale, max_input_image_size=args.max_size, separate_cfg_infer=True, use_kv_cache=True, offload_kv_cache=True, offload_model=args.offload_model, use_input_image_size_as_output=False, seed=current_seed, ) images[0].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()