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
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.
β οΈ 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:
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 as well as the SDXL OpenRAIL++ 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.
Model tree for chs35/improving-ca-lora-checkpoints
Base model
stabilityai/stable-diffusion-xl-base-1.0