You already know the learning loop
With the small network, we computed a prediction, determined its error, found the derivatives in the backward pass and changed the parameters a little. Here we do the same four steps for many text examples. What is new is the organization: which examples come up when, how large is a step, and when do we save the state?
An optimizer is the rule that turns derivatives into parameter changes. Our simple gradient step is an optimizer. AdamW additionally keeps running statistics of earlier gradients for each parameter and uses them to adjust the step sizes. It also subtracts a small fraction of every weight at each step (weight decay), so large weights shrink more in absolute terms. This does not replace the actual learning from the prediction error.
A learning rate schedule changes the learning rate over time. During warmup it rises gradually at the start. A subsequent cosine schedule lowers it along a smooth cosine curve. You don't need to know this curve by heart: first check whether your training run works with a small fixed learning rate.
A batch contains several prediction tasks
The trainer draws start positions from the prepared sequence of training tokens. From each start position it takes tokens. The first are the input, the last are the target. This produces arrays of equal length without padding.
starts = rng.integers(0, len(data) - context, size=batch_size)
chunks = np.stack([data[i:i+context+1] for i in starts])
x = torch.from_numpy(chunks[:, :-1].copy()).long().to(device)
y = torch.from_numpy(chunks[:, 1:].copy()).long().to(device)copy() prevents problems with certain array views. The data must be longer than the context window. A sensible trainer checks this early, with an understandable error message.
The five core lines
optimizer.zero_grad(set_to_none=True)
logits = model(x)
loss = F.cross_entropy(logits.reshape(-1, V), y.reshape(-1))
loss.backward()
optimizer.step()This is the heart of training, but not yet a complete training setup. You need error messages for NaNs, a controlled learning rate, checkpoints and validation. These additions are there so you can follow what happens, not to offer as many options as possible.
Learning rate and warmup
A fixed learning rate can work well for a small model. Our reference trainer uses a short linear warmup phase followed by a cosine decay down to a minimum fraction. Warmup reduces the first steps while activations and optimizer states settle.
Here s denotes the number of the step currently being run, starting at 0. For a planned run with steps and warmup steps, the rate first rises in proportion to . After that, a progress value between zero and one determines the cosine curve. The schedule is tied to the desired total number of steps. If you change the total when resuming, you also change the rest of the learning rate schedule.
The browser experiment shows the planned rate, not a measured loss. This distinction matters: a pretty cosine curve proves no learning progress.
Gradient accumulation
If a large batch doesn't fit in memory, you can run several small microbatches forward and backward one after another before calling the optimizer. For microbatches of equal size and equal weight, you divide each averaged loss by .
optimizer.zero_grad(set_to_none=True)
for _ in range(accum_steps):
x, y = get_batch()
loss = F.cross_entropy(model(x).reshape(-1, V), y.reshape(-1))
(loss / accum_steps).backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()B counts the examples per microbatch, T the tokens per example and A the accumulated microbatches. The effective batch size in tokens is . For example, 2 × 4 × 3 gives a total of 24 tokens per shared update. A microbatch's activations can be freed after its backward pass. The gradients accumulate in the same parameter fields.
Gradient clipping
Clipping here limits the global gradient norm. Otherwise a sudden, very large derivative can cause a destructive step. But clipping does not replace debugging: if almost every step is heavily clipped or NaNs appear, check the learning rate, data, initialization and numerical operations.
Clipping after every microbatch would not be the same as clipping the summed gradients. Our pretraining clips only after the complete accumulation step.
A concrete run
python train.py --data runs/data --out runs/base \
--steps 2000 --context 128 --dim 128 --layers 4 \
--heads 4 --batch 8 --accum 2 --lr 0.0003These parameters are a starting point, not a promised optimal setting. At tokens per optimizer step, 2000 steps come to about 4.096 million processed tokens. If the corpus has only 100,000 tokens, many examples are used multiple times.
The runtime depends heavily on backend, device, data size and implementation. Measure tokens per second on your machine. A blanket number of minutes based on the parameter count alone would be unreliable.
Validation without accidentally changing state
For validation, we set eval mode, disable gradients and compute the mean loss on separate data. The trainer uses its own random source for validation windows, so an evaluation doesn't change the order of future training batches.
A few windows only give a sample. The output therefore reports an estimated validation loss. For a final measurement, you evaluate the entire validation set with a defined context handling. Different windowing methods can offer different numbers of context tokens for the same targets and are therefore not fully interchangeable.
Checkpoints and resuming
last.pt stores the model, the configuration, tokenizer, optimizer, step count and RNG states. It is replaced atomically via a temporary file, so an interrupted write doesn't overwrite the previous checkpoint. best.pt stores the best observed validation checkpoint of the run.
python train.py --data runs/data --out runs/base --steps 2000 --resume runs/base/last.ptResuming checks that the data is identical. On the same environment, a very similar curve is achievable; across different accelerators or library versions, bit-exact repetition is not guaranteed. The project is not a distributed fault-tolerance system but a trainer you can follow locally.