HandEdit-LoRA / scripts /infer_omnigen_lora.py
HandEdit's picture
Add files using upload-large-folder tool
ce47bc4 verified
Raw History Blame Contribute Delete
4.82 kB
#!/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 = (
"<img><|image_1|></img> 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()