Post-training and Alignment

Formatting an SFT dataset

Supervised fine-tuning data has to be laid out precisely, and the single most important detail is which tokens the loss is allowed to see.

On this page 8
  1. The short answer
  2. The analogy you have already lived
  3. Why it exists
  4. How it works
  5. The three things that go wrong
  6. Where you have already seen this
  7. Remember this
  8. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

The short answer

Fine-tuning data is a list of question-and-answer pairs, and the model is scored only on the answer half.

The analogy you have already lived

Think of a school notebook where you copy the question and then write your answer. The teacher marks the answer. She does not give you marks for copying the question correctly.

You still need the question on the page. Without it your answer makes no sense. It is context, not work.

A fine-tuning example is that page. The question is context. The answer is what gets marked.

Why it exists

Training a language model works by scoring every position in the text. Left alone, it would score the user's question too.

That is wasted effort, and it teaches the wrong thing. You do not want a model that is good at inventing plausible user questions. You want one that is good at replying.

So you mark the question part as "do not score this". The model still reads it. It is not graded on it.

How it works

Each example has three pieces.

   the wrapper       who is speaking, and where each turn ends
   the question      given to the model, not scored
   the answer        given to the model, and scored

Written out, one example looks like this:

   <user>  Name two Indian classical dance forms.  <end>     <- read, not scored
   <robot> Bharatanatyam and Kathak.               <end>     <- read AND scored
                                                     ^
                                          the model must learn to produce
                                          this end marker, or it never stops

That end marker on the answer is not decoration. It is how the model learns where a reply finishes. Leave it out and your fine-tuned model rambles forever.

The three things that go wrong

The wrapper is written by hand. Every model family uses different markers. Copy them from a blog post, get one space wrong, and quality drops with no error message. Always ask the model's own tokenizer to write the wrapper.

A long example gets cut in half. If your examples are longer than the model's limit, the end gets chopped off. Now you are training the model to answer without ever finishing. Check your length distribution before training.

The scoring covers the question. Easy to do by accident, and the training loss looks completely normal. The next section shows exactly what this costs.

Where you have already seen this

  • A school notebook where only the answer is marked.
  • A form where some boxes are pre-filled and some are yours to fill.
  • Subtitles that show who is speaking before each line.

Remember this

  • Every example is a question and an answer, wrapped in speaker markers.
  • The model reads the question and is scored only on the answer.
  • The end-of-answer marker is what teaches the model to stop.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Runs on a CPU in about twenty seconds.

Loss masking, measured

loss_masking.py
import math
import torch
import torch.nn as nn
import torch.nn.functional as F

# A toy "chat" example. The prompt is unpredictable noise, drawn uniformly from
# 18 symbols. The answer is a rule the model CAN learn. This mirrors real SFT:
# the user's question is not yours to predict, the assistant's reply is.
V, PROMPT_LEN, ANS_LEN = 40, 12, 3
BOS, SEP = 38, 39


def make_batch(n, gen):
    p = torch.randint(0, 18, (n, PROMPT_LEN), generator=gen)
    ans = 20 + (p[:, :1] % 8).repeat(1, ANS_LEN) + torch.arange(ANS_LEN)
    return torch.cat([torch.full((n, 1), BOS), p, torch.full((n, 1), SEP), ans], 1)


gen = torch.Generator().manual_seed(0)
train, test = make_batch(512, gen), make_batch(128, gen)
print("one training row:", train[0].tolist())
print(f"  BOS, then {PROMPT_LEN} noise tokens, then SEP, then a {ANS_LEN}-token answer")


def labels_for(seq, mask_prompt):
    lab = seq[:, 1:].clone()                      # shift-by-one targets
    if mask_prompt:
        lab[:, :PROMPT_LEN + 1] = -100            # -100 = "do not train on this"
    return lab


print("\nlabels with prompt masking (-100 is ignored by cross_entropy):")
print(" ", labels_for(train[:1], True)[0].tolist())
print("labels without masking:")
print(" ", labels_for(train[:1], False)[0].tolist())

n_prompt = PROMPT_LEN + 1
n_ans = train.shape[1] - 1 - n_prompt
print(f"\nof {n_prompt + n_ans} supervised positions per row, {n_prompt} are prompt "
      f"and {n_ans} are answer")
print(f"the prompt tokens are uniform over 18 symbols, so their loss cannot fall "
      f"below ln(18) = {math.log(18):.3f}")


