Clover-Image-Tiny-Inpaint / inpainting /evaluate_semantic.py
neonforestmist's picture
Document and package context-aware inpainting v2
95bf78b verified
Raw
History Blame
18.7 kB
#!/usr/bin/env python3
"""Compare semantic inpainting quality on controlled, context-rich scenes."""
from __future__ import annotations
import argparse
import gc
import json
from pathlib import Path
from typing import Any
import numpy as np
import torch
from diffusers import (
DPMSolverMultistepScheduler,
StableDiffusionInpaintPipeline,
StableDiffusionPipeline,
)
from PIL import Image, ImageDraw, ImageFilter
from transformers import CLIPModel, CLIPProcessor
SOURCE_CASES = (
{
"name": "park-bench-cat",
"source_prompt": (
"a detailed photograph of an empty wooden park bench centered in a leafy "
"park, no people, no animals"
),
"edit_prompt": (
"a tabby cat sitting naturally on the wooden park bench, detailed photography"
),
"score_prompt": "a detailed photograph of a tabby cat",
"mask": (148, 224, 366, 400),
"shape": "ellipse",
},
{
"name": "kitchen-kettle",
"source_prompt": (
"a detailed photograph of an empty kitchen countertop viewed straight on, "
"warm daylight, no objects in the center"
),
"edit_prompt": (
"a glossy red enamel kettle resting naturally on the kitchen countertop, "
"detailed photography"
),
"score_prompt": "a detailed photograph of a glossy red enamel kettle",
"mask": (166, 236, 350, 414),
"shape": "rounded_rectangle",
},
{
"name": "garden-greenhouse",
"source_prompt": (
"a realistic moonlit garden with an empty grassy clearing in the center, "
"lush plants around the clearing"
),
"edit_prompt": (
"a tiny glass greenhouse glowing warmly in the moonlit garden clearing, "
"detailed photography"
),
"score_prompt": "a detailed photograph of a tiny glass greenhouse",
"mask": ((170, 400), (147, 245), (193, 161), (321, 154), (372, 246), (350, 404)),
"shape": "polygon",
},
{
"name": "street-bicycle",
"source_prompt": (
"a detailed photograph of a quiet city street with an empty road in the "
"foreground, late afternoon"
),
"edit_prompt": (
"a bright red bicycle standing naturally on the city street, detailed photography"
),
"score_prompt": "a detailed photograph of a bright red bicycle",
"mask": ((146, 378), (206, 318), (288, 385), (366, 327)),
"shape": "brush",
"width": 92,
},
{
"name": "living-room-dog",
"source_prompt": (
"a detailed photograph of a cozy living room with an empty rug centered on the "
"floor, soft window light"
),
"edit_prompt": (
"a small corgi sitting naturally on the living room rug, detailed photography"
),
"score_prompt": "a detailed photograph of a small corgi dog",
"mask": (156, 252, 360, 450),
"shape": "ellipse",
},
{
"name": "lake-swan",
"source_prompt": (
"a detailed photograph of a calm lake with empty water near the foreground, "
"mountains in the distance"
),
"edit_prompt": (
"a white swan floating naturally on the calm lake water, detailed photography"
),
"score_prompt": "a detailed photograph of a white swan",
"mask": ((151, 367), (208, 314), (272, 382), (355, 328)),
"shape": "brush",
"width": 96,
},
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--base_model", required=True)
parser.add_argument("--base_revision")
parser.add_argument("--baseline_model", required=True)
parser.add_argument("--baseline_revision")
parser.add_argument("--teacher_model", required=True)
parser.add_argument("--teacher_revision")
parser.add_argument("--teacher_variant")
parser.add_argument("--candidate_model", required=True)
parser.add_argument("--clip_model", default="openai/clip-vit-base-patch32")
parser.add_argument("--clip_revision")
parser.add_argument("--steps", type=int, default=30)
parser.add_argument("--guidance_scale", type=float, default=7.5)
parser.add_argument("--mask_crop_padding", type=int, default=0)
parser.add_argument("--seed", type=int, default=20260811)
parser.add_argument("--output_dir", type=Path, required=True)
return parser.parse_args()
def _scheduler(pipeline: Any) -> DPMSolverMultistepScheduler:
return DPMSolverMultistepScheduler.from_config(
pipeline.scheduler.config,
algorithm_type="dpmsolver++",
)
def _release_cuda(model: Any) -> None:
if isinstance(model, torch.nn.Module):
model.to("cpu")
else:
for component in getattr(model, "components", {}).values():
if isinstance(component, torch.nn.Module):
component.to("cpu")
del model
gc.collect()
torch.cuda.empty_cache()
def _mask(case: dict[str, Any]) -> Image.Image:
mask = Image.new("L", (512, 512), 0)
draw = ImageDraw.Draw(mask)
if case["shape"] == "ellipse":
draw.ellipse(case["mask"], fill=255)
elif case["shape"] == "rounded_rectangle":
draw.rounded_rectangle(case["mask"], radius=28, fill=255)
elif case["shape"] == "polygon":
draw.polygon(case["mask"], fill=255)
elif case["shape"] == "brush":
points = case["mask"]
width = int(case["width"])
draw.line(points, fill=255, width=width, joint="curve")
radius = width // 2
for x, y in points:
draw.ellipse((x - radius, y - radius, x + radius, y + radius), fill=255)
else:
raise ValueError(f"Unsupported semantic mask shape: {case['shape']}")
return mask
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, _feather_inside(mask))
def _square_crop_bounds(mask: Image.Image, padding: int) -> tuple[int, int, int, int]:
bounds = mask.convert("L").getbbox()
if bounds is None:
raise ValueError("Cannot crop around an empty mask")
left, top, right, bottom = bounds
side = min(
max(mask.size),
max(right - left, bottom - top) + max(0, padding) * 2,
)
center_x = (left + right) / 2
center_y = (top + bottom) / 2
crop_left = round(center_x - side / 2)
crop_top = round(center_y - side / 2)
crop_left = min(max(0, crop_left), mask.width - side)
crop_top = min(max(0, crop_top), mask.height - side)
return crop_left, crop_top, crop_left + side, crop_top + side
def _generate_sources(args: argparse.Namespace, cases: list[dict[str, Any]]) -> None:
kwargs = {"revision": args.base_revision} if args.base_revision else {}
pipeline = StableDiffusionPipeline.from_pretrained(
args.base_model,
torch_dtype=torch.float16,
safety_checker=None,
requires_safety_checker=False,
**kwargs,
).to("cuda")
pipeline.scheduler = _scheduler(pipeline)
negative = "people, animals, object in the center, blurry, distorted, low detail"
for index, case in enumerate(cases):
generator = torch.Generator(device="cuda").manual_seed(args.seed + index)
response = pipeline(
prompt=case["source_prompt"],
negative_prompt=negative,
num_inference_steps=args.steps,
guidance_scale=args.guidance_scale,
width=512,
height=512,
generator=generator,
)
case["source"] = response.images[0].convert("RGB")
case["mask_image"] = _mask(case)
case["source"].save(args.output_dir / f"source-{index:02d}-{case['name']}.png")
case["mask_image"].save(args.output_dir / f"mask-{index:02d}-{case['name']}.png")
_release_cuda(pipeline)
def _run_inpainting_model(
*,
model_name: str,
revision: str | None,
variant: str | None,
cases: list[dict[str, Any]],
args: argparse.Namespace,
output_dir: Path,
) -> list[Image.Image]:
kwargs = {"revision": revision} if revision else {}
if variant:
kwargs.update({"variant": variant, "use_safetensors": True})
pipeline = StableDiffusionInpaintPipeline.from_pretrained(
model_name,
torch_dtype=torch.float16,
safety_checker=None,
requires_safety_checker=False,
**kwargs,
).to("cuda")
pipeline.scheduler = _scheduler(pipeline)
results = []
for index, case in enumerate(cases):
source = case["source"]
mask_image = case["mask_image"]
crop_bounds = None
pipeline_image = source
pipeline_mask = mask_image
if args.mask_crop_padding > 0:
crop_bounds = _square_crop_bounds(mask_image, args.mask_crop_padding)
pipeline_image = source.crop(crop_bounds).resize(
(512, 512), Image.Resampling.LANCZOS
)
pipeline_mask = mask_image.crop(crop_bounds).resize(
(512, 512), Image.Resampling.NEAREST
)
generator = torch.Generator(device="cuda").manual_seed(args.seed + 10_000 + index)
response = pipeline(
prompt=case["edit_prompt"],
negative_prompt="black patch, blurry, distorted, low detail",
image=pipeline_image,
mask_image=pipeline_mask,
num_inference_steps=args.steps,
guidance_scale=args.guidance_scale,
width=512,
height=512,
generator=generator,
)
generated = response.images[0]
if crop_bounds is not None:
generated = generated.resize(
(crop_bounds[2] - crop_bounds[0], crop_bounds[3] - crop_bounds[1]),
Image.Resampling.LANCZOS,
)
full_generated = source.copy()
full_generated.paste(generated, crop_bounds[:2])
generated = full_generated
result = _composite(generated, source, mask_image)
result.save(output_dir / f"{index:02d}-{case['name']}.png")
results.append(result)
_release_cuda(pipeline)
return results
def _crop_around_mask(image: Image.Image, mask: Image.Image, padding: int = 48) -> Image.Image:
bounds = mask.getbbox()
if bounds is None:
return image
left, top, right, bottom = bounds
return image.crop(
(
max(0, left - padding),
max(0, top - padding),
min(image.width, right + padding),
min(image.height, bottom + padding),
)
)
def _image_metrics(
image: Image.Image,
source: Image.Image,
mask: Image.Image,
) -> dict[str, float | int]:
image_array = np.asarray(image.convert("RGB"), dtype=np.int16)
source_array = np.asarray(source.convert("RGB"), dtype=np.int16)
selected = np.asarray(mask.convert("L")) >= 128
changed = np.any(image_array != source_array, axis=2)
black = np.all(image_array <= 8, axis=2)
absolute_change = np.abs(image_array - source_array).mean(axis=2) / 255.0
eroded = np.asarray(
mask.convert("L").filter(ImageFilter.MinFilter(size=17))
) >= 128
inner_boundary = selected & ~eroded
return {
"mask_area_fraction": float(selected.mean()),
"masked_change_fraction": float(changed[selected].mean()),
"masked_mean_absolute_change": float(absolute_change[selected].mean()),
"boundary_mean_absolute_change": float(
absolute_change[inner_boundary].mean()
),
"masked_black_fraction": float(black[selected].mean()),
"outside_changed_pixels": int(changed[~selected].sum()),
}
def _clip_scores(
*,
model_name: str,
revision: str | None,
cases: list[dict[str, Any]],
outputs: dict[str, list[Image.Image]],
) -> tuple[dict[str, list[float]], dict[str, list[float]]]:
kwargs = {"revision": revision} if revision else {}
processor = CLIPProcessor.from_pretrained(model_name, **kwargs)
model = CLIPModel.from_pretrained(model_name, **kwargs).to("cuda")
scores: dict[str, list[float]] = {}
for label, images in outputs.items():
label_scores = []
for case, image in zip(cases, images):
inputs = processor(
text=[case["score_prompt"]],
images=[_crop_around_mask(image, case["mask_image"])],
return_tensors="pt",
padding=True,
).to("cuda")
with torch.inference_mode():
vision = model.get_image_features(pixel_values=inputs["pixel_values"])
text = model.get_text_features(
input_ids=inputs["input_ids"],
attention_mask=inputs["attention_mask"],
)
vision = vision / vision.norm(dim=-1, keepdim=True)
text = text / text.norm(dim=-1, keepdim=True)
label_scores.append(float((vision @ text.T).item()))
scores[label] = label_scores
teacher_scores = {label: [] for label in outputs}
for index, case in enumerate(cases):
labels = list(outputs)
crops = [
_crop_around_mask(outputs[label][index], case["mask_image"])
for label in labels
]
inputs = processor(images=crops, return_tensors="pt").to("cuda")
with torch.inference_mode():
features = model.get_image_features(pixel_values=inputs["pixel_values"])
features = features / features.norm(dim=-1, keepdim=True)
teacher_index = labels.index("teacher")
similarities = features @ features[teacher_index]
for label, similarity in zip(labels, similarities):
teacher_scores[label].append(float(similarity.item()))
_release_cuda(model)
return scores, teacher_scores
def _make_sheet(
cases: list[dict[str, Any]],
outputs: dict[str, list[Image.Image]],
destination: Path,
) -> None:
labels = ["source + mask", "current", "teacher", "candidate"]
cell = 512
label_height = 34
sheet = Image.new("RGB", (cell * len(labels), (cell + label_height) * len(cases)), "white")
draw = ImageDraw.Draw(sheet)
for index, case in enumerate(cases):
row_y = index * (cell + label_height)
mask_overlay = Image.new("RGB", case["source"].size, (255, 255, 255))
source_mask = Image.blend(
case["source"],
Image.composite(mask_overlay, case["source"], case["mask_image"]),
0.55,
)
images = [source_mask, outputs["baseline"][index], outputs["teacher"][index], outputs["candidate"][index]]
for column, (label, image) in enumerate(zip(labels, images)):
sheet.paste(image, (column * cell, row_y + label_height))
draw.text((column * cell + 8, row_y + 9), label, fill="black")
draw.text((cell + 88, row_y + 9), f"{case['name']}: {case['edit_prompt'][:52]}", fill="black")
sheet.save(destination)
def main() -> None:
args = parse_args()
args.output_dir.mkdir(parents=True, exist_ok=False)
cases = [dict(case) for case in SOURCE_CASES]
_generate_sources(args, cases)
outputs: dict[str, list[Image.Image]] = {}
model_specs = {
"baseline": (args.baseline_model, args.baseline_revision, None),
"teacher": (
args.teacher_model,
args.teacher_revision,
args.teacher_variant,
),
"candidate": (args.candidate_model, None, None),
}
for label, (model_name, revision, variant) in model_specs.items():
output_dir = args.output_dir / label
output_dir.mkdir()
outputs[label] = _run_inpainting_model(
model_name=model_name,
revision=revision,
variant=variant,
cases=cases,
args=args,
output_dir=output_dir,
)
clip_scores, teacher_image_scores = _clip_scores(
model_name=args.clip_model,
revision=args.clip_revision,
cases=cases,
outputs=outputs,
)
records = []
for index, case in enumerate(cases):
record: dict[str, Any] = {
"index": index,
"name": case["name"],
"source_prompt": case["source_prompt"],
"edit_prompt": case["edit_prompt"],
"score_prompt": case["score_prompt"],
}
for label, images in outputs.items():
record[label] = {
**_image_metrics(images[index], case["source"], case["mask_image"]),
"clip_similarity": clip_scores[label][index],
"clip_similarity_to_teacher": teacher_image_scores[label][index],
}
records.append(record)
summary = {
"base_model": args.base_model,
"base_revision": args.base_revision,
"baseline_model": args.baseline_model,
"baseline_revision": args.baseline_revision,
"teacher_model": args.teacher_model,
"teacher_revision": args.teacher_revision,
"teacher_variant": args.teacher_variant,
"candidate_model": args.candidate_model,
"clip_model": args.clip_model,
"clip_revision": args.clip_revision,
"steps": args.steps,
"guidance_scale": args.guidance_scale,
"mask_crop_padding": args.mask_crop_padding,
"seed": args.seed,
"cases": records,
"mean_clip_similarity": {
label: float(np.mean(scores)) for label, scores in clip_scores.items()
},
"mean_clip_similarity_to_teacher": {
label: float(np.mean(scores))
for label, scores in teacher_image_scores.items()
},
"candidate_clip_wins_over_baseline": int(
sum(
candidate > baseline
for candidate, baseline in zip(
clip_scores["candidate"], clip_scores["baseline"]
)
)
),
}
(args.output_dir / "metrics.json").write_text(json.dumps(summary, indent=2) + "\n")
_make_sheet(cases, outputs, args.output_dir / "comparison.png")
print(json.dumps(summary, indent=2))
if __name__ == "__main__":
main()