Course Content
Attention Mechanisms and Transformers
4 sections · 11 lessons
Training Transformers — Masking and Optimization
A training run reaches validation perplexity 4.2 on a translation task. Excellent — better than anything the team has seen. Then someone runs the model on a real sentence and it emits <sos> <sos> <sos> <sos> forever.
The bug is one line. The causal mask was built as torch.triu(ones) instead of torch.tril(ones), so every position could see the entire future and none of the past. During training, predicting token i while looking at token i is trivial, so the loss went to almost nothing. At inference the future column is empty and the model has never learned anything else.
Nothing in the loss curve warns you. This is the defining property of masking bugs: they make training look better, not worse. The same is true of most of the failures in this lesson. Transformer training is not hard because the maths is hard — it is hard because the ways it goes wrong are quiet, and you need to know in advance which numbers to watch.
Part 1: masking
The causal mask
Position i may attend to positions j≤i. Build it once and cache it:
1import torch23def causal_mask(T, device=None):4 """True where attention is ALLOWED. Lower-triangular including the diagonal."""5 return torch.tril(torch.ones(T, T, dtype=torch.bool, device=device))67print(causal_mask(4).int())8# tensor([[1, 0, 0, 0],9# [1, 1, 0, 0],10# [1, 1, 1, 0],11# [1, 1, 1, 1]])Read row i as "what query i can see". Row 0 sees only itself; row 3 sees everything. If your matrix is upper-triangular, every position is looking forwards and you have the bug from the opening.
Application must happen on the scores, before softmax:
scores = Q @ K.transpose(-2, -1) / math.sqrt(d_k) # (B, h, T, T)scores = scores.masked_fill(~mask, float('-inf')) # mask: True = allowedattn = torch.softmax(scores, dim=-1)Why before, and why −∞ rather than zeroing afterwards? Because softmax normalises over whatever it is given. Take position 3 of a four-token sequence with scaled scores [2.0, 1.0, 3.0, 5.0]:
| Approach | w1 | w2 | w3 | w4 | Row sum |
|---|---|---|---|---|---|
| No mask | 0.0414 | 0.0152 | 0.1124 | 0.8310 | 1.000 |
| Softmax then zero the future | 0.0414 | 0.0152 | 0.1124 | 0 | 0.169 |
| −∞ then softmax | 0.2447 | 0.0900 | 0.6652 | 0 | 1.000 |
Check the arithmetic on the correct row: e2=7.389, e1=2.718, e3=20.086, e−∞=0, sum =30.193. Then 7.389/30.193=0.2447, 2.718/30.193=0.0900, 20.086/30.193=0.6652. Sum =0.9999.
The middle row is the trap. Zeroing after softmax leaves the output scaled down by 0.169 at this position and by a different factor at every other position — the model still trains, but every token's contribution is arbitrarily attenuated depending on how much probability mass the mask removed.
A masked score must be excluded from the normalisation, not removed after it. That is the difference between a correct distribution and a randomly attenuated one.
One numerical caution: float('-inf') is correct in float32 but produces nan if an entire row is masked, since softmax of all-−∞ is 0/0. This happens with fully-padded rows. Using a large finite negative number such as -1e9 avoids the nan. In fp16 that value does not fit (fp16's maximum is 65504), and masked_fill with -1e9 on an fp16 tensor raises an overflow error; use -1e4, or torch.finfo(scores.dtype).min, which is right for every dtype.
The padding mask
Batching requires equal lengths, so short sequences get [PAD] tokens. Those positions carry no information and must not be attended to.
def padding_mask(tokens, pad_id=0): """tokens: (B, T) -> (B, 1, 1, T), True where the token is real.""" return (tokens != pad_id).unsqueeze(1).unsqueeze(2)The shape (B,1,1,T) is deliberate. It broadcasts against scores of shape (B,h,Tq,Tk): the two singleton axes expand over heads and over queries, and the final axis lines up with keys. It masks columns — you are removing padding from the set of things that can be attended to.
Concretely, for a batch of three sequences padded to length 6:
tokens (pad_id = 0) mask (True = real token)[ 5, 12, 9, 0, 0, 0] [ T, T, T, F, F, F][ 3, 7, 11, 22, 4, 0] [ T, T, T, T, T, F][ 8, 1, 0, 0, 0, 0] [ T, T, F, F, F, F]Skip this and the model computes attention over embedding vectors for a token that means nothing. In a batch where the longest sequence is 512 and the median is 40, the majority of every attention row is spent on padding, and the model learns to use it — then behaves differently at inference where padding ratios differ.
Note what this mask does not do: rows corresponding to padding queries still produce output. Those outputs are garbage but harmless, provided your loss ignores them. That is what ignore_index=PAD_ID in cross_entropy is for, and forgetting it means a large fraction of your loss is the model being graded on predicting padding.
Combining the two
The decoder needs both. Combine with a logical AND, letting broadcasting align the shapes:
1def make_target_mask(tgt, pad_id=0):2 B, T = tgt.shape3 pad = (tgt != pad_id).unsqueeze(1).unsqueeze(2) # (B, 1, 1, T)4 causal = torch.tril(5 torch.ones(T, T, dtype=torch.bool, device=tgt.device)6 ).unsqueeze(0).unsqueeze(0) # (1, 1, T, T)7 return pad & causal # (B, 1, T, T)For a target [5, 12, 9, 0] with pad_id = 0, the combined mask is:
key: t0 t1 t2 PADquery t0: T F F Fquery t1: T T F Fquery t2: T T T Fquery t3: T T T F <- padding column blocked even though causal allows itThree masks are needed in a full encoder-decoder model, and mixing them up is a common source of quiet failure:
| Mask | Used in | Shape | Blocks |
|---|---|---|---|
| Source padding | Encoder self-attention | (B,1,1,Tsrc) | Padded source columns |
| Target padding + causal | Decoder self-attention | (B,1,Ttgt,Ttgt) | Padded target columns and all future positions |
| Source padding | Decoder cross-attention | (B,1,1,Tsrc) | Padded source columns — never causal, the source is fully available |
Applying a causal mask in cross-attention is a real bug people write. It restricts target position i to source positions ≤i, which is meaningless — the two sequences have no positional correspondence — and quietly caps translation quality.
The leakage test
One test catches essentially every causal masking error:
1@torch.no_grad()2def test_no_future_leak(model, seq_len=8, vocab=100):3 model.eval()4 x = torch.randint(1, vocab, (1, seq_len))5 out_a = model(x)67 x2 = x.clone()8 x2[0, -1] = (x2[0, -1] + 1) % vocab # change ONLY the last token9 out_b = model(x2)1011 # every position before the last must be bit-identical12 assert torch.equal(out_a[0, :-1], out_b[0, :-1]), "future information is leaking"Run it in CI. It takes milliseconds and it is the difference between the opening scenario and a working model.
Part 2: optimisation
The learning rate problem
Transformers are unusually sensitive to learning rate early in training. At step 1 the attention weights are near-uniform and layer norm statistics are meaningless, so gradients are large and poorly conditioned. A learning rate that is correct at step 10,000 will destroy the model at step 10.
The original solution is the Noam schedule: increase the learning rate linearly for warmup steps, then decay it as the inverse square root of the step count.
Work out the actual values for dmodel=512, warmup =4000. First, 512−0.5=1/22.627=0.044194 and 40001.5=252,982.
| Step t | t−0.5 | t⋅4000−1.5 | min | Learning rate |
|---|---|---|---|---|
| 100 | 0.1000 | 0.000395 | 0.000395 | 1.75×10−5 |
| 1,000 | 0.0316 | 0.00395 | 0.00395 | 1.75×10−4 |
| 4,000 | 0.01581 | 0.01581 | 0.01581 | 6.99×10−4 (peak) |
| 16,000 | 0.00791 | 0.0632 | 0.00791 | 3.49×10−4 |
| 100,000 | 0.00316 | 0.395 | 0.00316 | 1.40×10−4 |
The two branches cross exactly at t=4000, which is what makes the schedule continuous. Note the shape: at step 100 the rate is 40 times smaller than at the peak. That gentle start is what keeps the model alive through the unstable early steps.
1import math2from torch.optim.lr_scheduler import LambdaLR34def noam_schedule(optimizer, d_model, warmup=4000):5 def fn(step):6 step = max(step, 1)7 return (d_model ** -0.5) * min(step ** -0.5, step * warmup ** -1.5)8 return LambdaLR(optimizer, fn)910optimizer = torch.optim.Adam(11 model.parameters(), lr=1.0, # lr=1.0 because the lambda IS the rate12 betas=(0.9, 0.98), eps=1e-9,13)14scheduler = noam_schedule(optimizer, d_model=512, warmup=4000)Setting lr=1.0 in the optimiser is not a typo. LambdaLR multiplies the base rate by the lambda, so the lambda must return the absolute rate. Leave the base at the PyTorch default of 10−3 and every rate above is a thousand times too small.
The betas are also non-default. Adam's usual β2=0.999 averages the squared gradient over roughly 1000 steps; β2=0.98 shortens that to about 50, which adapts faster to the rapidly changing gradient scale of early transformer training.
| Schedule | Shape | Best for |
|---|---|---|
| Noam (inverse sqrt) | Linear warmup, then t−0.5 decay | Training from scratch with no fixed step budget |
| Linear warmup + linear decay | Warmup, then straight line to zero | Fine-tuning; the default in Hugging Face |
| Warmup + cosine decay | Warmup, then half a cosine to a small floor | Large pre-training runs with a known total step count |
| Constant | Flat | Nothing, really — it either diverges early or plateaus late |
Gradient clipping
Warmup handles the systematic instability. Clipping handles the occasional catastrophic batch — a very long sequence, a rare token, a label error — that produces a gradient hundreds of times the usual size.
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)This computes the global norm across all parameters and rescales everything if it exceeds the threshold. Suppose two parameter groups have gradient norms 3.0 and 12.0. The global norm is
With max_norm=1.0, every gradient is multiplied by 1.0/12.369=0.0808, giving norms 0.2425 and 0.9701 — global norm exactly 1.0. Crucially the direction is preserved; only the magnitude is capped. Clipping each parameter separately would change the direction and is not what you want.
Log the pre-clip norm (clip_grad_norm_ returns it). A healthy run shows it settling into a stable band, say 0.3 to 1.5. Occasional spikes to 10 are normal and are exactly what clipping is for. A norm that grows steadily over thousands of steps means the learning rate is too high.
Gradient accumulation
Transformers train better with large batches, and large batches do not fit in memory. Accumulation simulates them: run several small batches, sum the gradients, then step once.
1accum_steps = 4 # micro-batch 8 -> effective batch 3223optimizer.zero_grad(set_to_none=True)4for i, batch in enumerate(loader):5 loss = compute_loss(model, batch)6 (loss / accum_steps).backward() # divide BEFORE backward78 if (i + 1) % accum_steps == 0:9 torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)10 optimizer.step()11 scheduler.step()12 optimizer.zero_grad(set_to_none=True)The division by accum_steps is the part people omit. Gradients accumulate additively, so without it the summed gradient is 4 times larger than the gradient of the true batch mean — equivalent to quadrupling the learning rate at every optimiser step, with all the instability that implies.
Also note where scheduler.step() sits. It belongs inside the if, alongside optimizer.step(). Outside it, your schedule advances 4 times per actual update and warmup finishes in a quarter of the intended steps.
Mixed precision
Half precision halves memory and roughly doubles throughput on tensor cores. It also introduces a range problem.
| float32 | float16 | bfloat16 | |
|---|---|---|---|
| Exponent bits | 8 | 5 | 8 |
| Mantissa bits | 23 | 10 | 7 |
| Largest value | 3.4×1038 | 65,504 | 3.4×1038 |
| Smallest normal | 1.2×10−38 | 6.1×10−5 | 1.2×10−38 |
| Loss scaling needed | No | Yes | No |
The problem is that smallest-normal row. Transformer gradients are routinely around 10−7 or smaller. In fp16 anything below 6.1×10−5 becomes subnormal and below about 6×10−8 becomes exactly zero — the gradient silently disappears.
Loss scaling fixes it. Multiply the loss by a large constant S before backward(); by linearity every gradient is scaled by S too. Then divide the gradients by S before the optimiser step. With S=65536, a gradient of 1×10−8 becomes 6.55×10−4 — comfortably inside fp16's normal range. GradScaler chooses S automatically, halving it whenever it detects an overflow and doubling it after a stretch of clean steps.
1from torch.amp import autocast, GradScaler23scaler = GradScaler('cuda')45for batch in loader:6 optimizer.zero_grad(set_to_none=True)78 with autocast('cuda', dtype=torch.float16):9 loss = compute_loss(model, batch)1011 scaler.scale(loss).backward()1213 scaler.unscale_(optimizer) # undo scaling FIRST14 torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) # ...then clip1516 scaler.step(optimizer)17 scaler.update()18 scheduler.step()The scaler.unscale_ call before clipping is mandatory. Clip a gradient that is still scaled by 65536 and every batch looks like an outlier, so every gradient gets crushed to 1/65536 of its size and the model does not move.
On hardware that supports it, bfloat16 avoids all of this: same exponent range as fp32, so no scaler at all. The cost is 7 mantissa bits instead of 10, which in practice matters far less than range.
Label smoothing
The original transformer also used label smoothing of 0.1. Instead of a one-hot target, the correct token gets probability 1−ε and the remaining ε is spread over the rest of the vocabulary. For ε=0.1 and a 32,000-token vocabulary, each incorrect token receives
This deliberately makes the training loss worse — the model can never reach zero loss, because the target is never achievable. What it buys is calibration: the model stops driving correct logits towards +∞, which reduces overconfidence and, in the original paper, improved BLEU despite raising perplexity. If you see training loss plateau at a value well above zero and wonder why, check whether label smoothing is on.
A complete step
1def train_epoch(model, loader, optimizer, scheduler, scaler,2 accum_steps=1, clip=1.0, pad_id=0, device='cuda'):3 model.train()4 total_loss, total_tokens = 0.0, 05 optimizer.zero_grad(set_to_none=True)67 for i, (src, tgt) in enumerate(loader):8 src, tgt = src.to(device), tgt.to(device)9 tgt_in, tgt_out = tgt[:, :-1], tgt[:, 1:]1011 src_mask = padding_mask(src, pad_id)12 tgt_mask = make_target_mask(tgt_in, pad_id)1314 with autocast('cuda', dtype=torch.float16):15 logits = model(src, tgt_in, src_mask, tgt_mask)16 loss = F.cross_entropy(17 logits.reshape(-1, logits.size(-1)),18 tgt_out.reshape(-1),19 ignore_index=pad_id,20 label_smoothing=0.1,21 )2223 scaler.scale(loss / accum_steps).backward()2425 if (i + 1) % accum_steps == 0:26 scaler.unscale_(optimizer)27 grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), clip)28 scaler.step(optimizer)29 scaler.update()30 scheduler.step()31 optimizer.zero_grad(set_to_none=True)3233 n = (tgt_out != pad_id).sum().item()34 total_loss += loss.item() * n35 total_tokens += n3637 return total_loss / total_tokensWeighting the loss by real token count rather than averaging per batch matters when sequence lengths vary. Averaging per batch gives a 5-token sequence the same weight as a 200-token one.
What to watch, and what each signal means
| Signal | Healthy | What a deviation means |
|---|---|---|
| Initial loss | ≈lnV — 10.4 for a 32k vocabulary | Much higher: bad initialisation or missing normalisation. Much lower: labels are leaking into the input |
| Pre-clip gradient norm | Settles into a band, e.g. 0.3–1.5 | Steadily rising: learning rate too high. Near zero: dead layers or a saturated softmax |
| Fraction of steps clipped | Under ~5% after warmup | Consistently high: reduce the learning rate rather than lowering the clip threshold |
| Loss scale (fp16) | Stabilises around 104–105 | Repeatedly halving: overflow somewhere; suspect the attention softmax or an unmasked −1e9 in fp16 |
| Attention entropy | Falls from lnT towards 1–3 nats | Stuck at lnT: heads are uniform and doing nothing. At 0: saturated, check dk |
| Train vs validation gap | Grows slowly | Validation far worse and diverging: overfitting. Validation better than train: dropout is on in eval, or you are computing them on different data |
What this means when you build something
Order your debugging by how quiet the failure is, not by how likely it seems. The loud failures — nan, out-of-memory, shape errors — announce themselves and you will fix them anyway. The quiet ones need deliberate tests.
| Observation | Cause | Fix |
|---|---|---|
| Validation loss suspiciously good; generation is degenerate | Causal mask inverted or absent; or decoder input not shifted | Run the leakage test; print the mask; confirm tgt[:, :-1] vs tgt[:, 1:] |
| Loss diverges in the first few hundred steps | No warmup, or post-norm without warmup | Add warmup over ~4000 steps; consider pre-norm |
| Model does not improve at all with fp16 | Clipping applied before unscale_ | Call scaler.unscale_(optimizer) first |
| Effective learning rate 4× too high; loss noisy | Loss not divided by accum_steps | Divide before backward() |
| Warmup ends far too early | scheduler.step() outside the accumulation guard | Move it next to optimizer.step() |
| Large fraction of loss coming from padding | ignore_index not set | Pass ignore_index=pad_id to cross_entropy |
Before any long run, overfit a single batch. Take eight examples, turn off dropout and weight decay, and train until the loss is essentially zero — under 0.01. It should take a couple of hundred steps. If it cannot, no hyperparameter will save the full run, and you have just localised the bug to the model or the loss rather than the data pipeline. If it can, you have proved the forward and backward passes are wired correctly, and every remaining problem is about scale, data, or regularisation.
That one test costs two minutes and saves days.