Simo76 commited on
Commit
ec1654a
·
1 Parent(s): 41c6456

Add Unified-LoRA benchmark script for GLUE tasks

Browse files

This script benchmarks the Unified-LoRA method for fine-tuning models on GLUE tasks, including MRPC, SST-2, CoLA, and RTE. It implements an adaptive per-layer rank controller for LoRA, allowing dynamic adjustments based on gradient stress trends.

Files changed (1) hide show
  1. benchmark.py +309 -0
benchmark.py ADDED
@@ -0,0 +1,309 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Unified-LoRA Benchmark
3
+ Adaptive per-layer rank controller for LoRA fine-tuning.
4
+
5
+ Runs 4 GLUE tasks (MRPC, SST-2, CoLA, RTE) comparing:
6
+ - Baseline: fixed rank=16
7
+ - Adaptive: per-layer gradient-stress rank controller
8
+
9
+ Requirements:
10
+ pip install transformers datasets evaluate accelerate scikit-learn
11
+
12
+ Hardware: GPU recommended (tested on T4, ~30 min total)
13
+ """
14
+
15
+ import copy, torch, time, gc
16
+ import torch.nn as nn
17
+ from datasets import load_dataset
18
+ from transformers import (
19
+ AutoTokenizer,
20
+ AutoModelForSequenceClassification,
21
+ DataCollatorWithPadding,
22
+ )
23
+ from torch.utils.data import DataLoader
24
+ import evaluate
25
+
26
+ # ================================================================
27
+ # CONFIG
28
+ # ================================================================
29
+ DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
30
+ MODEL_NAME = "distilbert-base-uncased"
31
+
32
+ BATCH_SIZE = 16
33
+ EPOCHS = 3
34
+ LR = 5e-4
35
+ MAX_RANK = 16
36
+ MIN_RANK = 4
37
+ ALPHA = 16
38
+ GRAD_CLIP = 1.0
39
+
40
+ TASKS = {
41
+ "mrpc": {"num_labels": 2, "metric_key": "f1",
42
+ "paired": True, "keys": ("sentence1", "sentence2")},
43
+ "sst2": {"num_labels": 2, "metric_key": "accuracy",
44
+ "paired": False, "keys": ("sentence",)},
45
+ "cola": {"num_labels": 2, "metric_key": "matthews_correlation",
46
+ "paired": False, "keys": ("sentence",)},
47
+ "rte": {"num_labels": 2, "metric_key": "accuracy",
48
+ "paired": True, "keys": ("sentence1", "sentence2")},
49
+ }
50
+
51
+ # ================================================================
52
+ # DATA
53
+ # ================================================================
54
+ def load_task(task_name):
55
+ cfg = TASKS[task_name]
56
+ tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
57
+ ds = load_dataset("glue", task_name)
58
+
59
+ if cfg["paired"]:
60
+ def preprocess(x):
61
+ return tokenizer(x[cfg["keys"][0]], x[cfg["keys"][1]], truncation=True)
62
+ else:
63
+ def preprocess(x):
64
+ return tokenizer(x[cfg["keys"][0]], truncation=True)
65
+
66
+ ds = ds.map(preprocess, batched=True)
67
+ ds = ds.rename_column("label", "labels")
68
+ ds.set_format(type="torch", columns=["input_ids", "attention_mask", "labels"])
69
+
70
+ collator = DataCollatorWithPadding(tokenizer)
71
+ train_loader = DataLoader(
72
+ ds["train"], batch_size=BATCH_SIZE, shuffle=True, collate_fn=collator
73
+ )
74
+ val_loader = DataLoader(
75
+ ds["validation"], batch_size=32, collate_fn=collator
76
+ )
77
+ metric = evaluate.load("glue", task_name)
78
+
79
+ return train_loader, val_loader, metric, cfg
80
+
81
+ # ================================================================
82
+ # LoRA MODULE — per-layer adaptive rank
83
+ # ================================================================
84
+ class LoRALinear(nn.Module):
85
+ """
86
+ LoRA adapter with:
87
+ - Per-layer gradient stress tracking (EMA)
88
+ - Dynamic rank adjustment based on stress trend
89
+ - Standard alpha/r scaling
90
+ """
91
+
92
+ def __init__(self, base, max_r=16, layer_name=""):
93
+ super().__init__()
94
+ self.base = copy.deepcopy(base)
95
+ for p in self.base.parameters():
96
+ p.requires_grad = False
97
+
98
+ self.max_r = max_r
99
+ self.layer_name = layer_name
100
+ self.A = nn.Parameter(torch.randn(max_r, base.in_features) * 0.01)
101
+ self.B = nn.Parameter(torch.zeros(base.out_features, max_r))
102
+ self.active_r = MIN_RANK
103
+
104
+ # Stress tracking
105
+ self.grad_ema = None
106
+ self.prev_grad_ema = None
107
+
108
+ def set_rank(self, r):
109
+ self.active_r = max(MIN_RANK, min(r, self.max_r))
110
+
111
+ def update_rank(self):
112
+ """Adapt rank based on gradient stress trend."""
113
+ if self.A.grad is None:
114
+ return
115
+
116
+ grad_norm = self.A.grad[:self.active_r].norm().item()
117
+
118
+ if self.grad_ema is None:
119
+ self.grad_ema = grad_norm
120
+ self.prev_grad_ema = grad_norm
121
+ return
122
+
123
+ self.prev_grad_ema = self.grad_ema
124
+ self.grad_ema = 0.9 * self.grad_ema + 0.1 * grad_norm
125
+
126
+ delta = self.grad_ema - self.prev_grad_ema
127
+ threshold = 0.01 * self.grad_ema if self.grad_ema > 0 else 0.01
128
+
129
+ if delta > threshold: # stress increasing -> more capacity
130
+ self.active_r = min(self.max_r, self.active_r + 2)
131
+ elif delta < -threshold: # stress decreasing -> reduce
132
+ self.active_r = max(MIN_RANK, self.active_r - 2)
133
+
134
+ def forward(self, x):
135
+ base_out = self.base(x)
136
+ A = self.A[:self.active_r]
137
+ B = self.B[:, :self.active_r]
138
+ lora_out = x @ A.t() @ B.t()
139
+ scale = ALPHA / self.active_r
140
+ return base_out + scale * lora_out
141
+
142
+ # ================================================================
143
+ # HELPERS
144
+ # ================================================================
145
+ def inject_lora(model):
146
+ for i, layer in enumerate(model.distilbert.transformer.layer):
147
+ layer.attention.q_lin = LoRALinear(
148
+ layer.attention.q_lin, MAX_RANK, layer_name=f"layer{i}.q"
149
+ )
150
+ layer.attention.v_lin = LoRALinear(
151
+ layer.attention.v_lin, MAX_RANK, layer_name=f"layer{i}.v"
152
+ )
153
+ return model
154
+
155
+
156
+ def get_lora_modules(model):
157
+ return [m for m in model.modules() if isinstance(m, LoRALinear)]
158
+
159
+
160
+ def setup_trainable(model):
161
+ for p in model.parameters():
162
+ p.requires_grad = False
163
+ for m in get_lora_modules(model):
164
+ m.A.requires_grad = True
165
+ m.B.requires_grad = True
166
+ for n, p in model.named_parameters():
167
+ if "classifier" in n or "pre_classifier" in n:
168
+ p.requires_grad = True
169
+ return model
170
+
171
+
172
+ def evaluate_model(model, val_loader, metric):
173
+ model.eval()
174
+ preds, labels = [], []
175
+ with torch.no_grad():
176
+ for batch in val_loader:
177
+ batch = {k: v.to(DEVICE) for k, v in batch.items()}
178
+ logits = model(**batch).logits
179
+ p = torch.argmax(logits, dim=1)
180
+ preds += p.cpu().tolist()
181
+ labels += batch["labels"].cpu().tolist()
182
+ return metric.compute(predictions=preds, references=labels)
183
+
184
+ # ================================================================
185
+ # TRAINING
186
+ # ================================================================
187
+ def train(task_name, adaptive=True):
188
+ train_loader, val_loader, metric, cfg = load_task(task_name)
189
+
190
+ model = AutoModelForSequenceClassification.from_pretrained(
191
+ MODEL_NAME, num_labels=cfg["num_labels"]
192
+ )
193
+ model = inject_lora(model)
194
+
195
+ if not adaptive:
196
+ for m in get_lora_modules(model):
197
+ m.set_rank(MAX_RANK)
198
+
199
+ model = setup_trainable(model).to(DEVICE)
200
+
201
+ opt = torch.optim.AdamW(
202
+ filter(lambda p: p.requires_grad, model.parameters()), lr=LR
203
+ )
204
+
205
+ rank_history = {m.layer_name: [] for m in get_lora_modules(model)}
206
+
207
+ t0 = time.time()
208
+
209
+ for epoch in range(EPOCHS):
210
+ model.train()
211
+ for step, batch in enumerate(train_loader):
212
+ batch = {k: v.to(DEVICE) for k, v in batch.items()}
213
+
214
+ loss = model(**batch).loss
215
+ loss.backward()
216
+ torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)
217
+
218
+ if adaptive:
219
+ for m in get_lora_modules(model):
220
+ m.update_rank()
221
+ rank_history[m.layer_name].append(m.active_r)
222
+
223
+ opt.step()
224
+ opt.zero_grad()
225
+
226
+ elapsed = time.time() - t0
227
+ res = evaluate_model(model, val_loader, metric)
228
+
229
+ # Stats
230
+ all_ranks = []
231
+ layer_avg = {}
232
+ for name, ranks in rank_history.items():
233
+ if ranks:
234
+ layer_avg[name] = sum(ranks) / len(ranks)
235
+ all_ranks.extend(ranks)
236
+
237
+ global_avg_rank = sum(all_ranks) / len(all_ranks) if all_ranks else MAX_RANK
238
+
239
+ if adaptive:
240
+ print(f"\n Per-layer rank ({task_name}):")
241
+ for name in sorted(layer_avg.keys()):
242
+ print(f" {name}: {layer_avg[name]:.1f}")
243
+
244
+ del model, opt
245
+ gc.collect()
246
+ if torch.cuda.is_available():
247
+ torch.cuda.empty_cache()
248
+
249
+ return {**res, "avg_rank": global_avg_rank, "time": elapsed}
250
+
251
+ # ================================================================
252
+ # RUN
253
+ # ================================================================
254
+ def main():
255
+ results = {}
256
+
257
+ for task_name in TASKS:
258
+ print(f"\n{'='*50}")
259
+ print(f" {task_name.upper()}")
260
+ print(f"{'='*50}")
261
+
262
+ results[task_name] = {}
263
+
264
+ print(f"\n Baseline (fixed rank=16)...")
265
+ results[task_name]["baseline"] = train(task_name, adaptive=False)
266
+
267
+ print(f"\n Adaptive (per-layer controller)...")
268
+ results[task_name]["adaptive"] = train(task_name, adaptive=True)
269
+
270
+ # Results table
271
+ print("\n" + "=" * 65)
272
+ print(" RESULTS")
273
+ print("=" * 65)
274
+
275
+ print(f"\n{'Task':<8} {'Method':<12} {'Metric':>10} {'Avg Rank':>10} {'Time':>8}")
276
+ print("-" * 50)
277
+
278
+ for task_name in TASKS:
279
+ metric_key = TASKS[task_name]["metric_key"]
280
+
281
+ for method in ["baseline", "adaptive"]:
282
+ r = results[task_name][method]
283
+ val = r.get(metric_key, r.get("accuracy", -1))
284
+ rank = r.get("avg_rank", -1)
285
+ t = r.get("time", -1)
286
+ print(f"{task_name:<8} {method:<12} {val:>10.4f} {rank:>10.1f} {t:>7.1f}s")
287
+ print()
288
+
289
+ # Summary
290
+ print("=" * 65)
291
+ print(" SUMMARY")
292
+ print("=" * 65)
293
+
294
+ for task_name in TASKS:
295
+ metric_key = TASKS[task_name]["metric_key"]
296
+ b = results[task_name]["baseline"]
297
+ a = results[task_name]["adaptive"]
298
+
299
+ b_val = b.get(metric_key, b.get("accuracy", 0))
300
+ a_val = a.get(metric_key, a.get("accuracy", 0))
301
+ a_rank = a.get("avg_rank", 16)
302
+
303
+ rank_red = 100 * (1 - a_rank / 16)
304
+
305
+ print(f" {task_name:<8} delta: {a_val - b_val:+.4f} rank: {a_rank:.1f}/16 reduction: {rank_red:.0f}%")
306
+
307
+
308
+ if __name__ == "__main__":
309
+ main()