In supervised fine-tuning, why is the loss computed on completion tokens only?
answer
- condition on it, don't reproduce it
- which half earns gradient
- the model writes both sides
- curve looks fine, samples do not
- stop token must stay supervised
basics
~20 sMasking restricts gradients to the assistant's answer, so the model learns to produce completions rather than reproduce prompts. Train on the whole sequence and it also learns to write user turns and system text, wasting capacity and corrupting behaviour.
solid answer
~50 sA supervised fine-tuning example is a pair: a prompt (system text plus the user turn) and the completion you want the model to produce. Both are fed through the model, but the training objective should only score the completion tokens. Concretely, the prompt positions are marked as ignored in the label sequence, so they still act as context through attention but contribute no gradient. The reason is that you are teaching a conditional distribution — answer given prompt — not the joint distribution of the whole transcript. If you leave the mask off, the model spends capacity modelling your prompt boilerplate and, worse, learns to *emit* it: a support model trained on full transcripts will happily continue past its own answer and write the customer's next message. Completion-only masking is standard, but it is also one of the easiest things to get silently wrong, because a mis-masked run still shows a smoothly falling loss curve.
code
python · 11 linesIGNORE = -100 # cross-entropy ignore_index
prompt_ids = tokenizer(rendered_prompt, add_special_tokens=False).input_ids
completion_ids = tokenizer(rendered_completion, add_special_tokens=False).input_ids
completion_ids = completion_ids + [tokenizer.eos_token_id]
input_ids = prompt_ids + completion_ids
labels = [IGNORE] * len(prompt_ids) + completion_ids
assert len(input_ids) == len(labels)
assert labels[len(prompt_ids)] == completion_ids[0] # first scored tokengo deeper
Know that an SFT example is a prompt plus the completion you want back, and that only the completion is scored. Say plainly that the prompt is context, not a target.
Explain the mechanics: the full sequence is fed forward, prompt positions are set to an ignore label, and cross-entropy skips them. Name the failure — the model generating the user's turn — and note that the end-of-turn token must stay inside the supervised span.
Show how you would catch this in a real run: decode one fully rendered example, verify the first supervised position, and compare per-token loss across prompt and completion spans rather than trusting the aggregate curve. Be ready to diagnose a deployed model that will not stop talking.
Own the guardrail rather than the fix. Argue for a dataset-rendering step that is shared by training and serving, with a golden rendered example checked into the repo, so masking boundaries cannot drift per experiment across a team running many fine-tunes.
## What the objective actually is Supervised fine-tuning (SFT) trains a language model on labelled pairs: an input prompt and the completion a good assistant would produce. The model is still doing next-token prediction — nothing exotic — but the *supervision* is selective. You want to maximise the probability of the completion tokens **conditioned on** the prompt tokens, not the probability of the entire rendered transcript. Mechanically, the whole rendered sequence (system text, user turn, assistant answer, end-of-turn marker) is tokenised and fed forward in one pass. What changes is the label tensor. Positions belonging to the prompt are set to an ignore value that the cross-entropy loss skips; positions belonging to the completion keep their real token ids. The prompt tokens are still fully visible to attention — they are the condition — they simply produce no gradient of their own. ## Why not just train on everything Training on the full sequence is a form of unsupervised language modelling over your dataset, and it has three concrete costs. **Capacity is spent on text you will never need to generate.** In a typical contact-centre dataset the prompt might be a 600-token system instruction plus a 3,000-token call transcript, and the completion a 120-token structured disposition summary. Unmasked, over 95% of the gradient signal is teaching the model to reproduce transcripts and boilerplate it will always be *given*, not asked to write. **The model learns the wrong turn-taking behaviour.** This is the failure people actually hit. If assistant answers and user turns are both in the loss, nothing in the objective says "stop after the assistant turn." The model has been trained that a plausible continuation of an assistant answer is another user message, so at inference it writes its answer and then hallucinates the customer's reply, then its own reply again, until it hits the token cap. Reviewers usually blame the sampling settings; the cause is the mask. **Repeated boilerplate becomes a strong prior.** If every example carries the same system prompt, an unmasked run sees that text thousands of times and drives its probability toward one. The model starts leaking chunks of the system prompt into user-visible output. ## Why the bug is silent The training loss of a mis-masked run looks *better*, not worse. Prompt tokens — long, formulaic, highly predictable — are easy to model, so averaging them into the loss drags the reported number down and makes the run look like it converged nicely. Nothing in the curve flags the defect. It shows up only when you sample from the model and read the output, or when you compare per-token loss on prompt versus completion spans. That is why the standard hygiene step is to decode one fully rendered training example, print which token positions carry a real label, and eyeball that the first non-ignored position is exactly the first token of the assistant answer. ## The edges worth knowing **The end-of-turn token must be inside the mask.** The stop token is part of the target: if it never appears as a supervised label, you have trained a model that does not know how to end. Off-by-one boundaries here are common when the mask is computed by string length rather than by token offsets. **Multi-turn examples have several completion spans.** A three-turn conversation has three assistant turns; the usual practice is to unmask all of them and train on each, which extracts more signal per example than keeping only the final turn. Either choice is defensible, but it must be deliberate — and if you keep only the last turn, the earlier assistant turns still stay in the input as context. **Masking is orthogonal to how many parameters you update.** Completion-only loss applies identically whether you are updating every weight or a small adapter; it is a property of the objective, not of the parameterisation. **Some tasks legitimately want more.** If you are doing continued pretraining on raw domain text rather than instruction tuning, there is no prompt/completion split and you train on everything. The masking rule belongs to instruction-shaped SFT specifically. ## What an interviewer is checking They want to hear that you know SFT is conditional, that you can name the concrete symptom of getting it wrong (the model generating both sides of the conversation), and that you know the loss curve will not tell you. Being able to say "I decode one rendered example and verify the label boundary before launching the run" is the answer that reads as having actually done it.
- Are the masked prompt tokens removed from the input entirely?No. They are fed through the model normally and are fully visible to attention — they are the condition the completion is predicted from. Only their entries in the label sequence are marked ignore, so they contribute no term to the cross-entropy loss. Deleting them would change the task; masking them only changes what is scored.
- In a multi-turn conversation example, which assistant turns do you train on?Usually all of them: each assistant turn is a completion given everything before it, so unmasking every assistant span extracts several supervised targets from one example. Keeping only the final turn is also valid and is sometimes preferred when earlier turns come from a weaker source, but the earlier turns must still remain in the input as context either way.
- Why doesn't the training loss reveal a broken mask?Prompt tokens are long, repetitive and highly predictable, so including them lowers the average loss. A mis-masked run therefore reports a smoother, lower curve than a correct one. Detection has to come from decoding a rendered example and checking the label boundary, or from comparing per-token loss on prompt spans versus completion spans.
saying these in an interview costs you the question
- Thinks masked prompt tokens are deleted from the input
- Believes a falling loss curve proves the mask is correct
- Leaves the end-of-turn token out of the supervised span
- Says loss masking only matters for full-parameter fine-tuning
- Treats the model generating user turns as a sampling-parameter problem