How Large Language Models Work

Next-Token Prediction - The Core Task


Here is the entire training objective of a large language model, stated honestly: given some text, guess the next token. Compare the guess to the real next token. Adjust the weights slightly. Repeat a few trillion times.

That is it. There is no separate module trained to reason, no grammar component, no fact database, no objective that rewards being helpful or correct. A single loss function, applied to raw text, over and over.

The reasonable first reaction is that this cannot possibly be enough — it sounds like a fancier version of phone keyboard autocomplete. The interesting question is not whether the objective is simple. It obviously is. The question is why such a simple objective ends up demanding so much, and what actually happens numerically on each of those trillions of steps.

The probabilities after "The cat sat on the"0.610.180.090.070.0501234mat —predictedroof — actualLoss is minus the log of the probability given to the true token: minus log 0.18 equals 1.71.
The model is never asked to be right, only to have put probability on what happened — that is the whole objective.

Why "predict the next token" is a harder task than it sounds

Consider a ladder of sentences, each ending with a blank the model must fill.

TextTo predict the blank you must have learned…
"the cat sat on the ___"Which words follow which — surface statistics
"The cats that live near the river ___"Subject–verb agreement across an intervening clause — syntax
"The capital of Australia is ___"A fact about the world (and that it is Canberra, not Sydney)
"Alice put the keys in her bag. Later, Bob asked where they were. Alice said they were in her ___"Who knows what, and that a state persists across sentences
"137 × 4 = ___"An arithmetic procedure
"The murderer, it turned out, was the ___"Everything the preceding chapters implied — inference over long context
"def is_prime(n):\n if n < 2:\n return ___"Program semantics

Nobody wrote a syntax objective or a fact objective. There is only one loss. But the corpus contains all of these sentences, and the only way to reduce loss on all of them simultaneously is to build internal machinery that handles all of them. Capability is not a separate goal; it is the cheapest way to lower the number.

Next-token prediction is not a weak objective that happens to work. It is a maximally demanding one: any regularity that helps predict text, the loss will pay you to learn.

Formally, the model learns the probability of a whole sequence by factorising it with the chain rule of probability:

P(x1,x2,…,xn)=∏t=1nP(xt∣x1,…,xt−1)P(x_1, x_2, \dots, x_n) = \prod_{t=1}^{n} P(x_t \mid x_1, \dots, x_{t-1})

That factorisation is exact, not an approximation. Modelling every conditional term is equivalent to modelling the full joint distribution over text — which is why the objective can, in principle, capture anything expressible in language.

The loss, computed by hand

Take the context "The cat sat on the" with the true continuation " mat". Suppose a toy vocabulary of five tokens, and the model's final layer emits these raw scores (logits):

Text
token      logitmat         3.2floor       2.1table       1.5sky        -0.4running    -1.2

Logits are unbounded real numbers, so first convert them to a probability distribution with softmax:

pi=ezi∑jezjp_i = \frac{e^{z_i}}{\sum_j e^{z_j}}

Text
exp(3.2)  = 24.532exp(2.1)  =  8.166exp(1.5)  =  4.482exp(-0.4) =  0.670exp(-1.2) =  0.301                sum = 38.151p(mat)     = 24.532 / 38.151 = 0.643p(floor)   =  8.166 / 38.151 = 0.214p(table)   =  4.482 / 38.151 = 0.117p(sky)     =  0.670 / 38.151 = 0.018p(running) =  0.301 / 38.151 = 0.008

Cross-entropy loss for one token is simply the negative log of the probability assigned to the correct answer:

L=−log⁡pcorrect=−log⁡(0.643)=0.442\mathcal{L} = -\log p_{\text{correct}} = -\log(0.643) = 0.442

Three cases make the shape of this function clear:

Model's pp for the true tokenLossReading
0.990.010Nearly certain and right — almost no penalty
0.6430.442Right but with real uncertainty
0.2001.609Uniform guess over 5 options — exactly log⁡5\log 5
0.0104.605Confidently wrong — punished hard
0.0016.908Punished harder still; the penalty is unbounded

The unbounded penalty on the right is the important design property. A loss that merely counted mistakes would give the model no reason to distinguish "wrong but uncertain" from "wrong and adamant". Cross-entropy makes overconfidence catastrophically expensive, which is what teaches a model to represent genuine ambiguity — to put mass on both mat and floor when both are plausible.

Perplexity: the loss in units you can reason about

