primer.ml.fine_tuning

Fine-tuning in practice: preparing data, forgetting old skills, merging models

Run: python -m primer.ml.fine_tuning

New to the notation? primer.notation explains every symbol used here (Σ, ‖x‖, λ, subscripts, and so on) from zero.

This lesson builds on primer.ml.training_stages, which explains what supervised fine-tuning (SFT), preference tuning and LoRA are. Here we deal with what happens when you actually do it: deciding whether to fine-tune at all, preparing the data, the old skills the model loses along the way, the small dataset it memorises, and how two fine-tuned models can be merged into one by plain arithmetic on their weights.

Level 1: The practitioner's guide

In one sentence. Fine-tuning in practice is the work around the training run: deciding with an evaluation set whether to train at all, preparing examples that teach what you mean, keeping the skills the model already had, stopping before it memorises your data, and sometimes combining two fine-tunes by arithmetic on their weights instead of training a third.

When you need it. You need this lesson when a fine-tune is on the table: a behaviour the best prompt cannot make consistent, or a long prompt sent so often that its tokens are most of the bill. The tell for the first is an eval score that stops improving however the prompt is reworded; the tell for the second is arithmetic. In this lesson's worked example, a 3,000-token prompt at \$2 per million tokens costs \$0.0060 per request, a tuned model that needs 300 tokens at \$4 per million costs \$0.0012, and a \$600 fine-tune pays for itself after 125,000 requests (25 days at 5,000 a day; the prices are illustrative). You don't need fine-tuning to add facts, which go stale the day they change and belong in retrieval, and you don't need it while a clearer instruction with two examples still moves the eval. Nothing in this lesson matters until the eval exists: it is the first box of the decision, before any model.

Your options. From the cheapest to the most committed:

Option What it does What it guarantees What it costs Where it lives
The best prompt, scored on the eval Instructions and examples in the prompt, measured on a frozen held-out set A baseline every other option must beat; survives a base-model upgrade Tokens on every call Your prompt
Hosted supervised fine-tuning Upload chat examples; the vendor applies the chat template, masks the user turns, trains and hosts Consistent format and style without the long prompt Curated examples, a training job, often a higher per-token price The vendor's fine-tuning API
LoRA on an open model Train a small adapter beside frozen weights on your examples Learns the task with less forgetting of the base model's skills (Biderman et al., 2024) A GPU, data, a serving stack; may learn a hard new domain less completely Your training stack
Full fine-tune, with replay Update every weight on your task mixed with a little general data The largest change the model can make, and the old skills kept if you replay them Multi-GPU training, the most forgetting when you don't replay, a copy of the model per variant Your training stack
Merge existing fine-tunes Add each fine-tune's task vector (its change from the base) to the base, no training Two separable skills in one model at no extra inference cost An eval on every task; fails when the two changes fight Your weights, with a merging toolkit

How to choose. Climb down only when the eval proves the rung above falls short.

  • A format or a house style: a few dozen to a few hundred excellent examples, hosted or LoRA. Start with 50 to 100 and double only while a doubling still beats the eval's margin of error (section 2e).
  • A behaviour that contradicts the base model ("always JSON" against "chat naturally", "be terse" against "explain"): expect the old behaviour to vanish wherever your data overrides it, and mix in general instruction-following data so the model stays an assistant (section 3c).
  • A harder specialised skill: thousands of examples, and consider a full fine-tune with replay, since LoRA learns less on a demanding domain.
  • Two skills trained separately, by different teams or on data that cannot be pooled: try a merge first, and check the cosine between the task vectors before you trust it (section 5b).
  • Whatever you pick, the fine-tune ships only if it beats the prompt on the same held-out set by more than that set's margin of error, and a general eval sits beside the task eval.

