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
Model card: evaluation results, gallery, figures, eval script
Browse files- .gitattributes +12 -0
- README.md +423 -90
- eval/eval_job.py +826 -0
- images/closest-pairs.jpg +3 -0
- images/fail-010.jpg +3 -0
- images/fail-013.jpg +3 -0
- images/fail-021.jpg +3 -0
- images/lora-strength.jpg +3 -0
- images/pair-002.jpg +3 -0
- images/pair-006.jpg +3 -0
- images/pair-015.jpg +3 -0
- images/pair-033.jpg +3 -0
- images/pair-036.jpg +3 -0
- images/pair-049.jpg +3 -0
- images/progression.jpg +3 -0
- images/val-loss.png +0 -0
.gitattributes
CHANGED
|
@@ -84,3 +84,15 @@ samples/base/047.png filter=lfs diff=lfs merge=lfs -text
|
|
| 84 |
samples/base/048.png filter=lfs diff=lfs merge=lfs -text
|
| 85 |
samples/base/049.png filter=lfs diff=lfs merge=lfs -text
|
| 86 |
samples/base/grid.jpg filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 84 |
samples/base/048.png filter=lfs diff=lfs merge=lfs -text
|
| 85 |
samples/base/049.png filter=lfs diff=lfs merge=lfs -text
|
| 86 |
samples/base/grid.jpg filter=lfs diff=lfs merge=lfs -text
|
| 87 |
+
images/closest-pairs.jpg filter=lfs diff=lfs merge=lfs -text
|
| 88 |
+
images/fail-010.jpg filter=lfs diff=lfs merge=lfs -text
|
| 89 |
+
images/fail-013.jpg filter=lfs diff=lfs merge=lfs -text
|
| 90 |
+
images/fail-021.jpg filter=lfs diff=lfs merge=lfs -text
|
| 91 |
+
images/lora-strength.jpg filter=lfs diff=lfs merge=lfs -text
|
| 92 |
+
images/pair-002.jpg filter=lfs diff=lfs merge=lfs -text
|
| 93 |
+
images/pair-006.jpg filter=lfs diff=lfs merge=lfs -text
|
| 94 |
+
images/pair-015.jpg filter=lfs diff=lfs merge=lfs -text
|
| 95 |
+
images/pair-033.jpg filter=lfs diff=lfs merge=lfs -text
|
| 96 |
+
images/pair-036.jpg filter=lfs diff=lfs merge=lfs -text
|
| 97 |
+
images/pair-049.jpg filter=lfs diff=lfs merge=lfs -text
|
| 98 |
+
images/progression.jpg filter=lfs diff=lfs merge=lfs -text
|
README.md
CHANGED
|
@@ -1,134 +1,467 @@
|
|
| 1 |
---
|
| 2 |
-
base_model:
|
|
|
|
| 3 |
library_name: diffusers
|
| 4 |
-
license: creativeml-openrail-m
|
| 5 |
pipeline_tag: text-to-image
|
|
|
|
|
|
|
|
|
|
| 6 |
tags:
|
| 7 |
- text-to-image
|
| 8 |
- stable-diffusion
|
|
|
|
| 9 |
- lora
|
| 10 |
- diffusers
|
| 11 |
-
|
| 12 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 13 |
---
|
| 14 |
|
| 15 |
# diffusiondb-sd15-lora
|
| 16 |
|
| 17 |
-
A
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
|
| 21 |
-
> **Status: not yet trained.** This repository was created ahead of the first
|
| 22 |
-
> run. Checkpoints and sample renders appear here as training proceeds.
|
| 23 |
|
| 24 |
-
|
| 25 |
|
| 26 |
-
|
| 27 |
-
|
| 28 |
-
toward the **DiffusionDB aesthetic** — the keyword-heavy `artstation /
|
| 29 |
-
intricate / octane render` idiom that its users prompted with. It does not
|
| 30 |
-
push image quality past what SD 1.5 already does, and it is not intended to.
|
| 31 |
-
Judge it on style adherence, not on "is it better than the base model".
|
| 32 |
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
[`whosouravsharma/text-to-image-diffusiondb-2M`](https://huggingface.co/datasets/whosouravsharma/text-to-image-diffusiondb-2M)
|
| 36 |
-
at revision `v2-clean` — 14,598 images from `part_id` 1–20.
|
| 37 |
-
|
| 38 |
-
| split | examples |
|
| 39 |
|---|---|
|
| 40 |
-
|
|
| 41 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 42 |
|
| 43 |
-
|
| 44 |
-
DiffusionDB users re-roll the same prompt at many seeds, so a random row split
|
| 45 |
-
leaks: in an earlier revision, 21% of validation prompts had already been seen
|
| 46 |
-
in training. Every image sharing a prompt now lands on the same side.
|
| 47 |
|
| 48 |
-
|
| 49 |
-
|
| 50 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 51 |
|
| 52 |
-
##
|
| 53 |
|
| 54 |
| | |
|
| 55 |
|---|---|
|
| 56 |
-
|
|
| 57 |
-
|
|
| 58 |
-
|
|
| 59 |
-
|
|
| 60 |
-
|
|
| 61 |
-
|
|
| 62 |
-
|
|
| 63 |
-
| caption dropout | 10%, for classifier-free guidance |
|
| 64 |
-
| precision | fp16 |
|
| 65 |
|
| 66 |
-
|
| 67 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 68 |
|
| 69 |
-
##
|
| 70 |
|
| 71 |
-
|
| 72 |
-
checkpoints/
|
| 73 |
-
checkpoint-<step>/
|
| 74 |
-
pytorch_lora_weights.safetensors the adapter
|
| 75 |
-
optimizer.pt optimizer + scaler state, for resuming
|
| 76 |
-
state.json step, epoch, hyperparameters
|
| 77 |
-
training/ the scripts that produced all of this
|
| 78 |
-
samples/
|
| 79 |
-
base/ vanilla SD 1.5, the comparison baseline
|
| 80 |
-
checkpoint-<step>/
|
| 81 |
-
grid.jpg contact sheet, all eval prompts
|
| 82 |
-
000.png … 049.png individual renders
|
| 83 |
-
prompts.json prompt list, seed, steps, guidance
|
| 84 |
-
```
|
| 85 |
|
| 86 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 87 |
|
| 88 |
-
|
| 89 |
-
rendered at every checkpoint with a fixed per-prompt seed, so differences
|
| 90 |
-
between contact sheets come from the weights rather than from noise.
|
| 91 |
-
`samples/base/` is vanilla SD 1.5 on the same prompts.
|
| 92 |
|
| 93 |
-
|
| 94 |
-
|
| 95 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 96 |
|
| 97 |
-
##
|
| 98 |
|
| 99 |
```python
|
| 100 |
import torch
|
| 101 |
from diffusers import StableDiffusionPipeline
|
| 102 |
|
| 103 |
pipe = StableDiffusionPipeline.from_pretrained(
|
| 104 |
-
"
|
|
|
|
|
|
|
| 105 |
).to("cuda")
|
|
|
|
| 106 |
pipe.load_lora_weights(
|
| 107 |
-
"whosouravsharma/diffusiondb-sd15-lora",
|
|
|
|
|
|
|
|
|
|
| 108 |
)
|
|
|
|
| 109 |
|
| 110 |
image = pipe(
|
| 111 |
"a steampunk owl inside a glass jar, intricate detail",
|
| 112 |
-
num_inference_steps=
|
|
|
|
|
|
|
| 113 |
).images[0]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 114 |
```
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 115 |
|
| 116 |
-
|
| 117 |
-
|
| 118 |
-
- Trained on ~13.6k images — fine-tuning scale, not from-scratch scale.
|
| 119 |
-
- Source images are SD 1.x generations, so artifacts of that model are
|
| 120 |
-
reproduced along with its style.
|
| 121 |
-
- 512×512 centre crop discards roughly 16% of the source image area; ~10% of
|
| 122 |
-
images lose more than 40% of their frame.
|
| 123 |
-
- NSFW filtering relies on DiffusionDB's own classifier scores at a 0.2
|
| 124 |
-
threshold. That classifier is noisy, so the training set is filtered, not
|
| 125 |
-
guaranteed clean.
|
| 126 |
-
- Prompts carry heavy style boilerplate (~22% mention `artstation`), which the
|
| 127 |
-
adapter will have learned as part of the aesthetic.
|
| 128 |
-
|
| 129 |
-
## License & attribution
|
| 130 |
-
|
| 131 |
-
Adapter released under CreativeML OpenRAIL-M, matching the SD 1.5 base model.
|
| 132 |
-
Training data derives from
|
| 133 |
-
[poloclub/diffusiondb](https://huggingface.co/datasets/poloclub/diffusiondb),
|
| 134 |
-
CC0-1.0. Wang et al., 2022, [arXiv:2210.14896](https://arxiv.org/abs/2210.14896).
|
|
|
|
| 1 |
---
|
| 2 |
+
base_model: stable-diffusion-v1-5/stable-diffusion-v1-5
|
| 3 |
+
base_model_relation: adapter
|
| 4 |
library_name: diffusers
|
|
|
|
| 5 |
pipeline_tag: text-to-image
|
| 6 |
+
license: creativeml-openrail-m
|
| 7 |
+
datasets:
|
| 8 |
+
- whosouravsharma/text-to-image-diffusiondb-2M
|
| 9 |
tags:
|
| 10 |
- text-to-image
|
| 11 |
- stable-diffusion
|
| 12 |
+
- stable-diffusion-diffusers
|
| 13 |
- lora
|
| 14 |
- diffusers
|
| 15 |
+
thumbnail: https://huggingface.co/whosouravsharma/diffusiondb-sd15-lora/resolve/main/images/pair-002.jpg
|
| 16 |
+
widget:
|
| 17 |
+
- text: 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
|
| 18 |
+
output:
|
| 19 |
+
url: images/pair-002.jpg
|
| 20 |
+
- text: a painting of a happy frog under the rain wearing a rainy coat by kazuo oga
|
| 21 |
+
output:
|
| 22 |
+
url: images/pair-006.jpg
|
| 23 |
+
- text: mushrooms growing through a human skull in the woods
|
| 24 |
+
output:
|
| 25 |
+
url: images/pair-036.jpg
|
| 26 |
+
- text: robots creating new robots in a factory, oil painting by justin gerard, deviantart, hd, 8 k
|
| 27 |
+
output:
|
| 28 |
+
url: images/pair-015.jpg
|
| 29 |
+
- text: a painting of a silhouette in the water in front of a eclipse!!! with a red sky, by jeffrey smith, noah bradley, peter mohrbacher, behance contest winner, symbolism, darksynth, poster art, apocalypse art, hellish background
|
| 30 |
+
output:
|
| 31 |
+
url: images/pair-033.jpg
|
| 32 |
+
- text: a giant skull with intricate rune carvings and glowing eyes with symmetrically braided lovecraftian tentacles haunting the cosmos by dan mumford, twirling smoke trail, a twisting vortex of dying galaxies, digital art, vivid colors, highly detailed
|
| 33 |
+
output:
|
| 34 |
+
url: images/pair-049.jpg
|
| 35 |
+
model-index:
|
| 36 |
+
- name: diffusiondb-sd15-lora (checkpoint-4240)
|
| 37 |
+
results:
|
| 38 |
+
- task:
|
| 39 |
+
type: text-to-image
|
| 40 |
+
dataset:
|
| 41 |
+
name: DiffusionDB held-out prompts (report split, 506 images)
|
| 42 |
+
type: whosouravsharma/text-to-image-diffusiondb-2M
|
| 43 |
+
split: validation
|
| 44 |
+
revision: v2-clean
|
| 45 |
+
metrics:
|
| 46 |
+
- name: KID (×10³, vs. real held-out images)
|
| 47 |
+
type: kid
|
| 48 |
+
value: 2.88
|
| 49 |
+
- name: FID (vs. real held-out images)
|
| 50 |
+
type: fid
|
| 51 |
+
value: 98.0
|
| 52 |
+
- name: CLIP score (ViT-L/14)
|
| 53 |
+
type: clip_score
|
| 54 |
+
value: 28.98
|
| 55 |
---
|
| 56 |
|
| 57 |
# diffusiondb-sd15-lora
|
| 58 |
|
| 59 |
+
A rank-32 **style LoRA** for Stable Diffusion 1.5, trained on 13,598 prompt–image
|
| 60 |
+
pairs from [DiffusionDB](https://huggingface.co/datasets/poloclub/diffusiondb).
|
| 61 |
+
It gives SD 1.5 a **more saturated, higher-contrast, illustrative look** while
|
| 62 |
+
following prompts just as well as the base model.
|
|
|
|
|
|
|
| 63 |
|
| 64 |
+
<Gallery />
|
| 65 |
|
| 66 |
+
*Each image: plain SD 1.5 on the left, the same prompt and seed with the LoRA on
|
| 67 |
+
the right.*
|
|
|
|
|
|
|
|
|
|
|
|
|
| 68 |
|
| 69 |
+
| | |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 70 |
|---|---|
|
| 71 |
+
| **What it does** | Adds bolder colour, stronger contrast and more dramatic lighting. Painterly prompts come out more like illustration |
|
| 72 |
+
| **Prompt-following** | Unchanged: CLIP score 28.98 vs. 28.94 for base |
|
| 73 |
+
| **Base model** | [`stable-diffusion-v1-5/stable-diffusion-v1-5`](https://huggingface.co/stable-diffusion-v1-5/stable-diffusion-v1-5) |
|
| 74 |
+
| **Recommended checkpoint** | `checkpoints/checkpoint-4240` (25.5 MB) |
|
| 75 |
+
| **Trigger word** | None. It applies to every prompt |
|
| 76 |
+
| **Try it** | [DiffusionDB SD 1.5 LoRA Space](https://huggingface.co/spaces/whosouravsharma/diffusiondb-sd15-lora). The strength slider gives a live before/after |
|
| 77 |
+
| **License** | CreativeML OpenRAIL-M |
|
| 78 |
+
|
| 79 |
+
> **Finding worth knowing.** The adapter was trained on DiffusionDB in order to
|
| 80 |
+
> reproduce DiffusionDB's look. It doesn't. DiffusionDB is itself Stable
|
| 81 |
+
> Diffusion 1.x output, and plain SD 1.5 already matches it statistically
|
| 82 |
+
> (KID ≈ 0). The LoRA instead adds a distinct style of its own, and measurably
|
| 83 |
+
> moves outputs *away* from DiffusionDB. See [Evaluation](#evaluation).
|
| 84 |
|
| 85 |
+
## Table of contents
|
|
|
|
|
|
|
|
|
|
| 86 |
|
| 87 |
+
- [Model details](#model-details)
|
| 88 |
+
- [Uses](#uses)
|
| 89 |
+
- [How to get started](#how-to-get-started)
|
| 90 |
+
- [Evaluation](#evaluation)
|
| 91 |
+
- [Bias, risks, and limitations](#bias-risks-and-limitations)
|
| 92 |
+
- [Training details](#training-details)
|
| 93 |
+
- [Environmental impact](#environmental-impact)
|
| 94 |
+
- [Technical specifications](#technical-specifications)
|
| 95 |
+
- [Reproduce](#reproduce)
|
| 96 |
+
- [Citation](#citation)
|
| 97 |
|
| 98 |
+
## Model details
|
| 99 |
|
| 100 |
| | |
|
| 101 |
|---|---|
|
| 102 |
+
| **Developed by** | [whosouravsharma](https://huggingface.co/whosouravsharma) |
|
| 103 |
+
| **Model type** | LoRA adapter for a latent diffusion text-to-image model |
|
| 104 |
+
| **Base model** | `stable-diffusion-v1-5/stable-diffusion-v1-5` (trained under its former name, `runwayml/stable-diffusion-v1-5`) |
|
| 105 |
+
| **Adapter** | LoRA, rank 32, alpha 32, on the UNet attention projections (`to_q`, `to_k`, `to_v`, `to_out.0`) |
|
| 106 |
+
| **Trainable parameters** | About 6.4M, against a UNet of about 860M |
|
| 107 |
+
| **Language** | English prompts |
|
| 108 |
+
| **License** | CreativeML OpenRAIL-M, the same as the base model |
|
|
|
|
|
|
|
| 109 |
|
| 110 |
+
| Resource | Link |
|
| 111 |
+
|---|---|
|
| 112 |
+
| Demo | [Space: diffusiondb-sd15-lora](https://huggingface.co/spaces/whosouravsharma/diffusiondb-sd15-lora) |
|
| 113 |
+
| Inference backend | [Space: diffusiondb-sd15-lora-inference](https://huggingface.co/spaces/whosouravsharma/diffusiondb-sd15-lora-inference) |
|
| 114 |
+
| Training data | [whosouravsharma/text-to-image-diffusiondb-2M](https://huggingface.co/datasets/whosouravsharma/text-to-image-diffusiondb-2M) |
|
| 115 |
+
| Training and evaluation code | [`training/`](./training) and [`eval/`](./eval) in this repository |
|
| 116 |
|
| 117 |
+
## Uses
|
| 118 |
|
| 119 |
+
### Direct use
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 120 |
|
| 121 |
+
- Giving SD 1.5 images a punchier, more saturated, illustrative finish, at an
|
| 122 |
+
adjustable strength.
|
| 123 |
+
- Studying how a fine-tune differs from its base. At strength 0 the output is
|
| 124 |
+
pixel-identical to plain SD 1.5 at the same seed, so every difference comes
|
| 125 |
+
from the adapter.
|
| 126 |
+
- As a documented example of a LoRA fine-tune evaluated end to end: a
|
| 127 |
+
leak-free split, a fixed-seed evaluation set, distribution metrics, a safety
|
| 128 |
+
check and a memorization check.
|
| 129 |
|
| 130 |
+
### Out-of-scope use
|
|
|
|
|
|
|
|
|
|
| 131 |
|
| 132 |
+
- Any use that the CreativeML OpenRAIL-M license prohibits.
|
| 133 |
+
- Reproducing the DiffusionDB distribution. Plain SD 1.5 already does that
|
| 134 |
+
better (see [Evaluation](#evaluation)).
|
| 135 |
+
- Muted, pastel or low-contrast palettes. The adapter tends to override them
|
| 136 |
+
(see [Failure cases](#failure-cases)).
|
| 137 |
+
- Resolutions other than 512×512.
|
| 138 |
+
- Photorealistic images of real, identifiable people.
|
| 139 |
+
- Any setting where unfiltered output reaches users with no moderation step.
|
| 140 |
|
| 141 |
+
## How to get started
|
| 142 |
|
| 143 |
```python
|
| 144 |
import torch
|
| 145 |
from diffusers import StableDiffusionPipeline
|
| 146 |
|
| 147 |
pipe = StableDiffusionPipeline.from_pretrained(
|
| 148 |
+
"stable-diffusion-v1-5/stable-diffusion-v1-5",
|
| 149 |
+
torch_dtype=torch.float16,
|
| 150 |
+
variant="fp16",
|
| 151 |
).to("cuda")
|
| 152 |
+
|
| 153 |
pipe.load_lora_weights(
|
| 154 |
+
"whosouravsharma/diffusiondb-sd15-lora",
|
| 155 |
+
subfolder="checkpoints/checkpoint-4240",
|
| 156 |
+
weight_name="pytorch_lora_weights.safetensors",
|
| 157 |
+
adapter_name="diffusiondb",
|
| 158 |
)
|
| 159 |
+
pipe.set_adapters(["diffusiondb"], adapter_weights=[1.0]) # 0.0 = plain SD 1.5
|
| 160 |
|
| 161 |
image = pipe(
|
| 162 |
"a steampunk owl inside a glass jar, intricate detail",
|
| 163 |
+
num_inference_steps=25,
|
| 164 |
+
guidance_scale=7.5,
|
| 165 |
+
generator=torch.Generator("cuda").manual_seed(42),
|
| 166 |
).images[0]
|
| 167 |
+
image.save("owl.png")
|
| 168 |
+
```
|
| 169 |
+
|
| 170 |
+
This snippet is run verbatim as part of the evaluation (diffusers 0.31.0,
|
| 171 |
+
torch 2.4.0).
|
| 172 |
+
|
| 173 |
+
**Strength.** 1.0 gives the full effect. Around 0.5 keeps the style while toning
|
| 174 |
+
down the colour when it gets too strong. 1.5 exaggerates the style and starts to change the
|
| 175 |
+
composition.
|
| 176 |
+
|
| 177 |
+

|
| 178 |
+
|
| 179 |
+
*The same seed at strength 0, 0.5, 1.0 and 1.5. Strength 0 is plain SD 1.5.*
|
| 180 |
+
|
| 181 |
+
## Evaluation
|
| 182 |
+
|
| 183 |
+
Everything below comes from one evaluation run on a single A10G (1 h 53 min).
|
| 184 |
+
The full outputs, including every render and metric, can be reproduced with
|
| 185 |
+
[`eval/eval_job.py`](./eval/eval_job.py).
|
| 186 |
+
|
| 187 |
+
### Protocol
|
| 188 |
+
|
| 189 |
+
- **Held-out data.** The 1,000 validation images were split by prompt group,
|
| 190 |
+
so the model never trained on these prompts. Those 1,000 were then split
|
| 191 |
+
again by prompt hash: a **selection half** (494 images) to choose a
|
| 192 |
+
checkpoint, and a **report half** (506 images) for the numbers below. The
|
| 193 |
+
reported numbers are therefore free of selection bias.
|
| 194 |
+
- **Rendering.** For every model: 30 steps, guidance 7.5, the pipeline's
|
| 195 |
+
default PNDM scheduler, and seed 42 + the validation row index. Every model
|
| 196 |
+
sees the same prompts and seeds.
|
| 197 |
+
- **Metrics.**
|
| 198 |
+
- **KID and FID** against the real held-out images (centre-cropped to 512,
|
| 199 |
+
the same view as training). Lower means closer to DiffusionDB. KID is the
|
| 200 |
+
primary metric because FID is biased at this sample size.
|
| 201 |
+
- **CLIP score** (ViT-L/14) of each image against its own prompt, which
|
| 202 |
+
measures prompt-following.
|
| 203 |
+
|
| 204 |
+
### Results
|
| 205 |
+
|
| 206 |
+
| Model (report half, n = 506) | KID ×10³ ↓ | FID ↓ | CLIP score |
|
| 207 |
+
|---|---|---|---|
|
| 208 |
+
| SD 1.5 base | **0.14 ± 0.32** | **95.4** | 28.94 |
|
| 209 |
+
| + LoRA `checkpoint-2000` | 2.87 ± 0.64 | 98.7 | 28.79 |
|
| 210 |
+
| + LoRA `checkpoint-4240` | 2.88 ± 0.61 | 98.0 | 28.98 |
|
| 211 |
+
| Real held-out images (reference) | — | — | 28.95 |
|
| 212 |
+
|
| 213 |
+
*KID is the mean ± standard deviation over 100 subsets of 253 images each.
|
| 214 |
+
CLIP scores have a standard deviation of about 3.9 per image, which makes the
|
| 215 |
+
standard error of each mean about 0.17.*
|
| 216 |
+
|
| 217 |
+
**What the numbers say**
|
| 218 |
+
|
| 219 |
+
1. **Base SD 1.5 already matches DiffusionDB.** Its KID is indistinguishable
|
| 220 |
+
from zero. That's expected, because DiffusionDB was generated with SD 1.x.
|
| 221 |
+
2. **The LoRA moves outputs away from DiffusionDB.** KID rises by about
|
| 222 |
+
2.7×10⁻³, roughly four times the combined spread of the two KID estimates. This is the style shift visible in
|
| 223 |
+
the gallery, and it is not an approach to the training distribution.
|
| 224 |
+
3. **Prompt-following is unchanged.** All CLIP scores are within noise of each
|
| 225 |
+
other and of the real images.
|
| 226 |
+
4. **Checkpoints from step 2,000 onward are statistically identical.** On the
|
| 227 |
+
selection half, step 2,000 won narrowly (KID 2.50 vs. 2.62 for step 4,240
|
| 228 |
+
and 2.90 for step 3,000, all ± about 0.6). On the report half the two are
|
| 229 |
+
tied. `checkpoint-4240` stays the recommended checkpoint.
|
| 230 |
+
|
| 231 |
+
**A likely cause, not yet tested.** The added saturation looks like
|
| 232 |
+
classifier-free guidance overshooting. The adapter may have sharpened the gap
|
| 233 |
+
between the conditional and unconditional predictions, so guidance 7.5
|
| 234 |
+
behaves like a higher value. Sweeping the guidance scale with the adapter would
|
| 235 |
+
confirm or rule this out.
|
| 236 |
+
|
| 237 |
+
### Validation loss
|
| 238 |
+
|
| 239 |
+

|
| 240 |
+
|
| 241 |
+
This is noise-prediction MSE on the 1,000 held-out images, with identical noise
|
| 242 |
+
and timesteps for every model. The loss falls 0.9% below the base model
|
| 243 |
+
(0.1496 → 0.1482), mostly in the first 2,000 steps, and is flat after about
|
| 244 |
+
3,000. There's no sign of divergence or overfitting. As is common for diffusion
|
| 245 |
+
fine-tunes, a lower denoising loss did not translate into samples closer to the
|
| 246 |
+
data (see KID above).
|
| 247 |
+
|
| 248 |
+
### How the style develops during training
|
| 249 |
+
|
| 250 |
+

|
| 251 |
+
|
| 252 |
+
*The same prompts and seeds at every checkpoint. Most of the style is in place
|
| 253 |
+
by step 1,500. Later checkpoints refine it rather than change it.*
|
| 254 |
+
|
| 255 |
+
### Failure cases
|
| 256 |
+
|
| 257 |
+

|
| 258 |
+
*"shattering of the moon's surface, digital art, illustration". The detailed
|
| 259 |
+
scene collapses into a flat, sparse composition.*
|
| 260 |
+
|
| 261 |
+

|
| 262 |
+
*"glass vodka bottle by shusei nagaoka, kaws, david rudnick, airbrush on
|
| 263 |
+
canvas, **pastel colors**, cell-shaded, 8 k". The adapter's saturation
|
| 264 |
+
overrides the requested pastel palette.*
|
| 265 |
+
|
| 266 |
+

|
| 267 |
+
*"long distance shot of a tiny cute polar bear on a tiny iceberg in the middle
|
| 268 |
+
of the ocean, sunset, atmospheric, hazy". The animal's anatomy degrades, and
|
| 269 |
+
the hazy mood is lost.*
|
| 270 |
+
|
| 271 |
+
### Correctness checks
|
| 272 |
+
|
| 273 |
+
| Check | Result |
|
| 274 |
+
|---|---|
|
| 275 |
+
| Strength 0 equals plain SD 1.5 (same seed) | ✅ Max pixel difference 0 |
|
| 276 |
+
| Unloading the adapter restores plain SD 1.5 | ✅ Max pixel difference 0 |
|
| 277 |
+
| Same seed gives the same image | ✅ Max pixel difference 0 |
|
| 278 |
+
| The adapter changes the output | ✅ Mean pixel difference 41.7 / 255 |
|
| 279 |
+
| The usage snippet above runs as written | ✅ |
|
| 280 |
+
|
| 281 |
+
## Bias, risks, and limitations
|
| 282 |
+
|
| 283 |
+
### Safety
|
| 284 |
+
|
| 285 |
+
The SD safety checker was run afterwards over the report-half renders; it is
|
| 286 |
+
disabled during generation in the demo.
|
| 287 |
+
|
| 288 |
+
| Images (n = 506) | Flagged |
|
| 289 |
+
|---|---|
|
| 290 |
+
| Real held-out images | 7 (1.4%) |
|
| 291 |
+
| SD 1.5 base | 13 (2.6%) |
|
| 292 |
+
| + LoRA `checkpoint-2000` | 10 (2.0%) |
|
| 293 |
+
| + LoRA `checkpoint-4240` | 13 (2.6%) |
|
| 294 |
+
|
| 295 |
+
The adapter doesn't raise the flag rate over the base model. The counts are
|
| 296 |
+
small, so treat these as rough rates. The training data was filtered with
|
| 297 |
+
DiffusionDB's own NSFW scores (below 0.2), but that classifier is noisy, so the
|
| 298 |
+
filtering is not a guarantee.
|
| 299 |
+
|
| 300 |
+
### Memorization
|
| 301 |
+
|
| 302 |
+
For each report-half render, we found the most similar of the 13,598 training
|
| 303 |
+
images by CLIP ViT-L/14 image-embedding cosine. Real held-out images, which
|
| 304 |
+
have different prompts but the same style, give the baseline for how similar
|
| 305 |
+
two unrelated DiffusionDB images can be.
|
| 306 |
+
|
| 307 |
+
| Images (n = 506) | Median nearest-neighbour cosine | ≥ 0.95 | Max |
|
| 308 |
+
|---|---|---|---|
|
| 309 |
+
| Real held-out images | 0.845 | 12 | 0.995 |
|
| 310 |
+
| SD 1.5 base | 0.833 | 3 | 0.971 |
|
| 311 |
+
| + LoRA `checkpoint-2000` | 0.829 | 2 | 0.959 |
|
| 312 |
+
| + LoRA `checkpoint-4240` | 0.833 | 2 | 0.961 |
|
| 313 |
+
|
| 314 |
+
The LoRA's renders are no closer to the training images than the base model's.
|
| 315 |
+
|
| 316 |
+

|
| 317 |
+
|
| 318 |
+
*The 8 closest pairs for `checkpoint-2000`. They share a subject and style, not
|
| 319 |
+
a composition, so none is a copy. CLIP similarity is a proxy for copying, not
|
| 320 |
+
proof either way.*
|
| 321 |
+
|
| 322 |
+
### Limitations
|
| 323 |
+
|
| 324 |
+
- **The style is fixed.** Saturation and contrast go up on every prompt,
|
| 325 |
+
including ones that ask for muted or pastel colours.
|
| 326 |
+
- **Occasional composition loss.** Some detailed scenes are simplified (see
|
| 327 |
+
[Failure cases](#failure-cases)).
|
| 328 |
+
- **Small training set.** 13,598 images is enough to fine-tune, not to train
|
| 329 |
+
from scratch.
|
| 330 |
+
- **It inherits SD 1.x flaws.** The training images were themselves SD 1.x
|
| 331 |
+
outputs, including their artifacts.
|
| 332 |
+
- **Centre crop only.** There was no aspect-ratio bucketing, and roughly half
|
| 333 |
+
the source images aren't square, so their edges were lost in training.
|
| 334 |
+
- **The text encoder was frozen,** so the model understands prompts exactly as
|
| 335 |
+
SD 1.5 does.
|
| 336 |
+
|
| 337 |
+
### Risks and recommendations
|
| 338 |
+
|
| 339 |
+
- **Turn the safety checker back on.** The demo Space disables it. Put a
|
| 340 |
+
safety checker or other moderation step in front of any public-facing use.
|
| 341 |
+
- **It inherits SD 1.5's biases.** SD 1.5 and its LAION training data carry
|
| 342 |
+
social and cultural biases. This adapter doesn't reduce them.
|
| 343 |
+
- **Artist names in prompts.** Many DiffusionDB prompts name living artists,
|
| 344 |
+
and the adapter learned from those prompt–image pairs.
|
| 345 |
+
|
| 346 |
+
## Training details
|
| 347 |
+
|
| 348 |
+
### Training data
|
| 349 |
+
|
| 350 |
+
[`whosouravsharma/text-to-image-diffusiondb-2M`](https://huggingface.co/datasets/whosouravsharma/text-to-image-diffusiondb-2M)
|
| 351 |
+
at revision `v2-clean`: 13,598 train and 1,000 validation images, built from
|
| 352 |
+
parts 1–20 of [`poloclub/diffusiondb`](https://huggingface.co/datasets/poloclub/diffusiondb).
|
| 353 |
+
The validation split is made **by prompt group**: DiffusionDB users often ran
|
| 354 |
+
the same prompt at many seeds, so every image sharing a normalized prompt goes
|
| 355 |
+
to the same side.
|
| 356 |
+
|
| 357 |
+
<details>
|
| 358 |
+
<summary>Filtering, stage by stage (from the dataset's <code>manifest.json</code>)</summary>
|
| 359 |
+
|
| 360 |
+
| Stage | Rows kept | Removed |
|
| 361 |
+
|---|---|---|
|
| 362 |
+
| Source metadata | 20,000 | — |
|
| 363 |
+
| Prompt has at least 4 words | 19,513 | 487 |
|
| 364 |
+
| Short side ≥ 384 px and area ≥ 262,144 px | 19,426 | 87 |
|
| 365 |
+
| Image and prompt NSFW scores below 0.2 | 15,451 | 3,975 |
|
| 366 |
+
| At most 2 images per normalized prompt | 14,606 | 845 |
|
| 367 |
+
| Exact SHA-256 duplicates removed | **14,598** | 8 |
|
| 368 |
+
|
| 369 |
+
</details>
|
| 370 |
+
|
| 371 |
+
### Training procedure
|
| 372 |
+
|
| 373 |
+
Every image was encoded once with
|
| 374 |
+
[`stabilityai/sd-vae-ft-mse`](https://huggingface.co/stabilityai/sd-vae-ft-mse)
|
| 375 |
+
after resizing and a centre crop to 512×512. The results are stored as fp16
|
| 376 |
+
latents: both the posterior mean and its log-variance, so each training step
|
| 377 |
+
samples a fresh latent. The text embeddings were not cached, which keeps 10%
|
| 378 |
+
caption dropout possible for classifier-free guidance.
|
| 379 |
+
|
| 380 |
+
<details>
|
| 381 |
+
<summary>Hyperparameters</summary>
|
| 382 |
+
|
| 383 |
+
| | |
|
| 384 |
+
|---|---|
|
| 385 |
+
| Objective | ε-prediction (MSE on the noise), DDPM noise schedule |
|
| 386 |
+
| Epochs | 10 (424 optimizer steps per epoch, 4,240 in total) |
|
| 387 |
+
| Batch size | 8 per step × 4 gradient-accumulation steps = 32 |
|
| 388 |
+
| Optimizer | AdamW, lr 1e-4, weight decay 1e-2, grad-norm clip 1.0 |
|
| 389 |
+
| Schedule | 500 linear warmup steps, then cosine decay |
|
| 390 |
+
| Caption dropout | 10% |
|
| 391 |
+
| LoRA init | Gaussian |
|
| 392 |
+
| Precision | Frozen weights in fp16; LoRA weights in fp32; forward pass under fp16 autocast |
|
| 393 |
+
| Seed | 42 |
|
| 394 |
+
| Checkpoints | Every 500 steps, plus the final step (4,240) |
|
| 395 |
+
|
| 396 |
+
Each checkpoint has a `state.json` recording the exact values it was trained
|
| 397 |
+
with.
|
| 398 |
+
|
| 399 |
+
</details>
|
| 400 |
+
|
| 401 |
+
## Environmental impact
|
| 402 |
+
|
| 403 |
+
| | |
|
| 404 |
+
|---|---|
|
| 405 |
+
| Hardware | 1× NVIDIA A10G (Hugging Face Jobs, `a10g-small`) |
|
| 406 |
+
| Training | 2.68 h (plus 0.25 h caching latents and 0.04 h baseline samples) |
|
| 407 |
+
| Evaluation | 1.88 h |
|
| 408 |
+
| Total | About 4.9 GPU-hours |
|
| 409 |
+
| Carbon emitted | Not measured |
|
| 410 |
+
|
| 411 |
+
## Technical specifications
|
| 412 |
+
|
| 413 |
+
SD 1.5 latent diffusion: CLIP ViT-L/14 text encoder (frozen), a UNet of about
|
| 414 |
+
860M parameters (frozen, with LoRA on the attention projections), and a VAE
|
| 415 |
+
with 8× downsampling. Every stage runs as a standalone PEP 723 script on
|
| 416 |
+
Hugging Face Jobs. Each job pulls its inputs from the Hub, and training can
|
| 417 |
+
resume from any checkpoint with `RESUME_FROM`.
|
| 418 |
+
|
| 419 |
+
<details>
|
| 420 |
+
<summary>Repository layout</summary>
|
| 421 |
+
|
| 422 |
```
|
| 423 |
+
checkpoints/
|
| 424 |
+
checkpoint-<step>/ step = 500, 1000, … 4000, 4240
|
| 425 |
+
pytorch_lora_weights.safetensors the adapter
|
| 426 |
+
optimizer.pt optimizer and grad-scaler state, for resuming
|
| 427 |
+
state.json step, epoch, hyperparameters
|
| 428 |
+
samples/base/ plain SD 1.5 on the 50 eval prompts
|
| 429 |
+
images/ figures used in this card
|
| 430 |
+
training/ data caching, training and sampling scripts
|
| 431 |
+
eval/eval_job.py the evaluation that produced the numbers above
|
| 432 |
+
```
|
| 433 |
+
|
| 434 |
+
</details>
|
| 435 |
+
|
| 436 |
+
## Reproduce
|
| 437 |
+
|
| 438 |
+
```bash
|
| 439 |
+
# training (run from training/)
|
| 440 |
+
python3 main.py baseline # plain SD 1.5 samples, for comparison
|
| 441 |
+
python3 main.py latents # VAE-encode the dataset once
|
| 442 |
+
python3 main.py train # LoRA fine-tune
|
| 443 |
+
python3 main.py sample checkpoint-4240
|
| 444 |
+
|
| 445 |
+
# evaluation: needs no token, writes only to the mounted bucket
|
| 446 |
+
hf jobs uv run --flavor a10g-small --timeout 4h \
|
| 447 |
+
-v hf://buckets/<you>/<bucket>:/out eval/eval_job.py
|
| 448 |
+
```
|
| 449 |
+
|
| 450 |
+
## Citation
|
| 451 |
+
|
| 452 |
+
Training data comes from DiffusionDB (CC0 1.0):
|
| 453 |
+
|
| 454 |
+
```bibtex
|
| 455 |
+
@article{wangDiffusionDBLargescalePrompt2022,
|
| 456 |
+
title = {DiffusionDB: A Large-Scale Prompt Gallery Dataset for Text-to-Image Generative Models},
|
| 457 |
+
author = {Wang, Zijie J. and Montoya, Evan and Munechika, David and Yang, Haoyang and Hoover, Benjamin and Chau, Duen Horng},
|
| 458 |
+
journal = {arXiv:2210.14896 [cs]},
|
| 459 |
+
year = {2022},
|
| 460 |
+
url = {https://arxiv.org/abs/2210.14896}
|
| 461 |
+
}
|
| 462 |
+
```
|
| 463 |
+
|
| 464 |
+
## Model card contact
|
| 465 |
|
| 466 |
+
Open a discussion in the
|
| 467 |
+
[Community tab](https://huggingface.co/whosouravsharma/diffusiondb-sd15-lora/discussions).
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
eval/eval_job.py
ADDED
|
@@ -0,0 +1,826 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# /// script
|
| 2 |
+
# requires-python = ">=3.10"
|
| 3 |
+
# dependencies = [
|
| 4 |
+
# "torch==2.4.0",
|
| 5 |
+
# "diffusers==0.31.0",
|
| 6 |
+
# "transformers==4.44.2",
|
| 7 |
+
# "peft==0.13.2",
|
| 8 |
+
# "accelerate==0.34.2",
|
| 9 |
+
# "safetensors",
|
| 10 |
+
# "huggingface_hub>=0.24,<1.0",
|
| 11 |
+
# "datasets>=2.19,<4",
|
| 12 |
+
# "torchmetrics[image]>=1.4",
|
| 13 |
+
# "torch-fidelity",
|
| 14 |
+
# "scipy",
|
| 15 |
+
# "numpy<2",
|
| 16 |
+
# "pillow>=10.1",
|
| 17 |
+
# "matplotlib",
|
| 18 |
+
# ]
|
| 19 |
+
# ///
|
| 20 |
+
"""Evaluate the DiffusionDB SD 1.5 LoRA without writing to any repo.
|
| 21 |
+
|
| 22 |
+
hf jobs uv run --flavor a10g-small --timeout 4h \\
|
| 23 |
+
-v hf://buckets/whosouravsharma/jobs-artifacts:/out \\
|
| 24 |
+
eval/eval_job.py
|
| 25 |
+
|
| 26 |
+
Deliberately launched WITHOUT --secrets HF_TOKEN. Every input is public, so
|
| 27 |
+
the job has no credentials that could write to the model repo, the dataset
|
| 28 |
+
repo or either Space. The only writable location is the mounted bucket, and
|
| 29 |
+
everything lands under /out/eval/<RUN_NAME>/.
|
| 30 |
+
|
| 31 |
+
Reads (all public, read-only):
|
| 32 |
+
whosouravsharma/diffusiondb-sd15-lora checkpoints
|
| 33 |
+
whosouravsharma/text-to-image-diffusiondb-2M v2-clean images + prompts,
|
| 34 |
+
latents-512 validation latents
|
| 35 |
+
stable-diffusion-v1-5/stable-diffusion-v1-5 base model, safety checker
|
| 36 |
+
openai/clip-vit-large-patch14 CLIP score + similarity
|
| 37 |
+
|
| 38 |
+
Stages (STAGES env var, comma-separated; default: all, in this order):
|
| 39 |
+
snippet run the model card's usage snippet verbatim
|
| 40 |
+
checks LoRA-scale-0 == base, unload restores base, same seed == same image
|
| 41 |
+
loss validation loss for base + every checkpoint (fixed noise)
|
| 42 |
+
metrics select a checkpoint on half the validation set (KID), then report
|
| 43 |
+
KID / FID / CLIP score on the other half, base vs selected (+ final)
|
| 44 |
+
grids the 50 eval prompts, base vs selected (+ final), same seeds
|
| 45 |
+
sweep LoRA strength 0 / 0.5 / 1.0 / 1.5 on a few prompts
|
| 46 |
+
progression base + every checkpoint on a few prompts
|
| 47 |
+
safety SD safety-checker flag rate on the report-half renders and real images
|
| 48 |
+
memorization nearest training image (CLIP cosine) for each report-half render
|
| 49 |
+
|
| 50 |
+
SMOKE=1 shrinks every stage to a few examples, to catch bugs cheaply.
|
| 51 |
+
Each stage is independent: a failure is recorded in status.json and the next
|
| 52 |
+
stage still runs.
|
| 53 |
+
"""
|
| 54 |
+
|
| 55 |
+
from __future__ import annotations
|
| 56 |
+
|
| 57 |
+
import hashlib
|
| 58 |
+
import io
|
| 59 |
+
import json
|
| 60 |
+
import math
|
| 61 |
+
import os
|
| 62 |
+
import re
|
| 63 |
+
import time
|
| 64 |
+
import traceback
|
| 65 |
+
from pathlib import Path
|
| 66 |
+
|
| 67 |
+
import numpy as np
|
| 68 |
+
import torch
|
| 69 |
+
import torch.nn.functional as F
|
| 70 |
+
from PIL import Image, ImageDraw, ImageFont
|
| 71 |
+
|
| 72 |
+
# ---------------------------------------------------------------------------
|
| 73 |
+
# configuration
|
| 74 |
+
# ---------------------------------------------------------------------------
|
| 75 |
+
|
| 76 |
+
MODEL_REPO = "whosouravsharma/diffusiondb-sd15-lora"
|
| 77 |
+
DATASET_REPO = "whosouravsharma/text-to-image-diffusiondb-2M"
|
| 78 |
+
DATASET_REVISION = "v2-clean"
|
| 79 |
+
LATENTS_REVISION = "latents-512"
|
| 80 |
+
BASE_MODEL = "stable-diffusion-v1-5/stable-diffusion-v1-5"
|
| 81 |
+
CLIP_MODEL = "openai/clip-vit-large-patch14"
|
| 82 |
+
CHECKPOINT_DIR = "checkpoints"
|
| 83 |
+
ADAPTER = "diffusiondb"
|
| 84 |
+
FINAL = "checkpoint-4240"
|
| 85 |
+
|
| 86 |
+
# Same render settings as training/sample_job.py, so grids line up with it.
|
| 87 |
+
SEED = 42
|
| 88 |
+
STEPS = 30
|
| 89 |
+
GUIDANCE = 7.5
|
| 90 |
+
RENDER_BATCH = int(os.environ.get("RENDER_BATCH", 4))
|
| 91 |
+
LOSS_BATCH = 8
|
| 92 |
+
|
| 93 |
+
SMOKE = os.environ.get("SMOKE") == "1"
|
| 94 |
+
RUN_NAME = os.environ.get("RUN_NAME") or time.strftime("%Y%m%dT%H%M%S") + ("-smoke" if SMOKE else "")
|
| 95 |
+
OUT = Path(os.environ.get("OUT_ROOT", "/out/eval")) / RUN_NAME
|
| 96 |
+
|
| 97 |
+
ALL_STAGES = ["snippet", "checks", "loss", "metrics", "grids", "sweep",
|
| 98 |
+
"progression", "safety", "memorization"]
|
| 99 |
+
STAGES = [s.strip() for s in os.environ.get("STAGES", ",".join(ALL_STAGES)).split(",") if s.strip()]
|
| 100 |
+
|
| 101 |
+
CANDIDATES = os.environ.get(
|
| 102 |
+
"CANDIDATES", "checkpoint-2000,checkpoint-3000,checkpoint-4240").split(",")
|
| 103 |
+
SWEEP_SCALES = [0.0, 0.5, 1.0, 1.5]
|
| 104 |
+
# Indices into eval_prompts.json.
|
| 105 |
+
SWEEP_PROMPTS = [int(i) for i in os.environ.get("SWEEP_PROMPTS", "2,6,8,0").split(",")]
|
| 106 |
+
PROGRESSION_PROMPTS = [int(i) for i in os.environ.get("PROGRESSION_PROMPTS", "2,6,8").split(",")]
|
| 107 |
+
|
| 108 |
+
# Smoke-test limits.
|
| 109 |
+
LIMIT_LOSS = 32 if SMOKE else None
|
| 110 |
+
LIMIT_HALF = 16 if SMOKE else None
|
| 111 |
+
LIMIT_GRID = 4 if SMOKE else None
|
| 112 |
+
LIMIT_TRAIN = 256 if SMOKE else None
|
| 113 |
+
LIMIT_CKPTS = 2 if SMOKE else None
|
| 114 |
+
|
| 115 |
+
DEVICE = "cuda"
|
| 116 |
+
DTYPE = torch.float16
|
| 117 |
+
|
| 118 |
+
STATUS: dict = {"run": RUN_NAME, "smoke": SMOKE, "stages": {}, "started": time.time()}
|
| 119 |
+
STATE: dict = {} # shared between stages: selected checkpoint, report renders, ...
|
| 120 |
+
|
| 121 |
+
|
| 122 |
+
def log(message: str) -> None:
|
| 123 |
+
print(f"[{time.strftime('%H:%M:%S')}] {message}", flush=True)
|
| 124 |
+
|
| 125 |
+
|
| 126 |
+
def write_json(path: Path, payload) -> None:
|
| 127 |
+
path.parent.mkdir(parents=True, exist_ok=True)
|
| 128 |
+
path.write_text(json.dumps(payload, indent=2, default=str))
|
| 129 |
+
|
| 130 |
+
|
| 131 |
+
def save_status() -> None:
|
| 132 |
+
write_json(OUT / "status.json", STATUS)
|
| 133 |
+
|
| 134 |
+
|
| 135 |
+
# ---------------------------------------------------------------------------
|
| 136 |
+
# inputs
|
| 137 |
+
# ---------------------------------------------------------------------------
|
| 138 |
+
|
| 139 |
+
def list_checkpoints() -> list[str]:
|
| 140 |
+
from huggingface_hub import HfApi
|
| 141 |
+
|
| 142 |
+
names = {
|
| 143 |
+
f.split("/")[1] for f in HfApi().list_repo_files(MODEL_REPO)
|
| 144 |
+
if f.startswith(f"{CHECKPOINT_DIR}/checkpoint-")
|
| 145 |
+
}
|
| 146 |
+
ordered = sorted(names, key=lambda n: int(n.split("-")[1]))
|
| 147 |
+
return ordered[:LIMIT_CKPTS] if LIMIT_CKPTS else ordered
|
| 148 |
+
|
| 149 |
+
|
| 150 |
+
def eval_prompts() -> list[str]:
|
| 151 |
+
from huggingface_hub import hf_hub_download
|
| 152 |
+
|
| 153 |
+
path = hf_hub_download(DATASET_REPO, "eval_prompts.json",
|
| 154 |
+
repo_type="dataset", revision=DATASET_REVISION)
|
| 155 |
+
return json.loads(Path(path).read_text())
|
| 156 |
+
|
| 157 |
+
|
| 158 |
+
def center_crop(image: Image.Image, size: int = 512) -> Image.Image:
|
| 159 |
+
"""Same view as training: resize the short side, crop the centre square."""
|
| 160 |
+
image = image.convert("RGB")
|
| 161 |
+
width, height = image.size
|
| 162 |
+
scale = size / min(width, height)
|
| 163 |
+
image = image.resize((max(size, round(width * scale)), max(size, round(height * scale))),
|
| 164 |
+
Image.BICUBIC)
|
| 165 |
+
width, height = image.size
|
| 166 |
+
left, top = (width - size) // 2, (height - size) // 2
|
| 167 |
+
return image.crop((left, top, left + size, top + size))
|
| 168 |
+
|
| 169 |
+
|
| 170 |
+
def normalize_prompt(prompt: str) -> str:
|
| 171 |
+
return re.sub(r"[^a-z0-9]+", " ", prompt.lower()).strip()
|
| 172 |
+
|
| 173 |
+
|
| 174 |
+
def validation_halves():
|
| 175 |
+
"""Split the 1,000 validation rows into a selection half and a report half.
|
| 176 |
+
|
| 177 |
+
Split by prompt hash so both images of a prompt group land on the same
|
| 178 |
+
side. The halves are used for different purposes: choosing a checkpoint
|
| 179 |
+
on one and reporting numbers on the other keeps the reported numbers
|
| 180 |
+
free of selection bias.
|
| 181 |
+
"""
|
| 182 |
+
from datasets import load_dataset
|
| 183 |
+
|
| 184 |
+
data = load_dataset(DATASET_REPO, split="validation", revision=DATASET_REVISION)
|
| 185 |
+
halves = {"select": [], "report": []}
|
| 186 |
+
for index, row in enumerate(data):
|
| 187 |
+
key = normalize_prompt(row["prompt"])
|
| 188 |
+
half = "select" if int(hashlib.sha1(key.encode()).hexdigest(), 16) % 2 == 0 else "report"
|
| 189 |
+
halves[half].append({
|
| 190 |
+
"index": index,
|
| 191 |
+
"prompt": row["prompt"],
|
| 192 |
+
"image": center_crop(row["image"]),
|
| 193 |
+
"seed": SEED + index, # same seed for every model on this row
|
| 194 |
+
})
|
| 195 |
+
if LIMIT_HALF:
|
| 196 |
+
halves = {k: v[:LIMIT_HALF] for k, v in halves.items()}
|
| 197 |
+
log(f"validation halves: select {len(halves['select'])}, report {len(halves['report'])}")
|
| 198 |
+
return halves
|
| 199 |
+
|
| 200 |
+
|
| 201 |
+
# ---------------------------------------------------------------------------
|
| 202 |
+
# pipeline
|
| 203 |
+
# ---------------------------------------------------------------------------
|
| 204 |
+
|
| 205 |
+
class Runner:
|
| 206 |
+
"""One SD 1.5 pipeline; LoRA checkpoints are swapped in and out of it."""
|
| 207 |
+
|
| 208 |
+
def __init__(self) -> None:
|
| 209 |
+
from diffusers import StableDiffusionPipeline
|
| 210 |
+
|
| 211 |
+
self.pipe = StableDiffusionPipeline.from_pretrained(
|
| 212 |
+
BASE_MODEL, torch_dtype=DTYPE, variant="fp16", use_safetensors=True,
|
| 213 |
+
safety_checker=None, requires_safety_checker=False,
|
| 214 |
+
).to(DEVICE)
|
| 215 |
+
self.pipe.set_progress_bar_config(disable=True)
|
| 216 |
+
self.current: str | None = None
|
| 217 |
+
|
| 218 |
+
def use(self, checkpoint: str | None, scale: float = 1.0) -> None:
|
| 219 |
+
if checkpoint in (None, "base"):
|
| 220 |
+
if self.current:
|
| 221 |
+
self.pipe.unload_lora_weights()
|
| 222 |
+
self.current = None
|
| 223 |
+
return
|
| 224 |
+
if self.current != checkpoint:
|
| 225 |
+
if self.current:
|
| 226 |
+
self.pipe.unload_lora_weights()
|
| 227 |
+
self.pipe.load_lora_weights(
|
| 228 |
+
MODEL_REPO, subfolder=f"{CHECKPOINT_DIR}/{checkpoint}",
|
| 229 |
+
weight_name="pytorch_lora_weights.safetensors", adapter_name=ADAPTER,
|
| 230 |
+
)
|
| 231 |
+
self.current = checkpoint
|
| 232 |
+
self.pipe.set_adapters([ADAPTER], adapter_weights=[float(scale)])
|
| 233 |
+
|
| 234 |
+
@torch.no_grad()
|
| 235 |
+
def render(self, prompts: list[str], seeds: list[int]) -> list[Image.Image]:
|
| 236 |
+
images: list[Image.Image] = []
|
| 237 |
+
for start in range(0, len(prompts), RENDER_BATCH):
|
| 238 |
+
batch = prompts[start:start + RENDER_BATCH]
|
| 239 |
+
generators = [torch.Generator(DEVICE).manual_seed(int(s))
|
| 240 |
+
for s in seeds[start:start + RENDER_BATCH]]
|
| 241 |
+
images += self.pipe(batch, num_inference_steps=STEPS, guidance_scale=GUIDANCE,
|
| 242 |
+
generator=generators).images
|
| 243 |
+
return images
|
| 244 |
+
|
| 245 |
+
|
| 246 |
+
_RUNNER: Runner | None = None
|
| 247 |
+
|
| 248 |
+
|
| 249 |
+
def runner() -> Runner:
|
| 250 |
+
global _RUNNER
|
| 251 |
+
if _RUNNER is None:
|
| 252 |
+
log(f"loading {BASE_MODEL}")
|
| 253 |
+
_RUNNER = Runner()
|
| 254 |
+
return _RUNNER
|
| 255 |
+
|
| 256 |
+
|
| 257 |
+
def render_model(model: str, prompts: list[str], seeds: list[int], scale: float = 1.0):
|
| 258 |
+
run = runner()
|
| 259 |
+
run.use(model, scale)
|
| 260 |
+
started = time.time()
|
| 261 |
+
images = run.render(prompts, seeds)
|
| 262 |
+
log(f" rendered {len(images)} with {model} (scale {scale}) "
|
| 263 |
+
f"in {time.time() - started:.0f}s")
|
| 264 |
+
return images
|
| 265 |
+
|
| 266 |
+
|
| 267 |
+
def selected() -> str:
|
| 268 |
+
return STATE.get("selected", FINAL)
|
| 269 |
+
|
| 270 |
+
|
| 271 |
+
def report_models() -> list[str]:
|
| 272 |
+
models = ["base", selected()]
|
| 273 |
+
if selected() != FINAL:
|
| 274 |
+
models.append(FINAL)
|
| 275 |
+
return models
|
| 276 |
+
|
| 277 |
+
|
| 278 |
+
# ---------------------------------------------------------------------------
|
| 279 |
+
# image helpers
|
| 280 |
+
# ---------------------------------------------------------------------------
|
| 281 |
+
|
| 282 |
+
def font(size: int = 16):
|
| 283 |
+
try:
|
| 284 |
+
return ImageFont.load_default(size=size)
|
| 285 |
+
except TypeError:
|
| 286 |
+
return ImageFont.load_default()
|
| 287 |
+
|
| 288 |
+
|
| 289 |
+
def labelled_grid(rows: list[list[Image.Image]], column_labels: list[str],
|
| 290 |
+
row_labels: list[str] | None = None, thumb: int = 256) -> Image.Image:
|
| 291 |
+
top, left, pad = 32, (220 if row_labels else 0), 6
|
| 292 |
+
width = left + len(column_labels) * (thumb + pad) + pad
|
| 293 |
+
height = top + len(rows) * (thumb + pad) + pad
|
| 294 |
+
sheet = Image.new("RGB", (width, height), (252, 252, 251))
|
| 295 |
+
draw = ImageDraw.Draw(sheet)
|
| 296 |
+
for c, label in enumerate(column_labels):
|
| 297 |
+
draw.text((left + pad + c * (thumb + pad) + 4, 8), label, fill=(11, 11, 11), font=font(16))
|
| 298 |
+
for r, row in enumerate(rows):
|
| 299 |
+
y = top + pad + r * (thumb + pad)
|
| 300 |
+
if row_labels:
|
| 301 |
+
text = row_labels[r]
|
| 302 |
+
lines = [text[i:i + 26] for i in range(0, min(len(text), 26 * 6), 26)]
|
| 303 |
+
draw.multiline_text((8, y + 4), "\n".join(lines), fill=(82, 81, 78), font=font(13))
|
| 304 |
+
for c, image in enumerate(row):
|
| 305 |
+
sheet.paste(image.resize((thumb, thumb), Image.LANCZOS), (left + pad + c * (thumb + pad), y))
|
| 306 |
+
return sheet
|
| 307 |
+
|
| 308 |
+
|
| 309 |
+
def contact_sheet(images: list[Image.Image], columns: int = 5, thumb: int = 320) -> Image.Image:
|
| 310 |
+
rows = (len(images) + columns - 1) // columns
|
| 311 |
+
sheet = Image.new("RGB", (columns * thumb, rows * thumb), (18, 18, 22))
|
| 312 |
+
for i, image in enumerate(images):
|
| 313 |
+
sheet.paste(image.resize((thumb, thumb), Image.LANCZOS), ((i % columns) * thumb, (i // columns) * thumb))
|
| 314 |
+
return sheet
|
| 315 |
+
|
| 316 |
+
|
| 317 |
+
def to_tensor(images: list[Image.Image]) -> torch.Tensor:
|
| 318 |
+
"""PIL list -> float tensor in [0, 1], shape (N, 3, H, W)."""
|
| 319 |
+
array = np.stack([np.asarray(im.convert("RGB"), dtype=np.uint8) for im in images])
|
| 320 |
+
return torch.from_numpy(array).permute(0, 3, 1, 2).float() / 255.0
|
| 321 |
+
|
| 322 |
+
|
| 323 |
+
def max_pixel_diff(a: Image.Image, b: Image.Image) -> tuple[int, float]:
|
| 324 |
+
x = np.asarray(a, dtype=np.int16)
|
| 325 |
+
y = np.asarray(b, dtype=np.int16)
|
| 326 |
+
d = np.abs(x - y)
|
| 327 |
+
return int(d.max()), float(d.mean())
|
| 328 |
+
|
| 329 |
+
|
| 330 |
+
# ---------------------------------------------------------------------------
|
| 331 |
+
# CLIP
|
| 332 |
+
# ---------------------------------------------------------------------------
|
| 333 |
+
|
| 334 |
+
_CLIP = None
|
| 335 |
+
|
| 336 |
+
|
| 337 |
+
def clip():
|
| 338 |
+
global _CLIP
|
| 339 |
+
if _CLIP is None:
|
| 340 |
+
from transformers import CLIPModel, CLIPProcessor
|
| 341 |
+
|
| 342 |
+
_CLIP = (CLIPModel.from_pretrained(CLIP_MODEL, torch_dtype=DTYPE).to(DEVICE).eval(),
|
| 343 |
+
CLIPProcessor.from_pretrained(CLIP_MODEL))
|
| 344 |
+
return _CLIP
|
| 345 |
+
|
| 346 |
+
|
| 347 |
+
@torch.no_grad()
|
| 348 |
+
def clip_image_embeddings(images: list[Image.Image], batch: int = 64) -> torch.Tensor:
|
| 349 |
+
model, processor = clip()
|
| 350 |
+
out = []
|
| 351 |
+
for start in range(0, len(images), batch):
|
| 352 |
+
pixels = processor(images=images[start:start + batch], return_tensors="pt")["pixel_values"]
|
| 353 |
+
emb = model.get_image_features(pixel_values=pixels.to(DEVICE, DTYPE))
|
| 354 |
+
out.append(F.normalize(emb.float(), dim=-1))
|
| 355 |
+
return torch.cat(out)
|
| 356 |
+
|
| 357 |
+
|
| 358 |
+
@torch.no_grad()
|
| 359 |
+
def clip_scores(images: list[Image.Image], prompts: list[str], batch: int = 64) -> np.ndarray:
|
| 360 |
+
"""CLIP score as in torchmetrics: 100 * max(cos(image, text), 0).
|
| 361 |
+
|
| 362 |
+
Prompts are truncated to CLIP's 77-token limit, as the SD text encoder does.
|
| 363 |
+
"""
|
| 364 |
+
model, processor = clip()
|
| 365 |
+
image_emb = clip_image_embeddings(images, batch)
|
| 366 |
+
scores = []
|
| 367 |
+
for start in range(0, len(prompts), batch):
|
| 368 |
+
tokens = processor(text=prompts[start:start + batch], return_tensors="pt",
|
| 369 |
+
padding=True, truncation=True, max_length=77)
|
| 370 |
+
text = model.get_text_features(input_ids=tokens["input_ids"].to(DEVICE),
|
| 371 |
+
attention_mask=tokens["attention_mask"].to(DEVICE))
|
| 372 |
+
text = F.normalize(text.float(), dim=-1)
|
| 373 |
+
cos = (image_emb[start:start + batch] * text).sum(-1)
|
| 374 |
+
scores.append((100 * cos.clamp(min=0)).cpu().numpy())
|
| 375 |
+
return np.concatenate(scores)
|
| 376 |
+
|
| 377 |
+
|
| 378 |
+
# ---------------------------------------------------------------------------
|
| 379 |
+
# stages
|
| 380 |
+
# ---------------------------------------------------------------------------
|
| 381 |
+
|
| 382 |
+
def stage_snippet() -> dict:
|
| 383 |
+
"""The model card's "How to get started" code, verbatim (plus a save path)."""
|
| 384 |
+
import diffusers
|
| 385 |
+
from diffusers import StableDiffusionPipeline
|
| 386 |
+
|
| 387 |
+
pipe = StableDiffusionPipeline.from_pretrained(
|
| 388 |
+
"stable-diffusion-v1-5/stable-diffusion-v1-5",
|
| 389 |
+
torch_dtype=torch.float16,
|
| 390 |
+
variant="fp16",
|
| 391 |
+
).to("cuda")
|
| 392 |
+
|
| 393 |
+
pipe.load_lora_weights(
|
| 394 |
+
"whosouravsharma/diffusiondb-sd15-lora",
|
| 395 |
+
subfolder="checkpoints/checkpoint-4240",
|
| 396 |
+
weight_name="pytorch_lora_weights.safetensors",
|
| 397 |
+
adapter_name="diffusiondb",
|
| 398 |
+
)
|
| 399 |
+
pipe.set_adapters(["diffusiondb"], adapter_weights=[1.0]) # 0.0 = plain SD 1.5
|
| 400 |
+
|
| 401 |
+
image = pipe(
|
| 402 |
+
"a steampunk owl inside a glass jar, intricate detail",
|
| 403 |
+
num_inference_steps=25,
|
| 404 |
+
guidance_scale=7.5,
|
| 405 |
+
generator=torch.Generator("cuda").manual_seed(42),
|
| 406 |
+
).images[0]
|
| 407 |
+
|
| 408 |
+
path = OUT / "snippet" / "owl.png"
|
| 409 |
+
path.parent.mkdir(parents=True, exist_ok=True)
|
| 410 |
+
image.save(path)
|
| 411 |
+
del pipe
|
| 412 |
+
torch.cuda.empty_cache()
|
| 413 |
+
return {"passed": True, "diffusers": diffusers.__version__, "torch": torch.__version__,
|
| 414 |
+
"image": str(path.relative_to(OUT))}
|
| 415 |
+
|
| 416 |
+
|
| 417 |
+
def stage_checks() -> dict:
|
| 418 |
+
prompt = "a steampunk owl inside a glass jar, intricate detail"
|
| 419 |
+
run = runner()
|
| 420 |
+
|
| 421 |
+
run.use("base")
|
| 422 |
+
base = run.render([prompt], [SEED])[0]
|
| 423 |
+
run.use(FINAL, 0.0)
|
| 424 |
+
scale_zero = run.render([prompt], [SEED])[0]
|
| 425 |
+
run.use(FINAL, 1.0)
|
| 426 |
+
lora_a = run.render([prompt], [SEED])[0]
|
| 427 |
+
lora_b = run.render([prompt], [SEED])[0]
|
| 428 |
+
run.use("base")
|
| 429 |
+
unloaded = run.render([prompt], [SEED])[0]
|
| 430 |
+
|
| 431 |
+
folder = OUT / "checks"
|
| 432 |
+
folder.mkdir(parents=True, exist_ok=True)
|
| 433 |
+
for name, image in [("base", base), ("scale0", scale_zero), ("lora_a", lora_a),
|
| 434 |
+
("lora_b", lora_b), ("after_unload", unloaded)]:
|
| 435 |
+
image.save(folder / f"{name}.png")
|
| 436 |
+
|
| 437 |
+
results = {}
|
| 438 |
+
for name, (x, y) in {
|
| 439 |
+
"scale0_equals_base": (base, scale_zero),
|
| 440 |
+
"unload_restores_base": (base, unloaded),
|
| 441 |
+
"same_seed_is_deterministic": (lora_a, lora_b),
|
| 442 |
+
"lora_changes_output": (base, lora_a),
|
| 443 |
+
}.items():
|
| 444 |
+
max_diff, mean_diff = max_pixel_diff(x, y)
|
| 445 |
+
expected_equal = name != "lora_changes_output"
|
| 446 |
+
passed = (max_diff <= 1) if expected_equal else (mean_diff > 1.0)
|
| 447 |
+
results[name] = {"max_pixel_diff": max_diff, "mean_pixel_diff": round(mean_diff, 4),
|
| 448 |
+
"passed": passed}
|
| 449 |
+
results["passed"] = all(r["passed"] for r in results.values() if isinstance(r, dict))
|
| 450 |
+
return results
|
| 451 |
+
|
| 452 |
+
|
| 453 |
+
@torch.no_grad()
|
| 454 |
+
def stage_loss() -> dict:
|
| 455 |
+
from datasets import load_dataset
|
| 456 |
+
from diffusers import DDPMScheduler
|
| 457 |
+
from huggingface_hub import hf_hub_download
|
| 458 |
+
|
| 459 |
+
manifest = json.loads(Path(hf_hub_download(
|
| 460 |
+
DATASET_REPO, "data/manifest.json", repo_type="dataset", revision=LATENTS_REVISION,
|
| 461 |
+
)).read_text())
|
| 462 |
+
scaling = manifest["scaling_factor"]
|
| 463 |
+
shape = tuple(manifest["latent_shape"])
|
| 464 |
+
|
| 465 |
+
data = load_dataset(DATASET_REPO, split="validation", revision=LATENTS_REVISION)
|
| 466 |
+
if LIMIT_LOSS:
|
| 467 |
+
data = data.select(range(LIMIT_LOSS))
|
| 468 |
+
|
| 469 |
+
def decode(blobs):
|
| 470 |
+
return torch.from_numpy(np.stack(
|
| 471 |
+
[np.frombuffer(b, dtype=np.float16).reshape(shape) for b in blobs]).copy())
|
| 472 |
+
|
| 473 |
+
scheduler = DDPMScheduler.from_pretrained(BASE_MODEL, subfolder="scheduler")
|
| 474 |
+
run = runner()
|
| 475 |
+
pipe = run.pipe
|
| 476 |
+
losses = {}
|
| 477 |
+
|
| 478 |
+
for model in ["base"] + list_checkpoints():
|
| 479 |
+
run.use(model, 1.0)
|
| 480 |
+
generator = torch.Generator(DEVICE).manual_seed(SEED) # identical noise per model
|
| 481 |
+
total, batches = 0.0, 0
|
| 482 |
+
for start in range(0, len(data), LOSS_BATCH):
|
| 483 |
+
rows = data[start:start + LOSS_BATCH]
|
| 484 |
+
mean = decode(rows["latent_mean"]).to(DEVICE, torch.float32)
|
| 485 |
+
logvar = decode(rows["latent_logvar"]).to(DEVICE, torch.float32).clamp(-30.0, 20.0)
|
| 486 |
+
eps = torch.randn(mean.shape, device=DEVICE, generator=generator)
|
| 487 |
+
latents = (mean + torch.exp(0.5 * logvar) * eps) * scaling
|
| 488 |
+
noise = torch.randn(latents.shape, device=DEVICE, generator=generator)
|
| 489 |
+
steps = torch.randint(0, scheduler.config.num_train_timesteps, (latents.shape[0],),
|
| 490 |
+
device=DEVICE, generator=generator)
|
| 491 |
+
noisy = scheduler.add_noise(latents, noise, steps)
|
| 492 |
+
tokens = pipe.tokenizer(rows["prompt"], padding="max_length", truncation=True,
|
| 493 |
+
max_length=pipe.tokenizer.model_max_length,
|
| 494 |
+
return_tensors="pt").input_ids.to(DEVICE)
|
| 495 |
+
encoded = pipe.text_encoder(tokens)[0]
|
| 496 |
+
predicted = pipe.unet(noisy.to(DTYPE), steps, encoded).sample
|
| 497 |
+
total += F.mse_loss(predicted.float(), noise.float()).item()
|
| 498 |
+
batches += 1
|
| 499 |
+
losses[model] = total / max(batches, 1)
|
| 500 |
+
log(f" val loss {model}: {losses[model]:.5f}")
|
| 501 |
+
|
| 502 |
+
write_json(OUT / "loss" / "val_loss.json", {"examples": len(data), "seed": SEED, "loss": losses})
|
| 503 |
+
plot_loss(losses, OUT / "loss" / "val_loss.png")
|
| 504 |
+
return {"examples": len(data), "loss": losses}
|
| 505 |
+
|
| 506 |
+
|
| 507 |
+
def plot_loss(losses: dict, path: Path) -> None:
|
| 508 |
+
import matplotlib
|
| 509 |
+
matplotlib.use("Agg")
|
| 510 |
+
import matplotlib.pyplot as plt
|
| 511 |
+
|
| 512 |
+
surface, ink, muted, grid, series = "#fcfcfb", "#0b0b0b", "#52514e", "#e6e5e1", "#2a78d6"
|
| 513 |
+
points = [(int(k.split("-")[1]), v) for k, v in losses.items() if k != "base"]
|
| 514 |
+
points.sort()
|
| 515 |
+
steps = [p[0] for p in points]
|
| 516 |
+
values = [p[1] for p in points]
|
| 517 |
+
|
| 518 |
+
fig, ax = plt.subplots(figsize=(8, 4.2), dpi=150)
|
| 519 |
+
fig.patch.set_facecolor(surface)
|
| 520 |
+
ax.set_facecolor(surface)
|
| 521 |
+
ax.plot(steps, values, color=series, linewidth=2, marker="o", markersize=6,
|
| 522 |
+
markeredgecolor=surface, markeredgewidth=2, zorder=3)
|
| 523 |
+
if "base" in losses:
|
| 524 |
+
ax.axhline(losses["base"], color=muted, linewidth=1.5, linestyle=(0, (4, 3)), zorder=2)
|
| 525 |
+
ax.annotate("SD 1.5 base (no LoRA)", xy=(steps[0] if steps else 0, losses["base"]),
|
| 526 |
+
xytext=(0, 6), textcoords="offset points", color=muted, fontsize=9)
|
| 527 |
+
if points:
|
| 528 |
+
ax.annotate(f"{values[-1]:.4f}", xy=(steps[-1], values[-1]), xytext=(6, -12),
|
| 529 |
+
textcoords="offset points", color=ink, fontsize=9)
|
| 530 |
+
ax.set_title("Validation loss by checkpoint (1,000 held-out images, fixed noise)",
|
| 531 |
+
color=ink, fontsize=11, loc="left")
|
| 532 |
+
ax.set_xlabel("training step", color=muted, fontsize=9)
|
| 533 |
+
ax.set_ylabel("MSE (noise prediction)", color=muted, fontsize=9)
|
| 534 |
+
ax.tick_params(colors=muted, labelsize=8)
|
| 535 |
+
ax.grid(axis="y", color=grid, linewidth=0.8)
|
| 536 |
+
for side in ("top", "right", "left"):
|
| 537 |
+
ax.spines[side].set_visible(False)
|
| 538 |
+
ax.spines["bottom"].set_color(grid)
|
| 539 |
+
fig.tight_layout()
|
| 540 |
+
path.parent.mkdir(parents=True, exist_ok=True)
|
| 541 |
+
fig.savefig(path, facecolor=surface)
|
| 542 |
+
plt.close(fig)
|
| 543 |
+
|
| 544 |
+
|
| 545 |
+
def inception_metrics(real: torch.Tensor):
|
| 546 |
+
from torchmetrics.image.fid import FrechetInceptionDistance
|
| 547 |
+
from torchmetrics.image.kid import KernelInceptionDistance
|
| 548 |
+
|
| 549 |
+
n = real.shape[0]
|
| 550 |
+
subset = max(2, min(1000, n // 2))
|
| 551 |
+
fid = FrechetInceptionDistance(feature=2048, normalize=True, reset_real_features=False).to(DEVICE)
|
| 552 |
+
kid = KernelInceptionDistance(subset_size=subset, subsets=100, normalize=True,
|
| 553 |
+
reset_real_features=False).to(DEVICE)
|
| 554 |
+
for start in range(0, n, 50):
|
| 555 |
+
chunk = real[start:start + 50].to(DEVICE)
|
| 556 |
+
fid.update(chunk, real=True)
|
| 557 |
+
kid.update(chunk, real=True)
|
| 558 |
+
return fid, kid
|
| 559 |
+
|
| 560 |
+
|
| 561 |
+
def score_fake(fid, kid, fake: torch.Tensor) -> dict:
|
| 562 |
+
fid.reset()
|
| 563 |
+
kid.reset()
|
| 564 |
+
for start in range(0, fake.shape[0], 50):
|
| 565 |
+
chunk = fake[start:start + 50].to(DEVICE)
|
| 566 |
+
fid.update(chunk, real=False)
|
| 567 |
+
kid.update(chunk, real=False)
|
| 568 |
+
kid_mean, kid_std = kid.compute()
|
| 569 |
+
return {"fid": float(fid.compute()), "kid": float(kid_mean), "kid_std": float(kid_std)}
|
| 570 |
+
|
| 571 |
+
|
| 572 |
+
def stage_metrics() -> dict:
|
| 573 |
+
halves = validation_halves()
|
| 574 |
+
result: dict = {"n_select": len(halves["select"]), "n_report": len(halves["report"])}
|
| 575 |
+
|
| 576 |
+
# --- selection half: pick the checkpoint closest to DiffusionDB (lowest KID)
|
| 577 |
+
rows = halves["select"]
|
| 578 |
+
prompts, seeds = [r["prompt"] for r in rows], [r["seed"] for r in rows]
|
| 579 |
+
fid, kid = inception_metrics(to_tensor([r["image"] for r in rows]))
|
| 580 |
+
candidates = CANDIDATES[:LIMIT_CKPTS] if LIMIT_CKPTS else CANDIDATES
|
| 581 |
+
selection = {}
|
| 582 |
+
for model in candidates:
|
| 583 |
+
images = render_model(model, prompts, seeds)
|
| 584 |
+
selection[model] = score_fake(fid, kid, to_tensor(images))
|
| 585 |
+
selection[model]["clip"] = float(clip_scores(images, prompts).mean())
|
| 586 |
+
log(f" select {model}: {selection[model]}")
|
| 587 |
+
winner = min(selection, key=lambda m: selection[m]["kid"])
|
| 588 |
+
STATE["selected"] = winner
|
| 589 |
+
result["selection"] = {"criterion": "lowest KID on the selection half",
|
| 590 |
+
"candidates": selection, "selected": winner}
|
| 591 |
+
del fid, kid
|
| 592 |
+
torch.cuda.empty_cache()
|
| 593 |
+
|
| 594 |
+
# --- report half: the numbers that go in the model card
|
| 595 |
+
rows = halves["report"]
|
| 596 |
+
prompts, seeds = [r["prompt"] for r in rows], [r["seed"] for r in rows]
|
| 597 |
+
real_images = [r["image"] for r in rows]
|
| 598 |
+
fid, kid = inception_metrics(to_tensor(real_images))
|
| 599 |
+
report = {"real_images": {"clip": float(clip_scores(real_images, prompts).mean()),
|
| 600 |
+
"clip_std": float(clip_scores(real_images, prompts).std())}}
|
| 601 |
+
renders = {"real": real_images}
|
| 602 |
+
for model in report_models():
|
| 603 |
+
images = render_model(model, prompts, seeds)
|
| 604 |
+
renders[model] = images
|
| 605 |
+
scores = clip_scores(images, prompts)
|
| 606 |
+
report[model] = score_fake(fid, kid, to_tensor(images))
|
| 607 |
+
report[model].update({"clip": float(scores.mean()), "clip_std": float(scores.std())})
|
| 608 |
+
folder = OUT / "renders" / model
|
| 609 |
+
folder.mkdir(parents=True, exist_ok=True)
|
| 610 |
+
for row, image in zip(rows, images):
|
| 611 |
+
image.save(folder / f"{row['index']:04}.jpg", quality=92)
|
| 612 |
+
log(f" report {model}: {report[model]}")
|
| 613 |
+
STATE["report_rows"] = rows
|
| 614 |
+
STATE["renders"] = renders
|
| 615 |
+
result["report"] = report
|
| 616 |
+
result["notes"] = {
|
| 617 |
+
"reference": "real validation images from the report half, centre-cropped to 512",
|
| 618 |
+
"kid_subset_size": max(2, min(1000, len(rows) // 2)),
|
| 619 |
+
"render": {"steps": STEPS, "guidance": GUIDANCE, "scheduler": "PNDM (pipeline default)",
|
| 620 |
+
"seed": "42 + validation row index"},
|
| 621 |
+
"clip_model": CLIP_MODEL,
|
| 622 |
+
}
|
| 623 |
+
write_json(OUT / "metrics" / "metrics.json", result)
|
| 624 |
+
return result
|
| 625 |
+
|
| 626 |
+
|
| 627 |
+
def stage_grids() -> dict:
|
| 628 |
+
prompts = eval_prompts()
|
| 629 |
+
if LIMIT_GRID:
|
| 630 |
+
prompts = prompts[:LIMIT_GRID]
|
| 631 |
+
seeds = [SEED + i for i in range(len(prompts))]
|
| 632 |
+
models = report_models()
|
| 633 |
+
outputs = {}
|
| 634 |
+
for model in models:
|
| 635 |
+
images = render_model(model, prompts, seeds)
|
| 636 |
+
outputs[model] = images
|
| 637 |
+
folder = OUT / "samples" / model
|
| 638 |
+
folder.mkdir(parents=True, exist_ok=True)
|
| 639 |
+
for i, image in enumerate(images):
|
| 640 |
+
image.save(folder / f"{i:03}.png")
|
| 641 |
+
contact_sheet(images).save(folder / "grid.jpg", quality=92)
|
| 642 |
+
write_json(folder / "prompts.json", {"checkpoint": model, "steps": STEPS, "guidance": GUIDANCE,
|
| 643 |
+
"seed": SEED, "prompts": prompts})
|
| 644 |
+
pairs = OUT / "samples" / "pairs"
|
| 645 |
+
pairs.mkdir(parents=True, exist_ok=True)
|
| 646 |
+
for i in range(len(prompts)):
|
| 647 |
+
labelled_grid([[outputs["base"][i], outputs[selected()][i]]],
|
| 648 |
+
["SD 1.5 base", f"+ LoRA ({selected()})"], thumb=384).save(
|
| 649 |
+
pairs / f"{i:03}.jpg", quality=92)
|
| 650 |
+
return {"prompts": len(prompts), "models": models}
|
| 651 |
+
|
| 652 |
+
|
| 653 |
+
def stage_sweep() -> dict:
|
| 654 |
+
prompts = eval_prompts()
|
| 655 |
+
chosen = [prompts[i] for i in SWEEP_PROMPTS][:1 if SMOKE else None]
|
| 656 |
+
rows = []
|
| 657 |
+
for p_index, prompt in zip(SWEEP_PROMPTS, chosen):
|
| 658 |
+
row = []
|
| 659 |
+
for scale in SWEEP_SCALES:
|
| 660 |
+
row += render_model(selected(), [prompt], [SEED + p_index], scale)
|
| 661 |
+
rows.append(row)
|
| 662 |
+
path = OUT / "sweep" / "lora_strength.jpg"
|
| 663 |
+
path.parent.mkdir(parents=True, exist_ok=True)
|
| 664 |
+
labelled_grid(rows, [f"strength {s}" for s in SWEEP_SCALES], chosen).save(path, quality=92)
|
| 665 |
+
return {"checkpoint": selected(), "scales": SWEEP_SCALES, "prompts": chosen}
|
| 666 |
+
|
| 667 |
+
|
| 668 |
+
def stage_progression() -> dict:
|
| 669 |
+
prompts = eval_prompts()
|
| 670 |
+
chosen = [prompts[i] for i in PROGRESSION_PROMPTS][:1 if SMOKE else None]
|
| 671 |
+
models = ["base"] + list_checkpoints()
|
| 672 |
+
columns = {m: render_model(m, chosen, [SEED + i for i in PROGRESSION_PROMPTS[:len(chosen)]])
|
| 673 |
+
for m in models}
|
| 674 |
+
rows = [[columns[m][r] for m in models] for r in range(len(chosen))]
|
| 675 |
+
labels = ["base"] + [m.split("-")[1] for m in models[1:]]
|
| 676 |
+
path = OUT / "progression" / "checkpoints.jpg"
|
| 677 |
+
path.parent.mkdir(parents=True, exist_ok=True)
|
| 678 |
+
labelled_grid(rows, labels, chosen, thumb=192).save(path, quality=92)
|
| 679 |
+
return {"prompts": chosen, "columns": labels}
|
| 680 |
+
|
| 681 |
+
|
| 682 |
+
def report_renders() -> dict:
|
| 683 |
+
"""Report-half renders from the metrics stage, or reloaded from disk."""
|
| 684 |
+
if "renders" in STATE:
|
| 685 |
+
return STATE["renders"]
|
| 686 |
+
raise RuntimeError("safety and memorization need the metrics stage in the same run")
|
| 687 |
+
|
| 688 |
+
|
| 689 |
+
@torch.no_grad()
|
| 690 |
+
def stage_safety() -> dict:
|
| 691 |
+
from diffusers.pipelines.stable_diffusion.safety_checker import StableDiffusionSafetyChecker
|
| 692 |
+
from transformers import CLIPImageProcessor
|
| 693 |
+
|
| 694 |
+
checker = StableDiffusionSafetyChecker.from_pretrained(BASE_MODEL, subfolder="safety_checker").to(DEVICE).eval()
|
| 695 |
+
processor = CLIPImageProcessor.from_pretrained(BASE_MODEL, subfolder="feature_extractor")
|
| 696 |
+
result = {}
|
| 697 |
+
for name, images in report_renders().items():
|
| 698 |
+
flagged = 0
|
| 699 |
+
for start in range(0, len(images), 32):
|
| 700 |
+
chunk = images[start:start + 32]
|
| 701 |
+
pixels = processor(chunk, return_tensors="pt").pixel_values.to(DEVICE)
|
| 702 |
+
_, has_nsfw = checker(images=np.zeros((len(chunk), 1, 1, 3)), clip_input=pixels)
|
| 703 |
+
flagged += int(sum(bool(x) for x in has_nsfw))
|
| 704 |
+
result[name] = {"flagged": flagged, "total": len(images),
|
| 705 |
+
"rate": round(flagged / max(len(images), 1), 4)}
|
| 706 |
+
log(f" safety {name}: {result[name]}")
|
| 707 |
+
del checker
|
| 708 |
+
torch.cuda.empty_cache()
|
| 709 |
+
write_json(OUT / "safety" / "safety.json", result)
|
| 710 |
+
return result
|
| 711 |
+
|
| 712 |
+
|
| 713 |
+
@torch.no_grad()
|
| 714 |
+
def stage_memorization() -> dict:
|
| 715 |
+
"""Nearest training image for every report-half render, by CLIP cosine.
|
| 716 |
+
|
| 717 |
+
CLIP similarity is a proxy for copying, not proof either way. Real
|
| 718 |
+
validation images (distinct prompts, same style) give the baseline for
|
| 719 |
+
what "similar" means in this dataset.
|
| 720 |
+
"""
|
| 721 |
+
from datasets import load_dataset
|
| 722 |
+
|
| 723 |
+
renders = report_renders()
|
| 724 |
+
names = list(renders)
|
| 725 |
+
queries = {n: clip_image_embeddings(renders[n]) for n in names}
|
| 726 |
+
best = {n: torch.full((len(renders[n]),), -1.0, device=DEVICE) for n in names}
|
| 727 |
+
best_thumb = {n: [None] * len(renders[n]) for n in names}
|
| 728 |
+
best_prompt = {n: [None] * len(renders[n]) for n in names}
|
| 729 |
+
|
| 730 |
+
stream = load_dataset(DATASET_REPO, split="train", revision=DATASET_REVISION, streaming=True)
|
| 731 |
+
batch_images, batch_prompts, seen = [], [], 0
|
| 732 |
+
|
| 733 |
+
def flush():
|
| 734 |
+
nonlocal batch_images, batch_prompts
|
| 735 |
+
if not batch_images:
|
| 736 |
+
return
|
| 737 |
+
emb = clip_image_embeddings(batch_images)
|
| 738 |
+
for n in names:
|
| 739 |
+
sims = queries[n] @ emb.T
|
| 740 |
+
top, arg = sims.max(dim=1)
|
| 741 |
+
improved = (top > best[n]).nonzero().flatten().tolist()
|
| 742 |
+
best[n] = torch.maximum(best[n], top)
|
| 743 |
+
for q in improved:
|
| 744 |
+
thumb = batch_images[arg[q].item()].resize((256, 256), Image.LANCZOS)
|
| 745 |
+
buffer = io.BytesIO()
|
| 746 |
+
thumb.save(buffer, format="JPEG", quality=85)
|
| 747 |
+
best_thumb[n][q] = buffer.getvalue()
|
| 748 |
+
best_prompt[n][q] = batch_prompts[arg[q].item()]
|
| 749 |
+
batch_images, batch_prompts = [], []
|
| 750 |
+
|
| 751 |
+
for row in stream:
|
| 752 |
+
batch_images.append(center_crop(row["image"]))
|
| 753 |
+
batch_prompts.append(row["prompt"])
|
| 754 |
+
seen += 1
|
| 755 |
+
if len(batch_images) == 64:
|
| 756 |
+
flush()
|
| 757 |
+
if LIMIT_TRAIN and seen >= LIMIT_TRAIN:
|
| 758 |
+
break
|
| 759 |
+
if seen % 2000 == 0:
|
| 760 |
+
log(f" memorization: {seen} training images embedded")
|
| 761 |
+
flush()
|
| 762 |
+
|
| 763 |
+
result = {"training_images": seen, "similarity": "CLIP ViT-L/14 image cosine"}
|
| 764 |
+
for n in names:
|
| 765 |
+
values = best[n].cpu().numpy()
|
| 766 |
+
result[n] = {
|
| 767 |
+
"mean": round(float(values.mean()), 4), "median": round(float(np.median(values)), 4),
|
| 768 |
+
"p95": round(float(np.percentile(values, 95)), 4), "max": round(float(values.max()), 4),
|
| 769 |
+
"count_ge_0.90": int((values >= 0.90).sum()), "count_ge_0.95": int((values >= 0.95).sum()),
|
| 770 |
+
}
|
| 771 |
+
log(f" memorization {n}: {result[n]}")
|
| 772 |
+
|
| 773 |
+
# The closest render/training pairs for the selected checkpoint, to look at.
|
| 774 |
+
model = selected()
|
| 775 |
+
order = np.argsort(-best[model].cpu().numpy())[:8]
|
| 776 |
+
rows = [[renders[model][q], Image.open(io.BytesIO(best_thumb[model][q]))] for q in order]
|
| 777 |
+
labels = [f"cos {best[model][q].item():.3f}" for q in order]
|
| 778 |
+
path = OUT / "memorization" / "closest_pairs.jpg"
|
| 779 |
+
path.parent.mkdir(parents=True, exist_ok=True)
|
| 780 |
+
labelled_grid(rows, ["render", "nearest training image"], labels, thumb=256).save(path, quality=90)
|
| 781 |
+
write_json(OUT / "memorization" / "memorization.json", result)
|
| 782 |
+
return result
|
| 783 |
+
|
| 784 |
+
|
| 785 |
+
STAGE_FUNCTIONS = {
|
| 786 |
+
"snippet": stage_snippet, "checks": stage_checks, "loss": stage_loss,
|
| 787 |
+
"metrics": stage_metrics, "grids": stage_grids, "sweep": stage_sweep,
|
| 788 |
+
"progression": stage_progression, "safety": stage_safety,
|
| 789 |
+
"memorization": stage_memorization,
|
| 790 |
+
}
|
| 791 |
+
|
| 792 |
+
|
| 793 |
+
def main() -> None:
|
| 794 |
+
if not torch.cuda.is_available():
|
| 795 |
+
raise SystemExit("No GPU. Run this on a GPU flavor, e.g. a10g-small.")
|
| 796 |
+
if os.environ.get("HF_TOKEN"):
|
| 797 |
+
log("NOTE: HF_TOKEN is set. This job needs no token; launch it without --secrets HF_TOKEN.")
|
| 798 |
+
OUT.mkdir(parents=True, exist_ok=True)
|
| 799 |
+
log(f"run {RUN_NAME} -> {OUT} | stages: {', '.join(STAGES)} | smoke={SMOKE}")
|
| 800 |
+
STATUS["config"] = {"base": BASE_MODEL, "model": MODEL_REPO, "dataset": f"{DATASET_REPO}@{DATASET_REVISION}",
|
| 801 |
+
"steps": STEPS, "guidance": GUIDANCE, "seed": SEED, "render_batch": RENDER_BATCH,
|
| 802 |
+
"candidates": CANDIDATES, "gpu": torch.cuda.get_device_name(0)}
|
| 803 |
+
save_status()
|
| 804 |
+
|
| 805 |
+
for name in STAGES:
|
| 806 |
+
started = time.time()
|
| 807 |
+
log(f"=== {name}")
|
| 808 |
+
try:
|
| 809 |
+
result = STAGE_FUNCTIONS[name]()
|
| 810 |
+
STATUS["stages"][name] = {"ok": True, "seconds": round(time.time() - started), "result": result}
|
| 811 |
+
except Exception as error:
|
| 812 |
+
traceback.print_exc()
|
| 813 |
+
STATUS["stages"][name] = {"ok": False, "seconds": round(time.time() - started),
|
| 814 |
+
"error": f"{type(error).__name__}: {error}"}
|
| 815 |
+
save_status()
|
| 816 |
+
log(f"=== {name} {'ok' if STATUS['stages'][name]['ok'] else 'FAILED'} "
|
| 817 |
+
f"({STATUS['stages'][name]['seconds']}s)")
|
| 818 |
+
|
| 819 |
+
STATUS["seconds"] = round(time.time() - STATUS["started"])
|
| 820 |
+
save_status()
|
| 821 |
+
failed = [n for n, s in STATUS["stages"].items() if not s["ok"]]
|
| 822 |
+
log(f"done in {STATUS['seconds']}s; failed stages: {failed or 'none'}; results in {OUT}")
|
| 823 |
+
|
| 824 |
+
|
| 825 |
+
if __name__ == "__main__":
|
| 826 |
+
main()
|
images/closest-pairs.jpg
ADDED
|
Git LFS Details
|
images/fail-010.jpg
ADDED
|
Git LFS Details
|
images/fail-013.jpg
ADDED
|
Git LFS Details
|
images/fail-021.jpg
ADDED
|
Git LFS Details
|
images/lora-strength.jpg
ADDED
|
Git LFS Details
|
images/pair-002.jpg
ADDED
|
Git LFS Details
|
images/pair-006.jpg
ADDED
|
Git LFS Details
|
images/pair-015.jpg
ADDED
|
Git LFS Details
|
images/pair-033.jpg
ADDED
|
Git LFS Details
|
images/pair-036.jpg
ADDED
|
Git LFS Details
|
images/pair-049.jpg
ADDED
|
Git LFS Details
|
images/progression.jpg
ADDED
|
Git LFS Details
|
images/val-loss.png
ADDED
|