Perplexity is simply the exponential of the average loss:

PPL=eL\text{PPL} = e^{\mathcal{L}}

Its value is interpretability. A perplexity of kk means the model is, on average, as uncertain as if it were choosing uniformly among kk equally likely options — an effective branching factor.

Average loss (nats)PerplexityWhat that corresponds to
10.8250,257Uniform guessing over a full GPT-2 vocabulary
4.61100A weak model, or an unfamiliar domain
3.0020Solid general English modelling
2.3010Strong; typical of a well-trained large model on in-domain text
0.692Highly predictable text — boilerplate, repeated structure

One caution that trips people up constantly: perplexity is not comparable across tokenisers. A model that splits text into more, smaller tokens will show lower per-token perplexity simply because each individual token is easier. Comparing perplexities is only meaningful when the tokenisation and the evaluation text are identical.

Training: one pass, thousands of predictions

The efficiency of this setup comes from a detail that is easy to miss. A single sequence of 2,048 tokens is not one training example. It is 2,048 training examples, all evaluated in one forward pass.

This works because of the causal mask inside attention: position tt can only see positions 1..t1..t. So the prediction made at every position is already correctly conditioned on exactly its own prefix, and nothing later leaks in. The labels are just the inputs shifted by one:

Text
input:   [The] [cat] [sat] [on] [the]label:   [cat] [sat] [on] [the] [mat]position 1 sees "The"                    -> must predict "cat"position 2 sees "The cat"                -> must predict "sat"position 3 sees "The cat sat"            -> must predict "on"position 4 sees "The cat sat on"         -> must predict "the"position 5 sees "The cat sat on the"     -> must predict "mat"
Python
import torch, torch.nn.functional as Fdef lm_loss(logits, input_ids):    """logits: (batch, seq, vocab)   input_ids: (batch, seq)"""    # drop the last prediction (nothing follows it) and the first label    shift_logits = logits[:, :-1, :].contiguous()    shift_labels = input_ids[:, 1:].contiguous()    return F.cross_entropy(        shift_logits.view(-1, shift_logits.size(-1)),        shift_labels.view(-1),        ignore_index=-100,          # masked positions contribute nothing    )

Feeding the true previous tokens rather than the model's own guesses is called teacher forcing. It is what makes the pass parallel, and it comes with a cost worth naming.

Exposure bias: the mismatch nobody can fully fix

During training, every prediction is conditioned on a perfect prefix — real human text. During generation, the model conditions on its own previous outputs. Once it emits a slightly odd token, it is now operating on a prefix unlike anything in its training distribution, and errors compound.

This is a large part of why long generations drift: a model can produce five excellent paragraphs and then wander, because paragraph six is conditioned on five paragraphs of machine text rather than human text. It is also why generation quality degrades much faster than perplexity would suggest — perplexity is measured under teacher forcing, which never puts the model in this position at all.

Inference: the same model, a completely different cost profile

Generation is inherently sequential. Token 100 cannot be computed before token 99 exists.

Python
def generate(model, ids, n_new):    past = None    for _ in range(n_new):        out = model(ids if past is None else ids[:, -1:], past_key_values=past)        past = out.past_key_values                # reuse cached K and V        next_id = out.logits[:, -1, :].argmax(-1, keepdim=True)        ids = torch.cat([ids, next_id], dim=1)    return ids

The past_key_values cache is what makes this tractable. Without it, generating token nn would require recomputing attention over the whole prefix from scratch — total work growing as O(n2)O(n^2) across the generation. With it, each new token only computes its own key and value and attends against the stored ones.

The result is two distinct phases with completely different bottlenecks:

Prefill (processing the prompt)Decode (generating each token)
Tokens processed per passAll of them at onceExactly one
ParallelismHighNone across time
Limited byArithmetic throughputMemory bandwidth — the whole weight matrix must be read to produce one token
Practical leverSend fewer prompt tokensBatch many users together; quantise weights

That decode row explains something counterintuitive about serving: generating one token for one user uses a tiny fraction of a GPU's arithmetic capacity, because the machine spends its time reading 13 GB of weights out of memory to perform a comparatively small amount of maths. Batching 32 users together reads those weights once and serves all 32 — which is why throughput per user can improve dramatically under load.

Batch size, and why gradient accumulation exists

Large models are trained with very large batches — millions of tokens per optimiser step — because gradients estimated from a small sample are noisy, and noise at this scale destabilises training.

