Upload README.md with huggingface_hub
Browse files
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
|
| 4 |
-
-
|
| 5 |
-
-
|
| 6 |
|
| 7 |
-
A small mapping `f` is learned from the
|
| 8 |
-
`(
|
|
|
|
| 9 |
|
| 10 |
-
Inspired by Sakana AI's **Text-to-LoRA** hypernetwork (arXiv 2506.06105) and
|
| 11 |
-
|
| 12 |
-
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 23 |
-
|
|
| 24 |
-
|
|
| 25 |
-
|
|
| 26 |
-
|
| 27 |
-
## Mapping function `f`
|
| 28 |
|
| 29 |
-
|
| 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 |
-
|
| 34 |
-
|
| 35 |
-
3. Predict `Ŷ_D = mean(Y) + α · Y_c`.
|
| 36 |
|
| 37 |
-
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
|
|
|
|
| 41 |
|
| 42 |
-
|
| 43 |
|
| 44 |
-
|
|
| 45 |
-
|---|---:|
|
| 46 |
-
|
|
| 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 |
-
|
| 67 |
-
plus a small negative pull away from the SST-2 anchor).
|
| 68 |
|
| 69 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 70 |
|
| 71 |
-
|
| 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 |
-
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 87 |
|
| 88 |
```
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
Y/
|
| 92 |
-
Y/
|
| 93 |
-
|
| 94 |
-
|
| 95 |
-
|
| 96 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 104 |
```
|
| 105 |
|
| 106 |
-
## Use
|
| 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 |
-
|
|
|
|
|
|
|
| 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
|