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