The problem is that a batch that large does not fit in memory. Gradient accumulation solves it by splitting one logical batch into micro-batches, summing gradients across them, and stepping the optimiser only once at the end:

Text
Target batch:      1,048,576 tokens per optimiser stepSequence length:   2,048=> sequences needed: 1,048,576 / 2,048 = 512 sequencesGPU memory allows:  8 sequences at a time=> accumulation steps: 512 / 8 = 64 micro-batches per optimiser stepOn 8 GPUs in parallel: 64 / 8 = 8 accumulation steps each
Python
accum_steps = 64optimizer.zero_grad()for i, batch in enumerate(loader):    loss = lm_loss(model(batch).logits, batch)    (loss / accum_steps).backward()          # scale so the sum is a true mean    if (i + 1) % accum_steps == 0:        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)        optimizer.step()        optimizer.zero_grad()

The division by accum_steps is the bug people ship. Omit it and your effective learning rate is 64× larger than intended, which usually shows up as a loss that looks fine for a hundred steps and then diverges.

How the loss actually falls, and how much compute it takes

Training curves for language models follow a recognisable shape. There is a steep initial drop in the first fraction of a percent of training, as the model learns token frequencies and stops assigning mass to obvious nonsense. After that, the loss decreases as a slow power law in compute — a straight line on log–log axes, with no sharp elbow.

The practical meaning of a power law is that improvements get expensive predictably. Each further reduction in loss costs roughly a constant multiple more compute than the last. There is no point at which the model "finishes"; there is only the point at which further compute stops being worth it.

The compute required is well approximated by a simple formula:

C≈6NDC \approx 6ND

where NN is parameter count and DD is training tokens. The 6 comes from arithmetic: each parameter contributes roughly one multiply and one add per token in the forward pass (2 FLOPs), and the backward pass costs about twice the forward (4 more).

Work it through for a 7-billion-parameter model trained on the compute-efficient ratio of roughly 20 tokens per parameter:

Text
N = 7e9 parametersD = 20 x 7e9 = 1.4e11 tokens  (140 billion)C = 6 x 7e9 x 1.4e11 = 5.88e21 FLOPsAt an effective 1.2e14 FLOP/s per accelerator (about 40% of an A100's312 TFLOP/s bf16 peak; newer accelerators are several times faster):  5.88e21 / 1.2e14 = 4.9e7 seconds = 567 accelerator-daysOn 256 accelerators: about 2.2 days of wall-clock time.

That "20 tokens per parameter" figure is the headline result of compute-optimal scaling work, and it corrected a real mistake in the field. Earlier large models were badly undertrained — parameters were scaled up far faster than data. Given a fixed compute budget, a smaller model trained on more tokens beats a larger model trained on fewer. And for anyone serving a model rather than training it, the incentive pushes further still: training well past the compute-optimal point costs more up front but produces a smaller model that is cheaper for every inference call thereafter, which is why widely deployed open models are trained on hundreds or thousands of tokens per parameter. Meta reports pretraining Llama 3 on over 15 trillion tokens; for the 8B model that is nearly 1,900 tokens per parameter, almost a hundred times the compute-optimal ratio.

What follows from all this in practice

When you evaluate a model, remember what the loss actually measured: average surprise on human text under a perfect prefix. Low perplexity does not mean truthful, helpful, or safe. A model can achieve excellent loss by faithfully reproducing the confident wrongness that exists in its training data — that is not a failure of the objective, it is the objective working exactly as specified.

When a long generation goes off the rails, suspect exposure bias before suspecting the model's knowledge. Shorter generations, chained with fresh grounding in between, keep the model closer to the distribution it was trained on than one enormous open-ended completion does.

And when you are estimating cost, separate the two phases. Prompt tokens are processed in parallel and are comparatively cheap per token; generated tokens are produced one at a time and are the expensive half. A request with a 4,000-token prompt and a 50-token answer has a very different cost profile from one with a 200-token prompt and a 2,000-token answer, even though both move around 4,000 tokens in total. Two current details sharpen this. Reasoning models generate their thinking as ordinary output tokens before the answer, so a 50-token reply can carry thousands of billed, sequentially generated tokens behind it. And providers now cache the processed prefix of repeated prompts (prompt or prefix caching), so a long system prompt that is identical on every call is far cheaper after the first one — another reason to keep the fixed part of a prompt at the front and byte-identical.