Training

Next-token loss

cross-entropy · the signal every weight learns from

shown
real GPT-2 toy last step
examples
5 per sentence
mean loss
4.26

All steps

  1. Inputs and targets

    Training reads real text. Every position predicts the token after it, so one sentence is five examples scored at once; the causal mask keeps each position from seeing its answer.

    inputs = tokens[:-1] · targets = tokens[1:]

  2. Probability of the right token

    For each example, how much probability does GPT-2 small give the right next token? It knows little from “The” alone and much more by “on the”.

    p(target | context)

  3. Cross-entropy loss

    The loss for one example is −log p: 0 for a certain right answer, large for a surprised one. Training minimises the average over billions of tokens, the cross-entropy.

    loss = −(1/N) Σ log p(target)

  4. Gradient on the logits

    The gradient of the loss with respect to the logits is simple: the probabilities minus a one-hot of the target. It says: raise the right token, lower the rest by as much as they took.

    ∂L/∂z = p − y

  5. A step downhill

    Stepping against the gradient raises the right token’s probability and lowers the loss. Here the logits are stepped directly as a toy; real training reaches them through the weights, one batch at a time.

    z ← z − η (p − y)

Code

x, y = idx[:, :-1], idx[:, 1:]             # inputs and targets, shifted by one
logits = model(x)                           # (B, T, 50257): every position at once
probs = F.softmax(logits, dim=-1)
loss = F.cross_entropy(logits.view(-1, logits.size(-1)), y.view(-1))   # mean of −log p
loss.backward()                             # dloss/dlogits = (probs − onehot(y)) / N, then back through every layer
optimizer.step(); optimizer.zero_grad()     # nudge all 124M weights

Go deeper