Spaces:
Runtime error
Runtime error
Add LangSmith tracing, structured logging, checkpoint sweep; fix base repo, fp16 variant, batch seeds, set_adapters, attention slicing
80cb2a6 verified Download scripts/compare_checkpoints.py from whosouravsharma/diffusiondb-sd15-lora-inference: direct link, hf CLI and curl.
- Browser
- Download file 6.57 kB
-
https://huggingface.co/spaces/whosouravsharma/diffusiondb-sd15-lora-inference/resolve/main/scripts/compare_checkpoints.py
- Command line
-
hf download hf://spaces/whosouravsharma/diffusiondb-sd15-lora-inference/scripts/compare_checkpoints.py
-
curl -L -o compare_checkpoints.py https://huggingface.co/spaces/whosouravsharma/diffusiondb-sd15-lora-inference/resolve/main/scripts/compare_checkpoints.py
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 | |
| 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()) | |