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