Fine-Tuning LLMs

Course Content

Fine-Tuning LLMs

6 sections · 52 lessons

How do you use training/validation loss curves to spot fine-tuning problems early?


The clause classifier's run, read as a pair2.102.120.310.480.520.810.210.410.870.040.630.84Train lossVal lossVal F1step 0step 190step 380step 570Training loss keeps falling after step 380; validation turns up.
The best checkpoint is where validation bottoms out, not where training loss is lowest — so it must have been saved.

What you need to know

What each shape means

ShapeLikely causeFix
Both fall, then flatten togetherHealthyStop soon after validation stops improving
Train keeps falling, validation turns upOverfittingKeep the best checkpoint; fewer epochs, more data, lower LR or rank
Both flat and high from step 1LR too low, LoRA not attached, or all labels maskedCheck print_trainable_parameters(), decode the labels of one batch
Starts near 11–12For a 128k vocabulary, ln(128,000) is about 11.8: the model is guessing at randomWrong chat template, wrong tokenizer or broken weights
Spikes, then NaNLR too high or fp16 overflowLower LR, use bf16, clip gradients at 1.0
Drops in steps at each epoch startThe model is recognising repeated examplesFewer epochs; more or deduplicated data
Validation below trainingDropout active only in training, or an easier validation splitCheck the split before celebrating
Sudden jump mid-runA bad data shard or format changeFind the batch, inspect the data

Set it up so you can act on it

Python
from trl import SFTConfigfrom transformers import EarlyStoppingCallbackargs = SFTConfig(    output_dir="clause-clf",    eval_strategy="steps", eval_steps=50,    save_strategy="steps", save_steps=50,    load_best_model_at_end=True, metric_for_best_model="eval_loss",    logging_steps=10,                      # logs loss, learning rate and grad norm)# trainer = SFTTrainer(..., args=args, callbacks=[EarlyStoppingCallback(early_stopping_patience=3)])

save_steps must match eval_steps so there is a saved checkpoint at every evaluation. load_best_model_at_end reloads the best one, not the last one.

Watch the first 50 steps

Most broken runs show it early: a starting loss that is far too high, a loss that does not move, or a gradient norm that keeps hitting the clipping limit. Stopping a broken run at step 50 saves the other 1,800 steps.

Loss is not the task

Loss measures how likely the reference text is, not whether answers are right. Run a small task eval (accuracy, F1, a judge score) at each checkpoint. For DPO in TRL, also watch rewards/margins, rewards/accuracies and logps/chosen: DPO loss can fall while the chosen answer's own probability also falls, which often shows up as worse outputs.

A real-life example

A Mumbai law firm fine-tunes a legal-clause classifier on 6,000 labelled clauses from Indian commercial contracts (indemnity, termination, arbitration, governing law and so on). Three epochs, LoRA learning rate 2e-4, about 190 optimizer steps per epoch.

StepTrain lossVal lossVal macro-F1
02.102.120.31
190 (end of epoch 1)0.480.520.81
380 (end of epoch 2)0.210.410.87
570 (end of epoch 3)0.040.630.84

Training loss drops sharply at the start of epochs 2 and 3 — the model is recognising clauses it has seen. Validation loss and F1 are best at the end of epoch 2. With load_best_model_at_end, the firm ships the step-380 checkpoint, and the next run uses two epochs. (Made-up numbers for illustration.)

Follow-up questions to expect

  • "Validation loss went up but accuracy also went up. Which do you trust?" — The task metric. Loss can rise because the model is very confident on a few wrong answers while more answers are right. Pick the checkpoint by the metric you ship on.
  • "How big should the validation set be?" — At least a few hundred examples from the same distribution, deduplicated against training, so noise does not decide the checkpoint.
  • "What does the gradient norm tell you?" — Spikes often come just before loss spikes. A norm that sits at the clip value all the time means the learning rate is too high.