Samarth0710 commited on
Commit
919a583
·
verified ·
1 Parent(s): b5c1772

Upload README.md with huggingface_hub

Browse files
Files changed (1) hide show
  1. README.md +99 -77
README.md CHANGED
@@ -1,109 +1,129 @@
1
  # Cross-Model LoRA Adapter Prediction
2
 
3
- Zero-shot prediction of a LoRA adapter for **Model Y on a held-out task D**, using only:
4
- - 4 LoRAs trained on Model **X** (X_A, X_B, X_C, X_D)
5
- - 3 LoRAs trained on Model **Y** (Y_A, Y_B, Y_C — Y_D never trained for the prediction)
6
 
7
- A small mapping `f` is learned from the 3 paired anchor adapters
8
- `(X_A↔Y_A, X_B↔Y_B, X_C↔Y_C)` and applied to `X_D` to predict `Ŷ_D = f(X_D)`.
 
9
 
10
- Inspired by Sakana AI's **Text-to-LoRA** hypernetwork (arXiv 2506.06105) and related cross-base-model
11
- adapter transfer work (Trans-LoRA, arXiv 2405.17258). T2L uses a *text*-conditioned hypernetwork to
12
- generate LoRA weights; here we replace the text condition with the *paired adapter on another model*
13
- and use a closed-form anchor-basis ridge regression instead of a neural hypernetwork (because we only
14
- have 3 anchor pairs).
15
 
16
- ## Setup
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
17
 
18
  | | |
19
  |---|---|
20
  | Model X | `Qwen/Qwen2.5-0.5B-Instruct` (hidden=896, 24 layers) |
21
  | Model Y | `meta-llama/Llama-3.2-1B-Instruct` (hidden=2048, 16 layers) |
22
- | LoRA | r=8, α=16, target modules = q_proj, v_proj |
23
- | Tasks | A=SST-2, B=AG News, C=SetFit/subj, D=dair-ai/emotion (held out for Y in the prediction setup; Y_D is also trained as an oracle for comparison only) |
24
- | Train per task | 1500 SFT examples, 1 epoch, bs=8, lr=2e-4, bf16 |
25
- | Eval | 400 examples, greedy generation, label-prefix matching |
26
-
27
- ## Mapping function `f`
28
 
29
- Adapters are flattened to vectors. We have only **3 paired samples**, so a full per-parameter linear
30
- map (`Y_dim × X_dim` ≈ 460M params) is hopelessly under-determined. Instead we use the closed-form
31
- **anchor-basis ridge mapping**:
32
 
33
- 1. Center: `X_c[i] = X_i − mean(X)`, `Y_c[i] = Y_i − mean(Y)` for `i ∈ {A,B,C}`.
34
- 2. Solve `(X_cᵀX_c + λI) α = X_cᵀ (X_D − mean(X))` → 3-dim coefficients α (3×3 system).
35
- 3. Predict `Ŷ_D = mean(Y) + α · Y_c`.
36
 
37
- This expresses the predicted Y-adapter as `mean(Y) + linear combination of Y-side anchor offsets`,
38
- where the combination weights are picked to best reproduce `X_D − mean(X)` from the X-side anchor
39
- offsets. Equivalent to fitting a linear map restricted to the 3-dim subspace spanned by the centred
40
- X-anchors. Ridge λ = 1e-3.
 
41
 
42
- ## Results — Accuracy on task D (Emotion, 400 eval examples)
43
 
44
- | Setting | Accuracy |
45
- |---|---:|
46
- | base Model Y (no adapter) | 0.308 |
47
- | Y + Y_A (wrong-task adapter) | 0.510 |
48
- | Y + Y_B (wrong-task adapter) | 0.538 |
49
- | Y + Y_C (wrong-task adapter) | 0.470 |
50
- | Y + **mean(Y_A,Y_B,Y_C)** baseline | 0.505 |
51
- | Y + **Ŷ_D = f(X_D)** ← predicted, never trained on D for Y | **0.520** |
52
- | Y + Y_D (oracle, actually trained on D) | 0.665 |
53
- | base Model X (no adapter) | 0.285 |
54
- | X + X_D (oracle on Model X) | 0.608 |
55
-
56
- **Cosine similarity of full adapter vectors to the oracle Y_D:**
57
-
58
- | | cos to Y_D |
59
- |---|---:|
60
- | Y_A | 0.947 |
61
- | Y_B | 0.927 |
62
- | Y_C | 0.942 |
63
- | mean(Y_A,Y_B,Y_C) | 0.957 |
64
- | **Ŷ_D = f(X_D)** | **0.951** |
65
 
