headless-start commited on
Commit
883b1c1
·
verified ·
1 Parent(s): 1ad5d61

add trained checkpoints, results and model card

Browse files
.gitattributes CHANGED
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ results/pet_samples.png filter=lfs diff=lfs merge=lfs -text
README.md ADDED
@@ -0,0 +1,148 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: mit
3
+ library_name: timm
4
+ pipeline_tag: image-classification
5
+ base_model: timm/vit_base_patch16_224.augreg_in21k_ft_in1k
6
+ datasets:
7
+ - timm/oxford-iiit-pet
8
+ tags:
9
+ - lora
10
+ - parameter-efficient-fine-tuning
11
+ - vision-transformer
12
+ - pytorch
13
+ metrics:
14
+ - accuracy
15
+ ---
16
+
17
+ # LoRA Fine-Tuning of ViT-B/16 on Oxford-IIIT Pets
18
+
19
+ Trained checkpoints for [github.com/headless-start/peft-lora-vit](https://github.com/headless-start/peft-lora-vit),
20
+ a hand-written LoRA implementation on a frozen ViT-B/16 (`vit_base_patch16_224`, timm).
21
+ LoRA matrices are added to the attention query and value projections
22
+ (`alpha = 2r`, `B` initialised to zero) and only they and the classification
23
+ head are trained.
24
+
25
+ The repository holds every checkpoint behind the results in the GitHub README:
26
+ the headline run, the linear-probe / LoRA / full fine-tuning comparison, the
27
+ placement study and the rank study. Code, training scripts and figures live on GitHub.
28
+
29
+ ![Dataset samples](results/pet_samples.png)
30
+
31
+ ## Results
32
+
33
+ Top-1 accuracy on the Oxford-IIIT Pets test split (3,669 images, 37 breeds).
34
+
35
+ **Headline run** (LoRA rank 8 on q and v, 25 epochs): **95.2%** with 323K trainable
36
+ parameters out of 86.1M (0.38%). File: `checkpoints/best.pt`.
37
+
38
+ ### Baselines
39
+
40
+ | Method | Accuracy | Trainable parameters | Checkpoint file | Size |
41
+ |---|---|---|---|---|
42
+ | Linear probe | 93.5% | 28K (0.03%) | `checkpoints/best_head.pt` | 0.1 MB |
43
+ | LoRA r=8, q+v | **94.9%** | 323K (0.38%) | `checkpoints/best_lora.pt` | 1.3 MB |
44
+ | Full fine-tuning | 93.9% | 85.8M (100%) | `checkpoints/best_full.pt` | 343 MB |
45
+
46
+ ![Baselines](results/baselines.png)
47
+
48
+ ### Placement study (rank 8)
49
+
50
+ | Placement | Accuracy | Trainable parameters | Checkpoint file |
51
+ |---|---|---|---|
52
+ | q | 94.3% | 176K | `checkpoints/best_r8_q.pt` |
53
+ | k | 94.1% | 176K | `checkpoints/best_r8_k.pt` |
54
+ | v | 94.7% | 176K | `checkpoints/best_r8_v.pt` |
55
+ | q + k | 94.3% | 323K | `checkpoints/best_r8_qk.pt` |
56
+ | q + v | **94.9%** | 323K | `checkpoints/best_r8_qv.pt` |
57
+ | q + k + v | 94.7% | 471K | `checkpoints/best_r8_qkv.pt` |
58
+
59
+ ![Placement study](results/placement.png)
60
+
61
+ ### Rank study (q + v)
62
+
63
+ | Rank | Accuracy | Trainable parameters | Checkpoint file |
64
+ |---|---|---|---|
65
+ | 4 | 94.8% | 176K | `checkpoints/best_r4_qv.pt` |
66
+ | 8 | 94.9% | 323K | `checkpoints/best_r8_qv.pt` |
67
+ | 16 | 94.6% | 618K | `checkpoints/best_r16_qv.pt` |
68
+ | 32 | 94.9% | 1.21M | `checkpoints/best_r32_qv.pt` |
69
+
70
+ ![Rank study](results/ablation.png)
71
+
72
+ ### Notes on the numbers
73
+
74
+ - `best_lora.pt` and `best_r8_qv.pt` are the same weights; the same run appears in
75
+ the baseline, placement and rank tables.
76
+ - `best.pt` (95.2%) is a separate run of the same configuration. The 0.3-point gap
77
+ to 94.9% is the run-to-run variation described in the GitHub README.
78
+ - Every number is a single run with seed 42. Each checkpoint is the epoch with the
79
+ highest accuracy on the test split, which is also the split reported here, so
80
+ the figures are best-epoch results rather than estimates from a held-out validation set.
81
+ - All checkpoints were re-evaluated on the test split before upload and reproduce the stored accuracies.
82
+
83
+ ## Files
84
+
85
+ ```text
86
+ checkpoints/
87
+ best.pt headline run, LoRA r=8 on q+v
88
+ best_head.pt linear probe (classification head only)
89
+ best_lora.pt LoRA r=8 on q+v, as used in the comparison tables
90
+ best_full.pt full fine-tuning (all weights)
91
+ best_r8_<placement>.pt placement study
92
+ best_r<rank>_qv.pt rank study
93
+ results/ the JSON results and figures from the GitHub repository
94
+ ```
95
+
96
+ The LoRA and linear-probe checkpoints store only the trained tensors (LoRA
97
+ matrices and head); the frozen backbone comes from the public timm weights.
98
+ `best_full.pt` stores the whole network. Every file is a PyTorch dictionary
99
+ with the keys `model`, `epoch` and `val_acc`.
100
+
101
+ ## Usage
102
+
103
+ Clone the code, download a checkpoint and run the prediction script:
104
+
105
+ ```bash
106
+ git clone https://github.com/headless-start/peft-lora-vit.git
107
+ cd peft-lora-vit
108
+ pip install -r requirements.txt
109
+
110
+ hf download headless-start/peft-lora-vit checkpoints/best.pt --local-dir .
111
+ python predict.py path/to/pet.jpg --ckpt checkpoints/best.pt
112
+ ```
113
+
114
+ For another LoRA checkpoint pass its rank and placement, for example
115
+ `--ckpt checkpoints/best_r16_qv.pt --lora-r 16` or
116
+ `--ckpt checkpoints/best_r8_k.pt --placement k`.
117
+
118
+ In Python:
119
+
120
+ ```python
121
+ import torch
122
+ from huggingface_hub import hf_hub_download
123
+ from predict import load_model
124
+
125
+ path = hf_hub_download("headless-start/peft-lora-vit", "checkpoints/best.pt")
126
+ model = load_model(path, "vit_base_patch16_224", r=8, alpha_factor=2,
127
+ device=torch.device("cpu"), placement="qv")
128
+ ```
129
+
130
+ Inputs are RGB images resized to 256, centre-cropped to 224 and normalised with
131
+ ImageNet statistics (`build_transforms` in `src/data.py`).
132
+
133
+ ## Training setup
134
+
135
+ | Setting | Value |
136
+ |---|---|
137
+ | Backbone | `vit_base_patch16_224` (timm, ImageNet pretrained), frozen for LoRA and the linear probe |
138
+ | Data | Oxford-IIIT Pets, `trainval` split for training, `test` split for evaluation |
139
+ | Epochs | 25 |
140
+ | Optimiser | AdamW, learning rate 3e-4 (3e-5 for full fine-tuning), weight decay 0.05 |
141
+ | Schedule | 2 warmup epochs, then cosine decay to 1e-7 |
142
+ | Batch size | 64 (16 for full fine-tuning) |
143
+ | Other | mixed precision, drop-path 0.1, random resized crop and horizontal flip |
144
+
145
+ ## Licence
146
+
147
+ Released under the MIT licence, as is the code. The pretrained backbone is
148
+ Apache-2.0 and Oxford-IIIT Pets is CC BY-SA 4.0; their terms continue to apply.
checkpoints/best.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:5bc0e71268fe7240473f7283cfb877bea169f3a465d915a9eef2c407ad95b779
3
+ size 1309645
checkpoints/best_full.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:37dfab4415663792f15786c71e259a6e9035a5fd44a08d88f1dbbe0568e993d4
3
+ size 343356619
checkpoints/best_head.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3e10e79b57066bca08495834431ab4d85e91be1c119233777be33b5ae7b0d6c6
3
+ size 115765
checkpoints/best_lora.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:331f8b5eeb057ad726c84ec404d4e8606971f0e8cc7c90e6a62c4bff9b36654d
3
+ size 1310333
checkpoints/best_r16_qv.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:7c8616f8c3d987b1a392be0aa272f27a69bfb6552f4ca319e7522c3f15df0c70
3
+ size 2490093
checkpoints/best_r32_qv.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:f58a6dee8264d39944b06fe76fd057eed29ba0c629fa62137c78da7f37e6488a
3
+ size 4849389
checkpoints/best_r4_qv.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:af70ae6e7010db860e9cd3f3c254bbcc0acc658374e1f38e36a937be711eb671
3
+ size 720565
checkpoints/best_r8_k.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:5134da558f4c31785e0a914745e2bf70a28f8104868611203cc0b5c0f0811390
3
+ size 712725
checkpoints/best_r8_q.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:dfd69e142156e9f3ec7c6a4ae5dbbd94dc47a97a20bb8cc7df35f7e315b1f789
3
+ size 712725
checkpoints/best_r8_qk.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d267261f6bb11b2de07e20de0e0e0b2cbb1256ab309deb4adb87be1314ba2cac
3
+ size 1310389
checkpoints/best_r8_qkv.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:9426fa4949649844f52db4daa4fcec143a7a0621661503f11801954dfc2b76c8
3
+ size 1908101
checkpoints/best_r8_qv.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:bce4374b5cb50af48a75e46a55b0420125c34982260c01f7b3d6687480fa5844
3
+ size 1310389
checkpoints/best_r8_v.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:8ebac657298058652deb78d68f41e7fbc1ca7b54ea0517afe67ce09ca4f97d89
3
+ size 712725
results/ablation.json ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [
2
+ {
3
+ "r": 4,
4
+ "placement": "qv",
5
+ "top1_acc": 0.9479,
6
+ "trainable_params": 175909,
7
+ "trainable_pct": 0.205
8
+ },
9
+ {
10
+ "r": 8,
11
+ "placement": "qv",
12
+ "top1_acc": 0.9488,
13
+ "trainable_params": 323365,
14
+ "trainable_pct": 0.375
15
+ },
16
+ {
17
+ "r": 16,
18
+ "placement": "qv",
19
+ "top1_acc": 0.946,
20
+ "trainable_params": 618277,
21
+ "trainable_pct": 0.715
22
+ },
23
+ {
24
+ "r": 32,
25
+ "placement": "qv",
26
+ "top1_acc": 0.949,
27
+ "trainable_params": 1208101,
28
+ "trainable_pct": 1.389
29
+ }
30
+ ]
results/ablation.png ADDED
results/baselines.json ADDED
@@ -0,0 +1,29 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [
2
+ {
3
+ "method": "head",
4
+ "top1_acc": 0.9351,
5
+ "trainable_params": 28453,
6
+ "trainable_pct": 0.033,
7
+ "ckpt_mb": 0.1,
8
+ "epoch_sec": 19.7,
9
+ "peak_vram_gb": 0.66
10
+ },
11
+ {
12
+ "method": "lora",
13
+ "top1_acc": 0.9488,
14
+ "trainable_params": 323365,
15
+ "trainable_pct": 0.375,
16
+ "ckpt_mb": 1.2,
17
+ "epoch_sec": 30.8,
18
+ "peak_vram_gb": 3.73
19
+ },
20
+ {
21
+ "method": "full",
22
+ "top1_acc": 0.9392,
23
+ "trainable_params": 85827109,
24
+ "trainable_pct": 100.0,
25
+ "ckpt_mb": 327.5,
26
+ "epoch_sec": 40.6,
27
+ "peak_vram_gb": 2.49
28
+ }
29
+ ]
results/baselines.png ADDED
results/metrics.json ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "backbone": "vit_base_patch16_224",
3
+ "dataset": "oxford_pets",
4
+ "epochs": 25,
5
+ "top1_acc": 0.9518,
6
+ "trainable_params": 323365,
7
+ "total_params": 86122021,
8
+ "trainable_pct": 0.375,
9
+ "lora": {
10
+ "r": 8,
11
+ "alpha_factor": 2,
12
+ "dropout": 0.0
13
+ }
14
+ }
results/pet_samples.png ADDED

