21. Fine-tuning: classifiers and instruction following
In this chapter
- What fine-tuning changes in a pretrained model, and what it can't add.
- Classification fine-tuning: replacing the vocabulary head and reading the last real token.
- Instruction fine-tuning: prompt templates, response-only loss, label shifting and padding, done exactly right.
- Fine-tuning Qwen3-0.6B, the memory full fine-tuning really needs, and how to evaluate the result.
You will build
instruction_batch in engine/lora.py and last_token_logits in engine/sft.py. You'll train a classifier, instruction-tune a GPT on 1,100 examples, and fine-tune a real Qwen3 checkpoint.
Time: 5-7 hours. GPU: recommended for GPT-2 and Qwen3 (the from-scratch demos run on a CPU in about a minute).
What fine-tuning changes
A pretrained model has learned language, facts and patterns from trillions of tokens, but its only skill is continuing text. Ask a base model “What is the capital of France?” and it may continue with three more quiz questions. Fine-tuning keeps training the same weights on a smaller, targeted dataset so the model does a specific job.
| goal | data | objective | what changes |
|---|---|---|---|
| continued pretraining | domain text | next-token loss on every token | knowledge and style of a domain |
| classification | labelled texts | cross-entropy over classes, from one position | a new output head, top layers |
| instruction tuning (SFT) | prompt-response pairs | next-token loss on response tokens only | format and behavior: answering, stopping |
| preference tuning (DPO, RLHF) | prompts with better and worse responses | a preference objective | which of several plausible answers is preferred |
The important intuition: fine-tuning mostly teaches format and behavior, and draws on knowledge already in the weights. A thousand examples can teach a model to answer in one sentence and stop. They can’t teach it chemistry. You’ll see this directly below, when a model trained from scratch learns the answer format perfectly and the facts not at all.
BALLM Chapters 6 and 7 build both classification and instruction tuning on GPT-2; this chapter follows them closely, then moves the same ideas to Qwen3.
Classification: replace the head
A language model ends in a head that maps the final hidden state to vocabulary scores. For a classifier, replace it with a head that maps to class scores: for spam detection, nn.Linear(width, 2) instead of nn.Linear(width, 50257).
Which position’s hidden state should be classified? Under a causal mask, position $t$ has seen tokens $0 \dots t$ only. The last token is the only one that has seen the whole message. With right-padded batches, the last real token is at length - 1, not at -1:
def replace_head(model, num_classes, train_last_blocks=1):
"""Freeze the model, swap the vocabulary head for a num_classes head, and unfreeze the
last few blocks and the final norm (BALLM Chapter 6's recipe)."""
for parameter in model.parameters():
parameter.requires_grad_(False)
width = model.head.in_features
p = model.head.weight
model.head = nn.Linear(width, num_classes, device=p.device, dtype=p.dtype) # new, trainable
for block in list(model.blocks)[len(model.blocks) - train_last_blocks:]:
block.requires_grad_(True)
model.norm.requires_grad_(True)
return model
def last_token_logits(model, ids, lengths):
"""Class scores read at each sequence's last REAL token. (Your engine: Chapter 21)
Under a causal mask only the last position has seen the whole text, and with right padding
that position is lengths - 1, not -1.
"""
logits = model(ids)
return logits[torch.arange(ids.shape[0], device=ids.device), lengths - 1]
def pad_sequences(sequences, pad_id, device="cpu"):
width = max(len(s) for s in sequences)
ids = torch.full((len(sequences), width), pad_id, dtype=torch.long)
for row, seq in enumerate(sequences):
ids[row, :len(seq)] = torch.tensor(seq)
return ids.to(device), torch.tensor([len(s) for s in sequences], device=device)
@torch.no_grad()
def accuracy(model, sequences, labels, pad_id, batch_size=32):
model.eval()
device = next(model.parameters()).device
correct = 0
for i in range(0, len(sequences), batch_size):
ids, lengths = pad_sequences(sequences[i:i + batch_size], pad_id, device)
predicted = last_token_logits(model, ids, lengths).argmax(-1).cpu()
correct += int((predicted == torch.tensor(labels[i:i + batch_size])).sum())
return correct / len(sequences)
def train_classifier(model, sequences, labels, pad_id, steps, batch_size=16, lr=5e-4, seed=0):
device = next(model.parameters()).device
rng = random.Random(seed)
trainable = [p for p in model.parameters() if p.requires_grad]
optimizer = torch.optim.AdamW(trainable, lr=lr, weight_decay=0.1)
for _ in range(steps):
model.train()
index = rng.sample(range(len(sequences)), min(batch_size, len(sequences)))
ids, lengths = pad_sequences([sequences[i] for i in index], pad_id, device)
target = torch.tensor([labels[i] for i in index], device=device)
loss = F.cross_entropy(last_token_logits(model, ids, lengths), target)
optimizer.zero_grad(set_to_none=True)
loss.backward()
optimizer.step()
return loss.item()
replace_head follows BALLM’s recipe: freeze everything, add a fresh trainable head, and unfreeze the last transformer block and the final norm. Training only the top of the network is fast, needs little data, and keeps the pretrained features intact. Train more blocks when you have more data and the task differs more from pretraining.
Run it
Without any download, the demo trains a small GPT on a synthetic task that requires reading the whole sequence (does the token 7 appear anywhere?):
python run.py classify --steps 300
{"before": 0.49166666666666664}
{"train_accuracy": 1.0, "test_accuracy": 1.0}
That took 12 seconds on a laptop CPU. For a real task, download the SMS Spam Collection (5,574 messages, tab-separated label<TAB>text) and pass --data SMSSpamCollection. With pretrained GPT-2 small, frozen except the last block, BALLM reports 95.7% test accuracy after five epochs (BALLM §6.7); balance the classes first, as the book does, so that “always say ham” doesn’t score 87%.
Note
The demo trains its GPT from scratch with all blocks unfrozen (
train_last_blocks=layers), because there are no pretrained features to preserve. With GPT-2 weights, keep the default of one block.
Instruction tuning
Instruction tuning is still next-token prediction. What changes is the data: each example is a prompt and a desired response, joined by a template, and the loss covers only the response.
A template
The template marks where the instruction ends and the answer begins. BALLM uses the Alpaca format:
Below is an instruction that describes a task. Write a response that appropriately completes the request.
### Instruction:
Rewrite the sentence using a simile.
### Input:
The car is very fast.
### Response:
The car is as fast as lightning.<|endoftext|>
Chat models use their own chat template with special tokens instead (Qwen’s <|im_start|> and <|im_end|>, Chapter 4). Two rules matter more than the template you choose: training and inference must use exactly the same format, and the response must end with the stop token the engine stops on. A model never trained to emit <|endoftext|> never stops.
def format_prompt(entry):
"""The Alpaca-style template from BALLM Chapter 7. The response follows the final header."""
text = ("Below is an instruction that describes a task. "
"Write a response that appropriately completes the request."
f"\n\n### Instruction:\n{entry['instruction']}")
if entry.get("input"):
text += f"\n\n### Input:\n{entry['input']}"
return text + "\n\n### Response:\n"
def load_instructions(path):
with open(path, encoding="utf-8") as f:
return json.load(f)
def split_entries(entries, valid_fraction=0.1, test_fraction=0.1, seed=0):
"""Shuffle once, then split by item: every version of an item stays in one split."""
order = list(range(len(entries)))
random.Random(seed).shuffle(order)
n_test, n_valid = int(len(entries) * test_fraction), int(len(entries) * valid_fraction)
pick = lambda idx: [entries[i] for i in idx]
return pick(order[n_test + n_valid:]), pick(order[n_test:n_test + n_valid]), pick(order[:n_test])
def encode_examples(entries, encode, eos_id, max_length):
"""(prompt_ids, response_ids + [eos]) pairs. Over-long examples are dropped and counted, never
truncated: truncation would silently cut the response, the only part that is learned."""
examples, dropped = [], 0
for entry in entries:
prompt, response = encode(format_prompt(entry)), encode(entry["output"]) + [eos_id]
if len(prompt) + len(response) > max_length + 1: # inputs are one shorter than the pair
dropped += 1
continue
examples.append((prompt, response))
return examples, dropped
The dataset is BALLM’s 1,100 instruction-response pairs, in data/instruction-data.json. split_entries splits it into 935 training, 55 validation and 110 test examples before any processing, so no item leaks between splits. encode_examples drops over-long examples instead of truncating them; truncation would cut off the response, which is the only part that’s learned.
Response-only loss and the label shift
Should the model learn to predict the prompt tokens too? Usually not: the prompt is given at inference time, and learning to generate instructions wastes capacity. So prompt positions get the target $-100$, which F.cross_entropy ignores by default. Padding positions get $-100$ too.
The subtle part is the shift. For a prompt [10, 11] and a response [20, 21, EOS], the complete sequence is [10, 11, 20, 21, EOS]. Inputs are all tokens but the last; targets are all tokens but the first:
inputs 10 11 20 21
targets -100 20 21 EOS
The input 11, the last prompt token, must predict 20, the first response token. That’s the most important prediction in the example, and the one most often lost by masking on the wrong side of the shift.
def instruction_batch(examples, pad_id, device="cpu"):
"""examples: list of (prompt_ids, response_ids_ending_in_eos). (Your engine: Chapter 21)
Returns inputs x and targets y, both [B, L-1], already shifted for next-token loss.
Prompt and padding targets are -100 so cross_entropy ignores them; only response tokens
(including the final EOS) are learned.
"""
if not examples or any(not p or not r for p, r in examples):
raise ValueError("Each example needs a non-empty prompt and response")
width = max(len(p) + len(r) for p, r in examples)
ids = torch.full((len(examples), width), pad_id, dtype=torch.long)
labels = torch.full_like(ids, -100)
for row, (prompt, response) in enumerate(examples):
complete = list(prompt) + list(response)
ids[row, :len(complete)] = torch.tensor(complete)
labels[row, len(prompt):len(complete)] = torch.tensor(list(response))
return ids[:, :-1].to(device), labels[:, 1:].to(device)
Warning
Hugging Face models shift internally: you pass unshifted
input_idsandlabelsof the same length, and the model drops the first label. If you shift in your data pipeline and pass the result to a Hugging Face model, every target moves one token too far and the model learns to predict two tokens ahead. Your own GPT expects shifted targets;finetune.pyfor Qwen3 passes unshifted ones. Know which convention each model uses.
When the pad token equals the EOS token (common: GPT-2 has no pad token), mask padding by position, not by ID. Masking every occurrence of the EOS ID would also mask the real end of the response, and the model would never learn to stop.
The training loop
def response_loss(model, x, y):
"""Mean cross-entropy over response targets only; -100 marks prompt and padding."""
logits = model(x)
return F.cross_entropy(logits.reshape(-1, logits.shape[-1]), y.reshape(-1), ignore_index=-100)
@torch.no_grad()
def evaluate_responses(model, examples, pad_id, batch_size=8):
"""Average loss per response token (not per batch), so long and short answers weigh fairly."""
model.eval()
device = next(model.parameters()).device
total, count = 0.0, 0
for i in range(0, len(examples), batch_size):
x, y = instruction_batch(examples[i:i + batch_size], pad_id, device)
logits = model(x)
total += F.cross_entropy(logits.reshape(-1, logits.shape[-1]), y.reshape(-1),
ignore_index=-100, reduction="sum").item()
count += int((y != -100).sum())
return total / max(count, 1)
def finetune(model, train_examples, valid_examples, pad_id, steps, batch_size=8, lr=5e-5,
weight_decay=0.1, warmup=10, eval_every=25, seed=0, log=print):
"""AdamW on response-only loss, each batch padded only to its own longest example."""
device = next(model.parameters()).device
rng = random.Random(seed)
trainable = [p for p in model.parameters() if p.requires_grad]
optimizer = torch.optim.AdamW(trainable, lr=lr, weight_decay=weight_decay)
history = []
for step in range(steps):
scale = (step + 1) / warmup if step < warmup else 0.5 * (1 + math.cos(math.pi * (step - warmup) / max(1, steps - warmup)))
for group in optimizer.param_groups:
group["lr"] = lr * scale
model.train()
x, y = instruction_batch(rng.sample(train_examples, min(batch_size, len(train_examples))), pad_id, device)
loss = response_loss(model, x, y)
optimizer.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(trainable, 1.0)
optimizer.step()
if step % eval_every == 0 or step + 1 == steps:
record = {"step": step + 1, "train_loss": round(loss.item(), 4),
"valid_loss": round(evaluate_responses(model, valid_examples, pad_id), 4)}
history.append(record)
if log:
log(json.dumps(record))
return history
Two details differ from Chapter 7’s loop:
- Dynamic padding. Each batch is padded to its own longest example, not to the dataset’s. Instruction lengths vary a lot, and padding to the maximum wastes compute.
- Token-weighted evaluation. Validation loss is the sum over response tokens divided by their count, not an average of batch averages, which would weight a three-word answer the same as a three-sentence one.
Run it
Without pretrained weights, run.py sft trains its own BPE and a 4-layer GPT from scratch on the 935 training examples only:
python run.py sft --steps 600 --width 192 --layers 4
On a laptop CPU, in 55 seconds:
{"train": 935, "valid": 55, "test": 110, "dropped_overlong": 0}
{"step": 1, "train_loss": 6.9229, "valid_loss": 6.8699}
{"step": 301, "train_loss": 3.9705, "valid_loss": 4.2244}
{"step": 600, "train_loss": 2.8521, "valid_loss": 3.8971}
### What is the contraction for 'it is'?
expected: The contraction for 'it is' is 'it's.'
model: The gas' is 'bua'.
### Convert the number 5 from decimal to binary.
expected: The binary equivalent of the decimal number 5 is 101.
model: The approximately 14 is 24.
### Rewrite the sentence using a simile. The baby is very cute.
expected: The baby is as cute as a button.
model: The ma is a pa.
The model has learned the shape of an answer (“The X is Y.”, then stop) and nothing else. It has never read enough text to know words, let alone facts. That’s fine-tuning without pretraining.
Now the real thing. Load GPT-2 (Chapter 9) and run the same command with --model-dir models/gpt2-medium (a GPU makes this a couple of minutes; BALLM reports 15.8 minutes for GPT-2 medium on an M3 MacBook Air CPU). BALLM’s GPT-2 medium (355M), after two epochs on this dataset, answered (BALLM §7.7):
### Rewrite the sentence using a simile. The car is very fast.
expected: The car is as fast as lightning.
model: The car is as fast as a bullet.
### What type of cloud is typically associated with thunderstorms?
expected: The type of cloud typically associated with thunderstorms is cumulonimbus.
model: The type of cloud associated with thunderstorms is a cumulus cloud.
Same data, same loop; the difference is entirely pretraining. Note the second answer: fluent, confident and wrong. Evaluating instruction-tuned models is hard.
Fine-tuning Qwen3-0.6B
For the real checkpoint, finetune.py uses Hugging Face Transformers for the model, so you can focus on the data and the measurements. The same instruction data is in chat format in data/instruction-{train,valid,test}.jsonl:
uv pip install -r optional-requirements.txt
python finetune.py train --mode full --model-dir models/Qwen3-0.6B \
--train-file data/instruction-train.jsonl --valid-file data/instruction-valid.jsonl \
--output runs/qwen3-sft --steps 200 --lr 1e-5 --batch-size 4 --accumulation 4
python finetune.py evaluate --model-dir models/Qwen3-0.6B \
--valid-file data/instruction-valid.jsonl --output runs/qwen3-sft
What the script does, each of which you’ve now seen from first principles:
- Formats every example with Qwen’s own chat template (
enable_thinking=False), and checks that the prompt-with-generation-prompt is an exact prefix of the full conversation, so the mask boundary is right. - Masks the prompt and padding, keeps the
<|im_end|>that ends the response, and refuses over-long examples. - Accumulates gradients over
--accumulationmicro-batches, dividing by the total number of response tokens in the update (not by the number of micro-batches). - Saves the model with a
manifest.jsonrecording the base config’s checksum, dataset checksums, arguments and library versions, andevaluatereloads it in a fresh process.
Qwen3-0.6B is already instruction-tuned, so this run teaches it the dataset’s terse answer style. To see the full effect of SFT, start from a base checkpoint such as Qwen/Qwen3-0.6B-Base.
The memory of full fine-tuning
Inference needs the weights. Full fine-tuning with AdamW in mixed precision needs, per parameter:
| item | bytes |
|---|---|
| BF16 weights used in the forward pass | 2 |
| FP32 master copy of the weights | 4 |
| gradients (FP32) | 4 |
| AdamW first and second moments (FP32) | 8 |
| total | 18 |
For 0.6B parameters, that’s about 11 GB before activations; for an 8B model, 144 GB. Activations add more, growing with batch size × sequence length × layers (gradient checkpointing trades recomputation for most of it). That’s why the next chapter’s methods exist.
Evaluating a fine-tuned model
A falling training loss proves only that the model memorizes the training set. In increasing order of cost and value:
- Held-out loss on response tokens, computed identically before and after.
- Exact checks where answers are checkable: classification accuracy, arithmetic, format compliance (“did it stop?”, “is it valid JSON?”).
- Regression checks on prompts unrelated to the fine-tuning data: did general ability survive? Small datasets and high learning rates cause catastrophic forgetting.
- Model-as-judge: a stronger model scores responses against references. BALLM §7.8 does this with Llama 3 8B through Ollama, scoring each response 0-100. Useful, but biased toward long and confident answers; spot-check it.
- Human review of a sample, especially of failures.
Build it
Engine milestone 21: fine-tuning. Implement instruction_batch in engine/lora.py and last_token_logits in engine/sft.py (formatting, the training loops and replace_head are provided).
pytest tests/test_ch21_finetune.py
python run.py classify --impl engine
python run.py sft --impl engine --steps 600 --width 192
The tests check the shifted targets and masks of a padded batch, that masked positions contribute nothing to the loss, that classification reads each sequence’s last real token, that replace_head freezes the right parameters, and that over-long examples are dropped rather than truncated. Then fine-tune GPT-2 (or Qwen3-0.6B) and record held-out loss before and after, and three test responses.
Stretch exercises
- ★ Train on all tokens (prompt included) instead of response-only, with the same budget. Compare held-out response loss and the test answers. Where: the response mask built by
instruction_batchinengine/lora.py, consumed byengine.sft.response_loss. - ★★ Classify SMS spam with pretrained GPT-2, comparing three settings: head only, head + last block, all layers. Plot test accuracy against trainable parameters. Where:
experiments/ch21.py(create it), adaptingrun.py’scmd_classifyand callingengine.sft.train_classifier. - ★★ Implement sequence packing: concatenate several examples into one row, with a block-diagonal causal mask and positions restarting at 0 for each example. Your
causal_attentionalready takes anallowedmask (Chapter 5). Measure the speedup over dynamic padding. Where: add a packed batch builder inengine/lora.py; pass its masks/positions throughengine/sft.pyandengine/gpt.pytoengine.attention.causal_attention. - ★★★ Implement DPO (Rafailov et al., 2023) on BALLM’s preference dataset (
ch07/04_preference-tuning-with-dpoin Raschka’s repository): the loss compares the policy’s and a frozen reference model’s log-probabilities of the chosen and rejected responses. Where: add a DPO loss/training helper inengine/sft.py; load chosen/rejected pairs inexperiments/ch21.py(create it).
Check your understanding
- Why does a classifier read the last token’s hidden state, not the first?
- In the shifted inputs and targets, which input position predicts the first response token?
- Why must padding be masked by position when the pad ID equals the EOS ID?
- Why does full fine-tuning need many times the memory of inference?
- Why did the from-scratch model learn the answer format but none of the facts?
Going deeper
- BALLM Chapter 6 (classification fine-tuning: §6.5 adding the head, §6.6 the last token, §6.7 training) and Chapter 7 (instruction fine-tuning: §7.3 batching and masking, §7.6 training, §7.8 evaluation with a judge model). Raschka’s repository has bonus material on DPO, LoRA variants and larger models.
- Ouyang et al., Training language models to follow instructions with human feedback (InstructGPT, 2022); Taori et al., Alpaca (2023); Zhou et al., LIMA: Less Is More for Alignment (2023), on how little data format-teaching needs.
- Rafailov et al., Direct Preference Optimization (2023).
- Hugging Face TRL’s
SFTTrainerdocumentation, the production version of this chapter’s loop.