Text-to-Image
Diffusers
stable-diffusion
stable-diffusion-diffusers
lora
Eval Results (legacy)
diffusiondb-sd15-lora / eval /eval_job.py
whosouravsharma's picture
Model card: evaluation results, gallery, figures, eval script
8cb43ae verified
Raw History Blame Contribute Delete
33.9 kB
# /// script
# requires-python = ">=3.10"
# dependencies = [
# "torch==2.4.0",
# "diffusers==0.31.0",
# "transformers==4.44.2",
# "peft==0.13.2",
# "accelerate==0.34.2",
# "safetensors",
# "huggingface_hub>=0.24,<1.0",
# "datasets>=2.19,<4",
# "torchmetrics[image]>=1.4",
# "torch-fidelity",
# "scipy",
# "numpy<2",
# "pillow>=10.1",
# "matplotlib",
# ]
# ///
"""Evaluate the DiffusionDB SD 1.5 LoRA without writing to any repo.
hf jobs uv run --flavor a10g-small --timeout 4h \\
-v hf://buckets/whosouravsharma/jobs-artifacts:/out \\
eval/eval_job.py
Deliberately launched WITHOUT --secrets HF_TOKEN. Every input is public, so
the job has no credentials that could write to the model repo, the dataset
repo or either Space. The only writable location is the mounted bucket, and
everything lands under /out/eval/<RUN_NAME>/.
Reads (all public, read-only):
whosouravsharma/diffusiondb-sd15-lora checkpoints
whosouravsharma/text-to-image-diffusiondb-2M v2-clean images + prompts,
latents-512 validation latents
stable-diffusion-v1-5/stable-diffusion-v1-5 base model, safety checker
openai/clip-vit-large-patch14 CLIP score + similarity
Stages (STAGES env var, comma-separated; default: all, in this order):
snippet run the model card's usage snippet verbatim
checks LoRA-scale-0 == base, unload restores base, same seed == same image
loss validation loss for base + every checkpoint (fixed noise)
metrics select a checkpoint on half the validation set (KID), then report
KID / FID / CLIP score on the other half, base vs selected (+ final)
grids the 50 eval prompts, base vs selected (+ final), same seeds
sweep LoRA strength 0 / 0.5 / 1.0 / 1.5 on a few prompts
progression base + every checkpoint on a few prompts
safety SD safety-checker flag rate on the report-half renders and real images
memorization nearest training image (CLIP cosine) for each report-half render
SMOKE=1 shrinks every stage to a few examples, to catch bugs cheaply.
Each stage is independent: a failure is recorded in status.json and the next
stage still runs.
"""
from __future__ import annotations
import hashlib
import io
import json
import math
import os
import re
import time
import traceback
from pathlib import Path
import numpy as np
import torch
import torch.nn.functional as F
from PIL import Image, ImageDraw, ImageFont
# ---------------------------------------------------------------------------
# configuration
# ---------------------------------------------------------------------------
MODEL_REPO = "whosouravsharma/diffusiondb-sd15-lora"
DATASET_REPO = "whosouravsharma/text-to-image-diffusiondb-2M"
DATASET_REVISION = "v2-clean"
LATENTS_REVISION = "latents-512"
BASE_MODEL = "stable-diffusion-v1-5/stable-diffusion-v1-5"
CLIP_MODEL = "openai/clip-vit-large-patch14"
CHECKPOINT_DIR = "checkpoints"
ADAPTER = "diffusiondb"
FINAL = "checkpoint-4240"
# Same render settings as training/sample_job.py, so grids line up with it.
SEED = 42
STEPS = 30
GUIDANCE = 7.5
RENDER_BATCH = int(os.environ.get("RENDER_BATCH", 4))
LOSS_BATCH = 8
SMOKE = os.environ.get("SMOKE") == "1"
RUN_NAME = os.environ.get("RUN_NAME") or time.strftime("%Y%m%dT%H%M%S") + ("-smoke" if SMOKE else "")
OUT = Path(os.environ.get("OUT_ROOT", "/out/eval")) / RUN_NAME
ALL_STAGES = ["snippet", "checks", "loss", "metrics", "grids", "sweep",
"progression", "safety", "memorization"]
STAGES = [s.strip() for s in os.environ.get("STAGES", ",".join(ALL_STAGES)).split(",") if s.strip()]
CANDIDATES = os.environ.get(
"CANDIDATES", "checkpoint-2000,checkpoint-3000,checkpoint-4240").split(",")
SWEEP_SCALES = [0.0, 0.5, 1.0, 1.5]
# Indices into eval_prompts.json.
SWEEP_PROMPTS = [int(i) for i in os.environ.get("SWEEP_PROMPTS", "2,6,8,0").split(",")]
PROGRESSION_PROMPTS = [int(i) for i in os.environ.get("PROGRESSION_PROMPTS", "2,6,8").split(",")]
# Smoke-test limits.
LIMIT_LOSS = 32 if SMOKE else None
LIMIT_HALF = 16 if SMOKE else None
LIMIT_GRID = 4 if SMOKE else None
LIMIT_TRAIN = 256 if SMOKE else None
LIMIT_CKPTS = 2 if SMOKE else None
DEVICE = "cuda"
DTYPE = torch.float16
STATUS: dict = {"run": RUN_NAME, "smoke": SMOKE, "stages": {}, "started": time.time()}
STATE: dict = {} # shared between stages: selected checkpoint, report renders, ...
def log(message: str) -> None:
print(f"[{time.strftime('%H:%M:%S')}] {message}", flush=True)
def write_json(path: Path, payload) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(json.dumps(payload, indent=2, default=str))
def save_status() -> None:
write_json(OUT / "status.json", STATUS)
# ---------------------------------------------------------------------------
# inputs
# ---------------------------------------------------------------------------
def list_checkpoints() -> list[str]:
from huggingface_hub import HfApi
names = {
f.split("/")[1] for f in HfApi().list_repo_files(MODEL_REPO)
if f.startswith(f"{CHECKPOINT_DIR}/checkpoint-")
}
ordered = sorted(names, key=lambda n: int(n.split("-")[1]))
return ordered[:LIMIT_CKPTS] if LIMIT_CKPTS else ordered
def eval_prompts() -> list[str]:
from huggingface_hub import hf_hub_download
path = hf_hub_download(DATASET_REPO, "eval_prompts.json",
repo_type="dataset", revision=DATASET_REVISION)
return json.loads(Path(path).read_text())
def center_crop(image: Image.Image, size: int = 512) -> Image.Image:
"""Same view as training: resize the short side, crop the centre square."""
image = image.convert("RGB")
width, height = image.size
scale = size / min(width, height)
image = image.resize((max(size, round(width * scale)), max(size, round(height * scale))),
Image.BICUBIC)
width, height = image.size
left, top = (width - size) // 2, (height - size) // 2
return image.crop((left, top, left + size, top + size))
def normalize_prompt(prompt: str) -> str:
return re.sub(r"[^a-z0-9]+", " ", prompt.lower()).strip()
def validation_halves():
"""Split the 1,000 validation rows into a selection half and a report half.
Split by prompt hash so both images of a prompt group land on the same
side. The halves are used for different purposes: choosing a checkpoint
on one and reporting numbers on the other keeps the reported numbers
free of selection bias.
"""
from datasets import load_dataset
data = load_dataset(DATASET_REPO, split="validation", revision=DATASET_REVISION)
halves = {"select": [], "report": []}
for index, row in enumerate(data):
key = normalize_prompt(row["prompt"])
half = "select" if int(hashlib.sha1(key.encode()).hexdigest(), 16) % 2 == 0 else "report"
halves[half].append({
"index": index,
"prompt": row["prompt"],
"image": center_crop(row["image"]),
"seed": SEED + index, # same seed for every model on this row
})
if LIMIT_HALF:
halves = {k: v[:LIMIT_HALF] for k, v in halves.items()}
log(f"validation halves: select {len(halves['select'])}, report {len(halves['report'])}")
return halves
# ---------------------------------------------------------------------------
# pipeline
# ---------------------------------------------------------------------------
class Runner:
"""One SD 1.5 pipeline; LoRA checkpoints are swapped in and out of it."""
def __init__(self) -> None:
from diffusers import StableDiffusionPipeline
self.pipe = StableDiffusionPipeline.from_pretrained(
BASE_MODEL, torch_dtype=DTYPE, variant="fp16", use_safetensors=True,
safety_checker=None, requires_safety_checker=False,
).to(DEVICE)
self.pipe.set_progress_bar_config(disable=True)
self.current: str | None = None
def use(self, checkpoint: str | None, scale: float = 1.0) -> None:
if checkpoint in (None, "base"):
if self.current:
self.pipe.unload_lora_weights()
self.current = None
return
if self.current != checkpoint:
if self.current:
self.pipe.unload_lora_weights()
self.pipe.load_lora_weights(
MODEL_REPO, subfolder=f"{CHECKPOINT_DIR}/{checkpoint}",
weight_name="pytorch_lora_weights.safetensors", adapter_name=ADAPTER,
)
self.current = checkpoint
self.pipe.set_adapters([ADAPTER], adapter_weights=[float(scale)])
@torch.no_grad()
def render(self, prompts: list[str], seeds: list[int]) -> list[Image.Image]:
images: list[Image.Image] = []
for start in range(0, len(prompts), RENDER_BATCH):
batch = prompts[start:start + RENDER_BATCH]
generators = [torch.Generator(DEVICE).manual_seed(int(s))
for s in seeds[start:start + RENDER_BATCH]]
images += self.pipe(batch, num_inference_steps=STEPS, guidance_scale=GUIDANCE,
generator=generators).images
return images
_RUNNER: Runner | None = None
def runner() -> Runner:
global _RUNNER
if _RUNNER is None:
log(f"loading {BASE_MODEL}")
_RUNNER = Runner()
return _RUNNER
def render_model(model: str, prompts: list[str], seeds: list[int], scale: float = 1.0):
run = runner()
run.use(model, scale)
started = time.time()
images = run.render(prompts, seeds)
log(f" rendered {len(images)} with {model} (scale {scale}) "
f"in {time.time() - started:.0f}s")
return images
def selected() -> str:
return STATE.get("selected", FINAL)
def report_models() -> list[str]:
models = ["base", selected()]
if selected() != FINAL:
models.append(FINAL)
return models
# ---------------------------------------------------------------------------
# image helpers
# ---------------------------------------------------------------------------
def font(size: int = 16):
try:
return ImageFont.load_default(size=size)
except TypeError:
return ImageFont.load_default()
def labelled_grid(rows: list[list[Image.Image]], column_labels: list[str],
row_labels: list[str] | None = None, thumb: int = 256) -> Image.Image:
top, left, pad = 32, (220 if row_labels else 0), 6
width = left + len(column_labels) * (thumb + pad) + pad
height = top + len(rows) * (thumb + pad) + pad
sheet = Image.new("RGB", (width, height), (252, 252, 251))
draw = ImageDraw.Draw(sheet)
for c, label in enumerate(column_labels):
draw.text((left + pad + c * (thumb + pad) + 4, 8), label, fill=(11, 11, 11), font=font(16))
for r, row in enumerate(rows):
y = top + pad + r * (thumb + pad)
if row_labels:
text = row_labels[r]
lines = [text[i:i + 26] for i in range(0, min(len(text), 26 * 6), 26)]
draw.multiline_text((8, y + 4), "\n".join(lines), fill=(82, 81, 78), font=font(13))
for c, image in enumerate(row):
sheet.paste(image.resize((thumb, thumb), Image.LANCZOS), (left + pad + c * (thumb + pad), y))
return sheet
def contact_sheet(images: list[Image.Image], columns: int = 5, thumb: int = 320) -> Image.Image:
rows = (len(images) + columns - 1) // columns
sheet = Image.new("RGB", (columns * thumb, rows * thumb), (18, 18, 22))
for i, image in enumerate(images):
sheet.paste(image.resize((thumb, thumb), Image.LANCZOS), ((i % columns) * thumb, (i // columns) * thumb))
return sheet
def to_tensor(images: list[Image.Image]) -> torch.Tensor:
"""PIL list -> float tensor in [0, 1], shape (N, 3, H, W)."""
array = np.stack([np.asarray(im.convert("RGB"), dtype=np.uint8) for im in images])
return torch.from_numpy(array).permute(0, 3, 1, 2).float() / 255.0
def max_pixel_diff(a: Image.Image, b: Image.Image) -> tuple[int, float]:
x = np.asarray(a, dtype=np.int16)
y = np.asarray(b, dtype=np.int16)
d = np.abs(x - y)
return int(d.max()), float(d.mean())
# ---------------------------------------------------------------------------
# CLIP
# ---------------------------------------------------------------------------
_CLIP = None
def clip():
global _CLIP
if _CLIP is None:
from transformers import CLIPModel, CLIPProcessor
_CLIP = (CLIPModel.from_pretrained(CLIP_MODEL, torch_dtype=DTYPE).to(DEVICE).eval(),
CLIPProcessor.from_pretrained(CLIP_MODEL))
return _CLIP
@torch.no_grad()
def clip_image_embeddings(images: list[Image.Image], batch: int = 64) -> torch.Tensor:
model, processor = clip()
out = []
for start in range(0, len(images), batch):
pixels = processor(images=images[start:start + batch], return_tensors="pt")["pixel_values"]
emb = model.get_image_features(pixel_values=pixels.to(DEVICE, DTYPE))
out.append(F.normalize(emb.float(), dim=-1))
return torch.cat(out)
@torch.no_grad()
def clip_scores(images: list[Image.Image], prompts: list[str], batch: int = 64) -> np.ndarray:
"""CLIP score as in torchmetrics: 100 * max(cos(image, text), 0).
Prompts are truncated to CLIP's 77-token limit, as the SD text encoder does.
"""
model, processor = clip()
image_emb = clip_image_embeddings(images, batch)
scores = []
for start in range(0, len(prompts), batch):
tokens = processor(text=prompts[start:start + batch], return_tensors="pt",
padding=True, truncation=True, max_length=77)
text = model.get_text_features(input_ids=tokens["input_ids"].to(DEVICE),
attention_mask=tokens["attention_mask"].to(DEVICE))
text = F.normalize(text.float(), dim=-1)
cos = (image_emb[start:start + batch] * text).sum(-1)
scores.append((100 * cos.clamp(min=0)).cpu().numpy())
return np.concatenate(scores)
# ---------------------------------------------------------------------------
# stages
# ---------------------------------------------------------------------------
def stage_snippet() -> dict:
"""The model card's "How to get started" code, verbatim (plus a save path)."""
import diffusers
from diffusers import StableDiffusionPipeline
pipe = StableDiffusionPipeline.from_pretrained(
"stable-diffusion-v1-5/stable-diffusion-v1-5",
torch_dtype=torch.float16,
variant="fp16",
).to("cuda")
pipe.load_lora_weights(
"whosouravsharma/diffusiondb-sd15-lora",
subfolder="checkpoints/checkpoint-4240",
weight_name="pytorch_lora_weights.safetensors",
adapter_name="diffusiondb",
)
pipe.set_adapters(["diffusiondb"], adapter_weights=[1.0]) # 0.0 = plain SD 1.5
image = pipe(
"a steampunk owl inside a glass jar, intricate detail",
num_inference_steps=25,
guidance_scale=7.5,
generator=torch.Generator("cuda").manual_seed(42),
).images[0]
path = OUT / "snippet" / "owl.png"
path.parent.mkdir(parents=True, exist_ok=True)
image.save(path)
del pipe
torch.cuda.empty_cache()
return {"passed": True, "diffusers": diffusers.__version__, "torch": torch.__version__,
"image": str(path.relative_to(OUT))}
def stage_checks() -> dict:
prompt = "a steampunk owl inside a glass jar, intricate detail"
run = runner()
run.use("base")
base = run.render([prompt], [SEED])[0]
run.use(FINAL, 0.0)
scale_zero = run.render([prompt], [SEED])[0]
run.use(FINAL, 1.0)
lora_a = run.render([prompt], [SEED])[0]
lora_b = run.render([prompt], [SEED])[0]
run.use("base")
unloaded = run.render([prompt], [SEED])[0]
folder = OUT / "checks"
folder.mkdir(parents=True, exist_ok=True)
for name, image in [("base", base), ("scale0", scale_zero), ("lora_a", lora_a),
("lora_b", lora_b), ("after_unload", unloaded)]:
image.save(folder / f"{name}.png")
results = {}
for name, (x, y) in {
"scale0_equals_base": (base, scale_zero),
"unload_restores_base": (base, unloaded),
"same_seed_is_deterministic": (lora_a, lora_b),
"lora_changes_output": (base, lora_a),
}.items():
max_diff, mean_diff = max_pixel_diff(x, y)
expected_equal = name != "lora_changes_output"
passed = (max_diff <= 1) if expected_equal else (mean_diff > 1.0)
results[name] = {"max_pixel_diff": max_diff, "mean_pixel_diff": round(mean_diff, 4),
"passed": passed}
results["passed"] = all(r["passed"] for r in results.values() if isinstance(r, dict))
return results
@torch.no_grad()
def stage_loss() -> dict:
from datasets import load_dataset
from diffusers import DDPMScheduler
from huggingface_hub import hf_hub_download
manifest = json.loads(Path(hf_hub_download(
DATASET_REPO, "data/manifest.json", repo_type="dataset", revision=LATENTS_REVISION,
)).read_text())
scaling = manifest["scaling_factor"]
shape = tuple(manifest["latent_shape"])
data = load_dataset(DATASET_REPO, split="validation", revision=LATENTS_REVISION)
if LIMIT_LOSS:
data = data.select(range(LIMIT_LOSS))
def decode(blobs):
return torch.from_numpy(np.stack(
[np.frombuffer(b, dtype=np.float16).reshape(shape) for b in blobs]).copy())
scheduler = DDPMScheduler.from_pretrained(BASE_MODEL, subfolder="scheduler")
run = runner()
pipe = run.pipe
losses = {}
for model in ["base"] + list_checkpoints():
run.use(model, 1.0)
generator = torch.Generator(DEVICE).manual_seed(SEED) # identical noise per model
total, batches = 0.0, 0
for start in range(0, len(data), LOSS_BATCH):
rows = data[start:start + LOSS_BATCH]
mean = decode(rows["latent_mean"]).to(DEVICE, torch.float32)
logvar = decode(rows["latent_logvar"]).to(DEVICE, torch.float32).clamp(-30.0, 20.0)
eps = torch.randn(mean.shape, device=DEVICE, generator=generator)
latents = (mean + torch.exp(0.5 * logvar) * eps) * scaling
noise = torch.randn(latents.shape, device=DEVICE, generator=generator)
steps = torch.randint(0, scheduler.config.num_train_timesteps, (latents.shape[0],),
device=DEVICE, generator=generator)
noisy = scheduler.add_noise(latents, noise, steps)
tokens = pipe.tokenizer(rows["prompt"], padding="max_length", truncation=True,
max_length=pipe.tokenizer.model_max_length,
return_tensors="pt").input_ids.to(DEVICE)
encoded = pipe.text_encoder(tokens)[0]
predicted = pipe.unet(noisy.to(DTYPE), steps, encoded).sample
total += F.mse_loss(predicted.float(), noise.float()).item()
batches += 1
losses[model] = total / max(batches, 1)
log(f" val loss {model}: {losses[model]:.5f}")
write_json(OUT / "loss" / "val_loss.json", {"examples": len(data), "seed": SEED, "loss": losses})
plot_loss(losses, OUT / "loss" / "val_loss.png")
return {"examples": len(data), "loss": losses}
def plot_loss(losses: dict, path: Path) -> None:
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
surface, ink, muted, grid, series = "#fcfcfb", "#0b0b0b", "#52514e", "#e6e5e1", "#2a78d6"
points = [(int(k.split("-")[1]), v) for k, v in losses.items() if k != "base"]
points.sort()
steps = [p[0] for p in points]
values = [p[1] for p in points]
fig, ax = plt.subplots(figsize=(8, 4.2), dpi=150)
fig.patch.set_facecolor(surface)
ax.set_facecolor(surface)
ax.plot(steps, values, color=series, linewidth=2, marker="o", markersize=6,
markeredgecolor=surface, markeredgewidth=2, zorder=3)
if "base" in losses:
ax.axhline(losses["base"], color=muted, linewidth=1.5, linestyle=(0, (4, 3)), zorder=2)
ax.annotate("SD 1.5 base (no LoRA)", xy=(steps[0] if steps else 0, losses["base"]),
xytext=(0, 6), textcoords="offset points", color=muted, fontsize=9)
if points:
ax.annotate(f"{values[-1]:.4f}", xy=(steps[-1], values[-1]), xytext=(6, -12),
textcoords="offset points", color=ink, fontsize=9)
ax.set_title("Validation loss by checkpoint (1,000 held-out images, fixed noise)",
color=ink, fontsize=11, loc="left")
ax.set_xlabel("training step", color=muted, fontsize=9)
ax.set_ylabel("MSE (noise prediction)", color=muted, fontsize=9)
ax.tick_params(colors=muted, labelsize=8)
ax.grid(axis="y", color=grid, linewidth=0.8)
for side in ("top", "right", "left"):
ax.spines[side].set_visible(False)
ax.spines["bottom"].set_color(grid)
fig.tight_layout()
path.parent.mkdir(parents=True, exist_ok=True)
fig.savefig(path, facecolor=surface)
plt.close(fig)
def inception_metrics(real: torch.Tensor):
from torchmetrics.image.fid import FrechetInceptionDistance
from torchmetrics.image.kid import KernelInceptionDistance
n = real.shape[0]
subset = max(2, min(1000, n // 2))
fid = FrechetInceptionDistance(feature=2048, normalize=True, reset_real_features=False).to(DEVICE)
kid = KernelInceptionDistance(subset_size=subset, subsets=100, normalize=True,
reset_real_features=False).to(DEVICE)
for start in range(0, n, 50):
chunk = real[start:start + 50].to(DEVICE)
fid.update(chunk, real=True)
kid.update(chunk, real=True)
return fid, kid
def score_fake(fid, kid, fake: torch.Tensor) -> dict:
fid.reset()
kid.reset()
for start in range(0, fake.shape[0], 50):
chunk = fake[start:start + 50].to(DEVICE)
fid.update(chunk, real=False)
kid.update(chunk, real=False)
kid_mean, kid_std = kid.compute()
return {"fid": float(fid.compute()), "kid": float(kid_mean), "kid_std": float(kid_std)}
def stage_metrics() -> dict:
halves = validation_halves()
result: dict = {"n_select": len(halves["select"]), "n_report": len(halves["report"])}
# --- selection half: pick the checkpoint closest to DiffusionDB (lowest KID)
rows = halves["select"]
prompts, seeds = [r["prompt"] for r in rows], [r["seed"] for r in rows]
fid, kid = inception_metrics(to_tensor([r["image"] for r in rows]))
candidates = CANDIDATES[:LIMIT_CKPTS] if LIMIT_CKPTS else CANDIDATES
selection = {}
for model in candidates:
images = render_model(model, prompts, seeds)
selection[model] = score_fake(fid, kid, to_tensor(images))
selection[model]["clip"] = float(clip_scores(images, prompts).mean())
log(f" select {model}: {selection[model]}")
winner = min(selection, key=lambda m: selection[m]["kid"])
STATE["selected"] = winner
result["selection"] = {"criterion": "lowest KID on the selection half",
"candidates": selection, "selected": winner}
del fid, kid
torch.cuda.empty_cache()
# --- report half: the numbers that go in the model card
rows = halves["report"]
prompts, seeds = [r["prompt"] for r in rows], [r["seed"] for r in rows]
real_images = [r["image"] for r in rows]
fid, kid = inception_metrics(to_tensor(real_images))
report = {"real_images": {"clip": float(clip_scores(real_images, prompts).mean()),
"clip_std": float(clip_scores(real_images, prompts).std())}}
renders = {"real": real_images}
for model in report_models():
images = render_model(model, prompts, seeds)
renders[model] = images
scores = clip_scores(images, prompts)
report[model] = score_fake(fid, kid, to_tensor(images))
report[model].update({"clip": float(scores.mean()), "clip_std": float(scores.std())})
folder = OUT / "renders" / model
folder.mkdir(parents=True, exist_ok=True)
for row, image in zip(rows, images):
image.save(folder / f"{row['index']:04}.jpg", quality=92)
log(f" report {model}: {report[model]}")
STATE["report_rows"] = rows
STATE["renders"] = renders
result["report"] = report
result["notes"] = {
"reference": "real validation images from the report half, centre-cropped to 512",
"kid_subset_size": max(2, min(1000, len(rows) // 2)),
"render": {"steps": STEPS, "guidance": GUIDANCE, "scheduler": "PNDM (pipeline default)",
"seed": "42 + validation row index"},
"clip_model": CLIP_MODEL,
}
write_json(OUT / "metrics" / "metrics.json", result)
return result
def stage_grids() -> dict:
prompts = eval_prompts()
if LIMIT_GRID:
prompts = prompts[:LIMIT_GRID]
seeds = [SEED + i for i in range(len(prompts))]
models = report_models()
outputs = {}
for model in models:
images = render_model(model, prompts, seeds)
outputs[model] = images
folder = OUT / "samples" / model
folder.mkdir(parents=True, exist_ok=True)
for i, image in enumerate(images):
image.save(folder / f"{i:03}.png")
contact_sheet(images).save(folder / "grid.jpg", quality=92)
write_json(folder / "prompts.json", {"checkpoint": model, "steps": STEPS, "guidance": GUIDANCE,
"seed": SEED, "prompts": prompts})
pairs = OUT / "samples" / "pairs"
pairs.mkdir(parents=True, exist_ok=True)
for i in range(len(prompts)):
labelled_grid([[outputs["base"][i], outputs[selected()][i]]],
["SD 1.5 base", f"+ LoRA ({selected()})"], thumb=384).save(
pairs / f"{i:03}.jpg", quality=92)
return {"prompts": len(prompts), "models": models}
def stage_sweep() -> dict:
prompts = eval_prompts()
chosen = [prompts[i] for i in SWEEP_PROMPTS][:1 if SMOKE else None]
rows = []
for p_index, prompt in zip(SWEEP_PROMPTS, chosen):
row = []
for scale in SWEEP_SCALES:
row += render_model(selected(), [prompt], [SEED + p_index], scale)
rows.append(row)
path = OUT / "sweep" / "lora_strength.jpg"
path.parent.mkdir(parents=True, exist_ok=True)
labelled_grid(rows, [f"strength {s}" for s in SWEEP_SCALES], chosen).save(path, quality=92)
return {"checkpoint": selected(), "scales": SWEEP_SCALES, "prompts": chosen}
def stage_progression() -> dict:
prompts = eval_prompts()
chosen = [prompts[i] for i in PROGRESSION_PROMPTS][:1 if SMOKE else None]
models = ["base"] + list_checkpoints()
columns = {m: render_model(m, chosen, [SEED + i for i in PROGRESSION_PROMPTS[:len(chosen)]])
for m in models}
rows = [[columns[m][r] for m in models] for r in range(len(chosen))]
labels = ["base"] + [m.split("-")[1] for m in models[1:]]
path = OUT / "progression" / "checkpoints.jpg"
path.parent.mkdir(parents=True, exist_ok=True)
labelled_grid(rows, labels, chosen, thumb=192).save(path, quality=92)
return {"prompts": chosen, "columns": labels}
def report_renders() -> dict:
"""Report-half renders from the metrics stage, or reloaded from disk."""
if "renders" in STATE:
return STATE["renders"]
raise RuntimeError("safety and memorization need the metrics stage in the same run")
@torch.no_grad()
def stage_safety() -> dict:
from diffusers.pipelines.stable_diffusion.safety_checker import StableDiffusionSafetyChecker
from transformers import CLIPImageProcessor
checker = StableDiffusionSafetyChecker.from_pretrained(BASE_MODEL, subfolder="safety_checker").to(DEVICE).eval()
processor = CLIPImageProcessor.from_pretrained(BASE_MODEL, subfolder="feature_extractor")
result = {}
for name, images in report_renders().items():
flagged = 0
for start in range(0, len(images), 32):
chunk = images[start:start + 32]
pixels = processor(chunk, return_tensors="pt").pixel_values.to(DEVICE)
_, has_nsfw = checker(images=np.zeros((len(chunk), 1, 1, 3)), clip_input=pixels)
flagged += int(sum(bool(x) for x in has_nsfw))
result[name] = {"flagged": flagged, "total": len(images),
"rate": round(flagged / max(len(images), 1), 4)}
log(f" safety {name}: {result[name]}")
del checker
torch.cuda.empty_cache()
write_json(OUT / "safety" / "safety.json", result)
return result
@torch.no_grad()
def stage_memorization() -> dict:
"""Nearest training image for every report-half render, by CLIP cosine.
CLIP similarity is a proxy for copying, not proof either way. Real
validation images (distinct prompts, same style) give the baseline for
what "similar" means in this dataset.
"""
from datasets import load_dataset
renders = report_renders()
names = list(renders)
queries = {n: clip_image_embeddings(renders[n]) for n in names}
best = {n: torch.full((len(renders[n]),), -1.0, device=DEVICE) for n in names}
best_thumb = {n: [None] * len(renders[n]) for n in names}
best_prompt = {n: [None] * len(renders[n]) for n in names}
stream = load_dataset(DATASET_REPO, split="train", revision=DATASET_REVISION, streaming=True)
batch_images, batch_prompts, seen = [], [], 0
def flush():
nonlocal batch_images, batch_prompts
if not batch_images:
return
emb = clip_image_embeddings(batch_images)
for n in names:
sims = queries[n] @ emb.T
top, arg = sims.max(dim=1)
improved = (top > best[n]).nonzero().flatten().tolist()
best[n] = torch.maximum(best[n], top)
for q in improved:
thumb = batch_images[arg[q].item()].resize((256, 256), Image.LANCZOS)
buffer = io.BytesIO()
thumb.save(buffer, format="JPEG", quality=85)
best_thumb[n][q] = buffer.getvalue()
best_prompt[n][q] = batch_prompts[arg[q].item()]
batch_images, batch_prompts = [], []
for row in stream:
batch_images.append(center_crop(row["image"]))
batch_prompts.append(row["prompt"])
seen += 1
if len(batch_images) == 64:
flush()
if LIMIT_TRAIN and seen >= LIMIT_TRAIN:
break
if seen % 2000 == 0:
log(f" memorization: {seen} training images embedded")
flush()
result = {"training_images": seen, "similarity": "CLIP ViT-L/14 image cosine"}
for n in names:
values = best[n].cpu().numpy()
result[n] = {
"mean": round(float(values.mean()), 4), "median": round(float(np.median(values)), 4),
"p95": round(float(np.percentile(values, 95)), 4), "max": round(float(values.max()), 4),
"count_ge_0.90": int((values >= 0.90).sum()), "count_ge_0.95": int((values >= 0.95).sum()),
}
log(f" memorization {n}: {result[n]}")
# The closest render/training pairs for the selected checkpoint, to look at.
model = selected()
order = np.argsort(-best[model].cpu().numpy())[:8]
rows = [[renders[model][q], Image.open(io.BytesIO(best_thumb[model][q]))] for q in order]
labels = [f"cos {best[model][q].item():.3f}" for q in order]
path = OUT / "memorization" / "closest_pairs.jpg"
path.parent.mkdir(parents=True, exist_ok=True)
labelled_grid(rows, ["render", "nearest training image"], labels, thumb=256).save(path, quality=90)
write_json(OUT / "memorization" / "memorization.json", result)
return result
STAGE_FUNCTIONS = {
"snippet": stage_snippet, "checks": stage_checks, "loss": stage_loss,
"metrics": stage_metrics, "grids": stage_grids, "sweep": stage_sweep,
"progression": stage_progression, "safety": stage_safety,
"memorization": stage_memorization,
}
def main() -> None:
if not torch.cuda.is_available():
raise SystemExit("No GPU. Run this on a GPU flavor, e.g. a10g-small.")
if os.environ.get("HF_TOKEN"):
log("NOTE: HF_TOKEN is set. This job needs no token; launch it without --secrets HF_TOKEN.")
OUT.mkdir(parents=True, exist_ok=True)
log(f"run {RUN_NAME} -> {OUT} | stages: {', '.join(STAGES)} | smoke={SMOKE}")
STATUS["config"] = {"base": BASE_MODEL, "model": MODEL_REPO, "dataset": f"{DATASET_REPO}@{DATASET_REVISION}",
"steps": STEPS, "guidance": GUIDANCE, "seed": SEED, "render_batch": RENDER_BATCH,
"candidates": CANDIDATES, "gpu": torch.cuda.get_device_name(0)}
save_status()
for name in STAGES:
started = time.time()
log(f"=== {name}")
try:
result = STAGE_FUNCTIONS[name]()
STATUS["stages"][name] = {"ok": True, "seconds": round(time.time() - started), "result": result}
except Exception as error:
traceback.print_exc()
STATUS["stages"][name] = {"ok": False, "seconds": round(time.time() - started),
"error": f"{type(error).__name__}: {error}"}
save_status()
log(f"=== {name} {'ok' if STATUS['stages'][name]['ok'] else 'FAILED'} "
f"({STATUS['stages'][name]['seconds']}s)")
STATUS["seconds"] = round(time.time() - STATUS["started"])
save_status()
failed = [n for n, s in STATUS["stages"].items() if not s["ok"]]
log(f"done in {STATUS['seconds']}s; failed stages: {failed or 'none'}; results in {OUT}")
if __name__ == "__main__":
main()