14

Transfer Learning & Fine-Tuning

Cantonese podcast title: 遷移學習與微調

Learning Objectives

  1. Derive the bias–variance decomposition of the target risk as a function
  2. Compute the parameter count and FLOPs of a LoRA adapter and contrast them
  3. Justify, from the linearisation of the pre-training loss around
  4. Diagnose a fine-tuning failure by separating dataset-shift (covariate,
  5. Apply the warm-up heuristic for the learning-rate schedule and derive
Transfer Learning & Fine-Tuning — visual guide
Lesson 14 — LoRA's low-rank update Delta W = (alpha / r) B A A diagram showing how LoRA wraps a frozen d x k pretrained weight matrix with two trainable low-rank factors A in R^{r x k} and B in R^{d x r}. The frozen matrix is drawn shaded, the adapter unshaded, and the scaling constant alpha / r is annotated. LoRA · ΔW = (α / r) · B · A frozen d×k base + trainable rank-r factorisation x k-dim frozen W₀ ∈ ℝ^{d×k} d × k params · no gradient flow (frozen at pretraining θ₀) A ∈ ℝ^{r×k} trainable · k·r params B d·r α / r fixed at init y d-dim + example · d = k = 4096, r = 8 full fine-tune : 16 777 216 params LoRA (B + A) : 65 536 params ≈ 0.4 % of the base matrix B initialised to 0 → ΔW = 0 at step 0 A initialised to 𝒩(0, 1/r) (Gaussian) α/r scaling fixed — gradient would vanish at step 0 if it were learned first-order equivalence with full fine-tune: both methods have ΔW = 0 and share ∇_W ℓ(θ₀) at step 0.

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 ℓ2\ell_2 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

  1. Derive the bias–variance decomposition of the target risk as a function of the source risk, the source–target divergence, and the capacity term ∥θS−θT∥2\lVert\theta_S - \theta_T\rVert^2, and identify the assumption whose violation breaks the bound.
  2. Compute the parameter count and FLOPs of a LoRA adapter and contrast them with full fine-tuning for the same rank rr, naming the exact proportionality that determines whether the saving is one or two orders of magnitude.
  3. Justify, from the linearisation of the pre-training loss around θ0\theta_0, why LoRA's update ΔW=BA\Delta W = BA recovers the full fine-tune to first order in the learning rate, and explain why the scaling constant α/r\alpha / r is set at initialisation rather than learned.
  4. 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.
  5. 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 PSP_S be the source distribution over inputs and PTP_T the target. The target risk of a model fθf_\theta decomposes under bounded-loss assumptions as

RT(θ)  ≤  RS(θ)  +  dist⁡(PS,PT)  +  ∥θ−θS⋆∥2,R_T(\theta) \;\le\; R_S(\theta) \;+\; \operatorname{dist}(P_S, P_T) \;+\; \lVert\theta - \theta_S^{\star}\rVert^2,

where RSR_S is the source risk, θS⋆\theta_S^{\star} is the source-domain optimum, and dist⁡\operatorname{dist} 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 θS⋆\theta_S^{\star} for a target-domain optimum θT⋆\theta_T^{\star} via gradient descent on a (small) target set.

The implicit assumption is that the source risk RS(θ)R_S(\theta) and the distance term dist⁡(PS,PT)\operatorname{dist}(P_S, P_T) are well-behaved as functions of θ\theta 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 gg 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.

StrategyUpdated paramsMemoryUse case
Full fine-tuningall θ\thetahighestlarge target set, distribution shift
Head-onlylast layer onlyminimalnear-source target, very few labels
Adaptersmall bottleneck per layermoderatemany targets, multi-tenant serving
LoRA / low-rankB∈Rd×r,A∈Rr×kB \in \mathbb{R}^{d \times r}, A \in \mathbb{R}^{r \times k}lowLLM 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 W0∈Rd×kW_0 \in \mathbb{R}^{d \times k}, full fine-tuning produces W=W0+ΔWW = W_0 + \Delta W for some ΔW\Delta W of the same shape. LoRA constrains ΔW\Delta W to be a low-rank product

ΔW  =  αr B A,B∈Rd×r,  A∈Rr×k,\Delta W \;=\; \frac{\alpha}{r}\, B\,A, \qquad B \in \mathbb{R}^{d \times r},\; A \in \mathbb{R}^{r \times k},

with rank rr much smaller than min⁡(d,k)\min(d, k) (typical: r=8r = 8 against d=k=4096d = k = 4096, a 512×512\times reduction in updated parameters for that matrix). The forward pass becomes y=W0x+αrBAxy = W_0 x + \tfrac{\alpha}{r} B A x, and only AA and BB receive gradient updates. W0W_0 stays frozen, and at inference the product BABA can be folded back into W0W_0 so the deployed network has zero per-token overhead.

The empirical observation that motivates this parameterisation is that ΔW\Delta W 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 α/r\alpha / r is set at initialisation rather than learned. The reason is that AA is initialised to a Gaussian and BB to zero, so the product BABA is exactly zero at step 0, and the gradient with respect to BB at step 0 is well-defined and nonzero. If we replaced α/r\alpha / r by a learned scalar λ\lambda, the gradient with respect to λ\lambda would be zero at step 0 (because the upstream gradient flows through BABA, which is zero), and learning would stall until either BB or AA moved away from their initialisations — a chicken-and-egg problem that is avoided by fixing α/r\alpha / r.

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 θ0\theta_0. At step 0, ΔW=0\Delta W = 0, so the linearised update direction (the gradient of L\mathcal{L} with respect to WW at θ0\theta_0) is the same for LoRA and for full fine-tuning — both methods move in the direction of −∇WL(θ0)-\nabla_W \mathcal{L}(\theta_0). For higher-order equivalence, you would need the entire ΔW\Delta W to live in the column-space of BB, which is not guaranteed by random initialisation; in practice, AA and BB 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

∣A∣+∣B∣∣W∣  =  r k+r dd k  =  rd+rk.\frac{\lvert A\rvert + \lvert B\rvert}{\lvert W\rvert} \;=\; \frac{r\,k + r\,d}{d\,k} \;=\; \frac{r}{d} + \frac{r}{k}.

For d=k=4096d = k = 4096 and r=8r = 8, this is 84096+84096=1256\tfrac{8}{4096} + \tfrac{8}{4096} = \tfrac{1}{256}, i.e. 0.4% of the original matrix. For a 7B-parameter LLM with 32 transformer blocks and one 4096×40964096 \times 4096 matrix per attention projection, full fine-tuning updates ≈7×109\approx 7 \times 10^9 parameters; LoRA at r=8r = 8 updates ≈32×4×(4096×8+8×4096)≈16.8×106\approx 32 \times 4 \times (4096 \times 8 + 8 \times 4096) \approx 16.8 \times 10^6 parameters — a ≈420×\approx 420\times reduction in trainable parameter count.

The saving in FLOPs during the forward pass is much smaller because the frozen W0xW_0 x and the adapter BAxBAx both run on the same hardware. The forward-pass FLOPs are (2 d k)+(2 d r+2 r k)(2\,d\,k) + (2\,d\,r + 2\,r\,k), an additive overhead of rk+rd\tfrac{r}{k} + \tfrac{r}{d} relative to W0W_0 — 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 W0W_0 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 W0W_0'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 420×420\times reduction in trainable parameters is therefore a 420×420\times 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 r=8r = 8) 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 θt+1=θt−ηmt/(vt+ϵ)\theta_{t+1} = \theta_t - \eta m_t / (\sqrt{v_t} + \epsilon) does not have a built-in prior toward θ0\theta_0. If the source loss was already low at θ0\theta_0, the fine-tune gradient on the target set can push θ\theta away from θ0\theta_0 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:

  1. ℓ2\ell_2 anchoring (sometimes called "weight decay regularised fine-tuning"): add λ∥θ−θ0∥2\lambda \lVert \theta - \theta_0 \rVert^2 to the loss. With AdamW this is the canonical recipe; λ=0.01\lambda = 0.01 is a common starting point.
  2. LoRA itself: because W0W_0 is frozen and only ΔW\Delta W learns, the distance ∥W−W0∥F=∥ΔW∥F=αr∥BA∥F\lVert W - W_0 \rVert_F = \lVert \Delta W \rVert_F = \tfrac{\alpha}{r}\lVert BA\rVert_F is implicitly bounded by the rank rr and the scale α\alpha.

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 θ0\theta_0.

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 γ\gamma (e.g. γ=0.95\gamma = 0.95), so layer ℓ\ell has learning rate ηγL−ℓ\eta \gamma^{L - \ell}.

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 ηmax⁡\eta_{\max} over the first TwT_w steps; cosine decay then ramps it back down to a floor ηmin⁡\eta_{\min} (often zero) over the remaining T−TwT - T_w steps. The full schedule is

ηt  =  {ηmax⁡⋅tTw0≤t≤Twηmin⁡+12(ηmax⁡−ηmin⁡)(1+cos⁡ ⁣(π t−TwT−Tw))Tw<t≤T.\eta_t \;=\; \begin{cases} \eta_{\max} \cdot \dfrac{t}{T_w} & 0 \le t \le T_w \\[6pt] \eta_{\min} + \tfrac{1}{2}(\eta_{\max} - \eta_{\min})\left(1 + \cos\!\left(\pi\,\dfrac{t - T_w}{T - T_w}\right)\right) & T_w < t \le T. \end{cases}

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 vtv_t is initialised to zero and has a slow warm-up of its own, and a too-large learning rate at t=0t = 0 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 t=0t = 0), 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 ηmin⁡>0\eta_{\min} > 0 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 3×3\times to 10×10\times the full-fine-tune learning rate when using LoRA, because the adapter's ΔW\Delta W 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 ΔW=(α/r) BA\Delta W = (\alpha / r)\,BA works because fine-tuning-induced updates empirically concentrate in a low-rank subspace; the scaling constant α/r\alpha / r 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 420×420\times reduction in trainable parameters is a 420×420\times reduction in optimiser state.
  • Catastrophic forgetting is the dominant fine-tuning failure mode and is mitigated by ℓ2\ell_2 anchoring to θ0\theta_0 (AdamW with λ≈0.01\lambda \approx 0.01), by LoRA's implicit bound on ∥ΔW∥F\lVert\Delta W\rVert_F, 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.

Check your understanding

7 questions · 80% to complete the lesson

1 / 7

6 correct to pass

LoRA's adapter for a weight matrix W_0 in R^{d x k} introduces the additional trainable parameters

0 of 7 answered

Pick a lesson to start the audio.