पाठ 10 / 27
LoRA और Parameter-Efficient Fine-Tuning
पूरे weight matrix की जगह छोटा low-rank अद्यतन प्रशिक्षित करें।
आवश्यक बदलाव अक्सर low-rank होता है
हर weight अद्यतन करना महँगा है और प्रति कार्य पूरे मॉडल की प्रति देता है। Parameter-efficient fine-tuning (PEFT) सिर्फ़ थोड़े जोड़े parameters प्रशिक्षित करती है। LoRA चुने हर weight matrix W को जमा देता है और ΔW = A·B सीखता है, जहाँ A d×r और B r×d है और rank r छोटा है (जैसे 8)। यह d² की जगह r·2d प्रशिक्षण-योग्य संख्याएँ हैं: d = 4096, r = 8 के लिए प्रति matrix 1.68 करोड़ की जगह लगभग 65,000 (0.39%)। दाँव है कि आवश्यक बदलाव low-rank है। Adapter megabytes का है, एक बेस मॉडल पर प्रति कार्य बदला या serving के लिए merge हो सकता है; QLoRA जमा बेस को 4-bit रूप में भी रखता है।
Rank-2 adapter से rank-2 बदलाव पाना, चलाकर
मैंने यह Python, numpy 2.5.3 और scikit-learn 1.9.1 के साथ, निश्चित random seeds पर चलाया। यह छोटे क्लासिकल मॉडल प्रशिक्षित करता है, language model नहीं: तंत्र (gradient descent, learning rate, overfitting, forgetting, low-rank अद्यतन) वही विचार हैं जो LLM के fine-tuning पर लागू होते हैं, पर संख्याएँ LLM के परिणाम नहीं हैं। जमा 32x32 matrix को rank-2 बदलाव चाहिए। बिना अद्यतन त्रुटि 3.29; rank-1 adapter (64 संख्याएँ, 6%) इसे 1.64 करता है; rank 2 (128 संख्याएँ, 12%) लगभग शून्य पर पहुँचता है, और rank 4 भी शून्य पर।
import numpy as np
rng = np.random.default_rng(1)
d, true_rank = 32, 2
W0 = rng.normal(size=(d, d)) / np.sqrt(d) # frozen "pretrained" weight
delta = (rng.normal(size=(d, true_rank)) @ rng.normal(size=(true_rank, d))) * 0.3 # the change the new task needs: rank 2
W_target = W0 + delta
X = rng.normal(size=(400, d)); Y = X @ W_target.T # task data produced by the target weights
def train_lora(r, steps=3000, lr=0.02):
A = rng.normal(size=(d, r)) * 0.1; B = np.zeros((r, d)) # W = W0 + A @ B ; only A and B are trained
for _ in range(steps):
W = W0 + A @ B
err = X @ W.T - Y # (n, d)
G = err.T @ X / len(X) # gradient wrt W (d, d)
gA, gB = G @ B.T, A.T @ G
A -= lr * gA; B -= lr * gB
return float(np.mean((X @ (W0 + A @ B).T - Y) ** 2)), A.size + B.size
print("frozen W0 only: loss", round(float(np.mean((X @ W0.T - Y) ** 2)), 4), "| full matrix has", d * d, "parameters")
for r in (1, 2, 4):
loss, params = train_lora(r)
print(f"LoRA rank {r}: loss {loss:.5f} with {params} trainable parameters ({params / (d * d):.0%} of the full matrix)")
Output:
frozen W0 only: loss 3.2905 | full matrix has 1024 parameters LoRA rank 1: loss 1.63657 with 64 trainable parameters (6% of the full matrix) LoRA rank 2: loss 0.00000 with 128 trainable parameters (12% of the full matrix) LoRA rank 4: loss 0.00000 with 256 trainable parameters (25% of the full matrix)
प्रशिक्षण memory: पूर्ण fine-tune बनाम LoRA, चलाकर
मैंने यह सादे Python 3 (सिर्फ़ standard library) से उदाहरण संख्याओं के साथ चलाया। प्रति प्रशिक्षित parameter 16 bytes के मोटे अनुमान से 7B मॉडल को पूरी तरह fine-tune करने को लगभग 112 GB चाहिए, पर 0.4% adapter प्रशिक्षित और base जमा (2 bytes) हो तो लगभग 14 GB। Activations और overhead से पहले के मोटे योजना-अंक।
# Rough training memory (mixed precision with Adam): weights 2 B + gradients 2 B + fp32 master weights 4 B + 2 optimizer states 8 B
# = about 16 bytes per TRAINED parameter, plus frozen weights at 2 B per parameter. Activations come on top.
def gb(x): return x / 1e9
def full_ft(params): return params * 16
def lora(params, trainable): return params * 2 + trainable * 16 # frozen base at 16-bit + trained adapters
for name, params in (("7B", 7e9), ("13B", 13e9), ("70B", 70e9)):
trainable = params * 0.004 # ~0.4% trainable with a small-rank adapter (example)
print(f"{name:4} full fine-tune ~ {gb(full_ft(params)):7.0f} GB LoRA ~ {gb(lora(params, trainable)):6.0f} GB (before activations)")
Output:
7B full fine-tune ~ 112 GB LoRA ~ 14 GB (before activations) 13B full fine-tune ~ 208 GB LoRA ~ 27 GB (before activations) 70B full fine-tune ~ 1120 GB LoRA ~ 144 GB (before activations)
त्वरित जाँच: LoRA क्या प्रशिक्षित करती है?
- जमा weights में जोड़े छोटे low-rank matrices
- मॉडल का हर weight
- सिर्फ़ tokenizer
- User interface
Answer
जमा weights में जोड़े छोटे low-rank matrices — Adapter संक्षिप्त अद्यतन सीखता है जबकि बेस जमा रहता है।