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())