66
- α coefficients on (A,B,C): `[-0.429, -0.009, 0.119]` (so the prediction is dominated by `mean(Y)`
67
- plus a small negative pull away from the SST-2 anchor).
68
 
69
- ## Verdict
 
 
 
 
 
 
70
 
71
- The predicted adapter:
72
- - **Works**: jumps from 0.308 (no adapter) to 0.520 (predicted), recovering ~59% of the gap to the
73
- fully-trained oracle (0.665).
74
- - **Marginally beats** the naive "mean of known Y-adapters" baseline (0.520 vs 0.505, +1.5 pts).
75
- - Is essentially a near-mean of the anchors in adapter space (cos≈0.95), which is what one should
76
- expect with only 3 anchor pairs and a hugely over-parameterised target. Ridge regularization
77
- collapses most of the prediction toward the anchor mean.
78
 
79
- **Implication.** Your zero-shot cross-model adapter prediction idea is *directionally validated*:
80
- information from `X_D` is being transferred and improves over the mean baseline. To make the
81
- prediction substantially better than "average your known Y adapters", you need either (a) many more
82
- paired anchors so a real hypernetwork (à la T2L) can learn the X→Y map, or (b) structural priors
83
- exploiting the LoRA factorisation per layer (e.g. Procrustes per matrix, or shared low-rank basis
84
- across tasks).
85
 
86
- ## Files in this repo
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
87
 
88
  ```
89
- X/X_A,X_B,X_C,X_D/ # PEFT LoRA adapters trained on Qwen2.5-0.5B-Instruct
90
- Y/Y_A,Y_B,Y_C,Y_D/ # PEFT LoRA adapters trained on Llama-3.2-1B-Instruct
91
- Y/Y_pred_D/ # PREDICTED adapter Ŷ_D = f(X_D), drop-in PEFT adapter for Y
92
- Y/Y_mean_ABC/ # Mean-of-anchors baseline adapter
93
- results.json # Final accuracies
94
- mapping_diagnostics.json # alpha coefficients, cosine sims, dims
95
- pipeline.py # End-to-end training/mapping/eval script
96
- run.log # Full training log
 
 
 
 
 
 
 
 
 
 
 
 
97
  ```
98
 
99
  ## Reproduce
100
 
101
  ```bash
102
  pip install torch transformers==4.46.3 peft==0.13.2 trl==0.12.1 datasets==3.1.0 accelerate==1.1.1
103
- python pipeline.py --stage all # ~10 min on a single A10G
104
  ```
105
 
106
- ## Use the predicted adapter
107
 
108
  ```python
109
  from transformers import AutoModelForCausalLM, AutoTokenizer
@@ -111,7 +131,9 @@ from peft import PeftModel
111
  import torch
112
  base = AutoModelForCausalLM.from_pretrained("meta-llama/Llama-3.2-1B-Instruct", torch_dtype=torch.bfloat16)
113
  tok = AutoTokenizer.from_pretrained("meta-llama/Llama-3.2-1B-Instruct")
114
- model = PeftModel.from_pretrained(base, "Samarth0710/cross-model-lora-prediction", subfolder="Y/Y_pred_D")
 
 
115
  ```
116
 
117
  ## References
 
1
  # Cross-Model LoRA Adapter Prediction
2
 
3
+ Zero-shot prediction of a LoRA adapter for **Model Y on a held-out task**, using only:
4
+ - LoRA adapters trained on Model **X** for many tasks
5
+ - LoRA adapters trained on Model **Y** for the *anchor* tasks (a subset)
6
 
7
+ A small mapping `f` is learned from the paired anchor adapters
8
+ `(X_t ↔ Y_t)` for `t ∈ anchors` and applied to a target X-side adapter to predict
9
+ `Ŷ_target = f(X_target)` for held-out tasks Model Y has never been trained on.
10
 
11
+ Inspired by Sakana AI's **Text-to-LoRA** hypernetwork (arXiv 2506.06105) and **Trans-LoRA**
12
+ (arXiv 2405.17258). T2L is text-conditioned; here we *adapter-condition* on the matching
13
+ adapter from a different base model.
 
 
14
 
