Instructions to use chs35/improving-ca-lora-checkpoints with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Inference
- Notebooks
- Google Colab
- Kaggle
File size: 6,574 Bytes
0a3894a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 | ---
license: openrail++
base_model: stabilityai/stable-diffusion-xl-base-1.0
tags:
- lora
- sdxl
- stable-diffusion-xl
- text-to-image
- cityscapes
- concept-aware-lora
- research
library_name: calora
pipeline_tag: text-to-image
---
# Improving CA-LoRA — head-sliced LoRA checkpoints
LoRA adapters from the reproduction-and-extension study of **CA-LoRA (Concept-Aware LoRA)** on
SDXL / Cityscapes. Every adapter here is a **head-sliced** LoRA: instead of adapting whole
attention projection matrices, it adapts only the ~2% of *individual attention heads* that a
concept-attribution (CA) measurement picked out at a chosen diffusion timestep.
These are the exact adapters that produced the numbers in Tables 3 and 5 of the report, and the
images in the companion dataset repo
[`chs35/improving-ca-lora-artifacts`](https://huggingface.co/datasets/chs35/improving-ca-lora-artifacts).
## ⚠️ Custom format — not loadable with `peft` or `diffusers`
Off-the-shelf adapter libraries cannot express attachment to a *partial slice* of a weight matrix,
so these checkpoints use a custom layout. `PeftModel.from_pretrained(...)` and
`pipe.load_lora_weights(...)` will **not** work.
Each `.safetensors` file contains, per adapted projection:
| tensor | dtype | meaning |
| --- | --- | --- |
| `<module>.lora_A` | f32 | LoRA down-projection, shaped to the selected slice |
| `<module>.lora_B` | f32 | LoRA up-projection, shaped to the selected slice |
| `<module>.lora_head_indices` | **int64** | which attention heads of that projection this adapter covers — the source of truth at load time |
plus safetensors metadata: `granularity` (`head`), `axis` (`style` / `viewpoint`), `strategy`,
`lora_rank` (64), `lora_alpha` (64), `steps` (10000).
Slice direction follows the single definition in the code (`HEAD_SLICE_DIMS`): **type A** =
`attn.to_out.0`, sliced along its *input columns*; **type B** = `to_q` / `to_k` / `to_v`, sliced
along their *output rows*. Partial selection matrices are never silently widened to the full
matrix, and a load under a mismatched `granularity` or `axis` is refused rather than coerced.
### Loading
Use the `calora` library from the **Improving-CA-LoRA** code repository:
```bash
python scripts/generate_images.py \
--lora <RUN>/lora-final.safetensors \
--run-name <RUN> \
--mode both
```
(`calora.train.lora.load_lora` is the underlying entry point.)
## Runs, and the report rows they back
Each `<RUN>/` directory is one training arm. `_style` / `_viewpoint` is the **concept axis** the CA
measurement was taken on; the letter/timestep prefix is the **timestep selection criterion**.
| directory | selection criterion | axis | report Table 3 row |
| --- | --- | --- | --- |
| `A_t81_style` | t = 81 (the original paper's timestep) | style | "t = 81 (paper) / style" |
| `A_t81_viewpoint` | t = 81 (the original paper's timestep) | viewpoint | "t = 81 (paper) / viewpoint" |
| `B_top3_style` | multi-t, aggregated over the top-3 timesteps [41, 1, 81] | style | "multi-t [41, 1, 81] / style" |
| `B_top3_viewpoint` | multi-t, aggregated over the top-3 timesteps [41, 1, 81] | viewpoint | "multi-t [41, 1, 81] / viewpoint" |
| `C_t41_style` | t = 41 (this study's top-ranked timestep) | style | "t = 41 / style" |
| `C_t41_viewpoint` | t = 41 (this study's top-ranked timestep) | viewpoint | "t = 41 / viewpoint" |
| `D_t1_style` | t = 1 | style | "t = 1 / style" |
| `D_t1_viewpoint` | t = 1 | viewpoint | "t = 1 / viewpoint" |
| `E_random` | **control**: a random 2% of heads, no CA measurement | — | "random 2% (no CA)" |
| `E_all` | **control**: all attention projections, full LoRA, no CA selection | — | "full attention (no CA)" |
The same ten directories back the drift rows of **Table 5** (per-snapshot Jaccard / Spearman of the
selected set against the trained set). The report's third control, `control_base` (0%, base SDXL
with no adapter), needs no checkpoint and so has none here; its generated images are in the dataset
repo.
## Checkpoints per run
Five files per run — intermediate checkpoints every 2,000 steps plus the final one at step 10,000:
```
<RUN>/lora-step002000.safetensors
<RUN>/lora-step004000.safetensors
<RUN>/lora-step006000.safetensors
<RUN>/lora-step008000.safetensors
<RUN>/lora-final.safetensors # step 10000 — the adapter evaluated in Table 3
```
`lora-final.safetensors` is the one used for every reported metric; the intermediates exist so that
the CA-drift-during-training analysis can be re-run at other points on the trajectory.
## `pilot_module_granularity/A_t81/` — granularity ablation
An **undocumented earlier pilot run at *module* granularity** (whole projection matrices selected,
not individual heads), kept here as a granularity-ablation data point. It is *not* the source of any
number in the report — the reported `A_t81_*` runs are the head-granularity ones above.
It differs in format as well as in granularity: it took the `peft` path (`granularity=module` is the
only mode `peft` supports here), so its keys are peft-style `...lora_A.default.weight`, it has no
`lora_head_indices` tensors, and its metadata predates the `granularity`/`axis` fields. 20 tensors
vs. the 168–1680 of the head-sliced runs; 34 MB for all five checkpoints.
## Training setup
SDXL base 1.0 in **fp32**; Cityscapes `train` split (2,975 images) at 1024²; **10,000 steps**,
batch size 1, constant lr **1e-4** (AdamW, β = 0.9/0.999, weight decay 0.01, grad clip 1.0), LoRA
**rank 64 / alpha 64**, diffusion loss only, random crop + random horizontal flip, seed 0. Selection
budget: 2% of attention units by unit count. Full config is released as `config.yaml` in the code
repository's `results/sdxl_cityscapes/`.
## Intended use and limitations
Research artifacts for reproducing and auditing the report — not production models. They are
fine-tuned on **Cityscapes** (German urban street scenes, daytime, vehicle-mounted camera), so they
absorb that domain narrowly and measurably lose prompt controllability relative to base SDXL; that
loss is the subject of the study, not a defect to work around. Anyone using these adapters must
respect the [Cityscapes dataset terms](https://www.cityscapes-dataset.com/license/) as well as the
SDXL [OpenRAIL++](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/blob/main/LICENSE.md)
use restrictions inherited from the base model.
## Verifying integrity
`results/sha256_checkpoints.txt` in the code repository lists the sha256 of all 55 files here,
keyed by the same repo-relative paths.
|