Git LFS Details

  • SHA256: 47f680f127ecf1abd04608a5d6dead3c10cd52d985ff8c87e2305a78a14b9998
  • Pointer size: 132 Bytes
  • Size of remote file: 1.22 MB
results/placement.json ADDED
@@ -0,0 +1,44 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [
2
+ {
3
+ "r": 8,
4
+ "placement": "q",
5
+ "top1_acc": 0.943,
6
+ "trainable_params": 175909,
7
+ "trainable_pct": 0.205
8
+ },
9
+ {
10
+ "r": 8,
11
+ "placement": "k",
12
+ "top1_acc": 0.9411,
13
+ "trainable_params": 175909,
14
+ "trainable_pct": 0.205
15
+ },
16
+ {
17
+ "r": 8,
18
+ "placement": "v",
19
+ "top1_acc": 0.9471,
20
+ "trainable_params": 175909,
21
+ "trainable_pct": 0.205
22
+ },
23
+ {
24
+ "r": 8,
25
+ "placement": "qk",
26
+ "top1_acc": 0.943,
27
+ "trainable_params": 323365,
28
+ "trainable_pct": 0.375
29
+ },
30
+ {
31
+ "r": 8,
32
+ "placement": "qv",
33
+ "top1_acc": 0.9488,
34
+ "trainable_params": 323365,
35
+ "trainable_pct": 0.375
36
+ },
37
+ {
38
+ "r": 8,
39
+ "placement": "qkv",
40
+ "top1_acc": 0.9466,
41
+ "trainable_params": 470821,
42
+ "trainable_pct": 0.546
43
+ }
44
+ ]
results/placement.png ADDED
results/training_curve.png ADDED