Text-to-Image
Diffusers
stable-diffusion
stable-diffusion-diffusers
lora
Eval Results (legacy)
whosouravsharma commited on
Commit
8cb43ae
·
verified ·
1 Parent(s): 5801142

Model card: evaluation results, gallery, figures, eval script

Browse files
.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: runwayml/stable-diffusion-v1-5
 
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
- datasets:
12
- - whosouravsharma/text-to-image-diffusiondb-2M
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
13
  ---
14
 
15
  # diffusiondb-sd15-lora
16
 
17
- A LoRA adapter for [Stable Diffusion 1.5](https://huggingface.co/runwayml/stable-diffusion-v1-5),
18
- fine-tuned on a cleaned, safety-filtered slice of
19
- [DiffusionDB](https://huggingface.co/datasets/poloclub/diffusiondb).
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
- ## What this is for
25
 
26
- DiffusionDB is itself Stable Diffusion 1.x output, collected from the official
27
- Stable Diffusion Discord. Fine-tuning SD 1.5 on it therefore shifts the model
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
- ## Training data
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
- | train | 13,598 |
41
- | validation | 1,000 |
 
 
 
 
 
 
 
 
 
 
 
42
 
43
- The validation split is separated by **normalized prompt group**, not by row.
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
- Filters applied upstream: prompts under 4 words dropped, short side ≥ 384 px
49
- and area ≥ 262,144 px, `image_nsfw` and `prompt_nsfw` below 0.2, at most 2
50
- images per normalized prompt, exact SHA-256 duplicates removed.
 
 
 
 
 
 
 
51
 
52
- ## Configuration
53
 
54
  | | |
55
  |---|---|
56
- | base | `runwayml/stable-diffusion-v1-5` |
57
- | adapter | LoRA rank 32, alpha 32, on UNet `to_q`/`to_k`/`to_v`/`to_out.0` |
58
- | text encoder | frozen |
59
- | resolution | 512×512, centre crop |
60
- | VAE | `stabilityai/sd-vae-ft-mse` (latents cached ahead of training) |
61
- | effective batch | 32 (8 × 4 grad accumulation) |
62
- | optimizer | AdamW, lr 1e-4, cosine schedule, 500 warmup |
63
- | caption dropout | 10%, for classifier-free guidance |
64
- | precision | fp16 |
65
 
66
- Each checkpoint carries a `state.json` recording the exact values it was
67
- trained with, so the table above can be checked rather than trusted.
 
 
 
 
68
 
69
- ## Repository layout
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
- ## Evaluation
 
 
 
 
 
 
 
87
 
88
- 50 prompts held out from training entirely (drawn from validation groups) are
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
- Validation loss is logged per epoch, but on a diffusion fine-tune it tracks
94
- output quality only loosely — it is there to catch divergence, not to rank
95
- checkpoints.
 
 
 
 
 
96
 
97
- ## Usage
98
 
99
  ```python
100
  import torch
101
  from diffusers import StableDiffusionPipeline
102
 
103
  pipe = StableDiffusionPipeline.from_pretrained(
104
- "runwayml/stable-diffusion-v1-5", torch_dtype=torch.float16
 
 
105
  ).to("cuda")
 
106
  pipe.load_lora_weights(
107
- "whosouravsharma/diffusiondb-sd15-lora", subfolder="checkpoints/checkpoint-4240"
 
 
 
108
  )
 
109
 
110
  image = pipe(
111
  "a steampunk owl inside a glass jar, intricate detail",
112
- num_inference_steps=30, guidance_scale=7.5,
 
 
113
  ).images[0]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
114
  ```
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
115
 
116
- ## Limitations
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
+ ![LoRA strength sweep](images/lora-strength.jpg)
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
+ ![Validation loss by checkpoint](images/val-loss.png)
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
+ ![Samples at each checkpoint](images/progression.jpg)
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
+ ![Moon prompt](images/fail-013.jpg)
258
+ *"shattering of the moon's surface, digital art, illustration". The detailed
259
+ scene collapses into a flat, sparse composition.*
260
+
261
+ ![Vodka bottle prompt](images/fail-010.jpg)
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
+ ![Polar bear prompt](images/fail-021.jpg)
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
+ ![Closest render / training-image pairs](images/closest-pairs.jpg)
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

  • SHA256: 58e0e65e7908d2bfd159aaf2cbdd1dafbfbdb7b65e430a155aa739569452abb5
  • Pointer size: 131 Bytes
  • Size of remote file: 492 kB
images/fail-010.jpg ADDED

Git LFS Details

  • SHA256: 196636be136c66a766c543d01fed24c120a18a74bd636af8b9a98cad4ff8464b
  • Pointer size: 131 Bytes
  • Size of remote file: 171 kB
images/fail-013.jpg ADDED

Git LFS Details

  • SHA256: 316d897b824c2085c9f8c3e277144d8fd4a0e8d389b17ba26acace68bdd28342
  • Pointer size: 131 Bytes
  • Size of remote file: 148 kB
images/fail-021.jpg ADDED

Git LFS Details

  • SHA256: 5a39a8754fe27886a67600cbbba2930cbf5aa242633fbac44ee453d2c9370a06
  • Pointer size: 131 Bytes
  • Size of remote file: 112 kB
images/lora-strength.jpg ADDED

Git LFS Details

  • SHA256: ee909dca1ef10804665c1a50b01a36efa852afb02a88b4471a6ecaaa408f6f79
  • Pointer size: 131 Bytes
  • Size of remote file: 463 kB
images/pair-002.jpg ADDED

Git LFS Details

  • SHA256: 35bceed3318f62ea9054dd5db3c53c44ace0ce3ae2748d9706acef55e5802060
  • Pointer size: 131 Bytes
  • Size of remote file: 143 kB
images/pair-006.jpg ADDED

Git LFS Details

  • SHA256: ea746b3fa97770739e630a58d8cecebfdfc4957e1934258d0920cb2f3a37824b
  • Pointer size: 131 Bytes
  • Size of remote file: 193 kB
images/pair-015.jpg ADDED

Git LFS Details

  • SHA256: 2d692b4abee5e3f069d95c76253417a5571aca3f9e79fc37ba2ea15c240f661f
  • Pointer size: 131 Bytes
  • Size of remote file: 180 kB
images/pair-033.jpg ADDED

Git LFS Details

  • SHA256: de241008ef513190669d8213aad1d77e4df143e17c0db5b4d1da2ba745b7ffaf
  • Pointer size: 131 Bytes
  • Size of remote file: 152 kB
images/pair-036.jpg ADDED

Git LFS Details

  • SHA256: d92397b87f9a7eff8534ac0dec51c1f30f3a05b8f93798a123cc24535ba5e28c
  • Pointer size: 131 Bytes
  • Size of remote file: 240 kB
images/pair-049.jpg ADDED

Git LFS Details

  • SHA256: da0f16702619123122ad74591b0c6008f35765361eddbcf1a7976fed9c545ebe
  • Pointer size: 131 Bytes
  • Size of remote file: 320 kB
images/progression.jpg ADDED

Git LFS Details

  • SHA256: f505a771b6987abf9902308a208e6289fa357d1751d487b00ba9d13063501eeb
  • Pointer size: 131 Bytes
  • Size of remote file: 561 kB
images/val-loss.png ADDED