Domain Adaptation, LoRA, and Catastrophic Forgetting

machine-learning
adaptation
Specialize a posttrained model with full fine-tuning and low-rank adapters, and measure what general behavior each method gives up.

Specialization can improve a domain score while damaging general instruction following, calibration, or tool use, and the damage is invisible in the adaptation loss because domain data contains no examples of the damaged behaviors. This chapter compares full fine-tuning with frozen-base adapters and low-rank updates, and treats the general-suite regression measurement as part of the method.

Low-rank adaptation

For a weight matrix W\in\mathbb{R}^{d\times k}, LoRA learns

W' = W + \frac{\alpha}{r}BA, \qquad A\in\mathbb{R}^{r\times k},\; B\in\mathbb{R}^{d\times r},\; r \ll \min(d,k),

with W frozen. Initializing B at zero makes W' = W at step zero, so the adapted model starts as the exact base model and the adapter is a controlled perturbation. Trainable parameters drop from dk to r(d+k); for a 4096\times 4096 layer at rank 8 that is 16.8\mathrm{M} down to 65.5\mathrm{k}. Implement adapter save/load, merge (W \mathrel{+}= \frac{\alpha}{r}BA), unmerge, and multiple named adapters per base checkpoint.

Full fine-tuning and LoRA

Adapt the posttrained model to a narrow corpus. Sweep rank, learning rate, the fraction of general instruction data mixed into the adaptation set, and domain-data size. For every configuration, report three numbers on identical prompts: domain accuracy, general-suite score from Chapter 07, and the tool-calling protocol from Chapter 11. The forgetting measurement is the delta of the last two scores before and after adaptation, not the domain score alone.

Adaptation reliability

A low-rank adapter updates few parameters, but the update applies on every forward pass in every context, including prompts far outside the domain corpus. Formatting, calibration, refusal-like behavior, and tool-call boundaries can all shift even when the adaptation data mentions none of them, which is why the regression suite runs on general prompts, not domain ones. Mixing general data into adaptation reduces the shift at a cost in domain accuracy; that exchange rate is itself a result worth reporting.

Summary

LoRA limits how many parameters move and enables swapping adapters per task at deployment. It does not limit which behaviors move: the update is applied unconditionally at inference, so general behavior can still shift, and the general-suite delta is the measurement that catches it.

Shared LoRA implementation

The equations above are implemented as frozen base linears plus zero-initialized low-rank residuals. The smoke check verifies the step-zero invariant and that merging the adapter into the base weight is numerically reversible. The adapter count is reported separately from the frozen model, so the parameter-efficiency claim is auditable.

import torch

from proof_lm.lora import (
    adapter_state_dict,
    inject_lora,
    merge_lora,
    trainable_parameter_count,
    unmerge_lora,
)
from proof_lm.model import DecoderConfig, ProofLM

torch.manual_seed(12)
adaptation_model = ProofLM(
    DecoderConfig(vocab_size=23, context_length=8, n_layers=1, d_model=16, n_heads=2, d_ff=32)
).eval()
for parameter in adaptation_model.parameters():
    parameter.requires_grad_(False)
inputs = torch.tensor([[1, 2, 3, 4]])
before = adaptation_model(inputs)[0]
target = ("blocks.0.attention.output", "blocks.0.feed_forward.0", "blocks.0.feed_forward.2")
injected = inject_lora(adaptation_model, rank=2, alpha=4, target_names=target)
after = adaptation_model(inputs)[0]
torch.testing.assert_close(before, after)
adapter_count = trainable_parameter_count(adaptation_model)
merge_lora(adaptation_model)
merged = adaptation_model(inputs)[0]
unmerge_lora(adaptation_model)
restored = adaptation_model(inputs)[0]
torch.testing.assert_close(after, merged)
torch.testing.assert_close(after, restored)
print("injected modules:", injected)
print("trainable adapter parameters:", adapter_count)
print("adapter tensors:", len(adapter_state_dict(adaptation_model)))
print("zero-impact initialization: verified")
injected modules: ('blocks.0.attention.output', 'blocks.0.feed_forward.0', 'blocks.0.feed_forward.2')
trainable adapter parameters: 256
adapter tensors: 6
zero-impact initialization: verified

Exercises

Use the exercises to test the chapter’s invariants and connect the derivations to the reusable implementation. Solutions are hidden in the notebook source and are available through the course tooling when needed.

[P12.1] Zero-initialized LoRA

Why initialize one LoRA factor at zero? What invariant does this provide at step zero?

[P12.2] LoRA parameter count

A base linear layer has 4096 input and 4096 output features. Compare the trainable parameter count of full fine-tuning against rank-8 LoRA, and list what the comparison omits.

[P12.3]

What two checks establish that a LoRA adapter is a controlled perturbation rather than an accidental model rewrite?

Back to top