Skip to content
AI-grafen
FAI engineeringLab· about 75 min· server sandbox

Lab: LoRA on a small network — train only the adapters

Implement a LoRA adapter (W + BA, rank r), freeze the base model, train only the adapters on a new task and compare parameter count and results with full fine-tuning.

Theory

ΔW = BA where B is d×r and A is r×k. With r ≪ min(d,k) the number of trained parameters is r(d+k) instead of dk. B is initialised to zero so that the model starts unchanged.

Sub-tasks

  1. LoRALinear — LoRALinear(base: nn.Linear, r, alpha) — freezes base, adds A (r×in, small normal) and B (out×r, zeros), forward = base(x) + (alpha/r)·(x Aᵀ) Bᵀ.
  2. apply_lora — apply_lora(model, r, alpha) replaces every nn.Linear in an nn.Sequential; trainable_params(model).
  3. fine-tune — finetune(model, X, y, steps, lr) trains only parameters with requires_grad; merge(lora_layer) returns an nn.Linear with W + (alpha/r)BA.

Passes when: lora_acc >= 0.85

The starter code

runs in an isolated sandbox on the server
import torch
import torch.nn as nn


class LoRALinear(nn.Module):
    def __init__(self, base: nn.Linear, r=4, alpha=8):
        super().__init__()
        self.base = base
        for p in self.base.parameters():
            p.requires_grad = False
        self.r, self.alpha = r, alpha
        # TODO: self.A = nn.Parameter(torch.randn(r, base.in_features) * 0.01); self.B = nn.Parameter(torch.zeros(base.out_features, r))
        ...

    def forward(self, x):
        # TODO: base(x) + (alpha / r) * (x @ A.T) @ B.T
        ...


def apply_lora(model: nn.Sequential, r=4, alpha=8):
    # TODO: ersätt varje nn.Linear i model med LoRALinear
    ...


def trainable_params(model):
    return sum(p.numel() for p in model.parameters() if p.requires_grad)


def total_params(model):
    return sum(p.numel() for p in model.parameters())


def finetune(model, X, y, steps=300, lr=1e-2):
    # TODO: Adam över [p for p in model.parameters() if p.requires_grad], CrossEntropyLoss
    ...


def merge(layer: LoRALinear) -> nn.Linear:
    # TODO: ny nn.Linear med weight = base.weight + (alpha/r) * B @ A, bias = base.bias
    ...

You write the code; tests you cannot see decide whether it holds up. Create a free account to run the lab.

Try the diagnosticCreate a free account

Expected results

LoRA with r=4 trains < 10 % of the parameters and reaches ≥ 0.85 accuracy on the shifted task; merge gives the same output as the adapter model.

Common mistakes

  • B initialised randomly → the model changes before training.
  • Forgets to freeze the base weights (requires_grad=False).
  • The alpha/r scaling is missing in merge.