diffusiondb-sd15-lora-inference / scripts /compare_checkpoints.py
whosouravsharma's picture
Add LangSmith tracing, structured logging, checkpoint sweep; fix base repo, fp16 variant, batch seeds, set_adapters, attention slicing
80cb2a6 verified
Raw History Blame Contribute Delete
6.57 kB
#!/usr/bin/env python3
"""Render a contact sheet comparing checkpoints and LoRA scales.
This is the thing that actually answers "did checkpoint-4240 overfit?" -- a
fixed prompt set, a fixed seed, rendered across every (checkpoint, scale)
combination and tiled into one labelled PNG you can look at side by side.
python3 -m scripts.compare_checkpoints \
--checkpoints base checkpoint-4240 \
--scales 0.0 0.5 1.0 \
--seed 42 --out sweep.png
Rows are prompts, columns are (checkpoint, scale). Because the seed is held
fixed, every image in a row differs only by the adapter -- which is the whole
point of the lora_scale slider, made legible.
With LANGSMITH_TRACING=true the sweep appears as one parent run in LangSmith
with a child run per image, so the per-image latency and parameters are
grouped under a single comparable experiment.
"""
from __future__ import annotations
import argparse
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from PIL import Image, ImageDraw, ImageFont # noqa: E402
from src.pipeline import LORA_CHECKPOINT, ModelRunner # noqa: E402
from src.tracing import configure_logging, logger, traceable # noqa: E402
DEFAULT_PROMPTS = [
"a steampunk owl inside a glass jar, intricate detail",
"a cyberpunk portrait of a woman, neon rim light",
"an oil painting of a lighthouse in a storm",
]
LABEL_HEIGHT = 28
ROW_LABEL_WIDTH = 220
PADDING = 8
def _font(size: int = 14):
try:
return ImageFont.load_default(size=size)
except TypeError: # Pillow < 10 has no size argument
return ImageFont.load_default()
def _wrap(text: str, width: int) -> list[str]:
lines, current = [], ""
for word in text.split():
candidate = f"{current} {word}".strip()
if len(candidate) * 6 > width and current:
lines.append(current)
current = word
else:
current = candidate
if current:
lines.append(current)
return lines[:4]
def build_sheet(grid: list[list[Image.Image]], row_labels: list[str],
col_labels: list[str]) -> Image.Image:
"""Tile the rendered grid into one labelled image."""
cell_w, cell_h = grid[0][0].size
columns, rows = len(col_labels), len(row_labels)
sheet_w = ROW_LABEL_WIDTH + columns * (cell_w + PADDING) + PADDING
sheet_h = LABEL_HEIGHT + rows * (cell_h + PADDING) + PADDING
sheet = Image.new("RGB", (sheet_w, sheet_h), "white")
draw = ImageDraw.Draw(sheet)
header_font, body_font = _font(15), _font(12)
for index, label in enumerate(col_labels):
x = ROW_LABEL_WIDTH + index * (cell_w + PADDING)
draw.text((x + 4, 6), label, fill="black", font=header_font)
for row_index, images in enumerate(grid):
y = LABEL_HEIGHT + row_index * (cell_h + PADDING)
for line_index, line in enumerate(_wrap(row_labels[row_index],
ROW_LABEL_WIDTH - 16)):
draw.text((8, y + 4 + line_index * 15), line,
fill="black", font=body_font)
for col_index, image in enumerate(images):
x = ROW_LABEL_WIDTH + col_index * (cell_w + PADDING)
sheet.paste(image, (x, y))
return sheet
@traceable(run_type="chain", name="checkpoint_sweep")
def run_sweep(prompts: list[str], checkpoints: list[str], scales: list[float],
seed: int, steps: int, guidance: float, negative: str) -> dict:
"""Render every (prompt, checkpoint, scale) cell. Returns a summary dict."""
col_labels = [f"{checkpoint} @ {scale:g}"
for checkpoint in checkpoints for scale in scales]
grid: list[list[Image.Image]] = [[] for _ in prompts]
total_ms = 0.0
# Loop checkpoints outermost: each one is a full pipeline reload, so this
# pays that cost len(checkpoints) times rather than once per cell.
for checkpoint in checkpoints:
logger.info("loading checkpoint=%s", checkpoint)
runner = ModelRunner(checkpoint=checkpoint)
for scale in scales:
for row_index, prompt in enumerate(prompts):
result = runner.generate(
prompt=prompt, negative_prompt=negative, steps=steps,
guidance=guidance, lora_scale=scale, seed=seed, count=1,
)
grid[row_index].append(result.images[0])
total_ms += result.duration_ms
del runner
# Columns were filled checkpoint-major but scale-minor per checkpoint,
# which already matches col_labels' ordering.
sheet = build_sheet(grid, prompts, col_labels)
return {"sheet": sheet, "cells": len(prompts) * len(col_labels),
"total_ms": total_ms, "columns": col_labels}
def main() -> int:
parser = argparse.ArgumentParser(
description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter,
)
parser.add_argument("--prompts", nargs="+", default=DEFAULT_PROMPTS)
parser.add_argument("--prompts-file", type=Path,
help="one prompt per line; overrides --prompts")
parser.add_argument("--checkpoints", nargs="+",
default=["base", LORA_CHECKPOINT],
help="'base' renders vanilla SD 1.5")
parser.add_argument("--scales", nargs="+", type=float, default=[0.0, 1.0])
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--steps", type=int, default=25)
parser.add_argument("--guidance", type=float, default=7.5)
parser.add_argument("--negative", default="")
parser.add_argument("--out", type=Path, default=Path("sweep.png"))
args = parser.parse_args()
configure_logging()
prompts = args.prompts
if args.prompts_file:
prompts = [line.strip() for line in
args.prompts_file.read_text().splitlines() if line.strip()]
if not prompts:
parser.error("no prompts")
cells = len(prompts) * len(args.checkpoints) * len(args.scales)
logger.info("sweep cells=%d seed=%d steps=%d", cells, args.seed, args.steps)
summary = run_sweep(
prompts=prompts, checkpoints=args.checkpoints, scales=args.scales,
seed=args.seed, steps=args.steps, guidance=args.guidance,
negative=args.negative,
)
summary["sheet"].save(args.out)
logger.info("wrote %s (%d cells, %.0f ms of generation)",
args.out, summary["cells"], summary["total_ms"])
return 0
if __name__ == "__main__":
raise SystemExit(main())