#!/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()