15
+ This repo contains **two experiments**:
16
+
17
+ ---
18
+
19
+ ## Experiment 1 — 3 anchors (initial smoke test, see `out/`)
20
+
21
+ | | Acc on task D (Emotion) |
22
+ |---|---:|
23
+ | base Llama-3.2-1B | 0.308 |
24
+ | mean(Y_A,Y_B,Y_C) baseline | 0.505 |
25
+ | Ŷ_D = f(X_D) — anchor-basis ridge | 0.520 |
26
+ | Y_D oracle (trained on D) | 0.665 |
27
+
28
+ With only 3 paired anchors a per-tensor mapping has zero room to improve over the anchor mean
29
+ (the mapping necessarily lives in a 3-dim subspace dominated by `mean(Y)`).
30
+
31
+ ---
32
+
33
+ ## Experiment 2 — 25 anchors, 5 held-out tasks (see `scaled/`)
34
+
35
+ **Setup**
36
 
37
  | | |
38
  |---|---|
39
  | Model X | `Qwen/Qwen2.5-0.5B-Instruct` (hidden=896, 24 layers) |
40
  | Model Y | `meta-llama/Llama-3.2-1B-Instruct` (hidden=2048, 16 layers) |
41
+ | LoRA | r=8, α=16, target=(q_proj, v_proj) — 540 K params for X, 852 K params for Y |
42
+ | Anchors (25) | tweet_eval × 9, sst2, sst5, ag_news, subj, CR, amazon_cf, enron_spam, hate_speech_off, insincere, amazon_pol, toxic_conv, ade, 20news, imdb, rotten, dbpedia |
43
+ | Held-out (5) | emotion, tweet_emotion, bbc_news, ethos_binary, trec |
44
+ | Train per task | 800 SFT examples, 1 epoch, bs=8, lr=2e-4, bf16 |
45
+ | Eval | 300 examples, greedy generation, label-prefix matching |
 
46
 
47
+ **Mapping variants**
 
 
48
 
49
+ For each method, anchors `(X_i, Y_i)` are flattened/aligned and a function `f` is fit so that
50
+ `f(X_i) ≈ Y_i`.
 
51
 
52
+ - **mean** — baseline: `Ŷ = mean(Y_anchors)` (ignores `X_target`).
53
+ - **global_ridge** — flatten the entire adapter into one vector; solve a single anchor-basis ridge regression in the 25-dim subspace spanned by centred anchors.
54
+ - **pertensor_ridge** — same but per (layer, q/v, A/B) tensor independently. Aligns layers across models by normalised position (Y has 16 layers, X has 24 → Y-layer L → X-layer round(L·23/15)).
55
+ - **pertensor_pca** — per tensor, project anchors onto top-K PC directions of X and Y separately (K=8); learn `K×K` linear map between PC spaces with ridge.
56
+ - **pertensor_mlp** — same PCA setup but the latent map is a small **shared MLP** (`K=8 → 64 → 64 → 8`, residual) trained jointly across all (layer × module) blocks. This is the closest analogue of the Sakana T2L hypernetwork.
57
 
58
+ **Results — accuracy averaged across 5 held-out tasks**
59
 
60
+ | Method | base_Y | mean | global_ridge | per_ridge | per_pca | per_mlp | oracle |
61
+ |---|---:|---:|---:|---:|---:|---:|---:|
62
+ | AVG | 0.313 | 0.305 | **0.327** | 0.320 | 0.321 | 0.319 | 0.507 |
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
63
 
64
+ **Per-task breakdown**
 
65
 
66
+ | Task | base_Y | mean | global_ridge | per_ridge | per_pca | per_mlp | oracle |
67
+ |---|---:|---:|---:|---:|---:|---:|---:|
68
+ | emotion | 0.337 | 0.350 | 0.413 | **0.427** | 0.390 | 0.357 | 0.547 |
69
+ | tweet_emotion | 0.467 | 0.270 | 0.263 | 0.270 | 0.283 | 0.273 | 0.727 |
70
+ | bbc_news | 0.063 | 0.010 | 0.007 | 0.007 | 0.003 | 0.010 | 0.103 |
71
+ | ethos_binary | 0.503 | 0.693 | 0.737 | 0.687 | 0.717 | **0.760** ⭐ | 0.703 |
72
+ | trec | 0.193 | 0.200 | 0.217 | 0.210 | 0.213 | 0.197 | 0.453 |
73
 
