How LoRA is Implemented for Efficient Transformer Fine-Tuning in PyTorch
LoRA (Low-Rank Adaptation) enables efficient transformer fine-tuning by freezing pre-trained weight matrices and injecting trainable low-rank matrices into linear projections, reducing trainable parameters by orders of magnitude while maintaining full model capacity.
The labmlai/annotated_deep_learning_paper_implementations repository demonstrates exactly how LoRA is implemented for efficient transformer fine-tuning in PyTorch through a minimal, educational codebase. This implementation modifies standard nn.Linear and nn.Embedding layers to freeze original weights and learn low-rank decomposition matrices instead.
Core LoRA Layer Mechanics
The foundation resides in labml_nn/lora/__init__.py, which reimplements PyTorch's linear and embedding layers with frozen base weights and trainable adapters.
Freezing Pre-Trained Weights
In labml_nn/lora/__init__.py lines 64-68 (for Linear) and lines 26-29 (for Embedding), the original weight parameters are initialized with requires_grad=False:
# Linear layer implementation (simplified)
self.weight = nn.Parameter(torch.empty((out_features, in_features)), requires_grad=False)
This ensures gradient computation skips the massive pre-trained matrices during backpropagation.
Injecting Low-Rank Adapters
Two small trainable matrices replace the full weight update. Lines 79-83 (Linear) and 33-35 (Embedding) initialize:
self.lora_awith shape(r, in_features)— initialized with Gaussian noiseself.lora_bwith shape(out_features, r)— initialized with zeros
Where r is the LoRA rank (typically 4-64), dramatically smaller than the original dimensions.
Scaling the Low-Rank Update
The implementation applies a scaling factor alpha / r to control adapter influence. Lines 60-63 (Linear) and 22-25 (Embedding) store:
self.scaling = alpha / r
This keeps the magnitude of the adaptation comparable to the original weight updates, preventing instability during fine-tuning.
The Forward Pass Computation
The critical logic appears in lines 90-96 (Linear) and 43-48 (Embedding), computing:
# Original frozen path + low-rank adapter path
original = F.linear(x, self.weight, bias=self.bias)
adapter = F.linear(F.linear(x, self.lora_a), self.lora_b) * self.scaling
return original + adapter
Mathematically, this implements (h = xW_0 + \frac{\alpha}{r}xA^{\top}B^{\top}), where (W_0) remains frozen and (A, B) are trainable.
Integrating LoRA into GPT-2 Transformers
The labml_nn/lora/gpt2.py file adapts the GPT-2 architecture by replacing every standard projection layer with LoRA-enabled equivalents. Within each transformer block:
- Query, Key, Value projections: use
Linear(d_model, d_model, r=lora_rank) - Embedding layers: use
Embedding(vocab_size, d_model, r=lora_rank) - Layer normalization and attention mechanisms: remain unchanged
This surgical replacement ensures only specific weight matrices receive gradient updates, while the bulk of the 124M+ parameter model stays frozen.
Loading Pre-Trained Weights and Optimization Strategy
The labml_nn/lora/experiment.py script handles weight loading and training configuration. Lines 64-68, 98-113, and 122-128 demonstrate loading Hugging Face GPT-2 checkpoints:
# Load pre-trained state dict
hf_model = AutoModelForCausalLM.from_pretrained("gpt2")
state_dict = hf_model.state_dict()
# Copy weights with strict=False to allow missing LoRA parameters
model.load_state_dict(mapped_weights, strict=False)
The strict=False parameter is essential — it ignores the newly initialized lora_a and lora_b parameters that don't exist in the original checkpoint.
Training Only the Adapters
Line 38 creates an optimizer receiving all parameters via model.parameters(), but only the LoRA matrices compute gradients:
optimizer = Adam(model.parameters(), lr=learning_rate)
During the training loop (lines 48-64), loss.backward() propagates gradients exclusively through lora_a and lora_b, leaving the frozen base weights untouched. This reduces training memory requirements and checkpoint sizes from gigabytes to megabytes.
Complete Fine-Tuning Workflow
Combine these components to fine-tune GPT-2 on a custom dataset:
from labml_nn.lora.gpt2 import GPTModel
from torch.optim import Adam
import torch.nn.functional as F
# 1. Initialize LoRA-augmented model
model = GPTModel(
d_model=768,
n_heads=12,
n_layers=12,
vocab_size=50257,
r=32, # LoRA rank
).to(device)
# 2. Load pre-trained weights (base weights frozen, LoRA params initialized)
# Implementation details in labml_nn/lora/experiment.py lines 64-113
# 3. Setup optimizer — only LoRA parameters receive gradients
optimizer = Adam(model.parameters(), lr=1e-4)
# 4. Training loop
for epoch in range(num_epochs):
for batch in dataloader:
inputs = batch[0].to(device)
logits = model(inputs[:, :-1])
loss = F.cross_entropy(
logits.view(-1, logits.size(-1)),
inputs[:, 1:].reshape(-1)
)
optimizer.zero_grad()
loss.backward() # Gradients flow only through lora_a and lora_b
optimizer.step()
This workflow fine-tunes the model on new tasks while preserving the original pre-trained knowledge in the frozen weights.
Summary
- LoRA modifies linear projections in
labml_nn/lora/__init__.pyby freezingself.weightat lines 64-68 and injecting trainablelora_a/lora_bmatrices at lines 79-83. - Scaling factor
alpha/r(lines 60-63) controls the magnitude of low-rank updates during the forward pass (lines 90-96). - GPT-2 integration in
labml_nn/lora/gpt2.pyreplaces standard layers while keeping attention mechanisms and layer norms unchanged. - Memory efficiency comes from loading pre-trained weights with
strict=False(lines 64-113 inexperiment.py) and optimizing only the small adapter matrices. - Training speed improves because backpropagation skips the massive frozen base weights, computing gradients only for the low-rank decomposition matrices.
Frequently Asked Questions
What makes LoRA more parameter-efficient than full fine-tuning?
Full fine-tuning updates all weight matrices in a transformer (millions or billions of parameters), whereas LoRA freezes these matrices and trains only the low-rank decomposition matrices (A) and (B). As implemented in labml_nn/lora/__init__.py, if the original weight has shape (768, 768) and rank r=32, LoRA trains only 768×32 + 32×768 = 49,152 parameters instead of the full 589,824 — a 12x reduction for that single layer.
Which layers in the transformer should use LoRA adapters?
According to the labml_nn/lora/gpt2.py implementation, LoRA adapters are applied to query, key, value projections and embedding layers — essentially every linear transformation that projects between the model dimension and itself. Layer normalization parameters and attention softmax operations remain frozen, as these typically contain less task-specific information.
How does the scaling factor alpha/r affect model performance?
The scaling factor, defined in lines 60-63 of labml_nn/lora/__init__.py, multiplies the low-rank update by alpha/r (defaulting to 1.0 when alpha=r). This hyperparameter controls the learning rate of the adaptation relative to the frozen pre-trained weights. Higher alpha values increase the influence of the LoRA adapters during the forward pass (lines 90-96), allowing faster adaptation but potentially destabilizing training if set too high.
Can this implementation be applied to architectures other than GPT-2?
Yes. The Linear and Embedding classes in labml_nn/lora/__init__.py are drop-in replacements for standard PyTorch layers. Any transformer architecture using nn.Linear for projections — including BERT, T5, or Vision Transformers — can integrate these classes by replacing the standard layers, exactly as demonstrated in labml_nn/lora/gpt2.py for the GPT-2 architecture.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →