← All Posts
Deep Learning · Transformers· Training

Training a Transformer: Targets, Gradients, Batches, and Evaluation

Training connects a precise prediction task to parameter updates. Start by verifying what each logit is allowed to see and which token it must predict. A sophisticated optimizer cannot repair a leaking mask or a shifted target bug.

A complete next-token training example

Take the sequence [BOS, the, cat, slept, EOS]. Feed [BOS, the, cat, slept] into a causal decoder. The corresponding targets are [the, cat, slept, EOS]. Position 2 receives the prefix through “cat” and predicts “slept.” It must not receive the later “slept” input through attention.

Given logits $z_t\in\mathbb R^V$ and target ID $y_t$, the per-position loss is

$$\ell_t=-z_{t,y_t}+\log\sum_{v=1}^V e^{z_{t,v}}.$$

Evaluate the log-sum-exp stably. If the correct-token probability is 0.25, the loss is $-\log0.25\approx1.3863$ nats. Raising it to 0.5 lowers the loss to about 0.6931. Cross-entropy rewards probability assigned to the observed target; it does not require the target to be the only plausible continuation.

Try the mechanism Inspect the inputs, targets, and visible prefix

Averaging the two batch means would give 1.5. Count valid targets across the effective batch when reducing the objective. Enable JavaScript to change the example inputs; the complete calculation remains in the article.

The vocabulary gradient has a simple form

$$\frac{\partial\ell_t}{\partial z_{t,v}}=p_{t,v}-\mathbf1[v=y_t].$$

The target logit receives a negative derivative unless its probability is already one; other logits receive positive derivatives. Gradient descent therefore increases the target relative to alternatives. Backpropagation carries this signal through the output projection, residual stack, and embeddings. Tied input/output embeddings receive gradients through both uses.

Attention masks and loss masks are different

An attention mask controls information flow. A loss mask controls which predictions count. In supervised instruction tuning, prompt positions may provide context while only response targets contribute to the objective. Ignoring prompt loss does not stop the model from attending to the prompt.

For valid-target indicators $m_t\in\{0,1\}$, reduce as $\mathcal L=\sum_tm_t\ell_t/\sum_tm_t$. Guard against a batch with zero valid targets. When packing unrelated documents, decide whether attention may cross boundaries; an EOS token alone is not a hard attention barrier.

# inputs and targets are already shifted; model returns [B, T, V].
logits = model(inputs, attention_mask=allowed)
per_token = cross_entropy(logits, targets, reduction="none")
loss_sum = (per_token * target_mask).sum()
valid_count = target_mask.sum()
loss = loss_sum / valid_count  # require valid_count > 0
loss.backward()

This is framework-neutral pseudocode: a real cross-entropy function may expect flattened logits or a different axis order. Do not pass already-softmaxed probabilities to an API that expects logits.

Accumulate the intended token-weighted objective

Suppose one microbatch has 10 valid tokens and mean loss 2, while another has 30 valid tokens and mean loss 1. The combined token mean is $(10\times2+30\times1)/40=1.25$. Averaging the two means gives 1.5, which gives the short microbatch disproportionate weight.

Gradient accumulation should reproduce the intended denominator across the whole effective batch. With equal valid-token counts, averaging microbatch means works. With unequal counts, scale by valid-token counts or accumulate sums with the correct global denominator. Distributed gradient averaging introduces another scaling factor that must agree with this convention.

AdamW and the update scale

Adam tracks moving averages of gradients and squared gradients, with bias correction. A simplified AdamW update is

$$\theta_{s+1}=(1-\eta_s\lambda)\theta_s-\eta_s\frac{\hat m_s}{\sqrt{\hat v_s}+\epsilon}.$$

The weight-decay term is applied separately from the adaptive gradient preconditioner; this is the distinction in Decoupled Weight Decay Regularization. Parameter groups may choose different decay rules for biases or normalization scales. Learning rate, warmup, decay schedule, and batch size jointly affect update behavior.

Gradient clipping limits an update's gradient norm before the optimizer step. With mixed precision and loss scaling, unscale gradients before inspecting or clipping them. Clipping does not make nonfinite gradients safe; detect those separately and follow the numerical training policy.

Training memory is more than model weights

Budget parameters, gradients, optimizer states, saved activations, temporary workspace, and any master precision copies. KV caches used for incremental inference are a different memory category from the activations needed by a full training backward pass.

Activation checkpointing saves selected intermediate states and recomputes others during backward, trading additional work for memory. FlashAttention changes which attention intermediates must be stored. Neither removes the MLP activations or optimizer state from the budget.

Evaluate a frozen prediction problem

Use held-out documents, disable training dropout, and keep tokenization and loss masking consistent. Aggregate total loss and valid-token counts before computing mean loss or perplexity. The perplexity of pooled tokens is not generally the arithmetic mean of per-batch perplexities.

Track contamination, duplicate documents across splits, and prompt formatting. Language-model loss is useful but does not fully measure factuality, reasoning, or retrieval under long contexts. Evaluate those behaviors with a protocol that specifies prompts, sampling, scoring, and uncertainty.

Compute-allocation studies such as Training Compute-Optimal Large Language Models examine how parameters and training tokens interact under a budget. Their fitted prescriptions are empirical results under particular setups, not a replacement for data quality or an invariant law for every architecture.

Debug a tiny run before a large one

Overfit a small clean batch; verify finite losses and gradients; perturb a future token and confirm earlier logits remain unchanged; compare cached and uncached evaluation; and resume from a checkpoint with optimizer and random states restored. These checks isolate data, masking, arithmetic, and state management before scaling makes failures expensive to inspect.

Try it: Two batches have different numbers of non-padding targets. Is averaging their mean losses always the dataset token loss?

No. Weight each batch mean by its valid-token count, or add loss sums and divide by the total valid-token count.

Connect training to generation

Read GPT-2 as a complete decoder and prefill versus decode. For a different prediction objective using bidirectional attention, continue to BERT.