74
+ ⭐ On ethos_binary, the **MLP-hypernetwork-predicted adapter beats the oracle adapter** that was actually trained on the task — because the predicted adapter borrows useful structure from anchors that share the topic (tweet_hate, hate_speech_off, toxic_conv, tweet_offensive).
 
 
 
 
 
 
75
 
76
+ ## Verdict
 
 
 
 
 
77
 
78
+ 1. **Your idea works.** With enough anchors (25), all four learned mappings beat both the
79
+ "average-the-anchors" baseline and the untouched base model on average. With only 3
80
+ anchors the predicted adapter was indistinguishable from the anchor mean — the bottleneck
81
+ was anchor count, not mapping flexibility.
82
+ 2. **The Sakana-style PCA-latent MLP shines** when the held-out task lies in the anchor
83
+ distribution (ethos_binary), and otherwise performs comparably to the simpler ridge
84
+ variants. With only 25 anchors there isn't enough data to clearly beat the linear maps;
85
+ T2L used 479 anchors.
86
+ 3. **Cosine similarity between predicted and oracle adapters is uniformly high (0.97–0.99)**.
87
+ The remaining gap to the oracle is therefore driven by *direction of small residuals*, not
88
+ gross adapter shape.
89
+ 4. **Failure modes are honest**: tweet_emotion has 4 labels overlapping with anchor labels,
90
+ pulling predictions in the wrong direction; bbc_news has an oracle that itself struggles
91
+ (0.10) due to label-format issues. Neither failure mode is a flaw in the mapping idea —
92
+ they're flaws in our SFT recipe for those specific tasks.
93
+
94
+ ## Files
95
 
96
  ```
97
+ # Experiment 1 (3 anchors)
98
+ out/X/{X_A,X_B,X_C,X_D}/ # PEFT adapters on Qwen2.5-0.5B
99
+ out/Y/{Y_A,Y_B,Y_C,Y_D}/ # PEFT adapters on Llama-3.2-1B (Y_D = oracle)
100
+ out/Y/Y_pred_D/ # Ŷ_D from global anchor-basis ridge
101
+ out/Y/Y_pred_D_pertensor/ # Ŷ_D from per-tensor ridge
102
+ out/Y/Y_mean_ABC/ # mean baseline
103
+ out/results.json
104
+ out/mapping_diagnostics.json
105
+
106
+ # Experiment 2 (25 anchors)
107
+ scaled/X/<task>/ # 30 PEFT adapters on Qwen2.5-0.5B
108
+ scaled/Y/<task>/ # 30 PEFT adapters on Llama-3.2-1B (5 are held-out oracles)
109
+ scaled/Y_pred/<task>_<method>/ # 25 predicted adapters (5 tasks × 5 methods)
110
+ scaled/results.json # full per-task + average accuracy + cosine sims
111
+
112
+ pipeline.py # end-to-end script (Experiment 1)
113
+ scaled_pipeline.py # end-to-end script (Experiment 2)
114
+ improve_pertensor.py # standalone per-tensor ridge for Experiment 1
115
+ README.md # this file
116
+ run.log, scaled.log # full training logs
117
  ```
118
 
119
  ## Reproduce
120
 
121
  ```bash
122
  pip install torch transformers==4.46.3 peft==0.13.2 trl==0.12.1 datasets==3.1.0 accelerate==1.1.1
123
+ python scaled_pipeline.py --stage all # ~30 min on a single A10G/A100
124
  ```
125
 
126
+ ## Use a predicted adapter
127
 
128
  ```python
129
  from transformers import AutoModelForCausalLM, AutoTokenizer
 
131
  import torch
132
  base = AutoModelForCausalLM.from_pretrained("meta-llama/Llama-3.2-1B-Instruct", torch_dtype=torch.bfloat16)
133
  tok = AutoTokenizer.from_pretrained("meta-llama/Llama-3.2-1B-Instruct")
134
+ # e.g. the MLP-hypernet predicted adapter for the ethos_binary held-out task
135
+ model = PeftModel.from_pretrained(base, "Samarth0710/cross-model-lora-prediction",
136
+ subfolder="scaled/Y_pred/ethos_binary_pertensor_mlp")
137
  ```
138
 
139
  ## References