Course Content
Fine-Tuning LLMs
6 sections · 52 lessons
How do you use training/validation loss curves to spot fine-tuning problems early?
What you need to know
What each shape means
| Shape | Likely cause | Fix |
|---|---|---|
| Both fall, then flatten together | Healthy | Stop soon after validation stops improving |
| Train keeps falling, validation turns up | Overfitting | Keep the best checkpoint; fewer epochs, more data, lower LR or rank |
| Both flat and high from step 1 | LR too low, LoRA not attached, or all labels masked | Check print_trainable_parameters(), decode the labels of one batch |
| Starts near 11–12 | For a 128k vocabulary, ln(128,000) is about 11.8: the model is guessing at random | Wrong chat template, wrong tokenizer or broken weights |
| Spikes, then NaN | LR too high or fp16 overflow | Lower LR, use bf16, clip gradients at 1.0 |
| Drops in steps at each epoch start | The model is recognising repeated examples | Fewer epochs; more or deduplicated data |
| Validation below training | Dropout active only in training, or an easier validation split | Check the split before celebrating |
| Sudden jump mid-run | A bad data shard or format change | Find the batch, inspect the data |
Set it up so you can act on it
1from trl import SFTConfig2from transformers import EarlyStoppingCallback34args = SFTConfig(5 output_dir="clause-clf",6 eval_strategy="steps", eval_steps=50,7 save_strategy="steps", save_steps=50,8 load_best_model_at_end=True, metric_for_best_model="eval_loss",9 logging_steps=10, # logs loss, learning rate and grad norm10)11# 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.
| Step | Train loss | Val loss | Val macro-F1 |
|---|---|---|---|
| 0 | 2.10 | 2.12 | 0.31 |
| 190 (end of epoch 1) | 0.48 | 0.52 | 0.81 |
| 380 (end of epoch 2) | 0.21 | 0.41 | 0.87 |
| 570 (end of epoch 3) | 0.04 | 0.63 | 0.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.