def run(mask_prompt, steps=150):
    torch.manual_seed(1)
    layer = nn.TransformerEncoderLayer(48, 4, 96, batch_first=True, dropout=0.0)
    m = nn.ModuleDict({"emb": nn.Embedding(V, 48), "pos": nn.Embedding(24, 48),
                       "blocks": nn.TransformerEncoder(layer, 2), "head": nn.Linear(48, V)})
    opt = torch.optim.AdamW(m.parameters(), lr=3e-3)

    def fwd(seq):
        h = m["emb"](seq) + m["pos"](torch.arange(seq.shape[1]))
        mask = nn.Transformer.generate_square_subsequent_mask(seq.shape[1])
        return m["head"](m["blocks"](h, mask=mask, is_causal=True))

    for _ in range(steps):
        logits = fwd(train)[:, :-1]
        loss = F.cross_entropy(logits.reshape(-1, V),
                               labels_for(train, mask_prompt).reshape(-1))
        opt.zero_grad()
        loss.backward()
        opt.step()

    with torch.no_grad():
        logits = fwd(test)[:, :-1]
        full = labels_for(test, False)
        per_tok = F.cross_entropy(logits.reshape(-1, V), full.reshape(-1),
                                  reduction="none").reshape(full.shape)
        prompt_loss = per_tok[:, :n_prompt].mean().item()
        ans_loss = per_tok[:, n_prompt:].mean().item()
        acc = (logits.argmax(-1)[:, n_prompt:] == full[:, n_prompt:]).float().mean().item()
    return prompt_loss, ans_loss, acc


print(f"\n{'training signal':<26} {'prompt loss':>12} {'answer loss':>12} {'answer acc':>11}")
for mask_prompt in (False, True):
    pl, al, acc = run(mask_prompt)
    name = "every token" if not mask_prompt else "answer tokens only"
    print(f"{name:<26} {pl:>12.4f} {al:>12.4f} {acc:>10.1%}")
Output
one training row: [38, 8, 9, 11, 6, 7, 15, 7, 1, 1, 9, 17, 2, 39, 20, 21, 22]
  BOS, then 12 noise tokens, then SEP, then a 3-token answer

labels with prompt masking (-100 is ignored by cross_entropy):
  [-100, -100, -100, -100, -100, -100, -100, -100, -100, -100, -100, -100, -100, 20, 21, 22]
labels without masking:
  [8, 9, 11, 6, 7, 15, 7, 1, 1, 9, 17, 2, 39, 20, 21, 22]

of 16 supervised positions per row, 13 are prompt and 3 are answer
the prompt tokens are uniform over 18 symbols, so their loss cannot fall below ln(18) = 2.890

training signal             prompt loss  answer loss  answer acc
every token                      3.9462       0.0244     100.0%
answer tokens only               6.8450       0.0033     100.0%

Written against PyTorch 2.5.1, CPU, fixed seeds — reproducible on this build. Fourth-decimal differences are possible elsewhere.

Read this output honestly, in both directions

13 of 16 supervised positions were prompt tokens. Without masking, 81% of the gradient signal went into predicting uniform noise. The prompt loss reached 3.9462 against an irreducible floor of 2.8900 — effort spent, nothing gained, because the prompt genuinely is unpredictable.

On the thing you care about, masking won by 7×. Answer loss 0.0033 against 0.0244. That is the case for masking, and it is a real case.

But answer accuracy was identical: 100% both ways. On an easy task with plenty of steps, both models learned the rule. This is the part usually left out of tutorials.

The published evidence agrees with that nuance. Shi et al., 2024 found that including the instruction in the loss helps in exactly the fragile regimes — few examples, long prompts, short completions — where it acts as a regulariser. Masking is the right default. It is not a law.

How the layout is expressed in practice

TRL (v1.12.0) accepts four dataset shapes and picks the masking behaviour from the shape:

python
# 1. language modelling - loss on everything
{"text": "The sky is blue."}

# 2. conversational language modelling
{"messages": [{"role": "user", "content": "What color is the sky?"},
              {"role": "assistant", "content": "It is blue."}]}

# 3. prompt-completion - loss on the completion only, BY DEFAULT
{"prompt": "The sky is", "completion": " blue."}

# 4. conversational prompt-completion
{"prompt": [{"role": "user", "content": "What color is the sky?"}],
 "completion": [{"role": "assistant", "content": "It is blue."}]}

The switches that matter, all on SFTConfig:

SettingDefaultEffect
completion_only_lossNoneTrue for prompt-completion data, False for plain text
assistant_only_lossFalsemask everything except assistant turns, in multi-turn data
packingFalsepack several examples per row
packing_strategy"bfd""bfd" best-fit-decreasing, "bfd_split", "wrapped"
max_length1024truncation length — check this against your data
learning_rate2e-5two orders of magnitude below pretraining

assistant_only_loss=True needs the chat template to contain {% generation %} markers so TRL knows which spans are assistant text. TRL patches templates for known families; for anything else, check first.

Check your data before you train it

inspect_lengths.py
import statistics

# stand-in for your tokenised dataset: (prompt_len, completion_len) per example
rows = [(31, 88), (410, 22), (77, 640), (18, 9), (1203, 40), (95, 210),
        (64, 55), (890, 700), (12, 5), (150, 120)]

MAX = 1024
total = [p + c for p, c in rows]
print(f"examples: {len(rows)}")
print(f"median total length: {statistics.median(total):.0f} tokens")
print(f"longest: {max(total)} tokens")
print(f"would be truncated at max_length={MAX}: "
      f"{sum(t > MAX for t in total)} of {len(rows)}")
print(f"completion is under 10 tokens in {sum(c < 10 for _, c in rows)} examples")
print(f"mean fraction of tokens that are completion: "
      f"{statistics.mean(c / (p + c) for p, c in rows):.2f}")
