Supervised fine-tuning (SFT) is the step that turns a model which continues text into one that follows instructions in your format. The method is plain imitation: show the model conversations in which the assistant replies the way you want, and minimise the negative log-likelihood of those replies, token by token. There is no reward model and no comparison between answers; the data is the specification.
That simplicity hides a set of details that decide whether the run works. The text the model trains on has to match what it will see at inference. The loss has to fall on the right tokens, including the one that ends a reply. Packing has to keep examples apart, and the loss has to be averaged so that each target token counts the same. Get any of these wrong and training still runs, the loss curve still falls, and the model ships subtly broken. This article walks the data path end to end with code, then works a run and lists the failure modes.
What SFT optimises
Given a conversation rendered to tokens x_1 ... x_T and a set of target positions S (the assistant's tokens), SFT minimises -sum over t in S of log p(x_t | x_<t), divided by the number of targets. Every other token is context: the model reads it but is not graded on predicting it. In code the non-target positions carry the label -100, which PyTorch's cross-entropy ignores, and the model shifts labels by one position so that position t predicts token t+1.
What this can teach is behaviour: an output format, a tone, a task decomposition, when to call a tool, when to refuse, and how long to answer. What it teaches poorly is new knowledge. A few thousand examples that mention a fact are a weak and unreliable way to install it, and models fine-tuned on facts they did not already know tend to learn to state things confidently whether or not they are true. Use retrieval for knowledge and SFT for behaviour. SFT is also where preference tuning starts: DPO and its relatives assume a model that already answers in the right shape, which is what SFT provides; DPO alignment for small models picks up from there.
The template is part of the model
A chat model never sees role names. It sees a string built by the chat template: special tokens marking the start of a turn, the role, the content and an end-of-turn marker. The model learns the exact token pattern it is trained on, so the template used for training must be the template the serving stack uses, byte for byte. A different system-prompt default, a missing newline after the role, or a template that strips reasoning blocks from earlier turns produces a distribution shift you will not see in the training loss.
The end-of-turn marker deserves special attention. The model learns to stop only if the token that ends each reply is a target. If your labels cover the reply text but not the marker, or if the trainer's end-of-sequence token is not the marker the template emits, the model learns to keep going: replies that run on into an invented user turn are the symptom. TRL documents this case directly: for base models whose tokenizer already carries a template, set eos_token to the template's turn-end token. The check below renders a conversation both ways and fails if the inference prompt is not a prefix of the training text.
from transformers import AutoTokenizer
tok = AutoTokenizer.from_pretrained("your-org/your-model")
msgs = [
{"role": "system", "content": "You are a support triage assistant."},
{"role": "user", "content": "My invoice shows two charges for May."},
{"role": "assistant", "content": "category=billing; priority=P2"},
]
# What training sees: the whole conversation, rendered.
train_text = tok.apply_chat_template(msgs, tokenize=False)
# What inference sends: everything before the reply, plus the assistant header.
prompt_text = tok.apply_chat_template(msgs[:-1], tokenize=False, add_generation_prompt=True)
assert train_text.startswith(prompt_text), "template renders the prefix differently at train and inference time"
completion = train_text[len(prompt_text):]
print(repr(completion)) # must END with the template's end-of-turn marker, or the model never learns to stopIf the assertion fails, find out why before training. Some templates render earlier turns differently once a later assistant turn exists, for example by removing reasoning text, which is exactly the train-inference mismatch you are trying to rule out. The Llama chat template walks through one template token by token.
Labels: train on replies, not prompts
Training on every token, prompts included, wastes capacity on predicting user messages and system prompts, and with long shared system prompts it can dominate the gradient. The usual choice is assistant-only loss: label the assistant replies, including their end-of-turn marker, and mask everything else. The robust way to build those labels is from character offsets. Render the conversation once, tokenise it once with offsets, and mark each token whose start offset falls inside an assistant reply. Tokenising the pieces separately and concatenating them is a common bug, because tokenisers merge across boundaries and the pieces then disagree with the whole.
IGNORE = -100
def assistant_spans(tok, msgs):
"""Character spans of each assistant reply in the rendered text, end-of-turn marker included."""
spans = []
for i, m in enumerate(msgs):
if m["role"] == "assistant":
start = len(tok.apply_chat_template(msgs[:i], tokenize=False, add_generation_prompt=True))
end = len(tok.apply_chat_template(msgs[: i + 1], tokenize=False))
spans.append((start, end))
return spans
def build_example(tok, msgs, max_len):
text = tok.apply_chat_template(msgs, tokenize=False)
enc = tok(text, add_special_tokens=False, return_offsets_mapping=True) # template already adds BOS
spans = assistant_spans(tok, msgs)
labels = [tid if any(s <= a < e for s, e in spans) else IGNORE
for tid, (a, _) in zip(enc["input_ids"], enc["offset_mapping"])]
ids, labels = enc["input_ids"][:max_len], labels[:max_len]
if all(l == IGNORE for l in labels):
return None # truncation removed every target: drop the row rather than train on nothing
return {"input_ids": ids, "labels": labels}
# Audit: decode only the trained tokens of a few rows and read them.
# print(tok.decode([t for t, l in zip(ex["input_ids"], ex["labels"]) if l != IGNORE]))Two habits catch most label bugs. First, decode only the trained tokens of a handful of rows and read them: you should see the replies and their end markers, and nothing else. Second, count rows whose targets were all truncated away; max_length cuts from the end, which is where the reply lives. TRL implements the same idea: assistant_only_loss=True for conversational data, which requires {% generation %} and {% endgeneration %} tags in the template, and completion-only loss by default for prompt-completion data. JSONL format for fine-tuning covers the record shapes on disk that feed this step.
Packing without leakage
SFT examples vary wildly in length, so padding every row to the longest wastes most of the compute. Packing places several whole examples in one fixed-length row. Done naively, by concatenating the dataset and cutting every N tokens, it causes two problems. Examples are cut in half, so a reply loses its question. And tokens of the second example can attend to the first, so the model learns from context that will never exist at inference.
The fix has three parts. Choose whole examples per row with a bin-packing heuristic such as best-fit decreasing. Restart position ids at zero for each example, and run attention in variable-length mode so each example attends only to itself; FlashAttention's variable-length kernels support this, and Transformers can derive the boundaries from position ids that restart. Finally, mask the first label of each example, so the last token of one conversation is never trained to predict the first token of the next.
def pack(examples, seq_len):
"""Concatenate whole examples into one row. Position ids restart at 0 per example and the
first label of each example is masked, so nothing is predicted across a boundary."""
ids, labels, pos = [], [], []
for ex in examples: # chosen by best-fit decreasing so they fit without cutting
n = len(ex["input_ids"])
ids += ex["input_ids"]
labels += [IGNORE] + ex["labels"][1:]
pos += list(range(n))
assert len(ids) <= seq_len
return {"input_ids": ids, "labels": labels, "position_ids": pos}In TRL, packing=True with the default packing_strategy="bfd" does best-fit decreasing and switches on padding-free batches, which need a FlashAttention backend. The "wrapped" strategy is the aggressive concatenate-and-cut variant; avoid it for chat data.
Normalise the loss over target tokens
Averaging seems trivial until rows hold different numbers of targets. If each micro-batch computes a mean over its own targets and gradient accumulation then averages those means, a micro-batch with 40 target tokens counts as much as one with 4,000. Short replies are over-weighted and long ones under-weighted, and the effective objective changes with batch composition. Several popular trainers shipped this behaviour until it was widely reported in 2024. The correct form sums the per-token loss across the whole optimizer step and divides once by the total number of targets, counted across every device.
import torch
import torch.nn.functional as F
def train_step(model, micro_batches, optimizer, scheduler, max_grad_norm=1.0):
"""One optimizer step over several micro-batches. Every target token weighs the same,
however the tokens are spread across rows and micro-batches."""
total = sum(int((mb["labels"][:, 1:] != -100).sum()) for mb in micro_batches)
# Multi-GPU: all_reduce(total) here so every rank divides by the global count.
for mb in micro_batches:
logits = model(input_ids=mb["input_ids"], position_ids=mb.get("position_ids")).logits
loss_sum = F.cross_entropy(
logits[:, :-1].reshape(-1, logits.size(-1)).float(),
mb["labels"][:, 1:].reshape(-1),
ignore_index=-100, reduction="sum")
(loss_sum / total).backward()
gnorm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_grad_norm)
optimizer.step(); scheduler.step(); optimizer.zero_grad(set_to_none=True)
return total, float(gnorm)TRL's average_tokens_across_devices defaults to true for the multi-device half of this. Whatever you use, log the target-token count per step next to the loss: a step with few targets has a noisy loss, and a sudden drop in targets usually means a masking or truncation change rather than learning.
A configuration that encodes all of this
For a small model, full fine-tuning in bf16 with a learning rate around 1e-5 to 2e-5 (TRL's default is 2e-5), a short warmup, cosine or linear decay, one to three epochs and gradient clipping at 1.0 is a sound starting point. LoRA adapters typically use a higher rate, around 1e-4; LoRA, in depth covers that path. Measure batch size in target tokens per step rather than rows. The TRL v1.14.1 configuration below sets the details from the previous sections explicitly. Two defaults are easy to miss: max_length is 1,024 tokens, and a model passed by name is loaded in float32 unless you set the dtype.
import torch
from trl import SFTConfig, SFTTrainer
args = SFTConfig(
output_dir="triage-sft",
model_init_kwargs={"dtype": torch.bfloat16, # by name, TRL otherwise loads float32
"attn_implementation": "flash_attention_2"}, # padding-free needs FlashAttention
assistant_only_loss=True, # needs {% generation %} tags in the chat template
packing=True, packing_strategy="bfd",
max_length=2048, # default is 1024; long replies are silently cut
learning_rate=1e-5, num_train_epochs=2,
per_device_train_batch_size=4, gradient_accumulation_steps=8,
eval_strategy="steps", eval_steps=200,
)
trainer = SFTTrainer(model="your-org/your-1.5b-model", args=args,
train_dataset=train_ds, eval_dataset=eval_ds)
trainer.train()
Worked example: a ticket-triage model
A team wants a 1.5-billion-parameter model to read support tickets and reply with a category, a priority and one sentence of reasoning, in a fixed format a downstream system parses. They have 18,000 historical tickets labelled by agents. The numbers here are illustrative; the sequence of checks is the point.
First pass: they render with the instruct model's template, train with full-sequence loss and naive packing, and the loss falls nicely. In evaluation the format is right 91% of the time, but 7% of replies continue past the answer with an invented follow-up question from the customer. Decoding the trained tokens shows why: their data pipeline appended the tokenizer's generic end-of-sequence token, not the template's end-of-turn marker, so the model had never been trained to emit the marker the server stops on.
Second pass: assistant-only labels from offsets, the template check passing, best-fit packing with position resets, and token-normalised loss. Format compliance rises to 99.4% and runaway replies disappear. Priority accuracy improves most on long tickets, which the old per-micro-batch mean had under-weighted. At four epochs validation loss rises and outputs turn formulaic, so they keep the two-epoch checkpoint.
Failure modes and trade-offs
| Symptom | Likely cause | Check or fix |
|---|---|---|
| Replies run on or invent the next user turn | End-of-turn marker not a target, or wrong EOS | Decode trained tokens; set eos_token to the turn-end token |
| Good eval loss, worse live behaviour | Template differs between training and serving | Prefix assertion; ship the template with the weights |
| Long answers truncated or vague | max_length too small; targets cut off | Measure length percentiles; raise max_length or split |
| Odd phrases leaking between examples | Naive packing without boundaries | Best-fit packing, position resets, varlen attention |
| Short-answer style dominates | Per-micro-batch mean loss | Sum then divide by targets per step, across devices |
| General ability drops | Too many epochs or a narrow mix | Fewer epochs, mix in general data, regression suite |
| Confident wrong facts | Teaching knowledge through SFT | Move knowledge to retrieval; keep SFT for behaviour |
Trade-offs: assistant-only loss gives a cleaner signal but fewer trained tokens; packing saves compute but needs FlashAttention; more epochs fit the format faster and erode general behaviour sooner. Evaluate on held-out tasks, as evaluating small language models describes.
What to do next
- Run the template prefix check against the exact template your server uses, and confirm each reply ends with the end-of-turn marker.
- Build labels from character offsets, then decode the trained tokens of 20 random rows and read them.
- Measure token-length percentiles and set max_length so fewer than 1% of rows lose targets to truncation.
- Pack with best-fit decreasing, position ids that restart and the first label of each example masked; avoid concatenate-and-cut packing.
- Normalise the loss by total target tokens per optimizer step across all devices, and log the target count beside the loss.
- Gate releases on held-out task metrics, stop rate and a general regression suite, and ship the template with the weights.