Assumes you know from ML-101
This lesson builds on ML-101 Lesson 14 (Neural Networks & Backprop) for the chain-rule machinery behind gradient flow through frozen and unfrozen parameters, and on ML-101 Lesson 10 (Overfitting, Bias & Variance) for the inductive-bias trade-off between source-domain and target-domain risk. We will not re-derive the chain rule or define what a feature representation is.
The reader is also expected to know ML-101 Lesson 12 (Ensemble Learning): specifically, that dropout is equivalent — under linear activation — to a Gaussian-augmented feature map, and that the resulting regulariser matches a particular rate. We will use this result once, in passing, when we connect the LoRA scaling constant to a Gaussian-equivalent prior on the weight perturbation.
Learning Objectives
- Derive the bias–variance decomposition of the target risk as a function of the source risk, the source–target divergence, and the capacity term , and identify the assumption whose violation breaks the bound.
- Compute the parameter count and FLOPs of a LoRA adapter and contrast them with full fine-tuning for the same rank , naming the exact proportionality that determines whether the saving is one or two orders of magnitude.
- Justify, from the linearisation of the pre-training loss around , why LoRA's update recovers the full fine-tune to first order in the learning rate, and explain why the scaling constant is set at initialisation rather than learned.
- Diagnose a fine-tuning failure by separating dataset-shift (covariate, label, or concept) from optimisation pathologies (catastrophic forgetting, learning-rate mismatch, dead adapters), and design the smallest experiment that distinguishes them.
- Apply the warm-up heuristic for the learning-rate schedule and derive why linear warm-up followed by cosine decay is the canonical recipe for language-model fine-tuning at scale.
Why transfer works at all
The empirical claim "a network pretrained on ImageNet helps when fine-tuned on a small medical-imaging dataset" is not a miracle — it is a consequence of a measure-theoretic statement about source and target distributions. Let be the source distribution over inputs and the target. The target risk of a model decomposes under bounded-loss assumptions as
where is the source risk, is the source-domain optimum, and is an appropriate divergence on input distributions (total variation for deterministic bounds, MMD or KL for distributional ones). The decomposition says that the target risk is bounded by three positive terms, each of which we have a handle on. Fine-tuning reduces the third term by trading in for a target-domain optimum via gradient descent on a (small) target set.
The implicit assumption is that the source risk and the distance term are well-behaved as functions of near the fine-tuning starting point. This is the assumption whose violation breaks transfer: if the source representation has learned features that are useless or actively harmful on the target, then the distance term is large enough that no amount of fine-tuning on a small target set can recover. The 101-level misconception is that "transfer always helps" — it helps when the distance term is small relative to the capacity you can afford to spend on the source, and it silently hurts when the target requires new features the source never learned.
A useful diagnostic for whether transfer will help is the linear-probe test: freeze the source encoder and train only a logistic-regression head on top of its features. If the linear probe reaches most of the end-to-end accuracy on a moderate target set, the source representation already contains the target's relevant features and full fine-tuning is wasted capacity. If the linear probe plateaus far below the end-to-end baseline, the source features are insufficient and you need either a much larger target set or a fundamentally different source domain.
The four flavours of fine-tuning
Fine-tuning is not a single technique; it is a family of four related strategies whose trade-offs can be derived from first principles. They differ along two axes: which parameters are updated, and how the learning-rate is distributed across parameter groups.
| Strategy | Updated params | Memory | Use case |
|---|---|---|---|
| Full fine-tuning | all | highest | large target set, distribution shift |
| Head-only | last layer only | minimal | near-source target, very few labels |
| Adapter | small bottleneck per layer | moderate | many targets, multi-tenant serving |
| LoRA / low-rank | low | LLM fine-tuning at scale |
The four strategies form a Pareto frontier on three axes: target-set size, distribution shift, and serving cost. Head-only fine-tuning is the cheapest and is also the least flexible: it can only adapt the decision boundary on the source's feature representation. Full fine-tuning is the most flexible but the most expensive; its memory cost is dominated by optimiser state, which for Adam is two FP32 tensors per parameter (the first and second moments), so the memory ratio of Adam-state to weights is 8:1 for FP16 weights.
LoRA sits between adapter and full fine-tuning on the cost axis and can match full fine-tuning on the quality axis for moderate target-set sizes. The reason it can match full fine-tuning is the subject of the next section.
LoRA: the low-rank update
The LoRA hypothesis is that the change in weights induced by fine-tuning lives in a low-dimensional subspace. Formally, given a pretrained weight matrix , full fine-tuning produces for some of the same shape. LoRA constrains to be a low-rank product
with rank much smaller than (typical: against , a reduction in updated parameters for that matrix). The forward pass becomes , and only and receive gradient updates. stays frozen, and at inference the product can be folded back into so the deployed network has zero per-token overhead.
The empirical observation that motivates this parameterisation is that from full fine-tuning has singular values that decay sharply: the top few singular values capture most of the operator norm. The deep-learning community knew this informally from "intrinsic-dimension" studies on large language models (Aghajanyan et al., 2020); LoRA made it into a parameter-efficient inference framework. The hypothesis is not a theorem — there is no closed-form proof that the top-r subspace always captures the relevant update — but the empirical record is consistent across many LLM fine-tuning tasks.
The scaling factor is set at initialisation rather than learned. The reason is that is initialised to a Gaussian and to zero, so the product is exactly zero at step 0, and the gradient with respect to at step 0 is well-defined and nonzero. If we replaced by a learned scalar , the gradient with respect to would be zero at step 0 (because the upstream gradient flows through , which is zero), and learning would stall until either or moved away from their initialisations — a chicken-and-egg problem that is avoided by fixing .
import torch
import torch.nn as nn
class LoRALinear(nn.Module):
"""Wrap a frozen linear layer with a low-rank adapter."""
def __init__(self, base: nn.Linear, r: int = 8, alpha: float = 16.0):
super().__init__()
self.base = base
# Freeze the pretrained weights — the LoRA hypothesis says
# we should not waste capacity updating them.
for p in self.base.parameters():
p.requires_grad_(False)
d_out, d_in = base.weight.shape
self.A = nn.Parameter(torch.randn(r, d_in) * (1.0 / r ** 0.5))
self.B = nn.Parameter(torch.zeros(d_out, r)) # B = 0 at init
self.scale = alpha / r # fixed, not learned
def forward(self, x):
# x: (..., d_in) y: (..., d_out)
return self.base(x) + (x @ self.A.T) @ self.B.T * self.scale
The first-order equivalence with full fine-tuning is recovered by linearising the pre-training loss around . At step 0, , so the linearised update direction (the gradient of with respect to at ) is the same for LoRA and for full fine-tuning — both methods move in the direction of . For higher-order equivalence, you would need the entire to live in the column-space of , which is not guaranteed by random initialisation; in practice, and both move during fine-tuning and span a richer subspace than the initial random-A column-space would suggest. The full-fine-tune equivalence is an empirical observation about the trajectory, not a structural guarantee about the parameter space.
Parameter-count and FLOP accounting
The parameter-count saving for LoRA on a single linear layer is the ratio
For and , this is , i.e. 0.4% of the original matrix. For a 7B-parameter LLM with 32 transformer blocks and one matrix per attention projection, full fine-tuning updates parameters; LoRA at updates parameters — a reduction in trainable parameter count.
The saving in FLOPs during the forward pass is much smaller because the
frozen and the adapter both run on the same hardware. The
forward-pass FLOPs are , an additive
overhead of relative to — small but not
zero. The backward-pass FLOPs are the dominant cost of fine-tuning, and
here LoRA's saving is closer to the parameter-count ratio because the
backward pass over the frozen can be skipped entirely if the forward
activations are not stored — only the adapter's activations and gradients
need to be saved. In practice, with PyTorch's autograd, the frozen 's
backward is a no-op (its requires_grad is False), so the backward cost
is exactly proportional to the adapter parameter count.
The deeper reason for the FLOPs saving is that the optimiser state is proportional to the trainable parameter count, not to the FLOP count. For Adam, optimiser state is two FP32 tensors per trainable parameter. A reduction in trainable parameters is therefore a reduction in optimiser state, and on a 7B-parameter model that is the difference between fine-tuning fitting on a 32 GB GPU (LoRA at ) and requiring an 80 GB GPU (full fine-tuning, even with gradient checkpointing). The hardware bill is what has driven the industry adoption of LoRA — not the parameter count itself.
Catastrophic forgetting and how to mitigate it
When fine-tuning a pretrained model on a small target set, the loss on the source task typically rises as the target loss falls. This is catastrophic forgetting, and it is the single most common fine-tuning failure mode. It is not the same as overfitting: an over-fit model fails on the held-out target set, while a forgetting model fails on a held-out source set.
The mechanism is that Adam's update does not have a built-in prior toward . If the source loss was already low at , the fine-tune gradient on the target set can push away from in directions that look locally good for the target but globally bad for the source. Mitigations therefore work by adding such a prior. The two standard ones are:
- anchoring (sometimes called "weight decay regularised fine-tuning"): add to the loss. With AdamW this is the canonical recipe; is a common starting point.
- LoRA itself: because is frozen and only learns, the distance is implicitly bounded by the rank and the scale .
A diagnostic for whether forgetting is the failure mode is to log both the source validation loss and the target validation loss during fine-tuning. If the target loss falls but the source loss rises, you are forgetting; if both losses plateau high, you are not fitting; if both fall, you are on the right track. The 101-level misconception is that "fine-tuning replaces the pretrained behaviour" — it does not, and a deployment that requires the source behaviour after fine-tuning must either include the source task in the fine-tune data or use a parameter-efficient method that bounds the divergence from .
def evaluate_both(model, source_loader, target_loader, device="cuda"):
"""Used at every eval tick to detect catastrophic forgetting."""
model.eval()
src_loss = tgt_loss = 0.0
with torch.no_grad():
for x, y in source_loader:
src_loss += loss_fn(model(x.to(device)), y.to(device)).item()
for x, y in target_loader:
tgt_loss += loss_fn(model(x.to(device)), y.to(device)).item()
return src_loss / len(source_loader), tgt_loss / len(target_loader)
A second mitigation, useful when the target distribution is genuinely near the source, is layer-wise learning-rate decay: assign a smaller learning rate to early layers (closer to the input) and a larger one to later layers. The intuition, which has a Bayesian reading, is that early layers encode generic features (edges, motifs) that transfer well, while later layers encode task-specific features that need more adaptation. The canonical decay schedule is geometric with ratio (e.g. ), so layer has learning rate .
import torch
def split_param_groups(model, *, base_lr: float, decay: float = 0.95):
"""
Build per-layer param groups for layer-wise LR. The n-th layer from
the END gets the full base_lr; earlier layers decay geometrically.
"""
groups = []
# Assume model.layers is an nn.ModuleList of N blocks, ordered input -> output.
n = len(model.layers)
for i, block in enumerate(model.layers):
lr = base_lr * (decay ** (n - 1 - i))
groups.append({"params": block.parameters(), "lr": lr})
groups.append({"params": model.head.parameters(), "lr": base_lr})
return groups
opt = torch.optim.AdamW(split_param_groups(model, base_lr=2e-5, decay=0.95))
Learning-rate schedules: warm-up and cosine decay
The default learning-rate schedule for LLM fine-tuning is linear warm-up followed by cosine decay. Linear warm-up ramps the learning rate from zero to a peak over the first steps; cosine decay then ramps it back down to a floor (often zero) over the remaining steps. The full schedule is
The reason for warm-up is that the gradient statistics of the early training steps are poor estimates of the gradient statistics later in training: the model is moving quickly through the loss landscape, the Adam second moment is initialised to zero and has a slow warm-up of its own, and a too-large learning rate at can therefore land the optimiser in a region from which it cannot recover. Warm-up is not strictly necessary for full-batch gradient descent (where the gradient is well-defined at ), but for stochastic mini-batch methods it is essential in the first few hundred steps.
The reason for cosine decay rather than linear or step decay is that the expected loss curvature near a good minimum is approximately symmetric in the learning rate, and the cosine schedule's gradual descent matches the slow isocurvature movement that Adam takes near a flat basin. Linear decay tends to overshoot the minimum at the end of training, which is why a small floor is sometimes kept to allow continued exploration.
A useful piece of practical advice: when fine-tuning with LoRA, the learning rate that worked for full fine-tuning is too large for the LoRA adapters, because the gradient flow through the adapter has a smaller effective batch (the adapter sees fewer gradient signals per step than the full network does) and the optimiser statistics are therefore noisier. A common rule of thumb is to start at to the full-fine-tune learning rate when using LoRA, because the adapter's updates are smaller in magnitude but their gradients are larger relative to the parameter scale. This rule is empirical admission, not a theorem, and tuning is required for each task.
Key Takeaways
- Transfer learning is not a miracle: it works because the target risk is bounded by the source risk plus a divergence term plus a capacity term, and fine-tuning reduces the capacity term at the cost of the divergence term. The decomposition fails when the source features are useless or harmful on the target — a situation the linear-probe test can diagnose.
- LoRA's parameterisation works because fine-tuning-induced updates empirically concentrate in a low-rank subspace; the scaling constant is fixed at initialisation because a learned scaling would stall at step 0 by symmetry.
- The hardware win of LoRA is optimiser-state memory, not parameter count: Adam state is 8 bytes per trainable parameter, so a reduction in trainable parameters is a reduction in optimiser state.
- Catastrophic forgetting is the dominant fine-tuning failure mode and is mitigated by anchoring to (AdamW with ), by LoRA's implicit bound on , or by layer-wise learning-rate decay.
- The linear-warm-up + cosine-decay learning-rate schedule is canonical because warm-up protects against poor Adam-second-moment estimates in the first few hundred steps, and cosine decay's gradual descent matches the curvature near a good minimum.