HandEdit-LoRA / scripts /infer_flux2_lora.py
HandEdit's picture
Add files using upload-large-folder tool
ce47bc4 verified
Raw History Blame Contribute Delete
5.11 kB
#!/usr/bin/env python3
"""Batch image editing with FLUX.2 Klein and a Diffusers-compatible LoRA file."""
import argparse
from pathlib import Path
import torch
from PIL import Image
from diffusers import Flux2KleinPipeline
IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
DEFAULT_BASE = "black-forest-labs/FLUX.2-klein-base-4B"
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 load_lora(pipe, lora_path: str, scale: float):
path = Path(lora_path).expanduser()
adapter_name = "handedit"
if path.is_file():
pipe.load_lora_weights(
str(path.parent),
weight_name=path.name,
adapter_name=adapter_name,
)
else:
pipe.load_lora_weights(str(path), adapter_name=adapter_name)
pipe.set_adapters(adapter_name, adapter_weights=scale)
def main():
parser = argparse.ArgumentParser(
description="Batch inference for FLUX.2 Klein with a LoRA checkpoint."
)
parser.add_argument(
"--base",
default=DEFAULT_BASE,
help=f"FLUX.2 Klein model directory or model ID (default: {DEFAULT_BASE}).",
)
parser.add_argument(
"--lora",
default="./checkpoints/flux2/handedit_flux2_klein4b_lora.safetensors",
help="LoRA .safetensors file or adapter directory.",
)
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=4.0)
parser.add_argument("--lora_scale", type=float, default=1.0)
parser.add_argument("--seed", type=int, default=43)
parser.add_argument("--seed_mode", choices=["fixed", "increment"], default="fixed")
parser.add_argument("--offload", choices=["model", "sequential", "none"], default="model")
parser.add_argument("--local_files_only", 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 FLUX.2 inference.")
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 FLUX.2 Klein...")
pipe = Flux2KleinPipeline.from_pretrained(
args.base,
torch_dtype=torch.bfloat16,
local_files_only=args.local_files_only,
)
print("[2/3] Loading LoRA adapter...")
load_lora(pipe, args.lora, args.lora_scale)
if args.offload == "model":
pipe.enable_model_cpu_offload()
elif args.offload == "sequential":
pipe.enable_sequential_cpu_offload()
else:
pipe.to("cuda")
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:
image = Image.open(input_path).convert("RGB")
generator = torch.Generator("cpu").manual_seed(current_seed)
result = pipe(
prompt=args.prompt,
image=image,
num_inference_steps=args.steps,
guidance_scale=args.guidance_scale,
generator=generator,
).images[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()