Output
examples: 10
median total length: 288 tokens
longest: 1590 tokens
would be truncated at max_length=1024: 2 of 10
completion is under 10 tokens in 2 examples
mean fraction of tokens that are completion: 0.44

Deterministic arithmetic. Run the equivalent on your real dataset before every fine-tune. The two numbers that predict trouble are the truncation count — anything above zero deserves an explanation — and the completion fraction, because a very low value means most of your compute is being spent reading rather than learning.

Common mistakes

Truncating from the end. truncation_mode="keep_start" keeps the beginning, which drops the answer. For long-prompt data you want the answer to survive, so shorten the prompt yourself instead of relying on truncation.

Blind packing of instruction data. Cutting mid-example splits a question from its answer. Use packing_strategy="bfd", which packs whole examples, not "wrapped", which cuts.

Missing EOS. If the model's end-of-turn token is not present at the end of every completion, generation never stops. With a base model plus a borrowed chat template you must set eos_token explicitly.

Training multi-turn data as a single completion. In a five-turn conversation, the user's turns 2 through 5 should be masked too. That is what assistant_only_loss=True is for.

Leaving the system prompt out of training and using one at inference. The model has then never seen that position. Train with the system prompts you intend to deploy with.

Duplicate or near-duplicate examples. SFT sets are small enough that 200 copies of one example measurably distorts the model. Deduplicate — see building a pretraining corpus for the method.

Try it yourself

In loss_masking.py, change ANS_LEN to 12 so prompt and answer are equal length. Re-run. The gap between the two rows shrinks sharply. That is the completion fraction from the second script, showing up as a training outcome.

What to learn next

Researcher — Mathematics and papers.

The objective, with the mask made explicit

$$ \mathcal{L}_{\text{SFT}}(\theta) = -\frac{1}{\sum_t m_t}\sum_{t=1}^{T} m_t \log p_\theta(y_t \mid y_{<t}) $$

$m_t \in {0,1}$ is the loss mask, $y_t$ the token at position $t$. Setting $m_t = 1$ everywhere gives ordinary language modelling. Setting $m_t = 0$ on prompt tokens gives completion-only SFT.

The denominator matters. Normalising by $\sum_t m_t$ per example weights every example equally; normalising by the batch's total unmasked tokens weights every token equally. The two give different gradients whenever completion lengths vary, and the second is what distributed trainers use, which is why average_tokens_across_devices exists.

Does masking help?

The theoretical argument is clean: the prompt distribution $p(x)$ is not the target, so gradient spent on it is spent on a nuisance objective. The empirical picture is more interesting.

Shi et al., 2024 (Instruction Tuning With Loss Over Instructions, NeurIPS 2024) introduce Instruction Modelling (IM) — loss over instructions as well as outputs — and identify when it wins:

  • Low ratio of completion length to prompt length. Long instructions, short answers.
  • Small training sets. Their result holds notably on the 1,000-example LIMA regime.

They frame the gain as regularisation: SFT on few examples overfits the output distribution, and the instruction-side loss constrains it. Where completions are long and datasets large, masking wins as expected.

The practical rule: treat completion_only_loss as a hyperparameter you sweep once per dataset family, not as a setting you inherit.

Packing without contamination

The packing analysis applies here with one difference: examples must not be split. This is bin packing.

TRL's "bfd" strategy is best-fit decreasing: sort by length descending, place each example in the fullest bin that still fits. BFD is a classic $\tfrac{11}{9}\mathrm{OPT} + \tfrac{6}{9}$ approximation for bin packing, and on realistic SFT length distributions it reaches well above 90% utilisation.

Two correctness requirements accompany it, and both are frequently missed:

  1. Block-diagonal attention. Example $i$ must not attend to example $j$ in the same row. Without it, the model conditions its answer on an unrelated question and its answer.
  2. Position id reset. Each packed example must start at position 0.

TRL's padding_free path uses FlashAttention's varlen kernels, which handle both by construction. A hand-rolled packer that does neither is a silent quality regression.

Data scale and composition

Published SFT mixtures vary over three orders of magnitude, and the disagreement is real rather than accidental:

DatasetSizeSource
LIMA1,000hand-curated
Alpaca52,000Self-Instruct from GPT-3.5
Tülu 3 SFT mix~939,000curated multi-source, decontaminated
OpenHermes 2.5~1,000,000aggregated synthetic

Lambert et al., 2024 (Tülu 3) is the most useful public reference here, because it documents the full recipe and the decontamination procedure rather than only the result. Their finding, consistent across the field, is that mixture composition — how much maths, code, safety, multilingual — moves benchmark scores more than raw example count.

Format sensitivity

A fine-tuned model is sensitive to the exact template string. Sclar et al., 2024 (Quantifying Language Models' Sensitivity to Spurious Features in Prompt Design) showed accuracy swings of tens of points from formatting changes as small as a separator character, on models of every size. The consequence for SFT is direct: render the template with apply_chat_template at training time and at inference time, from the same tokenizer revision, or your evaluation measures template mismatch rather than model quality.

Papers

What to learn next