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_a with shape (r, in_features) — initialized with Gaussian noise
  • self.lora_b with 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__.py by freezing self.weight at lines 64-68 and injecting trainable lora_a/lora_b matrices 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.py replaces standard layers while keeping attention mechanisms and layer norms unchanged.
  • Memory efficiency comes from loading pre-trained weights with strict=False (lines 64-113 in experiment.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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →