Spaces:
Runtime error
Runtime error
File size: 6,574 Bytes
80cb2a6 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 | #!/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())
|