Instructions to use whosouravsharma/diffusiondb-sd15-lora with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use whosouravsharma/diffusiondb-sd15-lora with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("stable-diffusion-v1-5/stable-diffusion-v1-5", dtype=torch.bfloat16, device_map="cuda") pipe.load_lora_weights("whosouravsharma/diffusiondb-sd15-lora") prompt = "a anthropomorphic lion wizard, diffuse lighting, fantasy, intricate, elegant, highly detailed, lifelike, photorealistic, digital painting, artstation, illustration, concept art, smooth, sharp focus, naturalism, trending on byron's - muse, by greg rutkowski and greg staples" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- Draw Things
- DiffusionBee
Download eval/eval_job.py from whosouravsharma/diffusiondb-sd15-lora: direct link, hf CLI and curl.
- Browser
- Download file 33.9 kB
-
https://huggingface.co/whosouravsharma/diffusiondb-sd15-lora/resolve/main/eval/eval_job.py
- Command line
-
hf download hf://whosouravsharma/diffusiondb-sd15-lora/eval/eval_job.py
-
curl -L -o eval_job.py https://huggingface.co/whosouravsharma/diffusiondb-sd15-lora/resolve/main/eval/eval_job.py
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)]) | |
| 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 | |
| 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) | |
| 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 | |
| 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") | |
| 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 | |
| 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() | |