File size: 6,104 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
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
#!/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()