Keyboard shortcuts

Press ← or → to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

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.

goaldataobjectivewhat changes
continued pretrainingdomain textnext-token loss on every tokenknowledge and style of a domain
classificationlabelled textscross-entropy over classes, from one positiona new output head, top layers
instruction tuning (SFT)prompt-response pairsnext-token loss on response tokens onlyformat and behavior: answering, stopping
preference tuning (DPO, RLHF)prompts with better and worse responsesa preference objectivewhich 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.

Explore: inputs, targets and the mask

Edit the prompt and response tokens and the batch's padded width. See which positions are learned, and what goes wrong if the mask is applied before the shift, or if labels are shifted twice.

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_ids and labels of 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.py for 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 --accumulation micro-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.json recording the base config’s checksum, dataset checksums, arguments and library versions, and evaluate reloads 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:

itembytes
BF16 weights used in the forward pass2
FP32 master copy of the weights4
gradients (FP32)4
AdamW first and second moments (FP32)8
total18

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:

  1. Held-out loss on response tokens, computed identically before and after.
  2. Exact checks where answers are checkable: classification accuracy, arithmetic, format compliance (“did it stop?”, “is it valid JSON?”).
  3. Regression checks on prompts unrelated to the fine-tuning data: did general ability survive? Small datasets and high learning rates cause catastrophic forgetting.
  4. 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.
  5. 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

  1. ★ 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_batch in engine/lora.py, consumed by engine.sft.response_loss.
  2. ★★ 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), adapting run.py’s cmd_classify and calling engine.sft.train_classifier.
  3. ★★ 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_attention already takes an allowed mask (Chapter 5). Measure the speedup over dynamic padding. Where: add a packed batch builder in engine/lora.py; pass its masks/positions through engine/sft.py and engine/gpt.py to engine.attention.causal_attention.
  4. ★★★ Implement DPO (Rafailov et al., 2023) on BALLM’s preference dataset (ch07/04_preference-tuning-with-dpo in 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 in engine/sft.py; load chosen/rejected pairs in experiments/ch21.py (create it).

Check your understanding

  1. Why does a classifier read the last token’s hidden state, not the first?
  2. In the shifted inputs and targets, which input position predicts the first response token?
  3. Why must padding be masked by position when the pad ID equals the EOS ID?
  4. Why does full fine-tuning need many times the memory of inference?
  5. 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 SFTTrainer documentation, the production version of this chapter’s loop.