How Models Are Actually Trained
The next-token objective
A language model is trained on one task only — guess the next piece of text — and that single task is enough to teach it grammar, facts and reasoning.
- 13 min read
- 3 reading levels
- Updated
Read these first
On this page 9
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
A language model is trained to guess the next bit of text, over and over, billions of times.
The analogy you have already lived
Open the messaging app on your phone and type "Happy birthday to". Look at the three words the keyboard offers you above the letters. It suggests "you" because it has seen that phrase a thousand times before.
Now imagine that keyboard practising alone, all night, on every book ever written. It reads a few words, covers the next one with its thumb, guesses, then peeks to check. Every wrong guess nudges it a little.
That is the entire training task of a language model. There is no second task.
Why this is such a good idea
Most machine learning needs labels — a correct answer written down by a human for each example. Labels are slow and expensive. Someone has to sit and tag ten thousand photos as cat or dog.
Text does not need this. The correct answer for "the cat sat on the ___" is already sitting in the sentence. The text is its own answer key.
This is called self-supervised learning — learning where the correct answers come from the data itself, with nobody labelling anything.
That is why models can train on trillions of words. Nobody had to label a single one.
How it works
Take a sentence. Feed the model a prefix. Ask for the next piece.
fed to the model the model must guess
------------------ --------------------
"the" → "cat"
"the cat" → "sat"
"the cat sat" → "on"
"the cat sat on" → "the"
"the cat sat on the" → "mat"Here is the part people miss. Those five questions come from one sentence, and the model answers all five in a single pass. Every position in the text is a training example.
The model does not output one word. It outputs a score for every word it knows, all at once. Higher score means "more likely to come next".
"the cat sat on the ___"
|
v
mat ████████████ most likely
floor ████
roof ██
tomato · least likelyTraining pushes the score of the word that actually came next upward, and everything else downward. That is one nudge. Repeat a few trillion times.
Why guessing words teaches so much more than words
This is the part that surprises everyone, so read it slowly.
To guess the next word well, you are forced to learn other things first.
- To finish "the boys ___ playing", you need grammar. "are", not "is".
- To finish "the capital of France is ___", you need a fact.
- To finish "she opened the umbrella because it started to ___", you need cause and effect.
- To finish "2 apples plus 3 apples makes ___", you need a little arithmetic.
Nobody taught the model grammar, geography or arithmetic. They came along for the ride, because they were useful for the guessing game.
Where you have already seen this
- Your phone keyboard's word suggestions.
- Gmail's Smart Compose finishing your sentence in grey.
- Code editors that complete a whole line for you.
- ChatGPT, which is this same guessing game run one piece at a time.
What is honestly hard here
This part confuses almost everyone the first time, so read it twice.
The model is never taught what is true. It is taught what is likely to come next in text. Those two things overlap a great deal, and that is why the model is useful.
But they are not the same thing. When a plausible sentence is also a false one, the model has no built-in reason to prefer the truth. That is the root of hallucination, and no amount of extra training data removes it completely.
Remember this
- The model has exactly one training task: guess the next piece of text.
- The answers are already in the text, so no human labelling is needed.
- Grammar, facts and reasoning are side effects of getting good at guessing.
What to learn next
- Cross-entropy and perplexity — turning this loss number into something you can interpret.
- Transformers — the architecture that computes all those predictions in parallel.
- Tokenization — what "the next piece of text" actually means.
Developer — Code and libraries.
Setup
pip install torchEverything below runs on a CPU in under twenty seconds. No downloads, no GPU.
The whole objective in one line of tensor code
The label tensor is the input tensor shifted left by one. That is the entire supervision signal.
import torch
data = torch.tensor([10, 11, 12, 13, 14, 15, 16]) # a tokenised sentence
x = data[:-1] # everything except the last token
y = data[1:] # everything except the first token
for a, b in zip(x.tolist(), y.tolist()):
print(f"given ...{a} predict {b}")given ...10 predict 11 given ...11 predict 12 given ...12 predict 13 given ...13 predict 14 given ...14 predict 15 given ...15 predict 16
Six training examples from seven tokens. A causal mask — a rule stopping each position from seeing tokens to its right — is what makes this legal. Without it, position 3 could peek at token 4 and the task would be free.
A real language model, trained from scratch, on your CPU
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
torch.manual_seed(0)
TEXT = "the cat sat on the mat. the cat ate the rat. the rat sat on the mat. " * 40
chars = sorted(set(TEXT))
stoi = {c: i for i, c in enumerate(chars)}
itos = {i: c for c, i in stoi.items()}
V = len(chars)
data = torch.tensor([stoi[c] for c in TEXT])
# The whole trick: the labels ARE the input, moved one step to the left.
CTX = 16
x = torch.stack([data[i:i + CTX] for i in range(0, len(data) - CTX - 1, 3)])
y = torch.stack([data[i + 1:i + CTX + 1] for i in range(0, len(data) - CTX - 1, 3)])
print(f"vocabulary size V = {V}")
print(f"{len(x)} training windows of {CTX} characters")
print("input :", repr("".join(itos[i] for i in x[0].tolist())))
print("target :", repr("".join(itos[i] for i in y[0].tolist())))
class TinyLM(nn.Module):
def __init__(self):
super().__init__()
self.emb = nn.Embedding(V, 32)
self.pos = nn.Embedding(CTX, 32)
layer = nn.TransformerEncoderLayer(32, 4, 64, batch_first=True, dropout=0.0)
self.block = nn.TransformerEncoder(layer, 2)
self.head = nn.Linear(32, V)
def forward(self, idx):
t = idx.shape[1]
h = self.emb(idx) + self.pos(torch.arange(t))
# causal mask: position t may not look at anything after t
mask = nn.Transformer.generate_square_subsequent_mask(t)
return self.head(self.block(h, mask=mask, is_causal=True))
model = TinyLM()
opt = torch.optim.AdamW(model.parameters(), lr=3e-3)
print(f"\nloss of a model that knows nothing = ln(V) = {math.log(V):.4f}")
for step in range(301):
logits = model(x)
# flatten every position of every window into one big classification problem
loss = F.cross_entropy(logits.reshape(-1, V), y.reshape(-1))
opt.zero_grad()
loss.backward()
opt.step()
if step % 100 == 0:
print(f"step {step:>3} loss {loss.item():.4f}")
model.eval()
prompt = "the cat s"
idx = torch.tensor([[stoi[c] for c in prompt]])
with torch.no_grad():
for _ in range(14):
nxt = model(idx[:, -CTX:])[0, -1].argmax()
idx = torch.cat([idx, nxt.view(1, 1)], dim=1)
print("\nprompt :", repr(prompt))
print("continuation:", repr("".join(itos[i] for i in idx[0].tolist())))vocabulary size V = 12 915 training windows of 16 characters input : 'the cat sat on t' target : 'he cat sat on th' loss of a model that knows nothing = ln(V) = 2.4849 step 0 loss 2.6746 step 100 loss 0.1280 step 200 loss 0.1163 step 300 loss 0.1146 prompt : 'the cat s' continuation: 'the cat sat on the mat.'
Written against PyTorch 2.5.1 on CPU. Loss values in the fourth decimal place can differ on another machine or PyTorch build — floating-point reduction order is not identical everywhere. The shape of the curve is what matters.
Reading that output
ln(V) = 2.4849 is the score to beat. A model that has learned nothing spreads its probability evenly over all twelve characters. Cross-entropy loss for that is the natural log of the vocabulary size. Step 0 sits at 2.6746, marginally worse than that, which is what random initialisation looks like.
Loss 0.1146 is not zero, and cannot be. After "the cat " the text sometimes continues "sat" and sometimes "ate". No model can score both perfectly. That residue is irreducible uncertainty — real ambiguity in the data, not a bug in your training.
reshape(-1, V) is where the free supervision cashes in. 915 windows times 16 positions gives 14,640 classification problems, all backpropagated from one forward pass. This is why language modelling scales: compute per token is fixed, and every token counts.
The generated text is greedy. argmax takes the top-scoring character every time, so the output is deterministic. Real generation samples instead — see temperature and sampling.
Common mistakes
Forgetting the causal mask. Without mask=, every position sees the whole window, including its own answer. Loss crashes to near zero in twenty steps and generation produces garbage. If your training loss looks impossibly good, check the mask first.
Shifting the wrong tensor. x = data[1:], y = data[:-1] trains the model to predict backwards. It trains fine and the loss looks normal. Print one input–target pair as text before you trust any run.
Passing probabilities to cross_entropy. PyTorch's F.cross_entropy expects raw logits — unnormalised scores straight from the final layer. Apply softmax yourself first and you get a wrong, quietly-plausible loss. More on this trap in logits and loss pitfalls.
Comparing loss across different tokenisers. Loss is per token, and a different tokenizer cuts the same sentence into a different number of tokens. Two models with different vocabularies have losses that are not comparable.
Try it yourself
Change TEXT so the second sentence becomes "the cat ate the mat.". Now "the cat " is followed by "sat" and "ate" equally often. Predict what the final loss does before you run it, then check.
What to learn next
- Cross-entropy and perplexity — turning this loss number into something you can interpret.
- Transformers — the architecture that computes all those predictions in parallel.
- Tokenization — what "the next piece of text" actually means.
Researcher — Mathematics and papers.
The objective
Autoregressive language modelling factorises the joint probability of a token sequence by the chain rule:
$$ p_\theta(x_1, \dots, x_T) = \prod_{t=1}^{T} p_\theta(x_t \mid x_{<t}) $$
Here $x_t$ is the token at position $t$, $x_{<t}$ is every token before it, and $\theta$ are the model parameters. The factorisation is exact — no independence assumption is being made.
Training minimises the negative log-likelihood, which for a discrete vocabulary is cross-entropy against the one-hot empirical distribution:
$$ \mathcal{L}(\theta) = -\frac{1}{T}\sum_{t=1}^{T} \log p_\theta(x_t \mid x_{<t}) $$
$T$ is the number of predicted tokens. The 1/T normalisation makes the loss comparable across sequence lengths, and is why a loss of 2.5 means roughly the same thing on a 512-token batch and a 4096-token batch.
What the minimum actually is
Expanding cross-entropy against the true data distribution $q$:
$$ \mathbb{E}{x \sim q}!\left[-\log p\theta(x)\right] = H(q) + D_{\mathrm{KL}}(q \,|\, p_\theta) $$
$H(q)$ is the entropy of the data — the genuine unpredictability of language. $D_{\mathrm{KL}}$ is the Kullback–Leibler divergence, which is zero only when the model matches the data distribution exactly.
The consequence is worth stating plainly: the loss floor is $H(q) > 0$, not zero. Shannon's classic estimate put English at roughly 0.6 to 1.3 bits per character (Shannon, 1951, Prediction and Entropy of Printed English). Reported per-token losses near 1.8 nats on web text are within a few tenths of nats of plausible estimates of that floor, which is why raw loss improvements have become so small and so expensive.
Teacher forcing and exposure bias
Training conditions every prediction on the ground-truth prefix, never on the model's own output. This is teacher forcing. It parallelises perfectly: all $T$ positions are computed in one forward pass under a causal mask.
At inference the model conditions on its own samples, so the conditioning distribution shifts. This mismatch is exposure bias (Ranzato et al., 2016, Sequence Level Training with Recurrent Neural Networks). Its practical severity is contested. Scheduled sampling (Bengio et al., 2015) was the classic fix and is an inconsistent estimator (Huszár, 2015); modern practice instead handles the gap in post-training, with on-policy methods such as DPO and GRPO.
Why other objectives lost
| Objective | Signal per sequence | Generation | Where it is used now |
|---|---|---|---|
| Causal LM (next token) | every position | native | all frontier LLMs |
| Masked LM (BERT) | ~15% of positions | not native | encoders, retrieval |
| Span corruption (T5) | corrupted spans only | encoder–decoder | some seq2seq work |
| Permutation LM (XLNet) | every position | awkward | largely abandoned |
Clark et al., 2020 (ELECTRA) made the sample-efficiency argument explicit: masked language modelling wastes most of the sequence. Causal language modelling extracts a gradient signal from every single token, which matters enormously once compute rather than data is the binding constraint.
Tay et al., 2023 (UL2) and the Raffel et al., 2020 (T5) ablations both found denoising objectives competitive at fixed compute for understanding tasks, and worse for generation. Decoder-only causal LM won on the strength of one property: it is the objective whose training-time computation exactly matches its inference-time computation.
Multi-token prediction
Predicting only the next token gives a myopic training signal. Gloeckle et al., 2024 (Better & Faster Large Language Models via Multi-token Prediction) add $n$ independent output heads predicting positions $t{+}1 \dots t{+}n$ from a shared trunk:
$$ \mathcal{L}{\text{MTP}} = -\sum{t}\sum_{i=1}^{n} \log p_\theta(x_{t+i} \mid x_{\leq t}) $$
Gains are negligible at small scale and grow with model size, particularly on code. DeepSeek-V3 (DeepSeek-AI, 2024) adopted a sequential variant, and it doubles as a built-in draft model for self-speculative decoding at inference time.
Cost
For a dense decoder with $N$ non-embedding parameters, forward and backward together cost approximately
$$ C \approx 6 N D $$
floating-point operations, where $D$ is the number of training tokens. The factor 6 is 2 for the forward multiply–accumulate, 4 for the backward pass. This estimate (Kaplan et al., 2020) ignores attention's quadratic term, which stays a minority of the cost while sequence length is well below model width. It is the arithmetic behind every scaling law budget.
Papers
- Shannon, Prediction and Entropy of Printed English, 1951 — Bell System Technical Journal.
- Bengio et al., A Neural Probabilistic Language Model, 2003 — jmlr.org/papers/v3/bengio03a
- Radford et al., Improving Language Understanding by Generative Pre-Training (GPT-1), 2018.
- Radford et al., Language Models are Unsupervised Multitask Learners (GPT-2), 2019.
- Raffel et al., Exploring the Limits of Transfer Learning with a Unified Text-to-Text Transformer (T5), 2020 — arxiv.org/abs/1910.10683
- Clark et al., ELECTRA, 2020 — arxiv.org/abs/2003.10555
- Tay et al., UL2: Unifying Language Learning Paradigms, 2023 — arxiv.org/abs/2205.05131
- Gloeckle et al., Better & Faster Large Language Models via Multi-token Prediction, 2024 — arxiv.org/abs/2404.19737
What to learn next
- Cross-entropy and perplexity — turning this loss number into something you can interpret.
- Transformers — the architecture that computes all those predictions in parallel.
- Tokenization — what "the next piece of text" actually means.