File size: 4,820 Bytes
ce47bc4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
#!/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()