neonforestmist's picture
Document and package context-aware inpainting v2
95bf78b verified
Raw
History Blame Contribute Delete
9.08 kB
#!/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()