What it costs. Money is the one-off cost of writing and checking data plus the run, against a per-request saving that may be zero if the tuned model is priced like the base one; the break-even rule in section 1 puts a number on it. Every new base model repeats the one-off cost. Data is the expensive part, and the eval is dearer than people expect: 100 held-out examples measured at 80% mean "somewhere between 72% and 88%", and it takes 400 to halve that margin to ±3.9 points (this lesson's margin_of_error). Labels cost accuracy too: with 10% of reference labels wrong, a perfect model scores 90% and a 90% model scores 82%. Training itself is short. In the lesson's toy the new task is learned in 5 steps and every step after only erodes the general skill, and fine-tunes on small data run for a few epochs, with the best checkpoint shipped rather than the last.

What breaks.

  • The wrong chat template. Training succeeds, the loss falls, and the deployed model sees role markers it never learned. Use the exact template the base model was trained with, and the system prompt you will deploy.
  • Training on the user's turns. Forget the loss mask and the model learns to write questions. Hosted APIs mask for you; render_chat shows which segments count.
  • A leaking eval. A training example that nearly copies a held-out one turns the eval into a memory test. Freeze the held-out set first and remove near-copies from training (Jaccard on word pairs, threshold 0.7 in the toy).
  • Forgetting. Fine-tuning on one task takes the toy's general skill from 0.99 to 0.655, and a task that contradicts an earlier one takes the earlier one from 0.98 to 0.025. A lower learning rate only slows the slide; 10 replayed examples in 210 bring it back to 0.945.
  • Memorising a small dataset. On 16 examples with 3 wrong labels, validation loss bottoms out at epoch 71 (0.319) and has quadrupled (1.277) by epoch 1,500. Checkpoint every epoch and ship the best.
  • A merge that cancels. Task vectors with a clearly negative cosine (−0.25 in the toy) leave at least one skill at a coin flip whatever the scale. Retrain jointly, or use a conflict-resolving merge.

In the wild. Hosted fine-tuning APIs take chat-formatted examples and handle the template and the mask; OpenAI's model optimization guide lists supervised fine-tuning, DPO and reinforcement fine-tuning. For open models, Hugging Face TRL's SFTTrainer runs the supervised loop and PEFT supplies LoRA. Merging has its own toolkit: mergekit implements linear averaging, SLERP, task arithmetic, TIES, DARE and more. The evidence behind the advice is in the papers at the end of the lesson: LIMA's 1,000 curated examples, model soups for averaging fine-tunes of one task, task arithmetic for adding and subtracting them, TIES for resolving interference, elastic weight consolidation for when the old data is gone, and Biderman et al. on LoRA learning less and forgetting less.

Go deeper. Level 2 puts every claim above on a 97-weight model you can train in a fraction of a second: the break-even formula, the chat mask, Jaccard deduplication, the margin of error and the label-noise ceiling, forgetting measured as distance from the base, the overfitting curve with early stopping, and task vectors merged at every scale. If you only needed to decide whether and how to fine-tune, you are done.

Level 2: How it works, from scratch.

Level 2: How it works, from scratch

Picture a skilled cook who joins your restaurant. They already know how to cook (that is the pretrained model). You want them to cook your menu, your way. You can hand them a note with every order (a prompt), give them the recipe binder to look things up in (retrieval, RAG), or send them on a course about your kitchen (fine-tuning). The course is the only option that changes the cook, and it comes with the risks every teacher knows: the lessons can be badly written, the cook can cram the practice exam instead of learning, and a month of drilling one cuisine can make them rusty at everything else.

Everything below is one of those risks, measured on a model small enough to train in a fraction of a second.

1. Should you fine-tune at all?

Everyday picture. Giving the cook a note with every order costs a little each time. The course costs a lot once, and it has to be repeated whenever you hire a new cook (switch to a newer base model). The course pays off only when the notes would otherwise be long, sent millions of times, or simply not enough to make the cook consistent.

Tiny worked example. A support bot sends a 3,000-token prompt full of instructions and examples on every request. A fine-tuned model has learned that behaviour and needs only 300 tokens of prompt. Take these illustrative prices (real ones vary by provider and change often): \$2 per million input tokens for the general model, \$4 per million for the tuned one, because hosting a custom model usually costs more per token.

Tokens per request Price per million Cost per request
prompted 3,000 \$2 3,000 × 2 / 1,000,000 = \$0.0060
fine-tuned 300 \$4 300 × 4 / 1,000,000 = \$0.0012

Each request saves \$0.0048. Writing and checking 1,000 training examples plus the training run costs, say, \$600 once. After 600 / 0.0048 = 125,000 requests the course has paid for itself: 25 days at 5,000 requests a day.

flowchart TD E[Build the eval set first] --> P[Best prompt you can write] P --> M1{Good enough<br/>on the eval?} M1 -->|yes| SHIP[Ship the prompt] M1 -->|"no: missing facts"| R[Add retrieval, RAG] R --> M2{Good enough<br/>on the eval?} M2 -->|yes| SHIP M2 -->|"no: wrong behaviour, format or style"| FT[Fine-tune, often with LoRA] FT --> M3{Beats the prompt<br/>on the SAME eval?} M3 -->|yes, and pays off| SHIPFT[Ship the fine-tune] M3 -->|no| P

Reading it: start at the top, and notice that the first box is not a model at all: it is the evaluation set, the fixed exam every option is scored on. Each rung is cheaper and faster to change than the one below it, so you climb down only when the eval proves the rung above falls short. Missing knowledge goes to retrieval, because fine-tuning is a poor way to add facts; wrong behaviour that no prompt can pin down goes to fine-tuning. The last diamond matters most: a fine-tune ships only if it beats the prompt on the same eval. primer.ml.training_stages has the same decision as a flowchart of approaches; this one adds the measurement at every step.

Level 3: the formula and its symbols

$$ N^\star = \frac{C_{\text{once}}}{c_{\text{prompt}} - c_{\text{tuned}}}, \qquad c = \frac{t \cdot p}{10^6} $$

Symbols

Symbol Meaning here In the example
$N^\star$ "N-star": the break-even number of requests 125,000
$C_{\text{once}}$ one-off cost: writing and checking data, the training run \$600
$c$ cost of one request
$c_{\text{prompt}}, c_{\text{tuned}}$ cost of one request with the long prompt, and with the fine-tuned model \$0.0060, \$0.0012
$t$ input tokens sent per request 3,000 and 300
$p$ price in dollars per million input tokens \$2 and \$4
$10^6$ one million: prices are quoted per million tokens

In words: "divide what the course costs once by what it saves on each request; that is how many requests it takes to earn the cost back."

With the numbers: $c_{\text{prompt}}$ = 3,000 × 2 / 10⁶ = 0.006, $c_{\text{tuned}}$ = 300 × 4 / 10⁶ = 0.0012, so $N^\star$ = 600 / 0.0048 = 125,000.

Level 3: in Python

In Python:

# c = t·p / 1,000,000 for each option
c_prompt = 3000 * 2.0 / 1_000_000
c_tuned = 300 * 4.0 / 1_000_000
c_prompt, c_tuned  # → (0.006, 0.0012)
C_once = 600
# N* = C_once / (c_prompt − c_tuned)
N_star = C_once / (c_prompt - c_tuned)
round(N_star)  # → 125000
# days to break even at 5,000 requests a day
round(N_star / 5000, 1)  # → 25.0

The prompted line starts at zero and climbs steeply; the fine-tuned line starts at 600 dollars and climbs slowly; they cross at 125,000 requests

Reading it: the x-axis counts requests served and the y-axis is the total money spent so far. The prompted line starts at zero but climbs steeply, because every request pays for 3,000 tokens. The fine-tuned line starts at \$600 (the one-off cost) and climbs slowly. Left of the dashed line, prompting is cheaper; right of it, the fine-tune is. If the tuned model cost as much per request as the prompt, the two lines would never cross, and the only reason left to fine-tune would be quality.

In code: per_request_cost prices one request and break_even_requests returns $N^\star$, or infinity when the tuned model saves nothing per request.

Why it matters in practice. Three costs hide outside this formula. Every new base model means repeating the fine-tune, so $C_{\text{once}}$ recurs. Facts learned by fine-tuning go stale and are hard to update, which is why knowledge belongs in retrieval (primer.agents.rag). And the engineering time to build an eval is spent whichever way you go, which is why it comes first.

2. Preparing the data

Everyday picture. Before the course, someone writes the course book. Every recipe must be written in the same layout, no recipe may appear five times, some recipes are locked in a drawer for the final exam before the cook ever sees the book, and the recipes must actually be right. Most fine-tuning failures are course-book failures.

flowchart LR RAW[Raw examples<br/>logs, experts, drafts] --> FMT[Format as chat<br/>messages] FMT --> VAL[Validate<br/>roles, empty turns] VAL --> DD[Deduplicate<br/>near-copies] DD --> SPLIT[Set aside the<br/>held-out set, frozen] SPLIT --> LEAK[Drop training examples<br/>that copy held-out ones] LEAK --> AUDIT[Audit a sample<br/>of the labels] AUDIT --> TRAIN[Training set] SPLIT --> EVAL[Held-out set]

Reading it: raw examples enter on the left and pass through five filters. The branch at "Set aside the held-out set" is the important one: the held-out examples leave the pipeline before anything is trained or tuned, and the step after it removes any training example that nearly copies one of them. The subsections below build each box.

2a. Formatting chat examples

Everyday picture. A play script: every line starts with who speaks it. The model learns the play by reading scripts, and it is graded only on the lines of the character it will play, the assistant.

Tiny worked example. One example as a list of role-tagged messages:

Turn Role Content Trained on?
1 system Be brief. no
2 user Capital of France? no
3 assistant Paris. yes

Rendered into text, it becomes <|system|>Be brief.<|end|><|user|>Capital of France?<|end|><|assistant|>Paris.<|end|>, and only the last segment counts towards the loss.

flowchart LR S["system: Be brief."] --> U["user: Capital of France?"] --> A["assistant: Paris."] S -.->|"context only, no loss"| L[Loss] U -.->|"context only, no loss"| L A ==>|"every token scored"| L

Reading it: the solid arrows are the order the model reads the turns. The dotted arrows carry no training signal: the system and user turns are context the model conditions on. Only the thick arrow feeds the loss, so the model learns to answer, not to write questions. This is the loss mask from primer.ml.training_stages, section 2.

The rules that matter. Use the exact chat template the base model was trained with; the role markers above are illustrative, and each model family has its own. Train with the same system prompt you will deploy with. Reject examples whose last turn is not the assistant's (nothing to learn), whose turns are empty, or whose roles are unknown.

In code: chat_example builds the message list, render_chat flattens it into (segment, trained?) pairs, and validate_chat lists every problem with an example.

Why it matters in practice. A template mismatch is silent: training succeeds, the loss falls, and the deployed model sees markers it never learned, so it behaves like the base model or worse.

2b. Deduplication

Everyday picture. A flashcard deck with the same card in it five times. You study that card five times as often, and you start answering every question with it.

Tiny worked example. Compare questions by their word pairs (two neighbouring words; a window of k words is called a shingle), after lowercasing and dropping punctuation:

Text Word pairs
"How do I reset my password?" how do, do i, i reset, reset my, my password
"how do I reset my password, please" the same 5, plus password please
"How do I change my email?" how do, do i, i change, change my, my email

The first two share 5 of the 6 distinct pairs between them: a near-duplicate. The first and third share 2 of 8: different questions.

Level 3: the formula and its symbols

$$ J(A, B) = \frac{|A \cap B|}{|A \cup B|} $$

Symbols

Symbol Meaning here In the example
$A, B$ the sets of word pairs of two texts 5 pairs and 6 pairs
$\cap$ intersection: pairs in both sets 5 shared pairs
$\cup$ union: pairs in either set, each counted once 6 distinct pairs
$\lvert \cdot \rvert$ the number of items in a set
$J(A, B)$ Jaccard similarity, from 0 (nothing shared) to 1 (identical) 0.83

In words: "count the pairs the two texts share, and divide by the number of distinct pairs they have between them."

With the numbers: J = 5 / 6 = 0.83 for the two password questions and 2 / 8 = 0.25 for password against email. A threshold of 0.7 keeps the first password question, drops its rewording, and keeps the email question.

Level 3: in Python

In Python:

def pairs(text):
    words = text.lower().replace("?", "").replace(",", "").split()
    return {(a, b) for a, b in zip(words, words[1:])}
A = pairs("How do I reset my password?")
B = pairs("how do I reset my password, please")
C = pairs("How do I change my email?")
# |A ∩ B| and |A ∪ B|
len(A & B), len(A | B)  # → (5, 6)
# J(A, B)
round(len(A & B) / len(A | B), 2)  # → 0.83
# J(A, C)
len(A & C) / len(A | C)  # → 0.25
flowchart LR T[Next example] --> CMP{Jaccard with any<br/>kept example ≥ 0.7?} CMP -->|yes| DROP[Drop it:<br/>a near-copy] CMP -->|no| KEEP[Keep it] KEEP --> T DROP --> T

Reading it: examples arrive one at a time and are compared against everything already kept. A match above the threshold is dropped; anything else joins the kept list. The first copy of each group survives, so the order of the data decides which wording you keep.

In code: normalize lowercases and strips punctuation, shingles collects the word pairs, jaccard scores two texts and deduplicate returns the indices worth keeping. Comparing every pair is fine for thousands of examples; at web scale, MinHash estimates the same Jaccard without comparing every pair.

Why it matters in practice. Duplicates overweight a few examples, so the model parrots them. Worse, a near-copy of an eval question in the training set lets the model recite the answer, and the eval score becomes a memory test. Lee et al. found training sets with thousands of near-duplicates, and removing them made models memorise less.

2c. The held-out set comes first

Everyday picture. A good teacher writes the final exam before teaching the course and locks it in a drawer. If the exam were written afterwards, it would drift towards what the class happened to practise.

Tiny worked example. Fifty support questions are deduplicated, shuffled with a fixed seed, and 10 (20%) go in the drawer. Only then are the other 40 used, and any of the 40 that nearly copies one of the 10 is removed. Every later choice (prompt wording, learning rate, which checkpoint to ship) is scored on those 10. But how much can 10, or even 100, questions tell you?

Level 3: the formula and its symbols

$$ \text{margin} = z \sqrt{\frac{a\,(1 - a)}{n}} $$

Symbols

Symbol Meaning here In the example
$a$ the accuracy measured on the held-out set 0.8
$n$ the number of held-out examples 100
$\sqrt{a(1-a)/n}$ the standard error: how much the measured accuracy would wobble if you drew a different held-out set of the same size 0.04
$z$ how many standard errors to allow; 1.96 covers 95% of the wobble 1.96
margin the true accuracy is probably within ± this of the measured one 0.078

In words: "the measured accuracy is uncertain by about two standard errors, and the standard error shrinks with the square root of the number of examples."

With the numbers: with 100 examples at 80%, the margin is 1.96 × √(0.8 × 0.2 / 100) = 1.96 × 0.04 = 0.078, so "80%" means "somewhere from about 72% to 88%". A fine-tune that scores 83% has not been shown to beat a prompt that scores 80%. With 400 examples the margin halves to 0.039.

Level 3: in Python

In Python:

import math
a, n, z = 0.8, 100, 1.96
# the standard error sqrt(a(1 − a)/n)
round(math.sqrt(a * (1 - a) / n), 4)  # → 0.04
# the 95% margin
round(z * math.sqrt(a * (1 - a) / n), 4)  # → 0.0784
# four times as many examples halves it
round(z * math.sqrt(a * (1 - a) / 400), 4)  # → 0.0392
flowchart LR ALL[All examples,<br/>deduplicated] --> SH[Shuffle with<br/>a fixed seed] SH --> EV[20% held out<br/>frozen, never trained on] SH --> TR[80% training] EV --> CHK{Training example<br/>near-copies one?} TR --> CHK CHK -->|yes| X[Remove from training] CHK -->|no| OK[Keep for training]

Reading it: the fixed seed makes the split repeatable, so everyone on a team scores against the same held-out set. The held-out branch is frozen the moment it is made. Both branches then meet in the leak check, which only ever removes examples from the training side.

The margin of error falls from about 16 points at 25 examples to 1.4 points at 3,200; each fourfold increase halves it

Reading it: the x-axis is the number of held-out examples (log scale) and the y-axis is the ± margin, in percentage points, around a measured 80%. The marked points show the square-root law: 100 examples give ±7.8 points, 400 give ±3.9 and 1,600 give ±2.0. Each halving of the margin costs four times the examples, which is why held-out sets of a few hundred are common and a few dozen can only detect big differences.

In code: split_before_training deduplicates, shuffles with a seed, freezes the held-out set and removes leaks with remove_leaks; margin_of_error gives the ± for any accuracy and size.

Why it matters in practice. Every time you look at held-out scores and change something, the held-out set leaks a little into your decisions. A held-out set built first and used sparingly is the only honest measure of whether the fine-tune helped. See primer.ml.regularization for train, validation and test splits, and primer.agents.evals for building evals.

2d. Label quality

Everyday picture. An answer key with typos. A student who gets every question right is marked wrong wherever the key is wrong, and a student who makes the same mistake as the key is marked right.

Tiny worked example. A model is truly right 90% of the time, and 10% of the reference labels are wrong. It scores a point when it is right on a right label (0.9 × 0.9 = 0.81), or when it is wrong on a wrong label and the two mistakes cancel (0.1 × 0.1 = 0.01). The eval reports 82%, not 90%.

Level 3: the formula and its symbols

$$ a_{\text{measured}} = a\,(1 - \varepsilon) + (1 - a)\,\varepsilon $$

Symbols

Symbol Meaning here In the example
$a$ the model's true accuracy 0.9
$\varepsilon$ epsilon: the share of reference labels that are wrong 0.1
$a(1-\varepsilon)$ right answer, right label 0.81
$(1-a)\varepsilon$ wrong answer on a wrong label (yes/no labels only, so the two errors agree) 0.01
$a_{\text{measured}}$ what the eval reports 0.82

In words: "the eval gives credit for being right on a correct label, and by accident for being wrong on a wrong one."

With the numbers: 0.9 × 0.9 + 0.1 × 0.1 = 0.82. A perfect model ($a = 1$) scores only 1 − ε = 0.9: noisy labels put a ceiling on what the eval can show.

Level 3: in Python

In Python:

a, eps = 0.9, 0.1
# right on a right label, plus wrong on a wrong label
round(a * (1 - eps) + (1 - a) * eps, 2)  # → 0.82
# even a perfect model scores only 1 − ε
round(1.0 * (1 - eps), 2)  # → 0.9
# two labelers, five items: how often do they agree?
labeler_1 = ["yes", "no", "yes", "yes", "no"]
labeler_2 = ["yes", "no", "no", "yes", "no"]
sum(p == q for p, q in zip(labeler_1, labeler_2)) / len(labeler_1)  # → 0.8
flowchart TD X[One eval item] --> R{Model right?} R -->|"yes, 0.9"| LR{Label right?} R -->|"no, 0.1"| LW{Label right?} LR -->|"yes, 0.9"| P1["scored right: 0.81"] LR -->|"no, 0.1"| P2["scored wrong: 0.09"] LW -->|"yes, 0.9"| P3["scored wrong: 0.09"] LW -->|"no, 0.1"| P4["scored right: 0.01"]

Reading it: each item takes two coin flips: is the model right, and is the label right? Multiply along a path to get its share. Two paths are scored as right, 0.81 and 0.01, which add up to the 0.82 the eval reports. The 0.09 on the second path is a correct answer marked wrong by a bad label.

How to check label quality. Have two people label the same sample independently and measure how often they agree. Above, they agree on 4 of 5 items, 80%: one item in five is ambiguous or mislabelled, and your eval can't resolve differences smaller than that noise. Read every disagreement; they are usually unclear instructions, not careless labelers.

In code: measured_accuracy applies the formula and label_agreement compares two labelers.

Why it matters in practice. Bad training labels teach the model the mistakes (section 4 shows a small model memorising three of them). Bad eval labels hide real improvements. Zhou et al. (LIMA) fine-tuned a large model on just 1,000 carefully chosen examples and got a strong assistant: quality of examples beats quantity.

2e. How many examples?

Everyday picture. Teaching a house style is not teaching a language. A new cook learns how you plate a dish from a few dozen good examples; they don't need ten thousand.

Tiny worked example. Fine-tuning teaches a behaviour the base model can almost do already, so the useful sizes are small: a few dozen examples to show a format, hundreds for a consistent style or a narrow task, thousands for a harder specialised skill. Section 4 fine-tunes on 16 examples and shows the other side: with so few, the model soon memorises them, noise included.

flowchart LR S[Start with 50 to 100<br/>high-quality examples] --> T[Fine-tune] T --> E[Score on the<br/>held-out set] E --> Q{Gained more than<br/>the margin of error?} Q -->|yes| D[Double the data] --> T Q -->|no| STOP[Stop adding data:<br/>fix quality or approach]

Reading it: the loop grows the dataset only while each doubling still buys a real improvement, meaning one larger than the held-out margin from section 2c. When a doubling stops paying, more of the same data won't help; the fix is better examples or a different approach.

Why it matters in practice. Labelled data is the expensive part of fine-tuning. Doubling until the gains flatten spends it where it helps, and the held-out set tells you when to stop.

3. Catastrophic forgetting

Everyday picture. Someone who learned to drive in Britain, keeping left, moves to the United States and practises keeping right every day for a month. The new habit wins, and on a trip home they drift to the wrong side of the road. Nobody told them to forget the old rule; the new practice simply rewrote the reflex both rules use. Neural networks do this to an extreme: train on a new task alone and an old one can vanish. This is catastrophic forgetting.

The model we'll fine-tune. To watch it happen, we need a model small enough to train instantly. Each example is four numbers, and the answer is yes or no:

Task Inputs it uses Rule Plays the role of
general all four, anywhere in [−3, 3] yes when x₁ + x₃ > 0 the base model's pretraining
A x₁ in [−3, −1], x₂ in [−2, 2] yes when x₂ > 0 keep left in Britain
B x₁ in [1, 3], x₂ in [−2, 2] yes when x₂ < 0 keep right in the US
C x₃, x₄ in [−2, 2] yes when x₃ + x₄ > 0 an unrelated skill

A and B read the same two inputs and apply opposite rules in different regions, so one model can learn both, but only by paying attention to the region. C reads inputs that A and B never touch.

flowchart LR X["4 inputs<br/>x1 x2 x3 x4"] --> H["16 hidden units<br/>tanh"] H --> O["1 output<br/>probability of yes"] TH["θ: all 97 weights<br/>in one vector"] -.-> H TH -.-> O

Reading it: four numbers go in, 16 hidden units each mix all of them, and one output unit turns the hidden units into a probability. Every weight (64 in the first layer, 16 hidden biases, 16 output weights and one output bias) sits in a single vector θ of 97 numbers. Every task flows through the same hidden units, which is exactly why one task's training can damage another. primer.ml.neural_net builds this kind of network from scratch.

The base model is this network trained on the general task (99% accuracy on held-out examples). Every fine-tune below starts from it and runs plain gradient descent on 200 examples.

In code: TinyNet holds θ and computes predictions, loss and TinyNet.gradient; make_task draws examples for each of the TASKS; base_model pretrains the base and fine_tune trains a copy, leaving the starting model untouched.

3a. The general skill fades while you teach a new one

Everyday picture. A month of drilling one cuisine makes the cook a little rusty at everything else, and a second month of drilling it makes them rustier still, even though they had mastered the cuisine in the first week.

Tiny worked example. Fine-tune the base on task A (learning rate 0.5) and check both skills on held-out examples:

After step Accuracy on A General skill Distance from base
0 (the base) 0.53 0.99 0
5 0.975 0.825 2.60
10 0.995 0.79 2.83
300 0.98 0.655 5.19

Task A is learned in 5 steps. The next 295 steps teach nothing new about A but keep eroding the general skill, from 0.825 to 0.655. The last column explains why: the weights keep travelling away from the base.

Level 3: the formula and its symbols

$$ d_t = \lVert \theta_t - \theta_{\text{base}} \rVert = \sqrt{\sum_{i=1}^{P} \left(\theta_{t,i} - \theta_{\text{base},i}\right)^2} $$

Symbols

Symbol Meaning here In the example
$\theta_{\text{base}}$ theta: every weight of the base model, as one vector 97 numbers
$\theta_t$ every weight after $t$ fine-tuning steps
$P$ the number of weights 97 (3 in the hand example)
$i$ a counter over the weights 1 … P
$\theta_{t,i}$ the $i$-th weight after $t$ steps
$\lVert v \rVert$ the length of vector $v$: square each entry, add, take the square root
$d_t$ how far fine-tuning has moved the model 5.19 after 300 steps

In words: "the distance travelled is the length of the change in the weights: square every weight's change, add them up, take the square root."

With the numbers: for a three-weight model moving from (1, 0, 2) to (1.5, −1, 2), the changes are (0.5, −1, 0), and d = √(0.25 + 1 + 0) = √1.25 = 1.118.

Level 3: in Python

In Python:

import math
theta_base = [1.0, 0.0, 2.0]
theta_t = [1.5, -1.0, 2.0]
# each weight's change
[t - b for t, b in zip(theta_t, theta_base)]  # → [0.5, -1.0, 0.0]
# ‖θ_t − θ_base‖
round(math.sqrt(sum((t - b) ** 2 for t, b in zip(theta_t, theta_base))), 3)  # → 1.118
flowchart LR DA[Task A examples] -->|gradient| W["Shared weights θ"] W --> SA[Answers on task A] W --> SG[Answers on the<br/>general skill] SG -.->|"no examples,<br/>no gradient"| W

Reading it: the gradient from task A's examples moves the shared weights, and both skills read those same weights. The dotted arrow is what is missing: the general skill has no examples in this training run, so nothing pushes back when a change that helps A hurts it. Forgetting is not an event; it is the absence of a counterweight.

Left: accuracy on A jumps to 1 within a few steps while the general skill slides from 0.99 to 0.66 at learning rate 0.5 and only to 0.79 at 0.02. Right: the general skill falls as distance from the base grows

Reading it: on the left, the x-axis is the training step (log scale). Blue lines are accuracy on task A, red lines the general skill; solid is learning rate 0.5, dashed is 0.02. With the big rate, A is learned within a few steps and the general skill then slides for the rest of the run. With the small rate, A is learned later (around step 70) and the general skill ends at 0.79 instead of 0.655, with A just as good. On the right, every step of both runs is plotted as distance from the base against the general skill: the points fall along one downward curve. Forgetting tracks how far the weights move.

Mitigations that shorten the trip. Three knobs limit the distance, and the toy measures two of them:

  • Fewer steps. Stop once the held-out score on the new task stops improving. Here, stopping at step 10 keeps the general skill at 0.79 instead of 0.655, with A at 0.995.
  • A lower learning rate. Smaller steps travel less far for the same result: 0.79 instead of 0.655 after 300 steps.
  • A smaller update. LoRA (primer.ml.training_stages, section 4) freezes the base weights and allows only a low-rank change. Biderman et al. measured this on real language models and summed it up in their title: LoRA learns less and forgets less.

In code: general_skill_run fine-tunes on A and records, after every step, accuracy on A, the general skill and the distance from the base.

Why it matters in practice. A fine-tuned assistant that has become worse at everything outside its narrow task is the most common fine-tuning disappointment. Always score the general skills you care about alongside the new task, and prefer the earliest checkpoint that has learned the task.

3b. A new task that contradicts an old one

Everyday picture. Back to the driver. Practising "keep right" does not merely add a skill; it pushes directly against "keep left", because both use the same reflex.

Tiny worked example. Take the model fine-tuned on A (98% on A), then fine-tune it on B alone for 300 steps:

Second fine-tune A before A after New task after
on B (contradicts A) 0.98 0.025 1.00
on C (separate inputs) 0.98 0.97 0.98

After B, the model does not merely forget A; it answers A's questions backwards, because it learned "yes when x₂ < 0" everywhere. After C, A is barely touched: C's inputs never flowed through the weights A relies on most.

Level 3: the formula and its symbols

$$ F = \text{acc}_{\text{old}}^{\text{before}} - \text{acc}_{\text{old}}^{\text{after}} $$

Symbols

Symbol Meaning here In the example
$\text{acc}_{\text{old}}^{\text{before}}$ held-out accuracy on the old task before the new fine-tune 0.98
$\text{acc}_{\text{old}}^{\text{after}}$ the same, after the new fine-tune 0.025
$F$ forgetting: accuracy lost on the old task 0.955

In words: "forgetting is how much accuracy the old task lost."

With the numbers: F = 0.98 − 0.025 = 0.955 after B, and 0.98 − 0.97 = 0.01 after C.

Level 3: in Python

In Python:

acc_before, acc_after_B, acc_after_C = 0.98, 0.025, 0.97
# F after the contradicting task B
round(acc_before - acc_after_B, 3)  # → 0.955
# F after task C, on separate inputs
round(acc_before - acc_after_C, 3)  # → 0.01
flowchart LR BASE[Base] -->|fine-tune on A| MA["Model knows A<br/>A: 0.98"] MA -->|fine-tune on B only| MB["Model knows B<br/>A: 0.025, B: 1.00"] MA -->|fine-tune on C only| MC["Model knows A and C<br/>A: 0.97, C: 0.98"]

Reading it: both second fine-tunes start from the same model. The only difference is which weights the new task needs: B needs the very ones A uses, in the opposite direction; C mostly needs others. How much a model forgets depends less on how long you train than on how much the new task overlaps and conflicts with the old one.

Fine-tuning on B after A: accuracy on A falls from 0.98 to near 0 within a few steps while B rises to 1; with 10 replayed A examples, A dips and then recovers to 0.945

Reading it: the x-axis is the step of the second fine-tune (log scale), the y-axis held-out accuracy. Solid lines are plain fine-tuning on B: B (green) climbs to 1 while A (blue) falls to about half in the same few steps, and A keeps sliding towards zero for as long as training continues. Dashed lines add 10 replayed A examples (section 3c): A still dips at first (to about 0.2 around step 40), then climbs back to 0.945 while B reaches 0.99.

Does a lower learning rate help here? Only in the sense of slowing the slide. At learning rate 0.02, or stopping after 10 steps, B reaches 0.975 and A still drops to 0.355.

Accuracy on A against accuracy on B during the second fine-tune: every run on B alone traces the same curve whatever the learning rate and ends near zero on A, while the replay run climbs the right edge to the top-right corner

Reading it: each line traces one run: x is accuracy on B, y is accuracy on A, and each run starts at the top left (knows A, not B). The three learning rates (0.5, 0.1, 0.02) take very different numbers of steps but trace the same curve, close to the dotted line where A + B = 1: every point of B they gain costs about a point of A, and once B is learned, further steps slide A down the right edge towards zero. The replay run (dashed) follows the same curve at first, then climbs the right edge and ends in the top-right corner, knowing both. When the new data contradicts the old skill, slowing down can't help, because nothing in B's data says A still matters.

In code: sequential_run starts from the A fine-tune, trains on a new task (optionally with replay), and records both accuracies at every step; forgetting computes F.

Why it matters in practice. Real fine-tunes contradict the base model more often than you'd think: "always answer in JSON" contradicts "chat naturally"; "be terse" contradicts "explain in detail". Expect the old behaviour to vanish wherever your data overrides it, and test for it.

3c. Replay: keep practising the old skill

Everyday picture. A pianist learning a new piece plays one old piece at the start of every practice session. It costs a few minutes and keeps the old repertoire alive.

Tiny worked example. Add just 10 of task A's 200 training examples to B's 200: under 5% of the mix. The result: A 0.945, B 0.99 (up from A 0.025 without replay). Why can so few examples do so much? Look at the loss. Suppose that early in training the model already scores B well (loss 0.05 per example) but has started forgetting A (loss 3.0 on each replayed example):

Level 3: the formula and its symbols

$$ \mathcal{L}_{\text{mix}} = \frac{n_{\text{new}}\,\mathcal{L}_{\text{new}} + n_{\text{old}}\,\mathcal{L}_{\text{old}}}{n_{\text{new}} + n_{\text{old}}} $$

Symbols

Symbol Meaning here In the example
$n_{\text{new}}$ examples of the new task 200
$n_{\text{old}}$ replayed examples of the old task 10
$\mathcal{L}_{\text{new}}$ average loss on the new examples 0.05
$\mathcal{L}_{\text{old}}$ average loss on the replayed examples 3.0
$\mathcal{L}_{\text{mix}}$ the loss training actually minimises: the average over every example in the mix 0.19

In words: "the loss on a mixed dataset is the average over all its examples, so each group counts in proportion to its size times its loss."

With the numbers: (200 × 0.05 + 10 × 3.0) / 210 = (10 + 30) / 210 = 0.19. The 10 replayed examples are under 5% of the data but contribute 0.143 of the 0.19: three quarters of the loss, and so most of the gradient. The old examples shout loudest exactly when they are being forgotten.

Level 3: in Python

In Python:

n_new, L_new = 200, 0.05
n_old, L_old = 10, 3.0
# the replayed share of the data
round(n_old / (n_new + n_old), 3)  # → 0.048
# each group's contribution to the average loss
round(n_new * L_new / 210, 3), round(n_old * L_old / 210, 3)  # → (0.048, 0.143)
# L_mix
round((n_new * L_new + n_old * L_old) / (n_new + n_old), 2)  # → 0.19
flowchart LR NB["New task B<br/>200 examples"] --> MIX[Shuffle together<br/>210 examples] OA["Old task A<br/>10 kept examples"] --> MIX MIX --> FT[Fine-tune] FT --> BOTH["Knows B: 0.99<br/>and A: 0.945"]

Reading it: the old task's small sample joins the new data before training, so every gradient step sees both. This supplies exactly the counterweight that was missing in section 3a's diagram: when a change that helps B starts hurting A, the replayed examples' loss rises and pushes back.

In code: replay_mix appends the old examples, mixed_loss is the formula, and sequential_run with n_replay=10 runs the experiment.

Why it matters in practice. When fine-tuning a language model, mix some general instruction-following data into your task data, so the model keeps being a good assistant while it learns your task. When you can't replay (the old data is gone or private), methods such as elastic weight consolidation (Kirkpatrick et al.) instead penalise changes to the weights the old task relied on most.

4. Overfitting a small dataset

Everyday picture. A student with only 16 flashcards, three of which have the wrong answer on the back. For a while, studying teaches the pattern. Keep drilling and the student memorises every card word for word, including the three wrong answers, and gets worse on new questions. primer.ml.regularization builds this idea from scratch; here it is in a fine-tune.

Tiny worked example. Fine-tune the base on only 16 examples of task A, 3 of them deliberately mislabelled, for 1,500 epochs. (With full-batch training, one step is one pass over the data, one epoch.)

At the best epoch (71) At the end (1,500 epochs)
training loss 0.373 0.005
validation loss 0.319 1.277

At the best epoch, training loss is higher than validation loss: the model is refusing to fit the three wrong labels, which is exactly right. By the end, training loss is nearly zero, so the wrong labels have been memorised, and validation loss has quadrupled.

Level 3: the formula and its symbols

$$ g(t) = \mathcal{L}_{\text{val}}(t) - \mathcal{L}_{\text{train}}(t) $$

Symbols

Symbol Meaning here In the example
$t$ the epoch 71, then 1,500
$\mathcal{L}_{\text{train}}(t)$ average loss on the 16 training examples 0.373, then 0.005
$\mathcal{L}_{\text{val}}(t)$ average loss on 200 held-out examples 0.319, then 1.277
$g(t)$ the generalisation gap: how much worse the model does on data it hasn't seen −0.054, then 1.272

In words: "the gap is held-out loss minus training loss; a gap that keeps growing means the model is memorising rather than learning."

With the numbers: 0.319 − 0.373 = −0.054 at epoch 71; 1.277 − 0.005 = 1.272 at the end.

Level 3: in Python

In Python:

L_train = {"best": 0.373, "end": 0.005}
L_val = {"best": 0.319, "end": 1.277}
# g = L_val − L_train at each point
{t: round(L_val[t] - L_train[t], 3) for t in L_val}  # → {'best': -0.054, 'end': 1.272}

Training loss falls steadily to near zero while validation loss bottoms out at epoch 71 and then climbs to four times its best

Reading it: the x-axis is the epoch (log scale), the y-axis the loss. Both curves fall at first while the model learns the real rule. At epoch 71 (the dashed line) validation loss bottoms out; after that the training curve keeps falling as the model memorises the three wrong labels, and the validation curve climbs. The dotted line is where early stopping with a patience of 20 epochs would end the run, keeping the weights from epoch 71.

flowchart LR EP[Train one epoch] --> SV[Save a checkpoint] SV --> SC[Score it on the<br/>held-out set] SC --> Q{Best so far?} Q -->|yes| MARK[Mark it best] --> EP Q -->|"no, patience used up"| SHIP[Ship the best checkpoint,<br/>not the last] Q -->|"no, patience left"| EP

Reading it: every epoch ends with a checkpoint and a held-out score. The loop keeps going while scores improve or patience remains, and when it stops, the checkpoint that ships is the marked best, not whatever the last epoch left behind.

In code: overfitting_run fine-tunes on the 16 examples and records both losses every epoch; primer.ml.regularization.early_stopping replays the validation curve and returns the best epoch and the stopping epoch.

Why it matters in practice. Fine-tuning datasets are small compared to pretraining, and big models memorise quickly, so fine-tunes typically run for only a few epochs. Save checkpoints, score each on the held-out set, and ship the best one.

5. Model merging

5a. Weight averaging and task arithmetic

Everyday picture. Two editors each take a copy of the same draft and make tracked changes: one fixes the grammar, the other tightens the argument. You can apply both sets of changes to the original. If instead you "average" the two edited copies, each change is applied at half strength: half the grammar fixed, half the argument tightened.

Tiny worked example. A three-weight base (1, 0, 2). One fine-tune moves it to (1.5, 0, 2); another to (1, −1, 2).

Weights Change from the base
base (1, 0, 2)
fine-tune on A (1.5, 0, 2) τ_A = (0.5, 0, 0)
fine-tune on C (1, −1, 2) τ_C = (0, −1, 0)
base + τ_A + τ_C (1.5, −1, 2) both changes in full
average of the two fine-tunes (1.25, −0.5, 2) both changes at half strength

The change a fine-tune made, θ_ft − θ_base, is its task vector. Adding task vectors to the base is task arithmetic.

Level 3: the formula and its symbols

$$ \tau_t = \theta_t - \theta_{\text{base}}, \qquad \theta_{\text{merged}} = \theta_{\text{base}} + \lambda \sum_{t=1}^{T} \tau_t $$

Symbols

Symbol Meaning here In the example
$t$ which task: a counter over the fine-tunes A, C
$T$ how many fine-tunes are merged 2
$\theta_t$ all the weights of the model fine-tuned on task $t$ (1.5, 0, 2)
$\theta_{\text{base}}$ the weights they all started from (1, 0, 2)
$\tau_t$ tau: task $t$'s task vector, everything its fine-tune changed (0.5, 0, 0)
$\sum_{t=1}^{T}$ add up the task vectors τ_A + τ_C = (0.5, −1, 0)
$\lambda$ lambda: how strongly to apply the combined changes 1
$\theta_{\text{merged}}$ the merged model's weights (1.5, −1, 2)

In words: "each task vector is what its fine-tune changed; add the changes up, scale them by λ, and apply them to the base."

With the numbers: θ_merged = (1, 0, 2) + 1 × (0.5, −1, 0) = (1.5, −1, 2). With λ = 1/2: (1, 0, 2) + 0.5 × (0.5, −1, 0) = (1.25, −0.5, 2), which is exactly the average of the two fine-tunes. Averaging T fine-tunes is task arithmetic with λ = 1/T.

Level 3: in Python

In Python:

theta_base = [1.0, 0.0, 2.0]
theta_A = [1.5, 0.0, 2.0]
theta_C = [1.0, -1.0, 2.0]
# τ_t = θ_t − θ_base
tau_A = [a - b for a, b in zip(theta_A, theta_base)]
tau_C = [c - b for c, b in zip(theta_C, theta_base)]
tau_A, tau_C  # → ([0.5, 0.0, 0.0], [0.0, -1.0, 0.0])
# θ_merged at λ = 1
[b + 1.0 * (a + c) for b, a, c in zip(theta_base, tau_A, tau_C)]  # → [1.5, -1.0, 2.0]
# λ = 1/2 ...
[b + 0.5 * (a + c) for b, a, c in zip(theta_base, tau_A, tau_C)]  # → [1.25, -0.5, 2.0]
# ... is the plain average of the two fine-tunes
[(a + c) / 2 for a, c in zip(theta_A, theta_C)]  # → [1.25, -0.5, 2.0]
flowchart LR B[Base θ] -->|fine-tune on A| FA[θ_A] B -->|fine-tune on C| FC[θ_C] FA --> TA["τ_A = θ_A − θ_base"] FC --> TC["τ_C = θ_C − θ_base"] TA --> SUM["λ × (τ_A + τ_C)"] TC --> SUM B --> ADD((+)) SUM --> ADD ADD --> M[Merged model:<br/>no extra training]

Reading it: both fine-tunes start from the same base; that shared start is what makes their changes comparable. Subtracting the base turns each fine-tuned model into a task vector, the arrows are added and scaled, and the result is added back to the base. No data and no gradient steps are involved: merging is arithmetic on weights, and the merged model is the same size and speed as the base.

Merging the fine-tunes on A and C: at lambda 1 the merged model scores 0.97 on both, while plain averaging (lambda one half) scores only 0.635 on A

Reading it: the x-axis is λ and the y-axis held-out accuracy of the merged model on A (blue) and on C (green); the dotted line is 0.5, a coin flip. At λ = 1/2, which is plain averaging, each skill is diluted: 0.635 on A. Around λ = 1 the merged model does both tasks at 0.97, as well as either specialist on its own task. Push λ further and the changes overshoot. Their task vectors are nearly perpendicular (cosine −0.04), so adding one barely disturbs the other.

In code: task_vector subtracts the base, merge adds scaled task vectors back, and merge_run merges two fine-tunes at several λ and scores the result.

Why it matters in practice. Merging combines skills trained separately, by different teams or on data that can't be pooled, without any further training and at no extra inference cost. Averaging several fine-tunes of the same task ("model soups", Wortsman et al.) often beats the best single one. Task vectors can also be subtracted: Ilharco et al. negated a task vector learned from toxic text to make a model less toxic.

5b. Interference: when task vectors collide

Everyday picture. Two editors rewrote the same sentence in opposite directions, one making it warmer and one making it colder. Applying both sets of tracked changes gives a sentence neither intended.

Tiny worked example. Our three-weight changes again, plus a new one: τ_A = (0.5, 0, 0), τ_C = (0, −1, 0) and τ_B = (−0.4, 0, 0.3). τ_A and τ_C touch different weights: no conflict. τ_B pulls the first weight the other way from τ_A. The cosine measures how aligned two changes are:

Level 3: the formula and its symbols

$$ \cos(\tau_A, \tau_B) = \frac{\tau_A \cdot \tau_B}{\lVert \tau_A \rVert \, \lVert \tau_B \rVert} $$

Symbols

Symbol Meaning here In the example
$\tau_A \cdot \tau_B$ dot product: multiply matching entries, then add 0.5 × (−0.4) = −0.2
$\lVert \tau \rVert$ a vector's length 0.5 and 0.5
$\cos$ the cosine of the angle between the two changes: 1 same direction, 0 unrelated, −1 opposite −0.8

In words: "multiply the changes weight by weight and add, then divide by both lengths, so only the direction counts."

With the numbers: τ_A · τ_B = 0.5 × (−0.4) + 0 + 0 = −0.2; ‖τ_A‖ = 0.5, ‖τ_B‖ = √(0.16 + 0.09) = 0.5; cos = −0.2 / 0.25 = −0.8: strongly opposed. cos(τ_A, τ_C) = 0: independent. See primer.ml.embeddings.similarity for the cosine from scratch.

Level 3: in Python

In Python:

import math
tau_A = [0.5, 0.0, 0.0]
tau_B = [-0.4, 0.0, 0.3]
tau_C = [0.0, -1.0, 0.0]
def cos(u, v):
    dot = sum(a * b for a, b in zip(u, v))
    return dot / (math.sqrt(sum(a * a for a in u)) * math.sqrt(sum(b * b for b in v)))
# opposed changes
round(cos(tau_A, tau_B), 2)  # → -0.8
# changes to different weights
cos(tau_A, tau_C)  # → 0.0

In the toy, the fine-tunes on A and on B (which reverses A's rule) have task vectors with cosine −0.25, against −0.04 for A and C. Each fine-tune learned its rule everywhere, not just in its own region, so the two vectors rewrite the same weights in opposite directions.

Merging the fine-tunes on A and B: at every lambda at least one task stays at or below a coin flip, and the best the merge manages on both at once is 0.525

Reading it: the same axes as the previous figure, now merging A with B. No value of λ lifts both lines: wherever one task improves, the other sits near or below a coin flip, and from λ = 1 on both stay there, because the two changes cancel. A model can do both tasks (replay in section 3c got 0.945 and 0.99), but this merge can't reach it: that model needs to tell the regions apart, and neither fine-tune learned to.

flowchart TD TV[Task vectors from<br/>the same base] --> COS{Cosine between them} COS -->|"near 0: separate weights"| ADD[Add them:<br/>task arithmetic] COS -->|"clearly negative: conflict"| FIX{Can you retrain?} FIX -->|yes| JOINT[Train one model on<br/>both datasets, or replay] FIX -->|no| TIES["Resolve conflicts:<br/>trim small changes,<br/>agree on a sign per weight"] ADD --> EV[Score the merge on<br/>every task's held-out set] TIES --> EV JOINT --> EV

Reading it: a cheap check before merging is the cosine between task vectors. Near zero, the changes live in different weights and adding them usually works. Clearly negative, they fight: train on both datasets together if you can; if you can't, conflict-resolving merges such as TIES-merging (Yadav et al.) drop each task's smallest changes and, for each weight, keep only the changes that agree with the majority sign. Every path ends at the same place: score the merge on every task's held-out set.

In code: cosine measures the angle between two task vectors, and merge_run reports it alongside the merged accuracies.

Why it matters in practice. Merges are free to try, which makes them tempting to trust. A merged model can quietly lose a skill both parents had, so it is evaluated like any new model. The cosine tells you in advance which merges to be nervous about.

In 20 seconds

  • Decide with an eval: build a held-out set first; try prompting, then retrieval, and fine-tune only for behaviour a prompt can't pin down. It pays off after C / (c_prompt − c_tuned) requests.
  • Data: use the model's chat template and train only on assistant turns; deduplicate (Jaccard on word shingles); freeze the held-out set before training and remove near-copies of it; audit labels, because wrong labels cap what the eval can show. Quality beats quantity.
  • Forgetting: every weight is shared, so learning a new task erodes old ones, in proportion to how far the weights move and how much the tasks conflict. Fewer steps, a lower learning rate and LoRA shorten the trip; replaying a little old data is what keeps a contradicted skill.
  • Overfitting: on small data, validation loss bottoms out early; ship the best checkpoint, not the last.
  • Merging: a task vector is θ_ft − θ_base. Adding task vectors combines skills without training; averaging is λ = 1/T and dilutes them; vectors that point in opposite directions interfere.

Self-test questions

When is fine-tuning the wrong tool, and what should you try first? When the model lacks facts, or the facts change: retrieval supplies them and is easy to update. When a clearer prompt with a few examples fixes the behaviour: that is cheaper and survives a base-model upgrade. Fine-tuning earns its cost when behaviour stays inconsistent under the best prompt, or when a long prompt sent millions of times costs more than the fine-tune.

Why build the held-out set before training, and why check it against the training set? So that no choice (prompt, learning rate, checkpoint) is made by looking at it, and it stays an honest measure. A training example that nearly copies a held-out one lets the model recite the answer, turning the eval into a memory test; deduplicating across the split prevents it.

A held-out set has 100 examples and the fine-tune scores 83% against the prompt's 80%. Has it won? Not yet. At 80% on 100 examples the 95% margin is about ±7.8 points, so a 3-point difference is well inside the noise. You need a larger held-out set (400 examples halve the margin) or a bigger difference.

If 10% of the eval's labels are wrong, what is the best score a perfect model can get? 90%, because it is marked wrong on every mislabelled item. A 90%-accurate model would score 0.9 × 0.9 + 0.1 × 0.1 = 82%. Noisy labels shrink and blur the differences you are trying to measure.

What is catastrophic forgetting, and why does it happen? Training on a new task alone erodes, or wipes out, skills the model had. Every weight is shared between tasks, and only the new task's examples produce gradients, so nothing pushes back when a change that helps the new task hurts an old one. It grows with how far the weights move and with how much the new task conflicts with the old.

Lowering the learning rate did not stop task A being forgotten. Why not, and what works? Task B contradicts A on the same inputs, so any progress on B costs A; a lower rate only walks the same trade-off more slowly. Replay works: mixing even 5% of A's examples into B's data gives the model a reason to keep A, and those few examples carry most of the loss exactly when A is slipping.

Training loss keeps falling, but validation loss has risen since epoch 71. What is happening, and which checkpoint do you ship? The model has stopped learning the general rule and is memorising the training set, including its mislabelled examples. Ship the checkpoint from epoch 71, the best on the held-out set; early stopping automates exactly this.

What is a task vector, and why is averaging two fine-tunes the same as task arithmetic with λ = 1/2? A task vector is everything a fine-tune changed: θ_ft − θ_base. The average of two fine-tunes is (θ_base + τ_1 + θ_base + τ_2) / 2 = θ_base + ½(τ_1 + τ_2), which is task arithmetic with λ = 1/2, so each skill arrives at half strength.

When does merging fail, and how can you see it coming? When the task vectors change the same weights in opposite directions, so adding them cancels both skills. A clearly negative cosine between task vectors is the warning; the remedy is joint training or replay, or a conflict-resolving merge such as TIES, and a held-out check on every task either way.

The papers behind this lesson

  • Ilharco et al., Editing Models with Task Arithmetic (2022): https://arxiv.org/abs/2212.04089. Defined task vectors as fine-tuned minus pretrained weights and showed that adding them combines skills, and negating them removes a behaviour. Annotated companion
  • Wortsman et al., Model soups: averaging weights of multiple fine-tuned models improves accuracy without increasing inference time (2022): https://arxiv.org/abs/2203.05482. Showed that averaging the weights of several fine-tunes of one base often beats the best single fine-tune.
  • Yadav et al., TIES-Merging: Resolving Interference When Merging Models (2023): https://arxiv.org/abs/2306.01708. Traced failed merges to small redundant changes and sign conflicts, and fixed both by trimming and electing a sign per weight.
  • Kirkpatrick et al., Overcoming catastrophic forgetting in neural networks (2017): https://arxiv.org/abs/1612.00796. Introduced elastic weight consolidation, which slows learning on the weights most important to earlier tasks.
  • Biderman et al., LoRA Learns Less and Forgets Less (2024): https://arxiv.org/abs/2405.09673. Measured on real language models that LoRA keeps more of the base model's abilities than full fine-tuning, at the price of learning the new task less completely.
  • Zhou et al., LIMA: Less Is More for Alignment (2023): https://arxiv.org/abs/2305.11206. Fine-tuned a large base model on 1,000 carefully curated examples and got a strong assistant, evidence that example quality matters more than quantity. Annotated companion
  • Lee et al., Deduplicating Training Data Makes Language Models Better (2021): https://arxiv.org/abs/2107.06499. Found widespread near-duplicates in standard datasets, including between training and test sets, and showed that removing them reduces memorisation.

Further reading

on GitHub
   1r"""
   2# Fine-tuning in practice: preparing data, forgetting old skills, merging models
   3
   4Run: `python -m primer.ml.fine_tuning`
   5
   6New to the notation? `primer.notation` explains every symbol used here
   7(Σ, ‖x‖, λ, subscripts, and so on) from zero.
   8
   9This lesson builds on `primer.ml.training_stages`, which explains *what*
  10supervised fine-tuning (SFT), preference tuning and LoRA are. Here we deal
  11with what happens when you actually do it: deciding whether to fine-tune at
  12all, preparing the data, the old skills the model loses along the way, the
  13small dataset it memorises, and how two fine-tuned models can be merged into
  14one by plain arithmetic on their weights.
  15
  16## Level 1: The practitioner's guide
  17
  18**In one sentence.** Fine-tuning in practice is the work around the training
  19run: deciding with an evaluation set whether to train at all, preparing
  20examples that teach what you mean, keeping the skills the model already had,
  21stopping before it memorises your data, and sometimes combining two
  22fine-tunes by arithmetic on their weights instead of training a third.
  23
  24**When you need it.** You need this lesson when a fine-tune is on the table:
  25a behaviour the best prompt cannot make consistent, or a long prompt sent so
  26often that its tokens are most of the bill. The tell for the first is an
  27eval score that stops improving however the prompt is reworded; the tell for
  28the second is arithmetic. In this lesson's worked example, a 3,000-token
  29prompt at \$2 per million tokens costs \$0.0060 per request, a tuned model
  30that needs 300 tokens at \$4 per million costs \$0.0012, and a \$600
  31fine-tune pays for itself after 125,000 requests (25 days at 5,000 a day;
  32the prices are illustrative). You don't need fine-tuning to add facts,
  33which go stale the day they change and belong in retrieval, and you don't
  34need it while a clearer instruction with two examples still moves the eval.
  35Nothing in this lesson matters until the eval exists: it is the first box of
  36the decision, before any model.
  37
  38**Your options.** From the cheapest to the most committed:
  39
  40| Option | What it does | What it guarantees | What it costs | Where it lives |
  41|---|---|---|---|---|
  42| The best prompt, scored on the eval | Instructions and examples in the prompt, measured on a frozen held-out set | A baseline every other option must beat; survives a base-model upgrade | Tokens on every call | Your prompt |
  43| Hosted supervised fine-tuning | Upload chat examples; the vendor applies the chat template, masks the user turns, trains and hosts | Consistent format and style without the long prompt | Curated examples, a training job, often a higher per-token price | The vendor's fine-tuning API |
  44| LoRA on an open model | Train a small adapter beside frozen weights on your examples | Learns the task with less forgetting of the base model's skills (Biderman et al., 2024) | A GPU, data, a serving stack; may learn a hard new domain less completely | Your training stack |
  45| Full fine-tune, with replay | Update every weight on your task mixed with a little general data | The largest change the model can make, and the old skills kept if you replay them | Multi-GPU training, the most forgetting when you don't replay, a copy of the model per variant | Your training stack |
  46| Merge existing fine-tunes | Add each fine-tune's task vector (its change from the base) to the base, no training | Two separable skills in one model at no extra inference cost | An eval on every task; fails when the two changes fight | Your weights, with a merging toolkit |
  47
  48**How to choose.** Climb down only when the eval proves the rung above
  49falls short.
  50
  51- A format or a house style: a few dozen to a few hundred excellent
  52  examples, hosted or LoRA. Start with 50 to 100 and double only while a
  53  doubling still beats the eval's margin of error (section 2e).
  54- A behaviour that contradicts the base model ("always JSON" against "chat
  55  naturally", "be terse" against "explain"): expect the old behaviour to
  56  vanish wherever your data overrides it, and mix in general
  57  instruction-following data so the model stays an assistant (section 3c).
  58- A harder specialised skill: thousands of examples, and consider a full
  59  fine-tune with replay, since LoRA learns less on a demanding domain.
  60- Two skills trained separately, by different teams or on data that cannot
  61  be pooled: try a merge first, and check the cosine between the task
  62  vectors before you trust it (section 5b).
  63- Whatever you pick, the fine-tune ships only if it beats the prompt on the
  64  same held-out set by more than that set's margin of error, and a general
  65  eval sits beside the task eval.
  66
  67**What it costs.** Money is the one-off cost of writing and checking data
  68plus the run, against a per-request saving that may be zero if the tuned
  69model is priced like the base one; the break-even rule in section 1 puts a
  70number on it. Every new base model repeats the one-off cost. Data is the
  71expensive part, and the eval is dearer than people expect: 100 held-out
  72examples measured at 80% mean "somewhere between 72% and 88%", and it takes
  73400 to halve that margin to ±3.9 points (this lesson's `margin_of_error`).
  74Labels cost accuracy too: with 10% of reference labels wrong, a perfect
  75model scores 90% and a 90% model scores 82%. Training itself is short.
  76In the lesson's toy the new task is learned in 5 steps and every step after
  77only erodes the general skill, and fine-tunes on small data run for a few
  78epochs, with the best checkpoint shipped rather than the last.
  79
  80**What breaks.**
  81
  82- **The wrong chat template.** Training succeeds, the loss falls, and the
  83  deployed model sees role markers it never learned. Use the exact template
  84  the base model was trained with, and the system prompt you will deploy.
  85- **Training on the user's turns.** Forget the loss mask and the model
  86  learns to write questions. Hosted APIs mask for you; `render_chat` shows
  87  which segments count.
  88- **A leaking eval.** A training example that nearly copies a held-out one
  89  turns the eval into a memory test. Freeze the held-out set first and
  90  remove near-copies from training (Jaccard on word pairs, threshold 0.7 in
  91  the toy).
  92- **Forgetting.** Fine-tuning on one task takes the toy's general skill
  93  from 0.99 to 0.655, and a task that contradicts an earlier one takes the
  94  earlier one from 0.98 to 0.025. A lower learning rate only slows the
  95  slide; 10 replayed examples in 210 bring it back to 0.945.
  96- **Memorising a small dataset.** On 16 examples with 3 wrong labels,
  97  validation loss bottoms out at epoch 71 (0.319) and has quadrupled
  98  (1.277) by epoch 1,500. Checkpoint every epoch and ship the best.
  99- **A merge that cancels.** Task vectors with a clearly negative cosine
 100  (−0.25 in the toy) leave at least one skill at a coin flip whatever the
 101  scale. Retrain jointly, or use a conflict-resolving merge.
 102
 103**In the wild.** Hosted fine-tuning APIs take chat-formatted examples and
 104handle the template and the mask; OpenAI's model optimization guide lists
 105supervised fine-tuning, DPO and reinforcement fine-tuning. For open models,
 106Hugging Face TRL's SFTTrainer runs the supervised loop and PEFT supplies
 107LoRA. Merging has its own toolkit: mergekit implements linear averaging,
 108SLERP, task arithmetic, TIES, DARE and more. The evidence behind the advice
 109is in the papers at the end of the lesson: LIMA's 1,000 curated examples,
 110model soups for averaging fine-tunes of one task, task arithmetic for adding
 111and subtracting them, TIES for resolving interference, elastic weight
 112consolidation for when the old data is gone, and Biderman et al. on LoRA
 113learning less and forgetting less.
 114
 115**Go deeper.** Level 2 puts every claim above on a 97-weight model you can
 116train in a fraction of a second: the break-even formula, the chat mask,
 117Jaccard deduplication, the margin of error and the label-noise ceiling,
 118forgetting measured as distance from the base, the overfitting curve with
 119early stopping, and task vectors merged at every scale. If you only needed
 120to decide whether and how to fine-tune, you are done.
 121
 122## Level 2: How it works, from scratch
 123
 124Picture a skilled cook who joins your restaurant. They already know how to
 125cook (that is the pretrained model). You want them to cook *your* menu, *your*
 126way. You can hand them a note with every order (a prompt), give them the
 127recipe binder to look things up in (retrieval, RAG), or send them on a course
 128about your kitchen (fine-tuning). The course is the only option that changes
 129the cook, and it comes with the risks every teacher knows: the lessons can
 130be badly written, the cook can cram the practice exam instead of learning,
 131and a month of drilling one cuisine can make them rusty at everything else.
 132
 133Everything below is one of those risks, measured on a model small enough to
 134train in a fraction of a second.
 135
 136## 1. Should you fine-tune at all?
 137
 138**Everyday picture.** Giving the cook a note with every order costs a
 139little each time. The course costs a lot once, and it has to be repeated
 140whenever you hire a new cook (switch to a newer base model). The course pays
 141off only when the notes would otherwise be long, sent millions of times, or
 142simply not enough to make the cook consistent.
 143
 144**Tiny worked example.** A support bot sends a 3,000-token prompt full of
 145instructions and examples on every request. A fine-tuned model has learned
 146that behaviour and needs only 300 tokens of prompt. Take these illustrative
 147prices (real ones vary by provider and change often): \$2 per million input
 148tokens for the general model, \$4 per million for the tuned one, because
 149hosting a custom model usually costs more per token.
 150
 151| | Tokens per request | Price per million | Cost per request |
 152|---|---|---|---|
 153| prompted | 3,000 | \$2 | 3,000 × 2 / 1,000,000 = **\$0.0060** |
 154| fine-tuned | 300 | \$4 | 300 × 4 / 1,000,000 = **\$0.0012** |
 155
 156Each request saves \$0.0048. Writing and checking 1,000 training examples
 157plus the training run costs, say, \$600 once. After 600 / 0.0048 = **125,000
 158requests** the course has paid for itself: 25 days at 5,000 requests a day.
 159
 160```mermaid
 161flowchart TD
 162  E[Build the eval set first] --> P[Best prompt you can write]
 163  P --> M1{Good enough<br/>on the eval?}
 164  M1 -->|yes| SHIP[Ship the prompt]
 165  M1 -->|"no: missing facts"| R[Add retrieval, RAG]
 166  R --> M2{Good enough<br/>on the eval?}
 167  M2 -->|yes| SHIP
 168  M2 -->|"no: wrong behaviour, format or style"| FT[Fine-tune, often with LoRA]
 169  FT --> M3{Beats the prompt<br/>on the SAME eval?}
 170  M3 -->|yes, and pays off| SHIPFT[Ship the fine-tune]
 171  M3 -->|no| P
 172```
 173
 174**Reading it:** start at the top, and notice that the first box is not a
 175model at all: it is the evaluation set, the fixed exam every option is
 176scored on. Each rung is cheaper and faster to change than the one below it,
 177so you climb down only when the eval proves the rung above falls short.
 178Missing knowledge goes to retrieval, because fine-tuning is a poor way to
 179add facts; wrong behaviour that no prompt can pin down goes to fine-tuning.
 180The last diamond matters most: a fine-tune ships only if it beats the prompt
 181on the same eval. `primer.ml.training_stages` has the same decision as a
 182flowchart of approaches; this one adds the measurement at every step.
 183
 184$$
 185N^\star = \frac{C_{\text{once}}}{c_{\text{prompt}} - c_{\text{tuned}}},
 186\qquad c = \frac{t \cdot p}{10^6}
 187$$
 188
 189**Symbols**
 190
 191| Symbol | Meaning here | In the example |
 192|---|---|---|
 193| $N^\star$ | "N-star": the break-even number of requests | 125,000 |
 194| $C_{\text{once}}$ | one-off cost: writing and checking data, the training run | \$600 |
 195| $c$ | cost of one request | |
 196| $c_{\text{prompt}}, c_{\text{tuned}}$ | cost of one request with the long prompt, and with the fine-tuned model | \$0.0060, \$0.0012 |
 197| $t$ | input tokens sent per request | 3,000 and 300 |
 198| $p$ | price in dollars per million input tokens | \$2 and \$4 |
 199| $10^6$ | one million: prices are quoted per million tokens | |
 200
 201**In words:** "divide what the course costs once by what it saves on each
 202request; that is how many requests it takes to earn the cost back."
 203
 204**With the numbers:** $c_{\text{prompt}}$ = 3,000 × 2 / 10⁶ = 0.006,
 205$c_{\text{tuned}}$ = 300 × 4 / 10⁶ = 0.0012, so $N^\star$ = 600 / 0.0048 =
 206125,000.
 207
 208**In Python:**
 209
 210```python
 211# c = t·p / 1,000,000 for each option
 212c_prompt = 3000 * 2.0 / 1_000_000
 213c_tuned = 300 * 4.0 / 1_000_000
 214c_prompt, c_tuned  # → (0.006, 0.0012)
 215C_once = 600
 216# N* = C_once / (c_prompt − c_tuned)
 217N_star = C_once / (c_prompt - c_tuned)
 218round(N_star)  # → 125000
 219# days to break even at 5,000 requests a day
 220round(N_star / 5000, 1)  # → 25.0
 221```
 222
 223![The prompted line starts at zero and climbs steeply; the fine-tuned line starts at 600 dollars and climbs slowly; they cross at 125,000 requests](figures/primer.ml.fine_tuning.break_even.svg)
 224
 225**Reading it:** the x-axis counts requests served and the y-axis is the
 226total money spent so far. The prompted line starts at zero but climbs
 227steeply, because every request pays for 3,000 tokens. The fine-tuned line
 228starts at \$600 (the one-off cost) and climbs slowly. Left of the dashed
 229line, prompting is cheaper; right of it, the fine-tune is. If the tuned
 230model cost as much per request as the prompt, the two lines would never
 231cross, and the only reason left to fine-tune would be quality.
 232
 233**In code:** `per_request_cost` prices one request and
 234`break_even_requests` returns $N^\star$, or infinity when the tuned model
 235saves nothing per request.
 236
 237**Why it matters in practice.** Three costs hide outside this formula.
 238Every new base model means repeating the fine-tune, so $C_{\text{once}}$
 239recurs. Facts learned by fine-tuning go stale and are hard to update, which
 240is why knowledge belongs in retrieval (`primer.agents.rag`). And the
 241engineering time to build an eval is spent whichever way you go, which is
 242why it comes first.
 243
 244## 2. Preparing the data
 245
 246**Everyday picture.** Before the course, someone writes the course book.
 247Every recipe must be written in the same layout, no recipe may appear five
 248times, some recipes are locked in a drawer for the final exam *before* the
 249cook ever sees the book, and the recipes must actually be right. Most
 250fine-tuning failures are course-book failures.
 251
 252```mermaid
 253flowchart LR
 254  RAW[Raw examples<br/>logs, experts, drafts] --> FMT[Format as chat<br/>messages]
 255  FMT --> VAL[Validate<br/>roles, empty turns]
 256  VAL --> DD[Deduplicate<br/>near-copies]
 257  DD --> SPLIT[Set aside the<br/>held-out set, frozen]
 258  SPLIT --> LEAK[Drop training examples<br/>that copy held-out ones]
 259  LEAK --> AUDIT[Audit a sample<br/>of the labels]
 260  AUDIT --> TRAIN[Training set]
 261  SPLIT --> EVAL[Held-out set]
 262```
 263
 264**Reading it:** raw examples enter on the left and pass through five
 265filters. The branch at "Set aside the held-out set" is the important one:
 266the held-out examples leave the pipeline *before* anything is trained or
 267tuned, and the step after it removes any training example that nearly
 268copies one of them. The subsections below build each box.
 269
 270### 2a. Formatting chat examples
 271
 272**Everyday picture.** A play script: every line starts with who speaks it.
 273The model learns the play by reading scripts, and it is graded only on the
 274lines of the character it will play, the assistant.
 275
 276**Tiny worked example.** One example as a list of role-tagged messages:
 277
 278| Turn | Role | Content | Trained on? |
 279|---|---|---|---|
 280| 1 | system | Be brief. | no |
 281| 2 | user | Capital of France? | no |
 282| 3 | assistant | Paris. | **yes** |
 283
 284Rendered into text, it becomes
 285`<|system|>Be brief.<|end|><|user|>Capital of France?<|end|><|assistant|>Paris.<|end|>`,
 286and only the last segment counts towards the loss.
 287
 288```mermaid
 289flowchart LR
 290  S["system: Be brief."] --> U["user: Capital of France?"] --> A["assistant: Paris."]
 291  S -.->|"context only, no loss"| L[Loss]
 292  U -.->|"context only, no loss"| L
 293  A ==>|"every token scored"| L
 294```
 295
 296**Reading it:** the solid arrows are the order the model reads the turns.
 297The dotted arrows carry no training signal: the system and user turns are
 298context the model conditions on. Only the thick arrow feeds the loss, so the
 299model learns to *answer*, not to write questions. This is the loss mask from
 300`primer.ml.training_stages`, section 2.
 301
 302**The rules that matter.** Use the exact chat template the base model was
 303trained with; the role markers above are illustrative, and each model family
 304has its own. Train with the same system prompt you will deploy with. Reject
 305examples whose last turn is not the assistant's (nothing to learn), whose
 306turns are empty, or whose roles are unknown.
 307
 308**In code:** `chat_example` builds the message list, `render_chat` flattens
 309it into (segment, trained?) pairs, and `validate_chat` lists every problem
 310with an example.
 311
 312**Why it matters in practice.** A template mismatch is silent: training
 313succeeds, the loss falls, and the deployed model sees markers it never
 314learned, so it behaves like the base model or worse.
 315
 316### 2b. Deduplication
 317
 318**Everyday picture.** A flashcard deck with the same card in it five times.
 319You study that card five times as often, and you start answering every
 320question with it.
 321
 322**Tiny worked example.** Compare questions by their **word pairs** (two
 323neighbouring words; a window of k words is called a **shingle**), after
 324lowercasing and dropping punctuation:
 325
 326| Text | Word pairs |
 327|---|---|
 328| "How do I reset my password?" | how do, do i, i reset, reset my, my password |
 329| "how do I reset my password, please" | the same 5, plus password please |
 330| "How do I change my email?" | how do, do i, i change, change my, my email |
 331
 332The first two share 5 of the 6 distinct pairs between them: a
 333near-duplicate. The first and third share 2 of 8: different questions.
 334
 335$$
 336J(A, B) = \frac{|A \cap B|}{|A \cup B|}
 337$$
 338
 339**Symbols**
 340
 341| Symbol | Meaning here | In the example |
 342|---|---|---|
 343| $A, B$ | the sets of word pairs of two texts | 5 pairs and 6 pairs |
 344| $\cap$ | intersection: pairs in both sets | 5 shared pairs |
 345| $\cup$ | union: pairs in either set, each counted once | 6 distinct pairs |
 346| $\lvert \cdot \rvert$ | the number of items in a set | |
 347| $J(A, B)$ | Jaccard similarity, from 0 (nothing shared) to 1 (identical) | 0.83 |
 348
 349**In words:** "count the pairs the two texts share, and divide by the number
 350of distinct pairs they have between them."
 351
 352**With the numbers:** J = 5 / 6 = 0.83 for the two password questions and
 3532 / 8 = 0.25 for password against email. A threshold of 0.7 keeps the first
 354password question, drops its rewording, and keeps the email question.
 355
 356**In Python:**
 357
 358```python
 359def pairs(text):
 360    words = text.lower().replace("?", "").replace(",", "").split()
 361    return {(a, b) for a, b in zip(words, words[1:])}
 362A = pairs("How do I reset my password?")
 363B = pairs("how do I reset my password, please")
 364C = pairs("How do I change my email?")
 365# |A ∩ B| and |A ∪ B|
 366len(A & B), len(A | B)  # → (5, 6)
 367# J(A, B)
 368round(len(A & B) / len(A | B), 2)  # → 0.83
 369# J(A, C)
 370len(A & C) / len(A | C)  # → 0.25
 371```
 372
 373```mermaid
 374flowchart LR
 375  T[Next example] --> CMP{Jaccard with any<br/>kept example ≥ 0.7?}
 376  CMP -->|yes| DROP[Drop it:<br/>a near-copy]
 377  CMP -->|no| KEEP[Keep it]
 378  KEEP --> T
 379  DROP --> T
 380```
 381
 382**Reading it:** examples arrive one at a time and are compared against
 383everything already kept. A match above the threshold is dropped; anything
 384else joins the kept list. The first copy of each group survives, so the
 385order of the data decides which wording you keep.
 386
 387**In code:** `normalize` lowercases and strips punctuation, `shingles`
 388collects the word pairs, `jaccard` scores two texts and `deduplicate`
 389returns the indices worth keeping. Comparing every pair is fine for
 390thousands of examples; at web scale, MinHash estimates the same Jaccard
 391without comparing every pair.
 392
 393**Why it matters in practice.** Duplicates overweight a few examples, so the
 394model parrots them. Worse, a near-copy of an eval question in the training
 395set lets the model recite the answer, and the eval score becomes a memory
 396test. Lee et al. found training sets with thousands of near-duplicates, and
 397removing them made models memorise less.
 398
 399### 2c. The held-out set comes first
 400
 401**Everyday picture.** A good teacher writes the final exam before teaching
 402the course and locks it in a drawer. If the exam were written afterwards,
 403it would drift towards what the class happened to practise.
 404
 405**Tiny worked example.** Fifty support questions are deduplicated, shuffled
 406with a fixed seed, and 10 (20%) go in the drawer. Only then are the other 40
 407used, and any of the 40 that nearly copies one of the 10 is removed. Every
 408later choice (prompt wording, learning rate, which checkpoint to ship) is
 409scored on those 10. But how much can 10, or even 100, questions tell you?
 410
 411$$
 412\text{margin} = z \sqrt{\frac{a\,(1 - a)}{n}}
 413$$
 414
 415**Symbols**
 416
 417| Symbol | Meaning here | In the example |
 418|---|---|---|
 419| $a$ | the accuracy measured on the held-out set | 0.8 |
 420| $n$ | the number of held-out examples | 100 |
 421| $\sqrt{a(1-a)/n}$ | the **standard error**: how much the measured accuracy would wobble if you drew a different held-out set of the same size | 0.04 |
 422| $z$ | how many standard errors to allow; 1.96 covers 95% of the wobble | 1.96 |
 423| margin | the true accuracy is probably within ± this of the measured one | 0.078 |
 424
 425**In words:** "the measured accuracy is uncertain by about two standard
 426errors, and the standard error shrinks with the square root of the number of
 427examples."
 428
 429**With the numbers:** with 100 examples at 80%, the margin is
 4301.96 × √(0.8 × 0.2 / 100) = 1.96 × 0.04 = 0.078, so "80%" means "somewhere
 431from about 72% to 88%". A fine-tune that scores 83% has not been shown to
 432beat a prompt that scores 80%. With 400 examples the margin halves to 0.039.
 433
 434**In Python:**
 435
 436```python
 437import math
 438a, n, z = 0.8, 100, 1.96
 439# the standard error sqrt(a(1 − a)/n)
 440round(math.sqrt(a * (1 - a) / n), 4)  # → 0.04
 441# the 95% margin
 442round(z * math.sqrt(a * (1 - a) / n), 4)  # → 0.0784
 443# four times as many examples halves it
 444round(z * math.sqrt(a * (1 - a) / 400), 4)  # → 0.0392
 445```
 446
 447```mermaid
 448flowchart LR
 449  ALL[All examples,<br/>deduplicated] --> SH[Shuffle with<br/>a fixed seed]
 450  SH --> EV[20% held out<br/>frozen, never trained on]
 451  SH --> TR[80% training]
 452  EV --> CHK{Training example<br/>near-copies one?}
 453  TR --> CHK
 454  CHK -->|yes| X[Remove from training]
 455  CHK -->|no| OK[Keep for training]
 456```
 457
 458**Reading it:** the fixed seed makes the split repeatable, so everyone on a
 459team scores against the same held-out set. The held-out branch is frozen the
 460moment it is made. Both branches then meet in the leak check, which only
 461ever removes examples from the *training* side.
 462
 463![The margin of error falls from about 16 points at 25 examples to 1.4 points at 3,200; each fourfold increase halves it](figures/primer.ml.fine_tuning.eval_margin.svg)
 464
 465**Reading it:** the x-axis is the number of held-out examples (log scale)
 466and the y-axis is the ± margin, in percentage points, around a measured 80%.
 467The marked points show the square-root law: 100 examples give ±7.8 points,
 468400 give ±3.9 and 1,600 give ±2.0. Each halving of the margin costs four
 469times the examples, which is why held-out sets of a few hundred are common
 470and a few dozen can only detect big differences.
 471
 472**In code:** `split_before_training` deduplicates, shuffles with a seed,
 473freezes the held-out set and removes leaks with `remove_leaks`;
 474`margin_of_error` gives the ± for any accuracy and size.
 475
 476**Why it matters in practice.** Every time you look at held-out scores and
 477change something, the held-out set leaks a little into your decisions. A
 478held-out set built first and used sparingly is the only honest measure of
 479whether the fine-tune helped. See `primer.ml.regularization` for train,
 480validation and test splits, and `primer.agents.evals` for building evals.
 481
 482### 2d. Label quality
 483
 484**Everyday picture.** An answer key with typos. A student who gets every
 485question right is marked wrong wherever the key is wrong, and a student who
 486makes the *same* mistake as the key is marked right.
 487
 488**Tiny worked example.** A model is truly right 90% of the time, and 10% of
 489the reference labels are wrong. It scores a point when it is right on a
 490right label (0.9 × 0.9 = 0.81), or when it is wrong on a wrong label and the
 491two mistakes cancel (0.1 × 0.1 = 0.01). The eval reports **82%**, not 90%.
 492
 493$$
 494a_{\text{measured}} = a\,(1 - \varepsilon) + (1 - a)\,\varepsilon
 495$$
 496
 497**Symbols**
 498
 499| Symbol | Meaning here | In the example |
 500|---|---|---|
 501| $a$ | the model's true accuracy | 0.9 |
 502| $\varepsilon$ | epsilon: the share of reference labels that are wrong | 0.1 |
 503| $a(1-\varepsilon)$ | right answer, right label | 0.81 |
 504| $(1-a)\varepsilon$ | wrong answer on a wrong label (yes/no labels only, so the two errors agree) | 0.01 |
 505| $a_{\text{measured}}$ | what the eval reports | 0.82 |
 506
 507**In words:** "the eval gives credit for being right on a correct label, and
 508by accident for being wrong on a wrong one."
 509
 510**With the numbers:** 0.9 × 0.9 + 0.1 × 0.1 = 0.82. A perfect model
 511($a = 1$) scores only 1 − ε = 0.9: noisy labels put a ceiling on what the
 512eval can show.
 513
 514**In Python:**
 515
 516```python
 517a, eps = 0.9, 0.1
 518# right on a right label, plus wrong on a wrong label
 519round(a * (1 - eps) + (1 - a) * eps, 2)  # → 0.82
 520# even a perfect model scores only 1 − ε
 521round(1.0 * (1 - eps), 2)  # → 0.9
 522# two labelers, five items: how often do they agree?
 523labeler_1 = ["yes", "no", "yes", "yes", "no"]
 524labeler_2 = ["yes", "no", "no", "yes", "no"]
 525sum(p == q for p, q in zip(labeler_1, labeler_2)) / len(labeler_1)  # → 0.8
 526```
 527
 528```mermaid
 529flowchart TD
 530  X[One eval item] --> R{Model right?}
 531  R -->|"yes, 0.9"| LR{Label right?}
 532  R -->|"no, 0.1"| LW{Label right?}
 533  LR -->|"yes, 0.9"| P1["scored right: 0.81"]
 534  LR -->|"no, 0.1"| P2["scored wrong: 0.09"]
 535  LW -->|"yes, 0.9"| P3["scored wrong: 0.09"]
 536  LW -->|"no, 0.1"| P4["scored right: 0.01"]
 537```
 538
 539**Reading it:** each item takes two coin flips: is the model right, and is
 540the label right? Multiply along a path to get its share. Two paths are
 541scored as right, 0.81 and 0.01, which add up to the 0.82 the eval reports.
 542The 0.09 on the second path is a correct answer marked wrong by a bad label.
 543
 544**How to check label quality.** Have two people label the same sample
 545independently and measure how often they agree. Above, they agree on 4 of
 5465 items, 80%: one item in five is ambiguous or mislabelled, and your eval
 547can't resolve differences smaller than that noise. Read every disagreement;
 548they are usually unclear instructions, not careless labelers.
 549
 550**In code:** `measured_accuracy` applies the formula and `label_agreement`
 551compares two labelers.
 552
 553**Why it matters in practice.** Bad training labels teach the model the
 554mistakes (section 4 shows a small model memorising three of them). Bad eval
 555labels hide real improvements. Zhou et al. (LIMA) fine-tuned a large model
 556on just 1,000 carefully chosen examples and got a strong assistant: quality
 557of examples beats quantity.
 558
 559### 2e. How many examples?
 560
 561**Everyday picture.** Teaching a house style is not teaching a language.
 562A new cook learns how you plate a dish from a few dozen good examples; they
 563don't need ten thousand.
 564
 565**Tiny worked example.** Fine-tuning teaches a *behaviour* the base model
 566can almost do already, so the useful sizes are small: a few dozen examples
 567to show a format, hundreds for a consistent style or a narrow task,
 568thousands for a harder specialised skill. Section 4 fine-tunes on 16
 569examples and shows the other side: with so few, the model soon memorises
 570them, noise included.
 571
 572```mermaid
 573flowchart LR
 574  S[Start with 50 to 100<br/>high-quality examples] --> T[Fine-tune]
 575  T --> E[Score on the<br/>held-out set]
 576  E --> Q{Gained more than<br/>the margin of error?}
 577  Q -->|yes| D[Double the data] --> T
 578  Q -->|no| STOP[Stop adding data:<br/>fix quality or approach]
 579```
 580
 581**Reading it:** the loop grows the dataset only while each doubling still
 582buys a real improvement, meaning one larger than the held-out margin from
 583section 2c. When a doubling stops paying, more of the same data won't help;
 584the fix is better examples or a different approach.
 585
 586**Why it matters in practice.** Labelled data is the expensive part of
 587fine-tuning. Doubling until the gains flatten spends it where it helps, and
 588the held-out set tells you when to stop.
 589
 590## 3. Catastrophic forgetting
 591
 592**Everyday picture.** Someone who learned to drive in Britain, keeping
 593left, moves to the United States and practises keeping right every day for
 594a month. The new habit wins, and on a trip home they drift to the wrong side
 595of the road. Nobody told them to forget the old rule; the new practice
 596simply rewrote the reflex both rules use. Neural networks do this to an
 597extreme: train on a new task alone and an old one can vanish. This is
 598**catastrophic forgetting**.
 599
 600**The model we'll fine-tune.** To watch it happen, we need a model small
 601enough to train instantly. Each example is four numbers, and the answer is
 602yes or no:
 603
 604| Task | Inputs it uses | Rule | Plays the role of |
 605|---|---|---|---|
 606| general | all four, anywhere in [−3, 3] | yes when x₁ + x₃ > 0 | the base model's pretraining |
 607| A | x₁ in [−3, −1], x₂ in [−2, 2] | yes when x₂ > 0 | keep left in Britain |
 608| B | x₁ in [1, 3], x₂ in [−2, 2] | yes when x₂ < 0 | keep right in the US |
 609| C | x₃, x₄ in [−2, 2] | yes when x₃ + x₄ > 0 | an unrelated skill |
 610
 611A and B read the same two inputs and apply opposite rules in different
 612regions, so one model can learn both, but only by paying attention to the
 613region. C reads inputs that A and B never touch.
 614
 615```mermaid
 616flowchart LR
 617  X["4 inputs<br/>x1 x2 x3 x4"] --> H["16 hidden units<br/>tanh"]
 618  H --> O["1 output<br/>probability of yes"]
 619  TH["θ: all 97 weights<br/>in one vector"] -.-> H
 620  TH -.-> O
 621```
 622
 623**Reading it:** four numbers go in, 16 hidden units each mix all of them,
 624and one output unit turns the hidden units into a probability. Every weight
 625(64 in the first layer, 16 hidden biases, 16 output weights and one output
 626bias) sits in a single vector θ of 97 numbers. Every task flows through the
 627same hidden units, which is exactly why one task's training can damage
 628another. `primer.ml.neural_net` builds this kind of network from scratch.
 629
 630The base model is this network trained on the general task (99% accuracy
 631on held-out examples). Every fine-tune below starts from it and runs plain
 632gradient descent on 200 examples.
 633
 634**In code:** `TinyNet` holds θ and computes predictions, loss and
 635`TinyNet.gradient`; `make_task` draws examples for each of the `TASKS`;
 636`base_model` pretrains the base and `fine_tune` trains a copy, leaving the
 637starting model untouched.
 638
 639### 3a. The general skill fades while you teach a new one
 640
 641**Everyday picture.** A month of drilling one cuisine makes the cook a
 642little rusty at everything else, and a second month of drilling it makes
 643them rustier still, even though they had mastered the cuisine in the first
 644week.
 645
 646**Tiny worked example.** Fine-tune the base on task A (learning rate 0.5)
 647and check both skills on held-out examples:
 648
 649| After step | Accuracy on A | General skill | Distance from base |
 650|---|---|---|---|
 651| 0 (the base) | 0.53 | 0.99 | 0 |
 652| 5 | 0.975 | 0.825 | 2.60 |
 653| 10 | 0.995 | 0.79 | 2.83 |
 654| 300 | 0.98 | 0.655 | 5.19 |
 655
 656Task A is learned in 5 steps. The next 295 steps teach nothing new about A
 657but keep eroding the general skill, from 0.825 to 0.655. The last column
 658explains why: the weights keep travelling away from the base.
 659
 660$$
 661d_t = \lVert \theta_t - \theta_{\text{base}} \rVert = \sqrt{\sum_{i=1}^{P} \left(\theta_{t,i} - \theta_{\text{base},i}\right)^2}
 662$$
 663
 664**Symbols**
 665
 666| Symbol | Meaning here | In the example |
 667|---|---|---|
 668| $\theta_{\text{base}}$ | theta: every weight of the base model, as one vector | 97 numbers |
 669| $\theta_t$ | every weight after $t$ fine-tuning steps | |
 670| $P$ | the number of weights | 97 (3 in the hand example) |
 671| $i$ | a counter over the weights | 1 … P |
 672| $\theta_{t,i}$ | the $i$-th weight after $t$ steps | |
 673| $\lVert v \rVert$ | the length of vector $v$: square each entry, add, take the square root | |
 674| $d_t$ | how far fine-tuning has moved the model | 5.19 after 300 steps |
 675
 676**In words:** "the distance travelled is the length of the change in the
 677weights: square every weight's change, add them up, take the square root."
 678
 679**With the numbers:** for a three-weight model moving from (1, 0, 2) to
 680(1.5, −1, 2), the changes are (0.5, −1, 0), and
 681d = √(0.25 + 1 + 0) = √1.25 = 1.118.
 682
 683**In Python:**
 684
 685```python
 686import math
 687theta_base = [1.0, 0.0, 2.0]
 688theta_t = [1.5, -1.0, 2.0]
 689# each weight's change
 690[t - b for t, b in zip(theta_t, theta_base)]  # → [0.5, -1.0, 0.0]
 691# ‖θ_t − θ_base‖
 692round(math.sqrt(sum((t - b) ** 2 for t, b in zip(theta_t, theta_base))), 3)  # → 1.118
 693```
 694
 695```mermaid
 696flowchart LR
 697  DA[Task A examples] -->|gradient| W["Shared weights θ"]
 698  W --> SA[Answers on task A]
 699  W --> SG[Answers on the<br/>general skill]
 700  SG -.->|"no examples,<br/>no gradient"| W
 701```
 702
 703**Reading it:** the gradient from task A's examples moves the shared
 704weights, and both skills read those same weights. The dotted arrow is what
 705is missing: the general skill has no examples in this training run, so
 706nothing pushes back when a change that helps A hurts it. Forgetting is not
 707an event; it is the absence of a counterweight.
 708
 709![Left: accuracy on A jumps to 1 within a few steps while the general skill slides from 0.99 to 0.66 at learning rate 0.5 and only to 0.79 at 0.02. Right: the general skill falls as distance from the base grows](figures/primer.ml.fine_tuning.general_skill.svg)
 710
 711**Reading it:** on the left, the x-axis is the training step (log scale).
 712Blue lines are accuracy on task A, red lines the general skill; solid is
 713learning rate 0.5, dashed is 0.02. With the big rate, A is learned within a
 714few steps and the general skill then slides for the rest of the run. With
 715the small rate, A is learned later (around step 70) and the general skill
 716ends at 0.79 instead of 0.655, with A just as good. On the right, every
 717step of both runs is plotted as distance from the base against the general
 718skill: the points fall along one downward curve. **Forgetting tracks how
 719far the weights move.**
 720
 721**Mitigations that shorten the trip.** Three knobs limit the distance, and
 722the toy measures two of them:
 723
 724- **Fewer steps.** Stop once the held-out score on the new task stops
 725  improving. Here, stopping at step 10 keeps the general skill at 0.79
 726  instead of 0.655, with A at 0.995.
 727- **A lower learning rate.** Smaller steps travel less far for the same
 728  result: 0.79 instead of 0.655 after 300 steps.
 729- **A smaller update.** LoRA (`primer.ml.training_stages`, section 4)
 730  freezes the base weights and allows only a low-rank change. Biderman et
 731  al. measured this on real language models and summed it up in their
 732  title: LoRA learns less and forgets less.
 733
 734**In code:** `general_skill_run` fine-tunes on A and records, after every
 735step, accuracy on A, the general skill and the distance from the base.
 736
 737**Why it matters in practice.** A fine-tuned assistant that has become
 738worse at everything outside its narrow task is the most common
 739fine-tuning disappointment. Always score the general skills you care about
 740alongside the new task, and prefer the earliest checkpoint that has learned
 741the task.
 742
 743### 3b. A new task that contradicts an old one
 744
 745**Everyday picture.** Back to the driver. Practising "keep right" does not
 746merely add a skill; it pushes directly against "keep left", because both
 747use the same reflex.
 748
 749**Tiny worked example.** Take the model fine-tuned on A (98% on A), then
 750fine-tune it on B alone for 300 steps:
 751
 752| Second fine-tune | A before | A after | New task after |
 753|---|---|---|---|
 754| on B (contradicts A) | 0.98 | **0.025** | 1.00 |
 755| on C (separate inputs) | 0.98 | 0.97 | 0.98 |
 756
 757After B, the model does not merely forget A; it answers A's questions
 758backwards, because it learned "yes when x₂ < 0" everywhere. After C, A is
 759barely touched: C's inputs never flowed through the weights A relies on
 760most.
 761
 762$$
 763F = \text{acc}_{\text{old}}^{\text{before}} - \text{acc}_{\text{old}}^{\text{after}}
 764$$
 765
 766**Symbols**
 767
 768| Symbol | Meaning here | In the example |
 769|---|---|---|
 770| $\text{acc}_{\text{old}}^{\text{before}}$ | held-out accuracy on the old task before the new fine-tune | 0.98 |
 771| $\text{acc}_{\text{old}}^{\text{after}}$ | the same, after the new fine-tune | 0.025 |
 772| $F$ | forgetting: accuracy lost on the old task | 0.955 |
 773
 774**In words:** "forgetting is how much accuracy the old task lost."
 775
 776**With the numbers:** F = 0.98 − 0.025 = 0.955 after B, and
 7770.98 − 0.97 = 0.01 after C.
 778
 779**In Python:**
 780
 781```python
 782acc_before, acc_after_B, acc_after_C = 0.98, 0.025, 0.97
 783# F after the contradicting task B
 784round(acc_before - acc_after_B, 3)  # → 0.955
 785# F after task C, on separate inputs
 786round(acc_before - acc_after_C, 3)  # → 0.01
 787```
 788
 789```mermaid
 790flowchart LR
 791  BASE[Base] -->|fine-tune on A| MA["Model knows A<br/>A: 0.98"]
 792  MA -->|fine-tune on B only| MB["Model knows B<br/>A: 0.025, B: 1.00"]
 793  MA -->|fine-tune on C only| MC["Model knows A and C<br/>A: 0.97, C: 0.98"]
 794```
 795
 796**Reading it:** both second fine-tunes start from the same model. The only
 797difference is which weights the new task needs: B needs the very ones A
 798uses, in the opposite direction; C mostly needs others. How much a model
 799forgets depends less on how long you train than on how much the new task
 800overlaps and conflicts with the old one.
 801
 802![Fine-tuning on B after A: accuracy on A falls from 0.98 to near 0 within a few steps while B rises to 1; with 10 replayed A examples, A dips and then recovers to 0.945](figures/primer.ml.fine_tuning.sequential.svg)
 803
 804**Reading it:** the x-axis is the step of the second fine-tune (log scale),
 805the y-axis held-out accuracy. Solid lines are plain fine-tuning on B: B
 806(green) climbs to 1 while A (blue) falls to about half in the same few
 807steps, and A keeps sliding towards zero for as long as training continues.
 808Dashed lines add 10 replayed A examples (section 3c): A still dips
 809at first (to about 0.2 around step 40), then climbs back to 0.945 while B
 810reaches 0.99.
 811
 812**Does a lower learning rate help here?** Only in the sense of slowing the
 813slide. At learning rate 0.02, or stopping after 10 steps, B reaches 0.975
 814and A still drops to 0.355.
 815
 816![Accuracy on A against accuracy on B during the second fine-tune: every run on B alone traces the same curve whatever the learning rate and ends near zero on A, while the replay run climbs the right edge to the top-right corner](figures/primer.ml.fine_tuning.tradeoff.svg)
 817
 818**Reading it:** each line traces one run: x is accuracy on B, y is accuracy
 819on A, and each run starts at the top left (knows A, not B). The three
 820learning rates (0.5, 0.1, 0.02) take very different numbers of steps but
 821trace the same curve, close to the dotted line where A + B = 1: every point
 822of B they gain costs about a point of A, and once B is learned, further
 823steps slide A down the right edge towards zero. The replay run (dashed)
 824follows the same curve at first, then climbs the right edge and ends in
 825the top-right corner, knowing both. When the new data
 826contradicts the old skill, slowing down can't help, because nothing in B's
 827data says A still matters.
 828
 829**In code:** `sequential_run` starts from the A fine-tune, trains on a new
 830task (optionally with replay), and records both accuracies at every step;
 831`forgetting` computes F.
 832
 833**Why it matters in practice.** Real fine-tunes contradict the base model
 834more often than you'd think: "always answer in JSON" contradicts "chat
 835naturally"; "be terse" contradicts "explain in detail". Expect the old
 836behaviour to vanish wherever your data overrides it, and test for it.
 837
 838### 3c. Replay: keep practising the old skill
 839
 840**Everyday picture.** A pianist learning a new piece plays one old piece
 841at the start of every practice session. It costs a few minutes and keeps
 842the old repertoire alive.
 843
 844**Tiny worked example.** Add just 10 of task A's 200 training examples to
 845B's 200: under 5% of the mix. The result: A **0.945**, B 0.99 (up from A
 8460.025 without replay). Why can so few examples do so much? Look at the loss.
 847Suppose that early in training the model already scores B well (loss 0.05
 848per example) but has started forgetting A (loss 3.0 on each replayed
 849example):
 850
 851$$
 852\mathcal{L}_{\text{mix}} = \frac{n_{\text{new}}\,\mathcal{L}_{\text{new}} + n_{\text{old}}\,\mathcal{L}_{\text{old}}}{n_{\text{new}} + n_{\text{old}}}
 853$$
 854
 855**Symbols**
 856
 857| Symbol | Meaning here | In the example |
 858|---|---|---|
 859| $n_{\text{new}}$ | examples of the new task | 200 |
 860| $n_{\text{old}}$ | replayed examples of the old task | 10 |
 861| $\mathcal{L}_{\text{new}}$ | average loss on the new examples | 0.05 |
 862| $\mathcal{L}_{\text{old}}$ | average loss on the replayed examples | 3.0 |
 863| $\mathcal{L}_{\text{mix}}$ | the loss training actually minimises: the average over every example in the mix | 0.19 |
 864
 865**In words:** "the loss on a mixed dataset is the average over all its
 866examples, so each group counts in proportion to its size times its loss."
 867
 868**With the numbers:** (200 × 0.05 + 10 × 3.0) / 210 = (10 + 30) / 210 =
 8690.19. The 10 replayed examples are under 5% of the data but contribute
 8700.143 of the 0.19: three quarters of the loss, and so most of the gradient.
 871The old examples shout loudest exactly when they are being forgotten.
 872
 873**In Python:**
 874
 875```python
 876n_new, L_new = 200, 0.05
 877n_old, L_old = 10, 3.0
 878# the replayed share of the data
 879round(n_old / (n_new + n_old), 3)  # → 0.048
 880# each group's contribution to the average loss
 881round(n_new * L_new / 210, 3), round(n_old * L_old / 210, 3)  # → (0.048, 0.143)
 882# L_mix
 883round((n_new * L_new + n_old * L_old) / (n_new + n_old), 2)  # → 0.19
 884```
 885
 886```mermaid
 887flowchart LR
 888  NB["New task B<br/>200 examples"] --> MIX[Shuffle together<br/>210 examples]
 889  OA["Old task A<br/>10 kept examples"] --> MIX
 890  MIX --> FT[Fine-tune]
 891  FT --> BOTH["Knows B: 0.99<br/>and A: 0.945"]
 892```
 893
 894**Reading it:** the old task's small sample joins the new data before
 895training, so every gradient step sees both. This supplies exactly the
 896counterweight that was missing in section 3a's diagram: when a change that
 897helps B starts hurting A, the replayed examples' loss rises and pushes back.
 898
 899**In code:** `replay_mix` appends the old examples, `mixed_loss` is the
 900formula, and `sequential_run` with `n_replay=10` runs the experiment.
 901
 902**Why it matters in practice.** When fine-tuning a language model, mix some
 903general instruction-following data into your task data, so the model keeps
 904being a good assistant while it learns your task. When you can't replay
 905(the old data is gone or private), methods such as elastic weight
 906consolidation (Kirkpatrick et al.) instead penalise changes to the weights
 907the old task relied on most.
 908
 909## 4. Overfitting a small dataset
 910
 911**Everyday picture.** A student with only 16 flashcards, three of which
 912have the wrong answer on the back. For a while, studying teaches the
 913pattern. Keep drilling and the student memorises every card word for word,
 914including the three wrong answers, and gets worse on new questions.
 915`primer.ml.regularization` builds this idea from scratch; here it is in a
 916fine-tune.
 917
 918**Tiny worked example.** Fine-tune the base on only 16 examples of task A,
 9193 of them deliberately mislabelled, for 1,500 epochs. (With full-batch
 920training, one step is one pass over the data, one **epoch**.)
 921
 922| | At the best epoch (71) | At the end (1,500 epochs) |
 923|---|---|---|
 924| training loss | 0.373 | 0.005 |
 925| validation loss | **0.319** | 1.277 |
 926
 927At the best epoch, training loss is *higher* than validation loss: the
 928model is refusing to fit the three wrong labels, which is exactly right. By
 929the end, training loss is nearly zero, so the wrong labels have been
 930memorised, and validation loss has quadrupled.
 931
 932$$
 933g(t) = \mathcal{L}_{\text{val}}(t) - \mathcal{L}_{\text{train}}(t)
 934$$
 935
 936**Symbols**
 937
 938| Symbol | Meaning here | In the example |
 939|---|---|---|
 940| $t$ | the epoch | 71, then 1,500 |
 941| $\mathcal{L}_{\text{train}}(t)$ | average loss on the 16 training examples | 0.373, then 0.005 |
 942| $\mathcal{L}_{\text{val}}(t)$ | average loss on 200 held-out examples | 0.319, then 1.277 |
 943| $g(t)$ | the **generalisation gap**: how much worse the model does on data it hasn't seen | −0.054, then 1.272 |
 944
 945**In words:** "the gap is held-out loss minus training loss; a gap that
 946keeps growing means the model is memorising rather than learning."
 947
 948**With the numbers:** 0.319 − 0.373 = −0.054 at epoch 71; 1.277 − 0.005 =
 9491.272 at the end.
 950
 951**In Python:**
 952
 953```python
 954L_train = {"best": 0.373, "end": 0.005}
 955L_val = {"best": 0.319, "end": 1.277}
 956# g = L_val − L_train at each point
 957{t: round(L_val[t] - L_train[t], 3) for t in L_val}  # → {'best': -0.054, 'end': 1.272}
 958```
 959
 960![Training loss falls steadily to near zero while validation loss bottoms out at epoch 71 and then climbs to four times its best](figures/primer.ml.fine_tuning.overfitting.svg)
 961
 962**Reading it:** the x-axis is the epoch (log scale), the y-axis the loss.
 963Both curves fall at first while the model learns the real rule. At epoch 71
 964(the dashed line) validation loss bottoms out; after that the training
 965curve keeps falling as the model memorises the three wrong labels, and the
 966validation curve climbs. The dotted line is where early stopping with a
 967patience of 20 epochs would end the run, keeping the weights from epoch 71.
 968
 969```mermaid
 970flowchart LR
 971  EP[Train one epoch] --> SV[Save a checkpoint]
 972  SV --> SC[Score it on the<br/>held-out set]
 973  SC --> Q{Best so far?}
 974  Q -->|yes| MARK[Mark it best] --> EP
 975  Q -->|"no, patience used up"| SHIP[Ship the best checkpoint,<br/>not the last]
 976  Q -->|"no, patience left"| EP
 977```
 978
 979**Reading it:** every epoch ends with a checkpoint and a held-out score.
 980The loop keeps going while scores improve or patience remains, and when it
 981stops, the checkpoint that ships is the marked best, not whatever the last
 982epoch left behind.
 983
 984**In code:** `overfitting_run` fine-tunes on the 16 examples and records
 985both losses every epoch; `primer.ml.regularization.early_stopping` replays
 986the validation curve and returns the best epoch and the stopping epoch.
 987
 988**Why it matters in practice.** Fine-tuning datasets are small compared to
 989pretraining, and big models memorise quickly, so fine-tunes typically run
 990for only a few epochs. Save checkpoints, score each on the held-out set,
 991and ship the best one.
 992
 993## 5. Model merging
 994
 995### 5a. Weight averaging and task arithmetic
 996
 997**Everyday picture.** Two editors each take a copy of the same draft and
 998make tracked changes: one fixes the grammar, the other tightens the
 999argument. You can apply both sets of changes to the original. If instead
1000you "average" the two edited copies, each change is applied at half
1001strength: half the grammar fixed, half the argument tightened.
1002
1003**Tiny worked example.** A three-weight base (1, 0, 2). One fine-tune
1004moves it to (1.5, 0, 2); another to (1, −1, 2).
1005
1006| | Weights | Change from the base |
1007|---|---|---|
1008| base | (1, 0, 2) | |
1009| fine-tune on A | (1.5, 0, 2) | τ_A = (0.5, 0, 0) |
1010| fine-tune on C | (1, −1, 2) | τ_C = (0, −1, 0) |
1011| base + τ_A + τ_C | **(1.5, −1, 2)** | both changes in full |
1012| average of the two fine-tunes | (1.25, −0.5, 2) | both changes at half strength |
1013
1014The change a fine-tune made, θ_ft − θ_base, is its **task vector**. Adding
1015task vectors to the base is **task arithmetic**.
1016
1017$$
1018\tau_t = \theta_t - \theta_{\text{base}}, \qquad
1019\theta_{\text{merged}} = \theta_{\text{base}} + \lambda \sum_{t=1}^{T} \tau_t
1020$$
1021
1022**Symbols**
1023
1024| Symbol | Meaning here | In the example |
1025|---|---|---|
1026| $t$ | which task: a counter over the fine-tunes | A, C |
1027| $T$ | how many fine-tunes are merged | 2 |
1028| $\theta_t$ | all the weights of the model fine-tuned on task $t$ | (1.5, 0, 2) |
1029| $\theta_{\text{base}}$ | the weights they all started from | (1, 0, 2) |
1030| $\tau_t$ | tau: task $t$'s task vector, everything its fine-tune changed | (0.5, 0, 0) |
1031| $\sum_{t=1}^{T}$ | add up the task vectors | τ_A + τ_C = (0.5, −1, 0) |
1032| $\lambda$ | lambda: how strongly to apply the combined changes | 1 |
1033| $\theta_{\text{merged}}$ | the merged model's weights | (1.5, −1, 2) |
1034
1035**In words:** "each task vector is what its fine-tune changed; add the
1036changes up, scale them by λ, and apply them to the base."
1037
1038**With the numbers:** θ_merged = (1, 0, 2) + 1 × (0.5, −1, 0) = (1.5, −1, 2).
1039With λ = 1/2: (1, 0, 2) + 0.5 × (0.5, −1, 0) = (1.25, −0.5, 2), which is
1040exactly the average of the two fine-tunes. **Averaging T fine-tunes is task
1041arithmetic with λ = 1/T.**
1042
1043**In Python:**
1044
1045```python
1046theta_base = [1.0, 0.0, 2.0]
1047theta_A = [1.5, 0.0, 2.0]
1048theta_C = [1.0, -1.0, 2.0]
1049# τ_t = θ_t − θ_base
1050tau_A = [a - b for a, b in zip(theta_A, theta_base)]
1051tau_C = [c - b for c, b in zip(theta_C, theta_base)]
1052tau_A, tau_C  # → ([0.5, 0.0, 0.0], [0.0, -1.0, 0.0])
1053# θ_merged at λ = 1
1054[b + 1.0 * (a + c) for b, a, c in zip(theta_base, tau_A, tau_C)]  # → [1.5, -1.0, 2.0]
1055# λ = 1/2 ...
1056[b + 0.5 * (a + c) for b, a, c in zip(theta_base, tau_A, tau_C)]  # → [1.25, -0.5, 2.0]
1057# ... is the plain average of the two fine-tunes
1058[(a + c) / 2 for a, c in zip(theta_A, theta_C)]  # → [1.25, -0.5, 2.0]
1059```
1060
1061```mermaid
1062flowchart LR
1063  B[Base θ] -->|fine-tune on A| FA[θ_A]
1064  B -->|fine-tune on C| FC[θ_C]
1065  FA --> TA["τ_A = θ_A − θ_base"]
1066  FC --> TC["τ_C = θ_C − θ_base"]
1067  TA --> SUM["λ × (τ_A + τ_C)"]
1068  TC --> SUM
1069  B --> ADD((+))
1070  SUM --> ADD
1071  ADD --> M[Merged model:<br/>no extra training]
1072```
1073
1074**Reading it:** both fine-tunes start from the same base; that shared start
1075is what makes their changes comparable. Subtracting the base turns each
1076fine-tuned model into a task vector, the arrows are added and scaled, and
1077the result is added back to the base. No data and no gradient steps are
1078involved: merging is arithmetic on weights, and the merged model is the same
1079size and speed as the base.
1080
1081![Merging the fine-tunes on A and C: at lambda 1 the merged model scores 0.97 on both, while plain averaging (lambda one half) scores only 0.635 on A](figures/primer.ml.fine_tuning.merge_separate.svg)
1082
1083**Reading it:** the x-axis is λ and the y-axis held-out accuracy of the
1084merged model on A (blue) and on C (green); the dotted line is 0.5, a coin
1085flip. At λ = 1/2, which is plain averaging, each skill is diluted: 0.635 on
1086A. Around λ = 1 the merged model does both tasks at 0.97, as well as either
1087specialist on its own task. Push λ further and the changes overshoot. Their
1088task vectors are nearly perpendicular (cosine −0.04), so adding one barely
1089disturbs the other.
1090
1091**In code:** `task_vector` subtracts the base, `merge` adds scaled task
1092vectors back, and `merge_run` merges two fine-tunes at several λ and scores
1093the result.
1094
1095**Why it matters in practice.** Merging combines skills trained separately,
1096by different teams or on data that can't be pooled, without any further
1097training and at no extra inference cost. Averaging several fine-tunes of the
1098same task ("model soups", Wortsman et al.) often beats the best single one.
1099Task vectors can also be *subtracted*: Ilharco et al. negated a task vector
1100learned from toxic text to make a model less toxic.
1101
1102### 5b. Interference: when task vectors collide
1103
1104**Everyday picture.** Two editors rewrote the *same* sentence in opposite
1105directions, one making it warmer and one making it colder. Applying both
1106sets of tracked changes gives a sentence neither intended.
1107
1108**Tiny worked example.** Our three-weight changes again, plus a new one:
1109τ_A = (0.5, 0, 0), τ_C = (0, −1, 0) and τ_B = (−0.4, 0, 0.3). τ_A and τ_C
1110touch different weights: no conflict. τ_B pulls the first weight the other
1111way from τ_A. The **cosine** measures how aligned two changes are:
1112
1113$$
1114\cos(\tau_A, \tau_B) = \frac{\tau_A \cdot \tau_B}{\lVert \tau_A \rVert \, \lVert \tau_B \rVert}
1115$$
1116
1117**Symbols**
1118
1119| Symbol | Meaning here | In the example |
1120|---|---|---|
1121| $\tau_A \cdot \tau_B$ | dot product: multiply matching entries, then add | 0.5 × (−0.4) = −0.2 |
1122| $\lVert \tau \rVert$ | a vector's length | 0.5 and 0.5 |
1123| $\cos$ | the cosine of the angle between the two changes: 1 same direction, 0 unrelated, −1 opposite | −0.8 |
1124
1125**In words:** "multiply the changes weight by weight and add, then divide by
1126both lengths, so only the direction counts."
1127
1128**With the numbers:** τ_A · τ_B = 0.5 × (−0.4) + 0 + 0 = −0.2; ‖τ_A‖ = 0.5,
1129‖τ_B‖ = √(0.16 + 0.09) = 0.5; cos = −0.2 / 0.25 = **−0.8**: strongly
1130opposed. cos(τ_A, τ_C) = 0: independent. See `primer.ml.embeddings.similarity`
1131for the cosine from scratch.
1132
1133**In Python:**
1134
1135```python
1136import math
1137tau_A = [0.5, 0.0, 0.0]
1138tau_B = [-0.4, 0.0, 0.3]
1139tau_C = [0.0, -1.0, 0.0]
1140def cos(u, v):
1141    dot = sum(a * b for a, b in zip(u, v))
1142    return dot / (math.sqrt(sum(a * a for a in u)) * math.sqrt(sum(b * b for b in v)))
1143# opposed changes
1144round(cos(tau_A, tau_B), 2)  # → -0.8
1145# changes to different weights
1146cos(tau_A, tau_C)  # → 0.0
1147```
1148
1149In the toy, the fine-tunes on A and on B (which reverses A's rule) have task
1150vectors with cosine **−0.25**, against −0.04 for A and C. Each fine-tune
1151learned its rule everywhere, not just in its own region, so the two vectors
1152rewrite the same weights in opposite directions.
1153
1154![Merging the fine-tunes on A and B: at every lambda at least one task stays at or below a coin flip, and the best the merge manages on both at once is 0.525](figures/primer.ml.fine_tuning.merge_conflict.svg)
1155
1156**Reading it:** the same axes as the previous figure, now merging A with B.
1157No value of λ lifts both lines: wherever one task improves, the other
1158sits near or below a coin flip, and from λ = 1 on both stay there, because
1159the two changes cancel. A
1160model *can* do both tasks (replay in section 3c got 0.945 and 0.99), but
1161this merge can't reach it: that model needs to tell the regions apart, and
1162neither fine-tune learned to.
1163
1164```mermaid
1165flowchart TD
1166  TV[Task vectors from<br/>the same base] --> COS{Cosine between them}
1167  COS -->|"near 0: separate weights"| ADD[Add them:<br/>task arithmetic]
1168  COS -->|"clearly negative: conflict"| FIX{Can you retrain?}
1169  FIX -->|yes| JOINT[Train one model on<br/>both datasets, or replay]
1170  FIX -->|no| TIES["Resolve conflicts:<br/>trim small changes,<br/>agree on a sign per weight"]
1171  ADD --> EV[Score the merge on<br/>every task's held-out set]
1172  TIES --> EV
1173  JOINT --> EV
1174```
1175
1176**Reading it:** a cheap check before merging is the cosine between task
1177vectors. Near zero, the changes live in different weights and adding them
1178usually works. Clearly negative, they fight: train on both datasets
1179together if you can; if you can't, conflict-resolving merges such as
1180TIES-merging (Yadav et al.) drop each task's smallest changes and, for each
1181weight, keep only the changes that agree with the majority sign. Every path
1182ends at the same place: score the merge on every task's held-out set.
1183
1184**In code:** `cosine` measures the angle between two task vectors, and
1185`merge_run` reports it alongside the merged accuracies.
1186
1187**Why it matters in practice.** Merges are free to try, which makes them
1188tempting to trust. A merged model can quietly lose a skill both parents
1189had, so it is evaluated like any new model. The cosine tells you in advance
1190which merges to be nervous about.
1191
1192## In 20 seconds
1193
1194- **Decide with an eval:** build a held-out set first; try prompting, then
1195  retrieval, and fine-tune only for behaviour a prompt can't pin down. It
1196  pays off after C / (c_prompt − c_tuned) requests.
1197- **Data:** use the model's chat template and train only on assistant turns;
1198  deduplicate (Jaccard on word shingles); freeze the held-out set before
1199  training and remove near-copies of it; audit labels, because wrong labels
1200  cap what the eval can show. Quality beats quantity.
1201- **Forgetting:** every weight is shared, so learning a new task erodes old
1202  ones, in proportion to how far the weights move and how much the tasks
1203  conflict. Fewer steps, a lower learning rate and LoRA shorten the trip;
1204  replaying a little old data is what keeps a contradicted skill.
1205- **Overfitting:** on small data, validation loss bottoms out early; ship
1206  the best checkpoint, not the last.
1207- **Merging:** a task vector is θ_ft − θ_base. Adding task vectors combines
1208  skills without training; averaging is λ = 1/T and dilutes them; vectors
1209  that point in opposite directions interfere.
1210
1211## Self-test questions
1212
1213**When is fine-tuning the wrong tool, and what should you try first?**
1214When the model lacks facts, or the facts change: retrieval supplies them and
1215is easy to update. When a clearer prompt with a few examples fixes the
1216behaviour: that is cheaper and survives a base-model upgrade. Fine-tuning
1217earns its cost when behaviour stays inconsistent under the best prompt, or
1218when a long prompt sent millions of times costs more than the fine-tune.
1219
1220**Why build the held-out set before training, and why check it against the training set?**
1221So that no choice (prompt, learning rate, checkpoint) is made by looking at
1222it, and it stays an honest measure. A training example that nearly copies a
1223held-out one lets the model recite the answer, turning the eval into a
1224memory test; deduplicating across the split prevents it.
1225
1226**A held-out set has 100 examples and the fine-tune scores 83% against the prompt's 80%. Has it won?**
1227Not yet. At 80% on 100 examples the 95% margin is about ±7.8 points, so a
12283-point difference is well inside the noise. You need a larger held-out set
1229(400 examples halve the margin) or a bigger difference.
1230
1231**If 10% of the eval's labels are wrong, what is the best score a perfect model can get?**
123290%, because it is marked wrong on every mislabelled item. A 90%-accurate
1233model would score 0.9 × 0.9 + 0.1 × 0.1 = 82%. Noisy labels shrink and blur
1234the differences you are trying to measure.
1235
1236**What is catastrophic forgetting, and why does it happen?**
1237Training on a new task alone erodes, or wipes out, skills the model had.
1238Every weight is shared between tasks, and only the new task's examples
1239produce gradients, so nothing pushes back when a change that helps the new
1240task hurts an old one. It grows with how far the weights move and with how
1241much the new task conflicts with the old.
1242
1243**Lowering the learning rate did not stop task A being forgotten. Why not, and what works?**
1244Task B contradicts A on the same inputs, so any progress on B costs A;
1245a lower rate only walks the same trade-off more slowly. Replay works: mixing
1246even 5% of A's examples into B's data gives the model a reason to keep A,
1247and those few examples carry most of the loss exactly when A is slipping.
1248
1249**Training loss keeps falling, but validation loss has risen since epoch 71. What is happening, and which checkpoint do you ship?**
1250The model has stopped learning the general rule and is memorising the
1251training set, including its mislabelled examples. Ship the checkpoint from
1252epoch 71, the best on the held-out set; early stopping automates exactly
1253this.
1254
1255**What is a task vector, and why is averaging two fine-tunes the same as task arithmetic with λ = 1/2?**
1256A task vector is everything a fine-tune changed: θ_ft − θ_base. The average
1257of two fine-tunes is (θ_base + τ_1 + θ_base + τ_2) / 2 = θ_base + ½(τ_1 + τ_2),
1258which is task arithmetic with λ = 1/2, so each skill arrives at half
1259strength.
1260
1261**When does merging fail, and how can you see it coming?**
1262When the task vectors change the same weights in opposite directions, so
1263adding them cancels both skills. A clearly negative cosine between task
1264vectors is the warning; the remedy is joint training or replay, or a
1265conflict-resolving merge such as TIES, and a held-out check on every task
1266either way.
1267
1268## The papers behind this lesson
1269
1270- **Ilharco et al., *Editing Models with Task Arithmetic* (2022)**:
1271  https://arxiv.org/abs/2212.04089. Defined task vectors as fine-tuned
1272  minus pretrained weights and showed that adding them combines skills,
1273  and negating them removes a behaviour.
1274  [Annotated companion](../../papers/task-arithmetic.html)
1275- **Wortsman et al., *Model soups: averaging weights of multiple fine-tuned
1276  models improves accuracy without increasing inference time* (2022)**:
1277  https://arxiv.org/abs/2203.05482. Showed that averaging the weights of
1278  several fine-tunes of one base often beats the best single fine-tune.
1279- **Yadav et al., *TIES-Merging: Resolving Interference When Merging
1280  Models* (2023)**: https://arxiv.org/abs/2306.01708. Traced failed merges
1281  to small redundant changes and sign conflicts, and fixed both by trimming
1282  and electing a sign per weight.
1283- **Kirkpatrick et al., *Overcoming catastrophic forgetting in neural
1284  networks* (2017)**: https://arxiv.org/abs/1612.00796. Introduced elastic
1285  weight consolidation, which slows learning on the weights most important
1286  to earlier tasks.
1287- **Biderman et al., *LoRA Learns Less and Forgets Less* (2024)**:
1288  https://arxiv.org/abs/2405.09673. Measured on real language models that
1289  LoRA keeps more of the base model's abilities than full fine-tuning, at
1290  the price of learning the new task less completely.
1291- **Zhou et al., *LIMA: Less Is More for Alignment* (2023)**:
1292  https://arxiv.org/abs/2305.11206. Fine-tuned a large base model on 1,000
1293  carefully curated examples and got a strong assistant, evidence that
1294  example quality matters more than quantity.
1295  [Annotated companion](../../papers/lima.html)
1296- **Lee et al., *Deduplicating Training Data Makes Language Models Better*
1297  (2021)**: https://arxiv.org/abs/2107.06499. Found widespread near-duplicates
1298  in standard datasets, including between training and test sets, and
1299  showed that removing them reduces memorisation.
1300
1301## Further reading
1302
1303- Goodfellow et al., *An Empirical Investigation of Catastrophic Forgetting in Gradient-Based Neural Networks* (2013): https://arxiv.org/abs/1312.6211
1304- Hu et al., *LoRA: Low-Rank Adaptation of Large Language Models* (2021): https://arxiv.org/abs/2106.09685
1305- Ilharco et al., *Editing Models with Task Arithmetic* (2022): https://arxiv.org/abs/2212.04089
1306- Yadav et al., *TIES-Merging* (2023): https://arxiv.org/abs/2306.01708
1307- Hugging Face TRL, supervised fine-tuning trainer: https://huggingface.co/docs/trl/sft_trainer
1308- Hugging Face PEFT, parameter-efficient fine-tuning (LoRA and friends): https://huggingface.co/docs/peft/index
1309- mergekit, an open-source toolkit for merging models: https://github.com/arcee-ai/mergekit
1310"""
1311
1312from __future__ import annotations
1313
1314import re
1315from functools import lru_cache
1316
1317import numpy as np
1318
1319from primer._show import banner, say, table, takeaway
1320
1321# ---------------------------------------------------------------------------
1322# 1. Should you fine-tune? The cost arithmetic
1323# ---------------------------------------------------------------------------
1324
1325
1326def per_request_cost(tokens: int, price_per_million: float) -> float:
1327    """Dollars for one request that sends `tokens` input tokens at `price_per_million` dollars per million."""
1328    return tokens * price_per_million / 1_000_000
1329
1330
1331def break_even_requests(fixed_cost: float, cost_prompted: float, cost_tuned: float) -> float:
1332    """How many requests before a fine-tune's one-off cost is paid back: C / (c_prompt - c_tuned).
1333
1334    If the tuned model is not cheaper per request, the saving never
1335    arrives, so the answer is infinity: fine-tune for quality, not cost.
1336    """
1337    saving = cost_prompted - cost_tuned
1338    return fixed_cost / saving if saving > 0 else float("inf")
1339
1340
1341# ---------------------------------------------------------------------------
1342# 2. Preparing data: chat format, deduplication, held-out set, label quality
1343# ---------------------------------------------------------------------------
1344
1345ROLES = ("system", "user", "assistant")
1346
1347
1348def chat_example(system: str, user: str, assistant: str) -> dict:
1349    """One training example in the common chat format: a list of role-tagged messages."""
1350    return {
1351        "messages": [
1352            {"role": "system", "content": system},
1353            {"role": "user", "content": user},
1354            {"role": "assistant", "content": assistant},
1355        ]
1356    }
1357
1358
1359def validate_chat(example: dict) -> list[str]:
1360    """Every problem that would make this example teach the wrong thing (empty list = fine)."""
1361    messages = example.get("messages", [])
1362    problems = []
1363    for i, m in enumerate(messages, 1):
1364        if m.get("role") not in ROLES:
1365            problems.append(f"turn {i} has unknown role {m.get('role')!r}")
1366        elif not str(m.get("content", "")).strip():
1367            problems.append(f"turn {i} ({m['role']}) is empty")
1368    # Training teaches the model to produce the final assistant turn; without one there is no target.
1369    if not messages or messages[-1].get("role") != "assistant":
1370        problems.append("the last turn must be the assistant's")
1371    return problems
1372
1373
1374def render_chat(example: dict) -> list[tuple[str, bool]]:
1375    """Flatten messages into the text the model sees, as (segment, trained?) pairs.
1376
1377    The role markers are illustrative; every model family has its own chat
1378    template, and you must use the one the base model was trained with.
1379    Only assistant segments are trained on: the loss mask of SFT (see
1380    `primer.ml.training_stages`).
1381    """
1382    return [(f"<|{m['role']}|>{m['content']}<|end|>", m["role"] == "assistant") for m in example["messages"]]
1383
1384
1385def normalize(text: str) -> str:
1386    """Lowercase, drop punctuation, collapse spaces: so trivial differences don't hide a duplicate."""
1387    return " ".join(re.sub(r"[^\w\s]", " ", text.lower()).split())
1388
1389
1390def shingles(text: str, k: int = 2) -> set[tuple[str, ...]]:
1391    """The set of k-word windows ("shingles") in the normalized text."""
1392    words = normalize(text).split()
1393    return {tuple(words[i : i + k]) for i in range(max(len(words) - k + 1, 1))}
1394
1395
1396def jaccard(a: str, b: str, k: int = 2) -> float:
1397    """|A ∩ B| / |A ∪ B| over word k-grams: 1 = same text, 0 = nothing shared."""
1398    sa, sb = shingles(a, k), shingles(b, k)
1399    return len(sa & sb) / len(sa | sb)
1400
1401
1402def deduplicate(texts: list[str], threshold: float = 0.7) -> list[int]:
1403    """Indices of the texts to keep: the first of every group of near-duplicates.
1404
1405    All pairs are compared here, which is fine for thousands of examples;
1406    at web scale MinHash estimates the same Jaccard without comparing every pair.
1407    """
1408    kept: list[int] = []
1409    for i, t in enumerate(texts):
1410        if all(jaccard(t, texts[j]) < threshold for j in kept):
1411            kept.append(i)
1412    return kept
1413
1414
1415def remove_leaks(train: list[str], held_out: list[str], threshold: float = 0.7) -> list[int]:
1416    """Indices of training texts that do NOT nearly copy any held-out text."""
1417    return [i for i, t in enumerate(train) if all(jaccard(t, e) < threshold for e in held_out)]
1418
1419
1420def split_before_training(texts: list[str], eval_fraction: float = 0.2, seed: int = 0) -> tuple[list[str], list[str]]:
1421    """Deduplicate, set aside a held-out set, then drop training texts that leak into it.
1422
1423    The held-out set is chosen first and frozen, before any training or
1424    prompt tuning, so no decision is ever made by looking at it.
1425    """
1426    unique = [texts[i] for i in deduplicate(texts)]
1427    order = np.random.default_rng(seed).permutation(len(unique))
1428    n_eval = round(eval_fraction * len(unique))
1429    held_out = [unique[i] for i in order[:n_eval]]
1430    train = [unique[i] for i in order[n_eval:]]
1431    return [train[i] for i in remove_leaks(train, held_out)], held_out
1432
1433
1434def margin_of_error(accuracy: float, n: int, z: float = 1.96) -> float:
1435    """Half-width of the 95% interval around an accuracy measured on n examples: z·sqrt(a(1-a)/n)."""
1436    return z * np.sqrt(accuracy * (1 - accuracy) / n)
1437
1438
1439def measured_accuracy(true_accuracy: float, label_error: float) -> float:
1440    """What a yes/no eval reports when a fraction `label_error` of its reference labels are wrong.
1441
1442    The model scores a point when it is right on a right label, or wrong on
1443    a wrong label (the two mistakes cancel): a(1 - ε) + (1 - a)ε.
1444    """
1445    a, e = true_accuracy, label_error
1446    return a * (1 - e) + (1 - a) * e
1447
1448
1449def label_agreement(labels_1: list, labels_2: list) -> float:
1450    """Share of items two labelers labelled the same way."""
1451    return float(np.mean([p == q for p, q in zip(labels_1, labels_2)]))
1452
1453
1454# ---------------------------------------------------------------------------
1455# 3. A tiny model to fine-tune: 4 inputs -> 16 tanh units -> 1 probability
1456# ---------------------------------------------------------------------------
1457
1458N_IN, N_HIDDEN = 4, 16
1459# Where each weight lives in the flat parameter vector θ (task arithmetic needs one vector).
1460_SHAPES = (("W1", (N_IN, N_HIDDEN)), ("b1", (N_HIDDEN,)), ("w2", (N_HIDDEN,)), ("b2", (1,)))
1461N_PARAMS = sum(int(np.prod(s)) for _, s in _SHAPES)
1462
1463
1464def _unpack(theta: np.ndarray) -> dict[str, np.ndarray]:
1465    out, i = {}, 0
1466    for name, shape in _SHAPES:
1467        n = int(np.prod(shape))
1468        out[name] = theta[i : i + n].reshape(shape)
1469        i += n
1470    return out
1471
1472
1473class TinyNet:
1474    """A two-layer network whose every weight lives in one flat vector `theta`.
1475
1476    Keeping θ flat is what makes fine-tuning arithmetic literal: a task
1477    vector is `theta_tuned - theta_base`, and merging is vector addition.
1478    """
1479
1480    def __init__(self, theta: np.ndarray):
1481        self.theta = np.asarray(theta, dtype=float)
1482
1483    @classmethod
1484    def random(cls, seed: int = 0) -> "TinyNet":
1485        rng = np.random.default_rng(seed)
1486        parts = {
1487            "W1": rng.normal(0, 1, (N_IN, N_HIDDEN)),
1488            "b1": np.zeros(N_HIDDEN),
1489            # 1/sqrt(fan-in) keeps the output's starting scale near 1 (see primer.ml.deep_nets).
1490            "w2": rng.normal(0, 1 / np.sqrt(N_HIDDEN), N_HIDDEN),
1491            "b2": np.zeros(1),
1492        }
1493        return cls(np.concatenate([parts[name].ravel() for name, _ in _SHAPES]))
1494
1495    def _forward(self, X: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
1496        p = _unpack(self.theta)
1497        h = np.tanh(X @ p["W1"] + p["b1"])  # (n, 16)
1498        z = h @ p["w2"] + p["b2"][0]  # (n,)
1499        return h, 1 / (1 + np.exp(-z))
1500
1501    def predict(self, X: np.ndarray) -> np.ndarray:
1502        """Probability of "yes" for each row of X (shape (n, 4))."""
1503        return self._forward(X)[1]
1504
1505    def accuracy(self, X: np.ndarray, y: np.ndarray) -> float:
1506        return float(np.mean((self.predict(X) > 0.5) == y))
1507
1508    def loss(self, X: np.ndarray, y: np.ndarray) -> float:
1509        """Average cross-entropy: -log of the probability given to the right answer."""
1510        q = np.clip(self.predict(X), 1e-12, 1 - 1e-12)
1511        return float(-np.mean(y * np.log(q) + (1 - y) * np.log(1 - q)))
1512
1513    def gradient(self, X: np.ndarray, y: np.ndarray) -> np.ndarray:
1514        """d loss / d θ by backpropagation, flattened like θ."""
1515        p = _unpack(self.theta)
1516        h, q = self._forward(X)
1517        dz = (q - y) / len(y)  # (n,): sigmoid + cross-entropy gives this simple error
1518        dh = np.outer(dz, p["w2"]) * (1 - h**2)  # (n, 16): back through tanh
1519        grads = {"W1": X.T @ dh, "b1": dh.sum(0), "w2": h.T @ dz, "b2": np.array([dz.sum()])}
1520        return np.concatenate([grads[name].ravel() for name, _ in _SHAPES])
1521
1522
1523def fine_tune(net: TinyNet, X: np.ndarray, y: np.ndarray, steps: int = 300, lr: float = 0.5, record=None):
1524    """Full-batch gradient descent from `net`, returning a NEW model (the start is untouched).
1525
1526    With full-batch training one step is one pass over the data: one epoch.
1527    If `record` is given, it is called on the model after every step and
1528    the list of its results is returned alongside the model.
1529    """
1530    theta = net.theta.copy()
1531    history = []
1532    for _ in range(steps):
1533        theta = theta - lr * TinyNet(theta).gradient(X, y)
1534        if record is not None:
1535            history.append(record(TinyNet(theta)))
1536    tuned = TinyNet(theta)
1537    return (tuned, history) if record is not None else tuned
1538
1539
1540# The tasks. Each example is four numbers; the answer is yes (1) or no (0).
1541TASKS = {
1542    "general": "the base model's broad skill: all four numbers anywhere in [-3, 3]; yes when x1 + x3 > 0",
1543    "A": "first pair only, left region (x1 in [-3, -1]); yes when x2 > 0",
1544    "B": "first pair only, right region (x1 in [1, 3]); yes when x2 < 0 (A's rule, reversed)",
1545    "C": "second pair only (x3, x4 in [-2, 2]); yes when x3 + x4 > 0",
1546}
1547TRAIN_SEEDS = {"A": 1, "B": 2, "C": 3, "general": 4}
1548EVAL_SEEDS = {name: seed + 10 for name, seed in TRAIN_SEEDS.items()}
1549
1550
1551def make_task(name: str, n: int = 200, seed: int | None = None) -> tuple[np.ndarray, np.ndarray]:
1552    """Examples (X of shape (n, 4), y of 0/1) for one of `TASKS`. Default seed = its training set."""
1553    rng = np.random.default_rng(TRAIN_SEEDS[name] if seed is None else seed)
1554    X = np.zeros((n, N_IN))
1555    s = rng.uniform(-2, 2, n)
1556    if name == "A":
1557        X[:, 0], X[:, 1], y = rng.uniform(-3, -1, n), s, s > 0
1558    elif name == "B":
1559        X[:, 0], X[:, 1], y = rng.uniform(1, 3, n), s, s < 0
1560    elif name == "C":
1561        X[:, 2], X[:, 3] = rng.uniform(-2, 2, n), s
1562        y = X[:, 2] + s > 0
1563    elif name == "general":
1564        X = rng.uniform(-3, 3, (n, N_IN))
1565        y = X[:, 0] + X[:, 2] > 0
1566    else:
1567        raise ValueError(f"unknown task {name!r}; choose from {list(TASKS)}")
1568    return X, y.astype(float)
1569
1570
1571@lru_cache(maxsize=None)
1572def _eval_set(name: str) -> tuple[np.ndarray, np.ndarray]:
1573    # Held-out examples, drawn with their own seed and never trained on.
1574    return make_task(name, seed=EVAL_SEEDS[name])
1575
1576
1577def eval_accuracy(net: TinyNet, name: str) -> float:
1578    """Accuracy on a task's held-out set."""
1579    return net.accuracy(*_eval_set(name))
1580
1581
1582@lru_cache(maxsize=1)
1583def base_model() -> TinyNet:
1584    """The "pretrained" base: a random network trained on the general skill."""
1585    return fine_tune(TinyNet.random(seed=0), *make_task("general"), steps=300, lr=0.5)
1586
1587
1588@lru_cache(maxsize=None)
1589def fine_tuned(name: str, steps: int = 300, lr: float = 0.5) -> TinyNet:
1590    """The base model fine-tuned on one task's training set."""
1591    return fine_tune(base_model(), *make_task(name), steps=steps, lr=lr)
1592
1593
1594# ---------------------------------------------------------------------------
1595# 4. Catastrophic forgetting and its mitigations
1596# ---------------------------------------------------------------------------
1597
1598
1599def forgetting(acc_before: float, acc_after: float) -> float:
1600    """How much accuracy on the old task was lost: F = acc_before - acc_after."""
1601    return acc_before - acc_after
1602
1603
1604@lru_cache(maxsize=None)
1605def general_skill_run(lr: float = 0.5, steps: int = 300) -> dict[str, list[float]]:
1606    """Fine-tune the base on task A, recording after every step: accuracy on A,
1607    accuracy on the base's general skill, and distance travelled from the base."""
1608    base = base_model()
1609
1610    def record(net: TinyNet):
1611        return eval_accuracy(net, "A"), eval_accuracy(net, "general"), float(np.linalg.norm(net.theta - base.theta))
1612
1613    _, history = fine_tune(base, *make_task("A"), steps=steps, lr=lr, record=record)
1614    task, general, distance = (list(col) for col in zip(*history))
1615    return {"task": task, "general": general, "distance": distance}
1616
1617
1618def replay_mix(new: tuple, old: tuple, n_old: int) -> tuple[np.ndarray, np.ndarray]:
1619    """The new task's examples plus the first `n_old` examples of an old task."""
1620    (X_new, y_new), (X_old, y_old) = new, old
1621    return np.concatenate([X_new, X_old[:n_old]]), np.concatenate([y_new, y_old[:n_old]])
1622
1623
1624def mixed_loss(loss_new: float, n_new: int, loss_old: float, n_old: int) -> float:
1625    """Average loss over a mixed dataset: (n_new·L_new + n_old·L_old) / (n_new + n_old)."""
1626    return (n_new * loss_new + n_old * loss_old) / (n_new + n_old)
1627
1628
1629@lru_cache(maxsize=None)
1630def sequential_run(new: str, lr: float = 0.5, steps: int = 300, n_replay: int = 0) -> dict:
1631    """Start from the task-A fine-tune, then fine-tune on `new` (optionally replaying A).
1632
1633    Returns accuracy on A before and after, accuracy on the new task after,
1634    and the (acc on A, acc on new) pair after every step.
1635    """
1636    start = fine_tuned("A")
1637    data = make_task(new)
1638    if n_replay:
1639        data = replay_mix(data, make_task("A"), n_replay)
1640    end, path = fine_tune(start, *data, steps=steps, lr=lr, record=lambda net: (eval_accuracy(net, "A"), eval_accuracy(net, new)))
1641    return {
1642        "A_before": eval_accuracy(start, "A"),
1643        "A_after": eval_accuracy(end, "A"),
1644        f"{new}_after": eval_accuracy(end, new),
1645        "path": path,
1646    }
1647
1648
1649# ---------------------------------------------------------------------------
1650# 5. Overfitting a small dataset
1651# ---------------------------------------------------------------------------
1652
1653
1654@lru_cache(maxsize=None)
1655def overfitting_run(n: int = 16, n_wrong: int = 3, epochs: int = 1500, lr: float = 0.5, seed: int = 2) -> dict:
1656    """Fine-tune on n task-A examples, `n_wrong` of them mislabelled, for many epochs.
1657
1658    Records training loss and held-out (validation) loss after every epoch.
1659    """
1660    X, y = make_task("A", n=n, seed=seed)
1661    wrong = np.random.default_rng(seed).choice(n, n_wrong, replace=False)
1662    y = y.copy()
1663    y[wrong] = 1 - y[wrong]  # real datasets have some bad labels; a big enough model memorises them
1664    X_val, y_val = _eval_set("A")
1665    _, history = fine_tune(base_model(), X, y, steps=epochs, lr=lr, record=lambda net: (net.loss(X, y), net.loss(X_val, y_val)))
1666    train, val = (list(col) for col in zip(*history))
1667    return {"train": train, "val": val, "best_epoch": int(np.argmin(val)) + 1}  # counted from 1, like the 1,500 epochs
1668
1669
1670# ---------------------------------------------------------------------------
1671# 6. Model merging: weight averaging and task arithmetic
1672# ---------------------------------------------------------------------------
1673
1674
1675def task_vector(tuned: TinyNet, base: TinyNet) -> np.ndarray:
1676    """τ = θ_tuned - θ_base: everything the fine-tune changed, as one vector."""
1677    return tuned.theta - base.theta
1678
1679
1680def merge(base: TinyNet, task_vectors: list[np.ndarray], lam: float = 1.0) -> TinyNet:
1681    """θ_base + λ·Σ τ. With λ = 1/T this is exactly the average of the T fine-tunes."""
1682    return TinyNet(base.theta + lam * np.sum(task_vectors, axis=0))
1683
1684
1685def cosine(a: np.ndarray, b: np.ndarray) -> float:
1686    """cos of the angle between two vectors: 1 same direction, 0 unrelated, -1 opposite."""
1687    return float(a @ b / (np.linalg.norm(a) * np.linalg.norm(b)))
1688
1689
1690LAMBDAS = (0.25, 0.5, 0.75, 1.0, 1.25, 1.5)
1691
1692
1693@lru_cache(maxsize=None)
1694def merge_run(first: str, second: str, lambdas: tuple = LAMBDAS) -> dict:
1695    """Merge the fine-tunes on two tasks by task arithmetic at several λ.
1696
1697    Returns the cosine between their task vectors and, for each λ, the
1698    merged model's accuracy on (first, second).
1699    """
1700    base = base_model()
1701    taus = [task_vector(fine_tuned(t), base) for t in (first, second)]
1702    by_lambda = {}
1703    for lam in lambdas:
1704        merged = merge(base, taus, lam)
1705        by_lambda[lam] = (eval_accuracy(merged, first), eval_accuracy(merged, second))
1706    return {"cosine": cosine(*taus), "by_lambda": by_lambda}
1707
1708
1709# ---------------------------------------------------------------------------
1710# 7. Figures (rendered into the HTML docs by `make figures`)
1711# ---------------------------------------------------------------------------
1712
1713
1714def figures() -> dict:
1715    """Plot this lesson's data. matplotlib is imported here, and only here,
1716    so the lesson itself needs nothing beyond NumPy."""
1717    import matplotlib
1718
1719    matplotlib.use("Agg")
1720    import matplotlib.pyplot as plt
1721
1722    BLUE, RED, GREEN, AMBER, MUTED = "#2563eb", "#dc2626", "#059669", "#d97706", "#9ca3af"
1723    figs = {}
1724
1725    # --- 1. Break-even: cumulative cost of prompting vs fine-tuning ---------
1726    c_prompt, c_tuned, fixed = per_request_cost(3000, 2.0), per_request_cost(300, 4.0), 600
1727    n = np.linspace(0, 250_000, 200)
1728    n_star = break_even_requests(fixed, c_prompt, c_tuned)
1729    fig, ax = plt.subplots(figsize=(6, 3.6))
1730    ax.plot(n, c_prompt * n, color=RED, label=r"long prompt: \$0.0060 per request")
1731    ax.plot(n, fixed + c_tuned * n, color=BLUE, label=r"fine-tuned: \$600 once + \$0.0012 per request")
1732    ax.axvline(n_star, color=MUTED, ls="--")
1733    ax.text(n_star * 1.03, 150, f"break-even\n{n_star:,.0f} requests", color="#4b5563")
1734    ax.set_xlabel("requests served")
1735    ax.set_ylabel("total cost so far ($)")
1736    ax.set_title("When does a fine-tune pay for itself?")
1737    ax.legend(frameon=False, loc="upper left")
1738    figs["break_even"] = fig
1739
1740    # --- 2. Margin of error vs held-out set size -----------------------------
1741    sizes = np.logspace(np.log10(25), np.log10(3200), 100)
1742    fig, ax = plt.subplots(figsize=(6, 3.4))
1743    ax.plot(sizes, 100 * margin_of_error(0.8, sizes), color=BLUE)
1744    for size in (100, 400, 1600):
1745        m = 100 * margin_of_error(0.8, size)
1746        ax.plot(size, m, "o", color=BLUE)
1747        ax.annotate(f"n = {size}: ±{m:.1f}", (size, m), textcoords="offset points", xytext=(6, 6))
1748    ax.set_xscale("log")
1749    ax.set_xlabel("held-out examples n (log scale)")
1750    ax.set_ylabel("95% margin (percentage points)")
1751    ax.set_title("How precise is an 80% score? Four times the data halves the margin")
1752    figs["eval_margin"] = fig
1753
1754    # --- 3. The general skill fades while task A is learned -----------------
1755    runs = {lr: general_skill_run(lr=lr, steps=300) for lr in (0.5, 0.02)}
1756    steps = np.arange(1, 301)
1757    fig, (a1, a2) = plt.subplots(1, 2, figsize=(9.5, 3.6))
1758    for lr, style in ((0.5, "-"), (0.02, "--")):
1759        a1.plot(steps, runs[lr]["task"], style, color=BLUE, label=f"task A, lr {lr}")
1760        a1.plot(steps, runs[lr]["general"], style, color=RED, label=f"general skill, lr {lr}")
1761        a2.plot(runs[lr]["distance"], runs[lr]["general"], style, color=RED, label=f"lr {lr}")
1762    a1.set_xscale("log")
1763    a1.set_xlabel("fine-tuning step (log scale)")
1764    a1.set_ylabel("held-out accuracy")
1765    a1.set_ylim(0.25, 1.02)
1766    a1.set_title("Learning A, forgetting the general skill")
1767    a1.legend(frameon=False, fontsize=8, loc="lower right")
1768    a2.set_xlabel("distance from the base  ‖θ − θ_base‖")
1769    a2.set_ylabel("general-skill accuracy")
1770    a2.set_title("Forgetting tracks how far the weights move")
1771    a2.legend(frameon=False)
1772    fig.tight_layout()
1773    figs["general_skill"] = fig
1774
1775    # --- 4. Sequential fine-tuning: A then B, with and without replay -------
1776    plain, replay = sequential_run("B"), sequential_run("B", n_replay=10)
1777    fig, ax = plt.subplots(figsize=(6, 3.6))
1778    for run, style, tag in ((plain, "-", "B only"), (replay, "--", "B + 10 replayed A")):
1779        acc_a, acc_b = zip(*run["path"])
1780        ax.plot(steps, acc_a, style, color=BLUE, label=f"accuracy on A ({tag})")
1781        ax.plot(steps, acc_b, style, color=GREEN, label=f"accuracy on B ({tag})")
1782    ax.set_xscale("log")
1783    ax.set_xlabel("step of the second fine-tune (log scale)")
1784    ax.set_ylabel("held-out accuracy")
1785    ax.set_title("Fine-tuning on B after A: A collapses unless replayed")
1786    ax.legend(frameon=False, fontsize=8, loc="center left", bbox_to_anchor=(0.3, 0.68))
1787    figs["sequential"] = fig
1788
1789    # --- 5. The trade-off: learning rate only walks the diagonal ------------
1790    fig, ax = plt.subplots(figsize=(5.2, 4.6))
1791    ax.plot([0, 1], [1, 0], color=MUTED, ls=":", label="A + B = 1")
1792    for lr, color in ((0.5, RED), (0.1, AMBER), (0.02, BLUE)):
1793        acc_a, acc_b = zip(*sequential_run("B", lr=lr)["path"])
1794        ax.plot(acc_b, acc_a, "-", color=color, label=f"B only, lr {lr}")
1795        ax.plot(acc_b[-1], acc_a[-1], "o", color=color)
1796    acc_a, acc_b = zip(*replay["path"])
1797    ax.plot(acc_b, acc_a, "--", color=GREEN, label="B + replay, lr 0.5")
1798    ax.plot(acc_b[-1], acc_a[-1], "o", color=GREEN)
1799    ax.set_xlabel("accuracy on the new task B")
1800    ax.set_ylabel("accuracy on the old task A")
1801    ax.set_xlim(0, 1.03)
1802    ax.set_ylim(0, 1.03)
1803    ax.set_title("Slower is not safer; replay escapes the trade-off")
1804    ax.legend(frameon=False, fontsize=8, loc="lower left")
1805    figs["tradeoff"] = fig
1806
1807    # --- 6. Overfitting 16 examples --------------------------------------------
1808    from primer.ml.regularization import early_stopping
1809
1810    run = overfitting_run()
1811    epochs = np.arange(1, len(run["train"]) + 1)
1812    best, stop = early_stopping(run["val"], patience=20)
1813    fig, ax = plt.subplots(figsize=(6, 3.6))
1814    ax.plot(epochs, run["train"], color=BLUE, label="training loss (16 examples, 3 mislabelled)")
1815    ax.plot(epochs, run["val"], color=RED, label="validation loss (200 held-out examples)")
1816    ax.axvline(best + 1, color=MUTED, ls="--")
1817    ax.axvline(stop + 1, color=MUTED, ls=":")
1818    ax.text((best + 1) * 1.08, 1.9, f"best epoch {best + 1}", color="#4b5563", zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1))
1819    ax.set_xscale("log")
1820    ax.set_xlabel("epoch (log scale)")
1821    ax.set_ylabel("loss")
1822    ax.set_title("A small dataset: learning, then memorising")
1823    ax.legend(frameon=False, loc="upper right", fontsize=8)
1824    figs["overfitting"] = fig
1825
1826    # --- 7 and 8. Merging by task arithmetic, a λ sweep -----------------------
1827    lams = tuple(float(x) for x in np.round(np.arange(0, 1.51, 0.125), 3))
1828    for key, second, color in (("merge_separate", "C", GREEN), ("merge_conflict", "B", AMBER)):
1829        merged = merge_run("A", second, lambdas=lams)
1830        acc_first, acc_second = zip(*merged["by_lambda"].values())
1831        fig, ax = plt.subplots(figsize=(6, 3.4))
1832        ax.plot(lams, acc_first, "o-", color=BLUE, label="accuracy on A")
1833        ax.plot(lams, acc_second, "o-", color=color, label=f"accuracy on {second}")
1834        ax.axhline(0.5, color=MUTED, ls=":")
1835        ax.axvline(0.5, color=MUTED, ls="--")
1836        ax.text(0.52, 0.2, "plain average\n(λ = 1/2)", color="#4b5563")
1837        ax.set_ylim(0, 1.05)
1838        ax.set_xlabel("λ: strength of the summed task vectors")
1839        ax.set_ylabel("held-out accuracy of the merge")
1840        ax.set_title(f"base + λ(τ_A + τ_{second}):  cosine(τ_A, τ_{second}) = {merged['cosine']:.2f}")
1841        ax.legend(frameon=False, loc="lower right")
1842        figs[key] = fig
1843
1844    return figs
1845
1846
1847# ---------------------------------------------------------------------------
1848# 8. Narrated walkthrough
1849# ---------------------------------------------------------------------------
1850
1851
1852def demo() -> None:
1853    banner("1. Should you fine-tune? Put numbers on it")
1854    c_prompt, c_tuned = per_request_cost(3000, 2.0), per_request_cost(300, 4.0)
1855    table(
1856        ["option", "tokens", "$ per million", "$ per request"],
1857        [("long prompt", 3000, 2.0, c_prompt), ("fine-tuned", 300, 4.0, c_tuned)],
1858    )
1859    n_star = break_even_requests(600, c_prompt, c_tuned)
1860    say(f"A $600 fine-tune saves ${c_prompt - c_tuned:.4f} per request, so it pays off after {n_star:,.0f} requests.")
1861    takeaway("Build the eval first; prompt, then retrieve, and fine-tune only for behaviour a prompt can't pin down.")
1862
1863    banner("2. Preparing the data")
1864    example = chat_example("Be brief.", "Capital of France?", "Paris.")
1865    table(["segment", "trained on?"], [(seg, "yes" if trained else "no") for seg, trained in render_chat(example)])
1866    say(f"Problems with an example that ends on the user's turn: {validate_chat({'messages': [{'role': 'user', 'content': 'Hi'}]})}")
1867    questions = ["How do I reset my password?", "how do I reset my password, please", "How do I change my email?"]
1868    table(
1869        ["pair", "Jaccard on word pairs"],
1870        [("password vs. its rewording", jaccard(questions[0], questions[1])), ("password vs. email", jaccard(questions[0], questions[2]))],
1871        floatfmt=".2f",
1872    )
1873    say(f"Deduplicating at 0.7 keeps indices {deduplicate(questions)}: the rewording is dropped.")
1874    table(
1875        ["held-out examples", "95% margin at 80%"],
1876        [(size, f"±{100 * margin_of_error(0.8, size):.1f} points") for size in (25, 100, 400, 1600)],
1877    )
1878    say(
1879        f"""
1880        With 10% wrong labels, a 90%-accurate model measures
1881        {measured_accuracy(0.9, 0.1):.0%}, and a perfect one only
1882        {measured_accuracy(1.0, 0.1):.0%}.
1883        """
1884    )
1885    takeaway("Freeze the held-out set first, keep its near-copies out of training, and check the labels.")
1886
1887    banner("3. Catastrophic forgetting")
1888    say(f"The base model knows its general skill: {eval_accuracy(base_model(), 'general'):.2f} on held-out examples.")
1889    run = general_skill_run(lr=0.5, steps=300)
1890    slow = general_skill_run(lr=0.02, steps=300)
1891    table(
1892        ["fine-tune on A", "accuracy on A", "general skill", "distance from base"],
1893        [
1894            ("lr 0.5, after 5 steps", run["task"][4], run["general"][4], run["distance"][4]),
1895            ("lr 0.5, after 10 steps", run["task"][9], run["general"][9], run["distance"][9]),
1896            ("lr 0.5, after 300 steps", run["task"][-1], run["general"][-1], run["distance"][-1]),
1897            ("lr 0.02, after 300 steps", slow["task"][-1], slow["general"][-1], slow["distance"][-1]),
1898        ],
1899        floatfmt=".3f",
1900    )
1901    say("Task A is learned within a few steps; every step after that only moves the weights further and erodes the general skill.")
1902    rows = []
1903    for label, kwargs in (
1904        ("then B (contradicts A)", dict(new="B")),
1905        ("then B at lr 0.02", dict(new="B", lr=0.02)),
1906        ("then B, 10 steps only", dict(new="B", steps=10)),
1907        ("then B + 10 replayed A", dict(new="B", n_replay=10)),
1908        ("then C (separate inputs)", dict(new="C")),
1909    ):
1910        r = sequential_run(**kwargs)
1911        rows.append((label, r["A_before"], r["A_after"], r[f"{kwargs['new']}_after"], forgetting(r["A_before"], r["A_after"])))
1912    table(["after fine-tuning on A", "A before", "A after", "new task", "forgetting F"], rows, floatfmt=".3f")
1913    takeaway(
1914        "Forgetting grows with how far the weights move and how much the new task conflicts. "
1915        "Slowing down only walks the trade-off; replaying a little old data keeps both."
1916    )
1917
1918    banner("4. Overfitting a small dataset")
1919    over = overfitting_run()
1920    best = over["best_epoch"]  # counted from 1; lists are indexed from 0
1921    table(
1922        ["", "training loss", "validation loss"],
1923        [(f"best epoch ({best})", over["train"][best - 1], over["val"][best - 1]), ("last epoch", over["train"][-1], over["val"][-1])],
1924        floatfmt=".3f",
1925    )
1926    takeaway("On 16 examples the model learns the rule, then memorises the 3 wrong labels. Ship the best checkpoint.")
1927
1928    banner("5. Merging models by task arithmetic")
1929    for second in ("C", "B"):
1930        merged = merge_run("A", second)
1931        say(f"A + {second}: cosine between task vectors = {merged['cosine']:.2f}")
1932        table(
1933            ["λ", "accuracy on A", f"accuracy on {second}"],
1934            [(lam, acc[0], acc[1]) for lam, acc in merged["by_lambda"].items()],
1935            floatfmt=".3f",
1936        )
1937    takeaway(
1938        "Nearly perpendicular task vectors add into a model that does both (0.97 and 0.97 at λ = 1); "
1939        "opposed ones cancel. Averaging is λ = 1/2 and dilutes each skill."
1940    )
1941
1942
1943if __name__ == "__main__":
1944    demo()
Level 3: the code, function by function.
def per_request_cost(tokens: int, price_per_million: float) -> float: on GitHub
1327def per_request_cost(tokens: int, price_per_million: float) -> float:
1328    """Dollars for one request that sends `tokens` input tokens at `price_per_million` dollars per million."""
1329    return tokens * price_per_million / 1_000_000

Dollars for one request that sends tokens input tokens at price_per_million dollars per million.

def break_even_requests(fixed_cost: float, cost_prompted: float, cost_tuned: float) -> float: on GitHub
1332def break_even_requests(fixed_cost: float, cost_prompted: float, cost_tuned: float) -> float:
1333    """How many requests before a fine-tune's one-off cost is paid back: C / (c_prompt - c_tuned).
1334
1335    If the tuned model is not cheaper per request, the saving never
1336    arrives, so the answer is infinity: fine-tune for quality, not cost.
1337    """
1338    saving = cost_prompted - cost_tuned
1339    return fixed_cost / saving if saving > 0 else float("inf")

How many requests before a fine-tune's one-off cost is paid back: C / (c_prompt - c_tuned).

If the tuned model is not cheaper per request, the saving never arrives, so the answer is infinity: fine-tune for quality, not cost.

ROLES = ('system', 'user', 'assistant')
def chat_example(system: str, user: str, assistant: str) -> dict: on GitHub
1349def chat_example(system: str, user: str, assistant: str) -> dict:
1350    """One training example in the common chat format: a list of role-tagged messages."""
1351    return {
1352        "messages": [
1353            {"role": "system", "content": system},
1354            {"role": "user", "content": user},
1355            {"role": "assistant", "content": assistant},
1356        ]
1357    }

One training example in the common chat format: a list of role-tagged messages.

def validate_chat(example: dict) -> list[str]: on GitHub
1360def validate_chat(example: dict) -> list[str]:
1361    """Every problem that would make this example teach the wrong thing (empty list = fine)."""
1362    messages = example.get("messages", [])
1363    problems = []
1364    for i, m in enumerate(messages, 1):
1365        if m.get("role") not in ROLES:
1366            problems.append(f"turn {i} has unknown role {m.get('role')!r}")
1367        elif not str(m.get("content", "")).strip():
1368            problems.append(f"turn {i} ({m['role']}) is empty")
1369    # Training teaches the model to produce the final assistant turn; without one there is no target.
1370    if not messages or messages[-1].get("role") != "assistant":
1371        problems.append("the last turn must be the assistant's")
1372    return problems

Every problem that would make this example teach the wrong thing (empty list = fine).

def render_chat(example: dict) -> list[tuple[str, bool]]: on GitHub
1375def render_chat(example: dict) -> list[tuple[str, bool]]:
1376    """Flatten messages into the text the model sees, as (segment, trained?) pairs.
1377
1378    The role markers are illustrative; every model family has its own chat
1379    template, and you must use the one the base model was trained with.
1380    Only assistant segments are trained on: the loss mask of SFT (see
1381    `primer.ml.training_stages`).
1382    """
1383    return [(f"<|{m['role']}|>{m['content']}<|end|>", m["role"] == "assistant") for m in example["messages"]]

Flatten messages into the text the model sees, as (segment, trained?) pairs.

The role markers are illustrative; every model family has its own chat template, and you must use the one the base model was trained with. Only assistant segments are trained on: the loss mask of SFT (see primer.ml.training_stages).

def normalize(text: str) -> str: on GitHub
1386def normalize(text: str) -> str:
1387    """Lowercase, drop punctuation, collapse spaces: so trivial differences don't hide a duplicate."""
1388    return " ".join(re.sub(r"[^\w\s]", " ", text.lower()).split())

Lowercase, drop punctuation, collapse spaces: so trivial differences don't hide a duplicate.

def shingles(text: str, k: int = 2) -> set[tuple[str, ...]]: on GitHub
1391def shingles(text: str, k: int = 2) -> set[tuple[str, ...]]:
1392    """The set of k-word windows ("shingles") in the normalized text."""
1393    words = normalize(text).split()
1394    return {tuple(words[i : i + k]) for i in range(max(len(words) - k + 1, 1))}

The set of k-word windows ("shingles") in the normalized text.

def jaccard(a: str, b: str, k: int = 2) -> float: on GitHub
1397def jaccard(a: str, b: str, k: int = 2) -> float:
1398    """|A ∩ B| / |A ∪ B| over word k-grams: 1 = same text, 0 = nothing shared."""
1399    sa, sb = shingles(a, k), shingles(b, k)
1400    return len(sa & sb) / len(sa | sb)

|A ∩ B| / |A ∪ B| over word k-grams: 1 = same text, 0 = nothing shared.

def deduplicate(texts: list[str], threshold: float = 0.7) -> list[int]: on GitHub
1403def deduplicate(texts: list[str], threshold: float = 0.7) -> list[int]:
1404    """Indices of the texts to keep: the first of every group of near-duplicates.
1405
1406    All pairs are compared here, which is fine for thousands of examples;
1407    at web scale MinHash estimates the same Jaccard without comparing every pair.
1408    """
1409    kept: list[int] = []
1410    for i, t in enumerate(texts):
1411        if all(jaccard(t, texts[j]) < threshold for j in kept):
1412            kept.append(i)
1413    return kept

Indices of the texts to keep: the first of every group of near-duplicates.

All pairs are compared here, which is fine for thousands of examples; at web scale MinHash estimates the same Jaccard without comparing every pair.

def remove_leaks( train: list[str], held_out: list[str], threshold: float = 0.7) -> list[int]: on GitHub
1416def remove_leaks(train: list[str], held_out: list[str], threshold: float = 0.7) -> list[int]:
1417    """Indices of training texts that do NOT nearly copy any held-out text."""
1418    return [i for i, t in enumerate(train) if all(jaccard(t, e) < threshold for e in held_out)]

Indices of training texts that do NOT nearly copy any held-out text.

def split_before_training( texts: list[str], eval_fraction: float = 0.2, seed: int = 0) -> tuple[list[str], list[str]]: on GitHub
1421def split_before_training(texts: list[str], eval_fraction: float = 0.2, seed: int = 0) -> tuple[list[str], list[str]]:
1422    """Deduplicate, set aside a held-out set, then drop training texts that leak into it.
1423
1424    The held-out set is chosen first and frozen, before any training or
1425    prompt tuning, so no decision is ever made by looking at it.
1426    """
1427    unique = [texts[i] for i in deduplicate(texts)]
1428    order = np.random.default_rng(seed).permutation(len(unique))
1429    n_eval = round(eval_fraction * len(unique))
1430    held_out = [unique[i] for i in order[:n_eval]]
1431    train = [unique[i] for i in order[n_eval:]]
1432    return [train[i] for i in remove_leaks(train, held_out)], held_out

Deduplicate, set aside a held-out set, then drop training texts that leak into it.

The held-out set is chosen first and frozen, before any training or prompt tuning, so no decision is ever made by looking at it.

def margin_of_error(accuracy: float, n: int, z: float = 1.96) -> float: on GitHub
1435def margin_of_error(accuracy: float, n: int, z: float = 1.96) -> float:
1436    """Half-width of the 95% interval around an accuracy measured on n examples: z·sqrt(a(1-a)/n)."""
1437    return z * np.sqrt(accuracy * (1 - accuracy) / n)

Half-width of the 95% interval around an accuracy measured on n examples: z·sqrt(a(1-a)/n).

def measured_accuracy(true_accuracy: float, label_error: float) -> float: on GitHub
1440def measured_accuracy(true_accuracy: float, label_error: float) -> float:
1441    """What a yes/no eval reports when a fraction `label_error` of its reference labels are wrong.
1442
1443    The model scores a point when it is right on a right label, or wrong on
1444    a wrong label (the two mistakes cancel): a(1 - ε) + (1 - a)ε.
1445    """
1446    a, e = true_accuracy, label_error
1447    return a * (1 - e) + (1 - a) * e

What a yes/no eval reports when a fraction label_error of its reference labels are wrong.

The model scores a point when it is right on a right label, or wrong on a wrong label (the two mistakes cancel): a(1 - ε) + (1 - a)ε.

def label_agreement(labels_1: list, labels_2: list) -> float: on GitHub
1450def label_agreement(labels_1: list, labels_2: list) -> float:
1451    """Share of items two labelers labelled the same way."""
1452    return float(np.mean([p == q for p, q in zip(labels_1, labels_2)]))

Share of items two labelers labelled the same way.

N_PARAMS = 97
class TinyNet: on GitHub
1474class TinyNet:
1475    """A two-layer network whose every weight lives in one flat vector `theta`.
1476
1477    Keeping θ flat is what makes fine-tuning arithmetic literal: a task
1478    vector is `theta_tuned - theta_base`, and merging is vector addition.
1479    """
1480
1481    def __init__(self, theta: np.ndarray):
1482        self.theta = np.asarray(theta, dtype=float)
1483
1484    @classmethod
1485    def random(cls, seed: int = 0) -> "TinyNet":
1486        rng = np.random.default_rng(seed)
1487        parts = {
1488            "W1": rng.normal(0, 1, (N_IN, N_HIDDEN)),
1489            "b1": np.zeros(N_HIDDEN),
1490            # 1/sqrt(fan-in) keeps the output's starting scale near 1 (see primer.ml.deep_nets).
1491            "w2": rng.normal(0, 1 / np.sqrt(N_HIDDEN), N_HIDDEN),
1492            "b2": np.zeros(1),
1493        }
1494        return cls(np.concatenate([parts[name].ravel() for name, _ in _SHAPES]))
1495
1496    def _forward(self, X: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
1497        p = _unpack(self.theta)
1498        h = np.tanh(X @ p["W1"] + p["b1"])  # (n, 16)
1499        z = h @ p["w2"] + p["b2"][0]  # (n,)
1500        return h, 1 / (1 + np.exp(-z))
1501
1502    def predict(self, X: np.ndarray) -> np.ndarray:
1503        """Probability of "yes" for each row of X (shape (n, 4))."""
1504        return self._forward(X)[1]
1505
1506    def accuracy(self, X: np.ndarray, y: np.ndarray) -> float:
1507        return float(np.mean((self.predict(X) > 0.5) == y))
1508
1509    def loss(self, X: np.ndarray, y: np.ndarray) -> float:
1510        """Average cross-entropy: -log of the probability given to the right answer."""
1511        q = np.clip(self.predict(X), 1e-12, 1 - 1e-12)
1512        return float(-np.mean(y * np.log(q) + (1 - y) * np.log(1 - q)))
1513
1514    def gradient(self, X: np.ndarray, y: np.ndarray) -> np.ndarray:
1515        """d loss / d θ by backpropagation, flattened like θ."""
1516        p = _unpack(self.theta)
1517        h, q = self._forward(X)
1518        dz = (q - y) / len(y)  # (n,): sigmoid + cross-entropy gives this simple error
1519        dh = np.outer(dz, p["w2"]) * (1 - h**2)  # (n, 16): back through tanh
1520        grads = {"W1": X.T @ dh, "b1": dh.sum(0), "w2": h.T @ dz, "b2": np.array([dz.sum()])}
1521        return np.concatenate([grads[name].ravel() for name, _ in _SHAPES])

A two-layer network whose every weight lives in one flat vector theta.

Keeping θ flat is what makes fine-tuning arithmetic literal: a task vector is theta_tuned - theta_base, and merging is vector addition.

TinyNet(theta: numpy.ndarray) on GitHub
1481    def __init__(self, theta: np.ndarray):
1482        self.theta = np.asarray(theta, dtype=float)
theta
@classmethod
def random(cls, seed: int = 0) -> TinyNet: on GitHub
1484    @classmethod
1485    def random(cls, seed: int = 0) -> "TinyNet":
1486        rng = np.random.default_rng(seed)
1487        parts = {
1488            "W1": rng.normal(0, 1, (N_IN, N_HIDDEN)),
1489            "b1": np.zeros(N_HIDDEN),
1490            # 1/sqrt(fan-in) keeps the output's starting scale near 1 (see primer.ml.deep_nets).
1491            "w2": rng.normal(0, 1 / np.sqrt(N_HIDDEN), N_HIDDEN),
1492            "b2": np.zeros(1),
1493        }
1494        return cls(np.concatenate([parts[name].ravel() for name, _ in _SHAPES]))
def predict(self, X: numpy.ndarray) -> numpy.ndarray: on GitHub
1502    def predict(self, X: np.ndarray) -> np.ndarray:
1503        """Probability of "yes" for each row of X (shape (n, 4))."""
1504        return self._forward(X)[1]

Probability of "yes" for each row of X (shape (n, 4)).

def accuracy(self, X: numpy.ndarray, y: numpy.ndarray) -> float: on GitHub
1506    def accuracy(self, X: np.ndarray, y: np.ndarray) -> float:
1507        return float(np.mean((self.predict(X) > 0.5) == y))
def loss(self, X: numpy.ndarray, y: numpy.ndarray) -> float: on GitHub
1509    def loss(self, X: np.ndarray, y: np.ndarray) -> float:
1510        """Average cross-entropy: -log of the probability given to the right answer."""
1511        q = np.clip(self.predict(X), 1e-12, 1 - 1e-12)
1512        return float(-np.mean(y * np.log(q) + (1 - y) * np.log(1 - q)))

Average cross-entropy: -log of the probability given to the right answer.

def gradient(self, X: numpy.ndarray, y: numpy.ndarray) -> numpy.ndarray: on GitHub
1514    def gradient(self, X: np.ndarray, y: np.ndarray) -> np.ndarray:
1515        """d loss / d θ by backpropagation, flattened like θ."""
1516        p = _unpack(self.theta)
1517        h, q = self._forward(X)
1518        dz = (q - y) / len(y)  # (n,): sigmoid + cross-entropy gives this simple error
1519        dh = np.outer(dz, p["w2"]) * (1 - h**2)  # (n, 16): back through tanh
1520        grads = {"W1": X.T @ dh, "b1": dh.sum(0), "w2": h.T @ dz, "b2": np.array([dz.sum()])}
1521        return np.concatenate([grads[name].ravel() for name, _ in _SHAPES])

d loss / d θ by backpropagation, flattened like θ.

def fine_tune( net: TinyNet, X: numpy.ndarray, y: numpy.ndarray, steps: int = 300, lr: float = 0.5, record=None): on GitHub
1524def fine_tune(net: TinyNet, X: np.ndarray, y: np.ndarray, steps: int = 300, lr: float = 0.5, record=None):
1525    """Full-batch gradient descent from `net`, returning a NEW model (the start is untouched).
1526
1527    With full-batch training one step is one pass over the data: one epoch.
1528    If `record` is given, it is called on the model after every step and
1529    the list of its results is returned alongside the model.
1530    """
1531    theta = net.theta.copy()
1532    history = []
1533    for _ in range(steps):
1534        theta = theta - lr * TinyNet(theta).gradient(X, y)
1535        if record is not None:
1536            history.append(record(TinyNet(theta)))
1537    tuned = TinyNet(theta)
1538    return (tuned, history) if record is not None else tuned

Full-batch gradient descent from net, returning a NEW model (the start is untouched).

With full-batch training one step is one pass over the data: one epoch. If record is given, it is called on the model after every step and the list of its results is returned alongside the model.

TASKS = {'general': "the base model's broad skill: all four numbers anywhere in [-3, 3]; yes when x1 + x3 > 0", 'A': 'first pair only, left region (x1 in [-3, -1]); yes when x2 > 0', 'B': "first pair only, right region (x1 in [1, 3]); yes when x2 < 0 (A's rule, reversed)", 'C': 'second pair only (x3, x4 in [-2, 2]); yes when x3 + x4 > 0'}
TRAIN_SEEDS = {'A': 1, 'B': 2, 'C': 3, 'general': 4}
EVAL_SEEDS = {'A': 11, 'B': 12, 'C': 13, 'general': 14}
def make_task( name: str, n: int = 200, seed: int | None = None) -> tuple[numpy.ndarray, numpy.ndarray]: on GitHub
1552def make_task(name: str, n: int = 200, seed: int | None = None) -> tuple[np.ndarray, np.ndarray]:
1553    """Examples (X of shape (n, 4), y of 0/1) for one of `TASKS`. Default seed = its training set."""
1554    rng = np.random.default_rng(TRAIN_SEEDS[name] if seed is None else seed)
1555    X = np.zeros((n, N_IN))
1556    s = rng.uniform(-2, 2, n)
1557    if name == "A":
1558        X[:, 0], X[:, 1], y = rng.uniform(-3, -1, n), s, s > 0
1559    elif name == "B":
1560        X[:, 0], X[:, 1], y = rng.uniform(1, 3, n), s, s < 0
1561    elif name == "C":
1562        X[:, 2], X[:, 3] = rng.uniform(-2, 2, n), s
1563        y = X[:, 2] + s > 0
1564    elif name == "general":
1565        X = rng.uniform(-3, 3, (n, N_IN))
1566        y = X[:, 0] + X[:, 2] > 0
1567    else:
1568        raise ValueError(f"unknown task {name!r}; choose from {list(TASKS)}")
1569    return X, y.astype(float)

Examples (X of shape (n, 4), y of 0/1) for one of TASKS. Default seed = its training set.

def eval_accuracy(net: TinyNet, name: str) -> float: on GitHub
1578def eval_accuracy(net: TinyNet, name: str) -> float:
1579    """Accuracy on a task's held-out set."""
1580    return net.accuracy(*_eval_set(name))

Accuracy on a task's held-out set.

@lru_cache(maxsize=1)
def base_model() -> TinyNet: on GitHub
1583@lru_cache(maxsize=1)
1584def base_model() -> TinyNet:
1585    """The "pretrained" base: a random network trained on the general skill."""
1586    return fine_tune(TinyNet.random(seed=0), *make_task("general"), steps=300, lr=0.5)

The "pretrained" base: a random network trained on the general skill.

@lru_cache(maxsize=None)
def fine_tuned( name: str, steps: int = 300, lr: float = 0.5) -> TinyNet: on GitHub
1589@lru_cache(maxsize=None)
1590def fine_tuned(name: str, steps: int = 300, lr: float = 0.5) -> TinyNet:
1591    """The base model fine-tuned on one task's training set."""
1592    return fine_tune(base_model(), *make_task(name), steps=steps, lr=lr)

The base model fine-tuned on one task's training set.

def forgetting(acc_before: float, acc_after: float) -> float: on GitHub
1600def forgetting(acc_before: float, acc_after: float) -> float:
1601    """How much accuracy on the old task was lost: F = acc_before - acc_after."""
1602    return acc_before - acc_after

How much accuracy on the old task was lost: F = acc_before - acc_after.

@lru_cache(maxsize=None)
def general_skill_run(lr: float = 0.5, steps: int = 300) -> dict[str, list[float]]: on GitHub
1605@lru_cache(maxsize=None)
1606def general_skill_run(lr: float = 0.5, steps: int = 300) -> dict[str, list[float]]:
1607    """Fine-tune the base on task A, recording after every step: accuracy on A,
1608    accuracy on the base's general skill, and distance travelled from the base."""
1609    base = base_model()
1610
1611    def record(net: TinyNet):
1612        return eval_accuracy(net, "A"), eval_accuracy(net, "general"), float(np.linalg.norm(net.theta - base.theta))
1613
1614    _, history = fine_tune(base, *make_task("A"), steps=steps, lr=lr, record=record)
1615    task, general, distance = (list(col) for col in zip(*history))
1616    return {"task": task, "general": general, "distance": distance}

Fine-tune the base on task A, recording after every step: accuracy on A, accuracy on the base's general skill, and distance travelled from the base.

def replay_mix( new: tuple, old: tuple, n_old: int) -> tuple[numpy.ndarray, numpy.ndarray]: on GitHub
1619def replay_mix(new: tuple, old: tuple, n_old: int) -> tuple[np.ndarray, np.ndarray]:
1620    """The new task's examples plus the first `n_old` examples of an old task."""
1621    (X_new, y_new), (X_old, y_old) = new, old
1622    return np.concatenate([X_new, X_old[:n_old]]), np.concatenate([y_new, y_old[:n_old]])

The new task's examples plus the first n_old examples of an old task.

def mixed_loss(loss_new: float, n_new: int, loss_old: float, n_old: int) -> float: on GitHub
1625def mixed_loss(loss_new: float, n_new: int, loss_old: float, n_old: int) -> float:
1626    """Average loss over a mixed dataset: (n_new·L_new + n_old·L_old) / (n_new + n_old)."""
1627    return (n_new * loss_new + n_old * loss_old) / (n_new + n_old)

Average loss over a mixed dataset: (n_new·L_new + n_old·L_old) / (n_new + n_old).

@lru_cache(maxsize=None)
def sequential_run(new: str, lr: float = 0.5, steps: int = 300, n_replay: int = 0) -> dict: on GitHub
1630@lru_cache(maxsize=None)
1631def sequential_run(new: str, lr: float = 0.5, steps: int = 300, n_replay: int = 0) -> dict:
1632    """Start from the task-A fine-tune, then fine-tune on `new` (optionally replaying A).
1633
1634    Returns accuracy on A before and after, accuracy on the new task after,
1635    and the (acc on A, acc on new) pair after every step.
1636    """
1637    start = fine_tuned("A")
1638    data = make_task(new)
1639    if n_replay:
1640        data = replay_mix(data, make_task("A"), n_replay)
1641    end, path = fine_tune(start, *data, steps=steps, lr=lr, record=lambda net: (eval_accuracy(net, "A"), eval_accuracy(net, new)))
1642    return {
1643        "A_before": eval_accuracy(start, "A"),
1644        "A_after": eval_accuracy(end, "A"),
1645        f"{new}_after": eval_accuracy(end, new),
1646        "path": path,
1647    }

Start from the task-A fine-tune, then fine-tune on new (optionally replaying A).

Returns accuracy on A before and after, accuracy on the new task after, and the (acc on A, acc on new) pair after every step.

@lru_cache(maxsize=None)
def overfitting_run( n: int = 16, n_wrong: int = 3, epochs: int = 1500, lr: float = 0.5, seed: int = 2) -> dict: on GitHub
1655@lru_cache(maxsize=None)
1656def overfitting_run(n: int = 16, n_wrong: int = 3, epochs: int = 1500, lr: float = 0.5, seed: int = 2) -> dict:
1657    """Fine-tune on n task-A examples, `n_wrong` of them mislabelled, for many epochs.
1658
1659    Records training loss and held-out (validation) loss after every epoch.
1660    """
1661    X, y = make_task("A", n=n, seed=seed)
1662    wrong = np.random.default_rng(seed).choice(n, n_wrong, replace=False)
1663    y = y.copy()
1664    y[wrong] = 1 - y[wrong]  # real datasets have some bad labels; a big enough model memorises them
1665    X_val, y_val = _eval_set("A")
1666    _, history = fine_tune(base_model(), X, y, steps=epochs, lr=lr, record=lambda net: (net.loss(X, y), net.loss(X_val, y_val)))
1667    train, val = (list(col) for col in zip(*history))
1668    return {"train": train, "val": val, "best_epoch": int(np.argmin(val)) + 1}  # counted from 1, like the 1,500 epochs

Fine-tune on n task-A examples, n_wrong of them mislabelled, for many epochs.

Records training loss and held-out (validation) loss after every epoch.

def task_vector( tuned: TinyNet, base: TinyNet) -> numpy.ndarray: on GitHub
1676def task_vector(tuned: TinyNet, base: TinyNet) -> np.ndarray:
1677    """τ = θ_tuned - θ_base: everything the fine-tune changed, as one vector."""
1678    return tuned.theta - base.theta

τ = θ_tuned - θ_base: everything the fine-tune changed, as one vector.

def merge( base: TinyNet, task_vectors: list[numpy.ndarray], lam: float = 1.0) -> TinyNet: on GitHub
1681def merge(base: TinyNet, task_vectors: list[np.ndarray], lam: float = 1.0) -> TinyNet:
1682    """θ_base + λ·Σ τ. With λ = 1/T this is exactly the average of the T fine-tunes."""
1683    return TinyNet(base.theta + lam * np.sum(task_vectors, axis=0))

θ_base + λ·Σ τ. With λ = 1/T this is exactly the average of the T fine-tunes.

def cosine(a: numpy.ndarray, b: numpy.ndarray) -> float: on GitHub
1686def cosine(a: np.ndarray, b: np.ndarray) -> float:
1687    """cos of the angle between two vectors: 1 same direction, 0 unrelated, -1 opposite."""
1688    return float(a @ b / (np.linalg.norm(a) * np.linalg.norm(b)))

cos of the angle between two vectors: 1 same direction, 0 unrelated, -1 opposite.

LAMBDAS = (0.25, 0.5, 0.75, 1.0, 1.25, 1.5)
@lru_cache(maxsize=None)
def merge_run( first: str, second: str, lambdas: tuple = (0.25, 0.5, 0.75, 1.0, 1.25, 1.5)) -> dict: on GitHub
1694@lru_cache(maxsize=None)
1695def merge_run(first: str, second: str, lambdas: tuple = LAMBDAS) -> dict:
1696    """Merge the fine-tunes on two tasks by task arithmetic at several λ.
1697
1698    Returns the cosine between their task vectors and, for each λ, the
1699    merged model's accuracy on (first, second).
1700    """
1701    base = base_model()
1702    taus = [task_vector(fine_tuned(t), base) for t in (first, second)]
1703    by_lambda = {}
1704    for lam in lambdas:
1705        merged = merge(base, taus, lam)
1706        by_lambda[lam] = (eval_accuracy(merged, first), eval_accuracy(merged, second))
1707    return {"cosine": cosine(*taus), "by_lambda": by_lambda}

Merge the fine-tunes on two tasks by task arithmetic at several λ.

Returns the cosine between their task vectors and, for each λ, the merged model's accuracy on (first, second).

def figures() -> dict: on GitHub
1715def figures() -> dict:
1716    """Plot this lesson's data. matplotlib is imported here, and only here,
1717    so the lesson itself needs nothing beyond NumPy."""
1718    import matplotlib
1719
1720    matplotlib.use("Agg")
1721    import matplotlib.pyplot as plt
1722
1723    BLUE, RED, GREEN, AMBER, MUTED = "#2563eb", "#dc2626", "#059669", "#d97706", "#9ca3af"
1724    figs = {}
1725
1726    # --- 1. Break-even: cumulative cost of prompting vs fine-tuning ---------
1727    c_prompt, c_tuned, fixed = per_request_cost(3000, 2.0), per_request_cost(300, 4.0), 600
1728    n = np.linspace(0, 250_000, 200)
1729    n_star = break_even_requests(fixed, c_prompt, c_tuned)
1730    fig, ax = plt.subplots(figsize=(6, 3.6))
1731    ax.plot(n, c_prompt * n, color=RED, label=r"long prompt: \$0.0060 per request")
1732    ax.plot(n, fixed + c_tuned * n, color=BLUE, label=r"fine-tuned: \$600 once + \$0.0012 per request")
1733    ax.axvline(n_star, color=MUTED, ls="--")
1734    ax.text(n_star * 1.03, 150, f"break-even\n{n_star:,.0f} requests", color="#4b5563")
1735    ax.set_xlabel("requests served")
1736    ax.set_ylabel("total cost so far ($)")
1737    ax.set_title("When does a fine-tune pay for itself?")
1738    ax.legend(frameon=False, loc="upper left")
1739    figs["break_even"] = fig
1740
1741    # --- 2. Margin of error vs held-out set size -----------------------------
1742    sizes = np.logspace(np.log10(25), np.log10(3200), 100)
1743    fig, ax = plt.subplots(figsize=(6, 3.4))
1744    ax.plot(sizes, 100 * margin_of_error(0.8, sizes), color=BLUE)
1745    for size in (100, 400, 1600):
1746        m = 100 * margin_of_error(0.8, size)
1747        ax.plot(size, m, "o", color=BLUE)
1748        ax.annotate(f"n = {size}: ±{m:.1f}", (size, m), textcoords="offset points", xytext=(6, 6))
1749    ax.set_xscale("log")
1750    ax.set_xlabel("held-out examples n (log scale)")
1751    ax.set_ylabel("95% margin (percentage points)")
1752    ax.set_title("How precise is an 80% score? Four times the data halves the margin")
1753    figs["eval_margin"] = fig
1754
1755    # --- 3. The general skill fades while task A is learned -----------------
1756    runs = {lr: general_skill_run(lr=lr, steps=300) for lr in (0.5, 0.02)}
1757    steps = np.arange(1, 301)
1758    fig, (a1, a2) = plt.subplots(1, 2, figsize=(9.5, 3.6))
1759    for lr, style in ((0.5, "-"), (0.02, "--")):
1760        a1.plot(steps, runs[lr]["task"], style, color=BLUE, label=f"task A, lr {lr}")
1761        a1.plot(steps, runs[lr]["general"], style, color=RED, label=f"general skill, lr {lr}")
1762        a2.plot(runs[lr]["distance"], runs[lr]["general"], style, color=RED, label=f"lr {lr}")
1763    a1.set_xscale("log")
1764    a1.set_xlabel("fine-tuning step (log scale)")
1765    a1.set_ylabel("held-out accuracy")
1766    a1.set_ylim(0.25, 1.02)
1767    a1.set_title("Learning A, forgetting the general skill")
1768    a1.legend(frameon=False, fontsize=8, loc="lower right")
1769    a2.set_xlabel("distance from the base  ‖θ − θ_base‖")
1770    a2.set_ylabel("general-skill accuracy")
1771    a2.set_title("Forgetting tracks how far the weights move")
1772    a2.legend(frameon=False)
1773    fig.tight_layout()
1774    figs["general_skill"] = fig
1775
1776    # --- 4. Sequential fine-tuning: A then B, with and without replay -------
1777    plain, replay = sequential_run("B"), sequential_run("B", n_replay=10)
1778    fig, ax = plt.subplots(figsize=(6, 3.6))
1779    for run, style, tag in ((plain, "-", "B only"), (replay, "--", "B + 10 replayed A")):
1780        acc_a, acc_b = zip(*run["path"])
1781        ax.plot(steps, acc_a, style, color=BLUE, label=f"accuracy on A ({tag})")
1782        ax.plot(steps, acc_b, style, color=GREEN, label=f"accuracy on B ({tag})")
1783    ax.set_xscale("log")
1784    ax.set_xlabel("step of the second fine-tune (log scale)")
1785    ax.set_ylabel("held-out accuracy")
1786    ax.set_title("Fine-tuning on B after A: A collapses unless replayed")
1787    ax.legend(frameon=False, fontsize=8, loc="center left", bbox_to_anchor=(0.3, 0.68))
1788    figs["sequential"] = fig
1789
1790    # --- 5. The trade-off: learning rate only walks the diagonal ------------
1791    fig, ax = plt.subplots(figsize=(5.2, 4.6))
1792    ax.plot([0, 1], [1, 0], color=MUTED, ls=":", label="A + B = 1")
1793    for lr, color in ((0.5, RED), (0.1, AMBER), (0.02, BLUE)):
1794        acc_a, acc_b = zip(*sequential_run("B", lr=lr)["path"])
1795        ax.plot(acc_b, acc_a, "-", color=color, label=f"B only, lr {lr}")
1796        ax.plot(acc_b[-1], acc_a[-1], "o", color=color)
1797    acc_a, acc_b = zip(*replay["path"])
1798    ax.plot(acc_b, acc_a, "--", color=GREEN, label="B + replay, lr 0.5")
1799    ax.plot(acc_b[-1], acc_a[-1], "o", color=GREEN)
1800    ax.set_xlabel("accuracy on the new task B")
1801    ax.set_ylabel("accuracy on the old task A")
1802    ax.set_xlim(0, 1.03)
1803    ax.set_ylim(0, 1.03)
1804    ax.set_title("Slower is not safer; replay escapes the trade-off")
1805    ax.legend(frameon=False, fontsize=8, loc="lower left")
1806    figs["tradeoff"] = fig
1807
1808    # --- 6. Overfitting 16 examples --------------------------------------------
1809    from primer.ml.regularization import early_stopping
1810
1811    run = overfitting_run()
1812    epochs = np.arange(1, len(run["train"]) + 1)
1813    best, stop = early_stopping(run["val"], patience=20)
1814    fig, ax = plt.subplots(figsize=(6, 3.6))
1815    ax.plot(epochs, run["train"], color=BLUE, label="training loss (16 examples, 3 mislabelled)")
1816    ax.plot(epochs, run["val"], color=RED, label="validation loss (200 held-out examples)")
1817    ax.axvline(best + 1, color=MUTED, ls="--")
1818    ax.axvline(stop + 1, color=MUTED, ls=":")
1819    ax.text((best + 1) * 1.08, 1.9, f"best epoch {best + 1}", color="#4b5563", zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1))
1820    ax.set_xscale("log")
1821    ax.set_xlabel("epoch (log scale)")
1822    ax.set_ylabel("loss")
1823    ax.set_title("A small dataset: learning, then memorising")
1824    ax.legend(frameon=False, loc="upper right", fontsize=8)
1825    figs["overfitting"] = fig
1826
1827    # --- 7 and 8. Merging by task arithmetic, a λ sweep -----------------------
1828    lams = tuple(float(x) for x in np.round(np.arange(0, 1.51, 0.125), 3))
1829    for key, second, color in (("merge_separate", "C", GREEN), ("merge_conflict", "B", AMBER)):
1830        merged = merge_run("A", second, lambdas=lams)
1831        acc_first, acc_second = zip(*merged["by_lambda"].values())
1832        fig, ax = plt.subplots(figsize=(6, 3.4))
1833        ax.plot(lams, acc_first, "o-", color=BLUE, label="accuracy on A")
1834        ax.plot(lams, acc_second, "o-", color=color, label=f"accuracy on {second}")
1835        ax.axhline(0.5, color=MUTED, ls=":")
1836        ax.axvline(0.5, color=MUTED, ls="--")
1837        ax.text(0.52, 0.2, "plain average\n(λ = 1/2)", color="#4b5563")
1838        ax.set_ylim(0, 1.05)
1839        ax.set_xlabel("λ: strength of the summed task vectors")
1840        ax.set_ylabel("held-out accuracy of the merge")
1841        ax.set_title(f"base + λ(τ_A + τ_{second}):  cosine(τ_A, τ_{second}) = {merged['cosine']:.2f}")
1842        ax.legend(frameon=False, loc="lower right")
1843        figs[key] = fig
1844
1845    return figs

Plot this lesson's data. matplotlib is imported here, and only here, so the lesson itself needs nothing beyond NumPy.

def demo() -> None: on GitHub
1853def demo() -> None:
1854    banner("1. Should you fine-tune? Put numbers on it")
1855    c_prompt, c_tuned = per_request_cost(3000, 2.0), per_request_cost(300, 4.0)
1856    table(
1857        ["option", "tokens", "$ per million", "$ per request"],
1858        [("long prompt", 3000, 2.0, c_prompt), ("fine-tuned", 300, 4.0, c_tuned)],
1859    )
1860    n_star = break_even_requests(600, c_prompt, c_tuned)
1861    say(f"A $600 fine-tune saves ${c_prompt - c_tuned:.4f} per request, so it pays off after {n_star:,.0f} requests.")
1862    takeaway("Build the eval first; prompt, then retrieve, and fine-tune only for behaviour a prompt can't pin down.")
1863
1864    banner("2. Preparing the data")
1865    example = chat_example("Be brief.", "Capital of France?", "Paris.")
1866    table(["segment", "trained on?"], [(seg, "yes" if trained else "no") for seg, trained in render_chat(example)])
1867    say(f"Problems with an example that ends on the user's turn: {validate_chat({'messages': [{'role': 'user', 'content': 'Hi'}]})}")
1868    questions = ["How do I reset my password?", "how do I reset my password, please", "How do I change my email?"]
1869    table(
1870        ["pair", "Jaccard on word pairs"],
1871        [("password vs. its rewording", jaccard(questions[0], questions[1])), ("password vs. email", jaccard(questions[0], questions[2]))],
1872        floatfmt=".2f",
1873    )
1874    say(f"Deduplicating at 0.7 keeps indices {deduplicate(questions)}: the rewording is dropped.")
1875    table(
1876        ["held-out examples", "95% margin at 80%"],
1877        [(size, f"±{100 * margin_of_error(0.8, size):.1f} points") for size in (25, 100, 400, 1600)],
1878    )
1879    say(
1880        f"""
1881        With 10% wrong labels, a 90%-accurate model measures
1882        {measured_accuracy(0.9, 0.1):.0%}, and a perfect one only
1883        {measured_accuracy(1.0, 0.1):.0%}.
1884        """
1885    )
1886    takeaway("Freeze the held-out set first, keep its near-copies out of training, and check the labels.")
1887
1888    banner("3. Catastrophic forgetting")
1889    say(f"The base model knows its general skill: {eval_accuracy(base_model(), 'general'):.2f} on held-out examples.")
1890    run = general_skill_run(lr=0.5, steps=300)
1891    slow = general_skill_run(lr=0.02, steps=300)
1892    table(
1893        ["fine-tune on A", "accuracy on A", "general skill", "distance from base"],
1894        [
1895            ("lr 0.5, after 5 steps", run["task"][4], run["general"][4], run["distance"][4]),
1896            ("lr 0.5, after 10 steps", run["task"][9], run["general"][9], run["distance"][9]),
1897            ("lr 0.5, after 300 steps", run["task"][-1], run["general"][-1], run["distance"][-1]),
1898            ("lr 0.02, after 300 steps", slow["task"][-1], slow["general"][-1], slow["distance"][-1]),
1899        ],
1900        floatfmt=".3f",
1901    )
1902    say("Task A is learned within a few steps; every step after that only moves the weights further and erodes the general skill.")
1903    rows = []
1904    for label, kwargs in (
1905        ("then B (contradicts A)", dict(new="B")),
1906        ("then B at lr 0.02", dict(new="B", lr=0.02)),
1907        ("then B, 10 steps only", dict(new="B", steps=10)),
1908        ("then B + 10 replayed A", dict(new="B", n_replay=10)),
1909        ("then C (separate inputs)", dict(new="C")),
1910    ):
1911        r = sequential_run(**kwargs)
1912        rows.append((label, r["A_before"], r["A_after"], r[f"{kwargs['new']}_after"], forgetting(r["A_before"], r["A_after"])))
1913    table(["after fine-tuning on A", "A before", "A after", "new task", "forgetting F"], rows, floatfmt=".3f")
1914    takeaway(
1915        "Forgetting grows with how far the weights move and how much the new task conflicts. "
1916        "Slowing down only walks the trade-off; replaying a little old data keeps both."
1917    )
1918
1919    banner("4. Overfitting a small dataset")
1920    over = overfitting_run()
1921    best = over["best_epoch"]  # counted from 1; lists are indexed from 0
1922    table(
1923        ["", "training loss", "validation loss"],
1924        [(f"best epoch ({best})", over["train"][best - 1], over["val"][best - 1]), ("last epoch", over["train"][-1], over["val"][-1])],
1925        floatfmt=".3f",
1926    )
1927    takeaway("On 16 examples the model learns the rule, then memorises the 3 wrong labels. Ship the best checkpoint.")
1928
1929    banner("5. Merging models by task arithmetic")
1930    for second in ("C", "B"):
1931        merged = merge_run("A", second)
1932        say(f"A + {second}: cosine between task vectors = {merged['cosine']:.2f}")
1933        table(
1934            ["λ", "accuracy on A", f"accuracy on {second}"],
1935            [(lam, acc[0], acc[1]) for lam, acc in merged["by_lambda"].items()],
1936            floatfmt=".3f",
1937        )
1938    takeaway(
1939        "Nearly perpendicular task vectors add into a model that does both (0.97 and 0.97 at λ = 1); "
1940        "opposed ones cancel. Averaging is λ = 1/2 and dilutes each skill."
1941    )