#!/usr/bin/env python3 """Compare two Clover inpainting checkpoints on held-out reconstruction masks.""" from __future__ import annotations import argparse import gc import json import random from pathlib import Path from typing import Any import numpy as np import torch from datasets import load_dataset from diffusers import DPMSolverMultistepScheduler, StableDiffusionInpaintPipeline from PIL import Image, ImageDraw, ImageFilter, ImageOps from inpainting.masks import mask_area_fraction, random_mask def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--baseline_model", required=True) parser.add_argument("--baseline_revision") parser.add_argument("--candidate_model", required=True) parser.add_argument("--dataset_name", required=True) parser.add_argument("--dataset_revision") parser.add_argument("--dataset_split", default="train") parser.add_argument("--image_column", default="image") parser.add_argument("--caption_column", default="caption") parser.add_argument("--validation_samples", type=int, default=128) parser.add_argument("--sample_count", type=int, default=6) parser.add_argument("--steps", type=int, default=30) parser.add_argument("--guidance_scale", type=float, default=7.5) parser.add_argument("--seed", type=int, default=20260811) parser.add_argument("--output_dir", type=Path, required=True) return parser.parse_args() def _caption(value: Any) -> str: if isinstance(value, list): value = value[0] if value else "" return " ".join(str(value or "").split()) def _prepare_image(image: Image.Image, resolution: int = 512) -> Image.Image: image = ImageOps.exif_transpose(image).convert("RGB") side = min(image.size) left = (image.width - side) // 2 top = (image.height - side) // 2 return image.crop((left, top, left + side, top + side)).resize( (resolution, resolution), Image.Resampling.LANCZOS ) def _feather_inside(mask: Image.Image, radius: int = 6) -> Image.Image: binary = mask.convert("L").point(lambda value: 255 if value >= 128 else 0) softened = binary.filter(ImageFilter.GaussianBlur(radius=radius)) return Image.composite(softened, Image.new("L", binary.size, 0), binary) def _composite(generated: Image.Image, source: Image.Image, mask: Image.Image) -> Image.Image: return Image.composite(generated.convert("RGB"), source.convert("RGB"), _feather_inside(mask)) def _masked_mae(result: Image.Image, target: Image.Image, mask: Image.Image) -> float: result_array = np.asarray(result, dtype=np.float32) / 255.0 target_array = np.asarray(target, dtype=np.float32) / 255.0 mask_array = np.asarray(mask.convert("L"), dtype=np.float32) / 255.0 denominator = max(1.0, mask_array.sum() * 3.0) return float((np.abs(result_array - target_array) * mask_array[..., None]).sum() / denominator) def _black_fraction(image: Image.Image, mask: Image.Image) -> float: pixels = np.asarray(image.convert("RGB"), dtype=np.uint8) selected = np.asarray(mask.convert("L")) >= 128 if not selected.any(): return 1.0 return float(np.all(pixels[selected] <= 8, axis=1).mean()) def _load_pipeline(model: str, *, revision: str | None = None): kwargs = {"revision": revision} if revision else {} pipeline = StableDiffusionInpaintPipeline.from_pretrained( model, torch_dtype=torch.float16, safety_checker=None, requires_safety_checker=False, **kwargs, ) pipeline.scheduler = DPMSolverMultistepScheduler.from_config( pipeline.scheduler.config, algorithm_type="dpmsolver++", ) return pipeline.to("cuda") def _run_model( *, model_name: str, revision: str | None, cases: list[dict[str, Any]], steps: int, guidance_scale: float, output_dir: Path, ) -> list[dict[str, Any]]: pipeline = _load_pipeline(model_name, revision=revision) records = [] for index, case in enumerate(cases): generator = torch.Generator(device="cuda").manual_seed(case["seed"]) response = pipeline( prompt=case["caption"], image=case["source"], mask_image=case["mask"], num_inference_steps=steps, guidance_scale=guidance_scale, generator=generator, height=512, width=512, ) raw = response.images[0].convert("RGB") result = _composite(raw, case["source"], case["mask"]) result.save(output_dir / f"{index:02d}.png") records.append( { "index": index, "masked_mae": _masked_mae(result, case["source"], case["mask"]), "black_fraction": _black_fraction(result, case["mask"]), "nsfw_content_detected": None, } ) del pipeline gc.collect() torch.cuda.empty_cache() return records def _make_sheet( cases: list[dict[str, Any]], baseline_dir: Path, candidate_dir: Path, destination: Path, ) -> None: cell = 512 label_height = 32 sheet = Image.new("RGB", (cell * 3, (cell + label_height) * len(cases)), "white") draw = ImageDraw.Draw(sheet) for index, case in enumerate(cases): row_y = index * (cell + label_height) source = case["source"].copy() overlay = Image.new("RGB", source.size, (255, 255, 255)) source = Image.blend(source, Image.composite(overlay, source, case["mask"]), 0.55) images = [ source, Image.open(baseline_dir / f"{index:02d}.png").convert("RGB"), Image.open(candidate_dir / f"{index:02d}.png").convert("RGB"), ] for column, image in enumerate(images): sheet.paste(image, (column * cell, row_y + label_height)) draw.text((8, row_y + 8), f"source + mask | {case['caption'][:58]}", fill="black") draw.text((cell + 8, row_y + 8), "current checkpoint", fill="black") draw.text((cell * 2 + 8, row_y + 8), "candidate", fill="black") sheet.save(destination) def main() -> None: args = parse_args() args.output_dir.mkdir(parents=True, exist_ok=False) baseline_dir = args.output_dir / "baseline" candidate_dir = args.output_dir / "candidate" baseline_dir.mkdir() candidate_dir.mkdir() dataset = load_dataset( args.dataset_name, split=args.dataset_split, revision=args.dataset_revision, ).shuffle(seed=args.seed) if args.validation_samples < args.sample_count or args.validation_samples >= len(dataset): raise ValueError("validation set must contain at least sample_count records") validation = dataset.select( range(len(dataset) - args.validation_samples, len(dataset)) ).select(range(args.sample_count)) cases = [] for index, example in enumerate(validation): image = example[args.image_column] if not isinstance(image, Image.Image): image = Image.fromarray(np.asarray(image)) source = _prepare_image(image) rng = random.Random(args.seed + index * 1009) mask = random_mask((512, 512), rng, min_area=0.08, max_area=0.45) caption = _caption(example[args.caption_column]) cases.append( { "source": source, "mask": mask, "caption": caption, "seed": args.seed + index, "mask_area": mask_area_fraction(mask), } ) source.save(args.output_dir / f"source-{index:02d}.png") mask.save(args.output_dir / f"mask-{index:02d}.png") baseline_records = _run_model( model_name=args.baseline_model, revision=args.baseline_revision, cases=cases, steps=args.steps, guidance_scale=args.guidance_scale, output_dir=baseline_dir, ) candidate_records = _run_model( model_name=args.candidate_model, revision=None, cases=cases, steps=args.steps, guidance_scale=args.guidance_scale, output_dir=candidate_dir, ) _make_sheet(cases, baseline_dir, candidate_dir, args.output_dir / "comparison.png") metrics = { "baseline_model": args.baseline_model, "baseline_revision": args.baseline_revision, "candidate_model": args.candidate_model, "dataset": args.dataset_name, "dataset_revision": args.dataset_revision, "sample_count": args.sample_count, "steps": args.steps, "baseline": baseline_records, "candidate": candidate_records, "baseline_mean_masked_mae": float( np.mean([record["masked_mae"] for record in baseline_records]) ), "candidate_mean_masked_mae": float( np.mean([record["masked_mae"] for record in candidate_records]) ), } (args.output_dir / "metrics.json").write_text(json.dumps(metrics, indent=2) + "\n") print(json.dumps(metrics, indent=2)) if __name__ == "__main__": main()